Team Ai
Apppublic

Jack1808/Claude_Code

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
test_streaming_errors.py563 linesDownload Raw Back to providers
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