Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_/ui-role-gate-test-org-list

This commit is contained in:
Yuneng Jiang 2026-08-11 21:50:11 -07:00
commit e40fe7a404
No known key found for this signature in database
22 changed files with 1325 additions and 321 deletions

View file

@ -472,6 +472,8 @@ EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE: Final = float(
### ANTHROPIC CONSTANTS ###
ANTHROPIC_TOKEN_COUNTING_BETA_VERSION = os.getenv("ANTHROPIC_TOKEN_COUNTING_BETA_VERSION", "token-counting-2024-11-01")
ANTHROPIC_SKILLS_API_BETA_VERSION: Final = "skills-2025-10-02"
ANTHROPIC_BATCHES_ROUTE: Final = "/v1/messages/batches"
VERTEX_BATCH_PREDICTION_JOBS_ROUTE: Final = "batchPredictionJobs"
ANTHROPIC_WEB_SEARCH_TOOL_MAX_USES: Final = {
"low": 1,
"medium": 5,

View file

@ -2,6 +2,7 @@
Helper utilities for tracking the cost of built-in tools.
"""
from collections.abc import Mapping
from typing import Any, Final, Literal
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:
"""
Helper class for tracking the cost of built-in tools
@ -351,6 +360,10 @@ class StandardBuiltInToolCostTracking:
# 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:
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
elif isinstance(response_object, ResponsesAPIResponse):
# response api explicitly includes web_search_call in the output
@ -370,6 +383,8 @@ class StandardBuiltInToolCostTracking:
)
):
return True
if _usage_reports_server_side_web_search_calls(usage):
return True
return False
@ -432,7 +447,9 @@ class StandardBuiltInToolCostTracking:
"""
output: Final = response_object.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:
return True
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
import httpx
@ -12,13 +12,15 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
strip_name_from_messages,
)
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.types.llms.openai import AllMessageValues
from litellm.types.utils import (
Choices,
ModelResponse,
ModelResponseStream,
PromptTokensDetailsWrapper,
Usage,
)
@ -248,7 +250,7 @@ class XAIChatConfig(OpenAIGPTConfig):
XAI API returns empty string for finish_reason when using tools,
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
@ -351,25 +353,20 @@ class XAIChatConfig(OpenAIGPTConfig):
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:
return
usage: Final[Usage] = model_response.usage
num_sources_used = None
response_usage: Final = raw_response_json.get("usage", {})
if isinstance(response_usage, dict) and "num_sources_used" in response_usage:
num_sources_used = response_usage.get("num_sources_used")
# Map num_sources_used to web_search_requests for cost detection
if num_sources_used is not None and num_sources_used > 0:
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)
response_usage: Final = raw_response_json.get("usage")
if not isinstance(response_usage, dict):
return
details: Final = response_usage.get("server_side_tool_usage_details")
if isinstance(details, Mapping):
apply_server_side_tool_usage_details_to_usage(usage, details)
verbose_logger.debug("X.AI server_side_tool_usage_details: %s", details)
@staticmethod
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)
"""
from collections.abc import Mapping
from typing import TYPE_CHECKING, Final
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:
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]:
"""
@ -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)
completion_tokens: Final = int(getattr(usage, "completion_tokens", 0) or 0)
total_tokens: Final = int(getattr(usage, "total_tokens", 0) or 0)
reasoning_tokens = 0
if hasattr(usage, "completion_tokens_details") and usage.completion_tokens_details:
reasoning_tokens = int(getattr(usage.completion_tokens_details, "reasoning_tokens", 0) or 0)
reasoning_tokens: Final = (
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
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
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:
"""
Calculate the cost of web search requests for X.AI models.
X.AI Live Search costs $25 per 1,000 sources used.
Each source costs $0.025.
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.
Counts invocations from usage.server_side_tool_usage_details.web_search_calls.
Per-call rate comes from model_info.search_context_cost_per_query when set,
otherwise the default xAI tools rate ($5 / 1k calls).
"""
# Cost per source used: $25 per 1,000 sources = $0.025 per source
cost_per_source: Final = 25.0 / 1000.0 # $0.025
num_sources_used = 0
if (
hasattr(usage, "prompt_tokens_details")
and usage.prompt_tokens_details is not None
and hasattr(usage.prompt_tokens_details, "web_search_requests")
and usage.prompt_tokens_details.web_search_requests is not None
):
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
details: Final = getattr(usage, "server_side_tool_usage_details", None)
if not isinstance(details, Mapping):
return 0.0
try:
web_search_calls: Final = int(details.get("web_search_calls") or 0)
except (TypeError, ValueError):
return 0.0
if web_search_calls <= 0:
return 0.0
return _web_search_cost_per_call_from_model_info(model_info) * web_search_calls

View file

@ -1,4 +1,4 @@
from typing import TYPE_CHECKING, Any, Final
from typing import Any, Final
import litellm
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.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):
"""

View file

@ -12940,6 +12940,103 @@
"supports_system_messages": true,
"supports_tool_choice": false
},
"dashscope/deepseek-v4-flash": {
"cache_read_input_token_cost": 4e-08,
"input_cost_per_token": 2e-07,
"litellm_provider": "dashscope",
"max_input_tokens": 1000000,
"max_output_tokens": 393216,
"max_tokens": 393216,
"mode": "chat",
"output_cost_per_token": 4e-07,
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"dashscope/deepseek-v4-flash-0731": {
"cache_read_input_token_cost": 4e-08,
"input_cost_per_token": 2e-07,
"litellm_provider": "dashscope",
"max_input_tokens": 1000000,
"max_output_tokens": 393216,
"max_tokens": 393216,
"mode": "chat",
"output_cost_per_token": 4e-07,
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"dashscope/deepseek-v4-pro": {
"cache_read_input_token_cost": 2e-07,
"input_cost_per_token": 2.4e-06,
"litellm_provider": "dashscope",
"max_input_tokens": 1000000,
"max_output_tokens": 393216,
"max_tokens": 393216,
"mode": "chat",
"output_cost_per_token": 4.8e-06,
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"dashscope/glm-5.1": {
"cache_read_input_token_cost": 2.6e-07,
"input_cost_per_token": 1.4e-06,
"litellm_provider": "dashscope",
"max_input_tokens": 202745,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.4e-06,
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"dashscope/glm-5.2": {
"cache_read_input_token_cost": 2.8e-07,
"input_cost_per_token": 1.4e-06,
"litellm_provider": "dashscope",
"max_input_tokens": 1048576,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.4e-06,
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"dashscope/kimi-k2.7-code": {
"cache_read_input_token_cost": 1.9e-07,
"input_cost_per_token": 9.5e-07,
"litellm_provider": "dashscope",
"max_input_tokens": 229376,
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
"output_cost_per_token": 4e-06,
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"dashscope/qwen-coder": {
"input_cost_per_token": 3e-07,
"litellm_provider": "dashscope",
@ -13733,6 +13830,23 @@
}
]
},
"dashscope/qwen3.8-max": {
"cache_read_input_token_cost": 2.5e-07,
"input_cost_per_token": 2e-06,
"litellm_provider": "dashscope",
"max_input_tokens": 991808,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 6e-06,
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"dashscope/qwq-plus": {
"input_cost_per_token": 8e-07,
"litellm_provider": "dashscope",

View file

@ -39,11 +39,10 @@ _UNATTRIBUTED_TRACKABLE_CALL_TYPES: Final[frozenset[str]] = frozenset(
CallTypes.pass_through.value,
CallTypes.llm_passthrough_route.value,
CallTypes.allm_passthrough_route.value,
# CheckBatchCost's synthetic logging_obj for a completed managed batch only ever
# carries user_api_key_user_id (from LiteLLM_ManagedObjectTable.created_by) and
# user_api_key_team_id (from .team_id) -- both are None for batches created with
# the master key or a team-less key, since the table never stores the raw key
# hash. The batch already incurred real provider cost, so track it regardless.
# CheckBatchCost's synthetic logging_obj for a completed managed batch carries
# whatever LiteLLM_ManagedObjectTable stored at create time, and all of it is
# None for a batch created before those columns were persisted, or by the master
# key. The batch already incurred real provider cost, so track it regardless.
CallTypes.aretrieve_batch.value,
}
)

View file

@ -1,3 +1,4 @@
import asyncio
import json
from collections.abc import Sequence
from datetime import datetime
@ -7,6 +8,7 @@ import httpx
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import ANTHROPIC_BATCHES_ROUTE
from litellm.litellm_core_utils.core_helpers import map_finish_reason
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.litellm_logging import use_custom_pricing_for_model
@ -20,6 +22,12 @@ from litellm.llms.anthropic.chat.handler import (
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
from litellm.proxy._types import PassThroughEndpointLoggingTypedDict
from litellm.proxy.auth.auth_utils import get_end_user_id_from_request_body
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.batch_attribution import (
is_collection_route,
log_batch_registration_result,
optional_str,
request_tags_from_metadata,
)
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
PassthroughStandardLoggingPayload,
)
@ -833,13 +841,14 @@ class AnthropicPassthroughLoggingHandler:
# Store the managed object for cost tracking
# This will be picked up by check_batch_cost polling mechanism
AnthropicPassthroughLoggingHandler._store_batch_managed_object(
unified_object_id=unified_object_id,
batch_object=litellm_batch_response,
model_object_id=batch_id,
logging_obj=logging_obj,
**kwargs,
)
if is_collection_route(url_route, ANTHROPIC_BATCHES_ROUTE):
AnthropicPassthroughLoggingHandler._store_batch_managed_object(
unified_object_id=unified_object_id,
batch_object=litellm_batch_response,
model_object_id=batch_id,
logging_obj=logging_obj,
**kwargs,
)
# Create a batch job response for logging
litellm_model_response = ModelResponse()
@ -964,8 +973,12 @@ class AnthropicPassthroughLoggingHandler:
**kwargs,
) -> None:
"""
Store batch managed object for cost tracking.
Register a newly created batch for cost tracking.
This will be picked up by the check_batch_cost polling mechanism.
Only the create reaches here, so the row records the creating key and its tags.
An id-scoped route cannot rebuild the unified object id anyway: the model comes
from the create's request body, which a retrieve does not have.
"""
try:
# Get the managed files hook from the logging object
@ -981,7 +994,7 @@ class AnthropicPassthroughLoggingHandler:
user_api_key_dict: Final = UserAPIKeyAuth(
user_id=_request_metadata.get("user_api_key_user_id", "default-user"),
api_key="",
api_key=optional_str(_request_metadata.get("user_api_key")),
team_id=_request_metadata.get("user_api_key_team_id"),
team_alias=None,
user_role=LitellmUserRoles.CUSTOMER, # Use proper enum value
@ -1003,9 +1016,7 @@ class AnthropicPassthroughLoggingHandler:
)
# Store the unified object for batch cost tracking
import asyncio
asyncio.create_task(
task: Final = asyncio.create_task(
managed_files_hook.store_unified_object_id(
unified_object_id=unified_object_id,
file_object=batch_object,
@ -1013,13 +1024,14 @@ class AnthropicPassthroughLoggingHandler:
model_object_id=model_object_id,
file_purpose="batch",
user_api_key_dict=user_api_key_dict,
request_tags=request_tags_from_metadata(_request_metadata),
persist_attribution=True,
)
)
verbose_proxy_logger.info(
"Stored Anthropic batch managed object with unified_object_id=%s, batch_id=%s",
unified_object_id,
model_object_id,
task.add_done_callback(
lambda finished: log_batch_registration_result(
finished, "Anthropic", unified_object_id, model_object_id, is_batch_create=True
)
)
else:
verbose_proxy_logger.warning(

View file

@ -0,0 +1,78 @@
"""Spend attribution for batches created through a passthrough endpoint.
The creating key and its tags are read off the passthrough request's metadata and
persisted on the managed object row, because the batch cost lands hours later in a
background poll that has no request to read them from.
"""
import asyncio
from collections.abc import Mapping, Sequence
from typing import Final
from litellm._logging import verbose_proxy_logger
def optional_str(value: object) -> str | None:
return value if isinstance(value, str) else None
def _optional_str_tuple(value: object) -> tuple[str, ...] | None:
if not isinstance(value, list):
return None
items: Final[Sequence[object]] = value
return tuple(tag for tag in items if isinstance(tag, str))
def is_collection_route(url_route: str, collection_suffix: str) -> bool:
"""Whether the route addresses the batch collection itself rather than one batch.
A POST to the collection is the create; every id-scoped route is a retrieve,
results or cancel.
"""
return url_route.split("?")[0].rstrip("/").endswith(collection_suffix)
def request_tags_from_metadata(request_metadata: Mapping[str, object]) -> tuple[str, ...] | None:
"""Tags for the batch-cost spend row: the request's own tags when it sent any,
otherwise the key's tags, which auth exposes as user_api_key_auth_metadata (a
tagged key does not put its tags in the top-level metadata "tags" on the
passthrough path)
"""
tags: Final = _optional_str_tuple(request_metadata.get("tags"))
if tags:
return tags
key_auth_metadata: Final = request_metadata.get("user_api_key_auth_metadata")
if isinstance(key_auth_metadata, dict):
return _optional_str_tuple(key_auth_metadata.get("tags"))
return None
def log_batch_registration_result(
finished: asyncio.Task[None],
provider: str,
unified_object_id: str,
model_object_id: str,
is_batch_create: bool,
) -> None:
"""Report the outcome of the fire-and-forget managed object write. A create that
fails is not retried by a later poll, so its cost is never tracked at all.
"""
error: Final = finished.exception() if not finished.cancelled() else None
if finished.cancelled() or error is not None:
consequence: Final = (
"its cost will not be tracked" if is_batch_create else "its status and output file may be stale"
)
verbose_proxy_logger.error(
"Failed to store %s batch managed object with unified_object_id=%s, batch_id=%s; %s: %s",
provider,
unified_object_id,
model_object_id,
consequence,
error,
)
return
verbose_proxy_logger.info(
"Stored %s batch managed object with unified_object_id=%s, batch_id=%s",
provider,
unified_object_id,
model_object_id,
)

View file

@ -1,6 +1,5 @@
import asyncio
import re
from collections.abc import Mapping
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, cast
from urllib.parse import urlparse
@ -9,6 +8,7 @@ import httpx
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import VERTEX_BATCH_PREDICTION_JOBS_ROUTE
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
ModelResponseIterator as VertexModelResponseIterator,
@ -18,6 +18,12 @@ from litellm.llms.vertex_ai.vector_stores.search_api.transformation import (
)
from litellm.llms.vertex_ai.videos.transformation import VertexAIVideoConfig
from litellm.proxy._types import PassThroughEndpointLoggingTypedDict
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.batch_attribution import (
is_collection_route,
log_batch_registration_result,
optional_str,
request_tags_from_metadata,
)
from litellm.types.utils import (
Choices,
EmbeddingResponse,
@ -41,32 +47,6 @@ else:
EndpointType = Any
def _optional_str(value: object) -> str | None:
return value if isinstance(value, str) else None
def _optional_str_tuple(value: object) -> tuple[str, ...] | None:
if not isinstance(value, list):
return None
items: Final = cast(list[object], value) # cast-ok: isinstance-narrowed; element type unknown
return tuple(tag for tag in items if isinstance(tag, str))
def _request_tags(request_metadata: Mapping[str, object]) -> tuple[str, ...] | None:
"""Tags for the batch-cost spend row: the request's own tags when it sent any,
otherwise the key's tags, which auth exposes as user_api_key_auth_metadata (a
tagged key does not put its tags in the top-level metadata "tags" on the
passthrough path)
"""
tags: Final = _optional_str_tuple(request_metadata.get("tags"))
if tags:
return tags
key_auth_metadata: Final = request_metadata.get("user_api_key_auth_metadata")
if isinstance(key_auth_metadata, dict):
return _optional_str_tuple(key_auth_metadata.get("tags"))
return None
class VertexPassthroughLoggingHandler:
@staticmethod
def vertex_passthrough_handler(
@ -685,7 +665,7 @@ class VertexPassthroughLoggingHandler:
# Store the managed object for cost tracking
# This will be picked up by check_batch_cost polling mechanism
is_batch_create: Final = url_route.split("?")[0].rstrip("/").endswith("batchPredictionJobs")
is_batch_create: Final = is_collection_route(url_route, VERTEX_BATCH_PREDICTION_JOBS_ROUTE)
VertexPassthroughLoggingHandler._store_batch_managed_object(
unified_object_id=unified_object_id,
batch_object=litellm_batch_response,
@ -809,29 +789,6 @@ class VertexPassthroughLoggingHandler:
"kwargs": kwargs,
}
@staticmethod
def _log_batch_registration_result(
finished: asyncio.Task, unified_object_id: str, model_object_id: str, is_batch_create: bool
) -> None:
error: Final = finished.exception() if not finished.cancelled() else None
if finished.cancelled() or error is not None:
consequence: Final = (
"its cost will not be tracked" if is_batch_create else "its status and output file may be stale"
)
verbose_proxy_logger.error(
"Failed to store batch managed object with unified_object_id=%s, batch_id=%s; %s: %s",
unified_object_id,
model_object_id,
consequence,
error,
)
return
verbose_proxy_logger.info(
"Stored batch managed object with unified_object_id=%s, batch_id=%s",
unified_object_id,
model_object_id,
)
@staticmethod
def _store_batch_managed_object(
unified_object_id: str,
@ -863,7 +820,7 @@ class VertexPassthroughLoggingHandler:
user_api_key_dict: Final = UserAPIKeyAuth(
user_id=_request_metadata.get("user_api_key_user_id", "default-user"),
api_key=_optional_str(_request_metadata.get("user_api_key")),
api_key=optional_str(_request_metadata.get("user_api_key")),
team_id=_request_metadata.get("user_api_key_team_id"),
team_alias=None,
user_role=LitellmUserRoles.CUSTOMER, # Use proper enum value
@ -893,14 +850,14 @@ class VertexPassthroughLoggingHandler:
model_object_id=model_object_id,
file_purpose="batch",
user_api_key_dict=user_api_key_dict,
request_tags=_request_tags(_request_metadata),
request_tags=request_tags_from_metadata(_request_metadata),
persist_attribution=is_batch_create,
create_if_missing=is_batch_create,
)
)
task.add_done_callback(
lambda finished: VertexPassthroughLoggingHandler._log_batch_registration_result(
finished, unified_object_id, model_object_id, is_batch_create
lambda finished: log_batch_registration_result(
finished, "Vertex AI", unified_object_id, model_object_id, is_batch_create
)
)
else:

View file

@ -82,6 +82,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
self.sent_output_item_done_event: bool = False
self.sent_annotation_events: bool = False
self.litellm_model_response: ModelResponse | TextCompletionResponse | None = None
self.completed_response: Any = None
self.final_text: str = ""
self._cached_item_id: str | None = None
self._cached_response_id: str | None = None

View file

@ -1033,13 +1033,18 @@ class ResponseAPILoggingUtils:
@staticmethod
def _transform_response_api_usage_to_chat_usage(
usage_input: dict | ResponseAPIUsage | None,
usage_input: Mapping[str, object] | ResponseAPIUsage | Usage | None,
) -> Usage:
"""
Transforms ResponseAPIUsage or ImageUsage to a Usage object.
Both have the same spec with input_tokens, output_tokens, and
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:
return Usage(
@ -1047,6 +1052,10 @@ class ResponseAPILoggingUtils:
completion_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
if isinstance(usage_input, dict):
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"]
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"]
total_tokens = usage_input.get("total_tokens")
if total_tokens is None:
if usage_input.get("total_tokens") is None:
input_tokens: Final = usage_input.get("input_tokens")
output_tokens: Final = usage_input.get("output_tokens")
if input_tokens is not None and output_tokens is not None:
total_tokens = input_tokens + output_tokens
usage_input["total_tokens"] = total_tokens
if isinstance(input_tokens, int) and isinstance(output_tokens, int):
usage_input["total_tokens"] = input_tokens + output_tokens
response_api_usage = ResponseAPIUsage(**usage_input)
else:
response_api_usage = usage_input
@ -1089,12 +1096,27 @@ class ResponseAPILoggingUtils:
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(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens,
prompt_tokens_details=prompt_tokens_details,
completion_tokens_details=completion_tokens_details,
**extra_usage_fields,
)
# Preserve cost attribute if it exists on ResponseAPIUsage

View file

@ -12940,6 +12940,103 @@
"supports_system_messages": true,
"supports_tool_choice": false
},
"dashscope/deepseek-v4-flash": {
"cache_read_input_token_cost": 4e-08,
"input_cost_per_token": 2e-07,
"litellm_provider": "dashscope",
"max_input_tokens": 1000000,
"max_output_tokens": 393216,
"max_tokens": 393216,
"mode": "chat",
"output_cost_per_token": 4e-07,
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"dashscope/deepseek-v4-flash-0731": {
"cache_read_input_token_cost": 4e-08,
"input_cost_per_token": 2e-07,
"litellm_provider": "dashscope",
"max_input_tokens": 1000000,
"max_output_tokens": 393216,
"max_tokens": 393216,
"mode": "chat",
"output_cost_per_token": 4e-07,
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"dashscope/deepseek-v4-pro": {
"cache_read_input_token_cost": 2e-07,
"input_cost_per_token": 2.4e-06,
"litellm_provider": "dashscope",
"max_input_tokens": 1000000,
"max_output_tokens": 393216,
"max_tokens": 393216,
"mode": "chat",
"output_cost_per_token": 4.8e-06,
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"dashscope/glm-5.1": {
"cache_read_input_token_cost": 2.6e-07,
"input_cost_per_token": 1.4e-06,
"litellm_provider": "dashscope",
"max_input_tokens": 202745,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.4e-06,
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"dashscope/glm-5.2": {
"cache_read_input_token_cost": 2.8e-07,
"input_cost_per_token": 1.4e-06,
"litellm_provider": "dashscope",
"max_input_tokens": 1048576,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 4.4e-06,
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"dashscope/kimi-k2.7-code": {
"cache_read_input_token_cost": 1.9e-07,
"input_cost_per_token": 9.5e-07,
"litellm_provider": "dashscope",
"max_input_tokens": 229376,
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
"output_cost_per_token": 4e-06,
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"dashscope/qwen-coder": {
"input_cost_per_token": 3e-07,
"litellm_provider": "dashscope",
@ -13733,6 +13830,23 @@
}
]
},
"dashscope/qwen3.8-max": {
"cache_read_input_token_cost": 2.5e-07,
"input_cost_per_token": 2e-06,
"litellm_provider": "dashscope",
"max_input_tokens": 991808,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 6e-06,
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"dashscope/qwq-plus": {
"input_cost_per_token": 8e-07,
"litellm_provider": "dashscope",

View file

@ -90,6 +90,35 @@ def test_extract_partial_responses_usage_no_completed_response():
assert usage is None
def test_extract_partial_responses_usage_bridge_iterator_no_completed_response():
"""
Regression for #35411: the bridge iterator
(LiteLLMCompletionStreamingIterator) overrides __init__ without calling
super().__init__(), so completed_response was never set until the stream
reached RESPONSE_COMPLETED. On a mid-stream provider error (before
completion) the fallback recovery path read source_iterator.completed_response
and raised AttributeError, masking the real provider error and bypassing
fallbacks. The attribute must always exist and default to None.
"""
from litellm.responses.litellm_completion_transformation.streaming_iterator import (
LiteLLMCompletionStreamingIterator,
)
wrapper = MagicMock()
wrapper.logging_obj = MagicMock()
iterator = LiteLLMCompletionStreamingIterator(
model="anthropic/claude-sonnet-4-5",
litellm_custom_stream_wrapper=wrapper,
request_input="hi",
responses_api_request={},
)
assert iterator.completed_response is None
# No chat chunks collected yet and no completed_response → must return
# None instead of raising AttributeError.
assert Router._extract_partial_responses_usage(iterator) is None
# -------- _combine_responses_fallback_usage --------

View file

@ -604,8 +604,8 @@ async def test_create_still_upserts_and_claims_attribution():
@pytest.mark.asyncio
async def test_default_callers_still_create_their_rows():
"""create_if_missing defaults to True, so the fine-tune, Responses and Anthropic
callers, none of which pass it, keep upserting exactly as before."""
"""create_if_missing defaults to True, so the fine-tune, Responses and managed
/v1/batches callers, none of which passes it, keep upserting exactly as before."""
managed_files, mock_prisma = _make_object_store_instance()
await managed_files.store_unified_object_id(

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.
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
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
# 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 sys
from unittest.mock import MagicMock
sys.path.insert(0, os.path.abspath("../../../../.."))
import pytest
import litellm
from litellm.llms.xai.responses.transformation import XAIResponsesAPIConfig
from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams
from litellm.types.utils import LlmProviders
from litellm.responses.utils import ResponseAPILoggingUtils
from litellm.types.llms.openai import (
ResponseAPIUsage,
ResponseCompletedEvent,
ResponsesAPIOptionalRequestParams,
ResponsesAPIResponse,
)
from litellm.types.utils import LlmProviders, Usage
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 isinstance(
config, XAIResponsesAPIConfig
), f"Expected XAIResponsesAPIConfig, got {type(config)}"
assert (
config.custom_llm_provider == LlmProviders.XAI
), "custom_llm_provider should be XAI"
assert isinstance(config, XAIResponsesAPIConfig), 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):
"""Test that container field is removed from code_interpreter tools"""
config = XAIResponsesAPIConfig()
params = ResponsesAPIOptionalRequestParams(
tools=[{"type": "code_interpreter", "container": {"type": "auto"}}]
)
params = ResponsesAPIOptionalRequestParams(tools=[{"type": "code_interpreter", "container": {"type": "auto"}}])
result = config.map_openai_params(
response_api_optional_params=params, model="grok-4-fast", drop_params=False
)
result = config.map_openai_params(response_api_optional_params=params, model="grok-4-fast", drop_params=False)
assert "tools" in result
assert len(result["tools"]) == 1
assert result["tools"][0]["type"] == "code_interpreter"
assert (
"container" not in result["tools"][0]
), "Container field should be removed"
assert "container" not in result["tools"][0], "Container field should be removed"
def test_instructions_parameter_dropped(self):
"""Test that instructions parameter is dropped for XAI"""
config = XAIResponsesAPIConfig()
params = ResponsesAPIOptionalRequestParams(
instructions="You are a helpful assistant.", temperature=0.7
)
params = ResponsesAPIOptionalRequestParams(instructions="You are a helpful assistant.", temperature=0.7)
result = config.map_openai_params(
response_api_optional_params=params, model="grok-4-fast", drop_params=False
)
result = config.map_openai_params(response_api_optional_params=params, model="grok-4-fast", drop_params=False)
assert "instructions" not in result, "Instructions should be dropped"
assert result.get("temperature") == 0.7, "Other params should be preserved"
@ -88,25 +82,15 @@ class TestXAIResponsesAPITransformation:
# Test with default XAI API base
url = config.get_complete_url(api_base=None, litellm_params={})
assert (
url == "https://api.x.ai/v1/responses"
), f"Expected XAI responses endpoint, got {url}"
assert url == "https://api.x.ai/v1/responses", f"Expected XAI responses endpoint, got {url}"
# Test with custom api_base
custom_url = config.get_complete_url(
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}"
custom_url = config.get_complete_url(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}"
# Test with trailing slash
url_with_slash = config.get_complete_url(
api_base="https://api.x.ai/v1/", litellm_params={}
)
assert (
url_with_slash == "https://api.x.ai/v1/responses"
), "Should handle trailing slash"
url_with_slash = config.get_complete_url(api_base="https://api.x.ai/v1/", litellm_params={})
assert url_with_slash == "https://api.x.ai/v1/responses", "Should handle trailing slash"
def test_web_search_tool_transformation(self):
"""Test that web_search tools are transformed to XAI format"""
@ -167,9 +151,7 @@ class TestXAIResponsesAPITransformation:
config = XAIResponsesAPIConfig()
params = ResponsesAPIOptionalRequestParams(
tools=[
{"type": "web_search", "excluded_domains": ["example.com", "test.com"]}
]
tools=[{"type": "web_search", "excluded_domains": ["example.com", "test.com"]}]
)
result = config.map_openai_params(
@ -309,3 +291,115 @@ class TestXAIResponsesAPITransformation:
# Verify function tool is unchanged
assert result["tools"][3]["type"] == "function"
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("../../../..")
) # Adds the parent directory to the system path
import pytest
import litellm
from litellm.llms.xai.chat.transformation import XAIChatConfig
from litellm.types.utils import (
CompletionTokensDetailsWrapper,
@ -135,3 +138,65 @@ class TestXAIUsageNormalization:
XAIChatConfig._normalize_openai_compatible_usage_totals(usage)
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("../../..")
) # 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:
@ -354,76 +363,53 @@ class TestXAICostCalculator:
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
def test_web_search_cost_calculation(self):
"""Test web search cost calculation for X.AI models."""
# Test with web_search_requests in prompt_tokens_details (primary path)
usage = Usage(
prompt_tokens=100,
completion_tokens=50,
total_tokens=150,
prompt_tokens_details=PromptTokensDetailsWrapper(
text_tokens=100,
web_search_requests=3, # 3 sources used
),
def test_web_search_cost_via_server_side_tool_usage_details(self):
"""usage.server_side_tool_usage_details.web_search_calls at default $5/1k."""
usage = Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150)
setattr(
usage,
"server_side_tool_usage_details",
{
"web_search_calls": 3,
"x_search_calls": 0,
"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={})
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
expected_cost = 3 * (25.0 / 1000.0) # 3 * $0.025
assert math.isclose(web_search_cost, expected_cost, rel_tol=1e-10)
assert math.isclose(web_search_cost, 0.075, rel_tol=1e-10)
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
usage = Usage(
prompt_tokens=100,
completion_tokens=50,
total_tokens=150,
def test_web_search_cost_uses_model_info_search_context_pricing(self):
usage = Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15)
setattr(usage, "server_side_tool_usage_details", {"web_search_calls": 2})
model_info = {
"search_context_cost_per_query": {
"search_context_size_medium": 0.01,
}
}
web_search_cost = cost_per_web_search_request(
usage=usage, model_info=model_info
)
# Manually set num_sources_used (as done by transformation layer)
setattr(usage, "num_sources_used", 5)
assert math.isclose(web_search_cost, 0.02, rel_tol=1e-10)
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
expected_cost = 5 * (25.0 / 1000.0) # 5 * $0.025
assert math.isclose(web_search_cost, expected_cost, rel_tol=1e-10)
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
),
def test_apply_details_sets_web_search_requests_for_cost_gate(self):
usage = Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15)
apply_server_side_tool_usage_details_to_usage(
usage, {"web_search_calls": 2, "x_search_calls": 0}
)
web_search_cost = cost_per_web_search_request(usage=usage, model_info={})
# Expected cost: 0 sources * $0.025 per source = $0.0
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,
assert usage.prompt_tokens_details is not None
assert usage.prompt_tokens_details.web_search_requests == 2
assert StandardBuiltInToolCostTracking.response_object_includes_web_search_call(
response_object=object(), usage=usage
)
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):
"""Test cost calculation for grok-4.20-beta-0309-reasoning model."""
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(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,3 +1,4 @@
import asyncio
import json
import os
import sys
@ -17,6 +18,13 @@ from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passth
)
async def _drain_tasks():
"""Await the fire-and-forget managed object write and let its done callback run."""
pending = [t for t in asyncio.all_tasks() if t is not asyncio.current_task()]
await asyncio.gather(*pending, return_exceptions=True)
await asyncio.sleep(0)
class TestAnthropicLoggingHandlerModelFallback:
"""Test the model fallback logic in the anthropic passthrough logging handler."""
@ -925,6 +933,114 @@ class TestAnthropicBatchPassthroughCostTracking:
assert call_kwargs["user_api_key_dict"].user_id == expected_user_id
assert call_kwargs["user_api_key_dict"].team_id == expected_team_id
async def _store_with_metadata(self, mock_logging_obj, metadata):
mock_managed_files_hook = MagicMock()
mock_managed_files_hook.store_unified_object_id = AsyncMock()
with (
patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_pl,
patch(
"litellm.proxy.pass_through_endpoints.llm_provider_handlers.batch_attribution.verbose_proxy_logger"
),
):
mock_pl.get_proxy_hook.return_value = mock_managed_files_hook
AnthropicPassthroughLoggingHandler._store_batch_managed_object(
unified_object_id="uoi",
batch_object={"id": "b1", "object": "batch", "status": "validating"},
model_object_id="b1",
logging_obj=mock_logging_obj,
litellm_params={"metadata": metadata},
)
await _drain_tasks()
mock_managed_files_hook.store_unified_object_id.assert_awaited_once()
return mock_managed_files_hook.store_unified_object_id.call_args[1]
@pytest.mark.asyncio
async def test_create_persists_key_hash_and_tags(self, mock_logging_obj):
"""Regression (LIT-5288): the batch create must persist the creating key's hashed
token and its tags so CheckBatchCost can attribute the batch-cost spend row to the
key, team and tags. Before this fix the stored api_key was always "" and no tags
were stored, so key/team/tag spend and budgets never moved for batch usage."""
call_kwargs = await self._store_with_metadata(
mock_logging_obj,
{
"user_api_key": "hashed-key-a",
"user_api_key_user_id": "alice",
"user_api_key_team_id": "team-alpha",
"user_api_key_auth_metadata": {"tags": ["env:prod", 7, "team:ml"]},
},
)
assert call_kwargs["user_api_key_dict"].api_key == "hashed-key-a"
assert call_kwargs["request_tags"] == ("env:prod", "team:ml")
assert call_kwargs["persist_attribution"] is True
@pytest.mark.asyncio
async def test_failed_create_write_is_reported_not_swallowed(self, mock_logging_obj):
"""The managed object write is fire-and-forget, and only the create writes the row,
so a failed create is never back-filled by a later retrieve and that batch's cost
is never tracked. The failure has to reach the log instead of being reported as a
success."""
mock_managed_files_hook = MagicMock()
mock_managed_files_hook.store_unified_object_id = AsyncMock(
side_effect=RuntimeError("db down")
)
with (
patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_pl,
patch(
"litellm.proxy.pass_through_endpoints.llm_provider_handlers.batch_attribution.verbose_proxy_logger"
) as mock_logger,
):
mock_pl.get_proxy_hook.return_value = mock_managed_files_hook
AnthropicPassthroughLoggingHandler._store_batch_managed_object(
unified_object_id="uoi",
batch_object={"id": "b1", "object": "batch", "status": "validating"},
model_object_id="b1",
logging_obj=mock_logging_obj,
litellm_params={"metadata": {"user_api_key": "hashed-key-a"}},
)
await _drain_tasks()
mock_logger.info.assert_not_called()
mock_logger.error.assert_called_once()
assert "its cost will not be tracked" in mock_logger.error.call_args[0]
assert "Anthropic" in mock_logger.error.call_args[0]
@pytest.mark.parametrize(
"url_route, registers",
[
("https://api.anthropic.com/v1/messages/batches", True),
("https://api.anthropic.com/v1/messages/batches/", True),
("https://api.anthropic.com/v1/messages/batches?limit=20", True),
("https://api.anthropic.com/v1/messages/batches/msgbatch_123", False),
("https://api.anthropic.com/v1/messages/batches/msgbatch_123/results", False),
("https://api.anthropic.com/v1/messages/batches/msgbatch_123/cancel", False),
],
)
def test_batch_is_registered_from_the_create_route_only(
self, mock_logging_obj, mock_httpx_response, mock_request_body, url_route, registers
):
"""Only a POST to the collection route registers the batch. Every id-scoped route
is a retrieve, results or cancel, and none of them can rebuild the unified object
id anyway: it embeds the model, which comes from the create's request body. Before
this gate an id-scoped route reached the store with a mismatched id, where it could
only either claim a row it did not create or fail the model_object_id unique
constraint."""
with patch.object(
AnthropicPassthroughLoggingHandler, "_store_batch_managed_object"
) as mock_store:
AnthropicPassthroughLoggingHandler.batch_creation_handler(
httpx_response=mock_httpx_response,
logging_obj=mock_logging_obj,
url_route=url_route,
result="success",
start_time=datetime.now(),
end_time=datetime.now(),
cache_hit=False,
request_body=mock_request_body,
)
assert mock_store.call_count == (1 if registers else 0)
def test_batch_creation_handler_failure_status_code(
self, mock_logging_obj, mock_request_body
):
@ -978,6 +1094,7 @@ class TestAnthropicBatchPassthroughCostTracking:
batch_object=batch_object,
model_object_id="msgbatch_123",
logging_obj=mock_logging_obj,
is_batch_create=True,
user_id="test-user",
)

View file

@ -0,0 +1,159 @@
import asyncio
from unittest.mock import patch
import pytest
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.batch_attribution import (
is_collection_route,
log_batch_registration_result,
optional_str,
request_tags_from_metadata,
)
@pytest.mark.parametrize(
"value, expected",
[("a", "a"), ("", ""), (None, None), (7, None), (["a"], None)],
)
def test_optional_str(value, expected):
assert optional_str(value) == expected
class TestRequestTagsFromMetadata:
"""Tags for the batch-cost spend row. These feed LiteLLM_ManagedObjectTable.request_tags,
which is the only record of the creating request's tags by the time CheckBatchCost bills
the batch hours later."""
@pytest.mark.parametrize(
"metadata, expected",
[
# a request that sent its own tags (x-litellm-tags header or body metadata)
({"tags": ["req:a", "req:b"]}, ("req:a", "req:b")),
# request tags win over the key's own tags
(
{"tags": ["req:a"], "user_api_key_auth_metadata": {"tags": ["key:b"]}},
("req:a",),
),
# no request tags: fall back to the tags the key itself carries, because a
# tagged key does not put its tags in the top-level metadata on this path
({"user_api_key_auth_metadata": {"tags": ["key:b"]}}, ("key:b",)),
# an empty request tag list is not a selection, so the key's tags still apply
(
{"tags": [], "user_api_key_auth_metadata": {"tags": ["key:b"]}},
("key:b",),
),
# neither: no tags on the spend row
({}, None),
# order is preserved, so the spend row is reproducible
({"tags": ["z", "a", "m"]}, ("z", "a", "m")),
],
)
def test_precedence(self, metadata, expected):
assert request_tags_from_metadata(metadata) == expected
@pytest.mark.parametrize(
"raw, expected",
[
# non-string entries are dropped rather than crashing the create
(["env:prod", 7, None, "team:ml"], ("env:prod", "team:ml")),
# nothing usable survives, so this is treated as no request tags at all
([7, None], None),
# a non-list is not a tag list
("env:prod", None),
({"env": "prod"}, None),
(None, None),
],
)
def test_malformed_tags_are_dropped(self, raw, expected):
assert request_tags_from_metadata({"tags": raw}) == expected
def test_malformed_key_auth_metadata_is_ignored(self):
assert request_tags_from_metadata({"user_api_key_auth_metadata": "nope"}) is None
@pytest.mark.parametrize(
"url_route, suffix, expected",
[
("https://api.anthropic.com/v1/messages/batches", "/v1/messages/batches", True),
("https://api.anthropic.com/v1/messages/batches/", "/v1/messages/batches", True),
("https://api.anthropic.com/v1/messages/batches?limit=20", "/v1/messages/batches", True),
("https://api.anthropic.com/v1/messages/batches/msgbatch_1", "/v1/messages/batches", False),
# a proxied base with a path prefix still resolves, because this is a suffix match
("https://gateway.internal/anthropic/v1/messages/batches", "/v1/messages/batches", True),
("https://aiplatform.googleapis.com/v1/projects/p/locations/l/batchPredictionJobs", "batchPredictionJobs", True),
("https://aiplatform.googleapis.com/v1/projects/p/locations/l/batchPredictionJobs/9", "batchPredictionJobs", False),
],
)
def test_is_collection_route(url_route, suffix, expected):
assert is_collection_route(url_route, suffix) is expected
class TestLogBatchRegistrationResult:
"""The managed object write is fire and forget, so its outcome only ever reaches an
operator through this log line."""
@staticmethod
async def _finished_task(coro):
task = asyncio.ensure_future(coro)
await asyncio.gather(task, return_exceptions=True)
return task
@pytest.mark.asyncio
async def test_success_names_the_provider(self):
async def ok():
return None
task = await self._finished_task(ok())
with patch(
"litellm.proxy.pass_through_endpoints.llm_provider_handlers.batch_attribution.verbose_proxy_logger"
) as logger:
log_batch_registration_result(task, "Anthropic", "uoi", "b1", is_batch_create=True)
logger.error.assert_not_called()
logger.info.assert_called_once()
assert "Anthropic" in logger.info.call_args[0]
@pytest.mark.asyncio
async def test_a_failed_create_says_the_cost_is_lost(self):
async def boom():
raise RuntimeError("db down")
task = await self._finished_task(boom())
with patch(
"litellm.proxy.pass_through_endpoints.llm_provider_handlers.batch_attribution.verbose_proxy_logger"
) as logger:
log_batch_registration_result(task, "Vertex AI", "uoi", "b1", is_batch_create=True)
logger.info.assert_not_called()
assert "its cost will not be tracked" in logger.error.call_args[0]
@pytest.mark.asyncio
async def test_a_failed_refresh_says_the_row_is_stale(self):
async def boom():
raise RuntimeError("db down")
task = await self._finished_task(boom())
with patch(
"litellm.proxy.pass_through_endpoints.llm_provider_handlers.batch_attribution.verbose_proxy_logger"
) as logger:
log_batch_registration_result(task, "Vertex AI", "uoi", "b1", is_batch_create=False)
logger.info.assert_not_called()
assert "its status and output file may be stale" in logger.error.call_args[0]
@pytest.mark.asyncio
async def test_a_cancelled_write_is_reported_not_reraised(self):
async def slow():
await asyncio.sleep(60)
task = asyncio.ensure_future(slow())
await asyncio.sleep(0)
task.cancel()
await asyncio.gather(task, return_exceptions=True)
with patch(
"litellm.proxy.pass_through_endpoints.llm_provider_handlers.batch_attribution.verbose_proxy_logger"
) as logger:
log_batch_registration_result(task, "Anthropic", "uoi", "b1", is_batch_create=True)
logger.info.assert_not_called()
logger.error.assert_called_once()

View file

@ -1,21 +1,16 @@
import base64
import json
import os
import sys
from unittest.mock import MagicMock, patch
import pytest
from fastapi.testclient import TestClient
sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
sys.path.insert(0, os.path.abspath("../../..")) # Adds the parent directory to the system path
import litellm
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
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
@ -54,9 +49,7 @@ class TestResponsesAPIRequestUtils:
# Setup
model = "gpt-4o"
config = OpenAIResponsesAPIConfig()
optional_params = ResponsesAPIOptionalRequestParams(
{"temperature": 0.7, "unsupported_param": "value"}
)
optional_params = ResponsesAPIOptionalRequestParams({"temperature": 0.7, "unsupported_param": "value"})
# Execute and Assert
with pytest.raises(litellm.UnsupportedParamsError) as excinfo:
@ -90,9 +83,7 @@ class TestResponsesAPIRequestUtils:
assert result == {"temperature": 0.7}
@pytest.mark.parametrize("request_drop_params", [None, False])
def test_get_optional_params_responses_api_still_raises_without_drop(
self, monkeypatch, request_drop_params
):
def test_get_optional_params_responses_api_still_raises_without_drop(self, monkeypatch, request_drop_params):
"""Absent or False request-level drop_params must not suppress the unsupported-param error"""
monkeypatch.setattr(litellm, "drop_params", False)
config = OpenAIResponsesAPIConfig()
@ -119,9 +110,7 @@ class TestResponsesAPIRequestUtils:
}
# Execute
result = ResponsesAPIRequestUtils.get_requested_response_api_optional_param(
params
)
result = ResponsesAPIRequestUtils.get_requested_response_api_optional_param(params)
# Assert
assert "temperature" in result
@ -147,40 +136,31 @@ class TestResponsesAPIRequestUtils:
)
# Execute
result = ResponsesAPIRequestUtils.decode_previous_response_id_to_original_previous_response_id(
encoded_id
)
result = ResponsesAPIRequestUtils.decode_previous_response_id_to_original_previous_response_id(encoded_id)
# Assert
assert result == original_response_id
# Test with a non-encoded ID
plain_id = "resp_xyz789"
result_plain = ResponsesAPIRequestUtils.decode_previous_response_id_to_original_previous_response_id(
plain_id
)
result_plain = ResponsesAPIRequestUtils.decode_previous_response_id_to_original_previous_response_id(plain_id)
assert result_plain == plain_id
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"""
responses_api_response = {"id": "resp_abc123"}
litellm_metadata = {"model_info": {"id": "gpt-4o"}}
updated = (
ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id(
responses_api_response=responses_api_response,
custom_llm_provider="openai",
litellm_metadata=litellm_metadata,
)
updated = ResponsesAPIRequestUtils._update_responses_api_response_id_with_model_id(
responses_api_response=responses_api_response,
custom_llm_provider="openai",
litellm_metadata=litellm_metadata,
)
assert updated["id"] != "resp_abc123"
decoded = ResponsesAPIRequestUtils._decode_responses_api_response_id(
updated["id"]
)
decoded = ResponsesAPIRequestUtils._decode_responses_api_response_id(updated["id"])
assert decoded.get("response_id") == "resp_abc123"
assert decoded.get("model_id") == "gpt-4o"
assert decoded.get("custom_llm_provider") == "openai"
def test_update_responses_api_response_id_with_model_id_is_idempotent_for_litellm_ids(self):
raw = "resp_" + "a" * 48
litellm_metadata = {"model_info": {"id": "model-123"}}
@ -207,9 +187,7 @@ class TestResponsesAPIRequestUtils:
model_id=None,
container_id="cntr_upstream_abc",
)
assert "None" not in base64.b64decode(
encoded.replace("cntr_", "").encode("utf-8")
).decode("utf-8")
assert "None" not in base64.b64decode(encoded.replace("cntr_", "").encode("utf-8")).decode("utf-8")
decoded = ResponsesAPIRequestUtils._decode_container_id(encoded)
assert decoded.get("custom_llm_provider") == "azure"
assert decoded.get("model_id") is None
@ -217,12 +195,8 @@ class TestResponsesAPIRequestUtils:
def test_decode_container_id_legacy_literal_none_model_id(self):
"""IDs encoded before the None fix should decode without a bogus model_id."""
legacy_inner = (
"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_inner = "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")
decoded = ResponsesAPIRequestUtils._decode_container_id(legacy_id)
assert decoded.get("model_id") is None
assert decoded.get("custom_llm_provider") == "azure"
@ -264,19 +238,14 @@ class TestResponseAPILoggingUtils:
}
# Execute
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
usage
)
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
# Assert
assert isinstance(result, Usage)
assert result.prompt_tokens == 10
assert result.completion_tokens == 20
assert result.total_tokens == 30
assert (
result.prompt_tokens_details
and result.prompt_tokens_details.cached_tokens == 2
)
assert result.prompt_tokens_details and result.prompt_tokens_details.cached_tokens == 2
def test_transform_response_api_usage_with_none_values(self):
"""Test transformation handles None values properly"""
@ -289,9 +258,7 @@ class TestResponseAPILoggingUtils:
}
# Execute
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
usage
)
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
# Assert
assert result.prompt_tokens == 0
@ -310,9 +277,7 @@ class TestResponseAPILoggingUtils:
}
# Execute
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
usage
)
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
# Assert
assert result.prompt_tokens == 15
@ -349,9 +314,7 @@ class TestResponseAPILoggingUtils:
}
# Execute
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
usage
)
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
# Assert - verify basic token counts
assert isinstance(result, Usage)
@ -386,9 +349,7 @@ class TestResponseAPILoggingUtils:
},
}
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
usage
)
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
assert result.prompt_tokens_details is not None
assert result.prompt_tokens_details.cache_write_tokens == 10059
@ -417,9 +378,7 @@ class TestResponseAPILoggingUtils:
}
# Execute
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
usage
)
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
# Assert - all token detail types should be preserved
assert result.prompt_tokens_details is not None
@ -451,9 +410,7 @@ class TestResponseAPILoggingUtils:
},
}
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
usage
)
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
assert result.prompt_tokens_details is not None
assert result.prompt_tokens_details.text_tokens == 8
@ -475,9 +432,7 @@ class TestResponseAPILoggingUtils:
"output_token_details": {"text_tokens": 2, "audio_tokens": 98},
}
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
usage
)
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage)
assert result.prompt_tokens_details is not None
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.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:
"""
@ -503,9 +545,7 @@ class TestResponsesAPIProviderSpecificParams:
}
# Should not raise any exception
result = ResponsesAPIRequestUtils.get_requested_response_api_optional_param(
params
)
result = ResponsesAPIRequestUtils.get_requested_response_api_optional_param(params)
assert "temperature" in result
def test_provider_specific_params_no_crash_with_openai(self):
@ -517,9 +557,7 @@ class TestResponsesAPIProviderSpecificParams:
}
# Should not raise any exception
result = ResponsesAPIRequestUtils.get_requested_response_api_optional_param(
params
)
result = ResponsesAPIRequestUtils.get_requested_response_api_optional_param(params)
assert "temperature" in result
def test_provider_specific_params_no_crash_with_vertex_ai(self):
@ -531,9 +569,7 @@ class TestResponsesAPIProviderSpecificParams:
}
# Should not raise any exception
result = ResponsesAPIRequestUtils.get_requested_response_api_optional_param(
params
)
result = ResponsesAPIRequestUtils.get_requested_response_api_optional_param(params)
assert "temperature" in result