mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
feat(agents): require explicit invocation grants for managed identities
This commit is contained in:
parent
055174f8d3
commit
eaadd317ad
4 changed files with 775 additions and 51 deletions
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue