refactor(trigger): share automation turn delivery

This commit is contained in:
chengyongru
2026-07-02 13:32:46 +08:00
committed by Xubin Ren
parent afef27dd6c
commit acb0e853ff
16 changed files with 413 additions and 299 deletions
+27 -14
View File
@@ -9,8 +9,8 @@ from typing import Any
from loguru import logger
from nanobot.agent.automation_turns import AutomationTurnError
from nanobot.bus.events import InboundMessage, OutboundMessage
from nanobot.bus.queue import MessageBus
from nanobot.triggers.local_session_turns import LOCAL_TRIGGER_META
from nanobot.triggers.local_store import LocalTriggerStore
from nanobot.triggers.local_types import LocalTrigger, TriggerDelivery
@@ -20,14 +20,13 @@ from nanobot.webui.metadata import WEBUI_MESSAGE_SOURCE_METADATA_KEY, WEBUI_TURN
async def run_local_trigger_queue(
*,
store: LocalTriggerStore,
bus: MessageBus | None = None,
submit_turn: Callable[[InboundMessage], Awaitable[OutboundMessage | None]] | None = None,
poll_interval_s: float = 0.5,
batch_size: int = 20,
) -> None:
"""Poll local trigger deliveries and publish them as normal inbound messages."""
if bus is None and submit_turn is None:
raise ValueError("run_local_trigger_queue requires bus or submit_turn")
"""Poll local trigger deliveries and submit them as session turns."""
if submit_turn is None:
raise ValueError("run_local_trigger_queue requires submit_turn")
logger.info("Local trigger queue started")
recovered = store.recover_processing_deliveries()
if recovered:
@@ -46,7 +45,6 @@ async def run_local_trigger_queue(
await _deliver_delivery(
store,
delivery,
bus=bus,
submit_turn=submit_turn,
)
store.complete_delivery(delivery)
@@ -67,6 +65,21 @@ async def run_local_trigger_queue(
delivery.trigger_id,
exc,
)
except AutomationTurnError as exc:
error = str(exc) or exc.__class__.__name__
store.record_delivery(
delivery.trigger_id,
status="error",
error=error,
run_at_ms=delivery.created_at_ms,
)
store.complete_delivery(delivery)
logger.warning(
"Trigger: delivery {} for {} reached the agent but failed: {}",
delivery.id,
delivery.trigger_id,
error,
)
except Exception as exc:
error = str(exc) or exc.__class__.__name__
retried = store.retry_delivery(delivery, error)
@@ -92,8 +105,7 @@ async def _deliver_delivery(
store: LocalTriggerStore,
delivery: TriggerDelivery,
*,
bus: MessageBus | None,
submit_turn: Callable[[InboundMessage], Awaitable[OutboundMessage | None]] | None,
submit_turn: Callable[[InboundMessage], Awaitable[OutboundMessage | None]],
) -> None:
trigger = store.get(delivery.trigger_id)
if trigger is None:
@@ -109,12 +121,7 @@ async def _deliver_delivery(
metadata=_delivery_metadata(trigger, delivery),
session_key_override=trigger.session_key,
)
if submit_turn is not None:
await submit_turn(msg)
else:
if bus is None:
raise RuntimeError("bus unavailable for local trigger delivery")
await bus.publish_inbound(msg)
await submit_turn(msg)
store.record_delivery(
trigger.id,
status="ok",
@@ -129,6 +136,7 @@ def _delivery_metadata(trigger: LocalTrigger, delivery: TriggerDelivery) -> dict
"trigger_name": trigger.name,
"delivery_id": delivery.id,
"created_at_ms": delivery.created_at_ms,
"persist_content": _history_content(trigger, delivery),
}
if trigger.channel == "websocket":
metadata.pop(WEBUI_TURN_METADATA_KEY, None)
@@ -138,3 +146,8 @@ def _delivery_metadata(trigger: LocalTrigger, delivery: TriggerDelivery) -> dict
source["label"] = trigger.name
metadata[WEBUI_MESSAGE_SOURCE_METADATA_KEY] = source
return metadata
def _history_content(trigger: LocalTrigger, delivery: TriggerDelivery) -> str:
label = trigger.name.strip() if trigger.name else trigger.id
return f"Local trigger received: {label}\n\n{delivery.content}"
+3
View File
@@ -14,6 +14,9 @@ LOCAL_TRIGGER_META = "_local_trigger"
def _local_trigger_history_text(trigger: Mapping[str, Any]) -> str:
persist_content = trigger.get("persist_content")
if isinstance(persist_content, str) and persist_content.strip():
return persist_content
name = trigger.get("trigger_name")
trigger_id = trigger.get("trigger_id")
label = name if isinstance(name, str) and name.strip() else trigger_id
+25 -106
View File
@@ -2,15 +2,14 @@
from __future__ import annotations
import asyncio
import dataclasses
from collections.abc import Awaitable, Callable, Iterable
from nanobot.bus.events import InboundMessage, OutboundMessage
from nanobot.agent.automation_turns import AutomationTurnCoordinator
from nanobot.bus.events import InboundMessage
from nanobot.triggers.local_session_turns import local_trigger, local_trigger_delivery_id
class LocalTriggerTurnCoordinator:
class LocalTriggerTurnCoordinator(AutomationTurnCoordinator):
"""Manage local trigger turns without mixing them into live injections."""
def __init__(
@@ -19,113 +18,33 @@ class LocalTriggerTurnCoordinator:
publish_inbound: Callable[[InboundMessage], Awaitable[None]],
dispatch: Callable[[InboundMessage], Awaitable[object]],
is_running: Callable[[], bool],
deferred_queues: dict[str, list[InboundMessage]] | None = None,
) -> None:
self._publish_inbound = publish_inbound
self._dispatch = dispatch
self._is_running = is_running
self.deferred_queues: dict[str, list[InboundMessage]] = {}
self._waiters: dict[str, asyncio.Future[OutboundMessage | None]] = {}
self._pending_messages_by_delivery_id: dict[str, InboundMessage] = {}
async def submit(self, msg: InboundMessage) -> OutboundMessage | None:
"""Submit a local trigger turn and wait for its session response."""
delivery_id = local_trigger_delivery_id(msg.metadata)
if not delivery_id:
raise ValueError("local trigger turn metadata must include a delivery_id")
if delivery_id in self._waiters:
raise RuntimeError(f"local trigger delivery {delivery_id!r} is already pending")
loop = asyncio.get_running_loop()
future: asyncio.Future[OutboundMessage | None] = loop.create_future()
self._waiters[delivery_id] = future
self._pending_messages_by_delivery_id[delivery_id] = msg
try:
if self._is_running():
await self._publish_inbound(msg)
else:
await self._dispatch(msg)
return await future
finally:
self._waiters.pop(delivery_id, None)
self._pending_messages_by_delivery_id.pop(delivery_id, None)
def should_defer(
self,
msg: InboundMessage,
*,
session_key: str,
active_session_keys: Iterable[str],
) -> bool:
return local_trigger(msg.metadata) is not None and session_key in active_session_keys
def defer_if_active(
self,
msg: InboundMessage,
*,
session_key: str,
active_session_keys: Iterable[str],
) -> bool:
"""Defer a local trigger turn when its target session is already active."""
if not self.should_defer(
msg,
session_key=session_key,
active_session_keys=active_session_keys,
):
return False
pending_msg = msg
if session_key != msg.session_key:
pending_msg = dataclasses.replace(
msg,
session_key_override=session_key,
)
self.defer(session_key, pending_msg)
return True
def complete(
self,
msg: InboundMessage,
*,
response: OutboundMessage | None = None,
error: BaseException | None = None,
) -> None:
delivery_id = local_trigger_delivery_id(msg.metadata)
if not delivery_id:
return
future = self._waiters.get(delivery_id)
if future is None or future.done():
return
if error is not None:
future.set_exception(error)
else:
future.set_result(response)
def defer(self, session_key: str, msg: InboundMessage) -> None:
self.deferred_queues.setdefault(session_key, []).append(msg)
super().__init__(
publish_inbound=publish_inbound,
dispatch=dispatch,
is_running=is_running,
turn_id=lambda msg: local_trigger_delivery_id(msg.metadata),
pending_id=_local_trigger_id,
should_defer_turn=_should_defer_local_trigger_turn,
missing_id_error="local trigger turn metadata must include a delivery_id",
duplicate_id_error=lambda delivery_id: (
f"local trigger delivery {delivery_id!r} is already pending"
),
deferred_queues=deferred_queues,
)
def pending_trigger_ids_for_session(self, session_key: str) -> set[str]:
"""Return local triggers waiting for or running in *session_key*."""
trigger_ids: set[str] = set()
for msg in self.deferred_queues.get(session_key, []):
trigger_id = _local_trigger_id(msg)
if trigger_id:
trigger_ids.add(trigger_id)
for msg in self._pending_messages_by_delivery_id.values():
if msg.session_key != session_key:
continue
trigger_id = _local_trigger_id(msg)
if trigger_id:
trigger_ids.add(trigger_id)
return trigger_ids
return self.pending_ids_for_session(session_key)
async def publish_next_deferred(self, session_key: str) -> bool:
queue = self.deferred_queues.get(session_key)
if not queue:
return False
msg = queue.pop(0)
if not queue:
self.deferred_queues.pop(session_key, None)
await self._publish_inbound(msg)
return True
def _should_defer_local_trigger_turn(
msg: InboundMessage,
session_key: str,
active_session_keys: Iterable[str],
) -> bool:
return local_trigger(msg.metadata) is not None and session_key in active_session_keys
def _local_trigger_id(msg: InboundMessage) -> str | None: