Skip to content

Commit 47244fc

Browse files
authored
Merge pull request #2220 from Jazzcort/add-abstract-safety-capability
LCORE-3201: Expand Shield Interface for direct running
2 parents ec29633 + eb34920 commit 47244fc

12 files changed

Lines changed: 218 additions & 57 deletions

File tree

src/models/common/moderation.py

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -20,7 +20,14 @@ class ShieldModerationBlocked(BaseModel):
2020
decision: Literal["blocked"] = "blocked"
2121
message: str
2222
moderation_id: str
23-
refusal_response: ResponseMessage
23+
24+
@property
25+
def refusal_response(self) -> ResponseMessage:
26+
"""Build a ResponseMessage carrying the shield's refusal text."""
27+
return ResponseMessage(
28+
role="assistant",
29+
content=self.message,
30+
)
2431

2532

2633
ShieldModerationResult = Annotated[
Lines changed: 18 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,18 @@
1+
"""Abstract base for safety capabilities with a standalone run interface."""
2+
3+
from abc import abstractmethod
4+
5+
from pydantic_ai.capabilities import AbstractCapability
6+
from typing_extensions import TypeVar
7+
8+
from models.common.moderation import ShieldModerationResult
9+
10+
T = TypeVar("T", default=object)
11+
12+
13+
class AbstractSafetyCapability(AbstractCapability[T]):
14+
"""Interface for safety/moderation that can be called directly."""
15+
16+
@abstractmethod
17+
async def run(self, input_text: str) -> ShieldModerationResult:
18+
"""Run moderation on input text."""

src/pydantic_ai_lightspeed/capabilities/question_validity/_capability.py

Lines changed: 23 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -13,20 +13,27 @@
1313
from dataclasses import dataclass, field
1414
from string import Template
1515
from typing import Optional
16+
from uuid import uuid4
1617

1718
from pydantic_ai import AgentRunResult, RunContext
1819
from pydantic_ai._agent_graph import GraphAgentState
19-
from pydantic_ai.capabilities import AbstractCapability, WrapRunHandler
20+
from pydantic_ai.capabilities import WrapRunHandler
2021
from pydantic_ai.direct import model_request
2122
from pydantic_ai.messages import ModelRequest, TextContent, UserContent
2223
from pydantic_ai.models import Model
2324
from pydantic_ai.models.openai import OpenAIResponsesModelSettings
2425

2526
from client import AsyncOgxClientHolder
2627
from log import get_logger
28+
from models.common.moderation import (
29+
ShieldModerationBlocked,
30+
ShieldModerationPassed,
31+
ShieldModerationResult,
32+
)
2733
from models.config import (
2834
QuestionValidityConfig,
2935
)
36+
from pydantic_ai_lightspeed.capabilities.base import AbstractSafetyCapability
3037
from pydantic_ai_lightspeed.llamastack import OgxResponsesModel
3138

3239
logger = get_logger(__name__)
@@ -56,7 +63,7 @@ def _extract_message_str_from_user_content(user_content: Sequence[UserContent])
5663

5764

5865
@dataclass
59-
class QuestionValidity(AbstractCapability[None]):
66+
class QuestionValidity(AbstractSafetyCapability):
6067
"""Block or modify user input based on a guardrail check.
6168
6269
The guard function receives the user prompt and returns True if safe.
@@ -140,3 +147,17 @@ async def wrap_run(
140147
return AgentRunResult(
141148
output=self.config.invalid_question_response, _state=state
142149
)
150+
151+
async def run(self, input_text: str) -> ShieldModerationResult:
152+
"""Run question-validity check and return a moderation result."""
153+
result = await model_request(
154+
model=self._model, messages=[ModelRequest.user_text_prompt(input_text)]
155+
)
156+
157+
if result.text is not None and result.text.strip() == SUBJECT_ALLOWED:
158+
return ShieldModerationPassed()
159+
160+
return ShieldModerationBlocked(
161+
message=self.config.invalid_question_response,
162+
moderation_id=f"modr-{uuid4()}",
163+
)

src/pydantic_ai_lightspeed/capabilities/redaction/_capability.py

Lines changed: 19 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -3,9 +3,9 @@
33
from collections.abc import Sequence
44
from dataclasses import dataclass, replace
55
from typing import Any, Optional
6+
from uuid import uuid4
67

78
from pydantic_ai import RunContext
8-
from pydantic_ai.capabilities import AbstractCapability
99
from pydantic_ai.messages import (
1010
ModelMessage,
1111
ModelRequest,
@@ -19,7 +19,13 @@
1919
)
2020
from pydantic_ai.models import ModelRequestContext
2121

22+
from models.common.moderation import (
23+
ShieldModerationBlocked,
24+
ShieldModerationPassed,
25+
ShieldModerationResult,
26+
)
2227
from models.config import RedactionConfig
28+
from pydantic_ai_lightspeed.capabilities.base import AbstractSafetyCapability
2329
from pydantic_ai_lightspeed.capabilities.redaction.core import (
2430
CompiledPatterns,
2531
redact_text,
@@ -257,7 +263,7 @@ def _redact_response(
257263

258264

259265
@dataclass
260-
class PiiRedactionCapability(AbstractCapability[Any]):
266+
class PiiRedactionCapability(AbstractSafetyCapability):
261267
"""Pydantic AI capability that redacts PII from agent messages.
262268
263269
Applies configurable regex-based redaction rules to user prompt
@@ -321,3 +327,14 @@ async def after_model_request(
321327
)
322328

323329
return new_response
330+
331+
async def run(self, input_text: str) -> ShieldModerationResult:
332+
"""Run PII redaction on input text and return a moderation result."""
333+
result = redact_text(input_text, self.config.compiled_patterns)
334+
335+
if result.redacted:
336+
return ShieldModerationBlocked(
337+
message="Sensitive content detected.", moderation_id=f"modr-{uuid4()}"
338+
)
339+
340+
return ShieldModerationPassed()

src/utils/shields.py

Lines changed: 0 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -3,7 +3,6 @@
33
from typing import Optional
44

55
from fastapi import HTTPException
6-
from ogx_api import OpenAIResponseMessage
76
from ogx_client import (
87
APIConnectionError,
98
AsyncOgxClient,
@@ -130,21 +129,6 @@ async def append_turn_to_conversation(
130129
raise HTTPException(**error_response.model_dump()) from e
131130

132131

133-
def create_refusal_response(refusal_message: str) -> OpenAIResponseMessage:
134-
"""Create a refusal response message object.
135-
136-
Args:
137-
refusal_message: The refusal message text.
138-
139-
Returns:
140-
OpenAIResponseMessage with refusal message.
141-
"""
142-
return OpenAIResponseMessage(
143-
role="assistant",
144-
content=refusal_message,
145-
)
146-
147-
148132
def get_shields_for_request(
149133
shields: list[ShieldConfiguration],
150134
shield_ids: Optional[list[str]] = None,

tests/integration/endpoints/test_responses_integration.py

Lines changed: 0 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,6 @@
1111
import pytest
1212
from fastapi import Request
1313
from fastapi.responses import StreamingResponse
14-
from ogx_api.openai_responses import OpenAIResponseMessage
1514
from ogx_client.types import ListModelsResponse
1615
from ogx_client.types.model import Model
1716
from pytest_mock import MockerFixture
@@ -172,10 +171,6 @@ def _configure_shield_blocked(
172171
blocked = ShieldModerationBlocked(
173172
message="Content blocked by safety shield",
174173
moderation_id=moderation_id,
175-
refusal_response=OpenAIResponseMessage(
176-
role="assistant",
177-
content="Content blocked by safety shield",
178-
),
179174
)
180175
mocker.patch(
181176
"app.endpoints.responses.run_shield_moderation",

tests/unit/app/endpoints/test_responses.py

Lines changed: 0 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -189,11 +189,6 @@ def _patch_moderation(mocker: MockerFixture, decision: str = "passed") -> Any:
189189
moderation_result = ShieldModerationBlocked(
190190
message="Content blocked",
191191
moderation_id="mod_blocked",
192-
refusal_response=OpenAIResponseMessage(
193-
role="assistant",
194-
content="Content blocked",
195-
type="message",
196-
),
197192
)
198193
else:
199194
moderation_result = ShieldModerationPassed()
@@ -628,9 +623,6 @@ async def test_responses_blocked_with_conversation_appends_refusal(
628623
mock_moderation = _patch_moderation(mocker, decision="blocked")
629624
mock_moderation.message = "Blocked"
630625
mock_moderation.moderation_id = "resp_blocked_123"
631-
mock_moderation.refusal_response = OpenAIResponseMessage(
632-
type="message", role="assistant", content="Blocked"
633-
)
634626
mock_append = mocker.patch(
635627
f"{MODULE}.append_turn_items_to_conversation",
636628
new=mocker.AsyncMock(),

tests/unit/app/endpoints/test_rlsapi_v1.py

Lines changed: 0 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,6 @@
1313

1414
import pytest
1515
from fastapi import HTTPException, status
16-
from ogx_api import OpenAIResponseMessage
1716
from ogx_client import APIConnectionError, APIStatusError
1817
from ogx_client.types import ListModelsResponse
1918
from ogx_client.types.model import Model
@@ -1254,10 +1253,6 @@ async def test_infer_quota_shield_blocked_does_not_consume_tokens(
12541253
blocked = ShieldModerationBlocked(
12551254
message="Blocked by moderation",
12561255
moderation_id="modr-test",
1257-
refusal_response=OpenAIResponseMessage(
1258-
role="assistant",
1259-
content="Blocked by moderation",
1260-
),
12611256
)
12621257
mocker.patch(
12631258
"app.endpoints.rlsapi_v1.run_shield_moderation",
@@ -1287,10 +1282,6 @@ def _create_blocked_moderation_result() -> ShieldModerationBlocked:
12871282
return ShieldModerationBlocked(
12881283
message="I can't answer that. Can I help with something else?",
12891284
moderation_id="modr-test-123",
1290-
refusal_response=OpenAIResponseMessage(
1291-
role="assistant",
1292-
content="I can't answer that. Can I help with something else?",
1293-
),
12941285
)
12951286

12961287

tests/unit/pydantic_ai_lightspeed/capabilities/question_validity/test_capability.py

Lines changed: 106 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,7 @@
1414
DEFAULT_INVALID_QUESTION_RESPONSE,
1515
DEFAULT_MODEL_PROMPT,
1616
)
17+
from models.common.moderation import ShieldModerationBlocked, ShieldModerationPassed
1718
from models.config import (
1819
QuestionValidityConfig,
1920
)
@@ -532,3 +533,108 @@ async def test_wrap_run_with_sequence_prompt(
532533
prompt_str = str(messages[0])
533534
assert "How to" in prompt_str
534535
assert "scale a deployment?" in prompt_str
536+
537+
538+
class TestQuestionValidityRun:
539+
"""Tests for QuestionValidity.run method."""
540+
541+
@pytest.fixture(autouse=True)
542+
def _mock_create_model(self, mocker: MockerFixture) -> None:
543+
"""Mock model creation for all tests."""
544+
mocker.patch(f"{_MODULE}.AsyncOgxClientHolder")
545+
mocker.patch(f"{_MODULE}.OgxResponsesModel.from_ogx_client")
546+
547+
@pytest.mark.asyncio
548+
async def test_allowed_returns_passed(self, mocker: MockerFixture) -> None:
549+
"""Test that an allowed response returns ShieldModerationPassed."""
550+
mock_response = ModelResponse(
551+
parts=[TextPart(content=SUBJECT_ALLOWED)],
552+
usage=RequestUsage(input_tokens=10, output_tokens=1),
553+
)
554+
mocker.patch(f"{_MODULE}.model_request", return_value=mock_response)
555+
556+
config = QuestionValidityConfig(model_id="test")
557+
qv = QuestionValidity(config=config)
558+
result = await qv.run("How do I create a pod?")
559+
560+
assert isinstance(result, ShieldModerationPassed)
561+
assert result.decision == "passed"
562+
563+
@pytest.mark.asyncio
564+
async def test_rejected_returns_blocked(self, mocker: MockerFixture) -> None:
565+
"""Test that a rejected response returns ShieldModerationBlocked."""
566+
mock_response = ModelResponse(
567+
parts=[TextPart(content=SUBJECT_REJECTED)],
568+
usage=RequestUsage(input_tokens=10, output_tokens=1),
569+
)
570+
mocker.patch(f"{_MODULE}.model_request", return_value=mock_response)
571+
572+
config = QuestionValidityConfig(model_id="test")
573+
qv = QuestionValidity(config=config)
574+
result = await qv.run("What is the meaning of life?")
575+
576+
assert isinstance(result, ShieldModerationBlocked)
577+
assert result.message == DEFAULT_INVALID_QUESTION_RESPONSE
578+
assert result.moderation_id.startswith("modr-")
579+
assert result.refusal_response.role == "assistant"
580+
assert result.refusal_response.content == DEFAULT_INVALID_QUESTION_RESPONSE
581+
582+
@pytest.mark.asyncio
583+
async def test_unexpected_response_returns_blocked(
584+
self, mocker: MockerFixture
585+
) -> None:
586+
"""Test that an unexpected model response is treated as blocked."""
587+
mock_response = ModelResponse(
588+
parts=[TextPart(content="I don't understand")],
589+
usage=RequestUsage(input_tokens=10, output_tokens=5),
590+
)
591+
mocker.patch(f"{_MODULE}.model_request", return_value=mock_response)
592+
593+
config = QuestionValidityConfig(model_id="test")
594+
qv = QuestionValidity(config=config)
595+
result = await qv.run("some input")
596+
597+
assert isinstance(result, ShieldModerationBlocked)
598+
assert result.message == DEFAULT_INVALID_QUESTION_RESPONSE
599+
600+
@pytest.mark.asyncio
601+
@pytest.mark.parametrize(
602+
"response_text",
603+
[" ALLOWED", "ALLOWED ", " ALLOWED ", "ALLOWED\n"],
604+
ids=["leading-space", "trailing-space", "both-spaces", "trailing-newline"],
605+
)
606+
async def test_allowed_with_whitespace_returns_passed(
607+
self, mocker: MockerFixture, response_text: str
608+
) -> None:
609+
"""Test that ALLOWED with surrounding whitespace still returns passed."""
610+
mock_response = ModelResponse(
611+
parts=[TextPart(content=response_text)],
612+
usage=RequestUsage(input_tokens=10, output_tokens=1),
613+
)
614+
mocker.patch(f"{_MODULE}.model_request", return_value=mock_response)
615+
616+
config = QuestionValidityConfig(model_id="test")
617+
qv = QuestionValidity(config=config)
618+
result = await qv.run("How do I scale pods?")
619+
620+
assert isinstance(result, ShieldModerationPassed)
621+
622+
@pytest.mark.asyncio
623+
async def test_custom_invalid_response_message(self, mocker: MockerFixture) -> None:
624+
"""Test that a custom rejection message is used in the blocked result."""
625+
mock_response = ModelResponse(
626+
parts=[TextPart(content=SUBJECT_REJECTED)],
627+
usage=RequestUsage(),
628+
)
629+
mocker.patch(f"{_MODULE}.model_request", return_value=mock_response)
630+
631+
config = QuestionValidityConfig(
632+
model_id="test", invalid_question_response="Custom rejection."
633+
)
634+
qv = QuestionValidity(config=config)
635+
result = await qv.run("off-topic question")
636+
637+
assert isinstance(result, ShieldModerationBlocked)
638+
assert result.message == "Custom rejection."
639+
assert result.refusal_response.role == "assistant"
640+
assert result.refusal_response.content == "Custom rejection."

0 commit comments

Comments
 (0)