mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-05 06:22:12 +00:00
Merge pull request #19548 from BerriAI/litellm_staging_01_22_2026
Litellm staging 01 22 2026
This commit is contained in:
@@ -234,6 +234,14 @@ db:
|
||||
# instance. See the "postgresql" top level key for additional configuration.
|
||||
deployStandalone: true
|
||||
|
||||
# Lifecycle hooks for the LiteLLM container
|
||||
# Example:
|
||||
# lifecycle:
|
||||
# preStop:
|
||||
# exec:
|
||||
# command: ["/bin/sh", "-c", "sleep 10"]
|
||||
lifecycle: {}
|
||||
|
||||
# Settings for Bitnami postgresql chart (if db.deployStandalone is true, ignored
|
||||
# otherwise)
|
||||
postgresql:
|
||||
|
||||
@@ -422,6 +422,10 @@ def cost_per_token( # noqa: PLR0915
|
||||
)
|
||||
|
||||
return dashscope_cost_per_token(model=model, usage=usage_block)
|
||||
elif custom_llm_provider == "azure_ai":
|
||||
return generic_cost_per_token(
|
||||
model=model, usage=usage_block, custom_llm_provider=custom_llm_provider
|
||||
)
|
||||
else:
|
||||
model_info = _cached_get_model_info_helper(
|
||||
model=model, custom_llm_provider=custom_llm_provider
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Union
|
||||
from copy import deepcopy
|
||||
|
||||
import httpx
|
||||
from openai.types.responses import ResponseReasoningItem
|
||||
@@ -43,7 +44,7 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
||||
"""
|
||||
Handle reasoning items to filter out the status field.
|
||||
Issue: https://github.com/BerriAI/litellm/issues/13484
|
||||
|
||||
|
||||
Azure OpenAI API does not accept 'status' field in reasoning input items.
|
||||
"""
|
||||
if item.get("type") == "reasoning":
|
||||
@@ -78,7 +79,7 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
||||
}
|
||||
return filtered_item
|
||||
return item
|
||||
|
||||
|
||||
def _validate_input_param(
|
||||
self, input: Union[str, ResponseInputParam]
|
||||
) -> Union[str, ResponseInputParam]:
|
||||
@@ -90,7 +91,7 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
||||
|
||||
# First call parent's validation
|
||||
validated_input = super()._validate_input_param(input)
|
||||
|
||||
|
||||
# Then filter out status from message items
|
||||
if isinstance(validated_input, list):
|
||||
filtered_input: List[Any] = []
|
||||
@@ -102,7 +103,7 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
||||
else:
|
||||
filtered_input.append(item)
|
||||
return cast(ResponseInputParam, filtered_input)
|
||||
|
||||
|
||||
return validated_input
|
||||
|
||||
def transform_responses_api_request(
|
||||
@@ -116,6 +117,21 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
||||
"""No transform applied since inputs are in OpenAI spec already"""
|
||||
stripped_model_name = self.get_stripped_model_name(model)
|
||||
|
||||
# Azure Responses API requires flattened tools (params at top level, not nested in 'function')
|
||||
if "tools" in response_api_optional_request_params and isinstance(
|
||||
response_api_optional_request_params["tools"], list
|
||||
):
|
||||
new_tools: List[Dict[str, Any]] = []
|
||||
for tool in response_api_optional_request_params["tools"]:
|
||||
if isinstance(tool, dict) and "function" in tool:
|
||||
new_tool: Dict[str, Any] = deepcopy(tool)
|
||||
function_data = new_tool.pop("function")
|
||||
new_tool.update(function_data)
|
||||
new_tools.append(new_tool)
|
||||
else:
|
||||
new_tools.append(tool)
|
||||
response_api_optional_request_params["tools"] = new_tools
|
||||
|
||||
return super().transform_responses_api_request(
|
||||
model=stripped_model_name,
|
||||
input=input,
|
||||
|
||||
@@ -1527,7 +1527,7 @@ class AWSEventStreamDecoder:
|
||||
]
|
||||
] = None
|
||||
|
||||
index = int(chunk_data.get("contentBlockIndex", 0))
|
||||
content_block_index = int(chunk_data.get("contentBlockIndex", 0))
|
||||
if "start" in chunk_data:
|
||||
start_obj = ContentBlockStartEvent(**chunk_data["start"])
|
||||
(
|
||||
@@ -1543,11 +1543,11 @@ class AWSEventStreamDecoder:
|
||||
provider_specific_fields,
|
||||
reasoning_content,
|
||||
thinking_blocks,
|
||||
) = self._handle_converse_delta_event(delta_obj, index)
|
||||
) = self._handle_converse_delta_event(delta_obj, content_block_index)
|
||||
elif (
|
||||
"contentBlockIndex" in chunk_data
|
||||
): # stop block, no 'start' or 'delta' object
|
||||
tool_use = self._handle_converse_stop_event(index)
|
||||
tool_use = self._handle_converse_stop_event(content_block_index)
|
||||
elif "stopReason" in chunk_data:
|
||||
finish_reason = map_finish_reason(chunk_data.get("stopReason", "stop"))
|
||||
elif "usage" in chunk_data:
|
||||
@@ -1561,7 +1561,7 @@ class AWSEventStreamDecoder:
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
finish_reason=finish_reason,
|
||||
index=index,
|
||||
index=0, # Always 0 - Bedrock never returns multiple choices
|
||||
delta=Delta(
|
||||
content=text,
|
||||
role="assistant",
|
||||
|
||||
@@ -42,6 +42,7 @@ from litellm.proxy._experimental.mcp_server.utils import (
|
||||
add_server_prefix_to_name,
|
||||
get_server_prefix,
|
||||
is_tool_name_prefixed,
|
||||
merge_mcp_headers,
|
||||
normalize_server_name,
|
||||
split_server_prefix_from_name,
|
||||
validate_mcp_server_name,
|
||||
@@ -372,7 +373,7 @@ class MCPServerManager:
|
||||
server_prefix = get_server_prefix(server)
|
||||
|
||||
# Build headers from server configuration
|
||||
headers = {}
|
||||
headers: Dict[str, str] = {}
|
||||
|
||||
# Add authentication headers if configured
|
||||
if server.authentication_token:
|
||||
@@ -385,10 +386,15 @@ class MCPServerManager:
|
||||
elif server.auth_type == MCPAuth.basic:
|
||||
headers["Authorization"] = f"Basic {server.authentication_token}"
|
||||
|
||||
# Add any extra headers from server config
|
||||
# Note: extra_headers is a List[str] of header names to forward, not a dict
|
||||
# For OpenAPI tools, we'll just use the authentication headers
|
||||
# If extra_headers were needed, they would be processed separately
|
||||
# Add any static headers from server config.
|
||||
#
|
||||
# Note: `extra_headers` on MCPServer is a List[str] of header names to forward
|
||||
# from the client request (not available in this OpenAPI tool generation step).
|
||||
# `static_headers` is a dict of concrete headers to always send.
|
||||
headers = merge_mcp_headers(
|
||||
extra_headers=headers,
|
||||
static_headers=server.static_headers,
|
||||
) or {}
|
||||
|
||||
verbose_logger.debug(
|
||||
f"Using headers for OpenAPI tools (excluding sensitive values): "
|
||||
|
||||
@@ -8,6 +8,7 @@ from litellm._logging import verbose_logger
|
||||
from litellm.proxy._experimental.mcp_server.ui_session_utils import (
|
||||
build_effective_auth_contexts,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.utils import merge_mcp_headers
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.types.mcp import MCPAuth
|
||||
@@ -438,16 +439,22 @@ if MCP_AVAILABLE:
|
||||
command=request.command,
|
||||
args=request.args,
|
||||
env=request.env,
|
||||
static_headers=request.static_headers,
|
||||
)
|
||||
|
||||
stdio_env = global_mcp_server_manager._build_stdio_env(
|
||||
server_model, raw_headers
|
||||
)
|
||||
|
||||
merged_headers = merge_mcp_headers(
|
||||
extra_headers=oauth2_headers,
|
||||
static_headers=request.static_headers,
|
||||
)
|
||||
|
||||
client = global_mcp_server_manager._create_mcp_client(
|
||||
server=server_model,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
extra_headers=oauth2_headers,
|
||||
extra_headers=merged_headers,
|
||||
stdio_env=stdio_env,
|
||||
)
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
"""
|
||||
MCP Server Utilities
|
||||
"""
|
||||
from typing import Tuple, Any
|
||||
from typing import Any, Dict, Mapping, Optional, Tuple
|
||||
|
||||
import os
|
||||
import importlib
|
||||
@@ -137,3 +137,31 @@ def validate_mcp_server_name(
|
||||
)
|
||||
else:
|
||||
raise Exception(error_message)
|
||||
|
||||
|
||||
def merge_mcp_headers(
|
||||
*,
|
||||
extra_headers: Optional[Mapping[str, str]] = None,
|
||||
static_headers: Optional[Mapping[str, str]] = None,
|
||||
) -> Optional[Dict[str, str]]:
|
||||
"""Merge outbound HTTP headers for MCP calls.
|
||||
|
||||
This is used when calling out to external MCP servers (or OpenAPI-based MCP tools).
|
||||
|
||||
Merge rules:
|
||||
- Start with `extra_headers` (typically OAuth2-derived headers)
|
||||
- Overlay `static_headers` (user-configured per MCP server)
|
||||
|
||||
If both contain the same key, `static_headers` wins. This matches the existing
|
||||
behavior in `MCPServerManager` where `server.static_headers` is applied after
|
||||
any caller-provided headers.
|
||||
"""
|
||||
merged: Dict[str, str] = {}
|
||||
|
||||
if extra_headers:
|
||||
merged.update({str(k): str(v) for k, v in extra_headers.items()})
|
||||
|
||||
if static_headers:
|
||||
merged.update({str(k): str(v) for k, v in static_headers.items()})
|
||||
|
||||
return merged or None
|
||||
|
||||
@@ -234,7 +234,10 @@ async def chat_completion_pass_through_endpoint( # noqa: PLR0915
|
||||
elif (
|
||||
llm_router is not None
|
||||
and data["model"] not in router_model_names
|
||||
and (llm_router.default_deployment is not None or len(llm_router.pattern_router.patterns) > 0)
|
||||
and (
|
||||
llm_router.default_deployment is not None
|
||||
or len(llm_router.pattern_router.patterns) > 0
|
||||
)
|
||||
): # check for wildcard routes or default deployment before checking deployment_names
|
||||
llm_response = asyncio.create_task(llm_router.aadapter_completion(**data))
|
||||
elif (
|
||||
@@ -443,10 +446,10 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
|
||||
|
||||
for field_name, field_value in form_data.items():
|
||||
if isinstance(field_value, (StarletteUploadFile, UploadFile)):
|
||||
files[field_name] = (
|
||||
await HttpPassThroughEndpointHelpers._build_request_files_from_upload_file(
|
||||
upload_file=field_value
|
||||
)
|
||||
files[
|
||||
field_name
|
||||
] = await HttpPassThroughEndpointHelpers._build_request_files_from_upload_file(
|
||||
upload_file=field_value
|
||||
)
|
||||
else:
|
||||
form_data_dict[field_name] = field_value
|
||||
@@ -539,9 +542,9 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
|
||||
"passthrough_logging_payload": passthrough_logging_payload,
|
||||
}
|
||||
|
||||
logging_obj.model_call_details["passthrough_logging_payload"] = (
|
||||
passthrough_logging_payload
|
||||
)
|
||||
logging_obj.model_call_details[
|
||||
"passthrough_logging_payload"
|
||||
] = passthrough_logging_payload
|
||||
|
||||
return kwargs
|
||||
|
||||
@@ -678,7 +681,7 @@ async def pass_through_request( # noqa: PLR0915
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
passthrough_guardrails_config=guardrails_config,
|
||||
)
|
||||
|
||||
|
||||
# Add guardrails to metadata if any should run
|
||||
if guardrails_to_run and len(guardrails_to_run) > 0:
|
||||
if _parsed_body is None:
|
||||
@@ -701,10 +704,10 @@ async def pass_through_request( # noqa: PLR0915
|
||||
litellm_call_id=litellm_call_id,
|
||||
function_id="1245",
|
||||
)
|
||||
|
||||
|
||||
# Store passthrough guardrails config on logging_obj for field targeting
|
||||
logging_obj.passthrough_guardrails_config = guardrails_config
|
||||
|
||||
|
||||
# Store logging_obj in data so guardrails can access it
|
||||
if _parsed_body is None:
|
||||
_parsed_body = {}
|
||||
@@ -739,7 +742,9 @@ async def pass_through_request( # noqa: PLR0915
|
||||
# Store custom_llm_provider in kwargs and logging object if provided
|
||||
if custom_llm_provider:
|
||||
logging_obj.model_call_details["custom_llm_provider"] = custom_llm_provider
|
||||
logging_obj.model_call_details["litellm_params"] = kwargs.get("litellm_params", {})
|
||||
logging_obj.model_call_details["litellm_params"] = kwargs.get(
|
||||
"litellm_params", {}
|
||||
)
|
||||
|
||||
# done for supporting 'parallel_request_limiter.py' with pass-through endpoints
|
||||
logging_obj.update_environment_variables(
|
||||
@@ -929,12 +934,16 @@ async def pass_through_request( # noqa: PLR0915
|
||||
if kwargs:
|
||||
for key, value in kwargs.items():
|
||||
request_payload[key] = value
|
||||
|
||||
if "model" not in request_payload and _parsed_body and isinstance(_parsed_body, dict):
|
||||
|
||||
if (
|
||||
"model" not in request_payload
|
||||
and _parsed_body
|
||||
and isinstance(_parsed_body, dict)
|
||||
):
|
||||
request_payload["model"] = _parsed_body.get("model", "")
|
||||
if "custom_llm_provider" not in request_payload and custom_llm_provider:
|
||||
request_payload["custom_llm_provider"] = custom_llm_provider
|
||||
|
||||
|
||||
await proxy_logging_obj.post_call_failure_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
original_exception=e,
|
||||
@@ -1443,9 +1452,9 @@ async def websocket_passthrough_request( # noqa: PLR0915
|
||||
)
|
||||
if extracted_model:
|
||||
kwargs["model"] = extracted_model
|
||||
kwargs["custom_llm_provider"] = (
|
||||
"vertex_ai-language-models"
|
||||
)
|
||||
kwargs[
|
||||
"custom_llm_provider"
|
||||
] = "vertex_ai-language-models"
|
||||
# Update logging object with correct model
|
||||
logging_obj.model = extracted_model
|
||||
logging_obj.model_call_details[
|
||||
@@ -1511,9 +1520,9 @@ async def websocket_passthrough_request( # noqa: PLR0915
|
||||
# Update logging object with correct model
|
||||
logging_obj.model = extracted_model
|
||||
logging_obj.model_call_details["model"] = extracted_model
|
||||
logging_obj.model_call_details["custom_llm_provider"] = (
|
||||
"vertex_ai_language_models"
|
||||
)
|
||||
logging_obj.model_call_details[
|
||||
"custom_llm_provider"
|
||||
] = "vertex_ai_language_models"
|
||||
verbose_proxy_logger.debug(
|
||||
f"WebSocket passthrough ({endpoint}): Successfully extracted model '{extracted_model}' and set provider to 'vertex_ai' from server setup response"
|
||||
)
|
||||
@@ -1841,10 +1850,9 @@ class InitPassThroughEndpointHelpers:
|
||||
# Check if this exact route is already registered
|
||||
if route_key in _registered_pass_through_routes:
|
||||
verbose_proxy_logger.debug(
|
||||
"Skipping duplicate exact pass through endpoint: %s (already registered)",
|
||||
"Updating duplicate exact pass through endpoint: %s (already registered)",
|
||||
path,
|
||||
)
|
||||
return
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
"adding exact pass through endpoint: %s, dependencies: %s",
|
||||
@@ -1853,7 +1861,7 @@ class InitPassThroughEndpointHelpers:
|
||||
)
|
||||
|
||||
# Use SafeRouteAdder to only add route if it doesn't exist on the app
|
||||
was_added = SafeRouteAdder.add_api_route_if_not_exists(
|
||||
SafeRouteAdder.add_api_route_if_not_exists(
|
||||
app=app,
|
||||
path=path,
|
||||
endpoint=create_pass_through_route( # type: ignore
|
||||
@@ -1870,22 +1878,21 @@ class InitPassThroughEndpointHelpers:
|
||||
dependencies=dependencies,
|
||||
)
|
||||
|
||||
# Register the route to prevent duplicates only if it was added
|
||||
if was_added:
|
||||
_registered_pass_through_routes[route_key] = {
|
||||
"endpoint_id": endpoint_id,
|
||||
"path": path,
|
||||
"type": "exact",
|
||||
"passthrough_params": {
|
||||
"target": target,
|
||||
"custom_headers": custom_headers,
|
||||
"forward_headers": forward_headers,
|
||||
"merge_query_params": merge_query_params,
|
||||
"dependencies": dependencies,
|
||||
"cost_per_request": cost_per_request,
|
||||
"guardrails": guardrails,
|
||||
},
|
||||
}
|
||||
# Always register/update the route metadata (headers, target) even if FastAPI route exists
|
||||
_registered_pass_through_routes[route_key] = {
|
||||
"endpoint_id": endpoint_id,
|
||||
"path": path,
|
||||
"type": "exact",
|
||||
"passthrough_params": {
|
||||
"target": target,
|
||||
"custom_headers": custom_headers,
|
||||
"forward_headers": forward_headers,
|
||||
"merge_query_params": merge_query_params,
|
||||
"dependencies": dependencies,
|
||||
"cost_per_request": cost_per_request,
|
||||
"guardrails": guardrails,
|
||||
},
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def add_subpath_route(
|
||||
@@ -1907,10 +1914,9 @@ class InitPassThroughEndpointHelpers:
|
||||
# Check if this subpath route is already registered
|
||||
if route_key in _registered_pass_through_routes:
|
||||
verbose_proxy_logger.debug(
|
||||
"Skipping duplicate wildcard pass through endpoint: %s (already registered)",
|
||||
"Updating duplicate wildcard pass through endpoint: %s (already registered)",
|
||||
wildcard_path,
|
||||
)
|
||||
return
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
"adding wildcard pass through endpoint: %s, dependencies: %s",
|
||||
@@ -1919,7 +1925,7 @@ class InitPassThroughEndpointHelpers:
|
||||
)
|
||||
|
||||
# Use SafeRouteAdder to only add route if it doesn't exist on the app
|
||||
was_added = SafeRouteAdder.add_api_route_if_not_exists(
|
||||
SafeRouteAdder.add_api_route_if_not_exists(
|
||||
app=app,
|
||||
path=wildcard_path,
|
||||
endpoint=create_pass_through_route( # type: ignore
|
||||
@@ -1938,21 +1944,20 @@ class InitPassThroughEndpointHelpers:
|
||||
)
|
||||
|
||||
# Register the route to prevent duplicates only if it was added
|
||||
if was_added:
|
||||
_registered_pass_through_routes[route_key] = {
|
||||
"endpoint_id": endpoint_id,
|
||||
"path": path,
|
||||
"type": "subpath",
|
||||
"passthrough_params": {
|
||||
"target": target,
|
||||
"custom_headers": custom_headers,
|
||||
"forward_headers": forward_headers,
|
||||
"merge_query_params": merge_query_params,
|
||||
"dependencies": dependencies,
|
||||
"cost_per_request": cost_per_request,
|
||||
"guardrails": guardrails,
|
||||
},
|
||||
}
|
||||
_registered_pass_through_routes[route_key] = {
|
||||
"endpoint_id": endpoint_id,
|
||||
"path": path,
|
||||
"type": "subpath",
|
||||
"passthrough_params": {
|
||||
"target": target,
|
||||
"custom_headers": custom_headers,
|
||||
"forward_headers": forward_headers,
|
||||
"merge_query_params": merge_query_params,
|
||||
"dependencies": dependencies,
|
||||
"cost_per_request": cost_per_request,
|
||||
"guardrails": guardrails,
|
||||
},
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def remove_endpoint_routes(endpoint_id: str):
|
||||
@@ -2149,7 +2154,7 @@ async def initialize_pass_through_endpoints(
|
||||
|
||||
# Get guardrails config if present
|
||||
_guardrails = endpoint.get("guardrails", None)
|
||||
|
||||
|
||||
# Add exact path route
|
||||
verbose_proxy_logger.debug(
|
||||
"Initializing pass through endpoint: %s (ID: %s)", _path, endpoint_id
|
||||
@@ -2328,6 +2333,7 @@ async def get_pass_through_endpoints(
|
||||
async def update_pass_through_endpoints(
|
||||
endpoint_id: str,
|
||||
data: PassThroughGenericEndpoint,
|
||||
request: Request,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
@@ -2418,6 +2424,37 @@ async def update_pass_through_endpoints(
|
||||
data=updated_data, user_api_key_dict=user_api_key_dict
|
||||
)
|
||||
|
||||
# Re-register the route with updated headers
|
||||
_custom_headers: Optional[dict] = updated_endpoint.headers or {}
|
||||
_custom_headers = await set_env_variables_in_header(custom_headers=_custom_headers)
|
||||
|
||||
if updated_endpoint.include_subpath:
|
||||
InitPassThroughEndpointHelpers.add_subpath_route(
|
||||
app=request.app,
|
||||
path=updated_endpoint.path,
|
||||
target=updated_endpoint.target,
|
||||
custom_headers=_custom_headers,
|
||||
forward_headers=None, # Defaults not available in model? assuming None logic handles it
|
||||
merge_query_params=None,
|
||||
dependencies=None,
|
||||
cost_per_request=updated_endpoint.cost_per_request,
|
||||
endpoint_id=updated_endpoint.id or endpoint_id or "",
|
||||
guardrails=getattr(updated_endpoint, "guardrails", None),
|
||||
)
|
||||
else:
|
||||
InitPassThroughEndpointHelpers.add_exact_path_route(
|
||||
app=request.app,
|
||||
path=updated_endpoint.path,
|
||||
target=updated_endpoint.target,
|
||||
custom_headers=_custom_headers,
|
||||
forward_headers=None,
|
||||
merge_query_params=None,
|
||||
dependencies=None,
|
||||
cost_per_request=updated_endpoint.cost_per_request,
|
||||
endpoint_id=updated_endpoint.id or endpoint_id or "",
|
||||
guardrails=getattr(updated_endpoint, "guardrails", None),
|
||||
)
|
||||
|
||||
return PassThroughEndpointResponse(
|
||||
endpoints=[updated_endpoint] if updated_endpoint else []
|
||||
)
|
||||
@@ -2429,6 +2466,7 @@ async def update_pass_through_endpoints(
|
||||
)
|
||||
async def create_pass_through_endpoints(
|
||||
data: PassThroughGenericEndpoint,
|
||||
request: Request,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
@@ -2473,6 +2511,38 @@ async def create_pass_through_endpoints(
|
||||
|
||||
# Return the created endpoint with the generated ID
|
||||
created_endpoint = PassThroughGenericEndpoint(**data_dict)
|
||||
|
||||
# Register the new route
|
||||
_custom_headers: Optional[dict] = created_endpoint.headers or {}
|
||||
_custom_headers = await set_env_variables_in_header(custom_headers=_custom_headers)
|
||||
|
||||
if created_endpoint.include_subpath:
|
||||
InitPassThroughEndpointHelpers.add_subpath_route(
|
||||
app=request.app,
|
||||
path=created_endpoint.path,
|
||||
target=created_endpoint.target,
|
||||
custom_headers=_custom_headers,
|
||||
forward_headers=None,
|
||||
merge_query_params=None,
|
||||
dependencies=None,
|
||||
cost_per_request=created_endpoint.cost_per_request,
|
||||
endpoint_id=created_endpoint.id or "",
|
||||
guardrails=getattr(created_endpoint, "guardrails", None),
|
||||
)
|
||||
else:
|
||||
InitPassThroughEndpointHelpers.add_exact_path_route(
|
||||
app=request.app,
|
||||
path=created_endpoint.path,
|
||||
target=created_endpoint.target,
|
||||
custom_headers=_custom_headers,
|
||||
forward_headers=None,
|
||||
merge_query_params=None,
|
||||
dependencies=None,
|
||||
cost_per_request=created_endpoint.cost_per_request,
|
||||
endpoint_id=created_endpoint.id or "",
|
||||
guardrails=getattr(created_endpoint, "guardrails", None),
|
||||
)
|
||||
|
||||
return PassThroughEndpointResponse(endpoints=[created_endpoint])
|
||||
|
||||
|
||||
|
||||
@@ -190,7 +190,7 @@ class PolicyValidator:
|
||||
PolicyValidationError(
|
||||
policy_name=policy_name,
|
||||
error_type=PolicyValidationErrorType.CIRCULAR_INHERITANCE,
|
||||
message=f"Inheritance chain too deep (exceeded max depth of 100)",
|
||||
message="Inheritance chain too deep (exceeded max depth of 100)",
|
||||
field="inherit",
|
||||
)
|
||||
)
|
||||
|
||||
@@ -546,7 +546,7 @@ except ImportError:
|
||||
enterprise_proxy_config = None
|
||||
###################
|
||||
|
||||
server_root_path = os.getenv("SERVER_ROOT_PATH", "")
|
||||
server_root_path = get_server_root_path()
|
||||
_license_check = LicenseCheck()
|
||||
premium_user: bool = _license_check.is_premium()
|
||||
premium_user_data: Optional["EnterpriseLicenseData"] = (
|
||||
@@ -824,7 +824,6 @@ app = FastAPI(
|
||||
title=_title,
|
||||
description=_description,
|
||||
version=version,
|
||||
root_path=server_root_path, # check if user passed root path, FastAPI defaults this value to ""
|
||||
lifespan=proxy_startup_event, # type: ignore[reportGeneralTypeIssues]
|
||||
)
|
||||
|
||||
|
||||
+25
-19
@@ -7,7 +7,7 @@ import smtplib
|
||||
import threading
|
||||
import time
|
||||
import traceback
|
||||
from datetime import datetime, timedelta
|
||||
from datetime import date, datetime, timedelta
|
||||
from email.mime.multipart import MIMEMultipart
|
||||
from email.mime.text import MIMEText
|
||||
from typing import (
|
||||
@@ -3891,11 +3891,15 @@ def _raise_failed_update_spend_exception(
|
||||
raise e
|
||||
|
||||
|
||||
def _get_month_end_date(today: date) -> date:
|
||||
if today.month == 12:
|
||||
return date(today.year + 1, 1, 1) - timedelta(days=1)
|
||||
return date(today.year, today.month + 1, 1) - timedelta(days=1)
|
||||
|
||||
|
||||
def _is_projected_spend_over_limit(
|
||||
current_spend: float, soft_budget_limit: Optional[float]
|
||||
):
|
||||
from datetime import date
|
||||
|
||||
if soft_budget_limit is None:
|
||||
# If there's no limit, we can't exceed it.
|
||||
return False
|
||||
@@ -3903,10 +3907,7 @@ def _is_projected_spend_over_limit(
|
||||
today = date.today()
|
||||
|
||||
# Finding the first day of the next month, then subtracting one day to get the end of the current month.
|
||||
if today.month == 12: # December edge case
|
||||
end_month = date(today.year + 1, 1, 1) - timedelta(days=1)
|
||||
else:
|
||||
end_month = date(today.year, today.month + 1, 1) - timedelta(days=1)
|
||||
end_month = _get_month_end_date(today)
|
||||
|
||||
remaining_days = (end_month - today).days
|
||||
|
||||
@@ -3928,25 +3929,30 @@ def _is_projected_spend_over_limit(
|
||||
def _get_projected_spend_over_limit(
|
||||
current_spend: float, soft_budget_limit: Optional[float]
|
||||
) -> Optional[tuple]:
|
||||
import datetime
|
||||
|
||||
if soft_budget_limit is None:
|
||||
return None
|
||||
|
||||
today = datetime.date.today()
|
||||
end_month = datetime.date(today.year, today.month + 1, 1) - datetime.timedelta(
|
||||
days=1
|
||||
)
|
||||
today = date.today()
|
||||
end_month = _get_month_end_date(today)
|
||||
remaining_days = (end_month - today).days
|
||||
|
||||
daily_spend = current_spend / (
|
||||
today.day - 1
|
||||
) # assuming the current spend till today (not including today)
|
||||
projected_spend = daily_spend * remaining_days
|
||||
# assuming the current spend till today (not including today)
|
||||
if today.day == 1:
|
||||
daily_spend = current_spend
|
||||
else:
|
||||
daily_spend = current_spend / (today.day - 1)
|
||||
projected_spend = current_spend + (daily_spend * remaining_days)
|
||||
|
||||
if projected_spend > soft_budget_limit:
|
||||
approx_days = soft_budget_limit / daily_spend
|
||||
limit_exceed_date = today + datetime.timedelta(days=approx_days)
|
||||
if daily_spend <= 0:
|
||||
limit_exceed_date = today
|
||||
else:
|
||||
remaining_budget = soft_budget_limit - current_spend
|
||||
if remaining_budget <= 0:
|
||||
limit_exceed_date = today
|
||||
else:
|
||||
approx_days = remaining_budget / daily_spend
|
||||
limit_exceed_date = today + timedelta(days=approx_days)
|
||||
|
||||
# return the projected spend and the date it will exceeded
|
||||
return projected_spend, limit_exceed_date
|
||||
|
||||
+12
-17
@@ -5,11 +5,6 @@ from typing import Any, Dict, List, Literal, Optional, Union
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator
|
||||
from typing_extensions import Required, TypedDict
|
||||
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
ChatCompletionToolCallChunk,
|
||||
ChatCompletionToolParam,
|
||||
)
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.enkryptai import (
|
||||
EnkryptAIGuardrailConfigs,
|
||||
)
|
||||
@@ -673,20 +668,20 @@ class LitellmParams(
|
||||
description="When to apply the guardrail (pre_call, post_call, during_call, logging_only)"
|
||||
)
|
||||
|
||||
@field_validator("default_action", mode="before", check_fields=False)
|
||||
@field_validator(
|
||||
"mode",
|
||||
"default_action",
|
||||
"on_disallowed_action",
|
||||
mode="before",
|
||||
check_fields=False,
|
||||
)
|
||||
@classmethod
|
||||
def normalize_default_action_litellm_params(cls, v):
|
||||
"""Normalize default_action to lowercase for ALL guardrail types."""
|
||||
if isinstance(v, str):
|
||||
return v.lower()
|
||||
return v
|
||||
|
||||
@field_validator("on_disallowed_action", mode="before", check_fields=False)
|
||||
@classmethod
|
||||
def normalize_on_disallowed_action_litellm_params(cls, v):
|
||||
"""Normalize on_disallowed_action to lowercase for ALL guardrail types."""
|
||||
def normalize_lowercase(cls, v):
|
||||
"""Normalize string and list fields to lowercase for ALL guardrail types."""
|
||||
if isinstance(v, str):
|
||||
return v.lower()
|
||||
if isinstance(v, list):
|
||||
return [x.lower() if isinstance(x, str) else x for x in v]
|
||||
return v
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
@@ -695,7 +690,7 @@ class LitellmParams(
|
||||
kwargs["default_on"] = default_on
|
||||
else:
|
||||
kwargs["default_on"] = False
|
||||
|
||||
|
||||
super().__init__(**kwargs)
|
||||
|
||||
def __contains__(self, key):
|
||||
|
||||
+1
-1
@@ -4722,7 +4722,7 @@ def calculate_max_parallel_requests(
|
||||
elif rpm is not None:
|
||||
return rpm
|
||||
elif tpm is not None:
|
||||
calculated_rpm = int(tpm / 1000 / 6)
|
||||
calculated_rpm = int(tpm / 1000 * 6)
|
||||
if calculated_rpm == 0:
|
||||
calculated_rpm = 1
|
||||
return calculated_rpm
|
||||
|
||||
@@ -2,8 +2,11 @@ import os
|
||||
import sys
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../..")
|
||||
) # Adds the parent directory to the system path
|
||||
0,
|
||||
os.path.abspath(
|
||||
os.path.join(os.path.dirname(__file__), "../../litellm-proxy-extras")
|
||||
),
|
||||
)
|
||||
|
||||
from litellm_proxy_extras.utils import ProxyExtrasDBManager
|
||||
|
||||
@@ -99,6 +102,11 @@ class TestIdempotentErrorDetection:
|
||||
error_message = "COLUMN 'ID' ALREADY EXISTS"
|
||||
assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True
|
||||
|
||||
def test_is_idempotent_error_does_not_exist(self):
|
||||
"""Test detection of 'does not exist' error"""
|
||||
error_message = "ERROR: index 'idx' does not exist"
|
||||
assert ProxyExtrasDBManager._is_idempotent_error(error_message) is True
|
||||
|
||||
def test_is_idempotent_error_negative(self):
|
||||
"""Test that non-idempotent errors are not detected as idempotent errors"""
|
||||
error_message = "Database error code: 42501 - permission denied"
|
||||
|
||||
@@ -114,7 +114,7 @@ def test_nova_invoke_streaming_chunk_parsing():
|
||||
}
|
||||
result = decoder._chunk_parser(nova_tool_start_chunk)
|
||||
assert result.choices[0].delta.content == ""
|
||||
assert result.choices[0].index == 1
|
||||
assert result.choices[0].index == 0
|
||||
assert result.choices[0].delta.tool_calls is not None
|
||||
assert result.choices[0].delta.tool_calls[0].type == "function"
|
||||
assert result.choices[0].delta.tool_calls[0].function.name == "get_weather"
|
||||
@@ -129,7 +129,7 @@ def test_nova_invoke_streaming_chunk_parsing():
|
||||
}
|
||||
result = decoder._chunk_parser(nova_tool_args_chunk)
|
||||
assert result.choices[0].delta.content == ""
|
||||
assert result.choices[0].index == 2
|
||||
assert result.choices[0].index == 0
|
||||
assert result.choices[0].delta.tool_calls is not None
|
||||
assert (
|
||||
result.choices[0].delta.tool_calls[0].function.arguments
|
||||
|
||||
@@ -54,7 +54,7 @@ def test_scenario(max_parallel_requests, tpm, rpm, default_max_parallel_requests
|
||||
elif rpm is not None:
|
||||
assert rpm == calculated_max_parallel_requests
|
||||
elif tpm is not None:
|
||||
calculated_rpm = int(tpm / 1000 / 6)
|
||||
calculated_rpm = int(tpm / 1000 * 6)
|
||||
if calculated_rpm == 0:
|
||||
calculated_rpm = 1
|
||||
print(
|
||||
@@ -107,7 +107,7 @@ def test_setting_mpr_limits_per_model(
|
||||
elif rpm is not None:
|
||||
assert rpm == mpr_client._value
|
||||
elif tpm is not None:
|
||||
calculated_rpm = int(tpm / 1000 / 6)
|
||||
calculated_rpm = int(tpm / 1000 * 6)
|
||||
if calculated_rpm == 0:
|
||||
calculated_rpm = 1
|
||||
print(
|
||||
|
||||
@@ -0,0 +1,145 @@
|
||||
from unittest.mock import MagicMock
|
||||
import asyncio
|
||||
|
||||
# Import the specific components we need to test
|
||||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
||||
InitPassThroughEndpointHelpers,
|
||||
_registered_pass_through_routes,
|
||||
)
|
||||
|
||||
|
||||
def test_update_pass_through_route_updates_registry():
|
||||
"""
|
||||
REGRESSION TEST: Verify that calling add_exact_path_route (or add_subpath_route)
|
||||
on an EXISTING route correctly updates the in-memory registry.
|
||||
"""
|
||||
|
||||
async def _async_test():
|
||||
# Setup - Unique IDs to avoid collision with other tests
|
||||
endpoint_id = "regression-test-endpoint"
|
||||
path = "/regression-test-path"
|
||||
route_key = f"{endpoint_id}:exact:{path}"
|
||||
target = "http://example.com"
|
||||
|
||||
# Cleanup: Ensure clean state before test
|
||||
if route_key in _registered_pass_through_routes:
|
||||
del _registered_pass_through_routes[route_key]
|
||||
|
||||
try:
|
||||
# 1. First Registration (Initial State)
|
||||
InitPassThroughEndpointHelpers.add_exact_path_route(
|
||||
app=MagicMock(),
|
||||
path=path,
|
||||
target=target,
|
||||
custom_headers={"Authorization": "Bearer INITIAL_TOKEN"},
|
||||
forward_headers=False,
|
||||
merge_query_params=False,
|
||||
dependencies=[],
|
||||
cost_per_request=0,
|
||||
endpoint_id=endpoint_id,
|
||||
)
|
||||
|
||||
# Verify Initial State
|
||||
assert route_key in _registered_pass_through_routes
|
||||
initial_headers = _registered_pass_through_routes[route_key][
|
||||
"passthrough_params"
|
||||
]["custom_headers"]
|
||||
assert initial_headers["Authorization"] == "Bearer INITIAL_TOKEN"
|
||||
|
||||
# 2. Perform Update (Simulate API Update)
|
||||
# This call should overwrite the existing entry
|
||||
InitPassThroughEndpointHelpers.add_exact_path_route(
|
||||
app=MagicMock(),
|
||||
path=path,
|
||||
target=target,
|
||||
custom_headers={
|
||||
"Authorization": "Bearer NEW_UPDATED_TOKEN"
|
||||
}, # Changed Header
|
||||
forward_headers=False,
|
||||
merge_query_params=False,
|
||||
dependencies=[],
|
||||
cost_per_request=0,
|
||||
endpoint_id=endpoint_id,
|
||||
)
|
||||
|
||||
# 3. Verify Update Occurred
|
||||
updated_headers = _registered_pass_through_routes[route_key][
|
||||
"passthrough_params"
|
||||
]["custom_headers"]
|
||||
|
||||
# This assertion protects against the regression
|
||||
assert (
|
||||
updated_headers["Authorization"] == "Bearer NEW_UPDATED_TOKEN"
|
||||
), "Registry failed to update! Old headers persisted despite update call."
|
||||
|
||||
finally:
|
||||
# Cleanup: Remove test entry
|
||||
if route_key in _registered_pass_through_routes:
|
||||
del _registered_pass_through_routes[route_key]
|
||||
|
||||
asyncio.run(_async_test())
|
||||
|
||||
|
||||
def test_update_subpath_route_updates_registry():
|
||||
"""
|
||||
REGRESSION TEST: Verify that calling add_subpath_route
|
||||
on an EXISTING route correctly updates the in-memory registry.
|
||||
"""
|
||||
|
||||
async def _async_test():
|
||||
# Setup
|
||||
endpoint_id = "regression-test-subpath"
|
||||
path = "/regression-test-wildcard"
|
||||
route_key = f"{endpoint_id}:subpath:{path}"
|
||||
target = "http://example.com"
|
||||
|
||||
if route_key in _registered_pass_through_routes:
|
||||
del _registered_pass_through_routes[route_key]
|
||||
|
||||
try:
|
||||
# 1. First Registration
|
||||
InitPassThroughEndpointHelpers.add_subpath_route(
|
||||
app=MagicMock(),
|
||||
path=path,
|
||||
target=target,
|
||||
custom_headers={"Authorization": "Bearer INITIAL_SUBPATH_TOKEN"},
|
||||
forward_headers=False,
|
||||
merge_query_params=False,
|
||||
dependencies=[],
|
||||
cost_per_request=0,
|
||||
endpoint_id=endpoint_id,
|
||||
)
|
||||
|
||||
assert (
|
||||
_registered_pass_through_routes[route_key]["passthrough_params"][
|
||||
"custom_headers"
|
||||
]["Authorization"]
|
||||
== "Bearer INITIAL_SUBPATH_TOKEN"
|
||||
)
|
||||
|
||||
# 2. Update
|
||||
InitPassThroughEndpointHelpers.add_subpath_route(
|
||||
app=MagicMock(),
|
||||
path=path,
|
||||
target=target,
|
||||
custom_headers={"Authorization": "Bearer NEW_SUBPATH_TOKEN"},
|
||||
forward_headers=False,
|
||||
merge_query_params=False,
|
||||
dependencies=[],
|
||||
cost_per_request=0,
|
||||
endpoint_id=endpoint_id,
|
||||
)
|
||||
|
||||
# 3. Verify
|
||||
updated_headers = _registered_pass_through_routes[route_key][
|
||||
"passthrough_params"
|
||||
]["custom_headers"]
|
||||
assert (
|
||||
updated_headers["Authorization"] == "Bearer NEW_SUBPATH_TOKEN"
|
||||
), "Subpath registry failed to update!"
|
||||
|
||||
finally:
|
||||
if route_key in _registered_pass_through_routes:
|
||||
del _registered_pass_through_routes[route_key]
|
||||
|
||||
asyncio.run(_async_test())
|
||||
@@ -1,5 +1,6 @@
|
||||
import os
|
||||
import sys
|
||||
from copy import deepcopy
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
@@ -191,12 +192,12 @@ def test_o_series_model_detection():
|
||||
config = AzureOpenAIOSeriesResponsesAPIConfig()
|
||||
|
||||
# Test explicit o_series naming
|
||||
assert config.is_o_series_model("o_series/gpt-o1") == True
|
||||
assert config.is_o_series_model("azure/o_series/gpt-o3") == True
|
||||
assert config.is_o_series_model("o_series/gpt-o1")
|
||||
assert config.is_o_series_model("azure/o_series/gpt-o3")
|
||||
|
||||
# Test regular models
|
||||
assert config.is_o_series_model("gpt-4o") == False
|
||||
assert config.is_o_series_model("gpt-3.5-turbo") == False
|
||||
assert not config.is_o_series_model("gpt-4o")
|
||||
assert not config.is_o_series_model("gpt-3.5-turbo")
|
||||
|
||||
|
||||
@pytest.mark.serial
|
||||
@@ -297,19 +298,19 @@ class TestAzureResponsesAPIConfig:
|
||||
def test_azure_cancel_response_api_request(self):
|
||||
"""Test Azure cancel response API request transformation"""
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
|
||||
response_id = "resp_test123"
|
||||
api_base = "https://test.openai.azure.com/openai/responses?api-version=2024-05-01-preview"
|
||||
litellm_params = GenericLiteLLMParams(api_version="2024-05-01-preview")
|
||||
headers = {"Authorization": "Bearer test-key"}
|
||||
|
||||
|
||||
url, data = self.config.transform_cancel_response_api_request(
|
||||
response_id=response_id,
|
||||
api_base=api_base,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
|
||||
expected_url = "https://test.openai.azure.com/openai/responses/resp_test123/cancel?api-version=2024-05-01-preview"
|
||||
assert url == expected_url
|
||||
assert data == {}
|
||||
@@ -318,7 +319,7 @@ class TestAzureResponsesAPIConfig:
|
||||
"""Test Azure cancel response API response transformation"""
|
||||
from unittest.mock import Mock
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
|
||||
|
||||
# Mock response
|
||||
mock_response = Mock()
|
||||
mock_response.json.return_value = {
|
||||
@@ -330,18 +331,164 @@ class TestAzureResponsesAPIConfig:
|
||||
"tool_choice": "auto",
|
||||
"tools": [],
|
||||
"top_p": 1.0,
|
||||
"status": "cancelled"
|
||||
"status": "cancelled",
|
||||
}
|
||||
mock_response.text = "test response"
|
||||
mock_response.status_code = 200
|
||||
|
||||
|
||||
# Mock logging object
|
||||
mock_logging_obj = Mock()
|
||||
|
||||
|
||||
result = self.config.transform_cancel_response_api_response(
|
||||
raw_response=mock_response,
|
||||
logging_obj=mock_logging_obj,
|
||||
)
|
||||
|
||||
|
||||
assert isinstance(result, ResponsesAPIResponse)
|
||||
assert result.id == "resp_test123"
|
||||
assert result.id == "resp_test123"
|
||||
|
||||
def test_azure_responses_api_tool_flattening_nested_to_flat(self):
|
||||
"""Test that nested tools are flattened correctly"""
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
# Setup
|
||||
nested_tools = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"description": "Get weather for a location",
|
||||
"parameters": {"type": "object", "properties": {}},
|
||||
},
|
||||
}
|
||||
]
|
||||
|
||||
response_api_params = {"tools": nested_tools}
|
||||
litellm_params = GenericLiteLLMParams()
|
||||
|
||||
# Execute
|
||||
self.config.transform_responses_api_request(
|
||||
model=self.model,
|
||||
input="test input",
|
||||
response_api_optional_request_params=response_api_params,
|
||||
litellm_params=litellm_params,
|
||||
headers={},
|
||||
)
|
||||
|
||||
# Verify
|
||||
expected_tools = [
|
||||
{
|
||||
"type": "function",
|
||||
"name": "get_weather",
|
||||
"description": "Get weather for a location",
|
||||
"parameters": {"type": "object", "properties": {}},
|
||||
}
|
||||
]
|
||||
assert response_api_params["tools"] == expected_tools
|
||||
|
||||
def test_azure_responses_api_tool_flattening_already_flat(self):
|
||||
"""Test that already flat tools are passed through unchanged"""
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
# Setup
|
||||
flat_tools = [
|
||||
{
|
||||
"type": "function",
|
||||
"name": "get_weather",
|
||||
"description": "Get weather for a location",
|
||||
"parameters": {"type": "object", "properties": {}},
|
||||
}
|
||||
]
|
||||
|
||||
# Make a copy to check it doesn't change
|
||||
response_api_params = {"tools": list(flat_tools)}
|
||||
litellm_params = GenericLiteLLMParams()
|
||||
|
||||
# Execute
|
||||
self.config.transform_responses_api_request(
|
||||
model=self.model,
|
||||
input="test input",
|
||||
response_api_optional_request_params=response_api_params,
|
||||
litellm_params=litellm_params,
|
||||
headers={},
|
||||
)
|
||||
|
||||
# Verify
|
||||
assert response_api_params["tools"] == flat_tools
|
||||
|
||||
def test_azure_responses_api_tool_flattening_preserves_original(self):
|
||||
"""Test that the original tool dictionary is not mutated"""
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
# Setup
|
||||
original_tool = {
|
||||
"type": "function",
|
||||
"function": {"name": "get_weather", "parameters": {}},
|
||||
}
|
||||
original_tool_copy = deepcopy(original_tool)
|
||||
|
||||
response_api_params = {"tools": [original_tool]}
|
||||
litellm_params = GenericLiteLLMParams()
|
||||
|
||||
# Execute
|
||||
self.config.transform_responses_api_request(
|
||||
model=self.model,
|
||||
input="test input",
|
||||
response_api_optional_request_params=response_api_params,
|
||||
litellm_params=litellm_params,
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert original_tool == original_tool_copy
|
||||
|
||||
def test_azure_responses_api_tool_flattening_mixed_tools(self):
|
||||
"""Test mixed nested and flat tools"""
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
# Setup
|
||||
nested_tool = {
|
||||
"type": "function",
|
||||
"function": {"name": "nested", "parameters": {}},
|
||||
}
|
||||
flat_tool = {"type": "function", "name": "flat", "parameters": {}}
|
||||
|
||||
response_api_params = {"tools": [nested_tool, flat_tool]}
|
||||
litellm_params = GenericLiteLLMParams()
|
||||
|
||||
# Execute
|
||||
self.config.transform_responses_api_request(
|
||||
model=self.model,
|
||||
input="test input",
|
||||
response_api_optional_request_params=response_api_params,
|
||||
litellm_params=litellm_params,
|
||||
headers={},
|
||||
)
|
||||
|
||||
# Verify
|
||||
assert len(response_api_params["tools"]) == 2
|
||||
|
||||
# First tool should be flattened
|
||||
assert "function" not in response_api_params["tools"][0]
|
||||
assert response_api_params["tools"][0]["name"] == "nested"
|
||||
|
||||
# Second tool should remain as is
|
||||
assert response_api_params["tools"][1] == flat_tool
|
||||
|
||||
def test_azure_responses_api_tool_flattening_no_tools(self):
|
||||
"""Test handling when no tools are present"""
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
# Setup
|
||||
response_api_params = {}
|
||||
litellm_params = GenericLiteLLMParams()
|
||||
|
||||
# Execute - should not crash
|
||||
self.config.transform_responses_api_request(
|
||||
model=self.model,
|
||||
input="test input",
|
||||
response_api_optional_request_params=response_api_params,
|
||||
litellm_params=litellm_params,
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert "tools" not in response_api_params
|
||||
|
||||
@@ -0,0 +1,114 @@
|
||||
"""
|
||||
Test that Bedrock streaming responses always use choice index 0,
|
||||
regardless of contentBlockIndex value.
|
||||
|
||||
Bedrock's contentBlockIndex identifies content blocks within a message (e.g.,
|
||||
text=0, toolUse=1), NOT parallel completions. Since Bedrock doesn't support
|
||||
n > 1, all chunks must use choice index 0.
|
||||
|
||||
References:
|
||||
- Bedrock InferenceConfiguration (no n parameter):
|
||||
https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_InferenceConfiguration.html
|
||||
- OpenAI choice.index (for n > 1):
|
||||
https://platform.openai.com/docs/api-reference/chat/object
|
||||
"""
|
||||
|
||||
from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder
|
||||
|
||||
|
||||
class TestBedrockStreamingChoiceIndex:
|
||||
"""Test that all streaming chunks use choice index 0."""
|
||||
|
||||
def test_tool_call_chunk_uses_choice_index_zero(self):
|
||||
"""
|
||||
Core regression test: tool call chunks must use choice index 0,
|
||||
not contentBlockIndex (which is 1 for tool calls).
|
||||
|
||||
This was the bug - contentBlockIndex was incorrectly used as choice.index,
|
||||
breaking OpenAI SDK's ChatCompletionAccumulator.
|
||||
"""
|
||||
handler = AWSEventStreamDecoder(model="anthropic.claude-3-sonnet-20240229-v1:0")
|
||||
|
||||
# First, simulate a tool use start event on contentBlockIndex 1
|
||||
start_chunk = {
|
||||
"start": {
|
||||
"toolUse": {
|
||||
"toolUseId": "tooluse_abc123",
|
||||
"name": "get_weather",
|
||||
}
|
||||
},
|
||||
"contentBlockIndex": 1, # Tool calls are on index 1
|
||||
}
|
||||
|
||||
start_result = handler.converse_chunk_parser(start_chunk)
|
||||
|
||||
# Choice index should be 0, NOT contentBlockIndex (1)
|
||||
assert start_result.choices[0].index == 0
|
||||
assert start_result.choices[0].delta.tool_calls is not None
|
||||
assert start_result.choices[0].delta.tool_calls[0]["id"] == "tooluse_abc123"
|
||||
|
||||
# Now simulate tool use delta on contentBlockIndex 1
|
||||
delta_chunk = {
|
||||
"delta": {
|
||||
"toolUse": {
|
||||
"input": '{"location": "San Francisco"}'
|
||||
}
|
||||
},
|
||||
"contentBlockIndex": 1, # Tool calls are on index 1
|
||||
}
|
||||
|
||||
delta_result = handler.converse_chunk_parser(delta_chunk)
|
||||
|
||||
# Choice index should still be 0, NOT contentBlockIndex (1)
|
||||
assert delta_result.choices[0].index == 0
|
||||
assert delta_result.choices[0].delta.tool_calls is not None
|
||||
assert delta_result.choices[0].delta.tool_calls[0]["function"]["arguments"] == '{"location": "San Francisco"}'
|
||||
|
||||
def test_mixed_content_blocks_all_use_choice_index_zero(self):
|
||||
"""
|
||||
Integration test simulating a realistic streaming session:
|
||||
text (contentBlockIndex=0) → tool call (contentBlockIndex=1) → finish.
|
||||
|
||||
All chunks must have choice.index=0 for OpenAI SDK compatibility.
|
||||
"""
|
||||
handler = AWSEventStreamDecoder(model="anthropic.claude-3-sonnet-20240229-v1:0")
|
||||
|
||||
# Chunk 1: Text on contentBlockIndex 0
|
||||
text_chunk = {
|
||||
"delta": {"text": "Let me check the weather."},
|
||||
"contentBlockIndex": 0,
|
||||
}
|
||||
result1 = handler.converse_chunk_parser(text_chunk)
|
||||
assert result1.choices[0].index == 0, "Text chunk should have index=0"
|
||||
|
||||
# Chunk 2: Tool call start on contentBlockIndex 1
|
||||
tool_start_chunk = {
|
||||
"start": {
|
||||
"toolUse": {
|
||||
"toolUseId": "tool_xyz",
|
||||
"name": "get_weather",
|
||||
}
|
||||
},
|
||||
"contentBlockIndex": 1,
|
||||
}
|
||||
result2 = handler.converse_chunk_parser(tool_start_chunk)
|
||||
assert result2.choices[0].index == 0, "Tool start should have index=0, not contentBlockIndex=1"
|
||||
|
||||
# Chunk 3: Tool call delta on contentBlockIndex 1
|
||||
tool_delta_chunk = {
|
||||
"delta": {
|
||||
"toolUse": {
|
||||
"input": '{"city": "NYC"}'
|
||||
}
|
||||
},
|
||||
"contentBlockIndex": 1,
|
||||
}
|
||||
result3 = handler.converse_chunk_parser(tool_delta_chunk)
|
||||
assert result3.choices[0].index == 0, "Tool delta should have index=0, not contentBlockIndex=1"
|
||||
|
||||
# Chunk 4: Finish reason
|
||||
finish_chunk = {
|
||||
"stopReason": "tool_use",
|
||||
}
|
||||
result4 = handler.converse_chunk_parser(finish_chunk)
|
||||
assert result4.choices[0].index == 0, "Finish reason should have index=0"
|
||||
@@ -1,3 +1,4 @@
|
||||
import json
|
||||
import importlib
|
||||
import logging
|
||||
import os
|
||||
@@ -989,6 +990,68 @@ class TestMCPServerManager:
|
||||
assert result.status == "healthy"
|
||||
assert result.health_check_error is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_register_openapi_tools_includes_static_headers(self, tmp_path):
|
||||
"""Ensure OpenAPI-to-MCP tool calls include server.static_headers (Issue #19341)."""
|
||||
manager = MCPServerManager()
|
||||
|
||||
spec_path = tmp_path / "openapi.json"
|
||||
spec_path.write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"openapi": "3.0.0",
|
||||
"info": {"title": "Demo", "version": "1.0.0"},
|
||||
"paths": {
|
||||
"/health": {
|
||||
"get": {
|
||||
"operationId": "health_check",
|
||||
"summary": "health",
|
||||
}
|
||||
}
|
||||
},
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
server = MCPServer(
|
||||
server_id="openapi-server",
|
||||
name="openapi-server",
|
||||
server_name="openapi-server",
|
||||
url="https://example.com",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.none,
|
||||
static_headers={"Authorization": "STATIC token"},
|
||||
)
|
||||
|
||||
captured: dict = {}
|
||||
|
||||
def fake_create_tool_function(path, method, operation, base_url, headers=None):
|
||||
captured["headers"] = headers
|
||||
|
||||
async def tool_func(**kwargs):
|
||||
return "ok"
|
||||
|
||||
return tool_func
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.create_tool_function",
|
||||
side_effect=fake_create_tool_function,
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.build_input_schema",
|
||||
return_value={"type": "object", "properties": {}, "required": []},
|
||||
), patch(
|
||||
"litellm.proxy._experimental.mcp_server.tool_registry.global_mcp_tool_registry.register_tool",
|
||||
return_value=None,
|
||||
):
|
||||
manager._register_openapi_tools(
|
||||
spec_path=str(spec_path),
|
||||
server=server,
|
||||
base_url="https://example.com",
|
||||
)
|
||||
|
||||
assert captured["headers"] is not None
|
||||
assert captured["headers"]["Authorization"] == "STATIC token"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_call_tool_check_allowed_tools_list_allows_tool(self):
|
||||
"""Test pre_call_tool_check allows tool when it's in allowed_tools list"""
|
||||
|
||||
@@ -101,6 +101,59 @@ class TestExecuteWithMcpClient:
|
||||
assert result["status"] == "error"
|
||||
assert "stack_trace" not in result
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_forwards_static_headers(self, monkeypatch):
|
||||
"""Ensure static_headers are forwarded to the MCP client during test calls.
|
||||
|
||||
This is required for `/mcp-rest/test/tools/list` (Issue #19341), where the UI
|
||||
sends `static_headers` but the backend must forward them during
|
||||
`session.initialize()` and tool discovery.
|
||||
"""
|
||||
captured: dict = {}
|
||||
|
||||
def fake_build_stdio_env(server, raw_headers):
|
||||
return None
|
||||
|
||||
def fake_create_client(*args, **kwargs):
|
||||
captured["extra_headers"] = kwargs.get("extra_headers")
|
||||
return object()
|
||||
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"_build_stdio_env",
|
||||
fake_build_stdio_env,
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"_create_mcp_client",
|
||||
fake_create_client,
|
||||
raising=False,
|
||||
)
|
||||
|
||||
async def ok_operation(client):
|
||||
return {"status": "ok"}
|
||||
|
||||
payload = NewMCPServerRequest(
|
||||
server_name="example",
|
||||
url="https://example.com",
|
||||
auth_type=MCPAuth.none,
|
||||
static_headers={"Authorization": "STATIC token"},
|
||||
)
|
||||
|
||||
result = await rest_endpoints._execute_with_mcp_client(
|
||||
payload,
|
||||
ok_operation,
|
||||
oauth2_headers={"X-OAuth": "1"},
|
||||
raw_headers={"x-test": "y"},
|
||||
)
|
||||
|
||||
assert result["status"] == "ok"
|
||||
assert captured["extra_headers"] == {
|
||||
"X-OAuth": "1",
|
||||
"Authorization": "STATIC token",
|
||||
}
|
||||
|
||||
|
||||
class TestTestConnection:
|
||||
def test_requires_auth_dependency(self):
|
||||
|
||||
@@ -7,7 +7,6 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi import Request, UploadFile
|
||||
from fastapi.testclient import TestClient
|
||||
from starlette.datastructures import Headers, QueryParams
|
||||
from starlette.datastructures import UploadFile as StarletteUploadFile
|
||||
|
||||
@@ -201,7 +200,6 @@ async def test_pass_through_request_failure_handler():
|
||||
|
||||
Critical Test: When a users pass through endpoint request fails, we must log the failure code, exception in litellm spend logs.
|
||||
"""
|
||||
print("running test_pass_through_request_failure_handler")
|
||||
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client"
|
||||
@@ -266,27 +264,27 @@ def test_is_langfuse_route():
|
||||
# Test positive cases
|
||||
assert (
|
||||
handler.is_langfuse_route("http://localhost:4000/langfuse/api/public/traces")
|
||||
== True
|
||||
is True
|
||||
)
|
||||
assert (
|
||||
handler.is_langfuse_route(
|
||||
"https://proxy.example.com/langfuse/api/public/sessions"
|
||||
)
|
||||
== True
|
||||
is True
|
||||
)
|
||||
assert handler.is_langfuse_route("/langfuse/api/public/ingestion") == True
|
||||
assert handler.is_langfuse_route("http://localhost:4000/langfuse/") == True
|
||||
assert handler.is_langfuse_route("/langfuse/api/public/ingestion") is True
|
||||
assert handler.is_langfuse_route("http://localhost:4000/langfuse/") is True
|
||||
|
||||
# Test negative cases
|
||||
assert (
|
||||
handler.is_langfuse_route("https://api.openai.com/v1/chat/completions") == False
|
||||
handler.is_langfuse_route("https://api.openai.com/v1/chat/completions") is False
|
||||
)
|
||||
assert (
|
||||
handler.is_langfuse_route("http://localhost:4000/anthropic/v1/messages")
|
||||
== False
|
||||
is False
|
||||
)
|
||||
assert handler.is_langfuse_route("https://example.com/other") == False
|
||||
assert handler.is_langfuse_route("") == False
|
||||
assert handler.is_langfuse_route("https://example.com/other") is False
|
||||
assert handler.is_langfuse_route("") is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -576,7 +574,6 @@ def test_set_cost_per_request():
|
||||
"""
|
||||
Test that _set_cost_per_request correctly sets the cost in logging object and kwargs
|
||||
"""
|
||||
from datetime import datetime
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
|
||||
@@ -687,7 +684,7 @@ async def test_pass_through_success_handler_with_cost_per_request():
|
||||
end_time = datetime.now()
|
||||
|
||||
# Call the success handler
|
||||
result = await handler.pass_through_async_success_handler(
|
||||
await handler.pass_through_async_success_handler(
|
||||
httpx_response=mock_response,
|
||||
response_body={"status": "success", "data": "test"},
|
||||
logging_obj=mock_logging_obj,
|
||||
@@ -719,8 +716,9 @@ async def test_create_pass_through_route_with_cost_per_request():
|
||||
)
|
||||
|
||||
# Create the endpoint function with cost_per_request
|
||||
unique_path = "/test/path/unique/cost_per_request"
|
||||
endpoint_func = create_pass_through_route(
|
||||
endpoint="/test/path",
|
||||
endpoint=unique_path,
|
||||
target="http://example.com",
|
||||
custom_headers={},
|
||||
_forward_headers=True,
|
||||
@@ -732,11 +730,19 @@ async def test_create_pass_through_route_with_cost_per_request():
|
||||
# Mock the pass_through_request function to capture its call
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_request"
|
||||
) as mock_pass_through:
|
||||
) as mock_pass_through, patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.InitPassThroughEndpointHelpers.is_registered_pass_through_route"
|
||||
) as mock_is_registered, patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.InitPassThroughEndpointHelpers.get_registered_pass_through_route"
|
||||
) as mock_get_registered:
|
||||
mock_pass_through.return_value = MagicMock()
|
||||
mock_is_registered.return_value = True
|
||||
mock_get_registered.return_value = None
|
||||
|
||||
# Create mock request
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.url = MagicMock()
|
||||
mock_request.url.path = unique_path
|
||||
mock_request.path_params = {}
|
||||
mock_request.query_params = QueryParams({})
|
||||
|
||||
@@ -817,7 +823,7 @@ def test_initialize_pass_through_endpoints_with_cost_per_request():
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pass_through_request_contains_proxy_server_request_in_kwargs():
|
||||
async def test_pass_through_request_contains_proxy_server_request_in_kwargs(): # noqa: PLR0915
|
||||
"""
|
||||
Test that pass_through_request (parent method) correctly includes proxy_server_request
|
||||
in kwargs passed to the success handler.
|
||||
@@ -825,8 +831,6 @@ async def test_pass_through_request_contains_proxy_server_request_in_kwargs():
|
||||
Critical Test: Ensures that when pass_through_request is called, the kwargs passed to
|
||||
downstream methods contain the proxy server request details (url, method, body).
|
||||
"""
|
||||
print("running test_pass_through_request_contains_proxy_server_request_in_kwargs")
|
||||
|
||||
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.HttpPassThroughEndpointHelpers.non_streaming_http_request_handler"
|
||||
@@ -891,7 +895,7 @@ async def test_pass_through_request_contains_proxy_server_request_in_kwargs():
|
||||
mock_user_api_key_dict.request_route = "/api/endpoint"
|
||||
|
||||
# Call pass_through_request (the parent method)
|
||||
result = await pass_through_request(
|
||||
await pass_through_request(
|
||||
request=mock_request,
|
||||
target="http://target-api.com/endpoint",
|
||||
custom_headers={"X-Custom": "header"},
|
||||
@@ -951,7 +955,6 @@ async def test_create_pass_through_endpoint():
|
||||
"""
|
||||
from litellm.proxy._types import (
|
||||
ConfigFieldInfo,
|
||||
ConfigFieldUpdate,
|
||||
PassThroughEndpointResponse,
|
||||
PassThroughGenericEndpoint,
|
||||
UserAPIKeyAuth,
|
||||
@@ -986,7 +989,9 @@ async def test_create_pass_through_endpoint():
|
||||
|
||||
# Call the create function
|
||||
result = await create_pass_through_endpoints(
|
||||
data=test_endpoint, user_api_key_dict=mock_user_api_key_dict
|
||||
data=test_endpoint,
|
||||
request=MagicMock(spec=Request),
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
)
|
||||
|
||||
# Verify the result
|
||||
@@ -1029,7 +1034,6 @@ async def test_update_pass_through_endpoint():
|
||||
"""
|
||||
from litellm.proxy._types import (
|
||||
ConfigFieldInfo,
|
||||
ConfigFieldUpdate,
|
||||
PassThroughEndpointResponse,
|
||||
PassThroughGenericEndpoint,
|
||||
UserAPIKeyAuth,
|
||||
@@ -1082,6 +1086,7 @@ async def test_update_pass_through_endpoint():
|
||||
result = await update_pass_through_endpoints(
|
||||
endpoint_id=existing_endpoint_id,
|
||||
data=update_data,
|
||||
request=MagicMock(spec=Request),
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
)
|
||||
|
||||
@@ -1165,6 +1170,7 @@ async def test_update_pass_through_endpoint_not_found():
|
||||
await update_pass_through_endpoints(
|
||||
endpoint_id="non-existent-endpoint-123",
|
||||
data=update_data,
|
||||
request=MagicMock(spec=Request),
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
)
|
||||
|
||||
@@ -1185,7 +1191,6 @@ async def test_delete_pass_through_endpoint():
|
||||
"""
|
||||
from litellm.proxy._types import (
|
||||
ConfigFieldInfo,
|
||||
ConfigFieldUpdate,
|
||||
PassThroughEndpointResponse,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
@@ -1421,7 +1426,7 @@ async def test_pass_through_request_query_params_forwarding():
|
||||
mock_user_api_key_dict.api_key = "sk-1234"
|
||||
|
||||
# Call pass_through_request
|
||||
result = await pass_through_request(
|
||||
await pass_through_request(
|
||||
request=mock_request,
|
||||
target="https://krris-m2f9a9i7-eastus2.openai.azure.com/openai/assistants",
|
||||
custom_headers={"Authorization": "Bearer azure_token"},
|
||||
@@ -1498,7 +1503,6 @@ async def test_pass_through_with_httpbin_redirect():
|
||||
|
||||
# httpbin.org/get returns JSON with info about the request
|
||||
assert '"url": "https://httpbin.org/get"' in response_content
|
||||
print("GOT A Response from HTTPBIN=", response_content)
|
||||
except Exception as e:
|
||||
# If httpbin.org is not accessible, skip the test
|
||||
import pytest
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import datetime as real_datetime
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
@@ -132,3 +133,48 @@ def test_join_paths_nested_path():
|
||||
"""Test path joining with nested paths"""
|
||||
result = join_paths(base_path="http://0.0.0.0:4000/v1", route="chat/completions")
|
||||
assert result == "http://0.0.0.0:4000/v1/chat/completions"
|
||||
|
||||
|
||||
def _patch_today(monkeypatch, year, month, day):
|
||||
class PatchedDate(real_datetime.date):
|
||||
@classmethod
|
||||
def today(cls):
|
||||
return real_datetime.date(year, month, day)
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.utils.date", PatchedDate)
|
||||
|
||||
|
||||
def test_get_projected_spend_over_limit_day_one(monkeypatch):
|
||||
from litellm.proxy.utils import _get_projected_spend_over_limit
|
||||
|
||||
_patch_today(monkeypatch, 2026, 1, 1)
|
||||
result = _get_projected_spend_over_limit(100.0, 1.0)
|
||||
|
||||
assert result is not None
|
||||
projected_spend, projected_exceeded_date = result
|
||||
assert projected_spend == 3100.0
|
||||
assert projected_exceeded_date == real_datetime.date(2026, 1, 1)
|
||||
|
||||
|
||||
def test_get_projected_spend_over_limit_december(monkeypatch):
|
||||
from litellm.proxy.utils import _get_projected_spend_over_limit
|
||||
|
||||
_patch_today(monkeypatch, 2026, 12, 15)
|
||||
result = _get_projected_spend_over_limit(100.0, 1.0)
|
||||
|
||||
assert result is not None
|
||||
projected_spend, projected_exceeded_date = result
|
||||
assert projected_spend == pytest.approx(214.28571428571428)
|
||||
assert projected_exceeded_date == real_datetime.date(2026, 12, 15)
|
||||
|
||||
|
||||
def test_get_projected_spend_over_limit_includes_current_spend(monkeypatch):
|
||||
from litellm.proxy.utils import _get_projected_spend_over_limit
|
||||
|
||||
_patch_today(monkeypatch, 2026, 4, 11)
|
||||
result = _get_projected_spend_over_limit(100.0, 200.0)
|
||||
|
||||
assert result is not None
|
||||
projected_spend, projected_exceeded_date = result
|
||||
assert projected_spend == 290.0
|
||||
assert projected_exceeded_date == real_datetime.date(2026, 4, 21)
|
||||
|
||||
@@ -803,6 +803,79 @@ def test_gemini_25_explicit_caching_cost_direct_usage():
|
||||
assert expected_actual_cost == total_cost
|
||||
|
||||
|
||||
def test_azure_ai_cache_cost_calculation():
|
||||
"""
|
||||
Test that azure_ai provider correctly calculates cache costs using generic_cost_per_token.
|
||||
|
||||
This verifies that azure_ai models with custom cache pricing in model_info
|
||||
will have their cache_creation_input_token_cost and cache_read_input_token_cost
|
||||
applied correctly.
|
||||
"""
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token
|
||||
from litellm.types.utils import (
|
||||
PromptTokensDetailsWrapper,
|
||||
Usage,
|
||||
)
|
||||
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
|
||||
# Register a custom azure_ai model with cache pricing
|
||||
test_model_id = "test-azure-ai-claude-model"
|
||||
litellm.register_model(
|
||||
model_cost={
|
||||
test_model_id: {
|
||||
"input_cost_per_token": 5.0e-06,
|
||||
"output_cost_per_token": 2.5e-05,
|
||||
"cache_creation_input_token_cost": 6.25e-06,
|
||||
"cache_read_input_token_cost": 5.0e-07,
|
||||
"litellm_provider": "azure_ai",
|
||||
"max_tokens": 200000,
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
# Create usage with cache tokens
|
||||
usage = Usage(
|
||||
completion_tokens=100,
|
||||
prompt_tokens=1000,
|
||||
total_tokens=1100,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
cached_tokens=800, # 800 cache read tokens
|
||||
text_tokens=100, # 100 regular text tokens
|
||||
),
|
||||
cache_creation_input_tokens=100, # 100 cache creation tokens
|
||||
)
|
||||
|
||||
input_cost, output_cost = generic_cost_per_token(
|
||||
model=test_model_id,
|
||||
usage=usage,
|
||||
custom_llm_provider="azure_ai",
|
||||
)
|
||||
|
||||
total_cost = input_cost + output_cost
|
||||
|
||||
# Calculate expected cost manually
|
||||
model_info = litellm.model_cost[test_model_id]
|
||||
expected_input_cost = (
|
||||
model_info["input_cost_per_token"] * 100 # text tokens
|
||||
+ model_info["cache_read_input_token_cost"] * 800 # cached tokens
|
||||
+ model_info["cache_creation_input_token_cost"] * 100 # cache creation tokens
|
||||
)
|
||||
expected_output_cost = model_info["output_cost_per_token"] * 100
|
||||
|
||||
print(f"Input cost: {input_cost}, Expected: {expected_input_cost}")
|
||||
print(f"Output cost: {output_cost}, Expected: {expected_output_cost}")
|
||||
print(f"Total cost: {total_cost}")
|
||||
|
||||
assert abs(input_cost - expected_input_cost) < 1e-10, (
|
||||
f"Input cost mismatch: got {input_cost}, expected {expected_input_cost}"
|
||||
)
|
||||
assert abs(output_cost - expected_output_cost) < 1e-10, (
|
||||
f"Output cost mismatch: got {output_cost}, expected {expected_output_cost}"
|
||||
)
|
||||
|
||||
|
||||
def test_cost_discount_vertex_ai():
|
||||
"""
|
||||
Test that cost discount is applied correctly for Vertex AI provider
|
||||
|
||||
Reference in New Issue
Block a user