mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
feat(tools): show tool logs
This commit is contained in:
parent
3284df3bfa
commit
0f8832f05b
19 changed files with 1272 additions and 836 deletions
|
|
@ -1,2 +0,0 @@
|
|||
-- This is an empty migration.
|
||||
|
||||
|
|
@ -1,2 +0,0 @@
|
|||
-- This is an empty migration.
|
||||
|
||||
|
|
@ -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");
|
||||
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_ObjectPermissionTable" ADD COLUMN "blocked_tools" TEXT[] DEFAULT ARRAY[]::TEXT[];
|
||||
|
|
@ -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");
|
||||
|
|
@ -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())
|
||||
|
|
|
|||
147
litellm/proxy/db/spend_log_tool_index.py
Normal file
147
litellm/proxy/db/spend_log_tool_index.py
Normal file
|
|
@ -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
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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<string | null>(null);
|
||||
const [blockKey, setBlockKey] = useState<KeyOption | null>(null);
|
||||
const [logsPage, setLogsPage] = useState(1);
|
||||
const [selectedRequestId, setSelectedRequestId] = useState<string | null>(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) {
|
|||
</Button>
|
||||
</div>
|
||||
</section>
|
||||
|
||||
<section className="bg-white rounded-lg border border-gray-200 p-5 shadow-sm">
|
||||
<h2 className="text-sm font-semibold text-gray-700 mb-3 flex items-center gap-2">
|
||||
<HistoryOutlined />
|
||||
Recent logs
|
||||
</h2>
|
||||
<p className="text-xs text-gray-500 mb-4">
|
||||
Requests that used this tool. Click a row to open full log details.
|
||||
</p>
|
||||
{logsLoading && (
|
||||
<div className="flex justify-center py-8">
|
||||
<Spin />
|
||||
</div>
|
||||
)}
|
||||
{!logsLoading && (!logsData?.logs?.length) && (
|
||||
<div className="py-8 text-center text-sm text-gray-500">
|
||||
No logs for this tool yet. Usage will appear here after requests that call this tool.
|
||||
</div>
|
||||
)}
|
||||
{!logsLoading && logsData && logsData.logs.length > 0 && (
|
||||
<>
|
||||
<Table<ToolUsageLogEntry>
|
||||
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 ? "…" : "") : "—"),
|
||||
},
|
||||
]}
|
||||
/>
|
||||
<div className="mt-4 flex justify-end">
|
||||
<Pagination
|
||||
current={logsPage}
|
||||
pageSize={LOGS_PAGE_SIZE}
|
||||
total={logsData.total}
|
||||
showSizeChanger={false}
|
||||
onChange={setLogsPage}
|
||||
/>
|
||||
</div>
|
||||
</>
|
||||
)}
|
||||
</section>
|
||||
</div>
|
||||
|
||||
<LogDetailsDrawer
|
||||
open={drawerOpen}
|
||||
onClose={() => {
|
||||
setDrawerOpen(false);
|
||||
setSelectedRequestId(null);
|
||||
}}
|
||||
logEntry={selectedLog}
|
||||
accessToken={accessToken}
|
||||
allLogs={selectedLog ? [selectedLog] : []}
|
||||
startTime={logsDateRange.start}
|
||||
/>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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<ToolUsageLogsResponse> => {
|
||||
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
|
||||
|
|
|
|||
|
|
@ -14,7 +14,7 @@
|
|||
"moduleResolution": "bundler",
|
||||
"resolveJsonModule": true,
|
||||
"isolatedModules": true,
|
||||
"jsx": "react-jsx",
|
||||
"jsx": "preserve",
|
||||
"incremental": true,
|
||||
"plugins": [
|
||||
{
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue