mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
Merge pull request #30817 from geraint0923/litellm_fix_xai_web_search_cost_billing
Some checks are pending
CI Coverage / assert-ci-coverage (push) Waiting to run
CodSpeed Benchmarks / benchmarks (push) Waiting to run
Publish basedpyright base counts / publish (push) Waiting to run
Code Quality Checks / code-quality (push) Waiting to run
UI Unit Tests / ui-unit-tests (push) Waiting to run
Unit Tests: Core Utilities / core-utils (push) Waiting to run
Unit Tests: Documentation Validation / documentation (push) Waiting to run
Unit Tests: Enterprise, Google GenAI & Routing / enterprise-routing (push) Waiting to run
Unit Tests: Integrations (Callbacks & Logging) / integrations (push) Waiting to run
Unit Tests: LLM Provider Transformations / Vertex AI (push) Waiting to run
Unit Tests: LLM Provider Transformations / All Other Providers (push) Waiting to run
Unit Tests: MCP, Secrets, Containers & Misc / misc (push) Waiting to run
Unit Tests: Proxy Auth & Key Management / proxy-auth (push) Waiting to run
Unit Tests: Proxy DB Operations / assert-shard-coverage (push) Waiting to run
Unit Tests: Proxy DB Operations / auth-checks (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / budgets (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / custom-logging (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / db-and-spend (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / endpoints-and-responses (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / guardrails-hooks (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-utils (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / jwt-and-keys (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / key-generation (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / logging-misc (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-runtime (push) Blocked by required conditions
Unit Tests: Proxy Infrastructure / proxy-infra (push) Waiting to run
Unit Tests: Proxy Legacy Tests / auth-and-jwt (push) Waiting to run
Unit Tests: Proxy Legacy Tests / key-generation (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-config (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-response-and-misc (push) Waiting to run
Unit Tests: Proxy DB Operations / proxy-server-core (push) Blocked by required conditions
Unit Tests: Proxy API Endpoints / proxy-endpoints (push) Waiting to run
Unit Tests: Proxy API Endpoints / proxy-server (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-server (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-server-extras (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-token-counter (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-user-auth-and-spend (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-utils (push) Waiting to run
Unit Tests: Responses, Caching & Types / responses-caching-types (push) Waiting to run
GitHub Actions Security Analysis / zizmor (push) Waiting to run
Some checks are pending
CI Coverage / assert-ci-coverage (push) Waiting to run
CodSpeed Benchmarks / benchmarks (push) Waiting to run
Publish basedpyright base counts / publish (push) Waiting to run
Code Quality Checks / code-quality (push) Waiting to run
UI Unit Tests / ui-unit-tests (push) Waiting to run
Unit Tests: Core Utilities / core-utils (push) Waiting to run
Unit Tests: Documentation Validation / documentation (push) Waiting to run
Unit Tests: Enterprise, Google GenAI & Routing / enterprise-routing (push) Waiting to run
Unit Tests: Integrations (Callbacks & Logging) / integrations (push) Waiting to run
Unit Tests: LLM Provider Transformations / Vertex AI (push) Waiting to run
Unit Tests: LLM Provider Transformations / All Other Providers (push) Waiting to run
Unit Tests: MCP, Secrets, Containers & Misc / misc (push) Waiting to run
Unit Tests: Proxy Auth & Key Management / proxy-auth (push) Waiting to run
Unit Tests: Proxy DB Operations / assert-shard-coverage (push) Waiting to run
Unit Tests: Proxy DB Operations / auth-checks (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / budgets (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / custom-logging (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / db-and-spend (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / endpoints-and-responses (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / guardrails-hooks (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-utils (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / jwt-and-keys (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / key-generation (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / logging-misc (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-runtime (push) Blocked by required conditions
Unit Tests: Proxy Infrastructure / proxy-infra (push) Waiting to run
Unit Tests: Proxy Legacy Tests / auth-and-jwt (push) Waiting to run
Unit Tests: Proxy Legacy Tests / key-generation (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-config (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-response-and-misc (push) Waiting to run
Unit Tests: Proxy DB Operations / proxy-server-core (push) Blocked by required conditions
Unit Tests: Proxy API Endpoints / proxy-endpoints (push) Waiting to run
Unit Tests: Proxy API Endpoints / proxy-server (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-server (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-server-extras (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-token-counter (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-user-auth-and-spend (push) Waiting to run
Unit Tests: Proxy Legacy Tests / proxy-utils (push) Waiting to run
Unit Tests: Responses, Caching & Types / responses-caching-types (push) Waiting to run
GitHub Actions Security Analysis / zizmor (push) Waiting to run
fix(xai): bill web_search from server_side_tool_usage_details
This commit is contained in:
commit
b4f5e46a44
10 changed files with 664 additions and 242 deletions
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,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