Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions litellm/proxy/guardrails/guardrail_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -487,6 +487,7 @@ def initialize_guardrail(
guardrail_id=guardrail.get("guardrail_id"),
guardrail_name=guardrail["guardrail_name"],
litellm_params=litellm_params,
guardrail_info=guardrail.get("guardrail_info"),
)

# store references to the guardrail in memory
Expand Down
80 changes: 50 additions & 30 deletions litellm/proxy/guardrails/usage_endpoints.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,8 @@

from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER
from litellm.types.guardrails import LitellmParams

router = APIRouter()

Expand Down Expand Up @@ -135,17 +137,29 @@ def _chart_from_metrics(metrics: Any) -> List[Dict[str, Any]]:
]


def _get_guardrail_field(g: Any, field: str) -> Any:
"""Read `field` off a guardrail (Prisma row attr or dict/TypedDict key)."""
if isinstance(g, dict):
return g.get(field)
return getattr(g, field, None)


def _get_guardrail_attrs(g: Any) -> tuple[Any, str]:
"""Get (guardrail_id, display_name) from guardrail - handles Prisma model or dict."""
gid = getattr(g, "guardrail_id", None) or (
g.get("guardrail_id") if isinstance(g, dict) else None
)
name = getattr(g, "guardrail_name", None) or (
g.get("guardrail_name") if isinstance(g, dict) else None
)
gid = _get_guardrail_field(g, "guardrail_id")
name = _get_guardrail_field(g, "guardrail_name")
return gid, (name or gid or "")


def _to_dict(value: Any) -> Dict[str, Any]:
"""Coerce a LitellmParams / dict / None into a plain dict."""
if isinstance(value, LitellmParams):
return value.model_dump(exclude_none=True)
if isinstance(value, dict):
return value
return {}
Comment on lines +154 to +160

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 _to_dict hardcodes LitellmParams type

_to_dict special-cases LitellmParams specifically, which means any other Pydantic model stored in these fields (e.g. a custom subclass or a future guardrail_info model) silently returns {} rather than its actual data. Using BaseModel from Pydantic as the isinstance check would be more resilient.

Suggested change
def _to_dict(value: Any) -> Dict[str, Any]:
"""Coerce a LitellmParams / dict / None into a plain dict."""
if isinstance(value, LitellmParams):
return value.model_dump(exclude_none=True)
if isinstance(value, dict):
return value
return {}
def _to_dict(value: Any) -> Dict[str, Any]:
"""Coerce a Pydantic BaseModel / dict / None into a plain dict."""
from pydantic import BaseModel as _BaseModel
if isinstance(value, _BaseModel):
return value.model_dump(exclude_none=True)
if isinstance(value, dict):
return value
return {}



def _guardrail_overview_rows(
guardrails: Any,
agg: Dict[str, Dict[str, Any]],
Expand All @@ -165,13 +179,9 @@ def _guardrail_overview_rows(
break
req, blocked = a["requests"], a["blocked"]
fail_rate = (100.0 * blocked / req) if req else 0.0
litellm_params = (
(g.litellm_params or {}) if isinstance(g.litellm_params, dict) else {}
)
litellm_params = _to_dict(_get_guardrail_field(g, "litellm_params"))
provider = str(litellm_params.get("guardrail", "Unknown"))
guardrail_info = (
(g.guardrail_info or {}) if isinstance(g.guardrail_info, dict) else {}
)
guardrail_info = _to_dict(_get_guardrail_field(g, "guardrail_info"))
gtype = str(guardrail_info.get("type", "Guardrail"))
prev_fail = 0.0
for k in lookup_keys:
Expand Down Expand Up @@ -271,8 +281,19 @@ async def guardrails_usage_overview(
start = start_date or (now - timedelta(days=7)).strftime("%Y-%m-%d")

try:
# Guardrails from DB
guardrails = await prisma_client.db.litellm_guardrailstable.find_many()
# Guardrails from DB unioned with YAML/in-memory guardrails (deduped by id).
db_guardrails = await prisma_client.db.litellm_guardrailstable.find_many()
seen_ids = {
getattr(g, "guardrail_id", None)
for g in db_guardrails
if getattr(g, "guardrail_id", None)
}
in_memory_guardrails = [
g
for g in IN_MEMORY_GUARDRAIL_HANDLER.list_in_memory_guardrails()
if g.get("guardrail_id") not in seen_ids
]
guardrails: List[Any] = list(db_guardrails) + in_memory_guardrails

# Daily metrics in range
metrics = await prisma_client.db.litellm_dailyguardrailmetrics.find_many(
Expand Down Expand Up @@ -338,15 +359,18 @@ async def guardrails_usage_detail(
guardrail = await prisma_client.db.litellm_guardrailstable.find_unique(
where={"guardrail_id": guardrail_id}
)
if not guardrail:
if guardrail is None:
# YAML-defined guardrails live only in the in-memory registry.
guardrail = IN_MEMORY_GUARDRAIL_HANDLER.get_guardrail_by_id(
guardrail_id=guardrail_id
)
if guardrail is None:
from fastapi import HTTPException

raise HTTPException(status_code=404, detail="Guardrail not found")

# Metrics are keyed by logical name (from spend log metadata), not UUID
logical_id = getattr(guardrail, "guardrail_name", None) or (
guardrail.get("guardrail_name") if isinstance(guardrail, dict) else None
)
logical_id = _get_guardrail_field(guardrail, "guardrail_name")
metric_ids = [i for i in (logical_id, guardrail_id) if i]

metrics = await prisma_client.db.litellm_dailyguardrailmetrics.find_many(
Expand Down Expand Up @@ -383,17 +407,9 @@ async def guardrails_usage_detail(
{"date": d, "passed": v["passed"], "blocked": v["blocked"], "score": None}
for d, v in sorted(ts_by_date.items())
]
_litellm_params = getattr(guardrail, "litellm_params", None) or (
guardrail.get("litellm_params") if isinstance(guardrail, dict) else None
)
litellm_params = _litellm_params if isinstance(_litellm_params, dict) else {}
_guardrail_info = getattr(guardrail, "guardrail_info", None) or (
guardrail.get("guardrail_info") if isinstance(guardrail, dict) else None
)
guardrail_info = _guardrail_info if isinstance(_guardrail_info, dict) else {}
_guardrail_name = getattr(guardrail, "guardrail_name", None) or (
guardrail.get("guardrail_name") if isinstance(guardrail, dict) else None
)
litellm_params = _to_dict(_get_guardrail_field(guardrail, "litellm_params"))
guardrail_info = _to_dict(_get_guardrail_field(guardrail, "guardrail_info"))
_guardrail_name = _get_guardrail_field(guardrail, "guardrail_name")

return UsageDetailResponse(
guardrail_id=guardrail_id,
Expand Down Expand Up @@ -577,8 +593,12 @@ async def guardrails_usage_logs(
guardrail = await prisma_client.db.litellm_guardrailstable.find_unique(
where={"guardrail_id": guardrail_id}
)
if guardrail is None:
guardrail = IN_MEMORY_GUARDRAIL_HANDLER.get_guardrail_by_id(
guardrail_id=guardrail_id
)
if guardrail:
logical_name = getattr(guardrail, "guardrail_name", None)
logical_name = _get_guardrail_field(guardrail, "guardrail_name")
if logical_name and logical_name not in effective_guardrail_ids:
effective_guardrail_ids.append(logical_name)

Expand Down
Loading
Loading