From c2a4c5a7d2cafdd691387f99cf99491af71dd215 Mon Sep 17 00:00:00 2001 From: Ali Khokhar <20476625+Alishahryar1@users.noreply.github.com> Date: Tue, 29 Sep 2026 12:14:10 -0700 Subject: [PATCH] patch: Log final inference request outcomes at INFO (#1943) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Why Default INFO logs do not show the final outcome of each inference request. An HTTP 200 access record can hide a failure inside a stream, and individual provider failures can be recovered by model fallback. ## How Emit one structured INFO `request.completed` record for each POST to `/v1/messages` or `/v1/responses`. Include the existing request ID, wire protocol, selected provider/model, elapsed milliseconds, HTTP status, outcome and failure reason. Provider/model remain unset when a request is rejected before routing. Carry routing and terminal failure details in request-local context, updating the selected route during fallback. Observe HTTP delivery and use the existing incremental SSE decoder to recognize wire error events, including split frames. An optional event-name filter skips JSON decoding for named non-error events; unnamed events are still decoded to inspect their payload type. Other decoder callers retain their default behavior. Finalize the record after the request's owned work returns, distinguishing success, failure and cancellation. Keep request and response bodies out of the summary. Provider error codes and types remain available as diagnostic reasons under the agreed local logging policy.

RetriggerConfidence Score: 3/5

Not safe to merge until the previously reported log-disclosure issue is addressed.
Summary The PR adds outcome logging for inference requests and now skips JSON parsing for named non-error stream events. No new findings are reported. A previously reported security issue remains unresolved.
Reviews (2) · Last reviewed commit: ["Skip JSON decoding for non-error outcome..."](https://github.com/alishahryar1/free-claude-code/commit/8a5650550089f90495352ed26f00f755a69461d5) --- src/free_claude_code/api/app.py | 2 + src/free_claude_code/api/handlers/messages.py | 5 + .../api/handlers/responses.py | 5 + src/free_claude_code/api/request_errors.py | 6 + src/free_claude_code/api/request_outcomes.py | 126 +++++ src/free_claude_code/api/response_streams.py | 8 + src/free_claude_code/application/execution.py | 2 + .../core/anthropic/stream_contracts.py | 18 +- .../core/anthropic/streaming/decoder.py | 11 +- src/free_claude_code/core/request_outcomes.py | 38 ++ tests/api/test_request_outcomes.py | 438 ++++++++++++++++++ tests/core/anthropic/test_sse_decoder.py | 41 ++ 12 files changed, 690 insertions(+), 10 deletions(-) create mode 100644 src/free_claude_code/api/request_outcomes.py create mode 100644 src/free_claude_code/core/request_outcomes.py create mode 100644 tests/api/test_request_outcomes.py diff --git a/src/free_claude_code/api/app.py b/src/free_claude_code/api/app.py index 215315fb55..77c8c67c42 100644 --- a/src/free_claude_code/api/app.py +++ b/src/free_claude_code/api/app.py @@ -37,6 +37,7 @@ get_request_id, ) from .request_lifetime import ClientRequestLifetimeMiddleware +from .request_outcomes import RequestOutcomeMiddleware from .routes import router from .validation_log import summarize_request_validation_body @@ -47,6 +48,7 @@ def create_app(services: ApiServices) -> FastAPI: app.state.services = services app.add_middleware(AdminNoStoreMiddleware) app.add_middleware(ClientRequestLifetimeMiddleware) + app.add_middleware(RequestOutcomeMiddleware) app.add_middleware(RequestCorrelationMiddleware) app.include_router(admin_router) diff --git a/src/free_claude_code/api/handlers/messages.py b/src/free_claude_code/api/handlers/messages.py index 2dfd9e5fa2..505a25652a 100644 --- a/src/free_claude_code/api/handlers/messages.py +++ b/src/free_claude_code/api/handlers/messages.py @@ -42,6 +42,7 @@ from free_claude_code.core.diagnostics import safe_exception_message from free_claude_code.core.failures import ExecutionFailure, find_execution_failure from free_claude_code.core.reasoning import ReasoningControl, ReasoningPolicy +from free_claude_code.core.request_outcomes import record_request_route from free_claude_code.core.trace import trace_event from .classifier_response import classifier_response @@ -103,6 +104,10 @@ async def create( require_non_empty_messages(request_data.messages) routed = self._model_router.resolve_messages_request(request_data) routed = self._apply_message_routing_policies(routed) + record_request_route( + routed.resolved.primary.provider_id, + routed.resolved.primary.provider_model, + ) tool_body = self._web_tools.try_stream_messages( routed, request_id=request_id ) diff --git a/src/free_claude_code/api/handlers/responses.py b/src/free_claude_code/api/handlers/responses.py index abd90f2e6f..8c10f100b7 100644 --- a/src/free_claude_code/api/handlers/responses.py +++ b/src/free_claude_code/api/handlers/responses.py @@ -29,6 +29,7 @@ openai_error_type_for_failure, openai_failure_payload, ) +from free_claude_code.core.request_outcomes import record_request_route class ResponsesHandler: @@ -77,6 +78,10 @@ async def create( try: routed = self._model_router.resolve_responses_request(request_data) + record_request_route( + routed.resolved.primary.provider_id, + routed.resolved.primary.provider_model, + ) streamed = self._provider_executor.stream_responses( routed, raw_log_payload=request_payload, diff --git a/src/free_claude_code/api/request_errors.py b/src/free_claude_code/api/request_errors.py index 4607250eae..83b093ec77 100644 --- a/src/free_claude_code/api/request_errors.py +++ b/src/free_claude_code/api/request_errors.py @@ -21,6 +21,10 @@ openai_error_payload, openai_error_type_for_failure, ) +from free_claude_code.core.request_outcomes import ( + record_request_exception, + record_request_failure, +) WireApi = Literal["messages", "responses"] @@ -37,6 +41,7 @@ def ordinary_application_error_response( request_id: str, ) -> JSONResponse: """Serialize a deterministic application error without terminal headers.""" + record_request_failure(error.kind.value) if wire_api == "responses": return JSONResponse( status_code=error.status_code, @@ -67,6 +72,7 @@ def log_unexpected_api_exception( request_id: str | None = None, ) -> None: """Log API failures without echoing exception text unless opted in.""" + record_request_exception(exc) if settings.log_api_error_tracebacks: if request_id is not None: logger.error( diff --git a/src/free_claude_code/api/request_outcomes.py b/src/free_claude_code/api/request_outcomes.py new file mode 100644 index 0000000000..c1d63d6831 --- /dev/null +++ b/src/free_claude_code/api/request_outcomes.py @@ -0,0 +1,126 @@ +"""Log one final inference outcome after HTTP delivery and owned cleanup.""" + +import asyncio +import codecs +from time import monotonic + +from loguru import logger +from starlette.datastructures import Headers +from starlette.types import ASGIApp, Message, Receive, Scope, Send + +from free_claude_code.core.anthropic.stream_contracts import SSEEvent +from free_claude_code.core.anthropic.streaming.decoder import AnthropicSSEDecoder +from free_claude_code.core.request_outcomes import ( + RequestOutcome, + current_request_outcome, + record_request_exception, + record_request_failure, +) + +_FAILURE_EVENTS = frozenset({"error", "response.error", "response.failed"}) + + +def _observe_event(event: SSEEvent) -> None: + kind = event.event or event.data.get("type") + if not isinstance(kind, str) or kind not in _FAILURE_EVENTS: + return + response = event.data.get("response") + payload = response if isinstance(response, dict) else event.data + error = payload.get("error") + error = error if isinstance(error, dict) else payload + reason = error.get("code") or error.get("type") + record_request_failure(reason if isinstance(reason, str) else str(kind)) + + +class RequestOutcomeMiddleware: + """Observe inference delivery without changing its body or control flow.""" + + def __init__(self, app: ASGIApp) -> None: + self._app = app + + async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: + if ( + scope["type"] != "http" + or scope.get("method") != "POST" + or scope.get("path") not in {"/v1/messages", "/v1/responses"} + ): + await self._app(scope, receive, send) + return + + started = monotonic() + outcome = RequestOutcome() + status_code: int | None = None + completed = disconnected = cancelled = False + decoder: AnthropicSSEDecoder | None = None + text_decoder = codecs.getincrementaldecoder("utf-8")(errors="replace") + + async def receive_observed() -> Message: + nonlocal disconnected + message = await receive() + if message["type"] == "http.disconnect": + disconnected = True + return message + + async def send_observed(message: Message) -> None: + nonlocal status_code, completed, decoder, disconnected + try: + await send(message) + except OSError: + disconnected = True + raise + if message["type"] == "http.response.start": + status_code = message["status"] + if ( + Headers(raw=message.get("headers", [])) + .get("content-type", "") + .startswith("text/event-stream") + ): + decoder = AnthropicSSEDecoder(event_names=_FAILURE_EVENTS) + elif message["type"] == "http.response.body": + completed = not message.get("more_body", False) + if decoder is not None: + text = text_decoder.decode( + message.get("body", b""), final=completed + ) + for event in decoder.feed(text): + _observe_event(event) + if completed: + for event in decoder.finish(): + _observe_event(event) + + token = current_request_outcome.set(outcome) + try: + await self._app(scope, receive_observed, send_observed) + except asyncio.CancelledError: + cancelled = True + raise + except BaseException as error: + if not disconnected: + record_request_exception(error) + if status_code is None: + status_code = 500 + raise + finally: + current_request_outcome.reset(token) + reason = outcome.failure_reason + if reason is not None or (status_code is not None and status_code >= 400): + result = "failure" + reason = reason or f"http_{status_code}" + elif (cancelled or disconnected) and not completed: + result = "cancelled" + elif not completed: + result, reason = "failure", "incomplete_response" + else: + result = "success" + logger.bind( + event="request.completed", + wire_api="responses" + if scope["path"] == "/v1/responses" + else "messages", + provider_id=outcome.provider_id, + model=outcome.model, + status_code=status_code, + outcome=result, + duration_ms=round((monotonic() - started) * 1000, 2), + failure_reason=reason, + ).info("Inference request {}", result) diff --git a/src/free_claude_code/api/response_streams.py b/src/free_claude_code/api/response_streams.py index 136b469ccc..2a1777af1e 100644 --- a/src/free_claude_code/api/response_streams.py +++ b/src/free_claude_code/api/response_streams.py @@ -28,6 +28,10 @@ committed_response_failure_frame, openai_error_type_for_failure, ) +from free_claude_code.core.request_outcomes import ( + record_request_exception, + record_request_failure, +) from free_claude_code.core.trace import close_stream_input, trace_event TERMINAL_EXECUTION_ERROR_HEADERS = {"x-should-retry": "false"} @@ -198,6 +202,10 @@ def trace_terminal_execution_error( error: BaseException | None = None, ) -> None: """Record one correlated terminal-execution decision at the HTTP boundary.""" + if error is not None: + record_request_exception(error) + else: + record_request_failure(error_type) fields: dict[str, object] = { "stage": "egress", "event": "free_claude_code.api.response.terminal_execution_error", diff --git a/src/free_claude_code/application/execution.py b/src/free_claude_code/application/execution.py index 0181bddb73..03c9582bc9 100644 --- a/src/free_claude_code/application/execution.py +++ b/src/free_claude_code/application/execution.py @@ -23,6 +23,7 @@ estimate_responses_input_tokens, ) from free_claude_code.core.reasoning import ReasoningPolicy +from free_claude_code.core.request_outcomes import record_request_route from free_claude_code.core.trace import ( close_stream_input, trace_event, @@ -341,6 +342,7 @@ async def provider_body() -> AsyncIterator[str]: loop = asyncio.get_running_loop() progress_deadline = loop.time() + self._progress_timeout_seconds for index, target in enumerate(candidates): + record_request_route(target.provider_id, target.provider_model) provider_stream: AsyncIterator[str] | None = None candidate_committed = False candidate_failure: ExecutionFailure | None = None diff --git a/src/free_claude_code/core/anthropic/stream_contracts.py b/src/free_claude_code/core/anthropic/stream_contracts.py index a5aa76d397..270917c5ea 100644 --- a/src/free_claude_code/core/anthropic/stream_contracts.py +++ b/src/free_claude_code/core/anthropic/stream_contracts.py @@ -47,7 +47,10 @@ class SSEEvent: raw: str -def parse_sse_lines(lines: Iterable[str]) -> list[SSEEvent]: +def parse_sse_lines( + lines: Iterable[str], *, event_names: frozenset[str] | None = None +) -> list[SSEEvent]: + """Decode selected named events and all unnamed events when a filter is supplied.""" events: list[SSEEvent] = [] current_event = "" data_parts: list[str] = [] @@ -56,7 +59,7 @@ def parse_sse_lines(lines: Iterable[str]) -> list[SSEEvent]: for line in lines: stripped = line.rstrip("\r\n") if stripped == "": - _append_event(events, current_event, data_parts, raw_parts) + _append_event(events, current_event, data_parts, raw_parts, event_names) current_event = "" data_parts = [] raw_parts = [] @@ -67,13 +70,15 @@ def parse_sse_lines(lines: Iterable[str]) -> list[SSEEvent]: elif stripped.startswith("data:"): data_parts.append(stripped.split(":", 1)[1].strip()) - _append_event(events, current_event, data_parts, raw_parts) + _append_event(events, current_event, data_parts, raw_parts, event_names) return events -def parse_sse_text(text: str) -> list[SSEEvent]: +def parse_sse_text( + text: str, *, event_names: frozenset[str] | None = None +) -> list[SSEEvent]: # SSE uses CR/LF framing; Unicode line separators can occur inside JSON text. - return parse_sse_lines(re.split(r"\r\n|\r|\n", text)) + return parse_sse_lines(re.split(r"\r\n|\r|\n", text), event_names=event_names) def _append_event( @@ -81,9 +86,12 @@ def _append_event( current_event: str, data_parts: list[str], raw_parts: list[str], + event_names: frozenset[str] | None, ) -> None: if not current_event and not data_parts: return + if event_names is not None and current_event and current_event not in event_names: + return data_text = "\n".join(data_parts) data: dict[str, Any] try: diff --git a/src/free_claude_code/core/anthropic/streaming/decoder.py b/src/free_claude_code/core/anthropic/streaming/decoder.py index 6c7f66a394..efabb81f92 100644 --- a/src/free_claude_code/core/anthropic/streaming/decoder.py +++ b/src/free_claude_code/core/anthropic/streaming/decoder.py @@ -8,14 +8,15 @@ class AnthropicSSEDecoder: - """Decode arbitrarily split SSE text without losing frame order.""" + """Decode split SSE text, optionally filtering named events before JSON parsing.""" - def __init__(self) -> None: + def __init__(self, *, event_names: frozenset[str] | None = None) -> None: + self._event_names = event_names self._parts: list[str] = [] self._boundary_tail = "" def feed(self, chunk: str) -> tuple[SSEEvent, ...]: - """Consume one wire chunk and return every complete event.""" + """Consume one wire chunk and return its selected complete events.""" events: list[SSEEvent] = [] probe = self._boundary_tail + chunk @@ -26,7 +27,7 @@ def feed(self, chunk: str) -> tuple[SSEEvent, ...]: self._parts.append(chunk[chunk_start:chunk_end]) raw = "".join(self._parts) self._parts.clear() - events.extend(parse_sse_text(raw)) + events.extend(parse_sse_text(raw, event_names=self._event_names)) chunk_start = chunk_end remainder = chunk[chunk_start:] @@ -46,4 +47,4 @@ def finish(self) -> tuple[SSEEvent, ...]: self._boundary_tail = "" if not remainder.strip(): return () - return tuple(parse_sse_text(remainder)) + return tuple(parse_sse_text(remainder, event_names=self._event_names)) diff --git a/src/free_claude_code/core/request_outcomes.py b/src/free_claude_code/core/request_outcomes.py new file mode 100644 index 0000000000..934c2b963c --- /dev/null +++ b/src/free_claude_code/core/request_outcomes.py @@ -0,0 +1,38 @@ +"""Request-scoped diagnostic fields shared by routing and HTTP delivery.""" + +from contextvars import ContextVar +from dataclasses import dataclass + +from .failures import find_execution_failure + + +@dataclass +class RequestOutcome: + """Mutable fields shared with the request's response and cleanup tasks.""" + + provider_id: str | None = None + model: str | None = None + failure_reason: str | None = None + + +current_request_outcome: ContextVar[RequestOutcome | None] = ContextVar( + "request_outcome", default=None +) + + +def record_request_route(provider_id: str, model: str) -> None: + outcome = current_request_outcome.get() + if outcome is not None: + outcome.provider_id = provider_id + outcome.model = model + + +def record_request_failure(reason: str) -> None: + outcome = current_request_outcome.get() + if outcome is not None and outcome.failure_reason is None: + outcome.failure_reason = reason + + +def record_request_exception(error: BaseException) -> None: + failure = find_execution_failure(error) + record_request_failure(failure.kind.value if failure else type(error).__name__) diff --git a/tests/api/test_request_outcomes.py b/tests/api/test_request_outcomes.py new file mode 100644 index 0000000000..e45a98b30f --- /dev/null +++ b/tests/api/test_request_outcomes.py @@ -0,0 +1,438 @@ +"""Final inference outcomes include failures hidden inside HTTP 200 streams.""" + +import asyncio +import json +from unittest.mock import patch + +import pytest +from fastapi.testclient import TestClient +from loguru import logger + +from free_claude_code.api.request_ids import RequestCorrelationMiddleware +from free_claude_code.api.request_outcomes import RequestOutcomeMiddleware +from free_claude_code.config.settings import Settings +from free_claude_code.core.failures import ExecutionFailure, FailureKind +from free_claude_code.core.request_outcomes import ( + current_request_outcome, + record_request_route, +) +from tests.api.support import create_test_app +from tests.api.test_request_lifetime import _http_scope + + +@pytest.fixture +def outcomes(): + records = [] + sink = logger.add( + lambda message: records.append(message.record), + level="INFO", + filter=lambda record: record["extra"].get("event") == "request.completed", + ) + try: + yield records + finally: + logger.remove(sink) + + +class OutcomeProvider: + def __init__(self, result): + self.result = result + self.started = asyncio.Event() + + async def stream_messages(self, request, **kwargs): + async for chunk in self._stream("messages"): + yield chunk + + async def stream_responses(self, request, **kwargs): + async for chunk in self._stream("responses"): + yield chunk + + async def _stream(self, wire_api): + failure = ExecutionFailure( + FailureKind.RATE_LIMIT, 429, "private error text", True + ) + if self.result == "wait_before_start": + self.started.set() + await asyncio.Event().wait() + if self.result == "pre_start": + raise failure + if wire_api == "responses": + yield 'event: response.created\ndata: {"type":"response.created","response":{"id":"resp_test","object":"response","status":"in_progress","output":[]}}\n\n' + else: + yield 'event: message_start\ndata: {"type":"message_start","message":{"id":"msg_test","type":"message","role":"assistant","content":[]}}\n\n' + if self.result == "exception": + raise failure + if self.result == "wait_after_start": + self.started.set() + await asyncio.Event().wait() + if self.result == "wire_error": + event = "error" if wire_api == "messages" else "response.failed" + data = {"type": event, "error": {"type": "rate_limit_error"}} + if wire_api == "responses": + data = { + "type": event, + "response": {"status": "failed", "error": data["error"]}, + } + frame = f"event: {event}\ndata: {json.dumps(data)}\n\n" + # Framing must work even when a wire event is split between chunks. + yield frame[:12] + yield frame[12:] + else: + yield 'event: message_stop\ndata: {"type":"message_stop"}\n\n' + + +@pytest.mark.parametrize("wire_api", ["messages", "responses"]) +@pytest.mark.parametrize("result", ["success", "pre_start", "exception", "wire_error"]) +def test_one_final_outcome_for_streamed_inference(outcomes, wire_api, result): + payload = {"model": "nvidia_nim/test-model", "stream": True} + if wire_api == "messages": + payload.update( + max_tokens=64, messages=[{"role": "user", "content": "private input"}] + ) + else: + payload["input"] = "private input" + with ( + patch( + "free_claude_code.api.routes.resolve_provider", + return_value=OutcomeProvider(result), + ), + TestClient(create_test_app(Settings())) as client, + ): + response = client.post(f"/v1/{wire_api}", json=payload) + assert response.status_code == (429 if result == "pre_start" else 200) + assert len(outcomes) == 1 + record = outcomes[0] + fields = record["extra"] + assert record["level"].name == "INFO" + assert fields["request_id"] == response.headers["request-id"] + assert fields["provider_id"] == "nvidia_nim" + assert fields["model"] == "test-model" + assert fields["wire_api"] == wire_api + assert fields["status_code"] == response.status_code + assert fields["outcome"] == ("success" if result == "success" else "failure") + assert fields["duration_ms"] >= 0 + assert ( + fields["failure_reason"] + == { + "success": None, + "pre_start": "rate_limit", + "exception": "rate_limit", + "wire_error": "rate_limit_error", + }[result] + ) + assert "private input" not in str(record) + assert "private error text" not in str(record) + + +def test_fallback_records_selected_provider_without_logging_failed_attempt(outcomes): + settings = Settings(model_fallbacks=["open_router/fallback-model"]) + providers = { + "nvidia_nim": OutcomeProvider("pre_start"), + "open_router": OutcomeProvider("success"), + } + with ( + patch( + "free_claude_code.api.routes.resolve_provider", + side_effect=lambda name, **kwargs: providers[name], + ), + TestClient(create_test_app(settings)) as client, + ): + response = client.post( + "/v1/messages", + json={ + "model": "nvidia_nim/test-model", + "messages": [{"role": "user", "content": "Hello"}], + "stream": True, + }, + ) + assert response.status_code == 200 + assert len(outcomes) == 1 + fields = outcomes[0]["extra"] + assert ( + fields["provider_id"], + fields["model"], + fields["outcome"], + fields["failure_reason"], + ) == ( + "open_router", + "fallback-model", + "success", + None, + ) + + +@pytest.mark.parametrize("result", ["success", "exception", "wire_error"]) +def test_non_streaming_messages_have_one_outcome(outcomes, result): + with ( + patch( + "free_claude_code.api.routes.resolve_provider", + return_value=OutcomeProvider(result), + ), + TestClient(create_test_app(Settings())) as client, + patch( + "free_claude_code.api.request_outcomes.monotonic", side_effect=[10.0, 12.5] + ), + ): + response = client.post( + "/v1/messages", + json={ + "model": "nvidia_nim/test-model", + "messages": [{"role": "user", "content": "Hello"}], + "stream": False, + }, + ) + assert response.status_code == (200 if result == "success" else 429) + assert len(outcomes) == 1 + assert outcomes[0]["extra"]["duration_ms"] == 2500 + assert outcomes[0]["extra"]["outcome"] == ( + "success" if result == "success" else "failure" + ) + + +@pytest.mark.parametrize( + "payload,status,reason", + [ + ({}, 422, "http_422"), + ( + {"model": "nvidia_nim/test-model", "input": "hello", "stream": False}, + 400, + "invalid_request", + ), + ], +) +def test_rejections_before_execution_are_recorded(outcomes, payload, status, reason): + with TestClient(create_test_app(Settings())) as client: + response = client.post("/v1/responses", json=payload) + assert response.status_code == status + assert len(outcomes) == 1 + fields = outcomes[0]["extra"] + assert fields["outcome"] == "failure" + assert fields["failure_reason"] == reason + assert fields["provider_id"] is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("wire_api", ["messages", "responses"]) +@pytest.mark.parametrize("stage", ["wait_before_start", "wait_after_start"]) +async def test_disconnect_logs_cancellation_after_cleanup(outcomes, wire_api, stage): + provider = OutcomeProvider(stage) + frames = asyncio.Queue() + payload = {"model": "nvidia_nim/test-model", "stream": True} + payload.update( + {"messages": [{"role": "user", "content": "hello"}]} + if wire_api == "messages" + else {"input": "hello"} + ) + await frames.put({"type": "http.request", "body": json.dumps(payload).encode()}) + sent = [] + + async def send(message): + sent.append(message) + + with patch("free_claude_code.api.routes.resolve_provider", return_value=provider): + app = create_test_app(Settings()) + scope = _http_scope(f"/v1/{wire_api}") + scope["headers"] = [(b"content-type", b"application/json")] + task = asyncio.create_task(app(scope, frames.get, send)) + try: + async with asyncio.timeout(5): + await provider.started.wait() + assert outcomes == [] + await frames.put({"type": "http.disconnect"}) + await task + finally: + task.cancel() + await asyncio.gather(task, return_exceptions=True) + assert len(outcomes) == 1 + fields = outcomes[0]["extra"] + assert fields["outcome"] == "cancelled" + assert fields["status_code"] == (200 if stage == "wait_after_start" else None) + assert fields["provider_id"] == "nvidia_nim" + assert current_request_outcome.get() is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "path,method", + [ + ("/health", "GET"), + ("/v1/models", "GET"), + ("/admin/api/status", "GET"), + ("/v1/messages/count_tokens", "POST"), + ("/v1/messages", "OPTIONS"), + ], +) +async def test_other_endpoints_do_not_emit_inference_outcomes(outcomes, path, method): + async def app(scope, receive, send): + assert current_request_outcome.get() is None + await send({"type": "http.response.start", "status": 200}) + await send({"type": "http.response.body", "body": b"{}"}) + + async def receive(): + return {"type": "http.disconnect"} + + async def send(message): + pass + + await RequestOutcomeMiddleware(app)(_http_scope(path, method=method), receive, send) + assert outcomes == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "ending", ["send_failure", "late_disconnect", "exception", "cancelled"] +) +async def test_delivery_outcomes_preserve_control_flow(outcomes, ending): + async def app(scope, receive, send): + if ending == "exception": + raise ValueError("private exception message") + if ending == "cancelled": + raise asyncio.CancelledError() + await send({"type": "http.response.start", "status": 200}) + await send({"type": "http.response.body", "body": b"{}"}) + await receive() + + async def receive(): + return {"type": "http.disconnect"} + + async def send(message): + if ending == "send_failure" and message["type"] == "http.response.body": + raise OSError("connection closed") + + middleware = RequestCorrelationMiddleware(RequestOutcomeMiddleware(app)) + if ending == "late_disconnect": + await middleware(_http_scope(), receive, send) + else: + with pytest.raises( + { + "exception": ValueError, + "cancelled": asyncio.CancelledError, + "send_failure": OSError, + }[ending] + ): + await middleware(_http_scope(), receive, send) + assert len(outcomes) == 1 + fields = outcomes[0]["extra"] + assert ( + fields["outcome"] + == { + "send_failure": "cancelled", + "late_disconnect": "success", + "exception": "failure", + "cancelled": "cancelled", + }[ending] + ) + assert fields["failure_reason"] == ("ValueError" if ending == "exception" else None) + assert "private exception message" not in str(outcomes) + assert current_request_outcome.get() is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("event_header", ["event: error\r\n", ""]) +async def test_wire_observation_preserves_split_utf8_body(outcomes, event_header): + frame = ( + event_header + + 'data: {"type":"error","error":{"type":"api_error","message":"☃"}}\r\n\r\n' + ).encode() + chunks = [frame[index : index + 2] for index in range(0, len(frame), 2)] + sent = [] + + async def app(scope, receive, send): + await send( + { + "type": "http.response.start", + "status": 200, + "headers": [(b"content-type", b"text/event-stream")], + } + ) + for chunk in chunks: + await send({"type": "http.response.body", "body": chunk, "more_body": True}) + await send({"type": "http.response.body", "body": b""}) + + async def receive(): + return {"type": "http.disconnect"} + + async def send(message): + sent.append(message) + + await RequestCorrelationMiddleware(RequestOutcomeMiddleware(app))( + _http_scope(), receive, send + ) + assert b"".join(message.get("body", b"") for message in sent) == frame + assert len(outcomes) == 1 + assert outcomes[0]["extra"]["outcome"] == "failure" + assert outcomes[0]["extra"]["failure_reason"] == "api_error" + assert "☃" not in str(outcomes) + + +@pytest.mark.asyncio +async def test_successful_named_events_skip_json_decoding(outcomes): + frame = ( + b'event: response.output_text.delta\ndata: {"type":"response.output_text.delta","delta":"hello"}\n\n' + b'event: response.completed\ndata: {"type":"response.completed","response":{"output":[]}}\n\n' + ) + + async def app(scope, receive, send): + await send( + { + "type": "http.response.start", + "status": 200, + "headers": [(b"content-type", b"text/event-stream")], + } + ) + await send({"type": "http.response.body", "body": frame}) + + async def receive(): + return {"type": "http.disconnect"} + + async def send(message): + pass + + with patch( + "free_claude_code.core.anthropic.stream_contracts.json.loads", wraps=json.loads + ) as loads: + await RequestCorrelationMiddleware(RequestOutcomeMiddleware(app))( + _http_scope("/v1/responses"), receive, send + ) + loads.assert_not_called() + assert len(outcomes) == 1 + assert outcomes[0]["extra"]["outcome"] == "success" + + +@pytest.mark.asyncio +async def test_overlapping_requests_keep_separate_outcomes(outcomes): + ready = asyncio.Event() + count = 0 + + async def app(scope, receive, send): + nonlocal count + name = scope["query_string"].decode() + record_request_route(name, f"{name}-model") + count += 1 + if count == 2: + ready.set() + await ready.wait() + await send({"type": "http.response.start", "status": 200}) + await send({"type": "http.response.body", "body": b"{}"}) + + async def receive(): + await asyncio.Event().wait() + + async def send(message): + pass + + middleware = RequestCorrelationMiddleware(RequestOutcomeMiddleware(app)) + scopes = [_http_scope(), _http_scope()] + for scope, name in zip(scopes, [b"first", b"second"], strict=True): + scope["query_string"] = name + async with asyncio.timeout(5): + await asyncio.gather(*(middleware(scope, receive, send) for scope in scopes)) + assert len(outcomes) == 2 + assert { + (row["extra"]["provider_id"], row["extra"]["model"]) for row in outcomes + } == { + ("first", "first-model"), + ("second", "second-model"), + } + assert len({row["extra"]["request_id"] for row in outcomes}) == 2 + assert current_request_outcome.get() is None diff --git a/tests/core/anthropic/test_sse_decoder.py b/tests/core/anthropic/test_sse_decoder.py index 2d9821ab21..50ee375e84 100644 --- a/tests/core/anthropic/test_sse_decoder.py +++ b/tests/core/anthropic/test_sse_decoder.py @@ -1,3 +1,8 @@ +import json +from unittest.mock import patch + +import pytest + from free_claude_code.core.anthropic.streaming import AnthropicSSEDecoder @@ -26,6 +31,42 @@ def test_decoder_returns_one_unterminated_final_event(): assert decoder.finish() == () +@pytest.mark.parametrize("ending", ["\n\n", ""]) +def test_event_filter_skips_json_and_preserves_unnamed_events_at_every_split(ending): + wire = ( + 'event: delta\r\ndata: {"text":"ignored"}\r\n\r\n' + 'data: {"type":"error"}\nevent: error\n\n' + 'data: {"type":"response.failed"}\r\n\r\n' + 'event: delta\ndata: {"text":"ignored final"}' + ending + ) + for split in range(len(wire) + 1): + decoder = AnthropicSSEDecoder(event_names=frozenset({"error"})) + with patch( + "free_claude_code.core.anthropic.stream_contracts.json.loads", + wraps=json.loads, + ) as loads: + events = ( + *decoder.feed(wire[:split]), + *decoder.feed(wire[split:]), + *decoder.finish(), + ) + assert [(event.event, event.data["type"]) for event in events] == [ + ("error", "error"), + ("", "response.failed"), + ] + assert [call.args[0] for call in loads.call_args_list] == [ + '{"type":"error"}', + '{"type":"response.failed"}', + ] + assert decoder.finish() == () + + +def test_event_filter_decodes_selected_unterminated_event(): + decoder = AnthropicSSEDecoder(event_names=frozenset({"error"})) + assert decoder.feed('event: error\ndata: {"type":"error"}') == () + assert [event.data for event in decoder.finish()] == [{"type": "error"}] + + def test_decoder_handles_many_tiny_fragments_without_losing_frames(): wire = "".join( f'event: delta\ndata: {{"index":{index}}}\n\n' for index in range(250)