diff --git a/CHANGELOG.md b/CHANGELOG.md index 92cc8de4a..cdfb09ab5 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -32,6 +32,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 * **code:** keep Python `from __future__` imports before executable code during AST compression and validate compressed Python with `compile(..., "exec")` so compile-time syntax rules are enforced ([#1233](https://github.com/chopratejas/headroom/issues/1233)). * **proxy:** report real input tokens on the streaming `message_start` event for LiteLLM/Bedrock-backed requests. LiteLLM streaming never surfaces prompt tokens mid-stream, so `message_start.usage.input_tokens` was always `0`; Anthropic clients (e.g. Claude Code) read input-token metrics from that event, underreporting token usage by ~99% in OTel/CloudWatch dashboards. The Bedrock streamer now backfills `input_tokens` with the count Headroom actually sent upstream when the backend leaves it unset, preserving any non-zero value the backend genuinely reports ([#1132](https://github.com/chopratejas/headroom/issues/1132)). +* **proxy:** give buffered Anthropic request paths their own longer read timeout, so long `/v1/messages` turns and Anthropic batch or passthrough reads no longer trip the generic proxy cap while unrelated request timeouts stay unchanged. * **proxy:** force Responses API `store=true` when Headroom injects memory tools so `previous_response_id` continuations work after memory tool calls from clients that requested `store=false` ([#1103](https://github.com/chopratejas/headroom/pull/1103)). * **proxy:** build SSL contexts for custom CA bundles so enterprise/private PKI roots work with Python/OpenSSL strict verification. * **tokenizers:** bound token-counting of oversized tool-content blobs instead of running `count_text` over the whole serialized string. `count_messages` runs on the proxy request path; serializing is cheap, but `count_text` over a multi-megabyte `tool_result` / `tool_use` string took seconds and could freeze `/health` and in-flight requests. For payloads over ~50KB serialized, `count_text` now runs on an even-spread sample of the string and scales by length; it stays model-accurate, bounded for any blob shape, and biased to under-count. Smaller payloads stay exact. diff --git a/headroom/cli/proxy.py b/headroom/cli/proxy.py index 1bb88adaf..1a8f73aa0 100644 --- a/headroom/cli/proxy.py +++ b/headroom/cli/proxy.py @@ -309,6 +309,17 @@ def dashboard(port: int, no_open: bool) -> None: "Env: HEADROOM_CONNECT_TIMEOUT_SECONDS." ), ) +@click.option( + "--anthropic-buffered-request-timeout-seconds", + type=click.IntRange(min=1), + default=None, + envvar="HEADROOM_ANTHROPIC_BUFFERED_REQUEST_TIMEOUT_SECONDS", + help=( + "Buffered Anthropic read timeout in seconds for non-streaming " + "message and batch paths (default: 600). " + "Env: HEADROOM_ANTHROPIC_BUFFERED_REQUEST_TIMEOUT_SECONDS." + ), +) @click.option( "--anthropic-pre-upstream-concurrency", type=int, @@ -765,6 +776,7 @@ def proxy( retry_max_attempts: int | None, request_timeout_seconds: int | None, connect_timeout_seconds: int | None, + anthropic_buffered_request_timeout_seconds: int | None, anthropic_pre_upstream_concurrency: int | None, anthropic_pre_upstream_acquire_timeout_seconds: float | None, anthropic_pre_upstream_memory_context_timeout_seconds: float | None, @@ -989,6 +1001,11 @@ def proxy( connect_timeout_seconds=connect_timeout_seconds if connect_timeout_seconds is not None else 10, + anthropic_buffered_request_timeout_seconds=( + anthropic_buffered_request_timeout_seconds + if anthropic_buffered_request_timeout_seconds is not None + else 600 + ), max_connections=max_connections, max_keepalive_connections=max_keepalive_connections, keepalive_expiry=keepalive_expiry, diff --git a/headroom/proxy/handlers/anthropic.py b/headroom/proxy/handlers/anthropic.py index fd9297162..b431b5567 100644 --- a/headroom/proxy/handlers/anthropic.py +++ b/headroom/proxy/handlers/anthropic.py @@ -130,6 +130,15 @@ class AnthropicHandlerMixin: int(cache_creation.get("ephemeral_1h_input_tokens", 0) or 0), ) + def _anthropic_buffered_request_timeout(self) -> httpx.Timeout: + """Timeout for buffered Anthropic reads.""" + return httpx.Timeout( + connect=self.config.connect_timeout_seconds, + read=self.config.anthropic_buffered_request_timeout_seconds, + write=self.config.request_timeout_seconds, + pool=self.config.connect_timeout_seconds, + ) + @classmethod def _sort_tools_deterministically( cls, tools: list[dict[str, Any]] | None @@ -2144,6 +2153,7 @@ class AnthropicHandlerMixin: request_id=request_id, forwarder_name="anthropic_messages", path_for_log="/v1/messages", + timeout=self._anthropic_buffered_request_timeout(), ) self.pipeline_extensions.emit( PipelineStage.POST_SEND, @@ -2346,7 +2356,7 @@ class AnthropicHandlerMixin: url, content=ccr_outbound_bytes, headers=ccr_outbound_headers, - timeout=httpx.Timeout(120.0), # Override timeout for CCR + timeout=self._anthropic_buffered_request_timeout(), ) logger.info( f"CCR: Got response status={cont_response.status_code}, " @@ -2448,7 +2458,11 @@ class AnthropicHandlerMixin: continuation_body["tools"] = tools cont_response = await self._retry_request( - "POST", url, headers, continuation_body + "POST", + url, + headers, + continuation_body, + timeout=self._anthropic_buffered_request_timeout(), ) # Update response with continuation @@ -2922,6 +2936,7 @@ class AnthropicHandlerMixin: request_id=request_id, forwarder_name="anthropic_batch", path_for_log="/v1/messages/batches", + timeout=self._anthropic_buffered_request_timeout(), ) # Batch create: tokens accumulated across all requests in @@ -3047,6 +3062,7 @@ class AnthropicHandlerMixin: url=url, headers=headers, content=body, + timeout=self._anthropic_buffered_request_timeout(), ) # Batch passthrough: no compression, no transforms — but we @@ -3163,7 +3179,11 @@ class AnthropicHandlerMixin: request_id=None, ) - response = await self.http_client.get(url, headers=headers) # type: ignore[union-attr] + response = await self.http_client.get( # type: ignore[union-attr] + url, + headers=headers, + timeout=self._anthropic_buffered_request_timeout(), + ) if response.status_code != 200: # Error - pass through diff --git a/headroom/proxy/models.py b/headroom/proxy/models.py index e39e1d91e..bc7a37dbc 100644 --- a/headroom/proxy/models.py +++ b/headroom/proxy/models.py @@ -262,6 +262,9 @@ class ProxyConfig: # Timeouts request_timeout_seconds: int = 300 connect_timeout_seconds: int = 10 + # Anthropic buffered reads can legitimately run longer than the generic + # proxy request cap. Keep the generic timeout unchanged elsewhere. + anthropic_buffered_request_timeout_seconds: int = 600 # Connection pool max_connections: int = 500 diff --git a/headroom/proxy/server.py b/headroom/proxy/server.py index 26d79b869..6b9ef06ab 100644 --- a/headroom/proxy/server.py +++ b/headroom/proxy/server.py @@ -1581,6 +1581,7 @@ class HeadroomProxy( request_id: str | None = None, forwarder_name: str = "server", path_for_log: str | None = None, + timeout: httpx.Timeout | float | None = None, ) -> httpx.Response: """Make request with retry and exponential backoff. @@ -1624,16 +1625,20 @@ class HeadroomProxy( source=source, ) + post_kwargs: dict = {"content": outbound_bytes, "headers": outbound_headers} + if timeout is not None: + post_kwargs["timeout"] = timeout + for attempt in range(self.config.retry_max_attempts): try: if stream: # For streaming, we return early - retry happens at higher level return await self.http_client.post( # type: ignore[union-attr] - url, content=outbound_bytes, headers=outbound_headers + url, **post_kwargs ) else: response = await self.http_client.post( # type: ignore[union-attr] - url, content=outbound_bytes, headers=outbound_headers + url, **post_kwargs ) # Don't retry client errors (4xx) @@ -3818,6 +3823,11 @@ def _proxy_config_from_env() -> ProxyConfig: port=_get_env_int("HEADROOM_PORT", 8787), openai_api_url=os.environ.get("OPENAI_TARGET_API_URL"), anthropic_api_url=os.environ.get("ANTHROPIC_TARGET_API_URL"), + anthropic_buffered_request_timeout_seconds=_get_env_int( + "HEADROOM_ANTHROPIC_BUFFERED_REQUEST_TIMEOUT_SECONDS", + 600, + min_value=1, + ), vertex_api_url=os.environ.get("VERTEX_TARGET_API_URL"), backend=_get_env_str("HEADROOM_BACKEND", "anthropic"), bedrock_region=_get_env_str("HEADROOM_BEDROCK_REGION", "us-west-2"), @@ -4023,15 +4033,25 @@ def _get_env_optional_bool(name: str) -> bool | None: return val.lower() in ("true", "1", "yes", "on") -def _get_env_int(name: str, default: int) -> int: +def _get_env_int(name: str, default: int, *, min_value: int | None = None) -> int: """Get integer from environment variable.""" val = os.environ.get(name) if val is None: return default try: - return int(val) + parsed = int(val) except ValueError: return default + if min_value is not None and parsed < min_value: + return default + return parsed + + +def _positive_int_arg(value: str) -> int: + parsed = int(value) + if parsed < 1: + raise argparse.ArgumentTypeError("must be >= 1") + return parsed def _get_env_float(name: str, default: float) -> float: @@ -4122,6 +4142,15 @@ if __name__ == "__main__": "--anthropic-api-url", help=f"Custom Anthropic API URL (default: {DEFAULT_ANTHROPIC_API_URL})", ) + parser.add_argument( + "--anthropic-buffered-request-timeout-seconds", + type=_positive_int_arg, + default=600, + help=( + "Anthropic buffered read timeout in seconds for non-streaming " + "message and batch paths (default: 600)" + ), + ) parser.add_argument( "--vertex-api-url", help=f"Custom Vertex AI regional API URL (default: {DEFAULT_VERTEX_API_URL})", @@ -4359,6 +4388,11 @@ if __name__ == "__main__": port=_get_env_int("HEADROOM_PORT", args.port), openai_api_url=_get_env_str("OPENAI_TARGET_API_URL", args.openai_api_url), anthropic_api_url=_get_env_str("ANTHROPIC_TARGET_API_URL", args.anthropic_api_url), + anthropic_buffered_request_timeout_seconds=_get_env_int( + "HEADROOM_ANTHROPIC_BUFFERED_REQUEST_TIMEOUT_SECONDS", + args.anthropic_buffered_request_timeout_seconds, + min_value=1, + ), vertex_api_url=_get_env_str("VERTEX_TARGET_API_URL", args.vertex_api_url), # Backend settings backend=_get_env_str("HEADROOM_BACKEND", args.backend), # type: ignore[arg-type] diff --git a/tests/test_anthropic_pre_upstream_backpressure.py b/tests/test_anthropic_pre_upstream_backpressure.py index 8b1496118..bbf47231e 100644 --- a/tests/test_anthropic_pre_upstream_backpressure.py +++ b/tests/test_anthropic_pre_upstream_backpressure.py @@ -243,13 +243,14 @@ class _DummyAnthropicHandler(AnthropicHandlerMixin): request_id: str | None = None, forwarder_name: str = "test_dummy", path_for_log: str | None = None, + timeout=None, ): # PR-A8 follow-up: A3 added byte-faithful kwargs to the real # ``_retry_request`` signature. The dummy stub doesn't need # to use them — just accept them so existing tests don't # break with TypeError on the new call sites. del original_body_bytes, body_mutated, mutation_reasons - del request_id, forwarder_name, path_for_log + del request_id, forwarder_name, path_for_log, timeout if self._raise_during_critical: raise RuntimeError("synthetic pre-upstream failure") enter = time.perf_counter() diff --git a/tests/test_cli_proxy_improvements.py b/tests/test_cli_proxy_improvements.py index b11bc90ba..ae1e1b6e1 100644 --- a/tests/test_cli_proxy_improvements.py +++ b/tests/test_cli_proxy_improvements.py @@ -14,6 +14,7 @@ Covers: from __future__ import annotations +import argparse from unittest.mock import patch import pytest @@ -329,6 +330,48 @@ class TestNewEnvVarWiring: assert result.exit_code == 0, result.output assert mock_run_server["config"].connect_timeout_seconds == 30 + def test_headroom_anthropic_buffered_timeout_from_env( + self, runner: CliRunner, mock_run_server: dict + ) -> None: + result = runner.invoke( + main, + ["proxy"], + env={"HEADROOM_ANTHROPIC_BUFFERED_REQUEST_TIMEOUT_SECONDS": "900"}, + catch_exceptions=False, + ) + assert result.exit_code == 0, result.output + assert mock_run_server["config"].anthropic_buffered_request_timeout_seconds == 900 + + def test_anthropic_buffered_timeout_cli_flag( + self, runner: CliRunner, mock_run_server: dict + ) -> None: + result = runner.invoke( + main, + ["proxy", "--anthropic-buffered-request-timeout-seconds", "901"], + catch_exceptions=False, + ) + assert result.exit_code == 0, result.output + assert mock_run_server["config"].anthropic_buffered_request_timeout_seconds == 901 + + def test_direct_server_env_timeout_zero_falls_back_to_default( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + import headroom.proxy.server as server_mod + + monkeypatch.delenv(server_mod._MULTI_WORKER_CONFIG_ENV, raising=False) + monkeypatch.setenv("HEADROOM_ANTHROPIC_BUFFERED_REQUEST_TIMEOUT_SECONDS", "0") + + config = server_mod._proxy_config_from_env() + assert config.anthropic_buffered_request_timeout_seconds == 600 + + def test_direct_server_timeout_parser_rejects_zero( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + import headroom.proxy.server as server_mod + + with pytest.raises(argparse.ArgumentTypeError): + server_mod._positive_int_arg("0") + def test_headroom_backend_from_env(self, runner: CliRunner, mock_run_server: dict) -> None: result = runner.invoke( main, diff --git a/tests/test_proxy/test_anthropic_buffered_timeout.py b/tests/test_proxy/test_anthropic_buffered_timeout.py new file mode 100644 index 000000000..69e44e7f5 --- /dev/null +++ b/tests/test_proxy/test_anthropic_buffered_timeout.py @@ -0,0 +1,412 @@ +from __future__ import annotations + +import pytest + +fastapi = pytest.importorskip("fastapi") +httpx = pytest.importorskip("httpx") + +from fastapi.testclient import TestClient # noqa: E402 + +from headroom.proxy.server import ProxyConfig, create_app # noqa: E402 + + +class _FakePrefixTracker: + def get_frozen_message_count(self) -> int: + return 0 + + def get_last_original_messages(self) -> list[dict]: + return [] + + def get_last_forwarded_messages(self) -> list[dict]: + return [] + + def update_from_response(self, **kwargs): # noqa: ANN003 + return None + + +class _BufferedPassthroughClient: + def __init__(self, response: httpx.Response) -> None: + self.response = response + self.calls: list[dict[str, object]] = [] + + async def request(self, method, url, headers=None, content=None, timeout=None): # noqa: ANN001 + self.calls.append( + { + "method": method, + "url": url, + "headers": headers, + "content": content, + "timeout": timeout, + } + ) + _assert_buffered_timeout(timeout) + return self.response + + async def get(self, url, headers=None, timeout=None): # noqa: ANN001 + return await self.request("GET", url, headers=headers, timeout=timeout) + + async def post(self, url, headers=None, content=None, timeout=None): # noqa: ANN001 + return await self.request("POST", url, headers=headers, content=content, timeout=timeout) + + async def aclose(self) -> None: + return None + + +def _assert_buffered_timeout(timeout: httpx.Timeout | None) -> None: + assert isinstance(timeout, httpx.Timeout) + assert timeout.connect == 3.0 + assert timeout.read == 19.0 + assert timeout.write == 7.0 + assert timeout.pool == 3.0 + + +def _make_config() -> ProxyConfig: + return ProxyConfig( + optimize=False, + cache_enabled=False, + rate_limit_enabled=False, + cost_tracking_enabled=False, + log_requests=False, + ccr_inject_tool=False, + ccr_handle_responses=False, + ccr_context_tracking=False, + image_optimize=False, + connect_timeout_seconds=3, + request_timeout_seconds=7, + anthropic_buffered_request_timeout_seconds=19, + ) + + +def _install_prefix_tracker(proxy) -> None: + tracker = _FakePrefixTracker() + proxy.session_tracker_store.compute_session_id = lambda request, model, messages: "s1" + proxy.session_tracker_store.get_or_create = lambda session_id, provider: tracker + + +def _anthropic_message_response() -> dict[str, object]: + return { + "id": "msg_test_1", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "ok"}], + "usage": { + "input_tokens": 12, + "output_tokens": 3, + "cache_read_input_tokens": 0, + "cache_creation_input_tokens": 0, + }, + } + + +def _anthropic_batch_response() -> dict[str, object]: + return {"id": "batch_test_1", "object": "batch", "status": "in_progress"} + + +def _anthropic_list_response() -> httpx.Response: + return httpx.Response( + 200, + json={"object": "list", "data": [], "first_id": None, "last_id": None}, + headers={"content-type": "application/json"}, + ) + + +def test_anthropic_messages_buffered_timeout_override_reaches_retry_request(): + config = _make_config() + app = create_app(config) + with TestClient(app) as client: + proxy = client.app.state.proxy + _install_prefix_tracker(proxy) + captured: dict[str, object] = {} + + async def _fake_retry(method, url, headers, body, stream=False, **kwargs): # noqa: ANN001 + timeout = kwargs.get("timeout") + captured["timeout"] = timeout + _assert_buffered_timeout(timeout) + return httpx.Response(200, json=_anthropic_message_response()) + + proxy._retry_request = _fake_retry # type: ignore[assignment] + + response = client.post( + "/v1/messages", + headers={ + "x-api-key": "test-key", + "anthropic-version": "2023-06-01", + "content-type": "application/json", + }, + json={ + "model": "claude-sonnet-4-6", + "max_tokens": 64, + "messages": [{"role": "user", "content": "hello"}], + }, + ) + + assert "timeout" in captured, response.text + assert isinstance(captured["timeout"], httpx.Timeout) + + +def test_anthropic_batch_create_buffered_timeout_override_reaches_retry_request(): + config = _make_config() + app = create_app(config) + with TestClient(app) as client: + proxy = client.app.state.proxy + captured: dict[str, object] = {} + + async def _fake_retry(method, url, headers, body, stream=False, **kwargs): # noqa: ANN001 + timeout = kwargs.get("timeout") + captured["timeout"] = timeout + _assert_buffered_timeout(timeout) + return httpx.Response(200, json=_anthropic_batch_response()) + + proxy._retry_request = _fake_retry # type: ignore[assignment] + + response = client.post( + "/v1/messages/batches", + headers={ + "x-api-key": "test-key", + "anthropic-version": "2023-06-01", + "content-type": "application/json", + }, + json={ + "requests": [ + { + "custom_id": "req-1", + "params": { + "model": "claude-sonnet-4-6", + "max_tokens": 64, + "messages": [{"role": "user", "content": "hello"}], + }, + } + ] + }, + ) + + assert response.status_code == 200, response.text + assert isinstance(captured["timeout"], httpx.Timeout) + + +def test_anthropic_batch_passthrough_buffered_timeout_override_reaches_http_client(): + config = _make_config() + app = create_app(config) + with TestClient(app) as client: + proxy = client.app.state.proxy + http_client = _BufferedPassthroughClient(_anthropic_list_response()) + proxy.http_client = http_client + + response = client.get( + "/v1/messages/batches", + headers={ + "x-api-key": "test-key", + "anthropic-version": "2023-06-01", + }, + ) + + assert response.status_code == 200, response.text + assert len(http_client.calls) == 1 + assert isinstance(http_client.calls[0]["timeout"], httpx.Timeout) + + +def test_anthropic_batch_results_buffered_timeout_override_reaches_http_client_get(): + config = _make_config() + app = create_app(config) + with TestClient(app) as client: + proxy = client.app.state.proxy + http_client = _BufferedPassthroughClient( + httpx.Response( + 200, + content=b'{"custom_id":"req-1","result":{"type":"succeeded"}}\n', + headers={"content-type": "application/jsonl"}, + ) + ) + proxy.http_client = http_client + + response = client.get( + "/v1/messages/batches/batch_test_1/results", + headers={ + "x-api-key": "test-key", + "anthropic-version": "2023-06-01", + }, + ) + + assert response.status_code == 200, response.text + assert len(http_client.calls) == 1 + assert http_client.calls[0]["method"] == "GET" + assert isinstance(http_client.calls[0]["timeout"], httpx.Timeout) + + +def test_anthropic_ccr_continuation_uses_buffered_timeout() -> None: + config = _make_config() + config.ccr_inject_tool = True + config.ccr_handle_responses = True + app = create_app(config) + + class _CCRHandler: + def has_ccr_tool_calls(self, response, provider): # noqa: ANN001 + return True + + async def handle_response( # noqa: ANN001 + self, + response, + optimized_messages, + tools, + api_call_fn, + provider, + ): + return await api_call_fn( + optimized_messages + + [{"role": "assistant", "content": response.get("content", [])}], + tools, + ) + + with TestClient(app) as client: + proxy = client.app.state.proxy + _install_prefix_tracker(proxy) + proxy.ccr_response_handler = _CCRHandler() + http_client = _BufferedPassthroughClient( + httpx.Response(200, json=_anthropic_message_response()) + ) + proxy.http_client = http_client + + async def _fake_retry(method, url, headers, body, stream=False, **kwargs): # noqa: ANN001 + return httpx.Response(200, json=_anthropic_message_response()) + + proxy._retry_request = _fake_retry # type: ignore[assignment] + + client.post( + "/v1/messages", + headers={ + "x-api-key": "test-key", + "anthropic-version": "2023-06-01", + "content-type": "application/json", + }, + json={ + "model": "claude-sonnet-4-6", + "max_tokens": 64, + "messages": [{"role": "user", "content": "hello"}], + }, + ) + + assert len(http_client.calls) == 1 + assert http_client.calls[0]["method"] == "POST" + assert isinstance(http_client.calls[0]["timeout"], httpx.Timeout) + + +def test_anthropic_memory_continuation_uses_buffered_timeout() -> None: + config = _make_config() + config.memory_enabled = True + app = create_app(config) + + class _MemoryHandler: + def __init__(self) -> None: + self.config = type( + "MemoryConfig", + (), + { + "inject_context": False, + "inject_tools": False, + "project_root_override": "", + }, + )() + self.initialized = False + self.backend = None + + def get_beta_headers(self) -> dict[str, str]: + return {} + + def has_memory_tool_calls(self, response, provider): # noqa: ANN001 + return True + + async def handle_memory_tool_calls( # noqa: ANN001 + self, + response, + user_id, + provider, + **kwargs, + ): + return [{"type": "tool_result", "tool_use_id": "mem_1", "content": "memory"}] + + with TestClient(app) as client: + proxy = client.app.state.proxy + _install_prefix_tracker(proxy) + proxy.memory_handler = _MemoryHandler() + captured_timeouts: list[httpx.Timeout | None] = [] + + async def _fake_retry(method, url, headers, body, stream=False, **kwargs): # noqa: ANN001 + timeout = kwargs.get("timeout") + captured_timeouts.append(timeout) + if len(body["messages"]) > 1: + _assert_buffered_timeout(timeout) + return httpx.Response(200, json=_anthropic_message_response()) + + proxy._retry_request = _fake_retry # type: ignore[assignment] + + client.post( + "/v1/messages", + headers={ + "x-api-key": "test-key", + "anthropic-version": "2023-06-01", + "content-type": "application/json", + "x-headroom-user-id": "user-1", + }, + json={ + "model": "claude-sonnet-4-6", + "max_tokens": 64, + "messages": [{"role": "user", "content": "hello"}], + }, + ) + + assert len(captured_timeouts) == 2 + assert all(isinstance(timeout, httpx.Timeout) for timeout in captured_timeouts) + + +def test_retry_request_without_override_uses_client_default_timeout(): + config = _make_config() + app = create_app(config) + with TestClient(app) as client: + proxy = client.app.state.proxy + captured: list[dict[str, object]] = [] + + async def _fake_retry(method, url, headers, body, stream=False, **kwargs): # noqa: ANN001 + captured.append(kwargs) + return httpx.Response( + 200, + json={ + "id": "chatcmpl-1", + "object": "chat.completion", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "ok"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 5, "completion_tokens": 1, "total_tokens": 6}, + }, + ) + + proxy._retry_request = _fake_retry # type: ignore[assignment] + + response = client.post( + "/v1/chat/completions", + headers={"authorization": "Bearer test-key", "content-type": "application/json"}, + json={ + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hello"}], + }, + ) + + assert response.status_code == 200, response.text + assert len(captured) == 1 + assert "timeout" not in captured[0], ( + "_retry_request without an override must not pass timeout=None to httpx" + ) + + +def test_generic_proxy_timeout_defaults_stay_unchanged(): + app = create_app(ProxyConfig()) + with TestClient(app) as client: + timeout = client.app.state.proxy.http_client.timeout + + assert timeout.connect == 10.0 + assert timeout.read == 300.0 + assert timeout.write == 300.0 + assert timeout.pool == 10.0 diff --git a/tests/test_proxy_handler_helpers.py b/tests/test_proxy_handler_helpers.py index f2bca5c99..daa0c42f7 100644 --- a/tests/test_proxy_handler_helpers.py +++ b/tests/test_proxy_handler_helpers.py @@ -194,10 +194,11 @@ class _RetryThenSuccessClient: def __init__(self) -> None: self.attempts = 0 - async def post(self, url, content, headers): # noqa: ANN001, ANN201 + async def post(self, url, content, headers, timeout=None): # noqa: ANN001, ANN201 self.attempts += 1 if self.attempts == 1: raise httpx.ConnectTimeout("connect timed out") + del timeout request = httpx.Request("POST", url, headers=headers, content=content) return httpx.Response(200, request=request, content=b"{}")