From fc8810a865b08127e6c67904ca004f75cd3343d0 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Tue, 3 Mar 2026 19:53:01 -0800 Subject: [PATCH] fix: tool mgmt endpoints --- litellm/proxy/db/tool_registry_writer.py | 121 ++++++++++++------ .../tool_policy/tool_policy_guardrail.py | 115 +++++++++++++---- .../tool_management_endpoints.py | 104 +++++++++++---- litellm/types/tool_management.py | 33 ++++- 4 files changed, 269 insertions(+), 104 deletions(-) diff --git a/litellm/proxy/db/tool_registry_writer.py b/litellm/proxy/db/tool_registry_writer.py index 5a6512f1374..0eda012d515 100644 --- a/litellm/proxy/db/tool_registry_writer.py +++ b/litellm/proxy/db/tool_registry_writer.py @@ -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 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 2d91febd5e4..368948414e9 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 @@ -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 diff --git a/litellm/proxy/management_endpoints/tool_management_endpoints.py b/litellm/proxy/management_endpoints/tool_management_endpoints.py index 304a973535f..7fdd3475c04 100644 --- a/litellm/proxy/management_endpoints/tool_management_endpoints.py +++ b/litellm/proxy/management_endpoints/tool_management_endpoints.py @@ -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: diff --git a/litellm/types/tool_management.py b/litellm/types/tool_management.py index ef8a488f4a0..1c5e1df9e9a 100644 --- a/litellm/types/tool_management.py +++ b/litellm/types/tool_management.py @@ -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)