Merge pull request #9469 from BerriAI/litellm_web_search_2

[Feat] Add testing for `litellm.supports_web_search()`  and render supports_web_search on model hub
This commit is contained in:
Ishaan Jaff 2025-03-22 19:47:25 -07:00 • committed by GitHub
commit 685400ff7c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
15 changed files with 931 additions and 3 deletions

View file

@ -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` <br/> - `/responses` |
| Supported Providers | `openai` |
| LiteLLM Cost Tracking | ✅ Supported |
| LiteLLM Version | `v1.63.15-nightly` or higher |
## `/chat/completions` (litellm.completion)
### Quick Start
<Tabs>
<TabItem value="sdk" label="SDK">
```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?",
}
],
)
```
</TabItem>
<TabItem value="proxy" label="PROXY">
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?"
}
]
)
```
</TabItem>
</Tabs>
### Search context size
<Tabs>
<TabItem value="sdk" label="SDK">
```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"
}
)
```
</TabItem>
<TabItem value="proxy" label="PROXY">
```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"
}
)
```
</TabItem>
</Tabs>
## `/responses` (litellm.responses)
### Quick Start
<Tabs>
<TabItem value="sdk" label="SDK">
```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
}]
)
```
</TabItem>
<TabItem value="proxy" label="PROXY">
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)
```
</TabItem>
</Tabs>
### Search context size
<Tabs>
<TabItem value="sdk" label="SDK">
```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"
}]
)
```
</TabItem>
<TabItem value="proxy" label="PROXY">
```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)
```
</TabItem>
</Tabs>
## Checking if a model supports web search
<Tabs>
<TabItem label="SDK" value="sdk">
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
```
</TabItem>
<TabItem label="PROXY" value="proxy">
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
}
]
}
```
</TabItem>
</Tabs>

View file

@ -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",

View file

@ -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,

View file

@ -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:

View file

@ -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)

View file

@ -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

View file

@ -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

View file

@ -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."""

View file

@ -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

View file

@ -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

View file

@ -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.

View file

@ -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

View file

@ -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`

View file

@ -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)

View file

@ -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),