From fec4318254402e529ad9af7f6b438e4e8d0764cf Mon Sep 17 00:00:00 2001 From: Jazzcort Date: Thu, 23 Jul 2026 15:18:48 -0400 Subject: [PATCH 1/2] 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/2] 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