diff --git a/src/pydantic_ai_lightspeed/capabilities/question_validity/_capability.py b/src/pydantic_ai_lightspeed/capabilities/question_validity/_capability.py index a4f69241a..0a7247a76 100644 --- a/src/pydantic_ai_lightspeed/capabilities/question_validity/_capability.py +++ b/src/pydantic_ai_lightspeed/capabilities/question_validity/_capability.py @@ -12,7 +12,7 @@ from collections.abc import Sequence from dataclasses import dataclass, field from string import Template -from typing import Optional +from typing import Any, Optional from pydantic_ai import AgentRunResult, RunContext from pydantic_ai._agent_graph import GraphAgentState @@ -56,7 +56,7 @@ def _extract_message_str_from_user_content(user_content: Sequence[UserContent]) @dataclass -class QuestionValidity(AbstractCapability[None]): +class QuestionValidity(AbstractCapability[Any]): """Block or modify user input based on a guardrail check. The guard function receives the user prompt and returns True if safe. diff --git a/src/utils/agents/query.py b/src/utils/agents/query.py index 16fa60100..b805fa983 100644 --- a/src/utils/agents/query.py +++ b/src/utils/agents/query.py @@ -316,7 +316,11 @@ async def retrieve_agent_response( ) try: agent = build_agent( - client, responses_params, configuration.skills, no_tools=no_tools + client, + responses_params, + configuration.skills, + shields=configuration.shields, + no_tools=no_tools, ) logger.debug("Starting agent non-streaming response processing") run_result = await agent.run(cast(str, responses_params.input)) diff --git a/src/utils/pydantic_ai_helpers.py b/src/utils/pydantic_ai_helpers.py index bc9515d69..6a61be4a4 100644 --- a/src/utils/pydantic_ai_helpers.py +++ b/src/utils/pydantic_ai_helpers.py @@ -13,7 +13,14 @@ from models.common.responses.responses_api_params import ResponsesApiParams from models.common.tools import CatalogTool, CatalogToolParameter -from models.config import SkillsConfiguration +from models.config import ( + QuestionValidityConfig, + RedactionConfig, + ShieldConfiguration, + SkillsConfiguration, +) +from pydantic_ai_lightspeed.capabilities import QuestionValidity +from pydantic_ai_lightspeed.capabilities.redaction import PiiRedactionCapability from pydantic_ai_lightspeed.llamastack import ( OgxResponsesModel, ) @@ -127,20 +134,51 @@ def get_agent_capability_tools( return tools +def _shield_capability(shield: ShieldConfiguration) -> AgentCapability[object]: + """Build the pydantic-ai capability instance for a single configured shield. + + Parameters: + shield: A single guardrail shield configuration entry. + + Returns: + A ``QuestionValidity`` capability when ``shield.type`` is + ``"question_validity"``, or a ``PiiRedactionCapability`` when it is + ``"redaction"``. + + Raises: + ValueError: If ``shield.config`` doesn't match a known shield config type. + """ + match shield.config: + case QuestionValidityConfig(): + return QuestionValidity(config=shield.config) + case RedactionConfig(): + return PiiRedactionCapability(config=shield.config) + case _: + raise ValueError( + f"Unsupported shield config type for shield '{shield.name}': " + f"{type(shield.config).__name__}" + ) + + def _agent_capabilities( skills: Optional[SkillsConfiguration], + shields: Optional[list[ShieldConfiguration]] = None, no_tools: bool = False, ) -> Optional[list[AgentCapability[object]]]: """Assemble pydantic-ai capabilities for an LCS agent. Args: skills: Agent skills configuration from LCS, or None when skills are disabled. + shields: Configured guardrail shields (question validity, redaction), or + None/empty when no shields are enabled. no_tools: When True, omit capabilities that expose a toolset via ``get_toolset()``. Returns: Configured capabilities, or None when no capabilities are enabled. """ capabilities: list[AgentCapability[object]] = [] + for shield in shields or []: + capabilities.append(_shield_capability(shield)) if skills_capability := _skills_capability(skills): capabilities.append(skills_capability) if no_tools: @@ -159,6 +197,7 @@ def build_agent( client: AsyncOgxClient | AsyncOGXAsLibraryClient, responses_params: ResponsesApiParams, skills: Optional[SkillsConfiguration], + shields: Optional[list[ShieldConfiguration]] = None, no_tools: bool = False, ) -> Agent[None, str]: """Build a Pydantic AI agent that mirrors ``responses_params`` on the Llama Stack backend. @@ -172,13 +211,15 @@ def build_agent( client: Initialized Llama Stack client from ``AsyncOgxClientHolder().get_client()``. responses_params: Parameters produced by ``prepare_responses_params`` for this turn. skills: Agent skills configuration from LCS, or None when skills are disabled. + shields: Configured guardrail shields (question validity, redaction), or + None/empty when no shields are enabled. no_tools: When True, omit capabilities that expose a toolset via ``get_toolset()``. Returns: ``Agent`` configured for ``await agent.run(...)`` (or streaming) against the same stack configuration as ``client.responses.create(**responses_params.model_dump())``. """ - capabilities = _agent_capabilities(skills, no_tools=no_tools) + capabilities = _agent_capabilities(skills, shields, no_tools=no_tools) model = OgxResponsesModel.from_ogx_client( responses_params.model, client, responses_params=responses_params diff --git a/tests/unit/utils/test_pydantic_ai.py b/tests/unit/utils/test_pydantic_ai.py index 4f2716d35..66e8124e0 100644 --- a/tests/unit/utils/test_pydantic_ai.py +++ b/tests/unit/utils/test_pydantic_ai.py @@ -3,20 +3,41 @@ # pylint: disable=protected-access import httpx +import pytest from ogx.core.library_client import AsyncOGXAsLibraryClient from ogx_client import AsyncOgxClient from pydantic_ai_skills import SkillsCapability from pytest_mock import MockerFixture from models.common.responses.responses_api_params import ResponsesApiParams -from models.config import SkillsConfiguration +from models.config import ( + QuestionValidityConfig, + QuestionValidityShieldConfiguration, + RedactionConfig, + RedactionShieldConfiguration, + SkillsConfiguration, +) +from pydantic_ai_lightspeed.capabilities import QuestionValidity +from pydantic_ai_lightspeed.capabilities.redaction import PiiRedactionCapability from utils.pydantic_ai_helpers import ( _agent_capabilities, + _shield_capability, _skills_capability, build_agent, get_agent_capability_tools, ) +_QUESTION_VALIDITY_MODULE = ( + "pydantic_ai_lightspeed.capabilities.question_validity._capability" +) + + +@pytest.fixture(autouse=True) +def _mock_question_validity_model(mocker: MockerFixture) -> None: + """Avoid constructing a real client/model when building QuestionValidity.""" + mocker.patch(f"{_QUESTION_VALIDITY_MODULE}.AsyncOgxClientHolder") + mocker.patch(f"{_QUESTION_VALIDITY_MODULE}.OgxResponsesModel.from_ogx_client") + class TestSkillsCapability: """Tests for _skills_capability.""" @@ -39,6 +60,49 @@ def test_returns_capability_for_configured_paths( assert list(capability.toolset.skills) == ["test-skill"] +class TestShieldCapability: + """Tests for _shield_capability.""" + + def test_question_validity_shield_builds_question_validity_capability( + self, + ) -> None: + """Test that a question_validity shield builds a QuestionValidity capability.""" + shield = QuestionValidityShieldConfiguration( + name="topic-guard", + type="question_validity", + config=QuestionValidityConfig(model_id="test-model"), + ) + + capability = _shield_capability(shield) + + assert isinstance(capability, QuestionValidity) + assert capability.config is shield.config + + def test_redaction_shield_builds_pii_redaction_capability(self) -> None: + """Test that a redaction shield builds a PiiRedactionCapability.""" + shield = RedactionShieldConfiguration( + name="pii-guard", + type="redaction", + config=RedactionConfig(rules=[]), + ) + + capability = _shield_capability(shield) + + assert isinstance(capability, PiiRedactionCapability) + assert capability.config is shield.config + + def test_unsupported_config_type_raises_value_error( + self, mocker: MockerFixture + ) -> None: + """Test that an unrecognized shield config type raises ValueError.""" + shield = mocker.Mock(name="bad-shield") + shield.name = "bad-shield" + shield.config = object() + + with pytest.raises(ValueError, match="Unsupported shield config type"): + _shield_capability(shield) + + class TestAgentCapabilities: """Tests for _agent_capabilities.""" @@ -46,6 +110,7 @@ def test_returns_none_when_no_capabilities_configured(self) -> None: """Test that missing configuration yields None for Agent construction.""" assert _agent_capabilities(None) is None assert _agent_capabilities(SkillsConfiguration(paths=[])) is None + assert _agent_capabilities(None, shields=[]) is None def test_returns_skills_capability_when_configured( self, mock_skills_configuration: SkillsConfiguration @@ -56,6 +121,46 @@ def test_returns_skills_capability_when_configured( assert len(capabilities) == 1 assert isinstance(capabilities[0], SkillsCapability) + def test_returns_shield_capabilities_when_configured(self) -> None: + """Test that configured shields are included in the capability list.""" + shields = [ + QuestionValidityShieldConfiguration( + name="topic-guard", + type="question_validity", + config=QuestionValidityConfig(model_id="test-model"), + ), + RedactionShieldConfiguration( + name="pii-guard", + type="redaction", + config=RedactionConfig(rules=[]), + ), + ] + + capabilities = _agent_capabilities(None, shields=shields) or [] + + assert len(capabilities) == 2 + assert isinstance(capabilities[0], QuestionValidity) + assert isinstance(capabilities[1], PiiRedactionCapability) + + def test_combines_shields_and_skills( + self, mock_skills_configuration: SkillsConfiguration + ) -> None: + """Test that shield and skill capabilities are both included together.""" + shields = [ + RedactionShieldConfiguration( + name="pii-guard", + type="redaction", + config=RedactionConfig(rules=[]), + ), + ] + + capabilities = ( + _agent_capabilities(mock_skills_configuration, shields=shields) or [] + ) + + capability_types = {type(capability) for capability in capabilities} + assert capability_types == {PiiRedactionCapability, SkillsCapability} + class TestBuildAgent: """Tests for the build_agent factory function.""" @@ -163,6 +268,47 @@ def test_agent_has_no_skills_capability_when_not_configured( } assert SkillsCapability not in capability_types + def test_agent_includes_shield_capabilities_when_configured( + self, + mock_client: AsyncOgxClient, + mock_params: ResponsesApiParams, + ) -> None: + """Test that build_agent attaches shield capabilities when shields are passed.""" + shields = [ + QuestionValidityShieldConfiguration( + name="topic-guard", + type="question_validity", + config=QuestionValidityConfig(model_id="test-model"), + ), + RedactionShieldConfiguration( + name="pii-guard", + type="redaction", + config=RedactionConfig(rules=[]), + ), + ] + + agent = build_agent(mock_client, mock_params, None, shields=shields) + + capability_types = { + type(capability) for capability in agent._root_capability.capabilities + } + assert QuestionValidity in capability_types + assert PiiRedactionCapability in capability_types + + def test_agent_has_no_shield_capabilities_when_not_configured( + self, + mock_client: AsyncOgxClient, + mock_params: ResponsesApiParams, + ) -> None: + """Test that build_agent omits shield capabilities when shields are not passed.""" + agent = build_agent(mock_client, mock_params, None) + + capability_types = { + type(capability) for capability in agent._root_capability.capabilities + } + assert QuestionValidity not in capability_types + assert PiiRedactionCapability not in capability_types + def test_agent_excludes_tool_capabilities_when_no_tools( self, mock_client: AsyncOgxClient,