mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
feat(agents): authenticate Entra identities and delegated requests (#43722)
* feat(agents): authenticate Entra identities and delegated requests * fix(agents): enforce target policy and preserve trusted authentication * fix(proxy): make inference model selection exhaustive * fix(agents): authorize targets against current database policy * fix(proxy): make exhaustive model resolution return explicit --------- Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
This commit is contained in:
parent
2bf0cddcc1
commit
405ed414cb
42 changed files with 2383 additions and 265 deletions
|
|
@ -170,6 +170,11 @@ async def identity_from_subject_token(
|
|||
return _refusal_for(denied, denied.message)
|
||||
except Exception as denied: # noqa: BLE001 # auth_jwt raises a plain Exception on signature and claim failures
|
||||
return _refusal_for(denied, denied)
|
||||
if result.get("agent_id") is not None:
|
||||
return SubjectTokenRefusal(
|
||||
error="invalid_request",
|
||||
description="Agent tokens require direct JWT authentication; this exchange supports users only",
|
||||
)
|
||||
user_id: Final = result["user_id"]
|
||||
if user_id is None:
|
||||
return SubjectTokenRefusal(error="invalid_request", description="subject_token names no user the gateway knows")
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ from dataclasses import dataclass
|
|||
from datetime import datetime
|
||||
from traceback import walk_tb
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict
|
||||
from uuid import uuid4
|
||||
|
||||
import anyio
|
||||
|
|
@ -14,6 +14,7 @@ import httpx2
|
|||
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
|
||||
from pydantic import ValidationError
|
||||
from starlette.datastructures import Headers
|
||||
from typing_extensions import ReadOnly
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import MCP_CLIENT_TIMEOUT, MCP_TOOL_LISTING_TIMEOUT
|
||||
|
|
@ -63,7 +64,27 @@ if TYPE_CHECKING:
|
|||
from litellm.proxy.utils import ProxyLogging
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.utils import CallTypes
|
||||
from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall
|
||||
|
||||
|
||||
class _MCPModelMetadata(TypedDict):
|
||||
model_group: ReadOnly[str]
|
||||
|
||||
|
||||
def _stamp_mcp_tool_metadata(logging_obj: "LiteLLMLoggingObj | None", server_id: str, tool_name: str) -> None:
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
if logging_obj is None:
|
||||
return
|
||||
server: Final = global_mcp_server_manager.get_mcp_server_by_id(
|
||||
server_id
|
||||
) or global_mcp_server_manager.get_mcp_server_by_name(server_id)
|
||||
metadata: Final[StandardLoggingMCPToolCall] = {
|
||||
"name": tool_name,
|
||||
"mcp_server_name": server.name if server is not None else server_id,
|
||||
}
|
||||
logging_obj.model_call_details["mcp_tool_call_metadata"] = metadata
|
||||
|
||||
|
||||
MCP_AVAILABLE: bool = True
|
||||
try:
|
||||
|
|
@ -1193,6 +1214,12 @@ if MCP_AVAILABLE:
|
|||
},
|
||||
)
|
||||
|
||||
data["model"] = f"MCP: {tool_name}"
|
||||
model_metadata: Final[_MCPModelMetadata] = {
|
||||
**(data.get("metadata") or MappingProxyType({})),
|
||||
"model_group": f"MCP: {tool_name}",
|
||||
}
|
||||
data["metadata"] = model_metadata
|
||||
proxy_base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data)
|
||||
_request_start_time: Final = datetime.now() # noqa: DTZ005 # naive to match the tool start time below
|
||||
try:
|
||||
|
|
@ -1226,6 +1253,8 @@ if MCP_AVAILABLE:
|
|||
if "metadata" in data and "user_api_key_auth" in data["metadata"]:
|
||||
data["user_api_key_auth"] = data["metadata"]["user_api_key_auth"]
|
||||
|
||||
_stamp_mcp_tool_metadata(logging_obj, server_id, tool_name)
|
||||
|
||||
# Resolve allowed MCP servers with IP filtering
|
||||
(
|
||||
allowed_mcp_servers,
|
||||
|
|
|
|||
|
|
@ -3320,6 +3320,7 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
|
|||
# single-owner so its meaning stays trustworthy.
|
||||
mcp_session_resource_server_id: str | None = Field(default=None, exclude=True)
|
||||
mcp_toolset_id: str | None = Field(default=None, exclude=True)
|
||||
authenticated_by_custom_auth: bool = Field(default=False, exclude=True)
|
||||
via_virtual_key: bool = Field(
|
||||
default=False,
|
||||
exclude=True,
|
||||
|
|
@ -3381,6 +3382,7 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
|
|||
values.pop("mcp_session_resource_server_id", None)
|
||||
values.pop("mcp_toolset_id", None)
|
||||
values.pop("via_virtual_key", None)
|
||||
values.pop("authenticated_by_custom_auth", None)
|
||||
values.pop("agent_caller", None)
|
||||
values.pop("managed_agent_context", None)
|
||||
values.pop("managed_agent_policy", None)
|
||||
|
|
|
|||
|
|
@ -722,6 +722,8 @@ async def invoke_agent_a2a(
|
|||
detail=f"Agent '{agent_id}' is not allowed for your key/team. Contact proxy admin for access.",
|
||||
)
|
||||
|
||||
user_api_key_dict.invoked_agent_id = agent.agent_id
|
||||
|
||||
_enforce_inbound_trace_id(agent, request)
|
||||
|
||||
# Get backend URL and agent name
|
||||
|
|
@ -759,6 +761,10 @@ async def invoke_agent_a2a(
|
|||
if "metadata" not in body:
|
||||
body["metadata"] = {}
|
||||
body["metadata"]["agent_id"] = agent.agent_id
|
||||
body["metadata"]["model_group"] = f"a2a_agent/{agent_name}"
|
||||
body["metadata"]["model_info"] = { # mutable-ok: request hooks mutate metadata before JSON logging
|
||||
"id": agent.agent_id
|
||||
}
|
||||
body["agent_id"] = agent.agent_id
|
||||
|
||||
body.update(
|
||||
|
|
@ -862,6 +868,7 @@ async def invoke_agent_a2a(
|
|||
# results written by the unified_guardrail hook are captured.
|
||||
logging_obj._defer_async_logging = True
|
||||
response = await asend_message(
|
||||
model=f"a2a_agent/{agent_name}",
|
||||
request=a2a_request,
|
||||
api_base=agent_url,
|
||||
litellm_params=litellm_params,
|
||||
|
|
|
|||
|
|
@ -57,7 +57,7 @@ async def route_a2a_agent_request(
|
|||
user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
|
||||
)
|
||||
if not is_admin:
|
||||
if not is_admin or agent.identity_managed:
|
||||
is_allowed: Final = await AgentRequestHandler.is_agent_allowed(
|
||||
agent_id=agent.agent_id,
|
||||
user_api_key_auth=user_api_key_dict,
|
||||
|
|
|
|||
|
|
@ -177,11 +177,10 @@ class AgentRequestHandler:
|
|||
|
||||
registered: Final = global_agent_registry.get_agent_by_id(agent_id)
|
||||
registry_managed: Final = isinstance(registered, AgentResponse) and registered.identity_managed
|
||||
if registry_managed or (registered is None and prisma_client is not None):
|
||||
if registry_managed or prisma_client is not None:
|
||||
target: Final = await AgentIdentityStore.from_client(prisma_client).agent(agent_id)
|
||||
if isinstance(target, AgentIdentityFailure):
|
||||
if registry_managed:
|
||||
raise_identity_failure(target)
|
||||
raise_identity_failure(target)
|
||||
elif target is None and registry_managed:
|
||||
return False
|
||||
elif isinstance(target, AgentResponse) and target.identity_managed:
|
||||
|
|
@ -200,6 +199,7 @@ class AgentRequestHandler:
|
|||
if key_hash
|
||||
and managed_agent_policy(user_api_key_auth) is None
|
||||
and not user_api_key_auth.is_session_token
|
||||
and not user_api_key_auth.authenticated_by_custom_auth
|
||||
else user_api_key_auth
|
||||
)
|
||||
fresh_auth: Final = authority.model_copy(
|
||||
|
|
@ -678,14 +678,44 @@ async def _managed_actor_agent_access(auth: UserAPIKeyAuth) -> AgentAccess:
|
|||
return RestrictedAgentAccess(capped.intersection(human_ids))
|
||||
|
||||
|
||||
async def verified_human_agent_grants(user_id: str | None, team_id: str | None = None) -> frozenset[str]:
|
||||
async def _verified_human_agent_sources(
|
||||
user_id: str | None, *, allowed_team_ids: frozenset[str] | None = None
|
||||
) -> tuple[tuple[str | None, frozenset[str]], ...]:
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
|
||||
|
||||
if user_id is None:
|
||||
return frozenset()
|
||||
return ()
|
||||
human: Final = await MCPRequestHandler.reload_admitted_user(user_id, requires_fresh_policy=True)
|
||||
sources: Final = await MCPRequestHandler.admitted_subject_sources(
|
||||
human, allowed_team_ids=frozenset((team_id,)) if team_id else frozenset()
|
||||
sources: Final = await MCPRequestHandler.admitted_subject_sources(human, allowed_team_ids=allowed_team_ids)
|
||||
access: Final = await asyncio.gather(*(_strict_agent_access(source) for source in sources))
|
||||
return tuple((source.team_id, _granted_ids(grant)) for source, grant in zip(sources, access, strict=True))
|
||||
|
||||
|
||||
async def verified_human_agent_grants(user_id: str | None, team_id: str | None = None) -> frozenset[str]:
|
||||
sources: Final = await _verified_human_agent_sources(
|
||||
user_id, allowed_team_ids=frozenset((team_id,)) if team_id else frozenset()
|
||||
)
|
||||
human_access: Final = await asyncio.gather(*(_strict_agent_access(source) for source in sources))
|
||||
return frozenset().union(*(_granted_ids(access) for access in human_access))
|
||||
return frozenset().union(*(grants for source, grants in sources if source is None or source == team_id))
|
||||
|
||||
|
||||
async def resolve_delegated_agent_team(
|
||||
user_id: str | None,
|
||||
agent_id: str,
|
||||
team_id: str | None,
|
||||
*,
|
||||
explicit_team: bool,
|
||||
allowed_team_ids: frozenset[str] | None = None,
|
||||
) -> str | None:
|
||||
sources: Final = await _verified_human_agent_sources(user_id)
|
||||
if any(source is None and agent_id in grants for source, grants in sources):
|
||||
return team_id
|
||||
granting_teams: Final = frozenset(
|
||||
source
|
||||
for source, grants in sources
|
||||
if source is not None and agent_id in grants and (allowed_team_ids is None or source in allowed_team_ids)
|
||||
)
|
||||
if team_id in granting_teams:
|
||||
return team_id
|
||||
if not explicit_team and granting_teams:
|
||||
return min(granting_teams)
|
||||
raise HTTPException(403, "Select a team that grants access to this agent using x-litellm-team-id")
|
||||
|
|
|
|||
|
|
@ -1,11 +1,129 @@
|
|||
from typing import Final
|
||||
from collections.abc import Mapping
|
||||
from itertools import product
|
||||
from types import MappingProxyType
|
||||
from typing import Annotated, Final, Literal
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from pydantic import Field, TypeAdapter, ValidationError
|
||||
|
||||
from litellm.proxy._types import LiteLLMRoutes, UserAPIKeyAuth
|
||||
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
|
||||
from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityFailure, ManagedAgentContext
|
||||
|
||||
_MANAGED_REALTIME_ROUTES: Final = frozenset(("/realtime", "/v1/realtime", "/openai/v1/realtime"))
|
||||
_MANAGED_MODEL_ROUTES: Final = frozenset(
|
||||
f"{prefix}/{operation}"
|
||||
for prefix, operation in product(
|
||||
("", "/v1"),
|
||||
(
|
||||
"chat/completions",
|
||||
"completions",
|
||||
"embeddings",
|
||||
"responses",
|
||||
"messages",
|
||||
"messages/count_tokens",
|
||||
"images/generations",
|
||||
"images/edits",
|
||||
"audio/transcriptions",
|
||||
"audio/speech",
|
||||
"moderations",
|
||||
"rerank",
|
||||
"ocr",
|
||||
),
|
||||
)
|
||||
) | frozenset(
|
||||
(
|
||||
"/openai/v1/responses",
|
||||
"/v2/rerank",
|
||||
"/claude_code_gateway/v1/messages",
|
||||
"/claude_code_gateway/v1/messages/count_tokens",
|
||||
"/cursor/chat/completions",
|
||||
)
|
||||
)
|
||||
_MANAGED_MODEL_PATHS: Final = (
|
||||
"/engines/{model:path}/chat/completions",
|
||||
"/engines/{model:path}/completions",
|
||||
"/engines/{model:path}/embeddings",
|
||||
"/openai/deployments/{model:path}/chat/completions",
|
||||
"/openai/deployments/{model:path}/completions",
|
||||
"/openai/deployments/{model:path}/embeddings",
|
||||
"/openai/deployments/{model:path}/images/generations",
|
||||
"/openai/deployments/{model:path}/images/edits",
|
||||
"/v1beta/models/{model_name:path}:countTokens",
|
||||
"/v1beta/models/{model_name:path}:generateContent",
|
||||
"/v1beta/models/{model_name:path}:streamGenerateContent",
|
||||
"/models/{model_name:path}:countTokens",
|
||||
"/models/{model_name:path}:generateContent",
|
||||
"/models/{model_name:path}:streamGenerateContent",
|
||||
)
|
||||
_MANAGED_MCP_ROUTES: Final = tuple(
|
||||
route for route in LiteLLMRoutes.mcp_inference_routes.value if route not in ("/token", "/introspect")
|
||||
)
|
||||
|
||||
|
||||
_MODEL_ROUTE_KINDS: Final[
|
||||
Mapping[str, Literal["image_generation", "image_edit", "moderation", "speech", "body", "path"]]
|
||||
] = MappingProxyType(
|
||||
{
|
||||
"/images/generations": "image_generation",
|
||||
"/images/edits": "image_edit",
|
||||
"/moderations": "moderation",
|
||||
"/audio/transcriptions": "moderation",
|
||||
"/audio/speech": "speech",
|
||||
"/rerank": "body",
|
||||
"/messages/count_tokens": "body",
|
||||
":countTokens": "path",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def managed_agent_route_allowed(route: str, method: str | None) -> bool:
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
|
||||
if route in ("/agents", "/v1/agents"):
|
||||
return method in (None, "GET", "HEAD")
|
||||
if route in _MANAGED_REALTIME_ROUTES:
|
||||
return method in (None, "GET")
|
||||
if route in _MANAGED_MODEL_ROUTES or RouteChecks.check_route_access(route, _MANAGED_MODEL_PATHS):
|
||||
return method in (None, "POST")
|
||||
return RouteChecks.check_route_access(route, _MANAGED_MCP_ROUTES) or RouteChecks.check_route_access(
|
||||
route, LiteLLMRoutes.agent_inference_routes.value
|
||||
)
|
||||
|
||||
|
||||
def managed_inference_request(
|
||||
route: str,
|
||||
body: Mapping[str, object],
|
||||
settings: Mapping[str, object],
|
||||
cli_model: str | None,
|
||||
path_model: object = None,
|
||||
query_model: object = None,
|
||||
) -> dict[str, object]:
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
|
||||
if route in _MANAGED_REALTIME_ROUTES:
|
||||
model: Final = query_model or body.get("model")
|
||||
if not isinstance(model, str) or not model:
|
||||
raise_identity_failure(
|
||||
AgentIdentityFailure(message="Managed inference requires an explicit or configured model")
|
||||
)
|
||||
return {**body, "model": model} # mutable-ok: centralized auth hooks add request tags and budget metadata
|
||||
if route not in _MANAGED_MODEL_ROUTES and not RouteChecks.check_route_access(route, _MANAGED_MODEL_PATHS):
|
||||
return dict(body) # mutable-ok: centralized auth hooks add request tags and budget metadata
|
||||
from litellm.proxy.common_utils.http_parsing_utils import resolve_inference_model
|
||||
|
||||
kind: Final = next((kind for suffix, kind in _MODEL_ROUTE_KINDS.items() if route.endswith(suffix)), "completion")
|
||||
endpoint_model: Final = path_model or (
|
||||
query_model if route.endswith(("/completions", "/embeddings", "/images/generations", "/images/edits")) else None
|
||||
)
|
||||
effective: Final = resolve_inference_model(body.get("model"), settings, cli_model, endpoint_model, kind=kind)
|
||||
if not isinstance(effective, str) or not effective:
|
||||
raise_identity_failure(
|
||||
AgentIdentityFailure(message="Managed inference requires an explicit or configured model")
|
||||
)
|
||||
return {**body, "model": effective} # mutable-ok: centralized auth hooks add request tags and budget metadata
|
||||
|
||||
|
||||
def managed_agent_policy(auth: "UserAPIKeyAuth | None") -> AgentResponse | None:
|
||||
"""The admitted managed policy, or ``None`` when the subject was never admitted as a managed agent.
|
||||
|
|
@ -82,3 +200,53 @@ def actor_admission_failure(
|
|||
if context.mode == "delegated" and not context.user_id:
|
||||
return AgentIdentityFailure(message="A verified human subject is required")
|
||||
return None
|
||||
|
||||
|
||||
_INVOCATION_COST: Final = TypeAdapter(Annotated[float, Field(ge=0, allow_inf_nan=False)])
|
||||
|
||||
|
||||
def invocation_target(route: str, body: Mapping[str, object]) -> str | None:
|
||||
components: Final = tuple(route.strip("/").split("/"))
|
||||
path: Final = components[1:] if components and components[0] == "v1" else components
|
||||
if len(path) >= 2 and path[0] == "a2a":
|
||||
return path[1] or None
|
||||
model: Final = body.get("model")
|
||||
return model.removeprefix("a2a/") or None if isinstance(model, str) and model.startswith("a2a/") else None
|
||||
|
||||
|
||||
async def prepare_agent_invocation(
|
||||
auth: UserAPIKeyAuth, target_name: str, store: AgentIdentityStore | None, *, billable: bool = True
|
||||
) -> None:
|
||||
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import AgentRequestHandler
|
||||
from litellm.proxy.common_utils.registry_read_through import get_agent_with_read_through
|
||||
|
||||
registered: Final = await get_agent_with_read_through(target_name)
|
||||
if registered is None:
|
||||
return
|
||||
registered_managed: Final = registered.identity_managed or registered.identity is not None
|
||||
if store is None and registered_managed:
|
||||
raise_identity_failure(
|
||||
AgentIdentityFailure(code="policy_unavailable", message="Managed agent policy requires a database")
|
||||
)
|
||||
target: Final = await store.agent(registered.agent_id) if store is not None else None
|
||||
if isinstance(target, AgentIdentityFailure):
|
||||
raise_identity_failure(target)
|
||||
if target is None and registered_managed:
|
||||
raise_identity_failure(AgentIdentityFailure(message="Invoked agent no longer exists"))
|
||||
effective: Final = target if target is not None else registered
|
||||
if not effective.identity_managed and auth.managed_agent_policy is None:
|
||||
return
|
||||
if not await AgentRequestHandler.is_agent_allowed(effective.agent_id, auth):
|
||||
raise_identity_failure(AgentIdentityFailure(message="The caller is not permitted to invoke this agent"))
|
||||
auth.invoked_agent_id = effective.agent_id
|
||||
auth.invoked_agent_policy = effective
|
||||
if auth.agent_id is None and effective.identity_managed:
|
||||
auth.billing_agent_policy = effective
|
||||
raw_fee: Final = (effective.litellm_params or MappingProxyType({})).get("cost_per_query", 0.0) if billable else 0.0
|
||||
try:
|
||||
fee: Final = _INVOCATION_COST.validate_python(raw_fee)
|
||||
except ValidationError:
|
||||
raise_identity_failure(
|
||||
AgentIdentityFailure(code="policy_unavailable", message="Agent invocation price is invalid")
|
||||
)
|
||||
auth.agent_invocation_cost = fee
|
||||
|
|
|
|||
|
|
@ -52,6 +52,10 @@ from litellm.proxy._types import (
|
|||
TeamMemberAddRequest,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_route_allowed
|
||||
from litellm.proxy.agent_endpoints.identity import has_legacy_identity
|
||||
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore, resolve_managed_agent
|
||||
from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
|
||||
from litellm.proxy.auth.auth_checks import can_team_access_model
|
||||
from litellm.proxy.auth.model_access_denied import (
|
||||
ModelAccessDeniedHTTPException,
|
||||
|
|
@ -67,6 +71,7 @@ from litellm.proxy.common_utils.user_api_key_cache import (
|
|||
from litellm.proxy.utils import PrismaClient, ProxyLogging
|
||||
from litellm.repositories.user_repository import UserRepository
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityFailure
|
||||
from litellm.types.proxy.auth.auth_checks import UserNotFoundError
|
||||
|
||||
from .auth_checks import (
|
||||
|
|
@ -157,6 +162,8 @@ class HeaderTeam:
|
|||
class AgentLookup(Protocol):
|
||||
"""The registered-agent lookups a JWT agent claim is matched against."""
|
||||
|
||||
def get_agent_list(self) -> Sequence[AgentResponse]: ...
|
||||
|
||||
def get_agent_by_id(self, agent_id: str) -> AgentResponse | None:
|
||||
"""The agent registered under ``agent_id``, if any."""
|
||||
|
||||
|
|
@ -167,6 +174,9 @@ class AgentLookup(Protocol):
|
|||
class _NoRegisteredAgents:
|
||||
"""The lookup in force until the proxy binds its agent registry: no agent is registered, so no claim matches."""
|
||||
|
||||
def get_agent_list(self) -> tuple[AgentResponse, ...]:
|
||||
return ()
|
||||
|
||||
def get_agent_by_id(self, agent_id: str) -> None:
|
||||
return None
|
||||
|
||||
|
|
@ -398,7 +408,7 @@ class JWTHandler:
|
|||
|
||||
return []
|
||||
|
||||
def get_all_jwt_team_ids(self, token: dict) -> list[str]:
|
||||
def get_all_jwt_team_ids(self, token: dict[str, object]) -> list[str]:
|
||||
"""
|
||||
Return team IDs from both the plural ``team_ids_jwt_field`` and the
|
||||
singular ``team_id_jwt_field`` claim (string or list of strings), as a
|
||||
|
|
@ -522,7 +532,7 @@ class JWTHandler:
|
|||
team_id = default_value
|
||||
return team_id
|
||||
|
||||
def get_team_alias(self, token: dict, default_value: str | None) -> str | None:
|
||||
def get_team_alias(self, token: dict[str, object], default_value: str | None) -> str | None:
|
||||
"""
|
||||
Extract team name/alias from JWT token using the configured team_alias_jwt_field.
|
||||
|
||||
|
|
@ -1096,6 +1106,15 @@ class JWTHandler:
|
|||
"options": options or None,
|
||||
}
|
||||
|
||||
def managed_issuer_is_trusted(self, issuer: object) -> bool:
|
||||
if not isinstance(issuer, str):
|
||||
return False
|
||||
configured: Final = self.litellm_jwtauth.issuers or ()
|
||||
for item in configured:
|
||||
if item.issuer == issuer:
|
||||
return bool(item.audience) and not item.disable_audience_validation
|
||||
return issuer == os.getenv("JWT_ISSUER") and bool(os.getenv("JWT_AUDIENCE"))
|
||||
|
||||
def _get_configured_issuer(self, token: str) -> JWTIssuerConfig | None:
|
||||
litellm_jwtauth: Final[_JWTAuthSettings | None] = getattr(self, "litellm_jwtauth", None)
|
||||
if litellm_jwtauth is None:
|
||||
|
|
@ -1488,7 +1507,12 @@ class JWTAuthManager:
|
|||
agent: Final = agent_registry.get_agent_by_id(agent_id=agent_claim) or agent_registry.get_agent_by_name(
|
||||
agent_name=agent_claim
|
||||
)
|
||||
if agent is None:
|
||||
if (
|
||||
agent is None
|
||||
or agent.identity_managed
|
||||
or agent.identity is not None
|
||||
or has_legacy_identity(agent.litellm_params)
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=f"No registered agent matches JWT claim {jwt_handler.litellm_jwtauth.agent_id_jwt_field}={agent_claim}",
|
||||
|
|
@ -2159,7 +2183,7 @@ class JWTAuthManager:
|
|||
parent_otel_span: Span | None,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
team_id_upsert: bool | None,
|
||||
) -> tuple:
|
||||
) -> tuple[str | None, LiteLLM_TeamTable | None, LiteLLM_TeamMembership | None]:
|
||||
"""
|
||||
If JWT did not resolve team_id, but the user belongs to exactly one team
|
||||
in LiteLLM, load that team (and membership when user_id is set) so that
|
||||
|
|
@ -2478,12 +2502,39 @@ class JWTAuthManager:
|
|||
"""Resolve and authorize JWT context; only normal admission supplies provisioning."""
|
||||
handler: Final = jwt_handler
|
||||
jwt_valid_token: Final = await JWTAuthManager.authenticate_jwt(api_key, handler)
|
||||
managed: Final = await resolve_managed_agent(jwt_valid_token, prisma_client, cache=user_api_key_cache)
|
||||
if managed is not None:
|
||||
if not handler.managed_issuer_is_trusted(jwt_valid_token.get("iss")):
|
||||
raise HTTPException(403, "Managed agents require trusted JWT issuer and audience validation")
|
||||
if not managed_agent_route_allowed(route, request_method):
|
||||
raise HTTPException(403, "Agent identities can only access inference and agent discovery routes")
|
||||
evidence: Final = await AgentIdentityStore.from_client(prisma_client).record_authentication(managed)
|
||||
if isinstance(evidence, AgentIdentityFailure):
|
||||
raise_identity_failure(evidence)
|
||||
if managed.mode == "autonomous":
|
||||
return JWTAuthBuilderResult(
|
||||
is_proxy_admin=False,
|
||||
team_id=None,
|
||||
team_object=None,
|
||||
user_id=None,
|
||||
user_email=None,
|
||||
user_object=None,
|
||||
org_id=None,
|
||||
org_object=None,
|
||||
end_user_id=None,
|
||||
end_user_object=None,
|
||||
token=api_key,
|
||||
team_membership=None,
|
||||
jwt_claims=jwt_valid_token,
|
||||
agent_id=managed.agent_id,
|
||||
managed_agent_context=managed,
|
||||
)
|
||||
team_id_upsert: Final = provisioning.team_id_upsert if provisioning is not None else False
|
||||
model: Final = request_data.get("model")
|
||||
requested_model: Final = model if isinstance(model, str) else None
|
||||
|
||||
# Check RBAC
|
||||
rbac_role: Final = handler.get_rbac_role(token=jwt_valid_token)
|
||||
rbac_role: Final = handler.get_rbac_role(token=jwt_valid_token) if managed is None else None
|
||||
await JWTAuthManager.check_rbac_role(handler, jwt_valid_token, general_settings, request_data, route, rbac_role)
|
||||
|
||||
# Check Scope Based Access
|
||||
|
|
@ -2499,7 +2550,11 @@ class JWTAuthManager:
|
|||
object_id = handler.get_object_id(token=jwt_valid_token, default_value=None)
|
||||
|
||||
# Get basic user info
|
||||
user_id, user_email, valid_user_email = await JWTAuthManager.get_user_info(handler, jwt_valid_token)
|
||||
user_id, user_email, valid_user_email = (
|
||||
(managed.user_id, None, None)
|
||||
if managed is not None
|
||||
else await JWTAuthManager.get_user_info(handler, jwt_valid_token)
|
||||
)
|
||||
|
||||
# Get IDs
|
||||
org_id: Final = handler.get_org_id(token=jwt_valid_token, default_value=None)
|
||||
|
|
@ -2514,23 +2569,31 @@ class JWTAuthManager:
|
|||
elif rbac_role == LitellmUserRoles.INTERNAL_USER:
|
||||
user_id = object_id
|
||||
|
||||
agent_id: Final = JWTAuthManager.resolve_agent_id(
|
||||
jwt_handler=handler,
|
||||
jwt_valid_token=jwt_valid_token,
|
||||
agent_registry=handler.agent_lookup,
|
||||
agent_id: Final = (
|
||||
managed.agent_id
|
||||
if managed is not None
|
||||
else JWTAuthManager.resolve_agent_id(
|
||||
jwt_handler=handler,
|
||||
jwt_valid_token=jwt_valid_token,
|
||||
agent_registry=handler.agent_lookup,
|
||||
)
|
||||
)
|
||||
|
||||
# Check admin access
|
||||
admin_result: Final = await JWTAuthManager.check_admin_access(
|
||||
handler,
|
||||
scopes,
|
||||
route,
|
||||
user_id,
|
||||
org_id,
|
||||
api_key,
|
||||
jwt_valid_token,
|
||||
user_email=user_email,
|
||||
agent_id=agent_id,
|
||||
admin_result: Final = (
|
||||
None
|
||||
if managed is not None
|
||||
else await JWTAuthManager.check_admin_access(
|
||||
handler,
|
||||
scopes,
|
||||
route,
|
||||
user_id,
|
||||
org_id,
|
||||
api_key,
|
||||
jwt_valid_token,
|
||||
user_email=user_email,
|
||||
agent_id=agent_id,
|
||||
)
|
||||
)
|
||||
if admin_result:
|
||||
await JWTAuthManager._attach_team_from_header_for_admin(
|
||||
|
|
@ -2673,8 +2736,47 @@ class JWTAuthManager:
|
|||
team_id_upsert=team_id_upsert,
|
||||
)
|
||||
|
||||
if team_id and not JWTAuthManager._team_has_passthrough_route_access(
|
||||
team_object=team_object,
|
||||
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import resolve_delegated_agent_team
|
||||
|
||||
claimed_teams: Final[frozenset[str]] = (
|
||||
frozenset(handler.get_all_jwt_team_ids(jwt_valid_token)) if managed is not None else frozenset()
|
||||
)
|
||||
scoped_teams: Final[frozenset[str] | None] = claimed_teams or (
|
||||
frozenset((team_id,))
|
||||
if managed is not None and team_id and handler.get_team_alias(jwt_valid_token, default_value=None)
|
||||
else None
|
||||
)
|
||||
granting_team: Final = (
|
||||
await resolve_delegated_agent_team(
|
||||
managed.user_id,
|
||||
managed.agent_id,
|
||||
team_id,
|
||||
explicit_team=header_team is not None,
|
||||
allowed_team_ids=None if handler.litellm_jwtauth.fallback_to_db_teams else scoped_teams,
|
||||
)
|
||||
if managed is not None
|
||||
else team_id
|
||||
)
|
||||
if granting_team is not None and granting_team != team_id:
|
||||
if not JWTAuthManager._is_team_route_allowed(route, request_method, handler):
|
||||
raise HTTPException(403, "The granting team is not allowed to access this route")
|
||||
|
||||
selected_team_id: Final[str | None] = granting_team if granting_team is not None else team_id
|
||||
selected_team_object: Final[LiteLLM_TeamTable | None] = (
|
||||
await get_team_object(
|
||||
team_id=selected_team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=True,
|
||||
)
|
||||
if selected_team_id is not None and selected_team_id != team_id
|
||||
else team_object
|
||||
)
|
||||
|
||||
if selected_team_id and not JWTAuthManager._team_has_passthrough_route_access(
|
||||
team_object=selected_team_object,
|
||||
route=route,
|
||||
request_method=request_method,
|
||||
team_allowed_routes=handler.litellm_jwtauth.team_allowed_routes,
|
||||
|
|
@ -2696,7 +2798,7 @@ class JWTAuthManager:
|
|||
user_email=user_email,
|
||||
org_id=org_id,
|
||||
end_user_id=end_user_id,
|
||||
team_id=team_id,
|
||||
team_id=selected_team_id,
|
||||
valid_user_email=valid_user_email,
|
||||
jwt_handler=handler,
|
||||
prisma_client=prisma_client,
|
||||
|
|
@ -2705,13 +2807,13 @@ class JWTAuthManager:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
route=route,
|
||||
org_alias=org_alias,
|
||||
user_id_upsert=provisioning.user_id_upsert if provisioning is not None else False,
|
||||
user_id_upsert=provisioning.user_id_upsert if provisioning is not None and managed is None else False,
|
||||
)
|
||||
|
||||
# Derive org_id from org_object if resolved by alias
|
||||
resolved_org_id: Final = org_object.organization_id if org_object else org_id
|
||||
|
||||
if provisioning is not None:
|
||||
if provisioning is not None and managed is None:
|
||||
await JWTAuthManager.sync_user_role_and_teams(
|
||||
jwt_handler=handler,
|
||||
jwt_valid_token=jwt_valid_token,
|
||||
|
|
@ -2721,7 +2823,7 @@ class JWTAuthManager:
|
|||
)
|
||||
|
||||
# If JWT did not resolve team_id, attempt a team fallback.
|
||||
if team_id is None and db_team_fallback:
|
||||
if selected_team_id is None and db_team_fallback:
|
||||
(
|
||||
team_id,
|
||||
team_object,
|
||||
|
|
@ -2750,7 +2852,7 @@ class JWTAuthManager:
|
|||
team_allowed_routes=handler.litellm_jwtauth.team_allowed_routes,
|
||||
):
|
||||
JWTAuthManager._raise_team_passthrough_route_denial(route=route)
|
||||
elif team_id is None:
|
||||
elif selected_team_id is None:
|
||||
(
|
||||
team_id,
|
||||
team_object,
|
||||
|
|
@ -2764,9 +2866,9 @@ class JWTAuthManager:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
team_id_upsert=team_id_upsert,
|
||||
)
|
||||
elif provisional_header_team is not None and team_id == provisional_header_team.team_id:
|
||||
elif provisional_header_team is not None and selected_team_id == provisional_header_team.team_id:
|
||||
JWTAuthManager._validate_header_team_in_db_membership(
|
||||
team_id=team_id,
|
||||
team_id=selected_team_id,
|
||||
user_object=user_object,
|
||||
header_value=provisional_header_team.header_value,
|
||||
)
|
||||
|
|
@ -2783,28 +2885,35 @@ class JWTAuthManager:
|
|||
),
|
||||
)
|
||||
|
||||
authorized_team_id: Final[str | None] = selected_team_id if selected_team_id is not None else team_id
|
||||
authorized_team_object: Final[LiteLLM_TeamTable | None] = (
|
||||
selected_team_object if selected_team_id is not None else team_object
|
||||
)
|
||||
|
||||
## MAP USER TO TEAMS
|
||||
if provisioning is not None:
|
||||
if provisioning is not None and managed is None:
|
||||
await JWTAuthManager.map_user_to_teams(
|
||||
user_object=user_object,
|
||||
team_object=team_object,
|
||||
team_object=authorized_team_object,
|
||||
)
|
||||
|
||||
# Validate that a valid rbac id is returned for spend tracking
|
||||
JWTAuthManager.validate_object_id(
|
||||
user_id=user_id,
|
||||
team_id=team_id,
|
||||
team_id=authorized_team_id,
|
||||
enforce_rbac=bool(general_settings.get("enforce_rbac", False)),
|
||||
is_proxy_admin=False,
|
||||
)
|
||||
|
||||
# check if user is proxy admin
|
||||
is_proxy_admin: Final = bool(user_object and user_object.user_role == LitellmUserRoles.PROXY_ADMIN)
|
||||
is_proxy_admin: Final = managed is None and bool(
|
||||
user_object and user_object.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
)
|
||||
|
||||
return JWTAuthBuilderResult(
|
||||
is_proxy_admin=is_proxy_admin,
|
||||
team_id=team_id,
|
||||
team_object=team_object,
|
||||
team_id=authorized_team_id,
|
||||
team_object=authorized_team_object,
|
||||
user_id=user_id,
|
||||
user_email=(user_object.user_email if user_object is not None and user_object.user_email else user_email),
|
||||
user_object=user_object,
|
||||
|
|
@ -2816,6 +2925,7 @@ class JWTAuthManager:
|
|||
team_membership=team_membership_object,
|
||||
jwt_claims=jwt_valid_token,
|
||||
agent_id=agent_id,
|
||||
managed_agent_context=managed,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -2826,11 +2936,13 @@ class JWTAuthManager:
|
|||
"""Keep JWT identity and permission attribution identical across consumers."""
|
||||
user: Final = result["user_object"]
|
||||
admin: Final = result["is_proxy_admin"]
|
||||
return UserAPIKeyAuth(
|
||||
auth: Final = UserAPIKeyAuth(
|
||||
api_key=None,
|
||||
user_role=(
|
||||
LitellmUserRoles.PROXY_ADMIN
|
||||
if admin
|
||||
else LitellmUserRoles.INTERNAL_USER
|
||||
if result.get("managed_agent_context") is not None
|
||||
else LitellmUserRoles(user.user_role)
|
||||
if user is not None and user.user_role is not None
|
||||
else LitellmUserRoles.INTERNAL_USER
|
||||
|
|
@ -2852,3 +2964,8 @@ class JWTAuthManager:
|
|||
user_id=result["user_id"],
|
||||
),
|
||||
)
|
||||
auth.managed_agent_context = result.get("managed_agent_context")
|
||||
auth._managed_delegation_verified = ( # pyright: ignore[reportPrivateUsage] # JWT admission produces the one-shot proof consumed by managed authorization
|
||||
auth.managed_agent_context is not None and auth.managed_agent_context.mode == "delegated"
|
||||
)
|
||||
return auth
|
||||
|
|
|
|||
|
|
@ -655,6 +655,8 @@ async def user_api_key_auth_websocket_for_model(websocket: WebSocket, model: str
|
|||
# never reaches the fallback.
|
||||
synthetic_scope: Final[dict[str, Any]] = {
|
||||
"type": "http",
|
||||
"method": "GET",
|
||||
"query_string": ws_scope.get("query_string", b""),
|
||||
"headers": scope_headers,
|
||||
"path": ws_scope.get("path", ""),
|
||||
"state": ws_scope.setdefault("state", {}), # mutable-ok: Starlette's socket state, shared with the request
|
||||
|
|
@ -1559,6 +1561,7 @@ async def _user_api_key_auth_builder(
|
|||
route=route,
|
||||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
validated.authenticated_by_custom_auth = True
|
||||
return validated
|
||||
elif response is not None and isinstance(response, str):
|
||||
api_key = response
|
||||
|
|
@ -1574,6 +1577,7 @@ async def _user_api_key_auth_builder(
|
|||
route=route,
|
||||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
validated.authenticated_by_custom_auth = True
|
||||
return validated
|
||||
|
||||
### LITELLM-DEFINED AUTH FUNCTION ###
|
||||
|
|
@ -1656,6 +1660,16 @@ async def _user_api_key_auth_builder(
|
|||
else:
|
||||
jwt_claims = await jwt_handler.auth_jwt(token=api_key)
|
||||
|
||||
from litellm.proxy.agent_endpoints.identity_store import resolve_managed_agent
|
||||
|
||||
if (
|
||||
jwt_claims
|
||||
and await resolve_managed_agent(jwt_claims, prisma_client, cache=user_api_key_cache) is not None
|
||||
):
|
||||
raise HTTPException(
|
||||
403, "Managed agents require direct JWT authentication without virtual-key mapping"
|
||||
)
|
||||
|
||||
resolve_result: Final = await _resolve_jwt_to_virtual_key(
|
||||
jwt_claims=jwt_claims,
|
||||
jwt_handler=jwt_handler,
|
||||
|
|
@ -3130,7 +3144,10 @@ async def _reserve_budget_after_common_checks(
|
|||
end_user_id=end_user_id,
|
||||
end_user_object=end_user_object,
|
||||
apply_user_budget_to_team_keys=general_settings.get("apply_user_budget_to_team_keys") is True,
|
||||
fail_closed_budget_enforcement=general_settings.get("fail_closed_budget_enforcement") is True,
|
||||
fail_closed_budget_enforcement=(
|
||||
general_settings.get("fail_closed_budget_enforcement") is True
|
||||
or user_api_key_auth_obj.billing_agent_policy is not None
|
||||
),
|
||||
raw_body=await read_raw_json_body(request=request),
|
||||
)
|
||||
if request is not None:
|
||||
|
|
@ -3204,19 +3221,48 @@ async def _authorize_authenticated_request(
|
|||
# admin-only-route / model-access / budget checks) surface as
|
||||
# ProxyException consistently with pre-refactor behavior.
|
||||
try:
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import admit_managed_actor
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import (
|
||||
admit_managed_actor,
|
||||
invocation_target,
|
||||
managed_agent_route_allowed,
|
||||
managed_inference_request,
|
||||
prepare_agent_invocation,
|
||||
)
|
||||
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
from litellm.proxy.proxy_server import general_settings, prisma_client, user_model
|
||||
|
||||
store: Final = AgentIdentityStore.from_client(prisma_client) if prisma_client is not None else None
|
||||
if user_api_key_auth_obj.agent_id is not None:
|
||||
await admit_managed_actor(
|
||||
await admit_managed_actor(user_api_key_auth_obj, store)
|
||||
if user_api_key_auth_obj.managed_agent_policy is not None and not managed_agent_route_allowed(
|
||||
route, request.method
|
||||
):
|
||||
raise HTTPException(403, "Agent identities can only access inference and agent discovery routes")
|
||||
authorized_data: Final = (
|
||||
managed_inference_request(
|
||||
route,
|
||||
request_data,
|
||||
general_settings,
|
||||
user_model,
|
||||
request.path_params.get("model") or request.path_params.get("model_name"),
|
||||
request.query_params.get("model"),
|
||||
)
|
||||
if user_api_key_auth_obj.managed_agent_policy is not None
|
||||
else request_data
|
||||
)
|
||||
target_name: Final = invocation_target(route, authorized_data)
|
||||
if target_name is not None:
|
||||
await prepare_agent_invocation(
|
||||
user_api_key_auth_obj,
|
||||
AgentIdentityStore.from_client(prisma_client) if prisma_client is not None else None,
|
||||
target_name,
|
||||
store,
|
||||
billable=request_data.get("method")
|
||||
in (None, "message/send", "message/stream", "SendMessage", "SendStreamingMessage"),
|
||||
)
|
||||
await _run_centralized_common_checks(
|
||||
user_api_key_auth_obj=user_api_key_auth_obj,
|
||||
request=request,
|
||||
request_data=request_data,
|
||||
request_data=authorized_data,
|
||||
route=route,
|
||||
)
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -93,6 +93,7 @@ from litellm.proxy.common_utils.error_body_call_id import JSON_OBJECT, error_bod
|
|||
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||
get_client_requested_model,
|
||||
get_tags_from_request_body,
|
||||
resolve_inference_model,
|
||||
)
|
||||
from litellm.proxy.common_utils.openai_error_payload import (
|
||||
LITELLM_CALL_ID_HEADER,
|
||||
|
|
@ -2068,11 +2069,12 @@ class ProxyBaseLLMRequestProcessing:
|
|||
if isinstance(model, str):
|
||||
reject_url_valued_destination("model", model)
|
||||
|
||||
self.data["model"] = (
|
||||
general_settings.get("completion_model", None) # server default
|
||||
or user_model # model name passed via cli args
|
||||
or model # for azure deployments
|
||||
or self.data.get("model", None) # default passed in http request
|
||||
self.data["model"] = resolve_inference_model(
|
||||
self.data.get("model"),
|
||||
general_settings,
|
||||
user_model,
|
||||
model,
|
||||
kind="image_edit" if route_type == "aimage_edit" else "completion",
|
||||
)
|
||||
|
||||
# override with user settings, these are params passed via cli
|
||||
|
|
|
|||
|
|
@ -2,11 +2,11 @@ import json
|
|||
import re
|
||||
from collections.abc import Collection, Mapping
|
||||
from types import MappingProxyType, UnionType
|
||||
from typing import Annotated, Any, Final, Union, get_args, get_origin
|
||||
from typing import Annotated, Any, Final, Literal, Union, get_args, get_origin
|
||||
|
||||
import orjson
|
||||
from fastapi import Request, UploadFile, status
|
||||
from typing_extensions import NotRequired, ReadOnly, Required
|
||||
from typing_extensions import NotRequired, ReadOnly, Required, assert_never
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import (
|
||||
|
|
@ -25,6 +25,40 @@ _FORM_CONTENT_TYPES: Final[frozenset[str]] = frozenset({"application/x-www-form-
|
|||
_ANNOTATION_QUALIFIERS: Final[frozenset[object]] = frozenset({Annotated, NotRequired, ReadOnly, Required})
|
||||
|
||||
|
||||
def resolve_inference_model(
|
||||
body_model: object,
|
||||
settings: Mapping[str, object],
|
||||
cli_model: str | None,
|
||||
endpoint_model: object = None,
|
||||
*,
|
||||
kind: Literal[
|
||||
"completion", "image_generation", "image_edit", "moderation", "speech", "body", "path"
|
||||
] = "completion",
|
||||
) -> object:
|
||||
match kind:
|
||||
case "image_generation":
|
||||
return cli_model or endpoint_model or settings.get("image_generation_model") or body_model
|
||||
case "image_edit":
|
||||
return (
|
||||
settings.get("completion_model")
|
||||
or cli_model
|
||||
or endpoint_model
|
||||
or settings.get("image_generation_model")
|
||||
or body_model
|
||||
)
|
||||
case "moderation":
|
||||
return cli_model or settings.get("moderation_model") or body_model
|
||||
case "speech":
|
||||
return cli_model or body_model
|
||||
case "body":
|
||||
return body_model
|
||||
case "path":
|
||||
return endpoint_model
|
||||
case "completion":
|
||||
return settings.get("completion_model") or cli_model or endpoint_model or body_model
|
||||
return assert_never(kind)
|
||||
|
||||
|
||||
def _normalize_media_type(content_type: str) -> str:
|
||||
"""Return the bare media type per RFC 7231: strip params, trim, lowercase."""
|
||||
if not content_type:
|
||||
|
|
|
|||
|
|
@ -174,7 +174,10 @@ async def _resync_agents(agent_id_or_name: str) -> bool:
|
|||
table: Final = agents_table(prisma_client)
|
||||
id_filter: Final[LiteLLM_AgentsTableWhereUniqueInput] = {"agent_id": agent_id_or_name}
|
||||
name_filter: Final[LiteLLM_AgentsTableWhereUniqueInput] = {"agent_name": agent_id_or_name}
|
||||
include_permission: Final[LiteLLM_AgentsTableInclude] = {"object_permission": True}
|
||||
include_permission: Final[LiteLLM_AgentsTableInclude] = {
|
||||
"object_permission": True,
|
||||
"identity": True,
|
||||
}
|
||||
async with AGENT_RECONCILE_LOCK:
|
||||
if _agent_from_registry(agent_id_or_name) is not None:
|
||||
return True
|
||||
|
|
|
|||
|
|
@ -3065,6 +3065,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
self,
|
||||
agent_id: str,
|
||||
data: dict,
|
||||
policy: "AgentResponse | None" = None,
|
||||
) -> list[RateLimitDescriptor]:
|
||||
"""
|
||||
Create rate limit descriptors for agent-level and session-level limits.
|
||||
|
|
@ -3074,7 +3075,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
"""
|
||||
descriptors: Final[list[RateLimitDescriptor]] = []
|
||||
|
||||
agent: Final = self._get_agent_from_registry(agent_id)
|
||||
agent: Final = policy if policy is not None else self._get_agent_from_registry(agent_id)
|
||||
if agent is None:
|
||||
return descriptors
|
||||
|
||||
|
|
@ -3269,14 +3270,19 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
descriptors=descriptors,
|
||||
)
|
||||
|
||||
# Agent-level and session-level rate limits
|
||||
resolved_agent_id: Final = self._get_resolved_agent_id(user_api_key_dict, data)
|
||||
|
||||
if resolved_agent_id:
|
||||
for agent_id in dict.fromkeys((resolved_agent_id, user_api_key_dict.invoked_agent_id)):
|
||||
if agent_id is None:
|
||||
continue
|
||||
descriptors.extend(
|
||||
self._create_agent_rate_limit_descriptors(
|
||||
agent_id=resolved_agent_id,
|
||||
agent_id=agent_id,
|
||||
data=data,
|
||||
policy=(
|
||||
user_api_key_dict.managed_agent_policy
|
||||
if agent_id == user_api_key_dict.agent_id
|
||||
else user_api_key_dict.invoked_agent_policy
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -4965,6 +4971,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
tpm_limited_tags=stash.tpm_limited_tags if stash is not None else frozenset(),
|
||||
model_group=reconcile_model.group if reconcile_model is not None else None,
|
||||
)
|
||||
targets.extend(
|
||||
scope
|
||||
for scope in sorted(reserved_scopes)
|
||||
if scope[0] in ("agent", "agent_session") and scope not in targets
|
||||
)
|
||||
charged_targets: Final = (
|
||||
[target for target in targets if target[0] != "model_per_team"]
|
||||
if self._key_owns_model_tpm_limit_from_request_metadata(request_metadata, reconcile_model)
|
||||
|
|
|
|||
|
|
@ -360,6 +360,7 @@ class _ProxyDBLogger(CustomLogger):
|
|||
team_id=team_id,
|
||||
end_user_id=end_user_id,
|
||||
call_type=call_type,
|
||||
agent_id=metadata.get("billing_agent_id") or metadata.get("agent_id"),
|
||||
):
|
||||
## UPDATE DATABASE
|
||||
charged: Final = await _update_database_and_spend_counters(
|
||||
|
|
@ -621,6 +622,7 @@ def _should_track_cost_callback(
|
|||
team_id: str | None,
|
||||
end_user_id: str | None,
|
||||
call_type: str | None = None,
|
||||
agent_id: str | None = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Determine if the cost callback should be tracked based on the kwargs
|
||||
|
|
@ -637,7 +639,13 @@ def _should_track_cost_callback(
|
|||
if ProxyUpdateSpend.disable_spend_updates() is True:
|
||||
return False
|
||||
|
||||
if user_api_key is not None or user_id is not None or team_id is not None or end_user_id is not None:
|
||||
if (
|
||||
agent_id is not None
|
||||
or user_api_key is not None
|
||||
or user_id is not None
|
||||
or team_id is not None
|
||||
or end_user_id is not None
|
||||
):
|
||||
return True
|
||||
return call_type in _UNATTRIBUTED_TRACKABLE_CALL_TYPES
|
||||
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ from litellm.proxy.common_request_processing import (
|
|||
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||
coerce_numeric_form_fields,
|
||||
numeric_form_fields,
|
||||
resolve_inference_model,
|
||||
)
|
||||
from litellm.proxy.common_utils.openai_error_payload import (
|
||||
error_status_code,
|
||||
|
|
@ -118,14 +119,9 @@ async def image_generation(
|
|||
if isinstance(model, str):
|
||||
reject_url_valued_destination("model", model)
|
||||
|
||||
data["model"] = (
|
||||
model
|
||||
or general_settings.get("image_generation_model", None) # server default
|
||||
or user_model # model name passed via cli args
|
||||
or data.get("model", None) # default passed in http request
|
||||
data["model"] = resolve_inference_model(
|
||||
data.get("model"), general_settings, user_model, model, kind="image_generation"
|
||||
)
|
||||
if user_model:
|
||||
data["model"] = user_model
|
||||
|
||||
### MODEL ALIAS MAPPING ###
|
||||
# check if model name in model alias map
|
||||
|
|
@ -324,12 +320,6 @@ async def image_edit_api(
|
|||
if "prompt" not in data:
|
||||
data["prompt"] = None
|
||||
|
||||
data["model"] = (
|
||||
model
|
||||
or general_settings.get("image_generation_model", None) # server default
|
||||
or user_model # model name passed via cli args
|
||||
or data.get("model", None) # default passed in http request
|
||||
)
|
||||
#########################################################
|
||||
# Process request
|
||||
#########################################################
|
||||
|
|
@ -346,7 +336,7 @@ async def image_edit_api(
|
|||
general_settings=general_settings,
|
||||
proxy_config=proxy_config,
|
||||
select_data_generator=select_data_generator,
|
||||
model=None,
|
||||
model=model,
|
||||
user_model=user_model,
|
||||
user_temperature=user_temperature,
|
||||
user_request_timeout=user_request_timeout,
|
||||
|
|
|
|||
|
|
@ -1664,7 +1664,19 @@ class LiteLLMProxyRequestSetup:
|
|||
_key_agent_id: Final = getattr(user_api_key_dict, "agent_id", None)
|
||||
_existing_agent_id: Final = data[_metadata_variable_name].get("agent_id")
|
||||
_resolved_agent_id: Final = _key_agent_id or _existing_agent_id
|
||||
data[_metadata_variable_name]["agent_id"] = _resolved_agent_id
|
||||
data[_metadata_variable_name]["agent_id"] = user_api_key_dict.invoked_agent_id or _resolved_agent_id
|
||||
managed_context: Final = user_api_key_dict.managed_agent_context
|
||||
data[_metadata_variable_name].update(
|
||||
MappingProxyType(
|
||||
{
|
||||
"actor_agent_id": user_api_key_dict.agent_id,
|
||||
"target_agent_id": user_api_key_dict.invoked_agent_id,
|
||||
"billing_agent_id": user_api_key_dict.agent_id or user_api_key_dict.invoked_agent_id,
|
||||
"agent_execution_mode": managed_context.mode if managed_context else None,
|
||||
"verified_human_user_id": managed_context.user_id if managed_context else None,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
data[_metadata_variable_name]["user_api_end_user_max_budget"] = getattr(
|
||||
user_api_key_dict, "end_user_max_budget", None
|
||||
|
|
|
|||
|
|
@ -0,0 +1,49 @@
|
|||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
from uuid import UUID
|
||||
|
||||
from litellm.types.proxy.agent_identity import MicrosoftInteractiveSubject
|
||||
|
||||
|
||||
def microsoft_interactive_subject(
|
||||
tenant: str | None,
|
||||
response: Mapping[str, object],
|
||||
endpoints: Mapping[str, str | None],
|
||||
) -> MicrosoftInteractiveSubject | None:
|
||||
if tenant is None:
|
||||
return None
|
||||
try:
|
||||
tenant_id: Final = str(UUID(tenant))
|
||||
object_id: Final = response.get("id")
|
||||
if not isinstance(object_id, str):
|
||||
return None
|
||||
oid: Final = str(UUID(object_id))
|
||||
except ValueError:
|
||||
return None
|
||||
expected: Final = MappingProxyType(
|
||||
{
|
||||
"MICROSOFT_AUTHORIZATION_ENDPOINT": f"https://login.microsoftonline.com/{tenant_id}/oauth2/v2.0/authorize",
|
||||
"MICROSOFT_TOKEN_ENDPOINT": f"https://login.microsoftonline.com/{tenant_id}/oauth2/v2.0/token",
|
||||
"MICROSOFT_USERINFO_ENDPOINT": "https://graph.microsoft.com/v1.0/me",
|
||||
}
|
||||
)
|
||||
if any(value and value != expected.get(name) for name, value in endpoints.items()):
|
||||
return None
|
||||
return MicrosoftInteractiveSubject(
|
||||
issuer=f"https://login.microsoftonline.com/{tenant_id}/v2.0",
|
||||
tenant_id=tenant_id,
|
||||
oid=oid,
|
||||
)
|
||||
|
||||
|
||||
async def enroll_microsoft_subject(subject: object, user_id: object, client: object) -> None:
|
||||
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
|
||||
from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityFailure
|
||||
|
||||
if not isinstance(subject, MicrosoftInteractiveSubject) or not isinstance(user_id, str) or not user_id:
|
||||
return
|
||||
result: Final = await AgentIdentityStore.from_client(client).enroll_interactive_human(subject, user_id)
|
||||
if isinstance(result, AgentIdentityFailure):
|
||||
raise_identity_failure(result)
|
||||
|
|
@ -2313,6 +2313,11 @@ async def _complete_cli_sso_callback_session(
|
|||
status_code=500,
|
||||
detail="Could not resolve team model grants for this login. Please try again",
|
||||
)
|
||||
from litellm.proxy.management_endpoints.sso.agent_subject_enrollment import enroll_microsoft_subject
|
||||
|
||||
await enroll_microsoft_subject(
|
||||
request.scope.get("litellm_microsoft_interactive_subject"), user_info.user_id, prisma_client
|
||||
)
|
||||
resolved_teams: Final = _cli_sso_session_teams(team_details)
|
||||
attribution_metadata: Final = build_cli_sso_attribution_metadata(result=result)
|
||||
if attribution_metadata:
|
||||
|
|
@ -3631,6 +3636,12 @@ class SSOAuthenticationHandler:
|
|||
},
|
||||
)
|
||||
|
||||
from litellm.proxy.management_endpoints.sso.agent_subject_enrollment import enroll_microsoft_subject
|
||||
|
||||
await enroll_microsoft_subject(
|
||||
request.scope.get("litellm_microsoft_interactive_subject"), user_id, prisma_client
|
||||
)
|
||||
|
||||
if isinstance(user_id, str) and user_id:
|
||||
await retain_sso_identity_assertion_for_ema(user_id=user_id, assertion=sso_assertion)
|
||||
await warn_if_id_jag_assertion_uncaptured(sso_assertion)
|
||||
|
|
@ -4300,6 +4311,22 @@ class MicrosoftSSOHandler:
|
|||
original_msft_result["app_roles"] = app_roles
|
||||
return original_msft_result or {}
|
||||
|
||||
from litellm.proxy.management_endpoints.sso.agent_subject_enrollment import microsoft_interactive_subject
|
||||
|
||||
request.scope["litellm_microsoft_interactive_subject"] = microsoft_interactive_subject(
|
||||
microsoft_tenant,
|
||||
original_msft_result,
|
||||
MappingProxyType(
|
||||
{
|
||||
name: os.getenv(name)
|
||||
for name in (
|
||||
"MICROSOFT_AUTHORIZATION_ENDPOINT",
|
||||
"MICROSOFT_TOKEN_ENDPOINT",
|
||||
"MICROSOFT_USERINFO_ENDPOINT",
|
||||
)
|
||||
}
|
||||
),
|
||||
)
|
||||
result: Final = MicrosoftSSOHandler.openid_from_response(
|
||||
response=original_msft_result,
|
||||
team_ids=user_team_ids,
|
||||
|
|
|
|||
|
|
@ -427,6 +427,7 @@ from litellm.proxy.common_utils.http_parsing_utils import (
|
|||
_safe_get_request_headers,
|
||||
check_file_size_under_limit,
|
||||
get_form_data,
|
||||
resolve_inference_model,
|
||||
)
|
||||
from litellm.proxy.common_utils.load_config_utils import get_config_from_bucket
|
||||
from litellm.proxy.common_utils.model_deprecation import collect_model_deprecations
|
||||
|
|
@ -12351,13 +12352,7 @@ async def moderations(
|
|||
proxy_config=proxy_config,
|
||||
)
|
||||
|
||||
data["model"] = (
|
||||
general_settings.get("moderation_model", None) # server default
|
||||
or user_model # model name passed via cli args
|
||||
or data.get("model") # default passed in http request
|
||||
)
|
||||
if user_model:
|
||||
data["model"] = user_model
|
||||
data["model"] = resolve_inference_model(data.get("model"), general_settings, user_model, kind="moderation")
|
||||
|
||||
### CALL HOOKS ### - modify incoming data / reject request before calling the model
|
||||
data = await proxy_logging_obj.pre_call_hook(
|
||||
|
|
@ -12611,13 +12606,7 @@ async def audio_transcriptions(
|
|||
if data.get("user", None) is None and user_api_key_dict.user_id is not None:
|
||||
data["user"] = user_api_key_dict.user_id
|
||||
|
||||
data["model"] = (
|
||||
general_settings.get("moderation_model", None) # server default
|
||||
or user_model # model name passed via cli args
|
||||
or data.get("model", None) # default passed in http request
|
||||
)
|
||||
if user_model:
|
||||
data["model"] = user_model
|
||||
data["model"] = resolve_inference_model(data.get("model"), general_settings, user_model, kind="moderation")
|
||||
|
||||
router_model_names: Final = llm_router.model_names if llm_router is not None else []
|
||||
|
||||
|
|
|
|||
|
|
@ -77,6 +77,11 @@ _SESSION_KEY_EXPR: Final = "COALESCE(NULLIF(session_id, ''), request_id)"
|
|||
_SESSION_GROUP_KEY_SQL: Final = f"{_SESSION_KEY_EXPR}, api_key"
|
||||
_MCP_CALL_TYPES_SQL: Final = "('call_mcp_tool', 'list_mcp_tools')"
|
||||
_AGENT_CALL_TYPE_SQL: Final = "'asend_message'"
|
||||
_SESSION_REPRESENTATIVE_ORDER_SQL: Final = (
|
||||
f"(call_type = {_AGENT_CALL_TYPE_SQL}) DESC, "
|
||||
f'CASE WHEN call_type = {_AGENT_CALL_TYPE_SQL} THEN "endTime" END DESC NULLS LAST, '
|
||||
f'call_type IN {_MCP_CALL_TYPES_SQL}, "startTime" DESC, request_id'
|
||||
)
|
||||
_BATCH_CALL_TYPES_SQL: Final = "('acreate_batch', 'create_batch', 'aretrieve_batch', 'retrieve_batch')"
|
||||
_SPAN_TYPE_SQL_CONDITIONS: Final[Mapping[str, str]] = MappingProxyType(
|
||||
{
|
||||
|
|
@ -2879,7 +2884,7 @@ async def ui_view_spend_logs(
|
|||
p += 1
|
||||
|
||||
# Status filter
|
||||
if status_filter is not None:
|
||||
if status_filter is not None and not (group_by_session is True and not is_search_lookup):
|
||||
if status_filter == "success":
|
||||
sql_conditions.append("(status = 'success' OR status IS NULL)")
|
||||
else:
|
||||
|
|
@ -2925,6 +2930,23 @@ async def ui_view_spend_logs(
|
|||
sql_params.append(f"%{error_message}%")
|
||||
p += 1
|
||||
|
||||
if status_filter is not None and group_by_session is True and not is_search_lookup:
|
||||
session_filter_conditions: Final = " AND ".join(sql_conditions) or "TRUE"
|
||||
sql_conditions.append(
|
||||
f"""({_SESSION_GROUP_KEY_SQL}) IN (
|
||||
SELECT session_key, api_key FROM (
|
||||
SELECT DISTINCT ON ({_SESSION_GROUP_KEY_SQL})
|
||||
{_SESSION_KEY_EXPR} AS session_key, api_key, status
|
||||
FROM "LiteLLM_SpendLogs"
|
||||
WHERE {session_filter_conditions}
|
||||
ORDER BY {_SESSION_GROUP_KEY_SQL}, {_SESSION_REPRESENTATIVE_ORDER_SQL}
|
||||
) AS session_outcomes
|
||||
WHERE COALESCE(status, 'success') = ${p}
|
||||
)"""
|
||||
)
|
||||
sql_params.append(status_filter)
|
||||
p += 1
|
||||
|
||||
if (
|
||||
group_by_session is True
|
||||
and not is_v2
|
||||
|
|
@ -2991,7 +3013,7 @@ async def ui_view_spend_logs(
|
|||
{_SPEND_LOG_LIST_COLUMNS}
|
||||
FROM "LiteLLM_SpendLogs"
|
||||
WHERE {joined_conditions}
|
||||
ORDER BY {_SESSION_GROUP_KEY_SQL}, call_type IN {_MCP_CALL_TYPES_SQL}, "startTime" DESC
|
||||
ORDER BY {_SESSION_GROUP_KEY_SQL}, {_SESSION_REPRESENTATIVE_ORDER_SQL}
|
||||
) AS session_representatives
|
||||
ORDER BY {exact_request_id_first}{_order_expr} {_sql_dir}{_nulls_clause}, request_id
|
||||
LIMIT ${p} OFFSET ${p + 1}
|
||||
|
|
@ -3063,7 +3085,7 @@ async def _fetch_session_representatives(
|
|||
next_param_index: int,
|
||||
session_keys: Sequence[tuple[str, str]],
|
||||
) -> list[dict[str, object]]: # mutable-ok: _build_ui_spend_logs_response writes session counts onto each row
|
||||
"""Fetch the newest non-MCP row of each ``(session_key, api_key)`` session, in ``session_keys`` order."""
|
||||
"""Fetch the final agent outcome, or newest non-MCP row, of each ``(session_key, api_key)`` session, in ``session_keys`` order."""
|
||||
rep_query: Final = f"""
|
||||
SELECT * FROM (
|
||||
SELECT DISTINCT ON ({_SESSION_GROUP_KEY_SQL})
|
||||
|
|
@ -3073,7 +3095,7 @@ async def _fetch_session_representatives(
|
|||
AND ({_SESSION_GROUP_KEY_SQL}) IN (
|
||||
SELECT * FROM unnest(${next_param_index}::text[], ${next_param_index + 1}::text[])
|
||||
)
|
||||
ORDER BY {_SESSION_GROUP_KEY_SQL}, call_type IN {_MCP_CALL_TYPES_SQL}, "startTime" DESC
|
||||
ORDER BY {_SESSION_GROUP_KEY_SQL}, {_SESSION_REPRESENTATIVE_ORDER_SQL}
|
||||
) AS session_representatives
|
||||
"""
|
||||
rep_rows: Final[Sequence[dict[str, object]]] = await _query_raw( # mutable-ok: rows are enriched in place
|
||||
|
|
@ -3140,7 +3162,7 @@ async def _ui_session_grouped_spend_logs(
|
|||
page_size``, trimmed to the end of the ``SPEND_LOGS_PAGINATION_COUNT_CAP``
|
||||
window the capped ``total`` promises, so a page never runs past that total
|
||||
and one starting at or past it returns no rows without a query. Each session is represented
|
||||
by its newest non-MCP row, enriched by ``_build_ui_spend_logs_response``
|
||||
by its final agent outcome (or newest non-MCP row), enriched by ``_build_ui_spend_logs_response``
|
||||
exactly like the flat listing, and the response carries
|
||||
``next_session_cursor`` / ``has_more`` while ``total`` counts sessions
|
||||
(capped like the flat total). A page that runs out of sessions while still
|
||||
|
|
|
|||
|
|
@ -796,6 +796,7 @@ def get_logging_payload(
|
|||
model_id=_model_id,
|
||||
mcp_namespaced_tool_name=mcp_namespaced_tool_name,
|
||||
agent_id=agent_id,
|
||||
billing_agent_id=clean_metadata.get("billing_agent_id"),
|
||||
requester_ip_address=clean_metadata.get("requester_ip_address", None),
|
||||
custom_llm_provider=custom_llm_provider or "",
|
||||
messages=_get_messages_for_spend_logs_payload(
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import logging
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -235,3 +236,13 @@ async def test_a_database_fault_retrying_cannot_clear_is_not_reported_as_a_trans
|
|||
assert refusal == SubjectTokenRefusal(error="temporarily_unavailable", description=SUBJECT_TOKEN_CHECK_FAULTED)
|
||||
assert "retrying will not help" in refusal.description
|
||||
assert "faulted: " in caplog.text and "query engine binary not found" in caplog.text
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("user_id", [None, "delegating-user"])
|
||||
async def test_agent_token_cannot_be_exchanged_for_a_user_identity(user_id: str | None) -> None:
|
||||
authorizer: Final = _Authorizer({**_authorized(user_id=user_id), "agent_id": "managed-agent"})
|
||||
result: Final = await _identity(authorizer)
|
||||
assert isinstance(result, SubjectTokenRefusal)
|
||||
assert result.error == "invalid_request"
|
||||
assert "direct JWT authentication" in result.description
|
||||
|
|
|
|||
|
|
@ -2244,7 +2244,7 @@ async def test_mcp_routing_chunked_initialize_to_stateful():
|
|||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(MagicMock(), None, ["progress_test"], None, None, None),
|
||||
return_value=(UserAPIKeyAuth(), None, ["progress_test"], None, None, None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.set_auth_context",
|
||||
|
|
@ -2356,7 +2356,7 @@ async def test_mcp_routing_caps_body_peek_for_oversized_chunked_body():
|
|||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(MagicMock(), None, ["progress_test"], None, None, None),
|
||||
return_value=(UserAPIKeyAuth(), None, ["progress_test"], None, None, None),
|
||||
),
|
||||
patch("litellm.proxy._experimental.mcp_server.server.set_auth_context"),
|
||||
patch(
|
||||
|
|
@ -2567,7 +2567,7 @@ async def test_mcp_routing_initialize_rejected_when_owner_at_session_cap():
|
|||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(MagicMock(), None, ["progress_test"], None, None, None),
|
||||
return_value=(UserAPIKeyAuth(), None, ["progress_test"], None, None, None),
|
||||
),
|
||||
patch("litellm.proxy._experimental.mcp_server.server.set_auth_context"),
|
||||
patch(
|
||||
|
|
|
|||
|
|
@ -9,9 +9,12 @@ they may send a stale `mcp-session-id` header. This test verifies that:
|
|||
|
||||
import asyncio
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
from litellm.types.mcp import MCPAuth
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.mcp import MCPAuth
|
||||
|
||||
|
||||
class TestHandleStaleMcpSession:
|
||||
"""Unit tests for the _handle_stale_mcp_session helper."""
|
||||
|
|
@ -260,7 +263,7 @@ async def test_stale_mcp_session_id_is_stripped():
|
|||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(MagicMock(), None, None, None, None, None),
|
||||
return_value=(UserAPIKeyAuth(), None, None, None, None, None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.set_auth_context",
|
||||
|
|
@ -337,7 +340,7 @@ async def test_delete_stale_mcp_session_returns_success():
|
|||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(MagicMock(), None, None, None, None, None),
|
||||
return_value=(UserAPIKeyAuth(), None, None, None, None, None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.set_auth_context",
|
||||
|
|
@ -386,7 +389,7 @@ async def test_failed_delete_preserves_stateful_session_tracking():
|
|||
pytest.skip("MCP server not available")
|
||||
|
||||
session_id = "delete-failure-session"
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.api_key = "sk-test"
|
||||
user_auth.user_id = "test-user"
|
||||
auth_context = MagicMock()
|
||||
|
|
@ -491,7 +494,7 @@ async def test_valid_mcp_session_id_is_preserved():
|
|||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(MagicMock(), None, None, None, None, None),
|
||||
return_value=(UserAPIKeyAuth(), None, None, None, None, None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.set_auth_context",
|
||||
|
|
@ -554,7 +557,7 @@ async def test_no_mcp_session_id_header_works_normally():
|
|||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(MagicMock(), None, None, None, None, None),
|
||||
return_value=(UserAPIKeyAuth(), None, None, None, None, None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.set_auth_context",
|
||||
|
|
@ -613,7 +616,7 @@ async def test_per_user_oauth_missing_stored_token_returns_preemptive_401():
|
|||
}
|
||||
receive = AsyncMock()
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = "test-user-id"
|
||||
oauth_server = MagicMock()
|
||||
oauth_server.auth_type = MCPAuth.oauth2
|
||||
|
|
@ -700,7 +703,7 @@ async def test_admitted_subject_missing_stored_token_challenged_with_resource_me
|
|||
}
|
||||
receive = AsyncMock()
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = "sso-user-42"
|
||||
user_auth.mcp_admitted_user_subject = True
|
||||
oauth_server = MagicMock()
|
||||
|
|
@ -806,7 +809,7 @@ async def test_client_credentials_server_is_not_preemptively_challenged(m2m_fiel
|
|||
}
|
||||
)
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = "test-user-id"
|
||||
m2m_server = MCPServer(
|
||||
server_id="m2m-server-id",
|
||||
|
|
@ -892,7 +895,7 @@ async def test_handle_streamable_http_mcp_delegated_server_surfaces_upstream_cha
|
|||
}
|
||||
)
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = None
|
||||
delegated_server = MagicMock()
|
||||
delegated_server.auth_type = MCPAuth.oauth2
|
||||
|
|
@ -996,7 +999,7 @@ async def test_per_user_oauth_with_stored_token_skips_preemptive_401():
|
|||
}
|
||||
)
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = "test-user-id"
|
||||
oauth_server = MagicMock()
|
||||
oauth_server.auth_type = MCPAuth.oauth2
|
||||
|
|
@ -1092,7 +1095,7 @@ async def test_handle_streamable_http_mcp_delegated_server_without_token_returns
|
|||
}
|
||||
)
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = None
|
||||
delegated_server = MagicMock()
|
||||
delegated_server.auth_type = MCPAuth.oauth2
|
||||
|
|
@ -1192,7 +1195,7 @@ async def test_handle_streamable_http_mcp_token_exchange_without_subject_returns
|
|||
}
|
||||
)
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = None
|
||||
obo_server = MagicMock()
|
||||
obo_server.auth_type = MCPAuth.oauth2_token_exchange
|
||||
|
|
@ -1301,7 +1304,7 @@ async def test_handle_streamable_http_mcp_oauth_delegate_without_token_returns_g
|
|||
}
|
||||
)
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = "u1"
|
||||
od_server = _build_passthrough_mode_server("od_server", MCPAuth.oauth_delegate)
|
||||
|
||||
|
|
@ -1366,7 +1369,7 @@ async def test_handle_streamable_http_mcp_oauth_delegate_with_forwarded_token_sk
|
|||
}
|
||||
)
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = "u1"
|
||||
od_server = _build_passthrough_mode_server("od_server", MCPAuth.oauth_delegate)
|
||||
|
||||
|
|
@ -1431,7 +1434,7 @@ async def _run_passthrough_connect(
|
|||
}
|
||||
)
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = "u1"
|
||||
server = _build_passthrough_mode_server(server_names[0], auth_type)
|
||||
|
||||
|
|
@ -1554,7 +1557,7 @@ async def test_handle_streamable_http_mcp_true_passthrough_without_token_surface
|
|||
}
|
||||
)
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = None
|
||||
tp_server = _build_passthrough_mode_server("tp_server", MCPAuth.true_passthrough)
|
||||
|
||||
|
|
@ -1620,7 +1623,7 @@ async def test_handle_streamable_http_mcp_true_passthrough_dcr_bridge_challenges
|
|||
}
|
||||
)
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = None
|
||||
bridge_server = _build_passthrough_mode_server("tp_bridge_server", MCPAuth.true_passthrough).model_copy(
|
||||
update={"dcr_bridge": True}
|
||||
|
|
@ -1691,7 +1694,7 @@ async def test_handle_streamable_http_mcp_true_passthrough_with_token_skips_prob
|
|||
}
|
||||
)
|
||||
send = AsyncMock()
|
||||
user_auth = MagicMock()
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.user_id = None
|
||||
tp_server = _build_passthrough_mode_server("tp_server", MCPAuth.true_passthrough)
|
||||
|
||||
|
|
|
|||
|
|
@ -810,6 +810,12 @@ class TestTestConnection:
|
|||
from litellm.proxy._types import LitellmUserRoles
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
from litellm.proxy.management_endpoints import mcp_management_endpoints
|
||||
|
||||
manager = MCPServerManager()
|
||||
monkeypatch.setattr(rest_endpoints, "global_mcp_server_manager", manager)
|
||||
monkeypatch.setattr(mcp_management_endpoints, "global_mcp_server_manager", manager)
|
||||
captured = self._capture_execute(monkeypatch)
|
||||
saved = MCPServer(
|
||||
server_id="saved-server-id",
|
||||
|
|
@ -1313,6 +1319,7 @@ class TestListToolsRestAPI:
|
|||
|
||||
async def fake_reload(user_id, *, requires_fresh_policy=False):
|
||||
assert user_id == "grant-user"
|
||||
assert requires_fresh_policy is False
|
||||
return admitted_auth
|
||||
|
||||
monkeypatch.setattr(
|
||||
|
|
@ -1480,9 +1487,12 @@ class TestListToolsRestAPI:
|
|||
from mcp.types import Tool as MCPTool
|
||||
|
||||
import litellm.experimental_mcp_client.client as mcp_client_module
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
from litellm.proxy._experimental.mcp_server.server import MCPServer
|
||||
from litellm.types.mcp import MCPTransport
|
||||
|
||||
monkeypatch.setattr(rest_endpoints, "global_mcp_server_manager", MCPServerManager())
|
||||
|
||||
async def fake_contexts(user_api_key_auth):
|
||||
return [user_api_key_auth]
|
||||
|
||||
|
|
@ -2414,6 +2424,7 @@ class TestCallToolRestAPI:
|
|||
|
||||
mock_server = MagicMock()
|
||||
mock_server.server_id = "server-1"
|
||||
mock_server.name = "Example server"
|
||||
|
||||
def fake_get_mcp_server_by_id(server_id):
|
||||
return mock_server if server_id == "server-1" else None
|
||||
|
|
@ -2431,6 +2442,11 @@ class TestCallToolRestAPI:
|
|||
raising=False,
|
||||
)
|
||||
|
||||
failure_log = AsyncMock()
|
||||
execute_tool = AsyncMock()
|
||||
monkeypatch.setattr(rest_endpoints, "_safe_fire_mcp_tool_call_failure_logging", failure_log)
|
||||
monkeypatch.setattr(rest_endpoints, "execute_mcp_tool", execute_tool)
|
||||
|
||||
request_payload = {
|
||||
"server_id": "server-1",
|
||||
"name": "demo-tool",
|
||||
|
|
@ -2452,6 +2468,16 @@ class TestCallToolRestAPI:
|
|||
assert exc_info.value.detail["error"] == "access_denied"
|
||||
assert "server server-1" in exc_info.value.detail["message"]
|
||||
|
||||
execute_tool.assert_not_awaited()
|
||||
failure_log.assert_awaited_once()
|
||||
logged_data = failure_log.await_args.args[4]
|
||||
assert logged_data["model"] == "MCP: demo-tool"
|
||||
assert logged_data["metadata"]["model_group"] == "MCP: demo-tool"
|
||||
logging_obj = failure_log.await_args.args[0]
|
||||
assert logging_obj.model_call_details["mcp_tool_call_metadata"] == {
|
||||
"name": "demo-tool", "mcp_server_name": "Example server",
|
||||
}
|
||||
|
||||
async def test_executes_tool_when_allowed(self, monkeypatch):
|
||||
async def fake_contexts(user_api_key_auth):
|
||||
return [user_api_key_auth]
|
||||
|
|
|
|||
|
|
@ -1043,3 +1043,119 @@ async def test_managed_target_preserves_ordinary_actor_ceilings_after_key_reload
|
|||
assert await AgentRequestHandler.is_agent_allowed("target", auth) is (permitted and ceiling != "group-without-grant")
|
||||
database.get_data.assert_awaited_once()
|
||||
assert auth.agent_caller == (AgentCaller(team_id="caller-team") if ceiling == "caller-team" else None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"direct,teams,selected,explicit,expected",
|
||||
[
|
||||
(False, ("a",), "b", True, "denied"),
|
||||
(False, ("a",), "a", True, "a"),
|
||||
(False, ("a",), None, False, "a"),
|
||||
(False, ("a",), "default-team", False, "a"),
|
||||
(False, ("a", "b"), "b", True, "b"),
|
||||
(False, ("b", "a"), None, False, "a"),
|
||||
(False, ("b", "a"), "default-team", False, "a"),
|
||||
(False, (), None, False, "denied"),
|
||||
(True, (), None, False, None),
|
||||
(True, ("a",), "b", True, "b"),
|
||||
],
|
||||
)
|
||||
async def test_delegated_team_selection_preserves_the_grant_source(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
direct: bool,
|
||||
teams: tuple[str, ...],
|
||||
selected: str | None,
|
||||
explicit: bool,
|
||||
expected: str | None,
|
||||
) -> None:
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy.agent_endpoints.auth import agent_permission_handler as permissions
|
||||
|
||||
sources: Final = [
|
||||
(None, frozenset({"actor"}) if direct else frozenset()),
|
||||
*((team, frozenset({"actor"})) for team in teams),
|
||||
]
|
||||
monkeypatch.setattr(permissions, "_verified_human_agent_sources", AsyncMock(return_value=sources))
|
||||
if expected == "denied":
|
||||
with pytest.raises(HTTPException) as error:
|
||||
await permissions.resolve_delegated_agent_team("human", "actor", selected, explicit_team=explicit)
|
||||
assert error.value.status_code == 403
|
||||
else:
|
||||
assert (
|
||||
await permissions.resolve_delegated_agent_team("human", "actor", selected, explicit_team=explicit)
|
||||
== expected
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"team_id,expected", [(None, {"direct"}), ("a", {"direct", "a-only"}), ("b", {"direct", "b-only"})]
|
||||
)
|
||||
async def test_delegated_target_grants_do_not_borrow_another_teams_authority(
|
||||
monkeypatch: pytest.MonkeyPatch, team_id: str | None, expected: set[str]
|
||||
) -> None:
|
||||
from litellm.proxy.agent_endpoints.auth import agent_permission_handler as permissions
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.agent_identity import ManagedAgentContext
|
||||
|
||||
sources: Final = [(None, frozenset({"direct"})), ("a", frozenset({"a-only"})), ("b", frozenset({"b-only"}))]
|
||||
monkeypatch.setattr(permissions, "_verified_human_agent_sources", AsyncMock(return_value=sources))
|
||||
auth: Final = UserAPIKeyAuth(agent_id="actor", team_id=team_id)
|
||||
auth.managed_agent_context = ManagedAgentContext(agent_id="actor", mode="delegated", user_id="human")
|
||||
auth.managed_agent_policy = AgentResponse(
|
||||
agent_id="actor",
|
||||
agent_name="Actor",
|
||||
agent_card_params={},
|
||||
object_permission={"object_permission_id": "own", "agents": ["direct", "a-only", "b-only"]},
|
||||
)
|
||||
assert await AgentRequestHandler.resolve_agent_access(auth) == RestrictedAgentAccess(frozenset(expected))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"managed,enabled,grant,outage,allowed",
|
||||
[
|
||||
(True, True, False, False, False),
|
||||
(True, True, True, False, True),
|
||||
(True, False, True, False, False),
|
||||
(False, True, False, False, True),
|
||||
(True, True, False, True, False),
|
||||
],
|
||||
)
|
||||
async def test_target_authorization_uses_live_policy_despite_stale_unmanaged_registry(
|
||||
monkeypatch: pytest.MonkeyPatch, managed: bool, enabled: bool, grant: bool, outage: bool, allowed: bool
|
||||
) -> None:
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.agent_endpoints import agent_registry
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityBinding
|
||||
|
||||
stale: Final = AgentResponse(agent_id="target", agent_name="Target", agent_card_params={})
|
||||
registry: Final = AgentRegistry()
|
||||
registry.register_agent(stale)
|
||||
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
|
||||
binding: Final = AgentIdentityBinding(
|
||||
agent_id="target", provider="microsoft_entra", tenant_id="tenant", client_id="client",
|
||||
issuer="issuer", revision="current",
|
||||
)
|
||||
current: Final = stale.model_copy(update={
|
||||
"identity_managed": managed, "identity": binding if managed else None, "enabled": enabled,
|
||||
})
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(
|
||||
return_value=current, side_effect=ConnectionError("writer unavailable") if outage else None,
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", database)
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="grant", agents=["target"])
|
||||
auth: Final = UserAPIKeyAuth(object_permission=permission if grant else None)
|
||||
|
||||
if outage:
|
||||
with pytest.raises(HTTPException) as denied:
|
||||
await AgentRequestHandler.is_agent_allowed("target", auth)
|
||||
assert denied.value.status_code == 503
|
||||
return
|
||||
assert await AgentRequestHandler.is_agent_allowed("target", auth) is allowed
|
||||
|
|
|
|||
|
|
@ -1,13 +1,15 @@
|
|||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy._types import LiteLLMRoutes, UserAPIKeyAuth
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import (
|
||||
actor_admission_failure,
|
||||
admit_managed_actor,
|
||||
invocation_target,
|
||||
)
|
||||
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
|
@ -68,6 +70,72 @@ def test_stale_binding_and_unverified_delegation_cannot_pass_admission(context:
|
|||
assert isinstance(actor_admission_failure(agent(), context), AgentIdentityFailure)
|
||||
|
||||
|
||||
def test_caller_cannot_construct_trusted_subject_or_policy() -> None:
|
||||
context: Final = ManagedAgentContext(
|
||||
agent_id="agent", binding_revision="current", mode="delegated", user_id="human"
|
||||
)
|
||||
auth: Final = UserAPIKeyAuth.model_validate(
|
||||
{
|
||||
"managed_agent_context": context,
|
||||
"requires_fresh_policy": True,
|
||||
"authenticated_by_custom_auth": True,
|
||||
"mcp_explicit_grants_only": True,
|
||||
"managed_agent_policy": agent(),
|
||||
"billing_agent_policy": agent(),
|
||||
"invoked_agent_id": "forged-target",
|
||||
"agent_invocation_cost": 0.0,
|
||||
}
|
||||
)
|
||||
assert auth.requires_fresh_policy is False
|
||||
assert auth.authenticated_by_custom_auth is False
|
||||
assert "authenticated_by_custom_auth" not in auth.model_dump()
|
||||
assert auth.mcp_explicit_grants_only is False
|
||||
assert "mcp_explicit_grants_only" not in auth.model_dump()
|
||||
assert auth.managed_agent_context is None
|
||||
assert auth.managed_agent_policy is None
|
||||
assert auth.billing_agent_policy is None
|
||||
assert auth.invoked_agent_id is None
|
||||
assert auth.agent_invocation_cost is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("autonomous", (True, False))
|
||||
async def test_invocation_prepares_target_fee_for_the_correct_agent(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
autonomous: bool,
|
||||
) -> None:
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
|
||||
from litellm.proxy.agent_endpoints import agent_registry
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation
|
||||
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
|
||||
|
||||
target: Final = agent(litellm_params={"cost_per_query": 0.25})
|
||||
registry: Final = agent_registry.AgentRegistry()
|
||||
registry.register_agent(target)
|
||||
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", database)
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="invoke-grant", agents=["agent"])
|
||||
auth: Final = UserAPIKeyAuth(
|
||||
agent_id="caller" if autonomous else None,
|
||||
user_id=None if autonomous else "human",
|
||||
object_permission=permission,
|
||||
)
|
||||
if autonomous:
|
||||
caller: Final = agent(agent_id="caller", object_permission=permission.model_dump())
|
||||
auth.managed_agent_policy = caller
|
||||
auth.billing_agent_policy = caller
|
||||
await prepare_agent_invocation(auth, "agent", AgentIdentityStore.from_client(database))
|
||||
assert auth.agent_invocation_cost == pytest.approx(0.25)
|
||||
assert auth.invoked_agent_id == "agent"
|
||||
assert auth.billing_agent_policy is not None
|
||||
assert auth.billing_agent_policy.agent_id == ("caller" if autonomous else "agent")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deleted_agent_key_cannot_fall_back_to_unmanaged_authentication() -> None:
|
||||
database: Final = MagicMock()
|
||||
|
|
@ -92,6 +160,24 @@ async def test_agent_history_outage_does_not_permit_legacy_fallback() -> None:
|
|||
assert failure.value.status_code == 503
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"route,body,expected",
|
||||
[
|
||||
("/a2a/agent", {}, "agent"),
|
||||
("/a2a/expensive", {"model": "a2a/cheap"}, "expensive"),
|
||||
("/a2a/expensive/message/send", {"model": "a2a/cheap"}, "expensive"),
|
||||
("/v1/a2a/expensive/message/send", {"model": "a2a/cheap"}, "expensive"),
|
||||
("/v1/a2a/agent/", {}, "agent"),
|
||||
("/v1/chat/completions", {"model": "a2a/Readable name"}, "Readable name"),
|
||||
("/v1/chat/completions", {"model": "a2a/"}, None),
|
||||
("/v1/chat/completions", {"model": "ordinary-model"}, None),
|
||||
("/a2a", {}, None),
|
||||
],
|
||||
)
|
||||
def test_invocation_routes_resolve_the_same_target(route: str, body: dict[str, object], expected: str | None) -> None:
|
||||
assert invocation_target(route, body) == expected
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_admission_database_outage_fails_closed() -> None:
|
||||
database: Final = MagicMock()
|
||||
|
|
@ -162,6 +248,36 @@ def test_execution_mode_must_match_verified_token_mode() -> None:
|
|||
assert "execution mode" in failure.message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("state,status", [("missing", 403), ("outage", 503), ("denied", 403), ("invalid-fee", 503)])
|
||||
async def test_invocation_cannot_bypass_missing_policy_permission_or_invalid_price(
|
||||
monkeypatch: pytest.MonkeyPatch, state: str, status: int
|
||||
) -> None:
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
|
||||
from litellm.proxy.agent_endpoints import agent_registry
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation
|
||||
|
||||
registered: Final = agent(litellm_params={"cost_per_query": -1 if state == "invalid-fee" else 0.25})
|
||||
registry: Final = agent_registry.AgentRegistry()
|
||||
registry.register_agent(registered)
|
||||
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(
|
||||
return_value=None if state == "missing" else registered,
|
||||
side_effect=RuntimeError("unavailable") if state == "outage" else None,
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", database)
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="grant", agents=[] if state == "denied" else ["agent"]
|
||||
)
|
||||
auth: Final = UserAPIKeyAuth(user_id="human", object_permission=permission)
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await prepare_agent_invocation(auth, "agent", AgentIdentityStore.from_client(database))
|
||||
assert failure.value.status_code == status
|
||||
assert auth.agent_invocation_cost is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_legacy_jwt_cannot_adopt_an_agent_bound_on_another_worker() -> None:
|
||||
database: Final = MagicMock()
|
||||
|
|
@ -192,6 +308,21 @@ async def test_managed_context_or_binding_requires_database(monkeypatch: pytest.
|
|||
assert denied.value.status_code == 503
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("managed_flag", [False, True])
|
||||
async def test_managed_invocation_requires_database(monkeypatch: pytest.MonkeyPatch, managed_flag: bool) -> None:
|
||||
from litellm.proxy.agent_endpoints import agent_registry
|
||||
from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation
|
||||
|
||||
registry: Final = AgentRegistry()
|
||||
registry.register_agent(agent(identity_managed=managed_flag))
|
||||
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
|
||||
with pytest.raises(HTTPException) as denied:
|
||||
await prepare_agent_invocation(UserAPIKeyAuth(user_id="human"), "agent", None)
|
||||
assert denied.value.status_code == 503
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_autonomous_app_rejects_persisted_virtual_key_impersonation() -> None:
|
||||
policy: Final = agent(execution_mode="autonomous")
|
||||
|
|
@ -205,6 +336,121 @@ async def test_autonomous_app_rejects_persisted_virtual_key_impersonation() -> N
|
|||
assert auth.billing_agent_policy is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unknown_invocation_target_leaves_billing_unset(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.agent_endpoints import agent_registry
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation
|
||||
|
||||
monkeypatch.setattr(agent_registry, "global_agent_registry", agent_registry.AgentRegistry())
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
auth: Final = UserAPIKeyAuth(user_id="human")
|
||||
await prepare_agent_invocation(auth, "missing", None)
|
||||
assert auth.invoked_agent_id is None
|
||||
assert auth.billing_agent_policy is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"route,method,allowed",
|
||||
[
|
||||
("/v1/agents", "GET", True),
|
||||
("/v1/agents", "POST", False),
|
||||
("/v1/chat/completions", "POST", True),
|
||||
("/v1/chat/completions", "DELETE", False),
|
||||
("/openai/deployments/model/chat/completions", "POST", True),
|
||||
("/engines/openai/model/chat/completions", "POST", True),
|
||||
("/openai/deployments/openai/model/images/generations", "POST", True),
|
||||
("/openai/deployments/openai/model/images/edits", "POST", True),
|
||||
("/v1beta/models/gemini-model:generateContent", "POST", True),
|
||||
("/v1/realtime", "GET", True),
|
||||
("/v1/realtime", "POST", False),
|
||||
("/v1/realtime/client_secrets", "POST", False),
|
||||
("/mcp/tools/call", "POST", True),
|
||||
("/a2a/target/message/send", "POST", True),
|
||||
("/v1/a2a/target/message/send", "POST", True),
|
||||
("/v1/videos", "POST", False),
|
||||
("/v1/videos/other-video", "GET", False),
|
||||
("/v1/search", "POST", False),
|
||||
("/search", "POST", False),
|
||||
("/v1/agents/target", "PATCH", False),
|
||||
("/v1/responses/other-response", "GET", False),
|
||||
("/v1/files", "GET", False),
|
||||
("/v1/files", "POST", False),
|
||||
("/openai/v1/files", "GET", False),
|
||||
("/anthropic/v1/files", "GET", False),
|
||||
],
|
||||
)
|
||||
def test_managed_route_scope_excludes_provider_resources(route: str, method: str, allowed: bool) -> None:
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_route_allowed
|
||||
|
||||
assert managed_agent_route_allowed(route, method) is allowed
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"route,body,settings,cli_model,path_model,expected",
|
||||
[
|
||||
("/v1/chat/completions", {"model": "body"}, {"completion_model": "default"}, "cli", "path", "default"),
|
||||
("/v1/moderations", {"model": "body"}, {"moderation_model": "default"}, "cli", None, "cli"),
|
||||
("/v1/audio/speech", {"model": "body"}, {"completion_model": "ignored"}, None, None, "body"),
|
||||
("/openai/deployments/path/embeddings", {"model": "body"}, {}, None, "path", "path"),
|
||||
("/v1/messages/count_tokens", {"model": "body"}, {"completion_model": "ignored"}, "cli", None, "body"),
|
||||
("/mcp/tools/call", {}, {"completion_model": "ignored"}, "cli", None, None),
|
||||
("/v1/images/generations", {"model": "image"}, {"completion_model": "text"}, None, None, "image"),
|
||||
("/v1/images/generations", {}, {"image_generation_model": "image"}, None, None, "image"),
|
||||
("/v1/images/edits", {}, {"image_generation_model": "image"}, None, None, "image"),
|
||||
("/v1/rerank", {"model": "reranker"}, {"completion_model": "text"}, "cli", None, "reranker"),
|
||||
("/v1beta/models/path:countTokens", {"model": "body"}, {"completion_model": "text"}, "cli", "path", "path"),
|
||||
],
|
||||
)
|
||||
def test_managed_inference_resolves_dispatch_precedence(
|
||||
route: str,
|
||||
body: Mapping[str, object],
|
||||
settings: Mapping[str, object],
|
||||
cli_model: str | None,
|
||||
path_model: str | None,
|
||||
expected: str | None,
|
||||
) -> None:
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_inference_request
|
||||
|
||||
assert managed_inference_request(route, body, settings, cli_model, path_model).get("model") == expected
|
||||
|
||||
|
||||
def test_managed_inference_without_any_model_cannot_skip_model_grants():
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_inference_request
|
||||
|
||||
with pytest.raises(HTTPException, match="explicit or configured model"):
|
||||
managed_inference_request("/v1/moderations", {}, {}, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("route", ["/v1/chat/completions", "/v1/images/generations", "/v1/images/edits"])
|
||||
def test_managed_inference_query_model_takes_precedence_over_body(route: str):
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_inference_request
|
||||
|
||||
assert managed_inference_request(route, {"model": "body"}, {}, None, query_model="query")["model"] == "query"
|
||||
|
||||
|
||||
def test_managed_inference_ignores_unsupported_query_model():
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_inference_request
|
||||
|
||||
assert (
|
||||
managed_inference_request("/v1/messages", {"model": "body"}, {}, None, query_model="query")["model"] == "body"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("route", ["/realtime", "/v1/realtime", "/openai/v1/realtime"])
|
||||
def test_managed_realtime_requires_a_model_and_ignores_completion_defaults(route: str) -> None:
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_inference_request
|
||||
|
||||
with pytest.raises(HTTPException, match="explicit or configured model"):
|
||||
managed_inference_request(route, {}, {"completion_model": "allowed-default"}, "cli")
|
||||
assert (
|
||||
managed_inference_request(route, {"model": "requested"}, {"completion_model": "allowed-default"}, "cli")[
|
||||
"model"
|
||||
]
|
||||
== "requested"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("mode,user", [("autonomous", None), ("delegated", "verified-human")])
|
||||
def test_matching_identity_revision_and_execution_mode_pass_admission(mode: str, user: str | None) -> None:
|
||||
context: Final = ManagedAgentContext.model_validate(
|
||||
|
|
@ -213,6 +459,28 @@ def test_matching_identity_revision_and_execution_mode_pass_admission(mode: str,
|
|||
assert actor_admission_failure(agent(), context) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unmanaged_agent_invocation_retains_legacy_behavior(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.agent_endpoints import agent_registry
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import prepare_agent_invocation
|
||||
|
||||
legacy: Final = agent(identity=None, identity_managed=False)
|
||||
registry: Final = agent_registry.AgentRegistry()
|
||||
registry.register_agent(legacy)
|
||||
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
auth: Final = UserAPIKeyAuth(agent_id="agent")
|
||||
await admit_managed_actor(auth, None)
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=legacy)
|
||||
await admit_managed_actor(auth, AgentIdentityStore.from_client(database))
|
||||
await prepare_agent_invocation(auth, "agent", AgentIdentityStore.from_client(database))
|
||||
assert auth.managed_agent_policy is None
|
||||
assert auth.billing_agent_policy is None
|
||||
assert auth.invoked_agent_id is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bound_autonomous_actor_is_admitted_without_a_human() -> None:
|
||||
database: Final = MagicMock()
|
||||
|
|
@ -234,6 +502,8 @@ async def test_admitted_managed_actor_requires_fresh_policy_so_revocations_bind_
|
|||
auth: Final = UserAPIKeyAuth(agent_id="agent")
|
||||
auth.managed_agent_context = ManagedAgentContext(agent_id="agent", binding_revision="current", mode="autonomous")
|
||||
assert auth.requires_fresh_policy is False
|
||||
assert auth.authenticated_by_custom_auth is False
|
||||
assert "authenticated_by_custom_auth" not in auth.model_dump()
|
||||
await admit_managed_actor(auth, AgentIdentityStore.from_client(database))
|
||||
assert auth.requires_fresh_policy is True
|
||||
|
||||
|
|
@ -283,3 +553,27 @@ async def test_ordinary_agent_admission_preserves_legacy_authentication(
|
|||
assert auth.agent_id == "agent"
|
||||
assert auth.managed_agent_policy is None
|
||||
assert auth.requires_fresh_policy is False
|
||||
assert auth.authenticated_by_custom_auth is False
|
||||
assert "authenticated_by_custom_auth" not in auth.model_dump()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"route",
|
||||
tuple(dict.fromkeys(
|
||||
LiteLLMRoutes.openai_routes.value
|
||||
+ LiteLLMRoutes.anthropic_routes.value
|
||||
+ LiteLLMRoutes.google_routes.value
|
||||
)),
|
||||
)
|
||||
def test_registered_inference_routes_have_an_explicit_managed_access_decision(route: str) -> None:
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_route_allowed
|
||||
|
||||
normalized: Final = route.removeprefix("/openai").removeprefix("/v1beta").removeprefix("/v1")
|
||||
unsupported: Final = normalized.startswith((
|
||||
"/videos", "/batches", "/files", "/fine_tuning", "/assistants", "/threads", "/utils/",
|
||||
"/vector_stores", "/vector_store/", "/search", "/containers", "/skills", "/claude-code/",
|
||||
"/interactions", "/agents", "/responses/{", "/responses/input_tokens",
|
||||
"/realtime/client_secrets", "/realtime/calls", "/realtime/transcription_sessions",
|
||||
)) or normalized in ("/models", "/cursor/models", "/cursor/v1/models")
|
||||
concrete: Final = route.split("?")[0].replace("{model}", "model").replace("{model_name:path}", "model")
|
||||
assert managed_agent_route_allowed(concrete, None) is not unsupported, route
|
||||
|
|
|
|||
|
|
@ -59,6 +59,7 @@ async def test_invoke_agent_a2a_adds_litellm_data():
|
|||
|
||||
# Mock agent
|
||||
mock_agent = MagicMock()
|
||||
mock_agent.agent_id = "test-agent"
|
||||
mock_agent.agent_card_params = {
|
||||
"url": "http://backend-agent:10001",
|
||||
"name": "Test Agent",
|
||||
|
|
@ -72,6 +73,7 @@ async def test_invoke_agent_a2a_adds_litellm_data():
|
|||
"jsonrpc": "2.0",
|
||||
"id": "test-id",
|
||||
"method": "message/send",
|
||||
"metadata": {"model_info": {"id": "caller-supplied-id"}},
|
||||
"params": {
|
||||
"message": {
|
||||
"role": "user",
|
||||
|
|
@ -153,7 +155,7 @@ async def test_invoke_agent_a2a_adds_litellm_data():
|
|||
"litellm.a2a_protocol.asend_message",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_response,
|
||||
),
|
||||
) as mock_send_message,
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{},
|
||||
|
|
@ -190,6 +192,9 @@ async def test_invoke_agent_a2a_adds_litellm_data():
|
|||
mock_add_data.assert_called_once()
|
||||
|
||||
# Verify model and custom_llm_provider were set
|
||||
assert mock_send_message.await_args.kwargs["model"] == "a2a_agent/Test Agent"
|
||||
assert captured_data["metadata"]["model_group"] == "a2a_agent/Test Agent"
|
||||
assert captured_data["metadata"]["model_info"] == {"id": mock_agent.agent_id}
|
||||
assert captured_data.get("model") == "a2a_agent/Test Agent"
|
||||
assert captured_data.get("custom_llm_provider") == "a2a_agent"
|
||||
|
||||
|
|
|
|||
|
|
@ -350,6 +350,7 @@ class TestAgentByIdKeyRedaction:
|
|||
|
||||
test_client = _make_app_with_role(role)
|
||||
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
|
||||
mock_prisma.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=None)
|
||||
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
|
|
@ -412,6 +413,7 @@ class TestAgentRBACInternalUser:
|
|||
return_value=_sample_agent_response()
|
||||
)
|
||||
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
|
||||
mock_prisma.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=None)
|
||||
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
|
|
@ -1342,6 +1344,7 @@ def test_get_agent_redacts_kill_switch_secret_for_admins_and_hides_it_from_other
|
|||
|
||||
def _get_as(role: LitellmUserRoles):
|
||||
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
|
||||
mock_prisma.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=None)
|
||||
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(return_value=None)
|
||||
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
|
||||
return _make_app_with_role(role).get("/v1/agents/agent-123", headers={"Authorization": "Bearer k"})
|
||||
|
|
|
|||
|
|
@ -2,15 +2,14 @@ import asyncio
|
|||
import re
|
||||
import time
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Final, Optional
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from fastapi import HTTPException
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
import litellm
|
||||
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy._types import (
|
||||
DEFAULT_JWKS_STALE_TTL,
|
||||
JWTLiteLLMRoleMap,
|
||||
|
|
@ -26,7 +25,6 @@ from litellm.proxy._types import (
|
|||
RoleBasedPermissions,
|
||||
ScopeMapping,
|
||||
)
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry
|
||||
from litellm.proxy.auth.auth_checks import TeamNotFoundError
|
||||
from litellm.proxy.auth.handle_jwt import (
|
||||
|
|
@ -1637,7 +1635,6 @@ async def test_auth_builder_returns_team_membership_object():
|
|||
@pytest.mark.asyncio
|
||||
async def test_auth_builder_with_oidc_userinfo_enabled():
|
||||
"""Test that auth_builder uses OIDC UserInfo endpoint when enabled"""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.caching import DualCache
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
|
@ -1648,9 +1645,7 @@ async def test_auth_builder_with_oidc_userinfo_enabled():
|
|||
general_settings = {"enforce_rbac": False}
|
||||
route = "/chat/completions"
|
||||
|
||||
user_object = LiteLLM_UserTable(
|
||||
user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER)
|
||||
|
||||
# Create JWT handler with OIDC UserInfo enabled
|
||||
jwt_handler = JWTHandler()
|
||||
|
|
@ -1677,18 +1672,12 @@ async def test_auth_builder_with_oidc_userinfo_enabled():
|
|||
|
||||
# Mock all the dependencies
|
||||
with (
|
||||
patch.object(
|
||||
jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock
|
||||
) as mock_get_userinfo,
|
||||
patch.object(jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock) as mock_get_userinfo,
|
||||
patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt,
|
||||
patch.object(
|
||||
JWTAuthManager, "check_rbac_role", new_callable=AsyncMock
|
||||
) as mock_check_rbac,
|
||||
patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock) as mock_check_rbac,
|
||||
patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac,
|
||||
patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes,
|
||||
patch.object(
|
||||
jwt_handler, "get_object_id", return_value=None
|
||||
) as mock_get_object_id,
|
||||
patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id,
|
||||
patch.object(
|
||||
JWTAuthManager,
|
||||
"get_user_info",
|
||||
|
|
@ -1696,9 +1685,7 @@ async def test_auth_builder_with_oidc_userinfo_enabled():
|
|||
return_value=("test_user_1", "test@example.com", True),
|
||||
) as mock_get_user_info,
|
||||
patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id,
|
||||
patch.object(
|
||||
jwt_handler, "get_end_user_id", return_value=None
|
||||
) as mock_get_end_user_id,
|
||||
patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id,
|
||||
patch.object(
|
||||
JWTAuthManager,
|
||||
"check_admin_access",
|
||||
|
|
@ -1711,9 +1698,7 @@ async def test_auth_builder_with_oidc_userinfo_enabled():
|
|||
new_callable=AsyncMock,
|
||||
return_value=(None, None),
|
||||
) as mock_find_team,
|
||||
patch.object(
|
||||
JWTAuthManager, "get_all_team_ids", return_value=set()
|
||||
) as mock_get_all_team_ids,
|
||||
patch.object(JWTAuthManager, "get_all_team_ids", return_value=set()) as mock_get_all_team_ids,
|
||||
patch.object(
|
||||
JWTAuthManager,
|
||||
"find_team_with_model_access",
|
||||
|
|
@ -1726,15 +1711,9 @@ async def test_auth_builder_with_oidc_userinfo_enabled():
|
|||
new_callable=AsyncMock,
|
||||
return_value=(user_object, None, None, None, user_object.user_id),
|
||||
) as mock_get_objects,
|
||||
patch.object(
|
||||
JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock
|
||||
) as mock_map_user,
|
||||
patch.object(
|
||||
JWTAuthManager, "validate_object_id", return_value=True
|
||||
) as mock_validate_object,
|
||||
patch.object(
|
||||
JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock
|
||||
) as mock_sync_user,
|
||||
patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock) as mock_map_user,
|
||||
patch.object(JWTAuthManager, "validate_object_id", return_value=True) as mock_validate_object,
|
||||
patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock) as mock_sync_user,
|
||||
):
|
||||
# Set up mock return values
|
||||
mock_get_userinfo.return_value = userinfo_response
|
||||
|
|
@ -1764,7 +1743,6 @@ async def test_auth_builder_with_oidc_userinfo_enabled():
|
|||
@pytest.mark.asyncio
|
||||
async def test_auth_builder_with_oidc_userinfo_disabled():
|
||||
"""Test that auth_builder uses JWT validation when OIDC UserInfo is disabled"""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.caching import DualCache
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
|
@ -1775,9 +1753,7 @@ async def test_auth_builder_with_oidc_userinfo_disabled():
|
|||
general_settings = {"enforce_rbac": False}
|
||||
route = "/chat/completions"
|
||||
|
||||
user_object = LiteLLM_UserTable(
|
||||
user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER)
|
||||
|
||||
# Create JWT handler with OIDC UserInfo disabled
|
||||
jwt_handler = JWTHandler()
|
||||
|
|
@ -1801,18 +1777,12 @@ async def test_auth_builder_with_oidc_userinfo_disabled():
|
|||
|
||||
# Mock all the dependencies
|
||||
with (
|
||||
patch.object(
|
||||
jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock
|
||||
) as mock_get_userinfo,
|
||||
patch.object(jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock) as mock_get_userinfo,
|
||||
patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt,
|
||||
patch.object(
|
||||
JWTAuthManager, "check_rbac_role", new_callable=AsyncMock
|
||||
) as mock_check_rbac,
|
||||
patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock) as mock_check_rbac,
|
||||
patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac,
|
||||
patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes,
|
||||
patch.object(
|
||||
jwt_handler, "get_object_id", return_value=None
|
||||
) as mock_get_object_id,
|
||||
patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id,
|
||||
patch.object(
|
||||
JWTAuthManager,
|
||||
"get_user_info",
|
||||
|
|
@ -1820,9 +1790,7 @@ async def test_auth_builder_with_oidc_userinfo_disabled():
|
|||
return_value=("test_user_1", None, None),
|
||||
) as mock_get_user_info,
|
||||
patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id,
|
||||
patch.object(
|
||||
jwt_handler, "get_end_user_id", return_value=None
|
||||
) as mock_get_end_user_id,
|
||||
patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id,
|
||||
patch.object(
|
||||
JWTAuthManager,
|
||||
"check_admin_access",
|
||||
|
|
@ -1835,9 +1803,7 @@ async def test_auth_builder_with_oidc_userinfo_disabled():
|
|||
new_callable=AsyncMock,
|
||||
return_value=(None, None),
|
||||
) as mock_find_team,
|
||||
patch.object(
|
||||
JWTAuthManager, "get_all_team_ids", return_value=set()
|
||||
) as mock_get_all_team_ids,
|
||||
patch.object(JWTAuthManager, "get_all_team_ids", return_value=set()) as mock_get_all_team_ids,
|
||||
patch.object(
|
||||
JWTAuthManager,
|
||||
"find_team_with_model_access",
|
||||
|
|
@ -1850,15 +1816,9 @@ async def test_auth_builder_with_oidc_userinfo_disabled():
|
|||
new_callable=AsyncMock,
|
||||
return_value=(user_object, None, None, None, user_object.user_id),
|
||||
) as mock_get_objects,
|
||||
patch.object(
|
||||
JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock
|
||||
) as mock_map_user,
|
||||
patch.object(
|
||||
JWTAuthManager, "validate_object_id", return_value=True
|
||||
) as mock_validate_object,
|
||||
patch.object(
|
||||
JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock
|
||||
) as mock_sync_user,
|
||||
patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock) as mock_map_user,
|
||||
patch.object(JWTAuthManager, "validate_object_id", return_value=True) as mock_validate_object,
|
||||
patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock) as mock_sync_user,
|
||||
):
|
||||
# Set up mock return values
|
||||
mock_auth_jwt.return_value = jwt_response
|
||||
|
|
@ -2631,7 +2591,6 @@ async def test_find_and_validate_specific_team_id_with_team_alias():
|
|||
"""
|
||||
Test that find_and_validate_specific_team_id resolves team by name when team_id is not found
|
||||
"""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.caching import DualCache
|
||||
from litellm.proxy._types import LiteLLM_JWTAuth, LiteLLM_TeamTable
|
||||
|
|
@ -2654,9 +2613,7 @@ async def test_find_and_validate_specific_team_id_with_team_alias():
|
|||
# Mock team object returned by get_team_object_by_alias
|
||||
team_object = LiteLLM_TeamTable(team_id="resolved-team-id", team_alias="my-team")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock
|
||||
) as mock_get_by_alias:
|
||||
with patch("litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock) as mock_get_by_alias:
|
||||
mock_get_by_alias.return_value = team_object
|
||||
|
||||
team_id, result_team = await JWTAuthManager.find_and_validate_specific_team_id(
|
||||
|
|
@ -2685,7 +2642,6 @@ async def test_find_and_validate_team_id_takes_precedence_over_name():
|
|||
"""
|
||||
Test that team_id_jwt_field takes precedence over team_alias_jwt_field
|
||||
"""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.caching import DualCache
|
||||
from litellm.proxy._types import LiteLLM_JWTAuth, LiteLLM_TeamTable
|
||||
|
|
@ -2699,9 +2655,7 @@ async def test_find_and_validate_team_id_takes_precedence_over_name():
|
|||
jwt_handler.update_environment(
|
||||
prisma_client=None,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
litellm_jwtauth=LiteLLM_JWTAuth(
|
||||
team_id_jwt_field="team_id", team_alias_jwt_field="team_alias"
|
||||
),
|
||||
litellm_jwtauth=LiteLLM_JWTAuth(team_id_jwt_field="team_id", team_alias_jwt_field="team_alias"),
|
||||
)
|
||||
|
||||
# Token with both team_id and team name
|
||||
|
|
@ -2711,9 +2665,7 @@ async def test_find_and_validate_team_id_takes_precedence_over_name():
|
|||
team_object = LiteLLM_TeamTable(team_id="direct-team-id")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock
|
||||
) as mock_get_by_id,
|
||||
patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_by_id,
|
||||
patch(
|
||||
"litellm.proxy.auth.handle_jwt.get_team_object_by_alias",
|
||||
new_callable=AsyncMock,
|
||||
|
|
@ -2890,7 +2842,6 @@ async def test_get_objects_resolves_org_by_name():
|
|||
@pytest.mark.asyncio
|
||||
async def test_resolve_jwks_url_passthrough_for_direct_jwks_url():
|
||||
"""Non-discovery URLs are returned unchanged."""
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
|
||||
|
|
@ -3143,7 +3094,7 @@ async def test_find_and_validate_specific_team_id_no_hint_for_valid_field():
|
|||
When team_id_jwt_field is a normal field name (no dot-notation) the
|
||||
error message should not contain a spurious bracket-notation hint.
|
||||
"""
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
|
||||
|
|
@ -3230,8 +3181,8 @@ async def test_find_and_validate_specific_team_id_no_hint_for_valid_field():
|
|||
async def test_auth_builder_single_team_db_fallback_when_jwt_has_no_team(
|
||||
user_id: str,
|
||||
user_teams: list,
|
||||
get_team_object_return: Optional[str],
|
||||
expected_team_id: Optional[str],
|
||||
get_team_object_return: str | None,
|
||||
expected_team_id: str | None,
|
||||
expect_get_team_called: bool,
|
||||
expect_get_membership_called: bool,
|
||||
) -> None:
|
||||
|
|
@ -3244,9 +3195,7 @@ async def test_auth_builder_single_team_db_fallback_when_jwt_has_no_team(
|
|||
if len(user_teams) == 1 and get_team_object_return == "resolved_row":
|
||||
only = user_teams[0]
|
||||
team_table = LiteLLM_TeamTable(team_id=only)
|
||||
membership = LiteLLM_TeamMembership(
|
||||
user_id=user_id, team_id=only, litellm_budget_table=None
|
||||
)
|
||||
membership = LiteLLM_TeamMembership(user_id=user_id, team_id=only, litellm_budget_table=None)
|
||||
get_team_return_value = team_table
|
||||
membership_return_value = membership
|
||||
else:
|
||||
|
|
@ -3305,9 +3254,7 @@ async def test_auth_builder_single_team_db_fallback_when_jwt_has_no_team(
|
|||
),
|
||||
patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock),
|
||||
patch.object(JWTAuthManager, "validate_object_id", return_value=True),
|
||||
patch.object(
|
||||
JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock
|
||||
),
|
||||
patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock),
|
||||
patch(
|
||||
"litellm.proxy.auth.handle_jwt.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
|
|
@ -3324,9 +3271,7 @@ async def test_auth_builder_single_team_db_fallback_when_jwt_has_no_team(
|
|||
code = 404 if get_team_object_return == "http_404" else 500
|
||||
mock_get_team.side_effect = HTTPException(
|
||||
status_code=code,
|
||||
detail={
|
||||
"error": f"Team doesn't exist in db. Team={user_teams[0]}. Create team via `/team/new` call."
|
||||
},
|
||||
detail={"error": f"Team doesn't exist in db. Team={user_teams[0]}. Create team via `/team/new` call."},
|
||||
)
|
||||
else:
|
||||
mock_get_team.return_value = get_team_return_value
|
||||
|
|
@ -4047,7 +3992,7 @@ def _encode_rsa_jwt(
|
|||
issuer: str,
|
||||
audience: str,
|
||||
kid: str,
|
||||
extra_claims: Optional[dict] = None,
|
||||
extra_claims: dict | None = None,
|
||||
) -> str:
|
||||
import time
|
||||
|
||||
|
|
@ -4743,12 +4688,9 @@ async def test_get_objects_team_membership_uses_rebound_user_id():
|
|||
async def fake_get_team_membership(user_id, team_id, *args, **kwargs):
|
||||
captured["user_id"] = user_id
|
||||
captured["team_id"] = team_id
|
||||
return None
|
||||
|
||||
jwt_handler = JWTHandler()
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(
|
||||
user_id_jwt_field="email", user_id_upsert=True
|
||||
)
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_id_jwt_field="email", user_id_upsert=True)
|
||||
|
||||
with (
|
||||
patch(
|
||||
|
|
@ -5389,7 +5331,7 @@ async def test_find_team_with_model_access_defers_no_team_403_under_db_fallback(
|
|||
assert team_object is None
|
||||
|
||||
|
||||
def _db_fallback_handler(litellm_jwtauth: Optional[LiteLLM_JWTAuth] = None) -> JWTHandler:
|
||||
def _db_fallback_handler(litellm_jwtauth: LiteLLM_JWTAuth | None = None) -> JWTHandler:
|
||||
handler = JWTHandler()
|
||||
handler.litellm_jwtauth = litellm_jwtauth or LiteLLM_JWTAuth()
|
||||
return handler
|
||||
|
|
@ -5447,9 +5389,7 @@ async def test_resolve_db_team_fallback_skips_unresolvable_membership():
|
|||
"expect_403",
|
||||
),
|
||||
[
|
||||
pytest.param(
|
||||
True, ["team_solo"], None, "team_solo", False, id="flag_on_single_db_team"
|
||||
),
|
||||
pytest.param(True, ["team_solo"], None, "team_solo", False, id="flag_on_single_db_team"),
|
||||
pytest.param(
|
||||
True,
|
||||
["team_a", "team_b"],
|
||||
|
|
@ -5497,8 +5437,8 @@ async def test_resolve_db_team_fallback_skips_unresolvable_membership():
|
|||
async def test_auth_builder_db_team_fallback_when_jwt_has_no_team(
|
||||
fallback_to_db_teams: bool,
|
||||
user_teams: list,
|
||||
header_team_id: Optional[str],
|
||||
expected_team_id: Optional[str],
|
||||
header_team_id: str | None,
|
||||
expected_team_id: str | None,
|
||||
expect_403: bool,
|
||||
) -> None:
|
||||
"""End-to-end auth_builder behavior with no JWT team claims.
|
||||
|
|
@ -5527,9 +5467,7 @@ async def test_auth_builder_db_team_fallback_when_jwt_has_no_team(
|
|||
|
||||
async def call_auth_builder():
|
||||
with (
|
||||
patch.object(
|
||||
jwt_handler, "auth_jwt", new_callable=AsyncMock
|
||||
) as mock_auth_jwt,
|
||||
patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt,
|
||||
patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock),
|
||||
patch.object(jwt_handler, "get_rbac_role", return_value=None),
|
||||
patch.object(jwt_handler, "get_scopes", return_value=[]),
|
||||
|
|
@ -5569,9 +5507,7 @@ async def test_auth_builder_db_team_fallback_when_jwt_has_no_team(
|
|||
),
|
||||
patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock),
|
||||
patch.object(JWTAuthManager, "validate_object_id", return_value=True),
|
||||
patch.object(
|
||||
JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock
|
||||
),
|
||||
patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock),
|
||||
patch(
|
||||
"litellm.proxy.auth.handle_jwt.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
|
|
@ -6765,7 +6701,7 @@ async def test_auth_builder_provisional_header_team_is_not_upserted():
|
|||
team_id_upsert=True,
|
||||
)
|
||||
|
||||
upsert_by_team: dict[str, Optional[bool]] = {}
|
||||
upsert_by_team: dict[str, bool | None] = {}
|
||||
|
||||
async def spy_get_team(team_id, **kwargs):
|
||||
upsert_by_team[team_id] = kwargs.get("team_id_upsert")
|
||||
|
|
@ -6800,9 +6736,7 @@ async def test_auth_builder_provisional_header_team_is_not_upserted():
|
|||
),
|
||||
patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock),
|
||||
patch.object(JWTAuthManager, "validate_object_id", return_value=True),
|
||||
patch.object(
|
||||
JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock
|
||||
),
|
||||
patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock),
|
||||
patch(
|
||||
"litellm.proxy.auth.handle_jwt.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
|
|
@ -7806,6 +7740,58 @@ async def test_admin_jwt_team_header_only_provisions_during_admission(monkeypatc
|
|||
assert result["team_id"] is None
|
||||
|
||||
|
||||
def _explicit_identity_registry() -> AgentRegistry:
|
||||
registry: Final = AgentRegistry()
|
||||
registry.register_agent(AgentResponse(
|
||||
agent_id="explicit-agent-id",
|
||||
agent_name="Readable agent name",
|
||||
agent_card_params={},
|
||||
litellm_params={"identity": {
|
||||
"provider": "microsoft_entra",
|
||||
"tenant_id": "11111111-1111-4111-8111-111111111111",
|
||||
"client_id": "22222222-2222-4222-8222-222222222222",
|
||||
}},
|
||||
))
|
||||
return registry
|
||||
|
||||
|
||||
@pytest.mark.parametrize("claim_field", ["azp", None])
|
||||
def test_runtime_json_cannot_establish_a_managed_identity(claim_field: str | None) -> None:
|
||||
registry: Final = _explicit_identity_registry()
|
||||
handler: Final = _entra_agent_jwt_handler(claim_field)
|
||||
claims: Final = {
|
||||
"iss": "https://login.microsoftonline.com/11111111-1111-4111-8111-111111111111/v2.0",
|
||||
"tid": "11111111-1111-4111-8111-111111111111",
|
||||
"azp": "22222222-2222-4222-8222-222222222222",
|
||||
}
|
||||
if claim_field is None:
|
||||
assert JWTAuthManager.resolve_agent_id(handler, claims, registry) is None
|
||||
else:
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
JWTAuthManager.resolve_agent_id(handler, claims, registry)
|
||||
assert failure.value.status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.parametrize("override", [
|
||||
{"iss": "https://attacker.example"},
|
||||
{"tid": "33333333-3333-4333-8333-333333333333"},
|
||||
{"azp": "33333333-3333-4333-8333-333333333333"},
|
||||
{"azp": "explicit-agent-id"},
|
||||
{"azp": "Readable agent name"},
|
||||
])
|
||||
def test_explicit_entra_identity_cannot_be_claimed_via_legacy_lookup(override: Mapping[str, object]) -> None:
|
||||
registry: Final = _explicit_identity_registry()
|
||||
handler: Final = _entra_agent_jwt_handler("azp")
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
JWTAuthManager.resolve_agent_id(handler, {
|
||||
"iss": "https://login.microsoftonline.com/11111111-1111-4111-8111-111111111111/v2.0",
|
||||
"tid": "11111111-1111-4111-8111-111111111111",
|
||||
"azp": "22222222-2222-4222-8222-222222222222",
|
||||
**override,
|
||||
}, registry)
|
||||
assert failure.value.status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("existing_user", [False, True])
|
||||
@pytest.mark.parametrize("warm_cache", [False, True])
|
||||
|
|
@ -7853,3 +7839,389 @@ async def test_scope_admin_admission_resolves_existing_user_without_provisioning
|
|||
users.create.assert_not_awaited()
|
||||
if existing_user:
|
||||
assert users.find_unique.await_count == (0 if warm_cache else 1)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("mode", ["autonomous", "both", "delegated"])
|
||||
@pytest.mark.parametrize("audience_validation", (True, False))
|
||||
@pytest.mark.parametrize(
|
||||
"route,allowed",
|
||||
[
|
||||
("/chat/completions", True), ("/v1/messages", True), ("/v1/responses", True),
|
||||
("/mcp-rest/tools/call", True), ("/a2a/target", True),
|
||||
("/v1/files", False), ("/v1/batches", False), ("/v1/vector_stores", False),
|
||||
("/v1/containers", False), ("/openai/v1/files", False),
|
||||
("/v1/responses/other-response", False), ("/v1/realtime/client_secrets", False),
|
||||
],
|
||||
)
|
||||
async def test_managed_application_uses_persisted_identity_without_provisioning_human(
|
||||
monkeypatch: pytest.MonkeyPatch, mode: str, audience_validation: bool, route: str, allowed: bool
|
||||
) -> None:
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityBinding
|
||||
|
||||
tenant: Final = "11111111-1111-4111-8111-111111111111"
|
||||
client_id: Final = "22222222-2222-4222-8222-222222222222"
|
||||
principal: Final = "33333333-3333-4333-8333-333333333333"
|
||||
issuer: Final = f"https://login.microsoftonline.com/{tenant}/v2.0"
|
||||
jwks_url: Final = "https://login.microsoftonline.test/managed-keys"
|
||||
monkeypatch.setenv("JWT_PUBLIC_KEY_URL", jwks_url)
|
||||
monkeypatch.setenv("JWT_ISSUER", issuer)
|
||||
monkeypatch.setenv("JWT_AUDIENCE", "api://gateway")
|
||||
private_key, jwk = _get_rsa_key_and_jwk(kid="managed-key")
|
||||
cache: Final = DualCache()
|
||||
cache.set_cache(key=f"litellm_jwt_auth_keys_{jwks_url}", value=[jwk])
|
||||
handler: Final = JWTHandler()
|
||||
handler.update_environment(None, cache, LiteLLM_JWTAuth(user_id_upsert=True))
|
||||
binding: Final = AgentIdentityBinding(
|
||||
agent_id="stable-id",
|
||||
provider="microsoft_entra",
|
||||
issuer=issuer,
|
||||
tenant_id=tenant,
|
||||
client_id=client_id,
|
||||
service_principal_id=principal,
|
||||
revision="revision-one",
|
||||
required_roles=("Agent.Invoke",),
|
||||
)
|
||||
agent: Final = AgentResponse.model_validate(
|
||||
{
|
||||
"agent_id": "stable-id",
|
||||
"agent_name": "A readable name",
|
||||
"agent_card_params": {},
|
||||
"identity": binding,
|
||||
"identity_managed": True,
|
||||
"execution_mode": mode,
|
||||
}
|
||||
)
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_agentidentity.find_unique = AsyncMock(return_value=binding)
|
||||
database.writer_db.litellm_agentidentity.update_many = AsyncMock(return_value=1)
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent)
|
||||
database.writer_db.litellm_verifiedsubject.find_unique = AsyncMock(return_value=None)
|
||||
database.db.litellm_usertable.upsert = AsyncMock()
|
||||
token: Final = _encode_rsa_jwt(
|
||||
private_key,
|
||||
issuer=issuer,
|
||||
audience="api://gateway",
|
||||
kid="managed-key",
|
||||
extra_claims={
|
||||
"tid": tenant,
|
||||
"azp": client_id,
|
||||
"oid": principal,
|
||||
"roles": ["Agent.Invoke"],
|
||||
"idtyp": "app",
|
||||
},
|
||||
)
|
||||
arguments: Final = dict(
|
||||
api_key=token,
|
||||
jwt_handler=handler,
|
||||
request_data={},
|
||||
general_settings={},
|
||||
route=route,
|
||||
prisma_client=database,
|
||||
user_api_key_cache=cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
if not audience_validation:
|
||||
monkeypatch.delenv("JWT_AUDIENCE")
|
||||
if mode == "delegated" or not audience_validation or not allowed:
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await JWTAuthManager.auth_builder(**arguments)
|
||||
assert failure.value.status_code == 403
|
||||
else:
|
||||
result: Final = await JWTAuthManager.auth_builder(**arguments)
|
||||
auth: Final = JWTAuthManager.user_api_key_auth_from_result(result)
|
||||
assert auth.agent_id == "stable-id"
|
||||
assert auth.api_key is None
|
||||
assert auth.token is None
|
||||
assert auth.user_id is None
|
||||
assert auth.team_id is None
|
||||
assert auth.managed_agent_context is not None
|
||||
assert auth.managed_agent_context.mode == "autonomous"
|
||||
assert result["is_proxy_admin"] is False
|
||||
database.db.litellm_usertable.upsert.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("claim_value", ["managed", "Readable managed agent"])
|
||||
def test_legacy_claim_cannot_select_a_top_level_entra_binding(claim_value: str) -> None:
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityBinding
|
||||
|
||||
registry: Final = AgentRegistry()
|
||||
registry.register_agent(
|
||||
AgentResponse(
|
||||
agent_id="managed",
|
||||
agent_name="Readable managed agent",
|
||||
agent_card_params={},
|
||||
identity_managed=True,
|
||||
identity=AgentIdentityBinding(
|
||||
agent_id="managed",
|
||||
provider="microsoft_entra",
|
||||
tenant_id="tenant",
|
||||
client_id="client",
|
||||
service_principal_id="principal",
|
||||
issuer="issuer",
|
||||
revision="revision",
|
||||
),
|
||||
)
|
||||
)
|
||||
with pytest.raises(HTTPException) as denied:
|
||||
JWTAuthManager.resolve_agent_id(_entra_agent_jwt_handler("agent"), {"agent": claim_value}, registry)
|
||||
assert denied.value.status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("kind", ["human", "config-agent", "managed-agent"])
|
||||
async def test_database_free_jwt_admission_with_entra_shaped_claims(monkeypatch: pytest.MonkeyPatch, kind: str) -> None:
|
||||
issuer: Final = "https://login.microsoftonline.com/test-tenant/v2.0"
|
||||
jwks_url: Final = "https://login.microsoftonline.test/config-only-keys"
|
||||
monkeypatch.setenv("JWT_PUBLIC_KEY_URL", jwks_url)
|
||||
monkeypatch.setenv("JWT_ISSUER", issuer)
|
||||
monkeypatch.setenv("JWT_AUDIENCE", "api://gateway")
|
||||
private_key, jwk = _get_rsa_key_and_jwk(kid="config-key")
|
||||
cache: Final = DualCache()
|
||||
cache.set_cache(key=f"litellm_jwt_auth_keys_{jwks_url}", value=[jwk])
|
||||
registry: Final = AgentRegistry()
|
||||
registry.register_agent(
|
||||
AgentResponse(
|
||||
agent_id="configured",
|
||||
agent_name="Configured",
|
||||
agent_card_params={},
|
||||
identity_managed=kind == "managed-agent",
|
||||
)
|
||||
)
|
||||
handler: Final = JWTHandler()
|
||||
handler.update_environment(None, cache, LiteLLM_JWTAuth(agent_id_jwt_field="agent", admin_allowed_routes=["llm_api_routes"]))
|
||||
handler.bind_agent_lookup(registry)
|
||||
token: Final = _encode_rsa_jwt(
|
||||
private_key,
|
||||
issuer=issuer,
|
||||
audience="api://gateway",
|
||||
kid="config-key",
|
||||
extra_claims={
|
||||
"tid": "test-tenant",
|
||||
"azp": "application",
|
||||
"scope": "litellm_proxy_admin",
|
||||
**({"agent": "configured"} if kind != "human" else {}),
|
||||
},
|
||||
)
|
||||
arguments: Final = dict(
|
||||
api_key=token,
|
||||
jwt_handler=handler,
|
||||
request_data={},
|
||||
general_settings={},
|
||||
route="/chat/completions",
|
||||
prisma_client=None,
|
||||
user_api_key_cache=cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
if kind == "managed-agent":
|
||||
with pytest.raises(HTTPException) as denied:
|
||||
await JWTAuthManager.auth_builder(**arguments)
|
||||
assert denied.value.status_code == 403
|
||||
else:
|
||||
result: Final = await JWTAuthManager.auth_builder(**arguments)
|
||||
auth: Final = JWTAuthManager.user_api_key_auth_from_result(result)
|
||||
assert auth.agent_id == ("configured" if kind == "config-agent" else None)
|
||||
assert auth.managed_agent_context is None
|
||||
assert result["is_proxy_admin"] is True
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"issuer,audience,disabled,expected",
|
||||
[
|
||||
(None, "gateway", False, False),
|
||||
("trusted", "gateway", False, True),
|
||||
("trusted", None, True, False),
|
||||
("other", "gateway", False, False),
|
||||
],
|
||||
)
|
||||
def test_managed_issuer_requires_configured_audience_validation(
|
||||
monkeypatch: pytest.MonkeyPatch, issuer: str | None, audience: str | None, disabled: bool, expected: bool
|
||||
) -> None:
|
||||
from litellm.proxy._types import JWTIssuerConfig
|
||||
|
||||
monkeypatch.delenv("JWT_ISSUER", raising=False)
|
||||
monkeypatch.delenv("JWT_AUDIENCE", raising=False)
|
||||
handler: Final = JWTHandler()
|
||||
handler.update_environment(
|
||||
None,
|
||||
DualCache(),
|
||||
LiteLLM_JWTAuth(
|
||||
issuers=[
|
||||
JWTIssuerConfig(issuer="trusted", audience=audience, disable_audience_validation=disabled),
|
||||
]
|
||||
),
|
||||
)
|
||||
assert handler.managed_issuer_is_trusted(issuer) is expected
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("authentication_write", ["success", "revoked", "unavailable"])
|
||||
async def test_managed_jwt_reuses_binding_lookup_but_rechecks_disabled_policy(
|
||||
monkeypatch: pytest.MonkeyPatch, authentication_write: str
|
||||
) -> None:
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityBinding
|
||||
|
||||
issuer: Final = "https://login.microsoftonline.com/tenant/v2.0"
|
||||
jwks_url: Final = "https://identity.example/managed-jwks"
|
||||
private_key, jwk = _get_rsa_key_and_jwk("managed-cache")
|
||||
cache: Final = UserApiKeyCache()
|
||||
cache.set_cache(f"litellm_jwt_auth_keys_{jwks_url}", [jwk])
|
||||
monkeypatch.setenv("JWT_PUBLIC_KEY_URL", jwks_url)
|
||||
monkeypatch.setenv("JWT_ISSUER", issuer)
|
||||
monkeypatch.setenv("JWT_AUDIENCE", "gateway")
|
||||
binding: Final = AgentIdentityBinding(
|
||||
agent_id="managed", provider="microsoft_entra", issuer=issuer, tenant_id="tenant",
|
||||
client_id="client", service_principal_id="principal", revision="current",
|
||||
)
|
||||
agent: Final = AgentResponse(
|
||||
agent_id="managed", agent_name="Managed", agent_card_params={}, identity_managed=True, identity=binding,
|
||||
)
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_agentidentity.find_unique = AsyncMock(return_value=binding)
|
||||
database.writer_db.litellm_agentidentity.update_many = AsyncMock(return_value=1)
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent)
|
||||
handler: Final = JWTHandler()
|
||||
handler.update_environment(database, cache, LiteLLM_JWTAuth())
|
||||
token: Final = _encode_rsa_jwt(
|
||||
private_key, issuer, "gateway", "managed-cache", {"tid": "tenant", "azp": "client", "oid": "principal"}
|
||||
)
|
||||
arguments: Final = dict(
|
||||
api_key=token, jwt_handler=handler, request_data={}, general_settings={}, route="/chat/completions",
|
||||
prisma_client=database, user_api_key_cache=cache, parent_otel_span=None, proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
for _ in range(2):
|
||||
result: Final = await JWTAuthManager.authorize_jwt(**arguments)
|
||||
assert result["agent_id"] == "managed"
|
||||
database.writer_db.litellm_agentidentity.find_unique.assert_awaited_once()
|
||||
assert database.writer_db.litellm_agentstable.find_unique.await_count == 2
|
||||
assert database.writer_db.litellm_agentidentity.update_many.await_count == 2
|
||||
if authentication_write != "success":
|
||||
database.writer_db.litellm_agentidentity.update_many.return_value = 0
|
||||
database.writer_db.litellm_agentidentity.update_many.side_effect = (
|
||||
RuntimeError("storage unavailable") if authentication_write == "unavailable" else None
|
||||
)
|
||||
with pytest.raises(HTTPException) as failed_write:
|
||||
await JWTAuthManager.authorize_jwt(**arguments)
|
||||
assert failed_write.value.status_code == (503 if authentication_write == "unavailable" else 403)
|
||||
assert database.writer_db.litellm_agentidentity.update_many.await_count == 3
|
||||
return
|
||||
database.writer_db.litellm_agentstable.find_unique.return_value = agent.model_copy(update={"enabled": False})
|
||||
with pytest.raises(HTTPException) as denied:
|
||||
await JWTAuthManager.authorize_jwt(**arguments)
|
||||
assert denied.value.status_code == 403
|
||||
assert database.writer_db.litellm_agentidentity.update_many.await_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"team_route_allowed,team_claim,db_fallback",
|
||||
[
|
||||
(True, None, False),
|
||||
(False, None, False),
|
||||
(True, "other-team", False),
|
||||
(True, "granting-team", False),
|
||||
(True, "other-team", True),
|
||||
(True, "alias:other-team", False),
|
||||
(True, "alias:other-team", True),
|
||||
],
|
||||
)
|
||||
async def test_delegated_jwt_uses_granting_team_policy_before_route_authorization(
|
||||
monkeypatch: pytest.MonkeyPatch, team_route_allowed: bool, team_claim: str | None, db_fallback: bool
|
||||
) -> None:
|
||||
from litellm.proxy.agent_endpoints.auth import agent_permission_handler
|
||||
from litellm.proxy.auth import handle_jwt
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.types.proxy.agent_identity import ManagedAgentContext
|
||||
|
||||
issuer: Final = "https://login.microsoftonline.com/11111111-1111-4111-8111-111111111111/v2.0"
|
||||
jwks_url: Final = "https://identity.example/delegated-jwks"
|
||||
private_key, jwk = _get_rsa_key_and_jwk("delegated-team")
|
||||
cache: Final = UserApiKeyCache()
|
||||
cache.set_cache(f"litellm_jwt_auth_keys_{jwks_url}", [jwk])
|
||||
monkeypatch.setenv("JWT_PUBLIC_KEY_URL", jwks_url)
|
||||
monkeypatch.setenv("JWT_ISSUER", issuer)
|
||||
monkeypatch.setenv("JWT_AUDIENCE", "gateway")
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_agentidentity.update_many = AsyncMock(return_value=1)
|
||||
handler: Final = JWTHandler()
|
||||
handler.update_environment(
|
||||
database,
|
||||
cache,
|
||||
LiteLLM_JWTAuth(
|
||||
team_allowed_routes=["/chat/completions" if team_route_allowed else "/embeddings"],
|
||||
team_id_jwt_field="team" if team_claim is not None else None,
|
||||
team_alias_jwt_field="team_alias" if team_claim is not None else None,
|
||||
fallback_to_db_teams=db_fallback,
|
||||
),
|
||||
)
|
||||
context: Final = ManagedAgentContext(
|
||||
agent_id="delegated-agent", binding_revision="revision", mode="delegated", user_id="human"
|
||||
)
|
||||
monkeypatch.setattr(handle_jwt, "resolve_managed_agent", AsyncMock(return_value=context))
|
||||
monkeypatch.setattr(
|
||||
agent_permission_handler,
|
||||
"_verified_human_agent_sources",
|
||||
AsyncMock(return_value=(("granting-team", frozenset(("delegated-agent",))),)),
|
||||
)
|
||||
team: Final = LiteLLM_TeamTable(team_id="granting-team", models=["allowed-model"], max_budget=5)
|
||||
|
||||
async def team_policy(team_id: str, **kwargs: object) -> LiteLLM_TeamTable:
|
||||
return team if team_id == team.team_id else LiteLLM_TeamTable(team_id=team_id)
|
||||
|
||||
load_team: Final = AsyncMock(side_effect=team_policy)
|
||||
monkeypatch.setattr(handle_jwt, "get_team_object", load_team)
|
||||
monkeypatch.setattr(
|
||||
handle_jwt, "get_team_object_by_alias", AsyncMock(return_value=LiteLLM_TeamTable(team_id="other-team"))
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
handle_jwt,
|
||||
"get_user_object",
|
||||
AsyncMock(return_value=LiteLLM_UserTable(user_id="human", teams=["granting-team", "other-team"])),
|
||||
)
|
||||
monkeypatch.setattr(handle_jwt, "get_team_membership", AsyncMock(return_value=None))
|
||||
token: Final = _encode_rsa_jwt(
|
||||
private_key,
|
||||
issuer,
|
||||
"gateway",
|
||||
"delegated-team",
|
||||
{
|
||||
"sub": "human",
|
||||
**(
|
||||
{"team_alias": "other-team"}
|
||||
if team_claim == "alias:other-team"
|
||||
else {"team": team_claim}
|
||||
if team_claim
|
||||
else {}
|
||||
),
|
||||
},
|
||||
)
|
||||
pending: Final = JWTAuthManager.authorize_jwt(
|
||||
api_key=token,
|
||||
jwt_handler=handler,
|
||||
request_data={"model": "allowed-model"},
|
||||
general_settings={},
|
||||
route="/chat/completions",
|
||||
request_method="POST",
|
||||
prisma_client=database,
|
||||
user_api_key_cache=cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
if not team_route_allowed or (team_claim in ("other-team", "alias:other-team") and not db_fallback):
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await pending
|
||||
assert failure.value.status_code == 403
|
||||
if team_claim is None:
|
||||
assert "granting team" in failure.value.detail
|
||||
load_team.assert_not_awaited()
|
||||
return
|
||||
result: Final = await pending
|
||||
assert result["team_id"] == "granting-team"
|
||||
assert result["team_object"] == team
|
||||
assert result["user_id"] == "human"
|
||||
assert result["managed_agent_context"] == context
|
||||
if team_claim != "granting-team":
|
||||
assert any(call.kwargs.get("check_db_only") is True for call in load_team.call_args_list)
|
||||
|
|
|
|||
|
|
@ -9278,6 +9278,8 @@ async def test_websocket_auth_hands_the_reservation_to_the_socket_state():
|
|||
)
|
||||
|
||||
async def auth_that_reserves(request, api_key):
|
||||
assert request.method == "GET"
|
||||
assert request.query_params.get("model") == "gpt-realtime"
|
||||
request.state.budget_reservation = reservation
|
||||
return UserAPIKeyAuth(token="hashed", budget_reservation=reservation)
|
||||
|
||||
|
|
@ -9383,6 +9385,209 @@ def test_identity_prefetch_keys_match_what_auth_reads_for_the_request():
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("invoke", [False, True])
|
||||
async def test_centralized_authorization_preserves_database_free_config_agents(monkeypatch, invoke: bool):
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry
|
||||
from litellm.proxy.agent_endpoints import agent_registry
|
||||
from litellm.proxy.auth.user_api_key_auth import _authorize_authenticated_request
|
||||
|
||||
for name, value in {
|
||||
**_proxy_attrs_for_centralized_checks(),
|
||||
"prisma_client": None,
|
||||
"proxy_logging_obj": MagicMock(post_call_failure_hook=AsyncMock(return_value=None)),
|
||||
}.items():
|
||||
monkeypatch.setattr(proxy_server, name, value)
|
||||
registry = AgentRegistry()
|
||||
registry.load_agents_from_config(
|
||||
[{"agent_name": "config-agent", "agent_card_params": {"name": "Config", "url": "http://localhost:9999"}}]
|
||||
)
|
||||
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
|
||||
registered = registry.get_agent_by_name("config-agent")
|
||||
model = "a2a/config-agent" if invoke else "test-model"
|
||||
auth = UserAPIKeyAuth(agent_id=registered.agent_id, jwt_claims={"agent": "config-agent"}, models=[model])
|
||||
data = {"model": model, "messages": [{"role": "user", "content": "hi"}]}
|
||||
assert (
|
||||
await _authorize_authenticated_request(
|
||||
auth, _alias_request("/v1/chat/completions", data), data, "/v1/chat/completions", "jwt-token"
|
||||
)
|
||||
is None
|
||||
)
|
||||
assert auth.managed_agent_policy is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("verified_identity", [False, True])
|
||||
async def test_managed_actor_cannot_access_provider_resource_routes(monkeypatch, verified_identity: bool):
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.auth.user_api_key_auth import _authorize_authenticated_request
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityBinding
|
||||
|
||||
policy = AgentResponse(
|
||||
agent_id="managed",
|
||||
agent_name="Managed",
|
||||
agent_card_params={},
|
||||
identity_managed=True,
|
||||
identity=AgentIdentityBinding(
|
||||
agent_id="managed",
|
||||
provider="microsoft_entra",
|
||||
tenant_id="tenant",
|
||||
client_id="application",
|
||||
service_principal_id="principal",
|
||||
issuer="issuer",
|
||||
revision="revision",
|
||||
),
|
||||
object_permission={"models": ["test-model"]},
|
||||
)
|
||||
database = MagicMock()
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy)
|
||||
for name, value in {
|
||||
**_proxy_attrs_for_centralized_checks(),
|
||||
"prisma_client": database,
|
||||
"proxy_logging_obj": MagicMock(post_call_failure_hook=AsyncMock(return_value=None)),
|
||||
}.items():
|
||||
monkeypatch.setattr(proxy_server, name, value)
|
||||
request = _alias_request("/v1/files", {})
|
||||
request.scope["method"] = "GET"
|
||||
from litellm.types.proxy.agent_identity import ManagedAgentContext
|
||||
|
||||
auth = UserAPIKeyAuth(agent_id="managed", api_key="persisted-key", models=["test-model"])
|
||||
if verified_identity:
|
||||
auth.managed_agent_context = ManagedAgentContext(
|
||||
agent_id="managed", binding_revision="revision", mode="autonomous"
|
||||
)
|
||||
with pytest.raises(ProxyException) as denied:
|
||||
await _authorize_authenticated_request(auth, request, {}, "/v1/files", "persisted-key")
|
||||
assert denied.value.code == "403"
|
||||
if verified_identity:
|
||||
assert denied.value.message == "Agent identities can only access inference and agent discovery routes"
|
||||
else:
|
||||
assert denied.value.message == "This agent requires its bound identity provider token"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("requested", [None, "test-model"])
|
||||
@pytest.mark.parametrize("grant_default", [False, True])
|
||||
@pytest.mark.parametrize(
|
||||
"route,settings,cli_model",
|
||||
[
|
||||
("/v1/chat/completions", {"completion_model": "forbidden-model"}, None),
|
||||
("/v1/responses", {"completion_model": "forbidden-model"}, None),
|
||||
("/v1/messages", {"completion_model": "forbidden-model"}, None),
|
||||
("/v1/moderations", {"moderation_model": "forbidden-model"}, None),
|
||||
("/v1/audio/transcriptions", {"moderation_model": "forbidden-model"}, None),
|
||||
("/v1/audio/speech", {}, "forbidden-model"),
|
||||
("/v1/chat/completions", {}, "forbidden-model"),
|
||||
("/v1/images/generations", {"image_generation_model": "forbidden-model"}, None),
|
||||
("/v1/images/edits", {"image_generation_model": "forbidden-model"}, None),
|
||||
],
|
||||
)
|
||||
async def test_managed_agent_cannot_bypass_grants_with_server_default(
|
||||
monkeypatch, requested, route, settings, cli_model, grant_default
|
||||
):
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.auth.user_api_key_auth import _authorize_authenticated_request
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityBinding, ManagedAgentContext
|
||||
|
||||
policy = AgentResponse(
|
||||
agent_id="managed",
|
||||
agent_name="Managed",
|
||||
agent_card_params={},
|
||||
identity_managed=True,
|
||||
identity=AgentIdentityBinding(
|
||||
agent_id="managed",
|
||||
provider="microsoft_entra",
|
||||
tenant_id="tenant",
|
||||
client_id="application",
|
||||
service_principal_id="principal",
|
||||
issuer="issuer",
|
||||
revision="revision",
|
||||
),
|
||||
object_permission={"models": ["test-model", "forbidden-model"] if grant_default else ["test-model"]},
|
||||
)
|
||||
database = MagicMock()
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy)
|
||||
for name, value in {
|
||||
**_proxy_attrs_for_centralized_checks(),
|
||||
"prisma_client": database,
|
||||
"general_settings": settings,
|
||||
"user_model": cli_model,
|
||||
"proxy_logging_obj": MagicMock(post_call_failure_hook=AsyncMock(return_value=None)),
|
||||
}.items():
|
||||
monkeypatch.setattr(proxy_server, name, value)
|
||||
data = {"messages": [{"role": "user", "content": "hi"}], **({"model": requested} if requested else {})}
|
||||
auth = UserAPIKeyAuth(agent_id="managed")
|
||||
auth.managed_agent_context = ManagedAgentContext(
|
||||
agent_id="managed", binding_revision="revision", mode="autonomous"
|
||||
)
|
||||
if not grant_default:
|
||||
with pytest.raises(ProxyException) as denied:
|
||||
await _authorize_authenticated_request(auth, _alias_request(route, data), data, route, "persisted-key")
|
||||
assert denied.value.code == "403"
|
||||
assert "forbidden-model" in denied.value.message
|
||||
return
|
||||
with patch(
|
||||
"litellm.proxy.spend_tracking.budget_reservation.reserve_budget_for_request",
|
||||
new_callable=AsyncMock,
|
||||
) as reserve:
|
||||
reserve.return_value = None
|
||||
assert (
|
||||
await _authorize_authenticated_request(auth, _alias_request(route, data), data, route, "persisted-key")
|
||||
is None
|
||||
)
|
||||
reserve.assert_awaited_once()
|
||||
assert reserve.call_args.kwargs["request_body"]["model"] == "forbidden-model"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_managed_jwt_cannot_be_downgraded_into_virtual_key_mapping(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
from typing import Final
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityBinding
|
||||
|
||||
binding: Final = AgentIdentityBinding(
|
||||
agent_id="managed", provider="microsoft_entra", issuer="issuer", tenant_id="tenant",
|
||||
client_id="client", service_principal_id="principal", revision="current",
|
||||
)
|
||||
agent: Final = AgentResponse(
|
||||
agent_id="managed", agent_name="Managed", agent_card_params={},
|
||||
identity_managed=True, identity=binding, execution_mode="autonomous",
|
||||
)
|
||||
client: Final = MagicMock()
|
||||
client.writer_db.litellm_agentidentity.find_unique = AsyncMock(return_value=binding)
|
||||
client.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent)
|
||||
handler: Final = MagicMock()
|
||||
handler.is_jwt.return_value = True
|
||||
handler.litellm_jwtauth = LiteLLM_JWTAuth(virtual_key_claim_field="sub")
|
||||
handler.auth_jwt = AsyncMock(return_value={
|
||||
"iss": "issuer", "tid": "tenant", "azp": "client", "oid": "principal", "sub": "mapped-key",
|
||||
})
|
||||
for name, value in {
|
||||
**_proxy_attrs_for_centralized_checks(),
|
||||
"general_settings": {"enable_jwt_auth": True}, "premium_user": True,
|
||||
"prisma_client": client, "jwt_handler": handler, "user_api_key_cache": UserApiKeyCache(),
|
||||
"proxy_logging_obj": MagicMock(post_call_failure_hook=AsyncMock(return_value=None)),
|
||||
}.items():
|
||||
monkeypatch.setattr(proxy_server, name, value)
|
||||
for _ in range(2):
|
||||
with pytest.raises(ProxyException) as failure:
|
||||
await _user_api_key_auth_builder(
|
||||
request=_alias_request("/v1/chat/completions", {}), api_key="Bearer verified.jwt.token",
|
||||
azure_api_key_header="", anthropic_api_key_header=None, google_ai_studio_api_key_header=None,
|
||||
azure_apim_header=None, request_data={},
|
||||
)
|
||||
assert failure.value.code == "403"
|
||||
assert "without virtual-key mapping" in failure.value.message
|
||||
client.writer_db.litellm_agentidentity.find_unique.assert_awaited_once()
|
||||
assert client.writer_db.litellm_agentstable.find_unique.await_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_virtual_key_cannot_enter_checks_as_an_identity_managed_actor(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
from typing import Final
|
||||
|
|
@ -9410,3 +9615,76 @@ async def test_virtual_key_cannot_enter_checks_as_an_identity_managed_actor(monk
|
|||
UserAPIKeyAuth(agent_id="bound"), request, data, "/v1/chat/completions", "sk-test"
|
||||
)
|
||||
checks.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("enterprise", [False, True])
|
||||
@pytest.mark.parametrize("credential", ["custom-credential", "sk-custom-credential"])
|
||||
@pytest.mark.parametrize("granted", [False, True])
|
||||
async def test_custom_auth_grants_reach_managed_targets_without_a_virtual_key_row(
|
||||
monkeypatch: pytest.MonkeyPatch, enterprise: bool, credential: str, granted: bool
|
||||
) -> None:
|
||||
import importlib
|
||||
from typing import Final
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.agent_endpoints import agent_registry
|
||||
from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry
|
||||
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import AgentRequestHandler
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityBinding
|
||||
|
||||
target: Final = AgentResponse(
|
||||
agent_id="target", agent_name="Target", agent_card_params={}, identity_managed=True,
|
||||
identity=AgentIdentityBinding(
|
||||
agent_id="target", provider="microsoft_entra", tenant_id="tenant", client_id="client",
|
||||
issuer="issuer", revision="current",
|
||||
),
|
||||
)
|
||||
registry: Final = AgentRegistry()
|
||||
registry.register_agent(target)
|
||||
trusted: Final = UserAPIKeyAuth(
|
||||
api_key=credential, object_permission={"object_permission_id": "custom", "agents": ["target"] if granted else ["other"]}
|
||||
)
|
||||
custom: Final = AsyncMock(return_value=trusted)
|
||||
database: Final = MagicMock()
|
||||
database.get_data = AsyncMock(return_value=None)
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target)
|
||||
for name, value in {
|
||||
**_proxy_server_attrs_for_custom_auth(user_custom_auth=None if enterprise else custom),
|
||||
"prisma_client": database,
|
||||
}.items():
|
||||
monkeypatch.setattr(proxy_server, name, value)
|
||||
module: Final = importlib.import_module("litellm.proxy.auth.user_api_key_auth")
|
||||
monkeypatch.setattr(module, "enterprise_custom_auth", custom if enterprise else None)
|
||||
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
|
||||
monkeypatch.setattr(litellm, "enable_post_custom_auth_checks", False, raising=False)
|
||||
admitted: Final = await _user_api_key_auth_builder(
|
||||
request=_alias_request("/a2a/target/message/send", {}), api_key=f"Bearer {credential}",
|
||||
azure_api_key_header="", anthropic_api_key_header=None, google_ai_studio_api_key_header=None,
|
||||
azure_apim_header=None, request_data={},
|
||||
)
|
||||
assert await AgentRequestHandler.is_agent_allowed("target", admitted) is granted
|
||||
custom.assert_awaited_once()
|
||||
database.get_data.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_enterprise_custom_auth_key_return_stays_a_proxy_validated_key(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
import importlib
|
||||
from typing import Final
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
custom: Final = AsyncMock(return_value="sk-master-key")
|
||||
for name, value in _proxy_server_attrs_for_custom_auth(user_custom_auth=custom).items():
|
||||
monkeypatch.setattr(proxy_server, name, value)
|
||||
module: Final = importlib.import_module("litellm.proxy.auth.user_api_key_auth")
|
||||
monkeypatch.setattr(module, "enterprise_custom_auth", custom)
|
||||
admitted: Final = await _user_api_key_auth_builder(
|
||||
request=_alias_request("/v1/chat/completions", {}), api_key="Bearer external-credential",
|
||||
azure_api_key_header="", anthropic_api_key_header=None, google_ai_studio_api_key_header=None,
|
||||
azure_apim_header=None, request_data={},
|
||||
)
|
||||
assert admitted.authenticated_by_custom_auth is False
|
||||
assert admitted.via_virtual_key is True
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import io
|
||||
import json
|
||||
from typing import get_type_hints
|
||||
from collections.abc import Mapping
|
||||
from typing import Literal, get_type_hints
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import orjson
|
||||
|
|
@ -1210,3 +1211,42 @@ class TestCoerceNumericFormFields:
|
|||
numeric_fields=self.numeric_fields,
|
||||
)
|
||||
assert result == {"n": 3, "temperature": None, "image": buffer}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"kind,settings,cli,path,body,expected",
|
||||
[
|
||||
("completion", {"completion_model": "default"}, "cli", "path", "body", "default"),
|
||||
("completion", {}, "cli", "path", "body", "cli"),
|
||||
("completion", {}, None, "path", "body", "path"),
|
||||
("completion", {}, None, None, "body", "body"),
|
||||
(
|
||||
"image_generation",
|
||||
{"completion_model": "text", "image_generation_model": "image"},
|
||||
None,
|
||||
None,
|
||||
"body",
|
||||
"image",
|
||||
),
|
||||
("image_generation", {"image_generation_model": "image"}, "cli", "path", "body", "cli"),
|
||||
("image_generation", {"image_generation_model": "image"}, None, "path", "body", "path"),
|
||||
("image_edit", {"completion_model": "text", "image_generation_model": "image"}, None, None, "body", "text"),
|
||||
("image_edit", {"image_generation_model": "image"}, None, "path", "body", "path"),
|
||||
("image_edit", {"image_generation_model": "image"}, None, None, "body", "image"),
|
||||
("moderation", {"moderation_model": "mod"}, "cli", None, "body", "cli"),
|
||||
("speech", {"completion_model": "text"}, None, None, "body", "body"),
|
||||
("body", {"completion_model": "text"}, "cli", None, "body", "body"),
|
||||
("path", {"completion_model": "text"}, "cli", "path", "body", "path"),
|
||||
],
|
||||
)
|
||||
def test_shared_inference_model_selection_preserves_handler_precedence(
|
||||
kind: Literal["completion", "image_generation", "image_edit", "moderation", "speech", "body", "path"],
|
||||
settings: Mapping[str, object],
|
||||
cli: str | None,
|
||||
path: str | None,
|
||||
body: str,
|
||||
expected: str,
|
||||
) -> None:
|
||||
from litellm.proxy.common_utils.http_parsing_utils import resolve_inference_model
|
||||
|
||||
assert resolve_inference_model(body, settings, cli, path, kind=kind) == expected
|
||||
|
|
|
|||
|
|
@ -177,7 +177,7 @@ async def test_get_agent_with_read_through_recovers_agent_created_on_sibling_rep
|
|||
assert agent.agent_id == agent_id
|
||||
prisma_client.db.litellm_agentstable.find_unique.assert_awaited_once_with(
|
||||
where={"agent_id": agent_id},
|
||||
include={"object_permission": True},
|
||||
include={"object_permission": True, "identity": True},
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -202,7 +202,7 @@ async def test_get_agent_with_read_through_recovers_agent_by_name(clean_agent_re
|
|||
assert agent.agent_name == agent_name
|
||||
prisma_client.db.litellm_agentstable.find_unique.assert_awaited_with(
|
||||
where={"agent_name": agent_name},
|
||||
include={"object_permission": True},
|
||||
include={"object_permission": True, "identity": True},
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -521,3 +521,33 @@ async def test_resync_agents_waits_for_agent_reload_and_skips_duplicate_registra
|
|||
|
||||
assert await resync_task is True
|
||||
assert len(clean_agent_registry.agent_list) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("lookup", ["agent-id", "Agent name"])
|
||||
async def test_agent_read_through_hydrates_identity_binding(lookup, clean_agent_registry, fresh_agent_read_through, monkeypatch):
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from litellm.proxy.common_utils.registry_read_through import get_agent_with_read_through
|
||||
|
||||
binding = {
|
||||
"agent_id": "agent-id", "provider": "microsoft_entra", "tenant_id": "tenant", "client_id": "client",
|
||||
"issuer": "https://login.microsoftonline.com/tenant/v2.0", "revision": "revision",
|
||||
}
|
||||
|
||||
async def load_row(*, where, include):
|
||||
if where == {"agent_id": "Agent name"}:
|
||||
return None
|
||||
row = FakeAgentRow("agent-id", "Agent name").model_dump()
|
||||
return SimpleNamespace(model_dump=lambda: {**row, "identity": binding if include.get("identity") else None})
|
||||
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_agentstable.find_unique = AsyncMock(side_effect=load_row)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
|
||||
agent = await get_agent_with_read_through(lookup)
|
||||
assert agent is not None
|
||||
assert agent.identity is not None
|
||||
assert agent.identity.model_dump(include=set(binding)) == binding
|
||||
assert clean_agent_registry.get_agent_by_id(agent_id="agent-id").identity == agent.identity
|
||||
|
|
|
|||
|
|
@ -7592,3 +7592,117 @@ async def test_success_tpm_accounting_keeps_the_admission_target_after_an_alias_
|
|||
charged: Final = {op["key"]: op["increment_value"] for op in ops}
|
||||
assert charged[admission_bucket] == 150 - stash.reserved_tokens
|
||||
assert not any(":target-b" in key for key in charged)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("self_call", [False, True])
|
||||
async def test_managed_invocations_enforce_actor_and_target_rate_policies(
|
||||
monkeypatch: pytest.MonkeyPatch, self_call: bool
|
||||
) -> None:
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
actor: Final = AgentResponse(
|
||||
agent_id="actor", agent_name="Actor", agent_card_params={}, rpm_limit=10, tpm_limit=1000
|
||||
)
|
||||
target: Final = AgentResponse(
|
||||
agent_id="target",
|
||||
agent_name="Target",
|
||||
agent_card_params={},
|
||||
rpm_limit=1,
|
||||
tpm_limit=1000,
|
||||
session_rpm_limit=1,
|
||||
session_tpm_limit=1000,
|
||||
)
|
||||
auth: Final = UserAPIKeyAuth(agent_id="actor")
|
||||
auth.managed_agent_policy = actor
|
||||
auth.invoked_agent_id = "actor" if self_call else "target"
|
||||
auth.invoked_agent_policy = actor if self_call else target
|
||||
cache: Final = DualCache()
|
||||
handler: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache))
|
||||
monkeypatch.setattr(handler, "_get_agent_from_registry", lambda _: None)
|
||||
descriptors: Final = handler._create_rate_limit_descriptors(
|
||||
user_api_key_dict=auth,
|
||||
data={"model": "a2a/target", "litellm_session_id": "session"},
|
||||
rpm_limit_type=None,
|
||||
tpm_limit_type=None,
|
||||
model_has_failures=False,
|
||||
)
|
||||
limits: Final = {(item["key"], item["value"]): item["rate_limit"]["requests_per_unit"] for item in descriptors}
|
||||
assert limits == (
|
||||
{("agent", "actor"): 10}
|
||||
if self_call
|
||||
else {("agent", "actor"): 10, ("agent", "target"): 1, ("agent_session", "target:session"): 1}
|
||||
)
|
||||
assert len(descriptors) == len(limits)
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=auth,
|
||||
cache=cache,
|
||||
data={
|
||||
"model": "gpt-4o-mini",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"max_tokens": 20,
|
||||
"litellm_session_id": "session",
|
||||
},
|
||||
call_type="acompletion",
|
||||
)
|
||||
stash: Final = get_request_stash()
|
||||
assert stash is not None and stash.reserved_tokens > 3
|
||||
response: Final = ModelResponse(usage=Usage(prompt_tokens=2, completion_tokens=1, total_tokens=3))
|
||||
operations: Final = handler._build_success_event_pipeline_operations(
|
||||
kwargs={"standard_logging_object": {"metadata": {"agent_id": auth.invoked_agent_id, "session_id": "session"}}},
|
||||
response_obj=response,
|
||||
rate_limit_type="total",
|
||||
)
|
||||
increments: Final = {op["key"]: op["increment_value"] for op in operations}
|
||||
for scope in stash.reserved_scopes:
|
||||
if scope[0] in ("agent", "agent_session"):
|
||||
assert increments[handler.create_rate_limit_keys(*scope, "tokens")] == 3 - stash.reserved_tokens
|
||||
|
||||
|
||||
@pytest.mark.parametrize("route", ["/a2a/expensive", "/a2a/expensive/message/send", "/v1/a2a/expensive/message/send"])
|
||||
async def test_a2a_url_target_owns_invocation_fee_and_request_limit(
|
||||
monkeypatch: pytest.MonkeyPatch, route: str
|
||||
) -> None:
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.agent_endpoints import agent_registry
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import invocation_target, prepare_agent_invocation
|
||||
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
expensive: Final = AgentResponse(
|
||||
agent_id="expensive", agent_name="Expensive", agent_card_params={}, rpm_limit=1,
|
||||
litellm_params={"cost_per_query": 0.25},
|
||||
)
|
||||
cheap: Final = AgentResponse(
|
||||
agent_id="cheap", agent_name="Cheap", agent_card_params={}, rpm_limit=100,
|
||||
litellm_params={"cost_per_query": 0.01},
|
||||
)
|
||||
registry: Final = agent_registry.AgentRegistry()
|
||||
registry.register_agent(expensive)
|
||||
registry.register_agent(cheap)
|
||||
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(
|
||||
side_effect=lambda where, include: {"expensive": expensive, "cheap": cheap}[where["agent_id"]]
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", database)
|
||||
auth: Final = UserAPIKeyAuth(agent_id="caller")
|
||||
auth.managed_agent_policy = AgentResponse(
|
||||
agent_id="caller", agent_name="Caller", agent_card_params={},
|
||||
object_permission={"object_permission_id": "both-targets", "agents": ["expensive", "cheap"]},
|
||||
)
|
||||
body: Final = {"model": "a2a/cheap"}
|
||||
target: Final = invocation_target(route, body)
|
||||
assert target is not None
|
||||
await prepare_agent_invocation(auth, target, AgentIdentityStore.from_client(database))
|
||||
assert auth.invoked_agent_id == "expensive"
|
||||
assert auth.invoked_agent_policy == expensive
|
||||
assert auth.agent_invocation_cost == pytest.approx(0.25)
|
||||
cache: Final = DualCache()
|
||||
limiter: Final = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache))
|
||||
await _rpm_request(limiter, cache, auth, "a2a/cheap")
|
||||
with pytest.raises(HTTPException) as denied:
|
||||
await _rpm_request(limiter, cache, auth, "a2a/cheap")
|
||||
assert denied.value.status_code == 429
|
||||
assert "expensive" in str(denied.value.detail)
|
||||
|
|
|
|||
|
|
@ -2714,3 +2714,34 @@ async def test_track_cost_callback_failure_alert_never_carries_request_metadata_
|
|||
assert "headers" in failure_debug_lines[0]
|
||||
else:
|
||||
assert failure_debug_lines == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("identity_field", ["agent_id", "billing_agent_id"])
|
||||
async def test_autonomous_llm_callback_persists_without_human_or_key(identity_field: str) -> None: # test-quality-ok: verifies anonymous-agent charges reach the persistence boundary; no injection seam
|
||||
kwargs: Final = {
|
||||
"call_type": "acompletion",
|
||||
"model": "test-model",
|
||||
"response_cost": 0.01,
|
||||
"litellm_params": {"metadata": {identity_field: "autonomous-agent"}},
|
||||
}
|
||||
with patch(
|
||||
"litellm.proxy.hooks.proxy_track_cost_callback._update_database_and_spend_counters",
|
||||
new_callable=AsyncMock,
|
||||
return_value=False,
|
||||
) as persist:
|
||||
await _ProxyDBLogger()._PROXY_track_cost_callback(
|
||||
kwargs=kwargs, completion_response=ModelResponse(), start_time=datetime.now(), end_time=datetime.now()
|
||||
)
|
||||
persist.assert_awaited_once()
|
||||
assert persist.call_args.kwargs["response_cost"] == 0.01
|
||||
assert persist.call_args.kwargs["user_id"] is None
|
||||
assert persist.call_args.kwargs["user_api_key"] is None
|
||||
assert persist.call_args.kwargs["kwargs"]["litellm_params"]["metadata"][identity_field] == "autonomous-agent"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("agent_id,expected", [(None, False), ("autonomous-agent", True)])
|
||||
def test_autonomous_agent_cost_tracking_needs_no_human_or_virtual_key(agent_id: str | None, expected: bool) -> None:
|
||||
assert _should_track_cost_callback(
|
||||
user_api_key=None, user_id=None, team_id=None, end_user_id=None, call_type="acompletion", agent_id=agent_id
|
||||
) is expected
|
||||
|
|
|
|||
|
|
@ -0,0 +1,114 @@
|
|||
from types import SimpleNamespace
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy.management_endpoints.sso.agent_subject_enrollment import (
|
||||
enroll_microsoft_subject,
|
||||
microsoft_interactive_subject,
|
||||
)
|
||||
|
||||
TENANT: Final = "11111111-1111-4111-8111-111111111111"
|
||||
OID: Final = "22222222-2222-4222-8222-222222222222"
|
||||
|
||||
|
||||
def test_enrollment_uses_provider_object_id_and_configured_tenant() -> None:
|
||||
subject: Final = microsoft_interactive_subject(
|
||||
TENANT, {"id": OID, "mail": "alias@example.com", "tid": "untrusted"}, {}
|
||||
)
|
||||
assert subject is not None
|
||||
assert subject.oid == OID
|
||||
assert subject.tenant_id == TENANT
|
||||
assert subject.issuer == f"https://login.microsoftonline.com/{TENANT}/v2.0"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("tenant", [None, "common", "organizations", "invalid"])
|
||||
def test_multitenant_sso_does_not_guess_the_subject_tenant(tenant: str | None) -> None:
|
||||
assert microsoft_interactive_subject(tenant, {"id": OID, "tid": TENANT}, {}) is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize("response", [{"mail": "user@example.com"}, {"id": "user@example.com"}, {"id": 42}])
|
||||
def test_email_and_configurable_aliases_are_not_human_subject_proof(response: dict[str, object]) -> None:
|
||||
assert microsoft_interactive_subject(TENANT, response, {}) is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"endpoint", ["MICROSOFT_USERINFO_ENDPOINT", "MICROSOFT_TOKEN_ENDPOINT", "MICROSOFT_AUTHORIZATION_ENDPOINT"]
|
||||
)
|
||||
def test_custom_provider_endpoints_do_not_enroll_trusted_microsoft_subjects(endpoint: str) -> None:
|
||||
assert microsoft_interactive_subject(TENANT, {"id": OID}, {endpoint: "https://custom.example"}) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_interactive_enrollment_preserves_the_canonical_local_user() -> None:
|
||||
table: Final = AsyncMock()
|
||||
table.upsert.return_value = SimpleNamespace(kind="human", user_id="canonical", verified_via="sso_interactive")
|
||||
client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_verifiedsubject=table))
|
||||
subject: Final = microsoft_interactive_subject(TENANT, {"id": OID}, {})
|
||||
assert subject is not None
|
||||
await enroll_microsoft_subject(subject, "canonical", client)
|
||||
table.upsert.assert_awaited_once_with(
|
||||
where={"issuer_tenant_id_oid": {"issuer": subject.issuer, "tenant_id": TENANT, "oid": OID}},
|
||||
data={
|
||||
"create": {
|
||||
"issuer": subject.issuer,
|
||||
"tenant_id": TENANT,
|
||||
"oid": OID,
|
||||
"user_id": "canonical",
|
||||
"verified_via": "sso_interactive",
|
||||
},
|
||||
"update": {},
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("user_id,verified_via", [("another-user", "sso_interactive"), ("canonical", "untrusted")])
|
||||
async def test_interactive_enrollment_does_not_reassign_an_existing_subject(user_id: str, verified_via: str) -> None:
|
||||
table: Final = AsyncMock()
|
||||
table.upsert.return_value = SimpleNamespace(kind="human", user_id=user_id, verified_via=verified_via)
|
||||
client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_verifiedsubject=table))
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await enroll_microsoft_subject(microsoft_interactive_subject(TENANT, {"id": OID}, {}), "canonical", client)
|
||||
assert failure.value.status_code == 403
|
||||
assert table.upsert.call_args.kwargs["data"]["update"] == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_enrollment_storage_failure_is_not_a_successful_login() -> None:
|
||||
table: Final = AsyncMock()
|
||||
table.upsert.side_effect = RuntimeError("database unavailable")
|
||||
client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_verifiedsubject=table))
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await enroll_microsoft_subject(microsoft_interactive_subject(TENANT, {"id": OID}, {}), "canonical", client)
|
||||
assert failure.value.status_code == 503
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("user_id", [None, "", 42])
|
||||
async def test_enrollment_requires_a_canonical_local_user(user_id: object) -> None:
|
||||
table: Final = AsyncMock()
|
||||
client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_verifiedsubject=table))
|
||||
await enroll_microsoft_subject(microsoft_interactive_subject(TENANT, {"id": OID}, {}), user_id, client)
|
||||
table.upsert.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_untrusted_metadata_cannot_enroll_a_human() -> None:
|
||||
table: Final = AsyncMock()
|
||||
client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_verifiedsubject=table))
|
||||
await enroll_microsoft_subject({"issuer": "forged", "tenant_id": TENANT, "oid": OID}, "canonical", client)
|
||||
table.upsert.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_scim_agent_subject_cannot_be_enrolled_as_a_human() -> None:
|
||||
table: Final = AsyncMock()
|
||||
table.upsert.return_value = SimpleNamespace(kind="agent_user", user_id=None, verified_via="scim")
|
||||
client: Final = SimpleNamespace(writer_db=SimpleNamespace(litellm_verifiedsubject=table))
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await enroll_microsoft_subject(microsoft_interactive_subject(TENANT, {"id": OID}, {}), "canonical", client)
|
||||
assert failure.value.status_code == 403
|
||||
assert table.upsert.call_args.kwargs["data"]["update"] == {}
|
||||
|
|
@ -206,6 +206,7 @@ def test_microsoft_sso_handler_openid_from_response_with_custom_attributes():
|
|||
def test_get_microsoft_callback_response():
|
||||
# Arrange
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.scope = {}
|
||||
mock_response = {
|
||||
"mail": "microsoft_user@example.com",
|
||||
"displayName": "Microsoft User",
|
||||
|
|
@ -2995,6 +2996,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
from litellm.proxy.management_endpoints.ui_sso import cli_sso_callback
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.scope = {}
|
||||
mock_request.base_url = "https://proxy.example.com/"
|
||||
|
||||
mock_user_info = LiteLLM_UserTable(
|
||||
|
|
@ -3158,6 +3160,7 @@ class TestCLIKeyRegenerationFlow:
|
|||
|
||||
# Mock request
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.scope = {}
|
||||
mock_request.base_url = "http://internal-proxy.local/"
|
||||
|
||||
# Test data
|
||||
|
|
@ -7106,6 +7109,7 @@ class TestCliSsoAttributionMetadata:
|
|||
from litellm.proxy.management_endpoints.types import CustomOpenID
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.scope = {}
|
||||
mock_request.base_url = "http://internal-proxy.local/"
|
||||
session_key = "cli-session-new-user"
|
||||
mock_user_info = LiteLLM_UserTable(
|
||||
|
|
@ -7220,6 +7224,7 @@ class TestCliSsoAttributionMetadata:
|
|||
)
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.scope = {}
|
||||
mock_request.base_url = "http://internal-proxy.local/"
|
||||
session_key = "cli-session-4567890"
|
||||
mock_user_info = LiteLLM_UserTable(
|
||||
|
|
@ -8751,6 +8756,7 @@ async def test_redirect_from_openid_persists_assertion_under_canonical_user_id()
|
|||
assertion = assertion_from_sso_login(_ema_id_token(), "rt_1")
|
||||
assert assertion is not None
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.scope = {}
|
||||
mock_request.base_url = "http://localhost:4000/"
|
||||
mock_request.cookies = {}
|
||||
|
||||
|
|
@ -8822,6 +8828,7 @@ async def test_cli_completion_persists_assertion_under_db_user_id():
|
|||
assertion = assertion_from_sso_login(_ema_id_token(), None)
|
||||
assert assertion is not None
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.scope = {}
|
||||
mock_request.base_url = "http://localhost:4000/"
|
||||
|
||||
user_info = MagicMock()
|
||||
|
|
@ -8989,6 +8996,7 @@ async def test_browser_funnel_reports_an_uncaptured_assertion(monkeypatch, caplo
|
|||
"""Wiring: the browser login path must reach the diagnostic, not just define it."""
|
||||
monkeypatch.setenv("GOOGLE_CLIENT_ID", "cid")
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.scope = {}
|
||||
mock_request.base_url = "http://localhost:4000/"
|
||||
mock_request.cookies = {}
|
||||
|
||||
|
|
@ -9059,6 +9067,7 @@ async def test_cli_funnel_reports_an_uncaptured_assertion(monkeypatch, caplog):
|
|||
|
||||
monkeypatch.setenv("MICROSOFT_CLIENT_ID", "cid")
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.scope = {}
|
||||
mock_request.base_url = "http://localhost:4000/"
|
||||
|
||||
user_info = MagicMock()
|
||||
|
|
@ -9134,6 +9143,7 @@ def _cli_callback_kwargs(flow):
|
|||
|
||||
def _cli_callback_request():
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.scope = {}
|
||||
mock_request.base_url = "http://localhost:4000/"
|
||||
return mock_request
|
||||
|
||||
|
|
@ -9438,3 +9448,45 @@ class TestSessionTokenCookie:
|
|||
resp = Response()
|
||||
set_session_token_cookie(resp, _make_http_request(), "jwt-token-value")
|
||||
assert "Secure" in self._cookie(resp)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("trusted", [False, True])
|
||||
@pytest.mark.parametrize("storage_available", [False, True])
|
||||
async def test_cli_sign_in_enrolls_only_verified_subjects_before_completing(
|
||||
monkeypatch: pytest.MonkeyPatch, trusted: bool, storage_available: bool
|
||||
) -> None:
|
||||
from typing import Final
|
||||
|
||||
from litellm.proxy.management_endpoints import ui_sso
|
||||
from litellm.types.proxy.agent_identity import MicrosoftInteractiveSubject
|
||||
|
||||
flow: Final[dict[str, object]] = {}
|
||||
kwargs: Final = _cli_callback_kwargs(flow)
|
||||
subject: Final = MicrosoftInteractiveSubject(issuer="issuer", tenant_id="tenant", oid="subject")
|
||||
kwargs["request"].scope = {"litellm_microsoft_interactive_subject": subject if trusted else subject.model_dump()}
|
||||
table: Final = kwargs["prisma_client"].writer_db.litellm_verifiedsubject
|
||||
table.upsert = AsyncMock(
|
||||
return_value=SimpleNamespace(kind="human", user_id="cli-user-id", verified_via="sso_interactive"),
|
||||
side_effect=None if storage_available else RuntimeError("storage unavailable"),
|
||||
)
|
||||
monkeypatch.setattr(ui_sso, "get_user_info_from_db", AsyncMock(return_value=_cli_callback_user_info([])))
|
||||
monkeypatch.setattr(ui_sso, "fetch_cli_sso_team_details", AsyncMock(return_value=()))
|
||||
monkeypatch.setattr(ui_sso, "retain_sso_identity_assertion_for_ema", AsyncMock())
|
||||
if trusted and not storage_available:
|
||||
with pytest.raises(HTTPException) as error:
|
||||
await ui_sso._complete_cli_sso_callback_session(**kwargs)
|
||||
assert error.value.status_code == 503
|
||||
assert "sso_complete" not in flow
|
||||
return
|
||||
response: Final = await ui_sso._complete_cli_sso_callback_session(**kwargs)
|
||||
assert response.status_code == 200
|
||||
assert flow["session_data"]["user_id"] == "cli-user-id"
|
||||
if trusted:
|
||||
table.upsert.assert_awaited_once_with(
|
||||
where={"issuer_tenant_id_oid": {"issuer": "issuer", "tenant_id": "tenant", "oid": "subject"}},
|
||||
data={"create": {"issuer": "issuer", "tenant_id": "tenant", "oid": "subject",
|
||||
"user_id": "cli-user-id", "verified_via": "sso_interactive"}, "update": {}},
|
||||
)
|
||||
else:
|
||||
table.upsert.assert_not_awaited()
|
||||
|
|
|
|||
|
|
@ -4699,24 +4699,28 @@ def _config_agent(agent_name: str) -> Dict[str, Any]:
|
|||
}
|
||||
|
||||
|
||||
class _FakeAgentRow:
|
||||
"""Stand-in for a prisma agent record: supports dict() and .object_permission."""
|
||||
def _agent_db_row(agent_id: str, agent_name: str):
|
||||
import json
|
||||
from datetime import datetime, timezone
|
||||
|
||||
def __init__(self, agent_id: str, agent_name: str) -> None:
|
||||
self.agent_id = agent_id
|
||||
self.agent_name = agent_name
|
||||
self.object_permission = None
|
||||
self.spend = 0.0
|
||||
from prisma.models import LiteLLM_AgentsTable
|
||||
|
||||
def __iter__(self):
|
||||
return iter(
|
||||
{
|
||||
"agent_id": self.agent_id,
|
||||
"agent_name": self.agent_name,
|
||||
"agent_card_params": {"name": self.agent_name, "url": "http://db-agent"},
|
||||
"litellm_params": {},
|
||||
}.items()
|
||||
)
|
||||
return LiteLLM_AgentsTable(
|
||||
agent_id=agent_id,
|
||||
agent_name=agent_name,
|
||||
agent_card_params=json.dumps({"name": agent_name, "url": "http://db-agent"}),
|
||||
extra_headers=[],
|
||||
agent_access_groups=[],
|
||||
access_group_ids=[],
|
||||
spend=0.0,
|
||||
identity_managed=False,
|
||||
enabled=True,
|
||||
execution_mode="autonomous",
|
||||
created_at=datetime.now(timezone.utc),
|
||||
updated_at=datetime.now(timezone.utc),
|
||||
created_by="admin",
|
||||
updated_by="admin",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -4740,7 +4744,7 @@ async def test_ProxyConfig__init_agents_in_db_keeps_config_defined_agents(clean_
|
|||
)
|
||||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_agentstable.find_many = AsyncMock(return_value=[_FakeAgentRow("db-id", "db-agent")])
|
||||
prisma_client.db.litellm_agentstable.find_many = AsyncMock(return_value=[_agent_db_row("db-id", "db-agent")])
|
||||
|
||||
await ProxyConfig()._init_agents_in_db(prisma_client=prisma_client)
|
||||
|
||||
|
|
@ -4777,7 +4781,7 @@ async def test_ProxyStartupEvent_jwt_auth_resolves_agent_claims_against_live_reg
|
|||
elif agents_source == "db":
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_agentstable.find_many = AsyncMock(
|
||||
return_value=[_FakeAgentRow("db-id", "loaded-agent")]
|
||||
return_value=[_agent_db_row("db-id", "loaded-agent")]
|
||||
)
|
||||
await ProxyConfig()._init_agents_in_db(prisma_client=prisma_client)
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ import datetime
|
|||
import hashlib
|
||||
import json
|
||||
import re
|
||||
import sqlite3
|
||||
from datetime import timezone
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
|
|
@ -3785,6 +3786,7 @@ class TestSpendLogsPayload:
|
|||
"status": "success",
|
||||
"mcp_namespaced_tool_name": None,
|
||||
"agent_id": None,
|
||||
"billing_agent_id": None,
|
||||
}
|
||||
)
|
||||
|
||||
|
|
@ -6590,9 +6592,7 @@ def test_key_spend_report_scopes_to_caller_key(client, monkeypatch):
|
|||
|
||||
|
||||
def test_key_spend_report_scopes_a_cli_session_to_the_per_user_alias_not_the_login_token(client, monkeypatch):
|
||||
mock_prisma = _spend_report_mock_prisma(
|
||||
query_raw_returns=[{"api_key": "cli-session-alice", "total_cost": 1.5}]
|
||||
)
|
||||
mock_prisma = _spend_report_mock_prisma(query_raw_returns=[{"api_key": "cli-session-alice", "total_cost": 1.5}])
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
|
||||
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||||
|
|
@ -7142,9 +7142,8 @@ async def test_ui_view_spend_logs_group_by_session_first_page(client, monkeypatc
|
|||
rep_call = emitted[2]
|
||||
assert f"DISTINCT ON ({SESSION_GROUP_KEY_SQL})" in rep_call[0]
|
||||
assert (
|
||||
f"ORDER BY {SESSION_GROUP_KEY_SQL}, call_type IN ('call_mcp_tool', 'list_mcp_tools'), \"startTime\" DESC"
|
||||
in rep_call[0]
|
||||
), "the session representative must prefer the newest non-MCP call"
|
||||
f"ORDER BY {SESSION_GROUP_KEY_SQL}, " + spend_management_endpoints._SESSION_REPRESENTATIVE_ORDER_SQL
|
||||
) in rep_call[0]
|
||||
assert rep_call[-2] == ["sess-1", "req-solo"]
|
||||
assert rep_call[-1] == ["hashed-key", "hashed-key"]
|
||||
finally:
|
||||
|
|
@ -7634,6 +7633,39 @@ def test_ui_view_request_response_internal_user_missing_row_forbidden(client, mo
|
|||
app.dependency_overrides.pop(ps.user_api_key_auth, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("parent_status", "child_status", "expected"),
|
||||
[("failure", "success", "failure"), ("success", "failure", "success")],
|
||||
)
|
||||
def test_session_representative_uses_completed_agent_outcome(parent_status, child_status, expected):
|
||||
with sqlite3.connect(":memory:") as connection:
|
||||
connection.execute(
|
||||
'CREATE TABLE logs (request_id TEXT, call_type TEXT, status TEXT, "startTime" TEXT, "endTime" TEXT)'
|
||||
)
|
||||
connection.executemany(
|
||||
"INSERT INTO logs VALUES (?, ?, ?, ?, ?)",
|
||||
(
|
||||
("parent", "asend_message", parent_status, "10:00:00", "10:00:05"),
|
||||
("nested-agent", "asend_message", child_status, "10:00:01", "10:00:03"),
|
||||
("llm", "acompletion", "success", "10:00:02", "10:00:04"),
|
||||
("tool", "call_mcp_tool", child_status, "10:00:04", "10:00:04"),
|
||||
),
|
||||
)
|
||||
result = connection.execute(
|
||||
"SELECT request_id, status FROM logs ORDER BY "
|
||||
+ spend_management_endpoints._SESSION_REPRESENTATIVE_ORDER_SQL
|
||||
+ " LIMIT 1"
|
||||
).fetchone()
|
||||
assert result == ("parent", expected)
|
||||
connection.execute("DELETE FROM logs WHERE call_type = 'asend_message'")
|
||||
fallback = connection.execute(
|
||||
"SELECT request_id, status FROM logs ORDER BY "
|
||||
+ spend_management_endpoints._SESSION_REPRESENTATIVE_ORDER_SQL
|
||||
+ " LIMIT 1"
|
||||
).fetchone()
|
||||
assert fallback == ("llm", "success")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_calculate_spend_unpriced_model_returns_400():
|
||||
model = "openrouter/unit-test-unpriced-model"
|
||||
|
|
|
|||
|
|
@ -539,7 +539,7 @@ async def test_spend_logs_ui_group_by_session_paginates_sessions(monkeypatch):
|
|||
async def mock_query_raw(sql_query, *params):
|
||||
if "COUNT(*) AS total_count" in sql_query:
|
||||
return [{"total_count": 60}]
|
||||
if "DISTINCT ON" in sql_query:
|
||||
if "AS session_representatives" in sql_query:
|
||||
return representative_rows
|
||||
return session_rows
|
||||
|
||||
|
|
@ -584,9 +584,11 @@ async def test_spend_logs_ui_group_by_session_paginates_sessions(monkeypatch):
|
|||
|
||||
rep_sql = emitted[2][0]
|
||||
assert f"DISTINCT ON ({group_key})" in rep_sql, f"page must return one row per session. SQL was:\n{rep_sql}"
|
||||
assert f"ORDER BY {group_key}, call_type IN ('call_mcp_tool', 'list_mcp_tools'), \"startTime\" DESC" in rep_sql, (
|
||||
"the session representative must prefer the newest non-MCP call"
|
||||
)
|
||||
assert (
|
||||
f"ORDER BY {group_key}, (call_type = 'asend_message') DESC, "
|
||||
"CASE WHEN call_type = 'asend_message' THEN \"endTime\" END DESC NULLS LAST, "
|
||||
"call_type IN ('call_mcp_tool', 'list_mcp_tools'), \"startTime\" DESC"
|
||||
) in rep_sql, "the session representative must prefer the final agent outcome, then the newest non-MCP call"
|
||||
assert "COUNT(*) OVER ()" not in rep_sql
|
||||
|
||||
assert [row["request_id"] for row in response["data"]] == ["req-1", "req-2"]
|
||||
|
|
|
|||
|
|
@ -5254,6 +5254,24 @@ def test_spend_log_request_id_is_the_response_id_a_bridged_messages_caller_recei
|
|||
)
|
||||
|
||||
|
||||
def test_failed_agent_request_keeps_registered_display_name():
|
||||
agent_model: Final = "a2a_agent/Research Agent"
|
||||
payload: Final = get_logging_payload(
|
||||
kwargs={
|
||||
"model": agent_model,
|
||||
"call_type": "asend_message",
|
||||
"litellm_params": {
|
||||
"metadata": {"model_group": agent_model, "model_info": {"id": "registered-agent"}, "status": "failure"}
|
||||
},
|
||||
},
|
||||
response_obj=ValueError("Agent action denied"),
|
||||
start_time=datetime.datetime.now(timezone.utc),
|
||||
end_time=datetime.datetime.now(timezone.utc),
|
||||
)
|
||||
assert payload["model"] == agent_model
|
||||
assert payload["status"] == "failure"
|
||||
assert payload["model_id"] == "registered-agent"
|
||||
|
||||
_CLI_SESSION_ALIAS: Final = "cli-session-alice"
|
||||
_CLI_SESSION_TOKEN: Final = "cli-session-Qm7xJ2kP9sLw4vT1nR8yAa"
|
||||
|
||||
|
|
@ -5379,3 +5397,21 @@ def test_baseline_estimate_metadata_comes_from_the_logging_stamp() -> None:
|
|||
assert result["autorouter_savings_estimate"] == recorded
|
||||
absent: Final = _get_spend_logs_metadata({"autorouter_savings_estimate": supplied}) # mutable-ok: legacy metadata helper accepts dicts
|
||||
assert absent["autorouter_savings_estimate"] is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize("billing_agent", [None, "authenticated-agent"])
|
||||
def test_untrusted_agent_label_cannot_replace_verified_billing_identity(billing_agent: str | None) -> None:
|
||||
kwargs = {
|
||||
"model": "gpt-4",
|
||||
"litellm_params": {"metadata": {
|
||||
"user_api_key": "test-key",
|
||||
"agent_id": "header-selected-agent",
|
||||
"billing_agent_id": billing_agent,
|
||||
}},
|
||||
}
|
||||
payload = get_logging_payload(
|
||||
kwargs=kwargs, response_obj={"id": "request"},
|
||||
start_time=datetime.datetime.now(timezone.utc), end_time=datetime.datetime.now(timezone.utc),
|
||||
)
|
||||
assert payload["agent_id"] == "header-selected-agent"
|
||||
assert payload["billing_agent_id"] == billing_agent
|
||||
|
|
|
|||
|
|
@ -149,6 +149,9 @@ async def test_route_a2a_model_read_through_recovers_agent_created_on_sibling_re
|
|||
prisma_client.db.litellm_agentstable.find_unique = AsyncMock(
|
||||
side_effect=[None, _DbAgentRow("a2a-sibling-replica-agent-id", agent_name)]
|
||||
)
|
||||
prisma_client.writer_db.litellm_agentstable.find_unique = AsyncMock(
|
||||
return_value=_DbAgentRow("a2a-sibling-replica-agent-id", agent_name)
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma_client)
|
||||
monkeypatch.setattr(proxy_server, "store_model_in_db", True)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue