Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/db/tool_registry_writer.py: 47%
179 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""
2DB helpers for LiteLLM_ToolTable — the global tool registry.
4Tools are auto-discovered from LLM responses and upserted here.
5Admins use the management endpoints to read and update input_policy / output_policy.
6"""
8import uuid
9from collections.abc import Mapping, Sequence
10from datetime import datetime, timezone
11from typing import TYPE_CHECKING, Final, Protocol
13from pydantic import TypeAdapter
15from litellm._logging import verbose_proxy_logger
16from litellm.proxy._types import ToolDiscoveryQueueItem
17from litellm.proxy.db.exception_handler import call_with_db_reconnect_retry
18from litellm.repositories.object_permission_repository import ObjectPermissionRepository
19from litellm.repositories.prisma_protocols import TableActions
20from litellm.repositories.table_repositories import ToolRepository
21from litellm.types.tool_management import (
22 LiteLLM_ToolTableRow,
23 ToolPolicyOverrideRow,
24)
26if TYPE_CHECKING: 26 ↛ 27line 26 didn't jump to line 27 because the condition on line 26 was never true
27 from prisma import models as prisma_db_models
29 from litellm.proxy.utils import PrismaClient
32class _ModelDumpMethod(Protocol):
33 def __call__(self) -> Mapping: ... 33 ↛ exitline 33 didn't return from function '__call__' because
36_ROW_DICT: Final = TypeAdapter(dict)
39def _tool_table_actions(prisma_client: "PrismaClient") -> "TableActions[prisma_db_models.LiteLLM_ToolTable]":
40 table: Final[TableActions[prisma_db_models.LiteLLM_ToolTable]] = ToolRepository(prisma_client).table
41 return table
44def _object_permission_table_actions(
45 prisma_client: "PrismaClient",
46) -> "TableActions[prisma_db_models.LiteLLM_ObjectPermissionTable]":
47 table: Final[TableActions[prisma_db_models.LiteLLM_ObjectPermissionTable]] = ObjectPermissionRepository(
48 prisma_client
49 ).table
50 return table
53def _row_to_model(row: object) -> LiteLLM_ToolTableRow:
54 """Convert a Prisma model instance or dict to LiteLLM_ToolTableRow."""
55 model_dump: Final[_ModelDumpMethod | None] = getattr(row, "model_dump", None)
56 if callable(model_dump): 56 ↛ 58line 56 didn't jump to line 58 because the condition on line 56 was always true
57 row = model_dump()
58 elif not isinstance(row, dict):
59 row = _ROW_DICT.validate_python(
60 {
61 k: getattr(row, k, None)
62 for k in (
63 "tool_id",
64 "tool_name",
65 "origin",
66 "input_policy",
67 "output_policy",
68 "call_count",
69 "assignments",
70 "key_hash",
71 "team_id",
72 "key_alias",
73 "user_agent",
74 "last_used_at",
75 "created_at",
76 "updated_at",
77 "created_by",
78 "updated_by",
79 )
80 }
81 )
82 return LiteLLM_ToolTableRow(
83 tool_id=row.get("tool_id", ""),
84 tool_name=row.get("tool_name", ""),
85 origin=row.get("origin"),
86 input_policy=row.get("input_policy") or "untrusted",
87 output_policy=row.get("output_policy") or "untrusted",
88 call_count=int(row.get("call_count") or 0),
89 assignments=row.get("assignments"),
90 key_hash=row.get("key_hash"),
91 team_id=row.get("team_id"),
92 key_alias=row.get("key_alias"),
93 user_agent=row.get("user_agent"),
94 last_used_at=row.get("last_used_at"),
95 created_at=row.get("created_at"),
96 updated_at=row.get("updated_at"),
97 created_by=row.get("created_by"),
98 updated_by=row.get("updated_by"),
99 )
102async def batch_upsert_tools(
103 prisma_client: "PrismaClient",
104 items: list[ToolDiscoveryQueueItem],
105) -> None:
106 """
107 Batch-upsert tool registry rows via Prisma.
109 On first insert: sets input_policy/output_policy = "untrusted" (default), call_count = 1.
110 On conflict: increments call_count; preserves existing policies.
111 """
112 if not items:
113 return
114 try:
115 data: Final = [item for item in items if item.get("tool_name")]
116 if not data:
117 return
118 now: Final = datetime.now(timezone.utc)
119 table: Final = _tool_table_actions(prisma_client)
120 for item in data:
121 tool_name = item.get("tool_name", "")
122 origin = item.get("origin") or "user_defined"
123 created_by = item.get("created_by") or "system"
124 key_hash = item.get("key_hash")
125 team_id = item.get("team_id")
126 key_alias = item.get("key_alias")
127 user_agent = item.get("user_agent")
128 await table.upsert(
129 where={"tool_name": tool_name},
130 data={
131 "create": {
132 "tool_id": str(uuid.uuid4()),
133 "tool_name": tool_name,
134 "origin": origin,
135 "input_policy": "untrusted",
136 "output_policy": "untrusted",
137 "call_count": 1,
138 "created_by": created_by,
139 "updated_by": created_by,
140 "key_hash": key_hash,
141 "team_id": team_id,
142 "key_alias": key_alias,
143 "user_agent": user_agent,
144 "last_used_at": now,
145 },
146 "update": {
147 "call_count": {"increment": 1},
148 "updated_at": now,
149 "last_used_at": now,
150 },
151 },
152 )
153 verbose_proxy_logger.debug("tool_registry_writer: upserted %d tool(s)", len(data))
154 except Exception as e:
155 verbose_proxy_logger.error("tool_registry_writer batch_upsert_tools error: %s", e)
158async def list_tools(
159 prisma_client: "PrismaClient",
160 input_policy: str | None = None,
161) -> list[LiteLLM_ToolTableRow]:
162 """Return all tools, optionally filtered by input_policy."""
163 try:
164 where: Final[Mapping[str, str]] = {"input_policy": input_policy} if input_policy is not None else {}
165 rows: Final = await _tool_table_actions(prisma_client).find_many(
166 where=where,
167 order={"created_at": "desc"},
168 )
169 return [_row_to_model(row) for row in rows]
170 except Exception as e:
171 verbose_proxy_logger.error("tool_registry_writer list_tools error: %s", e)
172 return []
175async def get_tool(
176 prisma_client: "PrismaClient",
177 tool_name: str,
178) -> LiteLLM_ToolTableRow | None:
179 """Return a single tool row by tool_name."""
180 try:
181 row: Final = await _tool_table_actions(prisma_client).find_unique(
182 where={"tool_name": tool_name},
183 )
184 if row is None:
185 return None
186 return _row_to_model(row)
187 except Exception as e:
188 verbose_proxy_logger.error("tool_registry_writer get_tool error: %s", e)
189 return None
192async def update_tool_policy(
193 prisma_client: "PrismaClient",
194 tool_name: str,
195 updated_by: str | None,
196 input_policy: str | None = None,
197 output_policy: str | None = None,
198) -> LiteLLM_ToolTableRow | None:
199 """Update input_policy and/or output_policy for a tool. Upserts the row if it does not exist yet."""
200 try:
201 _updated_by: Final = updated_by or "system"
202 now: Final = datetime.now(timezone.utc)
204 create_data: Final[Mapping[str, str | datetime]] = {
205 "tool_id": str(uuid.uuid4()),
206 "tool_name": tool_name,
207 "input_policy": input_policy or "untrusted",
208 "output_policy": output_policy or "untrusted",
209 "created_by": _updated_by,
210 "updated_by": _updated_by,
211 "created_at": now,
212 "updated_at": now,
213 }
214 update_data: Final[Mapping[str, str | datetime]] = {
215 key: value
216 for key, value in (
217 ("updated_by", _updated_by),
218 ("updated_at", now),
219 ("input_policy", input_policy),
220 ("output_policy", output_policy),
221 )
222 if value is not None
223 }
225 await _tool_table_actions(prisma_client).upsert(
226 where={"tool_name": tool_name},
227 data={
228 "create": create_data,
229 "update": update_data,
230 },
231 )
232 return await get_tool(prisma_client, tool_name)
233 except Exception as e:
234 verbose_proxy_logger.error("tool_registry_writer update_tool_policy error: %s", e)
235 return None
238async def get_tools_by_names(
239 prisma_client: "PrismaClient",
240 tool_names: list[str],
241) -> dict[str, tuple[str, str]]:
242 """
243 Return a {tool_name: (input_policy, output_policy)} map for the given tool names.
244 """
245 if not tool_names:
246 return {}
247 try:
248 rows: Final = await _tool_table_actions(prisma_client).find_many(
249 where={"tool_name": {"in": tool_names}},
250 )
251 return {
252 row.tool_name: (
253 getattr(row, "input_policy", "untrusted") or "untrusted",
254 getattr(row, "output_policy", "untrusted") or "untrusted",
255 )
256 for row in rows
257 }
258 except Exception as e:
259 verbose_proxy_logger.error("tool_registry_writer get_tools_by_names error: %s", e)
260 return {}
263async def list_overrides_for_tool(
264 prisma_client: "PrismaClient",
265 tool_name: str,
266) -> list[ToolPolicyOverrideRow]:
267 """
268 Return override-like rows for a tool by finding object permissions that have
269 this tool in blocked_tools, then resolving each permission to key/team scope for display.
270 """
271 out: Final[list[ToolPolicyOverrideRow]] = []
272 try:
273 perms: Final = await _object_permission_table_actions(prisma_client).find_many(
274 where={"blocked_tools": {"has": tool_name}},
275 include={
276 "verification_tokens": True,
277 "teams": True,
278 },
279 )
280 for perm in perms: 280 ↛ 281line 280 didn't jump to line 281 because the loop on line 280 never started
281 op_id = getattr(perm, "object_permission_id", None) or ""
282 tokens = getattr(perm, "verification_tokens", []) or []
283 teams = getattr(perm, "teams", []) or []
284 for t in tokens:
285 out.append(
286 ToolPolicyOverrideRow(
287 override_id=op_id,
288 tool_name=tool_name,
289 team_id=None,
290 key_hash=getattr(t, "token", None),
291 input_policy="blocked",
292 key_alias=getattr(t, "key_alias", None),
293 created_at=None,
294 updated_at=None,
295 )
296 )
297 for team in teams:
298 out.append(
299 ToolPolicyOverrideRow(
300 override_id=op_id,
301 tool_name=tool_name,
302 team_id=getattr(team, "team_id", None),
303 key_hash=None,
304 input_policy="blocked",
305 key_alias=getattr(team, "team_alias", None),
306 created_at=None,
307 updated_at=None,
308 )
309 )
310 return out
311 except Exception as e:
312 verbose_proxy_logger.error("tool_registry_writer list_overrides_for_tool error: %s", e)
313 return []
316class ToolPolicyRegistry:
317 """
318 In-memory registry of tool policies synced from DB.
319 Hot path uses get_effective_policies only — no DB, no cache.
320 """
322 def __init__(self) -> None:
323 self._tool_input_policies: dict[str, str] = {}
324 self._tool_output_policies: dict[str, str] = {}
325 self._blocked_tools_by_op_id: dict[str, list[str]] = {}
326 self._initialized: bool = False
328 def is_initialized(self) -> bool:
329 return self._initialized
331 async def sync_tool_policy_from_db(self, prisma_client: "PrismaClient") -> None:
332 """Load all tool policies and object-permission blocked_tools from DB."""
333 try:
334 tools: Final = await call_with_db_reconnect_retry(
335 prisma_client,
336 lambda: _tool_table_actions(prisma_client).find_many(),
337 reason="sync_tool_policy_from_db_tools_lookup_failure",
338 )
339 self._tool_input_policies = {
340 row.tool_name: getattr(row, "input_policy", "untrusted") or "untrusted" for row in tools
341 }
342 self._tool_output_policies = {
343 row.tool_name: getattr(row, "output_policy", "untrusted") or "untrusted" for row in tools
344 }
346 perms: Final = await call_with_db_reconnect_retry(
347 prisma_client,
348 lambda: _object_permission_table_actions(prisma_client).find_many(),
349 reason="sync_tool_policy_from_db_perms_lookup_failure",
350 )
351 self._blocked_tools_by_op_id = {}
352 for row in perms:
353 op_id = getattr(row, "object_permission_id", None)
354 blocked: Sequence[str] = getattr(row, "blocked_tools", None) or []
355 if op_id: 355 ↛ 352line 355 didn't jump to line 352 because the condition on line 355 was always true
356 self._blocked_tools_by_op_id[op_id] = list(blocked)
358 self._initialized = True
359 verbose_proxy_logger.info(
360 "ToolPolicyRegistry: synced %d tool policies and %d object permissions from DB",
361 len(self._tool_input_policies),
362 len(self._blocked_tools_by_op_id),
363 )
364 except Exception as e:
365 verbose_proxy_logger.exception("ToolPolicyRegistry sync_tool_policy_from_db error: %s", e)
366 raise
368 def get_input_policy(self, tool_name: str) -> str:
369 return self._tool_input_policies.get(tool_name, "untrusted")
371 def get_output_policy(self, tool_name: str) -> str:
372 return self._tool_output_policies.get(tool_name, "untrusted")
374 def get_effective_policies(
375 self,
376 tool_names: list[str],
377 object_permission_id: str | None = None,
378 team_object_permission_id: str | None = None,
379 ) -> dict[str, str]:
380 """
381 Return effective input_policy per tool from in-memory state.
382 If tool is in key or team blocked_tools -> "blocked", else global input_policy or "untrusted".
383 """
384 if not tool_names:
385 return {}
386 blocked: Final[frozenset[str]] = frozenset(
387 tool
388 for op_id in (object_permission_id, team_object_permission_id)
389 if op_id and op_id.strip()
390 for tool in self._blocked_tools_by_op_id.get(op_id.strip(), [])
391 )
392 result: Final[dict[str, str]] = {}
393 for name in tool_names:
394 if name in blocked:
395 result[name] = "blocked"
396 else:
397 result[name] = self._tool_input_policies.get(name, "untrusted")
398 return result
401_tool_policy_registry: ToolPolicyRegistry | None = None
404def get_tool_policy_registry() -> ToolPolicyRegistry:
405 """Return the global ToolPolicyRegistry singleton."""
406 global _tool_policy_registry
407 if _tool_policy_registry is None:
408 _tool_policy_registry = ToolPolicyRegistry()
409 return _tool_policy_registry
412async def add_tool_to_object_permission_blocked(
413 prisma_client: "PrismaClient",
414 object_permission_id: str,
415 tool_name: str,
416) -> bool:
417 """Add tool_name to the permission's blocked_tools if not already present."""
418 if not object_permission_id or not tool_name:
419 return False
420 try:
421 row: Final = await _object_permission_table_actions(prisma_client).find_unique(
422 where={"object_permission_id": object_permission_id},
423 )
424 if row is None:
425 return False
426 current: Final[Sequence[str]] = getattr(row, "blocked_tools", []) or []
427 if tool_name in current:
428 return True
429 await _object_permission_table_actions(prisma_client).update(
430 where={"object_permission_id": object_permission_id},
431 data={"blocked_tools": [*current, tool_name]},
432 )
433 return True
434 except Exception as e:
435 verbose_proxy_logger.error("tool_registry_writer add_tool_to_object_permission_blocked error: %s", e)
436 return False
439async def remove_tool_from_object_permission_blocked(
440 prisma_client: "PrismaClient",
441 object_permission_id: str,
442 tool_name: str,
443) -> bool:
444 """Remove tool_name from the permission's blocked_tools. Returns False if tool was not in list."""
445 if not object_permission_id or not tool_name: 445 ↛ 446line 445 didn't jump to line 446 because the condition on line 445 was never true
446 return False
447 try:
448 row: Final = await _object_permission_table_actions(prisma_client).find_unique(
449 where={"object_permission_id": object_permission_id},
450 )
451 if row is None: 451 ↛ 452line 451 didn't jump to line 452 because the condition on line 451 was never true
452 return False
453 current: Final[Sequence[str]] = getattr(row, "blocked_tools", []) or []
454 if tool_name not in current: 454 ↛ 456line 454 didn't jump to line 456 because the condition on line 454 was always true
455 return False
456 await _object_permission_table_actions(prisma_client).update(
457 where={"object_permission_id": object_permission_id},
458 data={"blocked_tools": [t for t in current if t != tool_name]},
459 )
460 return True
461 except Exception as e:
462 verbose_proxy_logger.error(
463 "tool_registry_writer remove_tool_from_object_permission_blocked error: %s",
464 e,
465 )
466 return False