diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index d7b37f8a3fe..e64fbe9084e 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -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( diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index b5ae8437d86..8aa8ba0570b 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -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 diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 1fc8bc377e7..e57820e1248 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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: diff --git a/ui/litellm-dashboard/src/components/guardrails/content_filter/ContentFilterManager.test.tsx b/ui/litellm-dashboard/src/components/guardrails/content_filter/ContentFilterManager.test.tsx new file mode 100644 index 00000000000..6d879e1c544 --- /dev/null +++ b/ui/litellm-dashboard/src/components/guardrails/content_filter/ContentFilterManager.test.tsx @@ -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: () =>
Mock Content Filter Configuration
+})); + +vi.mock("./ContentFilterDisplay", () => ({ + default: () =>
Mock Content Filter Display
+})); + +vi.mock("antd", () => ({ + Divider: ({ children }: { children: React.ReactNode }) =>
{children}
+})); + +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( + + ); + + // 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(); + }); +}); + diff --git a/ui/litellm-dashboard/src/components/guardrails/content_filter/ContentFilterManager.tsx b/ui/litellm-dashboard/src/components/guardrails/content_filter/ContentFilterManager.tsx index 65a27388a22..f81aa478b93 100644 --- a/ui/litellm-dashboard/src/components/guardrails/content_filter/ContentFilterManager.tsx +++ b/ui/litellm-dashboard/src/components/guardrails/content_filter/ContentFilterManager.tsx @@ -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 = ({ @@ -46,9 +47,12 @@ const ContentFilterManager: React.FC = ({ isEditing, accessToken, onDataChange, + onUnsavedChanges, }) => { const [selectedPatterns, setSelectedPatterns] = useState([]); const [blockedWords, setBlockedWords] = useState([]); + const [originalPatterns, setOriginalPatterns] = useState([]); + const [originalBlockedWords, setOriginalBlockedWords] = useState([]); // Load data from guardrail on mount or when guardrailData changes useEffect(() => { @@ -62,8 +66,10 @@ const ContentFilterManager: React.FC = ({ 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 = ({ description: w.description, })); setBlockedWords(words); + setOriginalBlockedWords(words); } else { setBlockedWords([]); + setOriginalBlockedWords([]); } }, [guardrailData]); @@ -86,6 +94,19 @@ const ContentFilterManager: React.FC = ({ } }, [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 = ({ return ( <> Content Filter Configuration + {hasUnsavedChanges && ( +
+

+ ⚠️ You have unsaved changes to patterns or keywords. Remember to click "Save Changes" at the bottom. +

+
+ )}
{guardrailSettings && guardrailSettings.content_filter_settings && ( = ({ guardrailId, onClose, }; } | null>(null); const [copiedStates, setCopiedStates] = useState>({}); + 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 = ({ 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 = ({ guardrailId, onClose, guardrailSettings={guardrailSettings} isEditing={true} accessToken={accessToken} - onDataChange={(patterns, blockedWords) => { - contentFilterDataRef.current = { patterns, blockedWords }; - }} + onDataChange={handleContentFilterDataChange} + onUnsavedChanges={setHasUnsavedContentFilterChanges} /> Provider Settings @@ -608,7 +615,10 @@ const GuardrailInfoView: React.FC = ({ guardrailId, onClose,
- + Save Changes