Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/policy_engine/policy_resolve_endpoints.py: 68%

168 statements  

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

1""" 

2Policy resolve and attachment impact estimation endpoints. 

3 

4- /policies/resolve — debug which guardrails apply for a given context 

5- /policies/attachments/estimate-impact — preview blast radius before creating an attachment 

6""" 

7 

8import json 

9from collections.abc import Sequence 

10from typing import TYPE_CHECKING, Final 

11 

12from fastapi import APIRouter, Depends, HTTPException, Query 

13 

14from litellm._logging import verbose_proxy_logger 

15from litellm.constants import MAX_POLICY_ESTIMATE_IMPACT_ROWS 

16from litellm.proxy._types import UserAPIKeyAuth 

17from litellm.proxy.auth.route_checks import RouteChecks 

18from litellm.proxy.auth.user_api_key_auth import user_api_key_auth 

19from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry 

20from litellm.proxy.policy_engine.policy_registry import get_policy_registry 

21from litellm.repositories.team_repository import TeamRepository 

22from litellm.repositories.verification_token_repository import ( 

23 VerificationTokenRepository, 

24) 

25from litellm.types.proxy.policy_engine import ( 

26 AttachmentImpactResponse, 

27 PolicyAttachmentCreateRequest, 

28 PolicyMatchContext, 

29 PolicyMatchDetail, 

30 PolicyResolveRequest, 

31 PolicyResolveResponse, 

32) 

33 

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

35 from prisma import models as prisma_models 

36 

37router: Final = APIRouter() 

38 

39 

40def _build_alias_where(field: str, patterns: Sequence[str]) -> dict[str, object]: 

41 """Build a Prisma ``where`` clause for alias patterns. 

42 

43 Supports exact matches and suffix wildcards (``prefix*``). 

44 Returns something like: 

45 {"OR": [{"field": {"in": ["a","b"]}}, {"field": {"startsWith": "dev-"}}]} 

46 """ 

47 exact: Final[list[str]] = [] 

48 prefix_conditions: Final[list[dict[str, object]]] = [] 

49 for pat in patterns: 

50 if pat.endswith("*"): 50 ↛ 51line 50 didn't jump to line 51 because the condition on line 50 was never true

51 prefix_conditions.append({field: {"startsWith": pat[:-1]}}) 

52 else: 

53 exact.append(pat) 

54 

55 conditions: Final[list[dict[str, object]]] = [] 

56 if exact: 56 ↛ 58line 56 didn't jump to line 58 because the condition on line 56 was always true

57 conditions.append({field: {"in": exact}}) 

58 conditions.extend(prefix_conditions) 

59 

60 if not conditions: 60 ↛ 61line 60 didn't jump to line 61 because the condition on line 60 was never true

61 return {field: {"not": None}} 

62 if len(conditions) == 1: 62 ↛ 64line 62 didn't jump to line 64 because the condition on line 62 was always true

63 return conditions[0] 

64 return {"OR": conditions} 

65 

66 

67def _parse_metadata(raw_metadata: object) -> dict: 

68 """Parse metadata that may be a dict, JSON string, or None.""" 

69 if raw_metadata is None: 69 ↛ 70line 69 didn't jump to line 70 because the condition on line 69 was never true

70 return {} 

71 if isinstance(raw_metadata, str): 71 ↛ 72line 71 didn't jump to line 72 because the condition on line 71 was never true

72 try: 

73 return json.loads(raw_metadata) 

74 except (json.JSONDecodeError, TypeError): 

75 return {} 

76 return raw_metadata if isinstance(raw_metadata, dict) else {} 

77 

78 

79def _get_tags_from_metadata(metadata: object, json_metadata: object = None) -> list: 

80 """Extract tags list from a metadata field (or metadata_json fallback).""" 

81 raw: Final = json_metadata if json_metadata is not None else metadata 

82 parsed: Final = _parse_metadata(raw) 

83 return parsed.get("tags", []) or [] 

84 

85 

86async def _fetch_all_teams(prisma_client: object) -> "Sequence[prisma_models.LiteLLM_TeamTable]": 

87 """Fetch teams from DB once. Reuse the result across tag and alias lookups.""" 

88 return await TeamRepository(prisma_client).table.find_many( 

89 where={}, 

90 order={"created_at": "desc"}, 

91 take=MAX_POLICY_ESTIMATE_IMPACT_ROWS, 

92 ) 

93 

94 

95def _filter_keys_by_tags( 

96 keys: "Sequence[prisma_models.LiteLLM_VerificationToken]", tag_patterns: Sequence[str] 

97) -> tuple[list[str], int]: 

98 """Filter key rows whose metadata.tags match any of the given patterns. 

99 

100 Returns (named_aliases, unnamed_count). 

101 """ 

102 

103 affected: Final[list[str]] = [] 

104 unnamed_count = 0 

105 for key in keys: 

106 key_alias = key.key_alias or "" 

107 key_tags = _get_tags_from_metadata(key.metadata, getattr(key, "metadata_json", None)) 

108 if key_tags and any( 108 ↛ 113line 108 didn't jump to line 113 because the condition on line 108 was never true

109 RouteChecks.route_matches_wildcard_pattern(route=tag, pattern=pat) 

110 for tag in key_tags 

111 for pat in tag_patterns 

112 ): 

113 if key_alias: 

114 affected.append(key_alias) 

115 else: 

116 unnamed_count += 1 

117 return affected, unnamed_count 

118 

119 

120def _filter_teams_by_tags( 

121 teams: "Sequence[prisma_models.LiteLLM_TeamTable]", tag_patterns: Sequence[str] 

122) -> tuple[list[str], int]: 

123 """Filter pre-fetched team rows whose metadata.tags match any patterns. 

124 

125 Returns (named_aliases, unnamed_count). 

126 """ 

127 

128 affected: Final[list[str]] = [] 

129 unnamed_count = 0 

130 for team in teams: 

131 team_alias = team.team_alias or "" 

132 team_tags = _get_tags_from_metadata(team.metadata) 

133 if team_tags and any( 133 ↛ 138line 133 didn't jump to line 138 because the condition on line 133 was never true

134 RouteChecks.route_matches_wildcard_pattern(route=tag, pattern=pat) 

135 for tag in team_tags 

136 for pat in tag_patterns 

137 ): 

138 if team_alias: 

139 affected.append(team_alias) 

140 else: 

141 unnamed_count += 1 

142 return affected, unnamed_count 

143 

144 

145async def _find_affected_by_team_patterns( 

146 prisma_client: object, 

147 all_teams: "Sequence[prisma_models.LiteLLM_TeamTable]", 

148 team_patterns: Sequence[str], 

149 existing_teams: Sequence[str], 

150 existing_keys: Sequence[str], 

151) -> tuple[list[str], list[str], int]: 

152 """Filter pre-fetched teams by alias patterns, then fetch their keys. 

153 

154 Returns (new_teams, new_keys, unnamed_keys_count). 

155 """ 

156 

157 new_teams: Final[list[str]] = [] 

158 matched_team_ids: Final[list[str]] = [] 

159 

160 for team in all_teams: 

161 team_alias = team.team_alias or "" 

162 if team_alias and any( 162 ↛ 165line 162 didn't jump to line 165 because the condition on line 162 was never true

163 RouteChecks.route_matches_wildcard_pattern(route=team_alias, pattern=pat) for pat in team_patterns 

164 ): 

165 if team_alias not in existing_teams: 

166 new_teams.append(team_alias) 

167 matched_team_ids.append(str(team.team_id)) 

168 

169 new_keys: Final[list[str]] = [] 

170 unnamed_keys_count = 0 

171 if matched_team_ids: 171 ↛ 172line 171 didn't jump to line 172 because the condition on line 171 was never true

172 keys: Final = await VerificationTokenRepository(prisma_client).table.find_many( 

173 where={"team_id": {"in": matched_team_ids}}, 

174 order={"created_at": "desc"}, 

175 take=MAX_POLICY_ESTIMATE_IMPACT_ROWS, 

176 ) 

177 for key in keys: 

178 key_alias = key.key_alias or "" 

179 if key_alias: 

180 if key_alias not in existing_keys: 

181 new_keys.append(key_alias) 

182 else: 

183 unnamed_keys_count += 1 

184 

185 return new_teams, new_keys, unnamed_keys_count 

186 

187 

188async def _find_affected_keys_by_alias( 

189 prisma_client: object, key_patterns: Sequence[str], existing_keys: Sequence[str] 

190) -> list[str]: 

191 """Find keys whose alias matches the given patterns.""" 

192 

193 affected: Final[list[str]] = [] 

194 

195 keys: Final = await VerificationTokenRepository(prisma_client).table.find_many( 

196 where=_build_alias_where("key_alias", key_patterns), 

197 order={"created_at": "desc"}, 

198 take=MAX_POLICY_ESTIMATE_IMPACT_ROWS, 

199 ) 

200 for key in keys: 200 ↛ 201line 200 didn't jump to line 201 because the loop on line 200 never started

201 key_alias = key.key_alias or "" 

202 if key_alias and any( 

203 RouteChecks.route_matches_wildcard_pattern(route=key_alias, pattern=pat) for pat in key_patterns 

204 ): 

205 if key_alias not in existing_keys: 

206 affected.append(key_alias) 

207 return affected 

208 

209 

210# ───────────────────────────────────────────────────────────────────────────── 

211# Policy Resolve Endpoint 

212# ───────────────────────────────────────────────────────────────────────────── 

213 

214 

215@router.post( 

216 "/policies/resolve", 

217 tags=["Policies"], 

218 dependencies=[Depends(user_api_key_auth)], 

219 response_model=PolicyResolveResponse, 

220) 

221async def resolve_policies_for_context( 

222 request: PolicyResolveRequest, 

223 force_sync: bool = Query( 

224 default=False, 

225 description="Force a DB sync before resolving. Default uses in-memory cache.", 

226 ), 

227 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), 

228): 

229 """ 

230 Resolve which policies and guardrails apply for a given context. 

231 

232 Use this endpoint to debug "what guardrails would apply to a request 

233 with this team/key/model/tags combination?" 

234 

235 Example Request: 

236 ```bash 

237 curl -X POST "http://localhost:4000/policies/resolve" \\ 

238 -H "Authorization: Bearer <your_api_key>" \\ 

239 -H "Content-Type: application/json" \\ 

240 -d '{ 

241 "tags": ["healthcare"], 

242 "model": "gpt-4" 

243 }' 

244 ``` 

245 """ 

246 from litellm.proxy.policy_engine.policy_matcher import PolicyMatcher 

247 from litellm.proxy.policy_engine.policy_resolver import PolicyResolver 

248 from litellm.proxy.proxy_server import prisma_client 

249 

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

251 raise HTTPException(status_code=500, detail="Database not connected") 

252 

253 try: 

254 # Only sync from DB when explicitly requested; otherwise use in-memory cache 

255 if force_sync: 

256 await get_policy_registry().sync_policies_from_db(prisma_client) 

257 await get_attachment_registry().sync_attachments_from_db(prisma_client) 

258 

259 # Build context from request 

260 context: Final = PolicyMatchContext( 

261 team_alias=request.team_alias, 

262 key_alias=request.key_alias, 

263 model=request.model, 

264 tags=request.tags, 

265 ) 

266 

267 # Get matching policies with reasons 

268 match_results: Final = get_attachment_registry().get_attached_policies_with_reasons( 

269 context=context, policy_applies=PolicyMatcher.policy_applies(context) 

270 ) 

271 

272 if not match_results: 

273 return PolicyResolveResponse( 

274 effective_guardrails=[], 

275 matched_policies=[], 

276 ) 

277 

278 # Filter by conditions 

279 policy_names: Final = [r["policy_name"] for r in match_results] 

280 applied_policy_names: Final = PolicyMatcher.get_policies_with_matching_conditions( 

281 policy_names=policy_names, 

282 context=context, 

283 ) 

284 

285 # Resolve guardrails for each applied policy 

286 matched_policies: Final = [] 

287 all_guardrails: Final[set] = set() 

288 for result in match_results: 

289 pname = result["policy_name"] 

290 if pname not in applied_policy_names: 290 ↛ 291line 290 didn't jump to line 291 because the condition on line 290 was never true

291 continue 

292 resolved = PolicyResolver.resolve_policy_guardrails( 

293 policy_name=pname, 

294 policies=get_policy_registry().get_all_policies(), 

295 context=context, 

296 ) 

297 guardrails = resolved.guardrails if resolved else [] 

298 all_guardrails.update(guardrails) 

299 matched_policies.append( 

300 PolicyMatchDetail( 

301 policy_name=pname, 

302 matched_via=result["matched_via"], 

303 guardrails_added=guardrails, 

304 ) 

305 ) 

306 

307 return PolicyResolveResponse( 

308 effective_guardrails=sorted(all_guardrails), 

309 matched_policies=matched_policies, 

310 ) 

311 except HTTPException: 

312 raise 

313 except Exception as e: 

314 verbose_proxy_logger.exception("Error resolving policies: %s", e) 

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

316 

317 

318# ───────────────────────────────────────────────────────────────────────────── 

319# Attachment Impact Estimation Endpoint 

320# ───────────────────────────────────────────────────────────────────────────── 

321 

322 

323@router.post( 

324 "/policies/attachments/estimate-impact", 

325 tags=["Policies"], 

326 dependencies=[Depends(user_api_key_auth)], 

327 response_model=AttachmentImpactResponse, 

328) 

329async def estimate_attachment_impact( 

330 request: PolicyAttachmentCreateRequest, 

331 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), 

332): 

333 """ 

334 Estimate how many keys and teams would be affected by a policy attachment. 

335 

336 Use this before creating an attachment to preview the blast radius. 

337 

338 Example Request: 

339 ```bash 

340 curl -X POST "http://localhost:4000/policies/attachments/estimate-impact" \\ 

341 -H "Authorization: Bearer <your_api_key>" \\ 

342 -H "Content-Type: application/json" \\ 

343 -d '{ 

344 "policy_name": "hipaa-compliance", 

345 "tags": ["healthcare", "health-*"] 

346 }' 

347 ``` 

348 """ 

349 from litellm.proxy.proxy_server import prisma_client 

350 

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

352 raise HTTPException(status_code=500, detail="Database not connected") 

353 

354 try: 

355 # If global scope, everything is affected — not useful to enumerate 

356 if request.scope == "*": 356 ↛ 357line 356 didn't jump to line 357 because the condition on line 356 was never true

357 return AttachmentImpactResponse( 

358 affected_keys_count=-1, 

359 affected_teams_count=-1, 

360 sample_keys=["(global scope — affects all keys)"], 

361 sample_teams=["(global scope — affects all teams)"], 

362 ) 

363 

364 affected_keys: list[str] = [] 

365 affected_teams: list[str] = [] 

366 unnamed_keys = 0 

367 unnamed_teams = 0 

368 

369 tag_patterns: Final = request.tags or [] 

370 team_patterns: Final = request.teams or [] 

371 

372 # Fetch teams once — reused by both tag-based and alias-based lookups 

373 all_teams: Sequence[prisma_models.LiteLLM_TeamTable] = [] 

374 if tag_patterns or team_patterns: 

375 all_teams = await _fetch_all_teams(prisma_client) 

376 

377 # Tag-based impact 

378 if tag_patterns: 

379 keys: Final = await VerificationTokenRepository(prisma_client).table.find_many( 

380 where={}, 

381 order={"created_at": "desc"}, 

382 take=MAX_POLICY_ESTIMATE_IMPACT_ROWS, 

383 ) 

384 affected_keys, unnamed_keys = _filter_keys_by_tags(keys, tag_patterns) 

385 affected_teams, unnamed_teams = _filter_teams_by_tags( 

386 all_teams, 

387 tag_patterns, 

388 ) 

389 

390 # Team-based impact (alias matching + keys belonging to those teams) 

391 if team_patterns: 

392 new_teams, new_keys, new_unnamed = await _find_affected_by_team_patterns( 

393 prisma_client, 

394 all_teams, 

395 team_patterns, 

396 affected_teams, 

397 affected_keys, 

398 ) 

399 affected_teams.extend(new_teams) 

400 affected_keys.extend(new_keys) 

401 unnamed_keys += new_unnamed 

402 

403 # Key-based impact (direct alias matching) 

404 key_patterns: Final = request.keys or [] 

405 if key_patterns: 

406 new_keys = await _find_affected_keys_by_alias( 

407 prisma_client, 

408 key_patterns, 

409 affected_keys, 

410 ) 

411 affected_keys.extend(new_keys) 

412 

413 return AttachmentImpactResponse( 

414 affected_keys_count=len(affected_keys) + unnamed_keys, 

415 affected_teams_count=len(affected_teams) + unnamed_teams, 

416 unnamed_keys_count=unnamed_keys, 

417 unnamed_teams_count=unnamed_teams, 

418 sample_keys=affected_keys[:10], 

419 sample_teams=affected_teams[:10], 

420 ) 

421 except HTTPException: 

422 raise 

423 except Exception as e: 

424 verbose_proxy_logger.exception("Error estimating attachment impact: %s", e) 

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