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
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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.
Expand Down
6 changes: 5 additions & 1 deletion src/utils/agents/query.py
Original file line number Diff line number Diff line change
Expand Up @@ -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))
Expand Down
45 changes: 43 additions & 2 deletions src/utils/pydantic_ai_helpers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down Expand Up @@ -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:
Expand All @@ -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.
Expand All @@ -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
Expand Down
148 changes: 147 additions & 1 deletion tests/unit/utils/test_pydantic_ai.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand All @@ -39,13 +60,57 @@ 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."""

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
Expand All @@ -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."""
Expand Down Expand Up @@ -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,
Expand Down
Loading