mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
refactor(e2e): type claude_code suite core modules under strict basedpyright
This commit is contained in:
parent
ff4a40f017
commit
949aa53096
10 changed files with 270 additions and 128 deletions
|
|
@ -1,7 +1,7 @@
|
|||
{
|
||||
"include": ["litellm"],
|
||||
"ignore": [],
|
||||
"exclude": ["**/node_modules", "**/__pycache__", "tests/e2e/claude_code", "litellm/types/utils.py", "litellm/proxy/_types.py"],
|
||||
"exclude": ["**/node_modules", "**/__pycache__", "litellm/types/utils.py", "litellm/proxy/_types.py"],
|
||||
"pythonVersion": "3.12",
|
||||
"typeCheckingMode": "strict",
|
||||
"enableTypeIgnoreComments": false,
|
||||
|
|
|
|||
|
|
@ -28,7 +28,7 @@ collecting this module as a test file.
|
|||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from typing import Any, Mapping, Sequence
|
||||
from typing import Mapping, Sequence
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -37,6 +37,8 @@ from claude_code.cli_driver import (
|
|||
failure_diagnostic,
|
||||
run_claude_models_parallel,
|
||||
)
|
||||
from claude_code.conftest import CompatResult
|
||||
from claude_code.json_types import JSONValue
|
||||
|
||||
PROXY_BASE_URL_ENV = "LITELLM_PROXY_BASE_URL"
|
||||
PROXY_API_KEY_ENV = "LITELLM_PROXY_API_KEY"
|
||||
|
|
@ -53,7 +55,7 @@ PROXY_API_KEY_ENV = "LITELLM_PROXY_API_KEY"
|
|||
MIN_STREAM_DELTA_EVENTS = 2
|
||||
|
||||
|
||||
def _count_stream_event_deltas(events: Sequence[Mapping[str, Any]]) -> int:
|
||||
def _count_stream_event_deltas(events: Sequence[Mapping[str, JSONValue]]) -> int:
|
||||
"""Count `stream_event` records that carry an SSE event payload.
|
||||
|
||||
With `--include-partial-messages`, Claude Code wraps every upstream
|
||||
|
|
@ -75,7 +77,7 @@ def _count_stream_event_deltas(events: Sequence[Mapping[str, Any]]) -> int:
|
|||
|
||||
def run_basic_messaging_cell(
|
||||
*,
|
||||
compat_result,
|
||||
compat_result: CompatResult,
|
||||
models: Sequence[str],
|
||||
prompt: str,
|
||||
verify_streaming: bool = False,
|
||||
|
|
@ -128,7 +130,7 @@ def run_basic_messaging_cell(
|
|||
extra_args=extra_args,
|
||||
)
|
||||
|
||||
failures = []
|
||||
failures: list[str] = []
|
||||
for model in models:
|
||||
outcome = outcomes[model]
|
||||
if isinstance(outcome, ClaudeCLIError):
|
||||
|
|
|
|||
|
|
@ -13,7 +13,6 @@ live in `matrix_builder.py`.
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import shutil
|
||||
import subprocess
|
||||
|
|
@ -22,8 +21,11 @@ import tempfile
|
|||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Callable, Dict, List, Mapping, Optional, Sequence, Tuple, Union
|
||||
from typing import Mapping, Optional, Protocol, Sequence, Union
|
||||
|
||||
from pydantic import ValidationError
|
||||
|
||||
from claude_code.json_types import JSON_OBJECT_ADAPTER, JSONObject, JSONValue
|
||||
from claude_code.rate_limiter import (
|
||||
RateLimiter,
|
||||
get_default_limiter,
|
||||
|
|
@ -58,7 +60,7 @@ DEFAULT_TIMEOUT_SECONDS = float(
|
|||
# or `~/.bash_history` on the cron VM (and on the CircleCI executor
|
||||
# the same isolation prevents accidentally exposing checkout-adjacent
|
||||
# files even though the runner home is ephemeral there).
|
||||
_CLI_ENV_ALLOWLIST: tuple = (
|
||||
_CLI_ENV_ALLOWLIST: tuple[str, ...] = (
|
||||
"PATH",
|
||||
"USER",
|
||||
"LOGNAME",
|
||||
|
|
@ -113,13 +115,51 @@ class DriverResult:
|
|||
"""
|
||||
|
||||
text: str
|
||||
events: List[Dict[str, Any]] = field(default_factory=list)
|
||||
events: list[JSONObject] = field(default_factory=list)
|
||||
exit_code: int = 0
|
||||
stderr: str = ""
|
||||
usage: Optional[Dict[str, Any]] = None
|
||||
usage: Optional[JSONObject] = None
|
||||
duration_ms: Optional[int] = None
|
||||
|
||||
|
||||
class CommandRunner(Protocol):
|
||||
def __call__(
|
||||
self,
|
||||
args: Sequence[str],
|
||||
/,
|
||||
*,
|
||||
env: Mapping[str, str],
|
||||
input: Optional[str],
|
||||
capture_output: bool,
|
||||
text: bool,
|
||||
timeout: float,
|
||||
check: bool,
|
||||
) -> subprocess.CompletedProcess[str]: ...
|
||||
|
||||
|
||||
def _run_subprocess(
|
||||
args: Sequence[str],
|
||||
/,
|
||||
*,
|
||||
env: Mapping[str, str],
|
||||
input: Optional[str],
|
||||
capture_output: bool,
|
||||
text: bool,
|
||||
timeout: float,
|
||||
check: bool,
|
||||
) -> subprocess.CompletedProcess[str]:
|
||||
del text
|
||||
return subprocess.run(
|
||||
args,
|
||||
env=env,
|
||||
input=input,
|
||||
capture_output=capture_output,
|
||||
text=True,
|
||||
timeout=timeout,
|
||||
check=check,
|
||||
)
|
||||
|
||||
|
||||
def run_claude(
|
||||
*,
|
||||
prompt: Optional[str],
|
||||
|
|
@ -131,7 +171,7 @@ def run_claude(
|
|||
stdin_input: Optional[str] = None,
|
||||
cli_path: str = CLAUDE_CLI_DEFAULT,
|
||||
timeout: float = DEFAULT_TIMEOUT_SECONDS,
|
||||
runner: Optional[Any] = None,
|
||||
runner: Optional[CommandRunner] = None,
|
||||
rate_limiter: Optional[RateLimiter] = None,
|
||||
) -> DriverResult:
|
||||
"""Invoke `claude` once in headless stream-JSON mode and return the result.
|
||||
|
|
@ -179,7 +219,7 @@ def run_claude(
|
|||
# the prompt terminates option parsing and leaves the prompt as a
|
||||
# plain positional, which works for variadic and non-variadic flags
|
||||
# alike.
|
||||
cmd: List[str] = [
|
||||
cmd: list[str] = [
|
||||
cli_path,
|
||||
"--print",
|
||||
"--output-format",
|
||||
|
|
@ -198,7 +238,7 @@ def run_claude(
|
|||
# process-runtime vars from os.environ, plus the explicit proxy
|
||||
# creds, plus any caller-supplied overrides. See _CLI_ENV_ALLOWLIST
|
||||
# above for the security rationale.
|
||||
env: Dict[str, str] = {
|
||||
env: dict[str, str] = {
|
||||
key: os.environ[key] for key in _CLI_ENV_ALLOWLIST if key in os.environ
|
||||
}
|
||||
env["ANTHROPIC_BASE_URL"] = base_url
|
||||
|
|
@ -221,7 +261,7 @@ def run_claude(
|
|||
provider = infer_provider(model)
|
||||
limiter.acquire(provider)
|
||||
|
||||
run_fn = runner or subprocess.run
|
||||
run_fn = runner or _run_subprocess
|
||||
try:
|
||||
try:
|
||||
completed = run_fn(
|
||||
|
|
@ -276,8 +316,8 @@ def run_claude_models_parallel(
|
|||
stdin_input: Optional[str] = None,
|
||||
cli_path: str = CLAUDE_CLI_DEFAULT,
|
||||
timeout: float = DEFAULT_TIMEOUT_SECONDS,
|
||||
runner: Optional[Callable[..., Any]] = None,
|
||||
) -> Dict[str, ModelResult]:
|
||||
runner: Optional[CommandRunner] = None,
|
||||
) -> dict[str, ModelResult]:
|
||||
"""Invoke `run_claude` for every `models[i]` concurrently and collect outcomes.
|
||||
|
||||
Each `claude` CLI invocation is a long-lived subprocess that spends
|
||||
|
|
@ -300,7 +340,7 @@ def run_claude_models_parallel(
|
|||
if not models:
|
||||
raise ValueError("models must be a non-empty sequence")
|
||||
|
||||
def _one(model: str) -> Tuple[str, ModelResult, float]:
|
||||
def _one(model: str) -> tuple[str, ModelResult, float]:
|
||||
# Per-model wall clock: this is what the matrix run actually pays for.
|
||||
# We record it whether the run succeeded or raised so the breakdown
|
||||
# log below covers both code paths and surfaces "which model is the
|
||||
|
|
@ -342,8 +382,8 @@ def run_claude_models_parallel(
|
|||
wrapped.__cause__ = exc
|
||||
return model, wrapped, elapsed
|
||||
|
||||
outcomes: Dict[str, ModelResult] = {}
|
||||
durations: Dict[str, float] = {}
|
||||
outcomes: dict[str, ModelResult] = {}
|
||||
durations: dict[str, float] = {}
|
||||
overall_started = time.monotonic()
|
||||
with ThreadPoolExecutor(max_workers=len(models)) as pool:
|
||||
futures = [pool.submit(_one, model) for model in models]
|
||||
|
|
@ -380,11 +420,11 @@ def _log_parallel_breakdown(
|
|||
"why didn't this get faster?".
|
||||
"""
|
||||
sequential_total = sum(durations.values())
|
||||
slowest_model = max(durations, key=durations.get) if durations else None
|
||||
slowest_model = max(durations, key=lambda model: durations[model]) if durations else None
|
||||
slowest = durations[slowest_model] if slowest_model else 0.0
|
||||
speedup = sequential_total / overall_elapsed if overall_elapsed > 0 else 0.0
|
||||
|
||||
lines: List[str] = []
|
||||
lines: list[str] = []
|
||||
lines.append("[parallel] per-model wall time:")
|
||||
for model in models:
|
||||
elapsed = durations.get(model, 0.0)
|
||||
|
|
@ -406,28 +446,27 @@ def _log_parallel_breakdown(
|
|||
print("\n".join(lines), file=sys.stderr, flush=True)
|
||||
|
||||
|
||||
def _parse_stream_json(stdout: str) -> List[Dict[str, Any]]:
|
||||
def _parse_stream_json(stdout: str) -> list[JSONObject]:
|
||||
"""Parse newline-delimited JSON emitted by `claude --output-format stream-json`.
|
||||
|
||||
Lines that don't parse as JSON are silently skipped — the CLI occasionally
|
||||
emits debug output we don't care about, and a single malformed line should
|
||||
not abort the whole run. Real failure modes surface via exit code.
|
||||
"""
|
||||
events: List[Dict[str, Any]] = []
|
||||
events: list[JSONObject] = []
|
||||
for line in stdout.splitlines():
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
try:
|
||||
obj = json.loads(line)
|
||||
except json.JSONDecodeError:
|
||||
event = JSON_OBJECT_ADAPTER.validate_json(line)
|
||||
except ValidationError:
|
||||
continue
|
||||
if isinstance(obj, dict):
|
||||
events.append(obj)
|
||||
events.append(event)
|
||||
return events
|
||||
|
||||
|
||||
def _extract_assistant_text(events: Sequence[Mapping[str, Any]]) -> str:
|
||||
def _extract_assistant_text(events: Sequence[Mapping[str, JSONValue]]) -> str:
|
||||
"""Concatenate the text content of every `assistant` event in order.
|
||||
|
||||
The non-streaming `--print` path emits a single `assistant` event whose
|
||||
|
|
@ -435,12 +474,12 @@ def _extract_assistant_text(events: Sequence[Mapping[str, Any]]) -> str:
|
|||
join every `text` block — the CLI prints other block types (e.g.
|
||||
`tool_use`) which we ignore for the basic-messaging case.
|
||||
"""
|
||||
chunks: List[str] = []
|
||||
chunks: list[str] = []
|
||||
for event in events:
|
||||
if event.get("type") != "assistant":
|
||||
continue
|
||||
message = event.get("message") or {}
|
||||
content = message.get("content")
|
||||
message = event.get("message")
|
||||
content = message.get("content") if isinstance(message, dict) else None
|
||||
if isinstance(content, str):
|
||||
chunks.append(content)
|
||||
continue
|
||||
|
|
@ -449,8 +488,9 @@ def _extract_assistant_text(events: Sequence[Mapping[str, Any]]) -> str:
|
|||
for block in content:
|
||||
if not isinstance(block, dict):
|
||||
continue
|
||||
if block.get("type") == "text" and isinstance(block.get("text"), str):
|
||||
chunks.append(block["text"])
|
||||
text = block.get("text")
|
||||
if block.get("type") == "text" and isinstance(text, str):
|
||||
chunks.append(text)
|
||||
return "".join(chunks)
|
||||
|
||||
|
||||
|
|
@ -478,7 +518,7 @@ def failure_diagnostic(result: "DriverResult", *, max_len: int = 800) -> str:
|
|||
page from a misbehaving load balancer doesn't blow up the matrix
|
||||
JSON.
|
||||
"""
|
||||
pieces: List[str] = [f"exit={result.exit_code}"]
|
||||
pieces: list[str] = [f"exit={result.exit_code}"]
|
||||
|
||||
# api_error_status only appears on the final `result` event when the
|
||||
# CLI received an HTTP error from the upstream API. Surfacing it
|
||||
|
|
@ -503,7 +543,7 @@ def failure_diagnostic(result: "DriverResult", *, max_len: int = 800) -> str:
|
|||
|
||||
|
||||
def _extract_api_error_status(
|
||||
events: Sequence[Mapping[str, Any]],
|
||||
events: Sequence[Mapping[str, JSONValue]],
|
||||
) -> Optional[int]:
|
||||
"""Return the `api_error_status` from the last `result` event, if any."""
|
||||
for event in reversed(list(events)):
|
||||
|
|
@ -521,14 +561,14 @@ def _truncate(s: str, max_len: int) -> str:
|
|||
return s[:max_len] + "...(truncated)"
|
||||
|
||||
|
||||
def _extract_usage(events: Sequence[Mapping[str, Any]]) -> Optional[Dict[str, Any]]:
|
||||
def _extract_usage(events: Sequence[Mapping[str, JSONValue]]) -> Optional[JSONObject]:
|
||||
"""Return the most recent `usage` block seen on any event, if any.
|
||||
|
||||
The CLI surfaces token + cache usage on the final `result` event for
|
||||
non-streaming runs, but earlier events also carry partial usage in some
|
||||
versions; taking the last non-empty one is the safe default.
|
||||
"""
|
||||
last: Optional[Dict[str, Any]] = None
|
||||
last: Optional[JSONObject] = None
|
||||
for event in events:
|
||||
usage = event.get("usage")
|
||||
if isinstance(usage, dict) and usage:
|
||||
|
|
|
|||
|
|
@ -38,10 +38,14 @@ import sys
|
|||
from collections import Counter, defaultdict
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, FrozenSet, List, Optional, Tuple
|
||||
from typing import Generator, Mapping, Optional, Sequence, TypedDict
|
||||
|
||||
import pluggy
|
||||
import pytest
|
||||
import yaml
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
from claude_code.json_types import JSON_OBJECT_ADAPTER, JSONValue
|
||||
|
||||
VALID_STATUSES = {"pass", "fail", "not_applicable", "not_tested"}
|
||||
RESULTS_ARTIFACT_ENV = "COMPAT_RESULTS_PATH"
|
||||
|
|
@ -86,14 +90,14 @@ class CompatResult:
|
|||
pass" rule.
|
||||
"""
|
||||
|
||||
value: Optional[Dict[str, Any]] = None
|
||||
values: List[Dict[str, Any]] = field(default_factory=list)
|
||||
value: Optional[dict[str, JSONValue]] = None
|
||||
values: list[dict[str, JSONValue]] = field(default_factory=list)
|
||||
|
||||
def set(self, result: Dict[str, Any]) -> None:
|
||||
def set(self, result: Mapping[str, JSONValue]) -> None:
|
||||
validated = self._validate(result)
|
||||
self.value = validated
|
||||
|
||||
def add(self, result: Dict[str, Any]) -> None:
|
||||
def add(self, result: Mapping[str, JSONValue]) -> None:
|
||||
"""Append one model's outcome to the per-test results list.
|
||||
|
||||
Use this when a single test exercises multiple Claude tiers
|
||||
|
|
@ -104,7 +108,7 @@ class CompatResult:
|
|||
self.values.append(validated)
|
||||
|
||||
@staticmethod
|
||||
def _validate(result: Dict[str, Any]) -> Dict[str, Any]:
|
||||
def _validate(result: Mapping[str, JSONValue]) -> dict[str, JSONValue]:
|
||||
if not isinstance(result, dict):
|
||||
raise TypeError("compat_result requires a dict")
|
||||
status = result.get("status")
|
||||
|
|
@ -121,7 +125,7 @@ class CompatResult:
|
|||
)
|
||||
return dict(result)
|
||||
|
||||
def collected(self) -> List[Dict[str, Any]]:
|
||||
def collected(self) -> list[dict[str, JSONValue]]:
|
||||
"""Return every result reported during the test, preserving order.
|
||||
|
||||
Multi-model tests use `.add(...)` per model; legacy tests use
|
||||
|
|
@ -140,12 +144,12 @@ class _CollectedResult:
|
|||
feature_id: str
|
||||
provider: str
|
||||
nodeid: str
|
||||
result: Dict[str, Any]
|
||||
result: dict[str, JSONValue]
|
||||
|
||||
|
||||
@dataclass
|
||||
class _Collector:
|
||||
items: List[_CollectedResult] = field(default_factory=list)
|
||||
items: list[_CollectedResult] = field(default_factory=list)
|
||||
|
||||
|
||||
_COLLECTOR = _Collector()
|
||||
|
|
@ -164,7 +168,7 @@ def compat_result() -> CompatResult:
|
|||
|
||||
|
||||
@functools.lru_cache(maxsize=1)
|
||||
def _manifest_feature_ids() -> FrozenSet[str]:
|
||||
def _manifest_feature_ids() -> frozenset[str]:
|
||||
"""Return the set of feature_ids declared in `manifest.yaml`.
|
||||
|
||||
Used as a positive filter so only directories that correspond to a
|
||||
|
|
@ -179,22 +183,20 @@ def _manifest_feature_ids() -> FrozenSet[str]:
|
|||
"""
|
||||
manifest_path = Path(__file__).resolve().parent / "manifest.yaml"
|
||||
try:
|
||||
raw = yaml.safe_load(manifest_path.read_text())
|
||||
except (OSError, yaml.YAMLError):
|
||||
return frozenset()
|
||||
if not isinstance(raw, dict):
|
||||
raw = JSON_OBJECT_ADAPTER.validate_python(yaml.safe_load(manifest_path.read_text()))
|
||||
except (OSError, yaml.YAMLError, ValidationError):
|
||||
return frozenset()
|
||||
features = raw.get("features")
|
||||
if not isinstance(features, list):
|
||||
return frozenset()
|
||||
return frozenset(
|
||||
entry["id"]
|
||||
feature_id
|
||||
for entry in features
|
||||
if isinstance(entry, dict) and isinstance(entry.get("id"), str)
|
||||
if isinstance(entry, dict) and isinstance(feature_id := entry.get("id"), str)
|
||||
)
|
||||
|
||||
|
||||
def _infer_feature_and_provider(node_path: Path) -> Optional[tuple]:
|
||||
def _infer_feature_and_provider(node_path: Path) -> Optional[tuple[str, str]]:
|
||||
"""Infer (feature_id, provider) from a test file path.
|
||||
|
||||
Path shape: tests/e2e/claude_code/<feature_id>/test_<provider>.py
|
||||
|
|
@ -216,7 +218,9 @@ def _infer_feature_and_provider(node_path: Path) -> Optional[tuple]:
|
|||
|
||||
|
||||
@pytest.hookimpl(hookwrapper=True)
|
||||
def pytest_runtest_makereport(item, call):
|
||||
def pytest_runtest_makereport(
|
||||
item: pytest.Item, call: pytest.CallInfo[None]
|
||||
) -> Generator[None, pluggy.Result[pytest.TestReport], None]:
|
||||
"""Capture compat_result reports at end-of-test and remember them for the artifact.
|
||||
|
||||
A single test may report multiple results (one per Claude tier when
|
||||
|
|
@ -256,8 +260,8 @@ def pytest_runtest_makereport(item, call):
|
|||
return
|
||||
feature_id, provider = inferred
|
||||
|
||||
fixture = item.funcargs.get("compat_result") if hasattr(item, "funcargs") else None
|
||||
collected: List[Dict[str, Any]] = (
|
||||
fixture = item.funcargs.get("compat_result") if isinstance(item, pytest.Function) else None
|
||||
collected: list[dict[str, JSONValue]] = (
|
||||
fixture.collected() if isinstance(fixture, CompatResult) else []
|
||||
)
|
||||
|
||||
|
|
@ -298,7 +302,7 @@ def pytest_runtest_makereport(item, call):
|
|||
)
|
||||
|
||||
|
||||
def _is_xdist_worker(session) -> bool:
|
||||
def _is_xdist_worker(session: pytest.Session) -> bool:
|
||||
"""Return True iff the current pytest session is an xdist worker.
|
||||
|
||||
The standard idiom is to look up `workerinput` on the config; the
|
||||
|
|
@ -309,11 +313,19 @@ def _is_xdist_worker(session) -> bool:
|
|||
return hasattr(session.config, "workerinput")
|
||||
|
||||
|
||||
def _xdist_worker_id(session) -> Optional[str]:
|
||||
info = getattr(session.config, "workerinput", None)
|
||||
_WORKERINPUT_ADAPTER = TypeAdapter[dict[str, object]](dict[str, object])
|
||||
|
||||
|
||||
def _xdist_worker_id(session: pytest.Session) -> Optional[str]:
|
||||
info: object = getattr(session.config, "workerinput", None)
|
||||
if not info:
|
||||
return None
|
||||
return info.get("workerid")
|
||||
try:
|
||||
workerinput = _WORKERINPUT_ADAPTER.validate_python(info)
|
||||
except ValidationError:
|
||||
return None
|
||||
worker_id = workerinput.get("workerid")
|
||||
return worker_id if isinstance(worker_id, str) else None
|
||||
|
||||
|
||||
def _shard_dir(artifact_path: Path) -> Path:
|
||||
|
|
@ -327,7 +339,7 @@ def _shard_dir(artifact_path: Path) -> Path:
|
|||
return artifact_path.with_name(artifact_path.name + ".shards")
|
||||
|
||||
|
||||
def _serialize_items(items: List["_CollectedResult"]) -> List[Dict[str, Any]]:
|
||||
def _serialize_items(items: Sequence[_CollectedResult]) -> list[dict[str, JSONValue]]:
|
||||
return [
|
||||
{
|
||||
"feature_id": item.feature_id,
|
||||
|
|
@ -339,9 +351,15 @@ def _serialize_items(items: List["_CollectedResult"]) -> List[Dict[str, Any]]:
|
|||
]
|
||||
|
||||
|
||||
class _RateLimitSummary(TypedDict):
|
||||
totals: dict[str, int]
|
||||
per_provider: dict[str, dict[str, int]]
|
||||
rate_limited_examples: list[dict[str, JSONValue]]
|
||||
|
||||
|
||||
def _build_rate_limit_summary(
|
||||
rows: List[Dict[str, Any]],
|
||||
) -> Dict[str, Any]:
|
||||
rows: Sequence[JSONValue],
|
||||
) -> _RateLimitSummary:
|
||||
"""Aggregate per-provider rate-limit signals from the result rows.
|
||||
|
||||
We classify any failure whose error string matches `_RATE_LIMIT_RE`
|
||||
|
|
@ -363,14 +381,19 @@ def _build_rate_limit_summary(
|
|||
],
|
||||
}
|
||||
"""
|
||||
totals: Counter = Counter()
|
||||
per_provider: Dict[str, Counter] = defaultdict(Counter)
|
||||
rate_limited_examples: List[Dict[str, Any]] = []
|
||||
totals: Counter[str] = Counter()
|
||||
per_provider: defaultdict[str, Counter[str]] = defaultdict(Counter)
|
||||
rate_limited_examples: list[dict[str, JSONValue]] = []
|
||||
|
||||
for row in rows:
|
||||
result = row.get("result") or {}
|
||||
status = result.get("status") or "unknown"
|
||||
provider = row.get("provider") or "unknown"
|
||||
if not isinstance(row, dict):
|
||||
continue
|
||||
raw_result = row.get("result")
|
||||
result = raw_result if isinstance(raw_result, dict) else {}
|
||||
raw_status = result.get("status")
|
||||
status = raw_status if isinstance(raw_status, str) and raw_status else "unknown"
|
||||
raw_provider = row.get("provider")
|
||||
provider = raw_provider if isinstance(raw_provider, str) and raw_provider else "unknown"
|
||||
totals[status] += 1
|
||||
per_provider[provider][status] += 1
|
||||
|
||||
|
|
@ -397,7 +420,7 @@ def _build_rate_limit_summary(
|
|||
}
|
||||
|
||||
|
||||
def _print_rate_limit_summary(summary: Dict[str, Any]) -> None:
|
||||
def _print_rate_limit_summary(summary: _RateLimitSummary) -> None:
|
||||
"""Emit a human-readable per-provider table to stderr.
|
||||
|
||||
Pytest only captures stderr when `-s` isn't set; we deliberately
|
||||
|
|
@ -405,9 +428,9 @@ def _print_rate_limit_summary(summary: Dict[str, Any]) -> None:
|
|||
with `-q` and grep-checks the structured JSON artifact, while a
|
||||
human running locally with `-s` sees the same numbers inline.
|
||||
"""
|
||||
totals = summary.get("totals", {})
|
||||
per_provider = summary.get("per_provider", {})
|
||||
lines: List[str] = []
|
||||
totals = summary["totals"]
|
||||
per_provider = summary["per_provider"]
|
||||
lines: list[str] = []
|
||||
lines.append("[compat] session totals:")
|
||||
for status in ("pass", "fail", "rate_limited", "not_applicable", "not_tested"):
|
||||
if status in totals:
|
||||
|
|
@ -430,7 +453,7 @@ def _print_rate_limit_summary(summary: Dict[str, Any]) -> None:
|
|||
print("\n".join(lines), file=sys.stderr, flush=True)
|
||||
|
||||
|
||||
def pytest_sessionstart(session):
|
||||
def pytest_sessionstart(session: pytest.Session) -> None:
|
||||
"""Reset per-session state before tests run.
|
||||
|
||||
Two responsibilities:
|
||||
|
|
@ -469,7 +492,7 @@ def pytest_sessionstart(session):
|
|||
continue
|
||||
|
||||
|
||||
def pytest_sessionfinish(session, exitstatus):
|
||||
def pytest_sessionfinish(session: pytest.Session, exitstatus: int) -> None:
|
||||
"""Write the per-process results shard, then merge if we're the controller.
|
||||
|
||||
Worker processes (xdist `gw0`, `gw1`, ...) only write their shard
|
||||
|
|
@ -517,10 +540,10 @@ def pytest_sessionfinish(session, exitstatus):
|
|||
if _is_xdist_worker(session):
|
||||
return
|
||||
|
||||
merged_rows: List[Dict[str, Any]] = []
|
||||
merged_rows: list[JSONValue] = []
|
||||
for shard_file in sorted(shard_dir.glob("*.json")):
|
||||
try:
|
||||
shard = json.loads(shard_file.read_text())
|
||||
shard = JSON_OBJECT_ADAPTER.validate_json(shard_file.read_text())
|
||||
except (OSError, ValueError):
|
||||
continue
|
||||
rows = shard.get("results")
|
||||
|
|
|
|||
|
|
@ -15,13 +15,28 @@ from __future__ import annotations
|
|||
import argparse
|
||||
import datetime
|
||||
import sys
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2]))
|
||||
|
||||
from claude_code.matrix_builder import build_from_paths # noqa: E402 # import needs the sys.path bootstrap above
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Args:
|
||||
manifest: Path
|
||||
results: Path
|
||||
output: Path
|
||||
litellm_version: str
|
||||
claude_code_version: str
|
||||
|
||||
|
||||
_ARGS_ADAPTER = TypeAdapter[_Args](_Args)
|
||||
|
||||
|
||||
def main() -> int:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--manifest", type=Path, required=True)
|
||||
|
|
@ -29,7 +44,7 @@ def main() -> int:
|
|||
parser.add_argument("--output", type=Path, required=True)
|
||||
parser.add_argument("--litellm-version", required=True)
|
||||
parser.add_argument("--claude-code-version", required=True)
|
||||
args = parser.parse_args()
|
||||
args = _ARGS_ADAPTER.validate_python(vars(parser.parse_args()))
|
||||
|
||||
generated_at = datetime.datetime.now(datetime.timezone.utc).strftime(
|
||||
"%Y-%m-%dT%H:%M:%SZ"
|
||||
|
|
|
|||
|
|
@ -26,12 +26,13 @@ bug).
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Mapping, Optional
|
||||
from typing import Mapping, Optional
|
||||
|
||||
import httpx
|
||||
from pydantic import ValidationError
|
||||
|
||||
from claude_code.json_types import JSON_VALUE_ADAPTER, JSONValue
|
||||
from claude_code.rate_limiter import (
|
||||
RateLimiter,
|
||||
get_default_limiter,
|
||||
|
|
@ -55,7 +56,7 @@ class ProbeResult:
|
|||
|
||||
status_code: int
|
||||
body: str
|
||||
payload: Optional[Mapping[str, Any]] = None
|
||||
payload: Optional[JSONValue] = None
|
||||
error: Optional[str] = None
|
||||
|
||||
|
||||
|
|
@ -108,8 +109,8 @@ def probe_count_tokens(
|
|||
|
||||
body = response.text or ""
|
||||
try:
|
||||
payload = response.json() if body else None
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
payload = JSON_VALUE_ADAPTER.validate_json(body) if body else None
|
||||
except ValidationError:
|
||||
payload = None
|
||||
|
||||
return ProbeResult(
|
||||
|
|
@ -211,8 +212,8 @@ def probe_tool_search(
|
|||
|
||||
body = response.text or ""
|
||||
try:
|
||||
payload_out = response.json() if body else None
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
payload_out = JSON_VALUE_ADAPTER.validate_json(body) if body else None
|
||||
except ValidationError:
|
||||
payload_out = None
|
||||
|
||||
return ProbeResult(
|
||||
|
|
|
|||
18
tests/e2e/claude_code/json_types.py
Normal file
18
tests/e2e/claude_code/json_types.py
Normal file
|
|
@ -0,0 +1,18 @@
|
|||
"""Shared JSON typing vocabulary for the Claude Code compat suite.
|
||||
|
||||
Everything the suite parses off a wire or a file (stream-json events from
|
||||
the `claude` CLI, npm packuments, manifest YAML, results artifacts) is
|
||||
JSON-shaped. `JSONValue` models that shape recursively so strict type
|
||||
checking can narrow payloads with plain `isinstance` checks, and the
|
||||
`TypeAdapter`s validate untyped input (`json.loads`, `yaml.safe_load`,
|
||||
HTTP bodies) into the typed shape at the boundary instead of letting
|
||||
`Any` leak through the suite.
|
||||
"""
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
type JSONValue = str | int | float | bool | None | list[JSONValue] | dict[str, JSONValue]
|
||||
type JSONObject = dict[str, JSONValue]
|
||||
|
||||
JSON_VALUE_ADAPTER = TypeAdapter[JSONValue](JSONValue)
|
||||
JSON_OBJECT_ADAPTER = TypeAdapter[JSONObject](JSONObject)
|
||||
|
|
@ -15,9 +15,12 @@ from __future__ import annotations
|
|||
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Mapping, Optional, Sequence
|
||||
from typing import Mapping, Optional, Sequence
|
||||
|
||||
import yaml
|
||||
from pydantic import ValidationError
|
||||
|
||||
from claude_code.json_types import JSON_OBJECT_ADAPTER, JSONValue
|
||||
|
||||
SCHEMA_VERSION = "1"
|
||||
VALID_STATUSES = {"pass", "fail", "not_applicable", "not_tested"}
|
||||
|
|
@ -31,14 +34,15 @@ class ResultsError(ValueError):
|
|||
"""Raised when the pytest results artifact is malformed."""
|
||||
|
||||
|
||||
def load_manifest(path: Path) -> Dict[str, Any]:
|
||||
def load_manifest(path: Path) -> dict[str, JSONValue]:
|
||||
"""Load and validate `manifest.yaml`.
|
||||
|
||||
Returns a dict with keys: schema_version, providers, features. Raises
|
||||
ManifestError on missing fields or schema mismatch.
|
||||
"""
|
||||
raw = yaml.safe_load(path.read_text())
|
||||
if not isinstance(raw, dict):
|
||||
try:
|
||||
raw = JSON_OBJECT_ADAPTER.validate_python(yaml.safe_load(path.read_text()))
|
||||
except ValidationError:
|
||||
raise ManifestError(f"manifest at {path} is not a mapping")
|
||||
schema_version = str(raw.get("schema_version", ""))
|
||||
if schema_version != SCHEMA_VERSION:
|
||||
|
|
@ -60,22 +64,26 @@ def load_manifest(path: Path) -> Dict[str, Any]:
|
|||
return raw
|
||||
|
||||
|
||||
def load_results(path: Path) -> List[Dict[str, Any]]:
|
||||
def load_results(path: Path) -> list[JSONValue]:
|
||||
"""Load the pytest results artifact and return its `results` list."""
|
||||
raw = json.loads(path.read_text())
|
||||
if not isinstance(raw, dict) or not isinstance(raw.get("results"), list):
|
||||
try:
|
||||
raw = JSON_OBJECT_ADAPTER.validate_python(json.loads(path.read_text()))
|
||||
except ValidationError:
|
||||
raise ResultsError(f"results artifact at {path} has no `results` list")
|
||||
return raw["results"]
|
||||
results = raw.get("results")
|
||||
if not isinstance(results, list):
|
||||
raise ResultsError(f"results artifact at {path} has no `results` list")
|
||||
return results
|
||||
|
||||
|
||||
def build_matrix(
|
||||
*,
|
||||
manifest: Mapping[str, Any],
|
||||
results: Sequence[Mapping[str, Any]],
|
||||
manifest: Mapping[str, JSONValue],
|
||||
results: Sequence[JSONValue],
|
||||
litellm_version: str,
|
||||
claude_code_version: str,
|
||||
generated_at: str,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, JSONValue]:
|
||||
"""Build the published matrix JSON from pre-loaded inputs.
|
||||
|
||||
Empty cells (no test ran for a (feature, provider) and no
|
||||
|
|
@ -85,26 +93,40 @@ def build_matrix(
|
|||
to `pass` only if every model passed; otherwise `fail` with the first
|
||||
breaking model surfaced in the error.
|
||||
"""
|
||||
providers: List[str] = list(manifest["providers"])
|
||||
feature_specs: List[Dict[str, Any]] = list(manifest["features"])
|
||||
providers_value = manifest["providers"]
|
||||
providers: list[str] = (
|
||||
[str(provider) for provider in providers_value]
|
||||
if isinstance(providers_value, list)
|
||||
else []
|
||||
)
|
||||
features_value = manifest["features"]
|
||||
feature_specs: list[dict[str, JSONValue]] = (
|
||||
[spec for spec in features_value if isinstance(spec, dict)]
|
||||
if isinstance(features_value, list)
|
||||
else []
|
||||
)
|
||||
|
||||
grouped: Dict[tuple, List[Dict[str, Any]]] = {}
|
||||
grouped: dict[tuple[str, str], list[dict[str, JSONValue]]] = {}
|
||||
for entry in results:
|
||||
if not isinstance(entry, Mapping):
|
||||
if not isinstance(entry, dict):
|
||||
continue
|
||||
feature_id = entry.get("feature_id")
|
||||
provider = entry.get("provider")
|
||||
result = entry.get("result")
|
||||
if not feature_id or not provider or not isinstance(result, Mapping):
|
||||
if not isinstance(feature_id, str) or not feature_id:
|
||||
continue
|
||||
if not isinstance(provider, str) or not provider:
|
||||
continue
|
||||
if not isinstance(result, dict):
|
||||
continue
|
||||
if result.get("status") not in VALID_STATUSES:
|
||||
continue
|
||||
grouped.setdefault((feature_id, provider), []).append(dict(result))
|
||||
|
||||
features_out: List[Dict[str, Any]] = []
|
||||
features_out: list[JSONValue] = []
|
||||
for spec in feature_specs:
|
||||
feature_id = spec["id"]
|
||||
cells: Dict[str, Dict[str, Any]] = {}
|
||||
feature_id = str(spec["id"])
|
||||
cells: dict[str, JSONValue] = {}
|
||||
for provider in providers:
|
||||
cell_results = grouped.get((feature_id, provider), [])
|
||||
cells[provider] = _aggregate_cell(cell_results)
|
||||
|
|
@ -121,12 +143,12 @@ def build_matrix(
|
|||
"generated_at": generated_at,
|
||||
"litellm_version": litellm_version,
|
||||
"claude_code_version": claude_code_version,
|
||||
"providers": providers,
|
||||
"providers": list[JSONValue](providers),
|
||||
"features": features_out,
|
||||
}
|
||||
|
||||
|
||||
def _aggregate_cell(results: Sequence[Mapping[str, Any]]) -> Dict[str, Any]:
|
||||
def _aggregate_cell(results: Sequence[Mapping[str, JSONValue]]) -> dict[str, JSONValue]:
|
||||
"""Aggregate a list of per-model results into a single cell status.
|
||||
|
||||
Order of precedence (most informative wins):
|
||||
|
|
@ -182,7 +204,7 @@ def build_from_paths(
|
|||
claude_code_version: str,
|
||||
generated_at: str,
|
||||
output_path: Optional[Path] = None,
|
||||
) -> Dict[str, Any]:
|
||||
) -> dict[str, JSONValue]:
|
||||
"""I/O wrapper around build_matrix used by the publisher script."""
|
||||
manifest = load_manifest(manifest_path)
|
||||
results = load_results(results_path)
|
||||
|
|
|
|||
|
|
@ -22,12 +22,14 @@ Code version is logged in the CI output").
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import sys
|
||||
import urllib.request
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from http.client import HTTPResponse
|
||||
from typing import Callable, Mapping, Optional
|
||||
|
||||
from claude_code.json_types import JSON_OBJECT_ADAPTER, JSONValue
|
||||
|
||||
PACKAGE_NAME = "@anthropic-ai/claude-code"
|
||||
NPM_REGISTRY_URL = "https://registry.npmjs.org/{package}"
|
||||
DEFAULT_MIN_AGE = timedelta(days=3)
|
||||
|
|
@ -51,7 +53,13 @@ def _parse_npm_timestamp(value: str) -> datetime:
|
|||
return parsed
|
||||
|
||||
|
||||
def _default_fetcher(package_name: str) -> dict:
|
||||
def _urlopen(request: urllib.request.Request) -> HTTPResponse:
|
||||
return urllib.request.urlopen( # noqa: S310 — registry URL is constant # pyright: ignore[reportAny] # urlopen is typed as Any in typeshed; https URLs yield HTTPResponse
|
||||
request, timeout=DEFAULT_FETCH_TIMEOUT_SECONDS
|
||||
)
|
||||
|
||||
|
||||
def _default_fetcher(package_name: str) -> dict[str, JSONValue]:
|
||||
"""Fetch the npm packument for ``package_name`` over HTTPS.
|
||||
|
||||
Uses urllib (stdlib) so this module has no extra dependencies in the
|
||||
|
|
@ -61,17 +69,15 @@ def _default_fetcher(package_name: str) -> dict:
|
|||
# npm registry expects literally; do a minimal hand-roll instead.
|
||||
url = NPM_REGISTRY_URL.format(package=package_name.replace("/", "%2F"))
|
||||
req = urllib.request.Request(url, headers={"Accept": "application/json"})
|
||||
with urllib.request.urlopen( # noqa: S310 — registry URL is constant
|
||||
req, timeout=DEFAULT_FETCH_TIMEOUT_SECONDS
|
||||
) as response:
|
||||
with _urlopen(req) as response:
|
||||
body = response.read().decode("utf-8")
|
||||
return json.loads(body)
|
||||
return JSON_OBJECT_ADAPTER.validate_json(body)
|
||||
|
||||
|
||||
def resolve_pr_gate_version(
|
||||
*,
|
||||
metadata: Optional[Mapping] = None,
|
||||
fetcher: Optional[Callable[[str], Mapping]] = None,
|
||||
metadata: Optional[Mapping[str, JSONValue]] = None,
|
||||
fetcher: Optional[Callable[[str], Mapping[str, JSONValue]]] = None,
|
||||
as_of: Optional[datetime] = None,
|
||||
min_age: timedelta = DEFAULT_MIN_AGE,
|
||||
package_name: str = PACKAGE_NAME,
|
||||
|
|
@ -100,7 +106,8 @@ def resolve_pr_gate_version(
|
|||
fetch = fetcher or _default_fetcher
|
||||
metadata = fetch(package_name)
|
||||
|
||||
times = metadata.get("time") or {}
|
||||
times_value = metadata.get("time")
|
||||
times: dict[str, JSONValue] = times_value if isinstance(times_value, dict) else {}
|
||||
if as_of is None:
|
||||
as_of = datetime.now(timezone.utc)
|
||||
cutoff = as_of - min_age
|
||||
|
|
|
|||
|
|
@ -43,13 +43,16 @@ from __future__ import annotations
|
|||
|
||||
import contextlib
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import tempfile
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Dict, Iterator, Mapping, Optional
|
||||
from typing import Callable, Dict, Generator, Mapping, Optional
|
||||
|
||||
from pydantic import ValidationError
|
||||
|
||||
from claude_code.json_types import JSON_OBJECT_ADAPTER, JSONValue
|
||||
|
||||
# `fcntl` is POSIX-only; the suite is Linux/macOS only, so we don't
|
||||
# attempt a Windows fallback. Importing at module load fails fast on
|
||||
|
|
@ -165,6 +168,17 @@ def _state_dir(env: Optional[Mapping[str, str]] = None) -> Path:
|
|||
return Path(tempfile.gettempdir()) / DEFAULT_STATE_DIR_NAME
|
||||
|
||||
|
||||
def _state_float(value: JSONValue) -> float:
|
||||
"""Convert a JSON state value the way `float(...)` would, for the
|
||||
caller's corrupt-state handling: numeric and numeric-string values
|
||||
convert, everything else raises into the caller's except clause."""
|
||||
if isinstance(value, (int, float)):
|
||||
return float(value)
|
||||
if isinstance(value, str):
|
||||
return float(value)
|
||||
raise TypeError(f"cannot convert {type(value).__name__} to float")
|
||||
|
||||
|
||||
class RateLimiter:
|
||||
"""Cross-process token bucket per provider.
|
||||
|
||||
|
|
@ -192,8 +206,8 @@ class RateLimiter:
|
|||
self,
|
||||
config: Optional[Mapping[str, ProviderConfig]] = None,
|
||||
state_dir: Optional[Path] = None,
|
||||
clock: Optional[callable] = None,
|
||||
sleep: Optional[callable] = None,
|
||||
clock: Optional[Callable[[], float]] = None,
|
||||
sleep: Optional[Callable[[float], None]] = None,
|
||||
) -> None:
|
||||
self._config = dict(config) if config is not None else load_config()
|
||||
self._state_dir = Path(state_dir) if state_dir is not None else _state_dir()
|
||||
|
|
@ -269,7 +283,7 @@ class RateLimiter:
|
|||
os.close(fd)
|
||||
|
||||
@staticmethod
|
||||
def _read_state(fd: int, cfg: ProviderConfig, now: float) -> tuple:
|
||||
def _read_state(fd: int, cfg: ProviderConfig, now: float) -> tuple[float, float]:
|
||||
"""Read {tokens, last_refill} from `fd`, defaulting to a full
|
||||
bucket on a missing/empty/corrupt file.
|
||||
|
||||
|
|
@ -283,17 +297,17 @@ class RateLimiter:
|
|||
if not raw.strip():
|
||||
return cfg.burst, now
|
||||
try:
|
||||
obj = json.loads(raw)
|
||||
tokens = float(obj.get("tokens", cfg.burst))
|
||||
last_refill = float(obj.get("last_refill", now))
|
||||
obj = JSON_OBJECT_ADAPTER.validate_json(raw)
|
||||
tokens = _state_float(obj.get("tokens", cfg.burst))
|
||||
last_refill = _state_float(obj.get("last_refill", now))
|
||||
return tokens, last_refill
|
||||
except (ValueError, TypeError):
|
||||
except (ValidationError, ValueError, TypeError):
|
||||
return cfg.burst, now
|
||||
|
||||
@staticmethod
|
||||
def _refill(
|
||||
tokens: float, last_refill: float, now: float, cfg: ProviderConfig
|
||||
) -> tuple:
|
||||
) -> tuple[float, float]:
|
||||
"""Apply elapsed time to the bucket, capped at burst.
|
||||
|
||||
Negative elapsed (clock went backward, e.g. across host
|
||||
|
|
@ -341,7 +355,7 @@ def reset_default_limiter() -> None:
|
|||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def use_limiter(limiter: RateLimiter) -> Iterator[RateLimiter]:
|
||||
def use_limiter(limiter: RateLimiter) -> Generator[RateLimiter]:
|
||||
"""Temporarily install `limiter` as the process default.
|
||||
|
||||
The driver's `run_claude` calls `get_default_limiter()`; tests that
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue