From 66318cb829a3e30666ed1fdf77392cb4cde000ae Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 25 Feb 2026 20:35:42 -0800 Subject: [PATCH 01/12] fix: fix ui render --- .../src/components/ToolPolicies.tsx | 19 ++----------------- 1 file changed, 2 insertions(+), 17 deletions(-) diff --git a/ui/litellm-dashboard/src/components/ToolPolicies.tsx b/ui/litellm-dashboard/src/components/ToolPolicies.tsx index 82785496eae..6e62370b0c6 100644 --- a/ui/litellm-dashboard/src/components/ToolPolicies.tsx +++ b/ui/litellm-dashboard/src/components/ToolPolicies.tsx @@ -2,19 +2,7 @@ import React, { useCallback, useDeferredValue, useEffect, useState } from "react"; import { Select, Switch, Tooltip } from "antd"; -<<<<<<< cursor/development-environment-setup-13a7 -// @ts-ignore - duplicate import removed -import { - Table, - TableHead, - TableHeaderCell, - TableBody, - TableRow, - TableCell, -} from "@tremor/react"; -======= import { Table, TableHead, TableHeaderCell, TableBody, TableRow, TableCell } from "@tremor/react"; ->>>>>>> main import { TimeCell } from "./view_logs/time_cell"; import { TableHeaderSortDropdown } from "./common_components/TableHeaderSortDropdown/TableHeaderSortDropdown"; import type { SortState } from "./common_components/TableHeaderSortDropdown/TableHeaderSortDropdown"; @@ -60,8 +48,7 @@ const PolicySelect: React.FC<{ minWidth: 110, fontWeight: 500, }} -<<<<<<< cursor/development-environment-setup-13a7 - {...{styles: { + styles={{ selector: { backgroundColor: style.bg, borderColor: style.border, @@ -72,9 +59,7 @@ const PolicySelect: React.FC<{ paddingLeft: 8, paddingRight: 4, }, - }} as any} -======= ->>>>>>> main + }} popupMatchSelectWidth={false} options={POLICY_OPTIONS.map((o) => ({ value: o.value, From b0439611f647e9ef2905c5afa902a2a91edb4e1f Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 25 Feb 2026 21:06:42 -0800 Subject: [PATCH 02/12] fix: fix minor bugs --- .../migration.sql | 2 + litellm/proxy/db/db_spend_update_writer.py | 155 ++++++++++-------- litellm/proxy/db/tool_registry_writer.py | 34 ++-- 3 files changed, 106 insertions(+), 85 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260225210135_baseline_diff/migration.sql diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260225210135_baseline_diff/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260225210135_baseline_diff/migration.sql new file mode 100644 index 00000000000..2f725d83806 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260225210135_baseline_diff/migration.sql @@ -0,0 +1,2 @@ +-- This is an empty migration. + diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 37fac8b56ed..9bd12c7d19a 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -13,49 +13,36 @@ import random import time import traceback from datetime import datetime, timedelta, timezone -from typing import ( - TYPE_CHECKING, - Any, - Dict, - List, - Literal, - Optional, - Union, - cast, - overload, -) +from typing import (TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, + cast, overload) import litellm from litellm._logging import verbose_proxy_logger from litellm.caching import DualCache, RedisCache from litellm.constants import DB_SPEND_UPDATE_JOB_NAME from litellm.litellm_core_utils.safe_json_loads import safe_json_loads -from litellm.proxy._types import ( - DB_CONNECTION_ERROR_TYPES, - BaseDailySpendTransaction, - DailyAgentSpendTransaction, - DailyEndUserSpendTransaction, - DailyOrganizationSpendTransaction, - DailyTagSpendTransaction, - DailyTeamSpendTransaction, - DailyUserSpendTransaction, - DBSpendUpdateTransactions, - Litellm_EntityType, - LiteLLM_UserTable, - SpendLogsMetadata, - SpendLogsPayload, - SpendUpdateQueueItem, - ToolDiscoveryQueueItem, -) -from litellm.proxy.db.db_transaction_queue.daily_spend_update_queue import ( - DailySpendUpdateQueue, -) -from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager -from litellm.proxy.db.db_transaction_queue.redis_update_buffer import RedisUpdateBuffer -from litellm.proxy.db.db_transaction_queue.spend_update_queue import SpendUpdateQueue -from litellm.proxy.db.db_transaction_queue.tool_discovery_queue import ( - ToolDiscoveryQueue, -) +from litellm.proxy._types import (DB_CONNECTION_ERROR_TYPES, + BaseDailySpendTransaction, + DailyAgentSpendTransaction, + DailyEndUserSpendTransaction, + DailyOrganizationSpendTransaction, + DailyTagSpendTransaction, + DailyTeamSpendTransaction, + DailyUserSpendTransaction, + DBSpendUpdateTransactions, + Litellm_EntityType, LiteLLM_UserTable, + SpendLogsMetadata, SpendLogsPayload, + SpendUpdateQueueItem, ToolDiscoveryQueueItem) +from litellm.proxy.db.db_transaction_queue.daily_spend_update_queue import \ + DailySpendUpdateQueue +from litellm.proxy.db.db_transaction_queue.pod_lock_manager import \ + PodLockManager +from litellm.proxy.db.db_transaction_queue.redis_update_buffer import \ + RedisUpdateBuffer +from litellm.proxy.db.db_transaction_queue.spend_update_queue import \ + SpendUpdateQueue +from litellm.proxy.db.db_transaction_queue.tool_discovery_queue import \ + ToolDiscoveryQueue from litellm.proxy.route_llm_request import ROUTE_ENDPOINT_MAPPING if TYPE_CHECKING: @@ -104,12 +91,10 @@ class DBSpendUpdateWriter: end_time: Optional[datetime], response_cost: Optional[float], ): - from litellm.proxy.proxy_server import ( - disable_spend_logs, - litellm_proxy_budget_name, - prisma_client, - user_api_key_cache, - ) + from litellm.proxy.proxy_server import (disable_spend_logs, + litellm_proxy_budget_name, + prisma_client, + user_api_key_cache) from litellm.proxy.utils import ProxyUpdateSpend, hash_token try: @@ -124,9 +109,8 @@ class DBSpendUpdateWriter: hashed_token = token ## CREATE SPEND LOG PAYLOAD ## - from litellm.proxy.spend_tracking.spend_tracking_utils import ( - get_logging_payload, - ) + from litellm.proxy.spend_tracking.spend_tracking_utils import \ + get_logging_payload payload = get_logging_payload( kwargs=kwargs, @@ -245,11 +229,13 @@ class DBSpendUpdateWriter: # --- MCP tool calls --- sl_object = kwargs.get("standard_logging_object") if sl_object is not None: - mcp_metadata = ( - sl_object.get("metadata", {}) or {} - ).get("mcp_tool_call_metadata") + mcp_metadata = (sl_object.get("metadata", {}) or {}).get( + "mcp_tool_call_metadata" + ) if mcp_metadata and isinstance(mcp_metadata, dict): - tool_name = mcp_metadata.get("namespaced_tool_name") or mcp_metadata.get("name") + tool_name = mcp_metadata.get( + "namespaced_tool_name" + ) or mcp_metadata.get("name") mcp_server_name = mcp_metadata.get("mcp_server_name") if tool_name: _enqueue(tool_name, origin=mcp_server_name or "user_defined") @@ -280,7 +266,9 @@ class DBSpendUpdateWriter: _enqueue(name) # --- Response tool_calls (OpenAI format; Anthropic pass-through converts tool_use here) --- - if completion_response is not None and hasattr(completion_response, "choices"): + if completion_response is not None and hasattr( + completion_response, "choices" + ): for choice in completion_response.choices or []: message = getattr(choice, "message", None) if message is None: @@ -312,7 +300,7 @@ class DBSpendUpdateWriter: prisma_client: Optional[PrismaClient], user_api_key_cache: DualCache, litellm_proxy_budget_name: Optional[str], - payload_copy: dict, + payload_copy: SpendLogsPayload, request_tags: Optional[Any], ): """ @@ -768,19 +756,46 @@ class DBSpendUpdateWriter: daily_end_user_spend_update_transactions, daily_agent_spend_update_transactions, daily_tag_spend_update_transactions, - ) = await self.redis_update_buffer.get_all_transactions_from_redis_buffer_pipeline() + ) = ( + await self.redis_update_buffer.get_all_transactions_from_redis_buffer_pipeline() + ) if db_spend_update_transactions is not None: verbose_proxy_logger.info( "Spend tracking - committing spend updates from Redis to DB: " "keys=%d, users=%d, teams=%d, orgs=%d, end_users=%d, team_members=%d, tags=%d", - len(db_spend_update_transactions.get("key_list_transactions") or {}), - len(db_spend_update_transactions.get("user_list_transactions") or {}), - len(db_spend_update_transactions.get("team_list_transactions") or {}), - len(db_spend_update_transactions.get("org_list_transactions") or {}), - len(db_spend_update_transactions.get("end_user_list_transactions") or {}), - len(db_spend_update_transactions.get("team_member_list_transactions") or {}), - len(db_spend_update_transactions.get("tag_list_transactions") or {}), + len( + db_spend_update_transactions.get("key_list_transactions") + or {} + ), + len( + db_spend_update_transactions.get("user_list_transactions") + or {} + ), + len( + db_spend_update_transactions.get("team_list_transactions") + or {} + ), + len( + db_spend_update_transactions.get("org_list_transactions") + or {} + ), + len( + db_spend_update_transactions.get( + "end_user_list_transactions" + ) + or {} + ), + len( + db_spend_update_transactions.get( + "team_member_list_transactions" + ) + or {} + ), + len( + db_spend_update_transactions.get("tag_list_transactions") + or {} + ), ) await self._commit_spend_updates_to_db( prisma_client=prisma_client, @@ -985,10 +1000,8 @@ class DBSpendUpdateWriter: Commits all the spend `UPDATE` transactions to the Database """ - from litellm.proxy.utils import ( - ProxyUpdateSpend, - _raise_failed_update_spend_exception, - ) + from litellm.proxy.utils import (ProxyUpdateSpend, + _raise_failed_update_spend_exception) ### UPDATE USER TABLE ### user_list_transactions = db_spend_update_transactions["user_list_transactions"] @@ -1523,14 +1536,14 @@ class DBSpendUpdateWriter: # Add cache-related fields if they exist if "cache_read_input_tokens" in transaction: - common_data[ - "cache_read_input_tokens" - ] = transaction.get("cache_read_input_tokens", 0) + common_data["cache_read_input_tokens"] = ( + transaction.get("cache_read_input_tokens", 0) + ) if "cache_creation_input_tokens" in transaction: - common_data[ - "cache_creation_input_tokens" - ] = transaction.get( - "cache_creation_input_tokens", 0 + common_data["cache_creation_input_tokens"] = ( + transaction.get( + "cache_creation_input_tokens", 0 + ) ) if entity_type == "tag" and "request_id" in transaction: diff --git a/litellm/proxy/db/tool_registry_writer.py b/litellm/proxy/db/tool_registry_writer.py index 4e0a8095a08..b12cf68f1be 100644 --- a/litellm/proxy/db/tool_registry_writer.py +++ b/litellm/proxy/db/tool_registry_writer.py @@ -66,10 +66,10 @@ async def batch_upsert_tools( await prisma_client.db.execute_raw( 'INSERT INTO "LiteLLM_ToolTable" ' "(tool_id, tool_name, origin, call_policy, call_count, created_by, updated_by, key_hash, team_id, key_alias, created_at, updated_at) " - "VALUES ($7, $1, $2, 'untrusted', 1, $3, $3, $4, $5, $6, $8, $8) " + "VALUES ($7, $1, $2, 'untrusted', 1, $3, $3, $4, $5, $6, $8::timestamp, $8::timestamp) " "ON CONFLICT (tool_name) DO UPDATE SET " - "call_count = \"LiteLLM_ToolTable\".call_count + 1, " - "updated_at = $8", + 'call_count = "LiteLLM_ToolTable".call_count + 1, ' + "updated_at = $8::timestamp", tool_name, origin, created_by, @@ -83,7 +83,9 @@ async def batch_upsert_tools( "tool_registry_writer: upserted %d tool(s)", len(data) ) except Exception as e: - verbose_proxy_logger.error("tool_registry_writer batch_upsert_tools error: %s", e) + verbose_proxy_logger.error( + "tool_registry_writer batch_upsert_tools error: %s", e + ) async def list_tools( @@ -94,15 +96,15 @@ async def list_tools( try: if call_policy is not None: rows = await prisma_client.db.query_raw( - 'SELECT tool_id, tool_name, origin, call_policy, call_count, assignments, ' - 'key_hash, team_id, key_alias, created_at, updated_at, created_by, updated_by ' + "SELECT tool_id, tool_name, origin, call_policy, call_count, assignments, " + "key_hash, team_id, key_alias, created_at, updated_at, created_by, updated_by " 'FROM "LiteLLM_ToolTable" WHERE call_policy = $1 ORDER BY created_at DESC', call_policy, ) else: rows = await prisma_client.db.query_raw( - 'SELECT tool_id, tool_name, origin, call_policy, call_count, assignments, ' - 'key_hash, team_id, key_alias, created_at, updated_at, created_by, updated_by ' + "SELECT tool_id, tool_name, origin, call_policy, call_count, assignments, " + "key_hash, team_id, key_alias, created_at, updated_at, created_by, updated_by " 'FROM "LiteLLM_ToolTable" ORDER BY created_at DESC', ) return [_row_to_model(row) for row in rows] @@ -118,8 +120,8 @@ async def get_tool( """Return a single tool row by tool_name.""" try: rows = await prisma_client.db.query_raw( - 'SELECT tool_id, tool_name, origin, call_policy, call_count, assignments, ' - 'key_hash, team_id, key_alias, created_at, updated_at, created_by, updated_by ' + "SELECT tool_id, tool_name, origin, call_policy, call_count, assignments, " + "key_hash, team_id, key_alias, created_at, updated_at, created_by, updated_by " 'FROM "LiteLLM_ToolTable" WHERE tool_name = $1', tool_name, ) @@ -143,8 +145,8 @@ async def update_tool_policy( now = datetime.now(timezone.utc).isoformat() await prisma_client.db.execute_raw( 'INSERT INTO "LiteLLM_ToolTable" (tool_id, tool_name, call_policy, created_by, updated_by, created_at, updated_at) ' - "VALUES ($4, $1, $2, $3, $3, $5, $5) " - "ON CONFLICT (tool_name) DO UPDATE SET call_policy = $2, updated_by = $3, updated_at = $5", + "VALUES ($4, $1, $2, $3, $3, $5::timestamp, $5::timestamp) " + "ON CONFLICT (tool_name) DO UPDATE SET call_policy = $2, updated_by = $3, updated_at = $5::timestamp", tool_name, call_policy, _updated_by, @@ -153,7 +155,9 @@ async def update_tool_policy( ) return await get_tool(prisma_client, tool_name) except Exception as e: - verbose_proxy_logger.error("tool_registry_writer update_tool_policy error: %s", e) + verbose_proxy_logger.error( + "tool_registry_writer update_tool_policy error: %s", e + ) return None @@ -175,5 +179,7 @@ async def get_tools_by_names( ) return {row["tool_name"]: row["call_policy"] for row in rows} except Exception as e: - verbose_proxy_logger.error("tool_registry_writer get_tools_by_names error: %s", e) + verbose_proxy_logger.error( + "tool_registry_writer get_tools_by_names error: %s", e + ) return {} From 4b4018b0b2552a34c820f991b569bbb4016bd125 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 25 Feb 2026 21:13:33 -0800 Subject: [PATCH 03/12] refactor: use prisma functions instead of raw sql (safer) --- AGENTS.md | 3 + CLAUDE.md | 4 + .../migration.sql | 2 + litellm/proxy/db/tool_registry_writer.py | 137 +++++++------ .../proxy/db/test_tool_registry_writer.py | 191 ++++++++++-------- 5 files changed, 190 insertions(+), 147 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260225211151_baseline_diff/migration.sql diff --git a/AGENTS.md b/AGENTS.md index 96776a3fae7..fe0eb15ae4b 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -109,6 +109,8 @@ Key files: - `litellm/proxy/auth/` - Authentication logic - `litellm/proxy/management_endpoints/` - Admin API endpoints +**Database (proxy)**: Use Prisma model methods (`prisma_client.db..upsert`, `.find_many`, `.find_unique`, etc.), not raw SQL (`execute_raw`/`query_raw`). See COMMON PITFALLS for details. + ## MCP (MODEL CONTEXT PROTOCOL) SUPPORT LiteLLM supports MCP for agent workflows: @@ -176,6 +178,7 @@ When opening issues or pull requests, follow these templates: 5. **Dependencies**: Keep dependencies minimal and well-justified 6. **UI/Backend Contract Mismatch**: When adding a new entity type to the UI, always check whether the backend endpoint accepts a single value or an array. Match the UI control accordingly (single-select vs. multi-select) to avoid silently dropping user selections 7. **Missing Tests for New Entity Types**: When adding a new entity type (e.g., in `EntityUsage`, `UsageViewSelect`), always add corresponding tests in the existing test files and update any icon/component mocks +8. **Raw SQL in proxy DB code**: Do not use `execute_raw` or `query_raw` for proxy database access. Use Prisma model methods (e.g. `prisma_client.db.litellm_tooltable.upsert()`, `.find_many()`, `.find_unique()`) so behavior stays consistent with the schema, the client stays mockable in tests, and you avoid the pitfalls of hand-written SQL (parameter ordering, type casting, schema drift) ## HELPFUL RESOURCES diff --git a/CLAUDE.md b/CLAUDE.md index 3b597fb8a90..c1eb75d2515 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -107,6 +107,10 @@ LiteLLM is a unified interface for 100+ LLM providers with two main components: - Migration files auto-generated with `prisma migrate dev` - Always test migrations against both PostgreSQL and SQLite +### Proxy database access +- **Do not write raw SQL** for proxy DB operations. Use Prisma model methods instead of `execute_raw` / `query_raw`. +- Use the generated client: `prisma_client.db.` (e.g. `litellm_tooltable`, `litellm_usertable`) with `.upsert()`, `.find_many()`, `.find_unique()`, `.update()`, `.update_many()` as appropriate. This avoids schema/client drift, keeps code testable with simple mocks, and matches patterns used in spend logs and other proxy code. + ### Enterprise Features - Enterprise-specific code in `enterprise/` directory - Optional features enabled via environment variables diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260225211151_baseline_diff/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260225211151_baseline_diff/migration.sql new file mode 100644 index 00000000000..2f725d83806 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260225211151_baseline_diff/migration.sql @@ -0,0 +1,2 @@ +-- This is an empty migration. + diff --git a/litellm/proxy/db/tool_registry_writer.py b/litellm/proxy/db/tool_registry_writer.py index b12cf68f1be..332607f0305 100644 --- a/litellm/proxy/db/tool_registry_writer.py +++ b/litellm/proxy/db/tool_registry_writer.py @@ -3,15 +3,11 @@ DB helpers for LiteLLM_ToolTable — the global tool registry. Tools are auto-discovered from LLM responses and upserted here. Admins use the management endpoints to read and update call_policy. - -NOTE: Uses raw SQL (query_raw / execute_raw) instead of Prisma model methods -because the generated Prisma Python client may not have LiteLLM_ToolTable -when running against an older generated schema. """ import uuid from datetime import datetime, timezone -from typing import TYPE_CHECKING, Dict, List, Optional +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union from litellm._logging import verbose_proxy_logger from litellm.proxy._types import ToolDiscoveryQueueItem @@ -21,7 +17,30 @@ if TYPE_CHECKING: from litellm.proxy.utils import PrismaClient -def _row_to_model(row: dict) -> LiteLLM_ToolTableRow: +def _row_to_model(row: Union[dict, Any]) -> LiteLLM_ToolTableRow: + """Convert a Prisma model instance or dict to LiteLLM_ToolTableRow.""" + model_dump = getattr(row, "model_dump", None) + if callable(model_dump): + row = model_dump() + elif not isinstance(row, dict): + row = { + k: getattr(row, k, None) + for k in ( + "tool_id", + "tool_name", + "origin", + "call_policy", + "call_count", + "assignments", + "key_hash", + "team_id", + "key_alias", + "created_at", + "updated_at", + "created_by", + "updated_by", + ) + } return LiteLLM_ToolTableRow( tool_id=row.get("tool_id", ""), tool_name=row.get("tool_name", ""), @@ -44,7 +63,7 @@ async def batch_upsert_tools( items: List[ToolDiscoveryQueueItem], ) -> None: """ - Batch-upsert tool registry rows via raw SQL. + Batch-upsert tool registry rows via Prisma. On first insert: sets call_policy = "untrusted" (schema default), call_count = 1. On conflict: increments call_count; preserves existing call_policy. @@ -55,6 +74,8 @@ async def batch_upsert_tools( data = [item for item in items if item.get("tool_name")] if not data: return + now = datetime.now(timezone.utc) + table = prisma_client.db.litellm_tooltable for item in data: tool_name = item.get("tool_name", "") origin = item.get("origin") or "user_defined" @@ -62,22 +83,26 @@ async def batch_upsert_tools( key_hash = item.get("key_hash") team_id = item.get("team_id") key_alias = item.get("key_alias") - now = datetime.now(timezone.utc).isoformat() - await prisma_client.db.execute_raw( - 'INSERT INTO "LiteLLM_ToolTable" ' - "(tool_id, tool_name, origin, call_policy, call_count, created_by, updated_by, key_hash, team_id, key_alias, created_at, updated_at) " - "VALUES ($7, $1, $2, 'untrusted', 1, $3, $3, $4, $5, $6, $8::timestamp, $8::timestamp) " - "ON CONFLICT (tool_name) DO UPDATE SET " - 'call_count = "LiteLLM_ToolTable".call_count + 1, ' - "updated_at = $8::timestamp", - tool_name, - origin, - created_by, - key_hash, - team_id, - key_alias, - str(uuid.uuid4()), - now, + await table.upsert( + where={"tool_name": tool_name}, + data={ + "create": { + "tool_id": str(uuid.uuid4()), + "tool_name": tool_name, + "origin": origin, + "call_policy": "untrusted", + "call_count": 1, + "created_by": created_by, + "updated_by": created_by, + "key_hash": key_hash, + "team_id": team_id, + "key_alias": key_alias, + }, + "update": { + "call_count": {"increment": 1}, + "updated_at": now, + }, + }, ) verbose_proxy_logger.debug( "tool_registry_writer: upserted %d tool(s)", len(data) @@ -94,19 +119,11 @@ async def list_tools( ) -> List[LiteLLM_ToolTableRow]: """Return all tools, optionally filtered by call_policy.""" try: - if call_policy is not None: - rows = await prisma_client.db.query_raw( - "SELECT tool_id, tool_name, origin, call_policy, call_count, assignments, " - "key_hash, team_id, key_alias, created_at, updated_at, created_by, updated_by " - 'FROM "LiteLLM_ToolTable" WHERE call_policy = $1 ORDER BY created_at DESC', - call_policy, - ) - else: - rows = await prisma_client.db.query_raw( - "SELECT tool_id, tool_name, origin, call_policy, call_count, assignments, " - "key_hash, team_id, key_alias, created_at, updated_at, created_by, updated_by " - 'FROM "LiteLLM_ToolTable" ORDER BY created_at DESC', - ) + where = {"call_policy": call_policy} if call_policy is not None else {} + rows = await prisma_client.db.litellm_tooltable.find_many( + where=where, + order={"created_at": "desc"}, + ) return [_row_to_model(row) for row in rows] except Exception as e: verbose_proxy_logger.error("tool_registry_writer list_tools error: %s", e) @@ -119,15 +136,12 @@ async def get_tool( ) -> Optional[LiteLLM_ToolTableRow]: """Return a single tool row by tool_name.""" try: - rows = await prisma_client.db.query_raw( - "SELECT tool_id, tool_name, origin, call_policy, call_count, assignments, " - "key_hash, team_id, key_alias, created_at, updated_at, created_by, updated_by " - 'FROM "LiteLLM_ToolTable" WHERE tool_name = $1', - tool_name, + row = await prisma_client.db.litellm_tooltable.find_unique( + where={"tool_name": tool_name}, ) - if not rows: + if row is None: return None - return _row_to_model(rows[0]) + return _row_to_model(row) except Exception as e: verbose_proxy_logger.error("tool_registry_writer get_tool error: %s", e) return None @@ -142,16 +156,25 @@ async def update_tool_policy( """Update the call_policy for a tool. Upserts the row if it does not exist yet.""" try: _updated_by = updated_by or "system" - now = datetime.now(timezone.utc).isoformat() - await prisma_client.db.execute_raw( - 'INSERT INTO "LiteLLM_ToolTable" (tool_id, tool_name, call_policy, created_by, updated_by, created_at, updated_at) ' - "VALUES ($4, $1, $2, $3, $3, $5::timestamp, $5::timestamp) " - "ON CONFLICT (tool_name) DO UPDATE SET call_policy = $2, updated_by = $3, updated_at = $5::timestamp", - tool_name, - call_policy, - _updated_by, - str(uuid.uuid4()), - now, + now = datetime.now(timezone.utc) + await prisma_client.db.litellm_tooltable.upsert( + where={"tool_name": tool_name}, + data={ + "create": { + "tool_id": str(uuid.uuid4()), + "tool_name": tool_name, + "call_policy": call_policy, + "created_by": _updated_by, + "updated_by": _updated_by, + "created_at": now, + "updated_at": now, + }, + "update": { + "call_policy": call_policy, + "updated_by": _updated_by, + "updated_at": now, + }, + }, ) return await get_tool(prisma_client, tool_name) except Exception as e: @@ -172,12 +195,10 @@ async def get_tools_by_names( if not tool_names: return {} try: - placeholders = ", ".join(f"${i+1}" for i in range(len(tool_names))) - rows = await prisma_client.db.query_raw( - f'SELECT tool_name, call_policy FROM "LiteLLM_ToolTable" WHERE tool_name IN ({placeholders})', - *tool_names, + rows = await prisma_client.db.litellm_tooltable.find_many( + where={"tool_name": {"in": tool_names}}, ) - return {row["tool_name"]: row["call_policy"] for row in rows} + return {row.tool_name: row.call_policy for row in rows} except Exception as e: verbose_proxy_logger.error( "tool_registry_writer get_tools_by_names error: %s", e diff --git a/tests/test_litellm/proxy/db/test_tool_registry_writer.py b/tests/test_litellm/proxy/db/test_tool_registry_writer.py index 44f9e32058a..f1b829b483f 100644 --- a/tests/test_litellm/proxy/db/test_tool_registry_writer.py +++ b/tests/test_litellm/proxy/db/test_tool_registry_writer.py @@ -1,6 +1,6 @@ """ Unit tests for tool_registry_writer.py — uses a mock prisma client -that exposes execute_raw / query_raw (matching the actual raw-SQL implementation). +that exposes litellm_tooltable.upsert / find_many / find_unique. """ import os @@ -12,18 +12,20 @@ import pytest sys.path.insert(0, os.path.abspath("../../..")) -from litellm.proxy.db.tool_registry_writer import ( - batch_upsert_tools, - get_tool, - get_tools_by_names, - list_tools, - update_tool_policy, -) +from litellm.proxy.db.tool_registry_writer import (batch_upsert_tools, + get_tool, + get_tools_by_names, + list_tools, + update_tool_policy) -def _make_prisma(query_rows=None): - """Return a minimal mock prisma_client with execute_raw / query_raw.""" - default_row = { +def _mock_row(**kwargs): + """Build a row-like object with real attributes (no MagicMock) for _row_to_model.""" + + class Row: + pass + + default = { "tool_id": "uuid-1", "tool_name": "my_tool", "origin": "user_defined", @@ -38,31 +40,53 @@ def _make_prisma(query_rows=None): "created_by": None, "updated_by": None, } - rows = query_rows if query_rows is not None else [default_row] + default.update(kwargs) + row = Row() + for k, v in default.items(): + setattr(row, k, v) + return row + +def _make_prisma( + *, + upsert_return=None, + find_many_rows=None, + find_unique_row=None, +): + """Return a mock prisma_client with litellm_tooltable.upsert, find_many, find_unique.""" prisma = MagicMock() - prisma.db.execute_raw = AsyncMock(return_value=None) - prisma.db.query_raw = AsyncMock(return_value=rows) + prisma.db.litellm_tooltable = MagicMock() + prisma.db.litellm_tooltable.upsert = AsyncMock(return_value=upsert_return) + prisma.db.litellm_tooltable.find_many = AsyncMock( + 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 + ) return prisma @pytest.mark.asyncio -async def test_batch_upsert_tools_calls_execute_raw(): +async def test_batch_upsert_tools_calls_upsert(): prisma = _make_prisma() items = [{"tool_name": "tool_a", "origin": "mcp_server", "created_by": None}] await batch_upsert_tools(prisma, items) - prisma.db.execute_raw.assert_awaited_once() - call_args = prisma.db.execute_raw.call_args - sql = call_args.args[0] - assert "LiteLLM_ToolTable" in sql - assert "ON CONFLICT" in sql + prisma.db.litellm_tooltable.upsert.assert_awaited_once() + call_kw = prisma.db.litellm_tooltable.upsert.call_args.kwargs + assert call_kw["where"] == {"tool_name": "tool_a"} + assert call_kw["data"]["create"]["tool_name"] == "tool_a" + assert call_kw["data"]["create"]["origin"] == "mcp_server" + assert call_kw["data"]["create"]["call_policy"] == "untrusted" + assert call_kw["data"]["create"]["call_count"] == 1 + assert call_kw["data"]["update"]["call_count"] == {"increment": 1} + assert "updated_at" in call_kw["data"]["update"] @pytest.mark.asyncio async def test_batch_upsert_tools_empty_list(): prisma = _make_prisma() await batch_upsert_tools(prisma, []) - prisma.db.execute_raw.assert_not_awaited() + prisma.db.litellm_tooltable.upsert.assert_not_awaited() @pytest.mark.asyncio @@ -70,123 +94,112 @@ async def test_batch_upsert_tools_skips_empty_names(): prisma = _make_prisma() items = [{"tool_name": "", "origin": None}, {"tool_name": None}] # type: ignore[list-item] await batch_upsert_tools(prisma, items) - prisma.db.execute_raw.assert_not_awaited() + prisma.db.litellm_tooltable.upsert.assert_not_awaited() @pytest.mark.asyncio -async def test_batch_upsert_multiple_tools_calls_execute_raw_per_tool(): +async def test_batch_upsert_multiple_tools_calls_upsert_per_tool(): prisma = _make_prisma() items = [ {"tool_name": "tool_a", "origin": "mcp_server", "created_by": None}, {"tool_name": "tool_b", "origin": "user_defined", "created_by": "alice"}, ] await batch_upsert_tools(prisma, items) - assert prisma.db.execute_raw.await_count == 2 + assert prisma.db.litellm_tooltable.upsert.await_count == 2 + calls = prisma.db.litellm_tooltable.upsert.call_args_list + assert calls[0].kwargs["where"]["tool_name"] == "tool_a" + assert calls[1].kwargs["where"]["tool_name"] == "tool_b" @pytest.mark.asyncio async def test_list_tools_no_filter(): - row = { - "tool_id": "id1", - "tool_name": "tool_a", - "origin": "mcp", - "call_policy": "untrusted", - "call_count": 5, - "assignments": {}, - "key_hash": None, - "team_id": None, - "key_alias": None, - "created_at": datetime.now(timezone.utc), - "updated_at": datetime.now(timezone.utc), - "created_by": None, - "updated_by": None, - } - prisma = _make_prisma(query_rows=[row]) + row = _mock_row( + tool_id="id1", + tool_name="tool_a", + origin="mcp", + call_policy="untrusted", + call_count=5, + ) + prisma = _make_prisma(find_many_rows=[row]) result = await list_tools(prisma) assert len(result) == 1 assert result[0].tool_name == "tool_a" assert result[0].call_count == 5 - prisma.db.query_raw.assert_awaited_once() + prisma.db.litellm_tooltable.find_many.assert_awaited_once() + call_kw = prisma.db.litellm_tooltable.find_many.call_args.kwargs + assert call_kw["where"] == {} + assert call_kw["order"] == {"created_at": "desc"} @pytest.mark.asyncio async def test_list_tools_with_policy_filter(): - row = { - "tool_id": "id1", - "tool_name": "blocked_tool", - "origin": None, - "call_policy": "blocked", - "call_count": 2, - "assignments": None, - "key_hash": None, - "team_id": None, - "key_alias": None, - "created_at": datetime.now(timezone.utc), - "updated_at": datetime.now(timezone.utc), - "created_by": None, - "updated_by": None, - } - prisma = _make_prisma(query_rows=[row]) + row = _mock_row( + tool_id="id1", + tool_name="blocked_tool", + origin=None, + call_policy="blocked", + call_count=2, + assignments=None, + ) + prisma = _make_prisma(find_many_rows=[row]) result = await list_tools(prisma, call_policy="blocked") assert result[0].call_policy == "blocked" - call_args = prisma.db.query_raw.call_args - sql = call_args.args[0] - assert "WHERE call_policy" in sql + call_kw = prisma.db.litellm_tooltable.find_many.call_args.kwargs + assert call_kw["where"] == {"call_policy": "blocked"} @pytest.mark.asyncio async def test_get_tool_found(): - prisma = _make_prisma() + row = _mock_row(tool_name="my_tool") + prisma = _make_prisma(find_unique_row=row) result = await get_tool(prisma, "my_tool") assert result is not None assert result.tool_name == "my_tool" - prisma.db.query_raw.assert_awaited_once() + prisma.db.litellm_tooltable.find_unique.assert_awaited_once_with( + where={"tool_name": "my_tool"} + ) @pytest.mark.asyncio async def test_get_tool_not_found(): - prisma = _make_prisma(query_rows=[]) + prisma = _make_prisma(find_unique_row=None) result = await get_tool(prisma, "nonexistent") assert result is None @pytest.mark.asyncio -async def test_update_tool_policy_calls_execute_raw(): - row = { - "tool_id": "uuid-1", - "tool_name": "my_tool", - "origin": "user_defined", - "call_policy": "blocked", - "call_count": 1, - "assignments": {}, - "key_hash": None, - "team_id": None, - "key_alias": None, - "created_at": datetime.now(timezone.utc), - "updated_at": datetime.now(timezone.utc), - "created_by": None, - "updated_by": "admin", - } - prisma = _make_prisma(query_rows=[row]) +async def test_update_tool_policy_calls_upsert_then_get_tool(): + row = _mock_row( + tool_name="my_tool", + call_policy="blocked", + updated_by="admin", + ) + prisma = _make_prisma(find_unique_row=row) result = await update_tool_policy(prisma, "my_tool", "blocked", "admin") assert result is not None assert result.call_policy == "blocked" - prisma.db.execute_raw.assert_awaited_once() - call_args = prisma.db.execute_raw.call_args - sql = call_args.args[0] - assert "ON CONFLICT" in sql - assert "call_policy" in sql + prisma.db.litellm_tooltable.upsert.assert_awaited_once() + call_kw = prisma.db.litellm_tooltable.upsert.call_args.kwargs + assert call_kw["where"] == {"tool_name": "my_tool"} + assert call_kw["data"]["update"]["call_policy"] == "blocked" + assert call_kw["data"]["update"]["updated_by"] == "admin" + prisma.db.litellm_tooltable.find_unique.assert_awaited_with( + where={"tool_name": "my_tool"} + ) @pytest.mark.asyncio async def test_get_tools_by_names_returns_policy_map(): rows = [ - {"tool_name": "tool_a", "call_policy": "trusted"}, - {"tool_name": "tool_b", "call_policy": "blocked"}, + _mock_row(tool_name="tool_a", call_policy="trusted"), + _mock_row(tool_name="tool_b", call_policy="blocked"), ] - prisma = _make_prisma(query_rows=rows) + prisma = _make_prisma(find_many_rows=rows) result = await get_tools_by_names(prisma, ["tool_a", "tool_b"]) assert result == {"tool_a": "trusted", "tool_b": "blocked"} + prisma.db.litellm_tooltable.find_many.assert_awaited_once_with( + where={"tool_name": {"in": ["tool_a", "tool_b"]}} + ) @pytest.mark.asyncio @@ -194,4 +207,4 @@ async def test_get_tools_by_names_empty_list(): prisma = _make_prisma() result = await get_tools_by_names(prisma, []) assert result == {} - prisma.db.query_raw.assert_not_awaited() + prisma.db.litellm_tooltable.find_many.assert_not_awaited() From 2487943846b5c2e8cd7cfd5a53df9da6ca9468d8 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 25 Feb 2026 21:53:42 -0800 Subject: [PATCH 04/12] fix(add-new-tiles-to-tool-policies): allow developer to see what's available --- .../proxy/test_tools_allowlist_enforcement.py | 548 ++++++++++++++++++ .../src/components/ToolPolicies.tsx | 133 ++++- 2 files changed, 679 insertions(+), 2 deletions(-) create mode 100644 tests/test_litellm/proxy/test_tools_allowlist_enforcement.py diff --git a/tests/test_litellm/proxy/test_tools_allowlist_enforcement.py b/tests/test_litellm/proxy/test_tools_allowlist_enforcement.py new file mode 100644 index 00000000000..09d56be30ba --- /dev/null +++ b/tests/test_litellm/proxy/test_tools_allowlist_enforcement.py @@ -0,0 +1,548 @@ +""" +Tests for tool allowlist enforcement by team/key (metadata.allowed_tools). + +No implementation yet; these tests define expected behavior. When check_tools_allowlist +is implemented in common_checks, disallowed-tool tests should raise; allowed and +no-allowlist tests should pass. +""" + +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from litellm.proxy._types import ProxyException, UserAPIKeyAuth +from litellm.proxy.auth.auth_checks import common_checks + + +class MockRequest: + """Mock request with method attribute.""" + + def __init__(self, method: str = "POST"): + self.method = method + + +def get_mock_user_token(metadata=None, team_metadata=None) -> UserAPIKeyAuth: + """Build UserAPIKeyAuth with optional metadata and team_metadata for allowlist.""" + kwargs = { + "api_key": "test-key", + "user_id": "test-user", + "team_id": "test-team", + "org_id": "test-org", + "models": ["*"], + "metadata": metadata or {}, + } + if team_metadata is not None: + kwargs["team_metadata"] = team_metadata + return UserAPIKeyAuth(**kwargs) + + +def _tools_allowlist_patches(): + """Patches so only tool-allowlist behavior is under test; heavy/DB parts no-op.""" + p1 = patch( + "litellm.proxy.auth.auth_checks._is_api_route_allowed", + new_callable=AsyncMock, + return_value=True, + ) + p2 = patch( + "litellm.proxy.auth.auth_checks.vector_store_access_check", + new_callable=AsyncMock, + return_value=None, + ) + p3 = patch( + "litellm.proxy.auth.auth_checks._run_project_checks", + new_callable=AsyncMock, + return_value=None, + ) + return p1, p2, p3 + + +class TestOpenAIChatCompletionsToolsAllowlist: + """Tool allowlist enforcement for /v1/chat/completions.""" + + @pytest.mark.asyncio + async def test_chat_completions_allowed_tool_passes(self): + """Request with tools in allowed_tools passes.""" + route = "/v1/chat/completions" + request_body = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hi"}], + "tools": [{"type": "function", "function": {"name": "get_weather"}}], + } + token = get_mock_user_token(metadata={"allowed_tools": ["get_weather"]}) + request = MockRequest("POST") + + p1, p2, p3 = _tools_allowlist_patches() + with p1, p2, p3: + result = await common_checks( + request_body=request_body, + team_object=None, + user_object=None, + end_user_object=None, + global_proxy_spend=None, + general_settings={}, + route=route, + llm_router=None, + proxy_logging_obj=MagicMock(), + valid_token=token, + request=request, + ) + assert result is True + + @pytest.mark.asyncio + async def test_chat_completions_disallowed_tool_raises(self): + """Request with tool not in allowed_tools raises.""" + route = "/v1/chat/completions" + request_body = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hi"}], + "tools": [{"type": "function", "function": {"name": "get_weather"}}], + } + token = get_mock_user_token(metadata={"allowed_tools": ["other_tool"]}) + request = MockRequest("POST") + + p1, p2, p3 = _tools_allowlist_patches() + with p1, p2, p3: + with pytest.raises((Exception, ProxyException)) as exc_info: + await common_checks( + request_body=request_body, + team_object=None, + user_object=None, + end_user_object=None, + global_proxy_spend=None, + general_settings={}, + route=route, + llm_router=None, + proxy_logging_obj=MagicMock(), + valid_token=token, + request=request, + ) + msg = str(exc_info.value).lower() + assert "tool" in msg or "allowed" in msg + + @pytest.mark.asyncio + async def test_chat_completions_legacy_functions_allowed(self): + """Legacy 'functions' (no tools) with allowed name passes.""" + route = "/v1/chat/completions" + request_body = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hi"}], + "functions": [{"name": "get_weather"}], + } + token = get_mock_user_token(metadata={"allowed_tools": ["get_weather"]}) + request = MockRequest("POST") + + p1, p2, p3 = _tools_allowlist_patches() + with p1, p2, p3: + result = await common_checks( + request_body=request_body, + team_object=None, + user_object=None, + end_user_object=None, + global_proxy_spend=None, + general_settings={}, + route=route, + llm_router=None, + proxy_logging_obj=MagicMock(), + valid_token=token, + request=request, + ) + assert result is True + + @pytest.mark.asyncio + async def test_chat_completions_no_allowlist_passes(self): + """Request with tools but no metadata.allowed_tools / team_metadata passes.""" + route = "/v1/chat/completions" + request_body = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hi"}], + "tools": [{"type": "function", "function": {"name": "get_weather"}}], + } + token = get_mock_user_token(metadata={}, team_metadata={}) + request = MockRequest("POST") + + p1, p2, p3 = _tools_allowlist_patches() + with p1, p2, p3: + result = await common_checks( + request_body=request_body, + team_object=None, + user_object=None, + end_user_object=None, + global_proxy_spend=None, + general_settings={}, + route=route, + llm_router=None, + proxy_logging_obj=MagicMock(), + valid_token=token, + request=request, + ) + assert result is True + + +class TestOpenAIResponsesAPIToolsAllowlist: + """Tool allowlist enforcement for /v1/responses.""" + + @pytest.mark.asyncio + async def test_responses_function_tool_allowed_passes(self): + """Responses request with function tool in allowed_tools passes.""" + route = "/v1/responses" + request_body = { + "model": "gpt-4", + "input": "What is the weather?", + "tools": [ + { + "type": "function", + "name": "get_current_weather", + "description": "Get current weather", + "parameters": {"type": "object"}, + } + ], + } + token = get_mock_user_token(metadata={"allowed_tools": ["get_current_weather"]}) + request = MockRequest("POST") + + p1, p2, p3 = _tools_allowlist_patches() + with p1, p2, p3: + result = await common_checks( + request_body=request_body, + team_object=None, + user_object=None, + end_user_object=None, + global_proxy_spend=None, + general_settings={}, + route=route, + llm_router=None, + proxy_logging_obj=MagicMock(), + valid_token=token, + request=request, + ) + assert result is True + + @pytest.mark.asyncio + async def test_responses_function_tool_disallowed_raises(self): + """Responses request with function tool not in allowed_tools raises.""" + route = "/v1/responses" + request_body = { + "model": "gpt-4", + "input": "What is the weather?", + "tools": [ + { + "type": "function", + "name": "get_current_weather", + "description": "Get current weather", + "parameters": {"type": "object"}, + } + ], + } + token = get_mock_user_token(metadata={"allowed_tools": ["other"]}) + request = MockRequest("POST") + + p1, p2, p3 = _tools_allowlist_patches() + with p1, p2, p3: + with pytest.raises((Exception, ProxyException)) as exc_info: + await common_checks( + request_body=request_body, + team_object=None, + user_object=None, + end_user_object=None, + global_proxy_spend=None, + general_settings={}, + route=route, + llm_router=None, + proxy_logging_obj=MagicMock(), + valid_token=token, + request=request, + ) + msg = str(exc_info.value).lower() + assert "tool" in msg or "allowed" in msg + + @pytest.mark.asyncio + async def test_responses_mcp_server_allowed_passes(self): + """Responses request with MCP server in allowed_tools passes.""" + route = "/v1/responses" + request_body = { + "model": "gpt-4", + "input": "Hi", + "tools": [ + { + "type": "mcp", + "server_label": "dmcp", + "server_description": "Example MCP server", + "server_url": "https://example.com", + "require_approval": "never", + } + ], + } + token = get_mock_user_token(metadata={"allowed_tools": ["dmcp"]}) + request = MockRequest("POST") + + p1, p2, p3 = _tools_allowlist_patches() + with p1, p2, p3: + result = await common_checks( + request_body=request_body, + team_object=None, + user_object=None, + end_user_object=None, + global_proxy_spend=None, + general_settings={}, + route=route, + llm_router=None, + proxy_logging_obj=MagicMock(), + valid_token=token, + request=request, + ) + assert result is True + + @pytest.mark.asyncio + async def test_responses_mcp_server_disallowed_raises(self): + """Responses request with MCP server not in allowed_tools raises.""" + route = "/v1/responses" + request_body = { + "model": "gpt-4", + "input": "Hi", + "tools": [ + { + "type": "mcp", + "server_label": "dmcp", + "server_description": "Example MCP server", + "server_url": "https://example.com", + "require_approval": "never", + } + ], + } + token = get_mock_user_token(metadata={"allowed_tools": ["other"]}) + request = MockRequest("POST") + + p1, p2, p3 = _tools_allowlist_patches() + with p1, p2, p3: + with pytest.raises((Exception, ProxyException)): + await common_checks( + request_body=request_body, + team_object=None, + user_object=None, + end_user_object=None, + global_proxy_spend=None, + general_settings={}, + route=route, + llm_router=None, + proxy_logging_obj=MagicMock(), + valid_token=token, + request=request, + ) + + +class TestAnthropicMessagesToolsAllowlist: + """Tool allowlist enforcement for Anthropic /v1/messages.""" + + @pytest.mark.asyncio + async def test_anthropic_allowed_tool_passes(self): + """Request with Anthropic-style tools in allowed_tools passes.""" + route = "/v1/messages" + request_body = { + "model": "claude-3-5-sonnet-20241022", + "max_tokens": 1024, + "messages": [{"role": "user", "content": "Hi"}], + "tools": [{"name": "get_weather", "description": "Get weather"}], + } + token = get_mock_user_token(metadata={"allowed_tools": ["get_weather"]}) + request = MockRequest("POST") + + p1, p2, p3 = _tools_allowlist_patches() + with p1, p2, p3: + result = await common_checks( + request_body=request_body, + team_object=None, + user_object=None, + end_user_object=None, + global_proxy_spend=None, + general_settings={}, + route=route, + llm_router=None, + proxy_logging_obj=MagicMock(), + valid_token=token, + request=request, + ) + assert result is True + + @pytest.mark.asyncio + async def test_anthropic_disallowed_tool_raises(self): + """Request with Anthropic-style tool not in allowed_tools raises.""" + route = "/v1/messages" + request_body = { + "model": "claude-3-5-sonnet-20241022", + "max_tokens": 1024, + "messages": [{"role": "user", "content": "Hi"}], + "tools": [{"name": "get_weather", "description": "Get weather"}], + } + token = get_mock_user_token(metadata={"allowed_tools": ["other_tool"]}) + request = MockRequest("POST") + + p1, p2, p3 = _tools_allowlist_patches() + with p1, p2, p3: + with pytest.raises((Exception, ProxyException)) as exc_info: + await common_checks( + request_body=request_body, + team_object=None, + user_object=None, + end_user_object=None, + global_proxy_spend=None, + general_settings={}, + route=route, + llm_router=None, + proxy_logging_obj=MagicMock(), + valid_token=token, + request=request, + ) + msg = str(exc_info.value).lower() + assert "tool" in msg or "allowed" in msg + + +class TestGoogleGenerateContentToolsAllowlist: + """Tool allowlist enforcement for Google generateContent.""" + + @pytest.mark.asyncio + async def test_google_allowed_tool_passes(self): + """Request with tools[].functionDeclarations[].name in allowed_tools passes.""" + route = "/v1beta/models/gemini-3-flash-preview:generateContent" + request_body = { + "contents": [ + {"role": "user", "parts": [{"text": "Schedule a meeting"}]} + ], + "tools": [ + { + "functionDeclarations": [ + { + "name": "schedule_meeting", + "description": "Schedules a meeting", + "parameters": {"type": "object", "properties": {}}, + } + ] + } + ], + } + token = get_mock_user_token(metadata={"allowed_tools": ["schedule_meeting"]}) + request = MockRequest("POST") + + p1, p2, p3 = _tools_allowlist_patches() + with p1, p2, p3: + result = await common_checks( + request_body=request_body, + team_object=None, + user_object=None, + end_user_object=None, + global_proxy_spend=None, + general_settings={}, + route=route, + llm_router=None, + proxy_logging_obj=MagicMock(), + valid_token=token, + request=request, + ) + assert result is True + + @pytest.mark.asyncio + async def test_google_disallowed_tool_raises(self): + """Request with tools[].functionDeclarations[].name not in allowed_tools raises.""" + route = "/v1beta/models/gemini-3-flash-preview:generateContent" + request_body = { + "contents": [ + {"role": "user", "parts": [{"text": "Schedule a meeting"}]} + ], + "tools": [ + { + "functionDeclarations": [ + { + "name": "schedule_meeting", + "description": "Schedules a meeting", + "parameters": {"type": "object", "properties": {}}, + } + ] + } + ], + } + token = get_mock_user_token(metadata={"allowed_tools": ["other_tool"]}) + request = MockRequest("POST") + + p1, p2, p3 = _tools_allowlist_patches() + with p1, p2, p3: + with pytest.raises((Exception, ProxyException)) as exc_info: + await common_checks( + request_body=request_body, + team_object=None, + user_object=None, + end_user_object=None, + global_proxy_spend=None, + general_settings={}, + route=route, + llm_router=None, + proxy_logging_obj=MagicMock(), + valid_token=token, + request=request, + ) + msg = str(exc_info.value).lower() + assert "tool" in msg or "allowed" in msg + + +# MCP REST tools/call body shape: server_id, name (tool name), arguments. +# See litellm/proxy/_experimental/mcp_server/rest_endpoints.py call_tool_rest_api. +# The exact field for tool name in the request body should match the implementation. +MCP_TOOL_CALL_BODY_ALLOWED = { + "server_id": "srv", + "name": "roll_dice", + "arguments": {}, +} + + +class TestMCPToolCallToolsAllowlist: + """Test that MCP tool call routes (/mcp/tools/call, /mcp-rest/tools/call) enforce token allowed_tools via common_checks.""" + + @pytest.mark.asyncio + async def test_mcp_tool_call_allowed_passes(self): + """Route /mcp-rest/tools/call with tool in token allowed_tools passes common_checks.""" + request = MockRequest("POST") + request_body = dict(MCP_TOOL_CALL_BODY_ALLOWED) + valid_token = get_mock_user_token(metadata={"allowed_tools": ["roll_dice"]}) + + p1, p2, p3 = _tools_allowlist_patches() + with p1, p2, p3: + result = await common_checks( + request_body=request_body, + team_object=None, + user_object=None, + end_user_object=None, + global_proxy_spend=None, + general_settings={}, + route="/mcp-rest/tools/call", + llm_router=None, + proxy_logging_obj=MagicMock(), + valid_token=valid_token, + request=request, + ) + assert result is True + + @pytest.mark.asyncio + async def test_mcp_tool_call_disallowed_raises(self): + """Route /mcp-rest/tools/call with tool not in token allowed_tools raises.""" + request = MockRequest("POST") + request_body = dict(MCP_TOOL_CALL_BODY_ALLOWED) + valid_token = get_mock_user_token(metadata={"allowed_tools": ["other"]}) + + p1, p2, p3 = _tools_allowlist_patches() + with p1, p2, p3: + with pytest.raises((Exception, ProxyException)) as exc_info: + await common_checks( + request_body=request_body, + team_object=None, + user_object=None, + end_user_object=None, + global_proxy_spend=None, + general_settings={}, + route="/mcp-rest/tools/call", + llm_router=None, + proxy_logging_obj=MagicMock(), + valid_token=valid_token, + request=request, + ) + exc_str = ( + getattr(exc_info.value, "message", None) or str(exc_info.value) or "" + ).lower() + assert "tool" in exc_str or "allowed" in exc_str diff --git a/ui/litellm-dashboard/src/components/ToolPolicies.tsx b/ui/litellm-dashboard/src/components/ToolPolicies.tsx index 6e62370b0c6..aa005495575 100644 --- a/ui/litellm-dashboard/src/components/ToolPolicies.tsx +++ b/ui/litellm-dashboard/src/components/ToolPolicies.tsx @@ -1,14 +1,41 @@ "use client"; -import React, { useCallback, useDeferredValue, useEffect, useState } from "react"; +import React, { useCallback, useDeferredValue, useEffect, useMemo, useState } from "react"; import { Select, Switch, Tooltip } from "antd"; import { Table, TableHead, TableHeaderCell, TableBody, TableRow, TableCell } from "@tremor/react"; import { TimeCell } from "./view_logs/time_cell"; import { TableHeaderSortDropdown } from "./common_components/TableHeaderSortDropdown/TableHeaderSortDropdown"; import type { SortState } from "./common_components/TableHeaderSortDropdown/TableHeaderSortDropdown"; import FilterComponent, { FilterOption } from "./molecules/filter"; +import { MetricCard } from "./GuardrailsMonitor/MetricCard"; import { fetchToolsList, updateToolPolicy, ToolRow } from "./networking"; +// --- Date helpers (UTC) for "new tools" counts --- +function getUTCDateKey(date: Date): string { + return `${date.getUTCFullYear()}-${String(date.getUTCMonth() + 1).padStart(2, "0")}-${String(date.getUTCDate()).padStart(2, "0")}`; +} + +function isCreatedInUTCDay(createdAt: string | undefined, utcDateKey: string): boolean { + if (!createdAt) return false; + try { + const d = new Date(createdAt); + return getUTCDateKey(d) === utcDateKey; + } catch { + return false; + } +} + +function countToolsInUTCDay(tools: ToolRow[], utcDateKey: string): number { + return tools.filter((t) => isCreatedInUTCDay(t.created_at, utcDateKey)).length; +} + +function getTrendSubtitle(newToday: number, newYesterday: number): string | undefined { + const diff = newToday - newYesterday; + if (diff === 0) return undefined; + if (diff > 0) return `+${diff} since yesterday`; + return `${diff} since yesterday`; +} + const POLICY_OPTIONS = [ { value: "trusted", label: "trusted", color: "#065f46", bg: "#d1fae5", border: "#6ee7b7" }, { value: "blocked", label: "blocked", color: "#991b1b", bg: "#fee2e2", border: "#fca5a5" }, @@ -197,6 +224,41 @@ export const ToolPolicies: React.FC = ({ accessToken }) => { }, ]; + // Derived counts for summary cards and "Needs Review" (UTC today/yesterday) + const { newToday, newYesterday, trendSubtitle, totalTools, blockedCount, activeTeamsCount, needsReviewTools } = + useMemo(() => { + const now = new Date(); + const todayKey = getUTCDateKey(now); + const yesterday = new Date(now); + yesterday.setUTCDate(yesterday.getUTCDate() - 1); + const yesterdayKey = getUTCDateKey(yesterday); + + const newToday = countToolsInUTCDay(tools, todayKey); + const newYesterday = countToolsInUTCDay(tools, yesterdayKey); + const trendSubtitle = getTrendSubtitle(newToday, newYesterday); + + const totalTools = tools.length; + const blockedCount = tools.filter((t) => t.call_policy === "blocked").length; + const activeTeamsCount = new Set(tools.map((t) => t.team_id).filter(Boolean)).size; + + // New in period (today) and not yet decided — untrusted or dual_llm + const needsReviewTools = tools.filter( + (t) => + isCreatedInUTCDay(t.created_at, todayKey) && + (t.call_policy === "untrusted" || t.call_policy === "dual_llm") + ); + + return { + newToday, + newYesterday, + trendSubtitle, + totalTools, + blockedCount, + activeTeamsCount, + needsReviewTools, + }; + }, [tools]); + const SortHeader = ({ label, field }: { label: string; field: SortField }) => (
{label} @@ -235,9 +297,76 @@ export const ToolPolicies: React.FC = ({ accessToken }) => { const totalPages = Math.max(1, Math.ceil(sorted.length / pageSize)); const paginated = sorted.slice((currentPage - 1) * pageSize, currentPage * pageSize); + const scrollToToolRow = (toolId: string) => { + const idx = sorted.findIndex((t) => t.tool_id === toolId); + if (idx >= 0) { + const page = Math.floor(idx / pageSize) + 1; + if (page !== currentPage) setCurrentPage(page); + // Scroll after a short delay so the table has re-rendered with the new page + requestAnimationFrame(() => { + setTimeout(() => { + document.getElementById(`tool-row-${toolId}`)?.scrollIntoView({ behavior: "smooth", block: "center" }); + }, 100); + }); + } + }; + return (

Tool Policies

+ + {/* Summary cards */} +
+ + + + } + /> + + 0 ? "text-red-600" : undefined} + /> + 0 ? activeTeamsCount : "—"} /> +
+ + {/* Needs Review */} + {needsReviewTools.length > 0 && ( +
+

Needs Review

+

+ {needsReviewTools.length} new tool{needsReviewTools.length !== 1 ? "s" : ""} discovered that require + policy decisions. +

+
+ {needsReviewTools.map((t) => ( + + + {t.tool_name} + + + + ))} +
+
+ )} +
{/* Toolbar */}
@@ -389,7 +518,7 @@ export const ToolPolicies: React.FC = ({ accessToken }) => { ) : ( paginated.map((tool) => ( - + From fffe8253ea4e68ab6664dde24a147d219d987a21 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 25 Feb 2026 22:12:23 -0800 Subject: [PATCH 05/12] feat: ensure tool allowlist runs correctly for tool names + mcp's --- .../migration.sql | 27 + .../litellm_proxy_extras/schema.prisma | 19 + .../chat/guardrail_translation/handler.py | 10 +- .../guardrail_translation/base_translation.py | 7 + .../chat/guardrail_translation/handler.py | 13 + .../guardrail_translation/handler.py | 39 +- litellm/proxy/_types.py | 80 ++- litellm/proxy/auth/auth_checks.py | 139 ++-- litellm/proxy/db/tool_registry_writer.py | 255 ++++++- .../tool_policy/tool_policy_guardrail.py | 151 ++-- .../proxy/guardrails/tool_name_extraction.py | 85 +++ litellm/proxy/litellm_pre_call_utils.py | 50 +- .../tool_management_endpoints.py | 128 +++- litellm/proxy/schema.prisma | 19 + litellm/types/tool_management.py | 23 +- scripts/test_tool_allowlist_script.py | 116 +++ .../proxy/test_tools_allowlist_enforcement.py | 674 +++++------------- .../src/components/ToolPolicies.tsx | 266 ++++++- .../src/components/networking.tsx | 86 ++- 19 files changed, 1477 insertions(+), 710 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260225220000_add_tool_policy_override_table/migration.sql create mode 100644 litellm/proxy/guardrails/tool_name_extraction.py create mode 100644 scripts/test_tool_allowlist_script.py diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260225220000_add_tool_policy_override_table/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260225220000_add_tool_policy_override_table/migration.sql new file mode 100644 index 00000000000..642508b7b20 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260225220000_add_tool_policy_override_table/migration.sql @@ -0,0 +1,27 @@ +-- CreateTable +CREATE TABLE "LiteLLM_ToolPolicyOverrideTable" ( + "override_id" TEXT NOT NULL, + "tool_name" TEXT NOT NULL, + "team_id" TEXT NOT NULL DEFAULT '', + "key_hash" TEXT NOT NULL DEFAULT '', + "call_policy" TEXT NOT NULL DEFAULT 'blocked', + "key_alias" TEXT, + "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "created_by" TEXT, + "updated_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "updated_by" TEXT, + + CONSTRAINT "LiteLLM_ToolPolicyOverrideTable_pkey" PRIMARY KEY ("override_id") +); + +-- CreateIndex +CREATE INDEX "LiteLLM_ToolPolicyOverrideTable_tool_name_idx" ON "LiteLLM_ToolPolicyOverrideTable"("tool_name"); + +-- CreateIndex +CREATE INDEX "LiteLLM_ToolPolicyOverrideTable_team_id_idx" ON "LiteLLM_ToolPolicyOverrideTable"("team_id"); + +-- CreateIndex +CREATE INDEX "LiteLLM_ToolPolicyOverrideTable_key_hash_idx" ON "LiteLLM_ToolPolicyOverrideTable"("key_hash"); + +-- CreateIndex +CREATE UNIQUE INDEX "LiteLLM_ToolPolicyOverrideTable_tool_name_team_id_key_hash_key" ON "LiteLLM_ToolPolicyOverrideTable"("tool_name", "team_id", "key_hash"); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 440c9c1d829..e48c0fe3027 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -1076,6 +1076,25 @@ model LiteLLM_ToolTable { @@index([team_id]) } +// Per-(tool, team/key) policy overrides. When present, override replaces global tool policy for that scope. +model LiteLLM_ToolPolicyOverrideTable { + override_id String @id @default(uuid()) + tool_name String + team_id String @default("") // "" = not scoped to team; non-empty = override for this team only + key_hash String @default("") // "" = not scoped to key; non-empty = override for this key only + call_policy String @default("blocked") + key_alias String? // human-readable key alias for UI + created_at DateTime @default(now()) + created_by String? + updated_at DateTime @default(now()) @updatedAt + updated_by String? + + @@unique([tool_name, team_id, key_hash]) + @@index([tool_name]) + @@index([team_id]) + @@index([key_hash]) +} + //Unified Access Groups table for storing unified access groups model LiteLLM_AccessGroupTable { access_group_id String @id @default(uuid()) diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py index 98650a238e9..a6df346e8a8 100644 --- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py @@ -75,7 +75,7 @@ class AnthropicMessagesHandler(BaseTranslation): if messages is None: return data - chat_completion_compatible_request, tool_name_mapping = ( + chat_completion_compatible_request, _tool_name_mapping = ( LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai( # Use a shallow copy to avoid mutating request data (pop on litellm_metadata). anthropic_message_request=cast(AnthropicMessagesRequest, data.copy()) @@ -141,6 +141,14 @@ class AnthropicMessagesHandler(BaseTranslation): return data + def extract_request_tool_names(self, data: dict) -> List[str]: + """Extract tool names from Anthropic messages request (tools[].name).""" + names: List[str] = [] + for tool in data.get("tools") or []: + if isinstance(tool, dict) and tool.get("name"): + names.append(str(tool["name"])) + return names + def _extract_input_text_and_images( self, message: Dict[str, Any], diff --git a/litellm/llms/base_llm/guardrail_translation/base_translation.py b/litellm/llms/base_llm/guardrail_translation/base_translation.py index 7106c207bd6..a7982cb606e 100644 --- a/litellm/llms/base_llm/guardrail_translation/base_translation.py +++ b/litellm/llms/base_llm/guardrail_translation/base_translation.py @@ -98,3 +98,10 @@ class BaseTranslation(ABC): Optional to override in subclasses. """ return responses_so_far + + def extract_request_tool_names(self, data: dict) -> List[str]: + """ + Extract tool names from the request body for allowlist/policy checks. + Override in tool-capable handlers; default returns []. + """ + return [] diff --git a/litellm/llms/openai/chat/guardrail_translation/handler.py b/litellm/llms/openai/chat/guardrail_translation/handler.py index 683e165c315..0658a953318 100644 --- a/litellm/llms/openai/chat/guardrail_translation/handler.py +++ b/litellm/llms/openai/chat/guardrail_translation/handler.py @@ -135,6 +135,19 @@ class OpenAIChatCompletionsHandler(BaseTranslation): return data + def extract_request_tool_names(self, data: dict) -> List[str]: + """Extract tool names from OpenAI chat completions request (tools[].function.name, functions[].name).""" + names: List[str] = [] + for tool in data.get("tools") or []: + if isinstance(tool, dict) and tool.get("type") == "function": + fn = tool.get("function") + if isinstance(fn, dict) and fn.get("name"): + names.append(str(fn["name"])) + for fn in data.get("functions") or []: + if isinstance(fn, dict) and fn.get("name"): + names.append(str(fn["name"])) + return names + def _extract_inputs( self, message: Dict[str, Any], diff --git a/litellm/llms/openai/responses/guardrail_translation/handler.py b/litellm/llms/openai/responses/guardrail_translation/handler.py index 6b092911d3c..7c3354cf88e 100644 --- a/litellm/llms/openai/responses/guardrail_translation/handler.py +++ b/litellm/llms/openai/responses/guardrail_translation/handler.py @@ -30,27 +30,22 @@ Output: response.output is List[GenericResponseOutputItem] where each has: from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast -from openai.types.responses.response_function_tool_call import ResponseFunctionToolCall +from openai.types.responses.response_function_tool_call import \ + ResponseFunctionToolCall from pydantic import BaseModel from litellm._logging import verbose_proxy_logger from litellm.completion_extras.litellm_responses_transformation.transformation import ( LiteLLMResponsesTransformationHandler, - OpenAiResponsesToChatCompletionStreamIterator, -) -from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation -from litellm.responses.litellm_completion_transformation.transformation import ( - LiteLLMCompletionResponsesConfig, -) -from litellm.types.llms.openai import ( - ChatCompletionToolCallChunk, - ChatCompletionToolParam, -) -from litellm.types.responses.main import ( - GenericResponseOutputItem, - OutputFunctionToolCall, - OutputText, -) + OpenAiResponsesToChatCompletionStreamIterator) +from litellm.llms.base_llm.guardrail_translation.base_translation import \ + BaseTranslation +from litellm.responses.litellm_completion_transformation.transformation import \ + LiteLLMCompletionResponsesConfig +from litellm.types.llms.openai import (ChatCompletionToolCallChunk, + ChatCompletionToolParam) +from litellm.types.responses.main import (GenericResponseOutputItem, + OutputFunctionToolCall, OutputText) from litellm.types.utils import GenericGuardrailAPIInputs if TYPE_CHECKING: @@ -188,6 +183,18 @@ class OpenAIResponsesHandler(BaseTranslation): return data + def extract_request_tool_names(self, data: dict) -> List[str]: + """Extract tool names from Responses API request (tools[].name for function, tools[].server_label for mcp).""" + names: List[str] = [] + for tool in data.get("tools") or []: + if not isinstance(tool, dict): + continue + if tool.get("type") == "function" and tool.get("name"): + names.append(str(tool["name"])) + elif tool.get("type") == "mcp" and tool.get("server_label"): + names.append(str(tool["server_label"])) + return names + def _extract_and_transform_tools( self, tools: List[Dict[str, Any]], diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 88cfa5f6c11..65ea9cb42d8 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1,40 +1,60 @@ import enum import json from datetime import datetime -from typing import (TYPE_CHECKING, Any, Callable, Dict, List, Literal, - Optional, Union) +from typing import TYPE_CHECKING, Any, Callable, Dict, List, Literal, Optional, Union import httpx -from pydantic import (BaseModel, ConfigDict, Field, Json, field_validator, - model_validator) +from pydantic import ( + BaseModel, + ConfigDict, + Field, + Json, + field_validator, + model_validator, +) from typing_extensions import Required, TypedDict from litellm._uuid import uuid from litellm.types.integrations.slack_alerting import AlertType -from litellm.types.llms.openai import (AllMessageValues, OpenAIFileObject, - ResponsesAPIResponse) -from litellm.types.mcp import (MCPAuth, MCPAuthType, MCPCredentials, - MCPTransport, MCPTransportType) +from litellm.types.llms.openai import ( + AllMessageValues, + OpenAIFileObject, + ResponsesAPIResponse, +) +from litellm.types.mcp import ( + MCPAuth, + MCPAuthType, + MCPCredentials, + MCPTransport, + MCPTransportType, +) from litellm.types.mcp_server.mcp_server_manager import MCPInfo from litellm.types.router import RouterErrors, UpdateRouterConfig from litellm.types.secret_managers.main import KeyManagementSystem -from litellm.types.utils import (CallTypes, CostBreakdown, EmbeddingResponse, - GenericBudgetConfigType, ImageResponse, - LiteLLMBatch, LiteLLMFineTuningJob, - LiteLLMPydanticObjectBase, ModelResponse, - ProviderField, StandardCallbackDynamicParams, - StandardLoggingGuardrailInformation, - StandardLoggingMCPToolCall, - StandardLoggingModelInformation, - StandardLoggingPayloadErrorInformation, - StandardLoggingPayloadStatus, - StandardLoggingVectorStoreRequest, - StandardPassThroughResponseObject, - TextCompletionResponse) +from litellm.types.utils import ( + CallTypes, + CostBreakdown, + EmbeddingResponse, + GenericBudgetConfigType, + ImageResponse, + LiteLLMBatch, + LiteLLMFineTuningJob, + LiteLLMPydanticObjectBase, + ModelResponse, + ProviderField, + StandardCallbackDynamicParams, + StandardLoggingGuardrailInformation, + StandardLoggingMCPToolCall, + StandardLoggingModelInformation, + StandardLoggingPayloadErrorInformation, + StandardLoggingPayloadStatus, + StandardLoggingVectorStoreRequest, + StandardPassThroughResponseObject, + TextCompletionResponse, +) from litellm.types.videos.main import VideoObject -from .types_utils.utils import (get_instance_fn, - validate_custom_validate_return_type) +from .types_utils.utils import get_instance_fn, validate_custom_validate_return_type if TYPE_CHECKING: from opentelemetry.trace import Span as _Span @@ -2349,8 +2369,7 @@ class UserAPIKeyAuth( This is used to track number of requests/spend for health check calls. """ - from litellm.constants import \ - LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME + from litellm.constants import LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME return cls( api_key=LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME, @@ -2382,8 +2401,7 @@ class UserAPIKeyAuth( This is used to track actions performed by automated system jobs. """ - from litellm.constants import \ - LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME + from litellm.constants import LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME return cls( api_key=LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME, @@ -2774,8 +2792,7 @@ class LiteLLM_AuditLogs(LiteLLMPydanticObjectBase): @model_validator(mode="after") def mask_api_keys(self): - from litellm.litellm_core_utils.sensitive_data_masker import \ - SensitiveDataMasker + from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker masker = SensitiveDataMasker(sensitive_patterns={"key"}) @@ -3324,6 +3341,11 @@ class ProxyErrorTypes(str, enum.Enum): Team member is already in team """ + tool_access_denied = "tool_access_denied" + """ + Tool is not in the allowed tools list for this key/team + """ + @classmethod def get_model_access_error_type_for_object( cls, object_type: Literal["key", "user", "team", "org", "project"] diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 500a39d9455..7b44699825c 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -11,7 +11,8 @@ Run checks for: import asyncio import re import time -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, cast +from typing import (TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, + cast) from fastapi import HTTPException, Request, status from pydantic import BaseModel @@ -20,44 +21,33 @@ import litellm from litellm._logging import verbose_proxy_logger from litellm.caching.caching import DualCache from litellm.caching.dual_cache import LimitedSizeOrderedDict -from litellm.constants import ( - CLI_JWT_EXPIRATION_HOURS, - CLI_JWT_TOKEN_NAME, - DEFAULT_ACCESS_GROUP_CACHE_TTL, - DEFAULT_IN_MEMORY_TTL, - DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, - DEFAULT_MAX_RECURSE_DEPTH, - EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE, -) +from litellm.constants import (CLI_JWT_EXPIRATION_HOURS, CLI_JWT_TOKEN_NAME, + DEFAULT_ACCESS_GROUP_CACHE_TTL, + DEFAULT_IN_MEMORY_TTL, + DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, + DEFAULT_MAX_RECURSE_DEPTH, + EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE) from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider -from litellm.proxy._types import ( - RBAC_ROLES, - CallInfo, - LiteLLM_AccessGroupTable, - LiteLLM_BudgetTable, - LiteLLM_EndUserTable, - Litellm_EntityType, - LiteLLM_JWTAuth, - LiteLLM_ObjectPermissionTable, - LiteLLM_OrganizationMembershipTable, - LiteLLM_OrganizationTable, - LiteLLM_ProjectTableCachedObj, - LiteLLM_TagTable, - LiteLLM_TeamMembership, - LiteLLM_TeamTable, - LiteLLM_TeamTableCachedObj, - LiteLLM_UserTable, - LiteLLMRoutes, - LitellmUserRoles, - NewTeamRequest, - ProxyErrorTypes, - ProxyException, - RoleBasedPermissions, - SpecialModelNames, - UserAPIKeyAuth, -) +from litellm.proxy._types import (RBAC_ROLES, CallInfo, + LiteLLM_AccessGroupTable, + LiteLLM_BudgetTable, LiteLLM_EndUserTable, + Litellm_EntityType, LiteLLM_JWTAuth, + LiteLLM_ObjectPermissionTable, + LiteLLM_OrganizationMembershipTable, + LiteLLM_OrganizationTable, + LiteLLM_ProjectTableCachedObj, + LiteLLM_TagTable, LiteLLM_TeamMembership, + LiteLLM_TeamTable, + LiteLLM_TeamTableCachedObj, + LiteLLM_UserTable, LiteLLMRoutes, + LitellmUserRoles, NewTeamRequest, + ProxyErrorTypes, ProxyException, + RoleBasedPermissions, SpecialModelNames, + UserAPIKeyAuth) from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler +from litellm.proxy.guardrails.tool_name_extraction import ( + TOOL_CAPABLE_CALL_TYPES, extract_request_tool_names) from litellm.proxy.route_llm_request import route_request from litellm.proxy.utils import PrismaClient, ProxyLogging, log_db_metrics from litellm.router import Router @@ -220,7 +210,47 @@ async def _run_project_checks( ) -async def common_checks( +async def check_tools_allowlist( + request_body: dict, + valid_token: Optional[UserAPIKeyAuth], + team_object: Optional[LiteLLM_TeamTable], + route: str, +) -> None: + """ + Enforce key/team tool allowlist (metadata.allowed_tools). No DB in hot path — + effective allowlist is read from valid_token.metadata and valid_token.team_metadata. + Raises ProxyException with tool_access_denied if a tool is not allowed. + """ + from litellm.litellm_core_utils.api_route_to_call_types import \ + get_call_types_for_route + + if valid_token is None: + return + call_types = get_call_types_for_route(route) + if not call_types or not any(ct.value in TOOL_CAPABLE_CALL_TYPES for ct in call_types): + return + tool_names = extract_request_tool_names(route, request_body) + if not tool_names: + return + key_meta = (valid_token.metadata or {}) if isinstance(valid_token.metadata, dict) else {} + team_meta = (valid_token.team_metadata or {}) if isinstance(valid_token.team_metadata, dict) else {} + key_allowed = key_meta.get("allowed_tools") + team_allowed = team_meta.get("allowed_tools") + effective = key_allowed if (isinstance(key_allowed, list) and len(key_allowed) > 0) else team_allowed + if not isinstance(effective, list) or len(effective) == 0: + return + allowed_set = {str(t) for t in effective} + disallowed = [n for n in tool_names if n not in allowed_set] + if disallowed: + raise ProxyException( + message=f"Tool(s) {disallowed} are not in the allowed tools list for this key/team.", + type=ProxyErrorTypes.tool_access_denied, + param="tools", + code=status.HTTP_403_FORBIDDEN, + ) + + +async def common_checks( # noqa: PLR0915 request_body: dict, team_object: Optional[LiteLLM_TeamTable], user_object: Optional[LiteLLM_UserTable], @@ -435,7 +465,8 @@ async def common_checks( _request_metadata: dict = request_body.get("metadata", {}) or {} if _request_metadata.get("guardrails"): # check if team allowed to modify guardrails - from litellm.proxy.guardrails.guardrail_helpers import can_modify_guardrails + from litellm.proxy.guardrails.guardrail_helpers import \ + can_modify_guardrails can_modify: bool = can_modify_guardrails(team_object) if can_modify is False: @@ -473,6 +504,14 @@ async def common_checks( valid_token=valid_token, ) + # 12. [OPTIONAL] Tool allowlist - key/team allowed_tools (no DB in hot path) + await check_tools_allowlist( + request_body=request_body, + valid_token=valid_token, + team_object=team_object, + route=route, + ) + return True @@ -1877,9 +1916,8 @@ class ExperimentalUIJWTToken: def get_experimental_ui_login_jwt_auth_token(user_info: LiteLLM_UserTable) -> str: from datetime import timedelta - from litellm.proxy.common_utils.encrypt_decrypt_utils import ( - encrypt_value_helper, - ) + from litellm.proxy.common_utils.encrypt_decrypt_utils import \ + encrypt_value_helper if user_info.user_role is None: raise Exception("User role is required for experimental UI login") @@ -1925,9 +1963,8 @@ class ExperimentalUIJWTToken: """ from datetime import timedelta - from litellm.proxy.common_utils.encrypt_decrypt_utils import ( - encrypt_value_helper, - ) + from litellm.proxy.common_utils.encrypt_decrypt_utils import \ + encrypt_value_helper if user_info.user_role is None: raise Exception("User role is required for CLI JWT login") @@ -1966,9 +2003,8 @@ class ExperimentalUIJWTToken: import json from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth - from litellm.proxy.common_utils.encrypt_decrypt_utils import ( - decrypt_value_helper, - ) + from litellm.proxy.common_utils.encrypt_decrypt_utils import \ + decrypt_value_helper decrypted_token = decrypt_value_helper( hashed_token, key="ui_hash_key", exception_type="debug" @@ -2263,8 +2299,10 @@ async def _get_resources_from_access_groups( # Lazy import to avoid circular imports if prisma_client is None or user_api_key_cache is None: from litellm.proxy.proxy_server import prisma_client as _prisma_client - from litellm.proxy.proxy_server import proxy_logging_obj as _proxy_logging_obj - from litellm.proxy.proxy_server import user_api_key_cache as _user_api_key_cache + from litellm.proxy.proxy_server import \ + proxy_logging_obj as _proxy_logging_obj + from litellm.proxy.proxy_server import \ + user_api_key_cache as _user_api_key_cache prisma_client = prisma_client or _prisma_client user_api_key_cache = user_api_key_cache or _user_api_key_cache @@ -3220,7 +3258,8 @@ async def _tag_max_budget_check( BudgetExceededError if any tag is over its max budget. Triggers a budget alert if any tag is over its max budget. """ - from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body + from litellm.proxy.common_utils.http_parsing_utils import \ + get_tags_from_request_body if prisma_client is None: return diff --git a/litellm/proxy/db/tool_registry_writer.py b/litellm/proxy/db/tool_registry_writer.py index 332607f0305..a0dffefdc59 100644 --- a/litellm/proxy/db/tool_registry_writer.py +++ b/litellm/proxy/db/tool_registry_writer.py @@ -10,12 +10,23 @@ from datetime import datetime, timezone from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union from litellm._logging import verbose_proxy_logger +from litellm.caching.dual_cache import DualCache +from litellm.constants import TOOL_POLICY_CACHE_TTL_SECONDS from litellm.proxy._types import ToolDiscoveryQueueItem -from litellm.types.tool_management import LiteLLM_ToolTableRow, ToolCallPolicy +from litellm.types.tool_management import ( + LiteLLM_ToolTableRow, + ToolCallPolicy, + ToolPolicyOverrideRow, +) if TYPE_CHECKING: from litellm.proxy.utils import PrismaClient +# Sentinel for "not scoped" in override table (DB unique constraint needs non-null) +_TOOL_OVERRIDE_ANY = "" + +TOOL_POLICY_CACHE_KEY_PREFIX = "tool_policy:" + def _row_to_model(row: Union[dict, Any]) -> LiteLLM_ToolTableRow: """Convert a Prisma model instance or dict to LiteLLM_ToolTableRow.""" @@ -204,3 +215,245 @@ async def get_tools_by_names( "tool_registry_writer get_tools_by_names error: %s", e ) return {} + + +def _override_row_to_model(row: Any) -> ToolPolicyOverrideRow: + """Convert a Prisma override row to ToolPolicyOverrideRow.""" + model_dump = getattr(row, "model_dump", None) + if callable(model_dump): + row = model_dump() + elif not isinstance(row, dict): + row = { + k: getattr(row, k, None) + for k in ( + "override_id", + "tool_name", + "team_id", + "key_hash", + "call_policy", + "key_alias", + "created_at", + "updated_at", + ) + } + def _norm(s: Optional[str]) -> Optional[str]: + if s is None or s == _TOOL_OVERRIDE_ANY: + return None + return s or None + return ToolPolicyOverrideRow( + override_id=row.get("override_id", ""), + tool_name=row.get("tool_name", ""), + team_id=_norm(row.get("team_id")), + key_hash=_norm(row.get("key_hash")), + call_policy=row.get("call_policy", "blocked"), + key_alias=row.get("key_alias"), + created_at=row.get("created_at"), + updated_at=row.get("updated_at"), + ) + + +async def list_overrides_for_tool( + prisma_client: "PrismaClient", + tool_name: str, +) -> List[ToolPolicyOverrideRow]: + """Return all policy overrides for a tool.""" + try: + rows = await prisma_client.db.litellm_toolpolicyoverridetable.find_many( + where={"tool_name": tool_name}, + order={"created_at": "desc"}, + ) + return [_override_row_to_model(row) for row in rows] + except Exception as e: + verbose_proxy_logger.error( + "tool_registry_writer list_overrides_for_tool error: %s", e + ) + return [] + + +async def upsert_tool_policy_override( + prisma_client: "PrismaClient", + tool_name: str, + call_policy: ToolCallPolicy, + team_id: Optional[str] = None, + key_hash: Optional[str] = None, + key_alias: Optional[str] = None, + updated_by: Optional[str] = None, +) -> Optional[ToolPolicyOverrideRow]: + """Create or update a per-(tool, team, key) policy override.""" + try: + _team = (team_id or "").strip() or _TOOL_OVERRIDE_ANY + _key = (key_hash or "").strip() or _TOOL_OVERRIDE_ANY + _updated_by = updated_by or "system" + now = datetime.now(timezone.utc) + table = prisma_client.db.litellm_toolpolicyoverridetable + await table.upsert( + where={ + "tool_name_team_id_key_hash": { + "tool_name": tool_name, + "team_id": _team, + "key_hash": _key, + } + }, + data={ + "create": { + "override_id": str(uuid.uuid4()), + "tool_name": tool_name, + "team_id": _team, + "key_hash": _key, + "call_policy": call_policy, + "key_alias": key_alias, + "created_by": _updated_by, + "updated_by": _updated_by, + "created_at": now, + "updated_at": now, + }, + "update": { + "call_policy": call_policy, + "key_alias": key_alias, + "updated_by": _updated_by, + "updated_at": now, + }, + }, + ) + row = await table.find_unique( + where={ + "tool_name_team_id_key_hash": { + "tool_name": tool_name, + "team_id": _team, + "key_hash": _key, + } + } + ) + if row is None: + return None + return _override_row_to_model(row) + except Exception as e: + verbose_proxy_logger.error( + "tool_registry_writer upsert_tool_policy_override error: %s", e + ) + return None + + +async def delete_tool_policy_override( + prisma_client: "PrismaClient", + tool_name: str, + team_id: Optional[str] = None, + key_hash: Optional[str] = None, +) -> bool: + """Remove a policy override. Exactly one of team_id or key_hash should be set for a specific override.""" + try: + _team = (team_id or "").strip() or _TOOL_OVERRIDE_ANY + _key = (key_hash or "").strip() or _TOOL_OVERRIDE_ANY + result = await prisma_client.db.litellm_toolpolicyoverridetable.delete_many( + where={ + "tool_name": tool_name, + "team_id": _team, + "key_hash": _key, + } + ) + return (result or 0) > 0 + except Exception as e: + verbose_proxy_logger.error( + "tool_registry_writer delete_tool_policy_override error: %s", e + ) + return False + + +async def get_effective_policies( + prisma_client: "PrismaClient", + tool_names: List[str], + team_id: Optional[str], + key_hash: Optional[str], +) -> Dict[str, str]: + """ + Return effective call_policy per tool: override for (tool, team_id, key_hash) if present, + otherwise global policy from LiteLLM_ToolTable. + """ + if not tool_names: + return {} + _team = (team_id or "").strip() or _TOOL_OVERRIDE_ANY + _key = (key_hash or "").strip() or _TOOL_OVERRIDE_ANY + try: + # Fetch overrides for (tool_name, team_id, key_hash) - exact match + overrides = await prisma_client.db.litellm_toolpolicyoverridetable.find_many( + where={ + "tool_name": {"in": tool_names}, + "team_id": _team, + "key_hash": _key, + } + ) + override_map = {row.tool_name: row.call_policy for row in overrides} + # Global policies for tools that have no override + missing = [n for n in tool_names if n not in override_map] + if not missing: + return override_map + global_map = await get_tools_by_names( + prisma_client=prisma_client, tool_names=missing + ) + override_map.update(global_map) + return override_map + except Exception as e: + verbose_proxy_logger.error( + "tool_registry_writer get_effective_policies error: %s", e + ) + return {} + + +def _effective_cache_suffix(team_id: Optional[str], key_hash: Optional[str]) -> str: + """Cache key suffix so different request contexts (team/key) get correct policies.""" + return f":{team_id or ''}:{key_hash or ''}" + + +async def get_tool_policies_cached( + tool_names: List[str], + cache: DualCache, + prisma_client: Optional["PrismaClient"], + team_id: Optional[str] = None, + key_hash: Optional[str] = None, +) -> Dict[str, str]: + """ + Return effective call_policy per tool (override for team/key if present, else global). + Cache-first; cache key includes team_id and key_hash when provided. + """ + if not tool_names: + return {} + suffix = _effective_cache_suffix(team_id, key_hash) + result: Dict[str, str] = {} + cache_misses: List[str] = [] + for name in tool_names: + key = f"{TOOL_POLICY_CACHE_KEY_PREFIX}{name}{suffix}" + cached = await cache.async_get_cache(key=key) + if cached is not None and isinstance(cached, str): + result[name] = cached + else: + cache_misses.append(name) + if cache_misses and prisma_client is not None: + try: + if team_id is not None or key_hash is not None: + fetched = await get_effective_policies( + prisma_client=prisma_client, + tool_names=cache_misses, + team_id=team_id, + key_hash=key_hash, + ) + else: + fetched = await get_tools_by_names( + prisma_client=prisma_client, tool_names=cache_misses + ) + for name, policy in fetched.items(): + result[name] = policy + await cache.async_set_cache( + key=f"{TOOL_POLICY_CACHE_KEY_PREFIX}{name}{suffix}", + value=policy, + ttl=TOOL_POLICY_CACHE_TTL_SECONDS, + ) + verbose_proxy_logger.debug( + "get_tool_policies_cached: fetched %d from DB (hits: %d)", + len(cache_misses), + len(tool_names) - len(cache_misses), + ) + except Exception as e: + verbose_proxy_logger.error( + "tool_registry_writer get_tool_policies_cached error: %s", e + ) + return result diff --git a/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py index 87558566c42..213b248929e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py @@ -23,32 +23,70 @@ or both pre and post call: mode: during_call # runs before LLM and on response """ -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional +from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple from fastapi import HTTPException from litellm._logging import verbose_proxy_logger from litellm.caching.dual_cache import DualCache -from litellm.constants import TOOL_POLICY_CACHE_TTL_SECONDS -from litellm.integrations.custom_guardrail import ( - CustomGuardrail, - log_guardrail_information, -) +from litellm.integrations.custom_guardrail import (CustomGuardrail, + log_guardrail_information) from litellm.types.guardrails import GuardrailEventHooks from litellm.types.utils import GenericGuardrailAPIInputs if TYPE_CHECKING: - from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.litellm_core_utils.litellm_logging import \ + Logging as LiteLLMLoggingObj GUARDRAIL_NAME = "tool_policy" +def _get_request_team_and_key(request_data: dict) -> Tuple[Optional[str], Optional[str]]: + """Extract team_id and key hash from request_data (litellm_metadata or metadata).""" + if not request_data: + return None, None + for key in ("litellm_metadata", "metadata"): + meta = request_data.get(key) + if not isinstance(meta, dict): + continue + team_id = meta.get("user_api_key_team_id") + key_hash = meta.get("user_api_key_api_key") or meta.get("user_api_key") + if team_id is not None or key_hash is not None: + return ( + str(team_id).strip() if team_id else None, + str(key_hash).strip() if key_hash else None, + ) + return None, None + + +def _get_request_route_from_data(request_data: dict) -> Optional[str]: + """Get request route from request_data (metadata or top-level).""" + route = request_data.get("user_api_key_request_route") + if route: + return route + meta = request_data.get("metadata") or request_data.get("litellm_metadata") or {} + return meta.get("user_api_key_request_route") + + +def _get_effective_allowed_tools_from_request(request_data: dict) -> Optional[List[str]]: + """Key allowed_tools overrides team; empty/missing means no restriction.""" + meta = request_data.get("metadata") or request_data.get("litellm_metadata") or {} + key_meta = meta.get("user_api_key_metadata") or {} + team_meta = meta.get("user_api_key_team_metadata") or {} + key_allowed = key_meta.get("allowed_tools") if isinstance(key_meta, dict) else None + team_allowed = team_meta.get("allowed_tools") if isinstance(team_meta, dict) else None + if isinstance(key_allowed, list) and len(key_allowed) > 0: + return key_allowed + if isinstance(team_allowed, list) and len(team_allowed) > 0: + return team_allowed + return None + + class ToolPolicyGuardrail(CustomGuardrail): """ - Guardrail that enforces per-tool call policies stored in LiteLLM_ToolTable. - - Tools with call_policy="blocked" are rejected before/after the LLM call. - Tools with call_policy="trusted" or "untrusted" pass through unchanged. + Guardrail that enforces per-tool call policies stored in LiteLLM_ToolTable + and key/team allowed_tools (allowlist). No DB in hot path for policy lookup — + uses shared cache (user_api_key_cache). """ def __init__(self, **kwargs: Any) -> None: @@ -70,12 +108,7 @@ class ToolPolicyGuardrail(CustomGuardrail): logging_obj: Optional["LiteLLMLoggingObj"] = None, ) -> GenericGuardrailAPIInputs: """ - Enforce tool policies on both request tools and response tool_calls. - - - input_type="request": check inputs["tools"] (tool definitions in the LLM request) - - input_type="response": check inputs["tool_calls"] (tool_calls in the LLM response) - - Raises HTTPException (400) if any tool is "blocked". + Enforce key/team allowlist then DB call_policy on request tools / response tool_calls. """ if input_type == "request": tools = inputs.get("tools") or [] @@ -86,6 +119,12 @@ class ToolPolicyGuardrail(CustomGuardrail): and isinstance(t.get("function"), dict) and t["function"].get("name") ] + if not tool_names: + route = _get_request_route_from_data(request_data) + if route: + from litellm.proxy.guardrails.tool_name_extraction import \ + extract_request_tool_names + tool_names = extract_request_tool_names(route, request_data) else: # response tool_calls = inputs.get("tool_calls") or [] tool_names = [] @@ -101,8 +140,26 @@ class ToolPolicyGuardrail(CustomGuardrail): if not tool_names: return inputs - policy_map = await self._get_policies_cached(tool_names) + allowed_tools = _get_effective_allowed_tools_from_request(request_data) + if isinstance(allowed_tools, list) and len(allowed_tools) > 0: + allowed_set = {str(t) for t in allowed_tools} + disallowed = [n for n in tool_names if n not in allowed_set] + if disallowed: + verbose_proxy_logger.warning( + "ToolPolicyGuardrail: tool(s) %s not in key/team allowed_tools", + disallowed, + ) + raise HTTPException( + status_code=400, + detail={ + "error": "Violated tool allowlist", + "disallowed_tools": disallowed, + "message": f"Tool(s) {disallowed} are not in the allowed tools list for this key/team.", + }, + ) + team_id, key_hash = _get_request_team_and_key(request_data) + policy_map = await self._get_policies_cached(tool_names, team_id, key_hash) blocked = [name for name in tool_names if policy_map.get(name) == "blocked"] if blocked: verbose_proxy_logger.warning( @@ -119,45 +176,27 @@ class ToolPolicyGuardrail(CustomGuardrail): return inputs - async def _get_policies_cached(self, tool_names: List[str]) -> Dict[str, str]: + async def _get_policies_cached( + self, + tool_names: List[str], + team_id: Optional[str] = None, + key_hash: Optional[str] = None, + ) -> Dict[str, str]: """ - Batch-fetch call_policy for the given tool names. - - Caches per individual tool name (not per combination) so that adding - a new tool to a request doesn't invalidate the cached policies for all - the other tools already in the cache. + Fetch effective call_policy (override for team/key if present, else global) + via shared cache to avoid DB in hot path. """ - from litellm.proxy.db.tool_registry_writer import get_tools_by_names - from litellm.proxy.proxy_server import prisma_client + from litellm.proxy.db.tool_registry_writer import \ + get_tool_policies_cached + from litellm.proxy.proxy_server import (prisma_client, + user_api_key_cache) - if not tool_names or prisma_client is None: + if not tool_names: return {} - - result: Dict[str, str] = {} - cache_misses: List[str] = [] - - for name in tool_names: - cached = await self._policy_cache.async_get_cache(f"tool_policy:{name}") - if cached is not None and isinstance(cached, str): - result[name] = cached - else: - cache_misses.append(name) - - if cache_misses: - fetched = await get_tools_by_names( - prisma_client=prisma_client, tool_names=cache_misses - ) - for name, policy in fetched.items(): - result[name] = policy - await self._policy_cache.async_set_cache( - key=f"tool_policy:{name}", - value=policy, - ttl=TOOL_POLICY_CACHE_TTL_SECONDS, - ) - verbose_proxy_logger.debug( - "ToolPolicyGuardrail: fetched %d policies from DB (cache hits: %d)", - len(cache_misses), - len(tool_names) - len(cache_misses), - ) - - return result + return await get_tool_policies_cached( + tool_names=tool_names, + cache=user_api_key_cache, + prisma_client=prisma_client, + team_id=team_id, + key_hash=key_hash, + ) diff --git a/litellm/proxy/guardrails/tool_name_extraction.py b/litellm/proxy/guardrails/tool_name_extraction.py new file mode 100644 index 00000000000..db24fa2277c --- /dev/null +++ b/litellm/proxy/guardrails/tool_name_extraction.py @@ -0,0 +1,85 @@ +""" +Extract tool names from request body by route/call type. + +Used by auth (check_tools_allowlist) and ToolPolicyGuardrail so tool-format +knowledge lives in one place. Uses guardrail translation handlers where available, +with standalone extractors for generate_content and MCP. +""" + +from typing import Any, Dict, List + +from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route +from litellm.llms import load_guardrail_translation_mappings +from litellm.types.utils import CallTypes + +# Call types that have no guardrail translation handler; we use standalone extractors +STANDALONE_EXTRACTORS: Dict[str, Any] = {} + + +def _extract_generate_content_tool_names(data: dict) -> List[str]: + """Google generateContent: tools[].functionDeclarations[].name""" + names: List[str] = [] + for tool in data.get("tools") or []: + if not isinstance(tool, dict): + continue + for decl in tool.get("functionDeclarations") or []: + if isinstance(decl, dict) and decl.get("name"): + names.append(str(decl["name"])) + return names + + +def _extract_mcp_tool_names(data: dict) -> List[str]: + """MCP call_tool: name or mcp_tool_name in body""" + names: List[str] = [] + name = data.get("name") or data.get("mcp_tool_name") + if name: + names.append(str(name)) + return names + + +def _register_standalone_extractors() -> None: + if STANDALONE_EXTRACTORS: + return + STANDALONE_EXTRACTORS[CallTypes.generate_content.value] = _extract_generate_content_tool_names + STANDALONE_EXTRACTORS[CallTypes.agenerate_content.value] = _extract_generate_content_tool_names + STANDALONE_EXTRACTORS[CallTypes.call_mcp_tool.value] = _extract_mcp_tool_names + + +# Tool-capable call types (routes that can send tools in the request) +TOOL_CAPABLE_CALL_TYPES = frozenset({ + CallTypes.completion.value, + CallTypes.acompletion.value, + CallTypes.responses.value, + CallTypes.aresponses.value, + CallTypes.anthropic_messages.value, + CallTypes.generate_content.value, + CallTypes.agenerate_content.value, + CallTypes.call_mcp_tool.value, +}) + + +def extract_request_tool_names(route: str, data: dict) -> List[str]: + """ + Extract tool names from the request body for the given route. + Uses guardrail translation handlers when available, else standalone extractors + for generate_content and MCP. Returns [] for non-tool-capable routes or when + no tools are present. + """ + call_types = get_call_types_for_route(route) + if not call_types: + return [] + _register_standalone_extractors() + mappings = load_guardrail_translation_mappings() + for call_type in call_types: + if not isinstance(call_type, CallTypes): + continue + if call_type.value not in TOOL_CAPABLE_CALL_TYPES: + continue + if call_type.value in STANDALONE_EXTRACTORS: + return STANDALONE_EXTRACTORS[call_type.value](data) + handler_cls = mappings.get(call_type) + if handler_cls is not None: + names = handler_cls().extract_request_tool_names(data) + if names: + return names + return [] diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index d2312a00c3b..c3d58df83ea 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -10,16 +10,12 @@ import litellm from litellm._logging import verbose_logger, verbose_proxy_logger from litellm._service_logger import ServiceLogging from litellm.litellm_core_utils.safe_json_loads import safe_json_loads -from litellm.proxy._types import ( - AddTeamCallback, - CommonProxyErrors, - LitellmDataForBackendLLMCall, - LitellmUserRoles, - SpecialHeaders, - TeamCallbackMetadata, - UserAPIKeyAuth, -) -from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers +from litellm.proxy._types import (AddTeamCallback, CommonProxyErrors, + LitellmDataForBackendLLMCall, + LitellmUserRoles, SpecialHeaders, + TeamCallbackMetadata, UserAPIKeyAuth) +from litellm.proxy.common_utils.http_parsing_utils import \ + _safe_get_request_headers # Cache special headers as a frozenset for O(1) lookup performance _SPECIAL_HEADERS_CACHE = frozenset( @@ -28,12 +24,9 @@ _SPECIAL_HEADERS_CACHE = frozenset( from litellm.router import Router from litellm.types.llms.anthropic import ANTHROPIC_API_HEADERS from litellm.types.services import ServiceTypes -from litellm.types.utils import ( - LlmProviders, - ProviderSpecificHeader, - StandardLoggingUserAPIKeyMetadata, - SupportedCacheControls, -) +from litellm.types.utils import (LlmProviders, ProviderSpecificHeader, + StandardLoggingUserAPIKeyMetadata, + SupportedCacheControls) service_logger_obj = ServiceLogging() # used for tracking latency on OTEL @@ -667,8 +660,7 @@ class LiteLLMProxyRequestSetup: return data from litellm.proxy._types import ( LiteLLM_ManagementEndpoint_MetadataFields, - LiteLLM_ManagementEndpoint_MetadataFields_Premium, - ) + LiteLLM_ManagementEndpoint_MetadataFields_Premium) # ignore any special fields added_metadata = {} @@ -1058,6 +1050,9 @@ async def add_litellm_data_to_request( # noqa: PLR0915 ] = user_api_key_dict.user_max_budget data[_metadata_variable_name]["user_api_key_metadata"] = user_api_key_dict.metadata + data[_metadata_variable_name]["user_api_key_team_metadata"] = ( + user_api_key_dict.team_metadata + ) data[_metadata_variable_name]["headers"] = _headers data[_metadata_variable_name]["endpoint"] = str(request.url) @@ -1501,7 +1496,8 @@ async def move_guardrails_to_metadata( # Only check policy engine if no local config (avoid import + registry lookup) if not (has_key_config or has_team_config or has_request_config): - from litellm.proxy.policy_engine.policy_registry import get_policy_registry + from litellm.proxy.policy_engine.policy_registry import \ + get_policy_registry if not get_policy_registry().is_initialized(): # Nothing configured anywhere - clean up request body fields and return @@ -1565,14 +1561,16 @@ async def move_guardrails_to_metadata( def _is_policy_version_id(s: str) -> bool: """Return True if string is a policy version ID (starts with policy_ prefix).""" - from litellm.proxy.policy_engine.policy_registry import POLICY_VERSION_ID_PREFIX + from litellm.proxy.policy_engine.policy_registry import \ + POLICY_VERSION_ID_PREFIX return isinstance(s, str) and s.startswith(POLICY_VERSION_ID_PREFIX) def _extract_policy_id(s: str) -> Optional[str]: """Extract raw UUID from policy_ string, or None if not a valid version ID.""" - from litellm.proxy.policy_engine.policy_registry import POLICY_VERSION_ID_PREFIX + from litellm.proxy.policy_engine.policy_registry import \ + POLICY_VERSION_ID_PREFIX if not _is_policy_version_id(s): return None @@ -1593,10 +1591,9 @@ def _match_and_track_policies( """ from litellm._logging import verbose_proxy_logger from litellm.proxy.common_utils.callback_utils import ( - add_policy_sources_to_metadata, - add_policy_to_applied_policies_header, - ) - from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry + add_policy_sources_to_metadata, add_policy_to_applied_policies_header) + from litellm.proxy.policy_engine.attachment_registry import \ + get_attachment_registry from litellm.proxy.policy_engine.policy_matcher import PolicyMatcher # Get matching policies via attachments (with match reasons for attribution) @@ -1741,7 +1738,8 @@ async def add_guardrails_from_policy_engine( user_api_key_dict: The user's API key authentication info """ from litellm._logging import verbose_proxy_logger - from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body + from litellm.proxy.common_utils.http_parsing_utils import \ + get_tags_from_request_body from litellm.proxy.policy_engine.policy_registry import get_policy_registry from litellm.types.proxy.policy_engine import PolicyMatchContext diff --git a/litellm/proxy/management_endpoints/tool_management_endpoints.py b/litellm/proxy/management_endpoints/tool_management_endpoints.py index 89880c9a4ec..bc8fc9165bb 100644 --- a/litellm/proxy/management_endpoints/tool_management_endpoints.py +++ b/litellm/proxy/management_endpoints/tool_management_endpoints.py @@ -10,7 +10,7 @@ POST /v1/tool/policy - Update the call_policy for a tool from typing import Optional -from fastapi import APIRouter, Depends, HTTPException +from fastapi import APIRouter, Depends, HTTPException, Query from litellm._logging import verbose_proxy_logger from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth @@ -18,6 +18,7 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.types.tool_management import ( LiteLLM_ToolTableRow, ToolCallPolicy, + ToolDetailResponse, ToolListResponse, ToolPolicyUpdateRequest, ToolPolicyUpdateResponse, @@ -58,6 +59,48 @@ async def list_tools( raise HTTPException(status_code=500, detail=str(e)) +@router.get( + "/v1/tool/{tool_name:path}/detail", + tags=["tool management"], + dependencies=[Depends(user_api_key_auth)], + response_model=ToolDetailResponse, +) +async def get_tool_detail( + tool_name: str, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Get a single tool with its policy overrides (for UI detail view). + + Parameters: + - tool_name: The tool name (supports namespaced names with slashes) + """ + from litellm.proxy.db.tool_registry_writer import get_tool as db_get_tool + from litellm.proxy.db.tool_registry_writer import list_overrides_for_tool + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException( + status_code=500, detail=CommonProxyErrors.db_not_connected_error.value + ) + + try: + tool = await db_get_tool(prisma_client=prisma_client, tool_name=tool_name) + if tool is None: + raise HTTPException( + status_code=404, detail=f"Tool '{tool_name}' not found" + ) + overrides = await list_overrides_for_tool( + prisma_client=prisma_client, tool_name=tool_name + ) + return ToolDetailResponse(tool=tool, overrides=overrides) + except HTTPException: + raise + except Exception as e: + verbose_proxy_logger.exception("Error getting tool detail: %s", e) + raise HTTPException(status_code=500, detail=str(e)) + + @router.get( "/v1/tool/{tool_name:path}", tags=["tool management"], @@ -107,18 +150,23 @@ async def update_tool_policy( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ - Set the call policy for a tool. + Set the call policy for a tool (global) or for a specific team/key (override). Parameters: - tool_name: str - The tool to update - call_policy: "trusted" | "untrusted" | "dual_llm" | "blocked" + - team_id: optional - if set, create/update override for this team only + - key_hash: optional - if set, create/update override for this key only + - key_alias: optional - human-readable key alias for UI - Setting a tool to "blocked" will cause the ToolPolicyGuardrail to remove - that tool_call from LLM responses before returning them to the client. + If both team_id and key_hash are omitted, updates the global tool policy. + Setting a tool to "blocked" will cause the ToolPolicyGuardrail to reject + that tool_call for the relevant scope. """ from litellm.proxy.db.tool_registry_writer import ( update_tool_policy as db_update_tool_policy, ) + from litellm.proxy.db.tool_registry_writer import upsert_tool_policy_override from litellm.proxy.proxy_server import prisma_client if prisma_client is None: @@ -127,6 +175,28 @@ async def update_tool_policy( ) try: + if data.team_id is not None or data.key_hash is not None: + override = await upsert_tool_policy_override( + prisma_client=prisma_client, + tool_name=data.tool_name, + call_policy=data.call_policy, + team_id=data.team_id, + key_hash=data.key_hash, + key_alias=data.key_alias, + updated_by=user_api_key_dict.user_id, + ) + if override is None: + raise HTTPException( + status_code=500, + detail=f"Failed to update policy override for tool '{data.tool_name}'", + ) + return ToolPolicyUpdateResponse( + tool_name=override.tool_name, + call_policy=override.call_policy, + updated=True, + team_id=override.team_id, + key_hash=override.key_hash, + ) updated = await db_update_tool_policy( prisma_client=prisma_client, tool_name=data.tool_name, @@ -135,7 +205,8 @@ async def update_tool_policy( ) if updated is None: raise HTTPException( - status_code=500, detail=f"Failed to update policy for tool '{data.tool_name}'" + status_code=500, + detail=f"Failed to update policy for tool '{data.tool_name}'", ) return ToolPolicyUpdateResponse( tool_name=updated.tool_name, @@ -147,3 +218,50 @@ async def update_tool_policy( except Exception as e: verbose_proxy_logger.exception("Error updating tool policy: %s", e) raise HTTPException(status_code=500, detail=str(e)) + + +@router.delete( + "/v1/tool/{tool_name:path}/overrides", + tags=["tool management"], + dependencies=[Depends(user_api_key_auth)], +) +async def delete_tool_policy_override( + tool_name: str, + team_id: Optional[str] = Query(None, description="Team ID of the override to remove"), + key_hash: Optional[str] = Query(None, description="Key hash of the override to remove"), + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Remove a policy override for a tool. Specify the override by team_id and/or key_hash + (must match the override that was created; use empty string for unscoped dimension). + """ + from litellm.proxy.db.tool_registry_writer import delete_tool_policy_override + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException( + status_code=500, detail=CommonProxyErrors.db_not_connected_error.value + ) + if team_id is None and key_hash is None: + raise HTTPException( + status_code=400, + detail="At least one of team_id or key_hash is required to identify the override", + ) + try: + deleted = await delete_tool_policy_override( + prisma_client=prisma_client, + tool_name=tool_name, + team_id=team_id, + key_hash=key_hash, + ) + if not deleted: + raise HTTPException( + status_code=404, + detail=f"No override found for tool '{tool_name}' with the given scope", + ) + return {"deleted": True, "tool_name": tool_name} + except HTTPException: + raise + except Exception as e: + verbose_proxy_logger.exception("Error deleting tool policy override: %s", e) + raise HTTPException(status_code=500, detail=str(e)) diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 440c9c1d829..e48c0fe3027 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -1076,6 +1076,25 @@ model LiteLLM_ToolTable { @@index([team_id]) } +// Per-(tool, team/key) policy overrides. When present, override replaces global tool policy for that scope. +model LiteLLM_ToolPolicyOverrideTable { + override_id String @id @default(uuid()) + tool_name String + team_id String @default("") // "" = not scoped to team; non-empty = override for this team only + key_hash String @default("") // "" = not scoped to key; non-empty = override for this key only + call_policy String @default("blocked") + key_alias String? // human-readable key alias for UI + created_at DateTime @default(now()) + created_by String? + updated_at DateTime @default(now()) @updatedAt + updated_by String? + + @@unique([tool_name, team_id, key_hash]) + @@index([tool_name]) + @@index([team_id]) + @@index([key_hash]) +} + //Unified Access Groups table for storing unified access groups model LiteLLM_AccessGroupTable { access_group_id String @id @default(uuid()) diff --git a/litellm/types/tool_management.py b/litellm/types/tool_management.py index 8704ff27759..90fc4a6ec7d 100644 --- a/litellm/types/tool_management.py +++ b/litellm/types/tool_management.py @@ -5,7 +5,7 @@ Pydantic models for Tool Policy management endpoints. from datetime import datetime from typing import Dict, List, Literal, Optional -from pydantic import BaseModel +from pydantic import BaseModel, Field ToolCallPolicy = Literal["trusted", "untrusted", "dual_llm", "blocked"] @@ -34,9 +34,30 @@ class ToolListResponse(BaseModel): class ToolPolicyUpdateRequest(BaseModel): tool_name: str call_policy: ToolCallPolicy + team_id: Optional[str] = None # if set, create/update override for this team + key_hash: Optional[str] = None # if set, create/update override for this key + key_alias: Optional[str] = None # human-readable key alias for UI class ToolPolicyUpdateResponse(BaseModel): tool_name: str call_policy: ToolCallPolicy updated: bool + team_id: Optional[str] = None + key_hash: Optional[str] = None + + +class ToolPolicyOverrideRow(BaseModel): + override_id: str + tool_name: str + team_id: Optional[str] = None + key_hash: Optional[str] = None + call_policy: ToolCallPolicy + key_alias: Optional[str] = None + created_at: Optional[datetime] = None + updated_at: Optional[datetime] = None + + +class ToolDetailResponse(BaseModel): + tool: LiteLLM_ToolTableRow + overrides: List[ToolPolicyOverrideRow] = Field(default_factory=list) diff --git a/scripts/test_tool_allowlist_script.py b/scripts/test_tool_allowlist_script.py new file mode 100644 index 00000000000..75a50d09b84 --- /dev/null +++ b/scripts/test_tool_allowlist_script.py @@ -0,0 +1,116 @@ +#!/usr/bin/env python3 +""" +Standalone script to test tool allowlist enforcement and tool name extraction. + +Run from repo root: + poetry run python scripts/test_tool_allowlist_script.py + +Or run the unit tests: + poetry run pytest tests/test_litellm/proxy/test_tools_allowlist_enforcement.py -v +""" + +import asyncio +import sys +from pathlib import Path + +# Ensure repo root is on path +repo_root = Path(__file__).resolve().parent.parent +if str(repo_root) not in sys.path: + sys.path.insert(0, str(repo_root)) + + +def test_extraction(): + """Test extract_request_tool_names for each API shape.""" + from litellm.proxy.guardrails.tool_name_extraction import extract_request_tool_names + + cases = [ + ("OpenAI chat tools", "/v1/chat/completions", {"tools": [{"type": "function", "function": {"name": "get_weather"}}]}), + ("OpenAI chat functions", "/v1/chat/completions", {"functions": [{"name": "run_sql"}]}), + ("OpenAI responses function", "/v1/responses", {"tools": [{"type": "function", "name": "get_current_weather"}]}), + ("OpenAI responses MCP", "/v1/responses", {"tools": [{"type": "mcp", "server_label": "dmcp"}]}), + ("Anthropic", "/v1/messages", {"tools": [{"name": "get_weather"}, {"name": "run_sql"}]}), + ("Google generateContent", "/generate_content", {"tools": [{"functionDeclarations": [{"name": "schedule_meeting"}]}]}), + ("MCP call_tool", "/mcp/call_tool", {"name": "my_tool", "arguments": {}}), + ("Non-tool route", "/v1/embeddings", {"tools": [{"type": "function", "function": {"name": "x"}}]}), + ] + print("=== extract_request_tool_names(route, data) ===\n") + for label, route, data in cases: + names = extract_request_tool_names(route, data) + print(f" {label}: {names}") + print() + + +async def test_check_tools_allowlist(): + """Test check_tools_allowlist with mock tokens.""" + from litellm.proxy._types import ProxyErrorTypes, ProxyException, UserAPIKeyAuth + from litellm.proxy.auth.auth_checks import check_tools_allowlist + + def token(metadata=None, team_metadata=None): + return UserAPIKeyAuth( + api_key="test-key", + user_id="user", + team_id="team", + org_id=None, + models=["*"], + metadata=metadata or {}, + team_metadata=team_metadata or {}, + ) + + print("=== check_tools_allowlist (auth) ===\n") + + # No allowlist -> pass + await check_tools_allowlist( + request_body={"tools": [{"type": "function", "function": {"name": "get_weather"}}]}, + valid_token=token(), + team_object=None, + route="/v1/chat/completions", + ) + print(" No allowlist, body has tools: PASS") + + # Allowed tool -> pass + await check_tools_allowlist( + request_body={"tools": [{"type": "function", "function": {"name": "get_weather"}}]}, + valid_token=token(metadata={"allowed_tools": ["get_weather"]}), + team_object=None, + route="/v1/chat/completions", + ) + print(" allowed_tools=['get_weather'], body has get_weather: PASS") + + # Disallowed tool -> raise + try: + await check_tools_allowlist( + request_body={"tools": [{"type": "function", "function": {"name": "get_weather"}}]}, + valid_token=token(metadata={"allowed_tools": ["other_tool"]}), + team_object=None, + route="/v1/chat/completions", + ) + print(" DISALLOWED: expected ProxyException") + except ProxyException as e: + if e.type == ProxyErrorTypes.tool_access_denied: + print(" allowed_tools=['other_tool'], body has get_weather: PASS (raised tool_access_denied)") + else: + print(f" Unexpected ProxyException type: {e.type}") + except Exception as e: + print(f" Unexpected: {e}") + + # Team allowlist when key empty + await check_tools_allowlist( + request_body={"tools": [{"type": "function", "function": {"name": "get_weather"}}]}, + valid_token=token(team_metadata={"allowed_tools": ["get_weather"]}), + team_object=None, + route="/v1/chat/completions", + ) + print(" team_metadata.allowed_tools=['get_weather']: PASS") + print() + + +def main(): + print("Tool allowlist / tool name extraction – script checks\n") + test_extraction() + asyncio.run(test_check_tools_allowlist()) + print("Done. For full unit tests run:") + print(" poetry run pytest tests/test_litellm/proxy/test_tools_allowlist_enforcement.py -v") + + +if __name__ == "__main__": + main() diff --git a/tests/test_litellm/proxy/test_tools_allowlist_enforcement.py b/tests/test_litellm/proxy/test_tools_allowlist_enforcement.py index 09d56be30ba..4adc5acde8b 100644 --- a/tests/test_litellm/proxy/test_tools_allowlist_enforcement.py +++ b/tests/test_litellm/proxy/test_tools_allowlist_enforcement.py @@ -1,548 +1,200 @@ """ -Tests for tool allowlist enforcement by team/key (metadata.allowed_tools). +Tests for tool allowlist enforcement (key/team metadata.allowed_tools). -No implementation yet; these tests define expected behavior. When check_tools_allowlist -is implemented in common_checks, disallowed-tool tests should raise; allowed and -no-allowlist tests should pass. +Covers: +- check_tools_allowlist: allowed, disallowed, no allowlist, non-tool routes +- extract_request_tool_names: OpenAI chat, responses, Anthropic, generate_content, MCP """ from unittest.mock import AsyncMock, MagicMock, patch import pytest -from litellm.proxy._types import ProxyException, UserAPIKeyAuth -from litellm.proxy.auth.auth_checks import common_checks +from litellm.proxy._types import (ProxyErrorTypes, ProxyException, + UserAPIKeyAuth) +from litellm.proxy.auth.auth_checks import check_tools_allowlist +from litellm.proxy.guardrails.tool_name_extraction import ( + TOOL_CAPABLE_CALL_TYPES, extract_request_tool_names) -class MockRequest: - """Mock request with method attribute.""" - - def __init__(self, method: str = "POST"): - self.method = method - - -def get_mock_user_token(metadata=None, team_metadata=None) -> UserAPIKeyAuth: - """Build UserAPIKeyAuth with optional metadata and team_metadata for allowlist.""" - kwargs = { - "api_key": "test-key", - "user_id": "test-user", - "team_id": "test-team", - "org_id": "test-org", - "models": ["*"], - "metadata": metadata or {}, - } - if team_metadata is not None: - kwargs["team_metadata"] = team_metadata - return UserAPIKeyAuth(**kwargs) - - -def _tools_allowlist_patches(): - """Patches so only tool-allowlist behavior is under test; heavy/DB parts no-op.""" - p1 = patch( - "litellm.proxy.auth.auth_checks._is_api_route_allowed", - new_callable=AsyncMock, - return_value=True, +def _token(metadata=None, team_metadata=None): + return UserAPIKeyAuth( + api_key="test-key", + user_id="user", + team_id="team", + org_id=None, + models=["*"], + metadata=metadata or {}, + team_metadata=team_metadata or {}, ) - p2 = patch( - "litellm.proxy.auth.auth_checks.vector_store_access_check", - new_callable=AsyncMock, - return_value=None, - ) - p3 = patch( - "litellm.proxy.auth.auth_checks._run_project_checks", - new_callable=AsyncMock, - return_value=None, - ) - return p1, p2, p3 -class TestOpenAIChatCompletionsToolsAllowlist: - """Tool allowlist enforcement for /v1/chat/completions.""" +class TestExtractRequestToolNames: + """Test tool name extraction per API format.""" - @pytest.mark.asyncio - async def test_chat_completions_allowed_tool_passes(self): - """Request with tools in allowed_tools passes.""" - route = "/v1/chat/completions" - request_body = { - "model": "gpt-4", - "messages": [{"role": "user", "content": "Hi"}], - "tools": [{"type": "function", "function": {"name": "get_weather"}}], - } - token = get_mock_user_token(metadata={"allowed_tools": ["get_weather"]}) - request = MockRequest("POST") - - p1, p2, p3 = _tools_allowlist_patches() - with p1, p2, p3: - result = await common_checks( - request_body=request_body, - team_object=None, - user_object=None, - end_user_object=None, - global_proxy_spend=None, - general_settings={}, - route=route, - llm_router=None, - proxy_logging_obj=MagicMock(), - valid_token=token, - request=request, - ) - assert result is True - - @pytest.mark.asyncio - async def test_chat_completions_disallowed_tool_raises(self): - """Request with tool not in allowed_tools raises.""" - route = "/v1/chat/completions" - request_body = { - "model": "gpt-4", - "messages": [{"role": "user", "content": "Hi"}], - "tools": [{"type": "function", "function": {"name": "get_weather"}}], - } - token = get_mock_user_token(metadata={"allowed_tools": ["other_tool"]}) - request = MockRequest("POST") - - p1, p2, p3 = _tools_allowlist_patches() - with p1, p2, p3: - with pytest.raises((Exception, ProxyException)) as exc_info: - await common_checks( - request_body=request_body, - team_object=None, - user_object=None, - end_user_object=None, - global_proxy_spend=None, - general_settings={}, - route=route, - llm_router=None, - proxy_logging_obj=MagicMock(), - valid_token=token, - request=request, - ) - msg = str(exc_info.value).lower() - assert "tool" in msg or "allowed" in msg - - @pytest.mark.asyncio - async def test_chat_completions_legacy_functions_allowed(self): - """Legacy 'functions' (no tools) with allowed name passes.""" - route = "/v1/chat/completions" - request_body = { - "model": "gpt-4", - "messages": [{"role": "user", "content": "Hi"}], - "functions": [{"name": "get_weather"}], - } - token = get_mock_user_token(metadata={"allowed_tools": ["get_weather"]}) - request = MockRequest("POST") - - p1, p2, p3 = _tools_allowlist_patches() - with p1, p2, p3: - result = await common_checks( - request_body=request_body, - team_object=None, - user_object=None, - end_user_object=None, - global_proxy_spend=None, - general_settings={}, - route=route, - llm_router=None, - proxy_logging_obj=MagicMock(), - valid_token=token, - request=request, - ) - assert result is True - - @pytest.mark.asyncio - async def test_chat_completions_no_allowlist_passes(self): - """Request with tools but no metadata.allowed_tools / team_metadata passes.""" - route = "/v1/chat/completions" - request_body = { - "model": "gpt-4", - "messages": [{"role": "user", "content": "Hi"}], - "tools": [{"type": "function", "function": {"name": "get_weather"}}], - } - token = get_mock_user_token(metadata={}, team_metadata={}) - request = MockRequest("POST") - - p1, p2, p3 = _tools_allowlist_patches() - with p1, p2, p3: - result = await common_checks( - request_body=request_body, - team_object=None, - user_object=None, - end_user_object=None, - global_proxy_spend=None, - general_settings={}, - route=route, - llm_router=None, - proxy_logging_obj=MagicMock(), - valid_token=token, - request=request, - ) - assert result is True - - -class TestOpenAIResponsesAPIToolsAllowlist: - """Tool allowlist enforcement for /v1/responses.""" - - @pytest.mark.asyncio - async def test_responses_function_tool_allowed_passes(self): - """Responses request with function tool in allowed_tools passes.""" - route = "/v1/responses" - request_body = { - "model": "gpt-4", - "input": "What is the weather?", + def test_openai_chat_tools(self): + data = { "tools": [ - { - "type": "function", - "name": "get_current_weather", - "description": "Get current weather", - "parameters": {"type": "object"}, - } - ], + {"type": "function", "function": {"name": "get_weather"}}, + {"type": "function", "function": {"name": "run_sql"}}, + ] } - token = get_mock_user_token(metadata={"allowed_tools": ["get_current_weather"]}) - request = MockRequest("POST") + assert extract_request_tool_names("/v1/chat/completions", data) == [ + "get_weather", + "run_sql", + ] - p1, p2, p3 = _tools_allowlist_patches() - with p1, p2, p3: - result = await common_checks( - request_body=request_body, - team_object=None, - user_object=None, - end_user_object=None, - global_proxy_spend=None, - general_settings={}, - route=route, - llm_router=None, - proxy_logging_obj=MagicMock(), - valid_token=token, - request=request, - ) - assert result is True + def test_openai_chat_functions_legacy(self): + data = {"functions": [{"name": "get_weather"}, {"name": "run_sql"}]} + assert extract_request_tool_names("/v1/chat/completions", data) == [ + "get_weather", + "run_sql", + ] - @pytest.mark.asyncio - async def test_responses_function_tool_disallowed_raises(self): - """Responses request with function tool not in allowed_tools raises.""" - route = "/v1/responses" - request_body = { - "model": "gpt-4", - "input": "What is the weather?", + def test_openai_responses_function_tools(self): + data = { "tools": [ - { - "type": "function", - "name": "get_current_weather", - "description": "Get current weather", - "parameters": {"type": "object"}, - } - ], + {"type": "function", "name": "get_current_weather", "description": "x"}, + ] } - token = get_mock_user_token(metadata={"allowed_tools": ["other"]}) - request = MockRequest("POST") + assert extract_request_tool_names("/v1/responses", data) == [ + "get_current_weather" + ] - p1, p2, p3 = _tools_allowlist_patches() - with p1, p2, p3: - with pytest.raises((Exception, ProxyException)) as exc_info: - await common_checks( - request_body=request_body, - team_object=None, - user_object=None, - end_user_object=None, - global_proxy_spend=None, - general_settings={}, - route=route, - llm_router=None, - proxy_logging_obj=MagicMock(), - valid_token=token, - request=request, - ) - msg = str(exc_info.value).lower() - assert "tool" in msg or "allowed" in msg - - @pytest.mark.asyncio - async def test_responses_mcp_server_allowed_passes(self): - """Responses request with MCP server in allowed_tools passes.""" - route = "/v1/responses" - request_body = { - "model": "gpt-4", - "input": "Hi", + def test_openai_responses_mcp_tools(self): + data = { "tools": [ - { - "type": "mcp", - "server_label": "dmcp", - "server_description": "Example MCP server", - "server_url": "https://example.com", - "require_approval": "never", - } - ], + {"type": "mcp", "server_label": "dmcp", "server_url": "http://x"}, + ] } - token = get_mock_user_token(metadata={"allowed_tools": ["dmcp"]}) - request = MockRequest("POST") + assert extract_request_tool_names("/v1/responses", data) == ["dmcp"] - p1, p2, p3 = _tools_allowlist_patches() - with p1, p2, p3: - result = await common_checks( - request_body=request_body, - team_object=None, - user_object=None, - end_user_object=None, - global_proxy_spend=None, - general_settings={}, - route=route, - llm_router=None, - proxy_logging_obj=MagicMock(), - valid_token=token, - request=request, - ) - assert result is True + def test_anthropic_tools(self): + data = {"tools": [{"name": "get_weather"}, {"name": "run_sql"}]} + assert extract_request_tool_names("/v1/messages", data) == [ + "get_weather", + "run_sql", + ] - @pytest.mark.asyncio - async def test_responses_mcp_server_disallowed_raises(self): - """Responses request with MCP server not in allowed_tools raises.""" - route = "/v1/responses" - request_body = { - "model": "gpt-4", - "input": "Hi", - "tools": [ - { - "type": "mcp", - "server_label": "dmcp", - "server_description": "Example MCP server", - "server_url": "https://example.com", - "require_approval": "never", - } - ], - } - token = get_mock_user_token(metadata={"allowed_tools": ["other"]}) - request = MockRequest("POST") - - p1, p2, p3 = _tools_allowlist_patches() - with p1, p2, p3: - with pytest.raises((Exception, ProxyException)): - await common_checks( - request_body=request_body, - team_object=None, - user_object=None, - end_user_object=None, - global_proxy_spend=None, - general_settings={}, - route=route, - llm_router=None, - proxy_logging_obj=MagicMock(), - valid_token=token, - request=request, - ) - - -class TestAnthropicMessagesToolsAllowlist: - """Tool allowlist enforcement for Anthropic /v1/messages.""" - - @pytest.mark.asyncio - async def test_anthropic_allowed_tool_passes(self): - """Request with Anthropic-style tools in allowed_tools passes.""" - route = "/v1/messages" - request_body = { - "model": "claude-3-5-sonnet-20241022", - "max_tokens": 1024, - "messages": [{"role": "user", "content": "Hi"}], - "tools": [{"name": "get_weather", "description": "Get weather"}], - } - token = get_mock_user_token(metadata={"allowed_tools": ["get_weather"]}) - request = MockRequest("POST") - - p1, p2, p3 = _tools_allowlist_patches() - with p1, p2, p3: - result = await common_checks( - request_body=request_body, - team_object=None, - user_object=None, - end_user_object=None, - global_proxy_spend=None, - general_settings={}, - route=route, - llm_router=None, - proxy_logging_obj=MagicMock(), - valid_token=token, - request=request, - ) - assert result is True - - @pytest.mark.asyncio - async def test_anthropic_disallowed_tool_raises(self): - """Request with Anthropic-style tool not in allowed_tools raises.""" - route = "/v1/messages" - request_body = { - "model": "claude-3-5-sonnet-20241022", - "max_tokens": 1024, - "messages": [{"role": "user", "content": "Hi"}], - "tools": [{"name": "get_weather", "description": "Get weather"}], - } - token = get_mock_user_token(metadata={"allowed_tools": ["other_tool"]}) - request = MockRequest("POST") - - p1, p2, p3 = _tools_allowlist_patches() - with p1, p2, p3: - with pytest.raises((Exception, ProxyException)) as exc_info: - await common_checks( - request_body=request_body, - team_object=None, - user_object=None, - end_user_object=None, - global_proxy_spend=None, - general_settings={}, - route=route, - llm_router=None, - proxy_logging_obj=MagicMock(), - valid_token=token, - request=request, - ) - msg = str(exc_info.value).lower() - assert "tool" in msg or "allowed" in msg - - -class TestGoogleGenerateContentToolsAllowlist: - """Tool allowlist enforcement for Google generateContent.""" - - @pytest.mark.asyncio - async def test_google_allowed_tool_passes(self): - """Request with tools[].functionDeclarations[].name in allowed_tools passes.""" - route = "/v1beta/models/gemini-3-flash-preview:generateContent" - request_body = { - "contents": [ - {"role": "user", "parts": [{"text": "Schedule a meeting"}]} - ], + def test_generate_content_tools(self): + data = { "tools": [ { "functionDeclarations": [ - { - "name": "schedule_meeting", - "description": "Schedules a meeting", - "parameters": {"type": "object", "properties": {}}, - } + {"name": "schedule_meeting", "description": "x"}, ] - } - ], + }, + ] } - token = get_mock_user_token(metadata={"allowed_tools": ["schedule_meeting"]}) - request = MockRequest("POST") + assert extract_request_tool_names("/generate_content", data) == [ + "schedule_meeting" + ] - p1, p2, p3 = _tools_allowlist_patches() - with p1, p2, p3: - result = await common_checks( - request_body=request_body, - team_object=None, - user_object=None, - end_user_object=None, - global_proxy_spend=None, - general_settings={}, - route=route, - llm_router=None, - proxy_logging_obj=MagicMock(), + def test_mcp_call_tool_name(self): + data = {"name": "my_tool", "arguments": {}} + assert extract_request_tool_names("/mcp/call_tool", data) == ["my_tool"] + + def test_mcp_call_tool_mcp_tool_name(self): + data = {"mcp_tool_name": "other_tool"} + assert extract_request_tool_names("/mcp/call_tool", data) == ["other_tool"] + + def test_non_tool_route_returns_empty(self): + data = {"tools": [{"type": "function", "function": {"name": "x"}}]} + assert extract_request_tool_names("/v1/embeddings", data) == [] + + +class TestCheckToolsAllowlist: + """Test allowlist enforcement in auth (no DB in hot path).""" + + @pytest.mark.asyncio + async def test_no_allowlist_passes(self): + token = _token(metadata={}, team_metadata={}) + body = { + "tools": [{"type": "function", "function": {"name": "get_weather"}}] + } + await check_tools_allowlist( + request_body=body, + valid_token=token, + team_object=None, + route="/v1/chat/completions", + ) + + @pytest.mark.asyncio + async def test_allowed_tool_passes(self): + token = _token(metadata={"allowed_tools": ["get_weather"]}) + body = { + "tools": [{"type": "function", "function": {"name": "get_weather"}}] + } + await check_tools_allowlist( + request_body=body, + valid_token=token, + team_object=None, + route="/v1/chat/completions", + ) + + @pytest.mark.asyncio + async def test_disallowed_tool_raises(self): + token = _token(metadata={"allowed_tools": ["other_tool"]}) + body = { + "tools": [{"type": "function", "function": {"name": "get_weather"}}] + } + with pytest.raises(ProxyException) as exc_info: + await check_tools_allowlist( + request_body=body, valid_token=token, - request=request, - ) - assert result is True - - @pytest.mark.asyncio - async def test_google_disallowed_tool_raises(self): - """Request with tools[].functionDeclarations[].name not in allowed_tools raises.""" - route = "/v1beta/models/gemini-3-flash-preview:generateContent" - request_body = { - "contents": [ - {"role": "user", "parts": [{"text": "Schedule a meeting"}]} - ], - "tools": [ - { - "functionDeclarations": [ - { - "name": "schedule_meeting", - "description": "Schedules a meeting", - "parameters": {"type": "object", "properties": {}}, - } - ] - } - ], - } - token = get_mock_user_token(metadata={"allowed_tools": ["other_tool"]}) - request = MockRequest("POST") - - p1, p2, p3 = _tools_allowlist_patches() - with p1, p2, p3: - with pytest.raises((Exception, ProxyException)) as exc_info: - await common_checks( - request_body=request_body, - team_object=None, - user_object=None, - end_user_object=None, - global_proxy_spend=None, - general_settings={}, - route=route, - llm_router=None, - proxy_logging_obj=MagicMock(), - valid_token=token, - request=request, - ) - msg = str(exc_info.value).lower() - assert "tool" in msg or "allowed" in msg - - -# MCP REST tools/call body shape: server_id, name (tool name), arguments. -# See litellm/proxy/_experimental/mcp_server/rest_endpoints.py call_tool_rest_api. -# The exact field for tool name in the request body should match the implementation. -MCP_TOOL_CALL_BODY_ALLOWED = { - "server_id": "srv", - "name": "roll_dice", - "arguments": {}, -} - - -class TestMCPToolCallToolsAllowlist: - """Test that MCP tool call routes (/mcp/tools/call, /mcp-rest/tools/call) enforce token allowed_tools via common_checks.""" - - @pytest.mark.asyncio - async def test_mcp_tool_call_allowed_passes(self): - """Route /mcp-rest/tools/call with tool in token allowed_tools passes common_checks.""" - request = MockRequest("POST") - request_body = dict(MCP_TOOL_CALL_BODY_ALLOWED) - valid_token = get_mock_user_token(metadata={"allowed_tools": ["roll_dice"]}) - - p1, p2, p3 = _tools_allowlist_patches() - with p1, p2, p3: - result = await common_checks( - request_body=request_body, team_object=None, - user_object=None, - end_user_object=None, - global_proxy_spend=None, - general_settings={}, - route="/mcp-rest/tools/call", - llm_router=None, - proxy_logging_obj=MagicMock(), - valid_token=valid_token, - request=request, + route="/v1/chat/completions", ) - assert result is True + assert exc_info.value.type == ProxyErrorTypes.tool_access_denied + assert "get_weather" in str(exc_info.value.message) @pytest.mark.asyncio - async def test_mcp_tool_call_disallowed_raises(self): - """Route /mcp-rest/tools/call with tool not in token allowed_tools raises.""" - request = MockRequest("POST") - request_body = dict(MCP_TOOL_CALL_BODY_ALLOWED) - valid_token = get_mock_user_token(metadata={"allowed_tools": ["other"]}) + async def test_team_allowlist_used_when_key_empty(self): + token = _token( + metadata={}, + team_metadata={"allowed_tools": ["get_weather"]}, + ) + body = { + "tools": [{"type": "function", "function": {"name": "get_weather"}}] + } + await check_tools_allowlist( + request_body=body, + valid_token=token, + team_object=None, + route="/v1/chat/completions", + ) - p1, p2, p3 = _tools_allowlist_patches() - with p1, p2, p3: - with pytest.raises((Exception, ProxyException)) as exc_info: - await common_checks( - request_body=request_body, - team_object=None, - user_object=None, - end_user_object=None, - global_proxy_spend=None, - general_settings={}, - route="/mcp-rest/tools/call", - llm_router=None, - proxy_logging_obj=MagicMock(), - valid_token=valid_token, - request=request, - ) - exc_str = ( - getattr(exc_info.value, "message", None) or str(exc_info.value) or "" - ).lower() - assert "tool" in exc_str or "allowed" in exc_str + @pytest.mark.asyncio + async def test_key_allowlist_overrides_team(self): + token = _token( + metadata={"allowed_tools": ["get_weather"]}, + team_metadata={"allowed_tools": ["other_tool"]}, + ) + body = { + "tools": [{"type": "function", "function": {"name": "get_weather"}}] + } + await check_tools_allowlist( + request_body=body, + valid_token=token, + team_object=None, + route="/v1/chat/completions", + ) + + @pytest.mark.asyncio + async def test_valid_token_none_skips(self): + await check_tools_allowlist( + request_body={"tools": [{"type": "function", "function": {"name": "x"}}]}, + valid_token=None, + team_object=None, + route="/v1/chat/completions", + ) + + @pytest.mark.asyncio + async def test_no_tools_in_body_passes(self): + token = _token(metadata={"allowed_tools": ["get_weather"]}) + await check_tools_allowlist( + request_body={"messages": []}, + valid_token=token, + team_object=None, + route="/v1/chat/completions", + ) diff --git a/ui/litellm-dashboard/src/components/ToolPolicies.tsx b/ui/litellm-dashboard/src/components/ToolPolicies.tsx index aa005495575..f39b21aed76 100644 --- a/ui/litellm-dashboard/src/components/ToolPolicies.tsx +++ b/ui/litellm-dashboard/src/components/ToolPolicies.tsx @@ -1,14 +1,24 @@ "use client"; import React, { useCallback, useDeferredValue, useEffect, useMemo, useState } from "react"; -import { Select, Switch, Tooltip } from "antd"; +import { Button, Modal, Select, Switch, Tooltip } from "antd"; import { Table, TableHead, TableHeaderCell, TableBody, TableRow, TableCell } from "@tremor/react"; import { TimeCell } from "./view_logs/time_cell"; import { TableHeaderSortDropdown } from "./common_components/TableHeaderSortDropdown/TableHeaderSortDropdown"; import type { SortState } from "./common_components/TableHeaderSortDropdown/TableHeaderSortDropdown"; import FilterComponent, { FilterOption } from "./molecules/filter"; import { MetricCard } from "./GuardrailsMonitor/MetricCard"; -import { fetchToolsList, updateToolPolicy, ToolRow } from "./networking"; +import TeamDropdown from "./common_components/team_dropdown"; +import { + fetchToolDetail, + fetchToolsList, + updateToolPolicy, + deleteToolPolicyOverride, + ToolRow, + ToolDetailResponse, + ToolPolicyOverrideRow, +} from "./networking"; +import { teamListCall, keyListCall } from "./networking"; // --- Date helpers (UTC) for "new tools" counts --- function getUTCDateKey(date: Date): string { @@ -119,6 +129,16 @@ const PolicySelect: React.FC<{ ); }; +interface TeamOption { + team_id: string; + team_alias?: string; +} + +interface KeyOption { + token: string; + key_alias?: string; +} + export const ToolPolicies: React.FC = ({ accessToken }) => { const [tools, setTools] = useState([]); const [loading, setLoading] = useState(true); @@ -126,6 +146,17 @@ export const ToolPolicies: React.FC = ({ accessToken }) => { const [error, setError] = useState(null); const [saving, setSaving] = useState(null); + const [detailModalOpen, setDetailModalOpen] = useState(false); + const [detailToolName, setDetailToolName] = useState(null); + const [detail, setDetail] = useState(null); + const [detailLoading, setDetailLoading] = useState(false); + const [teams, setTeams] = useState([]); + const [keys, setKeys] = useState([]); + const [overrideSaving, setOverrideSaving] = useState(false); + const [blockScope, setBlockScope] = useState<"team" | "key">("team"); + const [blockTeamId, setBlockTeamId] = useState(null); + const [blockKey, setBlockKey] = useState<{ token: string; key_alias?: string } | null>(null); + const [searchTerm, setSearchTerm] = useState(""); const [sortField, setSortField] = useState("created_at"); const [sortOrder, setSortOrder] = useState<"asc" | "desc">("desc"); @@ -168,6 +199,9 @@ export const ToolPolicies: React.FC = ({ accessToken }) => { try { await updateToolPolicy(accessToken, toolName, newPolicy); setTools((prev) => prev.map((t) => (t.tool_name === toolName ? { ...t, call_policy: newPolicy } : t))); + if (detailToolName === toolName && detail) { + setDetail((d) => (d ? { ...d, tool: { ...d.tool, call_policy: newPolicy } } : null)); + } } catch (e: any) { alert(`Failed to update policy: ${e.message}`); } finally { @@ -175,6 +209,91 @@ export const ToolPolicies: React.FC = ({ accessToken }) => { } }; + const openDetailModal = useCallback( + async (toolName: string) => { + if (!accessToken) return; + setDetailToolName(toolName); + setDetailModalOpen(true); + setDetail(null); + setDetailLoading(true); + setBlockTeamId(null); + setBlockKey(null); + try { + const [detailRes, teamsRes, keysRes] = await Promise.all([ + fetchToolDetail(accessToken, toolName), + teamListCall(accessToken, null, null), + keyListCall(accessToken, null, null, null, null, null, 1, 100), + ]); + setDetail(detailRes); + const teamsArray = Array.isArray(teamsRes) ? teamsRes : teamsRes?.data ?? []; + setTeams( + teamsArray.map((t: any) => ({ team_id: t.team_id ?? t.id, team_alias: t.team_alias ?? t.team_id })) + ); + const keysArray = keysRes?.keys ?? keysRes?.data ?? []; + setKeys( + keysArray.map((k: any) => ({ + token: k.token ?? k.api_key ?? k.key_hash ?? "", + key_alias: k.key_alias ?? k.token?.substring?.(0, 8), + })) + ); + } catch (e: any) { + setError(e.message ?? "Failed to load tool detail"); + } finally { + setDetailLoading(false); + } + }, + [accessToken] + ); + + const closeDetailModal = useCallback(() => { + setDetailModalOpen(false); + setDetailToolName(null); + setDetail(null); + }, []); + + const handleAddOverride = useCallback(async () => { + if (!accessToken || !detailToolName) return; + const isTeam = blockScope === "team"; + if (isTeam && !blockTeamId) return; + if (!isTeam && !blockKey?.token) return; + setOverrideSaving(true); + try { + await updateToolPolicy(accessToken, detailToolName, "blocked", { + team_id: isTeam ? blockTeamId! : undefined, + key_hash: !isTeam ? blockKey!.token : undefined, + key_alias: !isTeam ? blockKey!.key_alias : undefined, + }); + const refreshed = await fetchToolDetail(accessToken, detailToolName); + setDetail(refreshed); + setBlockTeamId(null); + setBlockKey(null); + } catch (e: any) { + alert(`Failed to add override: ${e.message}`); + } finally { + setOverrideSaving(false); + } + }, [accessToken, detailToolName, blockScope, blockTeamId, blockKey]); + + const handleRemoveOverride = useCallback( + async (override: ToolPolicyOverrideRow) => { + if (!accessToken || !detailToolName) return; + setOverrideSaving(true); + try { + await deleteToolPolicyOverride(accessToken, detailToolName, { + team_id: override.team_id ?? undefined, + key_hash: override.key_hash ?? undefined, + }); + const refreshed = await fetchToolDetail(accessToken, detailToolName); + setDetail(refreshed); + } catch (e: any) { + alert(`Failed to remove override: ${e.message}`); + } finally { + setOverrideSaving(false); + } + }, + [accessToken, detailToolName] + ); + const handleSortChange = (field: SortField, newState: SortState) => { if (newState === false) { setSortField("created_at"); @@ -523,11 +642,15 @@ export const ToolPolicies: React.FC = ({ accessToken }) => { - - - {tool.tool_name} - - + = ({ accessToken }) => {
)}
+ + {/* Tool detail modal: view tool, global policy, overrides, block for team/key */} + + {detailLoading ? ( +

Loading…

+ ) : detail ? ( +
+
+ + Origin: {detail.tool.origin ?? "—"} + + + # Calls: {(detail.tool.call_count ?? 0).toLocaleString()} + +
+
+ Global policy + +
+ + {detail.overrides.length > 0 && ( +
+ Blocked for team/key +
    + {detail.overrides.map((ov) => ( +
  • + + {ov.team_id ? `Team: ${ov.team_id}` : ""} + {ov.team_id && ov.key_hash ? " · " : ""} + {ov.key_hash ? `Key: ${ov.key_alias || ov.key_hash.substring(0, 8)}` : ""} + {!ov.team_id && !ov.key_hash ? "—" : ""} + + +
  • + ))} +
+
+ )} + +
+ Block for team or key +
+
+ + +
+ {blockScope === "team" ? ( +
+ setBlockTeamId(id ?? null)} + /> +
+ ) : ( + setBlockScope("team")} + className="align-middle" + /> + Team + + +
+
+
+ + {blockScope === "team" ? "Team" : "Key"} + + {blockScope === "team" ? ( + setBlockTeamId(id || null)} + /> + ) : ( + onChange(toolName, v)} - onClick={(e) => e.stopPropagation()} - style={{ - minWidth: 110, - fontWeight: 500, - }} - styles={{ - selector: { - backgroundColor: style.bg, - borderColor: style.border, - color: style.color, - borderRadius: 999, - fontSize: 11, - fontWeight: 600, - paddingLeft: 8, - paddingRight: 4, - }, - }} - popupMatchSelectWidth={false} - options={POLICY_OPTIONS.map((o) => ({ - value: o.value, - label: ( - - - {o.label} - - ), - }))} - /> - ); -}; - -interface TeamOption { - team_id: string; - team_alias?: string; -} - -interface KeyOption { - token: string; - key_alias?: string; -} - -export const ToolPolicies: React.FC = ({ accessToken }) => { +export const ToolPolicies: React.FC = ({ accessToken, onSelectTool }) => { const [tools, setTools] = useState([]); const [loading, setLoading] = useState(true); const [isFetching, setIsFetching] = useState(false); const [error, setError] = useState(null); const [saving, setSaving] = useState(null); - const [detailModalOpen, setDetailModalOpen] = useState(false); - const [detailToolName, setDetailToolName] = useState(null); - const [detail, setDetail] = useState(null); - const [detailLoading, setDetailLoading] = useState(false); - const [teams, setTeams] = useState([]); - const [keys, setKeys] = useState([]); - const [overrideSaving, setOverrideSaving] = useState(false); - const [blockScope, setBlockScope] = useState<"team" | "key">("team"); - const [blockTeamId, setBlockTeamId] = useState(null); - const [blockKey, setBlockKey] = useState<{ token: string; key_alias?: string } | null>(null); - const [searchTerm, setSearchTerm] = useState(""); const [sortField, setSortField] = useState("created_at"); const [sortOrder, setSortOrder] = useState<"asc" | "desc">("desc"); @@ -199,9 +102,6 @@ export const ToolPolicies: React.FC = ({ accessToken }) => { try { await updateToolPolicy(accessToken, toolName, newPolicy); setTools((prev) => prev.map((t) => (t.tool_name === toolName ? { ...t, call_policy: newPolicy } : t))); - if (detailToolName === toolName && detail) { - setDetail((d) => (d ? { ...d, tool: { ...d.tool, call_policy: newPolicy } } : null)); - } } catch (e: any) { alert(`Failed to update policy: ${e.message}`); } finally { @@ -209,91 +109,6 @@ export const ToolPolicies: React.FC = ({ accessToken }) => { } }; - const openDetailModal = useCallback( - async (toolName: string) => { - if (!accessToken) return; - setDetailToolName(toolName); - setDetailModalOpen(true); - setDetail(null); - setDetailLoading(true); - setBlockTeamId(null); - setBlockKey(null); - try { - const [detailRes, teamsRes, keysRes] = await Promise.all([ - fetchToolDetail(accessToken, toolName), - teamListCall(accessToken, null, null), - keyListCall(accessToken, null, null, null, null, null, 1, 100), - ]); - setDetail(detailRes); - const teamsArray = Array.isArray(teamsRes) ? teamsRes : teamsRes?.data ?? []; - setTeams( - teamsArray.map((t: any) => ({ team_id: t.team_id ?? t.id, team_alias: t.team_alias ?? t.team_id })) - ); - const keysArray = keysRes?.keys ?? keysRes?.data ?? []; - setKeys( - keysArray.map((k: any) => ({ - token: k.token ?? k.api_key ?? k.key_hash ?? "", - key_alias: k.key_alias ?? k.token?.substring?.(0, 8), - })) - ); - } catch (e: any) { - setError(e.message ?? "Failed to load tool detail"); - } finally { - setDetailLoading(false); - } - }, - [accessToken] - ); - - const closeDetailModal = useCallback(() => { - setDetailModalOpen(false); - setDetailToolName(null); - setDetail(null); - }, []); - - const handleAddOverride = useCallback(async () => { - if (!accessToken || !detailToolName) return; - const isTeam = blockScope === "team"; - if (isTeam && !blockTeamId) return; - if (!isTeam && !blockKey?.token) return; - setOverrideSaving(true); - try { - await updateToolPolicy(accessToken, detailToolName, "blocked", { - team_id: isTeam ? blockTeamId! : undefined, - key_hash: !isTeam ? blockKey!.token : undefined, - key_alias: !isTeam ? blockKey!.key_alias : undefined, - }); - const refreshed = await fetchToolDetail(accessToken, detailToolName); - setDetail(refreshed); - setBlockTeamId(null); - setBlockKey(null); - } catch (e: any) { - alert(`Failed to add override: ${e.message}`); - } finally { - setOverrideSaving(false); - } - }, [accessToken, detailToolName, blockScope, blockTeamId, blockKey]); - - const handleRemoveOverride = useCallback( - async (override: ToolPolicyOverrideRow) => { - if (!accessToken || !detailToolName) return; - setOverrideSaving(true); - try { - await deleteToolPolicyOverride(accessToken, detailToolName, { - team_id: override.team_id ?? undefined, - key_hash: override.key_hash ?? undefined, - }); - const refreshed = await fetchToolDetail(accessToken, detailToolName); - setDetail(refreshed); - } catch (e: any) { - alert(`Failed to remove override: ${e.message}`); - } finally { - setOverrideSaving(false); - } - }, - [accessToken, detailToolName] - ); - const handleSortChange = (field: SortField, newState: SortState) => { if (newState === false) { setSortField("created_at"); @@ -431,7 +246,7 @@ export const ToolPolicies: React.FC = ({ accessToken }) => { }; return ( -
+

Tool Policies

{/* Summary cards */} @@ -644,10 +459,10 @@ export const ToolPolicies: React.FC = ({ accessToken }) => { @@ -660,8 +475,10 @@ export const ToolPolicies: React.FC = ({ accessToken }) => { onChange={handlePolicyChange} /> - - {(tool.call_count ?? 0).toLocaleString()} + +
+ {(tool.call_count ?? 0).toLocaleString()} +
@@ -718,134 +535,6 @@ export const ToolPolicies: React.FC = ({ accessToken }) => { )}
- {/* Tool detail modal: view tool, global policy, overrides, block for team/key */} - - {detailLoading ? ( -

Loading…

- ) : detail ? ( -
-
- - Origin: {detail.tool.origin ?? "—"} - - - # Calls: {(detail.tool.call_count ?? 0).toLocaleString()} - -
-
- Global policy - -
- - {detail.overrides.length > 0 && ( -
- Blocked for team/key -
    - {detail.overrides.map((ov) => ( -
  • - - {ov.team_id ? `Team: ${ov.team_id}` : ""} - {ov.team_id && ov.key_hash ? " · " : ""} - {ov.key_hash ? `Key: ${ov.key_alias || ov.key_hash.substring(0, 8)}` : ""} - {!ov.team_id && !ov.key_hash ? "—" : ""} - - -
  • - ))} -
-
- )} - -
- Block for team or key -
-
- - -
- {blockScope === "team" ? ( -
- setBlockTeamId(id ?? null)} - /> -
- ) : ( - onChange(toolName, v)} + onClick={(e) => stopPropagation && e.stopPropagation()} + style={{ + minWidth, + fontWeight: 500, + }} + styles={{ + selector: { + backgroundColor: style.bg, + borderColor: style.border, + color: style.color, + borderRadius: 999, + fontSize: size === "small" ? 11 : 12, + fontWeight: 600, + paddingLeft: 8, + paddingRight: 4, + }, + }} + popupMatchSelectWidth={false} + options={POLICY_OPTIONS.map((o) => ({ + value: o.value, + label: ( + + + {o.label} + + ), + }))} + /> + ); +}; diff --git a/ui/litellm-dashboard/src/components/ToolPoliciesView.tsx b/ui/litellm-dashboard/src/components/ToolPoliciesView.tsx new file mode 100644 index 00000000000..3e4964f7e36 --- /dev/null +++ b/ui/litellm-dashboard/src/components/ToolPoliciesView.tsx @@ -0,0 +1,44 @@ +"use client"; + +import React, { useState } from "react"; +import { ToolDetail } from "@/components/ToolDetail"; +import { ToolPolicies } from "@/components/ToolPolicies"; + +type View = + | { type: "overview" } + | { type: "detail"; toolName: string }; + +interface ToolPoliciesViewProps { + accessToken: string | null; + userRole?: string; +} + +export default function ToolPoliciesView({ accessToken, userRole }: ToolPoliciesViewProps) { + const [view, setView] = useState({ type: "overview" }); + + const handleSelectTool = (toolName: string) => { + setView({ type: "detail", toolName }); + }; + + const handleBack = () => { + setView({ type: "overview" }); + }; + + return ( +
+ {view.type === "detail" ? ( + + ) : ( + + )} +
+ ); +} From 3284df3bfa9f010007c1aa2e467182daf8baffaf Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 25 Feb 2026 23:52:54 -0800 Subject: [PATCH 07/12] feat: working key tool blocking --- litellm/proxy/_new_secret_config.yaml | 30 +------- litellm/proxy/_types.py | 75 +++++++------------ litellm/proxy/db/tool_registry_writer.py | 10 +-- .../tool_policy/tool_policy_guardrail.py | 17 +++-- 4 files changed, 50 insertions(+), 82 deletions(-) diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index 6b84d90a327..508c1c94659 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -23,33 +23,11 @@ model_list: guardrails: - - guardrail_name: "airline-competitor-intent" - guardrail_id: "airline-competitor-intent" + - guardrail_name: "tool_policy" litellm_params: - guardrail: litellm_content_filter - mode: pre_call - default_on: false - competitor_intent_config: - brand_self: - - emirates - - ek - competitors: - - qatar airways - - qatar - - etihad - locations: - - qatar - - doha - - doh - competitor_aliases: - qatar airways: [qr, doha airline] - qatar: [qr] - policy: - competitor_comparison: refuse - possible_competitor_comparison: reframe - threshold_high: 0.70 - threshold_medium: 0.45 - threshold_low: 0.30 + guardrail: tool_policy + mode: [pre_call, post_call] + default_on: true mcp_servers: my_http_server: diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 65ea9cb42d8..2735b3780f5 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1,60 +1,40 @@ import enum import json from datetime import datetime -from typing import TYPE_CHECKING, Any, Callable, Dict, List, Literal, Optional, Union +from typing import (TYPE_CHECKING, Any, Callable, Dict, List, Literal, + Optional, Union) import httpx -from pydantic import ( - BaseModel, - ConfigDict, - Field, - Json, - field_validator, - model_validator, -) +from pydantic import (BaseModel, ConfigDict, Field, Json, field_validator, + model_validator) from typing_extensions import Required, TypedDict from litellm._uuid import uuid from litellm.types.integrations.slack_alerting import AlertType -from litellm.types.llms.openai import ( - AllMessageValues, - OpenAIFileObject, - ResponsesAPIResponse, -) -from litellm.types.mcp import ( - MCPAuth, - MCPAuthType, - MCPCredentials, - MCPTransport, - MCPTransportType, -) +from litellm.types.llms.openai import (AllMessageValues, OpenAIFileObject, + ResponsesAPIResponse) +from litellm.types.mcp import (MCPAuthType, MCPCredentials, MCPTransport, + MCPTransportType) from litellm.types.mcp_server.mcp_server_manager import MCPInfo from litellm.types.router import RouterErrors, UpdateRouterConfig from litellm.types.secret_managers.main import KeyManagementSystem -from litellm.types.utils import ( - CallTypes, - CostBreakdown, - EmbeddingResponse, - GenericBudgetConfigType, - ImageResponse, - LiteLLMBatch, - LiteLLMFineTuningJob, - LiteLLMPydanticObjectBase, - ModelResponse, - ProviderField, - StandardCallbackDynamicParams, - StandardLoggingGuardrailInformation, - StandardLoggingMCPToolCall, - StandardLoggingModelInformation, - StandardLoggingPayloadErrorInformation, - StandardLoggingPayloadStatus, - StandardLoggingVectorStoreRequest, - StandardPassThroughResponseObject, - TextCompletionResponse, -) +from litellm.types.utils import (CallTypes, CostBreakdown, EmbeddingResponse, + GenericBudgetConfigType, ImageResponse, + LiteLLMBatch, LiteLLMFineTuningJob, + LiteLLMPydanticObjectBase, ModelResponse, + ProviderField, StandardCallbackDynamicParams, + StandardLoggingGuardrailInformation, + StandardLoggingMCPToolCall, + StandardLoggingModelInformation, + StandardLoggingPayloadErrorInformation, + StandardLoggingPayloadStatus, + StandardLoggingVectorStoreRequest, + StandardPassThroughResponseObject, + TextCompletionResponse) from litellm.types.videos.main import VideoObject -from .types_utils.utils import get_instance_fn, validate_custom_validate_return_type +from .types_utils.utils import (get_instance_fn, + validate_custom_validate_return_type) if TYPE_CHECKING: from opentelemetry.trace import Span as _Span @@ -2369,7 +2349,8 @@ class UserAPIKeyAuth( This is used to track number of requests/spend for health check calls. """ - from litellm.constants import LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME + from litellm.constants import \ + LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME return cls( api_key=LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME, @@ -2401,7 +2382,8 @@ class UserAPIKeyAuth( This is used to track actions performed by automated system jobs. """ - from litellm.constants import LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME + from litellm.constants import \ + LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME return cls( api_key=LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME, @@ -2792,7 +2774,8 @@ class LiteLLM_AuditLogs(LiteLLMPydanticObjectBase): @model_validator(mode="after") def mask_api_keys(self): - from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker + from litellm.litellm_core_utils.sensitive_data_masker import \ + SensitiveDataMasker masker = SensitiveDataMasker(sensitive_patterns={"key"}) diff --git a/litellm/proxy/db/tool_registry_writer.py b/litellm/proxy/db/tool_registry_writer.py index a0dffefdc59..93d17685ccc 100644 --- a/litellm/proxy/db/tool_registry_writer.py +++ b/litellm/proxy/db/tool_registry_writer.py @@ -13,11 +13,9 @@ from litellm._logging import verbose_proxy_logger from litellm.caching.dual_cache import DualCache from litellm.constants import TOOL_POLICY_CACHE_TTL_SECONDS from litellm.proxy._types import ToolDiscoveryQueueItem -from litellm.types.tool_management import ( - LiteLLM_ToolTableRow, - ToolCallPolicy, - ToolPolicyOverrideRow, -) +from litellm.types.tool_management import (LiteLLM_ToolTableRow, + ToolCallPolicy, + ToolPolicyOverrideRow) if TYPE_CHECKING: from litellm.proxy.utils import PrismaClient @@ -236,10 +234,12 @@ def _override_row_to_model(row: Any) -> ToolPolicyOverrideRow: "updated_at", ) } + def _norm(s: Optional[str]) -> Optional[str]: if s is None or s == _TOOL_OVERRIDE_ANY: return None return s or None + return ToolPolicyOverrideRow( override_id=row.get("override_id", ""), tool_name=row.get("tool_name", ""), diff --git a/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py index 213b248929e..08b28824cc0 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py @@ -31,6 +31,8 @@ from litellm._logging import verbose_proxy_logger from litellm.caching.dual_cache import DualCache from litellm.integrations.custom_guardrail import (CustomGuardrail, log_guardrail_information) +from litellm.proxy.guardrails.tool_name_extraction import \ + extract_request_tool_names from litellm.types.guardrails import GuardrailEventHooks from litellm.types.utils import GenericGuardrailAPIInputs @@ -41,7 +43,9 @@ if TYPE_CHECKING: GUARDRAIL_NAME = "tool_policy" -def _get_request_team_and_key(request_data: dict) -> Tuple[Optional[str], Optional[str]]: +def _get_request_team_and_key( + request_data: dict, +) -> Tuple[Optional[str], Optional[str]]: """Extract team_id and key hash from request_data (litellm_metadata or metadata).""" if not request_data: return None, None @@ -68,13 +72,17 @@ def _get_request_route_from_data(request_data: dict) -> Optional[str]: return meta.get("user_api_key_request_route") -def _get_effective_allowed_tools_from_request(request_data: dict) -> Optional[List[str]]: +def _get_effective_allowed_tools_from_request( + request_data: dict, +) -> Optional[List[str]]: """Key allowed_tools overrides team; empty/missing means no restriction.""" meta = request_data.get("metadata") or request_data.get("litellm_metadata") or {} key_meta = meta.get("user_api_key_metadata") or {} team_meta = meta.get("user_api_key_team_metadata") or {} key_allowed = key_meta.get("allowed_tools") if isinstance(key_meta, dict) else None - team_allowed = team_meta.get("allowed_tools") if isinstance(team_meta, dict) else None + team_allowed = ( + team_meta.get("allowed_tools") if isinstance(team_meta, dict) else None + ) if isinstance(key_allowed, list) and len(key_allowed) > 0: return key_allowed if isinstance(team_allowed, list) and len(team_allowed) > 0: @@ -122,8 +130,7 @@ class ToolPolicyGuardrail(CustomGuardrail): if not tool_names: route = _get_request_route_from_data(request_data) if route: - from litellm.proxy.guardrails.tool_name_extraction import \ - extract_request_tool_names + tool_names = extract_request_tool_names(route, request_data) else: # response tool_calls = inputs.get("tool_calls") or [] From 0f8832f05b74c7dabafcee72d46280c26f51a4c1 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Thu, 26 Feb 2026 00:26:07 -0800 Subject: [PATCH 08/12] feat(tools): show tool logs --- .../migration.sql | 2 - .../migration.sql | 2 - .../migration.sql | 27 - .../migration.sql | 2 + .../migration.sql | 11 + .../litellm_proxy_extras/schema.prisma | 29 +- litellm/proxy/db/spend_log_tool_index.py | 147 +++ litellm/proxy/db/tool_registry_writer.py | 326 ++++--- .../tool_policy/tool_policy_guardrail.py | 45 +- litellm/proxy/litellm_pre_call_utils.py | 6 + .../tool_management_endpoints.py | 318 +++++- litellm/proxy/proxy_server.py | 905 +++++++----------- litellm/proxy/schema.prisma | 30 +- litellm/proxy/utils.py | 32 +- litellm/types/tool_management.py | 18 + schema.prisma | 10 + .../src/components/ToolDetail.tsx | 151 ++- .../src/components/networking.tsx | 45 + ui/litellm-dashboard/tsconfig.json | 2 +- 19 files changed, 1272 insertions(+), 836 deletions(-) delete mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260225210135_baseline_diff/migration.sql delete mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260225211151_baseline_diff/migration.sql delete mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260225220000_add_tool_policy_override_table/migration.sql create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260226000000_add_blocked_tools_to_object_permission/migration.sql create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260226120000_add_spend_log_tool_index/migration.sql create mode 100644 litellm/proxy/db/spend_log_tool_index.py diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260225210135_baseline_diff/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260225210135_baseline_diff/migration.sql deleted file mode 100644 index 2f725d83806..00000000000 --- a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260225210135_baseline_diff/migration.sql +++ /dev/null @@ -1,2 +0,0 @@ --- This is an empty migration. - diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260225211151_baseline_diff/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260225211151_baseline_diff/migration.sql deleted file mode 100644 index 2f725d83806..00000000000 --- a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260225211151_baseline_diff/migration.sql +++ /dev/null @@ -1,2 +0,0 @@ --- This is an empty migration. - diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260225220000_add_tool_policy_override_table/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260225220000_add_tool_policy_override_table/migration.sql deleted file mode 100644 index 642508b7b20..00000000000 --- a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260225220000_add_tool_policy_override_table/migration.sql +++ /dev/null @@ -1,27 +0,0 @@ --- CreateTable -CREATE TABLE "LiteLLM_ToolPolicyOverrideTable" ( - "override_id" TEXT NOT NULL, - "tool_name" TEXT NOT NULL, - "team_id" TEXT NOT NULL DEFAULT '', - "key_hash" TEXT NOT NULL DEFAULT '', - "call_policy" TEXT NOT NULL DEFAULT 'blocked', - "key_alias" TEXT, - "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, - "created_by" TEXT, - "updated_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, - "updated_by" TEXT, - - CONSTRAINT "LiteLLM_ToolPolicyOverrideTable_pkey" PRIMARY KEY ("override_id") -); - --- CreateIndex -CREATE INDEX "LiteLLM_ToolPolicyOverrideTable_tool_name_idx" ON "LiteLLM_ToolPolicyOverrideTable"("tool_name"); - --- CreateIndex -CREATE INDEX "LiteLLM_ToolPolicyOverrideTable_team_id_idx" ON "LiteLLM_ToolPolicyOverrideTable"("team_id"); - --- CreateIndex -CREATE INDEX "LiteLLM_ToolPolicyOverrideTable_key_hash_idx" ON "LiteLLM_ToolPolicyOverrideTable"("key_hash"); - --- CreateIndex -CREATE UNIQUE INDEX "LiteLLM_ToolPolicyOverrideTable_tool_name_team_id_key_hash_key" ON "LiteLLM_ToolPolicyOverrideTable"("tool_name", "team_id", "key_hash"); diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226000000_add_blocked_tools_to_object_permission/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226000000_add_blocked_tools_to_object_permission/migration.sql new file mode 100644 index 00000000000..cba06684193 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226000000_add_blocked_tools_to_object_permission/migration.sql @@ -0,0 +1,2 @@ +-- AlterTable +ALTER TABLE "LiteLLM_ObjectPermissionTable" ADD COLUMN "blocked_tools" TEXT[] DEFAULT ARRAY[]::TEXT[]; diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226120000_add_spend_log_tool_index/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226120000_add_spend_log_tool_index/migration.sql new file mode 100644 index 00000000000..e3199679ce2 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226120000_add_spend_log_tool_index/migration.sql @@ -0,0 +1,11 @@ +-- CreateTable +CREATE TABLE "LiteLLM_SpendLogToolIndex" ( + "request_id" TEXT NOT NULL, + "tool_name" TEXT NOT NULL, + "start_time" TIMESTAMP(3) NOT NULL, + + CONSTRAINT "LiteLLM_SpendLogToolIndex_pkey" PRIMARY KEY ("request_id","tool_name") +); + +-- CreateIndex +CREATE INDEX "LiteLLM_SpendLogToolIndex_tool_name_start_time_idx" ON "LiteLLM_SpendLogToolIndex"("tool_name", "start_time"); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 08f756e4fb0..ecaa43fbd7d 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -260,6 +260,7 @@ model LiteLLM_ObjectPermissionTable { vector_stores String[] @default([]) agents String[] @default([]) agent_access_groups String[] @default([]) + blocked_tools String[] @default([]) // Tool names blocked for any key/team/user with this permission teams LiteLLM_TeamTable[] projects LiteLLM_ProjectTable[] verification_tokens LiteLLM_VerificationToken[] @@ -920,6 +921,16 @@ model LiteLLM_SpendLogGuardrailIndex { @@index([policy_id, start_time]) } +// Index for fast "last N logs for tool" from SpendLogs – see how a tool is called in production +model LiteLLM_SpendLogToolIndex { + request_id String + tool_name String // matches LiteLLM_ToolTable.tool_name; join for call_policy etc. + start_time DateTime + + @@id([request_id, tool_name]) + @@index([tool_name, start_time]) +} + // Prompt table for storing prompt configurations model LiteLLM_PromptTable { id String @id @default(uuid()) @@ -1078,24 +1089,6 @@ model LiteLLM_ToolTable { } // Per-(tool, team/key) policy overrides. When present, override replaces global tool policy for that scope. -model LiteLLM_ToolPolicyOverrideTable { - override_id String @id @default(uuid()) - tool_name String - team_id String @default("") // "" = not scoped to team; non-empty = override for this team only - key_hash String @default("") // "" = not scoped to key; non-empty = override for this key only - call_policy String @default("blocked") - key_alias String? // human-readable key alias for UI - created_at DateTime @default(now()) - created_by String? - updated_at DateTime @default(now()) @updatedAt - updated_by String? - - @@unique([tool_name, team_id, key_hash]) - @@index([tool_name]) - @@index([team_id]) - @@index([key_hash]) -} - //Unified Access Groups table for storing unified access groups model LiteLLM_AccessGroupTable { access_group_id String @id @default(uuid()) diff --git a/litellm/proxy/db/spend_log_tool_index.py b/litellm/proxy/db/spend_log_tool_index.py new file mode 100644 index 00000000000..6e8c63675e6 --- /dev/null +++ b/litellm/proxy/db/spend_log_tool_index.py @@ -0,0 +1,147 @@ +""" +Track tool usage for the dashboard: insert into SpendLogToolIndex when spend logs +are written, so "last N requests for tool X" and "how is this tool called in production" +queries are fast. +""" + +from datetime import datetime, timezone +from typing import Any, Dict, List, Set + +from litellm._logging import verbose_proxy_logger +from litellm.litellm_core_utils.safe_json_loads import safe_json_loads +from litellm.proxy.utils import PrismaClient + + +def _add_tool_calls_to_set(tool_calls: Any, out: Set[str]) -> None: + """Extract tool names from OpenAI-style tool_calls list into out.""" + if not isinstance(tool_calls, list): + return + for tc in tool_calls: + if not isinstance(tc, dict): + continue + fn = tc.get("function") + if isinstance(fn, dict): + name = fn.get("name") + if name and isinstance(name, str) and name.strip(): + out.add(name.strip()) + + +def _parse_tool_names_from_payload(payload: Dict[str, Any]) -> Set[str]: + """ + Extract deduplicated tool names from a spend log payload. + Sources: mcp_namespaced_tool_name, response (tool_calls), proxy_server_request (tools). + """ + tool_names: Set[str] = set() + + # Top-level MCP tool name (single tool per request for that flow) + mcp_name = payload.get("mcp_namespaced_tool_name") + if mcp_name and isinstance(mcp_name, str) and mcp_name.strip(): + tool_names.add(mcp_name.strip()) + + # Response: OpenAI-style tool_calls[].function.name or choices[0].message.tool_calls + response_raw = payload.get("response") + if response_raw: + response_obj = ( + safe_json_loads(response_raw, default=None) + if isinstance(response_raw, str) + else response_raw + ) + if isinstance(response_obj, dict): + _add_tool_calls_to_set(response_obj.get("tool_calls"), tool_names) + choices = response_obj.get("choices") + if isinstance(choices, list) and choices: + msg = choices[0].get("message") if isinstance(choices[0], dict) else None + if isinstance(msg, dict): + _add_tool_calls_to_set(msg.get("tool_calls"), tool_names) + + # Request body: tools[].function.name + request_raw = payload.get("proxy_server_request") + if request_raw: + request_obj = ( + safe_json_loads(request_raw, default=None) + if isinstance(request_raw, str) + else request_raw + ) + if isinstance(request_obj, dict): + body = request_obj.get("body", request_obj) + if isinstance(body, dict): + request_obj = body + if isinstance(request_obj, dict): + tools = request_obj.get("tools") + if isinstance(tools, list): + for t in tools: + if isinstance(t, dict): + fn = t.get("function") + if isinstance(fn, dict): + name = fn.get("name") + if name and isinstance(name, str) and name.strip(): + tool_names.add(name.strip()) + + return tool_names + + +async def process_spend_logs_tool_usage( + prisma_client: PrismaClient, + logs_to_process: List[Dict[str, Any]], +) -> None: + """ + After spend logs are written: insert SpendLogToolIndex rows from each payload. + Extracts tool names from mcp_namespaced_tool_name, response tool_calls, and + proxy_server_request tools. + """ + if not logs_to_process: + return + + index_rows: List[Dict[str, Any]] = [] + + for payload in logs_to_process: + request_id = payload.get("request_id") + start_time = payload.get("startTime") + if not request_id or not start_time: + continue + if isinstance(start_time, str): + try: + start_time = datetime.fromisoformat( + start_time.replace("Z", "+00:00") + ) + except (ValueError, TypeError): + continue + if start_time.tzinfo is None: + start_time = start_time.replace(tzinfo=timezone.utc) + + tool_names = _parse_tool_names_from_payload(payload) + for tool_name in tool_names: + index_rows.append({ + "request_id": request_id, + "tool_name": tool_name, + "start_time": start_time, + }) + + if not index_rows: + return + + try: + index_data = [] + for r in index_rows: + st = r["start_time"] + if isinstance(st, str): + try: + st = datetime.fromisoformat(st.replace("Z", "+00:00")) + except (ValueError, TypeError): + continue + if st.tzinfo is None: + st = st.replace(tzinfo=timezone.utc) + index_data.append({ + "request_id": r["request_id"], + "tool_name": r["tool_name"], + "start_time": st, + }) + if index_data: + await prisma_client.db.litellm_spendlogtoolindex.create_many( + data=index_data, + skip_duplicates=True, + ) + except Exception as e: + verbose_proxy_logger.warning( + "Tool usage tracking (SpendLogToolIndex) failed (non-fatal): %s", e + ) diff --git a/litellm/proxy/db/tool_registry_writer.py b/litellm/proxy/db/tool_registry_writer.py index 93d17685ccc..172e05cec9d 100644 --- a/litellm/proxy/db/tool_registry_writer.py +++ b/litellm/proxy/db/tool_registry_writer.py @@ -20,9 +20,6 @@ from litellm.types.tool_management import (LiteLLM_ToolTableRow, if TYPE_CHECKING: from litellm.proxy.utils import PrismaClient -# Sentinel for "not scoped" in override table (DB unique constraint needs non-null) -_TOOL_OVERRIDE_ANY = "" - TOOL_POLICY_CACHE_KEY_PREFIX = "tool_policy:" @@ -215,54 +212,54 @@ async def get_tools_by_names( return {} -def _override_row_to_model(row: Any) -> ToolPolicyOverrideRow: - """Convert a Prisma override row to ToolPolicyOverrideRow.""" - model_dump = getattr(row, "model_dump", None) - if callable(model_dump): - row = model_dump() - elif not isinstance(row, dict): - row = { - k: getattr(row, k, None) - for k in ( - "override_id", - "tool_name", - "team_id", - "key_hash", - "call_policy", - "key_alias", - "created_at", - "updated_at", - ) - } - - def _norm(s: Optional[str]) -> Optional[str]: - if s is None or s == _TOOL_OVERRIDE_ANY: - return None - return s or None - - return ToolPolicyOverrideRow( - override_id=row.get("override_id", ""), - tool_name=row.get("tool_name", ""), - team_id=_norm(row.get("team_id")), - key_hash=_norm(row.get("key_hash")), - call_policy=row.get("call_policy", "blocked"), - key_alias=row.get("key_alias"), - created_at=row.get("created_at"), - updated_at=row.get("updated_at"), - ) - - async def list_overrides_for_tool( prisma_client: "PrismaClient", tool_name: str, ) -> List[ToolPolicyOverrideRow]: - """Return all policy overrides for a tool.""" + """ + Return override-like rows for a tool by finding object permissions that have + this tool in blocked_tools, then resolving each permission to key/team scope for display. + """ + out: List[ToolPolicyOverrideRow] = [] try: - rows = await prisma_client.db.litellm_toolpolicyoverridetable.find_many( - where={"tool_name": tool_name}, - order={"created_at": "desc"}, + perms = await prisma_client.db.litellm_objectpermissiontable.find_many( + where={"blocked_tools": {"has": tool_name}}, + include={ + "verification_tokens": True, + "teams": True, + }, ) - return [_override_row_to_model(row) for row in rows] + for perm in perms: + op_id = getattr(perm, "object_permission_id", None) or "" + tokens = getattr(perm, "verification_tokens", []) or [] + teams = getattr(perm, "teams", []) or [] + for t in tokens: + out.append( + ToolPolicyOverrideRow( + override_id=op_id, + tool_name=tool_name, + team_id=None, + key_hash=getattr(t, "token", None), + call_policy="blocked", + key_alias=getattr(t, "key_alias", None), + created_at=None, + updated_at=None, + ) + ) + for team in teams: + out.append( + ToolPolicyOverrideRow( + override_id=op_id, + tool_name=tool_name, + team_id=getattr(team, "team_id", None), + key_hash=None, + call_policy="blocked", + key_alias=getattr(team, "team_alias", None), + created_at=None, + updated_at=None, + ) + ) + return out except Exception as e: verbose_proxy_logger.error( "tool_registry_writer list_overrides_for_tool error: %s", e @@ -270,128 +267,61 @@ async def list_overrides_for_tool( return [] -async def upsert_tool_policy_override( +async def _get_merged_blocked_tools( prisma_client: "PrismaClient", - tool_name: str, - call_policy: ToolCallPolicy, - team_id: Optional[str] = None, - key_hash: Optional[str] = None, - key_alias: Optional[str] = None, - updated_by: Optional[str] = None, -) -> Optional[ToolPolicyOverrideRow]: - """Create or update a per-(tool, team, key) policy override.""" - try: - _team = (team_id or "").strip() or _TOOL_OVERRIDE_ANY - _key = (key_hash or "").strip() or _TOOL_OVERRIDE_ANY - _updated_by = updated_by or "system" - now = datetime.now(timezone.utc) - table = prisma_client.db.litellm_toolpolicyoverridetable - await table.upsert( - where={ - "tool_name_team_id_key_hash": { - "tool_name": tool_name, - "team_id": _team, - "key_hash": _key, - } - }, - data={ - "create": { - "override_id": str(uuid.uuid4()), - "tool_name": tool_name, - "team_id": _team, - "key_hash": _key, - "call_policy": call_policy, - "key_alias": key_alias, - "created_by": _updated_by, - "updated_by": _updated_by, - "created_at": now, - "updated_at": now, - }, - "update": { - "call_policy": call_policy, - "key_alias": key_alias, - "updated_by": _updated_by, - "updated_at": now, - }, - }, - ) - row = await table.find_unique( - where={ - "tool_name_team_id_key_hash": { - "tool_name": tool_name, - "team_id": _team, - "key_hash": _key, - } - } - ) - if row is None: - return None - return _override_row_to_model(row) - except Exception as e: - verbose_proxy_logger.error( - "tool_registry_writer upsert_tool_policy_override error: %s", e - ) - return None - - -async def delete_tool_policy_override( - prisma_client: "PrismaClient", - tool_name: str, - team_id: Optional[str] = None, - key_hash: Optional[str] = None, -) -> bool: - """Remove a policy override. Exactly one of team_id or key_hash should be set for a specific override.""" - try: - _team = (team_id or "").strip() or _TOOL_OVERRIDE_ANY - _key = (key_hash or "").strip() or _TOOL_OVERRIDE_ANY - result = await prisma_client.db.litellm_toolpolicyoverridetable.delete_many( - where={ - "tool_name": tool_name, - "team_id": _team, - "key_hash": _key, - } - ) - return (result or 0) > 0 - except Exception as e: - verbose_proxy_logger.error( - "tool_registry_writer delete_tool_policy_override error: %s", e - ) - return False + object_permission_id: Optional[str], + team_object_permission_id: Optional[str], +) -> set: + """Return union of blocked_tools from key and team object permissions.""" + blocked: set = set() + for op_id in (object_permission_id, team_object_permission_id): + if not op_id or not op_id.strip(): + continue + try: + row = await prisma_client.db.litellm_objectpermissiontable.find_unique( + where={"object_permission_id": op_id.strip()}, + select={"blocked_tools": True}, + ) + if row is not None and getattr(row, "blocked_tools", None): + blocked.update(row.blocked_tools) + except Exception as e: + verbose_proxy_logger.debug( + "tool_registry_writer _get_merged_blocked_tools error for %s: %s", + op_id, + e, + ) + return blocked async def get_effective_policies( prisma_client: "PrismaClient", tool_names: List[str], - team_id: Optional[str], - key_hash: Optional[str], + object_permission_id: Optional[str] = None, + team_object_permission_id: Optional[str] = None, ) -> Dict[str, str]: """ - Return effective call_policy per tool: override for (tool, team_id, key_hash) if present, - otherwise global policy from LiteLLM_ToolTable. + Return effective call_policy per tool: if tool is in key/team object permission + blocked_tools then "blocked", otherwise global policy from LiteLLM_ToolTable. """ if not tool_names: return {} - _team = (team_id or "").strip() or _TOOL_OVERRIDE_ANY - _key = (key_hash or "").strip() or _TOOL_OVERRIDE_ANY try: - # Fetch overrides for (tool_name, team_id, key_hash) - exact match - overrides = await prisma_client.db.litellm_toolpolicyoverridetable.find_many( - where={ - "tool_name": {"in": tool_names}, - "team_id": _team, - "key_hash": _key, - } + blocked = await _get_merged_blocked_tools( + prisma_client=prisma_client, + object_permission_id=object_permission_id, + team_object_permission_id=team_object_permission_id, ) - override_map = {row.tool_name: row.call_policy for row in overrides} - # Global policies for tools that have no override - missing = [n for n in tool_names if n not in override_map] - if not missing: - return override_map - global_map = await get_tools_by_names( - prisma_client=prisma_client, tool_names=missing - ) - override_map.update(global_map) - return override_map + result: Dict[str, str] = {} + for name in tool_names: + if name in blocked: + result[name] = "blocked" + missing = [n for n in tool_names if n not in result] + if missing: + global_map = await get_tools_by_names( + prisma_client=prisma_client, tool_names=missing + ) + result.update(global_map) + return result except Exception as e: verbose_proxy_logger.error( "tool_registry_writer get_effective_policies error: %s", e @@ -399,25 +329,28 @@ async def get_effective_policies( return {} -def _effective_cache_suffix(team_id: Optional[str], key_hash: Optional[str]) -> str: - """Cache key suffix so different request contexts (team/key) get correct policies.""" - return f":{team_id or ''}:{key_hash or ''}" +def _effective_cache_suffix( + object_permission_id: Optional[str], + team_object_permission_id: Optional[str], +) -> str: + """Cache key suffix so different request contexts get correct policies.""" + return f":{object_permission_id or ''}:{team_object_permission_id or ''}" async def get_tool_policies_cached( tool_names: List[str], cache: DualCache, prisma_client: Optional["PrismaClient"], - team_id: Optional[str] = None, - key_hash: Optional[str] = None, + object_permission_id: Optional[str] = None, + team_object_permission_id: Optional[str] = None, ) -> Dict[str, str]: """ - Return effective call_policy per tool (override for team/key if present, else global). - Cache-first; cache key includes team_id and key_hash when provided. + Return effective call_policy per tool (blocked if in object permission blocked_tools, + else global). Cache-first; cache key includes object_permission_id(s). """ if not tool_names: return {} - suffix = _effective_cache_suffix(team_id, key_hash) + suffix = _effective_cache_suffix(object_permission_id, team_object_permission_id) result: Dict[str, str] = {} cache_misses: List[str] = [] for name in tool_names: @@ -429,12 +362,12 @@ async def get_tool_policies_cached( cache_misses.append(name) if cache_misses and prisma_client is not None: try: - if team_id is not None or key_hash is not None: + if object_permission_id or team_object_permission_id: fetched = await get_effective_policies( prisma_client=prisma_client, tool_names=cache_misses, - team_id=team_id, - key_hash=key_hash, + object_permission_id=object_permission_id, + team_object_permission_id=team_object_permission_id, ) else: fetched = await get_tools_by_names( @@ -457,3 +390,66 @@ async def get_tool_policies_cached( "tool_registry_writer get_tool_policies_cached error: %s", e ) return result + + +async def add_tool_to_object_permission_blocked( + prisma_client: "PrismaClient", + object_permission_id: str, + tool_name: str, +) -> bool: + """Add tool_name to the permission's blocked_tools if not already present.""" + if not object_permission_id or not tool_name: + return False + try: + row = await prisma_client.db.litellm_objectpermissiontable.find_unique( + where={"object_permission_id": object_permission_id}, + select={"blocked_tools": True}, + ) + if row is None: + return False + current = list(getattr(row, "blocked_tools", []) or []) + if tool_name in current: + return True + current.append(tool_name) + await prisma_client.db.litellm_objectpermissiontable.update( + where={"object_permission_id": object_permission_id}, + data={"blocked_tools": current}, + ) + return True + except Exception as e: + verbose_proxy_logger.error( + "tool_registry_writer add_tool_to_object_permission_blocked error: %s", e + ) + return False + + +async def remove_tool_from_object_permission_blocked( + prisma_client: "PrismaClient", + object_permission_id: str, + tool_name: str, +) -> bool: + """Remove tool_name from the permission's blocked_tools. Returns False if tool was not in list.""" + if not object_permission_id or not tool_name: + return False + try: + row = await prisma_client.db.litellm_objectpermissiontable.find_unique( + where={"object_permission_id": object_permission_id}, + select={"blocked_tools": True}, + ) + if row is None: + return False + current = list(getattr(row, "blocked_tools", []) or []) + if tool_name not in current: + return False + current = [t for t in current if t != tool_name] + await prisma_client.db.litellm_objectpermissiontable.update( + where={"object_permission_id": object_permission_id}, + data={"blocked_tools": current}, + ) + return True + except Exception as e: + verbose_proxy_logger.error( + "tool_registry_writer remove_tool_from_object_permission_blocked error: %s", + e, + ) + return False diff --git a/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py index 08b28824cc0..c8763f162c1 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py @@ -43,22 +43,31 @@ if TYPE_CHECKING: GUARDRAIL_NAME = "tool_policy" -def _get_request_team_and_key( +def _get_request_object_permission_ids( request_data: dict, ) -> Tuple[Optional[str], Optional[str]]: - """Extract team_id and key hash from request_data (litellm_metadata or metadata).""" + """Extract object_permission_id and team_object_permission_id from request_data.""" if not request_data: return None, None for key in ("litellm_metadata", "metadata"): meta = request_data.get(key) if not isinstance(meta, dict): continue - team_id = meta.get("user_api_key_team_id") - key_hash = meta.get("user_api_key_api_key") or meta.get("user_api_key") - if team_id is not None or key_hash is not None: + auth = meta.get("user_api_key_auth") + if auth is not None and hasattr(auth, "object_permission_id"): + key_op = getattr(auth, "object_permission_id", None) + team_op = getattr(auth, "team_object_permission_id", None) + if key_op is not None or team_op is not None: + return ( + str(key_op).strip() if key_op else None, + str(team_op).strip() if team_op else None, + ) + key_op = meta.get("user_api_key_object_permission_id") + team_op = meta.get("user_api_key_team_object_permission_id") + if key_op is not None or team_op is not None: return ( - str(team_id).strip() if team_id else None, - str(key_hash).strip() if key_hash else None, + str(key_op).strip() if key_op else None, + str(team_op).strip() if team_op else None, ) return None, None @@ -165,8 +174,14 @@ class ToolPolicyGuardrail(CustomGuardrail): }, ) - team_id, key_hash = _get_request_team_and_key(request_data) - policy_map = await self._get_policies_cached(tool_names, team_id, key_hash) + object_permission_id, team_object_permission_id = ( + _get_request_object_permission_ids(request_data) + ) + policy_map = await self._get_policies_cached( + tool_names, + object_permission_id=object_permission_id, + team_object_permission_id=team_object_permission_id, + ) blocked = [name for name in tool_names if policy_map.get(name) == "blocked"] if blocked: verbose_proxy_logger.warning( @@ -186,12 +201,12 @@ class ToolPolicyGuardrail(CustomGuardrail): async def _get_policies_cached( self, tool_names: List[str], - team_id: Optional[str] = None, - key_hash: Optional[str] = None, + object_permission_id: Optional[str] = None, + team_object_permission_id: Optional[str] = None, ) -> Dict[str, str]: """ - Fetch effective call_policy (override for team/key if present, else global) - via shared cache to avoid DB in hot path. + Fetch effective call_policy (blocked if in object permission blocked_tools, + else global) via shared cache to avoid DB in hot path. """ from litellm.proxy.db.tool_registry_writer import \ get_tool_policies_cached @@ -204,6 +219,6 @@ class ToolPolicyGuardrail(CustomGuardrail): tool_names=tool_names, cache=user_api_key_cache, prisma_client=prisma_client, - team_id=team_id, - key_hash=key_hash, + object_permission_id=object_permission_id, + team_object_permission_id=team_object_permission_id, ) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index c3d58df83ea..76d715fc350 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -1053,6 +1053,12 @@ async def add_litellm_data_to_request( # noqa: PLR0915 data[_metadata_variable_name]["user_api_key_team_metadata"] = ( user_api_key_dict.team_metadata ) + data[_metadata_variable_name]["user_api_key_object_permission_id"] = ( + getattr(user_api_key_dict, "object_permission_id", None) + ) + data[_metadata_variable_name]["user_api_key_team_object_permission_id"] = ( + getattr(user_api_key_dict, "team_object_permission_id", None) + ) data[_metadata_variable_name]["headers"] = _headers data[_metadata_variable_name]["endpoint"] = str(request.url) diff --git a/litellm/proxy/management_endpoints/tool_management_endpoints.py b/litellm/proxy/management_endpoints/tool_management_endpoints.py index bc8fc9165bb..84e8a12138a 100644 --- a/litellm/proxy/management_endpoints/tool_management_endpoints.py +++ b/litellm/proxy/management_endpoints/tool_management_endpoints.py @@ -8,21 +8,24 @@ GET /v1/tool/{tool_name} - Get a single tool's details POST /v1/tool/policy - Update the call_policy for a tool """ -from typing import Optional +from datetime import datetime, timezone +from typing import TYPE_CHECKING, Any, List, Optional from fastapi import APIRouter, Depends, HTTPException, Query +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient + from litellm._logging import verbose_proxy_logger from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.types.tool_management import ( - LiteLLM_ToolTableRow, - ToolCallPolicy, - ToolDetailResponse, - ToolListResponse, - ToolPolicyUpdateRequest, - ToolPolicyUpdateResponse, -) +from litellm.types.tool_management import (LiteLLM_ToolTableRow, + ToolCallPolicy, ToolDetailResponse, + ToolListResponse, + ToolPolicyUpdateRequest, + ToolPolicyUpdateResponse, + ToolUsageLogEntry, + ToolUsageLogsResponse) router = APIRouter() @@ -43,7 +46,8 @@ async def list_tools( Parameters: - call_policy: Optional filter — one of "trusted", "untrusted", "dual_llm", "blocked" """ - from litellm.proxy.db.tool_registry_writer import list_tools as db_list_tools + from litellm.proxy.db.tool_registry_writer import \ + list_tools as db_list_tools from litellm.proxy.proxy_server import prisma_client if prisma_client is None: @@ -101,6 +105,154 @@ async def get_tool_detail( raise HTTPException(status_code=500, detail=str(e)) +def _input_snippet_for_tool_log(sl: Any, max_len: int = 200) -> Optional[str]: + """Short snippet from messages or proxy_server_request for tool usage log row.""" + if sl is None: + return None + messages = getattr(sl, "messages", None) + if messages is not None: + s = _snippet_str(messages, max_len) + if s: + return s + psr = getattr(sl, "proxy_server_request", None) + if not psr: + return None + if isinstance(psr, str): + import json + try: + psr = json.loads(psr) + except Exception: + return _snippet_str(psr, max_len) + if isinstance(psr, dict): + msgs = psr.get("messages") + if msgs is None and isinstance(psr.get("body"), dict): + msgs = psr["body"].get("messages") + s = _snippet_str(msgs, max_len) + if s: + return s + return _snippet_str(psr, max_len) + + +def _snippet_str(text: Any, max_len: int = 200) -> Optional[str]: + if text is None: + return None + if isinstance(text, str): + s = text + elif isinstance(text, list): + parts = [] + for item in text: + if isinstance(item, dict) and "content" in item: + c = item["content"] + parts.append(c if isinstance(c, str) else str(c)) + else: + parts.append(str(item)) + s = " ".join(parts) + else: + s = str(text) + if not s or s == "{}": + return None + return (s[:max_len] + "...") if len(s) > max_len else s + + +@router.get( + "/v1/tool/{tool_name:path}/logs", + tags=["tool management"], + dependencies=[Depends(user_api_key_auth)], + response_model=ToolUsageLogsResponse, +) +async def get_tool_usage_logs( + tool_name: str, + page: int = Query(1, ge=1), + page_size: int = Query(50, ge=1, le=100), + start_date: Optional[str] = Query(None, description="YYYY-MM-DD"), + end_date: Optional[str] = Query(None, description="YYYY-MM-DD"), + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Return paginated spend logs for requests that used this tool (from SpendLogToolIndex). + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException( + status_code=500, detail=CommonProxyErrors.db_not_connected_error.value + ) + + try: + where: dict = {"tool_name": tool_name} + if start_date or end_date: + start_time_filter: Optional[datetime] = None + end_time_filter: Optional[datetime] = None + if start_date: + try: + start_time_filter = datetime.strptime( + start_date + "T00:00:00", "%Y-%m-%dT%H:%M:%S" + ).replace(tzinfo=timezone.utc) + except ValueError: + pass + if end_date: + try: + end_time_filter = datetime.strptime( + end_date + "T23:59:59", "%Y-%m-%dT%H:%M:%S" + ).replace(tzinfo=timezone.utc) + except ValueError: + pass + if start_time_filter is not None or end_time_filter is not None: + where["start_time"] = {} + if start_time_filter is not None: + where["start_time"]["gte"] = start_time_filter + if end_time_filter is not None: + where["start_time"]["lte"] = end_time_filter + + total = await prisma_client.db.litellm_spendlogtoolindex.count(where=where) + index_rows = await prisma_client.db.litellm_spendlogtoolindex.find_many( + where=where, + order={"start_time": "desc"}, + skip=(page - 1) * page_size, + take=page_size, + ) + request_ids = [r.request_id for r in index_rows] + if not request_ids: + return ToolUsageLogsResponse( + logs=[], total=total, page=page, page_size=page_size + ) + + spend_logs = await prisma_client.db.litellm_spendlogs.find_many( + where={"request_id": {"in": request_ids}} + ) + log_by_id = {s.request_id: s for s in spend_logs} + + logs_out: List[ToolUsageLogEntry] = [] + for r in index_rows: + sl = log_by_id.get(r.request_id) + if not sl: + continue + ts = ( + sl.startTime.isoformat() + if hasattr(sl.startTime, "isoformat") + else str(sl.startTime) + ) + logs_out.append( + ToolUsageLogEntry( + id=sl.request_id, + timestamp=ts, + model=getattr(sl, "model", None) or None, + spend=getattr(sl, "spend", None), + total_tokens=getattr(sl, "total_tokens", None), + input_snippet=_input_snippet_for_tool_log(sl), + ) + ) + + return ToolUsageLogsResponse( + logs=logs_out, total=total, page=page, page_size=page_size + ) + except HTTPException: + raise + except Exception as e: + verbose_proxy_logger.exception("Error getting tool usage logs: %s", e) + raise HTTPException(status_code=500, detail=str(e)) + + @router.get( "/v1/tool/{tool_name:path}", tags=["tool management"], @@ -139,6 +291,66 @@ async def get_tool( raise HTTPException(status_code=500, detail=str(e)) +async def _resolve_key_hash_to_object_permission_id( + prisma_client: "PrismaClient", + key_hash: str, +) -> Optional[str]: + """Resolve key (hash or raw) to object_permission_id; create permission if key has none.""" + from litellm.proxy.proxy_server import hash_token + + hashed = key_hash if "sk-" not in (key_hash or "") else hash_token(key_hash) + if not hashed: + return None + row = await prisma_client.db.litellm_verificationtoken.find_unique( + where={"token": hashed}, + select={"object_permission_id": True}, + ) + if row is None: + return None + op_id = getattr(row, "object_permission_id", None) + if op_id: + return op_id + # Create new object permission and assign to key + import uuid as _uuid + new_id = str(_uuid.uuid4()) + await prisma_client.db.litellm_objectpermissiontable.create( + data={"object_permission_id": new_id, "blocked_tools": []} + ) + await prisma_client.db.litellm_verificationtoken.update( + where={"token": hashed}, + data={"object_permission_id": new_id}, + ) + return new_id + + +async def _resolve_team_id_to_object_permission_id( + prisma_client: "PrismaClient", + team_id: str, +) -> Optional[str]: + """Resolve team_id to object_permission_id; create permission if team has none.""" + if not team_id or not team_id.strip(): + return None + row = await prisma_client.db.litellm_teamtable.find_unique( + where={"team_id": team_id.strip()}, + select={"object_permission_id": True}, + ) + if row is None: + return None + op_id = getattr(row, "object_permission_id", None) + if op_id: + return op_id + import uuid as _uuid + new_id = str(_uuid.uuid4()) + await prisma_client.db.litellm_objectpermissiontable.create( + data={"object_permission_id": new_id, "blocked_tools": []} + ) + await prisma_client.db.litellm_teamtable.update( + where={"team_id": team_id.strip()}, + data={"object_permission_id": new_id}, + ) + return new_id + + @router.post( "/v1/tool/policy", tags=["tool management"], @@ -164,9 +376,10 @@ async def update_tool_policy( that tool_call for the relevant scope. """ from litellm.proxy.db.tool_registry_writer import ( - update_tool_policy as db_update_tool_policy, - ) - from litellm.proxy.db.tool_registry_writer import upsert_tool_policy_override + add_tool_to_object_permission_blocked, + remove_tool_from_object_permission_blocked) + from litellm.proxy.db.tool_registry_writer import \ + update_tool_policy as db_update_tool_policy from litellm.proxy.proxy_server import prisma_client if prisma_client is None: @@ -176,26 +389,47 @@ async def update_tool_policy( try: if data.team_id is not None or data.key_hash is not None: - override = await upsert_tool_policy_override( - prisma_client=prisma_client, - tool_name=data.tool_name, - call_policy=data.call_policy, - team_id=data.team_id, - key_hash=data.key_hash, - key_alias=data.key_alias, - updated_by=user_api_key_dict.user_id, - ) - if override is None: + if data.team_id is not None and data.key_hash is not None: + raise HTTPException( + status_code=400, + detail="Provide either team_id or key_hash, not both", + ) + if data.key_hash is not None: + op_id = await _resolve_key_hash_to_object_permission_id( + prisma_client, data.key_hash + ) + else: + op_id = await _resolve_team_id_to_object_permission_id( + prisma_client, data.team_id or "" + ) + if op_id is None: + raise HTTPException( + status_code=404, + detail="Key or team not found for the given identifier", + ) + if data.call_policy == "blocked": + ok = await add_tool_to_object_permission_blocked( + prisma_client=prisma_client, + object_permission_id=op_id, + tool_name=data.tool_name, + ) + else: + ok = await remove_tool_from_object_permission_blocked( + prisma_client=prisma_client, + object_permission_id=op_id, + tool_name=data.tool_name, + ) + if not ok: raise HTTPException( status_code=500, detail=f"Failed to update policy override for tool '{data.tool_name}'", ) return ToolPolicyUpdateResponse( - tool_name=override.tool_name, - call_policy=override.call_policy, + tool_name=data.tool_name, + call_policy=data.call_policy, updated=True, - team_id=override.team_id, - key_hash=override.key_hash, + team_id=data.team_id, + key_hash=data.key_hash, ) updated = await db_update_tool_policy( prisma_client=prisma_client, @@ -232,10 +466,11 @@ async def delete_tool_policy_override( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ - Remove a policy override for a tool. Specify the override by team_id and/or key_hash - (must match the override that was created; use empty string for unscoped dimension). + Remove a policy override for a tool. Specify the override by team_id or key_hash + (exactly one required). """ - from litellm.proxy.db.tool_registry_writer import delete_tool_policy_override + from litellm.proxy.db.tool_registry_writer import \ + remove_tool_from_object_permission_blocked from litellm.proxy.proxy_server import prisma_client if prisma_client is None: @@ -247,12 +482,29 @@ async def delete_tool_policy_override( status_code=400, detail="At least one of team_id or key_hash is required to identify the override", ) + if team_id is not None and key_hash is not None: + raise HTTPException( + status_code=400, + detail="Provide either team_id or key_hash, not both", + ) try: - deleted = await delete_tool_policy_override( + if key_hash is not None: + op_id = await _resolve_key_hash_to_object_permission_id( + prisma_client, key_hash + ) + else: + op_id = await _resolve_team_id_to_object_permission_id( + prisma_client, team_id or "" + ) + if op_id is None: + raise HTTPException( + status_code=404, + detail="Key or team not found for the given identifier", + ) + deleted = await remove_tool_from_object_permission_blocked( prisma_client=prisma_client, + object_permission_id=op_id, tool_name=tool_name, - team_id=team_id, - key_hash=key_hash, ) if not deleted: raise HTTPException( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index f0b1e66818c..cbafe5bb390 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -14,95 +14,47 @@ import time import traceback import warnings from datetime import datetime, timedelta, timezone -from typing import ( - TYPE_CHECKING, - Any, - AsyncGenerator, - Dict, - List, - Literal, - Optional, - Set, - Tuple, - Union, - cast, - get_args, - get_origin, - get_type_hints, -) +from typing import (TYPE_CHECKING, Any, AsyncGenerator, Dict, List, Literal, + Optional, Set, Tuple, Union, cast, get_args, get_origin, + get_type_hints) import anyio from pydantic import BaseModel, Json from litellm._uuid import uuid from litellm.constants import ( - AIOHTTP_CONNECTOR_LIMIT, - AIOHTTP_CONNECTOR_LIMIT_PER_HOST, - AIOHTTP_KEEPALIVE_TIMEOUT, - AIOHTTP_NEEDS_CLEANUP_CLOSED, - AIOHTTP_TTL_DNS_CACHE, - AUDIO_SPEECH_CHUNK_SIZE, - BASE_MCP_ROUTE, - DEFAULT_MAX_RECURSE_DEPTH, - DEFAULT_SHARED_HEALTH_CHECK_LOCK_TTL, - DEFAULT_SHARED_HEALTH_CHECK_TTL, - DEFAULT_SLACK_ALERTING_THRESHOLD, + AIOHTTP_CONNECTOR_LIMIT, AIOHTTP_CONNECTOR_LIMIT_PER_HOST, + AIOHTTP_KEEPALIVE_TIMEOUT, AIOHTTP_NEEDS_CLEANUP_CLOSED, + AIOHTTP_TTL_DNS_CACHE, AUDIO_SPEECH_CHUNK_SIZE, BASE_MCP_ROUTE, + DEFAULT_MAX_RECURSE_DEPTH, DEFAULT_SHARED_HEALTH_CHECK_LOCK_TTL, + DEFAULT_SHARED_HEALTH_CHECK_TTL, DEFAULT_SLACK_ALERTING_THRESHOLD, LITELLM_EMBEDDING_PROVIDERS_SUPPORTING_INPUT_ARRAY_OF_TOKENS, - LITELLM_SETTINGS_SAFE_DB_OVERRIDES, - LITELLM_UI_ALLOW_HEADERS, -) -from litellm.litellm_core_utils.litellm_logging import ( - _init_custom_logger_compatible_class, -) + LITELLM_SETTINGS_SAFE_DB_OVERRIDES, LITELLM_UI_ALLOW_HEADERS) +from litellm.litellm_core_utils.litellm_logging import \ + _init_custom_logger_compatible_class from litellm.litellm_core_utils.safe_json_dumps import safe_dumps -from litellm.proxy._types import ( - CallbackDelete, - CallInfo, - CommonProxyErrors, - ConfigFieldDelete, - ConfigFieldInfo, - ConfigFieldUpdate, - ConfigGeneralSettings, - ConfigList, - ConfigYAML, - EnterpriseLicenseData, - FieldDetail, - InvitationClaim, - InvitationDelete, - InvitationModel, - InvitationNew, - InvitationUpdate, - Litellm_EntityType, - LiteLLM_JWTAuth, - LiteLLM_TeamTable, - LiteLLM_UserTable, - LitellmUserRoles, - PassThroughGenericEndpoint, - ProxyErrorTypes, - ProxyException, - RoleBasedPermissions, - SpecialModelNames, - SupportedDBObjectType, - TeamDefaultSettings, - TokenCountRequest, - TransformRequestBody, - UserAPIKeyAuth, -) +from litellm.proxy._types import (CallbackDelete, CallInfo, CommonProxyErrors, + ConfigFieldDelete, ConfigFieldInfo, + ConfigFieldUpdate, ConfigGeneralSettings, + ConfigList, ConfigYAML, + EnterpriseLicenseData, FieldDetail, + InvitationClaim, InvitationDelete, + InvitationModel, InvitationNew, + InvitationUpdate, Litellm_EntityType, + LiteLLM_JWTAuth, LiteLLM_TeamTable, + LiteLLM_UserTable, LitellmUserRoles, + PassThroughGenericEndpoint, ProxyErrorTypes, + ProxyException, RoleBasedPermissions, + SpecialModelNames, SupportedDBObjectType, + TeamDefaultSettings, TokenCountRequest, + TransformRequestBody, UserAPIKeyAuth) from litellm.proxy.common_utils.callback_utils import ( - normalize_callback_names, - process_callback, -) + normalize_callback_names, process_callback) from litellm.proxy.common_utils.realtime_utils import _realtime_request_body -from litellm.types.utils import ( - ModelResponse, - ModelResponseStream, - TextCompletionResponse, - TokenCountResponse, -) -from litellm.utils import ( - _invalidate_model_cost_lowercase_map, - load_credentials_from_list, -) +from litellm.types.utils import (ModelResponse, ModelResponseStream, + TextCompletionResponse, TokenCountResponse) +from litellm.utils import (_invalidate_model_cost_lowercase_map, + load_credentials_from_list) if TYPE_CHECKING: from aiohttp import ClientSession @@ -199,350 +151,266 @@ from litellm import Router from litellm._logging import verbose_proxy_logger, verbose_router_logger from litellm.caching.caching import DualCache, RedisCache from litellm.caching.redis_cluster_cache import RedisClusterCache -from litellm.constants import ( - _REALTIME_BODY_CACHE_SIZE, - APSCHEDULER_COALESCE, - APSCHEDULER_MAX_INSTANCES, - APSCHEDULER_MISFIRE_GRACE_TIME, - APSCHEDULER_REPLACE_EXISTING, - DAYS_IN_A_MONTH, - DEFAULT_HEALTH_CHECK_INTERVAL, - DEFAULT_MODEL_CREATED_AT_TIME, - LITELLM_PROXY_ADMIN_NAME, - PROMETHEUS_FALLBACK_STATS_SEND_TIME_HOURS, - PROXY_BATCH_POLLING_INTERVAL, - PROXY_BATCH_WRITE_AT, - PROXY_BUDGET_RESCHEDULER_MAX_TIME, - PROXY_BUDGET_RESCHEDULER_MIN_TIME, -) +from litellm.constants import (_REALTIME_BODY_CACHE_SIZE, APSCHEDULER_COALESCE, + APSCHEDULER_MAX_INSTANCES, + APSCHEDULER_MISFIRE_GRACE_TIME, + APSCHEDULER_REPLACE_EXISTING, DAYS_IN_A_MONTH, + DEFAULT_HEALTH_CHECK_INTERVAL, + DEFAULT_MODEL_CREATED_AT_TIME, + LITELLM_PROXY_ADMIN_NAME, + PROMETHEUS_FALLBACK_STATS_SEND_TIME_HOURS, + PROXY_BATCH_POLLING_INTERVAL, + PROXY_BATCH_WRITE_AT, + PROXY_BUDGET_RESCHEDULER_MAX_TIME, + PROXY_BUDGET_RESCHEDULER_MIN_TIME) from litellm.exceptions import RejectedRequestError from litellm.integrations.custom_guardrail import ModifyResponseException from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting from litellm.litellm_core_utils.core_helpers import ( - _get_parent_otel_span_from_kwargs, - get_litellm_metadata_from_kwargs, -) + _get_parent_otel_span_from_kwargs, get_litellm_metadata_from_kwargs) from litellm.litellm_core_utils.credential_accessor import CredentialAccessor -from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker -from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.litellm_core_utils.litellm_logging import \ + Logging as LiteLLMLoggingObj +from litellm.litellm_core_utils.sensitive_data_masker import \ + SensitiveDataMasker +from litellm.llms.custom_httpx.http_handler import (AsyncHTTPHandler, + HTTPHandler) from litellm.llms.vertex_ai.vertex_llm_base import VertexBase -from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( - router as mcp_discoverable_endpoints_router, -) -from litellm.proxy._experimental.mcp_server.rest_endpoints import ( - router as mcp_rest_endpoints_router, -) +from litellm.proxy._experimental.mcp_server.discoverable_endpoints import \ + router as mcp_discoverable_endpoints_router +from litellm.proxy._experimental.mcp_server.rest_endpoints import \ + router as mcp_rest_endpoints_router from litellm.proxy._experimental.mcp_server.server import app as mcp_app -from litellm.proxy._experimental.mcp_server.tool_registry import ( - global_mcp_tool_registry, -) +from litellm.proxy._experimental.mcp_server.tool_registry import \ + global_mcp_tool_registry from litellm.proxy._types import * from litellm.proxy.agent_endpoints.a2a_endpoints import router as a2a_router from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry -from litellm.proxy.agent_endpoints.endpoints import router as agent_endpoints_router +from litellm.proxy.agent_endpoints.endpoints import \ + router as agent_endpoints_router from litellm.proxy.agent_endpoints.model_list_helpers import ( - append_agents_to_model_group, - append_agents_to_model_info, -) -from litellm.proxy.analytics_endpoints.analytics_endpoints import ( - router as analytics_router, -) -from litellm.proxy.anthropic_endpoints.claude_code_endpoints import ( - claude_code_marketplace_router, -) -from litellm.proxy.anthropic_endpoints.endpoints import router as anthropic_router -from litellm.proxy.anthropic_endpoints.skills_endpoints import ( - router as anthropic_skills_router, -) -from litellm.proxy.auth.auth_checks import ( - ExperimentalUIJWTToken, - get_team_object, - log_db_metrics, -) + append_agents_to_model_group, append_agents_to_model_info) +from litellm.proxy.analytics_endpoints.analytics_endpoints import \ + router as analytics_router +from litellm.proxy.anthropic_endpoints.claude_code_endpoints import \ + claude_code_marketplace_router +from litellm.proxy.anthropic_endpoints.endpoints import \ + router as anthropic_router +from litellm.proxy.anthropic_endpoints.skills_endpoints import \ + router as anthropic_skills_router +from litellm.proxy.auth.auth_checks import (ExperimentalUIJWTToken, + get_team_object, log_db_metrics) from litellm.proxy.auth.auth_utils import check_response_size_is_safe from litellm.proxy.auth.handle_jwt import JWTHandler from litellm.proxy.auth.litellm_license import LicenseCheck -from litellm.proxy.auth.model_checks import ( - get_all_fallbacks, - get_complete_model_list, - get_key_models, - get_mcp_server_ids, - get_team_models, -) +from litellm.proxy.auth.model_checks import (get_all_fallbacks, + get_complete_model_list, + get_key_models, + get_mcp_server_ids, + get_team_models) from litellm.proxy.auth.user_api_key_auth import ( - _fetch_global_spend_with_event_coordination, - user_api_key_auth, - user_api_key_auth_websocket, -) + _fetch_global_spend_with_event_coordination, user_api_key_auth, + user_api_key_auth_websocket) from litellm.proxy.batches_endpoints.endpoints import router as batches_router - ## Import All Misc routes here ## from litellm.proxy.caching_routes import router as caching_router from litellm.proxy.common_request_processing import ( - ProxyBaseLLMRequestProcessing, - create_response, -) -from litellm.proxy.common_utils.callback_utils import initialize_callbacks_on_proxy + ProxyBaseLLMRequestProcessing, create_response) +from litellm.proxy.common_utils.callback_utils import \ + initialize_callbacks_on_proxy from litellm.proxy.common_utils.debug_utils import init_verbose_loggers -from litellm.proxy.common_utils.debug_utils import router as debugging_endpoints_router +from litellm.proxy.common_utils.debug_utils import \ + router as debugging_endpoints_router from litellm.proxy.common_utils.encrypt_decrypt_utils import ( - decrypt_value_helper, - encrypt_value_helper, -) + decrypt_value_helper, encrypt_value_helper) from litellm.proxy.common_utils.html_forms.ui_login import build_ui_login_form from litellm.proxy.common_utils.http_parsing_utils import ( - _read_request_body, - _safe_get_request_headers, - check_file_size_under_limit, - get_form_data, -) + _read_request_body, _safe_get_request_headers, check_file_size_under_limit, + get_form_data) from litellm.proxy.common_utils.load_config_utils import ( - get_config_file_contents_from_gcs, - get_file_contents_from_s3, -) -from litellm.proxy.common_utils.openai_endpoint_utils import ( - remove_sensitive_info_from_deployment, -) + get_config_file_contents_from_gcs, get_file_contents_from_s3) +from litellm.proxy.common_utils.openai_endpoint_utils import \ + remove_sensitive_info_from_deployment from litellm.proxy.common_utils.proxy_state import ProxyState from litellm.proxy.common_utils.reset_budget_job import ResetBudgetJob from litellm.proxy.common_utils.swagger_utils import ERROR_RESPONSES -from litellm.proxy.container_endpoints.endpoints import router as container_router -from litellm.proxy.credential_endpoints.endpoints import router as credential_router -from litellm.proxy.db.db_transaction_queue.spend_log_cleanup import SpendLogCleanup +from litellm.proxy.container_endpoints.endpoints import \ + router as container_router +from litellm.proxy.credential_endpoints.endpoints import \ + router as credential_router +from litellm.proxy.db.db_transaction_queue.spend_log_cleanup import \ + SpendLogCleanup from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler from litellm.proxy.discovery_endpoints import ui_discovery_endpoints_router -from litellm.proxy.fine_tuning_endpoints.endpoints import router as fine_tuning_router -from litellm.proxy.fine_tuning_endpoints.endpoints import set_fine_tuning_config +from litellm.proxy.fine_tuning_endpoints.endpoints import \ + router as fine_tuning_router +from litellm.proxy.fine_tuning_endpoints.endpoints import \ + set_fine_tuning_config from litellm.proxy.google_endpoints.endpoints import router as google_router -from litellm.proxy.guardrails.guardrail_endpoints import router as guardrails_router -from litellm.proxy.guardrails.init_guardrails import ( - init_guardrails_v2, - initialize_guardrails, -) +from litellm.proxy.guardrails.guardrail_endpoints import \ + router as guardrails_router +from litellm.proxy.guardrails.init_guardrails import (init_guardrails_v2, + initialize_guardrails) from litellm.proxy.health_check import perform_health_check -from litellm.proxy.health_endpoints._health_endpoints import router as health_router -from litellm.proxy.hooks.model_max_budget_limiter import ( - _PROXY_VirtualKeyModelMaxBudgetLimiter, -) -from litellm.proxy.hooks.prompt_injection_detection import ( - _OPTIONAL_PromptInjectionDetection, -) +from litellm.proxy.health_endpoints._health_endpoints import \ + router as health_router +from litellm.proxy.hooks.model_max_budget_limiter import \ + _PROXY_VirtualKeyModelMaxBudgetLimiter +from litellm.proxy.hooks.prompt_injection_detection import \ + _OPTIONAL_PromptInjectionDetection from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger from litellm.proxy.image_endpoints.endpoints import router as image_router from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request -from litellm.proxy.management_endpoints.access_group_endpoints import ( - router as access_group_router, -) -from litellm.proxy.management_endpoints.budget_management_endpoints import ( - router as budget_management_router, -) -from litellm.proxy.management_endpoints.cache_settings_endpoints import ( - router as cache_settings_router, -) -from litellm.proxy.management_endpoints.callback_management_endpoints import ( - router as callback_management_endpoints_router, -) +from litellm.proxy.management_endpoints.access_group_endpoints import \ + router as access_group_router +from litellm.proxy.management_endpoints.budget_management_endpoints import \ + router as budget_management_router +from litellm.proxy.management_endpoints.cache_settings_endpoints import \ + router as cache_settings_router +from litellm.proxy.management_endpoints.callback_management_endpoints import \ + router as callback_management_endpoints_router from litellm.proxy.management_endpoints.common_utils import ( - _user_has_admin_privileges, - admin_can_invite_user, -) -from litellm.proxy.management_endpoints.compliance_endpoints import ( - router as compliance_router, -) -from litellm.proxy.management_endpoints.cost_tracking_settings import ( - router as cost_tracking_settings_router, -) -from litellm.proxy.management_endpoints.customer_endpoints import ( - router as customer_router, -) -from litellm.proxy.management_endpoints.fallback_management_endpoints import ( - router as fallback_management_router, -) -from litellm.proxy.management_endpoints.internal_user_endpoints import ( - router as internal_user_router, -) -from litellm.proxy.management_endpoints.internal_user_endpoints import ( - user_update, -) + _user_has_admin_privileges, admin_can_invite_user) +from litellm.proxy.management_endpoints.compliance_endpoints import \ + router as compliance_router +from litellm.proxy.management_endpoints.cost_tracking_settings import \ + router as cost_tracking_settings_router +from litellm.proxy.management_endpoints.customer_endpoints import \ + router as customer_router +from litellm.proxy.management_endpoints.fallback_management_endpoints import \ + router as fallback_management_router +from litellm.proxy.management_endpoints.internal_user_endpoints import \ + router as internal_user_router +from litellm.proxy.management_endpoints.internal_user_endpoints import \ + user_update from litellm.proxy.management_endpoints.key_management_endpoints import ( - delete_verification_tokens, - duration_in_seconds, - generate_key_helper_fn, -) -from litellm.proxy.management_endpoints.key_management_endpoints import ( - router as key_management_router, -) -from litellm.proxy.management_endpoints.mcp_management_endpoints import ( - router as mcp_management_router, -) -from litellm.proxy.management_endpoints.model_access_group_management_endpoints import ( - router as model_access_group_management_router, -) + delete_verification_tokens, duration_in_seconds, generate_key_helper_fn) +from litellm.proxy.management_endpoints.key_management_endpoints import \ + router as key_management_router +from litellm.proxy.management_endpoints.mcp_management_endpoints import \ + router as mcp_management_router +from litellm.proxy.management_endpoints.model_access_group_management_endpoints import \ + router as model_access_group_management_router from litellm.proxy.management_endpoints.model_management_endpoints import ( - _add_model_to_db, - _add_team_model_to_db, - _deduplicate_litellm_router_models, -) -from litellm.proxy.management_endpoints.model_management_endpoints import ( - router as model_management_router, -) -from litellm.proxy.management_endpoints.organization_endpoints import ( - router as organization_router, -) -from litellm.proxy.management_endpoints.policy_endpoints import router as policy_router -from litellm.proxy.management_endpoints.usage_endpoints import router as usage_ai_router -from litellm.proxy.management_endpoints.project_endpoints import ( - router as project_router, -) -from litellm.proxy.management_endpoints.router_settings_endpoints import ( - router as router_settings_router, -) + _add_model_to_db, _add_team_model_to_db, + _deduplicate_litellm_router_models) +from litellm.proxy.management_endpoints.model_management_endpoints import \ + router as model_management_router +from litellm.proxy.management_endpoints.organization_endpoints import \ + router as organization_router +from litellm.proxy.management_endpoints.policy_endpoints import \ + router as policy_router +from litellm.proxy.management_endpoints.project_endpoints import \ + router as project_router +from litellm.proxy.management_endpoints.router_settings_endpoints import \ + router as router_settings_router from litellm.proxy.management_endpoints.scim.scim_v2 import scim_router -from litellm.proxy.management_endpoints.tag_management_endpoints import ( - router as tag_management_router, -) -from litellm.proxy.management_endpoints.team_callback_endpoints import ( - router as team_callback_router, -) -from litellm.proxy.management_endpoints.team_endpoints import router as team_router +from litellm.proxy.management_endpoints.tag_management_endpoints import \ + router as tag_management_router +from litellm.proxy.management_endpoints.team_callback_endpoints import \ + router as team_callback_router +from litellm.proxy.management_endpoints.team_endpoints import \ + router as team_router from litellm.proxy.management_endpoints.team_endpoints import ( - update_team, - validate_membership, -) -from litellm.proxy.management_endpoints.tool_management_endpoints import ( - router as tool_management_router, -) -from litellm.proxy.management_endpoints.ui_sso import ( - get_disabled_non_admin_personal_key_creation, -) + update_team, validate_membership) +from litellm.proxy.management_endpoints.tool_management_endpoints import \ + router as tool_management_router +from litellm.proxy.management_endpoints.ui_sso import \ + get_disabled_non_admin_personal_key_creation from litellm.proxy.management_endpoints.ui_sso import router as ui_sso_router -from litellm.proxy.management_endpoints.user_agent_analytics_endpoints import ( - router as user_agent_analytics_router, -) -from litellm.proxy.management_helpers.audit_logs import create_audit_log_for_update -from litellm.proxy.middleware.prometheus_auth_middleware import PrometheusAuthMiddleware +from litellm.proxy.management_endpoints.usage_endpoints import \ + router as usage_ai_router +from litellm.proxy.management_endpoints.user_agent_analytics_endpoints import \ + router as user_agent_analytics_router +from litellm.proxy.management_helpers.audit_logs import \ + create_audit_log_for_update +from litellm.proxy.middleware.prometheus_auth_middleware import \ + PrometheusAuthMiddleware from litellm.proxy.ocr_endpoints.endpoints import router as ocr_router -from litellm.proxy.openai_evals_endpoints.endpoints import router as evals_router -from litellm.proxy.openai_files_endpoints.files_endpoints import ( - router as openai_files_router, -) -from litellm.proxy.openai_files_endpoints.files_endpoints import ( - set_files_config, -) -from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( - passthrough_endpoint_router, -) -from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( - router as llm_passthrough_router, -) -from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( - vertex_ai_live_websocket_passthrough, -) -from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( - initialize_pass_through_endpoints, -) -from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( - router as pass_through_router, -) -from litellm.proxy.policy_engine.policy_endpoints import router as policy_crud_router -from litellm.proxy.policy_engine.policy_resolve_endpoints import ( - router as policy_resolve_router, -) +from litellm.proxy.openai_evals_endpoints.endpoints import \ + router as evals_router +from litellm.proxy.openai_files_endpoints.files_endpoints import \ + router as openai_files_router +from litellm.proxy.openai_files_endpoints.files_endpoints import \ + set_files_config +from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import \ + passthrough_endpoint_router +from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import \ + router as llm_passthrough_router +from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import \ + vertex_ai_live_websocket_passthrough +from litellm.proxy.pass_through_endpoints.pass_through_endpoints import \ + initialize_pass_through_endpoints +from litellm.proxy.pass_through_endpoints.pass_through_endpoints import \ + router as pass_through_router +from litellm.proxy.policy_engine.policy_endpoints import \ + router as policy_crud_router +from litellm.proxy.policy_engine.policy_resolve_endpoints import \ + router as policy_resolve_router from litellm.proxy.prompts.prompt_endpoints import router as prompts_router from litellm.proxy.public_endpoints import router as public_endpoints_router from litellm.proxy.rag_endpoints.endpoints import router as rag_router from litellm.proxy.rerank_endpoints.endpoints import router as rerank_router -from litellm.proxy.response_api_endpoints.endpoints import router as response_router +from litellm.proxy.response_api_endpoints.endpoints import \ + router as response_router from litellm.proxy.route_llm_request import route_request from litellm.proxy.search_endpoints.endpoints import router as search_router -from litellm.proxy.search_endpoints.search_tool_management import ( - router as search_tool_management_router, -) -from litellm.proxy.spend_tracking.cloudzero_endpoints import router as cloudzero_router -from litellm.proxy.spend_tracking.spend_management_endpoints import ( - router as spend_management_router, -) -from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload +from litellm.proxy.search_endpoints.search_tool_management import \ + router as search_tool_management_router +from litellm.proxy.spend_tracking.cloudzero_endpoints import \ + router as cloudzero_router +from litellm.proxy.spend_tracking.spend_management_endpoints import \ + router as spend_management_router +from litellm.proxy.spend_tracking.spend_tracking_utils import \ + get_logging_payload from litellm.proxy.types_utils.utils import get_instance_fn -from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import ( - router as ui_crud_endpoints_router, -) -from litellm.proxy.utils import ( - PrismaClient, - ProxyLogging, - ProxyUpdateSpend, - _cache_user_row, - _get_docs_url, - _get_projected_spend_over_limit, - _get_redoc_url, - _is_projected_spend_over_limit, - _is_valid_team_configs, - get_custom_url, - get_error_message_str, - get_server_root_path, - handle_exception_on_proxy, - hash_token, - model_dump_with_preserved_fields, - update_spend, -) -from litellm.proxy.vector_store_endpoints.endpoints import router as vector_store_router -from litellm.proxy.vector_store_endpoints.management_endpoints import ( - router as vector_store_management_router, -) -from litellm.proxy.vector_store_files_endpoints.endpoints import ( - router as vector_store_files_router, -) -from litellm.proxy.vertex_ai_endpoints.langfuse_endpoints import ( - router as langfuse_router, -) +from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import \ + router as ui_crud_endpoints_router +from litellm.proxy.utils import (PrismaClient, ProxyLogging, ProxyUpdateSpend, + _cache_user_row, _get_docs_url, + _get_projected_spend_over_limit, + _get_redoc_url, + _is_projected_spend_over_limit, + _is_valid_team_configs, get_custom_url, + get_error_message_str, get_server_root_path, + handle_exception_on_proxy, hash_token, + model_dump_with_preserved_fields, + update_spend) +from litellm.proxy.vector_store_endpoints.endpoints import \ + router as vector_store_router +from litellm.proxy.vector_store_endpoints.management_endpoints import \ + router as vector_store_management_router +from litellm.proxy.vector_store_files_endpoints.endpoints import \ + router as vector_store_files_router +from litellm.proxy.vertex_ai_endpoints.langfuse_endpoints import \ + router as langfuse_router from litellm.proxy.video_endpoints.endpoints import router as video_router -from litellm.router import ( - AssistantsTypedDict, - Deployment, - LiteLLM_Params, - ModelGroupInfo, -) +from litellm.router import (AssistantsTypedDict, Deployment, LiteLLM_Params, + ModelGroupInfo) from litellm.scheduler import FlowItem, Scheduler from litellm.secret_managers.aws_secret_manager import load_aws_kms from litellm.secret_managers.google_kms import load_google_kms -from litellm.secret_managers.main import ( - get_secret, - get_secret_bool, - get_secret_str, - str_to_bool, -) +from litellm.secret_managers.main import (get_secret, get_secret_bool, + get_secret_str, str_to_bool) from litellm.types.integrations.slack_alerting import SlackAlertingArgs -from litellm.types.llms.anthropic import ( - AnthropicMessagesRequest, - AnthropicResponse, - AnthropicResponseContentBlockText, - AnthropicResponseUsageBlock, -) +from litellm.types.llms.anthropic import (AnthropicMessagesRequest, + AnthropicResponse, + AnthropicResponseContentBlockText, + AnthropicResponseUsageBlock) from litellm.types.llms.openai import HttpxBinaryResponseContent -from litellm.types.proxy.management_endpoints.model_management_endpoints import ( - ModelGroupInfoProxy, -) +from litellm.types.proxy.management_endpoints.model_management_endpoints import \ + ModelGroupInfoProxy from litellm.types.proxy.management_endpoints.ui_sso import ( - DefaultTeamSSOParams, - LiteLLM_UpperboundKeyGenerateParams, -) + DefaultTeamSSOParams, LiteLLM_UpperboundKeyGenerateParams) from litellm.types.realtime import RealtimeQueryParams -from litellm.types.router import ( - DeploymentTypedDict, -) +from litellm.types.router import DeploymentTypedDict from litellm.types.router import ModelInfo as RouterModelInfo -from litellm.types.router import ( - RouterGeneralSettings, - SearchToolTypedDict, - updateDeployment, -) +from litellm.types.router import (RouterGeneralSettings, SearchToolTypedDict, + updateDeployment) from litellm.types.scheduler import DefaultPriorities -from litellm.types.secret_managers.main import ( - KeyManagementSettings, - KeyManagementSystem, -) +from litellm.types.secret_managers.main import (KeyManagementSettings, + KeyManagementSystem) from litellm.types.utils import CredentialItem, CustomHuggingfaceTokenizer from litellm.types.utils import ModelInfo as ModelMapInfo from litellm.types.utils import RawRequestTypedDict, StandardLoggingPayload @@ -556,34 +424,15 @@ litellm.suppress_debug_info = True import json from typing import Union -from fastapi import ( - Depends, - FastAPI, - File, - Form, - Header, - HTTPException, - Path, - Query, - Request, - Response, - UploadFile, - WebSocket, - WebSocketDisconnect, - applications, - status, -) +from fastapi import (Depends, FastAPI, File, Form, Header, HTTPException, Path, + Query, Request, Response, UploadFile, WebSocket, + WebSocketDisconnect, applications, status) from fastapi.encoders import jsonable_encoder from fastapi.middleware.cors import CORSMiddleware from fastapi.openapi.docs import get_swagger_ui_html from fastapi.openapi.utils import get_openapi -from fastapi.responses import ( - FileResponse, - JSONResponse, - ORJSONResponse, - RedirectResponse, - StreamingResponse, -) +from fastapi.responses import (FileResponse, JSONResponse, ORJSONResponse, + RedirectResponse, StreamingResponse) from fastapi.routing import APIRouter from fastapi.security import OAuth2PasswordBearer from fastapi.security.api_key import APIKeyHeader @@ -606,7 +455,8 @@ except Exception: ################### # Import enterprise routes try: - from litellm_enterprise.proxy.enterprise_routes import router as _enterprise_router + from litellm_enterprise.proxy.enterprise_routes import \ + router as _enterprise_router from litellm_enterprise.proxy.proxy_server import EnterpriseProxyConfig enterprise_router = _enterprise_router @@ -948,9 +798,8 @@ def get_openapi_schema(): return app.openapi_schema # Use compatibility wrapper for FastAPI 0.120+ schema generation - from litellm.proxy.common_utils.openapi_schema_compat import ( - get_openapi_schema_with_compat, - ) + from litellm.proxy.common_utils.openapi_schema_compat import \ + get_openapi_schema_with_compat openapi_schema = get_openapi_schema_with_compat( get_openapi_func=get_openapi, @@ -1004,7 +853,8 @@ def get_openapi_schema(): } # Add LLM API request schema bodies for documentation - from litellm.proxy.common_utils.custom_openapi_spec import CustomOpenAPISpec + from litellm.proxy.common_utils.custom_openapi_spec import \ + CustomOpenAPISpec openapi_schema = CustomOpenAPISpec.add_llm_api_request_schema_body(openapi_schema) @@ -1030,7 +880,8 @@ def custom_openapi(): openapi_schema["paths"] = paths_to_include # Add LLM API request schema bodies for documentation - from litellm.proxy.common_utils.custom_openapi_spec import CustomOpenAPISpec + from litellm.proxy.common_utils.custom_openapi_spec import \ + CustomOpenAPISpec openapi_schema = CustomOpenAPISpec.add_llm_api_request_schema_body(openapi_schema) @@ -2023,9 +1874,8 @@ def _schedule_background_health_check_db_save( return import time as time_module - from litellm.proxy.health_endpoints._health_endpoints import ( - _save_background_health_checks_to_db, - ) + from litellm.proxy.health_endpoints._health_endpoints import \ + _save_background_health_checks_to_db checked_by = ( shared_health_manager.pod_id @@ -2085,9 +1935,8 @@ async def _run_background_health_check(): # Initialize shared health check manager if Redis is available and feature is enabled shared_health_manager = None if use_shared_health_check and redis_usage_cache is not None: - from litellm.proxy.health_check_utils.shared_health_check_manager import ( - SharedHealthCheckManager, - ) + from litellm.proxy.health_check_utils.shared_health_check_manager import \ + SharedHealthCheckManager shared_health_manager = SharedHealthCheckManager( redis_cache=redis_usage_cache, @@ -2769,25 +2618,24 @@ class ProxyConfig: litellm.guardrail_name_config_map = guardrail_name_config_map elif key == "global_prompt_directory": - from litellm.integrations.dotprompt import ( - set_global_prompt_directory, - ) + from litellm.integrations.dotprompt import \ + set_global_prompt_directory set_global_prompt_directory(value) verbose_proxy_logger.info( f"{blue_color_code}Set Global Prompt Directory on LiteLLM Proxy{reset_color_code}" ) elif key == "global_bitbucket_config": - from litellm.integrations.bitbucket import ( - set_global_bitbucket_config, - ) + from litellm.integrations.bitbucket import \ + set_global_bitbucket_config set_global_bitbucket_config(value) verbose_proxy_logger.info( f"{blue_color_code}Set Global BitBucket Config on LiteLLM Proxy{reset_color_code}" ) elif key == "global_gitlab_config": - from litellm.integrations.gitlab import set_global_gitlab_config + from litellm.integrations.gitlab import \ + set_global_gitlab_config set_global_gitlab_config(value) verbose_proxy_logger.info( @@ -2858,9 +2706,8 @@ class ProxyConfig: callback ) if "prometheus" in callback: - from litellm.integrations.prometheus import ( - PrometheusLogger, - ) + from litellm.integrations.prometheus import \ + PrometheusLogger if PrometheusLogger is not None: verbose_proxy_logger.debug( @@ -3271,9 +3118,8 @@ class ProxyConfig: mcp_servers_config = config.get("mcp_servers", None) if mcp_servers_config: - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - global_mcp_server_manager, - ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import \ + global_mcp_server_manager # Get mcp_aliases from litellm_settings if available litellm_settings = config.get("litellm_settings", {}) @@ -3286,7 +3132,8 @@ class ProxyConfig: ## VECTOR STORES vector_store_registry_config = config.get("vector_store_registry", None) if vector_store_registry_config: - from litellm.vector_stores.vector_store_registry import VectorStoreRegistry + from litellm.vector_stores.vector_store_registry import \ + VectorStoreRegistry if litellm.vector_store_registry is None: litellm.vector_store_registry = VectorStoreRegistry() @@ -3313,7 +3160,8 @@ class ProxyConfig: """ from litellm.proxy.policy_engine.init_policies import init_policies - from litellm.proxy.policy_engine.policy_validator import PolicyValidator + from litellm.proxy.policy_engine.policy_validator import \ + PolicyValidator if config is None: verbose_proxy_logger.debug("Policy engine: config is None, skipping") @@ -3402,9 +3250,8 @@ class ProxyConfig: key_management_system == KeyManagementSystem.AWS_SECRET_MANAGER.value # noqa: F405 ): - from litellm.secret_managers.aws_secret_manager_v2 import ( - AWSSecretsManagerV2, - ) + from litellm.secret_managers.aws_secret_manager_v2 import \ + AWSSecretsManagerV2 AWSSecretsManagerV2.load_aws_secret_manager( use_aws_secret_manager=True, @@ -3415,28 +3262,24 @@ class ProxyConfig: elif ( key_management_system == KeyManagementSystem.GOOGLE_SECRET_MANAGER.value ): - from litellm.secret_managers.google_secret_manager import ( - GoogleSecretManager, - ) + from litellm.secret_managers.google_secret_manager import \ + GoogleSecretManager GoogleSecretManager() elif key_management_system == KeyManagementSystem.HASHICORP_VAULT.value: - from litellm.secret_managers.hashicorp_secret_manager import ( - HashicorpSecretManager, - ) + from litellm.secret_managers.hashicorp_secret_manager import \ + HashicorpSecretManager HashicorpSecretManager() elif key_management_system == KeyManagementSystem.CYBERARK.value: - from litellm.secret_managers.cyberark_secret_manager import ( - CyberArkSecretManager, - ) + from litellm.secret_managers.cyberark_secret_manager import \ + CyberArkSecretManager CyberArkSecretManager() elif key_management_system == KeyManagementSystem.CUSTOM.value: ### LOAD CUSTOM SECRET MANAGER ### - from litellm.secret_managers.custom_secret_manager_loader import ( - load_custom_secret_manager, - ) + from litellm.secret_managers.custom_secret_manager_loader import \ + load_custom_secret_manager load_custom_secret_manager(config_file_path=config_file_path) else: @@ -3998,9 +3841,8 @@ class ProxyConfig: # Schedule new job if retention period is set (not None) retention_period = general_settings.get("maximum_spend_logs_retention_period") if retention_period is not None: - from litellm.proxy.db.db_transaction_queue.spend_log_cleanup import ( - SpendLogCleanup, - ) + from litellm.proxy.db.db_transaction_queue.spend_log_cleanup import \ + SpendLogCleanup spend_log_cleanup = SpendLogCleanup() cleanup_cron = general_settings.get("maximum_spend_logs_cleanup_cron") @@ -4027,9 +3869,8 @@ class ProxyConfig: ) else: # Interval-based scheduling (existing behavior) - from litellm.litellm_core_utils.duration_parser import ( - duration_in_seconds, - ) + from litellm.litellm_core_utils.duration_parser import \ + duration_in_seconds retention_interval = general_settings.get( "maximum_spend_logs_retention_interval", "1d" @@ -4393,9 +4234,8 @@ class ProxyConfig: if self._should_load_db_object(object_type="sso_settings"): await self._init_sso_settings_in_db(prisma_client=prisma_client) if self._should_load_db_object(object_type="cache_settings"): - from litellm.proxy.management_endpoints.cache_settings_endpoints import ( - CacheSettingsManager, - ) + from litellm.proxy.management_endpoints.cache_settings_endpoints import \ + CacheSettingsManager await CacheSettingsManager.init_cache_settings_in_db( prisma_client=prisma_client, proxy_config=self @@ -4412,7 +4252,8 @@ class ProxyConfig: import json import litellm - from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook + from litellm.proxy.hooks.mcp_semantic_filter import \ + SemanticToolFilterHook try: # Load litellm_settings from DB @@ -4577,9 +4418,8 @@ class ProxyConfig: if should_reload: # Perform the reload - from litellm.litellm_core_utils.get_model_cost_map import ( - get_model_cost_map, - ) + from litellm.litellm_core_utils.get_model_cost_map import \ + get_model_cost_map model_cost_map_url = litellm.model_cost_map_url new_model_cost_map = get_model_cost_map(url=model_cost_map_url) @@ -4685,9 +4525,8 @@ class ProxyConfig: if should_reload: # Perform the reload - from litellm.anthropic_beta_headers_manager import ( - reload_beta_headers_config, - ) + from litellm.anthropic_beta_headers_manager import \ + reload_beta_headers_config new_config = reload_beta_headers_config() @@ -4738,12 +4577,14 @@ class ProxyConfig: Returns: The PromptSpec object """ - from litellm.proxy.prompts.prompt_endpoints import create_versioned_prompt_spec + from litellm.proxy.prompts.prompt_endpoints import \ + create_versioned_prompt_spec return create_versioned_prompt_spec(db_prompt=db_prompt) async def _init_prompts_in_db(self, prisma_client: PrismaClient): - from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY + from litellm.proxy.prompts.prompt_registry import \ + IN_MEMORY_PROMPT_REGISTRY from litellm.types.prompts.init_prompts import PromptSpec try: @@ -4761,10 +4602,7 @@ class ProxyConfig: async def _init_guardrails_in_db(self, prisma_client: PrismaClient): from litellm.proxy.guardrails.guardrail_registry import ( - IN_MEMORY_GUARDRAIL_HANDLER, - Guardrail, - GuardrailRegistry, - ) + IN_MEMORY_GUARDRAIL_HANDLER, Guardrail, GuardrailRegistry) try: guardrails_in_db: List[ @@ -4790,10 +4628,10 @@ class ProxyConfig: """ Initialize policies and policy attachments from database into the in-memory registries. """ - from litellm.proxy.policy_engine.attachment_registry import ( - get_attachment_registry, - ) - from litellm.proxy.policy_engine.policy_registry import get_policy_registry + from litellm.proxy.policy_engine.attachment_registry import \ + get_attachment_registry + from litellm.proxy.policy_engine.policy_registry import \ + get_policy_registry try: # Get the global singleton instances @@ -4819,7 +4657,8 @@ class ProxyConfig: ) async def _init_vector_stores_in_db(self, prisma_client: PrismaClient): - from litellm.vector_stores.vector_store_registry import VectorStoreRegistry + from litellm.vector_stores.vector_store_registry import \ + VectorStoreRegistry try: # read vector stores from db table @@ -4846,7 +4685,8 @@ class ProxyConfig: ) async def _init_vector_store_indexes_in_db(self, prisma_client: PrismaClient): - from litellm.vector_stores.vector_store_registry import VectorStoreIndexRegistry + from litellm.vector_stores.vector_store_registry import \ + VectorStoreIndexRegistry try: # read vector stores from db table @@ -4876,7 +4716,8 @@ class ProxyConfig: ) async def _init_mcp_servers_in_db(self): - from litellm.proxy._experimental.mcp_server.utils import is_mcp_available + from litellm.proxy._experimental.mcp_server.utils import \ + is_mcp_available if not is_mcp_available(): verbose_proxy_logger.debug( @@ -4884,9 +4725,8 @@ class ProxyConfig: ) return - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - global_mcp_server_manager, - ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import \ + global_mcp_server_manager try: await global_mcp_server_manager.reload_servers_from_database() @@ -4898,9 +4738,8 @@ class ProxyConfig: ) async def _init_agents_in_db(self, prisma_client: PrismaClient): - from litellm.proxy.agent_endpoints.agent_registry import ( - global_agent_registry as AGENT_REGISTRY, - ) + from litellm.proxy.agent_endpoints.agent_registry import \ + global_agent_registry as AGENT_REGISTRY try: db_agents = await AGENT_REGISTRY.get_all_agents_from_db( @@ -4923,9 +4762,8 @@ class ProxyConfig: """ global llm_router - from litellm.proxy.search_endpoints.search_tool_registry import ( - SearchToolRegistry, - ) + from litellm.proxy.search_endpoints.search_tool_registry import \ + SearchToolRegistry from litellm.router_utils.search_api_router import SearchAPIRouter try: @@ -4965,9 +4803,8 @@ class ProxyConfig: ) async def _init_pass_through_endpoints_in_db(self): - from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( - initialize_pass_through_endpoints_in_db, - ) + from litellm.proxy.pass_through_endpoints.pass_through_endpoints import \ + initialize_pass_through_endpoints_in_db await initialize_pass_through_endpoints_in_db() @@ -5065,11 +4902,8 @@ async def initialize( # noqa: PLR0915 if debug is True: # this needs to be first, so users can see Router init debugg import logging - from litellm._logging import ( - verbose_logger, - verbose_proxy_logger, - verbose_router_logger, - ) + from litellm._logging import (verbose_logger, verbose_proxy_logger, + verbose_router_logger) # this must ALWAYS remain logging.INFO, DO NOT MODIFY THIS verbose_logger.setLevel(level=logging.INFO) # sets package logs to info @@ -5078,11 +4912,8 @@ async def initialize( # noqa: PLR0915 if detailed_debug is True: import logging - from litellm._logging import ( - verbose_logger, - verbose_proxy_logger, - verbose_router_logger, - ) + from litellm._logging import (verbose_logger, verbose_proxy_logger, + verbose_router_logger) verbose_logger.setLevel(level=logging.DEBUG) # set package log to debug verbose_router_logger.setLevel(level=logging.DEBUG) # set router logs to debug @@ -5094,7 +4925,8 @@ async def initialize( # noqa: PLR0915 if litellm_log_setting.upper() == "INFO": import logging - from litellm._logging import verbose_proxy_logger, verbose_router_logger + from litellm._logging import (verbose_proxy_logger, + verbose_router_logger) # this must ALWAYS remain logging.INFO, DO NOT MODIFY THIS @@ -5107,11 +4939,9 @@ async def initialize( # noqa: PLR0915 elif litellm_log_setting.upper() == "DEBUG": import logging - from litellm._logging import ( - verbose_logger, - verbose_proxy_logger, - verbose_router_logger, - ) + from litellm._logging import (verbose_logger, + verbose_proxy_logger, + verbose_router_logger) verbose_logger.setLevel(level=logging.DEBUG) # set package log to debug verbose_router_logger.setLevel( @@ -5463,7 +5293,8 @@ class ProxyStartupEvent: litellm_settings: Dict[str, Any], ): """Initialize MCP semantic tool filter if configured""" - from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook + from litellm.proxy.hooks.mcp_semantic_filter import \ + SemanticToolFilterHook mcp_semantic_filter_config = litellm_settings.get( "mcp_semantic_tool_filter", None @@ -5575,9 +5406,8 @@ class ProxyStartupEvent: _teams = litellm.default_internal_user_params.get("teams") or [] if _teams and all(isinstance(team, dict) for team in _teams): - from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import ( - update_default_team_member_budget, - ) + from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import \ + update_default_team_member_budget teams_pydantic_obj = [NewUserRequestTeam(**team) for team in _teams] await update_default_team_member_budget( @@ -5798,9 +5628,8 @@ class ProxyStartupEvent: ### CHECK BATCH COST ### if llm_router is not None: try: - from litellm_enterprise.proxy.common_utils.check_batch_cost import ( - CheckBatchCost, - ) + from litellm_enterprise.proxy.common_utils.check_batch_cost import \ + CheckBatchCost check_batch_cost_job = CheckBatchCost( proxy_logging_obj=proxy_logging_obj, @@ -5829,9 +5658,8 @@ class ProxyStartupEvent: ### CHECK RESPONSES COST ### if llm_router is not None: try: - from litellm_enterprise.proxy.common_utils.check_responses_cost import ( - CheckResponsesCost, - ) + from litellm_enterprise.proxy.common_utils.check_responses_cost import \ + CheckResponsesCost check_responses_cost_job = CheckResponsesCost( proxy_logging_obj=proxy_logging_obj, @@ -5891,7 +5719,8 @@ class ProxyStartupEvent: ######################################################## from litellm.integrations.cloudzero.cloudzero import CloudZeroLogger from litellm.integrations.focus.focus_logger import FocusLogger - from litellm.proxy.spend_tracking.cloudzero_endpoints import is_cloudzero_setup + from litellm.proxy.spend_tracking.cloudzero_endpoints import \ + is_cloudzero_setup if await is_cloudzero_setup(): await CloudZeroLogger.init_cloudzero_background_job(scheduler=scheduler) @@ -5913,17 +5742,15 @@ class ProxyStartupEvent: ######################################################## from litellm.constants import ( LITELLM_KEY_ROTATION_CHECK_INTERVAL_SECONDS, - LITELLM_KEY_ROTATION_ENABLED, - ) + LITELLM_KEY_ROTATION_ENABLED) key_rotation_enabled: Optional[bool] = str_to_bool(LITELLM_KEY_ROTATION_ENABLED) verbose_proxy_logger.debug(f"key_rotation_enabled: {key_rotation_enabled}") if key_rotation_enabled is True: try: - from litellm.proxy.common_utils.key_rotation_manager import ( - KeyRotationManager, - ) + from litellm.proxy.common_utils.key_rotation_manager import \ + KeyRotationManager # Get prisma_client from global scope global prisma_client @@ -6084,9 +5911,7 @@ class ProxyStartupEvent: Doc: https://docs.datadoghq.com/tracing/trace_collection/automatic_instrumentation/dd_libraries/python/ """ from litellm.litellm_core_utils.dd_tracing import ( - _should_use_dd_profiler, - _should_use_dd_tracer, - ) + _should_use_dd_profiler, _should_use_dd_tracer) if _should_use_dd_tracer(): import ddtrace @@ -6198,13 +6023,10 @@ async def model_list( """ global llm_model_list, general_settings, llm_router, prisma_client, user_api_key_cache, proxy_logging_obj - from litellm.proxy.management_endpoints.common_utils import ( - _user_has_admin_privileges, - ) - from litellm.proxy.utils import ( - create_model_info_response, - get_available_models_for_user, - ) + from litellm.proxy.management_endpoints.common_utils import \ + _user_has_admin_privileges + from litellm.proxy.utils import (create_model_info_response, + get_available_models_for_user) # Validate scope parameter if provided if scope is not None and scope != "expand": @@ -6331,11 +6153,9 @@ async def model_info( """ global llm_model_list, general_settings, llm_router, prisma_client, user_api_key_cache, proxy_logging_obj - from litellm.proxy.utils import ( - create_model_info_response, - get_available_models_for_user, - validate_model_access, - ) + from litellm.proxy.utils import (create_model_info_response, + get_available_models_for_user, + validate_model_access) # Get available models for the user all_models = await get_available_models_for_user( @@ -8252,7 +8072,8 @@ def _get_provider_token_counter( if deployment is None: return None - from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider + from litellm.litellm_core_utils.get_llm_provider_logic import \ + get_llm_provider full_model = deployment.get("litellm_params", {}).get("model", "") model: Optional[str] = None @@ -10505,6 +10326,12 @@ async def async_queue_request( data["metadata"]["user_api_key_team_id"] = getattr( user_api_key_dict, "team_id", None ) + data["metadata"]["user_api_key_object_permission_id"] = getattr( + user_api_key_dict, "object_permission_id", None + ) + data["metadata"]["user_api_key_team_object_permission_id"] = getattr( + user_api_key_dict, "team_object_permission_id", None + ) data["metadata"]["endpoint"] = str(request.url) global user_temperature, user_request_timeout, user_max_tokens, user_api_base @@ -10598,7 +10425,8 @@ async def fallback_login(request: Request): ) # hidden since this is a helper for UI sso login async def login(request: Request): # noqa: PLR0915 global premium_user, general_settings, master_key - from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object + from litellm.proxy.auth.login_utils import (authenticate_user, + create_ui_token_object) from litellm.proxy.utils import get_custom_url form = await request.form() @@ -10648,7 +10476,8 @@ async def login(request: Request): # noqa: PLR0915 ) # hidden helper for UI logins via API async def login_v2(request: Request): # noqa: PLR0915 global premium_user, general_settings, master_key - from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object + from litellm.proxy.auth.login_utils import (authenticate_user, + create_ui_token_object) from litellm.proxy.utils import get_custom_url try: @@ -10979,7 +10808,8 @@ async def get_image(): if logo_path.startswith(("http://", "https://")): try: # Download the image and cache it - from litellm.llms.custom_httpx.http_handler import get_async_httpx_client + from litellm.llms.custom_httpx.http_handler import \ + get_async_httpx_client from litellm.types.llms.custom_http import httpxSpecialProvider async_client = get_async_httpx_client( @@ -11023,9 +10853,8 @@ async def get_favicon(): if favicon_url.startswith(("http://", "https://")): try: - from litellm.llms.custom_httpx.http_handler import ( - get_async_httpx_client, - ) + from litellm.llms.custom_httpx.http_handler import \ + get_async_httpx_client from litellm.types.llms.custom_http import httpxSpecialProvider async_client = get_async_httpx_client( @@ -11090,9 +10919,8 @@ async def new_invitation( ``` """ try: - from litellm.proxy.management_helpers.user_invitation import ( - create_invitation_for_user, - ) + from litellm.proxy.management_helpers.user_invitation import \ + create_invitation_for_user global prisma_client @@ -12179,7 +12007,8 @@ async def reload_model_cost_map( ) # Immediately reload the model cost map in the current pod - from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map + from litellm.litellm_core_utils.get_model_cost_map import \ + get_model_cost_map model_cost_map_url = litellm.model_cost_map_url new_model_cost_map = get_model_cost_map(url=model_cost_map_url) @@ -12471,9 +12300,8 @@ async def get_model_cost_map_source( ) try: - from litellm.litellm_core_utils.get_model_cost_map import ( - get_model_cost_map_source_info, - ) + from litellm.litellm_core_utils.get_model_cost_map import \ + get_model_cost_map_source_info source_info = get_model_cost_map_source_info() model_count = len(litellm.model_cost) if litellm.model_cost else 0 @@ -12525,7 +12353,8 @@ async def reload_anthropic_beta_headers( ) # Immediately reload the beta headers config in the current pod - from litellm.anthropic_beta_headers_manager import reload_beta_headers_config + from litellm.anthropic_beta_headers_manager import \ + reload_beta_headers_config new_config = reload_beta_headers_config() @@ -12925,9 +12754,8 @@ async def dynamic_mcp_route(mcp_server_name: str, request: Request): """Handle dynamic MCP server routes like /github_mcp/mcp""" try: # Validate that the MCP server exists - from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( - global_mcp_server_manager, - ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import \ + global_mcp_server_manager from litellm.proxy.auth.ip_address_utils import IPAddressUtils from litellm.types.mcp import MCPAuth @@ -12946,9 +12774,8 @@ async def dynamic_mcp_route(mcp_server_name: str, request: Request): scope["path"] = f"/mcp/{mcp_server_name}" # Import the MCP handler - from litellm.proxy._experimental.mcp_server.server import ( - handle_streamable_http_mcp, - ) + from litellm.proxy._experimental.mcp_server.server import \ + handle_streamable_http_mcp # Create a custom send function to capture the response response_started = False diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 08f756e4fb0..cd4f9a4d247 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -260,6 +260,7 @@ model LiteLLM_ObjectPermissionTable { vector_stores String[] @default([]) agents String[] @default([]) agent_access_groups String[] @default([]) + blocked_tools String[] @default([]) // Tool names blocked for any key/team/user with this permission teams LiteLLM_TeamTable[] projects LiteLLM_ProjectTable[] verification_tokens LiteLLM_VerificationToken[] @@ -920,6 +921,16 @@ model LiteLLM_SpendLogGuardrailIndex { @@index([policy_id, start_time]) } +// Index for fast "last N logs for tool" from SpendLogs – see how a tool is called in production +model LiteLLM_SpendLogToolIndex { + request_id String + tool_name String // matches LiteLLM_ToolTable.tool_name; join for call_policy etc. + start_time DateTime + + @@id([request_id, tool_name]) + @@index([tool_name, start_time]) +} + // Prompt table for storing prompt configurations model LiteLLM_PromptTable { id String @id @default(uuid()) @@ -1077,25 +1088,6 @@ model LiteLLM_ToolTable { @@index([team_id]) } -// Per-(tool, team/key) policy overrides. When present, override replaces global tool policy for that scope. -model LiteLLM_ToolPolicyOverrideTable { - override_id String @id @default(uuid()) - tool_name String - team_id String @default("") // "" = not scoped to team; non-empty = override for this team only - key_hash String @default("") // "" = not scoped to key; non-empty = override for this key only - call_policy String @default("blocked") - key_alias String? // human-readable key alias for UI - created_at DateTime @default(now()) - created_by String? - updated_at DateTime @default(now()) @updatedAt - updated_by String? - - @@unique([tool_name, team_id, key_hash]) - @@index([tool_name]) - @@index([team_id]) - @@index([key_hash]) -} - //Unified Access Groups table for storing unified access groups model LiteLLM_AccessGroupTable { access_group_id String @id @default(uuid()) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index f6613b5548f..e31ded98497 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -10,17 +10,8 @@ import traceback from datetime import date, datetime, timedelta, timezone from email.mime.multipart import MIMEMultipart from email.mime.text import MIMEText -from typing import ( - TYPE_CHECKING, - Any, - Dict, - List, - Literal, - Optional, - Union, - cast, - overload, -) +from typing import (TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, + cast, overload) from litellm import _custom_logger_compatible_callbacks_literal from litellm.constants import (DEFAULT_MODEL_CREATED_AT_TIME, @@ -3550,8 +3541,9 @@ class PrismaClient: def _get_engine_pid(self) -> int: try: engine = self.db._original_prisma._engine # type: ignore[attr-defined] - if engine is not None and engine.process is not None: - return engine.process.pid + process = getattr(engine, "process", None) if engine is not None else None + if process is not None: + return process.pid except (AttributeError, TypeError): pass return 0 @@ -4638,6 +4630,20 @@ async def update_spend_logs_job( guardrail_tracking_err, ) + # Tool usage tracking (same batch): SpendLogToolIndex for "last N requests for tool X" + try: + from litellm.proxy.db.spend_log_tool_index import \ + process_spend_logs_tool_usage + await process_spend_logs_tool_usage( + prisma_client=prisma_client, + logs_to_process=logs_to_process, + ) + except Exception as tool_tracking_err: + verbose_proxy_logger.warning( + "Spend tracking - tool usage tracking failed (non-fatal): %s", + tool_tracking_err, + ) + async def _monitor_spend_logs_queue( prisma_client: PrismaClient, diff --git a/litellm/types/tool_management.py b/litellm/types/tool_management.py index 90fc4a6ec7d..ef8a488f4a0 100644 --- a/litellm/types/tool_management.py +++ b/litellm/types/tool_management.py @@ -61,3 +61,21 @@ class ToolPolicyOverrideRow(BaseModel): class ToolDetailResponse(BaseModel): tool: LiteLLM_ToolTableRow overrides: List[ToolPolicyOverrideRow] = Field(default_factory=list) + + +class ToolUsageLogEntry(BaseModel): + """One spend log row for a tool call (for UI "recent logs" table).""" + + id: str # request_id + timestamp: str + model: Optional[str] = None + spend: Optional[float] = None + total_tokens: Optional[int] = None + input_snippet: Optional[str] = None + + +class ToolUsageLogsResponse(BaseModel): + logs: List[ToolUsageLogEntry] + total: int + page: int + page_size: int diff --git a/schema.prisma b/schema.prisma index 61de0073e0b..691883ef446 100644 --- a/schema.prisma +++ b/schema.prisma @@ -920,6 +920,16 @@ model LiteLLM_SpendLogGuardrailIndex { @@index([policy_id, start_time]) } +// Index for fast "last N logs for tool" from SpendLogs – see how a tool is called in production +model LiteLLM_SpendLogToolIndex { + request_id String + tool_name String // matches LiteLLM_ToolTable.tool_name; join for call_policy etc. + start_time DateTime + + @@id([request_id, tool_name]) + @@index([tool_name, start_time]) +} + // Prompt table for storing prompt configurations model LiteLLM_PromptTable { id String @id @default(uuid()) diff --git a/ui/litellm-dashboard/src/components/ToolDetail.tsx b/ui/litellm-dashboard/src/components/ToolDetail.tsx index fa83190d7e3..33bfa9b4357 100644 --- a/ui/litellm-dashboard/src/components/ToolDetail.tsx +++ b/ui/litellm-dashboard/src/components/ToolDetail.tsx @@ -1,20 +1,25 @@ "use client"; -import { ArrowLeftOutlined, ToolOutlined } from "@ant-design/icons"; +import { ArrowLeftOutlined, HistoryOutlined, ToolOutlined } from "@ant-design/icons"; import { useQuery, useQueryClient } from "@tanstack/react-query"; -import { Button, Select, Spin } from "antd"; +import { Button, Pagination, Select, Spin, Table } from "antd"; import React, { useCallback, useMemo, useState } from "react"; import TeamDropdown from "@/components/common_components/team_dropdown"; import { PolicySelect } from "@/components/ToolPolicies/PolicySelect"; import { deleteToolPolicyOverride, fetchToolDetail, + getToolUsageLogs, keyListCall, teamListCall, + uiSpendLogsCall, updateToolPolicy, type ToolPolicyOverrideRow, + type ToolUsageLogEntry, } from "@/components/networking"; import type { Team } from "@/components/key_team_helpers/key_list"; +import type { LogEntry } from "@/components/view_logs/columns"; +import { LogDetailsDrawer } from "@/components/view_logs/LogDetailsDrawer"; interface ToolDetailProps { toolName: string; @@ -34,6 +39,17 @@ interface KeyOption { const TOOL_DETAIL_QUERY_KEY = "tool-detail"; +const LOGS_PAGE_SIZE = 20; + +function getDefaultLogsDateRange(): { start: string; end: string } { + const end = new Date(); + const start = new Date(); + start.setDate(start.getDate() - 90); + const fmt = (d: Date) => + d.toISOString().slice(0, 19).replace("T", " "); + return { start: fmt(start), end: fmt(end) }; +} + export function ToolDetail({ toolName, onBack, accessToken }: ToolDetailProps) { const queryClient = useQueryClient(); const [overrideSaving, setOverrideSaving] = useState(false); @@ -41,6 +57,11 @@ export function ToolDetail({ toolName, onBack, accessToken }: ToolDetailProps) { const [blockScope, setBlockScope] = useState<"team" | "key">("team"); const [blockTeamId, setBlockTeamId] = useState(null); const [blockKey, setBlockKey] = useState(null); + const [logsPage, setLogsPage] = useState(1); + const [selectedRequestId, setSelectedRequestId] = useState(null); + const [drawerOpen, setDrawerOpen] = useState(false); + + const logsDateRange = useMemo(() => getDefaultLogsDateRange(), []); const { data: detail, isLoading: detailLoading, error: detailError } = useQuery({ queryKey: [TOOL_DETAIL_QUERY_KEY, toolName], @@ -60,6 +81,35 @@ export function ToolDetail({ toolName, onBack, accessToken }: ToolDetailProps) { enabled: !!accessToken, }); + const { data: logsData, isLoading: logsLoading } = useQuery({ + queryKey: ["tool-usage-logs", toolName, logsPage], + queryFn: () => + getToolUsageLogs(accessToken!, toolName, { + page: logsPage, + pageSize: LOGS_PAGE_SIZE, + }), + enabled: !!accessToken && !!toolName, + }); + + const { data: fullLogResponse } = useQuery({ + queryKey: ["spend-log-by-request-tool-detail", selectedRequestId, logsDateRange.start, logsDateRange.end], + queryFn: async () => { + if (!accessToken || !selectedRequestId) return null; + const res = await uiSpendLogsCall({ + accessToken, + start_date: logsDateRange.start, + end_date: logsDateRange.end, + page: 1, + page_size: 10, + params: { request_id: selectedRequestId }, + }); + return res as { data: LogEntry[]; total: number }; + }, + enabled: !!accessToken && !!selectedRequestId && drawerOpen, + }); + + const selectedLog: LogEntry | null = fullLogResponse?.data?.[0] ?? null; + const teams: Team[] = useMemo(() => { const arr = Array.isArray(teamsData) ? teamsData : teamsData?.data ?? []; return arr.map((t: { team_id?: string; id?: string; team_alias?: string }) => ({ @@ -311,7 +361,104 @@ export function ToolDetail({ toolName, onBack, accessToken }: ToolDetailProps) {
+ +
+

+ + Recent logs +

+

+ Requests that used this tool. Click a row to open full log details. +

+ {logsLoading && ( +
+ +
+ )} + {!logsLoading && (!logsData?.logs?.length) && ( +
+ No logs for this tool yet. Usage will appear here after requests that call this tool. +
+ )} + {!logsLoading && logsData && logsData.logs.length > 0 && ( + <> + + dataSource={logsData.logs} + rowKey="id" + size="small" + pagination={false} + onRow={(record) => ({ + onClick: () => { + setSelectedRequestId(record.id); + setDrawerOpen(true); + }, + style: { cursor: "pointer" }, + })} + columns={[ + { + title: "Time", + dataIndex: "timestamp", + key: "timestamp", + width: 200, + render: (ts: string) => + ts ? new Date(ts).toLocaleString(undefined, { dateStyle: "short", timeStyle: "short" }) : "—", + }, + { + title: "Model", + dataIndex: "model", + key: "model", + ellipsis: true, + render: (v: string | null) => v ?? "—", + }, + { + title: "Spend", + dataIndex: "spend", + key: "spend", + width: 90, + render: (v: number | null) => + v != null ? `$${Number(v).toFixed(4)}` : "—", + }, + { + title: "Tokens", + dataIndex: "total_tokens", + key: "total_tokens", + width: 90, + render: (v: number | null) => (v != null ? v.toLocaleString() : "—"), + }, + { + title: "Input", + dataIndex: "input_snippet", + key: "input_snippet", + ellipsis: true, + render: (v: string | null) => (v ? String(v).slice(0, 80) + (String(v).length > 80 ? "…" : "") : "—"), + }, + ]} + /> +
+ +
+ + )} +
+ + { + setDrawerOpen(false); + setSelectedRequestId(null); + }} + logEntry={selectedLog} + accessToken={accessToken} + allLogs={selectedLog ? [selectedLog] : []} + startTime={logsDateRange.start} + />
); } diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 65a01420a70..61a4daeec65 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -10011,6 +10011,51 @@ export interface ToolDetailResponse { overrides: ToolPolicyOverrideRow[]; } +export interface ToolUsageLogEntry { + id: string; + timestamp: string; + model?: string | null; + spend?: number | null; + total_tokens?: number | null; + input_snippet?: string | null; +} + +export interface ToolUsageLogsResponse { + logs: ToolUsageLogEntry[]; + total: number; + page: number; + page_size: number; +} + +export const getToolUsageLogs = async ( + accessToken: string, + toolName: string, + options: { page?: number; pageSize?: number; startDate?: string; endDate?: string } +): Promise => { + const encoded = encodeURIComponent(toolName); + const url = proxyBaseUrl + ? `${proxyBaseUrl}/v1/tool/${encoded}/logs` + : `/v1/tool/${encoded}/logs`; + const params = new URLSearchParams(); + if (options.page != null) params.append("page", String(options.page)); + if (options.pageSize != null) params.append("page_size", String(options.pageSize)); + if (options.startDate) params.append("start_date", options.startDate); + if (options.endDate) params.append("end_date", options.endDate); + const fullUrl = params.toString() ? `${url}?${params.toString()}` : url; + const response = await fetch(fullUrl, { + method: "GET", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + }); + if (!response.ok) { + const errorData = await response.json().catch(() => ({})); + throw new Error(deriveErrorMessage(errorData)); + } + return response.json(); +}; + export const fetchToolDetail = async ( accessToken: string, toolName: string diff --git a/ui/litellm-dashboard/tsconfig.json b/ui/litellm-dashboard/tsconfig.json index d24bdd340f7..5b0352feb98 100644 --- a/ui/litellm-dashboard/tsconfig.json +++ b/ui/litellm-dashboard/tsconfig.json @@ -14,7 +14,7 @@ "moduleResolution": "bundler", "resolveJsonModule": true, "isolatedModules": true, - "jsx": "react-jsx", + "jsx": "preserve", "incremental": true, "plugins": [ { From b6ba09e1b3ca01ae3fc4bd2b7bd760fad41f2268 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Thu, 26 Feb 2026 00:48:05 -0800 Subject: [PATCH 09/12] refactor: backend code improvements --- .../migration.sql | 2 ++ .../migration.sql | 2 ++ .../migration.sql | 2 ++ .../migration.sql | 2 ++ litellm/proxy/db/tool_registry_writer.py | 3 --- .../tool_management_endpoints.py | 26 +++++++++++-------- 6 files changed, 23 insertions(+), 14 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260226002749_baseline_diff/migration.sql create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260226003106_baseline_diff/migration.sql create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260226003346_baseline_diff/migration.sql create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260226003614_baseline_diff/migration.sql diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226002749_baseline_diff/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226002749_baseline_diff/migration.sql new file mode 100644 index 00000000000..2f725d83806 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226002749_baseline_diff/migration.sql @@ -0,0 +1,2 @@ +-- This is an empty migration. + diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226003106_baseline_diff/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226003106_baseline_diff/migration.sql new file mode 100644 index 00000000000..2f725d83806 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226003106_baseline_diff/migration.sql @@ -0,0 +1,2 @@ +-- This is an empty migration. + diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226003346_baseline_diff/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226003346_baseline_diff/migration.sql new file mode 100644 index 00000000000..2f725d83806 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226003346_baseline_diff/migration.sql @@ -0,0 +1,2 @@ +-- This is an empty migration. + diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226003614_baseline_diff/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226003614_baseline_diff/migration.sql new file mode 100644 index 00000000000..2f725d83806 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226003614_baseline_diff/migration.sql @@ -0,0 +1,2 @@ +-- This is an empty migration. + diff --git a/litellm/proxy/db/tool_registry_writer.py b/litellm/proxy/db/tool_registry_writer.py index 172e05cec9d..1cb1c8b9f18 100644 --- a/litellm/proxy/db/tool_registry_writer.py +++ b/litellm/proxy/db/tool_registry_writer.py @@ -280,7 +280,6 @@ async def _get_merged_blocked_tools( try: row = await prisma_client.db.litellm_objectpermissiontable.find_unique( where={"object_permission_id": op_id.strip()}, - select={"blocked_tools": True}, ) if row is not None and getattr(row, "blocked_tools", None): blocked.update(row.blocked_tools) @@ -403,7 +402,6 @@ async def add_tool_to_object_permission_blocked( try: row = await prisma_client.db.litellm_objectpermissiontable.find_unique( where={"object_permission_id": object_permission_id}, - select={"blocked_tools": True}, ) if row is None: return False @@ -434,7 +432,6 @@ async def remove_tool_from_object_permission_blocked( try: row = await prisma_client.db.litellm_objectpermissiontable.find_unique( where={"object_permission_id": object_permission_id}, - select={"blocked_tools": True}, ) if row is None: return False diff --git a/litellm/proxy/management_endpoints/tool_management_endpoints.py b/litellm/proxy/management_endpoints/tool_management_endpoints.py index 84e8a12138a..90115b13ab3 100644 --- a/litellm/proxy/management_endpoints/tool_management_endpoints.py +++ b/litellm/proxy/management_endpoints/tool_management_endpoints.py @@ -56,7 +56,9 @@ async def list_tools( ) try: - tools = await db_list_tools(prisma_client=prisma_client, call_policy=call_policy) + tools = await db_list_tools( + prisma_client=prisma_client, call_policy=call_policy + ) return ToolListResponse(tools=tools, total=len(tools)) except Exception as e: verbose_proxy_logger.exception("Error listing tools: %s", e) @@ -91,9 +93,7 @@ async def get_tool_detail( try: tool = await db_get_tool(prisma_client=prisma_client, tool_name=tool_name) if tool is None: - raise HTTPException( - status_code=404, detail=f"Tool '{tool_name}' not found" - ) + raise HTTPException(status_code=404, detail=f"Tool '{tool_name}' not found") overrides = await list_overrides_for_tool( prisma_client=prisma_client, tool_name=tool_name ) @@ -119,6 +119,7 @@ def _input_snippet_for_tool_log(sl: Any, max_len: int = 200) -> Optional[str]: return None if isinstance(psr, str): import json + try: psr = json.loads(psr) except Exception: @@ -280,9 +281,7 @@ async def get_tool( try: tool = await db_get_tool(prisma_client=prisma_client, tool_name=tool_name) if tool is None: - raise HTTPException( - status_code=404, detail=f"Tool '{tool_name}' not found" - ) + raise HTTPException(status_code=404, detail=f"Tool '{tool_name}' not found") return tool except HTTPException: raise @@ -302,8 +301,7 @@ async def _resolve_key_hash_to_object_permission_id( if not hashed: return None row = await prisma_client.db.litellm_verificationtoken.find_unique( - where={"token": hashed}, - select={"object_permission_id": True}, + where={"token": hashed} ) if row is None: return None @@ -312,6 +310,7 @@ async def _resolve_key_hash_to_object_permission_id( return op_id # Create new object permission and assign to key import uuid as _uuid + new_id = str(_uuid.uuid4()) await prisma_client.db.litellm_objectpermissiontable.create( data={"object_permission_id": new_id, "blocked_tools": []} @@ -340,6 +339,7 @@ async def _resolve_team_id_to_object_permission_id( if op_id: return op_id import uuid as _uuid + new_id = str(_uuid.uuid4()) await prisma_client.db.litellm_objectpermissiontable.create( data={"object_permission_id": new_id, "blocked_tools": []} @@ -461,8 +461,12 @@ async def update_tool_policy( ) async def delete_tool_policy_override( tool_name: str, - team_id: Optional[str] = Query(None, description="Team ID of the override to remove"), - key_hash: Optional[str] = Query(None, description="Key hash of the override to remove"), + team_id: Optional[str] = Query( + None, description="Team ID of the override to remove" + ), + key_hash: Optional[str] = Query( + None, description="Key hash of the override to remove" + ), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ From 4906f2a7612dbc26c1e3a91218c930508e0d213b Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Thu, 26 Feb 2026 00:58:18 -0800 Subject: [PATCH 10/12] refactor: improve log viewer for tools --- .../migration.sql | 2 - .../migration.sql | 2 - .../migration.sql | 2 - .../migration.sql | 2 - litellm/proxy/_types.py | 3 +- litellm/proxy/db/tool_registry_writer.py | 79 ++++++++++ .../tool_policy/tool_policy_guardrail.py | 52 ++----- litellm/proxy/proxy_server.py | 22 +++ .../proxy/db/test_tool_registry_writer.py | 66 +++++++- .../test_tool_policy_guardrail.py | 85 +++++----- .../src/components/ToolDetail.tsx | 147 ++++-------------- 11 files changed, 259 insertions(+), 203 deletions(-) delete mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260226002749_baseline_diff/migration.sql delete mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260226003106_baseline_diff/migration.sql delete mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260226003346_baseline_diff/migration.sql delete mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260226003614_baseline_diff/migration.sql diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226002749_baseline_diff/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226002749_baseline_diff/migration.sql deleted file mode 100644 index 2f725d83806..00000000000 --- a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226002749_baseline_diff/migration.sql +++ /dev/null @@ -1,2 +0,0 @@ --- This is an empty migration. - diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226003106_baseline_diff/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226003106_baseline_diff/migration.sql deleted file mode 100644 index 2f725d83806..00000000000 --- a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226003106_baseline_diff/migration.sql +++ /dev/null @@ -1,2 +0,0 @@ --- This is an empty migration. - diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226003346_baseline_diff/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226003346_baseline_diff/migration.sql deleted file mode 100644 index 2f725d83806..00000000000 --- a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226003346_baseline_diff/migration.sql +++ /dev/null @@ -1,2 +0,0 @@ --- This is an empty migration. - diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226003614_baseline_diff/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226003614_baseline_diff/migration.sql deleted file mode 100644 index 2f725d83806..00000000000 --- a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260226003614_baseline_diff/migration.sql +++ /dev/null @@ -1,2 +0,0 @@ --- This is an empty migration. - diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 2735b3780f5..b608956e3ff 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -58,6 +58,7 @@ class SupportedDBObjectType(str, enum.Enum): PASS_THROUGH_ENDPOINTS = "pass_through_endpoints" PROMPTS = "prompts" MODEL_COST_MAP = "model_cost_map" + TOOLS = "tools" def __str__(self): return str(self.value) @@ -2101,7 +2102,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): user_header_mappings: Optional[List[UserHeaderMapping]] = None supported_db_objects: Optional[List[SupportedDBObjectType]] = Field( None, - description="Fine-grained control over which object types to load from the database when store_model_in_db is True. Available types: 'models', 'mcp', 'guardrails', 'vector_stores', 'pass_through_endpoints', 'prompts', 'model_cost_map'. If not set, all objects are loaded (default behavior).", + description="Fine-grained control over which object types to load from the database when store_model_in_db is True. Available types: 'models', 'mcp', 'guardrails', 'vector_stores', 'pass_through_endpoints', 'prompts', 'model_cost_map', 'tools'. If not set, all objects are loaded (default behavior).", ) user_mcp_management_mode: Optional[UserMCPManagementMode] = Field( None, diff --git a/litellm/proxy/db/tool_registry_writer.py b/litellm/proxy/db/tool_registry_writer.py index 1cb1c8b9f18..1c106df675b 100644 --- a/litellm/proxy/db/tool_registry_writer.py +++ b/litellm/proxy/db/tool_registry_writer.py @@ -328,6 +328,85 @@ async def get_effective_policies( return {} +class ToolPolicyRegistry: + """ + In-memory registry of tool policies synced from DB. + Synced in _init_tool_policy_in_db (from add_deployment / _init_non_llm_objects_in_db). + Hot path uses get_effective_policies only — no DB, no cache. + """ + + def __init__(self) -> None: + self._global_tool_policies: Dict[str, str] = {} + self._blocked_tools_by_op_id: Dict[str, List[str]] = {} + self._initialized: bool = False + + def is_initialized(self) -> bool: + return self._initialized + + async def sync_tool_policy_from_db(self, prisma_client: "PrismaClient") -> None: + """Load all tool policies and object-permission blocked_tools from DB; replace in-memory state.""" + try: + tools = await prisma_client.db.litellm_tooltable.find_many() + self._global_tool_policies = {row.tool_name: row.call_policy for row in tools} + + perms = await prisma_client.db.litellm_objectpermissiontable.find_many() + self._blocked_tools_by_op_id = {} + for row in perms: + op_id = getattr(row, "object_permission_id", None) + blocked = getattr(row, "blocked_tools", None) or [] + if op_id: + self._blocked_tools_by_op_id[op_id] = list(blocked) + + self._initialized = True + verbose_proxy_logger.info( + "ToolPolicyRegistry: synced %d global tool policies and %d object permissions from DB", + len(self._global_tool_policies), + len(self._blocked_tools_by_op_id), + ) + except Exception as e: + verbose_proxy_logger.exception( + "ToolPolicyRegistry sync_tool_policy_from_db error: %s", e + ) + raise + + def get_effective_policies( + self, + tool_names: List[str], + object_permission_id: Optional[str] = None, + team_object_permission_id: Optional[str] = None, + ) -> Dict[str, str]: + """ + Return effective call_policy per tool from in-memory state. + If tool is in key or team blocked_tools -> "blocked", else global policy or "untrusted". + """ + if not tool_names: + return {} + blocked: set = set() + for op_id in (object_permission_id, team_object_permission_id): + if op_id and op_id.strip(): + blocked.update( + self._blocked_tools_by_op_id.get(op_id.strip(), []) + ) + result: Dict[str, str] = {} + for name in tool_names: + if name in blocked: + result[name] = "blocked" + else: + result[name] = self._global_tool_policies.get(name, "untrusted") + return result + + +_tool_policy_registry: Optional[ToolPolicyRegistry] = None + + +def get_tool_policy_registry() -> ToolPolicyRegistry: + """Return the global ToolPolicyRegistry singleton.""" + global _tool_policy_registry + if _tool_policy_registry is None: + _tool_policy_registry = ToolPolicyRegistry() + return _tool_policy_registry + + def _effective_cache_suffix( object_permission_id: Optional[str], team_object_permission_id: Optional[str], diff --git a/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py index c8763f162c1..b6cb4bba138 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py @@ -23,12 +23,11 @@ or both pre and post call: mode: during_call # runs before LLM and on response """ -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple +from typing import TYPE_CHECKING, Any, List, Literal, Optional, Tuple from fastapi import HTTPException from litellm._logging import verbose_proxy_logger -from litellm.caching.dual_cache import DualCache from litellm.integrations.custom_guardrail import (CustomGuardrail, log_guardrail_information) from litellm.proxy.guardrails.tool_name_extraction import \ @@ -101,9 +100,9 @@ def _get_effective_allowed_tools_from_request( class ToolPolicyGuardrail(CustomGuardrail): """ - Guardrail that enforces per-tool call policies stored in LiteLLM_ToolTable - and key/team allowed_tools (allowlist). No DB in hot path for policy lookup — - uses shared cache (user_api_key_cache). + Guardrail that enforces per-tool call policies from the in-memory + ToolPolicyRegistry (synced from DB). Key/team allowed_tools (allowlist) still + enforced. No DB or cache in hot path — registry lookups only. """ def __init__(self, **kwargs: Any) -> None: @@ -114,7 +113,6 @@ class ToolPolicyGuardrail(CustomGuardrail): GuardrailEventHooks.during_call, ] super().__init__(**kwargs) - self._policy_cache: DualCache = DualCache() @log_guardrail_information async def apply_guardrail( @@ -177,11 +175,18 @@ class ToolPolicyGuardrail(CustomGuardrail): object_permission_id, team_object_permission_id = ( _get_request_object_permission_ids(request_data) ) - policy_map = await self._get_policies_cached( - tool_names, - object_permission_id=object_permission_id, - team_object_permission_id=team_object_permission_id, - ) + from litellm.proxy.db.tool_registry_writer import \ + get_tool_policy_registry + + registry = get_tool_policy_registry() + if not registry.is_initialized(): + policy_map = {} + else: + policy_map = registry.get_effective_policies( + tool_names, + object_permission_id=object_permission_id, + team_object_permission_id=team_object_permission_id, + ) blocked = [name for name in tool_names if policy_map.get(name) == "blocked"] if blocked: verbose_proxy_logger.warning( @@ -197,28 +202,3 @@ class ToolPolicyGuardrail(CustomGuardrail): ) return inputs - - async def _get_policies_cached( - self, - tool_names: List[str], - object_permission_id: Optional[str] = None, - team_object_permission_id: Optional[str] = None, - ) -> Dict[str, str]: - """ - Fetch effective call_policy (blocked if in object permission blocked_tools, - else global) via shared cache to avoid DB in hot path. - """ - from litellm.proxy.db.tool_registry_writer import \ - get_tool_policies_cached - from litellm.proxy.proxy_server import (prisma_client, - user_api_key_cache) - - if not tool_names: - return {} - return await get_tool_policies_cached( - tool_names=tool_names, - cache=user_api_key_cache, - prisma_client=prisma_client, - object_permission_id=object_permission_id, - team_object_permission_id=team_object_permission_id, - ) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index cbafe5bb390..f9cc771e041 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -4223,6 +4223,9 @@ class ProxyConfig: if self._should_load_db_object(object_type="search_tools"): await self._init_search_tools_in_db(prisma_client=prisma_client) + if self._should_load_db_object(object_type="tools"): + await self._init_tool_policy_in_db(prisma_client=prisma_client) + if self._should_load_db_object(object_type="model_cost_map"): await self._check_and_reload_model_cost_map(prisma_client=prisma_client) @@ -4656,6 +4659,25 @@ class ProxyConfig: ) ) + async def _init_tool_policy_in_db(self, prisma_client: PrismaClient): + """ + Initialize tool policy from database into the in-memory registry. + Synced periodically by add_deployment -> _init_non_llm_objects_in_db. + """ + from litellm.proxy.db.tool_registry_writer import \ + get_tool_policy_registry + + try: + registry = get_tool_policy_registry() + await registry.sync_tool_policy_from_db(prisma_client=prisma_client) + verbose_proxy_logger.debug("Successfully synced tool policy from DB") + except Exception as e: + verbose_proxy_logger.exception( + "litellm.proxy.proxy_server.py::ProxyConfig:_init_tool_policy_in_db - {}".format( + str(e) + ) + ) + async def _init_vector_stores_in_db(self, prisma_client: PrismaClient): from litellm.vector_stores.vector_store_registry import \ VectorStoreRegistry diff --git a/tests/test_litellm/proxy/db/test_tool_registry_writer.py b/tests/test_litellm/proxy/db/test_tool_registry_writer.py index f1b829b483f..d7172b1f2b8 100644 --- a/tests/test_litellm/proxy/db/test_tool_registry_writer.py +++ b/tests/test_litellm/proxy/db/test_tool_registry_writer.py @@ -12,8 +12,10 @@ import pytest sys.path.insert(0, os.path.abspath("../../..")) -from litellm.proxy.db.tool_registry_writer import (batch_upsert_tools, +from litellm.proxy.db.tool_registry_writer import (ToolPolicyRegistry, + batch_upsert_tools, get_tool, + get_tool_policy_registry, get_tools_by_names, list_tools, update_tool_policy) @@ -208,3 +210,65 @@ async def test_get_tools_by_names_empty_list(): result = await get_tools_by_names(prisma, []) assert result == {} prisma.db.litellm_tooltable.find_many.assert_not_awaited() + + +# --- ToolPolicyRegistry --- + + +def _mock_tool_row(tool_name: str, call_policy: str = "untrusted"): + row = MagicMock() + row.tool_name = tool_name + row.call_policy = call_policy + return row + + +def _mock_perm_row(object_permission_id: str, blocked_tools: list): + row = MagicMock() + row.object_permission_id = object_permission_id + row.blocked_tools = blocked_tools + return row + + +@pytest.mark.asyncio +async def test_tool_policy_registry_sync_and_get_effective_policies(): + """Registry syncs from DB; get_effective_policies returns merged blocked + global.""" + prisma = MagicMock() + prisma.db.litellm_tooltable.find_many = AsyncMock( + return_value=[ + _mock_tool_row("tool_a", "trusted"), + _mock_tool_row("tool_b", "blocked"), + _mock_tool_row("tool_c", "untrusted"), + ] + ) + prisma.db.litellm_objectpermissiontable.find_many = AsyncMock( + return_value=[ + _mock_perm_row("op-key-1", ["tool_a"]), + _mock_perm_row("op-team-1", ["tool_c"]), + ] + ) + registry = get_tool_policy_registry() + await registry.sync_tool_policy_from_db(prisma) + assert registry.is_initialized() + # Key blocked: tool_a. Team blocked: tool_c. Global: tool_b blocked. + result = registry.get_effective_policies( + ["tool_a", "tool_b", "tool_c"], + object_permission_id="op-key-1", + team_object_permission_id="op-team-1", + ) + assert result["tool_a"] == "blocked" + assert result["tool_b"] == "blocked" + assert result["tool_c"] == "blocked" + # No op ids: only global + result_global = registry.get_effective_policies(["tool_a", "tool_b", "tool_c"]) + assert result_global["tool_a"] == "trusted" + assert result_global["tool_b"] == "blocked" + assert result_global["tool_c"] == "untrusted" + + +@pytest.mark.asyncio +async def test_tool_policy_registry_not_initialized_returns_untrusted(): + """When not synced, get_effective_policies still returns untrusted for unknown tools.""" + registry = ToolPolicyRegistry() + assert not registry.is_initialized() + result = registry.get_effective_policies(["unknown_tool"]) + assert result == {"unknown_tool": "untrusted"} diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_policy_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_policy_guardrail.py index c6a81efbf0b..943a8d4be75 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_policy_guardrail.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_tool_policy_guardrail.py @@ -12,9 +12,8 @@ from fastapi import HTTPException sys.path.insert(0, os.path.abspath("../../../../../..")) -from litellm.proxy.guardrails.guardrail_hooks.tool_policy.tool_policy_guardrail import ( - ToolPolicyGuardrail, -) +from litellm.proxy.guardrails.guardrail_hooks.tool_policy.tool_policy_guardrail import \ + ToolPolicyGuardrail from litellm.types.guardrails import GuardrailEventHooks @@ -70,10 +69,21 @@ async def test_no_tool_calls_in_response_passes_through(guardrail): assert result is inputs +def _registry_mock(policy_map: dict): + """Return a mock registry with is_initialized=True and get_effective_policies returning policy_map.""" + reg = MagicMock() + reg.is_initialized.return_value = True + reg.get_effective_policies.return_value = policy_map + return reg + + @pytest.mark.asyncio async def test_untrusted_tools_pass_through(guardrail): policy_map = {"search": "untrusted", "read_file": "trusted"} - with patch.object(guardrail, "_get_policies_cached", new=AsyncMock(return_value=policy_map)): + with patch( + "litellm.proxy.db.tool_registry_writer.get_tool_policy_registry", + return_value=_registry_mock(policy_map), + ): inputs: Any = _tool_request_inputs(["search", "read_file"]) result = await guardrail.apply_guardrail( inputs=inputs, request_data={}, input_type="request" @@ -84,7 +94,10 @@ async def test_untrusted_tools_pass_through(guardrail): @pytest.mark.asyncio async def test_blocked_tool_in_request_raises_http_exception(guardrail): policy_map = {"dangerous_tool": "blocked"} - with patch.object(guardrail, "_get_policies_cached", new=AsyncMock(return_value=policy_map)): + with patch( + "litellm.proxy.db.tool_registry_writer.get_tool_policy_registry", + return_value=_registry_mock(policy_map), + ): inputs: Any = _tool_request_inputs(["dangerous_tool"]) with pytest.raises(HTTPException) as exc_info: await guardrail.apply_guardrail( @@ -97,7 +110,10 @@ async def test_blocked_tool_in_request_raises_http_exception(guardrail): @pytest.mark.asyncio async def test_blocked_tool_in_response_raises_http_exception(guardrail): policy_map = {"exfil_tool": "blocked"} - with patch.object(guardrail, "_get_policies_cached", new=AsyncMock(return_value=policy_map)): + with patch( + "litellm.proxy.db.tool_registry_writer.get_tool_policy_registry", + return_value=_registry_mock(policy_map), + ): inputs: Any = _tool_response_inputs(["exfil_tool"]) with pytest.raises(HTTPException) as exc_info: await guardrail.apply_guardrail( @@ -110,7 +126,10 @@ async def test_blocked_tool_in_response_raises_http_exception(guardrail): @pytest.mark.asyncio async def test_mixed_blocked_and_allowed_raises_for_blocked(guardrail): policy_map = {"safe_tool": "trusted", "bad_tool": "blocked"} - with patch.object(guardrail, "_get_policies_cached", new=AsyncMock(return_value=policy_map)): + with patch( + "litellm.proxy.db.tool_registry_writer.get_tool_policy_registry", + return_value=_registry_mock(policy_map), + ): inputs: Any = _tool_request_inputs(["safe_tool", "bad_tool"]) with pytest.raises(HTTPException) as exc_info: await guardrail.apply_guardrail( @@ -123,8 +142,11 @@ async def test_mixed_blocked_and_allowed_raises_for_blocked(guardrail): @pytest.mark.asyncio async def test_tool_not_in_db_passes_through(guardrail): - """Tools not found in the DB (no entry) should not be blocked.""" - with patch.object(guardrail, "_get_policies_cached", new=AsyncMock(return_value={})): + """When registry returns no policy (or empty), tools are not blocked.""" + with patch( + "litellm.proxy.db.tool_registry_writer.get_tool_policy_registry", + return_value=_registry_mock({}), + ): inputs: Any = _tool_request_inputs(["unknown_tool"]) result = await guardrail.apply_guardrail( inputs=inputs, request_data={}, input_type="request" @@ -133,43 +155,30 @@ async def test_tool_not_in_db_passes_through(guardrail): @pytest.mark.asyncio -async def test_get_policies_cached_uses_cache(guardrail): - """Second call with same tool names should return the cached result.""" - policy_map = {"tool_a": "trusted"} +async def test_registry_not_initialized_passes_through(guardrail): + """When registry is not initialized, no tools are blocked (empty policy map).""" + reg = MagicMock() + reg.is_initialized.return_value = False with patch( - "litellm.proxy.db.tool_registry_writer.get_tools_by_names", - new=AsyncMock(return_value=policy_map), - ) as mock_db, patch( - "litellm.proxy.proxy_server.prisma_client", - new=MagicMock(), + "litellm.proxy.db.tool_registry_writer.get_tool_policy_registry", + return_value=reg, ): - # first call — should hit DB - result1 = await guardrail._get_policies_cached(["tool_a"]) - assert result1 == policy_map - - # second call — should hit cache, not DB again - result2 = await guardrail._get_policies_cached(["tool_a"]) - assert result2 == policy_map - - assert mock_db.call_count == 1 - - -@pytest.mark.asyncio -async def test_get_policies_cached_no_prisma(guardrail): - """Without a prisma client, returns empty dict.""" - with patch( - "litellm.proxy.proxy_server.prisma_client", - None, - ): - result = await guardrail._get_policies_cached(["tool_a"]) - assert result == {} + inputs: Any = _tool_request_inputs(["any_tool"]) + result = await guardrail.apply_guardrail( + inputs=inputs, request_data={}, input_type="request" + ) + assert result is inputs + reg.get_effective_policies.assert_not_called() @pytest.mark.asyncio async def test_response_tool_calls_as_objects(guardrail): """tool_calls that are objects (not dicts) with .function.name should work.""" policy_map = {"obj_tool": "blocked"} - with patch.object(guardrail, "_get_policies_cached", new=AsyncMock(return_value=policy_map)): + with patch( + "litellm.proxy.db.tool_registry_writer.get_tool_policy_registry", + return_value=_registry_mock(policy_map), + ): fn = MagicMock() fn.name = "obj_tool" tc = MagicMock() diff --git a/ui/litellm-dashboard/src/components/ToolDetail.tsx b/ui/litellm-dashboard/src/components/ToolDetail.tsx index 33bfa9b4357..4665a72eaea 100644 --- a/ui/litellm-dashboard/src/components/ToolDetail.tsx +++ b/ui/litellm-dashboard/src/components/ToolDetail.tsx @@ -2,9 +2,11 @@ import { ArrowLeftOutlined, HistoryOutlined, ToolOutlined } from "@ant-design/icons"; import { useQuery, useQueryClient } from "@tanstack/react-query"; -import { Button, Pagination, Select, Spin, Table } from "antd"; +import { Button, Select, Spin } from "antd"; import React, { useCallback, useMemo, useState } from "react"; import TeamDropdown from "@/components/common_components/team_dropdown"; +import { LogViewer } from "@/components/GuardrailsMonitor/LogViewer"; +import type { LogEntry } from "@/components/GuardrailsMonitor/mockData"; import { PolicySelect } from "@/components/ToolPolicies/PolicySelect"; import { deleteToolPolicyOverride, @@ -12,14 +14,10 @@ import { getToolUsageLogs, keyListCall, teamListCall, - uiSpendLogsCall, updateToolPolicy, type ToolPolicyOverrideRow, - type ToolUsageLogEntry, } from "@/components/networking"; import type { Team } from "@/components/key_team_helpers/key_list"; -import type { LogEntry } from "@/components/view_logs/columns"; -import { LogDetailsDrawer } from "@/components/view_logs/LogDetailsDrawer"; interface ToolDetailProps { toolName: string; @@ -39,7 +37,7 @@ interface KeyOption { const TOOL_DETAIL_QUERY_KEY = "tool-detail"; -const LOGS_PAGE_SIZE = 20; +const LOGS_PAGE_SIZE = 50; function getDefaultLogsDateRange(): { start: string; end: string } { const end = new Date(); @@ -57,9 +55,6 @@ export function ToolDetail({ toolName, onBack, accessToken }: ToolDetailProps) { const [blockScope, setBlockScope] = useState<"team" | "key">("team"); const [blockTeamId, setBlockTeamId] = useState(null); const [blockKey, setBlockKey] = useState(null); - const [logsPage, setLogsPage] = useState(1); - const [selectedRequestId, setSelectedRequestId] = useState(null); - const [drawerOpen, setDrawerOpen] = useState(false); const logsDateRange = useMemo(() => getDefaultLogsDateRange(), []); @@ -82,33 +77,27 @@ export function ToolDetail({ toolName, onBack, accessToken }: ToolDetailProps) { }); const { data: logsData, isLoading: logsLoading } = useQuery({ - queryKey: ["tool-usage-logs", toolName, logsPage], + queryKey: ["tool-usage-logs", toolName, logsDateRange.start, logsDateRange.end], queryFn: () => getToolUsageLogs(accessToken!, toolName, { - page: logsPage, + page: 1, pageSize: LOGS_PAGE_SIZE, + startDate: logsDateRange.start, + endDate: logsDateRange.end, }), enabled: !!accessToken && !!toolName, }); - const { data: fullLogResponse } = useQuery({ - queryKey: ["spend-log-by-request-tool-detail", selectedRequestId, logsDateRange.start, logsDateRange.end], - queryFn: async () => { - if (!accessToken || !selectedRequestId) return null; - const res = await uiSpendLogsCall({ - accessToken, - start_date: logsDateRange.start, - end_date: logsDateRange.end, - page: 1, - page_size: 10, - params: { request_id: selectedRequestId }, - }); - return res as { data: LogEntry[]; total: number }; - }, - enabled: !!accessToken && !!selectedRequestId && drawerOpen, - }); - - const selectedLog: LogEntry | null = fullLogResponse?.data?.[0] ?? null; + const logs: LogEntry[] = useMemo(() => { + const list = logsData?.logs ?? []; + return list.map((l) => ({ + id: l.id, + timestamp: l.timestamp, + action: "passed" as const, + model: l.model ?? undefined, + input_snippet: l.input_snippet ?? undefined, + })); + }, [logsData?.logs]); const teams: Team[] = useMemo(() => { const arr = Array.isArray(teamsData) ? teamsData : teamsData?.data ?? []; @@ -367,98 +356,18 @@ export function ToolDetail({ toolName, onBack, accessToken }: ToolDetailProps) { Recent logs -

- Requests that used this tool. Click a row to open full log details. -

- {logsLoading && ( -
- -
- )} - {!logsLoading && (!logsData?.logs?.length) && ( -
- No logs for this tool yet. Usage will appear here after requests that call this tool. -
- )} - {!logsLoading && logsData && logsData.logs.length > 0 && ( - <> - - dataSource={logsData.logs} - rowKey="id" - size="small" - pagination={false} - onRow={(record) => ({ - onClick: () => { - setSelectedRequestId(record.id); - setDrawerOpen(true); - }, - style: { cursor: "pointer" }, - })} - columns={[ - { - title: "Time", - dataIndex: "timestamp", - key: "timestamp", - width: 200, - render: (ts: string) => - ts ? new Date(ts).toLocaleString(undefined, { dateStyle: "short", timeStyle: "short" }) : "—", - }, - { - title: "Model", - dataIndex: "model", - key: "model", - ellipsis: true, - render: (v: string | null) => v ?? "—", - }, - { - title: "Spend", - dataIndex: "spend", - key: "spend", - width: 90, - render: (v: number | null) => - v != null ? `$${Number(v).toFixed(4)}` : "—", - }, - { - title: "Tokens", - dataIndex: "total_tokens", - key: "total_tokens", - width: 90, - render: (v: number | null) => (v != null ? v.toLocaleString() : "—"), - }, - { - title: "Input", - dataIndex: "input_snippet", - key: "input_snippet", - ellipsis: true, - render: (v: string | null) => (v ? String(v).slice(0, 80) + (String(v).length > 80 ? "…" : "") : "—"), - }, - ]} - /> -
- -
- - )} +
- - { - setDrawerOpen(false); - setSelectedRequestId(null); - }} - logEntry={selectedLog} - accessToken={accessToken} - allLogs={selectedLog ? [selectedLog] : []} - startTime={logsDateRange.start} - />
); } From 0f20cdd10d991ee7e6ed8fd4ad9b3281403df17b Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Thu, 26 Feb 2026 16:41:04 -0800 Subject: [PATCH 11/12] fix: address PR review feedback for tool access control - Add missing blocked_tools column to root schema.prisma (schema drift) - Invalidate ToolPolicyRegistry after policy mutations so changes take effect immediately - Remove dead code: unused get_effective_policies, get_tool_policies_cached, and helpers Co-Authored-By: Claude Opus 4.6 --- litellm/proxy/db/tool_registry_writer.py | 136 +----------------- .../tool_management_endpoints.py | 44 ++++-- schema.prisma | 1 + 3 files changed, 36 insertions(+), 145 deletions(-) diff --git a/litellm/proxy/db/tool_registry_writer.py b/litellm/proxy/db/tool_registry_writer.py index 1c106df675b..5a6512f1374 100644 --- a/litellm/proxy/db/tool_registry_writer.py +++ b/litellm/proxy/db/tool_registry_writer.py @@ -10,18 +10,16 @@ from datetime import datetime, timezone from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union from litellm._logging import verbose_proxy_logger -from litellm.caching.dual_cache import DualCache -from litellm.constants import TOOL_POLICY_CACHE_TTL_SECONDS from litellm.proxy._types import ToolDiscoveryQueueItem -from litellm.types.tool_management import (LiteLLM_ToolTableRow, - ToolCallPolicy, - ToolPolicyOverrideRow) +from litellm.types.tool_management import ( + LiteLLM_ToolTableRow, + ToolCallPolicy, + ToolPolicyOverrideRow, +) if TYPE_CHECKING: from litellm.proxy.utils import PrismaClient -TOOL_POLICY_CACHE_KEY_PREFIX = "tool_policy:" - def _row_to_model(row: Union[dict, Any]) -> LiteLLM_ToolTableRow: """Convert a Prisma model instance or dict to LiteLLM_ToolTableRow.""" @@ -267,67 +265,6 @@ async def list_overrides_for_tool( return [] -async def _get_merged_blocked_tools( - prisma_client: "PrismaClient", - object_permission_id: Optional[str], - team_object_permission_id: Optional[str], -) -> set: - """Return union of blocked_tools from key and team object permissions.""" - blocked: set = set() - for op_id in (object_permission_id, team_object_permission_id): - if not op_id or not op_id.strip(): - continue - try: - row = await prisma_client.db.litellm_objectpermissiontable.find_unique( - where={"object_permission_id": op_id.strip()}, - ) - if row is not None and getattr(row, "blocked_tools", None): - blocked.update(row.blocked_tools) - except Exception as e: - verbose_proxy_logger.debug( - "tool_registry_writer _get_merged_blocked_tools error for %s: %s", - op_id, - e, - ) - return blocked - - -async def get_effective_policies( - prisma_client: "PrismaClient", - tool_names: List[str], - object_permission_id: Optional[str] = None, - team_object_permission_id: Optional[str] = None, -) -> Dict[str, str]: - """ - Return effective call_policy per tool: if tool is in key/team object permission - blocked_tools then "blocked", otherwise global policy from LiteLLM_ToolTable. - """ - if not tool_names: - return {} - try: - blocked = await _get_merged_blocked_tools( - prisma_client=prisma_client, - object_permission_id=object_permission_id, - team_object_permission_id=team_object_permission_id, - ) - result: Dict[str, str] = {} - for name in tool_names: - if name in blocked: - result[name] = "blocked" - missing = [n for n in tool_names if n not in result] - if missing: - global_map = await get_tools_by_names( - prisma_client=prisma_client, tool_names=missing - ) - result.update(global_map) - return result - except Exception as e: - verbose_proxy_logger.error( - "tool_registry_writer get_effective_policies error: %s", e - ) - return {} - - class ToolPolicyRegistry: """ In-memory registry of tool policies synced from DB. @@ -407,69 +344,6 @@ def get_tool_policy_registry() -> ToolPolicyRegistry: return _tool_policy_registry -def _effective_cache_suffix( - object_permission_id: Optional[str], - team_object_permission_id: Optional[str], -) -> str: - """Cache key suffix so different request contexts get correct policies.""" - return f":{object_permission_id or ''}:{team_object_permission_id or ''}" - - -async def get_tool_policies_cached( - tool_names: List[str], - cache: DualCache, - prisma_client: Optional["PrismaClient"], - object_permission_id: Optional[str] = None, - team_object_permission_id: Optional[str] = None, -) -> Dict[str, str]: - """ - Return effective call_policy per tool (blocked if in object permission blocked_tools, - else global). Cache-first; cache key includes object_permission_id(s). - """ - if not tool_names: - return {} - suffix = _effective_cache_suffix(object_permission_id, team_object_permission_id) - result: Dict[str, str] = {} - cache_misses: List[str] = [] - for name in tool_names: - key = f"{TOOL_POLICY_CACHE_KEY_PREFIX}{name}{suffix}" - cached = await cache.async_get_cache(key=key) - if cached is not None and isinstance(cached, str): - result[name] = cached - else: - cache_misses.append(name) - if cache_misses and prisma_client is not None: - try: - if object_permission_id or team_object_permission_id: - fetched = await get_effective_policies( - prisma_client=prisma_client, - tool_names=cache_misses, - object_permission_id=object_permission_id, - team_object_permission_id=team_object_permission_id, - ) - else: - fetched = await get_tools_by_names( - prisma_client=prisma_client, tool_names=cache_misses - ) - for name, policy in fetched.items(): - result[name] = policy - await cache.async_set_cache( - key=f"{TOOL_POLICY_CACHE_KEY_PREFIX}{name}{suffix}", - value=policy, - ttl=TOOL_POLICY_CACHE_TTL_SECONDS, - ) - verbose_proxy_logger.debug( - "get_tool_policies_cached: fetched %d from DB (hits: %d)", - len(cache_misses), - len(tool_names) - len(cache_misses), - ) - except Exception as e: - verbose_proxy_logger.error( - "tool_registry_writer get_tool_policies_cached error: %s", e - ) - return result - - async def add_tool_to_object_permission_blocked( prisma_client: "PrismaClient", object_permission_id: str, diff --git a/litellm/proxy/management_endpoints/tool_management_endpoints.py b/litellm/proxy/management_endpoints/tool_management_endpoints.py index 90115b13ab3..90c4fff5729 100644 --- a/litellm/proxy/management_endpoints/tool_management_endpoints.py +++ b/litellm/proxy/management_endpoints/tool_management_endpoints.py @@ -19,13 +19,16 @@ if TYPE_CHECKING: from litellm._logging import verbose_proxy_logger from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.types.tool_management import (LiteLLM_ToolTableRow, - ToolCallPolicy, ToolDetailResponse, - ToolListResponse, - ToolPolicyUpdateRequest, - ToolPolicyUpdateResponse, - ToolUsageLogEntry, - ToolUsageLogsResponse) +from litellm.types.tool_management import ( + LiteLLM_ToolTableRow, + ToolCallPolicy, + ToolDetailResponse, + ToolListResponse, + ToolPolicyUpdateRequest, + ToolPolicyUpdateResponse, + ToolUsageLogEntry, + ToolUsageLogsResponse, +) router = APIRouter() @@ -46,8 +49,7 @@ async def list_tools( Parameters: - call_policy: Optional filter — one of "trusted", "untrusted", "dual_llm", "blocked" """ - from litellm.proxy.db.tool_registry_writer import \ - list_tools as db_list_tools + from litellm.proxy.db.tool_registry_writer import list_tools as db_list_tools from litellm.proxy.proxy_server import prisma_client if prisma_client is None: @@ -377,9 +379,12 @@ async def update_tool_policy( """ from litellm.proxy.db.tool_registry_writer import ( add_tool_to_object_permission_blocked, - remove_tool_from_object_permission_blocked) - from litellm.proxy.db.tool_registry_writer import \ - update_tool_policy as db_update_tool_policy + get_tool_policy_registry, + remove_tool_from_object_permission_blocked, + ) + from litellm.proxy.db.tool_registry_writer import ( + update_tool_policy as db_update_tool_policy, + ) from litellm.proxy.proxy_server import prisma_client if prisma_client is None: @@ -424,6 +429,9 @@ async def update_tool_policy( status_code=500, detail=f"Failed to update policy override for tool '{data.tool_name}'", ) + registry = get_tool_policy_registry() + if registry.is_initialized(): + await registry.sync_tool_policy_from_db(prisma_client) return ToolPolicyUpdateResponse( tool_name=data.tool_name, call_policy=data.call_policy, @@ -442,6 +450,9 @@ async def update_tool_policy( status_code=500, detail=f"Failed to update policy for tool '{data.tool_name}'", ) + registry = get_tool_policy_registry() + if registry.is_initialized(): + await registry.sync_tool_policy_from_db(prisma_client) return ToolPolicyUpdateResponse( tool_name=updated.tool_name, call_policy=updated.call_policy, @@ -473,8 +484,10 @@ async def delete_tool_policy_override( Remove a policy override for a tool. Specify the override by team_id or key_hash (exactly one required). """ - from litellm.proxy.db.tool_registry_writer import \ - remove_tool_from_object_permission_blocked + from litellm.proxy.db.tool_registry_writer import ( + get_tool_policy_registry, + remove_tool_from_object_permission_blocked, + ) from litellm.proxy.proxy_server import prisma_client if prisma_client is None: @@ -515,6 +528,9 @@ async def delete_tool_policy_override( status_code=404, detail=f"No override found for tool '{tool_name}' with the given scope", ) + registry = get_tool_policy_registry() + if registry.is_initialized(): + await registry.sync_tool_policy_from_db(prisma_client) return {"deleted": True, "tool_name": tool_name} except HTTPException: raise diff --git a/schema.prisma b/schema.prisma index 691883ef446..cd4f9a4d247 100644 --- a/schema.prisma +++ b/schema.prisma @@ -260,6 +260,7 @@ model LiteLLM_ObjectPermissionTable { vector_stores String[] @default([]) agents String[] @default([]) agent_access_groups String[] @default([]) + blocked_tools String[] @default([]) // Tool names blocked for any key/team/user with this permission teams LiteLLM_TeamTable[] projects LiteLLM_ProjectTable[] verification_tokens LiteLLM_VerificationToken[] From 5ff48b3bab13fe4aeb08a21818b29cac7ef55bc6 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Thu, 26 Feb 2026 19:53:30 -0800 Subject: [PATCH 12/12] fix: race condition in permission resolution and remove duplicate allowlist check - Use atomic update_many with object_permission_id=None to prevent concurrent requests from creating orphaned permission rows and losing tool blocks - Remove duplicate allowed_tools enforcement from guardrail (already enforced in auth layer via check_tools_allowlist) - Move inline uuid import to module level Co-Authored-By: Claude Opus 4.6 --- .../tool_policy/tool_policy_guardrail.py | 60 ++++--------------- .../tool_management_endpoints.py | 43 +++++++++---- 2 files changed, 43 insertions(+), 60 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py index b6cb4bba138..2d91febd5e4 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/tool_policy/tool_policy_guardrail.py @@ -23,21 +23,21 @@ or both pre and post call: mode: during_call # runs before LLM and on response """ -from typing import TYPE_CHECKING, Any, List, Literal, Optional, Tuple +from typing import TYPE_CHECKING, Any, Literal, Optional, Tuple from fastapi import HTTPException from litellm._logging import verbose_proxy_logger -from litellm.integrations.custom_guardrail import (CustomGuardrail, - log_guardrail_information) -from litellm.proxy.guardrails.tool_name_extraction import \ - extract_request_tool_names +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + log_guardrail_information, +) +from litellm.proxy.guardrails.tool_name_extraction import extract_request_tool_names from litellm.types.guardrails import GuardrailEventHooks from litellm.types.utils import GenericGuardrailAPIInputs if TYPE_CHECKING: - from litellm.litellm_core_utils.litellm_logging import \ - Logging as LiteLLMLoggingObj + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj GUARDRAIL_NAME = "tool_policy" @@ -80,29 +80,12 @@ def _get_request_route_from_data(request_data: dict) -> Optional[str]: return meta.get("user_api_key_request_route") -def _get_effective_allowed_tools_from_request( - request_data: dict, -) -> Optional[List[str]]: - """Key allowed_tools overrides team; empty/missing means no restriction.""" - meta = request_data.get("metadata") or request_data.get("litellm_metadata") or {} - key_meta = meta.get("user_api_key_metadata") or {} - team_meta = meta.get("user_api_key_team_metadata") or {} - key_allowed = key_meta.get("allowed_tools") if isinstance(key_meta, dict) else None - team_allowed = ( - team_meta.get("allowed_tools") if isinstance(team_meta, dict) else None - ) - if isinstance(key_allowed, list) and len(key_allowed) > 0: - return key_allowed - if isinstance(team_allowed, list) and len(team_allowed) > 0: - return team_allowed - return None - - class ToolPolicyGuardrail(CustomGuardrail): """ Guardrail that enforces per-tool call policies from the in-memory - ToolPolicyRegistry (synced from DB). Key/team allowed_tools (allowlist) still - enforced. No DB or cache in hot path — registry lookups only. + ToolPolicyRegistry (synced from DB). Key/team allowed_tools (allowlist) is + enforced in the auth layer (check_tools_allowlist). No DB or cache in hot + path — registry lookups only. """ def __init__(self, **kwargs: Any) -> None: @@ -123,7 +106,7 @@ class ToolPolicyGuardrail(CustomGuardrail): logging_obj: Optional["LiteLLMLoggingObj"] = None, ) -> GenericGuardrailAPIInputs: """ - Enforce key/team allowlist then DB call_policy on request tools / response tool_calls. + Enforce DB call_policy on request tools / response tool_calls. """ if input_type == "request": tools = inputs.get("tools") or [] @@ -154,29 +137,10 @@ class ToolPolicyGuardrail(CustomGuardrail): if not tool_names: return inputs - allowed_tools = _get_effective_allowed_tools_from_request(request_data) - if isinstance(allowed_tools, list) and len(allowed_tools) > 0: - allowed_set = {str(t) for t in allowed_tools} - disallowed = [n for n in tool_names if n not in allowed_set] - if disallowed: - verbose_proxy_logger.warning( - "ToolPolicyGuardrail: tool(s) %s not in key/team allowed_tools", - disallowed, - ) - raise HTTPException( - status_code=400, - detail={ - "error": "Violated tool allowlist", - "disallowed_tools": disallowed, - "message": f"Tool(s) {disallowed} are not in the allowed tools list for this key/team.", - }, - ) - object_permission_id, team_object_permission_id = ( _get_request_object_permission_ids(request_data) ) - from litellm.proxy.db.tool_registry_writer import \ - get_tool_policy_registry + from litellm.proxy.db.tool_registry_writer import get_tool_policy_registry registry = get_tool_policy_registry() if not registry.is_initialized(): diff --git a/litellm/proxy/management_endpoints/tool_management_endpoints.py b/litellm/proxy/management_endpoints/tool_management_endpoints.py index 90c4fff5729..304a973535f 100644 --- a/litellm/proxy/management_endpoints/tool_management_endpoints.py +++ b/litellm/proxy/management_endpoints/tool_management_endpoints.py @@ -8,6 +8,7 @@ GET /v1/tool/{tool_name} - Get a single tool's details POST /v1/tool/policy - Update the call_policy for a tool """ +import uuid from datetime import datetime, timezone from typing import TYPE_CHECKING, Any, List, Optional @@ -310,17 +311,26 @@ async def _resolve_key_hash_to_object_permission_id( op_id = getattr(row, "object_permission_id", None) if op_id: return op_id - # Create new object permission and assign to key - import uuid as _uuid - - new_id = str(_uuid.uuid4()) + # Create new object permission and atomically assign to key. + # Uses update_many with object_permission_id=None to prevent race conditions: + # only one concurrent request wins; the loser cleans up its orphaned row. + new_id = str(uuid.uuid4()) await prisma_client.db.litellm_objectpermissiontable.create( data={"object_permission_id": new_id, "blocked_tools": []} ) - await prisma_client.db.litellm_verificationtoken.update( - where={"token": hashed}, + updated_count = await prisma_client.db.litellm_verificationtoken.update_many( + where={"token": hashed, "object_permission_id": None}, data={"object_permission_id": new_id}, ) + if updated_count == 0: + # Another request already assigned a permission; clean up orphan + await prisma_client.db.litellm_objectpermissiontable.delete( + where={"object_permission_id": new_id} + ) + row = await prisma_client.db.litellm_verificationtoken.find_unique( + where={"token": hashed} + ) + return getattr(row, "object_permission_id", None) if row else None return new_id @@ -331,8 +341,9 @@ async def _resolve_team_id_to_object_permission_id( """Resolve team_id to object_permission_id; create permission if team has none.""" if not team_id or not team_id.strip(): return None + team_id_clean = team_id.strip() row = await prisma_client.db.litellm_teamtable.find_unique( - where={"team_id": team_id.strip()}, + where={"team_id": team_id_clean}, select={"object_permission_id": True}, ) if row is None: @@ -340,16 +351,24 @@ async def _resolve_team_id_to_object_permission_id( op_id = getattr(row, "object_permission_id", None) if op_id: return op_id - import uuid as _uuid - - new_id = str(_uuid.uuid4()) + # Same atomic pattern as _resolve_key_hash_to_object_permission_id + new_id = str(uuid.uuid4()) await prisma_client.db.litellm_objectpermissiontable.create( data={"object_permission_id": new_id, "blocked_tools": []} ) - await prisma_client.db.litellm_teamtable.update( - where={"team_id": team_id.strip()}, + updated_count = await prisma_client.db.litellm_teamtable.update_many( + where={"team_id": team_id_clean, "object_permission_id": None}, data={"object_permission_id": new_id}, ) + if updated_count == 0: + await prisma_client.db.litellm_objectpermissiontable.delete( + where={"object_permission_id": new_id} + ) + row = await prisma_client.db.litellm_teamtable.find_unique( + where={"team_id": team_id_clean}, + select={"object_permission_id": True}, + ) + return getattr(row, "object_permission_id", None) if row else None return new_id