diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json
index 32559b98aea..663a8e0d6d8 100644
--- a/litellm/proxy/_lazy_openapi_snapshot.json
+++ b/litellm/proxy/_lazy_openapi_snapshot.json
@@ -52122,6 +52122,16 @@
],
"title": "Updated By"
},
+ "user": {
+ "anyOf": [
+ {
+ "$ref": "#/components/schemas/ToolDiscoveryUser"
+ },
+ {
+ "type": "null"
+ }
+ ]
+ },
"user_agent": {
"anyOf": [
{
@@ -52160,6 +52170,41 @@
"title": "ToolDetailResponse",
"type": "object"
},
+ "ToolDiscoveryUser": {
+ "properties": {
+ "user_alias": {
+ "anyOf": [
+ {
+ "type": "string"
+ },
+ {
+ "type": "null"
+ }
+ ],
+ "title": "User Alias"
+ },
+ "user_email": {
+ "anyOf": [
+ {
+ "type": "string"
+ },
+ {
+ "type": "null"
+ }
+ ],
+ "title": "User Email"
+ },
+ "user_id": {
+ "title": "User Id",
+ "type": "string"
+ }
+ },
+ "required": [
+ "user_id"
+ ],
+ "title": "ToolDiscoveryUser",
+ "type": "object"
+ },
"ToolListResponse": {
"properties": {
"tools": {
diff --git a/litellm/proxy/db/tool_registry_writer.py b/litellm/proxy/db/tool_registry_writer.py
index cd0aa75b859..cef90eb89c2 100644
--- a/litellm/proxy/db/tool_registry_writer.py
+++ b/litellm/proxy/db/tool_registry_writer.py
@@ -8,6 +8,7 @@ Admins use the management endpoints to read and update input_policy / output_pol
import uuid
from collections.abc import Mapping, Sequence
from datetime import datetime, timezone
+from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Protocol
from pydantic import TypeAdapter
@@ -18,8 +19,11 @@ from litellm.proxy.db.exception_handler import call_with_db_reconnect_retry
from litellm.repositories.object_permission_repository import ObjectPermissionRepository
from litellm.repositories.prisma_protocols import TableActions
from litellm.repositories.table_repositories import ToolRepository
+from litellm.repositories.user_repository import UserRepository
+from litellm.repositories.verification_token_repository import VerificationTokenRepository
from litellm.types.tool_management import (
LiteLLM_ToolTableRow,
+ ToolDiscoveryUser,
ToolPolicyOverrideRow,
)
@@ -155,18 +159,65 @@ async def batch_upsert_tools(
verbose_proxy_logger.error("tool_registry_writer batch_upsert_tools error: %s", e)
+_NO_OWNERS: Final[Mapping[str, ToolDiscoveryUser]] = MappingProxyType({})
+
+
+async def _key_owners(prisma_client: "PrismaClient", key_hashes: frozenset[str]) -> Mapping[str, ToolDiscoveryUser]:
+ """Map each key hash to the user that owns the key, skipping keys without an owner or an unknown owner."""
+ if not key_hashes:
+ return _NO_OWNERS
+ keys: Final = await VerificationTokenRepository(prisma_client).find_many_in("token", sorted(key_hashes))
+ owner_ids: Final = frozenset(key.user_id for key in keys if key.user_id)
+ if not owner_ids:
+ return _NO_OWNERS
+ users: Final = await UserRepository(prisma_client).find_many_in("user_id", sorted(owner_ids))
+ users_by_id: Final = MappingProxyType(
+ {
+ user.user_id: ToolDiscoveryUser(
+ user_id=user.user_id, user_email=user.user_email, user_alias=user.user_alias
+ )
+ for user in users
+ }
+ )
+ return MappingProxyType(
+ {key.token: users_by_id[key.user_id] for key in keys if key.token and key.user_id in users_by_id}
+ )
+
+
+async def _key_owners_or_none(
+ prisma_client: "PrismaClient", key_hashes: frozenset[str]
+) -> Mapping[str, ToolDiscoveryUser]:
+ from prisma.errors import PrismaError
+
+ try:
+ return await _key_owners(prisma_client, key_hashes)
+ except PrismaError as e:
+ verbose_proxy_logger.error("tool_registry_writer owner lookup error: %s", e)
+ return _NO_OWNERS
+
+
+async def _with_owners(
+ prisma_client: "PrismaClient", tools: Sequence[LiteLLM_ToolTableRow]
+) -> tuple[LiteLLM_ToolTableRow, ...]:
+ """Attach to each tool the user owning the key that discovered it; tools stay listed when that lookup fails."""
+ owners: Final = await _key_owners_or_none(
+ prisma_client, frozenset(tool.key_hash for tool in tools if tool.key_hash)
+ )
+ return tuple(tool.model_copy(update=MappingProxyType({"user": owners.get(tool.key_hash or "")})) for tool in tools)
+
+
async def list_tools(
prisma_client: "PrismaClient",
input_policy: str | None = None,
) -> list[LiteLLM_ToolTableRow]:
- """Return all tools, optionally filtered by input_policy."""
+ """Return all tools, optionally filtered by input_policy, each with the user owning the key that discovered it."""
try:
where: Final[Mapping[str, str]] = {"input_policy": input_policy} if input_policy is not None else {}
rows: Final = await _tool_table_actions(prisma_client).find_many(
where=where,
order={"created_at": "desc"},
)
- return [_row_to_model(row) for row in rows]
+ return list(await _with_owners(prisma_client, tuple(_row_to_model(row) for row in rows)))
except Exception as e:
verbose_proxy_logger.error("tool_registry_writer list_tools error: %s", e)
return []
@@ -176,14 +227,14 @@ async def get_tool(
prisma_client: "PrismaClient",
tool_name: str,
) -> LiteLLM_ToolTableRow | None:
- """Return a single tool row by tool_name."""
+ """Return a single tool row by tool_name, with the user owning the key that discovered it."""
try:
row: Final = await _tool_table_actions(prisma_client).find_unique(
where={"tool_name": tool_name},
)
if row is None:
return None
- return _row_to_model(row)
+ return (await _with_owners(prisma_client, (_row_to_model(row),)))[0]
except Exception as e:
verbose_proxy_logger.error("tool_registry_writer get_tool error: %s", e)
return None
diff --git a/litellm/repositories/base_repository.py b/litellm/repositories/base_repository.py
index 065842b39e2..81fba770b70 100644
--- a/litellm/repositories/base_repository.py
+++ b/litellm/repositories/base_repository.py
@@ -3,11 +3,12 @@ Base repository class with common functionality.
"""
from abc import ABC, abstractmethod
-from collections.abc import Iterable, Mapping, Sequence
+from collections.abc import Hashable, Iterable, Mapping, Sequence
from typing import Any, Final, Generic, Protocol, TypeVar, runtime_checkable
from pydantic import BaseModel
+from litellm.repositories.chunked_in import find_many_in
from litellm.repositories.prisma_protocols import TableActions
T = TypeVar("T", bound=BaseModel)
@@ -92,6 +93,10 @@ class BaseRepository(ABC, Generic[T]):
)
return self._to_model_list(records)
+ async def find_many_in(self, field: str, values: Iterable[Hashable]) -> list[T]:
+ """Records whose `field` is one of `values`, queried in chunks that stay under the bind-parameter cap."""
+ return self._to_model_list(await find_many_in(self.table, field, values))
+
async def create(self, data: Mapping[str, object]) -> T:
"""Create a new record."""
record: Final = await self.table.create(data=data)
diff --git a/litellm/types/tool_management.py b/litellm/types/tool_management.py
index 6fc19250ae9..13553dbecc6 100644
--- a/litellm/types/tool_management.py
+++ b/litellm/types/tool_management.py
@@ -13,6 +13,12 @@ ToolInputPolicy = Literal["trusted", "untrusted", "blocked"]
ToolOutputPolicy = Literal["trusted", "untrusted"]
+class ToolDiscoveryUser(BaseModel):
+ user_id: str
+ user_email: str | None = None
+ user_alias: str | None = None
+
+
class LiteLLM_ToolTableRow(BaseModel):
tool_id: str
tool_name: str
@@ -25,6 +31,7 @@ class LiteLLM_ToolTableRow(BaseModel):
team_id: str | None = None
key_alias: str | None = None
user_agent: str | None = None
+ user: ToolDiscoveryUser | None = None
last_used_at: datetime | None = None
created_at: datetime | None = None
updated_at: datetime | None = None
diff --git a/tests/e2e/ui/fixtures/pages.ts b/tests/e2e/ui/fixtures/pages.ts
index ba5887f3113..8210334c166 100644
--- a/tests/e2e/ui/fixtures/pages.ts
+++ b/tests/e2e/ui/fixtures/pages.ts
@@ -26,6 +26,7 @@ export enum Page {
Logs = "logs",
McpServers = "mcp-servers",
SearchTools = "search-tools",
+ ToolPolicies = "tool-policies",
TagManagement = "tag-management",
VectorStores = "vector-stores",
NewUsage = "new_usage",
diff --git a/tests/e2e/ui/tests/integrationCritical/expected.json b/tests/e2e/ui/tests/integrationCritical/expected.json
index 1614b188188..c6ee6051cd4 100644
--- a/tests/e2e/ui/tests/integrationCritical/expected.json
+++ b/tests/e2e/ui/tests/integrationCritical/expected.json
@@ -8,5 +8,6 @@
"tests/e2e/ui/tests/integrationCritical/mcpUserEnvVars.spec.ts::a server without per-user variables shows no credential row",
"tests/e2e/ui/tests/integrationCritical/mcpUserEnvVars.spec.ts::clearing credentials for a server deleted underneath the modal reports the failure without losing the page",
"tests/e2e/ui/tests/integrationCritical/costOptimizationModelGroups.spec.ts::cache leakage by model merges a deployment's resolved and requested model names into its model group",
- "tests/e2e/ui/tests/integrationCritical/logsDrawerCredentialCanary.spec.ts::the Logs drawer renders the stored request without the deployment api_key"
+ "tests/e2e/ui/tests/integrationCritical/logsDrawerCredentialCanary.spec.ts::the Logs drawer renders the stored request without the deployment api_key",
+ "tests/e2e/ui/tests/integrationCritical/toolPoliciesUserColumn.spec.ts::the Tool Policies page names the user behind the key that discovered a tool"
]
diff --git a/tests/e2e/ui/tests/integrationCritical/toolPoliciesUserColumn.spec.ts b/tests/e2e/ui/tests/integrationCritical/toolPoliciesUserColumn.spec.ts
new file mode 100644
index 00000000000..c65c8774d89
--- /dev/null
+++ b/tests/e2e/ui/tests/integrationCritical/toolPoliciesUserColumn.spec.ts
@@ -0,0 +1,186 @@
+import { test, expect, type APIRequestContext } from "@playwright/test";
+import { randomUUID } from "node:crypto";
+import { execFileSync } from "node:child_process";
+import * as path from "node:path";
+import { Page } from "../../fixtures/pages";
+import { dismissFeedbackPopup, navigateToPage } from "../../helpers/navigation";
+
+/**
+ * The Tool Policies table gets a User column: the owner of the key that discovered the tool, shown
+ * as alias (then email, then id) linking to the user's page, and a plain dash when the key has no
+ * owner. Both rows are produced the way a customer produces them, a chat completion carrying a
+ * tool through the proxy, so the column is read from the same registry the proxy writes.
+ */
+const unhex = (): string => randomUUID().replaceAll("-", "");
+
+const toolCall = (model: string, toolName: string) => ({
+ model,
+ messages: [{ role: "user", content: "tool policy user column" }],
+ tools: [
+ {
+ type: "function",
+ function: {
+ name: toolName,
+ description: "integration tool",
+ parameters: { type: "object", properties: {} },
+ },
+ },
+ ],
+});
+
+test("the Tool Policies page names the user behind the key that discovered a tool", async ({
+ page,
+ request,
+}) => {
+ const master = process.env.LITELLM_MASTER_KEY ?? "sk-integration-master";
+ const upstream = (
+ process.env.INTEGRATION_UPSTREAM_URL ?? "http://127.0.0.1:8190"
+ ).replace(/\/+$/, "");
+ const auth = { Authorization: `Bearer ${master}` };
+ const marker = unhex();
+ const alias = `ui-owner-${marker}`;
+ const model = `ui-tool-policies-${marker}`;
+ const ownedTool = `ui_owned_tool_${marker}`;
+ const unownedTool = `ui_unowned_tool_${marker}`;
+ const support = (...args: string[]) =>
+ execFileSync(
+ process.env.INTEGRATION_PYTHON ?? "python",
+ [
+ path.resolve(
+ __dirname,
+ "../../../../integration/_support/tool_rows.py",
+ ),
+ ...args,
+ ],
+ { encoding: "utf8", timeout: 10_000, killSignal: "SIGKILL" },
+ );
+
+ const post = async (api: APIRequestContext, route: string, data: object) => {
+ const response = await api.post(route, { headers: auth, data });
+ expect(response.status(), `POST ${route}: ${await response.text()}`).toBe(
+ 200,
+ );
+ return response.json();
+ };
+
+ let modelId = "";
+ let userId = "";
+ const keys: string[] = [];
+ try {
+ modelId = (
+ await post(request, "/model/new", {
+ model_name: model,
+ litellm_params: {
+ model: `openai/${model}`,
+ api_key: "sk-upstream",
+ api_base: `${upstream}/v1`,
+ },
+ })
+ ).model_id;
+ userId = (
+ await post(request, "/user/new", {
+ user_id: `ui-user-${marker}`,
+ user_alias: alias,
+ user_email: `${alias}@integration.example`,
+ auto_create_key: false,
+ })
+ ).user_id;
+ const ownedKey = (
+ await post(request, "/key/generate", { user_id: userId, models: [model] })
+ ).key;
+ const unownedKey = (
+ await post(request, "/key/generate", { models: [model] })
+ ).key;
+ keys.push(ownedKey, unownedKey);
+ for (const [key, toolName] of [
+ [ownedKey, ownedTool],
+ [unownedKey, unownedTool],
+ ]) {
+ const response = await request.post("/v1/chat/completions", {
+ headers: { Authorization: `Bearer ${key}` },
+ data: toolCall(model, toolName),
+ });
+ expect(response.status(), await response.text()).toBe(200);
+ }
+ await expect
+ .poll(
+ async () => {
+ const response = await request.get("/v1/tool/list", {
+ headers: auth,
+ });
+ if (response.status() !== 200) return [];
+ const names = (
+ (await response.json()).tools as { tool_name: string }[]
+ ).map((tool) => tool.tool_name);
+ return [ownedTool, unownedTool].filter((name) =>
+ names.includes(name),
+ );
+ },
+ {
+ timeout: 70_000,
+ message: "the discovered tools never reached the registry",
+ },
+ )
+ .toEqual([ownedTool, unownedTool]);
+
+ await page.goto("/ui/login");
+ await page.getByPlaceholder("Enter your username").fill("admin");
+ await page.getByPlaceholder("Enter your password").fill(master);
+ await page.getByRole("button", { name: "Login", exact: true }).click();
+ await expect(page).toHaveURL(
+ (url) =>
+ url.pathname.startsWith("/ui") && !url.pathname.includes("login"),
+ );
+ await navigateToPage(page, Page.ToolPolicies);
+ await dismissFeedbackPopup(page);
+
+ const table = page.locator("table").filter({ visible: true }).first();
+ const headers = table.getByRole("columnheader");
+ await expect(headers.filter({ hasText: /^User$/ })).toHaveCount(1, {
+ timeout: 20_000,
+ });
+ const headerTexts = (await headers.allInnerTexts()).map((text) =>
+ text.trim(),
+ );
+ const userColumn = headerTexts.indexOf("User");
+ expect(userColumn, `columns: ${headerTexts.join(", ")}`).toBeGreaterThan(
+ -1,
+ );
+
+ const search = page
+ .getByTestId("datatable-search")
+ .filter({ visible: true });
+ await expect(search).toBeVisible({ timeout: 20_000 });
+ await search.fill(unownedTool);
+ const unownedRow = table
+ .locator("tbody tr")
+ .filter({ hasText: unownedTool });
+ await expect(unownedRow).toHaveCount(1, { timeout: 30_000 });
+ const unownedCell = unownedRow.getByRole("cell").nth(userColumn);
+ await expect(unownedCell).toHaveText("-");
+ await expect(unownedCell.getByRole("link")).toHaveCount(0);
+
+ await search.fill(ownedTool);
+ const ownedRow = table.locator("tbody tr").filter({ hasText: ownedTool });
+ await expect(ownedRow).toHaveCount(1, { timeout: 30_000 });
+ const ownerLink = ownedRow
+ .getByRole("cell")
+ .nth(userColumn)
+ .getByRole("link", { name: alias, exact: true });
+ await expect(ownerLink).toBeVisible();
+ expect(await ownerLink.getAttribute("href")).toContain(
+ `user=${encodeURIComponent(userId)}`,
+ );
+ await ownerLink.click();
+ await expect(page).toHaveURL(
+ (url) =>
+ url.searchParams.get("user") === userId ||
+ url.pathname.includes(userId),
+ );
+ } finally {
+ support("clear", ownedTool, unownedTool);
+ if (keys.length) await post(request, "/key/delete", { keys });
+ if (userId) await post(request, "/user/delete", { user_ids: [userId] });
+ if (modelId) await post(request, "/model/delete", { id: modelId });
+ }
+});
diff --git a/tests/integration/_support/tool_rows.py b/tests/integration/_support/tool_rows.py
new file mode 100644
index 00000000000..bb460022242
--- /dev/null
+++ b/tests/integration/_support/tool_rows.py
@@ -0,0 +1,19 @@
+import json
+import sys
+from typing import Final, LiteralString
+
+from integration._support.database import write_rows
+
+CLEAR_QUERY: Final[LiteralString] = 'DELETE FROM "LiteLLM_ToolTable" WHERE tool_name = %s'
+
+
+def clear(tool_names: tuple[str, ...]) -> None:
+ for tool_name in tool_names:
+ write_rows(CLEAR_QUERY, (tool_name,))
+
+
+if __name__ == "__main__":
+ if sys.argv[1] != "clear":
+ raise SystemExit(f"unknown command: {sys.argv[1]}")
+ clear(tuple(sys.argv[2:]))
+ sys.stdout.write(json.dumps({"cleared": sys.argv[2:]}) + "\n")
diff --git a/tests/integration/management/test_tool_policy_user.py b/tests/integration/management/test_tool_policy_user.py
new file mode 100644
index 00000000000..00b4605ee2e
--- /dev/null
+++ b/tests/integration/management/test_tool_policy_user.py
@@ -0,0 +1,486 @@
+import json
+import time
+import uuid
+from collections.abc import Mapping
+from concurrent.futures import ThreadPoolExecutor
+from contextlib import ExitStack
+from hashlib import sha256
+from pathlib import Path
+from typing import Final, NamedTuple
+
+import jwt
+from cryptography.hazmat.primitives.asymmetric import rsa
+from integration._support.client import JSON_OBJECT, Gateway, Scenario, eventually, object_value, string_value
+from integration._support.database import read_rows, scratch_database, write_rows
+from integration._support.process import owned_proxy, owned_proxy_process
+from integration._support.wire import Reply, Request, wire_server
+from jwt.algorithms import RSAAlgorithm
+from pydantic import JsonValue
+
+from litellm.repositories.chunked_in import IN_LIST_CHUNK_SIZE
+
+AUDIENCE: Final = "litellm-integration"
+KEY_ID: Final = "integration-signing-key"
+CLIENT_CLAIM: Final = "client_id"
+
+
+def _tool_call_request(model: str, tool_name: str) -> dict[str, JsonValue]:
+ return {
+ "model": model,
+ "messages": [{"role": "user", "content": "tool policy user control"}],
+ "tools": [
+ {
+ "type": "function",
+ "function": {
+ "name": tool_name,
+ "description": "integration tool",
+ "parameters": {"type": "object", "properties": {}},
+ },
+ }
+ ],
+ }
+
+
+def _forget_tool(tool_name: str) -> None:
+ write_rows('DELETE FROM "LiteLLM_ToolTable" WHERE tool_name = %s', (tool_name,))
+
+
+def _discovered_tool(gateway: Gateway, tool_name: str) -> dict[str, JsonValue]:
+ def rows() -> list[dict[str, JsonValue]]:
+ tools: Final = gateway.get("/v1/tool/list")["tools"]
+ assert isinstance(tools, list)
+ return [object_value(tool) for tool in tools if object_value(tool)["tool_name"] == tool_name]
+
+ return eventually(rows, lambda found: len(found) == 1, seconds=70)[0]
+
+
+def test_tool_list_reports_the_user_that_owns_the_discovering_key(gateway: Gateway) -> None:
+ with gateway.scenario() as scenario:
+ model: Final = scenario.model()
+ alias: Final = "integration-alias-" + uuid.uuid4().hex
+ user: Final = scenario.user(user_alias=alias, user_email=f"{alias}@integration.example")
+ key: Final = scenario.key(user_id=user, models=[model])
+ tool_name: Final = "integration_tool_" + uuid.uuid4().hex
+ scenario.cleanups.callback(_forget_tool, tool_name)
+ response: Final = gateway.request("POST", "/v1/chat/completions", _tool_call_request(model, tool_name), key=key)
+ assert response.status_code == 200, response.text
+ tool: Final = _discovered_tool(gateway, tool_name)
+ assert tool["key_hash"] == sha256(key.encode()).hexdigest(), tool
+ assert tool["user"] == {"user_id": user, "user_email": f"{alias}@integration.example", "user_alias": alias}, (
+ tool
+ )
+
+
+def test_tool_list_reports_no_user_for_a_key_without_an_owner(gateway: Gateway) -> None:
+ with gateway.scenario() as scenario:
+ model: Final = scenario.model()
+ key: Final = scenario.key(models=[model])
+ tool_name: Final = "integration_tool_" + uuid.uuid4().hex
+ scenario.cleanups.callback(_forget_tool, tool_name)
+ response: Final = gateway.request("POST", "/v1/chat/completions", _tool_call_request(model, tool_name), key=key)
+ assert response.status_code == 200, response.text
+ tool: Final = _discovered_tool(gateway, tool_name)
+ assert tool["key_hash"] == sha256(key.encode()).hexdigest(), tool
+ assert tool["user"] is None, tool
+
+
+JWT_SETTINGS: Final[Mapping[str, JsonValue]] = {
+ "enable_jwt_auth": True,
+ "litellm_jwtauth": {
+ "user_id_jwt_field": "sub",
+ "user_email_jwt_field": "email",
+ "user_id_upsert": True,
+ "virtual_key_claim_field": CLIENT_CLAIM,
+ "unregistered_jwt_client_behavior": "auto_register",
+ },
+}
+
+
+def _proxy_config(
+ directory: Path, model: str, upstream_url: str, general_settings: Mapping[str, JsonValue] = JWT_SETTINGS
+) -> Path:
+ config: Final = directory / "tool_policy_user_config.yaml"
+ config.write_text(
+ json.dumps(
+ {
+ "model_list": [
+ {
+ "model_name": model,
+ "litellm_params": {
+ "model": "openai/" + model,
+ "api_base": upstream_url + "/v1",
+ "api_key": "sk-upstream",
+ },
+ }
+ ],
+ "general_settings": {
+ "master_key": "os.environ/LITELLM_MASTER_KEY",
+ "database_url": "os.environ/DATABASE_URL",
+ "store_model_in_db": True,
+ "proxy_batch_write_at": 1,
+ "proxy_batch_polling_interval": 1,
+ **general_settings,
+ },
+ "router_settings": {"disable_cooldowns": True},
+ }
+ )
+ )
+ return config
+
+
+def _signed_token(private_key: rsa.RSAPrivateKey, user_id: str, email: str, client_id: str) -> str:
+ now: Final = int(time.time())
+ return jwt.encode(
+ {"sub": user_id, "email": email, CLIENT_CLAIM: client_id, "aud": AUDIENCE, "iat": now, "exp": now + 300},
+ private_key,
+ algorithm="RS256",
+ headers={"kid": KEY_ID},
+ )
+
+
+def _forget_auto_registered_client(client_id: str, user_id: str) -> None:
+ write_rows(
+ 'DELETE FROM "LiteLLM_VerificationToken" WHERE token IN '
+ '(SELECT token FROM "LiteLLM_JWTKeyMapping" WHERE jwt_claim_value = %s)',
+ (client_id,),
+ )
+ write_rows('DELETE FROM "LiteLLM_JWTKeyMapping" WHERE jwt_claim_value = %s', (client_id,))
+ write_rows('DELETE FROM "LiteLLM_UserTable" WHERE user_id = %s', (user_id,))
+
+
+def test_tool_list_reports_the_jwt_user_behind_an_auto_registered_key(gateway: Gateway, tmp_path: Path) -> None:
+ private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048)
+ public_jwk: Final = json.loads(RSAAlgorithm.to_jwk(private_key.public_key()))
+ jwks: Final = json.dumps({"keys": [{**public_jwk, "kid": KEY_ID, "use": "sig", "alg": "RS256"}]}).encode()
+
+ def respond(request: Request) -> Reply:
+ assert request.target == "/jwks", request
+ return Reply(body=jwks)
+
+ model: Final = "integration-jwt-" + uuid.uuid4().hex
+ with wire_server(respond) as issuer:
+ config: Final = _proxy_config(tmp_path, model, gateway.upstream_url)
+ overrides: Final = {"JWT_PUBLIC_KEY_URL": issuer.url + "/jwks", "JWT_AUDIENCE": AUDIENCE}
+ with owned_proxy(gateway, tmp_path, overrides, config=config) as candidate, candidate.scenario() as scenario:
+ user: Final = "integration-jwt-user-" + uuid.uuid4().hex
+ email: Final = f"{user}@integration.example"
+ client_id: Final = "integration-client-" + uuid.uuid4().hex
+ tool_name: Final = "integration_tool_" + uuid.uuid4().hex
+ scenario.cleanups.callback(_forget_tool, tool_name)
+ scenario.cleanups.callback(_forget_auto_registered_client, client_id, user)
+ response: Final = candidate.request(
+ "POST",
+ "/v1/chat/completions",
+ _tool_call_request(model, tool_name),
+ key=_signed_token(private_key, user, email, client_id),
+ )
+ assert response.status_code == 200, response.text
+ mapped: Final = read_rows(
+ 'SELECT token FROM "LiteLLM_JWTKeyMapping" WHERE jwt_claim_name = %s AND jwt_claim_value = %s',
+ (CLIENT_CLAIM, client_id),
+ )
+ assert len(mapped) == 1, mapped
+ assert read_rows(
+ 'SELECT user_id FROM "LiteLLM_VerificationToken" WHERE token = %s', (mapped[0]["token"],)
+ ) == [{"user_id": user}]
+ tool: Final = _discovered_tool(candidate, tool_name)
+ assert tool["key_hash"] == mapped[0]["token"], tool
+ assert tool["user"] == {"user_id": user, "user_email": email, "user_alias": None}, tool
+
+
+def _owner(user_id: str, email: str | None, alias: str | None) -> dict[str, JsonValue]:
+ return {"user_id": user_id, "user_email": email, "user_alias": alias}
+
+
+def _discover(gateway: Gateway, cleanups: ExitStack, model: str, key: str) -> str:
+ tool_name: Final = "integration_tool_" + uuid.uuid4().hex
+ cleanups.callback(_forget_tool, tool_name)
+ response: Final = gateway.request("POST", "/v1/chat/completions", _tool_call_request(model, tool_name), key=key)
+ assert response.status_code == 200, response.text
+ return tool_name
+
+
+class Owned(NamedTuple):
+ tool_name: str
+ model: str
+ key: str
+ owner: dict[str, JsonValue]
+
+
+def _owned_tool(gateway: Gateway, scenario: Scenario, alias: str | None = None) -> Owned:
+ """A discovered tool, the model and key that discovered it, and the owner the tool routes must report."""
+ model: Final = scenario.model()
+ email: Final = f"{uuid.uuid4().hex}@integration.example"
+ fields: Final[Mapping[str, JsonValue]] = {"user_alias": alias} if alias else {}
+ user: Final = scenario.user(user_email=email, **fields)
+ key: Final = scenario.key(user_id=user, models=[model])
+ return Owned(_discover(gateway, scenario.cleanups, model, key), model, key, _owner(user, email, alias))
+
+
+def _single(gateway: Gateway, tool_name: str) -> dict[str, JsonValue]:
+ return gateway.get(f"/v1/tool/{tool_name}")
+
+
+def _detail_tool(gateway: Gateway, tool_name: str) -> dict[str, JsonValue]:
+ return object_value(gateway.get(f"/v1/tool/{tool_name}/detail")["tool"])
+
+
+def _listed_tools(gateway: Gateway, prefix: str, params: Mapping[str, str] | None = None) -> list[dict[str, JsonValue]]:
+ tools: Final = gateway.get("/v1/tool/list", params)["tools"]
+ assert isinstance(tools, list)
+ return [object_value(tool) for tool in tools if str(object_value(tool)["tool_name"]).startswith(prefix)]
+
+
+def test_tool_get_reports_the_owner_and_null_for_an_unowned_key(gateway: Gateway) -> None:
+ with gateway.scenario() as scenario:
+ owned, model, _, owner = _owned_tool(gateway, scenario, alias="alias-" + uuid.uuid4().hex)
+ unowned: Final = _discover(gateway, scenario.cleanups, model, scenario.key(models=[model]))
+ assert _discovered_tool(gateway, owned)["user"] == owner
+ _discovered_tool(gateway, unowned)
+ assert _single(gateway, owned)["user"] == owner
+ assert _single(gateway, unowned)["user"] is None
+
+
+def test_tool_detail_carries_the_owner_inside_the_tool(gateway: Gateway) -> None:
+ with gateway.scenario() as scenario:
+ tool_name, _, _, owner = _owned_tool(gateway, scenario, alias="alias-" + uuid.uuid4().hex)
+ assert _discovered_tool(gateway, tool_name)["user"] == owner
+ assert _detail_tool(gateway, tool_name)["user"] == owner
+
+
+def test_filtered_tool_list_keeps_the_owner(gateway: Gateway) -> None:
+ with gateway.scenario() as scenario:
+ tool_name, _, _, owner = _owned_tool(gateway, scenario)
+ listed: Final = _discovered_tool(gateway, tool_name)
+ assert listed["input_policy"] == "untrusted", listed
+ filtered: Final = _listed_tools(gateway, tool_name, {"input_policy": "untrusted"})
+ assert [tool["user"] for tool in filtered] == [owner], filtered
+ assert _listed_tools(gateway, tool_name, {"input_policy": "blocked"}) == []
+
+
+def test_two_tools_discovered_by_the_same_key_share_the_owner(gateway: Gateway) -> None:
+ with gateway.scenario() as scenario:
+ first, model, key, owner = _owned_tool(gateway, scenario)
+ second: Final = _discover(gateway, scenario.cleanups, model, key)
+ assert [_discovered_tool(gateway, name)["user"] for name in (first, second)] == [owner, owner]
+
+
+def test_owner_without_alias_or_email_reports_only_the_user_id(gateway: Gateway) -> None:
+ with gateway.scenario() as scenario:
+ model: Final = scenario.model()
+ user: Final = scenario.user()
+ key: Final = scenario.key(user_id=user, models=[model])
+ tool_name: Final = _discover(gateway, scenario.cleanups, model, key)
+ assert _discovered_tool(gateway, tool_name)["user"] == _owner(user, None, None)
+
+
+def test_missing_tool_is_404_on_get_and_detail(gateway: Gateway) -> None:
+ missing: Final = "integration_missing_" + uuid.uuid4().hex
+ for path in (f"/v1/tool/{missing}", f"/v1/tool/{missing}/detail"):
+ response: Final = gateway.request("GET", path)
+ assert response.status_code == 404, response.text
+ assert response.json() == {"detail": f"Tool '{missing}' not found"}
+
+
+def test_non_admin_keys_are_rejected_on_every_tool_read_route(gateway: Gateway) -> None:
+ with gateway.scenario() as scenario:
+ tool_name, _, _, owner = _owned_tool(gateway, scenario)
+ assert _discovered_tool(gateway, tool_name)["user"] == owner
+ internal: Final = scenario.key(user_id=scenario.user(user_role="internal_user"))
+ plain: Final = scenario.key()
+ for key in (internal, plain):
+ for path in ("/v1/tool/list", f"/v1/tool/{tool_name}", f"/v1/tool/{tool_name}/detail"):
+ response: Final = gateway.request("GET", path, key=key)
+ assert response.status_code == 401, (path, response.text)
+ assert string_value(owner["user_email"]) not in response.text, response.text
+
+
+def test_unauthenticated_tool_reads_are_rejected(gateway: Gateway) -> None:
+ for path in ("/v1/tool/list", "/v1/tool/some_tool", "/v1/tool/some_tool/detail"):
+ response: Final = gateway.client.get(path)
+ assert response.status_code == 401, (path, response.text)
+ assert "No api key passed in" in response.text, response.text
+
+
+def test_deleting_the_owner_keeps_the_tool_row_without_a_user(gateway: Gateway) -> None:
+ with gateway.scenario() as scenario:
+ model: Final = scenario.model()
+ user: Final = uuid.uuid4().hex
+ gateway.post("/user/new", {"user_id": user, "auto_create_key": False})
+ key: Final = string_value(gateway.post("/key/generate", {"user_id": user, "models": [model]})["key"])
+ tool_name: Final = _discover(gateway, scenario.cleanups, model, key)
+ assert _discovered_tool(gateway, tool_name)["user"] == _owner(user, None, None)
+ deleted: Final = gateway.request("POST", "/user/delete", {"user_ids": [user]})
+ assert deleted.status_code == 200, deleted.text
+ assert read_rows('SELECT token FROM "LiteLLM_VerificationToken" WHERE user_id = %s', (user,)) == []
+ tool: Final = _discovered_tool(gateway, tool_name)
+ assert tool["user"] is None, tool
+ assert tool["key_hash"] == sha256(key.encode()).hexdigest()
+ assert _single(gateway, tool_name)["user"] is None
+
+
+def test_deleting_the_key_keeps_the_tool_row_without_a_user(gateway: Gateway) -> None:
+ with gateway.scenario() as scenario:
+ model: Final = scenario.model()
+ user: Final = scenario.user()
+ key: Final = string_value(gateway.post("/key/generate", {"user_id": user, "models": [model]})["key"])
+ tool_name: Final = _discover(gateway, scenario.cleanups, model, key)
+ assert _discovered_tool(gateway, tool_name)["user"] == _owner(user, None, None)
+ scenario.delete_key(key)
+ tool: Final = _discovered_tool(gateway, tool_name)
+ assert tool["user"] is None, tool
+ assert tool["key_hash"] == sha256(key.encode()).hexdigest()
+
+
+def test_tool_row_without_a_key_hash_is_listed_without_a_user(gateway: Gateway) -> None:
+ tool_name: Final = "integration_tool_" + uuid.uuid4().hex
+ with ExitStack() as cleanups:
+ cleanups.callback(_forget_tool, tool_name)
+ write_rows(
+ 'INSERT INTO "LiteLLM_ToolTable" (tool_id, tool_name) VALUES (gen_random_uuid()::text, %s)', (tool_name,)
+ )
+ tool: Final = _discovered_tool(gateway, tool_name)
+ assert tool["key_hash"] is None, tool
+ assert tool["user"] is None, tool
+ assert _single(gateway, tool_name)["user"] is None
+
+
+def test_tool_row_with_an_unknown_key_hash_is_listed_without_a_user(gateway: Gateway) -> None:
+ tool_name: Final = "integration_tool_" + uuid.uuid4().hex
+ key_hash: Final = "integration-unknown-" + uuid.uuid4().hex
+ with ExitStack() as cleanups:
+ cleanups.callback(_forget_tool, tool_name)
+ write_rows(
+ 'INSERT INTO "LiteLLM_ToolTable" (tool_id, tool_name, key_hash) VALUES (gen_random_uuid()::text, %s, %s)',
+ (tool_name, key_hash),
+ )
+ tool: Final = _discovered_tool(gateway, tool_name)
+ assert tool["key_hash"] == key_hash, tool
+ assert tool["user"] is None, tool
+
+
+def _forget_prefixed(prefix: str) -> None:
+ write_rows('DELETE FROM "LiteLLM_ToolTable" WHERE tool_name LIKE %s', (prefix + "%",))
+ write_rows('DELETE FROM "LiteLLM_VerificationToken" WHERE token LIKE %s', (prefix + "%",))
+
+
+def test_owner_lookup_spans_more_keys_than_one_chunk(gateway: Gateway) -> None:
+ prefix: Final = "integration_chunk_" + uuid.uuid4().hex + "_"
+ count: Final = IN_LIST_CHUNK_SIZE + 1
+ with gateway.scenario() as scenario:
+ user: Final = scenario.user(user_alias="chunk-owner-" + uuid.uuid4().hex)
+ scenario.cleanups.callback(_forget_prefixed, prefix)
+ write_rows(
+ 'INSERT INTO "LiteLLM_VerificationToken" (token, user_id) '
+ "SELECT %s || g, %s FROM generate_series(1, %s::int) AS g",
+ (prefix, user, str(count)),
+ )
+ write_rows(
+ 'INSERT INTO "LiteLLM_ToolTable" (tool_id, tool_name, key_hash) '
+ "SELECT gen_random_uuid()::text, %s || g, %s || g FROM generate_series(1, %s::int) AS g",
+ (prefix, prefix, str(count)),
+ )
+ listed: Final = _listed_tools(gateway, prefix)
+ assert len(listed) == count, len(listed)
+ owners: Final = {json.dumps(tool["user"], sort_keys=True) for tool in listed}
+ assert len(owners) == 1, owners
+ assert object_value(listed[0]["user"])["user_id"] == user, listed[0]
+
+
+def test_repeated_tool_list_reads_are_identical_and_leave_rows_unchanged(gateway: Gateway) -> None:
+ with gateway.scenario() as scenario:
+ tool_name, _, _, _ = _owned_tool(gateway, scenario)
+ first: Final = _discovered_tool(gateway, tool_name)
+ before: Final = read_rows(
+ 'SELECT tool_name, key_hash, call_count, updated_at::text FROM "LiteLLM_ToolTable" WHERE tool_name = %s',
+ (tool_name,),
+ )
+ second: Final = _discovered_tool(gateway, tool_name)
+ after: Final = read_rows(
+ 'SELECT tool_name, key_hash, call_count, updated_at::text FROM "LiteLLM_ToolTable" WHERE tool_name = %s',
+ (tool_name,),
+ )
+ assert first == second, (first, second)
+ assert before == after and len(before) == 1, (before, after)
+
+
+def test_tool_list_total_matches_the_rows_in_postgres(gateway: Gateway) -> None:
+ with gateway.scenario() as scenario:
+ tool_name, _, _, _ = _owned_tool(gateway, scenario)
+ _discovered_tool(gateway, tool_name)
+ body: Final = gateway.get("/v1/tool/list")
+ tools: Final = body["tools"]
+ assert isinstance(tools, list)
+ names: Final = sorted(str(object_value(tool)["tool_name"]) for tool in tools)
+ stored: Final = sorted(
+ str(row["tool_name"]) for row in read_rows('SELECT tool_name FROM "LiteLLM_ToolTable"', ())
+ )
+ assert body["total"] == len(tools) == len(stored), body["total"]
+ assert names == stored
+
+
+def test_concurrent_tool_reads_on_two_workers_stay_consistent_during_discovery(
+ gateway: Gateway, tmp_path: Path
+) -> None:
+ model: Final = "integration-workers-" + uuid.uuid4().hex
+ config: Final = _proxy_config(tmp_path, model, gateway.upstream_url, {})
+ with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned:
+ candidate: Final = owned.gateway
+ with candidate.scenario() as scenario:
+ alias: Final = "burst-owner-" + uuid.uuid4().hex
+ user: Final = scenario.user(user_alias=alias)
+ key: Final = scenario.key(user_id=user, models=[model])
+ steady: Final = _discover(candidate, scenario.cleanups, model, key)
+ assert _discovered_tool(candidate, steady)["user"] == _owner(user, None, alias)
+ paths: Final = tuple(
+ ("/v1/tool/list", f"/v1/tool/{steady}", f"/v1/tool/{steady}/detail")[index % 3] for index in range(40)
+ )
+
+ def read(index: int) -> tuple[int, dict[str, JsonValue], str | None]:
+ burst: Final = _discover(candidate, scenario.cleanups, model, key) if index == 20 else None
+ response: Final = candidate.request("GET", paths[index])
+ assert response.status_code == 200, (paths[index], response.text)
+ return index, JSON_OBJECT.validate_json(response.content), burst
+
+ with ThreadPoolExecutor(max_workers=16) as pool:
+ results: Final = tuple(pool.map(read, range(40)))
+ for index, body, _ in results:
+ tool: Final = (
+ next(object_value(t) for t in body["tools"] if object_value(t)["tool_name"] == steady)
+ if paths[index].endswith("/list")
+ else object_value(body["tool"])
+ if paths[index].endswith("/detail")
+ else body
+ )
+ assert tool["user"] == _owner(user, None, alias), (paths[index], tool)
+ burst: Final = next(name for _, _, name in results if name)
+ assert _discovered_tool(candidate, burst)["user"] == _owner(user, None, alias)
+
+
+def test_owner_lookup_failure_keeps_tools_listed_without_a_user(gateway: Gateway, tmp_path: Path) -> None:
+ model: Final = "integration-fault-" + uuid.uuid4().hex
+ config: Final = _proxy_config(tmp_path, model, gateway.upstream_url, {})
+ with (
+ scratch_database() as database_url,
+ owned_proxy(gateway, tmp_path, {"DATABASE_URL": database_url}, config=config) as candidate,
+ ):
+ alias: Final = "fault-owner-" + uuid.uuid4().hex
+ user: Final = string_value(
+ candidate.post("/user/new", {"user_alias": alias, "auto_create_key": False})["user_id"]
+ )
+ key: Final = string_value(candidate.post("/key/generate", {"user_id": user, "models": [model]})["key"])
+ with ExitStack() as cleanups:
+ tool_name: Final = _discover(candidate, cleanups, model, key)
+ cleanups.pop_all()
+ assert _discovered_tool(candidate, tool_name)["user"] == _owner(user, None, alias)
+ write_rows('ALTER TABLE "LiteLLM_UserTable" RENAME TO "LiteLLM_UserTable_away"', (), database_url=database_url)
+ try:
+ degraded: Final = _discovered_tool(candidate, tool_name)
+ assert degraded["user"] is None, degraded
+ assert degraded["key_hash"] == sha256(key.encode()).hexdigest(), degraded
+ assert _single(candidate, tool_name)["user"] is None
+ finally:
+ write_rows(
+ 'ALTER TABLE "LiteLLM_UserTable_away" RENAME TO "LiteLLM_UserTable"', (), database_url=database_url
+ )
+ assert _discovered_tool(candidate, tool_name)["user"] == _owner(user, None, alias)
diff --git a/tests/unit/proxy/db/test_tool_registry_writer.py b/tests/unit/proxy/db/test_tool_registry_writer.py
index 6318e4422cf..c9df665741d 100644
--- a/tests/unit/proxy/db/test_tool_registry_writer.py
+++ b/tests/unit/proxy/db/test_tool_registry_writer.py
@@ -7,6 +7,7 @@ from datetime import datetime, timezone
from unittest.mock import AsyncMock, MagicMock
import pytest
+from prisma.errors import PrismaError
from litellm.proxy.db.tool_registry_writer import (
@@ -54,6 +55,8 @@ def _make_prisma(
upsert_return=None,
find_many_rows=None,
find_unique_row=None,
+ key_rows=(),
+ user_rows=(),
):
"""Return a mock prisma_client with litellm_tooltable.upsert, find_many, find_unique."""
prisma = MagicMock()
@@ -63,6 +66,10 @@ def _make_prisma(
return_value=find_many_rows if find_many_rows is not None else []
)
prisma.db.litellm_tooltable.find_unique = AsyncMock(return_value=find_unique_row)
+ prisma.db.litellm_verificationtoken = MagicMock()
+ prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=list(key_rows))
+ prisma.db.litellm_usertable = MagicMock()
+ prisma.db.litellm_usertable.find_many = AsyncMock(return_value=list(user_rows))
return prisma
@@ -133,6 +140,56 @@ async def test_list_tools_no_filter():
assert call_kw["order"] == {"created_at": "desc"}
+@pytest.mark.asyncio
+async def test_list_tools_attaches_the_owner_of_the_discovering_key():
+ owned = _mock_row(tool_id="id1", tool_name="owned_tool", key_hash="hash-owned")
+ orphan = _mock_row(tool_id="id2", tool_name="orphan_tool", key_hash="hash-orphan")
+ unknown_owner = _mock_row(tool_id="id3", tool_name="unknown_owner_tool", key_hash="hash-unknown-owner")
+ keyless = _mock_row(tool_id="id4", tool_name="keyless_tool", key_hash=None)
+ prisma = _make_prisma(
+ find_many_rows=[owned, orphan, unknown_owner, keyless],
+ key_rows=[
+ {"token": "hash-owned", "user_id": "user-1"},
+ {"token": "hash-orphan", "user_id": None},
+ {"token": "hash-unknown-owner", "user_id": "user-gone"},
+ ],
+ user_rows=[{"user_id": "user-1", "user_email": "one@example.com", "user_alias": "One"}],
+ )
+ result = await list_tools(prisma)
+ assert [tool.model_dump(include={"tool_name", "user"}) for tool in result] == [
+ {
+ "tool_name": "owned_tool",
+ "user": {"user_id": "user-1", "user_email": "one@example.com", "user_alias": "One"},
+ },
+ {"tool_name": "orphan_tool", "user": None},
+ {"tool_name": "unknown_owner_tool", "user": None},
+ {"tool_name": "keyless_tool", "user": None},
+ ]
+ key_where = prisma.db.litellm_verificationtoken.find_many.call_args.kwargs["where"]
+ assert key_where == {"token": {"in": ["hash-orphan", "hash-owned", "hash-unknown-owner"]}}
+ user_where = prisma.db.litellm_usertable.find_many.call_args.kwargs["where"]
+ assert user_where == {"user_id": {"in": ["user-1", "user-gone"]}}
+
+
+@pytest.mark.asyncio
+async def test_list_tools_keeps_tools_without_owners_when_the_owner_lookup_fails():
+ prisma = _make_prisma(find_many_rows=[_mock_row(tool_name="my_tool", key_hash="hash-owned")])
+ prisma.db.litellm_verificationtoken.find_many = AsyncMock(side_effect=PrismaError("verification token table down"))
+ result = await list_tools(prisma)
+ assert [tool.model_dump(include={"tool_name", "user"}) for tool in result] == [
+ {"tool_name": "my_tool", "user": None}
+ ]
+
+
+@pytest.mark.asyncio
+async def test_list_tools_skips_owner_lookup_when_no_tool_has_a_key_hash():
+ prisma = _make_prisma(find_many_rows=[_mock_row(key_hash=None)])
+ result = await list_tools(prisma)
+ assert [tool.user for tool in result] == [None]
+ prisma.db.litellm_verificationtoken.find_many.assert_not_awaited()
+ prisma.db.litellm_usertable.find_many.assert_not_awaited()
+
+
@pytest.mark.asyncio
async def test_list_tools_with_input_policy_filter():
row = _mock_row(
@@ -163,6 +220,24 @@ async def test_get_tool_found():
)
+@pytest.mark.asyncio
+async def test_get_tool_attaches_the_owner_of_the_discovering_key():
+ row = _mock_row(tool_name="my_tool", key_hash="hash-owned")
+ prisma = _make_prisma(
+ find_unique_row=row,
+ key_rows=[{"token": "hash-owned", "user_id": "user-1"}],
+ user_rows=[{"user_id": "user-1", "user_email": "one@example.com", "user_alias": "One"}],
+ )
+ result = await get_tool(prisma, "my_tool")
+ assert result is not None
+ assert result.model_dump(include={"tool_name", "user"}) == {
+ "tool_name": "my_tool",
+ "user": {"user_id": "user-1", "user_email": "one@example.com", "user_alias": "One"},
+ }
+ key_where = prisma.db.litellm_verificationtoken.find_many.call_args.kwargs["where"]
+ assert key_where == {"token": {"in": ["hash-owned"]}}
+
+
@pytest.mark.asyncio
async def test_get_tool_not_found():
prisma = _make_prisma(find_unique_row=None)
diff --git a/tests/unit/repositories/test_repositories.py b/tests/unit/repositories/test_repositories.py
index bd0f194b326..bae6db9ee88 100644
--- a/tests/unit/repositories/test_repositories.py
+++ b/tests/unit/repositories/test_repositories.py
@@ -196,6 +196,19 @@ class TestBaseRepository:
budgets = await repo.find_many(where={"budget_id": "b1"}, skip=0, take=10, order={"budget_id": "asc"})
assert len(budgets) == 1
+ @pytest.mark.asyncio
+ async def test_find_many_in_returns_models_from_every_chunk(self, prisma_client):
+ budget_ids: Final = tuple(f"b{i}" for i in range(IN_LIST_CHUNK_SIZE + 1))
+
+ async def find_many(where: dict[str, Any]) -> list[MockRecord]:
+ return [MockRecord({"budget_id": budget_id, "max_budget": 1.0}) for budget_id in where["budget_id"]["in"]]
+
+ prisma_client.db.litellm_budgettable.find_many = AsyncMock(side_effect=find_many)
+ budgets = await BudgetRepository(prisma_client).find_many_in("budget_id", budget_ids)
+ assert [budget.budget_id for budget in budgets] == list(budget_ids)
+ assert all(isinstance(budget, LiteLLM_BudgetTable) for budget in budgets)
+ assert prisma_client.db.litellm_budgettable.find_many.await_count == 2
+
def test_record_to_dict_branches(self):
from litellm.repositories.base_repository import record_to_dict
diff --git a/ui/litellm-dashboard/src/components/ToolPolicies/ToolPoliciesTableColumns.test.tsx b/ui/litellm-dashboard/src/components/ToolPolicies/ToolPoliciesTableColumns.test.tsx
index bb4cd8a463a..4cb6da90cd1 100644
--- a/ui/litellm-dashboard/src/components/ToolPolicies/ToolPoliciesTableColumns.test.tsx
+++ b/ui/litellm-dashboard/src/components/ToolPolicies/ToolPoliciesTableColumns.test.tsx
@@ -1,10 +1,12 @@
-import { render, screen } from "@testing-library/react";
+import { render, screen, within } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { describe, it, expect, vi } from "vitest";
import { flexRender, getCoreRowModel, useReactTable, type ColumnDef } from "@tanstack/react-table";
import { getToolPoliciesTableColumns } from "./ToolPoliciesTableColumns";
import type { ToolRow } from "@/components/networking";
+vi.mock("next/navigation", () => ({ useRouter: () => ({ push: vi.fn() }) }));
+
const row: ToolRow = {
tool_name: "search_docs",
input_policy: "untrusted",
@@ -60,10 +62,38 @@ describe("getToolPoliciesTableColumns", () => {
"team_id",
"key_hash",
"key_alias",
+ "user",
"user_agent",
]);
});
+ it("shows the owning user's alias, linking to their detail page", () => {
+ renderTable({}, [{ ...row, user: { user_id: "user-1", user_email: "one@example.com", user_alias: "Team One" } }]);
+
+ const link = screen.getByRole("link", { name: "Team One" });
+ expect(link).toHaveAttribute("href", expect.stringContaining("user-1"));
+ expect(screen.queryByText("one@example.com")).not.toBeInTheDocument();
+ });
+
+ it("falls back to the owning user's email, then id, when no alias is set", () => {
+ renderTable({}, [
+ { ...row, tool_name: "by_email", user: { user_id: "user-1", user_email: "one@example.com", user_alias: null } },
+ { ...row, tool_name: "by_id", user: { user_id: "user-2", user_email: null, user_alias: null } },
+ ]);
+
+ expect(screen.getByRole("link", { name: "one@example.com" })).toBeInTheDocument();
+ expect(screen.getByRole("link", { name: "user-2" })).toBeInTheDocument();
+ });
+
+ it("renders a dash without a link when the discovering key has no owner", () => {
+ renderTable({}, [{ ...row, user: null }]);
+
+ const userIndex = getToolPoliciesTableColumns(defaultDeps).findIndex((c) => c.id === "user");
+ const userCell = screen.getAllByRole("cell")[userIndex];
+ expect(userCell).toHaveTextContent("-");
+ expect(within(userCell).queryByRole("link")).not.toBeInTheDocument();
+ });
+
it("renders the row's identifying fields", () => {
renderTable();
diff --git a/ui/litellm-dashboard/src/components/ToolPolicies/ToolPoliciesTableColumns.tsx b/ui/litellm-dashboard/src/components/ToolPolicies/ToolPoliciesTableColumns.tsx
index d3d822759e2..5f38c58b1b4 100644
--- a/ui/litellm-dashboard/src/components/ToolPolicies/ToolPoliciesTableColumns.tsx
+++ b/ui/litellm-dashboard/src/components/ToolPolicies/ToolPoliciesTableColumns.tsx
@@ -4,7 +4,7 @@ import { ColumnDef } from "@tanstack/react-table";
import { ToolRow } from "@/components/networking";
import { DataTableSortHeader } from "@/components/shared/DataTable";
-import { DateCell, IdCell, IdentityCell } from "@/components/shared/table_cells";
+import { DateCell, IdCell, IdentityCell, UserPopoverCell } from "@/components/shared/table_cells";
import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip";
import { PolicySelect } from "./PolicySelect";
@@ -127,6 +127,32 @@ export const getToolPoliciesTableColumns = ({
meta: { title: "Key Name" },
cell: ({ row }) => ,
},
+ {
+ id: "user",
+ accessorFn: (row) => row.user?.user_alias ?? row.user?.user_email ?? row.user?.user_id ?? "",
+ header: () => (
+
+
+ User} />
+
+ The user who owns the key that discovered this tool. Displays the first available value: User Alias, User
+ Email, or User ID.
+
+
+
+ ),
+ size: 160,
+ enableSorting: false,
+ meta: { title: "User" },
+ cell: ({ row }) => (
+
+ ),
+ },
{
id: "user_agent",
accessorFn: (row) => row.user_agent ?? "",
diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx
index 129b1089e9b..d12d0a3219c 100644
--- a/ui/litellm-dashboard/src/components/networking.tsx
+++ b/ui/litellm-dashboard/src/components/networking.tsx
@@ -7614,6 +7614,11 @@ export interface ToolRow {
created_by?: string;
updated_by?: string;
user_agent?: string;
+ user?: {
+ user_id: string;
+ user_email: string | null;
+ user_alias: string | null;
+ } | null;
last_used_at?: string;
}
diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts
index fcb2805c767..ab164a61ca1 100644
--- a/ui/litellm-dashboard/src/lib/http/schema.d.ts
+++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts
@@ -34736,6 +34736,7 @@ export interface components {
updated_at?: string | null;
/** Updated By */
updated_by?: string | null;
+ user?: components["schemas"]["ToolDiscoveryUser"] | null;
/** User Agent */
user_agent?: string | null;
};
@@ -45621,6 +45622,15 @@ export interface components {
overrides?: components["schemas"]["ToolPolicyOverrideRow"][];
tool: components["schemas"]["LiteLLM_ToolTableRow"];
};
+ /** ToolDiscoveryUser */
+ ToolDiscoveryUser: {
+ /** User Alias */
+ user_alias?: string | null;
+ /** User Email */
+ user_email?: string | null;
+ /** User Id */
+ user_id: string;
+ };
/** ToolFunction */
ToolFunction: {
/** Defer Loading */