fix(trigger): defer local triggers until session idle

This commit is contained in:
chengyongru
2026-07-02 13:32:46 +08:00
committed by Xubin Ren
parent 09bde468eb
commit f32007c83f
13 changed files with 543 additions and 35 deletions
+3 -2
View File
@@ -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
View File
@@ -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."""
+5 -1
View File
@@ -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",
),
]
+29 -14
View File
@@ -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",
+8
View File
@@ -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]]:
+136
View File
@@ -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