mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
add policy_resolve_router
This commit is contained in:
parent
62f9172ecd
commit
602b81d149
3 changed files with 310 additions and 19 deletions
|
|
@ -23,10 +23,6 @@ from litellm.types.proxy.policy_engine import (
|
|||
|
||||
router = APIRouter()
|
||||
|
||||
# Get singleton instances
|
||||
POLICY_REGISTRY = get_policy_registry()
|
||||
ATTACHMENT_REGISTRY = get_attachment_registry()
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# Policy CRUD Endpoints
|
||||
|
|
@ -75,7 +71,7 @@ async def list_policies():
|
|||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
|
||||
try:
|
||||
policies = await POLICY_REGISTRY.get_all_policies_from_db(prisma_client)
|
||||
policies = await get_policy_registry().get_all_policies_from_db(prisma_client)
|
||||
return PolicyListDBResponse(policies=policies, total_count=len(policies))
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error listing policies: {e}")
|
||||
|
|
@ -130,7 +126,7 @@ async def create_policy(
|
|||
|
||||
try:
|
||||
created_by = user_api_key_dict.user_id
|
||||
result = await POLICY_REGISTRY.add_policy_to_db(
|
||||
result = await get_policy_registry().add_policy_to_db(
|
||||
policy_request=request,
|
||||
prisma_client=prisma_client,
|
||||
created_by=created_by,
|
||||
|
|
@ -168,7 +164,7 @@ async def get_policy(policy_id: str):
|
|||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
|
||||
try:
|
||||
result = await POLICY_REGISTRY.get_policy_by_id_from_db(
|
||||
result = await get_policy_registry().get_policy_by_id_from_db(
|
||||
policy_id=policy_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
|
@ -216,7 +212,7 @@ async def update_policy(
|
|||
|
||||
try:
|
||||
# Check if policy exists
|
||||
existing = await POLICY_REGISTRY.get_policy_by_id_from_db(
|
||||
existing = await get_policy_registry().get_policy_by_id_from_db(
|
||||
policy_id=policy_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
|
@ -226,7 +222,7 @@ async def update_policy(
|
|||
)
|
||||
|
||||
updated_by = user_api_key_dict.user_id
|
||||
result = await POLICY_REGISTRY.update_policy_in_db(
|
||||
result = await get_policy_registry().update_policy_in_db(
|
||||
policy_id=policy_id,
|
||||
policy_request=request,
|
||||
prisma_client=prisma_client,
|
||||
|
|
@ -269,7 +265,7 @@ async def delete_policy(policy_id: str):
|
|||
|
||||
try:
|
||||
# Check if policy exists
|
||||
existing = await POLICY_REGISTRY.get_policy_by_id_from_db(
|
||||
existing = await get_policy_registry().get_policy_by_id_from_db(
|
||||
policy_id=policy_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
|
@ -278,7 +274,7 @@ async def delete_policy(policy_id: str):
|
|||
status_code=404, detail=f"Policy with ID {policy_id} not found"
|
||||
)
|
||||
|
||||
result = await POLICY_REGISTRY.delete_policy_from_db(
|
||||
result = await get_policy_registry().delete_policy_from_db(
|
||||
policy_id=policy_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
|
@ -324,7 +320,7 @@ async def get_resolved_guardrails(policy_id: str):
|
|||
|
||||
try:
|
||||
# Get the policy
|
||||
policy = await POLICY_REGISTRY.get_policy_by_id_from_db(
|
||||
policy = await get_policy_registry().get_policy_by_id_from_db(
|
||||
policy_id=policy_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
|
@ -334,7 +330,7 @@ async def get_resolved_guardrails(policy_id: str):
|
|||
)
|
||||
|
||||
# Resolve guardrails
|
||||
resolved = await POLICY_REGISTRY.resolve_guardrails_from_db(
|
||||
resolved = await get_policy_registry().resolve_guardrails_from_db(
|
||||
policy_name=policy.policy_name,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
|
@ -399,7 +395,7 @@ async def list_policy_attachments():
|
|||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
|
||||
try:
|
||||
attachments = await ATTACHMENT_REGISTRY.get_all_attachments_from_db(
|
||||
attachments = await get_attachment_registry().get_all_attachments_from_db(
|
||||
prisma_client
|
||||
)
|
||||
return PolicyAttachmentListResponse(
|
||||
|
|
@ -466,7 +462,7 @@ async def create_policy_attachment(
|
|||
|
||||
try:
|
||||
# Verify the policy exists
|
||||
policy = await POLICY_REGISTRY.get_all_policies_from_db(prisma_client)
|
||||
policy = await get_policy_registry().get_all_policies_from_db(prisma_client)
|
||||
policy_names = [p.policy_name for p in policy]
|
||||
if request.policy_name not in policy_names:
|
||||
raise HTTPException(
|
||||
|
|
@ -475,7 +471,7 @@ async def create_policy_attachment(
|
|||
)
|
||||
|
||||
created_by = user_api_key_dict.user_id
|
||||
result = await ATTACHMENT_REGISTRY.add_attachment_to_db(
|
||||
result = await get_attachment_registry().add_attachment_to_db(
|
||||
attachment_request=request,
|
||||
prisma_client=prisma_client,
|
||||
created_by=created_by,
|
||||
|
|
@ -510,7 +506,7 @@ async def get_policy_attachment(attachment_id: str):
|
|||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
|
||||
try:
|
||||
result = await ATTACHMENT_REGISTRY.get_attachment_by_id_from_db(
|
||||
result = await get_attachment_registry().get_attachment_by_id_from_db(
|
||||
attachment_id=attachment_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
|
@ -556,7 +552,7 @@ async def delete_policy_attachment(attachment_id: str):
|
|||
|
||||
try:
|
||||
# Check if attachment exists
|
||||
existing = await ATTACHMENT_REGISTRY.get_attachment_by_id_from_db(
|
||||
existing = await get_attachment_registry().get_attachment_by_id_from_db(
|
||||
attachment_id=attachment_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
|
@ -566,7 +562,7 @@ async def delete_policy_attachment(attachment_id: str):
|
|||
detail=f"Attachment with ID {attachment_id} not found",
|
||||
)
|
||||
|
||||
result = await ATTACHMENT_REGISTRY.delete_attachment_from_db(
|
||||
result = await get_attachment_registry().delete_attachment_from_db(
|
||||
attachment_id=attachment_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
|
|
|||
291
litellm/proxy/policy_engine/policy_resolve_endpoints.py
Normal file
291
litellm/proxy/policy_engine/policy_resolve_endpoints.py
Normal file
|
|
@ -0,0 +1,291 @@
|
|||
"""
|
||||
Policy resolve and attachment impact estimation endpoints.
|
||||
|
||||
- /policies/resolve — debug which guardrails apply for a given context
|
||||
- /policies/attachments/estimate-impact — preview blast radius before creating an attachment
|
||||
"""
|
||||
|
||||
import json
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry
|
||||
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
|
||||
from litellm.types.proxy.policy_engine import (
|
||||
AttachmentImpactResponse,
|
||||
PolicyAttachmentCreateRequest,
|
||||
PolicyMatchContext,
|
||||
PolicyMatchDetail,
|
||||
PolicyResolveRequest,
|
||||
PolicyResolveResponse,
|
||||
)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# Policy Resolve Endpoint
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.post(
|
||||
"/policies/resolve",
|
||||
tags=["Policies"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=PolicyResolveResponse,
|
||||
)
|
||||
async def resolve_policies_for_context(
|
||||
request: PolicyResolveRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Resolve which policies and guardrails apply for a given context.
|
||||
|
||||
Use this endpoint to debug "what guardrails would apply to a request
|
||||
with this team/key/model/tags combination?"
|
||||
|
||||
Example Request:
|
||||
```bash
|
||||
curl -X POST "http://localhost:4000/policies/resolve" \\
|
||||
-H "Authorization: Bearer <your_api_key>" \\
|
||||
-H "Content-Type: application/json" \\
|
||||
-d '{
|
||||
"tags": ["healthcare"],
|
||||
"model": "gpt-4"
|
||||
}'
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.policy_engine.policy_matcher import PolicyMatcher
|
||||
from litellm.proxy.policy_engine.policy_resolver import PolicyResolver
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
|
||||
try:
|
||||
# Sync from DB to ensure in-memory state is current
|
||||
await get_policy_registry().sync_policies_from_db(prisma_client)
|
||||
await get_attachment_registry().sync_attachments_from_db(prisma_client)
|
||||
|
||||
# Build context from request
|
||||
context = PolicyMatchContext(
|
||||
team_alias=request.team_alias,
|
||||
key_alias=request.key_alias,
|
||||
model=request.model,
|
||||
tags=request.tags,
|
||||
)
|
||||
|
||||
# Get matching policies with reasons
|
||||
match_results = get_attachment_registry().get_attached_policies_with_reasons(
|
||||
context=context
|
||||
)
|
||||
|
||||
if not match_results:
|
||||
return PolicyResolveResponse(
|
||||
effective_guardrails=[],
|
||||
matched_policies=[],
|
||||
)
|
||||
|
||||
# Filter by conditions
|
||||
policy_names = [r["policy_name"] for r in match_results]
|
||||
applied_policy_names = PolicyMatcher.get_policies_with_matching_conditions(
|
||||
policy_names=policy_names,
|
||||
context=context,
|
||||
)
|
||||
|
||||
# Resolve guardrails for each applied policy
|
||||
matched_policies = []
|
||||
all_guardrails: set = set()
|
||||
for result in match_results:
|
||||
pname = result["policy_name"]
|
||||
if pname not in applied_policy_names:
|
||||
continue
|
||||
resolved = PolicyResolver.resolve_policy_guardrails(
|
||||
policy_name=pname,
|
||||
policies=get_policy_registry().get_all_policies(),
|
||||
context=context,
|
||||
)
|
||||
guardrails = resolved.guardrails if resolved else []
|
||||
all_guardrails.update(guardrails)
|
||||
matched_policies.append(
|
||||
PolicyMatchDetail(
|
||||
policy_name=pname,
|
||||
matched_via=result["matched_via"],
|
||||
guardrails_added=guardrails,
|
||||
)
|
||||
)
|
||||
|
||||
return PolicyResolveResponse(
|
||||
effective_guardrails=sorted(all_guardrails),
|
||||
matched_policies=matched_policies,
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error resolving policies: {e}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# Attachment Impact Estimation Endpoint
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.post(
|
||||
"/policies/attachments/estimate-impact",
|
||||
tags=["Policies"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=AttachmentImpactResponse,
|
||||
)
|
||||
async def estimate_attachment_impact(
|
||||
request: PolicyAttachmentCreateRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Estimate how many keys and teams would be affected by a policy attachment.
|
||||
|
||||
Use this before creating an attachment to preview the blast radius.
|
||||
|
||||
Example Request:
|
||||
```bash
|
||||
curl -X POST "http://localhost:4000/policies/attachments/estimate-impact" \\
|
||||
-H "Authorization: Bearer <your_api_key>" \\
|
||||
-H "Content-Type: application/json" \\
|
||||
-d '{
|
||||
"policy_name": "hipaa-compliance",
|
||||
"tags": ["healthcare", "health-*"]
|
||||
}'
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
|
||||
try:
|
||||
tag_patterns = request.tags or []
|
||||
team_patterns = request.teams or []
|
||||
|
||||
affected_keys: list = []
|
||||
affected_teams: list = []
|
||||
|
||||
# If global scope, everything is affected — not useful to enumerate
|
||||
if request.scope == "*":
|
||||
return AttachmentImpactResponse(
|
||||
affected_keys_count=-1,
|
||||
affected_teams_count=-1,
|
||||
sample_keys=["(global scope — affects all keys)"],
|
||||
sample_teams=["(global scope — affects all teams)"],
|
||||
)
|
||||
|
||||
# Check tag-based impact: find keys/teams whose metadata.tags match the patterns
|
||||
if tag_patterns:
|
||||
# Query keys with metadata
|
||||
keys = await prisma_client.db.litellm_verificationtoken.find_many(
|
||||
where={},
|
||||
order={"created_at": "desc"},
|
||||
)
|
||||
for key in keys:
|
||||
key_alias = key.key_alias or ""
|
||||
key_metadata = key.metadata_json if hasattr(key, "metadata_json") else (key.metadata or {})
|
||||
if isinstance(key_metadata, str):
|
||||
try:
|
||||
key_metadata = json.loads(key_metadata)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
key_metadata = {}
|
||||
key_tags = key_metadata.get("tags", []) if isinstance(key_metadata, dict) else []
|
||||
if key_tags and any(
|
||||
RouteChecks._route_matches_wildcard_pattern(
|
||||
route=tag, pattern=pattern
|
||||
)
|
||||
for tag in key_tags
|
||||
for pattern in tag_patterns
|
||||
):
|
||||
affected_keys.append(key_alias or str(key.token)[:8] + "...")
|
||||
|
||||
# Query teams with metadata
|
||||
teams = await prisma_client.db.litellm_teamtable.find_many(
|
||||
where={},
|
||||
order={"created_at": "desc"},
|
||||
)
|
||||
for team in teams:
|
||||
team_alias = team.team_alias or ""
|
||||
team_metadata = team.metadata or {}
|
||||
if isinstance(team_metadata, str):
|
||||
try:
|
||||
team_metadata = json.loads(team_metadata)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
team_metadata = {}
|
||||
team_tags = team_metadata.get("tags", []) if isinstance(team_metadata, dict) else []
|
||||
if team_tags and any(
|
||||
RouteChecks._route_matches_wildcard_pattern(
|
||||
route=tag, pattern=pattern
|
||||
)
|
||||
for tag in team_tags
|
||||
for pattern in tag_patterns
|
||||
):
|
||||
affected_teams.append(team_alias or str(team.team_id)[:8] + "...")
|
||||
|
||||
# Check team-based impact
|
||||
matched_team_ids: list = []
|
||||
if team_patterns:
|
||||
teams = await prisma_client.db.litellm_teamtable.find_many(
|
||||
where={},
|
||||
order={"created_at": "desc"},
|
||||
)
|
||||
for team in teams:
|
||||
team_alias = team.team_alias or ""
|
||||
if team_alias and any(
|
||||
RouteChecks._route_matches_wildcard_pattern(
|
||||
route=team_alias, pattern=pattern
|
||||
)
|
||||
for pattern in team_patterns
|
||||
):
|
||||
if team_alias not in affected_teams:
|
||||
affected_teams.append(team_alias)
|
||||
matched_team_ids.append(str(team.team_id))
|
||||
|
||||
# Also find keys belonging to matched teams
|
||||
if matched_team_ids:
|
||||
keys = await prisma_client.db.litellm_verificationtoken.find_many(
|
||||
where={"team_id": {"in": matched_team_ids}},
|
||||
order={"created_at": "desc"},
|
||||
)
|
||||
for key in keys:
|
||||
key_alias = key.key_alias or str(key.token)[:8] + "..."
|
||||
if key_alias not in affected_keys:
|
||||
affected_keys.append(key_alias)
|
||||
|
||||
# Check key-based impact (direct key alias matching)
|
||||
key_patterns = request.keys or []
|
||||
if key_patterns:
|
||||
keys = await prisma_client.db.litellm_verificationtoken.find_many(
|
||||
where={},
|
||||
order={"created_at": "desc"},
|
||||
)
|
||||
for key in keys:
|
||||
key_alias = key.key_alias or ""
|
||||
if key_alias and any(
|
||||
RouteChecks._route_matches_wildcard_pattern(
|
||||
route=key_alias, pattern=pattern
|
||||
)
|
||||
for pattern in key_patterns
|
||||
):
|
||||
if key_alias not in affected_keys:
|
||||
affected_keys.append(key_alias)
|
||||
|
||||
return AttachmentImpactResponse(
|
||||
affected_keys_count=len(affected_keys),
|
||||
affected_teams_count=len(affected_teams),
|
||||
sample_keys=affected_keys[:10],
|
||||
sample_teams=affected_teams[:10],
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error estimating attachment impact: {e}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
|
@ -427,6 +427,9 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
|||
router as pass_through_router,
|
||||
)
|
||||
from litellm.proxy.policy_engine.policy_endpoints import router as policy_crud_router
|
||||
from litellm.proxy.policy_engine.policy_resolve_endpoints import (
|
||||
router as policy_resolve_router,
|
||||
)
|
||||
from litellm.proxy.prompts.prompt_endpoints import router as prompts_router
|
||||
from litellm.proxy.public_endpoints import router as public_endpoints_router
|
||||
from litellm.proxy.rag_endpoints.endpoints import router as rag_router
|
||||
|
|
@ -11746,6 +11749,7 @@ app.include_router(analytics_router)
|
|||
app.include_router(guardrails_router)
|
||||
app.include_router(policy_router)
|
||||
app.include_router(policy_crud_router)
|
||||
app.include_router(policy_resolve_router)
|
||||
app.include_router(search_tool_management_router)
|
||||
app.include_router(prompts_router)
|
||||
app.include_router(callback_management_endpoints_router)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue