From c024bf6892a22040addced9512292d6356303756 Mon Sep 17 00:00:00 2001 From: Alishahryar1 Date: Wed, 17 Jun 2026 20:43:38 -0700 Subject: [PATCH] Fix stream cleanup context handling. Avoid contextvar-based log context in SSE generators and treat GeneratorExit as quiet teardown. --- api/services.py | 156 +++++++++--------- core/trace.py | 8 +- providers/transports/openai_chat/transport.py | 20 +-- pyproject.toml | 2 +- tests/core/test_trace.py | 101 +++++++++++- uv.lock | 2 +- 6 files changed, 191 insertions(+), 98 deletions(-) diff --git a/api/services.py b/api/services.py index d1d0ac81..fd0bc824 100644 --- a/api/services.py +++ b/api/services.py @@ -187,45 +187,45 @@ class ClaudeProxyService: ) request_id = f"req_{uuid.uuid4().hex[:12]}" - with logger.contextualize(request_id=request_id): - trace_event( - stage="ingress", - event="api.request.received", - source="api", - message_count=len(routed.request.messages), - snapshot=api_messages_request_snapshot(routed.request), + trace_event( + stage="ingress", + event="api.request.received", + source="api", + message_count=len(routed.request.messages), + snapshot=api_messages_request_snapshot(routed.request), + request_id=request_id, + ) + + if self._settings.log_raw_api_payloads: + logger.debug( + "FULL_PAYLOAD [{}]: {}", request_id, routed.request.model_dump() ) - if self._settings.log_raw_api_payloads: - logger.debug( - "FULL_PAYLOAD [{}]: {}", request_id, routed.request.model_dump() - ) + input_tokens = self._token_counter( + routed.request.messages, + routed.request.system, + routed.request.tools, + ) - input_tokens = self._token_counter( - routed.request.messages, - routed.request.system, - routed.request.tools, - ) - - streamed = traced_async_stream( - provider.stream_response( - routed.request, - input_tokens=input_tokens, - request_id=request_id, - thinking_enabled=routed.resolved.thinking_enabled, - ), - stage="egress", - source="api", - complete_event="api.response.stream_completed", - interrupted_event="api.response.stream_interrupted", - chunk_event=None, - extra={ - "request_id": request_id, - "provider_id": routed.resolved.provider_id, - "gateway_model": routed.request.model, - }, - ) - return anthropic_sse_streaming_response(streamed) + streamed = traced_async_stream( + provider.stream_response( + routed.request, + input_tokens=input_tokens, + request_id=request_id, + thinking_enabled=routed.resolved.thinking_enabled, + ), + stage="egress", + source="api", + complete_event="api.response.stream_completed", + interrupted_event="api.response.stream_interrupted", + chunk_event=None, + extra={ + "request_id": request_id, + "provider_id": routed.resolved.provider_id, + "gateway_model": routed.request.model, + }, + ) + return anthropic_sse_streaming_response(streamed) except ProviderError: raise @@ -289,52 +289,52 @@ class ClaudeProxyService: ) request_id = f"req_{uuid.uuid4().hex[:12]}" - with logger.contextualize(request_id=request_id): - trace_event( - stage="ingress", - event="api.responses.request.received", - source="api", - message_count=len(routed.request.messages), - snapshot=api_messages_request_snapshot(routed.request), + trace_event( + stage="ingress", + event="api.responses.request.received", + source="api", + message_count=len(routed.request.messages), + snapshot=api_messages_request_snapshot(routed.request), + request_id=request_id, + ) + + if self._settings.log_raw_api_payloads: + logger.debug( + "FULL_RESPONSES_PAYLOAD [{}]: {}", + request_id, + request_payload, ) - if self._settings.log_raw_api_payloads: - logger.debug( - "FULL_RESPONSES_PAYLOAD [{}]: {}", - request_id, - request_payload, - ) + input_tokens = self._token_counter( + routed.request.messages, + routed.request.system, + routed.request.tools, + ) - input_tokens = self._token_counter( - routed.request.messages, - routed.request.system, - routed.request.tools, - ) - - streamed = traced_async_stream( - provider.stream_response( - routed.request, - input_tokens=input_tokens, - request_id=request_id, - thinking_enabled=routed.resolved.thinking_enabled, - ), - stage="egress", - source="api", - complete_event="api.responses.stream_completed", - interrupted_event="api.responses.stream_interrupted", - chunk_event=None, - extra={ - "request_id": request_id, - "provider_id": routed.resolved.provider_id, - "gateway_model": routed.request.model, - }, - ) - return openai_responses_sse_streaming_response( - self._responses_adapter.iter_sse_from_anthropic( - streamed, - request_payload, - ) + streamed = traced_async_stream( + provider.stream_response( + routed.request, + input_tokens=input_tokens, + request_id=request_id, + thinking_enabled=routed.resolved.thinking_enabled, + ), + stage="egress", + source="api", + complete_event="api.responses.stream_completed", + interrupted_event="api.responses.stream_interrupted", + chunk_event=None, + extra={ + "request_id": request_id, + "provider_id": routed.resolved.provider_id, + "gateway_model": routed.request.model, + }, + ) + return openai_responses_sse_streaming_response( + self._responses_adapter.iter_sse_from_anthropic( + streamed, + request_payload, ) + ) except OpenAIResponsesAdapter.ConversionError as exc: invalid_request = InvalidRequestError(str(exc)) return JSONResponse( diff --git a/core/trace.py b/core/trace.py index a077d889..77a9458a 100644 --- a/core/trace.py +++ b/core/trace.py @@ -8,7 +8,7 @@ sanitized credential keys (e.g. ``api_key``, ``authorization``). from __future__ import annotations import asyncio -from collections.abc import AsyncIterator, Mapping +from collections.abc import AsyncGenerator, AsyncIterator, Mapping from typing import Any from loguru import logger @@ -116,7 +116,7 @@ async def traced_async_stream( chunk_event: str | None = None, chunk_interval: int = 250, extra: Mapping[str, Any] | None = None, -) -> AsyncIterator[str]: +) -> AsyncGenerator[str]: """Emit TRACE rows when a text stream completes, fails, cancels, or periodically.""" common = dict(extra or {}) count = 0 @@ -136,6 +136,8 @@ async def traced_async_stream( **common, ) yield chunk + except GeneratorExit: + raise except asyncio.CancelledError: interrupted = True trace_event( @@ -161,7 +163,7 @@ async def traced_async_stream( **common, ) raise - except BaseException as exc: + except Exception as exc: interrupted = True trace_event( stage=stage, diff --git a/providers/transports/openai_chat/transport.py b/providers/transports/openai_chat/transport.py index 2c83396a..978586dc 100644 --- a/providers/transports/openai_chat/transport.py +++ b/providers/transports/openai_chat/transport.py @@ -7,7 +7,6 @@ from collections.abc import AsyncIterator, Iterator from typing import Any import httpx -from loguru import logger from openai import AsyncOpenAI from core.anthropic import SSEBuilder @@ -146,13 +145,12 @@ class OpenAIChatTransport(BaseProvider): thinking_enabled: bool | None = None, ) -> AsyncIterator[str]: """Stream response in Anthropic SSE format.""" - with logger.contextualize(request_id=request_id): - runner = OpenAIChatStreamRunner( - self, - request=request, - input_tokens=input_tokens, - request_id=request_id, - thinking_enabled=thinking_enabled, - ) - async for event in runner.run(): - yield event + runner = OpenAIChatStreamRunner( + self, + request=request, + input_tokens=input_tokens, + request_id=request_id, + thinking_enabled=thinking_enabled, + ) + async for event in runner.run(): + yield event diff --git a/pyproject.toml b/pyproject.toml index 24f1471d..574d334b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "hatchling.build" [project] name = "free-claude-code" -version = "2.3.4" +version = "2.3.5" description = "Middleware between Claude Code CLI (Anthropic API) and NVIDIA NIM" readme = "README.md" requires-python = ">=3.14.0" diff --git a/tests/core/test_trace.py b/tests/core/test_trace.py index eb4c147d..1268922b 100644 --- a/tests/core/test_trace.py +++ b/tests/core/test_trace.py @@ -5,19 +5,26 @@ from __future__ import annotations import json from pathlib import Path +import pytest from loguru import logger from config.logging_config import configure_logging -from core.trace import TRACE_PAYLOAD_BINDING, trace_event +from core.trace import TRACE_PAYLOAD_BINDING, trace_event, traced_async_stream + + +def _json_log_rows(log_file: str) -> list[dict]: + logger.complete() + text = Path(log_file).read_text(encoding="utf-8").strip() + if not text: + return [] + return [json.loads(line) for line in text.split("\n")] def test_trace_payload_merged_into_json_line(tmp_path) -> None: log_file = str(tmp_path / "t.log") configure_logging(log_file, force=True) trace_event(stage="s", event="e.v1", source="unit", hello="world", n=42) - logger.complete() - text = Path(log_file).read_text(encoding="utf-8").strip().split("\n")[-1] - row = json.loads(text) + row = _json_log_rows(log_file)[-1] assert row["trace"] is True assert row["stage"] == "s" assert row["event"] == "e.v1" @@ -36,3 +43,89 @@ def test_sanitize_masks_nested_api_key_strings() -> None: ) assert out["outer"]["api_key"] == "" assert out["outer"]["text"] == "visible" + + +@pytest.mark.asyncio +async def test_traced_async_stream_logs_completion(tmp_path) -> None: + log_file = str(tmp_path / "complete.log") + configure_logging(log_file, force=True) + + async def source(): + yield "hello" + yield " world" + + chunks = [ + chunk + async for chunk in traced_async_stream( + source(), + stage="egress", + source="unit", + complete_event="stream.completed", + interrupted_event="stream.interrupted", + extra={"request_id": "req_complete"}, + ) + ] + + assert chunks == ["hello", " world"] + rows = _json_log_rows(log_file) + completed = [row for row in rows if row.get("event") == "stream.completed"] + assert len(completed) == 1 + assert completed[0]["request_id"] == "req_complete" + assert completed[0]["stream_chunks"] == 2 + assert completed[0]["outcome"] == "ok" + + +@pytest.mark.asyncio +async def test_traced_async_stream_logs_real_exception(tmp_path) -> None: + log_file = str(tmp_path / "error.log") + configure_logging(log_file, force=True) + + async def source(): + yield "before" + raise RuntimeError("boom") + + with pytest.raises(RuntimeError, match="boom"): + async for _chunk in traced_async_stream( + source(), + stage="egress", + source="unit", + complete_event="stream.completed", + interrupted_event="stream.interrupted", + extra={"request_id": "req_error"}, + ): + pass + + rows = _json_log_rows(log_file) + interrupted = [row for row in rows if row.get("event") == "stream.interrupted"] + assert len(interrupted) == 1 + assert interrupted[0]["request_id"] == "req_error" + assert interrupted[0]["stream_chunks"] == 1 + assert interrupted[0]["outcome"] == "error" + assert interrupted[0]["exc_type"] == "RuntimeError" + + +@pytest.mark.asyncio +async def test_traced_async_stream_closes_quietly_on_generator_exit(tmp_path) -> None: + log_file = str(tmp_path / "generator_exit.log") + configure_logging(log_file, force=True) + + async def source(): + yield "first" + yield "second" + + stream = traced_async_stream( + source(), + stage="egress", + source="unit", + complete_event="stream.completed", + interrupted_event="stream.interrupted", + extra={"request_id": "req_closed"}, + ) + + assert await anext(stream) == "first" + await stream.aclose() + + rows = _json_log_rows(log_file) + events = {row.get("event") for row in rows} + assert "stream.completed" not in events + assert "stream.interrupted" not in events diff --git a/uv.lock b/uv.lock index b57f6370..8284ce2a 100644 --- a/uv.lock +++ b/uv.lock @@ -561,7 +561,7 @@ wheels = [ [[package]] name = "free-claude-code" -version = "2.3.4" +version = "2.3.5" source = { editable = "." } dependencies = [ { name = "aiohttp" },