mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(guardrails): remove callbacks when scope validation fails
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
d392424670
commit
419f362294
4 changed files with 77 additions and 4 deletions
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue