diff --git a/litellm/proxy/_experimental/mcp_server/auth/managed_agent_access.py b/litellm/proxy/_experimental/mcp_server/auth/managed_agent_access.py index ebb434495bb..096b7eb3c77 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/managed_agent_access.py +++ b/litellm/proxy/_experimental/mcp_server/auth/managed_agent_access.py @@ -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) diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 89a1d3781d8..0cc21cc01a4 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -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 ( diff --git a/litellm/proxy/agent_endpoints/auth/agent_access_groups.py b/litellm/proxy/agent_endpoints/auth/agent_access_groups.py index 6f0cc0393ad..4e5d85dea63 100644 --- a/litellm/proxy/agent_endpoints/auth/agent_access_groups.py +++ b/litellm/proxy/agent_endpoints/auth/agent_access_groups.py @@ -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 ) diff --git a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py index 4fd5df6cb44..779268fdaae 100644 --- a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py +++ b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py @@ -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": diff --git a/litellm/proxy/agent_endpoints/identity_store.py b/litellm/proxy/agent_endpoints/identity_store.py index 0d9d21108e5..3c8163a8838 100644 --- a/litellm/proxy/agent_endpoints/identity_store.py +++ b/litellm/proxy/agent_endpoints/identity_store.py @@ -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") diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index de1f04790b1..567310309c7 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -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, diff --git a/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_access_groups.py b/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_access_groups.py index e744e84d671..75eed80ab69 100644 --- a/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_access_groups.py +++ b/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_access_groups.py @@ -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 diff --git a/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_permission_handler.py b/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_permission_handler.py index 2b8aaea8981..9bce2ccacd5 100644 --- a/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_permission_handler.py +++ b/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_permission_handler.py @@ -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) diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index db50263a675..0f0147b0c6c 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -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"]) == []