Skip to content

Commit 02afe55

Browse files
committed
fix(shields): fixed QuestionValidity shield
1 parent e6bdda7 commit 02afe55

2 files changed

Lines changed: 215 additions & 12 deletions

File tree

src/pydantic_ai_lightspeed/capabilities/question_validity/_capability.py

Lines changed: 78 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -19,7 +19,13 @@
1919
from pydantic_ai._agent_graph import GraphAgentState
2020
from pydantic_ai.capabilities import WrapRunHandler
2121
from pydantic_ai.direct import model_request
22-
from pydantic_ai.messages import ModelRequest, TextContent, UserContent
22+
from pydantic_ai.messages import (
23+
ModelRequest,
24+
ModelResponse,
25+
TextContent,
26+
TextPart,
27+
UserContent,
28+
)
2329
from pydantic_ai.models import Model
2430
from pydantic_ai.models.openai import OpenAIResponsesModelSettings
2531

@@ -35,6 +41,7 @@
3541
)
3642
from pydantic_ai_lightspeed.capabilities.base import AbstractSafetyCapability
3743
from pydantic_ai_lightspeed.llamastack import OgxResponsesModel
44+
from utils.shields import append_turn_to_conversation
3845

3946
logger = get_logger(__name__)
4047

@@ -62,6 +69,47 @@ def _extract_message_str_from_user_content(user_content: Sequence[UserContent])
6269
return "\n".join(str_arr)
6370

6471

72+
def _message_to_str(message: Optional[str | Sequence[UserContent]]) -> str:
73+
"""Convert a user message (string, content sequence, or None) to plain text.
74+
75+
Parameters:
76+
message: The user input as a string, sequence of user content, or None.
77+
78+
Returns:
79+
A plain-text representation of the message, or an empty string for None.
80+
"""
81+
match message:
82+
case str() as s:
83+
return s
84+
case Sequence() as seq:
85+
return _extract_message_str_from_user_content(seq)
86+
case None:
87+
return ""
88+
89+
90+
def _extract_conversation_id(model: Model) -> Optional[str]:
91+
"""Extract the Llama Stack conversation ID from the agent's model settings.
92+
93+
The main agent's model is built with ``conversation`` in its
94+
``extra_body`` model settings (see ``OgxResponsesModel.from_ogx_client``).
95+
This pulls it back out so the capability can persist the rejected turn
96+
to the same conversation.
97+
98+
Parameters:
99+
model: The model bound to the current agent run (``ctx.model``).
100+
101+
Returns:
102+
The conversation ID, or None if the model has no such setting
103+
(e.g. when used outside a Llama Stack-backed agent).
104+
"""
105+
extra_body = (model.settings or {}).get("extra_body")
106+
if not isinstance(extra_body, dict):
107+
return None
108+
109+
conversation_id = extra_body.get("conversation")
110+
return conversation_id if isinstance(conversation_id, str) else None
111+
112+
65113
@dataclass
66114
class QuestionValidity(AbstractSafetyCapability):
67115
"""Block or modify user input based on a guardrail check.
@@ -100,16 +148,10 @@ def _build_prompt(self, message: Optional[str | Sequence[UserContent]]) -> str:
100148
Returns:
101149
The rendered prompt string ready to send to the validity model.
102150
"""
103-
match message:
104-
case str() as s:
105-
_message = s
106-
case Sequence() as seq:
107-
_message = _extract_message_str_from_user_content(seq)
108-
case None:
109-
_message = ""
110-
111151
return Template(self.config.model_prompt).substitute(
112-
message=_message, allowed=SUBJECT_ALLOWED, rejected=SUBJECT_REJECTED
152+
message=_message_to_str(message),
153+
allowed=SUBJECT_ALLOWED,
154+
rejected=SUBJECT_REJECTED,
113155
)
114156

115157
async def wrap_run(
@@ -143,7 +185,32 @@ async def wrap_run(
143185
return await handler() # proceed with the real run
144186

145187
# short-circuit: return the rejection message with shield usage tracked
146-
state = GraphAgentState(usage=ctx.usage)
188+
user_message = _message_to_str(ctx.prompt)
189+
state = GraphAgentState(
190+
usage=ctx.usage,
191+
message_history=[
192+
ModelRequest.user_text_prompt(user_message),
193+
ModelResponse(
194+
[TextPart(self.config.invalid_question_response)],
195+
finish_reason="stop",
196+
),
197+
],
198+
)
199+
200+
conversation_id = _extract_conversation_id(ctx.model)
201+
if conversation_id is not None:
202+
await append_turn_to_conversation(
203+
AsyncOgxClientHolder().get_client(),
204+
conversation_id,
205+
user_message,
206+
self.config.invalid_question_response,
207+
)
208+
else:
209+
logger.warning(
210+
"Unable to determine conversation ID from model settings; "
211+
"skipping v1/conversation persistence for rejected question."
212+
)
213+
147214
return AgentRunResult(
148215
output=self.config.invalid_question_response, _state=state
149216
)

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

Lines changed: 137 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,7 @@
2222
SUBJECT_ALLOWED,
2323
SUBJECT_REJECTED,
2424
QuestionValidity,
25+
_extract_conversation_id,
2526
_extract_message_str_from_user_content,
2627
)
2728

@@ -65,6 +66,47 @@ def test_sequence_with_non_text_content(self) -> None:
6566
assert result == "keep"
6667

6768

69+
class TestExtractConversationId:
70+
"""Tests for _extract_conversation_id helper."""
71+
72+
def test_extracts_conversation_id(self, mocker: MockerFixture) -> None:
73+
"""Test extraction when extra_body.conversation is set."""
74+
model = mocker.Mock()
75+
model.settings = {"extra_body": {"conversation": "conv_123"}}
76+
77+
assert _extract_conversation_id(model) == "conv_123"
78+
79+
def test_returns_none_when_settings_missing(self, mocker: MockerFixture) -> None:
80+
"""Test that None settings yields None."""
81+
model = mocker.Mock()
82+
model.settings = None
83+
84+
assert _extract_conversation_id(model) is None
85+
86+
def test_returns_none_when_extra_body_missing(self, mocker: MockerFixture) -> None:
87+
"""Test that missing extra_body yields None."""
88+
model = mocker.Mock()
89+
model.settings = {}
90+
91+
assert _extract_conversation_id(model) is None
92+
93+
def test_returns_none_when_extra_body_not_dict(self, mocker: MockerFixture) -> None:
94+
"""Test that a non-dict extra_body yields None instead of raising."""
95+
model = mocker.Mock()
96+
model.settings = {"extra_body": "not-a-dict"}
97+
98+
assert _extract_conversation_id(model) is None
99+
100+
def test_returns_none_when_conversation_not_string(
101+
self, mocker: MockerFixture
102+
) -> None:
103+
"""Test that a non-string conversation value yields None."""
104+
model = mocker.Mock()
105+
model.settings = {"extra_body": {"conversation": 123}}
106+
107+
assert _extract_conversation_id(model) is None
108+
109+
68110
class TestQuestionValidityConfigInit:
69111
"""Tests for QuestionValidityConfig initialization."""
70112

@@ -216,12 +258,21 @@ def _mock_create_model(self, mocker: MockerFixture) -> None:
216258
mocker.patch(f"{_MODULE}.AsyncOgxClientHolder")
217259
mocker.patch(f"{_MODULE}.OgxResponsesModel.from_ogx_client")
218260

261+
@pytest.fixture(name="mock_append_turn", autouse=True)
262+
def mock_append_turn_fixture(self, mocker: MockerFixture) -> MockType:
263+
"""Mock the conversation-persistence call used on rejection."""
264+
return mocker.patch(
265+
f"{_MODULE}.append_turn_to_conversation", new_callable=mocker.AsyncMock
266+
)
267+
219268
@pytest.fixture(name="mock_ctx")
220269
def mock_ctx_fixture(self, mocker: MockerFixture) -> RunContext:
221-
"""Create a mock RunContext."""
270+
"""Create a mock RunContext bound to a model with a conversation ID."""
222271
ctx = mocker.Mock(spec=RunContext)
223272
ctx.prompt = "How do I create a pod?"
224273
ctx.usage = RunUsage()
274+
ctx.model = mocker.Mock()
275+
ctx.model.settings = {"extra_body": {"conversation": "conv_test"}}
225276
return ctx
226277

227278
@pytest.fixture(name="mock_handler")
@@ -280,6 +331,89 @@ async def test_rejected_question_returns_rejection(
280331
assert isinstance(result, AgentRunResult)
281332
assert result.output == DEFAULT_INVALID_QUESTION_RESPONSE
282333

334+
@pytest.mark.asyncio
335+
async def test_rejected_question_persists_turn_to_conversation(
336+
self,
337+
mocker: MockerFixture,
338+
mock_ctx: RunContext,
339+
mock_handler: MockType,
340+
mock_append_turn: MockType,
341+
) -> None:
342+
"""Test that a rejection appends the user question and refusal to the conversation."""
343+
mock_client = mocker.Mock()
344+
mocker.patch(
345+
f"{_MODULE}.AsyncOgxClientHolder"
346+
).return_value.get_client.return_value = mock_client
347+
mock_response = ModelResponse(
348+
parts=[TextPart(content=SUBJECT_REJECTED)],
349+
usage=RequestUsage(input_tokens=10, output_tokens=1),
350+
)
351+
mocker.patch(
352+
"pydantic_ai_lightspeed.capabilities.question_validity._capability.model_request",
353+
return_value=mock_response,
354+
)
355+
356+
config = QuestionValidityConfig(model_id="test")
357+
qv = QuestionValidity(config=config)
358+
await qv.wrap_run(mock_ctx, handler=mock_handler)
359+
360+
mock_append_turn.assert_awaited_once_with(
361+
mock_client,
362+
"conv_test",
363+
"How do I create a pod?",
364+
DEFAULT_INVALID_QUESTION_RESPONSE,
365+
)
366+
367+
@pytest.mark.asyncio
368+
async def test_rejection_skips_persistence_when_conversation_id_missing(
369+
self,
370+
mocker: MockerFixture,
371+
mock_ctx: RunContext,
372+
mock_handler: MockType,
373+
mock_append_turn: MockType,
374+
) -> None:
375+
"""Test that persistence is skipped (not crashed) without a conversation ID."""
376+
mock_ctx.model = mocker.Mock(settings={})
377+
mock_response = ModelResponse(
378+
parts=[TextPart(content=SUBJECT_REJECTED)],
379+
usage=RequestUsage(),
380+
)
381+
mocker.patch(
382+
"pydantic_ai_lightspeed.capabilities.question_validity._capability.model_request",
383+
return_value=mock_response,
384+
)
385+
386+
config = QuestionValidityConfig(model_id="test")
387+
qv = QuestionValidity(config=config)
388+
result = await qv.wrap_run(mock_ctx, handler=mock_handler)
389+
390+
mock_append_turn.assert_not_awaited()
391+
assert result.output == DEFAULT_INVALID_QUESTION_RESPONSE
392+
393+
@pytest.mark.asyncio
394+
async def test_allowed_question_does_not_persist_turn(
395+
self,
396+
mocker: MockerFixture,
397+
mock_ctx: RunContext,
398+
mock_handler: MockType,
399+
mock_append_turn: MockType,
400+
) -> None:
401+
"""Test that an allowed question does not touch the conversation."""
402+
mock_response = ModelResponse(
403+
parts=[TextPart(content=SUBJECT_ALLOWED)],
404+
usage=RequestUsage(input_tokens=10, output_tokens=1),
405+
)
406+
mocker.patch(
407+
"pydantic_ai_lightspeed.capabilities.question_validity._capability.model_request",
408+
return_value=mock_response,
409+
)
410+
411+
config = QuestionValidityConfig(model_id="test")
412+
qv = QuestionValidity(config=config)
413+
await qv.wrap_run(mock_ctx, handler=mock_handler)
414+
415+
mock_append_turn.assert_not_awaited()
416+
283417
@pytest.mark.asyncio
284418
async def test_unexpected_response_treated_as_rejected(
285419
self,
@@ -469,6 +603,8 @@ async def test_wrap_run_with_none_prompt(
469603
ctx = mocker.Mock(spec=RunContext)
470604
ctx.prompt = None
471605
ctx.usage = RunUsage()
606+
ctx.model = mocker.Mock()
607+
ctx.model.settings = {"extra_body": {"conversation": "conv_test"}}
472608

473609
mock_response = ModelResponse(
474610
parts=[TextPart(content=SUBJECT_REJECTED)],

0 commit comments

Comments
 (0)