From fec4318254402e529ad9af7f6b438e4e8d0764cf Mon Sep 17 00:00:00 2001 From: Jazzcort Date: Thu, 23 Jul 2026 15:18:48 -0400 Subject: [PATCH 1/3] Refactor ShieldModerationBlocked.refusal_response to a computed property Eliminates redundant stored state by deriving refusal_response from message at access time, removing the now-unnecessary create_refusal_response helper and all constructor-site arguments. --- src/models/common/moderation.py | 9 ++++++++- src/utils/shields.py | 16 ---------------- tests/unit/app/endpoints/test_responses.py | 5 ----- tests/unit/app/endpoints/test_rlsapi_v1.py | 9 --------- tests/unit/utils/agents/test_query.py | 7 ------- tests/unit/utils/agents/test_streaming.py | 13 +++---------- 6 files changed, 11 insertions(+), 48 deletions(-) diff --git a/src/models/common/moderation.py b/src/models/common/moderation.py index 7672e23c5..82bde17e3 100644 --- a/src/models/common/moderation.py +++ b/src/models/common/moderation.py @@ -20,7 +20,14 @@ class ShieldModerationBlocked(BaseModel): decision: Literal["blocked"] = "blocked" message: str moderation_id: str - refusal_response: ResponseMessage + + @property + def refusal_response(self) -> ResponseMessage: + + return ResponseMessage( + role="assistant", + content=self.message, + ) ShieldModerationResult = Annotated[ diff --git a/src/utils/shields.py b/src/utils/shields.py index e930c009f..a4470b4dd 100644 --- a/src/utils/shields.py +++ b/src/utils/shields.py @@ -3,7 +3,6 @@ from typing import Optional from fastapi import HTTPException -from ogx_api import OpenAIResponseMessage from ogx_client import ( APIConnectionError, AsyncOgxClient, @@ -130,21 +129,6 @@ async def append_turn_to_conversation( raise HTTPException(**error_response.model_dump()) from e -def create_refusal_response(refusal_message: str) -> OpenAIResponseMessage: - """Create a refusal response message object. - - Args: - refusal_message: The refusal message text. - - Returns: - OpenAIResponseMessage with refusal message. - """ - return OpenAIResponseMessage( - role="assistant", - content=refusal_message, - ) - - def get_shields_for_request( shields: list[ShieldConfiguration], shield_ids: Optional[list[str]] = None, diff --git a/tests/unit/app/endpoints/test_responses.py b/tests/unit/app/endpoints/test_responses.py index 0ec345038..c214f3a63 100644 --- a/tests/unit/app/endpoints/test_responses.py +++ b/tests/unit/app/endpoints/test_responses.py @@ -189,11 +189,6 @@ def _patch_moderation(mocker: MockerFixture, decision: str = "passed") -> Any: moderation_result = ShieldModerationBlocked( message="Content blocked", moderation_id="mod_blocked", - refusal_response=OpenAIResponseMessage( - role="assistant", - content="Content blocked", - type="message", - ), ) else: moderation_result = ShieldModerationPassed() diff --git a/tests/unit/app/endpoints/test_rlsapi_v1.py b/tests/unit/app/endpoints/test_rlsapi_v1.py index 43c56ab93..2338ba1ca 100644 --- a/tests/unit/app/endpoints/test_rlsapi_v1.py +++ b/tests/unit/app/endpoints/test_rlsapi_v1.py @@ -13,7 +13,6 @@ import pytest from fastapi import HTTPException, status -from ogx_api import OpenAIResponseMessage from ogx_client import APIConnectionError, APIStatusError from ogx_client.types import ListModelsResponse from ogx_client.types.model import Model @@ -1254,10 +1253,6 @@ async def test_infer_quota_shield_blocked_does_not_consume_tokens( blocked = ShieldModerationBlocked( message="Blocked by moderation", moderation_id="modr-test", - refusal_response=OpenAIResponseMessage( - role="assistant", - content="Blocked by moderation", - ), ) mocker.patch( "app.endpoints.rlsapi_v1.run_shield_moderation", @@ -1287,10 +1282,6 @@ def _create_blocked_moderation_result() -> ShieldModerationBlocked: return ShieldModerationBlocked( message="I can't answer that. Can I help with something else?", moderation_id="modr-test-123", - refusal_response=OpenAIResponseMessage( - role="assistant", - content="I can't answer that. Can I help with something else?", - ), ) diff --git a/tests/unit/utils/agents/test_query.py b/tests/unit/utils/agents/test_query.py index c6524e662..94f72cf48 100644 --- a/tests/unit/utils/agents/test_query.py +++ b/tests/unit/utils/agents/test_query.py @@ -5,9 +5,6 @@ import pytest from fastapi import HTTPException -from ogx_api.openai_responses import ( - OpenAIResponseMessage as ResponseMessage, -) from ogx_client import APIConnectionError, APIStatusError from pydantic_ai.messages import ( FinishReason, @@ -109,10 +106,6 @@ def blocked_moderation_fixture() -> ShieldModerationBlocked: return ShieldModerationBlocked( message="Content blocked by shield.", moderation_id="modr-test-456", - refusal_response=ResponseMessage( - role="assistant", - content="Content blocked by shield.", - ), ) diff --git a/tests/unit/utils/agents/test_streaming.py b/tests/unit/utils/agents/test_streaming.py index 0e7aa3afe..7386eef65 100644 --- a/tests/unit/utils/agents/test_streaming.py +++ b/tests/unit/utils/agents/test_streaming.py @@ -9,9 +9,6 @@ import pytest from fastapi import HTTPException -from ogx_api.openai_responses import ( - OpenAIResponseMessage as ResponseMessage, -) from ogx_client import APIStatusError from pydantic_ai import AgentRunResultEvent from pydantic_ai.exceptions import AgentRunError @@ -118,10 +115,6 @@ def blocked_moderation_fixture() -> ShieldModerationBlocked: return ShieldModerationBlocked( message="Content blocked by shield.", moderation_id="modr-test-456", - refusal_response=ResponseMessage( - role="assistant", - content="Content blocked by shield.", - ), ) @@ -1055,9 +1048,9 @@ async def __aexit__(self, *_args: object) -> None: num_chunks = len(chunk_ids) assert chunk_ids == sorted(chunk_ids), "chunk_ids must be monotonically ordered" assert all(cid >= 0 for cid in chunk_ids), "all chunk_ids must be non-negative" - assert num_chunks == len( - set(chunk_ids) - ), "chunk_ids must not contain duplicates" + assert num_chunks == len(set(chunk_ids)), ( + "chunk_ids must not contain duplicates" + ) assert chunk_ids[-1] == num_chunks - 1 @pytest.mark.asyncio From eb3492063e439e382b4e9f3874791e022c35d4d0 Mon Sep 17 00:00:00 2001 From: Jazzcort Date: Wed, 22 Jul 2026 14:28:32 -0400 Subject: [PATCH 2/3] Add AbstractSafetyCapability with standalone run() interface Introduce AbstractSafetyCapability(AbstractCapability[T]) in capabilities/base.py, requiring subclasses to implement a run() method that accepts raw text and returns ShieldModerationResult. This enables safety capabilities to be invoked outside the pydantic-ai agent lifecycle. Migrate QuestionValidity and PiiRedactionCapability to extend the new base class and implement run(): - QuestionValidity.run() delegates to model_request and maps ALLOWED/REJECTED to ShieldModerationPassed/Blocked - PiiRedactionCapability.run() applies regex redaction and returns Blocked when PII is detected Add unit tests for both run() implementations. --- src/models/common/moderation.py | 2 +- .../capabilities/base.py | 18 +++ .../question_validity/_capability.py | 25 ++++- .../capabilities/redaction/_capability.py | 21 +++- .../endpoints/test_responses_integration.py | 5 - tests/unit/app/endpoints/test_responses.py | 3 - .../question_validity/test_capability.py | 106 ++++++++++++++++++ .../capabilities/redaction/test_capability.py | 44 ++++++++ tests/unit/utils/agents/test_streaming.py | 6 +- 9 files changed, 214 insertions(+), 16 deletions(-) create mode 100644 src/pydantic_ai_lightspeed/capabilities/base.py diff --git a/src/models/common/moderation.py b/src/models/common/moderation.py index 82bde17e3..575d6e535 100644 --- a/src/models/common/moderation.py +++ b/src/models/common/moderation.py @@ -23,7 +23,7 @@ class ShieldModerationBlocked(BaseModel): @property def refusal_response(self) -> ResponseMessage: - + """Build a ResponseMessage carrying the shield's refusal text.""" return ResponseMessage( role="assistant", content=self.message, diff --git a/src/pydantic_ai_lightspeed/capabilities/base.py b/src/pydantic_ai_lightspeed/capabilities/base.py new file mode 100644 index 000000000..25b5dc8b4 --- /dev/null +++ b/src/pydantic_ai_lightspeed/capabilities/base.py @@ -0,0 +1,18 @@ +"""Abstract base for safety capabilities with a standalone run interface.""" + +from abc import abstractmethod + +from pydantic_ai.capabilities import AbstractCapability +from typing_extensions import TypeVar + +from models.common.moderation import ShieldModerationResult + +T = TypeVar("T", default=object) + + +class AbstractSafetyCapability(AbstractCapability[T]): + """Interface for safety/moderation that can be called directly.""" + + @abstractmethod + async def run(self, input_text: str) -> ShieldModerationResult: + """Run moderation on input text.""" diff --git a/src/pydantic_ai_lightspeed/capabilities/question_validity/_capability.py b/src/pydantic_ai_lightspeed/capabilities/question_validity/_capability.py index a4f69241a..55a9ae526 100644 --- a/src/pydantic_ai_lightspeed/capabilities/question_validity/_capability.py +++ b/src/pydantic_ai_lightspeed/capabilities/question_validity/_capability.py @@ -13,10 +13,11 @@ from dataclasses import dataclass, field from string import Template from typing import Optional +from uuid import uuid4 from pydantic_ai import AgentRunResult, RunContext from pydantic_ai._agent_graph import GraphAgentState -from pydantic_ai.capabilities import AbstractCapability, WrapRunHandler +from pydantic_ai.capabilities import WrapRunHandler from pydantic_ai.direct import model_request from pydantic_ai.messages import ModelRequest, TextContent, UserContent from pydantic_ai.models import Model @@ -24,9 +25,15 @@ from client import AsyncOgxClientHolder from log import get_logger +from models.common.moderation import ( + ShieldModerationBlocked, + ShieldModerationPassed, + ShieldModerationResult, +) from models.config import ( QuestionValidityConfig, ) +from pydantic_ai_lightspeed.capabilities.base import AbstractSafetyCapability from pydantic_ai_lightspeed.llamastack import OgxResponsesModel logger = get_logger(__name__) @@ -56,7 +63,7 @@ def _extract_message_str_from_user_content(user_content: Sequence[UserContent]) @dataclass -class QuestionValidity(AbstractCapability[None]): +class QuestionValidity(AbstractSafetyCapability): """Block or modify user input based on a guardrail check. The guard function receives the user prompt and returns True if safe. @@ -140,3 +147,17 @@ async def wrap_run( return AgentRunResult( output=self.config.invalid_question_response, _state=state ) + + async def run(self, input_text: str) -> ShieldModerationResult: + """Run question-validity check and return a moderation result.""" + result = await model_request( + model=self._model, messages=[ModelRequest.user_text_prompt(input_text)] + ) + + if result.text is not None and result.text.strip() == SUBJECT_ALLOWED: + return ShieldModerationPassed() + + return ShieldModerationBlocked( + message=self.config.invalid_question_response, + moderation_id=f"modr-{uuid4()}", + ) diff --git a/src/pydantic_ai_lightspeed/capabilities/redaction/_capability.py b/src/pydantic_ai_lightspeed/capabilities/redaction/_capability.py index 0bdedb174..29c3ef910 100644 --- a/src/pydantic_ai_lightspeed/capabilities/redaction/_capability.py +++ b/src/pydantic_ai_lightspeed/capabilities/redaction/_capability.py @@ -3,9 +3,9 @@ from collections.abc import Sequence from dataclasses import dataclass, replace from typing import Any, Optional +from uuid import uuid4 from pydantic_ai import RunContext -from pydantic_ai.capabilities import AbstractCapability from pydantic_ai.messages import ( ModelMessage, ModelRequest, @@ -19,7 +19,13 @@ ) from pydantic_ai.models import ModelRequestContext +from models.common.moderation import ( + ShieldModerationBlocked, + ShieldModerationPassed, + ShieldModerationResult, +) from models.config import RedactionConfig +from pydantic_ai_lightspeed.capabilities.base import AbstractSafetyCapability from pydantic_ai_lightspeed.capabilities.redaction.core import ( CompiledPatterns, redact_text, @@ -257,7 +263,7 @@ def _redact_response( @dataclass -class PiiRedactionCapability(AbstractCapability[Any]): +class PiiRedactionCapability(AbstractSafetyCapability): """Pydantic AI capability that redacts PII from agent messages. Applies configurable regex-based redaction rules to user prompt @@ -321,3 +327,14 @@ async def after_model_request( ) return new_response + + async def run(self, input_text: str) -> ShieldModerationResult: + """Run PII redaction on input text and return a moderation result.""" + result = redact_text(input_text, self.config.compiled_patterns) + + if result.redacted: + return ShieldModerationBlocked( + message="Sensitive content detected.", moderation_id=f"modr-{uuid4()}" + ) + + return ShieldModerationPassed() diff --git a/tests/integration/endpoints/test_responses_integration.py b/tests/integration/endpoints/test_responses_integration.py index 638277646..aae7b9f26 100644 --- a/tests/integration/endpoints/test_responses_integration.py +++ b/tests/integration/endpoints/test_responses_integration.py @@ -11,7 +11,6 @@ import pytest from fastapi import Request from fastapi.responses import StreamingResponse -from ogx_api.openai_responses import OpenAIResponseMessage from ogx_client.types import ListModelsResponse from ogx_client.types.model import Model from pytest_mock import MockerFixture @@ -172,10 +171,6 @@ def _configure_shield_blocked( blocked = ShieldModerationBlocked( message="Content blocked by safety shield", moderation_id=moderation_id, - refusal_response=OpenAIResponseMessage( - role="assistant", - content="Content blocked by safety shield", - ), ) mocker.patch( "app.endpoints.responses.run_shield_moderation", diff --git a/tests/unit/app/endpoints/test_responses.py b/tests/unit/app/endpoints/test_responses.py index c214f3a63..57020665b 100644 --- a/tests/unit/app/endpoints/test_responses.py +++ b/tests/unit/app/endpoints/test_responses.py @@ -623,9 +623,6 @@ async def test_responses_blocked_with_conversation_appends_refusal( mock_moderation = _patch_moderation(mocker, decision="blocked") mock_moderation.message = "Blocked" mock_moderation.moderation_id = "resp_blocked_123" - mock_moderation.refusal_response = OpenAIResponseMessage( - type="message", role="assistant", content="Blocked" - ) mock_append = mocker.patch( f"{MODULE}.append_turn_items_to_conversation", new=mocker.AsyncMock(), diff --git a/tests/unit/pydantic_ai_lightspeed/capabilities/question_validity/test_capability.py b/tests/unit/pydantic_ai_lightspeed/capabilities/question_validity/test_capability.py index 94b12471d..cae0d2916 100644 --- a/tests/unit/pydantic_ai_lightspeed/capabilities/question_validity/test_capability.py +++ b/tests/unit/pydantic_ai_lightspeed/capabilities/question_validity/test_capability.py @@ -14,6 +14,7 @@ DEFAULT_INVALID_QUESTION_RESPONSE, DEFAULT_MODEL_PROMPT, ) +from models.common.moderation import ShieldModerationBlocked, ShieldModerationPassed from models.config import ( QuestionValidityConfig, ) @@ -532,3 +533,108 @@ async def test_wrap_run_with_sequence_prompt( prompt_str = str(messages[0]) assert "How to" in prompt_str assert "scale a deployment?" in prompt_str + + +class TestQuestionValidityRun: + """Tests for QuestionValidity.run method.""" + + @pytest.fixture(autouse=True) + def _mock_create_model(self, mocker: MockerFixture) -> None: + """Mock model creation for all tests.""" + mocker.patch(f"{_MODULE}.AsyncOgxClientHolder") + mocker.patch(f"{_MODULE}.OgxResponsesModel.from_ogx_client") + + @pytest.mark.asyncio + async def test_allowed_returns_passed(self, mocker: MockerFixture) -> None: + """Test that an allowed response returns ShieldModerationPassed.""" + mock_response = ModelResponse( + parts=[TextPart(content=SUBJECT_ALLOWED)], + usage=RequestUsage(input_tokens=10, output_tokens=1), + ) + mocker.patch(f"{_MODULE}.model_request", return_value=mock_response) + + config = QuestionValidityConfig(model_id="test") + qv = QuestionValidity(config=config) + result = await qv.run("How do I create a pod?") + + assert isinstance(result, ShieldModerationPassed) + assert result.decision == "passed" + + @pytest.mark.asyncio + async def test_rejected_returns_blocked(self, mocker: MockerFixture) -> None: + """Test that a rejected response returns ShieldModerationBlocked.""" + mock_response = ModelResponse( + parts=[TextPart(content=SUBJECT_REJECTED)], + usage=RequestUsage(input_tokens=10, output_tokens=1), + ) + mocker.patch(f"{_MODULE}.model_request", return_value=mock_response) + + config = QuestionValidityConfig(model_id="test") + qv = QuestionValidity(config=config) + result = await qv.run("What is the meaning of life?") + + assert isinstance(result, ShieldModerationBlocked) + assert result.message == DEFAULT_INVALID_QUESTION_RESPONSE + assert result.moderation_id.startswith("modr-") + assert result.refusal_response.role == "assistant" + assert result.refusal_response.content == DEFAULT_INVALID_QUESTION_RESPONSE + + @pytest.mark.asyncio + async def test_unexpected_response_returns_blocked( + self, mocker: MockerFixture + ) -> None: + """Test that an unexpected model response is treated as blocked.""" + mock_response = ModelResponse( + parts=[TextPart(content="I don't understand")], + usage=RequestUsage(input_tokens=10, output_tokens=5), + ) + mocker.patch(f"{_MODULE}.model_request", return_value=mock_response) + + config = QuestionValidityConfig(model_id="test") + qv = QuestionValidity(config=config) + result = await qv.run("some input") + + assert isinstance(result, ShieldModerationBlocked) + assert result.message == DEFAULT_INVALID_QUESTION_RESPONSE + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "response_text", + [" ALLOWED", "ALLOWED ", " ALLOWED ", "ALLOWED\n"], + ids=["leading-space", "trailing-space", "both-spaces", "trailing-newline"], + ) + async def test_allowed_with_whitespace_returns_passed( + self, mocker: MockerFixture, response_text: str + ) -> None: + """Test that ALLOWED with surrounding whitespace still returns passed.""" + mock_response = ModelResponse( + parts=[TextPart(content=response_text)], + usage=RequestUsage(input_tokens=10, output_tokens=1), + ) + mocker.patch(f"{_MODULE}.model_request", return_value=mock_response) + + config = QuestionValidityConfig(model_id="test") + qv = QuestionValidity(config=config) + result = await qv.run("How do I scale pods?") + + assert isinstance(result, ShieldModerationPassed) + + @pytest.mark.asyncio + async def test_custom_invalid_response_message(self, mocker: MockerFixture) -> None: + """Test that a custom rejection message is used in the blocked result.""" + mock_response = ModelResponse( + parts=[TextPart(content=SUBJECT_REJECTED)], + usage=RequestUsage(), + ) + mocker.patch(f"{_MODULE}.model_request", return_value=mock_response) + + config = QuestionValidityConfig( + model_id="test", invalid_question_response="Custom rejection." + ) + qv = QuestionValidity(config=config) + result = await qv.run("off-topic question") + + assert isinstance(result, ShieldModerationBlocked) + assert result.message == "Custom rejection." + assert result.refusal_response.role == "assistant" + assert result.refusal_response.content == "Custom rejection." diff --git a/tests/unit/pydantic_ai_lightspeed/capabilities/redaction/test_capability.py b/tests/unit/pydantic_ai_lightspeed/capabilities/redaction/test_capability.py index c112caba5..cc7f6fe68 100644 --- a/tests/unit/pydantic_ai_lightspeed/capabilities/redaction/test_capability.py +++ b/tests/unit/pydantic_ai_lightspeed/capabilities/redaction/test_capability.py @@ -13,6 +13,7 @@ from pydantic_ai.models import ModelRequestContext from pytest_mock import MockerFixture +from models.common.moderation import ShieldModerationBlocked, ShieldModerationPassed from models.config import ( RedactionConfig, RedactionRule, @@ -314,3 +315,46 @@ async def test_after_model_request_no_match( ) assert result is resp assert resp.parts[0].content == "clean response" + + +class TestPiiRedactionCapabilityRun: + """Tests for PiiRedactionCapability.run method.""" + + @pytest.fixture(name="capability") + def capability_fixture(self) -> PiiRedactionCapability: + """Create a PiiRedactionCapability with an email redaction rule. + + Returns: + A configured PiiRedactionCapability instance. + """ + config = RedactionConfig( + rules=[ + RedactionRule( + pattern=r"[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}", + replacement="[REDACTED_EMAIL]", + ) + ], + case_sensitive=True, + ) + return PiiRedactionCapability(config=config) + + @pytest.mark.asyncio() + async def test_clean_text_returns_passed( + self, capability: PiiRedactionCapability + ) -> None: + """Test that clean text returns ShieldModerationPassed.""" + result = await capability.run("no sensitive content here") + + assert isinstance(result, ShieldModerationPassed) + assert result.decision == "passed" + + @pytest.mark.asyncio() + async def test_pii_text_returns_blocked( + self, capability: PiiRedactionCapability + ) -> None: + """Test that text with PII returns ShieldModerationBlocked.""" + result = await capability.run("contact user@example.com for details") + + assert isinstance(result, ShieldModerationBlocked) + assert result.message == "Sensitive content detected." + assert result.moderation_id.startswith("modr-") diff --git a/tests/unit/utils/agents/test_streaming.py b/tests/unit/utils/agents/test_streaming.py index 7386eef65..3190feb26 100644 --- a/tests/unit/utils/agents/test_streaming.py +++ b/tests/unit/utils/agents/test_streaming.py @@ -1048,9 +1048,9 @@ async def __aexit__(self, *_args: object) -> None: num_chunks = len(chunk_ids) assert chunk_ids == sorted(chunk_ids), "chunk_ids must be monotonically ordered" assert all(cid >= 0 for cid in chunk_ids), "all chunk_ids must be non-negative" - assert num_chunks == len(set(chunk_ids)), ( - "chunk_ids must not contain duplicates" - ) + assert num_chunks == len( + set(chunk_ids) + ), "chunk_ids must not contain duplicates" assert chunk_ids[-1] == num_chunks - 1 @pytest.mark.asyncio From b448aacd6a8e29ef43caad213423f3d062732f92 Mon Sep 17 00:00:00 2001 From: Jazzcort Date: Fri, 24 Jul 2026 16:08:49 -0400 Subject: [PATCH 3/3] Wire run_shield_moderation_v2 into Responses API endpoint Add run_shield_moderation_v2 and build_shield to utils/shields.py to run shield moderation through AbstractSafetyCapability instances instead of the Llama Stack client. Update the responses endpoint to call the new function with shield configs directly. Include unit tests covering pass, block, filtering, and error handling paths. --- src/app/endpoints/responses.py | 10 +- src/utils/shields.py | 86 ++++++++- .../endpoints/test_responses_integration.py | 2 +- tests/unit/app/endpoints/test_responses.py | 2 +- tests/unit/utils/test_shields.py | 168 +++++++++++++++++- 5 files changed, 251 insertions(+), 17 deletions(-) diff --git a/src/app/endpoints/responses.py b/src/app/endpoints/responses.py index 0b8a6cc65..47035eb9b 100644 --- a/src/app/endpoints/responses.py +++ b/src/app/endpoints/responses.py @@ -105,7 +105,7 @@ select_model_for_responses, ) from utils.rh_identity import get_rh_identity_context -from utils.shields import run_shield_moderation +from utils.shields import run_shield_moderation_v2 from utils.suid import ( normalize_conversation_id, ) @@ -424,11 +424,11 @@ async def responses_endpoint_handler( attachments_text = extract_attachments_text(original_request.input) endpoint_path = ENDPOINT_PATH_RESPONSES - moderation_result = await run_shield_moderation( - client, + + moderation_result = await run_shield_moderation_v2( input_text + "\n\n" + attachments_text, - endpoint_path, - original_request.shield_ids, + configuration.configuration.shields, + responses_request.shield_ids, ) filter_server_tools = ( diff --git a/src/utils/shields.py b/src/utils/shields.py index a4470b4dd..e6aafdc0e 100644 --- a/src/utils/shields.py +++ b/src/utils/shields.py @@ -10,8 +10,10 @@ from ogx_client import ( APIStatusError as LLSApiStatusError, ) +from pydantic_ai import ModelAPIError, ModelHTTPError, UnexpectedModelBehavior from configuration import AppConfig +from log import get_logger from models.api.requests import QueryRequest from models.api.responses.error import ( InternalServerErrorResponse, @@ -19,8 +21,20 @@ ServiceUnavailableResponse, UnprocessableEntityResponse, ) -from models.common import ShieldModerationPassed, ShieldModerationResult -from models.config import ShieldConfiguration +from models.common.moderation import ( + ShieldModerationPassed, + ShieldModerationResult, +) +from models.config import QuestionValidityConfig, RedactionConfig, ShieldConfiguration +from pydantic_ai_lightspeed.capabilities.base import AbstractSafetyCapability +from pydantic_ai_lightspeed.capabilities.question_validity._capability import ( + QuestionValidity, +) +from pydantic_ai_lightspeed.capabilities.redaction._capability import ( + PiiRedactionCapability, +) + +logger = get_logger(__name__) def validate_shield_ids_override( @@ -59,6 +73,74 @@ def validate_shield_ids_override( raise HTTPException(**response.model_dump()) +async def run_shield_moderation_v2( + input_text: str, + shield_configs: list[ShieldConfiguration], + selected_shield_ids: Optional[list[str]] = None, +) -> ShieldModerationResult: + """Run v2 shield moderation on input text. + + Iterates through configured shields and runs moderation checks. + + Parameters: + input_text: The text to moderate. + shield_configs: List of shield configurations to evaluate. + selected_shield_ids: Optional list of shield names to filter by. + + Returns: + Result indicating if content was blocked or passed. + + Raises: + HTTPException: 503 if the shield model is unreachable or returns an + HTTP error, 500 if the model returns an unexpected response. + """ + selected_shield_configs = ( + [c for c in shield_configs if c.name in selected_shield_ids] + if selected_shield_ids is not None + else shield_configs + ) + + try: + for shield_config in selected_shield_configs: + shield = build_shield(shield_config) + shield_result = await shield.run(input_text) + + if shield_result.decision == "blocked": + return shield_result + except (ModelHTTPError, ModelAPIError) as e: + logger.error("Shield moderation model request failed: %s", e) + error_response = ServiceUnavailableResponse( + backend_name="shield model", + cause=str(e), + ) + raise HTTPException(**error_response.model_dump()) from e + except UnexpectedModelBehavior as e: + logger.error("Shield moderation received unexpected model response: %s", e) + error_response = InternalServerErrorResponse( + response="Shield moderation failed", + cause=str(e), + ) + raise HTTPException(**error_response.model_dump()) from e + + return ShieldModerationPassed() + + +def build_shield(shield_config: ShieldConfiguration) -> AbstractSafetyCapability: + """Build a safety capability instance from a shield configuration. + + Parameters: + shield_config: The shield configuration to build from. + + Returns: + The constructed safety capability. + """ + match shield_config.config: + case QuestionValidityConfig(): + return QuestionValidity(shield_config.config) + case RedactionConfig(): + return PiiRedactionCapability(shield_config.config) + + async def run_shield_moderation( _client: AsyncOgxClient, _input_text: str, diff --git a/tests/integration/endpoints/test_responses_integration.py b/tests/integration/endpoints/test_responses_integration.py index aae7b9f26..9ef34711f 100644 --- a/tests/integration/endpoints/test_responses_integration.py +++ b/tests/integration/endpoints/test_responses_integration.py @@ -173,7 +173,7 @@ def _configure_shield_blocked( moderation_id=moderation_id, ) mocker.patch( - "app.endpoints.responses.run_shield_moderation", + "app.endpoints.responses.run_shield_moderation_v2", return_value=blocked, ) diff --git a/tests/unit/app/endpoints/test_responses.py b/tests/unit/app/endpoints/test_responses.py index 57020665b..63b4c203f 100644 --- a/tests/unit/app/endpoints/test_responses.py +++ b/tests/unit/app/endpoints/test_responses.py @@ -193,7 +193,7 @@ def _patch_moderation(mocker: MockerFixture, decision: str = "passed") -> Any: else: moderation_result = ShieldModerationPassed() mocker.patch( - f"{MODULE}.run_shield_moderation", + f"{MODULE}.run_shield_moderation_v2", new=mocker.AsyncMock(return_value=moderation_result), ) return moderation_result diff --git a/tests/unit/utils/test_shields.py b/tests/unit/utils/test_shields.py index 7183e3e89..3bd0d5287 100644 --- a/tests/unit/utils/test_shields.py +++ b/tests/unit/utils/test_shields.py @@ -2,17 +2,24 @@ import pytest from fastapi import HTTPException, status +from pydantic_ai import ModelAPIError, ModelHTTPError, UnexpectedModelBehavior from pytest_mock import MockerFixture -from models.config import QuestionValidityConfig, QuestionValidityShieldConfiguration +from models.common.moderation import ShieldModerationBlocked, ShieldModerationPassed +from models.config import ( + QuestionValidityConfig, + QuestionValidityShieldConfiguration, + ShieldConfiguration, +) from utils.shields import ( append_turn_to_conversation, get_shields_for_request, + run_shield_moderation_v2, validate_shield_ids_override, ) -def _shield(name: str) -> QuestionValidityShieldConfiguration: +def _shield_config(name: str) -> QuestionValidityShieldConfiguration: """Build a minimal question-validity shield configuration for tests.""" return QuestionValidityShieldConfiguration( name=name, @@ -133,12 +140,153 @@ def test_raises_422_when_empty_list_shield_ids_and_override_disabled( assert exc_info.value.status_code == status.HTTP_422_UNPROCESSABLE_ENTITY +class TestRunShieldModerationV2: + """Tests for run_shield_moderation_v2 function.""" + + @pytest.mark.asyncio + async def test_returns_passed_when_no_shields(self) -> None: + """Return ShieldModerationPassed when shield list is empty.""" + result = await run_shield_moderation_v2("test input", []) + assert isinstance(result, ShieldModerationPassed) + + @pytest.mark.asyncio + async def test_returns_passed_when_all_shields_pass( + self, mocker: MockerFixture + ) -> None: + """Return ShieldModerationPassed when every shield passes.""" + mock_shield = mocker.Mock() + mock_shield.run = mocker.AsyncMock(return_value=ShieldModerationPassed()) + mocker.patch("utils.shields.build_shield", return_value=mock_shield) + + shields: list[ShieldConfiguration] = [ + _shield_config("s1"), + _shield_config("s2"), + ] + result = await run_shield_moderation_v2("test input", shields) + + assert isinstance(result, ShieldModerationPassed) + assert mock_shield.run.call_count == 2 + + @pytest.mark.asyncio + async def test_returns_blocked_on_first_block(self, mocker: MockerFixture) -> None: + """Return blocked result from first shield that blocks.""" + blocked = ShieldModerationBlocked(message="rejected", moderation_id="modr-123") + mock_shield = mocker.Mock() + mock_shield.run = mocker.AsyncMock(return_value=blocked) + mocker.patch("utils.shields.build_shield", return_value=mock_shield) + + shields: list[ShieldConfiguration] = [ + _shield_config("s1"), + _shield_config("s2"), + ] + result = await run_shield_moderation_v2("test input", shields) + + assert isinstance(result, ShieldModerationBlocked) + assert result.message == "rejected" + mock_shield.run.assert_called_once() + + @pytest.mark.asyncio + async def test_filters_by_selected_shield_ids(self, mocker: MockerFixture) -> None: + """Only run shields matching the selected IDs.""" + mock_shield = mocker.Mock() + mock_shield.run = mocker.AsyncMock(return_value=ShieldModerationPassed()) + mocker.patch("utils.shields.build_shield", return_value=mock_shield) + + shields: list[ShieldConfiguration] = [ + _shield_config("s1"), + _shield_config("s2"), + _shield_config("s3"), + ] + result = await run_shield_moderation_v2( + "test input", shields, selected_shield_ids=["s2"] + ) + + assert isinstance(result, ShieldModerationPassed) + mock_shield.run.assert_called_once() + + @pytest.mark.asyncio + async def test_shields_stops_on_first_block(self, mocker: MockerFixture) -> None: + """Stop at the first blocking shield.""" + blocked = ShieldModerationBlocked(message="rejected", moderation_id="modr-789") + mock_qv_shield = mocker.Mock() + mock_qv_shield.run = mocker.AsyncMock(return_value=blocked) + + mock_redact_shield = mocker.Mock() + mock_redact_shield.run = mocker.AsyncMock(return_value=ShieldModerationPassed()) + + mocker.patch( + "utils.shields.build_shield", + side_effect=[mock_qv_shield, mock_redact_shield], + ) + + shields: list[ShieldConfiguration] = [ + _shield_config("s-1"), + _shield_config("s-2"), + ] + result = await run_shield_moderation_v2("test input", shields) + + assert isinstance(result, ShieldModerationBlocked) + mock_qv_shield.run.assert_called_once() + mock_redact_shield.run.assert_not_called() + + @pytest.mark.asyncio + async def test_raises_503_on_model_http_error(self, mocker: MockerFixture) -> None: + """Raise HTTP 503 when the shield model returns an HTTP error.""" + mock_shield = mocker.Mock() + mock_shield.run = mocker.AsyncMock( + side_effect=ModelHTTPError( + status_code=500, model_name="test-model", body=None + ) + ) + mocker.patch("utils.shields.build_shield", return_value=mock_shield) + + with pytest.raises(HTTPException) as exc_info: + await run_shield_moderation_v2("test input", [_shield_config("s1")]) + + assert exc_info.value.status_code == status.HTTP_503_SERVICE_UNAVAILABLE + + @pytest.mark.asyncio + async def test_raises_503_on_model_api_error(self, mocker: MockerFixture) -> None: + """Raise HTTP 503 when the shield model is unreachable.""" + mock_shield = mocker.Mock() + mock_shield.run = mocker.AsyncMock( + side_effect=ModelAPIError( + model_name="test-model", message="Connection refused" + ) + ) + mocker.patch("utils.shields.build_shield", return_value=mock_shield) + + with pytest.raises(HTTPException) as exc_info: + await run_shield_moderation_v2("test input", [_shield_config("s1")]) + + assert exc_info.value.status_code == status.HTTP_503_SERVICE_UNAVAILABLE + + @pytest.mark.asyncio + async def test_raises_500_on_unexpected_model_behavior( + self, mocker: MockerFixture + ) -> None: + """Raise HTTP 500 when the shield model returns an unexpected response.""" + mock_shield = mocker.Mock() + mock_shield.run = mocker.AsyncMock( + side_effect=UnexpectedModelBehavior("bad response format") + ) + mocker.patch("utils.shields.build_shield", return_value=mock_shield) + + with pytest.raises(HTTPException) as exc_info: + await run_shield_moderation_v2("test input", [_shield_config("s1")]) + + assert exc_info.value.status_code == status.HTTP_500_INTERNAL_SERVER_ERROR + + class TestGetShieldsForRequest: """Tests for get_shields_for_request function.""" def test_returns_all_shields_when_shield_ids_none(self) -> None: """Return all configured shields when shield_ids is None.""" - shields = [_shield("shield-1"), _shield("shield-2")] + shields = [ + _shield_config("shield-1"), + _shield_config("shield-2"), + ] result = get_shields_for_request(shields, shield_ids=None) @@ -146,7 +294,10 @@ def test_returns_all_shields_when_shield_ids_none(self) -> None: def test_returns_empty_list_when_shield_ids_empty(self) -> None: """Return no shields when an empty shield_ids list is provided.""" - shields = [_shield("shield-1"), _shield("shield-2")] + shields = [ + _shield_config("shield-1"), + _shield_config("shield-2"), + ] result = get_shields_for_request(shields, shield_ids=[]) @@ -154,9 +305,9 @@ def test_returns_empty_list_when_shield_ids_empty(self) -> None: def test_filters_to_requested_shields_when_all_exist(self) -> None: """Return only shields whose names appear in shield_ids.""" - shield1 = _shield("shield-1") - shield2 = _shield("shield-2") - shield3 = _shield("shield-3") + shield1 = _shield_config("shield-1") + shield2 = _shield_config("shield-2") + shield3 = _shield_config("shield-3") result = get_shields_for_request( [shield1, shield2, shield3], shield_ids=["shield-1", "shield-3"] @@ -168,7 +319,8 @@ def test_raises_404_when_requested_shield_not_configured(self) -> None: """Raise 404 when a requested shield name is not configured.""" with pytest.raises(HTTPException) as exc_info: get_shields_for_request( - [_shield("shield-1")], shield_ids=["shield-1", "missing-shield"] + [_shield_config("shield-1")], + shield_ids=["shield-1", "missing-shield"], ) assert exc_info.value.status_code == status.HTTP_404_NOT_FOUND