diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 2407c0032a3..b1291d663bb 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -943,7 +943,7 @@ class CustomGuardrail(CustomLogger): result: object, call_type: str, ) -> tuple[dict, object]: # mutable-ok: CustomLogger.async_logging_hook contract - """logging_only: run apply_guardrail on copies of the logged request/response and record the verdict.""" + """logging_only: scan copies of the logged request and/or response according to logging_only_scope.""" from litellm.llms import get_guardrail_translation_mapping if not self.uses_apply_guardrail_interface(): diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index 6527f6be50b..0258b291d21 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -558,8 +558,13 @@ class InMemoryGuardrailHandler: config_file_path=config_file_path, llm_router=llm_router, ) - for custom_guardrail_callback in created_callbacks: - _configure_callback_scoping(custom_guardrail_callback, guardrail["guardrail_name"], litellm_params) + try: + for custom_guardrail_callback in created_callbacks: + _configure_callback_scoping(custom_guardrail_callback, guardrail["guardrail_name"], litellm_params) + except Exception: + for custom_guardrail_callback in created_callbacks: + litellm.logging_callback_manager.remove_callback_from_all_lists(custom_guardrail_callback) + raise parsed_guardrail: Final = Guardrail( guardrail_id=guardrail.get("guardrail_id"), diff --git a/tests/integration/observability/test_guardrail_effects.py b/tests/integration/observability/test_guardrail_effects.py index 4e9691bd4e9..ad487b1f718 100644 --- a/tests/integration/observability/test_guardrail_effects.py +++ b/tests/integration/observability/test_guardrail_effects.py @@ -1391,3 +1391,62 @@ def test_logging_only_scope_observes_only_the_configured_direction_without_block return_last_on_timeout=True, ) assert detail["requestsEvaluated"] == len(scanned_directions), detail + + +def test_logging_only_scope_without_logging_only_mode_is_skipped_at_load_and_never_blocks( + gateway: Gateway, tmp_path: Path +) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + prompt: Final = "synthetic invalid-scope prompt " + identity + reply: Final = "synthetic invalid-scope reply " + identity + + def guardrail(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api" + return Reply(body=json.dumps({"action": "BLOCKED", "blocked_reason": "synthetic policy denial"}).encode()) + + def provider(request: Request) -> Reply: + assert request.target == "/v1/chat/completions" + assert json.loads(request.body)["messages"] == [{"role": "user", "content": prompt}] + return Reply( + body=json.dumps( + { + "id": identity, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": reply}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 9, "completion_tokens": 5, "total_tokens": 14}, + } + ).encode() + ) + + with wire_server(guardrail) as policy, wire_server(provider) as upstream: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["guardrails"] = [ + { + "guardrail_name": identity, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "logging_only_scope": "input", + "default_on": True, + "api_base": policy.url, + "api_key": "synthetic-guardrail-key", + }, + } + ] + path: Final = tmp_path / "invalid-scope-pre-call.yaml" + path.write_text(yaml.safe_dump(config)) + with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(api_base=upstream.url + "/v1") + response: Final = candidate.request( + "POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": prompt}]} + ) + assert response.status_code == 200, response.text + assert response.json()["choices"][0]["message"]["content"] == reply, response.text + assert len(policy.drain()) == 0 + assert len(upstream.drain()) == 1 + guardrails: Final = candidate.get("/v2/guardrails/list")["guardrails"] + assert all(object_value(row)["guardrail_name"] != identity for row in guardrails), guardrails diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py index 45e652bff72..ddfdabf689d 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py @@ -967,19 +967,24 @@ class TestLoggingOnlyScopeValidation: scope: LoggingOnlyScope | None, callback_type: type[CustomGuardrail] = _LoggingOnlyScopeSupportedGuardrail, ) -> CustomGuardrail: + import litellm from litellm.proxy.guardrails import guardrail_registry as registry_module guardrail_type: Final = "logging_only_scope_test" + created_callbacks: Final[list[CustomGuardrail]] = [] def _initializer(litellm_params: LitellmParams, guardrail: Guardrail) -> CustomGuardrail: supported_event_hooks: Final = ( [GuardrailEventHooks.logging_only] if callback_type.use_native_lifecycle_hooks else None ) - return callback_type( + callback: Final = callback_type( guardrail_name=guardrail["guardrail_name"], event_hook=litellm_params.mode, supported_event_hooks=supported_event_hooks, ) + litellm.logging_callback_manager.add_litellm_callback(callback) + created_callbacks.append(callback) + return callback registry_module.guardrail_initializer_registry[guardrail_type] = _initializer lists: Final = _all_callback_lists() @@ -1000,6 +1005,10 @@ class TestLoggingOnlyScopeValidation: callback: Final = handler.guardrail_id_to_custom_guardrail[result["guardrail_id"]] assert callback is not None return callback + except ValueError: + callback: Final = created_callbacks[0] + assert all(callback not in callback_list for callback_list in lists) + raise finally: for callback_list, snapshot in zip(lists, snapshots): callback_list[:] = snapshot