refactor(trigger): share automation turn delivery
This commit is contained in:
@@ -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}"
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user