diff --git a/engine/src/agent_control_engine/selectors.py b/engine/src/agent_control_engine/selectors.py index 92ee0e15..1fda2022 100644 --- a/engine/src/agent_control_engine/selectors.py +++ b/engine/src/agent_control_engine/selectors.py @@ -17,6 +17,8 @@ def select_data(step: Step, path: str) -> Any: """ if not path or path == "*": return step.model_dump(mode="json") + if path == "canonical_name": + return step.canonical_name or step.name parts = path.split(".") current: Any = step diff --git a/engine/tests/test_selectors.py b/engine/tests/test_selectors.py index 577922ff..59dedde6 100644 --- a/engine/tests/test_selectors.py +++ b/engine/tests/test_selectors.py @@ -31,6 +31,7 @@ def llm_step_payload() -> Step: "path,expected", [ ("name", "search_database"), + ("canonical_name", "search_database"), ("input.query", "SELECT * FROM users"), ("input.limit", 10), ("input.nested.key", "value"), @@ -85,6 +86,17 @@ def test_select_data_none_handling(): assert result is None +def test_select_data_prefers_explicit_canonical_name() -> None: + payload = Step( + type="tool", + name="writer.web_search", + canonical_name="web_search", + input={}, + ) + + assert select_data(payload, "canonical_name") == "web_search" + + def test_list_selection(): """Test that selecting a path pointing to a list returns the whole list.""" # Given: a payload with a list in the output diff --git a/models/src/agent_control_models/agent.py b/models/src/agent_control_models/agent.py index 6a0eedba..540b050e 100644 --- a/models/src/agent_control_models/agent.py +++ b/models/src/agent_control_models/agent.py @@ -150,6 +150,15 @@ class Step(BaseModel): name: str = Field( ..., min_length=1, description="Step name (tool name or model/chain id)" ) + canonical_name: str | None = Field( + default=None, + min_length=1, + exclude_if=lambda value: value is None, + description=( + "Optional integration-independent identity for a qualified step name " + "(for example, 'web_search' for 'writer.web_search')." + ), + ) input: JSONValue = Field( ..., description="Input content for this step" ) diff --git a/models/src/agent_control_models/controls.py b/models/src/agent_control_models/controls.py index 1e2bb9e9..3a4729d9 100644 --- a/models/src/agent_control_models/controls.py +++ b/models/src/agent_control_models/controls.py @@ -27,7 +27,8 @@ class ControlSelector(BaseModel): default="*", description=( "Path to data using dot notation. " - "Examples: 'input', 'output', 'context.user_id', 'name', 'type', '*'" + "Examples: 'input', 'output', 'context.user_id', 'name', " + "'canonical_name', 'type', '*'" ), ) @@ -43,7 +44,15 @@ def validate_path(cls, v: str | None) -> str: ) # Valid root fields - valid_roots = {"input", "output", "name", "type", "context", "*"} + valid_roots = { + "input", + "output", + "name", + "canonical_name", + "type", + "context", + "*", + } root = v.split(".")[0] if root not in valid_roots: @@ -61,6 +70,7 @@ def validate_path(cls, v: str | None) -> str: {"path": "input"}, {"path": "*"}, {"path": "name"}, + {"path": "canonical_name"}, {"path": "output"}, ] } diff --git a/sdks/python/src/agent_control/evaluation.py b/sdks/python/src/agent_control/evaluation.py index 767a3e02..bfd2348f 100644 --- a/sdks/python/src/agent_control/evaluation.py +++ b/sdks/python/src/agent_control/evaluation.py @@ -517,6 +517,7 @@ def _with_parse_errors(result: EvaluationResult) -> EvaluationResult: async def evaluate_controls( step_name: str, *, + canonical_step_name: str | None = None, input: Any | None = None, output: Any | None = None, context: dict[str, Any] | None = None, @@ -547,6 +548,7 @@ async def evaluate_controls( step_dict: dict[str, Any] = { "type": step_type, "name": step_name, + "canonical_name": canonical_step_name, "input": input if input is not None else default_value, "output": output if output is not None else default_value, } diff --git a/sdks/python/src/agent_control/integrations/_core.py b/sdks/python/src/agent_control/integrations/_core.py index 27693dbf..318ca7fd 100644 --- a/sdks/python/src/agent_control/integrations/_core.py +++ b/sdks/python/src/agent_control/integrations/_core.py @@ -48,6 +48,7 @@ async def _evaluate_and_enforce( agent_name: str, step_name: str, *, + canonical_step_name: str | None = None, input: Any | None = None, output: Any | None = None, context: dict[str, Any] | None = None, @@ -58,6 +59,7 @@ async def _evaluate_and_enforce( result = await agent_control.evaluate_controls( step_name=step_name, + canonical_step_name=canonical_step_name, input=input, output=output, context=context, diff --git a/sdks/python/src/agent_control/integrations/google_adk/plugin.py b/sdks/python/src/agent_control/integrations/google_adk/plugin.py index 28e59698..870a08ce 100644 --- a/sdks/python/src/agent_control/integrations/google_adk/plugin.py +++ b/sdks/python/src/agent_control/integrations/google_adk/plugin.py @@ -268,6 +268,7 @@ async def before_tool_callback( return None step_name = self._resolve_tool_step_name(tool, tool_context=tool_context) + canonical_step_name = resolve_tool_name(tool) self._ensure_step_known(self._build_tool_step_schema(tool, step_name)) context = self._safe_context( step_type="tool", @@ -281,6 +282,7 @@ async def before_tool_callback( await _evaluate_and_enforce( self.agent_name, step_name, + canonical_step_name=canonical_step_name, input=tool_args, context=context, step_type="tool", @@ -311,6 +313,7 @@ async def after_tool_callback( return None step_name = self._resolve_tool_step_name(tool, tool_context=tool_context) + canonical_step_name = resolve_tool_name(tool) self._ensure_step_known(self._build_tool_step_schema(tool, step_name)) context = self._safe_context( step_type="tool", @@ -325,6 +328,7 @@ async def after_tool_callback( await _evaluate_and_enforce( self.agent_name, step_name, + canonical_step_name=canonical_step_name, input=tool_args, output=result, context=context, diff --git a/sdks/python/src/agent_control/integrations/strands/plugin.py b/sdks/python/src/agent_control/integrations/strands/plugin.py index 1aa503cd..9e9c51eb 100644 --- a/sdks/python/src/agent_control/integrations/strands/plugin.py +++ b/sdks/python/src/agent_control/integrations/strands/plugin.py @@ -115,6 +115,7 @@ async def _evaluate_and_enforce( ) -> None: result = await agent_control.evaluate_controls( step_name=step_name, + canonical_step_name=step_name if step_type == "tool" else None, input=input, output=output, context=context, diff --git a/sdks/python/tests/test_evaluation.py b/sdks/python/tests/test_evaluation.py index 2fb92555..cb3d24f1 100644 --- a/sdks/python/tests/test_evaluation.py +++ b/sdks/python/tests/test_evaluation.py @@ -57,9 +57,9 @@ def json(self) -> dict[str, object]: json={ "agent_name": "agent-example_01", "step": { - "type": "llm", - "name": "chat", - "input": "hello", + "type": "llm", + "name": "chat", + "input": "hello", "output": None, "context": None, }, diff --git a/sdks/python/tests/test_google_adk_plugin.py b/sdks/python/tests/test_google_adk_plugin.py index f68bd341..1a056f57 100644 --- a/sdks/python/tests/test_google_adk_plugin.py +++ b/sdks/python/tests/test_google_adk_plugin.py @@ -10,7 +10,6 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest - from agent_control import ControlSteerError, ControlViolationError from agent_control._state import state @@ -369,6 +368,7 @@ async def test_tool_callbacks_scope_step_name_by_agent(plugin_module): ) assert mock_eval.await_args.args[1] == "writer.get_weather" + assert mock_eval.await_args.kwargs["canonical_step_name"] == "get_weather" @pytest.mark.asyncio diff --git a/sdks/typescript/src/generated/models/control-selector.ts b/sdks/typescript/src/generated/models/control-selector.ts index 8144bb20..1e8283c9 100644 --- a/sdks/typescript/src/generated/models/control-selector.ts +++ b/sdks/typescript/src/generated/models/control-selector.ts @@ -18,7 +18,7 @@ import { SDKValidationError } from "./errors/sdk-validation-error.js"; */ export type ControlSelector = { /** - * Path to data using dot notation. Examples: 'input', 'output', 'context.user_id', 'name', 'type', '*' + * Path to data using dot notation. Examples: 'input', 'output', 'context.user_id', 'name', 'canonical_name', 'type', '*' */ path?: string | null | undefined; }; diff --git a/sdks/typescript/src/generated/models/step.ts b/sdks/typescript/src/generated/models/step.ts index 132cf9c9..db7d5746 100644 --- a/sdks/typescript/src/generated/models/step.ts +++ b/sdks/typescript/src/generated/models/step.ts @@ -3,11 +3,16 @@ */ import * as z from "zod/v4-mini"; +import { remap as remap$ } from "../lib/primitives.js"; /** * Runtime payload for an agent step invocation. */ export type Step = { + /** + * Optional integration-independent identity for a qualified step name (for example, 'web_search' for 'writer.web_search'). + */ + canonicalName?: string | null | undefined; /** * Optional context (conversation history, metadata, etc.) */ @@ -32,6 +37,7 @@ export type Step = { /** @internal */ export type Step$Outbound = { + canonical_name?: string | null | undefined; context?: { [k: string]: any } | null | undefined; input: any; name: string; @@ -40,14 +46,20 @@ export type Step$Outbound = { }; /** @internal */ -export const Step$outboundSchema: z.ZodMiniType = z.object( - { +export const Step$outboundSchema: z.ZodMiniType = z.pipe( + z.object({ + canonicalName: z.optional(z.nullable(z.string())), context: z.optional(z.nullable(z.record(z.string(), z.any()))), input: z.any(), name: z.string(), output: z.optional(z.nullable(z.any())), type: z.string(), - }, + }), + z.transform((v) => { + return remap$(v, { + canonicalName: "canonical_name", + }); + }), ); export function stepToJSON(step: Step): string { diff --git a/server/alembic/versions/f3a1c8d7e2b4_out_of_box_control_seed_identity.py b/server/alembic/versions/f3a1c8d7e2b4_out_of_box_control_seed_identity.py index f976a2d9..d9fb6d31 100644 --- a/server/alembic/versions/f3a1c8d7e2b4_out_of_box_control_seed_identity.py +++ b/server/alembic/versions/f3a1c8d7e2b4_out_of_box_control_seed_identity.py @@ -17,6 +17,8 @@ branch_labels = None depends_on = None +_CANONICAL_NAME_SEED_SOURCE_ID = "oob-only-approved-tools-may-run" + def upgrade() -> None: op.add_column("controls", sa.Column("seed_source_id", sa.String(length=255), nullable=True)) @@ -35,6 +37,36 @@ def upgrade() -> None: def downgrade() -> None: + # Older servers reject ``canonical_name`` selectors. Retire only the seeded + # control that still uses that selector, and make both its current payload + # and historical snapshots parseable before rolling the application back. + op.execute( + f""" + UPDATE control_versions AS version + SET snapshot = jsonb_set( + version.snapshot, + '{{data,condition,selector,path}}', + '"name"'::jsonb + ) + FROM controls AS control + WHERE version.control_id = control.id + AND control.seed_source_id = '{_CANONICAL_NAME_SEED_SOURCE_ID}' + AND version.snapshot #>> '{{data,condition,selector,path}}' = 'canonical_name' + """ + ) + op.execute( + f""" + UPDATE controls + SET data = jsonb_set( + data, + '{{condition,selector,path}}', + '"name"'::jsonb + ), + deleted_at = COALESCE(deleted_at, CURRENT_TIMESTAMP) + WHERE seed_source_id = '{_CANONICAL_NAME_SEED_SOURCE_ID}' + AND data #>> '{{condition,selector,path}}' = 'canonical_name' + """ + ) with op.get_context().autocommit_block(): op.execute("DROP INDEX CONCURRENTLY IF EXISTS idx_controls_namespace_seed_source") op.drop_column("controls", "seed_opted_out_at") diff --git a/server/src/agent_control_server/bootstrap/out_of_box_controls.py b/server/src/agent_control_server/bootstrap/out_of_box_controls.py index 2368565d..d57ba3df 100644 --- a/server/src/agent_control_server/bootstrap/out_of_box_controls.py +++ b/server/src/agent_control_server/bootstrap/out_of_box_controls.py @@ -1,9 +1,5 @@ """Startup bootstrap for out-of-box controls. -Phase 1 provides the tooling needed to seed controls safely, but does not -register the static out-of-box control catalog yet. Phase 2 should add those -definitions to ``OUT_OF_BOX_CONTROL_TEMPLATES``. - Namespace rule: - Standalone Agent Control seeds into ``DEFAULT_NAMESPACE_KEY``. - Galileo-integrated Agent Control should call the same helper with @@ -36,6 +32,11 @@ _CONTROL_SEED_UNIQUE_CONSTRAINT = "idx_controls_namespace_seed_source" _INITIAL_VERSION_NOTE = "Out-of-box control seed" _SLUG_NAME_ADAPTER = TypeAdapter(SlugName) +_OUT_OF_BOX_TAGS = ["out-of-box"] +_SQL_TOOL_NAME_PATTERN = ( + r"(?i)(?:^|[._-])(?:sql|execute[_-]?sql|run[_-]?sql|sql[_-]?query|" + r"query[_-]?database|execute[_-]?query)(?:$|[._-])" +) @dataclass(frozen=True, slots=True) @@ -109,7 +110,300 @@ def skipped_count(self) -> int: ) -OUT_OF_BOX_CONTROL_TEMPLATES: tuple[OutOfBoxControlTemplate, ...] = () +def _leaf_control_payload( + *, + description: str, + selector_path: str, + evaluator_name: str, + evaluator_config: Mapping[str, object], + step_types: list[str], + stages: list[str], + decision: str, + tags: list[str], + steering_message: str | None = None, + step_name_regex: str | None = None, +) -> dict[str, object]: + action: dict[str, object] = {"decision": decision} + if steering_message is not None: + action["steering_context"] = {"message": steering_message} + + scope: dict[str, object] = {"step_types": step_types, "stages": stages} + if step_name_regex is not None: + scope["step_name_regex"] = step_name_regex + + return { + "description": description, + "enabled": True, + "execution": "server", + "scope": scope, + "condition": { + "selector": {"path": selector_path}, + "evaluator": { + "name": evaluator_name, + "config": dict(evaluator_config), + }, + }, + "action": action, + "tags": [*_OUT_OF_BOX_TAGS, *tags], + } + + +OUT_OF_BOX_CONTROL_TEMPLATES: tuple[OutOfBoxControlTemplate, ...] = ( + OutOfBoxControlTemplate.from_payload( + source_id="oob-ssn-match", + name="oob-ssn-match", + data=_leaf_control_payload( + description="Block LLM output containing US Social Security Numbers.", + selector_path="output", + evaluator_name="regex", + evaluator_config={"pattern": r"\b\d{3}-\d{2}-\d{4}\b"}, + step_types=["llm"], + stages=["post"], + decision="deny", + tags=["pii", "regex"], + ), + ), + OutOfBoxControlTemplate.from_payload( + source_id="oob-credit-card-number-match", + name="oob-credit-card-number-match", + data=_leaf_control_payload( + description="Block LLM output containing common credit-card-like numbers.", + selector_path="output", + evaluator_name="regex", + evaluator_config={"pattern": r"\b(?:\d[ -]?){13,19}\b"}, + step_types=["llm"], + stages=["post"], + decision="deny", + tags=["pii", "payment", "regex"], + ), + ), + OutOfBoxControlTemplate.from_payload( + source_id="oob-phone-number-match", + name="oob-phone-number-match", + data=_leaf_control_payload( + description="Block LLM output containing common US phone number formats.", + selector_path="output", + evaluator_name="regex", + evaluator_config={ + "pattern": ( + r"\b(?:\+?1[-.\s]?)?(?:\(?[2-9]\d{2}\)?[-.\s]?)?" + r"[2-9]\d{2}[-.\s]?\d{4}\b" + ) + }, + step_types=["llm"], + stages=["post"], + decision="deny", + tags=["pii", "regex"], + ), + ), + OutOfBoxControlTemplate.from_payload( + source_id="oob-dangerous-shell-command-match", + name="oob-dangerous-shell-command-match", + data=_leaf_control_payload( + description="Block tool commands matching common destructive shell operations.", + selector_path="input.command", + evaluator_name="regex", + evaluator_config={ + "pattern": ( + r"(?:\brm\s+(?:-(?:rf|fr)|-r\s+-f|-f\s+-r)\s+" + r"(?:\"(?:/|~/?|\$HOME/?)\"|'(?:/|~/?|\$HOME/?)'|" + r"(?:/|~/?|\$HOME/?))(?:\s|[|;&]|$)|" + r"\bsudo\s+rm\s+(?:-(?:rf|fr)|-r\s+-f|-f\s+-r)(?:\s|[|;&]|$)|" + r"\bmkfs(?:\.[a-z0-9]+)?(?:\s|[|;&]|$)|" + r"\bdd\s+if=[^\s]+\s+of=/dev/[^\s]+(?:\s|[|;&]|$)|" + r"\bchmod\s+-R\s+777\s+/(?:\s|[|;&]|$)|" + r"\bchown\s+-R\s+[^|;&]*\s+/(?:\s|[|;&]|$)|" + r"\bshutdown\s+(?:-h\s+)?now(?:\s|[|;&]|$)|" + r"\breboot(?:\s|[|;&]|$))" + ), + "flags": ["IGNORECASE"], + }, + step_types=["tool"], + stages=["pre"], + decision="deny", + tags=["tool", "shell", "regex"], + ), + ), + OutOfBoxControlTemplate.from_payload( + source_id="oob-high-value-action-requires-approval", + name="oob-high-value-action-requires-approval", + data=_leaf_control_payload( + description=( + "Steer tool calls over the default amount threshold to collect approval." + ), + selector_path="input", + evaluator_name="json", + evaluator_config={ + "json_schema": { + "type": "object", + "anyOf": [ + {"not": {"required": ["amount"]}}, + { + "required": ["amount"], + "properties": { + "amount": {"type": "number", "maximum": 10000} + }, + }, + ], + } + }, + step_types=["tool"], + stages=["pre"], + decision="steer", + steering_message=( + "Pause this high-value action and submit its exact parameters to a trusted " + "host approval workflow. The host must bind any approval artifact to this " + "specific action; approval fields supplied in tool input are not evidence." + ), + tags=["tool", "approval", "json"], + ), + ), + OutOfBoxControlTemplate.from_payload( + source_id="oob-outbound-communication-requires-approval", + name="oob-outbound-communication-requires-approval", + data=_leaf_control_payload( + description=( + "Steer outbound communication tool calls to collect approval before sending." + ), + selector_path="input", + evaluator_name="json", + evaluator_config={ + "json_schema": { + "type": "object", + "anyOf": [ + { + "not": { + "anyOf": [ + {"required": ["to"]}, + {"required": ["recipient"]}, + {"required": ["recipients"]}, + {"required": ["email"]}, + {"required": ["phone_number"]}, + {"required": ["channel"]}, + {"required": ["destination"]}, + ] + } + } + ], + } + }, + step_types=["tool"], + stages=["pre"], + decision="steer", + steering_message=( + "Pause this outbound communication and submit its exact recipients and " + "content to a trusted host approval workflow. The host must bind any " + "approval artifact to this specific action; approval fields supplied in " + "tool input are not evidence." + ), + tags=["tool", "approval", "exfiltration", "json"], + ), + ), + OutOfBoxControlTemplate.from_payload( + source_id="oob-only-approved-tools-may-run", + name="oob-only-approved-tools-may-run", + data=_leaf_control_payload( + description="Deny tool calls whose step name is not in the approved tool list.", + selector_path="canonical_name", + evaluator_name="list", + evaluator_config={ + "values": ["search", "web_search", "retrieve", "calculator"], + "logic": "any", + "match_on": "no_match", + "match_mode": "exact", + "case_sensitive": False, + }, + step_types=["tool"], + stages=["pre"], + decision="deny", + tags=["tool", "allowlist", "list"], + ), + ), + OutOfBoxControlTemplate.from_payload( + source_id="oob-owasp-llm05-read-only-sql", + name="oob-owasp-llm05-read-only-sql", + data=_leaf_control_payload( + description=("Block SQL tool calls that are not a single read-only SELECT statement."), + selector_path="input.query", + evaluator_name="sql", + evaluator_config={ + "allowed_operations": ["SELECT"], + "allow_multi_statements": False, + "block_ddl": True, + "block_dcl": True, + }, + step_types=["tool"], + stages=["pre"], + decision="deny", + tags=["owasp", "owasp-llm05", "owasp-asi02", "tool", "sql"], + step_name_regex=_SQL_TOOL_NAME_PATTERN, + ), + ), + OutOfBoxControlTemplate.from_payload( + source_id="oob-owasp-llm10-bounded-sql-query", + name="oob-owasp-llm10-bounded-sql-query", + data=_leaf_control_payload( + description=("Block SQL queries without bounded results or with excessive complexity."), + selector_path="input.query", + evaluator_name="sql", + evaluator_config={ + "require_limit": True, + "max_limit": 1000, + "max_result_window": 1000, + "max_subquery_depth": 3, + "max_joins": 5, + "max_union_count": 2, + }, + step_types=["tool"], + stages=["pre"], + decision="deny", + tags=["owasp", "owasp-llm10", "tool", "sql", "resource-limit"], + step_name_regex=_SQL_TOOL_NAME_PATTERN, + ), + ), + OutOfBoxControlTemplate.from_payload( + source_id="oob-owasp-llm02-common-credential-output-match", + name="oob-owasp-llm02-common-credential-output-match", + data=_leaf_control_payload( + description=("Block LLM output containing common private-key or API-token formats."), + selector_path="output", + evaluator_name="regex", + evaluator_config={ + "pattern": ( + r"(?:-----BEGIN (?:RSA |EC |DSA |OPENSSH )?PRIVATE KEY-----|" + r"\b(?:AKIA|ASIA)[A-Z0-9]{16}\b|" + r"\bgh[pousr]_[A-Za-z0-9]{36,255}\b|" + r"\bAIza[0-9A-Za-z_-]{35}\b|" + r"\bxox[baprs]-[A-Za-z0-9-]{10,}\b)" + ) + }, + step_types=["llm"], + stages=["post"], + decision="deny", + tags=["owasp", "owasp-llm02", "credential", "secret", "regex"], + ), + ), + OutOfBoxControlTemplate.from_payload( + source_id="oob-owasp-llm05-dangerous-uri-output-match", + name="oob-owasp-llm05-dangerous-uri-output-match", + data=_leaf_control_payload( + description=("Block LLM output containing executable or active-content URI schemes."), + selector_path="output", + evaluator_name="regex", + evaluator_config={ + "pattern": ( + r"(?:\b(?:javascript|vbscript)\s*:|" + r"\bdata\s*:\s*(?:text/html|application/xhtml\+xml|image/svg\+xml))" + ), + "flags": ["IGNORECASE"], + }, + step_types=["llm"], + stages=["post"], + decision="deny", + tags=["owasp", "owasp-llm05", "output-handling", "uri", "regex"], + ), + ), +) def default_out_of_box_namespace_key() -> str: @@ -150,6 +444,7 @@ async def seed_out_of_box_controls( available_evaluator_names = set(available_evaluators) async with session_factory() as session: + eligible_templates: list[OutOfBoxControlTemplate] = [] for template in templates: missing = missing_required_evaluators( template.required_evaluators, @@ -164,6 +459,19 @@ async def seed_out_of_box_controls( ) continue + eligible_templates.append(template) + + control_service = ControlService(session) + existing_source_ids, active_names = await control_service.find_existing_seed_controls( + namespace_key=namespace_key, + source_ids={template.source_id for template in eligible_templates}, + names={template.name for template in eligible_templates}, + ) + for template in eligible_templates: + if template.source_id in existing_source_ids or template.name in active_names: + skipped_existing.append(template.name) + continue + outcome = await _seed_one_control( session, namespace_key=namespace_key, @@ -191,14 +499,6 @@ async def _seed_one_control( template: OutOfBoxControlTemplate, ) -> str: control_service = ControlService(session) - if await control_service.seed_source_exists( - template.source_id, - namespace_key=namespace_key, - ): - return "existing" - if await control_service.active_control_name_exists(template.name, namespace_key=namespace_key): - return "existing" - control = control_service.create_control( namespace_key=namespace_key, name=template.name, diff --git a/server/src/agent_control_server/config.py b/server/src/agent_control_server/config.py index 00335611..5e342222 100644 --- a/server/src/agent_control_server/config.py +++ b/server/src/agent_control_server/config.py @@ -214,7 +214,6 @@ class Settings(BaseSettings): "AGENT_CONTROL_ALLOW_HEADERS", "ALLOW_HEADERS", ) - def get_cors_origins(self) -> list[str]: """Parse CORS origins from string or list.""" return self._parse_list_setting(self.cors_origins) diff --git a/server/src/agent_control_server/endpoints/controls.py b/server/src/agent_control_server/endpoints/controls.py index d328c7f9..dce876ad 100644 --- a/server/src/agent_control_server/endpoints/controls.py +++ b/server/src/agent_control_server/endpoints/controls.py @@ -1,6 +1,8 @@ +import asyncio import datetime as dt import uuid from copy import deepcopy +from functools import partial from typing import Any from agent_control_engine import list_evaluators @@ -43,7 +45,8 @@ from sqlalchemy.ext.asyncio import AsyncSession from ..auth_framework import Operation, Principal, get_authorizer, require_operation -from ..db import get_async_db +from ..bootstrap.out_of_box_controls import seed_out_of_box_controls +from ..db import AsyncSessionLocal, get_async_db from ..errors import ( APIError, APIValidationError, @@ -94,6 +97,9 @@ _GENERATED_CLONE_NAME_ATTEMPTS = 5 _TRUE_QUERY_VALUES = {"1", "true", "t", "yes", "y", "on"} _SLUG_NAME_ADAPTER = TypeAdapter(SlugName) +_OUT_OF_BOX_RECONCILIATION_TIMEOUT_SECONDS = 3.0 +_MAX_PENDING_OUT_OF_BOX_RECONCILIATIONS = 128 +_out_of_box_reconciliation_tasks: dict[str, asyncio.Task[None]] = {} def _is_target_context_value(value: object) -> bool: @@ -257,6 +263,95 @@ def _validate_attachment_filters( ) +async def _run_out_of_box_controls_reconciliation( + *, + namespace_key: str, +) -> None: + """Run one bounded, best-effort namespace reconciliation.""" + try: + async with asyncio.timeout(_OUT_OF_BOX_RECONCILIATION_TIMEOUT_SECONDS): + await seed_out_of_box_controls( + session_factory=AsyncSessionLocal, + namespace_key=namespace_key, + available_evaluators=set(list_evaluators().keys()), + ) + except TimeoutError: + _logger.warning( + "Out-of-box control reconciliation timed out for namespace '%s'; " + "continuing request", + namespace_key, + ) + except Exception: + _logger.warning( + "Out-of-box control seed failed for namespace '%s'; continuing request", + namespace_key, + exc_info=True, + ) + + +def _remove_out_of_box_reconciliation_task( + namespace_key: str, + task: asyncio.Future[None], +) -> None: + if _out_of_box_reconciliation_tasks.get(namespace_key) is task: + _out_of_box_reconciliation_tasks.pop(namespace_key, None) + + +async def _seed_out_of_box_controls_for_namespace( + *, + namespace_key: str, +) -> None: + """Join or start the bounded reconciliation for a namespace.""" + task = _out_of_box_reconciliation_tasks.get(namespace_key) + if task is None: + if len(_out_of_box_reconciliation_tasks) >= _MAX_PENDING_OUT_OF_BOX_RECONCILIATIONS: + _logger.warning( + "Out-of-box control reconciliation queue is full; skipping namespace '%s'", + namespace_key, + ) + return + task = asyncio.create_task( + _run_out_of_box_controls_reconciliation(namespace_key=namespace_key) + ) + _out_of_box_reconciliation_tasks[namespace_key] = task + task.add_done_callback( + partial(_remove_out_of_box_reconciliation_task, namespace_key) + ) + + await asyncio.shield(task) + + +def _should_seed_out_of_box_controls_on_list( + *, + cursor: int | None, + name: str | None, + enabled: bool | None, + template_backed: bool | None, + cloned: bool | None, + step_type: str | None, + stage: str | None, + execution: str | None, + tag: str | None, + include_attachments: bool, + attachment_target_type: str | None, + attachment_target_id: str | None, +) -> bool: + return ( + cursor is None + and name is None + and enabled is None + and template_backed is None + and cloned is not True + and step_type is None + and stage is None + and execution is None + and tag is None + and not include_attachments + and attachment_target_type is None + and attachment_target_id is None + ) + + def _serialize_control_data( control_data: ControlDefinition | UnrenderedTemplateControl, ) -> dict[str, object]: @@ -1234,6 +1329,21 @@ async def list_controls( control_service = ControlService(db) namespace_key = principal.namespace_key + if _should_seed_out_of_box_controls_on_list( + cursor=cursor, + name=name, + enabled=enabled, + template_backed=template_backed, + cloned=cloned, + step_type=step_type, + stage=stage, + execution=execution, + tag=tag, + include_attachments=include_attachments, + attachment_target_type=attachment_target_type, + attachment_target_id=attachment_target_id, + ): + await _seed_out_of_box_controls_for_namespace(namespace_key=namespace_key) filter_by_attachment = target_principal is not None and ( attachment_target_type is not None or attachment_target_id is not None ) diff --git a/server/src/agent_control_server/endpoints/evaluation.py b/server/src/agent_control_server/endpoints/evaluation.py index a31d757d..adc465a9 100644 --- a/server/src/agent_control_server/endpoints/evaluation.py +++ b/server/src/agent_control_server/endpoints/evaluation.py @@ -117,6 +117,21 @@ def _sanitize_evaluation_response(response: EvaluationResponse) -> EvaluationRes ) +def _normalize_legacy_qualified_step_name(request: EvaluationRequest) -> EvaluationRequest: + """Derive a canonical tool name for clients that predate ``canonical_name``.""" + step = request.step + if step.type != "tool" or step.canonical_name is not None: + return request + + _, separator, canonical_name = step.name.rpartition(".") + if not separator or not canonical_name: + return request + + return request.model_copy( + update={"step": step.model_copy(update={"canonical_name": canonical_name})} + ) + + async def _evaluation_context(request: Request) -> dict[str, object]: """Surface target identifiers to the runtime authorizer.""" try: @@ -196,6 +211,7 @@ async def evaluate( on the server; SDKs reconstruct and emit those events separately through the observability ingestion endpoint. """ + request = _normalize_legacy_qualified_step_name(request) engine_controls = await _load_engine_controls(request, principal) engine = ControlEngine(engine_controls) try: diff --git a/server/src/agent_control_server/services/controls.py b/server/src/agent_control_server/services/controls.py index 619cbed1..6af39721 100644 --- a/server/src/agent_control_server/services/controls.py +++ b/server/src/agent_control_server/services/controls.py @@ -1,7 +1,7 @@ from __future__ import annotations import datetime as dt -from collections.abc import Sequence +from collections.abc import Collection, Sequence from dataclasses import dataclass from typing import Any, Literal, cast @@ -163,19 +163,32 @@ def mark_control_deleted(control: Control, *, deleted_at: dt.datetime) -> None: if control.seed_source_id is not None: control.seed_opted_out_at = deleted_at - async def seed_source_exists( + async def find_existing_seed_controls( self, - seed_source_id: str, *, namespace_key: str, - ) -> bool: - """Return whether a control has claimed an immutable seed identity.""" - stmt = select(Control.id).where( + source_ids: Collection[str], + names: Collection[str], + ) -> tuple[frozenset[str], frozenset[str]]: + """Bulk-load claimed seed identities and conflicting active names.""" + if not source_ids and not names: + return frozenset(), frozenset() + + stmt = select(Control.seed_source_id, Control.name, Control.deleted_at).where( Control.namespace_key == namespace_key, - Control.seed_source_id == seed_source_id, + or_( + Control.seed_source_id.in_(source_ids), + Control.name.in_(names), + ), ) - result = await self._db.execute(stmt) - return result.first() is not None + rows = (await self._db.execute(stmt)).all() + existing_source_ids = frozenset( + source_id for source_id, _, _ in rows if source_id is not None + ) + active_names = frozenset( + name for _, name, deleted_at in rows if deleted_at is None + ) + return existing_source_ids, active_names async def get_control_or_404( self, diff --git a/server/tests/test_controls_additional.py b/server/tests/test_controls_additional.py index cf7aa4b0..8cefc8f7 100644 --- a/server/tests/test_controls_additional.py +++ b/server/tests/test_controls_additional.py @@ -1,5 +1,6 @@ from __future__ import annotations +import asyncio import json import uuid from collections.abc import AsyncGenerator @@ -12,6 +13,13 @@ from agent_control_evaluators import RegexEvaluatorConfig from agent_control_models import ConditionNode from agent_control_models.errors import ErrorCode, ErrorReason +from fastapi.testclient import TestClient +from sqlalchemy import select, text +from sqlalchemy.exc import IntegrityError +from sqlalchemy.ext.asyncio import AsyncSession +from sqlalchemy.orm import Session +from starlette.requests import Request + from agent_control_server.auth_framework import Operation, Principal, set_authorizer from agent_control_server.db import get_async_db from agent_control_server.endpoints import controls as controls_module @@ -23,11 +31,6 @@ ControlBinding, ControlVersion, ) -from fastapi.testclient import TestClient -from sqlalchemy import select, text -from sqlalchemy.exc import IntegrityError -from sqlalchemy.ext.asyncio import AsyncSession -from sqlalchemy.orm import Session from .conftest import engine from .utils import VALID_CONTROL_PAYLOAD @@ -40,6 +43,22 @@ def _make_integrity_error(constraint_name: str) -> IntegrityError: return IntegrityError("statement", {}, orig) +def _request(*, query: str = "", body: bytes = b"") -> Request: + async def receive() -> dict[str, object]: + return {"type": "http.request", "body": body, "more_body": False} + + return Request( + { + "type": "http", + "method": "GET", + "path": "/", + "headers": [], + "query_string": query.encode(), + }, + receive, + ) + + def _create_control( client: TestClient, name: str | None = None, @@ -473,6 +492,138 @@ def test_clone_and_bind_context_tolerates_invalid_body_shapes( assert bad_target_resp.status_code == 422 +@pytest.mark.asyncio +async def test_clone_and_bind_context_returns_empty_for_malformed_json() -> None: + malformed_request = _request(body=b"{") + invalid_target_request = _request( + body=json.dumps( + { + "target_binding": { + "target_type": "log_stream", + "target_id": "", + } + } + ).encode() + ) + + assert await controls_module._clone_and_bind_context(malformed_request) == {} + assert await controls_module._clone_and_bind_context(invalid_target_request) == {} + + +def test_attachment_target_context_rejects_invalid_values() -> None: + invalid_type = _request(query="attachment_target_type=&attachment_target_id=target") + invalid_id = _request(query="attachment_target_type=log_stream&attachment_target_id=") + + assert controls_module._attachment_target_context(invalid_type) == {} + assert controls_module._attachment_target_context(invalid_id) == {} + + +@pytest.mark.asyncio +async def test_optional_attachment_authorization_skips_false_flag() -> None: + request = _request(query="include_attachments=false") + + assert await controls_module._optional_attachment_target_principal(request) is None + + +def test_enabled_from_stored_payload_defaults_for_non_mapping() -> None: + assert controls_module._enabled_from_stored_payload("invalid") is True + + +@pytest.mark.asyncio +async def test_seed_out_of_box_controls_failure_is_best_effort( + monkeypatch: pytest.MonkeyPatch, +) -> None: + seed = AsyncMock(side_effect=RuntimeError("database unavailable")) + monkeypatch.setattr(controls_module, "seed_out_of_box_controls", seed) + + await controls_module._seed_out_of_box_controls_for_namespace( + namespace_key="seed-failure-namespace" + ) + + seed.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_seed_out_of_box_controls_deduplicates_concurrent_namespace_requests( + monkeypatch: pytest.MonkeyPatch, +) -> None: + # Given: one namespace reconciliation that remains in flight + started = asyncio.Event() + release = asyncio.Event() + calls = 0 + + async def seed(**_kwargs: object) -> None: + nonlocal calls + calls += 1 + started.set() + await release.wait() + + monkeypatch.setattr(controls_module, "seed_out_of_box_controls", seed) + + # When: two list requests reconcile the same namespace concurrently + first = asyncio.create_task( + controls_module._seed_out_of_box_controls_for_namespace( + namespace_key="shared-reconciliation-namespace" + ) + ) + await started.wait() + second = asyncio.create_task( + controls_module._seed_out_of_box_controls_for_namespace( + namespace_key="shared-reconciliation-namespace" + ) + ) + await asyncio.sleep(0) + release.set() + await asyncio.gather(first, second) + + # Then: both requests joined one reconciliation + assert calls == 1 + + +@pytest.mark.asyncio +async def test_seed_out_of_box_controls_bounds_reconciliation_wait( + monkeypatch: pytest.MonkeyPatch, +) -> None: + # Given: a reconciliation that cannot complete within the read-path deadline + cancelled = asyncio.Event() + + async def seed(**_kwargs: object) -> None: + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + cancelled.set() + raise + + monkeypatch.setattr(controls_module, "seed_out_of_box_controls", seed) + monkeypatch.setattr(controls_module, "_OUT_OF_BOX_RECONCILIATION_TIMEOUT_SECONDS", 0.01) + + # When: a list request starts reconciliation + await controls_module._seed_out_of_box_controls_for_namespace( + namespace_key="bounded-reconciliation-namespace" + ) + + # Then: the deadline cancels the database work and returns control to the request + assert cancelled.is_set() + + +@pytest.mark.asyncio +async def test_resolve_clone_name_reports_generated_name_exhaustion() -> None: + control_service = MagicMock() + control_service.active_control_name_exists = AsyncMock(return_value=True) + + with pytest.raises(APIError) as exc_info: + await controls_module._resolve_clone_name( + control_service, + namespace_key=DEFAULT_NAMESPACE_KEY, + source_id=1, + source_name="source-control", + requested_name=None, + ) + + assert exc_info.value.error_code == ErrorCode.CONTROL_NAME_CONFLICT + assert control_service.active_control_name_exists.await_count == 5 + + def test_clone_and_bind_context_drops_invalid_target_fields( client: TestClient, ) -> None: diff --git a/server/tests/test_data_model_v1_alembic_migration.py b/server/tests/test_data_model_v1_alembic_migration.py index 53334732..00c27e56 100644 --- a/server/tests/test_data_model_v1_alembic_migration.py +++ b/server/tests/test_data_model_v1_alembic_migration.py @@ -2,22 +2,23 @@ from __future__ import annotations +import json import uuid from pathlib import Path import pytest +from agent_control_server.config import db_config +from alembic import command from alembic.config import Config from sqlalchemy import create_engine, inspect, text from sqlalchemy.engine import Engine, make_url -from agent_control_server.config import db_config -from alembic import command - SERVER_DIR = Path(__file__).resolve().parents[1] PRE_MIGRATION_REVISION = "c1e9f9c4a1d2" MIGRATION_REVISION = "a7f3b1e0d9c5" OBSERVABILITY_NAMESPACE_REVISION = "b6f4c2d8e9a1" CLONE_LINEAGE_REVISION = "e2b7f4a9c6d1" +SEED_IDENTITY_REVISION = "f3a1c8d7e2b4" _BASE_DB_URL = make_url(db_config.get_url()) pytestmark = pytest.mark.skipif( @@ -396,6 +397,100 @@ def test_control_clone_lineage_migration_adds_composite_fk_and_partial_index( assert "ix_events_agent_time" not in indexes +def test_seed_identity_downgrade_retires_canonical_name_control( + alembic_config: Config, + temp_engine: Engine, +) -> None: + # Given: the seeded control and its initial version use the new selector + command.upgrade(alembic_config, SEED_IDENTITY_REVISION) + control_data = { + "condition": { + "selector": {"path": "canonical_name"}, + "evaluator": {"name": "list", "config": {"values": ["web_search"]}}, + }, + "action": {"decision": "deny"}, + } + with temp_engine.begin() as conn: + control_id = conn.execute( + text( + """ + INSERT INTO controls (namespace_key, name, data, seed_source_id) + VALUES ( + 'default', + 'oob-only-approved-tools-may-run', + CAST(:data AS jsonb), + 'oob-only-approved-tools-may-run' + ) + RETURNING id + """ + ), + {"data": json.dumps(control_data)}, + ).scalar_one() + conn.execute( + text( + """ + INSERT INTO control_versions ( + control_id, version_num, event_type, snapshot, note + ) + VALUES ( + :control_id, + 1, + 'created', + CAST(:snapshot AS jsonb), + 'Out-of-box control seed' + ) + """ + ), + { + "control_id": control_id, + "snapshot": json.dumps( + { + "name": "oob-only-approved-tools-may-run", + "data": control_data, + } + ), + }, + ) + + # When: rolling back to the last revision that rejects canonical_name + command.downgrade(alembic_config, CLONE_LINEAGE_REVISION) + + # Then: the seed is inactive and all persisted definitions use a supported selector + with temp_engine.begin() as conn: + control = conn.execute( + text( + """ + SELECT + deleted_at, + data #>> '{condition,selector,path}' AS selector_path + FROM controls + WHERE id = :control_id + """ + ), + {"control_id": control_id}, + ).mappings().one() + snapshot_selector_path = conn.execute( + text( + """ + SELECT snapshot #>> '{data,condition,selector,path}' + FROM control_versions + WHERE control_id = :control_id + """ + ), + {"control_id": control_id}, + ).scalar_one() + + assert control["deleted_at"] is not None + assert control["selector_path"] == "name" + assert snapshot_selector_path == "name" + assert "seed_source_id" not in _column_names(temp_engine, "controls") + assert "seed_opted_out_at" not in _column_names(temp_engine, "controls") + assert "idx_controls_namespace_seed_source" not in _index_names( + temp_engine, + "controls", + ) + + def test_downgrade_rejects_cross_namespace_agents_duplicates( alembic_config: Config, temp_engine: Engine ) -> None: diff --git a/server/tests/test_evaluation_legacy_canonical_name.py b/server/tests/test_evaluation_legacy_canonical_name.py new file mode 100644 index 00000000..e6a7d6f4 --- /dev/null +++ b/server/tests/test_evaluation_legacy_canonical_name.py @@ -0,0 +1,58 @@ +"""Compatibility coverage for canonical tool names in evaluation requests.""" + +from agent_control_models import EvaluationRequest, Step +from fastapi.testclient import TestClient + +from .utils import create_and_assign_policy + + +def test_canonical_name_allowlist_supports_mixed_sdk_versions(client: TestClient) -> None: + # Given: an allowlist using the canonical tool identity sent by current SDKs + control_data = { + "description": "Allow web search", + "enabled": True, + "execution": "server", + "scope": {"step_types": ["tool"], "stages": ["pre"]}, + "selector": {"path": "canonical_name"}, + "evaluator": { + "name": "list", + "config": { + "values": ["web_search"], + "logic": "any", + "match_on": "no_match", + "match_mode": "exact", + "case_sensitive": False, + }, + }, + "action": {"decision": "deny"}, + } + agent_name, _ = create_and_assign_policy( + client, + control_data, + agent_name="MixedVersionAgent", + ) + legacy_request = EvaluationRequest( + agent_name=agent_name, + step=Step(type="tool", name="writer.web_search", input={}), + stage="pre", + ) + current_request = EvaluationRequest( + agent_name=agent_name, + step=Step( + type="tool", + name="writer.web_search", + canonical_name="web_search", + input={}, + ), + stage="pre", + ) + + # When: legacy and current SDK payloads are evaluated by the same server + responses = [ + client.post("/api/v1/evaluation", json=request.model_dump(mode="json")) + for request in (legacy_request, current_request) + ] + + # Then: the approved tool is allowed for both client versions + assert [response.status_code for response in responses] == [200, 200] + assert [response.json()["is_safe"] for response in responses] == [True, True] diff --git a/server/tests/test_out_of_box_controls_bootstrap.py b/server/tests/test_out_of_box_controls_bootstrap.py index 46386635..da44201e 100644 --- a/server/tests/test_out_of_box_controls_bootstrap.py +++ b/server/tests/test_out_of_box_controls_bootstrap.py @@ -6,7 +6,16 @@ from typing import cast import pytest +from agent_control_evaluators.json.config import JSONEvaluatorConfig +from agent_control_evaluators.json.evaluator import JSONEvaluator +from agent_control_evaluators.list.config import ListEvaluatorConfig +from agent_control_evaluators.list.evaluator import ListEvaluator +from agent_control_evaluators.regex.config import RegexEvaluatorConfig +from agent_control_evaluators.regex.evaluator import RegexEvaluator +from agent_control_evaluators.sql import SQLEvaluator, SQLEvaluatorConfig +from agent_control_models import EvaluatorSpec from agent_control_server.bootstrap.out_of_box_controls import ( + OUT_OF_BOX_CONTROL_TEMPLATES, OutOfBoxControlTemplate, default_out_of_box_namespace_key, missing_required_evaluators, @@ -22,10 +31,25 @@ ) from agent_control_server.services.controls import ControlService from pydantic import ValidationError -from sqlalchemy import Table, func, select +from sqlalchemy import Table, event, func, select from sqlalchemy.orm import Session -from .conftest import AsyncSessionTest, engine +from .conftest import AsyncSessionTest, async_engine, engine + +_EXPECTED_OOB_CONTROL_NAMES = ( + "oob-ssn-match", + "oob-credit-card-number-match", + "oob-phone-number-match", + "oob-dangerous-shell-command-match", + "oob-high-value-action-requires-approval", + "oob-outbound-communication-requires-approval", + "oob-only-approved-tools-may-run", + "oob-owasp-llm05-read-only-sql", + "oob-owasp-llm10-bounded-sql-query", + "oob-owasp-llm02-common-credential-output-match", + "oob-owasp-llm05-dangerous-uri-output-match", +) +_AVAILABLE_PHASE_2_EVALUATORS = {"regex", "json", "list", "sql"} def _control_payload(*, evaluator_name: str = "regex") -> dict[str, object]: @@ -75,10 +99,46 @@ def _count_table_rows(table: Table) -> int: return cast(int, session.scalar(select(func.count()).select_from(table))) +def _oob_evaluator_spec(name: str) -> EvaluatorSpec: + template = next(template for template in OUT_OF_BOX_CONTROL_TEMPLATES if template.name == name) + leaf = template.control.primary_leaf() + assert leaf is not None + leaf_parts = leaf.leaf_parts() + assert leaf_parts is not None + _, evaluator = leaf_parts + return evaluator + + def test_default_namespace_key_uses_standalone_namespace() -> None: assert default_out_of_box_namespace_key() == DEFAULT_NAMESPACE_KEY +def test_out_of_box_catalog_contains_phase_2_templates() -> None: + assert tuple(template.name for template in OUT_OF_BOX_CONTROL_TEMPLATES) == ( + _EXPECTED_OOB_CONTROL_NAMES + ) + assert { + evaluator + for template in OUT_OF_BOX_CONTROL_TEMPLATES + for evaluator in template.required_evaluators + } == _AVAILABLE_PHASE_2_EVALUATORS + approved_tools = next( + template + for template in OUT_OF_BOX_CONTROL_TEMPLATES + if template.name == "oob-only-approved-tools-may-run" + ) + approved_tools_leaf = approved_tools.control.primary_leaf() + assert approved_tools_leaf is not None + assert approved_tools_leaf.selector.path == "canonical_name" + sql_controls = [ + template + for template in OUT_OF_BOX_CONTROL_TEMPLATES + if "sql" in template.control.tags + ] + assert len(sql_controls) == 2 + assert all(template.control.scope.step_name_regex for template in sql_controls) + + def test_missing_required_evaluators_returns_sorted_names() -> None: missing = missing_required_evaluators( {"galileo.luna", "regex", "json"}, @@ -185,6 +245,85 @@ async def test_seed_creates_control_version_in_namespace_without_bindings() -> N assert _count_table_rows(ControlBinding.__table__) == 0 +@pytest.mark.asyncio +async def test_seed_default_catalog_creates_all_controls_without_bindings() -> None: + result = await seed_out_of_box_controls( + session_factory=AsyncSessionTest, + namespace_key=DEFAULT_NAMESPACE_KEY, + available_evaluators=_AVAILABLE_PHASE_2_EVALUATORS, + ) + + assert result.created == _EXPECTED_OOB_CONTROL_NAMES + assert result.skipped_existing == () + assert result.skipped_missing_evaluator == () + assert result.skipped_conflict == () + + controls = _fetch_controls() + assert tuple(control.name for control in controls) == _EXPECTED_OOB_CONTROL_NAMES + assert {control.namespace_key for control in controls} == {DEFAULT_NAMESPACE_KEY} + assert len(_fetch_versions()) == len(_EXPECTED_OOB_CONTROL_NAMES) + assert _count_table_rows(policy_controls) == 0 + assert _count_table_rows(agent_controls) == 0 + assert _count_table_rows(ControlBinding.__table__) == 0 + + +@pytest.mark.asyncio +async def test_seed_default_catalog_is_idempotent() -> None: + await seed_out_of_box_controls( + session_factory=AsyncSessionTest, + namespace_key=DEFAULT_NAMESPACE_KEY, + available_evaluators=_AVAILABLE_PHASE_2_EVALUATORS, + ) + + result = await seed_out_of_box_controls( + session_factory=AsyncSessionTest, + namespace_key=DEFAULT_NAMESPACE_KEY, + available_evaluators=_AVAILABLE_PHASE_2_EVALUATORS, + ) + + assert result.created == () + assert result.skipped_existing == _EXPECTED_OOB_CONTROL_NAMES + assert len(_fetch_controls()) == len(_EXPECTED_OOB_CONTROL_NAMES) + assert len(_fetch_versions()) == len(_EXPECTED_OOB_CONTROL_NAMES) + + +@pytest.mark.asyncio +async def test_seed_existing_catalog_uses_one_bulk_lookup() -> None: + # Given: a namespace whose complete catalog is already seeded + await seed_out_of_box_controls( + session_factory=AsyncSessionTest, + namespace_key=DEFAULT_NAMESPACE_KEY, + available_evaluators=_AVAILABLE_PHASE_2_EVALUATORS, + ) + statements: list[str] = [] + + def record_statement( + _conn: object, + _cursor: object, + statement: str, + _parameters: object, + _context: object, + _executemany: bool, + ) -> None: + statements.append(statement) + + event.listen(async_engine.sync_engine, "before_cursor_execute", record_statement) + try: + # When: reconciling the already-seeded namespace + result = await seed_out_of_box_controls( + session_factory=AsyncSessionTest, + namespace_key=DEFAULT_NAMESPACE_KEY, + available_evaluators=_AVAILABLE_PHASE_2_EVALUATORS, + ) + finally: + event.remove(async_engine.sync_engine, "before_cursor_execute", record_statement) + + # Then: all seed identities are resolved by one database statement + assert result.skipped_existing == _EXPECTED_OOB_CONTROL_NAMES + assert len(statements) == 1 + assert statements[0].lstrip().startswith("SELECT") + + @pytest.mark.asyncio async def test_seed_is_idempotent_for_existing_active_control_names() -> None: template = _template(name="oob-idempotent-control") @@ -277,25 +416,16 @@ async def test_seed_treats_duplicate_insert_integrity_error_as_skip( templates=(template,), ) - async def active_control_name_exists( + async def find_existing_seed_controls( self: ControlService, - name: str, *, namespace_key: str, - exclude_control_id: int | None = None, - ) -> bool: - return False + source_ids: object, + names: object, + ) -> tuple[frozenset[str], frozenset[str]]: + return frozenset(), frozenset() - async def seed_source_exists( - self: ControlService, - seed_source_id: str, - *, - namespace_key: str, - ) -> bool: - return False - - monkeypatch.setattr(ControlService, "active_control_name_exists", active_control_name_exists) - monkeypatch.setattr(ControlService, "seed_source_exists", seed_source_exists) + monkeypatch.setattr(ControlService, "find_existing_seed_controls", find_existing_seed_controls) result = await seed_out_of_box_controls( session_factory=AsyncSessionTest, @@ -309,3 +439,174 @@ async def seed_source_exists( assert result.skipped_conflict == ("oob-race-control",) assert len(_fetch_controls()) == 1 assert len(_fetch_versions()) == 1 + + +@pytest.mark.asyncio +async def test_regex_out_of_box_controls_match_representative_payloads() -> None: + ssn_spec = _oob_evaluator_spec("oob-ssn-match") + ssn_evaluator = RegexEvaluator(RegexEvaluatorConfig.model_validate(ssn_spec.config)) + ssn_result = await ssn_evaluator.evaluate("Customer SSN is 123-45-6789.") + assert ssn_result.matched is True + + shell_spec = _oob_evaluator_spec("oob-dangerous-shell-command-match") + shell_evaluator = RegexEvaluator(RegexEvaluatorConfig.model_validate(shell_spec.config)) + for command in ( + "sudo rm -rf /", + "rm -rf /", + "rm -rf ~", + "chmod -R 777 /", + "chown -R root /", + ): + shell_result = await shell_evaluator.evaluate(command) + assert shell_result.matched is True, command + + +@pytest.mark.asyncio +async def test_dangerous_shell_control_matches_equivalent_recursive_rm_forms() -> None: + # Given: the destructive shell command control + shell_spec = _oob_evaluator_spec("oob-dangerous-shell-command-match") + shell_evaluator = RegexEvaluator(RegexEvaluatorConfig.model_validate(shell_spec.config)) + + # When: evaluating equivalent recursive deletion spellings and a scoped deletion + destructive_results = [ + await shell_evaluator.evaluate(command) + for command in ( + 'rm -rf "$HOME"', + "rm -rf ~/", + "rm -fr /", + "rm -r -f /", + "rm -f -r '$HOME/'", + ) + ] + scoped_result = await shell_evaluator.evaluate("rm -rf /tmp/build-output") + + # Then: equivalent root/home deletions are blocked without blocking scoped deletion + assert all(result.matched is True for result in destructive_results) + assert scoped_result.matched is False + + +@pytest.mark.asyncio +async def test_json_out_of_box_controls_ignore_caller_controlled_approval_flags() -> None: + high_value_spec = _oob_evaluator_spec("oob-high-value-action-requires-approval") + high_value_evaluator = JSONEvaluator( + JSONEvaluatorConfig.model_validate(high_value_spec.config) + ) + + high_value_result = await high_value_evaluator.evaluate({"amount": 25000}) + low_value_result = await high_value_evaluator.evaluate({"amount": 250}) + caller_approved_results = [ + await high_value_evaluator.evaluate({"amount": 25000, "approved": True}), + await high_value_evaluator.evaluate( + {"amount": 25000, "approval": {"approved": True}} + ), + ] + + assert high_value_result.matched is True + assert low_value_result.matched is False + assert all(result.matched is True for result in caller_approved_results) + + outbound_spec = _oob_evaluator_spec("oob-outbound-communication-requires-approval") + outbound_evaluator = JSONEvaluator(JSONEvaluatorConfig.model_validate(outbound_spec.config)) + + outbound_result = await outbound_evaluator.evaluate( + {"to": "customer@example.com", "message": "Hello"} + ) + internal_result = await outbound_evaluator.evaluate({"query": "customer history"}) + caller_approved_outbound_results = [ + await outbound_evaluator.evaluate( + {"to": "customer@example.com", "message": "Hello", "approved": True} + ), + await outbound_evaluator.evaluate( + { + "to": "customer@example.com", + "message": "Hello", + "approval": {"approved": True}, + } + ), + ] + + assert outbound_result.matched is True + assert internal_result.matched is False + assert all(result.matched is True for result in caller_approved_outbound_results) + + +@pytest.mark.asyncio +async def test_list_out_of_box_control_matches_unapproved_tools() -> None: + tool_spec = _oob_evaluator_spec("oob-only-approved-tools-may-run") + tool_evaluator = ListEvaluator(ListEvaluatorConfig.model_validate(tool_spec.config)) + + delete_result = await tool_evaluator.evaluate("delete_user") + search_result = await tool_evaluator.evaluate("web_search") + + assert delete_result.matched is True + assert search_result.matched is False + + +@pytest.mark.asyncio +async def test_owasp_credential_control_matches_common_secret_formats() -> None: + # Given: the OWASP-aligned common credential output control + spec = _oob_evaluator_spec("oob-owasp-llm02-common-credential-output-match") + evaluator = RegexEvaluator(RegexEvaluatorConfig.model_validate(spec.config)) + + # When: evaluating representative secret and non-secret output + private_key_result = await evaluator.evaluate("-----BEGIN OPENSSH PRIVATE KEY-----\nredacted") + aws_key_result = await evaluator.evaluate("Credential: AKIAIOSFODNN7EXAMPLE") + safe_result = await evaluator.evaluate("The operation completed successfully.") + + # Then: recognizable credentials are blocked while ordinary output passes + assert private_key_result.matched is True + assert aws_key_result.matched is True + assert safe_result.matched is False + + +@pytest.mark.asyncio +async def test_owasp_dangerous_uri_control_matches_active_content_schemes() -> None: + # Given: the OWASP-aligned dangerous URI output control + spec = _oob_evaluator_spec("oob-owasp-llm05-dangerous-uri-output-match") + evaluator = RegexEvaluator(RegexEvaluatorConfig.model_validate(spec.config)) + + # When: evaluating executable, active-content, and ordinary HTTPS links + javascript_result = await evaluator.evaluate( + 'click' + ) + data_uri_result = await evaluator.evaluate("data:image/svg+xml,") + safe_result = await evaluator.evaluate("https://docs.example.com/safety") + + # Then: active-content schemes are blocked while HTTPS passes + assert javascript_result.matched is True + assert data_uri_result.matched is True + assert safe_result.matched is False + + +@pytest.mark.asyncio +async def test_owasp_read_only_sql_control_blocks_mutation_and_multiple_statements() -> None: + # Given: the OWASP-aligned read-only SQL control + spec = _oob_evaluator_spec("oob-owasp-llm05-read-only-sql") + evaluator = SQLEvaluator(SQLEvaluatorConfig.model_validate(spec.config)) + + # When: evaluating read-only, mutating, and multi-statement SQL + select_result = await evaluator.evaluate("SELECT id FROM users") + delete_result = await evaluator.evaluate("DELETE FROM users") + multiple_result = await evaluator.evaluate("SELECT id FROM users; DROP TABLE users") + + # Then: only the single read-only query passes + assert select_result.matched is False + assert delete_result.matched is True + assert multiple_result.matched is True + + +@pytest.mark.asyncio +async def test_owasp_bounded_sql_control_enforces_result_and_complexity_limits() -> None: + # Given: the OWASP-aligned bounded SQL query control + spec = _oob_evaluator_spec("oob-owasp-llm10-bounded-sql-query") + evaluator = SQLEvaluator(SQLEvaluatorConfig.model_validate(spec.config)) + + # When: evaluating bounded, unbounded, and oversized result windows + bounded_result = await evaluator.evaluate("SELECT id FROM users LIMIT 100") + missing_limit_result = await evaluator.evaluate("SELECT id FROM users") + oversized_window_result = await evaluator.evaluate("SELECT id FROM users LIMIT 1000 OFFSET 1") + + # Then: only the bounded query within the configured result window passes + assert bounded_result.matched is False + assert missing_limit_result.matched is True + assert oversized_window_result.matched is True diff --git a/server/tests/test_principal_namespace_flow.py b/server/tests/test_principal_namespace_flow.py index 8f16a795..295c6573 100644 --- a/server/tests/test_principal_namespace_flow.py +++ b/server/tests/test_principal_namespace_flow.py @@ -11,9 +11,13 @@ Principal, set_authorizer, ) +from agent_control_server.bootstrap.out_of_box_controls import OUT_OF_BOX_CONTROL_TEMPLATES +from agent_control_server.models import Control from fastapi import FastAPI, Request from fastapi.testclient import TestClient +from sqlalchemy.orm import Session +from .conftest import engine from .utils import VALID_CONTROL_PAYLOAD @@ -39,6 +43,24 @@ async def authorize( ) +class ControlsReadOnlyAuthorizer(HeaderNamespaceAuthorizer): + """Allow the controls list read and record every authorization operation.""" + + def __init__(self) -> None: + self.operations: list[Operation] = [] + + async def authorize( + self, + request: Request, + operation: Operation, + context: dict[str, Any] | None = None, + ) -> Principal: + self.operations.append(operation) + if operation is not Operation.CONTROLS_READ: + raise AssertionError(f"Unexpected authorization operation: {operation}") + return await super().authorize(request, operation, context) + + def _client(app: FastAPI, namespace_key: str) -> TestClient: return TestClient( app, @@ -73,6 +95,103 @@ def _evaluation_payload(agent_name: str) -> dict[str, Any]: } +def test_controls_list_seeds_out_of_box_controls_for_principal_namespace( + app: FastAPI, +) -> None: + authorizer = ControlsReadOnlyAuthorizer() + set_authorizer(authorizer) + + namespace_client = _client(app, "org-oob-controls") + filtered = namespace_client.get("/api/v1/controls", params={"name": "oob"}) + assert filtered.status_code == 200, filtered.text + assert filtered.json()["controls"] == [] + + resp = namespace_client.get( + "/api/v1/controls", + params={"limit": len(OUT_OF_BOX_CONTROL_TEMPLATES)}, + ) + assert resp.status_code == 200, resp.text + + expected_names = {template.name for template in OUT_OF_BOX_CONTROL_TEMPLATES} + returned_names = {control["name"] for control in resp.json()["controls"]} + assert expected_names.issubset(returned_names) + assert authorizer.operations == [ + Operation.CONTROLS_READ, + Operation.CONTROLS_READ, + ] + + +def test_controls_list_seeds_out_of_box_controls_alongside_custom_control( + app: FastAPI, +) -> None: + set_authorizer(HeaderNamespaceAuthorizer()) + namespace_client = _client(app, "org-with-custom-control") + custom_name = f"custom-{uuid.uuid4().hex[:12]}" + + created = namespace_client.put( + "/api/v1/controls", + json={"name": custom_name, "data": VALID_CONTROL_PAYLOAD}, + ) + assert created.status_code == 200, created.text + + response = namespace_client.get( + "/api/v1/controls", + params={"limit": 20, "cloned": "false"}, + ) + assert response.status_code == 200, response.text + + returned_names = {control["name"] for control in response.json()["controls"]} + expected_names = {template.name for template in OUT_OF_BOX_CONTROL_TEMPLATES} + assert returned_names == {*expected_names, custom_name} + + +def test_controls_list_completes_partially_seeded_namespace(app: FastAPI) -> None: + set_authorizer(HeaderNamespaceAuthorizer()) + namespace_key = "org-partially-seeded" + first_template = OUT_OF_BOX_CONTROL_TEMPLATES[0] + with Session(engine) as session: + session.add( + Control( + namespace_key=namespace_key, + name=first_template.name, + data=first_template.control.model_dump( + mode="json", + by_alias=True, + exclude_none=True, + exclude_unset=True, + ), + seed_source_id=first_template.source_id, + ) + ) + session.commit() + + response = _client(app, namespace_key).get( + "/api/v1/controls", + params={"limit": 20, "cloned": "false"}, + ) + assert response.status_code == 200, response.text + + returned_names = {control["name"] for control in response.json()["controls"]} + expected_names = {template.name for template in OUT_OF_BOX_CONTROL_TEMPLATES} + assert returned_names == expected_names + + +def test_controls_list_retries_out_of_box_seeding_for_default_namespace( + app: FastAPI, +) -> None: + set_authorizer(HeaderNamespaceAuthorizer()) + + response = _client(app, "default").get( + "/api/v1/controls", + params={"limit": 20, "cloned": "false"}, + ) + assert response.status_code == 200, response.text + + returned_names = {control["name"] for control in response.json()["controls"]} + expected_names = {template.name for template in OUT_OF_BOX_CONTROL_TEMPLATES} + assert returned_names == expected_names + + def test_principal_namespace_scopes_management_and_runtime(app: FastAPI) -> None: set_authorizer(HeaderNamespaceAuthorizer())