mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
refactor: use prisma functions instead of raw sql (safer)
This commit is contained in:
parent
b0439611f6
commit
4b4018b0b2
5 changed files with 190 additions and 147 deletions
|
|
@ -109,6 +109,8 @@ Key files:
|
|||
- `litellm/proxy/auth/` - Authentication logic
|
||||
- `litellm/proxy/management_endpoints/` - Admin API endpoints
|
||||
|
||||
**Database (proxy)**: Use Prisma model methods (`prisma_client.db.<model>.upsert`, `.find_many`, `.find_unique`, etc.), not raw SQL (`execute_raw`/`query_raw`). See COMMON PITFALLS for details.
|
||||
|
||||
## MCP (MODEL CONTEXT PROTOCOL) SUPPORT
|
||||
|
||||
LiteLLM supports MCP for agent workflows:
|
||||
|
|
@ -176,6 +178,7 @@ When opening issues or pull requests, follow these templates:
|
|||
5. **Dependencies**: Keep dependencies minimal and well-justified
|
||||
6. **UI/Backend Contract Mismatch**: When adding a new entity type to the UI, always check whether the backend endpoint accepts a single value or an array. Match the UI control accordingly (single-select vs. multi-select) to avoid silently dropping user selections
|
||||
7. **Missing Tests for New Entity Types**: When adding a new entity type (e.g., in `EntityUsage`, `UsageViewSelect`), always add corresponding tests in the existing test files and update any icon/component mocks
|
||||
8. **Raw SQL in proxy DB code**: Do not use `execute_raw` or `query_raw` for proxy database access. Use Prisma model methods (e.g. `prisma_client.db.litellm_tooltable.upsert()`, `.find_many()`, `.find_unique()`) so behavior stays consistent with the schema, the client stays mockable in tests, and you avoid the pitfalls of hand-written SQL (parameter ordering, type casting, schema drift)
|
||||
|
||||
## HELPFUL RESOURCES
|
||||
|
||||
|
|
|
|||
|
|
@ -107,6 +107,10 @@ LiteLLM is a unified interface for 100+ LLM providers with two main components:
|
|||
- Migration files auto-generated with `prisma migrate dev`
|
||||
- Always test migrations against both PostgreSQL and SQLite
|
||||
|
||||
### Proxy database access
|
||||
- **Do not write raw SQL** for proxy DB operations. Use Prisma model methods instead of `execute_raw` / `query_raw`.
|
||||
- Use the generated client: `prisma_client.db.<model>` (e.g. `litellm_tooltable`, `litellm_usertable`) with `.upsert()`, `.find_many()`, `.find_unique()`, `.update()`, `.update_many()` as appropriate. This avoids schema/client drift, keeps code testable with simple mocks, and matches patterns used in spend logs and other proxy code.
|
||||
|
||||
### Enterprise Features
|
||||
- Enterprise-specific code in `enterprise/` directory
|
||||
- Optional features enabled via environment variables
|
||||
|
|
|
|||
|
|
@ -0,0 +1,2 @@
|
|||
-- This is an empty migration.
|
||||
|
||||
|
|
@ -3,15 +3,11 @@ 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.
|
||||
|
||||
NOTE: Uses raw SQL (query_raw / execute_raw) instead of Prisma model methods
|
||||
because the generated Prisma Python client may not have LiteLLM_ToolTable
|
||||
when running against an older generated schema.
|
||||
"""
|
||||
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING, Dict, List, Optional
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import ToolDiscoveryQueueItem
|
||||
|
|
@ -21,7 +17,30 @@ if TYPE_CHECKING:
|
|||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
|
||||
def _row_to_model(row: dict) -> LiteLLM_ToolTableRow:
|
||||
def _row_to_model(row: Union[dict, Any]) -> LiteLLM_ToolTableRow:
|
||||
"""Convert a Prisma model instance or dict to LiteLLM_ToolTableRow."""
|
||||
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 (
|
||||
"tool_id",
|
||||
"tool_name",
|
||||
"origin",
|
||||
"call_policy",
|
||||
"call_count",
|
||||
"assignments",
|
||||
"key_hash",
|
||||
"team_id",
|
||||
"key_alias",
|
||||
"created_at",
|
||||
"updated_at",
|
||||
"created_by",
|
||||
"updated_by",
|
||||
)
|
||||
}
|
||||
return LiteLLM_ToolTableRow(
|
||||
tool_id=row.get("tool_id", ""),
|
||||
tool_name=row.get("tool_name", ""),
|
||||
|
|
@ -44,7 +63,7 @@ async def batch_upsert_tools(
|
|||
items: List[ToolDiscoveryQueueItem],
|
||||
) -> None:
|
||||
"""
|
||||
Batch-upsert tool registry rows via raw SQL.
|
||||
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.
|
||||
|
|
@ -55,6 +74,8 @@ async def batch_upsert_tools(
|
|||
data = [item for item in items if item.get("tool_name")]
|
||||
if not data:
|
||||
return
|
||||
now = datetime.now(timezone.utc)
|
||||
table = prisma_client.db.litellm_tooltable
|
||||
for item in data:
|
||||
tool_name = item.get("tool_name", "")
|
||||
origin = item.get("origin") or "user_defined"
|
||||
|
|
@ -62,22 +83,26 @@ async def batch_upsert_tools(
|
|||
key_hash = item.get("key_hash")
|
||||
team_id = item.get("team_id")
|
||||
key_alias = item.get("key_alias")
|
||||
now = datetime.now(timezone.utc).isoformat()
|
||||
await prisma_client.db.execute_raw(
|
||||
'INSERT INTO "LiteLLM_ToolTable" '
|
||||
"(tool_id, tool_name, origin, call_policy, call_count, created_by, updated_by, key_hash, team_id, key_alias, created_at, updated_at) "
|
||||
"VALUES ($7, $1, $2, 'untrusted', 1, $3, $3, $4, $5, $6, $8::timestamp, $8::timestamp) "
|
||||
"ON CONFLICT (tool_name) DO UPDATE SET "
|
||||
'call_count = "LiteLLM_ToolTable".call_count + 1, '
|
||||
"updated_at = $8::timestamp",
|
||||
tool_name,
|
||||
origin,
|
||||
created_by,
|
||||
key_hash,
|
||||
team_id,
|
||||
key_alias,
|
||||
str(uuid.uuid4()),
|
||||
now,
|
||||
await table.upsert(
|
||||
where={"tool_name": tool_name},
|
||||
data={
|
||||
"create": {
|
||||
"tool_id": str(uuid.uuid4()),
|
||||
"tool_name": tool_name,
|
||||
"origin": origin,
|
||||
"call_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,
|
||||
},
|
||||
"update": {
|
||||
"call_count": {"increment": 1},
|
||||
"updated_at": now,
|
||||
},
|
||||
},
|
||||
)
|
||||
verbose_proxy_logger.debug(
|
||||
"tool_registry_writer: upserted %d tool(s)", len(data)
|
||||
|
|
@ -94,19 +119,11 @@ async def list_tools(
|
|||
) -> List[LiteLLM_ToolTableRow]:
|
||||
"""Return all tools, optionally filtered by call_policy."""
|
||||
try:
|
||||
if call_policy is not None:
|
||||
rows = await prisma_client.db.query_raw(
|
||||
"SELECT tool_id, tool_name, origin, call_policy, call_count, assignments, "
|
||||
"key_hash, team_id, key_alias, created_at, updated_at, created_by, updated_by "
|
||||
'FROM "LiteLLM_ToolTable" WHERE call_policy = $1 ORDER BY created_at DESC',
|
||||
call_policy,
|
||||
)
|
||||
else:
|
||||
rows = await prisma_client.db.query_raw(
|
||||
"SELECT tool_id, tool_name, origin, call_policy, call_count, assignments, "
|
||||
"key_hash, team_id, key_alias, created_at, updated_at, created_by, updated_by "
|
||||
'FROM "LiteLLM_ToolTable" ORDER BY created_at DESC',
|
||||
)
|
||||
where = {"call_policy": call_policy} if call_policy is not None else {}
|
||||
rows = await prisma_client.db.litellm_tooltable.find_many(
|
||||
where=where,
|
||||
order={"created_at": "desc"},
|
||||
)
|
||||
return [_row_to_model(row) for row in rows]
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("tool_registry_writer list_tools error: %s", e)
|
||||
|
|
@ -119,15 +136,12 @@ async def get_tool(
|
|||
) -> Optional[LiteLLM_ToolTableRow]:
|
||||
"""Return a single tool row by tool_name."""
|
||||
try:
|
||||
rows = await prisma_client.db.query_raw(
|
||||
"SELECT tool_id, tool_name, origin, call_policy, call_count, assignments, "
|
||||
"key_hash, team_id, key_alias, created_at, updated_at, created_by, updated_by "
|
||||
'FROM "LiteLLM_ToolTable" WHERE tool_name = $1',
|
||||
tool_name,
|
||||
row = await prisma_client.db.litellm_tooltable.find_unique(
|
||||
where={"tool_name": tool_name},
|
||||
)
|
||||
if not rows:
|
||||
if row is None:
|
||||
return None
|
||||
return _row_to_model(rows[0])
|
||||
return _row_to_model(row)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("tool_registry_writer get_tool error: %s", e)
|
||||
return None
|
||||
|
|
@ -142,16 +156,25 @@ async def update_tool_policy(
|
|||
"""Update the call_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).isoformat()
|
||||
await prisma_client.db.execute_raw(
|
||||
'INSERT INTO "LiteLLM_ToolTable" (tool_id, tool_name, call_policy, created_by, updated_by, created_at, updated_at) '
|
||||
"VALUES ($4, $1, $2, $3, $3, $5::timestamp, $5::timestamp) "
|
||||
"ON CONFLICT (tool_name) DO UPDATE SET call_policy = $2, updated_by = $3, updated_at = $5::timestamp",
|
||||
tool_name,
|
||||
call_policy,
|
||||
_updated_by,
|
||||
str(uuid.uuid4()),
|
||||
now,
|
||||
now = datetime.now(timezone.utc)
|
||||
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,
|
||||
},
|
||||
},
|
||||
)
|
||||
return await get_tool(prisma_client, tool_name)
|
||||
except Exception as e:
|
||||
|
|
@ -172,12 +195,10 @@ async def get_tools_by_names(
|
|||
if not tool_names:
|
||||
return {}
|
||||
try:
|
||||
placeholders = ", ".join(f"${i+1}" for i in range(len(tool_names)))
|
||||
rows = await prisma_client.db.query_raw(
|
||||
f'SELECT tool_name, call_policy FROM "LiteLLM_ToolTable" WHERE tool_name IN ({placeholders})',
|
||||
*tool_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: row.call_policy for row in rows}
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(
|
||||
"tool_registry_writer get_tools_by_names error: %s", e
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
"""
|
||||
Unit tests for tool_registry_writer.py — uses a mock prisma client
|
||||
that exposes execute_raw / query_raw (matching the actual raw-SQL implementation).
|
||||
that exposes litellm_tooltable.upsert / find_many / find_unique.
|
||||
"""
|
||||
|
||||
import os
|
||||
|
|
@ -12,18 +12,20 @@ import pytest
|
|||
|
||||
sys.path.insert(0, os.path.abspath("../../.."))
|
||||
|
||||
from litellm.proxy.db.tool_registry_writer import (
|
||||
batch_upsert_tools,
|
||||
get_tool,
|
||||
get_tools_by_names,
|
||||
list_tools,
|
||||
update_tool_policy,
|
||||
)
|
||||
from litellm.proxy.db.tool_registry_writer import (batch_upsert_tools,
|
||||
get_tool,
|
||||
get_tools_by_names,
|
||||
list_tools,
|
||||
update_tool_policy)
|
||||
|
||||
|
||||
def _make_prisma(query_rows=None):
|
||||
"""Return a minimal mock prisma_client with execute_raw / query_raw."""
|
||||
default_row = {
|
||||
def _mock_row(**kwargs):
|
||||
"""Build a row-like object with real attributes (no MagicMock) for _row_to_model."""
|
||||
|
||||
class Row:
|
||||
pass
|
||||
|
||||
default = {
|
||||
"tool_id": "uuid-1",
|
||||
"tool_name": "my_tool",
|
||||
"origin": "user_defined",
|
||||
|
|
@ -38,31 +40,53 @@ def _make_prisma(query_rows=None):
|
|||
"created_by": None,
|
||||
"updated_by": None,
|
||||
}
|
||||
rows = query_rows if query_rows is not None else [default_row]
|
||||
default.update(kwargs)
|
||||
row = Row()
|
||||
for k, v in default.items():
|
||||
setattr(row, k, v)
|
||||
return row
|
||||
|
||||
|
||||
def _make_prisma(
|
||||
*,
|
||||
upsert_return=None,
|
||||
find_many_rows=None,
|
||||
find_unique_row=None,
|
||||
):
|
||||
"""Return a mock prisma_client with litellm_tooltable.upsert, find_many, find_unique."""
|
||||
prisma = MagicMock()
|
||||
prisma.db.execute_raw = AsyncMock(return_value=None)
|
||||
prisma.db.query_raw = AsyncMock(return_value=rows)
|
||||
prisma.db.litellm_tooltable = MagicMock()
|
||||
prisma.db.litellm_tooltable.upsert = AsyncMock(return_value=upsert_return)
|
||||
prisma.db.litellm_tooltable.find_many = AsyncMock(
|
||||
return_value=find_many_rows if find_many_rows is not None else []
|
||||
)
|
||||
prisma.db.litellm_tooltable.find_unique = AsyncMock(
|
||||
return_value=find_unique_row
|
||||
)
|
||||
return prisma
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_batch_upsert_tools_calls_execute_raw():
|
||||
async def test_batch_upsert_tools_calls_upsert():
|
||||
prisma = _make_prisma()
|
||||
items = [{"tool_name": "tool_a", "origin": "mcp_server", "created_by": None}]
|
||||
await batch_upsert_tools(prisma, items)
|
||||
prisma.db.execute_raw.assert_awaited_once()
|
||||
call_args = prisma.db.execute_raw.call_args
|
||||
sql = call_args.args[0]
|
||||
assert "LiteLLM_ToolTable" in sql
|
||||
assert "ON CONFLICT" in sql
|
||||
prisma.db.litellm_tooltable.upsert.assert_awaited_once()
|
||||
call_kw = prisma.db.litellm_tooltable.upsert.call_args.kwargs
|
||||
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"]["call_count"] == 1
|
||||
assert call_kw["data"]["update"]["call_count"] == {"increment": 1}
|
||||
assert "updated_at" in call_kw["data"]["update"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_batch_upsert_tools_empty_list():
|
||||
prisma = _make_prisma()
|
||||
await batch_upsert_tools(prisma, [])
|
||||
prisma.db.execute_raw.assert_not_awaited()
|
||||
prisma.db.litellm_tooltable.upsert.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -70,123 +94,112 @@ async def test_batch_upsert_tools_skips_empty_names():
|
|||
prisma = _make_prisma()
|
||||
items = [{"tool_name": "", "origin": None}, {"tool_name": None}] # type: ignore[list-item]
|
||||
await batch_upsert_tools(prisma, items)
|
||||
prisma.db.execute_raw.assert_not_awaited()
|
||||
prisma.db.litellm_tooltable.upsert.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_batch_upsert_multiple_tools_calls_execute_raw_per_tool():
|
||||
async def test_batch_upsert_multiple_tools_calls_upsert_per_tool():
|
||||
prisma = _make_prisma()
|
||||
items = [
|
||||
{"tool_name": "tool_a", "origin": "mcp_server", "created_by": None},
|
||||
{"tool_name": "tool_b", "origin": "user_defined", "created_by": "alice"},
|
||||
]
|
||||
await batch_upsert_tools(prisma, items)
|
||||
assert prisma.db.execute_raw.await_count == 2
|
||||
assert prisma.db.litellm_tooltable.upsert.await_count == 2
|
||||
calls = prisma.db.litellm_tooltable.upsert.call_args_list
|
||||
assert calls[0].kwargs["where"]["tool_name"] == "tool_a"
|
||||
assert calls[1].kwargs["where"]["tool_name"] == "tool_b"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_tools_no_filter():
|
||||
row = {
|
||||
"tool_id": "id1",
|
||||
"tool_name": "tool_a",
|
||||
"origin": "mcp",
|
||||
"call_policy": "untrusted",
|
||||
"call_count": 5,
|
||||
"assignments": {},
|
||||
"key_hash": None,
|
||||
"team_id": None,
|
||||
"key_alias": None,
|
||||
"created_at": datetime.now(timezone.utc),
|
||||
"updated_at": datetime.now(timezone.utc),
|
||||
"created_by": None,
|
||||
"updated_by": None,
|
||||
}
|
||||
prisma = _make_prisma(query_rows=[row])
|
||||
row = _mock_row(
|
||||
tool_id="id1",
|
||||
tool_name="tool_a",
|
||||
origin="mcp",
|
||||
call_policy="untrusted",
|
||||
call_count=5,
|
||||
)
|
||||
prisma = _make_prisma(find_many_rows=[row])
|
||||
result = await list_tools(prisma)
|
||||
assert len(result) == 1
|
||||
assert result[0].tool_name == "tool_a"
|
||||
assert result[0].call_count == 5
|
||||
prisma.db.query_raw.assert_awaited_once()
|
||||
prisma.db.litellm_tooltable.find_many.assert_awaited_once()
|
||||
call_kw = prisma.db.litellm_tooltable.find_many.call_args.kwargs
|
||||
assert call_kw["where"] == {}
|
||||
assert call_kw["order"] == {"created_at": "desc"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_tools_with_policy_filter():
|
||||
row = {
|
||||
"tool_id": "id1",
|
||||
"tool_name": "blocked_tool",
|
||||
"origin": None,
|
||||
"call_policy": "blocked",
|
||||
"call_count": 2,
|
||||
"assignments": None,
|
||||
"key_hash": None,
|
||||
"team_id": None,
|
||||
"key_alias": None,
|
||||
"created_at": datetime.now(timezone.utc),
|
||||
"updated_at": datetime.now(timezone.utc),
|
||||
"created_by": None,
|
||||
"updated_by": None,
|
||||
}
|
||||
prisma = _make_prisma(query_rows=[row])
|
||||
row = _mock_row(
|
||||
tool_id="id1",
|
||||
tool_name="blocked_tool",
|
||||
origin=None,
|
||||
call_policy="blocked",
|
||||
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"
|
||||
call_args = prisma.db.query_raw.call_args
|
||||
sql = call_args.args[0]
|
||||
assert "WHERE call_policy" in sql
|
||||
call_kw = prisma.db.litellm_tooltable.find_many.call_args.kwargs
|
||||
assert call_kw["where"] == {"call_policy": "blocked"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_tool_found():
|
||||
prisma = _make_prisma()
|
||||
row = _mock_row(tool_name="my_tool")
|
||||
prisma = _make_prisma(find_unique_row=row)
|
||||
result = await get_tool(prisma, "my_tool")
|
||||
assert result is not None
|
||||
assert result.tool_name == "my_tool"
|
||||
prisma.db.query_raw.assert_awaited_once()
|
||||
prisma.db.litellm_tooltable.find_unique.assert_awaited_once_with(
|
||||
where={"tool_name": "my_tool"}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_tool_not_found():
|
||||
prisma = _make_prisma(query_rows=[])
|
||||
prisma = _make_prisma(find_unique_row=None)
|
||||
result = await get_tool(prisma, "nonexistent")
|
||||
assert result is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_tool_policy_calls_execute_raw():
|
||||
row = {
|
||||
"tool_id": "uuid-1",
|
||||
"tool_name": "my_tool",
|
||||
"origin": "user_defined",
|
||||
"call_policy": "blocked",
|
||||
"call_count": 1,
|
||||
"assignments": {},
|
||||
"key_hash": None,
|
||||
"team_id": None,
|
||||
"key_alias": None,
|
||||
"created_at": datetime.now(timezone.utc),
|
||||
"updated_at": datetime.now(timezone.utc),
|
||||
"created_by": None,
|
||||
"updated_by": "admin",
|
||||
}
|
||||
prisma = _make_prisma(query_rows=[row])
|
||||
async def test_update_tool_policy_calls_upsert_then_get_tool():
|
||||
row = _mock_row(
|
||||
tool_name="my_tool",
|
||||
call_policy="blocked",
|
||||
updated_by="admin",
|
||||
)
|
||||
prisma = _make_prisma(find_unique_row=row)
|
||||
result = await update_tool_policy(prisma, "my_tool", "blocked", "admin")
|
||||
assert result is not None
|
||||
assert result.call_policy == "blocked"
|
||||
prisma.db.execute_raw.assert_awaited_once()
|
||||
call_args = prisma.db.execute_raw.call_args
|
||||
sql = call_args.args[0]
|
||||
assert "ON CONFLICT" in sql
|
||||
assert "call_policy" in sql
|
||||
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"]["updated_by"] == "admin"
|
||||
prisma.db.litellm_tooltable.find_unique.assert_awaited_with(
|
||||
where={"tool_name": "my_tool"}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_tools_by_names_returns_policy_map():
|
||||
rows = [
|
||||
{"tool_name": "tool_a", "call_policy": "trusted"},
|
||||
{"tool_name": "tool_b", "call_policy": "blocked"},
|
||||
_mock_row(tool_name="tool_a", call_policy="trusted"),
|
||||
_mock_row(tool_name="tool_b", call_policy="blocked"),
|
||||
]
|
||||
prisma = _make_prisma(query_rows=rows)
|
||||
prisma = _make_prisma(find_many_rows=rows)
|
||||
result = await get_tools_by_names(prisma, ["tool_a", "tool_b"])
|
||||
assert result == {"tool_a": "trusted", "tool_b": "blocked"}
|
||||
prisma.db.litellm_tooltable.find_many.assert_awaited_once_with(
|
||||
where={"tool_name": {"in": ["tool_a", "tool_b"]}}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -194,4 +207,4 @@ async def test_get_tools_by_names_empty_list():
|
|||
prisma = _make_prisma()
|
||||
result = await get_tools_by_names(prisma, [])
|
||||
assert result == {}
|
||||
prisma.db.query_raw.assert_not_awaited()
|
||||
prisma.db.litellm_tooltable.find_many.assert_not_awaited()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue