diff --git a/src/ucode/agents/claude.py b/src/ucode/agents/claude.py index 375ea64d..b9a7dc4b 100644 --- a/src/ucode/agents/claude.py +++ b/src/ucode/agents/claude.py @@ -16,9 +16,6 @@ from typing import cast from ucode import gateway_proxy -from ucode.anthropic_model_discovery_proxy import ( - start_proxy as start_anthropic_model_discovery_proxy, -) from ucode.config_io import ( APP_DIR, ToolSpec, @@ -1311,51 +1308,6 @@ def _launch_relayed(state: dict, binary: str, tool_args: list[str]) -> None: raise SystemExit(returncode) -def _launch_claude_with_gateway_proxy( - state: dict, binary: str, tool_args: list[str], *, smart_routing: bool -) -> None: - """Launch Claude through the gateway model-alias proxy.""" - workspace = state["workspace"] - server, client = start_anthropic_model_discovery_proxy(workspace, 0) - os.environ["ANTHROPIC_BASE_URL"] = f"http://{LOOPBACK_HOST}:{server.server_address[1]}" - os.environ["CLAUDE_CODE_USE_GATEWAY"] = "1" - - server_thread = threading.Thread(target=server.serve_forever, daemon=True) - server_thread.start() - settings_override = {"env": {"ANTHROPIC_BASE_URL": os.environ["ANTHROPIC_BASE_URL"]}} - try: - if smart_routing: - - def compose_gateway_settings(args: list[str]) -> tuple[dict, list[str]]: - settings, remaining = _compose_v2_settings(args) - return _merge_claude_settings(settings, settings_override), remaining - - smart_routing_v2.launch_claude( - state, - tool_args, - binary=binary, - user_settings_path=CLAUDE_USER_SETTINGS_PATH, - launch_model=_original_launch_model(state), - compose_settings=compose_gateway_settings, - launch_model_args=_launch_model_args, - model_name=_maybe_add_1m_suffix, - ) - return - - proc = subprocess.Popen( - _build_claude_argv(binary, tool_args, settings_override=settings_override) - ) - try: - returncode = proc.wait() - except KeyboardInterrupt: - proc.send_signal(signal.SIGINT) - returncode = proc.wait() - finally: - server.shutdown() - client.close() - raise SystemExit(returncode) - - def launch(state: dict, tool_args: list[str]) -> None: binary = SPEC["binary"] workspace = state.get("workspace") @@ -1377,7 +1329,16 @@ def launch(state: dict, tool_args: list[str]) -> None: "Please use Codex or disable smart routing." ) if first_prompt_routing: - _launch_claude_with_gateway_proxy(state, binary, tool_args, smart_routing=True) + smart_routing_v2.launch_claude( + state, + tool_args, + binary=binary, + user_settings_path=CLAUDE_USER_SETTINGS_PATH, + launch_model=_original_launch_model(state), + compose_settings=_compose_v2_settings, + launch_model_args=_launch_model_args, + model_name=_maybe_add_1m_suffix, + ) return if ( workspace @@ -1387,8 +1348,6 @@ def launch(state: dict, tool_args: list[str]) -> None: # Discovery is launch-scoped. Pass it in the process environment rather # than persisting it in Claude's private or OS-managed settings. os.environ["CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY"] = "1" - _launch_claude_with_gateway_proxy(state, binary, tool_args, smart_routing=False) - return if workspace: os.environ["OAUTH_TOKEN"] = get_databricks_token(workspace, state.get("profile")) exec_or_spawn(_build_claude_argv(binary, tool_args)) diff --git a/src/ucode/anthropic_model_discovery_proxy.py b/src/ucode/anthropic_model_discovery_proxy.py deleted file mode 100644 index 61a66065..00000000 --- a/src/ucode/anthropic_model_discovery_proxy.py +++ /dev/null @@ -1,369 +0,0 @@ -"""Loopback proxy for Claude gateway model discovery. - -The proxy forwards Claude Code's apiKeyHelper credential, streams inference -responses verbatim, and rewrites model discovery responses when needed. - -Security invariants (mirroring `databricks.py` token handling): - - Binds 127.0.0.1 only; never exposed off-host. - - Never logs header values or bodies. -""" - -from __future__ import annotations - -import json -import threading -import time -import uuid -from collections.abc import Iterable -from http import HTTPStatus -from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer -from typing import cast -from urllib.parse import parse_qsl, urlencode, urlsplit, urlunsplit - -import httpx - -from ucode.constants import LOOPBACK_HOST -from ucode.databricks import _http_get_retry_delay -from ucode.gateway_proxy import ( - HOP_BY_HOP_HEADERS, - UPSTREAM_TIMEOUT, - log_proxy_diagnostic, -) - -# Claude Code abandons model discovery after roughly three seconds. One retry -# leaves enough time for the normal one-second backoff and the upstream request. -_ANTHROPIC_MODEL_DISCOVERY_MAX_RETRIES = 1 - - -class _ProxyHandler(BaseHTTPRequestHandler): - # Set by the server factory. - client: httpx.Client - - def log_message(self, format: str, *args: object) -> None: - return - - def _safe_send_error(self, code: int, message: str) -> None: - # The client (Claude Code) may already have disconnected, in which case - # reporting the error writes to a dead socket and raises again; swallow it. - try: - self.send_error(code, message) - except OSError: - pass - - def _transform_request(self, body: bytes | None) -> tuple[str, bytes | None]: - return self.path.lstrip("/"), body - - def _response_chunks(self, resp: httpx.Response) -> tuple[Iterable[bytes], frozenset[str]]: - return resp.iter_raw(), frozenset() - - def _should_retry_model_discovery(self, resp: httpx.Response) -> bool: - return ( - self.command == "GET" - and urlsplit(self.path).path == _ANTHROPIC_MODELS_PATH - and resp.status_code == HTTPStatus.TOO_MANY_REQUESTS - ) - - def _retry_model_discovery( - self, - url: str, - body: bytes | None, - diagnostic_id: str, - started: float, - retry_after: str | None, - ) -> None: - for retry_index in range(_ANTHROPIC_MODEL_DISCOVERY_MAX_RETRIES): - delay = _http_get_retry_delay(retry_after, retry_index) - log_proxy_diagnostic( - "model_discovery_retry_scheduled", - request_id=diagnostic_id, - attempt=retry_index + 2, - delay_ms=round(delay * 1000), - ) - time.sleep(delay) - headers = { - key: value - for key, value in self.headers.items() - if key.lower() not in HOP_BY_HOP_HEADERS - } - with self.client.stream(self.command, url, headers=headers, content=body) as resp: - log_proxy_diagnostic( - "model_discovery_upstream_headers", - request_id=diagnostic_id, - attempt=retry_index + 2, - status=resp.status_code, - elapsed_ms=round((time.monotonic() - started) * 1000), - ) - if ( - not self._should_retry_model_discovery(resp) - or retry_index == _ANTHROPIC_MODEL_DISCOVERY_MAX_RETRIES - 1 - ): - self._relay_response(resp, diagnostic_id=diagnostic_id, started=started) - return - retry_after = resp.headers.get("Retry-After") - resp.read() - - def _handle(self) -> None: - diagnostic_id = uuid.uuid4().hex[:12] - started = time.monotonic() - length = int(self.headers.get("Content-Length", 0) or 0) - body = self.rfile.read(length) if length else None - url, body = self._transform_request(body) - log_proxy_diagnostic( - "model_discovery_request_start", - request_id=diagnostic_id, - method=self.command, - path=self.path.split("?", 1)[0], - ) - try: - headers = { - key: value - for key, value in self.headers.items() - if key.lower() not in HOP_BY_HOP_HEADERS - } - with self.client.stream(self.command, url, headers=headers, content=body) as resp: - log_proxy_diagnostic( - "model_discovery_upstream_headers", - request_id=diagnostic_id, - attempt=1, - status=resp.status_code, - elapsed_ms=round((time.monotonic() - started) * 1000), - ) - if self._should_retry_model_discovery(resp): - retry_after = resp.headers.get("Retry-After") - resp.read() - self._retry_model_discovery( - url, - body, - diagnostic_id, - started, - retry_after, - ) - return - self._relay_response(resp, diagnostic_id=diagnostic_id, started=started) - except (BrokenPipeError, ConnectionResetError): - # Client closed before/while we relayed headers — routine on cancel. - log_proxy_diagnostic( - "model_discovery_client_disconnect", - request_id=diagnostic_id, - phase="request", - elapsed_ms=round((time.monotonic() - started) * 1000), - ) - return - except httpx.HTTPError as exc: - # Upstream failed before any bytes reached the client; a 502 is still - # sendable. (An HTTP *status* like 429 is not an error here — httpx - # only raises for transport failures — so real gateway errors are - # relayed verbatim by `_relay_response`.) - log_proxy_diagnostic( - "model_discovery_upstream_request_error", - request_id=diagnostic_id, - error_type=type(exc).__name__, - elapsed_ms=round((time.monotonic() - started) * 1000), - ) - self._safe_send_error(502, "gateway proxy upstream error") - - # Streaming passthrough: forward chunks as they arrive so SSE token streaming - # is not buffered (buffering would add full-response latency to first token). - # `iter_raw` preserves any Content-Encoding verbatim (we relay that header), - # so the proxy stays byte-transparent. - def _relay_response( - self, - resp: httpx.Response, - *, - diagnostic_id: str | None = None, - started: float | None = None, - ) -> None: - started = started if started is not None else time.monotonic() - chunks = 0 - bytes_relayed = 0 - first_byte_ms: int | None = None - try: - # The upstream request has completed through response headers before - # this hook selects raw streaming or a buffered response body. - response_chunks, dropped_headers = self._response_chunks(resp) - self.send_response(resp.status_code) - for key, value in resp.headers.items(): - header_name = key.lower() - if header_name not in HOP_BY_HOP_HEADERS and header_name not in dropped_headers: - self.send_header(key, value) - self.end_headers() - # Do not pass a fixed chunk size here. httpx accumulates bytes until - # that size is reached, which can hide small SSE heartbeat frames - # from Claude Code for minutes during a slow artifact/tool call. - # With ``chunk_size=None`` (the default), raw upstream chunks are - # yielded as they arrive and pings keep the downstream connection - # alive even before the model produces a large content block. - for chunk in response_chunks: - if chunk: - if first_byte_ms is None: - first_byte_ms = round((time.monotonic() - started) * 1000) - self.wfile.write(chunk) - self.wfile.flush() - chunks += 1 - bytes_relayed += len(chunk) - log_proxy_diagnostic( - "model_discovery_response_complete", - request_id=diagnostic_id, - status=resp.status_code, - chunks=chunks, - bytes=bytes_relayed, - first_byte_ms=first_byte_ms, - elapsed_ms=round((time.monotonic() - started) * 1000), - ) - except (BrokenPipeError, ConnectionResetError): - # Client (Claude Code) closed the connection mid-response — routine on - # cancelled turns / SSE teardown. Nothing left to relay to, so stop - # quietly rather than crashing the handler thread. - log_proxy_diagnostic( - "model_discovery_client_disconnect", - request_id=diagnostic_id, - phase="response", - chunks=chunks, - bytes=bytes_relayed, - elapsed_ms=round((time.monotonic() - started) * 1000), - ) - return - except httpx.HTTPError as exc: - # Upstream dropped mid-stream. Headers (and status) may already be - # sent, so we can't reliably signal a fresh error — stop and let the - # client see a truncated stream rather than corrupt the framing. - log_proxy_diagnostic( - "model_discovery_upstream_stream_error", - request_id=diagnostic_id, - error_type=type(exc).__name__, - status=resp.status_code, - chunks=chunks, - bytes=bytes_relayed, - elapsed_ms=round((time.monotonic() - started) * 1000), - ) - return - - # Forward every method: this is a transparent pass-through, so routing any - # `do_` lookup to `_handle` lets the gateway reject unsupported methods. - def __getattr__(self, name: str): - if name.startswith("do_"): - return self._handle - raise AttributeError(name) - - -_MODEL_ALIAS_PREFIX = "anthropic-aigw-" -_ANTHROPIC_MODELS_PATH = "/v1/models" -_ANTHROPIC_MESSAGES_PATH = "/v1/messages" - - -class _AnthropicModelAliases: - """Maps Claude-compatible discovery IDs back to their gateway model IDs.""" - - def __init__(self) -> None: - self._original_by_alias: dict[str, str] = {} - self._lock = threading.Lock() - - def prefix_model_ids(self, body: bytes) -> bytes: - try: - payload = json.loads(body) - models = payload["data"] - if not isinstance(models, list): - return body - except (UnicodeDecodeError, json.JSONDecodeError, KeyError, TypeError): - return body - - aliases: dict[str, str] = {} - for model in models: - if not isinstance(model, dict) or not isinstance(model.get("id"), str): - continue - model_id = model["id"] - lowered = model_id.lower() - if "claude" in lowered or "anthropic" in lowered: - continue - alias = f"{_MODEL_ALIAS_PREFIX}{model_id}" - model["id"] = alias - aliases[alias] = model_id - - with self._lock: - self._original_by_alias.update(aliases) - - for cursor in ("first_id", "last_id"): - model_id = payload.get(cursor) - alias = f"{_MODEL_ALIAS_PREFIX}{model_id}" - if alias in aliases: - payload[cursor] = alias - return json.dumps(payload, separators=(",", ":")).encode() - - def original_id(self, model_id: str) -> str: - with self._lock: - return self._original_by_alias.get(model_id, model_id) - - def rewrite_path(self, path: str) -> str: - parsed = urlsplit(path) - if parsed.path != _ANTHROPIC_MODELS_PATH: - return path - query = [ - (key, self.original_id(value) if key == "after_id" else value) - for key, value in parse_qsl(parsed.query, keep_blank_values=True) - ] - return urlunsplit( - (parsed.scheme, parsed.netloc, parsed.path, urlencode(query), parsed.fragment) - ) - - def rewrite_body(self, path: str, body: bytes | None) -> bytes | None: - if urlsplit(path).path != _ANTHROPIC_MESSAGES_PATH or body is None: - return body - try: - payload = json.loads(body) - model_id = payload.get("model") - if not isinstance(model_id, str): - return body - except (UnicodeDecodeError, json.JSONDecodeError, AttributeError): - return body - original_id = self.original_id(model_id) - if original_id == model_id: - return body - payload["model"] = original_id - return json.dumps(payload, separators=(",", ":")).encode() - - -class _AnthropicModelDiscoveryHandler(_ProxyHandler): - anthropic_model_aliases: _AnthropicModelAliases - - def _transform_request(self, body: bytes | None) -> tuple[str, bytes | None]: - body = self.anthropic_model_aliases.rewrite_body(self.path, body) - url = self.anthropic_model_aliases.rewrite_path(self.path).lstrip("/") - return url, body - - def _response_chunks(self, resp: httpx.Response) -> tuple[Iterable[bytes], frozenset[str]]: - should_prefix_model_ids = ( - self.command == "GET" - and urlsplit(self.path).path == _ANTHROPIC_MODELS_PATH - and HTTPStatus.OK <= resp.status_code < HTTPStatus.MULTIPLE_CHOICES - ) - if not should_prefix_model_ids: - return super()._response_chunks(resp) - body = self.anthropic_model_aliases.prefix_model_ids(resp.read()) - # resp.read() decodes compression; rewritten JSON is uncompressed. - return (body,), frozenset({"content-encoding"}) - - -def start_proxy( - workspace: str, - port: int, -) -> tuple[ThreadingHTTPServer, httpx.Client]: - """Start the Anthropic model discovery proxy.""" - upstream_base = f"{workspace.rstrip('/')}/ai-gateway/anthropic/" - client = httpx.Client(base_url=upstream_base, timeout=UPSTREAM_TIMEOUT, follow_redirects=False) - handler = cast( - type[BaseHTTPRequestHandler], - type( - "BoundProxyHandler", - (_AnthropicModelDiscoveryHandler,), - { - "client": client, - "anthropic_model_aliases": _AnthropicModelAliases(), - }, - ), - ) - try: - server = ThreadingHTTPServer((LOOPBACK_HOST, port), handler) - except OSError: - server = ThreadingHTTPServer((LOOPBACK_HOST, 0), handler) - - return server, client diff --git a/tests/test_agent_claude.py b/tests/test_agent_claude.py index 8563687a..1a7b8fbc 100644 --- a/tests/test_agent_claude.py +++ b/tests/test_agent_claude.py @@ -922,7 +922,7 @@ def boom(name, entry, scope=mcp_mod.MCP_USER_SCOPE): class TestClaudeLaunch: - def test_relayed_launch_uses_refresh_proxy_not_discovery_proxy(self, monkeypatch): + def test_relayed_launch_uses_refresh_proxy(self, monkeypatch): calls: list[tuple] = [] class Server: @@ -965,11 +965,6 @@ def start_proxy(workspace, profile, port, token_header, force_refresh_near_expir monkeypatch.setattr(claude, "_managed_relayed_conflicts", lambda: None) monkeypatch.setattr(claude, "_ensure_subscription_login", lambda: None) monkeypatch.setattr(claude.gateway_proxy, "start_proxy", start_proxy) - monkeypatch.setattr( - claude, - "start_anthropic_model_discovery_proxy", - lambda *_args: pytest.fail("relayed auth must not use the discovery proxy"), - ) monkeypatch.setattr(claude.subprocess, "Popen", Process) with pytest.raises(SystemExit) as exc: @@ -1093,12 +1088,20 @@ def test_v2_noninteractive_launch_bypasses_first_prompt_routing(self, monkeypatc def test_v2_positional_prompt_uses_first_prompt_routing(self, monkeypatch, tool_args): monkeypatch.setenv(v2.ENV_VAR, "1") launch_v2 = Mock() - monkeypatch.setattr(claude, "_launch_claude_with_gateway_proxy", launch_v2) + monkeypatch.setattr(claude, "_original_launch_model", lambda _state: None) + monkeypatch.setattr(v2, "launch_claude", launch_v2) claude.launch({"workspace": WS}, tool_args) launch_v2.assert_called_once_with( - {"workspace": WS}, "claude", tool_args, smart_routing=True + {"workspace": WS}, + tool_args, + binary="claude", + user_settings_path=claude.CLAUDE_USER_SETTINGS_PATH, + launch_model=None, + compose_settings=claude._compose_v2_settings, + launch_model_args=claude._launch_model_args, + model_name=claude._maybe_add_1m_suffix, ) def test_v2_does_not_treat_option_value_as_positional_argument(self): @@ -1107,123 +1110,19 @@ def test_v2_does_not_treat_option_value_as_positional_argument(self): def test_v2_treats_optional_option_value_as_interactive(self): assert claude._uses_interactive_tui(["--resume", "session-id"]) is True - def test_gateway_discovery_uses_anthropic_proxy(self, monkeypatch): - calls: list[tuple] = [] - + def test_gateway_discovery_uses_direct_gateway(self, monkeypatch): + calls: list[list[str]] = [] monkeypatch.delenv(v2.ENV_VAR, raising=False) - - class Server: - server_address = ("127.0.0.1", 12345) - - def serve_forever(self): - calls.append(("serve",)) - - def shutdown(self): - calls.append(("shutdown",)) - - class Client: - def close(self): - calls.append(("close",)) - - class Process: - def __init__(self, argv): - calls.append(("popen", argv)) - - def wait(self): - return 0 - - def start_proxy(workspace, port): - calls.append(("proxy", workspace, port)) - return Server(), Client() - monkeypatch.setenv(claude.GATEWAY_MODEL_DISCOVERY_ENV_VAR, "1") monkeypatch.delenv("OAUTH_TOKEN", raising=False) - monkeypatch.delenv("ANTHROPIC_AUTH_TOKEN", raising=False) - monkeypatch.delenv("ANTHROPIC_BASE_URL", raising=False) - monkeypatch.delenv("CLAUDE_CODE_USE_GATEWAY", raising=False) - monkeypatch.setattr( - claude, - "get_databricks_token", - lambda *_args: pytest.fail("model discovery must rely on apiKeyHelper"), - ) - monkeypatch.setattr( - claude, - "start_anthropic_model_discovery_proxy", - start_proxy, - ) - monkeypatch.setattr(claude.subprocess, "Popen", Process) + monkeypatch.setattr(claude, "get_databricks_token", lambda *_args: "token") + monkeypatch.setattr(claude, "exec_or_spawn", lambda argv: calls.append(argv)) - with pytest.raises(SystemExit) as exc: - claude.launch({"workspace": WS, "profile": "test"}, ["--debug"]) + claude.launch({"workspace": WS, "profile": "test"}, ["--debug"]) - assert exc.value.code == 0 - assert "OAUTH_TOKEN" not in os.environ - assert "ANTHROPIC_AUTH_TOKEN" not in os.environ - assert os.environ["ANTHROPIC_BASE_URL"] == "http://127.0.0.1:12345" - assert os.environ["CLAUDE_CODE_USE_GATEWAY"] == "1" + assert os.environ["OAUTH_TOKEN"] == "token" assert os.environ["CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY"] == "1" - assert calls[:2] == [ - ("proxy", WS, 0), - ("serve",), - ] - assert calls[2][0] == "popen" - argv = calls[2][1] - assert argv[:2] == ["claude", "--settings"] - assert json.loads(argv[2])["env"]["ANTHROPIC_BASE_URL"] == "http://127.0.0.1:12345" - assert "CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY" not in json.loads(argv[2])["env"] - assert argv[3:] == ["--debug"] - assert calls[3:] == [("shutdown",), ("close",)] - - def test_smart_routing_uses_anthropic_proxy(self, monkeypatch): - calls: list[tuple] = [] - captured: dict = {} - - class Server: - server_address = ("127.0.0.1", 12345) - - def serve_forever(self): - calls.append(("serve",)) - - def shutdown(self): - calls.append(("shutdown",)) - - class Client: - def close(self): - calls.append(("close",)) - - def start_proxy(workspace, port): - calls.append(("proxy", workspace, port)) - return Server(), Client() - - def launch_v2(state, tool_args, **kwargs): - captured["settings"] = kwargs["compose_settings"](["--debug"]) - raise SystemExit(0) - - monkeypatch.setenv(v2.ENV_VAR, "1") - monkeypatch.delenv("OAUTH_TOKEN", raising=False) - monkeypatch.delenv("ANTHROPIC_AUTH_TOKEN", raising=False) - monkeypatch.setattr(claude, "start_anthropic_model_discovery_proxy", start_proxy) - monkeypatch.setattr( - claude, - "_compose_v2_settings", - lambda args: ({"env": {"ANTHROPIC_BASE_URL": "https://direct"}}, args), - ) - monkeypatch.setattr(v2, "launch_claude", launch_v2) - - with pytest.raises(SystemExit) as exc: - claude.launch({"workspace": WS, "profile": "test"}, ["--debug"]) - - assert exc.value.code == 0 - assert "OAUTH_TOKEN" not in os.environ - assert "ANTHROPIC_AUTH_TOKEN" not in os.environ - assert calls[:2] == [ - ("proxy", WS, 0), - ("serve",), - ] - settings, remaining = captured["settings"] - assert settings["env"]["ANTHROPIC_BASE_URL"] == "http://127.0.0.1:12345" - assert remaining == ["--debug"] - assert calls[2:] == [("shutdown",), ("close",)] + assert calls == [["claude", "--settings", str(claude.CLAUDE_SETTINGS_PATH), "--debug"]] class TestWriteToolConfigPrunesStaleModelEnv: diff --git a/tests/test_anthropic_model_discovery_proxy.py b/tests/test_anthropic_model_discovery_proxy.py deleted file mode 100644 index 1d8996e1..00000000 --- a/tests/test_anthropic_model_discovery_proxy.py +++ /dev/null @@ -1,288 +0,0 @@ -"""Tests for Anthropic model discovery transformations.""" - -from __future__ import annotations - -import io -import json - -from ucode import anthropic_model_discovery_proxy - - -class _FakeResponse: - def __init__(self, status_code: int, headers: dict[str, str], body: bytes): - self.status_code = status_code - self.headers = headers - self._body = body - self.read_calls = 0 - self.iter_raw_calls = 0 - - def read(self): - self.read_calls += 1 - return self._body - - def iter_raw(self): - self.iter_raw_calls += 1 - yield self._body - - def __enter__(self): - return self - - def __exit__(self, *_args): - return False - - -class _FakeClient: - def __init__(self, response): - self.responses = list(response) if isinstance(response, list) else [response] - self.request = None - self.requests = [] - - def stream(self, method, url, headers, content): - self.request = (method, url, headers, content) - self.requests.append(self.request) - return self.responses.pop(0) - - -class _Collect(io.RawIOBase): - def __init__(self): - self.data = bytearray() - - def write(self, body): # type: ignore[override] - self.data += bytes(body) - return len(body) - - def flush(self): - return None - - -def _handler(wfile, path="/v1/models", command="GET"): - handler = object.__new__(anthropic_model_discovery_proxy._AnthropicModelDiscoveryHandler) - handler.wfile = wfile - handler.request_version = "HTTP/1.1" - handler.requestline = f"{command} {path} HTTP/1.1" - handler.command = command - handler.path = path - handler._headers_buffer = [] - handler.anthropic_model_aliases = anthropic_model_discovery_proxy._AnthropicModelAliases() - return handler - - -class TestAnthropicModelAliases: - def test_prefixes_custom_model_ids_without_changing_display_name(self): - aliases = anthropic_model_discovery_proxy._AnthropicModelAliases() - body = json.dumps( - { - "data": [ - {"id": "catalog.schema.custom", "display_name": "Custom model"}, - {"id": "system.ai.claude-sonnet"}, - {"id": "claude-sonnet"}, - {"id": "catalog.schema.anthropic-provider"}, - {"id": "anthropic-provider"}, - ], - "first_id": "catalog.schema.custom", - "last_id": "catalog.schema.anthropic-provider", - } - ).encode() - - payload = json.loads(aliases.prefix_model_ids(body)) - - assert payload == { - "data": [ - { - "id": "anthropic-aigw-catalog.schema.custom", - "display_name": "Custom model", - }, - {"id": "system.ai.claude-sonnet"}, - {"id": "claude-sonnet"}, - {"id": "catalog.schema.anthropic-provider"}, - {"id": "anthropic-provider"}, - ], - "first_id": "anthropic-aigw-catalog.schema.custom", - "last_id": "catalog.schema.anthropic-provider", - } - - def test_rewrites_known_alias_in_messages_body(self): - aliases = anthropic_model_discovery_proxy._AnthropicModelAliases() - aliases.prefix_model_ids(b'{"data":[{"id":"catalog.schema.custom"}]}') - - body = aliases.rewrite_body( - "/v1/messages", - b'{"model":"anthropic-aigw-catalog.schema.custom","messages":[]}', - ) - - assert json.loads(body) == {"model": "catalog.schema.custom", "messages": []} - - def test_rewrites_known_alias_in_pagination_cursor(self): - aliases = anthropic_model_discovery_proxy._AnthropicModelAliases() - aliases.prefix_model_ids(b'{"data":[{"id":"catalog.schema.custom"}]}') - - assert ( - aliases.rewrite_path( - "/v1/models?limit=1000&after_id=anthropic-aigw-catalog.schema.custom" - ) - == "/v1/models?limit=1000&after_id=catalog.schema.custom" - ) - - def test_ignores_non_anthropic_models_path(self): - aliases = anthropic_model_discovery_proxy._AnthropicModelAliases() - path = "/codex/v1/models?after_id=anthropic-aigw-catalog.schema.custom" - - assert aliases.rewrite_path(path) == path - - def test_does_not_strip_unknown_prefixed_id(self): - aliases = anthropic_model_discovery_proxy._AnthropicModelAliases() - unknown = "anthropic-aigw-legitimate-upstream-id" - - assert aliases.rewrite_path(f"/v1/models?after_id={unknown}") == ( - f"/v1/models?after_id={unknown}" - ) - assert ( - aliases.rewrite_body("/v1/messages", json.dumps({"model": unknown}).encode()) - == json.dumps({"model": unknown}).encode() - ) - - def test_leaves_malformed_discovery_response_unchanged(self): - aliases = anthropic_model_discovery_proxy._AnthropicModelAliases() - assert aliases.prefix_model_ids(b"not-json") == b"not-json" - - -class TestAnthropicModelDiscoveryHandler: - def test_forwards_api_key_helper_credential_without_refresh_auth(self): - out = _Collect() - handler = _handler(out) - handler.headers = {"X-Api-Key": "api-key-helper-token"} - handler.rfile = io.BytesIO() - handler.client = _FakeClient(_FakeResponse(200, {}, b'{"data":[]}')) - - handler._handle() - - _method, _url, headers, _body = handler.client.request - assert headers["X-Api-Key"] == "api-key-helper-token" - assert "Authorization" not in headers - assert "X-Databricks-AI-Gateway-Token" not in headers - - def test_prefixes_models_and_strips_hop_by_hop_headers(self): - out = _Collect() - handler = _handler(out) - handler.headers = { - "X-Api-Key": "api-key-helper-token", - "Connection": "keep-alive", - "Host": "127.0.0.1", - } - handler.rfile = io.BytesIO() - handler.client = _FakeClient(_FakeResponse(200, {}, b'{"data":[{"id":"custom-model"}]}')) - - handler._handle() - - method, url, headers, body = handler.client.request - assert (method, url, body) == ("GET", "v1/models", None) - assert headers == {"X-Api-Key": "api-key-helper-token"} - assert b"anthropic-aigw-custom-model" in bytes(out.data) - - def test_retries_rate_limited_model_discovery(self, monkeypatch): - out = _Collect() - handler = _handler(out) - handler.headers = {"X-Api-Key": "api-key-helper-token"} - handler.rfile = io.BytesIO() - rate_limited = _FakeResponse(429, {"Retry-After": "0"}, b"rate limited") - success = _FakeResponse(200, {}, b'{"data":[{"id":"custom-model"}]}') - handler.client = _FakeClient([rate_limited, success]) - monkeypatch.setattr(anthropic_model_discovery_proxy.time, "sleep", lambda _delay: None) - - handler._handle() - - assert len(handler.client.requests) == 2 - assert all( - request[2] == {"X-Api-Key": "api-key-helper-token"} - for request in handler.client.requests - ) - assert rate_limited.read_calls == 1 - assert b"429 Too Many Requests" not in bytes(out.data) - assert b"anthropic-aigw-custom-model" in bytes(out.data) - - def test_relays_rate_limit_after_model_discovery_retries_are_exhausted(self, monkeypatch): - out = _Collect() - handler = _handler(out) - handler.headers = {"X-Api-Key": "api-key-helper-token"} - handler.rfile = io.BytesIO() - responses = [_FakeResponse(429, {"Retry-After": "0"}, b"rate limited") for _ in range(2)] - handler.client = _FakeClient(responses) - monkeypatch.setattr(anthropic_model_discovery_proxy.time, "sleep", lambda _delay: None) - - handler._handle() - - assert len(handler.client.requests) == 2 - assert responses[0].read_calls == 1 - assert responses[1].read_calls == 0 - assert b"429 Too Many Requests" in bytes(out.data) - assert b"rate limited" in bytes(out.data) - - def test_prefixes_successful_model_response_and_drops_content_encoding(self): - out = _Collect() - handler = _handler(out) - response = _FakeResponse( - 200, - {"Content-Encoding": "gzip"}, - b'{"data":[{"id":"custom-model"}]}', - ) - - handler._relay_response(response) - - assert b"Content-Encoding" not in bytes(out.data) - assert b"anthropic-aigw-custom-model" in bytes(out.data) - - def test_keeps_content_encoding_for_unchanged_error(self): - out = _Collect() - handler = _handler(out) - response = _FakeResponse(400, {"Content-Encoding": "gzip"}, b"compressed-error") - - handler._relay_response(response) - - assert b"Content-Encoding: gzip" in bytes(out.data) - assert b"compressed-error" in bytes(out.data) - assert response.read_calls == 0 - assert response.iter_raw_calls == 1 - - def test_streams_inference_response_without_buffering(self): - out = _Collect() - handler = _handler(out, path="/v1/messages", command="POST") - handler.headers = {"X-Api-Key": "api-key-helper-token", "Content-Length": "2"} - handler.rfile = io.BytesIO(b"{}") - response = _FakeResponse(200, {"Content-Type": "text/event-stream"}, b"data: event\n\n") - handler.client = _FakeClient(response) - - handler._handle() - - _method, _url, headers, _body = handler.client.request - assert headers == {"X-Api-Key": "api-key-helper-token"} - assert response.read_calls == 0 - assert response.iter_raw_calls == 1 - assert b"data: event\n\n" in bytes(out.data) - - def test_strips_known_alias_from_message_request(self): - handler = _handler(_Collect(), path="/v1/messages", command="POST") - handler.anthropic_model_aliases.prefix_model_ids( - b'{"data":[{"id":"catalog.schema.custom"}]}' - ) - - url, body = handler._transform_request(b'{"model":"anthropic-aigw-catalog.schema.custom"}') - - assert url == "v1/messages" - assert json.loads(body) == {"model": "catalog.schema.custom"} - - -def test_start_proxy_uses_discovery_handler(): - server, client = anthropic_model_discovery_proxy.start_proxy("https://workspace.example.com", 0) - try: - handler = server.RequestHandlerClass - assert issubclass( - handler, - anthropic_model_discovery_proxy._AnthropicModelDiscoveryHandler, - ) - assert isinstance( - handler.anthropic_model_aliases, - anthropic_model_discovery_proxy._AnthropicModelAliases, - ) - finally: - server.server_close() - client.close()