|
|
@@ -0,0 +1,401 @@
|
|
|
+from __future__ import annotations
|
|
|
+
|
|
|
+from typing import Any
|
|
|
+
|
|
|
+import pytest
|
|
|
+
|
|
|
+from agent_lab.application.events import (
|
|
|
+ ConfirmationPolicy,
|
|
|
+ EventDefinition,
|
|
|
+ EventExecutionContext,
|
|
|
+ EventKernel,
|
|
|
+ EventRegistry,
|
|
|
+ EventRequest,
|
|
|
+ EventSource,
|
|
|
+ EventStatus,
|
|
|
+ ResultPolicy,
|
|
|
+ RiskLevel,
|
|
|
+)
|
|
|
+from agent_lab.application.tools import (
|
|
|
+ ToolDefinition,
|
|
|
+ ToolExecutionContext,
|
|
|
+ ToolRegistry,
|
|
|
+)
|
|
|
+from agent_lab.domain.events import ToolCallEvent
|
|
|
+from agent_lab.domain.messages import ChatMessage
|
|
|
+
|
|
|
+
|
|
|
+def _definition(
|
|
|
+ name: str = "example.lookup",
|
|
|
+ **overrides: Any,
|
|
|
+) -> EventDefinition:
|
|
|
+ values: dict[str, Any] = {
|
|
|
+ "name": name,
|
|
|
+ "description": "Look up an example value.",
|
|
|
+ "parameters": {
|
|
|
+ "type": "object",
|
|
|
+ "properties": {"query": {"type": "string"}},
|
|
|
+ "required": ["query"],
|
|
|
+ },
|
|
|
+ "handler": lambda request: {
|
|
|
+ "event": request.name,
|
|
|
+ "query": request.arguments["query"],
|
|
|
+ },
|
|
|
+ }
|
|
|
+ values.update(overrides)
|
|
|
+ return EventDefinition(**values)
|
|
|
+
|
|
|
+
|
|
|
+def test_registry_registers_flat_definitions_and_filters_enabled_catalog():
|
|
|
+ registry = EventRegistry(
|
|
|
+ [_definition("example.lookup"), _definition("device.inspect")]
|
|
|
+ )
|
|
|
+
|
|
|
+ assert [item["name"] for item in registry.catalog()] == [
|
|
|
+ "example.lookup",
|
|
|
+ "device.inspect",
|
|
|
+ ]
|
|
|
+ assert registry.catalog(["device.inspect"]) == [
|
|
|
+ {
|
|
|
+ "name": "device.inspect",
|
|
|
+ "description": "Look up an example value.",
|
|
|
+ "parameters": {
|
|
|
+ "type": "object",
|
|
|
+ "properties": {"query": {"type": "string"}},
|
|
|
+ "required": ["query"],
|
|
|
+ },
|
|
|
+ }
|
|
|
+ ]
|
|
|
+ assert registry.tool_schema("example.lookup")["function"]["name"] == (
|
|
|
+ "example.lookup"
|
|
|
+ )
|
|
|
+ assert registry.tool_schema("missing") is None
|
|
|
+
|
|
|
+
|
|
|
+def test_registry_rejects_duplicate_definition_names():
|
|
|
+ with pytest.raises(ValueError, match="duplicate event definition: example.lookup"):
|
|
|
+ EventRegistry([_definition(), _definition()])
|
|
|
+
|
|
|
+
|
|
|
+def test_tool_registry_public_api_remains_compatible():
|
|
|
+ registry = ToolRegistry(
|
|
|
+ [
|
|
|
+ ToolDefinition(
|
|
|
+ name="compat.lookup",
|
|
|
+ description="Look up compatibility data.",
|
|
|
+ parameters={
|
|
|
+ "type": "object",
|
|
|
+ "properties": {"query": {"type": "string"}},
|
|
|
+ "required": ["query"],
|
|
|
+ },
|
|
|
+ handler=lambda event: {
|
|
|
+ "tool": event.name,
|
|
|
+ "query": event.arguments["query"],
|
|
|
+ },
|
|
|
+ argument_resolver=lambda event, context: {
|
|
|
+ "query": context.history[-1].content
|
|
|
+ },
|
|
|
+ )
|
|
|
+ ]
|
|
|
+ )
|
|
|
+ event = ToolCallEvent(
|
|
|
+ id="call-1",
|
|
|
+ name="compat.lookup",
|
|
|
+ arguments={"query": "provider"},
|
|
|
+ raw_arguments='{"query":"provider"}',
|
|
|
+ )
|
|
|
+
|
|
|
+ assert registry.available_tools() == [
|
|
|
+ {
|
|
|
+ "name": "compat.lookup",
|
|
|
+ "description": "Look up compatibility data.",
|
|
|
+ "parameters": {
|
|
|
+ "type": "object",
|
|
|
+ "properties": {"query": {"type": "string"}},
|
|
|
+ "required": ["query"],
|
|
|
+ },
|
|
|
+ }
|
|
|
+ ]
|
|
|
+ assert "- compat.lookup: Look up compatibility data." in (
|
|
|
+ registry.chat_event_system_message(["compat.lookup"])
|
|
|
+ )
|
|
|
+ assert registry.tool_schema("compat.lookup")["function"]["name"] == (
|
|
|
+ "compat.lookup"
|
|
|
+ )
|
|
|
+ assert registry.handle(
|
|
|
+ event,
|
|
|
+ ToolExecutionContext(history=[ChatMessage(role="user", content="history")]),
|
|
|
+ ) == {"tool": "compat.lookup", "query": "history"}
|
|
|
+ assert registry.execute(event) == {"tool": "compat.lookup", "query": "provider"}
|
|
|
+
|
|
|
+
|
|
|
+@pytest.mark.asyncio
|
|
|
+async def test_kernel_executes_complete_deterministic_arguments_without_fallback():
|
|
|
+ fallback_calls: list[str] = []
|
|
|
+
|
|
|
+ async def fallback(*args: Any) -> dict[str, Any]:
|
|
|
+ fallback_calls.append("called")
|
|
|
+ return {"query": "fallback"}
|
|
|
+
|
|
|
+ registry = EventRegistry(
|
|
|
+ [_definition(resolver=lambda request, context: {"query": "deterministic"})]
|
|
|
+ )
|
|
|
+
|
|
|
+ result = await EventKernel(registry, 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": "deterministic"}
|
|
|
+ assert result.payload == {"event": "example.lookup", "query": "deterministic"}
|
|
|
+ assert result.used_fallback is False
|
|
|
+ assert fallback_calls == []
|
|
|
+
|
|
|
+
|
|
|
+@pytest.mark.asyncio
|
|
|
+async def test_kernel_calls_fallback_once_when_required_arguments_are_incomplete():
|
|
|
+ fallback_calls: list[dict[str, Any]] = []
|
|
|
+
|
|
|
+ async def fallback(
|
|
|
+ definition: EventDefinition,
|
|
|
+ request: EventRequest,
|
|
|
+ context: EventExecutionContext,
|
|
|
+ ) -> dict[str, Any]:
|
|
|
+ fallback_calls.append(dict(request.arguments))
|
|
|
+ return {"query": "resolved once"}
|
|
|
+
|
|
|
+ registry = EventRegistry([_definition(resolver=lambda request, context: {})])
|
|
|
+
|
|
|
+ result = await EventKernel(registry, 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 once"}
|
|
|
+ assert result.used_fallback is True
|
|
|
+ assert fallback_calls == [{}]
|
|
|
+
|
|
|
+
|
|
|
+@pytest.mark.asyncio
|
|
|
+async def test_kernel_does_not_fallback_when_definition_disallows_it():
|
|
|
+ fallback_calls = 0
|
|
|
+
|
|
|
+ async def fallback(*args: Any) -> dict[str, Any]:
|
|
|
+ nonlocal fallback_calls
|
|
|
+ fallback_calls += 1
|
|
|
+ return {"query": "not allowed"}
|
|
|
+
|
|
|
+ registry = EventRegistry(
|
|
|
+ [
|
|
|
+ _definition(
|
|
|
+ resolver=lambda request, context: {},
|
|
|
+ fallback_allowed=False,
|
|
|
+ )
|
|
|
+ ]
|
|
|
+ )
|
|
|
+
|
|
|
+ result = await EventKernel(registry, argument_fallback=fallback).execute(
|
|
|
+ EventRequest(id="event-1", name="example.lookup"),
|
|
|
+ enabled_names=["example.lookup"],
|
|
|
+ )
|
|
|
+
|
|
|
+ assert result.status is EventStatus.INVALID_ARGUMENTS
|
|
|
+ assert result.error == "missing required arguments: query"
|
|
|
+ assert result.used_fallback is False
|
|
|
+ assert fallback_calls == 0
|
|
|
+
|
|
|
+
|
|
|
+@pytest.mark.asyncio
|
|
|
+async def test_provider_resolved_arguments_are_not_rewritten_or_fallen_back():
|
|
|
+ resolver_calls = 0
|
|
|
+ fallback_calls = 0
|
|
|
+
|
|
|
+ def resolver(*args: Any) -> dict[str, Any]:
|
|
|
+ nonlocal resolver_calls
|
|
|
+ resolver_calls += 1
|
|
|
+ return {"query": "rewritten"}
|
|
|
+
|
|
|
+ async def fallback(*args: Any) -> dict[str, Any]:
|
|
|
+ nonlocal fallback_calls
|
|
|
+ fallback_calls += 1
|
|
|
+ return {"query": "fallback"}
|
|
|
+
|
|
|
+ registry = EventRegistry([_definition(resolver=resolver)])
|
|
|
+
|
|
|
+ result = await EventKernel(registry, argument_fallback=fallback).execute(
|
|
|
+ EventRequest(
|
|
|
+ id="event-1",
|
|
|
+ name="example.lookup",
|
|
|
+ arguments={"query": "provider value"},
|
|
|
+ source=EventSource.PROVIDER_RESOLVED,
|
|
|
+ ),
|
|
|
+ enabled_names=["example.lookup"],
|
|
|
+ )
|
|
|
+
|
|
|
+ assert result.status is EventStatus.SUCCESS
|
|
|
+ assert result.arguments == {"query": "provider value"}
|
|
|
+ assert resolver_calls == 0
|
|
|
+ assert fallback_calls == 0
|
|
|
+
|
|
|
+
|
|
|
+@pytest.mark.asyncio
|
|
|
+async def test_provider_resolved_missing_arguments_return_invalid_without_fallback():
|
|
|
+ fallback_calls = 0
|
|
|
+
|
|
|
+ async def fallback(*args: Any) -> dict[str, Any]:
|
|
|
+ nonlocal fallback_calls
|
|
|
+ fallback_calls += 1
|
|
|
+ return {"query": "fallback"}
|
|
|
+
|
|
|
+ result = await EventKernel(
|
|
|
+ EventRegistry([_definition()]), argument_fallback=fallback
|
|
|
+ ).execute(
|
|
|
+ EventRequest(
|
|
|
+ id="event-1",
|
|
|
+ name="example.lookup",
|
|
|
+ arguments={},
|
|
|
+ source=EventSource.PROVIDER_RESOLVED,
|
|
|
+ ),
|
|
|
+ enabled_names=["example.lookup"],
|
|
|
+ )
|
|
|
+
|
|
|
+ assert result.status is EventStatus.INVALID_ARGUMENTS
|
|
|
+ assert result.error == "missing required arguments: query"
|
|
|
+ assert fallback_calls == 0
|
|
|
+
|
|
|
+
|
|
|
+@pytest.mark.asyncio
|
|
|
+@pytest.mark.parametrize(
|
|
|
+ ("event_request", "enabled_names", "expected_status", "expected_error"),
|
|
|
+ [
|
|
|
+ (
|
|
|
+ EventRequest(id="event-1", name="missing"),
|
|
|
+ ["missing"],
|
|
|
+ EventStatus.UNKNOWN,
|
|
|
+ "unknown event",
|
|
|
+ ),
|
|
|
+ (
|
|
|
+ EventRequest(id="event-1", name="example.lookup"),
|
|
|
+ [],
|
|
|
+ EventStatus.DISABLED,
|
|
|
+ "event disabled",
|
|
|
+ ),
|
|
|
+ (
|
|
|
+ EventRequest(
|
|
|
+ id="event-1",
|
|
|
+ name="example.lookup",
|
|
|
+ arguments={"query": 42},
|
|
|
+ source=EventSource.PROVIDER_RESOLVED,
|
|
|
+ ),
|
|
|
+ ["example.lookup"],
|
|
|
+ EventStatus.INVALID_ARGUMENTS,
|
|
|
+ "invalid argument type for query: expected string",
|
|
|
+ ),
|
|
|
+ ],
|
|
|
+)
|
|
|
+async def test_kernel_normalizes_lookup_and_validation_failures(
|
|
|
+ event_request: EventRequest,
|
|
|
+ enabled_names: list[str],
|
|
|
+ expected_status: EventStatus,
|
|
|
+ expected_error: str,
|
|
|
+):
|
|
|
+ result = await EventKernel(EventRegistry([_definition()])).execute(
|
|
|
+ event_request,
|
|
|
+ enabled_names=enabled_names,
|
|
|
+ )
|
|
|
+
|
|
|
+ assert result.status is expected_status
|
|
|
+ assert result.error == expected_error
|
|
|
+
|
|
|
+
|
|
|
+@pytest.mark.asyncio
|
|
|
+async def test_kernel_normalizes_handler_exceptions():
|
|
|
+ def fail(request: EventRequest) -> dict[str, Any]:
|
|
|
+ raise RuntimeError("boom")
|
|
|
+
|
|
|
+ result = await EventKernel(
|
|
|
+ EventRegistry([_definition(handler=fail)])
|
|
|
+ ).execute(
|
|
|
+ EventRequest(
|
|
|
+ id="event-1",
|
|
|
+ name="example.lookup",
|
|
|
+ arguments={"query": "value"},
|
|
|
+ source=EventSource.PROVIDER_RESOLVED,
|
|
|
+ ),
|
|
|
+ enabled_names=["example.lookup"],
|
|
|
+ )
|
|
|
+
|
|
|
+ assert result.status is EventStatus.HANDLER_ERROR
|
|
|
+ assert result.error == "event handler failed: boom"
|
|
|
+
|
|
|
+
|
|
|
+@pytest.mark.asyncio
|
|
|
+async def test_kernel_rejects_boolean_for_json_number_arguments():
|
|
|
+ definition = _definition(
|
|
|
+ parameters={
|
|
|
+ "type": "object",
|
|
|
+ "properties": {"query": {"type": "number"}},
|
|
|
+ "required": ["query"],
|
|
|
+ }
|
|
|
+ )
|
|
|
+
|
|
|
+ result = await EventKernel(EventRegistry([definition])).execute(
|
|
|
+ EventRequest(
|
|
|
+ id="event-1",
|
|
|
+ name="example.lookup",
|
|
|
+ arguments={"query": True},
|
|
|
+ source=EventSource.PROVIDER_RESOLVED,
|
|
|
+ ),
|
|
|
+ enabled_names=["example.lookup"],
|
|
|
+ )
|
|
|
+
|
|
|
+ assert result.status is EventStatus.INVALID_ARGUMENTS
|
|
|
+ assert result.error == "invalid argument type for query: expected number"
|
|
|
+
|
|
|
+
|
|
|
+@pytest.mark.asyncio
|
|
|
+async def test_definition_metadata_survives_registration_and_result_creation():
|
|
|
+ definition = _definition(
|
|
|
+ result_policy=ResultPolicy.TEMPLATE_FOLLOW_UP,
|
|
|
+ confirmation_policy=ConfirmationPolicy.REQUIRED,
|
|
|
+ risk_level=RiskLevel.HIGH,
|
|
|
+ idempotency_key_fields=("session_id", "event_id"),
|
|
|
+ concurrency_class="device-write",
|
|
|
+ conflict_keys=("device",),
|
|
|
+ timeout_seconds=1.5,
|
|
|
+ terminal=True,
|
|
|
+ )
|
|
|
+ registry = EventRegistry([definition])
|
|
|
+
|
|
|
+ result = await EventKernel(registry).execute(
|
|
|
+ EventRequest(
|
|
|
+ id="event-1",
|
|
|
+ name="example.lookup",
|
|
|
+ arguments={"query": "value"},
|
|
|
+ source=EventSource.PROVIDER_RESOLVED,
|
|
|
+ ),
|
|
|
+ enabled_names=["example.lookup"],
|
|
|
+ )
|
|
|
+
|
|
|
+ registered = registry.definition("example.lookup")
|
|
|
+ assert registered is definition
|
|
|
+ assert result.result_policy is ResultPolicy.TEMPLATE_FOLLOW_UP
|
|
|
+ assert result.confirmation_policy is ConfirmationPolicy.REQUIRED
|
|
|
+ assert result.risk_level is RiskLevel.HIGH
|
|
|
+ assert result.idempotency_key_fields == ("session_id", "event_id")
|
|
|
+ assert result.concurrency_class == "device-write"
|
|
|
+ assert result.conflict_keys == ("device",)
|
|
|
+ assert result.timeout_seconds == 1.5
|
|
|
+ assert result.terminal is True
|
|
|
+
|
|
|
+
|
|
|
+def test_kernel_source_has_no_builtin_event_name_branches():
|
|
|
+ from pathlib import Path
|
|
|
+
|
|
|
+ source = Path("src/agent_lab/application/events/kernel.py").read_text()
|
|
|
+
|
|
|
+ assert "handoff_note" not in source
|
|
|
+ assert "mock_search" not in source
|
|
|
+ assert "mock_ticket" not in source
|