项目文件夹

文件
Ali Khokhar af12e7b2bb Coordinate provider recovery under concurrent load (#1205)
## Problem

Concurrent transient failures could start independent retry, replay,
continuation, and repair loops while holding provider concurrency slots.
This multiplied upstream attempts and could delay or strand terminal
errors under fan-out.

## Changes

| Before | After |
| --- | --- |
| Retry paths owned separate attempt budgets. | One logical-execution
session caps all upstream work at five attempts. |
| Concurrent failures backed off independently. | One provider-owned
recovery episode elects a single half-open probe while followers
coalesce. |
| Backoff occupied stream concurrency. | Concurrency is held only while
an upstream operation or stream is active. |
| Provider catalog calls and stream creation used separate admission
paths. | Every upstream operation uses one provider-generation admission
controller. |
| Cancellation could leave recovery ownership or follower state
unresolved. | Cancellation releases permits, transfers probe ownership,
and unregisters waiting followers. |
| Late in-flight failures could cross an exhausted episode boundary. |
Every coalesced execution retains that generation's terminal outcome. |
| Replay tests allowed loose lifecycle assertions. | Exact SSE contracts
prove retries and continuations emit one unduplicated response. |
| Recovery wrappers could mask final diagnostics. | Final responses and
traces retain the raw provider failure and request ID. |

<!-- greptile_comment -->

<details open><summary><h3>Greptile Summary</h3></summary>

This PR coordinates provider recovery and retry work under concurrent
load. The main changes are:

- One five-attempt budget for each logical execution.
- Provider-wide recovery episodes with one elected probe.
- Shared admission for streams, catalog calls, rate limits, and
concurrency.
- Concurrency permits held only during active upstream work.
- Cancellation-safe probe ownership and preserved final diagnostics.
</details>

<h3>Confidence Score: 5/5</h3>

This looks safe to merge.

No blocking issues found in the changed code.

<details><summary><h3><a href="https://www.greptile.com/trex"><img
alt="T-Rex"
src="https://greptile-static-assets.s3.amazonaws.com/trex/trex_green.svg"
height="20" align="absmiddle"></a> T-Rex Logs</h3></summary>

**What T-Rex did**
- Reviewed the coordinated-recovery-01-before.log to understand how the
exhausted generation outcome was not preserved in a late in-flight
failure.
- Reviewed the coordinated-recovery-02-after.log to confirm that the
updated implementation preserves the exhausted generation outcome for
the same focused contract set.
- Validated that the provider-admission-full-current.log shows the
complete requested test file passed under Python 3.14 with uv run pytest
-n 0.

<a
href="https://app.greptile.com/trex/runs/15050270/artifacts"><picture><source
media="(prefers-color-scheme: dark)"
srcset="https://greptile-static-assets.s3.amazonaws.com/badges/ViewAllArtifactsDark.svg?v=4"><source
media="(prefers-color-scheme: light)"
srcset="https://greptile-static-assets.s3.amazonaws.com/badges/ViewAllArtifacts.svg?v=4"><img
alt="View all artifacts"
src="https://greptile-static-assets.s3.amazonaws.com/badges/ViewAllArtifacts.svg?v=4"></picture></a>

<sub><a href="https://www.greptile.com/trex"><img alt="T-Rex"
src="https://greptile-static-assets.s3.amazonaws.com/trex/trex_green.svg"
height="14" align="absmiddle"></a> Ran code and verified through
T-Rex</sub>
</details>

<details open><summary><h3>Important Files Changed</h3></summary>

| Filename | Overview |
|----------|----------|
| src/free_claude_code/providers/admission.py | Adds shared admission,
retry budgets, recovery episodes, probe election, and cancellation
handling. |
| src/free_claude_code/providers/openai_chat/provider.py | Moves stream
creation, replay, continuation, and repair onto one admission-owned
retry session. |
| src/free_claude_code/providers/stream_recovery.py | Selects replay,
continuation, repair, or final failure using the remaining shared
attempt budget. |
| src/free_claude_code/providers/failure_policy.py | Adds recovery
exhaustion handling and preserves the underlying provider error for
final classification. |
| src/free_claude_code/providers/runtime/factory.py | Creates one
admission controller per provider generation and passes it through
provider factories. |

</details>

<details open><summary><h3>Sequence Diagram</h3></summary>

<a href="#gh-light-mode-only">

```mermaid
%%{init: {'theme': 'neutral'}}%%
sequenceDiagram
participant E as Execution
participant A as Admission controller
participant P as Provider
participant F as Concurrent follower

E->>A: Open attempt
A->>P: Send upstream request
P-->>E: Retryable failure
E->>A: Open recovery episode
F->>A: Request admission
A-->>F: Coalesce and wait
E->>A: Claim probe
A->>P: Send half-open probe
alt Probe succeeds
    P-->>E: Valid response
    E->>A: Close recovery episode
    A-->>F: Release waiter
else Probe fails
    P-->>E: Retryable failure
    E->>A: Schedule next probe or finalize error
end
```

</a>
<a href="#gh-dark-mode-only">

```mermaid
%%{init: {'theme': 'base', 'themeVariables': {"darkMode": true, "background": "#0d1117", "primaryColor": "#21262d", "primaryTextColor": "#e6edf3", "primaryBorderColor": "#8b949e", "lineColor": "#8b949e", "textColor": "#e6edf3", "edgeLabelBackground": "#161b22", "actorBkg": "#21262d", "actorBorder": "#8b949e", "actorTextColor": "#e6edf3", "actorLineColor": "#8b949e", "signalColor": "#8b949e", "signalTextColor": "#e6edf3", "noteBkgColor": "#373320", "noteBorderColor": "#d4a72c", "noteTextColor": "#f0e6c0", "labelBoxBkgColor": "#21262d", "labelBoxBorderColor": "#8b949e", "labelTextColor": "#e6edf3", "loopTextColor": "#e6edf3", "activationBkgColor": "#30363d", "activationBorderColor": "#8b949e"}}}%%
sequenceDiagram
participant E as Execution
participant A as Admission controller
participant P as Provider
participant F as Concurrent follower

E->>A: Open attempt
A->>P: Send upstream request
P-->>E: Retryable failure
E->>A: Open recovery episode
F->>A: Request admission
A-->>F: Coalesce and wait
E->>A: Claim probe
A->>P: Send half-open probe
alt Probe succeeds
    P-->>E: Valid response
    E->>A: Close recovery episode
    A-->>F: Release waiter
else Probe fails
    P-->>E: Retryable failure
    E->>A: Schedule next probe or finalize error
end
```

</a>
</details>

<sub>Reviews (2): Last reviewed commit: ["Harden coordinated retry
lifecycle
invar..."](https://github.com/alishahryar1/free-claude-code/commit/2e871c8649d148b5eb71d21f80bf870ae2d11708)
| [Re-trigger
Greptile](https://app.greptile.com/api/retrigger?id=45554917)</sub>

<!-- /greptile_comment -->
2026-07-19 22:43:29 -07:00

1574 行
56 KiB
Python

"""Tests for streaming error handling in providers/nvidia_nim/client.py."""
import json
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import openai
import pytest
from free_claude_code.config.nim import NimSettings
from free_claude_code.core.anthropic.stream_contracts import (
parse_sse_text,
)
from free_claude_code.core.anthropic.streaming import (
AnthropicStreamLedger,
make_text_recovery_body,
)
from free_claude_code.core.failures import ExecutionFailure
from free_claude_code.core.reasoning import DEFAULT_REASONING_POLICY, ReasoningPolicy
from free_claude_code.providers.admission import UPSTREAM_TRANSIENT_TOTAL_ATTEMPTS
from free_claude_code.providers.base import ProviderConfig
from free_claude_code.providers.nvidia_nim import NvidiaNimProvider
from free_claude_code.providers.openai_chat.provider import (
_OpenAIChatStreamRunner,
)
from free_claude_code.providers.openai_chat.tool_calls import (
OpenAIToolCallAssembler,
has_committed_sse_output,
iter_heuristic_tool_use_sse,
)
from free_claude_code.providers.stream_recovery import TruncatedProviderStreamError
from tests.providers.request_factory import make_messages_request
from tests.providers.support import REASONING_OFF, immediate_admission
class AsyncStreamMock:
"""Async iterable mock that yields chunks then optionally raises."""
def __init__(self, chunks, error=None):
self._chunks = chunks
self._error = error
def __aiter__(self):
return self._aiter()
async def _aiter(self):
for chunk in self._chunks:
yield chunk
if self._error:
raise self._error
class ClosableAsyncStreamMock(AsyncStreamMock):
"""Async stream mock that records cleanup."""
def __init__(self, chunks, error=None):
super().__init__(chunks, error=error)
self.closed = False
async def aclose(self):
self.closed = True
def _make_provider():
"""Create a provider instance for testing."""
config = ProviderConfig(
api_key="test_key",
base_url="https://test.api.nvidia.com/v1",
rate_limit=10,
rate_window=60,
)
return NvidiaNimProvider(
config,
nim_settings=NimSettings(),
admission=immediate_admission(),
)
def _make_tool_assembler(provider: NvidiaNimProvider) -> OpenAIToolCallAssembler:
return OpenAIToolCallAssembler(
record_extra_content=provider._record_tool_call_extra_content
)
def _make_request(model: str = "test-model", stream: bool = True, **overrides: object):
"""Create a concrete request matching the original streaming-test defaults."""
request_overrides: dict[str, object] = {
"messages": [],
"max_tokens": 4096,
"temperature": None,
"top_p": None,
"system": None,
"tools": None,
"extra_body": None,
"thinking": None,
"stream": stream,
}
request_overrides.update(overrides)
return make_messages_request(model, **request_overrides)
def _make_stream_runner(
provider: NvidiaNimProvider,
*,
request=None,
request_id: str | None = None,
) -> _OpenAIChatStreamRunner:
return _OpenAIChatStreamRunner(
provider,
request=request or _make_request(),
input_tokens=0,
request_id=request_id,
reasoning=DEFAULT_REASONING_POLICY,
)
def _make_chunk(
content=None, finish_reason=None, tool_calls=None, reasoning_content=None
):
"""Create a mock streaming chunk."""
delta = MagicMock()
delta.content = content
delta.tool_calls = tool_calls
delta.reasoning_content = reasoning_content
choice = MagicMock()
choice.delta = delta
choice.finish_reason = finish_reason
chunk = MagicMock()
chunk.choices = [choice]
chunk.usage = None
return chunk
def _make_tool_calls_chunk(*, name: str, arguments: str, tool_id: str, index: int = 0):
"""Single OpenAI-style tool_calls delta (starts a native streamed tool block)."""
tc = MagicMock()
tc.index = index
tc.id = tool_id
fn = MagicMock()
fn.name = name
fn.arguments = arguments
tc.function = fn
return _make_chunk(tool_calls=[tc])
async def _collect_stream(
provider,
request,
*,
reasoning: ReasoningPolicy = DEFAULT_REASONING_POLICY,
):
"""Collect all SSE events from a stream."""
return [e async for e in provider.stream_response(request, reasoning=reasoning)]
async def _collect_stream_error(provider, request, **kwargs) -> ExecutionFailure:
with pytest.raises(ExecutionFailure) as exc_info:
[e async for e in provider.stream_response(request, **kwargs)]
return exc_info.value
async def _collect_stream_and_error(
provider, request, **kwargs
) -> tuple[list[str], ExecutionFailure]:
events: list[str] = []
with pytest.raises(ExecutionFailure) as exc_info:
async for event in provider.stream_response(request, **kwargs):
events.extend((event,))
return events, exc_info.value
def _assert_no_content_deltas_after_error_text(
events: list[str], error_substr: str
) -> None:
"""After the error text delta, only block close + message tail events may follow."""
parsed = parse_sse_text("".join(events))
first_error_idx = None
for i, ev in enumerate(parsed):
if ev.event != "content_block_delta":
continue
delta = ev.data.get("delta", {})
if delta.get("type") == "text_delta" and error_substr in str(
delta.get("text", "")
):
first_error_idx = i
break
assert first_error_idx is not None, (error_substr, "".join(events))
for ev in parsed[first_error_idx + 1 :]:
assert ev.event in ("content_block_stop", "message_delta", "message_stop"), (
ev.event,
ev.data,
)
def _assert_error_not_in_text_deltas_after_tool(
events: list[str], error_substr: str
) -> None:
"""Transport errors after a native tool call must not use assistant text_delta (issue #206)."""
blob = "".join(events)
for ev in parse_sse_text(blob):
if ev.event != "content_block_delta":
continue
delta = ev.data.get("delta", {})
if delta.get("type") == "text_delta" and error_substr in str(
delta.get("text", "")
):
raise AssertionError(
f"error leaked as text_delta after tool_use: {ev.data!r} full={blob!r}"
)
class TestStreamingExceptionHandling:
@pytest.mark.asyncio
async def test_stream_normalization_failure_closes_raw_stream(self):
provider = _make_provider()
stream = ClosableAsyncStreamMock([])
retry_session = provider._admission.new_retry_session()
with (
patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
return_value=stream,
),
patch.object(
provider,
"_normalize_stream",
side_effect=ValueError("invalid stream wrapper"),
),
pytest.raises(ValueError, match="invalid stream wrapper"),
):
await provider._create_stream({"messages": []}, retry_session)
assert stream.closed
"""Tests for error paths during stream_response."""
@pytest.mark.asyncio
async def test_pre_start_api_error_raises_provider_error(self):
"""Before holdback commit, provider failures raise for API-level non-200."""
provider = _make_provider()
request = _make_request()
with (
patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
side_effect=RuntimeError("API failed"),
),
):
error = await _collect_stream_error(provider, request)
assert "API failed" in error.message
@pytest.mark.asyncio
async def test_read_timeout_with_empty_message_raises_fallback(self):
"""ReadTimeout(TimeoutError()) should raise a non-empty timeout message."""
provider = _make_provider()
request = _make_request()
with (
patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
side_effect=httpx.ReadTimeout(""),
),
patch("asyncio.sleep", new_callable=AsyncMock),
):
error = await _collect_stream_error(
provider,
request,
request_id="req_timeout123",
)
assert "timed out after" in error.message
assert "Request ID: req_timeout123" in error.message
@pytest.mark.asyncio
async def test_error_after_precommit_partial_content_raises(self):
"""Precommit partial text is discarded so the API can return non-200."""
provider = _make_provider()
request = _make_request()
chunk1 = _make_chunk(content="Hello ")
stream_mock = AsyncStreamMock([chunk1], error=RuntimeError("Connection lost"))
with (
patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
return_value=stream_mock,
),
):
error = await _collect_stream_error(provider, request)
assert "Connection lost" in error.message
@pytest.mark.asyncio
async def test_error_after_native_tool_call_closes_block_then_raises(self):
"""A provider closes tool state, then leaves terminal serialization to API."""
provider = _make_provider()
request = _make_request()
tool_chunk = _make_tool_calls_chunk(
name="echo_smoke", arguments="{}", tool_id="call_206", index=0
)
stream_mock = AsyncStreamMock(
[tool_chunk], error=RuntimeError("Connection lost after tool")
)
with (
patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
return_value=stream_mock,
),
):
events, error = await _collect_stream_and_error(provider, request)
event_text = "".join(events)
parsed = parse_sse_text(event_text)
assert "tool_use" in event_text
assert parsed[-1].event == "content_block_stop"
assert "Connection lost after tool" in error.message
assert "Connection lost after tool" not in event_text
assert "event: error\n" not in event_text
assert "message_stop" not in event_text
_assert_error_not_in_text_deltas_after_tool(
events, "Connection lost after tool"
)
@pytest.mark.asyncio
async def test_empty_response_gets_space(self):
"""Empty response with no text/tools gets a single space text block."""
provider = _make_provider()
request = _make_request()
empty_chunk = _make_chunk(finish_reason="stop")
stream_mock = AsyncStreamMock([empty_chunk])
with (
patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
return_value=stream_mock,
),
):
events = await _collect_stream(provider, request)
event_text = "".join(events)
assert '"text_delta"' in event_text
assert "message_stop" in event_text
@pytest.mark.asyncio
async def test_upstream_completion_tokens_null_emits_int_usage(self):
"""NIM/GLM may send usage.completion_tokens=null; final SSE must not use JSON null."""
provider = _make_provider()
request = _make_request()
delta = SimpleNamespace(
content="hello",
tool_calls=None,
reasoning_content=None,
)
choice = SimpleNamespace(delta=delta, finish_reason="stop")
usage = SimpleNamespace(completion_tokens=None, prompt_tokens=None)
chunk = SimpleNamespace(choices=[choice], usage=usage)
stream_mock = AsyncStreamMock([chunk])
with (
patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
return_value=stream_mock,
),
):
events = await _collect_stream(provider, request)
parsed = parse_sse_text("".join(events))
delta_events = [e for e in parsed if e.event == "message_delta"]
assert len(delta_events) == 1
usage_out = delta_events[0].data.get("usage", {})
assert isinstance(usage_out.get("output_tokens"), int)
assert usage_out["output_tokens"] is not None
assert '"output_tokens": null' not in "".join(events)
@pytest.mark.asyncio
async def test_reasoning_only_stream_emits_placeholder_text(self):
"""When the model streams only ``reasoning_content`` (no ``content``), add text block.
NIM / some templates may emit no main ``content``; a minimal text block matches
the empty-body placeholder and helps clients that expect a text segment.
"""
provider = _make_provider()
request = _make_request()
chunk1 = _make_chunk(reasoning_content="reasoning only from provider")
chunk2 = _make_chunk(finish_reason="stop")
stream_mock = AsyncStreamMock([chunk1, chunk2])
with (
patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
return_value=stream_mock,
),
):
events = await _collect_stream(provider, request)
event_text = "".join(events)
assert "thinking_delta" in event_text
assert '"text_delta"' in event_text
assert "message_stop" in event_text
@pytest.mark.asyncio
async def test_stream_with_thinking_content(self):
"""Thinking content via think tags is emitted as thinking blocks."""
provider = _make_provider()
request = _make_request()
chunk1 = _make_chunk(content="<think>reasoning</think>answer")
chunk2 = _make_chunk(finish_reason="stop")
stream_mock = AsyncStreamMock([chunk1, chunk2])
with (
patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
return_value=stream_mock,
),
):
events = await _collect_stream(provider, request)
event_text = "".join(events)
assert "thinking" in event_text
assert "reasoning" in event_text
assert "answer" in event_text
@pytest.mark.asyncio
async def test_stream_with_reasoning_content_field(self):
"""reasoning_content delta field is emitted as thinking block."""
provider = _make_provider()
request = _make_request()
chunk1 = _make_chunk(reasoning_content="I think...")
chunk2 = _make_chunk(content="The answer")
chunk3 = _make_chunk(finish_reason="stop")
stream_mock = AsyncStreamMock([chunk1, chunk2, chunk3])
with (
patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
return_value=stream_mock,
),
):
events = await _collect_stream(provider, request)
event_text = "".join(events)
assert "thinking_delta" in event_text
assert "I think..." in event_text
assert "The answer" in event_text
@pytest.mark.asyncio
async def test_stream_with_empty_reasoning_content_starts_thinking_block_only(self):
"""Empty reasoning_content is stateful but must not emit visible thinking text."""
provider = _make_provider()
request = _make_request()
chunk1 = _make_chunk(reasoning_content="")
chunk2 = _make_chunk(finish_reason="stop")
stream_mock = AsyncStreamMock([chunk1, chunk2])
with (
patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
return_value=stream_mock,
),
):
events = await _collect_stream(provider, request)
parsed = parse_sse_text("".join(events))
thinking_starts = [
event
for event in parsed
if event.event == "content_block_start"
and event.data["content_block"]["type"] == "thinking"
]
thinking_deltas = [
event
for event in parsed
if event.event == "content_block_delta"
and event.data["delta"]["type"] == "thinking_delta"
]
assert len(thinking_starts) == 1
assert thinking_deltas == []
assert parsed[-1].event == "message_stop"
@pytest.mark.asyncio
async def test_stream_with_reasoning_content_suppressed_when_disabled(self):
"""reasoning deltas are stripped while normal text still streams."""
provider = _make_provider()
request = _make_request()
chunk1 = _make_chunk(reasoning_content="I think...")
chunk2 = _make_chunk(content="<think>secret</think>The answer")
chunk3 = _make_chunk(finish_reason="stop")
stream_mock = AsyncStreamMock([chunk1, chunk2, chunk3])
with (
patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
return_value=stream_mock,
),
):
events = await _collect_stream(provider, request, reasoning=REASONING_OFF)
event_text = "".join(events)
assert "thinking_delta" not in event_text
assert "I think..." not in event_text
assert "secret" not in event_text
assert "The answer" in event_text
@pytest.mark.asyncio
async def test_stream_with_upstream_405_mentions_provider_name(self):
"""HTTP 405s are surfaced as upstream method/endpoint rejections."""
provider = _make_provider()
request = _make_request()
response = httpx.Response(
status_code=405,
request=httpx.Request("POST", "https://example.com/v1/chat/completions"),
)
error = httpx.HTTPStatusError(
"Method Not Allowed",
request=response.request,
response=response,
)
with patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
side_effect=error,
):
stream_error = await _collect_stream_error(
provider,
request,
request_id="REQ405",
)
assert (
"Upstream provider NIM rejected the request method or endpoint (HTTP 405)."
in stream_error.message
)
assert "Request ID: REQ405" in stream_error.message
@pytest.mark.asyncio
async def test_stream_with_openai_bad_request_surfaces_upstream_body(self):
"""OpenAI SDK bodies should be raised so users can copy exact provider errors."""
provider = _make_provider()
request = _make_request()
response = httpx.Response(
status_code=400,
request=httpx.Request("POST", "https://example.com/v1/chat/completions"),
)
body = {
"error": {
"type": "BadRequest",
"message": "Thinking mode does not support this tool_choice",
}
}
error = openai.BadRequestError("Bad Request", response=response, body=body)
with patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
side_effect=error,
):
stream_error = await _collect_stream_error(
provider,
request,
request_id="REQ_BODY",
)
assert "Upstream provider NIM returned HTTP 400." in stream_error.message
assert "Category: BadRequest" in stream_error.message
assert "Thinking mode does not support this tool_choice" in stream_error.message
assert (
'{"error":{"type":"BadRequest","message":"Thinking mode does not support this tool_choice"}}'
in stream_error.message
)
assert "Request ID: REQ_BODY" in stream_error.message
@pytest.mark.asyncio
async def test_error_after_native_tool_call_failure_includes_body(self):
"""Detailed failure data survives after the provider closes tool state."""
provider = _make_provider()
request = _make_request()
tool_chunk = _make_tool_calls_chunk(
name="echo_smoke", arguments="{}", tool_id="call_body", index=0
)
response = httpx.Response(
status_code=400,
request=httpx.Request("POST", "https://example.com/v1/chat/completions"),
)
body = {"error": {"message": "bad after tool"}}
error = openai.BadRequestError("Bad Request", response=response, body=body)
stream_mock = AsyncStreamMock([tool_chunk], error=error)
with patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
return_value=stream_mock,
):
events, stream_error = await _collect_stream_and_error(
provider,
request,
request_id="REQ_TOOL_BODY",
)
event_text = "".join(events)
parsed = parse_sse_text(event_text)
assert "tool_use" in event_text
assert parsed[-1].event == "content_block_stop"
assert "event: error\n" not in event_text
assert "bad after tool" not in event_text
assert "Request ID: REQ_TOOL_BODY" not in event_text
assert "message_stop" not in event_text
assert "bad after tool" in stream_error.message
assert "Request ID: REQ_TOOL_BODY" in stream_error.message
_assert_error_not_in_text_deltas_after_tool(events, "bad after tool")
@pytest.mark.asyncio
async def test_clean_eof_after_complete_tool_call_salvages_tool_use(self):
"""A complete tool JSON payload missing finish_reason is committed as tool_use."""
provider = _make_provider()
request = _make_request()
tool_chunk = _make_tool_calls_chunk(
name="echo_smoke", arguments='{"message":"ok"}', tool_id="call_eof"
)
stream_mock = AsyncStreamMock([tool_chunk])
with patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
return_value=stream_mock,
):
events = await _collect_stream(provider, request)
parsed = parse_sse_text("".join(events))
assert parsed[-1].event == "message_stop"
assert any(
event.event == "message_delta"
and event.data.get("delta", {}).get("stop_reason") == "tool_use"
for event in parsed
)
assert not any(event.event == "error" for event in parsed)
@pytest.mark.asyncio
@pytest.mark.parametrize("finish_reason", ["tool_calls", "stop"])
async def test_heuristic_only_tool_stream_does_not_emit_fallback_text(
self, finish_reason
):
"""Text-parsed tool calls count as emitted tool output when finalizing."""
provider = _make_provider()
request = _make_request()
heuristic_tool = (
"● <function=Read><parameter=path>test.py</parameter>"
"<parameter=limit>10</parameter>"
)
stream_mock = AsyncStreamMock(
[
_make_chunk(content=heuristic_tool),
_make_chunk(finish_reason=finish_reason),
]
)
with patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
return_value=stream_mock,
):
events = await _collect_stream(provider, request)
parsed = parse_sse_text("".join(events))
assert any(
event.event == "content_block_start"
and event.data.get("content_block", {}).get("type") == "tool_use"
for event in parsed
)
assert not any(
event.event == "content_block_delta"
and event.data.get("delta", {}).get("type") == "text_delta"
and event.data.get("delta", {}).get("text") == " "
for event in parsed
)
assert any(
event.event == "message_delta"
and event.data.get("delta", {}).get("stop_reason") == "tool_use"
for event in parsed
)
@pytest.mark.asyncio
async def test_precommit_retry_emits_one_unduplicated_downstream_lifecycle(self):
"""An abandoned attempt contributes no frame to the successful replay."""
provider = _make_provider()
request = _make_request()
first_stream = AsyncStreamMock(
[_make_chunk(content="hidden")],
error=httpx.ReadError("early cutoff"),
)
second_stream = AsyncStreamMock(
[
_make_chunk(content="visible"),
_make_chunk(finish_reason="stop"),
]
)
with patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
side_effect=[first_stream, second_stream],
) as mock_create:
events = await _collect_stream(provider, request)
event_text = "".join(events)
assert mock_create.await_count == 2
assert "hidden" not in event_text
parsed = parse_sse_text(event_text)
text_deltas = [
event.data.get("delta", {}).get("text", "")
for event in parsed
if event.event == "content_block_delta"
]
assert text_deltas == ["visible"]
assert sum(event.event == "message_start" for event in parsed) == 1
assert sum(event.event == "content_block_start" for event in parsed) == 1
assert sum(event.event == "content_block_stop" for event in parsed) == 1
assert sum(event.event == "message_delta" for event in parsed) == 1
assert sum(event.event == "message_stop" for event in parsed) == 1
assert parsed[0].event == "message_start"
assert parsed[-1].event == "message_stop"
@pytest.mark.asyncio
async def test_primary_replay_and_continuation_share_five_attempts(self):
"""Four replays plus continuation emit one unduplicated response."""
provider = _make_provider()
request = _make_request()
primary_streams = [
AsyncStreamMock([_make_chunk(content="hello")]) for _ in range(4)
]
continuation = AsyncStreamMock(
[
_make_chunk(content="hello world"),
_make_chunk(finish_reason="stop"),
]
)
with patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
side_effect=[*primary_streams, continuation],
) as create:
events = await _collect_stream(provider, request)
assert create.await_count == UPSTREAM_TRANSIENT_TOTAL_ATTEMPTS
assert all(
call.kwargs["messages"] == create.await_args_list[0].kwargs["messages"]
for call in create.await_args_list[:4]
)
assert (
create.await_args_list[4].kwargs["messages"]
!= (create.await_args_list[0].kwargs["messages"])
)
parsed = parse_sse_text("".join(events))
text = "".join(
event.data.get("delta", {}).get("text", "")
for event in parsed
if event.event == "content_block_delta"
)
assert text == "hello world"
assert sum(event.event == "message_start" for event in parsed) == 1
assert sum(event.event == "message_delta" for event in parsed) == 1
assert sum(event.event == "message_stop" for event in parsed) == 1
@pytest.mark.asyncio
async def test_clean_eof_after_text_continues_with_overlap_trim(self):
"""A truncated text stream is continued and duplicate overlap is trimmed."""
provider = _make_provider()
request = _make_request()
stream_mock = AsyncStreamMock([_make_chunk(content="hello wor")])
with (
patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
return_value=stream_mock,
),
patch.object(
_OpenAIChatStreamRunner,
"_collect_recovery_text",
new_callable=AsyncMock,
return_value=("world", ""),
),
):
events = await _collect_stream(provider, request)
parsed = parse_sse_text("".join(events))
text_deltas = [
event.data.get("delta", {}).get("text", "")
for event in parsed
if event.event == "content_block_delta"
]
assert text_deltas == ["hello wor", "ld"]
assert "".join(text_deltas) == "hello world"
assert sum(event.event == "message_start" for event in parsed) == 1
assert sum(event.event == "content_block_start" for event in parsed) == 1
assert sum(event.event == "content_block_stop" for event in parsed) == 1
assert sum(event.event == "message_delta" for event in parsed) == 1
assert sum(event.event == "message_stop" for event in parsed) == 1
assert any(
event.event == "message_delta"
and event.data.get("delta", {}).get("stop_reason") == "end_turn"
for event in parsed
)
assert not any(event.event == "error" for event in parsed)
@pytest.mark.asyncio
async def test_disabled_thinking_recovery_discards_reasoning(self):
provider = _make_provider()
request = _make_request()
initial_stream = AsyncStreamMock([_make_chunk(content="hello")])
recovery_stream = AsyncStreamMock(
[
_make_chunk(reasoning_content="hidden reasoning"),
_make_chunk(content="hello world"),
_make_chunk(finish_reason="stop"),
]
)
with patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
side_effect=[initial_stream, recovery_stream],
):
events = await _collect_stream(provider, request, reasoning=REASONING_OFF)
parsed = parse_sse_text("".join(events))
text = "".join(
event.data.get("delta", {}).get("text", "")
for event in parsed
if event.event == "content_block_delta"
)
assert text == "hello world"
assert "hidden reasoning" not in "".join(events)
assert not any(
event.data.get("delta", {}).get("type") == "thinking_delta"
for event in parsed
)
@pytest.mark.asyncio
async def test_recovery_collect_text_requires_finish_reason(self):
"""Recovery collectors reject truncated OpenAI-chat continuation streams."""
streams = [
ClosableAsyncStreamMock([_make_chunk(content=f"world {index}")])
for index in range(UPSTREAM_TRANSIENT_TOTAL_ATTEMPTS)
]
provider = _make_provider()
runner = _make_stream_runner(provider)
retry_session = provider._admission.new_retry_session()
with (
patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
side_effect=streams,
) as create,
pytest.raises(TruncatedProviderStreamError),
):
await runner._collect_recovery_text(
{"messages": []},
include_reasoning=True,
retry_session=retry_session,
)
assert create.await_count == UPSTREAM_TRANSIENT_TOTAL_ATTEMPTS
assert all(stream.closed for stream in streams)
@pytest.mark.asyncio
async def test_recovery_collect_text_closes_retryable_failed_streams(self):
"""Recovery collectors close failed stream attempts before retrying."""
streams = [
ClosableAsyncStreamMock(
[_make_chunk(content=f"partial {index}")],
error=TimeoutError("recovery cutoff"),
)
for index in range(UPSTREAM_TRANSIENT_TOTAL_ATTEMPTS)
]
provider = _make_provider()
runner = _make_stream_runner(provider)
retry_session = provider._admission.new_retry_session()
with (
patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
side_effect=streams,
) as create,
pytest.raises(TimeoutError),
):
await runner._collect_recovery_text(
{"messages": []},
include_reasoning=True,
retry_session=retry_session,
)
assert create.await_count == UPSTREAM_TRANSIENT_TOTAL_ATTEMPTS
assert all(stream.closed for stream in streams)
@pytest.mark.asyncio
async def test_recovery_collect_text_accepts_finish_reason(self):
"""Recovery collectors return text only after the upstream terminal marker."""
stream = ClosableAsyncStreamMock(
[
_make_chunk(content="world"),
_make_chunk(finish_reason="stop"),
]
)
provider = _make_provider()
runner = _make_stream_runner(provider)
retry_session = provider._admission.new_retry_session()
with patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
return_value=stream,
):
result = await runner._collect_recovery_text(
{"messages": []},
include_reasoning=True,
retry_session=retry_session,
)
assert result == ("world", "")
assert stream.closed is True
@pytest.mark.asyncio
async def test_recovery_collect_text_honors_provider_retry_classification(self):
"""Provider semantics apply before the first recovery chunk as well."""
provider = _make_provider()
runner = _make_stream_runner(provider)
retry_session = provider._admission.new_retry_session()
request = httpx.Request(
"POST", "https://test.api.nvidia.com/v1/chat/completions"
)
degraded = openai.BadRequestError(
"Bad Request",
response=httpx.Response(400, request=request),
body={
"status": 400,
"detail": (
"Function id 'test-function': DEGRADED function cannot be invoked"
),
},
)
rejected = ClosableAsyncStreamMock([], error=degraded)
recovered = ClosableAsyncStreamMock(
[
_make_chunk(content="world"),
_make_chunk(finish_reason="stop"),
]
)
with patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
side_effect=[rejected, recovered],
) as create:
result = await runner._collect_recovery_text(
{"messages": []},
include_reasoning=True,
retry_session=retry_session,
)
assert result == ("world", "")
assert create.await_count == 2
assert rejected.closed
assert recovered.closed
def test_text_recovery_body_preserves_thinking_context(self):
"""Continuation prompts include emitted thinking without provider-specific fields."""
body = {
"messages": [{"role": "user", "content": "hello"}],
"tools": [{"name": "Read"}],
"tool_choice": {"type": "auto"},
}
recovery_body = make_text_recovery_body(
body,
partial_text="visible answer",
partial_thinking="hidden reasoning",
)
assert "tools" not in recovery_body
assert "tool_choice" not in recovery_body
assert "stream" not in recovery_body
assert recovery_body["messages"][-2] == {
"role": "assistant",
"content": "visible answer",
}
recovery_prompt = recovery_body["messages"][-1]
assert recovery_prompt["role"] == "user"
assert "hidden reasoning" in recovery_prompt["content"]
assert "reasoning_content" not in recovery_prompt
@pytest.mark.asyncio
async def test_openai_text_recovery_passes_thinking_context(self):
"""OpenAI-chat recovery call sites seed emitted thinking in the prompt."""
runner = _make_stream_runner(
_make_provider(), request=_make_request(), request_id="req_recovery"
)
ledger = AnthropicStreamLedger("msg_recovery", "model")
ledger.start_thinking_block()
ledger.emit_thinking_delta("hidden reasoning")
list(ledger.ensure_text_block())
ledger.emit_text_delta("visible answer")
with patch.object(
runner,
"_collect_recovery_text",
new_callable=AsyncMock,
return_value=("visible answer done", "hidden reasoning more"),
) as mock_collect:
retry_session = runner._provider._admission.new_retry_session()
events = await runner._recovery_events(
body={"messages": [{"role": "user", "content": "hello"}]},
ledger=ledger,
error=TimeoutError("cutoff"),
tool_argument_alias_buffers={},
output_reasoning=True,
retry_session=retry_session,
)
assert events is not None
assert mock_collect.await_args is not None
recovery_body = mock_collect.await_args.args[0]
assert "hidden reasoning" in recovery_body["messages"][-1]["content"]
assert mock_collect.await_args.kwargs["include_reasoning"] is True
@pytest.mark.asyncio
async def test_primary_stream_closes_when_iteration_fails(self):
"""OpenAI-chat main streams close after iterator failures."""
provider = _make_provider()
request = _make_request()
stream = ClosableAsyncStreamMock(
[_make_chunk(content="partial")],
error=ValueError("provider stream failed"),
)
with patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
return_value=stream,
):
error = await _collect_stream_error(provider, request)
assert stream.closed is True
assert "provider stream failed" in error.message.lower()
@pytest.mark.asyncio
async def test_truncated_recovery_stream_closes_block_then_raises(self):
"""Partial recovery bytes never become success or provider-owned wire errors."""
provider = _make_provider()
request = _make_request()
original_text = "hello wor" + ("x" * 70_000)
original_stream = AsyncStreamMock([_make_chunk(content=original_text)])
with (
patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
return_value=original_stream,
) as mock_create,
patch.object(
_OpenAIChatStreamRunner,
"_collect_recovery_text",
new_callable=AsyncMock,
side_effect=TruncatedProviderStreamError(
"Recovery stream ended without finish_reason."
),
) as mock_collect,
):
events, error = await _collect_stream_and_error(provider, request)
event_text = "".join(events)
assert mock_create.await_count == 1
assert mock_collect.await_count == 1
assert original_text in event_text
assert "world" not in event_text
assert "Provider stream ended without finish_reason." in error.message
assert "Provider stream ended without finish_reason." not in event_text
parsed = parse_sse_text(event_text)
assert parsed[-1].event == "content_block_stop"
assert not any(event.event == "error" for event in parsed)
assert not any(event.event == "message_stop" for event in parsed)
assert not any(
event.event == "content_block_delta"
and event.data.get("delta", {}).get("text") == "ld"
for event in parse_sse_text(event_text)
)
@pytest.mark.asyncio
async def test_incomplete_tool_call_repair_appends_schema_valid_suffix(self):
"""A truncated tool JSON prefix is repaired append-only before tool_use tail."""
provider = _make_provider()
request = _make_request(
tools=[
{
"name": "echo_smoke",
"description": "Echo",
"input_schema": {
"type": "object",
"properties": {"message": {"type": "string"}},
"required": ["message"],
"additionalProperties": False,
},
}
]
)
tool_chunk = _make_tool_calls_chunk(
name="echo_smoke", arguments='{"message":', tool_id="call_repair"
)
stream_mock = AsyncStreamMock([tool_chunk])
with (
patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
return_value=stream_mock,
),
patch.object(
_OpenAIChatStreamRunner,
"_collect_recovery_text",
new_callable=AsyncMock,
return_value=('"ok"}', ""),
),
):
events = await _collect_stream(provider, request)
event_text = "".join(events)
parsed = parse_sse_text(event_text)
assert '"partial_json": "\\"ok\\"}"' in event_text
assert any(
event.event == "message_delta"
and event.data.get("delta", {}).get("stop_reason") == "tool_use"
for event in parsed
)
assert not any(event.event == "error" for event in parsed)
@pytest.mark.asyncio
async def test_stream_rate_limit_uses_the_execution_retry_session(self):
"""A create-time 429 consumes one attempt before a successful retry."""
provider = _make_provider()
request = _make_request()
chunk1 = _make_chunk(content="Response")
chunk2 = _make_chunk(finish_reason="stop")
stream_mock = AsyncStreamMock([chunk1, chunk2])
response = httpx.Response(
429,
request=httpx.Request(
"POST", "https://test.api.nvidia.com/v1/chat/completions"
),
)
error = httpx.HTTPStatusError(
"rate limited",
request=response.request,
response=response,
)
with patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
side_effect=[error, stream_mock],
) as create:
events = await _collect_stream(provider, request)
event_text = "".join(events)
assert create.await_count == 2
assert "Response" in event_text
class TestProcessToolCall:
"""Tests for OpenAI tool-call assembly."""
def test_heuristic_tool_use_sse_marks_committed_tool_output(self):
"""Heuristic tool blocks are emitted content, even without OpenAI tool state."""
from free_claude_code.core.anthropic import AnthropicStreamLedger
ledger = AnthropicStreamLedger("msg_test", "test-model")
events = list(
iter_heuristic_tool_use_sse(
ledger,
{
"id": "toolu_heuristic",
"name": "Read",
"input": {"path": "test.py"},
},
)
)
event_text = "".join(events)
assert "tool_use" in event_text
assert ledger.has_emitted_tool_block()
assert has_committed_sse_output(ledger)
def test_tool_call_with_id(self):
"""Tool call with id starts a tool block."""
provider = _make_provider()
from free_claude_code.core.anthropic import AnthropicStreamLedger
sse = AnthropicStreamLedger("msg_test", "test-model")
tc = {
"index": 0,
"id": "call_123",
"function": {"name": "search", "arguments": '{"q": "test"}'},
}
events = list(_make_tool_assembler(provider).process_tool_call(tc, sse))
event_text = "".join(events)
assert "tool_use" in event_text
assert "search" in event_text
assert "call_123" in event_text
def test_tool_call_id_arrives_before_name_still_emits_id_and_name(self):
"""Split-stream tool: id (no name) then name then args; id preserved on start."""
provider = _make_provider()
from free_claude_code.core.anthropic import AnthropicStreamLedger
sse = AnthropicStreamLedger("msg_test", "test-model")
t1 = {
"index": 0,
"id": "call_split",
"function": {"name": None, "arguments": ""},
}
t2 = {
"index": 0,
"id": "call_split",
"function": {"name": "Grep", "arguments": ""},
}
t3 = {
"index": 0,
"id": "call_split",
"function": {"name": None, "arguments": "{}"},
}
b1 = "".join(_make_tool_assembler(provider).process_tool_call(t1, sse))
b2 = "".join(_make_tool_assembler(provider).process_tool_call(t2, sse))
b3 = "".join(_make_tool_assembler(provider).process_tool_call(t3, sse))
combined = b1 + b2 + b3
assert "call_split" in combined
assert "Grep" in combined
assert b1 == ""
def test_tool_call_arguments_buffered_until_name(self):
"""Argument deltas before tool name are emitted after the block starts."""
provider = _make_provider()
from free_claude_code.core.anthropic import AnthropicStreamLedger
sse = AnthropicStreamLedger("msg_test", "test-model")
t1 = {
"index": 0,
"id": "call_buf",
"function": {"name": None, "arguments": '{"x":'},
}
t2 = {
"index": 0,
"id": "call_buf",
"function": {"name": "Read", "arguments": "1}"},
}
b1 = "".join(_make_tool_assembler(provider).process_tool_call(t1, sse))
b2 = "".join(_make_tool_assembler(provider).process_tool_call(t2, sse))
assert b1 == ""
combined = b2
assert "Read" in combined
assert "call_buf" in combined
assert '{"x":' in combined or "partial_json" in combined
def test_tool_call_without_id_generates_uuid(self):
"""Tool call without id generates a uuid-based id."""
provider = _make_provider()
from free_claude_code.core.anthropic import AnthropicStreamLedger
sse = AnthropicStreamLedger("msg_test", "test-model")
tc = {
"index": 0,
"id": None,
"function": {"name": "test", "arguments": "{}"},
}
events = list(_make_tool_assembler(provider).process_tool_call(tc, sse))
event_text = "".join(events)
assert "tool_" in event_text
def test_task_tool_forces_background_false(self):
"""Task tool with run_in_background=true is forced to false."""
provider = _make_provider()
from free_claude_code.core.anthropic import AnthropicStreamLedger
sse = AnthropicStreamLedger("msg_test", "test-model")
args = json.dumps({"run_in_background": True, "prompt": "test"})
tc = {
"index": 0,
"id": "call_task",
"function": {"name": "Task", "arguments": args},
}
events = list(_make_tool_assembler(provider).process_tool_call(tc, sse))
event_text = "".join(events)
# The intercepted args should have run_in_background=false
assert "false" in event_text.lower()
def test_task_tool_chunked_args_forces_background_false(self):
"""Chunked Task args are buffered until valid JSON, then forced to false."""
provider = _make_provider()
from free_claude_code.core.anthropic import AnthropicStreamLedger
sse = AnthropicStreamLedger("msg_test", "test-model")
tc1 = {
"index": 0,
"id": "call_task_chunked",
"function": {"name": "Task", "arguments": '{"run_in_background": true,'},
}
tc2 = {
"index": 0,
"id": "call_task_chunked",
"function": {"name": None, "arguments": ' "prompt": "test"}'},
}
events1 = list(_make_tool_assembler(provider).process_tool_call(tc1, sse))
assert len(events1) > 0
assert "false" not in "".join(events1).lower()
events2 = list(_make_tool_assembler(provider).process_tool_call(tc2, sse))
event_text = "".join(events1 + events2)
assert "false" in event_text.lower()
def test_task_tool_invalid_json_logs_warning_on_flush(self, caplog):
"""Invalid JSON args for Task tool emits {} on flush and logs a warning."""
provider = _make_provider()
from free_claude_code.core.anthropic import AnthropicStreamLedger
sse = AnthropicStreamLedger("msg_test", "test-model")
tc = {
"index": 0,
"id": "call_task2",
"function": {"name": "Task", "arguments": "not json"},
}
events = list(_make_tool_assembler(provider).process_tool_call(tc, sse))
assert len(events) > 0
with caplog.at_level("WARNING"):
flushed = list(_make_tool_assembler(provider).flush_task_arg_buffers(sse))
assert len(flushed) > 0
assert "{}" in "".join(flushed)
assert any("Task args invalid JSON" in r.message for r in caplog.records)
def test_negative_tool_index_fallback(self):
"""tc_index < 0 uses len(tool_indices) as fallback."""
provider = _make_provider()
from free_claude_code.core.anthropic import AnthropicStreamLedger
sse = AnthropicStreamLedger("msg_test", "test-model")
tc = {
"index": -1,
"id": "call_neg",
"function": {"name": "test", "arguments": "{}"},
}
events = list(_make_tool_assembler(provider).process_tool_call(tc, sse))
# Should not crash, should still emit events
assert len(events) > 0
def test_none_tool_index_defaults_to_zero(self):
"""Gemini may stream tool_call deltas with a null index."""
provider = _make_provider()
from free_claude_code.core.anthropic import AnthropicStreamLedger
sse = AnthropicStreamLedger("msg_test", "test-model")
tc = {
"index": None,
"id": "call_none",
"function": {"name": "test", "arguments": "{}"},
}
events = list(_make_tool_assembler(provider).process_tool_call(tc, sse))
event_text = "".join(events)
assert "tool_use" in event_text
assert "call_none" in event_text
def test_tool_args_emitted_as_delta(self):
"""Arguments are emitted as input_json_delta events."""
provider = _make_provider()
from free_claude_code.core.anthropic import AnthropicStreamLedger
sse = AnthropicStreamLedger("msg_test", "test-model")
tc = {
"index": 0,
"id": "call_args",
"function": {"name": "grep", "arguments": '{"pattern": "test"}'},
}
events = list(_make_tool_assembler(provider).process_tool_call(tc, sse))
event_text = "".join(events)
assert "input_json_delta" in event_text
class TestStreamChunkEdgeCases:
"""Tests for edge cases in stream chunk handling."""
@pytest.mark.asyncio
async def test_stream_chunk_with_empty_choices_skipped(self):
"""Chunk with choices=[] is skipped without crashing."""
provider = _make_provider()
request = _make_request()
empty_choices_chunk = MagicMock()
empty_choices_chunk.choices = []
empty_choices_chunk.usage = None
finish_chunk = _make_chunk(finish_reason="stop")
stream_mock = AsyncStreamMock([empty_choices_chunk, finish_chunk])
with (
patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
return_value=stream_mock,
),
):
events = await _collect_stream(provider, request)
event_text = "".join(events)
assert "message_start" in event_text
assert "message_stop" in event_text
@pytest.mark.asyncio
async def test_stream_chunk_with_none_delta_handled(self):
"""Chunk with choice.delta=None is handled defensively."""
provider = _make_provider()
request = _make_request()
none_delta_chunk = MagicMock()
none_delta_chunk.usage = None
choice = MagicMock()
choice.delta = None
choice.finish_reason = None
none_delta_chunk.choices = [choice]
finish_chunk = _make_chunk(finish_reason="stop")
stream_mock = AsyncStreamMock([none_delta_chunk, finish_chunk])
with (
patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
return_value=stream_mock,
),
):
events = await _collect_stream(provider, request)
event_text = "".join(events)
assert "message_start" in event_text
assert "message_stop" in event_text
@pytest.mark.asyncio
async def test_stream_generator_cleanup_on_exception(self):
"""When stream raises mid-iteration, message_stop still emitted."""
provider = _make_provider()
request = _make_request()
chunk1 = _make_chunk(content="Partial")
stream_mock = AsyncStreamMock(
[chunk1], error=ConnectionResetError("Connection reset")
)
with (
patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
return_value=stream_mock,
),
):
error = await _collect_stream_error(provider, request)
assert "Connection reset" in error.message
def test_stream_malformed_tool_args_chunked(self):
"""Chunked tool args that never form valid JSON are flushed with {}."""
provider = _make_provider()
from free_claude_code.core.anthropic import AnthropicStreamLedger
sse = AnthropicStreamLedger("msg_test", "test-model")
tc1 = {
"index": 0,
"id": "call_malformed",
"function": {"name": "Task", "arguments": '{"broken":'},
}
tc2 = {
"index": 0,
"id": "call_malformed",
"function": {"name": None, "arguments": " never valid }"},
}
events1 = list(_make_tool_assembler(provider).process_tool_call(tc1, sse))
events2 = list(_make_tool_assembler(provider).process_tool_call(tc2, sse))
flushed = list(_make_tool_assembler(provider).flush_task_arg_buffers(sse))
event_text = "".join(events1 + events2 + flushed)
assert "tool_use" in event_text
assert "{}" in event_text
@pytest.mark.asyncio
async def test_openai_compat_stream_ends_with_contract_when_tool_name_never_arrives() -> (
None
):
"""Nameless / incomplete tool-call buffer must not break Anthropic stream contract."""
provider = _make_provider()
request = _make_request()
tc0 = SimpleNamespace(
index=0,
id="call_inc",
function=SimpleNamespace(name=None, arguments="{}"),
)
stream_mock = AsyncStreamMock([_make_chunk(tool_calls=[tc0])])
with (
patch.object(
provider._client.chat.completions,
"create",
new_callable=AsyncMock,
return_value=stream_mock,
),
):
error = await _collect_stream_error(provider, request)
assert "Provider stream ended without finish_reason." in error.message