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: () =>
+ ⚠️ You have unsaved changes to patterns or keywords. Remember to click "Save Changes" at the bottom. +
+