diff --git a/pycodeloop/providers/generic.py b/pycodeloop/providers/generic.py index f2854ec..754e5cd 100644 --- a/pycodeloop/providers/generic.py +++ b/pycodeloop/providers/generic.py @@ -383,7 +383,8 @@ def _stream( body = {**body, "stream": True} text = "" pending: dict[int, dict] = {} - stop_reason = "stop" + stop_reason: str | None = None + saw_terminal_marker = False usage = Usage() with self._open(body, config) as response: @@ -393,8 +394,14 @@ def _stream( continue payload = line[len("data: ") :] if payload == "[DONE]": + saw_terminal_marker = True + break + try: + chunk = json.loads(payload) + except json.JSONDecodeError: + if stop_reason is None: + stop_reason = "malformed_stream" break - chunk = json.loads(payload) if chunk.get("usage"): usage = Usage( @@ -451,6 +458,10 @@ def _stream( if choice.get("finish_reason"): stop_reason = choice["finish_reason"] + saw_terminal_marker = True + + if stop_reason is None: + stop_reason = "stop" if saw_terminal_marker else "connection_lost" tool_calls = [ ToolCall( diff --git a/tests/providers/test_generic.py b/tests/providers/test_generic.py index 2c71abc..903846b 100644 --- a/tests/providers/test_generic.py +++ b/tests/providers/test_generic.py @@ -267,6 +267,77 @@ def test_streaming_response_also_captures_extra_tool_call_fields(self): {"extra_content": {"google": {"thought_signature": "xyz789"}}}, ) + def test_streaming_keeps_already_shown_text_on_malformed_chunk(self): + path = self._write_config( + {"url": "http://fake/v1/chat/completions", "model": "my-model"} + ) + provider = GenericProvider.from_json(path) + + chunks = [{"choices": [{"delta": {"content": "hello there"}}]}] + sse_body = ( + "".join(f"data: {json.dumps(c)}\n" for c in chunks) + + "data: {not valid json\n" + + "data: [DONE]\n" + ).encode() + + deltas = [] + with mock.patch( + "pycodeloop.providers.generic.urllib.request.urlopen", + return_value=_FakeResponse(sse_body), + ): + result = provider.complete("sys", [], [], on_delta=deltas.append) + + self.assertEqual(result.text, "hello there") + self.assertEqual(result.stop_reason, "malformed_stream") + self.assertEqual("".join(deltas), "hello there") + + def test_streaming_prefers_finish_reason_over_trailing_malformed_chunk( + self, + ): + path = self._write_config( + {"url": "http://fake/v1/chat/completions", "model": "my-model"} + ) + provider = GenericProvider.from_json(path) + + chunks = [ + { + "choices": [ + {"delta": {"content": "done"}, "finish_reason": "stop"} + ] + } + ] + sse_body = ( + "".join(f"data: {json.dumps(c)}\n" for c in chunks) + + "data: {junk after finish\n" + ).encode() + + with mock.patch( + "pycodeloop.providers.generic.urllib.request.urlopen", + return_value=_FakeResponse(sse_body), + ): + result = provider.complete("sys", [], [], on_delta=lambda _: None) + + self.assertEqual(result.text, "done") + self.assertEqual(result.stop_reason, "stop") + + def test_streaming_flags_a_connection_dropped_mid_response(self): + path = self._write_config( + {"url": "http://fake/v1/chat/completions", "model": "my-model"} + ) + provider = GenericProvider.from_json(path) + + chunks = [{"choices": [{"delta": {"content": "cut off mid"}}]}] + sse_body = "".join(f"data: {json.dumps(c)}\n" for c in chunks).encode() + + with mock.patch( + "pycodeloop.providers.generic.urllib.request.urlopen", + return_value=_FakeResponse(sse_body), + ): + result = provider.complete("sys", [], [], on_delta=lambda _: None) + + self.assertEqual(result.text, "cut off mid") + self.assertEqual(result.stop_reason, "connection_lost") + def test_streaming_cuts_a_looping_response_short(self): path = self._write_config( {"url": "http://fake/v1/chat/completions", "model": "my-model"}