From 72049427569f314a23d743dca05d735b611e60b5 Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Wed, 30 Sep 2026 10:31:48 -0700 Subject: [PATCH 1/6] fix(proxy): keep request-body credentials out of stored spend-log requests (#43635) * fix(proxy): keep request-body aws credentials out of stored spend-log requests * fix(proxy): redact every credential-named request-body field in stored spend-log requests Replace the hard-coded AWS key check in the spend-log request-body sanitizer with SensitiveDataMasker's key classification, so Azure, Vertex, watsonx, OCI, GigaChat, Gemini and header credentials are redacted too. Proxy-stamped key identity metadata is kept. * fix(proxy): keep request identifiers named like keys in stored spend-log requests * refactor(proxy): drop the AWS-only snapshot exclusion now that spend-log redaction is name-based * refactor(proxy): use SensitiveDataMasker's key classification without an exclusion list * refactor(proxy): always redact credential-named fields in stored spend-log payloads --- .../spend_tracking/spend_tracking_utils.py | 17 ++- .../test_spend_tracking_utils.py | 104 +++++++++++++++++- 2 files changed, 115 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index 1c51fb21d6e..e69c0c80420 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -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 diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py index 00223f192ec..f3991c0e494 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py @@ -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): """ From 0736a143267b23c6f9c66b09a278b1e425adeb74 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 30 Sep 2026 10:59:53 -0700 Subject: [PATCH 2/6] fix(ui): right-align money and count columns across tables (#37889) * fix(ui): right-align money and count columns across tables Make numeric the one alignment token for the three shared table wrappers (DataTable, MemberTable, SimpleTable) via NUMERIC_CELL_CLASS, and flag every money, cost, and bare-count column that was still left-aligned. Raw ui/table usages that render spend or budgets get the same class on their header and cell. * test(ui): render the organizations alignment test with providers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(ui): query alignment cells by role instead of DOM traversal Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../agents/_components/AgentsTable.test.tsx | 6 +++++ .../agents/_components/AgentsTableColumns.tsx | 2 +- .../provider_discount_table.test.tsx | 2 ++ .../_components/provider_discount_table.tsx | 3 ++- .../provider_margin_table.test.tsx | 2 ++ .../_components/provider_margin_table.tsx | 3 ++- .../components/AllModelsTable.test.tsx | 2 ++ .../components/ModelsTableColumns.tsx | 2 +- .../old-usage/_components/usage.tsx | 22 ++++++++++------ .../_components/OrganizationsTable.test.tsx | 10 ++++++++ .../_components/OrganizationsTableColumns.tsx | 6 ++--- .../users/_components/BulkEditUsers.tsx | 15 ++++++++--- .../AIHub/ModelHubTableColumns.test.tsx | 3 +++ .../components/AIHub/ModelHubTableColumns.tsx | 4 +-- .../components/TeamsPage/TeamsTable.test.tsx | 6 +++++ .../components/TeamsPage/teamTableColumns.tsx | 6 ++--- .../VirtualKeysPage/VirtualKeysTable.test.tsx | 6 +++++ .../VirtualKeysPage/keyTableColumns.tsx | 2 +- .../components/bulk_create_users_button.tsx | 14 ++++++++--- .../src/components/chat/KeysPanel.tsx | 21 +++++++++++++--- .../common_components/MemberTable.test.tsx | 17 +++++++++++++ .../common_components/MemberTable.tsx | 3 +++ .../common_components/simple_table.test.tsx | 25 +++++++++++++++++++ .../common_components/simple_table.tsx | 19 +++++++++++--- .../organization/organization_view.tsx | 1 + .../shared/DataTable/DataTable.test.tsx | 25 +++++++++++++++++++ .../components/shared/DataTable/DataTable.tsx | 5 ++-- .../components/team/TeamMemberTab.test.tsx | 3 +++ .../src/components/team/TeamMemberTab.tsx | 5 +++- .../team/TeamVirtualKeysTable.test.tsx | 8 ++++++ .../components/team/TeamVirtualKeysTable.tsx | 4 +-- .../src/components/ui/table.tsx | 4 ++- 32 files changed, 217 insertions(+), 39 deletions(-) create mode 100644 ui/litellm-dashboard/src/components/common_components/simple_table.test.tsx diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTable.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTable.test.tsx index bef938cd31c..17bb8bbfec4 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTable.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTable.test.tsx @@ -33,6 +33,12 @@ describe("AgentsTable", () => { } }); + it("right-aligns the Spend (USD) column", () => { + render(); + 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(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTableColumns.tsx index a8fe3973a42..9ec1eb097d2 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTableColumns.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTableColumns.tsx @@ -90,7 +90,7 @@ export const getAgentsTableColumns = ({ { id: "spend", accessorKey: "spend", - meta: { title: "Spend (USD)" }, + meta: { title: "Spend (USD)", numeric: true }, header: ({ column }) => , size: 130, enableSorting: true, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.test.tsx index 6606a4e6aaf..8d1ee100a2a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.test.tsx @@ -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", () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.tsx index fcc4c2af935..3d8be33fc4a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.tsx @@ -80,10 +80,11 @@ const ProviderDiscountTable: React.FC = ({ }, { header: "Discount Percentage", + numeric: true, cell: (row) => { const { displayName } = getProviderLogoAndName(row.provider); return ( -
+
{editingProvider === row.provider ? ( <> { 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", () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_margin_table.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_margin_table.tsx index 04823ac4aa0..5352695ef0a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_margin_table.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_margin_table.tsx @@ -123,10 +123,11 @@ const ProviderMarginTable: React.FC = ({ }, { header: "Margin", + numeric: true, cell: (row) => { const displayName = marginRowDisplayName(row.provider); return ( -
+
{editingProvider === row.provider ? ( <>
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTable.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTable.test.tsx index cc0169d745c..be6130b0288 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTable.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTable.test.tsx @@ -174,6 +174,8 @@ describe("AllModelsTable", () => { const { rerender } = render(); 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(); expect(screen.queryByText(/^\$/)).not.toBeInTheDocument(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/ModelsTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/ModelsTableColumns.tsx index cbc31747688..9581d3db198 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/ModelsTableColumns.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/ModelsTableColumns.tsx @@ -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 }) => , enableSorting: true, size: 130, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.tsx index 889a17bc88d..15b2a30b50c 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.tsx @@ -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 = ({ accessToken, token, userRole, use Provider - Spend + Spend {spendByProvider.map((provider) => ( {provider.provider} - + @@ -840,8 +848,8 @@ const UsagePage: React.FC = ({ accessToken, token, userRole, use Customer - Spend - Total Events + Spend + Total Events @@ -849,10 +857,10 @@ const UsagePage: React.FC = ({ accessToken, token, userRole, use {topUsers?.map((user: any, index: number) => ( {user.end_user} - + - {user.total_count} + {user.total_count} ))} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsTable.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsTable.test.tsx index 9d163fe2c08..839eb406200 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsTable.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsTable.test.tsx @@ -82,6 +82,16 @@ describe("OrganizationsTable", () => { } }); + it("right-aligns the money and count columns only", () => { + renderWithProviders(); + 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(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsTableColumns.tsx index 0fea6c6606e..5f170a32941 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsTableColumns.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsTableColumns.tsx @@ -129,7 +129,7 @@ export const getOrganizationsTableColumns = ({ { id: "spend", accessorKey: "spend", - meta: { title: "Spend (USD)" }, + meta: { title: "Spend (USD)", numeric: true }, header: ({ column }) => , 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, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/users/_components/BulkEditUsers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/users/_components/BulkEditUsers.tsx index 459e3fd8c92..67a648dacd1 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/users/_components/BulkEditUsers.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/users/_components/BulkEditUsers.tsx @@ -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 = ({ User ID Email Current Role - Budget + Budget @@ -263,7 +272,7 @@ const BulkEditUserModal: React.FC = ({ {possibleUIRoles?.[user.user_role]?.ui_label || user.user_role} - + diff --git a/ui/litellm-dashboard/src/components/AIHub/ModelHubTableColumns.test.tsx b/ui/litellm-dashboard/src/components/AIHub/ModelHubTableColumns.test.tsx index 3a2fea66b0c..678b811921c 100644 --- a/ui/litellm-dashboard/src/components/AIHub/ModelHubTableColumns.test.tsx +++ b/ui/litellm-dashboard/src/components/AIHub/ModelHubTableColumns.test.tsx @@ -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", () => { diff --git a/ui/litellm-dashboard/src/components/AIHub/ModelHubTableColumns.tsx b/ui/litellm-dashboard/src/components/AIHub/ModelHubTableColumns.tsx index 9f74771f3b1..fc71a9340ac 100644 --- a/ui/litellm-dashboard/src/components/AIHub/ModelHubTableColumns.tsx +++ b/ui/litellm-dashboard/src/components/AIHub/ModelHubTableColumns.tsx @@ -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 }) => , 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 }) => , size: 110, enableSorting: true, diff --git a/ui/litellm-dashboard/src/components/TeamsPage/TeamsTable.test.tsx b/ui/litellm-dashboard/src/components/TeamsPage/TeamsTable.test.tsx index 9699ea2b7d1..ad89b62d611 100644 --- a/ui/litellm-dashboard/src/components/TeamsPage/TeamsTable.test.tsx +++ b/ui/litellm-dashboard/src/components/TeamsPage/TeamsTable.test.tsx @@ -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(); diff --git a/ui/litellm-dashboard/src/components/TeamsPage/teamTableColumns.tsx b/ui/litellm-dashboard/src/components/TeamsPage/teamTableColumns.tsx index 84369a58307..05bfccf6417 100644 --- a/ui/litellm-dashboard/src/components/TeamsPage/teamTableColumns.tsx +++ b/ui/litellm-dashboard/src/components/TeamsPage/teamTableColumns.tsx @@ -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, diff --git a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx index 7ef7f1fcb09..423723a5938 100644 --- a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx +++ b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx @@ -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(); + 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(); await waitFor(() => { diff --git a/ui/litellm-dashboard/src/components/VirtualKeysPage/keyTableColumns.tsx b/ui/litellm-dashboard/src/components/VirtualKeysPage/keyTableColumns.tsx index 477d7b0ecb4..d69e90b1882 100644 --- a/ui/litellm-dashboard/src/components/VirtualKeysPage/keyTableColumns.tsx +++ b/ui/litellm-dashboard/src/components/VirtualKeysPage/keyTableColumns.tsx @@ -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 }) => , size: 180, enableSorting: true, diff --git a/ui/litellm-dashboard/src/components/bulk_create_users_button.tsx b/ui/litellm-dashboard/src/components/bulk_create_users_button.tsx index 8669f067d6d..b5029e98f9c 100644 --- a/ui/litellm-dashboard/src/components/bulk_create_users_button.tsx +++ b/ui/litellm-dashboard/src/components/bulk_create_users_button.tsx @@ -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 = ({ Email Role Teams - Budget + Budget Status @@ -809,7 +817,7 @@ const BulkCreateUsersButton: React.FC = ({ {record.user_email} {record.user_role} {record.teams} - {record.max_budget} + {record.max_budget} {renderStatusCell(record)} ))} diff --git a/ui/litellm-dashboard/src/components/chat/KeysPanel.tsx b/ui/litellm-dashboard/src/components/chat/KeysPanel.tsx index 8b8bd6c8189..fe37102eaba 100644 --- a/ui/litellm-dashboard/src/components/chat/KeysPanel.tsx +++ b/ui/litellm-dashboard/src/components/chat/KeysPanel.tsx @@ -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 = ({ accessToken, userId, premiumUser }) => { Key - Spend + + Spend + Expires Created {premiumUser && ( @@ -220,7 +231,9 @@ const KeysPanel: React.FC = ({ accessToken, userId, premiumUser }) => { Key - Spend + + Spend + Expires Created {premiumUser && ( @@ -237,7 +250,7 @@ const KeysPanel: React.FC = ({ accessToken, userId, premiumUser }) => { {maskKey(record.key_name)} {record.key_alias &&
{record.key_alias}
}
- + ${record.spend?.toFixed(2) ?? "0.00"} {record.max_budget != null && record.max_budget > 0 && ( / ${record.max_budget.toFixed(2)} diff --git a/ui/litellm-dashboard/src/components/common_components/MemberTable.test.tsx b/ui/litellm-dashboard/src/components/common_components/MemberTable.test.tsx index a16f11eab60..418f27c8188 100644 --- a/ui/litellm-dashboard/src/components/common_components/MemberTable.test.tsx +++ b/ui/litellm-dashboard/src/components/common_components/MemberTable.test.tsx @@ -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: () => $1.50 }, + { title: "Joined", key: "joined", render: () => Aug 1 }, + ], + }); + + 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"); + }); +}); diff --git a/ui/litellm-dashboard/src/components/common_components/MemberTable.tsx b/ui/litellm-dashboard/src/components/common_components/MemberTable.tsx index 6efc82e2d35..3cbe59bd839 100644 --- a/ui/litellm-dashboard/src/components/common_components/MemberTable.tsx +++ b/ui/litellm-dashboard/src/components/common_components/MemberTable.tsx @@ -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 => { header: () => {column.title}, enableSorting: false, enableGlobalFilter: false, + meta: { numeric: column.numeric }, cell: ({ row }) => column.render(row.original), }; } @@ -97,6 +99,7 @@ const extraColumnDef = (column: MemberTableColumn): ColumnDef => { sortDescFirst: false, sortUndefined: "last", enableGlobalFilter: false, + meta: { numeric: column.numeric }, cell: ({ row }) => column.render(row.original), }; }; diff --git a/ui/litellm-dashboard/src/components/common_components/simple_table.test.tsx b/ui/litellm-dashboard/src/components/common_components/simple_table.test.tsx new file mode 100644 index 00000000000..885d0ed51ac --- /dev/null +++ b/ui/litellm-dashboard/src/components/common_components/simple_table.test.tsx @@ -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[] = [ + { 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(); + + 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"); + }); +}); diff --git a/ui/litellm-dashboard/src/components/common_components/simple_table.tsx b/ui/litellm-dashboard/src/components/common_components/simple_table.tsx index 6a30a2e0273..a4b84d28801 100644 --- a/ui/litellm-dashboard/src/components/common_components/simple_table.tsx +++ b/ui/litellm-dashboard/src/components/common_components/simple_table.tsx @@ -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 { header: string; accessor?: keyof T; cell?: (row: T) => React.ReactNode; width?: string; + numeric?: boolean; } interface SimpleTableProps { @@ -34,7 +43,11 @@ export function SimpleTable({ {columns.map((column, index) => ( - + {column.header} ))} @@ -51,7 +64,7 @@ export function SimpleTable({ data.map((row, rowIndex) => ( {columns.map((column, colIndex) => ( - + {column.cell ? column.cell(row) : String(row[column.accessor as keyof T] ?? "")} ))} diff --git a/ui/litellm-dashboard/src/components/organization/organization_view.tsx b/ui/litellm-dashboard/src/components/organization/organization_view.tsx index c800d12ad62..f325e92d2b5 100644 --- a/ui/litellm-dashboard/src/components/organization/organization_view.tsx +++ b/ui/litellm-dashboard/src/components/organization/organization_view.tsx @@ -134,6 +134,7 @@ const OrganizationInfoView: React.FC = ({ { title: "Spend (USD)", key: "spend", + numeric: true, sortValue: (record: Member) => orgMemberFor(record)?.spend ?? null, render: (record: Member) => , }, diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test.tsx b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test.tsx index a7e4befa6ac..df88287b369 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test.tsx +++ b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test.tsx @@ -141,8 +141,33 @@ const expansionColumns: ColumnDef[] = [ }, ]; +const numericColumns: ColumnDef[] = [ + { + accessorKey: "name", + header: "Name", + cell: ({ row }) => {row.original.name}, + }, + { + id: "spend", + header: ({ column }) => , + meta: { numeric: true }, + cell: () => $1.50, + }, +]; + 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(); + + 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(); diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.tsx b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.tsx index 340f8d4f44f..feb1615b4e1 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.tsx +++ b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.tsx @@ -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({ 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({ 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, )} diff --git a/ui/litellm-dashboard/src/components/team/TeamMemberTab.test.tsx b/ui/litellm-dashboard/src/components/team/TeamMemberTab.test.tsx index bd7398bc6c8..1d50a9a4670 100644 --- a/ui/litellm-dashboard/src/components/team/TeamMemberTab.test.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamMemberTab.test.tsx @@ -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(); }); diff --git a/ui/litellm-dashboard/src/components/team/TeamMemberTab.tsx b/ui/litellm-dashboard/src/components/team/TeamMemberTab.tsx index e8576ce6dfe..a24d1b1e7cd 100644 --- a/ui/litellm-dashboard/src/components/team/TeamMemberTab.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamMemberTab.tsx @@ -186,6 +186,7 @@ export default function TeamMemberTab({ ), key: "spend", + numeric: true, sortValue: (record: Member) => getUserCurrentCycleSpend(record.user_id), render: (record: Member) => , }, @@ -199,6 +200,7 @@ export default function TeamMemberTab({ ), key: "total_spend", + numeric: true, sortValue: (record: Member) => getUserTotalSpend(record.user_id), render: (record: Member) => , }, @@ -212,11 +214,12 @@ export default function TeamMemberTab({ ), key: "budget", + numeric: true, sortValue: (record: Member) => getUserBudget(record.user_id), render: (record: Member) => { const source = getUserBudgetSource(record.user_id); return ( - + {source !== "none" && ( diff --git a/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.test.tsx b/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.test.tsx index 08408fb4ff4..7561264ea65 100644 --- a/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.test.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.test.tsx @@ -131,6 +131,14 @@ describe("TeamVirtualKeysTable", () => { }); }); + it("right-aligns the Spend (USD) and Budget (USD) columns", async () => { + renderWithProviders(); + + 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: { diff --git a/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.tsx b/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.tsx index 5b1b71e060e..4e4fa3f4bb2 100644 --- a/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.tsx @@ -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 }) => , 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 }) => , size: 110, enableSorting: true, diff --git a/ui/litellm-dashboard/src/components/ui/table.tsx b/ui/litellm-dashboard/src/components/ui/table.tsx index 6271a9e89ac..1c4c1a981de 100644 --- a/ui/litellm-dashboard/src/components/ui/table.tsx +++ b/ui/litellm-dashboard/src/components/ui/table.tsx @@ -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>( ({ className, ...props }, ref) => (
@@ -96,4 +98,4 @@ const TableCaption = React.forwardRef Date: Wed, 30 Sep 2026 11:11:37 -0700 Subject: [PATCH 3/6] feat(agents): enforce authoritative agent permissions (#43721) * feat(agents): authoritative permissions * fix: enforce authoritative managed agent permissions * fix(agents): only consult the identity store for managed targets is_agent_allowed entered the identity-store path whenever a prisma client was configured, so an ordinary agent paired with an internal user returned 503 instead of 200. Classify the target from the registry first and fall back to the store only when the registry has no entry, so an unmanaged target never depends on the store being reachable. * fix(agents): gate the managed path on an admitted policy object Ten call sites branched on `managed_agent_policy is not None`, which any MagicMock attribute satisfies, so the managed path fired on unmanaged subjects and died in Pydantic validation as a 503. Route every check through a shared helper that requires a real AgentResponse. * test(mcp): stub the writer replica the fresh-policy reads use reload_admitted_user now passes check_db_only through to get_user_object, so the user row is read from writer_db. Point the mocks at the replica the code actually reads and give each parametrized case its own user id. * fix(agents): cap a managed agent at the invoking team's agents resolve_agent_access returned the managed policy's grants before the agent_caller ceiling was applied, so a managed agent acting on behalf of a user reached agents that user's team was never granted. Intersect with the caller ceiling the unmanaged path already honours. * fix(agents): restore token narrowing and scope the private-access suppressions The managed-model check lost its valid_token narrowing when it moved to the shared helper. Make the caller-access resolver public rather than reaching into it from module scope, and give each remaining private access a reason. * docs(agents): drop the comment claiming admins skip the A2A permission check The check has never had an admin bypass on this path, so the comment described behaviour the code does not implement. * test(proxy): stub the writer reads and restore the MCP manager singleton Fresh-policy user lookups read writer_db, so the team and rest-endpoint mocks stubbed a replica the code no longer reads, and the dashboard session fake still had the pre-kwarg signature. The manager reload also rebound global_mcp_server_manager in every MCP module without restoring it, leaking an empty manager into later files. * style: sort imports under the litellm package ruff config * fix(mcp): cap a managed agent's servers and tools at the invoking caller managed_agent_servers and managed_agent_tools returned the agent's own grants without the agent_caller ceiling the unmanaged resolvers apply, so a managed agent reached MCP servers and tools the echoed caller could not. Call the existing ceiling helpers on both axes. * refactor(mcp): return the caller-capped tools without an interim list The ceiling helper already returns a sequence, so materializing it into a list added a mutable collection for nothing. Sort at the return sites instead, which also makes the tool order stable across both branches. * fix(agents): preserve actor ceilings during managed target checks * fix(agents): keep managed permission ceilings authoritative * fix(mcp): fail closed on authoritative caller team outages --------- Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> --- .../mcp_server/auth/managed_agent_access.py | 74 +++ .../mcp_server/auth/user_api_key_auth_mcp.py | 263 ++++++-- .../mcp_server/mcp_server_manager.py | 30 +- .../_experimental/mcp_server/toolset_db.py | 10 +- .../mcp_server/ui_session_utils.py | 5 +- litellm/proxy/_types.py | 2 + .../proxy/agent_endpoints/a2a_endpoints.py | 1 - .../auth/agent_access_groups.py | 32 +- .../agent_endpoints/auth/agent_caller.py | 3 +- .../auth/agent_permission_handler.py | 209 ++++++- .../auth/managed_authorization.py | 84 +++ .../proxy/agent_endpoints/identity_store.py | 4 +- litellm/proxy/auth/auth_checks.py | 161 +++-- litellm/proxy/auth/user_api_key_auth.py | 9 + litellm/proxy/utils.py | 18 +- .../object_permission_repository.py | 7 +- litellm/repositories/team_repository.py | 9 +- litellm/repositories/user_repository.py | 7 +- .../auth/test_managed_agent_access.py | 559 ++++++++++++++++++ .../auth/test_user_api_key_auth_mcp.py | 106 +++- .../mcp_server/test_discoverable_endpoints.py | 2 +- .../mcp_server/test_mcp_server_manager.py | 124 ++-- .../mcp_server/test_proxy_api_credentials.py | 4 +- .../mcp_server/test_rest_endpoints.py | 2 +- .../mcp_server/test_ui_session_utils.py | 4 +- .../auth/test_agent_access_groups.py | 24 + .../auth/test_agent_permission_handler.py | 445 +++++++++++++- .../auth/test_managed_authorization.py | 285 +++++++++ .../proxy/auth/test_auth_checks.py | 228 ++++++- .../proxy/auth/test_user_api_key_auth.py | 29 + .../test_mcp_management_endpoints.py | 4 +- .../test_team_endpoints.py | 12 +- .../test_prisma_client_get_data.py | 25 + 33 files changed, 2532 insertions(+), 249 deletions(-) create mode 100644 litellm/proxy/_experimental/mcp_server/auth/managed_agent_access.py create mode 100644 litellm/proxy/agent_endpoints/auth/managed_authorization.py create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/auth/test_managed_agent_access.py create mode 100644 tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py diff --git a/litellm/proxy/_experimental/mcp_server/auth/managed_agent_access.py b/litellm/proxy/_experimental/mcp_server/auth/managed_agent_access.py new file mode 100644 index 00000000000..096b7eb3c77 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/auth/managed_agent_access.py @@ -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") + ) diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index a93ffaeac9f..457c9b1680b 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -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") diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index c0792c32de2..ec2db433911 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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 {} diff --git a/litellm/proxy/_experimental/mcp_server/toolset_db.py b/litellm/proxy/_experimental/mcp_server/toolset_db.py index 48bad178927..dcbd0064514 100644 --- a/litellm/proxy/_experimental/mcp_server/toolset_db.py +++ b/litellm/proxy/_experimental/mcp_server/toolset_db.py @@ -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 [] diff --git a/litellm/proxy/_experimental/mcp_server/ui_session_utils.py b/litellm/proxy/_experimental/mcp_server/ui_session_utils.py index 107a4818de1..901259c18ad 100644 --- a/litellm/proxy/_experimental/mcp_server/ui_session_utils.py +++ b/litellm/proxy/_experimental/mcp_server/ui_session_utils.py @@ -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 diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index d421363ee92..be06d2e7321 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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( diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py index 2a189a76545..e94d5e7ea78 100644 --- a/litellm/proxy/agent_endpoints/a2a_endpoints.py +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -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, diff --git a/litellm/proxy/agent_endpoints/auth/agent_access_groups.py b/litellm/proxy/agent_endpoints/auth/agent_access_groups.py index 49e5407ff88..67547e82f24 100644 --- a/litellm/proxy/agent_endpoints/auth/agent_access_groups.py +++ b/litellm/proxy/agent_endpoints/auth/agent_access_groups.py @@ -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 () diff --git a/litellm/proxy/agent_endpoints/auth/agent_caller.py b/litellm/proxy/agent_endpoints/auth/agent_caller.py index 47d43e8f71b..1ff6f1ffe04 100644 --- a/litellm/proxy/agent_endpoints/auth/agent_caller.py +++ b/litellm/proxy/agent_endpoints/auth/agent_caller.py @@ -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: diff --git a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py index 9fe74bfee3f..75b99b0ab79 100644 --- a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py +++ b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py @@ -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)) diff --git a/litellm/proxy/agent_endpoints/auth/managed_authorization.py b/litellm/proxy/agent_endpoints/auth/managed_authorization.py new file mode 100644 index 00000000000..a15cf074ad5 --- /dev/null +++ b/litellm/proxy/agent_endpoints/auth/managed_authorization.py @@ -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 diff --git a/litellm/proxy/agent_endpoints/identity_store.py b/litellm/proxy/agent_endpoints/identity_store.py index 0d9d21108e5..3c8163a8838 100644 --- a/litellm/proxy/agent_endpoints/identity_store.py +++ b/litellm/proxy/agent_endpoints/identity_store.py @@ -28,6 +28,7 @@ if TYPE_CHECKING: LiteLLM_AgentIdentityWhereUniqueInput, LiteLLM_AgentsTableInclude, LiteLLM_AgentsTableWhereUniqueInput, + LiteLLM_RetiredAgentWhereUniqueInput, LiteLLM_VerifiedSubjectCreateInput, LiteLLM_VerifiedSubjectUpsertInput, LiteLLM_VerifiedSubjectWhereUniqueInput, @@ -183,7 +184,8 @@ class AgentIdentityStore: if self.retired_agents is None: return AgentIdentityFailure(code="policy_unavailable", message="Agent history is unavailable") try: - return await self.retired_agents.table.find_unique(where={"original_agent_id": agent_id}) is not None + where: Final[LiteLLM_RetiredAgentWhereUniqueInput] = {"original_agent_id": agent_id} + return await self.retired_agents.table.find_unique(where=where) is not None except Exception: return AgentIdentityFailure(code="policy_unavailable", message="Agent history is unavailable") diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index e8f335cb348..3ec430332ee 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -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 diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index ed3ec7b4dde..36c6a4c476b 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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, diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 43ad433c19b..c0e36e6e172 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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") diff --git a/litellm/repositories/object_permission_repository.py b/litellm/repositories/object_permission_repository.py index 7736939c696..b732d2ff94c 100644 --- a/litellm/repositories/object_permission_repository.py +++ b/litellm/repositories/object_permission_repository.py @@ -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]: diff --git a/litellm/repositories/team_repository.py b/litellm/repositories/team_repository.py index cbe263699c9..57f6fd33c11 100644 --- a/litellm/repositories/team_repository.py +++ b/litellm/repositories/team_repository.py @@ -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"]: diff --git a/litellm/repositories/user_repository.py b/litellm/repositories/user_repository.py index 87eb45f262d..4a2aea46197 100644 --- a/litellm/repositories/user_repository.py +++ b/litellm/repositories/user_repository.py @@ -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]: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_managed_agent_access.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_managed_agent_access.py new file mode 100644 index 00000000000..90403e5553f --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_managed_agent_access.py @@ -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"} diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index fc4d7b45785..0d0c65e3650 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index b1f0b3fa67e..4e27ec134d4 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index a1dc0e779da..a900ad50dfb 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_proxy_api_credentials.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_proxy_api_credentials.py index ed3e5f48516..8484dfdde72 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_proxy_api_credentials.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_proxy_api_credentials.py @@ -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) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index e82ab28bb4c..0573fc0d144 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_ui_session_utils.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_ui_session_utils.py index a5f6994b1a7..816ccc5e7e6 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_ui_session_utils.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_ui_session_utils.py @@ -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 diff --git a/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_access_groups.py b/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_access_groups.py index e744e84d671..08cceb0d967 100644 --- a/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_access_groups.py +++ b/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_access_groups.py @@ -144,3 +144,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 diff --git a/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_permission_handler.py b/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_permission_handler.py index a87716375e8..4b8d28e2406 100644 --- a/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_permission_handler.py +++ b/tests/test_litellm/proxy/agent_endpoints/auth/test_agent_permission_handler.py @@ -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) diff --git a/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py b/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py new file mode 100644 index 00000000000..7747e5eff71 --- /dev/null +++ b/tests/test_litellm/proxy/agent_endpoints/auth/test_managed_authorization.py @@ -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 diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 9803371c180..353249dddf0 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -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"]) == [] diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index d1973e1693b..b2da7f30926 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -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() diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index 160cf8be4e0..f5fc5ae24d4 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -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, diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index b066b3b80e6..a53894fcd1b 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -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"]) ) diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_get_data.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_get_data.py index 672dd1eb674..05c4f9d8a67 100644 --- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_get_data.py +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_get_data.py @@ -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) From 3930c5bab664dc506af05be6a1bc09575870bf0d Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 30 Sep 2026 11:22:25 -0700 Subject: [PATCH 4/6] fix(proxy): strip caller credentials from websocket passthrough (#43855) * fix(proxy): strip caller credentials from websocket passthrough Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): cover configured x-api-key in websocket passthrough credential test Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: oliver Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../pass_through_endpoints.py | 2 +- .../test_pass_through_endpoints.py | 63 ++++++++++++++++++- 2 files changed, 63 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index e0a4184291e..2e7f9c4a41c 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -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]: diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 3469df082e0..89feb2b6426 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -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: From 253627f484967a4088c40456570e2d6755f23a65 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 30 Sep 2026 18:43:32 +0000 Subject: [PATCH 5/6] fix(ui): render access group MCP and agent selections as wrapping chips (#41228) * fix(ui): render access group MCP and agent selections as wrapping chips Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(ui): cover 20 selected MCP servers rendering as separate chips Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(ui): cover MCP and agent chip selection in access group create dialog Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: ryan-crabbe-berri --- .../AccessGroupsModal/AccessGroupBaseForm.tsx | 61 ++----------------- .../AccessGroupEditModal.integration.test.tsx | 50 ++++++++++++++- .../AccessGroupCreateDialog.test.tsx | 25 ++++++++ .../AccessGroupCreateDialog.tsx | 61 ++----------------- 4 files changed, 83 insertions(+), 114 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsModal/AccessGroupBaseForm.tsx b/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsModal/AccessGroupBaseForm.tsx index f8ec3b5e1e7..33094565d6c 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsModal/AccessGroupBaseForm.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsModal/AccessGroupBaseForm.tsx @@ -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) => ( - -); - interface AccessGroupBaseFormProps { form: UseFormReturn; isNameDisabled?: boolean; @@ -145,15 +98,13 @@ export function AccessGroupBaseForm({ - {({ id, value, onChange, "aria-invalid": ariaInvalid, "aria-describedby": ariaDescribedBy }) => ( + {({ id, value, onChange }) => ( )} @@ -161,15 +112,13 @@ export function AccessGroupBaseForm({ - {({ id, value, onChange, "aria-invalid": ariaInvalid, "aria-describedby": ariaDescribedBy }) => ( + {({ id, value, onChange }) => ( )} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsModal/AccessGroupEditModal.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsModal/AccessGroupEditModal.integration.test.tsx index bd77ad8e897..7c4e218261f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsModal/AccessGroupEditModal.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsModal/AccessGroupEditModal.integration.test.tsx @@ -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: [] }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/access-group-create/AccessGroupCreateDialog.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/access-group-create/AccessGroupCreateDialog.test.tsx index 1ea7286c686..97afcca51c3 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/access-group-create/AccessGroupCreateDialog.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/access-group-create/AccessGroupCreateDialog.test.tsx @@ -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({ diff --git a/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/access-group-create/AccessGroupCreateDialog.tsx b/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/access-group-create/AccessGroupCreateDialog.tsx index a7f2ee18521..9965884728a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/access-group-create/AccessGroupCreateDialog.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/access-group-create/AccessGroupCreateDialog.tsx @@ -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) => ( - -); - const defaultCreateAccessGroup = async (body: AccessGroupCreateBody): Promise => { const { data } = await fetchClient.POST("/v1/access_group", { body }); return data; @@ -193,15 +146,13 @@ export const AccessGroupCreateDialog = ({ - {({ id, value, onChange, "aria-invalid": ariaInvalid, "aria-describedby": ariaDescribedBy }) => ( + {({ id, value, onChange }) => ( )} @@ -209,15 +160,13 @@ export const AccessGroupCreateDialog = ({ - {({ id, value, onChange, "aria-invalid": ariaInvalid, "aria-describedby": ariaDescribedBy }) => ( + {({ id, value, onChange }) => ( )} From 264b09ac8d5753f157ad65b529adf6f52ce869b7 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 30 Sep 2026 11:44:35 -0700 Subject: [PATCH 6/6] fix(responses): scan and mask top-level instructions with guardrails (#43629) * fix(responses): scan and mask top-level instructions with guardrails The Responses guardrail translation handler put a non-empty top-level instructions field into structured_messages as a system row but never into the flat texts list, so guardrails that scan texts skipped it, flat-text masking could not rewrite it, and PANW latest-only selection failed its alignment guard whenever instructions were present. Seed texts with the instructions row, carry that offset into the flat-text write-back so a rewritten row lands on data["instructions"], and account for the leading row in the PANW Responses alignment. Resolves LIT-8931 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(responses): reject empty guardrail rewrites instead of forwarding raw input An explicit texts=[] answer from a guardrail now fails the count check and raises UnappliableRequestRewrite like any other misaligned rewrite; only a missing texts key means no rewrite. Types the out-param as dict[str, object] and adds integration coverage for instructions blocking, masking, empty instructions, tool loops, latest-only and concurrent workers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(responses): type the texts-replacing guardrail helper explicitly Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(responses): honor skip_system_message_in_guardrail for instructions and system input items Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(responses): cover skip_system_message_in_guardrail on the live proxy Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(responses): keep skipped rows through full-coverage rewrites and align latest-only with skip_system Trust a guardrail's structured_messages_cover_full_request claim only when it returns as many rows as the full normalized request, otherwise merge the scoped rows back so skipped instructions and system items survive the write-back. Make PANW's Responses reasoning alignment skip-aware so latest-only still picks the latest user turn when system content is excluded from texts. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(responses): annotate new guardrail tests with return types Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(responses): treat an empty guardrail texts answer as no rewrite like chat completions Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(responses): type the guardrail test doubles explicitly Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yucheng Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../guardrail_translation/handler.py | 114 +++- .../panw_prisma_airs/panw_prisma_airs.py | 42 +- .../observability/test_guardrail_effects.py | 499 ++++++++++++++++++ .../guardrail_hooks/test_crowdstrike_aidr.py | 48 +- .../guardrail_hooks/test_panw_prisma_airs.py | 69 ++- ...test_openai_responses_guardrail_handler.py | 320 ++++++++++- 6 files changed, 1001 insertions(+), 91 deletions(-) diff --git a/litellm/llms/openai/responses/guardrail_translation/handler.py b/litellm/llms/openai/responses/guardrail_translation/handler.py index d6d68e0607a..620d0554bb1 100644 --- a/litellm/llms/openai/responses/guardrail_translation/handler.py +++ b/litellm/llms/openai/responses/guardrail_translation/handler.py @@ -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]: diff --git a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py index e4822195bec..d51e7c8b8fb 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py +++ b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py @@ -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) diff --git a/tests/integration/observability/test_guardrail_effects.py b/tests/integration/observability/test_guardrail_effects.py index c448473391f..d377afb206c 100644 --- a/tests/integration/observability/test_guardrail_effects.py +++ b/tests/integration/observability/test_guardrail_effects.py @@ -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, "")}} 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, ""), f"{name}: sent {sent}" + assert sent["input"] == shape, f"{name}: sent {sent}" + + +_SSN: Final = "123-45-6789" +_MASKED_SSN: Final = "" +_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 diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py index a1aae119d56..c95f7123221 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py @@ -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 "), ) assert cast(list, responses_input[0]["content"])[0]["text"] == "my ssn is " + assert data.get("instructions") == (None if instructions is None else "my ssn is ") @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py index 5db3e11ac06..dba67e7b7bc 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py @@ -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 diff --git a/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py b/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py index a6b930db7a9..88d8169e196 100644 --- a/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py +++ b/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py @@ -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: