Sfoglia il codice sorgente

fix: enforce event fallback semantics

Problem: complete-invalid deterministic arguments triggered unnecessary LLM fallback, raw and parsed JSON equality was type-insensitive, and validator failures looked like user errors.

Risk: stricter completeness and schema isolation may expose resolver or plugin-definition bugs previously masked by fallback or post-registration mutation.
zhenyu.hu 2 settimane fa
parent
commit
4b163509c3

+ 121 - 48
src/agent_lab/application/events/kernel.py

@@ -3,13 +3,14 @@ from __future__ import annotations
 import inspect
 import json
 from collections.abc import Awaitable, Callable, Iterable
-from dataclasses import replace
+from dataclasses import dataclass, replace
 from typing import Any
 
 from jsonschema.exceptions import ValidationError
 
 from agent_lab.application.events.models import (
     EventDefinition,
+    EventArgumentResolution,
     EventExecutionContext,
     EventRequest,
     EventResult,
@@ -27,6 +28,13 @@ ArgumentFallback = Callable[
 ]
 
 
+@dataclass(frozen=True)
+class _ValidationResult:
+    status: EventStatus | None = None
+    error: str | None = None
+    missing_required: bool = False
+
+
 class EventKernel:
     def __init__(
         self,
@@ -49,19 +57,27 @@ class EventKernel:
         assert definition is not None
         resolved_context = context or EventExecutionContext()
 
-        resolved, resolution_error = self._resolve_deterministic(
-            definition,
-            request,
-            resolved_context,
+        resolved, resolution_complete, resolution_error = (
+            self._resolve_deterministic(
+                definition,
+                request,
+                resolved_context,
+            )
         )
         if resolution_error is not None:
             return resolution_error
         assert resolved is not None
 
-        validation_error = self._validation_error(definition, resolved.arguments)
+        validation = self._validate(definition, resolved.arguments)
+        if not resolution_complete and validation.status is None:
+            validation = _ValidationResult(
+                status=EventStatus.INVALID_ARGUMENTS,
+                error="event arguments incomplete",
+            )
         used_fallback = False
         if (
-            validation_error is not None
+            validation.status is not EventStatus.DEFINITION_ERROR
+            and (validation.missing_required or not resolution_complete)
             and request.source is not EventSource.PROVIDER_RESOLVED
             and definition.fallback_allowed
             and self.argument_fallback is not None
@@ -97,18 +113,18 @@ class EventKernel:
                 if fallback_error is not None:
                     return fallback_error
                 assert resolved is not None
-            validation_error = self._validation_error(
-                definition,
-                resolved.arguments,
-            )
+                validation = self._validate(
+                    definition,
+                    resolved.arguments,
+                )
 
-        if validation_error is not None:
+        if validation.status is not None:
             return self._result(
                 definition,
                 request,
-                EventStatus.INVALID_ARGUMENTS,
+                validation.status,
                 resolved=resolved,
-                error=validation_error,
+                error=validation.error,
                 used_fallback=used_fallback,
             )
 
@@ -159,7 +175,7 @@ class EventKernel:
         if early_result is not None:
             return early_result
         assert definition is not None
-        resolved, resolution_error = self._resolve_deterministic(
+        resolved, _, resolution_error = self._resolve_deterministic(
             definition,
             request,
             context or EventExecutionContext(),
@@ -167,14 +183,14 @@ class EventKernel:
         if resolution_error is not None:
             return resolution_error
         assert resolved is not None
-        validation_error = self._validation_error(definition, resolved.arguments)
-        if validation_error is not None:
+        validation = self._validate(definition, resolved.arguments)
+        if validation.status is not None:
             return self._result(
                 definition,
                 request,
-                EventStatus.INVALID_ARGUMENTS,
+                validation.status,
                 resolved=resolved,
-                error=validation_error,
+                error=validation.error,
             )
         if inspect.iscoroutinefunction(definition.handler):
             return self._result(
@@ -260,7 +276,7 @@ class EventKernel:
         definition: EventDefinition,
         request: EventRequest,
         context: EventExecutionContext,
-    ) -> tuple[ResolvedEventArguments | None, EventResult | None]:
+    ) -> tuple[ResolvedEventArguments | None, bool, EventResult | None]:
         if request.source is EventSource.PROVIDER_RESOLVED:
             return (
                 ResolvedEventArguments(
@@ -268,22 +284,29 @@ class EventKernel:
                     arguments=dict(request.arguments),
                     raw_arguments=request.raw_arguments,
                 ),
+                True,
                 None,
             )
         try:
-            arguments = (
+            value = (
                 definition.resolver(request, context)
                 if definition.resolver is not None
                 else dict(request.arguments)
             )
         except Exception as exc:
-            return None, self._resolution_error(
+            return None, False, self._resolution_error(
                 definition,
                 request,
                 f"event argument resolver failed: {exc}",
             )
-        if not isinstance(arguments, dict):
-            return None, self._resolution_error(
+        if isinstance(value, EventArgumentResolution):
+            arguments = value.arguments
+            complete = value.complete
+        else:
+            arguments = value
+            complete = True
+        if not isinstance(arguments, dict) or not isinstance(complete, bool):
+            return None, False, self._resolution_error(
                 definition,
                 request,
                 "event argument resolver returned invalid payload",
@@ -291,7 +314,7 @@ class EventKernel:
         try:
             raw_arguments = self._json_arguments(arguments)
         except (TypeError, ValueError) as exc:
-            return None, self._resolution_error(
+            return None, False, self._resolution_error(
                 definition,
                 request,
                 f"event argument resolver failed to serialize: {exc}",
@@ -302,6 +325,7 @@ class EventKernel:
                 arguments=dict(arguments),
                 raw_arguments=raw_arguments,
             ),
+            complete,
             None,
         )
 
@@ -334,8 +358,11 @@ class EventKernel:
                 used_fallback=True,
             )
         try:
-            raw_arguments = json.loads(value.raw_arguments)
-        except (TypeError, json.JSONDecodeError):
+            raw_arguments = json.loads(
+                value.raw_arguments,
+                parse_constant=_reject_json_constant,
+            )
+        except (TypeError, ValueError, json.JSONDecodeError):
             return None, self._resolution_error(
                 definition,
                 request,
@@ -343,7 +370,17 @@ class EventKernel:
                 resolved=value,
                 used_fallback=True,
             )
-        if raw_arguments != value.arguments:
+        try:
+            self._json_arguments(value.arguments)
+        except (TypeError, ValueError):
+            return None, self._resolution_error(
+                definition,
+                request,
+                "fallback parsed arguments are not valid JSON",
+                resolved=value,
+                used_fallback=True,
+            )
+        if not _json_values_equal(raw_arguments, value.arguments):
             return None, self._resolution_error(
                 definition,
                 request,
@@ -353,32 +390,44 @@ class EventKernel:
             )
         return value, None
 
-    def _validation_error(
+    def _validate(
         self,
         definition: EventDefinition,
         arguments: dict[str, Any],
-    ) -> str | None:
+    ) -> _ValidationResult:
         validator = self.registry.validator(definition.name)
         assert validator is not None
         try:
-            validator.validate(arguments)
-        except ValidationError as exc:
-            if exc.validator == "required":
-                missing = [
-                    name
-                    for name in exc.validator_value
-                    if name not in exc.instance
-                ]
-                return f"missing required arguments: {', '.join(missing)}"
-            if exc.validator == "type" and exc.path:
-                return (
-                    f"invalid argument type for {exc.path[-1]}: "
-                    f"expected {exc.validator_value}"
-                )
-            return f"invalid event arguments: {exc.message}"
+            errors = list(validator.iter_errors(arguments))
         except Exception as exc:
-            return f"event argument validation failed: {exc}"
-        return None
+            return _ValidationResult(
+                status=EventStatus.DEFINITION_ERROR,
+                error=f"event argument validation failed: {exc}",
+            )
+        if not errors:
+            return _ValidationResult()
+        missing_errors = [error for error in errors if error.validator == "required"]
+        error = missing_errors[0] if missing_errors else errors[0]
+        return _ValidationResult(
+            status=EventStatus.INVALID_ARGUMENTS,
+            error=self._format_validation_error(error),
+            missing_required=bool(missing_errors),
+        )
+
+    def _format_validation_error(self, error: ValidationError) -> str:
+        if error.validator == "required":
+            missing = [
+                name
+                for name in error.validator_value
+                if name not in error.instance
+            ]
+            return f"missing required arguments: {', '.join(missing)}"
+        if error.validator == "type" and error.path:
+            return (
+                f"invalid argument type for {error.path[-1]}: "
+                f"expected {error.validator_value}"
+            )
+        return f"invalid event arguments: {error.message}"
 
     def _resolution_error(
         self,
@@ -435,4 +484,28 @@ class EventKernel:
         )
 
     def _json_arguments(self, arguments: dict[str, Any]) -> str:
-        return json.dumps(arguments, ensure_ascii=False, separators=(",", ":"))
+        return json.dumps(
+            arguments,
+            ensure_ascii=False,
+            separators=(",", ":"),
+            allow_nan=False,
+        )
+
+
+def _reject_json_constant(value: str) -> None:
+    raise ValueError(f"non-standard JSON constant: {value}")
+
+
+def _json_values_equal(left: Any, right: Any) -> bool:
+    if type(left) is not type(right):
+        return False
+    if isinstance(left, dict):
+        return left.keys() == right.keys() and all(
+            _json_values_equal(left[key], right[key]) for key in left
+        )
+    if isinstance(left, list):
+        return len(left) == len(right) and all(
+            _json_values_equal(left_item, right_item)
+            for left_item, right_item in zip(left, right, strict=True)
+        )
+    return left == right

+ 8 - 1
src/agent_lab/application/events/models.py

@@ -18,6 +18,7 @@ class EventStatus(StrEnum):
     UNKNOWN = "unknown"
     DISABLED = "disabled"
     INVALID_ARGUMENTS = "invalid_arguments"
+    DEFINITION_ERROR = "definition_error"
     RESOLUTION_ERROR = "resolution_error"
     HANDLER_ERROR = "handler_error"
 
@@ -63,11 +64,17 @@ class ResolvedEventArguments:
     raw_arguments: str
 
 
+@dataclass(frozen=True)
+class EventArgumentResolution:
+    arguments: dict[str, Any]
+    complete: bool = True
+
+
 EventHandlerResult = dict[str, Any] | Awaitable[dict[str, Any]]
 EventHandler = Callable[[EventRequest], EventHandlerResult]
 EventArgumentResolver = Callable[
     [EventRequest, EventExecutionContext],
-    dict[str, Any],
+    dict[str, Any] | EventArgumentResolution,
 ]
 
 

+ 30 - 5
src/agent_lab/application/events/registry.py

@@ -2,6 +2,9 @@ from __future__ import annotations
 
 from collections.abc import Iterable
 from copy import deepcopy
+from dataclasses import replace
+from types import MappingProxyType
+from typing import Any
 
 from jsonschema import Draft202012Validator
 from jsonschema.exceptions import SchemaError
@@ -19,15 +22,19 @@ class EventRegistry:
     def register(self, definition: EventDefinition) -> None:
         if definition.name in self._definitions:
             raise ValueError(f"duplicate event definition: {definition.name}")
+        parameters = deepcopy(definition.parameters)
         try:
-            Draft202012Validator.check_schema(definition.parameters)
+            Draft202012Validator.check_schema(parameters)
         except SchemaError as exc:
             raise ValueError(
                 f"invalid event schema for {definition.name}: {exc.message}"
             ) from None
-        self._definitions[definition.name] = definition
+        self._definitions[definition.name] = replace(
+            definition,
+            parameters=_freeze_json(parameters),
+        )
         self._validators[definition.name] = Draft202012Validator(
-            definition.parameters
+            deepcopy(parameters)
         )
 
     def definition(self, name: str) -> EventDefinition | None:
@@ -42,7 +49,7 @@ class EventRegistry:
             {
                 "name": definition.name,
                 "description": definition.description,
-                "parameters": deepcopy(definition.parameters),
+                "parameters": _thaw_json(definition.parameters),
             }
             for definition in self._definitions.values()
             if enabled is None or definition.name in enabled
@@ -57,6 +64,24 @@ class EventRegistry:
             "function": {
                 "name": definition.name,
                 "description": definition.description,
-                "parameters": deepcopy(definition.parameters),
+                "parameters": _thaw_json(definition.parameters),
             },
         }
+
+
+def _freeze_json(value: Any) -> Any:
+    if isinstance(value, dict):
+        return MappingProxyType(
+            {key: _freeze_json(item) for key, item in value.items()}
+        )
+    if isinstance(value, list):
+        return tuple(_freeze_json(item) for item in value)
+    return value
+
+
+def _thaw_json(value: Any) -> Any:
+    if isinstance(value, dict | MappingProxyType):
+        return {key: _thaw_json(item) for key, item in value.items()}
+    if isinstance(value, tuple):
+        return [_thaw_json(item) for item in value]
+    return deepcopy(value)

+ 6 - 0
src/agent_lab/application/tools.py

@@ -152,6 +152,12 @@ class ToolRegistry:
             error = "unknown tool"
         elif result.status is EventStatus.DISABLED:
             error = "tool disabled"
+        elif result.status is EventStatus.DEFINITION_ERROR:
+            error = error.replace(
+                "event argument validation failed:",
+                "tool definition validation failed:",
+                1,
+            )
         elif result.status is EventStatus.HANDLER_ERROR:
             error = error.replace("event handler failed:", "tool handler failed:", 1)
             error = error.replace("EventKernel.execute", "execute_async")

+ 239 - 2
tests/test_event_kernel.py

@@ -20,6 +20,7 @@ from agent_lab.application.events import (
     ResultPolicy,
     RiskLevel,
 )
+from agent_lab.application.events.models import EventArgumentResolution
 from agent_lab.application.tools import (
     ToolDefinition,
     ToolExecutionContext,
@@ -193,6 +194,51 @@ async def test_kernel_calls_fallback_once_when_required_arguments_are_incomplete
     assert fallback_calls == [{}]
 
 
+@pytest.mark.asyncio
+@pytest.mark.parametrize(
+    "arguments",
+    [
+        {"query": 42},
+        {"query": "unsupported"},
+        {"query": "valid", "unexpected": True},
+    ],
+)
+async def test_kernel_does_not_fallback_for_complete_invalid_arguments(
+    arguments: dict[str, Any],
+):
+    fallback_calls = 0
+
+    async def fallback(*args: Any) -> ResolvedEventArguments:
+        nonlocal fallback_calls
+        fallback_calls += 1
+        return ResolvedEventArguments(
+            event_name="example.lookup",
+            arguments={"query": "valid"},
+            raw_arguments='{"query":"valid"}',
+        )
+
+    definition = _definition(
+        parameters={
+            "type": "object",
+            "properties": {"query": {"type": "string", "enum": ["valid"]}},
+            "required": ["query"],
+            "additionalProperties": False,
+        },
+        resolver=lambda request, context: arguments,
+    )
+
+    result = await EventKernel(
+        EventRegistry([definition]), argument_fallback=fallback
+    ).execute(
+        EventRequest(id="event-1", name="example.lookup"),
+        enabled_names=["example.lookup"],
+    )
+
+    assert result.status is EventStatus.INVALID_ARGUMENTS
+    assert result.used_fallback is False
+    assert fallback_calls == 0
+
+
 @pytest.mark.asyncio
 async def test_kernel_does_not_fallback_when_definition_disallows_it():
     fallback_calls = 0
@@ -226,6 +272,96 @@ async def test_kernel_does_not_fallback_when_definition_disallows_it():
     assert fallback_calls == 0
 
 
+@pytest.mark.asyncio
+async def test_structured_resolution_can_mark_optional_arguments_incomplete():
+    fallback_calls = 0
+
+    async def fallback(*args: Any) -> ResolvedEventArguments:
+        nonlocal fallback_calls
+        fallback_calls += 1
+        return ResolvedEventArguments(
+            event_name="example.lookup",
+            arguments={"query": "resolved"},
+            raw_arguments='{"query":"resolved"}',
+        )
+
+    definition = _definition(
+        parameters={
+            "type": "object",
+            "properties": {"query": {"type": "string"}},
+            "additionalProperties": False,
+        },
+        resolver=lambda request, context: EventArgumentResolution(
+            arguments={},
+            complete=False,
+        ),
+    )
+
+    result = await EventKernel(
+        EventRegistry([definition]), argument_fallback=fallback
+    ).execute(
+        EventRequest(id="event-1", name="example.lookup"),
+        enabled_names=["example.lookup"],
+    )
+
+    assert result.status is EventStatus.SUCCESS
+    assert result.arguments == {"query": "resolved"}
+    assert result.used_fallback is True
+    assert fallback_calls == 1
+
+
+@pytest.mark.asyncio
+async def test_structured_incomplete_optional_arguments_fail_without_fallback():
+    definition = _definition(
+        parameters={
+            "type": "object",
+            "properties": {"query": {"type": "string"}},
+        },
+        resolver=lambda request, context: EventArgumentResolution(
+            arguments={},
+            complete=False,
+        ),
+        handler=lambda request: {"ok": True},
+    )
+
+    result = await EventKernel(EventRegistry([definition])).execute(
+        EventRequest(id="event-1", name="example.lookup"),
+        enabled_names=["example.lookup"],
+    )
+
+    assert result.status is EventStatus.INVALID_ARGUMENTS
+    assert result.error == "event arguments incomplete"
+
+
+@pytest.mark.asyncio
+async def test_plain_dict_resolution_remains_complete_for_optional_schema():
+    fallback_calls = 0
+
+    async def fallback(*args: Any) -> ResolvedEventArguments:
+        nonlocal fallback_calls
+        fallback_calls += 1
+        raise AssertionError("fallback should not run")
+
+    definition = _definition(
+        parameters={
+            "type": "object",
+            "properties": {"query": {"type": "string"}},
+        },
+        resolver=lambda request, context: {},
+        handler=lambda request: {"ok": True},
+    )
+
+    result = await EventKernel(
+        EventRegistry([definition]), argument_fallback=fallback
+    ).execute(
+        EventRequest(id="event-1", name="example.lookup"),
+        enabled_names=["example.lookup"],
+    )
+
+    assert result.status is EventStatus.SUCCESS
+    assert fallback_calls == 0
+
+
 @pytest.mark.asyncio
 async def test_provider_resolved_arguments_are_not_rewritten_or_fallen_back():
     resolver_calls = 0
@@ -430,10 +566,36 @@ async def test_kernel_normalizes_validator_runtime_exceptions():
         enabled_names=[definition.name],
     )
 
-    assert result.status is EventStatus.INVALID_ARGUMENTS
+    assert result.status is EventStatus.DEFINITION_ERROR
     assert result.error.startswith("event argument validation failed:")
 
 
+@pytest.mark.asyncio
+async def test_tool_registry_maps_definition_errors_to_tool_compatibility_payload():
+    registry = ToolRegistry(
+        [
+            ToolDefinition(
+                name="broken.lookup",
+                description="Broken lookup.",
+                parameters={"$ref": "urn:agent-lab:missing-schema"},
+                handler=lambda event: {"tool": event.name},
+            )
+        ]
+    )
+
+    payload = await registry.execute_async(
+        ToolCallEvent(
+            id="call-1",
+            name="broken.lookup",
+            arguments={},
+            raw_arguments="{}",
+        )
+    )
+
+    assert payload["tool"] == "broken.lookup"
+    assert payload["error"].startswith("tool definition validation failed:")
+
+
 @pytest.mark.asyncio
 @pytest.mark.parametrize("boundary", ["resolver", "fallback"])
 async def test_kernel_normalizes_resolution_boundary_exceptions(boundary: str):
@@ -570,6 +732,81 @@ async def test_kernel_rejects_inconsistent_structured_fallback(
     assert result.error == expected_error
 
 
+@pytest.mark.asyncio
+@pytest.mark.parametrize(
+    ("arguments", "raw_arguments"),
+    [
+        ({"query": 1}, '{"query":true}'),
+        ({"query": 1.0}, '{"query":1}'),
+        ({"query": {"nested": [1]}}, '{"query":{"nested":[true]}}'),
+        ({"query": float("nan")}, '{"query":NaN}'),
+        ({"query": float("inf")}, '{"query":Infinity}'),
+    ],
+)
+async def test_kernel_rejects_noncanonical_fallback_json(
+    arguments: dict[str, Any],
+    raw_arguments: str,
+):
+    async def fallback(*args: Any) -> ResolvedEventArguments:
+        return ResolvedEventArguments(
+            event_name="example.lookup",
+            arguments=arguments,
+            raw_arguments=raw_arguments,
+        )
+
+    result = await EventKernel(
+        EventRegistry([_definition(resolver=lambda request, context: {})]),
+        argument_fallback=fallback,
+    ).execute(
+        EventRequest(id="event-1", name="example.lookup"),
+        enabled_names=["example.lookup"],
+    )
+
+    assert result.status is EventStatus.RESOLUTION_ERROR
+
+
+@pytest.mark.asyncio
+async def test_registry_schema_is_isolated_from_caller_mutation():
+    parameters = {
+        "type": "object",
+        "properties": {"query": {"type": "string"}},
+        "required": ["query"],
+        "additionalProperties": False,
+    }
+    registry = EventRegistry([_definition(parameters=parameters)])
+
+    parameters["properties"]["query"]["type"] = "integer"
+    parameters["required"].clear()
+    parameters["additionalProperties"] = True
+
+    assert registry.catalog()[0]["parameters"] == {
+        "type": "object",
+        "properties": {"query": {"type": "string"}},
+        "required": ["query"],
+        "additionalProperties": False,
+    }
+    registered = registry.definition("example.lookup")
+    assert registered is not None
+    with pytest.raises(TypeError):
+        registered.parameters["additionalProperties"] = True
+    with pytest.raises(TypeError):
+        registered.parameters["properties"]["query"]["type"] = "integer"
+    with pytest.raises(AttributeError):
+        registered.parameters["required"].append("unexpected")
+
+    result = await EventKernel(registry).execute(
+        EventRequest(
+            id="event-1",
+            name="example.lookup",
+            arguments={"query": 42, "unexpected": True},
+            source=EventSource.PROVIDER_RESOLVED,
+        ),
+        enabled_names=["example.lookup"],
+    )
+
+    assert result.status is EventStatus.INVALID_ARGUMENTS
+
+
 @pytest.mark.asyncio
 async def test_kernel_passes_consistent_fallback_arguments_and_raw_json_to_handler():
     captured: list[EventRequest] = []
@@ -764,7 +1001,7 @@ async def test_definition_metadata_survives_registration_and_result_creation():
     )
 
     registered = registry.definition("example.lookup")
-    assert registered is definition
+    assert registered is not definition
     assert result.result_policy is ResultPolicy.TEMPLATE_FOLLOW_UP
     assert result.confirmation_policy is ConfirmationPolicy.REQUIRED
     assert result.risk_level is RiskLevel.HIGH