fix(trigger): defer local triggers until session idle
This commit is contained in:
@@ -124,14 +124,15 @@ class CronTurnCoordinator:
|
||||
job_ids.add(job_id)
|
||||
return job_ids
|
||||
|
||||
async def publish_next_deferred(self, session_key: str) -> None:
|
||||
async def publish_next_deferred(self, session_key: str) -> bool:
|
||||
queue = self.deferred_queues.get(session_key)
|
||||
if not queue:
|
||||
return
|
||||
return False
|
||||
msg = queue.pop(0)
|
||||
if not queue:
|
||||
self.deferred_queues.pop(session_key, None)
|
||||
await self._publish_inbound(msg)
|
||||
return True
|
||||
|
||||
|
||||
def _cron_job_id(msg: InboundMessage) -> str | None:
|
||||
|
||||
+35
-2
@@ -67,6 +67,7 @@ from nanobot.session.manager import (
|
||||
SessionManager,
|
||||
replay_max_messages_for_context,
|
||||
)
|
||||
from nanobot.triggers.local_turns import LocalTriggerTurnCoordinator
|
||||
from nanobot.utils.document import extract_documents, reference_non_image_attachments
|
||||
from nanobot.utils.helpers import image_placeholder_text
|
||||
from nanobot.utils.helpers import truncate_text as truncate_text_fn
|
||||
@@ -320,6 +321,11 @@ class AgentLoop:
|
||||
dispatch=self._dispatch,
|
||||
is_running=lambda: self._running,
|
||||
)
|
||||
self._local_trigger_turns = LocalTriggerTurnCoordinator(
|
||||
publish_inbound=self.bus.publish_inbound,
|
||||
dispatch=self._dispatch,
|
||||
is_running=lambda: self._running,
|
||||
)
|
||||
# NANOBOT_MAX_CONCURRENT_REQUESTS: <=0 means unlimited; default 3.
|
||||
_max = int(os.environ.get("NANOBOT_MAX_CONCURRENT_REQUESTS", "3"))
|
||||
self._concurrency_gate: asyncio.Semaphore | None = (
|
||||
@@ -589,9 +595,20 @@ class AgentLoop:
|
||||
async def submit_cron_turn(self, msg: InboundMessage) -> OutboundMessage | None:
|
||||
return await self._cron_turns.submit(msg)
|
||||
|
||||
async def submit_local_trigger_turn(self, msg: InboundMessage) -> OutboundMessage | None:
|
||||
return await self._local_trigger_turns.submit(msg)
|
||||
|
||||
def pending_cron_job_ids_for_session(self, session_key: str) -> set[str]:
|
||||
return self._cron_turns.pending_job_ids_for_session(session_key)
|
||||
|
||||
def pending_local_trigger_ids_for_session(self, session_key: str) -> set[str]:
|
||||
return self._local_trigger_turns.pending_trigger_ids_for_session(session_key)
|
||||
|
||||
async def _publish_next_deferred_automation_turn(self, session_key: str) -> None:
|
||||
if await self._cron_turns.publish_next_deferred(session_key):
|
||||
return
|
||||
await self._local_trigger_turns.publish_next_deferred(session_key)
|
||||
|
||||
def _persist_user_message_early(
|
||||
self,
|
||||
msg: InboundMessage,
|
||||
@@ -931,6 +948,16 @@ class AgentLoop:
|
||||
effective_key,
|
||||
)
|
||||
continue
|
||||
if self._local_trigger_turns.defer_if_active(
|
||||
msg,
|
||||
session_key=effective_key,
|
||||
active_session_keys=self._pending_queues.keys(),
|
||||
):
|
||||
logger.info(
|
||||
"Deferred local trigger turn for active session {}",
|
||||
effective_key,
|
||||
)
|
||||
continue
|
||||
# If this session already has an active pending queue (i.e. a task
|
||||
# is processing this session), route the message there for mid-turn
|
||||
# injection instead of creating a competing task.
|
||||
@@ -1053,11 +1080,16 @@ class AgentLoop:
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
self._cron_turns.complete(msg, response=response)
|
||||
self._local_trigger_turns.complete(msg, response=response)
|
||||
except asyncio.CancelledError:
|
||||
self._cron_turns.complete(
|
||||
msg,
|
||||
error=asyncio.CancelledError(),
|
||||
)
|
||||
self._local_trigger_turns.complete(
|
||||
msg,
|
||||
error=asyncio.CancelledError(),
|
||||
)
|
||||
logger.info("Task cancelled for session {}", session_key)
|
||||
# Preserve partial context from the interrupted turn so
|
||||
# the user does not lose tool results and assistant
|
||||
@@ -1097,6 +1129,7 @@ class AgentLoop:
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
self._cron_turns.complete(msg, error=exc)
|
||||
self._local_trigger_turns.complete(msg, error=exc)
|
||||
finally:
|
||||
# Drain any messages still in the pending queue and re-publish
|
||||
# them to the bus so they are processed as fresh inbound messages
|
||||
@@ -1127,14 +1160,14 @@ class AgentLoop:
|
||||
msg, session_key, "idle"
|
||||
)
|
||||
self._runtime_events().clear_turn(session_key)
|
||||
await self._cron_turns.publish_next_deferred(session_key)
|
||||
await self._publish_next_deferred_automation_turn(session_key)
|
||||
finally:
|
||||
if pending is None:
|
||||
await self._runtime_events().run_status_changed(
|
||||
msg, session_key, "idle"
|
||||
)
|
||||
self._runtime_events().clear_turn(session_key)
|
||||
await self._cron_turns.publish_next_deferred(session_key)
|
||||
await self._publish_next_deferred_automation_turn(session_key)
|
||||
|
||||
async def close_mcp(self) -> None:
|
||||
"""Drain pending background archives, then close MCP connections."""
|
||||
|
||||
@@ -1295,7 +1295,11 @@ def _run_gateway(
|
||||
asyncio.create_task(agent.run(), name="nanobot-agent-loop"),
|
||||
asyncio.create_task(channels.start_all(), name="nanobot-channels"),
|
||||
asyncio.create_task(
|
||||
run_local_trigger_queue(store=trigger_store, bus=bus),
|
||||
run_local_trigger_queue(
|
||||
store=trigger_store,
|
||||
bus=bus,
|
||||
submit_turn=getattr(agent, "submit_local_trigger_turn", None),
|
||||
),
|
||||
name="nanobot-local-triggers",
|
||||
),
|
||||
]
|
||||
|
||||
@@ -4,11 +4,12 @@ from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import uuid
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Any
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.bus.events import InboundMessage
|
||||
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
|
||||
@@ -19,11 +20,14 @@ from nanobot.webui.metadata import WEBUI_MESSAGE_SOURCE_METADATA_KEY, WEBUI_TURN
|
||||
async def run_local_trigger_queue(
|
||||
*,
|
||||
store: LocalTriggerStore,
|
||||
bus: MessageBus,
|
||||
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")
|
||||
logger.info("Local trigger queue started")
|
||||
recovered = store.recover_processing_deliveries()
|
||||
if recovered:
|
||||
@@ -39,7 +43,12 @@ async def run_local_trigger_queue(
|
||||
|
||||
for delivery in deliveries:
|
||||
try:
|
||||
await _publish_delivery(store, bus, delivery)
|
||||
await _deliver_delivery(
|
||||
store,
|
||||
delivery,
|
||||
bus=bus,
|
||||
submit_turn=submit_turn,
|
||||
)
|
||||
store.complete_delivery(delivery)
|
||||
except asyncio.CancelledError as exc:
|
||||
store.retry_delivery(delivery, str(exc) or exc.__class__.__name__)
|
||||
@@ -79,10 +88,12 @@ class _TerminalDeliveryError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
async def _publish_delivery(
|
||||
async def _deliver_delivery(
|
||||
store: LocalTriggerStore,
|
||||
bus: MessageBus,
|
||||
delivery: TriggerDelivery,
|
||||
*,
|
||||
bus: MessageBus | None,
|
||||
submit_turn: Callable[[InboundMessage], Awaitable[OutboundMessage | None]] | None,
|
||||
) -> None:
|
||||
trigger = store.get(delivery.trigger_id)
|
||||
if trigger is None:
|
||||
@@ -90,16 +101,20 @@ async def _publish_delivery(
|
||||
if not trigger.enabled:
|
||||
raise _TerminalDeliveryError("trigger is disabled")
|
||||
|
||||
await bus.publish_inbound(
|
||||
InboundMessage(
|
||||
channel=trigger.channel,
|
||||
sender_id=trigger.sender_id,
|
||||
chat_id=trigger.chat_id,
|
||||
content=delivery.content,
|
||||
metadata=_delivery_metadata(trigger, delivery),
|
||||
session_key_override=trigger.session_key,
|
||||
)
|
||||
msg = InboundMessage(
|
||||
channel=trigger.channel,
|
||||
sender_id=trigger.sender_id,
|
||||
chat_id=trigger.chat_id,
|
||||
content=delivery.content,
|
||||
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)
|
||||
store.record_delivery(
|
||||
trigger.id,
|
||||
status="ok",
|
||||
|
||||
@@ -41,6 +41,14 @@ def local_trigger(metadata: Mapping[str, Any] | None) -> dict[str, Any] | None:
|
||||
return automation_trigger(metadata, LOCAL_TRIGGER_AUTOMATION_SPEC)
|
||||
|
||||
|
||||
def local_trigger_delivery_id(metadata: Mapping[str, Any] | None) -> str | None:
|
||||
trigger = local_trigger(metadata)
|
||||
if not trigger:
|
||||
return None
|
||||
value = trigger.get("delivery_id")
|
||||
return value if isinstance(value, str) and value else None
|
||||
|
||||
|
||||
def local_trigger_history_overrides(
|
||||
metadata: Mapping[str, Any] | None,
|
||||
) -> tuple[str | None, dict[str, Any]]:
|
||||
|
||||
@@ -0,0 +1,136 @@
|
||||
"""Coordination for local trigger turns."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import dataclasses
|
||||
from collections.abc import Awaitable, Callable, Iterable
|
||||
|
||||
from nanobot.bus.events import InboundMessage, OutboundMessage
|
||||
from nanobot.triggers.local_session_turns import local_trigger, local_trigger_delivery_id
|
||||
|
||||
|
||||
class LocalTriggerTurnCoordinator:
|
||||
"""Manage local trigger turns without mixing them into live injections."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
publish_inbound: Callable[[InboundMessage], Awaitable[None]],
|
||||
dispatch: Callable[[InboundMessage], Awaitable[object]],
|
||||
is_running: Callable[[], bool],
|
||||
) -> 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)
|
||||
|
||||
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
|
||||
|
||||
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 _local_trigger_id(msg: InboundMessage) -> str | None:
|
||||
trigger = local_trigger(msg.metadata)
|
||||
if not trigger:
|
||||
return None
|
||||
value = trigger.get("trigger_id")
|
||||
return value if isinstance(value, str) and value else None
|
||||
Reference in New Issue
Block a user