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.
# =========================================================
# 1. LiteLLM Abstraction Keys
# Provide keys for whatever endpoints your pipelines target
# 1. LLM Backend
#
# 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...
# OPENAI_API_KEY=sk-...
+1 -1
View File
@@ -13,7 +13,7 @@ name: loldrivers
description: windows kernel driver vulnerability research pipeline
version: "2.0"
model: vertex_ai/gemini-2.5-pro
model: claude-code/sonnet
settings:
work_dir: work
@@ -360,10 +360,12 @@ def main():
# write per-ioctl file
# 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:
f.write("// IOCTL Code: 0x%08X\n" % code)
f.write("// Method: %d\n" % (code & 0x3))
f.write("// Device Type: 0x%04X\n" % ((code >> 16) & 0xFFFF))
f.write("// Function: 0x%03X\n\n" % ((code >> 2) & 0xFFF))
# unicode() cast required: a jython io text stream rejects str,
# and "..." % int yields str (raising "can't write str to text stream")
f.write(unicode("// IOCTL Code: 0x%08X\n" % code))
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))
result["success"] = True
+57 -17
View File
@@ -19,28 +19,42 @@ from deepzero.engine.stage import (
class SemgrepScanner(BulkMapProcessor):
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]:
errors = []
if not shutil.which("semgrep"):
errors.append("semgrep CLI not found in PATH - install with: pip install semgrep")
rules_dir = self.config.get("rules_dir")
if not rules_dir:
if not self.config.get("rules_dir"):
errors.append("semgrep_scanner requires 'rules_dir' in config")
else:
rules_path = (Path.cwd() / rules_dir).resolve()
if not rules_path.exists():
rules_path = (ctx.pipeline_dir / rules_dir).resolve()
if not rules_path.exists():
errors.append(f"rules_dir does not exist: {rules_dir}")
rules_path = self._resolve_rules_path(ctx)
if rules_path is None or not rules_path.exists():
errors.append(f"rules_dir does not exist: {self.config.get('rules_dir')}")
return errors
def process(
self, ctx: ProcessorContext, entries: list[ProcessorEntry]
) -> list[ProcessorResult]:
rules_dir = self.config.get("rules_dir", "")
rules_path = (ctx.pipeline_dir / rules_dir).resolve()
rules_path = self._resolve_rules_path(ctx)
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")
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")
return [r for r in results if r is not None]
if proc.returncode not in (0, 1):
err = stderr_bytes.decode("utf-8", errors="replace")[:500]
out_str = stdout_bytes.decode("utf-8", errors="replace")
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:
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]
try:
out_str = stdout_bytes.decode("utf-8", errors="replace")
output = json.loads(out_str) if out_str.strip() else {}
except json.JSONDecodeError:
# semgrep reports rule/config load failures in an "errors" array while
# still exiting 0 with empty results - which would otherwise look like
# "no vulnerabilities found". fail loudly when nothing scanned.
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:
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]
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)
return [r for r in results if r is not None]
+5
View File
@@ -36,6 +36,11 @@ pythonpath = ["src", "."]
[tool.ruff]
line-length = 100
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]
select = ["E", "F", "W", "I"]
+16
View File
@@ -1,6 +1,7 @@
from __future__ import annotations
import logging
import sys
import time
from pathlib import Path
from types import MappingProxyType
@@ -14,6 +15,21 @@ from rich.table import Table
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()
+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
from typing import Any
from deepzero.engine.backends import create_backend
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:
# 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):
self.model = model
self.default_kwargs = kwargs
self._ensure_litellm()
self.backend = create_backend(model, **kwargs)
def _ensure_litellm(self) -> None:
try:
import litellm
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
@property
def _litellm(self) -> Any:
# retained for backward compatibility with existing callers/tests
return getattr(self.backend, "litellm", None)
def complete(
self,
@@ -66,23 +37,22 @@ class LLMProvider:
) -> str:
"""send messages to the llm and return the response text.
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}
backoff = initial_backoff
for attempt in range(max_retries + 1):
try:
response = self._litellm.completion(
model=self.model,
messages=messages,
**merged,
)
content = response.choices[0].message.content or ""
content = backend.raw_complete(messages, **merged)
# decay backoff toward minimum on success
backoff = max(initial_backoff, backoff * backoff_decay)
return content
except self._rate_limit_error:
except backend.rate_limit_error:
if attempt == max_retries:
raise
backoff = min(max_backoff, backoff * 2.0)
@@ -94,11 +64,15 @@ class LLMProvider:
)
time.sleep(backoff)
except self._context_window_error:
except backend.context_window_error:
# context window errors won't be fixed by retry
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:
raise
wait = min(max_backoff, 2**attempt)
@@ -115,12 +89,8 @@ class LLMProvider:
@property
def provider_name(self) -> str:
if "/" in self.model:
return self.model.split("/")[0]
return "unknown"
return self.backend.provider_name
@property
def model_name(self) -> str:
if "/" in self.model:
return self.model.split("/", 1)[1]
return self.model
return self.backend.model_name
+4 -1
View File
@@ -271,7 +271,10 @@ def _expand_string(s: str) -> str:
if ":-" in var:
name, default = var.split(":-", 1)
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)
+4 -12
View File
@@ -34,21 +34,13 @@ class GenericLLM(MapProcessor):
if not prompt_path.exists():
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")
if model:
try:
import litellm
from deepzero.engine.backends import validate_model_binding
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:
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")
errors.extend(validate_model_binding(model))
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
@@ -6,3 +11,159 @@ def test_semgrep_scanner_init():
spec = StageSpec(name="test_scanner", processor="semgrep", config={"rules": []})
scanner = SemgrepScanner(spec)
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