diff --git a/nanobot/agent/loop.py b/nanobot/agent/loop.py index bc807092..c1f52117 100644 --- a/nanobot/agent/loop.py +++ b/nanobot/agent/loop.py @@ -33,7 +33,6 @@ from nanobot.config.schema import AgentDefaults, ModelPresetConfig from nanobot.providers.base import LLMProvider from nanobot.providers.factory import ProviderSnapshot from nanobot.session.goal_state import ( - goal_state_ws_blob, runner_wall_llm_timeout_s, ) from nanobot.session.manager import Session, SessionManager @@ -42,10 +41,14 @@ from nanobot.utils.document import extract_documents from nanobot.utils.helpers import image_placeholder_text from nanobot.utils.helpers import truncate_text as truncate_text_fn from nanobot.utils.image_generation_intent import image_generation_prompt +from nanobot.utils.llm_runtime import LLMRuntime from nanobot.utils.runtime import EMPTY_FINAL_RESPONSE_MESSAGE from nanobot.utils.session_attachments import merge_turn_media_into_last_assistant -from nanobot.utils.webui_titles import mark_webui_session, maybe_generate_webui_title_after_turn -from nanobot.utils.webui_turn_helpers import publish_turn_run_status +from nanobot.utils.webui_turn_helpers import ( + WebuiTurnCoordinator, + build_bus_progress_callback, + mark_webui_session, +) if TYPE_CHECKING: from nanobot.config.schema import ( @@ -136,6 +139,11 @@ class AgentLoop: def tool_names(self) -> list[str]: return self.tools.tool_names + def llm_runtime(self) -> LLMRuntime: + """Return the current provider/model pair owned by this loop.""" + self._refresh_provider_snapshot() + return LLMRuntime(self.provider, self.model) + _RUNTIME_CHECKPOINT_KEY = "runtime_checkpoint" _PENDING_USER_TURN_KEY = "pending_user_turn" @@ -237,6 +245,11 @@ class AgentLoop: self.context = ContextBuilder(workspace, timezone=timezone, disabled_skills=disabled_skills) self.sessions = session_manager or SessionManager(workspace) + self._webui_turns = WebuiTurnCoordinator( + bus=self.bus, + sessions=self.sessions, + schedule_background=lambda coro: self._schedule_background(coro), + ) self.tools = ToolRegistry() # One file-read/write tracker per logical session. The tool registry is # shared by this loop, so tools resolve the active state via contextvars. @@ -524,34 +537,7 @@ class AgentLoop: self, msg: InboundMessage ) -> Callable[..., Awaitable[None]]: """Build a progress callback that publishes to the message bus.""" - - async def _bus_progress( - content: str, - *, - tool_hint: bool = False, - tool_events: list[dict[str, Any]] | None = None, - reasoning: bool = False, - reasoning_end: bool = False, - ) -> None: - meta = dict(msg.metadata or {}) - meta["_progress"] = True - meta["_tool_hint"] = tool_hint - if reasoning: - meta["_reasoning_delta"] = True - if reasoning_end: - meta["_reasoning_end"] = True - if tool_events: - meta["_tool_events"] = tool_events - await self.bus.publish_outbound( - OutboundMessage( - channel=msg.channel, - chat_id=msg.chat_id, - content=content, - metadata=meta, - ) - ) - - return _bus_progress + return build_bus_progress_callback(self.bus, msg) async def _build_retry_wait_callback( self, msg: InboundMessage @@ -938,38 +924,12 @@ class AgentLoop: content="", metadata=msg.metadata or {}, )) if msg.channel == "websocket": - # Signal that the turn is fully complete (all tools executed, - # final text streamed). This lets WS clients know when to - # definitively stop the loading indicator. turn_lat = self._pending_turn_latency_ms.pop(session_key, None) - turn_metadata: dict[str, Any] = {**msg.metadata, "_turn_end": True} - if turn_lat is not None: - turn_metadata["latency_ms"] = int(turn_lat) - sess_turn = self.sessions.get_or_create(session_key) - turn_metadata["goal_state"] = goal_state_ws_blob(sess_turn.metadata) - await self.bus.publish_outbound(OutboundMessage( - channel=msg.channel, chat_id=msg.chat_id, - content="", metadata=turn_metadata, - )) - if msg.metadata.get("webui") is True: - async def _generate_title_and_notify() -> None: - generated = await maybe_generate_webui_title_after_turn( - channel=msg.channel, - metadata=msg.metadata, - sessions=self.sessions, - session_key=session_key, - provider=self.provider, - model=self.model, - ) - if generated: - await self.bus.publish_outbound(OutboundMessage( - channel=msg.channel, - chat_id=msg.chat_id, - content="", - metadata={**msg.metadata, "_session_updated": True}, - )) - - self._schedule_background(_generate_title_and_notify()) + await self._webui_turns.handle_turn_end( + msg, + session_key=session_key, + latency_ms=turn_lat, + ) except asyncio.CancelledError: logger.info("Task cancelled for session {}", session_key) # Preserve partial context from the interrupted turn so @@ -1021,8 +981,9 @@ class AgentLoop: "Re-published {} leftover message(s) to bus for session {}", leftover, session_key, ) - await publish_turn_run_status(self.bus, msg, "idle") + await self._webui_turns.publish_run_status(msg, "idle") self._pending_turn_latency_ms.pop(session_key, None) + self._webui_turns.discard(session_key) async def close_mcp(self) -> None: """Drain pending background archives, then close MCP connections.""" @@ -1338,6 +1299,11 @@ class AgentLoop: "include_timestamps": True, } ctx.history = ctx.session.get_history(**_hist_kwargs) + self._webui_turns.capture_title_context( + ctx.session_key, + ctx.msg, + self.llm_runtime(), + ) ctx.initial_messages = self._build_initial_messages( ctx.msg, ctx.session, ctx.history, ctx.pending_summary @@ -1354,7 +1320,7 @@ class AgentLoop: return "ok" async def _state_run(self, ctx: TurnContext) -> str: - await publish_turn_run_status(self.bus, ctx.msg, "running") + await self._webui_turns.publish_run_status(ctx.msg, "running") result = await self._run_agent_loop( ctx.initial_messages, on_progress=ctx.on_progress, diff --git a/nanobot/agent/runner.py b/nanobot/agent/runner.py index 56482f75..776885ec 100644 --- a/nanobot/agent/runner.py +++ b/nanobot/agent/runner.py @@ -15,6 +15,12 @@ from loguru import logger from nanobot.agent.hook import AgentHook, AgentHookContext from nanobot.agent.tools.registry import ToolRegistry from nanobot.providers.base import LLMProvider, LLMResponse, ToolCallRequest +from nanobot.utils.file_edit_events import ( + build_file_edit_end_event, + build_file_edit_error_event, + build_file_edit_start_event, + prepare_file_edit_tracker, +) from nanobot.utils.helpers import ( IncrementalThinkExtractor, build_assistant_message, @@ -26,6 +32,10 @@ from nanobot.utils.helpers import ( strip_think, truncate_text, ) +from nanobot.utils.progress_events import ( + invoke_file_edit_progress, + on_progress_accepts_file_edit_events, +) from nanobot.utils.prompt_templates import render_template from nanobot.utils.runtime import ( EMPTY_FINAL_RESPONSE_MESSAGE, @@ -813,6 +823,30 @@ class AgentRunner: return prep_error + hint, event, ( RuntimeError(prep_error) if spec.fail_on_tool_error else None ) + emit_file_edit_events = ( + spec.progress_callback is not None + and on_progress_accepts_file_edit_events(spec.progress_callback) + ) + progress_callback = spec.progress_callback if emit_file_edit_events else None + file_edit_tracker = ( + prepare_file_edit_tracker( + call_id=tool_call.id, + tool_name=tool_call.name, + tool=tool, + workspace=spec.workspace, + params=params if isinstance(params, dict) else None, + ) + if progress_callback is not None + else None + ) + if file_edit_tracker is not None and progress_callback is not None: + await invoke_file_edit_progress( + progress_callback, + [build_file_edit_start_event( + file_edit_tracker, + params if isinstance(params, dict) else None, + )], + ) try: if tool is not None: result = await tool.execute(**params) @@ -821,6 +855,11 @@ class AgentRunner: except asyncio.CancelledError: raise except BaseException as exc: + if file_edit_tracker is not None and progress_callback is not None: + await invoke_file_edit_progress( + progress_callback, + [build_file_edit_error_event(file_edit_tracker, str(exc))], + ) event = { "name": tool_call.name, "status": "error", @@ -842,6 +881,11 @@ class AgentRunner: return payload, event, None if isinstance(result, str) and result.startswith("Error"): + if file_edit_tracker is not None and progress_callback is not None: + await invoke_file_edit_progress( + progress_callback, + [build_file_edit_error_event(file_edit_tracker, result)], + ) event = { "name": tool_call.name, "status": "error", @@ -860,6 +904,12 @@ class AgentRunner: return result + hint, event, RuntimeError(result) return result + hint, event, None + if file_edit_tracker is not None and progress_callback is not None: + await invoke_file_edit_progress( + progress_callback, + [build_file_edit_end_event(file_edit_tracker)], + ) + detail = "" if result is None else str(result) detail = detail.replace("\n", " ").strip() if not detail: diff --git a/nanobot/channels/websocket.py b/nanobot/channels/websocket.py index 26e00ff6..0202bd33 100644 --- a/nanobot/channels/websocket.py +++ b/nanobot/channels/websocket.py @@ -230,6 +230,25 @@ def _mask_secret_hint(secret: str | None) -> str | None: return f"{secret[:4]}••••{secret[-4:]}" +def _provider_requires_api_key(spec: Any) -> bool: + if spec.backend == "azure_openai": + return True + if spec.is_local or spec.is_direct: + return False + return True + + +def _provider_configured_for_settings(spec: Any, provider_config: Any) -> bool: + if _provider_requires_api_key(spec): + return bool(provider_config.api_key) + return bool( + provider_config.api_key + or provider_config.api_base + or getattr(provider_config, "region", None) + or getattr(provider_config, "profile", None) + ) + + _WEB_SEARCH_PROVIDER_OPTIONS: tuple[dict[str, str], ...] = ( {"name": "duckduckgo", "label": "DuckDuckGo", "credential": "none"}, {"name": "brave", "label": "Brave Search", "credential": "api_key"}, @@ -786,13 +805,14 @@ class WebSocketChannel(BaseChannel): providers = [] for spec in PROVIDERS: provider_config = getattr(config.providers, spec.name, None) - if provider_config is None or spec.is_oauth or spec.is_local: + if provider_config is None or spec.is_oauth: continue providers.append( { "name": spec.name, "label": spec.label, - "configured": bool(provider_config.api_key), + "configured": _provider_configured_for_settings(spec, provider_config), + "api_key_required": _provider_requires_api_key(spec), "api_key_hint": _mask_secret_hint(provider_config.api_key), "api_base": provider_config.api_base, "default_api_base": spec.default_api_base or None, @@ -862,7 +882,12 @@ class WebSocketChannel(BaseChannel): if find_by_name(provider) is None: return _http_error(400, "unknown provider") provider_config = getattr(config.providers, provider, None) - if provider_config is None or not provider_config.api_key: + spec = find_by_name(provider) + if ( + provider_config is None + or spec is None + or not _provider_configured_for_settings(spec, provider_config) + ): return _http_error(400, "provider is not configured") if defaults.provider != provider: defaults.provider = provider @@ -885,7 +910,7 @@ class WebSocketChannel(BaseChannel): if not provider_name: return _http_error(400, "provider is required") spec = find_by_name(provider_name) - if spec is None or spec.is_oauth or spec.is_local: + if spec is None or spec.is_oauth: return _http_error(400, "unknown provider") config = load_config() @@ -1581,6 +1606,7 @@ class WebSocketChannel(BaseChannel): if not conns: if ( msg.metadata.get("_progress") + or msg.metadata.get("_file_edit_events") or msg.metadata.get("_turn_end") or msg.metadata.get("_session_updated") or msg.metadata.get("_goal_status") @@ -1613,7 +1639,22 @@ class WebSocketChannel(BaseChannel): await self.send_turn_end(msg.chat_id, latency_ms=lat_i, goal_state=gs_blob) return if msg.metadata.get("_session_updated"): - await self.send_session_updated(msg.chat_id) + scope = msg.metadata.get("_session_update_scope") + await self.send_session_updated( + msg.chat_id, + scope=scope if isinstance(scope, str) else None, + ) + return + if msg.metadata.get("_file_edit_events"): + payload: dict[str, Any] = { + "event": "file_edit", + "chat_id": msg.chat_id, + "edits": msg.metadata["_file_edit_events"], + } + self._try_append_webui_transcript(msg.chat_id, payload) + raw = json.dumps(payload, ensure_ascii=False) + for connection in conns: + await self._safe_send_to(connection, raw, label=" ") return text = msg.content payload: dict[str, Any] = { @@ -1780,12 +1821,14 @@ class WebSocketChannel(BaseChannel): for connection in conns: await self._safe_send_to(connection, raw, label=" goal_status ") - async def send_session_updated(self, chat_id: str) -> None: + async def send_session_updated(self, chat_id: str, *, scope: str | None = None) -> None: """Notify clients that session metadata changed outside the main turn.""" conns = list(self._subs.get(chat_id, ())) if not conns: return body: dict[str, Any] = {"event": "session_updated", "chat_id": chat_id} + if scope: + body["scope"] = scope raw = json.dumps(body, ensure_ascii=False) for connection in conns: await self._safe_send_to(connection, raw, label=" session_updated ") diff --git a/nanobot/cli/commands.py b/nanobot/cli/commands.py index e561906f..69420543 100644 --- a/nanobot/cli/commands.py +++ b/nanobot/cli/commands.py @@ -968,8 +968,7 @@ def _run_gateway( hb_cfg = config.gateway.heartbeat heartbeat = HeartbeatService( workspace=config.workspace_path, - provider=agent.provider, - model=agent.model, + llm_runtime=agent.llm_runtime, on_execute=on_heartbeat_execute, on_notify=on_heartbeat_notify, interval_s=hb_cfg.interval_s, diff --git a/nanobot/heartbeat/service.py b/nanobot/heartbeat/service.py index b41ee7a1..55d26cf1 100644 --- a/nanobot/heartbeat/service.py +++ b/nanobot/heartbeat/service.py @@ -4,12 +4,12 @@ from __future__ import annotations import asyncio from pathlib import Path -from typing import TYPE_CHECKING, Any, Callable, Coroutine +from typing import Any, Callable, Coroutine from loguru import logger -if TYPE_CHECKING: - from nanobot.providers.base import LLMProvider +from nanobot.providers.base import LLMProvider +from nanobot.utils.llm_runtime import LLMRuntimeResolver, static_llm_runtime _HEARTBEAT_TOOL = [ { @@ -53,17 +53,21 @@ class HeartbeatService: def __init__( self, workspace: Path, - provider: LLMProvider, - model: str, + provider: LLMProvider | None = None, + model: str | None = None, on_execute: Callable[[str], Coroutine[Any, Any, str]] | None = None, on_notify: Callable[[str], Coroutine[Any, Any, None]] | None = None, interval_s: int = 30 * 60, enabled: bool = True, timezone: str | None = None, + llm_runtime: LLMRuntimeResolver | None = None, ): self.workspace = workspace - self.provider = provider - self.model = model + if llm_runtime is None: + if provider is None or model is None: + raise ValueError("HeartbeatService requires either llm_runtime or provider/model") + llm_runtime = static_llm_runtime(provider, model) + self._llm_runtime = llm_runtime self.on_execute = on_execute self.on_notify = on_notify self.interval_s = interval_s @@ -91,7 +95,9 @@ class HeartbeatService: """ from nanobot.utils.helpers import current_time_str - response = await self.provider.chat_with_retry( + llm = self._llm_runtime() + + response = await llm.provider.chat_with_retry( messages=[ {"role": "system", "content": "You are a heartbeat agent. Call the heartbeat tool to report your decision."}, {"role": "user", "content": ( @@ -101,7 +107,7 @@ class HeartbeatService: )}, ], tools=_HEARTBEAT_TOOL, - model=self.model, + model=llm.model, ) if not response.should_execute_tools: @@ -214,8 +220,9 @@ class HeartbeatService: ) return + llm = self._llm_runtime() should_notify = await evaluate_response( - response, tasks, self.provider, self.model, + response, tasks, llm.provider, llm.model, ) if should_notify and self.on_notify: logger.info("Heartbeat: completed, delivering response") diff --git a/nanobot/providers/registry.py b/nanobot/providers/registry.py index 4dba0c46..e6f02218 100644 --- a/nanobot/providers/registry.py +++ b/nanobot/providers/registry.py @@ -396,7 +396,7 @@ PROVIDERS: tuple[ProviderSpec, ...] = ( name="vllm", keywords=("vllm",), env_key="HOSTED_VLLM_API_KEY", - display_name="vLLM/Local", + display_name="vLLM", backend="openai_compat", is_local=True, ), diff --git a/nanobot/utils/file_edit_events.py b/nanobot/utils/file_edit_events.py new file mode 100644 index 00000000..8164aa18 --- /dev/null +++ b/nanobot/utils/file_edit_events.py @@ -0,0 +1,311 @@ +"""File-edit activity helpers for WebUI progress events.""" + +from __future__ import annotations + +import difflib +import json +from dataclasses import dataclass +from pathlib import Path +from typing import Any + + +TRACKED_FILE_EDIT_TOOLS = frozenset({"write_file", "edit_file", "notebook_edit"}) +_MAX_SNAPSHOT_BYTES = 2 * 1024 * 1024 + + +@dataclass(slots=True) +class FileSnapshot: + path: Path + exists: bool + text: str | None + unreadable: bool = False + binary: bool = False + oversized: bool = False + + @property + def countable(self) -> bool: + return ( + self.text is not None + and not self.binary + and not self.oversized + and not self.unreadable + ) + + +@dataclass(slots=True) +class FileEditTracker: + call_id: str + tool: str + path: Path + display_path: str + before: FileSnapshot + + +def is_file_edit_tool(tool_name: str | None) -> bool: + return bool(tool_name) and tool_name in TRACKED_FILE_EDIT_TOOLS + + +def resolve_file_edit_path( + tool: Any, + workspace: Path | None, + params: dict[str, Any] | None, +) -> Path | None: + """Resolve the target file path after tool argument preparation.""" + if not isinstance(params, dict): + return None + raw_path = params.get("path") + if not isinstance(raw_path, str) or not raw_path.strip(): + return None + resolver = getattr(tool, "_resolve", None) + if callable(resolver): + try: + resolved = resolver(raw_path) + if isinstance(resolved, Path): + return resolved + if resolved: + return Path(resolved) + except Exception: + return None + if workspace is None: + return Path(raw_path).expanduser().resolve() + return (workspace / raw_path).expanduser().resolve() + + +def display_file_edit_path(path: Path, workspace: Path | None) -> str: + if workspace is not None: + try: + return path.resolve().relative_to(workspace.resolve()).as_posix() + except Exception: + pass + return path.as_posix() + + +def read_file_snapshot(path: Path, *, max_bytes: int = _MAX_SNAPSHOT_BYTES) -> FileSnapshot: + try: + if not path.exists() or not path.is_file(): + return FileSnapshot(path=path, exists=False, text="") + size = path.stat().st_size + if size > max_bytes: + return FileSnapshot(path=path, exists=True, text=None, oversized=True) + raw = path.read_bytes() + except OSError: + return FileSnapshot(path=path, exists=path.exists(), text=None, unreadable=True) + if b"\x00" in raw: + return FileSnapshot(path=path, exists=True, text=None, binary=True) + try: + text = raw.decode("utf-8") + except UnicodeDecodeError: + return FileSnapshot(path=path, exists=True, text=None, binary=True) + return FileSnapshot(path=path, exists=True, text=text.replace("\r\n", "\n")) + + +def line_diff_stats(before: str | None, after: str | None) -> tuple[int, int]: + """Return ``(added, deleted)`` for a UTF-8 text line-level diff.""" + if before is None or after is None: + return 0, 0 + before_lines = before.replace("\r\n", "\n").splitlines() + after_lines = after.replace("\r\n", "\n").splitlines() + added = 0 + deleted = 0 + matcher = difflib.SequenceMatcher(a=before_lines, b=after_lines, autojunk=False) + for tag, i1, i2, j1, j2 in matcher.get_opcodes(): + if tag == "equal": + continue + if tag in ("replace", "delete"): + deleted += i2 - i1 + if tag in ("replace", "insert"): + added += j2 - j1 + return added, deleted + + +def prepare_file_edit_tracker( + *, + call_id: str, + tool_name: str, + tool: Any, + workspace: Path | None, + params: dict[str, Any] | None, +) -> FileEditTracker | None: + if not is_file_edit_tool(tool_name): + return None + path = resolve_file_edit_path(tool, workspace, params) + if path is None: + return None + before = read_file_snapshot(path) + return FileEditTracker( + call_id=str(call_id or ""), + tool=tool_name, + path=path, + display_path=display_file_edit_path(path, workspace), + before=before, + ) + + +def build_file_edit_start_event( + tracker: FileEditTracker, + params: dict[str, Any] | None, +) -> dict[str, Any]: + predicted_after = _predict_after_text(tracker.tool, params or {}, tracker.before) + if tracker.before.countable and predicted_after is not None: + added, deleted = line_diff_stats(tracker.before.text, predicted_after) + else: + added, deleted = 0, 0 + return _event_payload( + tracker, + phase="start", + status="editing", + added=added, + deleted=deleted, + approximate=True, + ) + + +def build_file_edit_end_event(tracker: FileEditTracker) -> dict[str, Any]: + after = read_file_snapshot(tracker.path) + if tracker.before.countable and after.countable: + added, deleted = line_diff_stats(tracker.before.text, after.text) + else: + added, deleted = 0, 0 + return _event_payload( + tracker, + phase="end", + status="done", + added=added, + deleted=deleted, + approximate=False, + binary=after.binary or after.oversized or after.unreadable, + ) + + +def build_file_edit_error_event(tracker: FileEditTracker, error: str | None = None) -> dict[str, Any]: + payload = _event_payload( + tracker, + phase="error", + status="error", + added=0, + deleted=0, + approximate=False, + ) + if error: + payload["error"] = error.strip()[:240] + return payload + + +def _event_payload( + tracker: FileEditTracker, + *, + phase: str, + status: str, + added: int, + deleted: int, + approximate: bool, + binary: bool = False, +) -> dict[str, Any]: + payload: dict[str, Any] = { + "version": 1, + "call_id": tracker.call_id, + "tool": tracker.tool, + "path": tracker.display_path, + "phase": phase, + "added": max(0, int(added)), + "deleted": max(0, int(deleted)), + "approximate": bool(approximate), + "status": status, + } + if binary: + payload["binary"] = True + return payload + + +def _predict_after_text( + tool_name: str, + params: dict[str, Any], + before: FileSnapshot, +) -> str | None: + if not before.countable: + return None + before_text = before.text or "" + if tool_name == "write_file": + content = params.get("content") + return content if isinstance(content, str) else "" + if tool_name == "edit_file": + old_text = params.get("old_text") + new_text = params.get("new_text") + if not isinstance(old_text, str) or not isinstance(new_text, str): + return None + replace_all = bool(params.get("replace_all")) + if old_text == "": + return new_text if not before.exists else before_text + if old_text in before_text: + if replace_all: + return before_text.replace(old_text, new_text) + return before_text.replace(old_text, new_text, 1) + return None + if tool_name == "notebook_edit": + return _predict_notebook_after_text(params, before_text) + return None + + +def _predict_notebook_after_text(params: dict[str, Any], before_text: str) -> str | None: + try: + nb = json.loads(before_text) if before_text.strip() else _empty_notebook() + except Exception: + return None + cells = nb.get("cells") + if not isinstance(cells, list): + return None + try: + cell_index = int(params.get("cell_index", 0)) + except (TypeError, ValueError): + return None + new_source = params.get("new_source") + source = new_source if isinstance(new_source, str) else "" + cell_type = params.get("cell_type") if params.get("cell_type") in ("code", "markdown") else "code" + mode = params.get("edit_mode") if params.get("edit_mode") in ("replace", "insert", "delete") else "replace" + if mode == "delete": + if 0 <= cell_index < len(cells): + cells.pop(cell_index) + else: + return None + elif mode == "insert": + insert_at = min(max(cell_index + 1, 0), len(cells)) + cells.insert(insert_at, _new_notebook_cell(source, str(cell_type))) + else: + if not (0 <= cell_index < len(cells)): + return None + cell = cells[cell_index] + if not isinstance(cell, dict): + return None + cell["source"] = source + cell["cell_type"] = cell_type + if cell_type == "code": + cell.setdefault("outputs", []) + cell.setdefault("execution_count", None) + else: + cell.pop("outputs", None) + cell.pop("execution_count", None) + nb["cells"] = cells + try: + return json.dumps(nb, indent=1, ensure_ascii=False) + except Exception: + return None + + +def _empty_notebook() -> dict[str, Any]: + return { + "nbformat": 4, + "nbformat_minor": 5, + "metadata": { + "kernelspec": {"display_name": "Python 3", "language": "python", "name": "python3"}, + "language_info": {"name": "python"}, + }, + "cells": [], + } + + +def _new_notebook_cell(source: str, cell_type: str) -> dict[str, Any]: + cell: dict[str, Any] = {"cell_type": cell_type, "source": source, "metadata": {}} + if cell_type == "code": + cell["outputs"] = [] + cell["execution_count"] = None + return cell diff --git a/nanobot/utils/llm_runtime.py b/nanobot/utils/llm_runtime.py new file mode 100644 index 00000000..a74f0d8c --- /dev/null +++ b/nanobot/utils/llm_runtime.py @@ -0,0 +1,22 @@ +"""Small helpers for passing the active LLM provider/model together.""" + +from __future__ import annotations + +from collections.abc import Callable +from dataclasses import dataclass + +from nanobot.providers.base import LLMProvider + + +@dataclass(frozen=True) +class LLMRuntime: + provider: LLMProvider + model: str + + +LLMRuntimeResolver = Callable[[], LLMRuntime] + + +def static_llm_runtime(provider: LLMProvider, model: str) -> LLMRuntimeResolver: + runtime = LLMRuntime(provider=provider, model=model) + return lambda: runtime diff --git a/nanobot/utils/progress_events.py b/nanobot/utils/progress_events.py index 10a282b9..ccf125ec 100644 --- a/nanobot/utils/progress_events.py +++ b/nanobot/utils/progress_events.py @@ -10,13 +10,21 @@ from nanobot.agent.hook import AgentHookContext def on_progress_accepts_tool_events(cb: Callable[..., Any]) -> bool: + return _on_progress_accepts(cb, "tool_events") + + +def on_progress_accepts_file_edit_events(cb: Callable[..., Any]) -> bool: + return _on_progress_accepts(cb, "file_edit_events") + + +def _on_progress_accepts(cb: Callable[..., Any], name: str) -> bool: try: sig = inspect.signature(cb) except (TypeError, ValueError): return False if any(p.kind == inspect.Parameter.VAR_KEYWORD for p in sig.parameters.values()): return True - return "tool_events" in sig.parameters + return name in sig.parameters async def invoke_on_progress( @@ -32,6 +40,15 @@ async def invoke_on_progress( await on_progress(content, tool_hint=tool_hint) +async def invoke_file_edit_progress( + on_progress: Callable[..., Awaitable[None]], + file_edit_events: list[dict[str, Any]], +) -> None: + if not file_edit_events or not on_progress_accepts_file_edit_events(on_progress): + return + await on_progress("", file_edit_events=file_edit_events) + + def build_tool_event_start_payload(tool_call: Any) -> dict[str, Any]: return { "version": 1, diff --git a/nanobot/utils/webui_titles.py b/nanobot/utils/webui_titles.py deleted file mode 100644 index 2d363f92..00000000 --- a/nanobot/utils/webui_titles.py +++ /dev/null @@ -1,138 +0,0 @@ -"""Helpers for WebUI chat title generation.""" - -from __future__ import annotations - -import re -from typing import Any - -from loguru import logger - -from nanobot.providers.base import LLMProvider -from nanobot.session.manager import Session, SessionManager -from nanobot.utils.helpers import truncate_text - -WEBUI_SESSION_METADATA_KEY = "webui" -WEBUI_TITLE_METADATA_KEY = "title" -WEBUI_TITLE_USER_EDITED_METADATA_KEY = "title_user_edited" -TITLE_MAX_CHARS = 60 - - -def mark_webui_session(session: Session, metadata: dict[str, Any]) -> bool: - """Persist a WebUI marker only when the inbound websocket frame opted in.""" - if metadata.get(WEBUI_SESSION_METADATA_KEY) is not True: - return False - session.metadata[WEBUI_SESSION_METADATA_KEY] = True - return True - - -def clean_generated_title(raw: str | None) -> str: - text = (raw or "").strip() - if not text: - return "" - text = re.sub(r"^\s*(title|标题)\s*[::]\s*", "", text, flags=re.IGNORECASE) - text = text.strip().strip("\"'`“”‘’") - text = re.sub(r"\s+", " ", text).strip() - text = text.rstrip("。.!!??,,;;:") - if len(text) > TITLE_MAX_CHARS: - text = text[: TITLE_MAX_CHARS - 1].rstrip() + "…" - return text - - -def _title_inputs(session: Session) -> tuple[str, str]: - user_text = "" - assistant_text = "" - for message in session.messages: - role = message.get("role") - content = message.get("content") - if not isinstance(content, str) or not content.strip(): - continue - if role == "user" and not user_text: - user_text = content.strip() - elif role == "assistant" and not assistant_text: - assistant_text = content.strip() - if user_text and assistant_text: - break - return user_text, assistant_text - - -async def maybe_generate_webui_title( - *, - sessions: SessionManager, - session_key: str, - provider: LLMProvider, - model: str, -) -> bool: - """Generate and persist a short title for WebUI-owned sessions only.""" - 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: - return False - current_title = session.metadata.get(WEBUI_TITLE_METADATA_KEY) - if isinstance(current_title, str) and current_title.strip(): - return False - - user_text, assistant_text = _title_inputs(session) - if not user_text: - return False - - prompt = ( - "Generate a concise title for this chat.\n" - "Rules:\n" - "- Use the same language as the user when practical.\n" - "- 3 to 8 words.\n" - "- No quotes.\n" - "- No punctuation at the end.\n" - "- Return only the title.\n\n" - f"User: {truncate_text(user_text, 1_000)}" - ) - if assistant_text: - prompt += f"\nAssistant: {truncate_text(assistant_text, 1_000)}" - - try: - response = await provider.chat_with_retry( - [ - { - "role": "system", - "content": ( - "You write short, neutral chat titles. " - "Return only the title text." - ), - }, - {"role": "user", "content": prompt}, - ], - tools=None, - model=model, - max_tokens=32, - temperature=0.2, - retry_mode="standard", - ) - except Exception: - logger.debug("Failed to generate webui session title for {}", session_key, exc_info=True) - return False - - title = clean_generated_title(response.content) - if not title or title.lower().startswith("error"): - return False - session.metadata[WEBUI_TITLE_METADATA_KEY] = title - sessions.save(session) - return True - - -async def maybe_generate_webui_title_after_turn( - *, - channel: str, - metadata: dict[str, Any], - sessions: SessionManager, - session_key: str, - provider: LLMProvider, - model: str, -) -> bool: - if channel != "websocket" or metadata.get(WEBUI_SESSION_METADATA_KEY) is not True: - return False - return await maybe_generate_webui_title( - sessions=sessions, - session_key=session_key, - provider=provider, - model=model, - ) diff --git a/nanobot/utils/webui_transcript.py b/nanobot/utils/webui_transcript.py index dde0e916..bee71c54 100644 --- a/nanobot/utils/webui_transcript.py +++ b/nanobot/utils/webui_transcript.py @@ -125,11 +125,25 @@ def replay_transcript_to_ui_messages( buffer_message_id: str | None = None buffer_parts: list[str] = [] suppress_until_turn_end = False + active_activity_segment_id: str | None = None + active_file_edit_segment_id: str | None = None + activity_segment_counter = 0 _ts_base = int(time.time() * 1000) def _new_id(prefix: str, idx: int) -> str: return f"{prefix}-{idx}-{uuid.uuid4().hex[:8]}" + def _new_activity_segment(*, activate: bool = True) -> str: + nonlocal active_activity_segment_id, activity_segment_counter + activity_segment_counter += 1 + segment_id = f"activity-{activity_segment_counter}" + if activate: + active_activity_segment_id = segment_id + return segment_id + + def _ensure_activity_segment() -> str: + return active_activity_segment_id or _new_activity_segment() + def attach_reasoning_chunk(prev: list[dict[str, Any]], chunk: str, idx: int) -> None: for i in range(len(prev) - 1, -1, -1): candidate = prev[i] @@ -151,12 +165,19 @@ def replay_transcript_to_ui_messages( **candidate, "reasoning": (str(candidate.get("reasoning") or "")) + chunk, "reasoningStreaming": True, + "activitySegmentId": candidate.get("activitySegmentId") or _ensure_activity_segment(), } return if not has_answer and candidate.get("isStreaming"): - prev[i] = {**candidate, "reasoning": chunk, "reasoningStreaming": True} + prev[i] = { + **candidate, + "reasoning": chunk, + "reasoningStreaming": True, + "activitySegmentId": candidate.get("activitySegmentId") or _ensure_activity_segment(), + } return break + segment = _ensure_activity_segment() prev.append( { "id": _new_id("as", idx), @@ -165,6 +186,7 @@ def replay_transcript_to_ui_messages( "isStreaming": True, "reasoning": chunk, "reasoningStreaming": True, + "activitySegmentId": segment, "createdAt": _ts_base + idx, }, ) @@ -221,6 +243,7 @@ def replay_transcript_to_ui_messages( return def absorb_complete(extra: dict[str, Any], idx: int) -> None: + nonlocal active_activity_segment_id last = messages[-1] if messages else None if last and is_reasoning_only_placeholder(last): messages[-1] = { @@ -238,10 +261,76 @@ def replay_transcript_to_ui_messages( **extra, }, ) + active_activity_segment_id = None + + def _file_edit_key(edit: dict[str, Any]) -> str: + return "|".join( + str(edit.get(k) or "") + for k in ("call_id", "tool", "path") + ) + + def upsert_file_edits(edits: list[dict[str, Any]], idx: int) -> None: + nonlocal active_file_edit_segment_id + if not edits: + return + last = messages[-1] if messages else None + if ( + active_file_edit_segment_id + and last + and last.get("kind") == "trace" + and last.get("fileEdits") + ): + segment = active_file_edit_segment_id + else: + segment = _new_activity_segment(activate=False) + active_file_edit_segment_id = segment + if not ( + last + and last.get("kind") == "trace" + and not last.get("isStreaming") + and last.get("fileEdits") + and last.get("activitySegmentId") == segment + ): + messages.append( + { + "id": _new_id("tr", idx), + "role": "tool", + "kind": "trace", + "content": "", + "traces": [], + "fileEdits": [], + "activitySegmentId": segment, + "createdAt": _ts_base + idx, + }, + ) + last = messages[-1] + existing = list(last.get("fileEdits") or []) + index_by_key = { + _file_edit_key(edit): pos + for pos, edit in enumerate(existing) + if isinstance(edit, dict) + } + for edit in edits: + if not isinstance(edit, dict): + continue + key = _file_edit_key(edit) + if key in index_by_key: + pos = index_by_key[key] + existing[pos] = {**existing[pos], **edit} + else: + index_by_key[key] = len(existing) + existing.append(dict(edit)) + messages[-1] = { + **last, + "fileEdits": existing, + "activitySegmentId": last.get("activitySegmentId") or segment, + } for idx, rec in enumerate(lines): ev = rec.get("event") if ev == "user": + active_activity_segment_id = None + active_file_edit_segment_id = None text = rec.get("text") text_s = text if isinstance(text, str) else "" media_paths = rec.get("media_paths") @@ -264,6 +353,12 @@ def replay_transcript_to_ui_messages( messages.append(row) continue + if ev == "file_edit": + raw_edits = rec.get("edits") + if isinstance(raw_edits, list): + upsert_file_edits([e for e in raw_edits if isinstance(e, dict)], idx) + continue + if ev == "delta": if suppress_until_turn_end: continue @@ -338,14 +433,21 @@ def replay_transcript_to_ui_messages( trace_lines = structured if structured else ([text] if isinstance(text, str) and text else []) if not trace_lines: continue + segment = _ensure_activity_segment() last = messages[-1] if messages else None - if last and last.get("kind") == "trace" and not last.get("isStreaming"): + if ( + last + and last.get("kind") == "trace" + and not last.get("isStreaming") + and (last.get("activitySegmentId") in (None, segment)) + ): prev_traces = list(last.get("traces") or [last.get("content")]) merged_traces = prev_traces + trace_lines messages[-1] = { **last, "traces": merged_traces, "content": trace_lines[-1], + "activitySegmentId": last.get("activitySegmentId") or segment, } else: messages.append( @@ -355,6 +457,7 @@ def replay_transcript_to_ui_messages( "kind": "trace", "content": trace_lines[-1], "traces": trace_lines, + "activitySegmentId": segment, "createdAt": _ts_base + idx, }, ) @@ -389,6 +492,8 @@ def replay_transcript_to_ui_messages( if ev == "turn_end": suppress_until_turn_end = False + active_activity_segment_id = None + active_file_edit_segment_id = None for i, m in enumerate(messages): if m.get("isStreaming"): messages[i] = {**m, "isStreaming": False} diff --git a/nanobot/utils/webui_turn_helpers.py b/nanobot/utils/webui_turn_helpers.py index 3fbca372..6a3ac2ba 100644 --- a/nanobot/utils/webui_turn_helpers.py +++ b/nanobot/utils/webui_turn_helpers.py @@ -6,17 +6,163 @@ AgentLoop uses these without importing a concrete channel plugin; only from __future__ import annotations +import re import time +from collections.abc import Awaitable, Callable +from dataclasses import dataclass, field from typing import Any +from loguru import logger + from nanobot.bus.events import InboundMessage, OutboundMessage from nanobot.bus.queue import MessageBus +from nanobot.providers.base import LLMProvider +from nanobot.session.goal_state import goal_state_ws_blob +from nanobot.session.manager import Session, SessionManager +from nanobot.utils.helpers import truncate_text +from nanobot.utils.llm_runtime import LLMRuntime + +WEBUI_SESSION_METADATA_KEY = "webui" +WEBUI_TITLE_METADATA_KEY = "title" +WEBUI_TITLE_USER_EDITED_METADATA_KEY = "title_user_edited" +TITLE_MAX_CHARS = 60 +TITLE_GENERATION_MAX_TOKENS = 96 +TITLE_GENERATION_REASONING_EFFORT = "none" # Wall-clock turn start per ``chat_id`` (websocket only). Survives browser refresh while the # gateway process stays up; cleared on idle/stop and implicitly dropped on restart. _WEBSOCKET_TURN_WALL_STARTED_AT: dict[str, float] = {} +def mark_webui_session(session: Session, metadata: dict[str, Any]) -> bool: + """Persist a WebUI marker only when the inbound websocket frame opted in.""" + if metadata.get(WEBUI_SESSION_METADATA_KEY) is not True: + return False + session.metadata[WEBUI_SESSION_METADATA_KEY] = True + return True + + +def clean_generated_title(raw: str | None) -> str: + text = (raw or "").strip() + if not text: + return "" + text = re.sub(r"^\s*(title|标题)\s*[::]\s*", "", text, flags=re.IGNORECASE) + text = text.strip().strip("\"'`“”‘’") + text = re.sub(r"\s+", " ", text).strip() + text = text.rstrip("。.!!??,,;;:") + if len(text) > TITLE_MAX_CHARS: + text = text[: TITLE_MAX_CHARS - 1].rstrip() + "…" + return text + + +def _title_inputs(session: Session) -> tuple[str, str]: + user_text = "" + assistant_text = "" + for message in session.messages: + if message.get("_command") is True: + continue + role = message.get("role") + content = message.get("content") + if not isinstance(content, str) or not content.strip(): + continue + if role == "user" and not user_text: + user_text = content.strip() + elif role == "assistant" and not assistant_text: + assistant_text = content.strip() + if user_text and assistant_text: + break + return user_text, assistant_text + + +async def maybe_generate_webui_title( + *, + sessions: SessionManager, + session_key: str, + provider: LLMProvider, + model: str, +) -> bool: + """Generate and persist a short title for WebUI-owned sessions only.""" + 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: + return False + current_title = session.metadata.get(WEBUI_TITLE_METADATA_KEY) + if isinstance(current_title, str) and current_title.strip(): + return False + + user_text, assistant_text = _title_inputs(session) + if not user_text: + return False + + prompt = ( + "Generate a concise title for this chat.\n" + "Rules:\n" + "- Use the same language as the user when practical.\n" + "- 3 to 8 words.\n" + "- No quotes.\n" + "- No punctuation at the end.\n" + "- Return only the title.\n\n" + f"User: {truncate_text(user_text, 1_000)}" + ) + if assistant_text: + prompt += f"\nAssistant: {truncate_text(assistant_text, 1_000)}" + + try: + response = await provider.chat_with_retry( + [ + { + "role": "system", + "content": ( + "You write short, neutral chat titles. " + "Return only the title text." + ), + }, + {"role": "user", "content": prompt}, + ], + tools=None, + model=model, + max_tokens=TITLE_GENERATION_MAX_TOKENS, + temperature=0.2, + reasoning_effort=TITLE_GENERATION_REASONING_EFFORT, + retry_mode="standard", + ) + except Exception: + logger.debug("Failed to generate webui session title for {}", session_key, exc_info=True) + return False + + title = clean_generated_title(response.content) + if not title or title.lower().startswith("error"): + logger.debug( + "WebUI title generation returned no usable title for {} (finish_reason={})", + session_key, + response.finish_reason, + ) + return False + session.metadata[WEBUI_TITLE_METADATA_KEY] = title + sessions.save(session) + return True + + +async def maybe_generate_webui_title_after_turn( + *, + channel: str, + metadata: dict[str, Any], + sessions: SessionManager, + session_key: str, + provider: LLMProvider, + model: str, +) -> bool: + if channel != "websocket" or metadata.get(WEBUI_SESSION_METADATA_KEY) is not True: + return False + return await maybe_generate_webui_title( + sessions=sessions, + session_key=session_key, + provider=provider, + model=model, + ) + + def websocket_turn_wall_started_at(chat_id: str) -> float | None: """Return ``time.time()`` when the active user turn began, if still running.""" return _WEBSOCKET_TURN_WALL_STARTED_AT.get(chat_id) @@ -46,3 +192,156 @@ async def publish_turn_run_status(bus: MessageBus, msg: InboundMessage, status: metadata=meta, ), ) + + +def build_bus_progress_callback( + bus: MessageBus, + msg: InboundMessage, +) -> Callable[..., Awaitable[None]]: + """Return the bus progress callback for agent runtime events.""" + + async def _publish_progress( + content: str, + *, + tool_hint: bool = False, + tool_events: list[dict[str, Any]] | None = None, + file_edit_events: list[dict[str, Any]] | None = None, + reasoning: bool = False, + reasoning_end: bool = False, + ) -> None: + meta = dict(msg.metadata or {}) + meta["_progress"] = True + meta["_tool_hint"] = tool_hint + if reasoning: + meta["_reasoning_delta"] = True + if reasoning_end: + meta["_reasoning_end"] = True + if tool_events: + meta["_tool_events"] = tool_events + if file_edit_events: + meta["_file_edit_events"] = file_edit_events + await bus.publish_outbound( + OutboundMessage( + channel=msg.channel, + chat_id=msg.chat_id, + content=content, + metadata=meta, + ) + ) + + if msg.channel == "websocket": + async def _websocket_progress( + content: str, + *, + tool_hint: bool = False, + tool_events: list[dict[str, Any]] | None = None, + file_edit_events: list[dict[str, Any]] | None = None, + reasoning: bool = False, + reasoning_end: bool = False, + ) -> None: + await _publish_progress( + content, + tool_hint=tool_hint, + tool_events=tool_events, + file_edit_events=file_edit_events, + reasoning=reasoning, + reasoning_end=reasoning_end, + ) + + return _websocket_progress + + async def _bus_progress( + content: str, + *, + tool_hint: bool = False, + tool_events: list[dict[str, Any]] | None = None, + reasoning: bool = False, + reasoning_end: bool = False, + ) -> None: + await _publish_progress( + content, + tool_hint=tool_hint, + tool_events=tool_events, + reasoning=reasoning, + reasoning_end=reasoning_end, + ) + + return _bus_progress + + +@dataclass +class WebuiTurnCoordinator: + """Own the WebUI/WebSocket wire details that hang off AgentLoop turns.""" + + bus: MessageBus + sessions: SessionManager + schedule_background: Callable[[Awaitable[None]], None] + _title_contexts: dict[str, LLMRuntime] = field(default_factory=dict) + + def capture_title_context( + self, + session_key: str, + msg: InboundMessage, + llm: LLMRuntime, + ) -> None: + if msg.channel == "websocket" and msg.metadata.get("webui") is True: + self._title_contexts[session_key] = llm + + def discard(self, session_key: str) -> None: + self._title_contexts.pop(session_key, None) + + async def publish_run_status(self, msg: InboundMessage, status: str) -> None: + await publish_turn_run_status(self.bus, msg, status) + + async def handle_turn_end( + self, + msg: InboundMessage, + *, + session_key: str, + latency_ms: int | None, + ) -> None: + if msg.channel != "websocket": + return + + turn_metadata: dict[str, Any] = {**msg.metadata, "_turn_end": True} + if latency_ms is not None: + turn_metadata["latency_ms"] = int(latency_ms) + session = self.sessions.get_or_create(session_key) + turn_metadata["goal_state"] = goal_state_ws_blob(session.metadata) + await self.bus.publish_outbound(OutboundMessage( + channel=msg.channel, + chat_id=msg.chat_id, + content="", + metadata=turn_metadata, + )) + self._schedule_title_update(msg, session_key=session_key) + + def _schedule_title_update(self, msg: InboundMessage, *, session_key: str) -> None: + title_context = self._title_contexts.pop(session_key, None) + if msg.metadata.get("webui") is not True or title_context is None: + return + + async def _generate_title_and_notify( + title_llm: LLMRuntime = title_context, + ) -> None: + generated = await maybe_generate_webui_title_after_turn( + channel=msg.channel, + metadata=msg.metadata, + sessions=self.sessions, + session_key=session_key, + provider=title_llm.provider, + model=title_llm.model, + ) + if generated: + await self.bus.publish_outbound(OutboundMessage( + channel=msg.channel, + chat_id=msg.chat_id, + content="", + metadata={ + **msg.metadata, + "_session_updated": True, + "_session_update_scope": "metadata", + }, + )) + + self.schedule_background(_generate_title_and_notify()) diff --git a/tests/agent/test_heartbeat_service.py b/tests/agent/test_heartbeat_service.py index 8f563cff..fe7b5425 100644 --- a/tests/agent/test_heartbeat_service.py +++ b/tests/agent/test_heartbeat_service.py @@ -4,6 +4,7 @@ import pytest from nanobot.heartbeat.service import HeartbeatService from nanobot.providers.base import LLMProvider, LLMResponse, ToolCallRequest +from nanobot.utils.llm_runtime import LLMRuntime class DummyProvider(LLMProvider): @@ -11,9 +12,11 @@ class DummyProvider(LLMProvider): super().__init__() self._responses = list(responses) self.calls = 0 + self.models: list[str | None] = [] async def chat(self, *args, **kwargs) -> LLMResponse: self.calls += 1 + self.models.append(kwargs.get("model")) if self._responses: return self._responses.pop(0) return LLMResponse(content="", tool_calls=[]) @@ -215,6 +218,51 @@ async def test_tick_suppresses_when_evaluator_says_no(tmp_path, monkeypatch) -> assert notified == [] +def test_tick_uses_runtime_provider_and_model(tmp_path, monkeypatch) -> None: + """Preset changes must apply to heartbeat decision and post-run evaluation.""" + (tmp_path / "HEARTBEAT.md").write_text("- [ ] check runtime model", encoding="utf-8") + + runtime_provider = DummyProvider([ + LLMResponse( + content="", + tool_calls=[ + ToolCallRequest( + id="hb_1", + name="heartbeat", + arguments={"action": "run", "tasks": "check runtime model"}, + ) + ], + ), + ]) + runtime_model = "openai/gpt-4.1" + + executed: list[str] = [] + evaluated: list[tuple[LLMProvider, str]] = [] + + async def _on_execute(tasks: str) -> str: + executed.append(tasks) + return "runtime model produced a user-facing update" + + async def _eval_capture(response, tasks, provider, model): + evaluated.append((provider, model)) + return False + + service = HeartbeatService( + workspace=tmp_path, + llm_runtime=lambda: LLMRuntime(runtime_provider, runtime_model), + on_execute=_on_execute, + ) + + monkeypatch.setattr("nanobot.utils.evaluator.evaluate_response", _eval_capture) + + asyncio.run(service._tick()) + + assert runtime_provider.calls == 1 + assert runtime_provider.models == [runtime_model] + assert executed == ["check runtime model"] + assert evaluated == [(runtime_provider, runtime_model)] + + @pytest.mark.asyncio async def test_decide_retries_transient_error_then_succeeds(tmp_path, monkeypatch) -> None: provider = DummyProvider([ @@ -286,4 +334,3 @@ async def test_decide_prompt_includes_current_time(tmp_path) -> None: user_msg = captured_messages[1] assert user_msg["role"] == "user" assert "Current Time:" in user_msg["content"] - diff --git a/tests/agent/test_loop_progress.py b/tests/agent/test_loop_progress.py index fcf6198c..43a69143 100644 --- a/tests/agent/test_loop_progress.py +++ b/tests/agent/test_loop_progress.py @@ -6,10 +6,15 @@ from unittest.mock import AsyncMock, MagicMock import pytest +import nanobot.agent.runner as runner_module from nanobot.agent.loop import AgentLoop from nanobot.bus.events import InboundMessage from nanobot.bus.queue import MessageBus from nanobot.providers.base import LLMResponse, ToolCallRequest +from nanobot.utils.progress_events import ( + invoke_file_edit_progress, + on_progress_accepts_file_edit_events, +) def _make_loop(tmp_path: Path) -> AgentLoop: @@ -82,6 +87,142 @@ class TestToolEventProgress: ), ] + @pytest.mark.asyncio + async def test_write_file_emits_file_edit_progress(self, tmp_path: Path) -> None: + loop = _make_loop(tmp_path) + target = tmp_path / "foo.txt" + target.write_text("old\n", encoding="utf-8") + tool_call = ToolCallRequest( + id="call-write", + name="write_file", + arguments={"path": "foo.txt", "content": "new\nextra\n"}, + ) + calls = iter([ + LLMResponse(content="", tool_calls=[tool_call]), + LLMResponse(content="Done", tool_calls=[]), + ]) + loop.provider.chat_with_retry = AsyncMock(side_effect=lambda *a, **kw: next(calls)) + loop.tools.get_definitions = MagicMock(return_value=[]) + loop.tools.prepare_call = MagicMock( + return_value=(None, {"path": "foo.txt", "content": "new\nextra\n"}, None), + ) + + async def execute(name: str, params: dict) -> str: + target.write_text(params["content"], encoding="utf-8") + return "ok" + + loop.tools.execute = AsyncMock(side_effect=execute) + file_events: list[dict] = [] + + async def on_progress( + content: str, + *, + tool_hint: bool = False, + tool_events: list[dict] | None = None, + file_edit_events: list[dict] | None = None, + ) -> None: + if file_edit_events: + file_events.extend(file_edit_events) + + final_content, _, _, _, _ = await loop._run_agent_loop([], on_progress=on_progress) + + assert final_content == "Done" + assert [event["phase"] for event in file_events] == ["start", "end"] + assert file_events[0] == { + "version": 1, + "call_id": "call-write", + "tool": "write_file", + "path": "foo.txt", + "phase": "start", + "added": 2, + "deleted": 1, + "approximate": True, + "status": "editing", + } + assert file_events[1]["status"] == "done" + assert file_events[1]["approximate"] is False + assert (file_events[1]["added"], file_events[1]["deleted"]) == (2, 1) + + @pytest.mark.asyncio + async def test_file_edit_snapshot_skipped_when_progress_callback_cannot_emit_file_edits( + self, + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + loop = _make_loop(tmp_path) + target = tmp_path / "foo.txt" + target.write_text("old\n", encoding="utf-8") + tool_call = ToolCallRequest( + id="call-write", + name="write_file", + arguments={"path": "foo.txt", "content": "new\n"}, + ) + calls = iter([ + LLMResponse(content="", tool_calls=[tool_call]), + LLMResponse(content="Done", tool_calls=[]), + ]) + loop.provider.chat_with_retry = AsyncMock(side_effect=lambda *a, **kw: next(calls)) + loop.tools.get_definitions = MagicMock(return_value=[]) + loop.tools.prepare_call = MagicMock( + return_value=(None, {"path": "foo.txt", "content": "new\n"}, None), + ) + + async def execute(name: str, params: dict) -> str: + target.write_text(params["content"], encoding="utf-8") + return "ok" + + loop.tools.execute = AsyncMock(side_effect=execute) + prepare_tracker = MagicMock(side_effect=AssertionError("unexpected file snapshot")) + monkeypatch.setattr(runner_module, "prepare_file_edit_tracker", prepare_tracker) + + async def on_progress( + content: str, + *, + tool_hint: bool = False, + tool_events: list[dict] | None = None, + ) -> None: + pass + + final_content, _, _, _, _ = await loop._run_agent_loop([], on_progress=on_progress) + + assert final_content == "Done" + assert target.read_text(encoding="utf-8") == "new\n" + prepare_tracker.assert_not_called() + + @pytest.mark.asyncio + async def test_exec_does_not_emit_file_edit_progress(self, tmp_path: Path) -> None: + loop = _make_loop(tmp_path) + tool_call = ToolCallRequest( + id="call-exec", + name="exec", + arguments={"command": "printf hi > foo.txt"}, + ) + calls = iter([ + LLMResponse(content="", tool_calls=[tool_call]), + LLMResponse(content="Done", tool_calls=[]), + ]) + loop.provider.chat_with_retry = AsyncMock(side_effect=lambda *a, **kw: next(calls)) + loop.tools.get_definitions = MagicMock(return_value=[]) + loop.tools.prepare_call = MagicMock( + return_value=(None, {"command": "printf hi > foo.txt"}, None), + ) + loop.tools.execute = AsyncMock(return_value="ok") + file_events: list[dict] = [] + + async def on_progress( + content: str, + *, + tool_hint: bool = False, + tool_events: list[dict] | None = None, + file_edit_events: list[dict] | None = None, + ) -> None: + if file_edit_events: + file_events.extend(file_edit_events) + + await loop._run_agent_loop([], on_progress=on_progress) + + assert file_events == [] + @pytest.mark.asyncio async def test_bus_progress_forwards_tool_events_to_outbound_metadata(self, tmp_path: Path) -> None: """When run() handles a bus message, _tool_events lands in OutboundMessage metadata.""" @@ -130,6 +271,44 @@ class TestToolEventProgress: assert finish["phase"] == "end" assert finish["result"] == "file.txt" + @pytest.mark.asyncio + async def test_bus_progress_forwards_file_edit_events_for_websocket_only(self, tmp_path: Path) -> None: + bus = MessageBus() + provider = MagicMock() + provider.get_default_model.return_value = "test-model" + loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model") + edit_events = [{ + "call_id": "call-write", + "tool": "write_file", + "path": "foo.txt", + "phase": "start", + "added": 1, + "deleted": 0, + "approximate": True, + "status": "editing", + }] + + websocket_progress = await loop._build_bus_progress_callback(InboundMessage( + channel="websocket", + sender_id="u1", + chat_id="chat1", + content="edit", + )) + assert on_progress_accepts_file_edit_events(websocket_progress) is True + await websocket_progress("", file_edit_events=edit_events) + outbound = await bus.consume_outbound() + assert outbound.metadata["_file_edit_events"] == edit_events + + telegram_progress = await loop._build_bus_progress_callback(InboundMessage( + channel="telegram", + sender_id="u1", + chat_id="chat2", + content="edit", + )) + assert on_progress_accepts_file_edit_events(telegram_progress) is False + await invoke_file_edit_progress(telegram_progress, edit_events) + assert bus.outbound_size == 0 + @pytest.mark.asyncio async def test_non_streaming_channel_does_not_publish_codex_progress_deltas( self, @@ -353,8 +532,93 @@ class TestToolEventProgress: assert session_updated is not None assert (session_updated.metadata or {}).get("_session_updated") is True + assert (session_updated.metadata or {}).get("_session_update_scope") == "metadata" assert provider.chat_with_retry.await_count == 2 + @pytest.mark.asyncio + async def test_webui_title_generation_uses_turn_model_snapshot( + self, + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + bus = MessageBus() + provider = MagicMock() + provider.get_default_model.return_value = "test-model" + provider.chat_with_retry = AsyncMock(return_value=LLMResponse(content="Done", tool_calls=[])) + loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model") + loop.tools.get_definitions = MagicMock(return_value=[]) + loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign] + + captured: dict[str, object] = {} + + async def fake_title_after_turn(**kwargs: object) -> bool: + captured.update(kwargs) + return False + + monkeypatch.setattr( + "nanobot.utils.webui_turn_helpers.maybe_generate_webui_title_after_turn", + fake_title_after_turn, + ) + scheduled_title: list[object] = [] + + def schedule_background(coro: object) -> None: + name = getattr(coro, "__qualname__", "") + if "_generate_title_and_notify" in name: + scheduled_title.append(coro) + elif hasattr(coro, "close"): + coro.close() + + loop._schedule_background = schedule_background # type: ignore[method-assign] + + await loop._dispatch(InboundMessage( + channel="websocket", + sender_id="u1", + chat_id="chat1", + content="say hello", + metadata={"webui": True}, + )) + + assert len(scheduled_title) == 1 + loop.provider = MagicMock() + loop.model = "switched-after-turn" + + await scheduled_title[0] # type: ignore[misc] + + assert captured["provider"] is provider + assert captured["model"] == "test-model" + + @pytest.mark.asyncio + async def test_webui_command_turn_does_not_schedule_title_generation( + self, + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + bus = MessageBus() + provider = MagicMock() + provider.get_default_model.return_value = "test-model" + provider.chat_with_retry = AsyncMock(return_value=LLMResponse(content="Done", tool_calls=[])) + loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model") + + async def fake_title_after_turn(**_kwargs: object) -> bool: + raise AssertionError("command-only turns should not generate titles") + + monkeypatch.setattr( + "nanobot.utils.webui_turn_helpers.maybe_generate_webui_title_after_turn", + fake_title_after_turn, + ) + scheduled: list[object] = [] + loop._schedule_background = scheduled.append # type: ignore[method-assign] + + await loop._dispatch(InboundMessage( + channel="websocket", + sender_id="u1", + chat_id="chat1", + content="/model", + metadata={"webui": True}, + )) + + assert scheduled == [] + @pytest.mark.asyncio async def test_non_websocket_dispatch_does_not_publish_turn_end_marker(self, tmp_path: Path) -> None: bus = MessageBus() diff --git a/tests/agent/test_loop_save_turn.py b/tests/agent/test_loop_save_turn.py index ed78e719..9814c386 100644 --- a/tests/agent/test_loop_save_turn.py +++ b/tests/agent/test_loop_save_turn.py @@ -10,12 +10,16 @@ from nanobot.bus.events import InboundMessage from nanobot.bus.queue import MessageBus from nanobot.providers.base import LLMResponse from nanobot.session.goal_state import GOAL_STATE_KEY -from nanobot.session.manager import Session -from nanobot.utils.webui_titles import ( +from nanobot.session.manager import Session, SessionManager +from nanobot.utils.webui_turn_helpers import ( + TITLE_GENERATION_MAX_TOKENS, + TITLE_GENERATION_REASONING_EFFORT, WEBUI_SESSION_METADATA_KEY, WEBUI_TITLE_METADATA_KEY, + WebuiTurnCoordinator, maybe_generate_webui_title, ) +from nanobot.utils.llm_runtime import LLMRuntime def _mk_loop() -> AgentLoop: @@ -33,6 +37,22 @@ def _make_full_loop(tmp_path: Path) -> AgentLoop: return AgentLoop(bus=MessageBus(), provider=provider, workspace=tmp_path, model="test-model") +def test_agent_loop_llm_runtime_reflects_current_provider_and_model(tmp_path: Path) -> None: + loop = _make_full_loop(tmp_path) + runtime = loop.llm_runtime() + + assert runtime.provider is loop.provider + assert runtime.model == "test-model" + + next_provider = MagicMock() + loop.provider = next_provider + loop.model = "next-model" + runtime = loop.llm_runtime() + + assert runtime.provider is next_provider + assert runtime.model == "next-model" + + @pytest.mark.asyncio async def test_generate_webui_title_only_for_marked_webui_sessions(tmp_path: Path) -> None: loop = _make_full_loop(tmp_path) @@ -55,6 +75,11 @@ async def test_generate_webui_title_only_for_marked_webui_sessions(tmp_path: Pat assert generated is True assert session.metadata[WEBUI_TITLE_METADATA_KEY] == "优化 WebUI 侧边栏" loop.provider.chat_with_retry.assert_awaited_once() + assert loop.provider.chat_with_retry.await_args.kwargs["max_tokens"] == TITLE_GENERATION_MAX_TOKENS + assert ( + loop.provider.chat_with_retry.await_args.kwargs["reasoning_effort"] + == TITLE_GENERATION_REASONING_EFFORT + ) @pytest.mark.asyncio @@ -79,6 +104,80 @@ async def test_generate_webui_title_skips_plain_websocket_sessions(tmp_path: Pat loop.provider.chat_with_retry.assert_not_awaited() +@pytest.mark.asyncio +async def test_generate_webui_title_ignores_command_only_sessions(tmp_path: Path) -> None: + loop = _make_full_loop(tmp_path) + session = loop.sessions.get_or_create("websocket:command-title") + session.metadata[WEBUI_SESSION_METADATA_KEY] = True + session.add_message("user", "/model deep", _command=True) + session.add_message( + "assistant", + "Switched model preset to `deep`.\n- Model: `deepseek-v4-pro`", + _command=True, + ) + loop.sessions.save(session) + + generated = await maybe_generate_webui_title( + sessions=loop.sessions, + session_key="websocket:command-title", + provider=loop.provider, + model=loop.model, + ) + + assert generated is False + assert WEBUI_TITLE_METADATA_KEY not in session.metadata + loop.provider.chat_with_retry.assert_not_awaited() + + +def test_webui_title_update_uses_captured_llm_runtime( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + bus = MessageBus() + sessions = SessionManager(tmp_path) + scheduled: list[object] = [] + captured: dict[str, object] = {} + + async def fake_title_after_turn(**kwargs: object) -> bool: + captured.update(kwargs) + return False + + monkeypatch.setattr( + "nanobot.utils.webui_turn_helpers.maybe_generate_webui_title_after_turn", + fake_title_after_turn, + ) + coordinator = WebuiTurnCoordinator( + bus=bus, + sessions=sessions, + schedule_background=lambda coro: scheduled.append(coro), + ) + provider = MagicMock() + msg = InboundMessage( + channel="websocket", + sender_id="u1", + chat_id="chat1", + content="say hello", + metadata={"webui": True}, + ) + + coordinator.capture_title_context( + "websocket:chat1", + msg, + LLMRuntime(provider, "turn-model"), + ) + asyncio.run(coordinator.handle_turn_end( + msg, + session_key="websocket:chat1", + latency_ms=None, + )) + + assert len(scheduled) == 1 + asyncio.run(scheduled[0]) # type: ignore[arg-type] + + assert captured["provider"] is provider + assert captured["model"] == "turn-model" + + def test_save_turn_skips_multimodal_user_when_only_runtime_context() -> None: loop = _mk_loop() session = Session(key="test:runtime-only") diff --git a/tests/agent/test_runtime_refresh.py b/tests/agent/test_runtime_refresh.py index a6b19a9d..b36b1899 100644 --- a/tests/agent/test_runtime_refresh.py +++ b/tests/agent/test_runtime_refresh.py @@ -47,3 +47,28 @@ def test_provider_refresh_updates_all_model_dependents(tmp_path: Path) -> None: assert loop.dream.provider is new_provider assert loop.dream.model == "new-model" assert loop.dream._runner.provider is new_provider + + +def test_llm_runtime_refreshes_provider_snapshot(tmp_path: Path) -> None: + old_provider = _provider("old-model") + new_provider = _provider("new-model", max_tokens=456) + loop = AgentLoop( + bus=MessageBus(), + provider=old_provider, + workspace=tmp_path, + model="old-model", + context_window_tokens=1000, + provider_snapshot_loader=lambda: ProviderSnapshot( + provider=new_provider, + model="new-model", + context_window_tokens=2000, + signature=("new-model",), + ), + ) + + runtime = loop.llm_runtime() + + assert runtime.provider is new_provider + assert runtime.model == "new-model" + assert loop.provider is new_provider + assert loop.runner.provider is new_provider diff --git a/tests/channels/test_websocket_channel.py b/tests/channels/test_websocket_channel.py index 9b481e25..c6f9d66a 100644 --- a/tests/channels/test_websocket_channel.py +++ b/tests/channels/test_websocket_channel.py @@ -370,6 +370,55 @@ async def test_send_progress_includes_structured_tool_events() -> None: ] +@pytest.mark.asyncio +async def test_send_file_edit_progress_uses_file_edit_event() -> None: + bus = MagicMock() + channel = WebSocketChannel({"enabled": True, "allowFrom": ["*"]}, bus) + mock_ws = AsyncMock() + channel._attach(mock_ws, "chat-1") + + await channel.send(OutboundMessage( + channel="websocket", + chat_id="chat-1", + content="", + metadata={ + "_progress": True, + "_file_edit_events": [ + { + "version": 1, + "phase": "start", + "call_id": "call-1", + "tool": "write_file", + "path": "src/app.py", + "added": 12, + "deleted": 2, + "approximate": True, + "status": "editing", + } + ], + }, + )) + + payload = json.loads(mock_ws.send.await_args.args[0]) + assert payload == { + "event": "file_edit", + "chat_id": "chat-1", + "edits": [ + { + "version": 1, + "phase": "start", + "call_id": "call-1", + "tool": "write_file", + "path": "src/app.py", + "added": 12, + "deleted": 2, + "approximate": True, + "status": "editing", + } + ], + } + + @pytest.mark.asyncio async def test_send_progress_includes_agent_ui_blob() -> None: bus = MagicMock() @@ -758,6 +807,25 @@ async def test_send_session_updated_emits_session_updated_event() -> None: assert body == {"event": "session_updated", "chat_id": "chat-1"} +@pytest.mark.asyncio +async def test_send_session_updated_includes_scope_when_present() -> None: + bus = MagicMock() + channel = WebSocketChannel({"enabled": True, "allowFrom": ["*"]}, bus) + mock_ws = AsyncMock() + channel._attach(mock_ws, "chat-1") + + await channel.send(OutboundMessage( + channel="websocket", + chat_id="chat-1", + content="", + metadata={"_session_updated": True, "_session_update_scope": "metadata"}, + )) + + mock_ws.send.assert_awaited_once() + body = json.loads(mock_ws.send.await_args.args[0]) + assert body == {"event": "session_updated", "chat_id": "chat-1", "scope": "metadata"} + + @pytest.mark.asyncio async def test_send_non_connection_closed_exception_is_raised() -> None: bus = MagicMock() @@ -946,7 +1014,12 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist( providers = {provider["name"]: provider for provider in body["providers"]} assert providers["openai"]["configured"] is True assert providers["openai"]["api_key_hint"] == "secr••••-key" + assert providers["azure_openai"]["api_key_required"] is True assert providers["openrouter"]["configured"] is False + assert providers["openrouter"]["api_key_required"] is True + assert providers["atomic_chat"]["configured"] is False + assert providers["atomic_chat"]["api_key_required"] is False + assert providers["atomic_chat"]["default_api_base"] == "http://localhost:1337/v1" assert body["agent"]["has_api_key"] is True assert body["web_search"]["provider"] == "brave" assert body["web_search"]["api_key_hint"] == "brav••••cret" @@ -969,10 +1042,24 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist( assert provider_rows["openrouter"]["configured"] is True assert "sk-or-test" not in provider_updated.text + local_provider_updated = await _http_get( + "http://127.0.0.1:" + f"{port}/api/settings/provider/update?provider=atomic_chat" + "&api_base=http%3A%2F%2Flocalhost%3A1337%2Fv1", + headers={"Authorization": "Bearer tok"}, + ) + assert local_provider_updated.status_code == 200 + local_provider_body = local_provider_updated.json() + local_provider_rows = { + provider["name"]: provider for provider in local_provider_body["providers"] + } + assert local_provider_rows["atomic_chat"]["configured"] is True + assert "localhost:1337" in local_provider_updated.text + updated = await _http_get( "http://127.0.0.1:" - f"{port}/api/settings/update?model=openrouter/test" - "&provider=openrouter", + f"{port}/api/settings/update?model=atomic_chat/test" + "&provider=atomic_chat", headers={"Authorization": "Bearer tok"}, ) assert updated.status_code == 200 @@ -992,10 +1079,11 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist( assert search_body["web_search"]["base_url"] == "https://search.example.com" saved = load_config(config_path) - assert saved.agents.defaults.model == "openrouter/test" - assert saved.agents.defaults.provider == "openrouter" + assert saved.agents.defaults.model == "atomic_chat/test" + assert saved.agents.defaults.provider == "atomic_chat" assert saved.providers.openrouter.api_key == "sk-or-test" assert saved.providers.openrouter.api_base == "https://openrouter.ai/api/v1" + assert saved.providers.atomic_chat.api_base == "http://localhost:1337/v1" assert saved.tools.web.search.provider == "searxng" assert saved.tools.web.search.api_key == "" assert saved.tools.web.search.base_url == "https://search.example.com" diff --git a/tests/cli/test_commands.py b/tests/cli/test_commands.py index 90c2ce87..2778ddbb 100644 --- a/tests/cli/test_commands.py +++ b/tests/cli/test_commands.py @@ -1170,6 +1170,7 @@ def test_gateway_cron_evaluator_receives_scheduled_reminder_context( self.model = "test-model" self.provider = kwargs.get("provider", object()) self.tools = {} + seen["agent"] = self async def process_direct(self, *_args, **_kwargs): return OutboundMessage( @@ -1218,6 +1219,11 @@ def test_gateway_cron_evaluator_receives_scheduled_reminder_context( assert isinstance(cron, _FakeCron) assert cron.on_job is not None + runtime_provider = object() + agent = seen["agent"] + agent.provider = runtime_provider + agent.model = "runtime-model" + job = CronJob( id="cron-1", name="stretch", @@ -1233,8 +1239,8 @@ def test_gateway_cron_evaluator_receives_scheduled_reminder_context( assert response == "Time to stretch." assert seen["response"] == "Time to stretch." - assert seen["provider"] is provider - assert seen["model"] == "test-model" + assert seen["provider"] is runtime_provider + assert seen["model"] == "runtime-model" assert seen["task_context"] == ( "The scheduled time has arrived. Deliver this reminder to the user now, " "as a brief and natural message in their language. Speak directly to them — " @@ -1543,6 +1549,9 @@ def test_gateway_health_endpoint_binds_and_serves_expected_responses( self.dream = _FakeDream() self.sessions = _FakeSessionManager() + def llm_runtime(self) -> None: + return None + async def run(self) -> None: await asyncio.Event().wait() diff --git a/tests/utils/test_file_edit_events.py b/tests/utils/test_file_edit_events.py new file mode 100644 index 00000000..6176a5e3 --- /dev/null +++ b/tests/utils/test_file_edit_events.py @@ -0,0 +1,83 @@ +from __future__ import annotations + +from pathlib import Path + +from nanobot.utils.file_edit_events import ( + build_file_edit_end_event, + build_file_edit_start_event, + line_diff_stats, + prepare_file_edit_tracker, + read_file_snapshot, +) + + +def test_line_diff_stats_counts_replacements_insertions_and_deletions() -> None: + added, deleted = line_diff_stats("a\nb\nc\n", "a\nB\nc\nd\n") + assert (added, deleted) == (2, 1) + + +def test_line_diff_stats_normalizes_crlf() -> None: + assert line_diff_stats("a\r\nb\r\n", "a\nb\nc\n") == (1, 0) + + +def test_write_file_start_predicts_and_end_calibrates_exact_diff(tmp_path: Path) -> None: + target = tmp_path / "notes.txt" + target.write_text("old\nkeep\n", encoding="utf-8") + params = {"path": "notes.txt", "content": "new\nkeep\nextra\n"} + tracker = prepare_file_edit_tracker( + call_id="call-write", + tool_name="write_file", + tool=None, + workspace=tmp_path, + params=params, + ) + + assert tracker is not None + start = build_file_edit_start_event(tracker, params) + assert start == { + "version": 1, + "call_id": "call-write", + "tool": "write_file", + "path": "notes.txt", + "phase": "start", + "added": 2, + "deleted": 1, + "approximate": True, + "status": "editing", + } + + target.write_text("new\nkeep\nextra\n", encoding="utf-8") + end = build_file_edit_end_event(tracker) + assert end["phase"] == "end" + assert end["status"] == "done" + assert end["approximate"] is False + assert (end["added"], end["deleted"]) == (2, 1) + + +def test_binary_file_is_reported_but_not_counted(tmp_path: Path) -> None: + target = tmp_path / "data.bin" + target.write_bytes(b"\x00\x01before") + tracker = prepare_file_edit_tracker( + call_id="call-bin", + tool_name="edit_file", + tool=None, + workspace=tmp_path, + params={"path": "data.bin", "old_text": "before", "new_text": "after"}, + ) + + assert tracker is not None + assert not read_file_snapshot(target).countable + target.write_bytes(b"\x00\x01after") + event = build_file_edit_end_event(tracker) + assert event["binary"] is True + assert (event["added"], event["deleted"]) == (0, 0) + + +def test_untracked_tools_do_not_prepare_file_edit_tracker(tmp_path: Path) -> None: + assert prepare_file_edit_tracker( + call_id="call-exec", + tool_name="exec", + tool=None, + workspace=tmp_path, + params={"path": "created-by-shell.txt"}, + ) is None diff --git a/tests/utils/test_webui_transcript.py b/tests/utils/test_webui_transcript.py index 419abbfc..f13380f4 100644 --- a/tests/utils/test_webui_transcript.py +++ b/tests/utils/test_webui_transcript.py @@ -42,6 +42,62 @@ def test_replay_delta_and_turn_end(tmp_path, monkeypatch) -> None: assert msgs[1]["latencyMs"] == 42 +def test_replay_file_edit_event_creates_file_activity(tmp_path, monkeypatch) -> None: + monkeypatch.setattr("nanobot.config.paths.get_data_dir", lambda: tmp_path) + key = "websocket:t-file" + for ev in ( + {"event": "user", "chat_id": "t-file", "text": "edit"}, + { + "event": "message", + "chat_id": "t-file", + "text": 'write_file({"path":"foo.txt"})', + "kind": "tool_hint", + }, + { + "event": "file_edit", + "chat_id": "t-file", + "edits": [ + { + "version": 1, + "call_id": "call-write", + "tool": "write_file", + "path": "foo.txt", + "phase": "end", + "added": 2, + "deleted": 1, + "approximate": False, + "status": "done", + }, + ], + }, + ): + append_transcript_object(key, ev) + + msgs = replay_transcript_to_ui_messages(read_transcript_lines(key)) + + assert len(msgs) == 3 + assert msgs[1]["kind"] == "trace" + assert msgs[1]["traces"] == ['write_file({"path":"foo.txt"})'] + assert "fileEdits" not in msgs[1] + assert msgs[2]["kind"] == "trace" + assert msgs[2]["traces"] == [] + assert msgs[2]["fileEdits"] == [ + { + "version": 1, + "call_id": "call-write", + "tool": "write_file", + "path": "foo.txt", + "phase": "end", + "added": 2, + "deleted": 1, + "approximate": False, + "status": "done", + }, + ] + assert msgs[2]["activitySegmentId"] + assert msgs[2]["activitySegmentId"] != msgs[1]["activitySegmentId"] + + def test_build_response_schema(monkeypatch, tmp_path) -> None: from nanobot.utils.webui_transcript import build_webui_thread_response diff --git a/webui/src/App.tsx b/webui/src/App.tsx index e8dc0722..7ff9bae2 100644 --- a/webui/src/App.tsx +++ b/webui/src/App.tsx @@ -7,7 +7,8 @@ import { ThreadShell } from "@/components/thread/ThreadShell"; import { Sheet, SheetContent } from "@/components/ui/sheet"; import { useSessions } from "@/hooks/useSessions"; -import { useTheme } from "@/hooks/useTheme"; +import { useDeferredTitleRefresh } from "@/hooks/useDeferredTitleRefresh"; +import { ThemeProvider, useTheme } from "@/hooks/useTheme"; import { cn } from "@/lib/utils"; import { clearSavedSecret, @@ -16,6 +17,7 @@ import { loadSavedSecret, saveSecret, } from "@/lib/bootstrap"; +import { deriveTitle } from "@/lib/format"; import { NanobotClient } from "@/lib/nanobot-client"; import { ClientProvider, useClient } from "@/providers/ClientProvider"; import type { ChatSummary } from "@/lib/types"; @@ -30,14 +32,30 @@ type BootState = status: "ready"; client: NanobotClient; token: string; + tokenExpiresAt: number; modelName: string | null; }; const SIDEBAR_STORAGE_KEY = "nanobot-webui.sidebar"; const RESTART_STARTED_KEY = "nanobot-webui.restartStartedAt"; const SIDEBAR_WIDTH = 272; +const TOKEN_REFRESH_MARGIN_MS = 30_000; +const TOKEN_REFRESH_MIN_DELAY_MS = 5_000; type ShellView = "chat" | "settings"; +function bootstrapTokenExpiresAt(expiresInSeconds: number): number { + return Date.now() + Math.max(0, expiresInSeconds) * 1000; +} + +function tokenRefreshDelayMs(expiresAt: number): number { + const remaining = Math.max(0, expiresAt - Date.now()); + const margin = Math.min( + TOKEN_REFRESH_MARGIN_MS, + Math.max(1_000, remaining / 2), + ); + return Math.max(TOKEN_REFRESH_MIN_DELAY_MS, remaining - margin); +} + function AuthForm({ failed, onSecret, @@ -106,6 +124,7 @@ function readSidebarOpen(): boolean { export default function App() { const { t } = useTranslation(); const [state, setState] = useState({ status: "loading" }); + const bootstrapSecretRef = useRef(""); const bootstrapWithSecret = useCallback( (secret: string) => { @@ -117,22 +136,37 @@ export default function App() { if (cancelled) return; if (secret) saveSecret(secret); const url = deriveWsUrl(boot.ws_path, boot.token); - const client = new NanobotClient({ + let client: NanobotClient; + client = new NanobotClient({ url, onReauth: async () => { try { - const refreshed = await fetchBootstrap("", secret); - return deriveWsUrl(refreshed.ws_path, refreshed.token); + const refreshed = await fetchBootstrap("", bootstrapSecretRef.current); + const refreshedUrl = deriveWsUrl(refreshed.ws_path, refreshed.token); + const tokenExpiresAt = bootstrapTokenExpiresAt(refreshed.expires_in); + setState((current) => + current.status === "ready" && current.client === client + ? { + ...current, + token: refreshed.token, + tokenExpiresAt, + modelName: refreshed.model_name ?? current.modelName, + } + : current, + ); + return refreshedUrl; } catch { return null; } }, }); + bootstrapSecretRef.current = secret; client.connect(); setState({ status: "ready", client, token: boot.token, + tokenExpiresAt: bootstrapTokenExpiresAt(boot.expires_in), modelName: boot.model_name ?? null, }); } catch (e) { @@ -152,6 +186,35 @@ export default function App() { [], ); + useEffect(() => { + if (state.status !== "ready") return; + const client = state.client; + const timer = window.setTimeout(async () => { + try { + const boot = await fetchBootstrap("", bootstrapSecretRef.current); + const url = deriveWsUrl(boot.ws_path, boot.token); + const tokenExpiresAt = bootstrapTokenExpiresAt(boot.expires_in); + client.updateUrl(url); + setState((current) => + current.status === "ready" && current.client === client + ? { + ...current, + token: boot.token, + tokenExpiresAt, + modelName: boot.model_name ?? current.modelName, + } + : current, + ); + } catch (e) { + const msg = (e as Error).message; + if (msg.includes("HTTP 401") || msg.includes("HTTP 403")) { + setState({ status: "auth", failed: true }); + } + } + }, tokenRefreshDelayMs(state.tokenExpiresAt)); + return () => window.clearTimeout(timer); + }, [state]); + useEffect(() => { const saved = loadSavedSecret(); return bootstrapWithSecret(saved); @@ -219,7 +282,13 @@ export default function App() { ); } -function Shell({ onModelNameChange, onLogout }: { onModelNameChange: (modelName: string | null) => void; onLogout: () => void }) { +function Shell({ + onModelNameChange, + onLogout, +}: { + onModelNameChange: (modelName: string | null) => void; + onLogout: () => void; +}) { const { t, i18n } = useTranslation(); const { client } = useClient(); const { theme, toggle } = useTheme(); @@ -362,9 +431,7 @@ function Shell({ onModelNameChange, onLogout }: { onModelNameChange: (modelName: }); }, [client, t]); - const onTurnEnd = useCallback(() => { - void refresh(); - }, [refresh]); + const onTurnEnd = useDeferredTitleRefresh(activeSession, refresh); const onConfirmDelete = useCallback(async () => { if (!pendingDelete) return; @@ -386,8 +453,7 @@ function Shell({ onModelNameChange, onLogout }: { onModelNameChange: (modelName: const headerTitle = activeSession ? activeSession.title || - activeSession.preview || - t("chat.fallbackTitle", { id: activeSession.chatId.slice(0, 6) }) + deriveTitle(activeSession.preview, t("chat.newChat")) : t("app.brand"); useEffect(() => { @@ -415,93 +481,95 @@ function Shell({ onModelNameChange, onLogout }: { onModelNameChange: (modelName: const showMainSidebar = view !== "settings"; return ( -
- {/* Desktop sidebar: in normal flow, so the thread area width stays honest. */} - {showMainSidebar ? ( - - ) : null} - - {showMainSidebar ? ( - setMobileSidebarOpen(open)} - > - - - - - ) : null} - -
-
- -
- {view === "settings" && ( -
-
- )} -
+ {view === "settings" && ( +
+ +
+ )} + - setPendingDelete(null)} - onConfirm={onConfirmDelete} - /> - {restartToast ? ( -
- {restartToast} -
- ) : null} -
+ setPendingDelete(null)} + onConfirm={onConfirmDelete} + /> + {restartToast ? ( +
+ {restartToast} +
+ ) : null} + + ); } diff --git a/webui/src/components/ChatList.tsx b/webui/src/components/ChatList.tsx index fc667883..a5107651 100644 --- a/webui/src/components/ChatList.tsx +++ b/webui/src/components/ChatList.tsx @@ -7,6 +7,7 @@ import { DropdownMenuItem, DropdownMenuTrigger, } from "@/components/ui/dropdown-menu"; +import { deriveTitle } from "@/lib/format"; import { cn } from "@/lib/utils"; import type { ChatSummary } from "@/lib/types"; @@ -64,8 +65,11 @@ export function ChatList({ const fallbackTitle = t("chat.fallbackTitle", { id: s.chatId.slice(0, 6), }); - const rawLabel = (s.title || s.preview)?.trim(); - const title = rawLabel || fallbackTitle; + const generatedTitle = s.title?.trim() || ""; + const title = + generatedTitle || deriveTitle(s.preview, t("chat.newChat")); + const tooltipTitle = + generatedTitle || deriveTitle(s.preview, fallbackTitle); return (
  • onSelect(s.key)} - title={rawLabel || fallbackTitle} + title={tooltipTitle} className="min-w-0 flex-1 overflow-hidden py-1.5 text-left" > {title} diff --git a/webui/src/components/CodeBlock.tsx b/webui/src/components/CodeBlock.tsx index c19a7864..4e3b8b73 100644 --- a/webui/src/components/CodeBlock.tsx +++ b/webui/src/components/CodeBlock.tsx @@ -1,44 +1,75 @@ -import { useCallback, useEffect, useState } from "react"; +import { Suspense, lazy, useCallback, useState } from "react"; import { Check, Copy } from "lucide-react"; import { useTranslation } from "react-i18next"; -import { Prism as SyntaxHighlighter } from "react-syntax-highlighter"; -import { - oneDark, - oneLight, -} from "react-syntax-highlighter/dist/esm/styles/prism"; +import { useThemeValue } from "@/hooks/useTheme"; import { cn } from "@/lib/utils"; interface CodeBlockProps { language?: string; code: string; className?: string; + highlight?: boolean; } -/** Read dark mode straight from the DOM — stays in sync with Tailwind's `dark:`. */ -function useIsDark() { - const [isDark, setIsDark] = useState(() => - typeof document !== "undefined" - ? document.documentElement.classList.contains("dark") - : true, +interface HighlightedCodeProps { + language?: string; + code: string; + isDark: boolean; +} + +const LazyHighlightedCode = lazy(async () => { + const [ + { default: SyntaxHighlighter }, + { default: oneDark }, + { default: oneLight }, + ] = await Promise.all([ + import("react-syntax-highlighter/dist/esm/prism-async-light"), + import("react-syntax-highlighter/dist/esm/styles/prism/one-dark"), + import("react-syntax-highlighter/dist/esm/styles/prism/one-light"), + ]); + + return { + default({ language, code, isDark }: HighlightedCodeProps) { + return ( + + {code} + + ); + }, + }; +}); + +function PlainCodeFallback({ code }: { code: string }) { + return ( +
    +      {code}
    +    
    ); - - useEffect(() => { - const el = document.documentElement; - const observer = new MutationObserver(() => { - setIsDark(el.classList.contains("dark")); - }); - observer.observe(el, { attributeFilter: ["class"] }); - return () => observer.disconnect(); - }, []); - - return isDark; } -export function CodeBlock({ language, code, className }: CodeBlockProps) { +export function CodeBlock({ + language, + code, + className, + highlight = true, +}: CodeBlockProps) { const { t } = useTranslation(); const [copied, setCopied] = useState(false); - const isDark = useIsDark(); + const isDark = useThemeValue() === "dark"; const onCopy = useCallback(() => { if (!navigator.clipboard) return; @@ -86,20 +117,13 @@ export function CodeBlock({ language, code, className }: CodeBlockProps) { {copied ? t("code.copied") : t("code.copy")}
    - - {code} - + {highlight ? ( + }> + + + ) : ( + + )} ); } diff --git a/webui/src/components/ConnectionBadge.tsx b/webui/src/components/ConnectionBadge.tsx index 7616ddbe..a09aadd2 100644 --- a/webui/src/components/ConnectionBadge.tsx +++ b/webui/src/components/ConnectionBadge.tsx @@ -36,21 +36,25 @@ export function ConnectionBadge() { status === "connecting" || status === "reconnecting" || status === "error"; + const label = t(`connection.${status}`); return ( - + {pulsing && ( )} - + - {t(`connection.${status}`)} + {label} ); } diff --git a/webui/src/components/FileReferenceChip.tsx b/webui/src/components/FileReferenceChip.tsx new file mode 100644 index 00000000..18e63d1c --- /dev/null +++ b/webui/src/components/FileReferenceChip.tsx @@ -0,0 +1,220 @@ +import { + Tooltip, + TooltipContent, + TooltipProvider, + TooltipTrigger, +} from "@/components/ui/tooltip"; +import { cn } from "@/lib/utils"; + +type FileReferenceKind = + | "default" + | "css" + | "html" + | "json" + | "markdown" + | "notebook" + | "python" + | "react" + | "typescript"; + +interface FileReferenceChipProps { + path: string; + display?: "name" | "path"; + active?: boolean; + className?: string; + textClassName?: string; + testId?: string; +} + +export function FileReferenceChip({ + path, + display = "name", + active = false, + className, + textClassName, + testId = "inline-file-path", +}: FileReferenceChipProps) { + const { name } = splitFilePath(path); + const kind = fileKindForPath(path); + const displayText = display === "path" ? path.replace(/\\/g, "/") : name; + return ( + + + + + + + + {displayText} + + + + + + {path} + + + + ); +} + +export function isLikelyFilePath(value: string): boolean { + const raw = value.trim(); + if (!raw || raw.includes("\n")) return false; + if (/^[a-z][a-z0-9+.-]*:\/\//i.test(raw)) return false; + if (!/[\\/]/.test(raw) && !/^(dockerfile|makefile|readme|package-lock\.json)$/i.test(raw)) { + return false; + } + const normalized = raw.replace(/\\/g, "/"); + const name = normalized.split("/").filter(Boolean).pop() ?? normalized; + if (!name || name === "." || name === "..") return false; + if (/^(dockerfile|makefile|readme|package-lock\.json)$/i.test(name)) return true; + return /\.[a-z0-9][a-z0-9_-]{0,12}$/i.test(name); +} + +function splitFilePath(path: string): { directory: string; name: string } { + const normalized = path.replace(/\\/g, "/"); + const slash = normalized.lastIndexOf("/"); + if (slash < 0) return { directory: "", name: path }; + return { + directory: normalized.slice(0, slash + 1), + name: normalized.slice(slash + 1) || normalized, + }; +} + +function fileKindForPath(path: string): FileReferenceKind { + const normalized = path.toLowerCase(); + const name = normalized.split(/[\\/]/).pop() ?? normalized; + const ext = name.includes(".") ? name.split(".").pop() ?? "" : ""; + if (name === "dockerfile") { + return "default"; + } + switch (ext) { + case "py": + case "pyi": + return "python"; + case "jsx": + case "tsx": + return "react"; + case "ts": + return "typescript"; + case "html": + case "htm": + return "html"; + case "css": + case "scss": + case "sass": + return "css"; + case "json": + case "jsonl": + return "json"; + case "md": + case "mdx": + return "markdown"; + case "ipynb": + return "notebook"; + default: + return "default"; + } +} + +function FileReferenceIcon({ kind }: { kind: FileReferenceKind }) { + if (kind === "react") { + return ( + + + + + + + ); + } + if (kind === "default") { + return ( + + + + + ); + } + const label = fileKindLabel(kind); + return ( + + {label} + + ); +} + +function fileKindLabel(kind: FileReferenceKind): string { + switch (kind) { + case "css": + return "#"; + case "html": + return "H"; + case "json": + return "{}"; + case "markdown": + return "M"; + case "notebook": + return "N"; + case "python": + return "PY"; + case "typescript": + return "TS"; + default: + return ""; + } +} diff --git a/webui/src/components/MarkdownText.tsx b/webui/src/components/MarkdownText.tsx index 11115896..076ad55d 100644 --- a/webui/src/components/MarkdownText.tsx +++ b/webui/src/components/MarkdownText.tsx @@ -1,15 +1,46 @@ -import { Suspense, lazy } from "react"; +import { + Suspense, + lazy, + memo, + startTransition, + useCallback, + useEffect, + useLayoutEffect, + useRef, + useState, +} from "react"; import { cn } from "@/lib/utils"; interface MarkdownTextProps { children: string; className?: string; + streaming?: boolean; } const loadMarkdownRenderer = () => import("@/components/MarkdownTextRenderer"); const LazyMarkdownRenderer = lazy(loadMarkdownRenderer); +const MemoizedMarkdownRenderer = memo(function MemoizedMarkdownRenderer({ + source, + className, + highlightCode, +}: { + source: string; + className?: string; + highlightCode: boolean; +}) { + return ( + + {source} + + ); +}); + +const SHORT_STREAM_COMMIT_MS = 80; +const MEDIUM_STREAM_COMMIT_MS = 140; +const LONG_STREAM_COMMIT_MS = 220; + export function preloadMarkdownText(): void { void loadMarkdownRenderer(); } @@ -19,7 +50,18 @@ export function preloadMarkdownText(): void { * ``remark-math`` / ``rehype-katex``, and fenced code blocks delegated to * ``CodeBlock`` for copy-to-clipboard and syntax highlighting. */ -export function MarkdownText({ children, className }: MarkdownTextProps) { +export function MarkdownText({ + children, + className, + streaming = false, +}: MarkdownTextProps) { + const renderedSource = useStreamingMarkdownSource(children, streaming); + const highlightCode = !streaming && renderedSource === children; + + useEffect(() => { + if (streaming) preloadMarkdownText(); + }, [streaming]); + return ( - {children} + {renderedSource} } > - {children} + ); } + +function useStreamingMarkdownSource(source: string, streaming: boolean): string { + const [renderedSource, setRenderedSource] = useState(source); + const latestSourceRef = useRef(source); + const renderedSourceRef = useRef(source); + const timerRef = useRef(null); + + const clearPendingCommit = useCallback(() => { + if (timerRef.current !== null) { + window.clearTimeout(timerRef.current); + timerRef.current = null; + } + }, []); + + const commitSource = useCallback((next: string, urgent: boolean) => { + if (renderedSourceRef.current === next) return; + renderedSourceRef.current = next; + if (urgent) { + setRenderedSource(next); + return; + } + startTransition(() => setRenderedSource(next)); + }, []); + + const scheduleCommit = useCallback(() => { + if (timerRef.current !== null) return; + timerRef.current = window.setTimeout(() => { + timerRef.current = null; + commitSource(latestSourceRef.current, false); + }, streamingCommitDelay(latestSourceRef.current.length)); + }, [commitSource]); + + latestSourceRef.current = source; + + useLayoutEffect(() => { + latestSourceRef.current = source; + if (!streaming) { + clearPendingCommit(); + commitSource(source, true); + } + }, [clearPendingCommit, commitSource, source, streaming]); + + useEffect(() => { + latestSourceRef.current = source; + if (!streaming) return; + scheduleCommit(); + }, [scheduleCommit, source, streaming]); + + useEffect(() => clearPendingCommit, [clearPendingCommit]); + + return renderedSource; +} + +function streamingCommitDelay(length: number): number { + if (length > 24_000) return LONG_STREAM_COMMIT_MS; + if (length > 8_000) return MEDIUM_STREAM_COMMIT_MS; + return SHORT_STREAM_COMMIT_MS; +} diff --git a/webui/src/components/MarkdownTextRenderer.tsx b/webui/src/components/MarkdownTextRenderer.tsx index 17a7dc53..aa757ff0 100644 --- a/webui/src/components/MarkdownTextRenderer.tsx +++ b/webui/src/components/MarkdownTextRenderer.tsx @@ -1,10 +1,12 @@ -import { Children, isValidElement } from "react"; +import { Children, isValidElement, useMemo } from "react"; +import type { Components } from "react-markdown"; import ReactMarkdown from "react-markdown"; import rehypeKatex from "rehype-katex"; import remarkGfm from "remark-gfm"; import remarkMath from "remark-math"; import { CodeBlock } from "@/components/CodeBlock"; +import { FileReferenceChip, isLikelyFilePath } from "@/components/FileReferenceChip"; import { cn } from "@/lib/utils"; import "katex/dist/katex.min.css"; @@ -12,8 +14,12 @@ import "katex/dist/katex.min.css"; interface MarkdownTextRendererProps { children: string; className?: string; + highlightCode?: boolean; } +const remarkPlugins = [remarkGfm, remarkMath]; +const rehypePlugins = [rehypeKatex]; + /** * Heavy markdown stack (GFM, math, KaTeX, syntax highlighting) kept in a * separate chunk so the app shell can paint sooner on refresh. @@ -21,7 +27,91 @@ interface MarkdownTextRendererProps { export default function MarkdownTextRenderer({ children, className, + highlightCode = true, }: MarkdownTextRendererProps) { + const components = useMemo( + () => ({ + code({ className: cls, children: kids, ...props }) { + const match = /language-(\w+)/.exec(cls || ""); + if (match) { + const code = String(kids).replace(/\n$/, ""); + return ( + + ); + } + const raw = String(kids).replace(/\n$/, ""); + if (isLikelyFilePath(raw)) { + return ; + } + /** Plain fenced ``` blocks (no language) & wide one-liners: block monospace, not inline pill. */ + const widePlainBlock = raw.includes("\n") || raw.length > 120; + if (widePlainBlock) { + return ( + + {kids} + + ); + } + return ( + + {kids} + + ); + }, + pre({ children: markdownChildren }) { + const kids = Children.toArray(markdownChildren); + const lone = kids.length === 1 ? kids[0] : null; + /** Highlighted fences render ``CodeBlock`` (block shell); skip invalid ``
    ``. */ + if (lone != null && isValidElement(lone) && lone.type === CodeBlock) { + return <>{markdownChildren}; + } + return ( +
    +            {markdownChildren}
    +          
    + ); + }, + a({ href, children: markdownChildren, ...props }) { + return ( + + {markdownChildren} + + ); + }, + }), + [highlightCode], + ); + return (
    ; - } - const raw = String(kids).replace(/\n$/, ""); - /** Plain fenced ``` blocks (no language) & wide one-liners: block monospace, not inline pill. */ - const widePlainBlock = raw.includes("\n") || raw.length > 120; - if (widePlainBlock) { - return ( - - {kids} - - ); - } - return ( - - {kids} - - ); - }, - pre({ children: markdownChildren }) { - const kids = Children.toArray(markdownChildren); - const lone = kids.length === 1 ? kids[0] : null; - /** Highlighted fences render ``CodeBlock`` (block shell); skip invalid ``
    ``. */ - if (lone != null && isValidElement(lone) && lone.type === CodeBlock) { - return <>{markdownChildren}; - } - return ( -
    -                {markdownChildren}
    -              
    - ); - }, - a({ href, children: markdownChildren, ...props }) { - return ( - - {markdownChildren} - - ); - }, - }} + remarkPlugins={remarkPlugins} + rehypePlugins={rehypePlugins} + components={components} > {children} diff --git a/webui/src/components/MessageBubble.tsx b/webui/src/components/MessageBubble.tsx index ae15ced6..98ab0c94 100644 --- a/webui/src/components/MessageBubble.tsx +++ b/webui/src/components/MessageBubble.tsx @@ -1,6 +1,5 @@ import { useCallback, - useDeferredValue, useEffect, useRef, useState, @@ -120,7 +119,7 @@ export function MessageBubble({ ) : empty && message.isStreaming ? null : ( <> - {message.content} + {message.content} {media.length > 0 ? : null} {showAssistantFooterRow ? (
    @@ -167,10 +166,15 @@ function MessageMedia({ align: "left" | "right"; }) { if (media.length === 0) return null; - const images = media - .filter((item) => item.kind === "image") - .map(({ url, name }) => ({ url, name })); - const nonImages = media.filter((item) => item.kind !== "image"); + const images: UIImage[] = []; + const nonImages: UIMediaAttachment[] = []; + for (const item of media) { + if (item.kind === "image") { + images.push({ url: item.url, name: item.name }); + } else { + nonImages.push(item); + } + } return (
    ({ img, i })) - .filter(({ img }) => typeof img.url === "string" && img.url.length > 0); - const viewableImages = viewable.map(({ img }) => img); - const originalToViewable = new Map( - viewable.map(({ i }, v) => [i, v]), - ); + const viewableImages: UIImage[] = []; + const originalToViewable = new Map(); + for (let i = 0; i < images.length; i += 1) { + const img = images[i]; + if (typeof img.url !== "string" || img.url.length === 0) continue; + originalToViewable.set(i, viewableImages.length); + viewableImages.push(img); + } const [lightboxIndex, setLightboxIndex] = useState(null); @@ -416,7 +421,7 @@ function Dot({ delay }: { delay: string }) { ); } -/** L→R sheen overlay on label text; base copy stays solid ``text-muted-foreground``. */ +/** L→R sheen on the glyphs themselves; inactive labels stay solid muted text. */ export function StreamingLabelSheen({ children, active, @@ -426,21 +431,21 @@ export function StreamingLabelSheen({ active: boolean; className?: string; }) { + const sheenText = + typeof children === "string" || typeof children === "number" + ? String(children) + : undefined; return ( - + {children} - {active ? ( - - - - ) : null} ); } @@ -474,8 +479,6 @@ export function ReasoningBubble({ embeddedInCluster = false, }: ReasoningBubbleProps) { const { t } = useTranslation(); - const deferredText = useDeferredValue(text); - const markdownSource = streaming ? deferredText : text; const [userToggled, setUserToggled] = useState(false); const [openLocal, setOpenLocal] = useState(true); const open = userToggled ? openLocal : streaming; @@ -531,6 +534,7 @@ export function ReasoningBubble({ )} > - {markdownSource} + {text}
    )} diff --git a/webui/src/components/Sidebar.tsx b/webui/src/components/Sidebar.tsx index cf21c886..cd55475f 100644 --- a/webui/src/components/Sidebar.tsx +++ b/webui/src/components/Sidebar.tsx @@ -117,12 +117,12 @@ export function Sidebar(props: SidebarProps) { />
    -
    +
    ); } + +function shortFileName(path: string): string { + return path.split(/[\\/]/).pop() || path; +} + +function fileActivityVerb(editing: boolean, failed: boolean): string { + if (failed) return "Failed"; + return editing ? "Editing" : "Edited"; +} + +function fileActivitySummaryKey(editing: boolean, failed: boolean): string { + if (failed) return "message.fileActivityFailedOne"; + return editing ? "message.fileActivityEditingOne" : "message.fileActivityEditedOne"; +} + +function fileActivityManySummaryKey(editing: boolean, failed: boolean): string { + if (failed) return "message.fileActivityFailedMany"; + return editing ? "message.fileActivityEditingMany" : "message.fileActivityEditedMany"; +} + +function fileEditCallKey(edit: UIFileEdit): string { + return `${edit.call_id}|${edit.tool}|${edit.path}`; +} + +function collectFileEdits(messages: UIMessage[]): UIFileEdit[] { + const edits: UIFileEdit[] = []; + for (const message of messages) { + if (message.kind === "trace" && message.fileEdits?.length) { + edits.push(...message.fileEdits); + } + } + return edits; +} + +function latestFileEditEvents(edits: UIFileEdit[]): UIFileEdit[] { + const order: string[] = []; + const byKey = new Map(); + for (const edit of edits) { + const key = fileEditCallKey(edit); + if (!byKey.has(key)) order.push(key); + byKey.set(key, edit); + } + return order.map((key) => byKey.get(key)).filter(Boolean) as UIFileEdit[]; +} + +function summarizeFileEdits(edits: UIFileEdit[], active: boolean): FileEditSummary[] { + interface MutableSummary { + key: string; + path: string; + added: number; + deleted: number; + approximate: boolean; + binary: boolean; + hasSuccessfulChange: boolean; + hasActiveEditing: boolean; + hasFailed: boolean; + error?: string; + } + + const order: string[] = []; + const byPath = new Map(); + for (const edit of latestFileEditEvents(edits)) { + const key = edit.path; + let summary = byPath.get(key); + if (!summary) { + summary = { + key, + path: edit.path, + added: 0, + deleted: 0, + approximate: false, + binary: false, + hasSuccessfulChange: false, + hasActiveEditing: false, + hasFailed: false, + }; + byPath.set(key, summary); + order.push(key); + } + + if (active && edit.status === "editing") { + summary.hasActiveEditing = true; + summary.binary = summary.binary || !!edit.binary; + summary.approximate = summary.approximate || !!edit.approximate; + if (!edit.binary) { + summary.added += edit.added; + summary.deleted += edit.deleted; + } + continue; + } + + if (edit.status === "error") { + summary.hasFailed = true; + summary.error = edit.error ?? summary.error; + continue; + } + + summary.hasSuccessfulChange = true; + summary.binary = summary.binary || !!edit.binary; + summary.approximate = active && (summary.approximate || !!edit.approximate); + if (!edit.binary) { + summary.added += edit.added; + summary.deleted += edit.deleted; + } + } + + return order.map((key) => { + const summary = byPath.get(key)!; + const status: UIFileEdit["status"] = summary.hasActiveEditing + ? "editing" + : summary.hasSuccessfulChange + ? "done" + : summary.hasFailed + ? "error" + : "done"; + return { + key: summary.key, + path: summary.path, + added: summary.added, + deleted: summary.deleted, + approximate: summary.approximate, + binary: summary.binary, + status, + error: summary.error, + }; + }); +} + +function FileEditGroup({ edits }: { edits: FileEditSummary[] }) { + if (edits.length === 0) return null; + return ( +
      + {edits.map((edit) => ( + + ))} +
    + ); +} + +function FileEditRow({ edit }: { edit: FileEditSummary }) { + const { t } = useTranslation(); + const editing = edit.status === "editing"; + const failed = edit.status === "error"; + const hasCountedDiff = !failed && !edit.binary; + return ( +
  • +
    + + {failed ? ( + + + {t("message.fileEditFailed", { defaultValue: "Failed" })} + + ) : null} + {edit.approximate && !failed ? ( + + {t("message.fileEditApproximate", { defaultValue: "estimated" })} + + ) : null} +
    + {hasCountedDiff ? ( + + ) : null} +
  • + ); +} + +function DiffPair({ added, deleted }: { added: number; deleted: number }) { + return ( + + + + + + + - + + + ); +} + +function AnimatedNumber({ value }: { value: number }) { + const safeValue = Number.isFinite(value) ? Math.max(0, Math.round(value)) : 0; + const [display, setDisplay] = useState(0); + const displayRef = useRef(0); + + const setAnimatedDisplay = useCallback((next: number) => { + displayRef.current = next; + setDisplay(next); + }, []); + + useEffect(() => { + const reduceMotion = window.matchMedia?.("(prefers-reduced-motion: reduce)").matches; + if (reduceMotion) { + setAnimatedDisplay(safeValue); + return; + } + const start = displayRef.current; + const delta = safeValue - start; + if (delta === 0) { + setAnimatedDisplay(safeValue); + return; + } + const duration = 260; + const startedAt = performance.now(); + let frame = 0; + const tick = (now: number) => { + const progress = Math.min(1, (now - startedAt) / duration); + const eased = 1 - Math.pow(1 - progress, 3); + setAnimatedDisplay(Math.round(start + delta * eased)); + if (progress < 1) { + frame = window.requestAnimationFrame(tick); + return; + } + displayRef.current = safeValue; + }; + frame = window.requestAnimationFrame(tick); + return () => window.cancelAnimationFrame(frame); + }, [safeValue, setAnimatedDisplay]); + + return <>{display}; +} diff --git a/webui/src/components/thread/ThreadMessages.tsx b/webui/src/components/thread/ThreadMessages.tsx index 95f1ac42..869d282f 100644 --- a/webui/src/components/thread/ThreadMessages.tsx +++ b/webui/src/components/thread/ThreadMessages.tsx @@ -1,3 +1,6 @@ +import { useMemo } from "react"; +import { useTranslation } from "react-i18next"; + import { MessageBubble } from "@/components/MessageBubble"; import { AgentActivityCluster, @@ -9,6 +12,8 @@ interface ThreadMessagesProps { messages: UIMessage[]; /** When true, agent turn still in flight — keeps activity cluster expanded. */ isStreaming?: boolean; + hiddenMessageCount?: number; + onLoadEarlier?: () => void; } export type DisplayUnit = @@ -30,31 +35,160 @@ export function isFinalAssistantSliceBeforeNextUser( return true; } -function buildDisplayUnits(messages: UIMessage[]): DisplayUnit[] { +export function buildDisplayUnits(messages: UIMessage[]): DisplayUnit[] { const out: DisplayUnit[] = []; let i = 0; while (i < messages.length) { const m = messages[i]; if (isAgentActivityMember(m)) { const cluster: UIMessage[] = []; - while (i < messages.length && isAgentActivityMember(messages[i])) { - cluster.push(messages[i]); + let segmentId: string | undefined = m.activitySegmentId; + let clusterHasFileEdits = hasFileEdits(m); + while ( + i < messages.length + && isAgentActivityMember(messages[i]) + && canJoinActivityCluster(segmentId, clusterHasFileEdits, messages[i]) + ) { + const current = messages[i]; + if (!segmentId && current.activitySegmentId) { + segmentId = current.activitySegmentId; + } + clusterHasFileEdits = clusterHasFileEdits || hasFileEdits(current); + cluster.push(current); i += 1; } out.push({ type: "cluster", messages: cluster }); continue; } + const previous = out[out.length - 1]; + if ( + previous?.type === "cluster" + && assistantHasInlineReasoning(m) + && canFoldInlineReasoning(previous.messages, m) + ) { + previous.messages.push(reasoningOnlyMessageFromAnswer(m)); + out.push({ type: "single", message: stripInlineReasoning(m) }); + i += 1; + continue; + } + if (assistantHasInlineReasoning(m)) { + out.push({ type: "cluster", messages: [reasoningOnlyMessageFromAnswer(m)] }); + out.push({ type: "single", message: stripInlineReasoning(m) }); + i += 1; + continue; + } out.push({ type: "single", message: m }); i += 1; } return out; } -export function ThreadMessages({ messages, isStreaming = false }: ThreadMessagesProps) { - const units = buildDisplayUnits(messages); +function clusterSegmentId(messages: UIMessage[]): string | undefined { + return messages.find((message) => message.activitySegmentId)?.activitySegmentId; +} + +function hasFileEdits(message: UIMessage): boolean { + return !!message.fileEdits?.length; +} + +function clusterHasFileEdits(messages: UIMessage[]): boolean { + return messages.some(hasFileEdits); +} + +function canJoinActivityCluster( + clusterSegmentId: string | undefined, + clusterIncludesFileEdits: boolean, + message: UIMessage, +): boolean { + const messageHasFileEdits = hasFileEdits(message); + if (!clusterIncludesFileEdits && !messageHasFileEdits) return true; + if (!clusterSegmentId || !message.activitySegmentId) return true; + return clusterSegmentId === message.activitySegmentId; +} + +function canFoldInlineReasoning(cluster: UIMessage[], message: UIMessage): boolean { + if (!clusterHasFileEdits(cluster) && !hasFileEdits(message)) return true; + const segmentId = clusterSegmentId(cluster); + if (!segmentId || !message.activitySegmentId) return true; + return segmentId === message.activitySegmentId; +} + +function assistantHasInlineReasoning(message: UIMessage): boolean { + return ( + message.role === "assistant" + && message.kind !== "trace" + && message.content.trim().length > 0 + && (!!message.reasoning?.trim() || !!message.reasoningStreaming) + ); +} + +function reasoningOnlyMessageFromAnswer(message: UIMessage): UIMessage { + return { + id: `${message.id}-reasoning`, + role: "assistant", + content: "", + createdAt: message.createdAt, + reasoning: message.reasoning, + reasoningStreaming: message.reasoningStreaming, + isStreaming: message.reasoningStreaming, + activitySegmentId: message.activitySegmentId, + }; +} + +function stripInlineReasoning(message: UIMessage): UIMessage { + const next = { ...message }; + delete next.reasoning; + delete next.reasoningStreaming; + return next; +} + +export function assistantCopyFlags(units: DisplayUnit[]): boolean[] { + const flags = new Array(units.length).fill(true); + let hasLaterUnitBeforeUser = false; + for (let i = units.length - 1; i >= 0; i -= 1) { + const unit = units[i]; + if (unit.type === "single" && unit.message.role === "user") { + hasLaterUnitBeforeUser = false; + continue; + } + if (unit.type === "single" && unit.message.role === "assistant") { + flags[i] = !hasLaterUnitBeforeUser; + } + hasLaterUnitBeforeUser = true; + } + return flags; +} + +export function ThreadMessages({ + messages, + isStreaming = false, + hiddenMessageCount = 0, + onLoadEarlier, +}: ThreadMessagesProps) { + const { t } = useTranslation(); + const units = useMemo(() => buildDisplayUnits(messages), [messages]); + const copyFlags = useMemo(() => assistantCopyFlags(units), [units]); + const liveActivityClusterIndex = useMemo( + () => isStreaming ? currentActivityClusterIndex(units) : -1, + [isStreaming, units], + ); return (
    + {hiddenMessageCount > 0 && onLoadEarlier ? ( +
    + +
    + ) : null} {units.map((unit, index) => { const prev = units[index - 1]; const marginTop = @@ -72,7 +206,7 @@ export function ThreadMessages({ messages, isStreaming = false }: ThreadMessages {unit.type === "cluster" ? ( ) : ( @@ -80,7 +214,7 @@ export function ThreadMessages({ messages, isStreaming = false }: ThreadMessages message={unit.message} showAssistantCopyAction={ unit.message.role === "assistant" - ? isFinalAssistantSliceBeforeNextUser(units, index) + ? copyFlags[index] : true } /> @@ -92,6 +226,11 @@ export function ThreadMessages({ messages, isStreaming = false }: ThreadMessages ); } +function currentActivityClusterIndex(units: DisplayUnit[]): number { + const last = units.length - 1; + return units[last]?.type === "cluster" ? last : -1; +} + function unitKey(unit: DisplayUnit, index: number): string { if (unit.type === "cluster") { const anchor = unit.messages[0]?.id; diff --git a/webui/src/components/thread/ThreadShell.tsx b/webui/src/components/thread/ThreadShell.tsx index 309f206c..5711b6ce 100644 --- a/webui/src/components/thread/ThreadShell.tsx +++ b/webui/src/components/thread/ThreadShell.tsx @@ -167,8 +167,9 @@ export function ThreadShell({ useEffect(() => { if (!chatId) return; - return client.onSessionUpdate((updatedChatId) => { + return client.onSessionUpdate((updatedChatId, scope) => { if (updatedChatId !== chatId) return; + if (scope === "metadata") return; pendingCanonicalHydrateRef.current.add(chatId); refreshHistory(); }); @@ -389,6 +390,7 @@ export function ThreadShell({ composer={composer} scrollToBottomSignal={scrollToBottomSignal} conversationKey={historyKey} + showScrollToBottomButton={!!session} /> ); diff --git a/webui/src/components/thread/ThreadViewport.tsx b/webui/src/components/thread/ThreadViewport.tsx index 38b64340..3f84da68 100644 --- a/webui/src/components/thread/ThreadViewport.tsx +++ b/webui/src/components/thread/ThreadViewport.tsx @@ -1,8 +1,17 @@ -import { type ReactNode, useCallback, useEffect, useLayoutEffect, useRef, useState } from "react"; +import { + type ReactNode, + useCallback, + useEffect, + useLayoutEffect, + useMemo, + useRef, + useState, +} from "react"; import { ArrowDown } from "lucide-react"; import { useTranslation } from "react-i18next"; import { ThreadMessages } from "@/components/thread/ThreadMessages"; +import { isAgentActivityMember } from "@/components/thread/AgentActivityCluster"; import { Button } from "@/components/ui/button"; import { cn } from "@/lib/utils"; import type { UIMessage } from "@/lib/types"; @@ -14,9 +23,27 @@ interface ThreadViewportProps { emptyState?: ReactNode; scrollToBottomSignal?: number; conversationKey?: string | null; + showScrollToBottomButton?: boolean; } const NEAR_BOTTOM_PX = 48; +const DEFAULT_SCROLL_BUTTON_BOTTOM_PX = 192; +const SCROLL_BUTTON_COMPOSER_GAP_PX = 16; +export const INITIAL_HISTORY_WINDOW = 160; +export const HISTORY_WINDOW_INCREMENT = 120; + +export function windowMessages(messages: UIMessage[], visibleCount: number): UIMessage[] { + if (messages.length <= visibleCount) return messages; + let start = Math.max(0, messages.length - visibleCount); + while ( + start > 0 + && isAgentActivityMember(messages[start]) + && isAgentActivityMember(messages[start - 1]) + ) { + start -= 1; + } + return messages.slice(start); +} export function ThreadViewport({ messages, @@ -25,18 +52,33 @@ export function ThreadViewport({ emptyState, scrollToBottomSignal = 0, conversationKey = null, + showScrollToBottomButton = true, }: ThreadViewportProps) { const { t } = useTranslation(); const scrollRef = useRef(null); const contentRef = useRef(null); + const composerDockRef = useRef(null); const bottomRef = useRef(null); const lastConversationKeyRef = useRef(conversationKey); const pendingConversationScrollRef = useRef(true); const scrollFrameIdsRef = useRef([]); + const restoreScrollAfterPrependRef = + useRef<{ height: number; top: number } | null>(null); /** User scrolled away from the bottom; do not auto-yank until they return or we reset (new chat / send). */ const userReadingHistoryRef = useRef(false); const [atBottom, setAtBottom] = useState(true); + const [composerDockHeight, setComposerDockHeight] = useState(0); + const [visibleMessageCount, setVisibleMessageCount] = + useState(INITIAL_HISTORY_WINDOW); const hasMessages = messages.length > 0; + const visibleMessages = useMemo( + () => windowMessages(messages, visibleMessageCount), + [messages, visibleMessageCount], + ); + const hiddenMessageCount = messages.length - visibleMessages.length; + const scrollButtonBottom = composerDockHeight > 0 + ? composerDockHeight + SCROLL_BUTTON_COMPOSER_GAP_PX + : DEFAULT_SCROLL_BUTTON_BOTTOM_PX; const cancelScheduledBottomScroll = useCallback(() => { for (const id of scrollFrameIdsRef.current) { @@ -77,6 +119,30 @@ export function ThreadViewport({ [cancelScheduledBottomScroll, scrollToBottomNow], ); + const loadEarlierMessages = useCallback(() => { + const el = scrollRef.current; + if (el) { + restoreScrollAfterPrependRef.current = { + height: el.scrollHeight, + top: el.scrollTop, + }; + } + userReadingHistoryRef.current = true; + setAtBottom(false); + setVisibleMessageCount((count) => + Math.min(messages.length, count + HISTORY_WINDOW_INCREMENT), + ); + }, [messages.length]); + + const measureComposerDock = useCallback(() => { + const el = composerDockRef.current; + if (!el) return; + const height = el.getBoundingClientRect().height || el.offsetHeight; + setComposerDockHeight((current) => + Math.abs(current - height) < 1 ? current : height, + ); + }, []); + useEffect(() => { if (!atBottom) return; // Instant jump: CSS scroll-smooth + behavior "auto" still animates in some @@ -96,8 +162,19 @@ export function ThreadViewport({ pendingConversationScrollRef.current = true; userReadingHistoryRef.current = false; setAtBottom(true); + setVisibleMessageCount(INITIAL_HISTORY_WINDOW); }, [conversationKey]); + useLayoutEffect(() => { + const pending = restoreScrollAfterPrependRef.current; + if (!pending) return; + const el = scrollRef.current; + restoreScrollAfterPrependRef.current = null; + if (!el) return; + const delta = el.scrollHeight - pending.height; + el.scrollTop = pending.top + delta; + }, [visibleMessages.length]); + useLayoutEffect(() => { if (!pendingConversationScrollRef.current) return; if (!conversationKey) { @@ -110,6 +187,10 @@ export function ThreadViewport({ pendingConversationScrollRef.current = false; }, [conversationKey, hasMessages, messages, scrollToBottom]); + useLayoutEffect(() => { + measureComposerDock(); + }, [composer, hasMessages, measureComposerDock]); + useEffect(() => cancelScheduledBottomScroll, [cancelScheduledBottomScroll]); useEffect(() => { @@ -123,6 +204,14 @@ export function ThreadViewport({ return () => observer.disconnect(); }, [hasMessages, scrollToBottom]); + useEffect(() => { + const target = composerDockRef.current; + if (!target || typeof ResizeObserver === "undefined") return; + const observer = new ResizeObserver(() => measureComposerDock()); + observer.observe(target); + return () => observer.disconnect(); + }, [hasMessages, measureComposerDock]); + useEffect(() => { const el = scrollRef.current; if (!el) return; @@ -155,11 +244,20 @@ export function ThreadViewport({
    - +
    -
    +
    {composer}
    @@ -183,17 +281,18 @@ export function ThreadViewport({ className="pointer-events-none absolute inset-x-0 top-0 h-6 bg-gradient-to-b from-background to-transparent" /> - {!atBottom && ( + {showScrollToBottomButton && !atBottom && (
    } + />, + ); + const scroller = container.firstElementChild?.firstElementChild as HTMLElement; + Object.defineProperties(scroller, { + scrollHeight: { configurable: true, value: 2400 }, + clientHeight: { configurable: true, value: 600 }, + scrollTop: { configurable: true, value: 0 }, + }); + + act(() => { + scroller.dispatchEvent(new Event("scroll")); + }); + + const button = screen.getByRole("button", { name: "Scroll to bottom" }); + expect(button).toHaveStyle({ bottom: "192px" }); + + const composerDock = screen.getByTestId("thread-composer-dock"); + composerDock.getBoundingClientRect = () => + ({ + height: 240, + width: 800, + top: 0, + right: 800, + bottom: 240, + left: 0, + x: 0, + y: 0, + toJSON: () => ({}), + }) as DOMRect; + + const composerObserver = resizeObservers.find( + (observer) => observer.element === composerDock, + ); + expect(composerObserver).toBeDefined(); + + act(() => { + composerObserver!.callback([], composerObserver as unknown as ResizeObserver); + }); + + expect(button).toHaveStyle({ bottom: "256px" }); + } finally { + vi.stubGlobal("ResizeObserver", originalResizeObserver); + } + }); + + it("hides the scroll-to-bottom button when disabled for the welcome view", () => { + const { container } = render( + composer
    } + emptyState={
    welcome
    } + showScrollToBottomButton={false} + />, + ); + const scroller = container.firstElementChild?.firstElementChild as HTMLElement; + Object.defineProperties(scroller, { + scrollHeight: { configurable: true, value: 2400 }, + clientHeight: { configurable: true, value: 600 }, + scrollTop: { configurable: true, value: 0 }, + }); + + act(() => { + scroller.dispatchEvent(new Event("scroll")); + }); + + expect(screen.queryByRole("button", { name: "Scroll to bottom" })).not.toBeInTheDocument(); + }); + + it("renders only the tail window for long history by default", () => { + const longMessages = makeLongMessages(300); + + render( + } + />, + ); + + expect(screen.queryByText("message 139")).not.toBeInTheDocument(); + expect(screen.getByText("message 140")).toBeInTheDocument(); + expect(screen.getByText("message 299")).toBeInTheDocument(); + expect(screen.getByRole("button", { name: "Load earlier messages" })).toBeInTheDocument(); + }); + + it("loads earlier history in fixed increments without rendering the whole transcript", () => { + const longMessages = makeLongMessages(300); + + render( + } + />, + ); + + fireEvent.click(screen.getByRole("button", { name: "Load earlier messages" })); + + const firstVisible = + 300 - INITIAL_HISTORY_WINDOW - HISTORY_WINDOW_INCREMENT; + + expect( + screen.queryByText(`message ${firstVisible - 1}`), + ).not.toBeInTheDocument(); + expect(screen.getByText(`message ${firstVisible}`)).toBeInTheDocument(); + expect(screen.getByText("message 299")).toBeInTheDocument(); + }); + + it("expands the window start to avoid cutting an agent activity cluster", () => { + const clustered = makeLongMessages(200); + clustered.splice( + 38, + 3, + { + id: "r0", + role: "assistant", + content: "", + reasoning: "first reasoning", + createdAt: 38, + }, + { + id: "t0", + role: "tool", + kind: "trace", + content: "tool()", + traces: ["tool()"], + createdAt: 39, + }, + { + id: "r1", + role: "assistant", + content: "", + reasoning: "second reasoning", + createdAt: 40, + }, + ); + + const visible = windowMessages(clustered, INITIAL_HISTORY_WINDOW); + + expect(visible[0].id).toBe("r0"); + expect(visible).toHaveLength(INITIAL_HISTORY_WINDOW + 2); + }); + it("resets to the bottom when opening a different conversation", async () => { const scrollIntoView = vi.fn(); const originalScrollIntoView = HTMLElement.prototype.scrollIntoView; diff --git a/webui/src/tests/useDeferredTitleRefresh.test.tsx b/webui/src/tests/useDeferredTitleRefresh.test.tsx new file mode 100644 index 00000000..a823e504 --- /dev/null +++ b/webui/src/tests/useDeferredTitleRefresh.test.tsx @@ -0,0 +1,110 @@ +import { act, renderHook } from "@testing-library/react"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; + +import { useDeferredTitleRefresh } from "@/hooks/useDeferredTitleRefresh"; +import type { ChatSummary } from "@/lib/types"; + +function session(overrides: Partial = {}): ChatSummary { + return { + key: "websocket:chat-a", + channel: "websocket", + chatId: "chat-a", + createdAt: null, + updatedAt: null, + title: "", + preview: "First user message", + ...overrides, + }; +} + +describe("useDeferredTitleRefresh", () => { + beforeEach(() => { + vi.useFakeTimers(); + }); + + afterEach(() => { + vi.useRealTimers(); + }); + + it("retries refreshing untitled sessions after turn_end", () => { + const refresh = vi.fn().mockResolvedValue(undefined); + const { result } = renderHook(() => + useDeferredTitleRefresh(session(), refresh, [100, 300]), + ); + + act(() => { + result.current(); + }); + + expect(refresh).toHaveBeenCalledTimes(1); + + act(() => { + vi.advanceTimersByTime(100); + }); + expect(refresh).toHaveBeenCalledTimes(2); + + act(() => { + vi.advanceTimersByTime(200); + }); + expect(refresh).toHaveBeenCalledTimes(3); + }); + + it("stops pending retries once a generated title arrives", () => { + const refresh = vi.fn().mockResolvedValue(undefined); + const { result, rerender } = renderHook( + ({ activeSession }) => + useDeferredTitleRefresh(activeSession, refresh, [100, 300]), + { initialProps: { activeSession: session() } }, + ); + + act(() => { + result.current(); + }); + rerender({ activeSession: session({ title: "Generated title" }) }); + + act(() => { + vi.advanceTimersByTime(300); + }); + + expect(refresh).toHaveBeenCalledTimes(1); + }); + + it("does not retry when the active session already has a title", () => { + const refresh = vi.fn().mockResolvedValue(undefined); + const { result } = renderHook(() => + useDeferredTitleRefresh(session({ title: "Existing title" }), refresh, [100]), + ); + + act(() => { + result.current(); + vi.advanceTimersByTime(100); + }); + + expect(refresh).toHaveBeenCalledTimes(1); + }); + + it("clears pending retries when the active chat changes", () => { + const refresh = vi.fn().mockResolvedValue(undefined); + const { result, rerender } = renderHook( + ({ activeSession }) => + useDeferredTitleRefresh(activeSession, refresh, [100]), + { initialProps: { activeSession: session() } }, + ); + + act(() => { + result.current(); + }); + rerender({ + activeSession: session({ + key: "websocket:chat-b", + chatId: "chat-b", + }), + }); + + act(() => { + vi.advanceTimersByTime(100); + }); + + expect(refresh).toHaveBeenCalledTimes(1); + }); +}); diff --git a/webui/src/tests/useNanobotStream.test.tsx b/webui/src/tests/useNanobotStream.test.tsx index 57ecccd9..925102da 100644 --- a/webui/src/tests/useNanobotStream.test.tsx +++ b/webui/src/tests/useNanobotStream.test.tsx @@ -83,7 +83,112 @@ function wrap(client: ReturnType["client"]) { }; } +async function flushStreamFrame() { + await act(async () => { + await new Promise((resolve) => { + requestAnimationFrame(() => resolve()); + }); + }); +} + describe("useNanobotStream", () => { + it("batches answer deltas into one animation-frame update", async () => { + const fake = fakeClient(); + const requestFrame = vi.spyOn(window, "requestAnimationFrame"); + const { result } = renderHook(() => useNanobotStream("chat-batch", EMPTY_MESSAGES), { + wrapper: wrap(fake.client), + }); + + act(() => { + fake.emit("chat-batch", { + event: "delta", + chat_id: "chat-batch", + text: "Hello", + }); + fake.emit("chat-batch", { + event: "delta", + chat_id: "chat-batch", + text: " world", + }); + }); + + expect(requestFrame).toHaveBeenCalledTimes(1); + expect(result.current.messages).toHaveLength(0); + + await flushStreamFrame(); + + expect(result.current.messages).toHaveLength(1); + expect(result.current.messages[0]).toMatchObject({ + role: "assistant", + content: "Hello world", + isStreaming: true, + }); + requestFrame.mockRestore(); + }); + + it("flushes pending delta text before turn_end finalizes the turn", () => { + const fake = fakeClient(); + const { result } = renderHook(() => useNanobotStream("chat-flush", EMPTY_MESSAGES), { + wrapper: wrap(fake.client), + }); + + act(() => { + fake.emit("chat-flush", { + event: "delta", + chat_id: "chat-flush", + text: "final chunk", + }); + fake.emit("chat-flush", { + event: "turn_end", + chat_id: "chat-flush", + }); + }); + + expect(result.current.messages).toHaveLength(1); + expect(result.current.messages[0]).toMatchObject({ + role: "assistant", + content: "final chunk", + isStreaming: false, + }); + expect(result.current.isStreaming).toBe(false); + }); + + it("drops pending stream work when switching chats", async () => { + const fake = fakeClient(); + const { result, rerender } = renderHook( + ({ chatId }: { chatId: string }) => useNanobotStream(chatId, EMPTY_MESSAGES), + { + wrapper: wrap(fake.client), + initialProps: { chatId: "chat-old" }, + }, + ); + + act(() => { + fake.emit("chat-old", { + event: "delta", + chat_id: "chat-old", + text: "stale", + }); + }); + + rerender({ chatId: "chat-new" }); + + act(() => { + fake.emit("chat-new", { + event: "delta", + chat_id: "chat-new", + text: "fresh", + }); + }); + await flushStreamFrame(); + + expect(result.current.messages).toHaveLength(1); + expect(result.current.messages[0]).toMatchObject({ + role: "assistant", + content: "fresh", + }); + }); + it("starts in streaming mode when history shows pending tool calls", () => { const fake = fakeClient(); const initialMessages = [{ @@ -203,7 +308,174 @@ describe("useNanobotStream", () => { ); }); - it("accumulates reasoning_delta chunks on a placeholder until reasoning_end", () => { + it("renders live file_edit events as their own activity trace", () => { + const fake = fakeClient(); + const { result } = renderHook(() => useNanobotStream("chat-file-edit", EMPTY_MESSAGES), { + wrapper: wrap(fake.client), + }); + + act(() => { + fake.emit("chat-file-edit", { + event: "message", + chat_id: "chat-file-edit", + text: 'write_file({"path":"foo.txt"})', + kind: "tool_hint", + }); + fake.emit("chat-file-edit", { + event: "file_edit", + chat_id: "chat-file-edit", + edits: [{ + call_id: "call-write", + tool: "write_file", + path: "foo.txt", + phase: "start", + added: 1, + deleted: 0, + approximate: true, + status: "editing", + }], + }); + fake.emit("chat-file-edit", { + event: "file_edit", + chat_id: "chat-file-edit", + edits: [{ + call_id: "call-write", + tool: "write_file", + path: "foo.txt", + phase: "end", + added: 3, + deleted: 1, + approximate: false, + status: "done", + }], + }); + }); + + expect(result.current.messages).toHaveLength(2); + expect(result.current.messages[0]).toMatchObject({ + role: "tool", + kind: "trace", + traces: ['write_file({"path":"foo.txt"})'], + }); + expect(result.current.messages[1]).toMatchObject({ + role: "tool", + kind: "trace", + fileEdits: [{ + call_id: "call-write", + status: "done", + added: 3, + deleted: 1, + approximate: false, + }], + }); + expect(result.current.messages[1].activitySegmentId).toBeTruthy(); + expect(result.current.messages[1].activitySegmentId).not.toBe( + result.current.messages[0].activitySegmentId, + ); + }); + + it("starts a new assistant bubble for deltas after stream_end and activity", async () => { + const fake = fakeClient(); + const { result } = renderHook(() => useNanobotStream("chat-stream-segments", EMPTY_MESSAGES), { + wrapper: wrap(fake.client), + }); + + act(() => { + fake.emit("chat-stream-segments", { + event: "delta", + chat_id: "chat-stream-segments", + text: "I created the files.", + }); + fake.emit("chat-stream-segments", { + event: "stream_end", + chat_id: "chat-stream-segments", + }); + fake.emit("chat-stream-segments", { + event: "message", + chat_id: "chat-stream-segments", + text: 'write_file({"path":"minecraft-fps/options.txt"})', + kind: "tool_hint", + }); + fake.emit("chat-stream-segments", { + event: "delta", + chat_id: "chat-stream-segments", + text: "Now I will summarize the edits.", + }); + }); + + await flushStreamFrame(); + + expect(result.current.messages).toHaveLength(3); + expect(result.current.messages[0]).toMatchObject({ + role: "assistant", + content: "I created the files.", + }); + expect(result.current.messages[1]).toMatchObject({ + role: "tool", + kind: "trace", + traces: ['write_file({"path":"minecraft-fps/options.txt"})'], + }); + expect(result.current.messages[2]).toMatchObject({ + role: "assistant", + content: "Now I will summarize the edits.", + }); + }); + + it("opens a new activity segment for reasoning after file edit activity", async () => { + const fake = fakeClient(); + const { result } = renderHook(() => useNanobotStream("chat-file-segments", EMPTY_MESSAGES), { + wrapper: wrap(fake.client), + }); + + act(() => { + fake.emit("chat-file-segments", { + event: "reasoning_delta", + chat_id: "chat-file-segments", + text: "Plan.", + }); + fake.emit("chat-file-segments", { + event: "reasoning_end", + chat_id: "chat-file-segments", + }); + fake.emit("chat-file-segments", { + event: "message", + chat_id: "chat-file-segments", + text: 'edit_file({"path":"foo.txt"})', + kind: "tool_hint", + }); + fake.emit("chat-file-segments", { + event: "file_edit", + chat_id: "chat-file-segments", + edits: [{ + call_id: "call-edit", + tool: "edit_file", + path: "foo.txt", + phase: "start", + added: 1, + deleted: 1, + approximate: true, + status: "editing", + }], + }); + fake.emit("chat-file-segments", { + event: "reasoning_delta", + chat_id: "chat-file-segments", + text: "Review result.", + }); + }); + + await flushStreamFrame(); + + expect(result.current.messages).toHaveLength(4); + const firstSegment = result.current.messages[0].activitySegmentId; + expect(firstSegment).toBeTruthy(); + expect(result.current.messages[1].activitySegmentId).toBe(firstSegment); + expect(result.current.messages[2].activitySegmentId).toBeTruthy(); + expect(result.current.messages[2].activitySegmentId).not.toBe(firstSegment); + expect(result.current.messages[3].activitySegmentId).toBe(firstSegment); + }); + + it("accumulates reasoning_delta chunks on a placeholder until reasoning_end", async () => { const fake = fakeClient(); const { result } = renderHook(() => useNanobotStream("chat-r", EMPTY_MESSAGES), { wrapper: wrap(fake.client), @@ -222,6 +494,8 @@ describe("useNanobotStream", () => { }); }); + await flushStreamFrame(); + expect(result.current.messages).toHaveLength(1); expect(result.current.messages[0].role).toBe("assistant"); expect(result.current.messages[0].reasoning).toBe("Let me think step by step."); @@ -328,7 +602,7 @@ describe("useNanobotStream", () => { expect(result.current.messages[0].reasoningStreaming).toBe(false); }); - it("does not attach a new turn's reasoning across the latest user boundary", () => { + it("does not attach a new turn's reasoning across the latest user boundary", async () => { const fake = fakeClient(); const initialMessages = [ { @@ -358,6 +632,8 @@ describe("useNanobotStream", () => { }); }); + await flushStreamFrame(); + expect(result.current.messages).toHaveLength(3); expect(result.current.messages[0].reasoning).toBe("Previous thought."); expect(result.current.messages[2].role).toBe("assistant"); @@ -366,7 +642,7 @@ describe("useNanobotStream", () => { expect(result.current.messages[2].reasoningStreaming).toBe(true); }); - it("does not attach reasoning across a tool trace boundary", () => { + it("does not attach reasoning across a tool trace boundary", async () => { const fake = fakeClient(); const { result } = renderHook(() => useNanobotStream("chat-r7", EMPTY_MESSAGES), { wrapper: wrap(fake.client), @@ -392,6 +668,8 @@ describe("useNanobotStream", () => { }); }); + await flushStreamFrame(); + expect(result.current.messages).toHaveLength(3); expect(result.current.messages.map((m) => m.kind ?? "message")).toEqual([ "message", @@ -651,7 +929,7 @@ describe("useNanobotStream", () => { expect(result.current.messages[0].content).toBe("long task"); }); - it("keeps streaming alive across stream_end and completes on turn_end", () => { + it("keeps streaming alive across stream_end and completes on turn_end", async () => { const fake = fakeClient(); const onTurnEnd = vi.fn(); const { result } = renderHook(() => useNanobotStream("chat-s", EMPTY_MESSAGES, false, onTurnEnd), { @@ -666,6 +944,8 @@ describe("useNanobotStream", () => { }); }); + await flushStreamFrame(); + expect(result.current.isStreaming).toBe(true); expect(result.current.messages[0]).toMatchObject({ role: "assistant", diff --git a/webui/src/tests/useSessions.test.tsx b/webui/src/tests/useSessions.test.tsx index 9e340a66..72df813e 100644 --- a/webui/src/tests/useSessions.test.tsx +++ b/webui/src/tests/useSessions.test.tsx @@ -2,7 +2,7 @@ import { act, renderHook, waitFor } from "@testing-library/react"; import type { ReactNode } from "react"; import { beforeEach, describe, expect, it, vi } from "vitest"; -import { useSessionHistory, useSessions } from "@/hooks/useSessions"; +import { sessionTitle, useSessionHistory, useSessions } from "@/hooks/useSessions"; import * as api from "@/lib/api"; import { ClientProvider } from "@/providers/ClientProvider"; @@ -17,7 +17,7 @@ vi.mock("@/lib/api", async (importOriginal) => { }); function fakeClient() { - const sessionUpdateHandlers = new Set<(chatId: string) => void>(); + const sessionUpdateHandlers = new Set<(chatId: string, scope?: string) => void>(); return { status: "open" as const, defaultChatId: null as string | null, @@ -25,12 +25,12 @@ function fakeClient() { onError: () => () => {}, onChat: () => () => {}, getRunStartedAt: () => null, - onSessionUpdate: (handler: (chatId: string) => void) => { + onSessionUpdate: (handler: (chatId: string, scope?: string) => void) => { sessionUpdateHandlers.add(handler); return () => sessionUpdateHandlers.delete(handler); }, - emitSessionUpdate: (chatId: string) => { - for (const handler of sessionUpdateHandlers) handler(chatId); + emitSessionUpdate: (chatId: string, scope?: string) => { + for (const handler of sessionUpdateHandlers) handler(chatId, scope); }, sendMessage: vi.fn(), newChat: vi.fn(), @@ -61,6 +61,28 @@ describe("useSessions", () => { vi.mocked(api.fetchWebuiThread).mockReset(); }); + it("does not use low-information greetings as fallback session titles", () => { + expect(sessionTitle({ + key: "websocket:chat-hi", + channel: "websocket", + chatId: "chat-hi", + createdAt: "2026-04-16T10:00:00Z", + updatedAt: "2026-04-16T10:00:00Z", + title: "", + preview: "hi", + })).toBe("New chat"); + + expect(sessionTitle({ + key: "websocket:chat-work", + channel: "websocket", + chatId: "chat-work", + createdAt: "2026-04-16T10:00:00Z", + updatedAt: "2026-04-16T10:00:00Z", + title: "", + preview: "帮我优化 WebUI 性能", + })).toBe("帮我优化 WebUI 性能"); + }); + it("removes a session from the local list after delete succeeds", async () => { vi.mocked(api.listSessions).mockResolvedValue([ { diff --git a/webui/src/types/react-syntax-highlighter-subpaths.d.ts b/webui/src/types/react-syntax-highlighter-subpaths.d.ts new file mode 100644 index 00000000..57639f72 --- /dev/null +++ b/webui/src/types/react-syntax-highlighter-subpaths.d.ts @@ -0,0 +1,22 @@ +declare module "react-syntax-highlighter/dist/esm/prism-async-light" { + import * as React from "react"; + import type { SyntaxHighlighterProps } from "react-syntax-highlighter"; + + export default class SyntaxHighlighter extends React.Component { + static registerLanguage(name: string, func: unknown): void; + } +} + +declare module "react-syntax-highlighter/dist/esm/styles/prism/one-dark" { + import type * as React from "react"; + + const style: { [key: string]: React.CSSProperties }; + export default style; +} + +declare module "react-syntax-highlighter/dist/esm/styles/prism/one-light" { + import type * as React from "react"; + + const style: { [key: string]: React.CSSProperties }; + export default style; +} diff --git a/webui/vite.config.ts b/webui/vite.config.ts index 7a2c9edb..fb5dfe37 100644 --- a/webui/vite.config.ts +++ b/webui/vite.config.ts @@ -25,6 +25,36 @@ export default defineConfig(({ mode }) => { outDir: path.resolve(__dirname, "../nanobot/web/dist"), emptyOutDir: true, sourcemap: false, + rollupOptions: { + output: { + manualChunks(id) { + if (id.includes("node_modules/refractor/lang/")) { + return; + } + if ( + id.includes("node_modules/react-syntax-highlighter") + || id.includes("node_modules/refractor/core") + ) { + return "syntax-highlight"; + } + if ( + id.includes("node_modules/react-markdown") + || id.includes("node_modules/remark-") + || id.includes("node_modules/rehype-") + || id.includes("node_modules/unified") + || id.includes("node_modules/mdast-") + || id.includes("node_modules/hast-") + || id.includes("node_modules/micromark") + || id.includes("node_modules/unist-") + ) { + return "markdown-vendor"; + } + if (id.includes("node_modules/katex")) { + return "katex"; + } + }, + }, + }, }, server: { host: "127.0.0.1",