From 6a5e818edf74f2a64a26ed260ca671535a656c07 Mon Sep 17 00:00:00 2001 From: Odin Date: Sun, 27 Sep 2026 19:45:29 -0400 Subject: [PATCH 01/12] Restore request-bound strict Codex tools with canonical validation --- docs/strict-wire-lowerings.md | 68 ++++ src/discord/background_task.py | 23 ++ src/discord/native_tools/agents_tasks.py | 19 ++ src/discord/native_tools/scheduling.py | 32 +- src/discord/native_tools/skills_tools.py | 7 + src/discord/scheduled_events.py | 18 ++ src/discord/wiring.py | 2 +- src/llm/openai_codex.py | 376 ++++++++++++---------- src/llm/strict_tool_adapter.py | 381 +++++++++++++++++++++++ src/scheduler/scheduler.py | 5 + src/tools/defs/integrations_email.py | 18 +- src/tools/defs/media_scheduling.py | 8 + src/tools/handlers/browser_web.py | 10 +- src/tools/http_probe_ops.py | 29 +- src/tools/nested_payload.py | 123 ++++++++ src/tools/post_validation.py | 22 ++ tests/test_openai_codex_client.py | 7 +- tests/test_strict_nested_payload.py | 56 ++++ tests/test_strict_tool_adapter.py | 120 +++++++ tests/test_strict_wire_defs.py | 68 ++++ 20 files changed, 1211 insertions(+), 181 deletions(-) create mode 100644 docs/strict-wire-lowerings.md create mode 100644 src/llm/strict_tool_adapter.py create mode 100644 src/tools/nested_payload.py create mode 100644 tests/test_strict_nested_payload.py create mode 100644 tests/test_strict_tool_adapter.py create mode 100644 tests/test_strict_wire_defs.py diff --git a/docs/strict-wire-lowerings.md b/docs/strict-wire-lowerings.md new file mode 100644 index 000000000..7b847e148 --- /dev/null +++ b/docs/strict-wire-lowerings.md @@ -0,0 +1,68 @@ +# Strict wire lowerings + +This document inventories built-in wire shapes that differ from their existing +runtime interfaces. Lowerings happen before effectful dispatch and must be +tested at both the wire and runtime boundaries. + +## `http_probe.headers` + +- Canonical path: `http_probe.headers`, historically an arbitrary-key object. +- Wire form: array of closed `{name: string, value: string}` objects. +- Runtime: `build_http_probe_command` converts records to the existing header + dictionary before curl command construction; dictionary callers remain + supported. Case-folded duplicate names, malformed records, and non-string + names/values are rejected before request execution. Existing curl header + validation and output redaction remain in place. +- Fixtures: `tests/test_strict_wire_defs.py` covers schema, list and legacy + dict positive cases, case-insensitive duplicate names and malformed records. + +## Schedule triggers + +- Canonical paths: `schedule_task.trigger` and `update_schedule.trigger`. +- Wire form: shared closed object with `source`, `event`, `repo`, and + `alert_name`; source retains the five supported values. +- Semantics remain AND across supplied conditions; partial conditions without + `source` remain valid. Existing scheduler runtime rejects an empty trigger. +- Fixtures: `tests/test_strict_wire_defs.py` asserts both definitions use the + same closed four-field shape. + +## `validate_action.checks[].expected` + +- Canonical path: `validate_action.checks[].expected`. +- Wire form: integer, string, integer array, or string array, with typed items. +- Runtime: `parse_checks` rejects unsupported types and incompatible typed + expectations before any validation command runs. HTTP expectations are + status integers or digit strings; service status lists contain strings. + Existing scalar stringification paths for command checks are preserved. +- Omitted expectations, explicit expectations, and check-type defaults remain + distinct: only an omitted `expected` selects the existing built-in default. +- Fixtures: `tests/test_strict_wire_defs.py` covers integer lists, invalid + mixed HTTP values, invalid service list values, and command compatibility. + +## `computer_*` explicit lowerings + +The canonical computer contracts remain in `src/tools/defs/computer.py`. +The request-local adapter owns the wrapper lowering: a closed outer object +contains `payload`, whose schema is a nested per-operation union with const +operation selectors. It must not use a root `anyOf`, and must reject ambiguous +matches rather than choosing a first branch. Canonical schemas stay unchanged. + +The following unsupported wire constraints require explicit computer-only +lowering. Each wire relaxation must be enforced by the canonical JSON Schema +tripwire before dispatch; these are not generic recursive keyword deletions. + +| Canonical schema constraint | Wire representation | Pre-effect validator | +| --- | --- | --- | +| `uniqueItems` (notably modifier arrays) | Omitted from wire array schema | Canonical computer payload validator rejects duplicates, including repeated key modifiers | +| `minProperties` (task context) | Omitted from wire object schema | Canonical validator rejects empty task context where minimum properties are required | +| `not` | Omitted from wire branch | Canonical validator rejects the forbidden shape | +| `allOf` | Omitted from wire branch | Canonical validator validates all constituent schemas | +| `oneOf` | Represented by nested operation branches only when match is unambiguous | Canonical validator and ambiguity check reject zero/multiple distinct matches | +| `false` property schemas | Property omitted from wire branch | Canonical validator rejects forbidden operation fields | + +Expected fixture coverage includes duplicate modifiers, repeated key +modifiers, empty task context, coordinate-versus-region conflicts, forbidden +operation fields, and nested sequence/stroke contracts. Fixtures and mocked +dispatch only: no desktop input is performed. The corresponding adapter test +module should assert positive and negative fixtures against compiled wire form +and canonical pre-effect validation. diff --git a/src/discord/background_task.py b/src/discord/background_task.py index 11a02bc55..ed3ba1c1c 100644 --- a/src/discord/background_task.py +++ b/src/discord/background_task.py @@ -102,6 +102,7 @@ class BackgroundTask: channel: discord.abc.Messageable requester: str requester_id: str = "" + nested_payload_validated: bool = False created_at: str = field(default_factory=lambda: datetime.now().isoformat()) status: str = "running" # running, completed, failed, cancelled results: list[StepResult] = field(default_factory=list) @@ -188,6 +189,28 @@ async def run_background_task( # Variable substitution in tool_input string values tool_input = _substitute_vars(tool_input, variables, prev_output) + if task.nested_payload_validated: + try: + from ..tools.nested_payload import validate_nested_payload + catalog = getattr(executor, "_tool_catalog", None) + definitions = catalog.merged_definitions() if catalog is not None else [] + validate_nested_payload("delegate_task", {"steps": [{**step, "tool_input": tool_input}]}, + definitions, allow_placeholders=False) + denial = executor.check_permission(tool_name, task.requester_id) + if denial: + raise ValueError(denial) + if tool_name == "invoke_skill": + target = tool_input.get("name") + denial = executor.check_permission(target, task.requester_id) + if denial: + raise ValueError(denial) + except ValueError as exc: + task.results.append(StepResult(index=i, tool_name=tool_name, + description=step_desc, status="error", output=f"Invalid concrete step payload: {exc}")) + if on_failure == "abort": + task.status = "failed" + break + continue # Evaluate condition if condition and prev_output: diff --git a/src/discord/native_tools/agents_tasks.py b/src/discord/native_tools/agents_tasks.py index 08de15b00..91c5ea15f 100644 --- a/src/discord/native_tools/agents_tasks.py +++ b/src/discord/native_tools/agents_tasks.py @@ -28,6 +28,7 @@ from ...odin_log import get_logger from ...tools.defs.agents import SPAWN_NEUTRAL_REASONING_OPTIONS from ...tools.result_validator import ToolResult +from ...tools.nested_payload import validate_nested_payload from ..background_task import ( MAX_STEPS, BackgroundTask, @@ -752,6 +753,23 @@ def __init__(self, deps: AgentTaskDeps) -> None: async def _handle_delegate_task(self, message: discord.Message, inp: dict) -> str: """Create and start a background task.""" + try: + inp = validate_nested_payload( + "delegate_task", inp, + self._tool_catalog.merged_definitions() if self._tool_catalog else [], + ) + except ValueError as e: + return f"Invalid background task payload: {e}" + for i, step in enumerate(inp.get("steps", []), 1): + if isinstance(step, dict) and step.get("tool_name"): + denied = self._tool_executor.check_permission(step["tool_name"], str(message.author.id)) + if denied: + return f"Step {i}: {denied}" + if step["tool_name"] == "invoke_skill" and isinstance(step.get("tool_input"), dict): + target = step["tool_input"].get("name") + denied = self._tool_executor.check_permission(target, str(message.author.id)) + if denied: + return f"Step {i}: {denied}" description = inp.get("description", "Background task") steps = inp.get("steps", []) @@ -787,6 +805,7 @@ async def _handle_delegate_task(self, message: discord.Message, inp: dict) -> st channel=message.channel, requester=str(message.author), requester_id=str(message.author.id), + nested_payload_validated=True, ) # Prune old completed tasks diff --git a/src/discord/native_tools/scheduling.py b/src/discord/native_tools/scheduling.py index b1ea0900c..100fecce8 100644 --- a/src/discord/native_tools/scheduling.py +++ b/src/discord/native_tools/scheduling.py @@ -10,13 +10,15 @@ from ...odin_log import get_logger from ...scheduler.scheduler import ScheduleConnectionUnavailableError +from ...tools.nested_payload import validate_nested_payload log = get_logger("discord") class SchedulingTools: - def __init__(self, *, scheduler) -> None: + def __init__(self, *, scheduler, tool_catalog=None) -> None: self.scheduler = scheduler + self.tool_catalog = tool_catalog # -- creation-time validation --------------------------------------------- @@ -100,6 +102,10 @@ def _extract_tool_input_from_steps(inp: dict) -> dict | None: async def _handle_schedule_task(self, message, inp: dict) -> str: """Create a scheduled task.""" + try: + inp = validate_nested_payload("schedule_task", inp, self._nested_catalog()) + except ValueError as e: + return f"Failed to create schedule: {e}" validation_error = self._validate_schedule_payload(inp) if validation_error: return f"Failed to create schedule: {validation_error}" @@ -118,6 +124,7 @@ async def _handle_schedule_task(self, message, inp: dict) -> str: cron_timezone=inp.get("cron_timezone"), requester_id=str(message.author.id), report_format=inp.get("report_format"), + nested_payload_validated=True, ) if schedule.get("trigger"): trigger_desc = ", ".join(f"{k}={v}" for k, v in schedule["trigger"].items()) @@ -164,6 +171,18 @@ def _handle_list_schedules(self) -> str: async def _handle_update_schedule(self, inp: dict) -> str: """Update an existing schedule.""" + try: + inp = validate_nested_payload("update_schedule", inp, self._nested_catalog()) + if isinstance(inp.get("tool_input"), dict) and not inp.get("tool_name"): + current = next((s for s in self.scheduler.list_all() if s.get("id") == inp.get("schedule_id")), None) + target_name = (current or {}).get("tool_name") + if target_name: + validate_nested_payload("schedule_task", { + "action": "check", "tool_name": target_name, + "tool_input": inp["tool_input"], + }, self._nested_catalog()) + except ValueError as e: + return f"Error: {e}" schedule_id = inp.get("schedule_id", "") if not schedule_id: return "Error: 'schedule_id' is required." @@ -185,6 +204,7 @@ async def _handle_update_schedule(self, inp: dict) -> str: trigger = inp.get("trigger") if trigger is not None: kwargs["trigger"] = trigger + kwargs["nested_payload_validated"] = True if "paused" in inp: val = inp["paused"] if not isinstance(val, bool): @@ -207,6 +227,16 @@ async def _handle_update_schedule(self, inp: dict) -> str: ) return f"Updated schedule {schedule_id}." + def _nested_catalog(self): + """Canonical tool definitions for selected target validation.""" + catalog = getattr(self, "tool_catalog", None) + if catalog is not None: + return catalog.merged_definitions() + # ToolExecutor is available via the scheduler handler owner in production; + # static definitions still validate direct unit/legacy entry points. + from ...tools.registry import get_tool_definitions + return get_tool_definitions() + async def _handle_delete_schedule(self, inp: dict) -> str: """Delete a scheduled task.""" schedule_id = inp.get("schedule_id", "") diff --git a/src/discord/native_tools/skills_tools.py b/src/discord/native_tools/skills_tools.py index 252b10f78..d97ab57dd 100644 --- a/src/discord/native_tools/skills_tools.py +++ b/src/discord/native_tools/skills_tools.py @@ -24,6 +24,7 @@ import discord from ...odin_log import get_logger +from ...tools.nested_payload import validate_nested_payload from ..response_guards import scrub_response_secrets if TYPE_CHECKING: @@ -144,6 +145,12 @@ async def dispatch( "Use list_skills to see available skills.", effects, ) + try: + tool_input = validate_nested_payload( + "invoke_skill", tool_input, self.tool_catalog.merged_definitions() + ) + except ValueError as exc: + return f"Error: {exc}", effects skill_input = tool_input.get("input") or {} if not isinstance(skill_input, dict): return "Error: invoke_skill 'input' must be an object.", effects diff --git a/src/discord/scheduled_events.py b/src/discord/scheduled_events.py index 399df0a91..71559b068 100644 --- a/src/discord/scheduled_events.py +++ b/src/discord/scheduled_events.py @@ -24,6 +24,7 @@ from ..odin_log import get_logger from ..scheduler.scheduler import NonRetryableScheduleError from ..tools import ToolResult +from ..tools.nested_payload import validate_nested_payload from .delivery import close_open_fence from .mcp_dispatch import uncertain_outcome as mcp_uncertain_outcome from .response_guards import scrub_response_secrets @@ -306,6 +307,8 @@ async def _run_scheduled_workflow( results: list[str] = [] prev_output = "" workflow_ok = True + nested_validated = bool(schedule.get("_nested_payload_validated")) + catalog = self._tool_loop._tool_catalog.merged_definitions() if nested_validated else [] # Scheduled work runs under the identity of whoever created the schedule, so # host-access scoping / tier limits apply (None = unrestricted system task). req_id = schedule.get("requester_id") or None @@ -313,6 +316,21 @@ async def _run_scheduled_workflow( for i, step in enumerate(steps): tool_name = step["tool_name"] tool_input = step.get("tool_input", {}) + if nested_validated: + try: + validate_nested_payload( + "delegate_task", {"steps": [{**step, "tool_input": tool_input}]}, + catalog, allow_placeholders=False, + ) + except ValueError as exc: + results.append(f"**Step {i + 1}** (`{step.get('description', tool_name)}`): invalid payload: {exc}") + workflow_ok = False + break + denial = self._tool_executor.check_permission(tool_name, req_id) + if denial: + results.append(f"**Step {i + 1}**: {denial}") + workflow_ok = False + break condition = step.get("condition") step_desc = step.get("description", tool_name) diff --git a/src/discord/wiring.py b/src/discord/wiring.py index d5d184953..4f1fd0548 100644 --- a/src/discord/wiring.py +++ b/src/discord/wiring.py @@ -845,7 +845,7 @@ def build_components(bot, services: BotServices) -> BotComponents: # Domain handler bundles (P5b) — built BEFORE the dispatcher so they can # be its owners (RFC-002 P5). - scheduling_tools = SchedulingTools(scheduler=services.scheduler) + scheduling_tools = SchedulingTools(scheduler=services.scheduler, tool_catalog=tool_catalog) knowledge_tools = KnowledgeTools( sessions=services.sessions, # live: swappable at runtime via the bot's `knowledge` property diff --git a/src/llm/openai_codex.py b/src/llm/openai_codex.py index 5fdfc2afe..7139055d0 100644 --- a/src/llm/openai_codex.py +++ b/src/llm/openai_codex.py @@ -5,6 +5,7 @@ import re import string import unicodedata +from contextvars import ContextVar import aiohttp @@ -27,6 +28,9 @@ from .types import LLMResponse, ToolCall log = get_logger("codex") +_request_tool_adapter: ContextVar[object | None] = ContextVar( + "codex_request_tool_adapter", default=None +) def _reject_known_bad_pair(model: str | None, effort: str | None) -> None: @@ -40,6 +44,7 @@ def _reject_known_bad_pair(model: str | None, effort: str | None) -> None: if err: raise LLMRequestError(err) + CODEX_API_URL = "https://chatgpt.com/backend-api/codex/responses" # Streaming transport timeouts (config-overridable via the ctor). @@ -89,9 +94,7 @@ async def _read_error_body_bounded(content) -> tuple[bytes, bool]: return b"".join(chunks), overflowed -_ASCII_MIME_TOKEN_CHARS = frozenset( - string.ascii_letters + string.digits + "!#$%&'*+-.^_`|~" -) +_ASCII_MIME_TOKEN_CHARS = frozenset(string.ascii_letters + string.digits + "!#$%&'*+-.^_`|~") def _safe_mime(content_type: str | None) -> str: @@ -147,15 +150,9 @@ def _clean_error_field(value: str, limit: int = 200) -> str: error's own fields).""" stripped = value.strip() line = stripped.splitlines()[0].strip() if stripped else "" - line = "".join( - ch for ch in line if ch == "\t" or not unicodedata.category(ch).startswith("C") - ) + line = "".join(ch for ch in line if ch == "\t" or not unicodedata.category(ch).startswith("C")) low = line.lower() - if ( - "]*>", line) - ): + if "]*>", line): return "" line = line.replace("@everyone", "@\u200beveryone").replace("@here", "@\u200bhere") return scrub_output_secrets(line)[:limit] @@ -184,9 +181,7 @@ def _sanitized_error_fields(container: dict) -> list[str]: return list(dict.fromkeys(fields)) -def _describe_error_body( - content_type: str | None, raw: bytes, structured: dict | None -) -> str: +def _describe_error_body(content_type: str | None, raw: bytes, structured: dict | None) -> str: """Bounded, structure-aware descriptor of a non-2xx response body. Raise sites embed THIS instead of raw body excerpts (``errors.py`` @@ -207,11 +202,13 @@ def _describe_error_body( # them; matched against both error.type and error.code. Observed live # (2026-07-29/30 sol degradation): type=service_unavailable_error with # code=server_is_overloaded, plus bare server_error. -_CAPACITY_ERROR_MARKERS = frozenset({ - "service_unavailable_error", - "server_is_overloaded", - "server_error", -}) +_CAPACITY_ERROR_MARKERS = frozenset( + { + "service_unavailable_error", + "server_is_overloaded", + "server_error", + } +) class CodexStreamError(RuntimeError): @@ -246,8 +243,7 @@ def __init__( @property def is_capacity(self) -> bool: return ( - self.error_type in _CAPACITY_ERROR_MARKERS - or self.error_code in _CAPACITY_ERROR_MARKERS + self.error_type in _CAPACITY_ERROR_MARKERS or self.error_code in _CAPACITY_ERROR_MARKERS ) @@ -401,9 +397,7 @@ def eligible_account_keys_snapshot(self) -> frozenset[str]: raw = self.auth.get_account_id() raw_ids = frozenset({raw}) if isinstance(raw, str) and raw else frozenset() return frozenset( - key - for account_id in raw_ids - if (key := opaque_account_key(account_id)) is not None + key for account_id in raw_ids if (key := opaque_account_key(account_id)) is not None ) except Exception: log.exception("Could not resolve eligible Codex account keys") @@ -487,7 +481,9 @@ def _auth_headers(token: str, account_id: str | None) -> dict: @leased_call async def chat( - self, messages: list[dict], system: str, + self, + messages: list[dict], + system: str, max_tokens: int | None = None, ) -> str: """Send a chat request via the Codex backend API (streaming). @@ -533,10 +529,12 @@ def _convert_messages(self, messages: list[dict]) -> list[dict]: if isinstance(source, dict) and source.get("type") == "base64": media_type = source.get("media_type", "image/png") data = source.get("data", "") - image_parts.append({ - "type": "input_image", - "image_url": f"data:{media_type};base64,{data}", - }) + image_parts.append( + { + "type": "input_image", + "image_url": f"data:{media_type};base64,{data}", + } + ) elif block.get("type") == "tool_use": text_parts.append(f"[Used tool: {block.get('name', 'unknown')}]") elif block.get("type") == "tool_result": @@ -559,11 +557,13 @@ def _convert_messages(self, messages: list[dict]) -> list[dict]: msg_content.append({"type": "input_text", "text": " ".join(text_parts)}) msg_content.extend(image_parts) if msg_content: - codex_messages.append({ - "type": "message", - "role": "user", - "content": msg_content, - }) + codex_messages.append( + { + "type": "message", + "role": "user", + "content": msg_content, + } + ) continue content = " ".join(text_parts) if not content: @@ -579,11 +579,13 @@ def _convert_messages(self, messages: list[dict]) -> list[dict]: # User messages use input_text, assistant messages use output_text content_type = "output_text" if role == "assistant" else "input_text" - codex_messages.append({ - "type": "message", - "role": role, - "content": [{"type": content_type, "text": content}], - }) + codex_messages.append( + { + "type": "message", + "role": role, + "content": [{"type": content_type, "text": content}], + } + ) return codex_messages # ------------------------------------------------------------------ @@ -592,26 +594,10 @@ def _convert_messages(self, messages: list[dict]) -> list[dict]: @staticmethod def _convert_tools(tools: list[dict]) -> list[dict]: - """Convert internal tool definitions to OpenAI function format. + """Compile the request-local catalog without mutating canonical tools.""" + from .strict_tool_adapter import compile_catalog - Internal: {"name": ..., "description": ..., "input_schema": {...}} - OpenAI: {"type": "function", "name": ..., "description": ..., "parameters": {...}} - - Codex transport strict mode materializes every declared property as a - required model-facing argument. Odin's schemas use ``required`` to - distinguish mandatory fields from optional ones, so non-strict is the - safe default. An explicit boolean override remains authoritative. - """ - return [ - { - "type": "function", - "name": t["name"], - "description": t.get("description", ""), - "parameters": t.get("input_schema", {"type": "object", "properties": {}}), - "strict": t["strict"] if type(t.get("strict")) is bool else False, - } - for t in tools - ] + return compile_catalog(tools).wire_tools def _convert_messages_with_tools(self, messages: list[dict]) -> list[dict]: """Convert internal message format to Codex Responses API format with tool support. @@ -635,12 +621,15 @@ def _convert_messages_with_tools(self, messages: list[dict]) -> list[dict]: if not content: continue ct = "output_text" if role == "assistant" else "input_text" - codex_input.append({ - "type": "message", - "role": (role - if role in ("user", "assistant", "developer", "system") else "user"), - "content": [{"type": ct, "text": content}], - }) + codex_input.append( + { + "type": "message", + "role": ( + role if role in ("user", "assistant", "developer", "system") else "user" + ), + "content": [{"type": ct, "text": content}], + } + ) continue if not isinstance(content, list): @@ -661,21 +650,28 @@ def _convert_messages_with_tools(self, messages: list[dict]) -> list[dict]: elif btype == "tool_use": # Flush any accumulated text first if text_parts: - codex_input.append({ - "type": "message", - "role": "assistant", - "content": [{"type": "output_text", "text": " ".join(text_parts)}], - }) + codex_input.append( + { + "type": "message", + "role": "assistant", + "content": [{"type": "output_text", "text": " ".join(text_parts)}], + } + ) text_parts = [] # Convert to OpenAI function_call item tool_input = block.get("input", {}) - codex_input.append({ - "type": "function_call", - "call_id": block.get("id", ""), - "name": block.get("name", ""), - "arguments": (json.dumps(tool_input) - if isinstance(tool_input, dict) else str(tool_input)), - }) + codex_input.append( + { + "type": "function_call", + "call_id": block.get("id", ""), + "name": block.get("name", ""), + "arguments": ( + json.dumps(tool_input) + if isinstance(tool_input, dict) + else str(tool_input) + ), + } + ) elif btype == "tool_result": # Convert to OpenAI function_call_output item @@ -690,11 +686,13 @@ def _convert_messages_with_tools(self, messages: list[dict]) -> list[dict]: output = result_content else: output = str(result_content) - codex_input.append({ - "type": "function_call_output", - "call_id": block.get("tool_use_id", ""), - "output": output, - }) + codex_input.append( + { + "type": "function_call_output", + "call_id": block.get("tool_use_id", ""), + "output": output, + } + ) elif btype == "image": # Convert internal base64 image to OpenAI input_image format @@ -702,10 +700,12 @@ def _convert_messages_with_tools(self, messages: list[dict]) -> list[dict]: if isinstance(source, dict) and source.get("type") == "base64": media_type = source.get("media_type", "image/png") data = source.get("data", "") - image_parts.append({ - "type": "input_image", - "image_url": f"data:{media_type};base64,{data}", - }) + image_parts.append( + { + "type": "input_image", + "image_url": f"data:{media_type};base64,{data}", + } + ) # Flush remaining text/image parts if text_parts or image_parts: @@ -715,13 +715,17 @@ def _convert_messages_with_tools(self, messages: list[dict]) -> list[dict]: msg_content.append({"type": ct, "text": " ".join(text_parts)}) msg_content.extend(image_parts) if msg_content: - codex_input.append({ - "type": "message", - "role": (role - if role in ("user", "assistant", "developer", "system") - else "user"), - "content": msg_content, - }) + codex_input.append( + { + "type": "message", + "role": ( + role + if role in ("user", "assistant", "developer", "system") + else "user" + ), + "content": msg_content, + } + ) return codex_input @@ -745,7 +749,10 @@ def _convert_tools_cached(self, tools: list[dict]) -> list[dict]: iteration. This avoids re-converting 70+ tool definitions each time. """ if tools is not self._last_tools_list: - self._last_tools_converted = self._convert_tools(tools) + from .strict_tool_adapter import compile_catalog + + self._tool_adapter = compile_catalog(tools) + self._last_tools_converted = self._tool_adapter.wire_tools self._last_tools_list = tools return self._last_tools_converted @@ -782,11 +789,14 @@ async def chat_with_tools( effort = reasoning_effort if reasoning_effort is not None else self.reasoning_effort resolved_model = model if model else self.model _reject_known_bad_pair(resolved_model, effort) + from .strict_tool_adapter import compile_catalog + + adapter = compile_catalog(tools) body = { "model": resolved_model, "instructions": system, "input": self._convert_messages_with_tools(messages), - "tools": self._convert_tools_cached(tools), + "tools": adapter.wire_tools, "tool_choice": "auto", "store": False, "stream": True, @@ -795,10 +805,14 @@ async def chat_with_tools( body["reasoning"] = {"effort": effort} input_tokens = self._estimate_body_input_tokens(body) - if progress_observer is None: - result = await self._stream_tool_request(body) - else: - result = await self._stream_tool_request(body, progress_observer=progress_observer) + token = _request_tool_adapter.set(adapter) + try: + if progress_observer is None: + result = await self._stream_tool_request(body) + else: + result = await self._stream_tool_request(body, progress_observer=progress_observer) + finally: + _request_tool_adapter.reset(token) output_chars = len(result.text) for tc in result.tool_calls: output_chars += len(tc.name) + len(json.dumps(tc.input)) @@ -833,13 +847,19 @@ def observe(event): try: return await self._read_tool_stream(resp, progress_observer=observe) except (CodexStreamError, TimeoutError, aiohttp.ClientError): - emit_progress(progress_observer, GenerationProgress( - "discarded", "codex", discarded_text_chars=text_chars, - discarded_tool_argument_chars=argument_chars, - )) + emit_progress( + progress_observer, + GenerationProgress( + "discarded", + "codex", + discarded_text_chars=text_chars, + discarded_tool_argument_chars=argument_chars, + ), + ) log.warning( "Codex discarded partial stream: text_chars=%d; tool_argument_chars=%d", - text_chars, argument_chars, + text_chars, + argument_chars, ) raise @@ -859,7 +879,11 @@ async def _stream_request(self, body: dict) -> str: ) async def _send_with_retries( - self, body: dict, reader, result_is_empty, *, + self, + body: dict, + reader, + result_is_empty, + *, progress_observer: GenerationProgressObserver | None = None, ): """Shared retry/rotation/breaker engine for both streaming paths. @@ -963,7 +987,10 @@ async def _send_with_retries( ) log.warning( "Codex stream failed (attempt %d/%d): %s. Retrying in %.1fs...", - attempt + 1, self.max_retries, last_error, wait, + attempt + 1, + self.max_retries, + last_error, + wait, ) await asyncio.sleep(wait) continue @@ -984,7 +1011,8 @@ async def _send_with_retries( return result log.warning( "Codex returned 200 with empty response (attempt %d/%d)", - attempt + 1, self.max_retries, + attempt + 1, + self.max_retries, ) if attempt < self.max_retries - 1: wait = compute_backoff( @@ -999,9 +1027,7 @@ async def _send_with_retries( raw_error, body_overflowed = await _read_error_body_bounded(resp.content) error_body = raw_error.decode("utf-8", errors="replace") - structured = ( - None if body_overflowed else _parse_structured_error(error_body) - ) + structured = None if body_overflowed else _parse_structured_error(error_body) descriptor = _describe_error_body( resp.headers.get("Content-Type"), raw_error, structured ) @@ -1017,9 +1043,13 @@ async def _send_with_retries( # actually exercise the refresh token — merely dropping # the cached token re-serves the same unexpired bearer — # then retry the SAME account once. - if attempt == 0 and not invalidated and await self._force_refresh( - acct_idx, - token, + if ( + attempt == 0 + and not invalidated + and await self._force_refresh( + acct_idx, + token, + ) ): log.warning("Codex auth 401, token refreshed, retrying...") token, account_id = await self._token_for(acct_idx) @@ -1052,7 +1082,10 @@ async def _send_with_retries( log.warning( "Codex rate limited (attempt %d/%d): %s. " "Rotating + retry in %.1fs...", - attempt + 1, self.max_retries, last_error, wait, + attempt + 1, + self.max_retries, + last_error, + wait, ) await asyncio.sleep(wait) token, account_id, acct_idx = await self._acquire_auth() @@ -1088,7 +1121,10 @@ async def _send_with_retries( ) log.warning( "Codex API error (attempt %d/%d): %s. Retrying in %.1fs...", - attempt + 1, self.max_retries, last_error, wait, + attempt + 1, + self.max_retries, + last_error, + wait, ) await asyncio.sleep(wait) continue @@ -1119,7 +1155,10 @@ async def _send_with_retries( wait = compute_backoff(attempt, self.retry_base_delay, self.retry_max_delay) log.warning( "Codex connection error (attempt %d/%d): %s. Retrying in %.1fs...", - attempt + 1, self.max_retries, last_error, wait, + attempt + 1, + self.max_retries, + last_error, + wait, ) await asyncio.sleep(wait) else: @@ -1159,6 +1198,31 @@ async def _read_tool_stream( # Track in-progress function calls by output_index pending_calls: dict[int, dict] = {} # {index: {"call_id": ..., "name": ..., "args": ""}} event_types_seen: list[str] = [] + adapter = _request_tool_adapter.get() + if adapter is not None: + adapter.record_resolution(None) + + def finish_call(call_id: str, name: str, raw_args: str) -> ToolCall: + try: + arguments = json.loads(raw_args) if raw_args else {} + except json.JSONDecodeError: + return ToolCall( + id=call_id, + name=name, + input={}, + parse_error="malformed tool arguments (invalid JSON)", + ) + if adapter is not None: + try: + arguments = adapter.accept(name, arguments) + except ValueError as exc: + return ToolCall( + id=call_id, + name=name, + input={}, + parse_error=f"invalid tool arguments: {exc}", + ) + return ToolCall(id=call_id, name=name, input=arguments) async for raw_line in resp.content: emit_progress(progress_observer, GenerationProgress("wire", "codex")) @@ -1178,8 +1242,11 @@ async def _read_tool_stream( event_type = event.get("type", "") event_types_seen.append(event_type) + if event_type == "response.created" and adapter is not None: + adapter.record_resolution(event) if ( - event_type in { + event_type + in { "response.reasoning_text.delta", "response.reasoning_summary_text.delta", } @@ -1208,9 +1275,7 @@ async def _read_tool_stream( elif event_type == "response.output_item.added": item = event.get("item", {}) if item.get("type") == "function_call": - emit_progress( - progress_observer, GenerationProgress("substantive", "codex") - ) + emit_progress(progress_observer, GenerationProgress("substantive", "codex")) idx = event.get("output_index", 0) pending_calls[idx] = { "call_id": item.get("call_id", ""), @@ -1227,7 +1292,8 @@ async def _read_tool_stream( emit_progress( progress_observer, GenerationProgress( - "substantive", "codex", + "substantive", + "codex", tool_argument_chars=len(event["delta"]), ), ) @@ -1237,23 +1303,14 @@ async def _read_tool_stream( idx = event.get("output_index", 0) if idx in pending_calls: call_info = pending_calls[idx] - parse_error = None - try: - parsed_args = json.loads(call_info["args"]) if call_info["args"] else {} - except json.JSONDecodeError: - parsed_args = {} - parse_error = (f"malformed tool arguments (invalid JSON): " - f"{call_info['args'][:200]}") - log.warning( - "Failed to parse function call arguments: %s", - call_info["args"][:200], + if not any(tc.id == call_info["call_id"] for tc in tool_calls): + tool_calls.append( + finish_call( + call_info["call_id"], + call_info["name"], + event.get("arguments", call_info["args"]), + ) ) - tool_calls.append(ToolCall( - id=call_info["call_id"], - name=call_info["name"], - input=parsed_args, - parse_error=parse_error, - )) # Output item done — finalize any remaining pending call at this index elif event_type == "response.output_item.done": @@ -1266,19 +1323,9 @@ async def _read_tool_stream( call_info = pending_calls.pop(idx, None) # type: ignore[arg-type] if call_info and not any(tc.id == call_info["call_id"] for tc in tool_calls): args_str = item.get("arguments", call_info.get("args", "")) - parse_error = None - try: - parsed_args = json.loads(args_str) if args_str else {} - except json.JSONDecodeError: - parsed_args = {} - parse_error = (f"malformed tool arguments (invalid JSON): " - f"{args_str[:200]}") - tool_calls.append(ToolCall( - id=call_info["call_id"], - name=call_info["name"], - input=parsed_args, - parse_error=parse_error, - )) + tool_calls.append( + finish_call(call_info["call_id"], call_info["name"], args_str) + ) # Terminal failure events: the HTTP 200 turned out to be a failed # generation — surface it so the retry engine treats it as an @@ -1293,8 +1340,9 @@ async def _read_tool_stream( elif event_type == "response.incomplete": terminal_received = True incomplete = True - reason = ((event.get("response") or {}).get("incomplete_details") - or {}).get("reason") or "unknown" + reason = ((event.get("response") or {}).get("incomplete_details") or {}).get( + "reason" + ) or "unknown" log.warning( "Codex stream incomplete (reason: %s) — returning partial output", reason, @@ -1324,19 +1372,7 @@ async def _read_tool_stream( call_id = item.get("call_id", "") if not any(tc.id == call_id for tc in tool_calls): args_str = item.get("arguments", "") - parse_error = None - try: - parsed_args = json.loads(args_str) if args_str else {} - except json.JSONDecodeError: - parsed_args = {} - parse_error = (f"malformed tool arguments (invalid JSON): " - f"{args_str[:200]}") - tool_calls.append(ToolCall( - id=call_id, - name=item.get("name", ""), - input=parsed_args, - parse_error=parse_error, - )) + tool_calls.append(finish_call(call_id, item.get("name", ""), args_str)) if not terminal_received: # Argument/item completion and [DONE] are not response acceptance. @@ -1349,8 +1385,11 @@ async def _read_tool_stream( ) text = "".join(text_parts) if not text and not tool_calls: - log.warning("Codex tool stream empty (events: %s, pending: %s)", - event_types_seen, list(pending_calls.keys())) + log.warning( + "Codex tool stream empty (events: %s, pending: %s)", + event_types_seen, + list(pending_calls.keys()), + ) if incomplete: stop_reason = "incomplete" @@ -1414,8 +1453,9 @@ async def _read_stream(self, resp: aiohttp.ClientResponse) -> str: elif event_type == "response.incomplete": terminal_received = True - reason = ((event.get("response") or {}).get("incomplete_details") - or {}).get("reason") or "unknown" + reason = ((event.get("response") or {}).get("incomplete_details") or {}).get( + "reason" + ) or "unknown" log.warning( "Codex stream incomplete (reason: %s) — returning partial output", reason, diff --git a/src/llm/strict_tool_adapter.py b/src/llm/strict_tool_adapter.py new file mode 100644 index 000000000..e76622b3a --- /dev/null +++ b/src/llm/strict_tool_adapter.py @@ -0,0 +1,381 @@ +"""Request-local Codex tool schema compiler and canonical acceptance boundary.""" + +from __future__ import annotations + +import hashlib +import json +import logging +from copy import deepcopy + +from jsonschema import Draft202012Validator + +log = logging.getLogger(__name__) +PAYLOADS = {"schedule_task", "update_schedule", "delegate_task", "invoke_skill"} +LOWERED = {"uniqueItems", "minProperties", "oneOf", "allOf", "not", "default"} +WIRE = { + "type", + "description", + "enum", + "const", + "properties", + "required", + "additionalProperties", + "items", + "anyOf", + "$defs", + "$ref", + "format", + "minimum", + "maximum", + "exclusiveMinimum", + "exclusiveMaximum", + "minLength", + "maxLength", + "pattern", + "minItems", + "maxItems", + "maxProperties", + "multipleOf", +} + + +def _nullable(node): + return {"anyOf": [node, {"type": "null"}]} + + +def _optional(node, original): + # A canonical nullable property must distinguish explicit null from wire + # omission. The sentinel is a disjoint closed object for non-object types, + # or a closed object with a reserved member for canonical closed objects. + if _allows_null(original): + return { + "anyOf": [ + node, + { + "type": "object", + "properties": {"__odin_omitted__": {"type": "boolean", "const": True}}, + "required": ["__odin_omitted__"], + "additionalProperties": False, + }, + ] + } + return _nullable(node) + + +def _compile(node, name, builtin, path=()): + if not isinstance(node, dict): + raise ValueError(f"{name}: unsupported schema at {path}") + unsupported = set(node) - WIRE - (LOWERED if builtin else set()) + if unsupported: + raise ValueError(f"{name}: unsupported keywords {sorted(unsupported)} at {path}") + if not name.startswith("computer_") and set(node) & (LOWERED - {"default"}): + raise ValueError( + f"{name}: unsupported constraint outside audited computer contract at {path}" + ) + if any(key in node for key in ("oneOf", "allOf", "not")) and name != "computer_act": + raise ValueError(f"{name}: unaudited combinator at {path}") + if "$ref" in node or "$defs" in node or isinstance(node.get("type"), list): + raise ValueError(f"{name}: reference or type union needs an envelope at {path}") + if ( + node.get("type") == "object" + and node.get("additionalProperties") is not None + and node.get("additionalProperties") is not False + ): + raise ValueError(f"{name}: open object cannot be closed at {path}") + if ( + node.get("type") == "object" + and not builtin + and node.get("additionalProperties") is not False + ): + raise ValueError(f"{name}: open external object at {path}") + if path in ( + {("tool_input",), ("steps", "*", "tool_input")} + if name in {"schedule_task", "update_schedule"} + else {("steps", "*", "tool_input")} + if name == "delegate_task" + else {("input",)} + if name == "invoke_skill" + else set() + ): + return { + "type": "string", + "description": ( + "JSON object encoded as text; target schema and authorization " + "are checked before execution." + ), + } + if "anyOf" in node: + if set(node) - {"anyOf", "description"}: + raise ValueError(f"{name}: anyOf siblings require explicit lowering at {path}") + result = {"anyOf": [_compile(branch, name, builtin, path) for branch in node["anyOf"]]} + if "description" in node: + result["description"] = node["description"] + return result + out = { + key: deepcopy(value) + for key, value in node.items() + if key + in ( + "type", + "description", + "enum", + "const", + "$ref", + "format", + "minimum", + "maximum", + "exclusiveMinimum", + "exclusiveMaximum", + "minLength", + "maxLength", + "pattern", + "minItems", + "maxItems", + "maxProperties", + "multipleOf", + ) + } + if node.get("type") == "object": + props = node.get("properties", {}) + required = set(node.get("required", ())) + out["properties"] = { + key: (compiled if key in required else _optional(compiled, child)) + for key, child in props.items() + for compiled in [_compile(child, name, builtin, path + (key,))] + } + out["required"] = list(props) + out["additionalProperties"] = False + if node.get("type") == "array": + out["items"] = _compile(node["items"], name, builtin, path + ("*",)) + return out + + +def _computer_branches(schema): + branches = [] + for case in schema["oneOf"]: + op = case["properties"]["operation"]["const"] + props = { + key: value + for key, value in schema["properties"].items() + if case["properties"].get(key) is not False + } + for key, override in case["properties"].items(): + if isinstance(override, dict) and "type" in override: + props[key] = override + props["operation"] = {"type": "string", "const": op} + branch = deepcopy(schema) + branch.pop("oneOf") + branch["properties"] = props + branch["required"] = sorted(set(schema["required"]) | set(case.get("required", ()))) + if "steps" in props: + branch["properties"]["steps"] = deepcopy(props["steps"]) + branch["properties"]["steps"]["items"] = { + "anyOf": _computer_branches(props["steps"]["items"]) + } + branches.append(_compile(branch, "computer_act", True)) + return branches + + +def _allows_null(schema): + return ( + schema is True + or isinstance(schema, dict) + and ( + schema.get("type") == "null" + or isinstance(schema.get("type"), list) + and "null" in schema["type"] + or any(_allows_null(branch) for branch in schema.get("anyOf", [])) + ) + ) + + +def _normalize(value, canonical, wire): + if isinstance(value, dict) and "anyOf" in wire: + matches = [ + idx + for idx, branch in enumerate(wire["anyOf"]) + if Draft202012Validator(branch).is_valid(value) + ] + if len(matches) != 1: + raise ValueError("ambiguous or invalid wire branch") + idx = matches[0] + if "oneOf" in canonical: + case = canonical["oneOf"][idx] + canonical = deepcopy(canonical) + canonical.pop("oneOf") + canonical["properties"] = { + key: val + for key, val in canonical["properties"].items() + if case["properties"].get(key) is not False + } + elif "anyOf" in canonical: + canonical = canonical["anyOf"][idx] + return _normalize(value, canonical, wire["anyOf"][idx]) + if isinstance(value, dict) and "properties" in wire: + result = {} + for key, item in value.items(): + if key not in canonical.get("properties", {}): + result[key] = item + continue + child = canonical["properties"][key] + child_wire = wire["properties"][key] + if key not in canonical.get("required", ()): + if (item is None and not _allows_null(child)) or ( + _allows_null(child) and item == {"__odin_omitted__": True} + ): + continue + child_wire = child_wire["anyOf"][0] + result[key] = _normalize(item, child, child_wire) + return result + if isinstance(value, list) and "items" in wire: + return [_normalize(item, canonical["items"], wire["items"]) for item in value] + return value + + +class RequestToolAdapter: + def __init__(self, tools): + from ..tools.defs.computer import COMPUTER_TOOL_NAMES + from ..tools.registry import TOOL_MAP + + self.catalog = deepcopy(tools) + self.wire_tools = [] + self.report = {} + self._contracts = {} + builtins = set(TOOL_MAP) | set(COMPUTER_TOOL_NAMES) + for tool in self.catalog: + name = tool["name"] + if name in self._contracts: + raise ValueError(f"duplicate tool {name}") + canonical = tool.get("input_schema", {"type": "object", "properties": {}}) + builtin = name in builtins or tool.get("is_core") is True + try: + if not builtin and canonical.get("type") != "object": + raise ValueError("external root is not an object") + wire = ( + { + "type": "object", + "properties": {"payload": {"anyOf": _computer_branches(canonical)}}, + "required": ["payload"], + "additionalProperties": False, + } + if name == "computer_act" + else _compile(canonical, name, builtin) + ) + mode, reason = ("builtin_strict" if builtin else "external_compiled"), None + except ValueError as exc: + if builtin: + raise + reason = str(exc) + if isinstance(canonical, dict) and canonical.get("type") == "object": + wire = { + "type": "object", + "properties": { + "json": { + "type": "string", + "description": "JSON object conforming to canonical input schema: " + + json.dumps(canonical, sort_keys=True, ensure_ascii=False), + } + }, + "required": ["json"], + "additionalProperties": False, + } + mode = "external_envelope" + else: + # A JSON-object envelope cannot preserve a non-object root. + wire = deepcopy(canonical) + mode = "external_exception" + fingerprint = hashlib.sha256(json.dumps(canonical, sort_keys=True).encode()).hexdigest() + self._contracts[name] = (canonical, wire, mode) + self.wire_tools.append( + { + "type": "function", + "name": name, + "description": tool.get("description", ""), + "parameters": wire, + **( + {"strict": True} + if builtin + else {"strict": False} + if mode == "external_exception" + else {} + ), + } + ) + self.report[name] = { + "mode": mode, + "fingerprint": fingerprint, + "resolution": "unknown", + "reason": reason, + } + + def accept(self, name, arguments): + if name not in self._contracts: + raise ValueError(f"{name}: not present in request catalog") + canonical, wire, mode = self._contracts[name] + if mode == "external_exception": + return arguments # explicitly unchanged, named non-strict residual exposure + if not isinstance(arguments, dict) or not Draft202012Validator(wire).is_valid(arguments): + raise ValueError(f"{name}: invalid wire arguments") + if mode == "external_envelope": + from ..tools.nested_payload import decode_json_object + + value = decode_json_object(arguments["json"], name) + elif name == "computer_act": + value = _normalize(arguments["payload"], canonical, wire["properties"]["payload"]) + else: + value = _normalize(arguments, canonical, wire) + checked = deepcopy(canonical) + if mode == "builtin_strict": + checked["additionalProperties"] = False + if mode == "builtin_strict" and name in PAYLOADS: + from ..tools.nested_payload import decode_nested_payloads + + value = decode_nested_payloads(name, value) + errors = list(Draft202012Validator(checked).iter_errors(value)) + if errors: + error = min(errors, key=lambda e: (len(e.path), e.message)) + raise ValueError( + f"{name}: invalid {'/'.join(map(str, error.path)) or 'root'}: " + f"violates {error.validator} constraint" + ) + if name == "http_probe" and isinstance(value.get("headers"), list): + # Lower wire records only after their wire/canonical shape check; + # before host policy or curl. Dict callers never enter this path. + from ..tools.http_probe_ops import normalize_probe_headers + + value["headers"] = normalize_probe_headers(value["headers"]) + if name in PAYLOADS and mode == "builtin_strict": + from ..tools.nested_payload import validate_nested_payload + + value = validate_nested_payload(name, value, self.catalog, allow_placeholders=True) + return value + + def record_resolution(self, event=None): + response = ( + event.get("response", {}) + if isinstance(event, dict) and event.get("type") == "response.created" + else {} + ) + resolved = { + item.get("name"): item.get("strict") + for item in response.get("tools", []) + if isinstance(item, dict) + } + result = {} + for name, item in self.report.items(): + val = resolved.get(name) + item["resolution"] = "true" if val is True else "false" if val is False else "unknown" + result[name] = item["resolution"] + log.info( + "Codex schema name=%s fingerprint=%s mode=%s resolution=%s reason=%s", + name, + item["fingerprint"], + item["mode"], + item["resolution"], + item["reason"], + ) + return result + + +def compile_catalog(tools): + return RequestToolAdapter(tools) diff --git a/src/scheduler/scheduler.py b/src/scheduler/scheduler.py index 085b43afc..f8c1d8e5f 100644 --- a/src/scheduler/scheduler.py +++ b/src/scheduler/scheduler.py @@ -528,6 +528,7 @@ async def add( requester_id: str = "", cron_timezone: str | None = None, report_format: str | None = None, + nested_payload_validated: bool = False, ) -> dict: self._validate_report_format(report_format, action) if action == "digest": @@ -568,6 +569,7 @@ async def add( "action": action, "channel_id": channel_id, "requester_id": requester_id, + "_nested_payload_validated": bool(nested_payload_validated), "created_at": datetime.now(UTC).isoformat(), "last_run": None, } @@ -941,6 +943,7 @@ async def update( paused: bool | None = None, cron_timezone: str | None = None, report_format: str | None = None, + nested_payload_validated: bool = False, ) -> dict | None: """Update mutable fields on an existing schedule. @@ -969,6 +972,8 @@ async def update( # persist them. Commit only after every supplied field is valid. original = self._schedules[target_index] target = copy.deepcopy(original) + if nested_payload_validated: + target["_nested_payload_validated"] = True action = target["action"] if steps is not None and action == "workflow": self._validate_workflow_steps(steps) diff --git a/src/tools/defs/integrations_email.py b/src/tools/defs/integrations_email.py index b8b102c31..df9756bbb 100644 --- a/src/tools/defs/integrations_email.py +++ b/src/tools/defs/integrations_email.py @@ -42,10 +42,13 @@ ], }, "headers": { - "type": "object", - "description": ( - 'Request headers as key-value pairs (e.g. {"Authorization": "Bearer tok"})' - ), + "type": "array", + "description": "Request headers as name/value entries. Case-insensitive duplicate names are rejected.", + "items": { + "type": "object", + "properties": {"name": {"type": "string"}, "value": {"type": "string"}}, + "required": ["name", "value"], + }, }, "body": { "type": "string", @@ -178,7 +181,12 @@ }, "target": {"type": "string"}, "expected": { - "description": "Type-specific expectation (int, string, list)" + "description": "Type-specific expectation: integer, string, integer list, or string list", + "anyOf": [ + {"type": "integer"}, {"type": "string"}, + {"type": "array", "items": {"type": "integer"}}, + {"type": "array", "items": {"type": "string"}}, + ], }, "severity": { "type": "string", diff --git a/src/tools/defs/media_scheduling.py b/src/tools/defs/media_scheduling.py index 03a872f72..2bd9ebea4 100644 --- a/src/tools/defs/media_scheduling.py +++ b/src/tools/defs/media_scheduling.py @@ -148,6 +148,7 @@ "description": "Grafana alert name substring (case-insensitive)", }, }, + "additionalProperties": False, }, "action": { "type": "string", @@ -283,6 +284,13 @@ "trigger": { "type": "object", "description": "New webhook trigger (replaces previous timing)", + "properties": { + "source": {"type": "string", "enum": ["gitea", "grafana", "generic", "github", "gitlab"]}, + "event": {"type": "string"}, + "repo": {"type": "string"}, + "alert_name": {"type": "string"}, + }, + "additionalProperties": False, }, "message": { "type": "string", diff --git a/src/tools/handlers/browser_web.py b/src/tools/handlers/browser_web.py index dedfd684f..67d377d35 100644 --- a/src/tools/handlers/browser_web.py +++ b/src/tools/handlers/browser_web.py @@ -85,7 +85,15 @@ async def _handle_fetch_url(self, inp: dict) -> str: return await fetch_url(inp["url"]) async def _handle_http_probe(self, inp: dict) -> str | tuple[str, int]: - from ..http_probe_ops import build_http_probe_command + from ..http_probe_ops import build_http_probe_command, normalize_probe_headers + + try: + # Decode strict wire headers before host authorization and retain + # the canonical dict for the existing execution interface. + if "headers" in inp: + inp = {**inp, "headers": normalize_probe_headers(inp["headers"])} + except ValueError as e: + return f"http_probe error: {e}", 1 host = inp.get("host", "") if host: diff --git a/src/tools/http_probe_ops.py b/src/tools/http_probe_ops.py index 6f0cda4ed..d36de92c4 100644 --- a/src/tools/http_probe_ops.py +++ b/src/tools/http_probe_ops.py @@ -95,6 +95,30 @@ def _clamp_int(value, default: int, minimum: int, maximum: int) -> int: return max(minimum, min(v, maximum)) +def normalize_probe_headers(headers): + """Decode the strict wire header-record form, preserving legacy dict callers.""" + if headers is None or isinstance(headers, dict): + return headers + if not isinstance(headers, list): + # Preserve the legacy builder contract: unsupported direct-call types + # were ignored. The strict request boundary rejects them by schema. + return headers + normalized = {} + seen = set() + for i, entry in enumerate(headers): + if not isinstance(entry, dict) or set(entry) != {"name", "value"}: + raise ValueError(f"Header entry {i} must contain exactly name and value") + name, value = entry["name"], entry["value"] + if not isinstance(name, str) or not isinstance(value, str): + raise ValueError(f"Header entry {i} name and value must be strings") + folded = name.casefold() + if folded in seen: + raise ValueError(f"Duplicate HTTP header name (case-insensitive): {name}") + seen.add(folded) + normalized[name] = value + return normalized + + def build_http_probe_command(params: dict) -> str: """Build a curl command for HTTP probing. @@ -168,8 +192,9 @@ def build_http_probe_command(params: dict) -> str: retry_delay = _clamp_int(params.get("retry_delay"), DEFAULT_RETRY_DELAY, 0, MAX_RETRY_DELAY) parts.append(f"--retry-delay {retry_delay}") - # Custom headers - headers = params.get("headers") + # Custom headers. The wire uses name/value records because strict JSON + # schemas cannot express arbitrary object keys. Keep canonical dict callers. + headers = normalize_probe_headers(params.get("headers")) if isinstance(headers, dict): for name, value in headers.items(): # curl interprets -H @file as a request to read local file contents. diff --git a/src/tools/nested_payload.py b/src/tools/nested_payload.py new file mode 100644 index 000000000..3faf85cd7 --- /dev/null +++ b/src/tools/nested_payload.py @@ -0,0 +1,123 @@ +"""Canonical decoding and validation for tool inputs nested as JSON strings on strict wire calls.""" +from __future__ import annotations + +import json +import re +from typing import Any + +from jsonschema.validators import validator_for + +_NESTED_FIELDS = { + "schedule_task": ("tool_input", "steps[].tool_input"), + "update_schedule": ("tool_input", "steps[].tool_input"), + "delegate_task": ("steps[].tool_input",), + "invoke_skill": ("input",), +} +_PLACEHOLDER = re.compile(r"\{(?:var\.[^{}]+|prev_output)\}") + + +def _object_no_duplicates(pairs): + result = {} + for key, value in pairs: + if key in result: + raise ValueError(f"duplicate JSON key in nested payload: {key!r}") + result[key] = value + return result + + +def decode_json_object(value: Any, label: str) -> dict: + """Decode one wire JSON string or accept an already-canonical dict.""" + if isinstance(value, dict): + return value + if not isinstance(value, str): + raise ValueError(f"{label} must be a JSON object string or object") + try: + result = json.loads(value, object_pairs_hook=_object_no_duplicates, + parse_constant=lambda token: (_ for _ in ()).throw(ValueError(f"invalid JSON constant {token}"))) + except (json.JSONDecodeError, ValueError) as exc: + raise ValueError(f"{label} must contain valid JSON without duplicate keys: {exc}") from exc + if not isinstance(result, dict): + raise ValueError(f"{label} must decode to a JSON object") + return result + + +def decode_nested_payloads(tool_name: str, arguments: dict) -> dict: + """Return copied arguments with nested wire JSON strings decoded exactly once.""" + if tool_name not in _NESTED_FIELDS: + return arguments + result = dict(arguments) + if tool_name in ("schedule_task", "update_schedule") and "tool_input" in result and isinstance(result["tool_input"], str): + result["tool_input"] = decode_json_object(result["tool_input"], "tool_input") + if tool_name == "invoke_skill" and "input" in result and isinstance(result["input"], str): + result["input"] = decode_json_object(result["input"], "input") + if isinstance(result.get("steps"), list): + steps = [] + for index, original in enumerate(result["steps"], 1): + if not isinstance(original, dict): + steps.append(original); continue + step = dict(original) + if isinstance(step.get("tool_input"), str): + step["tool_input"] = decode_json_object(step["tool_input"], f"step {index} tool_input") + steps.append(step) + result["steps"] = steps + return result + + +def _schema_for(name: str, catalog: list[dict]) -> dict | None: + return next((t.get("input_schema") for t in catalog if t.get("name") == name), None) + + +def _validate(schema: dict, payload: dict, *, allow_placeholders: bool) -> None: + validator_cls = validator_for(schema) + validator_cls.check_schema(schema) + validator = validator_cls(schema) + import copy + candidate = copy.deepcopy(payload) + deferred: set[tuple] = set() + if allow_placeholders: + # The established workflow substitution contract touches direct string + # values in tool_input only, not nested objects or array members. + if isinstance(candidate, dict): + schema = copy.deepcopy(schema) + for key, value in candidate.items(): + if isinstance(value, str) and _PLACEHOLDER.search(value): + schema.setdefault("properties", {})[key] = {} + errors = sorted(validator_cls(schema).iter_errors(candidate), key=lambda e: list(map(str, e.path))) + if errors: + err = errors[0] + path = ".".join(map(str, err.absolute_path)) or "" + raise ValueError(f"invalid input for selected tool at {path}: {err.message}") + + +def validate_nested_payload(tool_name: str, canonical_arguments: dict, catalog: list[dict], *, allow_placeholders: bool = True) -> dict: + """Decode nested JSON fields, validate selected canonical targets, return canonical args. + + This validates payload shape only. Callers must also apply their existing authorization + checks to the selected target; a name in a payload is not permission. + """ + from .registry import TOOLS + args = decode_nested_payloads(tool_name, canonical_arguments) + def check(target, payload, label): + schema = _schema_for(target, catalog) or _schema_for(target, TOOLS) + if schema is None: + raise ValueError(f"{label} selects unknown tool {target!r}") + if not isinstance(payload, dict): + raise ValueError(f"{label} must be an object") + _validate(schema, payload, allow_placeholders=allow_placeholders) + if tool_name in ("schedule_task", "update_schedule"): + if args.get("tool_name") and isinstance(args.get("tool_input"), dict) and args.get("action", "check") == "check": + check(args["tool_name"], args["tool_input"], "tool_input") + if tool_name == "update_schedule" and isinstance(args.get("tool_input"), dict) and args.get("tool_name"): + check(args["tool_name"], args["tool_input"], "tool_input") + for i, step in enumerate(args.get("steps") or [], 1): + if isinstance(step, dict) and step.get("tool_name") and isinstance(step.get("tool_input"), dict): + check(step["tool_name"], step["tool_input"], f"step {i} tool_input") + elif tool_name == "delegate_task": + for i, step in enumerate(args.get("steps") or [], 1): + if isinstance(step, dict) and step.get("tool_name") and isinstance(step.get("tool_input"), dict): + check(step["tool_name"], step["tool_input"], f"step {i} tool_input") + elif tool_name == "invoke_skill": + target = args.get("name") + if target and isinstance(args.get("input"), dict): + check(target, args["input"], "input") + return args diff --git a/src/tools/post_validation.py b/src/tools/post_validation.py index 07e4c2523..35c29ad19 100644 --- a/src/tools/post_validation.py +++ b/src/tools/post_validation.py @@ -211,6 +211,28 @@ def parse_checks(raw_checks: list[dict]) -> tuple[list[Check], list[str]]: ) continue expected = raw.get("expected") + if "expected" in raw: + # Keep the public union deliberately small, then enforce the + # check-specific interpretation before any check is launched. + valid_scalar = isinstance(expected, (int, str)) and not isinstance(expected, bool) + valid_list = isinstance(expected, list) and bool(expected) and all( + isinstance(item, (int, str)) and not isinstance(item, bool) for item in expected + ) + if not (valid_scalar or valid_list): + errors.append(f"check[{i}]: expected must be an integer, string, or non-empty list of integers/strings") + continue + if c_type == "http": + values = expected if isinstance(expected, list) else [expected] + if not all(isinstance(v, int) or (isinstance(v, str) and v.isdigit()) for v in values): + errors.append(f"check[{i}]: http expected values must be status-code integers or digit strings") + continue + elif c_type == "service": + if isinstance(expected, list) and not all(isinstance(v, str) for v in expected): + errors.append(f"check[{i}]: service expected lists must contain strings") + continue + elif c_type in {"port", "process"} and expected is not None: + errors.append(f"check[{i}]: expected is not used for {c_type} checks") + continue # Require 'expected' for compare ops that need a value — but only when # the caller explicitly chose the compare op. Defaults have built-in # fallback expectations (e.g. http default = 2xx/3xx, service default diff --git a/tests/test_openai_codex_client.py b/tests/test_openai_codex_client.py index 350d0c216..7d3cab7a5 100644 --- a/tests/test_openai_codex_client.py +++ b/tests/test_openai_codex_client.py @@ -155,12 +155,13 @@ class TestToolsAndEstimation: def test_convert_tools_format(self): out = CodexChatClient._convert_tools([ {"name": "grep", "description": "search", "input_schema": {"type": "object"}}]) - assert out[0] == {"type": "function", "name": "grep", "description": "search", - "parameters": {"type": "object"}, "strict": False} + assert out[0]["name"] == "grep" + assert "strict" not in out[0] # external schema: server resolves omitted strict + assert set(out[0]["parameters"]["properties"]) == {"json"} def test_convert_tools_defaults(self): out = CodexChatClient._convert_tools([{"name": "bare"}]) - assert out[0]["parameters"] == {"type": "object", "properties": {}} + assert set(out[0]["parameters"]["properties"]) == {"json"} def test_convert_tools_cached_identity(self): c = _client() diff --git a/tests/test_strict_nested_payload.py b/tests/test_strict_nested_payload.py new file mode 100644 index 000000000..1a3cac685 --- /dev/null +++ b/tests/test_strict_nested_payload.py @@ -0,0 +1,56 @@ +"""Strict wire nested JSON payload decoding and selected-target validation.""" +import json + +import pytest + +from src.tools.nested_payload import decode_nested_payloads, validate_nested_payload + + +def catalog(): + return [ + {"name": "run_command", "input_schema": { + "type": "object", "properties": {"command": {"type": "string"}, "host": {"type": ["string", "null"]}}, + "required": ["command"], "additionalProperties": False, + }}, + {"name": "example_skill", "input_schema": { + "type": "object", "properties": {"value": {"type": ["string", "null"]}}, + "required": ["value"], "additionalProperties": False, + }}, + ] + + +def test_wire_json_decodes_canonical_payload_and_preserves_meaningful_null(): + raw = {'action': 'check', 'tool_name': 'run_command', + 'tool_input': json.dumps({'command': 'printf "\\u03bb"', 'host': None})} + result = validate_nested_payload('schedule_task', raw, catalog()) + assert result['tool_input'] == {'command': 'printf "\\u03bb"', 'host': None} + assert isinstance(raw['tool_input'], str) + + +@pytest.mark.parametrize('payload', ['{"command":"a","command":"b"}', '[]', 'null', '{bad']) +def test_rejects_bad_wire_nested_json(payload): + with pytest.raises(ValueError): + decode_nested_payloads('delegate_task', {'steps': [{'tool_name': 'run_command', 'tool_input': payload}]}) + + +def test_selected_target_schema_is_enforced(): + with pytest.raises(ValueError, match='invalid input'): + validate_nested_payload('delegate_task', { + 'steps': [{'tool_name': 'run_command', 'tool_input': {'command': 'ok', 'surprise': True}}] + }, catalog()) + + +def test_template_field_can_defer_validation_and_null_is_not_dropped(): + item = {'tool_name': 'run_command', 'tool_input': {'command': '{prev_output}', 'host': None}} + result = validate_nested_payload('delegate_task', {'steps': [item]}, catalog()) + assert result['steps'][0]['tool_input'] == item['tool_input'] + invalid = {'tool_name': 'run_command', 'tool_input': {'command': None, 'host': None}} + with pytest.raises(ValueError): + validate_nested_payload('delegate_task', {'steps': [invalid]}, catalog(), allow_placeholders=False) + + +def test_invoke_skill_validates_named_skill_schema(): + result = validate_nested_payload('invoke_skill', {'name': 'example_skill', 'input': '{"value":null}'}, catalog()) + assert result['input'] == {'value': None} + with pytest.raises(ValueError): + validate_nested_payload('invoke_skill', {'name': 'example_skill', 'input': {'wrong': 1}}, catalog()) diff --git a/tests/test_strict_tool_adapter.py b/tests/test_strict_tool_adapter.py new file mode 100644 index 000000000..a8b215f3c --- /dev/null +++ b/tests/test_strict_tool_adapter.py @@ -0,0 +1,120 @@ +"""Strict adapter fixtures, with no endpoint or desktop input.""" + +import copy +import json + +import pytest + +from src.llm.strict_tool_adapter import compile_catalog +from src.tools.defs.computer import computer_definitions +from src.tools.registry import get_tool_definitions + + +def _wire_args(adapter, name, specified): + """Populate omitted wire optionals using their request-local null branch.""" + props = next( + wire["parameters"]["properties"] for wire in adapter.wire_tools if wire["name"] == name + ) + result = dict(specified) + for key in props: + result.setdefault(key, None) + if name == "validate_action": + schema = props["checks"]["items"]["properties"] + for check in result["checks"]: + for key in schema: + check.setdefault(key, None) + return result + + +def test_catalog_strict_and_request_local(): + catalog = get_tool_definitions() + computer_definitions() + original = copy.deepcopy(catalog) + adapter = compile_catalog(catalog) + assert catalog == original + assert len(adapter.wire_tools) == len(catalog) + assert all(wire["strict"] is True for wire in adapter.wire_tools) + computer = adapter.wire_tools[-1]["parameters"] + assert "anyOf" not in computer + assert len(computer["properties"]["payload"]["anyOf"]) == 14 + + +def test_forced_values_and_nested_omission(): + adapter = compile_catalog(get_tool_definitions()) + assert adapter.accept( + "update_schedule", + _wire_args( + adapter, "update_schedule", {"schedule_id": "s", "paused": False, "report_format": ""} + ), + ) == {"schedule_id": "s", "paused": False, "report_format": ""} + assert adapter.accept( + "browser_read_table", + _wire_args( + adapter, + "browser_read_table", + {"url": "https://example.com", "table_index": 0, "wait_seconds": None}, + ), + ) == {"url": "https://example.com", "table_index": 0} + checks = { + "checks": [ + { + "type": "http", + "target": "https://example.com", + "expected": [200, 204], + "severity": None, + } + ] + } + assert adapter.accept("validate_action", _wire_args(adapter, "validate_action", checks)) == { + "checks": [{"type": "http", "target": "https://example.com", "expected": [200, 204]}] + } + + +def test_external_modes_round_trip_and_resolution(): + adapter = compile_catalog( + [ + { + "name": "closed_ext", + "input_schema": { + "type": "object", + "properties": {"text": {"type": "string"}}, + "additionalProperties": False, + }, + }, + { + "name": "open_ext", + "input_schema": {"type": "object", "properties": {"text": {"type": "string"}}}, + }, + {"name": "odd_ext", "input_schema": {"type": "string"}}, + ] + ) + assert [r["mode"] for r in adapter.report.values()] == [ + "external_compiled", + "external_envelope", + "external_exception", + ] + assert "strict" not in adapter.wire_tools[0] + assert "strict" not in adapter.wire_tools[1] + assert adapter.wire_tools[2]["strict"] is False + assert adapter.accept("open_ext", {"json": json.dumps({"text": "a", "more": None})}) == { + "text": "a", + "more": None, + } + assert ( + "canonical input schema" + in adapter.wire_tools[1]["parameters"]["properties"]["json"]["description"] + ) + for payload in ('{"a":1,"a":2}', "[1,2]", "not-json"): + with pytest.raises(ValueError): + adapter.accept("open_ext", {"json": payload}) + assert set(adapter.record_resolution(None).values()) == {"unknown"} + assert adapter.record_resolution( + { + "type": "response.created", + "response": { + "tools": [ + {"name": "closed_ext", "strict": True}, + {"name": "open_ext", "strict": False}, + ] + }, + } + ) == {"closed_ext": "true", "open_ext": "false", "odd_ext": "unknown"} diff --git a/tests/test_strict_wire_defs.py b/tests/test_strict_wire_defs.py new file mode 100644 index 000000000..a0bb291a3 --- /dev/null +++ b/tests/test_strict_wire_defs.py @@ -0,0 +1,68 @@ +"""Focused tests for strict wire definitions and their runtime lowerings.""" +from __future__ import annotations + +import pytest + +from src.tools.defs.integrations_email import TOOLS_SECTION as INTEGRATION_TOOLS +from src.tools.defs.media_scheduling import TOOLS_SECTION as SCHEDULING_TOOLS +from src.tools.http_probe_ops import build_http_probe_command +from src.tools.post_validation import parse_checks + + +def _schema(tools, name): + return next(tool["input_schema"] for tool in tools if tool["name"] == name) + + +def test_http_probe_headers_are_name_value_records_and_legacy_dict_still_works(): + headers = _schema(INTEGRATION_TOOLS, "http_probe")["properties"]["headers"] + assert headers["type"] == "array" + assert headers["items"]["required"] == ["name", "value"] + for shape in ([{"name": "X-Test", "value": "one"}], {"X-Test": "one"}): + cmd = build_http_probe_command({"url": "https://example.com", "headers": shape}) + assert "X-Test: one" in cmd + + +@pytest.mark.parametrize("headers", [ + [{"name": "Accept", "value": "a"}, {"name": "accept", "value": "b"}], + [{"name": "X", "value": "a", "extra": "no"}], + [{"name": "X", "value": 1}], +]) +def test_http_probe_wire_headers_reject_duplicates_and_bad_entries(headers): + with pytest.raises(ValueError): + build_http_probe_command({"url": "https://example.com", "headers": headers}) + + +def test_schedule_triggers_share_closed_four_field_shape(): + for tool_name in ("schedule_task", "update_schedule"): + trigger = _schema(SCHEDULING_TOOLS, tool_name)["properties"]["trigger"] + assert trigger["additionalProperties"] is False + assert set(trigger["properties"]) == {"source", "event", "repo", "alert_name"} + assert trigger["properties"]["source"]["enum"] == ["gitea", "grafana", "generic", "github", "gitlab"] + + +def test_trigger_runtime_keeps_partial_conditions_and_rejects_empty(): + from src.scheduler.scheduler import Scheduler + + Scheduler._validate_trigger({"event": "push"}) + with pytest.raises(ValueError, match="at least one condition"): + Scheduler._validate_trigger({}) + with pytest.raises(ValueError, match="Unknown trigger keys"): + Scheduler._validate_trigger({"event": "push", "extra": "not allowed"}) + + +def test_validate_action_expected_union_and_runtime_typed_values(): + expected = _schema(INTEGRATION_TOOLS, "validate_action")["properties"]["checks"]["items"]["properties"]["expected"] + assert [x["type"] for x in expected["anyOf"]] == ["integer", "string", "array", "array"] + checks, errors = parse_checks([{"type": "http", "target": "https://example.com", "expected": [200, 204]}]) + assert not errors and checks[0].expected == [200, 204] + for raw in ( + {"type": "http", "target": "https://example.com", "expected": [200, "up"]}, + {"type": "service", "target": "sshd", "expected": [200]}, + ): + checks, errors = parse_checks([raw]) + assert checks == [] and errors + + +def test_legacy_validation_expectations_keep_stringification_paths(): + checks, errors = parse_checks([{"type": "command", "target": "printf 12", "compare": "equals", "expected": 12}]) + assert not errors and checks[0].expected == 12 From 5f88962c76bcfa11f9749d7fcdb764615028c66f Mon Sep 17 00:00:00 2001 From: Odin Date: Sun, 27 Sep 2026 20:19:53 -0400 Subject: [PATCH 02/12] Correct strict adapter computer branches and completion semantics --- docs/strict-wire-lowerings.md | 12 +-- src/llm/openai_codex.py | 6 +- src/llm/strict_tool_adapter.py | 53 ++++++++-- tests/test_openai_codex_client.py | 84 ++++++++++++++++ tests/test_strict_tool_adapter.py | 156 +++++++++++++++++++++++++++++- 5 files changed, 291 insertions(+), 20 deletions(-) diff --git a/docs/strict-wire-lowerings.md b/docs/strict-wire-lowerings.md index 7b847e148..499c6307b 100644 --- a/docs/strict-wire-lowerings.md +++ b/docs/strict-wire-lowerings.md @@ -60,9 +60,9 @@ tripwire before dispatch; these are not generic recursive keyword deletions. | `oneOf` | Represented by nested operation branches only when match is unambiguous | Canonical validator and ambiguity check reject zero/multiple distinct matches | | `false` property schemas | Property omitted from wire branch | Canonical validator rejects forbidden operation fields | -Expected fixture coverage includes duplicate modifiers, repeated key -modifiers, empty task context, coordinate-versus-region conflicts, forbidden -operation fields, and nested sequence/stroke contracts. Fixtures and mocked -dispatch only: no desktop input is performed. The corresponding adapter test -module should assert positive and negative fixtures against compiled wire form -and canonical pre-effect validation. +`tests/test_strict_tool_adapter.py` exercises positive focus, click-coordinate, +click-region, sequence and stroke forms; negative focus expectation, +coordinate/region conflicts, forbidden operation fields, duplicate modifiers, +repeated key modifiers, empty task context, malformed nested sequence and +duplicate modifiers in nested strokes. These are local schema fixtures only: +no desktop input is performed. diff --git a/src/llm/openai_codex.py b/src/llm/openai_codex.py index 7139055d0..3dc5f00df 100644 --- a/src/llm/openai_codex.py +++ b/src/llm/openai_codex.py @@ -1199,8 +1199,7 @@ async def _read_tool_stream( pending_calls: dict[int, dict] = {} # {index: {"call_id": ..., "name": ..., "args": ""}} event_types_seen: list[str] = [] adapter = _request_tool_adapter.get() - if adapter is not None: - adapter.record_resolution(None) + resolution_seen = False def finish_call(call_id: str, name: str, raw_args: str) -> ToolCall: try: @@ -1244,6 +1243,7 @@ def finish_call(call_id: str, name: str, raw_args: str) -> ToolCall: event_types_seen.append(event_type) if event_type == "response.created" and adapter is not None: adapter.record_resolution(event) + resolution_seen = True if ( event_type in { @@ -1374,6 +1374,8 @@ def finish_call(call_id: str, name: str, raw_args: str) -> ToolCall: args_str = item.get("arguments", "") tool_calls.append(finish_call(call_id, item.get("name", ""), args_str)) + if adapter is not None and not resolution_seen: + adapter.record_resolution(None) if not terminal_received: # Argument/item completion and [DONE] are not response acceptance. # Nothing has escaped this reader or executed; the existing transport diff --git a/src/llm/strict_tool_adapter.py b/src/llm/strict_tool_adapter.py index e76622b3a..a1ce73039 100644 --- a/src/llm/strict_tool_adapter.py +++ b/src/llm/strict_tool_adapter.py @@ -160,8 +160,18 @@ def _computer_branches(schema): if case["properties"].get(key) is not False } for key, override in case["properties"].items(): - if isinstance(override, dict) and "type" in override: - props[key] = override + if isinstance(override, dict) and key != "operation": + # Case constraints can refine a shared nested object (notably + # computer_act.expect.type). Retain its base type and properties + # while replacing only the overridden constraints. + refined = deepcopy(props.get(key, {})) + refined.update({k: v for k, v in override.items() if k != "properties"}) + if "properties" in override: + refined["properties"] = { + **refined.get("properties", {}), + **override["properties"], + } + props[key] = refined props["operation"] = {"type": "string", "const": op} branch = deepcopy(schema) branch.pop("oneOf") @@ -208,6 +218,19 @@ def _normalize(value, canonical, wire): for key, val in canonical["properties"].items() if case["properties"].get(key) is not False } + canonical["required"] = list( + set(canonical.get("required", ())) | set(case.get("required", ())) + ) + for key, override in case["properties"].items(): + if key in canonical["properties"] and isinstance(override, dict): + merged = deepcopy(canonical["properties"][key]) + merged.update({k: v for k, v in override.items() if k != "properties"}) + if "properties" in override: + merged["properties"] = { + **merged.get("properties", {}), + **override["properties"], + } + canonical["properties"][key] = merged elif "anyOf" in canonical: canonical = canonical["anyOf"][idx] return _normalize(value, canonical, wire["anyOf"][idx]) @@ -241,6 +264,7 @@ def __init__(self, tools): self.wire_tools = [] self.report = {} self._contracts = {} + self._resolution_logged = False builtins = set(TOOL_MAP) | set(COMPUTER_TOOL_NAMES) for tool in self.catalog: name = tool["name"] @@ -331,6 +355,11 @@ def accept(self, name, arguments): from ..tools.nested_payload import decode_nested_payloads value = decode_nested_payloads(name, value) + headers = name == "http_probe" and isinstance(value.get("headers"), list) + if headers and checked["properties"]["headers"].get("type") == "object": + from ..tools.http_probe_ops import normalize_probe_headers + + value["headers"] = normalize_probe_headers(value["headers"]) errors = list(Draft202012Validator(checked).iter_errors(value)) if errors: error = min(errors, key=lambda e: (len(e.path), e.message)) @@ -338,9 +367,7 @@ def accept(self, name, arguments): f"{name}: invalid {'/'.join(map(str, error.path)) or 'root'}: " f"violates {error.validator} constraint" ) - if name == "http_probe" and isinstance(value.get("headers"), list): - # Lower wire records only after their wire/canonical shape check; - # before host policy or curl. Dict callers never enter this path. + if headers and isinstance(value["headers"], list): from ..tools.http_probe_ops import normalize_probe_headers value["headers"] = normalize_probe_headers(value["headers"]) @@ -351,16 +378,21 @@ def accept(self, name, arguments): return value def record_resolution(self, event=None): + if self._resolution_logged: + return {name: item["resolution"] for name, item in self.report.items()} response = ( event.get("response", {}) if isinstance(event, dict) and event.get("type") == "response.created" else {} ) - resolved = { - item.get("name"): item.get("strict") - for item in response.get("tools", []) - if isinstance(item, dict) - } + resolved = {} + for tool in response.get("tools", []): + if not isinstance(tool, dict): + continue + name = tool.get("name") + expected = next((t for t in self.wire_tools if t["name"] == name), None) + if expected is not None and tool.get("parameters") == expected["parameters"]: + resolved[name] = tool.get("strict") result = {} for name, item in self.report.items(): val = resolved.get(name) @@ -374,6 +406,7 @@ def record_resolution(self, event=None): item["resolution"], item["reason"], ) + self._resolution_logged = True return result diff --git a/tests/test_openai_codex_client.py b/tests/test_openai_codex_client.py index 7d3cab7a5..b46f84aba 100644 --- a/tests/test_openai_codex_client.py +++ b/tests/test_openai_codex_client.py @@ -271,6 +271,90 @@ async def test_empty_stream_returns_empty(self): class TestReadToolStream: + @pytest.mark.asyncio + @pytest.mark.parametrize("path", ["arguments_done", "item_done", "completed", "duplicates"]) + async def test_adapter_accepts_once_on_each_finalization_path(self, path): + from src.llm.openai_codex import _request_tool_adapter + from src.llm.strict_tool_adapter import compile_catalog + from src.tools.registry import get_tool_definitions + + adapter = compile_catalog(get_tool_definitions()) + original = adapter.accept + accepted = [] + + def counting_accept(name, args): + accepted.append(name) + return original(name, args) + + adapter.accept = counting_accept + schema = next(t["parameters"]["properties"] for t in adapter.wire_tools + if t["name"] == "browser_read_table") + args = {key: None for key in schema} | { + "url": "https://example.org", "table_index": 0, + } + item = {"type": "function_call", "call_id": "c1", + "name": "browser_read_table", "arguments": json.dumps(args)} + events = [{"type": "response.created", "response": {"tools": [ + {**wire, "strict": True} for wire in adapter.wire_tools + ]}}] + if path != "completed": + events.append({"type": "response.output_item.added", "output_index": 0, + "item": {key: item[key] for key in ("type", "call_id", "name")}}) + if path in ("arguments_done", "duplicates"): + events.append({"type": "response.function_call_arguments.done", "output_index": 0, + "arguments": item["arguments"]}) + if path in ("item_done", "duplicates"): + events.append({"type": "response.output_item.done", "output_index": 0, + "item": item}) + events.append({"type": "response.completed", "response": {"output": [item]}}) + token = _request_tool_adapter.set(adapter) + try: + response = await _client()._read_tool_stream(_FakeResp([_sse(e) for e in events])) + finally: + _request_tool_adapter.reset(token) + assert accepted == ["browser_read_table"] + assert len(response.tool_calls) == 1 + assert response.tool_calls[0].input == { + "url": "https://example.org", "table_index": 0, + } + + @pytest.mark.asyncio + async def test_nested_payload_decoded_once_for_duplicate_events(self): + from src.llm.openai_codex import _request_tool_adapter + from src.llm.strict_tool_adapter import compile_catalog + from src.tools.registry import get_tool_definitions + + adapter = compile_catalog(get_tool_definitions() + [{ + "name": "fixture_skill", + "input_schema": {"type": "object", "additionalProperties": False, + "properties": {"text": {"type": "string"}}, + "required": ["text"]}, + }]) + props = next(t["parameters"]["properties"] for t in adapter.wire_tools + if t["name"] == "invoke_skill") + inner = {"text": 'a "quote"'} + payload = {key: None for key in props} | { + "name": "fixture_skill", "input": json.dumps(inner), + } + item = {"type": "function_call", "call_id": "skill-call", "name": "invoke_skill", + "arguments": json.dumps(payload)} + events = [ + {"type": "response.output_item.added", "output_index": 0, + "item": {key: item[key] for key in ("type", "call_id", "name")}}, + {"type": "response.function_call_arguments.done", "output_index": 0, + "arguments": item["arguments"]}, + {"type": "response.output_item.done", "output_index": 0, "item": item}, + {"type": "response.completed", "response": {"output": [item]}}, + ] + token = _request_tool_adapter.set(adapter) + try: + response = await _client()._read_tool_stream(_FakeResp([_sse(e) for e in events])) + finally: + _request_tool_adapter.reset(token) + assert len(response.tool_calls) == 1 + assert response.tool_calls[0].input == {"name": "fixture_skill", "input": inner} + assert adapter.report["invoke_skill"]["resolution"] == "unknown" + @pytest.mark.asyncio async def test_streamed_function_call(self): resp = _FakeResp([ diff --git a/tests/test_strict_tool_adapter.py b/tests/test_strict_tool_adapter.py index a8b215f3c..49d42c22e 100644 --- a/tests/test_strict_tool_adapter.py +++ b/tests/test_strict_tool_adapter.py @@ -107,14 +107,166 @@ def test_external_modes_round_trip_and_resolution(): with pytest.raises(ValueError): adapter.accept("open_ext", {"json": payload}) assert set(adapter.record_resolution(None).values()) == {"unknown"} + adapter = compile_catalog(adapter.catalog) assert adapter.record_resolution( { "type": "response.created", "response": { "tools": [ - {"name": "closed_ext", "strict": True}, - {"name": "open_ext", "strict": False}, + {**adapter.wire_tools[0], "strict": True}, + {**adapter.wire_tools[1], "strict": False}, ] }, } ) == {"closed_ext": "true", "open_ext": "false", "odd_ext": "unknown"} + + +def test_probe_headers_lowered_after_wire_validation(): + adapter = compile_catalog(get_tool_definitions()) + args = _wire_args(adapter, "http_probe", {"url": "https://example.org", "headers": [ + {"name": "Accept", "value": "application/json"}, + ]}) + assert adapter.accept("http_probe", args)["headers"] == {"Accept": "application/json"} + args["headers"].append({"name": "accept", "value": "text/plain"}) + with pytest.raises(ValueError): + adapter.accept("http_probe", args) + + +def _computer_wire(adapter, operation, specified): + branches = next(w["parameters"] for w in adapter.wire_tools if w["name"] == "computer_act") + branch = next( + b for b in branches["properties"]["payload"]["anyOf"] + if b["properties"]["operation"]["const"] == operation + ) + args = { + "session_id": "s", "generation": 1, "consent_generation": 1, + "source_id": "src", "source_revision": 1, + "action_id": "unique", "observation_id": "observed", "operation": operation, + **specified, + } + + def populate(fields, data): + result = {name: data.get(name) for name in fields["properties"]} + for name, value in list(result.items()): + child = fields["properties"][name] + if isinstance(value, dict) and "properties" in child: + result[name] = populate(child, value) + elif name == "steps" and isinstance(value, list): + options = child["items"]["anyOf"] + result[name] = [populate(next( + b for b in options if b["properties"]["operation"]["const"] == step["operation"] + ), step) for step in value] + elif name == "strokes" and isinstance(value, list): + result[name] = [populate(child["items"], stroke) for stroke in value] + return result + + return {"payload": populate(branch, args)} + + +def test_computer_branch_normalization_and_pre_effect_tripwire(): + adapter = compile_catalog(computer_definitions()) + focus = _computer_wire(adapter, "focus", {"x": 1, "y": 2, + "expect": {"type": "visual_change"}}) + assert adapter.accept("computer_act", focus)["expect"] == {"type": "visual_change"} + with pytest.raises(ValueError): + adapter.accept("computer_act", _computer_wire(adapter, "focus", { + "x": 1, "y": 2, "expect": {"type": "pointer_at"}})) + + click = _computer_wire(adapter, "click", {"x": 3, "y": 4, + "expect": {"type": "visual_change"}}) + assert adapter.accept("computer_act", click)["x"] == 3 + region = {"x": 3, "y": 4, "width": 10, "height": 10} + assert adapter.accept("computer_act", _computer_wire(adapter, "click", { + "region": region, "expect": {"type": "visual_change"}}))["region"] == region + for changes in ({"x": 3, "region": region}, {"x": 3}, {"points": [[1, 2], [3, 4]]}): + with pytest.raises(ValueError): + adapter.accept("computer_act", _computer_wire(adapter, "click", { + "expect": {"type": "visual_change"}, **changes})) + + sequence = _computer_wire(adapter, "sequence", {"steps": [ + {"action_id": "step", "operation": "click", "x": 1, "y": 2, + "expect": {"type": "visual_change"}}, + ]}) + assert adapter.accept("computer_act", sequence)["steps"][0]["x"] == 1 + strokes = _computer_wire(adapter, "strokes", {"strokes": [ + {"action_id": "stroke", "points": [[1, 1], [2, 2]], "duration": 0.1}, + ]}) + assert len(adapter.accept("computer_act", strokes)["strokes"]) == 1 + bad_step = _computer_wire(adapter, "sequence", {"steps": [ + {"action_id": "step", "operation": "drag", "points": [[1, 2]], + "duration": 0.1, "expect": {"type": "visual_change"}}, + ]}) + with pytest.raises(ValueError): + adapter.accept("computer_act", bad_step) + bad_stroke = _computer_wire(adapter, "strokes", {"strokes": [ + {"action_id": "stroke", "points": [[1, 1], [2, 2]], + "duration": 0.1, "modifiers": ["shift", "shift"]}, + ]}) + with pytest.raises(ValueError): + adapter.accept("computer_act", bad_stroke) + duplicate_modifiers = _computer_wire(adapter, "drag", { + "points": [[1, 1], [2, 2]], "duration": 0.1, + "modifiers": ["ctrl", "ctrl"], "expect": {"type": "visual_change"}, + }) + with pytest.raises(ValueError): + adapter.accept("computer_act", duplicate_modifiers) + key_repeat = _computer_wire(adapter, "key", { + "key": "ctrl+ctrl+a", "expect": {"type": "visual_change"}, + }) + with pytest.raises(ValueError): + adapter.accept("computer_act", key_repeat) + + +def test_computer_task_context_min_properties(): + adapter = compile_catalog(computer_definitions()) + props = next(w["parameters"]["properties"] for w in adapter.wire_tools + if w["name"] == "computer_observe") + context_props = props["task_context"]["anyOf"][0]["properties"] + args = {k: None for k in props} | { + "session_id": "s", "generation": 1, + "task_context": {k: None for k in context_props}, + } + with pytest.raises(ValueError): + adapter.accept("computer_observe", args) + args["task_context"]["goal"] = "draw a shape" + assert adapter.accept("computer_observe", args)["task_context"] == {"goal": "draw a shape"} + + +def test_spawn_policy_wire_omissions_preserve_presence_checks(): + adapter = compile_catalog(get_tool_definitions()) + schema = next(t["parameters"]["properties"] for t in adapter.wire_tools + if t["name"] == "spawn_agent") + base = {name: None for name in schema} | { + "label": "audit", "goal": "verify", "model": "gpt-6-luna", + } + assert adapter.accept("spawn_agent", base) == { + "label": "audit", "goal": "verify", "model": "gpt-6-luna", + } + for optional, value in (("reasoning_effort", "high"), ("parent_id", "parent")): + assert adapter.accept("spawn_agent", base | {optional: value})[optional] == value + + # Dynamic policy filtering happens before compilation. A forbidden field + # remains forbidden, even when its value is null (not an omission token). + filtered = copy.deepcopy(next(t for t in get_tool_definitions() + if t["name"] == "spawn_agent")) + filtered["input_schema"]["properties"].pop("reasoning_effort") + filtered["input_schema"]["properties"].pop("parent_id") + filtered_adapter = compile_catalog([filtered]) + allowed = {key: value for key, value in base.items() + if key in filtered["input_schema"]["properties"]} + assert "reasoning_effort" not in filtered_adapter.accept("spawn_agent", allowed) + for forbidden in ("reasoning_effort", "parent_id"): + with pytest.raises(ValueError): + filtered_adapter.accept("spawn_agent", allowed | {forbidden: None}) + + +def test_external_nullable_omission_distinguishes_meaningful_null(): + adapter = compile_catalog([{ + "name": "nullable_ext", + "input_schema": {"type": "object", "additionalProperties": False, + "properties": {"value": { + "anyOf": [{"type": "string"}, {"type": "null"}], + }}}, + }]) + assert adapter.accept("nullable_ext", {"value": None}) == {"value": None} + assert adapter.accept("nullable_ext", {"value": {"__odin_omitted__": True}}) == {} From 8f6bbc6e74b22f58b7db8e4e55894915cdb19aa9 Mon Sep 17 00:00:00 2001 From: Odin Date: Sun, 27 Sep 2026 20:22:36 -0400 Subject: [PATCH 03/12] Harden strict tool admission and preserve lint and type gates --- src/discord/background_task.py | 47 ++++++++++---- src/discord/native_tools/agents_tasks.py | 25 +++++--- src/discord/native_tools/scheduling.py | 24 ++++--- src/discord/scheduled_events.py | 11 +++- src/llm/openai_codex.py | 3 +- src/llm/strict_tool_adapter.py | 2 +- src/tools/defs/integrations_email.py | 11 +++- src/tools/defs/media_scheduling.py | 8 ++- src/tools/nested_payload.py | 75 +++++++++++++++++----- src/tools/post_validation.py | 22 +++++-- tests/test_strict_nested_payload.py | 82 ++++++++++++++++-------- tests/test_strict_wire_defs.py | 34 +++++++--- 12 files changed, 244 insertions(+), 100 deletions(-) diff --git a/src/discord/background_task.py b/src/discord/background_task.py index ed3ba1c1c..264df412a 100644 --- a/src/discord/background_task.py +++ b/src/discord/background_task.py @@ -192,21 +192,35 @@ async def run_background_task( if task.nested_payload_validated: try: from ..tools.nested_payload import validate_nested_payload + catalog = getattr(executor, "_tool_catalog", None) definitions = catalog.merged_definitions() if catalog is not None else [] - validate_nested_payload("delegate_task", {"steps": [{**step, "tool_input": tool_input}]}, - definitions, allow_placeholders=False) + validate_nested_payload( + "delegate_task", + {"steps": [{**step, "tool_input": tool_input}]}, + definitions, + allow_placeholders=False, + ) denial = executor.check_permission(tool_name, task.requester_id) if denial: raise ValueError(denial) if tool_name == "invoke_skill": target = tool_input.get("name") + if not isinstance(target, str) or not target: + raise ValueError("invoke_skill requires a selected skill name") denial = executor.check_permission(target, task.requester_id) if denial: raise ValueError(denial) except ValueError as exc: - task.results.append(StepResult(index=i, tool_name=tool_name, - description=step_desc, status="error", output=f"Invalid concrete step payload: {exc}")) + task.results.append( + StepResult( + index=i, + tool_name=tool_name, + description=step_desc, + status="error", + output=f"Invalid concrete step payload: {exc}", + ) + ) if on_failure == "abort": task.status = "failed" break @@ -471,10 +485,22 @@ async def _execute_tool( with execution_delivery_scope(requester_id): result = await _execute_tool_captured( - tool_name, tool_input, executor, skill_manager, knowledge_store, - embedder, requester, step_desc, mcp_manager, requester_id) + tool_name, + tool_input, + executor, + skill_manager, + knowledge_store, + embedder, + requester, + step_desc, + mcp_manager, + requester_id, + ) return deliver_runtime_result( - executor, result, tool_name=tool_name, tool_input=tool_input, + executor, + result, + tool_name=tool_name, + tool_input=tool_input, user_id=requester_id, ) @@ -580,8 +606,7 @@ async def _execute_tool_captured( return RankedOutput( formatted, matches=tuple( - f"[{r['source']}] (score: {r.get('score', r.get('rrf_score', 0))}): " - f"{r['content']}" + f"[{r['source']}] (score: {r.get('score', r.get('rrf_score', 0))}): {r['content']}" for r in results ), recovery_required=any(len(r["content"]) > 200 for r in results), @@ -626,9 +651,7 @@ async def _execute_tool_captured( target_name, skill_input, requester_id=requester_id or None ) if skill_manager.has_skill(tool_name): - return await skill_manager.execute( - tool_name, tool_input, requester_id=requester_id or None - ) + return await skill_manager.execute(tool_name, tool_input, requester_id=requester_id or None) # MCP tools (namespaced as mcp__) if mcp_manager is not None and mcp_manager.has_tool(tool_name): diff --git a/src/discord/native_tools/agents_tasks.py b/src/discord/native_tools/agents_tasks.py index 91c5ea15f..de0c10e1b 100644 --- a/src/discord/native_tools/agents_tasks.py +++ b/src/discord/native_tools/agents_tasks.py @@ -27,8 +27,8 @@ from ...llm.tool_history import normalize_tool_calls from ...odin_log import get_logger from ...tools.defs.agents import SPAWN_NEUTRAL_REASONING_OPTIONS -from ...tools.result_validator import ToolResult from ...tools.nested_payload import validate_nested_payload +from ...tools.result_validator import ToolResult from ..background_task import ( MAX_STEPS, BackgroundTask, @@ -690,10 +690,14 @@ def _agent_iteration_cap(agents_cfg, *, provider: str, scheduled: bool) -> int: if provider != "codex": return hard_max configured = ( - getattr(agents_cfg, "scheduled_max_iterations", 180) - if scheduled - else getattr(agents_cfg, "max_iterations", 120) - ) if agents_cfg else (180 if scheduled else 120) + ( + getattr(agents_cfg, "scheduled_max_iterations", 180) + if scheduled + else getattr(agents_cfg, "max_iterations", 120) + ) + if agents_cfg + else (180 if scheduled else 120) + ) return min(configured, hard_max) @@ -755,14 +759,17 @@ async def _handle_delegate_task(self, message: discord.Message, inp: dict) -> st """Create and start a background task.""" try: inp = validate_nested_payload( - "delegate_task", inp, + "delegate_task", + inp, self._tool_catalog.merged_definitions() if self._tool_catalog else [], ) except ValueError as e: return f"Invalid background task payload: {e}" for i, step in enumerate(inp.get("steps", []), 1): if isinstance(step, dict) and step.get("tool_name"): - denied = self._tool_executor.check_permission(step["tool_name"], str(message.author.id)) + denied = self._tool_executor.check_permission( + step["tool_name"], str(message.author.id) + ) if denied: return f"Step {i}: {denied}" if step["tool_name"] == "invoke_skill" and isinstance(step.get("tool_input"), dict): @@ -1160,9 +1167,7 @@ async def _handle_spawn_agent(self, message: object, inp: dict) -> str: # list so the ``model_reasoning_dialect`` consumer below sees a # clean ``str`` rather than ``Any | None``. if not native_choices and _model_mode != "auto": - native_choices = [ - configured_agent_model(self._get_config()) or DEFAULT_AGENT_MODEL - ] + native_choices = [configured_agent_model(self._get_config()) or DEFAULT_AGENT_MODEL] if native_choices and all( model_reasoning_dialect(self._get_config(), item) == "effort" for item in native_choices ): diff --git a/src/discord/native_tools/scheduling.py b/src/discord/native_tools/scheduling.py index 100fecce8..b7abb7426 100644 --- a/src/discord/native_tools/scheduling.py +++ b/src/discord/native_tools/scheduling.py @@ -174,13 +174,21 @@ async def _handle_update_schedule(self, inp: dict) -> str: try: inp = validate_nested_payload("update_schedule", inp, self._nested_catalog()) if isinstance(inp.get("tool_input"), dict) and not inp.get("tool_name"): - current = next((s for s in self.scheduler.list_all() if s.get("id") == inp.get("schedule_id")), None) + current = next( + (s for s in self.scheduler.list_all() if s.get("id") == inp.get("schedule_id")), + None, + ) target_name = (current or {}).get("tool_name") if target_name: - validate_nested_payload("schedule_task", { - "action": "check", "tool_name": target_name, - "tool_input": inp["tool_input"], - }, self._nested_catalog()) + validate_nested_payload( + "schedule_task", + { + "action": "check", + "tool_name": target_name, + "tool_input": inp["tool_input"], + }, + self._nested_catalog(), + ) except ValueError as e: return f"Error: {e}" schedule_id = inp.get("schedule_id", "") @@ -221,10 +229,7 @@ async def _handle_update_schedule(self, inp: dict) -> str: if result is None: return f"Schedule {schedule_id} not found." if result.get("inert_reason"): - return ( - f"Schedule {schedule_id} remains paused and inert: " - f"{result['inert_reason']}" - ) + return f"Schedule {schedule_id} remains paused and inert: {result['inert_reason']}" return f"Updated schedule {schedule_id}." def _nested_catalog(self): @@ -235,6 +240,7 @@ def _nested_catalog(self): # ToolExecutor is available via the scheduler handler owner in production; # static definitions still validate direct unit/legacy entry points. from ...tools.registry import get_tool_definitions + return get_tool_definitions() async def _handle_delete_schedule(self, inp: dict) -> str: diff --git a/src/discord/scheduled_events.py b/src/discord/scheduled_events.py index 71559b068..73dc6a8b9 100644 --- a/src/discord/scheduled_events.py +++ b/src/discord/scheduled_events.py @@ -319,11 +319,16 @@ async def _run_scheduled_workflow( if nested_validated: try: validate_nested_payload( - "delegate_task", {"steps": [{**step, "tool_input": tool_input}]}, - catalog, allow_placeholders=False, + "delegate_task", + {"steps": [{**step, "tool_input": tool_input}]}, + catalog, + allow_placeholders=False, ) except ValueError as exc: - results.append(f"**Step {i + 1}** (`{step.get('description', tool_name)}`): invalid payload: {exc}") + results.append( + f"**Step {i + 1}** (`{step.get('description', tool_name)}`): " + f"invalid payload: {exc}" + ) workflow_ok = False break denial = self._tool_executor.check_permission(tool_name, req_id) diff --git a/src/llm/openai_codex.py b/src/llm/openai_codex.py index 3dc5f00df..f9d3f6318 100644 --- a/src/llm/openai_codex.py +++ b/src/llm/openai_codex.py @@ -25,10 +25,11 @@ ) from .progress import GenerationProgress, GenerationProgressObserver, emit_progress from .secret_scrubber import scrub_output_secrets +from .strict_tool_adapter import RequestToolAdapter from .types import LLMResponse, ToolCall log = get_logger("codex") -_request_tool_adapter: ContextVar[object | None] = ContextVar( +_request_tool_adapter: ContextVar[RequestToolAdapter | None] = ContextVar( "codex_request_tool_adapter", default=None ) diff --git a/src/llm/strict_tool_adapter.py b/src/llm/strict_tool_adapter.py index a1ce73039..3473b05c5 100644 --- a/src/llm/strict_tool_adapter.py +++ b/src/llm/strict_tool_adapter.py @@ -263,7 +263,7 @@ def __init__(self, tools): self.catalog = deepcopy(tools) self.wire_tools = [] self.report = {} - self._contracts = {} + self._contracts: dict[str, tuple[dict, dict, str]] = {} self._resolution_logged = False builtins = set(TOOL_MAP) | set(COMPUTER_TOOL_NAMES) for tool in self.catalog: diff --git a/src/tools/defs/integrations_email.py b/src/tools/defs/integrations_email.py index df9756bbb..467c36ac0 100644 --- a/src/tools/defs/integrations_email.py +++ b/src/tools/defs/integrations_email.py @@ -43,7 +43,8 @@ }, "headers": { "type": "array", - "description": "Request headers as name/value entries. Case-insensitive duplicate names are rejected.", + "description": "Request headers as name/value entries. " + "Case-insensitive duplicate names are rejected.", "items": { "type": "object", "properties": {"name": {"type": "string"}, "value": {"type": "string"}}, @@ -181,9 +182,13 @@ }, "target": {"type": "string"}, "expected": { - "description": "Type-specific expectation: integer, string, integer list, or string list", + "description": ( + "Type-specific expectation: integer, string, integer list, " + "or string list" + ), "anyOf": [ - {"type": "integer"}, {"type": "string"}, + {"type": "integer"}, + {"type": "string"}, {"type": "array", "items": {"type": "integer"}}, {"type": "array", "items": {"type": "string"}}, ], diff --git a/src/tools/defs/media_scheduling.py b/src/tools/defs/media_scheduling.py index 2bd9ebea4..371dfbfac 100644 --- a/src/tools/defs/media_scheduling.py +++ b/src/tools/defs/media_scheduling.py @@ -277,15 +277,17 @@ "run_at": { "type": "string", "description": ( - "New offset-aware ISO datetime for one-time " - "(replaces previous timing)" + "New offset-aware ISO datetime for one-time (replaces previous timing)" ), }, "trigger": { "type": "object", "description": "New webhook trigger (replaces previous timing)", "properties": { - "source": {"type": "string", "enum": ["gitea", "grafana", "generic", "github", "gitlab"]}, + "source": { + "type": "string", + "enum": ["gitea", "grafana", "generic", "github", "gitlab"], + }, "event": {"type": "string"}, "repo": {"type": "string"}, "alert_name": {"type": "string"}, diff --git a/src/tools/nested_payload.py b/src/tools/nested_payload.py index 3faf85cd7..cc772c961 100644 --- a/src/tools/nested_payload.py +++ b/src/tools/nested_payload.py @@ -1,4 +1,5 @@ """Canonical decoding and validation for tool inputs nested as JSON strings on strict wire calls.""" + from __future__ import annotations import json @@ -20,7 +21,7 @@ def _object_no_duplicates(pairs): result = {} for key, value in pairs: if key in result: - raise ValueError(f"duplicate JSON key in nested payload: {key!r}") + raise ValueError("duplicate JSON key in nested payload") result[key] = value return result @@ -32,8 +33,13 @@ def decode_json_object(value: Any, label: str) -> dict: if not isinstance(value, str): raise ValueError(f"{label} must be a JSON object string or object") try: - result = json.loads(value, object_pairs_hook=_object_no_duplicates, - parse_constant=lambda token: (_ for _ in ()).throw(ValueError(f"invalid JSON constant {token}"))) + result = json.loads( + value, + object_pairs_hook=_object_no_duplicates, + parse_constant=lambda token: (_ for _ in ()).throw( + ValueError(f"invalid JSON constant {token}") + ), + ) except (json.JSONDecodeError, ValueError) as exc: raise ValueError(f"{label} must contain valid JSON without duplicate keys: {exc}") from exc if not isinstance(result, dict): @@ -46,7 +52,11 @@ def decode_nested_payloads(tool_name: str, arguments: dict) -> dict: if tool_name not in _NESTED_FIELDS: return arguments result = dict(arguments) - if tool_name in ("schedule_task", "update_schedule") and "tool_input" in result and isinstance(result["tool_input"], str): + if ( + tool_name in ("schedule_task", "update_schedule") + and "tool_input" in result + and isinstance(result["tool_input"], str) + ): result["tool_input"] = decode_json_object(result["tool_input"], "tool_input") if tool_name == "invoke_skill" and "input" in result and isinstance(result["input"], str): result["input"] = decode_json_object(result["input"], "input") @@ -54,10 +64,13 @@ def decode_nested_payloads(tool_name: str, arguments: dict) -> dict: steps = [] for index, original in enumerate(result["steps"], 1): if not isinstance(original, dict): - steps.append(original); continue + steps.append(original) + continue step = dict(original) if isinstance(step.get("tool_input"), str): - step["tool_input"] = decode_json_object(step["tool_input"], f"step {index} tool_input") + step["tool_input"] = decode_json_object( + step["tool_input"], f"step {index} tool_input" + ) steps.append(step) result["steps"] = steps return result @@ -70,10 +83,9 @@ def _schema_for(name: str, catalog: list[dict]) -> dict | None: def _validate(schema: dict, payload: dict, *, allow_placeholders: bool) -> None: validator_cls = validator_for(schema) validator_cls.check_schema(schema) - validator = validator_cls(schema) import copy + candidate = copy.deepcopy(payload) - deferred: set[tuple] = set() if allow_placeholders: # The established workflow substitution contract touches direct string # values in tool_input only, not nested objects or array members. @@ -82,21 +94,33 @@ def _validate(schema: dict, payload: dict, *, allow_placeholders: bool) -> None: for key, value in candidate.items(): if isinstance(value, str) and _PLACEHOLDER.search(value): schema.setdefault("properties", {})[key] = {} - errors = sorted(validator_cls(schema).iter_errors(candidate), key=lambda e: list(map(str, e.path))) + errors = sorted( + validator_cls(schema).iter_errors(candidate), key=lambda e: list(map(str, e.path)) + ) if errors: err = errors[0] path = ".".join(map(str, err.absolute_path)) or "" - raise ValueError(f"invalid input for selected tool at {path}: {err.message}") - - -def validate_nested_payload(tool_name: str, canonical_arguments: dict, catalog: list[dict], *, allow_placeholders: bool = True) -> dict: + raise ValueError( + f"invalid input for selected tool at {path}: violates {err.validator} constraint" + ) + + +def validate_nested_payload( + tool_name: str, + canonical_arguments: dict, + catalog: list[dict], + *, + allow_placeholders: bool = True, +) -> dict: """Decode nested JSON fields, validate selected canonical targets, return canonical args. This validates payload shape only. Callers must also apply their existing authorization checks to the selected target; a name in a payload is not permission. """ from .registry import TOOLS + args = decode_nested_payloads(tool_name, canonical_arguments) + def check(target, payload, label): schema = _schema_for(target, catalog) or _schema_for(target, TOOLS) if schema is None: @@ -104,17 +128,34 @@ def check(target, payload, label): if not isinstance(payload, dict): raise ValueError(f"{label} must be an object") _validate(schema, payload, allow_placeholders=allow_placeholders) + if tool_name in ("schedule_task", "update_schedule"): - if args.get("tool_name") and isinstance(args.get("tool_input"), dict) and args.get("action", "check") == "check": + if ( + args.get("tool_name") + and isinstance(args.get("tool_input"), dict) + and args.get("action", "check") == "check" + ): check(args["tool_name"], args["tool_input"], "tool_input") - if tool_name == "update_schedule" and isinstance(args.get("tool_input"), dict) and args.get("tool_name"): + if ( + tool_name == "update_schedule" + and isinstance(args.get("tool_input"), dict) + and args.get("tool_name") + ): check(args["tool_name"], args["tool_input"], "tool_input") for i, step in enumerate(args.get("steps") or [], 1): - if isinstance(step, dict) and step.get("tool_name") and isinstance(step.get("tool_input"), dict): + if ( + isinstance(step, dict) + and step.get("tool_name") + and isinstance(step.get("tool_input"), dict) + ): check(step["tool_name"], step["tool_input"], f"step {i} tool_input") elif tool_name == "delegate_task": for i, step in enumerate(args.get("steps") or [], 1): - if isinstance(step, dict) and step.get("tool_name") and isinstance(step.get("tool_input"), dict): + if ( + isinstance(step, dict) + and step.get("tool_name") + and isinstance(step.get("tool_input"), dict) + ): check(step["tool_name"], step["tool_input"], f"step {i} tool_input") elif tool_name == "invoke_skill": target = args.get("name") diff --git a/src/tools/post_validation.py b/src/tools/post_validation.py index 35c29ad19..ab0b2205d 100644 --- a/src/tools/post_validation.py +++ b/src/tools/post_validation.py @@ -215,16 +215,28 @@ def parse_checks(raw_checks: list[dict]) -> tuple[list[Check], list[str]]: # Keep the public union deliberately small, then enforce the # check-specific interpretation before any check is launched. valid_scalar = isinstance(expected, (int, str)) and not isinstance(expected, bool) - valid_list = isinstance(expected, list) and bool(expected) and all( - isinstance(item, (int, str)) and not isinstance(item, bool) for item in expected + valid_list = ( + isinstance(expected, list) + and bool(expected) + and all( + isinstance(item, (int, str)) and not isinstance(item, bool) for item in expected + ) ) if not (valid_scalar or valid_list): - errors.append(f"check[{i}]: expected must be an integer, string, or non-empty list of integers/strings") + errors.append( + f"check[{i}]: expected must be an integer, string, or non-empty list " + f"of integers/strings" + ) continue if c_type == "http": values = expected if isinstance(expected, list) else [expected] - if not all(isinstance(v, int) or (isinstance(v, str) and v.isdigit()) for v in values): - errors.append(f"check[{i}]: http expected values must be status-code integers or digit strings") + if not all( + isinstance(v, int) or (isinstance(v, str) and v.isdigit()) for v in values + ): + errors.append( + f"check[{i}]: http expected values must be status-code integers or digit " + f"strings" + ) continue elif c_type == "service": if isinstance(expected, list) and not all(isinstance(v, str) for v in expected): diff --git a/tests/test_strict_nested_payload.py b/tests/test_strict_nested_payload.py index 1a3cac685..f2709acd4 100644 --- a/tests/test_strict_nested_payload.py +++ b/tests/test_strict_nested_payload.py @@ -1,4 +1,5 @@ """Strict wire nested JSON payload decoding and selected-target validation.""" + import json import pytest @@ -8,49 +9,76 @@ def catalog(): return [ - {"name": "run_command", "input_schema": { - "type": "object", "properties": {"command": {"type": "string"}, "host": {"type": ["string", "null"]}}, - "required": ["command"], "additionalProperties": False, - }}, - {"name": "example_skill", "input_schema": { - "type": "object", "properties": {"value": {"type": ["string", "null"]}}, - "required": ["value"], "additionalProperties": False, - }}, + { + "name": "run_command", + "input_schema": { + "type": "object", + "properties": {"command": {"type": "string"}, "host": {"type": ["string", "null"]}}, + "required": ["command"], + "additionalProperties": False, + }, + }, + { + "name": "example_skill", + "input_schema": { + "type": "object", + "properties": {"value": {"type": ["string", "null"]}}, + "required": ["value"], + "additionalProperties": False, + }, + }, ] def test_wire_json_decodes_canonical_payload_and_preserves_meaningful_null(): - raw = {'action': 'check', 'tool_name': 'run_command', - 'tool_input': json.dumps({'command': 'printf "\\u03bb"', 'host': None})} - result = validate_nested_payload('schedule_task', raw, catalog()) - assert result['tool_input'] == {'command': 'printf "\\u03bb"', 'host': None} - assert isinstance(raw['tool_input'], str) + raw = { + "action": "check", + "tool_name": "run_command", + "tool_input": json.dumps({"command": 'printf "\\u03bb"', "host": None}), + } + result = validate_nested_payload("schedule_task", raw, catalog()) + assert result["tool_input"] == {"command": 'printf "\\u03bb"', "host": None} + assert isinstance(raw["tool_input"], str) -@pytest.mark.parametrize('payload', ['{"command":"a","command":"b"}', '[]', 'null', '{bad']) +@pytest.mark.parametrize("payload", ['{"command":"a","command":"b"}', "[]", "null", "{bad"]) def test_rejects_bad_wire_nested_json(payload): with pytest.raises(ValueError): - decode_nested_payloads('delegate_task', {'steps': [{'tool_name': 'run_command', 'tool_input': payload}]}) + decode_nested_payloads( + "delegate_task", {"steps": [{"tool_name": "run_command", "tool_input": payload}]} + ) def test_selected_target_schema_is_enforced(): - with pytest.raises(ValueError, match='invalid input'): - validate_nested_payload('delegate_task', { - 'steps': [{'tool_name': 'run_command', 'tool_input': {'command': 'ok', 'surprise': True}}] - }, catalog()) + with pytest.raises(ValueError, match="invalid input"): + validate_nested_payload( + "delegate_task", + { + "steps": [ + {"tool_name": "run_command", "tool_input": {"command": "ok", "surprise": True}} + ] + }, + catalog(), + ) def test_template_field_can_defer_validation_and_null_is_not_dropped(): - item = {'tool_name': 'run_command', 'tool_input': {'command': '{prev_output}', 'host': None}} - result = validate_nested_payload('delegate_task', {'steps': [item]}, catalog()) - assert result['steps'][0]['tool_input'] == item['tool_input'] - invalid = {'tool_name': 'run_command', 'tool_input': {'command': None, 'host': None}} + item = {"tool_name": "run_command", "tool_input": {"command": "{prev_output}", "host": None}} + result = validate_nested_payload("delegate_task", {"steps": [item]}, catalog()) + assert result["steps"][0]["tool_input"] == item["tool_input"] + invalid = {"tool_name": "run_command", "tool_input": {"command": None, "host": None}} with pytest.raises(ValueError): - validate_nested_payload('delegate_task', {'steps': [invalid]}, catalog(), allow_placeholders=False) + validate_nested_payload( + "delegate_task", {"steps": [invalid]}, catalog(), allow_placeholders=False + ) def test_invoke_skill_validates_named_skill_schema(): - result = validate_nested_payload('invoke_skill', {'name': 'example_skill', 'input': '{"value":null}'}, catalog()) - assert result['input'] == {'value': None} + result = validate_nested_payload( + "invoke_skill", {"name": "example_skill", "input": '{"value":null}'}, catalog() + ) + assert result["input"] == {"value": None} with pytest.raises(ValueError): - validate_nested_payload('invoke_skill', {'name': 'example_skill', 'input': {'wrong': 1}}, catalog()) + validate_nested_payload( + "invoke_skill", {"name": "example_skill", "input": {"wrong": 1}}, catalog() + ) diff --git a/tests/test_strict_wire_defs.py b/tests/test_strict_wire_defs.py index a0bb291a3..6873fcb84 100644 --- a/tests/test_strict_wire_defs.py +++ b/tests/test_strict_wire_defs.py @@ -1,4 +1,5 @@ """Focused tests for strict wire definitions and their runtime lowerings.""" + from __future__ import annotations import pytest @@ -22,11 +23,14 @@ def test_http_probe_headers_are_name_value_records_and_legacy_dict_still_works() assert "X-Test: one" in cmd -@pytest.mark.parametrize("headers", [ - [{"name": "Accept", "value": "a"}, {"name": "accept", "value": "b"}], - [{"name": "X", "value": "a", "extra": "no"}], - [{"name": "X", "value": 1}], -]) +@pytest.mark.parametrize( + "headers", + [ + [{"name": "Accept", "value": "a"}, {"name": "accept", "value": "b"}], + [{"name": "X", "value": "a", "extra": "no"}], + [{"name": "X", "value": 1}], + ], +) def test_http_probe_wire_headers_reject_duplicates_and_bad_entries(headers): with pytest.raises(ValueError): build_http_probe_command({"url": "https://example.com", "headers": headers}) @@ -37,7 +41,13 @@ def test_schedule_triggers_share_closed_four_field_shape(): trigger = _schema(SCHEDULING_TOOLS, tool_name)["properties"]["trigger"] assert trigger["additionalProperties"] is False assert set(trigger["properties"]) == {"source", "event", "repo", "alert_name"} - assert trigger["properties"]["source"]["enum"] == ["gitea", "grafana", "generic", "github", "gitlab"] + assert trigger["properties"]["source"]["enum"] == [ + "gitea", + "grafana", + "generic", + "github", + "gitlab", + ] def test_trigger_runtime_keeps_partial_conditions_and_rejects_empty(): @@ -51,9 +61,13 @@ def test_trigger_runtime_keeps_partial_conditions_and_rejects_empty(): def test_validate_action_expected_union_and_runtime_typed_values(): - expected = _schema(INTEGRATION_TOOLS, "validate_action")["properties"]["checks"]["items"]["properties"]["expected"] + expected = _schema(INTEGRATION_TOOLS, "validate_action")["properties"]["checks"]["items"][ + "properties" + ]["expected"] assert [x["type"] for x in expected["anyOf"]] == ["integer", "string", "array", "array"] - checks, errors = parse_checks([{"type": "http", "target": "https://example.com", "expected": [200, 204]}]) + checks, errors = parse_checks( + [{"type": "http", "target": "https://example.com", "expected": [200, 204]}] + ) assert not errors and checks[0].expected == [200, 204] for raw in ( {"type": "http", "target": "https://example.com", "expected": [200, "up"]}, @@ -64,5 +78,7 @@ def test_validate_action_expected_union_and_runtime_typed_values(): def test_legacy_validation_expectations_keep_stringification_paths(): - checks, errors = parse_checks([{"type": "command", "target": "printf 12", "compare": "equals", "expected": 12}]) + checks, errors = parse_checks( + [{"type": "command", "target": "printf 12", "compare": "equals", "expected": 12}] + ) assert not errors and checks[0].expected == 12 From 7cc8b701e7e8d6ee670e4f74ba438fd53bea6938 Mon Sep 17 00:00:00 2001 From: Odin Date: Sun, 27 Sep 2026 20:47:00 -0400 Subject: [PATCH 04/12] Preserve native provider behavior and refresh strict transport contracts --- docs/reference/tools.md | 14 +++++-- docs/strict-wire-lowerings.md | 28 +++++++++++--- src/discord/native_tools/agents_tasks.py | 43 +++++++++++----------- src/discord/native_tools/scheduling.py | 42 ++++++++++++--------- src/discord/native_tools/skills_tools.py | 9 +---- src/discord/scheduled_events.py | 25 ++++++++++++- src/scheduler/scheduler.py | 6 ++- src/tools/nested_payload.py | 10 ++++- tests/characterization/test_tool_parity.py | 17 ++++----- tests/test_computer_tool_schema_r5.py | 26 ++++++++----- tests/test_native_agents_tasks.py | 24 ++++++++++++ tests/test_native_scheduling.py | 24 ++++++++++++ tests/test_provider_stream_acceptance.py | 11 ++++-- 13 files changed, 197 insertions(+), 82 deletions(-) diff --git a/docs/reference/tools.md b/docs/reference/tools.md index 408acda78..78fb2e898 100644 --- a/docs/reference/tools.md +++ b/docs/reference/tools.md @@ -150,7 +150,7 @@ Source: [`src/tools/defs/media_scheduling.py`](https://github.com/Calmingstorm/O | cron | string | No | Cron expression for recurring tasks (e.g. '0 9 * * *' = daily 9am). Omit for one-time. | | cron_timezone | string | No | IANA timezone for the cron expression (e.g. 'America/New_York'). The task fires on that timezone's wall clock across DST. Defaults to UTC. | | run_at | string | No | Offset-aware ISO datetime for one-time tasks (e.g. '2026-03-20T09:00:00Z'). Use parse_time to convert natural language. Omit for recurring. | -| trigger | object | No | Webhook trigger (AND logic). E.g. {"source": "github", "event": "push", "repo": "myproject"}. | +| trigger | object | No | Webhook trigger (AND logic). E.g. {"source": "github", "event": "push", "repo": "myproject"}.
Constraints: {"additionalProperties":false} | | trigger.source | string | No | Webhook source to match
Constraints: {"enum":["gitea","grafana","generic","github","gitlab"]} | | trigger.event | string | No | Event type (e.g. 'push', 'pull_request', 'alert') | | trigger.repo | string | No | Repository name substring (case-insensitive) | @@ -194,7 +194,11 @@ No input properties. | cron | string | No | New cron expression (replaces previous timing) | | cron_timezone | string | No | IANA timezone for the cron expression (e.g. 'America/New_York'). Defaults to UTC. | | run_at | string | No | New offset-aware ISO datetime for one-time (replaces previous timing) | -| trigger | object | No | New webhook trigger (replaces previous timing) | +| trigger | object | No | New webhook trigger (replaces previous timing)
Constraints: {"additionalProperties":false} | +| trigger.source | string | No |
Constraints: {"enum":["gitea","grafana","generic","github","gitlab"]} | +| trigger.event | string | No | — | +| trigger.repo | string | No | — | +| trigger.alert_name | string | No | — | | message | string | No | New message (for reminder actions) | | tool_name | string | No | New tool name (for check actions) | | tool_input | object | No | New tool input parameters | @@ -899,7 +903,9 @@ Source: [`src/tools/defs/integrations_email.py`](https://github.com/Calmingstorm | url | string | Yes | URL to probe (http or https) | | host | string | No | Host alias to run curl from (omit to run locally) | | method | string | No | HTTP method (default GET)
Constraints: {"enum":["GET","POST","PUT","DELETE","PATCH","HEAD","OPTIONS"]} | -| headers | object | No | Request headers as key-value pairs (e.g. {"Authorization": "Bearer tok"}) | +| headers | array<object> | No | Request headers as name/value entries. Case-insensitive duplicate names are rejected. | +| headers[].name | string | Yes | — | +| headers[].value | string | Yes | — | | body | string | No | Request body string (for POST/PUT/PATCH). Max 50KB. | | timeout | integer | No | Request timeout in seconds (default 30, max 120) | | follow_redirects | boolean | No | Follow HTTP redirects (default true) | @@ -949,7 +955,7 @@ Severity 'critical' (default), 'warn', or 'info'. | checks | array<object> | Yes | List of validation checks (max 25). | | checks[].type | string | Yes | http|port|service|process|log_absent|log_present|command | | checks[].target | string | Yes | — | -| checks[].expected | any | No | Type-specific expectation (int, string, list) | +| checks[].expected | anyOf(integer, string, array<integer>, array<string>) | No | Type-specific expectation: integer, string, integer list, or string list
Constraints: {"anyOf":[{"type":"integer"},{"type":"string"},{"type":"array","items":{"type":"integer"}},{"type":"array","items":{"type":"string"}}]} | | checks[].severity | string | No | critical (default) | warn | info | | checks[].host | string | No | — | | checks[].compare | string | No | — | diff --git a/docs/strict-wire-lowerings.md b/docs/strict-wire-lowerings.md index 499c6307b..e7acf028f 100644 --- a/docs/strict-wire-lowerings.md +++ b/docs/strict-wire-lowerings.md @@ -39,6 +39,22 @@ tested at both the wire and runtime boundaries. - Fixtures: `tests/test_strict_wire_defs.py` covers integer lists, invalid mixed HTTP values, invalid service list values, and command compatibility. +## Nested tool payloads + +- Canonical paths: `schedule_task.tool_input`, `schedule_task.steps[].tool_input`, + `update_schedule.tool_input`, `update_schedule.steps[].tool_input`, + `delegate_task.steps[].tool_input`, and `invoke_skill.input`. +- Wire form: JSON-encoded object in a string field. The request-local acceptance + boundary decodes once, rejects malformed JSON, duplicate keys, and non-object + roots, then validates the selected tool's canonical schema and authorization. + Canonical objects, including meaningful nested nulls, reach persistence and + dispatch. Workflow templates are decoded before substitution; concrete fields + are checked at admission, unresolved placeholders after substitution. Existing + stored jobs do not acquire new retroactive validation. +- Fixtures: `tests/test_strict_nested_payload.py` and + `tests/test_strict_tool_adapter.py` cover round trips, malformed data, + target validation, and unresolved placeholders. + ## `computer_*` explicit lowerings The canonical computer contracts remain in `src/tools/defs/computer.py`. @@ -53,12 +69,12 @@ tripwire before dispatch; these are not generic recursive keyword deletions. | Canonical schema constraint | Wire representation | Pre-effect validator | | --- | --- | --- | -| `uniqueItems` (notably modifier arrays) | Omitted from wire array schema | Canonical computer payload validator rejects duplicates, including repeated key modifiers | -| `minProperties` (task context) | Omitted from wire object schema | Canonical validator rejects empty task context where minimum properties are required | -| `not` | Omitted from wire branch | Canonical validator rejects the forbidden shape | -| `allOf` | Omitted from wire branch | Canonical validator validates all constituent schemas | -| `oneOf` | Represented by nested operation branches only when match is unambiguous | Canonical validator and ambiguity check reject zero/multiple distinct matches | -| `false` property schemas | Property omitted from wire branch | Canonical validator rejects forbidden operation fields | +| `computer_act.properties.modifiers.uniqueItems`, `computer_act.properties.steps.items.properties.modifiers.uniqueItems`, `computer_act.properties.strokes.items.properties.modifiers.uniqueItems` | Omitted from wire array schemas | Canonical computer payload validator rejects duplicate modifiers | +| `computer_observe.properties.task_context.minProperties` | Omitted from wire object schema | Canonical validator rejects an empty task context | +| `computer_act.properties.key.allOf[*].not`, and the matching nested `computer_act.properties.steps.items.properties.key.allOf[*].not` | Omitted from wire key schema | Canonical validator rejects repeated key modifiers | +| `computer_act.properties.key.allOf`, and the matching nested `computer_act.properties.steps.items.properties.key.allOf` | Omitted from wire key schema | Canonical validator applies every key-chord restriction | +| `computer_act.oneOf`, `computer_act.properties.steps.items.oneOf` and their coordinate-versus-region sub-unions | Nested operation branches with const selectors; coordinate and region remain distinct branches | Canonical validator and adapter ambiguity check reject conflicting and unmatched branches | +| `false` property schemas under `computer_act.oneOf[*].properties` and `computer_act.properties.steps.items.oneOf[*].properties`, including coordinate-versus-region sub-unions | Forbidden property removed from each wire branch | Canonical validator rejects forbidden fields, even if supplied as null | `tests/test_strict_tool_adapter.py` exercises positive focus, click-coordinate, click-region, sequence and stroke forms; negative focus expectation, diff --git a/src/discord/native_tools/agents_tasks.py b/src/discord/native_tools/agents_tasks.py index de0c10e1b..4c1ce1824 100644 --- a/src/discord/native_tools/agents_tasks.py +++ b/src/discord/native_tools/agents_tasks.py @@ -27,7 +27,7 @@ from ...llm.tool_history import normalize_tool_calls from ...odin_log import get_logger from ...tools.defs.agents import SPAWN_NEUTRAL_REASONING_OPTIONS -from ...tools.nested_payload import validate_nested_payload +from ...tools.nested_payload import ValidatedNestedPayload from ...tools.result_validator import ToolResult from ..background_task import ( MAX_STEPS, @@ -757,26 +757,10 @@ def __init__(self, deps: AgentTaskDeps) -> None: async def _handle_delegate_task(self, message: discord.Message, inp: dict) -> str: """Create and start a background task.""" - try: - inp = validate_nested_payload( - "delegate_task", - inp, - self._tool_catalog.merged_definitions() if self._tool_catalog else [], - ) - except ValueError as e: - return f"Invalid background task payload: {e}" - for i, step in enumerate(inp.get("steps", []), 1): - if isinstance(step, dict) and step.get("tool_name"): - denied = self._tool_executor.check_permission( - step["tool_name"], str(message.author.id) - ) - if denied: - return f"Step {i}: {denied}" - if step["tool_name"] == "invoke_skill" and isinstance(step.get("tool_input"), dict): - target = step["tool_input"].get("name") - denied = self._tool_executor.check_permission(target, str(message.author.id)) - if denied: - return f"Step {i}: {denied}" + # RequestToolAdapter validates Codex payloads before dispatch. Legacy + # providers supply canonical objects directly and retain their existing + # permissive delegation contract, including deferred execution checks. + nested_validated = isinstance(inp, ValidatedNestedPayload) description = inp.get("description", "Background task") steps = inp.get("steps", []) @@ -805,6 +789,21 @@ async def _handle_delegate_task(self, message: discord.Message, inp: dict) -> st f"Rebuild the steps with proper tool_input and retry." ) + if nested_validated: + for i, step in enumerate(steps, 1): + denied = self._tool_executor.check_permission( + step["tool_name"], str(message.author.id) + ) + if isinstance(denied, str) and denied: + return f"Step {i}: {denied}" + if step["tool_name"] == "invoke_skill": + target = (step.get("tool_input") or {}).get("name") + if not isinstance(target, str) or not target: + return f"Step {i}: invoke_skill requires a selected skill name" + denied = self._tool_executor.check_permission(target, str(message.author.id)) + if isinstance(denied, str) and denied: + return f"Step {i}: {denied}" + task = BackgroundTask( task_id=create_task_id(), description=description, @@ -812,7 +811,7 @@ async def _handle_delegate_task(self, message: discord.Message, inp: dict) -> st channel=message.channel, requester=str(message.author), requester_id=str(message.author.id), - nested_payload_validated=True, + nested_payload_validated=nested_validated, ) # Prune old completed tasks diff --git a/src/discord/native_tools/scheduling.py b/src/discord/native_tools/scheduling.py index b7abb7426..573308c64 100644 --- a/src/discord/native_tools/scheduling.py +++ b/src/discord/native_tools/scheduling.py @@ -10,7 +10,7 @@ from ...odin_log import get_logger from ...scheduler.scheduler import ScheduleConnectionUnavailableError -from ...tools.nested_payload import validate_nested_payload +from ...tools.nested_payload import ValidatedNestedPayload, validate_nested_payload log = get_logger("discord") @@ -102,10 +102,9 @@ def _extract_tool_input_from_steps(inp: dict) -> dict | None: async def _handle_schedule_task(self, message, inp: dict) -> str: """Create a scheduled task.""" - try: - inp = validate_nested_payload("schedule_task", inp, self._nested_catalog()) - except ValueError as e: - return f"Failed to create schedule: {e}" + # Codex inputs arrive already decoded and checked by RequestToolAdapter. + # Other providers keep the historical schedule input behavior. + nested_validated = isinstance(inp, ValidatedNestedPayload) validation_error = self._validate_schedule_payload(inp) if validation_error: return f"Failed to create schedule: {validation_error}" @@ -124,7 +123,7 @@ async def _handle_schedule_task(self, message, inp: dict) -> str: cron_timezone=inp.get("cron_timezone"), requester_id=str(message.author.id), report_format=inp.get("report_format"), - nested_payload_validated=True, + **({"nested_payload_validated": True} if nested_validated else {}), ) if schedule.get("trigger"): trigger_desc = ", ".join(f"{k}={v}" for k, v in schedule["trigger"].items()) @@ -171,15 +170,19 @@ def _handle_list_schedules(self) -> str: async def _handle_update_schedule(self, inp: dict) -> str: """Update an existing schedule.""" - try: - inp = validate_nested_payload("update_schedule", inp, self._nested_catalog()) - if isinstance(inp.get("tool_input"), dict) and not inp.get("tool_name"): - current = next( - (s for s in self.scheduler.list_all() if s.get("id") == inp.get("schedule_id")), - None, - ) - target_name = (current or {}).get("tool_name") - if target_name: + nested_validated = isinstance(inp, ValidatedNestedPayload) + if ( + nested_validated + and isinstance(inp.get("tool_input"), dict) + and not inp.get("tool_name") + ): + current = next( + (s for s in self.scheduler.list_all() if s.get("id") == inp.get("schedule_id")), + None, + ) + target_name = (current or {}).get("tool_name") + if target_name: + try: validate_nested_payload( "schedule_task", { @@ -189,8 +192,8 @@ async def _handle_update_schedule(self, inp: dict) -> str: }, self._nested_catalog(), ) - except ValueError as e: - return f"Error: {e}" + except ValueError as e: + return f"Error: {e}" schedule_id = inp.get("schedule_id", "") if not schedule_id: return "Error: 'schedule_id' is required." @@ -212,7 +215,6 @@ async def _handle_update_schedule(self, inp: dict) -> str: trigger = inp.get("trigger") if trigger is not None: kwargs["trigger"] = trigger - kwargs["nested_payload_validated"] = True if "paused" in inp: val = inp["paused"] if not isinstance(val, bool): @@ -220,6 +222,10 @@ async def _handle_update_schedule(self, inp: dict) -> str: kwargs["paused"] = val if not kwargs: return "Error: no fields to update." + # Updating a legacy schedule's description/format must not retroactively + # mark its old, never-validated workflow steps as adapter-validated. + if nested_validated and ("steps" in kwargs or "tool_input" in kwargs): + kwargs["nested_payload_validated"] = True try: result = await self.scheduler.update(schedule_id, **kwargs) except ScheduleConnectionUnavailableError as e: diff --git a/src/discord/native_tools/skills_tools.py b/src/discord/native_tools/skills_tools.py index d97ab57dd..38ffbd9a7 100644 --- a/src/discord/native_tools/skills_tools.py +++ b/src/discord/native_tools/skills_tools.py @@ -24,7 +24,6 @@ import discord from ...odin_log import get_logger -from ...tools.nested_payload import validate_nested_payload from ..response_guards import scrub_response_secrets if TYPE_CHECKING: @@ -145,12 +144,8 @@ async def dispatch( "Use list_skills to see available skills.", effects, ) - try: - tool_input = validate_nested_payload( - "invoke_skill", tool_input, self.tool_catalog.merged_definitions() - ) - except ValueError as exc: - return f"Error: {exc}", effects + # Strict Codex requests are decoded and validated before dispatch. + # Preserve the legacy non-Codex object/error contract here. skill_input = tool_input.get("input") or {} if not isinstance(skill_input, dict): return "Error: invoke_skill 'input' must be an object.", effects diff --git a/src/discord/scheduled_events.py b/src/discord/scheduled_events.py index 73dc6a8b9..b8ff685e2 100644 --- a/src/discord/scheduled_events.py +++ b/src/discord/scheduled_events.py @@ -332,10 +332,21 @@ async def _run_scheduled_workflow( workflow_ok = False break denial = self._tool_executor.check_permission(tool_name, req_id) - if denial: + if isinstance(denial, str) and denial: results.append(f"**Step {i + 1}**: {denial}") workflow_ok = False break + if tool_name == "invoke_skill": + target = tool_input.get("name") + if not isinstance(target, str) or not target: + results.append(f"**Step {i + 1}**: invoke_skill requires a skill name") + workflow_ok = False + break + denial = self._tool_executor.check_permission(target, req_id) + if isinstance(denial, str) and denial: + results.append(f"**Step {i + 1}**: {denial}") + workflow_ok = False + break condition = step.get("condition") step_desc = step.get("description", tool_name) @@ -504,6 +515,18 @@ async def _on_scheduled_task(self, schedule: dict) -> None: tool_input = schedule.get("tool_input", {}) req_id = schedule.get("requester_id") or None req_name = schedule.get("requester") or schedule.get("created_by") or "scheduler" + if schedule.get("_nested_payload_validated"): + # The selected tool may have changed since the request was + # admitted. Recheck the persisted canonical input at fire time. + try: + validate_nested_payload( + "schedule_task", + {"action": "check", "tool_name": tool_name, "tool_input": tool_input}, + self._tool_loop._tool_catalog.merged_definitions(), + allow_placeholders=False, + ) + except ValueError as exc: + raise ValueError(f"Invalid scheduled check payload: {exc}") from exc try: result = await self._execute_scheduled_tool( tool_name, # type: ignore[arg-type] # creation-time validation guarantees tool_name diff --git a/src/scheduler/scheduler.py b/src/scheduler/scheduler.py index f8c1d8e5f..26e9927aa 100644 --- a/src/scheduler/scheduler.py +++ b/src/scheduler/scheduler.py @@ -972,7 +972,11 @@ async def update( # persist them. Commit only after every supplied field is valid. original = self._schedules[target_index] target = copy.deepcopy(original) - if nested_payload_validated: + # Certify only replaced nested data, not unrelated legacy steps. + if nested_payload_validated and ( + (target["action"] == "check" and tool_input is not None) + or (target["action"] == "workflow" and steps is not None) + ): target["_nested_payload_validated"] = True action = target["action"] if steps is not None and action == "workflow": diff --git a/src/tools/nested_payload.py b/src/tools/nested_payload.py index cc772c961..d42ca81a6 100644 --- a/src/tools/nested_payload.py +++ b/src/tools/nested_payload.py @@ -17,6 +17,14 @@ _PLACEHOLDER = re.compile(r"\{(?:var\.[^{}]+|prev_output)\}") +class ValidatedNestedPayload(dict): + """Request-adapter-validated arguments, without a synthetic public field. + + Deferred handlers persist this provenance to validate concrete inputs at + execution time. Legacy provider dicts are deliberately unmarked. + """ + + def _object_no_duplicates(pairs): result = {} for key, value in pairs: @@ -161,4 +169,4 @@ def check(target, payload, label): target = args.get("name") if target and isinstance(args.get("input"), dict): check(target, args["input"], "input") - return args + return ValidatedNestedPayload(args) diff --git a/tests/characterization/test_tool_parity.py b/tests/characterization/test_tool_parity.py index cc1c01241..487e57f0b 100644 --- a/tests/characterization/test_tool_parity.py +++ b/tests/characterization/test_tool_parity.py @@ -62,11 +62,11 @@ "purge_messages": "db35efc321c205b1", "post_file": "6860faab30251338", "generate_file": "2f4687a63e985fdd", - # Updated 2026-09-18: removed unimplemented Discord event trigger sources. - "schedule_task": "68b99e031c52c71a", + # Updated for the intentionally closed four-field trigger shape (B2). + "schedule_task": "6571dd243ee3f137", "list_schedules": "6f72cb95cee9eb6c", - # Updated with schedule_task: report_format may be changed or cleared. - "update_schedule": "4635df8029e5e548", + # Updated with schedule_task: shared trigger shape, report_format clearing (B2). + "update_schedule": "ba5b0ca8c5b1f96a", "delete_schedule": "01e54d37b70471a8", "parse_time": "6ae3f4c04138a2cd", "search_history": "72aaa6b1024b0fc0", @@ -118,9 +118,9 @@ "kill_agent": "2543a3eeb5720fdf", "get_agent_results": "4742878b8c825633", "wait_for_agents": "c6c21343f9b82b90", - "http_probe": "dfc3b04b36c5e7f9", + "http_probe": "96d48f83a43142a8", # B1: wire header records, not arbitrary-key object "generate_image": "e9347378f7e4ccbb", - "validate_action": "ebe7843d7c4125c8", # raw affordance wording moved to generated footer + "validate_action": "019b8720735849c8", # B3: typed expected union "email_send": "1282279440e34e6f", "email_search": "3a7584b725d1c134", "email_read": "c88d947b915f9cf0", @@ -156,10 +156,7 @@ def test_schema_deep_equality(self): t["name"] for t in TOOLS if _canonical_hash(t) != EXPECTED_TOOL_HASHES[t["name"]] ] - assert not changed, ( - f"tool definitions changed content: {changed} — " - "schema/description edits are out of scope for the carve" - ) + assert not changed, f"tool definitions changed content: {changed} — review schema drift" def test_tool_map_completeness_and_identity(self): assert set(TOOL_MAP) == set(EXPECTED_TOOL_ORDER) diff --git a/tests/test_computer_tool_schema_r5.py b/tests/test_computer_tool_schema_r5.py index 45f570d46..135b4ae7a 100644 --- a/tests/test_computer_tool_schema_r5.py +++ b/tests/test_computer_tool_schema_r5.py @@ -3,25 +3,33 @@ import pytest from src.llm.openai_codex import CodexChatClient +from src.llm.strict_tool_adapter import compile_catalog from src.tools.defs.computer import computer_definitions -def test_computer_schema_explicitly_disables_transport_strict_normalization(): +def test_computer_schema_uses_closed_strict_wrapper_without_changing_canonical_contract(): tools = computer_definitions() converted = CodexChatClient._convert_tools(tools) assert len(converted) == 3 - assert all(tool["strict"] is False for tool in converted) - assert all(tool["parameters"] == source["input_schema"] - for tool, source in zip(converted, tools, strict=True)) + assert all(tool["strict"] is True for tool in converted) + action = next(tool for tool in converted if tool["name"] == "computer_act") + assert "anyOf" not in action["parameters"] + assert action["parameters"]["additionalProperties"] is False + assert len(action["parameters"]["properties"]["payload"]["anyOf"]) == 14 + assert tools == computer_definitions() # request-local lowering only -def test_ordinary_tool_conversion_defaults_non_strict_to_preserve_optional_fields(): +def test_external_tool_open_schema_uses_envelope_without_claiming_strict_resolution(): tool = {"name": "ordinary", "description": "ordinary tool", "input_schema": { "type": "object", "properties": {"value": {"type": "string"}}}} - assert CodexChatClient._convert_tools([tool]) == [{ - "type": "function", "name": "ordinary", "description": "ordinary tool", - "parameters": tool["input_schema"], "strict": False}] - assert CodexChatClient._convert_tools([{**tool, "strict": "false"}])[0]["strict"] is False + converted = CodexChatClient._convert_tools([tool])[0] + assert "strict" not in converted # server resolves external strictness + assert converted["parameters"]["required"] == ["json"] + assert converted["parameters"]["additionalProperties"] is False + assert '"value"' in converted["parameters"]["properties"]["json"]["description"] + assert compile_catalog([tool]).accept("ordinary", {"json": '{"value":"ok","other":1}'}) == { + "value": "ok", "other": 1, + } def test_public_key_vocabulary_is_executable_by_controller_and_private_backend(): diff --git a/tests/test_native_agents_tasks.py b/tests/test_native_agents_tasks.py index 9604ceee5..b95a76570 100644 --- a/tests/test_native_agents_tasks.py +++ b/tests/test_native_agents_tasks.py @@ -37,6 +37,7 @@ ) from src.llm.model_breaker import ModelBreakerRegistry from src.llm.recovery import RecoveryPolicy +from src.tools.nested_payload import ValidatedNestedPayload def _fake_gateway(client): @@ -116,6 +117,29 @@ def _task(status="running", results=None, steps=None, tid="T1", desc="job"): # delegate_task # --------------------------------------------------------------------------- # class TestDelegateTask: + async def test_strict_provenance_checks_selected_tool_before_launch(self): + executor = MagicMock() + executor.check_permission.return_value = "Permission denied" + t = _tools(tool_executor=executor) + out = await t._handle_delegate_task(_message(), ValidatedNestedPayload({ + "steps": [{"tool_name": "web_search", "tool_input": {"query": "test"}}], + })) + assert out == "Step 1: Permission denied" + assert not t._channel_state.background_tasks + + async def test_legacy_payload_retains_deferred_authorization(self): + executor = MagicMock() + executor.check_permission.return_value = "Permission denied" + t = _tools(tool_executor=executor) + with patch("src.discord.native_tools.agents_tasks.run_background_task", new=AsyncMock()): + out = await t._handle_delegate_task(_message(), { + "steps": [{"tool_name": "web_search"}], + }) + await asyncio.sleep(0) + assert "Background task started" in out + task = next(iter(t._channel_state.background_tasks.values())) + assert task.nested_payload_validated is False + async def test_validation(self): t = _tools() msg = _message() diff --git a/tests/test_native_scheduling.py b/tests/test_native_scheduling.py index 1174e9e85..02f79f3a0 100644 --- a/tests/test_native_scheduling.py +++ b/tests/test_native_scheduling.py @@ -11,6 +11,7 @@ from unittest.mock import AsyncMock, MagicMock, patch from src.discord.native_tools.scheduling import SchedulingTools +from src.tools.nested_payload import ValidatedNestedPayload def _tools(scheduler=None): @@ -70,6 +71,15 @@ def test_extract_from_steps_edge_cases(self): class TestScheduleTask: + async def test_adapter_validated_schedule_marks_only_codex_payloads(self): + sched = MagicMock() + sched.add = AsyncMock(return_value={"id": "S1", "description": "d"}) + payload = {"action": "reminder", "message": "m"} + await _tools(sched)._handle_schedule_task(_message(), payload) + assert "nested_payload_validated" not in sched.add.await_args.kwargs + await _tools(sched)._handle_schedule_task(_message(), ValidatedNestedPayload(payload)) + assert sched.add.await_args.kwargs["nested_payload_validated"] is True + async def test_validation_error(self): out = await _tools()._handle_schedule_task(_message(), {"action": "reminder"}) assert "Failed to create schedule" in out @@ -126,6 +136,20 @@ def test_formats_all_types(self): class TestUpdateSchedule: + async def test_legacy_schedule_update_not_retroactively_validated(self): + sched = MagicMock() + sched.update = AsyncMock(return_value={"id": "S1"}) + t = _tools(sched) + await t._handle_update_schedule(ValidatedNestedPayload({ + "schedule_id": "S1", "description": "new label", + })) + sched.update.assert_awaited_once_with("S1", description="new label") + sched.update.reset_mock() + await t._handle_update_schedule(ValidatedNestedPayload({ + "schedule_id": "S1", "steps": [{"tool_name": "web_search", "tool_input": {}}], + })) + assert sched.update.await_args.kwargs["nested_payload_validated"] is True + async def test_requires_id_and_fields(self): assert "'schedule_id' is required" in await _tools()._handle_update_schedule({}) assert "no fields to update" in await _tools()._handle_update_schedule( diff --git a/tests/test_provider_stream_acceptance.py b/tests/test_provider_stream_acceptance.py index e5d449c6d..99c9065be 100644 --- a/tests/test_provider_stream_acceptance.py +++ b/tests/test_provider_stream_acceptance.py @@ -6,6 +6,7 @@ from src.agents.manager import _run_agent from src.llm.openai_codex import CodexStreamError +from src.tools.registry import get_tool_definitions from src.tools.result_validator import ToolResult from tests.characterization.test_autonomous_loop import build as build_loop from tests.characterization.test_autonomous_loop import run_iteration @@ -22,14 +23,16 @@ def events(stage): output += [ {"type": "response.output_item.added", "output_index": 0, "item": {"type": "function_call", "call_id": "once", "name": "run_command"}}, - {"type": "response.function_call_arguments.delta", "output_index": 0, "delta": "{}"}, + {"type": "response.function_call_arguments.delta", "output_index": 0, + "delta": '{"host":null,"command":"echo ok"}'}, ] if stage == "delta": return output output.append({"type": "response.function_call_arguments.done", "output_index": 0}) if stage == "item": output.append({"type": "response.output_item.done", "output_index": 0, - "item": {"type": "function_call", "arguments": "{}"}}) + "item": {"type": "function_call", + "arguments": '{"host":null,"command":"echo ok"}'}}) return output @@ -87,7 +90,9 @@ async def callback(messages, system, tools, **kwargs): result = await client.chat_with_tools(messages, system, tools) return {"text": result.text, "tool_calls": result.tool_calls} - await _run_agent(agent(), "", [], callback, effect, max_iterations=3) + await _run_agent(agent(), "", [next(tool for tool in get_tool_definitions() + if tool["name"] == "run_command")], callback, effect, + max_iterations=3) else: bot, _ = (build if entry == "chat" else build_loop)([]) bot.llm_gateway.codex_client = client From 7987dc1b67c6d93bcf7771fd0d0732c1127211c5 Mon Sep 17 00:00:00 2001 From: Odin Date: Sun, 27 Sep 2026 21:11:52 -0400 Subject: [PATCH 05/12] Cover strict payload admission and selected tool checks --- ...test_background_task_failure_visibility.py | 81 ++++-- tests/test_http_probe_ops.py | 31 +++ tests/test_native_agents_tasks.py | 22 ++ tests/test_native_scheduling.py | 237 +++++++++++++----- tests/test_post_validation.py | 36 +++ tests/test_scheduled_events.py | 45 ++++ tests/test_scheduler.py | 9 + 7 files changed, 383 insertions(+), 78 deletions(-) diff --git a/tests/test_background_task_failure_visibility.py b/tests/test_background_task_failure_visibility.py index b55ba2cba..9aa63a3ff 100644 --- a/tests/test_background_task_failure_visibility.py +++ b/tests/test_background_task_failure_visibility.py @@ -20,6 +20,7 @@ from src.discord.background_task import ( BackgroundTask, + _execute_tool_captured, _is_error_output, create_task_id, run_background_task, @@ -69,6 +70,52 @@ async def run(task, executor): class TestStructuredFailureVisibility: + async def test_nested_validated_concrete_payload_is_rechecked_before_execution(self): + + executor = _FakeExecutor([]) + task = make_task([{"tool_name": "run_command", "tool_input": {"command": "x"}}]) + task.nested_payload_validated = True + # A selected tool's required field was removed after adapter admission. + task.steps[0]["tool_input"] = {} + catalog = MagicMock() + catalog.merged_definitions.return_value = [ + { + "name": "run_command", + "input_schema": { + "type": "object", + "required": ["command"], + "properties": {"command": {"type": "string"}}, + }, + } + ] + executor._tool_catalog = catalog + await run_background_task(task, executor, _FakeSkillManager()) + assert task.status == "failed" + assert "Invalid concrete step payload" in task.results[0].output + assert executor.calls == [] + + async def test_disabled_builtin_is_rejected_before_executor_effect(self): + from types import SimpleNamespace + + from src.tools.builtin_policy import BuiltinToolPolicy + + executor = _FakeExecutor([]) + executor._builtin_policy = BuiltinToolPolicy( + lambda: SimpleNamespace(tools=SimpleNamespace(disabled_tools=["run_command"])) + ) + result = await _execute_tool_captured( + "run_command", + {"command": "must not run"}, + executor, + _FakeSkillManager(), + None, + None, + "tester", + requester_id="4242", + ) + assert "disabled" in str(result).lower() + assert executor.calls == [] + @pytest.mark.parametrize( ("outcome", "expected"), [ @@ -78,24 +125,31 @@ class TestStructuredFailureVisibility: ], ) async def test_dedup_outcomes_succeed_and_do_not_abort_workflow( - self, outcome, expected, + self, + outcome, + expected, ): store = MagicMock() store.ingest = AsyncMock(return_value=outcome) executor = _FakeExecutor([]) - task = make_task([ - { - "tool_name": "ingest_document", - "tool_input": {"source": "doc.md", "content": "document body"}, - }, - { - "tool_name": "ingest_document", - "tool_input": {"source": "next.md", "content": "next document body"}, - }, - ]) + task = make_task( + [ + { + "tool_name": "ingest_document", + "tool_input": {"source": "doc.md", "content": "document body"}, + }, + { + "tool_name": "ingest_document", + "tool_input": {"source": "next.md", "content": "next document body"}, + }, + ] + ) await run_background_task( - task, executor, _FakeSkillManager(), knowledge_store=store, + task, + executor, + _FakeSkillManager(), + knowledge_store=store, embedder=object(), ) @@ -234,8 +288,6 @@ async def test_audit_error_field_populated_for_ok_false(self): kwargs = audit.log_execution.await_args.kwargs assert kwargs.get("error"), "audit entry must carry the error field for a failed step" - - async def test_mcp_audit_metadata_survives_background_path(self): metadata = { "mcp_server": "srv", @@ -271,7 +323,6 @@ async def test_mcp_audit_metadata_survives_background_path(self): assert "opaque-background-secret" not in str(kwargs["tool_input"]) assert kwargs["tool_input"]["password"] == "[redacted:sensitive-key]" - async def test_ingest_durability_failure_fails_real_background_step(self, tmp_path): fts = FullTextIndex(str(tmp_path / "fts.db")) store = KnowledgeStore(str(tmp_path / "knowledge.db"), fts_index=fts) diff --git a/tests/test_http_probe_ops.py b/tests/test_http_probe_ops.py index 5b1ba6a79..f17ab587d 100644 --- a/tests/test_http_probe_ops.py +++ b/tests/test_http_probe_ops.py @@ -839,6 +839,37 @@ async def test_command_failure_no_output(self, executor): assert "exit 1" in message assert code == 1 + @pytest.mark.asyncio + async def test_headers_must_be_validated_before_host_acquisition(self, executor): + """Malformed wire headers fail closed without acquiring/using a host.""" + from unittest.mock import Mock + + executor.browser_web_tools._acquire_host = Mock( + side_effect=AssertionError("host must not be acquired") + ) + message, code = await executor.browser_web_tools._handle_http_probe({ + "url": "https://example.com", + "host": "web", + "headers": [{"name": "X-Test", "value": "one"}, + {"name": "x-test", "value": "two"}], + }) + assert "http_probe error" in message + assert code != 0 + executor._exec_command.assert_not_called() + + @pytest.mark.asyncio + async def test_valid_wire_headers_are_normalized_for_execution(self, executor): + await executor.browser_web_tools._handle_http_probe({ + "url": "https://example.com", + "headers": [ + {"name": "X-Request-Id", "value": "probe-123"}, + {"name": "Accept", "value": "text/plain"}, + ], + }) + command = executor._exec_command.call_args.args[1] + assert "X-Request-Id: probe-123" in command + assert "Accept: text/plain" in command + @pytest.mark.asyncio async def test_empty_success(self, executor): executor._exec_command.return_value = (0, "") diff --git a/tests/test_native_agents_tasks.py b/tests/test_native_agents_tasks.py index b95a76570..1a7138510 100644 --- a/tests/test_native_agents_tasks.py +++ b/tests/test_native_agents_tasks.py @@ -127,6 +127,28 @@ async def test_strict_provenance_checks_selected_tool_before_launch(self): assert out == "Step 1: Permission denied" assert not t._channel_state.background_tasks + async def test_strict_invoke_skill_missing_name_denied_before_launch(self): + executor = MagicMock() + t = _tools(tool_executor=executor) + out = await t._handle_delegate_task(_message(), ValidatedNestedPayload({ + "steps": [{"tool_name": "invoke_skill", "tool_input": {}}], + })) + assert out == "Step 1: invoke_skill requires a selected skill name" + executor.check_permission.assert_called_once_with("invoke_skill", "7") + assert not t._channel_state.background_tasks + + async def test_strict_invoke_skill_target_permission_checked_before_launch(self): + executor = MagicMock() + executor.check_permission.side_effect = ["", "skill denied"] + t = _tools(tool_executor=executor) + out = await t._handle_delegate_task(_message(), ValidatedNestedPayload({ + "steps": [{"tool_name": "invoke_skill", "tool_input": {"name": "private"}}], + })) + assert out == "Step 1: skill denied" + assert [c.args for c in executor.check_permission.call_args_list] == [ + ("invoke_skill", "7"), ("private", "7")] + assert not t._channel_state.background_tasks + async def test_legacy_payload_retains_deferred_authorization(self): executor = MagicMock() executor.check_permission.return_value = "Permission denied" diff --git a/tests/test_native_scheduling.py b/tests/test_native_scheduling.py index 02f79f3a0..938feb271 100644 --- a/tests/test_native_scheduling.py +++ b/tests/test_native_scheduling.py @@ -4,6 +4,7 @@ the schedule CRUD handlers on SchedulingTools with a faked scheduler. parse_time is patched for determinism. """ + from __future__ import annotations from types import SimpleNamespace @@ -19,8 +20,7 @@ def _tools(scheduler=None): def _message(): - return SimpleNamespace( - channel=SimpleNamespace(id=42), author=SimpleNamespace(id=7)) + return SimpleNamespace(channel=SimpleNamespace(id=42), author=SimpleNamespace(id=7)) class TestValidatePayload: @@ -41,12 +41,16 @@ def test_check_tool_name_and_command_shortcut(self): def test_check_missing_tool_input(self): t = _tools() assert "requires 'tool_input'" in t._validate_schedule_payload( - {"action": "check", "tool_name": "some_tool"}) + {"action": "check", "tool_name": "some_tool"} + ) def test_check_tool_input_from_single_step(self): t = _tools() - inp = {"action": "check", "tool_name": "run_command", - "steps": [{"tool_input": {"command": "ls"}}]} + inp = { + "action": "check", + "tool_name": "run_command", + "steps": [{"tool_input": {"command": "ls"}}], + } assert t._validate_schedule_payload(inp) is None assert inp["tool_input"] == {"command": "ls"} @@ -54,20 +58,33 @@ def test_workflow_validation(self): t = _tools() assert "non-empty 'steps'" in t._validate_schedule_payload({"action": "workflow"}) assert "must be an object" in t._validate_schedule_payload( - {"action": "workflow", "steps": ["notadict"]}) + {"action": "workflow", "steps": ["notadict"]} + ) assert "missing 'tool_name'" in t._validate_schedule_payload( - {"action": "workflow", "steps": [{}]}) + {"action": "workflow", "steps": [{}]} + ) assert "non-empty" in t._validate_schedule_payload( - {"action": "workflow", "steps": [{"tool_name": "run_command", "tool_input": {}}]}) - assert t._validate_schedule_payload( - {"action": "workflow", - "steps": [{"tool_name": "run_command", "tool_input": {"command": "ls"}}]}) is None + {"action": "workflow", "steps": [{"tool_name": "run_command", "tool_input": {}}]} + ) + assert ( + t._validate_schedule_payload( + { + "action": "workflow", + "steps": [{"tool_name": "run_command", "tool_input": {"command": "ls"}}], + } + ) + is None + ) def test_extract_from_steps_edge_cases(self): t = _tools() assert t._extract_tool_input_from_steps({"steps": None}) is None - assert t._extract_tool_input_from_steps( # two populated → ambiguous → None - {"steps": [{"tool_input": {"a": 1}}, {"tool_input": {"b": 2}}]}) is None + assert ( + t._extract_tool_input_from_steps( # two populated → ambiguous → None + {"steps": [{"tool_input": {"a": 1}}, {"tool_input": {"b": 2}}]} + ) + is None + ) class TestScheduleTask: @@ -86,32 +103,39 @@ async def test_validation_error(self): async def test_success_variants(self): sched = MagicMock() - sched.add = AsyncMock(return_value={ - "id": "S1", "description": "d", "trigger": {"webhook": "x"}}) + sched.add = AsyncMock( + return_value={"id": "S1", "description": "d", "trigger": {"webhook": "x"}} + ) out = await _tools(sched)._handle_schedule_task( - _message(), {"action": "reminder", "message": "m", "trigger": {"webhook": "x"}}) + _message(), {"action": "reminder", "message": "m", "trigger": {"webhook": "x"}} + ) assert "webhook-triggered" in out and "S1" in out - sched.add = AsyncMock(return_value={ - "id": "S2", "description": "d", "cron": "* * * * *", "next_run": "soon"}) + sched.add = AsyncMock( + return_value={"id": "S2", "description": "d", "cron": "* * * * *", "next_run": "soon"} + ) out = await _tools(sched)._handle_schedule_task( - _message(), {"action": "reminder", "message": "m"}) + _message(), {"action": "reminder", "message": "m"} + ) assert "recurring" in out and "Next run: soon" in out sched.add = AsyncMock(return_value={"id": "S3", "description": "d"}) out = await _tools(sched)._handle_schedule_task( - _message(), {"action": "reminder", "message": "m"}) + _message(), {"action": "reminder", "message": "m"} + ) assert "one-time" in out async def test_value_error_and_generic(self): sched = MagicMock() sched.add = AsyncMock(side_effect=ValueError("bad cron")) out = await _tools(sched)._handle_schedule_task( - _message(), {"action": "reminder", "message": "m"}) + _message(), {"action": "reminder", "message": "m"} + ) assert "Failed to create schedule: bad cron" in out sched.add = AsyncMock(side_effect=RuntimeError("boom")) out = await _tools(sched)._handle_schedule_task( - _message(), {"action": "reminder", "message": "m"}) + _message(), {"action": "reminder", "message": "m"} + ) assert "Error creating schedule" in out @@ -126,8 +150,12 @@ def test_formats_all_types(self): sched.list_all.return_value = [ {"id": "A", "description": "trig", "trigger": {"webhook": "x"}}, {"id": "B", "description": "cronjob", "cron": "* * * * *", "paused": True}, - {"id": "C", "description": "once", "paused": True, - "inert_reason": "expired one-time schedule; set run_at"}, + { + "id": "C", + "description": "once", + "paused": True, + "inert_reason": "expired one-time schedule; set run_at", + }, ] out = _tools(sched)._handle_list_schedules() assert "3" in out and "trigger:" in out and "cron `* * * * *`" in out @@ -136,46 +164,106 @@ def test_formats_all_types(self): class TestUpdateSchedule: + async def test_strict_update_revalidates_replaced_input_against_persisted_tool(self): + sched = MagicMock() + sched.list_all.return_value = [{"id": "S1", "tool_name": "run_command"}] + sched.update = AsyncMock(return_value={"id": "S1"}) + catalog = SimpleNamespace( + merged_definitions=lambda: [ + { + "name": "run_command", + "input_schema": { + "type": "object", + "required": ["command"], + "properties": {"command": {"type": "string"}}, + }, + } + ] + ) + tools = SchedulingTools(scheduler=sched, tool_catalog=catalog) + result = await tools._handle_update_schedule( + ValidatedNestedPayload({"schedule_id": "S1", "tool_input": {}}) + ) + assert "invalid input" in result + sched.update.assert_not_awaited() + + async def test_strict_update_without_existing_selected_tool_remains_safe(self): + sched = MagicMock() + sched.list_all.return_value = [] + sched.update = AsyncMock(return_value=None) + result = await _tools(sched)._handle_update_schedule( + ValidatedNestedPayload( + { + "schedule_id": "missing", + "tool_input": {"opaque": True}, + } + ) + ) + assert "not found" in result + assert sched.update.await_args.kwargs["nested_payload_validated"] is True + async def test_legacy_schedule_update_not_retroactively_validated(self): sched = MagicMock() sched.update = AsyncMock(return_value={"id": "S1"}) t = _tools(sched) - await t._handle_update_schedule(ValidatedNestedPayload({ - "schedule_id": "S1", "description": "new label", - })) + await t._handle_update_schedule( + ValidatedNestedPayload( + { + "schedule_id": "S1", + "description": "new label", + } + ) + ) sched.update.assert_awaited_once_with("S1", description="new label") sched.update.reset_mock() - await t._handle_update_schedule(ValidatedNestedPayload({ - "schedule_id": "S1", "steps": [{"tool_name": "web_search", "tool_input": {}}], - })) + await t._handle_update_schedule( + ValidatedNestedPayload( + { + "schedule_id": "S1", + "steps": [{"tool_name": "web_search", "tool_input": {}}], + } + ) + ) assert sched.update.await_args.kwargs["nested_payload_validated"] is True async def test_requires_id_and_fields(self): assert "'schedule_id' is required" in await _tools()._handle_update_schedule({}) assert "no fields to update" in await _tools()._handle_update_schedule( - {"schedule_id": "S1"}) + {"schedule_id": "S1"} + ) async def test_paused_must_be_bool(self): assert "must be a boolean" in await _tools()._handle_update_schedule( - {"schedule_id": "S1", "paused": "yes"}) + {"schedule_id": "S1", "paused": "yes"} + ) async def test_value_error_not_found_and_success(self): sched = MagicMock() sched.update = AsyncMock(side_effect=ValueError("bad")) assert "Error: bad" in await _tools(sched)._handle_update_schedule( - {"schedule_id": "S1", "description": "d"}) + {"schedule_id": "S1", "description": "d"} + ) sched.update = AsyncMock(return_value=None) assert "not found" in await _tools(sched)._handle_update_schedule( - {"schedule_id": "S1", "paused": True}) + {"schedule_id": "S1", "paused": True} + ) sched.update = AsyncMock(return_value={"id": "S1"}) assert "Updated schedule S1" in await _tools(sched)._handle_update_schedule( - {"schedule_id": "S1", "cron": "* * * * *", "trigger": {"webhook": "x"}}) - sched.update = AsyncMock(return_value={ - "id": "S1", "paused": True, "inert_reason": "expired while paused", - }) - out = await _tools(sched)._handle_update_schedule({ - "schedule_id": "S1", "paused": False, - }) + {"schedule_id": "S1", "cron": "* * * * *", "trigger": {"webhook": "x"}} + ) + sched.update = AsyncMock( + return_value={ + "id": "S1", + "paused": True, + "inert_reason": "expired while paused", + } + ) + out = await _tools(sched)._handle_update_schedule( + { + "schedule_id": "S1", + "paused": False, + } + ) assert "remains paused and inert" in out assert "expired while paused" in out @@ -185,10 +273,10 @@ async def test_delete(self): sched = MagicMock() sched.delete = AsyncMock(return_value=True) assert "Deleted schedule S1" in await _tools(sched)._handle_delete_schedule( - {"schedule_id": "S1"}) + {"schedule_id": "S1"} + ) sched.delete = AsyncMock(return_value=False) - assert "not found" in await _tools(sched)._handle_delete_schedule( - {"schedule_id": "S1"}) + assert "not found" in await _tools(sched)._handle_delete_schedule({"schedule_id": "S1"}) def test_parse_time(self): t = _tools() @@ -202,16 +290,26 @@ def test_parse_time(self): class TestReportFormatNativeParity: async def test_create_passes_generic_format(self): scheduler = MagicMock() - scheduler.add = AsyncMock(return_value={ - "id": "S1", "description": "structured", "cron": "0 * * * *", - "next_run": "soon", "report_format": "paginated_embed_v1", - }) + scheduler.add = AsyncMock( + return_value={ + "id": "S1", + "description": "structured", + "cron": "0 * * * *", + "next_run": "soon", + "report_format": "paginated_embed_v1", + } + ) result = await _tools(scheduler)._handle_schedule_task( - _message(), { - "description": "structured", "action": "check", "cron": "0 * * * *", - "tool_name": "run_command", "tool_input": {"command": "status"}, + _message(), + { + "description": "structured", + "action": "check", + "cron": "0 * * * *", + "tool_name": "run_command", + "tool_input": {"command": "status"}, "report_format": "paginated_embed_v1", - }) + }, + ) assert "Scheduled recurring" in result assert scheduler.add.await_args.kwargs["report_format"] == "paginated_embed_v1" @@ -219,7 +317,8 @@ async def test_update_passes_and_can_clear_format(self): scheduler = MagicMock() scheduler.update = AsyncMock(return_value={"id": "S1"}) result = await _tools(scheduler)._handle_update_schedule( - {"schedule_id": "S1", "report_format": ""}) + {"schedule_id": "S1", "report_format": ""} + ) assert result == "Updated schedule S1." scheduler.update.assert_awaited_once_with("S1", report_format="") @@ -236,22 +335,34 @@ def _scheduler(tmp_path): async def test_native_add_rejects_unknown_format(self, tmp_path): scheduler = self._scheduler(tmp_path) result = await _tools(scheduler)._handle_schedule_task( - _message(), { - "description": "structured", "action": "check", "cron": "0 * * * *", - "tool_name": "run_command", "tool_input": {"command": "status"}, + _message(), + { + "description": "structured", + "action": "check", + "cron": "0 * * * *", + "tool_name": "run_command", + "tool_input": {"command": "status"}, "report_format": "paginated_embed_v2", - }) + }, + ) assert "Unsupported scheduled report format: paginated_embed_v2" in result assert scheduler.list_all() == [] async def test_native_update_rejects_unknown_format(self, tmp_path): scheduler = self._scheduler(tmp_path) created = await scheduler.add( - description="plain", action="check", channel_id="42", cron="0 * * * *", - tool_name="run_command", tool_input={"command": "status"}) - result = await _tools(scheduler)._handle_update_schedule({ - "schedule_id": created["id"], - "report_format": "paginated_embed_v2", - }) + description="plain", + action="check", + channel_id="42", + cron="0 * * * *", + tool_name="run_command", + tool_input={"command": "status"}, + ) + result = await _tools(scheduler)._handle_update_schedule( + { + "schedule_id": created["id"], + "report_format": "paginated_embed_v2", + } + ) assert "Unsupported scheduled report format: paginated_embed_v2" in result assert "report_format" not in scheduler.list_all()[0] diff --git a/tests/test_post_validation.py b/tests/test_post_validation.py index e44067761..dcbeaf4c7 100644 --- a/tests/test_post_validation.py +++ b/tests/test_post_validation.py @@ -94,6 +94,42 @@ def test_default_compare_does_not_require_expected(self): assert errors == [] assert len(checks) == 2 + def test_explicit_compare_rejects_omitted_expected_value(self): + """Explicit comparators never silently fall back to defaults.""" + _, errors = parse_checks([{ + "type": "command", "target": "printf ok", "compare": "contains", + }]) + assert any("requires a non-empty 'expected'" in error for error in errors) + + def test_expected_field_with_null_remains_invalid_without_compare(self): + """Presence of expected is distinct from omission, even for defaults.""" + checks, errors = parse_checks([{ + "type": "http", "target": "https://example.invalid", "expected": None, + }]) + assert checks == [] + assert any("expected must be an integer" in error for error in errors) + + def test_expected_values_are_validated_for_the_check_type(self): + checks, errors = parse_checks([ + {"type": "http", "target": "x", "expected": [200, "204"]}, + {"type": "service", "target": "nginx", "expected": ["active", "activating"]}, + ]) + assert errors == [] + assert [check.expected for check in checks] == [[200, "204"], ["active", "activating"]] + + @pytest.mark.parametrize( + ("check", "message"), + [ + ({"type": "http", "expected": "healthy"}, "http expected values"), + ({"type": "service", "expected": ["active", 2]}, "service expected lists"), + ({"type": "port", "expected": "open"}, "expected is not used for port"), + ({"type": "command", "expected": True}, "expected must be an integer"), + ], + ) + def test_invalid_expected_values_fail_closed(self, check, message): + _, errors = parse_checks([{**check, "target": check.get("target", "safe-target")}]) + assert any(message in error for error in errors) + def test_compare_exit_zero_ok_without_expected(self): """exit_zero doesn't need an expected value.""" checks, errors = parse_checks([ diff --git a/tests/test_scheduled_events.py b/tests/test_scheduled_events.py index 7287dddab..02685c4ff 100644 --- a/tests/test_scheduled_events.py +++ b/tests/test_scheduled_events.py @@ -235,6 +235,51 @@ async def test_dispatch_exception(self): class TestWorkflow: + async def test_strict_workflow_rejects_invalid_step_before_dispatch(self): + h = _handlers() + h._tool_loop._tool_catalog = SimpleNamespace(merged_definitions=lambda: [{ + "name": "run_command", + "input_schema": {"type": "object", "required": ["command"], + "properties": {"command": {"type": "string"}}}, + }]) + result = await h._run_scheduled_workflow(_channel(), { + "description": "strict", "_nested_payload_validated": True, + "steps": [{"tool_name": "run_command", "tool_input": {}}], + }) + assert result is False + h._tool_executor.check_permission.assert_not_called() + h._tool_loop.dispatch_loop_tool_inner.assert_not_awaited() + + async def test_strict_workflow_checks_skill_target_permission(self): + executor = MagicMock() + executor.check_permission.side_effect = ["", "skill denied"] + h = _handlers(tool_executor=executor) + h._tool_loop._tool_catalog = SimpleNamespace(merged_definitions=lambda: [{ + "name": "invoke_skill", "input_schema": {"type": "object"}, + }]) + result = await h._run_scheduled_workflow(_channel(), { + "description": "strict", "requester_id": "u", "_nested_payload_validated": True, + "steps": [{"tool_name": "invoke_skill", "tool_input": {"name": "private"}}], + }) + assert result is False + assert [c.args for c in executor.check_permission.call_args_list] == [ + ("invoke_skill", "u"), ("private", "u")] + h._tool_loop.dispatch_loop_tool_inner.assert_not_awaited() + + async def test_strict_check_revalidates_persisted_input_before_dispatch(self): + h = _handlers() + h._tool_loop._tool_catalog = SimpleNamespace(merged_definitions=lambda: [{ + "name": "run_command", + "input_schema": {"type": "object", "required": ["command"], + "properties": {"command": {"type": "string"}}}, + }]) + with pytest.raises(ValueError, match="Invalid scheduled check payload"): + await h._on_scheduled_task({ + "id": "S1", "description": "strict", "channel_id": "1", "action": "check", + "tool_name": "run_command", "tool_input": {}, "_nested_payload_validated": True, + }) + h._tool_loop.dispatch_loop_tool_inner.assert_not_awaited() + async def test_uncertain_mcp_workflow_requires_manual_resolution(self): ch = _channel() uncertain = ToolResult( diff --git a/tests/test_scheduler.py b/tests/test_scheduler.py index 3219e63df..5bdcea2e2 100644 --- a/tests/test_scheduler.py +++ b/tests/test_scheduler.py @@ -627,6 +627,15 @@ async def test_update_tool_input(self, tmp_path): updated = await s.update(sched["id"], tool_input={"command": "free -m", "host": "server1"}) assert updated["tool_input"]["command"] == "free -m" + async def test_strict_metadata_update_does_not_certify_legacy_nested_payload(self, tmp_path): + s = _make_scheduler(tmp_path) + legacy_steps = [{"tool_name": "web_search", "tool_input": {"query": "old"}}] + sched = await s.add("legacy", "workflow", "chan1", cron="0 * * * *", steps=legacy_steps) + assert sched["_nested_payload_validated"] is False + updated = await s.update(sched["id"], description="renamed", nested_payload_validated=True) + assert updated["_nested_payload_validated"] is False + assert updated["steps"] == legacy_steps + async def test_update_no_fields_still_persists(self, tmp_path): """Calling update with no changed fields returns the schedule unchanged.""" s = _make_scheduler(tmp_path) From 44c004a0137aba65ff67bf80937f25d13cbed1a5 Mon Sep 17 00:00:00 2001 From: Odin Date: Sun, 27 Sep 2026 21:41:21 -0400 Subject: [PATCH 06/12] Cover strict background and scheduled payload admission paths --- ...test_background_task_failure_visibility.py | 92 +++++++++++++++++++ tests/test_native_scheduling.py | 11 +++ tests/test_scheduled_events.py | 37 ++++++++ tests/test_scheduler.py | 31 +++++++ 4 files changed, 171 insertions(+) diff --git a/tests/test_background_task_failure_visibility.py b/tests/test_background_task_failure_visibility.py index 9aa63a3ff..670a4d6ec 100644 --- a/tests/test_background_task_failure_visibility.py +++ b/tests/test_background_task_failure_visibility.py @@ -22,6 +22,7 @@ BackgroundTask, _execute_tool_captured, _is_error_output, + _send_progress, create_task_id, run_background_task, ) @@ -70,6 +71,14 @@ async def run(task, executor): class TestStructuredFailureVisibility: + async def test_empty_workflow_progress_reports_no_steps(self): + channel = FakeChannel(id=555) + task = make_task([], channel=channel) + + await _send_progress(task, None) + + assert "No steps" in channel.sent_texts[0] + async def test_nested_validated_concrete_payload_is_rechecked_before_execution(self): executor = _FakeExecutor([]) @@ -94,6 +103,89 @@ async def test_nested_validated_concrete_payload_is_rechecked_before_execution(s assert "Invalid concrete step payload" in task.results[0].output assert executor.calls == [] + async def test_nested_validated_tool_permission_is_rechecked_before_execution(self): + executor = _FakeExecutor([]) + executor.check_permission = MagicMock(return_value="Permission denied: run_command") + executor._tool_catalog = MagicMock() + executor._tool_catalog.merged_definitions.return_value = [ + { + "name": "run_command", + "input_schema": { + "type": "object", + "required": ["command"], + "properties": {"command": {"type": "string"}}, + }, + } + ] + task = make_task( + [{"tool_name": "run_command", "tool_input": {"command": "echo unsafe"}}] + ) + task.nested_payload_validated = True + + await run_background_task(task, executor, _FakeSkillManager()) + + assert task.status == "failed" + assert "Invalid concrete step payload: Permission denied" in task.results[0].output + executor.check_permission.assert_called_once_with("run_command", "4242") + assert executor.calls == [] + + async def test_nested_invoke_skill_permission_is_checked_for_selected_skill(self): + executor = _FakeExecutor([]) + executor.check_permission = MagicMock( + side_effect=[None, "Permission denied: selected skill"] + ) + executor._tool_catalog = MagicMock() + executor._tool_catalog.merged_definitions.return_value = [] + skill_manager = MagicMock() + skill_manager.execute = AsyncMock() + skill_manager.has_skill.return_value = True + task = make_task( + [{"tool_name": "invoke_skill", "tool_input": {"name": "selected"}}] + ) + task.nested_payload_validated = True + + await run_background_task(task, executor, skill_manager) + + assert task.status == "failed" + assert "Invalid concrete step payload: Permission denied" in task.results[0].output + assert executor.check_permission.call_args_list == [ + (("invoke_skill", "4242"),), + (("selected", "4242"),), + ] + skill_manager.execute.assert_not_awaited() + + async def test_workflow_substitutes_outputs_and_skips_unmatched_condition(self): + executor = _FakeExecutor(["ready", "second result"]) + task = make_task( + [ + { + "tool_name": "run_command", + "tool_input": {"command": "first"}, + "store_as": "first_result", + }, + { + "tool_name": "run_command", + "tool_input": { + "command": "echo {var.first_result} then {prev_output}" + }, + "condition": "not-present", + }, + ] + ) + + await run_background_task(task, executor, _FakeSkillManager()) + + assert task.status == "completed" + assert [(r.status, r.output) for r in task.results] == [ + ("ok", "ready"), + ("skipped", "Condition not met: not-present"), + ] + # The condition is evaluated after placeholder expansion; no second tool call. + assert executor.calls == [ + ("run_command", {"command": "first"}, "4242"), + ] + + async def test_disabled_builtin_is_rejected_before_executor_effect(self): from types import SimpleNamespace diff --git a/tests/test_native_scheduling.py b/tests/test_native_scheduling.py index 938feb271..1db000928 100644 --- a/tests/test_native_scheduling.py +++ b/tests/test_native_scheduling.py @@ -164,6 +164,17 @@ def test_formats_all_types(self): class TestUpdateSchedule: + async def test_strict_update_without_catalog_validates_persisted_tool(self): + """Legacy/direct handler construction still checks the selected tool schema.""" + sched = MagicMock() + sched.list_all.return_value = [{"id": "S1", "tool_name": "run_command"}] + sched.update = AsyncMock(return_value={"id": "S1"}) + result = await _tools(sched)._handle_update_schedule( + ValidatedNestedPayload({"schedule_id": "S1", "tool_input": {}}) + ) + assert "invalid input for selected tool" in result + sched.update.assert_not_awaited() + async def test_strict_update_revalidates_replaced_input_against_persisted_tool(self): sched = MagicMock() sched.list_all.return_value = [{"id": "S1", "tool_name": "run_command"}] diff --git a/tests/test_scheduled_events.py b/tests/test_scheduled_events.py index 02685c4ff..01ff96b63 100644 --- a/tests/test_scheduled_events.py +++ b/tests/test_scheduled_events.py @@ -235,6 +235,43 @@ async def test_dispatch_exception(self): class TestWorkflow: + async def test_strict_workflow_stops_on_step_permission_denial(self): + executor = MagicMock() + executor.check_permission.return_value = "run_command denied" + h = _handlers(tool_executor=executor) + h._tool_loop._tool_catalog = SimpleNamespace(merged_definitions=lambda: [{ + "name": "run_command", "input_schema": {"type": "object"}, + }]) + channel = _channel() + + result = await h._run_scheduled_workflow(channel, { + "description": "strict", "requester_id": "u", "_nested_payload_validated": True, + "steps": [{"tool_name": "run_command", "tool_input": {}}], + }) + + assert result is False + assert "run_command denied" in channel.send.await_args.args[0] + executor.check_permission.assert_called_once_with("run_command", "u") + h._tool_loop.dispatch_loop_tool_inner.assert_not_awaited() + + @pytest.mark.parametrize("skill_name", [None, ""]) + async def test_strict_workflow_rejects_missing_skill_name(self, skill_name): + h = _handlers() + h._tool_loop._tool_catalog = SimpleNamespace(merged_definitions=lambda: [{ + "name": "invoke_skill", "input_schema": {"type": "object"}, + }]) + channel = _channel() + + result = await h._run_scheduled_workflow(channel, { + "description": "strict", "_nested_payload_validated": True, + "steps": [{"tool_name": "invoke_skill", "tool_input": {"name": skill_name}}], + }) + + assert result is False + assert "invoke_skill requires a skill name" in channel.send.await_args.args[0] + h._tool_executor.check_permission.assert_called_once_with("invoke_skill", None) + h._tool_loop.dispatch_loop_tool_inner.assert_not_awaited() + async def test_strict_workflow_rejects_invalid_step_before_dispatch(self): h = _handlers() h._tool_loop._tool_catalog = SimpleNamespace(merged_definitions=lambda: [{ diff --git a/tests/test_scheduler.py b/tests/test_scheduler.py index 5bdcea2e2..a2b76d478 100644 --- a/tests/test_scheduler.py +++ b/tests/test_scheduler.py @@ -636,6 +636,37 @@ async def test_strict_metadata_update_does_not_certify_legacy_nested_payload(sel assert updated["_nested_payload_validated"] is False assert updated["steps"] == legacy_steps + async def test_replacing_legacy_nested_payload_certifies_only_its_new_content(self, tmp_path): + s = _make_scheduler(tmp_path) + check = await s.add( + "legacy check", "check", "chan1", cron="0 * * * *", + tool_name="run_command", tool_input={"command": "old"}, + ) + updated_check = await s.update( + check["id"], tool_input={"command": "new"}, nested_payload_validated=True, + ) + assert updated_check["_nested_payload_validated"] is True + assert updated_check["tool_input"] == {"command": "new"} + + workflow = await s.add( + "legacy workflow", "workflow", "chan1", cron="0 * * * *", + steps=[{"tool_name": "web_search", "tool_input": {"query": "old"}}], + ) + new_steps = [{"tool_name": "web_search", "tool_input": {"query": "new"}}] + updated_workflow = await s.update( + workflow["id"], steps=new_steps, nested_payload_validated=True, + ) + assert updated_workflow["_nested_payload_validated"] is True + assert updated_workflow["steps"] == new_steps + + # A legacy caller's replacement never gains certification by itself. + legacy = await s.add( + "uncertified", "check", "chan1", cron="0 * * * *", + tool_name="run_command", tool_input={"command": "old"}, + ) + unchanged = await s.update(legacy["id"], tool_input={"command": "again"}) + assert unchanged["_nested_payload_validated"] is False + async def test_update_no_fields_still_persists(self, tmp_path): """Calling update with no changed fields returns the schedule unchanged.""" s = _make_scheduler(tmp_path) From eabf2407c7e8d73dbd6799448e192b3834650c86 Mon Sep 17 00:00:00 2001 From: Odin Date: Sun, 27 Sep 2026 21:47:00 -0400 Subject: [PATCH 07/12] Cover digest failure and scheduler host shortcut branches --- tests/test_native_scheduling.py | 8 +++++ tests/test_scheduled_events.py | 52 +++++++++++++++++++++++++++++++++ 2 files changed, 60 insertions(+) diff --git a/tests/test_native_scheduling.py b/tests/test_native_scheduling.py index 1db000928..f3d22eef9 100644 --- a/tests/test_native_scheduling.py +++ b/tests/test_native_scheduling.py @@ -38,6 +38,14 @@ def test_check_tool_name_and_command_shortcut(self): assert inp["tool_name"] == "run_command" assert inp["tool_input"]["command"] == "uname -r" + def test_check_shortcut_carries_an_explicit_host(self): + inp: dict[str, Any] = { + "action": "check", "tool_name": "run_command", + "command": "df -h", "host": "configured-host", + } + assert _tools()._validate_schedule_payload(inp) is None + assert inp["tool_input"] == {"command": "df -h", "host": "configured-host"} + def test_check_missing_tool_input(self): t = _tools() assert "requires 'tool_input'" in t._validate_schedule_payload( diff --git a/tests/test_scheduled_events.py b/tests/test_scheduled_events.py index 01ff96b63..b626cfe2b 100644 --- a/tests/test_scheduled_events.py +++ b/tests/test_scheduled_events.py @@ -57,6 +57,39 @@ def _handlers(**ov): class TestDigest: + async def test_every_failed_check_notifies_then_fails_without_summarizing(self): + channel = _channel() + failure = ToolResult( + output="probe failed", ok=False, error="execution_error", tool_name="run_command" + ) + loop = MagicMock(dispatch_loop_tool_inner=AsyncMock(return_value=failure)) + gateway = SimpleNamespace(active_client=SimpleNamespace(chat=AsyncMock())) + handler = _handlers( + get_channel=lambda _id: channel, + tool_loop=loop, + llm_gateway=gateway, + ) + + with pytest.raises(RuntimeError, match="all 2 checks failed"): + await handler._on_scheduled_digest({"id": "all-failed", "channel_id": "1"}) + + channel.send.assert_awaited_once() + notice = channel.send.await_args.args[0] + assert "Collection failed for every check (2 of 2)" in notice + assert "probe failed" in notice + gateway.active_client.chat.assert_not_awaited() + + async def test_audit_failure_after_digest_delivery_does_not_duplicate_notice(self): + channel = _channel() + audit = MagicMock(log_execution=AsyncMock(side_effect=RuntimeError("audit offline"))) + handler = _handlers(get_channel=lambda _id: channel, audit=audit) + + await handler._on_scheduled_digest({"id": "delivered", "channel_id": "1"}) + + channel.send.assert_awaited_once() + assert "LLM summary" in channel.send.await_args.args[0] + audit.log_execution.assert_awaited_once() + async def test_no_channel_id(self): h = _handlers() with pytest.raises(RuntimeError, match="has no channel_id"): @@ -114,6 +147,25 @@ async def test_collect_failure_and_notice_delivery_failure_propagate(self): class TestFormatDigestRaw: + async def test_gather_exception_marks_the_probe_failed(self): + # CancelledError is a BaseException: dispatch's Exception handler cannot + # consume it, but gather(return_exceptions=True) must record its failure. + import asyncio + + loop = MagicMock( + dispatch_loop_tool_inner=AsyncMock( + side_effect=[asyncio.CancelledError(), "memory available"] + ) + ) + handler = _handlers(tool_loop=loop) + + raw, failed, total = await handler._format_digest_raw({}, _channel()) + + assert total == 2 + assert failed == ["Disk (srv)"] + assert "### Disk (srv)\nCollection failed:" in raw + assert "### Memory (srv)\nmemory available" in raw + async def test_no_configured_hosts_fails_instead_of_reporting_empty_digest(self): h = _handlers(get_config=lambda: SimpleNamespace(tools=SimpleNamespace(hosts={}))) From fbdad0ffe4ffca5f78cfa538f6ffc8b1cef443f2 Mon Sep 17 00:00:00 2001 From: Odin Date: Sun, 27 Sep 2026 22:05:43 -0400 Subject: [PATCH 08/12] Preserve nested computer schema types in strict branch refinements --- src/llm/strict_tool_adapter.py | 37 +++++++++++++++++++--------------- 1 file changed, 21 insertions(+), 16 deletions(-) diff --git a/src/llm/strict_tool_adapter.py b/src/llm/strict_tool_adapter.py index 3473b05c5..5784dae88 100644 --- a/src/llm/strict_tool_adapter.py +++ b/src/llm/strict_tool_adapter.py @@ -164,14 +164,7 @@ def _computer_branches(schema): # Case constraints can refine a shared nested object (notably # computer_act.expect.type). Retain its base type and properties # while replacing only the overridden constraints. - refined = deepcopy(props.get(key, {})) - refined.update({k: v for k, v in override.items() if k != "properties"}) - if "properties" in override: - refined["properties"] = { - **refined.get("properties", {}), - **override["properties"], - } - props[key] = refined + props[key] = _refine_schema(props.get(key, {}), override) props["operation"] = {"type": "string", "const": op} branch = deepcopy(schema) branch.pop("oneOf") @@ -186,6 +179,23 @@ def _computer_branches(schema): return branches +def _refine_schema(base, override): + """Apply operation-specific constraints without losing nested base types. + + The canonical oneOf branches refine expect.type with an enum or const, + while its actual string type lives in the shared property definition. + """ + result = deepcopy(base) + for key, value in override.items(): + if key == "properties": + properties = result.setdefault("properties", {}) + for child, refinement in value.items(): + properties[child] = _refine_schema(properties.get(child, {}), refinement) + else: + result[key] = deepcopy(value) + return result + + def _allows_null(schema): return ( schema is True @@ -223,14 +233,9 @@ def _normalize(value, canonical, wire): ) for key, override in case["properties"].items(): if key in canonical["properties"] and isinstance(override, dict): - merged = deepcopy(canonical["properties"][key]) - merged.update({k: v for k, v in override.items() if k != "properties"}) - if "properties" in override: - merged["properties"] = { - **merged.get("properties", {}), - **override["properties"], - } - canonical["properties"][key] = merged + canonical["properties"][key] = _refine_schema( + canonical["properties"][key], override + ) elif "anyOf" in canonical: canonical = canonical["anyOf"][idx] return _normalize(value, canonical, wire["anyOf"][idx]) From 9b083eef2466542537a723ce167008519786884f Mon Sep 17 00:00:00 2001 From: Odin Date: Sun, 27 Sep 2026 22:08:11 -0400 Subject: [PATCH 09/12] Lower unsupported computer key lookaround with canonical enforcement --- docs/strict-wire-lowerings.md | 1 + src/llm/strict_tool_adapter.py | 7 +++++++ tests/test_strict_tool_adapter.py | 5 +++++ 3 files changed, 13 insertions(+) diff --git a/docs/strict-wire-lowerings.md b/docs/strict-wire-lowerings.md index e7acf028f..fc31dd44f 100644 --- a/docs/strict-wire-lowerings.md +++ b/docs/strict-wire-lowerings.md @@ -73,6 +73,7 @@ tripwire before dispatch; these are not generic recursive keyword deletions. | `computer_observe.properties.task_context.minProperties` | Omitted from wire object schema | Canonical validator rejects an empty task context | | `computer_act.properties.key.allOf[*].not`, and the matching nested `computer_act.properties.steps.items.properties.key.allOf[*].not` | Omitted from wire key schema | Canonical validator rejects repeated key modifiers | | `computer_act.properties.key.allOf`, and the matching nested `computer_act.properties.steps.items.properties.key.allOf` | Omitted from wire key schema | Canonical validator applies every key-chord restriction | +| `computer_act.properties.key.pattern` and nested `computer_act.properties.steps.items.properties.key.pattern` (Python lookaround) | Key string type retained; incompatible lookaround omitted on wire | Canonical computer payload validator applies the exact original regex before dispatch; controller key parsing independently rejects malformed chords | | `computer_act.oneOf`, `computer_act.properties.steps.items.oneOf` and their coordinate-versus-region sub-unions | Nested operation branches with const selectors; coordinate and region remain distinct branches | Canonical validator and adapter ambiguity check reject conflicting and unmatched branches | | `false` property schemas under `computer_act.oneOf[*].properties` and `computer_act.properties.steps.items.oneOf[*].properties`, including coordinate-versus-region sub-unions | Forbidden property removed from each wire branch | Canonical validator rejects forbidden fields, even if supplied as null | diff --git a/src/llm/strict_tool_adapter.py b/src/llm/strict_tool_adapter.py index 5784dae88..1f66f5ded 100644 --- a/src/llm/strict_tool_adapter.py +++ b/src/llm/strict_tool_adapter.py @@ -72,6 +72,12 @@ def _compile(node, name, builtin, path=()): raise ValueError( f"{name}: unsupported constraint outside audited computer contract at {path}" ) + # computer_act.key uses a Python lookaround that the server's regex + # dialect rejects. The canonical validator enforces the original pattern + # (and the action parser checks the chord) before any desktop dispatch. + computer_key = name == "computer_act" and path == ("key",) + if "pattern" in node and computer_key and "(?" not in node["pattern"]: + raise ValueError(f"{name}: unexpected key pattern at {path}") if any(key in node for key in ("oneOf", "allOf", "not")) and name != "computer_act": raise ValueError(f"{name}: unaudited combinator at {path}") if "$ref" in node or "$defs" in node or isinstance(node.get("type"), list): @@ -114,6 +120,7 @@ def _compile(node, name, builtin, path=()): out = { key: deepcopy(value) for key, value in node.items() + if key != "pattern" or not computer_key if key in ( "type", diff --git a/tests/test_strict_tool_adapter.py b/tests/test_strict_tool_adapter.py index 49d42c22e..2eb1dd74b 100644 --- a/tests/test_strict_tool_adapter.py +++ b/tests/test_strict_tool_adapter.py @@ -36,6 +36,11 @@ def test_catalog_strict_and_request_local(): computer = adapter.wire_tools[-1]["parameters"] assert "anyOf" not in computer assert len(computer["properties"]["payload"]["anyOf"]) == 14 + key_branch = next( + branch for branch in computer["properties"]["payload"]["anyOf"] + if branch["properties"]["operation"]["const"] == "key" + ) + assert "pattern" not in key_branch["properties"]["key"] def test_forced_values_and_nested_omission(): From 9db338bfaf378d36369754801d75ed340d3a5589 Mon Sep 17 00:00:00 2001 From: Odin Date: Sun, 27 Sep 2026 22:12:44 -0400 Subject: [PATCH 10/12] Validate nested HTTP headers in canonical dictionary form --- src/tools/nested_payload.py | 16 +++++++++++++++- tests/test_strict_nested_payload.py | 22 ++++++++++++++++++++++ 2 files changed, 37 insertions(+), 1 deletion(-) diff --git a/src/tools/nested_payload.py b/src/tools/nested_payload.py index d42ca81a6..d84e6f34a 100644 --- a/src/tools/nested_payload.py +++ b/src/tools/nested_payload.py @@ -135,7 +135,21 @@ def check(target, payload, label): raise ValueError(f"{label} selects unknown tool {target!r}") if not isinstance(payload, dict): raise ValueError(f"{label} must be an object") - _validate(schema, payload, allow_placeholders=allow_placeholders) + # http_probe's new wire schema uses header records, while its public + # canonical API and persisted legacy inputs still accept a dictionary. + # Validate an equivalent record view, keeping the original dict for + # execution. Never let duplicate names bypass the header parser. + if target == "http_probe" and isinstance(payload.get("headers"), dict): + from .http_probe_ops import normalize_probe_headers + + normalized = normalize_probe_headers(payload["headers"]) + validation_view = dict(payload) + validation_view["headers"] = [ + {"name": name, "value": value} for name, value in normalized.items() + ] + _validate(schema, validation_view, allow_placeholders=allow_placeholders) + else: + _validate(schema, payload, allow_placeholders=allow_placeholders) if tool_name in ("schedule_task", "update_schedule"): if ( diff --git a/tests/test_strict_nested_payload.py b/tests/test_strict_nested_payload.py index f2709acd4..074f0282b 100644 --- a/tests/test_strict_nested_payload.py +++ b/tests/test_strict_nested_payload.py @@ -82,3 +82,25 @@ def test_invoke_skill_validates_named_skill_schema(): validate_nested_payload( "invoke_skill", {"name": "example_skill", "input": {"wrong": 1}}, catalog() ) + + +def test_nested_http_probe_preserves_canonical_header_dict_and_checks_record_shape(): + from src.tools.registry import TOOL_MAP + + definitions = [{"name": "http_probe", "input_schema": TOOL_MAP["http_probe"]["input_schema"]}] + args = { + "steps": [ + { + "tool_name": "http_probe", + "tool_input": { + "url": "https://example.com/?q={prev_output}", + "headers": {"X-Test": "quoted"}, + }, + } + ] + } + result = validate_nested_payload("delegate_task", args, definitions) + assert result["steps"][0]["tool_input"]["headers"] == {"X-Test": "quoted"} + args["steps"][0]["tool_input"]["headers"] = {"X-Test": 123} + with pytest.raises(ValueError): + validate_nested_payload("delegate_task", args, definitions) From d0f3cbd30a7f95b91dc6ebd14ba3637b415db273 Mon Sep 17 00:00:00 2001 From: Odin Date: Sun, 27 Sep 2026 22:45:45 -0400 Subject: [PATCH 11/12] Harden external schema fallback and bound resolution logging --- src/llm/openai_codex.py | 18 ++++++ src/llm/strict_tool_adapter.py | 34 +++++++--- tests/test_openai_codex_client.py | 34 ++++++++++ tests/test_strict_tool_adapter.py | 79 ++++++++++++++++++++++++ tests/test_tool_lifecycle_correlation.py | 66 +++++++++++++++++++- 5 files changed, 220 insertions(+), 11 deletions(-) diff --git a/src/llm/openai_codex.py b/src/llm/openai_codex.py index f9d3f6318..84be54565 100644 --- a/src/llm/openai_codex.py +++ b/src/llm/openai_codex.py @@ -1222,6 +1222,24 @@ def finish_call(call_id: str, name: str, raw_args: str) -> ToolCall: input={}, parse_error=f"invalid tool arguments: {exc}", ) + except Exception as exc: + log.exception( + "Codex tool adapter failure for name=%s type=%s", + name, type(exc).__name__, + # Traceback source lines and exception messages may + # contain argument text. Emit only sanitized exception + # metadata, not the original traceback or message. + exc_info=(RuntimeError, RuntimeError("redacted adapter failure"), None), + ) + return ToolCall( + id=call_id, + name=name, + input={}, + parse_error=( + "invalid tool arguments: internal adapter error " + f"({type(exc).__name__})" + ), + ) return ToolCall(id=call_id, name=name, input=arguments) async for raw_line in resp.content: diff --git a/src/llm/strict_tool_adapter.py b/src/llm/strict_tool_adapter.py index 1f66f5ded..ba246f3f9 100644 --- a/src/llm/strict_tool_adapter.py +++ b/src/llm/strict_tool_adapter.py @@ -6,10 +6,13 @@ import json import logging from copy import deepcopy +from threading import Lock from jsonschema import Draft202012Validator log = logging.getLogger(__name__) +_resolution_lock = Lock() +_resolution_seen: set[tuple[str, str, str, str]] = set() PAYLOADS = {"schedule_task", "update_schedule", "delegate_task", "invoke_skill"} LOWERED = {"uniqueItems", "minProperties", "oneOf", "allOf", "not", "default"} WIRE = { @@ -65,7 +68,7 @@ def _optional(node, original): def _compile(node, name, builtin, path=()): if not isinstance(node, dict): raise ValueError(f"{name}: unsupported schema at {path}") - unsupported = set(node) - WIRE - (LOWERED if builtin else set()) + unsupported = set(node) - WIRE - (LOWERED if builtin else {"default"}) if unsupported: raise ValueError(f"{name}: unsupported keywords {sorted(unsupported)} at {path}") if not name.startswith("computer_") and set(node) & (LOWERED - {"default"}): @@ -153,6 +156,8 @@ def _compile(node, name, builtin, path=()): out["required"] = list(props) out["additionalProperties"] = False if node.get("type") == "array": + if "items" not in node: + raise ValueError(f"{name}: array without items at {path}") out["items"] = _compile(node["items"], name, builtin, path + ("*",)) return out @@ -298,10 +303,13 @@ def __init__(self, tools): else _compile(canonical, name, builtin) ) mode, reason = ("builtin_strict" if builtin else "external_compiled"), None - except ValueError as exc: + except Exception as exc: if builtin: raise - reason = str(exc) + reason = ( + str(exc) if isinstance(exc, ValueError) + else f"compiler {type(exc).__name__}" + ) if isinstance(canonical, dict) and canonical.get("type") == "object": wire = { "type": "object", @@ -410,13 +418,19 @@ def record_resolution(self, event=None): val = resolved.get(name) item["resolution"] = "true" if val is True else "false" if val is False else "unknown" result[name] = item["resolution"] - log.info( - "Codex schema name=%s fingerprint=%s mode=%s resolution=%s reason=%s", - name, - item["fingerprint"], - item["mode"], - item["resolution"], - item["reason"], + key = (name, item["fingerprint"], item["mode"], item["resolution"]) + with _resolution_lock: + first = key not in _resolution_seen + _resolution_seen.add(key) + if item["mode"] == "builtin_strict" and item["resolution"] != "true": + level = logging.WARNING + elif first: + level = logging.INFO + else: + level = logging.DEBUG + log.log( + level, "Codex schema name=%s fingerprint=%s mode=%s resolution=%s reason=%s", + name, item["fingerprint"], item["mode"], item["resolution"], item["reason"], ) self._resolution_logged = True return result diff --git a/tests/test_openai_codex_client.py b/tests/test_openai_codex_client.py index b46f84aba..f0e50d567 100644 --- a/tests/test_openai_codex_client.py +++ b/tests/test_openai_codex_client.py @@ -271,6 +271,40 @@ async def test_empty_stream_returns_empty(self): class TestReadToolStream: + @pytest.mark.asyncio + async def test_adapter_unexpected_error_is_tool_error_without_aborting_stream(self, caplog): + from src.llm.openai_codex import _request_tool_adapter + from src.llm.strict_tool_adapter import compile_catalog + + adapter = compile_catalog([{"name": "broken_ext", "input_schema": { + "type": "object", "properties": {}, "additionalProperties": False, + }}]) + + def broken_accept(name, arguments): + raise RuntimeError("secret argument content must not escape") + + adapter.accept = broken_accept + item = {"type": "function_call", "call_id": "broken", "name": "broken_ext", + "arguments": "{}"} + events = [ + {"type": "response.completed", "response": {"output": [item, { + "type": "message", "content": [{"text": "turn continued"}], + }]}} + ] + token = _request_tool_adapter.set(adapter) + try: + out = await _client()._read_tool_stream(_FakeResp([_sse(e) for e in events])) + finally: + _request_tool_adapter.reset(token) + assert out.tool_calls[0].parse_error == ( + "invalid tool arguments: internal adapter error (RuntimeError)" + ) + assert out.tool_calls[0].input == {} + assert out.text == "turn continued" + assert "secret argument content" not in out.tool_calls[0].parse_error + assert "name=broken_ext type=RuntimeError" in caplog.text + assert "secret argument content" not in caplog.text + @pytest.mark.asyncio @pytest.mark.parametrize("path", ["arguments_done", "item_done", "completed", "duplicates"]) async def test_adapter_accepts_once_on_each_finalization_path(self, path): diff --git a/tests/test_strict_tool_adapter.py b/tests/test_strict_tool_adapter.py index 2eb1dd74b..a05564044 100644 --- a/tests/test_strict_tool_adapter.py +++ b/tests/test_strict_tool_adapter.py @@ -126,6 +126,85 @@ def test_external_modes_round_trip_and_resolution(): ) == {"closed_ext": "true", "open_ext": "false", "odd_ext": "unknown"} +def test_external_unsupported_subschemas_never_break_catalog(monkeypatch): + import src.llm.strict_tool_adapter as strict_adapter + + def external(name, property_schema): + return {"name": name, "input_schema": { + "type": "object", "properties": {"tags": property_schema}, + "required": ["tags"], "additionalProperties": False, + }} + + tools = [external("itemless", {"type": "array"}), + external("boolean", True)] + adapter = compile_catalog(tools) + assert [item["mode"] for item in adapter.report.values()] == [ + "external_envelope", "external_envelope", + ] + assert adapter.accept("itemless", {"json": '{"tags":[1,"x"]}'}) == { + "tags": [1, "x"], + } + assert adapter.accept("boolean", {"json": '{"tags":42}'}) == {"tags": 42} + assert "array without items" in adapter.report["itemless"]["reason"] + + original = strict_adapter._compile + + def broken(node, name, builtin, path=()): + if name == "broken_external": + raise RuntimeError("do not leak this text") + return original(node, name, builtin, path) + + monkeypatch.setattr(strict_adapter, "_compile", broken) + injected = compile_catalog([external("broken_external", {"type": "array"})]) + assert injected.report["broken_external"]["mode"] == "external_envelope" + assert injected.report["broken_external"]["reason"] == "compiler RuntimeError" + assert injected.accept("broken_external", {"json": '{"tags":[]}'}) == {"tags": []} + with pytest.raises(ValueError, match="array without items"): + compile_catalog([{ + "name": "run_command", "input_schema": tools[0]["input_schema"], + }]) + + +def test_external_defaults_lowered_without_mutating_canonical(): + schema = {"type": "object", "properties": { + "name": {"type": "string", "default": "unused"}, + }, "additionalProperties": False} + adapter = compile_catalog([{"name": "defaults_ext", "input_schema": schema}]) + assert adapter.report["defaults_ext"]["mode"] == "external_compiled" + assert "default" not in str(adapter.wire_tools[0]["parameters"]) + assert schema["properties"]["name"]["default"] == "unused" + assert adapter.accept("defaults_ext", {"name": None}) == {} + + +def test_resolution_logging_is_process_memoized_and_builtin_warnings_repeat(caplog, monkeypatch): + import logging + + import src.llm.strict_tool_adapter as strict_adapter + + monkeypatch.setattr(strict_adapter, "_resolution_seen", set()) + catalog = [{"name": "fixture_ext", "input_schema": { + "type": "object", "properties": {}, "additionalProperties": False, + }}] + with caplog.at_level(logging.INFO, logger=strict_adapter.__name__): + for _ in range(3): + adapter = compile_catalog(catalog) + adapter.record_resolution({"type": "response.created", "response": { + "tools": [{**adapter.wire_tools[0], "strict": True}], + }}) + assert len([r for r in caplog.records if r.levelno == logging.INFO]) == 1 + changed = copy.deepcopy(catalog) + changed[0]["input_schema"]["properties"]["extra"] = {"type": "string"} + compile_catalog(changed).record_resolution(None) + compile_catalog(catalog).record_resolution(None) + assert len([r for r in caplog.records if r.levelno == logging.INFO]) == 3 + builtin = next(t for t in get_tool_definitions() if t["name"] == "browser_read_table") + for _ in range(2): + compile_catalog([builtin]).record_resolution(None) + warnings = [r for r in caplog.records if r.levelno == logging.WARNING] + assert len(warnings) == 2 + assert all("browser_read_table" in r.message for r in warnings) + + def test_probe_headers_lowered_after_wire_validation(): adapter = compile_catalog(get_tool_definitions()) args = _wire_args(adapter, "http_probe", {"url": "https://example.org", "headers": [ diff --git a/tests/test_tool_lifecycle_correlation.py b/tests/test_tool_lifecycle_correlation.py index 9d2aea177..af92b97ca 100644 --- a/tests/test_tool_lifecycle_correlation.py +++ b/tests/test_tool_lifecycle_correlation.py @@ -3,7 +3,7 @@ import asyncio import json from types import SimpleNamespace -from unittest.mock import AsyncMock, Mock +from unittest.mock import AsyncMock, MagicMock, Mock import pytest @@ -12,8 +12,13 @@ from src.audit.logger import AuditLogger from src.audit.tool_context import _pending_observers from src.config.schema import ToolsConfig +from src.discord.native_tools.registry import NativeToolDispatcher from src.discord.tool_loop import ToolLoopRunner +from src.llm.strict_tool_adapter import compile_catalog +from src.llm.tool_history import assistant_content, normalize_tool_calls from src.observability.correlation import reset_turn, set_turn +from src.tools.nested_payload import ValidatedNestedPayload +from src.tools.registry import get_tool_definitions from src.tools.result_validator import ToolResult @@ -91,6 +96,65 @@ def correlation(record): ) +@pytest.mark.parametrize("route", ["foreground", "autonomous", "agent"]) +@pytest.mark.parametrize("name", ["schedule_task", "update_schedule", "delegate_task"]) +async def test_adapter_marker_reaches_native_dispatch_by_identity(tmp_path, route, name): + """Storage/audit snapshots may copy, but the executable argument must not.""" + adapter = compile_catalog(get_tool_definitions()) + wire = next(t["parameters"]["properties"] for t in adapter.wire_tools + if t["name"] == name) + specified = { + "schedule_task": {"description": "proof", "action": "reminder", "message": "test"}, + "update_schedule": {"schedule_id": "existing", "description": "proof"}, + "delegate_task": {"description": "proof", "steps": [{ + "tool_name": "run_command", "tool_input": '{"command":"uptime"}', + }]}, + }[name] + args = {key: specified.get(key) for key in wire} + if name == "delegate_task": + step_schema = wire["steps"]["items"]["properties"] + args["steps"] = [{key: specified["steps"][0].get(key) for key in step_schema}] + marked = adapter.accept(name, args) + assert isinstance(marked, ValidatedNestedPayload) + runner, st, _ = harness(tmp_path, native=True) + call = SimpleNamespace(id="marker", name=name, input=marked, parse_error=None) + owner = SimpleNamespace( + _handle_schedule_task=AsyncMock(return_value="ok"), + _handle_update_schedule=AsyncMock(return_value="ok"), + _handle_delegate_task=AsyncMock(return_value="ok"), + ) + native = NativeToolDispatcher( + owners={"scheduling": owner, "agents": owner}, skill_manager=MagicMock(), + tool_catalog=MagicMock(), prompt_builder=MagicMock(), channel_state=MagicMock(), + ) + native.register(name, "agents" if name == "delegate_task" else "scheduling", + f"_handle_{name}", "input" if name == "update_schedule" else "msg_input") + runner._native_tools = native + if route == "foreground": + await runner._run_one_tool(st, call) + elif route == "autonomous": + await runner._run_one_loop_tool(st, call) + else: + agent = AgentInfo(id="marker-agent", label="proof", goal="proof", channel_id="c", + requester_id="u", requester_name="User", turn_id="marker-turn") + accepted_calls = normalize_tool_calls([call]) + assistant_content("", accepted_calls) # checkpoint copy is not the executable input + assert isinstance(accepted_calls[0]["input"], ValidatedNestedPayload) + + async def execute(tool_name, incoming): + return await runner.dispatch_loop_tool(tool_name, incoming, st.msg_proxy, "u") + + await execute_cycle( + agent, accepted_calls, execute, [], + timeouts={}, default_timeout=1, + ) + invoked = getattr(owner, f"_handle_{name}") + invoked.assert_awaited_once() + assert isinstance(invoked.await_args.args[-1], ValidatedNestedPayload) + if route != "agent": + assert invoked.await_args.args[-1] is marked + + @pytest.mark.parametrize("route", ["foreground", "autonomous", "agent"]) @pytest.mark.parametrize("native", [False, True]) @pytest.mark.parametrize("failure", [None, ValueError("ordinary failure")]) From ff362f52c92cdafb04ca74218807639990c400b5 Mon Sep 17 00:00:00 2001 From: Odin Date: Sun, 27 Sep 2026 22:48:27 -0400 Subject: [PATCH 12/12] Type resolution memo key explicitly --- src/llm/strict_tool_adapter.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/src/llm/strict_tool_adapter.py b/src/llm/strict_tool_adapter.py index ba246f3f9..a424e7cab 100644 --- a/src/llm/strict_tool_adapter.py +++ b/src/llm/strict_tool_adapter.py @@ -418,7 +418,9 @@ def record_resolution(self, event=None): val = resolved.get(name) item["resolution"] = "true" if val is True else "false" if val is False else "unknown" result[name] = item["resolution"] - key = (name, item["fingerprint"], item["mode"], item["resolution"]) + key = ( + name, str(item["fingerprint"]), str(item["mode"]), str(item["resolution"]), + ) with _resolution_lock: first = key not in _resolution_seen _resolution_seen.add(key)