Merge pull request #30817 from geraint0923/litellm_fix_xai_web_search_cost_billing
Some checks are pending
CI Coverage / assert-ci-coverage (push) Waiting to run
CodSpeed Benchmarks / benchmarks (push) Waiting to run
Publish basedpyright base counts / publish (push) Waiting to run
Code Quality Checks / code-quality (push) Waiting to run
UI Unit Tests / ui-unit-tests (push) Waiting to run
Unit Tests: Core Utilities / core-utils (push) Waiting to run
Unit Tests: Documentation Validation / documentation (push) Waiting to run
Unit Tests: Enterprise, Google GenAI & Routing / enterprise-routing (push) Waiting to run
Unit Tests: Integrations (Callbacks & Logging) / integrations (push) Waiting to run
Unit Tests: LLM Provider Transformations / Vertex AI (push) Waiting to run
Unit Tests: LLM Provider Transformations / All Other Providers (push) Waiting to run
Unit Tests: MCP, Secrets, Containers & Misc / misc (push) Waiting to run
Unit Tests: Proxy Auth & Key Management / proxy-auth (push) Waiting to run
Unit Tests: Proxy DB Operations / assert-shard-coverage (push) Waiting to run
Unit Tests: Proxy DB Operations / auth-checks (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / budgets (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / custom-logging (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / db-and-spend (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / endpoints-and-responses (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / guardrails-hooks (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-utils (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / jwt-and-keys (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / key-generation (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / logging-misc (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-runtime (push) Blocked by required conditions
Unit Tests: Proxy Infrastructure / proxy-infra (push) Waiting to run
Unit Tests: Proxy Legacy Tests / auth-and-jwt (push) Waiting to run
Unit Tests: Proxy Legacy Tests / key-generation (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-config (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-response-and-misc (push) Waiting to run
Unit Tests: Proxy DB Operations / proxy-server-core (push) Blocked by required conditions
Unit Tests: Proxy API Endpoints / proxy-endpoints (push) Waiting to run
Unit Tests: Proxy API Endpoints / proxy-server (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-server (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-server-extras (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-token-counter (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-user-auth-and-spend (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-utils (push) Waiting to run
Unit Tests: Responses, Caching & Types / responses-caching-types (push) Waiting to run
GitHub Actions Security Analysis / zizmor (push) Waiting to run

fix(xai): bill web_search from server_side_tool_usage_details
This commit is contained in:
Mateo Wang 2026-08-11 19:40:27 -07:00 • committed by GitHub
commit b4f5e46a44
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 664 additions and 242 deletions

View file

@ -2,6 +2,7 @@
Helper utilities for tracking the cost of built-in tools. Helper utilities for tracking the cost of built-in tools.
""" """
from collections.abc import Mapping
from typing import Any, Final, Literal from typing import Any, Final, Literal
import litellm import litellm
@ -23,6 +24,14 @@ from litellm.types.utils import (
) )
def _usage_reports_server_side_web_search_calls(usage: Usage) -> bool:
details: Final = getattr(usage, "server_side_tool_usage_details", None)
if not isinstance(details, Mapping):
return False
calls: Final = details.get("web_search_calls")
return isinstance(calls, int) and calls > 0
class StandardBuiltInToolCostTracking: class StandardBuiltInToolCostTracking:
""" """
Helper class for tracking the cost of built-in tools Helper class for tracking the cost of built-in tools
@ -351,6 +360,10 @@ class StandardBuiltInToolCostTracking:
# and _handle_web_search_cost() is never called. # and _handle_web_search_cost() is never called.
if hasattr(usage, "server_tool_use") and _get_web_search_requests(usage.server_tool_use) is not None: if hasattr(usage, "server_tool_use") and _get_web_search_requests(usage.server_tool_use) is not None:
return True return True
# xAI reports usage.server_side_tool_usage_details.web_search_calls; a searched
# answer with no url_citation annotations has no other chat-path signal
if _usage_reports_server_side_web_search_calls(usage):
return True
return False return False
elif isinstance(response_object, ResponsesAPIResponse): elif isinstance(response_object, ResponsesAPIResponse):
# response api explicitly includes web_search_call in the output # response api explicitly includes web_search_call in the output
@ -370,6 +383,8 @@ class StandardBuiltInToolCostTracking:
) )
): ):
return True return True
if _usage_reports_server_side_web_search_calls(usage):
return True
return False return False
@ -432,7 +447,9 @@ class StandardBuiltInToolCostTracking:
""" """
output: Final = response_object.output output: Final = response_object.output
for output_item in output: for output_item in output:
_output_type: str | None = getattr(output_item, "type", None) _output_type: str | None = (
output_item.get("type") if isinstance(output_item, dict) else getattr(output_item, "type", None)
)
if _output_type == output_type: if _output_type == output_type:
return True return True
return False return False

View file

@ -1,4 +1,4 @@
from collections.abc import AsyncIterator, Iterator from collections.abc import AsyncIterator, Iterator, Mapping
from typing import Any, Final from typing import Any, Final
import httpx import httpx
@ -12,13 +12,15 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
strip_name_from_messages, strip_name_from_messages,
) )
from litellm.llms.xai.common_utils import XAIModelInfo from litellm.llms.xai.common_utils import XAIModelInfo
from litellm.llms.xai.cost_calculator import (
apply_server_side_tool_usage_details_to_usage,
)
from litellm.secret_managers.main import get_secret_str from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import ( from litellm.types.utils import (
Choices, Choices,
ModelResponse, ModelResponse,
ModelResponseStream, ModelResponseStream,
PromptTokensDetailsWrapper,
Usage, Usage,
) )
@ -248,7 +250,7 @@ class XAIChatConfig(OpenAIGPTConfig):
XAI API returns empty string for finish_reason when using tools, XAI API returns empty string for finish_reason when using tools,
so we need to fix this after the standard OpenAI transformation. so we need to fix this after the standard OpenAI transformation.
Also handles X.AI web search usage tracking by extracting num_sources_used. Also handles X.AI web search usage tracking.
""" """
# First, let the parent class handle the standard transformation # First, let the parent class handle the standard transformation
@ -351,25 +353,20 @@ class XAIChatConfig(OpenAIGPTConfig):
def _enhance_usage_with_xai_web_search_fields(self, model_response: ModelResponse, raw_response_json: dict) -> None: def _enhance_usage_with_xai_web_search_fields(self, model_response: ModelResponse, raw_response_json: dict) -> None:
""" """
Extract num_sources_used from X.AI response and map it to web_search_requests. Copy usage.server_side_tool_usage_details from the provider usage block
onto model_response.usage for tool cost calculation.
""" """
if not hasattr(model_response, "usage") or model_response.usage is None: if not hasattr(model_response, "usage") or model_response.usage is None:
return return
usage: Final[Usage] = model_response.usage usage: Final[Usage] = model_response.usage
num_sources_used = None response_usage: Final = raw_response_json.get("usage")
response_usage: Final = raw_response_json.get("usage", {}) if not isinstance(response_usage, dict):
if isinstance(response_usage, dict) and "num_sources_used" in response_usage: return
num_sources_used = response_usage.get("num_sources_used") details: Final = response_usage.get("server_side_tool_usage_details")
if isinstance(details, Mapping):
# Map num_sources_used to web_search_requests for cost detection apply_server_side_tool_usage_details_to_usage(usage, details)
if num_sources_used is not None and num_sources_used > 0: verbose_logger.debug("X.AI server_side_tool_usage_details: %s", details)
if usage.prompt_tokens_details is None:
usage.prompt_tokens_details = PromptTokensDetailsWrapper()
usage.prompt_tokens_details.web_search_requests = int(num_sources_used)
setattr(usage, "num_sources_used", int(num_sources_used))
verbose_logger.debug("X.AI web search sources used: %s", num_sources_used)
@staticmethod @staticmethod
def _normalize_openai_compatible_usage_totals( def _normalize_openai_compatible_usage_totals(

View file

@ -4,14 +4,37 @@ Helper util for handling XAI-specific cost calculation
- Handles XAI-specific reasoning token billing (billed as part of completion tokens) - Handles XAI-specific reasoning token billing (billed as part of completion tokens)
""" """
from collections.abc import Mapping
from typing import TYPE_CHECKING, Final from typing import TYPE_CHECKING, Final
from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token
from litellm.types.utils import Usage from litellm.types.utils import PromptTokensDetailsWrapper, Usage
if TYPE_CHECKING: if TYPE_CHECKING:
from litellm.types.utils import ModelInfo from litellm.types.utils import ModelInfo
# https://docs.x.ai/developers/pricing#tools-pricing — default when unset in model map
_DEFAULT_WEB_SEARCH_COST_PER_CALL: Final = 5.0 / 1000.0
def apply_server_side_tool_usage_details_to_usage(usage: Usage, details: Mapping[str, object] | None) -> None:
"""
Attach server_side_tool_usage_details and mirror web_search_calls onto
prompt_tokens_details.web_search_requests for built-in tool cost gating.
"""
if details is None:
return
usage.server_side_tool_usage_details = details # pyright: ignore[reportAttributeAccessIssue] # extra # rebind-ok: extras
try:
web_search_calls: Final = int(details.get("web_search_calls") or 0)
except (TypeError, ValueError):
return
if web_search_calls <= 0:
return
prompt_tokens_details: Final = usage.prompt_tokens_details or PromptTokensDetailsWrapper()
prompt_tokens_details.web_search_requests = web_search_calls
usage.prompt_tokens_details = prompt_tokens_details # rebind-ok: write details onto caller usage
def cost_per_token(model: str, usage: Usage) -> tuple[float, float]: def cost_per_token(model: str, usage: Usage) -> tuple[float, float]:
""" """
@ -32,9 +55,11 @@ def cost_per_token(model: str, usage: Usage) -> tuple[float, float]:
prompt_tokens: Final = int(getattr(usage, "prompt_tokens", 0) or 0) prompt_tokens: Final = int(getattr(usage, "prompt_tokens", 0) or 0)
completion_tokens: Final = int(getattr(usage, "completion_tokens", 0) or 0) completion_tokens: Final = int(getattr(usage, "completion_tokens", 0) or 0)
total_tokens: Final = int(getattr(usage, "total_tokens", 0) or 0) total_tokens: Final = int(getattr(usage, "total_tokens", 0) or 0)
reasoning_tokens = 0 reasoning_tokens: Final = (
if hasattr(usage, "completion_tokens_details") and usage.completion_tokens_details: int(getattr(usage.completion_tokens_details, "reasoning_tokens", 0) or 0)
reasoning_tokens = int(getattr(usage.completion_tokens_details, "reasoning_tokens", 0) or 0) if hasattr(usage, "completion_tokens_details") and usage.completion_tokens_details
else 0
)
already_normalised: Final = total_tokens == prompt_tokens + completion_tokens already_normalised: Final = total_tokens == prompt_tokens + completion_tokens
total_completion_tokens: Final = completion_tokens if already_normalised else completion_tokens + reasoning_tokens total_completion_tokens: Final = completion_tokens if already_normalised else completion_tokens + reasoning_tokens
@ -52,33 +77,48 @@ def cost_per_token(model: str, usage: Usage) -> tuple[float, float]:
return prompt_cost, completion_cost return prompt_cost, completion_cost
def _web_search_cost_per_call_from_model_info(model_info: "ModelInfo") -> float:
"""
Per-invocation web_search price from model_info when configured.
Prefer ``search_context_cost_per_query`` (same shape as Gemini/Anthropic web
search pricing in the model cost map). Fall back to current xAI list pricing.
"""
search_costs: Final = model_info.get("search_context_cost_per_query")
if not isinstance(search_costs, Mapping):
return _DEFAULT_WEB_SEARCH_COST_PER_CALL
for key in (
"search_context_size_medium",
"search_context_size_low",
"search_context_size_high",
):
value = search_costs.get(key)
if value is None:
continue
try:
cost = float(value)
except (TypeError, ValueError):
continue
if cost > 0:
return cost
return _DEFAULT_WEB_SEARCH_COST_PER_CALL
def cost_per_web_search_request(usage: "Usage", model_info: "ModelInfo") -> float: def cost_per_web_search_request(usage: "Usage", model_info: "ModelInfo") -> float:
""" """
Calculate the cost of web search requests for X.AI models. Calculate the cost of web search requests for X.AI models.
X.AI Live Search costs $25 per 1,000 sources used. Counts invocations from usage.server_side_tool_usage_details.web_search_calls.
Each source costs $0.025. Per-call rate comes from model_info.search_context_cost_per_query when set,
otherwise the default xAI tools rate ($5 / 1k calls).
The number of sources is stored in prompt_tokens_details.web_search_requests
by the transformation layer to be compatible with the existing detection system.
""" """
# Cost per source used: $25 per 1,000 sources = $0.025 per source details: Final = getattr(usage, "server_side_tool_usage_details", None)
cost_per_source: Final = 25.0 / 1000.0 # $0.025 if not isinstance(details, Mapping):
return 0.0
num_sources_used = 0 try:
web_search_calls: Final = int(details.get("web_search_calls") or 0)
if ( except (TypeError, ValueError):
hasattr(usage, "prompt_tokens_details") return 0.0
and usage.prompt_tokens_details is not None if web_search_calls <= 0:
and hasattr(usage.prompt_tokens_details, "web_search_requests") return 0.0
and usage.prompt_tokens_details.web_search_requests is not None return _web_search_cost_per_call_from_model_info(model_info) * web_search_calls
):
num_sources_used = int(usage.prompt_tokens_details.web_search_requests)
# Fallback: try to get from num_sources_used if set directly
elif hasattr(usage, "num_sources_used") and usage.num_sources_used is not None:
num_sources_used = int(usage.num_sources_used)
total_cost: Final = cost_per_source * num_sources_used
return total_cost

View file

@ -1,4 +1,4 @@
from typing import TYPE_CHECKING, Any, Final from typing import Any, Final
import litellm import litellm
from litellm._logging import verbose_logger from litellm._logging import verbose_logger
@ -12,13 +12,6 @@ from litellm.types.llms.xai import XAIWebSearchTool, XAIXSearchTool
from litellm.types.router import GenericLiteLLMParams from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import LlmProviders from litellm.types.utils import LlmProviders
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig): class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
""" """

View file

@ -1033,13 +1033,18 @@ class ResponseAPILoggingUtils:
@staticmethod @staticmethod
def _transform_response_api_usage_to_chat_usage( def _transform_response_api_usage_to_chat_usage(
usage_input: dict | ResponseAPIUsage | None, usage_input: Mapping[str, object] | ResponseAPIUsage | Usage | None,
) -> Usage: ) -> Usage:
""" """
Transforms ResponseAPIUsage or ImageUsage to a Usage object. Transforms ResponseAPIUsage or ImageUsage to a Usage object.
Both have the same spec with input_tokens, output_tokens, and Both have the same spec with input_tokens, output_tokens, and
input_tokens_details (text_tokens, image_tokens). input_tokens_details (text_tokens, image_tokens).
Usage inputs are returned as-is so re-running this helper never drops
fields. Non-standard provider fields (e.g. xAI's
server_side_tool_usage_details) are carried onto the returned Usage so
provider cost calculators can read them after normalization.
""" """
if usage_input is None: if usage_input is None:
return Usage( return Usage(
@ -1047,6 +1052,10 @@ class ResponseAPILoggingUtils:
completion_tokens=0, completion_tokens=0,
total_tokens=0, total_tokens=0,
) )
if isinstance(usage_input, Usage):
return usage_input
if isinstance(usage_input, dict) and not ResponseAPILoggingUtils._is_response_api_usage(usage_input):
return Usage(**usage_input)
response_api_usage: ResponseAPIUsage response_api_usage: ResponseAPIUsage
if isinstance(usage_input, dict): if isinstance(usage_input, dict):
usage_input = dict(usage_input) # shallow copy; avoid mutating caller usage_input = dict(usage_input) # shallow copy; avoid mutating caller
@ -1055,13 +1064,11 @@ class ResponseAPILoggingUtils:
usage_input["input_tokens_details"] = usage_input["input_token_details"] usage_input["input_tokens_details"] = usage_input["input_token_details"]
if usage_input.get("output_tokens_details") is None and "output_token_details" in usage_input: if usage_input.get("output_tokens_details") is None and "output_token_details" in usage_input:
usage_input["output_tokens_details"] = usage_input["output_token_details"] usage_input["output_tokens_details"] = usage_input["output_token_details"]
total_tokens = usage_input.get("total_tokens") if usage_input.get("total_tokens") is None:
if total_tokens is None:
input_tokens: Final = usage_input.get("input_tokens") input_tokens: Final = usage_input.get("input_tokens")
output_tokens: Final = usage_input.get("output_tokens") output_tokens: Final = usage_input.get("output_tokens")
if input_tokens is not None and output_tokens is not None: if isinstance(input_tokens, int) and isinstance(output_tokens, int):
total_tokens = input_tokens + output_tokens usage_input["total_tokens"] = input_tokens + output_tokens
usage_input["total_tokens"] = total_tokens
response_api_usage = ResponseAPIUsage(**usage_input) response_api_usage = ResponseAPIUsage(**usage_input)
else: else:
response_api_usage = usage_input response_api_usage = usage_input
@ -1089,12 +1096,27 @@ class ResponseAPILoggingUtils:
audio_tokens=getattr(output_tokens_details, "audio_tokens", None), audio_tokens=getattr(output_tokens_details, "audio_tokens", None),
) )
extra_usage_fields: Final = {
key: value
for key, value in (response_api_usage.model_extra or {}).items()
if key
not in (
"input_token_details",
"output_token_details",
"prompt_tokens",
"completion_tokens",
"total_tokens",
"prompt_tokens_details",
"completion_tokens_details",
)
}
chat_usage: Final = Usage( chat_usage: Final = Usage(
prompt_tokens=prompt_tokens, prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens, completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens, total_tokens=prompt_tokens + completion_tokens,
prompt_tokens_details=prompt_tokens_details, prompt_tokens_details=prompt_tokens_details,
completion_tokens_details=completion_tokens_details, completion_tokens_details=completion_tokens_details,
**extra_usage_fields,
) )
# Preserve cost attribute if it exists on ResponseAPIUsage # Preserve cost attribute if it exists on ResponseAPIUsage

View file

@ -339,7 +339,7 @@ def test_get_cost_for_vertex_ai_gemini_web_search(model, custom_llm_provider):
for url_citation annotations, not usage.prompt_tokens_details.web_search_requests. for url_citation annotations, not usage.prompt_tokens_details.web_search_requests.
This causes Vertex AI grounding costs to not be tracked. This causes Vertex AI grounding costs to not be tracked.
""" """
from litellm.types.utils import PromptTokensDetailsWrapper, Usage, Choices, Message from litellm.types.utils import Choices, Message, PromptTokensDetailsWrapper, Usage
# Create a realistic ModelResponse like what Vertex AI returns # Create a realistic ModelResponse like what Vertex AI returns
response = ModelResponse( response = ModelResponse(
@ -604,3 +604,66 @@ def test_web_search_provider_prefix_fallback_does_not_misprice_non_gemini_model(
# Note: File search integration test removed due to complex annotation detection logic # Note: File search integration test removed due to complex annotation detection logic
# The unit tests in test_azure_assistant_cost_tracking.py provide comprehensive coverage # The unit tests in test_azure_assistant_cost_tracking.py provide comprehensive coverage
def test_response_includes_output_type_reads_dict_output_items():
"""
Regression: output items that fail OpenAI SDK validation (e.g. xAI web_search_call
items without an "action" field) stay plain dicts in the output union. The gate must
read their "type" key instead of returning False and skipping the web search fee.
"""
from litellm.types.llms.openai import ResponsesAPIResponse
response = ResponsesAPIResponse.model_validate(
{
"id": "resp_1",
"created_at": 1754900000,
"model": "grok-4",
"object": "response",
"status": "completed",
"output": [{"type": "web_search_call", "id": "ws_1", "status": "completed"}],
}
)
assert isinstance(response.output[0], dict)
assert StandardBuiltInToolCostTracking.response_includes_output_type(
response_object=response, output_type="web_search_call"
)
assert not StandardBuiltInToolCostTracking.response_includes_output_type(
response_object=response, output_type="file_search_call"
)
def test_web_search_gate_reads_server_side_tool_usage_details_without_citations():
"""
Regression: xAI chat responses bridged from the Responses API only carry
usage.server_side_tool_usage_details; a searched answer with no url_citation
annotations must still be billed for its web search calls.
"""
from litellm.llms.xai.cost_calculator import _DEFAULT_WEB_SEARCH_COST_PER_CALL
from litellm.types.utils import Usage
usage = Usage(
prompt_tokens=10,
completion_tokens=20,
total_tokens=30,
server_side_tool_usage_details={"web_search_calls": 3},
)
response = ModelResponse(model="xai/grok-4.5")
assert StandardBuiltInToolCostTracking.response_object_includes_web_search_call(
response_object=response, usage=usage
)
assert not StandardBuiltInToolCostTracking.response_object_includes_web_search_call(
response_object=response,
usage=Usage(prompt_tokens=10, completion_tokens=20, total_tokens=30),
)
cost = StandardBuiltInToolCostTracking.get_cost_for_built_in_tools(
model="xai/grok-4.5",
response_object=response,
usage=usage,
custom_llm_provider="xai",
standard_built_in_tools_params=None,
)
assert cost == 3 * _DEFAULT_WEB_SEARCH_COST_PER_CALL

View file

@ -9,14 +9,22 @@ Source: litellm/llms/xai/responses/transformation.py
import os import os
import sys import sys
from unittest.mock import MagicMock
sys.path.insert(0, os.path.abspath("../../../../..")) sys.path.insert(0, os.path.abspath("../../../../.."))
import pytest import pytest
import litellm
from litellm.llms.xai.responses.transformation import XAIResponsesAPIConfig from litellm.llms.xai.responses.transformation import XAIResponsesAPIConfig
from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams from litellm.responses.utils import ResponseAPILoggingUtils
from litellm.types.utils import LlmProviders from litellm.types.llms.openai import (
ResponseAPIUsage,
ResponseCompletedEvent,
ResponsesAPIOptionalRequestParams,
ResponsesAPIResponse,
)
from litellm.types.utils import LlmProviders, Usage
from litellm.utils import ProviderConfigManager from litellm.utils import ProviderConfigManager
@ -31,43 +39,29 @@ class TestXAIResponsesAPITransformation:
) )
assert config is not None, "Config should not be None for XAI provider" assert config is not None, "Config should not be None for XAI provider"
assert isinstance( assert isinstance(config, XAIResponsesAPIConfig), f"Expected XAIResponsesAPIConfig, got {type(config)}"
config, XAIResponsesAPIConfig assert config.custom_llm_provider == LlmProviders.XAI, "custom_llm_provider should be XAI"
), f"Expected XAIResponsesAPIConfig, got {type(config)}"
assert (
config.custom_llm_provider == LlmProviders.XAI
), "custom_llm_provider should be XAI"
def test_code_interpreter_container_field_removed(self): def test_code_interpreter_container_field_removed(self):
"""Test that container field is removed from code_interpreter tools""" """Test that container field is removed from code_interpreter tools"""
config = XAIResponsesAPIConfig() config = XAIResponsesAPIConfig()
params = ResponsesAPIOptionalRequestParams( params = ResponsesAPIOptionalRequestParams(tools=[{"type": "code_interpreter", "container": {"type": "auto"}}])
tools=[{"type": "code_interpreter", "container": {"type": "auto"}}]
)
result = config.map_openai_params( result = config.map_openai_params(response_api_optional_params=params, model="grok-4-fast", drop_params=False)
response_api_optional_params=params, model="grok-4-fast", drop_params=False
)
assert "tools" in result assert "tools" in result
assert len(result["tools"]) == 1 assert len(result["tools"]) == 1
assert result["tools"][0]["type"] == "code_interpreter" assert result["tools"][0]["type"] == "code_interpreter"
assert ( assert "container" not in result["tools"][0], "Container field should be removed"
"container" not in result["tools"][0]
), "Container field should be removed"
def test_instructions_parameter_dropped(self): def test_instructions_parameter_dropped(self):
"""Test that instructions parameter is dropped for XAI""" """Test that instructions parameter is dropped for XAI"""
config = XAIResponsesAPIConfig() config = XAIResponsesAPIConfig()
params = ResponsesAPIOptionalRequestParams( params = ResponsesAPIOptionalRequestParams(instructions="You are a helpful assistant.", temperature=0.7)
instructions="You are a helpful assistant.", temperature=0.7
)
result = config.map_openai_params( result = config.map_openai_params(response_api_optional_params=params, model="grok-4-fast", drop_params=False)
response_api_optional_params=params, model="grok-4-fast", drop_params=False
)
assert "instructions" not in result, "Instructions should be dropped" assert "instructions" not in result, "Instructions should be dropped"
assert result.get("temperature") == 0.7, "Other params should be preserved" assert result.get("temperature") == 0.7, "Other params should be preserved"
@ -88,25 +82,15 @@ class TestXAIResponsesAPITransformation:
# Test with default XAI API base # Test with default XAI API base
url = config.get_complete_url(api_base=None, litellm_params={}) url = config.get_complete_url(api_base=None, litellm_params={})
assert ( assert url == "https://api.x.ai/v1/responses", f"Expected XAI responses endpoint, got {url}"
url == "https://api.x.ai/v1/responses"
), f"Expected XAI responses endpoint, got {url}"
# Test with custom api_base # Test with custom api_base
custom_url = config.get_complete_url( custom_url = config.get_complete_url(api_base="https://custom.x.ai/v1", litellm_params={})
api_base="https://custom.x.ai/v1", litellm_params={} assert custom_url == "https://custom.x.ai/v1/responses", f"Expected custom endpoint, got {custom_url}"
)
assert (
custom_url == "https://custom.x.ai/v1/responses"
), f"Expected custom endpoint, got {custom_url}"
# Test with trailing slash # Test with trailing slash
url_with_slash = config.get_complete_url( url_with_slash = config.get_complete_url(api_base="https://api.x.ai/v1/", litellm_params={})
api_base="https://api.x.ai/v1/", litellm_params={} assert url_with_slash == "https://api.x.ai/v1/responses", "Should handle trailing slash"
)
assert (
url_with_slash == "https://api.x.ai/v1/responses"
), "Should handle trailing slash"
def test_web_search_tool_transformation(self): def test_web_search_tool_transformation(self):
"""Test that web_search tools are transformed to XAI format""" """Test that web_search tools are transformed to XAI format"""
@ -167,9 +151,7 @@ class TestXAIResponsesAPITransformation:
config = XAIResponsesAPIConfig() config = XAIResponsesAPIConfig()
params = ResponsesAPIOptionalRequestParams( params = ResponsesAPIOptionalRequestParams(
tools=[ tools=[{"type": "web_search", "excluded_domains": ["example.com", "test.com"]}]
{"type": "web_search", "excluded_domains": ["example.com", "test.com"]}
]
) )
result = config.map_openai_params( result = config.map_openai_params(
@ -309,3 +291,115 @@ class TestXAIResponsesAPITransformation:
# Verify function tool is unchanged # Verify function tool is unchanged
assert result["tools"][3]["type"] == "function" assert result["tools"][3]["type"] == "function"
assert result["tools"][3]["name"] == "get_weather" assert result["tools"][3]["name"] == "get_weather"
class TestXAIResponsesWebSearchBilling:
"""Web search billing must not change the client-visible Responses usage schema."""
_TOOL_DETAILS = {
"web_search_calls": 2,
"x_search_calls": 0,
"code_interpreter_calls": 0,
"file_search_calls": 0,
"mcp_calls": 0,
"document_search_calls": 0,
}
def _raw_response_json(self, include_web_search: bool) -> dict:
web_search_output = (
[{
"type": "web_search_call",
"id": "ws_1",
"status": "completed",
"action": {"type": "search", "query": "grok"},
}] if include_web_search else []
)
tool_usage = {"server_side_tool_usage_details": self._TOOL_DETAILS} if include_web_search else {}
return {
"id": "resp_1",
"object": "response",
"created_at": 1754900000,
"model": "grok-4",
"status": "completed",
"parallel_tool_calls": True,
"tool_choice": "auto",
"tools": [],
"top_p": 1.0,
"output": web_search_output
+ [
{
"type": "message",
"id": "msg_1",
"role": "assistant",
"status": "completed",
"content": [{"type": "output_text", "text": "grok says hi", "annotations": []}],
}
],
"usage": {
"input_tokens": 100,
"output_tokens": 20,
"total_tokens": 120,
"input_tokens_details": {"cached_tokens": 0},
"output_tokens_details": {"reasoning_tokens": 0},
**tool_usage,
},
}
def _transform(self, include_web_search: bool) -> ResponsesAPIResponse:
raw_response = MagicMock()
raw_response.json.return_value = self._raw_response_json(include_web_search)
raw_response.text = "raw"
raw_response.headers = {}
return XAIResponsesAPIConfig().transform_response_api_response(
model="grok-4", raw_response=raw_response, logging_obj=MagicMock()
)
def test_response_usage_keeps_responses_api_schema(self):
response = self._transform(include_web_search=True)
assert isinstance(response.usage, ResponseAPIUsage)
assert response.usage.input_tokens == 100
assert response.usage.output_tokens == 20
assert response.usage.model_extra["server_side_tool_usage_details"] == self._TOOL_DETAILS
def test_bridged_usage_keeps_tool_details_for_billing(self):
response = self._transform(include_web_search=True)
bridged = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(response.usage)
assert isinstance(bridged, Usage)
assert bridged.prompt_tokens == 100
assert bridged.completion_tokens == 20
assert getattr(bridged, "server_side_tool_usage_details") == self._TOOL_DETAILS
def test_completion_cost_bills_web_search_calls(self):
with_search = litellm.completion_cost(
completion_response=self._transform(include_web_search=True),
model="xai/grok-4",
custom_llm_provider="xai",
)
without_search = litellm.completion_cost(
completion_response=self._transform(include_web_search=False),
model="xai/grok-4",
custom_llm_provider="xai",
)
assert with_search - without_search == pytest.approx(2 * 5.0 / 1000.0)
def test_streaming_terminal_event_keeps_schema_and_details(self):
parsed_chunk = {
"type": "response.completed",
"sequence_number": 7,
"response": self._raw_response_json(include_web_search=True),
}
event = XAIResponsesAPIConfig().transform_streaming_response(
model="grok-4", parsed_chunk=parsed_chunk, logging_obj=MagicMock()
)
assert isinstance(event, ResponseCompletedEvent)
assert isinstance(event.response.usage, ResponseAPIUsage)
assert event.response.usage.input_tokens == 100
bridged = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(event.response.usage)
assert getattr(bridged, "server_side_tool_usage_details") == self._TOOL_DETAILS

View file

@ -5,6 +5,9 @@ sys.path.insert(
0, os.path.abspath("../../../..") 0, os.path.abspath("../../../..")
) # Adds the parent directory to the system path ) # Adds the parent directory to the system path
import pytest
import litellm
from litellm.llms.xai.chat.transformation import XAIChatConfig from litellm.llms.xai.chat.transformation import XAIChatConfig
from litellm.types.utils import ( from litellm.types.utils import (
CompletionTokensDetailsWrapper, CompletionTokensDetailsWrapper,
@ -135,3 +138,65 @@ class TestXAIUsageNormalization:
XAIChatConfig._normalize_openai_compatible_usage_totals(usage) XAIChatConfig._normalize_openai_compatible_usage_totals(usage)
assert usage["total_tokens"] == 200 assert usage["total_tokens"] == 200
class TestXAIChatWebSearchBilling:
_TOOL_DETAILS = {
"web_search_calls": 3,
"x_search_calls": 0,
"code_interpreter_calls": 0,
"file_search_calls": 0,
"mcp_calls": 0,
"document_search_calls": 0,
}
@staticmethod
def _response_with_usage() -> ModelResponse:
response = ModelResponse(model="grok-4")
setattr(
response,
"usage",
Usage(prompt_tokens=100, completion_tokens=20, total_tokens=120),
)
return response
def test_enhance_copies_details_and_mirrors_web_search_requests(self):
response = self._response_with_usage()
XAIChatConfig()._enhance_usage_with_xai_web_search_fields(
response,
{"usage": {"server_side_tool_usage_details": self._TOOL_DETAILS}},
)
usage = response.usage
assert getattr(usage, "server_side_tool_usage_details") == self._TOOL_DETAILS
assert usage.prompt_tokens_details is not None
assert usage.prompt_tokens_details.web_search_requests == 3
def test_enhance_noop_without_details(self):
response = self._response_with_usage()
XAIChatConfig()._enhance_usage_with_xai_web_search_fields(
response, {"usage": {"prompt_tokens": 100}}
)
assert response.usage.prompt_tokens_details is None
assert getattr(response.usage, "server_side_tool_usage_details", None) is None
def test_completion_cost_bills_chat_web_search_calls(self):
billed = self._response_with_usage()
XAIChatConfig()._enhance_usage_with_xai_web_search_fields(
billed,
{"usage": {"server_side_tool_usage_details": self._TOOL_DETAILS}},
)
with_search = litellm.completion_cost(
completion_response=billed, model="xai/grok-4", custom_llm_provider="xai"
)
without_search = litellm.completion_cost(
completion_response=self._response_with_usage(),
model="xai/grok-4",
custom_llm_provider="xai",
)
assert with_search - without_search == pytest.approx(3 * 5.0 / 1000.0)

View file

@ -17,7 +17,16 @@ sys.path.insert(
0, os.path.abspath("../../..") 0, os.path.abspath("../../..")
) # Adds the parent directory to the system path ) # Adds the parent directory to the system path
from litellm.llms.xai.cost_calculator import cost_per_token, cost_per_web_search_request from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import (
StandardBuiltInToolCostTracking,
)
from litellm.llms.xai.cost_calculator import (
_DEFAULT_WEB_SEARCH_COST_PER_CALL,
_web_search_cost_per_call_from_model_info,
apply_server_side_tool_usage_details_to_usage,
cost_per_token,
cost_per_web_search_request,
)
class TestXAICostCalculator: class TestXAICostCalculator:
@ -354,76 +363,53 @@ class TestXAICostCalculator:
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10) assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10) assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
def test_web_search_cost_calculation(self): def test_web_search_cost_via_server_side_tool_usage_details(self):
"""Test web search cost calculation for X.AI models.""" """usage.server_side_tool_usage_details.web_search_calls at default $5/1k."""
# Test with web_search_requests in prompt_tokens_details (primary path) usage = Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150)
usage = Usage( setattr(
prompt_tokens=100, usage,
completion_tokens=50, "server_side_tool_usage_details",
total_tokens=150, {
prompt_tokens_details=PromptTokensDetailsWrapper( "web_search_calls": 3,
text_tokens=100, "x_search_calls": 0,
web_search_requests=3, # 3 sources used "code_interpreter_calls": 0,
), "file_search_calls": 0,
"mcp_calls": 0,
"document_search_calls": 0,
},
) )
web_search_cost = cost_per_web_search_request(usage=usage, model_info={}) web_search_cost = cost_per_web_search_request(usage=usage, model_info={})
assert math.isclose(web_search_cost, 3 * (5.0 / 1000.0), rel_tol=1e-10)
# Expected cost: 3 sources * $0.025 per source = $0.075 def test_web_search_cost_uses_model_info_search_context_pricing(self):
expected_cost = 3 * (25.0 / 1000.0) # 3 * $0.025 usage = Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15)
setattr(usage, "server_side_tool_usage_details", {"web_search_calls": 2})
assert math.isclose(web_search_cost, expected_cost, rel_tol=1e-10) model_info = {
assert math.isclose(web_search_cost, 0.075, rel_tol=1e-10) "search_context_cost_per_query": {
"search_context_size_medium": 0.01,
def test_web_search_cost_fallback_calculation(self): }
"""Test web search cost calculation using fallback num_sources_used.""" }
# Test fallback: num_sources_used on usage object web_search_cost = cost_per_web_search_request(
usage = Usage( usage=usage, model_info=model_info
prompt_tokens=100,
completion_tokens=50,
total_tokens=150,
) )
# Manually set num_sources_used (as done by transformation layer) assert math.isclose(web_search_cost, 0.02, rel_tol=1e-10)
setattr(usage, "num_sources_used", 5)
web_search_cost = cost_per_web_search_request(usage=usage, model_info={}) def test_web_search_cost_zero_without_details(self):
usage = Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150)
assert cost_per_web_search_request(usage=usage, model_info={}) == 0.0
# Expected cost: 5 sources * $0.025 per source = $0.125 def test_apply_details_sets_web_search_requests_for_cost_gate(self):
expected_cost = 5 * (25.0 / 1000.0) # 5 * $0.025 usage = Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15)
apply_server_side_tool_usage_details_to_usage(
assert math.isclose(web_search_cost, expected_cost, rel_tol=1e-10) usage, {"web_search_calls": 2, "x_search_calls": 0}
assert math.isclose(web_search_cost, 0.125, rel_tol=1e-10)
def test_web_search_no_sources_used(self):
"""Test web search cost calculation when no sources are used."""
usage = Usage(
prompt_tokens=100,
completion_tokens=50,
total_tokens=150,
prompt_tokens_details=PromptTokensDetailsWrapper(
text_tokens=100,
web_search_requests=0, # No web search
),
) )
assert usage.prompt_tokens_details is not None
web_search_cost = cost_per_web_search_request(usage=usage, model_info={}) assert usage.prompt_tokens_details.web_search_requests == 2
assert StandardBuiltInToolCostTracking.response_object_includes_web_search_call(
# Expected cost: 0 sources * $0.025 per source = $0.0 response_object=object(), usage=usage
assert web_search_cost == 0.0
def test_web_search_cost_without_prompt_tokens_details(self):
"""Test web search cost calculation when prompt_tokens_details is None."""
usage = Usage(
prompt_tokens=100,
completion_tokens=50,
total_tokens=150,
) )
web_search_cost = cost_per_web_search_request(usage=usage, model_info={})
# Expected cost: No web search data = $0.0
assert web_search_cost == 0.0
def test_grok_4_20_beta_reasoning_cost_calculation(self): def test_grok_4_20_beta_reasoning_cost_calculation(self):
"""Test cost calculation for grok-4.20-beta-0309-reasoning model.""" """Test cost calculation for grok-4.20-beta-0309-reasoning model."""
usage = Usage(prompt_tokens=100, completion_tokens=200, total_tokens=300) usage = Usage(prompt_tokens=100, completion_tokens=200, total_tokens=300)
@ -499,3 +485,112 @@ class TestXAICostCalculator:
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10) assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10) assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
class TestXAIWebSearchCostHelpers:
"""Focused coverage for web_search / tool-usage helpers in cost_calculator.py."""
def test_apply_details_noop_when_details_none(self):
usage = Usage(prompt_tokens=1, completion_tokens=1, total_tokens=2)
apply_server_side_tool_usage_details_to_usage(usage, None)
assert getattr(usage, "server_side_tool_usage_details", None) is None
def test_apply_details_sets_attr_but_skips_mirror_when_web_search_zero(self):
usage = Usage(prompt_tokens=1, completion_tokens=1, total_tokens=2)
details = {"web_search_calls": 0, "x_search_calls": 3}
apply_server_side_tool_usage_details_to_usage(usage, details)
assert getattr(usage, "server_side_tool_usage_details") == details
assert (
usage.prompt_tokens_details is None
or usage.prompt_tokens_details.web_search_requests is None
)
def test_apply_details_skips_mirror_when_web_search_calls_invalid(self):
usage = Usage(prompt_tokens=1, completion_tokens=1, total_tokens=2)
details = {"web_search_calls": "not-a-number"}
apply_server_side_tool_usage_details_to_usage(usage, details)
assert getattr(usage, "server_side_tool_usage_details") == details
assert usage.prompt_tokens_details is None
def test_apply_details_updates_existing_prompt_tokens_details(self):
usage = Usage(
prompt_tokens=1,
completion_tokens=1,
total_tokens=2,
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=7),
)
apply_server_side_tool_usage_details_to_usage(usage, {"web_search_calls": 4})
assert usage.prompt_tokens_details is not None
assert usage.prompt_tokens_details.cached_tokens == 7
assert usage.prompt_tokens_details.web_search_requests == 4
def test_web_search_cost_per_call_default_when_model_info_empty(self):
assert (
_web_search_cost_per_call_from_model_info({})
== _DEFAULT_WEB_SEARCH_COST_PER_CALL
)
def test_web_search_cost_per_call_prefers_medium_over_low(self):
model_info = {
"search_context_cost_per_query": {
"search_context_size_low": 0.001,
"search_context_size_medium": 0.009,
}
}
assert _web_search_cost_per_call_from_model_info(model_info) == 0.009
def test_web_search_cost_per_call_falls_back_to_low_then_high(self):
assert (
_web_search_cost_per_call_from_model_info(
{"search_context_cost_per_query": {"search_context_size_low": 0.003}}
)
== 0.003
)
assert (
_web_search_cost_per_call_from_model_info(
{"search_context_cost_per_query": {"search_context_size_high": 0.007}}
)
== 0.007
)
def test_web_search_cost_per_call_ignores_zero_and_invalid_values(self):
assert (
_web_search_cost_per_call_from_model_info(
{
"search_context_cost_per_query": {
"search_context_size_medium": 0,
"search_context_size_low": "bad",
}
}
)
== _DEFAULT_WEB_SEARCH_COST_PER_CALL
)
def test_cost_per_web_search_request_zero_when_details_not_mapping(self):
usage = Usage(prompt_tokens=1, completion_tokens=1, total_tokens=2)
setattr(usage, "server_side_tool_usage_details", "invalid")
assert cost_per_web_search_request(usage=usage, model_info={}) == 0.0
def test_cost_per_web_search_request_zero_when_web_search_calls_invalid(self):
usage = Usage(prompt_tokens=1, completion_tokens=1, total_tokens=2)
setattr(
usage,
"server_side_tool_usage_details",
{"web_search_calls": object()},
)
assert cost_per_web_search_request(usage=usage, model_info={}) == 0.0
def test_cost_per_web_search_request_zero_when_web_search_calls_zero(self):
usage = Usage(prompt_tokens=1, completion_tokens=1, total_tokens=2)
setattr(
usage,
"server_side_tool_usage_details",
{"web_search_calls": 0, "x_search_calls": 5},
)
assert cost_per_web_search_request(usage=usage, model_info={}) == 0.0
def test_cost_per_web_search_request_uses_default_rate_without_model_pricing(self):
usage = Usage(prompt_tokens=1, completion_tokens=1, total_tokens=2)
setattr(usage, "server_side_tool_usage_details", {"web_search_calls": 4})
cost = cost_per_web_search_request(usage=usage, model_info={})
assert math.isclose(cost, 4 * _DEFAULT_WEB_SEARCH_COST_PER_CALL, rel_tol=1e-10)

View file

@ -1,21 +1,16 @@
import base64 import base64
import json
import os import os
import sys import sys
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
import pytest import pytest
from fastapi.testclient import TestClient
sys.path.insert( sys.path.insert(0, os.path.abspath("../../..")) # Adds the parent directory to the system path
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
import litellm import litellm
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
from litellm.responses.utils import ResponseAPILoggingUtils, ResponsesAPIRequestUtils from litellm.responses.utils import ResponseAPILoggingUtils, ResponsesAPIRequestUtils
from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIOptionalRequestParams
from litellm.types.utils import Usage from litellm.types.utils import Usage
@ -54,9 +49,7 @@ class TestResponsesAPIRequestUtils:
# Setup # Setup
model = "gpt-4o" model = "gpt-4o"
config = OpenAIResponsesAPIConfig() config = OpenAIResponsesAPIConfig()
optional_params = ResponsesAPIOptionalRequestParams( optional_params = ResponsesAPIOptionalRequestParams({"temperature": 0.7, "unsupported_param": "value"})
{"temperature": 0.7, "unsupported_param": "value"}
)
# Execute and Assert # Execute and Assert
with pytest.raises(litellm.UnsupportedParamsError) as excinfo: with pytest.raises(litellm.UnsupportedParamsError) as excinfo:
@ -90,9 +83,7 @@ class TestResponsesAPIRequestUtils:
assert result == {"temperature": 0.7} assert result == {"temperature": 0.7}
@pytest.mark.parametrize("request_drop_params", [None, False]) @pytest.mark.parametrize("request_drop_params", [None, False])
def test_get_optional_params_responses_api_still_raises_without_drop( def test_get_optional_params_responses_api_still_raises_without_drop(self, monkeypatch, request_drop_params):
self, monkeypatch, request_drop_params
):
"""Absent or False request-level drop_params must not suppress the unsupported-param error""" """Absent or False request-level drop_params must not suppress the unsupported-param error"""
monkeypatch.setattr(litellm, "drop_params", False) monkeypatch.setattr(litellm, "drop_params", False)
config = OpenAIResponsesAPIConfig() config = OpenAIResponsesAPIConfig()
@ -119,9 +110,7 @@ class TestResponsesAPIRequestUtils:
} }
# Execute # Execute
result = ResponsesAPIRequestUtils.get_requested_response_api_optional_param( result = ResponsesAPIRequestUtils.get_requested_response_api_optional_param(params)
params
)
# Assert # Assert
assert "temperature" in result assert "temperature" in result
@ -147,40 +136,31 @@ class TestResponsesAPIRequestUtils:
) )
# Execute # Execute
result = ResponsesAPIRequestUtils.decode_previous_response_id_to_original_previous_response_id( result = ResponsesAPIRequestUtils.decode_previous_response_id_to_original_previous_response_id(encoded_id)
encoded_id
)
# Assert # Assert
assert result == original_response_id assert result == original_response_id
# Test with a non-encoded ID # Test with a non-encoded ID
plain_id = "resp_xyz789" plain_id = "resp_xyz789"
result_plain = ResponsesAPIRequestUtils.decode_previous_response_id_to_original_previous_response_id( result_plain = ResponsesAPIRequestUtils.decode_previous_response_id_to_original_previous_response_id(plain_id)
plain_id
)
assert result_plain == plain_id assert result_plain == plain_id
def test_update_responses_api_response_id_with_model_id_handles_dict(self): def test_update_responses_api_response_id_with_model_id_handles_dict(self):
"""Ensure _update_responses_api_response_id_with_model_id works with dict input""" """Ensure _update_responses_api_response_id_with_model_id works with dict input"""
responses_api_response = {"id": "resp_abc123"} responses_api_response = {"id": "resp_abc123"}
litellm_metadata = {"model_info": {"id": "gpt-4o"}} litellm_metadata = {"model_info": {"id": "gpt-4o"}}
updated = ( updated = ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id(
ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id(
responses_api_response=responses_api_response, responses_api_response=responses_api_response,
custom_llm_provider="openai", custom_llm_provider="openai",
litellm_metadata=litellm_metadata, litellm_metadata=litellm_metadata,
) )
)
assert updated["id"] != "resp_abc123" assert updated["id"] != "resp_abc123"
decoded = ResponsesAPIRequestUtils._decode_responses_api_response_id( decoded = ResponsesAPIRequestUtils._decode_responses_api_response_id(updated["id"])
updated["id"]
)
assert decoded.get("response_id") == "resp_abc123" assert decoded.get("response_id") == "resp_abc123"
assert decoded.get("model_id") == "gpt-4o" assert decoded.get("model_id") == "gpt-4o"
assert decoded.get("custom_llm_provider") == "openai" assert decoded.get("custom_llm_provider") == "openai"
def test_update_responses_api_response_id_with_model_id_is_idempotent_for_litellm_ids(self): def test_update_responses_api_response_id_with_model_id_is_idempotent_for_litellm_ids(self):
raw = "resp_" + "a" * 48 raw = "resp_" + "a" * 48
litellm_metadata = {"model_info": {"id": "model-123"}} litellm_metadata = {"model_info": {"id": "model-123"}}
@ -207,9 +187,7 @@ class TestResponsesAPIRequestUtils:
model_id=None, model_id=None,
container_id="cntr_upstream_abc", container_id="cntr_upstream_abc",
) )
assert "None" not in base64.b64decode( assert "None" not in base64.b64decode(encoded.replace("cntr_", "").encode("utf-8")).decode("utf-8")
encoded.replace("cntr_", "").encode("utf-8")
).decode("utf-8")
decoded = ResponsesAPIRequestUtils._decode_container_id(encoded) decoded = ResponsesAPIRequestUtils._decode_container_id(encoded)
assert decoded.get("custom_llm_provider") == "azure" assert decoded.get("custom_llm_provider") == "azure"
assert decoded.get("model_id") is None assert decoded.get("model_id") is None
@ -217,12 +195,8 @@ class TestResponsesAPIRequestUtils:
def test_decode_container_id_legacy_literal_none_model_id(self): def test_decode_container_id_legacy_literal_none_model_id(self):
"""IDs encoded before the None fix should decode without a bogus model_id.""" """IDs encoded before the None fix should decode without a bogus model_id."""
legacy_inner = ( legacy_inner = "litellm:custom_llm_provider:azure;model_id:None;container_id:cntr_x"
"litellm:custom_llm_provider:azure;model_id:None;container_id:cntr_x" legacy_id = "cntr_" + base64.b64encode(legacy_inner.encode("utf-8")).decode("utf-8")
)
legacy_id = "cntr_" + base64.b64encode(legacy_inner.encode("utf-8")).decode(
"utf-8"
)
decoded = ResponsesAPIRequestUtils._decode_container_id(legacy_id) decoded = ResponsesAPIRequestUtils._decode_container_id(legacy_id)
assert decoded.get("model_id") is None assert decoded.get("model_id") is None
assert decoded.get("custom_llm_provider") == "azure" assert decoded.get("custom_llm_provider") == "azure"
@ -264,19 +238,14 @@ class TestResponseAPILoggingUtils:
} }
# Execute # Execute
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
usage
)
# Assert # Assert
assert isinstance(result, Usage) assert isinstance(result, Usage)
assert result.prompt_tokens == 10 assert result.prompt_tokens == 10
assert result.completion_tokens == 20 assert result.completion_tokens == 20
assert result.total_tokens == 30 assert result.total_tokens == 30
assert ( assert result.prompt_tokens_details and result.prompt_tokens_details.cached_tokens == 2
result.prompt_tokens_details
and result.prompt_tokens_details.cached_tokens == 2
)
def test_transform_response_api_usage_with_none_values(self): def test_transform_response_api_usage_with_none_values(self):
"""Test transformation handles None values properly""" """Test transformation handles None values properly"""
@ -289,9 +258,7 @@ class TestResponseAPILoggingUtils:
} }
# Execute # Execute
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
usage
)
# Assert # Assert
assert result.prompt_tokens == 0 assert result.prompt_tokens == 0
@ -310,9 +277,7 @@ class TestResponseAPILoggingUtils:
} }
# Execute # Execute
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
usage
)
# Assert # Assert
assert result.prompt_tokens == 15 assert result.prompt_tokens == 15
@ -349,9 +314,7 @@ class TestResponseAPILoggingUtils:
} }
# Execute # Execute
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
usage
)
# Assert - verify basic token counts # Assert - verify basic token counts
assert isinstance(result, Usage) assert isinstance(result, Usage)
@ -386,9 +349,7 @@ class TestResponseAPILoggingUtils:
}, },
} }
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
usage
)
assert result.prompt_tokens_details is not None assert result.prompt_tokens_details is not None
assert result.prompt_tokens_details.cache_write_tokens == 10059 assert result.prompt_tokens_details.cache_write_tokens == 10059
@ -417,9 +378,7 @@ class TestResponseAPILoggingUtils:
} }
# Execute # Execute
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
usage
)
# Assert - all token detail types should be preserved # Assert - all token detail types should be preserved
assert result.prompt_tokens_details is not None assert result.prompt_tokens_details is not None
@ -451,9 +410,7 @@ class TestResponseAPILoggingUtils:
}, },
} }
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
usage
)
assert result.prompt_tokens_details is not None assert result.prompt_tokens_details is not None
assert result.prompt_tokens_details.text_tokens == 8 assert result.prompt_tokens_details.text_tokens == 8
@ -475,9 +432,7 @@ class TestResponseAPILoggingUtils:
"output_token_details": {"text_tokens": 2, "audio_tokens": 98}, "output_token_details": {"text_tokens": 2, "audio_tokens": 98},
} }
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
usage
)
assert result.prompt_tokens_details is not None assert result.prompt_tokens_details is not None
assert result.prompt_tokens_details.text_tokens == 10 assert result.prompt_tokens_details.text_tokens == 10
@ -487,6 +442,93 @@ class TestResponseAPILoggingUtils:
assert result.completion_tokens_details.text_tokens == 20 assert result.completion_tokens_details.text_tokens == 20
assert result.completion_tokens_details.audio_tokens is None assert result.completion_tokens_details.audio_tokens is None
def test_transform_response_api_usage_carries_extra_provider_fields(self):
"""Non-standard usage fields (e.g. xAI tool details) must survive chat normalization."""
details = {"web_search_calls": 2, "x_search_calls": 0}
usage = ResponseAPIUsage(
input_tokens=100,
output_tokens=20,
total_tokens=120,
server_side_tool_usage_details=details,
)
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
assert isinstance(result, Usage)
assert result.prompt_tokens == 100
assert result.completion_tokens == 20
assert getattr(result, "server_side_tool_usage_details") == details
def test_transform_response_api_usage_ignores_chat_shaped_extras(self):
"""Gemini image usage carries chat-shaped keys as extras; they must not collide with explicit kwargs."""
usage = ResponseAPIUsage(
input_tokens=35,
output_tokens=1716,
total_tokens=1751,
prompt_tokens=35,
prompt_tokens_details={"image_tokens": 5, "text_tokens": 30},
completion_tokens=1716,
completion_tokens_details={"image_tokens": 1120, "text_tokens": 596},
server_side_tool_usage_details={"web_search_calls": 1},
)
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
assert result.prompt_tokens == 35
assert result.completion_tokens == 1716
assert getattr(result, "server_side_tool_usage_details") == {"web_search_calls": 1}
def test_transform_already_chat_usage_passthrough_keeps_tool_details(self):
"""Re-running the bridge on an already-converted chat Usage must not drop fields."""
details = {"web_search_calls": 2, "x_search_calls": 0}
usage = Usage(
prompt_tokens=100,
completion_tokens=20,
total_tokens=120,
prompt_tokens_details={"web_search_requests": 2},
)
setattr(usage, "server_side_tool_usage_details", details)
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
assert result is usage
assert getattr(result, "server_side_tool_usage_details") == details
assert result.prompt_tokens_details is not None
assert result.prompt_tokens_details.web_search_requests == 2
def test_transform_chat_shaped_usage_dict_keeps_tool_details(self):
"""Streaming chat bridge dumps already-converted Usage as a prompt_tokens dict."""
details = {
"web_search_calls": 3,
"x_search_calls": 0,
"code_interpreter_calls": 0,
"file_search_calls": 0,
"mcp_calls": 0,
"document_search_calls": 0,
"image_generation_calls": 0,
}
usage = {
"prompt_tokens": 50,
"completion_tokens": 10,
"total_tokens": 60,
"prompt_tokens_details": {"web_search_requests": 3, "cached_tokens": 8},
"completion_tokens_details": {"reasoning_tokens": 4},
"server_side_tool_usage_details": details,
}
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
assert isinstance(result, Usage)
assert result.prompt_tokens == 50
assert result.completion_tokens == 10
assert result.total_tokens == 60
assert getattr(result, "server_side_tool_usage_details") == details
assert result.prompt_tokens_details is not None
assert result.prompt_tokens_details.web_search_requests == 3
assert result.prompt_tokens_details.cached_tokens == 8
assert result.completion_tokens_details is not None
assert result.completion_tokens_details.reasoning_tokens == 4
class TestResponsesAPIProviderSpecificParams: class TestResponsesAPIProviderSpecificParams:
""" """
@ -503,9 +545,7 @@ class TestResponsesAPIProviderSpecificParams:
} }
# Should not raise any exception # Should not raise any exception
result = ResponsesAPIRequestUtils.get_requested_response_api_optional_param( result = ResponsesAPIRequestUtils.get_requested_response_api_optional_param(params)
params
)
assert "temperature" in result assert "temperature" in result
def test_provider_specific_params_no_crash_with_openai(self): def test_provider_specific_params_no_crash_with_openai(self):
@ -517,9 +557,7 @@ class TestResponsesAPIProviderSpecificParams:
} }
# Should not raise any exception # Should not raise any exception
result = ResponsesAPIRequestUtils.get_requested_response_api_optional_param( result = ResponsesAPIRequestUtils.get_requested_response_api_optional_param(params)
params
)
assert "temperature" in result assert "temperature" in result
def test_provider_specific_params_no_crash_with_vertex_ai(self): def test_provider_specific_params_no_crash_with_vertex_ai(self):
@ -531,9 +569,7 @@ class TestResponsesAPIProviderSpecificParams:
} }
# Should not raise any exception # Should not raise any exception
result = ResponsesAPIRequestUtils.get_requested_response_api_optional_param( result = ResponsesAPIRequestUtils.get_requested_response_api_optional_param(params)
params
)
assert "temperature" in result assert "temperature" in result