diff --git a/nanobot/agent/context.py b/nanobot/agent/context.py index 031c1ab9e..19538a762 100644 --- a/nanobot/agent/context.py +++ b/nanobot/agent/context.py @@ -217,16 +217,18 @@ class ContextBuilder: include_memory_recent_history: bool = True, session_key: str | None = None, unified_session: bool = False, + conversation_only: bool = False, ) -> list[dict[str, Any]]: """Build the complete message list for an LLM call.""" - root = workspace or self.workspace - active_skill_names = ( - self.skills.get_explicitly_invoked_skills(current_message) - if current_role == "user" - else [] - ) - messages: list[dict[str, Any]] = [ - { + messages = list(history) + if not conversation_only: + root = workspace or self.workspace + active_skill_names = ( + self.skills.get_explicitly_invoked_skills(current_message) + if current_role == "user" + else [] + ) + messages.insert(0, { "role": "system", "content": self.build_system_prompt( active_skill_names=active_skill_names, @@ -237,16 +239,14 @@ class ContextBuilder: session_key=session_key, unified_session=unified_session, ), - }, - *history, - ] + }) current = self.build_current_message( current_message, media=media, current_role=current_role, runtime_context_blocks=runtime_context_blocks, ) - if messages[-1].get("role") == current_role: + if messages and messages[-1].get("role") == current_role: last = dict(messages[-1]) last["content"] = self._merge_message_content( last.get("content"), diff --git a/nanobot/agent/loop.py b/nanobot/agent/loop.py index a218451b8..b1604324a 100644 --- a/nanobot/agent/loop.py +++ b/nanobot/agent/loop.py @@ -723,6 +723,7 @@ class AgentLoop: include_memory_recent_history=not ctx.ephemeral, session_key=ctx.session.key, unified_session=self._unified_session, + conversation_only=ctx.session.transient is True, ) def _request_context_for_turn(self, ctx: TurnContext) -> RequestContext: @@ -750,10 +751,12 @@ class AgentLoop: self, ctx: TurnContext, ) -> list[RuntimeContextBlock]: + if ctx.require_session().transient is True: + return [] assert ctx.request_context is not None return await self._resolve_runtime_context_for_request( ctx.request_context, - ctx.tools or self.tools, + ctx.tools if ctx.tools is not None else self.tools, ) async def _resolve_runtime_context_for_request( @@ -784,18 +787,24 @@ class AgentLoop: else: logger.warning("Command '{}' matched but dispatch returned None", raw) - async def _cancel_active_tasks(self, key: str) -> int: - """Cancel and await all active tasks and subagents for *key*. + async def cancel_active_turn(self, key: str) -> int: + """Cancel active work and discard queued follow-ups for *key*. Returns the total number of cancelled tasks + subagents. """ + pending = self._pending_queues.pop(key, None) + queued = 0 + if pending is not None: + while not pending.empty(): + pending.get_nowait() + queued += 1 tasks = tuple(self._active_tasks.pop(key, set())) cancelled = sum(1 for t in tasks if not t.done() and t.cancel()) for t in tasks: with suppress(asyncio.CancelledError, Exception): await t sub_cancelled = await self.subagents.cancel_by_session(key) - return cancelled + sub_cancelled + return queued + cancelled + sub_cancelled def _effective_session_key(self, msg: InboundMessage) -> str: """Return the session key used for task routing and mid-turn injections.""" @@ -922,7 +931,10 @@ class AgentLoop: if isinstance(metadata_value, dict) else {} ) - if pending_msg.channel != "system": + if ( + pending_msg.channel != "system" + and not (session is not None and session.transient is True) + ): scope = self.workspace_scopes.for_turn( channel=pending_msg.channel, message_metadata=metadata, @@ -1002,7 +1014,7 @@ class AgentLoop: message_metadata=metadata, session_metadata=session.metadata if session is not None else None, ) - effective_tools = tools or self.tools + effective_tools = tools if tools is not None else self.tools request_ctx = request_context or RequestContext( channel=channel, chat_id=chat_id, @@ -1160,6 +1172,11 @@ class AgentLoop: effective_key = self._effective_session_key(msg) if await agent_context.handle_runtime_control(self, msg, self.tools): continue + if ( + msg.transient_session + and not self.sessions.is_transient_active(effective_key) + ): + continue if self.commands.is_priority(raw): await self._dispatch_command_inline( msg, effective_key, raw, @@ -1271,6 +1288,8 @@ class AgentLoop: session_key, exc_info=True, ) + if msg.transient_session: + raise # Preserve partial context from the interrupted turn so # the user does not lose tool results and assistant # messages accumulated before /stop. The checkpoint was @@ -1573,13 +1592,16 @@ class AgentLoop: if ctx.session is None: ctx.session = self.sessions.get_or_create(ctx.session_key) session = ctx.session + if session.transient is True: + ctx.ephemeral = True + ctx.tools = ToolRegistry() self._remember_unified_session_route( session, msg, is_user_turn=ctx.original_user_text is not None, ) await ctx.delivery.started() - if ctx.kind is TurnKind.USER: + if ctx.kind is TurnKind.USER and not session.transient: self.workspace_scopes.persist_message_scope(session, msg) if self._restore_runtime_checkpoint(session): @@ -1589,6 +1611,8 @@ class AgentLoop: async def _compact_session(self, ctx: TurnContext) -> None: session = ctx.require_session() + if session.transient is True: + return ctx.session, pending = self.auto_compact.prepare_session( session, ctx.session_key, diff --git a/nanobot/bus/events.py b/nanobot/bus/events.py index def7703f9..d2d0f2d6c 100644 --- a/nanobot/bus/events.py +++ b/nanobot/bus/events.py @@ -18,6 +18,7 @@ INBOUND_META_RUNTIME_CONTROL = "_runtime_control" RUNTIME_CONTROL_ACK = "_ack" RUNTIME_CONTROL_MCP_RELOAD = "mcp_reload" RUNTIME_CONTROL_IMAGE_GENERATION_RELOAD = "image_generation_reload" +INBOUND_META_TRANSIENT_SESSION = "_transient_session" @dataclass @@ -32,6 +33,7 @@ class InboundMessage: media: list[str] = field(default_factory=list) # Media URLs metadata: dict[str, Any] = field(default_factory=dict) # Channel-specific data session_key_override: str | None = None # Optional override for thread-scoped sessions + transient_session: bool = False # In-memory session whose lifetime is owned by the channel @property def session_key(self) -> str: diff --git a/nanobot/channels/base.py b/nanobot/channels/base.py index 1784d6671..4fa92c690 100644 --- a/nanobot/channels/base.py +++ b/nanobot/channels/base.py @@ -8,7 +8,11 @@ from typing import Any, cast from loguru import logger -from nanobot.bus.events import InboundMessage, OutboundMessage +from nanobot.bus.events import ( + INBOUND_META_TRANSIENT_SESSION, + InboundMessage, + OutboundMessage, +) from nanobot.bus.queue import MessageBus from nanobot.pairing import ( PAIRING_CODE_META_KEY, @@ -277,7 +281,8 @@ class BaseChannel(ABC): ) return - meta = metadata or {} + meta = dict(metadata or {}) + transient_session = meta.pop(INBOUND_META_TRANSIENT_SESSION, False) is True if self.supports_streaming: meta = {**meta, "_wants_stream": True} @@ -289,6 +294,7 @@ class BaseChannel(ABC): media=media or [], metadata=meta, session_key_override=session_key, + transient_session=transient_session, ) await self.bus.publish_inbound(msg) diff --git a/nanobot/channels/manager.py b/nanobot/channels/manager.py index 27d9352cb..8d217a159 100644 --- a/nanobot/channels/manager.py +++ b/nanobot/channels/manager.py @@ -5,7 +5,7 @@ from __future__ import annotations import asyncio import hashlib import inspect -from collections.abc import Callable, Iterable +from collections.abc import Awaitable, Callable, Iterable from contextlib import suppress from pathlib import Path from typing import TYPE_CHECKING, Any, cast @@ -97,6 +97,7 @@ class ChannelManager: webui_runtime_model_name: Callable[[], str | None] | None = None, webui_cron_pending_job_ids: Callable[[str], set[str]] | None = None, webui_local_trigger_pending_ids: Callable[[str], set[str]] | None = None, + webui_cancel_active_turn: Callable[[str], Awaitable[int]] | None = None, webui_static_dist: bool = True, webui_runtime_surface: str = "browser", webui_runtime_capabilities: dict[str, Any] | None = None, @@ -110,6 +111,7 @@ class ChannelManager: self._webui_runtime_model_name = webui_runtime_model_name self._webui_cron_pending_job_ids = webui_cron_pending_job_ids self._webui_local_trigger_pending_ids = webui_local_trigger_pending_ids + self._webui_cancel_active_turn = webui_cancel_active_turn self._webui_static_dist = webui_static_dist self._webui_runtime_surface = webui_runtime_surface self._webui_runtime_capabilities = dict(webui_runtime_capabilities or {}) @@ -178,6 +180,7 @@ class ChannelManager: local_trigger_store=self._local_trigger_store, cron_pending_job_ids=self._webui_cron_pending_job_ids, local_trigger_pending_ids=self._webui_local_trigger_pending_ids, + cancel_active_turn=self._webui_cancel_active_turn, channel_feature_action=self.apply_channel_feature_action, channel_runtime_status=self.get_status, skill_state_action=self._webui_skill_state_action, diff --git a/nanobot/channels/websocket/runtime.py b/nanobot/channels/websocket/runtime.py index 1e0f39904..998de5661 100644 --- a/nanobot/channels/websocket/runtime.py +++ b/nanobot/channels/websocket/runtime.py @@ -18,7 +18,11 @@ from websockets.asyncio.server import ServerConnection, serve, unix_serve from websockets.exceptions import ConnectionClosed from websockets.http11 import Request as WsRequest -from nanobot.bus.events import OUTBOUND_META_AGENT_UI, OutboundMessage +from nanobot.bus.events import ( + INBOUND_META_TRANSIENT_SESSION, + OUTBOUND_META_AGENT_UI, + OutboundMessage, +) from nanobot.bus.outbound_events import ( GoalStateSyncEvent, GoalStatusEvent, @@ -32,6 +36,10 @@ from nanobot.bus.outbound_events import ( ) from nanobot.bus.queue import MessageBus from nanobot.channels.base import BaseChannel +from nanobot.channels.websocket.temporary_chat import ( + TemporaryChatLifecycle, + TemporaryChatLifecycleError, +) from nanobot.command.builtin import builtin_command_starts_agent_turn from nanobot.config.schema import Base from nanobot.runtime_context import ( @@ -76,6 +84,8 @@ from nanobot.webui.websocket_logging import websockets_server_logger # Plain HTTP WebUI routes also run through websockets.process_request. _WEBUI_HTTP_OPEN_TIMEOUT_S = 360.0 +_TEMPORARY_CHAT_ID_PREFIX = "temporary-" +_TEMPORARY_COMMANDS = frozenset({"/model", "/stop"}) class WebSocketConfig(Base): @@ -215,6 +225,10 @@ def _is_valid_chat_id(value: Any) -> TypeGuard[str]: return isinstance(value, str) and _CHAT_ID_RE.match(value) is not None +def _is_temporary_chat_id(value: Any) -> TypeGuard[str]: + return _is_valid_chat_id(value) and value.startswith(_TEMPORARY_CHAT_ID_PREFIX) + + def _parse_envelope(raw: str) -> dict[str, Any] | None: """Return a typed envelope dict if the frame is a new-style JSON envelope, else None. @@ -286,6 +300,13 @@ class WebSocketChannel(BaseChannel): self._workspaces = gateway.workspaces self._stream_text_buffers: dict[tuple[str, str], list[str]] = {} + self._temporary_chats = TemporaryChatLifecycle( + sessions=gateway.session_manager, + cancel_active_turn=gateway.cancel_active_turn, + attach=self._attach, + detach=self._detach, + clear_stream_buffers=self._clear_stream_buffers, + ) # -- Subscription bookkeeping ------------------------------------------- @@ -297,6 +318,23 @@ class WebSocketChannel(BaseChannel): self._subs.setdefault(chat_id, set()).add(connection) self._conn_chats.setdefault(connection, set()).add(chat_id) + def _detach(self, connection: ServerConnection, chat_id: str) -> None: + chats = self._conn_chats.get(connection) + if chats is not None: + chats.discard(chat_id) + if not chats: + self._conn_chats.pop(connection, None) + subscribers = self._subs.get(chat_id) + if subscribers is not None: + subscribers.discard(connection) + if not subscribers: + self._subs.pop(chat_id, None) + + def _clear_stream_buffers(self, chat_id: str) -> None: + for key in tuple(self._stream_text_buffers): + if key[0] == chat_id: + self._stream_text_buffers.pop(key, None) + async def send_webui_protocol_error( self, connection: ServerConnection, @@ -325,18 +363,15 @@ class WebSocketChannel(BaseChannel): ) await self._hydrate_after_subscribe(fork_id) - def _cleanup_connection(self, connection: ServerConnection) -> None: + async def _cleanup_connection(self, connection: ServerConnection) -> None: """Remove *connection* from every subscription set; safe to call multiple times.""" - chat_ids = self._conn_chats.pop(connection, set()) - for cid in chat_ids: - subs = self._subs.get(cid) - if subs is None: - continue - subs.discard(connection) - if not subs: - self._subs.pop(cid, None) - self._conn_default.pop(connection, None) - self._webui_connections.discard(connection) + try: + await self._temporary_chats.discard_owner(connection) + finally: + for chat_id in tuple(self._conn_chats.get(connection, ())): + self._detach(connection, chat_id) + self._conn_default.pop(connection, None) + self._webui_connections.discard(connection) async def _maybe_push_active_goal_state(self, chat_id: str) -> None: """Replay an active sustained goal from session metadata after *chat_id* is subscribed. @@ -387,7 +422,7 @@ class WebSocketChannel(BaseChannel): try: await connection.send(raw) except ConnectionClosed: - self._cleanup_connection(connection) + await self._cleanup_connection(connection) except Exception as e: self.logger.warning("failed to send {} event: {}", event, e) @@ -609,7 +644,7 @@ class WebSocketChannel(BaseChannel): except Exception as e: self.logger.debug("connection ended: {}", e) finally: - self._cleanup_connection(connection) + await self._cleanup_connection(connection) # -- Inbound WebSocket envelopes --------------------------------------- @@ -647,11 +682,36 @@ class WebSocketChannel(BaseChannel): if t == "fork_chat": await handle_webui_fork_chat(self, connection, envelope) return + if t == "discard_temporary_chat": + cid = envelope.get("chat_id") + if not _is_temporary_chat_id(cid): + await self._send_event(connection, "error", detail="invalid temporary chat_id") + return + try: + await self._temporary_chats.discard(connection, cid) + except TemporaryChatLifecycleError as exc: + await self._send_event( + connection, + "error", + detail=exc.detail, + chat_id=cid, + ) + return + await self._send_event(connection, "temporary_chat_discarded", chat_id=cid) + return if t == "attach": cid = envelope.get("chat_id") if not _is_valid_chat_id(cid): await self._send_event(connection, "error", detail="invalid chat_id") return + if _is_temporary_chat_id(cid): + await self._send_event( + connection, + "error", + detail="temporary_chat_cannot_attach", + chat_id=cid, + ) + return self._attach(connection, cid) await self._send_event(connection, "attached", chat_id=cid) await self._hydrate_after_subscribe(cid) @@ -661,6 +721,14 @@ class WebSocketChannel(BaseChannel): if not _is_valid_chat_id(cid): await self._send_event(connection, "error", detail="invalid chat_id") return + if _is_temporary_chat_id(cid): + await self._send_event( + connection, + "error", + detail="temporary_chat_has_no_workspace", + chat_id=cid, + ) + return scope = await self._workspace_scope_or_error( connection, lambda: self._workspaces.scope_for_set_request( @@ -692,6 +760,15 @@ class WebSocketChannel(BaseChannel): if not _is_valid_chat_id(cid): await self._send_event(connection, "error", detail="invalid chat_id") return + temporary = envelope.get("temporary") is True + if _is_temporary_chat_id(cid) != temporary: + await self._send_event( + connection, + "error", + detail="temporary_chat_mismatch", + chat_id=cid, + ) + return raw_turn_id = envelope.get("turn_id") turn_id = raw_turn_id if isinstance(raw_turn_id, str) and raw_turn_id else None rejection_fields = { @@ -728,6 +805,17 @@ class WebSocketChannel(BaseChannel): **rejection_fields, ) return + if temporary: + await self._dispatch_temporary_message( + connection, + client_id=client_id, + chat_id=cid, + content=content, + turn_id=turn_id, + envelope=envelope, + rejection_fields=rejection_fields, + ) + return raw_media = envelope.get("media") media_paths: list[str] = [] @@ -849,6 +937,103 @@ class WebSocketChannel(BaseChannel): return await self._send_event(connection, "error", detail=f"unknown type: {t!r}") + async def _dispatch_temporary_message( + self, + connection: ServerConnection, + *, + client_id: str, + chat_id: str, + content: str, + turn_id: str | None, + envelope: dict[str, Any], + rejection_fields: dict[str, str], + ) -> None: + """Admit a WebUI-only message without durable or local-agent capabilities.""" + if connection not in self._webui_connections: + await self._send_event( + connection, + "error", + detail="temporary_chat_unavailable", + **rejection_fields, + ) + return + forbidden = ( + "media", + "cli_apps", + "mcp_presets", + "quoted_context", + "workspace_scope", + ) + if any(field in envelope for field in forbidden): + await self._send_event( + connection, + "error", + detail="temporary_chat_capability_rejected", + **rejection_fields, + ) + return + if not content.strip(): + await self._send_event( + connection, + "error", + detail="missing content", + **rejection_fields, + ) + return + command = content.strip().partition(" ")[0].lower() + if command.startswith("/") and command not in _TEMPORARY_COMMANDS: + await self._send_event( + connection, + "error", + detail="temporary_chat_command_rejected", + **rejection_fields, + ) + return + + try: + session_key = self._temporary_chats.claim(connection, chat_id) + except TemporaryChatLifecycleError as exc: + await self._send_event( + connection, + "error", + detail=exc.detail, + **rejection_fields, + ) + return + + metadata: dict[str, Any] = { + "remote": getattr(connection, "remote_address", None), + "webui": True, + INBOUND_META_TRANSIENT_SESSION: True, + **self._transcripts.client_turn_metadata(turn_id), + } + queued_owner = None + if builtin_command_starts_agent_turn(content): + queued_owner = register_queued_websocket_turn_if_idle(chat_id, turn_id) + if queued_owner is not None: + metadata[WEBSOCKET_TURN_OWNER_METADATA_KEY] = queued_owner + accepted = False + try: + await self._handle_message( + sender_id=client_id, + chat_id=chat_id, + content=content, + metadata=metadata, + session_key=session_key, + is_dm=False, + ) + accepted = True + finally: + if not accepted and queued_owner is not None: + clear_websocket_turn_if_current(chat_id, queued_owner) + if turn_id: + await self._send_event( + connection, + "message_accepted", + chat_id=chat_id, + turn_id=turn_id, + ) + async def _workspace_scope_or_error( self, connection: ServerConnection, @@ -889,6 +1074,8 @@ class WebSocketChannel(BaseChannel): except Exception as e: self.logger.warning("server task error during shutdown: {}", e) self._server_task = None + for connection in tuple(self._conn_chats): + await self._temporary_chats.discard_owner(connection) self._subs.clear() self._conn_chats.clear() self._conn_default.clear() @@ -906,7 +1093,7 @@ class WebSocketChannel(BaseChannel): try: await connection.send(raw) except ConnectionClosed: - self._cleanup_connection(connection) + await self._cleanup_connection(connection) self.logger.warning("connection gone{}", label) except Exception: self.logger.exception("send failed{}", label) @@ -923,6 +1110,8 @@ class WebSocketChannel(BaseChannel): transcript_overrides: dict[str, Any] | None = None, ) -> bool: """Persist one canonical turn event and retain unsafe owners on failure.""" + if _is_temporary_chat_id(chat_id): + return True persisted = self._transcripts.prepare_and_append( chat_id, event, diff --git a/nanobot/channels/websocket/temporary_chat.py b/nanobot/channels/websocket/temporary_chat.py new file mode 100644 index 000000000..007bfef72 --- /dev/null +++ b/nanobot/channels/websocket/temporary_chat.py @@ -0,0 +1,85 @@ +"""Connection-owned lifecycle for WebUI Temporary Chat sessions.""" + +from __future__ import annotations + +from collections.abc import Awaitable, Callable + +from websockets.asyncio.server import ServerConnection + +from nanobot.session.manager import SessionManager +from nanobot.session.webui_turns import clear_websocket_turns + + +class TemporaryChatLifecycleError(RuntimeError): + """A stable WebSocket protocol error raised by the temporary-chat lifecycle.""" + + def __init__(self, detail: str) -> None: + self.detail = detail + super().__init__(detail) + + +class TemporaryChatLifecycle: + """Own temporary session identity, cancellation, and cleanup ordering.""" + + def __init__( + self, + *, + sessions: SessionManager | None, + cancel_active_turn: Callable[[str], Awaitable[int]] | None, + attach: Callable[[ServerConnection, str], None], + detach: Callable[[ServerConnection, str], None], + clear_stream_buffers: Callable[[str], None], + ) -> None: + self._sessions = sessions + self._cancel_active_turn = cancel_active_turn + self._attach = attach + self._detach = detach + self._clear_stream_buffers = clear_stream_buffers + self._owners: dict[str, ServerConnection] = {} + + def claim(self, owner: ServerConnection, chat_id: str) -> str: + """Claim *chat_id* for *owner* and return its in-memory session key.""" + if self._sessions is None or self._cancel_active_turn is None: + raise TemporaryChatLifecycleError("temporary_chat_unavailable") + current = self._owners.get(chat_id) + if current is not None and current is not owner: + raise TemporaryChatLifecycleError("temporary_chat_not_owned") + + session_key = f"websocket:{chat_id}" + self._sessions.get_or_create_transient(session_key) + self._owners[chat_id] = owner + self._attach(owner, chat_id) + return session_key + + async def discard(self, owner: ServerConnection, chat_id: str) -> None: + """Discard an owned chat; an unused chat is already discarded.""" + current = self._owners.get(chat_id) + if current is None: + return + if current is not owner: + raise TemporaryChatLifecycleError("temporary_chat_not_owned") + await self._discard_owned(owner, chat_id) + + async def discard_owner(self, owner: ServerConnection) -> None: + """Discard every temporary chat held by a disconnected owner.""" + chat_ids = ( + chat_id + for chat_id, current in self._owners.items() + if current is owner + ) + for chat_id in tuple(chat_ids): + await self._discard_owned(owner, chat_id) + + async def _discard_owned(self, owner: ServerConnection, chat_id: str) -> None: + self._owners.pop(chat_id, None) + self._detach(owner, chat_id) + + session_key = f"websocket:{chat_id}" + assert self._sessions is not None + assert self._cancel_active_turn is not None + self._sessions.discard_transient(session_key) + try: + await self._cancel_active_turn(session_key) + finally: + clear_websocket_turns(chat_id) + self._clear_stream_buffers(chat_id) diff --git a/nanobot/channels/websocket/tests/test_websocket_channel.py b/nanobot/channels/websocket/tests/test_websocket_channel.py index 1e44b1e65..b381ec5c9 100644 --- a/nanobot/channels/websocket/tests/test_websocket_channel.py +++ b/nanobot/channels/websocket/tests/test_websocket_channel.py @@ -111,6 +111,7 @@ def _basic_handler(bus: Any, **kw: Any) -> GatewayServices: runtime_model_name=None, runtime_surface=kw.get("runtime_surface", "browser"), runtime_capabilities_overrides=kw.get("runtime_capabilities_overrides"), + cancel_active_turn=kw.get("cancel_active_turn"), ) @@ -190,6 +191,182 @@ def isolate_webui_workspace_state(tmp_path, monkeypatch) -> None: wth._WEBSOCKET_TURN_OWNERS.clear() +@pytest.mark.asyncio +async def test_temporary_message_registers_in_memory_session(bus, tmp_path) -> None: + sessions = SessionManager(tmp_path) + cancel = AsyncMock(return_value=0) + channel = WebSocketChannel( + {"enabled": True, "allowFrom": ["*"]}, + bus, + gateway=_basic_handler( + bus, + session_manager=sessions, + cancel_active_turn=cancel, + ), + ) + connection = AsyncMock() + connection.remote_address = None + channel._webui_connections.add(connection) + chat_id = "temporary-test" + + await channel._dispatch_envelope( + connection, + "client", + { + "type": "message", + "chat_id": chat_id, + "content": "hello", + "turn_id": "turn-1", + "temporary": True, + "webui": True, + }, + ) + + inbound = bus.publish_inbound.await_args.args[0] + assert inbound.session_key == f"websocket:{chat_id}" + assert inbound.transient_session is True + assert sessions.is_transient_active(inbound.session_key) is True + assert sessions.get_cached(inbound.session_key).transient is True + assert read_transcript_lines(inbound.session_key) == [] + assert json.loads(connection.send.await_args.args[0])["event"] == "message_accepted" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "envelope", + [ + {"type": "attach", "chat_id": "temporary-test"}, + { + "type": "set_workspace_scope", + "chat_id": "temporary-test", + "workspace_scope": {}, + }, + { + "type": "message", + "chat_id": "temporary-test", + "content": "hello", + "temporary": True, + "media": [], + }, + { + "type": "message", + "chat_id": "temporary-test", + "content": "/history", + "temporary": True, + }, + ], +) +async def test_temporary_chat_rejects_persistent_capabilities( + bus, + tmp_path, + envelope, +) -> None: + channel = WebSocketChannel( + {"enabled": True, "allowFrom": ["*"]}, + bus, + gateway=_basic_handler( + bus, + session_manager=SessionManager(tmp_path), + cancel_active_turn=AsyncMock(return_value=0), + ), + ) + connection = AsyncMock() + connection.remote_address = None + channel._webui_connections.add(connection) + + await channel._dispatch_envelope(connection, "client", envelope) + + payload = json.loads(connection.send.await_args.args[0]) + assert payload["event"] == "error" + bus.publish_inbound.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_discard_temporary_chat_cancels_then_forgets_session(bus, tmp_path) -> None: + sessions = SessionManager(tmp_path) + cancel = AsyncMock(return_value=1) + channel = WebSocketChannel( + {"enabled": True, "allowFrom": ["*"]}, + bus, + gateway=_basic_handler( + bus, + session_manager=sessions, + cancel_active_turn=cancel, + ), + ) + connection = AsyncMock() + connection.remote_address = None + channel._webui_connections.add(connection) + chat_id = "temporary-test" + session_key = channel._temporary_chats.claim(connection, chat_id) + sessions.get_cached(session_key).add_message("user", "private") + + await channel._dispatch_envelope( + connection, + "client", + {"type": "discard_temporary_chat", "chat_id": chat_id}, + ) + + cancel.assert_awaited_once_with(session_key) + assert sessions.get_cached(session_key) is None + assert chat_id not in channel._subs + assert json.loads(connection.send.await_args.args[0]) == { + "event": "temporary_chat_discarded", + "chat_id": chat_id, + } + + +@pytest.mark.asyncio +async def test_discard_unused_temporary_chat_is_idempotent(bus, tmp_path) -> None: + cancel = AsyncMock(return_value=0) + channel = WebSocketChannel( + {"enabled": True, "allowFrom": ["*"]}, + bus, + gateway=_basic_handler( + bus, + session_manager=SessionManager(tmp_path), + cancel_active_turn=cancel, + ), + ) + connection = AsyncMock() + + await channel._dispatch_envelope( + connection, + "client", + {"type": "discard_temporary_chat", "chat_id": "temporary-unused"}, + ) + + cancel.assert_not_awaited() + assert json.loads(connection.send.await_args.args[0]) == { + "event": "temporary_chat_discarded", + "chat_id": "temporary-unused", + } + + +@pytest.mark.asyncio +async def test_disconnect_discards_owned_temporary_chat(bus, tmp_path) -> None: + sessions = SessionManager(tmp_path) + cancel = AsyncMock(return_value=1) + channel = WebSocketChannel( + {"enabled": True, "allowFrom": ["*"]}, + bus, + gateway=_basic_handler( + bus, + session_manager=sessions, + cancel_active_turn=cancel, + ), + ) + connection = AsyncMock() + chat_id = "temporary-disconnect" + session_key = channel._temporary_chats.claim(connection, chat_id) + + await channel._cleanup_connection(connection) + + cancel.assert_awaited_once_with(session_key) + assert sessions.get_cached(session_key) is None + assert chat_id not in channel._subs + + @pytest.mark.asyncio async def test_send_session_updated_broadcasts_to_other_webui_connections(bus) -> None: class Conn: diff --git a/nanobot/cli/gateway_runtime.py b/nanobot/cli/gateway_runtime.py index 718b138ec..ada04014b 100644 --- a/nanobot/cli/gateway_runtime.py +++ b/nanobot/cli/gateway_runtime.py @@ -581,6 +581,7 @@ def _run_gateway( webui_runtime_model_name=_webui_runtime_model_name, webui_cron_pending_job_ids=agent.pending_cron_job_ids_for_session, webui_local_trigger_pending_ids=agent.pending_local_trigger_ids_for_session, + webui_cancel_active_turn=getattr(agent, "cancel_active_turn", None), webui_static_dist=webui_static_dist, webui_runtime_surface=webui_runtime_surface, webui_runtime_capabilities=webui_runtime_capabilities, diff --git a/nanobot/command/builtin.py b/nanobot/command/builtin.py index 15e1fdfa4..df003f7e4 100644 --- a/nanobot/command/builtin.py +++ b/nanobot/command/builtin.py @@ -203,16 +203,7 @@ async def cmd_stop(ctx: CommandContext) -> OutboundMessage: """Cancel all active tasks and subagents for the session.""" loop = ctx.loop msg = ctx.msg - total = await loop._cancel_active_tasks(ctx.key) # pyright: ignore[reportPrivateUsage] - # Also drain pending queue to prevent mid-turn injection deadlock - pending = loop._pending_queues.pop(ctx.key, None) # pyright: ignore[reportPrivateUsage] - if pending is not None: - while not pending.empty(): - try: - pending.get_nowait() - total += 1 - except Exception: - break + total = await loop.cancel_active_turn(ctx.key) content = f"Stopped {total} task(s)." if total else "No active task to stop." return OutboundMessage( channel=msg.channel, chat_id=msg.chat_id, content=content, @@ -301,7 +292,7 @@ async def cmd_status(ctx: CommandContext) -> OutboundMessage: async def cmd_new(ctx: CommandContext) -> OutboundMessage: """Stop active task and start a fresh session.""" loop = ctx.loop - await loop._cancel_active_tasks(ctx.key) # pyright: ignore[reportPrivateUsage] + await loop.cancel_active_turn(ctx.key) session = ctx.session or loop.sessions.get_or_create(ctx.key) snapshot = session.messages[session.last_consolidated:] runtime = None diff --git a/nanobot/session/manager.py b/nanobot/session/manager.py index 0c832de88..283ec7d90 100644 --- a/nanobot/session/manager.py +++ b/nanobot/session/manager.py @@ -157,6 +157,7 @@ class Session: metadata: dict[str, Any] = field(default_factory=dict) last_consolidated: int = 0 # Number of messages already consolidated to files provider_state: ProviderConversationState | None = field(default=None, repr=False) + transient: bool = field(default=False, repr=False, compare=False) def __post_init__(self) -> None: if not isinstance(cast(object, self.metadata), dict): @@ -964,6 +965,7 @@ class SessionManager: self._cache: OrderedDict[str, Session] = OrderedDict() # Preserve identity for sessions held by active callers without retaining idle ones. self._overflow_cache: WeakValueDictionary[str, Session] = WeakValueDictionary() + self._transient_sessions: dict[str, Session] = {} self._max_cached_sessions = SESSION_CACHE_MAX_SIZE self._file_cap_archiver: Callable[..., None] | None = None @@ -977,6 +979,10 @@ class SessionManager: self._overflow_cache[key] = evicted def _cached(self, key: str) -> Session | None: + transient = self._transient_sessions.get(key) + if transient is not None: + return transient + session = self._cache.get(key) if session is not None: self._cache.move_to_end(key) @@ -1053,6 +1059,24 @@ class SessionManager: self._remember(session) return session + def get_or_create_transient(self, key: str) -> Session: + """Return an active in-memory session that can never reach the store.""" + session = self._transient_sessions.get(key) + if session is None: + self._cache.pop(key, None) + self._overflow_cache.pop(key, None) + session = Session(key=key, transient=True) + self._transient_sessions[key] = session + return session + + def is_transient_active(self, key: str) -> bool: + """Return whether *key* still accepts transient turns.""" + return key in self._transient_sessions + + def discard_transient(self, key: str) -> bool: + """Forget all transient contents without retaining a discarded-key tombstone.""" + return self._transient_sessions.pop(key, None) is not None + def _load(self, key: str) -> Session | None: return self._store.load(key) @@ -1066,6 +1090,9 @@ class SessionManager: def save(self, session: Session, *, fsync: bool = False) -> None: """Persist a session and retain it in the cache.""" + if session.transient is True: + return + archiver = self._file_cap_archiver if archiver is not None: session.enforce_file_cap( @@ -1098,6 +1125,7 @@ class SessionManager: def invalidate(self, key: str) -> None: """Remove a session from the in-memory cache.""" + self._transient_sessions.pop(key, None) self._cache.pop(key, None) self._overflow_cache.pop(key, None) diff --git a/nanobot/session/webui_turns.py b/nanobot/session/webui_turns.py index 0e526891e..e6c977afe 100644 --- a/nanobot/session/webui_turns.py +++ b/nanobot/session/webui_turns.py @@ -334,6 +334,16 @@ def clear_websocket_turn_if_current( return False +def clear_websocket_turns(chat_id: str) -> int: + """Clear every in-memory lifecycle owner for a discarded chat.""" + turns = _WEBSOCKET_ACTIVE_TURNS.pop(chat_id, None) + count = len(turns) if turns is not None else 0 + _WEBSOCKET_TURN_WALL_STARTED_AT.pop(chat_id, None) + _WEBSOCKET_TURN_IDS.pop(chat_id, None) + _WEBSOCKET_TURN_OWNERS.pop(chat_id, None) + return count + + def build_bus_progress_callback( bus: MessageBus, msg: InboundMessage, diff --git a/nanobot/webui/gateway_services.py b/nanobot/webui/gateway_services.py index 5bb6702bc..a189cdec6 100644 --- a/nanobot/webui/gateway_services.py +++ b/nanobot/webui/gateway_services.py @@ -2,9 +2,10 @@ from __future__ import annotations +from collections.abc import Awaitable, Callable from dataclasses import dataclass from pathlib import Path -from typing import TYPE_CHECKING, Any, Callable +from typing import TYPE_CHECKING, Any from loguru import logger as default_logger @@ -38,6 +39,7 @@ class GatewayServices: local_trigger_store: LocalTriggerStore | None cron_pending_job_ids: Callable[[str], set[str]] | None local_trigger_pending_ids: Callable[[str], set[str]] | None + cancel_active_turn: Callable[[str], Awaitable[int]] | None def build_gateway_services( @@ -56,6 +58,7 @@ def build_gateway_services( local_trigger_store: LocalTriggerStore | None = None, cron_pending_job_ids: Callable[[str], set[str]] | None = None, local_trigger_pending_ids: Callable[[str], set[str]] | None = None, + cancel_active_turn: Callable[[str], Awaitable[int]] | None = None, channel_feature_action: Callable[..., Any] | None = None, channel_runtime_status: Callable[[], dict[str, Any]] | None = None, skill_state_action: Callable[[set[str]], None] | None = None, @@ -117,4 +120,5 @@ def build_gateway_services( local_trigger_store=local_trigger_store, cron_pending_job_ids=cron_pending_job_ids, local_trigger_pending_ids=local_trigger_pending_ids, + cancel_active_turn=cancel_active_turn, ) diff --git a/tests/agent/test_context_builder.py b/tests/agent/test_context_builder.py index 503602c5f..4314f8ea0 100644 --- a/tests/agent/test_context_builder.py +++ b/tests/agent/test_context_builder.py @@ -15,6 +15,20 @@ def _builder(tmp_path: Path, **kw) -> ContextBuilder: return ContextBuilder(workspace=tmp_path, **kw) +def test_conversation_only_messages_omit_the_system_prompt(tmp_path) -> None: + (tmp_path / "AGENTS.md").write_text("SECRET PROJECT INSTRUCTIONS", encoding="utf-8") + builder = _builder(tmp_path) + + messages = builder.build_messages( + [], + "hello", + conversation_only=True, + ) + + assert messages == [{"role": "user", "content": "hello"}] + assert "SECRET PROJECT INSTRUCTIONS" not in str(messages) + + # --------------------------------------------------------------------------- # _merge_message_content (static) # --------------------------------------------------------------------------- diff --git a/tests/agent/test_task_cancel.py b/tests/agent/test_task_cancel.py index 6acb41ff7..819551162 100644 --- a/tests/agent/test_task_cancel.py +++ b/tests/agent/test_task_cancel.py @@ -111,8 +111,50 @@ class TestHandleStop: assert all(e.is_set() for e in events) assert "2 task" in out.content + @pytest.mark.asyncio + async def test_cancel_active_turn_discards_pending_followups(self): + from nanobot.bus.events import InboundMessage + + loop, _ = _make_loop() + pending = asyncio.Queue() + pending.put_nowait( + InboundMessage(channel="test", sender_id="u1", chat_id="c1", content="next") + ) + loop._pending_queues["test:c1"] = pending + + assert await loop.cancel_active_turn("test:c1") == 1 + assert "test:c1" not in loop._pending_queues + class TestDispatch: + @pytest.mark.asyncio + async def test_run_drops_deactivated_transient_message(self): + from nanobot.bus.events import InboundMessage + + loop, bus = _make_loop() + msg = InboundMessage( + channel="websocket", + sender_id="u1", + chat_id="temporary-test", + content="private", + session_key_override="websocket:temporary-test", + transient_session=True, + ) + + async def consume_once(): + loop.stop() + return msg + + bus.consume_inbound = AsyncMock(side_effect=consume_once) + loop.sessions.is_transient_active.return_value = False + loop._dispatch = AsyncMock() + loop.close_mcp = AsyncMock() + loop._running = True + + await loop.run() + + loop._dispatch.assert_not_awaited() + @pytest.mark.asyncio async def test_run_logs_and_continues_after_leaked_cancelled_error(self, monkeypatch): loop, bus = _make_loop() diff --git a/tests/agent/test_temporary_chat.py b/tests/agent/test_temporary_chat.py new file mode 100644 index 000000000..cbc30c21f --- /dev/null +++ b/tests/agent/test_temporary_chat.py @@ -0,0 +1,169 @@ +import asyncio +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from nanobot.agent.loop import AgentLoop +from nanobot.agent.tools.registry import ToolRegistry +from nanobot.bus.events import InboundMessage +from nanobot.bus.queue import MessageBus +from nanobot.providers.base import GenerationSettings, LLMResponse +from nanobot.runtime_context import RuntimeContextBlock +from nanobot.session.manager import SessionManager + + +@pytest.mark.asyncio +async def test_temporary_chat_reuses_memory_only_history_without_tools(tmp_path) -> None: + (tmp_path / "AGENTS.md").write_text("private project instruction", encoding="utf-8") + provider = MagicMock() + provider.get_default_model.return_value = "test-model" + provider.generation = GenerationSettings() + provider.chat_with_retry = AsyncMock( + side_effect=[ + LLMResponse(content="first answer", usage={}), + LLMResponse(content="second answer", usage={}), + ] + ) + loop = AgentLoop( + bus=MessageBus(), + provider=provider, + workspace=tmp_path, + model="test-model", + unified_session=True, + ) + key = "websocket:temporary-test" + loop.sessions.get_or_create_transient(key) + + for content in ("first question", "second question"): + response = await loop._process_message( + InboundMessage( + channel="websocket", + sender_id="user", + chat_id="temporary-test", + content=content, + session_key_override=key, + transient_session=True, + ) + ) + assert response is not None + + first_call, second_call = provider.chat_with_retry.await_args_list + assert first_call.kwargs["tools"] == [] + assert second_call.kwargs["tools"] == [] + assert all( + message["role"] != "system" + for call in (first_call, second_call) + for message in call.kwargs["messages"] + ) + assert "private project instruction" not in str(first_call.kwargs["messages"]) + assert str(tmp_path) not in str(first_call.kwargs["messages"]) + assert "first answer" in str(second_call.kwargs["messages"]) + + transient = loop.sessions.get_cached(key) + assert transient is not None + assert [message["role"] for message in transient.messages] == [ + "user", + "assistant", + "user", + "assistant", + ] + assert loop.sessions.read_session_file(key) is None + assert SessionManager(tmp_path).read_session_file(key) is None + + +@pytest.mark.asyncio +async def test_temporary_follow_up_does_not_resolve_runtime_context(tmp_path) -> None: + provider = MagicMock() + provider.get_default_model.return_value = "test-model" + provider.generation = GenerationSettings() + provider.chat_with_retry = AsyncMock( + side_effect=[ + LLMResponse(content="first answer", usage={}), + LLMResponse(content="second answer", usage={}), + ] + ) + loop = AgentLoop( + bus=MessageBus(), + provider=provider, + workspace=tmp_path, + model="test-model", + ) + runtime_context_provider = AsyncMock( + return_value=RuntimeContextBlock( + source="project", + content="SECRET LOCAL PROJECT CONTEXT", + ) + ) + loop.register_runtime_context_provider(runtime_context_provider) + + key = "websocket:temporary-follow-up" + session = loop.sessions.get_or_create_transient(key) + pending_queue: asyncio.Queue[InboundMessage] = asyncio.Queue() + await pending_queue.put( + InboundMessage( + channel="websocket", + sender_id="user", + chat_id="temporary-follow-up", + content="follow up", + session_key_override=key, + transient_session=True, + ) + ) + + _, _, messages, _, _ = await loop._run_agent_loop( + [{"role": "user", "content": "first question"}], + runtime=loop.llm_runtime(), + session=session, + channel="websocket", + chat_id="temporary-follow-up", + session_key=key, + pending_queue=pending_queue, + tools=ToolRegistry(), + ) + + runtime_context_provider.assert_not_awaited() + assert "SECRET LOCAL PROJECT CONTEXT" not in str(messages) + + +@pytest.mark.asyncio +async def test_discarding_active_temporary_chat_does_not_create_durable_session( + tmp_path, +) -> None: + provider_started = asyncio.Event() + provider = MagicMock() + provider.get_default_model.return_value = "test-model" + provider.generation = GenerationSettings() + + async def block_provider(**_kwargs): + provider_started.set() + await asyncio.Event().wait() + + provider.chat_with_retry = AsyncMock(side_effect=block_provider) + loop = AgentLoop( + bus=MessageBus(), + provider=provider, + workspace=tmp_path, + model="test-model", + ) + key = "websocket:temporary-cancelled" + loop.sessions.get_or_create_transient(key) + message = InboundMessage( + channel="websocket", + sender_id="user", + chat_id="temporary-cancelled", + content="private", + session_key_override=key, + transient_session=True, + ) + task = asyncio.create_task(loop._dispatch(message)) + active_tasks = loop._active_tasks.setdefault(key, set()) + active_tasks.add(task) + task.add_done_callback(active_tasks.discard) + + await provider_started.wait() + assert loop.sessions.discard_transient(key) + assert await loop.cancel_active_turn(key) == 1 + + assert loop.sessions.get_cached(key) is None + assert loop.sessions.flush_all() == 0 + assert loop.sessions.read_session_file(key) is None diff --git a/tests/agent/test_unified_session.py b/tests/agent/test_unified_session.py index 7c0656ec2..334573de2 100644 --- a/tests/agent/test_unified_session.py +++ b/tests/agent/test_unified_session.py @@ -253,7 +253,7 @@ class TestCmdNewUnifiedSession: loop = SimpleNamespace( sessions=sessions, consolidator=SimpleNamespace(archive=AsyncMock(return_value=True)), - _cancel_active_tasks=AsyncMock(return_value=0), + cancel_active_turn=AsyncMock(return_value=0), llm_runtime=MagicMock(return_value=MagicMock()), schedule_background=lambda coro: asyncio.ensure_future(coro), ) @@ -301,7 +301,7 @@ class TestCmdNewUnifiedSession: loop = SimpleNamespace( sessions=sessions, consolidator=SimpleNamespace(archive=AsyncMock(return_value=True)), - _cancel_active_tasks=AsyncMock(return_value=0), + cancel_active_turn=AsyncMock(return_value=0), runtime_for_session=MagicMock(return_value=MagicMock()), schedule_background=lambda coro: asyncio.ensure_future(coro), ) diff --git a/tests/command/test_router_dispatchable.py b/tests/command/test_router_dispatchable.py index 837b55e23..885752d17 100644 --- a/tests/command/test_router_dispatchable.py +++ b/tests/command/test_router_dispatchable.py @@ -109,7 +109,7 @@ class TestMidTurnCommandDispatchedDirectly: loop.sessions.save = MagicMock() loop.sessions.invalidate = MagicMock() loop.schedule_background = MagicMock() - loop._cancel_active_tasks = AsyncMock(return_value=0) + loop.cancel_active_turn = AsyncMock(return_value=0) return loop @pytest.fixture() diff --git a/tests/command/test_stop_pending_queue.py b/tests/command/test_stop_pending_queue.py index d07de2feb..721cbca09 100644 --- a/tests/command/test_stop_pending_queue.py +++ b/tests/command/test_stop_pending_queue.py @@ -1,6 +1,5 @@ """Test cmd_stop drains pending queue to prevent mid-turn injection deadlock.""" -import asyncio from unittest.mock import AsyncMock, MagicMock import pytest @@ -14,13 +13,7 @@ from nanobot.command.router import CommandContext async def test_cmd_stop_drains_pending_queue(): """cmd_stop should drain pending queue in addition to cancelling active tasks.""" mock_loop = MagicMock() - mock_loop._cancel_active_tasks = AsyncMock(return_value=1) - mock_loop._pending_queues = {} - - pending = asyncio.Queue() - await pending.put("msg1") - await pending.put("msg2") - mock_loop._pending_queues["test-session"] = pending + mock_loop.cancel_active_turn = AsyncMock(return_value=3) ctx = CommandContext( msg=MagicMock(channel="websocket", chat_id="test-chat", metadata={}), @@ -34,18 +27,14 @@ async def test_cmd_stop_drains_pending_queue(): assert isinstance(result, OutboundMessage) assert "Stopped 3 task(s)" in result.content # 1 cancelled + 2 drained - assert "test-session" not in mock_loop._pending_queues + mock_loop.cancel_active_turn.assert_awaited_once_with("test-session") @pytest.mark.asyncio async def test_cmd_stop_with_empty_pending_queue(): """cmd_stop should work correctly when pending queue is empty.""" mock_loop = MagicMock() - mock_loop._cancel_active_tasks = AsyncMock(return_value=2) - mock_loop._pending_queues = {} - - pending = asyncio.Queue() - mock_loop._pending_queues["test-session"] = pending + mock_loop.cancel_active_turn = AsyncMock(return_value=2) ctx = CommandContext( msg=MagicMock(channel="websocket", chat_id="test-chat", metadata={}), @@ -58,15 +47,14 @@ async def test_cmd_stop_with_empty_pending_queue(): result = await cmd_stop(ctx) assert "Stopped 2 task(s)" in result.content - assert "test-session" not in mock_loop._pending_queues + mock_loop.cancel_active_turn.assert_awaited_once_with("test-session") @pytest.mark.asyncio async def test_cmd_stop_no_pending_queue(): """cmd_stop should work when no pending queue exists.""" mock_loop = MagicMock() - mock_loop._cancel_active_tasks = AsyncMock(return_value=0) - mock_loop._pending_queues = {} + mock_loop.cancel_active_turn = AsyncMock(return_value=0) ctx = CommandContext( msg=MagicMock(channel="websocket", chat_id="test-chat", metadata={}), diff --git a/tests/session/test_session_cache.py b/tests/session/test_session_cache.py index cf1445b44..7c5173184 100644 --- a/tests/session/test_session_cache.py +++ b/tests/session/test_session_cache.py @@ -73,3 +73,23 @@ def test_flush_all_includes_live_sessions_outside_strong_cache(tmp_path, monkeyp assert manager.flush_all() == 2 assert set(saved) == {("test:active", True), ("test:other", True)} + + +def test_transient_session_never_reaches_store(tmp_path) -> None: + manager = SessionManager(tmp_path) + session = manager.get_or_create_transient("websocket:temporary-test") + session.add_message("user", "private") + + manager.save(session, fsync=True) + + assert manager.get_cached(session.key) is session + assert manager.read_session_file(session.key) is None + + +def test_transient_session_becomes_inactive_when_discarded(tmp_path) -> None: + manager = SessionManager(tmp_path) + session = manager.get_or_create_transient("websocket:temporary-test") + + assert manager.discard_transient(session.key) is True + assert manager.is_transient_active(session.key) is False + assert manager.get_cached(session.key) is None diff --git a/webui/src/App.tsx b/webui/src/App.tsx index 90b550cb3..adf874fd5 100644 --- a/webui/src/App.tsx +++ b/webui/src/App.tsx @@ -8,7 +8,7 @@ import { useState, type ReactNode, } from "react"; -import { Moon, PanelLeft, ShieldCheck, Sun, X } from "lucide-react"; +import { Ghost, Moon, PanelLeft, ShieldCheck, Sun, X } from "lucide-react"; import { useTranslation } from "react-i18next"; import { channelUiPresentation } from "@/channel-plugins/registry"; import { Sidebar } from "@/components/Sidebar"; @@ -37,6 +37,13 @@ import { import { displayTitle } from "@/lib/chat-groups"; import { deriveTitle } from "@/lib/format"; import { NanobotClient } from "@/lib/nanobot-client"; +import { + createTemporaryChatSession, + isQuickChatKey, + QUICK_CHAT_ID, + QUICK_CHAT_KEY, + quickChatSession, +} from "@/lib/quick-chat"; import { ClientProvider, useClient } from "@/providers/ClientProvider"; import type { BootstrapResponse, @@ -225,6 +232,9 @@ function readShellRoute(): ShellRoute { if (path === "/skills") { return { view: "skills", activeKey, settingsSection: "skills" }; } + if (path === "/quick-chat") { + return { view: "chat", activeKey: QUICK_CHAT_KEY, settingsSection: "overview" }; + } if (path.startsWith("/chat/")) { const encoded = path.slice("/chat/".length); try { @@ -241,6 +251,7 @@ function readShellRoute(): ShellRoute { function shellRouteHash(route: ShellRoute): string { if (route.view === "chat") { + if (isQuickChatKey(route.activeKey)) return "#/quick-chat"; return route.activeKey ? `#/chat/${encodeURIComponent(route.activeKey)}` : "#/new"; @@ -947,14 +958,24 @@ function Shell({ deleteChat, getSessionAutomations, } = useSessions(); + const regularSessions = useMemo( + () => sessions.filter((session) => !isQuickChatKey(session.key)), + [sessions], + ); + const quickSession = useMemo( + () => quickChatSession(sessions.find((session) => isQuickChatKey(session.key))), + [sessions], + ); const { state: sidebarState, update: updateSidebarState } = - useSidebarState(sessions, !loading); + useSidebarState(regularSessions, !loading); const initialRouteRef = useRef(null); if (!initialRouteRef.current) initialRouteRef.current = readShellRoute(); const [activeKey, setActiveKey] = useState( initialRouteRef.current.activeKey, ); const [view, setView] = useState(initialRouteRef.current.view); + const [temporarySession, setTemporarySession] = useState(null); + const temporarySessionRef = useRef(null); const [settingsInitialSection, setSettingsInitialSection] = useState(initialRouteRef.current.settingsSection); const [hostSidebarOpen, setHostSidebarOpen] = @@ -1004,19 +1025,33 @@ function Shell({ const showHostChrome = effectiveRuntimeSurface === "native"; const showMainSidebar = view !== "settings"; + const discardTemporaryChat = useCallback(() => { + const current = temporarySessionRef.current; + if (!current) return; + temporarySessionRef.current = null; + client.discardTemporaryChat(current.chatId); + setTemporarySession(null); + }, [client]); + const navigate = useCallback( (route: ShellRoute, options?: { replace?: boolean }) => { + if (route.view !== "chat" || route.activeKey !== QUICK_CHAT_KEY) { + discardTemporaryChat(); + } setActiveKey(route.activeKey); setView(route.view); setSettingsInitialSection(route.settingsSection); writeShellRoute(route, options?.replace); }, - [], + [discardTemporaryChat], ); useEffect(() => { const applyRoute = () => { const route = readShellRoute(); + if (route.view !== "chat" || route.activeKey !== QUICK_CHAT_KEY) { + discardTemporaryChat(); + } setActiveKey(route.activeKey); setView(route.view); setSettingsInitialSection(route.settingsSection); @@ -1027,7 +1062,15 @@ function Shell({ }; window.addEventListener("hashchange", applyRoute); return () => window.removeEventListener("hashchange", applyRoute); - }, []); + }, [discardTemporaryChat]); + + useEffect(() => { + return client.onStatus((status) => { + if (status !== "open") discardTemporaryChat(); + }); + }, [client, discardTemporaryChat]); + + useEffect(() => () => discardTemporaryChat(), [discardTemporaryChat]); useEffect(() => { let cancelled = false; @@ -1114,8 +1157,11 @@ function Shell({ const activeSession = useMemo(() => { if (!activeKey) return null; + if (isQuickChatKey(activeKey)) return temporarySession ?? quickSession; return sessions.find((s) => s.key === activeKey) ?? null; - }, [sessions, activeKey]); + }, [sessions, activeKey, quickSession, temporarySession]); + const quickChatActive = isQuickChatKey(activeKey); + const temporaryChatActive = quickChatActive && temporarySession !== null; const runningChatIdList = useMemo(() => Array.from(runningChatIds), [runningChatIds]); const updatedChatIdList = useMemo(() => Array.from(updatedChatIds), [updatedChatIds]); const activeChatId = activeSession?.chatId ?? null; @@ -1130,6 +1176,12 @@ function Shell({ }); }, [activeChatId]); const activeWorkspaceScope = useMemo(() => { + if (temporaryChatActive) { + return null; + } + if (quickChatActive) { + return workspaces?.default_scope ?? null; + } if (activeChatId && workspaceOverrides[activeChatId]) { return workspaceOverrides[activeChatId]; } @@ -1141,6 +1193,8 @@ function Shell({ activeChatId, activeSession?.workspaceScope, draftWorkspaceScope, + quickChatActive, + temporaryChatActive, workspaceOverrides, workspaces?.default_scope, ]); @@ -1161,7 +1215,10 @@ function Shell({ useEffect(() => { if (loading) return; - const knownChatIds = new Set(sessions.map((session) => session.chatId)); + const knownChatIds = new Set([ + QUICK_CHAT_ID, + ...sessions.map((session) => session.chatId), + ]); setUpdatedChatIds((current) => { const next = new Set( Array.from(current).filter((chatId) => knownChatIds.has(chatId)), @@ -1176,6 +1233,7 @@ function Shell({ useEffect(() => { if (loading || !activeKey) return; + if (isQuickChatKey(activeKey)) return; if (sessions.some((session) => session.key === activeKey)) return; const currentRoute = readShellRoute(); navigate( @@ -1417,6 +1475,28 @@ function Shell({ setMobileSidebarOpen(false); }, [navigate]); + const onOpenQuickChat = useCallback(() => { + setDraftWorkspaceScope(null); + setWorkspaceError(null); + setSessionSearchOpen(false); + navigate({ + view: "chat", + activeKey: QUICK_CHAT_KEY, + settingsSection: "overview", + }); + setMobileSidebarOpen(false); + }, [navigate]); + + const onToggleTemporaryChat = useCallback(() => { + if (temporarySessionRef.current) { + discardTemporaryChat(); + return; + } + const session = createTemporaryChatSession(); + temporarySessionRef.current = session; + setTemporarySession(session); + }, [discardTemporaryChat]); + const onNewChatInProject = useCallback( (projectPath: string, projectName: string) => { const base = workspaces?.default_scope ?? activeWorkspaceScope; @@ -1682,6 +1762,7 @@ function Shell({ setMobileSidebarOpen(false); const nextKey = (() => { if (!activeKey) return null; + if (isQuickChatKey(activeKey)) return activeKey; if (sessions.some((session) => session.key === activeKey)) return activeKey; return sessions[0]?.key ?? null; })(); @@ -1773,7 +1854,10 @@ function Shell({ }); }, [client, t]); - const onTurnEnd = useDeferredTitleRefresh(activeSession, refresh); + const onTurnEnd = useDeferredTitleRefresh( + quickChatActive ? null : activeSession, + refresh, + ); const onConfirmDelete = useCallback(async () => { if (!pendingDelete) return; @@ -1863,11 +1947,39 @@ function Shell({ }); }, []); - const headerTitle = activeSession + const headerTitle = temporaryChatActive + ? t("quickChat.temporary.title") + : quickChatActive + ? t("sidebar.quickChat") + : activeSession ? sidebarState.title_overrides[activeSession.key] || activeSession.title || deriveTitle(activeSession.preview, t("chat.newChat")) - : t("app.brand"); + : t("app.brand"); + + const temporaryChatAction = quickChatActive ? ( + + ) : undefined; useEffect(() => { if (view === "settings") { @@ -1900,10 +2012,12 @@ function Shell({ }, [activeSession, headerTitle, i18n.resolvedLanguage, t, view]); const sidebarProps = { - sessions, + sessions: regularSessions, activeKey: view === "chat" ? activeKey : null, loading, + quickChatActive: view === "chat" && quickChatActive, newChatActive: view === "chat" && activeKey === null, + onOpenQuickChat, onNewChat, onSelect: onSelectChat, onRequestDelete, @@ -2066,7 +2180,7 @@ function Shell({ {view !== "chat" && ( diff --git a/webui/src/components/Sidebar.tsx b/webui/src/components/Sidebar.tsx index 485727569..a9ad84f82 100644 --- a/webui/src/components/Sidebar.tsx +++ b/webui/src/components/Sidebar.tsx @@ -8,6 +8,7 @@ import { Archive, Brain, CalendarClock, + MessageCircle, Menu, Search, Settings, @@ -33,7 +34,9 @@ interface SidebarProps { sessions: ChatSummary[]; activeKey: string | null; loading: boolean; + quickChatActive: boolean; newChatActive: boolean; + onOpenQuickChat: () => void; onNewChat: () => void; onSelect: (key: string) => void; onRequestDelete: (key: string, label: string) => void; @@ -93,11 +96,13 @@ export function Sidebar(props: SidebarProps) { const toggleLabel = t("thread.header.toggleSidebar"); const newChatShortcut = newChatShortcutLabel(); const activeActionRef = useRef(null); - const activeActionId = props.newChatActive - ? "new-chat" - : props.activeUtility - ? `utility:${props.activeUtility}` - : null; + const activeActionId = props.quickChatActive + ? "quick-chat" + : props.newChatActive + ? "new-chat" + : props.activeUtility + ? `utility:${props.activeUtility}` + : null; return (