mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(agents): preserve actor ceilings during managed target checks
This commit is contained in:
parent
1e6a76aa86
commit
41f149d5e2
9 changed files with 154 additions and 25 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"]) == []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue