From 58aa1c3bc9395f5935c19c046bbf2c1a4a61ce4a Mon Sep 17 00:00:00 2001 From: Krishna Date: Thu, 23 Jul 2026 09:06:44 +0530 Subject: [PATCH] fix(output_parsers): parse trailing category tokens for safe verdicts in is_content_safe --- nemoguardrails/llm/output_parsers.py | 12 +++++++++++- tests/test_content_safety_integration.py | 4 +--- 2 files changed, 12 insertions(+), 4 deletions(-) diff --git a/nemoguardrails/llm/output_parsers.py b/nemoguardrails/llm/output_parsers.py index b9c3b2c1eb..bf6d31c304 100644 --- a/nemoguardrails/llm/output_parsers.py +++ b/nemoguardrails/llm/output_parsers.py @@ -74,6 +74,16 @@ def verbose_v1_parser(s: str): return "\n".join(lines) +def _parse_safe_violations(response_text): + """Helper function to parse trailing category tokens from safe response.""" + lower_response = response_text.lower() + safe_pos = lower_response.find("safe") + if safe_pos != -1: + after_safe = response_text[safe_pos + len("safe"):].strip() + if after_safe: + violations = [v.strip() for v in after_safe.split() if v.strip()] + return violations + return [] def _parse_unsafe_violations(response_text): """Helper function to parse violations from unsafe response.""" @@ -125,7 +135,7 @@ def is_content_safe(response: str) -> Sequence[Union[bool, str]]: splited_response = response_lower.split(" ")[:2] response_actions = { - "safe": lambda: [True], + "safe": lambda: [True] + _parse_safe_violations(original_response), "unsafe": lambda: [False] + _parse_unsafe_violations(original_response), "yes": lambda: [False], "no": lambda: [True], diff --git a/tests/test_content_safety_integration.py b/tests/test_content_safety_integration.py index 930b84a029..ff0a8a7de6 100644 --- a/tests/test_content_safety_integration.py +++ b/tests/test_content_safety_integration.py @@ -109,9 +109,7 @@ async def test_content_safety_input_with_is_content_safe_parser_safe_with_violat ) assert result.is_blocked is False - # following assertion fails - # assert result.metadata["policy_violations"] == ["S1", "S8"] - assert result.metadata["policy_violations"] == [] + assert result.metadata["policy_violations"] == ["S1", "S8"] @pytest.mark.parametrize( "response,expected_allowed,expected_violations",