1515from prompt_toolkit .history import FileHistory
1616from prompt_toolkit .key_binding import KeyBindings
1717from prompt_toolkit .patch_stdout import patch_stdout
18+ from prompt_toolkit .utils import get_cwidth
1819from rich import get_console
1920from rich .spinner import SPINNERS
2021from 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 ()
0 commit comments