mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-09-04 10:11:46 +03:00
refactor: simplify cross-session messaging
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user