Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 8 additions & 1 deletion src/models/common/moderation.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,14 @@ class ShieldModerationBlocked(BaseModel):
decision: Literal["blocked"] = "blocked"
message: str
moderation_id: str
refusal_response: ResponseMessage

@property

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Perfect, this is even better than I suggested

def refusal_response(self) -> ResponseMessage:
"""Build a ResponseMessage carrying the shield's refusal text."""
return ResponseMessage(
role="assistant",
content=self.message,
)


ShieldModerationResult = Annotated[
Expand Down
18 changes: 18 additions & 0 deletions src/pydantic_ai_lightspeed/capabilities/base.py
Original file line number Diff line number Diff line change
@@ -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."""
Original file line number Diff line number Diff line change
Expand Up @@ -13,20 +13,27 @@
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
from pydantic_ai.models.openai import OpenAIResponsesModelSettings

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__)
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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()}",
)
21 changes: 19 additions & 2 deletions src/pydantic_ai_lightspeed/capabilities/redaction/_capability.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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()
16 changes: 0 additions & 16 deletions src/utils/shields.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,6 @@
from typing import Optional

from fastapi import HTTPException
from ogx_api import OpenAIResponseMessage
from ogx_client import (
APIConnectionError,
AsyncOgxClient,
Expand Down Expand Up @@ -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,
Expand Down
5 changes: 0 additions & 5 deletions tests/integration/endpoints/test_responses_integration.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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",
Expand Down
8 changes: 0 additions & 8 deletions tests/unit/app/endpoints/test_responses.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -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(),
Expand Down
9 changes: 0 additions & 9 deletions tests/unit/app/endpoints/test_rlsapi_v1.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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",
Expand Down Expand Up @@ -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?",
),
)


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
DEFAULT_INVALID_QUESTION_RESPONSE,
DEFAULT_MODEL_PROMPT,
)
from models.common.moderation import ShieldModerationBlocked, ShieldModerationPassed
from models.config import (
QuestionValidityConfig,
)
Expand Down Expand Up @@ -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."
Loading
Loading