Skip to content

Commit 95c6cea

Browse files
committed
feat(shields): wire shields into build_agent
1 parent b2ab480 commit 95c6cea

6 files changed

Lines changed: 338 additions & 21 deletions

File tree

src/app/endpoints/a2a.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -372,7 +372,7 @@ async def _process_task_streaming( # pylint: disable=too-many-locals
372372
)
373373
responses_params = compaction.params
374374

375-
agent = build_agent(client, responses_params, configuration.skills)
375+
agent = build_agent(client, responses_params, configuration)
376376
except (AgentRunError, APIStatusError, APIConnectionError, RuntimeError) as e:
377377
error_response = map_agent_inference_error(e, query_request.model or "")
378378
logger.error("Error preparing A2A agent: %s", str(e), exc_info=True)

src/utils/agents/query.py

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -287,6 +287,7 @@ async def retrieve_agent_response(
287287
endpoint_path: str,
288288
_original_input: Optional[ResponseInput] = None,
289289
no_tools: bool = False,
290+
shield_ids: Optional[list[str]] = None,
290291
) -> TurnSummary:
291292
"""Retrieve a turn summary from a blocking agent run.
292293
@@ -297,6 +298,8 @@ async def retrieve_agent_response(
297298
endpoint_path: Endpoint path used for metric labeling.
298299
_original_input: Original user input before the explicit-input rewrite.
299300
no_tools: Whether to skip tool processing.
301+
shield_ids: Optional list of shield names to run for this turn, mirroring
302+
``QueryRequest.shield_ids``. If ``None``, all configured shields run.
300303
Returns:
301304
Turn summary for the completed agent run.
302305
@@ -316,7 +319,11 @@ async def retrieve_agent_response(
316319
)
317320
try:
318321
agent = build_agent(
319-
client, responses_params, configuration.skills, no_tools=no_tools
322+
client,
323+
responses_params,
324+
configuration,
325+
shields=shield_ids,
326+
no_tools=no_tools,
320327
)
321328
logger.debug("Starting agent non-streaming response processing")
322329
run_result = await agent.run(cast(str, responses_params.input))

src/utils/agents/streaming.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -121,7 +121,7 @@ async def retrieve_agent_response_generator(
121121
)
122122

123123
agent = build_agent(
124-
context.client, responses_params, configuration.skills, no_tools=no_tools
124+
context.client, responses_params, configuration, no_tools=no_tools
125125
)
126126

127127
return (

src/utils/pydantic_ai_helpers.py

Lines changed: 51 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -11,12 +11,21 @@
1111
from pydantic_ai.capabilities import AbstractCapability, AgentCapability
1212
from pydantic_ai_skills import SkillsCapability
1313

14+
from configuration import AppConfig
1415
from models.common.responses.responses_api_params import ResponsesApiParams
1516
from models.common.tools import CatalogTool, CatalogToolParameter
16-
from models.config import SkillsConfiguration
17+
from models.config import (
18+
QuestionValidityConfig,
19+
RedactionConfig,
20+
ShieldConfiguration,
21+
SkillsConfiguration,
22+
)
23+
from pydantic_ai_lightspeed.capabilities import QuestionValidity
24+
from pydantic_ai_lightspeed.capabilities.redaction import PiiRedactionCapability
1725
from pydantic_ai_lightspeed.llamastack import (
1826
OgxResponsesModel,
1927
)
28+
from utils.shields import get_shields_for_request
2029

2130
_AGENT_SKILLS_PROVIDER_ID: Final[str] = "agent-skills"
2231
_AGENT_SKILLS_TOOLGROUP_ID: Final[str] = "builtin::agent-skills"
@@ -127,20 +136,51 @@ def get_agent_capability_tools(
127136
return tools
128137

129138

139+
def _shield_capability(shield: ShieldConfiguration) -> AgentCapability[object]:
140+
"""Build the pydantic-ai capability instance for a single configured shield.
141+
142+
Parameters:
143+
shield: A single guardrail shield configuration entry.
144+
145+
Returns:
146+
A ``QuestionValidity`` capability when ``shield.provider_id`` is
147+
``"question_validity"``, or a ``PiiRedactionCapability`` when it is
148+
``"redaction"``.
149+
150+
Raises:
151+
ValueError: If ``shield.config`` doesn't match a known shield config type.
152+
"""
153+
match shield.config:
154+
case QuestionValidityConfig():
155+
return QuestionValidity(config=shield.config)
156+
case RedactionConfig():
157+
return PiiRedactionCapability(config=shield.config)
158+
case _:
159+
raise ValueError(
160+
f"Unsupported shield config type for shield '{shield.name}': "
161+
f"{type(shield.config).__name__}"
162+
)
163+
164+
130165
def _agent_capabilities(
131166
skills: Optional[SkillsConfiguration],
167+
shields: Optional[list[ShieldConfiguration]] = None,
132168
no_tools: bool = False,
133169
) -> Optional[list[AgentCapability[object]]]:
134170
"""Assemble pydantic-ai capabilities for an LCS agent.
135171
136172
Args:
137173
skills: Agent skills configuration from LCS, or None when skills are disabled.
174+
shields: Configured guardrail shields (question validity, redaction), or
175+
None/empty when no shields are enabled.
138176
no_tools: When True, omit capabilities that expose a toolset via ``get_toolset()``.
139177
140178
Returns:
141179
Configured capabilities, or None when no capabilities are enabled.
142180
"""
143181
capabilities: list[AgentCapability[object]] = []
182+
for shield in shields or []:
183+
capabilities.append(_shield_capability(shield))
144184
if skills_capability := _skills_capability(skills):
145185
capabilities.append(skills_capability)
146186
if no_tools:
@@ -158,7 +198,8 @@ def _agent_capabilities(
158198
def build_agent(
159199
client: AsyncOgxClient | AsyncOGXAsLibraryClient,
160200
responses_params: ResponsesApiParams,
161-
skills: Optional[SkillsConfiguration],
201+
config: AppConfig,
202+
shields: Optional[list[str]] = None,
162203
no_tools: bool = False,
163204
) -> Agent[None, str]:
164205
"""Build a Pydantic AI agent that mirrors ``responses_params`` on the Llama Stack backend.
@@ -171,14 +212,20 @@ def build_agent(
171212
Parameters:
172213
client: Initialized Llama Stack client from ``AsyncOgxClientHolder().get_client()``.
173214
responses_params: Parameters produced by ``prepare_responses_params`` for this turn.
174-
skills: Agent skills configuration from LCS, or None when skills are disabled.
215+
config: Application configuration. Agent skills (``config.skills``) and the
216+
configured guardrail shields (``config.shields``) are extracted from it.
217+
shields: Optional list of shield names to run for this turn, matching each
218+
shield's configured ``name``. Mirrors ``QueryRequest.shield_ids``: if
219+
``None``, all shields configured in ``config.shields`` run; an empty
220+
list disables all shields.
175221
no_tools: When True, omit capabilities that expose a toolset via ``get_toolset()``.
176222
177223
Returns:
178224
``Agent`` configured for ``await agent.run(...)`` (or streaming) against the same
179225
stack configuration as ``client.responses.create(**responses_params.model_dump())``.
180226
"""
181-
capabilities = _agent_capabilities(skills, no_tools=no_tools)
227+
shield_configs = get_shields_for_request(config.shields, shields)
228+
capabilities = _agent_capabilities(config.skills, shield_configs, no_tools=no_tools)
182229

183230
model = OgxResponsesModel.from_ogx_client(
184231
responses_params.model, client, responses_params=responses_params

tests/unit/conftest.py

Lines changed: 26 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -3,8 +3,9 @@
33
from __future__ import annotations
44

55
import logging
6-
from collections.abc import Generator
6+
from collections.abc import Callable, Generator
77
from pathlib import Path
8+
from typing import Optional
89

910
import httpx
1011
import pytest
@@ -14,7 +15,7 @@
1415
from configuration import AppConfig
1516
from constants import DEFAULT_LOGGER_NAME
1617
from models.common.responses.responses_api_params import ResponsesApiParams
17-
from models.config import SkillsConfiguration
18+
from models.config import ShieldConfiguration, SkillsConfiguration
1819

1920
type AgentFixtures = Generator[
2021
tuple[
@@ -143,3 +144,26 @@ def mock_skills_configuration_fixture(tmp_path: Path) -> SkillsConfiguration:
143144
encoding="utf-8",
144145
)
145146
return SkillsConfiguration(paths=[skills_root])
147+
148+
149+
@pytest.fixture(name="make_agent_config")
150+
def make_agent_config_fixture(
151+
mocker: MockerFixture,
152+
) -> Callable[..., AppConfig]:
153+
"""Return a factory building a duck-typed AppConfig stand-in for build_agent.
154+
155+
``build_agent`` only reads ``config.skills`` and ``config.shields`` off the
156+
config object it receives, so tests can pass a lightweight mock instead of
157+
a fully-initialized ``AppConfig``.
158+
"""
159+
160+
def _make(
161+
skills: Optional[SkillsConfiguration] = None,
162+
shields: Optional[list[ShieldConfiguration]] = None,
163+
) -> AppConfig:
164+
config = mocker.Mock()
165+
config.skills = skills
166+
config.shields = shields or []
167+
return config
168+
169+
return _make

0 commit comments

Comments
 (0)