|
|
@@ -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]:
|