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
This commit is contained in:
Ishaan Jaffer 2026-03-03 20:08:29 -08:00
parent 05603791af
commit 4b88e05ae1
3 changed files with 63 additions and 43 deletions

View file

@ -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])

View file

@ -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])
}

View file

@ -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(