fix(agents): keep managed permission ceilings authoritative

This commit is contained in:
Joshua Valluru 2026-09-30 10:07:04 -07:00
parent 41f149d5e2
commit bd3a240d47
8 changed files with 73 additions and 9 deletions

View file

@ -65,13 +65,20 @@ async def resolve_agent_access_group_ceiling(
agent_id: str,
load_access_group_ids: AccessGroupIdsLoader = _registry_access_group_ids,
load_access_group: AccessGroupLoader = _load_access_group,
*,
check_db_only: bool = False,
) -> AgentAccessGroupCeiling | None:
"""``None`` when the agent has no access groups attached, so nothing is capped."""
access_group_ids: Final = await load_access_group_ids(agent_id)
if not access_group_ids:
return None
loaded: Final = await asyncio.gather(*(load_access_group(group_id) for group_id in access_group_ids))
loaded: Final = await asyncio.gather(
*(
_load_access_group(group_id, check_db_only=True) if check_db_only else load_access_group(group_id)
for group_id in access_group_ids
)
)
groups: Final = tuple(group for group in loaded if group is not None)
return AgentAccessGroupCeiling(
access_group_ids=access_group_ids,

View file

@ -8,6 +8,7 @@ can only narrow access and need no trust.
"""
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final
from litellm._logging import verbose_proxy_logger
@ -45,7 +46,7 @@ def agent_caller_auth(user_api_key_auth: UserAPIKeyAuth) -> UserAPIKeyAuth | Non
user_id=caller.user_id,
team_id=caller.team_id,
parent_otel_span=user_api_key_auth.parent_otel_span,
)
).model_copy(update=MappingProxyType({"requires_fresh_policy": user_api_key_auth.requires_fresh_policy}))
async def load_agent_caller_team(user_api_key_auth: UserAPIKeyAuth) -> LiteLLM_TeamTable | None:

View file

@ -101,7 +101,9 @@ class AgentRequestHandler:
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)
agent_ceiling: Final = await AgentRequestHandler._agent_access_group_ceiling(
user_api_key_auth, resolve_ceiling, strict=strict
)
if agent_ceiling is None:
return own_access
if isinstance(own_access, UnrestrictedAgentAccess):
@ -137,10 +139,16 @@ class AgentRequestHandler:
async def _agent_access_group_ceiling(
user_api_key_auth: UserAPIKeyAuth | None,
resolve_ceiling: CeilingResolver,
*,
strict: bool = False,
) -> frozenset[str] | None:
if user_api_key_auth is None or not user_api_key_auth.agent_id:
return None
ceiling: Final = await resolve_ceiling(user_api_key_auth.agent_id)
ceiling: Final = (
await resolve_agent_access_group_ceiling(user_api_key_auth.agent_id, check_db_only=True)
if strict
else await resolve_ceiling(user_api_key_auth.agent_id)
)
if ceiling is None:
return None
return _to_stable_ids(ceiling.agent_ids)

View file

@ -3407,8 +3407,12 @@ async def get_access_object(
access_group_id,
)
raise HTTPException(
status_code=404,
detail={"error": f"Access group doesn't exist in db. Access group={access_group_id}. Error: {e}"},
status_code=503 if check_db_only else 404,
detail=(
"Access group policy is unavailable"
if check_db_only
else {"error": f"Access group doesn't exist in db. Access group={access_group_id}. Error: {e}"}
),
)

View file

@ -490,3 +490,45 @@ async def test_managed_agent_mcp_access_is_capped_at_the_invoking_callers_grants
assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == {"slack"}
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == ["read"]
@pytest.mark.asyncio
@pytest.mark.parametrize("fresh", [False, True])
@pytest.mark.parametrize("caller_kind", ["team", "user"])
async def test_caller_mcp_revocation_uses_fresh_policy(
monkeypatch: pytest.MonkeyPatch, fresh: bool, caller_kind: str,
) -> None:
from litellm.proxy._types import LiteLLM_TeamTable
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key
from litellm.types.agents import AgentCaller
cached_permission: Final = LiteLLM_ObjectPermissionTable(
object_permission_id="caller-permission", mcp_servers=["slack", "linear"],
mcp_tool_permissions={"slack": ["read", "write"]},
)
current_permission: Final = LiteLLM_ObjectPermissionTable(
object_permission_id="caller-permission", mcp_servers=["slack"],
mcp_tool_permissions={"slack": ["read"]},
)
team: Final = LiteLLM_TeamTable(
team_id="caller", object_permission_id="caller-permission", object_permission=current_permission,
)
user: Final = LiteLLM_UserTable(
user_id="caller", teams=[], object_permission_id="caller-permission", object_permission=current_permission,
)
database: Final = MagicMock()
database.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=team)
database.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=user)
database.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=current_permission)
cache: Final = UserApiKeyCache()
cache.set_cache("team_id:caller", team.model_copy(update={"object_permission": cached_permission}))
cache.set_cache("caller", user.model_copy(update={"object_permission": cached_permission}))
cache.set_cache(object_permission_cache_key("caller-permission"), cached_permission)
monkeypatch.setattr(proxy_server, "prisma_client", database)
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
auth: Final = actor(("read", "write"))
auth.requires_fresh_policy = fresh
auth.agent_caller = AgentCaller(team_id="caller") if caller_kind == "team" else AgentCaller(user_id="caller")
assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == ({"slack"} if fresh else {"slack", "linear"})
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == (["read"] if fresh else ["read", "write"])

View file

@ -163,7 +163,8 @@ async def test_authoritative_group_ceiling_propagates_policy_outages(
monkeypatch.setattr(proxy_server, "prisma_client", database)
monkeypatch.setattr(proxy_server, "user_api_key_cache", UserApiKeyCache())
if strict:
with pytest.raises(HTTPException):
with pytest.raises(HTTPException) as failure:
await _load_access_group("group", check_db_only=True)
assert failure.value.status_code == 503
else:
assert await _load_access_group("group") is None

View file

@ -1034,7 +1034,7 @@ async def test_managed_target_preserves_ordinary_actor_ceilings_after_key_reload
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("access_group_id:actor-group", group.model_copy(update={"access_agent_ids": ["target"]}))
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)

View file

@ -10068,7 +10068,8 @@ async def test_authoritative_access_group_outage_does_not_use_cached_grants() ->
cache.async_get_cache = AsyncMock()
with pytest.raises(HTTPException) as failure:
await get_access_object("group", client, cache, check_db_only=True)
assert failure.value.status_code == 404
assert failure.value.status_code == 503
assert failure.value.detail == "Access group policy is unavailable"
cache.async_get_cache.assert_not_awaited()