From 949aa53096e0c14431a69c1e1e62304bb9f219b4 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 15 Jul 2026 16:50:32 -0700 Subject: [PATCH] refactor(e2e): type claude_code suite core modules under strict basedpyright --- pyrightconfig.json | 2 +- tests/e2e/claude_code/_basic_messaging.py | 10 +- tests/e2e/claude_code/cli_driver.py | 104 +++++++++++------ tests/e2e/claude_code/conftest.py | 105 +++++++++++------- tests/e2e/claude_code/cron_vm/build_matrix.py | 17 ++- tests/e2e/claude_code/http_probe.py | 15 +-- tests/e2e/claude_code/json_types.py | 18 +++ tests/e2e/claude_code/matrix_builder.py | 66 +++++++---- .../claude_code/pr_gate_version_resolver.py | 25 +++-- tests/e2e/claude_code/rate_limiter.py | 36 ++++-- 10 files changed, 270 insertions(+), 128 deletions(-) create mode 100644 tests/e2e/claude_code/json_types.py diff --git a/pyrightconfig.json b/pyrightconfig.json index eabfbf515c4..97f099d5b2c 100644 --- a/pyrightconfig.json +++ b/pyrightconfig.json @@ -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, diff --git a/tests/e2e/claude_code/_basic_messaging.py b/tests/e2e/claude_code/_basic_messaging.py index f6b82a38f6a..eb172dfce1e 100644 --- a/tests/e2e/claude_code/_basic_messaging.py +++ b/tests/e2e/claude_code/_basic_messaging.py @@ -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): diff --git a/tests/e2e/claude_code/cli_driver.py b/tests/e2e/claude_code/cli_driver.py index 97eaa0e6847..d9f81cc34fa 100644 --- a/tests/e2e/claude_code/cli_driver.py +++ b/tests/e2e/claude_code/cli_driver.py @@ -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: diff --git a/tests/e2e/claude_code/conftest.py b/tests/e2e/claude_code/conftest.py index d2bfa1a54bf..ab1e2ffe98e 100644 --- a/tests/e2e/claude_code/conftest.py +++ b/tests/e2e/claude_code/conftest.py @@ -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//test_.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") diff --git a/tests/e2e/claude_code/cron_vm/build_matrix.py b/tests/e2e/claude_code/cron_vm/build_matrix.py index 128f041cced..158e1c90195 100644 --- a/tests/e2e/claude_code/cron_vm/build_matrix.py +++ b/tests/e2e/claude_code/cron_vm/build_matrix.py @@ -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" diff --git a/tests/e2e/claude_code/http_probe.py b/tests/e2e/claude_code/http_probe.py index d95307db3b9..b3cfa622b7a 100644 --- a/tests/e2e/claude_code/http_probe.py +++ b/tests/e2e/claude_code/http_probe.py @@ -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( diff --git a/tests/e2e/claude_code/json_types.py b/tests/e2e/claude_code/json_types.py new file mode 100644 index 00000000000..afd64de9404 --- /dev/null +++ b/tests/e2e/claude_code/json_types.py @@ -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) diff --git a/tests/e2e/claude_code/matrix_builder.py b/tests/e2e/claude_code/matrix_builder.py index 5641e488da2..99c3bf07166 100644 --- a/tests/e2e/claude_code/matrix_builder.py +++ b/tests/e2e/claude_code/matrix_builder.py @@ -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) diff --git a/tests/e2e/claude_code/pr_gate_version_resolver.py b/tests/e2e/claude_code/pr_gate_version_resolver.py index 82e12a2bf15..a2c039e0e04 100644 --- a/tests/e2e/claude_code/pr_gate_version_resolver.py +++ b/tests/e2e/claude_code/pr_gate_version_resolver.py @@ -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 diff --git a/tests/e2e/claude_code/rate_limiter.py b/tests/e2e/claude_code/rate_limiter.py index 06d21b83832..48ff1ee8069 100644 --- a/tests/e2e/claude_code/rate_limiter.py +++ b/tests/e2e/claude_code/rate_limiter.py @@ -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