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
8 changes: 8 additions & 0 deletions src/configuration.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@
RerankerConfiguration,
RlsapiV1Configuration,
ServiceConfiguration,
ShieldConfiguration,
SkillsConfiguration,
SplunkConfiguration,
UserDataCollection,
Expand Down Expand Up @@ -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.
Expand Down
117 changes: 117 additions & 0 deletions src/models/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
PositiveInt,
PrivateAttr,
SecretStr,
ValidationError,
field_validator,
model_validator,
)
Expand Down Expand Up @@ -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", "<unnamed>")
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."""

Expand Down Expand Up @@ -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:
"""
Expand Down
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/streaming.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down
43 changes: 41 additions & 2 deletions src/utils/pydantic_ai_helpers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down Expand Up @@ -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:
Expand All @@ -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.
Expand All @@ -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
Expand Down
10 changes: 10 additions & 0 deletions tests/unit/models/config/test_dump_configuration.py
Original file line number Diff line number Diff line change
Expand Up @@ -231,6 +231,7 @@ def test_dump_configuration_minimal_cfg(tmp_path: Path) -> None:
},
"saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP,
"skills": None,
"shields": [],
}


Expand Down Expand Up @@ -454,6 +455,7 @@ def test_dump_configuration_valid_values(tmp_path: Path) -> None:
},
"saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP,
"skills": None,
"shields": [],
}


Expand Down Expand Up @@ -828,6 +830,7 @@ def test_dump_configuration_with_quota_limiters(tmp_path: Path) -> None:
},
"saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP,
"skills": None,
"shields": [],
}


Expand Down Expand Up @@ -1086,6 +1089,7 @@ def test_dump_configuration_with_quota_limiters_different_values(
},
"saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP,
"skills": None,
"shields": [],
}


Expand Down Expand Up @@ -1324,6 +1328,7 @@ def test_dump_configuration_byok(tmp_path: Path) -> None:
},
"saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP,
"skills": None,
"shields": [],
}


Expand Down Expand Up @@ -1542,6 +1547,7 @@ def test_dump_configuration_pg_namespace(tmp_path: Path) -> None:
},
"saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP,
"skills": None,
"shields": [],
}


Expand Down Expand Up @@ -1920,6 +1926,7 @@ def test_dump_configuration_allow_degraded_mode(tmp_path: Path) -> None:
},
"saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP,
"skills": None,
"shields": [],
}


Expand Down Expand Up @@ -2144,6 +2151,7 @@ def test_dump_configuration_max_retries_settings(tmp_path: Path) -> None:
},
"saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP,
"skills": None,
"shields": [],
}


Expand Down Expand Up @@ -2368,6 +2376,7 @@ def test_dump_configuration_retry_count_settings(tmp_path: Path) -> None:
},
"saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP,
"skills": None,
"shields": [],
}


Expand Down Expand Up @@ -2599,4 +2608,5 @@ def test_dump_configuration_specific_compaction_values(tmp_path: Path) -> None:
},
"saved_prompts": _DEFAULT_SAVED_PROMPTS_DUMP,
"skills": None,
"shields": [],
}
Loading
Loading