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:
joshua-berri 2026-09-30 13:05:22 -07:00 • committed by GitHub
parent 2bf0cddcc1
commit 405ed414cb
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
42 changed files with 2383 additions and 265 deletions

View file

@ -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")

View file

@ -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,

View file

@ -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)

View file

@ -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,

View file

@ -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,

View file

@ -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")

View file

@ -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

View file

@ -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

View file

@ -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:

View file

@ -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

View file

@ -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:

View file

@ -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

View file

@ -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)

View file

@ -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

View file

@ -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,

View file

@ -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

View file

@ -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)

View file

@ -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,

View file

@ -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 []

View file

@ -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

View file

@ -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(

View file

@ -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

View file

@ -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(

View file

@ -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)

View file

@ -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]

View file

@ -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

View file

@ -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

View file

@ -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"

View file

@ -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"})

View file

@ -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)

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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)

View file

@ -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

View file

@ -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"] == {}

View file

@ -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()

View file

@ -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:

View file

@ -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"

View file

@ -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"]

View file

@ -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

View file

@ -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)