| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401 |
- 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
|