mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(proxy): correct Bedrock guardrail_mode in spend logs (pre/during/post)
Bedrock ApplyGuardrail uses INPUT/OUTPUT for the API body; spend logs must use the proxy hook (pre_call, during_call, post_call). Pass logging_event_type from each hook into make_bedrock_api_request. During-call was still logged as pre_call because unified guardrails call apply_guardrail with input_type=request. BedrockGuardrail now sets use_native_during_call_hook so during_call runs async_moderation_hook instead. Includes a small test asserting the Bedrock class flag. Made-with: Cursor
This commit is contained in:
parent
e64d98f725
commit
840fdfa915
4 changed files with 65 additions and 14 deletions
|
|
@ -2,6 +2,7 @@ from datetime import datetime
|
|||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
ClassVar,
|
||||
Dict,
|
||||
List,
|
||||
Literal,
|
||||
|
|
@ -81,6 +82,9 @@ class ModifyResponseException(Exception):
|
|||
|
||||
|
||||
class CustomGuardrail(CustomLogger):
|
||||
# If True, during_call runs async_moderation_hook instead of the unified apply_guardrail path.
|
||||
use_native_during_call_hook: ClassVar[bool] = False
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
guardrail_name: Optional[str] = None,
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ from typing import (
|
|||
TYPE_CHECKING,
|
||||
Any,
|
||||
AsyncGenerator,
|
||||
ClassVar,
|
||||
Dict,
|
||||
List,
|
||||
Literal,
|
||||
|
|
@ -123,6 +124,10 @@ def _redact_pii_matches(response_json: dict) -> dict:
|
|||
|
||||
|
||||
class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
||||
# During-call must use async_moderation_hook (not unified apply_guardrail), otherwise
|
||||
# OpenAI translation always passes input_type="request" and spend/UI show PRE-CALL.
|
||||
use_native_during_call_hook: ClassVar[bool] = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
guardrailIdentifier: Optional[str] = None,
|
||||
|
|
@ -413,6 +418,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
messages: Optional[List[AllMessageValues]] = None,
|
||||
response: Optional[Union[Any, litellm.ModelResponse]] = None,
|
||||
request_data: Optional[dict] = None,
|
||||
logging_event_type: Optional[GuardrailEventHooks] = None,
|
||||
) -> BedrockGuardrailResponse:
|
||||
from datetime import datetime
|
||||
|
||||
|
|
@ -450,11 +456,17 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
prepared_request.headers,
|
||||
)
|
||||
|
||||
event_type = (
|
||||
GuardrailEventHooks.pre_call
|
||||
if source == "INPUT"
|
||||
else GuardrailEventHooks.post_call
|
||||
)
|
||||
# UI / spend logs use event_type. Bedrock's `source` is INPUT vs OUTPUT for the API
|
||||
# body, which must not be confused with the proxy hook (pre_call / during_call /
|
||||
# post_call). When omitted, keep legacy mapping for backward compatibility.
|
||||
if logging_event_type is not None:
|
||||
event_type = logging_event_type
|
||||
else:
|
||||
event_type = (
|
||||
GuardrailEventHooks.pre_call
|
||||
if source == "INPUT"
|
||||
else GuardrailEventHooks.post_call
|
||||
)
|
||||
|
||||
try:
|
||||
httpx_response = await self.async_handler.post(
|
||||
|
|
@ -944,7 +956,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
] = None
|
||||
try:
|
||||
bedrock_guardrail_response = await self.make_bedrock_api_request(
|
||||
source="INPUT", messages=filtered_messages, request_data=data
|
||||
source="INPUT",
|
||||
messages=filtered_messages,
|
||||
request_data=data,
|
||||
logging_event_type=GuardrailEventHooks.pre_call,
|
||||
)
|
||||
except GuardrailInterventionNormalStringError as e:
|
||||
bedrock_guardrail_response = e.message
|
||||
|
|
@ -1016,7 +1031,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
] = None
|
||||
try:
|
||||
bedrock_guardrail_response = await self.make_bedrock_api_request(
|
||||
source="INPUT", messages=filtered_messages, request_data=data
|
||||
source="INPUT",
|
||||
messages=filtered_messages,
|
||||
request_data=data,
|
||||
logging_event_type=GuardrailEventHooks.during_call,
|
||||
)
|
||||
except GuardrailInterventionNormalStringError as e:
|
||||
bedrock_guardrail_response = e.message
|
||||
|
|
@ -1120,9 +1138,13 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
source="INPUT",
|
||||
messages=input_messages,
|
||||
request_data=data,
|
||||
logging_event_type=GuardrailEventHooks.post_call,
|
||||
)
|
||||
output_task = self.make_bedrock_api_request(
|
||||
source="OUTPUT", response=response, request_data=data
|
||||
source="OUTPUT",
|
||||
response=response,
|
||||
request_data=data,
|
||||
logging_event_type=GuardrailEventHooks.post_call,
|
||||
)
|
||||
|
||||
# Execute both requests in parallel
|
||||
|
|
@ -1136,7 +1158,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
# Only run OUTPUT validation (INPUT was already validated in pre_call or during_call)
|
||||
try:
|
||||
output_content_bedrock = await self.make_bedrock_api_request(
|
||||
source="OUTPUT", response=response, request_data=data
|
||||
source="OUTPUT",
|
||||
response=response,
|
||||
request_data=data,
|
||||
logging_event_type=GuardrailEventHooks.post_call,
|
||||
)
|
||||
except GuardrailInterventionNormalStringError as e:
|
||||
output_content_bedrock = e.message
|
||||
|
|
@ -1263,9 +1288,12 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
source="INPUT",
|
||||
messages=input_messages,
|
||||
request_data=request_data,
|
||||
logging_event_type=GuardrailEventHooks.post_call,
|
||||
) # Only input messages
|
||||
output_task = self.make_bedrock_api_request(
|
||||
source="OUTPUT", response=assembled_model_response
|
||||
source="OUTPUT",
|
||||
response=assembled_model_response,
|
||||
logging_event_type=GuardrailEventHooks.post_call,
|
||||
) # Only response
|
||||
|
||||
# Execute both requests in parallel
|
||||
|
|
@ -1279,7 +1307,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
# Only run OUTPUT validation (INPUT was already validated in pre_call or during_call)
|
||||
try:
|
||||
output_guardrail_response = await self.make_bedrock_api_request(
|
||||
source="OUTPUT", response=assembled_model_response
|
||||
source="OUTPUT",
|
||||
response=assembled_model_response,
|
||||
logging_event_type=GuardrailEventHooks.post_call,
|
||||
)
|
||||
except GuardrailInterventionNormalStringError as e:
|
||||
output_guardrail_response = e.message
|
||||
|
|
@ -1554,10 +1584,16 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
|
||||
# Bedrock will throw an error if there is no text to process
|
||||
if filtered_messages:
|
||||
_log_hook = (
|
||||
GuardrailEventHooks.pre_call
|
||||
if input_type == "request"
|
||||
else GuardrailEventHooks.post_call
|
||||
)
|
||||
bedrock_response = await self.make_bedrock_api_request(
|
||||
source="INPUT",
|
||||
messages=filtered_messages,
|
||||
request_data=request_data,
|
||||
logging_event_type=_log_hook,
|
||||
)
|
||||
|
||||
# Apply any masking that was applied by the guardrail
|
||||
|
|
|
|||
|
|
@ -940,7 +940,11 @@ class ProxyLogging:
|
|||
Result from the guardrail execution
|
||||
"""
|
||||
# Use unified_guardrail if callback has apply_guardrail method
|
||||
use_unified = "apply_guardrail" in type(callback).__dict__
|
||||
has_apply_guardrail = "apply_guardrail" in type(callback).__dict__
|
||||
use_unified = has_apply_guardrail and not (
|
||||
hook_type == "during_call"
|
||||
and getattr(callback, "use_native_during_call_hook", False)
|
||||
)
|
||||
if use_unified:
|
||||
data["guardrail_to_apply"] = callback
|
||||
|
||||
|
|
@ -1537,10 +1541,12 @@ class ProxyLogging:
|
|||
else:
|
||||
user_api_key_auth_dict = user_api_key_dict
|
||||
# Add task to list for parallel execution
|
||||
if (
|
||||
use_unified_during = (
|
||||
"apply_guardrail" in type(callback).__dict__
|
||||
and user_api_key_dict is not None
|
||||
):
|
||||
and not getattr(callback, "use_native_during_call_hook", False)
|
||||
)
|
||||
if use_unified_during:
|
||||
data["guardrail_to_apply"] = callback
|
||||
guardrail_task = self._run_guardrail_task_with_enrichment(
|
||||
callback,
|
||||
|
|
|
|||
|
|
@ -1204,6 +1204,11 @@ def _make_guardrail() -> BedrockGuardrail:
|
|||
)
|
||||
|
||||
|
||||
def test_bedrock_guardrail_uses_native_during_call_hook():
|
||||
"""during_call must use async_moderation_hook, not unified apply_guardrail(input=request)."""
|
||||
assert BedrockGuardrail.use_native_during_call_hook is True
|
||||
|
||||
|
||||
def test_extract_blocked_assessments_pii_entity():
|
||||
"""L3: PII entity match (BLOCKED) is surfaced with category, type, and matched term."""
|
||||
g = _make_guardrail()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue