refactor(e2e): type claude_code suite core modules under strict basedpyright

This commit is contained in:
mateo-berri 2026-07-15 16:50:32 -07:00
parent ff4a40f017
commit 949aa53096
10 changed files with 270 additions and 128 deletions

View file

@ -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,

View file

@ -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):

View file

@ -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:

View file

@ -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")

View file

@ -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"

View file

@ -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(

View 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)

View file

@ -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)

View file

@ -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

View file

@ -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