Merge pull request #43 from suryawanshishantanu6/feature/fix-input-token

This commit is contained in:
Ali Khokhar
2026-02-18 18:43:14 -08:00
committed by GitHub
4 changed files with 9 additions and 12 deletions
+7 -2
View File
@@ -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,
},
},
)
-1
View File
@@ -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()
-3
View File
@@ -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 ---
+2 -6
View File
@@ -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 {}."""