Skip to content
Draft
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
10 changes: 5 additions & 5 deletions src/app/endpoints/responses.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down Expand Up @@ -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 = (
Expand Down
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
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()
102 changes: 84 additions & 18 deletions src/utils/shields.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,25 +3,38 @@
from typing import Optional

from fastapi import HTTPException
from ogx_api import OpenAIResponseMessage
from ogx_client import (
APIConnectionError,
AsyncOgxClient,
)
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,
NotFoundResponse,
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(
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down
7 changes: 1 addition & 6 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,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,
)

Expand Down
10 changes: 1 addition & 9 deletions tests/unit/app/endpoints/test_responses.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
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
Loading
Loading