fix(agents): preserve actor ceilings during managed target checks

This commit is contained in:
Joshua Valluru 2026-09-30 09:48:19 -07:00
parent 1e6a76aa86
commit 41f149d5e2
9 changed files with 154 additions and 25 deletions

View file

@ -30,7 +30,7 @@ async def managed_agent_servers(auth: UserAPIKeyAuth) -> tuple[str, ...]:
for ceiling in ceilings
)
grouped: Final = frozenset(server for server in base if all(server in ceiling for ceiling in expanded))
caller_capped, _ = await MCPRequestHandler._apply_agent_caller_ceiling(sorted(grouped), auth)
caller_capped, _ = await MCPRequestHandler.apply_agent_caller_ceiling(sorted(grouped), auth)
own: Final = frozenset(caller_capped)
context: Final = auth.managed_agent_context
if context is None or context.mode == "autonomous":
@ -55,7 +55,7 @@ async def managed_agent_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str]
return []
try:
granted: Final = await MCPRequestHandler.get_agent_tool_permissions_for_server(server_id, auth)
own: Final = await MCPRequestHandler._apply_agent_caller_tool_ceiling(granted, server_id, auth)
own: Final = await MCPRequestHandler.apply_agent_caller_tool_ceiling(granted, server_id, auth)
context: Final = auth.managed_agent_context
if context is None or context.mode == "autonomous":
return None if own is None else sorted(own)

View file

@ -1725,7 +1725,7 @@ class MCPRequestHandler:
#########################################################
# Cap an agent key at what the user and team that invoked the agent may reach
#########################################################
caller_capped, caller_restricts = await MCPRequestHandler._apply_agent_caller_ceiling(
caller_capped, caller_restricts = await MCPRequestHandler.apply_agent_caller_ceiling(
allowed_mcp_servers, user_api_key_auth
)
@ -2333,7 +2333,7 @@ class MCPRequestHandler:
)
allowed_tools = _as_list(
await MCPRequestHandler._apply_agent_caller_tool_ceiling(allowed_tools, server_id, user_api_key_auth)
await MCPRequestHandler.apply_agent_caller_tool_ceiling(allowed_tools, server_id, user_api_key_auth)
)
return await MCPRequestHandler._apply_agent_and_org_tool_ceilings(
@ -3169,7 +3169,7 @@ class MCPRequestHandler:
return capped, True
@staticmethod
async def _apply_agent_caller_ceiling(
async def apply_agent_caller_ceiling(
allowed_mcp_servers: Sequence[str],
user_api_key_auth: UserAPIKeyAuth | None = None,
) -> tuple[tuple[str, ...], bool]:
@ -3278,7 +3278,7 @@ class MCPRequestHandler:
return list(set(allowed_tools) & set(user_tools))
@staticmethod
async def _apply_agent_caller_tool_ceiling(
async def apply_agent_caller_tool_ceiling(
allowed_tools: Sequence[str] | None,
server_id: str,
user_api_key_auth: UserAPIKeyAuth | None = None,
@ -3286,7 +3286,7 @@ class MCPRequestHandler:
"""Narrow an agent key's tools on ``server_id`` to those the invoking user and team (echoed back
by the agent as ``x-litellm-user-id`` / ``x-litellm-team-id``) may call: the echoed team's tool
grants when it names any on this server, then the echoed user's own tool entitlement. The tools
axis twin of ``_apply_agent_caller_ceiling``, so the headers only ever narrow. Denies every tool
axis twin of ``apply_agent_caller_ceiling``, so the headers only ever narrow. Denies every tool
on the server when the caller's team cannot be loaded, since a caller we cannot resolve must not
read as unrestricted."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (

View file

@ -53,6 +53,8 @@ async def _load_access_group(access_group_id: str, *, check_db_only: bool = Fals
check_db_only=check_db_only,
)
except HTTPException as e:
if check_db_only:
raise
verbose_proxy_logger.warning(
"Agent access group %s could not be loaded, treating it as empty: %s", access_group_id, e.detail
)

View file

@ -87,13 +87,19 @@ class AgentRequestHandler:
async def resolve_agent_access(
user_api_key_auth: UserAPIKeyAuth | None = None,
resolve_ceiling: CeilingResolver = resolve_agent_access_group_ceiling,
*,
strict: bool = False,
) -> 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 managed_agent_policy(user_api_key_auth) 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)
key_team_access: Final = await AgentRequestHandler.resolve_key_team_agent_access(
user_api_key_auth, strict=strict
)
if strict and isinstance(key_team_access, UnrestrictedAgentAccess):
return RestrictedAgentAccess(frozenset())
caller_access: Final = await AgentRequestHandler.agent_caller_access(user_api_key_auth, strict=strict)
own_access: Final = _intersect_agent_access(key_team_access, caller_access)
agent_ceiling: Final = await AgentRequestHandler._agent_access_group_ceiling(user_api_key_auth, resolve_ceiling)
if agent_ceiling is None:
@ -103,11 +109,11 @@ class AgentRequestHandler:
return RestrictedAgentAccess(own_access.agent_ids & agent_ceiling)
@staticmethod
async def agent_caller_access(user_api_key_auth: UserAPIKeyAuth | None) -> AgentAccess:
async def agent_caller_access(user_api_key_auth: UserAPIKeyAuth | None, *, strict: bool = False) -> AgentAccess:
caller_auth: Final = agent_caller_auth(user_api_key_auth) if user_api_key_auth else None
if caller_auth is None:
return UnrestrictedAgentAccess()
return await AgentRequestHandler._get_allowed_agents_for_team(caller_auth)
return await AgentRequestHandler._get_allowed_agents_for_team(caller_auth, strict=strict)
@staticmethod
async def resolve_key_team_agent_access(
@ -188,7 +194,11 @@ class AgentRequestHandler:
and not user_api_key_auth.is_session_token
else user_api_key_auth
)
fresh_auth: Final = authority.model_copy(update={"requires_fresh_policy": True})
fresh_auth: Final = authority.model_copy(
update=MappingProxyType(
{"requires_fresh_policy": True, "agent_caller": user_api_key_auth.agent_caller}
)
)
explicit: Final = await _granted_agent_ids(
fresh_auth,
_strict_agent_access,
@ -292,7 +302,7 @@ class AgentRequestHandler:
access_group_agents: Final = (
tuple(
await AgentRequestHandler._get_agents_from_access_groups(
list(declared_access_groups), check_db_only=strict
declared_access_groups, check_db_only=strict
)
)
if declared_access_groups
@ -301,7 +311,7 @@ class AgentRequestHandler:
unified_agents: Final = (
tuple(
await AgentRequestHandler._get_unified_access_group_agents(
list(key_access_group_ids), check_db_only=strict
key_access_group_ids, check_db_only=strict
)
)
if key_access_group_ids
@ -375,7 +385,7 @@ class AgentRequestHandler:
access_group_agents: Final = (
tuple(
await AgentRequestHandler._get_agents_from_access_groups(
list(declared_access_groups), check_db_only=strict
declared_access_groups, check_db_only=strict
)
)
if declared_access_groups
@ -384,7 +394,7 @@ class AgentRequestHandler:
unified_agents: Final = (
tuple(
await AgentRequestHandler._get_unified_access_group_agents(
list(team_access_group_ids), check_db_only=strict
team_access_group_ids, check_db_only=strict
)
)
if team_access_group_ids
@ -403,7 +413,7 @@ class AgentRequestHandler:
@staticmethod
def _get_config_agent_ids_for_access_groups(
config_agents: Sequence[AgentResponse], access_groups: list[str]
config_agents: Sequence[AgentResponse], access_groups: Sequence[str]
) -> set[str]:
"""
Helper to get agent_ids from config-loaded agents that match any of the given access groups.
@ -418,7 +428,7 @@ class AgentRequestHandler:
@staticmethod
async def _get_db_agent_ids_for_access_groups(
prisma_client, access_groups: list[str], *, check_db_only: bool = False
prisma_client, access_groups: Sequence[str], *, check_db_only: bool = False
) -> set[str]:
"""
Helper to get agent_ids from DB agents that match any of the given access groups.
@ -436,7 +446,7 @@ class AgentRequestHandler:
@staticmethod
async def _get_unified_access_group_agents(
access_group_ids: list[str], *, check_db_only: bool = False
access_group_ids: Sequence[str], *, check_db_only: bool = False
) -> list[str]:
"""
Resolve unified access group ids to agent IDs.
@ -447,7 +457,7 @@ class AgentRequestHandler:
@staticmethod
async def _get_agents_from_access_groups(
access_groups: list[str],
access_groups: Sequence[str],
*,
check_db_only: bool = False,
) -> list[str]:
@ -634,9 +644,7 @@ async def accessible_agents(
async def _strict_agent_access(auth: UserAPIKeyAuth) -> AgentAccess:
if managed_agent_policy(auth) is not None:
return await _managed_actor_agent_access(auth)
return await AgentRequestHandler.resolve_key_team_agent_access(auth, strict=True)
return await AgentRequestHandler.resolve_agent_access(auth, strict=True)
async def _managed_actor_agent_access(auth: UserAPIKeyAuth) -> AgentAccess:
@ -651,7 +659,7 @@ async def _managed_actor_agent_access(auth: UserAPIKeyAuth) -> AgentAccess:
ceilings: Final = await resolve_managed_agent_ceilings(agent)
grouped: Final = frozenset(target for target in own if all(target in ceiling.agent_ids for ceiling in ceilings))
caller: Final = await AgentRequestHandler.agent_caller_access(auth)
caller: Final = await AgentRequestHandler.agent_caller_access(auth, strict=True)
capped: Final = grouped if isinstance(caller, UnrestrictedAgentAccess) else grouped & caller.agent_ids
context: Final = auth.managed_agent_context
if context is None or context.mode == "autonomous":

View file

@ -28,6 +28,7 @@ if TYPE_CHECKING:
LiteLLM_AgentIdentityWhereUniqueInput,
LiteLLM_AgentsTableInclude,
LiteLLM_AgentsTableWhereUniqueInput,
LiteLLM_RetiredAgentWhereUniqueInput,
LiteLLM_VerifiedSubjectCreateInput,
LiteLLM_VerifiedSubjectUpsertInput,
LiteLLM_VerifiedSubjectWhereUniqueInput,
@ -183,7 +184,8 @@ class AgentIdentityStore:
if self.retired_agents is None:
return AgentIdentityFailure(code="policy_unavailable", message="Agent history is unavailable")
try:
return await self.retired_agents.table.find_unique(where={"original_agent_id": agent_id}) is not None
where: Final[LiteLLM_RetiredAgentWhereUniqueInput] = {"original_agent_id": agent_id}
return await self.retired_agents.table.find_unique(where=where) is not None
except Exception:
return AgentIdentityFailure(code="policy_unavailable", message="Agent history is unavailable")

View file

@ -4280,6 +4280,8 @@ async def _get_resources_from_access_groups(
)
resources.extend(getattr(ag, resource_field, []))
except Exception:
if check_db_only:
raise
verbose_proxy_logger.debug(
"Could not fetch access group %s for resource field %s",
ag_id,

View file

@ -144,3 +144,26 @@ async def test_default_loader_returns_nothing_without_a_db(monkeypatch: pytest.M
monkeypatch.setattr(proxy_server, "prisma_client", None)
assert await _load_access_group("ag-1") is None
@pytest.mark.asyncio
@pytest.mark.parametrize("strict", [False, True])
async def test_authoritative_group_ceiling_propagates_policy_outages(
monkeypatch: pytest.MonkeyPatch, strict: bool
) -> None:
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy import proxy_server
from litellm.proxy.agent_endpoints.auth.agent_access_groups import _load_access_group
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
database: Final = MagicMock()
database.db.litellm_accessgrouptable.find_unique = AsyncMock(side_effect=RuntimeError("database unavailable"))
database.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(side_effect=RuntimeError("database unavailable"))
monkeypatch.setattr(proxy_server, "prisma_client", database)
monkeypatch.setattr(proxy_server, "user_api_key_cache", UserApiKeyCache())
if strict:
with pytest.raises(HTTPException):
await _load_access_group("group", check_db_only=True)
else:
assert await _load_access_group("group") is None

View file

@ -976,3 +976,70 @@ async def test_managed_target_rechecks_authoritative_key_after_peer_revocation(
await AgentRequestHandler.is_agent_allowed("target", warm)
else:
assert await AgentRequestHandler.is_agent_allowed("target", warm) is False
@pytest.mark.asyncio
@pytest.mark.parametrize("ceiling", ["agent-group", "caller-team", "group-without-grant"])
@pytest.mark.parametrize("permitted", [False, True])
async def test_managed_target_preserves_ordinary_actor_ceilings_after_key_reload(
monkeypatch: pytest.MonkeyPatch, ceiling: str, permitted: bool
) -> None:
from unittest.mock import MagicMock
from litellm.proxy import proxy_server
from litellm.proxy._types import LiteLLM_AccessGroupTable, LiteLLM_TeamTable
from litellm.proxy.agent_endpoints import agent_registry
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.types.agents import AgentResponse
from litellm.types.proxy.agent_identity import AgentIdentityBinding
target: Final = AgentResponse(
agent_id="target", agent_name="Target", agent_card_params={}, identity_managed=True,
identity=AgentIdentityBinding(
agent_id="target", provider="microsoft_entra", tenant_id="tenant", client_id="client",
issuer="issuer", revision="current",
),
)
actor: Final = AgentResponse(
agent_id="ordinary", agent_name="Ordinary", agent_card_params={},
access_group_ids=["actor-group"] if ceiling != "caller-team" else [],
)
registry: Final = AgentRegistry()
registry.register_agent(actor)
registry.register_agent(target)
permission: Final = LiteLLM_ObjectPermissionTable(
object_permission_id="key-grant", agents=[] if ceiling == "group-without-grant" else ["target"]
)
persisted: Final = UserAPIKeyAuth(
api_key="a" * 64, agent_id="ordinary", object_permission_id="key-grant", object_permission=permission,
)
auth: Final = persisted.model_copy()
auth.agent_caller = AgentCaller(team_id="caller-team") if ceiling == "caller-team" else None
group: Final = LiteLLM_AccessGroupTable(
access_group_id="actor-group", access_group_name="Actor group",
access_agent_ids=["target"] if permitted else ["other"],
)
team: Final = LiteLLM_TeamTable(
team_id="caller-team", object_permission_id="caller-grant",
object_permission=LiteLLM_ObjectPermissionTable(
object_permission_id="caller-grant", agents=["target"] if permitted else ["other"],
),
)
database: Final = MagicMock()
database.get_data = AsyncMock(return_value=persisted)
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target)
database.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(
side_effect=lambda where: permission if where["object_permission_id"] == "key-grant" else team.object_permission
)
database.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=team)
database.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(return_value=group)
cache: Final = UserApiKeyCache()
cache.set_cache("access_group_id:actor-group", group)
cache.set_cache("team_id:caller-team", team.model_copy(update={"object_permission": permission}))
monkeypatch.setattr(proxy_server, "prisma_client", database)
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
assert await AgentRequestHandler.is_agent_allowed("target", auth) is (permitted and ceiling != "group-without-grant")
database.get_data.assert_awaited_once()
assert auth.agent_caller == (AgentCaller(team_id="caller-team") if ceiling == "caller-team" else None)

View file

@ -10200,3 +10200,28 @@ async def test_authoritative_key_cannot_keep_grants_when_permission_is_unavailab
)
with pytest.raises(Exception, match=r"does not exist|unavailable"):
await get_key_object("hash", database, UserApiKeyCache(), check_db_only=True)
@pytest.mark.asyncio
@pytest.mark.parametrize("strict", [False, True])
async def test_authoritative_group_grants_propagate_policy_outages(
monkeypatch: pytest.MonkeyPatch, strict: bool
) -> None:
from unittest.mock import AsyncMock, MagicMock
from fastapi import HTTPException
from litellm.proxy import proxy_server
from litellm.proxy.auth.auth_checks import _get_agent_ids_from_access_groups
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
database: Final = MagicMock()
database.db.litellm_accessgrouptable.find_unique = AsyncMock(side_effect=RuntimeError("database unavailable"))
database.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(side_effect=RuntimeError("database unavailable"))
monkeypatch.setattr(proxy_server, "prisma_client", database)
monkeypatch.setattr(proxy_server, "user_api_key_cache", UserApiKeyCache())
if strict:
with pytest.raises(HTTPException):
await _get_agent_ids_from_access_groups(["group"], check_db_only=True)
else:
assert await _get_agent_ids_from_access_groups(["group"]) == []