mirror of
https://github.com/416rehman/DeepZero
synced 2026-08-09 11:54:39 +00:00
feat: run pipelines with your Claude Code subscription, no API key required
Pipelines can now use the Claude Code you already have installed and
signed in, so LLM analysis works without a third-party API key or
per-token billing:
model: claude-code # default model
model: claude-code/sonnet # or /opus, or a full model name
DeepZero never handles your credentials - it runs your own `claude`
binary in headless mode and inherits its existing sign-in. If the CLI is
missing or not signed in, validation says so before a run starts.
The pipeline feeds untrusted decompiled code to the model, so the
integration denies tools and external servers, sends the prompt over
stdin rather than the command line, and hides any ANTHROPIC_API_KEY from
the CLI so your subscription is used and not a metered API. Usage limits
retry with backoff; sign-in failures stop immediately with a clear
message.
Under the hood, LLM backends are selected from the model string by a
registry, so another agent CLI can be added later without touching the
engine, pipelines, or stages. Existing API-key models are unaffected.
Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 5
parent
c6ef8a9a7c
commit
d83d1711e0
+15
-2
@@ -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-...
|
||||
|
||||
@@ -0,0 +1,62 @@
|
||||
"""pluggable llm backends.
|
||||
|
||||
to add support for another agent cli (codex, gemini cli, ...):
|
||||
|
||||
# deepzero/engine/backends/mytool.py
|
||||
class MyToolBackend(CLIAgentBackend):
|
||||
scheme = "mytool"
|
||||
display_name = "my tool"
|
||||
binary_names = ("mytool",)
|
||||
|
||||
def build_argv(self, system): ...
|
||||
def parse_output(self, returncode, stdout, stderr): ...
|
||||
|
||||
then register it below (or call register_backend() from your own code - the
|
||||
registry is open, so third parties can add backends without patching deepzero).
|
||||
nothing in LLMProvider, the pipeline, or the stages needs to change.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from deepzero.engine.backends.base import (
|
||||
BackendAuthError,
|
||||
BackendContextWindowError,
|
||||
BackendError,
|
||||
BackendNotFoundError,
|
||||
BackendRateLimitError,
|
||||
CLIAgentBackend,
|
||||
LLMBackend,
|
||||
)
|
||||
from deepzero.engine.backends.claude_code import ClaudeCodeBackend
|
||||
from deepzero.engine.backends.litellm_backend import LiteLLMBackend
|
||||
from deepzero.engine.backends.registry import (
|
||||
create_backend,
|
||||
get_registered_backends,
|
||||
model_scheme,
|
||||
register_backend,
|
||||
resolve_backend_class,
|
||||
validate_model_binding,
|
||||
)
|
||||
|
||||
# -- built-in backends --
|
||||
register_backend(ClaudeCodeBackend)
|
||||
# litellm is the fallback for any scheme no other backend claims
|
||||
register_backend(LiteLLMBackend, default=True)
|
||||
|
||||
__all__ = [
|
||||
"BackendAuthError",
|
||||
"BackendContextWindowError",
|
||||
"BackendError",
|
||||
"BackendNotFoundError",
|
||||
"BackendRateLimitError",
|
||||
"CLIAgentBackend",
|
||||
"ClaudeCodeBackend",
|
||||
"LLMBackend",
|
||||
"LiteLLMBackend",
|
||||
"create_backend",
|
||||
"get_registered_backends",
|
||||
"model_scheme",
|
||||
"register_backend",
|
||||
"resolve_backend_class",
|
||||
"validate_model_binding",
|
||||
]
|
||||
@@ -0,0 +1,298 @@
|
||||
"""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)
|
||||
|
||||
# whether per-call generation kwargs (temperature, max_tokens, ...) are
|
||||
# forwarded to raw_complete. cli backends take their options via argv.
|
||||
accepts_generation_kwargs: ClassVar[bool] = True
|
||||
|
||||
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
|
||||
|
||||
accepts_generation_kwargs: ClassVar[bool] = False
|
||||
|
||||
def __init__(self, model: str, **kwargs: Any):
|
||||
super().__init__(model, **kwargs)
|
||||
self.alias = self.model_alias(model)
|
||||
self.timeout = int(kwargs.get("timeout", self.default_timeout))
|
||||
self.cwd = kwargs.get("cwd")
|
||||
# choosing an agent-cli backend signals intent to use its own login
|
||||
self.prefer_subscription_auth = bool(kwargs.get("prefer_subscription_auth", True))
|
||||
|
||||
self._binary = kwargs.get("binary") or self.find_binary()
|
||||
if not self._binary:
|
||||
raise BackendNotFoundError(self.not_found_message())
|
||||
|
||||
# -- discovery --------------------------------------------------------------
|
||||
|
||||
@classmethod
|
||||
def find_binary(cls) -> str | None:
|
||||
for name in cls.binary_names:
|
||||
found = shutil.which(name)
|
||||
if found:
|
||||
return found
|
||||
for candidate in cls.extra_search_paths():
|
||||
try:
|
||||
if candidate.is_file():
|
||||
return str(candidate)
|
||||
except OSError:
|
||||
continue
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def extra_search_paths(cls) -> list[Path]:
|
||||
"""non-PATH locations to probe. override for vendor-specific installs."""
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
def not_found_message(cls) -> str:
|
||||
label = cls.display_name or cls.scheme
|
||||
hint = f" {cls.install_hint}" if cls.install_hint else ""
|
||||
return (
|
||||
f"{label} cli not found (looked for: {', '.join(cls.binary_names) or 'n/a'})."
|
||||
f"{hint} alternatively use an api-key model instead."
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def validate_model(cls, model: str) -> list[str]:
|
||||
if not cls.find_binary():
|
||||
return [f"LLM backend '{model}' is unavailable: {cls.not_found_message()}"]
|
||||
return []
|
||||
|
||||
# -- prompt assembly --------------------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
def split_messages(messages: list[dict[str, str]]) -> tuple[str, str]:
|
||||
"""flatten a message list into (system_prompt, user_prompt).
|
||||
|
||||
agent clis are single-shot, so prior turns are rendered inline. a lone
|
||||
user message - the pipeline's normal case - passes through untouched.
|
||||
"""
|
||||
system_parts: list[str] = []
|
||||
convo: list[dict[str, str]] = []
|
||||
for msg in messages or []:
|
||||
role = (msg.get("role") or "").lower()
|
||||
content = msg.get("content") or ""
|
||||
if role == "system":
|
||||
if content:
|
||||
system_parts.append(content)
|
||||
else:
|
||||
convo.append({"role": role or "user", "content": content})
|
||||
|
||||
system = "\n\n".join(system_parts)
|
||||
|
||||
if len(convo) == 1:
|
||||
return system, convo[0]["content"]
|
||||
|
||||
rendered = [
|
||||
f"{'Assistant' if m['role'] == 'assistant' else 'User'}: {m['content']}" for m in convo
|
||||
]
|
||||
return system, "\n\n".join(rendered)
|
||||
|
||||
# -- subclass hooks ---------------------------------------------------------
|
||||
|
||||
@abstractmethod
|
||||
def build_argv(self, system: str) -> tuple[list[str], str]:
|
||||
"""returns (argv, system_text_to_inline).
|
||||
|
||||
a system prompt too large for argv should be returned as the second
|
||||
element so the base class routes it through stdin instead.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def parse_output(self, returncode: int, stdout: str, stderr: str) -> str:
|
||||
"""extract response text, or raise a BackendError subclass."""
|
||||
|
||||
def build_env(self) -> dict[str, str]:
|
||||
env = dict(os.environ)
|
||||
if self.prefer_subscription_auth:
|
||||
for var in self.subscription_auth_env_blocklist:
|
||||
env.pop(var, None)
|
||||
return env
|
||||
|
||||
def classify_error(self, detail: str, status: Any = None) -> BackendError:
|
||||
"""map a failure into the generic taxonomy. override to add vendor cases."""
|
||||
low = (detail or "").lower()
|
||||
try:
|
||||
code = int(status) if status is not None else None
|
||||
except (TypeError, ValueError):
|
||||
code = None
|
||||
|
||||
label = self.display_name or self.scheme
|
||||
if code in (401, 403) or any(k in low for k in ("authenticat", "oauth", "unauthorized")):
|
||||
return BackendAuthError(f"{label} is not authenticated ({detail})")
|
||||
if code == 429 or any(
|
||||
k in low for k in ("rate limit", "usage limit", "quota", "overloaded", "too many")
|
||||
):
|
||||
return BackendRateLimitError(detail)
|
||||
if "context" in low and any(k in low for k in ("too long", "exceed", "window")):
|
||||
return BackendContextWindowError(detail)
|
||||
if code is not None and 500 <= code < 600:
|
||||
return BackendError(f"{label} server error {code}: {detail}")
|
||||
return BackendError(detail)
|
||||
|
||||
# -- execution --------------------------------------------------------------
|
||||
|
||||
def raw_complete(self, messages: list[dict[str, str]], **kwargs: Any) -> str:
|
||||
system, prompt = self.split_messages(messages)
|
||||
argv, inline_system = self.build_argv(system)
|
||||
|
||||
stdin_text = f"{inline_system}\n\n{prompt}" if inline_system else prompt
|
||||
timeout = int(kwargs.get("timeout", self.timeout))
|
||||
|
||||
try:
|
||||
proc = subprocess.run(
|
||||
argv,
|
||||
input=stdin_text,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
encoding="utf-8",
|
||||
errors="replace",
|
||||
timeout=timeout,
|
||||
env=self.build_env(),
|
||||
cwd=self.cwd,
|
||||
)
|
||||
except subprocess.TimeoutExpired as exc:
|
||||
label = self.display_name or self.scheme
|
||||
raise BackendError(f"{label} timed out after {timeout}s") from exc
|
||||
except OSError as exc:
|
||||
label = self.display_name or self.scheme
|
||||
raise BackendError(f"failed to launch {label}: {exc}") from exc
|
||||
|
||||
return self.parse_output(proc.returncode, proc.stdout or "", proc.stderr or "")
|
||||
|
||||
def check_binary(self) -> str | None:
|
||||
"""returns None if the binary runs, else a description of the problem."""
|
||||
try:
|
||||
proc = subprocess.run(
|
||||
[str(self._binary), "--version"],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
encoding="utf-8",
|
||||
errors="replace",
|
||||
timeout=30,
|
||||
env=self.build_env(),
|
||||
)
|
||||
except (OSError, subprocess.SubprocessError) as exc:
|
||||
return f"could not run {self.display_name or self.scheme}: {exc}"
|
||||
if proc.returncode != 0:
|
||||
return f"exited {proc.returncode}: {(proc.stderr or '').strip()[:200]}"
|
||||
return None
|
||||
@@ -0,0 +1,129 @@
|
||||
"""claude code backend - drives the user's locally installed, already-signed-in
|
||||
`claude` cli in headless mode (`claude -p --output-format json`).
|
||||
|
||||
lets deepzero run against a claude code subscription instead of a metered api
|
||||
key. all subprocess handling lives in CLIAgentBackend; this module only declares
|
||||
claude-specific argv and output parsing.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Any, ClassVar
|
||||
|
||||
from deepzero.engine.backends.base import BackendError, CLIAgentBackend
|
||||
|
||||
log = logging.getLogger("deepzero.llm.claude_code")
|
||||
|
||||
# the prompt carries untrusted content (decompiled malware, attacker-controlled
|
||||
# strings). a completion never needs tools, so deny the ones that could turn a
|
||||
# prompt injection into code execution or exfiltration.
|
||||
_DEFAULT_DISALLOWED_TOOLS = (
|
||||
"Bash",
|
||||
"Edit",
|
||||
"Write",
|
||||
"NotebookEdit",
|
||||
"Read",
|
||||
"Glob",
|
||||
"Grep",
|
||||
"WebFetch",
|
||||
"WebSearch",
|
||||
"Agent",
|
||||
"Task",
|
||||
)
|
||||
|
||||
|
||||
class ClaudeCodeBackend(CLIAgentBackend):
|
||||
scheme: ClassVar[str] = "claude-code"
|
||||
display_name: ClassVar[str] = "claude code"
|
||||
binary_names: ClassVar[tuple[str, ...]] = ("claude",)
|
||||
install_hint: ClassVar[str] = "install it and sign in (https://claude.com/claude-code)."
|
||||
# never use --bare: it disables oauth/keychain reads, which is exactly the
|
||||
# auth this backend exists to use
|
||||
subscription_auth_env_blocklist: ClassVar[tuple[str, ...]] = (
|
||||
"ANTHROPIC_API_KEY",
|
||||
"ANTHROPIC_AUTH_TOKEN",
|
||||
)
|
||||
|
||||
def __init__(self, model: str, **kwargs: Any):
|
||||
super().__init__(model, **kwargs)
|
||||
self.max_turns = int(kwargs.get("max_turns", 1))
|
||||
disallowed = kwargs.get("disallowed_tools", _DEFAULT_DISALLOWED_TOOLS)
|
||||
self.disallowed_tools: tuple[str, ...] = tuple(disallowed) if disallowed else ()
|
||||
|
||||
@classmethod
|
||||
def extra_search_paths(cls) -> list[Path]:
|
||||
home = Path.home()
|
||||
paths = [
|
||||
home / ".claude" / "local" / "claude",
|
||||
home / ".local" / "bin" / "claude",
|
||||
home / "bin" / "claude",
|
||||
]
|
||||
if sys.platform == "win32":
|
||||
appdata = os.environ.get("APPDATA", str(home / "AppData" / "Roaming"))
|
||||
localappdata = os.environ.get("LOCALAPPDATA", str(home / "AppData" / "Local"))
|
||||
for base in (Path(appdata) / "npm", Path(localappdata) / "claude" / "bin"):
|
||||
paths += [base / "claude.cmd", base / "claude.exe", base / "claude"]
|
||||
paths.append(home / ".local" / "bin" / "claude.exe")
|
||||
return paths
|
||||
|
||||
def build_argv(self, system: str) -> tuple[list[str], str]:
|
||||
argv = [str(self._binary), "-p", "--output-format", "json"]
|
||||
|
||||
if self.alias:
|
||||
argv += ["--model", self.alias]
|
||||
if self.max_turns > 0:
|
||||
argv += ["--max-turns", str(self.max_turns)]
|
||||
if self.disallowed_tools:
|
||||
argv += ["--disallowed-tools", ",".join(self.disallowed_tools)]
|
||||
# ignore any globally configured mcp servers - a completion needs none,
|
||||
# and they would widen the blast radius of injected content
|
||||
argv.append("--strict-mcp-config")
|
||||
|
||||
inline_system = ""
|
||||
if system:
|
||||
if len(system) <= self.max_argv_text:
|
||||
argv += ["--append-system-prompt", system]
|
||||
else:
|
||||
inline_system = system
|
||||
return argv, inline_system
|
||||
|
||||
def parse_output(self, returncode: int, stdout: str, stderr: str) -> str:
|
||||
payload: dict[str, Any] | None = None
|
||||
text = stdout.strip()
|
||||
if text:
|
||||
try:
|
||||
parsed = json.loads(text)
|
||||
if isinstance(parsed, dict):
|
||||
payload = parsed
|
||||
except json.JSONDecodeError:
|
||||
payload = None
|
||||
|
||||
if payload is None:
|
||||
detail = (stderr or stdout or "").strip()[-500:]
|
||||
if returncode != 0:
|
||||
raise self.classify_error(detail)
|
||||
raise BackendError(f"could not parse claude code output: {detail}")
|
||||
|
||||
result = payload.get("result")
|
||||
# the cli reports subtype "success" even on failures - is_error is authoritative
|
||||
if payload.get("is_error") or returncode != 0:
|
||||
detail = result if isinstance(result, str) and result else (stderr or "").strip()
|
||||
raise self.classify_error(detail or "unknown error", payload.get("api_error_status"))
|
||||
|
||||
if not isinstance(result, str):
|
||||
raise BackendError("claude code returned no result text")
|
||||
|
||||
usage = payload.get("usage") or {}
|
||||
log.debug(
|
||||
"claude code ok (in=%s out=%s cost=%s session=%s)",
|
||||
usage.get("input_tokens"),
|
||||
usage.get("output_tokens"),
|
||||
payload.get("total_cost_usd"),
|
||||
payload.get("session_id"),
|
||||
)
|
||||
return result
|
||||
@@ -0,0 +1,88 @@
|
||||
"""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"
|
||||
accepts_generation_kwargs: ClassVar[bool] = True
|
||||
# litellm surfaces auth problems as APIError subclasses that are also in the
|
||||
# retryable set, so there is no distinct fail-fast class to declare
|
||||
non_retryable_errors: ClassVar[tuple[type[BaseException], ...]] = (_NeverRaised,)
|
||||
|
||||
def __init__(self, model: str, **kwargs: Any):
|
||||
super().__init__(model, **kwargs)
|
||||
try:
|
||||
import litellm
|
||||
except ImportError as exc:
|
||||
raise ImportError(
|
||||
"litellm is required for LLM support. install with: pip install litellm"
|
||||
) from exc
|
||||
|
||||
self.litellm = litellm
|
||||
# suppress litellm's noisy logging and traceback spam
|
||||
litellm.suppress_debug_info = True
|
||||
logging.getLogger("litellm").setLevel(logging.CRITICAL)
|
||||
|
||||
# instance-level overrides of the class retry roles, with safe fallbacks
|
||||
# for test mocks that lack the real exception classes
|
||||
self.rate_limit_error = _resolve_exc(litellm, "RateLimitError")
|
||||
self.context_window_error = _resolve_exc(litellm, "ContextWindowExceededError")
|
||||
api_errors = tuple(
|
||||
_resolve_exc(litellm, name) for name in ("APIConnectionError", "APIError")
|
||||
)
|
||||
self.retryable_errors = api_errors + (OSError, ValueError, RuntimeError)
|
||||
|
||||
def raw_complete(self, messages: list[dict[str, str]], **kwargs: Any) -> str:
|
||||
response = self.litellm.completion(model=self.model, messages=messages, **kwargs)
|
||||
return response.choices[0].message.content or ""
|
||||
|
||||
@classmethod
|
||||
def validate_model(cls, model: str) -> list[str]:
|
||||
try:
|
||||
import litellm
|
||||
except ImportError:
|
||||
return ["LLM configured, but 'litellm' framework is not installed"]
|
||||
|
||||
env_state = litellm.validate_environment(model=model)
|
||||
if not env_state.get("keys_in_environment", True):
|
||||
missing_keys = env_state.get("missing_keys", [])
|
||||
if missing_keys:
|
||||
return [
|
||||
f"LLM backend '{model}' missing credentials in environment. "
|
||||
f"Need: {missing_keys}"
|
||||
]
|
||||
return []
|
||||
|
||||
@property
|
||||
def provider_name(self) -> str:
|
||||
if "/" in self.model:
|
||||
return self.model.split("/")[0]
|
||||
return "unknown"
|
||||
|
||||
@property
|
||||
def model_name(self) -> str:
|
||||
if "/" in self.model:
|
||||
return self.model.split("/", 1)[1]
|
||||
return self.model
|
||||
@@ -0,0 +1,77 @@
|
||||
"""llm backend registry.
|
||||
|
||||
model strings are `<scheme>` or `<scheme>/<model>`. a registered scheme selects
|
||||
that backend; anything unrecognised falls through to the default backend
|
||||
(litellm), which covers every api-key endpoint it supports.
|
||||
|
||||
claude-code -> ClaudeCodeBackend (default model)
|
||||
claude-code/sonnet -> ClaudeCodeBackend (sonnet)
|
||||
openai/gpt-4o -> default backend (litellm)
|
||||
gpt-4o -> default backend (litellm)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from deepzero.engine.backends.base import LLMBackend
|
||||
|
||||
_BACKEND_REGISTRY: dict[str, type[LLMBackend]] = {}
|
||||
_DEFAULT_BACKEND: type[LLMBackend] | None = None
|
||||
|
||||
|
||||
def register_backend(cls: type[LLMBackend], *, default: bool = False) -> type[LLMBackend]:
|
||||
"""register a backend under its `scheme` (and any `aliases`).
|
||||
|
||||
usable as a decorator. pass default=True for the fallback backend that
|
||||
handles every unregistered scheme.
|
||||
"""
|
||||
global _DEFAULT_BACKEND
|
||||
|
||||
if default:
|
||||
_DEFAULT_BACKEND = cls
|
||||
|
||||
if cls.scheme:
|
||||
_BACKEND_REGISTRY[cls.scheme.lower()] = cls
|
||||
for alias in cls.aliases:
|
||||
_BACKEND_REGISTRY[alias.lower()] = cls
|
||||
elif not default:
|
||||
raise ValueError(f"{cls.__name__} must define a 'scheme' to be registered")
|
||||
|
||||
return cls
|
||||
|
||||
|
||||
def get_registered_backends() -> dict[str, type[LLMBackend]]:
|
||||
return dict(_BACKEND_REGISTRY)
|
||||
|
||||
|
||||
def model_scheme(model: str) -> str:
|
||||
m = (model or "").strip().lower()
|
||||
return m.split("/", 1)[0] if "/" in m else m
|
||||
|
||||
|
||||
def resolve_backend_class(model: str) -> type[LLMBackend]:
|
||||
"""map a model string to its backend class."""
|
||||
backend = _BACKEND_REGISTRY.get(model_scheme(model))
|
||||
if backend is not None:
|
||||
return backend
|
||||
|
||||
if _DEFAULT_BACKEND is None:
|
||||
raise ValueError(
|
||||
f"no backend can handle model '{model}' and no default backend is registered. "
|
||||
f"known schemes: {sorted(_BACKEND_REGISTRY)}"
|
||||
)
|
||||
return _DEFAULT_BACKEND
|
||||
|
||||
|
||||
def create_backend(model: str, **kwargs: Any) -> LLMBackend:
|
||||
return resolve_backend_class(model)(model, **kwargs)
|
||||
|
||||
|
||||
def validate_model_binding(model: str) -> list[str]:
|
||||
"""pre-flight check for a model string, delegated to its backend."""
|
||||
try:
|
||||
backend_cls = resolve_backend_class(model)
|
||||
except ValueError as exc:
|
||||
return [str(exc)]
|
||||
return backend_cls.validate_model(model)
|
||||
+25
-57
@@ -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,20 @@ class LLMProvider:
|
||||
) -> str:
|
||||
"""send messages to the llm and return the response text.
|
||||
handles rate limiting with adaptive backoff."""
|
||||
merged = {**self.default_kwargs, **kwargs}
|
||||
backend = self.backend
|
||||
# cli backends take their options via argv, not per-call kwargs
|
||||
merged = {**self.default_kwargs, **kwargs} if backend.accepts_generation_kwargs else {}
|
||||
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 +62,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 +87,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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -0,0 +1,310 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import subprocess
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
from deepzero.engine.backends import ClaudeCodeBackend
|
||||
from deepzero.engine.backends.base import (
|
||||
BackendAuthError,
|
||||
BackendContextWindowError,
|
||||
BackendError,
|
||||
BackendNotFoundError,
|
||||
BackendRateLimitError,
|
||||
)
|
||||
from deepzero.engine.llm import LLMProvider
|
||||
|
||||
_FIND = "deepzero.engine.backends.claude_code.ClaudeCodeBackend.find_binary"
|
||||
|
||||
|
||||
def _backend(**kwargs) -> ClaudeCodeBackend:
|
||||
kwargs.setdefault("binary", "/usr/bin/claude")
|
||||
return ClaudeCodeBackend(kwargs.pop("model", "claude-code/sonnet"), **kwargs)
|
||||
|
||||
|
||||
def _completed(stdout: str, returncode: int = 0, stderr: str = ""):
|
||||
return subprocess.CompletedProcess(args=[], returncode=returncode, stdout=stdout, stderr=stderr)
|
||||
|
||||
|
||||
def _result_json(**overrides) -> str:
|
||||
payload = {
|
||||
"type": "result",
|
||||
"subtype": "success",
|
||||
"is_error": False,
|
||||
"result": "hello from claude",
|
||||
"session_id": "abc",
|
||||
"total_cost_usd": 0.01,
|
||||
"usage": {"input_tokens": 10, "output_tokens": 5},
|
||||
}
|
||||
payload.update(overrides)
|
||||
return json.dumps(payload)
|
||||
|
||||
|
||||
class TestAliasParsing:
|
||||
def test_alias_parsing(self):
|
||||
assert _backend(model="claude-code/sonnet").alias == "sonnet"
|
||||
assert _backend(model="claude-code").alias == ""
|
||||
assert _backend(model="claude-code/claude-opus-5").alias == "claude-opus-5"
|
||||
|
||||
|
||||
class TestBinaryResolution:
|
||||
def test_raises_when_not_installed(self):
|
||||
with patch(_FIND, return_value=None):
|
||||
with pytest.raises(BackendNotFoundError) as exc:
|
||||
ClaudeCodeBackend("claude-code")
|
||||
assert "not found" in str(exc.value).lower()
|
||||
|
||||
def test_uses_explicit_binary(self):
|
||||
assert ClaudeCodeBackend("claude-code", binary="/custom/claude")._binary == "/custom/claude"
|
||||
|
||||
|
||||
class TestMessageFlattening:
|
||||
def test_single_user_message_passes_through(self):
|
||||
system, prompt = ClaudeCodeBackend.split_messages([{"role": "user", "content": "analyze"}])
|
||||
assert system == ""
|
||||
assert prompt == "analyze"
|
||||
|
||||
def test_system_is_separated(self):
|
||||
system, prompt = ClaudeCodeBackend.split_messages(
|
||||
[{"role": "system", "content": "be terse"}, {"role": "user", "content": "hi"}]
|
||||
)
|
||||
assert system == "be terse"
|
||||
assert prompt == "hi"
|
||||
|
||||
def test_multi_turn_is_rendered_inline(self):
|
||||
system, prompt = ClaudeCodeBackend.split_messages(
|
||||
[
|
||||
{"role": "system", "content": "sys"},
|
||||
{"role": "user", "content": "q1"},
|
||||
{"role": "assistant", "content": "a1"},
|
||||
{"role": "user", "content": "q2"},
|
||||
]
|
||||
)
|
||||
assert system == "sys"
|
||||
assert "User: q1" in prompt
|
||||
assert "Assistant: a1" in prompt
|
||||
assert "User: q2" in prompt
|
||||
|
||||
def test_empty(self):
|
||||
assert ClaudeCodeBackend.split_messages([]) == ("", "")
|
||||
|
||||
|
||||
class TestCommandConstruction:
|
||||
def test_core_flags(self):
|
||||
argv, inline = _backend().build_argv("")
|
||||
assert argv[:4] == ["/usr/bin/claude", "-p", "--output-format", "json"]
|
||||
assert "--model" in argv and "sonnet" in argv
|
||||
assert "--strict-mcp-config" in argv
|
||||
assert inline == ""
|
||||
|
||||
def test_no_model_flag_for_bare_alias(self):
|
||||
argv, _ = _backend(model="claude-code").build_argv("")
|
||||
assert "--model" not in argv
|
||||
|
||||
def test_dangerous_tools_denied_by_default(self):
|
||||
argv, _ = _backend().build_argv("")
|
||||
denied = argv[argv.index("--disallowed-tools") + 1]
|
||||
for tool in ("Bash", "Write", "Edit", "WebFetch"):
|
||||
assert tool in denied
|
||||
|
||||
def test_max_turns_limits_agent_loop(self):
|
||||
argv, _ = _backend().build_argv("")
|
||||
assert "--max-turns" in argv
|
||||
|
||||
def test_small_system_prompt_goes_to_argv(self):
|
||||
argv, inline = _backend().build_argv("be terse")
|
||||
assert "--append-system-prompt" in argv
|
||||
assert "be terse" in argv
|
||||
assert inline == ""
|
||||
|
||||
def test_huge_system_prompt_moves_to_stdin(self):
|
||||
big = "x" * 20000
|
||||
argv, inline = _backend().build_argv(big)
|
||||
assert "--append-system-prompt" not in argv
|
||||
assert inline == big
|
||||
|
||||
def test_never_uses_bare_flag(self):
|
||||
# --bare disables oauth/keychain reads, which would break subscription auth
|
||||
argv, _ = _backend().build_argv("sys")
|
||||
assert "--bare" not in argv
|
||||
|
||||
|
||||
class TestEnvHandling:
|
||||
def test_api_key_hidden_by_default(self):
|
||||
with patch.dict("os.environ", {"ANTHROPIC_API_KEY": "sk-ant-x"}, clear=False):
|
||||
env = _backend().build_env()
|
||||
assert "ANTHROPIC_API_KEY" not in env
|
||||
|
||||
def test_api_key_preserved_when_opted_out(self):
|
||||
with patch.dict("os.environ", {"ANTHROPIC_API_KEY": "sk-ant-x"}, clear=False):
|
||||
env = _backend(prefer_subscription_auth=False).build_env()
|
||||
assert env["ANTHROPIC_API_KEY"] == "sk-ant-x"
|
||||
|
||||
|
||||
class TestResponseParsing:
|
||||
def test_success(self):
|
||||
with patch("subprocess.run", return_value=_completed(_result_json())) as mock_run:
|
||||
out = _backend().raw_complete([{"role": "user", "content": "hi"}])
|
||||
assert out == "hello from claude"
|
||||
# the prompt travels via stdin, never argv - no shell/argv injection surface
|
||||
assert mock_run.call_args.kwargs["input"] == "hi"
|
||||
|
||||
def test_is_error_true_despite_success_subtype(self):
|
||||
# the cli reports subtype "success" even on failures; is_error is authoritative
|
||||
body = _result_json(is_error=True, subtype="success", result="boom")
|
||||
with patch("subprocess.run", return_value=_completed(body)):
|
||||
with pytest.raises(BackendError):
|
||||
_backend().raw_complete([{"role": "user", "content": "hi"}])
|
||||
|
||||
def test_auth_error_is_classified(self):
|
||||
body = _result_json(
|
||||
is_error=True,
|
||||
api_error_status=401,
|
||||
result="Failed to authenticate. API Error: 401 OAuth access token has been revoked.",
|
||||
)
|
||||
with patch("subprocess.run", return_value=_completed(body)):
|
||||
with pytest.raises(BackendAuthError) as exc:
|
||||
_backend().raw_complete([{"role": "user", "content": "hi"}])
|
||||
assert "not authenticated" in str(exc.value)
|
||||
|
||||
def test_rate_limit_is_classified(self):
|
||||
body = _result_json(is_error=True, api_error_status=429, result="rate limit exceeded")
|
||||
with patch("subprocess.run", return_value=_completed(body)):
|
||||
with pytest.raises(BackendRateLimitError):
|
||||
_backend().raw_complete([{"role": "user", "content": "hi"}])
|
||||
|
||||
def test_usage_limit_text_is_rate_limit(self):
|
||||
body = _result_json(is_error=True, result="Claude usage limit reached")
|
||||
with patch("subprocess.run", return_value=_completed(body)):
|
||||
with pytest.raises(BackendRateLimitError):
|
||||
_backend().raw_complete([{"role": "user", "content": "hi"}])
|
||||
|
||||
def test_context_window_is_classified(self):
|
||||
body = _result_json(is_error=True, result="prompt is too long: context window exceeded")
|
||||
with patch("subprocess.run", return_value=_completed(body)):
|
||||
with pytest.raises(BackendContextWindowError):
|
||||
_backend().raw_complete([{"role": "user", "content": "hi"}])
|
||||
|
||||
def test_unparseable_output(self):
|
||||
with patch("subprocess.run", return_value=_completed("not json")):
|
||||
with pytest.raises(BackendError):
|
||||
_backend().raw_complete([{"role": "user", "content": "hi"}])
|
||||
|
||||
def test_nonzero_exit_with_stderr(self):
|
||||
with patch("subprocess.run", return_value=_completed("", 1, "command failed")):
|
||||
with pytest.raises(BackendError) as exc:
|
||||
_backend().raw_complete([{"role": "user", "content": "hi"}])
|
||||
assert "command failed" in str(exc.value)
|
||||
|
||||
def test_timeout(self):
|
||||
with patch(
|
||||
"subprocess.run", side_effect=subprocess.TimeoutExpired(cmd="claude", timeout=5)
|
||||
):
|
||||
with pytest.raises(BackendError) as exc:
|
||||
_backend().raw_complete([{"role": "user", "content": "hi"}])
|
||||
assert "timed out" in str(exc.value)
|
||||
|
||||
def test_launch_failure(self):
|
||||
with patch("subprocess.run", side_effect=OSError("no exec")):
|
||||
with pytest.raises(BackendError):
|
||||
_backend().raw_complete([{"role": "user", "content": "hi"}])
|
||||
|
||||
|
||||
class TestProviderIntegration:
|
||||
def test_provider_selects_claude_code_without_litellm(self, monkeypatch):
|
||||
import sys
|
||||
|
||||
# litellm intentionally unavailable - the claude code path must not need it
|
||||
monkeypatch.setitem(sys.modules, "litellm", None)
|
||||
with patch(_FIND, return_value="/usr/bin/claude"):
|
||||
provider = LLMProvider("claude-code/sonnet")
|
||||
|
||||
assert isinstance(provider.backend, ClaudeCodeBackend)
|
||||
assert provider.provider_name == "claude-code"
|
||||
assert provider.model_name == "sonnet"
|
||||
|
||||
def test_provider_completes_through_backend(self):
|
||||
with patch(_FIND, return_value="/usr/bin/claude"):
|
||||
provider = LLMProvider("claude-code")
|
||||
with patch("subprocess.run", return_value=_completed(_result_json())):
|
||||
assert provider.complete([{"role": "user", "content": "hi"}]) == "hello from claude"
|
||||
|
||||
def test_generation_kwargs_not_forwarded_to_cli(self):
|
||||
# litellm-style kwargs must not leak into the subprocess call
|
||||
with patch(_FIND, return_value="/usr/bin/claude"):
|
||||
provider = LLMProvider("claude-code", temperature=0.7)
|
||||
with patch("subprocess.run", return_value=_completed(_result_json())) as mock_run:
|
||||
provider.complete([{"role": "user", "content": "hi"}])
|
||||
assert "temperature" not in mock_run.call_args.kwargs
|
||||
|
||||
def test_rate_limit_retries_then_succeeds(self):
|
||||
with patch(_FIND, return_value="/usr/bin/claude"):
|
||||
provider = LLMProvider("claude-code")
|
||||
|
||||
failure = _completed(_result_json(is_error=True, api_error_status=429, result="rate limit"))
|
||||
success = _completed(_result_json())
|
||||
with patch("subprocess.run", side_effect=[failure, success]):
|
||||
with patch("time.sleep"):
|
||||
assert provider.complete([{"role": "user", "content": "hi"}]) == "hello from claude"
|
||||
|
||||
def test_auth_error_fails_fast_without_retrying(self):
|
||||
with patch(_FIND, return_value="/usr/bin/claude"):
|
||||
provider = LLMProvider("claude-code")
|
||||
|
||||
failure = _completed(
|
||||
_result_json(is_error=True, api_error_status=401, result="authentication failed")
|
||||
)
|
||||
with patch("subprocess.run", return_value=failure) as mock_run:
|
||||
with patch("time.sleep"):
|
||||
with pytest.raises(BackendAuthError):
|
||||
provider.complete([{"role": "user", "content": "hi"}], max_retries=3)
|
||||
|
||||
assert mock_run.call_count == 1
|
||||
|
||||
def test_context_window_error_not_retried(self):
|
||||
with patch(_FIND, return_value="/usr/bin/claude"):
|
||||
provider = LLMProvider("claude-code")
|
||||
|
||||
failure = _completed(_result_json(is_error=True, result="context window exceeded"))
|
||||
with patch("subprocess.run", return_value=failure) as mock_run:
|
||||
with patch("time.sleep"):
|
||||
with pytest.raises(BackendContextWindowError):
|
||||
provider.complete([{"role": "user", "content": "hi"}], max_retries=3)
|
||||
|
||||
assert mock_run.call_count == 1
|
||||
|
||||
|
||||
class TestValidation:
|
||||
def _stage(self):
|
||||
from deepzero.engine.stage import StageSpec
|
||||
from deepzero.stages.llm import GenericLLM
|
||||
|
||||
return GenericLLM(StageSpec(name="assess", processor="generic_llm", config={}))
|
||||
|
||||
def _validate(self, model: str) -> list[str]:
|
||||
from deepzero.engine.stage import ProcessorContext
|
||||
|
||||
ctx = ProcessorContext(
|
||||
pipeline_dir=__import__("pathlib").Path("."),
|
||||
global_config={"model": model},
|
||||
llm=None,
|
||||
)
|
||||
# only inspect llm-binding errors, not the unrelated prompt-config error
|
||||
return [e for e in self._stage().validate(ctx) if "backend" in e.lower()]
|
||||
|
||||
def test_claude_code_binding_ok_when_installed(self):
|
||||
with patch(_FIND, return_value="/usr/bin/claude"):
|
||||
assert self._validate("claude-code/sonnet") == []
|
||||
|
||||
def test_claude_code_binding_reports_missing_cli(self):
|
||||
with patch(_FIND, return_value=None):
|
||||
errors = self._validate("claude-code")
|
||||
assert errors and "claude code cli not found" in errors[0].lower()
|
||||
|
||||
def test_claude_code_binding_does_not_require_api_keys(self):
|
||||
# no ANTHROPIC_API_KEY in env, yet validation passes
|
||||
with patch.dict("os.environ", {}, clear=True):
|
||||
with patch(_FIND, return_value="/usr/bin/claude"):
|
||||
assert self._validate("claude-code") == []
|
||||
@@ -0,0 +1,196 @@
|
||||
"""proves the extensibility contract: a new agent-cli backend can be added by
|
||||
subclassing CLIAgentBackend and registering it, with no changes to LLMProvider,
|
||||
the registry internals, the pipeline, or the stages.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import subprocess
|
||||
from typing import ClassVar
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
from deepzero.engine.backends import (
|
||||
ClaudeCodeBackend,
|
||||
CLIAgentBackend,
|
||||
LiteLLMBackend,
|
||||
get_registered_backends,
|
||||
model_scheme,
|
||||
register_backend,
|
||||
resolve_backend_class,
|
||||
validate_model_binding,
|
||||
)
|
||||
from deepzero.engine.backends.base import BackendError
|
||||
from deepzero.engine.backends.registry import _BACKEND_REGISTRY
|
||||
from deepzero.engine.llm import LLMProvider
|
||||
|
||||
|
||||
# a stand-in for a future backend (codex, gemini cli, ...). the only things it
|
||||
# has to supply are argv construction and output parsing.
|
||||
class FakeAgentBackend(CLIAgentBackend):
|
||||
scheme: ClassVar[str] = "faketool"
|
||||
aliases: ClassVar[tuple[str, ...]] = ("ft",)
|
||||
display_name: ClassVar[str] = "fake tool"
|
||||
binary_names: ClassVar[tuple[str, ...]] = ("faketool",)
|
||||
subscription_auth_env_blocklist: ClassVar[tuple[str, ...]] = ("FAKETOOL_API_KEY",)
|
||||
|
||||
def build_argv(self, system: str) -> tuple[list[str], str]:
|
||||
argv = [str(self._binary), "exec", "--json"]
|
||||
if self.alias:
|
||||
argv += ["--model", self.alias]
|
||||
if system:
|
||||
argv += ["--system", system]
|
||||
return argv, ""
|
||||
|
||||
def parse_output(self, returncode: int, stdout: str, stderr: str) -> str:
|
||||
try:
|
||||
payload = json.loads(stdout)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise BackendError(f"bad output: {stdout[:100]}") from exc
|
||||
if payload.get("error"):
|
||||
raise self.classify_error(payload["error"], payload.get("status"))
|
||||
return payload["text"]
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def registered_fake():
|
||||
"""register the fake backend, then restore the registry."""
|
||||
saved = dict(_BACKEND_REGISTRY)
|
||||
register_backend(FakeAgentBackend)
|
||||
try:
|
||||
yield FakeAgentBackend
|
||||
finally:
|
||||
_BACKEND_REGISTRY.clear()
|
||||
_BACKEND_REGISTRY.update(saved)
|
||||
|
||||
|
||||
def _completed(stdout: str, returncode: int = 0, stderr: str = ""):
|
||||
return subprocess.CompletedProcess(args=[], returncode=returncode, stdout=stdout, stderr=stderr)
|
||||
|
||||
|
||||
class TestSchemeParsing:
|
||||
@pytest.mark.parametrize(
|
||||
"model,expected",
|
||||
[
|
||||
("claude-code", "claude-code"),
|
||||
("claude-code/sonnet", "claude-code"),
|
||||
("CLAUDE-CODE/Sonnet", "claude-code"),
|
||||
("openai/gpt-4o", "openai"),
|
||||
("gpt-4o", "gpt-4o"),
|
||||
("", ""),
|
||||
],
|
||||
)
|
||||
def test_scheme(self, model, expected):
|
||||
assert model_scheme(model) == expected
|
||||
|
||||
|
||||
class TestBuiltInResolution:
|
||||
@pytest.mark.parametrize("model", ["claude-code", "claude-code/sonnet", "claude-code/opus"])
|
||||
def test_claude_code_schemes(self, model):
|
||||
assert resolve_backend_class(model) is ClaudeCodeBackend
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model", ["openai/gpt-4o", "anthropic/claude-opus-5", "gemini/pro", "gpt-4o", "test"]
|
||||
)
|
||||
def test_unclaimed_schemes_fall_back_to_litellm(self, model):
|
||||
assert resolve_backend_class(model) is LiteLLMBackend
|
||||
|
||||
def test_claude_code_prefix_is_not_over_matched(self):
|
||||
# "claude-code-extra" is a different scheme, not claude code
|
||||
assert resolve_backend_class("claude-code-extra/x") is LiteLLMBackend
|
||||
|
||||
def test_registry_lists_known_schemes(self):
|
||||
assert "claude-code" in get_registered_backends()
|
||||
|
||||
|
||||
class TestAddingANewBackend:
|
||||
def test_scheme_and_alias_resolve(self, registered_fake):
|
||||
assert resolve_backend_class("faketool") is FakeAgentBackend
|
||||
assert resolve_backend_class("faketool/big-model") is FakeAgentBackend
|
||||
assert resolve_backend_class("ft/big-model") is FakeAgentBackend
|
||||
|
||||
def test_does_not_disturb_existing_backends(self, registered_fake):
|
||||
assert resolve_backend_class("claude-code") is ClaudeCodeBackend
|
||||
assert resolve_backend_class("openai/gpt-4o") is LiteLLMBackend
|
||||
|
||||
def test_provider_drives_it_end_to_end(self, registered_fake):
|
||||
with patch.object(FakeAgentBackend, "find_binary", return_value="/usr/bin/faketool"):
|
||||
provider = LLMProvider("faketool/big-model")
|
||||
|
||||
assert provider.provider_name == "faketool"
|
||||
assert provider.model_name == "big-model"
|
||||
|
||||
with patch("subprocess.run", return_value=_completed('{"text": "fake says hi"}')) as run:
|
||||
assert provider.complete([{"role": "user", "content": "hi"}]) == "fake says hi"
|
||||
|
||||
argv = run.call_args.args[0]
|
||||
assert argv == ["/usr/bin/faketool", "exec", "--json", "--model", "big-model"]
|
||||
assert run.call_args.kwargs["input"] == "hi"
|
||||
|
||||
def test_inherits_retry_semantics_for_free(self, registered_fake):
|
||||
with patch.object(FakeAgentBackend, "find_binary", return_value="/usr/bin/faketool"):
|
||||
provider = LLMProvider("faketool")
|
||||
|
||||
limited = _completed(json.dumps({"error": "rate limit hit", "status": 429}))
|
||||
ok = _completed('{"text": "recovered"}')
|
||||
with patch("subprocess.run", side_effect=[limited, ok]):
|
||||
with patch("time.sleep"):
|
||||
assert provider.complete([{"role": "user", "content": "hi"}]) == "recovered"
|
||||
|
||||
def test_inherits_auth_fail_fast_for_free(self, registered_fake):
|
||||
with patch.object(FakeAgentBackend, "find_binary", return_value="/usr/bin/faketool"):
|
||||
provider = LLMProvider("faketool")
|
||||
|
||||
denied = _completed(json.dumps({"error": "unauthorized", "status": 401}))
|
||||
with patch("subprocess.run", return_value=denied) as run:
|
||||
with patch("time.sleep"):
|
||||
with pytest.raises(BackendError):
|
||||
provider.complete([{"role": "user", "content": "hi"}], max_retries=3)
|
||||
assert run.call_count == 1
|
||||
|
||||
def test_inherits_env_sanitization_for_free(self, registered_fake):
|
||||
with patch.dict("os.environ", {"FAKETOOL_API_KEY": "secret"}, clear=False):
|
||||
with patch.object(FakeAgentBackend, "find_binary", return_value="/usr/bin/faketool"):
|
||||
env = FakeAgentBackend("faketool").build_env()
|
||||
assert "FAKETOOL_API_KEY" not in env
|
||||
|
||||
def test_inherits_validation_wiring_for_free(self, registered_fake):
|
||||
with patch.object(FakeAgentBackend, "find_binary", return_value=None):
|
||||
errors = validate_model_binding("faketool")
|
||||
assert errors and "fake tool cli not found" in errors[0].lower()
|
||||
|
||||
with patch.object(FakeAgentBackend, "find_binary", return_value="/usr/bin/faketool"):
|
||||
assert validate_model_binding("faketool") == []
|
||||
|
||||
|
||||
class TestRegistrationRules:
|
||||
def test_backend_without_scheme_is_rejected(self):
|
||||
class Anonymous(CLIAgentBackend):
|
||||
def build_argv(self, system):
|
||||
return [], ""
|
||||
|
||||
def parse_output(self, returncode, stdout, stderr):
|
||||
return ""
|
||||
|
||||
with pytest.raises(ValueError, match="scheme"):
|
||||
register_backend(Anonymous)
|
||||
|
||||
def test_register_returns_class_for_decorator_use(self):
|
||||
saved = dict(_BACKEND_REGISTRY)
|
||||
try:
|
||||
|
||||
class Decorated(CLIAgentBackend):
|
||||
scheme = "decorated"
|
||||
|
||||
def build_argv(self, system):
|
||||
return [], ""
|
||||
|
||||
def parse_output(self, returncode, stdout, stderr):
|
||||
return ""
|
||||
|
||||
assert register_backend(Decorated) is Decorated
|
||||
finally:
|
||||
_BACKEND_REGISTRY.clear()
|
||||
_BACKEND_REGISTRY.update(saved)
|
||||
Reference in New Issue
Block a user