Jack1808/Claude_Code
0
1"""Tests for streaming error handling in providers/nvidia_nim/client.py."""2 3import json4from unittest.mock import AsyncMock, MagicMock, patch5 6import httpx7import pytest8 9from config.nim import NimSettings10from providers.base import ProviderConfig11from providers.nvidia_nim import NvidiaNimProvider12 13 14class AsyncStreamMock:15 """Async iterable mock that yields chunks then optionally raises."""16 17 def __init__(self, chunks, error=None):18 self._chunks = chunks19 self._error = error20 21 def __aiter__(self):22 return self._aiter()23 24 async def _aiter(self):25 for chunk in self._chunks:26 yield chunk27 if self._error:28 raise self._error29 30 31def _make_provider():32 """Create a provider instance for testing."""33 config = ProviderConfig(34 api_key="test_key",35 base_url="https://test.api.nvidia.com/v1",36 rate_limit=10,37 rate_window=60,38 )39 return NvidiaNimProvider(config, nim_settings=NimSettings())40 41 42def _make_request(model="test-model", stream=True):43 """Create a mock request with all fields build_request_body needs."""44 req = MagicMock()45 req.model = model46 req.stream = stream47 req.messages = []48 req.system = None49 req.tools = None50 req.tool_choice = None51 req.metadata = None52 req.max_tokens = 409653 req.temperature = None54 req.top_p = None55 req.top_k = None56 req.stop_sequences = None57 req.extra_body = None58 req.thinking = None59 return req60 61 62def _make_chunk(63 content=None, finish_reason=None, tool_calls=None, reasoning_content=None64):65 """Create a mock streaming chunk."""66 delta = MagicMock()67 delta.content = content68 delta.tool_calls = tool_calls69 delta.reasoning_content = reasoning_content if reasoning_content else None70 71 choice = MagicMock()72 choice.delta = delta73 choice.finish_reason = finish_reason74 75 chunk = MagicMock()76 chunk.choices = [choice]77 chunk.usage = None78 return chunk79 80 81async def _collect_stream(provider, request):82 """Collect all SSE events from a stream."""83 return [e async for e in provider.stream_response(request)]84 85 86class TestStreamingExceptionHandling:87 """Tests for error paths during stream_response."""88 89 @pytest.mark.asyncio90 async def test_api_error_emits_sse_error_event(self):91 """When API raises during streaming, SSE error event is emitted."""92 provider = _make_provider()93 request = _make_request()94 95 mock_stream = AsyncMock()96 mock_stream.__aiter__ = MagicMock(side_effect=RuntimeError("API failed"))97 98 with (99 patch.object(100 provider._client.chat.completions,101 "create",102 new_callable=AsyncMock,103 side_effect=RuntimeError("API failed"),104 ),105 patch.object(106 provider._global_rate_limiter,107 "wait_if_blocked",108 new_callable=AsyncMock,109 return_value=False,110 ),111 ):112 events = await _collect_stream(provider, request)113 114 # Should have message_start, error text block, close blocks, message_delta, message_stop115 event_text = "".join(events)116 assert "message_start" in event_text117 assert "API failed" in event_text118 assert "message_stop" in event_text119 120 @pytest.mark.asyncio121 async def test_read_timeout_with_empty_message_emits_fallback(self):122 """ReadTimeout(TimeoutError()) should emit a visible, non-empty timeout message."""123 provider = _make_provider()124 request = _make_request()125 126 with (127 patch.object(128 provider._client.chat.completions,129 "create",130 new_callable=AsyncMock,131 side_effect=httpx.ReadTimeout(""),132 ),133 patch.object(134 provider._global_rate_limiter,135 "wait_if_blocked",136 new_callable=AsyncMock,137 return_value=False,138 ),139 ):140 events = [141 e142 async for e in provider.stream_response(143 request,144 request_id="req_timeout123",145 )146 ]147 148 event_text = "".join(events)149 assert "timed out after" in event_text150 assert "request_id=req_timeout123" in event_text151 assert "message_stop" in event_text152 153 @pytest.mark.asyncio154 async def test_error_after_partial_content(self):155 """Error after partial content: blocks closed, error emitted."""156 provider = _make_provider()157 request = _make_request()158 159 chunk1 = _make_chunk(content="Hello ")160 stream_mock = AsyncStreamMock([chunk1], error=RuntimeError("Connection lost"))161 162 with (163 patch.object(164 provider._client.chat.completions,165 "create",166 new_callable=AsyncMock,167 return_value=stream_mock,168 ),169 patch.object(170 provider._global_rate_limiter,171 "wait_if_blocked",172 new_callable=AsyncMock,173 return_value=False,174 ),175 ):176 events = await _collect_stream(provider, request)177 178 event_text = "".join(events)179 assert "Hello" in event_text180 assert "Connection lost" in event_text181 assert "message_stop" in event_text182 183 @pytest.mark.asyncio184 async def test_empty_response_gets_space(self):185 """Empty response with no text/tools gets a single space text block."""186 provider = _make_provider()187 request = _make_request()188 189 empty_chunk = _make_chunk(finish_reason="stop")190 stream_mock = AsyncStreamMock([empty_chunk])191 192 with (193 patch.object(194 provider._client.chat.completions,195 "create",196 new_callable=AsyncMock,197 return_value=stream_mock,198 ),199 patch.object(200 provider._global_rate_limiter,201 "wait_if_blocked",202 new_callable=AsyncMock,203 return_value=False,204 ),205 ):206 events = await _collect_stream(provider, request)207 208 event_text = "".join(events)209 assert '"text_delta"' in event_text210 assert "message_stop" in event_text211 212 @pytest.mark.asyncio213 async def test_stream_with_thinking_content(self):214 """Thinking content via think tags is emitted as thinking blocks."""215 provider = _make_provider()216 request = _make_request()217 218 chunk1 = _make_chunk(content="<think>reasoning</think>answer")219 chunk2 = _make_chunk(finish_reason="stop")220 stream_mock = AsyncStreamMock([chunk1, chunk2])221 222 with (223 patch.object(224 provider._client.chat.completions,225 "create",226 new_callable=AsyncMock,227 return_value=stream_mock,228 ),229 patch.object(230 provider._global_rate_limiter,231 "wait_if_blocked",232 new_callable=AsyncMock,233 return_value=False,234 ),235 ):236 events = await _collect_stream(provider, request)237 238 event_text = "".join(events)239 assert "thinking" in event_text240 assert "reasoning" in event_text241 assert "answer" in event_text242 243 @pytest.mark.asyncio244 async def test_stream_with_reasoning_content_field(self):245 """reasoning_content delta field is emitted as thinking block."""246 provider = _make_provider()247 request = _make_request()248 249 chunk1 = _make_chunk(reasoning_content="I think...")250 chunk2 = _make_chunk(content="The answer")251 chunk3 = _make_chunk(finish_reason="stop")252 stream_mock = AsyncStreamMock([chunk1, chunk2, chunk3])253 254 with (255 patch.object(256 provider._client.chat.completions,257 "create",258 new_callable=AsyncMock,259 return_value=stream_mock,260 ),261 patch.object(262 provider._global_rate_limiter,263 "wait_if_blocked",264 new_callable=AsyncMock,265 return_value=False,266 ),267 ):268 events = await _collect_stream(provider, request)269 270 event_text = "".join(events)271 assert "thinking_delta" in event_text272 assert "I think..." in event_text273 assert "The answer" in event_text274 275 @pytest.mark.asyncio276 async def test_stream_rate_limited_retries_via_execute_with_retry(self):277 """When rate limited, execute_with_retry handles retries transparently."""278 provider = _make_provider()279 request = _make_request()280 281 chunk1 = _make_chunk(content="Response")282 chunk2 = _make_chunk(finish_reason="stop")283 stream_mock = AsyncStreamMock([chunk1, chunk2])284 285 with patch.object(286 provider._client.chat.completions,287 "create",288 new_callable=AsyncMock,289 return_value=stream_mock,290 ):291 # Mock execute_with_retry to pass through to the actual function292 async def _passthrough(fn, *args, **kwargs):293 return await fn(*args, **kwargs)294 295 with patch.object(296 provider._global_rate_limiter,297 "execute_with_retry",298 new_callable=AsyncMock,299 side_effect=_passthrough,300 ):301 events = await _collect_stream(provider, request)302 303 event_text = "".join(events)304 assert "Response" in event_text305 306 307class TestProcessToolCall:308 """Tests for _process_tool_call method."""309 310 def test_tool_call_with_id(self):311 """Tool call with id starts a tool block."""312 provider = _make_provider()313 from providers.common import SSEBuilder314 315 sse = SSEBuilder("msg_test", "test-model")316 tc = {317 "index": 0,318 "id": "call_123",319 "function": {"name": "search", "arguments": '{"q": "test"}'},320 }321 events = list(provider._process_tool_call(tc, sse))322 event_text = "".join(events)323 assert "tool_use" in event_text324 assert "search" in event_text325 assert "call_123" in event_text326 327 def test_tool_call_without_id_generates_uuid(self):328 """Tool call without id generates a uuid-based id."""329 provider = _make_provider()330 from providers.common import SSEBuilder331 332 sse = SSEBuilder("msg_test", "test-model")333 tc = {334 "index": 0,335 "id": None,336 "function": {"name": "test", "arguments": "{}"},337 }338 events = list(provider._process_tool_call(tc, sse))339 event_text = "".join(events)340 assert "tool_" in event_text341 342 def test_task_tool_forces_background_false(self):343 """Task tool with run_in_background=true is forced to false."""344 provider = _make_provider()345 from providers.common import SSEBuilder346 347 sse = SSEBuilder("msg_test", "test-model")348 args = json.dumps({"run_in_background": True, "prompt": "test"})349 tc = {350 "index": 0,351 "id": "call_task",352 "function": {"name": "Task", "arguments": args},353 }354 events = list(provider._process_tool_call(tc, sse))355 event_text = "".join(events)356 # The intercepted args should have run_in_background=false357 assert "false" in event_text.lower()358 359 def test_task_tool_chunked_args_forces_background_false(self):360 """Chunked Task args are buffered until valid JSON, then forced to false."""361 provider = _make_provider()362 from providers.common import SSEBuilder363 364 sse = SSEBuilder("msg_test", "test-model")365 tc1 = {366 "index": 0,367 "id": "call_task_chunked",368 "function": {"name": "Task", "arguments": '{"run_in_background": true,'},369 }370 tc2 = {371 "index": 0,372 "id": "call_task_chunked",373 "function": {"name": None, "arguments": ' "prompt": "test"}'},374 }375 376 events1 = list(provider._process_tool_call(tc1, sse))377 assert len(events1) > 0378 assert "false" not in "".join(events1).lower()379 380 events2 = list(provider._process_tool_call(tc2, sse))381 event_text = "".join(events1 + events2)382 assert "false" in event_text.lower()383 384 def test_task_tool_invalid_json_logs_warning_on_flush(self, caplog):385 """Invalid JSON args for Task tool emits {} on flush and logs a warning."""386 provider = _make_provider()387 from providers.common import SSEBuilder388 389 sse = SSEBuilder("msg_test", "test-model")390 tc = {391 "index": 0,392 "id": "call_task2",393 "function": {"name": "Task", "arguments": "not json"},394 }395 events = list(provider._process_tool_call(tc, sse))396 assert len(events) > 0397 398 with caplog.at_level("WARNING"):399 flushed = list(provider._flush_task_arg_buffers(sse))400 assert len(flushed) > 0401 assert "{}" in "".join(flushed)402 assert any("Task args invalid JSON" in r.message for r in caplog.records)403 404 def test_negative_tool_index_fallback(self):405 """tc_index < 0 uses len(tool_indices) as fallback."""406 provider = _make_provider()407 from providers.common import SSEBuilder408 409 sse = SSEBuilder("msg_test", "test-model")410 tc = {411 "index": -1,412 "id": "call_neg",413 "function": {"name": "test", "arguments": "{}"},414 }415 events = list(provider._process_tool_call(tc, sse))416 # Should not crash, should still emit events417 assert len(events) > 0418 419 def test_tool_args_emitted_as_delta(self):420 """Arguments are emitted as input_json_delta events."""421 provider = _make_provider()422 from providers.common import SSEBuilder423 424 sse = SSEBuilder("msg_test", "test-model")425 tc = {426 "index": 0,427 "id": "call_args",428 "function": {"name": "grep", "arguments": '{"pattern": "test"}'},429 }430 events = list(provider._process_tool_call(tc, sse))431 event_text = "".join(events)432 assert "input_json_delta" in event_text433 434 435class TestStreamChunkEdgeCases:436 """Tests for edge cases in stream chunk handling."""437 438 @pytest.mark.asyncio439 async def test_stream_chunk_with_empty_choices_skipped(self):440 """Chunk with choices=[] is skipped without crashing."""441 provider = _make_provider()442 request = _make_request()443 444 empty_choices_chunk = MagicMock()445 empty_choices_chunk.choices = []446 empty_choices_chunk.usage = None447 448 finish_chunk = _make_chunk(finish_reason="stop")449 stream_mock = AsyncStreamMock([empty_choices_chunk, finish_chunk])450 451 with (452 patch.object(453 provider._client.chat.completions,454 "create",455 new_callable=AsyncMock,456 return_value=stream_mock,457 ),458 patch.object(459 provider._global_rate_limiter,460 "wait_if_blocked",461 new_callable=AsyncMock,462 return_value=False,463 ),464 ):465 events = await _collect_stream(provider, request)466 467 event_text = "".join(events)468 assert "message_start" in event_text469 assert "message_stop" in event_text470 471 @pytest.mark.asyncio472 async def test_stream_chunk_with_none_delta_handled(self):473 """Chunk with choice.delta=None is handled defensively."""474 provider = _make_provider()475 request = _make_request()476 477 none_delta_chunk = MagicMock()478 none_delta_chunk.usage = None479 choice = MagicMock()480 choice.delta = None481 choice.finish_reason = None482 none_delta_chunk.choices = [choice]483 484 finish_chunk = _make_chunk(finish_reason="stop")485 stream_mock = AsyncStreamMock([none_delta_chunk, finish_chunk])486 487 with (488 patch.object(489 provider._client.chat.completions,490 "create",491 new_callable=AsyncMock,492 return_value=stream_mock,493 ),494 patch.object(495 provider._global_rate_limiter,496 "wait_if_blocked",497 new_callable=AsyncMock,498 return_value=False,499 ),500 ):501 events = await _collect_stream(provider, request)502 503 event_text = "".join(events)504 assert "message_start" in event_text505 assert "message_stop" in event_text506 507 @pytest.mark.asyncio508 async def test_stream_generator_cleanup_on_exception(self):509 """When stream raises mid-iteration, message_stop still emitted."""510 provider = _make_provider()511 request = _make_request()512 513 chunk1 = _make_chunk(content="Partial")514 stream_mock = AsyncStreamMock(515 [chunk1], error=ConnectionResetError("Connection reset")516 )517 518 with (519 patch.object(520 provider._client.chat.completions,521 "create",522 new_callable=AsyncMock,523 return_value=stream_mock,524 ),525 patch.object(526 provider._global_rate_limiter,527 "wait_if_blocked",528 new_callable=AsyncMock,529 return_value=False,530 ),531 ):532 events = await _collect_stream(provider, request)533 534 event_text = "".join(events)535 assert "Partial" in event_text536 assert "Connection reset" in event_text537 assert "message_stop" in event_text538 539 def test_stream_malformed_tool_args_chunked(self):540 """Chunked tool args that never form valid JSON are flushed with {}."""541 provider = _make_provider()542 from providers.common import SSEBuilder543 544 sse = SSEBuilder("msg_test", "test-model")545 tc1 = {546 "index": 0,547 "id": "call_malformed",548 "function": {"name": "Task", "arguments": '{"broken":'},549 }550 tc2 = {551 "index": 0,552 "id": "call_malformed",553 "function": {"name": None, "arguments": " never valid }"},554 }555 556 events1 = list(provider._process_tool_call(tc1, sse))557 events2 = list(provider._process_tool_call(tc2, sse))558 flushed = list(provider._flush_task_arg_buffers(sse))559 560 event_text = "".join(events1 + events2 + flushed)561 assert "tool_use" in event_text562 assert "{}" in event_text563 