mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix: enforce authoritative managed agent permissions
This commit is contained in:
parent
18576ee4f2
commit
cab5213bf3
14 changed files with 696 additions and 28 deletions
|
|
@ -36,7 +36,9 @@ async def managed_agent_servers(auth: UserAPIKeyAuth) -> tuple[str, ...]:
|
|||
if context.user_id is None:
|
||||
return ()
|
||||
human: Final = await _delegated_resource_subject(context.user_id)
|
||||
allowed: Final = await MCPRequestHandler.resolve_admitted_subject_servers(human)
|
||||
allowed: Final = await MCPRequestHandler.resolve_admitted_subject_servers(
|
||||
human, allowed_team_ids=frozenset((auth.team_id,)) if auth.team_id else frozenset()
|
||||
)
|
||||
return tuple(sorted(own.intersection(allowed)))
|
||||
except Exception: # noqa: BLE001 # Authorization boundary: every unresolved policy must deny access
|
||||
raise_identity_failure(
|
||||
|
|
@ -57,7 +59,9 @@ async def managed_agent_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str]
|
|||
if context.user_id is None:
|
||||
return []
|
||||
human: Final = await _delegated_resource_subject(context.user_id)
|
||||
human_tools: Final = await MCPRequestHandler.resolve_admitted_subject_tools(server_id, human)
|
||||
human_tools: Final = await MCPRequestHandler.resolve_admitted_subject_tools(
|
||||
server_id, human, allowed_team_ids=frozenset((auth.team_id,)) if auth.team_id else frozenset()
|
||||
)
|
||||
if own is None:
|
||||
return human_tools
|
||||
return own if human_tools is None else sorted(frozenset(own).intersection(human_tools))
|
||||
|
|
|
|||
|
|
@ -1204,7 +1204,7 @@ class MCPRequestHandler:
|
|||
return None
|
||||
|
||||
@staticmethod
|
||||
async def _reload_admitted_key(key_hash: str) -> UserAPIKeyAuth:
|
||||
async def _reload_admitted_key(key_hash: str, *, check_db_only: bool = False) -> UserAPIKeyAuth:
|
||||
"""Reload the live key record an admitted envelope references and re-check live policy.
|
||||
|
||||
Resolving the current ``UserAPIKeyAuth`` (cache first, then DB) is what stops the
|
||||
|
|
@ -1236,6 +1236,7 @@ class MCPRequestHandler:
|
|||
hashed_token=key_hash,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
check_db_only=check_db_only,
|
||||
)
|
||||
except (ProxyException, HTTPException):
|
||||
raise HTTPException(status_code=401, detail="Invalid or expired credential") from None
|
||||
|
|
@ -1841,7 +1842,9 @@ class MCPRequestHandler:
|
|||
return scoped
|
||||
|
||||
@staticmethod
|
||||
async def admitted_subject_sources(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]:
|
||||
async def admitted_subject_sources(
|
||||
auth: UserAPIKeyAuth, *, allowed_team_ids: frozenset[str] | None = None
|
||||
) -> list[UserAPIKeyAuth]:
|
||||
"""The independent sources a keyless admitted subject reaches MCP servers through: their own
|
||||
direct grants, plus every team they are a live roster member of.
|
||||
|
||||
|
|
@ -1858,6 +1861,8 @@ class MCPRequestHandler:
|
|||
if not auth.user_id or prisma_client is None:
|
||||
return sources
|
||||
for team_id in await MCPRequestHandler._resolve_user_team_ids(auth.user_id, auth):
|
||||
if allowed_team_ids is not None and team_id not in allowed_team_ids:
|
||||
continue
|
||||
team_obj = await MCPRequestHandler._roster_team_object(team_id, auth)
|
||||
if team_obj is None:
|
||||
continue
|
||||
|
|
@ -1942,7 +1947,9 @@ class MCPRequestHandler:
|
|||
return team_obj
|
||||
|
||||
@staticmethod
|
||||
async def admitted_source_grants(auth: UserAPIKeyAuth) -> list[tuple[UserAPIKeyAuth, set[str]]]:
|
||||
async def admitted_source_grants(
|
||||
auth: UserAPIKeyAuth, *, allowed_team_ids: frozenset[str] | None = None
|
||||
) -> list[tuple[UserAPIKeyAuth, set[str]]]:
|
||||
"""``(source, the servers that source grants)`` for every source of an admitted subject.
|
||||
|
||||
THE owner of "which source reaches which server". The reachable union, the per-team throttle
|
||||
|
|
@ -1951,15 +1958,17 @@ class MCPRequestHandler:
|
|||
roster instead of by grant charged unrelated teams' buckets)."""
|
||||
return [
|
||||
(source, set(await MCPRequestHandler.get_allowed_mcp_servers(source, keyless_source=True)))
|
||||
for source in await MCPRequestHandler.admitted_subject_sources(auth)
|
||||
for source in await MCPRequestHandler.admitted_subject_sources(auth, allowed_team_ids=allowed_team_ids)
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
async def resolve_admitted_subject_servers(auth: UserAPIKeyAuth) -> list[str]:
|
||||
async def resolve_admitted_subject_servers(
|
||||
auth: UserAPIKeyAuth, *, allowed_team_ids: frozenset[str] | None = None
|
||||
) -> list[str]:
|
||||
"""Union of what each of the admitted subject's sources reaches, each answered by the
|
||||
canonical resolver so no rule is reimplemented for this caller shape."""
|
||||
reachable: Final[set[str]] = set()
|
||||
for _source, granted in await MCPRequestHandler.admitted_source_grants(auth):
|
||||
for _source, granted in await MCPRequestHandler.admitted_source_grants(auth, allowed_team_ids=allowed_team_ids):
|
||||
reachable.update(granted)
|
||||
return list(reachable)
|
||||
|
||||
|
|
@ -2017,7 +2026,9 @@ class MCPRequestHandler:
|
|||
return min((source for source, _ in granting), key=lambda s: s.team_id or "")
|
||||
|
||||
@staticmethod
|
||||
async def resolve_admitted_subject_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str] | None:
|
||||
async def resolve_admitted_subject_tools(
|
||||
server_id: str, auth: UserAPIKeyAuth, *, allowed_team_ids: frozenset[str] | None = None
|
||||
) -> list[str] | None:
|
||||
"""Effective tool allowlist on ``server_id`` for an admitted subject, as the union over the
|
||||
sources that actually grant that server.
|
||||
|
||||
|
|
@ -2039,7 +2050,7 @@ class MCPRequestHandler:
|
|||
) or await MCPRequestHandler.admin_view_unscoped(auth)
|
||||
|
||||
allowed: Final[set[str]] = set()
|
||||
for source, granted in await MCPRequestHandler.admitted_source_grants(auth):
|
||||
for source, granted in await MCPRequestHandler.admitted_source_grants(auth, allowed_team_ids=allowed_team_ids):
|
||||
# The open channel is evaluated against the user's OWN source (team_id is None), so that
|
||||
# source's restrictions apply to it; a team's rules never ride an open-channel server.
|
||||
if server_id not in granted and not (reachable_via_open_channel and source.team_id is None):
|
||||
|
|
@ -3464,7 +3475,8 @@ class MCPRequestHandler:
|
|||
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(
|
||||
obj_perm, requires_fresh_policy=user_api_key_auth.requires_fresh_policy
|
||||
)
|
||||
return list({*expanded_direct_servers, *access_group_servers, *toolset_grants})
|
||||
inline_tools: Final = global_mcp_server_manager.expand_tool_permissions(obj_perm.mcp_tool_permissions)
|
||||
return list({*expanded_direct_servers, *access_group_servers, *toolset_grants, *inline_tools})
|
||||
except Exception as e:
|
||||
if user_api_key_auth.managed_agent_policy is not None or isinstance(e, UnloadableEntitlementError):
|
||||
raise
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ from pydantic import (
|
|||
Json,
|
||||
JsonValue,
|
||||
PositiveInt,
|
||||
PrivateAttr,
|
||||
field_validator,
|
||||
model_validator,
|
||||
)
|
||||
|
|
@ -3334,6 +3335,7 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob
|
|||
invoked_agent_policy: AgentResponse | None = Field(default=None, exclude=True)
|
||||
agent_invocation_cost: float | None = Field(default=None, exclude=True)
|
||||
billing_agent_policy: AgentResponse | None = Field(default=None, exclude=True)
|
||||
_managed_delegation_verified: bool = PrivateAttr(default=False)
|
||||
managed_agent_policy: AgentResponse | None = Field(default=None, exclude=True)
|
||||
managed_agent_context: ManagedAgentContext | None = Field(default=None, exclude=True)
|
||||
agent_caller: AgentCaller | None = Field(
|
||||
|
|
|
|||
|
|
@ -175,7 +175,17 @@ class AgentRequestHandler:
|
|||
or user_api_key_auth is None
|
||||
):
|
||||
return False
|
||||
fresh_auth: Final = user_api_key_auth.model_copy(update={"requires_fresh_policy": True})
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
|
||||
|
||||
key_hash: Final = user_api_key_auth.api_key or user_api_key_auth.token
|
||||
authority: Final = (
|
||||
await MCPRequestHandler._reload_admitted_key(key_hash, check_db_only=True)
|
||||
if key_hash
|
||||
and user_api_key_auth.managed_agent_policy is None
|
||||
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})
|
||||
explicit: Final = await _granted_agent_ids(
|
||||
fresh_auth,
|
||||
_strict_agent_access,
|
||||
|
|
@ -643,16 +653,18 @@ async def _managed_actor_agent_access(auth: UserAPIKeyAuth) -> AgentAccess:
|
|||
return RestrictedAgentAccess(capped)
|
||||
if context.user_id is None:
|
||||
return RestrictedAgentAccess(frozenset())
|
||||
human_ids: Final = await verified_human_agent_grants(context.user_id)
|
||||
human_ids: Final = await verified_human_agent_grants(context.user_id, auth.team_id)
|
||||
return RestrictedAgentAccess(capped.intersection(human_ids))
|
||||
|
||||
|
||||
async def verified_human_agent_grants(user_id: str | None) -> frozenset[str]:
|
||||
async def verified_human_agent_grants(user_id: str | None, team_id: str | None = None) -> frozenset[str]:
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
|
||||
|
||||
if user_id is None:
|
||||
return frozenset()
|
||||
human: Final = await MCPRequestHandler.reload_admitted_user(user_id, requires_fresh_policy=True)
|
||||
sources: Final = await MCPRequestHandler.admitted_subject_sources(human)
|
||||
sources: Final = await MCPRequestHandler.admitted_subject_sources(
|
||||
human, allowed_team_ids=frozenset((team_id,)) if team_id else frozenset()
|
||||
)
|
||||
human_access: Final = await asyncio.gather(*(_strict_agent_access(source) for source in sources))
|
||||
return frozenset().union(*(_granted_ids(access) for access in human_access))
|
||||
|
|
|
|||
74
litellm/proxy/agent_endpoints/auth/managed_authorization.py
Normal file
74
litellm/proxy/agent_endpoints/auth/managed_authorization.py
Normal file
|
|
@ -0,0 +1,74 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
|
||||
from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityFailure, ManagedAgentContext
|
||||
|
||||
|
||||
async def admit_managed_actor(auth: UserAPIKeyAuth, store: AgentIdentityStore | None) -> None:
|
||||
delegation_verified: Final = auth._managed_delegation_verified
|
||||
auth._managed_delegation_verified = False
|
||||
if auth.agent_id is None:
|
||||
return
|
||||
if store is None:
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
|
||||
registered: Final = global_agent_registry.get_agent_by_id(auth.agent_id)
|
||||
if auth.managed_agent_context is not None or (
|
||||
registered is not None and (registered.identity_managed or registered.identity is not None)
|
||||
):
|
||||
raise_identity_failure(
|
||||
AgentIdentityFailure(code="policy_unavailable", message="Managed agent policy requires a database")
|
||||
)
|
||||
return
|
||||
agent: Final = await store.agent(auth.agent_id)
|
||||
if isinstance(agent, AgentIdentityFailure):
|
||||
raise_identity_failure(agent)
|
||||
if agent is None:
|
||||
retired: Final = await store.retired_agent(auth.agent_id)
|
||||
if isinstance(retired, AgentIdentityFailure):
|
||||
raise_identity_failure(retired)
|
||||
if auth.managed_agent_context is not None or retired:
|
||||
raise_identity_failure(AgentIdentityFailure(message="Agent no longer exists"))
|
||||
return
|
||||
if not agent.identity_managed:
|
||||
return
|
||||
if auth.jwt_claims and auth.managed_agent_context is None:
|
||||
raise_identity_failure(AgentIdentityFailure(message="A managed agent requires a matching verified identity"))
|
||||
failure: Final = actor_admission_failure(agent, auth.managed_agent_context)
|
||||
if failure is not None:
|
||||
raise_identity_failure(failure)
|
||||
auth.managed_agent_policy = agent
|
||||
auth.billing_agent_policy = agent
|
||||
auth.requires_fresh_policy = True
|
||||
if (
|
||||
auth.managed_agent_context is not None
|
||||
and auth.managed_agent_context.mode == "delegated"
|
||||
and not delegation_verified
|
||||
):
|
||||
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import verified_human_agent_grants
|
||||
|
||||
grants: Final = await verified_human_agent_grants(auth.managed_agent_context.user_id, auth.team_id)
|
||||
if agent.agent_id not in grants:
|
||||
raise_identity_failure(
|
||||
AgentIdentityFailure(message="The delegated user is not permitted to invoke this agent")
|
||||
)
|
||||
|
||||
|
||||
def actor_admission_failure(
|
||||
agent: AgentResponse,
|
||||
context: ManagedAgentContext | None,
|
||||
) -> AgentIdentityFailure | None:
|
||||
if not agent.enabled or agent.identity is None or not agent.identity.active:
|
||||
return AgentIdentityFailure(message="Agent execution is disabled")
|
||||
if context is None:
|
||||
return AgentIdentityFailure(message="This agent requires its bound identity provider token")
|
||||
if context.agent_id != agent.agent_id or context.binding_revision != agent.identity.revision:
|
||||
return AgentIdentityFailure(message="Agent identity changed during authentication; retry")
|
||||
if agent.execution_mode not in (context.mode, "both"):
|
||||
return AgentIdentityFailure(message="Agent is not enabled for this execution mode")
|
||||
if context.mode == "delegated" and not context.user_id:
|
||||
return AgentIdentityFailure(message="A verified human subject is required")
|
||||
return None
|
||||
|
|
@ -14,6 +14,7 @@ import math
|
|||
import re
|
||||
import time
|
||||
from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence
|
||||
from functools import partial
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Generic, Literal, Optional, Protocol, TypeAlias
|
||||
|
||||
|
|
@ -3738,6 +3739,8 @@ async def _fetch_key_object_from_db_with_reconnect(
|
|||
parent_otel_span: Span | None,
|
||||
proxy_logging_obj: ProxyLogging | None,
|
||||
deadline_seconds: float | None = None,
|
||||
*,
|
||||
check_db_only: bool = False,
|
||||
) -> BaseModel | None:
|
||||
"""
|
||||
Fetch key object from DB and retry once if a DB connection error can be healed.
|
||||
|
|
@ -3751,6 +3754,7 @@ async def _fetch_key_object_from_db_with_reconnect(
|
|||
prisma_client=prisma_client,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=check_db_only,
|
||||
),
|
||||
name="key",
|
||||
deadline_seconds=deadline_seconds,
|
||||
|
|
@ -3762,10 +3766,13 @@ async def _fetch_key_object_from_db_unbounded(
|
|||
prisma_client: PrismaClient,
|
||||
parent_otel_span: Span | None,
|
||||
proxy_logging_obj: ProxyLogging | None,
|
||||
*,
|
||||
check_db_only: bool = False,
|
||||
) -> BaseModel | None:
|
||||
fetch: Final = partial(prisma_client.get_data, use_writer=True) if check_db_only else prisma_client.get_data
|
||||
async with db_lookup_gate.current():
|
||||
try:
|
||||
return await prisma_client.get_data(
|
||||
return await fetch(
|
||||
token=hashed_token,
|
||||
table_name="combined_view",
|
||||
parent_otel_span=parent_otel_span,
|
||||
|
|
@ -3787,7 +3794,7 @@ async def _fetch_key_object_from_db_unbounded(
|
|||
lock_timeout_seconds=auth_reconnect_lock_timeout,
|
||||
)
|
||||
if did_reconnect:
|
||||
return await prisma_client.get_data(
|
||||
return await fetch(
|
||||
token=hashed_token,
|
||||
table_name="combined_view",
|
||||
parent_otel_span=parent_otel_span,
|
||||
|
|
@ -3875,6 +3882,8 @@ async def get_key_object(
|
|||
parent_otel_span: Span | None = None,
|
||||
proxy_logging_obj: ProxyLogging | None = None,
|
||||
check_cache_only: bool | None = None,
|
||||
*,
|
||||
check_db_only: bool = False,
|
||||
) -> UserAPIKeyAuth:
|
||||
"""
|
||||
- Check if team id in proxy Team Table
|
||||
|
|
@ -3889,9 +3898,8 @@ async def get_key_object(
|
|||
|
||||
# Same flow as before: use cache only when we have a hit we can turn into UserAPIKeyAuth
|
||||
# (dict from Redis / model_dump, or UserAPIKeyAuth from in-memory). Otherwise fall through to DB.
|
||||
user_api_key_auth: Final = await user_api_key_cache.async_get_cache(
|
||||
key=key,
|
||||
model_type=UserAPIKeyAuth,
|
||||
user_api_key_auth: Final = (
|
||||
None if check_db_only else await user_api_key_cache.async_get_cache(key=key, model_type=UserAPIKeyAuth)
|
||||
)
|
||||
if user_api_key_auth is not None:
|
||||
return _copy_user_api_key_auth_for_cache(user_api_key_obj=user_api_key_auth)
|
||||
|
|
@ -3905,6 +3913,7 @@ async def get_key_object(
|
|||
prisma_client=prisma_client,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=check_db_only,
|
||||
)
|
||||
|
||||
if _valid_token is None:
|
||||
|
|
@ -3918,7 +3927,7 @@ async def get_key_object(
|
|||
_response: Final = UserAPIKeyAuth.model_validate(_valid_token.model_dump(exclude_none=True))
|
||||
|
||||
# Load object_permission if object_permission_id exists but object_permission is not loaded
|
||||
if _response.object_permission_id and not _response.object_permission:
|
||||
if _response.object_permission_id and (check_db_only or not _response.object_permission):
|
||||
try:
|
||||
_response.object_permission = await get_object_permission(
|
||||
object_permission_id=_response.object_permission_id,
|
||||
|
|
@ -3926,14 +3935,20 @@ async def get_key_object(
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=check_db_only,
|
||||
)
|
||||
except Exception as e:
|
||||
if check_db_only:
|
||||
raise
|
||||
verbose_proxy_logger.debug(
|
||||
"Failed to load object_permission for key with object_permission_id=%s: %s",
|
||||
_response.object_permission_id,
|
||||
e,
|
||||
)
|
||||
|
||||
if check_db_only:
|
||||
return _response
|
||||
|
||||
# save the key object to cache
|
||||
await _cache_key_object(
|
||||
hashed_token=hashed_token,
|
||||
|
|
|
|||
|
|
@ -3204,6 +3204,15 @@ async def _authorize_authenticated_request(
|
|||
# admin-only-route / model-access / budget checks) surface as
|
||||
# ProxyException consistently with pre-refactor behavior.
|
||||
try:
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import admit_managed_actor
|
||||
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if user_api_key_auth_obj.agent_id is not None:
|
||||
await admit_managed_actor(
|
||||
user_api_key_auth_obj,
|
||||
AgentIdentityStore.from_client(prisma_client) if prisma_client is not None else None,
|
||||
)
|
||||
await _run_centralized_common_checks(
|
||||
user_api_key_auth_obj=user_api_key_auth_obj,
|
||||
request=request,
|
||||
|
|
|
|||
|
|
@ -4194,6 +4194,8 @@ _PRISMA_DEFAULT_TX_TIMEOUT: Final = timedelta(seconds=5)
|
|||
async def _lookup_deprecated_key(
|
||||
db: PrismaWrapper | RoutingPrismaWrapper,
|
||||
hashed_token: str,
|
||||
*,
|
||||
check_db_only: bool = False,
|
||||
) -> str | None:
|
||||
"""
|
||||
Check if a token exists in the deprecated keys table and is still within its grace period.
|
||||
|
|
@ -4205,7 +4207,7 @@ async def _lookup_deprecated_key(
|
|||
now_ts: Final = now.timestamp()
|
||||
|
||||
# Check cache first
|
||||
cached: Final = _deprecated_key_cache.get(hashed_token)
|
||||
cached: Final = None if check_db_only else _deprecated_key_cache.get(hashed_token)
|
||||
if cached is not None:
|
||||
active_token_id, cache_expires_at_ts, revoke_at_ts = cached
|
||||
if now_ts < cache_expires_at_ts and now_ts < revoke_at_ts:
|
||||
|
|
@ -4873,6 +4875,7 @@ class PrismaClient:
|
|||
proxy_logging_obj: ProxyLogging | None = None,
|
||||
budget_id_list: list[str] | None = None,
|
||||
check_deprecated: bool = True,
|
||||
use_writer: bool = False,
|
||||
):
|
||||
args_passed_in: Final = locals()
|
||||
start_time: Final = time.time()
|
||||
|
|
@ -5171,12 +5174,20 @@ class PrismaClient:
|
|||
WHERE v.token = $1
|
||||
"""
|
||||
|
||||
response = await self._query_first_with_cached_plan_fallback(sql_query, hashed_token)
|
||||
response = (
|
||||
await self.writer_db.query_first(sql_query, hashed_token)
|
||||
if use_writer
|
||||
else await self._query_first_with_cached_plan_fallback(sql_query, hashed_token)
|
||||
)
|
||||
|
||||
# If not found in main table, check deprecated keys (grace period)
|
||||
# check_deprecated=False on the recursive call prevents unbounded chaining
|
||||
if response is None and hashed_token is not None and check_deprecated:
|
||||
active_token_id: Final = await _lookup_deprecated_key(db=self.db, hashed_token=hashed_token)
|
||||
active_token_id: Final = await _lookup_deprecated_key(
|
||||
db=self.writer_db if use_writer else self.db,
|
||||
hashed_token=hashed_token,
|
||||
check_db_only=use_writer,
|
||||
)
|
||||
if active_token_id:
|
||||
# The recursive call returns a finished
|
||||
# LiteLLM_VerificationTokenView; the dict
|
||||
|
|
@ -5188,6 +5199,7 @@ class PrismaClient:
|
|||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_deprecated=False,
|
||||
use_writer=use_writer,
|
||||
)
|
||||
if deprecated_response is not None:
|
||||
verbose_proxy_logger.debug("Deprecated key used during grace period")
|
||||
|
|
|
|||
|
|
@ -275,6 +275,7 @@ async def test_delegated_mcp_uses_explicit_team_grants_even_for_dashboard_admins
|
|||
client.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=permission)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", client)
|
||||
auth: Final = actor(agent_tools, delegated=True)
|
||||
auth.team_id = "team"
|
||||
assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == (["slack"] if has_grant else [])
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == (["read"] if has_grant else [])
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == []
|
||||
|
|
@ -365,3 +366,83 @@ async def test_manager_does_not_replace_managed_policy_failure_with_open_servers
|
|||
with pytest.raises(HTTPException) as failure:
|
||||
await manager.get_allowed_mcp_servers(actor(None, delegated=True))
|
||||
assert failure.value.status_code == 503
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_inline_tool_grant_admits_its_server_without_widening_tools() -> None:
|
||||
auth: Final = actor(("read",))
|
||||
assert auth.managed_agent_policy is not None
|
||||
auth.managed_agent_policy = auth.managed_agent_policy.model_copy(
|
||||
update={"object_permission": {"object_permission_id": "tools", "mcp_tool_permissions": {"slack": ["read"]}}}
|
||||
)
|
||||
assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == ["slack"]
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == ["read"]
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("selected_team", (None, "selected"))
|
||||
@pytest.mark.parametrize("selected_grant", (False, True))
|
||||
async def test_delegation_never_borrows_another_teams_server_or_tools(
|
||||
monkeypatch: pytest.MonkeyPatch, selected_team: str | None, selected_grant: bool
|
||||
) -> None:
|
||||
from litellm.proxy._types import LiteLLM_TeamTable
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
|
||||
user: Final = LiteLLM_UserTable(user_id="human", teams=["selected", "other"], organization_memberships=[])
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="selected-grant",
|
||||
mcp_servers=["slack"] if selected_grant else [],
|
||||
mcp_tool_permissions={"slack": ["read"]} if selected_grant else {},
|
||||
)
|
||||
teams: Final = {
|
||||
name: LiteLLM_TeamTable(
|
||||
team_id=name,
|
||||
models=[],
|
||||
members_with_roles=[{"user_id": "human", "role": "user"}],
|
||||
object_permission=permission if name == "selected" else LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="other-grant", mcp_servers=["slack", "linear"]
|
||||
),
|
||||
)
|
||||
for name in ("selected", "other")
|
||||
}
|
||||
|
||||
async def get_team(team_id: str, **kwargs: object) -> LiteLLM_TeamTable:
|
||||
return teams[team_id]
|
||||
|
||||
monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(return_value=user))
|
||||
monkeypatch.setattr(auth_checks, "get_team_object", get_team)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", UserApiKeyCache())
|
||||
auth: Final = actor(None, delegated=True)
|
||||
auth.team_id = selected_team
|
||||
expected: Final = ["slack"] if selected_team and selected_grant else []
|
||||
assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == expected
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == (["read"] if expected else [])
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == []
|
||||
ordinary: Final = await MCPRequestHandler.reload_admitted_user("human", requires_fresh_policy=True)
|
||||
assert set(await MCPRequestHandler.resolve_admitted_subject_servers(ordinary)) == {"slack", "linear"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("entitlement", ("group", "toolset"))
|
||||
async def test_managed_mcp_rejects_unavailable_authoritative_entitlements(
|
||||
monkeypatch: pytest.MonkeyPatch, entitlement: str
|
||||
) -> None:
|
||||
client: Final = MagicMock()
|
||||
client.writer_db.litellm_mcpservertable.find_many = AsyncMock(side_effect=RuntimeError("writer unavailable"))
|
||||
client.writer_db.litellm_mcptoolsettable.find_many = AsyncMock(side_effect=RuntimeError("writer unavailable"))
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", client)
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="entitlements",
|
||||
mcp_access_groups=["group"] if entitlement == "group" else [],
|
||||
mcp_toolsets=["toolset"] if entitlement == "toolset" else [],
|
||||
)
|
||||
auth: Final = actor(None)
|
||||
assert auth.managed_agent_policy is not None
|
||||
auth.managed_agent_policy = auth.managed_agent_policy.model_copy(update={"object_permission": permission.model_dump()})
|
||||
auth.requires_fresh_policy = True
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await MCPRequestHandler.get_allowed_tools_for_server("slack", auth)
|
||||
assert failure.value.status_code == 503
|
||||
client.db.litellm_mcpservertable.find_many.assert_not_called()
|
||||
client.db.litellm_mcptoolsettable.find_many.assert_not_called()
|
||||
|
|
|
|||
|
|
@ -680,7 +680,7 @@ async def test_managed_agent_invocation_grants_intersect_verified_user_grants(
|
|||
from litellm.proxy._types import LiteLLM_UserTable
|
||||
from litellm.proxy.auth import auth_checks
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.agent_identity import ManagedAgentContext
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityBinding, ManagedAgentContext
|
||||
|
||||
database: Final = MagicMock()
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", database)
|
||||
|
|
@ -688,7 +688,7 @@ async def test_managed_agent_invocation_grants_intersect_verified_user_grants(
|
|||
human_grants: Final = LiteLLM_ObjectPermissionTable(object_permission_id="human", agents=["shared", "human-only"])
|
||||
human: Final = LiteLLM_UserTable(user_id="human", teams=[], object_permission=human_grants)
|
||||
monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(return_value=human))
|
||||
auth: Final = UserAPIKeyAuth(agent_id="actor")
|
||||
auth: Final = UserAPIKeyAuth(agent_id="actor", api_key="verified-jwt")
|
||||
auth.managed_agent_policy = AgentResponse(
|
||||
agent_id="actor", agent_name="Actor", agent_card_params={}, object_permission=own.model_dump()
|
||||
)
|
||||
|
|
@ -698,6 +698,15 @@ async def test_managed_agent_invocation_grants_intersect_verified_user_grants(
|
|||
access: Final = await AgentRequestHandler.resolve_agent_access(auth)
|
||||
assert access == RestrictedAgentAccess(frozenset({"shared"} if delegated else {"shared", "agent-only"}))
|
||||
|
||||
target: Final = AgentResponse(
|
||||
agent_id="shared", agent_name="Shared", agent_card_params={}, identity_managed=True,
|
||||
identity=AgentIdentityBinding(
|
||||
agent_id="shared", provider="microsoft_entra", tenant_id="tenant", client_id="client", issuer="issuer", revision="current"
|
||||
),
|
||||
)
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target)
|
||||
assert await AgentRequestHandler.is_agent_allowed("shared", auth) is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("revoked", ["user", "team-member", "team-grant", "team-permission", "direct-grant", "access-group"])
|
||||
|
|
@ -742,7 +751,7 @@ async def test_delegated_grants_revoke_with_warm_user_team_and_permission_caches
|
|||
client.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(return_value=group)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", client)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
|
||||
assert await verified_human_agent_grants("human") == frozenset({"target"})
|
||||
assert await verified_human_agent_grants("human", "team") == frozenset({"target"})
|
||||
client.writer_db.litellm_usertable.find_unique.return_value = (
|
||||
human.model_copy(update={"teams": []}) if revoked == "user" else human
|
||||
)
|
||||
|
|
@ -759,7 +768,7 @@ async def test_delegated_grants_revoke_with_warm_user_team_and_permission_caches
|
|||
client.writer_db.litellm_accessgrouptable.find_unique.return_value = (
|
||||
group.model_copy(update={"access_agent_ids": []}) if grouped else group
|
||||
)
|
||||
assert await verified_human_agent_grants("human") == frozenset()
|
||||
assert await verified_human_agent_grants("human", "team") == frozenset()
|
||||
client.db.litellm_usertable.find_unique.assert_not_called()
|
||||
client.db.litellm_teamtable.find_unique.assert_not_called()
|
||||
client.db.litellm_objectpermissiontable.find_unique.assert_not_called()
|
||||
|
|
@ -885,3 +894,62 @@ async def test_delegation_without_a_verified_human_never_grants_agents(grant: bo
|
|||
auth.managed_agent_context = ManagedAgentContext(agent_id="actor", mode="delegated")
|
||||
assert await AgentRequestHandler.resolve_agent_access(auth) == RestrictedAgentAccess(frozenset())
|
||||
assert await verified_human_agent_grants(None) == frozenset()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("change", ("grant", "permission_reference", "groups", "team", "blocked", "expired", "deleted", "outage"))
|
||||
async def test_managed_target_rechecks_authoritative_key_after_peer_revocation(
|
||||
monkeypatch: pytest.MonkeyPatch, change: str
|
||||
) -> None:
|
||||
from unittest.mock import MagicMock
|
||||
from fastapi import HTTPException
|
||||
from litellm.proxy import proxy_server
|
||||
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
|
||||
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="grant", agents=["target"])
|
||||
warm: Final = UserAPIKeyAuth(api_key="a" * 64, token="a" * 64, object_permission_id="grant", object_permission=permission)
|
||||
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"),
|
||||
)
|
||||
client: Final = MagicMock()
|
||||
client.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target)
|
||||
client.get_data = AsyncMock(return_value=warm)
|
||||
client.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=permission)
|
||||
cache: Final = UserApiKeyCache()
|
||||
cache.set_cache("a" * 64, warm)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", client)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
|
||||
assert await AgentRequestHandler.is_agent_allowed("target", warm) is True
|
||||
client.get_data.return_value = warm.model_copy(update={
|
||||
"object_permission": None,
|
||||
"object_permission_id": "replacement" if change == "permission_reference" else "grant",
|
||||
"access_group_ids": [],
|
||||
"team_id": "new-team" if change == "team" else None,
|
||||
"blocked": change == "blocked",
|
||||
"expires": "2000-01-01T00:00:00+00:00" if change == "expired" else None,
|
||||
})
|
||||
client.writer_db.litellm_objectpermissiontable.find_unique.return_value = permission.model_copy(update={"agents": []})
|
||||
if change == "team":
|
||||
from litellm.proxy._types import LiteLLM_TeamTable
|
||||
from litellm.proxy.auth import auth_checks
|
||||
client.writer_db.litellm_objectpermissiontable.find_unique.return_value = permission
|
||||
monkeypatch.setattr(auth_checks, "get_team_object", AsyncMock(return_value=LiteLLM_TeamTable(
|
||||
team_id="new-team", object_permission=permission.model_copy(update={"agents": ["other"]})
|
||||
)))
|
||||
if change == "groups":
|
||||
warm.object_permission = None
|
||||
warm.access_group_ids = ["old-group"]
|
||||
from litellm.proxy.auth import auth_checks
|
||||
monkeypatch.setattr(auth_checks, "_get_agent_ids_from_access_groups", AsyncMock(return_value=["target"]))
|
||||
if change == "deleted":
|
||||
client.get_data.return_value = None
|
||||
if change == "outage":
|
||||
client.get_data.side_effect = RuntimeError("writer unavailable")
|
||||
if change in ("blocked", "expired", "deleted", "outage"):
|
||||
with pytest.raises((HTTPException, RuntimeError)):
|
||||
await AgentRequestHandler.is_agent_allowed("target", warm)
|
||||
else:
|
||||
assert await AgentRequestHandler.is_agent_allowed("target", warm) is False
|
||||
|
|
|
|||
|
|
@ -0,0 +1,285 @@
|
|||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import (
|
||||
actor_admission_failure,
|
||||
admit_managed_actor,
|
||||
)
|
||||
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityBinding, AgentIdentityFailure, ManagedAgentContext
|
||||
|
||||
BINDING: Final = AgentIdentityBinding(
|
||||
agent_id="agent",
|
||||
provider="microsoft_entra",
|
||||
tenant_id="tenant",
|
||||
client_id="client",
|
||||
service_principal_id="principal",
|
||||
issuer="issuer",
|
||||
revision="current",
|
||||
)
|
||||
|
||||
|
||||
def agent(**overrides: object) -> AgentResponse:
|
||||
return AgentResponse.model_validate(
|
||||
{
|
||||
"agent_id": "agent",
|
||||
"agent_name": "Agent",
|
||||
"agent_card_params": {},
|
||||
"identity": BINDING,
|
||||
"identity_managed": True,
|
||||
"execution_mode": "both",
|
||||
**overrides,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"state",
|
||||
[
|
||||
{"enabled": False},
|
||||
{"identity": None},
|
||||
{"identity": BINDING.model_copy(update={"active": False})},
|
||||
{"execution_mode": "delegated"},
|
||||
],
|
||||
)
|
||||
def test_keys_cannot_bypass_lifecycle_or_delegated_only_mode(state: dict[str, object]) -> None:
|
||||
assert isinstance(actor_admission_failure(agent(**state), None), AgentIdentityFailure)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("mode", ["autonomous", "both", "delegated"])
|
||||
def test_keys_cannot_impersonate_an_entra_bound_agent(mode: str) -> None:
|
||||
assert isinstance(actor_admission_failure(agent(execution_mode=mode), None), AgentIdentityFailure)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"context",
|
||||
[
|
||||
ManagedAgentContext(agent_id="agent", binding_revision="previous", mode="autonomous"),
|
||||
ManagedAgentContext(agent_id="another", binding_revision="current", mode="autonomous"),
|
||||
ManagedAgentContext(agent_id="agent", binding_revision="current", mode="delegated"),
|
||||
],
|
||||
)
|
||||
def test_stale_binding_and_unverified_delegation_cannot_pass_admission(context: ManagedAgentContext) -> None:
|
||||
assert isinstance(actor_admission_failure(agent(), context), AgentIdentityFailure)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deleted_agent_key_cannot_fall_back_to_unmanaged_authentication() -> None:
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=None)
|
||||
database.writer_db.litellm_retiredagent.find_unique = AsyncMock(return_value={"original_agent_id": "deleted"})
|
||||
with pytest.raises(HTTPException, match="Agent no longer exists"):
|
||||
await admit_managed_actor(UserAPIKeyAuth(agent_id="deleted"), AgentIdentityStore.from_client(database))
|
||||
database.writer_db.litellm_retiredagent.find_unique.return_value = None
|
||||
auth: Final = UserAPIKeyAuth(agent_id="legacy-attribution-label")
|
||||
await admit_managed_actor(auth, AgentIdentityStore.from_client(database))
|
||||
assert auth.managed_agent_policy is None
|
||||
database.db.litellm_agentstable.find_unique.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_history_outage_does_not_permit_legacy_fallback() -> None:
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=None)
|
||||
database.writer_db.litellm_retiredagent.find_unique = AsyncMock(side_effect=RuntimeError("unavailable"))
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await admit_managed_actor(UserAPIKeyAuth(agent_id="deleted"), AgentIdentityStore.from_client(database))
|
||||
assert failure.value.status_code == 503
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_admission_database_outage_fails_closed() -> None:
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(side_effect=RuntimeError("DB unavailable"))
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await admit_managed_actor(UserAPIKeyAuth(agent_id="agent"), AgentIdentityStore.from_client(database))
|
||||
assert failure.value.status_code == 503
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_human_authentication_does_not_load_an_agent() -> None:
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock()
|
||||
await admit_managed_actor(UserAPIKeyAuth(user_id="human"), AgentIdentityStore.from_client(database))
|
||||
database.writer_db.litellm_agentstable.find_unique.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disabled_agent_key_is_rejected_at_admission() -> None:
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent(enabled=False))
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await admit_managed_actor(UserAPIKeyAuth(agent_id="agent"), AgentIdentityStore.from_client(database))
|
||||
assert failure.value.status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("permitted", [True, False])
|
||||
async def test_verified_human_still_needs_an_explicit_agent_invocation_grant(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
permitted: bool,
|
||||
) -> None:
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_UserTable
|
||||
from litellm.proxy.auth import auth_checks
|
||||
|
||||
policy: Final = agent()
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", database)
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="human-grants",
|
||||
agents=["agent"] if permitted else [],
|
||||
)
|
||||
human: Final = LiteLLM_UserTable(user_id="human", teams=[], object_permission=permission)
|
||||
monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(return_value=human))
|
||||
auth: Final = UserAPIKeyAuth(agent_id="agent")
|
||||
auth.managed_agent_context = ManagedAgentContext(
|
||||
agent_id="agent",
|
||||
binding_revision="current",
|
||||
mode="delegated",
|
||||
user_id="human",
|
||||
)
|
||||
if permitted:
|
||||
await admit_managed_actor(auth, AgentIdentityStore.from_client(database))
|
||||
assert auth.managed_agent_policy == policy
|
||||
assert auth.billing_agent_policy == policy
|
||||
else:
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await admit_managed_actor(auth, AgentIdentityStore.from_client(database))
|
||||
assert failure.value.status_code == 403
|
||||
|
||||
|
||||
def test_execution_mode_must_match_verified_token_mode() -> None:
|
||||
context: Final = ManagedAgentContext(agent_id="agent", binding_revision="current", mode="autonomous")
|
||||
failure: Final = actor_admission_failure(agent(execution_mode="delegated"), context)
|
||||
assert isinstance(failure, AgentIdentityFailure)
|
||||
assert "execution mode" in failure.message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_legacy_jwt_cannot_adopt_an_agent_bound_on_another_worker() -> None:
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent(execution_mode="autonomous"))
|
||||
auth: Final = UserAPIKeyAuth(agent_id="agent", jwt_claims={"agent": "agent", "sub": "unrelated-subject"})
|
||||
with pytest.raises(HTTPException) as denied:
|
||||
await admit_managed_actor(auth, AgentIdentityStore.from_client(database))
|
||||
assert denied.value.status_code == 403
|
||||
assert auth.managed_agent_policy is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("bound", [False, True])
|
||||
async def test_managed_context_or_binding_requires_database(monkeypatch: pytest.MonkeyPatch, bound: bool) -> None:
|
||||
from litellm.proxy.agent_endpoints import agent_registry
|
||||
from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry
|
||||
|
||||
registry: Final = AgentRegistry()
|
||||
registry.register_agent(agent(identity_managed=bound, identity=BINDING if bound else None))
|
||||
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
|
||||
auth: Final = UserAPIKeyAuth(agent_id="agent")
|
||||
if not bound:
|
||||
auth.managed_agent_context = ManagedAgentContext(
|
||||
agent_id="agent", binding_revision="current", mode="autonomous"
|
||||
)
|
||||
with pytest.raises(HTTPException) as denied:
|
||||
await admit_managed_actor(auth, None)
|
||||
assert denied.value.status_code == 503
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_autonomous_app_rejects_persisted_virtual_key_impersonation() -> None:
|
||||
policy: Final = agent(execution_mode="autonomous")
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy)
|
||||
auth: Final = UserAPIKeyAuth(agent_id="agent", api_key="persisted-key")
|
||||
with pytest.raises(HTTPException, match="bound identity provider token") as denied:
|
||||
await admit_managed_actor(auth, AgentIdentityStore.from_client(database))
|
||||
assert denied.value.status_code == 403
|
||||
assert auth.managed_agent_policy is None
|
||||
assert auth.billing_agent_policy is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize("mode,user", [("autonomous", None), ("delegated", "verified-human")])
|
||||
def test_matching_identity_revision_and_execution_mode_pass_admission(mode: str, user: str | None) -> None:
|
||||
context: Final = ManagedAgentContext.model_validate(
|
||||
{"agent_id": "agent", "binding_revision": "current", "mode": mode, "user_id": user}
|
||||
)
|
||||
assert actor_admission_failure(agent(), context) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bound_autonomous_actor_is_admitted_without_a_human() -> None:
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent())
|
||||
auth: Final = UserAPIKeyAuth(agent_id="agent")
|
||||
auth.managed_agent_context = ManagedAgentContext(agent_id="agent", binding_revision="current", mode="autonomous")
|
||||
await admit_managed_actor(auth, AgentIdentityStore.from_client(database))
|
||||
assert auth.managed_agent_policy == agent()
|
||||
assert auth.billing_agent_policy == agent()
|
||||
assert auth.user_id is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_admitted_managed_actor_requires_fresh_policy_so_revocations_bind_next_request() -> None:
|
||||
"""Managed MCP grants (toolsets, access groups) are read through the shared resolvers, which only
|
||||
bypass the warm cache and the replica when the subject carries requires_fresh_policy"""
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=agent())
|
||||
auth: Final = UserAPIKeyAuth(agent_id="agent")
|
||||
auth.managed_agent_context = ManagedAgentContext(agent_id="agent", binding_revision="current", mode="autonomous")
|
||||
assert auth.requires_fresh_policy is False
|
||||
await admit_managed_actor(auth, AgentIdentityStore.from_client(database))
|
||||
assert auth.requires_fresh_policy is True
|
||||
|
||||
|
||||
async def test_jwt_delegation_verification_is_consumed_once_and_cannot_be_supplied_by_a_caller(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from litellm.proxy.agent_endpoints.auth import agent_permission_handler
|
||||
|
||||
policy: Final = agent()
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=policy)
|
||||
store: Final = AgentIdentityStore.from_client(database)
|
||||
grants: Final = AsyncMock(return_value=frozenset())
|
||||
monkeypatch.setattr(agent_permission_handler, "verified_human_agent_grants", grants)
|
||||
auth: Final = UserAPIKeyAuth.model_validate({"agent_id": "agent", "_managed_delegation_verified": True})
|
||||
assert auth._managed_delegation_verified is False
|
||||
auth.managed_agent_context = ManagedAgentContext(
|
||||
agent_id="agent", binding_revision="current", mode="delegated", user_id="human"
|
||||
)
|
||||
auth._managed_delegation_verified = True
|
||||
assert "_managed_delegation_verified" not in auth.model_dump()
|
||||
await admit_managed_actor(auth, store)
|
||||
grants.assert_not_awaited()
|
||||
assert auth._managed_delegation_verified is False
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await admit_managed_actor(auth, store)
|
||||
assert failure.value.status_code == 403
|
||||
grants.assert_awaited_once_with("human", None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("database_available", (False, True))
|
||||
async def test_ordinary_agent_admission_preserves_legacy_authentication(
|
||||
monkeypatch: pytest.MonkeyPatch, database_available: bool
|
||||
) -> None:
|
||||
from litellm.proxy.agent_endpoints import agent_registry
|
||||
|
||||
registry: Final = agent_registry.AgentRegistry()
|
||||
ordinary: Final = agent(identity_managed=False, identity=None)
|
||||
registry.register_agent(ordinary)
|
||||
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=ordinary)
|
||||
auth: Final = UserAPIKeyAuth(agent_id="agent")
|
||||
await admit_managed_actor(auth, AgentIdentityStore.from_client(database) if database_available else None)
|
||||
assert auth.agent_id == "agent"
|
||||
assert auth.managed_agent_policy is None
|
||||
assert auth.requires_fresh_policy is False
|
||||
|
|
@ -10160,3 +10160,43 @@ async def test_managed_agent_model_policy_checks_dispatched_model(
|
|||
with pytest.raises((HTTPException, ModelAccessDeniedProxyException)) as failure:
|
||||
await checks
|
||||
assert str(getattr(failure.value, "status_code", getattr(failure.value, "code", None))) == "403"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("reconnect", (False, True))
|
||||
async def test_authoritative_key_load_bypasses_warm_key_and_permission_caches(reconnect: bool) -> None:
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key
|
||||
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="current", agents=["allowed"])
|
||||
stale: Final = UserAPIKeyAuth(token="hash", team_id="old-team", object_permission_id="old")
|
||||
current: Final = UserAPIKeyAuth(token="hash", team_id="new-team", object_permission_id="current")
|
||||
cache: Final = UserApiKeyCache()
|
||||
cache.set_cache("hash", stale)
|
||||
cache.set_cache(object_permission_cache_key("current"), permission.model_copy(update={"agents": ["revoked"]}))
|
||||
database: Final = MagicMock()
|
||||
database.get_data = AsyncMock(side_effect=[httpx.ConnectError("reset"), current] if reconnect else [current])
|
||||
database.attempt_db_reconnect = AsyncMock(return_value=True)
|
||||
database.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=permission)
|
||||
fresh: Final = await get_key_object("hash", database, cache, check_db_only=True)
|
||||
assert fresh.team_id == "new-team"
|
||||
assert fresh.object_permission == permission
|
||||
assert all(call.kwargs["use_writer"] is True for call in database.get_data.await_args_list)
|
||||
database.db.litellm_objectpermissiontable.find_unique.assert_not_called()
|
||||
cached: Final = await get_key_object("hash", database, cache)
|
||||
assert cached.team_id == "old-team"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("missing", (False, True))
|
||||
async def test_authoritative_key_cannot_keep_grants_when_permission_is_unavailable(missing: bool) -> None:
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
|
||||
database: Final = MagicMock()
|
||||
database.get_data = AsyncMock(return_value=UserAPIKeyAuth(
|
||||
object_permission_id="grant", object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="grant", agents=["allowed"])
|
||||
))
|
||||
database.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(
|
||||
return_value=None, side_effect=None if missing else RuntimeError("writer unavailable")
|
||||
)
|
||||
with pytest.raises(Exception, match=r"does not exist|unavailable"):
|
||||
await get_key_object("hash", database, UserApiKeyCache(), check_db_only=True)
|
||||
|
|
|
|||
|
|
@ -9381,3 +9381,32 @@ def test_identity_prefetch_keys_match_what_auth_reads_for_the_request():
|
|||
assert _identity_cache_keys("sk-1234", end_user_id=None, key_is_resolved=True) == (
|
||||
model_access_group_registry_cache_key(),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_virtual_key_cannot_enter_checks_as_an_identity_managed_actor(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
from typing import Final
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.auth import user_api_key_auth as auth_module
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityBinding
|
||||
|
||||
target: Final = AgentResponse(
|
||||
agent_id="bound", agent_name="Bound", agent_card_params={}, identity_managed=True,
|
||||
identity=AgentIdentityBinding(
|
||||
agent_id="bound", provider="microsoft_entra", tenant_id="tenant", client_id="client", issuer="issuer", revision="current"
|
||||
),
|
||||
)
|
||||
client: Final = MagicMock()
|
||||
client.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", client)
|
||||
checks: Final = AsyncMock()
|
||||
monkeypatch.setattr(auth_module, "_run_centralized_common_checks", checks)
|
||||
monkeypatch.setattr(proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock(return_value=None)))
|
||||
data: Final = {"model": "allowed", "messages": [{"role": "user", "content": "hello"}]}
|
||||
request: Final = _alias_request("/v1/chat/completions", data)
|
||||
with pytest.raises(ProxyException):
|
||||
await auth_module._authorize_authenticated_request(
|
||||
UserAPIKeyAuth(agent_id="bound"), request, data, "/v1/chat/completions", "sk-test"
|
||||
)
|
||||
checks.assert_not_awaited()
|
||||
|
|
|
|||
|
|
@ -634,3 +634,28 @@ async def test_query_first_with_cached_plan_fallback_reports_the_reader_generati
|
|||
"reader_served_the_query": 2,
|
||||
"writer_served_the_query": 0,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("rotated", (False, True))
|
||||
async def test_authoritative_combined_key_view_uses_writer_through_rotation(
|
||||
prisma_client: PrismaClient, rotated: bool
|
||||
) -> None:
|
||||
writer: Final = MagicMock()
|
||||
reader: Final = MagicMock()
|
||||
active: Final = {
|
||||
"token": "current-token", "team_id": "current-team", "team_models": None,
|
||||
"team_blocked": None, "team_members_with_roles": None, "user_id": None, "expires": None,
|
||||
}
|
||||
writer.query_first = AsyncMock(side_effect=[None, active] if rotated else [active])
|
||||
reader.query_first = AsyncMock(return_value={**active, "team_id": "stale-team"})
|
||||
writer.litellm_deprecatedverificationtoken.find_first = AsyncMock(return_value=SimpleNamespace(
|
||||
active_token_id="current-token", revoke_at=datetime.now(timezone.utc) + timedelta(hours=1)
|
||||
))
|
||||
prisma_client.db = RoutingPrismaWrapper(writer=writer, reader=reader)
|
||||
response: Final = await prisma_client.get_data(token="original-token", table_name="combined_view", use_writer=True)
|
||||
assert isinstance(response, LiteLLM_VerificationTokenView)
|
||||
assert response.team_id == "current-team"
|
||||
assert response.token == "current-token"
|
||||
reader.query_first.assert_not_awaited()
|
||||
assert writer.query_first.await_count == (2 if rotated else 1)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue