diff --git a/src/bub/channels/cli/__init__.py b/src/bub/channels/cli/__init__.py index a56c0d82..2416526f 100644 --- a/src/bub/channels/cli/__init__.py +++ b/src/bub/channels/cli/__init__.py @@ -9,13 +9,15 @@ from loguru import logger from prompt_toolkit import PromptSession -from prompt_toolkit.application import run_in_terminal from prompt_toolkit.completion import WordCompleter -from prompt_toolkit.formatted_text import FormattedText +from prompt_toolkit.data_structures import Point +from prompt_toolkit.filters import Condition +from prompt_toolkit.formatted_text import ANSI, AnyFormattedText, FormattedText from prompt_toolkit.history import FileHistory from prompt_toolkit.key_binding import KeyBindings +from prompt_toolkit.layout import ConditionalContainer, FormattedTextControl, HSplit, Window +from prompt_toolkit.layout.dimension import Dimension from prompt_toolkit.patch_stdout import patch_stdout -from prompt_toolkit.utils import get_cwidth from rich import get_console from rich.spinner import SPINNERS from rich.text import Text @@ -26,7 +28,13 @@ from bub.builtin.tape import TapeInfo from bub.channels.admission import AdmitDecision, TurnSnapshot from bub.channels.base import Interface +from bub.channels.cli.ansi_bridge import render_to_ansi from bub.channels.cli.renderer import CliRenderer +from bub.channels.cli.terminal_output import ( + TerminalPresenter, + create_synchronized_output, +) +from bub.channels.cli.writers import MarkdownWriter from bub.channels.contracts import MessageHandler from bub.channels.message import ChannelMessage from bub.envelope import Envelope, field_of @@ -38,15 +46,25 @@ class _StreamPrinter: - def __init__(self, *, console, print_head: Callable[[], None], expand_thinking: bool) -> None: + def __init__( + self, + *, + console, + print_head: Callable[[], None], + expand_thinking: bool, + presenter: TerminalPresenter, + writer: MarkdownWriter | None = None, + invalidate: Callable[[], None] | None = None, + ) -> None: self._console = console self._print_head = print_head self._expand_thinking = expand_thinking + self._presenter = presenter self._reasoning_chars = 0 self._reasoning_streaming = False - self._current_text_line = "" - self._rendered_text_line: str | None = None - self._live_text_rows = 0 + self._writer = writer or MarkdownWriter() + self._invalidate = invalidate or (lambda: None) + self._ansi_cache: tuple[int, str] | None = None self.head_printed = False async def render(self, event: StreamEvent) -> bool: @@ -59,7 +77,7 @@ async def render(self, event: StreamEvent) -> bool: elif event.kind == "tool_call": await self._print_stream_boundary() elif event.kind == "final": - await self._print_end() + await self.finish() return True async def _record_reasoning(self, reasoning: str) -> None: @@ -67,6 +85,7 @@ async def _record_reasoning(self, reasoning: str) -> None: if self._reasoning_chars == 0: await self._ensure_head() self._reasoning_chars += len(reasoning) + self._invalidate() return await self._ensure_head() @@ -84,27 +103,23 @@ async def _print_content(self, content: str) -> bool: await self._write_text(content) return True - async def _print_end(self) -> None: + async def finish(self) -> None: + await self._close_reasoning_stream() if self._reasoning_chars: await self._ensure_head() await self._flush_reasoning() - if self._current_text_line: - await self._commit_text_line() - elif self.head_printed and not self._live_text_rows: - await self._print("") + if self._writer.has_content(): + await self._flush_text() async def _print_stream_boundary(self) -> None: - await self._close_reasoning_stream() - await self._flush_reasoning() - if self._current_text_line or self._live_text_rows: - await self._commit_text_line() + await self.finish() if self.head_printed: await self._print("") async def _ensure_head(self) -> None: if self.head_printed: return - await run_in_terminal(self._print_head, render_cli_done=False) + await self._presenter.write(self._print_head) self.head_printed = True async def _close_reasoning_stream(self) -> None: @@ -121,80 +136,59 @@ async def _flush_reasoning(self) -> None: self._reasoning_chars = 0 async def _write_text(self, text: str) -> None: - parts = text.split("\n") - for index, part in enumerate(parts): - self._current_text_line += part - if index < len(parts) - 1: - await self._commit_text_line() - - if self._current_text_line: - await self._render_live_text_line() - - async def _commit_text_line(self) -> None: - if self._live_text_rows and self._rendered_text_line == self._current_text_line: - self._current_text_line = "" - self._rendered_text_line = None - self._live_text_rows = 0 + self._writer.append(text) + self._ansi_cache = None + self._invalidate() + + def render_live_ansi(self, *, width: int) -> str: + if not self._writer.has_content(): + return "" + if self._ansi_cache is None or self._ansi_cache[0] != width: + rendered = render_to_ansi(self._writer.render_live(), width=width).rstrip("\n") + self._ansi_cache = (width, rendered) + return self._ansi_cache[1] + + def live_cursor_position(self, *, width: int) -> Point: + rendered = self.render_live_ansi(width=width) + return Point(x=0, y=max(0, len(rendered.splitlines()) - 1)) + + def has_live_content(self) -> bool: + return self._writer.has_content() + + async def _flush_text(self) -> None: + finished = self._writer.render_final() + if finished is None: return - self._live_text_rows = await self._render_text_line(self._current_text_line) - self._current_text_line = "" - self._rendered_text_line = None - self._live_text_rows = 0 - - async def commit_live_text(self) -> None: - if self._current_text_line or self._live_text_rows: - await self._commit_text_line() - - async def _render_live_text_line(self) -> None: - self._live_text_rows = await self._render_text_line(self._current_text_line) - self._rendered_text_line = self._current_text_line - - async def _render_text_line(self, text: str) -> int: - previous_rows = self._live_text_rows - rows = self._display_rows(text) - - def render() -> None: - self._rewind_live_text(previous_rows) - self._console.print(f"{text}\n", end="", highlight=False) - await run_in_terminal(render, render_cli_done=False) - return rows + def commit() -> None: + self._console.print(finished) + self._writer.clear() - def _display_rows(self, text: str) -> int: - columns = max(1, int(getattr(self._console, "width", 80) or 80)) - return max(1, (get_cwidth(text) + columns - 1) // columns) - - def _rewind_live_text(self, rows: int) -> None: - if rows <= 0: - return - output = getattr(self._console, "file", None) - if output is None: - return - output.write(f"\x1b[{rows}A\r") - for row in range(rows): - output.write("\x1b[2K") - if row < rows - 1: - output.write("\x1b[1B\r") - if rows > 1: - output.write(f"\x1b[{rows - 1}A\r") - output.flush() + await self._presenter.write(commit) + self._ansi_cache = None + self._invalidate() async def _print(self, *args: Any, **kwargs: Any) -> None: - await run_in_terminal(lambda: self._console.print(*args, **kwargs), render_cli_done=False) + await self._presenter.write(lambda: self._console.print(*args, **kwargs)) class _CliToolCallReporter: - def __init__(self, renderer: CliRenderer) -> None: + def __init__(self, renderer: CliRenderer, presenter: TerminalPresenter) -> None: self._renderer = renderer + self._presenter = presenter - def start(self, name: str, args: tuple[Any, ...], kwargs: dict[str, Any]) -> None: - self._renderer.tool_call_start(name=name, args=args, kwargs=kwargs) + async def start(self, name: str, args: tuple[Any, ...], kwargs: dict[str, Any]) -> None: + await self._presenter.write(lambda: self._renderer.tool_call_start(name=name, args=args, kwargs=kwargs)) - def success(self, name: str, result: object, elapsed_ms: float) -> None: - self._renderer.tool_call_success(name=name, result=result, elapsed_ms=elapsed_ms) + async def success(self, name: str, result: object, elapsed_ms: float) -> None: + await self._presenter.write( + lambda: self._renderer.tool_call_success(name=name, result=result, elapsed_ms=elapsed_ms) + ) - def error(self, name: str, error: BaseException, elapsed_ms: float) -> None: - self._renderer.tool_call_error(name=name, error=error, elapsed_ms=elapsed_ms) + async def error(self, name: str, error: BaseException, elapsed_ms: float) -> None: + await self._presenter.write( + lambda: self._renderer.tool_call_error(name=name, error=error, elapsed_ms=elapsed_ms) + ) class CliChannel(Interface): @@ -214,9 +208,11 @@ def __init__(self, on_receive: MessageHandler, agent: Agent) -> None: self._mode = "agent" # or "shell" self._expand_thinking = False self._llm_loop_running = False + self._generation_tick: asyncio.TimerHandle | None = None self._main_task: asyncio.Task | None = None self._stream_printer: _StreamPrinter | None = None self._renderer = CliRenderer(get_console()) + self._presenter = TerminalPresenter() self._last_tape_info: TapeInfo | None = None self._workspace = self._agent.framework.workspace self._prompt = self._build_prompt(self._workspace) @@ -242,6 +238,7 @@ async def start(self, stop_event: asyncio.Event) -> None: self._main_task = asyncio.create_task(self._main_loop()) async def stop(self) -> None: + self._stop_generation_animation() if self._main_task is not None: self._main_task.cancel() with contextlib.suppress(asyncio.CancelledError): @@ -250,23 +247,20 @@ async def stop(self) -> None: async def send(self, message: ChannelMessage) -> None: if message.kind != "error": return - self._renderer.error(message.content) + await self._presenter.write(lambda: self._renderer.error(message.content)) async def _main_loop(self) -> None: - self._renderer.welcome(model=self._agent.settings.model, workspace=str(self._workspace)) + await self._presenter.write( + lambda: self._renderer.welcome(model=self._agent.settings.model, workspace=str(self._workspace)) + ) await self._refresh_tape_info() while not self._stop_event.is_set(): try: with patch_stdout(raw=True): - raw = ( - await self._prompt.prompt_async( - self._prompt_message, - refresh_interval=_PROMPT_REFRESH_INTERVAL, - ) - ).strip() + raw = (await self._prompt.prompt_async(self._prompt_message)).strip() except KeyboardInterrupt: - self._renderer.info("Interrupted. Use ',quit' to exit.") + await self._presenter.write(lambda: self._renderer.info("Interrupted. Use ',quit' to exit.")) continue except EOFError: break @@ -277,7 +271,7 @@ async def _main_loop(self) -> None: break if raw == ",thinking": await self._echo_input(raw) - self._toggle_thinking() + await self._toggle_thinking() continue request = self._normalize_input(raw) @@ -297,7 +291,8 @@ async def _main_loop(self) -> None: self._set_llm_loop_running(False) raise - self._renderer.info("Bye.") + self._stop_generation_animation() + await self._presenter.write(lambda: self._renderer.info("Bye.")) self._stop_event.set() @contextlib.asynccontextmanager @@ -316,16 +311,32 @@ def _normalize_input(self, raw: str) -> str: return raw return f",{raw}" - def _prompt_message(self) -> FormattedText: - prompt = self._prompt_label() - if not self._llm_loop_running: - return FormattedText([("bold", prompt)]) + def _prompt_message(self) -> AnyFormattedText: + return FormattedText([("bold", self._prompt_label())]) + + def _live_output_message(self) -> AnyFormattedText: + stream_printer: _StreamPrinter | None = getattr(self, "_stream_printer", None) + if stream_printer is None: + return FormattedText([]) + return ANSI(stream_printer.render_live_ansi(width=get_console().width)) + + def _live_output_cursor(self) -> Point: + stream_printer: _StreamPrinter | None = getattr(self, "_stream_printer", None) + if stream_printer is None: + return Point(x=0, y=0) + return stream_printer.live_cursor_position(width=get_console().width) + + def _has_live_output(self) -> bool: + stream_printer: _StreamPrinter | None = getattr(self, "_stream_printer", None) + return stream_printer is not None and stream_printer.has_live_content() + + def _is_generating(self) -> bool: + return getattr(self, "_stream_printer", None) is not None or self._llm_loop_running + + def _generation_status(self) -> FormattedText: index = int(monotonic() / _PROMPT_REFRESH_INTERVAL) % len(_GENERATION_SPINNER) spinner = _GENERATION_SPINNER[index] - return FormattedText([ - ("blue", f"\n{spinner} Generating\n"), - ("bold", prompt), - ]) + return FormattedText([("blue", f"{spinner} Generating")]) def _prompt_label(self) -> str: cwd = Path.cwd().name @@ -333,10 +344,7 @@ def _prompt_label(self) -> str: return f"{cwd} {symbol} " async def _echo_input(self, raw: str, steering: bool = False) -> None: - stream_printer = getattr(self, "_stream_printer", None) - if stream_printer is not None: - await stream_printer.commit_live_text() - self._renderer.input_echo(self._prompt_label(), raw, steering=steering) + await self._presenter.write(lambda: self._renderer.input_echo(self._prompt_label(), raw, steering=steering)) async def stream_events( self, message: ChannelMessage, stream: AsyncIterable[StreamEvent] @@ -346,16 +354,23 @@ async def stream_events( console=console, print_head=lambda: self._renderer.print_head(message.kind), expand_thinking=self._expand_thinking, + presenter=self._presenter, + invalidate=self._invalidate_prompt, ) self._stream_printer = printer + self._invalidate_prompt() try: - with tool_call_reporter(_CliToolCallReporter(self._renderer)): + with tool_call_reporter(_CliToolCallReporter(self._renderer, self._presenter)): async for event in stream: if await printer.render(event): yield event finally: - if self._stream_printer is printer: - self._stream_printer = None + try: + await printer.finish() + finally: + if self._stream_printer is printer: + self._stream_printer = None + self._invalidate_prompt() def _build_prompt(self, workspace: Path) -> PromptSession[str]: kb = KeyBindings() @@ -374,14 +389,42 @@ def _tool_sort_key(tool_name: str) -> tuple[str, str]: history = FileHistory(str(history_file)) tool_names = sorted([*(f",{name}" for name in REGISTRY), ",thinking"], key=_tool_sort_key) completer = WordCompleter(tool_names, ignore_case=True, sentence=True) - return PromptSession( + prompt: PromptSession[str] = PromptSession( completer=completer, complete_while_typing=True, key_bindings=kb, history=history, bottom_toolbar=self._render_bottom_toolbar, erase_when_done=True, + output=create_synchronized_output(), + ) + self._attach_live_layout(prompt) + prompt.app.min_redraw_interval = _PROMPT_REFRESH_INTERVAL + return prompt + + def _attach_live_layout(self, prompt: PromptSession[str]) -> None: + root = prompt.layout.container + if not isinstance(root, HSplit): + raise TypeError("PromptSession root layout must be an HSplit") + live_output = Window( + FormattedTextControl( + self._live_output_message, + show_cursor=False, + get_cursor_position=self._live_output_cursor, + ), + height=Dimension(min=0, weight=1), + wrap_lines=False, + always_hide_cursor=True, + ) + generation_status = Window( + FormattedTextControl(self._generation_status), + height=1, + dont_extend_height=True, ) + root.children[0:0] = [ + ConditionalContainer(live_output, Condition(self._has_live_output)), + ConditionalContainer(generation_status, Condition(self._is_generating)), + ] def _render_bottom_toolbar(self) -> FormattedText: info = self._last_tape_info @@ -396,10 +439,10 @@ def _render_bottom_toolbar(self) -> FormattedText: ) return FormattedText([("", f"{left} {right}")]) - def _toggle_thinking(self) -> None: + async def _toggle_thinking(self) -> None: self._expand_thinking = not self._expand_thinking state = "expanded" if self._expand_thinking else "collapsed" - self._renderer.info(f"Thinking output is now {state}.") + await self._presenter.write(lambda: self._renderer.info(f"Thinking output is now {state}.")) def _invalidate_prompt(self) -> None: with contextlib.suppress(Exception): @@ -409,8 +452,33 @@ def _set_llm_loop_running(self, running: bool) -> None: if self._llm_loop_running == running: return self._llm_loop_running = running + if running: + self._schedule_generation_tick() + else: + self._stop_generation_animation() self._invalidate_prompt() + def _schedule_generation_tick(self) -> None: + if getattr(self, "_generation_tick", None) is not None: + return + self._generation_tick = asyncio.get_running_loop().call_later( + _PROMPT_REFRESH_INTERVAL, + self._tick_generation_status, + ) + + def _tick_generation_status(self) -> None: + self._generation_tick = None + if not self._llm_loop_running: + return + self._invalidate_prompt() + self._schedule_generation_tick() + + def _stop_generation_animation(self) -> None: + tick: asyncio.TimerHandle | None = getattr(self, "_generation_tick", None) + if tick is not None: + tick.cancel() + self._generation_tick = None + @staticmethod def _history_file(home: Path, workspace: Path) -> Path: workspace_hash = md5(str(workspace).encode("utf-8"), usedforsecurity=False).hexdigest() diff --git a/src/bub/channels/cli/ansi_bridge.py b/src/bub/channels/cli/ansi_bridge.py new file mode 100644 index 00000000..3a5460ca --- /dev/null +++ b/src/bub/channels/cli/ansi_bridge.py @@ -0,0 +1,23 @@ +from __future__ import annotations + +import re +from io import StringIO + +from rich.console import Console, RenderableType + +# prompt_toolkit's ANSI parser does not understand OSC 8 hyperlinks. Marking +# each OSC sequence as zero-width preserves the link without affecting layout. +_OSC8_RE = re.compile(r"\x1b\]8;[^\x07\x1b]*(?:\x1b\\|\x07)") + + +def render_to_ansi(renderable: RenderableType, *, width: int | None = None) -> str: + """Render a Rich object as ANSI text consumable by prompt_toolkit.""" + output = StringIO() + console = Console( + file=output, + force_terminal=True, + highlight=False, + width=width, + ) + console.print(renderable, end="") + return _OSC8_RE.sub(lambda match: f"\x01{match.group(0)}\x02", output.getvalue()) diff --git a/src/bub/channels/cli/terminal_output.py b/src/bub/channels/cli/terminal_output.py new file mode 100644 index 00000000..74fb6609 --- /dev/null +++ b/src/bub/channels/cli/terminal_output.py @@ -0,0 +1,140 @@ +from __future__ import annotations + +import asyncio +import sys +from collections.abc import Callable, Iterator +from contextlib import contextmanager, redirect_stderr, redirect_stdout +from typing import TextIO, cast + +from prompt_toolkit.application import get_app_or_none, run_in_terminal +from prompt_toolkit.output.base import Output +from prompt_toolkit.output.vt100 import Vt100_Output + +_BEGIN_SYNCHRONIZED_UPDATE = "\x1b[?2026h" +_END_SYNCHRONIZED_UPDATE = "\x1b[?2026l" + + +class _SynchronizedTextIO: + """Wrap complete terminal writes without reaching into prompt_toolkit buffers.""" + + def __init__(self, target: TextIO) -> None: + self._target = target + self._depth = 0 + + @property + def encoding(self) -> str | None: + return self._target.encoding + + def fileno(self) -> int: + return self._target.fileno() + + def isatty(self) -> bool: + return self._target.isatty() + + def write(self, data: str) -> int: + if self._depth: + return self._target.write(data) + return self._target.write(f"{_BEGIN_SYNCHRONIZED_UPDATE}{data}{_END_SYNCHRONIZED_UPDATE}") + + def flush(self) -> None: + self._target.flush() + + @contextmanager + def synchronized_update(self) -> Iterator[None]: + if self._depth == 0: + self._target.write(_BEGIN_SYNCHRONIZED_UPDATE) + self._target.flush() + self._depth += 1 + try: + yield + finally: + self._depth -= 1 + if self._depth == 0: + self._target.write(_END_SYNCHRONIZED_UPDATE) + self._target.flush() + + +class SynchronizedVt100Output(Vt100_Output): + """VT100 output that presents each rendered frame atomically.""" + + @contextmanager + def synchronized_update(self) -> Iterator[None]: + self.flush() + output = cast(_SynchronizedTextIO, self.stdout) + with output.synchronized_update(): + yield + self.flush() + + +def create_synchronized_output(stdout: TextIO | None = None) -> Output | None: + target = stdout if stdout is not None else sys.stdout + if sys.platform == "win32" or not target.isatty(): + return None + synchronized_stdout = cast(TextIO, _SynchronizedTextIO(target)) + return cast(SynchronizedVt100Output, SynchronizedVt100Output.from_pty(synchronized_stdout)) + + +def _original_stream(stream: TextIO) -> TextIO: + original = getattr(stream, "original_stdout", None) + return cast(TextIO, original) if original is not None else stream + + +@contextmanager +def direct_terminal_stdio() -> Iterator[None]: + """Bypass prompt_toolkit's deferred stdout proxy for an active terminal callback.""" + with redirect_stdout(_original_stream(sys.stdout)), redirect_stderr(_original_stream(sys.stderr)): + yield + + +async def restore_synchronized_prompt() -> None: + """Finish the CPR-dependent toolbar redraw before presenting a synchronized frame.""" + app = get_app_or_none() + if app is None or not app.is_running or not isinstance(app.output, SynchronizedVt100Output): + return + if app.renderer.waiting_for_cpr: + await app.renderer.wait_for_cpr_responses() + if app.is_running and not app.is_done and app.renderer.height_is_known: + prompt_finished = app.future + if prompt_finished is None: + return + rendered = asyncio.get_running_loop().create_future() + + def after_render(_) -> None: + if not rendered.done(): + rendered.set_result(None) + + app.after_render.add_handler(after_render) + try: + app.invalidate() + await asyncio.wait((rendered, prompt_finished), return_when=asyncio.FIRST_COMPLETED) + finally: + app.after_render.remove_handler(after_render) + + +@contextmanager +def synchronized_prompt_output() -> Iterator[None]: + app = get_app_or_none() + output = app.output if app is not None else None + if isinstance(output, SynchronizedVt100Output): + with output.synchronized_update(): + yield + return + yield + + +class TerminalPresenter: + """Serialize every write that temporarily interrupts the active prompt.""" + + def __init__(self) -> None: + self._lock = asyncio.Lock() + + async def write(self, function: Callable[[], None]) -> None: + async with self._lock: + + def write_directly() -> None: + with direct_terminal_stdio(): + function() + + with synchronized_prompt_output(): + await run_in_terminal(write_directly, render_cli_done=False) + await restore_synchronized_prompt() diff --git a/src/bub/channels/cli/writers.py b/src/bub/channels/cli/writers.py new file mode 100644 index 00000000..cc8d30b8 --- /dev/null +++ b/src/bub/channels/cli/writers.py @@ -0,0 +1,42 @@ +"""Response-scoped Markdown buffering for CLI streaming output.""" + +from __future__ import annotations + +from rich.markdown import Markdown +from rich.text import Text + +_MARKDOWN_CODE_THEME = "ansi_dark" + + +def _markdown(content: str) -> Markdown: + return Markdown(content, code_theme=_MARKDOWN_CODE_THEME) + + +class MarkdownWriter: + """Keep one response segment as a single Markdown document. + + Newlines are Markdown structure, not terminal commit boundaries. The + buffer is drained only at an explicit model or tool boundary. + """ + + def __init__(self) -> None: + self._buffer = "" + + def append(self, text: str) -> None: + self._buffer += text + + def render_live(self) -> Markdown | Text: + if not self._buffer.strip(): + return Text("") + return _markdown(self._buffer) + + def render_final(self) -> Markdown | None: + if not self._buffer.strip(): + return None + return _markdown(self._buffer.rstrip()) + + def clear(self) -> None: + self._buffer = "" + + def has_content(self) -> bool: + return bool(self._buffer.strip()) diff --git a/src/bub/tools.py b/src/bub/tools.py index 81019f16..071c6e62 100644 --- a/src/bub/tools.py +++ b/src/bub/tools.py @@ -6,7 +6,7 @@ import inspect import json import time -from collections.abc import Callable, Sequence +from collections.abc import Awaitable, Callable, Sequence from dataclasses import dataclass, field, replace from typing import TYPE_CHECKING, Any, Protocol, overload @@ -139,11 +139,11 @@ class ToolExecution: class ToolCallReporter(Protocol): - def start(self, name: str, args: tuple[Any, ...], kwargs: dict[str, Any]) -> None: ... + def start(self, name: str, args: tuple[Any, ...], kwargs: dict[str, Any]) -> Awaitable[None] | None: ... - def success(self, name: str, result: Any, elapsed_ms: float) -> None: ... + def success(self, name: str, result: Any, elapsed_ms: float) -> Awaitable[None] | None: ... - def error(self, name: str, error: BaseException, elapsed_ms: float) -> None: ... + def error(self, name: str, error: BaseException, elapsed_ms: float) -> Awaitable[None] | None: ... _TOOL_CALL_REPORTER: contextvars.ContextVar[ToolCallReporter | None] = contextvars.ContextVar( @@ -160,6 +160,11 @@ def tool_call_reporter(reporter: ToolCallReporter): _TOOL_CALL_REPORTER.reset(token) +async def _await_report(report: Awaitable[None] | None) -> None: + if report is not None: + await report + + class ToolExecutor: """Execute already-resolved Bub tool invocations.""" @@ -328,7 +333,7 @@ async def wrapped(*args, **kwargs): if reporter is None: _log_tool_call(tool.name, args, call_kwargs) else: - reporter.start(tool.name, args, call_kwargs) + await _await_report(reporter.start(tool.name, args, call_kwargs)) start = time.monotonic() try: @@ -340,14 +345,14 @@ async def wrapped(*args, **kwargs): if reporter is None: logger.exception("tool.call.error name={} elapsed_time={:.2f}ms", tool.name, elapsed_time) else: - reporter.error(tool.name, exc, elapsed_time) + await _await_report(reporter.error(tool.name, exc, elapsed_time)) raise else: elapsed_time = (time.monotonic() - start) * 1000 if reporter is None: logger.info("tool.call.success name={} elapsed_time={:.2f}ms", tool.name, elapsed_time) else: - reporter.success(tool.name, result, elapsed_time) + await _await_report(reporter.success(tool.name, result, elapsed_time)) return result return replace(tool, handler=wrapped) diff --git a/tests/test_channels.py b/tests/test_channels.py index 87b8aa2e..0d86c7c4 100644 --- a/tests/test_channels.py +++ b/tests/test_channels.py @@ -2,6 +2,7 @@ import asyncio import contextlib +import io import os import pty import re @@ -73,6 +74,11 @@ def _plain_terminal_text(raw: bytes) -> str: return ANSI_RE.sub("", text).replace("\r", "\n") +class _ImmediatePresenter: + async def write(self, function) -> None: + function() + + class _FakeChannelMixin: def __init__(self, name: str, *, needs_debounce: bool = False) -> None: self.name = name @@ -356,14 +362,11 @@ class FakePrompt: def __init__(self) -> None: self.inputs = iter(["first", "second", ",quit"]) self.refresh_intervals: list[float | None] = [] - self.messages: list[str] = [] self.received_callables: list[bool] = [] async def prompt_async(self, message, *, refresh_interval=None): self.refresh_intervals.append(refresh_interval) self.received_callables.append(callable(message)) - rendered = message() if callable(message) else message - self.messages.append("".join(part for _, part in rendered)) return next(self.inputs) async def on_receive(message: ChannelMessage) -> None: @@ -382,6 +385,7 @@ async def on_receive(message: ChannelMessage) -> None: channel._mode = "agent" channel._llm_loop_running = False channel._prompt = FakePrompt() + channel._presenter = _ImmediatePresenter() echoed: list[tuple[str, str]] = [] channel._renderer = SimpleNamespace( welcome=lambda **kwargs: None, @@ -392,23 +396,23 @@ async def on_receive(message: ChannelMessage) -> None: await asyncio.wait_for(channel._main_loop(), timeout=1) - import bub.channels.cli as cli_module - assert [message.content for message in received] == ["first", "second"] - assert channel._prompt.refresh_intervals == [cli_module._PROMPT_REFRESH_INTERVAL] * 3 + + assert channel._prompt.refresh_intervals == [None] * 3 assert channel._prompt.received_callables == [True, True, True] - assert "Generating\n" not in channel._prompt.messages[0] - assert "Generating\n" in channel._prompt.messages[1] assert echoed == [] assert all(message.lifespan is not None for message in received) def test_cli_channel_build_prompt_erases_submitted_prompt(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None: captured: dict[str, object] = {} + from prompt_toolkit.layout import HSplit class FakePromptSession: def __init__(self, **kwargs) -> None: captured.update(kwargs) + self.app = SimpleNamespace(min_redraw_interval=None) + self.layout = SimpleNamespace(container=HSplit([])) monkeypatch.setattr("bub.channels.cli.PromptSession", FakePromptSession) channel = CliChannel.__new__(CliChannel) @@ -421,38 +425,148 @@ def __init__(self, **kwargs) -> None: assert isinstance(prompt, FakePromptSession) assert captured["erase_when_done"] is True + assert len(prompt.layout.container.children) == 2 -def test_cli_channel_generating_spinner_renders_above_input_not_toolbar(monkeypatch: pytest.MonkeyPatch) -> None: +@pytest.mark.asyncio +async def test_cli_live_layout_keeps_markdown_tail_and_status_visible_when_output_exceeds_terminal( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from prompt_toolkit import PromptSession + from prompt_toolkit.formatted_text import FormattedText + from prompt_toolkit.input import create_pipe_input + from prompt_toolkit.output import DummyOutput + from prompt_toolkit.output.base import Size + from rich.console import Console + + from bub.channels.cli import _PROMPT_REFRESH_INTERVAL, _StreamPrinter + + class SizedOutput(DummyOutput): + def get_size(self) -> Size: + return Size(rows=35, columns=80) + + def rendered_screen_text(session: PromptSession[str]) -> str: + screen = session.app.renderer.last_rendered_screen + assert screen is not None + lines: list[str] = [] + for row_number in range(screen.height): + row = screen.data_buffer[row_number] + last_column = max(row.keys(), default=-1) + lines.append("".join(row[column].char for column in range(last_column + 1)).rstrip()) + return "\n".join(lines) + channel = CliChannel.__new__(CliChannel) - channel._llm_loop_running = True channel._mode = "agent" - channel._expand_thinking = False - channel._last_tape_info = None - channel._agent = SimpleNamespace(settings=SimpleNamespace(model="test-model")) + channel._llm_loop_running = True + channel._stream_printer = None + console = Console(file=io.StringIO(), force_terminal=True, width=80) + monkeypatch.setattr("bub.channels.cli.get_console", lambda: console) + + with create_pipe_input() as pipe_input: + prompt: PromptSession[str] = PromptSession( + input=pipe_input, + output=SizedOutput(), + bottom_toolbar=lambda: FormattedText([("", "toolbar")]), + erase_when_done=True, + ) + channel._prompt = prompt + channel._attach_live_layout(prompt) + prompt.app.min_redraw_interval = _PROMPT_REFRESH_INTERVAL + printer = _StreamPrinter( + console=console, + print_head=lambda: None, + expand_thinking=False, + presenter=_ImmediatePresenter(), + invalidate=prompt.app.invalidate, + ) + channel._stream_printer = printer + first_render = asyncio.get_running_loop().create_future() + + def after_first_render(_) -> None: + if not first_render.done(): + first_render.set_result(None) - prompt_text = "".join(part for _, part in channel._prompt_message()) - toolbar_text = "".join(part for _, part in channel._render_bottom_toolbar()) + prompt.app.after_render.add_handler(after_first_render) + prompt_task = asyncio.create_task(prompt.prompt_async(channel._prompt_message)) + after_live_render = None + try: + await asyncio.wait_for(first_render, timeout=1) + prompt.app.after_render.remove_handler(after_first_render) + live_render = asyncio.get_running_loop().create_future() - assert "\n" in prompt_text - assert "Generating\n" in prompt_text - assert prompt_text.endswith(f"{Path.cwd().name} > ") - assert "Generating" not in toolbar_text + def after_live_render(_) -> None: + if not live_render.done(): + live_render.set_result(None) - import bub.channels.cli as cli_module + prompt.app.after_render.add_handler(after_live_render) + paragraphs = "\n\n".join( + f"Paragraph {index}: terminal streaming content remains structured." for index in range(60) + ) + await printer.render(StreamEvent("text", {"delta": f"# Report\n\n{paragraphs}\n\nTAIL_MARKER"})) + await asyncio.wait_for(live_render, timeout=1) + prompt.app.after_render.remove_handler(after_live_render) + after_live_render = None + visible = rendered_screen_text(prompt) + finally: + with contextlib.suppress(ValueError): + prompt.app.after_render.remove_handler(after_first_render) + if after_live_render is not None: + prompt.app.after_render.remove_handler(after_live_render) + pipe_input.send_text("\n") + await prompt_task + + assert "TAIL_MARKER" in visible + assert "Generating" in visible + assert f"{Path.cwd().name} >" in visible + + +def test_cli_generation_spinner_refreshes_only_while_model_is_running( + monkeypatch: pytest.MonkeyPatch, +) -> None: + invalidations: list[None] = [] + callbacks: list[object] = [] - monkeypatch.setattr(cli_module, "monotonic", lambda: 0.0) - first_frame = "".join(part for _, part in channel._prompt_message()) - monkeypatch.setattr(cli_module, "monotonic", lambda: 0.2) - second_frame = "".join(part for _, part in channel._prompt_message()) + class FakeTimerHandle: + def __init__(self) -> None: + self._cancelled = False + + def cancel(self) -> None: + self._cancelled = True + + def cancelled(self) -> bool: + return self._cancelled + + class FakeLoop: + def call_later(self, delay, callback): + assert delay > 0 + callbacks.append(callback) + return FakeTimerHandle() + + monkeypatch.setattr("bub.channels.cli.asyncio.get_running_loop", FakeLoop) + channel = CliChannel.__new__(CliChannel) + channel._llm_loop_running = False + channel._generation_tick = None + channel._prompt = SimpleNamespace(app=SimpleNamespace(invalidate=lambda: invalidations.append(None))) - assert first_frame != second_frame + channel._set_llm_loop_running(True) + assert len(callbacks) == 1 + first_tick = channel._generation_tick + callbacks.pop()() + assert len(callbacks) == 1 + second_tick = channel._generation_tick + channel._set_llm_loop_running(False) + + assert first_tick is not second_tick + assert second_tick.cancelled() + assert len(invalidations) == 3 + assert channel._generation_tick is None @pytest.mark.asyncio async def test_cli_channel_admit_message_steers_when_turn_is_running() -> None: channel = CliChannel.__new__(CliChannel) channel._mode = "agent" + channel._presenter = _ImmediatePresenter() echoed: list[tuple[str, str, bool]] = [] channel._renderer = SimpleNamespace( input_echo=lambda prompt, text, steering=False: echoed.append((prompt, text, steering)), @@ -910,6 +1024,7 @@ async def test_cli_channel_stream_events_prints_stream_and_yields_events(monkeyp heads: list[str] = [] printed: list[tuple[str, str | None, bool | None]] = [] channel._renderer = SimpleNamespace(print_head=heads.append) + channel._presenter = _ImmediatePresenter() channel._expand_thinking = False monkeypatch.setattr( "bub.channels.cli.get_console", @@ -922,17 +1037,222 @@ async def test_cli_channel_stream_events_prints_stream_and_yields_events(monkeyp async def source() -> asyncio.AsyncIterator[StreamEvent]: yield StreamEvent("text", {"delta": " "}) - yield StreamEvent("text", {"delta": "hel"}) - yield StreamEvent("text", {"delta": "lo"}) + yield StreamEvent("text", {"delta": "first paragraph\n\n"}) + yield StreamEvent("text", {"delta": "second paragraph"}) yield StreamEvent("final", {}) yielded = [event async for event in channel.stream_events(message, source())] assert heads == ["command"] - assert printed == [("hel\n", "", False), ("hello\n", "", False)] + assert len(printed) == 1 + assert getattr(printed[0][0], "markup", None) == "first paragraph\n\nsecond paragraph" assert [event.kind for event in yielded] == ["text", "text", "final"] +@pytest.mark.asyncio +async def test_cli_channel_stream_error_preserves_partial_markdown(monkeypatch: pytest.MonkeyPatch) -> None: + channel = CliChannel.__new__(CliChannel) + printed: list[object] = [] + channel._renderer = SimpleNamespace(print_head=lambda kind: None) + channel._presenter = _ImmediatePresenter() + channel._expand_thinking = False + monkeypatch.setattr( + "bub.channels.cli.get_console", + lambda: SimpleNamespace( + width=80, + print=lambda content, end=None, highlight=None: printed.append(content), + ), + ) + + async def source() -> asyncio.AsyncIterator[StreamEvent]: + yield StreamEvent("text", {"delta": "# Partial response"}) + raise RuntimeError("stream failed") + + with pytest.raises(RuntimeError, match="stream failed"): + [event async for event in channel.stream_events(_message("ignored"), source())] + + assert any(getattr(item, "markup", None) == "# Partial response" for item in printed) + assert channel._stream_printer is None + + +@pytest.mark.asyncio +async def test_cli_channel_stream_cancellation_preserves_partial_markdown(monkeypatch: pytest.MonkeyPatch) -> None: + channel = CliChannel.__new__(CliChannel) + printed: list[object] = [] + partial_received = asyncio.Event() + channel._renderer = SimpleNamespace(print_head=lambda kind: None) + channel._presenter = _ImmediatePresenter() + channel._expand_thinking = False + monkeypatch.setattr( + "bub.channels.cli.get_console", + lambda: SimpleNamespace( + width=80, + print=lambda content, end=None, highlight=None: printed.append(content), + ), + ) + + async def source() -> asyncio.AsyncIterator[StreamEvent]: + yield StreamEvent("text", {"delta": "# Partial before cancellation"}) + partial_received.set() + await asyncio.Event().wait() + + async def consume() -> None: + [event async for event in channel.stream_events(_message("ignored"), source())] + + task = asyncio.create_task(consume()) + await partial_received.wait() + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + assert any(getattr(item, "markup", None) == "# Partial before cancellation" for item in printed) + assert channel._stream_printer is None + + +@pytest.mark.asyncio +async def test_cli_channel_final_write_cancellation_retries_partial_markdown(monkeypatch: pytest.MonkeyPatch) -> None: + channel = CliChannel.__new__(CliChannel) + printed: list[object] = [] + + class CancelFinalWriteOnce: + def __init__(self) -> None: + self.calls = 0 + + async def write(self, function) -> None: + self.calls += 1 + if self.calls == 2: + raise asyncio.CancelledError + function() + + presenter = CancelFinalWriteOnce() + channel._renderer = SimpleNamespace(print_head=lambda kind: None) + channel._presenter = presenter + channel._expand_thinking = False + monkeypatch.setattr( + "bub.channels.cli.get_console", + lambda: SimpleNamespace( + width=80, + print=lambda content, end=None, highlight=None: printed.append(content), + ), + ) + + async def source() -> asyncio.AsyncIterator[StreamEvent]: + yield StreamEvent("text", {"delta": "# Partial during final write"}) + yield StreamEvent("final", {}) + + with pytest.raises(asyncio.CancelledError): + [event async for event in channel.stream_events(_message("ignored"), source())] + + assert presenter.calls == 3 + assert any(getattr(item, "markup", None) == "# Partial during final write" for item in printed) + assert channel._stream_printer is None + + +@pytest.mark.asyncio +async def test_terminal_presenter_redraw_wait_stops_when_prompt_finishes(monkeypatch: pytest.MonkeyPatch) -> None: + from bub.channels.cli.terminal_output import SynchronizedVt100Output, restore_synchronized_prompt + + prompt_finished = asyncio.get_running_loop().create_future() + handlers: list[object] = [] + app = SimpleNamespace( + output=object.__new__(SynchronizedVt100Output), + is_running=True, + is_done=False, + future=prompt_finished, + renderer=SimpleNamespace(waiting_for_cpr=False, height_is_known=True), + after_render=SimpleNamespace( + add_handler=handlers.append, + remove_handler=handlers.remove, + ), + invalidate=lambda: prompt_finished.set_result(None), + ) + monkeypatch.setattr("bub.channels.cli.terminal_output.get_app_or_none", lambda: app) + + await asyncio.wait_for(restore_synchronized_prompt(), timeout=1) + + assert handlers == [] + + +@pytest.mark.asyncio +async def test_cli_tool_reporter_finishes_before_next_model_output() -> None: + from bub.channels.cli import _CliToolCallReporter + from bub.tools import REGISTRY, tool, tool_call_reporter + + events: list[str] = [] + + class OrderedPresenter: + async def write(self, function) -> None: + await asyncio.sleep(0) + function() + + renderer = SimpleNamespace( + tool_call_start=lambda **kwargs: events.append("tool-start"), + tool_call_success=lambda **kwargs: events.append("tool-success"), + tool_call_error=lambda **kwargs: events.append("tool-error"), + ) + presenter = OrderedPresenter() + reporter = _CliToolCallReporter(renderer, presenter) # type: ignore[arg-type] + tool_name = "tests.cli_ordered_tool" + REGISTRY.pop(tool_name, None) + + @tool(name=tool_name) + def ordered_tool() -> str: + events.append("tool-body") + return "done" + + try: + with tool_call_reporter(reporter): + assert await ordered_tool.run() == "done" + await presenter.write(lambda: events.append("next-model-text")) + finally: + REGISTRY.pop(tool_name, None) + + assert events == ["tool-start", "tool-body", "tool-success", "next-model-text"] + + +@pytest.mark.asyncio +async def test_cli_markdown_stream_keeps_and_caches_complete_live_response( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import bub.channels.cli as cli_module + from bub.channels.cli import _StreamPrinter + + printed: list[object] = [] + invalidations: list[None] = [] + render_calls: list[None] = [] + render_to_ansi = cli_module.render_to_ansi + + def count_render(*args, **kwargs) -> str: + render_calls.append(None) + return render_to_ansi(*args, **kwargs) + + monkeypatch.setattr(cli_module, "render_to_ansi", count_render) + console = SimpleNamespace( + width=80, + print=lambda content, end=None, highlight=None: printed.append(content), + ) + printer = _StreamPrinter( + console=console, + print_head=lambda: None, + expand_thinking=False, + presenter=_ImmediatePresenter(), + invalidate=lambda: invalidations.append(None), + ) + + await printer.render(StreamEvent("text", {"delta": "# Heading\n\n"})) + first_frame = printer.render_live_ansi(width=console.width) + assert printer.render_live_ansi(width=console.width) == first_frame + await printer.render(StreamEvent("text", {"delta": "Second paragraph"})) + second_frame = printer.render_live_ansi(width=console.width) + + assert "Heading" in first_frame + assert "Heading" in second_frame + assert "Second paragraph" in second_frame + assert printed == [] + assert len(invalidations) == 2 + assert len(render_calls) == 2 + + def test_cli_stream_output_does_not_overlap_active_pty_prompt() -> None: script = textwrap.dedent( """ @@ -942,20 +1262,39 @@ def test_cli_stream_output_does_not_overlap_active_pty_prompt() -> None: from prompt_toolkit.patch_stdout import patch_stdout from rich.console import Console - from bub.channels.cli import _StreamPrinter + import bub.channels.cli as cli_module + from bub.channels.cli import CliChannel, _StreamPrinter + from bub.channels.cli.terminal_output import TerminalPresenter, create_synchronized_output from bub.streaming import StreamEvent async def main(): console = Console(force_terminal=True, color_system=None, width=80) + cli_module.get_console = lambda: console + output = create_synchronized_output() + assert output is not None + session = PromptSession(erase_when_done=True, output=output) + session.app.min_redraw_interval = 0.08 + channel = CliChannel.__new__(CliChannel) + channel._mode = "agent" + channel._llm_loop_running = False + channel._generation_tick = None + channel._stream_printer = None + channel._prompt = session + channel._attach_live_layout(session) + presenter = TerminalPresenter() printer = _StreamPrinter( console=console, print_head=lambda: console.print("Assistant >"), expand_thinking=False, + presenter=presenter, + invalidate=session.app.invalidate, ) - session = PromptSession(erase_when_done=True) + channel._stream_printer = printer + channel._set_llm_loop_running(True) async def stream(): + await asyncio.sleep(0.35) chunks = [ "春风一夜入江城\\n", "细雨无声湿客", @@ -969,17 +1308,15 @@ async def stream(): await asyncio.sleep(0.03) await printer.render(StreamEvent("text", {"delta": chunk})) if index == 3: - await printer.commit_live_text() - console.print("bub > steer now") + await presenter.write(lambda: console.print("bub > steer now")) await asyncio.sleep(0.03) await printer.render(StreamEvent("final", {})) + channel._stream_printer = None + channel._set_llm_loop_running(False) task = asyncio.create_task(stream()) with patch_stdout(raw=True): - await session.prompt_async( - lambda: [("", "\\n* Generating\\nbub > ")], - refresh_interval=0.02, - ) + await session.prompt_async(channel._prompt_message) await task @@ -989,6 +1326,7 @@ async def stream(): master_fd, slave_fd = pty.openpty() env = os.environ.copy() env["PYTHONPATH"] = f"{Path.cwd() / 'src'}{os.pathsep}{env.get('PYTHONPATH', '')}" + env["TERM"] = "xterm-256color" process = subprocess.Popen( [sys.executable, "-c", script], stdin=slave_fd, @@ -1000,9 +1338,25 @@ async def stream(): ) os.close(slave_fd) try: - time.sleep(0.25) + before_input = bytearray() + deadline = time.monotonic() + 15 + final_text = "明朝山色满前庭".encode() + while time.monotonic() < deadline: + readable, _, _ = select.select([master_fd], [], [], 0.05) + if not readable: + continue + chunk = os.read(master_fd, 65536) + before_input.extend(chunk) + from bub.channels.cli import _GENERATION_SPINNER + + frames = {frame for frame in _GENERATION_SPINNER if frame.encode() in before_input} + if final_text in before_input and len(frames) >= 2: + break + else: + pytest.fail(before_input.decode(errors="replace")) + os.write(master_fd, b"next\n") - raw_output = _read_pty_until_exit(master_fd, process) + raw_output = bytes(before_input) + _read_pty_until_exit(master_fd, process, timeout=15) finally: if process.poll() is None: process.terminate() @@ -1022,23 +1376,29 @@ async def stream(): assert "明朝山色满前庭bub >" not in output assert "明朝山色满前庭* Generating" not in output + from bub.channels.cli import _GENERATION_SPINNER + + spinner_frames = {frame for frame in _GENERATION_SPINNER if frame.encode() in raw_output} + assert len(spinner_frames) >= 2 + @pytest.mark.asyncio -async def test_cli_channel_input_echo_commits_active_stream_line() -> None: +async def test_cli_channel_steering_echo_does_not_finish_active_markdown() -> None: channel = CliChannel.__new__(CliChannel) calls: list[str] = [] class FakeStreamPrinter: - async def commit_live_text(self) -> None: - calls.append("commit") + async def finish(self) -> None: + calls.append("finish") channel._stream_printer = FakeStreamPrinter() channel._mode = "agent" + channel._presenter = _ImmediatePresenter() channel._renderer = SimpleNamespace(input_echo=lambda prompt, text, steering=False: calls.append(f"echo:{text}")) - await channel._echo_input("steer now") + await channel._echo_input("steer now", steering=True) - assert calls == ["commit", "echo:steer now"] + assert calls == ["echo:steer now"] @pytest.mark.asyncio @@ -1047,6 +1407,7 @@ async def test_cli_channel_collapsed_reasoning_does_not_start_status_spinner( ) -> None: channel = CliChannel.__new__(CliChannel) channel._renderer = SimpleNamespace(print_head=lambda kind: None) + channel._presenter = _ImmediatePresenter() channel._expand_thinking = False printed: list[object] = [] @@ -1072,7 +1433,7 @@ async def source() -> asyncio.AsyncIterator[StreamEvent]: assert [event.kind for event in yielded] == ["reasoning", "text", "final"] assert printed - assert any("hello" in str(item) for item in printed) + assert any(getattr(item, "markup", None) == "hello" for item in printed) def test_cli_channel_history_file_uses_workspace_hash(tmp_path: Path) -> None: