From 24a5e4d9680603d1fb1302753c9bcce3d9734a7a Mon Sep 17 00:00:00 2001 From: Shantanu Suryawanshi Date: Wed, 18 Feb 2026 21:31:28 -0500 Subject: [PATCH] Fixing sse stream --- providers/common/sse_builder.py | 9 +++++++-- providers/openai_compat.py | 1 - tests/providers/test_lmstudio.py | 3 --- tests/providers/test_streaming_errors.py | 8 ++------ 4 files changed, 9 insertions(+), 12 deletions(-) diff --git a/providers/common/sse_builder.py b/providers/common/sse_builder.py index d446303f..11500073 100644 --- a/providers/common/sse_builder.py +++ b/providers/common/sse_builder.py @@ -146,6 +146,7 @@ class SSEBuilder: # Message lifecycle events def message_start(self) -> str: """Generate message_start event.""" + usage = {"input_tokens": self.input_tokens, "output_tokens": 1} return self._format_event( "message_start", { @@ -158,8 +159,9 @@ class SSEBuilder: "model": self.model, "stop_reason": None, "stop_sequence": None, - "usage": {"input_tokens": self.input_tokens, "output_tokens": 1}, + "usage": usage, }, + "usage": usage, }, ) @@ -170,7 +172,10 @@ class SSEBuilder: { "type": "message_delta", "delta": {"stop_reason": stop_reason, "stop_sequence": None}, - "usage": {"output_tokens": output_tokens}, + "usage": { + "input_tokens": self.input_tokens, + "output_tokens": output_tokens, + }, }, ) diff --git a/providers/openai_compat.py b/providers/openai_compat.py index ffb99cb1..7740788f 100644 --- a/providers/openai_compat.py +++ b/providers/openai_compat.py @@ -323,4 +323,3 @@ class OpenAICompatibleProvider(BaseProvider): ) yield sse.message_delta(map_stop_reason(finish_reason), output_tokens) yield sse.message_stop() - yield sse.done() diff --git a/tests/providers/test_lmstudio.py b/tests/providers/test_lmstudio.py index 7299a0f6..1a49e490 100644 --- a/tests/providers/test_lmstudio.py +++ b/tests/providers/test_lmstudio.py @@ -290,7 +290,6 @@ class TestLMStudioStreamingExceptionHandling: assert "message_start" in event_text assert "API failed" in event_text assert "message_stop" in event_text - assert "[DONE]" in event_text @pytest.mark.asyncio async def test_error_after_partial_content(self, lmstudio_provider): @@ -361,7 +360,6 @@ class TestLMStudioStreamChunkEdgeCases: event_text = "".join(events) assert "message_start" in event_text assert "message_stop" in event_text - assert "[DONE]" in event_text @pytest.mark.asyncio async def test_stream_chunk_with_none_delta_handled(self, lmstudio_provider): @@ -387,7 +385,6 @@ class TestLMStudioStreamChunkEdgeCases: event_text = "".join(events) assert "message_start" in event_text assert "message_stop" in event_text - assert "[DONE]" in event_text # --- Native Tool Calls --- diff --git a/tests/providers/test_streaming_errors.py b/tests/providers/test_streaming_errors.py index 9909426a..3e842520 100644 --- a/tests/providers/test_streaming_errors.py +++ b/tests/providers/test_streaming_errors.py @@ -110,12 +110,11 @@ class TestStreamingExceptionHandling: ): events = await _collect_stream(provider, request) - # Should have message_start, error text block, close blocks, message_delta, message_stop, done + # Should have message_start, error text block, close blocks, message_delta, message_stop event_text = "".join(events) assert "message_start" in event_text assert "API failed" in event_text assert "message_stop" in event_text - assert "[DONE]" in event_text @pytest.mark.asyncio async def test_error_after_partial_content(self): @@ -434,7 +433,6 @@ class TestStreamChunkEdgeCases: event_text = "".join(events) assert "message_start" in event_text assert "message_stop" in event_text - assert "[DONE]" in event_text @pytest.mark.asyncio async def test_stream_chunk_with_none_delta_handled(self): @@ -471,11 +469,10 @@ class TestStreamChunkEdgeCases: event_text = "".join(events) assert "message_start" in event_text assert "message_stop" in event_text - assert "[DONE]" in event_text @pytest.mark.asyncio async def test_stream_generator_cleanup_on_exception(self): - """When stream raises mid-iteration, message_stop and [DONE] still emitted.""" + """When stream raises mid-iteration, message_stop still emitted.""" provider = _make_provider() request = _make_request() @@ -504,7 +501,6 @@ class TestStreamChunkEdgeCases: assert "Partial" in event_text assert "Connection reset" in event_text assert "message_stop" in event_text - assert "[DONE]" in event_text def test_stream_malformed_tool_args_chunked(self): """Chunked tool args that never form valid JSON are flushed with {}."""