refactor: simplify cross-session messaging

This commit is contained in:
chengyongru
2026-08-19 01:15:56 +08:00
committed by chengyongru
parent 0e184965e8
commit 251a1ccd40
78 changed files with 1578 additions and 7569 deletions
+72 -98
View File
@@ -4,7 +4,7 @@ from __future__ import annotations
import re
import time
from collections.abc import Awaitable, Callable, Mapping
from collections.abc import Awaitable, Callable
from dataclasses import dataclass, replace
from typing import Any, cast
from uuid import uuid4
@@ -19,10 +19,10 @@ from nanobot.bus.outbound_events import (
GoalStateSyncEvent,
GoalStatusEvent,
RuntimeModelUpdatedEvent,
SessionMessageInputEvent,
SessionUpdatedEvent,
TurnEndEvent,
TurnModelUpdatedEvent,
UserInputEvent,
outbound_message_for_event,
)
from nanobot.bus.queue import MessageBus
@@ -35,6 +35,7 @@ from nanobot.bus.runtime_events import (
TurnCompleted,
TurnRunStatusChanged,
TurnRuntimeAdmitted,
UserInputAccepted,
)
from nanobot.providers.base import LLMProvider
from nanobot.providers.fallback_provider import FallbackModelObserver
@@ -42,17 +43,15 @@ from nanobot.runtime_context import public_history_message
from nanobot.session.goal_state import goal_state_ws_blob
from nanobot.session.history_visibility import is_hidden_history_message
from nanobot.session.manager import Session, SessionManager
from nanobot.session.session_handles import session_handle_for_key
from nanobot.session.session_messages import (
SESSION_MESSAGE_METADATA_KEY,
session_message_inbound,
session_message_public_metadata,
session_reply_timeout_inbound,
SessionMessageEnvelope,
session_message_envelope,
)
from nanobot.utils.helpers import strip_think, truncate_text
from nanobot.utils.llm_runtime import LLMRuntime
from nanobot.webui.metadata import (
WEBSOCKET_TURN_OWNER_METADATA_KEY,
WEBUI_MESSAGE_SOURCE_METADATA_KEY,
WEBUI_TURN_METADATA_KEY,
)
from nanobot.webui.transcript import append_session_message_input
@@ -83,6 +82,16 @@ class _WebsocketTurn:
_WEBSOCKET_ACTIVE_TURNS: dict[str, dict[str, _WebsocketTurn]] = {}
def _session_message_public_metadata(
envelope: SessionMessageEnvelope,
) -> dict[str, Any]:
source = session_handle_for_key(envelope["source_session_key"])
return {
"message_id": envelope["message_id"],
"session": source.public_payload(),
}
def _validated_llm_runtime(value: object) -> LLMRuntime | None:
"""Keep runtime-event consumers defensive if an external publisher violates the contract."""
return value if isinstance(value, LLMRuntime) else None
@@ -115,20 +124,6 @@ def mark_webui_session(session: Session, metadata: dict[str, Any]) -> bool:
return True
def _session_for_webui_lifecycle(
sessions: SessionManager,
msg: InboundMessage,
session_key: str,
) -> Session | None:
"""Resolve lifecycle state without reviving deleted internal-message targets."""
if (
session_message_inbound(msg) is not None
or session_reply_timeout_inbound(msg) is not None
):
return sessions.get_existing(session_key)
return sessions.get_or_create(session_key)
def clean_generated_title(raw: str | None) -> str:
text = (raw or "").strip()
if not text:
@@ -176,9 +171,7 @@ async def maybe_generate_webui_title(
model: str,
) -> bool:
"""Generate and persist a short title for WebUI-owned sessions only."""
session = sessions.get_existing(session_key)
if session is None:
return False
session = sessions.get_or_create(session_key)
if session.metadata.get(WEBUI_SESSION_METADATA_KEY) is not True:
return False
if session.metadata.get(WEBUI_TITLE_USER_EDITED_METADATA_KEY) is True:
@@ -426,8 +419,7 @@ class WebuiTurnRoutePolicy:
) -> TurnRoute:
"""Make an independently dispatched agent turn visible in WebUI."""
routed = route
session_message = session_message_inbound(msg)
reply_timeout = session_reply_timeout_inbound(msg)
internal_user_input = msg.channel == "system" and msg.is_user_input
if (
(
(
@@ -435,41 +427,19 @@ class WebuiTurnRoutePolicy:
and msg.sender_id == "subagent"
and msg.metadata.get("injected_event") == "subagent_result"
)
or session_message is not None
or reply_timeout is not None
or internal_user_input
)
and route.channel == "websocket"
):
if session_message is not None or reply_timeout is not None:
persisted = self.sessions.read_session_metadata(session_key)
raw_session_metadata = (
persisted.get("metadata") if persisted is not None else None
)
session_metadata: Mapping[str, Any] = (
cast(Mapping[str, Any], raw_session_metadata)
if isinstance(raw_session_metadata, Mapping)
else {}
)
else:
session_metadata = self.sessions.get_or_create(session_key).metadata
if session_metadata.get(WEBUI_SESSION_METADATA_KEY) is True:
session = self.sessions.get_or_create(session_key)
if session.metadata.get(WEBUI_SESSION_METADATA_KEY) is True:
metadata = dict(route.metadata)
turn_prefix = "subagent"
if session_message is not None:
turn_prefix = "session-message"
elif reply_timeout is not None:
turn_prefix = "session-reply-timeout"
turn_prefix = "session-input" if internal_user_input else "subagent"
metadata.update({
WEBUI_SESSION_METADATA_KEY: True,
"_wants_stream": True,
WEBUI_TURN_METADATA_KEY: f"{turn_prefix}:{uuid4().hex}",
})
if session_message is not None:
metadata[SESSION_MESSAGE_METADATA_KEY] = session_message
metadata[WEBUI_MESSAGE_SOURCE_METADATA_KEY] = {
"kind": "session",
"label": f"@{session_message['source']['name']}",
}
routed = replace(route, metadata=metadata, publish_lifecycle=True)
if routed.channel == "websocket" and routed.publish_lifecycle:
@@ -501,40 +471,6 @@ class WebuiTurnRoutePolicy:
return routed
async def project_session_message_input(
bus: MessageBus,
msg: InboundMessage,
session_key: str,
) -> None:
"""Persist and publish an incoming session message for WebUI clients."""
envelope = session_message_inbound(msg)
if envelope is None or msg.channel != "websocket":
return
public_metadata = session_message_public_metadata(envelope)
try:
append_session_message_input(
session_key,
content=msg.content,
created_at_ms=envelope["created_at_ms"],
session_message=public_metadata,
)
except (OSError, TypeError, ValueError):
logger.warning(
"Failed to persist session input {}",
envelope["message_id"],
exc_info=True,
)
await bus.publish_outbound(outbound_message_for_event(
channel="websocket",
chat_id=str(msg.chat_id),
event=SessionMessageInputEvent(
content=msg.content,
created_at_ms=envelope["created_at_ms"],
session_message=public_metadata,
),
))
def build_webui_fallback_model_observer(bus: MessageBus) -> FallbackModelObserver:
"""Translate provider fallback choices into chat-scoped WebUI events."""
@@ -575,6 +511,10 @@ class WebuiTurnCoordinator:
def subscribe(self, runtime_events: RuntimeEventBus) -> Callable[[], None]:
"""Subscribe this coordinator to runtime events."""
unsubscribe = [
runtime_events.subscribe(
self._handle_user_input_accepted,
UserInputAccepted,
),
runtime_events.subscribe(
self._handle_session_turn_started,
SessionTurnStarted,
@@ -622,17 +562,53 @@ class WebuiTurnCoordinator:
def _is_websocket_event(ctx: RuntimeEventContext) -> bool:
return ctx.channel == "websocket"
async def _handle_session_turn_started(self, event: SessionTurnStarted) -> None:
async def _handle_user_input_accepted(self, event: UserInputAccepted) -> None:
envelope = session_message_envelope(event.context.metadata)
session_key = event.context.session_key
if (
event.context.channel != "system"
or envelope is None
or envelope["target_session_key"] != session_key
or not session_key.startswith("websocket:")
):
return
persisted = self.sessions.read_session_metadata(session_key)
metadata_value: object = persisted.get("metadata") if persisted is not None else None
metadata = (
cast(dict[str, Any], metadata_value)
if isinstance(metadata_value, dict)
else None
)
if metadata is None or metadata.get(WEBUI_SESSION_METADATA_KEY) is not True:
return
public_metadata = _session_message_public_metadata(envelope)
try:
append_session_message_input(
session_key,
content=event.content,
created_at_ms=envelope["created_at_ms"],
session_message=public_metadata,
)
except (OSError, TypeError, ValueError):
logger.warning(
"Failed to persist session input {}",
envelope["message_id"],
exc_info=True,
)
await self.bus.publish_outbound(outbound_message_for_event(
channel="websocket",
chat_id=session_key.split(":", 1)[1],
event=UserInputEvent(
content=event.content,
created_at_ms=envelope["created_at_ms"],
provenance={"session_message": public_metadata},
),
))
def _handle_session_turn_started(self, event: SessionTurnStarted) -> None:
if not self._is_websocket_event(event.context):
return
msg = self._ctx_msg(event.context)
session = _session_for_webui_lifecycle(
self.sessions,
msg,
event.context.session_key,
)
if session is None:
return
session = self.sessions.get_or_create(event.context.session_key)
mark_webui_session(session, event.context.metadata)
async def _handle_run_status_changed(self, event: TurnRunStatusChanged) -> None:
@@ -726,9 +702,7 @@ class WebuiTurnCoordinator:
if msg.channel != "websocket":
return
session = _session_for_webui_lifecycle(self.sessions, msg, session_key)
if session is None:
return
session = self.sessions.get_or_create(session_key)
await self.bus.publish_outbound(
outbound_message_for_event(
channel=msg.channel,