feat(thread-picker): add ThreadPickerWidget for inline thread selection in TUI

This commit is contained in:
X-iZhang
2026-03-10 00:45:21 +00:00
parent aad431785c
commit 903ef4d158
4 changed files with 511 additions and 8 deletions
+108 -8
View File
@@ -275,6 +275,7 @@ def run_textual_interactive(
BINDINGS = [
Binding("ctrl+c", "request_quit", "Quit", show=False),
Binding("up", "edit_queued", show=False, priority=True),
Binding("down", "down_delegate", show=False, priority=True),
Binding("escape", "cancel_queued", show=False, priority=True),
]
@@ -306,6 +307,7 @@ def run_textual_interactive(
self._comp_index: int = -1
self._hitl_auto_approve: bool = False
self._approval_future: asyncio.Future | None = None
self._picker_future: asyncio.Future | None = None
self._history_suggester = HistorySuggester(get_config_dir() / "history")
# ── Layout ─────────────────────────────────────────────
@@ -416,6 +418,33 @@ def run_textual_interactive(
if self._approval_future and not self._approval_future.done():
self._approval_future.set_result(event)
async def _wait_for_thread_pick(self, picker_widget) -> str | None:
"""Wait for user to pick a thread from ThreadPickerWidget.
Returns the selected thread_id, or ``None`` on cancel/timeout.
"""
self._picker_future = asyncio.get_event_loop().create_future()
try:
return await asyncio.wait_for(self._picker_future, timeout=120)
except (asyncio.TimeoutError, asyncio.CancelledError):
return None
finally:
self._picker_future = None
try:
picker_widget.remove()
except Exception:
pass
def on_thread_picker_widget_picked(self, event) -> None: # type: ignore[override]
"""Handle ThreadPickerWidget.Picked message."""
if self._picker_future and not self._picker_future.done():
self._picker_future.set_result(event.thread_id)
def on_thread_picker_widget_cancelled(self, event) -> None: # type: ignore[override]
"""Handle ThreadPickerWidget.Cancelled message."""
if self._picker_future and not self._picker_future.done():
self._picker_future.set_result(None)
# ── Streaming core ─────────────────────────────────────
async def _stream_with_widgets(
@@ -1138,7 +1167,10 @@ def run_textual_interactive(
if text.startswith("/"):
self._hide_completions()
await self._handle_command(text)
# Launch as independent task to free the message pump.
# Commands like /resume mount interactive widgets that need
# the pump to process key events and message bubbling.
asyncio.ensure_future(self._handle_command(text))
return
self._history_suggester.append_entry(text)
@@ -1183,26 +1215,34 @@ def run_textual_interactive(
def action_cancel_queued(self) -> None:
"""Cancel the last queued message on Esc."""
# Delegate to ApprovalWidget if it has focus
# Delegate to ApprovalWidget or ThreadPickerWidget if focused
focused = self.focused
if focused is not None:
from .widgets.approval_widget import ApprovalWidget
from .widgets.thread_selector import ThreadPickerWidget
if isinstance(focused, ApprovalWidget):
focused.action_select_reject()
return
if isinstance(focused, ThreadPickerWidget):
focused.action_cancel()
return
if self._queued_messages:
self._queued_messages.pop()
self._render_queue_indicator()
def action_edit_queued(self) -> None:
"""Pop the last queued message back into input for editing."""
# Skip if an ApprovalWidget has focus — let it handle up/down
# Skip if an ApprovalWidget or ThreadPickerWidget has focus
focused = self.focused
if focused is not None:
from .widgets.approval_widget import ApprovalWidget
from .widgets.thread_selector import ThreadPickerWidget
if isinstance(focused, ApprovalWidget):
focused.action_move_up()
return
if isinstance(focused, ThreadPickerWidget):
focused.action_move_up()
return
if self._queued_messages:
last = self._queued_messages.pop()
prompt = self.query_one("#prompt", Input)
@@ -1211,6 +1251,19 @@ def run_textual_interactive(
prompt.focus()
self._render_queue_indicator()
def action_down_delegate(self) -> None:
"""Delegate down key to focused ApprovalWidget or ThreadPickerWidget."""
focused = self.focused
if focused is not None:
from .widgets.approval_widget import ApprovalWidget
from .widgets.thread_selector import ThreadPickerWidget
if isinstance(focused, ApprovalWidget):
focused.action_move_down()
return
if isinstance(focused, ThreadPickerWidget):
focused.action_move_down()
return
def on_key(self, event: Any) -> None:
comp_widget = self.query_one("#completions", Static)
if not (comp_widget.display and self._comp_items):
@@ -1446,9 +1499,32 @@ def run_textual_interactive(
async def _cmd_resume(self, arg: str) -> None:
if not arg:
self._append_system("Usage: /resume <thread-id-prefix>", style="yellow")
await self._cmd_threads()
return
# Show inline thread picker
threads = await list_threads(
limit=0,
include_message_count=True,
include_preview=True,
)
if not threads:
self._append_system("No sessions to resume.", style="yellow")
return
from .widgets.thread_selector import ThreadPickerWidget
container = self.query_one("#chat", VerticalScroll)
picker = ThreadPickerWidget(
threads,
current_thread=self._conversation_tid,
title=">>> Select session to resume <<<",
)
await container.mount(picker)
container.scroll_end(animate=False)
picker.focus()
selected = await self._wait_for_thread_pick(picker)
if selected is None:
return
arg = selected
resolved = await self._resolve_thread_id(arg)
if not resolved:
@@ -1474,8 +1550,32 @@ def run_textual_interactive(
async def _cmd_delete(self, arg: str) -> None:
if not arg:
self._append_system("Usage: /delete <thread-id-prefix>", style="yellow")
return
# Show inline thread picker for deletion
threads = await list_threads(
limit=0,
include_message_count=True,
include_preview=True,
)
if not threads:
self._append_system("No sessions to delete.", style="yellow")
return
from .widgets.thread_selector import ThreadPickerWidget
container = self.query_one("#chat", VerticalScroll)
picker = ThreadPickerWidget(
threads,
current_thread=self._conversation_tid,
title=">>> Select session to delete <<<",
)
await container.mount(picker)
container.scroll_end(animate=False)
picker.focus()
selected = await self._wait_for_thread_pick(picker)
if selected is None:
return
arg = selected
resolved = await self._resolve_thread_id(arg)
if not resolved:
+2
View File
@@ -10,6 +10,7 @@ from .user_message import UserMessage
from .system_message import SystemMessage
from .usage_widget import UsageWidget
from .approval_widget import ApprovalWidget
from .thread_selector import ThreadPickerWidget
__all__ = [
"LoadingWidget",
@@ -22,4 +23,5 @@ __all__ = [
"SystemMessage",
"UsageWidget",
"ApprovalWidget",
"ThreadPickerWidget",
]
+197
View File
@@ -0,0 +1,197 @@
"""Inline thread picker widget for /resume and /delete in TUI.
Keyboard-driven widget mounted directly into the chat container (like
ApprovalWidget). Posts ``ThreadPickerWidget.Picked`` when user selects
a thread, or ``ThreadPickerWidget.Cancelled`` on Esc.
"""
from __future__ import annotations
from typing import TYPE_CHECKING, Any, ClassVar
from rich.text import Text
from textual.binding import Binding, BindingType
from textual.containers import Container
from textual.message import Message
from textual.widget import Widget
from textual.widgets import Static
if TYPE_CHECKING:
from textual import events
from textual.app import ComposeResult
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def build_row_text(
thread: dict,
*,
selected: bool = False,
current: bool = False,
) -> Text:
"""Build a Rich Text object for a single thread row.
Pure function — no Textual app context required, safe for unit tests.
"""
from ...sessions import _format_relative_time
tid = thread["thread_id"]
preview = thread.get("preview", "") or ""
msgs = thread.get("message_count", 0)
model = thread.get("model", "") or ""
when = _format_relative_time(thread.get("updated_at"))
line = Text()
cursor = "\u25b8 " if selected else " "
line.append(cursor, style="bold cyan" if selected else "dim")
line.append(f"{tid}", style="bold" if selected else "")
if current:
line.append(" *", style="bold green")
line.append(" ")
if preview:
display_preview = preview[:40] + "\u2026" if len(preview) > 40 else preview
line.append(display_preview, style="")
line.append(" ", style="")
line.append(f"({msgs} msgs)", style="dim")
if model:
line.append(f" {model}", style="dim italic")
if when:
line.append(f" {when}", style="dim")
return line
# ---------------------------------------------------------------------------
# Widget
# ---------------------------------------------------------------------------
class ThreadPickerWidget(Widget):
"""Inline thread picker — mounts in chat, keyboard-driven.
Posts ``Picked(thread_id)`` on Enter, ``Cancelled()`` on Esc.
Follows the same pattern as ``ApprovalWidget``.
"""
can_focus = True
can_focus_children = False
DEFAULT_CSS = """
ThreadPickerWidget {
height: auto;
max-height: 22;
margin: 1 0;
padding: 0 1;
background: $surface;
border: solid $primary;
}
ThreadPickerWidget .picker-title {
height: 1;
text-style: bold;
color: $primary;
}
ThreadPickerWidget .picker-rows {
height: auto;
max-height: 16;
overflow-y: auto;
}
ThreadPickerWidget .picker-row {
height: 1;
padding: 0 1;
}
ThreadPickerWidget .picker-row-selected {
background: $primary;
text-style: bold;
}
ThreadPickerWidget .picker-help {
height: 1;
color: $text-muted;
text-style: italic;
}
"""
BINDINGS: ClassVar[list[BindingType]] = [
Binding("up", "move_up", "Up", show=False),
Binding("k", "move_up", "Up", show=False),
Binding("down", "move_down", "Down", show=False),
Binding("j", "move_down", "Down", show=False),
Binding("enter", "select", "Select", show=False),
Binding("escape", "cancel", "Cancel", show=False),
]
class Picked(Message):
"""Posted when user picks a thread."""
def __init__(self, thread_id: str) -> None:
super().__init__()
self.thread_id = thread_id
class Cancelled(Message):
"""Posted when user cancels selection."""
def __init__(
self,
threads: list[dict],
*,
current_thread: str | None = None,
title: str = "Select a session",
**kwargs: Any,
) -> None:
super().__init__(**kwargs)
self._threads = threads
self._current_thread = current_thread
self._title = title
self._selected = 0
self._row_widgets: list[Static] = []
def compose(self) -> ComposeResult:
yield Static(self._title, classes="picker-title")
with Container(classes="picker-rows"):
for _ in self._threads:
widget = Static("", classes="picker-row")
self._row_widgets.append(widget)
yield widget
yield Static(
"\u2191/\u2193 navigate \u00b7 Enter select \u00b7 Esc cancel",
classes="picker-help",
)
def on_mount(self) -> None:
self._update_rows()
self.call_later(self.focus)
def _update_rows(self) -> None:
for i, (thread, widget) in enumerate(zip(self._threads, self._row_widgets)):
is_current = thread["thread_id"] == self._current_thread
text = build_row_text(thread, selected=(i == self._selected), current=is_current)
widget.update(text)
widget.remove_class("picker-row-selected")
if i == self._selected:
widget.add_class("picker-row-selected")
widget.scroll_visible()
def action_move_up(self) -> None:
if not self._threads:
return
self._selected = (self._selected - 1) % len(self._threads)
self._update_rows()
def action_move_down(self) -> None:
if not self._threads:
return
self._selected = (self._selected + 1) % len(self._threads)
self._update_rows()
def action_select(self) -> None:
if not self._threads:
self.post_message(self.Cancelled())
return
self.post_message(self.Picked(self._threads[self._selected]["thread_id"]))
def action_cancel(self) -> None:
self.post_message(self.Cancelled())
def on_blur(self, event: events.Blur) -> None:
"""Re-focus to keep focus trapped until decision is made."""
self.call_after_refresh(self.focus)
+204
View File
@@ -0,0 +1,204 @@
"""Tests for EvoScientist.cli.widgets.thread_selector module."""
from __future__ import annotations
from unittest import mock
from rich.text import Text
from EvoScientist.cli.widgets.thread_selector import (
ThreadPickerWidget,
build_row_text,
)
# ---------------------------------------------------------------------------
# Sample data
# ---------------------------------------------------------------------------
_THREADS = [
{
"thread_id": "abc12345",
"preview": "Help me write a paper",
"message_count": 12,
"model": "claude-sonnet-4-6",
"updated_at": "2026-03-09T10:00:00+00:00",
},
{
"thread_id": "def67890",
"preview": "Run experiment pipeline",
"message_count": 5,
"model": "gpt-4o",
"updated_at": "2026-03-08T08:00:00+00:00",
},
{
"thread_id": "ghi11111",
"preview": "",
"message_count": 0,
"model": "",
"updated_at": None,
},
]
# ---------------------------------------------------------------------------
# build_row_text unit tests (pure function, no Textual app needed)
# ---------------------------------------------------------------------------
class TestBuildRowText:
def test_selected_has_cursor(self):
text = build_row_text(_THREADS[0], selected=True)
assert isinstance(text, Text)
assert "\u25b8" in text.plain
def test_not_selected_no_cursor(self):
text = build_row_text(_THREADS[0], selected=False)
assert "\u25b8" not in text.plain
def test_thread_id_shown(self):
text = build_row_text(_THREADS[0])
assert "abc12345" in text.plain
def test_current_marker(self):
text = build_row_text(_THREADS[0], current=True)
assert "*" in text.plain
def test_no_current_marker(self):
text = build_row_text(_THREADS[0], current=False)
assert " *" not in text.plain
def test_preview_shown(self):
text = build_row_text(_THREADS[0])
assert "Help me write a paper" in text.plain
def test_no_preview(self):
text = build_row_text(_THREADS[2])
assert "ghi11111" in text.plain
assert "(0 msgs)" in text.plain
def test_message_count(self):
text = build_row_text(_THREADS[0])
assert "(12 msgs)" in text.plain
def test_model_shown(self):
text = build_row_text(_THREADS[0])
assert "claude-sonnet-4-6" in text.plain
def test_empty_model_omitted(self):
text = build_row_text(_THREADS[2])
plain = text.plain
assert "ghi11111" in plain
def test_long_preview_truncated(self):
long_thread = {**_THREADS[0], "preview": "A" * 60}
text = build_row_text(long_thread)
assert "\u2026" in text.plain
assert "A" * 60 not in text.plain
def test_relative_time_shown(self):
text = build_row_text(_THREADS[0])
assert "abc12345" in text.plain
def test_no_time_for_none_updated_at(self):
text = build_row_text(_THREADS[2])
assert "ghi11111" in text.plain
# ---------------------------------------------------------------------------
# ThreadPickerWidget unit tests
# ---------------------------------------------------------------------------
class TestThreadPickerWidget:
def test_init_stores_threads(self):
picker = ThreadPickerWidget(_THREADS, current_thread="abc12345")
assert picker._threads == _THREADS
assert picker._current_thread == "abc12345"
assert picker._selected == 0
def test_init_empty_threads(self):
picker = ThreadPickerWidget([])
assert picker._threads == []
def test_custom_title(self):
picker = ThreadPickerWidget(_THREADS, title="Pick one")
assert picker._title == "Pick one"
def test_default_title(self):
picker = ThreadPickerWidget(_THREADS)
assert picker._title == "Select a session"
def test_action_move_down_wraps(self):
picker = ThreadPickerWidget(_THREADS)
picker._row_widgets = [mock.MagicMock() for _ in _THREADS]
picker._selected = len(_THREADS) - 1
# Mock _update_rows to avoid Textual update calls
picker._update_rows = mock.MagicMock()
picker.action_move_down()
assert picker._selected == 0
def test_action_move_up_wraps(self):
picker = ThreadPickerWidget(_THREADS)
picker._row_widgets = [mock.MagicMock() for _ in _THREADS]
picker._selected = 0
picker._update_rows = mock.MagicMock()
picker.action_move_up()
assert picker._selected == len(_THREADS) - 1
def test_action_move_down_increments(self):
picker = ThreadPickerWidget(_THREADS)
picker._row_widgets = [mock.MagicMock() for _ in _THREADS]
picker._selected = 0
picker._update_rows = mock.MagicMock()
picker.action_move_down()
assert picker._selected == 1
def test_action_move_up_decrements(self):
picker = ThreadPickerWidget(_THREADS)
picker._row_widgets = [mock.MagicMock() for _ in _THREADS]
picker._selected = 2
picker._update_rows = mock.MagicMock()
picker.action_move_up()
assert picker._selected == 1
def test_action_move_empty_noop(self):
picker = ThreadPickerWidget([])
picker.action_move_down()
assert picker._selected == 0
picker.action_move_up()
assert picker._selected == 0
def test_action_select_posts_picked(self):
picker = ThreadPickerWidget(_THREADS)
picker._selected = 1
picker.post_message = mock.MagicMock()
picker.action_select()
msg = picker.post_message.call_args[0][0]
assert isinstance(msg, ThreadPickerWidget.Picked)
assert msg.thread_id == "def67890"
def test_action_cancel_posts_cancelled(self):
picker = ThreadPickerWidget(_THREADS)
picker.post_message = mock.MagicMock()
picker.action_cancel()
msg = picker.post_message.call_args[0][0]
assert isinstance(msg, ThreadPickerWidget.Cancelled)
def test_action_select_empty_posts_cancelled(self):
picker = ThreadPickerWidget([])
picker.post_message = mock.MagicMock()
picker.action_select()
msg = picker.post_message.call_args[0][0]
assert isinstance(msg, ThreadPickerWidget.Cancelled)
def test_bindings_include_navigation(self):
bindings = {b.key for b in ThreadPickerWidget.BINDINGS}
assert "up" in bindings
assert "down" in bindings
assert "enter" in bindings
assert "escape" in bindings
assert "k" in bindings
assert "j" in bindings
def test_can_focus(self):
assert ThreadPickerWidget.can_focus is True
assert ThreadPickerWidget.can_focus_children is False