Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/management_endpoints/tool_management_endpoints.py: 56%

330 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1""" 

2TOOL POLICY MANAGEMENT 

3 

4All /tool management endpoints 

5 

6GET /v1/tool/list - List all discovered tools and their policies 

7GET /v1/tool/policy/options - List available input/output policy options with descriptions 

8GET /v1/tool/{tool_name} - Get a single tool's details 

9POST /v1/tool/policy - Update the input_policy / output_policy for a tool 

10""" 

11 

12import uuid 

13from collections.abc import Mapping, Sequence 

14from datetime import datetime, timedelta, timezone 

15from typing import TYPE_CHECKING, Annotated, Final, Protocol, TypeAlias, TypeVar, overload 

16 

17from fastapi import APIRouter, Depends, HTTPException, Query 

18from pydantic import BaseModel, Field, TypeAdapter 

19 

20if TYPE_CHECKING: 20 ↛ 21line 20 didn't jump to line 21 because the condition on line 20 was never true

21 from prisma.models import LiteLLM_DailyToolSpend as PrismaDailyToolSpendRow 

22 from prisma.models import LiteLLM_ObjectPermissionTable as PrismaObjectPermissionRow 

23 from prisma.models import LiteLLM_SpendLogs as PrismaSpendLogRow 

24 from prisma.models import LiteLLM_SpendLogToolIndex as PrismaSpendLogToolIndexRow 

25 from prisma.models import LiteLLM_TeamTable as PrismaTeamRow 

26 from prisma.models import LiteLLM_VerificationToken as PrismaVerificationTokenRow 

27 

28 from litellm.proxy.utils import PrismaClient 

29 

30from litellm._logging import verbose_proxy_logger 

31from litellm.constants import TOOL_SPEND_TOP_TOOLS 

32from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth 

33from litellm.proxy.auth.user_api_key_auth import user_api_key_auth 

34from litellm.repositories.object_permission_repository import ObjectPermissionRepository 

35from litellm.repositories.table_repositories import ( 

36 DailyToolSpendRepository, 

37 SpendLogsRepository, 

38 SpendLogToolIndexRepository, 

39) 

40from litellm.repositories.team_repository import TeamRepository 

41from litellm.repositories.verification_token_repository import ( 

42 VerificationTokenRepository, 

43) 

44from litellm.types.tool_management import ( 

45 LiteLLM_ToolTableRow, 

46 ToolDetailResponse, 

47 ToolInputPolicy, 

48 ToolListResponse, 

49 ToolPolicyOption, 

50 ToolPolicyOptionsResponse, 

51 ToolPolicyUpdateRequest, 

52 ToolPolicyUpdateResponse, 

53 ToolSpendDailyEntry, 

54 ToolSpendEntry, 

55 ToolSpendResponse, 

56 ToolUsageLogEntry, 

57 ToolUsageLogsResponse, 

58) 

59 

60_RowT_co: Final = TypeVar("_RowT_co", covariant=True) 

61 

62if TYPE_CHECKING: 62 ↛ 64line 62 didn't jump to line 64 because the condition on line 62 was never true

63 

64 class _TableOps(Protocol[_RowT_co]): 

65 async def find_many( 

66 self, 

67 where: Mapping[str, object] | None = None, 

68 order: Mapping[str, object] | Sequence[Mapping[str, object]] | None = None, 

69 skip: int | None = None, 

70 take: int | None = None, 

71 ) -> Sequence[_RowT_co]: ... 

72 

73 async def find_unique(self, where: Mapping[str, object]) -> _RowT_co | None: ... 

74 

75 async def count(self, where: Mapping[str, object] | None = None) -> int: ... 

76 

77 async def create(self, data: Mapping[str, object]) -> _RowT_co: ... 

78 

79 async def update_many( 

80 self, 

81 where: Mapping[str, object], 

82 data: Mapping[str, object], 

83 ) -> int: ... 

84 

85 async def delete(self, where: Mapping[str, object]) -> _RowT_co | None: ... 

86 

87 async def group_by( 

88 self, 

89 by: Sequence[str], 

90 sum: Mapping[str, bool] | None = None, 

91 where: Mapping[str, object] | None = None, 

92 order: Mapping[str, object] | None = None, 

93 take: int | None = None, 

94 ) -> Sequence[Mapping[str, object]]: ... 

95 

96 class _SpendLogRow(Protocol): 

97 @property 

98 def messages(self) -> object: ... 

99 @property 

100 def proxy_server_request(self) -> str | Mapping[str, object] | None: ... 

101 

102 

103@overload 

104def _typed_table(repo: DailyToolSpendRepository) -> "_TableOps[PrismaDailyToolSpendRow]": ... 104 ↛ exitline 104 didn't return from function '_typed_table' because

105@overload 

106def _typed_table(repo: SpendLogToolIndexRepository) -> "_TableOps[PrismaSpendLogToolIndexRow]": ... 106 ↛ exitline 106 didn't return from function '_typed_table' because

107@overload 

108def _typed_table(repo: SpendLogsRepository) -> "_TableOps[PrismaSpendLogRow]": ... 108 ↛ exitline 108 didn't return from function '_typed_table' because

109@overload 

110def _typed_table(repo: VerificationTokenRepository) -> "_TableOps[PrismaVerificationTokenRow]": ... 110 ↛ exitline 110 didn't return from function '_typed_table' because

111@overload 

112def _typed_table(repo: TeamRepository) -> "_TableOps[PrismaTeamRow]": ... 112 ↛ exitline 112 didn't return from function '_typed_table' because

113@overload 

114def _typed_table(repo: ObjectPermissionRepository) -> "_TableOps[PrismaObjectPermissionRow]": ... 114 ↛ exitline 114 didn't return from function '_typed_table' because

115def _typed_table( 

116 repo: DailyToolSpendRepository 

117 | SpendLogToolIndexRepository 

118 | SpendLogsRepository 

119 | VerificationTokenRepository 

120 | TeamRepository 

121 | ObjectPermissionRepository, 

122) -> object: 

123 return repo.table 

124 

125 

126router: Final = APIRouter() 

127 

128TOOL_POLICY_OPTIONS: Final = ToolPolicyOptionsResponse( 

129 input_policies=[ 

130 ToolPolicyOption( 

131 value="untrusted", 

132 label="Untrusted", 

133 description="Tool accepts any input, including data from untrusted tool outputs. Default for newly discovered tools.", 

134 ), 

135 ToolPolicyOption( 

136 value="trusted", 

137 label="Trusted", 

138 description="Tool requires trusted input. Blocked if the conversation contains output from any tool with output_policy=untrusted.", 

139 ), 

140 ToolPolicyOption( 

141 value="blocked", 

142 label="Blocked", 

143 description="Tool is completely prohibited. Any attempt to call it is rejected.", 

144 ), 

145 ], 

146 output_policies=[ 

147 ToolPolicyOption( 

148 value="untrusted", 

149 label="Untrusted", 

150 description="Tool output may contain unsafe content (prompt injection, risky code). Downstream tools with input_policy=trusted will be blocked.", 

151 ), 

152 ToolPolicyOption( 

153 value="trusted", 

154 label="Trusted", 

155 description="Tool output is verified safe. Will not trigger trust-chain blocks on downstream tools.", 

156 ), 

157 ], 

158) 

159 

160 

161@router.get( 

162 "/v1/tool/policy/options", 

163 tags=["tool management"], 

164 dependencies=[Depends(user_api_key_auth)], 

165 response_model=ToolPolicyOptionsResponse, 

166) 

167async def get_tool_policy_options( 

168 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), 

169): 

170 """ 

171 Return the available input and output policy options with descriptions. 

172 Static data — no DB call. 

173 """ 

174 return TOOL_POLICY_OPTIONS 

175 

176 

177@router.get( 

178 "/v1/tool/list", 

179 tags=["tool management"], 

180 dependencies=[Depends(user_api_key_auth)], 

181 response_model=ToolListResponse, 

182) 

183async def list_tools( 

184 input_policy: ToolInputPolicy | None = None, 

185 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), 

186): 

187 """ 

188 List all auto-discovered tools and their policies. 

189 

190 Parameters: 

191 - input_policy: Optional filter — one of "trusted", "untrusted", "blocked" 

192 """ 

193 from litellm.proxy.db.tool_registry_writer import list_tools as db_list_tools 

194 from litellm.proxy.proxy_server import prisma_client 

195 

196 if prisma_client is None: 196 ↛ 197line 196 didn't jump to line 197 because the condition on line 196 was never true

197 raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value) 

198 

199 try: 

200 tools: Final = await db_list_tools(prisma_client=prisma_client, input_policy=input_policy) 

201 return ToolListResponse(tools=tools, total=len(tools)) 

202 except Exception as e: 

203 verbose_proxy_logger.exception("Error listing tools: %s", e) 

204 raise HTTPException(status_code=500, detail=str(e)) 

205 

206 

207def _parse_day_start(value: str | None) -> datetime | None: 

208 if not value: 

209 return None 

210 try: 

211 return datetime.strptime(value.strip(), "%Y-%m-%d").replace(tzinfo=timezone.utc) 

212 except ValueError: 

213 raise HTTPException( 

214 status_code=400, 

215 detail=f"Invalid date format: {value}. Expected: 'YYYY-MM-DD'", 

216 ) 

217 

218 

219class _ToolSpendSums(BaseModel): 

220 spend: float = 0.0 

221 total_tokens: int = 0 

222 request_count: int = 0 

223 

224 

225class _TopToolRow(BaseModel): 

226 tool_name: str 

227 sums: _ToolSpendSums = Field(alias="_sum") 

228 

229 

230_TOP_TOOL_ROWS: Final = TypeAdapter(list[_TopToolRow]) 

231 

232 

233@router.get( 

234 "/v1/tool/spend", 

235 tags=["tool management"], 

236 dependencies=[Depends(user_api_key_auth)], 

237 response_model=ToolSpendResponse, 

238) 

239async def get_tool_spend( 

240 user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], 

241 start_date: Annotated[str | None, Query(description="YYYY-MM-DD (defaults to 30 days ago)")] = None, 

242 end_date: Annotated[str | None, Query(description="YYYY-MM-DD (defaults to today)")] = None, 

243): 

244 """ 

245 Spend attributed to each tool over a date range, for the Cost Optimization dashboard. 

246 

247 Reads the ``LiteLLM_DailyToolSpend`` rollup, written at request time from invoked 

248 tools only (MCP tool calls and response tool_calls; declaring a tool without 

249 invoking it does not count). A request that invoked multiple tools counts its 

250 full spend toward each of them, so per-tool numbers are attributions and do not 

251 sum to a deduplicated total. 

252 

253 ``by_tool`` is the top ``TOOL_SPEND_TOP_TOOLS`` tools by spend, aggregated in 

254 SQL, and ``daily`` covers only those tools, so the response is bounded by 

255 days x TOOL_SPEND_TOP_TOOLS regardless of the requested range or how many 

256 distinct tool names exist. 

257 """ 

258 from litellm.proxy.proxy_server import prisma_client 

259 

260 if user_api_key_dict.user_role not in ( 260 ↛ 264line 260 didn't jump to line 264 because the condition on line 260 was never true

261 LitellmUserRoles.PROXY_ADMIN, 

262 LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, 

263 ): 

264 raise HTTPException( 

265 status_code=403, 

266 detail="Only proxy admin roles can view tool spend across the deployment", 

267 ) 

268 

269 if prisma_client is None: 269 ↛ 270line 269 didn't jump to line 270 because the condition on line 269 was never true

270 raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value) 

271 

272 end_day: Final = _parse_day_start(end_date) or datetime.now(timezone.utc) 

273 start_day: Final = _parse_day_start(start_date) or end_day - timedelta(days=30) 

274 start_str: Final = start_day.strftime("%Y-%m-%d") 

275 end_str: Final = end_day.strftime("%Y-%m-%d") 

276 date_window: Final = {"date": {"gte": start_str, "lte": end_str}} 

277 

278 table: Final = _typed_table(DailyToolSpendRepository(prisma_client)) 

279 top_tools: Final = _TOP_TOOL_ROWS.validate_python( 

280 await table.group_by( 

281 by=["tool_name"], 

282 sum={"spend": True, "total_tokens": True, "request_count": True}, 

283 where=date_window, 

284 order={"_sum": {"spend": "desc"}}, 

285 take=TOOL_SPEND_TOP_TOOLS, 

286 ) 

287 or [] 

288 ) 

289 by_tool: Final = [ 

290 ToolSpendEntry( 

291 tool_name=row.tool_name, 

292 spend=row.sums.spend, 

293 call_count=row.sums.request_count, 

294 total_tokens=row.sums.total_tokens, 

295 ) 

296 for row in top_tools 

297 ] 

298 

299 daily_rows: Final[Sequence[PrismaDailyToolSpendRow]] = ( 

300 await table.find_many( 

301 where={**date_window, "tool_name": {"in": [row.tool_name for row in top_tools]}}, 

302 order=[{"date": "asc"}, {"spend": "desc"}], 

303 ) 

304 if top_tools 

305 else [] 

306 ) 

307 daily: Final = [ 

308 ToolSpendDailyEntry(date=row.date, tool_name=row.tool_name, spend=row.spend, call_count=row.request_count) 

309 for row in daily_rows 

310 ] 

311 return ToolSpendResponse(by_tool=by_tool, daily=daily, start_date=start_str, end_date=end_str) 

312 

313 

314@router.get( 

315 "/v1/tool/{tool_name:path}/detail", 

316 tags=["tool management"], 

317 dependencies=[Depends(user_api_key_auth)], 

318 response_model=ToolDetailResponse, 

319) 

320async def get_tool_detail( 

321 tool_name: str, 

322 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), 

323): 

324 """ 

325 Get a single tool with its policy overrides (for UI detail view). 

326 """ 

327 from litellm.proxy.db.tool_registry_writer import get_tool as db_get_tool 

328 from litellm.proxy.db.tool_registry_writer import list_overrides_for_tool 

329 from litellm.proxy.proxy_server import prisma_client 

330 

331 if prisma_client is None: 331 ↛ 332line 331 didn't jump to line 332 because the condition on line 331 was never true

332 raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value) 

333 

334 try: 

335 tool: Final = await db_get_tool(prisma_client=prisma_client, tool_name=tool_name) 

336 if tool is None: 

337 raise HTTPException(status_code=404, detail=f"Tool '{tool_name}' not found") 

338 overrides: Final = await list_overrides_for_tool(prisma_client=prisma_client, tool_name=tool_name) 

339 return ToolDetailResponse(tool=tool, overrides=overrides) 

340 except HTTPException: 

341 raise 

342 except Exception as e: 

343 verbose_proxy_logger.exception("Error getting tool detail: %s", e) 

344 raise HTTPException(status_code=500, detail=str(e)) 

345 

346 

347_ParsedJson: TypeAlias = dict[str, object] | list[object] | str | int | float | bool | None 

348_PARSED_JSON: Final[TypeAdapter[_ParsedJson]] = TypeAdapter(_ParsedJson) 

349_STR_OBJECT_DICT: Final = TypeAdapter(dict[str, object]) 

350 

351 

352def _input_snippet_for_tool_log(sl: "_SpendLogRow | None", max_len: int = 200) -> str | None: 

353 """Short snippet from messages or proxy_server_request for tool usage log row.""" 

354 if sl is None: 

355 return None 

356 messages: Final = sl.messages 

357 if messages is not None: 

358 s = _snippet_str(messages, max_len) 

359 if s: 

360 return s 

361 psr = sl.proxy_server_request 

362 if not psr: 

363 return None 

364 if isinstance(psr, str): 

365 import json 

366 

367 try: 

368 psr = _PARSED_JSON.validate_python(json.loads(psr)) 

369 except Exception: 

370 return _snippet_str(psr, max_len) 

371 if isinstance(psr, dict): 

372 msgs = psr.get("messages") 

373 if msgs is None: 

374 body: Final = psr.get("body") 

375 if isinstance(body, dict): 

376 msgs = _STR_OBJECT_DICT.validate_python(body).get("messages") 

377 s = _snippet_str(msgs, max_len) 

378 if s: 

379 return s 

380 return _snippet_str(psr, max_len) 

381 

382 

383def _snippet_str(text: object, max_len: int = 200) -> str | None: 

384 if text is None: 

385 return None 

386 if isinstance(text, str): 

387 s = text 

388 elif isinstance(text, list): 

389 parts: Final = [] 

390 for item in text: 

391 if isinstance(item, dict) and "content" in item: 

392 c = item["content"] 

393 parts.append(c if isinstance(c, str) else str(c)) 

394 else: 

395 parts.append(str(item)) 

396 s = " ".join(parts) 

397 else: 

398 s = str(text) 

399 if not s or s == "{}": 

400 return None 

401 return (s[:max_len] + "...") if len(s) > max_len else s 

402 

403 

404@router.get( 

405 "/v1/tool/{tool_name:path}/logs", 

406 tags=["tool management"], 

407 dependencies=[Depends(user_api_key_auth)], 

408 response_model=ToolUsageLogsResponse, 

409) 

410async def get_tool_usage_logs( 

411 tool_name: str, 

412 page: int = Query(1, ge=1), 

413 page_size: int = Query(50, ge=1, le=100), 

414 start_date: str | None = Query(None, description="YYYY-MM-DD"), 

415 end_date: str | None = Query(None, description="YYYY-MM-DD"), 

416 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), 

417): 

418 """ 

419 Return paginated spend logs for requests that invoked this tool (from SpendLogToolIndex). 

420 Declaring a tool in a request body without the model invoking it does not create an entry. 

421 """ 

422 from litellm.proxy.proxy_server import prisma_client 

423 

424 if prisma_client is None: 424 ↛ 425line 424 didn't jump to line 425 because the condition on line 424 was never true

425 raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value) 

426 

427 try: 

428 where: Final[dict[str, object]] = {"tool_name": tool_name} 

429 if start_date or end_date: 

430 start_time_filter: datetime | None = None 

431 end_time_filter: datetime | None = None 

432 if start_date: 

433 try: 

434 start_time_filter = datetime.strptime(start_date + "T00:00:00", "%Y-%m-%dT%H:%M:%S").replace( 

435 tzinfo=timezone.utc 

436 ) 

437 except ValueError: 

438 pass 

439 if end_date: 

440 try: 

441 end_time_filter = datetime.strptime(end_date + "T23:59:59", "%Y-%m-%dT%H:%M:%S").replace( 

442 tzinfo=timezone.utc 

443 ) 

444 except ValueError: 

445 pass 

446 if start_time_filter is not None or end_time_filter is not None: 

447 where["start_time"] = { 

448 key: value 

449 for key, value in (("gte", start_time_filter), ("lte", end_time_filter)) 

450 if value is not None 

451 } 

452 

453 total: Final = await _typed_table(SpendLogToolIndexRepository(prisma_client)).count(where=where) 

454 index_rows: Final = await _typed_table(SpendLogToolIndexRepository(prisma_client)).find_many( 

455 where=where, 

456 order={"start_time": "desc"}, 

457 skip=(page - 1) * page_size, 

458 take=page_size, 

459 ) 

460 request_ids: Final = [r.request_id for r in index_rows] 

461 if not request_ids: 461 ↛ 464line 461 didn't jump to line 464 because the condition on line 461 was always true

462 return ToolUsageLogsResponse(logs=[], total=total, page=page, page_size=page_size) 

463 

464 spend_logs = await _typed_table(SpendLogsRepository(prisma_client)).find_many( 

465 where={"request_id": {"in": request_ids}} 

466 ) 

467 log_by_id: Final = {s.request_id: s for s in spend_logs} 

468 

469 logs_out: Final[list[ToolUsageLogEntry]] = [] 

470 for r in index_rows: 

471 sl = log_by_id.get(r.request_id) 

472 if not sl: 

473 continue 

474 ts = sl.startTime.isoformat() if hasattr(sl.startTime, "isoformat") else str(sl.startTime) 

475 logs_out.append( 

476 ToolUsageLogEntry( 

477 id=sl.request_id, 

478 timestamp=ts, 

479 model=getattr(sl, "model", None) or None, 

480 spend=getattr(sl, "spend", None), 

481 total_tokens=getattr(sl, "total_tokens", None), 

482 input_snippet=_input_snippet_for_tool_log(sl), 

483 ) 

484 ) 

485 

486 return ToolUsageLogsResponse(logs=logs_out, total=total, page=page, page_size=page_size) 

487 except HTTPException: 

488 raise 

489 except Exception as e: 

490 verbose_proxy_logger.exception("Error getting tool usage logs: %s", e) 

491 raise HTTPException(status_code=500, detail=str(e)) 

492 

493 

494@router.get( 

495 "/v1/tool/{tool_name:path}", 

496 tags=["tool management"], 

497 dependencies=[Depends(user_api_key_auth)], 

498 response_model=LiteLLM_ToolTableRow, 

499) 

500async def get_tool( 

501 tool_name: str, 

502 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), 

503): 

504 """ 

505 Get details for a single tool. 

506 """ 

507 from litellm.proxy.db.tool_registry_writer import get_tool as db_get_tool 

508 from litellm.proxy.proxy_server import prisma_client 

509 

510 if prisma_client is None: 510 ↛ 511line 510 didn't jump to line 511 because the condition on line 510 was never true

511 raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value) 

512 

513 try: 

514 tool: Final = await db_get_tool(prisma_client=prisma_client, tool_name=tool_name) 

515 if tool is None: 

516 raise HTTPException(status_code=404, detail=f"Tool '{tool_name}' not found") 

517 return tool 

518 except HTTPException: 

519 raise 

520 except Exception as e: 

521 verbose_proxy_logger.exception("Error getting tool: %s", e) 

522 raise HTTPException(status_code=500, detail=str(e)) 

523 

524 

525async def _resolve_key_hash_to_object_permission_id( 

526 prisma_client: "PrismaClient", 

527 key_hash: str, 

528) -> str | None: 

529 """Resolve key (hash or raw) to object_permission_id; create permission if key has none.""" 

530 from litellm.proxy.proxy_server import hash_token 

531 

532 hashed: Final = key_hash if "sk-" not in (key_hash or "") else hash_token(key_hash) 

533 if not hashed: 

534 return None 

535 row = await _typed_table(VerificationTokenRepository(prisma_client)).find_unique(where={"token": hashed}) 

536 if row is None: 536 ↛ 538line 536 didn't jump to line 538 because the condition on line 536 was always true

537 return None 

538 op_id: Final = row.object_permission_id 

539 if op_id: 

540 return op_id 

541 new_id: Final = str(uuid.uuid4()) 

542 await _typed_table(ObjectPermissionRepository(prisma_client)).create( 

543 data={"object_permission_id": new_id, "blocked_tools": []} 

544 ) 

545 updated_count: Final = await _typed_table(VerificationTokenRepository(prisma_client)).update_many( 

546 where={"token": hashed, "object_permission_id": None}, 

547 data={"object_permission_id": new_id}, 

548 ) 

549 if updated_count == 0: 

550 await _typed_table(ObjectPermissionRepository(prisma_client)).delete(where={"object_permission_id": new_id}) 

551 row = await _typed_table(VerificationTokenRepository(prisma_client)).find_unique(where={"token": hashed}) 

552 return row.object_permission_id if row else None 

553 return new_id 

554 

555 

556async def _resolve_team_id_to_object_permission_id( 

557 prisma_client: "PrismaClient", 

558 team_id: str, 

559) -> str | None: 

560 """Resolve team_id to object_permission_id; create permission if team has none.""" 

561 if not team_id or not team_id.strip(): 

562 return None 

563 team_id_clean: Final = team_id.strip() 

564 row = await _typed_table(TeamRepository(prisma_client)).find_unique(where={"team_id": team_id_clean}) 

565 if row is None: 

566 return None 

567 op_id: Final = row.object_permission_id 

568 if op_id: 

569 return op_id 

570 new_id: Final = str(uuid.uuid4()) 

571 await _typed_table(ObjectPermissionRepository(prisma_client)).create( 

572 data={"object_permission_id": new_id, "blocked_tools": []} 

573 ) 

574 updated_count: Final = await _typed_table(TeamRepository(prisma_client)).update_many( 

575 where={"team_id": team_id_clean, "object_permission_id": None}, 

576 data={"object_permission_id": new_id}, 

577 ) 

578 if updated_count == 0: 578 ↛ 579line 578 didn't jump to line 579 because the condition on line 578 was never true

579 await _typed_table(ObjectPermissionRepository(prisma_client)).delete(where={"object_permission_id": new_id}) 

580 row = await _typed_table(TeamRepository(prisma_client)).find_unique(where={"team_id": team_id_clean}) 

581 return row.object_permission_id if row else None 

582 return new_id 

583 

584 

585@router.post( 

586 "/v1/tool/policy", 

587 tags=["tool management"], 

588 dependencies=[Depends(user_api_key_auth)], 

589 response_model=ToolPolicyUpdateResponse, 

590) 

591async def update_tool_policy( 

592 data: ToolPolicyUpdateRequest, 

593 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), 

594): 

595 """ 

596 Set the input_policy and/or output_policy for a tool (global), or block for a specific team/key (override). 

597 

598 Parameters: 

599 - tool_name: str - The tool to update 

600 - input_policy: optional - "trusted" | "untrusted" | "blocked" 

601 - output_policy: optional - "trusted" | "untrusted" 

602 - team_id: optional - if set, create/update override for this team only 

603 - key_hash: optional - if set, create/update override for this key only 

604 """ 

605 from litellm.proxy.db.tool_registry_writer import ( 

606 add_tool_to_object_permission_blocked, 

607 get_tool_policy_registry, 

608 remove_tool_from_object_permission_blocked, 

609 ) 

610 from litellm.proxy.db.tool_registry_writer import ( 

611 update_tool_policy as db_update_tool_policy, 

612 ) 

613 from litellm.proxy.proxy_server import prisma_client 

614 

615 if prisma_client is None: 615 ↛ 616line 615 didn't jump to line 616 because the condition on line 615 was never true

616 raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value) 

617 

618 try: 

619 if data.team_id is not None or data.key_hash is not None: 

620 if data.team_id is not None and data.key_hash is not None: 

621 raise HTTPException( 

622 status_code=400, 

623 detail="Provide either team_id or key_hash, not both", 

624 ) 

625 if data.key_hash is not None: 

626 op_id = await _resolve_key_hash_to_object_permission_id(prisma_client, data.key_hash) 

627 else: 

628 op_id = await _resolve_team_id_to_object_permission_id(prisma_client, data.team_id or "") 

629 if op_id is None: 

630 raise HTTPException( 

631 status_code=404, 

632 detail="Key or team not found for the given identifier", 

633 ) 

634 is_blocking: Final = data.input_policy == "blocked" 

635 if is_blocking: 635 ↛ 636line 635 didn't jump to line 636 because the condition on line 635 was never true

636 ok = await add_tool_to_object_permission_blocked( 

637 prisma_client=prisma_client, 

638 object_permission_id=op_id, 

639 tool_name=data.tool_name, 

640 ) 

641 else: 

642 ok = await remove_tool_from_object_permission_blocked( 

643 prisma_client=prisma_client, 

644 object_permission_id=op_id, 

645 tool_name=data.tool_name, 

646 ) 

647 if not ok: 647 ↛ 652line 647 didn't jump to line 652 because the condition on line 647 was always true

648 raise HTTPException( 

649 status_code=500, 

650 detail=f"Failed to update policy override for tool '{data.tool_name}'", 

651 ) 

652 registry = get_tool_policy_registry() 

653 if registry.is_initialized(): 

654 await registry.sync_tool_policy_from_db(prisma_client) 

655 return ToolPolicyUpdateResponse( 

656 tool_name=data.tool_name, 

657 input_policy=data.input_policy, 

658 output_policy=data.output_policy, 

659 updated=True, 

660 team_id=data.team_id, 

661 key_hash=data.key_hash, 

662 ) 

663 

664 if data.input_policy is None and data.output_policy is None: 

665 raise HTTPException( 

666 status_code=400, 

667 detail="At least one of input_policy or output_policy must be provided", 

668 ) 

669 

670 updated: Final = await db_update_tool_policy( 

671 prisma_client=prisma_client, 

672 tool_name=data.tool_name, 

673 updated_by=user_api_key_dict.user_id, 

674 input_policy=data.input_policy, 

675 output_policy=data.output_policy, 

676 ) 

677 if updated is None: 677 ↛ 678line 677 didn't jump to line 678 because the condition on line 677 was never true

678 raise HTTPException( 

679 status_code=500, 

680 detail=f"Failed to update policy for tool '{data.tool_name}'", 

681 ) 

682 registry = get_tool_policy_registry() 

683 if registry.is_initialized(): 683 ↛ 685line 683 didn't jump to line 685 because the condition on line 683 was always true

684 await registry.sync_tool_policy_from_db(prisma_client) 

685 return ToolPolicyUpdateResponse( 

686 tool_name=updated.tool_name, 

687 input_policy=updated.input_policy, 

688 output_policy=updated.output_policy, 

689 updated=True, 

690 ) 

691 except HTTPException: 

692 raise 

693 except Exception as e: 

694 verbose_proxy_logger.exception("Error updating tool policy: %s", e) 

695 raise HTTPException(status_code=500, detail=str(e)) 

696 

697 

698@router.delete( 

699 "/v1/tool/{tool_name:path}/overrides", 

700 tags=["tool management"], 

701 dependencies=[Depends(user_api_key_auth)], 

702) 

703async def delete_tool_policy_override( 

704 tool_name: str, 

705 team_id: str | None = Query(None, description="Team ID of the override to remove"), 

706 key_hash: str | None = Query(None, description="Key hash of the override to remove"), 

707 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), 

708): 

709 """ 

710 Remove a policy override for a tool. Specify the override by team_id or key_hash 

711 (exactly one required). 

712 """ 

713 from litellm.proxy.db.tool_registry_writer import ( 

714 get_tool_policy_registry, 

715 remove_tool_from_object_permission_blocked, 

716 ) 

717 from litellm.proxy.proxy_server import prisma_client 

718 

719 if prisma_client is None: 719 ↛ 720line 719 didn't jump to line 720 because the condition on line 719 was never true

720 raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value) 

721 if team_id is None and key_hash is None: 721 ↛ 722line 721 didn't jump to line 722 because the condition on line 721 was never true

722 raise HTTPException( 

723 status_code=400, 

724 detail="At least one of team_id or key_hash is required to identify the override", 

725 ) 

726 if team_id is not None and key_hash is not None: 

727 raise HTTPException( 

728 status_code=400, 

729 detail="Provide either team_id or key_hash, not both", 

730 ) 

731 try: 

732 if key_hash is not None: 

733 op_id = await _resolve_key_hash_to_object_permission_id(prisma_client, key_hash) 

734 else: 

735 op_id = await _resolve_team_id_to_object_permission_id(prisma_client, team_id or "") 

736 if op_id is None: 

737 raise HTTPException( 

738 status_code=404, 

739 detail="Key or team not found for the given identifier", 

740 ) 

741 deleted: Final = await remove_tool_from_object_permission_blocked( 

742 prisma_client=prisma_client, 

743 object_permission_id=op_id, 

744 tool_name=tool_name, 

745 ) 

746 if not deleted: 746 ↛ 751line 746 didn't jump to line 751 because the condition on line 746 was always true

747 raise HTTPException( 

748 status_code=404, 

749 detail=f"No override found for tool '{tool_name}' with the given scope", 

750 ) 

751 registry: Final = get_tool_policy_registry() 

752 if registry.is_initialized(): 

753 await registry.sync_tool_policy_from_db(prisma_client) 

754 return {"deleted": True, "tool_name": tool_name} 

755 except HTTPException: 

756 raise 

757 except Exception as e: 

758 verbose_proxy_logger.exception("Error deleting tool policy override: %s", e) 

759 raise HTTPException(status_code=500, detail=str(e))