Skip to content

Commit b12ada4

Browse files
authored
fix(channel/ci): fix stream output mix (#242)
Signed-off-by: Frost Ming <me@frostming.com>
1 parent 4448f10 commit b12ada4

2 files changed

Lines changed: 233 additions & 10 deletions

File tree

src/bub/channels/cli/__init__.py

Lines changed: 87 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,7 @@
1515
from prompt_toolkit.history import FileHistory
1616
from prompt_toolkit.key_binding import KeyBindings
1717
from prompt_toolkit.patch_stdout import patch_stdout
18+
from prompt_toolkit.utils import get_cwidth
1819
from rich import get_console
1920
from rich.spinner import SPINNERS
2021
from rich.text import Text
@@ -43,6 +44,9 @@ def __init__(self, *, console, print_head: Callable[[], None], expand_thinking:
4344
self._expand_thinking = expand_thinking
4445
self._reasoning_chars = 0
4546
self._reasoning_streaming = False
47+
self._current_text_line = ""
48+
self._rendered_text_line: str | None = None
49+
self._live_text_rows = 0
4650
self.head_printed = False
4751

4852
async def render(self, event: StreamEvent) -> bool:
@@ -77,19 +81,23 @@ async def _print_content(self, content: str) -> bool:
7781
await self._ensure_head()
7882
await self._close_reasoning_stream()
7983
await self._flush_reasoning()
80-
await self._print(content, end="", highlight=False)
84+
await self._write_text(content)
8185
return True
8286

8387
async def _print_end(self) -> None:
8488
if self._reasoning_chars:
8589
await self._ensure_head()
8690
await self._flush_reasoning()
87-
if self.head_printed:
91+
if self._current_text_line:
92+
await self._commit_text_line()
93+
elif self.head_printed and not self._live_text_rows:
8894
await self._print("")
8995

9096
async def _print_stream_boundary(self) -> None:
9197
await self._close_reasoning_stream()
9298
await self._flush_reasoning()
99+
if self._current_text_line or self._live_text_rows:
100+
await self._commit_text_line()
93101
if self.head_printed:
94102
await self._print("")
95103

@@ -112,6 +120,65 @@ async def _flush_reasoning(self) -> None:
112120
await self._print(Tree(label, guide_style="dim", expanded=False))
113121
self._reasoning_chars = 0
114122

123+
async def _write_text(self, text: str) -> None:
124+
parts = text.split("\n")
125+
for index, part in enumerate(parts):
126+
self._current_text_line += part
127+
if index < len(parts) - 1:
128+
await self._commit_text_line()
129+
130+
if self._current_text_line:
131+
await self._render_live_text_line()
132+
133+
async def _commit_text_line(self) -> None:
134+
if self._live_text_rows and self._rendered_text_line == self._current_text_line:
135+
self._current_text_line = ""
136+
self._rendered_text_line = None
137+
self._live_text_rows = 0
138+
return
139+
self._live_text_rows = await self._render_text_line(self._current_text_line)
140+
self._current_text_line = ""
141+
self._rendered_text_line = None
142+
self._live_text_rows = 0
143+
144+
async def commit_live_text(self) -> None:
145+
if self._current_text_line or self._live_text_rows:
146+
await self._commit_text_line()
147+
148+
async def _render_live_text_line(self) -> None:
149+
self._live_text_rows = await self._render_text_line(self._current_text_line)
150+
self._rendered_text_line = self._current_text_line
151+
152+
async def _render_text_line(self, text: str) -> int:
153+
previous_rows = self._live_text_rows
154+
rows = self._display_rows(text)
155+
156+
def render() -> None:
157+
self._rewind_live_text(previous_rows)
158+
self._console.print(f"{text}\n", end="", highlight=False)
159+
160+
await run_in_terminal(render, render_cli_done=False)
161+
return rows
162+
163+
def _display_rows(self, text: str) -> int:
164+
columns = max(1, int(getattr(self._console, "width", 80) or 80))
165+
return max(1, (get_cwidth(text) + columns - 1) // columns)
166+
167+
def _rewind_live_text(self, rows: int) -> None:
168+
if rows <= 0:
169+
return
170+
output = getattr(self._console, "file", None)
171+
if output is None:
172+
return
173+
output.write(f"\x1b[{rows}A\r")
174+
for row in range(rows):
175+
output.write("\x1b[2K")
176+
if row < rows - 1:
177+
output.write("\x1b[1B\r")
178+
if rows > 1:
179+
output.write(f"\x1b[{rows - 1}A\r")
180+
output.flush()
181+
115182
async def _print(self, *args: Any, **kwargs: Any) -> None:
116183
await run_in_terminal(lambda: self._console.print(*args, **kwargs), render_cli_done=False)
117184

@@ -148,6 +215,7 @@ def __init__(self, on_receive: MessageHandler, agent: Agent) -> None:
148215
self._expand_thinking = False
149216
self._llm_loop_running = False
150217
self._main_task: asyncio.Task | None = None
218+
self._stream_printer: _StreamPrinter | None = None
151219
self._renderer = CliRenderer(get_console())
152220
self._last_tape_info: TapeInfo | None = None
153221
self._workspace = self._agent.framework.workspace
@@ -208,12 +276,12 @@ async def _main_loop(self) -> None:
208276
if raw in {",quit", ",exit"}:
209277
break
210278
if raw == ",thinking":
211-
self._renderer.input_echo(self._prompt_label(), raw)
279+
await self._echo_input(raw)
212280
self._toggle_thinking()
213281
continue
214282

215283
request = self._normalize_input(raw)
216-
self._renderer.input_echo(self._prompt_label(), raw)
284+
await self._echo_input(raw)
217285

218286
message = ChannelMessage(
219287
session_id=self._message_template["session_id"],
@@ -264,6 +332,12 @@ def _prompt_label(self) -> str:
264332
symbol = ">" if self._mode == "agent" else ","
265333
return f"{cwd} {symbol} "
266334

335+
async def _echo_input(self, raw: str) -> None:
336+
stream_printer = getattr(self, "_stream_printer", None)
337+
if stream_printer is not None:
338+
await stream_printer.commit_live_text()
339+
self._renderer.input_echo(self._prompt_label(), raw)
340+
267341
async def stream_events(
268342
self, message: ChannelMessage, stream: AsyncIterable[StreamEvent]
269343
) -> AsyncIterable[StreamEvent]:
@@ -273,10 +347,15 @@ async def stream_events(
273347
print_head=lambda: self._renderer.print_head(message.kind),
274348
expand_thinking=self._expand_thinking,
275349
)
276-
with tool_call_reporter(_CliToolCallReporter(self._renderer)):
277-
async for event in stream:
278-
if await printer.render(event):
279-
yield event
350+
self._stream_printer = printer
351+
try:
352+
with tool_call_reporter(_CliToolCallReporter(self._renderer)):
353+
async for event in stream:
354+
if await printer.render(event):
355+
yield event
356+
finally:
357+
if self._stream_printer is printer:
358+
self._stream_printer = None
280359

281360
def _build_prompt(self, workspace: Path) -> PromptSession[str]:
282361
kb = KeyBindings()

tests/test_channels.py

Lines changed: 146 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,14 @@
22

33
import asyncio
44
import contextlib
5+
import os
6+
import pty
7+
import re
8+
import select
9+
import subprocess
10+
import sys
11+
import textwrap
12+
import time
513
from datetime import datetime
614
from pathlib import Path
715
from types import SimpleNamespace
@@ -18,6 +26,8 @@
1826
from bub.runtime import StreamEvent
1927
from bub.turn_admission import AdmitDecision, SessionTurnController, SteeringBuffer, TurnSnapshot
2028

29+
ANSI_RE = re.compile(r"\x1b(?:\[[0-?]*[ -/]*[@-~]|\][^\x07]*(?:\x07|\x1b\\)|[()][A-Za-z])")
30+
2131

2232
def _load_channel_config(
2333
load_config,
@@ -35,6 +45,32 @@ def _load_channel_config(
3545
load_config(content)
3646

3747

48+
def _read_pty_until_exit(master_fd: int, process: subprocess.Popen[bytes], *, timeout: float = 3.0) -> bytes:
49+
chunks: list[bytes] = []
50+
deadline = time.monotonic() + timeout
51+
while time.monotonic() < deadline:
52+
if process.poll() is not None:
53+
with contextlib.suppress(OSError):
54+
chunks.append(os.read(master_fd, 65536))
55+
break
56+
readable, _, _ = select.select([master_fd], [], [], 0.05)
57+
if not readable:
58+
continue
59+
try:
60+
chunk = os.read(master_fd, 65536)
61+
except OSError:
62+
break
63+
if not chunk:
64+
break
65+
chunks.append(chunk)
66+
return b"".join(chunks)
67+
68+
69+
def _plain_terminal_text(raw: bytes) -> str:
70+
text = raw.decode(errors="replace")
71+
return ANSI_RE.sub("", text).replace("\r", "\n")
72+
73+
3874
class _FakeChannelMixin:
3975
def __init__(self, name: str, *, needs_debounce: bool = False) -> None:
4076
self.name = name
@@ -803,10 +839,118 @@ async def source() -> asyncio.AsyncIterator[StreamEvent]:
803839
yielded = [event async for event in channel.stream_events(message, source())]
804840

805841
assert heads == ["command"]
806-
assert printed == [("hel", "", False), ("lo", "", False), ("", None, None)]
842+
assert printed == [("hel\n", "", False), ("hello\n", "", False)]
807843
assert [event.kind for event in yielded] == ["text", "text", "final"]
808844

809845

846+
def test_cli_stream_output_does_not_overlap_active_pty_prompt() -> None:
847+
script = textwrap.dedent(
848+
"""
849+
import asyncio
850+
851+
from prompt_toolkit import PromptSession
852+
from prompt_toolkit.patch_stdout import patch_stdout
853+
from rich.console import Console
854+
855+
from bub.channels.cli import _StreamPrinter
856+
from bub.runtime import StreamEvent
857+
858+
859+
async def main():
860+
console = Console(force_terminal=True, color_system=None, width=80)
861+
printer = _StreamPrinter(
862+
console=console,
863+
print_head=lambda: console.print("Assistant >"),
864+
expand_thinking=False,
865+
)
866+
session = PromptSession(erase_when_done=True)
867+
868+
async def stream():
869+
chunks = [
870+
"春风一夜入江城\\n",
871+
"细雨无声湿客",
872+
"程\\n",
873+
"莫问归帆何处",
874+
"去\\n",
875+
"明朝山色满",
876+
"前庭",
877+
]
878+
for index, chunk in enumerate(chunks):
879+
await asyncio.sleep(0.03)
880+
await printer.render(StreamEvent("text", {"delta": chunk}))
881+
if index == 3:
882+
await printer.commit_live_text()
883+
console.print("bub > steer now")
884+
await asyncio.sleep(0.03)
885+
await printer.render(StreamEvent("final", {}))
886+
887+
task = asyncio.create_task(stream())
888+
with patch_stdout(raw=True):
889+
await session.prompt_async(
890+
lambda: [("", "\\n* Generating\\nbub > ")],
891+
refresh_interval=0.02,
892+
)
893+
await task
894+
895+
896+
asyncio.run(main())
897+
"""
898+
)
899+
master_fd, slave_fd = pty.openpty()
900+
env = os.environ.copy()
901+
env["PYTHONPATH"] = f"{Path.cwd() / 'src'}{os.pathsep}{env.get('PYTHONPATH', '')}"
902+
process = subprocess.Popen(
903+
[sys.executable, "-c", script],
904+
stdin=slave_fd,
905+
stdout=slave_fd,
906+
stderr=slave_fd,
907+
cwd=Path.cwd(),
908+
env=env,
909+
close_fds=True,
910+
)
911+
os.close(slave_fd)
912+
try:
913+
time.sleep(0.25)
914+
os.write(master_fd, b"next\n")
915+
raw_output = _read_pty_until_exit(master_fd, process)
916+
finally:
917+
if process.poll() is None:
918+
process.terminate()
919+
with contextlib.suppress(subprocess.TimeoutExpired):
920+
process.wait(timeout=1)
921+
os.close(master_fd)
922+
923+
assert process.wait(timeout=1) == 0, raw_output.decode(errors="replace")
924+
output = _plain_terminal_text(raw_output)
925+
926+
assert "春风一夜入江城" in output
927+
assert "细雨无声湿客程" in output
928+
assert "莫问归帆何处" in output
929+
assert "去" in output
930+
assert "明朝山色满前庭" in output
931+
assert "bub > steer now" in output
932+
assert "明朝山色满前庭bub >" not in output
933+
assert "明朝山色满前庭* Generating" not in output
934+
935+
936+
@pytest.mark.asyncio
937+
async def test_cli_channel_input_echo_commits_active_stream_line() -> None:
938+
channel = CliChannel.__new__(CliChannel)
939+
calls: list[str] = []
940+
941+
class FakeStreamPrinter:
942+
async def commit_live_text(self) -> None:
943+
calls.append("commit")
944+
945+
channel._stream_printer = FakeStreamPrinter()
946+
channel._mode = "agent"
947+
channel._renderer = SimpleNamespace(input_echo=lambda prompt, text: calls.append(f"echo:{text}"))
948+
949+
await channel._echo_input("steer now")
950+
951+
assert calls == ["commit", "echo:steer now"]
952+
953+
810954
@pytest.mark.asyncio
811955
async def test_cli_channel_collapsed_reasoning_does_not_start_status_spinner(
812956
monkeypatch: pytest.MonkeyPatch,
@@ -838,7 +982,7 @@ async def source() -> asyncio.AsyncIterator[StreamEvent]:
838982

839983
assert [event.kind for event in yielded] == ["reasoning", "text", "final"]
840984
assert printed
841-
assert "hello" in [str(item) for item in printed]
985+
assert any("hello" in str(item) for item in printed)
842986

843987

844988
def test_cli_channel_history_file_uses_workspace_hash(tmp_path: Path) -> None:

0 commit comments

Comments
 (0)