diff --git a/litellm/litellm_core_utils/logging_callback_manager.py b/litellm/litellm_core_utils/logging_callback_manager.py index b7adda3a9a4..b8bb50fbe03 100644 --- a/litellm/litellm_core_utils/logging_callback_manager.py +++ b/litellm/litellm_core_utils/logging_callback_manager.py @@ -61,7 +61,9 @@ class LoggingCallbackManager: callback=callback, parent_list=litellm.service_callback ) - def add_litellm_callback(self, callback: Union[CustomLogger, str, Callable]): + def add_litellm_callback( + self, callback: Union[CustomLogger, str, Callable] + ) -> None: """ Add a callback to litellm.callbacks @@ -73,7 +75,7 @@ class LoggingCallbackManager: def add_litellm_success_callback( self, callback: Union[CustomLogger, str, Callable] - ): + ) -> None: """ Add a success callback to `litellm.success_callback`. Auto-routes async callbacks to litellm._async_success_callback. diff --git a/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/__init__.py index 774a0334072..9cbb7e6b5a1 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/__init__.py @@ -1,8 +1,12 @@ """Cisco AI Defense Guardrail Integration for LiteLLM.""" -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any, Dict, List, Optional from litellm.types.guardrails import SupportedGuardrailIntegrations +from litellm.types.proxy.guardrails.guardrail_hooks.cisco_ai_defense import ( + CiscoAIDefenseGuardrailConfigModelOptionalParams, + CiscoAIDefenseRule, +) from .cisco_ai_defense import ( CiscoAIDefenseGuardrail, @@ -14,47 +18,34 @@ if TYPE_CHECKING: from litellm.types.guardrails import Guardrail, LitellmParams -def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"): +def initialize_guardrail( + litellm_params: "LitellmParams", guardrail: "Guardrail" +) -> CiscoAIDefenseGuardrail: import litellm guardrail_name = guardrail.get("guardrail_name") if not guardrail_name: raise ValueError("Cisco AI Defense: guardrail_name is required") - optional_params = getattr(litellm_params, "optional_params", None) + optional_params = _resolve_optional_params(litellm_params) + # fmt: off + enabled_rules = _dump_rules(optional_params.enabled_rules) # any-ok: Pydantic model_dump returns dict[str, Any] across the typed/untyped boundary + # fmt: on _callback = CiscoAIDefenseGuardrail( guardrail_name=guardrail_name, api_key=litellm_params.api_key, api_base=litellm_params.api_base, - inspection_type=_get_optional_value( - litellm_params, optional_params, "inspection_type" - ), - inspect_path=_get_optional_value( - litellm_params, optional_params, "inspect_path" - ), - enabled_rules=_get_optional_value( - litellm_params, optional_params, "enabled_rules" - ), - integration_profile_id=_get_optional_value( - litellm_params, optional_params, "integration_profile_id" - ), - integration_profile_version=_get_optional_value( - litellm_params, optional_params, "integration_profile_version" - ), - integration_tenant_id=_get_optional_value( - litellm_params, optional_params, "integration_tenant_id" - ), - integration_type=_get_optional_value( - litellm_params, optional_params, "integration_type" - ), - on_flagged_action=_get_optional_value( - litellm_params, optional_params, "on_flagged_action" - ), - fallback_on_error=_get_optional_value( - litellm_params, optional_params, "fallback_on_error" - ), - timeout=_get_optional_value(litellm_params, optional_params, "timeout"), + inspection_type=optional_params.inspection_type, + inspect_path=optional_params.inspect_path, + enabled_rules=enabled_rules, # any-ok: Pydantic dict[str, Any] propagates from _dump_rules across the typed/untyped boundary + integration_profile_id=optional_params.integration_profile_id, + integration_profile_version=optional_params.integration_profile_version, + integration_tenant_id=optional_params.integration_tenant_id, + integration_type=optional_params.integration_type, + on_flagged_action=optional_params.on_flagged_action, + fallback_on_error=optional_params.fallback_on_error, + timeout=optional_params.timeout, event_hook=litellm_params.mode, default_on=litellm_params.default_on or False, ) @@ -66,26 +57,46 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" return _callback -def _get_optional_value(litellm_params, optional_params, attribute_name): - """Resolve Cisco optional params without inheriting sibling defaults.""" - if optional_params is not None: - if isinstance(optional_params, dict): - if attribute_name in optional_params: - return optional_params[attribute_name] - else: - nested_fields_set = getattr(optional_params, "model_fields_set", None) - if nested_fields_set is None or attribute_name in nested_fields_set: - value = getattr(optional_params, attribute_name, None) - if value is not None: - return value +def _resolve_optional_params( + litellm_params: "LitellmParams", +) -> CiscoAIDefenseGuardrailConfigModelOptionalParams: + """Resolve Cisco optional params by merging explicitly-set flat fields on + ``litellm_params`` with the nested ``optional_params``. Nested explicit + values override flat ones (matching the original lookup priority); only + fields the user actually set are forwarded so sibling guardrails' defaults + inherited via MRO on ``LitellmParams`` are not picked up.""" + return CiscoAIDefenseGuardrailConfigModelOptionalParams.model_validate( + _explicit_overrides(litellm_params) + ) - if litellm_params is None: + +def _explicit_overrides(litellm_params: "LitellmParams") -> Dict[str, object]: + """Collect Cisco optional-param overrides the user actually set, suitable + for Pydantic validation. Crosses a typed/untyped boundary because + ``model_dump`` and dynamic dict access on untyped sources return ``Any``.""" + cisco_fields = CiscoAIDefenseGuardrailConfigModelOptionalParams.model_fields.keys() + flat_keys = litellm_params.model_fields_set & cisco_fields + # fmt: off + merged: Dict[str, object] = dict(litellm_params.model_dump(include=flat_keys)) # any-ok: Pydantic model_dump returns dict[str, Any] across the typed/untyped boundary + nested = litellm_params.optional_params + if isinstance(nested, CiscoAIDefenseGuardrailConfigModelOptionalParams): + merged.update(nested.model_dump(exclude_unset=True, exclude_none=True)) # any-ok: Pydantic model_dump returns dict[str, Any] across the typed/untyped boundary + elif isinstance(nested, dict): + for key, value in nested.items(): + if key in cisco_fields: + merged[key] = value # any-ok: nested optional_params dict is untyped at the user-input boundary + # fmt: on + return merged + + +def _dump_rules( + rules: Optional[List[CiscoAIDefenseRule]], +) -> Optional[List[Dict[str, Any]]]: + if not rules: return None - # Only accept flattened values the caller explicitly set. - fields_set = getattr(litellm_params, "model_fields_set", None) - if fields_set is None or attribute_name not in fields_set: - return None - return getattr(litellm_params, attribute_name, None) + # fmt: off + return [rule.model_dump() for rule in rules] # any-ok: Pydantic model_dump returns dict[str, Any] across the typed/untyped boundary + # fmt: on guardrail_initializer_registry = { diff --git a/litellm/proxy/guardrails/guardrail_hooks/ovalix/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/ovalix/__init__.py index b73572e4ed5..083eda49525 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/ovalix/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/ovalix/__init__.py @@ -10,23 +10,19 @@ if TYPE_CHECKING: from litellm.types.guardrails import Guardrail, LitellmParams -def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"): +def initialize_guardrail( + litellm_params: "LitellmParams", guardrail: "Guardrail" +) -> OvalixGuardrail: """Create and register an Ovalix guardrail callback from proxy config.""" import litellm - tracker_api_base = getattr(litellm_params, "tracker_api_base", None) - tracker_api_key = getattr(litellm_params, "tracker_api_key", None) - application_id = getattr(litellm_params, "application_id", None) - pre_checkpoint_id = getattr(litellm_params, "pre_checkpoint_id", None) - post_checkpoint_id = getattr(litellm_params, "post_checkpoint_id", None) - _ovalix_callback = OvalixGuardrail( guardrail_name=guardrail.get("guardrail_name", ""), - tracker_api_base=tracker_api_base, - tracker_api_key=tracker_api_key, - application_id=application_id, - pre_checkpoint_id=pre_checkpoint_id, - post_checkpoint_id=post_checkpoint_id, + tracker_api_base=litellm_params.tracker_api_base, + tracker_api_key=litellm_params.tracker_api_key, + application_id=litellm_params.application_id, + pre_checkpoint_id=litellm_params.pre_checkpoint_id, + post_checkpoint_id=litellm_params.post_checkpoint_id, event_hook=litellm_params.mode, default_on=litellm_params.default_on, )