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/models/common/moderation.py b/src/models/common/moderation.py index 7672e23c5..575d6e535 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: + """Build a ResponseMessage carrying the shield's refusal text.""" + return ResponseMessage( + role="assistant", + content=self.message, + ) ShieldModerationResult = Annotated[ 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/src/utils/shields.py b/src/utils/shields.py index e930c009f..e6aafdc0e 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, @@ -11,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, @@ -20,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( @@ -60,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, @@ -130,21 +211,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/integration/endpoints/test_responses_integration.py b/tests/integration/endpoints/test_responses_integration.py index 638277646..9ef34711f 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,13 +171,9 @@ 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", + "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 0ec345038..63b4c203f 100644 --- a/tests/unit/app/endpoints/test_responses.py +++ b/tests/unit/app/endpoints/test_responses.py @@ -189,16 +189,11 @@ 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() mocker.patch( - f"{MODULE}.run_shield_moderation", + f"{MODULE}.run_shield_moderation_v2", new=mocker.AsyncMock(return_value=moderation_result), ) return moderation_result @@ -628,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/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/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_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..3190feb26 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.", - ), ) 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