Fix usage in xai

This commit is contained in:
Sameer Kankute
2026-02-19 18:48:30 +05:30
parent e00c181f0c
commit fc7fa73b6f
2 changed files with 98 additions and 3 deletions
+46 -3
View File
@@ -1,4 +1,4 @@
from typing import List, Optional, Tuple
from typing import Any, AsyncIterator, Iterator, List, Optional, Tuple, Union
import httpx
@@ -11,9 +11,18 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
)
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import Choices, ModelResponse, Usage, PromptTokensDetailsWrapper
from litellm.types.utils import (
Choices,
ModelResponse,
ModelResponseStream,
PromptTokensDetailsWrapper,
Usage,
)
from ...openai.chat.gpt_transformation import OpenAIGPTConfig
from ...openai.chat.gpt_transformation import (
OpenAIChatCompletionStreamingHandler,
OpenAIGPTConfig,
)
class XAIChatConfig(OpenAIGPTConfig):
@@ -119,6 +128,18 @@ class XAIChatConfig(OpenAIGPTConfig):
optional_params[param] = value
return optional_params
def get_model_response_iterator(
self,
streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse],
sync_stream: bool,
json_mode: Optional[bool] = False,
) -> Any:
return XAIChatCompletionStreamingHandler(
streaming_response=streaming_response,
sync_stream=sync_stream,
json_mode=json_mode,
)
def transform_request(
self,
model: str,
@@ -225,3 +246,25 @@ class XAIChatConfig(OpenAIGPTConfig):
usage.prompt_tokens_details.web_search_requests = int(num_sources_used)
setattr(usage, "num_sources_used", int(num_sources_used))
verbose_logger.debug(f"X.AI web search sources used: {num_sources_used}")
class XAIChatCompletionStreamingHandler(OpenAIChatCompletionStreamingHandler):
def chunk_parser(self, chunk: dict) -> ModelResponseStream:
"""
Handle xAI-specific streaming behavior.
xAI Grok sends a final chunk with empty choices array but with usage data
when stream_options={"include_usage": True} is set.
Example from xAI API:
{"id":"...","object":"chat.completion.chunk","created":...,"model":"grok-4-1-fast-non-reasoning",
"choices":[],"usage":{"prompt_tokens":171,"completion_tokens":2,"total_tokens":173,...}}
"""
# Handle chunks with empty choices but with usage data
choices = chunk.get("choices", [])
if len(choices) == 0 and "usage" in chunk:
# xAI sends usage in a chunk with empty choices array
# Add a dummy choice with empty delta to ensure proper processing
chunk["choices"] = [{"index": 0, "delta": {}, "finish_reason": None}]
return super().chunk_parser(chunk)
+52
View File
@@ -201,3 +201,55 @@ class TestXAIChat(BaseLLMChatTest):
)
assert response is not None
def test_xai_streaming_with_include_usage():
"""
Test that xAI streaming correctly handles usage in the last chunk
when stream_options={"include_usage": True} is set.
xAI sends usage in a chunk with empty choices array, which should be
handled by XAIChatCompletionStreamingHandler.
"""
try:
response = completion(
model="xai/grok-4-1-fast-non-reasoning",
messages=[
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Say hello in one word"}
],
stream=True,
stream_options={"include_usage": True},
max_tokens=10,
)
chunks = []
usage_chunk = None
for chunk in response:
chunks.append(chunk)
if hasattr(chunk, "usage") and chunk.usage is not None:
usage_chunk = chunk
# Verify we got chunks
assert len(chunks) > 0, "Should receive streaming chunks"
# Verify usage was included in one of the chunks
assert usage_chunk is not None, "Should receive usage in streaming chunks"
# Verify usage has expected fields
assert hasattr(usage_chunk.usage, "prompt_tokens"), "Usage should have prompt_tokens"
assert hasattr(usage_chunk.usage, "completion_tokens"), "Usage should have completion_tokens"
assert hasattr(usage_chunk.usage, "total_tokens"), "Usage should have total_tokens"
# Verify usage values are positive
assert usage_chunk.usage.prompt_tokens > 0, "prompt_tokens should be positive"
assert usage_chunk.usage.completion_tokens > 0, "completion_tokens should be positive"
assert usage_chunk.usage.total_tokens > 0, "total_tokens should be positive"
print(f"✓ Successfully received usage in streaming chunk: {usage_chunk.usage}")
except Exception as e:
if "API key" in str(e) or "authentication" in str(e).lower():
pytest.skip(f"Skipping test due to API key issue: {str(e)}")
raise