mirror of
https://github.com/416rehman/DeepZero-Agentic-Vulnerability-Research-Pipeline
synced 2026-08-09 11:55:01 +00:00
Merge PR #13: Claude Code LLM backend + end-to-end pipeline fixes
feat: Claude Code LLM backend (no API key) + fixes to run the loldrivers pipeline end-to-end on Windows
This commit is contained in:
+15
-2
@@ -5,8 +5,21 @@
|
|||||||
# DeepZero will automatically load '.env' using python-dotenv.
|
# DeepZero will automatically load '.env' using python-dotenv.
|
||||||
|
|
||||||
# =========================================================
|
# =========================================================
|
||||||
# 1. LiteLLM Abstraction Keys
|
# 1. LLM Backend
|
||||||
# Provide keys for whatever endpoints your pipelines target
|
#
|
||||||
|
# Two options - pick per pipeline via `model:` in pipeline.yaml
|
||||||
|
# (or override with `deepzero run -m ...`).
|
||||||
|
#
|
||||||
|
# a) Claude Code (no API key)
|
||||||
|
# Uses your own locally installed, already-signed-in Claude
|
||||||
|
# Code CLI. Nothing to configure here - just set:
|
||||||
|
# model: claude-code # default model
|
||||||
|
# model: claude-code/sonnet # or /opus, or a full model name
|
||||||
|
# Requires: `claude` on PATH and signed in (run `claude`).
|
||||||
|
#
|
||||||
|
# b) LiteLLM (API key, metered)
|
||||||
|
# Best for CI/servers. Provide keys for whatever endpoints
|
||||||
|
# your pipelines target, e.g. model: openai/gpt-4o
|
||||||
# =========================================================
|
# =========================================================
|
||||||
# GEMINI_API_KEY=AI...
|
# GEMINI_API_KEY=AI...
|
||||||
# OPENAI_API_KEY=sk-...
|
# OPENAI_API_KEY=sk-...
|
||||||
|
|||||||
@@ -13,7 +13,7 @@ name: loldrivers
|
|||||||
description: windows kernel driver vulnerability research pipeline
|
description: windows kernel driver vulnerability research pipeline
|
||||||
version: "2.0"
|
version: "2.0"
|
||||||
|
|
||||||
model: vertex_ai/gemini-2.5-pro
|
model: claude-code/sonnet
|
||||||
|
|
||||||
settings:
|
settings:
|
||||||
work_dir: work
|
work_dir: work
|
||||||
|
|||||||
@@ -360,10 +360,12 @@ def main():
|
|||||||
# write per-ioctl file
|
# write per-ioctl file
|
||||||
# FIX: Use io.open with utf-8 encoding and unicode() cast for Jython stability
|
# FIX: Use io.open with utf-8 encoding and unicode() cast for Jython stability
|
||||||
with io.open(os.path.join(ioctls_dir, "0x%08X.c" % code), "w", encoding="utf-8") as f:
|
with io.open(os.path.join(ioctls_dir, "0x%08X.c" % code), "w", encoding="utf-8") as f:
|
||||||
f.write("// IOCTL Code: 0x%08X\n" % code)
|
# unicode() cast required: a jython io text stream rejects str,
|
||||||
f.write("// Method: %d\n" % (code & 0x3))
|
# and "..." % int yields str (raising "can't write str to text stream")
|
||||||
f.write("// Device Type: 0x%04X\n" % ((code >> 16) & 0xFFFF))
|
f.write(unicode("// IOCTL Code: 0x%08X\n" % code))
|
||||||
f.write("// Function: 0x%03X\n\n" % ((code >> 2) & 0xFFF))
|
f.write(unicode("// Method: %d\n" % (code & 0x3)))
|
||||||
|
f.write(unicode("// Device Type: 0x%04X\n" % ((code >> 16) & 0xFFFF)))
|
||||||
|
f.write(unicode("// Function: 0x%03X\n\n" % ((code >> 2) & 0xFFF)))
|
||||||
f.write(unicode(dispatch_c))
|
f.write(unicode(dispatch_c))
|
||||||
|
|
||||||
result["success"] = True
|
result["success"] = True
|
||||||
|
|||||||
@@ -19,28 +19,42 @@ from deepzero.engine.stage import (
|
|||||||
class SemgrepScanner(BulkMapProcessor):
|
class SemgrepScanner(BulkMapProcessor):
|
||||||
description = "runs semgrep batch scan against decompiled source across all active samples"
|
description = "runs semgrep batch scan against decompiled source across all active samples"
|
||||||
|
|
||||||
|
def _resolve_rules_path(self, ctx: ProcessorContext) -> Path | None:
|
||||||
|
# resolve rules_dir consistently for validate() and process(): try
|
||||||
|
# cwd-relative first (how the shipped pipelines reference their rules),
|
||||||
|
# then relative to the pipeline directory. resolving these differently
|
||||||
|
# let validation pass while the scan silently ran with no rules loaded.
|
||||||
|
rules_dir = self.config.get("rules_dir", "")
|
||||||
|
if not rules_dir:
|
||||||
|
return None
|
||||||
|
cwd_path = (Path.cwd() / rules_dir).resolve()
|
||||||
|
if cwd_path.exists():
|
||||||
|
return cwd_path
|
||||||
|
return (ctx.pipeline_dir / rules_dir).resolve()
|
||||||
|
|
||||||
def validate(self, ctx: ProcessorContext) -> list[str]:
|
def validate(self, ctx: ProcessorContext) -> list[str]:
|
||||||
errors = []
|
errors = []
|
||||||
if not shutil.which("semgrep"):
|
if not shutil.which("semgrep"):
|
||||||
errors.append("semgrep CLI not found in PATH - install with: pip install semgrep")
|
errors.append("semgrep CLI not found in PATH - install with: pip install semgrep")
|
||||||
|
|
||||||
rules_dir = self.config.get("rules_dir")
|
if not self.config.get("rules_dir"):
|
||||||
if not rules_dir:
|
|
||||||
errors.append("semgrep_scanner requires 'rules_dir' in config")
|
errors.append("semgrep_scanner requires 'rules_dir' in config")
|
||||||
else:
|
else:
|
||||||
rules_path = (Path.cwd() / rules_dir).resolve()
|
rules_path = self._resolve_rules_path(ctx)
|
||||||
if not rules_path.exists():
|
if rules_path is None or not rules_path.exists():
|
||||||
rules_path = (ctx.pipeline_dir / rules_dir).resolve()
|
errors.append(f"rules_dir does not exist: {self.config.get('rules_dir')}")
|
||||||
if not rules_path.exists():
|
|
||||||
errors.append(f"rules_dir does not exist: {rules_dir}")
|
|
||||||
|
|
||||||
return errors
|
return errors
|
||||||
|
|
||||||
def process(
|
def process(
|
||||||
self, ctx: ProcessorContext, entries: list[ProcessorEntry]
|
self, ctx: ProcessorContext, entries: list[ProcessorEntry]
|
||||||
) -> list[ProcessorResult]:
|
) -> list[ProcessorResult]:
|
||||||
rules_dir = self.config.get("rules_dir", "")
|
rules_path = self._resolve_rules_path(ctx)
|
||||||
rules_path = (ctx.pipeline_dir / rules_dir).resolve()
|
if rules_path is None or not rules_path.exists():
|
||||||
|
# fail fast rather than pointing semgrep at a fallback path and
|
||||||
|
# scanning with no rules loaded (the silent-empty-results class)
|
||||||
|
reason = f"semgrep rules_dir not found: {self.config.get('rules_dir') or '(unset)'}"
|
||||||
|
return [ProcessorResult.fail(reason) for _ in entries]
|
||||||
|
|
||||||
target_subdir = self.config.get("target_dir", "decompiled")
|
target_subdir = self.config.get("target_dir", "decompiled")
|
||||||
timeout = self.config.get("timeout", 300)
|
timeout = self.config.get("timeout", 300)
|
||||||
@@ -170,20 +184,46 @@ class SemgrepScanner(BulkMapProcessor):
|
|||||||
results[idx] = ProcessorResult.fail(f"semgrep batch timed out after {timeout}s")
|
results[idx] = ProcessorResult.fail(f"semgrep batch timed out after {timeout}s")
|
||||||
return [r for r in results if r is not None]
|
return [r for r in results if r is not None]
|
||||||
|
|
||||||
if proc.returncode not in (0, 1):
|
out_str = stdout_bytes.decode("utf-8", errors="replace")
|
||||||
err = stderr_bytes.decode("utf-8", errors="replace")[:500]
|
err_str = stderr_bytes.decode("utf-8", errors="replace")
|
||||||
|
|
||||||
|
# semgrep emits a complete results document on stdout even when it exits
|
||||||
|
# with an unexpected code (observed on windows), so a parseable scan
|
||||||
|
# result is authoritative over the exit code. only treat the run as
|
||||||
|
# failed when no usable output came back - and always report the exit
|
||||||
|
# code, never an empty "semgrep error:".
|
||||||
|
output: dict[str, Any] | None = None
|
||||||
|
if out_str.strip():
|
||||||
|
try:
|
||||||
|
parsed = json.loads(out_str)
|
||||||
|
if isinstance(parsed, dict) and "results" in parsed:
|
||||||
|
output = parsed
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
output = None
|
||||||
|
|
||||||
|
if output is None:
|
||||||
|
detail = (err_str.strip() or "no parseable output on stdout")[:500]
|
||||||
for idx, _ in uncached_entries:
|
for idx, _ in uncached_entries:
|
||||||
results[idx] = ProcessorResult.fail(f"semgrep error: {err}")
|
results[idx] = ProcessorResult.fail(
|
||||||
|
f"semgrep failed (exit {proc.returncode}): {detail}"
|
||||||
|
)
|
||||||
return [r for r in results if r is not None]
|
return [r for r in results if r is not None]
|
||||||
|
|
||||||
try:
|
# semgrep reports rule/config load failures in an "errors" array while
|
||||||
out_str = stdout_bytes.decode("utf-8", errors="replace")
|
# still exiting 0 with empty results - which would otherwise look like
|
||||||
output = json.loads(out_str) if out_str.strip() else {}
|
# "no vulnerabilities found". fail loudly when nothing scanned.
|
||||||
except json.JSONDecodeError:
|
scan_errors = output.get("errors") or []
|
||||||
|
if scan_errors and not output.get("results"):
|
||||||
|
detail = "; ".join(
|
||||||
|
str(e.get("message") or e.get("type") or e) for e in scan_errors[:3]
|
||||||
|
)[:500]
|
||||||
for idx, _ in uncached_entries:
|
for idx, _ in uncached_entries:
|
||||||
results[idx] = ProcessorResult.fail("failed to parse semgrep json output")
|
results[idx] = ProcessorResult.fail(f"semgrep produced no results: {detail}")
|
||||||
return [r for r in results if r is not None]
|
return [r for r in results if r is not None]
|
||||||
|
|
||||||
|
if scan_errors:
|
||||||
|
self.log.warning("semgrep reported %d non-fatal error(s) during scan", len(scan_errors))
|
||||||
|
|
||||||
self._distribute_findings(output, file_to_sample, uncached_entries, results, min_findings)
|
self._distribute_findings(output, file_to_sample, uncached_entries, results, min_findings)
|
||||||
return [r for r in results if r is not None]
|
return [r for r in results if r is not None]
|
||||||
|
|
||||||
|
|||||||
@@ -36,6 +36,11 @@ pythonpath = ["src", "."]
|
|||||||
[tool.ruff]
|
[tool.ruff]
|
||||||
line-length = 100
|
line-length = 100
|
||||||
target-version = "py311"
|
target-version = "py311"
|
||||||
|
# docs contain illustrative python snippets in markdown. newer ruff versions
|
||||||
|
# format embedded code blocks, which made `ruff format --check` fail in CI
|
||||||
|
# purely from a ruff upgrade. documentation examples are prose, not build
|
||||||
|
# artifacts - keep them out of the code formatter's scope.
|
||||||
|
extend-exclude = ["docs"]
|
||||||
|
|
||||||
[tool.ruff.lint]
|
[tool.ruff.lint]
|
||||||
select = ["E", "F", "W", "I"]
|
select = ["E", "F", "W", "I"]
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
|
import sys
|
||||||
import time
|
import time
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from types import MappingProxyType
|
from types import MappingProxyType
|
||||||
@@ -14,6 +15,21 @@ from rich.table import Table
|
|||||||
|
|
||||||
from deepzero.engine.types import RunStatus
|
from deepzero.engine.types import RunStatus
|
||||||
|
|
||||||
|
|
||||||
|
def _ensure_utf8_streams() -> None:
|
||||||
|
# deepzero prints unicode status glyphs (checkmarks, arrows, box drawing).
|
||||||
|
# on a legacy windows console or when stdout is redirected to a file, the
|
||||||
|
# default cp1252 encoding raises UnicodeEncodeError and crashes the run.
|
||||||
|
# force utf-8 with a non-fatal error handler before any Console is built.
|
||||||
|
for stream in (sys.stdout, sys.stderr):
|
||||||
|
try:
|
||||||
|
stream.reconfigure(encoding="utf-8", errors="backslashreplace")
|
||||||
|
except (AttributeError, ValueError):
|
||||||
|
# stream is not a reconfigurable TextIOWrapper (e.g. already wrapped)
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
_ensure_utf8_streams()
|
||||||
console = Console()
|
console = Console()
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,62 @@
|
|||||||
|
"""pluggable llm backends.
|
||||||
|
|
||||||
|
to add support for another agent cli (codex, gemini cli, ...):
|
||||||
|
|
||||||
|
# deepzero/engine/backends/mytool.py
|
||||||
|
class MyToolBackend(CLIAgentBackend):
|
||||||
|
scheme = "mytool"
|
||||||
|
display_name = "my tool"
|
||||||
|
binary_names = ("mytool",)
|
||||||
|
|
||||||
|
def build_argv(self, system): ...
|
||||||
|
def parse_output(self, returncode, stdout, stderr): ...
|
||||||
|
|
||||||
|
then register it below (or call register_backend() from your own code - the
|
||||||
|
registry is open, so third parties can add backends without patching deepzero).
|
||||||
|
nothing in LLMProvider, the pipeline, or the stages needs to change.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from deepzero.engine.backends.base import (
|
||||||
|
BackendAuthError,
|
||||||
|
BackendContextWindowError,
|
||||||
|
BackendError,
|
||||||
|
BackendNotFoundError,
|
||||||
|
BackendRateLimitError,
|
||||||
|
CLIAgentBackend,
|
||||||
|
LLMBackend,
|
||||||
|
)
|
||||||
|
from deepzero.engine.backends.claude_code import ClaudeCodeBackend
|
||||||
|
from deepzero.engine.backends.litellm_backend import LiteLLMBackend
|
||||||
|
from deepzero.engine.backends.registry import (
|
||||||
|
create_backend,
|
||||||
|
get_registered_backends,
|
||||||
|
model_scheme,
|
||||||
|
register_backend,
|
||||||
|
resolve_backend_class,
|
||||||
|
validate_model_binding,
|
||||||
|
)
|
||||||
|
|
||||||
|
# -- built-in backends --
|
||||||
|
register_backend(ClaudeCodeBackend)
|
||||||
|
# litellm is the fallback for any scheme no other backend claims
|
||||||
|
register_backend(LiteLLMBackend, default=True)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"BackendAuthError",
|
||||||
|
"BackendContextWindowError",
|
||||||
|
"BackendError",
|
||||||
|
"BackendNotFoundError",
|
||||||
|
"BackendRateLimitError",
|
||||||
|
"CLIAgentBackend",
|
||||||
|
"ClaudeCodeBackend",
|
||||||
|
"LLMBackend",
|
||||||
|
"LiteLLMBackend",
|
||||||
|
"create_backend",
|
||||||
|
"get_registered_backends",
|
||||||
|
"model_scheme",
|
||||||
|
"register_backend",
|
||||||
|
"resolve_backend_class",
|
||||||
|
"validate_model_binding",
|
||||||
|
]
|
||||||
@@ -0,0 +1,292 @@
|
|||||||
|
"""llm backend interface and shared implementations.
|
||||||
|
|
||||||
|
two layers live here:
|
||||||
|
|
||||||
|
LLMBackend - the contract LLMProvider talks to. anything that can turn
|
||||||
|
messages into text can be a backend.
|
||||||
|
CLIAgentBackend - reusable plumbing for backends that shell out to a locally
|
||||||
|
installed coding-agent cli (claude code, codex, gemini cli,
|
||||||
|
...). subclasses declare *what* to run and *how to read the
|
||||||
|
output*; process handling, prompt flattening, env
|
||||||
|
sanitization and error classification are inherited.
|
||||||
|
|
||||||
|
adding a new agent cli should mean subclassing CLIAgentBackend and registering
|
||||||
|
it - never editing LLMProvider or the dispatch logic.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import shutil
|
||||||
|
import subprocess
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, ClassVar
|
||||||
|
|
||||||
|
log = logging.getLogger("deepzero.llm.backend")
|
||||||
|
|
||||||
|
|
||||||
|
# -- generic error taxonomy -----------------------------------------------------
|
||||||
|
#
|
||||||
|
# the retry loop reasons about these, never about vendor-specific classes.
|
||||||
|
# backends either raise these directly or declare which of their own exception
|
||||||
|
# classes map onto each role (see the class attributes on LLMBackend).
|
||||||
|
|
||||||
|
|
||||||
|
class BackendError(RuntimeError):
|
||||||
|
"""backend call failed."""
|
||||||
|
|
||||||
|
|
||||||
|
class BackendNotFoundError(BackendError):
|
||||||
|
"""the backend's binary or dependency is not installed."""
|
||||||
|
|
||||||
|
|
||||||
|
class BackendAuthError(BackendError):
|
||||||
|
"""backend is present but not authenticated - retrying will not help."""
|
||||||
|
|
||||||
|
|
||||||
|
class BackendRateLimitError(BackendError):
|
||||||
|
"""rate/usage limit hit - worth retrying after backoff."""
|
||||||
|
|
||||||
|
|
||||||
|
class BackendContextWindowError(BackendError):
|
||||||
|
"""prompt exceeded the context window - retrying will not help."""
|
||||||
|
|
||||||
|
|
||||||
|
class _NeverRaised(Exception):
|
||||||
|
"""placeholder so a backend can opt out of an error role."""
|
||||||
|
|
||||||
|
|
||||||
|
class LLMBackend(ABC):
|
||||||
|
"""turns messages into response text.
|
||||||
|
|
||||||
|
subclasses declare which exception classes fill each retry role, so
|
||||||
|
LLMProvider's backoff loop needs no knowledge of any specific backend.
|
||||||
|
"""
|
||||||
|
|
||||||
|
# registry identity. `scheme` is the part of the model string before the
|
||||||
|
# first "/" (e.g. "claude-code" in "claude-code/sonnet").
|
||||||
|
scheme: ClassVar[str] = ""
|
||||||
|
aliases: ClassVar[tuple[str, ...]] = ()
|
||||||
|
display_name: ClassVar[str] = ""
|
||||||
|
|
||||||
|
# retry roles - override with backend-specific classes where needed
|
||||||
|
rate_limit_error: ClassVar[type[BaseException]] = BackendRateLimitError
|
||||||
|
context_window_error: ClassVar[type[BaseException]] = BackendContextWindowError
|
||||||
|
non_retryable_errors: ClassVar[tuple[type[BaseException], ...]] = (BackendAuthError,)
|
||||||
|
retryable_errors: ClassVar[tuple[type[BaseException], ...]] = (BackendError, OSError)
|
||||||
|
|
||||||
|
def __init__(self, model: str, **kwargs: Any):
|
||||||
|
self.model = model
|
||||||
|
self.options = kwargs
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def raw_complete(self, messages: list[dict[str, str]], **kwargs: Any) -> str:
|
||||||
|
"""single completion attempt. raise on failure - retries are handled upstream."""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def validate_model(cls, model: str) -> list[str]:
|
||||||
|
"""pre-flight check run at pipeline load time. return problems, or []."""
|
||||||
|
return []
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def model_alias(cls, model: str) -> str:
|
||||||
|
"""the part after the scheme, e.g. "sonnet" for "claude-code/sonnet"."""
|
||||||
|
m = (model or "").strip()
|
||||||
|
return m.split("/", 1)[1].strip() if "/" in m else ""
|
||||||
|
|
||||||
|
@property
|
||||||
|
def provider_name(self) -> str:
|
||||||
|
return self.scheme or "unknown"
|
||||||
|
|
||||||
|
@property
|
||||||
|
def model_name(self) -> str:
|
||||||
|
return self.model_alias(self.model) or "default"
|
||||||
|
|
||||||
|
|
||||||
|
class CLIAgentBackend(LLMBackend):
|
||||||
|
"""drives a locally installed, already-authenticated agent cli.
|
||||||
|
|
||||||
|
deepzero never handles credentials for these: it invokes the user's own
|
||||||
|
binary and inherits whatever auth that binary already has.
|
||||||
|
"""
|
||||||
|
|
||||||
|
# -- subclass declares these --
|
||||||
|
binary_names: ClassVar[tuple[str, ...]] = ()
|
||||||
|
install_hint: ClassVar[str] = ""
|
||||||
|
default_timeout: ClassVar[int] = 900
|
||||||
|
# env vars to hide from the child so it uses its own login rather than a
|
||||||
|
# metered api key that happens to be in .env for another backend
|
||||||
|
subscription_auth_env_blocklist: ClassVar[tuple[str, ...]] = ()
|
||||||
|
|
||||||
|
# argv has hard length limits (~32k on windows); larger text goes via stdin
|
||||||
|
max_argv_text: ClassVar[int] = 8000
|
||||||
|
|
||||||
|
def __init__(self, model: str, **kwargs: Any):
|
||||||
|
super().__init__(model, **kwargs)
|
||||||
|
self.alias = self.model_alias(model)
|
||||||
|
self.timeout = int(kwargs.get("timeout", self.default_timeout))
|
||||||
|
self.cwd = kwargs.get("cwd")
|
||||||
|
# choosing an agent-cli backend signals intent to use its own login
|
||||||
|
self.prefer_subscription_auth = bool(kwargs.get("prefer_subscription_auth", True))
|
||||||
|
|
||||||
|
self._binary = kwargs.get("binary") or self.find_binary()
|
||||||
|
if not self._binary:
|
||||||
|
raise BackendNotFoundError(self.not_found_message())
|
||||||
|
|
||||||
|
# -- discovery --------------------------------------------------------------
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def find_binary(cls) -> str | None:
|
||||||
|
for name in cls.binary_names:
|
||||||
|
found = shutil.which(name)
|
||||||
|
if found:
|
||||||
|
return found
|
||||||
|
for candidate in cls.extra_search_paths():
|
||||||
|
try:
|
||||||
|
if candidate.is_file():
|
||||||
|
return str(candidate)
|
||||||
|
except OSError:
|
||||||
|
continue
|
||||||
|
return None
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def extra_search_paths(cls) -> list[Path]:
|
||||||
|
"""non-PATH locations to probe. override for vendor-specific installs."""
|
||||||
|
return []
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def not_found_message(cls) -> str:
|
||||||
|
label = cls.display_name or cls.scheme
|
||||||
|
hint = f" {cls.install_hint}" if cls.install_hint else ""
|
||||||
|
return (
|
||||||
|
f"{label} cli not found (looked for: {', '.join(cls.binary_names) or 'n/a'})."
|
||||||
|
f"{hint} alternatively use an api-key model instead."
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def validate_model(cls, model: str) -> list[str]:
|
||||||
|
if not cls.find_binary():
|
||||||
|
return [f"LLM backend '{model}' is unavailable: {cls.not_found_message()}"]
|
||||||
|
return []
|
||||||
|
|
||||||
|
# -- prompt assembly --------------------------------------------------------
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def split_messages(messages: list[dict[str, str]]) -> tuple[str, str]:
|
||||||
|
"""flatten a message list into (system_prompt, user_prompt).
|
||||||
|
|
||||||
|
agent clis are single-shot, so prior turns are rendered inline. a lone
|
||||||
|
user message - the pipeline's normal case - passes through untouched.
|
||||||
|
"""
|
||||||
|
system_parts: list[str] = []
|
||||||
|
convo: list[dict[str, str]] = []
|
||||||
|
for msg in messages or []:
|
||||||
|
role = (msg.get("role") or "").lower()
|
||||||
|
content = msg.get("content") or ""
|
||||||
|
if role == "system":
|
||||||
|
if content:
|
||||||
|
system_parts.append(content)
|
||||||
|
else:
|
||||||
|
convo.append({"role": role or "user", "content": content})
|
||||||
|
|
||||||
|
system = "\n\n".join(system_parts)
|
||||||
|
|
||||||
|
if len(convo) == 1:
|
||||||
|
return system, convo[0]["content"]
|
||||||
|
|
||||||
|
rendered = [
|
||||||
|
f"{'Assistant' if m['role'] == 'assistant' else 'User'}: {m['content']}" for m in convo
|
||||||
|
]
|
||||||
|
return system, "\n\n".join(rendered)
|
||||||
|
|
||||||
|
# -- subclass hooks ---------------------------------------------------------
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def build_argv(self, system: str) -> tuple[list[str], str]:
|
||||||
|
"""returns (argv, system_text_to_inline).
|
||||||
|
|
||||||
|
a system prompt too large for argv should be returned as the second
|
||||||
|
element so the base class routes it through stdin instead.
|
||||||
|
"""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def parse_output(self, returncode: int, stdout: str, stderr: str) -> str:
|
||||||
|
"""extract response text, or raise a BackendError subclass."""
|
||||||
|
|
||||||
|
def build_env(self) -> dict[str, str]:
|
||||||
|
env = dict(os.environ)
|
||||||
|
if self.prefer_subscription_auth:
|
||||||
|
for var in self.subscription_auth_env_blocklist:
|
||||||
|
env.pop(var, None)
|
||||||
|
return env
|
||||||
|
|
||||||
|
def classify_error(self, detail: str, status: Any = None) -> BackendError:
|
||||||
|
"""map a failure into the generic taxonomy. override to add vendor cases."""
|
||||||
|
low = (detail or "").lower()
|
||||||
|
try:
|
||||||
|
code = int(status) if status is not None else None
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
code = None
|
||||||
|
|
||||||
|
label = self.display_name or self.scheme
|
||||||
|
if code in (401, 403) or any(k in low for k in ("authenticat", "oauth", "unauthorized")):
|
||||||
|
return BackendAuthError(f"{label} is not authenticated ({detail})")
|
||||||
|
if code == 429 or any(
|
||||||
|
k in low for k in ("rate limit", "usage limit", "quota", "overloaded", "too many")
|
||||||
|
):
|
||||||
|
return BackendRateLimitError(detail)
|
||||||
|
if "context" in low and any(k in low for k in ("too long", "exceed", "window")):
|
||||||
|
return BackendContextWindowError(detail)
|
||||||
|
if code is not None and 500 <= code < 600:
|
||||||
|
return BackendError(f"{label} server error {code}: {detail}")
|
||||||
|
return BackendError(detail)
|
||||||
|
|
||||||
|
# -- execution --------------------------------------------------------------
|
||||||
|
|
||||||
|
def raw_complete(self, messages: list[dict[str, str]], **kwargs: Any) -> str:
|
||||||
|
system, prompt = self.split_messages(messages)
|
||||||
|
argv, inline_system = self.build_argv(system)
|
||||||
|
|
||||||
|
stdin_text = f"{inline_system}\n\n{prompt}" if inline_system else prompt
|
||||||
|
timeout = int(kwargs.get("timeout", self.timeout))
|
||||||
|
|
||||||
|
try:
|
||||||
|
proc = subprocess.run(
|
||||||
|
argv,
|
||||||
|
input=stdin_text,
|
||||||
|
capture_output=True,
|
||||||
|
text=True,
|
||||||
|
encoding="utf-8",
|
||||||
|
errors="replace",
|
||||||
|
timeout=timeout,
|
||||||
|
env=self.build_env(),
|
||||||
|
cwd=self.cwd,
|
||||||
|
)
|
||||||
|
except subprocess.TimeoutExpired as exc:
|
||||||
|
label = self.display_name or self.scheme
|
||||||
|
raise BackendError(f"{label} timed out after {timeout}s") from exc
|
||||||
|
except OSError as exc:
|
||||||
|
label = self.display_name or self.scheme
|
||||||
|
raise BackendError(f"failed to launch {label}: {exc}") from exc
|
||||||
|
|
||||||
|
return self.parse_output(proc.returncode, proc.stdout or "", proc.stderr or "")
|
||||||
|
|
||||||
|
def check_binary(self) -> str | None:
|
||||||
|
"""returns None if the binary runs, else a description of the problem."""
|
||||||
|
try:
|
||||||
|
proc = subprocess.run(
|
||||||
|
[str(self._binary), "--version"],
|
||||||
|
capture_output=True,
|
||||||
|
text=True,
|
||||||
|
encoding="utf-8",
|
||||||
|
errors="replace",
|
||||||
|
timeout=30,
|
||||||
|
env=self.build_env(),
|
||||||
|
)
|
||||||
|
except (OSError, subprocess.SubprocessError) as exc:
|
||||||
|
return f"could not run {self.display_name or self.scheme}: {exc}"
|
||||||
|
if proc.returncode != 0:
|
||||||
|
return f"exited {proc.returncode}: {(proc.stderr or '').strip()[:200]}"
|
||||||
|
return None
|
||||||
@@ -0,0 +1,129 @@
|
|||||||
|
"""claude code backend - drives the user's locally installed, already-signed-in
|
||||||
|
`claude` cli in headless mode (`claude -p --output-format json`).
|
||||||
|
|
||||||
|
lets deepzero run against a claude code subscription instead of a metered api
|
||||||
|
key. all subprocess handling lives in CLIAgentBackend; this module only declares
|
||||||
|
claude-specific argv and output parsing.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, ClassVar
|
||||||
|
|
||||||
|
from deepzero.engine.backends.base import BackendError, CLIAgentBackend
|
||||||
|
|
||||||
|
log = logging.getLogger("deepzero.llm.claude_code")
|
||||||
|
|
||||||
|
# the prompt carries untrusted content (decompiled malware, attacker-controlled
|
||||||
|
# strings). a completion never needs tools, so deny the ones that could turn a
|
||||||
|
# prompt injection into code execution or exfiltration.
|
||||||
|
_DEFAULT_DISALLOWED_TOOLS = (
|
||||||
|
"Bash",
|
||||||
|
"Edit",
|
||||||
|
"Write",
|
||||||
|
"NotebookEdit",
|
||||||
|
"Read",
|
||||||
|
"Glob",
|
||||||
|
"Grep",
|
||||||
|
"WebFetch",
|
||||||
|
"WebSearch",
|
||||||
|
"Agent",
|
||||||
|
"Task",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class ClaudeCodeBackend(CLIAgentBackend):
|
||||||
|
scheme: ClassVar[str] = "claude-code"
|
||||||
|
display_name: ClassVar[str] = "claude code"
|
||||||
|
binary_names: ClassVar[tuple[str, ...]] = ("claude",)
|
||||||
|
install_hint: ClassVar[str] = "install it and sign in (https://claude.com/claude-code)."
|
||||||
|
# never use --bare: it disables oauth/keychain reads, which is exactly the
|
||||||
|
# auth this backend exists to use
|
||||||
|
subscription_auth_env_blocklist: ClassVar[tuple[str, ...]] = (
|
||||||
|
"ANTHROPIC_API_KEY",
|
||||||
|
"ANTHROPIC_AUTH_TOKEN",
|
||||||
|
)
|
||||||
|
|
||||||
|
def __init__(self, model: str, **kwargs: Any):
|
||||||
|
super().__init__(model, **kwargs)
|
||||||
|
self.max_turns = int(kwargs.get("max_turns", 1))
|
||||||
|
disallowed = kwargs.get("disallowed_tools", _DEFAULT_DISALLOWED_TOOLS)
|
||||||
|
self.disallowed_tools: tuple[str, ...] = tuple(disallowed) if disallowed else ()
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def extra_search_paths(cls) -> list[Path]:
|
||||||
|
home = Path.home()
|
||||||
|
paths = [
|
||||||
|
home / ".claude" / "local" / "claude",
|
||||||
|
home / ".local" / "bin" / "claude",
|
||||||
|
home / "bin" / "claude",
|
||||||
|
]
|
||||||
|
if sys.platform == "win32":
|
||||||
|
appdata = os.environ.get("APPDATA", str(home / "AppData" / "Roaming"))
|
||||||
|
localappdata = os.environ.get("LOCALAPPDATA", str(home / "AppData" / "Local"))
|
||||||
|
for base in (Path(appdata) / "npm", Path(localappdata) / "claude" / "bin"):
|
||||||
|
paths += [base / "claude.cmd", base / "claude.exe", base / "claude"]
|
||||||
|
paths.append(home / ".local" / "bin" / "claude.exe")
|
||||||
|
return paths
|
||||||
|
|
||||||
|
def build_argv(self, system: str) -> tuple[list[str], str]:
|
||||||
|
argv = [str(self._binary), "-p", "--output-format", "json"]
|
||||||
|
|
||||||
|
if self.alias:
|
||||||
|
argv += ["--model", self.alias]
|
||||||
|
if self.max_turns > 0:
|
||||||
|
argv += ["--max-turns", str(self.max_turns)]
|
||||||
|
if self.disallowed_tools:
|
||||||
|
argv += ["--disallowed-tools", ",".join(self.disallowed_tools)]
|
||||||
|
# ignore any globally configured mcp servers - a completion needs none,
|
||||||
|
# and they would widen the blast radius of injected content
|
||||||
|
argv.append("--strict-mcp-config")
|
||||||
|
|
||||||
|
inline_system = ""
|
||||||
|
if system:
|
||||||
|
if len(system) <= self.max_argv_text:
|
||||||
|
argv += ["--append-system-prompt", system]
|
||||||
|
else:
|
||||||
|
inline_system = system
|
||||||
|
return argv, inline_system
|
||||||
|
|
||||||
|
def parse_output(self, returncode: int, stdout: str, stderr: str) -> str:
|
||||||
|
payload: dict[str, Any] | None = None
|
||||||
|
text = stdout.strip()
|
||||||
|
if text:
|
||||||
|
try:
|
||||||
|
parsed = json.loads(text)
|
||||||
|
if isinstance(parsed, dict):
|
||||||
|
payload = parsed
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
payload = None
|
||||||
|
|
||||||
|
if payload is None:
|
||||||
|
detail = (stderr or stdout or "").strip()[-500:]
|
||||||
|
if returncode != 0:
|
||||||
|
raise self.classify_error(detail)
|
||||||
|
raise BackendError(f"could not parse claude code output: {detail}")
|
||||||
|
|
||||||
|
result = payload.get("result")
|
||||||
|
# the cli reports subtype "success" even on failures - is_error is authoritative
|
||||||
|
if payload.get("is_error") or returncode != 0:
|
||||||
|
detail = result if isinstance(result, str) and result else (stderr or "").strip()
|
||||||
|
raise self.classify_error(detail or "unknown error", payload.get("api_error_status"))
|
||||||
|
|
||||||
|
if not isinstance(result, str):
|
||||||
|
raise BackendError("claude code returned no result text")
|
||||||
|
|
||||||
|
usage = payload.get("usage") or {}
|
||||||
|
log.debug(
|
||||||
|
"claude code ok (in=%s out=%s cost=%s session=%s)",
|
||||||
|
usage.get("input_tokens"),
|
||||||
|
usage.get("output_tokens"),
|
||||||
|
payload.get("total_cost_usd"),
|
||||||
|
payload.get("session_id"),
|
||||||
|
)
|
||||||
|
return result
|
||||||
@@ -0,0 +1,87 @@
|
|||||||
|
"""litellm backend - api-key access to every endpoint litellm supports.
|
||||||
|
|
||||||
|
registered as the default, so any model string whose scheme is not claimed by a
|
||||||
|
more specific backend lands here ("openai/gpt-4o", "gemini/...", "gpt-4o", ...).
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from typing import Any, ClassVar
|
||||||
|
|
||||||
|
from deepzero.engine.backends.base import LLMBackend, _NeverRaised
|
||||||
|
|
||||||
|
log = logging.getLogger("deepzero.llm.litellm")
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_exc(obj: Any, name: str) -> type[BaseException]:
|
||||||
|
cls = getattr(obj, name, None)
|
||||||
|
try:
|
||||||
|
if isinstance(cls, type) and issubclass(cls, BaseException):
|
||||||
|
return cls
|
||||||
|
except TypeError:
|
||||||
|
pass
|
||||||
|
return _NeverRaised
|
||||||
|
|
||||||
|
|
||||||
|
class LiteLLMBackend(LLMBackend):
|
||||||
|
display_name: ClassVar[str] = "litellm"
|
||||||
|
# litellm surfaces auth problems as APIError subclasses that are also in the
|
||||||
|
# retryable set, so there is no distinct fail-fast class to declare
|
||||||
|
non_retryable_errors: ClassVar[tuple[type[BaseException], ...]] = (_NeverRaised,)
|
||||||
|
|
||||||
|
def __init__(self, model: str, **kwargs: Any):
|
||||||
|
super().__init__(model, **kwargs)
|
||||||
|
try:
|
||||||
|
import litellm
|
||||||
|
except ImportError as exc:
|
||||||
|
raise ImportError(
|
||||||
|
"litellm is required for LLM support. install with: pip install litellm"
|
||||||
|
) from exc
|
||||||
|
|
||||||
|
self.litellm = litellm
|
||||||
|
# suppress litellm's noisy logging and traceback spam
|
||||||
|
litellm.suppress_debug_info = True
|
||||||
|
logging.getLogger("litellm").setLevel(logging.CRITICAL)
|
||||||
|
|
||||||
|
# instance-level overrides of the class retry roles, with safe fallbacks
|
||||||
|
# for test mocks that lack the real exception classes
|
||||||
|
self.rate_limit_error = _resolve_exc(litellm, "RateLimitError")
|
||||||
|
self.context_window_error = _resolve_exc(litellm, "ContextWindowExceededError")
|
||||||
|
api_errors = tuple(
|
||||||
|
_resolve_exc(litellm, name) for name in ("APIConnectionError", "APIError")
|
||||||
|
)
|
||||||
|
self.retryable_errors = api_errors + (OSError, ValueError, RuntimeError)
|
||||||
|
|
||||||
|
def raw_complete(self, messages: list[dict[str, str]], **kwargs: Any) -> str:
|
||||||
|
response = self.litellm.completion(model=self.model, messages=messages, **kwargs)
|
||||||
|
return response.choices[0].message.content or ""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def validate_model(cls, model: str) -> list[str]:
|
||||||
|
try:
|
||||||
|
import litellm
|
||||||
|
except ImportError:
|
||||||
|
return ["LLM configured, but 'litellm' framework is not installed"]
|
||||||
|
|
||||||
|
env_state = litellm.validate_environment(model=model)
|
||||||
|
if not env_state.get("keys_in_environment", True):
|
||||||
|
missing_keys = env_state.get("missing_keys", [])
|
||||||
|
if missing_keys:
|
||||||
|
return [
|
||||||
|
f"LLM backend '{model}' missing credentials in environment. "
|
||||||
|
f"Need: {missing_keys}"
|
||||||
|
]
|
||||||
|
return []
|
||||||
|
|
||||||
|
@property
|
||||||
|
def provider_name(self) -> str:
|
||||||
|
if "/" in self.model:
|
||||||
|
return self.model.split("/")[0]
|
||||||
|
return "unknown"
|
||||||
|
|
||||||
|
@property
|
||||||
|
def model_name(self) -> str:
|
||||||
|
if "/" in self.model:
|
||||||
|
return self.model.split("/", 1)[1]
|
||||||
|
return self.model
|
||||||
@@ -0,0 +1,77 @@
|
|||||||
|
"""llm backend registry.
|
||||||
|
|
||||||
|
model strings are `<scheme>` or `<scheme>/<model>`. a registered scheme selects
|
||||||
|
that backend; anything unrecognised falls through to the default backend
|
||||||
|
(litellm), which covers every api-key endpoint it supports.
|
||||||
|
|
||||||
|
claude-code -> ClaudeCodeBackend (default model)
|
||||||
|
claude-code/sonnet -> ClaudeCodeBackend (sonnet)
|
||||||
|
openai/gpt-4o -> default backend (litellm)
|
||||||
|
gpt-4o -> default backend (litellm)
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from deepzero.engine.backends.base import LLMBackend
|
||||||
|
|
||||||
|
_BACKEND_REGISTRY: dict[str, type[LLMBackend]] = {}
|
||||||
|
_DEFAULT_BACKEND: type[LLMBackend] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def register_backend(cls: type[LLMBackend], *, default: bool = False) -> type[LLMBackend]:
|
||||||
|
"""register a backend under its `scheme` (and any `aliases`).
|
||||||
|
|
||||||
|
usable as a decorator. pass default=True for the fallback backend that
|
||||||
|
handles every unregistered scheme.
|
||||||
|
"""
|
||||||
|
global _DEFAULT_BACKEND
|
||||||
|
|
||||||
|
if default:
|
||||||
|
_DEFAULT_BACKEND = cls
|
||||||
|
|
||||||
|
if cls.scheme:
|
||||||
|
_BACKEND_REGISTRY[cls.scheme.lower()] = cls
|
||||||
|
for alias in cls.aliases:
|
||||||
|
_BACKEND_REGISTRY[alias.lower()] = cls
|
||||||
|
elif not default:
|
||||||
|
raise ValueError(f"{cls.__name__} must define a 'scheme' to be registered")
|
||||||
|
|
||||||
|
return cls
|
||||||
|
|
||||||
|
|
||||||
|
def get_registered_backends() -> dict[str, type[LLMBackend]]:
|
||||||
|
return dict(_BACKEND_REGISTRY)
|
||||||
|
|
||||||
|
|
||||||
|
def model_scheme(model: str) -> str:
|
||||||
|
m = (model or "").strip().lower()
|
||||||
|
return m.split("/", 1)[0] if "/" in m else m
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_backend_class(model: str) -> type[LLMBackend]:
|
||||||
|
"""map a model string to its backend class."""
|
||||||
|
backend = _BACKEND_REGISTRY.get(model_scheme(model))
|
||||||
|
if backend is not None:
|
||||||
|
return backend
|
||||||
|
|
||||||
|
if _DEFAULT_BACKEND is None:
|
||||||
|
raise ValueError(
|
||||||
|
f"no backend can handle model '{model}' and no default backend is registered. "
|
||||||
|
f"known schemes: {sorted(_BACKEND_REGISTRY)}"
|
||||||
|
)
|
||||||
|
return _DEFAULT_BACKEND
|
||||||
|
|
||||||
|
|
||||||
|
def create_backend(model: str, **kwargs: Any) -> LLMBackend:
|
||||||
|
return resolve_backend_class(model)(model, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
def validate_model_binding(model: str) -> list[str]:
|
||||||
|
"""pre-flight check for a model string, delegated to its backend."""
|
||||||
|
try:
|
||||||
|
backend_cls = resolve_backend_class(model)
|
||||||
|
except ValueError as exc:
|
||||||
|
return [str(exc)]
|
||||||
|
return backend_cls.validate_model(model)
|
||||||
+26
-56
@@ -4,56 +4,27 @@ import logging
|
|||||||
import time
|
import time
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
|
from deepzero.engine.backends import create_backend
|
||||||
|
|
||||||
log = logging.getLogger("deepzero.llm")
|
log = logging.getLogger("deepzero.llm")
|
||||||
|
|
||||||
|
|
||||||
# sentinel exception that is never raised, used as a safe fallback
|
|
||||||
# when litellm exception classes are unavailable (e.g. in mock envs)
|
|
||||||
class _NeverRaised(Exception):
|
|
||||||
pass
|
|
||||||
|
|
||||||
|
|
||||||
def _resolve_exc(obj: Any, name: str) -> type[BaseException]:
|
|
||||||
cls = getattr(obj, name, None)
|
|
||||||
try:
|
|
||||||
if isinstance(cls, type) and issubclass(cls, BaseException):
|
|
||||||
return cls
|
|
||||||
except TypeError:
|
|
||||||
pass
|
|
||||||
return _NeverRaised
|
|
||||||
|
|
||||||
|
|
||||||
class LLMProvider:
|
class LLMProvider:
|
||||||
# litellm-backed llm provider with adaptive retry and backoff
|
# llm provider with adaptive retry and backoff.
|
||||||
|
#
|
||||||
|
# the backend is resolved from the model string by the backend registry
|
||||||
|
# (see deepzero.engine.backends). this class knows nothing about any
|
||||||
|
# specific backend - it only drives the retry roles each one declares.
|
||||||
|
|
||||||
def __init__(self, model: str, **kwargs: Any):
|
def __init__(self, model: str, **kwargs: Any):
|
||||||
self.model = model
|
self.model = model
|
||||||
self.default_kwargs = kwargs
|
self.default_kwargs = kwargs
|
||||||
self._ensure_litellm()
|
self.backend = create_backend(model, **kwargs)
|
||||||
|
|
||||||
def _ensure_litellm(self) -> None:
|
@property
|
||||||
try:
|
def _litellm(self) -> Any:
|
||||||
import litellm
|
# retained for backward compatibility with existing callers/tests
|
||||||
|
return getattr(self.backend, "litellm", None)
|
||||||
self._litellm = litellm
|
|
||||||
# suppress litellm's noisy logging and traceback spam
|
|
||||||
litellm.suppress_debug_info = True
|
|
||||||
logging.getLogger("litellm").setLevel(logging.CRITICAL)
|
|
||||||
|
|
||||||
# capture exception classes with safe fallbacks for test mocks
|
|
||||||
self._rate_limit_error = _resolve_exc(litellm, "RateLimitError")
|
|
||||||
self._context_window_error = _resolve_exc(litellm, "ContextWindowExceededError")
|
|
||||||
|
|
||||||
# build the retryable error tuple once at init
|
|
||||||
api_errors = tuple(
|
|
||||||
_resolve_exc(litellm, name) for name in ("APIConnectionError", "APIError")
|
|
||||||
)
|
|
||||||
self._retryable_errors = api_errors + (OSError, ValueError, RuntimeError)
|
|
||||||
|
|
||||||
except ImportError as exc:
|
|
||||||
raise ImportError(
|
|
||||||
"litellm is required for LLM support. install with: pip install litellm"
|
|
||||||
) from exc
|
|
||||||
|
|
||||||
def complete(
|
def complete(
|
||||||
self,
|
self,
|
||||||
@@ -66,23 +37,22 @@ class LLMProvider:
|
|||||||
) -> str:
|
) -> str:
|
||||||
"""send messages to the llm and return the response text.
|
"""send messages to the llm and return the response text.
|
||||||
handles rate limiting with adaptive backoff."""
|
handles rate limiting with adaptive backoff."""
|
||||||
|
backend = self.backend
|
||||||
|
# forward all options; each backend uses what applies (litellm passes
|
||||||
|
# generation kwargs to the api, cli backends read controls like timeout
|
||||||
|
# and ignore the rest) so e.g. timeout= is not silently dropped
|
||||||
merged = {**self.default_kwargs, **kwargs}
|
merged = {**self.default_kwargs, **kwargs}
|
||||||
backoff = initial_backoff
|
backoff = initial_backoff
|
||||||
|
|
||||||
for attempt in range(max_retries + 1):
|
for attempt in range(max_retries + 1):
|
||||||
try:
|
try:
|
||||||
response = self._litellm.completion(
|
content = backend.raw_complete(messages, **merged)
|
||||||
model=self.model,
|
|
||||||
messages=messages,
|
|
||||||
**merged,
|
|
||||||
)
|
|
||||||
content = response.choices[0].message.content or ""
|
|
||||||
|
|
||||||
# decay backoff toward minimum on success
|
# decay backoff toward minimum on success
|
||||||
backoff = max(initial_backoff, backoff * backoff_decay)
|
backoff = max(initial_backoff, backoff * backoff_decay)
|
||||||
return content
|
return content
|
||||||
|
|
||||||
except self._rate_limit_error:
|
except backend.rate_limit_error:
|
||||||
if attempt == max_retries:
|
if attempt == max_retries:
|
||||||
raise
|
raise
|
||||||
backoff = min(max_backoff, backoff * 2.0)
|
backoff = min(max_backoff, backoff * 2.0)
|
||||||
@@ -94,11 +64,15 @@ class LLMProvider:
|
|||||||
)
|
)
|
||||||
time.sleep(backoff)
|
time.sleep(backoff)
|
||||||
|
|
||||||
except self._context_window_error:
|
except backend.context_window_error:
|
||||||
# context window errors won't be fixed by retry
|
# context window errors won't be fixed by retry
|
||||||
raise
|
raise
|
||||||
|
|
||||||
except self._retryable_errors as e:
|
except backend.non_retryable_errors:
|
||||||
|
# e.g. auth failures - surface immediately instead of burning retries
|
||||||
|
raise
|
||||||
|
|
||||||
|
except backend.retryable_errors as e:
|
||||||
if attempt == max_retries:
|
if attempt == max_retries:
|
||||||
raise
|
raise
|
||||||
wait = min(max_backoff, 2**attempt)
|
wait = min(max_backoff, 2**attempt)
|
||||||
@@ -115,12 +89,8 @@ class LLMProvider:
|
|||||||
|
|
||||||
@property
|
@property
|
||||||
def provider_name(self) -> str:
|
def provider_name(self) -> str:
|
||||||
if "/" in self.model:
|
return self.backend.provider_name
|
||||||
return self.model.split("/")[0]
|
|
||||||
return "unknown"
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def model_name(self) -> str:
|
def model_name(self) -> str:
|
||||||
if "/" in self.model:
|
return self.backend.model_name
|
||||||
return self.model.split("/", 1)[1]
|
|
||||||
return self.model
|
|
||||||
|
|||||||
@@ -271,7 +271,10 @@ def _expand_string(s: str) -> str:
|
|||||||
if ":-" in var:
|
if ":-" in var:
|
||||||
name, default = var.split(":-", 1)
|
name, default = var.split(":-", 1)
|
||||||
return os.environ.get(name, default)
|
return os.environ.get(name, default)
|
||||||
return os.environ.get(var, match.group(0))
|
# an unset var expands to empty rather than leaking the literal "${VAR}",
|
||||||
|
# which would read as a truthy value downstream and produce confusing
|
||||||
|
# errors like "ghidra not found at ${GHIDRA_INSTALL_DIR}"
|
||||||
|
return os.environ.get(var, "")
|
||||||
|
|
||||||
return re.sub(r"\$\{([^}]+)\}", _replace, s)
|
return re.sub(r"\$\{([^}]+)\}", _replace, s)
|
||||||
|
|
||||||
|
|||||||
@@ -34,21 +34,13 @@ class GenericLLM(MapProcessor):
|
|||||||
if not prompt_path.exists():
|
if not prompt_path.exists():
|
||||||
errors.append(f"Prompt template does not exist: {prompt_ref}")
|
errors.append(f"Prompt template does not exist: {prompt_ref}")
|
||||||
|
|
||||||
# structurally validate LLM bindings early
|
# structurally validate LLM bindings early - each backend checks its own
|
||||||
|
# prerequisites (api keys, cli presence, ...) via the registry
|
||||||
model = ctx.global_config.get("model")
|
model = ctx.global_config.get("model")
|
||||||
if model:
|
if model:
|
||||||
try:
|
from deepzero.engine.backends import validate_model_binding
|
||||||
import litellm
|
|
||||||
|
|
||||||
env_state = litellm.validate_environment(model=model)
|
errors.extend(validate_model_binding(model))
|
||||||
if not env_state.get("keys_in_environment", True):
|
|
||||||
missing_keys = env_state.get("missing_keys", [])
|
|
||||||
if missing_keys:
|
|
||||||
errors.append(
|
|
||||||
f"LLM backend '{model}' missing credentials in environment. Need: {missing_keys}"
|
|
||||||
)
|
|
||||||
except ImportError:
|
|
||||||
errors.append("LLM configured, but 'litellm' framework is not installed")
|
|
||||||
|
|
||||||
return errors
|
return errors
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,310 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import subprocess
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from deepzero.engine.backends import ClaudeCodeBackend
|
||||||
|
from deepzero.engine.backends.base import (
|
||||||
|
BackendAuthError,
|
||||||
|
BackendContextWindowError,
|
||||||
|
BackendError,
|
||||||
|
BackendNotFoundError,
|
||||||
|
BackendRateLimitError,
|
||||||
|
)
|
||||||
|
from deepzero.engine.llm import LLMProvider
|
||||||
|
|
||||||
|
_FIND = "deepzero.engine.backends.claude_code.ClaudeCodeBackend.find_binary"
|
||||||
|
|
||||||
|
|
||||||
|
def _backend(**kwargs) -> ClaudeCodeBackend:
|
||||||
|
kwargs.setdefault("binary", "/usr/bin/claude")
|
||||||
|
return ClaudeCodeBackend(kwargs.pop("model", "claude-code/sonnet"), **kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
def _completed(stdout: str, returncode: int = 0, stderr: str = ""):
|
||||||
|
return subprocess.CompletedProcess(args=[], returncode=returncode, stdout=stdout, stderr=stderr)
|
||||||
|
|
||||||
|
|
||||||
|
def _result_json(**overrides) -> str:
|
||||||
|
payload = {
|
||||||
|
"type": "result",
|
||||||
|
"subtype": "success",
|
||||||
|
"is_error": False,
|
||||||
|
"result": "hello from claude",
|
||||||
|
"session_id": "abc",
|
||||||
|
"total_cost_usd": 0.01,
|
||||||
|
"usage": {"input_tokens": 10, "output_tokens": 5},
|
||||||
|
}
|
||||||
|
payload.update(overrides)
|
||||||
|
return json.dumps(payload)
|
||||||
|
|
||||||
|
|
||||||
|
class TestAliasParsing:
|
||||||
|
def test_alias_parsing(self):
|
||||||
|
assert _backend(model="claude-code/sonnet").alias == "sonnet"
|
||||||
|
assert _backend(model="claude-code").alias == ""
|
||||||
|
assert _backend(model="claude-code/claude-opus-5").alias == "claude-opus-5"
|
||||||
|
|
||||||
|
|
||||||
|
class TestBinaryResolution:
|
||||||
|
def test_raises_when_not_installed(self):
|
||||||
|
with patch(_FIND, return_value=None):
|
||||||
|
with pytest.raises(BackendNotFoundError) as exc:
|
||||||
|
ClaudeCodeBackend("claude-code")
|
||||||
|
assert "not found" in str(exc.value).lower()
|
||||||
|
|
||||||
|
def test_uses_explicit_binary(self):
|
||||||
|
assert ClaudeCodeBackend("claude-code", binary="/custom/claude")._binary == "/custom/claude"
|
||||||
|
|
||||||
|
|
||||||
|
class TestMessageFlattening:
|
||||||
|
def test_single_user_message_passes_through(self):
|
||||||
|
system, prompt = ClaudeCodeBackend.split_messages([{"role": "user", "content": "analyze"}])
|
||||||
|
assert system == ""
|
||||||
|
assert prompt == "analyze"
|
||||||
|
|
||||||
|
def test_system_is_separated(self):
|
||||||
|
system, prompt = ClaudeCodeBackend.split_messages(
|
||||||
|
[{"role": "system", "content": "be terse"}, {"role": "user", "content": "hi"}]
|
||||||
|
)
|
||||||
|
assert system == "be terse"
|
||||||
|
assert prompt == "hi"
|
||||||
|
|
||||||
|
def test_multi_turn_is_rendered_inline(self):
|
||||||
|
system, prompt = ClaudeCodeBackend.split_messages(
|
||||||
|
[
|
||||||
|
{"role": "system", "content": "sys"},
|
||||||
|
{"role": "user", "content": "q1"},
|
||||||
|
{"role": "assistant", "content": "a1"},
|
||||||
|
{"role": "user", "content": "q2"},
|
||||||
|
]
|
||||||
|
)
|
||||||
|
assert system == "sys"
|
||||||
|
assert "User: q1" in prompt
|
||||||
|
assert "Assistant: a1" in prompt
|
||||||
|
assert "User: q2" in prompt
|
||||||
|
|
||||||
|
def test_empty(self):
|
||||||
|
assert ClaudeCodeBackend.split_messages([]) == ("", "")
|
||||||
|
|
||||||
|
|
||||||
|
class TestCommandConstruction:
|
||||||
|
def test_core_flags(self):
|
||||||
|
argv, inline = _backend().build_argv("")
|
||||||
|
assert argv[:4] == ["/usr/bin/claude", "-p", "--output-format", "json"]
|
||||||
|
assert "--model" in argv and "sonnet" in argv
|
||||||
|
assert "--strict-mcp-config" in argv
|
||||||
|
assert inline == ""
|
||||||
|
|
||||||
|
def test_no_model_flag_for_bare_alias(self):
|
||||||
|
argv, _ = _backend(model="claude-code").build_argv("")
|
||||||
|
assert "--model" not in argv
|
||||||
|
|
||||||
|
def test_dangerous_tools_denied_by_default(self):
|
||||||
|
argv, _ = _backend().build_argv("")
|
||||||
|
denied = argv[argv.index("--disallowed-tools") + 1]
|
||||||
|
for tool in ("Bash", "Write", "Edit", "WebFetch"):
|
||||||
|
assert tool in denied
|
||||||
|
|
||||||
|
def test_max_turns_limits_agent_loop(self):
|
||||||
|
argv, _ = _backend().build_argv("")
|
||||||
|
assert "--max-turns" in argv
|
||||||
|
|
||||||
|
def test_small_system_prompt_goes_to_argv(self):
|
||||||
|
argv, inline = _backend().build_argv("be terse")
|
||||||
|
assert "--append-system-prompt" in argv
|
||||||
|
assert "be terse" in argv
|
||||||
|
assert inline == ""
|
||||||
|
|
||||||
|
def test_huge_system_prompt_moves_to_stdin(self):
|
||||||
|
big = "x" * 20000
|
||||||
|
argv, inline = _backend().build_argv(big)
|
||||||
|
assert "--append-system-prompt" not in argv
|
||||||
|
assert inline == big
|
||||||
|
|
||||||
|
def test_never_uses_bare_flag(self):
|
||||||
|
# --bare disables oauth/keychain reads, which would break subscription auth
|
||||||
|
argv, _ = _backend().build_argv("sys")
|
||||||
|
assert "--bare" not in argv
|
||||||
|
|
||||||
|
|
||||||
|
class TestEnvHandling:
|
||||||
|
def test_api_key_hidden_by_default(self):
|
||||||
|
with patch.dict("os.environ", {"ANTHROPIC_API_KEY": "sk-ant-x"}, clear=False):
|
||||||
|
env = _backend().build_env()
|
||||||
|
assert "ANTHROPIC_API_KEY" not in env
|
||||||
|
|
||||||
|
def test_api_key_preserved_when_opted_out(self):
|
||||||
|
with patch.dict("os.environ", {"ANTHROPIC_API_KEY": "sk-ant-x"}, clear=False):
|
||||||
|
env = _backend(prefer_subscription_auth=False).build_env()
|
||||||
|
assert env["ANTHROPIC_API_KEY"] == "sk-ant-x"
|
||||||
|
|
||||||
|
|
||||||
|
class TestResponseParsing:
|
||||||
|
def test_success(self):
|
||||||
|
with patch("subprocess.run", return_value=_completed(_result_json())) as mock_run:
|
||||||
|
out = _backend().raw_complete([{"role": "user", "content": "hi"}])
|
||||||
|
assert out == "hello from claude"
|
||||||
|
# the prompt travels via stdin, never argv - no shell/argv injection surface
|
||||||
|
assert mock_run.call_args.kwargs["input"] == "hi"
|
||||||
|
|
||||||
|
def test_is_error_true_despite_success_subtype(self):
|
||||||
|
# the cli reports subtype "success" even on failures; is_error is authoritative
|
||||||
|
body = _result_json(is_error=True, subtype="success", result="boom")
|
||||||
|
with patch("subprocess.run", return_value=_completed(body)):
|
||||||
|
with pytest.raises(BackendError):
|
||||||
|
_backend().raw_complete([{"role": "user", "content": "hi"}])
|
||||||
|
|
||||||
|
def test_auth_error_is_classified(self):
|
||||||
|
body = _result_json(
|
||||||
|
is_error=True,
|
||||||
|
api_error_status=401,
|
||||||
|
result="Failed to authenticate. API Error: 401 OAuth access token has been revoked.",
|
||||||
|
)
|
||||||
|
with patch("subprocess.run", return_value=_completed(body)):
|
||||||
|
with pytest.raises(BackendAuthError) as exc:
|
||||||
|
_backend().raw_complete([{"role": "user", "content": "hi"}])
|
||||||
|
assert "not authenticated" in str(exc.value)
|
||||||
|
|
||||||
|
def test_rate_limit_is_classified(self):
|
||||||
|
body = _result_json(is_error=True, api_error_status=429, result="rate limit exceeded")
|
||||||
|
with patch("subprocess.run", return_value=_completed(body)):
|
||||||
|
with pytest.raises(BackendRateLimitError):
|
||||||
|
_backend().raw_complete([{"role": "user", "content": "hi"}])
|
||||||
|
|
||||||
|
def test_usage_limit_text_is_rate_limit(self):
|
||||||
|
body = _result_json(is_error=True, result="Claude usage limit reached")
|
||||||
|
with patch("subprocess.run", return_value=_completed(body)):
|
||||||
|
with pytest.raises(BackendRateLimitError):
|
||||||
|
_backend().raw_complete([{"role": "user", "content": "hi"}])
|
||||||
|
|
||||||
|
def test_context_window_is_classified(self):
|
||||||
|
body = _result_json(is_error=True, result="prompt is too long: context window exceeded")
|
||||||
|
with patch("subprocess.run", return_value=_completed(body)):
|
||||||
|
with pytest.raises(BackendContextWindowError):
|
||||||
|
_backend().raw_complete([{"role": "user", "content": "hi"}])
|
||||||
|
|
||||||
|
def test_unparseable_output(self):
|
||||||
|
with patch("subprocess.run", return_value=_completed("not json")):
|
||||||
|
with pytest.raises(BackendError):
|
||||||
|
_backend().raw_complete([{"role": "user", "content": "hi"}])
|
||||||
|
|
||||||
|
def test_nonzero_exit_with_stderr(self):
|
||||||
|
with patch("subprocess.run", return_value=_completed("", 1, "command failed")):
|
||||||
|
with pytest.raises(BackendError) as exc:
|
||||||
|
_backend().raw_complete([{"role": "user", "content": "hi"}])
|
||||||
|
assert "command failed" in str(exc.value)
|
||||||
|
|
||||||
|
def test_timeout(self):
|
||||||
|
with patch(
|
||||||
|
"subprocess.run", side_effect=subprocess.TimeoutExpired(cmd="claude", timeout=5)
|
||||||
|
):
|
||||||
|
with pytest.raises(BackendError) as exc:
|
||||||
|
_backend().raw_complete([{"role": "user", "content": "hi"}])
|
||||||
|
assert "timed out" in str(exc.value)
|
||||||
|
|
||||||
|
def test_launch_failure(self):
|
||||||
|
with patch("subprocess.run", side_effect=OSError("no exec")):
|
||||||
|
with pytest.raises(BackendError):
|
||||||
|
_backend().raw_complete([{"role": "user", "content": "hi"}])
|
||||||
|
|
||||||
|
|
||||||
|
class TestProviderIntegration:
|
||||||
|
def test_provider_selects_claude_code_without_litellm(self, monkeypatch):
|
||||||
|
import sys
|
||||||
|
|
||||||
|
# litellm intentionally unavailable - the claude code path must not need it
|
||||||
|
monkeypatch.setitem(sys.modules, "litellm", None)
|
||||||
|
with patch(_FIND, return_value="/usr/bin/claude"):
|
||||||
|
provider = LLMProvider("claude-code/sonnet")
|
||||||
|
|
||||||
|
assert isinstance(provider.backend, ClaudeCodeBackend)
|
||||||
|
assert provider.provider_name == "claude-code"
|
||||||
|
assert provider.model_name == "sonnet"
|
||||||
|
|
||||||
|
def test_provider_completes_through_backend(self):
|
||||||
|
with patch(_FIND, return_value="/usr/bin/claude"):
|
||||||
|
provider = LLMProvider("claude-code")
|
||||||
|
with patch("subprocess.run", return_value=_completed(_result_json())):
|
||||||
|
assert provider.complete([{"role": "user", "content": "hi"}]) == "hello from claude"
|
||||||
|
|
||||||
|
def test_generation_kwargs_not_forwarded_to_cli(self):
|
||||||
|
# litellm-style kwargs must not leak into the subprocess call
|
||||||
|
with patch(_FIND, return_value="/usr/bin/claude"):
|
||||||
|
provider = LLMProvider("claude-code", temperature=0.7)
|
||||||
|
with patch("subprocess.run", return_value=_completed(_result_json())) as mock_run:
|
||||||
|
provider.complete([{"role": "user", "content": "hi"}])
|
||||||
|
assert "temperature" not in mock_run.call_args.kwargs
|
||||||
|
|
||||||
|
def test_rate_limit_retries_then_succeeds(self):
|
||||||
|
with patch(_FIND, return_value="/usr/bin/claude"):
|
||||||
|
provider = LLMProvider("claude-code")
|
||||||
|
|
||||||
|
failure = _completed(_result_json(is_error=True, api_error_status=429, result="rate limit"))
|
||||||
|
success = _completed(_result_json())
|
||||||
|
with patch("subprocess.run", side_effect=[failure, success]):
|
||||||
|
with patch("time.sleep"):
|
||||||
|
assert provider.complete([{"role": "user", "content": "hi"}]) == "hello from claude"
|
||||||
|
|
||||||
|
def test_auth_error_fails_fast_without_retrying(self):
|
||||||
|
with patch(_FIND, return_value="/usr/bin/claude"):
|
||||||
|
provider = LLMProvider("claude-code")
|
||||||
|
|
||||||
|
failure = _completed(
|
||||||
|
_result_json(is_error=True, api_error_status=401, result="authentication failed")
|
||||||
|
)
|
||||||
|
with patch("subprocess.run", return_value=failure) as mock_run:
|
||||||
|
with patch("time.sleep"):
|
||||||
|
with pytest.raises(BackendAuthError):
|
||||||
|
provider.complete([{"role": "user", "content": "hi"}], max_retries=3)
|
||||||
|
|
||||||
|
assert mock_run.call_count == 1
|
||||||
|
|
||||||
|
def test_context_window_error_not_retried(self):
|
||||||
|
with patch(_FIND, return_value="/usr/bin/claude"):
|
||||||
|
provider = LLMProvider("claude-code")
|
||||||
|
|
||||||
|
failure = _completed(_result_json(is_error=True, result="context window exceeded"))
|
||||||
|
with patch("subprocess.run", return_value=failure) as mock_run:
|
||||||
|
with patch("time.sleep"):
|
||||||
|
with pytest.raises(BackendContextWindowError):
|
||||||
|
provider.complete([{"role": "user", "content": "hi"}], max_retries=3)
|
||||||
|
|
||||||
|
assert mock_run.call_count == 1
|
||||||
|
|
||||||
|
|
||||||
|
class TestValidation:
|
||||||
|
def _stage(self):
|
||||||
|
from deepzero.engine.stage import StageSpec
|
||||||
|
from deepzero.stages.llm import GenericLLM
|
||||||
|
|
||||||
|
return GenericLLM(StageSpec(name="assess", processor="generic_llm", config={}))
|
||||||
|
|
||||||
|
def _validate(self, model: str) -> list[str]:
|
||||||
|
from deepzero.engine.stage import ProcessorContext
|
||||||
|
|
||||||
|
ctx = ProcessorContext(
|
||||||
|
pipeline_dir=__import__("pathlib").Path("."),
|
||||||
|
global_config={"model": model},
|
||||||
|
llm=None,
|
||||||
|
)
|
||||||
|
# only inspect llm-binding errors, not the unrelated prompt-config error
|
||||||
|
return [e for e in self._stage().validate(ctx) if "backend" in e.lower()]
|
||||||
|
|
||||||
|
def test_claude_code_binding_ok_when_installed(self):
|
||||||
|
with patch(_FIND, return_value="/usr/bin/claude"):
|
||||||
|
assert self._validate("claude-code/sonnet") == []
|
||||||
|
|
||||||
|
def test_claude_code_binding_reports_missing_cli(self):
|
||||||
|
with patch(_FIND, return_value=None):
|
||||||
|
errors = self._validate("claude-code")
|
||||||
|
assert errors and "claude code cli not found" in errors[0].lower()
|
||||||
|
|
||||||
|
def test_claude_code_binding_does_not_require_api_keys(self):
|
||||||
|
# no ANTHROPIC_API_KEY in env, yet validation passes
|
||||||
|
with patch.dict("os.environ", {}, clear=True):
|
||||||
|
with patch(_FIND, return_value="/usr/bin/claude"):
|
||||||
|
assert self._validate("claude-code") == []
|
||||||
@@ -0,0 +1,196 @@
|
|||||||
|
"""proves the extensibility contract: a new agent-cli backend can be added by
|
||||||
|
subclassing CLIAgentBackend and registering it, with no changes to LLMProvider,
|
||||||
|
the registry internals, the pipeline, or the stages.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import subprocess
|
||||||
|
from typing import ClassVar
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from deepzero.engine.backends import (
|
||||||
|
ClaudeCodeBackend,
|
||||||
|
CLIAgentBackend,
|
||||||
|
LiteLLMBackend,
|
||||||
|
get_registered_backends,
|
||||||
|
model_scheme,
|
||||||
|
register_backend,
|
||||||
|
resolve_backend_class,
|
||||||
|
validate_model_binding,
|
||||||
|
)
|
||||||
|
from deepzero.engine.backends.base import BackendError
|
||||||
|
from deepzero.engine.backends.registry import _BACKEND_REGISTRY
|
||||||
|
from deepzero.engine.llm import LLMProvider
|
||||||
|
|
||||||
|
|
||||||
|
# a stand-in for a future backend (codex, gemini cli, ...). the only things it
|
||||||
|
# has to supply are argv construction and output parsing.
|
||||||
|
class FakeAgentBackend(CLIAgentBackend):
|
||||||
|
scheme: ClassVar[str] = "faketool"
|
||||||
|
aliases: ClassVar[tuple[str, ...]] = ("ft",)
|
||||||
|
display_name: ClassVar[str] = "fake tool"
|
||||||
|
binary_names: ClassVar[tuple[str, ...]] = ("faketool",)
|
||||||
|
subscription_auth_env_blocklist: ClassVar[tuple[str, ...]] = ("FAKETOOL_API_KEY",)
|
||||||
|
|
||||||
|
def build_argv(self, system: str) -> tuple[list[str], str]:
|
||||||
|
argv = [str(self._binary), "exec", "--json"]
|
||||||
|
if self.alias:
|
||||||
|
argv += ["--model", self.alias]
|
||||||
|
if system:
|
||||||
|
argv += ["--system", system]
|
||||||
|
return argv, ""
|
||||||
|
|
||||||
|
def parse_output(self, returncode: int, stdout: str, stderr: str) -> str:
|
||||||
|
try:
|
||||||
|
payload = json.loads(stdout)
|
||||||
|
except json.JSONDecodeError as exc:
|
||||||
|
raise BackendError(f"bad output: {stdout[:100]}") from exc
|
||||||
|
if payload.get("error"):
|
||||||
|
raise self.classify_error(payload["error"], payload.get("status"))
|
||||||
|
return payload["text"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def registered_fake():
|
||||||
|
"""register the fake backend, then restore the registry."""
|
||||||
|
saved = dict(_BACKEND_REGISTRY)
|
||||||
|
register_backend(FakeAgentBackend)
|
||||||
|
try:
|
||||||
|
yield FakeAgentBackend
|
||||||
|
finally:
|
||||||
|
_BACKEND_REGISTRY.clear()
|
||||||
|
_BACKEND_REGISTRY.update(saved)
|
||||||
|
|
||||||
|
|
||||||
|
def _completed(stdout: str, returncode: int = 0, stderr: str = ""):
|
||||||
|
return subprocess.CompletedProcess(args=[], returncode=returncode, stdout=stdout, stderr=stderr)
|
||||||
|
|
||||||
|
|
||||||
|
class TestSchemeParsing:
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"model,expected",
|
||||||
|
[
|
||||||
|
("claude-code", "claude-code"),
|
||||||
|
("claude-code/sonnet", "claude-code"),
|
||||||
|
("CLAUDE-CODE/Sonnet", "claude-code"),
|
||||||
|
("openai/gpt-4o", "openai"),
|
||||||
|
("gpt-4o", "gpt-4o"),
|
||||||
|
("", ""),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_scheme(self, model, expected):
|
||||||
|
assert model_scheme(model) == expected
|
||||||
|
|
||||||
|
|
||||||
|
class TestBuiltInResolution:
|
||||||
|
@pytest.mark.parametrize("model", ["claude-code", "claude-code/sonnet", "claude-code/opus"])
|
||||||
|
def test_claude_code_schemes(self, model):
|
||||||
|
assert resolve_backend_class(model) is ClaudeCodeBackend
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"model", ["openai/gpt-4o", "anthropic/claude-opus-5", "gemini/pro", "gpt-4o", "test"]
|
||||||
|
)
|
||||||
|
def test_unclaimed_schemes_fall_back_to_litellm(self, model):
|
||||||
|
assert resolve_backend_class(model) is LiteLLMBackend
|
||||||
|
|
||||||
|
def test_claude_code_prefix_is_not_over_matched(self):
|
||||||
|
# "claude-code-extra" is a different scheme, not claude code
|
||||||
|
assert resolve_backend_class("claude-code-extra/x") is LiteLLMBackend
|
||||||
|
|
||||||
|
def test_registry_lists_known_schemes(self):
|
||||||
|
assert "claude-code" in get_registered_backends()
|
||||||
|
|
||||||
|
|
||||||
|
class TestAddingANewBackend:
|
||||||
|
def test_scheme_and_alias_resolve(self, registered_fake):
|
||||||
|
assert resolve_backend_class("faketool") is FakeAgentBackend
|
||||||
|
assert resolve_backend_class("faketool/big-model") is FakeAgentBackend
|
||||||
|
assert resolve_backend_class("ft/big-model") is FakeAgentBackend
|
||||||
|
|
||||||
|
def test_does_not_disturb_existing_backends(self, registered_fake):
|
||||||
|
assert resolve_backend_class("claude-code") is ClaudeCodeBackend
|
||||||
|
assert resolve_backend_class("openai/gpt-4o") is LiteLLMBackend
|
||||||
|
|
||||||
|
def test_provider_drives_it_end_to_end(self, registered_fake):
|
||||||
|
with patch.object(FakeAgentBackend, "find_binary", return_value="/usr/bin/faketool"):
|
||||||
|
provider = LLMProvider("faketool/big-model")
|
||||||
|
|
||||||
|
assert provider.provider_name == "faketool"
|
||||||
|
assert provider.model_name == "big-model"
|
||||||
|
|
||||||
|
with patch("subprocess.run", return_value=_completed('{"text": "fake says hi"}')) as run:
|
||||||
|
assert provider.complete([{"role": "user", "content": "hi"}]) == "fake says hi"
|
||||||
|
|
||||||
|
argv = run.call_args.args[0]
|
||||||
|
assert argv == ["/usr/bin/faketool", "exec", "--json", "--model", "big-model"]
|
||||||
|
assert run.call_args.kwargs["input"] == "hi"
|
||||||
|
|
||||||
|
def test_inherits_retry_semantics_for_free(self, registered_fake):
|
||||||
|
with patch.object(FakeAgentBackend, "find_binary", return_value="/usr/bin/faketool"):
|
||||||
|
provider = LLMProvider("faketool")
|
||||||
|
|
||||||
|
limited = _completed(json.dumps({"error": "rate limit hit", "status": 429}))
|
||||||
|
ok = _completed('{"text": "recovered"}')
|
||||||
|
with patch("subprocess.run", side_effect=[limited, ok]):
|
||||||
|
with patch("time.sleep"):
|
||||||
|
assert provider.complete([{"role": "user", "content": "hi"}]) == "recovered"
|
||||||
|
|
||||||
|
def test_inherits_auth_fail_fast_for_free(self, registered_fake):
|
||||||
|
with patch.object(FakeAgentBackend, "find_binary", return_value="/usr/bin/faketool"):
|
||||||
|
provider = LLMProvider("faketool")
|
||||||
|
|
||||||
|
denied = _completed(json.dumps({"error": "unauthorized", "status": 401}))
|
||||||
|
with patch("subprocess.run", return_value=denied) as run:
|
||||||
|
with patch("time.sleep"):
|
||||||
|
with pytest.raises(BackendError):
|
||||||
|
provider.complete([{"role": "user", "content": "hi"}], max_retries=3)
|
||||||
|
assert run.call_count == 1
|
||||||
|
|
||||||
|
def test_inherits_env_sanitization_for_free(self, registered_fake):
|
||||||
|
with patch.dict("os.environ", {"FAKETOOL_API_KEY": "secret"}, clear=False):
|
||||||
|
with patch.object(FakeAgentBackend, "find_binary", return_value="/usr/bin/faketool"):
|
||||||
|
env = FakeAgentBackend("faketool").build_env()
|
||||||
|
assert "FAKETOOL_API_KEY" not in env
|
||||||
|
|
||||||
|
def test_inherits_validation_wiring_for_free(self, registered_fake):
|
||||||
|
with patch.object(FakeAgentBackend, "find_binary", return_value=None):
|
||||||
|
errors = validate_model_binding("faketool")
|
||||||
|
assert errors and "fake tool cli not found" in errors[0].lower()
|
||||||
|
|
||||||
|
with patch.object(FakeAgentBackend, "find_binary", return_value="/usr/bin/faketool"):
|
||||||
|
assert validate_model_binding("faketool") == []
|
||||||
|
|
||||||
|
|
||||||
|
class TestRegistrationRules:
|
||||||
|
def test_backend_without_scheme_is_rejected(self):
|
||||||
|
class Anonymous(CLIAgentBackend):
|
||||||
|
def build_argv(self, system):
|
||||||
|
return [], ""
|
||||||
|
|
||||||
|
def parse_output(self, returncode, stdout, stderr):
|
||||||
|
return ""
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="scheme"):
|
||||||
|
register_backend(Anonymous)
|
||||||
|
|
||||||
|
def test_register_returns_class_for_decorator_use(self):
|
||||||
|
saved = dict(_BACKEND_REGISTRY)
|
||||||
|
try:
|
||||||
|
|
||||||
|
class Decorated(CLIAgentBackend):
|
||||||
|
scheme = "decorated"
|
||||||
|
|
||||||
|
def build_argv(self, system):
|
||||||
|
return [], ""
|
||||||
|
|
||||||
|
def parse_output(self, returncode, stdout, stderr):
|
||||||
|
return ""
|
||||||
|
|
||||||
|
assert register_backend(Decorated) is Decorated
|
||||||
|
finally:
|
||||||
|
_BACKEND_REGISTRY.clear()
|
||||||
|
_BACKEND_REGISTRY.update(saved)
|
||||||
@@ -1,4 +1,9 @@
|
|||||||
from deepzero.engine.stage import StageSpec
|
import asyncio
|
||||||
|
import json
|
||||||
|
from pathlib import Path
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
from deepzero.engine.stage import ProcessorContext, ProcessorEntry, StageSpec
|
||||||
from processors.semgrep_scanner.semgrep_scanner import SemgrepScanner
|
from processors.semgrep_scanner.semgrep_scanner import SemgrepScanner
|
||||||
|
|
||||||
|
|
||||||
@@ -6,3 +11,159 @@ def test_semgrep_scanner_init():
|
|||||||
spec = StageSpec(name="test_scanner", processor="semgrep", config={"rules": []})
|
spec = StageSpec(name="test_scanner", processor="semgrep", config={"rules": []})
|
||||||
scanner = SemgrepScanner(spec)
|
scanner = SemgrepScanner(spec)
|
||||||
assert scanner.description != ""
|
assert scanner.description != ""
|
||||||
|
|
||||||
|
|
||||||
|
def _ctx(pipeline_dir):
|
||||||
|
return ProcessorContext(pipeline_dir=pipeline_dir, global_config={}, llm=None)
|
||||||
|
|
||||||
|
|
||||||
|
class TestRulesPathResolution:
|
||||||
|
def _scanner(self, rules_dir):
|
||||||
|
return SemgrepScanner(
|
||||||
|
StageSpec(name="scan", processor="semgrep", config={"rules_dir": rules_dir})
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_validate_and_process_resolve_the_same_path(self, tmp_path, monkeypatch):
|
||||||
|
# regression: validate() approved a cwd-relative path while process()
|
||||||
|
# used a pipeline_dir-relative one, so the scan ran with no rules.
|
||||||
|
# pin cwd so the cwd-relative "rules" resolves deterministically.
|
||||||
|
monkeypatch.chdir(tmp_path)
|
||||||
|
rules = tmp_path / "rules"
|
||||||
|
rules.mkdir()
|
||||||
|
(rules / "r.yaml").write_text("rules: []")
|
||||||
|
scanner = self._scanner("rules")
|
||||||
|
ctx = _ctx(tmp_path)
|
||||||
|
|
||||||
|
validation = scanner.validate(ctx)
|
||||||
|
# [] when semgrep is installed, else only the "semgrep CLI not found" note
|
||||||
|
assert validation == [] or "semgrep CLI" in validation[0]
|
||||||
|
assert scanner._resolve_rules_path(ctx) == rules.resolve()
|
||||||
|
assert scanner._resolve_rules_path(ctx).exists()
|
||||||
|
|
||||||
|
def test_missing_rules_dir_flagged(self, tmp_path):
|
||||||
|
scanner = self._scanner("does_not_exist_anywhere")
|
||||||
|
ctx = _ctx(tmp_path)
|
||||||
|
errs = [e for e in scanner.validate(ctx) if "rules_dir" in e]
|
||||||
|
assert errs
|
||||||
|
|
||||||
|
def test_process_fails_fast_when_rules_unresolvable(self, tmp_path):
|
||||||
|
# must not fall back to scanning with an unintended config path
|
||||||
|
scanner = self._scanner("nope_missing_rules")
|
||||||
|
ctx = _ctx(tmp_path)
|
||||||
|
results = scanner.process(ctx, [_entry(tmp_path)])
|
||||||
|
assert len(results) == 1
|
||||||
|
assert results[0].status == "failed"
|
||||||
|
assert "rules_dir not found" in results[0].error
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeProc:
|
||||||
|
def __init__(self, returncode, stdout=b"", stderr=b""):
|
||||||
|
self.returncode = returncode
|
||||||
|
self._stdout = stdout
|
||||||
|
self._stderr = stderr
|
||||||
|
|
||||||
|
async def communicate(self):
|
||||||
|
return self._stdout, self._stderr
|
||||||
|
|
||||||
|
|
||||||
|
def _scanner(tmp_path):
|
||||||
|
return SemgrepScanner(StageSpec(name="scan", processor="semgrep", config={}))
|
||||||
|
|
||||||
|
|
||||||
|
def _entry(tmp_path, sample_id="s1"):
|
||||||
|
d = tmp_path / "samples" / sample_id
|
||||||
|
d.mkdir(parents=True, exist_ok=True)
|
||||||
|
return ProcessorEntry(
|
||||||
|
sample_id=sample_id,
|
||||||
|
source_path=tmp_path / f"{sample_id}.sys",
|
||||||
|
filename=f"{sample_id}.sys",
|
||||||
|
sample_dir=d,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _distribute(tmp_path, proc):
|
||||||
|
scanner = _scanner(tmp_path)
|
||||||
|
entry = _entry(tmp_path)
|
||||||
|
uncached = [(0, entry)]
|
||||||
|
file_to_sample = {"s1_dispatch.c": 0}
|
||||||
|
results = [None]
|
||||||
|
with patch("asyncio.create_subprocess_exec", return_value=proc):
|
||||||
|
return asyncio.run(
|
||||||
|
scanner._run_and_distribute(
|
||||||
|
Path("rules"), tmp_path / "bulk", 300, uncached, file_to_sample, results, 0
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_valid_json_with_findings_exit_1(tmp_path):
|
||||||
|
out = json.dumps(
|
||||||
|
{
|
||||||
|
"results": [
|
||||||
|
{
|
||||||
|
"check_id": "r1",
|
||||||
|
"path": "x/s1_dispatch.c",
|
||||||
|
"start": {"line": 5},
|
||||||
|
"end": {"line": 5},
|
||||||
|
"extra": {"severity": "ERROR", "message": "bad"},
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
).encode()
|
||||||
|
res = _distribute(tmp_path, _FakeProc(1, stdout=out))
|
||||||
|
assert res[0].status == "completed"
|
||||||
|
assert res[0].data["finding_count"] == 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_valid_json_survives_unexpected_exit_code(tmp_path):
|
||||||
|
# semgrep can emit a full results doc yet exit non-0/1 (observed on windows)
|
||||||
|
out = json.dumps({"results": [], "version": "1.0"}).encode()
|
||||||
|
res = _distribute(tmp_path, _FakeProc(2, stdout=out))
|
||||||
|
assert res[0].status == "completed"
|
||||||
|
assert res[0].data["finding_count"] == 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_no_output_fails_with_exit_code(tmp_path):
|
||||||
|
res = _distribute(tmp_path, _FakeProc(2, stdout=b"", stderr=b"boom"))
|
||||||
|
assert res[0].status == "failed"
|
||||||
|
# the opaque "semgrep error:" is gone - exit code and stderr are surfaced
|
||||||
|
assert "exit 2" in res[0].error
|
||||||
|
assert "boom" in res[0].error
|
||||||
|
|
||||||
|
|
||||||
|
def test_unparseable_output_fails_with_exit_code(tmp_path):
|
||||||
|
res = _distribute(tmp_path, _FakeProc(0, stdout=b"not json"))
|
||||||
|
assert res[0].status == "failed"
|
||||||
|
assert "exit 0" in res[0].error
|
||||||
|
|
||||||
|
|
||||||
|
def test_scan_errors_with_no_results_fail_loudly(tmp_path):
|
||||||
|
# a bad --config makes semgrep exit 0 with errors but no results; that must
|
||||||
|
# not look like "no vulnerabilities found"
|
||||||
|
out = json.dumps(
|
||||||
|
{"results": [], "errors": [{"message": "config error: rules not found"}]}
|
||||||
|
).encode()
|
||||||
|
res = _distribute(tmp_path, _FakeProc(0, stdout=out))
|
||||||
|
assert res[0].status == "failed"
|
||||||
|
assert "no results" in res[0].error
|
||||||
|
assert "config error" in res[0].error
|
||||||
|
|
||||||
|
|
||||||
|
def test_findings_present_tolerate_nonfatal_errors(tmp_path):
|
||||||
|
# per-file parse warnings alongside real results should not fail the scan
|
||||||
|
out = json.dumps(
|
||||||
|
{
|
||||||
|
"results": [
|
||||||
|
{
|
||||||
|
"check_id": "r1",
|
||||||
|
"path": "x/s1_dispatch.c",
|
||||||
|
"start": {"line": 1},
|
||||||
|
"end": {"line": 1},
|
||||||
|
"extra": {"severity": "ERROR", "message": "m"},
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"errors": [{"message": "could not parse one file"}],
|
||||||
|
}
|
||||||
|
).encode()
|
||||||
|
res = _distribute(tmp_path, _FakeProc(0, stdout=out))
|
||||||
|
assert res[0].status == "completed"
|
||||||
|
assert res[0].data["finding_count"] == 1
|
||||||
|
|||||||
Reference in New Issue
Block a user