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
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""
2Policy resolve and attachment impact estimation endpoints.
4- /policies/resolve — debug which guardrails apply for a given context
5- /policies/attachments/estimate-impact — preview blast radius before creating an attachment
6"""
8import json
9from collections.abc import Sequence
10from typing import TYPE_CHECKING, Final
12from fastapi import APIRouter, Depends, HTTPException, Query
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)
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
37router: Final = APIRouter()
40def _build_alias_where(field: str, patterns: Sequence[str]) -> dict[str, object]:
41 """Build a Prisma ``where`` clause for alias patterns.
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)
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)
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}
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 {}
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 []
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 )
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.
100 Returns (named_aliases, unnamed_count).
101 """
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
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.
125 Returns (named_aliases, unnamed_count).
126 """
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
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.
154 Returns (new_teams, new_keys, unnamed_keys_count).
155 """
157 new_teams: Final[list[str]] = []
158 matched_team_ids: Final[list[str]] = []
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))
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
185 return new_teams, new_keys, unnamed_keys_count
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."""
193 affected: Final[list[str]] = []
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
210# ─────────────────────────────────────────────────────────────────────────────
211# Policy Resolve Endpoint
212# ─────────────────────────────────────────────────────────────────────────────
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.
232 Use this endpoint to debug "what guardrails would apply to a request
233 with this team/key/model/tags combination?"
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
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")
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)
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 )
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 )
272 if not match_results:
273 return PolicyResolveResponse(
274 effective_guardrails=[],
275 matched_policies=[],
276 )
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 )
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 )
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))
318# ─────────────────────────────────────────────────────────────────────────────
319# Attachment Impact Estimation Endpoint
320# ─────────────────────────────────────────────────────────────────────────────
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.
336 Use this before creating an attachment to preview the blast radius.
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
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")
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 )
364 affected_keys: list[str] = []
365 affected_teams: list[str] = []
366 unnamed_keys = 0
367 unnamed_teams = 0
369 tag_patterns: Final = request.tags or []
370 team_patterns: Final = request.teams or []
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)
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 )
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
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)
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))