mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
test(agents): inject the access group ceiling resolver instead of patching it
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
322262db01
commit
d743e08432
7 changed files with 193 additions and 201 deletions
|
|
@ -44,6 +44,10 @@ from litellm.proxy._types import (
|
|||
UserAPIKeyAuth,
|
||||
user_api_key_has_admin_view,
|
||||
)
|
||||
from litellm.proxy.agent_endpoints.auth.agent_access_groups import (
|
||||
CeilingResolver,
|
||||
resolve_agent_access_group_ceiling,
|
||||
)
|
||||
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
|
||||
from litellm.proxy.auth.user_api_key_auth import (
|
||||
_get_bearer_token_or_received_api_key, # pyright: ignore[reportPrivateUsage] # shared x-litellm-api-key parser lives with user_api_key_auth
|
||||
|
|
@ -184,6 +188,24 @@ def _has_client_supplied_mcp_auth(
|
|||
return bool(mcp_auth_header) or bool(mcp_server_auth_headers)
|
||||
|
||||
|
||||
def _agent_capped_servers(
|
||||
allowed_mcp_servers: Sequence[str],
|
||||
agent_servers: Sequence[str],
|
||||
agent_access_group_servers: frozenset[str] | None,
|
||||
) -> tuple[str, ...] | None:
|
||||
"""Servers left once the agent's object_permission and attached access groups both cap the
|
||||
key/team result, or None when the agent restricts nothing. An attached group set naming no
|
||||
server is an empty ceiling, not an absent one, so it denies every server."""
|
||||
if not agent_servers and agent_access_group_servers is None:
|
||||
return None
|
||||
return tuple(
|
||||
s
|
||||
for s in allowed_mcp_servers
|
||||
if (not agent_servers or s in agent_servers)
|
||||
and (agent_access_group_servers is None or s in agent_access_group_servers)
|
||||
)
|
||||
|
||||
|
||||
def _is_mcp_admitted_user_subject(user_api_key_auth: UserAPIKeyAuth | None) -> bool:
|
||||
"""True when this auth is a keyless subject admitted by the gateway session / bridge user
|
||||
path, as opposed to a JWT or other keyless auth that merely lacks a ``team_id``.
|
||||
|
|
@ -1546,21 +1568,14 @@ class MCPRequestHandler:
|
|||
# Check agent permissions if agent_id is set on the key
|
||||
#########################################################
|
||||
if user_api_key_auth and user_api_key_auth.agent_id:
|
||||
allowed_mcp_servers_for_agent: Final = await MCPRequestHandler._get_allowed_mcp_servers_for_agent(
|
||||
user_api_key_auth
|
||||
agent_capped: Final = _agent_capped_servers(
|
||||
allowed_mcp_servers,
|
||||
await MCPRequestHandler._get_allowed_mcp_servers_for_agent(user_api_key_auth),
|
||||
await MCPRequestHandler._get_agent_access_group_server_ceiling(user_api_key_auth),
|
||||
)
|
||||
agent_access_group_servers: Final = await MCPRequestHandler._get_agent_access_group_server_ceiling(
|
||||
user_api_key_auth
|
||||
)
|
||||
if len(allowed_mcp_servers_for_agent) > 0 or agent_access_group_servers is not None:
|
||||
if agent_capped is not None:
|
||||
has_lower_level_mcp_restrictions = True
|
||||
# Intersect: agent can only use servers allowed by key/team AND agent config AND agent access groups
|
||||
allowed_mcp_servers = [
|
||||
s
|
||||
for s in allowed_mcp_servers
|
||||
if (len(allowed_mcp_servers_for_agent) == 0 or s in allowed_mcp_servers_for_agent)
|
||||
and (agent_access_group_servers is None or s in agent_access_group_servers)
|
||||
]
|
||||
allowed_mcp_servers = list(agent_capped)
|
||||
verbose_logger.debug(
|
||||
"Applied agent intersection filter. Final allowed servers: %s", allowed_mcp_servers
|
||||
)
|
||||
|
|
@ -3148,6 +3163,7 @@ class MCPRequestHandler:
|
|||
@staticmethod
|
||||
async def _get_agent_access_group_server_ceiling(
|
||||
user_api_key_auth: UserAPIKeyAuth,
|
||||
resolve_ceiling: CeilingResolver = resolve_agent_access_group_ceiling,
|
||||
) -> frozenset[str] | None:
|
||||
"""
|
||||
Server IDs the agent's attached unified access groups (``LiteLLM_AgentsTable.access_group_ids``)
|
||||
|
|
@ -3157,13 +3173,10 @@ class MCPRequestHandler:
|
|||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.proxy.agent_endpoints.auth.agent_access_groups import (
|
||||
resolve_agent_access_group_ceiling,
|
||||
)
|
||||
|
||||
if not user_api_key_auth.agent_id:
|
||||
return None
|
||||
ceiling: Final = await resolve_agent_access_group_ceiling(user_api_key_auth.agent_id)
|
||||
ceiling: Final = await resolve_ceiling(user_api_key_auth.agent_id)
|
||||
if ceiling is None:
|
||||
return None
|
||||
return frozenset(global_mcp_server_manager.expand_permission_list(sorted(ceiling.mcp_server_ids)))
|
||||
|
|
|
|||
|
|
@ -33,6 +33,9 @@ class AgentAccessGroupCeiling:
|
|||
agent_ids: frozenset[str]
|
||||
|
||||
|
||||
CeilingResolver: TypeAlias = Callable[[str], Awaitable[AgentAccessGroupCeiling | None]] # mutable-ok: Callable params
|
||||
|
||||
|
||||
async def _load_agent(agent_id: str) -> AgentResponse | None:
|
||||
from litellm.proxy.common_utils.registry_read_through import get_agent_with_read_through
|
||||
|
||||
|
|
|
|||
|
|
@ -19,6 +19,10 @@ from litellm.proxy._types import (
|
|||
LitellmUserRoles,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.agent_endpoints.auth.agent_access_groups import (
|
||||
CeilingResolver,
|
||||
resolve_agent_access_group_ceiling,
|
||||
)
|
||||
from litellm.repositories.table_repositories import AgentsRepository
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
|
|
@ -61,6 +65,7 @@ class AgentRequestHandler:
|
|||
@staticmethod
|
||||
async def resolve_agent_access(
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
resolve_ceiling: CeilingResolver = resolve_agent_access_group_ceiling,
|
||||
) -> AgentAccess:
|
||||
"""
|
||||
Resolve the agents the given user/key may reach.
|
||||
|
|
@ -71,7 +76,7 @@ class AgentRequestHandler:
|
|||
never widen what it reaches.
|
||||
"""
|
||||
key_team_access: Final = await AgentRequestHandler._resolve_key_team_agent_access(user_api_key_auth)
|
||||
agent_ceiling: Final = await AgentRequestHandler._agent_access_group_ceiling(user_api_key_auth)
|
||||
agent_ceiling: Final = await AgentRequestHandler._agent_access_group_ceiling(user_api_key_auth, resolve_ceiling)
|
||||
if agent_ceiling is None:
|
||||
return key_team_access
|
||||
match key_team_access:
|
||||
|
|
@ -104,13 +109,12 @@ class AgentRequestHandler:
|
|||
@staticmethod
|
||||
async def _agent_access_group_ceiling(
|
||||
user_api_key_auth: UserAPIKeyAuth | None,
|
||||
resolve_ceiling: CeilingResolver,
|
||||
) -> frozenset[str] | None:
|
||||
"""Stable IDs of the agents the calling agent's attached access groups allow; None when none attached."""
|
||||
from litellm.proxy.agent_endpoints.auth.agent_access_groups import resolve_agent_access_group_ceiling
|
||||
|
||||
if user_api_key_auth is None or not user_api_key_auth.agent_id:
|
||||
return None
|
||||
ceiling: Final = await resolve_agent_access_group_ceiling(user_api_key_auth.agent_id)
|
||||
ceiling: Final = await resolve_ceiling(user_api_key_auth.agent_id)
|
||||
if ceiling is None:
|
||||
return None
|
||||
return _to_stable_ids(ceiling.agent_ids)
|
||||
|
|
@ -119,6 +123,7 @@ class AgentRequestHandler:
|
|||
async def is_agent_allowed(
|
||||
agent_id: str,
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
resolve_ceiling: CeilingResolver = resolve_agent_access_group_ceiling,
|
||||
) -> bool:
|
||||
"""
|
||||
Check if a specific agent is allowed for the given user/key.
|
||||
|
|
@ -132,7 +137,7 @@ class AgentRequestHandler:
|
|||
"""
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
|
||||
match await AgentRequestHandler.resolve_agent_access(user_api_key_auth):
|
||||
match await AgentRequestHandler.resolve_agent_access(user_api_key_auth, resolve_ceiling):
|
||||
case UnrestrictedAgentAccess():
|
||||
return True
|
||||
case RestrictedAgentAccess(allowed_agent_ids):
|
||||
|
|
|
|||
|
|
@ -68,6 +68,10 @@ from litellm.proxy._types import (
|
|||
SpecialModelNames,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.agent_endpoints.auth.agent_access_groups import (
|
||||
CeilingResolver,
|
||||
resolve_agent_access_group_ceiling,
|
||||
)
|
||||
from litellm.proxy.auth.budget_throttle import (
|
||||
budget_throttle_percentage,
|
||||
should_throttle_budget_exceeded,
|
||||
|
|
@ -4199,15 +4203,14 @@ async def _check_agent_access_group_model_access(
|
|||
model: str | list[str] | None, # mutable-ok: _can_object_call_model and the client message helper take list[str]
|
||||
valid_token: UserAPIKeyAuth | None,
|
||||
llm_router: Router | None,
|
||||
resolve_ceiling: CeilingResolver = resolve_agent_access_group_ceiling,
|
||||
) -> Literal[True]:
|
||||
"""Raises when the key's agent has access groups attached and none of them names the model.
|
||||
Attached groups that name no model deny every model; ``_can_object_call_model`` would read
|
||||
an empty allowlist as unrestricted."""
|
||||
from litellm.proxy.agent_endpoints.auth.agent_access_groups import resolve_agent_access_group_ceiling
|
||||
|
||||
if not model or valid_token is None or not valid_token.agent_id:
|
||||
return True
|
||||
ceiling: Final = await resolve_agent_access_group_ceiling(valid_token.agent_id)
|
||||
ceiling: Final = await resolve_ceiling(valid_token.agent_id)
|
||||
if ceiling is None:
|
||||
return True
|
||||
if not ceiling.models:
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ from starlette.datastructures import Headers
|
|||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
|
||||
MCPRequestHandler,
|
||||
UnloadableEntitlementError,
|
||||
_agent_capped_servers,
|
||||
_is_mcp_admitted_user_subject,
|
||||
)
|
||||
from litellm.proxy._types import (
|
||||
|
|
@ -4169,6 +4170,27 @@ async def test_get_allowed_mcp_servers_for_key_prefers_in_memory_permission():
|
|||
global_mcp_server_manager.registry.pop("direct-server", None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("agent_servers", "group_ceiling", "expected"),
|
||||
[
|
||||
([], frozenset({"server_1"}), ("server_1",)),
|
||||
([], frozenset({"server_1", "server_2", "server_3"}), ("server_1", "server_2")),
|
||||
([], frozenset(), ()),
|
||||
(["server_2"], frozenset({"server_1", "server_2"}), ("server_2",)),
|
||||
(["server_1"], frozenset({"server_2"}), ()),
|
||||
(["server_1"], None, ("server_1",)),
|
||||
],
|
||||
)
|
||||
def test_agent_capped_servers_intersects_agent_config_and_access_groups(agent_servers, group_ceiling, expected):
|
||||
"""The agent's attached access groups cap the key/team servers alongside its own
|
||||
object_permission; groups naming no server deny all."""
|
||||
assert _agent_capped_servers(["server_1", "server_2"], agent_servers, group_ceiling) == expected
|
||||
|
||||
|
||||
def test_agent_capped_servers_without_agent_restrictions_is_uncapped():
|
||||
assert _agent_capped_servers(["server_1", "server_2"], [], None) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestAgentMCPPermissions:
|
||||
"""Test agent-level MCP server and tool permission intersection."""
|
||||
|
|
@ -4208,64 +4230,45 @@ class TestAgentMCPPermissions:
|
|||
assert sorted(result) == ["server_1", "server_2"]
|
||||
mock_agent.assert_called_once_with(user_api_key_auth)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("group_ceiling", "expected"),
|
||||
[
|
||||
(frozenset({"server_1"}), ["server_1"]),
|
||||
(frozenset({"server_1", "server_2", "server_3"}), ["server_1", "server_2"]),
|
||||
(frozenset(), []),
|
||||
],
|
||||
)
|
||||
async def test_get_allowed_mcp_servers_agent_access_group_ceiling(self, group_ceiling, expected):
|
||||
"""The agent's attached access groups cap the key/team servers; groups naming no server deny all."""
|
||||
user_api_key_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user", agent_id="agent-ag")
|
||||
with (
|
||||
patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_key", return_value=["server_1", "server_2"]),
|
||||
patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_team", return_value=[]),
|
||||
patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_agent", return_value=[]),
|
||||
patch.object(MCPRequestHandler, "_get_agent_access_group_server_ceiling", return_value=group_ceiling),
|
||||
):
|
||||
access = await MCPRequestHandler.get_mcp_server_access(user_api_key_auth=user_api_key_auth)
|
||||
assert sorted(access.server_ids) == expected
|
||||
assert access.scope == "scoped"
|
||||
|
||||
async def test_get_allowed_mcp_servers_agent_without_access_groups_is_uncapped(self):
|
||||
user_api_key_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user", agent_id="agent-ag")
|
||||
with (
|
||||
patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_key", return_value=["server_1", "server_2"]),
|
||||
patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_team", return_value=[]),
|
||||
patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_agent", return_value=[]),
|
||||
patch.object(MCPRequestHandler, "_get_agent_access_group_server_ceiling", return_value=None),
|
||||
):
|
||||
result = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth=user_api_key_auth)
|
||||
assert sorted(result) == ["server_1", "server_2"]
|
||||
|
||||
async def test_agent_access_group_server_ceiling_expands_group_servers(self):
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
from litellm.proxy.agent_endpoints.auth.agent_access_groups import AgentAccessGroupCeiling
|
||||
from litellm.types.mcp import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
ceiling = AgentAccessGroupCeiling(
|
||||
access_group_ids=("ag-1",),
|
||||
models=frozenset(),
|
||||
mcp_server_ids=frozenset({"server_1"}),
|
||||
agent_ids=frozenset(),
|
||||
)
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.agent_endpoints.auth.agent_access_groups.resolve_agent_access_group_ceiling",
|
||||
new=AsyncMock(return_value=ceiling),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
|
||||
) as mock_manager,
|
||||
):
|
||||
mock_manager.expand_permission_list.return_value = ["server_1"]
|
||||
result = await MCPRequestHandler._get_agent_access_group_server_ceiling(
|
||||
UserAPIKeyAuth(api_key="test-key", agent_id="agent-ag")
|
||||
asked: list[str] = []
|
||||
|
||||
async def resolve(agent_id: str) -> AgentAccessGroupCeiling | None:
|
||||
asked.append(agent_id)
|
||||
return AgentAccessGroupCeiling(
|
||||
access_group_ids=("ag-1",),
|
||||
models=frozenset(),
|
||||
mcp_server_ids=frozenset({"aliased-server"}),
|
||||
agent_ids=frozenset(),
|
||||
)
|
||||
assert result == frozenset({"server_1"})
|
||||
mock_manager.expand_permission_list.assert_called_once_with(["server_1"])
|
||||
|
||||
assert await MCPRequestHandler._get_agent_access_group_server_ceiling(UserAPIKeyAuth(api_key="k")) is None
|
||||
global_mcp_server_manager.registry["ag-server-id"] = MCPServer(
|
||||
server_id="ag-server-id",
|
||||
name="ag-server",
|
||||
server_name="ag-server",
|
||||
alias="aliased-server",
|
||||
url="https://ag-server.example.com",
|
||||
transport=MCPTransport.http,
|
||||
)
|
||||
try:
|
||||
result = await MCPRequestHandler._get_agent_access_group_server_ceiling(
|
||||
UserAPIKeyAuth(api_key="test-key", agent_id="agent-ag"), resolve
|
||||
)
|
||||
finally:
|
||||
global_mcp_server_manager.registry.pop("ag-server-id", None)
|
||||
|
||||
assert result == frozenset({"ag-server-id"})
|
||||
assert asked == ["agent-ag"]
|
||||
assert (
|
||||
await MCPRequestHandler._get_agent_access_group_server_ceiling(UserAPIKeyAuth(api_key="k"), resolve)
|
||||
is None
|
||||
)
|
||||
assert asked == ["agent-ag"]
|
||||
|
||||
async def test_get_allowed_mcp_servers_key_team_agent_intersection(self):
|
||||
"""Key allows [1, 2], agent allows [2, 3]. Result = [2]."""
|
||||
|
|
|
|||
|
|
@ -11,9 +11,9 @@ import pytest
|
|||
|
||||
|
||||
from litellm.constants import UI_SESSION_TOKEN_TEAM_ID
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry
|
||||
from litellm.proxy.agent_endpoints.auth.agent_access_groups import AgentAccessGroupCeiling
|
||||
from litellm.proxy.agent_endpoints.auth.agent_access_groups import AgentAccessGroupCeiling, CeilingResolver
|
||||
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import (
|
||||
AgentAccess,
|
||||
AgentRequestHandler,
|
||||
|
|
@ -159,86 +159,74 @@ class TestAgentRequestHandler:
|
|||
), agent_id
|
||||
|
||||
@staticmethod
|
||||
def _ceiling(agent_ids: frozenset[str]) -> AgentAccessGroupCeiling:
|
||||
return AgentAccessGroupCeiling(
|
||||
access_group_ids=("ag-1",),
|
||||
models=frozenset(),
|
||||
mcp_server_ids=frozenset(),
|
||||
agent_ids=agent_ids,
|
||||
def _ceiling_resolver(agent_ids: frozenset[str] | None) -> tuple[CeilingResolver, list[str]]:
|
||||
"""A resolver that records the agent ids it was asked about and answers with a fixed
|
||||
ceiling, or None when the agent has no access groups attached."""
|
||||
asked: Final[list[str]] = []
|
||||
|
||||
async def resolve(agent_id: str) -> AgentAccessGroupCeiling | None:
|
||||
asked.append(agent_id)
|
||||
if agent_ids is None:
|
||||
return None
|
||||
return AgentAccessGroupCeiling(
|
||||
access_group_ids=("ag-1",), models=frozenset(), mcp_server_ids=frozenset(), agent_ids=agent_ids
|
||||
)
|
||||
|
||||
return resolve, asked
|
||||
|
||||
@staticmethod
|
||||
def _key_granting(agent_ids: list[str], agent_id: str | None) -> UserAPIKeyAuth:
|
||||
return UserAPIKeyAuth(
|
||||
api_key="test-key",
|
||||
user_id="test-user",
|
||||
agent_id=agent_id,
|
||||
object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="obj-1", agents=agent_ids),
|
||||
)
|
||||
|
||||
async def test_agent_access_groups_cap_an_otherwise_unrestricted_key(self):
|
||||
"""A key with no agent grant of its own may still only reach the agents its
|
||||
agent's attached access groups name."""
|
||||
agent_key: Final = UserAPIKeyAuth(api_key="test-key", user_id="test-user", agent_id="caller-agent")
|
||||
resolve, asked = self._ceiling_resolver(frozenset({"agent-beta"}))
|
||||
|
||||
with (
|
||||
patch.object(AgentRequestHandler, "_get_allowed_agents_for_key", return_value=UnrestrictedAgentAccess()),
|
||||
patch.object(AgentRequestHandler, "_get_allowed_agents_for_team", return_value=UnrestrictedAgentAccess()),
|
||||
patch(
|
||||
"litellm.proxy.agent_endpoints.auth.agent_access_groups.resolve_agent_access_group_ceiling",
|
||||
new=AsyncMock(return_value=self._ceiling(frozenset({"agent-beta"}))),
|
||||
) as mock_ceiling,
|
||||
):
|
||||
assert await AgentRequestHandler.resolve_agent_access(agent_key) == RestrictedAgentAccess(
|
||||
frozenset({"agent-beta"})
|
||||
)
|
||||
assert await AgentRequestHandler.is_agent_allowed("agent-beta", agent_key) is True
|
||||
assert await AgentRequestHandler.is_agent_allowed("agent-alpha", agent_key) is False
|
||||
mock_ceiling.assert_called_with("caller-agent")
|
||||
|
||||
async def test_agent_access_groups_intersect_with_key_and_team_grants(self):
|
||||
agent_key: Final = UserAPIKeyAuth(
|
||||
api_key="test-key", user_id="test-user", team_id="test-team", agent_id="caller-agent"
|
||||
assert await AgentRequestHandler.resolve_agent_access(agent_key, resolve) == RestrictedAgentAccess(
|
||||
frozenset({"agent-beta"})
|
||||
)
|
||||
assert await AgentRequestHandler.is_agent_allowed("agent-beta", agent_key, resolve) is True
|
||||
assert await AgentRequestHandler.is_agent_allowed("agent-alpha", agent_key, resolve) is False
|
||||
assert asked == ["caller-agent"] * 3
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
AgentRequestHandler,
|
||||
"_get_allowed_agents_for_key",
|
||||
return_value=RestrictedAgentAccess(frozenset({"agent-alpha", "agent-beta"})),
|
||||
),
|
||||
patch.object(
|
||||
AgentRequestHandler,
|
||||
"_get_allowed_agents_for_team",
|
||||
return_value=RestrictedAgentAccess(frozenset({"agent-alpha", "agent-beta", "agent-gamma"})),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.agent_endpoints.auth.agent_access_groups.resolve_agent_access_group_ceiling",
|
||||
new=AsyncMock(return_value=self._ceiling(frozenset({"agent-beta", "agent-gamma"}))),
|
||||
),
|
||||
):
|
||||
assert await AgentRequestHandler.resolve_agent_access(agent_key) == RestrictedAgentAccess(
|
||||
frozenset({"agent-beta"})
|
||||
)
|
||||
async def test_agent_access_groups_intersect_with_key_grants(self):
|
||||
agent_key: Final = self._key_granting(["agent-alpha", "agent-beta"], agent_id="caller-agent")
|
||||
resolve, _ = self._ceiling_resolver(frozenset({"agent-beta", "agent-gamma"}))
|
||||
|
||||
assert await AgentRequestHandler.resolve_agent_access(agent_key, resolve) == RestrictedAgentAccess(
|
||||
frozenset({"agent-beta"})
|
||||
)
|
||||
assert await AgentRequestHandler.is_agent_allowed("agent-gamma", agent_key, resolve) is False
|
||||
|
||||
async def test_agent_access_groups_naming_no_agent_deny_every_agent(self):
|
||||
agent_key: Final = UserAPIKeyAuth(api_key="test-key", user_id="test-user", agent_id="caller-agent")
|
||||
resolve, _ = self._ceiling_resolver(frozenset())
|
||||
|
||||
with (
|
||||
patch.object(AgentRequestHandler, "_get_allowed_agents_for_key", return_value=UnrestrictedAgentAccess()),
|
||||
patch.object(AgentRequestHandler, "_get_allowed_agents_for_team", return_value=UnrestrictedAgentAccess()),
|
||||
patch(
|
||||
"litellm.proxy.agent_endpoints.auth.agent_access_groups.resolve_agent_access_group_ceiling",
|
||||
new=AsyncMock(return_value=self._ceiling(frozenset())),
|
||||
),
|
||||
):
|
||||
assert await AgentRequestHandler.resolve_agent_access(agent_key) == RestrictedAgentAccess(frozenset())
|
||||
assert await AgentRequestHandler.is_agent_allowed("agent-alpha", agent_key) is False
|
||||
assert await AgentRequestHandler.resolve_agent_access(agent_key, resolve) == RestrictedAgentAccess(frozenset())
|
||||
assert await AgentRequestHandler.is_agent_allowed("agent-alpha", agent_key, resolve) is False
|
||||
|
||||
async def test_agent_without_access_groups_keeps_key_grants(self):
|
||||
agent_key: Final = self._key_granting(["agent-alpha"], agent_id="caller-agent")
|
||||
resolve, asked = self._ceiling_resolver(None)
|
||||
|
||||
assert await AgentRequestHandler.resolve_agent_access(agent_key, resolve) == RestrictedAgentAccess(
|
||||
frozenset({"agent-alpha"})
|
||||
)
|
||||
assert asked == ["caller-agent"]
|
||||
|
||||
async def test_key_without_agent_never_consults_agent_access_groups(self):
|
||||
plain_key: Final = UserAPIKeyAuth(api_key="test-key", user_id="test-user")
|
||||
resolve, asked = self._ceiling_resolver(frozenset())
|
||||
|
||||
with (
|
||||
patch.object(AgentRequestHandler, "_get_allowed_agents_for_key", return_value=UnrestrictedAgentAccess()),
|
||||
patch.object(AgentRequestHandler, "_get_allowed_agents_for_team", return_value=UnrestrictedAgentAccess()),
|
||||
patch(
|
||||
"litellm.proxy.agent_endpoints.auth.agent_access_groups.resolve_agent_access_group_ceiling",
|
||||
new=AsyncMock(return_value=self._ceiling(frozenset())),
|
||||
) as mock_ceiling,
|
||||
):
|
||||
assert await AgentRequestHandler.resolve_agent_access(plain_key) == UnrestrictedAgentAccess()
|
||||
mock_ceiling.assert_not_called()
|
||||
assert await AgentRequestHandler.resolve_agent_access(plain_key, resolve) == UnrestrictedAgentAccess()
|
||||
assert asked == []
|
||||
|
||||
async def test_empty_access_group_denies_every_agent(self):
|
||||
"""LIT-5143: a key restricted to an access group that resolves to no agents is
|
||||
|
|
|
|||
|
|
@ -32,11 +32,13 @@ from litellm.proxy._types import (
|
|||
UserAPIKeyAuth,
|
||||
WebhookEvent,
|
||||
)
|
||||
from litellm.proxy.agent_endpoints.auth.agent_access_groups import AgentAccessGroupCeiling, CeilingResolver
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
ExperimentalUIJWTToken,
|
||||
_cache_management_object,
|
||||
_can_object_call_model,
|
||||
_can_object_call_vector_stores,
|
||||
_check_agent_access_group_model_access,
|
||||
_check_end_user_budget,
|
||||
_check_team_member_budget,
|
||||
_fetch_key_object_from_db_with_reconnect,
|
||||
|
|
@ -8466,89 +8468,64 @@ def test_request_skips_budget_checks_extends_route_rule_with_zero_cost_models()
|
|||
# Agent access group model ceiling
|
||||
|
||||
|
||||
def _agent_model_ceiling(models: frozenset[str]):
|
||||
from litellm.proxy.agent_endpoints.auth.agent_access_groups import AgentAccessGroupCeiling
|
||||
def _agent_model_ceiling_resolver(
|
||||
models: frozenset[str] | None,
|
||||
) -> tuple[CeilingResolver, list[str]]:
|
||||
"""Resolver that records the agent ids it was asked about and answers with a fixed model
|
||||
ceiling, or None when the agent has no access groups attached."""
|
||||
asked: Final[list[str]] = []
|
||||
|
||||
return AgentAccessGroupCeiling(
|
||||
access_group_ids=("ag-1",), models=models, mcp_server_ids=frozenset(), agent_ids=frozenset()
|
||||
)
|
||||
async def resolve(agent_id: str) -> AgentAccessGroupCeiling | None:
|
||||
asked.append(agent_id)
|
||||
if models is None:
|
||||
return None
|
||||
return AgentAccessGroupCeiling(
|
||||
access_group_ids=("ag-1",), models=models, mcp_server_ids=frozenset(), agent_ids=frozenset()
|
||||
)
|
||||
|
||||
|
||||
async def _run_common_checks_for_agent_key(model: str, valid_token: UserAPIKeyAuth):
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy.auth.auth_checks import common_checks
|
||||
|
||||
return await common_checks(
|
||||
request_body={"model": model, "messages": [{"role": "user", "content": "hi"}]},
|
||||
team_object=None,
|
||||
user_object=None,
|
||||
end_user_object=None,
|
||||
global_proxy_spend=None,
|
||||
general_settings={},
|
||||
route="/chat/completions",
|
||||
llm_router=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
valid_token=valid_token,
|
||||
request=MagicMock(spec=Request),
|
||||
)
|
||||
return resolve, asked
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_common_checks_agent_access_groups_cap_models_even_when_key_allows_them():
|
||||
async def test_agent_access_groups_cap_models_even_when_key_allows_them():
|
||||
agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", models=["gpt-5", "claude-sonnet"])
|
||||
resolve, asked = _agent_model_ceiling_resolver(frozenset({"gpt-5"}))
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.auth.agent_access_groups.resolve_agent_access_group_ceiling",
|
||||
new=AsyncMock(return_value=_agent_model_ceiling(frozenset({"gpt-5"}))),
|
||||
):
|
||||
assert await _run_common_checks_for_agent_key("gpt-5", agent_key) is True
|
||||
assert await _check_agent_access_group_model_access("gpt-5", agent_key, None, resolve) is True
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await _run_common_checks_for_agent_key("claude-sonnet", agent_key)
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await _check_agent_access_group_model_access("claude-sonnet", agent_key, None, resolve)
|
||||
|
||||
assert exc_info.value.type == ProxyErrorTypes.agent_model_access_denied
|
||||
assert exc_info.value.code == str(status.HTTP_403_FORBIDDEN)
|
||||
assert asked == ["agent-1", "agent-1"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_common_checks_agent_access_groups_naming_no_model_deny_every_model():
|
||||
async def test_agent_access_groups_naming_no_model_deny_every_model():
|
||||
agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", models=[])
|
||||
resolve, _ = _agent_model_ceiling_resolver(frozenset())
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.agent_endpoints.auth.agent_access_groups.resolve_agent_access_group_ceiling",
|
||||
new=AsyncMock(return_value=_agent_model_ceiling(frozenset())),
|
||||
),
|
||||
pytest.raises(ProxyException) as exc_info,
|
||||
):
|
||||
await _run_common_checks_for_agent_key("gpt-5", agent_key)
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await _check_agent_access_group_model_access("gpt-5", agent_key, None, resolve)
|
||||
|
||||
assert exc_info.value.type == ProxyErrorTypes.agent_model_access_denied
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_common_checks_agent_without_access_groups_adds_no_model_ceiling():
|
||||
async def test_agent_without_access_groups_adds_no_model_ceiling():
|
||||
agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", models=["gpt-5", "claude-sonnet"])
|
||||
resolve, asked = _agent_model_ceiling_resolver(None)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.auth.agent_access_groups.resolve_agent_access_group_ceiling",
|
||||
new=AsyncMock(return_value=None),
|
||||
) as mock_ceiling:
|
||||
assert await _run_common_checks_for_agent_key("gpt-5", agent_key) is True
|
||||
assert await _run_common_checks_for_agent_key("claude-sonnet", agent_key) is True
|
||||
|
||||
mock_ceiling.assert_called_with("agent-1")
|
||||
assert await _check_agent_access_group_model_access("gpt-5", agent_key, None, resolve) is True
|
||||
assert await _check_agent_access_group_model_access("claude-sonnet", agent_key, None, resolve) is True
|
||||
assert asked == ["agent-1", "agent-1"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_common_checks_key_without_agent_never_consults_agent_access_groups():
|
||||
async def test_key_without_agent_never_consults_agent_access_groups():
|
||||
plain_key: Final = UserAPIKeyAuth(token="plain-token", models=["gpt-5"])
|
||||
resolve, asked = _agent_model_ceiling_resolver(frozenset())
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.auth.agent_access_groups.resolve_agent_access_group_ceiling",
|
||||
new=AsyncMock(return_value=_agent_model_ceiling(frozenset())),
|
||||
) as mock_ceiling:
|
||||
assert await _run_common_checks_for_agent_key("gpt-5", plain_key) is True
|
||||
|
||||
mock_ceiling.assert_not_called()
|
||||
assert await _check_agent_access_group_model_access("gpt-5", plain_key, None, resolve) is True
|
||||
assert asked == []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue