184 lines
5.7 KiB
Python
184 lines
5.7 KiB
Python
"""Streaming terminal renderer.
|
|
|
|
In a TTY the assistant's reply renders live into a :class:`rich.live.Live`
|
|
frame: already-completed Markdown blocks are laid out once (boundary-flushed)
|
|
while the in-fight tail streams as bright-white plain text, and the thinking
|
|
shows as dim italic text above. On a non-TTY or with ``--plain`` the raw
|
|
Markdown source streams to stdout instead (reasoning stays silent there).
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import sys
|
|
from typing import Optional
|
|
|
|
from rich.console import Console, Group, RenderableType
|
|
from rich.live import Live
|
|
from rich.markdown import Markdown
|
|
from rich.rule import Rule
|
|
from rich.text import Text
|
|
|
|
REASONING_STYLE = "bright_black italic"
|
|
CONTENT_STYLE = "white"
|
|
WARN_STYLE = "yellow"
|
|
ERROR_STYLE = "red bold"
|
|
FORCE_FLUSH_CHARS = 4000
|
|
|
|
|
|
def is_tty() -> bool:
|
|
return sys.stdout.isatty()
|
|
|
|
|
|
class StreamRenderer:
|
|
"""Accumulates a turn and renders it live (or plainly)."""
|
|
|
|
def __init__(self, *, plain: Optional[bool] = None):
|
|
auto_plain = not is_tty()
|
|
self.plain = plain if plain is not None else auto_plain
|
|
self.reasoning: list[str] = []
|
|
self.content_parts: list[str] = []
|
|
self.flushed: list[RenderableType] = []
|
|
self.tail = ""
|
|
self._messages: list[RenderableType] = [] # sticky info/error lines
|
|
self._live: Optional[Live] = None
|
|
self._rendered = 0
|
|
self._console = Console(
|
|
color_system="auto",
|
|
markup=False,
|
|
highlight=None,
|
|
tab_size=4,
|
|
force_terminal=not self.plain and is_tty(),
|
|
)
|
|
|
|
# -- lifecycle -------------------------------------------------------
|
|
|
|
def start(self) -> None:
|
|
if not self.plain:
|
|
self._live = Live(
|
|
console=self._console,
|
|
refresh_per_second=12,
|
|
vertical_overflow="visible",
|
|
)
|
|
self._live.start(refresh=False)
|
|
|
|
def stop(self) -> None:
|
|
if self._live is not None:
|
|
self._live.stop()
|
|
self._live = None
|
|
|
|
def __enter__(self) -> "StreamRenderer":
|
|
self.start()
|
|
return self
|
|
|
|
def __exit__(self, *exc) -> None:
|
|
if not self.plain:
|
|
self.finish()
|
|
self.stop()
|
|
|
|
# -- content ---------------------------------------------------------
|
|
|
|
def add_reasoning(self, text: str) -> None:
|
|
if not text:
|
|
return
|
|
self.reasoning.append(text)
|
|
if not self.plain:
|
|
self.refresh()
|
|
|
|
def add_content(self, text: str) -> None:
|
|
if not text:
|
|
return
|
|
self.content_parts.append(text)
|
|
if self.plain:
|
|
sys.stdout.write(text)
|
|
sys.stdout.flush()
|
|
return
|
|
self.tail += text
|
|
self._maybe_flush()
|
|
self.refresh()
|
|
|
|
def add_info(self, message: str) -> None:
|
|
if not message:
|
|
return
|
|
if self.plain:
|
|
sys.stderr.write(f"shellbound: {message}\n")
|
|
sys.stderr.flush()
|
|
else:
|
|
self._messages.append(Text(str(message), style=WARN_STYLE))
|
|
self.refresh()
|
|
|
|
def add_error(self, message: str) -> None:
|
|
if self.plain:
|
|
sys.stderr.write(f"shellbound: {message}\n")
|
|
sys.stderr.flush()
|
|
else:
|
|
self._messages.append(Text(f"error: {message}", style=ERROR_STYLE))
|
|
self.refresh()
|
|
|
|
# -- flushing --------------------------------------------------------
|
|
|
|
def _fence_count(self) -> int:
|
|
return self.tail.count("```")
|
|
|
|
def _maybe_flush(self) -> None:
|
|
if not self.tail:
|
|
return
|
|
fences = self._fence_count()
|
|
if fences % 2 == 1:
|
|
# inside an open fenced block: only force-flush raw at the cap
|
|
if len(self.tail) > FORCE_FLUSH_CHARS:
|
|
self._flush(markdown=False)
|
|
return
|
|
if self.tail.rstrip().endswith("```"):
|
|
# fenced block just closed
|
|
self._flush(markdown=True)
|
|
return
|
|
if self.tail.endswith("\n\n"):
|
|
self._flush(markdown=True)
|
|
elif len(self.tail) > FORCE_FLUSH_CHARS:
|
|
self._flush(markdown=False)
|
|
|
|
def _flush(self, *, markdown: bool) -> None:
|
|
block = self.tail
|
|
self.tail = ""
|
|
if not block.strip():
|
|
return
|
|
if markdown:
|
|
self.flushed.append(Markdown(block))
|
|
else:
|
|
self.flushed.append(Text(block, style=CONTENT_STYLE))
|
|
|
|
# -- display ---------------------------------------------------------
|
|
|
|
def _renderable(self) -> RenderableType:
|
|
items: list[RenderableType] = []
|
|
if self.reasoning:
|
|
items.append(Text("".join(self.reasoning), style=REASONING_STYLE))
|
|
if self.content_parts or self.flushed or self.tail:
|
|
items.append(Rule(style="bright_black"))
|
|
items.extend(self.flushed)
|
|
if self.tail:
|
|
items.append(Text(self.tail, style=CONTENT_STYLE))
|
|
items.extend(self._messages)
|
|
if not items:
|
|
return Text("")
|
|
if len(items) == 1:
|
|
return items[0]
|
|
return Group(*items)
|
|
|
|
def refresh(self) -> None:
|
|
if self._live is not None:
|
|
self._live.update(self._renderable())
|
|
|
|
def finish(self) -> str:
|
|
"""Flush any pending tail and return the raw streamed reply text."""
|
|
if not self.plain and self.tail.strip():
|
|
self._flush(markdown=True)
|
|
if not self.plain:
|
|
self.refresh()
|
|
return "".join(self.content_parts)
|
|
|
|
# -- plain helpers ---------------------------------------------------
|
|
|
|
@property
|
|
def reasoning_text(self) -> str:
|
|
return "".join(self.reasoning) |