Fix stream cleanup context handling.

Avoid contextvar-based log context in SSE generators and treat GeneratorExit as quiet teardown.
This commit is contained in:
Alishahryar1
2026-06-17 20:43:38 -07:00
parent da672af337
commit c024bf6892
6 changed files with 191 additions and 98 deletions
+78 -78
View File
@@ -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(
+5 -3
View File
@@ -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,
+9 -11
View File
@@ -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
+1 -1
View File
@@ -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"
+97 -4
View File
@@ -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"] == "<redacted>"
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
Generated
+1 -1
View File
@@ -561,7 +561,7 @@ wheels = [
[[package]]
name = "free-claude-code"
version = "2.3.4"
version = "2.3.5"
source = { editable = "." }
dependencies = [
{ name = "aiohttp" },