Просмотр исходного кода

fix: contain mixed validation failures

Problem: mixed missing and substantive validation failures were treated as fallback-eligible, and arbitrary handler payload snapshot exceptions leaked across kernel integrations.

Risk: every top-level validation issue must now be missing-only before fallback; plugin dictionary subclasses that fail serialization are normalized while BaseException continues to propagate.
zhenyu.hu 2 недель назад
Родитель
Сommit
3cffbc4519
3 измененных файлов с 234 добавлено и 12 удалено
  1. 18 12
      src/agent_lab/application/events/kernel.py
  2. 37 0
      tests/test_event_agent.py
  3. 179 0
      tests/test_event_kernel.py

+ 18 - 12
src/agent_lab/application/events/kernel.py

@@ -155,7 +155,7 @@ class EventKernel:
             )
         try:
             payload = self._strict_json_copy(payload)
-        except (TypeError, ValueError):
+        except Exception:
             return self._result(
                 definition,
                 request,
@@ -249,7 +249,7 @@ class EventKernel:
             )
         try:
             payload = self._strict_json_copy(payload)
-        except (TypeError, ValueError):
+        except Exception:
             return self._result(
                 definition,
                 request,
@@ -455,20 +455,26 @@ class EventKernel:
             )
         if not errors:
             return _ValidationResult()
-        missing_error = next(
-            (
-                required_error
-                for error in errors
-                if (required_error := _fallback_required_error(error))
-                is not None
-            ),
-            None,
+        classified_errors = [
+            (error, _fallback_required_error(error)) for error in errors
+        ]
+        missing_required = all(
+            required_error is not None
+            for _, required_error in classified_errors
         )
-        error = missing_error or errors[0]
+        if missing_required:
+            error = classified_errors[0][1]
+            assert error is not None
+        else:
+            error = next(
+                issue
+                for issue, required_error in classified_errors
+                if required_error is None
+            )
         return _ValidationResult(
             status=EventStatus.INVALID_ARGUMENTS,
             error=self._format_validation_error(error),
-            missing_required=missing_error is not None,
+            missing_required=missing_required,
         )
 
     def _format_validation_error(self, error: ValidationIssue) -> str:

+ 37 - 0
tests/test_event_agent.py

@@ -486,6 +486,43 @@ async def test_event_agent_serializes_non_json_handler_payload_as_tool_error():
     }
 
 
+@pytest.mark.asyncio
+async def test_event_agent_normalizes_handler_payload_snapshot_exceptions():
+    class ExplodingItemsDict(dict):
+        def items(self):
+            raise RuntimeError("payload items failed")
+
+    registry = ToolRegistry(
+        [
+            ToolDefinition(
+                name="bad_payload",
+                description="Return a payload that fails during snapshot.",
+                parameters={"type": "object"},
+                handler=lambda event: ExplodingItemsDict(ok=True),
+            )
+        ]
+    )
+
+    reply = await EventAgent(
+        enabled_tools=["bad_payload"],
+        registry=registry,
+    ).handle(
+        ToolCallEvent(
+            id="call-1",
+            name="bad_payload",
+            arguments={},
+            raw_arguments="{}",
+        )
+    )
+
+    assert reply.role == "tool"
+    assert reply.tool_call_id == "call-1"
+    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(

+ 179 - 0
tests/test_event_kernel.py

@@ -265,6 +265,63 @@ async def test_kernel_does_not_fallback_for_complete_invalid_arguments(
     assert fallback_calls == 0
 
 
+@pytest.mark.asyncio
+@pytest.mark.parametrize(
+    ("arguments", "expected_error"),
+    [
+        (
+            {"mode": "invalid"},
+            "invalid event arguments: 'invalid' is not one of ['valid']",
+        ),
+        (
+            {"mode": "valid", "count": "invalid"},
+            "invalid argument type for count: expected integer",
+        ),
+        (
+            {"mode": "valid", "unexpected": True},
+            "invalid event arguments: Additional properties are not allowed "
+            "('unexpected' was unexpected)",
+        ),
+    ],
+)
+async def test_kernel_does_not_fallback_for_required_plus_substantive_error(
+    arguments: dict[str, Any],
+    expected_error: 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": {
+                "query": {"type": "string"},
+                "mode": {"enum": ["valid"]},
+                "count": {"type": "integer"},
+            },
+            "required": ["query"],
+            "additionalProperties": False,
+        },
+        resolver=lambda request, context: arguments,
+    )
+
+    result = await EventKernel(
+        EventRegistry([definition]), argument_fallback=fallback
+    ).execute(
+        EventRequest(id="event-1", name=definition.name),
+        enabled_names=[definition.name],
+    )
+
+    assert result.status is EventStatus.INVALID_ARGUMENTS
+    assert result.error == expected_error
+    assert result.used_fallback is False
+    assert fallback_calls == 0
+
+
 def _discriminated_composed_parameters(composition: str) -> dict[str, Any]:
     return {
         "type": "object",
@@ -392,6 +449,39 @@ async def test_kernel_does_not_fallback_when_composed_branch_is_ambiguous(
     assert fallback_calls == 0
 
 
+@pytest.mark.asyncio
+async def test_kernel_does_not_fallback_for_substantive_and_composed_missing_errors():
+    fallback_calls = 0
+
+    async def fallback(*args: Any) -> ResolvedEventArguments:
+        nonlocal fallback_calls
+        fallback_calls += 1
+        raise AssertionError("fallback should not run")
+
+    parameters = _discriminated_composed_parameters("anyOf")
+    parameters["properties"]["mode"] = {"enum": ["valid"]}
+    definition = _definition(
+        parameters=parameters,
+        resolver=lambda request, context: {
+            "mode": "invalid",
+            "choice": {"kind": "a"},
+        },
+        handler=lambda request: {"ok": True},
+    )
+
+    result = await EventKernel(
+        EventRegistry([definition]), argument_fallback=fallback
+    ).execute(
+        EventRequest(id="event-1", name=definition.name),
+        enabled_names=[definition.name],
+    )
+
+    assert result.status is EventStatus.INVALID_ARGUMENTS
+    assert result.error == "invalid event arguments: 'invalid' is not one of ['valid']"
+    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
@@ -1234,6 +1324,95 @@ async def test_kernel_rejects_non_json_handler_dictionary_payloads(
     assert result.error == "event handler returned non-JSON payload"
 
 
+class _ExplodingItemsDict(dict[str, Any]):
+    def items(self):
+        raise RuntimeError("payload items failed")
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("execution", ["async", "sync"])
+async def test_kernel_normalizes_handler_payload_snapshot_exceptions(
+    execution: str,
+):
+    definition = _definition(
+        handler=lambda request: _ExplodingItemsDict(ok=True),
+    )
+    request = EventRequest(
+        id="event-1",
+        name=definition.name,
+        arguments={"query": "value"},
+        raw_arguments='{"query":"value"}',
+        source=EventSource.PROVIDER_RESOLVED,
+    )
+    kernel = EventKernel(EventRegistry([definition]))
+
+    result = (
+        await kernel.execute(request, enabled_names=[definition.name])
+        if execution == "async"
+        else kernel.execute_sync(request, enabled_names=[definition.name])
+    )
+
+    assert result.status is EventStatus.HANDLER_ERROR
+    assert result.error == "event handler returned non-JSON payload"
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize("execution", ["async", "sync"])
+async def test_kernel_does_not_swallow_base_exception_from_payload_snapshot(
+    execution: str,
+):
+    class SnapshotAbort(BaseException):
+        pass
+
+    class AbortingItemsDict(dict[str, Any]):
+        def items(self):
+            raise SnapshotAbort
+
+    definition = _definition(
+        handler=lambda request: AbortingItemsDict(ok=True),
+    )
+    request = EventRequest(
+        id="event-1",
+        name=definition.name,
+        arguments={"query": "value"},
+        source=EventSource.PROVIDER_RESOLVED,
+    )
+    kernel = EventKernel(EventRegistry([definition]))
+
+    with pytest.raises(SnapshotAbort):
+        if execution == "async":
+            await kernel.execute(request, enabled_names=[definition.name])
+        else:
+            kernel.execute_sync(request, enabled_names=[definition.name])
+
+
+@pytest.mark.asyncio
+async def test_tool_registry_normalizes_handler_payload_snapshot_exceptions():
+    registry = ToolRegistry(
+        [
+            ToolDefinition(
+                name="bad_payload",
+                description="Return a payload that fails during snapshot.",
+                parameters={"type": "object"},
+                handler=lambda event: _ExplodingItemsDict(ok=True),
+            )
+        ]
+    )
+    event = ToolCallEvent(
+        id="call-1",
+        name="bad_payload",
+        arguments={},
+        raw_arguments="{}",
+    )
+    expected = {
+        "tool": "bad_payload",
+        "error": "event handler returned non-JSON payload",
+    }
+
+    assert registry.execute(event) == expected
+    assert await registry.execute_async(event) == expected
+
+
 @pytest.mark.asyncio
 async def test_async_kernel_supports_async_handler():
     async def handler(request: EventRequest) -> dict[str, Any]: