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()