fix(xai): fold reasoning tokens before normalizing usage in streaming chunks

The non-streaming transform_response folds xAI's reasoning_tokens into
completion_tokens before calling _normalize_openai_compatible_usage_totals,
preserving the OpenAI invariant total = prompt + completion. The streaming
chunk_parser only ran the normalization, so when xAI streamed usage with
reasoning tokens (total = prompt + completion + reasoning), the normalize
check (total < prompt + completion) was a no-op and the invariant remained
violated.

Refactor _fold_reasoning_tokens_into_completion to also accept a raw usage
dict (in addition to ModelResponse / Usage) and call it from the streaming
chunk_parser before normalization, so streaming and non-streaming paths
report usage consistently for reasoning models.

Co-authored-by: Yassin Kortam <yassin@berri.ai>
This commit is contained in:
Cursor Agent 2026-05-20 22:46:29 +00:00
parent 459b83d731
commit 3d0fc09f8a
No known key found for this signature in database

View file

@ -230,18 +230,53 @@ class XAIChatConfig(OpenAIGPTConfig):
return response
@staticmethod
def _fold_reasoning_tokens_into_completion(model_response: ModelResponse) -> None:
def _fold_reasoning_tokens_into_completion(
target: Union[ModelResponse, Usage, Dict[str, Any], None],
) -> None:
"""Reconcile xAI Usage to the OpenAI invariant.
xAI accounts ``reasoning_tokens`` separately from
``completion_tokens`` while still summing them into ``total_tokens``.
OpenAI's contract (o1/o3) folds reasoning into ``completion_tokens``,
so fold here to keep ``total = prompt + completion``. Idempotent.
Accepts a ``ModelResponse`` (non-streaming), a ``Usage`` object, or a
raw usage ``dict`` (streaming chunk) so streaming and non-streaming
paths stay in sync.
"""
usage = getattr(model_response, "usage", None)
if target is None:
return
if isinstance(target, ModelResponse):
usage: Union[Usage, Dict[str, Any], None] = getattr(target, "usage", None)
else:
usage = target
if usage is None:
return
if isinstance(usage, dict):
details = usage.get("completion_tokens_details") or {}
if isinstance(details, dict):
reasoning_tokens = int(details.get("reasoning_tokens") or 0)
else:
reasoning_tokens = int(getattr(details, "reasoning_tokens", 0) or 0)
if reasoning_tokens <= 0:
return
prompt_tokens = int(usage.get("prompt_tokens") or 0)
completion_tokens = int(usage.get("completion_tokens") or 0)
total_tokens = int(usage.get("total_tokens") or 0)
if total_tokens == prompt_tokens + completion_tokens:
return
# Guard against double-counting if xAI changes accounting.
if total_tokens != prompt_tokens + completion_tokens + reasoning_tokens:
return
usage["completion_tokens"] = completion_tokens + reasoning_tokens
return
details = getattr(usage, "completion_tokens_details", None)
reasoning_tokens = (
int(getattr(details, "reasoning_tokens", 0) or 0) if details else 0
@ -326,6 +361,7 @@ class XAIChatCompletionStreamingHandler(OpenAIChatCompletionStreamingHandler):
chunk["choices"] = [{"index": 0, "delta": {}, "finish_reason": None}]
if "usage" in chunk and chunk["usage"] is not None:
XAIChatConfig._fold_reasoning_tokens_into_completion(chunk["usage"])
XAIChatConfig._normalize_openai_compatible_usage_totals(chunk["usage"])
return super().chunk_parser(chunk)