[UI] Guardrails - allow updating guardrails through UI. Ensure litellm_params actually get updated in memory (#16384)

* fix safe dumps

* add patterns.json

* add PrebuiltPattern

* add test patterns

* fix edit and view

* fix backend handling

* fix CF ui edit

* fix init

* add _has_guardrail_params_changed

* fix sync_guardrail_from_db

* fix _has_guardrail_params_changed

* fix patch_guardrail

* add unsaved change check

* fix ContentFilterManager
This commit is contained in:
Ishaan Jaff 2025-11-07 18:15:59 -08:00 • committed by GitHub
parent 674d4b4cab
commit a978680714
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 227 additions and 10 deletions

View file

@ -564,8 +564,7 @@ async def patch_guardrail(guardrail_id: str, request: PatchGuardrailRequest):
guardrail_name = result.get("guardrail_name", "Unknown")
try:
IN_MEMORY_GUARDRAIL_HANDLER.update_in_memory_guardrail(
guardrail_id=guardrail_id,
IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(
guardrail=guardrail,
)
verbose_proxy_logger.info(

View file

@ -2,12 +2,12 @@
import importlib
import os
from litellm._uuid import uuid
from datetime import datetime, timezone
from typing import Dict, List, Optional, Type, cast
import litellm
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.proxy.utils import PrismaClient
@ -525,9 +525,19 @@ class InMemoryGuardrailHandler:
def delete_in_memory_guardrail(self, guardrail_id: str) -> None:
"""
Delete a guardrail in memory
Delete a guardrail in memory and remove from litellm callbacks.
"""
# Remove from in-memory storage
self.IN_MEMORY_GUARDRAILS.pop(guardrail_id, None)
# Remove the callback from litellm.callbacks
custom_guardrail_callback = self.guardrail_id_to_custom_guardrail.pop(guardrail_id, None)
if custom_guardrail_callback:
litellm.logging_callback_manager.remove_callback_from_list_by_object(
callback_list=litellm.callbacks,
obj=custom_guardrail_callback,
require_self=False
)
def list_in_memory_guardrails(self) -> List[Guardrail]:
"""
@ -541,6 +551,99 @@ class InMemoryGuardrailHandler:
"""
return self.IN_MEMORY_GUARDRAILS.get(guardrail_id)
def _has_guardrail_params_changed(
self, guardrail_id: str, new_guardrail: Guardrail
) -> bool:
"""
Check if guardrail params or name have changed compared to in-memory version.
Returns True if params/name changed or guardrail doesn't exist in memory.
"""
existing = self.IN_MEMORY_GUARDRAILS.get(guardrail_id)
if existing is None:
return True
# Compare guardrail_name
if existing.get("guardrail_name") != new_guardrail.get("guardrail_name"):
return True
# Compare litellm_params
existing_params = existing.get("litellm_params")
new_params = new_guardrail.get("litellm_params")
# Convert to dicts for comparison
existing_dict = (
existing_params.model_dump()
if isinstance(existing_params, LitellmParams)
else existing_params
)
new_dict = (
new_params.model_dump()
if isinstance(new_params, LitellmParams)
else new_params
)
# Compare and identify specific differences
changed_fields = {}
all_keys = set(existing_dict.keys()) | set(new_dict.keys())
for key in all_keys:
old_val = existing_dict.get(key)
new_val = new_dict.get(key)
if old_val != new_val:
changed_fields[key] = {"old": old_val, "new": new_val}
# Log differences if any found
if changed_fields:
verbose_proxy_logger.debug(
f"Guardrail params changed. Differences: {changed_fields}"
)
# Return True if any fields changed
return len(changed_fields) > 0
def reinitialize_guardrail(
self, guardrail: Guardrail, config_file_path: Optional[str] = None
) -> Optional[Guardrail]:
"""
Force re-initialization of a guardrail even if it exists in memory.
Removes old callback from litellm.callbacks and creates fresh instance.
"""
guardrail_id = guardrail.get("guardrail_id")
if not guardrail_id:
verbose_proxy_logger.error("Cannot reinitialize guardrail without guardrail_id")
return None
# Remove from memory if exists (also removes from callbacks)
if guardrail_id in self.IN_MEMORY_GUARDRAILS:
self.delete_in_memory_guardrail(guardrail_id)
# Initialize fresh (will add new callback to litellm.callbacks)
return self.initialize_guardrail(
guardrail=guardrail, config_file_path=config_file_path
)
def sync_guardrail_from_db(
self, guardrail: Guardrail, config_file_path: Optional[str] = None
) -> Optional[Guardrail]:
"""
Sync a guardrail from DB - initializes if new, re-initializes if changed.
This is the method to call during DB polling.
"""
guardrail_id = guardrail.get("guardrail_id")
if not guardrail_id:
verbose_proxy_logger.error("Cannot sync guardrail without guardrail_id")
return None
if self._has_guardrail_params_changed(guardrail_id, guardrail):
guardrail_name = guardrail.get("guardrail_name", "Unknown")
verbose_proxy_logger.info(
f"Guardrail '{guardrail_name}' (ID: {guardrail_id}) params changed, re-initializing..."
)
return self.reinitialize_guardrail(
guardrail=guardrail, config_file_path=config_file_path
)
return self.IN_MEMORY_GUARDRAILS.get(guardrail_id)
########################################################
# In Memory Guardrail Handler for LiteLLM Proxy

View file

@ -3595,7 +3595,7 @@ class ProxyConfig:
"guardrails from the DB %s", str(guardrails_in_db)
)
for guardrail in guardrails_in_db:
IN_MEMORY_GUARDRAIL_HANDLER.initialize_guardrail(
IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(
guardrail=cast(Guardrail, guardrail),
)
except Exception as e:

View file

@ -0,0 +1,77 @@
import { describe, it, expect, vi, beforeEach } from "vitest";
import { render, screen, waitFor } from "@testing-library/react";
import ContentFilterManager from "./ContentFilterManager";
import React from "react";
vi.mock("./ContentFilterConfiguration", () => ({
default: () => <div data-testid="content-filter-config">Mock Content Filter Configuration</div>
}));
vi.mock("./ContentFilterDisplay", () => ({
default: () => <div data-testid="content-filter-display">Mock Content Filter Display</div>
}));
vi.mock("antd", () => ({
Divider: ({ children }: { children: React.ReactNode }) => <div>{children}</div>
}));
describe("ContentFilterManager - Unsaved Changes Detection", () => {
beforeEach(() => {
vi.clearAllMocks();
});
it("should call onUnsavedChanges with false when component initializes with matching data", async () => {
/**
* Tests that the ContentFilterManager correctly initializes the unsaved changes
* detection and calls onUnsavedChanges(false) when the current state matches
* the original loaded state (no changes yet).
*/
const mockOnUnsavedChanges = vi.fn();
const mockOnDataChange = vi.fn();
const guardrailData = {
litellm_params: {
guardrail: "litellm_content_filter",
patterns: [
{ pattern_type: "prebuilt", pattern_name: "email", action: "BLOCK" }
],
blocked_words: [
{ keyword: "test", action: "BLOCK", description: null }
]
}
};
const guardrailSettings = {
content_filter_settings: {
prebuilt_patterns: [],
pattern_categories: ["PII"],
supported_actions: ["BLOCK", "MASK"]
}
};
render(
<ContentFilterManager
guardrailData={guardrailData}
guardrailSettings={guardrailSettings}
isEditing={true}
accessToken="test-token"
onDataChange={mockOnDataChange}
onUnsavedChanges={mockOnUnsavedChanges}
/>
);
// Wait for component to render in edit mode
await waitFor(() => {
expect(screen.getByTestId("content-filter-config")).toBeInTheDocument();
});
// Verify onUnsavedChanges was called with false (no changes initially)
await waitFor(() => {
expect(mockOnUnsavedChanges).toHaveBeenCalledWith(false);
});
// Verify onDataChange was called with initial data
expect(mockOnDataChange).toHaveBeenCalled();
});
});

View file

@ -38,6 +38,7 @@ interface ContentFilterManagerProps {
isEditing: boolean;
accessToken: string | null;
onDataChange?: (patterns: Pattern[], blockedWords: BlockedWord[]) => void;
onUnsavedChanges?: (hasChanges: boolean) => void;
}
const ContentFilterManager: React.FC<ContentFilterManagerProps> = ({
@ -46,9 +47,12 @@ const ContentFilterManager: React.FC<ContentFilterManagerProps> = ({
isEditing,
accessToken,
onDataChange,
onUnsavedChanges,
}) => {
const [selectedPatterns, setSelectedPatterns] = useState<Pattern[]>([]);
const [blockedWords, setBlockedWords] = useState<BlockedWord[]>([]);
const [originalPatterns, setOriginalPatterns] = useState<Pattern[]>([]);
const [originalBlockedWords, setOriginalBlockedWords] = useState<BlockedWord[]>([]);
// Load data from guardrail on mount or when guardrailData changes
useEffect(() => {
@ -62,8 +66,10 @@ const ContentFilterManager: React.FC<ContentFilterManagerProps> = ({
action: p.action || "BLOCK",
}));
setSelectedPatterns(patterns);
setOriginalPatterns(patterns);
} else {
setSelectedPatterns([]);
setOriginalPatterns([]);
}
if (guardrailData?.litellm_params?.blocked_words) {
@ -74,8 +80,10 @@ const ContentFilterManager: React.FC<ContentFilterManagerProps> = ({
description: w.description,
}));
setBlockedWords(words);
setOriginalBlockedWords(words);
} else {
setBlockedWords([]);
setOriginalBlockedWords([]);
}
}, [guardrailData]);
@ -86,6 +94,19 @@ const ContentFilterManager: React.FC<ContentFilterManagerProps> = ({
}
}, [selectedPatterns, blockedWords, onDataChange]);
// Detect unsaved changes
const hasUnsavedChanges = React.useMemo(() => {
const hasPatternChanges = JSON.stringify(selectedPatterns) !== JSON.stringify(originalPatterns);
const hasWordChanges = JSON.stringify(blockedWords) !== JSON.stringify(originalBlockedWords);
return hasPatternChanges || hasWordChanges;
}, [selectedPatterns, blockedWords, originalPatterns, originalBlockedWords]);
useEffect(() => {
if (isEditing && onUnsavedChanges) {
onUnsavedChanges(hasUnsavedChanges);
}
}, [hasUnsavedChanges, isEditing, onUnsavedChanges]);
// Check if this is a content filter guardrail
if (guardrailData?.litellm_params?.guardrail !== "litellm_content_filter") {
return null;
@ -106,6 +127,13 @@ const ContentFilterManager: React.FC<ContentFilterManagerProps> = ({
return (
<>
<Divider orientation="left">Content Filter Configuration</Divider>
{hasUnsavedChanges && (
<div className="mb-4 px-4 py-3 bg-yellow-50 border border-yellow-200 rounded-md">
<p className="text-sm text-yellow-800 font-medium">
⚠️ You have unsaved changes to patterns or keywords. Remember to click "Save Changes" at the bottom.
</p>
</div>
)}
<div className="mb-6">
{guardrailSettings && guardrailSettings.content_filter_settings && (
<ContentFilterConfiguration

View file

@ -1,4 +1,4 @@
import React, { useState, useEffect } from "react";
import React, { useState, useEffect, useCallback } from "react";
import {
Card,
Title,
@ -82,12 +82,19 @@ const GuardrailInfoView: React.FC<GuardrailInfoProps> = ({ guardrailId, onClose,
};
} | null>(null);
const [copiedStates, setCopiedStates] = useState<Record<string, boolean>>({});
const [hasUnsavedContentFilterChanges, setHasUnsavedContentFilterChanges] = useState(false);
// Content Filter data ref (managed by ContentFilterManager)
const contentFilterDataRef = React.useRef<{ patterns: any[]; blockedWords: any[] }>({
patterns: [],
blockedWords: [],
});
// Memoize onDataChange callback to prevent unnecessary re-renders
const handleContentFilterDataChange = useCallback((patterns: any[], blockedWords: any[]) => {
contentFilterDataRef.current = { patterns, blockedWords };
}, []);
const fetchGuardrailInfo = async () => {
try {
setLoading(true);
@ -329,6 +336,7 @@ const GuardrailInfoView: React.FC<GuardrailInfoProps> = ({ guardrailId, onClose,
await updateGuardrailCall(accessToken, guardrailId, updateData);
NotificationsManager.success("Guardrail updated successfully");
setHasUnsavedContentFilterChanges(false);
fetchGuardrailInfo();
setIsEditing(false);
} catch (error) {
@ -561,9 +569,8 @@ const GuardrailInfoView: React.FC<GuardrailInfoProps> = ({ guardrailId, onClose,
guardrailSettings={guardrailSettings}
isEditing={true}
accessToken={accessToken}
onDataChange={(patterns, blockedWords) => {
contentFilterDataRef.current = { patterns, blockedWords };
}}
onDataChange={handleContentFilterDataChange}
onUnsavedChanges={setHasUnsavedContentFilterChanges}
/>
<Divider orientation="left">Provider Settings</Divider>
@ -608,7 +615,10 @@ const GuardrailInfoView: React.FC<GuardrailInfoProps> = ({ guardrailId, onClose,
</Form.Item>
<div className="flex justify-end gap-2 mt-6">
<Button onClick={() => setIsEditing(false)}>Cancel</Button>
<Button onClick={() => {
setIsEditing(false);
setHasUnsavedContentFilterChanges(false);
}}>Cancel</Button>
<TremorButton>Save Changes</TremorButton>
</div>
</Form>