mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
init schema.prisma
This commit is contained in:
parent
bfa5d59dd2
commit
1a3cb2a1d3
9 changed files with 1303 additions and 7 deletions
|
|
@ -73,6 +73,7 @@ class SupportedDBObjectType(str, enum.Enum):
|
|||
MODELS = "models"
|
||||
MCP = "mcp"
|
||||
GUARDRAILS = "guardrails"
|
||||
POLICIES = "policies"
|
||||
VECTOR_STORES = "vector_stores"
|
||||
PASS_THROUGH_ENDPOINTS = "pass_through_endpoints"
|
||||
PROMPTS = "prompts"
|
||||
|
|
@ -844,6 +845,7 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase):
|
|||
model_rpm_limit: Optional[dict] = None
|
||||
model_tpm_limit: Optional[dict] = None
|
||||
guardrails: Optional[List[str]] = None
|
||||
policies: Optional[List[str]] = None
|
||||
prompts: Optional[List[str]] = None
|
||||
blocked: Optional[bool] = None
|
||||
aliases: Optional[dict] = {}
|
||||
|
|
@ -1477,6 +1479,7 @@ class NewTeamRequest(TeamBase):
|
|||
model_aliases: Optional[dict] = None
|
||||
tags: Optional[list] = None
|
||||
guardrails: Optional[List[str]] = None
|
||||
policies: Optional[List[str]] = None
|
||||
prompts: Optional[List[str]] = None
|
||||
object_permission: Optional[LiteLLM_ObjectPermissionBase] = None
|
||||
allowed_passthrough_routes: Optional[list] = None
|
||||
|
|
@ -1526,6 +1529,7 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase):
|
|||
blocked: Optional[bool] = None
|
||||
budget_duration: Optional[str] = None
|
||||
guardrails: Optional[List[str]] = None
|
||||
policies: Optional[List[str]] = None
|
||||
"""
|
||||
|
||||
team_id: str # required
|
||||
|
|
@ -1541,6 +1545,7 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase):
|
|||
tags: Optional[list] = None
|
||||
model_aliases: Optional[dict] = None
|
||||
guardrails: Optional[List[str]] = None
|
||||
policies: Optional[List[str]] = None
|
||||
object_permission: Optional[LiteLLM_ObjectPermissionBase] = None
|
||||
team_member_budget: Optional[float] = None
|
||||
team_member_budget_duration: Optional[str] = None
|
||||
|
|
@ -3499,6 +3504,7 @@ LiteLLM_ManagementEndpoint_MetadataFields = [
|
|||
|
||||
LiteLLM_ManagementEndpoint_MetadataFields_Premium = [
|
||||
"guardrails",
|
||||
"policies",
|
||||
"tags",
|
||||
"team_member_key_duration",
|
||||
"prompts",
|
||||
|
|
|
|||
|
|
@ -14,11 +14,11 @@ import copy
|
|||
import json
|
||||
import secrets
|
||||
import traceback
|
||||
import yaml
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any, Dict, List, Literal, Optional, Tuple, cast
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
|
||||
import fastapi
|
||||
import yaml
|
||||
from fastapi import APIRouter, Depends, Header, HTTPException, Query, Request, status
|
||||
|
||||
import litellm
|
||||
|
|
@ -31,6 +31,7 @@ from litellm.constants import (
|
|||
UI_SESSION_TOKEN_TEAM_ID,
|
||||
)
|
||||
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.proxy._experimental.mcp_server.db import (
|
||||
rotate_mcp_server_credentials_master_key,
|
||||
)
|
||||
|
|
@ -2077,6 +2078,7 @@ async def generate_key_helper_fn( # noqa: PLR0915
|
|||
model_rpm_limit: Optional[dict] = None,
|
||||
model_tpm_limit: Optional[dict] = None,
|
||||
guardrails: Optional[list] = None,
|
||||
policies: Optional[list] = None,
|
||||
prompts: Optional[list] = None,
|
||||
teams: Optional[list] = None,
|
||||
organization_id: Optional[str] = None,
|
||||
|
|
@ -2139,6 +2141,9 @@ async def generate_key_helper_fn( # noqa: PLR0915
|
|||
if guardrails is not None:
|
||||
metadata = metadata or {}
|
||||
metadata["guardrails"] = guardrails
|
||||
if policies is not None:
|
||||
metadata = metadata or {}
|
||||
metadata["policies"] = policies
|
||||
if prompts is not None:
|
||||
metadata = metadata or {}
|
||||
metadata["prompts"] = prompts
|
||||
|
|
|
|||
|
|
@ -22,11 +22,13 @@ from pydantic import BaseModel
|
|||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.proxy._types import (
|
||||
BlockTeamRequest,
|
||||
CommonProxyErrors,
|
||||
DeleteTeamRequest,
|
||||
LiteLLM_AuditLogs,
|
||||
LiteLLM_DeletedTeamTable,
|
||||
LiteLLM_ManagementEndpoint_MetadataFields,
|
||||
LiteLLM_ManagementEndpoint_MetadataFields_Premium,
|
||||
LiteLLM_ModelTable,
|
||||
|
|
@ -34,7 +36,6 @@ from litellm.proxy._types import (
|
|||
LiteLLM_OrganizationTableWithMembers,
|
||||
LiteLLM_TeamMembership,
|
||||
LiteLLM_TeamTable,
|
||||
LiteLLM_DeletedTeamTable,
|
||||
LiteLLM_TeamTableCachedObj,
|
||||
LiteLLM_UserTable,
|
||||
LiteLLM_VerificationToken,
|
||||
|
|
@ -102,7 +103,7 @@ from litellm.types.proxy.management_endpoints.team_endpoints import (
|
|||
TeamMemberAddResult,
|
||||
UpdateTeamMemberPermissionsRequest,
|
||||
)
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
|
|
@ -689,6 +690,7 @@ async def new_team( # noqa: PLR0915
|
|||
- organization_id: Optional[str] - The organization id of the team. Default is None. Create via `/organization/new`.
|
||||
- model_aliases: Optional[dict] - Model aliases for the team. [Docs](https://docs.litellm.ai/docs/proxy/team_based_routing#create-team-with-model-alias)
|
||||
- guardrails: Optional[List[str]] - Guardrails for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails)
|
||||
- policies: Optional[List[str]] - Policies for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails/guardrail_policies)
|
||||
- disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the key.
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - team-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"], "agents": ["agent_1", "agent_2"], "agent_access_groups": ["dev_group"]}. IF null or {} then no object permission.
|
||||
- team_member_budget: Optional[float] - The maximum budget allocated to an individual team member.
|
||||
|
|
@ -1228,6 +1230,7 @@ async def update_team( # noqa: PLR0915
|
|||
- organization_id: Optional[str] - The organization id of the team. Default is None. Create via `/organization/new`.
|
||||
- model_aliases: Optional[dict] - Model aliases for the team. [Docs](https://docs.litellm.ai/docs/proxy/team_based_routing#create-team-with-model-alias)
|
||||
- guardrails: Optional[List[str]] - Guardrails for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails)
|
||||
- policies: Optional[List[str]] - Policies for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails/guardrail_policies)
|
||||
- disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the key.
|
||||
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - team-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"], "agents": ["agent_1", "agent_2"], "agent_access_groups": ["dev_group"]}. IF null or {} then no object permission.
|
||||
- team_member_budget: Optional[float] - The maximum budget allocated to an individual team member.
|
||||
|
|
|
|||
|
|
@ -5,14 +5,20 @@ Attachments define WHERE policies apply, separate from the policy definitions.
|
|||
This allows the same policy to be attached to multiple scopes.
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, List, Optional
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.types.proxy.policy_engine import (
|
||||
PolicyAttachment,
|
||||
PolicyAttachmentCreateRequest,
|
||||
PolicyAttachmentDBResponse,
|
||||
PolicyMatchContext,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
|
||||
class AttachmentRegistry:
|
||||
"""
|
||||
|
|
@ -188,6 +194,238 @@ class AttachmentRegistry:
|
|||
)
|
||||
return removed_count
|
||||
|
||||
def remove_attachment_by_id(self, attachment_id: str) -> bool:
|
||||
"""
|
||||
Remove an attachment by its ID (for DB-synced attachments).
|
||||
|
||||
Args:
|
||||
attachment_id: The ID of the attachment to remove
|
||||
|
||||
Returns:
|
||||
True if removed, False if not found
|
||||
"""
|
||||
# Note: In-memory attachments don't have IDs, so this is primarily
|
||||
# for consistency after DB operations
|
||||
return False
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────
|
||||
# Database CRUD Methods
|
||||
# ─────────────────────────────────────────────────────────────────────────
|
||||
|
||||
async def add_attachment_to_db(
|
||||
self,
|
||||
attachment_request: PolicyAttachmentCreateRequest,
|
||||
prisma_client: "PrismaClient",
|
||||
created_by: Optional[str] = None,
|
||||
) -> PolicyAttachmentDBResponse:
|
||||
"""
|
||||
Add a policy attachment to the database.
|
||||
|
||||
Args:
|
||||
attachment_request: The attachment creation request
|
||||
prisma_client: The Prisma client instance
|
||||
created_by: User who created the attachment
|
||||
|
||||
Returns:
|
||||
PolicyAttachmentDBResponse with the created attachment
|
||||
"""
|
||||
try:
|
||||
created_attachment = (
|
||||
await prisma_client.db.litellm_policyattachmenttable.create(
|
||||
data={
|
||||
"policy_name": attachment_request.policy_name,
|
||||
"scope": attachment_request.scope,
|
||||
"teams": attachment_request.teams or [],
|
||||
"keys": attachment_request.keys or [],
|
||||
"models": attachment_request.models or [],
|
||||
"created_at": datetime.now(timezone.utc),
|
||||
"updated_at": datetime.now(timezone.utc),
|
||||
"created_by": created_by,
|
||||
"updated_by": created_by,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
# Also add to in-memory registry
|
||||
attachment = PolicyAttachment(
|
||||
policy=attachment_request.policy_name,
|
||||
scope=attachment_request.scope,
|
||||
teams=attachment_request.teams,
|
||||
keys=attachment_request.keys,
|
||||
models=attachment_request.models,
|
||||
)
|
||||
self.add_attachment(attachment)
|
||||
|
||||
return PolicyAttachmentDBResponse(
|
||||
attachment_id=created_attachment.attachment_id,
|
||||
policy_name=created_attachment.policy_name,
|
||||
scope=created_attachment.scope,
|
||||
teams=created_attachment.teams or [],
|
||||
keys=created_attachment.keys or [],
|
||||
models=created_attachment.models or [],
|
||||
created_at=created_attachment.created_at,
|
||||
updated_at=created_attachment.updated_at,
|
||||
created_by=created_attachment.created_by,
|
||||
updated_by=created_attachment.updated_by,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error adding attachment to DB: {e}")
|
||||
raise Exception(f"Error adding attachment to DB: {str(e)}")
|
||||
|
||||
async def delete_attachment_from_db(
|
||||
self,
|
||||
attachment_id: str,
|
||||
prisma_client: "PrismaClient",
|
||||
) -> Dict[str, str]:
|
||||
"""
|
||||
Delete a policy attachment from the database.
|
||||
|
||||
Args:
|
||||
attachment_id: The ID of the attachment to delete
|
||||
prisma_client: The Prisma client instance
|
||||
|
||||
Returns:
|
||||
Dict with success message
|
||||
"""
|
||||
try:
|
||||
# Get attachment before deleting
|
||||
attachment = (
|
||||
await prisma_client.db.litellm_policyattachmenttable.find_unique(
|
||||
where={"attachment_id": attachment_id}
|
||||
)
|
||||
)
|
||||
|
||||
if attachment is None:
|
||||
raise Exception(f"Attachment with ID {attachment_id} not found")
|
||||
|
||||
# Delete from DB
|
||||
await prisma_client.db.litellm_policyattachmenttable.delete(
|
||||
where={"attachment_id": attachment_id}
|
||||
)
|
||||
|
||||
# Note: In-memory attachments don't have IDs, so we need to sync from DB
|
||||
# to properly update in-memory state
|
||||
await self.sync_attachments_from_db(prisma_client)
|
||||
|
||||
return {"message": f"Attachment {attachment_id} deleted successfully"}
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error deleting attachment from DB: {e}")
|
||||
raise Exception(f"Error deleting attachment from DB: {str(e)}")
|
||||
|
||||
async def get_attachment_by_id_from_db(
|
||||
self,
|
||||
attachment_id: str,
|
||||
prisma_client: "PrismaClient",
|
||||
) -> Optional[PolicyAttachmentDBResponse]:
|
||||
"""
|
||||
Get a policy attachment by ID from the database.
|
||||
|
||||
Args:
|
||||
attachment_id: The ID of the attachment to retrieve
|
||||
prisma_client: The Prisma client instance
|
||||
|
||||
Returns:
|
||||
PolicyAttachmentDBResponse if found, None otherwise
|
||||
"""
|
||||
try:
|
||||
attachment = (
|
||||
await prisma_client.db.litellm_policyattachmenttable.find_unique(
|
||||
where={"attachment_id": attachment_id}
|
||||
)
|
||||
)
|
||||
|
||||
if attachment is None:
|
||||
return None
|
||||
|
||||
return PolicyAttachmentDBResponse(
|
||||
attachment_id=attachment.attachment_id,
|
||||
policy_name=attachment.policy_name,
|
||||
scope=attachment.scope,
|
||||
teams=attachment.teams or [],
|
||||
keys=attachment.keys or [],
|
||||
models=attachment.models or [],
|
||||
created_at=attachment.created_at,
|
||||
updated_at=attachment.updated_at,
|
||||
created_by=attachment.created_by,
|
||||
updated_by=attachment.updated_by,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error getting attachment from DB: {e}")
|
||||
raise Exception(f"Error getting attachment from DB: {str(e)}")
|
||||
|
||||
async def get_all_attachments_from_db(
|
||||
self,
|
||||
prisma_client: "PrismaClient",
|
||||
) -> List[PolicyAttachmentDBResponse]:
|
||||
"""
|
||||
Get all policy attachments from the database.
|
||||
|
||||
Args:
|
||||
prisma_client: The Prisma client instance
|
||||
|
||||
Returns:
|
||||
List of PolicyAttachmentDBResponse objects
|
||||
"""
|
||||
try:
|
||||
attachments = (
|
||||
await prisma_client.db.litellm_policyattachmenttable.find_many(
|
||||
order={"created_at": "desc"},
|
||||
)
|
||||
)
|
||||
|
||||
return [
|
||||
PolicyAttachmentDBResponse(
|
||||
attachment_id=a.attachment_id,
|
||||
policy_name=a.policy_name,
|
||||
scope=a.scope,
|
||||
teams=a.teams or [],
|
||||
keys=a.keys or [],
|
||||
models=a.models or [],
|
||||
created_at=a.created_at,
|
||||
updated_at=a.updated_at,
|
||||
created_by=a.created_by,
|
||||
updated_by=a.updated_by,
|
||||
)
|
||||
for a in attachments
|
||||
]
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error getting attachments from DB: {e}")
|
||||
raise Exception(f"Error getting attachments from DB: {str(e)}")
|
||||
|
||||
async def sync_attachments_from_db(
|
||||
self,
|
||||
prisma_client: "PrismaClient",
|
||||
) -> None:
|
||||
"""
|
||||
Sync policy attachments from the database to in-memory registry.
|
||||
|
||||
Args:
|
||||
prisma_client: The Prisma client instance
|
||||
"""
|
||||
try:
|
||||
attachments = await self.get_all_attachments_from_db(prisma_client)
|
||||
|
||||
# Clear existing attachments and reload from DB
|
||||
self._attachments = []
|
||||
|
||||
for attachment_response in attachments:
|
||||
attachment = PolicyAttachment(
|
||||
policy=attachment_response.policy_name,
|
||||
scope=attachment_response.scope,
|
||||
teams=attachment_response.teams if attachment_response.teams else None,
|
||||
keys=attachment_response.keys if attachment_response.keys else None,
|
||||
models=attachment_response.models if attachment_response.models else None,
|
||||
)
|
||||
self._attachments.append(attachment)
|
||||
|
||||
self._initialized = True
|
||||
verbose_proxy_logger.info(
|
||||
f"Synced {len(attachments)} attachments from DB to in-memory registry"
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error syncing attachments from DB: {e}")
|
||||
raise Exception(f"Error syncing attachments from DB: {str(e)}")
|
||||
|
||||
|
||||
# Global singleton instance
|
||||
_attachment_registry: Optional[AttachmentRegistry] = None
|
||||
|
|
|
|||
517
litellm/proxy/policy_engine/policy_endpoints.py
Normal file
517
litellm/proxy/policy_engine/policy_endpoints.py
Normal file
|
|
@ -0,0 +1,517 @@
|
|||
"""
|
||||
CRUD ENDPOINTS FOR POLICIES
|
||||
|
||||
Provides REST API endpoints for managing policies and policy attachments.
|
||||
"""
|
||||
|
||||
from typing import List, Optional
|
||||
|
||||
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 (
|
||||
PolicyAttachmentCreateRequest,
|
||||
PolicyAttachmentDBResponse,
|
||||
PolicyAttachmentListResponse,
|
||||
PolicyCreateRequest,
|
||||
PolicyDBResponse,
|
||||
PolicyListDBResponse,
|
||||
PolicyUpdateRequest,
|
||||
)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
# Get singleton instances
|
||||
POLICY_REGISTRY = get_policy_registry()
|
||||
ATTACHMENT_REGISTRY = get_attachment_registry()
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# Policy CRUD Endpoints
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.get(
|
||||
"/policies/list",
|
||||
tags=["Policies"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=PolicyListDBResponse,
|
||||
)
|
||||
async def list_policies():
|
||||
"""
|
||||
List all policies from the database.
|
||||
|
||||
Example Request:
|
||||
```bash
|
||||
curl -X GET "http://localhost:4000/policies/list" \\
|
||||
-H "Authorization: Bearer <your_api_key>"
|
||||
```
|
||||
|
||||
Example Response:
|
||||
```json
|
||||
{
|
||||
"policies": [
|
||||
{
|
||||
"policy_id": "123e4567-e89b-12d3-a456-426614174000",
|
||||
"policy_name": "global-baseline",
|
||||
"inherit": null,
|
||||
"description": "Base guardrails for all requests",
|
||||
"guardrails_add": ["pii_masking"],
|
||||
"guardrails_remove": [],
|
||||
"condition": null,
|
||||
"created_at": "2024-01-01T00:00:00Z",
|
||||
"updated_at": "2024-01-01T00:00:00Z"
|
||||
}
|
||||
],
|
||||
"total_count": 1
|
||||
}
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
|
||||
try:
|
||||
policies = await 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}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.post(
|
||||
"/policies",
|
||||
tags=["Policies"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=PolicyDBResponse,
|
||||
)
|
||||
async def create_policy(
|
||||
request: PolicyCreateRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Create a new policy.
|
||||
|
||||
Example Request:
|
||||
```bash
|
||||
curl -X POST "http://localhost:4000/policies" \\
|
||||
-H "Authorization: Bearer <your_api_key>" \\
|
||||
-H "Content-Type: application/json" \\
|
||||
-d '{
|
||||
"policy_name": "global-baseline",
|
||||
"description": "Base guardrails for all requests",
|
||||
"guardrails_add": ["pii_masking", "prompt_injection"],
|
||||
"guardrails_remove": []
|
||||
}'
|
||||
```
|
||||
|
||||
Example Response:
|
||||
```json
|
||||
{
|
||||
"policy_id": "123e4567-e89b-12d3-a456-426614174000",
|
||||
"policy_name": "global-baseline",
|
||||
"inherit": null,
|
||||
"description": "Base guardrails for all requests",
|
||||
"guardrails_add": ["pii_masking", "prompt_injection"],
|
||||
"guardrails_remove": [],
|
||||
"condition": null,
|
||||
"created_at": "2024-01-01T00:00:00Z",
|
||||
"updated_at": "2024-01-01T00:00:00Z"
|
||||
}
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
|
||||
try:
|
||||
created_by = user_api_key_dict.user_id
|
||||
result = await POLICY_REGISTRY.add_policy_to_db(
|
||||
policy_request=request,
|
||||
prisma_client=prisma_client,
|
||||
created_by=created_by,
|
||||
)
|
||||
return result
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error creating policy: {e}")
|
||||
if "unique constraint" in str(e).lower():
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Policy with name '{request.policy_name}' already exists",
|
||||
)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.get(
|
||||
"/policies/{policy_id}",
|
||||
tags=["Policies"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=PolicyDBResponse,
|
||||
)
|
||||
async def get_policy(policy_id: str):
|
||||
"""
|
||||
Get a policy by ID.
|
||||
|
||||
Example Request:
|
||||
```bash
|
||||
curl -X GET "http://localhost:4000/policies/123e4567-e89b-12d3-a456-426614174000" \\
|
||||
-H "Authorization: Bearer <your_api_key>"
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
|
||||
try:
|
||||
result = await POLICY_REGISTRY.get_policy_by_id_from_db(
|
||||
policy_id=policy_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
if result is None:
|
||||
raise HTTPException(
|
||||
status_code=404, detail=f"Policy with ID {policy_id} not found"
|
||||
)
|
||||
return result
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error getting policy: {e}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.put(
|
||||
"/policies/{policy_id}",
|
||||
tags=["Policies"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=PolicyDBResponse,
|
||||
)
|
||||
async def update_policy(
|
||||
policy_id: str,
|
||||
request: PolicyUpdateRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Update an existing policy.
|
||||
|
||||
Example Request:
|
||||
```bash
|
||||
curl -X PUT "http://localhost:4000/policies/123e4567-e89b-12d3-a456-426614174000" \\
|
||||
-H "Authorization: Bearer <your_api_key>" \\
|
||||
-H "Content-Type: application/json" \\
|
||||
-d '{
|
||||
"description": "Updated description",
|
||||
"guardrails_add": ["pii_masking", "toxicity_filter"]
|
||||
}'
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
|
||||
try:
|
||||
# Check if policy exists
|
||||
existing = await POLICY_REGISTRY.get_policy_by_id_from_db(
|
||||
policy_id=policy_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
if existing is None:
|
||||
raise HTTPException(
|
||||
status_code=404, detail=f"Policy with ID {policy_id} not found"
|
||||
)
|
||||
|
||||
updated_by = user_api_key_dict.user_id
|
||||
result = await POLICY_REGISTRY.update_policy_in_db(
|
||||
policy_id=policy_id,
|
||||
policy_request=request,
|
||||
prisma_client=prisma_client,
|
||||
updated_by=updated_by,
|
||||
)
|
||||
return result
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error updating policy: {e}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/policies/{policy_id}",
|
||||
tags=["Policies"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def delete_policy(policy_id: str):
|
||||
"""
|
||||
Delete a policy.
|
||||
|
||||
Example Request:
|
||||
```bash
|
||||
curl -X DELETE "http://localhost:4000/policies/123e4567-e89b-12d3-a456-426614174000" \\
|
||||
-H "Authorization: Bearer <your_api_key>"
|
||||
```
|
||||
|
||||
Example Response:
|
||||
```json
|
||||
{
|
||||
"message": "Policy 123e4567-e89b-12d3-a456-426614174000 deleted successfully"
|
||||
}
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
|
||||
try:
|
||||
# Check if policy exists
|
||||
existing = await POLICY_REGISTRY.get_policy_by_id_from_db(
|
||||
policy_id=policy_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
if existing is None:
|
||||
raise HTTPException(
|
||||
status_code=404, detail=f"Policy with ID {policy_id} not found"
|
||||
)
|
||||
|
||||
result = await POLICY_REGISTRY.delete_policy_from_db(
|
||||
policy_id=policy_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
return result
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error deleting policy: {e}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# Policy Attachment CRUD Endpoints
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.get(
|
||||
"/policies/attachments/list",
|
||||
tags=["Policies"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=PolicyAttachmentListResponse,
|
||||
)
|
||||
async def list_policy_attachments():
|
||||
"""
|
||||
List all policy attachments from the database.
|
||||
|
||||
Example Request:
|
||||
```bash
|
||||
curl -X GET "http://localhost:4000/policies/attachments/list" \\
|
||||
-H "Authorization: Bearer <your_api_key>"
|
||||
```
|
||||
|
||||
Example Response:
|
||||
```json
|
||||
{
|
||||
"attachments": [
|
||||
{
|
||||
"attachment_id": "123e4567-e89b-12d3-a456-426614174000",
|
||||
"policy_name": "global-baseline",
|
||||
"scope": "*",
|
||||
"teams": [],
|
||||
"keys": [],
|
||||
"models": [],
|
||||
"created_at": "2024-01-01T00:00:00Z",
|
||||
"updated_at": "2024-01-01T00:00:00Z"
|
||||
}
|
||||
],
|
||||
"total_count": 1
|
||||
}
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
|
||||
try:
|
||||
attachments = await ATTACHMENT_REGISTRY.get_all_attachments_from_db(
|
||||
prisma_client
|
||||
)
|
||||
return PolicyAttachmentListResponse(
|
||||
attachments=attachments, total_count=len(attachments)
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error listing policy attachments: {e}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.post(
|
||||
"/policies/attachments",
|
||||
tags=["Policies"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=PolicyAttachmentDBResponse,
|
||||
)
|
||||
async def create_policy_attachment(
|
||||
request: PolicyAttachmentCreateRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Create a new policy attachment.
|
||||
|
||||
Example Request:
|
||||
```bash
|
||||
curl -X POST "http://localhost:4000/policies/attachments" \\
|
||||
-H "Authorization: Bearer <your_api_key>" \\
|
||||
-H "Content-Type: application/json" \\
|
||||
-d '{
|
||||
"policy_name": "global-baseline",
|
||||
"scope": "*"
|
||||
}'
|
||||
```
|
||||
|
||||
Example with team-specific attachment:
|
||||
```bash
|
||||
curl -X POST "http://localhost:4000/policies/attachments" \\
|
||||
-H "Authorization: Bearer <your_api_key>" \\
|
||||
-H "Content-Type: application/json" \\
|
||||
-d '{
|
||||
"policy_name": "healthcare-compliance",
|
||||
"teams": ["healthcare-team", "medical-research"]
|
||||
}'
|
||||
```
|
||||
|
||||
Example Response:
|
||||
```json
|
||||
{
|
||||
"attachment_id": "123e4567-e89b-12d3-a456-426614174000",
|
||||
"policy_name": "global-baseline",
|
||||
"scope": "*",
|
||||
"teams": [],
|
||||
"keys": [],
|
||||
"models": [],
|
||||
"created_at": "2024-01-01T00:00:00Z",
|
||||
"updated_at": "2024-01-01T00:00:00Z"
|
||||
}
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
|
||||
try:
|
||||
# Verify the policy exists
|
||||
policy = await 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(
|
||||
status_code=404,
|
||||
detail=f"Policy '{request.policy_name}' not found. Create the policy first.",
|
||||
)
|
||||
|
||||
created_by = user_api_key_dict.user_id
|
||||
result = await ATTACHMENT_REGISTRY.add_attachment_to_db(
|
||||
attachment_request=request,
|
||||
prisma_client=prisma_client,
|
||||
created_by=created_by,
|
||||
)
|
||||
return result
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error creating policy attachment: {e}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.get(
|
||||
"/policies/attachments/{attachment_id}",
|
||||
tags=["Policies"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=PolicyAttachmentDBResponse,
|
||||
)
|
||||
async def get_policy_attachment(attachment_id: str):
|
||||
"""
|
||||
Get a policy attachment by ID.
|
||||
|
||||
Example Request:
|
||||
```bash
|
||||
curl -X GET "http://localhost:4000/policies/attachments/123e4567-e89b-12d3-a456-426614174000" \\
|
||||
-H "Authorization: Bearer <your_api_key>"
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
|
||||
try:
|
||||
result = await ATTACHMENT_REGISTRY.get_attachment_by_id_from_db(
|
||||
attachment_id=attachment_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
if result is None:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail=f"Attachment with ID {attachment_id} not found",
|
||||
)
|
||||
return result
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error getting policy attachment: {e}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/policies/attachments/{attachment_id}",
|
||||
tags=["Policies"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def delete_policy_attachment(attachment_id: str):
|
||||
"""
|
||||
Delete a policy attachment.
|
||||
|
||||
Example Request:
|
||||
```bash
|
||||
curl -X DELETE "http://localhost:4000/policies/attachments/123e4567-e89b-12d3-a456-426614174000" \\
|
||||
-H "Authorization: Bearer <your_api_key>"
|
||||
```
|
||||
|
||||
Example Response:
|
||||
```json
|
||||
{
|
||||
"message": "Attachment 123e4567-e89b-12d3-a456-426614174000 deleted successfully"
|
||||
}
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
|
||||
try:
|
||||
# Check if attachment exists
|
||||
existing = await ATTACHMENT_REGISTRY.get_attachment_by_id_from_db(
|
||||
attachment_id=attachment_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
if existing is None:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail=f"Attachment with ID {attachment_id} not found",
|
||||
)
|
||||
|
||||
result = await ATTACHMENT_REGISTRY.delete_attachment_from_db(
|
||||
attachment_id=attachment_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
return result
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error deleting policy attachment: {e}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
|
@ -7,15 +7,23 @@ Policies define WHAT guardrails to apply. WHERE they apply is defined
|
|||
by policy_attachments (see AttachmentRegistry).
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, List, Optional
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.types.proxy.policy_engine import (
|
||||
Policy,
|
||||
PolicyCondition,
|
||||
PolicyCreateRequest,
|
||||
PolicyDBResponse,
|
||||
PolicyGuardrails,
|
||||
PolicyUpdateRequest,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
|
||||
class PolicyRegistry:
|
||||
"""
|
||||
|
|
@ -178,6 +186,309 @@ class PolicyRegistry:
|
|||
return True
|
||||
return False
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────
|
||||
# Database CRUD Methods
|
||||
# ─────────────────────────────────────────────────────────────────────────
|
||||
|
||||
async def add_policy_to_db(
|
||||
self,
|
||||
policy_request: PolicyCreateRequest,
|
||||
prisma_client: "PrismaClient",
|
||||
created_by: Optional[str] = None,
|
||||
) -> PolicyDBResponse:
|
||||
"""
|
||||
Add a policy to the database.
|
||||
|
||||
Args:
|
||||
policy_request: The policy creation request
|
||||
prisma_client: The Prisma client instance
|
||||
created_by: User who created the policy
|
||||
|
||||
Returns:
|
||||
PolicyDBResponse with the created policy
|
||||
"""
|
||||
try:
|
||||
condition_json = None
|
||||
if policy_request.condition:
|
||||
condition_json = safe_dumps(policy_request.condition.model_dump())
|
||||
|
||||
created_policy = await prisma_client.db.litellm_policytable.create(
|
||||
data={
|
||||
"policy_name": policy_request.policy_name,
|
||||
"inherit": policy_request.inherit,
|
||||
"description": policy_request.description,
|
||||
"guardrails_add": policy_request.guardrails_add or [],
|
||||
"guardrails_remove": policy_request.guardrails_remove or [],
|
||||
"condition": condition_json,
|
||||
"created_at": datetime.now(timezone.utc),
|
||||
"updated_at": datetime.now(timezone.utc),
|
||||
"created_by": created_by,
|
||||
"updated_by": created_by,
|
||||
}
|
||||
)
|
||||
|
||||
# Also add to in-memory registry
|
||||
policy = self._parse_policy(
|
||||
policy_request.policy_name,
|
||||
{
|
||||
"inherit": policy_request.inherit,
|
||||
"description": policy_request.description,
|
||||
"guardrails": {
|
||||
"add": policy_request.guardrails_add,
|
||||
"remove": policy_request.guardrails_remove,
|
||||
},
|
||||
"condition": policy_request.condition.model_dump()
|
||||
if policy_request.condition
|
||||
else None,
|
||||
},
|
||||
)
|
||||
self.add_policy(policy_request.policy_name, policy)
|
||||
|
||||
return PolicyDBResponse(
|
||||
policy_id=created_policy.policy_id,
|
||||
policy_name=created_policy.policy_name,
|
||||
inherit=created_policy.inherit,
|
||||
description=created_policy.description,
|
||||
guardrails_add=created_policy.guardrails_add or [],
|
||||
guardrails_remove=created_policy.guardrails_remove or [],
|
||||
condition=created_policy.condition,
|
||||
created_at=created_policy.created_at,
|
||||
updated_at=created_policy.updated_at,
|
||||
created_by=created_policy.created_by,
|
||||
updated_by=created_policy.updated_by,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error adding policy to DB: {e}")
|
||||
raise Exception(f"Error adding policy to DB: {str(e)}")
|
||||
|
||||
async def update_policy_in_db(
|
||||
self,
|
||||
policy_id: str,
|
||||
policy_request: PolicyUpdateRequest,
|
||||
prisma_client: "PrismaClient",
|
||||
updated_by: Optional[str] = None,
|
||||
) -> PolicyDBResponse:
|
||||
"""
|
||||
Update a policy in the database.
|
||||
|
||||
Args:
|
||||
policy_id: The ID of the policy to update
|
||||
policy_request: The policy update request
|
||||
prisma_client: The Prisma client instance
|
||||
updated_by: User who updated the policy
|
||||
|
||||
Returns:
|
||||
PolicyDBResponse with the updated policy
|
||||
"""
|
||||
try:
|
||||
# Build update data - only include fields that are set
|
||||
update_data: Dict[str, Any] = {
|
||||
"updated_at": datetime.now(timezone.utc),
|
||||
"updated_by": updated_by,
|
||||
}
|
||||
|
||||
if policy_request.policy_name is not None:
|
||||
update_data["policy_name"] = policy_request.policy_name
|
||||
if policy_request.inherit is not None:
|
||||
update_data["inherit"] = policy_request.inherit
|
||||
if policy_request.description is not None:
|
||||
update_data["description"] = policy_request.description
|
||||
if policy_request.guardrails_add is not None:
|
||||
update_data["guardrails_add"] = policy_request.guardrails_add
|
||||
if policy_request.guardrails_remove is not None:
|
||||
update_data["guardrails_remove"] = policy_request.guardrails_remove
|
||||
if policy_request.condition is not None:
|
||||
update_data["condition"] = safe_dumps(
|
||||
policy_request.condition.model_dump()
|
||||
)
|
||||
|
||||
updated_policy = await prisma_client.db.litellm_policytable.update(
|
||||
where={"policy_id": policy_id},
|
||||
data=update_data,
|
||||
)
|
||||
|
||||
# Update in-memory registry
|
||||
policy = self._parse_policy(
|
||||
updated_policy.policy_name,
|
||||
{
|
||||
"inherit": updated_policy.inherit,
|
||||
"description": updated_policy.description,
|
||||
"guardrails": {
|
||||
"add": updated_policy.guardrails_add,
|
||||
"remove": updated_policy.guardrails_remove,
|
||||
},
|
||||
"condition": updated_policy.condition,
|
||||
},
|
||||
)
|
||||
self.add_policy(updated_policy.policy_name, policy)
|
||||
|
||||
return PolicyDBResponse(
|
||||
policy_id=updated_policy.policy_id,
|
||||
policy_name=updated_policy.policy_name,
|
||||
inherit=updated_policy.inherit,
|
||||
description=updated_policy.description,
|
||||
guardrails_add=updated_policy.guardrails_add or [],
|
||||
guardrails_remove=updated_policy.guardrails_remove or [],
|
||||
condition=updated_policy.condition,
|
||||
created_at=updated_policy.created_at,
|
||||
updated_at=updated_policy.updated_at,
|
||||
created_by=updated_policy.created_by,
|
||||
updated_by=updated_policy.updated_by,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error updating policy in DB: {e}")
|
||||
raise Exception(f"Error updating policy in DB: {str(e)}")
|
||||
|
||||
async def delete_policy_from_db(
|
||||
self,
|
||||
policy_id: str,
|
||||
prisma_client: "PrismaClient",
|
||||
) -> Dict[str, str]:
|
||||
"""
|
||||
Delete a policy from the database.
|
||||
|
||||
Args:
|
||||
policy_id: The ID of the policy to delete
|
||||
prisma_client: The Prisma client instance
|
||||
|
||||
Returns:
|
||||
Dict with success message
|
||||
"""
|
||||
try:
|
||||
# Get policy name before deleting
|
||||
policy = await prisma_client.db.litellm_policytable.find_unique(
|
||||
where={"policy_id": policy_id}
|
||||
)
|
||||
|
||||
if policy is None:
|
||||
raise Exception(f"Policy with ID {policy_id} not found")
|
||||
|
||||
# Delete from DB
|
||||
await prisma_client.db.litellm_policytable.delete(
|
||||
where={"policy_id": policy_id}
|
||||
)
|
||||
|
||||
# Remove from in-memory registry
|
||||
self.remove_policy(policy.policy_name)
|
||||
|
||||
return {"message": f"Policy {policy_id} deleted successfully"}
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error deleting policy from DB: {e}")
|
||||
raise Exception(f"Error deleting policy from DB: {str(e)}")
|
||||
|
||||
async def get_policy_by_id_from_db(
|
||||
self,
|
||||
policy_id: str,
|
||||
prisma_client: "PrismaClient",
|
||||
) -> Optional[PolicyDBResponse]:
|
||||
"""
|
||||
Get a policy by ID from the database.
|
||||
|
||||
Args:
|
||||
policy_id: The ID of the policy to retrieve
|
||||
prisma_client: The Prisma client instance
|
||||
|
||||
Returns:
|
||||
PolicyDBResponse if found, None otherwise
|
||||
"""
|
||||
try:
|
||||
policy = await prisma_client.db.litellm_policytable.find_unique(
|
||||
where={"policy_id": policy_id}
|
||||
)
|
||||
|
||||
if policy is None:
|
||||
return None
|
||||
|
||||
return PolicyDBResponse(
|
||||
policy_id=policy.policy_id,
|
||||
policy_name=policy.policy_name,
|
||||
inherit=policy.inherit,
|
||||
description=policy.description,
|
||||
guardrails_add=policy.guardrails_add or [],
|
||||
guardrails_remove=policy.guardrails_remove or [],
|
||||
condition=policy.condition,
|
||||
created_at=policy.created_at,
|
||||
updated_at=policy.updated_at,
|
||||
created_by=policy.created_by,
|
||||
updated_by=policy.updated_by,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error getting policy from DB: {e}")
|
||||
raise Exception(f"Error getting policy from DB: {str(e)}")
|
||||
|
||||
async def get_all_policies_from_db(
|
||||
self,
|
||||
prisma_client: "PrismaClient",
|
||||
) -> List[PolicyDBResponse]:
|
||||
"""
|
||||
Get all policies from the database.
|
||||
|
||||
Args:
|
||||
prisma_client: The Prisma client instance
|
||||
|
||||
Returns:
|
||||
List of PolicyDBResponse objects
|
||||
"""
|
||||
try:
|
||||
policies = await prisma_client.db.litellm_policytable.find_many(
|
||||
order={"created_at": "desc"},
|
||||
)
|
||||
|
||||
return [
|
||||
PolicyDBResponse(
|
||||
policy_id=p.policy_id,
|
||||
policy_name=p.policy_name,
|
||||
inherit=p.inherit,
|
||||
description=p.description,
|
||||
guardrails_add=p.guardrails_add or [],
|
||||
guardrails_remove=p.guardrails_remove or [],
|
||||
condition=p.condition,
|
||||
created_at=p.created_at,
|
||||
updated_at=p.updated_at,
|
||||
created_by=p.created_by,
|
||||
updated_by=p.updated_by,
|
||||
)
|
||||
for p in policies
|
||||
]
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error getting policies from DB: {e}")
|
||||
raise Exception(f"Error getting policies from DB: {str(e)}")
|
||||
|
||||
async def sync_policies_from_db(
|
||||
self,
|
||||
prisma_client: "PrismaClient",
|
||||
) -> None:
|
||||
"""
|
||||
Sync policies from the database to in-memory registry.
|
||||
|
||||
Args:
|
||||
prisma_client: The Prisma client instance
|
||||
"""
|
||||
try:
|
||||
policies = await self.get_all_policies_from_db(prisma_client)
|
||||
|
||||
for policy_response in policies:
|
||||
policy = self._parse_policy(
|
||||
policy_response.policy_name,
|
||||
{
|
||||
"inherit": policy_response.inherit,
|
||||
"description": policy_response.description,
|
||||
"guardrails": {
|
||||
"add": policy_response.guardrails_add,
|
||||
"remove": policy_response.guardrails_remove,
|
||||
},
|
||||
"condition": policy_response.condition,
|
||||
},
|
||||
)
|
||||
self.add_policy(policy_response.policy_name, policy)
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
f"Synced {len(policies)} policies from DB to in-memory registry"
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error syncing policies from DB: {e}")
|
||||
raise Exception(f"Error syncing policies from DB: {str(e)}")
|
||||
|
||||
|
||||
# Global singleton instance
|
||||
_policy_registry: Optional[PolicyRegistry] = None
|
||||
|
|
|
|||
|
|
@ -880,3 +880,36 @@ model LiteLLM_ClaudeCodePluginTable {
|
|||
@@index([name])
|
||||
@@map("litellm_claudecodeplugin")
|
||||
}
|
||||
|
||||
// Policy table for storing policy configurations
|
||||
model LiteLLM_PolicyTable {
|
||||
policy_id String @id @default(uuid())
|
||||
policy_name String @unique
|
||||
inherit String? // Parent policy name for inheritance
|
||||
description String?
|
||||
guardrails_add String[] @default([]) // Guardrails to add
|
||||
guardrails_remove String[] @default([]) // Guardrails to remove (from inherited)
|
||||
condition Json? // Condition for when policy applies, e.g. {model: "gpt-4.*"}
|
||||
created_at DateTime @default(now())
|
||||
updated_at DateTime @updatedAt
|
||||
created_by String?
|
||||
updated_by String?
|
||||
|
||||
@@index([policy_name])
|
||||
}
|
||||
|
||||
// Policy attachment table for storing where policies apply
|
||||
model LiteLLM_PolicyAttachmentTable {
|
||||
attachment_id String @id @default(uuid())
|
||||
policy_name String // Reference to policy
|
||||
scope String? // "*" for global scope
|
||||
teams String[] @default([]) // Team aliases or patterns
|
||||
keys String[] @default([]) // Key aliases or patterns
|
||||
models String[] @default([]) // Model names or patterns
|
||||
created_at DateTime @default(now())
|
||||
updated_at DateTime @updatedAt
|
||||
created_by String?
|
||||
updated_by String?
|
||||
|
||||
@@index([policy_name])
|
||||
}
|
||||
|
|
|
|||
|
|
@ -19,13 +19,21 @@ from litellm.types.proxy.policy_engine.policy_types import (
|
|||
PolicyScope,
|
||||
)
|
||||
from litellm.types.proxy.policy_engine.resolver_types import (
|
||||
PolicyAttachmentCreateRequest,
|
||||
PolicyAttachmentDBResponse,
|
||||
PolicyAttachmentListResponse,
|
||||
PolicyConditionRequest,
|
||||
PolicyCreateRequest,
|
||||
PolicyDBResponse,
|
||||
PolicyGuardrailsResponse,
|
||||
PolicyInfoResponse,
|
||||
PolicyListDBResponse,
|
||||
PolicyListResponse,
|
||||
PolicyMatchContext,
|
||||
PolicyScopeResponse,
|
||||
PolicySummaryItem,
|
||||
PolicyTestResponse,
|
||||
PolicyUpdateRequest,
|
||||
ResolvedPolicy,
|
||||
)
|
||||
from litellm.types.proxy.policy_engine.validation_types import (
|
||||
|
|
@ -58,4 +66,13 @@ __all__ = [
|
|||
"PolicyScopeResponse",
|
||||
"PolicySummaryItem",
|
||||
"PolicyTestResponse",
|
||||
# CRUD Request/Response types
|
||||
"PolicyConditionRequest",
|
||||
"PolicyCreateRequest",
|
||||
"PolicyUpdateRequest",
|
||||
"PolicyDBResponse",
|
||||
"PolicyListDBResponse",
|
||||
"PolicyAttachmentCreateRequest",
|
||||
"PolicyAttachmentDBResponse",
|
||||
"PolicyAttachmentListResponse",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -5,7 +5,8 @@ These types are used for matching requests to policies and resolving
|
|||
the final guardrails list.
|
||||
"""
|
||||
|
||||
from typing import Dict, List, Optional
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
|
|
@ -108,3 +109,168 @@ class PolicyTestResponse(BaseModel):
|
|||
matching_policies: List[str]
|
||||
resolved_guardrails: List[str]
|
||||
message: Optional[str] = None
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# CRUD Request/Response Types for Policy Endpoints
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class PolicyConditionRequest(BaseModel):
|
||||
"""Condition for when a policy applies."""
|
||||
|
||||
model: Optional[str] = Field(
|
||||
default=None,
|
||||
description="Model name pattern (exact match or regex) for when policy applies.",
|
||||
)
|
||||
|
||||
|
||||
class PolicyCreateRequest(BaseModel):
|
||||
"""Request body for creating a new policy."""
|
||||
|
||||
policy_name: str = Field(description="Unique name for the policy.")
|
||||
inherit: Optional[str] = Field(
|
||||
default=None,
|
||||
description="Name of parent policy to inherit from.",
|
||||
)
|
||||
description: Optional[str] = Field(
|
||||
default=None,
|
||||
description="Human-readable description of the policy.",
|
||||
)
|
||||
guardrails_add: Optional[List[str]] = Field(
|
||||
default=None,
|
||||
description="List of guardrail names to add.",
|
||||
)
|
||||
guardrails_remove: Optional[List[str]] = Field(
|
||||
default=None,
|
||||
description="List of guardrail names to remove (from inherited).",
|
||||
)
|
||||
condition: Optional[PolicyConditionRequest] = Field(
|
||||
default=None,
|
||||
description="Condition for when this policy applies.",
|
||||
)
|
||||
|
||||
|
||||
class PolicyUpdateRequest(BaseModel):
|
||||
"""Request body for updating a policy."""
|
||||
|
||||
policy_name: Optional[str] = Field(
|
||||
default=None,
|
||||
description="New name for the policy.",
|
||||
)
|
||||
inherit: Optional[str] = Field(
|
||||
default=None,
|
||||
description="Name of parent policy to inherit from.",
|
||||
)
|
||||
description: Optional[str] = Field(
|
||||
default=None,
|
||||
description="Human-readable description of the policy.",
|
||||
)
|
||||
guardrails_add: Optional[List[str]] = Field(
|
||||
default=None,
|
||||
description="List of guardrail names to add.",
|
||||
)
|
||||
guardrails_remove: Optional[List[str]] = Field(
|
||||
default=None,
|
||||
description="List of guardrail names to remove (from inherited).",
|
||||
)
|
||||
condition: Optional[PolicyConditionRequest] = Field(
|
||||
default=None,
|
||||
description="Condition for when this policy applies.",
|
||||
)
|
||||
|
||||
|
||||
class PolicyDBResponse(BaseModel):
|
||||
"""Response for a policy from the database."""
|
||||
|
||||
policy_id: str = Field(description="Unique ID of the policy.")
|
||||
policy_name: str = Field(description="Name of the policy.")
|
||||
inherit: Optional[str] = Field(default=None, description="Parent policy name.")
|
||||
description: Optional[str] = Field(default=None, description="Policy description.")
|
||||
guardrails_add: List[str] = Field(
|
||||
default_factory=list, description="Guardrails to add."
|
||||
)
|
||||
guardrails_remove: List[str] = Field(
|
||||
default_factory=list, description="Guardrails to remove."
|
||||
)
|
||||
condition: Optional[Dict[str, Any]] = Field(
|
||||
default=None, description="Policy condition."
|
||||
)
|
||||
created_at: Optional[datetime] = Field(
|
||||
default=None, description="When the policy was created."
|
||||
)
|
||||
updated_at: Optional[datetime] = Field(
|
||||
default=None, description="When the policy was last updated."
|
||||
)
|
||||
created_by: Optional[str] = Field(default=None, description="Who created the policy.")
|
||||
updated_by: Optional[str] = Field(
|
||||
default=None, description="Who last updated the policy."
|
||||
)
|
||||
|
||||
|
||||
class PolicyListDBResponse(BaseModel):
|
||||
"""Response for listing policies from the database."""
|
||||
|
||||
policies: List[PolicyDBResponse] = Field(
|
||||
default_factory=list, description="List of policies."
|
||||
)
|
||||
total_count: int = Field(default=0, description="Total number of policies.")
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# Policy Attachment CRUD Types
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class PolicyAttachmentCreateRequest(BaseModel):
|
||||
"""Request body for creating a policy attachment."""
|
||||
|
||||
policy_name: str = Field(description="Name of the policy to attach.")
|
||||
scope: Optional[str] = Field(
|
||||
default=None,
|
||||
description="Use '*' for global scope (applies to all requests).",
|
||||
)
|
||||
teams: Optional[List[str]] = Field(
|
||||
default=None,
|
||||
description="Team aliases or patterns this attachment applies to.",
|
||||
)
|
||||
keys: Optional[List[str]] = Field(
|
||||
default=None,
|
||||
description="Key aliases or patterns this attachment applies to.",
|
||||
)
|
||||
models: Optional[List[str]] = Field(
|
||||
default=None,
|
||||
description="Model names or patterns this attachment applies to.",
|
||||
)
|
||||
|
||||
|
||||
class PolicyAttachmentDBResponse(BaseModel):
|
||||
"""Response for a policy attachment from the database."""
|
||||
|
||||
attachment_id: str = Field(description="Unique ID of the attachment.")
|
||||
policy_name: str = Field(description="Name of the attached policy.")
|
||||
scope: Optional[str] = Field(default=None, description="Scope of the attachment.")
|
||||
teams: List[str] = Field(default_factory=list, description="Team patterns.")
|
||||
keys: List[str] = Field(default_factory=list, description="Key patterns.")
|
||||
models: List[str] = Field(default_factory=list, description="Model patterns.")
|
||||
created_at: Optional[datetime] = Field(
|
||||
default=None, description="When the attachment was created."
|
||||
)
|
||||
updated_at: Optional[datetime] = Field(
|
||||
default=None, description="When the attachment was last updated."
|
||||
)
|
||||
created_by: Optional[str] = Field(
|
||||
default=None, description="Who created the attachment."
|
||||
)
|
||||
updated_by: Optional[str] = Field(
|
||||
default=None, description="Who last updated the attachment."
|
||||
)
|
||||
|
||||
|
||||
class PolicyAttachmentListResponse(BaseModel):
|
||||
"""Response for listing policy attachments."""
|
||||
|
||||
attachments: List[PolicyAttachmentDBResponse] = Field(
|
||||
default_factory=list, description="List of policy attachments."
|
||||
)
|
||||
total_count: int = Field(default=0, description="Total number of attachments.")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue