From bd03d6b456f9058dc5bec38613404dedf2891850 Mon Sep 17 00:00:00 2001 From: JR Boos Date: Wed, 22 Jul 2026 14:13:55 -0400 Subject: [PATCH 1/3] feat(config): add shields config --- src/configuration.py | 8 + src/models/config.py | 117 ++++++++++++ .../models/config/test_dump_configuration.py | 10 + .../config/test_shields_configuration.py | 177 ++++++++++++++++++ tests/unit/test_configuration.py | 68 +++++++ 5 files changed, 380 insertions(+) create mode 100644 tests/unit/models/config/test_shields_configuration.py diff --git a/src/configuration.py b/src/configuration.py index ac8f88c7d..87553e668 100644 --- a/src/configuration.py +++ b/src/configuration.py @@ -32,6 +32,7 @@ RerankerConfiguration, RlsapiV1Configuration, ServiceConfiguration, + ShieldConfiguration, SkillsConfiguration, SplunkConfiguration, UserDataCollection, @@ -553,6 +554,13 @@ def skills(self) -> Optional[SkillsConfiguration]: raise LogicError("logic error: configuration is not loaded") return self._configuration.skills + @property + def shields(self) -> list[ShieldConfiguration]: + """Return the list of configured guardrail shields.""" + if self._configuration is None: + raise LogicError("logic error: configuration is not loaded") + return self._configuration.shields + @property def rag_id_mapping(self) -> dict[str, str]: """Return mapping from vector_db_id to rag_id from BYOK and OKP RAG config. diff --git a/src/models/config.py b/src/models/config.py index 488790e10..e38c67b92 100644 --- a/src/models/config.py +++ b/src/models/config.py @@ -22,6 +22,7 @@ PositiveInt, PrivateAttr, SecretStr, + ValidationError, field_validator, model_validator, ) @@ -2590,6 +2591,94 @@ def compiled_patterns(self) -> CompiledPatterns: return list(self._compiled_patterns) +class ShieldConfiguration(ConfigurationBase): + """Configuration for a single named pydantic-ai-lightspeed guardrail shield. + + Each entry configures one instance of an agent guardrail capability + implemented in ``pydantic_ai_lightspeed.capabilities``: question + validity filtering or PII redaction. Multiple shields of the same + ``type`` may be listed (each with a distinct ``name``) to run several + independently-configured instances, e.g. two question validity checks + against different topics, or two separate redaction rule sets. + + These guardrails run inside the pydantic-ai agent and are distinct from + Llama Stack's own safety shields, which are configured in Llama Stack's + ``run.yaml`` and exposed through the ``/v1/shields`` endpoint. Both kinds + of shields can be used together. + + Attributes: + name: Unique, user-facing name identifying this shield instance. + type: Which guardrail capability this shield configures. + config: Type-specific configuration matching ``type``: a + ``QuestionValidityConfig`` when ``type`` is + ``"question_validity"``, or a ``RedactionConfig`` when ``type`` + is ``"redaction"``. + """ + + name: str = Field( + ..., + title="Shield name", + description="Unique, user-facing name identifying this shield instance.", + ) + + type: Literal["question_validity", "redaction"] = Field( + ..., + title="Shield type", + description="Which guardrail capability this shield configures: " + "'question_validity' or 'redaction'.", + ) + + config: QuestionValidityConfig | RedactionConfig = Field( + ..., + title="Shield configuration", + description="Type-specific configuration for this shield. Must match " + "the schema for the selected 'type': QuestionValidityConfig for " + "'question_validity', or RedactionConfig for 'redaction'.", + ) + + @model_validator(mode="before") + @classmethod + def parse_config_by_type(cls, data: Any) -> Any: + """Parse ``config`` using the model that matches ``type``. + + Runs before standard field validation so ``config`` is built with + the concrete model selected by the sibling ``type`` field, rather + than relying on ambiguous union matching. + + Parameters: + data: Raw input data for this model. + + Returns: + The input data, with ``config`` replaced by a parsed instance + of the matching configuration model when possible. + """ + if not isinstance(data, dict): + return data + shield_type = data.get("type") + raw_config = data.get("config") + if not isinstance(raw_config, dict): + return data + + if shield_type == "question_validity": + model_cls: type[QuestionValidityConfig | RedactionConfig] = ( + QuestionValidityConfig + ) + elif shield_type == "redaction": + model_cls = RedactionConfig + else: + return data + + try: + parsed_config = model_cls(**raw_config) + except ValidationError as e: + shield_name = data.get("name", "") + raise ValueError( + f"Invalid config for shield '{shield_name}' of type " + f"'{shield_type}': {e}" + ) from e + return {**data, "config": parsed_config} + + class Configuration(ConfigurationBase): """Global service configuration.""" @@ -2775,6 +2864,34 @@ class Configuration(ConfigurationBase): "maximum prompts per user, display name length, and content length.", ) + shields: list[ShieldConfiguration] = Field( + default_factory=list, + title="Shields configuration", + description="List of pydantic-ai-lightspeed agent guardrail shields " + "(question validity and PII redaction). Each entry has a unique 'name', " + "a 'type' ('question_validity' or 'redaction'), and a type-specific " + "'config'. Distinct from Llama Stack's own safety shields configured " + "in run.yaml.", + ) + + @model_validator(mode="after") + def validate_shield_names_unique(self) -> Self: + """Reject shields lists containing duplicate names. + + Returns: + Self: The model instance after validation. + + Raises: + ValueError: If two or more shields share the same name. + """ + names = [shield.name for shield in self.shields] + duplicates = {name for name in names if names.count(name) > 1} + if duplicates: + raise ValueError( + f"Shield names must be unique, found duplicates: {sorted(duplicates)}" + ) + return self + @model_validator(mode="after") def validate_mcp_auth_headers(self) -> Self: """ diff --git a/tests/unit/models/config/test_dump_configuration.py b/tests/unit/models/config/test_dump_configuration.py index e1cd67816..6360c5599 100644 --- a/tests/unit/models/config/test_dump_configuration.py +++ b/tests/unit/models/config/test_dump_configuration.py @@ -231,6 +231,7 @@ def test_dump_configuration_minimal_cfg(tmp_path: Path) -> None: }, "saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP, "skills": None, + "shields": [], } @@ -454,6 +455,7 @@ def test_dump_configuration_valid_values(tmp_path: Path) -> None: }, "saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP, "skills": None, + "shields": [], } @@ -828,6 +830,7 @@ def test_dump_configuration_with_quota_limiters(tmp_path: Path) -> None: }, "saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP, "skills": None, + "shields": [], } @@ -1086,6 +1089,7 @@ def test_dump_configuration_with_quota_limiters_different_values( }, "saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP, "skills": None, + "shields": [], } @@ -1324,6 +1328,7 @@ def test_dump_configuration_byok(tmp_path: Path) -> None: }, "saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP, "skills": None, + "shields": [], } @@ -1542,6 +1547,7 @@ def test_dump_configuration_pg_namespace(tmp_path: Path) -> None: }, "saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP, "skills": None, + "shields": [], } @@ -1920,6 +1926,7 @@ def test_dump_configuration_allow_degraded_mode(tmp_path: Path) -> None: }, "saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP, "skills": None, + "shields": [], } @@ -2144,6 +2151,7 @@ def test_dump_configuration_max_retries_settings(tmp_path: Path) -> None: }, "saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP, "skills": None, + "shields": [], } @@ -2368,6 +2376,7 @@ def test_dump_configuration_retry_count_settings(tmp_path: Path) -> None: }, "saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP, "skills": None, + "shields": [], } @@ -2599,4 +2608,5 @@ def test_dump_configuration_specific_compaction_values(tmp_path: Path) -> None: }, "saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP, "skills": None, + "shields": [], } diff --git a/tests/unit/models/config/test_shields_configuration.py b/tests/unit/models/config/test_shields_configuration.py new file mode 100644 index 000000000..977dc5bea --- /dev/null +++ b/tests/unit/models/config/test_shields_configuration.py @@ -0,0 +1,177 @@ +"""Unit tests for ShieldConfiguration model and the Configuration.shields list.""" + +# pylint: disable=no-member + +import pytest +from pydantic import ValidationError + +from models.config import ( + CompactionConfiguration, + Configuration, + LlamaStackConfiguration, + QuestionValidityConfig, + RedactionConfig, + ServiceConfiguration, + ShieldConfiguration, + UserDataCollection, +) + + +class TestShieldConfiguration: + """Tests for the ShieldConfiguration model.""" + + def test_question_validity_shield(self) -> None: + """A question_validity shield parses config into QuestionValidityConfig.""" + shield = ShieldConfiguration( + name="topic-guard", + type="question_validity", + config={"model_id": "test-model"}, + ) + assert shield.name == "topic-guard" + assert shield.type == "question_validity" + assert isinstance(shield.config, QuestionValidityConfig) + assert shield.config.model_id == "test-model" + + def test_redaction_shield(self) -> None: + """A redaction shield parses config into RedactionConfig.""" + shield = ShieldConfiguration( + name="pii-guard", + type="redaction", + config={"rules": [{"pattern": r"\d+", "replacement": "[NUM]"}]}, + ) + assert shield.name == "pii-guard" + assert shield.type == "redaction" + assert isinstance(shield.config, RedactionConfig) + assert len(shield.config.compiled_patterns) == 1 + + def test_accepts_already_constructed_config_instance(self) -> None: + """config may be passed as an already-constructed model instance.""" + shield = ShieldConfiguration( + name="topic-guard", + type="question_validity", + config=QuestionValidityConfig(model_id="test-model"), + ) + assert isinstance(shield.config, QuestionValidityConfig) + + def test_rejects_config_mismatched_with_type(self) -> None: + """A redaction type with question_validity-shaped config is rejected.""" + with pytest.raises(ValidationError, match="Invalid config for shield"): + ShieldConfiguration( + name="bad", + type="redaction", + config={"model_id": "oops"}, + ) + + def test_rejects_unknown_type(self) -> None: + """An unrecognized shield type is rejected.""" + with pytest.raises(ValidationError): + ShieldConfiguration( + name="bad", + type="unknown_type", # type: ignore[arg-type] + config={"model_id": "test-model"}, + ) + + def test_rejects_unknown_fields(self) -> None: + """Unknown fields are forbidden on ShieldConfiguration.""" + with pytest.raises(ValidationError, match="Extra inputs are not permitted"): + ShieldConfiguration( + name="topic-guard", + type="question_validity", + config={"model_id": "test-model"}, + unknown_field="value", # type: ignore[call-arg] + ) + + +def _minimal_configuration_kwargs() -> dict: + return { + "name": "test", + "service": ServiceConfiguration(), + "llama_stack": LlamaStackConfiguration( + use_as_library_client=True, + library_client_config_path="tests/configuration/run.yaml", + ), + "user_data_collection": UserDataCollection( + feedback_enabled=False, feedback_storage=None + ), + "compaction": CompactionConfiguration(), + } + + +def test_root_configuration_has_shields_field() -> None: + """The root Configuration declares a shields list field, empty by default.""" + field_info = Configuration.model_fields.get("shields") + assert field_info is not None + + factory = field_info.default_factory + assert factory is not None + assert factory() == [] # type: ignore[call-arg] + + +def test_root_configuration_default_shields_is_empty() -> None: + """Configuration constructed without shields defaults to an empty list.""" + cfg = Configuration(**_minimal_configuration_kwargs()) + assert cfg.shields == [] + + +def test_root_configuration_accepts_multiple_shields_of_same_type() -> None: + """Multiple shields of the same type may be configured with distinct names.""" + cfg = Configuration( + **_minimal_configuration_kwargs(), + shields=[ + ShieldConfiguration( + name="topic-guard-a", + type="question_validity", + config={"model_id": "model-a"}, + ), + ShieldConfiguration( + name="topic-guard-b", + type="question_validity", + config={"model_id": "model-b"}, + ), + ], + ) + assert len(cfg.shields) == 2 + assert cfg.shields[0].name == "topic-guard-a" + assert cfg.shields[1].name == "topic-guard-b" + + +def test_root_configuration_accepts_mixed_shield_types() -> None: + """Shields of different types may be mixed in the same list.""" + cfg = Configuration( + **_minimal_configuration_kwargs(), + shields=[ + ShieldConfiguration( + name="topic-guard", + type="question_validity", + config={"model_id": "test-model"}, + ), + ShieldConfiguration( + name="pii-guard", + type="redaction", + config={"rules": [{"pattern": r"\d+", "replacement": "[NUM]"}]}, + ), + ], + ) + assert len(cfg.shields) == 2 + assert cfg.shields[0].type == "question_validity" + assert cfg.shields[1].type == "redaction" + + +def test_root_configuration_rejects_duplicate_shield_names() -> None: + """Shield names must be unique across the shields list.""" + with pytest.raises(ValidationError, match="Shield names must be unique"): + Configuration( + **_minimal_configuration_kwargs(), + shields=[ + ShieldConfiguration( + name="dup", + type="question_validity", + config={"model_id": "model-a"}, + ), + ShieldConfiguration( + name="dup", + type="redaction", + config={"rules": []}, + ), + ], + ) diff --git a/tests/unit/test_configuration.py b/tests/unit/test_configuration.py index fe1d1fc20..013dfbbe2 100644 --- a/tests/unit/test_configuration.py +++ b/tests/unit/test_configuration.py @@ -130,6 +130,10 @@ def test_default_configuration() -> None: # try to read property _ = cfg.deployment_environment # pylint: disable=pointless-statement + with pytest.raises(LogicError, match="logic error: configuration is not loaded"): + # try to read property + _ = cfg.shields # pylint: disable=pointless-statement + def test_configuration_is_singleton() -> None: """Test that configuration is singleton.""" @@ -242,6 +246,70 @@ def test_init_from_dict() -> None: # check token usage history assert cfg.token_usage_history is None + # check shields - not configured in config_dict, defaults to empty list + assert cfg.shields == [] + + +def test_init_from_dict_with_shields() -> None: + """Test initialization with guardrail shields configuration.""" + config_dict: dict[str, Any] = { + "name": "foo", + "service": { + "host": "localhost", + "port": 8080, + "auth_enabled": False, + "workers": 1, + "color_log": True, + "access_log": True, + }, + "llama_stack": { + "api_key": "xyzzy", + "url": "http://x.y.com:1234", + "use_as_library_client": False, + }, + "user_data_collection": { + "feedback_enabled": False, + }, + "mcp_servers": [], + "customization": None, + "authentication": { + "module": "noop", + }, + "shields": [ + { + "name": "topic-guard-a", + "type": "question_validity", + "config": {"model_id": "test-model"}, + }, + { + "name": "topic-guard-b", + "type": "question_validity", + "config": {"model_id": "test-model-2"}, + }, + { + "name": "pii-guard", + "type": "redaction", + "config": { + "rules": [ + {"pattern": r"\d+", "replacement": "[NUM]"}, + ], + }, + }, + ], + } + cfg = AppConfig() + cfg.init_from_dict(config_dict) + + assert len(cfg.shields) == 3 + assert cfg.shields[0].name == "topic-guard-a" + assert cfg.shields[0].type == "question_validity" + assert cfg.shields[0].config.model_id == "test-model" # type: ignore[union-attr] + assert cfg.shields[1].name == "topic-guard-b" + assert cfg.shields[1].config.model_id == "test-model-2" # type: ignore[union-attr] + assert cfg.shields[2].name == "pii-guard" + assert cfg.shields[2].type == "redaction" + assert len(cfg.shields[2].config.compiled_patterns) == 1 # type: ignore[union-attr] + def test_init_from_dict_with_mcp_servers() -> None: """Test initialization with MCP servers configuration.""" From 8b552c6de26d26e24ab075ee59dd5683eaaca898 Mon Sep 17 00:00:00 2001 From: JR Boos Date: Wed, 22 Jul 2026 14:39:07 -0400 Subject: [PATCH 2/3] feat(shields): wire shields into `build_agent` --- .../question_validity/_capability.py | 4 +- src/utils/pydantic_ai_helpers.py | 43 ++++- tests/unit/utils/test_pydantic_ai.py | 147 +++++++++++++++++- 3 files changed, 189 insertions(+), 5 deletions(-) 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/pydantic_ai_helpers.py b/src/utils/pydantic_ai_helpers.py index b9f410cfb..5211e5546 100644 --- a/src/utils/pydantic_ai_helpers.py +++ b/src/utils/pydantic_ai_helpers.py @@ -12,7 +12,14 @@ from pydantic_ai_skills import SkillsCapability from models.common.responses.responses_api_params import ResponsesApiParams -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, ) @@ -128,20 +135,49 @@ 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. + """ + if isinstance(shield.config, QuestionValidityConfig): + return QuestionValidity(config=shield.config) + if isinstance(shield.config, RedactionConfig): + return PiiRedactionCapability(config=shield.config) + 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: @@ -160,6 +196,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. @@ -173,13 +210,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 9bec4c051..154ae1398 100644 --- a/tests/unit/utils/test_pydantic_ai.py +++ b/tests/unit/utils/test_pydantic_ai.py @@ -3,20 +3,40 @@ # 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, + RedactionConfig, + ShieldConfiguration, + 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 +59,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 = ShieldConfiguration( + 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 = ShieldConfiguration( + 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 +109,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 +120,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 = [ + ShieldConfiguration( + name="topic-guard", + type="question_validity", + config=QuestionValidityConfig(model_id="test-model"), + ), + ShieldConfiguration( + 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 = [ + ShieldConfiguration( + 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 +267,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 = [ + ShieldConfiguration( + name="topic-guard", + type="question_validity", + config=QuestionValidityConfig(model_id="test-model"), + ), + ShieldConfiguration( + 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, From ce0b4e70cf6c162e53e4c71da29fc616f1eaf743 Mon Sep 17 00:00:00 2001 From: JR Boos Date: Wed, 22 Jul 2026 14:52:17 -0400 Subject: [PATCH 3/3] feat(shields): add shields to `/streaming_query` --- src/utils/agents/streaming.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/src/utils/agents/streaming.py b/src/utils/agents/streaming.py index 63c2d0375..c55d2bfbb 100644 --- a/src/utils/agents/streaming.py +++ b/src/utils/agents/streaming.py @@ -121,7 +121,11 @@ async def retrieve_agent_response_generator( ) agent = build_agent( - context.client, responses_params, configuration.skills, no_tools=no_tools + context.client, + responses_params, + configuration.skills, + shields=configuration.shields, + no_tools=no_tools, ) return (