From 0e2c246c4daeec05bc5a9442c47617ac7504e7fe Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A2=A8=E8=8F=8A?= <277378677+maiphucgiang@users.noreply.github.com> Date: Thu, 17 Sep 2026 01:18:59 +0800 Subject: [PATCH] Add opt-in scoped sessions and request-attempt tracing Keep legacy upstream identity behavior by default. Scoped mode fingerprints protocol-adapted input, validates explicit session hints, isolates account conversations and assigns distinct attempts under one logical request. Correlate response headers, text logs and bounded audits without trusting client IDs or changing replay rules. Verify 52 backend scripts, isolated deployment and the actual local gateway against zero-rate hy3; restore original runtime settings. Keep version 1.2.6. --- .env.example | 2 + app/audit_store.py | 6 +- app/inference_resources.py | 8 +- app/observability.py | 21 +- app/request_context.py | 173 ++++++++++++ app/runtime_management.py | 2 + app/settings.py | 2 + app/upstream_io.py | 5 +- converter.py | 70 ++++- docker-compose.yml | 1 + docs/advanced.md | 12 + docs/advanced.zh-CN.md | 12 + tests/test_environment_config.py | 19 +- tests/test_request_context.py | 441 +++++++++++++++++++++++++++++++ 14 files changed, 750 insertions(+), 24 deletions(-) create mode 100644 app/request_context.py create mode 100644 tests/test_request_context.py diff --git a/.env.example b/.env.example index 4f89cd9..0596f8f 100644 --- a/.env.example +++ b/.env.example @@ -39,6 +39,8 @@ CODEBUDDY2API_MAX_CONCURRENT=64 # CODEBUDDY2API_UPSTREAM_KEEPALIVE=false # Per-account in-flight limit; zero preserves unlimited account capacity. # CODEBUDDY2API_MAX_INFLIGHT_PER_ACCOUNT=0 +# Optional session/attempt tracing; legacy preserves existing upstream header behavior. +# CODEBUDDY2API_REQUEST_CONTEXT_MODE=legacy # SQLite audit defaults to the data directory; this optional path enables a separate text log. CODEBUDDY2API_LOG= diff --git a/app/audit_store.py b/app/audit_store.py index f6b3b6d..d6038f4 100644 --- a/app/audit_store.py +++ b/app/audit_store.py @@ -47,11 +47,13 @@ def safe_attempt(value: Any) -> dict: return {} result = {} for key in ("stage", "code", "error_code", "model", "upstream_model", "profile", - "credential", "usage_source", "outcome", "consent_source", "agreement_revision", "conversation_id", "request_id"): + "credential", "usage_source", "outcome", "consent_source", "agreement_revision", "conversation_id", "request_id", + "attempt_id", "upstream_request_id", "session_source"): clean = safe_label(value.get(key)) if clean is not None: result[key] = clean - for key in ("status_code", "duration_ms", "attempt", "retry_after", "max_attempts", "total_tokens"): + for key in ("status_code", "duration_ms", "attempt", "retry_after", "max_attempts", "total_tokens", + "attempt_index", "dropped"): clean = number(value.get(key)) if clean is not None: result[key] = clean diff --git a/app/inference_resources.py b/app/inference_resources.py index 40d780a..b59f14f 100644 --- a/app/inference_resources.py +++ b/app/inference_resources.py @@ -10,6 +10,7 @@ import httpx from app.site_routing import PROFILE_ENDPOINTS +from app.request_context import with_request_context class _RejectCookies(DefaultCookiePolicy): @@ -140,10 +141,13 @@ def close(self): class InferenceResourcesMiddleware: - def __init__(self, app): - self.app = app + def __init__(self, app, config=None): + self.app, self.config = app, config async def __call__(self, scope, receive, send): + await with_request_context(self._serve, scope, receive, send, self.config) + + async def _serve(self, scope, receive, send): if (scope["type"] != "http" or scope.get("method") != "POST" or scope.get("path") not in ("/v1/chat/completions", "/v1/responses", "/v1/messages")): return await self.app(scope, receive, send) diff --git a/app/observability.py b/app/observability.py index d5e71b9..b0e75a2 100644 --- a/app/observability.py +++ b/app/observability.py @@ -8,10 +8,10 @@ from dataclasses import dataclass, field import json import time -import uuid from typing import Any from app.audit_store import AuditStore, METRICS, number, safe_attempt, safe_label +from app.request_context import current_context, ensure_context _PATHS = {"/v1/chat/completions": "chat", "/v1/responses": "responses", "/v1/messages": "messages"} _PARSE_LIMIT = 16384 @@ -148,8 +148,17 @@ def observe_usage(usage): def observe_attempt(stage, **safe_metadata): observation = _current.get() - if observation is not None and len(observation.attempts) < 32: - observation.attempts.append(safe_attempt({**safe_metadata, "stage": stage})) + if observation is not None: + context = current_context() + metadata = context.attempt_metadata() if context is not None else {} + entry = safe_attempt({**metadata, **safe_metadata, "stage": stage}) + if len(observation.attempts) < 32: + observation.attempts.append(entry) + else: + # Reserve the last slot for a bounded overflow marker, not an unbounded trace. + previous = observation.attempts[-1] + dropped = previous.get("dropped", 0) + 1 if previous.get("stage") == "attempts_truncated" else 2 + observation.attempts[-1] = safe_attempt({"stage": "attempts_truncated", "dropped": dropped}) def observe_failure(code): @@ -177,8 +186,7 @@ def observe_recovery(through=None): code = observation.record.get("error_code") or "upstream_error" observation.failed = False observation.record["error_code"] = None - if len(observation.attempts) < 32: - observation.attempts.append(safe_attempt({"stage": "failover_recovered", "code": code})) + observe_attempt("failover_recovered", code=code) class _Parser: @@ -265,6 +273,7 @@ def feed(self, body, final=False): class AuditMiddleware: def __init__(self, app, config): self.app = app + self.config = config if isinstance(config, AuditStore) or callable(getattr(config, "record_request", None)): self.store = config elif isinstance(config, dict): @@ -301,7 +310,7 @@ async def lifespan_receive(): await self.app(scope, receive, send) return started, monotonic_start = time.time(), time.monotonic() - observation = _Observation({"id": uuid.uuid4().hex, "epoch": -1, + observation = _Observation({"id": ensure_context(scope, self.config).request_id, "epoch": -1, "detail_generation": -1, "started_at": started, "protocol": _PATHS[scope["path"]], **{key: None for key in METRICS}}, monotonic_start) diff --git a/app/request_context.py b/app/request_context.py new file mode 100644 index 0000000..2213a23 --- /dev/null +++ b/app/request_context.py @@ -0,0 +1,173 @@ +"""Separate local request identity, optional session hints and individual upstream attempts.""" +from contextvars import ContextVar +from dataclasses import dataclass +import hashlib +import json +import secrets +import threading +import uuid + + +PATHS = {"/v1/chat/completions": "chat", "/v1/responses": "responses", "/v1/messages": "messages"} +_SCOPE_KEY = "codebuddy.request_context" +_current = ContextVar("request_context", default=None) + + +class SessionIdentifierError(ValueError): + """Reject ambiguous or unsafe session hints without echoing their contents.""" + + +def _identifier(value): + if value is None or value == "": + return None + if not isinstance(value, str) or len(value) > 512 or not value.isprintable(): + raise SessionIdentifierError("session ID must be a printable string of at most 512 UTF-8 bytes") + try: + if len(value.encode("utf-8")) > 512: + raise SessionIdentifierError("session ID exceeds 512 UTF-8 bytes") + except UnicodeError: + raise SessionIdentifierError("session ID contains invalid Unicode") from None + return value.strip() or None + + +def _digest(value): + digest = hashlib.sha256() + for part in json.JSONEncoder(ensure_ascii=False, sort_keys=True, separators=(",", ":"), allow_nan=False).iterencode(value): + for offset in range(0, len(part), 65536): + digest.update(part[offset:offset + 65536].encode("utf-8")) + return digest.hexdigest() + + +def _content_parts(content): + if isinstance(content, str): + return [{"type": "text", "text": content}] if content else [] + if not isinstance(content, list): + return [] + parts = [] + for block in content: + if not isinstance(block, dict): + continue + if block.get("type") in ("text", "input_text", "output_text"): + if block.get("text"): + parts.append({"type": "text", "text": block["text"]}) + else: + parts.append({key: value for key, value in block.items() if key != "cache_control"}) + return parts + + +@dataclass(frozen=True) +class Attempt: + id: str + index: int + span: str + + +class RequestContext: + def __init__(self, protocol, mode="legacy", headers=()): + self.request_id = uuid.uuid4().hex + self.protocol = protocol + self.mode = mode + self._session_headers = tuple(value.decode("latin-1") for key, value in headers + if key.lower() == b"x-codebuddy-session-id") if self.scoped else () + self.session_key = "scoped:v1:temporary:" + self.request_id + self.session_source = "temporary" if self.scoped else "legacy" + self._bound = False + self._root_span = secrets.token_hex(8) + self._lock = threading.Lock() + self._attempt_index = 0 + self.attempt = None + + @property + def scoped(self): + return self.mode == "scoped" + + def bind_session(self, payload, messages): + """Fingerprint adapted messages before projection or desensitization, once per request.""" + if not self.scoped or self._bound: + return + hints = [("explicit_header", value) for value in self._session_headers] + metadata = payload.get("metadata") + for label, source in (("explicit_metadata", metadata), ("explicit_body", payload)): + if isinstance(source, dict): + hints.extend((label, source[key]) for key in ("conversation_id", "conversationId") if key in source) + clean = [(label, _identifier(value)) for label, value in hints] + clean = [(label, value) for label, value in clean if value is not None] + if len({value for _, value in clean}) > 1: + raise SessionIdentifierError("conflicting session IDs") + if clean: + self.session_key = "scoped:v1:" + _digest([self.protocol, "explicit", clean[0][1]]) + self.session_source = clean[0][0] + elif isinstance(messages, list): + instructions = [_content_parts(m.get("content")) for m in messages + if isinstance(m, dict) and m.get("role") in ("system", "developer")] + first = next((m for m in messages if isinstance(m, dict) and m.get("role") == "user"), {}) + content = _content_parts(first.get("content")) + if content: + self.session_key = "scoped:v1:" + _digest([self.protocol, "content", instructions, content]) + self.session_source = "fingerprint" + self._session_headers = () + self._bound = True + + def conversation_id(self, profile, account): + return str(uuid.UUID(hex=_digest([profile, account, self.session_key])[:32])) + + def start_attempt(self): + with self._lock: + self._attempt_index += 1 + self.attempt = Attempt(uuid.uuid4().hex, self._attempt_index, secrets.token_hex(8)) + return self.attempt + + def attempt_headers(self, headers, attempt): + if not self.scoped: + return dict(headers) + root, span = self.request_id, attempt.span + return {**headers, "X-Request-ID": attempt.id, "X-Conversation-Message-ID": attempt.id, + "X-Conversation-Request-ID": root, "X-Root-Request-ID": root, "X-Trace-ID": root, + "traceparent": f"00-{root}-{span}-01", "b3": f"{root}-{span}-1-{self._root_span}", + "X-B3-TraceId": root, "X-B3-SpanId": span, "X-B3-ParentSpanId": self._root_span, + "X-B3-Sampled": "1"} + + def attempt_metadata(self): + data = {"request_id": self.request_id, "session_source": self.session_source} + if self.attempt is not None: + data.update(attempt_id=self.attempt.id, attempt_index=self.attempt.index) + return data + + +def current_context(): + return _current.get() + + +def ensure_context(scope, config=None): + """Share one identity through middleware regardless of audit availability or ordering.""" + if _SCOPE_KEY not in scope: + values = config() if callable(config) else config + mode = values.get("request_context_mode", "legacy") if isinstance(values, dict) else "legacy" + scope[_SCOPE_KEY] = RequestContext(PATHS[scope["path"]], mode, scope.get("headers", ())) + return scope[_SCOPE_KEY] + + +async def with_request_context(app, scope, receive, send, config=None): + if scope.get("type") != "http" or scope.get("method") != "POST" or scope.get("path") not in PATHS: + return await app(scope, receive, send) + context = ensure_context(scope, config) + if current_context() is context: + return await app(scope, receive, send) + token = _current.set(context) + async def identified_send(message): + if message["type"] == "http.response.start": + headers = [(key, value) for key, value in message.get("headers", ()) if key.lower() != b"x-request-id"] + message = {**message, "headers": [*headers, (b"x-request-id", context.request_id.encode("ascii"))]} + await send(message) + try: + await app(scope, receive, identified_send) + finally: + _current.reset(token) + + +class RequestContextMiddleware: + def __init__(self, app, config): + self.app, self.config = app, config + + async def __call__(self, scope, receive, send): + await with_request_context(self.app, scope, receive, send, self.config) diff --git a/app/runtime_management.py b/app/runtime_management.py index 0eb6b2e..460de2a 100644 --- a/app/runtime_management.py +++ b/app/runtime_management.py @@ -104,6 +104,8 @@ def install(gateway): app.add_middleware(ConcurrencyLimitMiddleware, config=config) # Authenticate headers before consuming inference capacity or buffering request bodies. app.add_middleware(InferenceAuthMiddleware, config=config) + from .request_context import RequestContextMiddleware + app.add_middleware(RequestContextMiddleware, config=config) install_pages(app, Path(gateway.__file__).resolve().parent / "web" / "dist") diff --git a/app/settings.py b/app/settings.py index 9410ec2..51267e6 100644 --- a/app/settings.py +++ b/app/settings.py @@ -47,6 +47,8 @@ def _item(default, type_, label, *, mode="hot", env=None, minimum=None, maximum= env="CODEBUDDY2API_UPSTREAM_KEEPALIVE"), "max_inflight_per_account": _item(0, "integer", "单账号在途上限(0 不限制)", env="CODEBUDDY2API_MAX_INFLIGHT_PER_ACCOUNT", minimum=0, maximum=10000), + "request_context_mode": _item("legacy", "string", "请求上下文模式", + env="CODEBUDDY2API_REQUEST_CONTEXT_MODE", choices=["legacy", "scoped"]), "audit_max_bytes": _item(256 * 1024 * 1024, "integer", "审计明细预算", minimum=1024**2, maximum=1024**4), "audit_retention_days": _item(30, "integer", "审计明细保留天数", minimum=1, maximum=36500), "audit_diagnostic_bytes": _item(8192, "integer", "失败诊断最大字节", minimum=0, maximum=8192), diff --git a/app/upstream_io.py b/app/upstream_io.py index 226f171..a9d4c4d 100644 --- a/app/upstream_io.py +++ b/app/upstream_io.py @@ -223,7 +223,7 @@ async def _attempt_client(url, timeout, clients): @asynccontextmanager async def open_backend_stream(url, headers, body, *, read_timeout=300, on_retry=None, - retry_write_timeout=False, clients=None): + retry_write_timeout=False, clients=None, headers_for_attempt=None): """Retry connection failures once on a fresh client; write timeouts require explicit opt-in. Never replay after the upstream response opens. """ @@ -233,7 +233,8 @@ async def open_backend_stream(url, headers, body, *, read_timeout=300, on_retry= opened = False try: async with _attempt_client(url, timeout, clients if attempt == 0 else None) as client: - async with client.stream("POST", url, headers=headers, json=body, timeout=timeout) as response: + attempt_headers = headers_for_attempt() if headers_for_attempt is not None else headers + async with client.stream("POST", url, headers=attempt_headers, json=body, timeout=timeout) as response: opened = True yield response return diff --git a/converter.py b/converter.py index 8e77311..11689cd 100644 --- a/converter.py +++ b/converter.py @@ -61,6 +61,7 @@ def desensitize_body(body, roles=("system",), desensitize_harness_user=False, open_backend_stream, parse_retry_after, read_bounded_error) from app.inference_resources import (AccountCapacity, InferenceResourcesMiddleware, inference_lifespan, request_resources, release_credential) +from app.request_context import SessionIdentifierError, current_context 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 @@ -1435,7 +1436,7 @@ def _housekeeper_loop(pool: CredentialPool, ledger) -> None: # --------------------------------------------------------------------------- app = FastAPI(title="codebuddy2api", version=APP_VERSION, lifespan=inference_lifespan) -app.add_middleware(InferenceResourcesMiddleware) +app.add_middleware(InferenceResourcesMiddleware, config=lambda: CONFIG) # Anthropic error types: https://platform.claude.com/docs/en/api/errors _ANTHROPIC_ERROR_TYPES = { @@ -1485,6 +1486,7 @@ async def _protocol_http_exception(request: Request, exc: HTTPException): "max_inbound_bytes": 64 * 1024 * 1024, "max_collect_bytes": 8 * 1024 * 1024, "max_concurrent": 64, "upstream_keepalive": False, "max_inflight_per_account": 0, + "request_context_mode": "legacy", "failover_max": 0, # Credential failovers allowed before the first response byte "retry_write_timeout": False, # Opt-in replay after incomplete writes "usage_daily": None, # Usage aggregated by date and model @@ -1571,7 +1573,8 @@ def _check_admin_auth(authorization: Optional[str], x_api_key: Optional[str]): def _cred_for(payload: dict, model: str | None = None, *, region=None, tried=()): """Select a fresh credential lease and headers, excluding tried accounts; report unavailable capacity.""" - raw_key = session_key(payload) + context = current_context() + raw_key = context.session_key if context is not None and context.scoped else session_key(payload) skey = f"{region}:{raw_key}" if raw_key and region is not None else raw_key skey = model_policy.sticky_scope(CONFIG, skey, model) pool = CONFIG.get("cred_pool") @@ -1613,7 +1616,11 @@ def _cred_for(payload: dict, model: str | None = None, *, region=None, tried=()) profile = profile_for_headers(headers) if not _in_region(profile, region): raise HTTPException(status_code=503, detail={"error": {"message": "未找到指定地域凭据", "type": "auth_error"}}) - headers.update(_dynamic_request_headers(f"{profile}:{skey}" if skey else None)) + if context is not None and context.scoped: + identity = account_key(profile, headers.get("X-User-Id"), headers.get("X-Enterprise-Id")) + headers["X-Conversation-ID"] = context.conversation_id(profile, identity) + else: + headers.update(_dynamic_request_headers(f"{profile}:{skey}" if skey else None)) return cm, headers @@ -2302,8 +2309,28 @@ def _normalize_tool_choice(body): body["tools"], body["tool_choice"] = matches, "required" -def _prepare_chat_body(body: dict, *, region=None) -> dict: +def _bind_request_session(payload, body): + context = current_context() + if context is not None and context.scoped: + try: + context.bind_session(payload, body.get("messages")) + except SessionIdentifierError as error: + raise HTTPException(status_code=400, detail={"error": {"message": str(error), + "type": "invalid_request_error", "param": "session_id"}}) from None + except (ValueError, TypeError, UnicodeError, RecursionError): + raise HTTPException(status_code=400, detail={"error": {"message": "invalid session input", + "type": "invalid_request_error"}}) from None + + +def _request_id(): + context = current_context() + return context.request_id if context is not None else uuid.uuid4().hex + + +def _prepare_chat_body(body: dict, *, region=None, session_payload=None) -> dict: """Normalize models, system messages, streaming, desensitization and payload budgets.""" + if session_payload is not None: + _bind_request_session(session_payload, body) body = dict(body) body["model"] = model_policy.resolve(CONFIG, body.get("model", "auto")) guard_model(body["model"], region=region, resolved=True) @@ -2410,14 +2437,14 @@ 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 = await run_in_threadpool(_prepare_chat_body, body) + body = await run_in_threadpool(_prepare_chat_body, body, session_payload=payload) # Record request metadata. model_name = payload.get("model", "auto") tool_names = [t.get("function", {}).get("name") for t in (payload.get("tools") or []) if isinstance(t, dict)] last_user = _last_user_text(messages) - rid = os.urandom(4).hex() + rid = _request_id() _log(f"[{rid}] ▶ REQUEST {model_name} | stream={client_wants_stream} | msgs={len(messages)}" + (f" | tools={tool_names}" if tool_names else "") + (f" | last_user={_truncate(last_user, 60)!r}" if last_user else "")) @@ -2619,6 +2646,22 @@ def _public_sse_line(line, model_name): @asynccontextmanager async def _backend_stream(url, headers, body, *, timeout=300, rid="", model_name="?"): started, opened = time.monotonic(), False + context = current_context() + if context is not None: + context.attempt = None + + def attempt_headers(): + if context is None: + return dict(headers) + attempt = context.start_attempt() + outgoing = context.attempt_headers(headers, attempt) + profile = profile_for_headers(headers) + observe_attempt("upstream_attempt", profile=profile, upstream_model=body.get("model"), + credential=account_key(profile, headers.get("X-User-Id"), headers.get("X-Enterprise-Id")), + conversation_id=outgoing.get("X-Conversation-ID"), + upstream_request_id=outgoing.get("X-Request-ID")) + return outgoing + def retry(error): """Record connection retries and flag possible billing after write timeouts.""" timeout_on_write = isinstance(error, WRITE_TIMEOUT_TRANSPORT) @@ -2632,7 +2675,7 @@ def retry(error): clients = resources.clients if resources is not None and CONFIG.get("upstream_keepalive") else None async with open_backend_stream(url, headers, body, read_timeout=timeout, on_retry=retry, retry_write_timeout=bool(CONFIG.get("retry_write_timeout")), - clients=clients) as response: + clients=clients, headers_for_attempt=attempt_headers) as response: opened = True observe_attempt("upstream_http", status_code=response.status_code, duration_ms=(time.monotonic() - started) * 1000) @@ -3092,13 +3135,14 @@ async def create_response(request: Request, except Exception as e: raise HTTPException(status_code=400, detail={"error": {"message": f"request conversion error: {e}", "type": "invalid_request_error"}}) + await run_in_threadpool(_bind_request_session, payload, chat_body) chat_body, projection_stats = project_responses_chat_body( chat_body, keep_tool_metadata=CONFIG.get("keep_tool_metadata", False)) chat_body = await run_in_threadpool(_prepare_chat_body, chat_body) client_wants_stream = _client_wants_stream(payload) model_name = payload.get("model", "auto") - rid = os.urandom(4).hex() + rid = _request_id() _log(f"[{rid}] ▶ RESPONSES {model_name} | stream={client_wants_stream} | input_items={len(payload.get('input', []))}") _log( f"[{rid}] ── RESPONSES PROJECTION ── " @@ -3212,10 +3256,10 @@ 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 = await run_in_threadpool(_prepare_chat_body, chat_body) + chat_body = await run_in_threadpool(_prepare_chat_body, chat_body, session_payload=payload) model_name = payload.get("model", "auto") chat_messages = chat_body.get("messages", []) - rid = os.urandom(4).hex() + rid = _request_id() _log(f"[{rid}] ▶ ANTHROPIC {model_name} | msgs={len(chat_messages)} | anthropic_msgs={len(messages)}") # Keep blocking credential selection and refresh off the event loop. prepared = chat_body # Preserve canonical input for routing policy checks. @@ -3440,6 +3484,9 @@ def main(): ap.add_argument("--upstream-keepalive", type=_boolean_arg, nargs="?", const=True, default=os.environ.get("CODEBUDDY2API_UPSTREAM_KEEPALIVE", "false"), help="按上游入口复用有界连接池,默认 false;重启生效,不改变超时或重放规则") + ap.add_argument("--request-context-mode", choices=("legacy", "scoped"), + default=os.environ.get("CODEBUDDY2API_REQUEST_CONTEXT_MODE", "legacy"), + help="请求上下文:legacy 保持旧会话头,scoped 启用显式会话与逐尝试追踪;默认 legacy") ap.add_argument("--log-body-limit", type=_nonnegative_int, metavar="BYTES", default=os.environ.get("CODEBUDDY2API_LOG_BODY_LIMIT", "65536"), help="每条正文日志的预览字节上限,默认 64 KiB;0 只记录摘要") @@ -3469,7 +3516,8 @@ def main(): for key in ("max_images", "image_policy", "max_request_bytes", "log_body_limit", "tool_call_max_retry", "max_inbound_bytes", "max_collect_bytes", "max_concurrent", - "failover_max", "retry_write_timeout", "upstream_keepalive", "max_inflight_per_account"): + "failover_max", "retry_write_timeout", "upstream_keepalive", "max_inflight_per_account", + "request_context_mode"): CONFIG[key] = getattr(args, key) CONFIG["api_key"] = args.api_key CONFIG["desensitize"] = args.desensitize diff --git a/docker-compose.yml b/docker-compose.yml index 7913d8e..350d070 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -29,6 +29,7 @@ services: CODEBUDDY2API_MAX_CONCURRENT: ${CODEBUDDY2API_MAX_CONCURRENT:-64} CODEBUDDY2API_UPSTREAM_KEEPALIVE: CODEBUDDY2API_MAX_INFLIGHT_PER_ACCOUNT: + CODEBUDDY2API_REQUEST_CONTEXT_MODE: CODEBUDDY2API_TOOL_CALL_MAX_RETRY: ${CODEBUDDY2API_TOOL_CALL_MAX_RETRY:-3} CODEBUDDY2API_FAILOVER_MAX: CODEBUDDY2API_RETRY_WRITE_TIMEOUT: diff --git a/docs/advanced.md b/docs/advanced.md index f47d661..e113989 100644 --- a/docs/advanced.md +++ b/docs/advanced.md @@ -34,6 +34,7 @@ Compose explicitly passes some environment variables and CLI flags, so deleting | `--max-concurrent` | `64` | Concurrency limit for the three generation endpoints only; excess requests get 503 with Retry-After; token counting is unaffected; `0` disables | | `--max-inflight-per-account` | `0` | Per-process, per-account in-flight client inference limit; `0` disables, full accounts return 503 | | `--upstream-keepalive [true/false]` | `false` | Bounded connection reuse isolated by official origin; requires restart | +| `--request-context-mode` | `legacy` | `scoped` enables explicit sessions and per-attempt tracing; changes apply to new requests | | `--failover-max` | `0` | Extra credentials tried when a request fails before the first response byte reaches the client; `0` keeps the upstream behaviour of surfacing the failure directly | | `--retry-write-timeout` | `false` | Opt a request-body write timeout into replay (fresh connection and `--failover-max`), accepting that bytes already sent may have been processed | | `--max-request-bytes` | `33554432` | Positive byte limit for the processed upstream JSON | @@ -72,6 +73,17 @@ The account limit defaults to `0`. Positive limits skip full accounts within exi Before downgrading the source, remove the new CLI arguments and restore a control-store backup without these keys; disabling the features does not remove persisted settings. +### Request context + +Generation responses carry a gateway-generated `X-Request-ID` for correlation with text logs and available audit details. Response-body and tool-call IDs are unchanged; client request IDs are not trusted or used for deduplication. + +`request_context_mode` defaults to `legacy`, preserving existing session keys and upstream headers. Enable `scoped` in WebUI settings, with `--request-context-mode scoped`, or through `CODEBUDDY2API_REQUEST_CONTEXT_MODE=scoped`. Each client HTTP request keeps one root ID across existing retries/failover, with a new ID/span for every upstream attempt and account-isolated conversation IDs. This does not enable additional retries, server-side history or automatic cache keys. + +In scoped mode, optionally send `X-Codebuddy-Session-ID`, `metadata.conversation_id` / `metadata.conversationId`, or top-level `conversation_id` / `conversationId`. Values must agree; conflicting, non-string, control-character or over-512-UTF-8-byte values return 400. Empty values fall back to a fingerprint of adapted instructions and the first user input, including image references; URLs are not fetched. Without reliable input a temporary session is used. Identical inputs without explicit IDs remain indistinguishable; `user`, `metadata.user_id` and `prompt_cache_key` are not session IDs. Raw hints are neither logged nor forwarded upstream. + +Switch back to `legacy` to restore old behavior for new requests; in-flight requests retain their initial mode. Before source downgrade, remove the new startup option and restore a compatible control-store backup without `request_context_mode`. + + ## APIs and authentication | Client endpoint | Description | diff --git a/docs/advanced.zh-CN.md b/docs/advanced.zh-CN.md index 4be22de..2f1d358 100644 --- a/docs/advanced.zh-CN.md +++ b/docs/advanced.zh-CN.md @@ -34,6 +34,7 @@ Compose 会显式传入部分环境变量及 CLI 参数,删除 `.env` 中的 | `--max-concurrent` | `64` | 仅限制三个生成端点;占满立即 503(含 Retry-After),不限制 token 估算;`0` 不限制 | | `--max-inflight-per-account` | `0` | 每进程、每账号的客户端推理在途上限;`0` 不限制,满载立即 503 | | `--upstream-keepalive [true/false]` | `false` | 启用按官方入口隔离的有界连接复用;重启生效 | +| `--request-context-mode` | `legacy` | `scoped` 启用显式会话与逐尝试追踪;变更只影响新请求 | | `--failover-max` | `0` | 请求在「一个字节都还没发给下游」之前失败时,最多再换几个凭证就地重放;`0` 表示如实把失败回给下游 | | `--retry-write-timeout` | `false` | 让「写请求体超时」也参与重放(换新连接与 `--failover-max` 换凭证),代价是已发出的那半截正文可能已被上游处理 | | `--max-request-bytes` | `33554432` | 处理后的上游 JSON 字节上限,须为正整数 | @@ -72,6 +73,17 @@ WebUI 系统设置可配置这两项;环境变量为 `CODEBUDDY2API_UPSTREAM_K 源码降级前还需移除新增启动参数,并恢复不含这两个配置键的控制库备份;仅关闭开关不会删除持久化配置。 +### 请求上下文 + +生成接口响应带网关生成的 `X-Request-ID`,可关联文本日志及可用的审计明细;正文响应 ID、工具调用 ID 不变,不采用客户端请求 ID 进行鉴权或去重。 + +`request_context_mode` 默认 `legacy`,保留旧会话键和上游头。通过 WebUI、`--request-context-mode scoped` 或 `CODEBUDDY2API_REQUEST_CONTEXT_MODE=scoped` 启用新模式:每次客户端 HTTP 请求的根 ID 在既有重试/换号中保持不变,每次上游尝试生成独立 ID/span,会话 ID 按账号隔离。不增加重试,不保存服务端历史,不自动生成缓存键。 + +scoped 模式可选传入 `X-Codebuddy-Session-ID`、`metadata.conversation_id` / `metadata.conversationId` 或顶层 `conversation_id` / `conversationId`。多处值须一致;冲突、非字符串、控制字符或超过 512 UTF-8 字节时返回 400。空值回退为协议适配后的指令与首条用户输入指纹,包含图片引用但不抓取 URL;缺少可靠输入时使用临时会话。无显式 ID 的相同输入仍无法区分;`user`、`metadata.user_id`、`prompt_cache_key` 不是会话 ID。原始标识不记录、不转发上游。 + +设回 `legacy` 即恢复新请求的旧行为,在途请求保留入口模式;源码降级前移除新增启动参数,并恢复不含 `request_context_mode` 的兼容控制库备份。 + + ## API 与鉴权 | 客户端接口 | 说明 | diff --git a/tests/test_environment_config.py b/tests/test_environment_config.py index 5fbb2a6..61ae98e 100644 --- a/tests/test_environment_config.py +++ b/tests/test_environment_config.py @@ -118,6 +118,21 @@ def test_invalid_pooling_environment_fails_before_startup(self): self.start(env) + def test_request_context_mode_precedence_and_validation(self): + _, items, config = self.start(saved={'request_context_mode': 'scoped'}) + self.assertEqual(config['request_context_mode'], 'scoped') + self.assertEqual(items['request_context_mode']['source'], 'management') + _, items, config = self.start({'CODEBUDDY2API_REQUEST_CONTEXT_MODE': 'scoped'}) + self.assertEqual(config['request_context_mode'], 'scoped') + self.assertTrue(items['request_context_mode']['locked']) + _, items, config = self.start({'CODEBUDDY2API_REQUEST_CONTEXT_MODE': 'invalid'}, + cli=('--request-context-mode=legacy',), saved={'request_context_mode': 'scoped'}) + self.assertEqual(config['request_context_mode'], 'legacy') + self.assertEqual(items['request_context_mode']['source'], 'cli') + with self.assertRaises(ValueError): + self.start({'CODEBUDDY2API_REQUEST_CONTEXT_MODE': 'invalid'}) + + def test_example_covers_all_active_runtime_environment_names(self): example = (ROOT / '.env.example').read_text() documented = set(re.findall(r'(?m)^(?:# )?(CODEBUDDY[A-Z0-9_]+)=', example)) @@ -151,6 +166,7 @@ def test_compose_forwards_dotenv_limits_and_retries_without_changing_internal_bi 'CODEBUDDY2API_MAX_CONCURRENT': '2', 'CODEBUDDY2API_TOOL_CALL_MAX_RETRY': '1', 'CODEBUDDY2API_FAILOVER_MAX': '1', 'CODEBUDDY2API_RETRY_WRITE_TIMEOUT': 'true', 'CODEBUDDY2API_UPSTREAM_KEEPALIVE': 'true', 'CODEBUDDY2API_MAX_INFLIGHT_PER_ACCOUNT': '2', + 'CODEBUDDY2API_REQUEST_CONTEXT_MODE': 'scoped', 'CODEBUDDY2API_KEEP_TOOL_METADATA': 'false', 'CODEBUDDY_IMPORT_DIR': '/data/auth/incoming'} service = self.compose(values) port = service['ports'][0] @@ -164,7 +180,8 @@ def test_compose_forwards_dotenv_limits_and_retries_without_changing_internal_bi def test_compose_unset_optional_settings_do_not_override_webui(self): service = self.compose({}) for name in ('CODEBUDDY2API_KEEP_TOOL_METADATA', 'CODEBUDDY2API_FAILOVER_MAX', 'CODEBUDDY2API_RETRY_WRITE_TIMEOUT', - 'CODEBUDDY2API_UPSTREAM_KEEPALIVE', 'CODEBUDDY2API_MAX_INFLIGHT_PER_ACCOUNT'): + 'CODEBUDDY2API_UPSTREAM_KEEPALIVE', 'CODEBUDDY2API_MAX_INFLIGHT_PER_ACCOUNT', + 'CODEBUDDY2API_REQUEST_CONTEXT_MODE'): self.assertIsNone(service['environment'].get(name)) self.assertEqual(service['ports'][0]['host_ip'], '127.0.0.1') self.assertEqual(service['environment']['CODEBUDDY_IMPORT_DIR'], '/data/auth/imports') diff --git a/tests/test_request_context.py b/tests/test_request_context.py new file mode 100644 index 0000000..f1600dc --- /dev/null +++ b/tests/test_request_context.py @@ -0,0 +1,441 @@ +"""Verify request/session isolation and tracing with synthetic credentials and offline upstreams.""" +import asyncio +from concurrent.futures import ThreadPoolExecutor +from copy import deepcopy +import json +from pathlib import Path +import sys +import unittest +from unittest.mock import AsyncMock, patch +from types import SimpleNamespace + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +import httpx +from fastapi.testclient import TestClient +import converter as gateway +from app.audit_store import AuditStore +from app.observability import AuditMiddleware, observe_attempt +from app.request_context import (RequestContext, RequestContextMiddleware, SessionIdentifierError, + current_context, ensure_context) +from app.adapters.responses_adapter import responses_request_to_chat +from app.adapters.anthropic_adapter import anthropic_request_to_chat +import test_api_flow as fixtures +import test_workbuddy_filter as filters + +REAL_CLIENT = fixtures.REAL_ASYNC_CLIENT +USER = {"role": "user", "content": "first question"} + + +def context(payload, *, protocol="chat", headers=(), messages=None): + ctx = RequestContext(protocol, "scoped", headers) + ctx.bind_session(payload, payload.get("messages", [USER]) if messages is None else messages) + return ctx + + +class SessionTests(unittest.TestCase): + def test_explicit_aliases_and_header_share_a_key_without_disclosing_the_value(self): + value = "synthetic-private-conversation" + cases = [{"metadata": {key: value}} for key in ("conversation_id", "conversationId")] + cases += [{key: value} for key in ("conversation_id", "conversationId")] + keys = [context(body).session_key for body in cases] + keys.append(context({}, headers=[(b"X-Codebuddy-Session-ID", value.encode())]).session_key) + self.assertEqual(len(set(keys)), 1) + self.assertNotIn(value, keys[0]) + self.assertNotEqual(keys[0], context({"conversation_id": "different"}).session_key) + self.assertNotEqual(keys[0], context(cases[0], protocol="responses").session_key) + + def test_conflicting_and_malformed_identifiers_are_rejected_without_echo(self): + with self.assertRaises(SessionIdentifierError): + context({"metadata": {"conversation_id": "one"}, "conversationId": "two"}) + with self.assertRaises(SessionIdentifierError): + context({}, headers=[(b"x-codebuddy-session-id", b"one"), (b"x-codebuddy-session-id", b"two")]) + for value in (True, 1, [], {}, "canary\nprivate", "x" * 513, "é" * 257, "\ud800"): + with self.subTest(value=repr(value)[:20]), self.assertRaises(SessionIdentifierError) as error: + context({"conversation_id": value}) + self.assertNotIn("canary", str(error.exception)) + a = context({"conversation_id": " same ", "metadata": {"conversationId": "same"}}) + self.assertEqual(a.session_key, context({"conversation_id": "same"}).session_key) + self.assertEqual(context({"conversation_id": None}).session_key, context({}).session_key) + + def test_user_identity_and_cache_key_do_not_become_session_identity(self): + baseline = context({}).session_key + for payload in ({"user": "different"}, {"metadata": {"user_id": "different"}}, + {"prompt_cache_key": "different"}): + self.assertEqual(context(payload).session_key, baseline) + + def test_adapted_instructions_participate_for_each_protocol(self): + for protocol, adapt, field in (("responses", responses_request_to_chat, "instructions"), + ("messages", anthropic_request_to_chat, "system")): + keys = [] + for value in ("A", "B"): + body = {"input": "first question", "messages": [USER], field: value} + keys.append(context(body, protocol=protocol, messages=adapt(body)["messages"]).session_key) + self.assertNotEqual(*keys) + first = context({"messages": [{"role": "system", "content": "A"}, USER]}) + second = context({"messages": [{"role": "developer", "content": [{"type": "text", "text": "A"}]}, USER]}) + self.assertEqual(first.session_key, second.session_key) + + def test_followup_turns_preserve_key_but_instruction_boundaries_do_not_collapse(self): + first = {"messages": [{"role": "system", "content": "A"}, USER]} + continued = deepcopy(first) + continued["messages"] += [{"role": "assistant", "content": "answer"}, {"role": "user", "content": "next"}] + self.assertEqual(context(first).session_key, context(continued).session_key) + a = [{"role": "system", "content": x} for x in ("a", "bc")] + [USER] + b = [{"role": "system", "content": x} for x in ("ab", "c")] + [USER] + self.assertNotEqual(context({}, messages=a).session_key, context({}, messages=b).session_key) + + def test_images_order_and_contents_are_hashed_without_mutation(self): + image = {"type": "image_url", "image_url": {"url": "data:image/png;base64,QQ=="}} + text = {"type": "text", "text": "inspect"} + keys = [] + for blocks in ([image], [{**image, "image_url": {"url": "data:image/png;base64,Qg=="}}], + [image, text], [text, image]): + body = {"messages": [{"role": "user", "content": blocks}]} + original = deepcopy(body) + ctx = context(body) + self.assertEqual(ctx.session_source, "fingerprint") + self.assertEqual(body, original) + keys.append(ctx.session_key) + self.assertEqual(len(set(keys)), 4) + + def test_no_reliable_input_uses_temporary_session_and_binding_happens_once(self): + a, b = context({}, messages=[]), context({}, messages=[]) + self.assertEqual(a.session_source, "temporary") + self.assertNotEqual(a.session_key, b.session_key) + original = a.session_key + a.bind_session({"conversation_id": "late"}, [USER]) + self.assertEqual(a.session_key, original) + + def test_legacy_does_not_parse_hints_and_account_conversations_are_isolated(self): + ctx = RequestContext("chat", "legacy", [(b"x-codebuddy-session-id", b"bad\nvalue")]) + ctx.bind_session({"conversation_id": []}, [USER]) + self.assertEqual(ctx.session_source, "legacy") + ctx = context({"conversation_id": "same"}) + ids = [ctx.conversation_id(profile, account) for profile in fixtures.fixtures.PROFILES for account in ("A", "B")] + self.assertEqual(len(set(ids)), 8) + + def test_attempt_ids_are_unique_and_request_mode_is_a_snapshot(self): + scope = {"path": "/v1/responses", "headers": [(b"x-request-id", b"untrusted")]} + config = {"request_context_mode": "scoped"} + ctx = ensure_context(scope, config) + config["request_context_mode"] = "legacy" + self.assertIs(ensure_context(scope, config), ctx) + self.assertTrue(ctx.scoped) + self.assertNotEqual(ctx.request_id, "untrusted") + with ThreadPoolExecutor(max_workers=4) as executor: + attempts = list(executor.map(lambda _: ctx.start_attempt(), range(32))) + self.assertEqual({a.index for a in attempts}, set(range(1, 33))) + self.assertEqual(len({a.id for a in attempts}), 32) + base = {"Authorization": "synthetic", "X-Conversation-ID": "conversation"} + for attempt in attempts: + headers = ctx.attempt_headers(base, attempt) + self.assertEqual(headers["X-Root-Request-ID"], ctx.request_id) + self.assertEqual(headers["X-Request-ID"], attempt.id) + self.assertEqual(headers["traceparent"], f"00-{ctx.request_id}-{attempt.span}-01") + self.assertEqual(headers["X-B3-SpanId"], attempt.span) + self.assertEqual(base, {"Authorization": "synthetic", "X-Conversation-ID": "conversation"}) + + +class EndpointContextTests(fixtures.GatewayFixture, unittest.TestCase): + def setUp(self): + super().setUp() + self.enterContext(patch.dict(gateway.CONFIG, {"request_context_mode": "scoped", "max_inflight_per_account": 1})) + + def assert_attempts(self, response, seen, *, same_account=True): + self.assertEqual(response.status_code, 200, response.text) + root = response.headers["X-Request-ID"] + self.assertRegex(root, r"^[0-9a-f]{32}$") + self.assertEqual({r.headers["X-Root-Request-ID"] for r in seen}, {root}) + self.assertEqual(len({r.headers["X-Request-ID"] for r in seen}), len(seen)) + self.assertEqual(len({r.headers["X-Conversation-ID"] for r in seen}), 1 if same_account else len(seen)) + for request in seen: + headers = request.headers + self.assertEqual(headers["X-Conversation-Request-ID"], root) + self.assertEqual(headers["X-Trace-ID"], root) + self.assertEqual(headers["X-Request-ID"], headers["X-Conversation-Message-ID"]) + self.assertEqual(headers["traceparent"], f'00-{root}-{headers["X-B3-SpanId"]}-01') + self.assertEqual(self.fx.pool._capacity._counts, {}) + + def test_profiles_protocols_and_modes_preserve_body_ids_and_isolate_requests(self): + for profile in fixtures.fixtures.PROFILES: + self.fx.configure(profiles=(profile,)) + for route in fixtures.fixtures.GENERATIONS: + for stream in (False, True): + with self.subTest(profile=profile, route=route, stream=stream): + payload = self.fx.payload(route, stream=stream, text="same") + payload["metadata"] = {"conversation_id": "explicit"} + first = self.fx.client.post("/v1/" + route, json=payload) + request = self.fx.requests[-1] + self.assert_attempts(first, [request]) + second = self.fx.client.post("/v1/" + route, json=payload) + self.assert_attempts(second, [self.fx.requests[-1]]) + self.assertNotEqual(first.headers["X-Request-ID"], second.headers["X-Request-ID"]) + self.assertEqual(request.headers["X-Conversation-ID"], self.fx.requests[-1].headers["X-Conversation-ID"]) + self.assertIn("synthetic-completion" if route == "chat/completions" and stream else "ok", first.text) + if not stream: + prefix = {"chat/completions": "chatcmpl-", "responses": "resp_", "messages": "msg_"}[route] + self.assertTrue(first.json()["id"].startswith(prefix)) + self.assertNotEqual(first.json()["id"], first.headers["X-Request-ID"]) + self.assertNotIn("metadata", json.loads(request.content)) + + def test_pooling_and_capacity_settings_do_not_change_context_contract(self): + for keepalive in (False, True): + for capacity in (0, 1): + for mode in ("legacy", "scoped"): + with self.subTest(keepalive=keepalive, capacity=capacity, mode=mode), patch.dict(gateway.CONFIG, { + "upstream_keepalive": keepalive, "max_inflight_per_account": capacity, "request_context_mode": mode}): + self.fx.configure(profiles=("cn-cli",)) + response = self.fx.client.post("/v1/responses", json=self.fx.payload("responses", text="same")) + request = self.fx.requests[-1] + if mode == "scoped": + self.assert_attempts(response, [request]) + else: + self.assertEqual(response.status_code, 200) + self.assertNotEqual(response.headers["X-Request-ID"], request.headers["X-Root-Request-ID"]) + self.assertEqual(self.fx.pool._capacity._counts, {}) + + def test_hint_errors_do_not_send_upstream_and_legacy_ignores_them(self): + for route in fixtures.fixtures.GENERATIONS: + body = self.fx.payload(route) + body["metadata"] = {"conversation_id": "private-hint"} + response = self.fx.client.post("/v1/" + route, json=body, + headers={"X-Codebuddy-Session-ID": "different-private-hint"}) + self.assertEqual(response.status_code, 400, response.text) + self.assertNotIn("private-hint", response.text) + self.assertIn("X-Request-ID", response.headers) + self.assertEqual(self.fx.requests, []) + with patch.dict(gateway.CONFIG, {"request_context_mode": "legacy"}): + body = self.fx.payload() + body["conversation_id"] = [] + response = self.fx.client.post("/v1/chat/completions", json=body) + self.assertEqual(response.status_code, 200, response.text) + + def test_protocol_instructions_and_explicit_ids_change_conversation(self): + self.fx.configure(profiles=("cn-cli",)) + for route in fixtures.fixtures.GENERATIONS: + for field in ("instructions", "explicit"): + ids = [] + for value in ("A", "B"): + body = self.fx.payload(route, text="same", system=value if field == "instructions" else "constant") + if field == "explicit": + body["metadata"] = {"conversation_id": value} + response = self.fx.client.post("/v1/" + route, json=body) + self.assertEqual(response.status_code, 200, response.text) + ids.append(self.fx.requests[-1].headers["X-Conversation-ID"]) + self.assertNotEqual(*ids, (route, field)) + + def test_responses_fingerprint_is_bound_before_projection(self): + self.fx.configure(profiles=("cn-cli",)) + original = gateway.project_responses_chat_body + seen = [] + def projection(body, **kwargs): + projected, stats = original(body, **kwargs) + seen.append(len(seen)) + return {**projected, "messages": [{"role": "user", "content": str(len(seen))}]}, stats + ids = [] + with patch.object(gateway, "project_responses_chat_body", side_effect=projection): + for _ in range(2): + response = self.fx.client.post("/v1/responses", json=self.fx.payload("responses", text="same", system="original")) + self.assertEqual(response.status_code, 200, response.text) + ids.append(self.fx.requests[-1].headers["X-Conversation-ID"]) + self.assertEqual(*ids) + self.assertEqual(len(seen), 2) + + def test_failover_keeps_root_and_rotates_account_conversation_and_attempt(self): + for route in fixtures.fixtures.GENERATIONS: + for stream in (False, True): + self.fx.configure(profiles=("cn-cli", "cn-work")) + seen = [] + def respond(request): + seen.append(request) + return httpx.Response(429, json={"error": {"message": "quota"}}) if len(seen) == 1 else httpx.Response(200, content=fixtures.fixtures.success_sse()) + 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(len(seen), 2) + self.assert_attempts(response, seen, same_account=False) + + def test_scoped_hints_cannot_escape_a_full_strict_binding_or_free_tier(self): + for strict in (False, True): + with self.subTest(strict=strict): + self.fx.configure(profiles=("cn-cli", "cn-work")) + if strict: + identity = self.fx.entries["cn-cli"]["account_key"] + store = SimpleNamespace(snapshot=lambda: {"revision": 1, "credentials": {}, + "models": {"shared-model": {"credential_ids": [identity]}}}) + else: + store = None + self.fx.account_catalogs({ + "cn-cli": [fixtures.fixtures.model("shared-model", "x0.00")], + "cn-work": [fixtures.fixtures.model("shared-model", "x1.00")]}) + with patch.dict(gateway.CONFIG, {"control_store": store}): + lease, _ = self.fx.pool.headers_for("held-session", "shared-model", with_capacity=True) + before = len(self.fx.requests) + try: + for route in fixtures.fixtures.GENERATIONS: + body = self.fx.payload(route) + body["metadata"] = {"conversation_id": "different-session"} + response = self.fx.client.post("/v1/" + route, json=body) + self.assertEqual(response.status_code, 503, response.text) + self.assertEqual(response.json()["error"]["code"], "credential_concurrency_limit") + self.assertEqual(len(self.fx.requests), before) + finally: + lease.release() + self.assertEqual(self.fx.pool._capacity._counts, {}) + + def test_scoped_tracing_does_not_replay_partial_upstream_failures(self): + for route in fixtures.fixtures.GENERATIONS: + seen = [] + def respond(request): + seen.append(request) + return httpx.Response(200, content=filters.sse( + filters.event({"content": "partial"}), {"error": {"code": 429, "message": "late"}})) + 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.assertIn("late", response.text) + self.assertEqual(len(seen), 1) + self.assertEqual(seen[0].headers["X-Root-Request-ID"], response.headers["X-Request-ID"]) + self.assertEqual(self.fx.pool._capacity._counts, {}) + + + def test_connect_retries_get_new_attempts_but_legacy_headers_remain_unchanged(self): + self.fx.configure(profiles=("cn-cli",)) + for mode in ("legacy", "scoped"): + seen = [] + def respond(request): + seen.append(request) + if len(seen) == 1: + raise httpx.ConnectError("synthetic connect failure") + return httpx.Response(200, content=fixtures.fixtures.success_sse()) + with self.responder(respond), patch.dict(gateway.CONFIG, {"request_context_mode": mode}), patch.object(gateway.asyncio, "sleep", new=AsyncMock()): + response = self.fx.client.post("/v1/chat/completions", json=self.fx.payload()) + self.assertEqual(len(seen), 2) + if mode == "scoped": + self.assert_attempts(response, seen) + else: + self.assertEqual(dict(seen[0].headers), dict(seen[1].headers)) + + def test_tool_repair_and_filter_fallback_get_new_attempts(self): + self.fx.configure(profiles=("cn-cli",)) + for route in filters.ROUTES: + for filtered in (False, True): + seen = [] + def respond(request): + seen.append(request) + if filtered: + return filters.reply() if len(seen) == 1 else filters.reply("ok") + call = deepcopy(filters.TOOL) + if len(seen) == 1: + call["function"]["arguments"] = "{" + return httpx.Response(200, content=filters.sse(filters.event({"tool_calls": [call]}, "tool_calls"))) + body = filters.payload(route, tools=not filtered) + body["model"] = "shared-model" + with self.responder(respond), patch.dict(gateway.CONFIG, {"desensitize": True, "no_compact": True}): + response = self.fx.client.post(route, json=body) + self.assertEqual(len(seen), 2, response.text) + self.assert_attempts(response, seen) + if not filtered: + self.assertIn("call-test", response.text) + + def test_inflight_mode_does_not_change_during_a_retry(self): + self.fx.configure(profiles=("cn-cli",)) + seen = [] + def respond(request): + seen.append(request) + if len(seen) == 1: + gateway.CONFIG["request_context_mode"] = "legacy" + raise httpx.ConnectError("synthetic") + return httpx.Response(200, content=fixtures.fixtures.success_sse()) + with self.responder(respond), patch.object(gateway.asyncio, "sleep", new=AsyncMock()): + first = self.fx.client.post("/v1/chat/completions", json=self.fx.payload()) + self.assert_attempts(first, seen) + second = self.fx.client.post("/v1/chat/completions", json=self.fx.payload()) + self.assertEqual(second.status_code, 200) + self.assertNotEqual(second.headers["X-Request-ID"], seen[-1].headers["X-Root-Request-ID"]) + + def test_response_ids_match_audit_without_deduplicating_client_supplied_ids(self): + self.fx.configure(profiles=("cn-cli",)) + store = AuditStore(self.fx.root / "context-audit.sqlite3") + self.addCleanup(store.close) + app = AuditMiddleware(gateway.app, {**gateway.CONFIG, "audit_store": store}) + roots = [] + with TestClient(app) as client: + for _ in range(2): + body = self.fx.payload(text="same") + body["metadata"] = {"conversation_id": "private-session-canary"} + response = client.post("/v1/chat/completions", json=body, headers={"X-Request-ID": "client-reused-id"}) + roots.append(response.headers["X-Request-ID"]) + rows = store.list_records()["items"] + self.assertEqual({row["id"] for row in rows}, set(roots)) + self.assertEqual(len(set(roots)), 2) + self.assertEqual(store.dashboard()["summary"]["requests"], 2) + for row in rows: + attempt = next(a for a in row["attempts"] if a["stage"] == "upstream_attempt") + self.assertEqual(attempt["request_id"], row["id"]) + self.assertEqual(attempt["attempt_index"], 1) + self.assertEqual(attempt["attempt_id"], attempt["upstream_request_id"]) + self.assertNotIn("private-session-canary", json.dumps(rows)) + self.assertNotIn("synthetic-access", json.dumps(rows)) + self.assertNotIn("client-reused-id", json.dumps(rows)) + + +class AsyncContextTests(fixtures.GatewayFixture, unittest.IsolatedAsyncioTestCase): + async def test_parallel_requests_and_cancellation_leave_no_context_or_capacity(self): + self.fx.configure(profiles=("cn-cli",)) + reached, release = asyncio.Event(), asyncio.Event() + seen = [] + async def respond(request): + seen.append((request, current_context().request_id)) + if len(seen) == 4: + reached.set() + await release.wait() + return httpx.Response(200, content=fixtures.fixtures.success_sse()) + with self.responder(respond), patch.dict(gateway.CONFIG, {"request_context_mode": "scoped"}): + async with REAL_CLIENT(transport=httpx.ASGITransport(app=gateway.app), base_url="http://test") as client: + tasks = [asyncio.create_task(client.post("/v1/responses", json=self.fx.payload("responses"))) for _ in range(4)] + try: + await asyncio.wait_for(reached.wait(), 2) + self.assertEqual(sum(self.fx.pool._capacity._counts.values()), 4) + tasks[0].cancel() + await asyncio.gather(tasks[0], return_exceptions=True) + release.set() + responses = await asyncio.gather(*tasks[1:]) + finally: + release.set() + for task in tasks: + task.cancel() + await asyncio.gather(*tasks, return_exceptions=True) + roots = {root for _, root in seen} + self.assertEqual(len(roots), 4) + self.assertTrue(all(request.headers["X-Root-Request-ID"] == root for request, root in seen)) + self.assertTrue(all(r.status_code == 200 and r.headers["X-Request-ID"] in roots for r in responses)) + self.assertEqual(self.fx.pool._capacity._counts, {}) + self.assertIsNone(current_context()) + + async def test_bounded_audit_overflow_and_excluded_routes(self): + store = AuditStore(self.fx.root / "overflow.sqlite3") + self.addCleanup(store.close) + async def app(scope, receive, send): + for i in range(40): + observe_attempt("probe", attempt=i, token="private-canary") + await send({"type": "http.response.start", "status": 200, "headers": []}) + await send({"type": "http.response.body", "body": b'{}'}) + wrapped = RequestContextMiddleware(AuditMiddleware(app, {"audit_store": store}), {}) + sent = [] + async def receive(): + raise AssertionError("Middleware must not read the request body") + async def send(message): + sent.append(message) + await wrapped({"type": "http", "method": "POST", "path": "/v1/responses"}, receive, send) + row = store.list_records()["items"][0] + self.assertEqual(row["id"].encode(), dict(sent[0]["headers"])[b"x-request-id"]) + self.assertEqual(len(row["attempts"]), 32) + self.assertEqual(row["attempts"][-1], {"stage": "attempts_truncated", "dropped": 9}) + self.assertNotIn("private-canary", json.dumps(row)) + sent.clear() + await wrapped({"type": "http", "method": "POST", "path": "/admin/credentials"}, receive, send) + self.assertNotIn(b"x-request-id", dict(sent[0]["headers"])) + self.assertEqual(len(store.list_records()["items"]), 1) + + +if __name__ == "__main__": + unittest.main()