mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_mcp_server_env_vars
# Conflicts: # tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py
This commit is contained in:
commit
aa78a90a58
45 changed files with 4304 additions and 147 deletions
|
|
@ -42,7 +42,7 @@ When you must use real LLM models to, for example, write e2e tests, write a QA r
|
|||
|
||||
If you're an internal contributor, when creating a new PR, the typical flow is to branch off litellm_internal_staging and create a branch prefixed with litellm_. Do not create a branch prefixed with claude/ and generally do not have / in your branch names
|
||||
|
||||
Do not add `Co-Authored-By: Claude` or any Claude attribution to commit messages. Never use a `claude/` prefix or put a `/` in a branch name. Do not add "Generated with Claude Code" (or any similar attribution) to PR descriptions. Do not create a new PR/branch off the existing PR to fix/add something that is related and could've just been committed directly to the existing PR's branch
|
||||
Do not add `Co-Authored-By: Claude` or any Claude attribution to commit messages. Never use a `claude/` prefix or put a `/` in a branch name. Do not add "Generated with Claude Code" (or any similar attribution) to PR descriptions or comments. Do not create a new PR/branch off the existing PR to fix/add something that is related and could've just been committed directly to the existing PR's branch
|
||||
|
||||
When working on a PR, keep the PR description in sync with new commits being made
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "oauth2_flow" TEXT;
|
||||
|
|
@ -327,6 +327,7 @@ model LiteLLM_MCPServerTable {
|
|||
authorization_url String?
|
||||
token_url String?
|
||||
registration_url String?
|
||||
oauth2_flow String?
|
||||
allow_all_keys Boolean @default(false)
|
||||
available_on_public_internet Boolean @default(true)
|
||||
delegate_auth_to_upstream Boolean @default(false)
|
||||
|
|
|
|||
|
|
@ -2425,12 +2425,11 @@ class BaseTokenUsageProcessor:
|
|||
if not attr.startswith("_") and not callable(
|
||||
getattr(usage.completion_tokens_details, attr)
|
||||
):
|
||||
current_val = getattr(
|
||||
combined.completion_tokens_details, attr, 0
|
||||
current_val = (
|
||||
getattr(combined.completion_tokens_details, attr, 0) or 0
|
||||
)
|
||||
new_val = getattr(usage.completion_tokens_details, attr, 0)
|
||||
|
||||
if new_val is not None and current_val is not None:
|
||||
new_val = getattr(usage.completion_tokens_details, attr, 0) or 0
|
||||
if isinstance(new_val, (int, float)):
|
||||
setattr(
|
||||
combined.completion_tokens_details,
|
||||
attr,
|
||||
|
|
|
|||
|
|
@ -1062,3 +1062,37 @@ class GuardrailInterventionNormalStringError(
|
|||
|
||||
def __repr__(self):
|
||||
return self.__str__()
|
||||
|
||||
|
||||
class SensitiveDataRouteException(Exception):
|
||||
"""
|
||||
Exception raised when a guardrail detects sensitive data and wants to reroute the request.
|
||||
|
||||
Instead of blocking the request, this exception signals that the request should be
|
||||
routed to a different model (typically an on-premise model for data privacy).
|
||||
|
||||
The proxy catches this exception and:
|
||||
1. Reroutes the current request to the specified model
|
||||
2. When sticky_session_routing is True, stores the routing decision in session
|
||||
cache so all subsequent requests in the same session are routed to the same model
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
route_to_model: str,
|
||||
session_id: str,
|
||||
guardrail_name: Optional[str] = None,
|
||||
detection_info: Optional[Dict[str, Any]] = None,
|
||||
message: Optional[str] = None,
|
||||
sticky_session_routing: bool = True,
|
||||
):
|
||||
self.route_to_model = route_to_model
|
||||
self.session_id = session_id
|
||||
self.guardrail_name = guardrail_name
|
||||
self.detection_info = detection_info or {}
|
||||
self.sticky_session_routing = sticky_session_routing
|
||||
self.message = (
|
||||
message
|
||||
or f"Sensitive data detected by {guardrail_name}. Routing to model: {route_to_model}"
|
||||
)
|
||||
super().__init__(self.message)
|
||||
|
|
|
|||
|
|
@ -47,9 +47,29 @@ from litellm.exceptions import (
|
|||
BlockedPiiEntityError,
|
||||
GuardrailRaisedException,
|
||||
ModifyResponseException,
|
||||
SensitiveDataRouteException,
|
||||
)
|
||||
|
||||
|
||||
def get_session_id_from_request_data(request_data: Dict[str, Any]) -> Optional[str]:
|
||||
"""Extract session_id from request data (litellm_session_id or metadata)."""
|
||||
session_id = request_data.get("litellm_session_id")
|
||||
if session_id:
|
||||
return str(session_id)
|
||||
|
||||
metadata = request_data.get("metadata") or {}
|
||||
session_id = metadata.get("session_id")
|
||||
if session_id:
|
||||
return str(session_id)
|
||||
|
||||
litellm_metadata = request_data.get("litellm_metadata") or {}
|
||||
session_id = litellm_metadata.get("session_id")
|
||||
if session_id:
|
||||
return str(session_id)
|
||||
|
||||
return None
|
||||
|
||||
|
||||
class CustomGuardrail(CustomLogger):
|
||||
# If True, during_call runs async_moderation_hook instead of the unified apply_guardrail path.
|
||||
use_native_during_call_hook: ClassVar[bool] = False
|
||||
|
|
@ -68,6 +88,9 @@ class CustomGuardrail(CustomLogger):
|
|||
end_session_after_n_fails: Optional[int] = None,
|
||||
on_violation: Optional[str] = None,
|
||||
realtime_violation_message: Optional[str] = None,
|
||||
on_sensitive_data: Optional[str] = None,
|
||||
sensitive_data_route_to_model: Optional[str] = None,
|
||||
sticky_session_routing: bool = True,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
|
|
@ -83,6 +106,9 @@ class CustomGuardrail(CustomLogger):
|
|||
end_session_after_n_fails: For /v1/realtime sessions, end the session after this many violations
|
||||
on_violation: For /v1/realtime sessions, 'warn' or 'end_session'
|
||||
realtime_violation_message: Message the bot speaks aloud when a /v1/realtime guardrail fires
|
||||
on_sensitive_data: Action when sensitive data is detected. 'block' (default) or 'route'
|
||||
sensitive_data_route_to_model: Model to route to when on_sensitive_data='route'
|
||||
sticky_session_routing: When True, all subsequent requests in the session use the same model
|
||||
"""
|
||||
self.guardrail_name = guardrail_name
|
||||
self.supported_event_hooks = supported_event_hooks
|
||||
|
|
@ -96,6 +122,11 @@ class CustomGuardrail(CustomLogger):
|
|||
self.end_session_after_n_fails: Optional[int] = end_session_after_n_fails
|
||||
self.on_violation: Optional[str] = on_violation
|
||||
self.realtime_violation_message: Optional[str] = realtime_violation_message
|
||||
self.on_sensitive_data: Optional[str] = on_sensitive_data
|
||||
self.sensitive_data_route_to_model: Optional[str] = (
|
||||
sensitive_data_route_to_model
|
||||
)
|
||||
self.sticky_session_routing: bool = sticky_session_routing
|
||||
|
||||
if supported_event_hooks:
|
||||
## validate event_hook is in supported_event_hooks
|
||||
|
|
@ -167,6 +198,108 @@ class CustomGuardrail(CustomLogger):
|
|||
detection_info=detection_info,
|
||||
)
|
||||
|
||||
def raise_sensitive_data_route_exception(
|
||||
self,
|
||||
route_to_model: str,
|
||||
request_data: Dict[str, Any],
|
||||
detection_info: Optional[Dict[str, Any]] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Raise an exception to reroute the request to a different model.
|
||||
|
||||
Use this when sensitive data is detected and the guardrail is configured
|
||||
to route to an on-premise model instead of blocking.
|
||||
|
||||
The exception will reroute this request to the specified model. When
|
||||
sticky_session_routing is enabled (the default), it also stores the
|
||||
routing decision so subsequent requests in this session reuse the model.
|
||||
|
||||
Args:
|
||||
route_to_model: The model to route this request (and session) to
|
||||
request_data: The original request data dictionary
|
||||
detection_info: Optional non-sensitive detection metadata (e.g. matched
|
||||
entity types, rule ids, scores). This is surfaced in request metadata
|
||||
and logs, so it must not contain the raw detected sensitive values.
|
||||
|
||||
Raises:
|
||||
SensitiveDataRouteException: Always raises to trigger rerouting
|
||||
"""
|
||||
session_id = self._get_session_id_from_request_data(request_data)
|
||||
if not session_id:
|
||||
raise ValueError(
|
||||
"Cannot route sensitive data without a session_id. "
|
||||
"Ensure the request includes a session_id in metadata or headers."
|
||||
)
|
||||
|
||||
raise SensitiveDataRouteException(
|
||||
route_to_model=route_to_model,
|
||||
session_id=session_id,
|
||||
guardrail_name=self.guardrail_name,
|
||||
detection_info=detection_info,
|
||||
sticky_session_routing=self.sticky_session_routing,
|
||||
)
|
||||
|
||||
def _get_session_id_from_request_data(
|
||||
self, request_data: Dict[str, Any]
|
||||
) -> Optional[str]:
|
||||
"""Extract session_id from request data."""
|
||||
return get_session_id_from_request_data(request_data)
|
||||
|
||||
def should_route_on_sensitive_data(self) -> bool:
|
||||
"""
|
||||
Returns True if this guardrail is configured to route requests
|
||||
to a different model when sensitive data is detected.
|
||||
"""
|
||||
return (
|
||||
self.on_sensitive_data == "route"
|
||||
and self.sensitive_data_route_to_model is not None
|
||||
)
|
||||
|
||||
def handle_sensitive_data_detection(
|
||||
self,
|
||||
request_data: Dict[str, Any],
|
||||
detection_info: Optional[Dict[str, Any]] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Handle sensitive data detection based on guardrail configuration.
|
||||
|
||||
If on_sensitive_data='route', raises SensitiveDataRouteException to reroute.
|
||||
Otherwise, raises GuardrailRaisedException to block. When routing is
|
||||
configured but the request carries no session_id, routing is not possible
|
||||
so the request falls back to a graceful block.
|
||||
|
||||
Args:
|
||||
request_data: The request data dictionary
|
||||
detection_info: Optional non-sensitive detection metadata. When routing,
|
||||
this is surfaced in request metadata and logs, so it must not contain
|
||||
the raw detected sensitive values.
|
||||
|
||||
Raises:
|
||||
SensitiveDataRouteException: When configured to route and a session_id is present
|
||||
GuardrailRaisedException: When configured to block, or when routing is
|
||||
configured but no session_id is available
|
||||
"""
|
||||
if self.should_route_on_sensitive_data():
|
||||
try:
|
||||
self.raise_sensitive_data_route_exception(
|
||||
route_to_model=self.sensitive_data_route_to_model, # type: ignore
|
||||
request_data=request_data,
|
||||
detection_info=detection_info,
|
||||
)
|
||||
except ValueError:
|
||||
raise GuardrailRaisedException(
|
||||
message=(
|
||||
f"Sensitive data detected by {self.guardrail_name} "
|
||||
"(routing skipped: request has no session_id)"
|
||||
),
|
||||
guardrail_name=self.guardrail_name,
|
||||
)
|
||||
else:
|
||||
raise GuardrailRaisedException(
|
||||
message=f"Sensitive data detected by {self.guardrail_name}",
|
||||
guardrail_name=self.guardrail_name,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def get_config_model() -> Optional[Type["GuardrailConfigModel"]]:
|
||||
"""
|
||||
|
|
@ -753,12 +886,20 @@ class CustomGuardrail(CustomLogger):
|
|||
Guardrails signal intentional blocks by raising:
|
||||
- GuardrailRaisedException (generic guardrail API, tool permission)
|
||||
- BlockedPiiEntityError (Presidio PII detection)
|
||||
- SensitiveDataRouteException (sensitive-data reroute to on-premise model)
|
||||
- HTTPException with status 400 (content policy violation)
|
||||
- ModifyResponseException (passthrough mode violation)
|
||||
"""
|
||||
if isinstance(e, ModifyResponseException):
|
||||
return True
|
||||
if isinstance(e, (GuardrailRaisedException, BlockedPiiEntityError)):
|
||||
if isinstance(
|
||||
e,
|
||||
(
|
||||
GuardrailRaisedException,
|
||||
BlockedPiiEntityError,
|
||||
SensitiveDataRouteException,
|
||||
),
|
||||
):
|
||||
return True
|
||||
if (
|
||||
HTTPException is not None
|
||||
|
|
|
|||
|
|
@ -2690,7 +2690,7 @@ class PrometheusLogger(CustomLogger):
|
|||
Args:
|
||||
guardrail_name: Name of the guardrail
|
||||
latency_seconds: Execution latency in seconds
|
||||
status: "success" or "error"
|
||||
status: "success", "error", or "intervened"
|
||||
error_type: Type of error if any, None otherwise
|
||||
hook_type: "pre_call", "during_call", or "post_call"
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -81,7 +81,6 @@ from litellm.types.utils import (
|
|||
from litellm.utils import (
|
||||
ModelResponse,
|
||||
Usage,
|
||||
_supports_factory,
|
||||
add_dummy_tool,
|
||||
any_assistant_message_has_thinking_blocks,
|
||||
get_max_tokens,
|
||||
|
|
@ -337,50 +336,6 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
v in model_lower for v in ("opus-4-7", "opus_4_7", "opus-4.7", "opus_4.7")
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _supports_model_capability(model: str, key: str) -> bool:
|
||||
"""Check a boolean capability ``key`` in the model map.
|
||||
|
||||
Strips bedrock/vertex prefixes so a provider-routed Claude still
|
||||
resolves to the Anthropic model-map entry.
|
||||
"""
|
||||
try:
|
||||
if _supports_factory(
|
||||
model=model,
|
||||
custom_llm_provider="anthropic",
|
||||
key=key,
|
||||
):
|
||||
return True
|
||||
except Exception:
|
||||
pass
|
||||
candidates = [model]
|
||||
for prefix in (
|
||||
"bedrock/converse/",
|
||||
"bedrock/invoke/",
|
||||
"bedrock/",
|
||||
"vertex_ai/",
|
||||
):
|
||||
if model.startswith(prefix):
|
||||
candidates.append(model[len(prefix) :])
|
||||
try:
|
||||
from litellm.llms.bedrock.common_utils import BedrockModelInfo
|
||||
|
||||
base = BedrockModelInfo.get_base_model(model)
|
||||
if base:
|
||||
candidates.append(base)
|
||||
candidates.append(f"bedrock/{base}")
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
for cand in candidates:
|
||||
if cand in litellm.model_cost and (
|
||||
litellm.model_cost[cand].get(key) is True
|
||||
):
|
||||
return True
|
||||
except Exception:
|
||||
pass
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def _supports_effort_level(model: str, level: str) -> bool:
|
||||
"""Check ``supports_{level}_reasoning_effort`` in the model map."""
|
||||
|
|
@ -918,7 +873,39 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
anthropic_tools = []
|
||||
mcp_servers = []
|
||||
for tool in tools:
|
||||
if "input_schema" in tool: # assume in anthropic format
|
||||
if tool.get("type") == "namespace":
|
||||
# Namespace is a grouping container (e.g. codex's multi_agent_v1).
|
||||
# Extract its nested tools and map them individually.
|
||||
for nested in tool.get("tools") or []:
|
||||
if "input_schema" in nested:
|
||||
# Already in Anthropic format.
|
||||
anthropic_tools.append(nested)
|
||||
elif "function" not in nested and "name" in nested:
|
||||
# Flat format: {type, name, description, parameters, ...}.
|
||||
# Normalize to OpenAI-wrapped format before mapping.
|
||||
wrapped = cast(
|
||||
ChatCompletionToolParam,
|
||||
{
|
||||
"type": nested.get("type", "function"),
|
||||
"function": {
|
||||
k: v for k, v in nested.items() if k != "type"
|
||||
},
|
||||
},
|
||||
)
|
||||
nested_tool, nested_mcp = self._map_tool_helper(wrapped)
|
||||
if nested_tool is not None:
|
||||
anthropic_tools.append(nested_tool)
|
||||
if nested_mcp is not None:
|
||||
mcp_servers.append(nested_mcp)
|
||||
elif "function" in nested:
|
||||
nested_tool, nested_mcp = self._map_tool_helper(
|
||||
cast(ChatCompletionToolParam, nested)
|
||||
)
|
||||
if nested_tool is not None:
|
||||
anthropic_tools.append(nested_tool)
|
||||
if nested_mcp is not None:
|
||||
mcp_servers.append(nested_mcp)
|
||||
elif "input_schema" in tool: # assume in anthropic format
|
||||
anthropic_tools.append(tool)
|
||||
else: # assume openai tool call
|
||||
new_tool, mcp_server_tool = self._map_tool_helper(tool)
|
||||
|
|
@ -1978,6 +1965,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
# Remove internal LiteLLM parameters that should not be sent to Anthropic API
|
||||
optional_params.pop("is_vertex_request", None)
|
||||
optional_params.pop("client_metadata", None)
|
||||
|
||||
data = {
|
||||
"model": model,
|
||||
|
|
|
|||
|
|
@ -272,19 +272,63 @@ class AnthropicModelInfo(BaseLLMModelInfo):
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def _is_adaptive_thinking_model(model: str) -> bool:
|
||||
"""Claude 4.6+ models use adaptive thinking with ``output_config.effort``."""
|
||||
def _supports_model_capability(model: str, key: str) -> bool:
|
||||
"""Check a boolean capability ``key`` in the model map.
|
||||
|
||||
Strips bedrock/vertex prefixes so a provider-routed Claude still
|
||||
resolves to the Anthropic model-map entry.
|
||||
"""
|
||||
from litellm.utils import _supports_factory
|
||||
|
||||
try:
|
||||
if _supports_factory(
|
||||
model=model,
|
||||
custom_llm_provider=None,
|
||||
key="supports_adaptive_thinking",
|
||||
custom_llm_provider="anthropic",
|
||||
key=key,
|
||||
):
|
||||
return True
|
||||
except Exception:
|
||||
pass
|
||||
candidates = [model]
|
||||
for prefix in (
|
||||
"bedrock/converse/",
|
||||
"bedrock/invoke/",
|
||||
"bedrock/",
|
||||
"vertex_ai/",
|
||||
):
|
||||
if model.startswith(prefix):
|
||||
candidates.append(model[len(prefix) :])
|
||||
try:
|
||||
from litellm.llms.bedrock.common_utils import BedrockModelInfo
|
||||
|
||||
base = BedrockModelInfo.get_base_model(model)
|
||||
if base:
|
||||
candidates.append(base)
|
||||
candidates.append(f"bedrock/{base}")
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
for cand in candidates:
|
||||
if cand in litellm.model_cost and (
|
||||
litellm.model_cost[cand].get(key) is True
|
||||
):
|
||||
return True
|
||||
except Exception:
|
||||
pass
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def _is_adaptive_thinking_model(model: str) -> bool:
|
||||
"""Claude 4.6+ models use adaptive thinking with ``output_config.effort``.
|
||||
|
||||
Driven by the ``supports_adaptive_thinking`` flag in the model map; the
|
||||
4.6/4.7 name checks remain only as a fallback for provider-routed ids
|
||||
whose map entries predate the flag.
|
||||
"""
|
||||
if AnthropicModelInfo._supports_model_capability(
|
||||
model, "supports_adaptive_thinking"
|
||||
):
|
||||
return True
|
||||
return AnthropicModelInfo._is_claude_4_6_model(
|
||||
model
|
||||
) or AnthropicModelInfo._is_claude_4_7_model(model)
|
||||
|
|
|
|||
|
|
@ -2586,6 +2586,8 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers,
|
||||
)
|
||||
|
||||
headers.setdefault("Content-Type", "application/json")
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=input,
|
||||
|
|
@ -2676,6 +2678,8 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers,
|
||||
)
|
||||
|
||||
headers.setdefault("Content-Type", "application/json")
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=input,
|
||||
|
|
|
|||
|
|
@ -1319,6 +1319,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_adaptive_thinking": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
|
|
@ -1350,6 +1351,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_adaptive_thinking": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
|
|
@ -1381,6 +1383,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_adaptive_thinking": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
|
|
@ -1412,6 +1415,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_adaptive_thinking": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
|
|
@ -1443,6 +1447,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_adaptive_thinking": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
|
|
@ -2194,6 +2199,7 @@
|
|||
"cache_creation_input_token_cost": 6.25e-06,
|
||||
"cache_creation_input_token_cost_above_1hr": 1e-05,
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
"supports_adaptive_thinking": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
|
|
@ -33927,6 +33933,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_adaptive_thinking": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
|
|
@ -33955,6 +33962,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_adaptive_thinking": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
|
|
|
|||
|
|
@ -1406,6 +1406,7 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase):
|
|||
authorization_url: Optional[str] = None
|
||||
token_url: Optional[str] = None
|
||||
registration_url: Optional[str] = None
|
||||
oauth2_flow: Optional[Literal["client_credentials", "authorization_code"]] = None
|
||||
allow_all_keys: bool = False
|
||||
available_on_public_internet: bool = True
|
||||
delegate_auth_to_upstream: bool = False
|
||||
|
|
@ -1480,6 +1481,7 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase):
|
|||
authorization_url: Optional[str] = None
|
||||
token_url: Optional[str] = None
|
||||
registration_url: Optional[str] = None
|
||||
oauth2_flow: Optional[Literal["client_credentials", "authorization_code"]] = None
|
||||
allow_all_keys: bool = False
|
||||
available_on_public_internet: bool = True
|
||||
delegate_auth_to_upstream: bool = False
|
||||
|
|
|
|||
|
|
@ -15,6 +15,10 @@ from fastapi.responses import JSONResponse, StreamingResponse
|
|||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.url_utils import SSRFError, validate_url
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.agent_endpoints.databricks_oauth import (
|
||||
DATABRICKS_OAUTH_PARAM,
|
||||
resolve_databricks_app_auth_header,
|
||||
)
|
||||
from litellm.proxy.agent_endpoints.utils import merge_agent_headers
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.types.utils import all_litellm_params
|
||||
|
|
@ -677,6 +681,17 @@ async def invoke_agent_a2a( # noqa: PLR0915
|
|||
static_headers=static_headers or None,
|
||||
)
|
||||
|
||||
# Databricks App endpoints require a short-lived OAuth M2M token rather
|
||||
# than a static bearer. Only agents explicitly configured with a
|
||||
# ``databricks_oauth`` block get one; every other agent is left untouched.
|
||||
if litellm_params.get(DATABRICKS_OAUTH_PARAM):
|
||||
databricks_auth = await resolve_databricks_app_auth_header(litellm_params)
|
||||
if databricks_auth:
|
||||
agent_extra_headers = {
|
||||
**(agent_extra_headers or {}),
|
||||
**databricks_auth,
|
||||
}
|
||||
|
||||
# Merge agent-level guardrails into data so post_call_success_hook and
|
||||
# _handle_stream_message both pick them up. A2A agents use model
|
||||
# a2a_agent/*, which is not an llm_router deployment, so
|
||||
|
|
|
|||
250
litellm/proxy/agent_endpoints/databricks_oauth.py
Normal file
250
litellm/proxy/agent_endpoints/databricks_oauth.py
Normal file
|
|
@ -0,0 +1,250 @@
|
|||
"""
|
||||
OAuth M2M (client_credentials) support for A2A agents that target Databricks
|
||||
App endpoints.
|
||||
|
||||
Databricks Apps reject static bearer tokens; they require a short-lived OAuth
|
||||
access token minted from the workspace OIDC token endpoint. When an agent is
|
||||
registered with a ``databricks_oauth`` block in its ``litellm_params``, LiteLLM
|
||||
fetches that token via the client_credentials grant, caches it until shortly
|
||||
before expiry, and attaches it as the outbound ``Authorization`` header on every
|
||||
call the proxy makes to the agent.
|
||||
|
||||
Config example::
|
||||
|
||||
agents:
|
||||
- agent_name: my-databricks-app
|
||||
agent_card_params:
|
||||
url: https://my-app-1234.aws.databricksapps.com
|
||||
litellm_params:
|
||||
databricks_oauth:
|
||||
client_id: os.environ/DATABRICKS_CLIENT_ID
|
||||
client_secret: os.environ/DATABRICKS_CLIENT_SECRET
|
||||
workspace_url: https://dbc-abc123.cloud.databricks.com
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import hashlib
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, Optional, Tuple
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
|
||||
DATABRICKS_OAUTH_PARAM = "databricks_oauth"
|
||||
|
||||
_DEFAULT_SCOPE = "all-apis"
|
||||
_TOKEN_EXPIRY_BUFFER_SECONDS = 60
|
||||
_DEFAULT_TTL_SECONDS = 3600
|
||||
|
||||
|
||||
def _resolve_secret(value: Any) -> Optional[str]:
|
||||
"""Resolve a config value, expanding ``os.environ/`` references."""
|
||||
if not isinstance(value, str):
|
||||
return None
|
||||
if value.startswith("os.environ/"):
|
||||
return get_secret_str(value)
|
||||
return value
|
||||
|
||||
|
||||
def _token_url_from_workspace(workspace_url: str) -> str:
|
||||
"""Build the workspace OIDC token endpoint from a workspace URL."""
|
||||
base = workspace_url.strip().rstrip("/")
|
||||
if base.endswith("/serving-endpoints"):
|
||||
base = base[: -len("/serving-endpoints")]
|
||||
return f"{base}/oidc/v1/token"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class DatabricksAppOAuthConfig:
|
||||
client_id: str
|
||||
client_secret: str
|
||||
token_url: str
|
||||
scope: str
|
||||
|
||||
@property
|
||||
def cache_key(self) -> str:
|
||||
# Include a digest of the secret so a rotated client_secret yields a new
|
||||
# key and forces a fresh token instead of serving the stale one.
|
||||
secret_digest = hashlib.sha256(self.client_secret.encode()).hexdigest()[:16]
|
||||
return f"{self.token_url}|{self.client_id}|{self.scope}|{secret_digest}"
|
||||
|
||||
|
||||
def parse_databricks_oauth_config(
|
||||
litellm_params: Optional[Dict[str, Any]],
|
||||
) -> Optional[DatabricksAppOAuthConfig]:
|
||||
"""Build a Databricks App OAuth config from an agent's ``litellm_params``.
|
||||
|
||||
Returns ``None`` when the agent has no ``databricks_oauth`` block. Raises
|
||||
``ValueError`` when the block is present but incomplete, so misconfiguration
|
||||
surfaces loudly instead of silently sending an unauthenticated request.
|
||||
"""
|
||||
if not litellm_params:
|
||||
return None
|
||||
|
||||
raw = litellm_params.get(DATABRICKS_OAUTH_PARAM)
|
||||
if raw is None:
|
||||
return None
|
||||
if not isinstance(raw, dict):
|
||||
raise ValueError(
|
||||
f"'{DATABRICKS_OAUTH_PARAM}' must be a mapping of OAuth settings, "
|
||||
f"got {type(raw).__name__}"
|
||||
)
|
||||
|
||||
client_id = _resolve_secret(raw.get("client_id"))
|
||||
client_secret = _resolve_secret(raw.get("client_secret"))
|
||||
workspace_url = _resolve_secret(raw.get("workspace_url"))
|
||||
|
||||
missing = [
|
||||
name
|
||||
for name, value in (
|
||||
("client_id", client_id),
|
||||
("client_secret", client_secret),
|
||||
("workspace_url", workspace_url),
|
||||
)
|
||||
if not value
|
||||
]
|
||||
if missing:
|
||||
raise ValueError(
|
||||
f"Databricks App OAuth config is missing required field(s): "
|
||||
f"{', '.join(missing)}"
|
||||
)
|
||||
|
||||
scope = _resolve_secret(raw.get("scope")) or _DEFAULT_SCOPE
|
||||
|
||||
return DatabricksAppOAuthConfig(
|
||||
client_id=client_id, # type: ignore[arg-type]
|
||||
client_secret=client_secret, # type: ignore[arg-type]
|
||||
token_url=_token_url_from_workspace(workspace_url), # type: ignore[arg-type]
|
||||
scope=scope,
|
||||
)
|
||||
|
||||
|
||||
class DatabricksAppOAuthTokenCache(InMemoryCache):
|
||||
"""In-memory cache for Databricks App OAuth client_credentials tokens.
|
||||
|
||||
Keyed by token endpoint + client_id + scope so distinct agents and service
|
||||
principals never share a token. A per-key ``asyncio.Lock`` collapses
|
||||
concurrent fetches into a single token request.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__(default_ttl=_DEFAULT_TTL_SECONDS)
|
||||
self._locks: Dict[str, asyncio.Lock] = {}
|
||||
|
||||
def _get_lock(self, cache_key: str) -> asyncio.Lock:
|
||||
return self._locks.setdefault(cache_key, asyncio.Lock())
|
||||
|
||||
def _remove_key(self, key: str) -> None:
|
||||
# Drop the per-key lock alongside the cached token so ``_locks`` stays
|
||||
# bounded by the live key set rather than growing for every key ever seen.
|
||||
super()._remove_key(key)
|
||||
self._locks.pop(key, None)
|
||||
|
||||
def flush_cache(self) -> None:
|
||||
super().flush_cache()
|
||||
self._locks.clear()
|
||||
|
||||
async def async_get_token(self, config: DatabricksAppOAuthConfig) -> str:
|
||||
cache_key = config.cache_key
|
||||
|
||||
cached = self.get_cache(cache_key)
|
||||
if cached is not None:
|
||||
return cached
|
||||
|
||||
async with self._get_lock(cache_key):
|
||||
cached = self.get_cache(cache_key)
|
||||
if cached is not None:
|
||||
return cached
|
||||
|
||||
token, ttl = await self._fetch_token(config)
|
||||
# ttl == 0 means the token's own lifetime is shorter than the
|
||||
# refresh buffer; skip caching so we never hand out a stale token,
|
||||
# and drop the lock we just created since no cached entry will ever
|
||||
# trigger _remove_key to clean it up.
|
||||
if ttl > 0:
|
||||
self.set_cache(cache_key, token, ttl=ttl)
|
||||
else:
|
||||
self._locks.pop(cache_key, None)
|
||||
return token
|
||||
|
||||
async def _fetch_token(self, config: DatabricksAppOAuthConfig) -> Tuple[str, int]:
|
||||
client = get_async_httpx_client(llm_provider=httpxSpecialProvider.A2A)
|
||||
|
||||
verbose_logger.debug(
|
||||
"Fetching Databricks App OAuth token from %s", config.token_url
|
||||
)
|
||||
|
||||
basic_auth = base64.b64encode(
|
||||
f"{config.client_id}:{config.client_secret}".encode()
|
||||
).decode()
|
||||
try:
|
||||
response = await client.post(
|
||||
config.token_url,
|
||||
data={
|
||||
"grant_type": "client_credentials",
|
||||
"scope": config.scope,
|
||||
},
|
||||
headers={
|
||||
"Authorization": f"Basic {basic_auth}",
|
||||
"Content-Type": "application/x-www-form-urlencoded",
|
||||
},
|
||||
)
|
||||
except httpx.HTTPStatusError as exc:
|
||||
raise ValueError(
|
||||
"Databricks App OAuth token request failed with status "
|
||||
f"{exc.response.status_code}"
|
||||
) from exc
|
||||
except httpx.HTTPError as exc:
|
||||
raise ValueError(
|
||||
f"Databricks App OAuth token request failed: {exc}"
|
||||
) from exc
|
||||
|
||||
body = response.json()
|
||||
if not isinstance(body, dict):
|
||||
raise ValueError(
|
||||
"Databricks App OAuth token response returned non-object JSON "
|
||||
f"(got {type(body).__name__})"
|
||||
)
|
||||
|
||||
access_token = body.get("access_token")
|
||||
if not access_token:
|
||||
raise ValueError(
|
||||
"Databricks App OAuth token response missing 'access_token'"
|
||||
)
|
||||
|
||||
raw_expires_in = body.get("expires_in")
|
||||
try:
|
||||
expires_in = (
|
||||
int(raw_expires_in)
|
||||
if raw_expires_in is not None
|
||||
else _DEFAULT_TTL_SECONDS
|
||||
)
|
||||
except (TypeError, ValueError):
|
||||
expires_in = _DEFAULT_TTL_SECONDS
|
||||
|
||||
ttl = max(expires_in - _TOKEN_EXPIRY_BUFFER_SECONDS, 0)
|
||||
return access_token, ttl
|
||||
|
||||
|
||||
databricks_app_oauth_token_cache = DatabricksAppOAuthTokenCache()
|
||||
|
||||
|
||||
async def resolve_databricks_app_auth_header(
|
||||
litellm_params: Optional[Dict[str, Any]],
|
||||
) -> Optional[Dict[str, str]]:
|
||||
"""Return ``{"Authorization": "Bearer <token>"}`` for a Databricks App agent.
|
||||
|
||||
Returns ``None`` when the agent is not configured for Databricks App OAuth.
|
||||
"""
|
||||
config = parse_databricks_oauth_config(litellm_params)
|
||||
if config is None:
|
||||
return None
|
||||
|
||||
token = await databricks_app_oauth_token_cache.async_get_token(config)
|
||||
return {"Authorization": f"Bearer {token}"}
|
||||
|
|
@ -10,6 +10,7 @@ from .max_iterations_limiter import _PROXY_MaxIterationsHandler
|
|||
from .parallel_request_limiter import _PROXY_MaxParallelRequestsHandler
|
||||
from .parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3
|
||||
from .responses_id_security import ResponsesIDSecurity
|
||||
from .sensitive_data_routing import _PROXY_SensitiveDataRoutingHandler
|
||||
|
||||
# List of all available hooks that can be enabled.
|
||||
# Defined before the enterprise import below so that any module re-imported
|
||||
|
|
@ -23,6 +24,7 @@ PROXY_HOOKS = {
|
|||
"litellm_skills": SkillsInjectionHook,
|
||||
"max_iterations_limiter": _PROXY_MaxIterationsHandler,
|
||||
"max_budget_per_session_limiter": _PROXY_MaxBudgetPerSessionHandler,
|
||||
"sensitive_data_routing": _PROXY_SensitiveDataRoutingHandler,
|
||||
}
|
||||
|
||||
## FEATURE FLAG HOOKS ##
|
||||
|
|
|
|||
|
|
@ -32,6 +32,10 @@ from litellm.batches.batch_utils import (
|
|||
)
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy._types import SpecialModelNames, UserAPIKeyAuth
|
||||
from litellm.proxy.hooks.rate_limiter_utils import (
|
||||
ProxyHTTPRateLimitError,
|
||||
resolve_llm_provider_for_rate_limit,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span as _Span
|
||||
|
|
@ -375,6 +379,7 @@ class _PROXY_BatchRateLimiter(CustomLogger):
|
|||
descriptors: List["RateLimitDescriptor"],
|
||||
batch_usage: BatchFileUsage,
|
||||
limit_type: str,
|
||||
requested_model: Optional[str] = None,
|
||||
) -> None:
|
||||
"""Raise HTTPException for rate limit exceeded."""
|
||||
from datetime import datetime
|
||||
|
|
@ -419,7 +424,10 @@ class _PROXY_BatchRateLimiter(CustomLogger):
|
|||
f"Limit resets at: {reset_time_formatted}"
|
||||
)
|
||||
|
||||
raise HTTPException(
|
||||
resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(
|
||||
requested_model
|
||||
)
|
||||
raise ProxyHTTPRateLimitError(
|
||||
status_code=429,
|
||||
detail=detail,
|
||||
headers={
|
||||
|
|
@ -427,6 +435,8 @@ class _PROXY_BatchRateLimiter(CustomLogger):
|
|||
"rate_limit_type": limit_type,
|
||||
"reset_at": reset_time_formatted,
|
||||
},
|
||||
model=resolved_model,
|
||||
llm_provider=llm_provider,
|
||||
)
|
||||
|
||||
async def _check_and_increment_batch_counters(
|
||||
|
|
@ -470,6 +480,7 @@ class _PROXY_BatchRateLimiter(CustomLogger):
|
|||
)
|
||||
|
||||
if rate_limit_response["overall_code"] == "OVER_LIMIT":
|
||||
requested_model = data.get("model") if data else None
|
||||
for status in rate_limit_response["statuses"]:
|
||||
if status["code"] == "OVER_LIMIT":
|
||||
self._raise_rate_limit_error(
|
||||
|
|
@ -477,6 +488,7 @@ class _PROXY_BatchRateLimiter(CustomLogger):
|
|||
descriptors,
|
||||
batch_usage,
|
||||
status["rate_limit_type"],
|
||||
requested_model=requested_model,
|
||||
)
|
||||
|
||||
async def count_input_file_usage(
|
||||
|
|
|
|||
|
|
@ -6,20 +6,21 @@ import asyncio
|
|||
import os
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
import litellm
|
||||
from litellm import ModelResponse, Router
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.hooks.rate_limiter_utils import (
|
||||
ProxyHTTPRateLimitError,
|
||||
convert_priority_to_percent,
|
||||
resolve_llm_provider_for_rate_limit,
|
||||
)
|
||||
from litellm.types.router import ModelGroupInfo
|
||||
from litellm.types.utils import CallTypesLiteral
|
||||
from litellm.utils import get_utc_datetime
|
||||
|
||||
from .rate_limiter_utils import convert_priority_to_percent
|
||||
|
||||
|
||||
class DynamicRateLimiterCache:
|
||||
"""
|
||||
|
|
@ -218,7 +219,10 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger):
|
|||
)
|
||||
### CHECK TPM ###
|
||||
if available_tpm is not None and available_tpm == 0:
|
||||
raise HTTPException(
|
||||
resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(
|
||||
data.get("model")
|
||||
)
|
||||
raise ProxyHTTPRateLimitError(
|
||||
status_code=429,
|
||||
detail={
|
||||
"error": "Key={} over available TPM={}. Model TPM={}, Active keys={}".format(
|
||||
|
|
@ -228,10 +232,15 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger):
|
|||
active_projects,
|
||||
)
|
||||
},
|
||||
model=resolved_model,
|
||||
llm_provider=llm_provider,
|
||||
)
|
||||
### CHECK RPM ###
|
||||
elif available_rpm is not None and available_rpm == 0:
|
||||
raise HTTPException(
|
||||
resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(
|
||||
data.get("model")
|
||||
)
|
||||
raise ProxyHTTPRateLimitError(
|
||||
status_code=429,
|
||||
detail={
|
||||
"error": "Key={} over available RPM={}. Model RPM={}, Active keys={}".format(
|
||||
|
|
@ -241,6 +250,8 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger):
|
|||
active_projects,
|
||||
)
|
||||
},
|
||||
model=resolved_model,
|
||||
llm_provider=llm_provider,
|
||||
)
|
||||
elif available_rpm is not None or available_tpm is not None:
|
||||
## UPDATE CACHE WITH ACTIVE PROJECT
|
||||
|
|
|
|||
|
|
@ -19,7 +19,11 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import (
|
|||
RateLimitDescriptorRateLimitObject,
|
||||
_PROXY_MaxParallelRequestsHandler_v3,
|
||||
)
|
||||
from litellm.proxy.hooks.rate_limiter_utils import convert_priority_to_percent
|
||||
from litellm.proxy.hooks.rate_limiter_utils import (
|
||||
ProxyHTTPRateLimitError,
|
||||
convert_priority_to_percent,
|
||||
resolve_llm_provider_for_rate_limit,
|
||||
)
|
||||
from litellm.proxy.utils import InternalUsageCache
|
||||
from litellm.types.router import ModelGroupInfo
|
||||
from litellm.types.utils import CallTypesLiteral
|
||||
|
|
@ -487,12 +491,13 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
|
|||
)
|
||||
|
||||
if atomic_response["overall_code"] == "OVER_LIMIT":
|
||||
resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(model)
|
||||
for status in atomic_response["statuses"]:
|
||||
if status["code"] != "OVER_LIMIT":
|
||||
continue
|
||||
descriptor_key = status["descriptor_key"]
|
||||
if descriptor_key == "model_saturation_check":
|
||||
raise HTTPException(
|
||||
raise ProxyHTTPRateLimitError(
|
||||
status_code=429,
|
||||
detail={
|
||||
"error": f"Model capacity reached for {model}. "
|
||||
|
|
@ -507,13 +512,15 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
|
|||
"rate_limit_type": str(status["rate_limit_type"]),
|
||||
"x-litellm-priority": priority or "default",
|
||||
},
|
||||
model=resolved_model,
|
||||
llm_provider=llm_provider,
|
||||
)
|
||||
if descriptor_key == "priority_model":
|
||||
verbose_proxy_logger.debug(
|
||||
f"Enforcing priority limits for {model}, saturation: {saturation:.1%}, "
|
||||
f"priority: {priority}"
|
||||
)
|
||||
raise HTTPException(
|
||||
raise ProxyHTTPRateLimitError(
|
||||
status_code=429,
|
||||
detail={
|
||||
"error": f"Priority-based rate limit exceeded. "
|
||||
|
|
@ -531,6 +538,8 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
|
|||
"x-litellm-priority": priority or "default",
|
||||
"x-litellm-saturation": f"{saturation:.2%}",
|
||||
},
|
||||
model=resolved_model,
|
||||
llm_provider=llm_provider,
|
||||
)
|
||||
|
||||
# Fail-closed guard: overall_code says OVER_LIMIT but no status
|
||||
|
|
@ -547,7 +556,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
|
|||
f"Dynamic rate limiter: OVER_LIMIT response with unknown "
|
||||
f"descriptor_key(s) — refusing request. response={atomic_response}"
|
||||
)
|
||||
raise HTTPException(
|
||||
raise ProxyHTTPRateLimitError(
|
||||
status_code=429,
|
||||
detail={
|
||||
"error": "Rate limit exceeded",
|
||||
|
|
@ -562,6 +571,8 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger):
|
|||
"retry-after": str(self.v3_limiter.window_size),
|
||||
"x-litellm-priority": priority or "default",
|
||||
},
|
||||
model=resolved_model,
|
||||
llm_provider=llm_provider,
|
||||
)
|
||||
|
||||
# If priority is NOT enforced (saturation below threshold) but
|
||||
|
|
|
|||
|
|
@ -5,6 +5,10 @@ from litellm._logging import verbose_proxy_logger
|
|||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.hooks.rate_limiter_utils import (
|
||||
ProxyHTTPRateLimitError,
|
||||
resolve_llm_provider_for_rate_limit,
|
||||
)
|
||||
|
||||
|
||||
class _PROXY_MaxBudgetLimiter(CustomLogger):
|
||||
|
|
@ -63,7 +67,15 @@ class _PROXY_MaxBudgetLimiter(CustomLogger):
|
|||
|
||||
# CHECK IF REQUEST ALLOWED
|
||||
if curr_spend >= max_budget:
|
||||
raise HTTPException(status_code=429, detail="Max budget limit reached.")
|
||||
resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(
|
||||
data.get("model") if data else None
|
||||
)
|
||||
raise ProxyHTTPRateLimitError(
|
||||
status_code=429,
|
||||
detail="Max budget limit reached.",
|
||||
model=resolved_model,
|
||||
llm_provider=llm_provider,
|
||||
)
|
||||
except HTTPException as e:
|
||||
raise e
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -17,12 +17,14 @@ Follows the same pattern as max_iterations_limiter.py.
|
|||
import os
|
||||
from typing import TYPE_CHECKING, Any, Optional, Union
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm import DualCache
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.hooks.rate_limiter_utils import (
|
||||
ProxyHTTPRateLimitError,
|
||||
resolve_llm_provider_for_rate_limit,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache
|
||||
|
|
@ -112,13 +114,18 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger):
|
|||
)
|
||||
|
||||
if current_spend >= max_budget:
|
||||
raise HTTPException(
|
||||
resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(
|
||||
data.get("model") if data else None
|
||||
)
|
||||
raise ProxyHTTPRateLimitError(
|
||||
status_code=429,
|
||||
detail=(
|
||||
f"Session budget exceeded for session {session_id}. "
|
||||
f"Current spend: ${current_spend:.4f}, "
|
||||
f"max_budget_per_session: ${max_budget:.2f}."
|
||||
),
|
||||
model=resolved_model,
|
||||
llm_provider=llm_provider,
|
||||
)
|
||||
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -13,12 +13,14 @@ Follows the same pattern as parallel_request_limiter_v3.py.
|
|||
import os
|
||||
from typing import TYPE_CHECKING, Any, Optional, Union
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm import DualCache
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.hooks.rate_limiter_utils import (
|
||||
ProxyHTTPRateLimitError,
|
||||
resolve_llm_provider_for_rate_limit,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache
|
||||
|
|
@ -116,12 +118,17 @@ class _PROXY_MaxIterationsHandler(CustomLogger):
|
|||
current_count = await self._increment_and_get(cache_key)
|
||||
|
||||
if current_count > max_iterations:
|
||||
raise HTTPException(
|
||||
resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(
|
||||
data.get("model") if data else None
|
||||
)
|
||||
raise ProxyHTTPRateLimitError(
|
||||
status_code=429,
|
||||
detail=(
|
||||
f"Max iterations exceeded for session {session_id}. "
|
||||
f"Current count: {current_count}, max_iterations: {max_iterations}."
|
||||
),
|
||||
model=resolved_model,
|
||||
llm_provider=llm_provider,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
|
|
|
|||
|
|
@ -17,6 +17,10 @@ from litellm.proxy.auth.auth_utils import (
|
|||
get_key_model_rpm_limit,
|
||||
get_key_model_tpm_limit,
|
||||
)
|
||||
from litellm.proxy.hooks.rate_limiter_utils import (
|
||||
ProxyHTTPRateLimitError,
|
||||
resolve_llm_provider_for_rate_limit,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span as _Span
|
||||
|
|
@ -73,7 +77,8 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
|
|||
if max_parallel_requests == 0 or tpm_limit == 0 or rpm_limit == 0:
|
||||
# base case
|
||||
raise self.raise_rate_limit_error(
|
||||
additional_details=f"{CommonProxyErrors.max_parallel_request_limit_reached.value}. Hit limit for {rate_limit_type}. Current limits: max_parallel_requests: {max_parallel_requests}, tpm_limit: {tpm_limit}, rpm_limit: {rpm_limit}"
|
||||
additional_details=f"{CommonProxyErrors.max_parallel_request_limit_reached.value}. Hit limit for {rate_limit_type}. Current limits: max_parallel_requests: {max_parallel_requests}, tpm_limit: {tpm_limit}, rpm_limit: {rpm_limit}",
|
||||
requested_model=data.get("model") if data else None,
|
||||
)
|
||||
new_val = {
|
||||
"current_requests": 1,
|
||||
|
|
@ -95,10 +100,16 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
|
|||
values_to_update_in_cache.append((request_count_api_key, new_val))
|
||||
|
||||
else:
|
||||
raise HTTPException(
|
||||
requested_model = data.get("model") if data else None
|
||||
resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(
|
||||
requested_model
|
||||
)
|
||||
raise ProxyHTTPRateLimitError(
|
||||
status_code=429,
|
||||
detail=f"LiteLLM Rate Limit Handler for rate limit type = {rate_limit_type}. {CommonProxyErrors.max_parallel_request_limit_reached.value}. current rpm: {current['current_rpm']}, rpm limit: {rpm_limit}, current tpm: {current['current_tpm']}, tpm limit: {tpm_limit}, current max_parallel_requests: {current['current_requests']}, max_parallel_requests: {max_parallel_requests}",
|
||||
headers={"retry-after": str(self.time_to_next_minute())},
|
||||
model=resolved_model,
|
||||
llm_provider=llm_provider,
|
||||
)
|
||||
|
||||
await self.internal_usage_cache.async_batch_set_cache(
|
||||
|
|
@ -122,18 +133,31 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
|
|||
return seconds_to_next_minute
|
||||
|
||||
def raise_rate_limit_error(
|
||||
self, additional_details: Optional[str] = None
|
||||
self,
|
||||
additional_details: Optional[str] = None,
|
||||
requested_model: Optional[str] = None,
|
||||
) -> HTTPException:
|
||||
"""
|
||||
Raise an HTTPException with a 429 status code and a retry-after header
|
||||
Raise an HTTPException with a 429 status code and a retry-after header.
|
||||
|
||||
``requested_model`` is resolved via :func:`get_llm_provider` so the
|
||||
raised exception carries ``llm_provider`` for downstream loggers
|
||||
(Prometheus failure metric, observability callbacks). Falls back to
|
||||
``llm_provider="litellm_proxy"`` when the model is missing or
|
||||
unparseable — see ``resolve_llm_provider_for_rate_limit``.
|
||||
"""
|
||||
error_message = "Max parallel request limit reached"
|
||||
if additional_details is not None:
|
||||
error_message = error_message + " " + additional_details
|
||||
raise HTTPException(
|
||||
resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(
|
||||
requested_model
|
||||
)
|
||||
raise ProxyHTTPRateLimitError(
|
||||
status_code=429,
|
||||
detail=f"Max parallel request limit reached {additional_details}",
|
||||
detail=error_message,
|
||||
headers={"retry-after": str(self.time_to_next_minute())},
|
||||
model=resolved_model,
|
||||
llm_provider=llm_provider,
|
||||
)
|
||||
|
||||
async def get_all_cache_objects(
|
||||
|
|
@ -225,7 +249,8 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
|
|||
# if above -> raise error
|
||||
if current_global_requests >= global_max_parallel_requests:
|
||||
return self.raise_rate_limit_error(
|
||||
additional_details=f"Hit Global Limit: Limit={global_max_parallel_requests}, current: {current_global_requests}"
|
||||
additional_details=f"Hit Global Limit: Limit={global_max_parallel_requests}, current: {current_global_requests}",
|
||||
requested_model=data.get("model") if data else None,
|
||||
)
|
||||
# if below -> increment
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -23,8 +23,6 @@ from typing import (
|
|||
cast,
|
||||
)
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm import DualCache
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE
|
||||
|
|
@ -34,6 +32,10 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.auth_utils import get_model_rate_limit_from_metadata
|
||||
from litellm.proxy.hooks.rate_limiter_utils import (
|
||||
ProxyHTTPRateLimitError,
|
||||
resolve_llm_provider_for_rate_limit,
|
||||
)
|
||||
from litellm.types.caching import RedisPipelineIncrementOperation
|
||||
from litellm.types.llms.openai import BaseLiteLLMOpenAIResponseObject
|
||||
from litellm.types.utils import CallTypes, ModelResponse, Usage
|
||||
|
|
@ -1967,6 +1969,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
self,
|
||||
response: RateLimitResponse,
|
||||
descriptors: List[RateLimitDescriptor],
|
||||
requested_model: Optional[str] = None,
|
||||
) -> None:
|
||||
"""Handle rate limit exceeded error by raising HTTPException."""
|
||||
for status in response["statuses"]:
|
||||
|
|
@ -1999,7 +2002,10 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
f"Limit resets at: {reset_time_formatted}"
|
||||
)
|
||||
|
||||
raise HTTPException(
|
||||
resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(
|
||||
requested_model
|
||||
)
|
||||
raise ProxyHTTPRateLimitError(
|
||||
status_code=429,
|
||||
detail=detail,
|
||||
headers={
|
||||
|
|
@ -2007,6 +2013,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
"rate_limit_type": str(status["rate_limit_type"]),
|
||||
"reset_at": reset_time_formatted,
|
||||
},
|
||||
model=resolved_model,
|
||||
llm_provider=llm_provider,
|
||||
)
|
||||
|
||||
async def async_pre_call_hook(
|
||||
|
|
@ -2115,6 +2123,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
self._handle_rate_limit_error(
|
||||
response=response,
|
||||
descriptors=descriptors,
|
||||
requested_model=requested_model,
|
||||
)
|
||||
else:
|
||||
# add descriptors to request headers
|
||||
|
|
@ -2188,6 +2197,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
self._handle_rate_limit_error(
|
||||
response=tpm_response,
|
||||
descriptors=descriptors,
|
||||
requested_model=requested_model,
|
||||
)
|
||||
else:
|
||||
self._stash_value_in_metadata_channels(
|
||||
|
|
|
|||
|
|
@ -2,11 +2,105 @@
|
|||
Shared utility functions for rate limiter hooks.
|
||||
"""
|
||||
|
||||
from typing import Optional, Union
|
||||
from typing import Any, Optional, Tuple, Union
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.exceptions import RateLimitError
|
||||
from litellm.types.router import ModelGroupInfo
|
||||
from litellm.types.utils import PriorityReservationDict
|
||||
|
||||
PROXY_LLM_PROVIDER_FALLBACK = "litellm_proxy"
|
||||
|
||||
|
||||
def resolve_llm_provider_for_rate_limit(
|
||||
model: Optional[str],
|
||||
) -> Tuple[str, str]:
|
||||
"""
|
||||
Resolve ``(model, llm_provider)`` for a request being rejected by an
|
||||
internal proxy-side rate-limit hook.
|
||||
|
||||
These hooks fire from ``async_pre_call_hook`` — well before
|
||||
:func:`litellm.get_llm_provider` is invoked anywhere else in the request
|
||||
lifecycle — so the raised 429 would otherwise have an empty
|
||||
``llm_provider`` field, making the resulting Prometheus
|
||||
``litellm_proxy_failed_requests_metric`` show up with
|
||||
``exception_class="RateLimitError"`` and no provider attribution.
|
||||
|
||||
Wrapped defensively: if ``model`` is missing, malformed, or
|
||||
``get_llm_provider`` raises (unknown alias, router-only model, etc.) we
|
||||
fall back to ``("", "litellm_proxy")`` so we never break the request path
|
||||
by piling a second exception on top of the rate-limit one we're trying to
|
||||
raise.
|
||||
"""
|
||||
if not model:
|
||||
return "", PROXY_LLM_PROVIDER_FALLBACK
|
||||
try:
|
||||
resolved_model, custom_llm_provider, _, _ = litellm.get_llm_provider(
|
||||
model=model,
|
||||
)
|
||||
return (
|
||||
resolved_model or model,
|
||||
custom_llm_provider or PROXY_LLM_PROVIDER_FALLBACK,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug(
|
||||
"rate_limiter_utils.resolve_llm_provider_for_rate_limit: "
|
||||
"could not resolve provider for model=%s, falling back to %s. err=%s",
|
||||
model,
|
||||
PROXY_LLM_PROVIDER_FALLBACK,
|
||||
str(e),
|
||||
)
|
||||
return model, PROXY_LLM_PROVIDER_FALLBACK
|
||||
|
||||
|
||||
class ProxyHTTPRateLimitError(HTTPException, RateLimitError): # type: ignore[misc]
|
||||
"""
|
||||
HTTPException raised by proxy-side rate-limit hooks that *also* exposes
|
||||
``model`` and ``llm_provider`` attributes.
|
||||
|
||||
Why both base classes:
|
||||
|
||||
- The proxy server's exception handler keys off ``HTTPException`` to render
|
||||
a 429 response, so we must remain an ``HTTPException``.
|
||||
- Downstream loggers (Prometheus ``async_post_call_failure_hook``,
|
||||
structured logging, observability callbacks) read ``exception.llm_provider``
|
||||
via :meth:`litellm.integrations.prometheus.PrometheusLogger._get_exception_class_name`
|
||||
and ``isinstance(exc, RateLimitError)`` for category routing. Inheriting
|
||||
from :class:`litellm.exceptions.RateLimitError` keeps that wiring intact.
|
||||
|
||||
We intentionally do not call ``RateLimitError.__init__`` (which constructs
|
||||
an httpx.Response) — it isn't needed here and just adds failure surface.
|
||||
Attribute parity is what downstream consumers rely on.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
status_code: int,
|
||||
detail: Any = None,
|
||||
headers: Optional[dict] = None,
|
||||
*,
|
||||
model: str = "",
|
||||
llm_provider: str = PROXY_LLM_PROVIDER_FALLBACK,
|
||||
) -> None:
|
||||
HTTPException.__init__(
|
||||
self, status_code=status_code, detail=detail, headers=headers
|
||||
)
|
||||
self.status_code = status_code
|
||||
self.model = model or ""
|
||||
self.llm_provider = llm_provider or PROXY_LLM_PROVIDER_FALLBACK
|
||||
# `message` is what RateLimitError.__str__ would print and what some
|
||||
# observability callbacks log. Keep it human-readable.
|
||||
self.message = detail if isinstance(detail, str) else str(detail)
|
||||
# `RateLimitError.__str__` (resolved via MRO since Starlette's
|
||||
# HTTPException doesn't define `__str__`) unconditionally reads
|
||||
# these attributes. Set them so `str(exc)` doesn't raise
|
||||
# AttributeError from logging/traceback paths.
|
||||
self.num_retries: Optional[int] = None
|
||||
self.max_retries: Optional[int] = None
|
||||
|
||||
|
||||
def convert_priority_to_percent(
|
||||
value: Union[float, PriorityReservationDict], model_info: Optional[ModelGroupInfo]
|
||||
|
|
|
|||
206
litellm/proxy/hooks/sensitive_data_routing.py
Normal file
206
litellm/proxy/hooks/sensitive_data_routing.py
Normal file
|
|
@ -0,0 +1,206 @@
|
|||
"""
|
||||
Sensitive Data Routing Hook for LiteLLM Proxy.
|
||||
|
||||
When a guardrail detects sensitive data and is configured with on_sensitive_data='route',
|
||||
this hook manages:
|
||||
1. Storing the routing decision (session_id -> model) in cache
|
||||
2. Checking incoming requests for existing routing overrides
|
||||
3. Applying sticky routing so all subsequent requests in a session go to the same model
|
||||
|
||||
Works across multiple proxy instances via DualCache (in-memory + Redis).
|
||||
"""
|
||||
|
||||
import os
|
||||
from typing import TYPE_CHECKING, Any, Optional, Union
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_guardrail import get_session_id_from_request_data
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache
|
||||
|
||||
InternalUsageCache = _InternalUsageCache
|
||||
else:
|
||||
InternalUsageCache = Any
|
||||
|
||||
|
||||
SENSITIVE_ROUTING_CACHE_PREFIX = "sensitive_route"
|
||||
DEFAULT_SENSITIVE_ROUTING_TTL = 3600
|
||||
|
||||
|
||||
class _PROXY_SensitiveDataRoutingHandler(CustomLogger):
|
||||
"""
|
||||
Pre-call hook that checks for existing sensitive data routing overrides
|
||||
and applies them to incoming requests.
|
||||
|
||||
This hook runs early in the pre-call chain and modifies the request's
|
||||
model field if a routing override exists for the session.
|
||||
"""
|
||||
|
||||
def __init__(self, internal_usage_cache: InternalUsageCache):
|
||||
self.internal_usage_cache = internal_usage_cache
|
||||
self.ttl = int(
|
||||
os.getenv(
|
||||
"LITELLM_SENSITIVE_ROUTING_TTL",
|
||||
str(DEFAULT_SENSITIVE_ROUTING_TTL),
|
||||
)
|
||||
)
|
||||
|
||||
def _make_cache_key(self, session_id: str, tenant: str) -> str:
|
||||
return f"{{{SENSITIVE_ROUTING_CACHE_PREFIX}:{tenant}:{session_id}}}:model"
|
||||
|
||||
@staticmethod
|
||||
def _resolve_tenant(user_api_key_dict: Optional[UserAPIKeyAuth]) -> str:
|
||||
"""
|
||||
Identify the authenticated principal the routing override belongs to.
|
||||
|
||||
API-key auth is scoped by the hashed key. JWT (and other keyless) auth
|
||||
has no api_key, so fall back to a stable identity claim. Without this,
|
||||
every keyless caller would share the ``default`` namespace and could read
|
||||
or overwrite another principal's session routing.
|
||||
"""
|
||||
if user_api_key_dict is None:
|
||||
return "default"
|
||||
if user_api_key_dict.api_key:
|
||||
return user_api_key_dict.api_key
|
||||
principal = [
|
||||
f"{label}:{value}"
|
||||
for label, value in (
|
||||
("user", user_api_key_dict.user_id),
|
||||
("team", user_api_key_dict.team_id),
|
||||
("org", user_api_key_dict.org_id),
|
||||
)
|
||||
if value
|
||||
]
|
||||
return "|".join(principal) if principal else "default"
|
||||
|
||||
async def _get_routed_model(
|
||||
self, session_id: str, user_api_key_dict: Optional[UserAPIKeyAuth]
|
||||
) -> Optional[str]:
|
||||
"""Get the model this session should be routed to, if any."""
|
||||
cache_key = self._make_cache_key(
|
||||
session_id, self._resolve_tenant(user_api_key_dict)
|
||||
)
|
||||
|
||||
if self.internal_usage_cache.dual_cache.redis_cache is not None:
|
||||
try:
|
||||
result = await self.internal_usage_cache.dual_cache.redis_cache.async_get_cache(
|
||||
key=cache_key
|
||||
)
|
||||
if result is not None:
|
||||
routed_model = str(result)
|
||||
remaining_ttl = await self.internal_usage_cache.dual_cache.redis_cache.async_get_ttl(
|
||||
key=cache_key
|
||||
)
|
||||
await self.internal_usage_cache.async_set_cache(
|
||||
key=cache_key,
|
||||
value=routed_model,
|
||||
ttl=remaining_ttl if remaining_ttl is not None else self.ttl,
|
||||
litellm_parent_otel_span=None,
|
||||
local_only=True,
|
||||
)
|
||||
return routed_model
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning(
|
||||
"SensitiveDataRoutingHandler: Redis GET failed, falling back to in-memory: %s",
|
||||
str(e),
|
||||
)
|
||||
|
||||
result = await self.internal_usage_cache.async_get_cache(
|
||||
key=cache_key,
|
||||
litellm_parent_otel_span=None,
|
||||
local_only=True,
|
||||
)
|
||||
if result is not None:
|
||||
return str(result)
|
||||
return None
|
||||
|
||||
async def set_session_routing(
|
||||
self,
|
||||
session_id: str,
|
||||
model: str,
|
||||
user_api_key_dict: Optional[UserAPIKeyAuth] = None,
|
||||
guardrail_name: Optional[str] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Store a routing override for a session.
|
||||
|
||||
Called by guardrails when they detect sensitive data and want to
|
||||
route the session to a specific model. The override is scoped to the
|
||||
requesting principal so sessions from different tenants cannot collide.
|
||||
"""
|
||||
cache_key = self._make_cache_key(
|
||||
session_id, self._resolve_tenant(user_api_key_dict)
|
||||
)
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
"SensitiveDataRoutingHandler: Setting session routing session_id=%s model=%s guardrail=%s ttl=%s",
|
||||
session_id,
|
||||
model,
|
||||
guardrail_name,
|
||||
self.ttl,
|
||||
)
|
||||
|
||||
if self.internal_usage_cache.dual_cache.redis_cache is not None:
|
||||
try:
|
||||
await self.internal_usage_cache.dual_cache.redis_cache.async_set_cache(
|
||||
key=cache_key,
|
||||
value=model,
|
||||
ttl=self.ttl,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning(
|
||||
"SensitiveDataRoutingHandler: Redis SET failed, falling back to in-memory: %s",
|
||||
str(e),
|
||||
)
|
||||
|
||||
await self.internal_usage_cache.async_set_cache(
|
||||
key=cache_key,
|
||||
value=model,
|
||||
ttl=self.ttl,
|
||||
litellm_parent_otel_span=None,
|
||||
local_only=True,
|
||||
)
|
||||
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
cache: DualCache,
|
||||
data: dict,
|
||||
call_type: str,
|
||||
) -> Optional[Union[Exception, str, dict]]:
|
||||
"""
|
||||
Before each LLM call, check if this session has a routing override.
|
||||
If so, modify the request's model field.
|
||||
"""
|
||||
session_id = get_session_id_from_request_data(data)
|
||||
if session_id is None:
|
||||
return None
|
||||
|
||||
routed_model = await self._get_routed_model(session_id, user_api_key_dict)
|
||||
if routed_model is None:
|
||||
return None
|
||||
|
||||
original_model = data.get("model")
|
||||
if original_model == routed_model:
|
||||
return None
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
"SensitiveDataRoutingHandler: Applying session routing override "
|
||||
"session_id=%s original_model=%s routed_model=%s",
|
||||
session_id,
|
||||
original_model,
|
||||
routed_model,
|
||||
)
|
||||
|
||||
data["model"] = routed_model
|
||||
|
||||
metadata = data.get("metadata") or {}
|
||||
metadata["sensitive_data_routing_applied"] = True
|
||||
metadata["sensitive_data_routing_original_model"] = original_model
|
||||
data["metadata"] = metadata
|
||||
|
||||
return data
|
||||
|
|
@ -327,6 +327,7 @@ model LiteLLM_MCPServerTable {
|
|||
authorization_url String?
|
||||
token_url String?
|
||||
registration_url String?
|
||||
oauth2_flow String?
|
||||
allow_all_keys Boolean @default(false)
|
||||
available_on_public_internet Boolean @default(true)
|
||||
delegate_auth_to_upstream Boolean @default(false)
|
||||
|
|
|
|||
|
|
@ -89,11 +89,14 @@ from litellm._logging import _redact_string, verbose_proxy_logger
|
|||
from litellm._service_logger import ServiceLogging, ServiceTypes
|
||||
from litellm.caching.caching import DualCache, RedisCache
|
||||
from litellm.caching.dual_cache import LimitedSizeOrderedDict
|
||||
from litellm.exceptions import RejectedRequestError
|
||||
from litellm.exceptions import RejectedRequestError, SensitiveDataRouteException
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
ModifyResponseException,
|
||||
)
|
||||
from litellm.proxy.hooks.sensitive_data_routing import (
|
||||
_PROXY_SensitiveDataRoutingHandler,
|
||||
)
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
|
||||
from litellm.integrations.SlackAlerting.utils import _add_langfuse_trace_id_to_alert
|
||||
|
|
@ -1151,6 +1154,9 @@ class ProxyLogging:
|
|||
response=response, data=data, call_type=call_type
|
||||
)
|
||||
|
||||
except SensitiveDataRouteException:
|
||||
status = "intervened"
|
||||
raise
|
||||
except Exception as e:
|
||||
status = "error"
|
||||
error_type = type(e).__name__
|
||||
|
|
@ -1460,47 +1466,57 @@ class ProxyLogging:
|
|||
self._process_guardrail_metadata(data)
|
||||
return data
|
||||
|
||||
deferred_route_exc: Optional[SensitiveDataRouteException] = None
|
||||
for _callback in caps.resolved_callbacks:
|
||||
start_time = time.time()
|
||||
if isinstance(_callback, CustomGuardrail) and data is not None:
|
||||
# Skip guardrails managed by a pipeline
|
||||
if (
|
||||
_callback.guardrail_name
|
||||
and _callback.guardrail_name in pipeline_managed
|
||||
):
|
||||
continue
|
||||
try:
|
||||
if isinstance(_callback, CustomGuardrail) and data is not None:
|
||||
# Skip guardrails managed by a pipeline
|
||||
if (
|
||||
_callback.guardrail_name
|
||||
and _callback.guardrail_name in pipeline_managed
|
||||
):
|
||||
continue
|
||||
|
||||
result = await self._process_guardrail_callback(
|
||||
callback=_callback,
|
||||
data=data, # type: ignore
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_type=call_type,
|
||||
event_type=GuardrailEventHooks.pre_call,
|
||||
)
|
||||
if result is None:
|
||||
continue
|
||||
data = result
|
||||
|
||||
elif (
|
||||
_callback is not None
|
||||
and isinstance(_callback, CustomLogger)
|
||||
and "async_pre_call_hook" in vars(_callback.__class__)
|
||||
and _callback.__class__.async_pre_call_hook
|
||||
!= CustomLogger.async_pre_call_hook
|
||||
):
|
||||
if call_type == "call_mcp_tool" and user_api_key_dict is None:
|
||||
continue
|
||||
|
||||
response = await _callback.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=self.call_details["user_api_key_cache"],
|
||||
data=data, # type: ignore
|
||||
call_type=call_type, # type: ignore
|
||||
)
|
||||
if response is not None:
|
||||
data = await self.process_pre_call_hook_response(
|
||||
response=response, data=data, call_type=call_type
|
||||
result = await self._process_guardrail_callback(
|
||||
callback=_callback,
|
||||
data=data, # type: ignore
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_type=call_type,
|
||||
event_type=GuardrailEventHooks.pre_call,
|
||||
)
|
||||
if result is None:
|
||||
continue
|
||||
data = result
|
||||
|
||||
elif (
|
||||
_callback is not None
|
||||
and isinstance(_callback, CustomLogger)
|
||||
and "async_pre_call_hook" in vars(_callback.__class__)
|
||||
and _callback.__class__.async_pre_call_hook
|
||||
!= CustomLogger.async_pre_call_hook
|
||||
):
|
||||
if call_type == "call_mcp_tool" and user_api_key_dict is None:
|
||||
continue
|
||||
|
||||
response = await _callback.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=self.call_details["user_api_key_cache"],
|
||||
data=data, # type: ignore
|
||||
call_type=call_type, # type: ignore
|
||||
)
|
||||
if response is not None:
|
||||
data = await self.process_pre_call_hook_response(
|
||||
response=response, data=data, call_type=call_type
|
||||
)
|
||||
except SensitiveDataRouteException as e:
|
||||
# Defer the reroute until remaining guardrails have run so later
|
||||
# security checks are not skipped; the first reroute wins and a
|
||||
# later guardrail that blocks still propagates. Fall through to the
|
||||
# service-span recording below so the triggering guardrail is still
|
||||
# timed like every other callback.
|
||||
if deferred_route_exc is None:
|
||||
deferred_route_exc = e
|
||||
|
||||
end_time = time.time()
|
||||
duration = end_time - start_time
|
||||
|
|
@ -1516,13 +1532,76 @@ class ProxyLogging:
|
|||
end_time=end_time,
|
||||
)
|
||||
|
||||
if deferred_route_exc is not None and data is not None:
|
||||
data = await self._handle_sensitive_data_route_exception(
|
||||
deferred_route_exc, data, user_api_key_dict
|
||||
)
|
||||
|
||||
if data is not None:
|
||||
self._process_guardrail_metadata(data)
|
||||
|
||||
return data
|
||||
except SensitiveDataRouteException as e:
|
||||
data = await self._handle_sensitive_data_route_exception(
|
||||
e, data, user_api_key_dict
|
||||
)
|
||||
if data is not None:
|
||||
self._process_guardrail_metadata(data)
|
||||
return data
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
||||
async def _handle_sensitive_data_route_exception(
|
||||
self,
|
||||
exc: SensitiveDataRouteException,
|
||||
data: Optional[dict],
|
||||
user_api_key_dict: Optional[UserAPIKeyAuth],
|
||||
) -> Optional[dict]:
|
||||
"""
|
||||
Handle SensitiveDataRouteException by rerouting the current request to
|
||||
the target model and, when sticky_session_routing is enabled, persisting
|
||||
the session override so subsequent requests reuse the same model.
|
||||
"""
|
||||
if data is None:
|
||||
return None
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
"SensitiveDataRouteException caught: session_id=%s route_to_model=%s guardrail=%s sticky=%s",
|
||||
exc.session_id,
|
||||
exc.route_to_model,
|
||||
exc.guardrail_name,
|
||||
exc.sticky_session_routing,
|
||||
)
|
||||
|
||||
if exc.sticky_session_routing:
|
||||
sensitive_routing_hook = self.get_proxy_hook("sensitive_data_routing")
|
||||
if isinstance(sensitive_routing_hook, _PROXY_SensitiveDataRoutingHandler):
|
||||
await sensitive_routing_hook.set_session_routing(
|
||||
session_id=exc.session_id,
|
||||
model=exc.route_to_model,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
guardrail_name=exc.guardrail_name,
|
||||
)
|
||||
else:
|
||||
verbose_proxy_logger.warning(
|
||||
"SensitiveDataRouteException requested sticky routing for session_id=%s "
|
||||
"but the 'sensitive_data_routing' hook is not registered. Only this request "
|
||||
"will be rerouted; subsequent requests will not be sticky.",
|
||||
exc.session_id,
|
||||
)
|
||||
|
||||
original_model = data.get("model")
|
||||
data["model"] = exc.route_to_model
|
||||
|
||||
metadata = data.get("metadata") or {}
|
||||
metadata["sensitive_data_routing_applied"] = True
|
||||
metadata["sensitive_data_routing_original_model"] = original_model
|
||||
metadata["sensitive_data_routing_guardrail"] = exc.guardrail_name
|
||||
metadata["sensitive_data_routing_detection_info"] = exc.detection_info
|
||||
data["metadata"] = metadata
|
||||
|
||||
return data
|
||||
|
||||
@staticmethod
|
||||
async def _run_guardrail_task_with_enrichment(
|
||||
callback: Any, coro: Awaitable[Any]
|
||||
|
|
|
|||
|
|
@ -1006,6 +1006,20 @@ class ResponseAPILoggingUtils:
|
|||
)
|
||||
response_api_usage: ResponseAPIUsage
|
||||
if isinstance(usage_input, dict):
|
||||
usage_input = dict(usage_input) # shallow copy; avoid mutating caller
|
||||
# Realtime *_token_details → *_tokens_details when unset.
|
||||
if (
|
||||
usage_input.get("input_tokens_details") is None
|
||||
and "input_token_details" in usage_input
|
||||
):
|
||||
usage_input["input_tokens_details"] = usage_input["input_token_details"]
|
||||
if (
|
||||
usage_input.get("output_tokens_details") is None
|
||||
and "output_token_details" in usage_input
|
||||
):
|
||||
usage_input["output_tokens_details"] = usage_input[
|
||||
"output_token_details"
|
||||
]
|
||||
total_tokens = usage_input.get("total_tokens")
|
||||
if total_tokens is None:
|
||||
input_tokens = usage_input.get("input_tokens")
|
||||
|
|
@ -1050,6 +1064,7 @@ class ResponseAPILoggingUtils:
|
|||
),
|
||||
image_tokens=getattr(output_tokens_details, "image_tokens", None),
|
||||
text_tokens=getattr(output_tokens_details, "text_tokens", None),
|
||||
audio_tokens=getattr(output_tokens_details, "audio_tokens", None),
|
||||
)
|
||||
|
||||
chat_usage = Usage(
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ from datetime import datetime
|
|||
from enum import Enum
|
||||
from typing import Any, Dict, List, Literal, Optional, Union
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
||||
from typing_extensions import Required, TypedDict
|
||||
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.akto import (
|
||||
|
|
@ -771,6 +771,58 @@ class BaseLitellmParams(
|
|||
),
|
||||
)
|
||||
|
||||
on_sensitive_data: Optional[Literal["block", "route"]] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Action to take when sensitive data is detected. "
|
||||
"'block' raises an exception (default behavior). "
|
||||
"'route' reroutes the request to the model specified in sensitive_data_route_to_model."
|
||||
),
|
||||
)
|
||||
|
||||
sensitive_data_route_to_model: Optional[str] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Model to route requests to when sensitive data is detected and on_sensitive_data='route'. "
|
||||
"This is typically an on-premise model for data privacy. "
|
||||
"The routing decision persists for the entire session."
|
||||
),
|
||||
)
|
||||
|
||||
sticky_session_routing: Optional[bool] = Field(
|
||||
default=True,
|
||||
description=(
|
||||
"When True (default), after sensitive data is detected and routed, all subsequent "
|
||||
"requests in the same session will continue routing to the same model."
|
||||
),
|
||||
)
|
||||
|
||||
@field_validator(
|
||||
"mode",
|
||||
"default_action",
|
||||
"on_disallowed_action",
|
||||
"unreachable_fallback",
|
||||
"on_sensitive_data",
|
||||
mode="before",
|
||||
check_fields=False,
|
||||
)
|
||||
@classmethod
|
||||
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
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_sensitive_data_routing(self) -> "BaseLitellmParams":
|
||||
if self.on_sensitive_data == "route" and not self.sensitive_data_route_to_model:
|
||||
raise ValueError(
|
||||
"sensitive_data_route_to_model must be set when on_sensitive_data='route'"
|
||||
)
|
||||
return self
|
||||
|
||||
model_config = ConfigDict(extra="allow", protected_namespaces=())
|
||||
|
||||
|
||||
|
|
@ -811,23 +863,6 @@ class LitellmParams(
|
|||
description="When to apply the guardrail (pre_call, post_call, during_call, logging_only)"
|
||||
)
|
||||
|
||||
@field_validator(
|
||||
"mode",
|
||||
"default_action",
|
||||
"on_disallowed_action",
|
||||
"unreachable_fallback",
|
||||
mode="before",
|
||||
check_fields=False,
|
||||
)
|
||||
@classmethod
|
||||
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
|
||||
|
||||
@field_validator("timeout", mode="before", check_fields=False)
|
||||
@classmethod
|
||||
def coerce_timeout(cls, v):
|
||||
|
|
|
|||
|
|
@ -1319,6 +1319,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_adaptive_thinking": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
|
|
@ -1350,6 +1351,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_adaptive_thinking": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
|
|
@ -1381,6 +1383,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_adaptive_thinking": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
|
|
@ -1412,6 +1415,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_adaptive_thinking": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
|
|
@ -1443,6 +1447,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_adaptive_thinking": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
|
|
@ -2194,6 +2199,7 @@
|
|||
"cache_creation_input_token_cost": 6.25e-06,
|
||||
"cache_creation_input_token_cost_above_1hr": 1e-05,
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
"supports_adaptive_thinking": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
|
|
@ -33962,6 +33968,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_adaptive_thinking": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
|
|
@ -33990,6 +33997,7 @@
|
|||
"search_context_size_low": 0.01,
|
||||
"search_context_size_medium": 0.01
|
||||
},
|
||||
"supports_adaptive_thinking": true,
|
||||
"supports_assistant_prefill": false,
|
||||
"supports_computer_use": true,
|
||||
"supports_function_calling": true,
|
||||
|
|
|
|||
|
|
@ -327,6 +327,7 @@ model LiteLLM_MCPServerTable {
|
|||
authorization_url String?
|
||||
token_url String?
|
||||
registration_url String?
|
||||
oauth2_flow String?
|
||||
allow_all_keys Boolean @default(false)
|
||||
available_on_public_internet Boolean @default(true)
|
||||
delegate_auth_to_upstream Boolean @default(false)
|
||||
|
|
|
|||
|
|
@ -1249,3 +1249,47 @@ class TestCustomGuardrailSpendLogMatchRedaction:
|
|||
slg = request_data["metadata"]["standard_logging_guardrail_information"][0]
|
||||
assert slg["guardrail_response"]["filters"][0]["regex"] == "[REDACTED]"
|
||||
assert raw["filters"][0]["regex"] == r"\d{3}-\d{2}-\d{4}"
|
||||
|
||||
|
||||
class TestGuardrailInterventionClassification:
|
||||
"""A routing decision is a deliberate guardrail intervention, not a failure."""
|
||||
|
||||
def test_sensitive_data_route_exception_is_intervention(self):
|
||||
from litellm.exceptions import SensitiveDataRouteException
|
||||
|
||||
exc = SensitiveDataRouteException(
|
||||
route_to_model="on-prem-model",
|
||||
session_id="sess-1",
|
||||
guardrail_name="pii-rail",
|
||||
)
|
||||
assert CustomGuardrail._is_guardrail_intervention(exc) is True
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_routing_logged_as_intervened_not_failed(self):
|
||||
from litellm.exceptions import SensitiveDataRouteException
|
||||
from litellm.integrations.custom_guardrail import log_guardrail_information
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
class RoutingGuardrail(CustomGuardrail):
|
||||
def __init__(self):
|
||||
super().__init__(
|
||||
guardrail_name="pii-rail",
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
)
|
||||
|
||||
@log_guardrail_information
|
||||
async def async_pre_call_hook(self, data, **kwargs):
|
||||
raise SensitiveDataRouteException(
|
||||
route_to_model="on-prem-model",
|
||||
session_id="sess-1",
|
||||
guardrail_name=self.guardrail_name,
|
||||
)
|
||||
|
||||
guardrail = RoutingGuardrail()
|
||||
request_data: dict = {"metadata": {}}
|
||||
|
||||
with pytest.raises(SensitiveDataRouteException):
|
||||
await guardrail.async_pre_call_hook(data=request_data)
|
||||
|
||||
slg = request_data["metadata"]["standard_logging_guardrail_information"][0]
|
||||
assert slg["guardrail_status"] == "guardrail_intervened"
|
||||
|
|
|
|||
|
|
@ -5090,3 +5090,101 @@ def test_map_tool_helper_collision_prefers_definitions_over_components_schemas()
|
|||
# Cross-namespace ref *also* resolves to the `definitions` body because
|
||||
# ``unpack_defs`` keys by last path segment -- documented limitation.
|
||||
assert transformed["input_schema"]["properties"]["from_components"] == expected
|
||||
|
||||
|
||||
def test_namespace_tool_flat_nested_tools_are_extracted():
|
||||
"""Codex sends nested tools in flat format {type, name, description, parameters} with no 'function' wrapper.
|
||||
These must be normalized and mapped without raising KeyError: 'function'."""
|
||||
config = AnthropicConfig()
|
||||
tools = [
|
||||
{
|
||||
"type": "namespace",
|
||||
"name": "multi_agent_v1",
|
||||
"tools": [
|
||||
{
|
||||
"type": "function",
|
||||
"name": "close_agent",
|
||||
"description": "Close an agent.",
|
||||
"strict": False,
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"target": {"type": "string"}},
|
||||
"required": ["target"],
|
||||
"additionalProperties": False,
|
||||
},
|
||||
},
|
||||
],
|
||||
}
|
||||
]
|
||||
anthropic_tools, _ = config._map_tools(tools)
|
||||
assert len(anthropic_tools) == 1
|
||||
assert anthropic_tools[0]["name"] == "close_agent"
|
||||
|
||||
|
||||
def test_namespace_tool_nested_tools_are_extracted():
|
||||
"""Codex sends type='namespace' wrapping nested tools in Anthropic format.
|
||||
The namespace container must be dropped and its nested tools extracted individually.
|
||||
"""
|
||||
config = AnthropicConfig()
|
||||
tools = [
|
||||
{
|
||||
"type": "namespace",
|
||||
"name": "multi_agent_v1",
|
||||
"description": "Tools for spawning and managing sub-agents.",
|
||||
"tools": [
|
||||
{
|
||||
"name": "close_agent",
|
||||
"type": "custom",
|
||||
"description": "Close an agent.",
|
||||
"input_schema": {
|
||||
"type": "object",
|
||||
"properties": {"target": {"type": "string"}},
|
||||
"required": ["target"],
|
||||
},
|
||||
},
|
||||
{
|
||||
"name": "resume_agent",
|
||||
"type": "custom",
|
||||
"description": "Resume a closed agent.",
|
||||
"input_schema": {
|
||||
"type": "object",
|
||||
"properties": {"id": {"type": "string"}},
|
||||
"required": ["id"],
|
||||
},
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "exec_command",
|
||||
"description": "Run a command.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"cmd": {"type": "string"}},
|
||||
"required": ["cmd"],
|
||||
},
|
||||
},
|
||||
},
|
||||
]
|
||||
anthropic_tools, mcp_servers = config._map_tools(tools)
|
||||
names = [t["name"] for t in anthropic_tools]
|
||||
assert "close_agent" in names
|
||||
assert "resume_agent" in names
|
||||
assert "exec_command" in names
|
||||
assert "multi_agent_v1" not in names
|
||||
assert len(anthropic_tools) == 3
|
||||
assert mcp_servers == []
|
||||
|
||||
|
||||
def test_client_metadata_stripped_from_anthropic_request():
|
||||
"""client_metadata passed by codex must not reach the Anthropic (or Vertex Anthropic) payload."""
|
||||
config = AnthropicConfig()
|
||||
result = config.transform_request(
|
||||
model="claude-3-5-haiku-20241022",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
optional_params={"max_tokens": 10, "client_metadata": {"originator": "codex"}},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assert "client_metadata" not in result
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@
|
|||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.anthropic.common_utils import AnthropicError
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
|
||||
AnthropicMessagesConfig,
|
||||
|
|
@ -11,6 +12,22 @@ from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_tran
|
|||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def local_model_cost_map(monkeypatch):
|
||||
"""Force the bundled backup cost map so Opus 4.8 adaptive detection (driven
|
||||
by the ``supports_adaptive_thinking`` flag) doesn't depend on the
|
||||
network-fetched ``main`` copy, which lacks the flag until this branch merges."""
|
||||
original = litellm.model_cost
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
litellm.get_model_info.cache_clear()
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
litellm.model_cost = original
|
||||
litellm.get_model_info.cache_clear()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"reasoning_effort,expected_effort",
|
||||
[
|
||||
|
|
@ -298,6 +315,40 @@ def test_legacy_thinking_high_budget_keeps_xhigh_when_supported():
|
|||
assert result.get("output_config") == {"effort": "xhigh"}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
[
|
||||
"claude-opus-4-8",
|
||||
"bedrock/us.anthropic.claude-opus-4-8",
|
||||
"bedrock/invoke/us.anthropic.claude-opus-4-8",
|
||||
],
|
||||
)
|
||||
def test_legacy_thinking_translates_to_adaptive_for_opus_48(
|
||||
model, local_model_cost_map
|
||||
):
|
||||
"""Regression for issue #29188: Opus 4.8 requires adaptive thinking, but the
|
||||
legacy ``thinking.type='enabled'`` shape was passed through unchanged for
|
||||
Bedrock 4.8 (its cost-map entry lacked ``supports_adaptive_thinking`` and the
|
||||
lookup didn't strip the provider prefix), so Bedrock rejected the request. The
|
||||
reporter's reproducer used ``budget_tokens=24000``, the ``xhigh`` bucket."""
|
||||
config = AnthropicMessagesConfig()
|
||||
optional_params = {
|
||||
"max_tokens": 100,
|
||||
"thinking": {"type": "enabled", "budget_tokens": 24000},
|
||||
}
|
||||
|
||||
result = config.transform_anthropic_messages_request(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": "ping"}],
|
||||
anthropic_messages_optional_request_params=optional_params,
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert result.get("thinking") == {"type": "adaptive"}
|
||||
assert result.get("output_config") == {"effort": "xhigh"}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"budget_tokens,expected_effort",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -14,6 +14,8 @@ import os
|
|||
import sys
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../.."))
|
||||
)
|
||||
|
|
@ -1378,3 +1380,73 @@ class TestAnthropicThinkingSignatureSelfHeal:
|
|||
config.transform_anthropic_messages_request_on_http_error(err, data)
|
||||
assert "thinking" not in data
|
||||
assert data["messages"] == []
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def local_model_cost_map(monkeypatch):
|
||||
"""Force the bundled backup cost map so detection doesn't depend on the
|
||||
network-fetched ``main`` copy (which lacks this branch's flags until merge)."""
|
||||
import litellm
|
||||
|
||||
original = litellm.model_cost
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
litellm.get_model_info.cache_clear()
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
litellm.model_cost = original
|
||||
litellm.get_model_info.cache_clear()
|
||||
|
||||
|
||||
class TestClaudeOpus48AdaptiveThinking:
|
||||
"""Opus 4.8 requires adaptive thinking (``thinking.type='adaptive'`` +
|
||||
``output_config.effort``). Detection is driven by the
|
||||
``supports_adaptive_thinking`` cost-map flag, resolved through provider
|
||||
prefixes. Before the fix the Bedrock entries lacked the flag and the lookup
|
||||
didn't strip the ``us.anthropic.``/``invoke/`` prefixes, so a
|
||||
``bedrock/us.anthropic.claude-opus-4-8`` call sent the legacy
|
||||
``thinking.type='enabled'`` shape and Bedrock rejected it (issue #29188)."""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
[
|
||||
"claude-opus-4-8",
|
||||
"anthropic/claude-opus-4-8",
|
||||
"anthropic.claude-opus-4-8",
|
||||
"bedrock/us.anthropic.claude-opus-4-8",
|
||||
"bedrock/invoke/us.anthropic.claude-opus-4-8",
|
||||
"bedrock/eu.anthropic.claude-opus-4-8",
|
||||
"vertex_ai/claude-opus-4-8",
|
||||
"azure_ai/claude-opus-4-8",
|
||||
],
|
||||
)
|
||||
def test_adaptive_thinking_detected_for_opus_4_8(self, local_model_cost_map, model):
|
||||
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
|
||||
|
||||
assert AnthropicModelInfo._is_adaptive_thinking_model(model) is True
|
||||
|
||||
def test_resolver_reads_flag_through_bedrock_invoke_prefix(
|
||||
self, local_model_cost_map
|
||||
):
|
||||
"""The resolver fix: ``bedrock/invoke/...`` resolves to the flagged
|
||||
Bedrock entry. Pure ``_supports_factory`` without prefix-stripping
|
||||
returns False here, which is why the data-only fix alone was not enough."""
|
||||
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
|
||||
|
||||
assert (
|
||||
AnthropicModelInfo._supports_model_capability(
|
||||
"bedrock/invoke/us.anthropic.claude-opus-4-8",
|
||||
"supports_adaptive_thinking",
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
["claude-opus-4-5", "claude-3-7-sonnet", "claude-3-5-haiku-20241022"],
|
||||
)
|
||||
def test_non_adaptive_models_not_detected(self, local_model_cost_map, model):
|
||||
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
|
||||
|
||||
assert AnthropicModelInfo._is_adaptive_thinking_model(model) is False
|
||||
|
|
|
|||
|
|
@ -564,6 +564,46 @@ def test_sync_delete_responses_omits_body_for_azure():
|
|||
)
|
||||
|
||||
|
||||
def _content_type(headers: dict) -> str:
|
||||
for key, value in headers.items():
|
||||
if key.lower() == "content-type":
|
||||
return value
|
||||
return ""
|
||||
|
||||
|
||||
def test_async_delete_responses_sets_json_content_type():
|
||||
"""OpenAI rejects a responses DELETE with no Content-Type by treating it as
|
||||
application/octet-stream. The handler must declare application/json."""
|
||||
captured: dict = {}
|
||||
fake_async_delete, _ = _build_delete_response_mock(captured)
|
||||
|
||||
async def run():
|
||||
with patch.object(AsyncHTTPHandler, "delete", new=fake_async_delete):
|
||||
await litellm.adelete_responses(
|
||||
response_id="resp_xyz",
|
||||
custom_llm_provider="openai",
|
||||
api_key="test-key",
|
||||
)
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
assert _content_type(captured["headers"]) == "application/json"
|
||||
|
||||
|
||||
def test_sync_delete_responses_sets_json_content_type():
|
||||
captured: dict = {}
|
||||
_, fake_sync_delete = _build_delete_response_mock(captured)
|
||||
|
||||
with patch.object(HTTPHandler, "delete", new=fake_sync_delete):
|
||||
litellm.delete_responses(
|
||||
response_id="resp_xyz",
|
||||
custom_llm_provider="openai",
|
||||
api_key="test-key",
|
||||
)
|
||||
|
||||
assert _content_type(captured["headers"]) == "application/json"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Parity tests: request-body is serialized once and reused for the wire.
|
||||
# (_async_post_anthropic_messages_with_http_error_retry)
|
||||
|
|
|
|||
|
|
@ -14,7 +14,6 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||
|
||||
import pytest
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helper: build a minimal mock agent
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
@ -307,6 +306,97 @@ async def test_convention_unrelated_prefix_not_forwarded():
|
|||
assert headers is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Databricks App OAuth M2M injection
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _mock_databricks_token_client(access_token="dbx-oauth-token"):
|
||||
response = MagicMock()
|
||||
response.raise_for_status = MagicMock()
|
||||
response.json = MagicMock(
|
||||
return_value={"access_token": access_token, "expires_in": 3600}
|
||||
)
|
||||
client = MagicMock()
|
||||
client.post = AsyncMock(return_value=response)
|
||||
return client
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_databricks_oauth_header_injected():
|
||||
"""A databricks_oauth block mints an outbound Bearer Authorization header."""
|
||||
from litellm.proxy.agent_endpoints import databricks_oauth
|
||||
|
||||
databricks_oauth.databricks_app_oauth_token_cache.flush_cache()
|
||||
|
||||
mock_agent = _make_mock_agent()
|
||||
mock_agent.litellm_params = {
|
||||
"databricks_oauth": {
|
||||
"client_id": "cid",
|
||||
"client_secret": "secret",
|
||||
"workspace_url": "https://dbc.cloud.databricks.com",
|
||||
}
|
||||
}
|
||||
mock_request = _make_mock_request()
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
|
||||
return_value=_mock_databricks_token_client("minted-token"),
|
||||
):
|
||||
mock_asend = await _invoke(mock_agent, mock_request, None)
|
||||
|
||||
headers = mock_asend.call_args.kwargs.get("agent_extra_headers")
|
||||
assert headers is not None
|
||||
assert headers.get("Authorization") == "Bearer minted-token"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_databricks_oauth_overrides_static_authorization():
|
||||
"""The minted OAuth token wins over a statically configured Authorization."""
|
||||
from litellm.proxy.agent_endpoints import databricks_oauth
|
||||
|
||||
databricks_oauth.databricks_app_oauth_token_cache.flush_cache()
|
||||
|
||||
mock_agent = _make_mock_agent(static_headers={"Authorization": "Bearer static-pat"})
|
||||
mock_agent.litellm_params = {
|
||||
"databricks_oauth": {
|
||||
"client_id": "cid",
|
||||
"client_secret": "secret",
|
||||
"workspace_url": "https://dbc.cloud.databricks.com",
|
||||
}
|
||||
}
|
||||
mock_request = _make_mock_request()
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
|
||||
return_value=_mock_databricks_token_client("oauth-wins"),
|
||||
):
|
||||
mock_asend = await _invoke(mock_agent, mock_request, None)
|
||||
|
||||
headers = mock_asend.call_args.kwargs.get("agent_extra_headers")
|
||||
assert headers is not None
|
||||
assert headers.get("Authorization") == "Bearer oauth-wins"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_databricks_agent_skips_oauth_resolution():
|
||||
"""Agents without a databricks_oauth block never enter the OAuth path."""
|
||||
mock_agent = _make_mock_agent(static_headers={"x-custom": "v"})
|
||||
mock_agent.litellm_params = {"require_trace_id_on_calls_to_agent": False}
|
||||
mock_request = _make_mock_request()
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.a2a_endpoints.resolve_databricks_app_auth_header",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_resolve:
|
||||
mock_asend = await _invoke(mock_agent, mock_request, None)
|
||||
|
||||
mock_resolve.assert_not_called()
|
||||
headers = mock_asend.call_args.kwargs.get("agent_extra_headers")
|
||||
assert headers == {"x-custom": "v"}
|
||||
assert "Authorization" not in headers
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Direct unit test for the merge utility
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -0,0 +1,496 @@
|
|||
"""
|
||||
Unit tests for Databricks App OAuth M2M support for A2A agents.
|
||||
|
||||
Covers config parsing (including os.environ/ resolution and validation),
|
||||
workspace token-URL construction, client_credentials token fetching, caching
|
||||
with expiry buffering, and the public ``resolve_databricks_app_auth_header``
|
||||
helper.
|
||||
"""
|
||||
|
||||
import base64
|
||||
from unittest.mock import MagicMock, create_autospec, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.proxy.agent_endpoints.databricks_oauth import (
|
||||
DatabricksAppOAuthConfig,
|
||||
DatabricksAppOAuthTokenCache,
|
||||
parse_databricks_oauth_config,
|
||||
resolve_databricks_app_auth_header,
|
||||
)
|
||||
|
||||
|
||||
def _expected_basic_auth(client_id: str, client_secret: str) -> str:
|
||||
token = base64.b64encode(f"{client_id}:{client_secret}".encode()).decode()
|
||||
return f"Basic {token}"
|
||||
|
||||
|
||||
def _mock_http_handler(access_token="tok-abc", expires_in=3600, post_error=None):
|
||||
"""Return a mock that mirrors litellm's ``AsyncHTTPHandler`` contract.
|
||||
|
||||
Two properties of the real handler matter for these tests and were the
|
||||
source of a runtime bug the original suite missed:
|
||||
|
||||
1. ``post`` does not accept an ``auth`` kwarg. ``create_autospec`` enforces
|
||||
the real signature, so reintroducing HTTP Basic via ``auth=`` fails with
|
||||
``TypeError`` instead of silently passing.
|
||||
2. ``post`` calls ``raise_for_status`` internally and raises
|
||||
``httpx.HTTPStatusError`` itself on non-2xx; callers never inspect the
|
||||
returned response's status. Error-path tests therefore raise from
|
||||
``post`` rather than from ``response.raise_for_status``.
|
||||
"""
|
||||
handler = create_autospec(AsyncHTTPHandler, instance=True)
|
||||
if post_error is not None:
|
||||
handler.post.side_effect = post_error
|
||||
else:
|
||||
response = MagicMock()
|
||||
response.json = MagicMock(
|
||||
return_value={"access_token": access_token, "expires_in": expires_in}
|
||||
)
|
||||
handler.post.return_value = response
|
||||
return handler
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Config parsing
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_parse_returns_none_without_block():
|
||||
assert parse_databricks_oauth_config(None) is None
|
||||
assert parse_databricks_oauth_config({}) is None
|
||||
assert parse_databricks_oauth_config({"other": "value"}) is None
|
||||
|
||||
|
||||
def test_parse_builds_config_and_token_url():
|
||||
config = parse_databricks_oauth_config(
|
||||
{
|
||||
"databricks_oauth": {
|
||||
"client_id": "cid",
|
||||
"client_secret": "secret",
|
||||
"workspace_url": "https://dbc-abc.cloud.databricks.com",
|
||||
}
|
||||
}
|
||||
)
|
||||
assert config == DatabricksAppOAuthConfig(
|
||||
client_id="cid",
|
||||
client_secret="secret",
|
||||
token_url="https://dbc-abc.cloud.databricks.com/oidc/v1/token",
|
||||
scope="all-apis",
|
||||
)
|
||||
|
||||
|
||||
def test_parse_strips_serving_endpoints_and_trailing_slash():
|
||||
config = parse_databricks_oauth_config(
|
||||
{
|
||||
"databricks_oauth": {
|
||||
"client_id": "cid",
|
||||
"client_secret": "secret",
|
||||
"workspace_url": "https://dbc-abc.cloud.databricks.com/serving-endpoints/",
|
||||
}
|
||||
}
|
||||
)
|
||||
assert config is not None
|
||||
assert config.token_url == "https://dbc-abc.cloud.databricks.com/oidc/v1/token"
|
||||
|
||||
|
||||
def test_parse_custom_scope():
|
||||
config = parse_databricks_oauth_config(
|
||||
{
|
||||
"databricks_oauth": {
|
||||
"client_id": "cid",
|
||||
"client_secret": "secret",
|
||||
"workspace_url": "https://dbc-abc.cloud.databricks.com",
|
||||
"scope": "custom-scope",
|
||||
}
|
||||
}
|
||||
)
|
||||
assert config is not None
|
||||
assert config.scope == "custom-scope"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"missing_field", ["client_id", "client_secret", "workspace_url"]
|
||||
)
|
||||
def test_parse_raises_on_missing_field(missing_field):
|
||||
block = {
|
||||
"client_id": "cid",
|
||||
"client_secret": "secret",
|
||||
"workspace_url": "https://dbc-abc.cloud.databricks.com",
|
||||
}
|
||||
block.pop(missing_field)
|
||||
with pytest.raises(ValueError, match=missing_field):
|
||||
parse_databricks_oauth_config({"databricks_oauth": block})
|
||||
|
||||
|
||||
def test_parse_raises_on_non_mapping_block():
|
||||
with pytest.raises(ValueError, match="mapping"):
|
||||
parse_databricks_oauth_config({"databricks_oauth": "not-a-dict"})
|
||||
|
||||
|
||||
def test_parse_resolves_os_environ_references(monkeypatch):
|
||||
monkeypatch.setenv("MY_DBX_CLIENT_ID", "env-cid")
|
||||
monkeypatch.setenv("MY_DBX_SECRET", "env-secret")
|
||||
config = parse_databricks_oauth_config(
|
||||
{
|
||||
"databricks_oauth": {
|
||||
"client_id": "os.environ/MY_DBX_CLIENT_ID",
|
||||
"client_secret": "os.environ/MY_DBX_SECRET",
|
||||
"workspace_url": "https://dbc-abc.cloud.databricks.com",
|
||||
}
|
||||
}
|
||||
)
|
||||
assert config is not None
|
||||
assert config.client_id == "env-cid"
|
||||
assert config.client_secret == "env-secret"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Token fetching + caching
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_token_posts_client_credentials_with_basic_auth():
|
||||
cache = DatabricksAppOAuthTokenCache()
|
||||
config = DatabricksAppOAuthConfig(
|
||||
client_id="cid",
|
||||
client_secret="secret",
|
||||
token_url="https://dbc.cloud.databricks.com/oidc/v1/token",
|
||||
scope="all-apis",
|
||||
)
|
||||
client = _mock_http_handler(access_token="tok-1")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
|
||||
return_value=client,
|
||||
):
|
||||
token = await cache.async_get_token(config)
|
||||
|
||||
assert token == "tok-1"
|
||||
client.post.assert_awaited_once()
|
||||
call = client.post.call_args
|
||||
assert call.args[0] == config.token_url
|
||||
assert call.kwargs["data"] == {
|
||||
"grant_type": "client_credentials",
|
||||
"scope": "all-apis",
|
||||
}
|
||||
# Databricks authenticates the client with HTTP Basic; it must be sent as a
|
||||
# header because litellm's AsyncHTTPHandler.post has no ``auth`` parameter.
|
||||
assert call.kwargs["headers"]["Authorization"] == _expected_basic_auth(
|
||||
"cid", "secret"
|
||||
)
|
||||
assert "auth" not in call.kwargs
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_token_is_cached_across_calls():
|
||||
cache = DatabricksAppOAuthTokenCache()
|
||||
config = DatabricksAppOAuthConfig(
|
||||
client_id="cid",
|
||||
client_secret="secret",
|
||||
token_url="https://dbc.cloud.databricks.com/oidc/v1/token",
|
||||
scope="all-apis",
|
||||
)
|
||||
client = _mock_http_handler(access_token="tok-cached")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
|
||||
return_value=client,
|
||||
):
|
||||
first = await cache.async_get_token(config)
|
||||
second = await cache.async_get_token(config)
|
||||
|
||||
assert first == second == "tok-cached"
|
||||
client.post.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_distinct_clients_do_not_share_token():
|
||||
cache = DatabricksAppOAuthTokenCache()
|
||||
config_a = DatabricksAppOAuthConfig(
|
||||
client_id="cid-a",
|
||||
client_secret="secret",
|
||||
token_url="https://dbc.cloud.databricks.com/oidc/v1/token",
|
||||
scope="all-apis",
|
||||
)
|
||||
config_b = DatabricksAppOAuthConfig(
|
||||
client_id="cid-b",
|
||||
client_secret="secret",
|
||||
token_url="https://dbc.cloud.databricks.com/oidc/v1/token",
|
||||
scope="all-apis",
|
||||
)
|
||||
|
||||
clients = [_mock_http_handler("tok-a"), _mock_http_handler("tok-b")]
|
||||
|
||||
def _next_client(*args, **kwargs):
|
||||
return clients.pop(0)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
|
||||
side_effect=_next_client,
|
||||
):
|
||||
token_a = await cache.async_get_token(config_a)
|
||||
token_b = await cache.async_get_token(config_b)
|
||||
|
||||
assert token_a == "tok-a"
|
||||
assert token_b == "tok-b"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ttl_applies_expiry_buffer():
|
||||
cache = DatabricksAppOAuthTokenCache()
|
||||
config = DatabricksAppOAuthConfig(
|
||||
client_id="cid",
|
||||
client_secret="secret",
|
||||
token_url="https://dbc.cloud.databricks.com/oidc/v1/token",
|
||||
scope="all-apis",
|
||||
)
|
||||
client = _mock_http_handler(access_token="tok", expires_in=600)
|
||||
|
||||
captured = {}
|
||||
real_set = cache.set_cache
|
||||
|
||||
def _spy_set(key, value, **kwargs):
|
||||
captured["ttl"] = kwargs.get("ttl")
|
||||
return real_set(key, value, **kwargs)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
|
||||
return_value=client,
|
||||
),
|
||||
patch.object(cache, "set_cache", side_effect=_spy_set),
|
||||
):
|
||||
await cache.async_get_token(config)
|
||||
|
||||
assert captured["ttl"] == 600 - 60
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_access_token_raises():
|
||||
cache = DatabricksAppOAuthTokenCache()
|
||||
config = DatabricksAppOAuthConfig(
|
||||
client_id="cid",
|
||||
client_secret="secret",
|
||||
token_url="https://dbc.cloud.databricks.com/oidc/v1/token",
|
||||
scope="all-apis",
|
||||
)
|
||||
client = _mock_http_handler()
|
||||
client.post.return_value.json.return_value = {"not_a_token": "x"}
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
|
||||
return_value=client,
|
||||
):
|
||||
with pytest.raises(ValueError, match="access_token"):
|
||||
await cache.async_get_token(config)
|
||||
|
||||
|
||||
def _config():
|
||||
return DatabricksAppOAuthConfig(
|
||||
client_id="cid",
|
||||
client_secret="secret",
|
||||
token_url="https://dbc.cloud.databricks.com/oidc/v1/token",
|
||||
scope="all-apis",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_http_status_error_raises_value_error():
|
||||
cache = DatabricksAppOAuthTokenCache()
|
||||
request = httpx.Request("POST", _config().token_url)
|
||||
error_response = httpx.Response(status_code=401, request=request)
|
||||
client = _mock_http_handler(
|
||||
post_error=httpx.HTTPStatusError(
|
||||
"unauthorized", request=request, response=error_response
|
||||
)
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
|
||||
return_value=client,
|
||||
):
|
||||
with pytest.raises(ValueError, match="status 401"):
|
||||
await cache.async_get_token(_config())
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_transport_error_raises_value_error():
|
||||
cache = DatabricksAppOAuthTokenCache()
|
||||
client = _mock_http_handler(post_error=httpx.ConnectError("boom"))
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
|
||||
return_value=client,
|
||||
):
|
||||
with pytest.raises(ValueError, match="token request failed"):
|
||||
await cache.async_get_token(_config())
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_object_json_body_raises():
|
||||
cache = DatabricksAppOAuthTokenCache()
|
||||
client = _mock_http_handler()
|
||||
client.post.return_value.json.return_value = ["not", "an", "object"]
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
|
||||
return_value=client,
|
||||
):
|
||||
with pytest.raises(ValueError, match="non-object JSON"):
|
||||
await cache.async_get_token(_config())
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("expires_in", [None, "not-a-number"])
|
||||
async def test_invalid_expires_in_falls_back_to_default_ttl(expires_in):
|
||||
cache = DatabricksAppOAuthTokenCache()
|
||||
client = _mock_http_handler(expires_in=expires_in)
|
||||
|
||||
captured = {}
|
||||
real_set = cache.set_cache
|
||||
|
||||
def _spy_set(key, value, **kwargs):
|
||||
captured["ttl"] = kwargs.get("ttl")
|
||||
return real_set(key, value, **kwargs)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
|
||||
return_value=client,
|
||||
),
|
||||
patch.object(cache, "set_cache", side_effect=_spy_set),
|
||||
):
|
||||
await cache.async_get_token(_config())
|
||||
|
||||
# default TTL (3600) minus the 60s expiry buffer
|
||||
assert captured["ttl"] == 3600 - 60
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_short_lived_token_not_cached():
|
||||
"""A token whose lifetime is below the refresh buffer is never cached and
|
||||
leaves no per-key lock behind."""
|
||||
cache = DatabricksAppOAuthTokenCache()
|
||||
config = _config()
|
||||
client = _mock_http_handler(access_token="short", expires_in=30)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
|
||||
return_value=client,
|
||||
):
|
||||
await cache.async_get_token(config)
|
||||
await cache.async_get_token(config)
|
||||
|
||||
assert cache.get_cache(config.cache_key) is None
|
||||
assert config.cache_key not in cache._locks
|
||||
assert client.post.await_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rotated_secret_forces_new_token():
|
||||
"""Rotating client_secret changes the cache key so a fresh token is minted."""
|
||||
cache = DatabricksAppOAuthTokenCache()
|
||||
old = DatabricksAppOAuthConfig(
|
||||
client_id="cid",
|
||||
client_secret="old-secret",
|
||||
token_url="https://dbc.cloud.databricks.com/oidc/v1/token",
|
||||
scope="all-apis",
|
||||
)
|
||||
rotated = DatabricksAppOAuthConfig(
|
||||
client_id="cid",
|
||||
client_secret="new-secret",
|
||||
token_url="https://dbc.cloud.databricks.com/oidc/v1/token",
|
||||
scope="all-apis",
|
||||
)
|
||||
assert old.cache_key != rotated.cache_key
|
||||
|
||||
clients = [_mock_http_handler("old-token"), _mock_http_handler("new-token")]
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
|
||||
side_effect=lambda *a, **k: clients.pop(0),
|
||||
):
|
||||
assert await cache.async_get_token(old) == "old-token"
|
||||
assert await cache.async_get_token(rotated) == "new-token"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_lock_pruned_when_token_evicted():
|
||||
"""The per-key lock is removed when its cached token is deleted/evicted."""
|
||||
cache = DatabricksAppOAuthTokenCache()
|
||||
config = _config()
|
||||
client = _mock_http_handler("tok")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
|
||||
return_value=client,
|
||||
):
|
||||
await cache.async_get_token(config)
|
||||
|
||||
assert config.cache_key in cache._locks
|
||||
|
||||
cache.delete_cache(config.cache_key)
|
||||
|
||||
assert config.cache_key not in cache._locks
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_flush_cache_clears_locks():
|
||||
"""flush_cache drops the per-key locks alongside the cached tokens."""
|
||||
cache = DatabricksAppOAuthTokenCache()
|
||||
config = _config()
|
||||
client = _mock_http_handler("tok")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
|
||||
return_value=client,
|
||||
):
|
||||
await cache.async_get_token(config)
|
||||
|
||||
assert config.cache_key in cache._locks
|
||||
|
||||
cache.flush_cache()
|
||||
|
||||
assert cache._locks == {}
|
||||
assert cache.get_cache(config.cache_key) is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public helper
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_returns_none_when_not_configured():
|
||||
assert await resolve_databricks_app_auth_header(None) is None
|
||||
assert await resolve_databricks_app_auth_header({"foo": "bar"}) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_returns_bearer_header():
|
||||
from litellm.proxy.agent_endpoints.databricks_oauth import (
|
||||
databricks_app_oauth_token_cache,
|
||||
)
|
||||
|
||||
databricks_app_oauth_token_cache.flush_cache()
|
||||
|
||||
litellm_params = {
|
||||
"databricks_oauth": {
|
||||
"client_id": "resolve-cid",
|
||||
"client_secret": "secret",
|
||||
"workspace_url": "https://resolve.cloud.databricks.com",
|
||||
}
|
||||
}
|
||||
client = _mock_http_handler(access_token="resolved-token")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.databricks_oauth.get_async_httpx_client",
|
||||
return_value=client,
|
||||
):
|
||||
header = await resolve_databricks_app_auth_header(litellm_params)
|
||||
|
||||
assert header == {"Authorization": "Bearer resolved-token"}
|
||||
|
|
@ -0,0 +1,968 @@
|
|||
"""
|
||||
Regression tests for the "provider field missing" bug on proxy-side
|
||||
rate-limit errors.
|
||||
|
||||
Background
|
||||
----------
|
||||
The proxy's internal rate-limit hooks (parallel_request_limiter,
|
||||
parallel_request_limiter_v3, dynamic_rate_limiter, dynamic_rate_limiter_v3,
|
||||
batch_rate_limiter, max_budget_limiter, max_iterations_limiter,
|
||||
max_budget_per_session_limiter) all fire from ``async_pre_call_hook`` —
|
||||
*before* :func:`litellm.get_llm_provider` runs anywhere else in the request
|
||||
lifecycle.
|
||||
|
||||
Until now, those hooks raised a bare ``HTTPException(429, ...)`` which carries
|
||||
no ``llm_provider`` / ``model`` attribute. Downstream:
|
||||
|
||||
- The Prometheus ``litellm_proxy_failed_requests_metric`` reads
|
||||
``exception.llm_provider`` via ``_get_exception_class_name`` — it came back
|
||||
empty, so dashboards showed ``exception_class="HTTPException"`` with no
|
||||
provider attribution.
|
||||
- Observability callbacks that ``isinstance(e, RateLimitError)`` for
|
||||
category routing missed these entirely.
|
||||
|
||||
The fix wraps every internal raise site in
|
||||
:class:`ProxyHTTPRateLimitError` (an ``HTTPException`` *and* a
|
||||
``litellm.RateLimitError``), and resolves ``model`` / ``llm_provider`` from
|
||||
``data["model"]`` via :func:`get_llm_provider`. When the model is missing or
|
||||
unparseable we fall back to ``llm_provider="litellm_proxy"`` so we never break
|
||||
the request path with a second exception.
|
||||
|
||||
These tests pin both the happy path (provider correctly resolved) and the
|
||||
fallback path (unknown model, missing model) for every limiter.
|
||||
"""
|
||||
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
import litellm
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.exceptions import RateLimitError
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.hooks.batch_rate_limiter import (
|
||||
BatchFileUsage,
|
||||
_PROXY_BatchRateLimiter,
|
||||
)
|
||||
from litellm.proxy.hooks.dynamic_rate_limiter import _PROXY_DynamicRateLimitHandler
|
||||
from litellm.proxy.hooks.dynamic_rate_limiter_v3 import (
|
||||
_PROXY_DynamicRateLimitHandlerV3,
|
||||
)
|
||||
from litellm.proxy.hooks.max_budget_limiter import _PROXY_MaxBudgetLimiter
|
||||
from litellm.proxy.hooks.max_budget_per_session_limiter import (
|
||||
_PROXY_MaxBudgetPerSessionHandler,
|
||||
)
|
||||
from litellm.proxy.hooks.max_iterations_limiter import _PROXY_MaxIterationsHandler
|
||||
from litellm.proxy.hooks.parallel_request_limiter import (
|
||||
_PROXY_MaxParallelRequestsHandler,
|
||||
)
|
||||
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
|
||||
_PROXY_MaxParallelRequestsHandler_v3,
|
||||
)
|
||||
from litellm.proxy.hooks.rate_limiter_utils import (
|
||||
PROXY_LLM_PROVIDER_FALLBACK,
|
||||
ProxyHTTPRateLimitError,
|
||||
resolve_llm_provider_for_rate_limit,
|
||||
)
|
||||
from litellm.proxy.utils import InternalUsageCache
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helper class itself
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestProxyHTTPRateLimitErrorClass:
|
||||
"""Pin the dual ``HTTPException`` + ``RateLimitError`` shape."""
|
||||
|
||||
def test_is_both_http_exception_and_rate_limit_error(self):
|
||||
e = ProxyHTTPRateLimitError(
|
||||
status_code=429,
|
||||
detail="boom",
|
||||
model="gpt-4o-mini",
|
||||
llm_provider="openai",
|
||||
)
|
||||
# FastAPI handler keys off HTTPException to render the 429.
|
||||
assert isinstance(e, HTTPException)
|
||||
# Prometheus / observability key off RateLimitError + .llm_provider.
|
||||
assert isinstance(e, RateLimitError)
|
||||
assert e.status_code == 429
|
||||
assert e.model == "gpt-4o-mini"
|
||||
assert e.llm_provider == "openai"
|
||||
assert e.message == "boom"
|
||||
assert e.detail == "boom"
|
||||
|
||||
def test_dict_detail_is_stringified_for_message(self):
|
||||
# Some hooks pass a dict detail (e.g. dynamic_rate_limiter v1) — the
|
||||
# `message` attr (read by RateLimitError.__str__ and observability
|
||||
# callbacks) must still be a string.
|
||||
e = ProxyHTTPRateLimitError(
|
||||
status_code=429,
|
||||
detail={"error": "over rpm"},
|
||||
model="claude-3-5-sonnet",
|
||||
llm_provider="anthropic",
|
||||
)
|
||||
assert isinstance(e.message, str)
|
||||
assert "over rpm" in e.message
|
||||
|
||||
def test_defaults_to_litellm_proxy_provider(self):
|
||||
e = ProxyHTTPRateLimitError(status_code=429, detail="x")
|
||||
assert e.llm_provider == PROXY_LLM_PROVIDER_FALLBACK
|
||||
assert e.model == ""
|
||||
|
||||
def test_none_provider_normalized_to_fallback(self):
|
||||
e = ProxyHTTPRateLimitError(
|
||||
status_code=429,
|
||||
detail="x",
|
||||
model=None, # type: ignore[arg-type]
|
||||
llm_provider=None, # type: ignore[arg-type]
|
||||
)
|
||||
assert e.llm_provider == PROXY_LLM_PROVIDER_FALLBACK
|
||||
assert e.model == ""
|
||||
|
||||
|
||||
class TestResolveLLMProviderForRateLimit:
|
||||
@pytest.mark.parametrize(
|
||||
"model, expected_provider",
|
||||
[
|
||||
("gpt-4o-mini", "openai"),
|
||||
("anthropic/claude-3-5-sonnet", "anthropic"),
|
||||
("bedrock/meta.llama3-1-70b-instruct-v1:0", "bedrock"),
|
||||
],
|
||||
)
|
||||
def test_known_models_resolve_provider(self, model, expected_provider):
|
||||
resolved_model, provider = resolve_llm_provider_for_rate_limit(model)
|
||||
assert provider == expected_provider
|
||||
assert resolved_model # non-empty
|
||||
|
||||
@pytest.mark.parametrize("model", [None, "", "totally-not-a-real-model-name"])
|
||||
def test_missing_or_unknown_model_falls_back(self, model):
|
||||
# Must never raise — the resolver wraps `get_llm_provider` defensively
|
||||
# because raising here would mask the rate-limit error we're trying
|
||||
# to surface to the user.
|
||||
resolved_model, provider = resolve_llm_provider_for_rate_limit(model)
|
||||
assert provider == PROXY_LLM_PROVIDER_FALLBACK
|
||||
# Resolver returns the input model verbatim on the unknown branch so
|
||||
# the `.model` attribute is never silently swapped to a different one.
|
||||
if not model:
|
||||
assert resolved_model == ""
|
||||
else:
|
||||
assert resolved_model == model
|
||||
|
||||
def test_get_llm_provider_raising_is_swallowed(self):
|
||||
# If get_llm_provider itself blows up (unexpected error), we still
|
||||
# fall back rather than letting the secondary exception escape.
|
||||
with patch.object(
|
||||
litellm,
|
||||
"get_llm_provider",
|
||||
side_effect=RuntimeError("boom"),
|
||||
):
|
||||
resolved_model, provider = resolve_llm_provider_for_rate_limit("anything")
|
||||
assert provider == PROXY_LLM_PROVIDER_FALLBACK
|
||||
assert resolved_model == "anything"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# parallel_request_limiter v1
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_parallel_request_limiter_v1_populates_provider_when_at_rpm_limit():
|
||||
"""
|
||||
Trip the per-key RPM cap and assert the raised exception carries
|
||||
``model`` / ``llm_provider`` resolved from ``data["model"]``.
|
||||
"""
|
||||
handler = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(DualCache())
|
||||
)
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="sk-rl-test",
|
||||
max_parallel_requests=10,
|
||||
rpm_limit=1,
|
||||
tpm_limit=10,
|
||||
)
|
||||
data = {"model": "gpt-4o-mini"}
|
||||
|
||||
# First request consumes the budget.
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=DualCache(),
|
||||
data=data,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=DualCache(),
|
||||
data=data,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
exc = exc_info.value
|
||||
assert exc.status_code == 429
|
||||
assert isinstance(exc, RateLimitError)
|
||||
assert exc.llm_provider == "openai"
|
||||
assert exc.model == "gpt-4o-mini"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_parallel_request_limiter_v1_zero_limit_path_populates_provider():
|
||||
"""
|
||||
When tpm_limit / rpm_limit is 0 the limiter takes the
|
||||
``raise_rate_limit_error`` path. That path receives ``requested_model``
|
||||
via the call-site change and must pass it through.
|
||||
"""
|
||||
handler = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(DualCache())
|
||||
)
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="sk-rl-zero",
|
||||
max_parallel_requests=0,
|
||||
rpm_limit=10,
|
||||
tpm_limit=10,
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=DualCache(),
|
||||
data={"model": "anthropic/claude-3-5-sonnet"},
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
exc = exc_info.value
|
||||
assert exc.status_code == 429
|
||||
assert isinstance(exc, RateLimitError)
|
||||
assert exc.llm_provider == "anthropic"
|
||||
assert exc.model == "claude-3-5-sonnet"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_parallel_request_limiter_v1_global_limit_populates_provider():
|
||||
"""global_max_parallel_requests path also threads the model through."""
|
||||
handler = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(DualCache())
|
||||
)
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="sk-global")
|
||||
|
||||
# Pre-fill the global counter so the next call exceeds it.
|
||||
await handler.internal_usage_cache.async_set_cache(
|
||||
key="global_max_parallel_requests",
|
||||
value=5,
|
||||
local_only=True,
|
||||
litellm_parent_otel_span=None,
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=DualCache(),
|
||||
data={
|
||||
"model": "bedrock/meta.llama3-1-70b-instruct-v1:0",
|
||||
"metadata": {"global_max_parallel_requests": 1},
|
||||
},
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
exc = exc_info.value
|
||||
assert exc.status_code == 429
|
||||
assert exc.llm_provider == "bedrock"
|
||||
assert exc.model == "meta.llama3-1-70b-instruct-v1:0"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_parallel_request_limiter_v1_unknown_model_falls_back():
|
||||
"""
|
||||
When ``data["model"]`` is unparseable, the resolver falls back to
|
||||
``litellm_proxy`` — and crucially does *not* leak a secondary exception.
|
||||
"""
|
||||
handler = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(DualCache())
|
||||
)
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="sk-rl-unknown",
|
||||
max_parallel_requests=10,
|
||||
rpm_limit=1,
|
||||
tpm_limit=10,
|
||||
)
|
||||
data = {"model": "totally-not-a-real-model"}
|
||||
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=DualCache(),
|
||||
data=data,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=DualCache(),
|
||||
data=data,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
exc = exc_info.value
|
||||
assert exc.status_code == 429
|
||||
assert exc.llm_provider == PROXY_LLM_PROVIDER_FALLBACK
|
||||
# Resolver returns the input verbatim so we don't silently relabel the
|
||||
# model in the user-facing 429 detail.
|
||||
assert exc.model == "totally-not-a-real-model"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_parallel_request_limiter_v1_missing_model_falls_back():
|
||||
handler = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(DualCache())
|
||||
)
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="sk-rl-no-model",
|
||||
max_parallel_requests=10,
|
||||
rpm_limit=1,
|
||||
tpm_limit=10,
|
||||
)
|
||||
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=DualCache(),
|
||||
data={},
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=DualCache(),
|
||||
data={},
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
exc = exc_info.value
|
||||
assert exc.llm_provider == PROXY_LLM_PROVIDER_FALLBACK
|
||||
assert exc.model == ""
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# parallel_request_limiter v3
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _v3_over_limit_response(rate_limit_type: str = "rpm") -> dict:
|
||||
return {
|
||||
"overall_code": "OVER_LIMIT",
|
||||
"statuses": [
|
||||
{
|
||||
"code": "OVER_LIMIT",
|
||||
"descriptor_key": "key",
|
||||
"current_limit": 1,
|
||||
"limit_remaining": -1,
|
||||
"rate_limit_type": rate_limit_type,
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"model, expected_provider",
|
||||
[
|
||||
("gpt-4o-mini", "openai"),
|
||||
("anthropic/claude-3-5-sonnet", "anthropic"),
|
||||
],
|
||||
)
|
||||
async def test_parallel_request_limiter_v3_populates_provider(model, expected_provider):
|
||||
handler = _PROXY_MaxParallelRequestsHandler_v3(
|
||||
internal_usage_cache=InternalUsageCache(DualCache())
|
||||
)
|
||||
|
||||
descriptors = [{"key": "key", "value": "v", "rate_limit": {"requests_per_unit": 1}}]
|
||||
over = _v3_over_limit_response()
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
handler._handle_rate_limit_error(
|
||||
response=over,
|
||||
descriptors=descriptors,
|
||||
requested_model=model,
|
||||
)
|
||||
|
||||
exc = exc_info.value
|
||||
assert exc.status_code == 429
|
||||
assert isinstance(exc, RateLimitError)
|
||||
assert exc.llm_provider == expected_provider
|
||||
# v3 may strip the "anthropic/" prefix in the resolved model — accept
|
||||
# either; we only care that the provider field is correct and the model
|
||||
# is non-empty.
|
||||
assert exc.model
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_parallel_request_limiter_v3_unknown_model_falls_back():
|
||||
handler = _PROXY_MaxParallelRequestsHandler_v3(
|
||||
internal_usage_cache=InternalUsageCache(DualCache())
|
||||
)
|
||||
descriptors = [{"key": "key", "value": "v", "rate_limit": {"requests_per_unit": 1}}]
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
handler._handle_rate_limit_error(
|
||||
response=_v3_over_limit_response(),
|
||||
descriptors=descriptors,
|
||||
requested_model="totally-bogus",
|
||||
)
|
||||
|
||||
assert exc_info.value.llm_provider == PROXY_LLM_PROVIDER_FALLBACK
|
||||
assert exc_info.value.model == "totally-bogus"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_parallel_request_limiter_v3_missing_model_falls_back():
|
||||
handler = _PROXY_MaxParallelRequestsHandler_v3(
|
||||
internal_usage_cache=InternalUsageCache(DualCache())
|
||||
)
|
||||
descriptors = [{"key": "key", "value": "v", "rate_limit": {"requests_per_unit": 1}}]
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
handler._handle_rate_limit_error(
|
||||
response=_v3_over_limit_response(),
|
||||
descriptors=descriptors,
|
||||
requested_model=None,
|
||||
)
|
||||
|
||||
assert exc_info.value.llm_provider == PROXY_LLM_PROVIDER_FALLBACK
|
||||
assert exc_info.value.model == ""
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# dynamic_rate_limiter v1
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dynamic_rate_limiter_v1_tpm_zero_populates_provider():
|
||||
handler = _PROXY_DynamicRateLimitHandler(internal_usage_cache=DualCache())
|
||||
handler.check_available_usage = AsyncMock(return_value=(0, 5, 100, 5, 1))
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="sk-dyn")
|
||||
user_api_key_dict.metadata = {}
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=DualCache(),
|
||||
data={"model": "gpt-4o-mini"},
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
exc = exc_info.value
|
||||
assert exc.status_code == 429
|
||||
assert isinstance(exc, RateLimitError)
|
||||
assert exc.llm_provider == "openai"
|
||||
assert exc.model == "gpt-4o-mini"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dynamic_rate_limiter_v1_rpm_zero_populates_provider():
|
||||
handler = _PROXY_DynamicRateLimitHandler(internal_usage_cache=DualCache())
|
||||
handler.check_available_usage = AsyncMock(return_value=(5, 0, 5, 100, 1))
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="sk-dyn")
|
||||
user_api_key_dict.metadata = {}
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=DualCache(),
|
||||
data={"model": "anthropic/claude-3-5-sonnet"},
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
exc = exc_info.value
|
||||
assert exc.llm_provider == "anthropic"
|
||||
assert exc.model == "claude-3-5-sonnet"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dynamic_rate_limiter_v1_unknown_model_falls_back():
|
||||
handler = _PROXY_DynamicRateLimitHandler(internal_usage_cache=DualCache())
|
||||
handler.check_available_usage = AsyncMock(return_value=(0, 5, 100, 5, 1))
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="sk-dyn")
|
||||
user_api_key_dict.metadata = {}
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=DualCache(),
|
||||
data={"model": "no-such-model"},
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
assert exc_info.value.llm_provider == PROXY_LLM_PROVIDER_FALLBACK
|
||||
assert exc_info.value.model == "no-such-model"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# dynamic_rate_limiter v3 — exercise just the raise path via the helper, not
|
||||
# the full Redis/Lua stack.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dynamic_rate_limiter_v3_model_capacity_path_populates_provider():
|
||||
"""
|
||||
The v3 dynamic limiter has three raise sites: model_saturation_check,
|
||||
priority_model, and the fail-closed unknown-descriptor branch. We patch
|
||||
the atomic increment to short-circuit straight into the model_saturation
|
||||
path — that's the most common production trip — and confirm the
|
||||
raised exception carries provider info.
|
||||
"""
|
||||
from litellm.types.router import ModelGroupInfo
|
||||
|
||||
handler = _PROXY_DynamicRateLimitHandlerV3(internal_usage_cache=DualCache())
|
||||
handler.v3_limiter.atomic_check_and_increment_by_n = AsyncMock(
|
||||
return_value={
|
||||
"overall_code": "OVER_LIMIT",
|
||||
"statuses": [
|
||||
{
|
||||
"code": "OVER_LIMIT",
|
||||
"descriptor_key": "model_saturation_check",
|
||||
"current_limit": 100,
|
||||
"limit_remaining": 0,
|
||||
"rate_limit_type": "rpm",
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
handler._create_priority_based_descriptors = MagicMock(return_value=[])
|
||||
handler._create_model_tracking_descriptor = MagicMock(
|
||||
return_value={
|
||||
"key": "model_saturation_check",
|
||||
"value": "gpt-4o-mini",
|
||||
"rate_limit": {"requests_per_unit": 100},
|
||||
}
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="sk-dyn-v3")
|
||||
user_api_key_dict.metadata = {}
|
||||
model_info = ModelGroupInfo(model_group="gpt-4o-mini", providers=["openai"])
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handler._check_rate_limits(
|
||||
model="gpt-4o-mini",
|
||||
model_group_info=model_info,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
priority="default",
|
||||
saturation=1.0,
|
||||
data={"model": "gpt-4o-mini"},
|
||||
)
|
||||
|
||||
exc = exc_info.value
|
||||
assert exc.status_code == 429
|
||||
assert isinstance(exc, RateLimitError)
|
||||
assert exc.llm_provider == "openai"
|
||||
assert exc.model == "gpt-4o-mini"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dynamic_rate_limiter_v3_unknown_descriptor_path_populates_provider():
|
||||
"""Fail-closed unknown-descriptor branch must still attribute provider."""
|
||||
from litellm.types.router import ModelGroupInfo
|
||||
|
||||
handler = _PROXY_DynamicRateLimitHandlerV3(internal_usage_cache=DualCache())
|
||||
handler.v3_limiter.atomic_check_and_increment_by_n = AsyncMock(
|
||||
return_value={
|
||||
"overall_code": "OVER_LIMIT",
|
||||
"statuses": [
|
||||
{
|
||||
"code": "OVER_LIMIT",
|
||||
"descriptor_key": "something_we_dont_handle",
|
||||
"current_limit": 1,
|
||||
"limit_remaining": 0,
|
||||
"rate_limit_type": "rpm",
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
handler._create_priority_based_descriptors = MagicMock(return_value=[])
|
||||
handler._create_model_tracking_descriptor = MagicMock(
|
||||
return_value={
|
||||
"key": "model_saturation_check",
|
||||
"value": "gpt-4o-mini",
|
||||
"rate_limit": {"requests_per_unit": 1},
|
||||
}
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="sk-dyn-v3-unknown")
|
||||
user_api_key_dict.metadata = {}
|
||||
model_info = ModelGroupInfo(model_group="gpt-4o-mini", providers=["openai"])
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handler._check_rate_limits(
|
||||
model="gpt-4o-mini",
|
||||
model_group_info=model_info,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
priority="default",
|
||||
saturation=1.0,
|
||||
data={"model": "gpt-4o-mini"},
|
||||
)
|
||||
|
||||
assert exc_info.value.llm_provider == "openai"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# batch_rate_limiter
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _batch_over_limit_response() -> dict:
|
||||
return {
|
||||
"overall_code": "OVER_LIMIT",
|
||||
"statuses": [
|
||||
{
|
||||
"code": "OVER_LIMIT",
|
||||
"descriptor_key": "key",
|
||||
"current_limit": 10,
|
||||
"limit_remaining": -5,
|
||||
"rate_limit_type": "requests",
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_batch_rate_limiter_populates_provider():
|
||||
"""
|
||||
batch_rate_limiter trips when the file's request/token count exceeds the
|
||||
remaining window. The raise must thread `data["model"]` through the
|
||||
helper.
|
||||
"""
|
||||
parallel_limiter = MagicMock()
|
||||
parallel_limiter.window_size = 60
|
||||
parallel_limiter._create_rate_limit_descriptors = MagicMock(
|
||||
return_value=[
|
||||
{"key": "key", "value": "v", "rate_limit": {"requests_per_unit": 10}}
|
||||
]
|
||||
)
|
||||
parallel_limiter.atomic_check_and_increment_by_n = AsyncMock(
|
||||
return_value=_batch_over_limit_response()
|
||||
)
|
||||
|
||||
handler = _PROXY_BatchRateLimiter(
|
||||
internal_usage_cache=InternalUsageCache(DualCache()),
|
||||
parallel_request_limiter=parallel_limiter,
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handler._check_and_increment_batch_counters(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-batch"),
|
||||
data={"model": "gpt-4o-mini"},
|
||||
batch_usage=BatchFileUsage(total_tokens=100, request_count=15),
|
||||
)
|
||||
|
||||
exc = exc_info.value
|
||||
assert exc.status_code == 429
|
||||
assert isinstance(exc, RateLimitError)
|
||||
assert exc.llm_provider == "openai"
|
||||
assert exc.model == "gpt-4o-mini"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_batch_rate_limiter_unknown_model_falls_back():
|
||||
parallel_limiter = MagicMock()
|
||||
parallel_limiter.window_size = 60
|
||||
parallel_limiter._create_rate_limit_descriptors = MagicMock(
|
||||
return_value=[
|
||||
{"key": "key", "value": "v", "rate_limit": {"requests_per_unit": 10}}
|
||||
]
|
||||
)
|
||||
parallel_limiter.atomic_check_and_increment_by_n = AsyncMock(
|
||||
return_value=_batch_over_limit_response()
|
||||
)
|
||||
|
||||
handler = _PROXY_BatchRateLimiter(
|
||||
internal_usage_cache=InternalUsageCache(DualCache()),
|
||||
parallel_request_limiter=parallel_limiter,
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handler._check_and_increment_batch_counters(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-batch"),
|
||||
data={"model": "fake-model-xyz"},
|
||||
batch_usage=BatchFileUsage(total_tokens=100, request_count=15),
|
||||
)
|
||||
|
||||
assert exc_info.value.llm_provider == PROXY_LLM_PROVIDER_FALLBACK
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# max_budget_limiter
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_max_budget_limiter_populates_provider():
|
||||
handler = _PROXY_MaxBudgetLimiter()
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="sk-budget",
|
||||
user_id="user-1",
|
||||
user_max_budget=10.0,
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.get_current_spend",
|
||||
new=AsyncMock(return_value=10.0),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=DualCache(),
|
||||
data={"model": "gpt-4o-mini"},
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
exc = exc_info.value
|
||||
assert exc.status_code == 429
|
||||
assert isinstance(exc, RateLimitError)
|
||||
assert exc.llm_provider == "openai"
|
||||
assert exc.model == "gpt-4o-mini"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_max_budget_limiter_no_model_falls_back():
|
||||
handler = _PROXY_MaxBudgetLimiter()
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="sk-budget",
|
||||
user_id="user-1",
|
||||
user_max_budget=10.0,
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.get_current_spend",
|
||||
new=AsyncMock(return_value=10.0),
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=DualCache(),
|
||||
data={},
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
assert exc_info.value.llm_provider == PROXY_LLM_PROVIDER_FALLBACK
|
||||
assert exc_info.value.model == ""
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# max_iterations_limiter
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_iter_agent(max_iterations: int) -> AgentResponse:
|
||||
return AgentResponse(
|
||||
agent_id="agent-iter",
|
||||
agent_name="iter-agent",
|
||||
litellm_params={"max_iterations": max_iterations},
|
||||
agent_card_params={"name": "iter-agent", "version": "1.0.0"},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_max_iterations_limiter_populates_provider():
|
||||
local_cache = DualCache()
|
||||
handler = _PROXY_MaxIterationsHandler(
|
||||
internal_usage_cache=InternalUsageCache(local_cache)
|
||||
)
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="sk-iter", agent_id="agent-iter")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry"
|
||||
) as mock_registry:
|
||||
mock_registry.get_agent_by_id.return_value = _make_iter_agent(max_iterations=1)
|
||||
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=local_cache,
|
||||
data={
|
||||
"model": "gpt-4o-mini",
|
||||
"metadata": {"session_id": "session-iter-1"},
|
||||
},
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=local_cache,
|
||||
data={
|
||||
"model": "gpt-4o-mini",
|
||||
"metadata": {"session_id": "session-iter-1"},
|
||||
},
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
exc = exc_info.value
|
||||
assert exc.status_code == 429
|
||||
assert isinstance(exc, RateLimitError)
|
||||
assert exc.llm_provider == "openai"
|
||||
assert exc.model == "gpt-4o-mini"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_max_iterations_limiter_unknown_model_falls_back():
|
||||
local_cache = DualCache()
|
||||
handler = _PROXY_MaxIterationsHandler(
|
||||
internal_usage_cache=InternalUsageCache(local_cache)
|
||||
)
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="sk-iter", agent_id="agent-iter")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry"
|
||||
) as mock_registry:
|
||||
mock_registry.get_agent_by_id.return_value = _make_iter_agent(max_iterations=1)
|
||||
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=local_cache,
|
||||
data={
|
||||
"model": "no-such-model",
|
||||
"metadata": {"session_id": "session-iter-2"},
|
||||
},
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=local_cache,
|
||||
data={
|
||||
"model": "no-such-model",
|
||||
"metadata": {"session_id": "session-iter-2"},
|
||||
},
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
assert exc_info.value.llm_provider == PROXY_LLM_PROVIDER_FALLBACK
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# max_budget_per_session_limiter
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_session_budget_agent(max_budget: float) -> AgentResponse:
|
||||
return AgentResponse(
|
||||
agent_id="agent-session-budget",
|
||||
agent_name="session-budget-agent",
|
||||
litellm_params={"max_budget_per_session": max_budget},
|
||||
agent_card_params={"name": "session-budget-agent", "version": "1.0.0"},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_max_budget_per_session_limiter_populates_provider():
|
||||
handler = _PROXY_MaxBudgetPerSessionHandler(
|
||||
internal_usage_cache=InternalUsageCache(DualCache())
|
||||
)
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="sk-session-budget", agent_id="agent-session-budget"
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry"
|
||||
) as mock_registry:
|
||||
mock_registry.get_agent_by_id.return_value = _make_session_budget_agent(
|
||||
max_budget=1.0
|
||||
)
|
||||
with patch.object(
|
||||
handler, "_get_current_spend", new=AsyncMock(return_value=5.0)
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=DualCache(),
|
||||
data={
|
||||
"model": "anthropic/claude-3-5-sonnet",
|
||||
"metadata": {"session_id": "session-budget-1"},
|
||||
},
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
exc = exc_info.value
|
||||
assert exc.status_code == 429
|
||||
assert isinstance(exc, RateLimitError)
|
||||
assert exc.llm_provider == "anthropic"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_max_budget_per_session_limiter_unknown_model_falls_back():
|
||||
handler = _PROXY_MaxBudgetPerSessionHandler(
|
||||
internal_usage_cache=InternalUsageCache(DualCache())
|
||||
)
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="sk-session-budget", agent_id="agent-session-budget"
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry"
|
||||
) as mock_registry:
|
||||
mock_registry.get_agent_by_id.return_value = _make_session_budget_agent(
|
||||
max_budget=1.0
|
||||
)
|
||||
with patch.object(
|
||||
handler, "_get_current_spend", new=AsyncMock(return_value=5.0)
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=DualCache(),
|
||||
data={
|
||||
"model": "no-such-model",
|
||||
"metadata": {"session_id": "session-budget-2"},
|
||||
},
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
assert exc_info.value.llm_provider == PROXY_LLM_PROVIDER_FALLBACK
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Prometheus integration: failure metric reads exception.llm_provider
|
||||
# via _get_exception_class_name. With the fix, this returns
|
||||
# "Openai.RateLimitError" instead of plain "HTTPException" for proxy-side
|
||||
# 429s on a known model. Pin that contract — that's what dashboards see.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_prometheus_exception_class_name_includes_provider():
|
||||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
|
||||
exc = ProxyHTTPRateLimitError(
|
||||
status_code=429,
|
||||
detail="over limit",
|
||||
model="gpt-4o-mini",
|
||||
llm_provider="openai",
|
||||
)
|
||||
|
||||
name = PrometheusLogger._get_exception_class_name(exc)
|
||||
# Format is "{Provider.}{ClassName}" per `_get_exception_class_name`.
|
||||
assert name.startswith("Openai.")
|
||||
# And specifically: it ends in our exception class. (We don't pin the
|
||||
# full string to avoid coupling the test to PR #27687's parallel rename.)
|
||||
assert name.endswith("ProxyHTTPRateLimitError")
|
||||
|
||||
|
||||
def test_prometheus_exception_class_name_falls_back_when_no_model():
|
||||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
|
||||
exc = ProxyHTTPRateLimitError(status_code=429, detail="over limit")
|
||||
name = PrometheusLogger._get_exception_class_name(exc)
|
||||
# `litellm_proxy` -> `Litellm_proxy.` (capitalize first char only).
|
||||
assert name.startswith("Litellm_proxy.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(pytest.main([__file__, "-vv", "-x"]))
|
||||
1036
tests/test_litellm/proxy/hooks/test_sensitive_data_routing.py
Normal file
1036
tests/test_litellm/proxy/hooks/test_sensitive_data_routing.py
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -4398,3 +4398,55 @@ class TestMCPUserEnvVarsAccessControl:
|
|||
)
|
||||
assert exc.value.status_code == 403
|
||||
get_mcp_server_mock.assert_not_awaited()
|
||||
|
||||
|
||||
def test_oauth2_flow_accepted_on_create_request():
|
||||
"""NewMCPServerRequest carries oauth2_flow through to the persisted dict."""
|
||||
from litellm.proxy._experimental.mcp_server.db import _prepare_mcp_server_data
|
||||
|
||||
payload = NewMCPServerRequest(
|
||||
server_name="m2m-server",
|
||||
url="https://example.com/mcp",
|
||||
transport="http",
|
||||
auth_type="oauth2",
|
||||
token_url="https://idp.example.com/oauth/token",
|
||||
oauth2_flow="client_credentials",
|
||||
)
|
||||
data_dict = _prepare_mcp_server_data(payload)
|
||||
assert data_dict["oauth2_flow"] == "client_credentials"
|
||||
|
||||
|
||||
def test_oauth2_flow_round_trips_on_update_and_response_models():
|
||||
"""oauth2_flow survives UpdateMCPServerRequest and the LiteLLM_MCPServerTable
|
||||
response model. Before the fix these models dropped the field (no attribute),
|
||||
which is why a persisted value never round-tripped."""
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_MCPServerTable,
|
||||
UpdateMCPServerRequest,
|
||||
)
|
||||
|
||||
update = UpdateMCPServerRequest(
|
||||
server_id="srv-1", oauth2_flow="client_credentials"
|
||||
)
|
||||
assert update.oauth2_flow == "client_credentials"
|
||||
|
||||
row = LiteLLM_MCPServerTable(
|
||||
server_id="srv-1",
|
||||
transport="http",
|
||||
oauth2_flow="client_credentials",
|
||||
)
|
||||
assert row.oauth2_flow == "client_credentials"
|
||||
|
||||
|
||||
def test_oauth2_flow_defaults_to_none_when_omitted():
|
||||
"""Omitting oauth2_flow is valid and resolves to None (runtime infers it)."""
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_MCPServerTable,
|
||||
UpdateMCPServerRequest,
|
||||
)
|
||||
|
||||
assert UpdateMCPServerRequest(server_id="srv-1").oauth2_flow is None
|
||||
assert (
|
||||
LiteLLM_MCPServerTable(server_id="srv-1", transport="http").oauth2_flow
|
||||
is None
|
||||
)
|
||||
|
|
|
|||
|
|
@ -327,7 +327,8 @@ class TestResponseAPILoggingUtils:
|
|||
"output_tokens_details": {
|
||||
"reasoning_tokens": 30,
|
||||
"image_tokens": 100,
|
||||
"text_tokens": 70,
|
||||
"text_tokens": 50,
|
||||
"audio_tokens": 20,
|
||||
},
|
||||
}
|
||||
|
||||
|
|
@ -346,7 +347,61 @@ class TestResponseAPILoggingUtils:
|
|||
assert result.completion_tokens_details is not None
|
||||
assert result.completion_tokens_details.reasoning_tokens == 30
|
||||
assert result.completion_tokens_details.image_tokens == 100
|
||||
assert result.completion_tokens_details.text_tokens == 70
|
||||
assert result.completion_tokens_details.text_tokens == 50
|
||||
assert result.completion_tokens_details.audio_tokens == 20
|
||||
|
||||
def test_transform_response_api_usage_with_realtime_keys(self):
|
||||
"""Realtime input_token_details / output_token_details normalize for Usage."""
|
||||
usage = {
|
||||
"input_tokens": 10,
|
||||
"output_tokens": 20,
|
||||
"total_tokens": 30,
|
||||
"input_token_details": {
|
||||
"text_tokens": 8,
|
||||
"audio_tokens": 2,
|
||||
"cached_tokens": 0,
|
||||
},
|
||||
"output_token_details": {
|
||||
"text_tokens": 12,
|
||||
"audio_tokens": 8,
|
||||
},
|
||||
}
|
||||
|
||||
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
|
||||
usage
|
||||
)
|
||||
|
||||
assert result.prompt_tokens_details is not None
|
||||
assert result.prompt_tokens_details.text_tokens == 8
|
||||
assert result.prompt_tokens_details.audio_tokens == 2
|
||||
|
||||
assert result.completion_tokens_details is not None
|
||||
assert result.completion_tokens_details.text_tokens == 12
|
||||
assert result.completion_tokens_details.audio_tokens == 8
|
||||
|
||||
def test_transform_response_api_usage_tokens_details_keep_values(self):
|
||||
"""Keeps input_tokens_details / output_tokens_details when singular keys are also present."""
|
||||
usage = {
|
||||
"input_tokens": 10,
|
||||
"output_tokens": 20,
|
||||
"total_tokens": 30,
|
||||
"input_tokens_details": {"text_tokens": 10},
|
||||
"output_tokens_details": {"text_tokens": 20},
|
||||
"input_token_details": {"text_tokens": 1, "audio_tokens": 99},
|
||||
"output_token_details": {"text_tokens": 2, "audio_tokens": 98},
|
||||
}
|
||||
|
||||
result = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
|
||||
usage
|
||||
)
|
||||
|
||||
assert result.prompt_tokens_details is not None
|
||||
assert result.prompt_tokens_details.text_tokens == 10
|
||||
assert result.prompt_tokens_details.audio_tokens is None
|
||||
|
||||
assert result.completion_tokens_details is not None
|
||||
assert result.completion_tokens_details.text_tokens == 20
|
||||
assert result.completion_tokens_details.audio_tokens is None
|
||||
|
||||
|
||||
class TestResponsesAPIProviderSpecificParams:
|
||||
|
|
|
|||
|
|
@ -182,3 +182,24 @@ def test_opus_4_8_provider_resolves_via_model_info(local_model_cost_map):
|
|||
assert info["litellm_provider"] == "anthropic"
|
||||
assert info["max_input_tokens"] == 1000000
|
||||
assert info["max_output_tokens"] == 128000
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"cost_map",
|
||||
[_load_root_cost_map(), GetModelCostMap.load_local_model_cost_map()],
|
||||
ids=["root", "bundled_backup"],
|
||||
)
|
||||
def test_opus_4_8_all_variants_carry_adaptive_thinking_flag(cost_map):
|
||||
"""Every Opus 4.8 entry must advertise ``supports_adaptive_thinking``.
|
||||
|
||||
Adaptive-thinking detection is cost-map driven, so a single variant missing
|
||||
the flag silently sends the legacy ``thinking.type='enabled'`` shape and the
|
||||
provider 400s (issue #29188, which the Bedrock/Vertex/Azure variants hit
|
||||
because only the bare ``claude-opus-4-8`` entry carried the flag). This guards
|
||||
against a future variant being added without it."""
|
||||
variants = [k for k in cost_map if "claude-opus-4-8" in k]
|
||||
assert variants, "no claude-opus-4-8 entries found in cost map"
|
||||
missing = [
|
||||
k for k in variants if cost_map[k].get("supports_adaptive_thinking") is not True
|
||||
]
|
||||
assert not missing, f"missing supports_adaptive_thinking: {missing}"
|
||||
|
|
|
|||
|
|
@ -385,6 +385,51 @@ def test_handle_realtime_stream_cost_calculation():
|
|||
)
|
||||
assert cost == 0.0 # No usage, no cost
|
||||
|
||||
|
||||
def test_realtime_stream_combines_text_and_audio_token_details():
|
||||
"""Realtime response.done usage with input_token_details / output_token_details."""
|
||||
from litellm.cost_calculator import RealtimeAPITokenUsageProcessor
|
||||
|
||||
results: OpenAIRealtimeStreamList = [
|
||||
{"type": "session.created", "session": {"model": "gpt-4o-realtime-preview"}},
|
||||
{
|
||||
"type": "response.done",
|
||||
"response": {
|
||||
"usage": {
|
||||
"input_tokens": 10,
|
||||
"output_tokens": 20,
|
||||
"total_tokens": 30,
|
||||
"input_token_details": {"text_tokens": 8, "audio_tokens": 2},
|
||||
"output_token_details": {"text_tokens": 12, "audio_tokens": 8},
|
||||
}
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "response.done",
|
||||
"response": {
|
||||
"usage": {
|
||||
"input_tokens": 5,
|
||||
"output_tokens": 15,
|
||||
"total_tokens": 20,
|
||||
"input_token_details": {"text_tokens": 3, "audio_tokens": 2},
|
||||
"output_token_details": {"text_tokens": 5, "audio_tokens": 10},
|
||||
}
|
||||
},
|
||||
},
|
||||
]
|
||||
|
||||
combined = RealtimeAPITokenUsageProcessor.collect_and_combine_usage_from_realtime_stream_results(
|
||||
results=results,
|
||||
)
|
||||
|
||||
assert combined.prompt_tokens_details is not None
|
||||
assert combined.prompt_tokens_details.text_tokens == 11
|
||||
assert combined.prompt_tokens_details.audio_tokens == 4
|
||||
|
||||
assert combined.completion_tokens_details is not None
|
||||
assert combined.completion_tokens_details.text_tokens == 17
|
||||
assert combined.completion_tokens_details.audio_tokens == 18
|
||||
|
||||
|
||||
def test_realtime_logging_object_allows_null_transcript_in_conversation_item_added():
|
||||
results: OpenAIRealtimeStreamList = [
|
||||
|
|
|
|||
|
|
@ -3,7 +3,9 @@ Test case normalization in LitellmParams for all guardrail types
|
|||
"""
|
||||
|
||||
import pytest
|
||||
from litellm.types.guardrails import LitellmParams
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm.types.guardrails import BaseLitellmParams, LitellmParams
|
||||
|
||||
|
||||
class TestLitellmParamsCaseNormalization:
|
||||
|
|
@ -89,3 +91,66 @@ class TestLitellmParamsCaseNormalization:
|
|||
)
|
||||
assert params.on_disallowed_action in ["block", "rewrite"]
|
||||
assert params.on_disallowed_action.islower()
|
||||
|
||||
|
||||
class TestSensitiveDataRoutingValidation:
|
||||
"""on_sensitive_data='route' requires a target model to be set"""
|
||||
|
||||
def test_route_with_target_model_is_valid(self):
|
||||
params = LitellmParams(
|
||||
guardrail="presidio",
|
||||
mode="pre_call",
|
||||
on_sensitive_data="route",
|
||||
sensitive_data_route_to_model="on-prem-model",
|
||||
)
|
||||
assert params.on_sensitive_data == "route"
|
||||
assert params.sensitive_data_route_to_model == "on-prem-model"
|
||||
|
||||
def test_route_without_target_model_raises(self):
|
||||
with pytest.raises(ValidationError, match="sensitive_data_route_to_model"):
|
||||
LitellmParams(
|
||||
guardrail="presidio",
|
||||
mode="pre_call",
|
||||
on_sensitive_data="route",
|
||||
)
|
||||
|
||||
def test_base_params_route_without_target_model_raises(self):
|
||||
with pytest.raises(ValidationError, match="sensitive_data_route_to_model"):
|
||||
BaseLitellmParams(on_sensitive_data="route")
|
||||
|
||||
def test_base_params_normalize_on_sensitive_data_case(self):
|
||||
params = BaseLitellmParams(
|
||||
on_sensitive_data="Route",
|
||||
sensitive_data_route_to_model="on-prem-model",
|
||||
)
|
||||
assert params.on_sensitive_data == "route"
|
||||
|
||||
def test_base_params_capitalized_route_without_target_model_raises(self):
|
||||
with pytest.raises(ValidationError, match="sensitive_data_route_to_model"):
|
||||
BaseLitellmParams(on_sensitive_data="ROUTE")
|
||||
|
||||
def test_block_without_target_model_is_valid(self):
|
||||
params = LitellmParams(
|
||||
guardrail="presidio",
|
||||
mode="pre_call",
|
||||
on_sensitive_data="block",
|
||||
)
|
||||
assert params.on_sensitive_data == "block"
|
||||
assert params.sensitive_data_route_to_model is None
|
||||
|
||||
def test_on_sensitive_data_is_case_normalized(self):
|
||||
params = LitellmParams(
|
||||
guardrail="presidio",
|
||||
mode="pre_call",
|
||||
on_sensitive_data="Route",
|
||||
sensitive_data_route_to_model="on-prem-model",
|
||||
)
|
||||
assert params.on_sensitive_data == "route"
|
||||
|
||||
def test_on_sensitive_data_uppercase_block_normalized(self):
|
||||
params = LitellmParams(
|
||||
guardrail="presidio",
|
||||
mode="pre_call",
|
||||
on_sensitive_data="BLOCK",
|
||||
)
|
||||
assert params.on_sensitive_data == "block"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue