From 259fd485a2f3aadf173f163d7dd9c0b2a1a85ba7 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A2=A8=E8=8F=8A?= <277378677+maiphucgiang@users.noreply.github.com> Date: Wed, 16 Sep 2026 22:50:03 +0800 Subject: [PATCH] Keep API preparation responsive and preserve retry hints Move credential-dependent request preparation off the event loop, preserve bounded Retry-After hints through protocol errors and model cooldowns, and forward explicit Responses prompt cache keys. Add cross-profile and protocol regressions; verify all 50 backend test scripts and an isolated loopback deployment. Keep replay policy and version unchanged. --- app/adapters/responses_adapter.py | 2 +- app/upstream_io.py | 32 ++++ converter.py | 51 ++++-- docs/advanced.md | 4 +- docs/advanced.zh-CN.md | 4 +- tests/test_api_flow.py | 272 ++++++++++++++++++++++++++++++ 6 files changed, 344 insertions(+), 21 deletions(-) create mode 100644 tests/test_api_flow.py diff --git a/app/adapters/responses_adapter.py b/app/adapters/responses_adapter.py index bf46ee1..171c671 100644 --- a/app/adapters/responses_adapter.py +++ b/app/adapters/responses_adapter.py @@ -75,7 +75,7 @@ def responses_request_to_chat(body: dict) -> dict: # Forward supported parameters. for key in ("temperature", "top_p", "stop", "seed", "presence_penalty", "frequency_penalty", - "response_format", "reasoning_effort", "parallel_tool_calls"): + "response_format", "reasoning_effort", "parallel_tool_calls", "prompt_cache_key"): if key in body: chat[key] = body[key] diff --git a/app/upstream_io.py b/app/upstream_io.py index 0dcab36..8ee12f9 100644 --- a/app/upstream_io.py +++ b/app/upstream_io.py @@ -2,6 +2,10 @@ import asyncio from contextlib import asynccontextmanager +from datetime import timezone +from email.utils import parsedate_to_datetime +import math +import time import json import httpx @@ -21,6 +25,34 @@ def __init__(self, status, raw): class UpstreamHTTPError(UpstreamResponseError): """Distinguish actual upstream HTTP errors from failures synthesized while collecting a response.""" + def __init__(self, status, raw, *, retry_after=None): + super().__init__(status, raw) + self.retry_after = retry_after + self.headers = {"Retry-After": str(retry_after)} if retry_after is not None else {} + + +MAX_RETRY_AFTER = 86400 + + +def parse_retry_after(value, *, now=None) -> int | None: + """Normalize bounded Retry-After seconds or HTTP dates; ignore invalid or expired values.""" + if not isinstance(value, str) or len(value) > 128 or not value.isascii() or not value.isprintable(): + return None + value = value.strip() + if not value: + return None + try: + if value.isdecimal(): + delay = int(value) + else: + deadline = parsedate_to_datetime(value) + if deadline.tzinfo is None: + deadline = deadline.replace(tzinfo=timezone.utc) # Obsolete HTTP asctime uses GMT. + delay = deadline.timestamp() - (time.time() if now is None else now) + return math.ceil(delay) if 0 <= delay <= MAX_RETRY_AFTER else None + except (TypeError, ValueError, OverflowError): + return None + class ChatSSEAccumulator: """Collect Chat SSE and reject error events, empty output and incomplete streams.""" diff --git a/converter.py b/converter.py index 2dc7deb..3750dc1 100644 --- a/converter.py +++ b/converter.py @@ -7,6 +7,7 @@ import asyncio import hashlib import json +import math import os import re import secrets @@ -57,7 +58,7 @@ def desensitize_body(body, roles=("system",), desensitize_harness_user=False, from app.credential_io import (CredentialFileError, read_import_file, atomic_write_credential, credential_file_lock) from app.upstream_io import (ChatSSEAccumulator, UpstreamHTTPError, UpstreamResponseError, - open_backend_stream, read_bounded_error) + open_backend_stream, parse_retry_after, read_bounded_error) from app.inference_auth import require_api_key from app.content_filter import ContentFilterDetector, is_filter_error from app.request_limits import ImageLimitError, apply_image_policy @@ -868,7 +869,7 @@ def cooldown(self, cm: CredentialManager, reason: str = "", *, generation=None): _log(f"[cred] 凭证熔断 {CRED_COOLDOWN}s: {Path(cm.path).name} {reason}") def note_status(self, cm: CredentialManager | None, status: int, - model: str | None = None, raw: bytes = b"", *, generation=None): + model: str | None = None, raw: bytes = b"", *, generation=None, retry_after=None): """Apply credential-wide auth cooldowns, per-model 429 cooldowns and backend/model backoff.""" if cm is None: return @@ -882,7 +883,11 @@ def note_status(self, cm: CredentialManager | None, status: int, if status != 429 or not model: return now = time.time() - until = _parse_reset_time(raw) or now + MODEL_COOLDOWN + if retry_after is not None: + until = now + retry_after + else: + reset = _parse_reset_time(raw) + until = reset if reset is not None and reset > now else now + MODEL_COOLDOWN until = min(until, now + MODEL_COOLDOWN_MAX) with self._lock, (cm._lock if generation is not None else nullcontext()): if not self._lease_matches(cm, generation): @@ -891,7 +896,9 @@ def note_status(self, cm: CredentialManager | None, status: int, for e in self._entries: if e["cm"] is cm: routed_model = _upstream_model(model, self._entry_profile(e)) - self._model_fail[(e["id"], routed_model)] = until + key = (e["id"], routed_model) + until = max(until, self._model_fail.get(key, 0.0)) + self._model_fail[key] = until _log(f"[cred] 模型冷却 {model} @ {Path(cm.path).name} 至 " f"{time.strftime('%m-%d %H:%M:%S', time.localtime(until))} (HTTP 429)") @@ -1539,7 +1546,9 @@ def _cred_for(payload: dict, model: str | None = None, *, region=None, tried=()) until = pool.model_cooldown_until(model, region=region) if until: t = time.strftime("%Y-%m-%d %H:%M:%S", time.localtime(until)) - raise HTTPException(status_code=429, detail={"error": { + raise HTTPException(status_code=429, + headers={"Retry-After": str(max(1, math.ceil(until - time.time())))}, + detail={"error": { "message": f"模型 {model} 额度冷却中(全部凭证),预计 {t} 重置后恢复", "type": "rate_limit_error"}}) blocked = pool.model_block_until(model, region=region) @@ -1593,12 +1602,12 @@ def _note_cred_model_ok(cred, model: str | None) -> None: pool.note_model_ok(cm, model) -def _note_cred_status(cred, status: int, model: str | None = None, raw: bytes = b""): +def _note_cred_status(cred, status: int, model: str | None = None, raw: bytes = b"", *, retry_after=None): """Record generation-scoped authentication, quota and unsupported-model failures.""" pool = CONFIG.get("cred_pool") if pool is not None and cred is not None: cm, generation = cred if isinstance(cred, tuple) else (cred, None) - pool.note_status(cm, status, model=model, raw=raw, generation=generation) + pool.note_status(cm, status, model=model, raw=raw, generation=generation, retry_after=retry_after) @app.get("/health") def health(): @@ -2362,7 +2371,7 @@ async def chat_completions(request: Request, # Forward only supported request fields. client_wants_stream = _client_wants_stream(payload) body = {k: payload[k] for k in PASSTHROUGH_BODY_KEYS if k in payload} - body = _prepare_chat_body(body) + body = await run_in_threadpool(_prepare_chat_body, body) # Record request metadata. model_name = payload.get("model", "auto") @@ -2600,11 +2609,12 @@ def _safe_err_raw(raw: bytes, status: int) -> dict: return {"error": {"message": raw.decode("utf-8", "replace")[:500], "type": "upstream_error", "code": status}} -def _check_upstream_status(status, raw, cred, model): +def _check_upstream_status(status, raw, cred, model, *, headers=None): if status != 200: + retry_after = parse_retry_after((headers or {}).get("Retry-After")) if not is_filter_error(raw): - _note_cred_status(cred, status, model=model, raw=raw) - raise UpstreamHTTPError(status, raw) + _note_cred_status(cred, status, model=model, raw=raw, retry_after=retry_after) + raise UpstreamHTTPError(status, raw, retry_after=retry_after) def _upstream_failure(error, model_name, t0, rid): @@ -2642,7 +2652,8 @@ async def _fetch_checked_chat(url, headers, body, model_name, rid, cred=None, *, rejection = None async with _backend_stream(url, headers, body, rid=rid, model_name=model_name) as response: if response.status_code != 200: - _check_upstream_status(response.status_code, await read_bounded_error(response), cred, body.get("model")) + _check_upstream_status(response.status_code, await read_bounded_error(response), cred, body.get("model"), + headers=response.headers) else: _note_cred_model_ok(cred, body.get("model")) try: @@ -2708,7 +2719,8 @@ async def _chat_sse_lines(url, headers, body, model_name, t0, rid, cred=None, *, budget = CONFIG["log_body_limit"] if CONFIG.get("log_path") else 0 async with _backend_stream(url, headers, body, rid=rid, model_name=model_name) as response: if response.status_code != 200: - _check_upstream_status(response.status_code, await read_bounded_error(response), cred, body.get("model")) + _check_upstream_status(response.status_code, await read_bounded_error(response), cred, body.get("model"), + headers=response.headers) else: _note_cred_model_ok(cred, body.get("model")) async for line in response.aiter_lines(): @@ -2798,6 +2810,7 @@ def __init__(self, status, raw, error=None): self.status = status self.raw = raw self.error = error + self.headers = error.headers if isinstance(error, UpstreamHTTPError) else None super().__init__(f"stream failed before first byte (HTTP {status})") @@ -2918,7 +2931,7 @@ async def _stream_plan(payload, canonical, model_name, rid, t0, make, routed, cr recovered = observe_failure_seq() # Recover only this failure sequence. tried.append(cred) limit = _failover_limit() - surface = HTTPException(status_code=failure.status, + surface = HTTPException(status_code=failure.status, headers=failure.headers, detail=_safe_err_raw(failure.raw, failure.status)) if limit <= 0 or len(tried) > limit or not _failover_safe(failure.error, failure.raw): raise surface from None @@ -2967,7 +2980,8 @@ async def _routed_fetch(payload, canonical, model_name, rid, t0, fetch, routed, recovered = observe_failure_seq() tried.append(cred) limit = _failover_limit() - surface = HTTPException(status_code=status, detail=_safe_err_raw(raw, status)) + surface = HTTPException(status_code=status, detail=_safe_err_raw(raw, status), + headers=error.headers if isinstance(error, UpstreamHTTPError) else None) if limit <= 0 or len(tried) > limit or not _failover_safe(error, raw): raise surface from None try: @@ -3036,7 +3050,7 @@ async def create_response(request: Request, chat_body, projection_stats = project_responses_chat_body( chat_body, keep_tool_metadata=CONFIG.get("keep_tool_metadata", False)) - chat_body = _prepare_chat_body(chat_body) + chat_body = await run_in_threadpool(_prepare_chat_body, chat_body) client_wants_stream = _client_wants_stream(payload) model_name = payload.get("model", "auto") @@ -3085,7 +3099,8 @@ async def fetch(routed, cred, headers, url): converter.finish() except (httpx.HTTPError, UpstreamResponseError) as error: status, raw = _upstream_failure(error, model_name, t0, rid) - raise HTTPException(status_code=status, detail=_safe_err_raw(raw, status)) from None + raise HTTPException(status_code=status, detail=_safe_err_raw(raw, status), + headers=error.headers if isinstance(error, UpstreamHTTPError) else None) from None except ClientHungUp: return _hungup_response(rid, model_name, t0) result = converter.get_nonstream_response() @@ -3153,7 +3168,7 @@ async def create_message(request: Request, except Exception as e: raise HTTPException(status_code=400, detail={"error": {"message": f"request conversion error: {e}", "type": "invalid_request_error"}}) - chat_body = _prepare_chat_body(chat_body) + chat_body = await run_in_threadpool(_prepare_chat_body, chat_body) model_name = payload.get("model", "auto") chat_messages = chat_body.get("messages", []) rid = os.urandom(4).hex() diff --git a/docs/advanced.md b/docs/advanced.md index 50e53c0..c6dfb7c 100644 --- a/docs/advanced.md +++ b/docs/advanced.md @@ -157,7 +157,9 @@ Credential domain / token issuer determine the product identity. Chat and refres - `--image-policy error` returns local `413 / too_many_images`. JSON still over budget after processing returns `413 / request_too_large`, without further text truncation to fit the limit. - Image count does not guarantee acceptable individual image sizes or model vision support. URL/base64 images can be converted; Responses image `file_id` is unsupported. - When `stream` is omitted all three endpoints follow the protocol default and return a complete JSON response; `stream` must be a boolean. Streaming Responses and Chat/Messages with tools aggregate and validate before emitting SSE; not every path forwards tokens in real time. -- Inference errors are shaped per client protocol: OpenAI routes return a top-level `error` object and Messages returns `{"type": "error", ...}`; status codes and `Retry-After` are unchanged. +- Inference errors follow the client protocol: OpenAI routes return a top-level `error` object and Messages returns `{"type": "error", ...}`. Status codes are retained; errors after streaming starts are reported through SSE without replay. +- Valid upstream `Retry-After` values (0–86400 seconds or equivalent HTTP dates) are returned as seconds before streaming starts; 429 only cools the selected account/model. Invalid or expired values fall back to the body's reset time or 600 seconds. Pool-generated 429 responses include the remaining wait. +- Chat and Responses preserve an explicit client `prompt_cache_key` without generating one; cache hits and savings depend on the upstream. - Unsupported capabilities are rejected rather than silently degraded: chat `n` other than 1 and the Responses state fields `previous_response_id`/`conversation` (this gateway keeps no server-side response state) return 400; length-truncated or content-filtered Responses are reported as `incomplete`, never disguised as `completed`. - `/v1/messages/count_tokens` returns a character-based heuristic estimate for budgeting, not an exact count. - Text logs and SQLite auditing have separate budgets. Logs contain bounded, redacted previews, not complete original requests. Treat logs, credential exports and backups as private data. diff --git a/docs/advanced.zh-CN.md b/docs/advanced.zh-CN.md index a8e70d6..db5a831 100644 --- a/docs/advanced.zh-CN.md +++ b/docs/advanced.zh-CN.md @@ -154,7 +154,9 @@ WebUI 可以直接上传文件;以下限制针对 `POST /admin/credentials` - `--image-policy error` 在本地返回 `413 / too_many_images`。处理后仍超过字节上限则返回 `413 / request_too_large`,不为满足预算继续截断文本。 - 图片数量合规不保证单图大小或模型视觉能力满足上游要求。URL/base64 图片可转换,Responses 图片 `file_id` 不支持。 - 省略 `stream` 时三个端点都按协议默认返回完整 JSON(非流式);`stream` 必须是布尔值。Responses 流式以及带工具的 Chat / Messages 流式先聚合校验,再输出 SSE,并非所有路径都实时逐 token 转发。 -- 推理端点的错误体按客户端协议成形:OpenAI 路由为顶层 `error` 对象,Messages 路由为 `{"type": "error", ...}`;状态码与 `Retry-After` 保持不变。 +- 推理错误按客户端协议成形:OpenAI 路由为顶层 `error` 对象,Messages 路由为 `{"type": "error", ...}`;保留状态码,开流后的错误只用 SSE 报告,不重放。 +- 上游有效 `Retry-After`(0–86400 秒或对应 HTTP 日期)规范化为秒并在开流前返回;429 仅冷却对应账号/模型。无效或过期值回落正文重置时间或默认 600 秒;本地全凭据冷却的 429 返回剩余等待秒数。 +- Chat 与 Responses 保留客户端显式 `prompt_cache_key`,不自动生成;缓存命中和节费取决于上游。 - 不支持的能力显式拒绝而非静默降级:Chat 的 `n≠1`、Responses 的 `previous_response_id`/`conversation`(本网关不保存服务端响应状态)返回 400;长度截断或审核过滤的 Responses 标记为 `incomplete`,不伪装为 `completed`。 - `/v1/messages/count_tokens` 返回字符启发式估算值,仅作预算参考,不是精确计数。 - 兼容文本日志和 SQLite 审计使用独立预算;日志仅记录有界、脱敏预览,不是完整原始请求。日志、凭证导出和备份仍须按私有数据保管。 diff --git a/tests/test_api_flow.py b/tests/test_api_flow.py new file mode 100644 index 0000000..52302ec --- /dev/null +++ b/tests/test_api_flow.py @@ -0,0 +1,272 @@ +"""Keep request preparation responsive and preserve retry hints and explicit cache keys.""" +import sys +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +import asyncio +from copy import deepcopy +from email.utils import formatdate +import json +import threading +import time +import unittest +from unittest.mock import patch + +import httpx + +import converter as gateway +from app import upstream_io +from app.adapters.responses_adapter import responses_request_to_chat +import test_region_routing as fixtures + +REAL_ASYNC_CLIENT = httpx.AsyncClient + + +class RetryAfterParsingTests(unittest.TestCase): + def test_standard_seconds_and_http_dates(self): + now = 1_700_000_000 + cases = (("0", 0), ("2", 2), (" 17 ", 17), ("86400", 86400), + (formatdate(now + 30, usegmt=True), 30)) + for value, expected in cases: + with self.subTest(value=value): + self.assertEqual(upstream_io.parse_retry_after(value, now=now), expected) + self.assertEqual(upstream_io.parse_retry_after(formatdate(now + 1, usegmt=True), + now=now + 0.25), 1) + + def test_invalid_expired_and_excessive_values_are_ignored(self): + now = 1_700_000_000 + values = (None, "", "-1", "+2", "1.5", "nan", "inf", "two", "2, 3", "86401", + "9" * 200, "2", "2\r\nSet-Cookie: bad", "2\n", "2\x00", + formatdate(now - 60, usegmt=True), formatdate(now + 86401, usegmt=True)) + for value in values: + with self.subTest(value=value): + self.assertIsNone(upstream_io.parse_retry_after(value, now=now)) + + +class GatewayFixture: + def setUp(self): + super().setUp() + self.fx = fixtures.RegionRoutingTests() + self.addCleanup(self.fx.doCleanups) + self.fx.setUp() + self.fx.allowed_profiles = set(fixtures.PROFILES) + self.enterContext(patch.object(gateway, "_log")) + + def responder(self, callback): + transport = httpx.MockTransport(callback) + return patch.object(httpx, "AsyncClient", side_effect=lambda **kw: + REAL_ASYNC_CLIENT(transport=transport, **kw)) + + +class RetryAfterEndpointTests(GatewayFixture, unittest.TestCase): + def test_retry_hint_reaches_all_protocols_and_only_cools_the_selected_model(self): + for profile in fixtures.PROFILES: + for route in fixtures.GENERATIONS: + for stream in (False, True): + for status in (429, 503): + with self.subTest(profile=profile, route=route, stream=stream, status=status): + self.fx.configure(profiles=(profile,)) + seen = [] + + def reject(request): + seen.append(request) + return httpx.Response(status, headers={"Retry-After": "17", "X-Private": "hidden"}, + json={"error": {"message": "synthetic rejection", "code": "quota"}}) + + before = time.time() + with self.responder(reject): + response = self.fx.client.post("/v1/" + route, + json=self.fx.payload(route, stream=stream)) + self.assertEqual(response.status_code, status, response.text) + self.assertEqual(response.headers.get("Retry-After"), "17") + self.assertNotIn("X-Private", response.headers) + self.assertEqual(response.json()["error"]["code"], "quota") + self.assertEqual(len(seen), 1, "A retry hint must not enable replay") + entry = self.fx.entries[profile] + self.assertEqual(entry["fail_until"], 0) + if status == 429: + until = self.fx.pool._model_fail[(entry["id"], "shared-model")] + self.assertGreaterEqual(until, before + 17) + self.assertLessEqual(until, time.time() + 17) + self.assertTrue(self.fx.pool._model_healthy(entry, "other-model")) + else: + self.assertEqual(self.fx.pool._model_fail, {}) + + def test_http_date_is_normalized_to_seconds(self): + self.fx.configure(profiles=("cn-cli",)) + date = formatdate(time.time() + 60, usegmt=True) + with self.responder(lambda request: httpx.Response(429, headers={"Retry-After": date}, + json={"error": {"message": "slow down"}})): + response = self.fx.client.post("/v1/responses", json=self.fx.payload("responses", stream=True)) + self.assertEqual(response.status_code, 429) + self.assertLessEqual(int(response.headers["Retry-After"]), 60) + self.assertGreaterEqual(int(response.headers["Retry-After"]), 55) + + def test_invalid_header_falls_back_to_body_or_default_cooldown(self): + for reset_seconds in (None, -60, 90): + with self.subTest(reset_seconds=reset_seconds): + self.fx.configure(profiles=("cn-cli",)) + before = time.time() + message = "slow down" + if reset_seconds is not None: + message += time.strftime(" %Y-%m-%d %H:%M:%S UTC+0", time.gmtime(before + reset_seconds)) + with self.responder(lambda request: httpx.Response(429, headers={"Retry-After": "999999999"}, + json={"error": {"message": message}})): + response = self.fx.client.post("/v1/chat/completions", json=self.fx.payload()) + self.assertEqual(response.status_code, 429) + self.assertNotIn("Retry-After", response.headers) + until = next(iter(self.fx.pool._model_fail.values())) + expected = 90 if reset_seconds == 90 else gateway.MODEL_COOLDOWN + self.assertAlmostEqual(until - before, expected, delta=2) + + def test_pool_generated_429_includes_remaining_wait_without_sending_again(self): + self.fx.configure(profiles=("cn-cli",)) + entry = self.fx.entries["cn-cli"] + self.fx.pool.note_status(entry["cm"], 429, model="shared-model", retry_after=17) + for route in fixtures.GENERATIONS: + with self.subTest(route=route): + response = self.fx.client.post("/v1/" + route, json=self.fx.payload(route, stream=True)) + self.assertEqual(response.status_code, 429, response.text) + self.assertGreater(int(response.headers["Retry-After"]), 0) + self.assertLessEqual(int(response.headers["Retry-After"]), 17) + self.assertEqual(self.fx.requests, []) + + def test_explicit_zero_does_not_invent_a_cooldown(self): + self.fx.configure(profiles=("cn-cli",)) + entry = self.fx.entries["cn-cli"] + self.fx.pool.note_status(entry["cm"], 429, model="shared-model", retry_after=0) + self.assertTrue(self.fx.pool._model_healthy(entry, "shared-model")) + + def test_stale_lease_cannot_apply_retry_hint_to_replaced_credentials(self): + self.fx.configure(profiles=("cn-cli",)) + cm = self.fx.pool.first() + old = (cm, cm._generation) + cm.invalidate() + gateway._note_cred_status(old, 429, model="shared-model", retry_after=17) + self.assertEqual(self.fx.pool._model_fail, {}) + + def test_concurrent_shorter_hints_cannot_clear_an_active_cooldown(self): + self.fx.configure(profiles=("cn-cli",)) + cm = self.fx.pool.first() + self.fx.pool.note_status(cm, 429, model="shared-model", retry_after=31) + original = dict(self.fx.pool._model_fail) + for delay in (2, 0): + self.fx.pool.note_status(cm, 429, model="shared-model", retry_after=delay) + self.assertEqual(self.fx.pool._model_fail, original) + + def test_exhausted_pool_keeps_the_original_upstream_hint(self): + for route in fixtures.GENERATIONS: + for stream in (False, True): + with self.subTest(route=route, stream=stream): + self.fx.configure(profiles=("cn-cli",)) + with self.responder(lambda request: httpx.Response(429, headers={"Retry-After": "17"}, + json={"error": {"message": "original rejection"}})), patch.dict(gateway.CONFIG, {"failover_max": 1}): + response = self.fx.client.post("/v1/" + route, json=self.fx.payload(route, stream=stream)) + self.assertEqual(response.status_code, 429) + self.assertEqual(response.headers["Retry-After"], "17") + self.assertEqual(response.json()["error"]["message"], "original rejection") + + def test_success_headers_do_not_leak_into_sse_failures_or_enable_replay(self): + data = (b'data: {"choices":[{"index":0,"delta":{"content":"partial"}}]}\n\n' + b'data: {"error":{"message":"late rejection","code":429}}\n\n') + for route in fixtures.GENERATIONS: + with self.subTest(route=route): + self.fx.configure() + seen = [] + + def respond(request): + seen.append(request) + return httpx.Response(200, headers={"Retry-After": "17"}, content=data) + + with self.responder(respond), patch.dict(gateway.CONFIG, {"failover_max": 1}): + response = self.fx.client.post("/v1/" + route, json=self.fx.payload(route, stream=True)) + self.assertEqual(response.status_code, 502 if route == "responses" else 200) + self.assertIn("late rejection", response.text) + self.assertNotIn("Retry-After", response.headers) + self.assertEqual(len(seen), 1) + + + def test_failover_discards_old_hint_on_success_and_preserves_last_failure_hint(self): + for final_status in (200, 429): + for route in fixtures.GENERATIONS: + for stream in (False, True): + with self.subTest(final_status=final_status, route=route, stream=stream): + self.fx.configure() + seen = [] + + def respond(request): + seen.append(request) + if len(seen) == 2 and final_status == 200: + return httpx.Response(200, content=fixtures.success_sse()) + return httpx.Response(429, headers={"Retry-After": "17" if len(seen) == 1 else "31"}, + json={"error": {"message": "slow down"}}) + + with self.responder(respond), patch.dict(gateway.CONFIG, {"failover_max": 1}): + response = self.fx.client.post("/v1/" + route, + json=self.fx.payload(route, stream=stream)) + self.assertEqual(response.status_code, final_status, response.text) + self.assertEqual(len(seen), 2) + self.assertEqual(response.headers.get("Retry-After"), "31" if final_status == 429 else None) + + +class CacheKeyTests(GatewayFixture, unittest.TestCase): + def test_responses_adapter_preserves_explicit_key_without_mutating_input(self): + for key in ("client-cache-key", "", None): + with self.subTest(key=key): + payload = {"input": "hi", "prompt_cache_key": key} + original = deepcopy(payload) + self.assertEqual(responses_request_to_chat(payload).get("prompt_cache_key"), key) + self.assertIn("prompt_cache_key", responses_request_to_chat(payload)) + self.assertEqual(payload, original) + self.assertNotIn("prompt_cache_key", responses_request_to_chat({"input": "hi"})) + + def test_chat_and_responses_preserve_keys_through_projection_and_routing(self): + for profile in fixtures.PROFILES: + self.fx.configure(profiles=(profile,)) + for route in ("chat/completions", "responses"): + for stream in (False, True): + with self.subTest(profile=profile, route=route, stream=stream): + payload = self.fx.payload(route, stream=stream) + payload["prompt_cache_key"] = "synthetic-cache" + response = self.fx.client.post("/v1/" + route, json=payload) + self.assertEqual(response.status_code, 200, response.text) + self.assertEqual(json.loads(self.fx.requests[-1].content)["prompt_cache_key"], "synthetic-cache") + + +class PreparationResponsivenessTests(GatewayFixture, unittest.IsolatedAsyncioTestCase): + async def test_credential_lock_during_preparation_does_not_block_event_loop(self): + self.fx.configure(profiles=("cn-cli",)) + cm = self.fx.pool.first() + async with REAL_ASYNC_CLIENT(transport=httpx.ASGITransport(app=gateway.app), base_url="http://test") as client: + for route in fixtures.GENERATIONS: + for stream in (False, True): + with self.subTest(route=route, stream=stream): + owned, release, expired = threading.Event(), threading.Event(), threading.Event() + + def hold_refresh_lock(): + with cm._lock: + owned.set() + if not release.wait(2): + expired.set() + + thread = threading.Thread(target=hold_refresh_lock) + thread.start() + try: + self.assertTrue(await asyncio.to_thread(owned.wait, 1)) + heartbeat = asyncio.get_running_loop().call_later(0.05, release.set) + try: + response = await asyncio.wait_for(client.post("/v1/" + route, + json=self.fx.payload(route, stream=stream)), 4) + finally: + heartbeat.cancel() + self.assertFalse(expired.is_set(), "The event loop could not release the credential lock") + self.assertEqual(response.status_code, 200, response.text) + finally: + release.set() + await asyncio.to_thread(thread.join, 2) + self.assertFalse(thread.is_alive()) + + +if __name__ == "__main__": + unittest.main()