From 4b88e05ae18bc083db99f7003f9904c7508e8cb1 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Tue, 3 Mar 2026 20:08:29 -0800 Subject: [PATCH] fix: sync root schema.prisma and fix test_tool_registry_writer for input/output policy - Migrate root schema.prisma LiteLLM_ToolTable from call_policy to input_policy/output_policy, add missing user_agent and last_used_at columns (now consistent with litellm/proxy/schema.prisma and litellm-proxy-extras) - Fix SpendLogToolIndex comment across all three schema files - Fix all call_policy references in test_tool_registry_writer.py: swapped update_tool_policy arguments, wrong get_tools_by_names return type assertions, _mock_tool_row setting call_policy instead of input_policy Addresses Greptile review feedback on PR #22732. Made-with: Cursor --- .../litellm_proxy_extras/schema.prisma | 2 +- schema.prisma | 36 +++++----- .../proxy/db/test_tool_registry_writer.py | 68 ++++++++++++------- 3 files changed, 63 insertions(+), 43 deletions(-) diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index b6a2e37ae7c..ff96e99c148 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -924,7 +924,7 @@ model LiteLLM_SpendLogGuardrailIndex { // 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 input_policy etc. + tool_name String // matches LiteLLM_ToolTable.tool_name; join for input_policy/output_policy etc. start_time DateTime @@id([request_id, tool_name]) diff --git a/schema.prisma b/schema.prisma index cd4f9a4d247..d425f2e6ce3 100644 --- a/schema.prisma +++ b/schema.prisma @@ -924,7 +924,7 @@ model LiteLLM_SpendLogGuardrailIndex { // 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. + tool_name String // matches LiteLLM_ToolTable.tool_name; join for input_policy/output_policy etc. start_time DateTime @@id([request_id, tool_name]) @@ -1068,23 +1068,27 @@ model LiteLLM_PolicyAttachmentTable { updated_by String? } -// Global tool registry - auto-discovered from LLM responses; admins set call_policy here +// Global tool registry - auto-discovered from LLM responses; admins set input/output policies here model LiteLLM_ToolTable { - tool_id String @id @default(uuid()) - tool_name String @unique // e.g. "huggingface_remote-mcp__dynamic_space" - origin String? // MCP server name or "user_defined" - call_policy String @default("untrusted") // "trusted" | "untrusted" | "dual_llm" | "blocked" - call_count Int @default(0) // cumulative number of times this tool was seen - assignments Json? @default("{}") - key_hash String? // hash of the virtual key that first called this tool - team_id String? // team that first called this tool - key_alias String? // human-readable alias of the virtual key - created_at DateTime @default(now()) - created_by String? - updated_at DateTime @default(now()) @updatedAt - updated_by String? + tool_id String @id @default(uuid()) + tool_name String @unique // e.g. "huggingface_remote-mcp__dynamic_space" + origin String? // MCP server name or "user_defined" + input_policy String @default("untrusted") // "trusted" | "untrusted" | "blocked" + output_policy String @default("untrusted") // "trusted" | "untrusted" + call_count Int @default(0) // cumulative number of times this tool was seen + assignments Json? @default("{}") + key_hash String? // hash of the virtual key that first called this tool + team_id String? // team that first called this tool + key_alias String? // human-readable alias of the virtual key + user_agent String? // user-agent of the first request that discovered this tool + last_used_at DateTime? // timestamp of the most recent call + created_at DateTime @default(now()) + created_by String? + updated_at DateTime @default(now()) @updatedAt + updated_by String? - @@index([call_policy]) + @@index([input_policy]) + @@index([output_policy]) @@index([team_id]) } diff --git a/tests/test_litellm/proxy/db/test_tool_registry_writer.py b/tests/test_litellm/proxy/db/test_tool_registry_writer.py index cd4eb55e2e1..1b1ee7afcba 100644 --- a/tests/test_litellm/proxy/db/test_tool_registry_writer.py +++ b/tests/test_litellm/proxy/db/test_tool_registry_writer.py @@ -12,13 +12,15 @@ import pytest sys.path.insert(0, os.path.abspath("../../..")) -from litellm.proxy.db.tool_registry_writer import (ToolPolicyRegistry, - batch_upsert_tools, - get_tool, - get_tool_policy_registry, - get_tools_by_names, - list_tools, - update_tool_policy) +from litellm.proxy.db.tool_registry_writer import ( + ToolPolicyRegistry, + batch_upsert_tools, + get_tool, + get_tool_policy_registry, + get_tools_by_names, + list_tools, + update_tool_policy, +) def _mock_row(**kwargs): @@ -31,7 +33,8 @@ def _mock_row(**kwargs): "tool_id": "uuid-1", "tool_name": "my_tool", "origin": "user_defined", - "call_policy": "untrusted", + "input_policy": "untrusted", + "output_policy": "untrusted", "call_count": 1, "assignments": {}, "key_hash": None, @@ -78,7 +81,8 @@ async def test_batch_upsert_tools_calls_upsert(): assert call_kw["where"] == {"tool_name": "tool_a"} assert call_kw["data"]["create"]["tool_name"] == "tool_a" assert call_kw["data"]["create"]["origin"] == "mcp_server" - assert call_kw["data"]["create"]["call_policy"] == "untrusted" + assert call_kw["data"]["create"]["input_policy"] == "untrusted" + assert call_kw["data"]["create"]["output_policy"] == "untrusted" assert call_kw["data"]["create"]["call_count"] == 1 assert call_kw["data"]["update"]["call_count"] == {"increment": 1} assert "updated_at" in call_kw["data"]["update"] @@ -119,7 +123,8 @@ async def test_list_tools_no_filter(): tool_id="id1", tool_name="tool_a", origin="mcp", - call_policy="untrusted", + input_policy="untrusted", + output_policy="untrusted", call_count=5, ) prisma = _make_prisma(find_many_rows=[row]) @@ -134,20 +139,21 @@ async def test_list_tools_no_filter(): @pytest.mark.asyncio -async def test_list_tools_with_policy_filter(): +async def test_list_tools_with_input_policy_filter(): row = _mock_row( tool_id="id1", tool_name="blocked_tool", origin=None, - call_policy="blocked", + input_policy="blocked", + output_policy="untrusted", call_count=2, assignments=None, ) prisma = _make_prisma(find_many_rows=[row]) - result = await list_tools(prisma, call_policy="blocked") - assert result[0].call_policy == "blocked" + result = await list_tools(prisma, input_policy="blocked") + assert result[0].input_policy == "blocked" call_kw = prisma.db.litellm_tooltable.find_many.call_args.kwargs - assert call_kw["where"] == {"call_policy": "blocked"} + assert call_kw["where"] == {"input_policy": "blocked"} @pytest.mark.asyncio @@ -173,17 +179,20 @@ async def test_get_tool_not_found(): async def test_update_tool_policy_calls_upsert_then_get_tool(): row = _mock_row( tool_name="my_tool", - call_policy="blocked", + input_policy="blocked", + output_policy="untrusted", updated_by="admin", ) prisma = _make_prisma(find_unique_row=row) - result = await update_tool_policy(prisma, "my_tool", updated_by="admin", input_policy="blocked") + result = await update_tool_policy( + prisma, "my_tool", updated_by="admin", input_policy="blocked" + ) assert result is not None - assert result.call_policy == "blocked" + assert result.input_policy == "blocked" prisma.db.litellm_tooltable.upsert.assert_awaited_once() call_kw = prisma.db.litellm_tooltable.upsert.call_args.kwargs assert call_kw["where"] == {"tool_name": "my_tool"} - assert call_kw["data"]["update"]["call_policy"] == "blocked" + assert call_kw["data"]["update"]["input_policy"] == "blocked" assert call_kw["data"]["update"]["updated_by"] == "admin" prisma.db.litellm_tooltable.find_unique.assert_awaited_with( where={"tool_name": "my_tool"} @@ -193,12 +202,15 @@ async def test_update_tool_policy_calls_upsert_then_get_tool(): @pytest.mark.asyncio async def test_get_tools_by_names_returns_policy_map(): rows = [ - _mock_row(tool_name="tool_a", call_policy="trusted"), - _mock_row(tool_name="tool_b", call_policy="blocked"), + _mock_row(tool_name="tool_a", input_policy="trusted", output_policy="untrusted"), + _mock_row(tool_name="tool_b", input_policy="blocked", output_policy="untrusted"), ] prisma = _make_prisma(find_many_rows=rows) result = await get_tools_by_names(prisma, ["tool_a", "tool_b"]) - assert result == {"tool_a": ("trusted", "untrusted"), "tool_b": ("blocked", "untrusted")} + assert result == { + "tool_a": ("trusted", "untrusted"), + "tool_b": ("blocked", "untrusted"), + } prisma.db.litellm_tooltable.find_many.assert_awaited_once_with( where={"tool_name": {"in": ["tool_a", "tool_b"]}} ) @@ -215,7 +227,11 @@ async def test_get_tools_by_names_empty_list(): # --- ToolPolicyRegistry --- -def _mock_tool_row(tool_name: str, input_policy: str = "untrusted", output_policy: str = "untrusted"): +def _mock_tool_row( + tool_name: str, + input_policy: str = "untrusted", + output_policy: str = "untrusted", +): row = MagicMock() row.tool_name = tool_name row.input_policy = input_policy @@ -236,9 +252,9 @@ async def test_tool_policy_registry_sync_and_get_effective_policies(): prisma = MagicMock() prisma.db.litellm_tooltable.find_many = AsyncMock( return_value=[ - _mock_tool_row("tool_a", "trusted"), - _mock_tool_row("tool_b", "blocked"), - _mock_tool_row("tool_c", "untrusted"), + _mock_tool_row("tool_a", input_policy="trusted"), + _mock_tool_row("tool_b", input_policy="blocked"), + _mock_tool_row("tool_c", input_policy="untrusted"), ] ) prisma.db.litellm_objectpermissiontable.find_many = AsyncMock(