mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
chore(roi): sync main and remove unrelated changes
This commit is contained in:
commit
7eba31c719
85 changed files with 4107 additions and 623 deletions
|
|
@ -53,6 +53,10 @@ from litellm.llms.base_llm.guardrail_translation.base_translation import (
|
|||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import (
|
||||
blocked_responses_stream_usage,
|
||||
effective_skip_system_message_for_guardrail,
|
||||
merge_guardrailed_scoped_messages,
|
||||
role_out_of_guardrail_scope,
|
||||
scoped_structured_message_indices,
|
||||
stream_item_field,
|
||||
stream_item_fingerprint,
|
||||
stream_item_items,
|
||||
|
|
@ -376,6 +380,17 @@ class _RequestFields(NamedTuple):
|
|||
class _ExtractedInputs(NamedTuple):
|
||||
inputs: GenericGuardrailAPIInputs
|
||||
task_mappings: tuple[tuple[int, int | None], ...]
|
||||
instructions: str | None
|
||||
|
||||
|
||||
def scannable_instructions(data: Mapping[str, object], *, skip_system: bool = False) -> str | None:
|
||||
instructions: Final = data.get("instructions")
|
||||
return instructions if isinstance(instructions, str) and instructions and not skip_system else None
|
||||
|
||||
|
||||
def _input_item_role(item: object) -> str:
|
||||
role: Final = item.get("role") if isinstance(item, Mapping) else None
|
||||
return role.lower() if isinstance(role, str) else ""
|
||||
|
||||
|
||||
def _patched_request_fields(
|
||||
|
|
@ -494,7 +509,14 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
input_data: Final[str | ResponseInputParam | None] = data.get("input")
|
||||
if not isinstance(input_data, (str, list)):
|
||||
return data
|
||||
skip_system: Final = effective_skip_system_message_for_guardrail(guardrail_to_apply)
|
||||
structured_messages: Final = self.get_structured_messages(data)
|
||||
scoped_indices: Final = scoped_structured_message_indices(
|
||||
structured_messages or [], scan_only_tool_results=False, skip_system=skip_system, skip_tool=False
|
||||
)
|
||||
scoped_structured_messages: Final = (
|
||||
[structured_messages[index] for index in scoped_indices] if structured_messages else None
|
||||
)
|
||||
raw_tools: Final = data.get("tools")
|
||||
original_tools: Final[tuple[Mapping[str, object], ...]] = (
|
||||
tuple(raw_tools) if isinstance(raw_tools, list) else ()
|
||||
|
|
@ -502,11 +524,13 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
flattened_tool_groups: Final = tuple(
|
||||
form.chat_tools for form in LiteLLMCompletionResponsesConfig.responses_tools_to_chat_forms(original_tools)
|
||||
)
|
||||
extracted: Final = self._extract_guardrail_inputs(data, input_data, flattened_tool_groups)
|
||||
extracted: Final = self._extract_guardrail_inputs(
|
||||
data, input_data, flattened_tool_groups, skip_system=skip_system
|
||||
)
|
||||
if not extracted.inputs.get("texts"):
|
||||
return data
|
||||
if structured_messages:
|
||||
extracted.inputs["structured_messages"] = structured_messages
|
||||
if scoped_structured_messages:
|
||||
extracted.inputs["structured_messages"] = scoped_structured_messages
|
||||
guardrailed_inputs: Final = await guardrail_to_apply.apply_guardrail(
|
||||
inputs=extracted.inputs,
|
||||
request_data=data,
|
||||
|
|
@ -516,37 +540,63 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
self._apply_guardrailed_tools_to_data(
|
||||
data, original_tools, flattened_tool_groups, guardrailed_inputs.get("tools")
|
||||
)
|
||||
written_back: Final = self._written_back_request_fields(data, structured_messages, guardrailed_inputs)
|
||||
written_back: Final = self._written_back_request_fields(
|
||||
data,
|
||||
structured_messages or (),
|
||||
scoped_indices,
|
||||
scoped_structured_messages,
|
||||
guardrail_to_apply,
|
||||
guardrailed_inputs,
|
||||
)
|
||||
if written_back is not None:
|
||||
data["input"] = list(written_back.input) # mutable-ok: JSON body
|
||||
if written_back.instructions is None:
|
||||
data.pop("instructions", None)
|
||||
else:
|
||||
data["instructions"] = written_back.instructions # rebind-ok: data is an out-param
|
||||
elif isinstance(input_data, str):
|
||||
guardrailed_texts: Final = guardrailed_inputs.get("texts") or ()
|
||||
if len(guardrailed_texts) > 1:
|
||||
raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name)
|
||||
data["input"] = guardrailed_texts[0] if guardrailed_texts else input_data # rebind-ok: data is an out-param
|
||||
else:
|
||||
rewritten_texts: Final = guardrailed_inputs.get("texts") or ()
|
||||
if len(rewritten_texts) != len(extracted.task_mappings):
|
||||
raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name)
|
||||
await self._apply_guardrail_responses_to_input(
|
||||
messages=input_data,
|
||||
responses=rewritten_texts,
|
||||
task_mappings=extracted.task_mappings,
|
||||
)
|
||||
await self._apply_guardrailed_texts(data, input_data, extracted, guardrail_to_apply, guardrailed_inputs)
|
||||
verbose_proxy_logger.debug("OpenAI Responses API: Processed input messages: %s", data.get("input"))
|
||||
return data
|
||||
|
||||
async def _apply_guardrailed_texts(
|
||||
self,
|
||||
data: dict[str, object],
|
||||
input_data: "str | ResponseInputParam",
|
||||
extracted: _ExtractedInputs,
|
||||
guardrail_to_apply: "CustomGuardrail",
|
||||
guardrailed_inputs: GenericGuardrailAPIInputs,
|
||||
) -> None:
|
||||
returned_texts: Final = guardrailed_inputs.get("texts")
|
||||
if not returned_texts:
|
||||
return
|
||||
rewritten_texts: Final = tuple(returned_texts)
|
||||
offset: Final = 0 if extracted.instructions is None else 1
|
||||
input_texts: Final = rewritten_texts[offset:]
|
||||
expected: Final = 1 if isinstance(input_data, str) else len(extracted.task_mappings)
|
||||
if len(rewritten_texts) != offset + expected:
|
||||
raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name)
|
||||
if offset:
|
||||
data["instructions"] = rewritten_texts[0] # rebind-ok: data is an out-param
|
||||
if isinstance(input_data, str):
|
||||
data["input"] = input_texts[0] # rebind-ok: data is an out-param
|
||||
return
|
||||
await self._apply_guardrail_responses_to_input(
|
||||
messages=input_data,
|
||||
responses=input_texts,
|
||||
task_mappings=extracted.task_mappings,
|
||||
)
|
||||
|
||||
def _extract_guardrail_inputs(
|
||||
self,
|
||||
data: Mapping[str, object],
|
||||
input_data: "str | ResponseInputParam",
|
||||
flattened_tool_groups: Sequence[Sequence[Mapping[str, object]]],
|
||||
*,
|
||||
skip_system: bool = False,
|
||||
) -> _ExtractedInputs:
|
||||
texts_to_check: Final[list[str]] = []
|
||||
instructions: Final = scannable_instructions(data, skip_system=skip_system)
|
||||
texts_to_check: Final[list[str]] = [] if instructions is None else [instructions]
|
||||
images_to_check: Final[list[str]] = []
|
||||
task_mappings: Final[list[tuple[int, int | None]]] = []
|
||||
tools_to_check: Final[list[ChatCompletionToolParam]] = list( # mutable-ok: guardrail inputs want a list
|
||||
|
|
@ -562,6 +612,10 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
texts_to_check.append(input_data)
|
||||
else:
|
||||
for msg_idx, message in enumerate(input_data):
|
||||
if role_out_of_guardrail_scope(
|
||||
_input_item_role(message), skip_system_message=skip_system, skip_tool_message=False
|
||||
):
|
||||
continue
|
||||
self._extract_input_text_and_images(
|
||||
message=message,
|
||||
msg_idx=msg_idx,
|
||||
|
|
@ -577,22 +631,32 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
model: Final = data.get("model")
|
||||
if isinstance(model, str):
|
||||
inputs["model"] = model
|
||||
return _ExtractedInputs(inputs=inputs, task_mappings=tuple(task_mappings))
|
||||
return _ExtractedInputs(inputs=inputs, task_mappings=tuple(task_mappings), instructions=instructions)
|
||||
|
||||
@staticmethod
|
||||
def _written_back_request_fields(
|
||||
data: Mapping[str, object],
|
||||
structured_messages: Sequence[AllMessageValues] | None,
|
||||
structured_messages: Sequence[AllMessageValues],
|
||||
scoped_indices: Sequence[int],
|
||||
scoped_structured_messages: Sequence[AllMessageValues] | None,
|
||||
guardrail_to_apply: "CustomGuardrail",
|
||||
guardrailed_inputs: GenericGuardrailAPIInputs,
|
||||
) -> _RequestFields | None:
|
||||
guardrailed: Final = guardrailed_inputs.get("structured_messages")
|
||||
if guardrailed is None or guardrailed is structured_messages:
|
||||
if guardrailed is None or guardrailed is scoped_structured_messages:
|
||||
return None
|
||||
covers_full_request: Final = len(scoped_indices) == len(structured_messages) or (
|
||||
guardrail_to_apply.structured_messages_cover_full_request() and len(guardrailed) == len(structured_messages)
|
||||
)
|
||||
merged: Final = (
|
||||
guardrailed
|
||||
if covers_full_request
|
||||
else merge_guardrailed_scoped_messages(
|
||||
full_messages=structured_messages, scoped_indices=scoped_indices, guardrailed_scoped=guardrailed
|
||||
)
|
||||
)
|
||||
return _patch_or_convert_request_fields(
|
||||
data.get("input"),
|
||||
data.get("instructions"),
|
||||
structured_messages or (),
|
||||
guardrailed,
|
||||
data.get("input"), data.get("instructions"), structured_messages, merged
|
||||
)
|
||||
|
||||
def extract_request_tool_names(self, data: dict) -> list[str]:
|
||||
|
|
|
|||
|
|
@ -0,0 +1,74 @@
|
|||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.agent_endpoints.auth.agent_access_groups import resolve_managed_agent_ceilings
|
||||
from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityFailure
|
||||
|
||||
|
||||
async def _delegated_resource_subject(user_id: str) -> UserAPIKeyAuth:
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
|
||||
|
||||
human: Final = await MCPRequestHandler.reload_admitted_user(user_id, requires_fresh_policy=True)
|
||||
return human.model_copy(update=MappingProxyType({"mcp_explicit_grants_only": True}))
|
||||
|
||||
|
||||
async def managed_agent_servers(auth: UserAPIKeyAuth) -> tuple[str, ...]:
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
agent: Final = auth.managed_agent_policy
|
||||
if agent is None:
|
||||
return ()
|
||||
|
||||
try:
|
||||
base: Final = frozenset(await MCPRequestHandler.get_allowed_mcp_servers_for_agent(auth))
|
||||
ceilings: Final = await resolve_managed_agent_ceilings(agent)
|
||||
expanded: Final = tuple(
|
||||
frozenset(global_mcp_server_manager.expand_permission_list(sorted(ceiling.mcp_server_ids)))
|
||||
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)
|
||||
own: Final = frozenset(caller_capped)
|
||||
context: Final = auth.managed_agent_context
|
||||
if context is None or context.mode == "autonomous":
|
||||
return tuple(sorted(own))
|
||||
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_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(
|
||||
AgentIdentityFailure(code="policy_unavailable", message="Agent MCP policy is unavailable")
|
||||
)
|
||||
|
||||
|
||||
async def managed_agent_tools(server_id: str, auth: UserAPIKeyAuth) -> list[str] | None:
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
|
||||
|
||||
if server_id not in await managed_agent_servers(auth):
|
||||
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)
|
||||
context: Final = auth.managed_agent_context
|
||||
if context is None or context.mode == "autonomous":
|
||||
return None if own is None else sorted(own)
|
||||
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, allowed_team_ids=frozenset((auth.team_id,)) if auth.team_id else frozenset()
|
||||
)
|
||||
if own is None:
|
||||
return human_tools
|
||||
return sorted(own) if human_tools is None else sorted(frozenset(own).intersection(human_tools))
|
||||
except Exception: # noqa: BLE001 # Authorization boundary: every unresolved policy must deny access
|
||||
raise_identity_failure(
|
||||
AgentIdentityFailure(code="policy_unavailable", message="Agent tool policy is unavailable")
|
||||
)
|
||||
|
|
@ -51,6 +51,7 @@ from litellm.proxy.agent_endpoints.auth.agent_access_groups import (
|
|||
resolve_agent_access_group_ceiling,
|
||||
)
|
||||
from litellm.proxy.agent_endpoints.auth.agent_caller import agent_caller_auth
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_policy
|
||||
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
|
||||
|
|
@ -67,7 +68,6 @@ from litellm.repositories.table_repositories import (
|
|||
AgentsRepository,
|
||||
MCPServerRepository,
|
||||
)
|
||||
from litellm.repositories.user_repository import UserRepository
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -1086,7 +1086,7 @@ class MCPRequestHandler:
|
|||
assert_never(identity.subject_type)
|
||||
|
||||
@staticmethod
|
||||
async def reload_admitted_user(user_id: str) -> UserAPIKeyAuth:
|
||||
async def reload_admitted_user(user_id: str, *, requires_fresh_policy: bool = False) -> UserAPIKeyAuth:
|
||||
"""Reload the live user an interactively-minted envelope references and admit them as themselves.
|
||||
|
||||
The user's own object permission and ``org_id`` ride on the returned ``UserAPIKeyAuth``, and the
|
||||
|
|
@ -1111,6 +1111,7 @@ class MCPRequestHandler:
|
|||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
check_db_only=requires_fresh_policy,
|
||||
)
|
||||
# Resolve the user's own MCP object permission (get_user_object does not load it) so the shared
|
||||
# get_allowed_mcp_servers can grant the user their litellm-granted servers. Reuses the same
|
||||
|
|
@ -1119,6 +1120,7 @@ class MCPRequestHandler:
|
|||
if user_object is not None and object_permission is None and user_object.object_permission_id:
|
||||
object_permission = await get_object_permission(
|
||||
object_permission_id=user_object.object_permission_id,
|
||||
check_db_only=requires_fresh_policy,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
|
|
@ -1147,6 +1149,7 @@ class MCPRequestHandler:
|
|||
# Server-only marker, set AFTER construction: the before-validator strips it from any validated
|
||||
# input, so caller-supplied data (key metadata, JWT claims) can never forge it.
|
||||
admitted.mcp_admitted_user_subject = True
|
||||
admitted.requires_fresh_policy = requires_fresh_policy
|
||||
# Carry each granting team's per-server mcp_rpm_limit: this subject reaches servers through
|
||||
# several teams under its own identity, so without this a cross-team user outruns every team's
|
||||
# limit. Resolved from the same roster-checked sources as the grant union, so a team throttles
|
||||
|
|
@ -1202,7 +1205,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
|
||||
|
|
@ -1234,6 +1237,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
|
||||
|
|
@ -1597,6 +1601,11 @@ class MCPRequestHandler:
|
|||
"""
|
||||
from litellm.proxy.proxy_server import general_settings
|
||||
|
||||
if managed_agent_policy(user_api_key_auth) is not None:
|
||||
from litellm.proxy._experimental.mcp_server.auth.managed_agent_access import managed_agent_servers
|
||||
|
||||
return MCPServerAccess(server_ids=await managed_agent_servers(user_api_key_auth), scope="scoped")
|
||||
|
||||
key_object_permission: Final = MCPRequestHandler._get_key_object_permission(user_api_key_auth)
|
||||
|
||||
try:
|
||||
|
|
@ -1606,7 +1615,7 @@ class MCPRequestHandler:
|
|||
# independent; an opt-out silences only its own source, inside the recursive call).
|
||||
if _is_mcp_admitted_user_subject(user_api_key_auth) and user_api_key_auth is not None:
|
||||
return MCPServerAccess(
|
||||
server_ids=tuple(await MCPRequestHandler._resolve_admitted_subject_servers(user_api_key_auth)),
|
||||
server_ids=tuple(await MCPRequestHandler.resolve_admitted_subject_servers(user_api_key_auth)),
|
||||
)
|
||||
|
||||
# Get allowed servers from key and team
|
||||
|
|
@ -1703,7 +1712,7 @@ class MCPRequestHandler:
|
|||
if user_api_key_auth and user_api_key_auth.agent_id:
|
||||
agent_capped: Final = _agent_capped_servers(
|
||||
allowed_mcp_servers,
|
||||
await MCPRequestHandler._get_allowed_mcp_servers_for_agent(user_api_key_auth),
|
||||
await MCPRequestHandler.get_allowed_mcp_servers_for_agent(user_api_key_auth),
|
||||
await MCPRequestHandler._get_agent_access_group_server_ceiling(user_api_key_auth),
|
||||
)
|
||||
if agent_capped is not None:
|
||||
|
|
@ -1716,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
|
||||
)
|
||||
|
||||
|
|
@ -1829,10 +1838,14 @@ class MCPRequestHandler:
|
|||
scoped.object_permission = auth.object_permission
|
||||
scoped.object_permission_id = auth.object_permission_id
|
||||
scoped.access_group_ids = auth.access_group_ids
|
||||
scoped.requires_fresh_policy = auth.requires_fresh_policy
|
||||
scoped.mcp_explicit_grants_only = auth.mcp_explicit_grants_only
|
||||
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.
|
||||
|
||||
|
|
@ -1849,6 +1862,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
|
||||
|
|
@ -1886,6 +1901,7 @@ class MCPRequestHandler:
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(auth and auth.requires_fresh_policy),
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # per-source isolation: one team's blip must not deny the others
|
||||
# Fault isolation is per SOURCE: an unresolvable team contributes nothing (fail closed for
|
||||
|
|
@ -1932,7 +1948,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
|
||||
|
|
@ -1941,15 +1959,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)
|
||||
|
||||
|
|
@ -2007,7 +2027,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.
|
||||
|
||||
|
|
@ -2029,7 +2051,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):
|
||||
|
|
@ -2088,6 +2110,7 @@ class MCPRequestHandler:
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
|
||||
if not team_obj:
|
||||
|
|
@ -2098,6 +2121,8 @@ class MCPRequestHandler:
|
|||
@staticmethod
|
||||
async def _toolset_tool_permissions(
|
||||
object_permission: LiteLLM_ObjectPermissionTable | None,
|
||||
*,
|
||||
requires_fresh_policy: bool = False,
|
||||
) -> Mapping[str, Sequence[str]]:
|
||||
"""The ``server_id -> tool names`` grants of this permission row's toolsets, empty when it
|
||||
declares none. The shared resolver for the team, org, and internal-user levels, so a toolset
|
||||
|
|
@ -2114,7 +2139,8 @@ class MCPRequestHandler:
|
|||
if object_permission is None or not object_permission.mcp_toolsets:
|
||||
return _EMPTY_TOOLSET_GRANTS
|
||||
resolved: Final = await global_mcp_server_manager.resolve_toolset_tool_permissions(
|
||||
toolset_ids=object_permission.mcp_toolsets
|
||||
toolset_ids=object_permission.mcp_toolsets,
|
||||
requires_fresh_policy=requires_fresh_policy,
|
||||
)
|
||||
if not resolved:
|
||||
raise UnloadableEntitlementError(
|
||||
|
|
@ -2126,10 +2152,15 @@ class MCPRequestHandler:
|
|||
async def _toolset_tools_for_server(
|
||||
object_permission: LiteLLM_ObjectPermissionTable | None,
|
||||
server_id: str,
|
||||
*,
|
||||
requires_fresh_policy: bool = False,
|
||||
) -> Sequence[str] | None:
|
||||
"""Tool names this row's toolsets grant on ``server_id``, ``None`` when its toolsets place
|
||||
no restriction on that server (it declares no toolsets, or none of them name it)."""
|
||||
return (await MCPRequestHandler._toolset_tool_permissions(object_permission)).get(server_id)
|
||||
grants: Final = await MCPRequestHandler._toolset_tool_permissions(
|
||||
object_permission, requires_fresh_policy=requires_fresh_policy
|
||||
)
|
||||
return grants.get(server_id)
|
||||
|
||||
@staticmethod
|
||||
def _union_tool_grants(
|
||||
|
|
@ -2171,6 +2202,7 @@ class MCPRequestHandler:
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -2219,12 +2251,17 @@ class MCPRequestHandler:
|
|||
if not user_api_key_auth:
|
||||
return None
|
||||
|
||||
if managed_agent_policy(user_api_key_auth) is not None:
|
||||
from litellm.proxy._experimental.mcp_server.auth.managed_agent_access import managed_agent_tools
|
||||
|
||||
return await managed_agent_tools(server_id, user_api_key_auth)
|
||||
|
||||
try:
|
||||
# FIRST statement, mirroring get_allowed_mcp_servers: a keyless admitted subject resolves per
|
||||
# source and shares nothing with the single-credential prelude below. Ordering is the invariant:
|
||||
# sat after the prelude, a fault in a lookup the subject never uses denied tools its teams grant.
|
||||
if _is_mcp_admitted_user_subject(user_api_key_auth):
|
||||
return await MCPRequestHandler._resolve_admitted_subject_tools(server_id, user_api_key_auth)
|
||||
return await MCPRequestHandler.resolve_admitted_subject_tools(server_id, user_api_key_auth)
|
||||
|
||||
# Get key and team object permissions (already loaded in main auth flow)
|
||||
key_obj_perm: Final = MCPRequestHandler._get_key_object_permission(user_api_key_auth)
|
||||
|
|
@ -2249,9 +2286,12 @@ class MCPRequestHandler:
|
|||
# tool-level check sees the key's full effective tool scope
|
||||
key_toolset_ids: Final = (key_obj_perm.mcp_toolsets or []) if key_obj_perm else []
|
||||
key_toolset_tools: Final = (
|
||||
(await global_mcp_server_manager.resolve_toolset_tool_permissions(toolset_ids=key_toolset_ids)).get(
|
||||
server_id
|
||||
)
|
||||
(
|
||||
await global_mcp_server_manager.resolve_toolset_tool_permissions(
|
||||
toolset_ids=key_toolset_ids,
|
||||
requires_fresh_policy=user_api_key_auth.requires_fresh_policy,
|
||||
)
|
||||
).get(server_id)
|
||||
if key_toolset_ids
|
||||
else None
|
||||
)
|
||||
|
|
@ -2265,7 +2305,9 @@ class MCPRequestHandler:
|
|||
|
||||
# Tools granted through the team's toolsets restrict this server exactly
|
||||
# as the team's direct tool permissions do, mirroring the key path above
|
||||
team_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(team_obj_perm, server_id)
|
||||
team_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(
|
||||
team_obj_perm, server_id, requires_fresh_policy=user_api_key_auth.requires_fresh_policy
|
||||
)
|
||||
team_tools: Final = MCPRequestHandler._union_tool_grants(team_direct_tools, team_toolset_tools)
|
||||
|
||||
# Apply same inheritance logic as get_allowed_mcp_servers
|
||||
|
|
@ -2291,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(
|
||||
|
|
@ -2334,7 +2376,7 @@ class MCPRequestHandler:
|
|||
if user_api_key_auth.agent_id:
|
||||
# Pre-fetch agent object_permission once to avoid a duplicate DB query.
|
||||
agent_obj_perm: Final = await MCPRequestHandler._get_agent_object_permission(user_api_key_auth)
|
||||
agent_tools: Final = await MCPRequestHandler._get_agent_tool_permissions_for_server(
|
||||
agent_tools: Final = await MCPRequestHandler.get_agent_tool_permissions_for_server(
|
||||
server_id=server_id,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
agent_object_permission=agent_obj_perm,
|
||||
|
|
@ -2365,7 +2407,9 @@ class MCPRequestHandler:
|
|||
if org_obj_perm and org_obj_perm.mcp_tool_permissions
|
||||
else None
|
||||
)
|
||||
org_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(org_obj_perm, server_id)
|
||||
org_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(
|
||||
org_obj_perm, server_id, requires_fresh_policy=user_api_key_auth.requires_fresh_policy
|
||||
)
|
||||
org_tools: Final = MCPRequestHandler._union_tool_grants(org_direct_tools, org_toolset_tools)
|
||||
if org_tools is not None:
|
||||
allowed_tools = (
|
||||
|
|
@ -2456,6 +2500,7 @@ class MCPRequestHandler:
|
|||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
if not raw_server_ids:
|
||||
return []
|
||||
|
|
@ -2502,6 +2547,7 @@ class MCPRequestHandler:
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
if key_object_permission is None:
|
||||
return []
|
||||
|
|
@ -2518,7 +2564,8 @@ class MCPRequestHandler:
|
|||
|
||||
# Get MCP servers from access groups
|
||||
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
|
||||
key_object_permission.mcp_access_groups or []
|
||||
key_object_permission.mcp_access_groups or [],
|
||||
requires_fresh_policy=user_api_key_auth.requires_fresh_policy,
|
||||
)
|
||||
|
||||
# servers referenced in tool permissions should also be accessible
|
||||
|
|
@ -2531,7 +2578,14 @@ class MCPRequestHandler:
|
|||
# ceilings as any other key-level grant
|
||||
toolset_ids: Final = key_object_permission.mcp_toolsets or []
|
||||
toolset_servers: Final = (
|
||||
list((await global_mcp_server_manager.resolve_toolset_tool_permissions(toolset_ids=toolset_ids)).keys())
|
||||
list(
|
||||
(
|
||||
await global_mcp_server_manager.resolve_toolset_tool_permissions(
|
||||
toolset_ids=toolset_ids,
|
||||
requires_fresh_policy=user_api_key_auth.requires_fresh_policy,
|
||||
)
|
||||
).keys()
|
||||
)
|
||||
if toolset_ids
|
||||
else []
|
||||
)
|
||||
|
|
@ -2550,7 +2604,7 @@ class MCPRequestHandler:
|
|||
"""Get allowed MCP servers a caller inherits from the team it is pinned to.
|
||||
|
||||
Exactly one team, or none. A subject that reaches servers through SEVERAL teams does not
|
||||
fan out here: it is resolved one source per team in ``_resolve_admitted_subject_servers``,
|
||||
fan out here: it is resolved one source per team in ``resolve_admitted_subject_servers``,
|
||||
and each of those sources pins a single ``team_id`` before reaching this point. Keeping the
|
||||
fan-out here as well would be a second multi-team path to drift from that one.
|
||||
"""
|
||||
|
|
@ -2568,7 +2622,7 @@ class MCPRequestHandler:
|
|||
which must NOT silently gain the union across every team the user belongs to), and it covers
|
||||
each single-source auth an admitted subject fans out into — those pin a team_id, so they land
|
||||
on the first branch. The admitted subject itself never reaches here: it resolves per source
|
||||
in ``_resolve_admitted_subject_servers`` before this point. The ``UI_TEAM_ID`` sentinel
|
||||
in ``resolve_admitted_subject_servers`` before this point. The ``UI_TEAM_ID`` sentinel
|
||||
resolves to no teams exactly as before."""
|
||||
if user_api_key_auth is None or not user_api_key_auth.team_id:
|
||||
return []
|
||||
|
|
@ -2596,6 +2650,7 @@ class MCPRequestHandler:
|
|||
user_id_upsert=False,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # a team-resolution blip narrows access, never raises
|
||||
verbose_logger.warning("Failed to resolve user teams for MCP grant: %s", e)
|
||||
|
|
@ -2605,7 +2660,12 @@ class MCPRequestHandler:
|
|||
return list(dict.fromkeys(t for t in user_object.teams if t and t != UI_TEAM_ID))
|
||||
|
||||
@staticmethod
|
||||
async def _team_granted_servers(team_obj: LiteLLM_TeamTable, team_access_group_servers: list[str]) -> set[str]:
|
||||
async def _team_granted_servers(
|
||||
team_obj: LiteLLM_TeamTable,
|
||||
team_access_group_servers: list[str],
|
||||
*,
|
||||
requires_fresh_policy: bool = False,
|
||||
) -> set[str]:
|
||||
"""The raw MCP-server set a team grants (before any org ceiling): its object_permission (direct
|
||||
``mcp_servers``, the ``all_proxy_servers`` sentinel → the full registry, legacy access groups,
|
||||
tool-perm-referenced servers, toolset-referenced servers) unioned with its unified
|
||||
|
|
@ -2620,13 +2680,17 @@ class MCPRequestHandler:
|
|||
if SpecialMCPServerName.all_proxy_servers.value in (object_permissions.mcp_servers or []):
|
||||
return set(global_mcp_server_manager.get_registry().keys())
|
||||
legacy_access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
|
||||
object_permissions.mcp_access_groups or []
|
||||
object_permissions.mcp_access_groups or [],
|
||||
requires_fresh_policy=requires_fresh_policy,
|
||||
)
|
||||
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(
|
||||
object_permissions, requires_fresh_policy=requires_fresh_policy
|
||||
)
|
||||
return (
|
||||
set(global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or []))
|
||||
| set(legacy_access_group_servers)
|
||||
| set(global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys())
|
||||
| (await MCPRequestHandler._toolset_tool_permissions(object_permissions)).keys()
|
||||
| toolset_grants.keys()
|
||||
| set(team_access_group_servers)
|
||||
)
|
||||
|
||||
|
|
@ -2667,6 +2731,7 @@ class MCPRequestHandler:
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
if team_obj is None:
|
||||
return []
|
||||
|
|
@ -2680,12 +2745,19 @@ class MCPRequestHandler:
|
|||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
|
||||
servers: Final = await MCPRequestHandler._team_granted_servers(team_obj, team_access_group_servers)
|
||||
servers: Final = await MCPRequestHandler._team_granted_servers(
|
||||
team_obj,
|
||||
team_access_group_servers,
|
||||
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
return list(servers)
|
||||
except Exception as e:
|
||||
if isinstance(e, UnloadableEntitlementError):
|
||||
if isinstance(e, UnloadableEntitlementError) or (
|
||||
user_api_key_auth is not None and user_api_key_auth.requires_fresh_policy
|
||||
):
|
||||
raise
|
||||
verbose_logger.warning("Failed to get allowed MCP servers for team: %s", e)
|
||||
return []
|
||||
|
|
@ -2716,6 +2788,7 @@ class MCPRequestHandler:
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # a named entitlement we cannot read denies, whatever the read failed with
|
||||
raise unloadable from e
|
||||
|
|
@ -2811,7 +2884,8 @@ class MCPRequestHandler:
|
|||
direct_mcp_servers = global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or [])
|
||||
|
||||
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
|
||||
object_permissions.mcp_access_groups or []
|
||||
object_permissions.mcp_access_groups or [],
|
||||
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
|
||||
tool_perm_servers: Final = list(
|
||||
|
|
@ -2820,7 +2894,10 @@ class MCPRequestHandler:
|
|||
|
||||
# servers referenced by the org's toolset grants are part of the org ceiling,
|
||||
# exactly as servers referenced by its inline tool permissions are
|
||||
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(object_permissions)
|
||||
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(
|
||||
object_permissions,
|
||||
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
|
||||
all_servers: Final = tuple(
|
||||
{*direct_mcp_servers, *access_group_servers, *tool_perm_servers, *toolset_grants}
|
||||
|
|
@ -2912,7 +2989,8 @@ class MCPRequestHandler:
|
|||
|
||||
# Get MCP servers from access groups
|
||||
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
|
||||
object_permission.mcp_access_groups or []
|
||||
object_permission.mcp_access_groups or [],
|
||||
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
|
||||
# servers referenced in tool permissions should also be accessible
|
||||
|
|
@ -2961,7 +3039,9 @@ class MCPRequestHandler:
|
|||
return None
|
||||
|
||||
user_id: Final = user_api_key_auth.user_id
|
||||
object_permission_id: Final = await MCPRequestHandler._user_object_permission_id(user_id, prisma_client)
|
||||
object_permission_id: Final = await MCPRequestHandler._user_object_permission_id(
|
||||
user_id, prisma_client, check_db_only=user_api_key_auth.requires_fresh_policy
|
||||
)
|
||||
if object_permission_id is None:
|
||||
return None
|
||||
|
||||
|
|
@ -2971,6 +3051,7 @@ class MCPRequestHandler:
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
if object_permission is None:
|
||||
raise ValueError(
|
||||
|
|
@ -2979,7 +3060,9 @@ class MCPRequestHandler:
|
|||
return object_permission
|
||||
|
||||
@staticmethod
|
||||
async def _user_object_permission_id(user_id: str, prisma_client: "PrismaClient") -> str | None:
|
||||
async def _user_object_permission_id(
|
||||
user_id: str, prisma_client: "PrismaClient", *, check_db_only: bool = False
|
||||
) -> str | None:
|
||||
"""The permission row this human's user row links to, or None when they link none.
|
||||
|
||||
Caches the link (with a sentinel for "links none") so a human without an entitlement costs no
|
||||
|
|
@ -2988,16 +3071,23 @@ class MCPRequestHandler:
|
|||
whether someone is entitled is the state that existed before this level, so it places no
|
||||
ceiling. Only a link we DID resolve can make the caller deny.
|
||||
"""
|
||||
from litellm.proxy.auth.auth_checks import get_user_object
|
||||
from litellm.proxy.proxy_server import user_api_key_cache
|
||||
|
||||
cache_key: Final = user_object_permission_id_cache_key(user_id)
|
||||
try:
|
||||
cached: Final[object] = await user_api_key_cache.async_get_cache(key=cache_key)
|
||||
cached: Final[object] = None if check_db_only else await user_api_key_cache.async_get_cache(key=cache_key)
|
||||
if cached == USER_NO_MCP_PERMISSION_SENTINEL:
|
||||
return None
|
||||
if isinstance(cached, str) and cached:
|
||||
return cached
|
||||
user_row: Final = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id})
|
||||
user_row: Final = await get_user_object(
|
||||
user_id=user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
check_db_only=check_db_only,
|
||||
)
|
||||
linked: Final[object] = getattr(user_row, "object_permission_id", None) if user_row is not None else None
|
||||
object_permission_id: Final = linked if isinstance(linked, str) and linked else None
|
||||
await user_api_key_cache.async_set_cache(
|
||||
|
|
@ -3006,7 +3096,9 @@ class MCPRequestHandler:
|
|||
ttl=get_management_object_ttl(user_api_key_cache),
|
||||
)
|
||||
return object_permission_id
|
||||
except Exception as e: # noqa: BLE001 # unknown whether entitled at all: no ceiling, as before
|
||||
except Exception as e: # noqa: BLE001 # Legacy callers retain their existing optional user-ceiling behavior
|
||||
if check_db_only:
|
||||
raise HTTPException(503, "User policy is unavailable") from e
|
||||
verbose_logger.warning("MCP user entitlement: link for %r unresolved, no ceiling: %s", user_id, e)
|
||||
return None
|
||||
|
||||
|
|
@ -3031,13 +3123,17 @@ class MCPRequestHandler:
|
|||
return []
|
||||
|
||||
direct_mcp_servers = global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or [])
|
||||
fresh: Final = bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy)
|
||||
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
|
||||
object_permissions.mcp_access_groups or []
|
||||
object_permissions.mcp_access_groups or [],
|
||||
requires_fresh_policy=fresh,
|
||||
)
|
||||
tool_perm_servers: Final = list(
|
||||
global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys()
|
||||
)
|
||||
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(object_permissions)
|
||||
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(
|
||||
object_permissions, requires_fresh_policy=fresh
|
||||
)
|
||||
return tuple({*direct_mcp_servers, *access_group_servers, *tool_perm_servers, *toolset_grants})
|
||||
except Exception as e: # noqa: BLE001 # any resolution fault is an unresolved ceiling, never "no ceiling"
|
||||
verbose_logger.warning("Failed to get allowed MCP servers for user: %s", e)
|
||||
|
|
@ -3075,7 +3171,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]:
|
||||
|
|
@ -3119,9 +3215,13 @@ class MCPRequestHandler:
|
|||
(any non-empty entitlement, or an unresolved one, disqualifies), exactly as
|
||||
``operator_open_server_ids`` reads the same row. The one owner of this predicate: the
|
||||
server-axis registry resolution in ``get_allowed_mcp_servers`` and the tools-axis open
|
||||
channel in ``_resolve_admitted_subject_tools`` both consult it, so the two axes cannot
|
||||
channel in ``resolve_admitted_subject_tools`` both consult it, so the two axes cannot
|
||||
disagree."""
|
||||
if user_api_key_auth is None or not user_api_key_has_admin_view(user_api_key_auth):
|
||||
if (
|
||||
user_api_key_auth is None
|
||||
or user_api_key_auth.mcp_explicit_grants_only
|
||||
or not user_api_key_has_admin_view(user_api_key_auth)
|
||||
):
|
||||
return False
|
||||
object_permission: Final = user_api_key_auth.object_permission
|
||||
credential_scoped: Final = (
|
||||
|
|
@ -3167,7 +3267,11 @@ class MCPRequestHandler:
|
|||
user_direct_tools: Final = global_mcp_server_manager.expand_tool_permissions(
|
||||
object_permissions.mcp_tool_permissions
|
||||
).get(server_id)
|
||||
user_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(object_permissions, server_id)
|
||||
user_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(
|
||||
object_permissions,
|
||||
server_id,
|
||||
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
user_tools: Final = MCPRequestHandler._union_tool_grants(user_direct_tools, user_toolset_tools)
|
||||
if user_tools is None:
|
||||
return allowed_tools
|
||||
|
|
@ -3176,7 +3280,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,
|
||||
|
|
@ -3184,7 +3288,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 (
|
||||
|
|
@ -3196,7 +3300,9 @@ class MCPRequestHandler:
|
|||
return allowed_tools
|
||||
try:
|
||||
team_obj_perm: Final = await MCPRequestHandler._get_team_object_permission(caller_auth)
|
||||
team_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(team_obj_perm, server_id)
|
||||
team_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(
|
||||
team_obj_perm, server_id, requires_fresh_policy=caller_auth.requires_fresh_policy
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # an unresolved caller team must deny, not widen
|
||||
verbose_logger.warning(
|
||||
"MCP agent caller team tool ceiling unresolvable, denying tools on %r: %s", server_id, e
|
||||
|
|
@ -3241,7 +3347,11 @@ class MCPRequestHandler:
|
|||
end_user_direct_tools: Final = global_mcp_server_manager.expand_tool_permissions(
|
||||
object_permissions.mcp_tool_permissions
|
||||
).get(server_id)
|
||||
end_user_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(object_permissions, server_id)
|
||||
end_user_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(
|
||||
object_permissions,
|
||||
server_id,
|
||||
requires_fresh_policy=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
end_user_tools: Final = MCPRequestHandler._union_tool_grants(end_user_direct_tools, end_user_toolset_tools)
|
||||
if end_user_tools is None:
|
||||
return allowed_tools
|
||||
|
|
@ -3302,6 +3412,11 @@ class MCPRequestHandler:
|
|||
if not user_api_key_auth or not user_api_key_auth.agent_id:
|
||||
return None
|
||||
|
||||
managed: Final = managed_agent_policy(user_api_key_auth)
|
||||
if managed is not None:
|
||||
permission: Final = managed.object_permission
|
||||
return LiteLLM_ObjectPermissionTable.model_validate(permission) if permission is not None else None
|
||||
|
||||
if prisma_client is None:
|
||||
verbose_logger.debug("prisma_client is None")
|
||||
return None
|
||||
|
|
@ -3319,7 +3434,7 @@ class MCPRequestHandler:
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
async def _get_allowed_mcp_servers_for_agent(
|
||||
async def get_allowed_mcp_servers_for_agent(
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
agent_object_permission: LiteLLM_ObjectPermissionTable | None = None,
|
||||
) -> list[str]:
|
||||
|
|
@ -3358,12 +3473,16 @@ class MCPRequestHandler:
|
|||
obj_perm.mcp_servers or []
|
||||
)
|
||||
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
|
||||
obj_perm.mcp_access_groups or []
|
||||
obj_perm.mcp_access_groups or [],
|
||||
requires_fresh_policy=user_api_key_auth.requires_fresh_policy,
|
||||
)
|
||||
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(obj_perm)
|
||||
return list({*expanded_direct_servers, *access_group_servers, *toolset_grants})
|
||||
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(
|
||||
obj_perm, requires_fresh_policy=user_api_key_auth.requires_fresh_policy
|
||||
)
|
||||
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 isinstance(e, UnloadableEntitlementError):
|
||||
if managed_agent_policy(user_api_key_auth) is not None or isinstance(e, UnloadableEntitlementError):
|
||||
raise
|
||||
verbose_logger.warning("Failed to get allowed MCP servers for agent: %s", e)
|
||||
return []
|
||||
|
|
@ -3390,7 +3509,7 @@ class MCPRequestHandler:
|
|||
return frozenset(global_mcp_server_manager.expand_permission_list(sorted(ceiling.mcp_server_ids)))
|
||||
|
||||
@staticmethod
|
||||
async def _get_agent_tool_permissions_for_server(
|
||||
async def get_agent_tool_permissions_for_server(
|
||||
server_id: str,
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
agent_object_permission: LiteLLM_ObjectPermissionTable | None = None,
|
||||
|
|
@ -3430,11 +3549,13 @@ class MCPRequestHandler:
|
|||
if obj_perm.mcp_tool_permissions
|
||||
else None
|
||||
)
|
||||
toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(obj_perm, server_id)
|
||||
toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(
|
||||
obj_perm, server_id, requires_fresh_policy=user_api_key_auth.requires_fresh_policy
|
||||
)
|
||||
agent_tools: Final = MCPRequestHandler._union_tool_grants(direct_tools, toolset_tools)
|
||||
return list(agent_tools) if agent_tools else None
|
||||
return list(agent_tools) if agent_tools is not None else None
|
||||
except Exception as e:
|
||||
if isinstance(e, UnloadableEntitlementError):
|
||||
if managed_agent_policy(user_api_key_auth) is not None or isinstance(e, UnloadableEntitlementError):
|
||||
raise
|
||||
verbose_logger.warning("Failed to get agent tool permissions for server: %s", e)
|
||||
return None
|
||||
|
|
@ -3452,28 +3573,38 @@ class MCPRequestHandler:
|
|||
return server_ids
|
||||
|
||||
@staticmethod
|
||||
async def _get_db_server_ids_for_access_groups(prisma_client, access_groups: list[str]) -> set[str]:
|
||||
async def _get_db_server_ids_for_access_groups(
|
||||
prisma_client,
|
||||
access_groups: list[str],
|
||||
*,
|
||||
use_writer: bool = False,
|
||||
) -> set[str]:
|
||||
"""
|
||||
Helper to get server_ids from DB servers that match any of the given access groups.
|
||||
"""
|
||||
server_ids: Final[set[str]] = set()
|
||||
if access_groups and prisma_client is not None:
|
||||
try:
|
||||
mcp_servers: Final = await MCPServerRepository(prisma_client).table.find_many(
|
||||
mcp_servers: Final = await MCPServerRepository(prisma_client, use_writer=use_writer).table.find_many(
|
||||
where={"mcp_access_groups": {"hasSome": access_groups}}
|
||||
)
|
||||
for server in mcp_servers:
|
||||
server_ids.add(server.server_id)
|
||||
except Exception as e:
|
||||
if use_writer:
|
||||
raise
|
||||
verbose_logger.debug("Error getting MCP servers from access groups: %s", e)
|
||||
return server_ids
|
||||
|
||||
@staticmethod
|
||||
async def _get_mcp_servers_from_access_groups(
|
||||
access_groups: list[str],
|
||||
*,
|
||||
requires_fresh_policy: bool = False,
|
||||
) -> list[str]:
|
||||
"""
|
||||
Resolve MCP access groups to server IDs by querying BOTH the MCP server table (DB) AND config-loaded servers
|
||||
Resolve MCP access groups to server IDs by querying BOTH the MCP server table (DB) AND config-loaded servers.
|
||||
``requires_fresh_policy`` reads the writer and propagates a read fault instead of resolving to no servers.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
|
|
@ -3489,11 +3620,15 @@ class MCPRequestHandler:
|
|||
)
|
||||
|
||||
# Use the new helper for DB servers
|
||||
db_server_ids = await MCPRequestHandler._get_db_server_ids_for_access_groups(prisma_client, access_groups)
|
||||
db_server_ids = await MCPRequestHandler._get_db_server_ids_for_access_groups(
|
||||
prisma_client, access_groups, use_writer=requires_fresh_policy
|
||||
)
|
||||
server_ids.update(db_server_ids)
|
||||
|
||||
return list(server_ids)
|
||||
except Exception as e:
|
||||
if requires_fresh_policy:
|
||||
raise
|
||||
verbose_logger.warning("Failed to get MCP servers from access groups: %s", e)
|
||||
return []
|
||||
|
||||
|
|
@ -3548,6 +3683,7 @@ class MCPRequestHandler:
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
if key_object_permission is None:
|
||||
return []
|
||||
|
|
@ -3591,6 +3727,7 @@ class MCPRequestHandler:
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=bool(user_api_key_auth and user_api_key_auth.requires_fresh_policy),
|
||||
)
|
||||
if team_obj is None:
|
||||
verbose_logger.debug("team_obj is None")
|
||||
|
|
|
|||
|
|
@ -181,6 +181,7 @@ from litellm.proxy._types import (
|
|||
UserAPIKeyAuth,
|
||||
is_per_server_oauth_discovery_eligible,
|
||||
)
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_policy
|
||||
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
|
||||
from litellm.proxy.common_utils.user_api_key_cache import get_management_object_ttl
|
||||
|
|
@ -3428,7 +3429,9 @@ class MCPServerManager:
|
|||
|
||||
``allow_all_server_ids`` / ``submitted_server_ids`` are injectable so the server union,
|
||||
which precomputes both for its fallback path, does not compute them twice."""
|
||||
if user_api_key_auth is not None and user_api_key_auth.mcp_toolset_id is not None:
|
||||
if user_api_key_auth is not None and (
|
||||
user_api_key_auth.mcp_toolset_id is not None or user_api_key_auth.mcp_explicit_grants_only
|
||||
):
|
||||
return set()
|
||||
if allow_all_server_ids is None:
|
||||
allow_all_server_ids = self.get_allow_all_keys_server_ids()
|
||||
|
|
@ -3477,9 +3480,14 @@ class MCPServerManager:
|
|||
2. If admin and no object_permission, return all servers
|
||||
3. Otherwise, use standard permission checks
|
||||
"""
|
||||
if managed_agent_policy(user_api_key_auth) is not None:
|
||||
managed: Final = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth)
|
||||
return managed if access is None else [server for server in managed if server in access.server_ids]
|
||||
|
||||
from litellm.proxy.proxy_server import general_settings as proxy_general_settings
|
||||
|
||||
resolved_general_settings: Final = proxy_general_settings if general_settings is None else general_settings
|
||||
explicit_grants_only: Final = bool(user_api_key_auth and user_api_key_auth.mcp_explicit_grants_only)
|
||||
allow_all_server_ids: Final = self.get_allow_all_keys_server_ids()
|
||||
|
||||
# A keyless admitted subject is resolved per grant source, and channel decisions that are
|
||||
|
|
@ -3511,7 +3519,7 @@ class MCPServerManager:
|
|||
# only keys without their own mcp_servers list get submitted servers unioned in.
|
||||
submitted_server_ids: Final = (
|
||||
[]
|
||||
if has_explicit_object_permission
|
||||
if has_explicit_object_permission or explicit_grants_only
|
||||
else await self._get_active_submitted_mcp_server_ids_for_user(user_api_key_auth)
|
||||
)
|
||||
|
||||
|
|
@ -3580,12 +3588,14 @@ class MCPServerManager:
|
|||
return [
|
||||
server_id
|
||||
for server_id in dict.fromkeys(allow_all_server_ids + submitted_server_ids)
|
||||
if scope is None or server_id == scope
|
||||
if not explicit_grants_only and (scope is None or server_id == scope)
|
||||
]
|
||||
|
||||
async def resolve_toolset_tool_permissions(
|
||||
self,
|
||||
toolset_ids: list[str],
|
||||
*,
|
||||
requires_fresh_policy: bool = False,
|
||||
) -> dict[str, list[str]]:
|
||||
"""
|
||||
Resolve a list of toolset IDs into a mcp_tool_permissions dict.
|
||||
|
|
@ -3595,6 +3605,10 @@ class MCPServerManager:
|
|||
Redis-backed ``DualCache`` in production) so that cache entries are
|
||||
shared across workers and cold-cache DB hits are minimised.
|
||||
|
||||
``requires_fresh_policy`` bypasses the cache and reads the writer so a
|
||||
revocation is honoured on the very next request; a read fault then
|
||||
propagates instead of resolving to no grants.
|
||||
|
||||
A row names a tool on the server identified by ``server_id``, so the
|
||||
stored name is the tool's own name and is used as written. It is never
|
||||
reduced by the server's wire prefix: that prefix is added on the way out
|
||||
|
|
@ -3609,12 +3623,16 @@ class MCPServerManager:
|
|||
return {}
|
||||
|
||||
cache_key: Final = "toolset_perms:" + ",".join(sorted(toolset_ids))
|
||||
cached: Final[dict[str, list[str]] | None] = await user_api_key_cache.async_get_cache(key=cache_key)
|
||||
cached: Final[dict[str, list[str]] | None] = (
|
||||
None if requires_fresh_policy else await user_api_key_cache.async_get_cache(key=cache_key)
|
||||
)
|
||||
if cached is not None:
|
||||
return cached
|
||||
|
||||
try:
|
||||
toolsets: Final = await list_mcp_toolsets(prisma_client, toolset_ids=toolset_ids)
|
||||
toolsets: Final = await list_mcp_toolsets(
|
||||
prisma_client, toolset_ids=toolset_ids, use_writer=requires_fresh_policy
|
||||
)
|
||||
tool_permissions: Final[dict[str, list[str]]] = {}
|
||||
for toolset in toolsets:
|
||||
for tool in toolset.tools:
|
||||
|
|
@ -3628,6 +3646,8 @@ class MCPServerManager:
|
|||
)
|
||||
return tool_permissions
|
||||
except Exception as e:
|
||||
if requires_fresh_policy:
|
||||
raise
|
||||
verbose_logger.warning("Failed to resolve toolset permissions: %s", e)
|
||||
return {}
|
||||
|
||||
|
|
|
|||
|
|
@ -65,9 +65,9 @@ class MCPToolsetTable(Protocol):
|
|||
async def delete(self, where: Mapping[str, object]) -> MCPToolsetRow: ...
|
||||
|
||||
|
||||
def _toolset_table(prisma_client: PrismaClient) -> MCPToolsetTable:
|
||||
def _toolset_table(prisma_client: PrismaClient, *, use_writer: bool = False) -> MCPToolsetTable:
|
||||
"""The toolset table actions of the prisma client."""
|
||||
return MCPToolsetRepository(prisma_client).table
|
||||
return MCPToolsetRepository(prisma_client, use_writer=use_writer).table
|
||||
|
||||
|
||||
def _toolset_from_row(row: MCPToolsetRow) -> MCPToolset:
|
||||
|
|
@ -107,12 +107,16 @@ async def get_mcp_toolset(
|
|||
async def list_mcp_toolsets(
|
||||
prisma_client: PrismaClient,
|
||||
toolset_ids: Sequence[str] | None = None,
|
||||
*,
|
||||
use_writer: bool = False,
|
||||
) -> Sequence[MCPToolset]:
|
||||
try:
|
||||
where: Final[Mapping[str, object]] = {} if toolset_ids is None else {"toolset_id": {"in": toolset_ids}}
|
||||
rows: Final = await _toolset_table(prisma_client).find_many(where=where)
|
||||
rows: Final = await _toolset_table(prisma_client, use_writer=use_writer).find_many(where=where)
|
||||
return [_toolset_from_row(r) for r in rows]
|
||||
except Exception as e:
|
||||
if use_writer:
|
||||
raise
|
||||
verbose_proxy_logger.warning("litellm.proxy._experimental.mcp_server.toolset_db::list_mcp_toolsets - %s", e)
|
||||
return []
|
||||
|
||||
|
|
|
|||
|
|
@ -58,6 +58,7 @@ async def resolve_ui_session_team_ids(
|
|||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
check_db_only=user_api_key_auth.requires_fresh_policy,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
|
@ -92,7 +93,9 @@ async def admitted_user_context(user_api_key_auth: UserAPIKeyAuth) -> UserAPIKey
|
|||
)
|
||||
|
||||
try:
|
||||
admitted: Final = await MCPRequestHandler.reload_admitted_user(user_id)
|
||||
admitted: Final = await MCPRequestHandler.reload_admitted_user(
|
||||
user_id, requires_fresh_policy=user_api_key_auth.requires_fresh_policy
|
||||
)
|
||||
except HTTPException as e:
|
||||
verbose_logger.warning("MCP dashboard session: admitted-subject reload failed for %s: %s", user_id, e.detail)
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -597,7 +597,6 @@ async def get_agent_card(
|
|||
if agent is None:
|
||||
raise HTTPException(status_code=404, detail=f"Agent '{agent_id}' not found")
|
||||
|
||||
# Check agent permission (skip for admin users)
|
||||
is_allowed: Final = await AgentRequestHandler.is_agent_allowed(
|
||||
agent_id=agent.agent_id,
|
||||
user_api_key_auth=user_api_key_dict,
|
||||
|
|
|
|||
|
|
@ -1,13 +1,16 @@
|
|||
import asyncio
|
||||
from collections.abc import Awaitable, Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Final, TypeAlias
|
||||
from typing import TYPE_CHECKING, Final, TypeAlias
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import LiteLLM_AccessGroupTable
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
AccessGroupIds: TypeAlias = tuple[str, ...]
|
||||
AccessGroupIdsLoader: TypeAlias = Callable[[str], Awaitable[AccessGroupIds]] # mutable-ok: Callable params
|
||||
LoadedAccessGroup: TypeAlias = LiteLLM_AccessGroupTable | None
|
||||
|
|
@ -34,7 +37,7 @@ async def _registry_access_group_ids(agent_id: str) -> AccessGroupIds:
|
|||
return tuple(agent.access_group_ids or ()) if agent is not None else ()
|
||||
|
||||
|
||||
async def _load_access_group(access_group_id: str) -> LoadedAccessGroup:
|
||||
async def _load_access_group(access_group_id: str, *, check_db_only: bool = False) -> LoadedAccessGroup:
|
||||
from litellm.proxy.auth.auth_checks import get_access_object
|
||||
from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache
|
||||
|
||||
|
|
@ -47,8 +50,11 @@ async def _load_access_group(access_group_id: str) -> LoadedAccessGroup:
|
|||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
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
|
||||
)
|
||||
|
|
@ -59,13 +65,20 @@ async def resolve_agent_access_group_ceiling(
|
|||
agent_id: str,
|
||||
load_access_group_ids: AccessGroupIdsLoader = _registry_access_group_ids,
|
||||
load_access_group: AccessGroupLoader = _load_access_group,
|
||||
*,
|
||||
check_db_only: bool = False,
|
||||
) -> AgentAccessGroupCeiling | None:
|
||||
"""``None`` when the agent has no access groups attached, so nothing is capped."""
|
||||
access_group_ids: Final = await load_access_group_ids(agent_id)
|
||||
if not access_group_ids:
|
||||
return None
|
||||
|
||||
loaded: Final = await asyncio.gather(*(load_access_group(group_id) for group_id in access_group_ids))
|
||||
loaded: Final = await asyncio.gather(
|
||||
*(
|
||||
_load_access_group(group_id, check_db_only=True) if check_db_only else load_access_group(group_id)
|
||||
for group_id in access_group_ids
|
||||
)
|
||||
)
|
||||
groups: Final = tuple(group for group in loaded if group is not None)
|
||||
return AgentAccessGroupCeiling(
|
||||
access_group_ids=access_group_ids,
|
||||
|
|
@ -73,3 +86,16 @@ async def resolve_agent_access_group_ceiling(
|
|||
mcp_server_ids=frozenset(server_id for group in groups for server_id in group.access_mcp_server_ids),
|
||||
agent_ids=frozenset(target_id for group in groups for target_id in group.access_agent_ids),
|
||||
)
|
||||
|
||||
|
||||
async def resolve_managed_agent_ceilings(agent: "AgentResponse") -> tuple[AgentAccessGroupCeiling, ...]:
|
||||
async def authoritative_group(group_id: str) -> LoadedAccessGroup:
|
||||
return await _load_access_group(group_id, check_db_only=True)
|
||||
|
||||
async def manual_ids(_agent_id: str) -> AccessGroupIds:
|
||||
return tuple(agent.access_group_ids or ())
|
||||
|
||||
manual: Final = await resolve_agent_access_group_ceiling(
|
||||
agent.agent_id, load_access_group_ids=manual_ids, load_access_group=authoritative_group
|
||||
)
|
||||
return (manual,) if manual is not None else ()
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ can only narrow access and need no trust.
|
|||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -45,7 +46,7 @@ def agent_caller_auth(user_api_key_auth: UserAPIKeyAuth) -> UserAPIKeyAuth | Non
|
|||
user_id=caller.user_id,
|
||||
team_id=caller.team_id,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
)
|
||||
).model_copy(update=MappingProxyType({"requires_fresh_policy": user_api_key_auth.requires_fresh_policy}))
|
||||
|
||||
|
||||
async def load_agent_caller_team(user_api_key_auth: UserAPIKeyAuth) -> LiteLLM_TeamTable | None:
|
||||
|
|
|
|||
|
|
@ -8,8 +8,11 @@ Follows the same pattern as MCP permission handling.
|
|||
import asyncio
|
||||
from collections.abc import Awaitable, Callable, Sequence
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Final, TypeAlias
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.proxy._experimental.mcp_server.ui_session_utils import build_effective_auth_contexts
|
||||
from litellm.proxy._types import (
|
||||
|
|
@ -24,6 +27,7 @@ from litellm.proxy.agent_endpoints.auth.agent_access_groups import (
|
|||
resolve_agent_access_group_ceiling,
|
||||
)
|
||||
from litellm.proxy.agent_endpoints.auth.agent_caller import agent_caller_auth
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_policy
|
||||
from litellm.repositories.table_repositories import AgentsRepository
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
|
|
@ -83,13 +87,23 @@ 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."""
|
||||
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)
|
||||
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, 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)
|
||||
agent_ceiling: Final = await AgentRequestHandler._agent_access_group_ceiling(
|
||||
user_api_key_auth, resolve_ceiling, strict=strict
|
||||
)
|
||||
if agent_ceiling is None:
|
||||
return own_access
|
||||
if isinstance(own_access, UnrestrictedAgentAccess):
|
||||
|
|
@ -97,20 +111,26 @@ 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(
|
||||
async def resolve_key_team_agent_access(
|
||||
user_api_key_auth: UserAPIKeyAuth | None,
|
||||
*,
|
||||
strict: bool = False,
|
||||
) -> AgentAccess:
|
||||
try:
|
||||
key_access: Final = await AgentRequestHandler._get_allowed_agents_for_key(user_api_key_auth)
|
||||
team_access: Final = await AgentRequestHandler._get_allowed_agents_for_team(user_api_key_auth)
|
||||
key_access: Final = await AgentRequestHandler.get_allowed_agents_for_key(user_api_key_auth, strict=strict)
|
||||
team_access: Final = await AgentRequestHandler._get_allowed_agents_for_team(
|
||||
user_api_key_auth, strict=strict
|
||||
)
|
||||
except Exception as e:
|
||||
if strict:
|
||||
raise HTTPException(503, "Agent invocation policy is unavailable") from e
|
||||
verbose_logger.warning("Failed to get allowed agents: %s", e)
|
||||
return UnrestrictedAgentAccess()
|
||||
return _intersect_agent_access(key_access, team_access)
|
||||
|
|
@ -119,10 +139,16 @@ class AgentRequestHandler:
|
|||
async def _agent_access_group_ceiling(
|
||||
user_api_key_auth: UserAPIKeyAuth | None,
|
||||
resolve_ceiling: CeilingResolver,
|
||||
*,
|
||||
strict: bool = False,
|
||||
) -> frozenset[str] | None:
|
||||
if user_api_key_auth is None or not user_api_key_auth.agent_id:
|
||||
return None
|
||||
ceiling: Final = await resolve_ceiling(user_api_key_auth.agent_id)
|
||||
ceiling: Final = (
|
||||
await resolve_agent_access_group_ceiling(user_api_key_auth.agent_id, check_db_only=True)
|
||||
if strict
|
||||
else await resolve_ceiling(user_api_key_auth.agent_id)
|
||||
)
|
||||
if ceiling is None:
|
||||
return None
|
||||
return _to_stable_ids(ceiling.agent_ids)
|
||||
|
|
@ -144,6 +170,49 @@ class AgentRequestHandler:
|
|||
bool: True if agent is allowed, False otherwise
|
||||
"""
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore
|
||||
from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityFailure
|
||||
|
||||
registered: Final = global_agent_registry.get_agent_by_id(agent_id)
|
||||
registry_managed: Final = isinstance(registered, AgentResponse) and registered.identity_managed
|
||||
if registry_managed or (registered is None and prisma_client is not None):
|
||||
target: Final = await AgentIdentityStore.from_client(prisma_client).agent(agent_id)
|
||||
if isinstance(target, AgentIdentityFailure):
|
||||
if registry_managed:
|
||||
raise_identity_failure(target)
|
||||
elif target is None and registry_managed:
|
||||
return False
|
||||
elif isinstance(target, AgentResponse) and target.identity_managed:
|
||||
if (
|
||||
not target.enabled
|
||||
or target.identity is None
|
||||
or not target.identity.active
|
||||
or user_api_key_auth is None
|
||||
):
|
||||
return False
|
||||
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) # pyright: ignore[reportPrivateUsage] # the authoritative key reload has no public seam
|
||||
if key_hash
|
||||
and managed_agent_policy(user_api_key_auth) is None
|
||||
and not user_api_key_auth.is_session_token
|
||||
else user_api_key_auth
|
||||
)
|
||||
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,
|
||||
build_effective_auth_contexts,
|
||||
)
|
||||
return target.agent_id in explicit
|
||||
|
||||
match await AgentRequestHandler.resolve_agent_access(user_api_key_auth, resolve_ceiling):
|
||||
case UnrestrictedAgentAccess():
|
||||
|
|
@ -202,8 +271,10 @@ class AgentRequestHandler:
|
|||
return team_obj.object_permission
|
||||
|
||||
@staticmethod
|
||||
async def _get_allowed_agents_for_key(
|
||||
async def get_allowed_agents_for_key(
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
*,
|
||||
strict: bool = False,
|
||||
) -> AgentAccess:
|
||||
"""
|
||||
Get allowed agents for a key.
|
||||
|
|
@ -237,24 +308,36 @@ class AgentRequestHandler:
|
|||
return UnrestrictedAgentAccess()
|
||||
|
||||
access_group_agents: Final = (
|
||||
tuple(await AgentRequestHandler._get_agents_from_access_groups(list(declared_access_groups)))
|
||||
tuple(
|
||||
await AgentRequestHandler._get_agents_from_access_groups(
|
||||
declared_access_groups, check_db_only=strict
|
||||
)
|
||||
)
|
||||
if declared_access_groups
|
||||
else ()
|
||||
)
|
||||
unified_agents: Final = (
|
||||
tuple(await AgentRequestHandler._get_unified_access_group_agents(list(key_access_group_ids)))
|
||||
tuple(
|
||||
await AgentRequestHandler._get_unified_access_group_agents(
|
||||
key_access_group_ids, check_db_only=strict
|
||||
)
|
||||
)
|
||||
if key_access_group_ids
|
||||
else ()
|
||||
)
|
||||
|
||||
return RestrictedAgentAccess(frozenset(direct_agents + access_group_agents + unified_agents))
|
||||
except Exception as e:
|
||||
if strict:
|
||||
raise HTTPException(503, "Agent invocation policy is unavailable") from e
|
||||
verbose_logger.warning("Failed to get allowed agents for key: %s", e)
|
||||
return UnrestrictedAgentAccess()
|
||||
|
||||
@staticmethod
|
||||
async def _get_allowed_agents_for_team(
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
*,
|
||||
strict: bool = False,
|
||||
) -> AgentAccess:
|
||||
"""
|
||||
Get allowed agents for a team.
|
||||
|
|
@ -263,7 +346,7 @@ class AgentRequestHandler:
|
|||
2. Also includes agents from team's access_group_ids (unified access groups)
|
||||
|
||||
Fetches the team object once and reuses it for both permission sources.
|
||||
Declared-but-empty grants stay restricted; see `_get_allowed_agents_for_key`.
|
||||
Declared-but-empty grants stay restricted; see `get_allowed_agents_for_key`.
|
||||
"""
|
||||
if user_api_key_auth is None:
|
||||
return UnrestrictedAgentAccess()
|
||||
|
|
@ -280,7 +363,7 @@ class AgentRequestHandler:
|
|||
)
|
||||
|
||||
if not prisma_client:
|
||||
return UnrestrictedAgentAccess()
|
||||
return RestrictedAgentAccess(frozenset()) if strict else UnrestrictedAgentAccess()
|
||||
|
||||
# Fetch the team object once for both permission sources
|
||||
team_obj: Final = await get_team_object(
|
||||
|
|
@ -289,10 +372,11 @@ class AgentRequestHandler:
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=strict,
|
||||
)
|
||||
|
||||
if team_obj is None:
|
||||
return UnrestrictedAgentAccess()
|
||||
return RestrictedAgentAccess(frozenset()) if strict else UnrestrictedAgentAccess()
|
||||
|
||||
# 1. Get agents from object_permission (native permissions)
|
||||
object_permissions: Final = team_obj.object_permission
|
||||
|
|
@ -307,18 +391,28 @@ class AgentRequestHandler:
|
|||
return UnrestrictedAgentAccess()
|
||||
|
||||
access_group_agents: Final = (
|
||||
tuple(await AgentRequestHandler._get_agents_from_access_groups(list(declared_access_groups)))
|
||||
tuple(
|
||||
await AgentRequestHandler._get_agents_from_access_groups(
|
||||
declared_access_groups, check_db_only=strict
|
||||
)
|
||||
)
|
||||
if declared_access_groups
|
||||
else ()
|
||||
)
|
||||
unified_agents: Final = (
|
||||
tuple(await AgentRequestHandler._get_unified_access_group_agents(list(team_access_group_ids)))
|
||||
tuple(
|
||||
await AgentRequestHandler._get_unified_access_group_agents(
|
||||
team_access_group_ids, check_db_only=strict
|
||||
)
|
||||
)
|
||||
if team_access_group_ids
|
||||
else ()
|
||||
)
|
||||
|
||||
return RestrictedAgentAccess(frozenset(direct_agents + access_group_agents + unified_agents))
|
||||
except Exception as e:
|
||||
if strict:
|
||||
raise HTTPException(503, "Agent invocation policy is unavailable") from e
|
||||
# litellm-dashboard is the default UI team and will never have agents;
|
||||
# skip noisy warnings for it.
|
||||
if user_api_key_auth.team_id != UI_TEAM_ID:
|
||||
|
|
@ -326,7 +420,9 @@ class AgentRequestHandler:
|
|||
return UnrestrictedAgentAccess()
|
||||
|
||||
@staticmethod
|
||||
def _get_config_agent_ids_for_access_groups(config_agents: list, access_groups: list[str]) -> set[str]:
|
||||
def _get_config_agent_ids_for_access_groups(
|
||||
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.
|
||||
"""
|
||||
|
|
@ -339,7 +435,9 @@ class AgentRequestHandler:
|
|||
return server_ids
|
||||
|
||||
@staticmethod
|
||||
async def _get_db_agent_ids_for_access_groups(prisma_client, access_groups: list[str]) -> set[str]:
|
||||
async def _get_db_agent_ids_for_access_groups(
|
||||
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.
|
||||
|
||||
|
|
@ -349,23 +447,27 @@ class AgentRequestHandler:
|
|||
if not access_groups or prisma_client is None:
|
||||
return set()
|
||||
|
||||
agents: Final = await AgentsRepository(prisma_client).table.find_many(
|
||||
agents: Final = await AgentsRepository(prisma_client, use_writer=check_db_only).table.find_many(
|
||||
where={"agent_access_groups": {"hasSome": access_groups}}
|
||||
)
|
||||
return {agent.agent_id for agent in agents}
|
||||
|
||||
@staticmethod
|
||||
async def _get_unified_access_group_agents(access_group_ids: list[str]) -> list[str]:
|
||||
async def _get_unified_access_group_agents(
|
||||
access_group_ids: Sequence[str], *, check_db_only: bool = False
|
||||
) -> list[str]:
|
||||
"""
|
||||
Resolve unified access group ids to agent IDs.
|
||||
"""
|
||||
from litellm.proxy.auth.auth_checks import _get_agent_ids_from_access_groups
|
||||
|
||||
return await _get_agent_ids_from_access_groups(access_group_ids=access_group_ids)
|
||||
return await _get_agent_ids_from_access_groups(access_group_ids=access_group_ids, check_db_only=check_db_only)
|
||||
|
||||
@staticmethod
|
||||
async def _get_agents_from_access_groups(
|
||||
access_groups: list[str],
|
||||
access_groups: Sequence[str],
|
||||
*,
|
||||
check_db_only: bool = False,
|
||||
) -> list[str]:
|
||||
"""
|
||||
Resolve agent access groups to agent IDs by querying BOTH the agent table (DB) AND config-loaded agents.
|
||||
|
|
@ -373,14 +475,13 @@ class AgentRequestHandler:
|
|||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
# Use the helper for config-loaded agents
|
||||
config_agent_ids: Final = AgentRequestHandler._get_config_agent_ids_for_access_groups(
|
||||
global_agent_registry.agent_list, access_groups
|
||||
)
|
||||
|
||||
# Use the helper for DB agents
|
||||
db_agent_ids: Final = await AgentRequestHandler._get_db_agent_ids_for_access_groups(
|
||||
prisma_client, access_groups
|
||||
prisma_client, access_groups, check_db_only=check_db_only
|
||||
)
|
||||
|
||||
return list(config_agent_ids | db_agent_ids)
|
||||
|
|
@ -531,4 +632,60 @@ async def accessible_agents(
|
|||
AgentRequestHandler.resolve_agent_access if resolve_access is None else resolve_access,
|
||||
effective_contexts,
|
||||
)
|
||||
return tuple(agent for agent in agents if agent.agent_id in allowed_agent_ids)
|
||||
allowed: Final = await asyncio.gather(
|
||||
*(
|
||||
AgentRequestHandler.is_agent_allowed(agent.agent_id, user_api_key_auth)
|
||||
for agent in agents
|
||||
if agent.identity_managed
|
||||
)
|
||||
)
|
||||
managed_ids: Final = frozenset(
|
||||
agent.agent_id
|
||||
for agent, permitted in zip((agent for agent in agents if agent.identity_managed), allowed)
|
||||
if permitted
|
||||
)
|
||||
return tuple(
|
||||
agent
|
||||
for agent in agents
|
||||
if (agent.agent_id in managed_ids if agent.identity_managed else agent.agent_id in allowed_agent_ids)
|
||||
)
|
||||
|
||||
|
||||
async def _strict_agent_access(auth: UserAPIKeyAuth) -> AgentAccess:
|
||||
return await AgentRequestHandler.resolve_agent_access(auth, strict=True)
|
||||
|
||||
|
||||
async def _managed_actor_agent_access(auth: UserAPIKeyAuth) -> AgentAccess:
|
||||
agent: Final = managed_agent_policy(auth)
|
||||
if agent is None or not agent.object_permission:
|
||||
return RestrictedAgentAccess(frozenset())
|
||||
permission: Final = LiteLLM_ObjectPermissionTable.model_validate(agent.object_permission or MappingProxyType({}))
|
||||
own_auth: Final = UserAPIKeyAuth(object_permission=permission)
|
||||
own: Final = _granted_ids(await AgentRequestHandler.get_allowed_agents_for_key(own_auth, strict=True))
|
||||
|
||||
from litellm.proxy.agent_endpoints.auth.agent_access_groups import resolve_managed_agent_ceilings
|
||||
|
||||
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, 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":
|
||||
return RestrictedAgentAccess(capped)
|
||||
if context.user_id is None:
|
||||
return RestrictedAgentAccess(frozenset())
|
||||
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, 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, 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))
|
||||
|
|
|
|||
84
litellm/proxy/agent_endpoints/auth/managed_authorization.py
Normal file
84
litellm/proxy/agent_endpoints/auth/managed_authorization.py
Normal file
|
|
@ -0,0 +1,84 @@
|
|||
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
|
||||
|
||||
|
||||
def managed_agent_policy(auth: "UserAPIKeyAuth | None") -> AgentResponse | None:
|
||||
"""The admitted managed policy, or ``None`` when the subject was never admitted as a managed agent.
|
||||
|
||||
``admit_managed_actor`` only assigns ``managed_agent_policy`` after ``actor_admission_failure``
|
||||
has verified the bound context, so an ``AgentResponse`` here means admission succeeded.
|
||||
"""
|
||||
policy: Final = auth.managed_agent_policy if auth is not None else None
|
||||
return policy if isinstance(policy, AgentResponse) else None
|
||||
|
||||
|
||||
async def admit_managed_actor(auth: UserAPIKeyAuth, store: AgentIdentityStore | None) -> None:
|
||||
delegation_verified: Final = auth._managed_delegation_verified # pyright: ignore[reportPrivateUsage] # the one-shot delegation marker is a PrivateAttr by design
|
||||
auth._managed_delegation_verified = False # pyright: ignore[reportPrivateUsage] # consumed here so a replayed token cannot reuse it
|
||||
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
|
||||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
@ -77,6 +78,7 @@ from litellm.proxy.agent_endpoints.auth.agent_caller import (
|
|||
load_agent_caller_team,
|
||||
load_agent_caller_user,
|
||||
)
|
||||
from litellm.proxy.agent_endpoints.auth.managed_authorization import managed_agent_policy
|
||||
from litellm.proxy.auth.budget_throttle import (
|
||||
budget_throttle_percentage,
|
||||
should_throttle_budget_exceeded,
|
||||
|
|
@ -1057,6 +1059,20 @@ async def common_checks(
|
|||
code=status.HTTP_400_BAD_REQUEST,
|
||||
)
|
||||
|
||||
managed_policy: Final = managed_agent_policy(valid_token)
|
||||
if _model and valid_token is not None and managed_policy is not None:
|
||||
managed_models: Final = (managed_policy.object_permission or MappingProxyType({})).get("models", ())
|
||||
if not isinstance(managed_models, (list, tuple)) or not managed_models:
|
||||
raise HTTPException(403, "This agent has no model grants")
|
||||
_can_object_call_model(
|
||||
model=_resolve_team_alias(_model, valid_token.team_model_aliases, valid_token.team_id, llm_router),
|
||||
llm_router=llm_router,
|
||||
models=list(managed_models),
|
||||
team_id=valid_token.team_id,
|
||||
object_type="agent",
|
||||
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
|
||||
)
|
||||
|
||||
await _check_agent_access_group_model_access(model=_model, valid_token=valid_token, llm_router=llm_router)
|
||||
await _check_agent_caller_model_access(
|
||||
model=_model,
|
||||
|
|
@ -2642,7 +2658,7 @@ async def get_user_object(
|
|||
)
|
||||
|
||||
if should_check_db:
|
||||
response = await _user_table(UserRepository(prisma_client)).find_unique(
|
||||
response = await _user_table(UserRepository(prisma_client, use_writer=bool(check_db_only))).find_unique(
|
||||
where={"user_id": user_id}, include={"organization_memberships": True}
|
||||
)
|
||||
|
||||
|
|
@ -2680,7 +2696,7 @@ async def get_user_object(
|
|||
budget_duration=new_user_params["budget_duration"]
|
||||
)
|
||||
|
||||
response = await _user_table(UserRepository(prisma_client)).create(
|
||||
response = await _user_table(UserRepository(prisma_client, use_writer=bool(check_db_only))).create(
|
||||
data=new_user_params,
|
||||
include={"organization_memberships": True},
|
||||
)
|
||||
|
|
@ -3126,9 +3142,9 @@ class TeamNotFoundError(HTTPException):
|
|||
|
||||
@log_db_metrics
|
||||
async def _get_team_db_check(
|
||||
team_id: str, prisma_client: PrismaClient, team_id_upsert: bool | None = None
|
||||
team_id: str, prisma_client: PrismaClient, team_id_upsert: bool | None = None, *, use_writer: bool = False
|
||||
) -> "_PrismaTeamRow | None":
|
||||
response = await _team_table(TeamRepository(prisma_client)).find_unique(
|
||||
response = await _team_table(TeamRepository(prisma_client, use_writer=use_writer)).find_unique(
|
||||
where={"team_id": team_id}, include=_TEAM_GRANT_RELATIONS
|
||||
)
|
||||
|
||||
|
|
@ -3162,6 +3178,7 @@ async def _get_team_object_from_user_api_key_cache(
|
|||
proxy_logging_obj: ProxyLogging | None,
|
||||
key: str,
|
||||
team_id_upsert: bool | None = None,
|
||||
use_writer: bool = False,
|
||||
) -> LiteLLM_TeamTableCachedObj:
|
||||
db_access_time_key: Final = key
|
||||
should_check_db: Final = _should_check_db(
|
||||
|
|
@ -3170,7 +3187,9 @@ async def _get_team_object_from_user_api_key_cache(
|
|||
db_cache_expiry=db_cache_expiry,
|
||||
)
|
||||
if should_check_db:
|
||||
response = await _get_team_db_check(team_id=team_id, prisma_client=prisma_client, team_id_upsert=team_id_upsert)
|
||||
response = await _get_team_db_check(
|
||||
team_id=team_id, prisma_client=prisma_client, team_id_upsert=team_id_upsert, use_writer=use_writer
|
||||
)
|
||||
# The database answered and the row is not there. Distinct from every
|
||||
# other failure here, which leaves the team's grant unknown.
|
||||
if response is None:
|
||||
|
|
@ -3192,8 +3211,11 @@ async def _get_team_object_from_user_api_key_cache(
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=use_writer,
|
||||
)
|
||||
except Exception as e:
|
||||
if use_writer:
|
||||
raise
|
||||
verbose_proxy_logger.debug(
|
||||
"Failed to load object_permission for team %s with object_permission_id=%s: %s",
|
||||
team_id,
|
||||
|
|
@ -3283,6 +3305,7 @@ async def get_team_object(
|
|||
db_cache_expiry=db_cache_expiry,
|
||||
key=key,
|
||||
team_id_upsert=team_id_upsert,
|
||||
use_writer=bool(check_db_only),
|
||||
)
|
||||
except TeamNotFoundError:
|
||||
raise
|
||||
|
|
@ -3328,16 +3351,15 @@ async def get_access_object(
|
|||
prisma_client: DatabaseClient | None,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
proxy_logging_obj: ProxyLogging | None = None,
|
||||
*,
|
||||
check_db_only: bool = False,
|
||||
) -> LiteLLM_AccessGroupTable:
|
||||
"""
|
||||
- Check if access_group_id in proxy AccessGroupTable
|
||||
- Always checks cache first, then DB only when not found in cache
|
||||
- Checks cache first unless authoritative writer admission is requested
|
||||
- if valid, return LiteLLM_AccessGroupTable object
|
||||
- if not, then raise an error
|
||||
|
||||
Unlike get_team_object, this has no check_cache_only or check_db_only flags;
|
||||
it always follows cache-first-then-db semantics.
|
||||
|
||||
Raises:
|
||||
- HTTPException: If access group doesn't exist in db or cache (status_code=404)
|
||||
"""
|
||||
|
|
@ -3346,18 +3368,19 @@ async def get_access_object(
|
|||
|
||||
key: Final = f"access_group_id:{access_group_id}"
|
||||
|
||||
cached_access_obj: Final = await user_api_key_cache.async_get_cache(
|
||||
key=key,
|
||||
model_type=LiteLLM_AccessGroupTable,
|
||||
cached_access_obj: Final = (
|
||||
None
|
||||
if check_db_only
|
||||
else await user_api_key_cache.async_get_cache(key=key, model_type=LiteLLM_AccessGroupTable)
|
||||
)
|
||||
if cached_access_obj is not None:
|
||||
return cached_access_obj
|
||||
|
||||
# Not in cache - fetch from DB
|
||||
try:
|
||||
response: Final = await _dictable_table(AccessGroupRepository(prisma_client), "access_group").find_unique(
|
||||
where={"access_group_id": access_group_id}
|
||||
)
|
||||
response: Final = await _dictable_table(
|
||||
AccessGroupRepository(prisma_client, use_writer=check_db_only), "access_group"
|
||||
).find_unique(where={"access_group_id": access_group_id})
|
||||
|
||||
if response is None:
|
||||
raise HTTPException(
|
||||
|
|
@ -3384,8 +3407,12 @@ async def get_access_object(
|
|||
access_group_id,
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": f"Access group doesn't exist in db. Access group={access_group_id}. Error: {e}"},
|
||||
status_code=503 if check_db_only else 404,
|
||||
detail=(
|
||||
"Access group policy is unavailable"
|
||||
if check_db_only
|
||||
else {"error": f"Access group doesn't exist in db. Access group={access_group_id}. Error: {e}"}
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -3719,6 +3746,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.
|
||||
|
|
@ -3732,6 +3761,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,
|
||||
|
|
@ -3743,10 +3773,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,
|
||||
|
|
@ -3768,7 +3801,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,
|
||||
|
|
@ -3856,6 +3889,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
|
||||
|
|
@ -3870,9 +3905,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)
|
||||
|
|
@ -3886,6 +3920,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:
|
||||
|
|
@ -3899,7 +3934,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,
|
||||
|
|
@ -3907,14 +3942,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,
|
||||
|
|
@ -3944,6 +3985,7 @@ async def get_object_permission(
|
|||
user_api_key_cache: UserApiKeyCache,
|
||||
parent_otel_span: Span | None = None,
|
||||
proxy_logging_obj: ProxyLogging | None = None,
|
||||
check_db_only: bool = False,
|
||||
) -> LiteLLM_ObjectPermissionTable | None:
|
||||
"""
|
||||
- Check if object permission id in proxy ObjectPermissionTable
|
||||
|
|
@ -3955,9 +3997,13 @@ async def get_object_permission(
|
|||
|
||||
# check if in cache
|
||||
key: Final = object_permission_cache_key(object_permission_id)
|
||||
deserialized_perm: Final = await user_api_key_cache.async_get_cache(
|
||||
key=key,
|
||||
model_type=LiteLLM_ObjectPermissionTable,
|
||||
deserialized_perm: Final = (
|
||||
None
|
||||
if check_db_only
|
||||
else await user_api_key_cache.async_get_cache(
|
||||
key=key,
|
||||
model_type=LiteLLM_ObjectPermissionTable,
|
||||
)
|
||||
)
|
||||
if deserialized_perm is not None:
|
||||
return deserialized_perm
|
||||
|
|
@ -3965,10 +4011,12 @@ async def get_object_permission(
|
|||
# else, check db
|
||||
try:
|
||||
response: Final = await _dictable_table(
|
||||
ObjectPermissionRepository(prisma_client), "object_permission"
|
||||
ObjectPermissionRepository(prisma_client, use_writer=check_db_only), "object_permission"
|
||||
).find_unique(where={"object_permission_id": object_permission_id})
|
||||
|
||||
if response is None:
|
||||
if check_db_only:
|
||||
raise HTTPException(status_code=403, detail="Referenced object permission does not exist")
|
||||
return None
|
||||
|
||||
_perm_obj: Final = LiteLLM_ObjectPermissionTable.model_validate(response.dict())
|
||||
|
|
@ -3981,6 +4029,8 @@ async def get_object_permission(
|
|||
|
||||
return _perm_obj
|
||||
except Exception:
|
||||
if check_db_only:
|
||||
raise
|
||||
return None
|
||||
|
||||
|
||||
|
|
@ -4190,6 +4240,7 @@ async def _get_resources_from_access_groups(
|
|||
prisma_client: DatabaseClient | None = None,
|
||||
user_api_key_cache: UserApiKeyCache | None = None,
|
||||
proxy_logging_obj: ProxyLogging | None = None,
|
||||
check_db_only: bool = False,
|
||||
) -> list[str]:
|
||||
"""
|
||||
Fetch access groups by their IDs (from cache or DB) and collect
|
||||
|
|
@ -4232,9 +4283,12 @@ async def _get_resources_from_access_groups(
|
|||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=check_db_only,
|
||||
)
|
||||
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,
|
||||
|
|
@ -4267,6 +4321,7 @@ async def _get_mcp_server_ids_from_access_groups(
|
|||
prisma_client: PrismaClient | None = None,
|
||||
user_api_key_cache: UserApiKeyCache | None = None,
|
||||
proxy_logging_obj: ProxyLogging | None = None,
|
||||
check_db_only: bool = False,
|
||||
) -> list[str]:
|
||||
"""
|
||||
Collect MCP server IDs from unified access groups.
|
||||
|
|
@ -4278,6 +4333,7 @@ async def _get_mcp_server_ids_from_access_groups(
|
|||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=check_db_only,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -4286,6 +4342,7 @@ async def _get_agent_ids_from_access_groups(
|
|||
prisma_client: PrismaClient | None = None,
|
||||
user_api_key_cache: UserApiKeyCache | None = None,
|
||||
proxy_logging_obj: ProxyLogging | None = None,
|
||||
check_db_only: bool = False,
|
||||
) -> list[str]:
|
||||
"""
|
||||
Collect agent IDs from unified access groups.
|
||||
|
|
@ -4297,6 +4354,7 @@ async def _get_agent_ids_from_access_groups(
|
|||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=check_db_only,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -4496,26 +4554,37 @@ async def _check_agent_access_group_model_access(
|
|||
"""Attached groups naming no model deny every model; the empty allowlist in ``_can_object_call_model`` allows."""
|
||||
if not model or valid_token is None or not valid_token.agent_id:
|
||||
return True
|
||||
ceiling: Final = await resolve_ceiling(valid_token.agent_id)
|
||||
if ceiling is None:
|
||||
return True
|
||||
if not ceiling.models:
|
||||
raise ModelAccessDeniedProxyException(
|
||||
message=model_access_denied_client_message(model=model),
|
||||
internal_message=f"agent {valid_token.agent_id} access groups {ceiling.access_group_ids} grant no models",
|
||||
type=ProxyErrorTypes.agent_model_access_denied,
|
||||
param="model",
|
||||
code=status.HTTP_403_FORBIDDEN,
|
||||
)
|
||||
dispatched: Final = _resolve_team_alias(model, valid_token.team_model_aliases, valid_token.team_id, llm_router)
|
||||
return _can_object_call_model(
|
||||
model=dispatched,
|
||||
llm_router=llm_router,
|
||||
models=sorted(ceiling.models),
|
||||
team_id=valid_token.team_id,
|
||||
object_type="agent",
|
||||
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
|
||||
|
||||
from litellm.proxy.agent_endpoints.auth.agent_access_groups import resolve_managed_agent_ceilings
|
||||
|
||||
managed: Final = managed_agent_policy(valid_token)
|
||||
unmanaged: Final = await resolve_ceiling(valid_token.agent_id) if managed is None else None
|
||||
ceilings: Final = (
|
||||
await resolve_managed_agent_ceilings(managed)
|
||||
if managed is not None
|
||||
else (unmanaged,)
|
||||
if unmanaged is not None
|
||||
else ()
|
||||
)
|
||||
dispatched: Final = _resolve_team_alias(model, valid_token.team_model_aliases, valid_token.team_id, llm_router)
|
||||
for ceiling in ceilings:
|
||||
if not ceiling.models:
|
||||
raise ModelAccessDeniedProxyException(
|
||||
message=model_access_denied_client_message(model=model),
|
||||
internal_message=f"agent {valid_token.agent_id} access groups grant no models",
|
||||
type=ProxyErrorTypes.agent_model_access_denied,
|
||||
param="model",
|
||||
code=status.HTTP_403_FORBIDDEN,
|
||||
)
|
||||
_can_object_call_model(
|
||||
model=dispatched,
|
||||
llm_router=llm_router,
|
||||
models=sorted(ceiling.models),
|
||||
team_id=valid_token.team_id,
|
||||
object_type="agent",
|
||||
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
LoadedCallerTeam: TypeAlias = LiteLLM_TeamTable | None
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -140,7 +140,7 @@ def encrypt_value_helper(value: str, new_encryption_key: str | None = None):
|
|||
return encrypted_value
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
"Invalid value type passed to encrypt_value: %s. Value must be a string", type(value)
|
||||
"Invalid value type passed to encrypt_value: %s for Value: %s\n Value must be a string", type(value), value
|
||||
)
|
||||
# if it's not a string - do not encrypt it and return the value
|
||||
return value
|
||||
|
|
|
|||
|
|
@ -28,12 +28,15 @@ from litellm.integrations.custom_guardrail import (
|
|||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import (
|
||||
effective_scan_only_tool_results_for_guardrail,
|
||||
effective_skip_system_message_for_guardrail,
|
||||
role_out_of_guardrail_scope,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.llms.openai.responses.guardrail_translation.handler import scannable_instructions
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
add_guardrail_scan_id,
|
||||
|
|
@ -105,9 +108,14 @@ class _ResponsesInputItem(BaseModel):
|
|||
model_config = ConfigDict(extra="ignore")
|
||||
|
||||
type: str | None = None
|
||||
role: str | None = None
|
||||
content: str | tuple[_ResponsesContentPart, ...] | None = None
|
||||
|
||||
def text_count(self) -> int:
|
||||
def text_count(self, *, skip_system: bool) -> int:
|
||||
if role_out_of_guardrail_scope(
|
||||
(self.role or "").lower(), skip_system_message=skip_system, skip_tool_message=False
|
||||
):
|
||||
return 0
|
||||
if isinstance(self.content, str):
|
||||
return 1
|
||||
if self.content is None:
|
||||
|
|
@ -1636,10 +1644,10 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
|
||||
A message's texts are consumed only when they sit at the running position of
|
||||
``texts``; messages the translation handler added without a counterpart in
|
||||
``texts`` (Responses ``instructions``, ``function_call_output``, ``reasoning``)
|
||||
are skipped. The walk runs front-to-back and back-to-front and both must agree,
|
||||
so an added message whose text happens to equal a neighbouring real message's
|
||||
text cannot steal that text's attribution. Returns None otherwise.
|
||||
``texts`` (Responses ``function_call_output``, ``reasoning``) are skipped. The walk
|
||||
runs front-to-back and back-to-front and both must agree, so an added message whose
|
||||
text happens to equal a neighbouring real message's text cannot steal that text's
|
||||
attribution. Returns None otherwise.
|
||||
"""
|
||||
runs: Final = tuple(cls._message_texts(message) for message in messages)
|
||||
|
||||
|
|
@ -1660,17 +1668,19 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
)
|
||||
return forward if len(forward) == len(texts) and forward == backward else None
|
||||
|
||||
@classmethod
|
||||
@staticmethod
|
||||
def _reasoning_item_text_indices(
|
||||
cls,
|
||||
texts: Sequence[str],
|
||||
request_data: Mapping[str, object],
|
||||
*,
|
||||
skip_system: bool,
|
||||
) -> frozenset[int] | None:
|
||||
"""Return the ``texts`` indices flattened from Responses ``reasoning`` input items.
|
||||
|
||||
The Responses translation handler gives those model-authored items the default
|
||||
``user`` role, so the latest-turn selection must not mistake one for a human turn.
|
||||
Empty for requests without a Responses ``input`` item list; None when the raw items
|
||||
(after the leading ``instructions`` text, both minus whatever ``skip_system`` drops)
|
||||
do not account for every entry of ``texts``.
|
||||
"""
|
||||
try:
|
||||
|
|
@ -1679,10 +1689,11 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
return None
|
||||
if not isinstance(raw_input, tuple):
|
||||
return frozenset()
|
||||
counts: Final = tuple(item.text_count() for item in raw_input)
|
||||
if sum(counts) != len(texts):
|
||||
offset: Final = 0 if scannable_instructions(request_data, skip_system=skip_system) is None else 1
|
||||
counts: Final = tuple(item.text_count(skip_system=skip_system) for item in raw_input)
|
||||
if offset + sum(counts) != len(texts):
|
||||
return None
|
||||
starts: Final = itertools.accumulate(counts, initial=0)
|
||||
starts: Final = itertools.accumulate(counts, initial=offset)
|
||||
return frozenset(
|
||||
text_idx
|
||||
for item, count, start in zip(raw_input, counts, starts)
|
||||
|
|
@ -1690,9 +1701,8 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
for text_idx in range(start, start + count)
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _get_latest_user_text_indices(
|
||||
cls,
|
||||
self,
|
||||
texts: Sequence[str],
|
||||
messages: Sequence[AllMessageValues],
|
||||
request_data: Mapping[str, object],
|
||||
|
|
@ -1706,10 +1716,12 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
user/developer message exists, or the latest one carries text that never reached
|
||||
``texts`` (safety fallback to the role-filter scan).
|
||||
"""
|
||||
sources: Final = cls._text_source_message_indices(texts, messages)
|
||||
sources: Final = self._text_source_message_indices(texts, messages)
|
||||
if sources is None:
|
||||
return None
|
||||
reasoning: Final = cls._reasoning_item_text_indices(texts, request_data)
|
||||
reasoning: Final = self._reasoning_item_text_indices(
|
||||
texts, request_data, skip_system=effective_skip_system_message_for_guardrail(self)
|
||||
)
|
||||
if reasoning is None:
|
||||
return None
|
||||
reasoning_messages: Final = frozenset(sources[text_idx] for text_idx in reasoning)
|
||||
|
|
@ -1723,7 +1735,7 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
)
|
||||
if latest_human is None:
|
||||
return None
|
||||
if latest_human not in sources and cls._message_texts(messages[latest_human]):
|
||||
if latest_human not in sources and self._message_texts(messages[latest_human]):
|
||||
return None
|
||||
return frozenset(text_idx for text_idx, source in enumerate(sources) if source == latest_human)
|
||||
|
||||
|
|
|
|||
|
|
@ -2290,7 +2290,7 @@ def _upstream_close_to_relay(task_results: Iterable[object]) -> Close | None:
|
|||
return upstream_close
|
||||
|
||||
|
||||
_WEBSOCKET_FORWARDED_HEADERS: Final = frozenset(("authorization", "x-api-key", "x-goog-user-project"))
|
||||
_WEBSOCKET_FORWARDED_HEADERS: Final = frozenset(("x-goog-user-project",))
|
||||
|
||||
|
||||
def _with_trace_context(headers: Mapping[str, str], parent_span: object) -> dict[str, str]:
|
||||
|
|
|
|||
|
|
@ -44,6 +44,7 @@ from litellm.litellm_core_utils.litellm_logging import (
|
|||
)
|
||||
from litellm.litellm_core_utils.ptu_pricing import azure_spillover
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps, strip_null_bytes
|
||||
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
|
||||
from litellm.proxy._types import SpendLogsMetadata, SpendLogsPayload, SpendLogsRouterMetadata
|
||||
from litellm.proxy.route_llm_request import ProxyModelNotFoundError
|
||||
from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error
|
||||
|
|
@ -1083,6 +1084,11 @@ def _get_messages_for_spend_logs_payload(
|
|||
|
||||
|
||||
_SENSITIVE_REQUEST_BODY_KEYS: Final = frozenset({"secret_fields"})
|
||||
_REQUEST_BODY_CREDENTIAL_MASKER: Final = SensitiveDataMasker(extra_sensitive_patterns=frozenset({"apikey"}))
|
||||
|
||||
|
||||
def _is_request_body_credential(key: str, value: object) -> bool:
|
||||
return isinstance(value, str) and _REQUEST_BODY_CREDENTIAL_MASKER.is_sensitive_key(key)
|
||||
|
||||
|
||||
def _sanitize_request_body_for_spend_logs_payload(
|
||||
|
|
@ -1094,8 +1100,9 @@ def _sanitize_request_body_for_spend_logs_payload(
|
|||
Recursively sanitize request body to prevent logging large base64 strings or other large values.
|
||||
Truncates strings longer than MAX_STRING_LENGTH_PROMPT_IN_DB characters and handles nested dictionaries.
|
||||
|
||||
Also strips keys listed in _SENSITIVE_REQUEST_BODY_KEYS (e.g. secret_fields
|
||||
which contains raw HTTP headers including Authorization tokens).
|
||||
At every nesting level, also strips keys listed in _SENSITIVE_REQUEST_BODY_KEYS (e.g. secret_fields,
|
||||
which holds raw HTTP headers including Authorization tokens), and replaces string values under keys
|
||||
SensitiveDataMasker classifies as credentials with REDACTED_BY_LITELM_STRING.
|
||||
"""
|
||||
from litellm.constants import (
|
||||
LITELLM_TRUNCATED_PAYLOAD_FIELD,
|
||||
|
|
@ -1152,7 +1159,11 @@ def _sanitize_request_body_for_spend_logs_payload(
|
|||
return value
|
||||
return value
|
||||
|
||||
return {k: _sanitize_value(v) for k, v in request_body.items() if k not in _SENSITIVE_REQUEST_BODY_KEYS}
|
||||
return {
|
||||
k: REDACTED_BY_LITELM_STRING if _is_request_body_credential(k, v) else _sanitize_value(v)
|
||||
for k, v in request_body.items()
|
||||
if k not in _SENSITIVE_REQUEST_BODY_KEYS
|
||||
}
|
||||
|
||||
|
||||
# Quoted-key form: ``"input"`` / ``'messages'`` / ``"prompt"`` followed by
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -15,9 +15,14 @@ if TYPE_CHECKING:
|
|||
class ObjectPermissionRepository(BaseRepository[LiteLLM_ObjectPermissionTable]):
|
||||
"""Repository for object permission database operations."""
|
||||
|
||||
def __init__(self, prisma_client: object, *, use_writer: bool = False) -> None:
|
||||
super().__init__(prisma_client)
|
||||
self._use_writer = use_writer
|
||||
|
||||
@property
|
||||
def table(self) -> TableActions["prisma_models.LiteLLM_ObjectPermissionTable"]:
|
||||
return self.prisma_client.db.litellm_objectpermissiontable
|
||||
database: Final = self.prisma_client.writer_db if self._use_writer else self.prisma_client.db
|
||||
return database.litellm_objectpermissiontable
|
||||
|
||||
@property
|
||||
def model_class(self) -> type[LiteLLM_ObjectPermissionTable]:
|
||||
|
|
|
|||
|
|
@ -70,6 +70,9 @@ class _PrismaClientView(Protocol):
|
|||
@property
|
||||
def db(self) -> _PrismaTeamDb: ...
|
||||
|
||||
@property
|
||||
def writer_db(self) -> _PrismaTeamDb: ...
|
||||
|
||||
|
||||
_MEMBERS_WITH_ROLES_ADAPTER: Final = TypeAdapter(list[Member])
|
||||
_JSON_ENCODED_TEAM_FIELDS: Final = (
|
||||
|
|
@ -85,10 +88,14 @@ _JSON_ENCODED_TEAM_FIELDS: Final = (
|
|||
class TeamRepository(BaseRepository[LiteLLM_TeamTable]):
|
||||
"""Repository for team database operations."""
|
||||
|
||||
def __init__(self, prisma_client: object, *, use_writer: bool = False) -> None:
|
||||
super().__init__(prisma_client)
|
||||
self._use_writer = use_writer
|
||||
|
||||
@property
|
||||
def _db(self) -> _PrismaTeamDb:
|
||||
client: Final[_PrismaClientView] = self.prisma_client
|
||||
return client.db
|
||||
return client.writer_db if self._use_writer else client.db
|
||||
|
||||
@property
|
||||
def table(self) -> TableActions["prisma_models.LiteLLM_TeamTable"]:
|
||||
|
|
|
|||
|
|
@ -38,9 +38,14 @@ _PLACEHOLDER_ROWS_ADAPTER: Final = TypeAdapter(tuple[SCIMPlaceholder, ...])
|
|||
class UserRepository(BaseRepository[LiteLLM_UserTable]):
|
||||
"""Repository for user database operations."""
|
||||
|
||||
def __init__(self, prisma_client: object, *, use_writer: bool = False) -> None:
|
||||
super().__init__(prisma_client)
|
||||
self._use_writer = use_writer
|
||||
|
||||
@property
|
||||
def table(self) -> TableActions["prisma_models.LiteLLM_UserTable"]:
|
||||
return self.prisma_client.db.litellm_usertable
|
||||
database: Final = self.prisma_client.writer_db if self._use_writer else self.prisma_client.db
|
||||
return database.litellm_usertable
|
||||
|
||||
@property
|
||||
def model_class(self) -> type[LiteLLM_UserTable]:
|
||||
|
|
|
|||
|
|
@ -343,6 +343,505 @@ def test_panw_latest_role_message_only_scans_only_latest_turn_on_responses_input
|
|||
assert json.loads(upstream.drain()[0].body)["input"] == shape["input"]
|
||||
|
||||
|
||||
def test_panw_scans_and_masks_top_level_instructions_on_responses_input(gateway: Gateway, tmp_path: Path) -> None:
|
||||
identity: Final = "guardrail" + uuid.uuid4().hex
|
||||
ssn: Final = "123-45-6789"
|
||||
instructions: Final = "Never repeat the SSN " + ssn + " back " + uuid.uuid4().hex
|
||||
latest: Final = "latest turn " + uuid.uuid4().hex
|
||||
shapes: Final = {
|
||||
"list_input": ([{"role": "user", "content": "first turn"}, {"role": "user", "content": latest}], "first turn"),
|
||||
"string_input": (latest, None),
|
||||
}
|
||||
|
||||
def scanner(request: Request) -> Reply:
|
||||
assert request.target == "/v1/scan/sync/request"
|
||||
body: Final = json.loads(request.body)
|
||||
prompt: Final = body["contents"][0]["prompt"]
|
||||
masked: Final = {"prompt_masked_data": {"data": prompt.replace(ssn, "<US_SSN>")}} if ssn in prompt else {}
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"action": "allow",
|
||||
"category": "dlp" if masked else "benign",
|
||||
"profile_name": "synthetic-profile",
|
||||
"report_id": "R" + body["tr_id"],
|
||||
"scan_id": "S" + body["tr_id"],
|
||||
"tr_id": body["tr_id"],
|
||||
"prompt_detected": {"injection": False, "url_cats": False, "dlp": bool(masked)},
|
||||
"response_detected": {},
|
||||
**masked,
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
def provider(request: Request) -> Reply:
|
||||
assert request.target == "/v1/responses"
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"id": "resp_" + identity,
|
||||
"object": "response",
|
||||
"created_at": 1700000000,
|
||||
"status": "completed",
|
||||
"model": "gpt-4.1-mini",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_" + identity,
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "permitted response", "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {"input_tokens": 11, "output_tokens": 4, "total_tokens": 15},
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
with wire_server(scanner) as policy, wire_server(provider) as upstream:
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["guardrails"] = [
|
||||
{
|
||||
"guardrail_name": identity,
|
||||
"litellm_params": {
|
||||
"guardrail": "panw_prisma_airs",
|
||||
"mode": "pre_call",
|
||||
"default_on": True,
|
||||
"api_base": policy.url,
|
||||
"api_key": "synthetic-panw-key",
|
||||
"profile_name": "synthetic-profile",
|
||||
},
|
||||
}
|
||||
]
|
||||
path: Final = tmp_path / "panw.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario:
|
||||
model: Final = scenario.model(
|
||||
model="openai/gpt-4.1-mini", api_base=upstream.url + "/v1", api_key="synthetic-key"
|
||||
)
|
||||
for name, (shape, first_turn) in shapes.items():
|
||||
response = candidate.request(
|
||||
"POST", "/v1/responses", {"model": model, "instructions": instructions, "input": shape}
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.json()["output"][0]["content"][0]["text"] == "permitted response"
|
||||
scanned = [json.loads(scan.body)["contents"][0]["prompt"] for scan in policy.drain()]
|
||||
expected = [instructions, *([first_turn] if first_turn else []), latest]
|
||||
assert scanned == expected, f"{name}: scanned {scanned}"
|
||||
sent = json.loads(upstream.drain()[0].body)
|
||||
assert sent["instructions"] == instructions.replace(ssn, "<US_SSN>"), f"{name}: sent {sent}"
|
||||
assert sent["input"] == shape, f"{name}: sent {sent}"
|
||||
|
||||
|
||||
_SSN: Final = "123-45-6789"
|
||||
_MASKED_SSN: Final = "<US_SSN>"
|
||||
_DENIED_TERM: Final = "RIGBLOCKME"
|
||||
|
||||
|
||||
def _panw_scanner(request: Request) -> Reply:
|
||||
assert request.target == "/v1/scan/sync/request"
|
||||
body: Final = json.loads(request.body)
|
||||
prompt: Final = body["contents"][0]["prompt"]
|
||||
denied: Final = _DENIED_TERM in prompt
|
||||
masked: Final = {"prompt_masked_data": {"data": prompt.replace(_SSN, _MASKED_SSN)}} if _SSN in prompt else {}
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"action": "block" if denied else "allow",
|
||||
"category": "malicious" if denied else ("dlp" if masked else "benign"),
|
||||
"profile_name": "synthetic-profile",
|
||||
"report_id": "R" + body["tr_id"],
|
||||
"scan_id": "S" + body["tr_id"],
|
||||
"tr_id": body["tr_id"],
|
||||
"prompt_detected": {"injection": denied, "url_cats": False, "dlp": bool(masked)},
|
||||
"response_detected": {},
|
||||
**masked,
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
|
||||
def _responses_provider(request: Request) -> Reply:
|
||||
if request.method == "GET" and request.target.endswith("/models"):
|
||||
return Reply(body=json.dumps({"object": "list", "data": []}).encode())
|
||||
assert request.target == "/v1/responses", request.target
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"id": "resp_" + uuid.uuid4().hex,
|
||||
"object": "response",
|
||||
"created_at": 1700000000,
|
||||
"status": "completed",
|
||||
"model": "gpt-4.1-mini",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_synthetic",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "permitted response", "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {"input_tokens": 11, "output_tokens": 4, "total_tokens": 15},
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
|
||||
def _panw_config(tmp_path: Path, identity: str, policy_url: str, **flags: bool) -> Path:
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["guardrails"] = [
|
||||
{
|
||||
"guardrail_name": identity,
|
||||
"litellm_params": {
|
||||
"guardrail": "panw_prisma_airs",
|
||||
"mode": "pre_call",
|
||||
"default_on": True,
|
||||
"api_base": policy_url,
|
||||
"api_key": "synthetic-panw-key",
|
||||
"profile_name": "synthetic-profile",
|
||||
**flags,
|
||||
},
|
||||
}
|
||||
]
|
||||
path: Final = tmp_path / "panw.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
return path
|
||||
|
||||
|
||||
def _scanned_prompts(scans: tuple[Request, ...]) -> list[str]:
|
||||
return [json.loads(scan.body)["contents"][0]["prompt"] for scan in scans]
|
||||
|
||||
|
||||
def _forwarded_bodies(requests: tuple[Request, ...]) -> list[dict[str, object]]:
|
||||
return [json.loads(request.body) for request in requests if request.method == "POST"]
|
||||
|
||||
|
||||
def test_guardrail_denies_responses_request_whose_only_flagged_text_is_in_instructions(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
identity: Final = "guardrail" + uuid.uuid4().hex
|
||||
instructions: Final = "You are terse and say " + _DENIED_TERM + " " + uuid.uuid4().hex
|
||||
shapes: Final = {"string_input": "say hi", "list_input": [{"role": "user", "content": "say hi"}]}
|
||||
with wire_server(_panw_scanner) as policy, wire_server(_responses_provider) as upstream:
|
||||
config: Final = _panw_config(tmp_path, identity, policy.url)
|
||||
with (
|
||||
owned_proxy(gateway, tmp_path, {}, config=config, workers=2) as candidate,
|
||||
candidate.scenario() as scenario,
|
||||
):
|
||||
model: Final = scenario.model(
|
||||
model="openai/gpt-4.1-mini", api_base=upstream.url + "/v1", api_key="synthetic-key"
|
||||
)
|
||||
for name, shape in shapes.items():
|
||||
response = candidate.request(
|
||||
"POST", "/v1/responses", {"model": model, "instructions": instructions, "input": shape}
|
||||
)
|
||||
assert response.status_code == 400, f"{name}: {response.text}"
|
||||
assert "Prompt blocked by PANW Prisma AI Security policy" in response.text, response.text
|
||||
assert _scanned_prompts(policy.drain()) == [instructions], name
|
||||
assert _forwarded_bodies(upstream.drain()) == [], (
|
||||
f"{name}: denied instructions must not reach the provider"
|
||||
)
|
||||
|
||||
|
||||
def test_empty_instructions_are_not_scanned_while_input_and_chat_system_masking_are_unchanged(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
identity: Final = "guardrail" + uuid.uuid4().hex
|
||||
secret: Final = "my SSN is " + _SSN + " " + uuid.uuid4().hex
|
||||
masked: Final = secret.replace(_SSN, _MASKED_SSN)
|
||||
|
||||
def chat_provider(request: Request) -> Reply:
|
||||
assert request.target == "/v1/chat/completions"
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"id": "chatcmpl_" + identity,
|
||||
"object": "chat.completion",
|
||||
"created": 1700000000,
|
||||
"model": "gpt-4.1-mini",
|
||||
"choices": [
|
||||
{"index": 0, "finish_reason": "stop", "message": {"role": "assistant", "content": "ok"}}
|
||||
],
|
||||
"usage": {"prompt_tokens": 5, "completion_tokens": 1, "total_tokens": 6},
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
def provider(request: Request) -> Reply:
|
||||
return chat_provider(request) if request.target == "/v1/chat/completions" else _responses_provider(request)
|
||||
|
||||
with wire_server(_panw_scanner) as policy, wire_server(provider) as upstream:
|
||||
config: Final = _panw_config(tmp_path, identity, policy.url, mask_request_content=True)
|
||||
with (
|
||||
owned_proxy(gateway, tmp_path, {}, config=config, workers=2) as candidate,
|
||||
candidate.scenario() as scenario,
|
||||
):
|
||||
model: Final = scenario.model(
|
||||
model="openai/gpt-4.1-mini", api_base=upstream.url + "/v1", api_key="synthetic-key"
|
||||
)
|
||||
for instructions in ("", None):
|
||||
body = {"model": model, "input": secret, **({} if instructions is None else {"instructions": ""})}
|
||||
response = candidate.request("POST", "/v1/responses", body)
|
||||
assert response.status_code == 200, response.text
|
||||
assert _scanned_prompts(policy.drain()) == [secret], f"instructions={instructions!r}"
|
||||
(sent,) = _forwarded_bodies(upstream.drain())
|
||||
assert sent.get("instructions") == instructions, f"instructions={instructions!r}: sent {sent}"
|
||||
assert sent["input"] == masked, f"instructions={instructions!r}: sent {sent}"
|
||||
|
||||
response = candidate.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{
|
||||
"model": model,
|
||||
"messages": [{"role": "system", "content": secret}, {"role": "user", "content": "hi"}],
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert _scanned_prompts(policy.drain()) == [secret, "hi"]
|
||||
(sent_chat,) = _forwarded_bodies(upstream.drain())
|
||||
assert sent_chat["messages"] == [
|
||||
{"role": "system", "content": masked},
|
||||
{"role": "user", "content": "hi"},
|
||||
]
|
||||
|
||||
|
||||
def test_skip_system_message_leaves_instructions_and_system_items_unscanned_on_responses(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
identity: Final = "guardrail" + uuid.uuid4().hex
|
||||
instructions: Final = "Escalations go to " + _SSN + " " + uuid.uuid4().hex
|
||||
system_item: Final = "House rules: never share " + _SSN + " " + uuid.uuid4().hex
|
||||
developer_item: Final = "Developer note " + _SSN + " " + uuid.uuid4().hex
|
||||
latest: Final = "my contact is " + _SSN + " " + uuid.uuid4().hex
|
||||
|
||||
with wire_server(_panw_scanner) as policy, wire_server(_responses_provider) as upstream:
|
||||
config: Final = _panw_config(
|
||||
tmp_path, identity, policy.url, mask_request_content=True, skip_system_message_in_guardrail=True
|
||||
)
|
||||
with (
|
||||
owned_proxy(gateway, tmp_path, {}, config=config, workers=2) as candidate,
|
||||
candidate.scenario() as scenario,
|
||||
):
|
||||
model: Final = scenario.model(
|
||||
model="openai/gpt-4.1-mini", api_base=upstream.url + "/v1", api_key="synthetic-key"
|
||||
)
|
||||
response = candidate.request(
|
||||
"POST",
|
||||
"/v1/responses",
|
||||
{
|
||||
"model": model,
|
||||
"instructions": instructions,
|
||||
"input": [
|
||||
{"role": "system", "content": system_item},
|
||||
{"role": "developer", "content": developer_item},
|
||||
{"role": "user", "content": latest},
|
||||
],
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert _scanned_prompts(policy.drain()) == [developer_item, latest]
|
||||
(sent,) = _forwarded_bodies(upstream.drain())
|
||||
assert sent["instructions"] == instructions, f"sent {sent}"
|
||||
assert sent["input"] == [
|
||||
{"role": "system", "content": system_item},
|
||||
{"role": "developer", "content": developer_item.replace(_SSN, _MASKED_SSN)},
|
||||
{"role": "user", "content": latest.replace(_SSN, _MASKED_SSN)},
|
||||
], f"sent {sent}"
|
||||
|
||||
|
||||
def test_instructions_masking_lands_next_to_multimodal_and_tool_loop_input_items(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
identity: Final = "guardrail" + uuid.uuid4().hex
|
||||
instructions: Final = "Never repeat the SSN " + _SSN + " back " + uuid.uuid4().hex
|
||||
latest: Final = "latest turn with " + _SSN + " " + uuid.uuid4().hex
|
||||
image: Final = {"type": "input_image", "image_url": "https://example.test/receipt.png", "detail": "low"}
|
||||
shapes: Final = {
|
||||
"multimodal": [
|
||||
{"role": "user", "content": [{"type": "input_text", "text": "first turn"}, image]},
|
||||
{"role": "user", "content": [image, {"type": "input_text", "text": latest}]},
|
||||
],
|
||||
"tool_loop": [
|
||||
{"role": "user", "content": "first turn"},
|
||||
{"type": "function_call", "call_id": "call_1", "name": "lookup", "arguments": "{}"},
|
||||
{"type": "function_call_output", "call_id": "call_1", "output": "tool result with " + _SSN},
|
||||
{"role": "user", "content": latest},
|
||||
],
|
||||
}
|
||||
with wire_server(_panw_scanner) as policy, wire_server(_responses_provider) as upstream:
|
||||
config: Final = _panw_config(tmp_path, identity, policy.url, mask_request_content=True)
|
||||
with (
|
||||
owned_proxy(gateway, tmp_path, {}, config=config, workers=2) as candidate,
|
||||
candidate.scenario() as scenario,
|
||||
):
|
||||
model: Final = scenario.model(
|
||||
model="openai/gpt-4.1-mini", api_base=upstream.url + "/v1", api_key="synthetic-key"
|
||||
)
|
||||
for name, shape in shapes.items():
|
||||
response = candidate.request(
|
||||
"POST", "/v1/responses", {"model": model, "instructions": instructions, "input": shape}
|
||||
)
|
||||
assert response.status_code == 200, f"{name}: {response.text}"
|
||||
assert _scanned_prompts(policy.drain()) == [instructions, "first turn", latest], name
|
||||
(sent,) = _forwarded_bodies(upstream.drain())
|
||||
assert sent["instructions"] == instructions.replace(_SSN, _MASKED_SSN), f"{name}: sent {sent}"
|
||||
expected = json.loads(json.dumps(shape).replace(latest, latest.replace(_SSN, _MASKED_SSN)))
|
||||
assert sent["input"] == expected, f"{name}: sent {sent}"
|
||||
|
||||
|
||||
def test_panw_latest_only_with_instructions_masks_only_the_latest_turn(gateway: Gateway, tmp_path: Path) -> None:
|
||||
identity: Final = "guardrail" + uuid.uuid4().hex
|
||||
instructions: Final = "Keep " + _SSN + " confidential " + uuid.uuid4().hex
|
||||
latest: Final = "latest turn with " + _SSN + " " + uuid.uuid4().hex
|
||||
history: Final = ({"role": "user", "content": "first turn"}, {"role": "assistant", "content": "first reply"})
|
||||
shapes: Final = {
|
||||
"plain": [*history, {"role": "user", "content": latest}],
|
||||
"reasoning": [
|
||||
*history,
|
||||
{"type": "reasoning", "id": "rs_1", "summary": [{"type": "summary_text", "text": "thinking"}]},
|
||||
{"role": "user", "content": latest},
|
||||
],
|
||||
}
|
||||
with wire_server(_panw_scanner) as policy, wire_server(_responses_provider) as upstream:
|
||||
config: Final = _panw_config(
|
||||
tmp_path, identity, policy.url, mask_request_content=True, experimental_use_latest_role_message_only=True
|
||||
)
|
||||
with (
|
||||
owned_proxy(gateway, tmp_path, {}, config=config, workers=2) as candidate,
|
||||
candidate.scenario() as scenario,
|
||||
):
|
||||
model: Final = scenario.model(
|
||||
model="openai/gpt-4.1-mini", api_base=upstream.url + "/v1", api_key="synthetic-key"
|
||||
)
|
||||
for name, shape in shapes.items():
|
||||
response = candidate.request(
|
||||
"POST", "/v1/responses", {"model": model, "instructions": instructions, "input": shape}
|
||||
)
|
||||
assert response.status_code == 200, f"{name}: {response.text}"
|
||||
assert _scanned_prompts(policy.drain()) == [latest], name
|
||||
(sent,) = _forwarded_bodies(upstream.drain())
|
||||
assert sent["instructions"] == instructions, f"{name}: latest-only must leave instructions alone"
|
||||
assert sent["input"] == [*shape[:-1], {"role": "user", "content": latest.replace(_SSN, _MASKED_SSN)}], (
|
||||
f"{name}: sent {sent}"
|
||||
)
|
||||
|
||||
|
||||
def test_bedrock_latest_only_masks_latest_turn_on_responses_input_with_instructions(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
identity: Final = "guardrail" + uuid.uuid4().hex
|
||||
guardrail_id: Final = "synthetic" + uuid.uuid4().hex[:8]
|
||||
instructions: Final = "Keep " + _SSN + " confidential " + uuid.uuid4().hex
|
||||
latest: Final = "latest turn with " + _SSN + " " + uuid.uuid4().hex
|
||||
|
||||
def guardrail(request: Request) -> Reply:
|
||||
assert request.target == f"/guardrail/{guardrail_id}/version/DRAFT/apply", request.target
|
||||
body: Final = json.loads(request.body)
|
||||
assert body["source"] == "INPUT", body
|
||||
assert body["content"] == [{"text": {"text": latest}}], body
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"action": "GUARDRAIL_INTERVENED",
|
||||
"outputs": [{"text": latest.replace(_SSN, _MASKED_SSN)}],
|
||||
"assessments": [
|
||||
{
|
||||
"sensitiveInformationPolicy": {
|
||||
"piiEntities": [
|
||||
{"type": "US_SOCIAL_SECURITY_NUMBER", "match": _SSN, "action": "ANONYMIZED"}
|
||||
]
|
||||
}
|
||||
}
|
||||
],
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
with wire_server(guardrail) as policy, wire_server(_responses_provider) as upstream:
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["guardrails"] = [
|
||||
{
|
||||
"guardrail_name": identity,
|
||||
"litellm_params": {
|
||||
"guardrail": "bedrock",
|
||||
"mode": "pre_call",
|
||||
"default_on": True,
|
||||
"mask_request_content": True,
|
||||
"experimental_use_latest_role_message_only": True,
|
||||
"guardrailIdentifier": guardrail_id,
|
||||
"guardrailVersion": "DRAFT",
|
||||
"aws_region_name": "us-east-1",
|
||||
"aws_access_key_id": "AKIASYNTHETICGUARDRAIL",
|
||||
"aws_secret_access_key": "synthetic-secret",
|
||||
"aws_bedrock_runtime_endpoint": policy.url,
|
||||
},
|
||||
}
|
||||
]
|
||||
path: Final = tmp_path / "bedrock-instructions.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
with owned_proxy(gateway, tmp_path, {}, config=path, workers=2) as candidate, candidate.scenario() as scenario:
|
||||
model: Final = scenario.model(
|
||||
model="openai/gpt-4.1-mini", api_base=upstream.url + "/v1", api_key="synthetic-key"
|
||||
)
|
||||
response: Final = candidate.request(
|
||||
"POST",
|
||||
"/v1/responses",
|
||||
{
|
||||
"model": model,
|
||||
"instructions": instructions,
|
||||
"input": [
|
||||
{"role": "user", "content": "first turn"},
|
||||
{"role": "assistant", "content": "first reply"},
|
||||
{"role": "user", "content": latest},
|
||||
],
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert len(policy.drain()) == 1
|
||||
(sent,) = _forwarded_bodies(upstream.drain())
|
||||
assert sent["instructions"] == instructions, sent
|
||||
assert sent["input"] == [
|
||||
{"role": "user", "content": "first turn"},
|
||||
{"role": "assistant", "content": "first reply"},
|
||||
{"role": "user", "content": latest.replace(_SSN, _MASKED_SSN)},
|
||||
], sent
|
||||
|
||||
|
||||
def test_instructions_masking_holds_under_concurrent_load_across_two_workers(gateway: Gateway, tmp_path: Path) -> None:
|
||||
identity: Final = "guardrail" + uuid.uuid4().hex
|
||||
with wire_server(_panw_scanner) as policy, wire_server(_responses_provider) as upstream:
|
||||
config: Final = _panw_config(tmp_path, identity, policy.url, mask_request_content=True)
|
||||
with (
|
||||
owned_proxy(gateway, tmp_path, {}, config=config, workers=2) as candidate,
|
||||
candidate.scenario() as scenario,
|
||||
):
|
||||
model: Final = scenario.model(
|
||||
model="openai/gpt-4.1-mini", api_base=upstream.url + "/v1", api_key="synthetic-key"
|
||||
)
|
||||
tags: Final = tuple(uuid.uuid4().hex for _ in range(16))
|
||||
|
||||
def send(tag: str) -> httpx.Response:
|
||||
return candidate.request(
|
||||
"POST",
|
||||
"/v1/responses",
|
||||
{"model": model, "instructions": "Keep " + _SSN + " private " + tag, "input": "say hi " + tag},
|
||||
)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=8) as pool:
|
||||
responses: Final = tuple(pool.map(send, tags))
|
||||
assert [response.status_code for response in responses] == [200] * len(tags), [
|
||||
response.text for response in responses
|
||||
]
|
||||
sent: Final = {str(body["input"]): body for body in _forwarded_bodies(upstream.drain())}
|
||||
assert sorted(_scanned_prompts(policy.drain())) == sorted(
|
||||
[text for tag in tags for text in ("Keep " + _SSN + " private " + tag, "say hi " + tag)]
|
||||
)
|
||||
assert {tag: sent["say hi " + tag]["instructions"] for tag in tags} == {
|
||||
tag: "Keep " + _MASKED_SSN + " private " + tag for tag in tags
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.covers("other.observability.guardrails.bedrock_passthrough_converse_scans_only_caller_content")
|
||||
def test_bedrock_passthrough_converse_guardrail_ignores_denied_term_in_tool_definition(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
|
|
|
|||
|
|
@ -0,0 +1,559 @@
|
|||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._experimental.mcp_server import mcp_server_manager
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_UserTable, UserAPIKeyAuth
|
||||
from litellm.proxy.auth import auth_checks
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.agent_identity import ManagedAgentContext
|
||||
|
||||
|
||||
def actor(tools: tuple[str, ...] | None, *, delegated: bool = False) -> UserAPIKeyAuth:
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="agent-permissions",
|
||||
mcp_servers=["slack", "linear"],
|
||||
mcp_tool_permissions={"slack": list(tools)} if tools is not None else None,
|
||||
)
|
||||
agent: Final = AgentResponse(
|
||||
agent_id="publisher",
|
||||
agent_name="Publisher",
|
||||
agent_card_params={},
|
||||
object_permission=permission.model_dump(),
|
||||
identity_managed=True,
|
||||
)
|
||||
auth: Final = UserAPIKeyAuth(agent_id=agent.agent_id)
|
||||
auth.managed_agent_policy = agent
|
||||
auth.managed_agent_context = ManagedAgentContext(
|
||||
agent_id=agent.agent_id,
|
||||
mode="delegated" if delegated else "autonomous",
|
||||
user_id="human" if delegated else None,
|
||||
)
|
||||
return auth
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def isolated_manager(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", mcp_server_manager.MCPServerManager())
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", MagicMock())
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("tools", (None, (), ("read",), ("read", "write")))
|
||||
async def test_autonomous_agent_uses_only_its_own_tool_grants(tools: tuple[str, ...] | None) -> None:
|
||||
auth: Final = actor(tools)
|
||||
assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == {"slack", "linear"}
|
||||
actual: Final = await MCPRequestHandler.get_allowed_tools_for_server("slack", auth)
|
||||
assert (frozenset(actual) if actual is not None else None) == (frozenset(tools) if tools is not None else None)
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("ungranted-server", auth) == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"agent_tools,user_tools,expected",
|
||||
(
|
||||
(None, ("read",), ("read",)),
|
||||
(("read",), None, ("read",)),
|
||||
(("read", "write"), ("read",), ("read",)),
|
||||
(("read",), ("write",), ()),
|
||||
((), None, ()),
|
||||
),
|
||||
)
|
||||
async def test_delegated_server_and_tool_intersections(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
agent_tools: tuple[str, ...] | None,
|
||||
user_tools: tuple[str, ...] | None,
|
||||
expected: tuple[str, ...],
|
||||
) -> None:
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="user-permissions",
|
||||
mcp_servers=["slack", "user-only"],
|
||||
mcp_tool_permissions={"slack": list(user_tools)} if user_tools is not None else None,
|
||||
)
|
||||
user: Final = LiteLLM_UserTable(user_id="human", teams=[], object_permission=permission)
|
||||
monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(return_value=user))
|
||||
auth: Final = actor(agent_tools, delegated=True)
|
||||
assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == ["slack"]
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == list(expected)
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == []
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("user-only", auth) == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unavailable_delegated_user_never_leaves_agent_permissions_unrestricted(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(side_effect=RuntimeError("DB unavailable")))
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await MCPRequestHandler.get_allowed_tools_for_server("slack", actor(None, delegated=True))
|
||||
assert failure.value.status_code == 503
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("servers,expected", (((), ()), (("slack",), ("slack",)), (("user-only",), ())))
|
||||
async def test_access_groups_cap_agent_servers_without_granting_new_ones(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
servers: tuple[str, ...],
|
||||
expected: tuple[str, ...],
|
||||
) -> None:
|
||||
from litellm.proxy._types import LiteLLM_AccessGroupTable
|
||||
|
||||
group: Final = LiteLLM_AccessGroupTable(
|
||||
access_group_id="group", access_group_name="Restricted", access_mcp_server_ids=list(servers)
|
||||
)
|
||||
monkeypatch.setattr(auth_checks, "get_access_object", AsyncMock(return_value=group))
|
||||
auth: Final = actor(None)
|
||||
assert auth.managed_agent_policy is not None
|
||||
auth.managed_agent_policy = auth.managed_agent_policy.model_copy(update={"access_group_ids": ["group"]})
|
||||
assert tuple(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == expected
|
||||
if "slack" not in expected:
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("change", ["tools", "servers", "disabled", "outage"])
|
||||
async def test_delegated_mcp_revokes_warm_human_policy_before_tool_execution(
|
||||
monkeypatch: pytest.MonkeyPatch, change: str
|
||||
) -> None:
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key
|
||||
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="user-grant", mcp_servers=["slack"], mcp_tool_permissions={"slack": ["read", "write"]}
|
||||
)
|
||||
user: Final = LiteLLM_UserTable(
|
||||
user_id="human", teams=[], organization_memberships=[], object_permission_id="user-grant"
|
||||
)
|
||||
cache: Final = UserApiKeyCache()
|
||||
cache.set_cache("human", user)
|
||||
cache.set_cache(object_permission_cache_key("user-grant"), permission)
|
||||
client: Final = MagicMock()
|
||||
client.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=user)
|
||||
client.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=permission)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", client)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
|
||||
auth: Final = actor(("read", "write"), delegated=True)
|
||||
assert set(await MCPRequestHandler.get_allowed_tools_for_server("slack", auth)) == {"read", "write"}
|
||||
if change == "disabled":
|
||||
client.writer_db.litellm_usertable.find_unique.return_value = user.model_copy(
|
||||
update={"metadata": {"scim_active": False}}
|
||||
)
|
||||
elif change == "outage":
|
||||
client.writer_db.litellm_usertable.find_unique.side_effect = RuntimeError("writer unavailable")
|
||||
elif change == "servers":
|
||||
client.writer_db.litellm_objectpermissiontable.find_unique.return_value = permission.model_copy(
|
||||
update={"mcp_servers": [], "mcp_tool_permissions": {}}
|
||||
)
|
||||
else:
|
||||
client.writer_db.litellm_objectpermissiontable.find_unique.return_value = permission.model_copy(
|
||||
update={"mcp_tool_permissions": {"slack": ["read"]}}
|
||||
)
|
||||
if change in ("disabled", "outage"):
|
||||
with pytest.raises(HTTPException):
|
||||
await MCPRequestHandler.get_allowed_tools_for_server("slack", auth)
|
||||
else:
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == (
|
||||
["read"] if change == "tools" else []
|
||||
)
|
||||
client.db.litellm_usertable.find_unique.assert_not_called()
|
||||
client.db.litellm_objectpermissiontable.find_unique.assert_not_called()
|
||||
|
||||
|
||||
def _server_row(server_id: str, access_groups: tuple[str, ...]) -> MagicMock:
|
||||
row: Final = MagicMock()
|
||||
row.server_id = server_id
|
||||
row.mcp_access_groups = list(access_groups)
|
||||
return row
|
||||
|
||||
|
||||
def _toolset_row(server_id: str, tool_name: str) -> MagicMock:
|
||||
row: Final = MagicMock()
|
||||
row.tools = [{"server_id": server_id, "tool_name": tool_name}]
|
||||
return row
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("change", ["tool", "server", "outage"])
|
||||
async def test_autonomous_agent_toolset_and_access_group_revocations_bind_on_the_next_request(
|
||||
monkeypatch: pytest.MonkeyPatch, change: str
|
||||
) -> None:
|
||||
"""The agent's entitlements are read through the shared toolset and access-group resolvers. Once the
|
||||
writer revokes a tool or drops the server from the group, the next managed request must be denied
|
||||
even though the legacy cache still holds the warm grant and the replica still shows the old rows"""
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy._experimental.mcp_server import toolset_db
|
||||
|
||||
warm_toolset: Final = _toolset_row("slack", "read")
|
||||
list_toolsets: Final = AsyncMock(return_value=[warm_toolset])
|
||||
monkeypatch.setattr(toolset_db, "list_mcp_toolsets", list_toolsets)
|
||||
client: Final = MagicMock()
|
||||
client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[_server_row("linear", ("grp",))])
|
||||
client.writer_db.litellm_mcpservertable.find_many = AsyncMock(return_value=[_server_row("linear", ("grp",))])
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", client)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", DualCache())
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="agent-permissions", mcp_toolsets=["ts"], mcp_access_groups=["grp"]
|
||||
)
|
||||
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
|
||||
|
||||
assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == {"slack", "linear"}
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == ["read"]
|
||||
|
||||
if change == "tool":
|
||||
list_toolsets.return_value = [_toolset_row("slack", "other")]
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == ["other"]
|
||||
elif change == "server":
|
||||
client.writer_db.litellm_mcpservertable.find_many.return_value = []
|
||||
assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == {"slack"}
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("linear", auth) == []
|
||||
else:
|
||||
list_toolsets.side_effect = RuntimeError("writer unavailable")
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await MCPRequestHandler.get_allowed_tools_for_server("slack", auth)
|
||||
assert failure.value.status_code == 503
|
||||
for call in list_toolsets.await_args_list:
|
||||
assert call.kwargs["use_writer"] is True, "managed agent toolsets must be read from the writer"
|
||||
client.db.litellm_mcpservertable.find_many.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("role", ["proxy_admin", "proxy_admin_viewer", "internal_user"])
|
||||
@pytest.mark.parametrize("open_channel", ["none", "operator", "submitted"])
|
||||
@pytest.mark.parametrize("has_grant", [True, False])
|
||||
@pytest.mark.parametrize("agent_tools", [("read", "write"), None])
|
||||
async def test_delegated_mcp_uses_explicit_team_grants_even_for_dashboard_admins(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
role: str,
|
||||
open_channel: str,
|
||||
has_grant: bool,
|
||||
agent_tools: tuple[str, ...] | None,
|
||||
) -> None:
|
||||
from litellm.proxy._types import LiteLLM_TeamTable
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
manager: Final = mcp_server_manager.global_mcp_server_manager
|
||||
manager.registry = {
|
||||
name: MCPServer(
|
||||
server_id=name,
|
||||
name=name,
|
||||
transport="http",
|
||||
url="https://example.com/mcp",
|
||||
allow_all_keys=open_channel == "operator",
|
||||
)
|
||||
for name in ("slack", "linear")
|
||||
}
|
||||
from litellm.proxy._experimental.mcp_server import db
|
||||
|
||||
monkeypatch.setattr(
|
||||
db,
|
||||
"get_active_submitted_mcp_server_ids_for_user",
|
||||
AsyncMock(return_value=["slack", "linear"] if open_channel == "submitted" else []),
|
||||
)
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="team-grant", mcp_servers=["slack"], mcp_tool_permissions={"slack": ["read"]}
|
||||
)
|
||||
user: Final = LiteLLM_UserTable(
|
||||
user_id="human", user_role=role, teams=["team"] if has_grant else [], organization_memberships=[]
|
||||
)
|
||||
team: Final = LiteLLM_TeamTable(
|
||||
team_id="team",
|
||||
models=[],
|
||||
members_with_roles=[{"user_id": "human", "role": "user"}],
|
||||
object_permission_id="team-grant",
|
||||
)
|
||||
client: Final = MagicMock()
|
||||
client.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=user)
|
||||
client.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=team)
|
||||
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) == []
|
||||
admitted: Final = await MCPRequestHandler.reload_admitted_user("human", requires_fresh_policy=True)
|
||||
assert admitted.user_role == role
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_explicit_grants_never_fall_back_to_open_servers_on_resolution_failure(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from litellm.proxy._experimental.mcp_server import db
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
manager: Final = mcp_server_manager.global_mcp_server_manager
|
||||
manager.registry = {"slack": MCPServer(server_id="slack", name="slack", transport="http", allow_all_keys=True)}
|
||||
monkeypatch.setattr(db, "get_active_submitted_mcp_server_ids_for_user", AsyncMock(return_value=["slack"]))
|
||||
auth: Final = UserAPIKeyAuth(user_id="human")
|
||||
auth.mcp_explicit_grants_only = True
|
||||
with pytest.MonkeyPatch.context() as patcher:
|
||||
patcher.setattr(MCPRequestHandler, "get_mcp_server_access", AsyncMock(side_effect=RuntimeError("unavailable")))
|
||||
assert await manager.get_allowed_mcp_servers(auth) == []
|
||||
auth.mcp_explicit_grants_only = False
|
||||
assert await manager.get_allowed_mcp_servers(auth) == ["slack"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_absent_agent_policy_and_missing_delegated_subject_grant_no_servers() -> None:
|
||||
from litellm.proxy._experimental.mcp_server.auth.managed_agent_access import managed_agent_servers
|
||||
|
||||
assert await managed_agent_servers(UserAPIKeyAuth()) == ()
|
||||
auth: Final = actor(None, delegated=True)
|
||||
assert auth.managed_agent_context is not None
|
||||
auth.managed_agent_context = auth.managed_agent_context.model_copy(update={"user_id": None})
|
||||
assert await MCPRequestHandler.get_allowed_mcp_servers(auth) == []
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_policy_outage_after_server_admission_fails_closed(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="human-grant", mcp_servers=["slack"])
|
||||
user: Final = LiteLLM_UserTable(user_id="human", teams=[], object_permission=permission)
|
||||
monkeypatch.setattr(
|
||||
auth_checks, "get_user_object", AsyncMock(side_effect=[user, RuntimeError("tool lookup unavailable")])
|
||||
)
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await MCPRequestHandler.get_allowed_tools_for_server("slack", actor(None, delegated=True))
|
||||
assert failure.value.status_code == 503
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("role", (None, "proxy_admin", "internal_user"))
|
||||
@pytest.mark.parametrize("scoped", (False, True))
|
||||
async def test_manager_preserves_managed_server_grants_across_open_channels(
|
||||
monkeypatch: pytest.MonkeyPatch, role: str | None, scoped: bool
|
||||
) -> None:
|
||||
from litellm.proxy._experimental.mcp_server import db
|
||||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPServerAccess
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
manager: Final = mcp_server_manager.global_mcp_server_manager
|
||||
manager.registry = {
|
||||
"open": MCPServer(server_id="open", name="open", transport="http", allow_all_keys=True),
|
||||
"submitted": MCPServer(server_id="submitted", name="submitted", transport="http"),
|
||||
"passthrough": MCPServer(
|
||||
server_id="passthrough", name="passthrough", transport="http", auth_type="true_passthrough"
|
||||
),
|
||||
}
|
||||
monkeypatch.setattr(db, "get_active_submitted_mcp_server_ids_for_user", AsyncMock(return_value=["submitted"]))
|
||||
auth: Final = actor(None)
|
||||
auth.user_role = role
|
||||
assert not auth.mcp_explicit_grants_only
|
||||
access: Final = MCPServerAccess(server_ids=("slack", "open")) if scoped else None
|
||||
assert set(await manager.get_allowed_mcp_servers(auth, access=access)) == (
|
||||
{"slack"} if scoped else {"slack", "linear"}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_manager_does_not_replace_managed_policy_failure_with_open_servers(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
manager: Final = mcp_server_manager.global_mcp_server_manager
|
||||
manager.registry = {"open": MCPServer(server_id="open", name="open", transport="http", allow_all_keys=True)}
|
||||
monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(side_effect=RuntimeError("writer unavailable")))
|
||||
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()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_managed_agent_mcp_access_is_capped_at_the_invoking_callers_grants(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""The managed MCP path must honour the agent_caller ceiling the same way the unmanaged path does:
|
||||
the agent's own policy grants slack and linear, but the team echoed back on the request reaches
|
||||
only slack, so the agent may use slack alone."""
|
||||
from litellm.proxy._types import AgentCaller
|
||||
|
||||
monkeypatch.setattr(
|
||||
MCPRequestHandler,
|
||||
"_get_allowed_mcp_servers_for_team",
|
||||
AsyncMock(return_value=["slack"]),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
MCPRequestHandler,
|
||||
"_apply_user_server_ceiling",
|
||||
AsyncMock(side_effect=lambda servers, _auth: (tuple(servers), False)),
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
MCPRequestHandler,
|
||||
"_get_team_object_permission",
|
||||
AsyncMock(
|
||||
return_value=LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="caller-team-permissions",
|
||||
mcp_servers=["slack"],
|
||||
mcp_tool_permissions={"slack": ["read"]},
|
||||
)
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
MCPRequestHandler,
|
||||
"_apply_user_tool_ceiling",
|
||||
AsyncMock(side_effect=lambda tools, _server_id, _auth: tools),
|
||||
)
|
||||
|
||||
auth: Final = actor(("read", "write"))
|
||||
auth.agent_caller = AgentCaller(user_id="alice", team_id="callers")
|
||||
|
||||
assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == {"slack"}
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == ["read"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("fresh", [False, True])
|
||||
@pytest.mark.parametrize("caller_kind", ["team", "user"])
|
||||
async def test_caller_mcp_revocation_uses_fresh_policy(
|
||||
monkeypatch: pytest.MonkeyPatch, fresh: bool, caller_kind: str,
|
||||
) -> None:
|
||||
from litellm.proxy._types import LiteLLM_TeamTable
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key
|
||||
from litellm.types.agents import AgentCaller
|
||||
|
||||
cached_permission: Final = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="caller-permission", mcp_servers=["slack", "linear"],
|
||||
mcp_tool_permissions={"slack": ["read", "write"]},
|
||||
)
|
||||
current_permission: Final = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="caller-permission", mcp_servers=["slack"],
|
||||
mcp_tool_permissions={"slack": ["read"]},
|
||||
)
|
||||
team: Final = LiteLLM_TeamTable(
|
||||
team_id="caller", object_permission_id="caller-permission", object_permission=current_permission,
|
||||
)
|
||||
user: Final = LiteLLM_UserTable(
|
||||
user_id="caller", teams=[], object_permission_id="caller-permission", object_permission=current_permission,
|
||||
)
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=team)
|
||||
database.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=user)
|
||||
database.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=current_permission)
|
||||
cache: Final = UserApiKeyCache()
|
||||
cache.set_cache("team_id:caller", team.model_copy(update={"object_permission": cached_permission}))
|
||||
cache.set_cache("caller", user.model_copy(update={"object_permission": cached_permission}))
|
||||
cache.set_cache(object_permission_cache_key("caller-permission"), cached_permission)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", database)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
|
||||
auth: Final = actor(("read", "write"))
|
||||
auth.requires_fresh_policy = fresh
|
||||
auth.agent_caller = AgentCaller(team_id="caller") if caller_kind == "team" else AgentCaller(user_id="caller")
|
||||
|
||||
assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == ({"slack"} if fresh else {"slack", "linear"})
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("slack", auth) == (["read"] if fresh else ["read", "write"])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("fresh", [False, True])
|
||||
async def test_caller_team_outage_cannot_remove_authoritative_server_ceiling(
|
||||
monkeypatch: pytest.MonkeyPatch, fresh: bool,
|
||||
) -> None:
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.types.agents import AgentCaller
|
||||
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_teamtable.find_unique = AsyncMock(side_effect=RuntimeError("writer unavailable"))
|
||||
database.db.litellm_teamtable.find_unique = AsyncMock(side_effect=RuntimeError("reader unavailable"))
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", database)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", UserApiKeyCache())
|
||||
auth: Final = actor(("read",))
|
||||
auth.agent_caller = AgentCaller(team_id="caller")
|
||||
auth.requires_fresh_policy = fresh
|
||||
|
||||
if fresh:
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await MCPRequestHandler.get_allowed_mcp_servers(auth)
|
||||
assert failure.value.status_code == 503
|
||||
else:
|
||||
assert set(await MCPRequestHandler.get_allowed_mcp_servers(auth)) == {"slack", "linear"}
|
||||
|
|
@ -369,7 +369,9 @@ class TestMCPRequestHandler:
|
|||
result = await MCPRequestHandler._get_allowed_mcp_servers_for_key(user_api_key_auth)
|
||||
|
||||
assert result == ["server-a"]
|
||||
mock_manager.resolve_toolset_tool_permissions.assert_awaited_once_with(toolset_ids=["toolset-1"])
|
||||
mock_manager.resolve_toolset_tool_permissions.assert_awaited_once_with(
|
||||
toolset_ids=["toolset-1"], requires_fresh_policy=False
|
||||
)
|
||||
|
||||
async def test_get_allowed_mcp_servers_for_key_skips_toolset_resolution_when_none_granted(self):
|
||||
user_api_key_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user")
|
||||
|
|
@ -4147,7 +4149,7 @@ async def test_get_allowed_mcp_servers_for_team_uses_helper():
|
|||
"group-server2",
|
||||
}
|
||||
|
||||
mock_get_access_group_servers.assert_called_once_with(["dev-group"])
|
||||
mock_get_access_group_servers.assert_called_once_with(["dev-group"], requires_fresh_policy=False)
|
||||
finally:
|
||||
for sid in ("direct-server1", "direct-server2"):
|
||||
global_mcp_server_manager.registry.pop(sid, None)
|
||||
|
|
@ -4316,7 +4318,7 @@ async def test_get_allowed_mcp_servers_for_key_prefers_in_memory_permission():
|
|||
|
||||
assert set(result) == {"direct-server", "group-server"}
|
||||
mock_get_perm.assert_not_called()
|
||||
mock_access_groups.assert_called_once_with(["grp-alpha"])
|
||||
mock_access_groups.assert_called_once_with(["grp-alpha"], requires_fresh_policy=False)
|
||||
finally:
|
||||
global_mcp_server_manager.registry.pop("direct-server", None)
|
||||
|
||||
|
|
@ -4383,7 +4385,7 @@ class TestAgentMCPPermissions:
|
|||
self._team_servers({"callers": ["server_2", "server_3"]}),
|
||||
),
|
||||
patch.object( # test-quality-ok: agent object_permission lookup hits the DB, not under test here
|
||||
MCPRequestHandler, "_get_allowed_mcp_servers_for_agent", AsyncMock(return_value=[])
|
||||
MCPRequestHandler, "get_allowed_mcp_servers_for_agent", AsyncMock(return_value=[])
|
||||
),
|
||||
patch.object( # test-quality-ok: neither the agent's owner nor the caller has a personal grant
|
||||
MCPRequestHandler, "_get_allowed_mcp_servers_for_user", self._user_servers({})
|
||||
|
|
@ -4402,7 +4404,7 @@ class TestAgentMCPPermissions:
|
|||
MCPRequestHandler, "_get_allowed_mcp_servers_for_team", self._team_servers({})
|
||||
),
|
||||
patch.object( # test-quality-ok: agent object_permission lookup hits the DB, not under test here
|
||||
MCPRequestHandler, "_get_allowed_mcp_servers_for_agent", AsyncMock(return_value=[])
|
||||
MCPRequestHandler, "get_allowed_mcp_servers_for_agent", AsyncMock(return_value=[])
|
||||
),
|
||||
patch.object( # test-quality-ok: same seam, keyed by which user is being asked about
|
||||
MCPRequestHandler, "_get_allowed_mcp_servers_for_user", self._user_servers({"alice": ["server_1"]})
|
||||
|
|
@ -4421,7 +4423,7 @@ class TestAgentMCPPermissions:
|
|||
MCPRequestHandler, "_get_allowed_mcp_servers_for_team", self._team_servers({})
|
||||
),
|
||||
patch.object( # test-quality-ok: agent object_permission lookup hits the DB, not under test here
|
||||
MCPRequestHandler, "_get_allowed_mcp_servers_for_agent", AsyncMock(return_value=[])
|
||||
MCPRequestHandler, "get_allowed_mcp_servers_for_agent", AsyncMock(return_value=[])
|
||||
),
|
||||
patch.object( # test-quality-ok: None is the resolver's own "entitlement unresolvable" signal
|
||||
MCPRequestHandler, "_get_allowed_mcp_servers_for_user", self._user_servers({"alice": None})
|
||||
|
|
@ -4538,7 +4540,7 @@ class TestAgentMCPPermissions:
|
|||
)
|
||||
with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_key") as mock_key:
|
||||
with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_team") as mock_team:
|
||||
with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_agent") as mock_agent:
|
||||
with patch.object(MCPRequestHandler, "get_allowed_mcp_servers_for_agent") as mock_agent:
|
||||
mock_key.return_value = ["server_1", "server_2"]
|
||||
mock_team.return_value = []
|
||||
mock_agent.return_value = ["server_1"]
|
||||
|
|
@ -4555,7 +4557,7 @@ class TestAgentMCPPermissions:
|
|||
)
|
||||
with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_key") as mock_key:
|
||||
with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_team") as mock_team:
|
||||
with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_agent") as mock_agent:
|
||||
with patch.object(MCPRequestHandler, "get_allowed_mcp_servers_for_agent") as mock_agent:
|
||||
mock_key.return_value = ["server_1", "server_2"]
|
||||
mock_team.return_value = []
|
||||
mock_agent.return_value = [] # no agent-level restriction
|
||||
|
|
@ -4611,7 +4613,7 @@ class TestAgentMCPPermissions:
|
|||
)
|
||||
with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_key") as mock_key:
|
||||
with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_team") as mock_team:
|
||||
with patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_agent") as mock_agent:
|
||||
with patch.object(MCPRequestHandler, "get_allowed_mcp_servers_for_agent") as mock_agent:
|
||||
mock_key.return_value = ["server_1", "server_2"]
|
||||
mock_team.return_value = []
|
||||
mock_agent.return_value = ["server_2", "server_3"]
|
||||
|
|
@ -4637,7 +4639,7 @@ class TestAgentMCPPermissions:
|
|||
):
|
||||
with patch.object(
|
||||
MCPRequestHandler,
|
||||
"_get_agent_tool_permissions_for_server",
|
||||
"get_agent_tool_permissions_for_server",
|
||||
new_callable=AsyncMock,
|
||||
return_value=["tool_a"],
|
||||
) as mock_agent_tools:
|
||||
|
|
@ -4669,7 +4671,7 @@ class TestAgentMCPPermissions:
|
|||
):
|
||||
with patch.object(
|
||||
MCPRequestHandler,
|
||||
"_get_agent_tool_permissions_for_server",
|
||||
"get_agent_tool_permissions_for_server",
|
||||
new_callable=AsyncMock,
|
||||
return_value=None,
|
||||
):
|
||||
|
|
@ -4718,10 +4720,12 @@ class TestAgentMCPPermissions:
|
|||
with contextlib.ExitStack() as stack:
|
||||
for patcher in self._agent_toolset_patches(agent_object_permission, mock_manager):
|
||||
stack.enter_context(patcher)
|
||||
result = await MCPRequestHandler._get_allowed_mcp_servers_for_agent(user_api_key_auth)
|
||||
result = await MCPRequestHandler.get_allowed_mcp_servers_for_agent(user_api_key_auth)
|
||||
|
||||
assert sorted(result) == ["server-a", "server-direct"]
|
||||
mock_manager.resolve_toolset_tool_permissions.assert_awaited_once_with(toolset_ids=["toolset-1"])
|
||||
mock_manager.resolve_toolset_tool_permissions.assert_awaited_once_with(
|
||||
toolset_ids=["toolset-1"], requires_fresh_policy=False
|
||||
)
|
||||
|
||||
async def test_get_allowed_mcp_servers_toolset_only_agent_caps_key_servers(self):
|
||||
"""Regression: an agent whose only grant is a toolset used to resolve to [] and place
|
||||
|
|
@ -4760,7 +4764,7 @@ class TestAgentMCPPermissions:
|
|||
for patcher in self._agent_toolset_patches(agent_object_permission, mock_manager):
|
||||
stack.enter_context(patcher)
|
||||
with pytest.raises(UnloadableEntitlementError):
|
||||
await MCPRequestHandler._get_allowed_mcp_servers_for_agent(user_api_key_auth)
|
||||
await MCPRequestHandler.get_allowed_mcp_servers_for_agent(user_api_key_auth)
|
||||
stack.enter_context(
|
||||
patch.object( # test-quality-ok: key resolution has its own tests; pin its grants here
|
||||
MCPRequestHandler,
|
||||
|
|
@ -4789,13 +4793,13 @@ class TestAgentMCPPermissions:
|
|||
with contextlib.ExitStack() as stack:
|
||||
for patcher in self._agent_toolset_patches(agent_object_permission, mock_manager):
|
||||
stack.enter_context(patcher)
|
||||
server_a_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server(
|
||||
server_a_tools = await MCPRequestHandler.get_agent_tool_permissions_for_server(
|
||||
"server-a", user_api_key_auth
|
||||
)
|
||||
server_b_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server(
|
||||
server_b_tools = await MCPRequestHandler.get_agent_tool_permissions_for_server(
|
||||
"server-b", user_api_key_auth
|
||||
)
|
||||
server_c_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server(
|
||||
server_c_tools = await MCPRequestHandler.get_agent_tool_permissions_for_server(
|
||||
"server-c", user_api_key_auth
|
||||
)
|
||||
|
||||
|
|
@ -5833,7 +5837,7 @@ def test_expand_permission_list_does_not_honor_all_proxy_sentinel():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_allowed_mcp_servers_for_team_expands_all_proxy_sentinel_dynamically():
|
||||
async def test_get_allowed_mcp_servers_for_team_expands_all_proxy_sentinel_dynamically(monkeypatch):
|
||||
"""The TEAM resolver expands the all-proxy sentinel to every registered server and
|
||||
picks up a server registered later, so a team scoped to all-proxy tracks the live
|
||||
registry without any change to its stored permission. Reverting the team-side
|
||||
|
|
@ -5850,6 +5854,9 @@ async def test_get_allowed_mcp_servers_for_team_expands_all_proxy_sentinel_dynam
|
|||
from litellm.types.mcp import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
monkeypatch.setattr(global_mcp_server_manager, "registry", {})
|
||||
monkeypatch.setattr(global_mcp_server_manager, "config_mcp_servers", {})
|
||||
|
||||
for sid in ("srv-x", "srv-y"):
|
||||
global_mcp_server_manager.registry[sid] = MCPServer(
|
||||
server_id=sid,
|
||||
|
|
@ -8305,7 +8312,7 @@ class TestUserSubjectTeamUnion:
|
|||
) == ["t1"]
|
||||
# An admitted subject never fans out HERE: it resolves one source per team first, and each of
|
||||
# those pins a team_id, so this helper only ever answers the single-team question. The fan-out
|
||||
# itself is _admitted_subject_sources' job, asserted below.
|
||||
# itself is admitted_subject_sources' job, asserted below.
|
||||
with self._patch(teams_by_id={}, user_teams=["t2", "t3"]):
|
||||
assert await MCPRequestHandler._team_ids_for_mcp_grant(_make_admitted_subject("u")) == []
|
||||
# keyless, no user_id -> nothing
|
||||
|
|
@ -8868,7 +8875,7 @@ class TestUserSubjectTeamUnion:
|
|||
teams["t-member"].organization_id = "org-a"
|
||||
auth = _make_admitted_subject("sso-user")
|
||||
with self._patch(teams_by_id=teams, user_teams=["t-member", "t-stale"]):
|
||||
sources = await MCPRequestHandler._admitted_subject_sources(auth)
|
||||
sources = await MCPRequestHandler.admitted_subject_sources(auth)
|
||||
|
||||
assert [(s.team_id, s.org_id) for s in sources] == [(None, None), ("t-member", "org-a")]
|
||||
# The user's own source carries their grants; a team source must NOT, or the team would be
|
||||
|
|
@ -9673,7 +9680,10 @@ class TestGetUserObjectPermission:
|
|||
|
||||
def _prisma_with_user(self, user_row):
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row)
|
||||
from litellm.proxy._types import LiteLLM_UserTable
|
||||
|
||||
row = LiteLLM_UserTable(user_id="human", object_permission_id=user_row.object_permission_id) if user_row is not None else None
|
||||
prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=row)
|
||||
return prisma_client
|
||||
|
||||
async def test_resolves_through_the_shared_permission_cache(self):
|
||||
|
|
@ -9688,7 +9698,7 @@ class TestGetUserObjectPermission:
|
|||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(service_logging_obj=MagicMock(async_service_success_hook=AsyncMock()))),
|
||||
patch(
|
||||
"litellm.proxy.auth.auth_checks.get_object_permission",
|
||||
new_callable=AsyncMock,
|
||||
|
|
@ -9715,7 +9725,7 @@ class TestGetUserObjectPermission:
|
|||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(service_logging_obj=MagicMock(async_service_success_hook=AsyncMock()))),
|
||||
patch("litellm.proxy.auth.auth_checks.get_object_permission", new_callable=AsyncMock) as mock_get_perm,
|
||||
):
|
||||
assert await MCPRequestHandler._get_user_object_permission(auth) is None
|
||||
|
|
@ -9734,7 +9744,7 @@ class TestGetUserObjectPermission:
|
|||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(service_logging_obj=MagicMock(async_service_success_hook=AsyncMock()))),
|
||||
):
|
||||
assert await MCPRequestHandler._get_user_object_permission(auth) is None
|
||||
|
||||
|
|
@ -9748,7 +9758,7 @@ class TestGetUserObjectPermission:
|
|||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(service_logging_obj=MagicMock(async_service_success_hook=AsyncMock()))),
|
||||
):
|
||||
assert await MCPRequestHandler._get_user_object_permission(auth) is None
|
||||
|
||||
|
|
@ -9765,7 +9775,7 @@ class TestGetUserObjectPermission:
|
|||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(service_logging_obj=MagicMock(async_service_success_hook=AsyncMock()))),
|
||||
patch(
|
||||
"litellm.proxy.auth.auth_checks.get_object_permission",
|
||||
new_callable=AsyncMock,
|
||||
|
|
@ -10085,3 +10095,47 @@ class TestScopedSessionAdmission:
|
|||
def test_scope_field_cannot_be_forged_through_construction(self):
|
||||
forged = UserAPIKeyAuth(user_id="u1", mcp_session_resource_server_id="any-server")
|
||||
assert forged.mcp_session_resource_server_id is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fresh_mcp_user_permission_link_ignores_cached_and_replica_grants(monkeypatch):
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._types import LiteLLM_UserTable
|
||||
|
||||
cached = LiteLLM_UserTable(user_id="fresh-human", object_permission_id="revoked")
|
||||
current = LiteLLM_UserTable(user_id="fresh-human", object_permission_id="current")
|
||||
cache = DualCache()
|
||||
await cache.async_set_cache(key="fresh-human", value=cached)
|
||||
database = MagicMock()
|
||||
database.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=current)
|
||||
database.db.litellm_usertable.find_unique = AsyncMock(return_value=cached)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
|
||||
assert await MCPRequestHandler._user_object_permission_id("fresh-human", database, check_db_only=True) == "current"
|
||||
database.db.litellm_usertable.find_unique.assert_not_awaited()
|
||||
database.writer_db.litellm_usertable.find_unique.side_effect = RuntimeError("unavailable")
|
||||
with pytest.raises(HTTPException) as denied:
|
||||
await MCPRequestHandler._user_object_permission_id("fresh-human", database, check_db_only=True)
|
||||
assert denied.value.status_code == 503
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("operation", ["servers", "tools"])
|
||||
async def test_managed_agent_permission_resolution_outage_is_not_an_unrestricted_grant(monkeypatch, operation):
|
||||
from litellm.proxy._experimental.mcp_server import mcp_server_manager
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
auth = UserAPIKeyAuth(agent_id="managed")
|
||||
auth.managed_agent_policy = AgentResponse(agent_id="managed", agent_name="Managed", agent_card_params={})
|
||||
permission = LiteLLM_ObjectPermissionTable(object_permission_id="policy", mcp_toolsets=["unavailable"])
|
||||
manager = MagicMock()
|
||||
manager.expand_permission_list.return_value = []
|
||||
manager.resolve_toolset_tool_permissions = AsyncMock(side_effect=RuntimeError("policy unavailable"))
|
||||
monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager)
|
||||
resolution = (
|
||||
MCPRequestHandler.get_allowed_mcp_servers_for_agent(auth, permission)
|
||||
if operation == "servers"
|
||||
else MCPRequestHandler.get_agent_tool_permissions_for_server("slack", auth, permission)
|
||||
)
|
||||
with pytest.raises(RuntimeError, match="policy unavailable"):
|
||||
await resolution
|
||||
|
|
|
|||
|
|
@ -7847,7 +7847,7 @@ async def test_load_active_user_by_id_reads_the_row_from_the_database_not_the_ca
|
|||
key="fresh-jwt-user", value=LiteLLM_UserTable(user_id="fresh-jwt-user", teams=[]), model_type=LiteLLM_UserTable
|
||||
)
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_usertable.find_unique = AsyncMock(
|
||||
prisma.writer_db.litellm_usertable.find_unique = AsyncMock(
|
||||
return_value=LiteLLM_UserTable(user_id="fresh-jwt-user", teams=["team-a"])
|
||||
)
|
||||
proxy_globals.user_api_key_cache = cache
|
||||
|
|
|
|||
|
|
@ -211,15 +211,9 @@ def _reload_mcp_manager_module():
|
|||
manager_module = sys.modules["litellm.proxy._experimental.mcp_server.mcp_server_manager"]
|
||||
importlib.reload(utils_module)
|
||||
reloaded = importlib.reload(manager_module)
|
||||
# After reload, server.py still holds a stale reference to the old
|
||||
# global_mcp_server_manager. Update it so tests that exercise server.py
|
||||
# functions (e.g. _get_tools_from_mcp_servers) use the fresh instance.
|
||||
server_module = sys.modules.get("litellm.proxy._experimental.mcp_server.server")
|
||||
if server_module is not None and hasattr(server_module, "global_mcp_server_manager"):
|
||||
server_module.global_mcp_server_manager = reloaded.global_mcp_server_manager
|
||||
operations_module = sys.modules.get("litellm.proxy._experimental.mcp_server.operations")
|
||||
if operations_module is not None:
|
||||
operations_module.global_mcp_server_manager = reloaded.global_mcp_server_manager
|
||||
for name, module in tuple(sys.modules.items()):
|
||||
if name.startswith("litellm.proxy._experimental.mcp_server.") and hasattr(module, "global_mcp_server_manager"):
|
||||
module.global_mcp_server_manager = reloaded.global_mcp_server_manager
|
||||
return reloaded
|
||||
|
||||
|
||||
|
|
@ -230,6 +224,20 @@ def enable_eager_mcp_oauth_discovery(monkeypatch):
|
|||
monkeypatch.setenv("LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP", "1")
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def restore_mcp_manager_singleton():
|
||||
"""``_reload_mcp_manager_module`` rebinds ``global_mcp_server_manager`` in every MCP module, so
|
||||
without this the next test file inherits a manager that has none of its servers registered."""
|
||||
bound: Final = tuple(
|
||||
(module, module.global_mcp_server_manager)
|
||||
for name, module in tuple(sys.modules.items())
|
||||
if name.startswith("litellm.proxy._experimental.mcp_server.") and hasattr(module, "global_mcp_server_manager")
|
||||
)
|
||||
yield
|
||||
for module, manager in bound:
|
||||
module.global_mcp_server_manager = manager
|
||||
|
||||
|
||||
class TestMCPServerManager:
|
||||
"""Test MCP Server Manager stdio functionality"""
|
||||
|
||||
|
|
@ -5585,9 +5593,7 @@ class TestMCPServerManager:
|
|||
|
||||
# Mock dependencies - set object_permission and object_permission_id to None
|
||||
# so permission checks return None (no restrictions)
|
||||
user_api_key_auth = MagicMock()
|
||||
user_api_key_auth.object_permission = None
|
||||
user_api_key_auth.object_permission_id = None
|
||||
user_api_key_auth: Final = UserAPIKeyAuth()
|
||||
proxy_logging_obj = MagicMock()
|
||||
|
||||
# Mock the async methods that pre_call_tool_check calls
|
||||
|
|
@ -5654,9 +5660,7 @@ class TestMCPServerManager:
|
|||
|
||||
# Mock dependencies - set object_permission and object_permission_id to None
|
||||
# so permission checks return None (no restrictions)
|
||||
user_api_key_auth = MagicMock()
|
||||
user_api_key_auth.object_permission = None
|
||||
user_api_key_auth.object_permission_id = None
|
||||
user_api_key_auth: Final = UserAPIKeyAuth()
|
||||
proxy_logging_obj = MagicMock()
|
||||
|
||||
# Mock the async methods that pre_call_tool_check calls
|
||||
|
|
@ -5723,9 +5727,7 @@ class TestMCPServerManager:
|
|||
|
||||
# Mock dependencies - set object_permission and object_permission_id to None
|
||||
# so permission checks return None (no restrictions)
|
||||
user_api_key_auth = MagicMock()
|
||||
user_api_key_auth.object_permission = None
|
||||
user_api_key_auth.object_permission_id = None
|
||||
user_api_key_auth: Final = UserAPIKeyAuth()
|
||||
proxy_logging_obj = MagicMock()
|
||||
|
||||
# Mock the async methods that pre_call_tool_check calls
|
||||
|
|
@ -5760,9 +5762,7 @@ class TestMCPServerManager:
|
|||
|
||||
# Mock dependencies - set object_permission and object_permission_id to None
|
||||
# so permission checks return None (no restrictions)
|
||||
user_api_key_auth = MagicMock()
|
||||
user_api_key_auth.object_permission = None
|
||||
user_api_key_auth.object_permission_id = None
|
||||
user_api_key_auth: Final = UserAPIKeyAuth()
|
||||
proxy_logging_obj = MagicMock()
|
||||
|
||||
# Mock the async methods that pre_call_tool_check calls
|
||||
|
|
@ -6838,9 +6838,7 @@ class TestMCPServerManager:
|
|||
|
||||
# Mock dependencies - set object_permission and object_permission_id to None
|
||||
# so permission checks return None (no restrictions)
|
||||
user_api_key_auth = MagicMock()
|
||||
user_api_key_auth.object_permission = None
|
||||
user_api_key_auth.object_permission_id = None
|
||||
user_api_key_auth: Final = UserAPIKeyAuth()
|
||||
proxy_logging_obj = MagicMock()
|
||||
|
||||
# Mock the async methods that pre_call_tool_check calls
|
||||
|
|
@ -6924,9 +6922,7 @@ class TestMCPServerManager:
|
|||
manager._create_mcp_client = AsyncMock(return_value=mock_client)
|
||||
|
||||
# Mock user auth with no restrictions
|
||||
user_api_key_auth = MagicMock()
|
||||
user_api_key_auth.object_permission = None
|
||||
user_api_key_auth.object_permission_id = None
|
||||
user_api_key_auth: Final = UserAPIKeyAuth()
|
||||
|
||||
# Mock proxy logging
|
||||
proxy_logging_obj = MagicMock()
|
||||
|
|
@ -11175,6 +11171,72 @@ async def test_resolve_toolset_tool_permissions_single_db_fetch_across_checks():
|
|||
list_toolsets_mock.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_toolset_tool_permissions_fresh_policy_sees_writer_revocation_past_warm_cache():
|
||||
"""A managed agent's tool grant revoked in the writer DB must be gone on the very next fresh
|
||||
request even though the legacy cache still holds the old grant, and the fresh read must go to
|
||||
the writer, not the replica"""
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
MCPServerManager,
|
||||
)
|
||||
|
||||
manager = MCPServerManager()
|
||||
granted = MagicMock()
|
||||
granted.tools = [{"server_id": "server-a", "tool_name": "echo"}]
|
||||
revoked = MagicMock()
|
||||
revoked.tools = [{"server_id": "server-a", "tool_name": "other"}]
|
||||
list_toolsets_mock = AsyncMock(side_effect=[[granted], [revoked]])
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.toolset_db.list_mcp_toolsets",
|
||||
list_toolsets_mock,
|
||||
),
|
||||
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()),
|
||||
):
|
||||
warm = await manager.resolve_toolset_tool_permissions(toolset_ids=["ts-1"])
|
||||
legacy_after_revoke = await manager.resolve_toolset_tool_permissions(toolset_ids=["ts-1"])
|
||||
fresh_after_revoke = await manager.resolve_toolset_tool_permissions(
|
||||
toolset_ids=["ts-1"], requires_fresh_policy=True
|
||||
)
|
||||
|
||||
assert warm == {"server-a": ["echo"]}
|
||||
assert legacy_after_revoke == warm, "legacy callers keep the cached grant by design"
|
||||
assert fresh_after_revoke == {"server-a": ["other"]}
|
||||
assert list_toolsets_mock.await_count == 2
|
||||
assert list_toolsets_mock.await_args_list[0].kwargs["use_writer"] is False
|
||||
assert list_toolsets_mock.await_args_list[1].kwargs["use_writer"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_toolset_tool_permissions_fresh_policy_propagates_db_fault_instead_of_no_grants():
|
||||
"""A fresh read that fails must raise so the managed-agent boundary fails closed; the legacy
|
||||
path keeps its swallow-to-empty behaviour"""
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
MCPServerManager,
|
||||
)
|
||||
|
||||
manager = MCPServerManager()
|
||||
list_toolsets_mock = AsyncMock(side_effect=RuntimeError("relation does not exist"))
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.toolset_db.list_mcp_toolsets",
|
||||
list_toolsets_mock,
|
||||
),
|
||||
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", DualCache()),
|
||||
):
|
||||
legacy = await manager.resolve_toolset_tool_permissions(toolset_ids=["ts-1"])
|
||||
with pytest.raises(RuntimeError, match="relation does not exist"):
|
||||
await manager.resolve_toolset_tool_permissions(toolset_ids=["ts-1"], requires_fresh_policy=True)
|
||||
|
||||
assert legacy == {}
|
||||
|
||||
|
||||
class TestMaterializeAuthHeaders:
|
||||
"""_materialize_auth_headers drives one step of a resolved httpx.Auth's own flow to turn it
|
||||
into a header dict for the OpenAPI egress arm, which sends plain headers and cannot carry an
|
||||
|
|
@ -11496,12 +11558,8 @@ class TestDiscoveryFailureLogging:
|
|||
assert "unresolved" in caplog.text
|
||||
|
||||
|
||||
def _unrestricted_auth() -> MagicMock:
|
||||
"""A caller with no object_permission, so only server-level checks apply."""
|
||||
user_api_key_auth = MagicMock()
|
||||
user_api_key_auth.object_permission = None
|
||||
user_api_key_auth.object_permission_id = None
|
||||
return user_api_key_auth
|
||||
def _unrestricted_auth() -> UserAPIKeyAuth:
|
||||
return UserAPIKeyAuth()
|
||||
|
||||
|
||||
def _permissive_proxy_logging() -> MagicMock:
|
||||
|
|
|
|||
|
|
@ -139,7 +139,7 @@ async def test_mint_reads_the_users_teams_from_the_database_not_a_stale_cached_r
|
|||
key="stale-cache-user", value=_user(user_id="stale-cache-user", teams=[]), model_type=LiteLLM_UserTable
|
||||
)
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_usertable.find_unique = AsyncMock(
|
||||
prisma.writer_db.litellm_usertable.find_unique = AsyncMock(
|
||||
return_value=_user(user_id="stale-cache-user", teams=["team-a"])
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
|
||||
|
|
@ -164,7 +164,7 @@ async def test_mint_refuses_a_user_scim_deactivated_after_the_cache_last_saw_the
|
|||
key="deactivated-user", value=_user(user_id="deactivated-user", teams=["team-a"]), model_type=LiteLLM_UserTable
|
||||
)
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_usertable.find_unique = AsyncMock(
|
||||
prisma.writer_db.litellm_usertable.find_unique = AsyncMock(
|
||||
return_value=_user(user_id="deactivated-user", teams=["team-a"], metadata={"scim_active": False})
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
|
||||
|
|
|
|||
|
|
@ -1311,7 +1311,7 @@ class TestListToolsRestAPI:
|
|||
session_auth = UserAPIKeyAuth(team_id=UI_SESSION_TOKEN_TEAM_ID, user_id="grant-user", user_role="internal_user")
|
||||
admitted_auth = UserAPIKeyAuth(user_id="grant-user", org_id="admitted-org")
|
||||
|
||||
async def fake_reload(user_id):
|
||||
async def fake_reload(user_id, *, requires_fresh_policy=False):
|
||||
assert user_id == "grant-user"
|
||||
return admitted_auth
|
||||
|
||||
|
|
|
|||
|
|
@ -145,7 +145,7 @@ async def test_build_effective_auth_contexts_appends_admitted_user_context(monke
|
|||
|
||||
assert contexts[-1].user_id == "user-42" and contexts[-1].team_id is None
|
||||
assert [ctx.team_id for ctx in contexts[:-1]] == ["team-one"]
|
||||
reload_mock.assert_awaited_once_with("user-42")
|
||||
reload_mock.assert_awaited_once_with("user-42", requires_fresh_policy=False)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -198,7 +198,7 @@ async def test_acting_user_auth_returns_admitted_subject_for_non_admin_sessions(
|
|||
result = await acting_user_auth(user_auth)
|
||||
|
||||
assert result.user_id == "user-42" and result.team_id is None
|
||||
reload_mock.assert_awaited_once_with("user-42")
|
||||
reload_mock.assert_awaited_once_with("user-42", requires_fresh_policy=False)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -144,3 +144,27 @@ 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) as failure:
|
||||
await _load_access_group("group", check_db_only=True)
|
||||
assert failure.value.status_code == 503
|
||||
else:
|
||||
assert await _load_access_group("group") is None
|
||||
|
|
|
|||
|
|
@ -67,7 +67,7 @@ class TestAgentRequestHandler:
|
|||
|
||||
# Case 1: Both key and team have agents - intersection
|
||||
with patch.object(
|
||||
AgentRequestHandler, "_get_allowed_agents_for_key"
|
||||
AgentRequestHandler, "get_allowed_agents_for_key"
|
||||
) as mock_key:
|
||||
with patch.object(
|
||||
AgentRequestHandler, "_get_allowed_agents_for_team"
|
||||
|
|
@ -86,7 +86,7 @@ class TestAgentRequestHandler:
|
|||
|
||||
# Case 2: Team has agents, key has none - inherit from team
|
||||
with patch.object(
|
||||
AgentRequestHandler, "_get_allowed_agents_for_key"
|
||||
AgentRequestHandler, "get_allowed_agents_for_key"
|
||||
) as mock_key:
|
||||
with patch.object(
|
||||
AgentRequestHandler, "_get_allowed_agents_for_team"
|
||||
|
|
@ -105,7 +105,7 @@ class TestAgentRequestHandler:
|
|||
|
||||
# Case 3: Key has agents, team has none - key restrictions stand
|
||||
with patch.object(
|
||||
AgentRequestHandler, "_get_allowed_agents_for_key"
|
||||
AgentRequestHandler, "get_allowed_agents_for_key"
|
||||
) as mock_key:
|
||||
with patch.object(
|
||||
AgentRequestHandler, "_get_allowed_agents_for_team"
|
||||
|
|
@ -120,7 +120,7 @@ class TestAgentRequestHandler:
|
|||
|
||||
# Case 4: No grant anywhere - unrestricted (documented open-by-default)
|
||||
with patch.object(
|
||||
AgentRequestHandler, "_get_allowed_agents_for_key"
|
||||
AgentRequestHandler, "get_allowed_agents_for_key"
|
||||
) as mock_key:
|
||||
with patch.object(
|
||||
AgentRequestHandler, "_get_allowed_agents_for_team"
|
||||
|
|
@ -141,7 +141,7 @@ class TestAgentRequestHandler:
|
|||
api_key="test-key", user_id="test-user", team_id="test-team"
|
||||
)
|
||||
|
||||
with patch.object(AgentRequestHandler, "_get_allowed_agents_for_key") as mock_key:
|
||||
with patch.object(AgentRequestHandler, "get_allowed_agents_for_key") as mock_key:
|
||||
with patch.object(AgentRequestHandler, "_get_allowed_agents_for_team") as mock_team:
|
||||
mock_key.return_value = RestrictedAgentAccess(frozenset({"agent-alpha"}))
|
||||
mock_team.return_value = RestrictedAgentAccess(frozenset({"agent-beta"}))
|
||||
|
|
@ -198,7 +198,7 @@ class TestAgentRequestHandler:
|
|||
|
||||
@staticmethod
|
||||
def _team_grants(grants: dict[str, AgentAccess]) -> AsyncMock:
|
||||
async def by_team(user_api_key_auth: UserAPIKeyAuth | None = None) -> AgentAccess:
|
||||
async def by_team(user_api_key_auth: UserAPIKeyAuth | None = None, *, strict: bool = False) -> AgentAccess:
|
||||
assert user_api_key_auth is not None
|
||||
return grants.get(user_api_key_auth.team_id or "", UnrestrictedAgentAccess())
|
||||
|
||||
|
|
@ -237,6 +237,29 @@ class TestAgentRequestHandler:
|
|||
frozenset()
|
||||
)
|
||||
|
||||
async def test_managed_agent_acting_for_a_user_is_capped_at_the_invoking_teams_agents(self):
|
||||
"""The managed path must honour the invoking team's ceiling the same way the unmanaged path does:
|
||||
the agent's own policy grants alpha and beta, but the human who invoked it reaches only beta."""
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
managed: Final = UserAPIKeyAuth(api_key="test-key", user_id="test-user", agent_id="actor")
|
||||
managed.managed_agent_policy = AgentResponse(
|
||||
agent_id="actor",
|
||||
agent_name="Actor",
|
||||
agent_card_params={},
|
||||
object_permission={"object_permission_id": "own", "agents": ["agent-alpha", "agent-beta"]},
|
||||
)
|
||||
managed.agent_caller = AgentCaller(user_id="alice", team_id="callers")
|
||||
|
||||
with patch.object( # test-quality-ok: the team resolver reads proxy_server globals with no injection seam
|
||||
AgentRequestHandler,
|
||||
"_get_allowed_agents_for_team",
|
||||
self._team_grants({"callers": RestrictedAgentAccess(frozenset({"agent-beta"}))}),
|
||||
):
|
||||
assert await AgentRequestHandler.resolve_agent_access(managed) == RestrictedAgentAccess(
|
||||
frozenset({"agent-beta"})
|
||||
)
|
||||
|
||||
async def test_agent_key_acting_for_an_ungranted_caller_keeps_its_own_agents(self):
|
||||
agent_key: Final = self._key_granting(["agent-alpha"], agent_id="caller-agent")
|
||||
agent_key.agent_caller = AgentCaller(user_id="alice", team_id="callers")
|
||||
|
|
@ -249,7 +272,6 @@ class TestAgentRequestHandler:
|
|||
frozenset({"agent-alpha"})
|
||||
)
|
||||
|
||||
|
||||
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"}))
|
||||
|
|
@ -299,7 +321,7 @@ class TestAgentRequestHandler:
|
|||
) as mock_groups:
|
||||
mock_groups.return_value = []
|
||||
|
||||
assert await AgentRequestHandler._get_allowed_agents_for_key(
|
||||
assert await AgentRequestHandler.get_allowed_agents_for_key(
|
||||
user_api_key_auth=mock_user_auth
|
||||
) == RestrictedAgentAccess(frozenset())
|
||||
|
||||
|
|
@ -315,7 +337,7 @@ class TestAgentRequestHandler:
|
|||
) as mock_groups:
|
||||
mock_groups.side_effect = Exception("DB Error")
|
||||
|
||||
assert await AgentRequestHandler._get_allowed_agents_for_key(
|
||||
assert await AgentRequestHandler.get_allowed_agents_for_key(
|
||||
user_api_key_auth=mock_user_auth
|
||||
) == UnrestrictedAgentAccess()
|
||||
|
||||
|
|
@ -404,7 +426,7 @@ class TestAgentRequestHandler:
|
|||
)
|
||||
|
||||
with patch.object(
|
||||
AgentRequestHandler, "_get_allowed_agents_for_key"
|
||||
AgentRequestHandler, "get_allowed_agents_for_key"
|
||||
) as mock_key:
|
||||
with patch.object(
|
||||
AgentRequestHandler, "_get_allowed_agents_for_team"
|
||||
|
|
@ -489,9 +511,9 @@ class TestAgentRequestHandler:
|
|||
listed: Final = await accessible_agents(session, registry.get_agent_list(), resolve_access, effective_contexts)
|
||||
assert {agent.agent_name for agent in listed} == {"alpha", "beta"}
|
||||
|
||||
async def test_get_allowed_agents_for_key_via_access_group_ids(self):
|
||||
async def testget_allowed_agents_for_key_via_access_group_ids(self):
|
||||
"""
|
||||
Test that _get_allowed_agents_for_key includes agents from key's access_group_ids
|
||||
Test that get_allowed_agents_for_key includes agents from key's access_group_ids
|
||||
(unified access groups) when key has no native object_permission.
|
||||
"""
|
||||
mock_user_auth = UserAPIKeyAuth(
|
||||
|
|
@ -508,16 +530,16 @@ class TestAgentRequestHandler:
|
|||
new_callable=AsyncMock,
|
||||
return_value=["agent-from-ag-1", "agent-from-ag-2"],
|
||||
):
|
||||
result = await AgentRequestHandler._get_allowed_agents_for_key(
|
||||
result = await AgentRequestHandler.get_allowed_agents_for_key(
|
||||
user_api_key_auth=mock_user_auth
|
||||
)
|
||||
assert result == RestrictedAgentAccess(
|
||||
frozenset({"agent-from-ag-1", "agent-from-ag-2"})
|
||||
)
|
||||
|
||||
async def test_get_allowed_agents_for_key_combines_native_and_access_groups(self):
|
||||
async def testget_allowed_agents_for_key_combines_native_and_access_groups(self):
|
||||
"""
|
||||
Test that _get_allowed_agents_for_key combines agents from native object_permission
|
||||
Test that get_allowed_agents_for_key combines agents from native object_permission
|
||||
and key's access_group_ids (unified access groups).
|
||||
"""
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
|
||||
|
|
@ -540,7 +562,7 @@ class TestAgentRequestHandler:
|
|||
new_callable=AsyncMock,
|
||||
return_value=["agent-from-ag"],
|
||||
):
|
||||
result = await AgentRequestHandler._get_allowed_agents_for_key(
|
||||
result = await AgentRequestHandler.get_allowed_agents_for_key(
|
||||
user_api_key_auth=mock_user_auth
|
||||
)
|
||||
assert result == RestrictedAgentAccess(
|
||||
|
|
@ -611,7 +633,7 @@ class TestAgentRequestHandler:
|
|||
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry",
|
||||
registry,
|
||||
):
|
||||
with patch.object(AgentRequestHandler, "_get_allowed_agents_for_key") as mock_key:
|
||||
with patch.object(AgentRequestHandler, "get_allowed_agents_for_key") as mock_key:
|
||||
with patch.object(AgentRequestHandler, "_get_allowed_agents_for_team") as mock_team:
|
||||
for key_grant, team_grant in (
|
||||
(
|
||||
|
|
@ -632,3 +654,392 @@ class TestAgentRequestHandler:
|
|||
assert await AgentRequestHandler.resolve_agent_access(
|
||||
user_api_key_auth=mock_user_auth
|
||||
) == RestrictedAgentAccess(frozenset({agent.agent_id})), (key_grant, team_grant)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"state,allowed",
|
||||
[
|
||||
({}, True),
|
||||
({"enabled": False}, False),
|
||||
],
|
||||
)
|
||||
async def test_managed_invocation_requires_local_and_directory_admission(
|
||||
monkeypatch: pytest.MonkeyPatch, state: dict[str, object], allowed: bool
|
||||
) -> None:
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.agent_identity import AgentIdentityBinding
|
||||
|
||||
binding: Final = AgentIdentityBinding(
|
||||
agent_id="target",
|
||||
provider="microsoft_entra",
|
||||
tenant_id="tenant",
|
||||
client_id="client",
|
||||
issuer="issuer",
|
||||
revision="revision",
|
||||
)
|
||||
target: Final = AgentResponse(
|
||||
agent_id="target", agent_name="Target", agent_card_params={}, identity=binding, identity_managed=True
|
||||
).model_copy(update=state)
|
||||
client: Final = MagicMock()
|
||||
client.writer_db.litellm_agentstable.find_unique = AsyncMock(return_value=target)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", client)
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="human-grant", agents=["target"])
|
||||
auth: Final = UserAPIKeyAuth(user_id="human", object_permission=permission)
|
||||
assert await AgentRequestHandler.is_agent_allowed("target", auth) is allowed
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("delegated", [True, False])
|
||||
async def test_managed_agent_invocation_grants_intersect_verified_user_grants(
|
||||
monkeypatch: pytest.MonkeyPatch, delegated: bool
|
||||
) -> None:
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
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 AgentIdentityBinding, ManagedAgentContext
|
||||
|
||||
database: Final = MagicMock()
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", database)
|
||||
own: Final = LiteLLM_ObjectPermissionTable(object_permission_id="own", agents=["shared", "agent-only"])
|
||||
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", api_key="verified-jwt")
|
||||
auth.managed_agent_policy = AgentResponse(
|
||||
agent_id="actor", agent_name="Actor", agent_card_params={}, object_permission=own.model_dump()
|
||||
)
|
||||
auth.managed_agent_context = ManagedAgentContext(
|
||||
agent_id="actor", mode="delegated" if delegated else "autonomous", user_id="human" if delegated else None
|
||||
)
|
||||
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"])
|
||||
async def test_delegated_grants_revoke_with_warm_user_team_and_permission_caches(
|
||||
monkeypatch: pytest.MonkeyPatch, revoked: str
|
||||
) -> None:
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._types import LiteLLM_AccessGroupTable, LiteLLM_TeamTable, LiteLLM_UserTable
|
||||
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import verified_human_agent_grants
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key
|
||||
|
||||
direct: Final = revoked == "direct-grant"
|
||||
grouped: Final = revoked == "access-group"
|
||||
permission: Final = LiteLLM_ObjectPermissionTable(object_permission_id="grant", agents=["target"])
|
||||
human: Final = LiteLLM_UserTable(
|
||||
user_id="human",
|
||||
teams=[] if direct else ["team"],
|
||||
organization_memberships=[],
|
||||
object_permission_id="grant" if direct else None,
|
||||
)
|
||||
team: Final = LiteLLM_TeamTable(
|
||||
team_id="team",
|
||||
models=[],
|
||||
members_with_roles=[{"user_id": "human", "role": "user"}],
|
||||
object_permission_id=None if grouped else "grant",
|
||||
access_group_ids=["group"] if grouped else [],
|
||||
)
|
||||
group: Final = LiteLLM_AccessGroupTable(
|
||||
access_group_id="group", access_group_name="Group", access_agent_ids=["target"]
|
||||
)
|
||||
cache: Final = UserApiKeyCache()
|
||||
cache.set_cache("human", human)
|
||||
cache.set_cache("team_id:team", team)
|
||||
cache.set_cache(object_permission_cache_key("grant"), permission)
|
||||
cache.set_cache("access_group_id:group", group)
|
||||
client: Final = MagicMock()
|
||||
client.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=human)
|
||||
client.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=team)
|
||||
client.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=permission)
|
||||
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", "team") == frozenset({"target"})
|
||||
client.writer_db.litellm_usertable.find_unique.return_value = (
|
||||
human.model_copy(update={"teams": []}) if revoked == "user" else human
|
||||
)
|
||||
client.writer_db.litellm_teamtable.find_unique.return_value = (
|
||||
team.model_copy(update={"members_with_roles": []})
|
||||
if revoked == "team-member"
|
||||
else team.model_copy(update={"object_permission_id": None})
|
||||
if revoked == "team-grant"
|
||||
else team
|
||||
)
|
||||
client.writer_db.litellm_objectpermissiontable.find_unique.return_value = (
|
||||
permission.model_copy(update={"agents": []}) if direct or revoked == "team-permission" else permission
|
||||
)
|
||||
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", "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()
|
||||
client.db.litellm_accessgrouptable.find_unique.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_strict_legacy_group_grants_ignore_stale_replica(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.agent_endpoints import agent_registry
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
stale: Final = AgentResponse(agent_id="revoked", agent_name="Revoked", agent_card_params={})
|
||||
registry: Final = AgentRegistry()
|
||||
registry.register_agent(stale)
|
||||
database: Final = MagicMock()
|
||||
database.db.litellm_agentstable.find_many = AsyncMock(return_value=[stale])
|
||||
database.writer_db.litellm_agentstable.find_many = AsyncMock(return_value=[stale])
|
||||
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", database)
|
||||
auth: Final = UserAPIKeyAuth(
|
||||
object_permission=LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="permission", agent_access_groups=["group"]
|
||||
)
|
||||
)
|
||||
assert await AgentRequestHandler.get_allowed_agents_for_key(auth, strict=True) == RestrictedAgentAccess(
|
||||
frozenset({"revoked"})
|
||||
)
|
||||
database.writer_db.litellm_agentstable.find_many.return_value = []
|
||||
assert await AgentRequestHandler.get_allowed_agents_for_key(auth, strict=True) == RestrictedAgentAccess(
|
||||
frozenset()
|
||||
)
|
||||
database.db.litellm_agentstable.find_many.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("groups", [[], ["group"]])
|
||||
async def test_legacy_groups_without_database_grant_no_agents(groups: list[str]) -> None:
|
||||
assert await AgentRequestHandler._get_db_agent_ids_for_access_groups(None, groups, check_db_only=True) == set()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("team", [False, True])
|
||||
async def test_strict_invocation_policy_outage_denies_instead_of_allowing_all(
|
||||
monkeypatch: pytest.MonkeyPatch, team: bool
|
||||
) -> None:
|
||||
from fastapi import HTTPException
|
||||
from unittest.mock import MagicMock
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_teamtable.find_unique = AsyncMock(side_effect=ConnectionError("writer unavailable"))
|
||||
database.writer_db.litellm_agentstable.find_many = AsyncMock(side_effect=ConnectionError("writer unavailable"))
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", database)
|
||||
auth: Final = UserAPIKeyAuth(
|
||||
team_id="team" if team else None,
|
||||
object_permission=None if team else LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="grant", agent_access_groups=["group"]
|
||||
),
|
||||
)
|
||||
with pytest.raises(HTTPException, match="policy is unavailable") as denied:
|
||||
await AgentRequestHandler.resolve_key_team_agent_access(auth, strict=True)
|
||||
assert denied.value.status_code == 503
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("available", [False, True])
|
||||
async def test_missing_team_cannot_grant_strict_agent_access(monkeypatch: pytest.MonkeyPatch, available: bool) -> None:
|
||||
from unittest.mock import MagicMock
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.auth import auth_checks
|
||||
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", MagicMock() if available else None)
|
||||
monkeypatch.setattr(auth_checks, "get_team_object", AsyncMock(return_value=None))
|
||||
assert await AgentRequestHandler._get_allowed_agents_for_team(
|
||||
UserAPIKeyAuth(team_id="missing"), strict=True
|
||||
) == RestrictedAgentAccess(frozenset())
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("outage", [False, True])
|
||||
async def test_registered_managed_target_cannot_bypass_missing_or_unavailable_policy(
|
||||
monkeypatch: pytest.MonkeyPatch, outage: bool
|
||||
) -> None:
|
||||
from fastapi import HTTPException
|
||||
from unittest.mock import MagicMock
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.agent_endpoints import agent_registry
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
registry: Final = AgentRegistry()
|
||||
registry.register_agent(AgentResponse(
|
||||
agent_id="target", agent_name="Target", agent_card_params={}, identity_managed=True
|
||||
))
|
||||
monkeypatch.setattr(agent_registry, "global_agent_registry", registry)
|
||||
database: Final = MagicMock()
|
||||
database.writer_db.litellm_agentstable.find_unique = AsyncMock(
|
||||
return_value=None, side_effect=ConnectionError("unavailable") if outage else None
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", database)
|
||||
if outage:
|
||||
with pytest.raises(HTTPException, match="could not be loaded") as denied:
|
||||
await AgentRequestHandler.is_agent_allowed("target", UserAPIKeyAuth())
|
||||
assert denied.value.status_code == 503
|
||||
else:
|
||||
assert await AgentRequestHandler.is_agent_allowed("target", UserAPIKeyAuth()) is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("grant", [False, True])
|
||||
async def test_delegation_without_a_verified_human_never_grants_agents(grant: bool) -> None:
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.agent_identity import ManagedAgentContext
|
||||
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import verified_human_agent_grants
|
||||
|
||||
auth: Final = UserAPIKeyAuth(agent_id="actor")
|
||||
auth.managed_agent_policy = AgentResponse(
|
||||
agent_id="actor", agent_name="Actor", agent_card_params={},
|
||||
object_permission={"object_permission_id": "own", "agents": ["target"]} if grant else None,
|
||||
)
|
||||
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
|
||||
|
||||
|
||||
@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.model_copy(update={"access_agent_ids": ["target"]}))
|
||||
cache.set_cache("team_id:caller-team", team.model_copy(update={"object_permission": permission}))
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", database)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
|
||||
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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -1141,7 +1141,7 @@ async def test_get_user_object_check_db_only_ignores_recent_miss(monkeypatch):
|
|||
monkeypatch.setitem(auth_checks.last_db_access_time, f"user_id:{user_id}", (None, time.time()))
|
||||
db_row = LiteLLM_UserTable(user_id=user_id, user_email=None, user_role="internal_user")
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=db_row)
|
||||
mock_prisma_client.writer_db.litellm_usertable.find_unique = AsyncMock(return_value=db_row)
|
||||
|
||||
result = await get_user_object(
|
||||
user_id=user_id,
|
||||
|
|
@ -1153,7 +1153,7 @@ async def test_get_user_object_check_db_only_ignores_recent_miss(monkeypatch):
|
|||
|
||||
assert result is not None
|
||||
assert result.user_id == user_id
|
||||
mock_prisma_client.db.litellm_usertable.find_unique.assert_awaited_once()
|
||||
mock_prisma_client.writer_db.litellm_usertable.find_unique.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -3108,7 +3108,7 @@ async def test_get_team_object_raises_404_when_not_found():
|
|||
mock_prisma_client = MagicMock()
|
||||
mock_db = AsyncMock()
|
||||
mock_prisma_client.db = mock_db
|
||||
mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None)
|
||||
mock_prisma_client.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=None)
|
||||
|
||||
mock_cache = MagicMock()
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
||||
|
|
@ -3126,11 +3126,40 @@ async def test_get_team_object_raises_404_when_not_found():
|
|||
assert "Team doesn't exist in db" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_team_object_check_db_only_reads_writer_through_the_shared_loader():
|
||||
"""Management endpoints mock ``_get_team_object_from_user_api_key_cache`` and expect
|
||||
``check_db_only`` to still flow through it; only the table it reads moves to the writer."""
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from litellm.proxy.auth import auth_checks
|
||||
from litellm.proxy.auth.auth_checks import get_team_object
|
||||
|
||||
row = {"team_id": "team-writer", "models": ["gpt-4o"], "object_permission_id": None}
|
||||
prisma = MagicMock()
|
||||
prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=SimpleNamespace(dict=lambda: row))
|
||||
prisma.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=SimpleNamespace(dict=lambda: row))
|
||||
cache = MagicMock()
|
||||
cache.async_get_cache = AsyncMock(return_value=None)
|
||||
cache.async_set_cache = AsyncMock()
|
||||
shared_loader = AsyncMock(wraps=auth_checks._get_team_object_from_user_api_key_cache)
|
||||
|
||||
with patch.object(auth_checks, "_get_team_object_from_user_api_key_cache", shared_loader):
|
||||
team = await get_team_object("team-writer", prisma, cache, check_db_only=True)
|
||||
|
||||
assert team.team_id == "team-writer"
|
||||
assert shared_loader.await_args.kwargs["use_writer"] is True
|
||||
prisma.writer_db.litellm_teamtable.find_unique.assert_awaited_once()
|
||||
prisma.db.litellm_teamtable.find_unique.assert_not_awaited()
|
||||
cache.async_set_cache.assert_awaited_once()
|
||||
|
||||
|
||||
def _mock_prisma_for_team_lookup(find_unique):
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_teamtable.find_unique = find_unique
|
||||
mock_prisma_client.writer_db.litellm_teamtable.find_unique = find_unique
|
||||
return mock_prisma_client
|
||||
|
||||
|
||||
|
|
@ -10054,3 +10083,196 @@ def test_can_object_call_model_allows_listed_model_for_key():
|
|||
)
|
||||
|
||||
assert result is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("allowed", [True, False])
|
||||
async def test_authoritative_access_group_reads_writer_despite_stale_allow_cache(allowed: bool) -> None:
|
||||
from litellm.proxy._types import LiteLLM_AccessGroupTable
|
||||
from litellm.proxy.auth.auth_checks import get_access_object
|
||||
|
||||
stale: Final = LiteLLM_AccessGroupTable(access_group_id="group", access_group_name="Policy", access_model_names=["old"])
|
||||
current: Final = stale.model_copy(update={"access_model_names": ["new"] if allowed else []})
|
||||
client: Final = MagicMock()
|
||||
client.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(return_value=current)
|
||||
client.db.litellm_accessgrouptable.find_unique = AsyncMock(return_value=stale)
|
||||
cache: Final = MagicMock()
|
||||
cache.async_get_cache = AsyncMock(return_value=stale)
|
||||
cache.async_set_cache = AsyncMock()
|
||||
result: Final = await get_access_object("group", client, cache, check_db_only=True)
|
||||
assert result.access_model_names == (["new"] if allowed else [])
|
||||
cache.async_get_cache.assert_not_awaited()
|
||||
client.db.litellm_accessgrouptable.find_unique.assert_not_awaited()
|
||||
client.writer_db.litellm_accessgrouptable.find_unique.assert_awaited_once_with(where={"access_group_id": "group"})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authoritative_access_group_outage_does_not_use_cached_grants() -> None:
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy.auth.auth_checks import get_access_object
|
||||
|
||||
client: Final = MagicMock()
|
||||
client.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(side_effect=RuntimeError("writer unavailable"))
|
||||
cache: Final = MagicMock()
|
||||
cache.async_get_cache = AsyncMock()
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await get_access_object("group", client, cache, check_db_only=True)
|
||||
assert failure.value.status_code == 503
|
||||
assert failure.value.detail == "Access group policy is unavailable"
|
||||
cache.async_get_cache.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authoritative_team_permission_outage_cannot_drop_the_teams_restrictions() -> None:
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy.auth.auth_checks import get_team_object
|
||||
|
||||
row: Final = LiteLLM_TeamTable(team_id="team-policy-outage", object_permission_id="team-permission")
|
||||
client: Final = MagicMock()
|
||||
client.writer_db.litellm_teamtable.find_unique = AsyncMock(return_value=row)
|
||||
client.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(side_effect=RuntimeError("unavailable"))
|
||||
cache: Final = MagicMock()
|
||||
cache.async_get_cache = AsyncMock()
|
||||
cache.async_set_cache = AsyncMock()
|
||||
with pytest.raises(HTTPException) as failure:
|
||||
await get_team_object(row.team_id, client, cache, check_db_only=True)
|
||||
assert failure.value.status_code == 404
|
||||
client.writer_db.litellm_objectpermissiontable.find_unique.assert_awaited_once()
|
||||
cache.async_set_cache.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("strict", [True, False])
|
||||
@pytest.mark.parametrize("missing", [True, False])
|
||||
async def test_referenced_permission_failures_preserve_legacy_behavior_and_deny_strict_reads(strict, missing):
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy.auth.auth_checks import get_object_permission
|
||||
|
||||
client = MagicMock()
|
||||
lookup = AsyncMock(return_value=None, side_effect=None if missing else RuntimeError("unavailable"))
|
||||
client.writer_db.litellm_objectpermissiontable.find_unique = lookup
|
||||
client.db.litellm_objectpermissiontable.find_unique = lookup
|
||||
cache = MagicMock()
|
||||
cache.async_get_cache = AsyncMock(return_value=None)
|
||||
if strict:
|
||||
with pytest.raises(HTTPException if missing else RuntimeError):
|
||||
await get_object_permission("referenced", client, cache, check_db_only=True)
|
||||
cache.async_get_cache.assert_not_awaited()
|
||||
else:
|
||||
assert await get_object_permission("referenced", client, cache) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"models,key_aliases,team_aliases,allowed",
|
||||
[
|
||||
(["fast"], {}, {}, True),
|
||||
([], {}, {}, False),
|
||||
(["other"], {}, {}, False),
|
||||
(["target"], {"fast": "target"}, {}, True),
|
||||
(["target"], {}, {"fast": "target"}, True),
|
||||
(["fast"], {}, {"fast": "forbidden"}, False),
|
||||
],
|
||||
)
|
||||
async def test_managed_agent_model_policy_checks_dispatched_model(
|
||||
models: list[str], key_aliases: dict[str, str], team_aliases: dict[str, str], allowed: bool
|
||||
) -> None:
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy.auth.auth_checks import common_checks
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
agent: Final = AgentResponse(
|
||||
agent_id="managed", agent_name="Managed", agent_card_params={}, object_permission={"models": models}
|
||||
)
|
||||
auth: Final = UserAPIKeyAuth(
|
||||
token="test-token", team_id="team", aliases=key_aliases, team_model_aliases=team_aliases
|
||||
)
|
||||
auth.managed_agent_policy = agent
|
||||
checks: Final = common_checks(
|
||||
request_body={"model": "fast", "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=auth,
|
||||
request=MagicMock(spec=Request),
|
||||
)
|
||||
if allowed:
|
||||
assert await checks is True
|
||||
else:
|
||||
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)
|
||||
|
||||
|
||||
@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"]) == []
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
import json
|
||||
from collections.abc import AsyncIterator
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Final, cast
|
||||
import json
|
||||
from unittest.mock import patch
|
||||
|
||||
import httpx
|
||||
|
|
@ -12,9 +12,9 @@ from pydantic import ValidationError
|
|||
import litellm
|
||||
from litellm.exceptions import Timeout
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.llms.openai.responses.guardrail_translation.handler import OpenAIResponsesHandler
|
||||
from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket
|
||||
from litellm.proxy.guardrails.guardrail_hooks.crowdstrike_aidr import initialize_guardrail
|
||||
from litellm.proxy.guardrails.guardrail_hooks.crowdstrike_aidr.crowdstrike_aidr import (
|
||||
CrowdStrikeAIDRGuardrailMissingSecrets,
|
||||
|
|
@ -1805,36 +1805,22 @@ class _MessageShapedGuardrail(CustomGuardrail):
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("case", "instructions", "responses_input"),
|
||||
[
|
||||
(
|
||||
"instructions add a system message",
|
||||
"be terse",
|
||||
[{"role": "user", "content": [{"type": "input_text", "text": "my ssn is 078-05-1120"}]}],
|
||||
),
|
||||
(
|
||||
"tool items add messages that carry no text",
|
||||
None,
|
||||
[
|
||||
{"role": "user", "content": [{"type": "input_text", "text": "my ssn is 078-05-1120"}]},
|
||||
{"type": "function_call", "call_id": "c1", "name": "get_x", "arguments": "{}"},
|
||||
{"type": "function_call_output", "call_id": "c1", "output": "42"},
|
||||
],
|
||||
),
|
||||
],
|
||||
("case", "instructions"),
|
||||
[("tool items add messages that carry no text", None), ("instructions do not rescue the tool desync", "be terse")],
|
||||
)
|
||||
async def test_unalignable_rewrite_is_rejected_never_sent_unredacted(
|
||||
case: str,
|
||||
instructions: str | None,
|
||||
responses_input: list[dict[str, object]],
|
||||
) -> None:
|
||||
async def test_unalignable_rewrite_is_rejected_never_sent_unredacted(case: str, instructions: str | None) -> None:
|
||||
"""An unalignable rewrite must fail the request, not forward the raw prompt.
|
||||
|
||||
Skipping the write-back would hand the model the unredacted text, so a
|
||||
guardrail could be bypassed by adding ``instructions`` or a tool call.
|
||||
guardrail could be bypassed by adding a tool call.
|
||||
"""
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
|
||||
|
||||
responses_input: list[dict[str, object]] = [
|
||||
{"role": "user", "content": [{"type": "input_text", "text": "my ssn is 078-05-1120"}]},
|
||||
{"type": "function_call", "call_id": "c1", "name": "get_x", "arguments": "{}"},
|
||||
{"type": "function_call_output", "call_id": "c1", "output": "42"},
|
||||
]
|
||||
data: dict[str, object] = {"model": "gpt-4o", "input": responses_input}
|
||||
if instructions is not None:
|
||||
data["instructions"] = instructions
|
||||
|
|
@ -1846,21 +1832,27 @@ async def test_unalignable_rewrite_is_rejected_never_sent_unredacted(
|
|||
)
|
||||
|
||||
assert "078-05-1120" in str(responses_input), case
|
||||
assert data.get("instructions") == instructions, case
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aligned_rewrite_is_written_back() -> None:
|
||||
"""Matching counts must still redact the input in place."""
|
||||
@pytest.mark.parametrize("instructions", [None, "be terse"])
|
||||
async def test_aligned_rewrite_is_written_back(instructions: str | None) -> None:
|
||||
"""Matching counts must redact the input, and the instructions when present, in place."""
|
||||
responses_input: list[dict[str, object]] = [
|
||||
{"role": "user", "content": [{"type": "input_text", "text": "my ssn is 078-05-1120"}]}
|
||||
]
|
||||
data: dict[str, object] = {"model": "gpt-4o", "input": responses_input}
|
||||
if instructions is not None:
|
||||
data["instructions"] = instructions
|
||||
|
||||
await OpenAIResponsesHandler().process_input_messages(
|
||||
data={"model": "gpt-4o", "input": responses_input},
|
||||
data=data,
|
||||
guardrail_to_apply=_MessageShapedGuardrail("my ssn is <US_SSN>"),
|
||||
)
|
||||
|
||||
assert cast(list, responses_input[0]["content"])[0]["text"] == "my ssn is <US_SSN>"
|
||||
assert data.get("instructions") == (None if instructions is None else "my ssn is <US_SSN>")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -4867,7 +4867,39 @@ class TestPanwAirsLatestRoleMessageOnlyEveryRequestShape:
|
|||
assert result["input"][0]["content"] == "First user turn"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_flag_false_responses_scans_full_history(self):
|
||||
@pytest.mark.parametrize(
|
||||
"history_tail",
|
||||
[
|
||||
pytest.param((), id="plain"),
|
||||
pytest.param(
|
||||
({"type": "reasoning", "id": "rs_1", "summary": [{"type": "summary_text", "text": "thinking"}]},),
|
||||
id="reasoning",
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_flag_true_with_skip_system_still_scans_only_the_latest_turn_on_responses(
|
||||
self, history_tail: Sequence[Mapping[str, object]]
|
||||
) -> None:
|
||||
from litellm.llms.openai.responses.guardrail_translation.handler import (
|
||||
OpenAIResponsesHandler,
|
||||
)
|
||||
|
||||
handler = make_handler(experimental_use_latest_role_message_only=True)
|
||||
handler.skip_system_message_in_guardrail = True
|
||||
request_data = self._responses_request(
|
||||
{"role": "system", "content": "House rules"},
|
||||
*history_tail,
|
||||
{"role": "user", "content": self.LATEST},
|
||||
instructions="answer briefly",
|
||||
)
|
||||
patcher, mock_api = self._scan(handler)
|
||||
with patcher:
|
||||
await OpenAIResponsesHandler().process_input_messages(data=request_data, guardrail_to_apply=handler)
|
||||
|
||||
assert [call.kwargs["content"] for call in mock_api.call_args_list] == [self.LATEST]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_flag_false_responses_scans_instructions_and_full_history(self) -> None:
|
||||
from litellm.llms.openai.responses.guardrail_translation.handler import (
|
||||
OpenAIResponsesHandler,
|
||||
)
|
||||
|
|
@ -4878,7 +4910,11 @@ class TestPanwAirsLatestRoleMessageOnlyEveryRequestShape:
|
|||
with patcher:
|
||||
await OpenAIResponsesHandler().process_input_messages(data=request_data, guardrail_to_apply=handler)
|
||||
|
||||
assert [call.kwargs["content"] for call in mock_api.call_args_list] == ["First user turn", self.LATEST]
|
||||
assert [call.kwargs["content"] for call in mock_api.call_args_list] == [
|
||||
"answer briefly",
|
||||
"First user turn",
|
||||
self.LATEST,
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_flag_true_unalignable_texts_fall_back_to_scanning_everything(self):
|
||||
|
|
@ -4966,8 +5002,12 @@ class TestPanwAirsLatestRoleMessageOnlyEveryRequestShape:
|
|||
),
|
||||
],
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
"instructions",
|
||||
[pytest.param(None, id="no_instructions"), pytest.param("answer briefly", id="instructions")],
|
||||
)
|
||||
async def test_flag_true_reasoning_content_after_latest_user_turn_still_scans_that_turn(
|
||||
self, tail: Sequence[Mapping[str, object]]
|
||||
self, tail: Sequence[Mapping[str, object]], instructions: str | None
|
||||
):
|
||||
from litellm.llms.openai.responses.guardrail_translation.handler import (
|
||||
OpenAIResponsesHandler,
|
||||
|
|
@ -4983,6 +5023,7 @@ class TestPanwAirsLatestRoleMessageOnlyEveryRequestShape:
|
|||
"content": [{"type": "reasoning_text", "text": "model chain of thought"}],
|
||||
},
|
||||
*tail,
|
||||
**({"instructions": instructions} if instructions is not None else {}),
|
||||
)
|
||||
patcher, mock_api = self._scan(handler)
|
||||
with patcher:
|
||||
|
|
@ -5016,6 +5057,28 @@ class TestPanwAirsLatestRoleMessageOnlyEveryRequestShape:
|
|||
"thinking",
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_flag_true_texts_short_of_the_input_items_fall_back_to_scanning_everything(self) -> None:
|
||||
handler = make_handler(experimental_use_latest_role_message_only=True)
|
||||
reasoning = {"type": "reasoning", "id": "rs_1", "content": [{"type": "reasoning_text", "text": "thinking"}]}
|
||||
inputs: GenericGuardrailAPIInputs = {
|
||||
"texts": ["thinking", self.LATEST],
|
||||
"structured_messages": [{"role": "user", "content": "thinking"}, {"role": "user", "content": self.LATEST}],
|
||||
}
|
||||
request_data: dict[str, object] = {
|
||||
"litellm_call_id": "test-call-id",
|
||||
"input": [
|
||||
{"role": "user", "content": "First user turn"},
|
||||
reasoning,
|
||||
{"role": "user", "content": self.LATEST},
|
||||
],
|
||||
}
|
||||
patcher, mock_api = self._scan(handler)
|
||||
with patcher:
|
||||
await handler.apply_guardrail(inputs=inputs, request_data=request_data, input_type="request")
|
||||
|
||||
assert [call.kwargs["content"] for call in mock_api.call_args_list] == ["thinking", self.LATEST]
|
||||
|
||||
|
||||
class TestPanwAirsMcpToolCallWithoutCallId:
|
||||
"""Tests for MCP tool invocations flowing through apply_guardrail without
|
||||
|
|
|
|||
|
|
@ -7679,7 +7679,7 @@ class TestConnectedAppViewAnnotation:
|
|||
|
||||
flags = {server.server_id: server.connected_app_reachable for server in result}
|
||||
assert flags == {"server-1": True, "server-2": False}
|
||||
reload_mock.assert_awaited_once_with("test_user_id")
|
||||
reload_mock.assert_awaited_once_with("test_user_id", requires_fresh_policy=False)
|
||||
mock_manager.get_allowed_mcp_servers.assert_awaited_once_with(admitted_auth)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -9653,7 +9653,7 @@ class TestMCPServerResolutionCharacterization:
|
|||
server_id: str,
|
||||
) -> tuple[MagicMock, MCPServerManager, UserAPIKeyAuth]:
|
||||
team_id: Final = UI_SESSION_TOKEN_TEAM_ID if grant_route == "direct user object_permission" else "lit3974_team"
|
||||
user_id: Final = "lit3974_direct_user"
|
||||
user_id: Final = f"{server_id}:{grant_route}:user"
|
||||
key_permission: Final = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id=f"lit3974_{grant_route}_key_permission",
|
||||
mcp_servers=None,
|
||||
|
|
|
|||
|
|
@ -4355,7 +4355,7 @@ async def test_list_team_v2_org_admin_own_query_keeps_memberships_in_other_orgs(
|
|||
prisma_client.db.litellm_teamtable.find_many = AsyncMock(side_effect=find_many)
|
||||
prisma_client.db.litellm_teamtable.count = AsyncMock(side_effect=count)
|
||||
prisma_client.db.litellm_verificationtoken.group_by = AsyncMock(return_value=[])
|
||||
prisma_client.db.litellm_usertable.find_unique = AsyncMock(
|
||||
prisma_client.writer_db.litellm_usertable.find_unique = AsyncMock(
|
||||
return_value=LiteLLM_UserTable(
|
||||
user_id="org_admin_user",
|
||||
teams=["team_in_org_A", "team_in_org_B"],
|
||||
|
|
@ -4394,11 +4394,11 @@ async def test_list_team_v2_org_admin_own_query_keeps_memberships_in_other_orgs(
|
|||
assert await list_teams(None) == own_view
|
||||
assert await list_teams("org_admin_user", search="team_in_org_B") == ["team_in_org_B"]
|
||||
assert await list_teams("other_user") == ["other_team_in_org_A"]
|
||||
prisma_client.db.litellm_usertable.find_unique.assert_awaited_with(
|
||||
prisma_client.writer_db.litellm_usertable.find_unique.assert_awaited_with(
|
||||
where={"user_id": "org_admin_user"}, include={"organization_memberships": True}
|
||||
)
|
||||
|
||||
prisma_client.db.litellm_usertable.find_unique.side_effect = RuntimeError("db down")
|
||||
prisma_client.writer_db.litellm_usertable.find_unique.side_effect = RuntimeError("db down")
|
||||
with pytest.raises(ValueError, match="db down"):
|
||||
await list_teams("org_admin_user")
|
||||
|
||||
|
|
@ -15813,7 +15813,7 @@ async def test_get_team_spend_by_user_team_admin_sees_every_member(mock_db_clien
|
|||
alpha = _team_spend_by_user_team("team-alpha", "Team Alpha", Member(user_id="alice", role="admin"), [])
|
||||
mock_db_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[alpha])
|
||||
mock_db_client.db.query_raw = AsyncMock(return_value=[])
|
||||
mock_db_client.db.litellm_usertable.find_unique = AsyncMock(
|
||||
mock_db_client.writer_db.litellm_usertable.find_unique = AsyncMock(
|
||||
return_value=_team_spend_by_user_caller("alice", ["team-alpha"])
|
||||
)
|
||||
|
||||
|
|
@ -15835,7 +15835,7 @@ async def test_get_team_spend_by_user_plain_member_only_sees_own_row(mock_db_cli
|
|||
mock_db_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[alpha])
|
||||
mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
|
||||
mock_db_client.db.query_raw = AsyncMock(return_value=[_team_spend_by_user_db_row("team-alpha", "bob", 0.25, 2)])
|
||||
mock_db_client.db.litellm_usertable.find_unique = AsyncMock(
|
||||
mock_db_client.writer_db.litellm_usertable.find_unique = AsyncMock(
|
||||
return_value=_team_spend_by_user_caller("bob", ["team-alpha"])
|
||||
)
|
||||
|
||||
|
|
@ -15856,7 +15856,7 @@ async def test_get_team_spend_by_user_member_of_other_team_gets_404(mock_db_clie
|
|||
|
||||
caller = UserAPIKeyAuth(user_id="bob", user_role=LitellmUserRoles.INTERNAL_USER)
|
||||
mock_db_client.db.query_raw = AsyncMock(return_value=[])
|
||||
mock_db_client.db.litellm_usertable.find_unique = AsyncMock(
|
||||
mock_db_client.writer_db.litellm_usertable.find_unique = AsyncMock(
|
||||
return_value=_team_spend_by_user_caller("bob", ["team-alpha"])
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -6068,7 +6068,68 @@ async def test_websocket_passthrough_propagates_active_trace_context(
|
|||
propagated = get_current_span(TraceContextTextMapPropagator().extract(captured["headers"]))
|
||||
assert propagated.get_span_context().trace_id == span.get_span_context().trace_id
|
||||
assert propagated.get_span_context().span_id == span.get_span_context().span_id
|
||||
assert captured["headers"].get("authorization") == ("Bearer client" if forward_headers else None)
|
||||
assert "authorization" not in captured["headers"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_websocket_passthrough_never_forwards_caller_credentials_upstream(monkeypatch):
|
||||
from starlette.websockets import WebSocketState
|
||||
|
||||
captured: dict[str, dict[str, str]] = {}
|
||||
upstream_ws = FakeUpstreamWebSocket("{}")
|
||||
|
||||
def fake_connect(target, additional_headers):
|
||||
captured["headers"] = additional_headers
|
||||
return FakeUpstreamConnect(upstream_ws)
|
||||
|
||||
websocket = MagicMock()
|
||||
websocket.accept = AsyncMock()
|
||||
websocket.send_text = AsyncMock()
|
||||
websocket.send_bytes = AsyncMock()
|
||||
websocket.receive = AsyncMock(return_value={"type": "websocket.disconnect"})
|
||||
websocket.close = AsyncMock()
|
||||
websocket.headers = {
|
||||
"authorization": "Bearer sk-caller-virtual-key",
|
||||
"api-key": "sk-caller-virtual-key",
|
||||
"x-api-key": "sk-caller-virtual-key",
|
||||
"x-goog-api-key": "sk-caller-virtual-key",
|
||||
"x-goog-user-project": "caller-project",
|
||||
}
|
||||
websocket.client_state = WebSocketState.CONNECTED
|
||||
websocket.application_state = WebSocketState.CONNECTED
|
||||
|
||||
mock_proxy_logging = MagicMock()
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
mock_proxy_logging.post_call_success_hook = AsyncMock()
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock()
|
||||
mock_worker = MagicMock()
|
||||
mock_worker.ensure_initialized_and_enqueue = MagicMock(side_effect=lambda async_coroutine: async_coroutine.close())
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.connect",
|
||||
fake_connect,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.GLOBAL_LOGGING_WORKER",
|
||||
mock_worker,
|
||||
)
|
||||
await websocket_passthrough_request(
|
||||
websocket=websocket,
|
||||
target="wss://upstream.example.test/v1/realtime",
|
||||
custom_headers={
|
||||
"Authorization": "Bearer upstream-admin-secret",
|
||||
"x-api-key": "upstream-admin-key",
|
||||
},
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
forward_headers=True,
|
||||
endpoint="/realtime",
|
||||
accept_websocket=True,
|
||||
)
|
||||
|
||||
assert all("sk-caller-virtual-key" not in value for value in captured["headers"].values())
|
||||
assert captured["headers"]["Authorization"] == "Bearer upstream-admin-secret"
|
||||
assert captured["headers"]["x-api-key"] == "upstream-admin-key"
|
||||
assert captured["headers"]["x-goog-user-project"] == "caller-project"
|
||||
|
||||
|
||||
class ClosingUpstreamWebSocket:
|
||||
|
|
|
|||
|
|
@ -612,7 +612,7 @@ def test_sanitize_request_body_for_spend_logs_payload_mixed_types():
|
|||
request_body = {
|
||||
"text": long_string,
|
||||
"number": 42,
|
||||
"nested": {"list": ["short", long_string], "dict": {"key": long_string}},
|
||||
"nested": {"list": ["short", long_string], "dict": {"value": long_string}},
|
||||
}
|
||||
sanitized = _sanitize_request_body_for_spend_logs_payload(request_body)
|
||||
|
||||
|
|
@ -631,7 +631,7 @@ def test_sanitize_request_body_for_spend_logs_payload_mixed_types():
|
|||
assert sanitized["number"] == 42
|
||||
assert sanitized["nested"]["list"][0] == "short"
|
||||
assert len(sanitized["nested"]["list"][1]) == expected_length
|
||||
assert len(sanitized["nested"]["dict"]["key"]) == expected_length
|
||||
assert len(sanitized["nested"]["dict"]["value"]) == expected_length
|
||||
|
||||
|
||||
def test_sanitize_request_body_for_spend_logs_payload_uses_runtime_env_override(
|
||||
|
|
@ -1207,7 +1207,7 @@ def test_get_logging_payload_placeholders_the_metadata_copied_into_the_stored_re
|
|||
stored_request_body: Final = json.loads(payload["proxy_server_request"])
|
||||
assert stored_request_body["metadata"]["model_group"] == expected_stored_model_group
|
||||
assert stored_request_body["metadata"]["error_information"]["error_message"] == expected_stored_error_message
|
||||
assert stored_request_body["metadata"]["user_api_key"] == "sk-test"
|
||||
assert stored_request_body["metadata"]["user_api_key"] == REDACTED_BY_LITELM_STRING
|
||||
assert ("medical records" in payload["proxy_server_request"]) == bool(deployment_info)
|
||||
|
||||
|
||||
|
|
@ -2691,6 +2691,104 @@ def test_sanitize_request_body_strips_secret_fields():
|
|||
assert sanitized["messages"] == [{"role": "user", "content": "hi"}]
|
||||
|
||||
|
||||
@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs")
|
||||
def test_proxy_server_request_payload_strips_nested_aws_credentials(mock_should_store: MagicMock) -> None:
|
||||
mock_should_store.return_value = True
|
||||
credentials: Final = {
|
||||
"aws_access_key_id": "AKIA-canary",
|
||||
"aws_secret_access_key": "secret-canary",
|
||||
"aws_session_token": "token-canary",
|
||||
"aws_web_identity_token": "wit-canary",
|
||||
}
|
||||
tool_parameters: Final = {"type": "object", "properties": {"aws_secret_access_key": {"type": "string"}}}
|
||||
litellm_params: Final = {
|
||||
"proxy_server_request": {
|
||||
"body": {
|
||||
"model": "bedrock-claude",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"fallbacks": [{"model": "bedrock-b", "aws_region_name": "us-west-2", **credentials}],
|
||||
"extra_body": {"aws_role_name": "arn:aws:iam::123456789012:role/r", **credentials},
|
||||
"tools": [{"type": "function", "function": {"name": "f", "parameters": tool_parameters}}],
|
||||
**credentials,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
parsed: Final = json.loads(
|
||||
_get_proxy_server_request_for_spend_logs_payload(metadata={}, litellm_params=litellm_params, kwargs={})
|
||||
)
|
||||
|
||||
assert "canary" not in json.dumps(parsed)
|
||||
masked: Final = dict.fromkeys(credentials, REDACTED_BY_LITELM_STRING)
|
||||
assert parsed["fallbacks"] == [{"model": "bedrock-b", "aws_region_name": "us-west-2", **masked}]
|
||||
assert parsed["extra_body"] == {"aws_role_name": "arn:aws:iam::123456789012:role/r", **masked}
|
||||
assert {name: parsed[name] for name in credentials} == masked
|
||||
assert parsed["tools"][0]["function"]["parameters"] == tool_parameters
|
||||
assert parsed["messages"] == [{"role": "user", "content": "hello"}]
|
||||
|
||||
|
||||
@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs")
|
||||
def test_proxy_server_request_payload_redacts_provider_credentials(mock_should_store: MagicMock) -> None:
|
||||
mock_should_store.return_value = True
|
||||
credentials: Final = {
|
||||
"azure_password": "canary-azure-password",
|
||||
"client_secret": "canary-client-secret",
|
||||
"azure_ad_token": "canary-azure-ad-token",
|
||||
"vertex_credentials": "canary-vertex-credentials",
|
||||
"s3_secret_access_key": "canary-s3-secret",
|
||||
"token": "canary-watsonx-token",
|
||||
"apikey": "canary-watsonx-apikey",
|
||||
"zen_api_key": "canary-zen-api-key",
|
||||
"gemini_api_key": "canary-gemini-api-key",
|
||||
"gigachat_access_token": "canary-gigachat-token",
|
||||
"oci_key": "canary-oci-key",
|
||||
}
|
||||
metadata: Final = {"user_api_key": "custom-auth-raw-key", "requester_ip_address": "10.0.0.1"}
|
||||
tool_parameters: Final = {"type": "object", "properties": {"client_secret": {"type": "string"}}}
|
||||
litellm_params: Final = {
|
||||
"proxy_server_request": {
|
||||
"body": {
|
||||
"model": "azure-gpt",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
"max_tokens": 10,
|
||||
"prompt_cache_key": "user-123-cache",
|
||||
"vertex_credentials": {"private_key": "canary-private-key", "client_email": "sa@example.com"},
|
||||
"extra_headers": {"Authorization": "Bearer canary-extra-header"},
|
||||
"tools": [
|
||||
{"type": "function", "function": {"name": "f", "parameters": tool_parameters}},
|
||||
{"type": "mcp", "server_url": "https://mcp.example.com", "headers": {"Authorization": "canary-mcp"}},
|
||||
],
|
||||
"fallbacks": [{"model": "azure-b", **credentials}],
|
||||
"metadata": metadata,
|
||||
**credentials,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
parsed: Final = json.loads(
|
||||
_get_proxy_server_request_for_spend_logs_payload(metadata={}, litellm_params=litellm_params, kwargs={})
|
||||
)
|
||||
|
||||
assert "canary" not in json.dumps(parsed)
|
||||
assert {name: parsed[name] for name in credentials} == dict.fromkeys(credentials, REDACTED_BY_LITELM_STRING)
|
||||
assert parsed["vertex_credentials"] == REDACTED_BY_LITELM_STRING
|
||||
assert parsed["extra_headers"] == {"Authorization": REDACTED_BY_LITELM_STRING}
|
||||
assert parsed["tools"][0]["function"]["parameters"] == tool_parameters
|
||||
assert parsed["tools"][1]["server_url"] == "https://mcp.example.com"
|
||||
assert parsed["metadata"] == {"user_api_key": REDACTED_BY_LITELM_STRING, "requester_ip_address": "10.0.0.1"}
|
||||
assert parsed["max_tokens"] == 10
|
||||
assert parsed["prompt_cache_key"] == REDACTED_BY_LITELM_STRING
|
||||
assert parsed["messages"] == [{"role": "user", "content": "hello"}]
|
||||
|
||||
|
||||
def test_sanitize_response_redacts_credential_named_fields() -> None:
|
||||
response: Final = {"access_token": "canary-oauth-token", "usage": {"prompt_tokens": 1}}
|
||||
|
||||
assert _sanitize_request_body_for_spend_logs_payload({"response": response}) == {
|
||||
"response": {"access_token": REDACTED_BY_LITELM_STRING, "usage": {"prompt_tokens": 1}}
|
||||
}
|
||||
|
||||
|
||||
@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs")
|
||||
def test_proxy_server_request_payload_excludes_secret_fields(mock_should_store):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -9,16 +9,13 @@ Run with: pytest tests/unit/interactions/test_openapi_compliance.py -v
|
|||
|
||||
import json
|
||||
import os
|
||||
from typing import Any, Dict, Final
|
||||
from typing import Any, Dict
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from openapi_core import OpenAPI
|
||||
|
||||
from litellm.llms.gemini.interactions.transformation import GoogleAIStudioInteractionsConfig
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
OPENAPI_SPEC_URL = "https://ai.google.dev/static/api/interactions.openapi.json"
|
||||
|
||||
|
||||
|
|
@ -35,7 +32,8 @@ def _load_openapi_spec_dict() -> Dict[str, Any]:
|
|||
return response.json()
|
||||
except Exception as e: # pragma: no cover - defensive, env-dependent
|
||||
pytest.skip(
|
||||
f"Skipping Google Interactions OpenAPI compliance tests - unable to load spec from {OPENAPI_SPEC_URL}: {e}"
|
||||
f"Skipping Google Interactions OpenAPI compliance tests - "
|
||||
f"unable to load spec from {OPENAPI_SPEC_URL}: {e}"
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -46,20 +44,6 @@ def _declared_type_value(variant_schema: Dict[str, Any]) -> Any:
|
|||
return type_property.get("const") or (enum_values[0] if len(enum_values) == 1 else None)
|
||||
|
||||
|
||||
def _model_request_schema(spec: Dict[str, Any]) -> Dict[str, Any]:
|
||||
create: Final = next(
|
||||
methods["post"]
|
||||
for path, methods in spec["paths"].items()
|
||||
if path.endswith("/interactions") and "post" in methods
|
||||
)
|
||||
schema: Final = create["requestBody"]["content"]["application/json"]["schema"]
|
||||
variants: Final = schema.get("oneOf", [schema])
|
||||
resolved: Final = tuple(
|
||||
spec["components"]["schemas"][item["$ref"].split("/")[-1]] if "$ref" in item else item for item in variants
|
||||
)
|
||||
return next(item for item in resolved if "model" in item.get("properties", {}))
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def spec_dict() -> Dict[str, Any]:
|
||||
"""Load raw spec dict for manual validation."""
|
||||
|
|
@ -77,22 +61,11 @@ class TestRequestCompliance:
|
|||
|
||||
def test_create_model_interaction_request_schema(self, spec_dict):
|
||||
"""Verify CreateModelInteractionParams schema fields."""
|
||||
schema = _model_request_schema(spec_dict)
|
||||
schema = spec_dict["components"]["schemas"]["CreateModelInteractionParams"]
|
||||
|
||||
assert "model" in schema["properties"]
|
||||
assert "input" in schema["properties"]
|
||||
|
||||
request: Final = GoogleAIStudioInteractionsConfig().transform_request(
|
||||
model="gemini/test-model",
|
||||
agent=None,
|
||||
input="Hello",
|
||||
optional_params={},
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
assert request["model"] == "test-model"
|
||||
assert request["input"] == "Hello"
|
||||
assert "model" in schema.get("required", ())
|
||||
# Required fields per spec
|
||||
assert "model" in schema["required"]
|
||||
assert "input" in schema["required"]
|
||||
|
||||
# Check our supported optional fields exist in spec
|
||||
our_optional_fields = [
|
||||
|
|
@ -115,7 +88,7 @@ class TestRequestCompliance:
|
|||
|
||||
def test_input_types_match_spec(self, spec_dict):
|
||||
"""Verify input field supports string, Content, Content[], Turn[]."""
|
||||
schema = _model_request_schema(spec_dict)
|
||||
schema = spec_dict["components"]["schemas"]["CreateModelInteractionParams"]
|
||||
input_schema = schema["properties"]["input"]
|
||||
|
||||
# The input property may be inline oneOf or a $ref to InteractionsInput
|
||||
|
|
@ -152,18 +125,22 @@ class TestRequestCompliance:
|
|||
|
||||
discriminator = content_schema.get("discriminator")
|
||||
if discriminator is not None:
|
||||
assert discriminator.get("propertyName") == "type", (
|
||||
f"Content is discriminated on {discriminator.get('propertyName')!r}, not 'type'"
|
||||
)
|
||||
assert (
|
||||
discriminator.get("propertyName") == "type"
|
||||
), f"Content is discriminated on {discriminator.get('propertyName')!r}, not 'type'"
|
||||
|
||||
variant_names = [
|
||||
option["$ref"].split("/")[-1] for option in content_schema.get("oneOf", []) if "$ref" in option
|
||||
option["$ref"].split("/")[-1]
|
||||
for option in content_schema.get("oneOf", [])
|
||||
if "$ref" in option
|
||||
]
|
||||
assert variant_names, f"Content is not a union of named variants: {content_schema}"
|
||||
|
||||
mapping = (discriminator or {}).get("mapping") or {}
|
||||
type_values = {
|
||||
variant: mapping_value for mapping_value, ref in mapping.items() for variant in [ref.split("/")[-1]]
|
||||
variant: mapping_value
|
||||
for mapping_value, ref in mapping.items()
|
||||
for variant in [ref.split("/")[-1]]
|
||||
} or {
|
||||
variant: _declared_type_value(spec_dict["components"]["schemas"].get(variant, {}))
|
||||
for variant in variant_names
|
||||
|
|
@ -214,9 +191,7 @@ class TestRequestCompliance:
|
|||
for option in spec_dict["components"]["schemas"]["Step"]["oneOf"]
|
||||
if "$ref" in option
|
||||
}
|
||||
assert {"UserInputStep", "ModelOutputStep"} <= step_variants, (
|
||||
f"Step union is missing role steps: {step_variants}"
|
||||
)
|
||||
assert {"UserInputStep", "ModelOutputStep"} <= step_variants, f"Step union is missing role steps: {step_variants}"
|
||||
|
||||
for step_name, type_value in [("UserInputStep", "user_input"), ("ModelOutputStep", "model_output")]:
|
||||
step_schema = spec_dict["components"]["schemas"][step_name]
|
||||
|
|
@ -286,7 +261,9 @@ class TestResponseCompliance:
|
|||
expected_fields = ["total_input_tokens", "total_output_tokens", "total_tokens"]
|
||||
|
||||
for field in expected_fields:
|
||||
assert field in usage_schema["properties"], f"Usage field '{field}' not in spec"
|
||||
assert (
|
||||
field in usage_schema["properties"]
|
||||
), f"Usage field '{field}' not in spec"
|
||||
print(f"✓ Usage field '{field}' exists")
|
||||
|
||||
|
||||
|
|
@ -305,7 +282,9 @@ class TestToolsCompliance:
|
|||
"""Verify FunctionDeclaration schema for function tools."""
|
||||
if "FunctionDeclaration" in spec_dict["components"]["schemas"]:
|
||||
func_schema = spec_dict["components"]["schemas"]["FunctionDeclaration"]
|
||||
assert "name" in func_schema.get("properties", {}) or "name" in func_schema.get("required", [])
|
||||
assert "name" in func_schema.get(
|
||||
"properties", {}
|
||||
) or "name" in func_schema.get("required", [])
|
||||
print("✓ FunctionDeclaration schema found")
|
||||
else:
|
||||
print("⚠ FunctionDeclaration schema not found (may be nested)")
|
||||
|
|
@ -334,7 +313,7 @@ class TestEndpointCompliance:
|
|||
|
||||
get_path = None
|
||||
for path, methods in paths.items():
|
||||
if "/interactions/{" in path and path.endswith("}") and "get" in methods:
|
||||
if "{id}" in path and "interactions" in path and "get" in methods:
|
||||
get_path = path
|
||||
break
|
||||
|
||||
|
|
@ -347,7 +326,7 @@ class TestEndpointCompliance:
|
|||
|
||||
delete_path = None
|
||||
for path, methods in paths.items():
|
||||
if "/interactions/{" in path and path.endswith("}") and "delete" in methods:
|
||||
if "{id}" in path and "interactions" in path and "delete" in methods:
|
||||
delete_path = path
|
||||
break
|
||||
|
||||
|
|
@ -371,4 +350,6 @@ if __name__ == "__main__":
|
|||
if method in ["get", "post", "delete", "put", "patch"]:
|
||||
print(f" {method.upper()} {path}")
|
||||
|
||||
print(f"\nSchemas: {list(spec.get('components', {}).get('schemas', {}).keys())[:10]}...")
|
||||
print(
|
||||
f"\nSchemas: {list(spec.get('components', {}).get('schemas', {}).keys())[:10]}..."
|
||||
)
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ with guardrail transformations.
|
|||
|
||||
import copy
|
||||
from collections.abc import Callable
|
||||
from typing import Any, List, Literal, Optional, Tuple
|
||||
from typing import Any, Final, List, Literal, Optional, Tuple
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import logging
|
||||
|
|
@ -67,6 +67,55 @@ class MockGuardrail(CustomGuardrail):
|
|||
return inputs
|
||||
|
||||
|
||||
class RecordingMaskingGuardrail(MockGuardrail):
|
||||
"""MockGuardrail that also records the texts and structured message contents it was shown"""
|
||||
|
||||
def __init__(self, guardrail_name: str) -> None:
|
||||
super().__init__(guardrail_name=guardrail_name)
|
||||
self.seen_texts: list[list[str]] = []
|
||||
self.seen_message_contents: list[list[object]] = []
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict[str, object],
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: LiteLLMLoggingObj | None = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
self.seen_texts.append(list(inputs.get("texts", [])))
|
||||
self.seen_message_contents.append([m["content"] for m in inputs.get("structured_messages") or []])
|
||||
return await super().apply_guardrail(inputs, request_data, input_type, logging_obj)
|
||||
|
||||
|
||||
class LastTextDroppingGuardrail(CustomGuardrail):
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict[str, object],
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: LiteLLMLoggingObj | None = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
return {**inputs, "texts": list(inputs.get("texts", []))[:-1]}
|
||||
|
||||
|
||||
class TextsReplacingGuardrail(CustomGuardrail):
|
||||
"""Answers with the given texts list, or without a texts key at all when given None"""
|
||||
|
||||
def __init__(self, guardrail_name: str, texts: tuple[str, ...] | None) -> None:
|
||||
super().__init__(guardrail_name=guardrail_name)
|
||||
self.texts: Final = texts
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict[str, object],
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: LiteLLMLoggingObj | None = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
answer: Final = {key: value for key, value in inputs.items() if key != "texts"}
|
||||
return answer if self.texts is None else {**answer, "texts": list(self.texts)}
|
||||
|
||||
|
||||
class PersimmonMaskingGuardrail(CustomGuardrail):
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
|
|
@ -217,15 +266,9 @@ class TestOpenAIResponsesHandlerInputProcessing:
|
|||
|
||||
result = await handler.process_input_messages(data, guardrail)
|
||||
|
||||
assert (
|
||||
result["input"][0]["content"][0]["text"]
|
||||
== "Describe this image [GUARDRAILED]"
|
||||
)
|
||||
assert result["input"][0]["content"][0]["text"] == "Describe this image [GUARDRAILED]"
|
||||
# Image URL should remain unchanged
|
||||
assert (
|
||||
result["input"][0]["content"][1]["image_url"]["url"]
|
||||
== "https://example.com/image.jpg"
|
||||
)
|
||||
assert result["input"][0]["content"][1]["image_url"]["url"] == "https://example.com/image.jpg"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_input_with_empty_content(self):
|
||||
|
|
@ -248,6 +291,217 @@ class TestOpenAIResponsesHandlerInputProcessing:
|
|||
# Empty string should be processed
|
||||
assert result["input"][1]["content"] == " [GUARDRAILED]"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_instructions_over_string_input_are_scanned_first_and_rewritten_in_place(self) -> None:
|
||||
handler = OpenAIResponsesHandler()
|
||||
guardrail = RecordingMaskingGuardrail(guardrail_name="test")
|
||||
data = {"model": "gpt-4", "instructions": "Be terse", "input": "Hello"}
|
||||
|
||||
result = await handler.process_input_messages(data, guardrail)
|
||||
|
||||
assert guardrail.seen_texts == [["Be terse", "Hello"]]
|
||||
assert guardrail.seen_message_contents == [["Be terse", "Hello"]]
|
||||
assert result["instructions"] == "Be terse [GUARDRAILED]"
|
||||
assert result["input"] == "Hello [GUARDRAILED]"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_instructions_over_list_input_are_scanned_first_and_rewritten_in_place(self) -> None:
|
||||
handler = OpenAIResponsesHandler()
|
||||
guardrail = RecordingMaskingGuardrail(guardrail_name="test")
|
||||
data = {
|
||||
"model": "gpt-4",
|
||||
"instructions": "Be terse",
|
||||
"input": [
|
||||
{"role": "user", "content": "Hello"},
|
||||
{"role": "user", "content": [{"type": "input_text", "text": "World"}]},
|
||||
],
|
||||
}
|
||||
|
||||
result = await handler.process_input_messages(data, guardrail)
|
||||
|
||||
assert guardrail.seen_texts == [["Be terse", "Hello", "World"]]
|
||||
assert guardrail.seen_message_contents == [["Be terse", "Hello", [{"type": "text", "text": "World"}]]]
|
||||
assert result["instructions"] == "Be terse [GUARDRAILED]"
|
||||
assert result["input"] == [
|
||||
{"role": "user", "content": "Hello [GUARDRAILED]"},
|
||||
{"role": "user", "content": [{"type": "input_text", "text": "World [GUARDRAILED]"}]},
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_instructions_are_not_scanned(self) -> None:
|
||||
handler = OpenAIResponsesHandler()
|
||||
guardrail = RecordingMaskingGuardrail(guardrail_name="test")
|
||||
data = {"model": "gpt-4", "instructions": "", "input": "Hello"}
|
||||
|
||||
result = await handler.process_input_messages(data, guardrail)
|
||||
|
||||
assert guardrail.seen_texts == [["Hello"]]
|
||||
assert result["instructions"] == ""
|
||||
assert result["input"] == "Hello [GUARDRAILED]"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_text_answer_missing_the_instructions_row_is_rejected_and_leaves_request_untouched(self) -> None:
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
|
||||
|
||||
handler = OpenAIResponsesHandler()
|
||||
guardrail = LastTextDroppingGuardrail(guardrail_name="dropper")
|
||||
data = {"model": "gpt-4", "instructions": "Be terse", "input": [{"role": "user", "content": "Hello"}]}
|
||||
original = copy.deepcopy(data)
|
||||
|
||||
with pytest.raises(UnappliableRequestRewrite) as excinfo:
|
||||
await handler.process_input_messages(data, guardrail)
|
||||
|
||||
assert excinfo.value.guardrail_name == "dropper"
|
||||
assert data["instructions"] == original["instructions"]
|
||||
assert data["input"] == original["input"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("answered_texts", [None, ()], ids=["no_texts_key", "empty_texts"])
|
||||
@pytest.mark.parametrize("data_input", ["Hello", [{"role": "user", "content": "Hello"}]])
|
||||
async def test_answer_without_texts_leaves_instructions_and_input_untouched_like_chat_completions(
|
||||
self, answered_texts: tuple[str, ...] | None, data_input: str | list[dict[str, str]]
|
||||
) -> None:
|
||||
handler = OpenAIResponsesHandler()
|
||||
guardrail = TextsReplacingGuardrail(guardrail_name="silent", texts=answered_texts)
|
||||
data = {"model": "gpt-4", "instructions": "Be terse", "input": data_input}
|
||||
original = copy.deepcopy(data)
|
||||
|
||||
result = await handler.process_input_messages(data, guardrail)
|
||||
|
||||
assert result["instructions"] == original["instructions"]
|
||||
assert result["input"] == original["input"]
|
||||
|
||||
|
||||
def _skipping_system(guardrail: CustomGuardrail) -> CustomGuardrail:
|
||||
guardrail.skip_system_message_in_guardrail = True
|
||||
return guardrail
|
||||
|
||||
|
||||
class TestSkipSystemMessageScopesInstructions:
|
||||
"""skip_system_message_in_guardrail keeps the Responses system prompt out of the scan the same
|
||||
way it keeps chat `system` messages and Anthropic top-level `system` out: instructions and
|
||||
system-role input items leave both texts and structured_messages, and rewrites leave them verbatim."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("data_input", ["Hello", [{"role": "user", "content": "Hello"}]])
|
||||
async def test_instructions_are_neither_scanned_nor_rewritten(self, data_input: str | list[dict[str, str]]) -> None:
|
||||
handler = OpenAIResponsesHandler()
|
||||
guardrail = _skipping_system(RecordingMaskingGuardrail(guardrail_name="test"))
|
||||
data = {"model": "gpt-4", "instructions": "Be terse", "input": data_input}
|
||||
|
||||
result = await handler.process_input_messages(data, guardrail)
|
||||
|
||||
assert guardrail.seen_texts == [["Hello"]]
|
||||
assert guardrail.seen_message_contents == [["Hello"]]
|
||||
assert result["instructions"] == "Be terse"
|
||||
rewritten = result["input"][0]["content"] if isinstance(data_input, list) else result["input"]
|
||||
assert rewritten == "Hello [GUARDRAILED]"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_system_input_items_leave_scope_and_user_items_still_align_with_structured_messages(self) -> None:
|
||||
handler = OpenAIResponsesHandler()
|
||||
guardrail = _skipping_system(RecordingMaskingGuardrail(guardrail_name="test"))
|
||||
data = {
|
||||
"model": "gpt-4",
|
||||
"instructions": "Be terse",
|
||||
"input": [
|
||||
{"role": "system", "content": "House rules"},
|
||||
{"role": "developer", "content": "Dev note"},
|
||||
{"role": "user", "content": [{"type": "input_text", "text": "World"}]},
|
||||
],
|
||||
}
|
||||
|
||||
result = await handler.process_input_messages(data, guardrail)
|
||||
|
||||
assert guardrail.seen_texts == [["Dev note", "World"]]
|
||||
assert guardrail.seen_message_contents == [["Dev note", [{"type": "text", "text": "World"}]]]
|
||||
assert result["instructions"] == "Be terse"
|
||||
assert result["input"] == [
|
||||
{"role": "system", "content": "House rules"},
|
||||
{"role": "developer", "content": "Dev note [GUARDRAILED]"},
|
||||
{"role": "user", "content": [{"type": "input_text", "text": "World [GUARDRAILED]"}]},
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_only_system_content_means_nothing_is_scanned(self) -> None:
|
||||
handler = OpenAIResponsesHandler()
|
||||
guardrail = _skipping_system(RecordingMaskingGuardrail(guardrail_name="test"))
|
||||
data = {"model": "gpt-4", "instructions": "Be terse", "input": [{"role": "system", "content": "Rules"}]}
|
||||
original = copy.deepcopy(data)
|
||||
|
||||
result = await handler.process_input_messages(data, guardrail)
|
||||
|
||||
assert guardrail.seen_texts == []
|
||||
assert result == original
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_structured_rewrite_of_the_scoped_rows_keeps_the_skipped_system_prompt(self) -> None:
|
||||
handler = OpenAIResponsesHandler()
|
||||
data = {
|
||||
"model": "gpt-5.6",
|
||||
"instructions": "Answer from the memo only.",
|
||||
"input": [
|
||||
{"role": "system", "content": "House rules"},
|
||||
{"role": "user", "content": "memo " * 400},
|
||||
{"role": "assistant", "content": "Understood."},
|
||||
{"role": "user", "content": "What is the codename?"},
|
||||
],
|
||||
}
|
||||
|
||||
result = await handler.process_input_messages(data, _skipping_system(StructuredRewriteGuardrail()))
|
||||
|
||||
assert result["instructions"] == "Answer from the memo only."
|
||||
assert [(item["role"], _texts(item)) for item in result["input"]] == [
|
||||
("system", ["House rules"]),
|
||||
("user", [COMPRESSED_MARKER]),
|
||||
("assistant", ["Understood."]),
|
||||
("user", ["What is the codename?"]),
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_full_coverage_claim_over_only_the_scoped_rows_still_keeps_the_skipped_system_prompt(self) -> None:
|
||||
handler = OpenAIResponsesHandler()
|
||||
data = {
|
||||
"model": "gpt-5.6",
|
||||
"instructions": "Answer from the memo only.",
|
||||
"input": [
|
||||
{"role": "system", "content": "House rules"},
|
||||
{"role": "user", "content": "memo " * 400},
|
||||
{"role": "user", "content": "What is the codename?"},
|
||||
],
|
||||
}
|
||||
|
||||
result = await handler.process_input_messages(data, _skipping_system(ScopedRowsFullCoverageGuardrail()))
|
||||
|
||||
assert result["instructions"] == "Answer from the memo only."
|
||||
assert [(item["role"], _texts(item)) for item in result["input"]] == [
|
||||
("system", ["House rules"]),
|
||||
("user", [COMPRESSED_MARKER]),
|
||||
("user", ["What is the codename?"]),
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_full_coverage_claim_over_the_whole_request_is_installed_without_a_second_merge(self) -> None:
|
||||
handler = OpenAIResponsesHandler()
|
||||
data = {
|
||||
"model": "gpt-5.6",
|
||||
"instructions": "Answer from the memo only.",
|
||||
"input": [
|
||||
{"role": "system", "content": "House rules"},
|
||||
{"role": "user", "content": "memo " * 400},
|
||||
{"role": "user", "content": "What is the codename?"},
|
||||
],
|
||||
}
|
||||
|
||||
result = await handler.process_input_messages(data, _skipping_system(RebuildingFullCoverageGuardrail()))
|
||||
|
||||
assert result["instructions"] == "Answer from the memo only."
|
||||
assert [(item["role"], _texts(item)) for item in result["input"]] == [
|
||||
("system", ["House rules"]),
|
||||
("user", [COMPRESSED_MARKER]),
|
||||
("user", ["What is the codename?"]),
|
||||
]
|
||||
|
||||
|
||||
class TestOpenAIResponsesHandlerOutputProcessing:
|
||||
"""Test output processing functionality"""
|
||||
|
|
@ -2156,6 +2410,36 @@ class StructuredRewriteGuardrail(CustomGuardrail):
|
|||
return {**inputs, "structured_messages": rewritten}
|
||||
|
||||
|
||||
class ScopedRowsFullCoverageGuardrail(StructuredRewriteGuardrail):
|
||||
"""Claims its structured_messages span the whole request but, like CrowdStrike AIDR on a
|
||||
Responses body (no `messages` to rebuild from), only ever returns the scoped rows it was given."""
|
||||
|
||||
def structured_messages_cover_full_request(self) -> bool:
|
||||
return True
|
||||
|
||||
|
||||
class RebuildingFullCoverageGuardrail(CustomGuardrail):
|
||||
"""Claims full coverage and honours it: rebuilds every conversation row from the raw request,
|
||||
compressing the first user turn, the way CrowdStrike AIDR does on a chat body."""
|
||||
|
||||
def structured_messages_cover_full_request(self) -> bool:
|
||||
return True
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict[str, object],
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: LiteLLMLoggingObj | None = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
raw_input = request_data["input"]
|
||||
assert isinstance(raw_input, list)
|
||||
full: list[dict[str, object]] = [{"role": "system", "content": request_data["instructions"]}, *raw_input]
|
||||
first_user = next(i for i, m in enumerate(full) if m.get("role") == "user")
|
||||
rewritten = [{**m, "content": COMPRESSED_MARKER} if i == first_user else m for i, m in enumerate(full)]
|
||||
return {**inputs, "structured_messages": rewritten}
|
||||
|
||||
|
||||
class ToolOutputRewriteGuardrail(CustomGuardrail):
|
||||
"""Guardrail that compresses the first tool-result row, the way Headroom does."""
|
||||
|
||||
|
|
@ -2527,8 +2811,9 @@ def _string_input_request() -> dict:
|
|||
class TestPerMessageRewriteWriteBack:
|
||||
"""A guardrail that rewrites per chat row hands the rows back as
|
||||
structured_messages, and the handler lands them on the instructions and the
|
||||
input items they came from; the same rewrite handed back as texts alone has
|
||||
no item to land on and is rejected by name instead of sent unrewritten."""
|
||||
input items they came from; the same rewrite handed back as texts alone lands
|
||||
only where every row has a scanned text (instructions plus a string input) and
|
||||
is otherwise rejected by name instead of sent unrewritten."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_structured_rows_land_on_instructions_and_tool_output(self):
|
||||
|
|
@ -2576,20 +2861,15 @@ class TestPerMessageRewriteWriteBack:
|
|||
assert [_texts(item) for item in result["input"]] == [["My SSN is " + REDACTED_SSN + "."]]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_texts_only_per_message_answer_over_a_string_input_is_rejected_by_name(self):
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
|
||||
|
||||
async def test_texts_only_per_message_answer_over_a_string_input_lands_on_instructions_and_input(self) -> None:
|
||||
guardrail = _per_message_redactor()
|
||||
data = _string_input_request()
|
||||
original = copy.deepcopy(data)
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", side_effect=_per_message_guardrail_server(False)):
|
||||
with pytest.raises(UnappliableRequestRewrite) as excinfo:
|
||||
await OpenAIResponsesHandler().process_input_messages(data, guardrail)
|
||||
result = await OpenAIResponsesHandler().process_input_messages(data, guardrail)
|
||||
|
||||
assert excinfo.value.guardrail_name == "per-message-redactor"
|
||||
assert data["input"] == original["input"]
|
||||
assert data["instructions"] == original["instructions"]
|
||||
assert result["instructions"] == "Never repeat the SSN " + REDACTED_SSN + " back."
|
||||
assert result["input"] == "My SSN is " + REDACTED_SSN + "."
|
||||
|
||||
|
||||
class TestProvenancePatching:
|
||||
|
|
|
|||
|
|
@ -1,11 +1,11 @@
|
|||
import logging
|
||||
import os
|
||||
from typing import Final, cast
|
||||
|
||||
import pytest
|
||||
from dotenv import load_dotenv
|
||||
|
||||
load_dotenv()
|
||||
import io
|
||||
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
|
||||
|
|
@ -20,7 +20,9 @@ def test_encrypt_decrypt_with_master_key():
|
|||
assert decrypt_value_helper(encrypt_value_helper(10), key="test_key") == 10
|
||||
assert decrypt_value_helper(encrypt_value_helper(True), key="test_key") is True
|
||||
assert decrypt_value_helper(encrypt_value_helper(None), key="test_key") is None
|
||||
assert decrypt_value_helper(encrypt_value_helper({"rpm": 10}), key="test_key") == {"rpm": 10}
|
||||
assert decrypt_value_helper(encrypt_value_helper({"rpm": 10}), key="test_key") == {
|
||||
"rpm": 10
|
||||
}
|
||||
|
||||
# encryption should actually occur for strings
|
||||
assert encrypt_value_helper("test") != "test"
|
||||
|
|
@ -28,25 +30,16 @@ def test_encrypt_decrypt_with_master_key():
|
|||
|
||||
def test_encrypt_decrypt_with_salt_key():
|
||||
os.environ["LITELLM_SALT_KEY"] = "sk-salt-key2222"
|
||||
print(f"LITELLM_SALT_KEY: {os.environ['LITELLM_SALT_KEY']}")
|
||||
assert decrypt_value_helper(encrypt_value_helper("test"), key="test_key") == "test"
|
||||
assert decrypt_value_helper(encrypt_value_helper(10), key="test_key") == 10
|
||||
assert decrypt_value_helper(encrypt_value_helper(True), key="test_key") is True
|
||||
assert decrypt_value_helper(encrypt_value_helper(None), key="test_key") is None
|
||||
assert decrypt_value_helper(encrypt_value_helper({"rpm": 10}), key="test_key") == {"rpm": 10}
|
||||
assert decrypt_value_helper(encrypt_value_helper({"rpm": 10}), key="test_key") == {
|
||||
"rpm": 10
|
||||
}
|
||||
|
||||
# encryption should actually occur for strings
|
||||
assert encrypt_value_helper("test") != "test"
|
||||
|
||||
os.environ.pop("LITELLM_SALT_KEY", None)
|
||||
|
||||
|
||||
def test_encrypt_value_helper_does_not_log_invalid_value(caplog: pytest.LogCaptureFixture) -> None:
|
||||
caplog.set_level(logging.DEBUG, logger="LiteLLM Proxy")
|
||||
secret: Final[str] = "must-not-be-logged"
|
||||
value: Final[dict[str, str]] = {"token": secret}
|
||||
|
||||
encrypt_value_helper(cast(str, value))
|
||||
|
||||
messages: Final = tuple(record.getMessage() for record in caplog.records)
|
||||
assert any("Invalid value type passed to encrypt_value" in message for message in messages)
|
||||
assert all(secret not in message for message in messages)
|
||||
|
|
|
|||
98
ui/litellm-dashboard/package-lock.json
generated
98
ui/litellm-dashboard/package-lock.json
generated
|
|
@ -25,7 +25,7 @@
|
|||
"jwt-decode": "4.0.0",
|
||||
"lucide-react": "0.513.0",
|
||||
"moment": "2.31.0",
|
||||
"next": "16.3.6",
|
||||
"next": "16.3.3",
|
||||
"next-themes": "^0.4.6",
|
||||
"nuqs": "^2.9.4",
|
||||
"openai": "4.104.0",
|
||||
|
|
@ -62,7 +62,7 @@
|
|||
"@vitest/coverage-v8": "4.1.11",
|
||||
"@vitest/ui": "4.1.11",
|
||||
"eslint": "9.39.2",
|
||||
"eslint-config-next": "16.3.6",
|
||||
"eslint-config-next": "16.3.3",
|
||||
"eslint-config-prettier": "10.1.8",
|
||||
"eslint-plugin-jest-dom": "5.10.1",
|
||||
"eslint-plugin-testing-library": "7.16.2",
|
||||
|
|
@ -2061,15 +2061,15 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@next/env": {
|
||||
"version": "16.3.6",
|
||||
"resolved": "https://registry.npmjs.org/@next/env/-/env-16.3.6.tgz",
|
||||
"integrity": "sha512-x9Vblze1EbtltQYnNH38xCPWU3TVfBd1eXqA3+w9+BTpedkkdNpAaltXlGQ/nsc1+E0mVTNrtcbX3GoO09zeLQ==",
|
||||
"version": "16.3.3",
|
||||
"resolved": "https://registry.npmjs.org/@next/env/-/env-16.3.3.tgz",
|
||||
"integrity": "sha512-U2eYQRwXj+dsqxV79zFqExDdatnNY/ZWc2nsJU1p/OgT7fd3dXwlF6OjYaFQCfMoeTA19PWq+wVmYgimVA+V+g==",
|
||||
"license": "MIT"
|
||||
},
|
||||
"node_modules/@next/eslint-plugin-next": {
|
||||
"version": "16.3.6",
|
||||
"resolved": "https://registry.npmjs.org/@next/eslint-plugin-next/-/eslint-plugin-next-16.3.6.tgz",
|
||||
"integrity": "sha512-jowwDX+7DOlDIjJLgTMxudw+k37QnWu1JkZLkSi9MaJBfDYcfhAPMKBhXL0idYzFN/AGg//axnOR4cLkHX/Rng==",
|
||||
"version": "16.3.3",
|
||||
"resolved": "https://registry.npmjs.org/@next/eslint-plugin-next/-/eslint-plugin-next-16.3.3.tgz",
|
||||
"integrity": "sha512-pbEh30vvjKpDoTAmo1v3q2uM4JUi8QaEBpbmjWvGfoec2jLghy/WNtvzAT0bk+Ik9oz6etjt4YjXEk4BQnicCw==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
|
|
@ -2078,9 +2078,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@next/swc-darwin-arm64": {
|
||||
"version": "16.3.6",
|
||||
"resolved": "https://registry.npmjs.org/@next/swc-darwin-arm64/-/swc-darwin-arm64-16.3.6.tgz",
|
||||
"integrity": "sha512-E/7GEqaUkt8mk/T8v9lAnrhzR06kdq1ZBkC12F8tAMkdIadwNp3H1KqHynDHrpcTlGCUdq/qu6vUL2aYVyYBdw==",
|
||||
"version": "16.3.3",
|
||||
"resolved": "https://registry.npmjs.org/@next/swc-darwin-arm64/-/swc-darwin-arm64-16.3.3.tgz",
|
||||
"integrity": "sha512-8Hiv32QJPwdV6KYJ8meR9SBA061tQqnIKTJDocvOXlEQqib0xMFpzArosuffFUUc0sslbh7QQ8a3Yey1QV8EIw==",
|
||||
"cpu": [
|
||||
"arm64"
|
||||
],
|
||||
|
|
@ -2094,9 +2094,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@next/swc-darwin-x64": {
|
||||
"version": "16.3.6",
|
||||
"resolved": "https://registry.npmjs.org/@next/swc-darwin-x64/-/swc-darwin-x64-16.3.6.tgz",
|
||||
"integrity": "sha512-yBE893/nDWTlaiBD1p+qgt7NUen4U5R6FXyH0s67Npq1S3E0cVSef1WIXC2xBRgQvwAvJq6DnS6Y6PrY0cy4Ew==",
|
||||
"version": "16.3.3",
|
||||
"resolved": "https://registry.npmjs.org/@next/swc-darwin-x64/-/swc-darwin-x64-16.3.3.tgz",
|
||||
"integrity": "sha512-A1lgKgwVchRYmSe467zdwhxT9040dd8lH+o65sL5Jet8fjB4kegw/rDyPIpYVRb6jAqwXFOJpjIXJLxQKLiE3A==",
|
||||
"cpu": [
|
||||
"x64"
|
||||
],
|
||||
|
|
@ -2110,9 +2110,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@next/swc-linux-arm64-gnu": {
|
||||
"version": "16.3.6",
|
||||
"resolved": "https://registry.npmjs.org/@next/swc-linux-arm64-gnu/-/swc-linux-arm64-gnu-16.3.6.tgz",
|
||||
"integrity": "sha512-KJDpjBqBPYlvkivmyrp+Qys6k/7ksbqGQvRVc6ZEGfR+cjQxx+nUkJaWmNZJsmoOrqYNbaXByF8wa0lBwDhB3Q==",
|
||||
"version": "16.3.3",
|
||||
"resolved": "https://registry.npmjs.org/@next/swc-linux-arm64-gnu/-/swc-linux-arm64-gnu-16.3.3.tgz",
|
||||
"integrity": "sha512-bf0FIssMFueU2dm7vQEWWxk0c8UjKTdW0yzuh0sQsD8pf1+KCLDdaqhYZNMYGmXwEOiHAUzgBKudovIlcvvBjg==",
|
||||
"cpu": [
|
||||
"arm64"
|
||||
],
|
||||
|
|
@ -2129,9 +2129,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@next/swc-linux-arm64-musl": {
|
||||
"version": "16.3.6",
|
||||
"resolved": "https://registry.npmjs.org/@next/swc-linux-arm64-musl/-/swc-linux-arm64-musl-16.3.6.tgz",
|
||||
"integrity": "sha512-mqNg2K+hvWskSRb/QM+Ix412DvBsuSF0XV+frTSw5vmoucNnIlynFwKYew8D01bfATErMOM7Bujrf0BA5DRKFA==",
|
||||
"version": "16.3.3",
|
||||
"resolved": "https://registry.npmjs.org/@next/swc-linux-arm64-musl/-/swc-linux-arm64-musl-16.3.3.tgz",
|
||||
"integrity": "sha512-W7viwCk9JY/cAkdz/A273rd5bb3RgT/IHwR7Upv90tunjBWNtAAhGhoecHh+teRNRSinuAFmE+l7fwZ4YKkrXg==",
|
||||
"cpu": [
|
||||
"arm64"
|
||||
],
|
||||
|
|
@ -2148,9 +2148,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@next/swc-linux-x64-gnu": {
|
||||
"version": "16.3.6",
|
||||
"resolved": "https://registry.npmjs.org/@next/swc-linux-x64-gnu/-/swc-linux-x64-gnu-16.3.6.tgz",
|
||||
"integrity": "sha512-nFncBNGAYouRHjRVaITs9beZRfhX4ssVwpnvPIAbkZVH6LtGoAVlH4bJ8Cnf9SOo9bsXgPFer/GdHtEE3JNOkw==",
|
||||
"version": "16.3.3",
|
||||
"resolved": "https://registry.npmjs.org/@next/swc-linux-x64-gnu/-/swc-linux-x64-gnu-16.3.3.tgz",
|
||||
"integrity": "sha512-0W46zw1N3ODpI6n0GeivHvvob1pooozgZVqy65k0mh4/7vr+FbY9+WpHzNVXjHipJf/A3FDheBG19H1s5A25rA==",
|
||||
"cpu": [
|
||||
"x64"
|
||||
],
|
||||
|
|
@ -2167,9 +2167,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@next/swc-linux-x64-musl": {
|
||||
"version": "16.3.6",
|
||||
"resolved": "https://registry.npmjs.org/@next/swc-linux-x64-musl/-/swc-linux-x64-musl-16.3.6.tgz",
|
||||
"integrity": "sha512-5Mf3cHDGR/Iz0ng2Bj3zUR3p5QS9YK3Hn2QiAfavFmyF48zwThAjpFoiTKNIcOHLYS4zEk+gzyJ/9deQ2ZB8yQ==",
|
||||
"version": "16.3.3",
|
||||
"resolved": "https://registry.npmjs.org/@next/swc-linux-x64-musl/-/swc-linux-x64-musl-16.3.3.tgz",
|
||||
"integrity": "sha512-H4mBso8ZTMBPtdT0PN0pBx2ayTvQuTuvS6qT13d77yVFJXAPCxkyIhLTmdMaGTJs0krQYI/qpzdHijCeihXhbg==",
|
||||
"cpu": [
|
||||
"x64"
|
||||
],
|
||||
|
|
@ -2186,9 +2186,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@next/swc-win32-arm64-msvc": {
|
||||
"version": "16.3.6",
|
||||
"resolved": "https://registry.npmjs.org/@next/swc-win32-arm64-msvc/-/swc-win32-arm64-msvc-16.3.6.tgz",
|
||||
"integrity": "sha512-0jkJy0C2kbrJWTk4YLa3xk80pVBpx8FCHJym7CnUfDAXe/FWv5qT7SQJbR0KuemyxaEDlEx5WT4VQJoTW+/9Qw==",
|
||||
"version": "16.3.3",
|
||||
"resolved": "https://registry.npmjs.org/@next/swc-win32-arm64-msvc/-/swc-win32-arm64-msvc-16.3.3.tgz",
|
||||
"integrity": "sha512-cTMUJpcEGmeywofCUfhR+rSsoE33+rVPnPEYNTNdLNlsOeEg/vktOsKUSTb28vUGqD2jkm4Zaskcwn7OCI6FQg==",
|
||||
"cpu": [
|
||||
"arm64"
|
||||
],
|
||||
|
|
@ -2202,9 +2202,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@next/swc-win32-x64-msvc": {
|
||||
"version": "16.3.6",
|
||||
"resolved": "https://registry.npmjs.org/@next/swc-win32-x64-msvc/-/swc-win32-x64-msvc-16.3.6.tgz",
|
||||
"integrity": "sha512-/YXjI1e5OXcZ7YpxRwgP/1jAV/SBKTzeVKqN2mk7mLpcICsyn3Gl5+dIfDTJp70M0ccMhyMMRso4v6mPDCGepg==",
|
||||
"version": "16.3.3",
|
||||
"resolved": "https://registry.npmjs.org/@next/swc-win32-x64-msvc/-/swc-win32-x64-msvc-16.3.3.tgz",
|
||||
"integrity": "sha512-2VR4cTBzHXaBjnGsuH6GyJjENzQOmHeAh11uY1iUhjm3j5dEUrVJuUj+VL78jaGi/Dik8xS76zEj18BsFhlVZQ==",
|
||||
"cpu": [
|
||||
"x64"
|
||||
],
|
||||
|
|
@ -6077,13 +6077,13 @@
|
|||
}
|
||||
},
|
||||
"node_modules/eslint-config-next": {
|
||||
"version": "16.3.6",
|
||||
"resolved": "https://registry.npmjs.org/eslint-config-next/-/eslint-config-next-16.3.6.tgz",
|
||||
"integrity": "sha512-1Upt3U7BDwU+ilpe2byZjAfts9oNq4d4fv/zXEvs8/4yS+cwOQW/WCxUNy8gCDquX67SzeehDvKblVC6ZBMocQ==",
|
||||
"version": "16.3.3",
|
||||
"resolved": "https://registry.npmjs.org/eslint-config-next/-/eslint-config-next-16.3.3.tgz",
|
||||
"integrity": "sha512-teqtsR26tnlfXFHfVLTM/4tzEzU8DMu6GS1sddZzhfGzgd2f2ofbgDUcsk6cssSCzX6Tk6fmWifJcdANSdPJrw==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@next/eslint-plugin-next": "16.3.6",
|
||||
"@next/eslint-plugin-next": "16.3.3",
|
||||
"eslint-import-resolver-node": "^0.3.6",
|
||||
"eslint-import-resolver-typescript": "^3.5.2",
|
||||
"eslint-plugin-import": "^2.32.0",
|
||||
|
|
@ -9712,12 +9712,12 @@
|
|||
"license": "MIT"
|
||||
},
|
||||
"node_modules/next": {
|
||||
"version": "16.3.6",
|
||||
"resolved": "https://registry.npmjs.org/next/-/next-16.3.6.tgz",
|
||||
"integrity": "sha512-L+otWM/aQbYTx98aZhgEoMb4bZAXx1YVW4UMA/vuCyCoWG5HJyZUili8QAkqzrcC+5///tsz3s0M+SlyB5bLMw==",
|
||||
"version": "16.3.3",
|
||||
"resolved": "https://registry.npmjs.org/next/-/next-16.3.3.tgz",
|
||||
"integrity": "sha512-tuRTx1nQ/yVw83cwJBo9F+njGUgMn3UHQycreWHB8XsStvvAh1AthbI8/4IpKnFaF58F+iSiHejYOlMQ/eq83g==",
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@next/env": "16.3.6",
|
||||
"@next/env": "16.3.3",
|
||||
"@swc/helpers": "0.5.23",
|
||||
"baseline-browser-mapping": "^2.9.19",
|
||||
"caniuse-lite": "^1.0.30001579",
|
||||
|
|
@ -9731,15 +9731,15 @@
|
|||
"node": ">=20.9.0"
|
||||
},
|
||||
"optionalDependencies": {
|
||||
"@next/swc-darwin-arm64": "16.3.6",
|
||||
"@next/swc-darwin-x64": "16.3.6",
|
||||
"@next/swc-linux-arm64-gnu": "16.3.6",
|
||||
"@next/swc-linux-arm64-musl": "16.3.6",
|
||||
"@next/swc-linux-x64-gnu": "16.3.6",
|
||||
"@next/swc-linux-x64-musl": "16.3.6",
|
||||
"@next/swc-win32-arm64-msvc": "16.3.6",
|
||||
"@next/swc-win32-x64-msvc": "16.3.6",
|
||||
"sharp": "^0.35.4"
|
||||
"@next/swc-darwin-arm64": "16.3.3",
|
||||
"@next/swc-darwin-x64": "16.3.3",
|
||||
"@next/swc-linux-arm64-gnu": "16.3.3",
|
||||
"@next/swc-linux-arm64-musl": "16.3.3",
|
||||
"@next/swc-linux-x64-gnu": "16.3.3",
|
||||
"@next/swc-linux-x64-musl": "16.3.3",
|
||||
"@next/swc-win32-arm64-msvc": "16.3.3",
|
||||
"@next/swc-win32-x64-msvc": "16.3.3",
|
||||
"sharp": "^0.35.3"
|
||||
},
|
||||
"peerDependencies": {
|
||||
"@opentelemetry/api": "^1.1.0",
|
||||
|
|
|
|||
|
|
@ -41,7 +41,7 @@
|
|||
"jwt-decode": "4.0.0",
|
||||
"lucide-react": "0.513.0",
|
||||
"moment": "2.31.0",
|
||||
"next": "16.3.6",
|
||||
"next": "16.3.3",
|
||||
"next-themes": "^0.4.6",
|
||||
"nuqs": "^2.9.4",
|
||||
"openai": "4.104.0",
|
||||
|
|
@ -78,7 +78,7 @@
|
|||
"@vitest/coverage-v8": "4.1.11",
|
||||
"@vitest/ui": "4.1.11",
|
||||
"eslint": "9.39.2",
|
||||
"eslint-config-next": "16.3.6",
|
||||
"eslint-config-next": "16.3.3",
|
||||
"eslint-config-prettier": "10.1.8",
|
||||
"eslint-plugin-jest-dom": "5.10.1",
|
||||
"eslint-plugin-testing-library": "7.16.2",
|
||||
|
|
|
|||
|
|
@ -9,8 +9,8 @@ import { useMCPServers } from "@/app/(dashboard)/hooks/mcpServers/useMCPServers"
|
|||
import { ModelSelect } from "@/components/ModelSelect/ModelSelect";
|
||||
import { FieldGroup } from "@/components/ui/field";
|
||||
import { FormField } from "@/components/shared/form/FormField";
|
||||
import { MultiSelect } from "@/components/shared/MultiSelect";
|
||||
import { Input } from "@/components/ui/input";
|
||||
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
|
||||
import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs";
|
||||
import { Textarea } from "@/components/ui/textarea";
|
||||
|
||||
|
|
@ -29,53 +29,6 @@ export const MODELS_TAB = "models";
|
|||
export const MCP_SERVERS_TAB = "mcp-servers";
|
||||
export const AGENTS_TAB = "agents";
|
||||
|
||||
interface MultiSelectOption {
|
||||
value: string;
|
||||
label: string;
|
||||
}
|
||||
|
||||
interface MultiSelectProps {
|
||||
id: string;
|
||||
value: string[];
|
||||
onChange: (value: string[]) => void;
|
||||
options: MultiSelectOption[];
|
||||
placeholder: string;
|
||||
"aria-invalid": true | undefined;
|
||||
"aria-describedby": string | undefined;
|
||||
}
|
||||
|
||||
const MultiSelect = ({
|
||||
id,
|
||||
value,
|
||||
onChange,
|
||||
options,
|
||||
placeholder,
|
||||
"aria-invalid": ariaInvalid,
|
||||
"aria-describedby": ariaDescribedBy,
|
||||
}: MultiSelectProps) => (
|
||||
<Select multiple items={options} value={value} onValueChange={onChange}>
|
||||
<SelectTrigger id={id} aria-invalid={ariaInvalid} aria-describedby={ariaDescribedBy} className="w-full">
|
||||
<SelectValue placeholder={placeholder}>
|
||||
{(selected: string[]) =>
|
||||
selected.length === 0
|
||||
? placeholder
|
||||
: options
|
||||
.filter((option) => selected.includes(option.value))
|
||||
.map((option) => option.label)
|
||||
.join(", ")
|
||||
}
|
||||
</SelectValue>
|
||||
</SelectTrigger>
|
||||
<SelectContent>
|
||||
{options.map((option) => (
|
||||
<SelectItem key={option.value} value={option.value} title={option.label}>
|
||||
{option.label}
|
||||
</SelectItem>
|
||||
))}
|
||||
</SelectContent>
|
||||
</Select>
|
||||
);
|
||||
|
||||
interface AccessGroupBaseFormProps {
|
||||
form: UseFormReturn<AccessGroupFormValues>;
|
||||
isNameDisabled?: boolean;
|
||||
|
|
@ -145,15 +98,13 @@ export function AccessGroupBaseForm({
|
|||
|
||||
<TabsContent value={MCP_SERVERS_TAB} className="pt-4">
|
||||
<FormField control={form.control} name="mcpServerIds" label="Allowed MCP Servers">
|
||||
{({ id, value, onChange, "aria-invalid": ariaInvalid, "aria-describedby": ariaDescribedBy }) => (
|
||||
{({ id, value, onChange }) => (
|
||||
<MultiSelect
|
||||
id={id}
|
||||
value={value}
|
||||
onChange={onChange}
|
||||
onValueChange={onChange}
|
||||
options={mcpServerOptions}
|
||||
placeholder="Select MCP servers"
|
||||
aria-invalid={ariaInvalid}
|
||||
aria-describedby={ariaDescribedBy}
|
||||
/>
|
||||
)}
|
||||
</FormField>
|
||||
|
|
@ -161,15 +112,13 @@ export function AccessGroupBaseForm({
|
|||
|
||||
<TabsContent value={AGENTS_TAB} className="pt-4">
|
||||
<FormField control={form.control} name="agentIds" label="Allowed Agents">
|
||||
{({ id, value, onChange, "aria-invalid": ariaInvalid, "aria-describedby": ariaDescribedBy }) => (
|
||||
{({ id, value, onChange }) => (
|
||||
<MultiSelect
|
||||
id={id}
|
||||
value={value}
|
||||
onChange={onChange}
|
||||
onValueChange={onChange}
|
||||
options={agentOptions}
|
||||
placeholder="Select agents"
|
||||
aria-invalid={ariaInvalid}
|
||||
aria-describedby={ariaDescribedBy}
|
||||
/>
|
||||
)}
|
||||
</FormField>
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import { describe, it, expect, vi, beforeEach } from "vitest";
|
||||
import userEvent, { PointerEventsCheckLevel } from "@testing-library/user-event";
|
||||
import { fireEvent, renderWithProviders, screen, waitFor } from "../../../../../../tests/test-utils";
|
||||
import { fireEvent, renderWithProviders, screen, waitFor, within } from "../../../../../../tests/test-utils";
|
||||
import { AccessGroupEditModal } from "./AccessGroupEditModal";
|
||||
import { AccessGroupResponse } from "@/app/(dashboard)/hooks/accessGroups/useAccessGroups";
|
||||
|
||||
|
|
@ -14,8 +14,13 @@ vi.mock("@/app/(dashboard)/hooks/agents/useAgents", () => ({
|
|||
useAgents: () => ({ data: { agents: [{ agent_id: "agent-1", agent_name: "Support Bot" }] } }),
|
||||
}));
|
||||
|
||||
const manyServers = Array.from({ length: 20 }, (_, i) => ({
|
||||
server_id: `srv-${i + 1}`,
|
||||
server_name: `Server ${i + 1}`,
|
||||
}));
|
||||
|
||||
vi.mock("@/app/(dashboard)/hooks/mcpServers/useMCPServers", () => ({
|
||||
useMCPServers: () => ({ data: [{ server_id: "srv-1", server_name: "Files" }] }),
|
||||
useMCPServers: () => ({ data: [{ server_id: "srv-1", server_name: "Files" }, ...manyServers.slice(1)] }),
|
||||
}));
|
||||
|
||||
vi.mock("@/components/ModelSelect/ModelSelect", () => ({
|
||||
|
|
@ -164,6 +169,47 @@ describe("AccessGroupEditModal submit payload", () => {
|
|||
expect(mutate).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("renders each selected MCP server as its own removable chip and drops one on remove", async () => {
|
||||
const user = setup();
|
||||
renderModal();
|
||||
await screen.findByDisplayValue("Engineering");
|
||||
|
||||
await user.click(screen.getByRole("tab", { name: /MCP Servers/ }));
|
||||
const chip = await screen.findByLabelText("Files");
|
||||
expect(chip).toHaveAttribute("data-slot", "combobox-chip");
|
||||
expect(screen.queryByText("srv-1")).not.toBeInTheDocument();
|
||||
|
||||
await user.click(within(chip).getByRole("button"));
|
||||
await save(user);
|
||||
|
||||
await waitFor(() => expect(mutate).toHaveBeenCalled());
|
||||
expect(variables().params.access_mcp_server_ids).toStrictEqual([]);
|
||||
});
|
||||
|
||||
it("keeps 20 selected MCP servers as separate chips instead of one joined string", async () => {
|
||||
const user = setup();
|
||||
renderModal({ ...accessGroup, access_mcp_server_ids: manyServers.map((s) => s.server_id) });
|
||||
await screen.findByDisplayValue("Engineering");
|
||||
|
||||
await user.click(screen.getByRole("tab", { name: /MCP Servers/ }));
|
||||
await screen.findByLabelText("Server 20");
|
||||
const chips = screen.getAllByLabelText(/^(Files|Server \d+)$/);
|
||||
expect(chips).toHaveLength(20);
|
||||
expect(chips.map((chip) => chip.textContent)).toStrictEqual([
|
||||
"Files",
|
||||
...manyServers.slice(1).map((s) => s.server_name),
|
||||
]);
|
||||
expect(screen.queryByText(/Server 2, Server 3/)).not.toBeInTheDocument();
|
||||
|
||||
await user.click(within(screen.getByLabelText("Server 7")).getByRole("button"));
|
||||
await save(user);
|
||||
|
||||
await waitFor(() => expect(mutate).toHaveBeenCalled());
|
||||
expect(variables().params.access_mcp_server_ids).toStrictEqual(
|
||||
manyServers.map((s) => s.server_id).filter((id) => id !== "srv-7"),
|
||||
);
|
||||
});
|
||||
|
||||
it("sends models chosen on the Models tab", async () => {
|
||||
const user = setup();
|
||||
renderModal({ ...accessGroup, access_model_names: [] });
|
||||
|
|
|
|||
|
|
@ -98,6 +98,31 @@ describe("AccessGroupCreateDialog", () => {
|
|||
});
|
||||
});
|
||||
|
||||
it("sends MCP servers and agents picked from the chip selectors as ids", async () => {
|
||||
const user = userEvent.setup();
|
||||
const { createAccessGroup } = renderDialog();
|
||||
|
||||
await user.type(screen.getByLabelText("Group Name"), "mcp-group");
|
||||
await user.click(screen.getByRole("tab", { name: "MCP Servers" }));
|
||||
await user.click(screen.getByLabelText("Allowed MCP Servers"));
|
||||
await user.click(await screen.findByRole("option", { name: "GitHub MCP" }));
|
||||
expect(screen.getByLabelText("GitHub MCP")).toHaveAttribute("data-slot", "combobox-chip");
|
||||
await user.keyboard("{Escape}");
|
||||
|
||||
await user.click(screen.getByRole("tab", { name: "Agents" }));
|
||||
await user.click(screen.getByLabelText("Allowed Agents"));
|
||||
await user.click(await screen.findByRole("option", { name: "Support Agent" }));
|
||||
await user.keyboard("{Escape}");
|
||||
await user.click(screen.getByRole("button", { name: "Create Group" }));
|
||||
|
||||
await waitFor(() => expect(createAccessGroup).toHaveBeenCalledTimes(1));
|
||||
expect(createAccessGroup.mock.calls[0][0]).toStrictEqual({
|
||||
access_group_name: "mcp-group",
|
||||
access_mcp_server_ids: ["srv-1"],
|
||||
access_agent_ids: ["agent-1"],
|
||||
});
|
||||
});
|
||||
|
||||
it("keeps the dialog open with the entered values when the create fails", async () => {
|
||||
const user = userEvent.setup();
|
||||
const { createAccessGroup } = renderDialog({
|
||||
|
|
|
|||
|
|
@ -11,10 +11,10 @@ import { ModelSelect } from "@/components/ModelSelect/ModelSelect";
|
|||
import { toast } from "@/lib/toast";
|
||||
import { FieldGroup } from "@/components/ui/field";
|
||||
import { FormField } from "@/components/shared/form/FormField";
|
||||
import { MultiSelect } from "@/components/shared/MultiSelect";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import { Dialog, DialogContent, DialogFooter, DialogHeader, DialogTitle } from "@/components/ui/dialog";
|
||||
import { Input } from "@/components/ui/input";
|
||||
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
|
||||
import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs";
|
||||
import { Textarea } from "@/components/ui/textarea";
|
||||
import { useZodForm } from "@/lib/forms/useZodForm";
|
||||
|
|
@ -25,53 +25,6 @@ import { accessGroupCreateSchema } from "./schema";
|
|||
|
||||
const GENERAL_TAB = "general";
|
||||
|
||||
interface MultiSelectOption {
|
||||
value: string;
|
||||
label: string;
|
||||
}
|
||||
|
||||
interface MultiSelectProps {
|
||||
id: string;
|
||||
value: string[];
|
||||
onChange: (value: string[]) => void;
|
||||
options: MultiSelectOption[];
|
||||
placeholder: string;
|
||||
"aria-invalid": true | undefined;
|
||||
"aria-describedby": string | undefined;
|
||||
}
|
||||
|
||||
const MultiSelect = ({
|
||||
id,
|
||||
value,
|
||||
onChange,
|
||||
options,
|
||||
placeholder,
|
||||
"aria-invalid": ariaInvalid,
|
||||
"aria-describedby": ariaDescribedBy,
|
||||
}: MultiSelectProps) => (
|
||||
<Select multiple items={options} value={value} onValueChange={onChange}>
|
||||
<SelectTrigger id={id} aria-invalid={ariaInvalid} aria-describedby={ariaDescribedBy} className="w-full">
|
||||
<SelectValue placeholder={placeholder}>
|
||||
{(selected: string[]) =>
|
||||
selected.length === 0
|
||||
? placeholder
|
||||
: options
|
||||
.filter((option) => selected.includes(option.value))
|
||||
.map((option) => option.label)
|
||||
.join(", ")
|
||||
}
|
||||
</SelectValue>
|
||||
</SelectTrigger>
|
||||
<SelectContent>
|
||||
{options.map((option) => (
|
||||
<SelectItem key={option.value} value={option.value}>
|
||||
{option.label}
|
||||
</SelectItem>
|
||||
))}
|
||||
</SelectContent>
|
||||
</Select>
|
||||
);
|
||||
|
||||
const defaultCreateAccessGroup = async (body: AccessGroupCreateBody): Promise<unknown> => {
|
||||
const { data } = await fetchClient.POST("/v1/access_group", { body });
|
||||
return data;
|
||||
|
|
@ -193,15 +146,13 @@ export const AccessGroupCreateDialog = ({
|
|||
|
||||
<TabsContent value="mcp-servers" className="pt-4">
|
||||
<FormField control={form.control} name="mcpServerIds" label="Allowed MCP Servers">
|
||||
{({ id, value, onChange, "aria-invalid": ariaInvalid, "aria-describedby": ariaDescribedBy }) => (
|
||||
{({ id, value, onChange }) => (
|
||||
<MultiSelect
|
||||
id={id}
|
||||
value={value}
|
||||
onChange={onChange}
|
||||
onValueChange={onChange}
|
||||
options={mcpServerOptions}
|
||||
placeholder="Select MCP servers"
|
||||
aria-invalid={ariaInvalid}
|
||||
aria-describedby={ariaDescribedBy}
|
||||
/>
|
||||
)}
|
||||
</FormField>
|
||||
|
|
@ -209,15 +160,13 @@ export const AccessGroupCreateDialog = ({
|
|||
|
||||
<TabsContent value="agents" className="pt-4">
|
||||
<FormField control={form.control} name="agentIds" label="Allowed Agents">
|
||||
{({ id, value, onChange, "aria-invalid": ariaInvalid, "aria-describedby": ariaDescribedBy }) => (
|
||||
{({ id, value, onChange }) => (
|
||||
<MultiSelect
|
||||
id={id}
|
||||
value={value}
|
||||
onChange={onChange}
|
||||
onValueChange={onChange}
|
||||
options={agentOptions}
|
||||
placeholder="Select agents"
|
||||
aria-invalid={ariaInvalid}
|
||||
aria-describedby={ariaDescribedBy}
|
||||
/>
|
||||
)}
|
||||
</FormField>
|
||||
|
|
|
|||
|
|
@ -33,6 +33,12 @@ describe("AgentsTable", () => {
|
|||
}
|
||||
});
|
||||
|
||||
it("right-aligns the Spend (USD) column", () => {
|
||||
render(<AgentsTable agents={[]} {...baseProps} />);
|
||||
expect(screen.getByRole("columnheader", { name: "Spend (USD)" })).toHaveClass("text-right");
|
||||
expect(screen.getByRole("columnheader", { name: "Agent Name" })).not.toHaveClass("text-right");
|
||||
});
|
||||
|
||||
it("renders the agent's model and opens the detail view when the ID cell is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onAgentClick = vi.fn();
|
||||
|
|
|
|||
|
|
@ -90,7 +90,7 @@ export const getAgentsTableColumns = ({
|
|||
{
|
||||
id: "spend",
|
||||
accessorKey: "spend",
|
||||
meta: { title: "Spend (USD)" },
|
||||
meta: { title: "Spend (USD)", numeric: true },
|
||||
header: ({ column }) => <DataTableSortHeader column={column} title="Spend (USD)" />,
|
||||
size: 130,
|
||||
enableSorting: true,
|
||||
|
|
|
|||
|
|
@ -50,6 +50,8 @@ describe("ProviderDiscountTable", () => {
|
|||
expect(screen.getByRole("columnheader", { name: "Provider" })).toBeInTheDocument();
|
||||
expect(screen.getByRole("columnheader", { name: "Discount Percentage" })).toBeInTheDocument();
|
||||
expect(screen.getByRole("columnheader", { name: "Actions" })).toBeInTheDocument();
|
||||
expect(screen.getByRole("columnheader", { name: "Discount Percentage" })).toHaveClass("text-right");
|
||||
expect(screen.getByRole("columnheader", { name: "Provider" })).not.toHaveClass("text-right");
|
||||
});
|
||||
|
||||
it("should display provider display names in the table", () => {
|
||||
|
|
|
|||
|
|
@ -80,10 +80,11 @@ const ProviderDiscountTable: React.FC<ProviderDiscountTableProps> = ({
|
|||
},
|
||||
{
|
||||
header: "Discount Percentage",
|
||||
numeric: true,
|
||||
cell: (row) => {
|
||||
const { displayName } = getProviderLogoAndName(row.provider);
|
||||
return (
|
||||
<div className="flex items-center gap-2">
|
||||
<div className="flex items-center justify-end gap-2">
|
||||
{editingProvider === row.provider ? (
|
||||
<>
|
||||
<Input
|
||||
|
|
|
|||
|
|
@ -46,6 +46,8 @@ describe("ProviderMarginTable", () => {
|
|||
expect(screen.getByRole("columnheader", { name: "Provider" })).toBeInTheDocument();
|
||||
expect(screen.getByRole("columnheader", { name: "Margin" })).toBeInTheDocument();
|
||||
expect(screen.getByRole("columnheader", { name: "Actions" })).toBeInTheDocument();
|
||||
expect(screen.getByRole("columnheader", { name: "Margin" })).toHaveClass("text-right");
|
||||
expect(screen.getByRole("columnheader", { name: "Provider" })).not.toHaveClass("text-right");
|
||||
});
|
||||
|
||||
it("should display the provider display name", () => {
|
||||
|
|
|
|||
|
|
@ -123,10 +123,11 @@ const ProviderMarginTable: React.FC<ProviderMarginTableProps> = ({
|
|||
},
|
||||
{
|
||||
header: "Margin",
|
||||
numeric: true,
|
||||
cell: (row) => {
|
||||
const displayName = marginRowDisplayName(row.provider);
|
||||
return (
|
||||
<div className="flex items-center gap-2">
|
||||
<div className="flex items-center justify-end gap-2">
|
||||
{editingProvider === row.provider ? (
|
||||
<>
|
||||
<div className="flex items-center gap-2">
|
||||
|
|
|
|||
|
|
@ -174,6 +174,8 @@ describe("AllModelsTable", () => {
|
|||
const { rerender } = render(<AllModelsTable {...baseProps} />);
|
||||
expect(screen.getByText("$30")).toBeInTheDocument();
|
||||
expect(screen.getByText("$60")).toBeInTheDocument();
|
||||
expect(screen.getByRole("cell", { name: /\$30/ })).toHaveClass("text-right");
|
||||
expect(screen.getByRole("columnheader", { name: /costs/i })).toHaveClass("text-right");
|
||||
|
||||
rerender(<AllModelsTable {...baseProps} data={[makeModel({ input_cost: null, output_cost: null })]} />);
|
||||
expect(screen.queryByText(/^\$/)).not.toBeInTheDocument();
|
||||
|
|
|
|||
|
|
@ -437,7 +437,7 @@ export const getModelsTableColumns = ({
|
|||
{
|
||||
id: COSTS_COLUMN_ID,
|
||||
accessorFn: (row) => row.input_cost,
|
||||
meta: { title: "Costs" },
|
||||
meta: { title: "Costs", numeric: true },
|
||||
header: ({ column }) => <DataTableSortHeader column={column} title="Costs" />,
|
||||
enableSorting: true,
|
||||
size: 130,
|
||||
|
|
|
|||
|
|
@ -19,7 +19,15 @@ import {
|
|||
} from "@/components/ui/combobox";
|
||||
import { Meter, MeterIndicator, MeterTrack } from "@/components/shared/Meter";
|
||||
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
|
||||
import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table";
|
||||
import {
|
||||
NUMERIC_CELL_CLASS,
|
||||
Table,
|
||||
TableBody,
|
||||
TableCell,
|
||||
TableHead,
|
||||
TableHeader,
|
||||
TableRow,
|
||||
} from "@/components/ui/table";
|
||||
import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs";
|
||||
import { AreaChart, BarChart, DonutChart } from "@/components/shared/charts";
|
||||
|
||||
|
|
@ -651,14 +659,14 @@ const UsagePage: React.FC<UsagePageProps> = ({ accessToken, token, userRole, use
|
|||
<TableHeader>
|
||||
<TableRow>
|
||||
<TableHead>Provider</TableHead>
|
||||
<TableHead>Spend</TableHead>
|
||||
<TableHead className={NUMERIC_CELL_CLASS}>Spend</TableHead>
|
||||
</TableRow>
|
||||
</TableHeader>
|
||||
<TableBody>
|
||||
{spendByProvider.map((provider) => (
|
||||
<TableRow key={provider.provider}>
|
||||
<TableCell>{provider.provider}</TableCell>
|
||||
<TableCell>
|
||||
<TableCell className={NUMERIC_CELL_CLASS}>
|
||||
<MoneyCell value={provider.spend} decimals={2} />
|
||||
</TableCell>
|
||||
</TableRow>
|
||||
|
|
@ -840,8 +848,8 @@ const UsagePage: React.FC<UsagePageProps> = ({ accessToken, token, userRole, use
|
|||
<TableHeader>
|
||||
<TableRow>
|
||||
<TableHead>Customer</TableHead>
|
||||
<TableHead>Spend</TableHead>
|
||||
<TableHead>Total Events</TableHead>
|
||||
<TableHead className={NUMERIC_CELL_CLASS}>Spend</TableHead>
|
||||
<TableHead className={NUMERIC_CELL_CLASS}>Total Events</TableHead>
|
||||
</TableRow>
|
||||
</TableHeader>
|
||||
|
||||
|
|
@ -849,10 +857,10 @@ const UsagePage: React.FC<UsagePageProps> = ({ accessToken, token, userRole, use
|
|||
{topUsers?.map((user: any, index: number) => (
|
||||
<TableRow key={index}>
|
||||
<TableCell>{user.end_user}</TableCell>
|
||||
<TableCell>
|
||||
<TableCell className={NUMERIC_CELL_CLASS}>
|
||||
<MoneyCell value={user.total_spend} decimals={2} />
|
||||
</TableCell>
|
||||
<TableCell>{user.total_count}</TableCell>
|
||||
<TableCell className={NUMERIC_CELL_CLASS}>{user.total_count}</TableCell>
|
||||
</TableRow>
|
||||
))}
|
||||
</TableBody>
|
||||
|
|
|
|||
|
|
@ -82,6 +82,16 @@ describe("OrganizationsTable", () => {
|
|||
}
|
||||
});
|
||||
|
||||
it("right-aligns the money and count columns only", () => {
|
||||
renderWithProviders(<OrganizationsTable {...baseProps} organizations={[]} />);
|
||||
for (const header of ["Spend (USD)", "Budget (USD)", "Members"]) {
|
||||
expect(screen.getByRole("columnheader", { name: header })).toHaveClass("text-right");
|
||||
}
|
||||
for (const header of ["Organization Name", "TPM / RPM Limits"]) {
|
||||
expect(screen.getByRole("columnheader", { name: header })).not.toHaveClass("text-right");
|
||||
}
|
||||
});
|
||||
|
||||
it("opens the detail view when the organization ID cell is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onOrganizationClick = vi.fn();
|
||||
|
|
|
|||
|
|
@ -129,7 +129,7 @@ export const getOrganizationsTableColumns = ({
|
|||
{
|
||||
id: "spend",
|
||||
accessorKey: "spend",
|
||||
meta: { title: "Spend (USD)" },
|
||||
meta: { title: "Spend (USD)", numeric: true },
|
||||
header: ({ column }) => <DataTableSortHeader column={column} title="Spend (USD)" />,
|
||||
size: 120,
|
||||
enableSorting: true,
|
||||
|
|
@ -137,7 +137,7 @@ export const getOrganizationsTableColumns = ({
|
|||
},
|
||||
{
|
||||
id: "max_budget",
|
||||
meta: { title: "Budget (USD)" },
|
||||
meta: { title: "Budget (USD)", numeric: true },
|
||||
header: "Budget (USD)",
|
||||
size: 120,
|
||||
enableSorting: false,
|
||||
|
|
@ -163,7 +163,7 @@ export const getOrganizationsTableColumns = ({
|
|||
},
|
||||
{
|
||||
id: "members",
|
||||
meta: { title: "Members" },
|
||||
meta: { title: "Members", numeric: true },
|
||||
header: "Members",
|
||||
size: 100,
|
||||
enableSorting: false,
|
||||
|
|
|
|||
|
|
@ -9,7 +9,16 @@ import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card";
|
|||
import { Checkbox } from "@/components/ui/checkbox";
|
||||
import { Dialog, DialogContent, DialogHeader, DialogTitle } from "@/components/ui/dialog";
|
||||
import { Separator } from "@/components/ui/separator";
|
||||
import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table";
|
||||
import {
|
||||
NUMERIC_CELL_CLASS,
|
||||
Table,
|
||||
TableBody,
|
||||
TableCell,
|
||||
TableHead,
|
||||
TableHeader,
|
||||
TableRow,
|
||||
} from "@/components/ui/table";
|
||||
import { cn } from "@/lib/cva.config";
|
||||
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
||||
|
||||
interface BulkEditUserModalProps {
|
||||
|
|
@ -250,7 +259,7 @@ const BulkEditUserModal: React.FC<BulkEditUserModalProps> = ({
|
|||
<TableHead className="w-[30%]">User ID</TableHead>
|
||||
<TableHead className="w-[25%]">Email</TableHead>
|
||||
<TableHead className="w-[25%]">Current Role</TableHead>
|
||||
<TableHead className="w-[20%]">Budget</TableHead>
|
||||
<TableHead className={cn("w-[20%]", NUMERIC_CELL_CLASS)}>Budget</TableHead>
|
||||
</TableRow>
|
||||
</TableHeader>
|
||||
<TableBody>
|
||||
|
|
@ -263,7 +272,7 @@ const BulkEditUserModal: React.FC<BulkEditUserModalProps> = ({
|
|||
<TableCell className="text-xs text-foreground">
|
||||
{possibleUIRoles?.[user.user_role]?.ui_label || user.user_role}
|
||||
</TableCell>
|
||||
<TableCell>
|
||||
<TableCell className={NUMERIC_CELL_CLASS}>
|
||||
<MoneyCell value={user.max_budget} decimals={2} emptyText="Unlimited" showZero />
|
||||
</TableCell>
|
||||
</TableRow>
|
||||
|
|
|
|||
|
|
@ -50,6 +50,9 @@ describe("getModelHubTableColumns", () => {
|
|||
expect(screen.getByText("128.0K / 16.4K")).toBeInTheDocument();
|
||||
expect(screen.getByText("$2.50")).toBeInTheDocument();
|
||||
expect(screen.getByText("$10.00")).toBeInTheDocument();
|
||||
expect(screen.getByRole("cell", { name: "128.0K / 16.4K" })).toHaveClass("text-right");
|
||||
expect(screen.getByRole("cell", { name: /\$2\.50/ })).toHaveClass("text-right");
|
||||
expect(screen.getByRole("cell", { name: "gpt-4o" })).not.toHaveClass("text-right");
|
||||
});
|
||||
|
||||
it("shows capability badges only for supported features", () => {
|
||||
|
|
|
|||
|
|
@ -143,7 +143,7 @@ export const getModelHubTableColumns = ({ onModelClick }: ModelHubTableColumnsDe
|
|||
{
|
||||
id: "max_input_tokens",
|
||||
accessorKey: "max_input_tokens",
|
||||
meta: { title: "Tokens", className: "hidden lg:table-cell" },
|
||||
meta: { title: "Tokens", className: "hidden lg:table-cell", numeric: true },
|
||||
header: ({ column }) => <DataTableSortHeader column={column} title="Tokens" />,
|
||||
size: 110,
|
||||
enableSorting: true,
|
||||
|
|
@ -165,7 +165,7 @@ export const getModelHubTableColumns = ({ onModelClick }: ModelHubTableColumnsDe
|
|||
{
|
||||
id: "input_cost_per_token",
|
||||
accessorKey: "input_cost_per_token",
|
||||
meta: { title: "Cost/1M", skeleton: "twoLine" },
|
||||
meta: { title: "Cost/1M", skeleton: "twoLine", numeric: true },
|
||||
header: ({ column }) => <DataTableSortHeader column={column} title="Cost/1M" />,
|
||||
size: 110,
|
||||
enableSorting: true,
|
||||
|
|
|
|||
|
|
@ -161,6 +161,12 @@ describe("sort contract – only backend-sortable columns are sortable", () => {
|
|||
});
|
||||
});
|
||||
|
||||
it("right-aligns Spend / Budget but not Team", () => {
|
||||
renderTable();
|
||||
expect(screen.getByRole("columnheader", { name: "Spend / Budget" })).toHaveClass("text-right");
|
||||
expect(screen.getByRole("columnheader", { name: "Team" })).not.toHaveClass("text-right");
|
||||
});
|
||||
|
||||
it("does not make Spend / Budget sortable (the backend rejects sort_by=spend)", () => {
|
||||
renderTable();
|
||||
expect(screen.queryByText("Spend / Budget").closest("button")).toBeNull();
|
||||
|
|
|
|||
|
|
@ -210,7 +210,7 @@ export const getTeamTableColumns = ({
|
|||
{
|
||||
id: "spend",
|
||||
accessorKey: "spend",
|
||||
meta: { title: "Spend / Budget", skeleton: "meter" },
|
||||
meta: { title: "Spend / Budget", skeleton: "meter", numeric: true },
|
||||
header: "Spend / Budget",
|
||||
size: 200,
|
||||
enableSorting: false,
|
||||
|
|
@ -234,7 +234,7 @@ export const getTeamTableColumns = ({
|
|||
},
|
||||
{
|
||||
id: "members",
|
||||
meta: { title: "Members" },
|
||||
meta: { title: "Members", numeric: true },
|
||||
header: "Members",
|
||||
size: 110,
|
||||
enableSorting: false,
|
||||
|
|
@ -242,7 +242,7 @@ export const getTeamTableColumns = ({
|
|||
},
|
||||
{
|
||||
id: "models",
|
||||
meta: { title: "Models" },
|
||||
meta: { title: "Models", numeric: true },
|
||||
header: "Models",
|
||||
size: 100,
|
||||
enableSorting: false,
|
||||
|
|
|
|||
|
|
@ -207,6 +207,12 @@ it("should render VirtualKeysTable component", () => {
|
|||
expect(screen.getByText("Test Key Alias")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("right-aligns the Spend / Budget column", async () => {
|
||||
renderWithProviders(<VirtualKeysTable />);
|
||||
expect(await screen.findByRole("columnheader", { name: /^Spend/ })).toHaveClass("text-right");
|
||||
expect(screen.getByRole("columnheader", { name: /^Key$/ })).not.toHaveClass("text-right");
|
||||
});
|
||||
|
||||
it("shows the Budget Reset column by default", async () => {
|
||||
renderWithProviders(<VirtualKeysTable />);
|
||||
await waitFor(() => {
|
||||
|
|
|
|||
|
|
@ -267,7 +267,7 @@ export const getKeyTableColumns = ({
|
|||
{
|
||||
id: "spend",
|
||||
accessorKey: "spend",
|
||||
meta: { title: "Spend / Budget", skeleton: "meter" },
|
||||
meta: { title: "Spend / Budget", skeleton: "meter", numeric: true },
|
||||
header: ({ table }) => <DataTableMultiSortHeader table={table} fields={SPEND_BUDGET_SORT_FIELDS} />,
|
||||
size: 180,
|
||||
enableSorting: true,
|
||||
|
|
|
|||
|
|
@ -1,7 +1,15 @@
|
|||
import React, { useState, useEffect } from "react";
|
||||
import { Button, buttonVariants } from "@/components/ui/button";
|
||||
import { Dialog, DialogContent, DialogHeader, DialogTitle } from "@/components/ui/dialog";
|
||||
import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table";
|
||||
import {
|
||||
NUMERIC_CELL_CLASS,
|
||||
Table,
|
||||
TableBody,
|
||||
TableCell,
|
||||
TableHead,
|
||||
TableHeader,
|
||||
TableRow,
|
||||
} from "@/components/ui/table";
|
||||
import { Download, FileText, FileWarning, Trash2, TriangleAlert, Upload } from "lucide-react";
|
||||
import { userCreateCall, invitationCreateCall, getProxyUISettings } from "./networking";
|
||||
import Papa from "papaparse";
|
||||
|
|
@ -798,7 +806,7 @@ const BulkCreateUsersButton: React.FC<BulkCreateUsersProps> = ({
|
|||
<TableHead>Email</TableHead>
|
||||
<TableHead>Role</TableHead>
|
||||
<TableHead>Teams</TableHead>
|
||||
<TableHead>Budget</TableHead>
|
||||
<TableHead className={NUMERIC_CELL_CLASS}>Budget</TableHead>
|
||||
<TableHead>Status</TableHead>
|
||||
</TableRow>
|
||||
</TableHeader>
|
||||
|
|
@ -809,7 +817,7 @@ const BulkCreateUsersButton: React.FC<BulkCreateUsersProps> = ({
|
|||
<TableCell className="whitespace-normal break-words">{record.user_email}</TableCell>
|
||||
<TableCell className="whitespace-normal break-words">{record.user_role}</TableCell>
|
||||
<TableCell className="whitespace-normal break-words">{record.teams}</TableCell>
|
||||
<TableCell>{record.max_budget}</TableCell>
|
||||
<TableCell className={NUMERIC_CELL_CLASS}>{record.max_budget}</TableCell>
|
||||
<TableCell className="whitespace-normal break-words">{renderStatusCell(record)}</TableCell>
|
||||
</TableRow>
|
||||
))}
|
||||
|
|
|
|||
|
|
@ -10,7 +10,16 @@ import { Label } from "@/components/ui/label";
|
|||
import { Badge } from "@/components/ui/badge";
|
||||
import { Skeleton } from "@/components/ui/skeleton";
|
||||
import { Dialog, DialogContent, DialogFooter, DialogHeader, DialogTitle } from "@/components/ui/dialog";
|
||||
import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table";
|
||||
import {
|
||||
NUMERIC_CELL_CLASS,
|
||||
Table,
|
||||
TableBody,
|
||||
TableCell,
|
||||
TableHead,
|
||||
TableHeader,
|
||||
TableRow,
|
||||
} from "@/components/ui/table";
|
||||
import { cn } from "@/lib/cva.config";
|
||||
import { toast } from "@/lib/toast";
|
||||
import { keyListCall, regenerateKeyCall } from "../networking";
|
||||
import { KeyResponse } from "../key_team_helpers/key_list";
|
||||
|
|
@ -176,7 +185,9 @@ const KeysPanel: React.FC<Props> = ({ accessToken, userId, premiumUser }) => {
|
|||
<TableHeader>
|
||||
<TableRow className="bg-muted/50">
|
||||
<TableHead className="text-xs font-semibold uppercase tracking-wide">Key</TableHead>
|
||||
<TableHead className="text-xs font-semibold uppercase tracking-wide">Spend</TableHead>
|
||||
<TableHead className={cn("text-xs font-semibold uppercase tracking-wide", NUMERIC_CELL_CLASS)}>
|
||||
Spend
|
||||
</TableHead>
|
||||
<TableHead className="text-xs font-semibold uppercase tracking-wide">Expires</TableHead>
|
||||
<TableHead className="text-xs font-semibold uppercase tracking-wide">Created</TableHead>
|
||||
{premiumUser && (
|
||||
|
|
@ -220,7 +231,9 @@ const KeysPanel: React.FC<Props> = ({ accessToken, userId, premiumUser }) => {
|
|||
<TableHeader>
|
||||
<TableRow className="bg-muted/50">
|
||||
<TableHead className="text-xs font-semibold uppercase tracking-wide">Key</TableHead>
|
||||
<TableHead className="text-xs font-semibold uppercase tracking-wide">Spend</TableHead>
|
||||
<TableHead className={cn("text-xs font-semibold uppercase tracking-wide", NUMERIC_CELL_CLASS)}>
|
||||
Spend
|
||||
</TableHead>
|
||||
<TableHead className="text-xs font-semibold uppercase tracking-wide">Expires</TableHead>
|
||||
<TableHead className="text-xs font-semibold uppercase tracking-wide">Created</TableHead>
|
||||
{premiumUser && (
|
||||
|
|
@ -237,7 +250,7 @@ const KeysPanel: React.FC<Props> = ({ accessToken, userId, premiumUser }) => {
|
|||
<span className="font-mono text-[13px]">{maskKey(record.key_name)}</span>
|
||||
{record.key_alias && <div className="text-xs text-muted-foreground">{record.key_alias}</div>}
|
||||
</TableCell>
|
||||
<TableCell className="text-[13px]">
|
||||
<TableCell className={cn("text-[13px]", NUMERIC_CELL_CLASS)}>
|
||||
${record.spend?.toFixed(2) ?? "0.00"}
|
||||
{record.max_budget != null && record.max_budget > 0 && (
|
||||
<span className="text-muted-foreground"> / ${record.max_budget.toFixed(2)}</span>
|
||||
|
|
|
|||
|
|
@ -213,3 +213,20 @@ describe("MemberTable actions", () => {
|
|||
expect(screen.getByText("No members found")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
describe("MemberTable numeric columns", () => {
|
||||
it("right-aligns the header and cells of a numeric extra column only", () => {
|
||||
renderTable({
|
||||
members: [MEMBERS[0]],
|
||||
extraColumns: [
|
||||
{ title: "Spend (USD)", key: "spend", numeric: true, render: () => <span>$1.50</span> },
|
||||
{ title: "Joined", key: "joined", render: () => <span>Aug 1</span> },
|
||||
],
|
||||
});
|
||||
|
||||
expect(screen.getByRole("columnheader", { name: "Spend (USD)" })).toHaveClass("text-right", "tabular-nums");
|
||||
expect(screen.getByRole("cell", { name: "$1.50" })).toHaveClass("text-right", "tabular-nums");
|
||||
expect(screen.getByRole("columnheader", { name: "Joined" })).not.toHaveClass("text-right");
|
||||
expect(screen.getByRole("cell", { name: "Aug 1" })).not.toHaveClass("text-right");
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -24,6 +24,7 @@ export interface MemberTableColumn {
|
|||
key: string;
|
||||
render: (member: Member) => React.ReactNode;
|
||||
sortValue?: (member: Member) => MemberTableSortValue;
|
||||
numeric?: boolean;
|
||||
}
|
||||
|
||||
export interface MemberTableProps {
|
||||
|
|
@ -87,6 +88,7 @@ const extraColumnDef = (column: MemberTableColumn): ColumnDef<Member> => {
|
|||
header: () => <span className="font-medium">{column.title}</span>,
|
||||
enableSorting: false,
|
||||
enableGlobalFilter: false,
|
||||
meta: { numeric: column.numeric },
|
||||
cell: ({ row }) => column.render(row.original),
|
||||
};
|
||||
}
|
||||
|
|
@ -97,6 +99,7 @@ const extraColumnDef = (column: MemberTableColumn): ColumnDef<Member> => {
|
|||
sortDescFirst: false,
|
||||
sortUndefined: "last",
|
||||
enableGlobalFilter: false,
|
||||
meta: { numeric: column.numeric },
|
||||
cell: ({ row }) => column.render(row.original),
|
||||
};
|
||||
};
|
||||
|
|
|
|||
|
|
@ -0,0 +1,25 @@
|
|||
import { render, screen } from "@testing-library/react";
|
||||
import { describe, expect, it } from "vitest";
|
||||
|
||||
import { SimpleTable, type SimpleTableColumn } from "./simple_table";
|
||||
|
||||
interface Row {
|
||||
name: string;
|
||||
spend: number;
|
||||
}
|
||||
|
||||
const columns: SimpleTableColumn<Row>[] = [
|
||||
{ header: "Name", accessor: "name" },
|
||||
{ header: "Spend", accessor: "spend", numeric: true },
|
||||
];
|
||||
|
||||
describe("SimpleTable numeric columns", () => {
|
||||
it("right-aligns the header and cells of a numeric column only", () => {
|
||||
render(<SimpleTable data={[{ name: "Alice", spend: 42 }]} columns={columns} />);
|
||||
|
||||
expect(screen.getByRole("columnheader", { name: "Spend" })).toHaveClass("text-right", "tabular-nums");
|
||||
expect(screen.getByRole("cell", { name: "42" })).toHaveClass("text-right", "tabular-nums");
|
||||
expect(screen.getByRole("columnheader", { name: "Name" })).not.toHaveClass("text-right");
|
||||
expect(screen.getByRole("cell", { name: "Alice" })).not.toHaveClass("text-right");
|
||||
});
|
||||
});
|
||||
|
|
@ -1,11 +1,20 @@
|
|||
import React from "react";
|
||||
import { Table, TableHeader, TableRow, TableHead, TableBody, TableCell } from "@/components/ui/table";
|
||||
import {
|
||||
NUMERIC_CELL_CLASS,
|
||||
Table,
|
||||
TableHeader,
|
||||
TableRow,
|
||||
TableHead,
|
||||
TableBody,
|
||||
TableCell,
|
||||
} from "@/components/ui/table";
|
||||
|
||||
export interface SimpleTableColumn<T> {
|
||||
header: string;
|
||||
accessor?: keyof T;
|
||||
cell?: (row: T) => React.ReactNode;
|
||||
width?: string;
|
||||
numeric?: boolean;
|
||||
}
|
||||
|
||||
interface SimpleTableProps<T> {
|
||||
|
|
@ -34,7 +43,11 @@ export function SimpleTable<T>({
|
|||
<TableHeader>
|
||||
<TableRow>
|
||||
{columns.map((column, index) => (
|
||||
<TableHead key={index} style={{ width: column.width }}>
|
||||
<TableHead
|
||||
key={index}
|
||||
style={{ width: column.width }}
|
||||
className={column.numeric ? NUMERIC_CELL_CLASS : undefined}
|
||||
>
|
||||
{column.header}
|
||||
</TableHead>
|
||||
))}
|
||||
|
|
@ -51,7 +64,7 @@ export function SimpleTable<T>({
|
|||
data.map((row, rowIndex) => (
|
||||
<TableRow key={getRowKey ? getRowKey(row, rowIndex) : rowIndex}>
|
||||
{columns.map((column, colIndex) => (
|
||||
<TableCell key={colIndex}>
|
||||
<TableCell key={colIndex} className={column.numeric ? NUMERIC_CELL_CLASS : undefined}>
|
||||
{column.cell ? column.cell(row) : String(row[column.accessor as keyof T] ?? "")}
|
||||
</TableCell>
|
||||
))}
|
||||
|
|
|
|||
|
|
@ -134,6 +134,7 @@ const OrganizationInfoView: React.FC<OrganizationInfoProps> = ({
|
|||
{
|
||||
title: "Spend (USD)",
|
||||
key: "spend",
|
||||
numeric: true,
|
||||
sortValue: (record: Member) => orgMemberFor(record)?.spend ?? null,
|
||||
render: (record: Member) => <MoneyCell value={orgMemberFor(record)?.spend} decimals={4} />,
|
||||
},
|
||||
|
|
|
|||
|
|
@ -141,8 +141,33 @@ const expansionColumns: ColumnDef<Person, unknown>[] = [
|
|||
},
|
||||
];
|
||||
|
||||
const numericColumns: ColumnDef<Person, unknown>[] = [
|
||||
{
|
||||
accessorKey: "name",
|
||||
header: "Name",
|
||||
cell: ({ row }) => <span data-testid="name-cell">{row.original.name}</span>,
|
||||
},
|
||||
{
|
||||
id: "spend",
|
||||
header: ({ column }) => <DataTableSortHeader column={column} title="Spend" />,
|
||||
meta: { numeric: true },
|
||||
cell: () => <span>$1.50</span>,
|
||||
},
|
||||
];
|
||||
|
||||
const CHARLIE_ALICE_BOB: Person[] = [person("c", "Charlie"), person("a", "Alice"), person("b", "Bob")];
|
||||
|
||||
describe("DataTable numeric columns", () => {
|
||||
it("right-aligns the header and cells of a numeric column only", () => {
|
||||
render(<DataTable data={[person("a", "Alice")]} columns={numericColumns} sortingMode="client" />);
|
||||
|
||||
expect(screen.getByRole("columnheader", { name: "Spend" })).toHaveClass("text-right", "tabular-nums");
|
||||
expect(screen.getByRole("cell", { name: "$1.50" })).toHaveClass("text-right", "tabular-nums");
|
||||
expect(screen.getByRole("columnheader", { name: "Name" })).not.toHaveClass("text-right");
|
||||
expect(screen.getByRole("cell", { name: "Alice" })).not.toHaveClass("text-right");
|
||||
});
|
||||
});
|
||||
|
||||
describe("DataTable sorting", () => {
|
||||
it("client mode reorders rows when the sort header is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
|
|
|
|||
|
|
@ -31,6 +31,7 @@ import { Fragment, useEffect, useState } from "react";
|
|||
|
||||
import { Skeleton } from "@/components/ui/skeleton";
|
||||
import {
|
||||
NUMERIC_CELL_CLASS,
|
||||
Table as TableRoot,
|
||||
TableBody,
|
||||
TableCell,
|
||||
|
|
@ -193,7 +194,7 @@ function DataTableHeadCell<TData>({ header, size, stickyHeader, enableColumnResi
|
|||
className={cn(
|
||||
"relative text-muted-foreground",
|
||||
size === "compact" ? "h-8 px-2 py-1 text-xs" : "",
|
||||
meta?.numeric ? "text-right" : "",
|
||||
meta?.numeric ? NUMERIC_CELL_CLASS : "",
|
||||
meta?.className,
|
||||
meta?.headerClassName,
|
||||
sticky.className,
|
||||
|
|
@ -238,7 +239,7 @@ function DataTableBodyCell<TData>({ cell, size, stickyHeader, enableColumnResizi
|
|||
className={cn(
|
||||
"overflow-hidden text-ellipsis",
|
||||
size === "compact" ? "px-2 py-1 text-xs" : "",
|
||||
meta?.numeric ? "text-right tabular-nums" : "",
|
||||
meta?.numeric ? NUMERIC_CELL_CLASS : "",
|
||||
meta?.className,
|
||||
sticky.className,
|
||||
)}
|
||||
|
|
|
|||
|
|
@ -292,6 +292,9 @@ describe("TeamMembersComponent", () => {
|
|||
|
||||
expect(screen.getByText("$100.50")).toBeInTheDocument();
|
||||
expect(screen.getByText("$1,538.26")).toBeInTheDocument();
|
||||
expect(screen.getByRole("cell", { name: "$100.50" })).toHaveClass("text-right");
|
||||
expect(screen.getByRole("columnheader", { name: /^Team Member Budget \(USD\)/ })).toHaveClass("text-right");
|
||||
expect(screen.getByRole("columnheader", { name: "User Email" })).not.toHaveClass("text-right");
|
||||
expect(screen.getByText(/100 RPM/)).toBeInTheDocument();
|
||||
expect(screen.getByText(/10000 TPM/)).toBeInTheDocument();
|
||||
});
|
||||
|
|
|
|||
|
|
@ -186,6 +186,7 @@ export default function TeamMemberTab({
|
|||
</span>
|
||||
),
|
||||
key: "spend",
|
||||
numeric: true,
|
||||
sortValue: (record: Member) => getUserCurrentCycleSpend(record.user_id),
|
||||
render: (record: Member) => <MoneyCell value={getUserCurrentCycleSpend(record.user_id)} decimals={2} />,
|
||||
},
|
||||
|
|
@ -199,6 +200,7 @@ export default function TeamMemberTab({
|
|||
</span>
|
||||
),
|
||||
key: "total_spend",
|
||||
numeric: true,
|
||||
sortValue: (record: Member) => getUserTotalSpend(record.user_id),
|
||||
render: (record: Member) => <MoneyCell value={getUserTotalSpend(record.user_id)} decimals={2} />,
|
||||
},
|
||||
|
|
@ -212,11 +214,12 @@ export default function TeamMemberTab({
|
|||
</span>
|
||||
),
|
||||
key: "budget",
|
||||
numeric: true,
|
||||
sortValue: (record: Member) => getUserBudget(record.user_id),
|
||||
render: (record: Member) => {
|
||||
const source = getUserBudgetSource(record.user_id);
|
||||
return (
|
||||
<span className="flex items-center gap-2">
|
||||
<span className="flex items-center justify-end gap-2">
|
||||
<MoneyCell value={getUserBudget(record.user_id)} decimals={2} emptyText="Unlimited" showZero />
|
||||
{source !== "none" && (
|
||||
<Badge variant={source === "custom" ? "outline" : "secondary"} data-testid="member-budget-source">
|
||||
|
|
|
|||
|
|
@ -131,6 +131,14 @@ describe("TeamVirtualKeysTable", () => {
|
|||
});
|
||||
});
|
||||
|
||||
it("right-aligns the Spend (USD) and Budget (USD) columns", async () => {
|
||||
renderWithProviders(<TeamVirtualKeysTable {...defaultProps} />);
|
||||
|
||||
expect(await screen.findByRole("columnheader", { name: "Spend (USD)" })).toHaveClass("text-right");
|
||||
expect(screen.getByRole("columnheader", { name: "Budget (USD)" })).toHaveClass("text-right");
|
||||
expect(screen.getByRole("columnheader", { name: "Key ID" })).not.toHaveClass("text-right");
|
||||
});
|
||||
|
||||
it("should display keys in table when data is loaded", async () => {
|
||||
mockUseKeys.mockReturnValue({
|
||||
data: {
|
||||
|
|
|
|||
|
|
@ -285,7 +285,7 @@ export function TeamVirtualKeysTable({ teamId, teamAlias, organization }: TeamVi
|
|||
{
|
||||
id: "spend",
|
||||
accessorKey: "spend",
|
||||
meta: { title: "Spend (USD)" },
|
||||
meta: { title: "Spend (USD)", numeric: true },
|
||||
header: ({ column }) => <DataTableSortHeader column={column} title="Spend (USD)" variant="header-cycle" />,
|
||||
size: 100,
|
||||
enableSorting: true,
|
||||
|
|
@ -294,7 +294,7 @@ export function TeamVirtualKeysTable({ teamId, teamAlias, organization }: TeamVi
|
|||
{
|
||||
id: "max_budget",
|
||||
accessorKey: "max_budget",
|
||||
meta: { title: "Budget (USD)" },
|
||||
meta: { title: "Budget (USD)", numeric: true },
|
||||
header: ({ column }) => <DataTableSortHeader column={column} title="Budget (USD)" variant="header-cycle" />,
|
||||
size: 110,
|
||||
enableSorting: true,
|
||||
|
|
|
|||
|
|
@ -4,6 +4,8 @@ import * as React from "react";
|
|||
|
||||
import { cn } from "@/lib/cva.config";
|
||||
|
||||
const NUMERIC_CELL_CLASS = "text-right tabular-nums";
|
||||
|
||||
const Table = React.forwardRef<HTMLTableElement, React.ComponentPropsWithoutRef<"table">>(
|
||||
({ className, ...props }, ref) => (
|
||||
<div data-slot="table-container" className="relative w-full overflow-x-auto">
|
||||
|
|
@ -96,4 +98,4 @@ const TableCaption = React.forwardRef<HTMLTableCaptionElement, React.ComponentPr
|
|||
);
|
||||
TableCaption.displayName = "TableCaption";
|
||||
|
||||
export { Table, TableHeader, TableBody, TableFooter, TableHead, TableRow, TableCell, TableCaption };
|
||||
export { NUMERIC_CELL_CLASS, Table, TableHeader, TableBody, TableFooter, TableHead, TableRow, TableCell, TableCaption };
|
||||
|
|
|
|||
12
uv.lock
generated
12
uv.lock
generated
|
|
@ -7857,14 +7857,14 @@ wheels = [
|
|||
|
||||
[[package]]
|
||||
name = "pyjwt"
|
||||
version = "2.15.0"
|
||||
version = "2.14.0"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "typing-extensions", marker = "python_full_version < '3.11'" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/02/a5/5197bfd06417837ac079921c66fa6393f1dea3557272a263cebfef69e432/pyjwt-2.15.0.tar.gz", hash = "sha256:b11c5f9791d7bf51c2b39a81ed669f6b2dbbd669df2942f6c60167e9e3d1abe4", size = 120513, upload-time = "2026-09-23T16:56:00.689Z" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/af/c3/8a3b59c25070cc61dc517fbdfa5dc0904670c96f605cc69759dc09166b99/pyjwt-2.14.0.tar.gz", hash = "sha256:77283c83fb56ecf566a886c757a714bc83668e38156de2cce8263302f42e0b86", size = 113177, upload-time = "2026-09-11T13:11:54.638Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/e8/55/40e45bf052ee8ee12a4dfd785519660f8effa7b065442b91646ec6828619/pyjwt-2.15.0-py3-none-any.whl", hash = "sha256:7a3742debf6b879e912dbb9819ceec1594be812452b78c5f2e2dfc56564954f8", size = 33680, upload-time = "2026-09-23T16:55:59.241Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/9c/97/672cb32ce0dfea44b740cb7b4f97038463b9cf7c0ead1aacf595572851d6/pyjwt-2.14.0-py3-none-any.whl", hash = "sha256:ad0cef71c756a56e74863c2919cf0985f72decbcfcb550ee2f422e7c62b5eedc", size = 32896, upload-time = "2026-09-11T13:11:53.409Z" },
|
||||
]
|
||||
|
||||
[package.optional-dependencies]
|
||||
|
|
@ -10196,11 +10196,11 @@ wheels = [
|
|||
|
||||
[[package]]
|
||||
name = "urllib3"
|
||||
version = "2.8.0"
|
||||
version = "2.7.0"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/e3/05/b17359e1cefb4f909b5e40b1b90a496d987258916dbbf88e842c729f510e/urllib3-2.8.0.tar.gz", hash = "sha256:63bf2ead4c879426ebf22ef2a781eeb4aa3b4ae798a0435506f8687fd5bb9b63", size = 458972, upload-time = "2026-09-15T19:29:36.253Z" }
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/53/0c/06f8b233b8fd13b9e5ee11424ef85419ba0d8ba0b3138bf360be2ff56953/urllib3-2.7.0.tar.gz", hash = "sha256:231e0ec3b63ceb14667c67be60f2f2c40a518cb38b03af60abc813da26505f4c", size = 433602, upload-time = "2026-05-07T16:13:18.596Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/92/9d/c4e665119135114480843e7ab388fa94d8480650450e6f8e26b70d323a4c/urllib3-2.8.0-py3-none-any.whl", hash = "sha256:0cf3cae568d36aa9576b28dfb35f11328f1cb974ca7647d9475ebb86c75ac6e3", size = 135717, upload-time = "2026-09-15T19:29:34.577Z" },
|
||||
{ url = "https://files.pythonhosted.org/packages/7f/3e/5db95bcf282c52709639744ca2a8b149baccf648e39c8cc87553df9eae0c/urllib3-2.7.0-py3-none-any.whl", hash = "sha256:9fb4c81ebbb1ce9531cce37674bbc6f1360472bc18ca9a553ede278ef7276897", size = 131087, upload-time = "2026-05-07T16:13:17.151Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue