mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix StandardBuiltInToolsParams
This commit is contained in:
parent
0f36b8afc2
commit
8766b43a4b
4 changed files with 63 additions and 15 deletions
|
|
@ -386,7 +386,10 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
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(
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ Helper utilities for tracking the cost of built-in tools.
|
|||
from typing import Any, Dict, Optional
|
||||
|
||||
import litellm
|
||||
from litellm.types.llms.openai import WebSearchOptions
|
||||
from litellm.types.llms.openai import FileSearchTool, WebSearchOptions
|
||||
from litellm.types.utils import (
|
||||
ModelInfo,
|
||||
ModelResponse,
|
||||
|
|
@ -39,17 +39,22 @@ class StandardBuiltInToolCostTracking:
|
|||
model=model, custom_llm_provider=custom_llm_provider
|
||||
)
|
||||
|
||||
if (
|
||||
standard_built_in_tools_params is not None
|
||||
and standard_built_in_tools_params.get("web_search_options", None)
|
||||
is not None
|
||||
):
|
||||
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 is not None:
|
||||
if (
|
||||
standard_built_in_tools_params.get("web_search_options", None)
|
||||
is not None
|
||||
):
|
||||
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(
|
||||
|
|
@ -98,6 +103,19 @@ class StandardBuiltInToolCostTracking:
|
|||
) 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,
|
||||
|
|
@ -122,14 +140,31 @@ class StandardBuiltInToolCostTracking:
|
|||
# Look for web search tool in the tools array
|
||||
for tool in tools:
|
||||
if isinstance(tool, dict):
|
||||
if StandardBuiltInToolCostTracking._web_search_tool_call(tool):
|
||||
if StandardBuiltInToolCostTracking._is_web_search_tool_call(tool):
|
||||
return WebSearchOptions(**tool)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _web_search_tool_call(tool: Dict) -> bool:
|
||||
def _get_file_search_tool_call(kwargs: Dict) -> Optional[FileSearchTool]:
|
||||
if "tools" in kwargs:
|
||||
tools = kwargs.get("tools", [])
|
||||
# Look for web search tool in the tools array
|
||||
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
|
||||
|
|
|
|||
|
|
@ -421,6 +421,14 @@ class WebSearchOptions(TypedDict, total=False):
|
|||
"""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."""
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ from openai.types.completion_usage import (
|
|||
CompletionUsage,
|
||||
PromptTokensDetails,
|
||||
)
|
||||
from openai.types.llms.openai import FileSearchTool
|
||||
from openai.types.moderation import (
|
||||
Categories,
|
||||
CategoryAppliedInputTypes,
|
||||
|
|
@ -1633,6 +1634,7 @@ class StandardBuiltInToolsParams(TypedDict, total=False):
|
|||
"""
|
||||
|
||||
web_search_options: Optional[WebSearchOptions]
|
||||
file_search: Optional[FileSearchTool]
|
||||
|
||||
|
||||
class StandardLoggingPromptManagementMetadata(TypedDict):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue