Anthropic unified web search + tool cost tracking support (#10846)

* fix(duration_parser.py): support `mo` unit

* test(test_key_management_endpoints.py): add test confirming generate_key_helper_fn uses predictable budgets

Closes https://github.com/BerriAI/litellm/issues/10800

* fix(anthropic/chat/transformation.py): add tool use cost tracking

* fix(anthropic/): refactor how hosted tool usage tracking is done

keep it separate from prompt / completion token details

* fix(anthropic/): add web search tool cost tracking

accurate cost tracking

* feat(anthropic/chat/transformation.py): map openai 'web_search_options' param to anthropic hosted tool

Allows calling anthropic web search in same format as openai

* feat(anthropic/chat/transformation.py): support unified anthropic 'web_search_options' param

Allows calling anthropic's web search tool in the openai format

* feat(anthropic/chat/transformation.py): map openai 'search_context_size' to anthropic 'max_uses' param

Translate search effort across both providers

* fix: mark web_search_options param as supported by openai + azure

* fix: fix linting error

* fix: fix linting errors

* fix: fix linting error

* fix: check if usage hasattr

* fix: pass web search options param
This commit is contained in:
Krish Dholakia 2025-05-14 22:41:12 -07:00 • committed by GitHub
parent 11740ce144
commit 5146b2903f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
19 changed files with 451 additions and 26 deletions

View file

@ -153,6 +153,13 @@ FIREWORKS_AI_80_B = int(os.getenv("FIREWORKS_AI_80_B", 80))
#### Logging callback constants ####
REDACTED_BY_LITELM_STRING = "REDACTED_BY_LITELM"
### ANTHROPIC CONSTANTS ###
ANTHROPIC_WEB_SEARCH_TOOL_MAX_USES = {
"low": 1,
"medium": 5,
"high": 10,
}
LITELLM_CHAT_PROVIDERS = [
"openai",
"openai_like",
@ -259,6 +266,7 @@ OPENAI_CHAT_COMPLETION_PARAMS = [
"reasoning_effort",
"extra_headers",
"thinking",
"web_search_options",
]
openai_compatible_endpoints: List = [

View file

@ -909,6 +909,7 @@ def completion_cost( # noqa: PLR0915
StandardBuiltInToolCostTracking.get_cost_for_built_in_tools(
model=model,
response_object=completion_response,
usage=cost_per_token_usage_object,
standard_built_in_tools_params=standard_built_in_tools_params,
custom_llm_provider=custom_llm_provider,
)

View file

@ -138,6 +138,8 @@ def get_next_standardized_reset_time(
return _handle_minute_reset(current_time, base_midnight, value)
elif unit == "s":
return _handle_second_reset(current_time, base_midnight, value)
elif unit == "mo":
return _handle_month_reset(current_time, base_midnight, value)
else:
# Unrecognized unit, default to next midnight
return base_midnight + timedelta(days=1)
@ -343,3 +345,40 @@ def _handle_second_reset(
return current_time.replace(
hour=next_hour, minute=next_minute, second=next_second, microsecond=0
)
def _handle_month_reset(
current_time: datetime, base_midnight: datetime, value: int
) -> datetime:
"""
Handle monthly reset times. For monthly resets, we always reset at the start of the next month.
Args:
current_time: Current datetime
base_midnight: Midnight of current day
value: Number of months (currently only supports 1 month resets)
Returns:
datetime: First day of next month at midnight
"""
if value != 1:
raise ValueError("Monthly resets currently only support 1 month intervals")
# Get the first day of next month
if current_time.month == 12:
next_month = 1
next_year = current_time.year + 1
else:
next_month = current_time.month + 1
next_year = current_time.year
return datetime(
year=next_year,
month=next_month,
day=1,
hour=0,
minute=0,
second=0,
microsecond=0,
tzinfo=current_time.tzinfo,
)

View file

@ -17,6 +17,7 @@ from litellm.types.utils import (
ModelResponse,
SearchContextCostPerQuery,
StandardBuiltInToolsParams,
Usage,
)
@ -27,10 +28,46 @@ class StandardBuiltInToolCostTracking:
Example: Web Search
"""
@staticmethod
def get_cost_for_anthropic_web_search(
model_info: Optional[ModelInfo] = None,
usage: Optional[Usage] = None,
) -> float:
"""
Get the cost of using a web search tool for Anthropic.
"""
## Check if web search requests are in the usage object
if model_info is None:
return 0.0
if (
usage is None
or usage.server_tool_use is None
or usage.server_tool_use.web_search_requests is None
):
return 0.0
## Get the cost per web search request
search_context_pricing: SearchContextCostPerQuery = (
model_info.get("search_context_cost_per_query", {}) or {}
)
cost_per_web_search_request = search_context_pricing.get(
"search_context_size_medium", 0.0
)
if cost_per_web_search_request is None or cost_per_web_search_request == 0.0:
return 0.0
## Calculate the total cost
total_cost = (
cost_per_web_search_request * usage.server_tool_use.web_search_requests
)
return total_cost
@staticmethod
def get_cost_for_built_in_tools(
model: str,
response_object: Any,
usage: Optional[Usage] = None,
custom_llm_provider: Optional[str] = None,
standard_built_in_tools_params: Optional[StandardBuiltInToolsParams] = None,
) -> float:
@ -46,17 +83,26 @@ class StandardBuiltInToolCostTracking:
# Web Search
#########################################################
if StandardBuiltInToolCostTracking.response_object_includes_web_search_call(
response_object=response_object
response_object=response_object,
usage=usage,
):
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 custom_llm_provider == "anthropic":
return (
StandardBuiltInToolCostTracking.get_cost_for_anthropic_web_search(
model_info=model_info,
usage=usage,
)
)
else:
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
@ -72,7 +118,7 @@ class StandardBuiltInToolCostTracking:
@staticmethod
def response_object_includes_web_search_call(
response_object: Any,
response_object: Any, usage: Optional[Usage] = None
) -> bool:
"""
Check if the response object includes a web search call.
@ -91,6 +137,13 @@ class StandardBuiltInToolCostTracking:
return StandardBuiltInToolCostTracking.response_includes_output_type(
response_object=response_object, output_type="web_search_call"
)
elif (
usage is not None
and hasattr(usage, "server_tool_use")
and usage.server_tool_use is not None
and usage.server_tool_use.web_search_requests is not None
):
return True
return False
@staticmethod

View file

@ -6,6 +6,7 @@ import httpx
import litellm
from litellm.constants import (
ANTHROPIC_WEB_SEARCH_TOOL_MAX_USES,
DEFAULT_ANTHROPIC_CHAT_MAX_TOKENS,
DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET,
DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET,
@ -25,6 +26,8 @@ from litellm.types.llms.anthropic import (
AnthropicMessagesToolChoice,
AnthropicSystemMessageContent,
AnthropicThinkingParam,
AnthropicWebSearchTool,
AnthropicWebSearchUserLocation,
)
from litellm.types.llms.openai import (
REASONING_EFFORT,
@ -36,10 +39,11 @@ from litellm.types.llms.openai import (
ChatCompletionToolCallChunk,
ChatCompletionToolCallFunctionChunk,
ChatCompletionToolParam,
OpenAIWebSearchOptions,
)
from litellm.types.utils import CompletionTokensDetailsWrapper
from litellm.types.utils import Message as LitellmMessage
from litellm.types.utils import PromptTokensDetailsWrapper
from litellm.types.utils import PromptTokensDetailsWrapper, ServerToolUse
from litellm.utils import (
ModelResponse,
Usage,
@ -114,6 +118,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
"response_format",
"user",
"reasoning_effort",
"web_search_options",
]
if "claude-3-7-sonnet" in model:
@ -329,6 +334,37 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
return _tool
def map_web_search_tool(
self,
value: OpenAIWebSearchOptions,
) -> AnthropicWebSearchTool:
value_typed = cast(OpenAIWebSearchOptions, value)
hosted_web_search_tool = AnthropicWebSearchTool(
type="web_search_20250305",
name="web_search",
)
user_location = value_typed.get("user_location")
if user_location is not None:
anthropic_user_location = AnthropicWebSearchUserLocation(type="approximate")
anthropic_user_location_keys = (
AnthropicWebSearchUserLocation.__annotations__.keys()
)
user_location_approximate = user_location.get("approximate")
if user_location_approximate is not None:
for key, user_location_value in user_location_approximate.items():
if key in anthropic_user_location_keys and key != "type":
anthropic_user_location[key] = user_location_value # type: ignore
hosted_web_search_tool["user_location"] = anthropic_user_location
## MAP SEARCH CONTEXT SIZE
search_context_size = value_typed.get("search_context_size")
if search_context_size is not None:
hosted_web_search_tool["max_uses"] = ANTHROPIC_WEB_SEARCH_TOOL_MAX_USES[
search_context_size
]
return hosted_web_search_tool
def map_openai_params(
self,
non_default_params: dict,
@ -392,11 +428,19 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
optional_params["thinking"] = AnthropicConfig._map_reasoning_effort(
value
)
elif param == "web_search_options" and isinstance(value, dict):
hosted_web_search_tool = self.map_web_search_tool(
cast(OpenAIWebSearchOptions, value)
)
self._add_tools_to_optional_params(
optional_params=optional_params, tools=[hosted_web_search_tool]
)
## handle thinking tokens
self.update_optional_params_with_thinking_tokens(
non_default_params=non_default_params, optional_params=optional_params
)
return optional_params
def _create_json_tool_call_for_response_format(
@ -648,15 +692,20 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
_usage = usage_object
cache_creation_input_tokens: int = 0
cache_read_input_tokens: int = 0
web_search_requests: Optional[int] = None
if "cache_creation_input_tokens" in _usage:
cache_creation_input_tokens = _usage["cache_creation_input_tokens"]
if "cache_read_input_tokens" in _usage:
cache_read_input_tokens = _usage["cache_read_input_tokens"]
prompt_tokens += cache_read_input_tokens
if "server_tool_use" in _usage:
if "web_search_requests" in _usage["server_tool_use"]:
web_search_requests = cast(
int, _usage["server_tool_use"]["web_search_requests"]
)
prompt_tokens_details = PromptTokensDetailsWrapper(
cached_tokens=cache_read_input_tokens
cached_tokens=cache_read_input_tokens,
)
completion_token_details = (
CompletionTokensDetailsWrapper(
@ -668,6 +717,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
else None
)
total_tokens = prompt_tokens + completion_tokens
usage = Usage(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
@ -676,6 +726,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
cache_creation_input_tokens=cache_creation_input_tokens,
cache_read_input_tokens=cache_read_input_tokens,
completion_tokens_details=completion_token_details,
server_tool_use=ServerToolUse(web_search_requests=web_search_requests)
if web_search_requests is not None
else None,
)
return usage

View file

@ -105,6 +105,7 @@ class AzureOpenAIConfig(BaseConfig):
"prediction",
"modalities",
"audio",
"web_search_options",
]
def _is_response_format_supported_model(self, model: str) -> bool:

View file

@ -150,6 +150,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
"extra_headers",
"parallel_tool_calls",
"audio",
"web_search_options",
] # works across all models
model_specific_params = []

View file

@ -185,6 +185,7 @@ from .types.llms.openai import (
HttpxBinaryResponseContent,
ImageGenerationRequestQuality,
OpenAIModerationResponse,
OpenAIWebSearchOptions,
)
from .types.utils import (
LITELLM_IMAGE_VARIATION_PROVIDERS,
@ -353,6 +354,7 @@ async def acompletion(
extra_headers: Optional[dict] = None,
# Optional liteLLM function params
thinking: Optional[AnthropicThinkingParam] = None,
web_search_options: Optional[OpenAIWebSearchOptions] = None,
**kwargs,
) -> Union[ModelResponse, CustomStreamWrapper]:
"""
@ -472,6 +474,7 @@ async def acompletion(
"extra_headers": extra_headers,
"acompletion": True, # assuming this is a required parameter
"thinking": thinking,
"web_search_options": web_search_options,
}
if custom_llm_provider is None:
_, custom_llm_provider, _, _ = get_llm_provider(
@ -835,6 +838,7 @@ def completion( # type: ignore # noqa: PLR0915
logprobs: Optional[bool] = None,
top_logprobs: Optional[int] = None,
parallel_tool_calls: Optional[bool] = None,
web_search_options: Optional[OpenAIWebSearchOptions] = None,
deployment_id=None,
extra_headers: Optional[dict] = None,
# soon to be deprecated params by OpenAI
@ -1168,6 +1172,7 @@ def completion( # type: ignore # noqa: PLR0915
messages=messages,
reasoning_effort=reasoning_effort,
thinking=thinking,
web_search_options=web_search_options,
allowed_openai_params=kwargs.get("allowed_openai_params"),
**non_default_params,
)

View file

@ -1,6 +1,8 @@
from datetime import datetime, timezone
from litellm.litellm_core_utils.duration_parser import get_next_standardized_reset_time
from datetime import datetime, timezone
def get_budget_reset_timezone():
"""
Get the budget reset timezone from general_settings.
@ -8,11 +10,12 @@ def get_budget_reset_timezone():
"""
# Import at function level to avoid circular imports
from litellm.proxy.proxy_server import general_settings
if general_settings:
litellm_settings = general_settings.get("litellm_settings", {})
if litellm_settings and "timezone" in litellm_settings:
return litellm_settings["timezone"]
return "UTC"
@ -21,9 +24,10 @@ def get_budget_reset_time(budget_duration: str):
Get the budget reset time from general_settings.
Falls back to UTC if not specified.
"""
reset_at = get_next_standardized_reset_time(
duration=budget_duration,
duration=budget_duration,
current_time=datetime.now(timezone.utc),
timezone_str=get_budget_reset_timezone()
timezone_str=get_budget_reset_timezone(),
)
return reset_at
return reset_at

View file

@ -34,6 +34,7 @@ from litellm.proxy.auth.auth_checks import (
get_team_object,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
from litellm.proxy.hooks.key_management_event_hooks import KeyManagementEventHooks
from litellm.proxy.management_endpoints.common_utils import (
_is_user_team_admin,
@ -1208,8 +1209,7 @@ async def generate_key_helper_fn( # noqa: PLR0915
if budget_duration is None: # one-time budget
reset_at = None
else:
duration_s = duration_in_seconds(duration=budget_duration)
reset_at = datetime.now(timezone.utc) + timedelta(seconds=duration_s)
reset_at = get_budget_reset_time(budget_duration=budget_duration)
aliases_json = json.dumps(aliases)
config_json = json.dumps(config)

View file

@ -35,6 +35,22 @@ class AnthropicComputerTool(TypedDict, total=False):
name: Required[str]
class AnthropicWebSearchUserLocation(TypedDict, total=False):
city: Optional[str]
country: Optional[str]
region: Optional[str]
timezone: Optional[str]
type: Required[Literal["approximate"]]
class AnthropicWebSearchTool(TypedDict, total=False):
name: Required[Literal["web_search"]]
type: Required[str]
cache_control: Optional[Union[dict, ChatCompletionCachedContent]]
max_uses: Optional[int]
user_location: Optional[AnthropicWebSearchUserLocation]
class AnthropicHostedTools(TypedDict, total=False): # for bash_tool and text_editor
type: Required[str]
name: Required[str]
@ -42,7 +58,10 @@ class AnthropicHostedTools(TypedDict, total=False): # for bash_tool and text_ed
AllAnthropicToolsValues = Union[
AnthropicComputerTool, AnthropicHostedTools, AnthropicMessagesTool
AnthropicComputerTool,
AnthropicHostedTools,
AnthropicMessagesTool,
AnthropicWebSearchTool,
]

View file

@ -1355,3 +1355,20 @@ class OpenAIChatCompletionResponse(TypedDict, total=False):
OpenAIChatCompletionFinishReason = Literal[
"stop", "content_filter", "function_call", "tool_calls", "length"
]
class OpenAIWebSearchUserLocationApproximate(TypedDict):
city: str
country: str
region: str
timezone: str
class OpenAIWebSearchUserLocation(TypedDict):
approximate: OpenAIWebSearchUserLocationApproximate
type: Literal["approximate"]
class OpenAIWebSearchOptions(TypedDict, total=False):
search_context_size: Optional[Literal["low", "medium", "high"]]
user_location: Optional[OpenAIWebSearchUserLocation]

View file

@ -822,6 +822,9 @@ class PromptTokensDetailsWrapper(
image_tokens: Optional[int] = None
"""Image tokens sent to the model."""
web_search_requests: Optional[int] = None
"""Number of web search requests made by the tool call. Used for Anthropic to calculate web search cost."""
character_count: Optional[int] = None
"""Character count sent to the model. Used for Vertex AI multimodal embeddings."""
@ -839,6 +842,12 @@ class PromptTokensDetailsWrapper(
del self.image_count
if self.video_length_seconds is None:
del self.video_length_seconds
if self.web_search_requests is None:
del self.web_search_requests
class ServerToolUse(BaseModel):
web_search_requests: Optional[int]
class Usage(CompletionUsage):
@ -849,6 +858,8 @@ class Usage(CompletionUsage):
0
) # hidden param for prompt caching. Might change, once openai introduces their equivalent.
server_tool_use: Optional[ServerToolUse] = None
def __init__(
self,
prompt_tokens: Optional[int] = None,
@ -859,6 +870,7 @@ class Usage(CompletionUsage):
completion_tokens_details: Optional[
Union[CompletionTokensDetailsWrapper, dict]
] = None,
server_tool_use: Optional[ServerToolUse] = None,
**params,
):
# handle reasoning_tokens
@ -916,6 +928,11 @@ class Usage(CompletionUsage):
prompt_tokens_details=_prompt_tokens_details or None,
)
if server_tool_use is not None:
self.server_tool_use = server_tool_use
else: # maintain openai compatibility in usage object if possible
del self.server_tool_use
## ANTHROPIC MAPPING ##
if "cache_creation_input_tokens" in params and isinstance(
params["cache_creation_input_tokens"], int

View file

@ -143,6 +143,7 @@ from litellm.types.llms.openai import (
ChatCompletionToolParam,
ChatCompletionToolParamFunctionChunk,
OpenAITextCompletionUserMessage,
OpenAIWebSearchOptions,
)
from litellm.types.rerank import RerankResponse
from litellm.types.utils import FileTypes # type: ignore
@ -2682,6 +2683,7 @@ def get_optional_params( # noqa: PLR0915
additional_drop_params=None,
messages: Optional[List[AllMessageValues]] = None,
thinking: Optional[AnthropicThinkingParam] = None,
web_search_options: Optional[OpenAIWebSearchOptions] = None,
**kwargs,
):
# retrieve all parameters passed to the function
@ -2769,6 +2771,7 @@ def get_optional_params( # noqa: PLR0915
"messages": None,
"reasoning_effort": None,
"thinking": None,
"web_search_options": None,
}
# filter out those parameters that were passed with non-default values

View file

@ -1,6 +1,7 @@
import json
import os
import sys
from unittest.mock import MagicMock
import pytest
from fastapi.testclient import TestClient
@ -91,6 +92,7 @@ def test_get_cost_for_built_in_tools_web_search():
cost = StandardBuiltInToolCostTracking.get_cost_for_built_in_tools(
model=model,
usage=None,
response_object=None,
standard_built_in_tools_params=standard_built_in_tools_params,
)
@ -110,7 +112,25 @@ def test_get_cost_for_built_in_tools_file_search():
cost = StandardBuiltInToolCostTracking.get_cost_for_built_in_tools(
model=model,
response_object=None,
usage=None,
standard_built_in_tools_params=standard_built_in_tools_params,
)
assert cost == 0.00
def test_get_cost_for_anthropic_web_search():
"""
Test that the cost for a web search is 0.00 when no response object is provided
"""
from litellm.types.utils import ServerToolUse, Usage
model = "claude-3-7-sonnet-latest"
usage = Usage(server_tool_use=ServerToolUse(web_search_requests=1))
cost = StandardBuiltInToolCostTracking.get_cost_for_built_in_tools(
model=model,
usage=usage,
response_object=None,
standard_built_in_tools_params=None,
)
assert cost > 0.0

View file

@ -122,3 +122,69 @@ def test_map_tool_helper():
assert result is not None
assert result["name"] == "web_search"
assert result["max_uses"] == 5
def test_server_tool_use_usage():
config = AnthropicConfig()
usage_object = {
"input_tokens": 15956,
"cache_creation_input_tokens": 0,
"cache_read_input_tokens": 0,
"output_tokens": 567,
"server_tool_use": {"web_search_requests": 1},
}
usage = config.calculate_usage(usage_object=usage_object, reasoning_content=None)
assert usage.server_tool_use.web_search_requests == 1
def test_web_search_tool_transformation():
from litellm.types.llms.openai import OpenAIWebSearchOptions
config = AnthropicConfig()
openai_web_search_options = OpenAIWebSearchOptions(
user_location={
"type": "approximate",
"approximate": {
"city": "San Francisco",
},
}
)
anthropic_web_search_tool = config.map_web_search_tool(openai_web_search_options)
assert anthropic_web_search_tool is not None
assert anthropic_web_search_tool["user_location"] is not None
assert anthropic_web_search_tool["user_location"]["type"] == "approximate"
assert (
anthropic_web_search_tool["user_location"]["approximate"]["city"]
== "San Francisco"
)
@pytest.mark.parametrize(
"search_context_size, expected_max_uses", [("low", 1), ("medium", 5), ("high", 10)]
)
def test_web_search_tool_transformation_with_search_context_size(
search_context_size, expected_max_uses
):
from litellm.types.llms.openai import OpenAIWebSearchOptions
config = AnthropicConfig()
openai_web_search_options = OpenAIWebSearchOptions(
user_location={
"type": "approximate",
"approximate": {
"city": "San Francisco",
},
},
search_context_size=search_context_size,
)
anthropic_web_search_tool = config.map_web_search_tool(openai_web_search_options)
assert anthropic_web_search_tool is not None
assert anthropic_web_search_tool["user_location"] is not None
assert anthropic_web_search_tool["user_location"]["type"] == "approximate"
assert anthropic_web_search_tool["user_location"]["city"] == "San Francisco"
assert anthropic_web_search_tool["max_uses"] == expected_max_uses

View file

@ -0,0 +1,35 @@
import asyncio
import json
import os
import sys
import time
from datetime import datetime, timedelta, timezone
import pytest
from fastapi.testclient import TestClient
sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
def test_get_budget_reset_time():
"""
Test that the budget reset time is set to the first of the next month
"""
# Get the current date
now = datetime.now(timezone.utc)
# Calculate expected reset date (first of next month)
if now.month == 12:
expected_month = 1
expected_year = now.year + 1
else:
expected_month = now.month + 1
expected_year = now.year
expected_reset_at = datetime(expected_year, expected_month, 1, tzinfo=timezone.utc)
# Verify budget_reset_at is set to first of next month
assert get_budget_reset_time(budget_duration="1mo") == expected_reset_at

View file

@ -11,6 +11,8 @@ sys.path.insert(
from unittest.mock import AsyncMock, MagicMock
import pytest
from litellm.proxy.management_endpoints.key_management_endpoints import _list_key_helper
from litellm.proxy.proxy_server import app
@ -98,3 +100,70 @@ async def test_key_token_handling(monkeypatch):
assert (
response.token == response.token_id
), "Token should equal token_id if token_id exists"
@pytest.mark.asyncio
async def test_budget_reset_at_first_of_month(monkeypatch):
"""
Test that when budget_duration is "1mo", budget_reset_at is set to first of next month
"""
mock_prisma_client = AsyncMock()
mock_insert_data = AsyncMock(
return_value=MagicMock(token="hashed_token_123", litellm_budget_table=None)
)
mock_prisma_client.insert_data = mock_insert_data
mock_prisma_client.db = MagicMock()
mock_prisma_client.db.litellm_verificationtoken = MagicMock()
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=None
)
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[]
)
mock_prisma_client.db.litellm_verificationtoken.count = AsyncMock(return_value=0)
mock_prisma_client.db.litellm_verificationtoken.update = AsyncMock(
return_value=MagicMock(token="hashed_token_123", litellm_budget_table=None)
)
from datetime import datetime, timezone
import pytest
from litellm.proxy.management_endpoints.key_management_endpoints import (
generate_key_helper_fn,
)
from litellm.proxy.proxy_server import prisma_client
# Use monkeypatch to set the prisma_client
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
# Test key generation with budget_duration="1mo"
response = await generate_key_helper_fn(
request_type="user",
budget_duration="1mo",
user_id="test_user",
)
print(f"response: {response}\n")
# Get the current date
now = datetime.now(timezone.utc)
# Calculate expected reset date (first of next month)
if now.month == 12:
expected_month = 1
expected_year = now.year + 1
else:
expected_month = now.month + 1
expected_year = now.year
# Parse the response date
response_date = response["budget_reset_at"]
# Verify budget_reset_at is set to first of next month
assert (
response_date.year == expected_year
), f"Expected year {expected_year}, got {response_date.year}"
assert (
response_date.month == expected_month
), f"Expected month {expected_month}, got {response_date.month}"
assert response_date.day == 1, f"Expected day 1, got {response_date.day}"

View file

@ -1223,16 +1223,27 @@ async def test_anthropic_api_max_completion_tokens(model: str):
"model": model.split("/")[-1],
}
def test_anthropic_websearch():
@pytest.mark.parametrize(
"optional_params",
[
# {
# "tools": [{
# "type": "web_search_20250305",
# "name": "web_search",
# "max_uses": 5
# }]
# },
{
"web_search_options": {}
}
]
)
def test_anthropic_websearch(optional_params: dict):
litellm._turn_on_debug()
params = {
"model": "anthropic/claude-3-5-sonnet-latest",
"messages": [{"role": "user", "content": "What is the capital of France?"}],
"tools": [{
"type": "web_search_20250305",
"name": "web_search",
"max_uses": 5
}]
"messages": [{"role": "user", "content": "Who won the World Cup in 2022?"}],
**optional_params
}
try:
@ -1242,6 +1253,9 @@ def test_anthropic_websearch():
assert response is not None
print(f"response: {response}\n")
assert response.usage.server_tool_use.web_search_requests == 1
def test_anthropic_text_editor():
litellm._turn_on_debug()