mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(agents): keep managed permission ceilings authoritative
This commit is contained in:
parent
41f149d5e2
commit
bd3a240d47
8 changed files with 73 additions and 9 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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}"}
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue