diff --git a/docs/devel_doc/openapi.json b/docs/devel_doc/openapi.json index e903f7e22..680a3bf70 100644 --- a/docs/devel_doc/openapi.json +++ b/docs/devel_doc/openapi.json @@ -13295,11 +13295,15 @@ }, { "$ref": "#/components/schemas/RedactionShieldConfiguration" + }, + { + "$ref": "#/components/schemas/GraniteGuardianShieldConfiguration" } ], "discriminator": { "propertyName": "provider_id", "mapping": { + "granite_guardian": "#/components/schemas/GraniteGuardianShieldConfiguration", "question_validity": "#/components/schemas/QuestionValidityShieldConfiguration", "redaction": "#/components/schemas/RedactionShieldConfiguration" } @@ -14630,6 +14634,105 @@ } ] }, + "GraniteGuardianConfig": { + "properties": { + "url": { + "type": "string", + "title": "Base URL", + "description": "The model_id to use for the guard" + }, + "api_key": { + "anyOf": [ + { + "type": "string", + "format": "password", + "writeOnly": true + }, + { + "type": "null" + } + ], + "title": "Granite Guardian API key", + "description": "API key for the inference" + }, + "max_retries": { + "type": "integer", + "maximum": 5.0, + "minimum": 0.0, + "exclusiveMinimum": 0.0, + "title": "Max retries", + "description": "Maximun number of retires", + "default": 2 + }, + "timeout": { + "type": "integer", + "maximum": 300.0, + "minimum": 5.0, + "exclusiveMinimum": 0.0, + "title": "Timeout", + "description": "Request timeout in seconds", + "default": 30 + }, + "verify_ssl": { + "anyOf": [ + { + "type": "boolean" + }, + { + "type": "string" + } + ], + "title": "Verify SSL", + "description": "SSL certificate verification. Can be:\n - True: Verify using system CA bundle (default, recommended)\n - False: Disable verification (insecure, for dev only)\n - str: Path to custom CA bundle file (for internal PKI)", + "default": true + }, + "risks": { + "items": { + "$ref": "#/components/schemas/RiskDefinition" + }, + "type": "array", + "title": "Defined risks", + "description": "Risks to be considered while applying this guradrail" + } + }, + "additionalProperties": false, + "type": "object", + "required": [ + "url", + "risks" + ], + "title": "GraniteGuardianConfig", + "description": "Configuration for the Granite Guardian moderation guardrail." + }, + "GraniteGuardianShieldConfiguration": { + "properties": { + "name": { + "type": "string", + "title": "Shield name", + "description": "Unique, user-facing name identifying this shield instance." + }, + "provider_id": { + "type": "string", + "const": "granite_guardian", + "title": "Shield provider id", + "description": "Discriminator identifying this as a granite-guardian shield." + }, + "config": { + "$ref": "#/components/schemas/GraniteGuardianConfig", + "title": "Shield configuration", + "description": "Granite-guardian-specific configuration for this shield" + } + }, + "additionalProperties": false, + "type": "object", + "required": [ + "name", + "provider_id", + "config" + ], + "title": "GraniteGuardianShieldConfiguration", + "description": "Configuration for a named Granite Guardian guardrail shield.\n\nAttributes:\n name: Unique, user-facing name identifying this shield instance.\n provider_id: Discriminator identifying this as a granite-guardian shield.\n config: Granite-guardian-specific configuration." + }, "HTTPAuthSecurityScheme": { "properties": { "bearerFormat": { @@ -20718,6 +20821,69 @@ "title": "RetrievalStrategyConfiguration", "description": "Configuration for a single retrieval strategy (inline or tool)." }, + "RiskDefinition": { + "properties": { + "name": { + "type": "string", + "title": "Risk name", + "description": "Unique identifier for this risk (e.g., 'liability', 'competitor_mention')" + }, + "description": { + "type": "string", + "title": "Rist description", + "description": "Risk definition text passed to Granite Guardian as custom_criteria" + }, + "threshold": { + "type": "number", + "maximum": 1.0, + "minimum": 0.0, + "title": "Risk threshold", + "description": "Score threshold for flagging (lower = more sensitive)", + "default": 0.65 + }, + "enabled": { + "type": "boolean", + "title": "Risk enabled", + "description": "Whether to run this check", + "default": true + }, + "enable_thinking": { + "type": "boolean", + "title": "Risk enable thinking", + "description": "Internal field - set via ModerationConfig.thinking_enabled list, not directly. When True, Granite Guardian provides detailed reasoning before scoring.", + "default": false + }, + "points": { + "items": { + "type": "string", + "enum": [ + "input", + "output", + "tool" + ] + }, + "type": "array", + "minItems": 1, + "title": "Guardrail points", + "description": "Where this risk is evaluated: `input` (user message), `output` (model response), or `tool` (tool/MCP content)." + }, + "violation_message": { + "type": "string", + "title": "Violation message", + "description": "Message to be displayed when this risk is violated" + } + }, + "additionalProperties": false, + "type": "object", + "required": [ + "name", + "description", + "points", + "violation_message" + ], + "title": "RiskDefinition", + "description": "Definition for a custom risk category.\n\nCustom risks allow applications to add use-case-specific safety checks\nbeyond the standard harm, jailbreak, leetspeak, amnesia, and\nhistory_politics checks.\nExample:\n liability_risk = RiskDefinition(\n name=\"liability\",\n description=\"Content requesting legal, medical, or financial advice\",\n threshold=0.55,\n points=[\"input\"],\n )\n pii_risk = RiskDefinition(\n name=\"pii_request\",\n description=\"User is asking the AI to reveal personal information\",\n threshold=0.50,\n points=[\"input\", \"tool\"],\n )\nNote:\n To enable think mode (detailed reasoning) for a risk, add the risk name\n to the `thinking_enabled` list in `ModerationConfig`. Do not set\n `enable_thinking` directly - it is managed internally." + }, "RlsapiV1Attachment": { "properties": { "contents": { diff --git a/src/models/config.py b/src/models/config.py index f69b6f0d3..f9e977c80 100644 --- a/src/models/config.py +++ b/src/models/config.py @@ -3142,8 +3142,146 @@ class RedactionShieldConfiguration(ConfigurationBase): ) +class RiskDefinition(ConfigurationBase): + """ + Definition for a custom risk category. + + Custom risks allow applications to add use-case-specific safety checks + beyond the standard harm, jailbreak, leetspeak, amnesia, and + history_politics checks. + Example: + liability_risk = RiskDefinition( + name="liability", + description="Content requesting legal, medical, or financial advice", + threshold=0.55, + points=["input"], + ) + pii_risk = RiskDefinition( + name="pii_request", + description="User is asking the AI to reveal personal information", + threshold=0.50, + points=["input", "tool"], + ) + Note: + To enable think mode (detailed reasoning) for a risk, add the risk name + to the `thinking_enabled` list in `ModerationConfig`. Do not set + `enable_thinking` directly - it is managed internally. + """ + + name: str = Field( + ..., + title="Risk name", + description="Unique identifier for this risk (e.g., 'liability', 'competitor_mention')", + ) + description: str = Field( + ..., + title="Rist description", + description="Risk definition text passed to Granite Guardian as custom_criteria", + ) + threshold: float = Field( + default=0.65, + ge=0.0, + le=1.0, + title="Risk threshold", + description="Score threshold for flagging (lower = more sensitive)", + ) + enabled: bool = Field( + default=True, title="Risk enabled", description="Whether to run this check" + ) + enable_thinking: bool = Field( + default=False, + title="Risk enable thinking", + description=( + "Internal field - set via ModerationConfig.thinking_enabled list, " + "not directly. When True, Granite Guardian provides detailed " + "reasoning before scoring." + ), + ) + points: list[Literal["input", "output", "tool"]] = Field( + ..., + min_length=1, + title="Guardrail points", + description=( + "Where this risk is evaluated: `input` (user message), " + "`output` (model response), or `tool` (tool/MCP content)." + ), + ) + violation_message: str = Field( + ..., + title="Violation message", + description="Message to be displayed when this risk is violated", + ) + + +class GraniteGuardianConfig(ConfigurationBase): + """Configuration for the Granite Guardian moderation guardrail.""" + + url: str = Field( + ..., title="Base URL", description="The model_id to use for the guard" + ) + + api_key: Optional[SecretStr] = Field( + None, title="Granite Guardian API key", description="API key for the inference" + ) + + max_retries: PositiveInt = Field( + 2, ge=0, le=5, title="Max retries", description="Maximun number of retires" + ) + + timeout: PositiveInt = Field( + 30, ge=5, le=300, title="Timeout", description="Request timeout in seconds" + ) + + verify_ssl: bool | str = Field( + True, + title="Verify SSL", + description=( + "SSL certificate verification. Can be:\n" + " - True: Verify using system CA bundle (default, recommended)\n" + " - False: Disable verification (insecure, for dev only)\n" + " - str: Path to custom CA bundle file (for internal PKI)" + ), + ) + + risks: list[RiskDefinition] = Field( + ..., + title="Defined risks", + description="Risks to be considered while applying this guradrail", + ) + + +class GraniteGuardianShieldConfiguration(ConfigurationBase): + """Configuration for a named Granite Guardian guardrail shield. + + Attributes: + name: Unique, user-facing name identifying this shield instance. + provider_id: Discriminator identifying this as a granite-guardian shield. + config: Granite-guardian-specific configuration. + """ + + name: str = Field( + ..., + title="Shield name", + description="Unique, user-facing name identifying this shield instance.", + ) + + provider_id: Literal["granite_guardian"] = Field( + ..., + title="Shield provider id", + description="Discriminator identifying this as a granite-guardian shield.", + ) + + config: GraniteGuardianConfig = Field( + ..., + title="Shield configuration", + description="Granite-guardian-specific configuration for this shield", + ) + + ShieldConfiguration = Annotated[ - QuestionValidityShieldConfiguration | RedactionShieldConfiguration, + QuestionValidityShieldConfiguration + | RedactionShieldConfiguration + | GraniteGuardianShieldConfiguration, Field(discriminator="provider_id"), ] """Configuration for a single named guardrail shield (question validity or redaction). diff --git a/src/utils/pydantic_ai_helpers.py b/src/utils/pydantic_ai_helpers.py index 2c2940737..629334d98 100644 --- a/src/utils/pydantic_ai_helpers.py +++ b/src/utils/pydantic_ai_helpers.py @@ -16,6 +16,7 @@ from models.common.skills import SkillMetadata from models.common.tools import CatalogTool, CatalogToolParameter from models.config import ( + GraniteGuardianConfig, QuestionValidityConfig, RedactionConfig, ShieldConfiguration, @@ -174,6 +175,8 @@ def _shield_capability(shield: ShieldConfiguration) -> AgentCapability[object]: return QuestionValidity(config=shield.config) case RedactionConfig(): return PiiRedactionCapability(config=shield.config) + case GraniteGuardianConfig(): + raise NotImplementedError("Granite Guardian capability not implemented") case _: raise ValueError( f"Unsupported shield config type for shield '{shield.name}': " diff --git a/src/utils/shields.py b/src/utils/shields.py index d45cbcb23..d805fb481 100644 --- a/src/utils/shields.py +++ b/src/utils/shields.py @@ -21,7 +21,12 @@ ShieldModerationPassed, ShieldModerationResult, ) -from models.config import QuestionValidityConfig, RedactionConfig, ShieldConfiguration +from models.config import ( + GraniteGuardianConfig, + QuestionValidityConfig, + RedactionConfig, + ShieldConfiguration, +) from pydantic_ai_lightspeed.capabilities.base import AbstractSafetyCapability from pydantic_ai_lightspeed.capabilities.question_validity._capability import ( QuestionValidity, @@ -158,6 +163,13 @@ def build_shield(shield_config: ShieldConfiguration) -> AbstractSafetyCapability return QuestionValidity(shield_config.config) case RedactionConfig(): return PiiRedactionCapability(shield_config.config) + case GraniteGuardianConfig(): + raise NotImplementedError("Granite Guardian capability not implemented") + case _: + raise ValueError( + f"Unsupported shield config type for shield '{shield_config.name}': " + f"{type(shield_config.config).__name__}" + ) async def run_shield_moderation(