Skip to content
Open
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
9 changes: 8 additions & 1 deletion src/utils/agents/query.py
Original file line number Diff line number Diff line change
Expand Up @@ -287,6 +287,7 @@ async def retrieve_agent_response(
endpoint_path: str,
_original_input: Optional[ResponseInput] = None,
no_tools: bool = False,
shield_ids: Optional[list[str]] = None,
) -> TurnSummary:
"""Retrieve a turn summary from a blocking agent run.
Expand All @@ -297,6 +298,8 @@ async def retrieve_agent_response(
endpoint_path: Endpoint path used for metric labeling.
_original_input: Original user input before the explicit-input rewrite.
no_tools: Whether to skip tool processing.
shield_ids: Optional list of shield names to run for this turn, mirroring
``QueryRequest.shield_ids``. If ``None``, all configured shields run.
Returns:
Turn summary for the completed agent run.
Expand All @@ -316,7 +319,11 @@ async def retrieve_agent_response(
)
try:
agent = build_agent(
client, responses_params, configuration.skills, no_tools=no_tools
client,
responses_params,
configuration,
shields=shield_ids,
no_tools=no_tools,
)
logger.debug("Starting agent non-streaming response processing")
run_result = await agent.run(cast(str, responses_params.input))
Expand Down
55 changes: 51 additions & 4 deletions src/utils/pydantic_ai_helpers.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,12 +11,21 @@
from pydantic_ai.capabilities import AbstractCapability, AgentCapability
from pydantic_ai_skills import SkillsCapability

from configuration import AppConfig
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,
)
from utils.shields import get_shields_for_request

_AGENT_SKILLS_PROVIDER_ID: Final[str] = "agent-skills"
_AGENT_SKILLS_TOOLGROUP_ID: Final[str] = "builtin::agent-skills"
Expand Down Expand Up @@ -127,20 +136,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.provider_id`` 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 @@ -158,7 +198,8 @@ def _agent_capabilities(
def build_agent(
client: AsyncOgxClient | AsyncOGXAsLibraryClient,
responses_params: ResponsesApiParams,
skills: Optional[SkillsConfiguration],
config: AppConfig,
shields: Optional[list[str]] = None,
no_tools: bool = False,
) -> Agent[None, str]:
"""Build a Pydantic AI agent that mirrors ``responses_params`` on the Llama Stack backend.
Expand All @@ -171,14 +212,20 @@ def build_agent(
Parameters:
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.
config: Application configuration. Agent skills (``config.skills``) and the
configured guardrail shields (``config.shields``) are extracted from it.
shields: Optional list of shield names to run for this turn, matching each
shield's configured ``name``. Mirrors ``QueryRequest.shield_ids``: if
``None``, all shields configured in ``config.shields`` run; an empty
list disables all shields.
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)
shield_configs = get_shields_for_request(config.shields, shields)
capabilities = _agent_capabilities(config.skills, shield_configs, no_tools=no_tools)

model = OgxResponsesModel.from_ogx_client(
responses_params.model, client, responses_params=responses_params
Expand Down
28 changes: 26 additions & 2 deletions tests/unit/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,8 +3,9 @@
from __future__ import annotations

import logging
from collections.abc import Generator
from collections.abc import Callable, Generator
from pathlib import Path
from typing import Optional

import httpx
import pytest
Expand All @@ -14,7 +15,7 @@
from configuration import AppConfig
from constants import DEFAULT_LOGGER_NAME
from models.common.responses.responses_api_params import ResponsesApiParams
from models.config import SkillsConfiguration
from models.config import ShieldConfiguration, SkillsConfiguration

type AgentFixtures = Generator[
tuple[
Expand Down Expand Up @@ -143,3 +144,26 @@ def mock_skills_configuration_fixture(tmp_path: Path) -> SkillsConfiguration:
encoding="utf-8",
)
return SkillsConfiguration(paths=[skills_root])


@pytest.fixture(name="make_agent_config")
def make_agent_config_fixture(
mocker: MockerFixture,
) -> Callable[..., AppConfig]:
"""Return a factory building a duck-typed AppConfig stand-in for build_agent.

``build_agent`` only reads ``config.skills`` and ``config.shields`` off the
config object it receives, so tests can pass a lightweight mock instead of
a fully-initialized ``AppConfig``.
"""

def _make(
skills: Optional[SkillsConfiguration] = None,
shields: Optional[list[ShieldConfiguration]] = None,
) -> AppConfig:
config = mocker.Mock()
config.skills = skills
config.shields = shields or []
return config

return _make
Loading
Loading