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