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

1""" 

2DB helpers for LiteLLM_ToolTable — the global tool registry. 

3 

4Tools are auto-discovered from LLM responses and upserted here. 

5Admins use the management endpoints to read and update input_policy / output_policy. 

6""" 

7 

8import uuid 

9from collections.abc import Mapping, Sequence 

10from datetime import datetime, timezone 

11from typing import TYPE_CHECKING, Final, Protocol 

12 

13from pydantic import TypeAdapter 

14 

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) 

25 

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 

28 

29 from litellm.proxy.utils import PrismaClient 

30 

31 

32class _ModelDumpMethod(Protocol): 

33 def __call__(self) -> Mapping: ... 33 ↛ exitline 33 didn't return from function '__call__' because

34 

35 

36_ROW_DICT: Final = TypeAdapter(dict) 

37 

38 

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 

42 

43 

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 

51 

52 

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 ) 

100 

101 

102async def batch_upsert_tools( 

103 prisma_client: "PrismaClient", 

104 items: list[ToolDiscoveryQueueItem], 

105) -> None: 

106 """ 

107 Batch-upsert tool registry rows via Prisma. 

108 

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) 

156 

157 

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 [] 

173 

174 

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 

190 

191 

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) 

203 

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 } 

224 

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 

236 

237 

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 {} 

261 

262 

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 [] 

314 

315 

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 """ 

321 

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 

327 

328 def is_initialized(self) -> bool: 

329 return self._initialized 

330 

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 } 

345 

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) 

357 

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 

367 

368 def get_input_policy(self, tool_name: str) -> str: 

369 return self._tool_input_policies.get(tool_name, "untrusted") 

370 

371 def get_output_policy(self, tool_name: str) -> str: 

372 return self._tool_output_policies.get(tool_name, "untrusted") 

373 

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 

399 

400 

401_tool_policy_registry: ToolPolicyRegistry | None = None 

402 

403 

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 

410 

411 

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 

437 

438 

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