diff --git a/docs/my-website/docs/completion/web_search.md b/docs/my-website/docs/completion/web_search.md
new file mode 100644
index 00000000000..7a67dc265e4
--- /dev/null
+++ b/docs/my-website/docs/completion/web_search.md
@@ -0,0 +1,308 @@
+import Tabs from '@theme/Tabs';
+import TabItem from '@theme/TabItem';
+
+# Using Web Search
+
+Use web search with litellm
+
+| Feature | Details |
+|---------|---------|
+| Supported Endpoints | - `/chat/completions`
- `/responses` |
+| Supported Providers | `openai` |
+| LiteLLM Cost Tracking | ✅ Supported |
+| LiteLLM Version | `v1.63.15-nightly` or higher |
+
+
+## `/chat/completions` (litellm.completion)
+
+### Quick Start
+
+
+
+
+```python showLineNumbers
+from litellm import completion
+
+response = completion(
+ model="openai/gpt-4o-search-preview",
+ messages=[
+ {
+ "role": "user",
+ "content": "What was a positive news story from today?",
+ }
+ ],
+)
+```
+
+
+
+1. Setup config.yaml
+
+```yaml
+model_list:
+ - model_name: gpt-4o-search-preview
+ litellm_params:
+ model: openai/gpt-4o-search-preview
+ api_key: os.environ/OPENAI_API_KEY
+```
+
+2. Start the proxy
+
+```bash
+litellm --config /path/to/config.yaml
+```
+
+3. Test it!
+
+```python showLineNumbers
+from openai import OpenAI
+
+# Point to your proxy server
+client = OpenAI(
+ api_key="sk-1234",
+ base_url="http://0.0.0.0:4000"
+)
+
+response = client.chat.completions.create(
+ model="gpt-4o-search-preview",
+ messages=[
+ {
+ "role": "user",
+ "content": "What was a positive news story from today?"
+ }
+ ]
+)
+```
+
+
+
+### Search context size
+
+
+
+
+```python showLineNumbers
+from litellm import completion
+
+# Customize search context size
+response = completion(
+ model="openai/gpt-4o-search-preview",
+ messages=[
+ {
+ "role": "user",
+ "content": "What was a positive news story from today?",
+ }
+ ],
+ web_search_options={
+ "search_context_size": "low" # Options: "low", "medium" (default), "high"
+ }
+)
+```
+
+
+
+```python showLineNumbers
+from openai import OpenAI
+
+# Point to your proxy server
+client = OpenAI(
+ api_key="sk-1234",
+ base_url="http://0.0.0.0:4000"
+)
+
+# Customize search context size
+response = client.chat.completions.create(
+ model="gpt-4o-search-preview",
+ messages=[
+ {
+ "role": "user",
+ "content": "What was a positive news story from today?"
+ }
+ ],
+ web_search_options={
+ "search_context_size": "low" # Options: "low", "medium" (default), "high"
+ }
+)
+```
+
+
+
+## `/responses` (litellm.responses)
+
+### Quick Start
+
+
+
+
+```python showLineNumbers
+from litellm import responses
+
+response = responses(
+ model="openai/gpt-4o",
+ input=[
+ {
+ "role": "user",
+ "content": "What was a positive news story from today?"
+ }
+ ],
+ tools=[{
+ "type": "web_search_preview" # enables web search with default medium context size
+ }]
+)
+```
+
+
+
+1. Setup config.yaml
+
+```yaml
+model_list:
+ - model_name: gpt-4o
+ litellm_params:
+ model: openai/gpt-4o
+ api_key: os.environ/OPENAI_API_KEY
+```
+
+2. Start the proxy
+
+```bash
+litellm --config /path/to/config.yaml
+```
+
+3. Test it!
+
+```python showLineNumbers
+from openai import OpenAI
+
+# Point to your proxy server
+client = OpenAI(
+ api_key="sk-1234",
+ base_url="http://0.0.0.0:4000"
+)
+
+response = client.responses.create(
+ model="gpt-4o",
+ tools=[{
+ "type": "web_search_preview"
+ }],
+ input="What was a positive news story from today?",
+)
+
+print(response.output_text)
+```
+
+
+
+### Search context size
+
+
+
+
+```python showLineNumbers
+from litellm import responses
+
+# Customize search context size
+response = responses(
+ model="openai/gpt-4o",
+ input=[
+ {
+ "role": "user",
+ "content": "What was a positive news story from today?"
+ }
+ ],
+ tools=[{
+ "type": "web_search_preview",
+ "search_context_size": "low" # Options: "low", "medium" (default), "high"
+ }]
+)
+```
+
+
+
+```python showLineNumbers
+from openai import OpenAI
+
+# Point to your proxy server
+client = OpenAI(
+ api_key="sk-1234",
+ base_url="http://0.0.0.0:4000"
+)
+
+# Customize search context size
+response = client.responses.create(
+ model="gpt-4o",
+ tools=[{
+ "type": "web_search_preview",
+ "search_context_size": "low" # Options: "low", "medium" (default), "high"
+ }],
+ input="What was a positive news story from today?",
+)
+
+print(response.output_text)
+```
+
+
+
+
+
+
+
+
+## Checking if a model supports web search
+
+
+
+
+Use `litellm.supports_web_search(model="openai/gpt-4o-search-preview")` -> returns `True` if model can perform web searches
+
+```python showLineNumbers
+assert litellm.supports_web_search(model="openai/gpt-4o-search-preview") == True
+```
+
+
+
+
+1. Define OpenAI models in config.yaml
+
+```yaml
+model_list:
+ - model_name: gpt-4o-search-preview
+ litellm_params:
+ model: openai/gpt-4o-search-preview
+ api_key: os.environ/OPENAI_API_KEY
+ model_info:
+ supports_web_search: True
+```
+
+2. Run proxy server
+
+```bash
+litellm --config config.yaml
+```
+
+3. Call `/model_group/info` to check if a model supports web search
+
+```shell
+curl -X 'GET' \
+ 'http://localhost:4000/model_group/info' \
+ -H 'accept: application/json' \
+ -H 'x-api-key: sk-1234'
+```
+
+Expected Response
+
+```json showLineNumbers
+{
+ "data": [
+ {
+ "model_group": "gpt-4o-search-preview",
+ "providers": ["openai"],
+ "max_tokens": 128000,
+ "supports_web_search": true, # 👈 supports_web_search is true
+ }
+ ]
+}
+```
+
+
+
diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js
index 96267d09dab..b0f6db7c444 100644
--- a/docs/my-website/sidebars.js
+++ b/docs/my-website/sidebars.js
@@ -244,6 +244,7 @@ const sidebars = {
"completion/provider_specific_params",
"guides/finetuned_models",
"completion/audio",
+ "completion/web_search",
"completion/document_understanding",
"completion/vision",
"completion/json_mode",
diff --git a/litellm/__init__.py b/litellm/__init__.py
index 25da6504405..4f0b0a16be7 100644
--- a/litellm/__init__.py
+++ b/litellm/__init__.py
@@ -756,6 +756,7 @@ from .utils import (
create_pretrained_tokenizer,
create_tokenizer,
supports_function_calling,
+ supports_web_search,
supports_response_schema,
supports_parallel_function_calling,
supports_vision,
diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py
index e17a94c87ec..55736772af1 100644
--- a/litellm/cost_calculator.py
+++ b/litellm/cost_calculator.py
@@ -9,6 +9,9 @@ from pydantic import BaseModel
import litellm
import litellm._logging
from litellm import verbose_logger
+from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import (
+ StandardBuiltInToolCostTracking,
+)
from litellm.litellm_core_utils.llm_cost_calc.utils import _generic_cost_per_character
from litellm.llms.anthropic.cost_calculation import (
cost_per_token as anthropic_cost_per_token,
@@ -57,6 +60,7 @@ from litellm.types.utils import (
LlmProvidersSet,
ModelInfo,
PassthroughCallTypes,
+ StandardBuiltInToolsParams,
Usage,
)
from litellm.utils import (
@@ -524,6 +528,7 @@ def completion_cost( # noqa: PLR0915
optional_params: Optional[dict] = None,
custom_pricing: Optional[bool] = None,
base_model: Optional[str] = None,
+ standard_built_in_tools_params: Optional[StandardBuiltInToolsParams] = None,
) -> float:
"""
Calculate the cost of a given completion call fot GPT-3.5-turbo, llama2, any litellm supported llm.
@@ -802,6 +807,12 @@ def completion_cost( # noqa: PLR0915
rerank_billed_units=rerank_billed_units,
)
_final_cost = prompt_tokens_cost_usd_dollar + completion_tokens_cost_usd_dollar
+ _final_cost += StandardBuiltInToolCostTracking.get_cost_for_built_in_tools(
+ model=model,
+ response_object=completion_response,
+ standard_built_in_tools_params=standard_built_in_tools_params,
+ custom_llm_provider=custom_llm_provider,
+ )
return _final_cost
except Exception as e:
@@ -861,6 +872,7 @@ def response_cost_calculator(
base_model: Optional[str] = None,
custom_pricing: Optional[bool] = None,
prompt: str = "",
+ standard_built_in_tools_params: Optional[StandardBuiltInToolsParams] = None,
) -> float:
"""
Returns
@@ -890,6 +902,7 @@ def response_cost_calculator(
custom_pricing=custom_pricing,
base_model=base_model,
prompt=prompt,
+ standard_built_in_tools_params=standard_built_in_tools_params,
)
return response_cost
except Exception as e:
diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py
index 3e694220a54..67511968e21 100644
--- a/litellm/litellm_core_utils/litellm_logging.py
+++ b/litellm/litellm_core_utils/litellm_logging.py
@@ -35,6 +35,9 @@ from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.mlflow import MlflowLogger
from litellm.integrations.pagerduty.pagerduty import PagerDutyAlerting
from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
+from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import (
+ StandardBuiltInToolCostTracking,
+)
from litellm.litellm_core_utils.model_param_helper import ModelParamHelper
from litellm.litellm_core_utils.redact_messages import (
redact_message_input_output_from_custom_logger,
@@ -60,6 +63,7 @@ from litellm.types.utils import (
ModelResponse,
ModelResponseStream,
RawRequestTypedDict,
+ StandardBuiltInToolsParams,
StandardCallbackDynamicParams,
StandardLoggingAdditionalHeaders,
StandardLoggingHiddenParams,
@@ -264,7 +268,9 @@ class Logging(LiteLLMLoggingBaseClass):
self.standard_callback_dynamic_params: StandardCallbackDynamicParams = (
self.initialize_standard_callback_dynamic_params(kwargs)
)
-
+ self.standard_built_in_tools_params: StandardBuiltInToolsParams = (
+ self.initialize_standard_built_in_tools_params(kwargs)
+ )
## TIME TO FIRST TOKEN LOGGING ##
self.completion_start_time: Optional[datetime.datetime] = None
self._llm_caching_handler: Optional[LLMCachingHandler] = None
@@ -369,6 +375,23 @@ class Logging(LiteLLMLoggingBaseClass):
"""
return _initialize_standard_callback_dynamic_params(kwargs)
+ def initialize_standard_built_in_tools_params(
+ self, kwargs: Optional[Dict] = None
+ ) -> StandardBuiltInToolsParams:
+ """
+ Initialize the standard built-in tools params from the kwargs
+
+ checks if web_search_options in kwargs or tools and sets the corresponding attribute in StandardBuiltInToolsParams
+ """
+ return StandardBuiltInToolsParams(
+ web_search_options=StandardBuiltInToolCostTracking._get_web_search_options(
+ kwargs or {}
+ ),
+ file_search=StandardBuiltInToolCostTracking._get_file_search_tool_call(
+ kwargs or {}
+ ),
+ )
+
def update_environment_variables(
self,
litellm_params: Dict,
@@ -903,6 +926,7 @@ class Logging(LiteLLMLoggingBaseClass):
"optional_params": self.optional_params,
"custom_pricing": custom_pricing,
"prompt": prompt,
+ "standard_built_in_tools_params": self.standard_built_in_tools_params,
}
except Exception as e: # error creating kwargs for cost calculation
debug_info = StandardLoggingModelCostFailureDebugInformation(
@@ -1067,6 +1091,7 @@ class Logging(LiteLLMLoggingBaseClass):
end_time=end_time,
logging_obj=self,
status="success",
+ standard_built_in_tools_params=self.standard_built_in_tools_params,
)
)
elif isinstance(result, dict): # pass-through endpoints
@@ -1079,6 +1104,7 @@ class Logging(LiteLLMLoggingBaseClass):
end_time=end_time,
logging_obj=self,
status="success",
+ standard_built_in_tools_params=self.standard_built_in_tools_params,
)
)
elif standard_logging_object is not None:
@@ -1102,6 +1128,7 @@ class Logging(LiteLLMLoggingBaseClass):
prompt="",
completion=getattr(result, "content", ""),
total_time=float_diff,
+ standard_built_in_tools_params=self.standard_built_in_tools_params,
)
return start_time, end_time, result
@@ -1155,6 +1182,7 @@ class Logging(LiteLLMLoggingBaseClass):
end_time=end_time,
logging_obj=self,
status="success",
+ standard_built_in_tools_params=self.standard_built_in_tools_params,
)
)
callbacks = self.get_combined_callback_list(
@@ -1695,6 +1723,7 @@ class Logging(LiteLLMLoggingBaseClass):
end_time=end_time,
logging_obj=self,
status="success",
+ standard_built_in_tools_params=self.standard_built_in_tools_params,
)
)
callbacks = self.get_combined_callback_list(
@@ -1911,6 +1940,7 @@ class Logging(LiteLLMLoggingBaseClass):
status="failure",
error_str=str(exception),
original_exception=exception,
+ standard_built_in_tools_params=self.standard_built_in_tools_params,
)
)
return start_time, end_time
@@ -3367,6 +3397,7 @@ def get_standard_logging_object_payload(
status: StandardLoggingPayloadStatus,
error_str: Optional[str] = None,
original_exception: Optional[Exception] = None,
+ standard_built_in_tools_params: Optional[StandardBuiltInToolsParams] = None,
) -> Optional[StandardLoggingPayload]:
try:
kwargs = kwargs or {}
@@ -3542,6 +3573,7 @@ def get_standard_logging_object_payload(
guardrail_information=metadata.get(
"standard_logging_guardrail_information", None
),
+ standard_built_in_tools_params=standard_built_in_tools_params,
)
emit_standard_logging_payload(payload)
diff --git a/litellm/litellm_core_utils/llm_cost_calc/tool_call_cost_tracking.py b/litellm/litellm_core_utils/llm_cost_calc/tool_call_cost_tracking.py
new file mode 100644
index 00000000000..74d15e9a015
--- /dev/null
+++ b/litellm/litellm_core_utils/llm_cost_calc/tool_call_cost_tracking.py
@@ -0,0 +1,199 @@
+"""
+Helper utilities for tracking the cost of built-in tools.
+"""
+
+from typing import Any, Dict, List, Optional
+
+import litellm
+from litellm.types.llms.openai import FileSearchTool, WebSearchOptions
+from litellm.types.utils import (
+ ModelInfo,
+ ModelResponse,
+ SearchContextCostPerQuery,
+ StandardBuiltInToolsParams,
+)
+
+
+class StandardBuiltInToolCostTracking:
+ """
+ Helper class for tracking the cost of built-in tools
+
+ Example: Web Search
+ """
+
+ @staticmethod
+ def get_cost_for_built_in_tools(
+ model: str,
+ response_object: Any,
+ custom_llm_provider: Optional[str] = None,
+ standard_built_in_tools_params: Optional[StandardBuiltInToolsParams] = None,
+ ) -> float:
+ """
+ Get the cost of using built-in tools.
+
+ Supported tools:
+ - Web Search
+
+ """
+ if standard_built_in_tools_params is not None:
+ if (
+ standard_built_in_tools_params.get("web_search_options", None)
+ is not None
+ ):
+ model_info = StandardBuiltInToolCostTracking._safe_get_model_info(
+ model=model, custom_llm_provider=custom_llm_provider
+ )
+
+ return StandardBuiltInToolCostTracking.get_cost_for_web_search(
+ web_search_options=standard_built_in_tools_params.get(
+ "web_search_options", None
+ ),
+ model_info=model_info,
+ )
+
+ if standard_built_in_tools_params.get("file_search", None) is not None:
+ return StandardBuiltInToolCostTracking.get_cost_for_file_search(
+ file_search=standard_built_in_tools_params.get("file_search", None),
+ )
+
+ if isinstance(response_object, ModelResponse):
+ if StandardBuiltInToolCostTracking.chat_completion_response_includes_annotations(
+ response_object
+ ):
+ model_info = StandardBuiltInToolCostTracking._safe_get_model_info(
+ model=model, custom_llm_provider=custom_llm_provider
+ )
+ return StandardBuiltInToolCostTracking.get_default_cost_for_web_search(
+ model_info
+ )
+ return 0.0
+
+ @staticmethod
+ def _safe_get_model_info(
+ model: str, custom_llm_provider: Optional[str] = None
+ ) -> Optional[ModelInfo]:
+ try:
+ return litellm.get_model_info(
+ model=model, custom_llm_provider=custom_llm_provider
+ )
+ except Exception:
+ return None
+
+ @staticmethod
+ def get_cost_for_web_search(
+ web_search_options: Optional[WebSearchOptions] = None,
+ model_info: Optional[ModelInfo] = None,
+ ) -> float:
+ """
+ If request includes `web_search_options`, calculate the cost of the web search.
+ """
+ if web_search_options is None:
+ return 0.0
+ if model_info is None:
+ return 0.0
+
+ search_context_pricing: SearchContextCostPerQuery = (
+ model_info.get("search_context_cost_per_query", {}) or {}
+ )
+ if web_search_options.get("search_context_size", None) == "low":
+ return search_context_pricing.get("search_context_size_low", 0.0)
+ elif web_search_options.get("search_context_size", None) == "medium":
+ return search_context_pricing.get("search_context_size_medium", 0.0)
+ elif web_search_options.get("search_context_size", None) == "high":
+ return search_context_pricing.get("search_context_size_high", 0.0)
+ return StandardBuiltInToolCostTracking.get_default_cost_for_web_search(
+ model_info
+ )
+
+ @staticmethod
+ def get_default_cost_for_web_search(
+ model_info: Optional[ModelInfo] = None,
+ ) -> float:
+ """
+ If no web search options are provided, use the `search_context_size_medium` pricing.
+
+ https://platform.openai.com/docs/pricing#web-search
+ """
+ if model_info is None:
+ return 0.0
+ search_context_pricing: SearchContextCostPerQuery = (
+ model_info.get("search_context_cost_per_query", {}) or {}
+ ) or {}
+ return search_context_pricing.get("search_context_size_medium", 0.0)
+
+ @staticmethod
+ def get_cost_for_file_search(
+ file_search: Optional[FileSearchTool] = None,
+ ) -> float:
+ """ "
+ Charged at $2.50/1k calls
+
+ Doc: https://platform.openai.com/docs/pricing#built-in-tools
+ """
+ if file_search is None:
+ return 0.0
+ return 2.5 / 1000
+
+ @staticmethod
+ def chat_completion_response_includes_annotations(
+ response_object: ModelResponse,
+ ) -> bool:
+ for _choice in response_object.choices:
+ message = getattr(_choice, "message", None)
+ if (
+ message is not None
+ and hasattr(message, "annotations")
+ and message.annotations is not None
+ and len(message.annotations) > 0
+ ):
+ return True
+ return False
+
+ @staticmethod
+ def _get_web_search_options(kwargs: Dict) -> Optional[WebSearchOptions]:
+ if "web_search_options" in kwargs:
+ return WebSearchOptions(**kwargs.get("web_search_options", {}))
+
+ tools = StandardBuiltInToolCostTracking._get_tools_from_kwargs(
+ kwargs, "web_search_preview"
+ )
+ if tools:
+ # Look for web search tool in the tools array
+ for tool in tools:
+ if isinstance(tool, dict):
+ if StandardBuiltInToolCostTracking._is_web_search_tool_call(tool):
+ return WebSearchOptions(**tool)
+ return None
+
+ @staticmethod
+ def _get_tools_from_kwargs(kwargs: Dict, tool_type: str) -> Optional[List[Dict]]:
+ if "tools" in kwargs:
+ tools = kwargs.get("tools", [])
+ return tools
+ return None
+
+ @staticmethod
+ def _get_file_search_tool_call(kwargs: Dict) -> Optional[FileSearchTool]:
+ tools = StandardBuiltInToolCostTracking._get_tools_from_kwargs(
+ kwargs, "file_search"
+ )
+ if tools:
+ for tool in tools:
+ if isinstance(tool, dict):
+ if StandardBuiltInToolCostTracking._is_file_search_tool_call(tool):
+ return FileSearchTool(**tool)
+ return None
+
+ @staticmethod
+ def _is_web_search_tool_call(tool: Dict) -> bool:
+ if tool.get("type", None) == "web_search_preview":
+ return True
+ if "search_context_size" in tool:
+ return True
+ return False
+
+ @staticmethod
+ def _is_file_search_tool_call(tool: Dict) -> bool:
+ if tool.get("type", None) == "file_search":
+ return True
+ return False
diff --git a/litellm/router.py b/litellm/router.py
index af7b00e79d7..524c55539ff 100644
--- a/litellm/router.py
+++ b/litellm/router.py
@@ -4924,6 +4924,11 @@ class Router:
and model_info["supports_function_calling"] is True # type: ignore
):
model_group_info.supports_function_calling = True
+ if (
+ model_info.get("supports_web_search", None) is not None
+ and model_info["supports_web_search"] is True # type: ignore
+ ):
+ model_group_info.supports_web_search = True
if (
model_info.get("supported_openai_params", None) is not None
and model_info["supported_openai_params"] is not None
diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py
index e58f5732271..19899648f58 100644
--- a/litellm/types/llms/openai.py
+++ b/litellm/types/llms/openai.py
@@ -382,6 +382,53 @@ class ChatCompletionThinkingBlock(TypedDict, total=False):
cache_control: Optional[Union[dict, ChatCompletionCachedContent]]
+class WebSearchOptionsUserLocationApproximate(TypedDict, total=False):
+ city: str
+ """Free text input for the city of the user, e.g. `San Francisco`."""
+
+ country: str
+ """
+ The two-letter [ISO country code](https://en.wikipedia.org/wiki/ISO_3166-1) of
+ the user, e.g. `US`.
+ """
+
+ region: str
+ """Free text input for the region of the user, e.g. `California`."""
+
+ timezone: str
+ """
+ The [IANA timezone](https://timeapi.io/documentation/iana-timezones) of the
+ user, e.g. `America/Los_Angeles`.
+ """
+
+
+class WebSearchOptionsUserLocation(TypedDict, total=False):
+ approximate: Required[WebSearchOptionsUserLocationApproximate]
+ """Approximate location parameters for the search."""
+
+ type: Required[Literal["approximate"]]
+ """The type of location approximation. Always `approximate`."""
+
+
+class WebSearchOptions(TypedDict, total=False):
+ search_context_size: Literal["low", "medium", "high"]
+ """
+ High level guidance for the amount of context window space to use for the
+ search. One of `low`, `medium`, or `high`. `medium` is the default.
+ """
+
+ user_location: Optional[WebSearchOptionsUserLocation]
+ """Approximate location parameters for the search."""
+
+
+class FileSearchTool(TypedDict, total=False):
+ type: Literal["file_search"]
+ """The type of tool being defined: `file_search`"""
+
+ vector_store_ids: Optional[List[str]]
+ """The IDs of the vector stores to search."""
+
+
class ChatCompletionAnnotationURLCitation(TypedDict, total=False):
end_index: int
"""The index of the last character of the URL citation in the message."""
diff --git a/litellm/types/router.py b/litellm/types/router.py
index e34366aa229..dcd547def29 100644
--- a/litellm/types/router.py
+++ b/litellm/types/router.py
@@ -559,6 +559,7 @@ class ModelGroupInfo(BaseModel):
rpm: Optional[int] = None
supports_parallel_function_calling: bool = Field(default=False)
supports_vision: bool = Field(default=False)
+ supports_web_search: bool = Field(default=False)
supports_function_calling: bool = Field(default=False)
supported_openai_params: Optional[List[str]] = Field(default=[])
configurable_clientside_auth_params: CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS = None
diff --git a/litellm/types/utils.py b/litellm/types/utils.py
index 8821d2c80bd..2cc06eecbf3 100644
--- a/litellm/types/utils.py
+++ b/litellm/types/utils.py
@@ -32,7 +32,9 @@ from .llms.openai import (
ChatCompletionThinkingBlock,
ChatCompletionToolCallChunk,
ChatCompletionUsageBlock,
+ FileSearchTool,
OpenAIChatCompletionChunk,
+ WebSearchOptions,
)
from .rerank import RerankResponse
@@ -97,6 +99,13 @@ class ProviderSpecificModelInfo(TypedDict, total=False):
supports_pdf_input: Optional[bool]
supports_native_streaming: Optional[bool]
supports_parallel_function_calling: Optional[bool]
+ supports_web_search: Optional[bool]
+
+
+class SearchContextCostPerQuery(TypedDict, total=False):
+ search_context_size_low: float
+ search_context_size_medium: float
+ search_context_size_high: float
class ModelInfoBase(ProviderSpecificModelInfo, total=False):
@@ -135,6 +144,9 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
output_cost_per_video_per_second: Optional[float] # only for vertex ai models
output_cost_per_audio_per_second: Optional[float] # only for vertex ai models
output_cost_per_second: Optional[float] # for OpenAI Speech models
+ search_context_cost_per_query: Optional[
+ SearchContextCostPerQuery
+ ] # Cost for using web search tool
litellm_provider: Required[str]
mode: Required[
@@ -586,6 +598,11 @@ class Message(OpenAIObject):
# OpenAI compatible APIs like mistral API will raise an error if audio is passed in
del self.audio
+ if annotations is None:
+ # ensure default response matches OpenAI spec
+ # Some OpenAI compatible APIs raise an error if annotations are passed in
+ del self.annotations
+
if reasoning_content is None:
# ensure default response matches OpenAI spec
del self.reasoning_content
@@ -1612,6 +1629,19 @@ class StandardLoggingUserAPIKeyMetadata(TypedDict):
user_api_key_end_user_id: Optional[str]
+class StandardBuiltInToolsParams(TypedDict, total=False):
+ """
+ Standard built-in OpenAItools parameters
+
+ This is used to calculate the cost of built-in tools, insert any standard built-in tools parameters here
+
+ OpenAI charges users based on the `web_search_options` parameter
+ """
+
+ web_search_options: Optional[WebSearchOptions]
+ file_search: Optional[FileSearchTool]
+
+
class StandardLoggingPromptManagementMetadata(TypedDict):
prompt_id: str
prompt_variables: Optional[dict]
@@ -1729,6 +1759,7 @@ class StandardLoggingPayload(TypedDict):
model_parameters: dict
hidden_params: StandardLoggingHiddenParams
guardrail_information: Optional[StandardLoggingGuardrailInformation]
+ standard_built_in_tools_params: Optional[StandardBuiltInToolsParams]
from typing import AsyncIterator, Iterator
diff --git a/litellm/utils.py b/litellm/utils.py
index 03e69acf4e3..dc97c4d898f 100644
--- a/litellm/utils.py
+++ b/litellm/utils.py
@@ -1975,7 +1975,7 @@ def supports_system_messages(model: str, custom_llm_provider: Optional[str]) ->
)
-def supports_web_search(model: str, custom_llm_provider: Optional[str]) -> bool:
+def supports_web_search(model: str, custom_llm_provider: Optional[str] = None) -> bool:
"""
Check if the given model supports web search and return a boolean value.
@@ -4544,6 +4544,10 @@ def _get_model_info_helper( # noqa: PLR0915
supports_native_streaming=_model_info.get(
"supports_native_streaming", None
),
+ supports_web_search=_model_info.get("supports_web_search", False),
+ search_context_cost_per_query=_model_info.get(
+ "search_context_cost_per_query", None
+ ),
tpm=_model_info.get("tpm", None),
rpm=_model_info.get("rpm", None),
)
@@ -4612,6 +4616,7 @@ def get_model_info(model: str, custom_llm_provider: Optional[str] = None) -> Mod
supports_audio_input: Optional[bool]
supports_audio_output: Optional[bool]
supports_pdf_input: Optional[bool]
+ supports_web_search: Optional[bool]
Raises:
Exception: If the model is not mapped yet.
diff --git a/tests/litellm/litellm_core_utils/llm_cost_calc/test_tool_call_cost_tracking.py b/tests/litellm/litellm_core_utils/llm_cost_calc/test_tool_call_cost_tracking.py
new file mode 100644
index 00000000000..0d6bcb7cce0
--- /dev/null
+++ b/tests/litellm/litellm_core_utils/llm_cost_calc/test_tool_call_cost_tracking.py
@@ -0,0 +1,113 @@
+import json
+import os
+import sys
+
+import pytest
+from fastapi.testclient import TestClient
+
+import litellm
+from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import (
+ StandardBuiltInToolCostTracking,
+)
+from litellm.types.llms.openai import FileSearchTool, WebSearchOptions
+from litellm.types.utils import ModelInfo, ModelResponse, StandardBuiltInToolsParams
+
+sys.path.insert(
+ 0, os.path.abspath("../../..")
+) # Adds the parent directory to the system path
+
+
+# Test basic web search cost calculations
+def test_web_search_cost_low():
+ web_search_options = WebSearchOptions(search_context_size="low")
+ model_info = litellm.get_model_info("gpt-4o-search-preview")
+
+ cost = StandardBuiltInToolCostTracking.get_cost_for_web_search(
+ web_search_options=web_search_options, model_info=model_info
+ )
+
+ assert (
+ cost == model_info["search_context_cost_per_query"]["search_context_size_low"]
+ )
+
+
+def test_web_search_cost_medium():
+ web_search_options = WebSearchOptions(search_context_size="medium")
+ model_info = litellm.get_model_info("gpt-4o-search-preview")
+
+ cost = StandardBuiltInToolCostTracking.get_cost_for_web_search(
+ web_search_options=web_search_options, model_info=model_info
+ )
+
+ assert (
+ cost
+ == model_info["search_context_cost_per_query"]["search_context_size_medium"]
+ )
+
+
+def test_web_search_cost_high():
+ web_search_options = WebSearchOptions(search_context_size="high")
+ model_info = litellm.get_model_info("gpt-4o-search-preview")
+
+ cost = StandardBuiltInToolCostTracking.get_cost_for_web_search(
+ web_search_options=web_search_options, model_info=model_info
+ )
+
+ assert (
+ cost == model_info["search_context_cost_per_query"]["search_context_size_high"]
+ )
+
+
+# Test file search cost calculation
+def test_file_search_cost():
+ file_search = FileSearchTool(type="file_search")
+ cost = StandardBuiltInToolCostTracking.get_cost_for_file_search(
+ file_search=file_search
+ )
+ assert cost == 0.0025 # $2.50/1000 calls = 0.0025 per call
+
+
+# Test edge cases
+def test_none_inputs():
+ # Test with None inputs
+ assert (
+ StandardBuiltInToolCostTracking.get_cost_for_web_search(
+ web_search_options=None, model_info=None
+ )
+ == 0.0
+ )
+ assert (
+ StandardBuiltInToolCostTracking.get_cost_for_file_search(file_search=None)
+ == 0.0
+ )
+
+
+# Test the main get_cost_for_built_in_tools method
+def test_get_cost_for_built_in_tools_web_search():
+ model = "gpt-4"
+ standard_built_in_tools_params = StandardBuiltInToolsParams(
+ web_search_options=WebSearchOptions(search_context_size="medium")
+ )
+
+ cost = StandardBuiltInToolCostTracking.get_cost_for_built_in_tools(
+ model=model,
+ response_object=None,
+ standard_built_in_tools_params=standard_built_in_tools_params,
+ )
+
+ assert isinstance(cost, float)
+
+
+def test_get_cost_for_built_in_tools_file_search():
+ model = "gpt-4"
+ standard_built_in_tools_params = StandardBuiltInToolsParams(
+ file_search=FileSearchTool(type="file_search")
+ )
+
+ cost = StandardBuiltInToolCostTracking.get_cost_for_built_in_tools(
+ model=model,
+ response_object=None,
+ standard_built_in_tools_params=standard_built_in_tools_params,
+ )
+
+ assert cost == 0.0025
diff --git a/tests/litellm_utils_tests/test_utils.py b/tests/litellm_utils_tests/test_utils.py
index 42df7c495c2..fea225e4a3b 100644
--- a/tests/litellm_utils_tests/test_utils.py
+++ b/tests/litellm_utils_tests/test_utils.py
@@ -477,6 +477,25 @@ def test_supports_function_calling(model, expected_bool):
pytest.fail(f"Error occurred: {e}")
+@pytest.mark.parametrize(
+ "model, expected_bool",
+ [
+ ("gpt-4o-mini-search-preview", True),
+ ("openai/gpt-4o-mini-search-preview", True),
+ ("gpt-4o-search-preview", True),
+ ("openai/gpt-4o-search-preview", True),
+ ("groq/deepseek-r1-distill-llama-70b", False),
+ ("groq/llama-3.3-70b-versatile", False),
+ ("codestral/codestral-latest", False),
+ ],
+)
+def test_supports_web_search(model, expected_bool):
+ try:
+ assert litellm.supports_web_search(model=model) == expected_bool
+ except Exception as e:
+ pytest.fail(f"Error occurred: {e}")
+
+
def test_get_max_token_unit_test():
"""
More complete testing in `test_completion_cost.py`
diff --git a/tests/logging_callback_tests/test_built_in_tools_cost_tracking.py b/tests/logging_callback_tests/test_built_in_tools_cost_tracking.py
new file mode 100644
index 00000000000..d4013f2e8db
--- /dev/null
+++ b/tests/logging_callback_tests/test_built_in_tools_cost_tracking.py
@@ -0,0 +1,151 @@
+import os
+import sys
+import traceback
+import uuid
+import pytest
+from dotenv import load_dotenv
+from fastapi import Request
+from fastapi.routing import APIRoute
+
+load_dotenv()
+import io
+import os
+import time
+import json
+
+# this file is to test litellm/proxy
+
+sys.path.insert(
+ 0, os.path.abspath("../..")
+) # Adds the parent directory to the system path
+import litellm
+import asyncio
+from typing import Optional
+from litellm.types.utils import StandardLoggingPayload, Usage, ModelInfoBase
+from litellm.integrations.custom_logger import CustomLogger
+
+
+class TestCustomLogger(CustomLogger):
+ def __init__(self):
+ self.recorded_usage: Optional[Usage] = None
+ self.standard_logging_payload: Optional[StandardLoggingPayload] = None
+
+ async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
+ standard_logging_payload = kwargs.get("standard_logging_object")
+ self.standard_logging_payload = standard_logging_payload
+ print(
+ "standard_logging_payload",
+ json.dumps(standard_logging_payload, indent=4, default=str),
+ )
+
+ self.recorded_usage = Usage(
+ prompt_tokens=standard_logging_payload.get("prompt_tokens"),
+ completion_tokens=standard_logging_payload.get("completion_tokens"),
+ total_tokens=standard_logging_payload.get("total_tokens"),
+ )
+ pass
+
+
+async def _setup_web_search_test():
+ """Helper function to setup common test requirements"""
+ litellm._turn_on_debug()
+ test_custom_logger = TestCustomLogger()
+ litellm.callbacks = [test_custom_logger]
+ return test_custom_logger
+
+
+async def _verify_web_search_cost(test_custom_logger, expected_context_size):
+ """Helper function to verify web search costs"""
+ await asyncio.sleep(1)
+
+ standard_logging_payload = test_custom_logger.standard_logging_payload
+ response_cost = standard_logging_payload.get("response_cost")
+ assert response_cost is not None
+
+ # Calculate token cost
+ model_map_information = standard_logging_payload["model_map_information"]
+ model_map_value: ModelInfoBase = model_map_information["model_map_value"]
+ total_token_cost = (
+ standard_logging_payload["prompt_tokens"]
+ * model_map_value["input_cost_per_token"]
+ ) + (
+ standard_logging_payload["completion_tokens"]
+ * model_map_value["output_cost_per_token"]
+ )
+
+ # Verify total cost
+ assert (
+ response_cost
+ == total_token_cost
+ + model_map_value["search_context_cost_per_query"][expected_context_size]
+ )
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize(
+ "web_search_options,expected_context_size",
+ [
+ (None, "search_context_size_medium"),
+ ({"search_context_size": "low"}, "search_context_size_low"),
+ ({"search_context_size": "high"}, "search_context_size_high"),
+ ],
+)
+async def test_openai_web_search_logging_cost_tracking(
+ web_search_options, expected_context_size
+):
+ """Test web search cost tracking with different search context sizes"""
+ test_custom_logger = await _setup_web_search_test()
+
+ request_kwargs = {
+ "model": "openai/gpt-4o-search-preview",
+ "messages": [
+ {"role": "user", "content": "What was a positive news story from today?"}
+ ],
+ }
+ if web_search_options is not None:
+ request_kwargs["web_search_options"] = web_search_options
+
+ response = await litellm.acompletion(**request_kwargs)
+
+ await _verify_web_search_cost(test_custom_logger, expected_context_size)
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize(
+ "tools_config,expected_context_size,stream",
+ [
+ (
+ [{"type": "web_search_preview", "search_context_size": "high"}],
+ "search_context_size_high",
+ True,
+ ),
+ (
+ [{"type": "web_search_preview", "search_context_size": "high"}],
+ "search_context_size_high",
+ False,
+ ),
+ ([{"type": "web_search_preview"}], "search_context_size_medium", True),
+ ([{"type": "web_search_preview"}], "search_context_size_medium", False),
+ ],
+)
+async def test_openai_responses_api_web_search_cost_tracking(
+ tools_config, expected_context_size, stream
+):
+ """Test web search cost tracking with different search context sizes and streaming options"""
+ test_custom_logger = await _setup_web_search_test()
+
+ response = await litellm.aresponses(
+ model="openai/gpt-4o",
+ input=[
+ {"role": "user", "content": "What was a positive news story from today?"}
+ ],
+ tools=tools_config,
+ stream=stream,
+ )
+ if stream is True:
+ async for chunk in response:
+ print("chunk", chunk)
+ else:
+ print("response", response)
+
+ await _verify_web_search_cost(test_custom_logger, expected_context_size)
diff --git a/tests/logging_callback_tests/test_token_counting.py b/tests/logging_callback_tests/test_token_counting.py
index 341ef2a5451..ed778fa86b2 100644
--- a/tests/logging_callback_tests/test_token_counting.py
+++ b/tests/logging_callback_tests/test_token_counting.py
@@ -21,16 +21,18 @@ sys.path.insert(
import litellm
import asyncio
from typing import Optional
-from litellm.types.utils import StandardLoggingPayload, Usage
+from litellm.types.utils import StandardLoggingPayload, Usage, ModelInfoBase
from litellm.integrations.custom_logger import CustomLogger
class TestCustomLogger(CustomLogger):
def __init__(self):
self.recorded_usage: Optional[Usage] = None
+ self.standard_logging_payload: Optional[StandardLoggingPayload] = None
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
standard_logging_payload = kwargs.get("standard_logging_object")
+ self.standard_logging_payload = standard_logging_payload
print(
"standard_logging_payload",
json.dumps(standard_logging_payload, indent=4, default=str),