feat(agents): require explicit invocation grants for managed identities

This commit is contained in:
Joshua Valluru 2026-09-26 11:50:52 -07:00
parent 4d69db17c2
commit 7cbce50513
4 changed files with 775 additions and 51 deletions

View file

@ -8,8 +8,11 @@ Follows the same pattern as MCP permission handling.
import asyncio
from collections.abc import Awaitable, Callable, Sequence
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final, TypeAlias
from fastapi import HTTPException
from litellm._logging import verbose_logger
from litellm.proxy._experimental.mcp_server.ui_session_utils import build_effective_auth_contexts
from litellm.proxy._types import (
@ -86,6 +89,8 @@ class AgentRequestHandler:
) -> AgentAccess:
"""Agents the key may reach: key and team grants, intersected with the agent's access group ceiling
and, for an agent key acting on behalf of an invoking user, with that user's team grants."""
if user_api_key_auth is not None and user_api_key_auth.managed_agent_policy is not None:
return await _managed_actor_agent_access(user_api_key_auth)
key_team_access: Final = await AgentRequestHandler._resolve_key_team_agent_access(user_api_key_auth)
caller_access: Final = await AgentRequestHandler._agent_caller_access(user_api_key_auth)
own_access: Final = _intersect_agent_access(key_team_access, caller_access)
@ -106,11 +111,17 @@ class AgentRequestHandler:
@staticmethod
async def _resolve_key_team_agent_access(
user_api_key_auth: UserAPIKeyAuth | None,
*,
strict: bool = False,
) -> AgentAccess:
try:
key_access: Final = await AgentRequestHandler._get_allowed_agents_for_key(user_api_key_auth)
team_access: Final = await AgentRequestHandler._get_allowed_agents_for_team(user_api_key_auth)
key_access: Final = await AgentRequestHandler._get_allowed_agents_for_key(user_api_key_auth, strict=strict)
team_access: Final = await AgentRequestHandler._get_allowed_agents_for_team(
user_api_key_auth, strict=strict
)
except Exception as e:
if strict:
raise HTTPException(503, "Agent invocation policy is unavailable") from e
verbose_logger.warning("Failed to get allowed agents: %s", e)
return UnrestrictedAgentAccess()
return _intersect_agent_access(key_access, team_access)
@ -144,6 +155,33 @@ class AgentRequestHandler:
bool: True if agent is allowed, False otherwise
"""
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
from litellm.proxy.proxy_server import prisma_client
from litellm.types.proxy.agent_identity import AgentIdentityFailure
registered: Final = global_agent_registry.get_agent_by_id(agent_id)
if prisma_client is not None or (registered is not None and registered.identity_managed):
target: Final = await AgentIdentityStore.from_client(prisma_client).agent(agent_id)
if isinstance(target, AgentIdentityFailure):
raise_identity_failure(target)
if target is None and registered is not None and registered.identity_managed:
return False
if target is not None and target.identity_managed:
if (
not target.enabled
or target.identity is None
or not target.identity.active
or user_api_key_auth is None
):
return False
fresh_auth: Final = user_api_key_auth.model_copy(update={"requires_fresh_policy": True})
explicit: Final = await _granted_agent_ids(
fresh_auth,
_strict_agent_access,
build_effective_auth_contexts,
)
return target.agent_id in explicit
match await AgentRequestHandler.resolve_agent_access(user_api_key_auth, resolve_ceiling):
case UnrestrictedAgentAccess():
@ -204,6 +242,8 @@ class AgentRequestHandler:
@staticmethod
async def _get_allowed_agents_for_key(
user_api_key_auth: UserAPIKeyAuth | None = None,
*,
strict: bool = False,
) -> AgentAccess:
"""
Get allowed agents for a key.
@ -237,24 +277,36 @@ class AgentRequestHandler:
return UnrestrictedAgentAccess()
access_group_agents: Final = (
tuple(await AgentRequestHandler._get_agents_from_access_groups(list(declared_access_groups)))
tuple(
await AgentRequestHandler._get_agents_from_access_groups(
list(declared_access_groups), check_db_only=strict
)
)
if declared_access_groups
else ()
)
unified_agents: Final = (
tuple(await AgentRequestHandler._get_unified_access_group_agents(list(key_access_group_ids)))
tuple(
await AgentRequestHandler._get_unified_access_group_agents(
list(key_access_group_ids), check_db_only=strict
)
)
if key_access_group_ids
else ()
)
return RestrictedAgentAccess(frozenset(direct_agents + access_group_agents + unified_agents))
except Exception as e:
if strict:
raise HTTPException(503, "Agent invocation policy is unavailable") from e
verbose_logger.warning("Failed to get allowed agents for key: %s", e)
return UnrestrictedAgentAccess()
@staticmethod
async def _get_allowed_agents_for_team(
user_api_key_auth: UserAPIKeyAuth | None = None,
*,
strict: bool = False,
) -> AgentAccess:
"""
Get allowed agents for a team.
@ -280,7 +332,7 @@ class AgentRequestHandler:
)
if not prisma_client:
return UnrestrictedAgentAccess()
return RestrictedAgentAccess(frozenset()) if strict else UnrestrictedAgentAccess()
# Fetch the team object once for both permission sources
team_obj: Final = await get_team_object(
@ -289,10 +341,11 @@ class AgentRequestHandler:
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_db_only=strict,
)
if team_obj is None:
return UnrestrictedAgentAccess()
return RestrictedAgentAccess(frozenset()) if strict else UnrestrictedAgentAccess()
# 1. Get agents from object_permission (native permissions)
object_permissions: Final = team_obj.object_permission
@ -307,18 +360,28 @@ class AgentRequestHandler:
return UnrestrictedAgentAccess()
access_group_agents: Final = (
tuple(await AgentRequestHandler._get_agents_from_access_groups(list(declared_access_groups)))
tuple(
await AgentRequestHandler._get_agents_from_access_groups(
list(declared_access_groups), check_db_only=strict
)
)
if declared_access_groups
else ()
)
unified_agents: Final = (
tuple(await AgentRequestHandler._get_unified_access_group_agents(list(team_access_group_ids)))
tuple(
await AgentRequestHandler._get_unified_access_group_agents(
list(team_access_group_ids), check_db_only=strict
)
)
if team_access_group_ids
else ()
)
return RestrictedAgentAccess(frozenset(direct_agents + access_group_agents + unified_agents))
except Exception as e:
if strict:
raise HTTPException(503, "Agent invocation policy is unavailable") from e
# litellm-dashboard is the default UI team and will never have agents;
# skip noisy warnings for it.
if user_api_key_auth.team_id != UI_TEAM_ID:
@ -326,7 +389,9 @@ class AgentRequestHandler:
return UnrestrictedAgentAccess()
@staticmethod
def _get_config_agent_ids_for_access_groups(config_agents: list, access_groups: list[str]) -> set[str]:
def _get_config_agent_ids_for_access_groups(
config_agents: Sequence[AgentResponse], access_groups: list[str]
) -> set[str]:
"""
Helper to get agent_ids from config-loaded agents that match any of the given access groups.
"""
@ -339,7 +404,9 @@ class AgentRequestHandler:
return server_ids
@staticmethod
async def _get_db_agent_ids_for_access_groups(prisma_client, access_groups: list[str]) -> set[str]:
async def _get_db_agent_ids_for_access_groups(
prisma_client, access_groups: list[str], *, check_db_only: bool = False
) -> set[str]:
"""
Helper to get agent_ids from DB agents that match any of the given access groups.
@ -349,23 +416,27 @@ class AgentRequestHandler:
if not access_groups or prisma_client is None:
return set()
agents: Final = await AgentsRepository(prisma_client).table.find_many(
agents: Final = await AgentsRepository(prisma_client, use_writer=check_db_only).table.find_many(
where={"agent_access_groups": {"hasSome": access_groups}}
)
return {agent.agent_id for agent in agents}
@staticmethod
async def _get_unified_access_group_agents(access_group_ids: list[str]) -> list[str]:
async def _get_unified_access_group_agents(
access_group_ids: list[str], *, check_db_only: bool = False
) -> list[str]:
"""
Resolve unified access group ids to agent IDs.
"""
from litellm.proxy.auth.auth_checks import _get_agent_ids_from_access_groups
return await _get_agent_ids_from_access_groups(access_group_ids=access_group_ids)
return await _get_agent_ids_from_access_groups(access_group_ids=access_group_ids, check_db_only=check_db_only)
@staticmethod
async def _get_agents_from_access_groups(
access_groups: list[str],
*,
check_db_only: bool = False,
) -> list[str]:
"""
Resolve agent access groups to agent IDs by querying BOTH the agent table (DB) AND config-loaded agents.
@ -373,14 +444,13 @@ class AgentRequestHandler:
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
from litellm.proxy.proxy_server import prisma_client
# Use the helper for config-loaded agents
config_agent_ids: Final = AgentRequestHandler._get_config_agent_ids_for_access_groups(
global_agent_registry.agent_list, access_groups
)
# Use the helper for DB agents
db_agent_ids: Final = await AgentRequestHandler._get_db_agent_ids_for_access_groups(
prisma_client, access_groups
prisma_client, access_groups, check_db_only=check_db_only
)
return list(config_agent_ids | db_agent_ids)
@ -531,4 +601,58 @@ async def accessible_agents(
AgentRequestHandler.resolve_agent_access if resolve_access is None else resolve_access,
effective_contexts,
)
return tuple(agent for agent in agents if agent.agent_id in allowed_agent_ids)
allowed: Final = await asyncio.gather(
*(
AgentRequestHandler.is_agent_allowed(agent.agent_id, user_api_key_auth)
for agent in agents
if agent.identity_managed
)
)
managed_ids: Final = frozenset(
agent.agent_id
for agent, permitted in zip((agent for agent in agents if agent.identity_managed), allowed)
if permitted
)
return tuple(
agent
for agent in agents
if (agent.agent_id in managed_ids if agent.identity_managed else agent.agent_id in allowed_agent_ids)
)
async def _strict_agent_access(auth: UserAPIKeyAuth) -> AgentAccess:
if auth.managed_agent_policy is not None:
return await _managed_actor_agent_access(auth)
return await AgentRequestHandler._resolve_key_team_agent_access(auth, strict=True)
async def _managed_actor_agent_access(auth: UserAPIKeyAuth) -> AgentAccess:
agent: Final = auth.managed_agent_policy
if agent is None or not agent.object_permission:
return RestrictedAgentAccess(frozenset())
permission: Final = LiteLLM_ObjectPermissionTable.model_validate(agent.object_permission or MappingProxyType({}))
own_auth: Final = UserAPIKeyAuth(object_permission=permission)
own: Final = _granted_ids(await AgentRequestHandler._get_allowed_agents_for_key(own_auth, strict=True))
from litellm.proxy.agent_endpoints.auth.agent_access_groups import resolve_managed_agent_ceilings
ceilings: Final = await resolve_managed_agent_ceilings(agent)
capped: Final = frozenset(target for target in own if all(target in ceiling.agent_ids for ceiling in ceilings))
context: Final = auth.managed_agent_context
if context is None or context.mode == "autonomous":
return RestrictedAgentAccess(capped)
if context.user_id is None:
return RestrictedAgentAccess(frozenset())
human_ids: Final = await verified_human_agent_grants(context.user_id)
return RestrictedAgentAccess(capped.intersection(human_ids))
async def verified_human_agent_grants(user_id: 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()
human: Final = await MCPRequestHandler.reload_admitted_user(user_id, requires_fresh_policy=True)
sources: Final = await MCPRequestHandler._admitted_subject_sources(human)
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))

View file

@ -1,8 +1,12 @@
from collections.abc import Mapping
from itertools import product
from typing import Final
from types import MappingProxyType
from typing import Annotated, Final
from litellm.proxy._types import LiteLLMRoutes
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
@ -119,6 +123,49 @@ def managed_inference_request(
return {**body, "model": effective}
async def admit_managed_actor(auth: UserAPIKeyAuth, store: AgentIdentityStore | None) -> None:
if auth.agent_id is None:
return
if store is None:
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
registered: Final = global_agent_registry.get_agent_by_id(auth.agent_id)
if auth.managed_agent_context is not None or (
registered is not None and (registered.identity_managed or registered.identity is not None)
):
raise_identity_failure(
AgentIdentityFailure(code="policy_unavailable", message="Managed agent policy requires a database")
)
return
agent: Final = await store.agent(auth.agent_id)
if isinstance(agent, AgentIdentityFailure):
raise_identity_failure(agent)
if agent is None:
retired: Final = await store.retired_agent(auth.agent_id)
if isinstance(retired, AgentIdentityFailure):
raise_identity_failure(retired)
if auth.managed_agent_context is not None or retired:
raise_identity_failure(AgentIdentityFailure(message="Agent no longer exists"))
return
if not agent.identity_managed:
return
if auth.jwt_claims and auth.managed_agent_context is None:
raise_identity_failure(AgentIdentityFailure(message="A managed agent requires a matching verified identity"))
failure: Final = actor_admission_failure(agent, auth.managed_agent_context)
if failure is not None:
raise_identity_failure(failure)
auth.managed_agent_policy = agent
auth.billing_agent_policy = agent
if auth.managed_agent_context is not None and auth.managed_agent_context.mode == "delegated":
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import verified_human_agent_grants
grants: Final = await verified_human_agent_grants(auth.managed_agent_context.user_id)
if agent.agent_id not in grants:
raise_identity_failure(
AgentIdentityFailure(message="The delegated user is not permitted to invoke this agent")
)
def actor_admission_failure(
agent: AgentResponse,
context: ManagedAgentContext | None,
@ -136,6 +183,9 @@ def actor_admission_failure(
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:
model: Final = body.get("model")
if isinstance(model, str) and model.startswith("a2a/"):
@ -143,3 +193,40 @@ 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
return path[1] if len(path) >= 2 and path[0] == "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
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

@ -198,7 +198,7 @@ class TestAgentRequestHandler:
@staticmethod
def _team_grants(grants: dict[str, AgentAccess]) -> AsyncMock:
async def by_team(user_api_key_auth: UserAPIKeyAuth | None = None) -> AgentAccess:
async def by_team(user_api_key_auth: UserAPIKeyAuth | None = None, *, strict: bool = False) -> AgentAccess:
assert user_api_key_auth is not None
return grants.get(user_api_key_auth.team_id or "", UnrestrictedAgentAccess())
@ -249,7 +249,6 @@ class TestAgentRequestHandler:
frozenset({"agent-alpha"})
)
async def test_agent_access_groups_intersect_with_key_grants(self):
agent_key: Final = self._key_granting(["agent-alpha", "agent-beta"], agent_id="caller-agent")
resolve, _ = self._ceiling_resolver(frozenset({"agent-beta", "agent-gamma"}))
@ -632,3 +631,257 @@ class TestAgentRequestHandler:
assert await AgentRequestHandler.resolve_agent_access(
user_api_key_auth=mock_user_auth
) == RestrictedAgentAccess(frozenset({agent.agent_id})), (key_grant, team_grant)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"state,allowed",
[
({}, True),
({"enabled": False}, False),
],
)
async def test_managed_invocation_requires_local_and_directory_admission(
monkeypatch: pytest.MonkeyPatch, state: dict[str, object], allowed: bool
) -> None:
from unittest.mock import MagicMock
from litellm.proxy import proxy_server
from litellm.types.agents import AgentResponse
from litellm.types.proxy.agent_identity import AgentIdentityBinding
binding: Final = AgentIdentityBinding(
agent_id="target",
provider="microsoft_entra",
tenant_id="tenant",
client_id="client",
issuer="issuer",
revision="revision",
)
target: Final = AgentResponse(
agent_id="target", agent_name="Target", agent_card_params={}, identity=binding, identity_managed=True
).model_copy(update=state)
client: Final = MagicMock()
client.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target)
monkeypatch.setattr(proxy_server, "prisma_client", client)
permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="human-grant", agents=["target"])
auth: Final = UserAPIKeyAuth(user_id="human", object_permission=permission)
assert await AgentRequestHandler.is_agent_allowed("target", auth) is allowed
@pytest.mark.asyncio
@pytest.mark.parametrize("delegated", [True, False])
async def test_managed_agent_invocation_grants_intersect_verified_user_grants(
monkeypatch: pytest.MonkeyPatch, delegated: bool
) -> None:
from unittest.mock import MagicMock
from litellm.proxy import proxy_server
from litellm.proxy._types import LiteLLM_UserTable
from litellm.proxy.auth import auth_checks
from litellm.types.agents import AgentResponse
from litellm.types.proxy.agent_identity import ManagedAgentContext
database: Final = MagicMock()
monkeypatch.setattr(proxy_server, "prisma_client", database)
own: Final = LiteLLM_ObjectPermissionTable(object_permission_id="own", agents=["shared", "agent-only"])
human_grants: Final = LiteLLM_ObjectPermissionTable(object_permission_id="human", agents=["shared", "human-only"])
human: Final = LiteLLM_UserTable(user_id="human", teams=[], object_permission=human_grants)
monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(return_value=human))
auth: Final = UserAPIKeyAuth(agent_id="actor")
auth.managed_agent_policy = AgentResponse(
agent_id="actor", agent_name="Actor", agent_card_params={}, object_permission=own.model_dump()
)
auth.managed_agent_context = ManagedAgentContext(
agent_id="actor", mode="delegated" if delegated else "autonomous", user_id="human" if delegated else None
)
access: Final = await AgentRequestHandler.resolve_agent_access(auth)
assert access == RestrictedAgentAccess(frozenset({"shared"} if delegated else {"shared", "agent-only"}))
@pytest.mark.asyncio
@pytest.mark.parametrize("revoked", ["user", "team-member", "team-grant", "team-permission", "direct-grant", "access-group"])
async def test_delegated_grants_revoke_with_warm_user_team_and_permission_caches(
monkeypatch: pytest.MonkeyPatch, revoked: str
) -> None:
from unittest.mock import MagicMock
from litellm.proxy import proxy_server
from litellm.proxy._types import LiteLLM_AccessGroupTable, LiteLLM_TeamTable, LiteLLM_UserTable
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import verified_human_agent_grants
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key
direct: Final = revoked == "direct-grant"
grouped: Final = revoked == "access-group"
permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="grant", agents=["target"])
human: Final = LiteLLM_UserTable(
user_id="human",
teams=[] if direct else ["team"],
organization_memberships=[],
object_permission_id="grant" if direct else None,
)
team: Final = LiteLLM_TeamTable(
team_id="team",
models=[],
members_with_roles=[{"user_id": "human", "role": "user"}],
object_permission_id=None if grouped else "grant",
access_group_ids=["group"] if grouped else [],
)
group: Final = LiteLLM_AccessGroupTable(
access_group_id="group", access_group_name="Group", access_agent_ids=["target"]
)
cache: Final = UserApiKeyCache()
cache.set_cache("human", human)
cache.set_cache("team_id:team", team)
cache.set_cache(object_permission_cache_key("grant"), permission)
cache.set_cache("access_group_id:group", group)
client: Final = MagicMock()
client.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=human)
client.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=team)
client.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=permission)
client.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(return_value=group)
monkeypatch.setattr(proxy_server, "prisma_client", client)
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
assert await verified_human_agent_grants("human") == frozenset({"target"})
client.writer_db.litellm_usertable.find_unique.return_value = (
human.model_copy(update={"teams": []}) if revoked == "user" else human
)
client.writer_db.litellm_teamtable.find_unique.return_value = (
team.model_copy(update={"members_with_roles": []})
if revoked == "team-member"
else team.model_copy(update={"object_permission_id": None})
if revoked == "team-grant"
else team
)
client.writer_db.litellm_objectpermissiontable.find_unique.return_value = (
permission.model_copy(update={"agents": []}) if direct or revoked == "team-permission" else permission
)
client.writer_db.litellm_accessgrouptable.find_unique.return_value = (
group.model_copy(update={"access_agent_ids": []}) if grouped else group
)
assert await verified_human_agent_grants("human") == frozenset()
client.db.litellm_usertable.find_unique.assert_not_called()
client.db.litellm_teamtable.find_unique.assert_not_called()
client.db.litellm_objectpermissiontable.find_unique.assert_not_called()
client.db.litellm_accessgrouptable.find_unique.assert_not_called()
@pytest.mark.asyncio
async def test_strict_legacy_group_grants_ignore_stale_replica(monkeypatch: pytest.MonkeyPatch) -> None:
from unittest.mock import MagicMock
from litellm.proxy import proxy_server
from litellm.proxy.agent_endpoints import agent_registry
from litellm.types.agents import AgentResponse
stale: Final = AgentResponse(agent_id="revoked", agent_name="Revoked", agent_card_params={})
registry: Final = AgentRegistry()
registry.register_agent(stale)
database: Final = MagicMock()
database.db.litellm_agentstable.find_many = AsyncMock(return_value=[stale])
database.writer_db.litellm_agentstable.find_many = AsyncMock(return_value=[stale])
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
monkeypatch.setattr(proxy_server, "prisma_client", database)
auth: Final = UserAPIKeyAuth(
object_permission=LiteLLM_ObjectPermissionTable(
object_permission_id="permission", agent_access_groups=["group"]
)
)
assert await AgentRequestHandler._get_allowed_agents_for_key(auth, strict=True) == RestrictedAgentAccess(
frozenset({"revoked"})
)
database.writer_db.litellm_agentstable.find_many.return_value = []
assert await AgentRequestHandler._get_allowed_agents_for_key(auth, strict=True) == RestrictedAgentAccess(
frozenset()
)
database.db.litellm_agentstable.find_many.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("groups", [[], ["group"]])
async def test_legacy_groups_without_database_grant_no_agents(groups: list[str]) -> None:
assert await AgentRequestHandler._get_db_agent_ids_for_access_groups(None, groups, check_db_only=True) == set()
@pytest.mark.asyncio
@pytest.mark.parametrize("team", [False, True])
async def test_strict_invocation_policy_outage_denies_instead_of_allowing_all(
monkeypatch: pytest.MonkeyPatch, team: bool
) -> None:
from fastapi import HTTPException
from unittest.mock import MagicMock
from litellm.proxy import proxy_server
database: Final = MagicMock()
database.writer_db.litellm_teamtable.find_unique = AsyncMock(side_effect=ConnectionError("writer unavailable"))
database.writer_db.litellm_agentstable.find_many = AsyncMock(side_effect=ConnectionError("writer unavailable"))
monkeypatch.setattr(proxy_server, "prisma_client", database)
auth: Final = UserAPIKeyAuth(
team_id="team" if team else None,
object_permission=None if team else LiteLLM_ObjectPermissionTable(
object_permission_id="grant", agent_access_groups=["group"]
),
)
with pytest.raises(HTTPException, match="policy is unavailable") as denied:
await AgentRequestHandler._resolve_key_team_agent_access(auth, strict=True)
assert denied.value.status_code == 503
@pytest.mark.asyncio
@pytest.mark.parametrize("available", [False, True])
async def test_missing_team_cannot_grant_strict_agent_access(monkeypatch: pytest.MonkeyPatch, available: bool) -> None:
from unittest.mock import MagicMock
from litellm.proxy import proxy_server
from litellm.proxy.auth import auth_checks
monkeypatch.setattr(proxy_server, "prisma_client", MagicMock() if available else None)
monkeypatch.setattr(auth_checks, "get_team_object", AsyncMock(return_value=None))
assert await AgentRequestHandler._get_allowed_agents_for_team(
UserAPIKeyAuth(team_id="missing"), strict=True
) == RestrictedAgentAccess(frozenset())
@pytest.mark.asyncio
@pytest.mark.parametrize("outage", [False, True])
async def test_registered_managed_target_cannot_bypass_missing_or_unavailable_policy(
monkeypatch: pytest.MonkeyPatch, outage: bool
) -> None:
from fastapi import HTTPException
from unittest.mock import MagicMock
from litellm.proxy import proxy_server
from litellm.proxy.agent_endpoints import agent_registry
from litellm.types.agents import AgentResponse
registry: Final = AgentRegistry()
registry.register_agent(AgentResponse(
agent_id="target", agent_name="Target", agent_card_params={}, identity_managed=True
))
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
database: Final = MagicMock()
database.writer_db.litellm_agentstable.find_unique = AsyncMock(
return_value=None, side_effect=ConnectionError("unavailable") if outage else None
)
monkeypatch.setattr(proxy_server, "prisma_client", database)
if outage:
with pytest.raises(HTTPException, match="could not be loaded") as denied:
await AgentRequestHandler.is_agent_allowed("target", UserAPIKeyAuth())
assert denied.value.status_code == 503
else:
assert await AgentRequestHandler.is_agent_allowed("target", UserAPIKeyAuth()) is False
@pytest.mark.asyncio
@pytest.mark.parametrize("grant", [False, True])
async def test_delegation_without_a_verified_human_never_grants_agents(grant: bool) -> None:
from litellm.types.agents import AgentResponse
from litellm.types.proxy.agent_identity import ManagedAgentContext
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import verified_human_agent_grants
auth: Final = UserAPIKeyAuth(agent_id="actor")
auth.managed_agent_policy = AgentResponse(
agent_id="actor", agent_name="Actor", agent_card_params={},
object_permission={"object_permission_id": "own", "agents": ["target"]} if grant else None,
)
auth.managed_agent_context = ManagedAgentContext(agent_id="actor", mode="delegated")
assert await AgentRequestHandler.resolve_agent_access(auth) == RestrictedAgentAccess(frozenset())
assert await verified_human_agent_grants(None) == frozenset()

View file

@ -1,14 +1,16 @@
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.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
from litellm.types.proxy.agent_identity import AgentIdentityBinding, AgentIdentityFailure, ManagedAgentContext
@ -67,6 +69,93 @@ 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,
"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.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()
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=None)
database.writer_db.litellm_retiredagent.find_unique = AsyncMock(return_value={"original_agent_id": "deleted"})
with pytest.raises(HTTPException, match="Agent no longer exists"):
await admit_managed_actor(UserAPIKeyAuth(agent_id="deleted"), AgentIdentityStore.from_client(database))
database.writer_db.litellm_retiredagent.find_unique.return_value = None
auth: Final = UserAPIKeyAuth(agent_id="legacy-attribution-label")
await admit_managed_actor(auth, AgentIdentityStore.from_client(database))
assert auth.managed_agent_policy is None
database.db.litellm_agentstable.find_unique.assert_not_called()
@pytest.mark.asyncio
async def test_agent_history_outage_does_not_permit_legacy_fallback() -> None:
database: Final = MagicMock()
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=None)
database.writer_db.litellm_retiredagent.find_unique = AsyncMock(side_effect=RuntimeError("unavailable"))
with pytest.raises(HTTPException) as failure:
await admit_managed_actor(UserAPIKeyAuth(agent_id="deleted"), AgentIdentityStore.from_client(database))
assert failure.value.status_code == 503
@pytest.mark.parametrize(
"route,body,expected",
[
@ -82,6 +171,69 @@ def test_invocation_routes_resolve_the_same_target(route: str, body: dict[str, o
assert invocation_target(route, body) == expected
@pytest.mark.asyncio
async def test_agent_admission_database_outage_fails_closed() -> None:
database: Final = MagicMock()
database.writer_db.litellm_agentstable.find_unique = AsyncMock(side_effect=RuntimeError("DB unavailable"))
with pytest.raises(HTTPException) as failure:
await admit_managed_actor(UserAPIKeyAuth(agent_id="agent"), AgentIdentityStore.from_client(database))
assert failure.value.status_code == 503
@pytest.mark.asyncio
async def test_human_authentication_does_not_load_an_agent() -> None:
database: Final = MagicMock()
database.writer_db.litellm_agentstable.find_unique = AsyncMock()
await admit_managed_actor(UserAPIKeyAuth(user_id="human"), AgentIdentityStore.from_client(database))
database.writer_db.litellm_agentstable.find_unique.assert_not_awaited()
@pytest.mark.asyncio
async def test_disabled_agent_key_is_rejected_at_admission() -> None:
database: Final = MagicMock()
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent(enabled=False))
with pytest.raises(HTTPException) as failure:
await admit_managed_actor(UserAPIKeyAuth(agent_id="agent"), AgentIdentityStore.from_client(database))
assert failure.value.status_code == 403
@pytest.mark.asyncio
@pytest.mark.parametrize("permitted", [True, False])
async def test_verified_human_still_needs_an_explicit_agent_invocation_grant(
monkeypatch: pytest.MonkeyPatch,
permitted: bool,
) -> None:
from litellm.proxy import proxy_server
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_UserTable
from litellm.proxy.auth import auth_checks
policy: Final = agent()
database: Final = MagicMock()
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy)
monkeypatch.setattr(proxy_server, "prisma_client", database)
permission: Final = LiteLLM_ObjectPermissionTable(
object_permission_id="human-grants",
agents=["agent"] if permitted else [],
)
human: Final = LiteLLM_UserTable(user_id="human", teams=[], object_permission=permission)
monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(return_value=human))
auth: Final = UserAPIKeyAuth(agent_id="agent")
auth.managed_agent_context = ManagedAgentContext(
agent_id="agent",
binding_revision="current",
mode="delegated",
user_id="human",
)
if permitted:
await admit_managed_actor(auth, AgentIdentityStore.from_client(database))
assert auth.managed_agent_policy == policy
assert auth.billing_agent_policy == policy
else:
with pytest.raises(HTTPException) as failure:
await admit_managed_actor(auth, AgentIdentityStore.from_client(database))
assert failure.value.status_code == 403
def test_execution_mode_must_match_verified_token_mode() -> None:
context: Final = ManagedAgentContext(agent_id="agent", binding_revision="current", mode="autonomous")
failure: Final = actor_admission_failure(agent(execution_mode="delegated"), context)
@ -89,12 +241,106 @@ def test_execution_mode_must_match_verified_token_mode() -> None:
assert "execution mode" in failure.message
@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(
{"agent_id": "agent", "binding_revision": "current", "mode": mode, "user_id": user}
@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,
)
assert actor_admission_failure(agent(), context) is 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()
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent(execution_mode="autonomous"))
auth: Final = UserAPIKeyAuth(agent_id="agent", jwt_claims={"agent": "agent", "sub": "unrelated-subject"})
with pytest.raises(HTTPException) as denied:
await admit_managed_actor(auth, AgentIdentityStore.from_client(database))
assert denied.value.status_code == 403
assert auth.managed_agent_policy is None
@pytest.mark.asyncio
@pytest.mark.parametrize("bound", [False, True])
async def test_managed_context_or_binding_requires_database(monkeypatch: pytest.MonkeyPatch, bound: bool) -> None:
from litellm.proxy.agent_endpoints import agent_registry
from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry
registry: Final = AgentRegistry()
registry.register_agent(agent(identity_managed=bound, identity=BINDING if bound else None))
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
auth: Final = UserAPIKeyAuth(agent_id="agent")
if not bound:
auth.managed_agent_context = ManagedAgentContext(
agent_id="agent", binding_revision="current", mode="autonomous"
)
with pytest.raises(HTTPException) as denied:
await admit_managed_actor(auth, None)
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")
database: Final = MagicMock()
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy)
auth: Final = UserAPIKeyAuth(agent_id="agent", api_key="persisted-key")
with pytest.raises(HTTPException, match="bound identity provider token") as denied:
await admit_managed_actor(auth, AgentIdentityStore.from_client(database))
assert denied.value.status_code == 403
assert auth.managed_agent_policy is None
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(
@ -178,34 +424,48 @@ def test_managed_realtime_requires_a_model_and_ignores_completion_defaults(route
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"
)
assert managed_inference_request(
route, {"model": "requested"}, {"completion_model": "allowed-default"}, "cli"
)["model"] == "requested"
def test_caller_cannot_construct_trusted_subject_or_policy() -> None:
context: Final = ManagedAgentContext(
agent_id="agent", binding_revision="current", mode="delegated", user_id="human"
@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(
{"agent_id": "agent", "binding_revision": "current", "mode": mode, "user_id": user}
)
auth: Final = UserAPIKeyAuth.model_validate(
{
"managed_agent_context": context,
"requires_fresh_policy": 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.mcp_explicit_grants_only is False
assert "mcp_explicit_grants_only" not in auth.model_dump()
assert auth.managed_agent_context is None
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
assert auth.agent_invocation_cost is None
@pytest.mark.asyncio
async def test_bound_autonomous_actor_is_admitted_without_a_human() -> None:
database: Final = MagicMock()
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent())
auth: Final = UserAPIKeyAuth(agent_id="agent")
auth.managed_agent_context = ManagedAgentContext(agent_id="agent", binding_revision="current", mode="autonomous")
await admit_managed_actor(auth, AgentIdentityStore.from_client(database))
assert auth.managed_agent_policy == agent()
assert auth.billing_agent_policy == agent()
assert auth.user_id is None