mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
[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:
parent
674d4b4cab
commit
a978680714
6 changed files with 227 additions and 10 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue