mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(guardrails): keep guardrails enforcing when logging_only_scope is invalid at load
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
fa5f7904a9
commit
e755b9a994
5 changed files with 145 additions and 26 deletions
|
|
@ -403,7 +403,11 @@ async def create_guardrail(
|
|||
guardrail_id: Final = result.get("guardrail_id", "Unknown")
|
||||
|
||||
try:
|
||||
IN_MEMORY_GUARDRAIL_HANDLER.initialize_guardrail(guardrail=cast(Guardrail, result), source="db")
|
||||
IN_MEMORY_GUARDRAIL_HANDLER.initialize_guardrail(
|
||||
guardrail=cast(Guardrail, result),
|
||||
source="db",
|
||||
reject_invalid_logging_only_scope=True,
|
||||
)
|
||||
verbose_proxy_logger.info(
|
||||
"Immediate sync: Successfully initialized guardrail '%s' (ID: %s)", guardrail_name, guardrail_id
|
||||
)
|
||||
|
|
@ -526,7 +530,10 @@ async def update_guardrail(
|
|||
guardrail_name: Final = result.get("guardrail_name", "Unknown")
|
||||
|
||||
try:
|
||||
IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(guardrail=cast(Guardrail, result))
|
||||
IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(
|
||||
guardrail=cast(Guardrail, result),
|
||||
reject_invalid_logging_only_scope=True,
|
||||
)
|
||||
verbose_proxy_logger.info(
|
||||
"Immediate sync: Successfully updated guardrail '%s' (ID: %s)", guardrail_name, guardrail_id
|
||||
)
|
||||
|
|
@ -1246,6 +1253,7 @@ async def patch_guardrail(
|
|||
try:
|
||||
IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(
|
||||
guardrail=guardrail,
|
||||
reject_invalid_logging_only_scope=True,
|
||||
)
|
||||
verbose_proxy_logger.info(
|
||||
"Immediate sync: Successfully updated guardrail '%s' (ID: %s)", guardrail_name, guardrail_id
|
||||
|
|
|
|||
|
|
@ -437,23 +437,46 @@ def _as_callback_tuple(
|
|||
return (initialized,)
|
||||
|
||||
|
||||
def _configure_callback_scoping(
|
||||
def _logging_only_scope_error(
|
||||
custom_guardrail_callback: CustomGuardrail, guardrail_name: str, litellm_params: LitellmParams
|
||||
) -> None:
|
||||
) -> str | None:
|
||||
logging_only_scope: Final = litellm_params.logging_only_scope
|
||||
custom_guardrail_callback.logging_only_scope = logging_only_scope
|
||||
if logging_only_scope is not None and GuardrailEventHooks.logging_only.value not in _configured_event_hooks(
|
||||
litellm_params.mode
|
||||
):
|
||||
raise ValueError(
|
||||
return (
|
||||
f"Guardrail {guardrail_name}: logging_only_scope is set, but mode does not include logging_only, "
|
||||
"so it would never apply. Add logging_only to mode or remove logging_only_scope."
|
||||
)
|
||||
if logging_only_scope in ("input", "output") and not custom_guardrail_callback.supports_logging_only_scope():
|
||||
raise ValueError(
|
||||
return (
|
||||
f"Guardrail {guardrail_name}: logging_only_scope={logging_only_scope!r} is not supported by this "
|
||||
"guardrail, whose logging_only hook scans on its own. Remove logging_only_scope."
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
def _configure_callback_scoping(
|
||||
custom_guardrail_callback: CustomGuardrail,
|
||||
guardrail_name: str,
|
||||
litellm_params: LitellmParams,
|
||||
*,
|
||||
reject_invalid_logging_only_scope: bool = False,
|
||||
) -> None:
|
||||
logging_only_scope: Final = litellm_params.logging_only_scope
|
||||
logging_only_scope_error: Final = _logging_only_scope_error(
|
||||
custom_guardrail_callback, guardrail_name, litellm_params
|
||||
)
|
||||
if logging_only_scope_error is not None:
|
||||
if reject_invalid_logging_only_scope:
|
||||
raise ValueError(logging_only_scope_error)
|
||||
verbose_proxy_logger.error(
|
||||
"%s Ignoring logging_only_scope; the guardrail keeps its configured mode.",
|
||||
logging_only_scope_error,
|
||||
)
|
||||
custom_guardrail_callback.logging_only_scope = None
|
||||
else:
|
||||
custom_guardrail_callback.logging_only_scope = logging_only_scope
|
||||
for scoping_param in (
|
||||
"skip_system_message_in_guardrail",
|
||||
"skip_tool_message_in_guardrail",
|
||||
|
|
@ -512,6 +535,8 @@ class InMemoryGuardrailHandler:
|
|||
config_file_path: str | None = None,
|
||||
llm_router: Optional["Router"] = None,
|
||||
source: Literal["db", "config"] = "config",
|
||||
*,
|
||||
reject_invalid_logging_only_scope: bool = False,
|
||||
) -> Guardrail | None:
|
||||
"""
|
||||
Initialize a guardrail from a dictionary and add it to the litellm callback manager
|
||||
|
|
@ -560,7 +585,12 @@ class InMemoryGuardrailHandler:
|
|||
)
|
||||
try:
|
||||
for custom_guardrail_callback in created_callbacks:
|
||||
_configure_callback_scoping(custom_guardrail_callback, guardrail["guardrail_name"], litellm_params)
|
||||
_configure_callback_scoping(
|
||||
custom_guardrail_callback,
|
||||
guardrail["guardrail_name"],
|
||||
litellm_params,
|
||||
reject_invalid_logging_only_scope=reject_invalid_logging_only_scope,
|
||||
)
|
||||
except Exception:
|
||||
for custom_guardrail_callback in created_callbacks:
|
||||
litellm.logging_callback_manager.remove_callback_from_all_lists(custom_guardrail_callback)
|
||||
|
|
@ -672,6 +702,8 @@ class InMemoryGuardrailHandler:
|
|||
guardrail_id: str,
|
||||
guardrail: Guardrail,
|
||||
source: Literal["db", "config"] = "db",
|
||||
*,
|
||||
reject_invalid_logging_only_scope: bool = False,
|
||||
) -> None:
|
||||
"""
|
||||
Update a guardrail in memory: a changed name or litellm_params rebuilds the
|
||||
|
|
@ -680,7 +712,11 @@ class InMemoryGuardrailHandler:
|
|||
"""
|
||||
updated_guardrail: Final = cast(Guardrail, {**guardrail, "guardrail_id": guardrail_id})
|
||||
if self._has_guardrail_params_changed(guardrail_id, updated_guardrail):
|
||||
self.reinitialize_guardrail(guardrail=updated_guardrail, source=source)
|
||||
self.reinitialize_guardrail(
|
||||
guardrail=updated_guardrail,
|
||||
source=source,
|
||||
reject_invalid_logging_only_scope=reject_invalid_logging_only_scope,
|
||||
)
|
||||
return
|
||||
self.IN_MEMORY_GUARDRAILS[guardrail_id] = updated_guardrail
|
||||
self._sources[guardrail_id] = source
|
||||
|
|
@ -833,6 +869,8 @@ class InMemoryGuardrailHandler:
|
|||
guardrail: Guardrail,
|
||||
config_file_path: str | None = None,
|
||||
source: Literal["db", "config"] = "config",
|
||||
*,
|
||||
reject_invalid_logging_only_scope: bool = False,
|
||||
) -> Guardrail | None:
|
||||
"""
|
||||
Force re-initialization of a guardrail even if it exists in memory.
|
||||
|
|
@ -862,7 +900,12 @@ class InMemoryGuardrailHandler:
|
|||
# instance instead of leaving the guardrail silently removed: a guardrail
|
||||
# that was enforcing must never fail open because an update was bad.
|
||||
try:
|
||||
return self.initialize_guardrail(guardrail=guardrail, config_file_path=config_file_path, source=source)
|
||||
return self.initialize_guardrail(
|
||||
guardrail=guardrail,
|
||||
config_file_path=config_file_path,
|
||||
source=source,
|
||||
reject_invalid_logging_only_scope=reject_invalid_logging_only_scope,
|
||||
)
|
||||
except Exception as init_error:
|
||||
if previous_guardrail is not None:
|
||||
verbose_proxy_logger.exception(
|
||||
|
|
@ -877,7 +920,13 @@ class InMemoryGuardrailHandler:
|
|||
verbose_proxy_logger.exception("Restoring previous guardrail %s also failed", guardrail_id)
|
||||
raise ValueError(f"Guardrail initialization failed: {init_error}") from init_error
|
||||
|
||||
def sync_guardrail_from_db(self, guardrail: Guardrail, config_file_path: str | None = None) -> Guardrail | None:
|
||||
def sync_guardrail_from_db(
|
||||
self,
|
||||
guardrail: Guardrail,
|
||||
config_file_path: str | None = None,
|
||||
*,
|
||||
reject_invalid_logging_only_scope: bool = False,
|
||||
) -> Guardrail | None:
|
||||
"""
|
||||
Sync a guardrail from DB - initializes if new, re-initializes if changed.
|
||||
This is the method to call during DB polling.
|
||||
|
|
@ -896,6 +945,7 @@ class InMemoryGuardrailHandler:
|
|||
guardrail=guardrail,
|
||||
config_file_path=config_file_path,
|
||||
source="db",
|
||||
reject_invalid_logging_only_scope=reject_invalid_logging_only_scope,
|
||||
)
|
||||
|
||||
# Params unchanged but the entry is still DB-backed; make sure the
|
||||
|
|
|
|||
|
|
@ -1393,12 +1393,11 @@ def test_logging_only_scope_observes_only_the_configured_direction_without_block
|
|||
assert detail["requestsEvaluated"] == len(scanned_directions), detail
|
||||
|
||||
|
||||
def test_logging_only_scope_without_logging_only_mode_is_skipped_at_load_and_never_blocks(
|
||||
def test_logging_only_scope_without_logging_only_mode_is_ignored_at_load_and_keeps_blocking(
|
||||
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
|
||||
prompt: Final = "synthetic invalid-scope prompt pineapple " + identity
|
||||
|
||||
def guardrail(request: Request) -> Reply:
|
||||
assert request.target == "/beta/litellm_basic_guardrail_api"
|
||||
|
|
@ -1415,7 +1414,11 @@ def test_logging_only_scope_without_logging_only_mode_is_skipped_at_load_and_nev
|
|||
"created": 1,
|
||||
"model": "gpt-4o-mini",
|
||||
"choices": [
|
||||
{"index": 0, "message": {"role": "assistant", "content": reply}, "finish_reason": "stop"}
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "unchanged provider reply"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 9, "completion_tokens": 5, "total_tokens": 14},
|
||||
}
|
||||
|
|
@ -1432,6 +1435,7 @@ def test_logging_only_scope_without_logging_only_mode_is_skipped_at_load_and_nev
|
|||
"mode": "pre_call",
|
||||
"logging_only_scope": "input",
|
||||
"default_on": True,
|
||||
"blocked_words": [{"keyword": "pineapple", "action": "BLOCK"}],
|
||||
"api_base": policy.url,
|
||||
"api_key": "synthetic-guardrail-key",
|
||||
},
|
||||
|
|
@ -1444,9 +1448,9 @@ def test_logging_only_scope_without_logging_only_mode_is_skipped_at_load_and_nev
|
|||
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
|
||||
assert response.status_code == 400, response.text
|
||||
assert "synthetic policy denial" in response.text, response.text
|
||||
assert len(policy.drain()) == 1
|
||||
assert len(upstream.drain()) == 0
|
||||
guardrails: Final = candidate.get("/v2/guardrails/list")["guardrails"]
|
||||
assert all(object_value(row)["guardrail_name"] != identity for row in guardrails), guardrails
|
||||
assert any(object_value(row)["guardrail_name"] == identity for row in guardrails), guardrails
|
||||
|
|
|
|||
|
|
@ -1123,7 +1123,10 @@ async def test_update_guardrail_endpoint(
|
|||
prisma_client=mocker.ANY,
|
||||
)
|
||||
|
||||
mock_in_memory_handler.sync_guardrail_from_db.assert_called_once_with(guardrail=mocker.ANY)
|
||||
mock_in_memory_handler.sync_guardrail_from_db.assert_called_once_with(
|
||||
guardrail=mocker.ANY,
|
||||
reject_invalid_logging_only_scope=True,
|
||||
)
|
||||
|
||||
if scenario == "success_sync_fails_unexpected_error":
|
||||
assert mock_logger is not None
|
||||
|
|
@ -1252,7 +1255,10 @@ async def test_patch_guardrail_endpoint(
|
|||
|
||||
mock_guardrail_registry.update_guardrail_in_db.assert_called_once()
|
||||
|
||||
mock_in_memory_handler.sync_guardrail_from_db.assert_called_once_with(guardrail=mocker.ANY)
|
||||
mock_in_memory_handler.sync_guardrail_from_db.assert_called_once_with(
|
||||
guardrail=mocker.ANY,
|
||||
reject_invalid_logging_only_scope=True,
|
||||
)
|
||||
|
||||
if scenario == "success_sync_fails_unexpected_error":
|
||||
assert mock_logger is not None
|
||||
|
|
@ -1276,6 +1282,34 @@ async def test_patch_guardrail_rejects_mcp_only_on_violation_with_422(mocker, mo
|
|||
mock_guardrail_registry.update_guardrail_in_db.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_patch_guardrail_rejects_invalid_logging_only_scope_with_422(mocker, mock_guardrail_registry):
|
||||
mocker.patch("litellm.proxy.proxy_server.prisma_client", mocker.Mock())
|
||||
mocker.patch(
|
||||
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY",
|
||||
mock_guardrail_registry,
|
||||
)
|
||||
mock_in_memory_handler = mocker.Mock(spec=InMemoryGuardrailHandler)
|
||||
mock_in_memory_handler.sync_guardrail_from_db.side_effect = ValueError(
|
||||
"Guardrail test-db-guardrail: logging_only_scope is set, but mode does not include logging_only"
|
||||
)
|
||||
mocker.patch(
|
||||
"litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER",
|
||||
mock_in_memory_handler,
|
||||
)
|
||||
request = PatchGuardrailRequest(litellm_params=BaseLitellmParams(mode="pre_call", logging_only_scope="input"))
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await patch_guardrail("test-guardrail-id", request, user_api_key_dict=MOCK_ADMIN_USER)
|
||||
|
||||
assert exc_info.value.status_code == 422
|
||||
assert "update rejected" in str(exc_info.value.detail)
|
||||
mock_in_memory_handler.sync_guardrail_from_db.assert_called_once_with(
|
||||
guardrail=mocker.ANY,
|
||||
reject_invalid_logging_only_scope=True,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"scenario,expected_result,expected_exception",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -966,6 +966,8 @@ class TestLoggingOnlyScopeValidation:
|
|||
mode: str | list[str] | Mode,
|
||||
scope: LoggingOnlyScope | None,
|
||||
callback_type: type[CustomGuardrail] = _LoggingOnlyScopeSupportedGuardrail,
|
||||
reject_invalid_logging_only_scope: bool = False,
|
||||
assert_registered: bool = False,
|
||||
) -> CustomGuardrail:
|
||||
import litellm
|
||||
from litellm.proxy.guardrails import guardrail_registry as registry_module
|
||||
|
|
@ -980,6 +982,7 @@ class TestLoggingOnlyScopeValidation:
|
|||
callback: Final = callback_type(
|
||||
guardrail_name=guardrail["guardrail_name"],
|
||||
event_hook=litellm_params.mode,
|
||||
default_on=True,
|
||||
supported_event_hooks=supported_event_hooks,
|
||||
)
|
||||
litellm.logging_callback_manager.add_litellm_callback(callback)
|
||||
|
|
@ -999,11 +1002,14 @@ class TestLoggingOnlyScopeValidation:
|
|||
"mode": mode,
|
||||
"logging_only_scope": scope,
|
||||
},
|
||||
}
|
||||
},
|
||||
reject_invalid_logging_only_scope=reject_invalid_logging_only_scope,
|
||||
)
|
||||
assert result is not None
|
||||
callback: Final = handler.guardrail_id_to_custom_guardrail[result["guardrail_id"]]
|
||||
assert callback is not None
|
||||
if assert_registered:
|
||||
assert callback in lists[0]
|
||||
return callback
|
||||
except ValueError:
|
||||
callback: Final = created_callbacks[0]
|
||||
|
|
@ -1014,9 +1020,15 @@ class TestLoggingOnlyScopeValidation:
|
|||
callback_list[:] = snapshot
|
||||
registry_module.guardrail_initializer_registry.pop(guardrail_type, None)
|
||||
|
||||
def test_scope_requires_logging_only_mode(self) -> None:
|
||||
def test_scope_without_logging_only_mode_is_ignored_at_load(self) -> None:
|
||||
callback: Final = self._initialize(mode="pre_call", scope="input", assert_registered=True)
|
||||
|
||||
assert callback.logging_only_scope is None
|
||||
assert callback.should_run_guardrail(data={}, event_type=GuardrailEventHooks.pre_call) is True
|
||||
|
||||
def test_scope_without_logging_only_mode_is_rejected_for_api_writes(self) -> None:
|
||||
with pytest.raises(ValueError, match="logging_only_scope is set") as exc_info:
|
||||
self._initialize(mode="pre_call", scope="input")
|
||||
self._initialize(mode="pre_call", scope="input", reject_invalid_logging_only_scope=True)
|
||||
|
||||
assert str(exc_info.value) == (
|
||||
"Guardrail logging-only-scope-guardrail: logging_only_scope is set, but mode does not include "
|
||||
|
|
@ -1036,12 +1048,23 @@ class TestLoggingOnlyScopeValidation:
|
|||
|
||||
assert callback.logging_only_scope == "input"
|
||||
|
||||
def test_directional_scope_rejected_when_guardrail_owns_logging_hook(self) -> None:
|
||||
def test_directional_scope_is_ignored_at_load_when_guardrail_owns_logging_hook(self) -> None:
|
||||
callback: Final = self._initialize(
|
||||
mode="logging_only",
|
||||
scope="input",
|
||||
callback_type=_LoggingOnlyScopeUnsupportedGuardrail,
|
||||
assert_registered=True,
|
||||
)
|
||||
|
||||
assert callback.logging_only_scope is None
|
||||
|
||||
def test_directional_scope_rejected_for_api_writes_when_guardrail_owns_logging_hook(self) -> None:
|
||||
with pytest.raises(ValueError, match="logging_only_scope='input' is not supported") as exc_info:
|
||||
self._initialize(
|
||||
mode="logging_only",
|
||||
scope="input",
|
||||
callback_type=_LoggingOnlyScopeUnsupportedGuardrail,
|
||||
reject_invalid_logging_only_scope=True,
|
||||
)
|
||||
|
||||
assert str(exc_info.value) == (
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue