mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
commit
685400ff7c
15 changed files with 931 additions and 3 deletions
308
docs/my-website/docs/completion/web_search.md
Normal file
308
docs/my-website/docs/completion/web_search.md
Normal 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>
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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`
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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),
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue