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:
416rehman
2026-07-24 23:03:09 -06:00
committed by GitHub
17 changed files with 1449 additions and 94 deletions
+15 -2
View File
@@ -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-...
+1 -1
View File
@@ -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
+57 -17
View File
@@ -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]
+5
View File
@@ -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"]
+16
View File
@@ -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()
+62
View File
@@ -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",
]
+292
View File
@@ -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
+129
View File
@@ -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
+77
View File
@@ -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
View File
@@ -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
+4 -1
View File
@@ -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)
+4 -12
View File
@@ -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
+310
View File
@@ -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") == []
+196
View File
@@ -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)
+162 -1
View File
@@ -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