mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
fix(guardrails): log resolved guardrail hook in StandardLoggingGuardrailInformation
This commit is contained in:
parent
735a5bfc5b
commit
5916dc047f
2 changed files with 121 additions and 5 deletions
|
|
@ -725,12 +725,22 @@ class CustomGuardrail(CustomLogger):
|
|||
guardrail_json_response = str(guardrail_json_response)
|
||||
from litellm.types.utils import GuardrailMode
|
||||
|
||||
# Use event_type if provided, otherwise fall back to self.event_hook
|
||||
guardrail_mode: Union[GuardrailEventHooks, GuardrailMode, List[GuardrailEventHooks]]
|
||||
# Use the resolved event_type (the concrete hook that actually fired
|
||||
# for *this* invocation) when available. Fall back to self.event_hook
|
||||
# only as a last resort — and normalise its various shapes so
|
||||
# downstream loggers always receive a JSON-serialisable scalar or
|
||||
# dict, never a raw List[str] that looks like a broken value.
|
||||
guardrail_mode: Union[GuardrailEventHooks, GuardrailMode, List[GuardrailEventHooks], str]
|
||||
if event_type is not None:
|
||||
guardrail_mode = event_type
|
||||
elif isinstance(self.event_hook, Mode):
|
||||
guardrail_mode = GuardrailMode(**dict(self.event_hook.model_dump())) # type: ignore[typeddict-item]
|
||||
elif isinstance(self.event_hook, list):
|
||||
# Config lists like ["pre_call", "post_call"] are the *configured*
|
||||
# modes, not the one that fired. Join into a comma-separated
|
||||
# string so the logged value is unambiguously "raw config" rather
|
||||
# than something that looks like a single resolved hook.
|
||||
guardrail_mode = ",".join(str(h) for h in self.event_hook) # type: ignore[assignment]
|
||||
else:
|
||||
guardrail_mode = self.event_hook # type: ignore[assignment]
|
||||
|
||||
|
|
@ -1089,8 +1099,17 @@ def log_guardrail_information(func):
|
|||
|
||||
def _infer_event_type_from_function_name(
|
||||
func_name: str,
|
||||
kwargs: Optional[dict] = None,
|
||||
) -> Optional[GuardrailEventHooks]:
|
||||
"""Infer the actual event type from the function name"""
|
||||
"""Infer the actual event type from the function name.
|
||||
|
||||
For ``apply_guardrail`` the concrete hook is derived from the
|
||||
``input_type`` kwarg: ``"request"`` → ``pre_call``,
|
||||
``"response"`` → ``post_call``. This avoids logging the raw
|
||||
``self.event_hook`` config (which may be a ``List`` or ``Mode``
|
||||
dict) instead of the concrete hook that was resolved for this
|
||||
invocation.
|
||||
"""
|
||||
if func_name == "async_pre_call_hook":
|
||||
return GuardrailEventHooks.pre_call
|
||||
elif func_name == "async_moderation_hook":
|
||||
|
|
@ -1100,6 +1119,13 @@ def log_guardrail_information(func):
|
|||
"async_post_call_streaming_hook",
|
||||
):
|
||||
return GuardrailEventHooks.post_call
|
||||
elif func_name == "apply_guardrail" and kwargs:
|
||||
input_type = kwargs.get("input_type")
|
||||
if input_type == "request":
|
||||
return GuardrailEventHooks.pre_call
|
||||
elif input_type == "response":
|
||||
return GuardrailEventHooks.post_call
|
||||
return GuardrailEventHooks.during_call
|
||||
return None
|
||||
|
||||
def _count_recorded_guardrail_entries(request_data: dict) -> int:
|
||||
|
|
@ -1117,7 +1143,7 @@ def log_guardrail_information(func):
|
|||
start_time = datetime.now() # Move start_time inside the wrapper
|
||||
self: CustomGuardrail = args[0]
|
||||
request_data: dict = kwargs.get("data") or kwargs.get("request_data") or {}
|
||||
event_type = _infer_event_type_from_function_name(func.__name__)
|
||||
event_type = _infer_event_type_from_function_name(func.__name__, kwargs)
|
||||
|
||||
# Store original inputs for comparison (for apply_guardrail functions)
|
||||
original_inputs = None
|
||||
|
|
@ -1158,7 +1184,7 @@ def log_guardrail_information(func):
|
|||
start_time = datetime.now() # Move start_time inside the wrapper
|
||||
self: CustomGuardrail = args[0]
|
||||
request_data: dict = kwargs.get("data") or kwargs.get("request_data") or {}
|
||||
event_type = _infer_event_type_from_function_name(func.__name__)
|
||||
event_type = _infer_event_type_from_function_name(func.__name__, kwargs)
|
||||
|
||||
# Store original inputs for comparison (for apply_guardrail functions)
|
||||
original_inputs = None
|
||||
|
|
|
|||
|
|
@ -1423,6 +1423,96 @@ class TestEventTypeLogging:
|
|||
assert len(logged_info) == 1
|
||||
assert logged_info[0]["guardrail_mode"] == GuardrailEventHooks.pre_call
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_infers_pre_call_from_input_type_request(self):
|
||||
"""apply_guardrail with input_type='request' should resolve to pre_call."""
|
||||
from litellm.integrations.custom_guardrail import log_guardrail_information
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
class TestGuardrail(CustomGuardrail):
|
||||
def __init__(self):
|
||||
super().__init__(
|
||||
guardrail_name="test_resolve",
|
||||
event_hook=[
|
||||
GuardrailEventHooks.pre_call,
|
||||
GuardrailEventHooks.post_call,
|
||||
],
|
||||
)
|
||||
|
||||
@log_guardrail_information
|
||||
async def apply_guardrail(self, inputs, request_data, input_type, **kwargs):
|
||||
return inputs
|
||||
|
||||
guardrail = TestGuardrail()
|
||||
request_data = {"metadata": {}}
|
||||
|
||||
await guardrail.apply_guardrail(
|
||||
inputs={"texts": ["hi"]},
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
logged_info = request_data["metadata"]["standard_logging_guardrail_information"]
|
||||
assert len(logged_info) == 1
|
||||
assert logged_info[0]["guardrail_mode"] == GuardrailEventHooks.pre_call
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_infers_post_call_from_input_type_response(self):
|
||||
"""apply_guardrail with input_type='response' should resolve to post_call."""
|
||||
from litellm.integrations.custom_guardrail import log_guardrail_information
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
class TestGuardrail(CustomGuardrail):
|
||||
def __init__(self):
|
||||
super().__init__(
|
||||
guardrail_name="test_resolve",
|
||||
event_hook=[
|
||||
GuardrailEventHooks.pre_call,
|
||||
GuardrailEventHooks.post_call,
|
||||
],
|
||||
)
|
||||
|
||||
@log_guardrail_information
|
||||
async def apply_guardrail(self, inputs, request_data, input_type, **kwargs):
|
||||
return inputs
|
||||
|
||||
guardrail = TestGuardrail()
|
||||
request_data = {"metadata": {}}
|
||||
|
||||
await guardrail.apply_guardrail(
|
||||
inputs={"texts": ["hi"]},
|
||||
request_data=request_data,
|
||||
input_type="response",
|
||||
)
|
||||
|
||||
logged_info = request_data["metadata"]["standard_logging_guardrail_information"]
|
||||
assert len(logged_info) == 1
|
||||
assert logged_info[0]["guardrail_mode"] == GuardrailEventHooks.post_call
|
||||
|
||||
def test_list_event_hook_fallback_serialises_as_csv_string(self):
|
||||
"""When event_type is None and event_hook is a list, the logged
|
||||
guardrail_mode should be a comma-separated string, not a raw list."""
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
guardrail = CustomGuardrail(
|
||||
guardrail_name="test_csv",
|
||||
event_hook=[GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call],
|
||||
)
|
||||
|
||||
request_data = {"metadata": {}}
|
||||
guardrail.add_standard_logging_guardrail_information_to_request_data(
|
||||
guardrail_json_response={"result": "ok"},
|
||||
request_data=request_data,
|
||||
guardrail_status="success",
|
||||
event_type=None,
|
||||
)
|
||||
|
||||
logged_info = request_data["metadata"]["standard_logging_guardrail_information"]
|
||||
mode = logged_info[0]["guardrail_mode"]
|
||||
assert isinstance(mode, str), f"Expected string, got {type(mode)}"
|
||||
assert "pre_call" in mode
|
||||
assert "post_call" in mode
|
||||
|
||||
|
||||
class TestTracingFieldsPopulation:
|
||||
"""Verify add_standard_logging_guardrail_information_to_request_data passes tracing_detail fields."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue