mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_/ui-role-gate-test-org-list
This commit is contained in:
commit
e40fe7a404
22 changed files with 1325 additions and 321 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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 --------
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue