mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-17 23:51:30 +00:00
fix(model_armor): record logging_only verdicts without raising so both scans run
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
a5cc65f1a9
commit
f417d7f739
2 changed files with 47 additions and 18 deletions
|
|
@ -1245,7 +1245,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
armor_response: Final = await self.make_model_armor_request(
|
||||
content=content, source=source, request_data=request_data
|
||||
)
|
||||
except ModelArmorAPIError as e:
|
||||
except (ModelArmorAPIError, httpx.HTTPError) as e:
|
||||
error_end_time: Final = time.time()
|
||||
self.add_standard_logging_guardrail_information_to_request_data(
|
||||
guardrail_json_response=str(e),
|
||||
|
|
@ -1256,7 +1256,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
end_time=error_end_time,
|
||||
duration=error_end_time - start_time,
|
||||
)
|
||||
raise
|
||||
return inputs
|
||||
|
||||
flagged: Final = self._should_block_content(armor_response, allow_sanitization=False)
|
||||
end_time: Final = time.time()
|
||||
|
|
@ -1269,14 +1269,6 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
end_time=end_time,
|
||||
duration=end_time - start_time,
|
||||
)
|
||||
if flagged:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=self._build_block_error_detail(
|
||||
"Response blocked by Model Armor" if input_type == "response" else "Violated content safety policy",
|
||||
armor_response,
|
||||
),
|
||||
)
|
||||
return inputs
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -5203,18 +5203,55 @@ async def test_native_post_call_mode_ignores_logging_hook():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_raises_on_flagged_content():
|
||||
async def test_apply_guardrail_records_flagged_without_raising():
|
||||
guardrail = _logging_only_guardrail()
|
||||
guardrail.make_model_armor_request = AsyncMock(return_value=_flagged_armor_response())
|
||||
request_data = {"metadata": {}}
|
||||
inputs = {"texts": ["forbidden output"]}
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await guardrail.apply_guardrail(
|
||||
inputs={"texts": ["forbidden output"]},
|
||||
request_data=request_data,
|
||||
input_type="response",
|
||||
)
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type="response",
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert result == inputs
|
||||
entries = request_data["metadata"]["standard_logging_guardrail_information"]
|
||||
assert entries[-1]["guardrail_status"] == "guardrail_flagged"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_logging_only_records_transport_error():
|
||||
guardrail = _logging_only_guardrail()
|
||||
guardrail.make_model_armor_request = AsyncMock(side_effect=httpx.ConnectError("boom"))
|
||||
response = _chat_response("some output")
|
||||
kwargs = _logged_kwargs()
|
||||
|
||||
out_kwargs, out_result = await guardrail.async_logging_hook(
|
||||
kwargs=kwargs, result=response, call_type="acompletion"
|
||||
)
|
||||
|
||||
assert out_result is response
|
||||
entries = _metadata_entries(out_kwargs)
|
||||
failed = [e for e in entries if e["guardrail_status"] == "guardrail_failed_to_respond"]
|
||||
assert failed
|
||||
assert all(e["guardrail_provider"] == "model_armor" for e in failed)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_logging_only_flagged_prompt_still_scans_response():
|
||||
"""A flagged input scan must not abort the output scan; both verdicts are recorded."""
|
||||
guardrail = _logging_only_guardrail()
|
||||
guardrail.make_model_armor_request = AsyncMock(return_value=_flagged_armor_response())
|
||||
response = _chat_response("flagged output")
|
||||
kwargs = _logged_kwargs()
|
||||
|
||||
out_kwargs, _ = await guardrail.async_logging_hook(
|
||||
kwargs=kwargs, result=response, call_type="acompletion"
|
||||
)
|
||||
|
||||
sources = [call.kwargs.get("source") for call in guardrail.make_model_armor_request.await_args_list]
|
||||
assert sources == ["user_prompt", "model_response"]
|
||||
entries = _metadata_entries(out_kwargs)
|
||||
flagged = [e for e in entries if e["guardrail_status"] == "guardrail_flagged"]
|
||||
assert len(flagged) == 2
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue