Răsfoiți Sursa

fix: close event kernel quality boundaries

Problem: EventAgent discarded configured extra_body defaults, composed-schema required errors were misclassified, handler mutation could corrupt audit arguments, validator internals were exposed, and non-JSON handler payloads escaped the kernel.

Risk: recursive fallback classification plus defensive copying and strict serialization add CPU and memory cost; explicit empty extra_body remains an override, while complete-invalid and provider-resolved arguments must remain no-fallback.
zhenyu.hu 2 săptămâni în urmă
părinte
comite
58807c7d7c

+ 3 - 1
src/agent_lab/application/event_agent.py

@@ -78,7 +78,9 @@ class EventAgent:
             context=ToolExecutionContext(
                 history=history,
                 system_prompt=system_prompt,
-                extra_body=extra_body or {},
+                extra_body=(
+                    self.params.extra_body if extra_body is None else extra_body
+                ),
             ),
         )
         payload = self.registry.tool_payload(result)

+ 68 - 15
src/agent_lab/application/events/kernel.py

@@ -3,6 +3,7 @@ from __future__ import annotations
 import inspect
 import json
 from collections.abc import Awaitable, Callable, Iterable
+from copy import deepcopy
 from dataclasses import dataclass, replace
 from typing import Any
 
@@ -130,7 +131,7 @@ class EventKernel:
 
         resolved_request = replace(
             request,
-            arguments=resolved.arguments,
+            arguments=deepcopy(resolved.arguments),
             raw_arguments=resolved.raw_arguments,
         )
         try:
@@ -155,6 +156,17 @@ class EventKernel:
                 error="event handler returned non-object payload",
                 used_fallback=used_fallback,
             )
+        try:
+            payload = self._strict_json_copy(payload)
+        except (TypeError, ValueError):
+            return self._result(
+                definition,
+                request,
+                EventStatus.HANDLER_ERROR,
+                resolved=resolved,
+                error="event handler returned non-JSON payload",
+                used_fallback=used_fallback,
+            )
         return self._result(
             definition,
             request,
@@ -207,7 +219,7 @@ class EventKernel:
             )
         resolved_request = replace(
             request,
-            arguments=resolved.arguments,
+            arguments=deepcopy(resolved.arguments),
             raw_arguments=resolved.raw_arguments,
         )
         try:
@@ -238,6 +250,16 @@ class EventKernel:
                 resolved=resolved,
                 error="event handler returned non-object payload",
             )
+        try:
+            payload = self._strict_json_copy(payload)
+        except (TypeError, ValueError):
+            return self._result(
+                definition,
+                request,
+                EventStatus.HANDLER_ERROR,
+                resolved=resolved,
+                error="event handler returned non-JSON payload",
+            )
         return self._result(
             definition,
             request,
@@ -258,7 +280,7 @@ class EventKernel:
                 event_name=request.name,
                 status=EventStatus.UNKNOWN,
                 source=request.source,
-                arguments=dict(request.arguments),
+                arguments=deepcopy(request.arguments),
                 raw_arguments=request.raw_arguments,
                 error="unknown event",
             )
@@ -269,7 +291,7 @@ class EventKernel:
                 EventStatus.DISABLED,
                 resolved=ResolvedEventArguments(
                     event_name=request.name,
-                    arguments=dict(request.arguments),
+                    arguments=deepcopy(request.arguments),
                     raw_arguments=request.raw_arguments,
                 ),
                 error="event disabled",
@@ -286,7 +308,7 @@ class EventKernel:
             return (
                 ResolvedEventArguments(
                     event_name=request.name,
-                    arguments=dict(request.arguments),
+                    arguments=deepcopy(request.arguments),
                     raw_arguments=request.raw_arguments,
                 ),
                 True,
@@ -316,8 +338,9 @@ class EventKernel:
                 request,
                 "event argument resolver returned invalid payload",
             )
+        copied_arguments = deepcopy(arguments)
         try:
-            raw_arguments = self._json_arguments(arguments)
+            raw_arguments = self._json_arguments(copied_arguments)
         except (TypeError, ValueError) as exc:
             return None, False, self._resolution_error(
                 definition,
@@ -327,7 +350,7 @@ class EventKernel:
         return (
             ResolvedEventArguments(
                 event_name=request.name,
-                arguments=dict(arguments),
+                arguments=copied_arguments,
                 raw_arguments=raw_arguments,
             ),
             complete,
@@ -393,17 +416,27 @@ class EventKernel:
                 resolved=value,
                 used_fallback=True,
             )
-        return value, None
+        return (
+            ResolvedEventArguments(
+                event_name=value.event_name,
+                arguments=deepcopy(value.arguments),
+                raw_arguments=value.raw_arguments,
+            ),
+            None,
+        )
 
     def _validate(
         self,
         definition: EventDefinition,
         arguments: dict[str, Any],
     ) -> _ValidationResult:
-        validator = self.registry.validator(definition.name)
-        assert validator is not None
         try:
-            errors = list(validator.iter_errors(arguments))
+            errors = list(
+                self.registry.iter_validation_errors(
+                    definition.name,
+                    arguments,
+                )
+            )
         except Exception as exc:
             return _ValidationResult(
                 status=EventStatus.DEFINITION_ERROR,
@@ -411,7 +444,11 @@ class EventKernel:
             )
         if not errors:
             return _ValidationResult()
-        missing_errors = [error for error in errors if error.validator == "required"]
+        missing_errors = [
+            required_error
+            for error in errors
+            for required_error in _required_errors(error)
+        ]
         error = missing_errors[0] if missing_errors else errors[0]
         return _ValidationResult(
             status=EventStatus.INVALID_ARGUMENTS,
@@ -450,7 +487,7 @@ class EventKernel:
             resolved=resolved
             or ResolvedEventArguments(
                 event_name=request.name,
-                arguments=dict(request.arguments),
+                arguments=deepcopy(request.arguments),
                 raw_arguments=request.raw_arguments,
             ),
             error=error,
@@ -473,9 +510,9 @@ class EventKernel:
             event_name=request.name,
             status=status,
             source=request.source,
-            arguments=dict(resolved.arguments),
+            arguments=deepcopy(resolved.arguments),
             raw_arguments=resolved.raw_arguments,
-            payload=payload or {},
+            payload=deepcopy(payload or {}),
             error=error,
             used_fallback=used_fallback,
             result_policy=definition.result_policy,
@@ -496,11 +533,27 @@ class EventKernel:
             allow_nan=False,
         )
 
+    def _strict_json_copy(self, value: dict[str, Any]) -> dict[str, Any]:
+        copied = json.loads(
+            self._json_arguments(value),
+            parse_constant=_reject_json_constant,
+        )
+        if not isinstance(copied, dict) or not _json_values_equal(copied, value):
+            raise ValueError("JSON round-trip changed payload")
+        return copied
+
 
 def _reject_json_constant(value: str) -> None:
     raise ValueError(f"non-standard JSON constant: {value}")
 
 
+def _required_errors(error: ValidationError) -> Iterable[ValidationError]:
+    if error.validator == "required":
+        yield error
+    for nested_error in error.context:
+        yield from _required_errors(nested_error)
+
+
 def _json_values_equal(left: Any, right: Any) -> bool:
     if type(left) is not type(right):
         return False

+ 12 - 3
src/agent_lab/application/events/registry.py

@@ -7,7 +7,7 @@ from types import MappingProxyType
 from typing import Any
 
 from jsonschema import Draft202012Validator
-from jsonschema.exceptions import SchemaError
+from jsonschema.exceptions import SchemaError, ValidationError
 
 from agent_lab.application.events.models import EventDefinition
 
@@ -40,8 +40,17 @@ class EventRegistry:
     def definition(self, name: str) -> EventDefinition | None:
         return self._definitions.get(name)
 
-    def validator(self, name: str) -> Draft202012Validator | None:
-        return self._validators.get(name)
+    def iter_validation_errors(
+        self,
+        name: str,
+        arguments: dict[str, Any],
+    ) -> tuple[ValidationError, ...]:
+        validator = self._validators.get(name)
+        if validator is None:
+            raise LookupError(f"missing validator for event: {name}")
+        return tuple(
+            deepcopy(error) for error in validator.iter_errors(arguments)
+        )
 
     def catalog(self, enabled_names: Iterable[str] | None = None) -> list[dict]:
         enabled = set(enabled_names) if enabled_names is not None else None

+ 104 - 0
tests/test_event_agent.py

@@ -455,6 +455,37 @@ async def test_event_agent_returns_structured_error_when_tool_handler_raises():
     }
 
 
+@pytest.mark.asyncio
+async def test_event_agent_serializes_non_json_handler_payload_as_tool_error():
+    registry = ToolRegistry(
+        [
+            ToolDefinition(
+                name="bad_payload",
+                description="Return an invalid payload.",
+                parameters={"type": "object"},
+                handler=lambda event: {"invalid": object()},
+            )
+        ]
+    )
+
+    reply = await EventAgent(
+        enabled_tools=["bad_payload"],
+        registry=registry,
+    ).handle(
+        ToolCallEvent(
+            id="call-1",
+            name="bad_payload",
+            arguments={},
+            raw_arguments="{}",
+        )
+    )
+
+    assert json.loads(reply.content) == {
+        "tool": "bad_payload",
+        "error": "event handler returned non-JSON payload",
+    }
+
+
 @pytest.mark.asyncio
 async def test_event_agent_llm_receives_history_and_agent_config_context():
     chat_client = ToolCallingChatClient(
@@ -511,6 +542,79 @@ async def test_event_agent_llm_receives_history_and_agent_config_context():
     }
 
 
+@pytest.mark.asyncio
+async def test_event_agent_uses_configured_extra_body_when_override_is_omitted():
+    chat_client = ToolCallingChatClient({"message": "resolved"})
+    registry = ToolRegistry(
+        [
+            ToolDefinition(
+                name="handoff_note",
+                description="Send a note to the event agent.",
+                parameters={
+                    "type": "object",
+                    "properties": {"message": {"type": "string"}},
+                    "required": ["message"],
+                },
+                handler=lambda event: {"tool": event.name},
+                argument_resolver=lambda event, context: {},
+            )
+        ]
+    )
+
+    await EventAgent(
+        enabled_tools=["handoff_note"],
+        registry=registry,
+        chat_client=chat_client,
+        params=AgentParams(extra_body={"configured": True}),
+    ).handle(
+        ToolCallEvent(
+            id="call-1",
+            name="handoff_note",
+            arguments={},
+            raw_arguments="{}",
+        )
+    )
+
+    assert chat_client.calls[0]["params"].extra_body == {"configured": True}
+
+
+@pytest.mark.asyncio
+async def test_event_agent_preserves_explicit_empty_extra_body_override():
+    chat_client = ToolCallingChatClient({"message": "resolved"})
+    registry = ToolRegistry(
+        [
+            ToolDefinition(
+                name="handoff_note",
+                description="Send a note to the event agent.",
+                parameters={
+                    "type": "object",
+                    "properties": {"message": {"type": "string"}},
+                    "required": ["message"],
+                },
+                handler=lambda event: {"tool": event.name},
+                argument_resolver=lambda event, context: {},
+            )
+        ]
+    )
+
+    await EventAgent(
+        enabled_tools=["handoff_note"],
+        registry=registry,
+        chat_client=chat_client,
+        params=AgentParams(extra_body={"configured": True}),
+    ).handle(
+        ToolCallEvent(
+            id="call-1",
+            name="handoff_note",
+            arguments={},
+            raw_arguments="{}",
+        ),
+        extra_body={},
+    )
+
+    assert chat_client.calls[0]["params"].extra_body == {}
+
+
 @pytest.mark.asyncio
 async def test_event_agent_projects_complete_tool_round_to_visible_history():
     chat_client = ToolCallingChatClient({"message": "projected history"})

+ 214 - 0
tests/test_event_kernel.py

@@ -89,6 +89,28 @@ def test_registry_rejects_invalid_draft_2020_12_schema():
         EventRegistry([definition])
 
 
+def test_registry_does_not_expose_mutable_validator_instances():
+    registry = EventRegistry([_definition()])
+
+    assert not hasattr(registry, "validator")
+    errors = list(
+        registry.iter_validation_errors(
+            "example.lookup",
+            {"query": 42},
+        )
+    )
+    assert errors
+
+    errors[0].schema["type"] = "integer"
+
+    assert list(
+        registry.iter_validation_errors(
+            "example.lookup",
+            {"query": 42},
+        )
+    )
+
+
 def test_tool_registry_public_api_remains_compatible():
     registry = ToolRegistry(
         [
@@ -239,6 +261,108 @@ async def test_kernel_does_not_fallback_for_complete_invalid_arguments(
     assert fallback_calls == 0
 
 
+@pytest.mark.asyncio
+@pytest.mark.parametrize("composition", ["anyOf", "oneOf"])
+async def test_kernel_falls_back_for_required_missing_inside_composed_schema(
+    composition: str,
+):
+    fallback_calls = 0
+
+    async def fallback(*args: Any) -> ResolvedEventArguments:
+        nonlocal fallback_calls
+        fallback_calls += 1
+        return ResolvedEventArguments(
+            event_name="example.lookup",
+            arguments={"choice": {"mode": "alpha"}},
+            raw_arguments='{"choice":{"mode":"alpha"}}',
+        )
+
+    definition = _definition(
+        parameters={
+            "type": "object",
+            "properties": {
+                "choice": {
+                    composition: [
+                        {
+                            "type": "object",
+                            "properties": {"mode": {"const": "alpha"}},
+                            "required": ["mode"],
+                        },
+                        {
+                            "type": "object",
+                            "properties": {"mode": {"const": "beta"}},
+                            "required": ["mode"],
+                        },
+                    ]
+                }
+            },
+            "required": ["choice"],
+        },
+        resolver=lambda request, context: {"choice": {}},
+        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 result.used_fallback is True
+    assert fallback_calls == 1
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("composition", ["anyOf", "oneOf"])
+async def test_kernel_does_not_fallback_for_complete_invalid_composed_schema(
+    composition: str,
+):
+    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": {
+                "choice": {
+                    composition: [
+                        {
+                            "type": "object",
+                            "properties": {"mode": {"const": "alpha"}},
+                            "required": ["mode"],
+                        },
+                        {
+                            "type": "object",
+                            "properties": {"mode": {"const": "beta"}},
+                            "required": ["mode"],
+                        },
+                    ]
+                }
+            },
+            "required": ["choice"],
+        },
+        resolver=lambda request, context: {"choice": {"mode": "other"}},
+        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.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
@@ -905,6 +1029,64 @@ async def test_kernel_passes_consistent_fallback_arguments_and_raw_json_to_handl
     assert result.raw_arguments == captured[0].raw_arguments
 
 
+@pytest.mark.asyncio
+@pytest.mark.parametrize("execution", ["async", "sync"])
+async def test_kernel_isolates_nested_handler_mutation_from_audit_values(
+    execution: str,
+):
+    caller_arguments = {"nested": {"items": ["original"]}}
+    handler_arguments: list[dict[str, Any]] = []
+
+    def handler(request: EventRequest) -> dict[str, Any]:
+        handler_arguments.append(request.arguments)
+        request.arguments["nested"]["items"].append("handler")
+        return {"nested": request.arguments["nested"]}
+
+    definition = _definition(
+        parameters={
+            "type": "object",
+            "properties": {
+                "nested": {
+                    "type": "object",
+                    "properties": {
+                        "items": {"type": "array", "items": {"type": "string"}}
+                    },
+                    "required": ["items"],
+                }
+            },
+            "required": ["nested"],
+        },
+        handler=handler,
+    )
+    request = EventRequest(
+        id="event-1",
+        name="example.lookup",
+        arguments=caller_arguments,
+        raw_arguments='{"nested":{"items":["original"]}}',
+        source=EventSource.PROVIDER_RESOLVED,
+    )
+    kernel = EventKernel(EventRegistry([definition]))
+
+    result = (
+        await kernel.execute(request, enabled_names=["example.lookup"])
+        if execution == "async"
+        else kernel.execute_sync(request, enabled_names=["example.lookup"])
+    )
+
+    assert result.status is EventStatus.SUCCESS
+    assert caller_arguments == {"nested": {"items": ["original"]}}
+    assert request.arguments == {"nested": {"items": ["original"]}}
+    assert result.arguments == {"nested": {"items": ["original"]}}
+    assert json.loads(result.raw_arguments) == result.arguments
+    assert result.payload == {"nested": {"items": ["original", "handler"]}}
+
+    handler_arguments[0]["nested"]["items"].append("later")
+    assert result.payload == {"nested": {"items": ["original", "handler"]}}
+    result.payload["nested"]["items"].append("result")
+    assert result.arguments == {"nested": {"items": ["original"]}}
+    assert caller_arguments == {"nested": {"items": ["original"]}}
+
+
 @pytest.mark.asyncio
 @pytest.mark.parametrize("payload", [None, "text", 1, ["item"]])
 async def test_kernel_normalizes_invalid_handler_payloads(payload: Any):
@@ -924,6 +1106,38 @@ async def test_kernel_normalizes_invalid_handler_payloads(payload: Any):
     assert result.error == "event handler returned non-object payload"
 
 
+@pytest.mark.asyncio
+@pytest.mark.parametrize("execution", ["async", "sync"])
+@pytest.mark.parametrize(
+    "invalid_value",
+    [object(), {"set-item"}, float("nan"), float("inf")],
+)
+async def test_kernel_rejects_non_json_handler_dictionary_payloads(
+    execution: str,
+    invalid_value: Any,
+):
+    definition = _definition(
+        handler=lambda request: {"invalid": invalid_value},
+    )
+    request = EventRequest(
+        id="event-1",
+        name="example.lookup",
+        arguments={"query": "value"},
+        raw_arguments='{"query":"value"}',
+        source=EventSource.PROVIDER_RESOLVED,
+    )
+    kernel = EventKernel(EventRegistry([definition]))
+
+    result = (
+        await kernel.execute(request, enabled_names=["example.lookup"])
+        if execution == "async"
+        else kernel.execute_sync(request, enabled_names=["example.lookup"])
+    )
+
+    assert result.status is EventStatus.HANDLER_ERROR
+    assert result.error == "event handler returned non-JSON payload"
+
+
 @pytest.mark.asyncio
 async def test_async_kernel_supports_async_handler():
     async def handler(request: EventRequest) -> dict[str, Any]: