Add Xai websearch cost (#16001)

* Add xai websearch cost

* Add test for websearch cost in xai

* remove not required changes
This commit is contained in:
Sameer Kankute
2025-10-30 20:35:34 -07:00
committed by GitHub
parent 297c2a0bf4
commit 689b210c18
4 changed files with 144 additions and 3 deletions
+3
View File
@@ -43,6 +43,9 @@ def get_cost_for_web_search_request(
# Perplexity handles search costs internally in its own cost calculator
# Return 0.0 to indicate costs are already accounted for
return 0.0
elif custom_llm_provider == "xai":
from .xai.cost_calculator import cost_per_web_search_request
return cost_per_web_search_request(usage=usage, model_info=model_info)
else:
return None
+33 -1
View File
@@ -10,7 +10,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
)
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import Choices, ModelResponse
from litellm.types.utils import Choices, ModelResponse, Usage, PromptTokensDetailsWrapper
from ...openai.chat.gpt_transformation import OpenAIGPTConfig
@@ -170,6 +170,8 @@ 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.
"""
# First, let the parent class handle the standard transformation
@@ -193,4 +195,34 @@ class XAIChatConfig(OpenAIGPTConfig):
if isinstance(choice, Choices):
self._fix_choice_finish_reason_for_tool_calls(choice)
# Handle X.AI web search usage tracking
try:
raw_response_json = raw_response.json()
self._enhance_usage_with_xai_web_search_fields(response, raw_response_json)
except Exception as e:
verbose_logger.debug(f"Error extracting X.AI web search usage: {e}")
return response
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.
"""
if not hasattr(model_response, "usage") or model_response.usage is None:
return
usage: Usage = model_response.usage
num_sources_used = None
response_usage = 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(f"X.AI web search sources used: {num_sources_used}")
+36 -1
View File
@@ -4,11 +4,14 @@ Helper util for handling XAI-specific cost calculation
- Handles XAI-specific reasoning token billing (billed as part of completion tokens)
"""
from typing import Tuple
from typing import TYPE_CHECKING, Tuple
from litellm.types.utils import Usage
from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token
if TYPE_CHECKING:
from litellm.types.utils import ModelInfo
def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]:
"""
@@ -46,3 +49,35 @@ def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]:
)
return prompt_cost, completion_cost
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.
"""
# Cost per source used: $25 per 1,000 sources = $0.025 per source
cost_per_source = 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 = cost_per_source * num_sources_used
return total_cost
@@ -9,6 +9,7 @@ import sys
import litellm
from litellm.types.utils import (
CompletionTokensDetailsWrapper,
PromptTokensDetailsWrapper,
Usage,
)
@@ -16,7 +17,7 @@ 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
from litellm.llms.xai.cost_calculator import cost_per_token, cost_per_web_search_request
class TestXAICostCalculator:
@@ -320,3 +321,73 @@ 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
)
)
web_search_cost = cost_per_web_search_request(usage=usage, model_info={})
# 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,
)
# Manually set num_sources_used (as done by transformation layer)
setattr(usage, "num_sources_used", 5)
web_search_cost = cost_per_web_search_request(usage=usage, model_info={})
# 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
)
)
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,
)
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