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)