From 2828720c437a537705db26510e36ca959c40d8d2 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 2 May 2025 07:06:46 -0700 Subject: [PATCH] [Bug Fix] Ensure Web Search / File Search cost are only added when the response includes the too call (#10476) * only apply cost if response includes annotations for url/file * only apply cost if response includes annotations for url/file * testing fix tool cost tracking * fix response_includes_annotation_type --- .../llm_cost_calc/tool_call_cost_tracking.py | 153 ++++++++++++++---- 1 file changed, 120 insertions(+), 33 deletions(-) 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 index 34c370ffca7..53d658c5c34 100644 --- 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 @@ -2,12 +2,17 @@ Helper utilities for tracking the cost of built-in tools. """ -from typing import Any, Dict, List, Optional +from typing import Any, Dict, List, Literal, Optional import litellm from litellm.constants import OPENAI_FILE_SEARCH_COST_PER_1K_CALLS -from litellm.types.llms.openai import FileSearchTool, WebSearchOptions +from litellm.types.llms.openai import ( + FileSearchTool, + ResponsesAPIResponse, + WebSearchOptions, +) from litellm.types.utils import ( + Message, ModelInfo, ModelResponse, SearchContextCostPerQuery, @@ -36,39 +41,122 @@ class StandardBuiltInToolCostTracking: - 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 - ) + standard_built_in_tools_params = standard_built_in_tools_params or {} + ######################################################### + # Web Search + ######################################################### + if StandardBuiltInToolCostTracking.response_object_includes_web_search_call( + response_object=response_object + ): + 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, + ) - return StandardBuiltInToolCostTracking.get_cost_for_web_search( - web_search_options=standard_built_in_tools_params.get( - "web_search_options", None - ), - model_info=model_info, - ) + ######################################################### + # File Search + ######################################################### + elif StandardBuiltInToolCostTracking.response_object_includes_file_search_call( + response_object=response_object + ): + return StandardBuiltInToolCostTracking.get_cost_for_file_search( + file_search=standard_built_in_tools_params.get("file_search", None), + ) - 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 response_object_includes_web_search_call( + response_object: Any, + ) -> bool: + """ + Check if the response object includes a web search call. + + This covers: + - Chat Completion Response (ModelResponse) + - ResponsesAPIResponse (streaming + non-streaming) + """ + if isinstance(response_object, ModelResponse): + # chat completions only include url_citation annotations when a web search call is made + return StandardBuiltInToolCostTracking.response_includes_annotation_type( + response_object=response_object, annotation_type="url_citation" + ) + elif isinstance(response_object, ResponsesAPIResponse): + # response api explicitly includes web_search_call in the output + return StandardBuiltInToolCostTracking.response_includes_output_type( + response_object=response_object, output_type="web_search_call" + ) + return False + + @staticmethod + def response_object_includes_file_search_call( + response_object: Any, + ) -> bool: + """ + Check if the response object includes a file search call. + + This covers: + - Chat Completion Response (ModelResponse) + - ResponsesAPIResponse (streaming + non-streaming) + """ + if isinstance(response_object, ModelResponse): + # chat completions only include file_citation annotations when a file search call is made + return StandardBuiltInToolCostTracking.response_includes_annotation_type( + response_object=response_object, annotation_type="file_citation" + ) + elif isinstance(response_object, ResponsesAPIResponse): + # response api explicitly includes file_search_call in the output + return StandardBuiltInToolCostTracking.response_includes_output_type( + response_object=response_object, output_type="file_search_call" + ) + return False + + @staticmethod + def response_includes_annotation_type( + response_object: ModelResponse, + annotation_type: Literal["url_citation", "file_citation"], + ) -> bool: + if isinstance(response_object, ModelResponse): + for choice in response_object.choices: + message: Optional[Message] = getattr(choice, "message", None) + if message is None: + continue + if annotations := getattr(message, "annotations", None): + if len(annotations) > 0: + for annotation in annotations: + if annotation.get("type", None) == annotation_type: + return True + return False + + @staticmethod + def response_includes_output_type( + response_object: ResponsesAPIResponse, + output_type: Literal["web_search_call", "file_search_call"], + ) -> bool: + """ + Check if the ResponsesAPIResponse includes one of the specified output types. + + This is used for cost tracking of built-in tools. + + Args: + response_object: The ResponsesAPIResponse object to check. + output_type: The type of output to check for. + + Returns: + True if the ResponsesAPIResponse includes one of the specified output types, False otherwise. + """ + output = response_object.output + for output_item in output: + _output_type: Optional[str] = getattr(output_item, "type", None) + if _output_type == output_type: + return True + return False + @staticmethod def _safe_get_model_info( model: str, custom_llm_provider: Optional[str] = None @@ -88,8 +176,7 @@ class StandardBuiltInToolCostTracking: """ If request includes `web_search_options`, calculate the cost of the web search. """ - if web_search_options is None: - return 0.0 + web_search_options = web_search_options or {} if model_info is None: return 0.0