mirror of
https://github.com/headroomlabs-ai/headroom.git
synced 2026-08-27 14:17:10 -04:00
fix(proxy): add an Anthropic buffered read-timeout override (#1331)
## Description Buffered Anthropic `/v1/messages` requests still use Headroom's generic 300-second read timeout, which can produce proxy-generated `502 ReadTimeout` errors on long turns. This adds a dedicated buffered Anthropic timeout, keeps it applied across CCR and memory continuations plus batch paths, and makes the direct server entrypoint enforce the same positive-integer contract as the Click CLI. Closes #1261. ## Type of Change - [x] Bug fix (non-breaking change that fixes an issue) ## Changes Made - Added `anthropic_buffered_request_timeout_seconds` for buffered Anthropic reads. - Routed `/v1/messages`, CCR continuation, memory continuation, batch create, batch passthrough, and batch results through that timeout. - Enforced the same positive-integer validation for `HEADROOM_ANTHROPIC_BUFFERED_REQUEST_TIMEOUT_SECONDS` and `--anthropic-buffered-request-timeout-seconds` in both startup paths. - Added focused regressions and updated `CHANGELOG.md`. ## Testing - [x] `uv run pytest tests/test_proxy/test_anthropic_buffered_timeout.py tests/test_cli_proxy_improvements.py::TestNewEnvVarWiring` - [x] `uv run ruff check .` - [x] `uv run ruff format . --check` ### Test Output ```text $ uv run pytest tests/test_proxy/test_anthropic_buffered_timeout.py tests/test_cli_proxy_improvements.py::TestNewEnvVarWiring 17 passed in 3.42s $ uv run ruff check . All checks passed! $ uv run ruff format . --check 966 files already formatted ``` ## Real Behavior Proof - Environment: local FastAPI `TestClient` with stubbed retry and HTTP client seams - Exact command / steps: run `uv run pytest tests/test_proxy/test_anthropic_buffered_timeout.py tests/test_cli_proxy_improvements.py::TestNewEnvVarWiring`; the tests build `ProxyConfig(request_timeout_seconds=7, connect_timeout_seconds=3, anthropic_buffered_request_timeout_seconds=19)`, drive `/v1/messages`, `/v1/messages/batches`, `/v1/messages/batches/{batch_id}/results`, a CCR continuation, and a memory continuation through `TestClient`, then verify `HEADROOM_ANTHROPIC_BUFFERED_REQUEST_TIMEOUT_SECONDS=0` falls back to `600`, `--anthropic-buffered-request-timeout-seconds 0` is rejected, and default proxy timeouts stay `read=300` and `write=300` - Observed result: buffered Anthropic paths use `httpx.Timeout(connect=3, read=19, write=7, pool=3)`, continuation requests stay on that same budget, invalid zero-valued startup config is rejected or ignored back to the default, and unrelated proxy timeout defaults stay unchanged - Not tested: live upstream Anthropic latency beyond the focused stubbed-timeout regression ## Review Readiness - [x] I have performed a self-review - [x] This PR is ready for human review ## Checklist - [x] My code follows the project's style guidelines - [x] I have added tests that prove the fix - [x] New and existing unit tests pass locally with my changes - [x] I have updated the CHANGELOG.md if applicable
This commit is contained in:
parent
82384022bd
commit
3be2526b76
9 changed files with 541 additions and 9 deletions
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
412
tests/test_proxy/test_anthropic_buffered_timeout.py
Normal file
412
tests/test_proxy/test_anthropic_buffered_timeout.py
Normal file
|
|
@ -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
|
||||
|
|
@ -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"{}")
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue