From 0f8832f05b74c7dabafcee72d46280c26f51a4c1 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Thu, 26 Feb 2026 00:26:07 -0800 Subject: [PATCH] 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": [ {