mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix: tool mgmt endpoints
This commit is contained in:
parent
fe8a907169
commit
fc8810a865
4 changed files with 269 additions and 104 deletions
|
|
@ -2,18 +2,17 @@
|
|||
DB helpers for LiteLLM_ToolTable — the global tool registry.
|
||||
|
||||
Tools are auto-discovered from LLM responses and upserted here.
|
||||
Admins use the management endpoints to read and update call_policy.
|
||||
Admins use the management endpoints to read and update input_policy / output_policy.
|
||||
"""
|
||||
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import ToolDiscoveryQueueItem
|
||||
from litellm.types.tool_management import (
|
||||
LiteLLM_ToolTableRow,
|
||||
ToolCallPolicy,
|
||||
ToolPolicyOverrideRow,
|
||||
)
|
||||
|
||||
|
|
@ -33,12 +32,15 @@ def _row_to_model(row: Union[dict, Any]) -> LiteLLM_ToolTableRow:
|
|||
"tool_id",
|
||||
"tool_name",
|
||||
"origin",
|
||||
"call_policy",
|
||||
"input_policy",
|
||||
"output_policy",
|
||||
"call_count",
|
||||
"assignments",
|
||||
"key_hash",
|
||||
"team_id",
|
||||
"key_alias",
|
||||
"user_agent",
|
||||
"last_used_at",
|
||||
"created_at",
|
||||
"updated_at",
|
||||
"created_by",
|
||||
|
|
@ -49,12 +51,15 @@ def _row_to_model(row: Union[dict, Any]) -> LiteLLM_ToolTableRow:
|
|||
tool_id=row.get("tool_id", ""),
|
||||
tool_name=row.get("tool_name", ""),
|
||||
origin=row.get("origin"),
|
||||
call_policy=row.get("call_policy", "untrusted"),
|
||||
input_policy=row.get("input_policy") or "untrusted",
|
||||
output_policy=row.get("output_policy") or "untrusted",
|
||||
call_count=int(row.get("call_count") or 0),
|
||||
assignments=row.get("assignments"),
|
||||
key_hash=row.get("key_hash"),
|
||||
team_id=row.get("team_id"),
|
||||
key_alias=row.get("key_alias"),
|
||||
user_agent=row.get("user_agent"),
|
||||
last_used_at=row.get("last_used_at"),
|
||||
created_at=row.get("created_at"),
|
||||
updated_at=row.get("updated_at"),
|
||||
created_by=row.get("created_by"),
|
||||
|
|
@ -69,8 +74,8 @@ async def batch_upsert_tools(
|
|||
"""
|
||||
Batch-upsert tool registry rows via Prisma.
|
||||
|
||||
On first insert: sets call_policy = "untrusted" (schema default), call_count = 1.
|
||||
On conflict: increments call_count; preserves existing call_policy.
|
||||
On first insert: sets input_policy/output_policy = "untrusted" (default), call_count = 1.
|
||||
On conflict: increments call_count; preserves existing policies.
|
||||
"""
|
||||
if not items:
|
||||
return
|
||||
|
|
@ -87,6 +92,7 @@ async def batch_upsert_tools(
|
|||
key_hash = item.get("key_hash")
|
||||
team_id = item.get("team_id")
|
||||
key_alias = item.get("key_alias")
|
||||
user_agent = item.get("user_agent")
|
||||
await table.upsert(
|
||||
where={"tool_name": tool_name},
|
||||
data={
|
||||
|
|
@ -94,17 +100,21 @@ async def batch_upsert_tools(
|
|||
"tool_id": str(uuid.uuid4()),
|
||||
"tool_name": tool_name,
|
||||
"origin": origin,
|
||||
"call_policy": "untrusted",
|
||||
"input_policy": "untrusted",
|
||||
"output_policy": "untrusted",
|
||||
"call_count": 1,
|
||||
"created_by": created_by,
|
||||
"updated_by": created_by,
|
||||
"key_hash": key_hash,
|
||||
"team_id": team_id,
|
||||
"key_alias": key_alias,
|
||||
"user_agent": user_agent,
|
||||
"last_used_at": now,
|
||||
},
|
||||
"update": {
|
||||
"call_count": {"increment": 1},
|
||||
"updated_at": now,
|
||||
"last_used_at": now,
|
||||
},
|
||||
},
|
||||
)
|
||||
|
|
@ -119,11 +129,11 @@ async def batch_upsert_tools(
|
|||
|
||||
async def list_tools(
|
||||
prisma_client: "PrismaClient",
|
||||
call_policy: Optional[ToolCallPolicy] = None,
|
||||
input_policy: Optional[str] = None,
|
||||
) -> List[LiteLLM_ToolTableRow]:
|
||||
"""Return all tools, optionally filtered by call_policy."""
|
||||
"""Return all tools, optionally filtered by input_policy."""
|
||||
try:
|
||||
where = {"call_policy": call_policy} if call_policy is not None else {}
|
||||
where = {"input_policy": input_policy} if input_policy is not None else {}
|
||||
rows = await prisma_client.db.litellm_tooltable.find_many(
|
||||
where=where,
|
||||
order={"created_at": "desc"},
|
||||
|
|
@ -154,30 +164,39 @@ async def get_tool(
|
|||
async def update_tool_policy(
|
||||
prisma_client: "PrismaClient",
|
||||
tool_name: str,
|
||||
call_policy: ToolCallPolicy,
|
||||
updated_by: Optional[str],
|
||||
input_policy: Optional[str] = None,
|
||||
output_policy: Optional[str] = None,
|
||||
) -> Optional[LiteLLM_ToolTableRow]:
|
||||
"""Update the call_policy for a tool. Upserts the row if it does not exist yet."""
|
||||
"""Update input_policy and/or output_policy for a tool. Upserts the row if it does not exist yet."""
|
||||
try:
|
||||
_updated_by = updated_by or "system"
|
||||
now = datetime.now(timezone.utc)
|
||||
|
||||
create_data: dict = {
|
||||
"tool_id": str(uuid.uuid4()),
|
||||
"tool_name": tool_name,
|
||||
"input_policy": input_policy or "untrusted",
|
||||
"output_policy": output_policy or "untrusted",
|
||||
"created_by": _updated_by,
|
||||
"updated_by": _updated_by,
|
||||
"created_at": now,
|
||||
"updated_at": now,
|
||||
}
|
||||
update_data: dict = {
|
||||
"updated_by": _updated_by,
|
||||
"updated_at": now,
|
||||
}
|
||||
if input_policy is not None:
|
||||
update_data["input_policy"] = input_policy
|
||||
if output_policy is not None:
|
||||
update_data["output_policy"] = output_policy
|
||||
|
||||
await prisma_client.db.litellm_tooltable.upsert(
|
||||
where={"tool_name": tool_name},
|
||||
data={
|
||||
"create": {
|
||||
"tool_id": str(uuid.uuid4()),
|
||||
"tool_name": tool_name,
|
||||
"call_policy": call_policy,
|
||||
"created_by": _updated_by,
|
||||
"updated_by": _updated_by,
|
||||
"created_at": now,
|
||||
"updated_at": now,
|
||||
},
|
||||
"update": {
|
||||
"call_policy": call_policy,
|
||||
"updated_by": _updated_by,
|
||||
"updated_at": now,
|
||||
},
|
||||
"create": create_data,
|
||||
"update": update_data,
|
||||
},
|
||||
)
|
||||
return await get_tool(prisma_client, tool_name)
|
||||
|
|
@ -191,10 +210,9 @@ async def update_tool_policy(
|
|||
async def get_tools_by_names(
|
||||
prisma_client: "PrismaClient",
|
||||
tool_names: List[str],
|
||||
) -> Dict[str, str]:
|
||||
) -> Dict[str, Tuple[str, str]]:
|
||||
"""
|
||||
Return a {tool_name: call_policy} map for the given tool names.
|
||||
Used by the policy enforcement guardrail — single batch query, never N+1.
|
||||
Return a {tool_name: (input_policy, output_policy)} map for the given tool names.
|
||||
"""
|
||||
if not tool_names:
|
||||
return {}
|
||||
|
|
@ -202,7 +220,13 @@ async def get_tools_by_names(
|
|||
rows = await prisma_client.db.litellm_tooltable.find_many(
|
||||
where={"tool_name": {"in": tool_names}},
|
||||
)
|
||||
return {row.tool_name: row.call_policy for row in rows}
|
||||
return {
|
||||
row.tool_name: (
|
||||
getattr(row, "input_policy", "untrusted") or "untrusted",
|
||||
getattr(row, "output_policy", "untrusted") or "untrusted",
|
||||
)
|
||||
for row in rows
|
||||
}
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(
|
||||
"tool_registry_writer get_tools_by_names error: %s", e
|
||||
|
|
@ -238,7 +262,7 @@ async def list_overrides_for_tool(
|
|||
tool_name=tool_name,
|
||||
team_id=None,
|
||||
key_hash=getattr(t, "token", None),
|
||||
call_policy="blocked",
|
||||
input_policy="blocked",
|
||||
key_alias=getattr(t, "key_alias", None),
|
||||
created_at=None,
|
||||
updated_at=None,
|
||||
|
|
@ -251,7 +275,7 @@ async def list_overrides_for_tool(
|
|||
tool_name=tool_name,
|
||||
team_id=getattr(team, "team_id", None),
|
||||
key_hash=None,
|
||||
call_policy="blocked",
|
||||
input_policy="blocked",
|
||||
key_alias=getattr(team, "team_alias", None),
|
||||
created_at=None,
|
||||
updated_at=None,
|
||||
|
|
@ -268,12 +292,12 @@ async def list_overrides_for_tool(
|
|||
class ToolPolicyRegistry:
|
||||
"""
|
||||
In-memory registry of tool policies synced from DB.
|
||||
Synced in _init_tool_policy_in_db (from add_deployment / _init_non_llm_objects_in_db).
|
||||
Hot path uses get_effective_policies only — no DB, no cache.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._global_tool_policies: Dict[str, str] = {}
|
||||
self._tool_input_policies: Dict[str, str] = {}
|
||||
self._tool_output_policies: Dict[str, str] = {}
|
||||
self._blocked_tools_by_op_id: Dict[str, List[str]] = {}
|
||||
self._initialized: bool = False
|
||||
|
||||
|
|
@ -281,10 +305,17 @@ class ToolPolicyRegistry:
|
|||
return self._initialized
|
||||
|
||||
async def sync_tool_policy_from_db(self, prisma_client: "PrismaClient") -> None:
|
||||
"""Load all tool policies and object-permission blocked_tools from DB; replace in-memory state."""
|
||||
"""Load all tool policies and object-permission blocked_tools from DB."""
|
||||
try:
|
||||
tools = await prisma_client.db.litellm_tooltable.find_many()
|
||||
self._global_tool_policies = {row.tool_name: row.call_policy for row in tools}
|
||||
self._tool_input_policies = {
|
||||
row.tool_name: getattr(row, "input_policy", "untrusted") or "untrusted"
|
||||
for row in tools
|
||||
}
|
||||
self._tool_output_policies = {
|
||||
row.tool_name: getattr(row, "output_policy", "untrusted") or "untrusted"
|
||||
for row in tools
|
||||
}
|
||||
|
||||
perms = await prisma_client.db.litellm_objectpermissiontable.find_many()
|
||||
self._blocked_tools_by_op_id = {}
|
||||
|
|
@ -296,8 +327,8 @@ class ToolPolicyRegistry:
|
|||
|
||||
self._initialized = True
|
||||
verbose_proxy_logger.info(
|
||||
"ToolPolicyRegistry: synced %d global tool policies and %d object permissions from DB",
|
||||
len(self._global_tool_policies),
|
||||
"ToolPolicyRegistry: synced %d tool policies and %d object permissions from DB",
|
||||
len(self._tool_input_policies),
|
||||
len(self._blocked_tools_by_op_id),
|
||||
)
|
||||
except Exception as e:
|
||||
|
|
@ -306,6 +337,12 @@ class ToolPolicyRegistry:
|
|||
)
|
||||
raise
|
||||
|
||||
def get_input_policy(self, tool_name: str) -> str:
|
||||
return self._tool_input_policies.get(tool_name, "untrusted")
|
||||
|
||||
def get_output_policy(self, tool_name: str) -> str:
|
||||
return self._tool_output_policies.get(tool_name, "untrusted")
|
||||
|
||||
def get_effective_policies(
|
||||
self,
|
||||
tool_names: List[str],
|
||||
|
|
@ -313,8 +350,8 @@ class ToolPolicyRegistry:
|
|||
team_object_permission_id: Optional[str] = None,
|
||||
) -> Dict[str, str]:
|
||||
"""
|
||||
Return effective call_policy per tool from in-memory state.
|
||||
If tool is in key or team blocked_tools -> "blocked", else global policy or "untrusted".
|
||||
Return effective input_policy per tool from in-memory state.
|
||||
If tool is in key or team blocked_tools -> "blocked", else global input_policy or "untrusted".
|
||||
"""
|
||||
if not tool_names:
|
||||
return {}
|
||||
|
|
@ -329,7 +366,7 @@ class ToolPolicyRegistry:
|
|||
if name in blocked:
|
||||
result[name] = "blocked"
|
||||
else:
|
||||
result[name] = self._global_tool_policies.get(name, "untrusted")
|
||||
result[name] = self._tool_input_policies.get(name, "untrusted")
|
||||
return result
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,13 +1,16 @@
|
|||
"""
|
||||
Tool Policy Guardrail
|
||||
|
||||
Reads call_policy from LiteLLM_ToolTable and enforces it on LLM requests/responses.
|
||||
Reads input_policy / output_policy from LiteLLM_ToolTable and enforces them.
|
||||
|
||||
Policy values:
|
||||
"trusted" - allow through (no action)
|
||||
"untrusted" - allow through (no action; default for newly discovered tools)
|
||||
Input policy values:
|
||||
"untrusted" - allow through (default for newly discovered tools)
|
||||
"trusted" - only allow if conversation contains no untrusted tool output
|
||||
"blocked" - raise HTTPException, preventing the tool call
|
||||
"dual_llm" - (Phase 3) send to second LLM for verification; currently treated as allowed
|
||||
|
||||
Output policy values:
|
||||
"untrusted" - output may be tainted (default)
|
||||
"trusted" - output is verified safe
|
||||
|
||||
Configuration in proxy config YAML:
|
||||
guardrails:
|
||||
|
|
@ -15,15 +18,9 @@ Configuration in proxy config YAML:
|
|||
litellm_params:
|
||||
guardrail: tool_policy
|
||||
mode: post_call
|
||||
|
||||
or both pre and post call:
|
||||
- guardrail_name: "tool_policy"
|
||||
litellm_params:
|
||||
guardrail: tool_policy
|
||||
mode: during_call # runs before LLM and on response
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Literal, Optional, Tuple
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
|
|
@ -80,12 +77,33 @@ def _get_request_route_from_data(request_data: dict) -> Optional[str]:
|
|||
return meta.get("user_api_key_request_route")
|
||||
|
||||
|
||||
def _resolve_tool_names_from_messages(messages: List[dict]) -> Dict[str, str]:
|
||||
"""
|
||||
Build a map of tool_call_id -> tool_name from assistant messages' tool_calls.
|
||||
Used to resolve which tool produced each tool result in the conversation.
|
||||
"""
|
||||
mapping: Dict[str, str] = {}
|
||||
for msg in messages:
|
||||
if msg.get("role") != "assistant":
|
||||
continue
|
||||
tool_calls = msg.get("tool_calls") or []
|
||||
for tc in tool_calls:
|
||||
if isinstance(tc, dict):
|
||||
tc_id = tc.get("id")
|
||||
fn = (tc.get("function") or {}).get("name")
|
||||
else:
|
||||
tc_id = getattr(tc, "id", None)
|
||||
fn_obj = getattr(tc, "function", None)
|
||||
fn = getattr(fn_obj, "name", None) if fn_obj else None
|
||||
if tc_id and fn:
|
||||
mapping[tc_id] = fn
|
||||
return mapping
|
||||
|
||||
|
||||
class ToolPolicyGuardrail(CustomGuardrail):
|
||||
"""
|
||||
Guardrail that enforces per-tool call policies from the in-memory
|
||||
ToolPolicyRegistry (synced from DB). Key/team allowed_tools (allowlist) is
|
||||
enforced in the auth layer (check_tools_allowlist). No DB or cache in hot
|
||||
path — registry lookups only.
|
||||
Guardrail that enforces per-tool input/output policies from the in-memory
|
||||
ToolPolicyRegistry (synced from DB).
|
||||
"""
|
||||
|
||||
def __init__(self, **kwargs: Any) -> None:
|
||||
|
|
@ -106,7 +124,7 @@ class ToolPolicyGuardrail(CustomGuardrail):
|
|||
logging_obj: Optional["LiteLLMLoggingObj"] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
"""
|
||||
Enforce DB call_policy on request tools / response tool_calls.
|
||||
Enforce input_policy and output_policy trust chain on request tools / response tool_calls.
|
||||
"""
|
||||
if input_type == "request":
|
||||
tools = inputs.get("tools") or []
|
||||
|
|
@ -120,9 +138,8 @@ class ToolPolicyGuardrail(CustomGuardrail):
|
|||
if not tool_names:
|
||||
route = _get_request_route_from_data(request_data)
|
||||
if route:
|
||||
|
||||
tool_names = extract_request_tool_names(route, request_data)
|
||||
else: # response
|
||||
else:
|
||||
tool_calls = inputs.get("tool_calls") or []
|
||||
tool_names = []
|
||||
for tc in tool_calls:
|
||||
|
|
@ -144,17 +161,18 @@ class ToolPolicyGuardrail(CustomGuardrail):
|
|||
|
||||
registry = get_tool_policy_registry()
|
||||
if not registry.is_initialized():
|
||||
policy_map = {}
|
||||
else:
|
||||
policy_map = registry.get_effective_policies(
|
||||
tool_names,
|
||||
object_permission_id=object_permission_id,
|
||||
team_object_permission_id=team_object_permission_id,
|
||||
)
|
||||
return inputs
|
||||
|
||||
# Stage 1: Check for blocked tools (input_policy=blocked or per-key/team override)
|
||||
policy_map = registry.get_effective_policies(
|
||||
tool_names,
|
||||
object_permission_id=object_permission_id,
|
||||
team_object_permission_id=team_object_permission_id,
|
||||
)
|
||||
blocked = [name for name in tool_names if policy_map.get(name) == "blocked"]
|
||||
if blocked:
|
||||
verbose_proxy_logger.warning(
|
||||
"ToolPolicyGuardrail: blocking tool(s) %s (policy=blocked)", blocked
|
||||
"ToolPolicyGuardrail: blocking tool(s) %s (input_policy=blocked)", blocked
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
|
|
@ -165,4 +183,47 @@ class ToolPolicyGuardrail(CustomGuardrail):
|
|||
},
|
||||
)
|
||||
|
||||
# Stage 2: Trust chain enforcement (response path only)
|
||||
# For each tool with input_policy=trusted, check if conversation
|
||||
# contains output from tools with output_policy=untrusted
|
||||
if input_type == "response":
|
||||
trusted_input_tools = [
|
||||
name for name in tool_names if policy_map.get(name) == "trusted"
|
||||
]
|
||||
if trusted_input_tools:
|
||||
messages = request_data.get("messages") or []
|
||||
tc_id_to_name = _resolve_tool_names_from_messages(messages)
|
||||
|
||||
untrusted_sources: List[str] = []
|
||||
for msg in messages:
|
||||
if msg.get("role") != "tool":
|
||||
continue
|
||||
tool_call_id = msg.get("tool_call_id")
|
||||
source_tool = tc_id_to_name.get(tool_call_id, "") if tool_call_id else ""
|
||||
if not source_tool:
|
||||
continue
|
||||
if registry.get_output_policy(source_tool) == "untrusted":
|
||||
if source_tool not in untrusted_sources:
|
||||
untrusted_sources.append(source_tool)
|
||||
|
||||
if untrusted_sources:
|
||||
verbose_proxy_logger.warning(
|
||||
"ToolPolicyGuardrail: trust chain violation — %s require trusted input "
|
||||
"but conversation has untrusted output from %s",
|
||||
trusted_input_tools,
|
||||
untrusted_sources,
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "Violated tool policy",
|
||||
"blocked_tools": trusted_input_tools,
|
||||
"untrusted_sources": untrusted_sources,
|
||||
"message": (
|
||||
f"{', '.join(trusted_input_tools)} requires trusted input but "
|
||||
f"conversation contains untrusted output from {', '.join(untrusted_sources)}."
|
||||
),
|
||||
},
|
||||
)
|
||||
|
||||
return inputs
|
||||
|
|
|
|||
|
|
@ -4,8 +4,9 @@ TOOL POLICY MANAGEMENT
|
|||
All /tool management endpoints
|
||||
|
||||
GET /v1/tool/list - List all discovered tools and their policies
|
||||
GET /v1/tool/policy/options - List available input/output policy options with descriptions
|
||||
GET /v1/tool/{tool_name} - Get a single tool's details
|
||||
POST /v1/tool/policy - Update the call_policy for a tool
|
||||
POST /v1/tool/policy - Update the input_policy / output_policy for a tool
|
||||
"""
|
||||
|
||||
import uuid
|
||||
|
|
@ -22,9 +23,12 @@ 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,
|
||||
ToolInputPolicy,
|
||||
ToolListResponse,
|
||||
ToolOutputPolicy,
|
||||
ToolPolicyOption,
|
||||
ToolPolicyOptionsResponse,
|
||||
ToolPolicyUpdateRequest,
|
||||
ToolPolicyUpdateResponse,
|
||||
ToolUsageLogEntry,
|
||||
|
|
@ -33,6 +37,54 @@ from litellm.types.tool_management import (
|
|||
|
||||
router = APIRouter()
|
||||
|
||||
TOOL_POLICY_OPTIONS = ToolPolicyOptionsResponse(
|
||||
input_policies=[
|
||||
ToolPolicyOption(
|
||||
value="untrusted",
|
||||
label="Untrusted",
|
||||
description="Tool accepts any input, including data from untrusted tool outputs. Default for newly discovered tools.",
|
||||
),
|
||||
ToolPolicyOption(
|
||||
value="trusted",
|
||||
label="Trusted",
|
||||
description="Tool requires trusted input. Blocked if the conversation contains output from any tool with output_policy=untrusted.",
|
||||
),
|
||||
ToolPolicyOption(
|
||||
value="blocked",
|
||||
label="Blocked",
|
||||
description="Tool is completely prohibited. Any attempt to call it is rejected.",
|
||||
),
|
||||
],
|
||||
output_policies=[
|
||||
ToolPolicyOption(
|
||||
value="untrusted",
|
||||
label="Untrusted",
|
||||
description="Tool output may contain unsafe content (prompt injection, risky code). Downstream tools with input_policy=trusted will be blocked.",
|
||||
),
|
||||
ToolPolicyOption(
|
||||
value="trusted",
|
||||
label="Trusted",
|
||||
description="Tool output is verified safe. Will not trigger trust-chain blocks on downstream tools.",
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/v1/tool/policy/options",
|
||||
tags=["tool management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=ToolPolicyOptionsResponse,
|
||||
)
|
||||
async def get_tool_policy_options(
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Return the available input and output policy options with descriptions.
|
||||
Static data — no DB call.
|
||||
"""
|
||||
return TOOL_POLICY_OPTIONS
|
||||
|
||||
|
||||
@router.get(
|
||||
"/v1/tool/list",
|
||||
|
|
@ -41,14 +93,14 @@ router = APIRouter()
|
|||
response_model=ToolListResponse,
|
||||
)
|
||||
async def list_tools(
|
||||
call_policy: Optional[ToolCallPolicy] = None,
|
||||
input_policy: Optional[ToolInputPolicy] = None,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
List all auto-discovered tools and their call policies.
|
||||
List all auto-discovered tools and their policies.
|
||||
|
||||
Parameters:
|
||||
- call_policy: Optional filter — one of "trusted", "untrusted", "dual_llm", "blocked"
|
||||
- input_policy: Optional filter — one of "trusted", "untrusted", "blocked"
|
||||
"""
|
||||
from litellm.proxy.db.tool_registry_writer import list_tools as db_list_tools
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
|
@ -60,7 +112,7 @@ async def list_tools(
|
|||
|
||||
try:
|
||||
tools = await db_list_tools(
|
||||
prisma_client=prisma_client, call_policy=call_policy
|
||||
prisma_client=prisma_client, input_policy=input_policy
|
||||
)
|
||||
return ToolListResponse(tools=tools, total=len(tools))
|
||||
except Exception as e:
|
||||
|
|
@ -80,9 +132,6 @@ async def get_tool_detail(
|
|||
):
|
||||
"""
|
||||
Get a single tool with its policy overrides (for UI detail view).
|
||||
|
||||
Parameters:
|
||||
- tool_name: The tool name (supports namespaced names with slashes)
|
||||
"""
|
||||
from litellm.proxy.db.tool_registry_writer import get_tool as db_get_tool
|
||||
from litellm.proxy.db.tool_registry_writer import list_overrides_for_tool
|
||||
|
|
@ -269,9 +318,6 @@ async def get_tool(
|
|||
):
|
||||
"""
|
||||
Get details for a single tool.
|
||||
|
||||
Parameters:
|
||||
- tool_name: The tool name (supports namespaced names with slashes)
|
||||
"""
|
||||
from litellm.proxy.db.tool_registry_writer import get_tool as db_get_tool
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
|
@ -311,9 +357,6 @@ async def _resolve_key_hash_to_object_permission_id(
|
|||
op_id = getattr(row, "object_permission_id", None)
|
||||
if op_id:
|
||||
return op_id
|
||||
# Create new object permission and atomically assign to key.
|
||||
# Uses update_many with object_permission_id=None to prevent race conditions:
|
||||
# only one concurrent request wins; the loser cleans up its orphaned row.
|
||||
new_id = str(uuid.uuid4())
|
||||
await prisma_client.db.litellm_objectpermissiontable.create(
|
||||
data={"object_permission_id": new_id, "blocked_tools": []}
|
||||
|
|
@ -323,7 +366,6 @@ async def _resolve_key_hash_to_object_permission_id(
|
|||
data={"object_permission_id": new_id},
|
||||
)
|
||||
if updated_count == 0:
|
||||
# Another request already assigned a permission; clean up orphan
|
||||
await prisma_client.db.litellm_objectpermissiontable.delete(
|
||||
where={"object_permission_id": new_id}
|
||||
)
|
||||
|
|
@ -351,7 +393,6 @@ async def _resolve_team_id_to_object_permission_id(
|
|||
op_id = getattr(row, "object_permission_id", None)
|
||||
if op_id:
|
||||
return op_id
|
||||
# Same atomic pattern as _resolve_key_hash_to_object_permission_id
|
||||
new_id = str(uuid.uuid4())
|
||||
await prisma_client.db.litellm_objectpermissiontable.create(
|
||||
data={"object_permission_id": new_id, "blocked_tools": []}
|
||||
|
|
@ -383,18 +424,14 @@ async def update_tool_policy(
|
|||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Set the call policy for a tool (global) or for a specific team/key (override).
|
||||
Set the input_policy and/or output_policy for a tool (global), or block for a specific team/key (override).
|
||||
|
||||
Parameters:
|
||||
- tool_name: str - The tool to update
|
||||
- call_policy: "trusted" | "untrusted" | "dual_llm" | "blocked"
|
||||
- input_policy: optional - "trusted" | "untrusted" | "blocked"
|
||||
- output_policy: optional - "trusted" | "untrusted"
|
||||
- team_id: optional - if set, create/update override for this team only
|
||||
- key_hash: optional - if set, create/update override for this key only
|
||||
- key_alias: optional - human-readable key alias for UI
|
||||
|
||||
If both team_id and key_hash are omitted, updates the global tool policy.
|
||||
Setting a tool to "blocked" will cause the ToolPolicyGuardrail to reject
|
||||
that tool_call for the relevant scope.
|
||||
"""
|
||||
from litellm.proxy.db.tool_registry_writer import (
|
||||
add_tool_to_object_permission_blocked,
|
||||
|
|
@ -431,7 +468,8 @@ async def update_tool_policy(
|
|||
status_code=404,
|
||||
detail="Key or team not found for the given identifier",
|
||||
)
|
||||
if data.call_policy == "blocked":
|
||||
is_blocking = data.input_policy == "blocked"
|
||||
if is_blocking:
|
||||
ok = await add_tool_to_object_permission_blocked(
|
||||
prisma_client=prisma_client,
|
||||
object_permission_id=op_id,
|
||||
|
|
@ -453,16 +491,25 @@ async def update_tool_policy(
|
|||
await registry.sync_tool_policy_from_db(prisma_client)
|
||||
return ToolPolicyUpdateResponse(
|
||||
tool_name=data.tool_name,
|
||||
call_policy=data.call_policy,
|
||||
input_policy=data.input_policy,
|
||||
output_policy=data.output_policy,
|
||||
updated=True,
|
||||
team_id=data.team_id,
|
||||
key_hash=data.key_hash,
|
||||
)
|
||||
|
||||
if data.input_policy is None and data.output_policy is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="At least one of input_policy or output_policy must be provided",
|
||||
)
|
||||
|
||||
updated = await db_update_tool_policy(
|
||||
prisma_client=prisma_client,
|
||||
tool_name=data.tool_name,
|
||||
call_policy=data.call_policy,
|
||||
updated_by=user_api_key_dict.user_id,
|
||||
input_policy=data.input_policy,
|
||||
output_policy=data.output_policy,
|
||||
)
|
||||
if updated is None:
|
||||
raise HTTPException(
|
||||
|
|
@ -474,7 +521,8 @@ async def update_tool_policy(
|
|||
await registry.sync_tool_policy_from_db(prisma_client)
|
||||
return ToolPolicyUpdateResponse(
|
||||
tool_name=updated.tool_name,
|
||||
call_policy=updated.call_policy,
|
||||
input_policy=updated.input_policy,
|
||||
output_policy=updated.output_policy,
|
||||
updated=True,
|
||||
)
|
||||
except HTTPException:
|
||||
|
|
|
|||
|
|
@ -9,17 +9,23 @@ from pydantic import BaseModel, Field
|
|||
|
||||
ToolCallPolicy = Literal["trusted", "untrusted", "dual_llm", "blocked"]
|
||||
|
||||
ToolInputPolicy = Literal["trusted", "untrusted", "blocked"]
|
||||
ToolOutputPolicy = Literal["trusted", "untrusted"]
|
||||
|
||||
|
||||
class LiteLLM_ToolTableRow(BaseModel):
|
||||
tool_id: str
|
||||
tool_name: str
|
||||
origin: Optional[str] = None
|
||||
call_policy: ToolCallPolicy = "untrusted"
|
||||
input_policy: ToolInputPolicy = "untrusted"
|
||||
output_policy: ToolOutputPolicy = "untrusted"
|
||||
call_count: int = 0
|
||||
assignments: Optional[Dict] = None
|
||||
key_hash: Optional[str] = None
|
||||
team_id: Optional[str] = None
|
||||
key_alias: Optional[str] = None
|
||||
user_agent: Optional[str] = None
|
||||
last_used_at: Optional[datetime] = None
|
||||
created_at: Optional[datetime] = None
|
||||
updated_at: Optional[datetime] = None
|
||||
created_by: Optional[str] = None
|
||||
|
|
@ -33,15 +39,17 @@ class ToolListResponse(BaseModel):
|
|||
|
||||
class ToolPolicyUpdateRequest(BaseModel):
|
||||
tool_name: str
|
||||
call_policy: ToolCallPolicy
|
||||
team_id: Optional[str] = None # if set, create/update override for this team
|
||||
key_hash: Optional[str] = None # if set, create/update override for this key
|
||||
key_alias: Optional[str] = None # human-readable key alias for UI
|
||||
input_policy: Optional[ToolInputPolicy] = None
|
||||
output_policy: Optional[ToolOutputPolicy] = None
|
||||
team_id: Optional[str] = None
|
||||
key_hash: Optional[str] = None
|
||||
key_alias: Optional[str] = None
|
||||
|
||||
|
||||
class ToolPolicyUpdateResponse(BaseModel):
|
||||
tool_name: str
|
||||
call_policy: ToolCallPolicy
|
||||
input_policy: Optional[ToolInputPolicy] = None
|
||||
output_policy: Optional[ToolOutputPolicy] = None
|
||||
updated: bool
|
||||
team_id: Optional[str] = None
|
||||
key_hash: Optional[str] = None
|
||||
|
|
@ -52,12 +60,23 @@ class ToolPolicyOverrideRow(BaseModel):
|
|||
tool_name: str
|
||||
team_id: Optional[str] = None
|
||||
key_hash: Optional[str] = None
|
||||
call_policy: ToolCallPolicy
|
||||
input_policy: ToolInputPolicy = "blocked"
|
||||
key_alias: Optional[str] = None
|
||||
created_at: Optional[datetime] = None
|
||||
updated_at: Optional[datetime] = None
|
||||
|
||||
|
||||
class ToolPolicyOption(BaseModel):
|
||||
value: str
|
||||
label: str
|
||||
description: str
|
||||
|
||||
|
||||
class ToolPolicyOptionsResponse(BaseModel):
|
||||
input_policies: List[ToolPolicyOption]
|
||||
output_policies: List[ToolPolicyOption]
|
||||
|
||||
|
||||
class ToolDetailResponse(BaseModel):
|
||||
tool: LiteLLM_ToolTableRow
|
||||
overrides: List[ToolPolicyOverrideRow] = Field(default_factory=list)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue