From aa3dd004099e3933e5cc1c45437d8da95f84bfe2 Mon Sep 17 00:00:00 2001 From: Xi Zhang <106144707+X-iZhang@users.noreply.github.com> Date: Tue, 21 Apr 2026 22:26:40 +0200 Subject: [PATCH] =?UTF-8?q?feat:=20add=20support=20for=20session=20resumpt?= =?UTF-8?q?ion=20with=20--resume=20flag=20and=20enhan=E2=80=A6=20(#170)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * feat: add support for session resumption with --resume flag and enhance thread ID resolution * refactor(tests): streamline help output testing for --resume flag * feat: enhance session resume functionality with improved thread ID resolution and SQL wildcard handling * feat: improve error handling for resume hint retrieval in interactive modes * refactor: streamline logging for print_resume_hint failure in interactive mode * feat: implement deferred scrolling for Markdown-heavy content in interactive mode --- EvoScientist/cli/commands.py | 32 +++++- EvoScientist/cli/interactive.py | 50 +++++--- EvoScientist/cli/resume_hint.py | 19 ++++ EvoScientist/cli/tui_interactive.py | 107 ++++++++++++------ .../commands/implementation/general.py | 4 - EvoScientist/sessions.py | 27 ++++- tests/test_cli_resume_flag.py | 35 ++++++ tests/test_resume_hint.py | 35 ++++++ tests/test_sessions.py | 28 +++++ 9 files changed, 277 insertions(+), 60 deletions(-) create mode 100644 EvoScientist/cli/resume_hint.py create mode 100644 tests/test_cli_resume_flag.py create mode 100644 tests/test_resume_hint.py diff --git a/EvoScientist/cli/commands.py b/EvoScientist/cli/commands.py index c40e277..309804b 100644 --- a/EvoScientist/cli/commands.py +++ b/EvoScientist/cli/commands.py @@ -1088,7 +1088,10 @@ def _main_callback( None, "-p", "--prompt", help="Query to execute (single-shot mode)" ), thread_id: str | None = typer.Option( - None, "--thread-id", help="Thread ID for conversation persistence" + None, + "--resume", + "--thread-id", + help="Thread ID (or prefix) to resume a previous session.", ), workdir: str | None = typer.Option( None, "--workdir", help="Override workspace directory for this session" @@ -1278,17 +1281,40 @@ def _main_callback( # Single-shot mode: wrap in persistent checkpointer import asyncio - from ..sessions import generate_thread_id, get_checkpointer + from ..sessions import ( + generate_thread_id, + get_checkpointer, + resolve_thread_id_prefix, + ) async def _single_shot(): async with get_checkpointer() as checkpointer: + # Resolve resume target first so a bad --resume/--thread-id + # exits before the slow _load_agent() provider setup. + if thread_id: + resolved, matches = await resolve_thread_id_prefix(thread_id) + if resolved: + tid = resolved + elif matches: + console.print( + f"[yellow]Ambiguous thread ID '{escape(thread_id)}'. Matches:[/yellow]" + ) + for s in matches: + console.print(f" [cyan]{escape(s)}[/cyan]") + raise typer.Exit(1) + else: + console.print( + f"[red]Thread '{escape(thread_id)}' not found.[/red]" + ) + raise typer.Exit(1) + else: + tid = generate_thread_id() console.print("[dim]Loading agent...[/dim]") agent = _load_agent( workspace_dir=workspace_dir, checkpointer=checkpointer, config=config, ) - tid = thread_id or generate_thread_id() cmd_run( agent, prompt, diff --git a/EvoScientist/cli/interactive.py b/EvoScientist/cli/interactive.py index 90d5a79..28b725d 100644 --- a/EvoScientist/cli/interactive.py +++ b/EvoScientist/cli/interactive.py @@ -33,12 +33,12 @@ import EvoScientist.cli.channel as _ch_mod from ..sessions import ( _format_relative_time, delete_thread, - find_similar_threads, generate_thread_id, get_checkpointer, get_thread_messages, get_thread_metadata, list_threads, + resolve_thread_id_prefix, thread_exists, ) from ..stream.display import console @@ -432,16 +432,14 @@ def cmd_interactive( async def _resolve_thread_id(tid: str) -> str | None: """Resolve a (possibly partial) thread ID. Returns full ID or None.""" - if await thread_exists(tid): - return tid - similar = await find_similar_threads(tid) - if len(similar) == 1: - return similar[0] - if len(similar) > 1: + resolved, matches = await resolve_thread_id_prefix(tid) + if resolved: + return resolved + if matches: console.print( f"[yellow]Ambiguous thread ID '{escape(tid)}'. Matches:[/yellow]" ) - for s in similar: + for s in matches: console.print(f" [cyan]{s}[/cyan]") return None console.print(f"[red]Thread '{escape(tid)}' not found.[/red]") @@ -686,6 +684,12 @@ def cmd_interactive( state["status_last_input_tokens"] = None if ws: state["workspace_dir"] = ws + else: + # Resolution failed (ambiguous/not-found); the user's raw + # input is still seeded in state["thread_id"] from init. + # Replace with a fresh ID so a new session isn't + # checkpointed under the bad prefix. + state["thread_id"] = generate_thread_id() console.print("[dim]Loading agent...[/dim]") state["agent"] = _load_agent( @@ -709,6 +713,7 @@ def cmd_interactive( console.print( f"[green]Resumed session [yellow]{state['thread_id']}[/yellow][/green]\n" ) + await _render_history(state["thread_id"]) else: print_banner( state["thread_id"], @@ -937,7 +942,6 @@ def cmd_interactive( # Special commands if user_input.lower() in ("/exit", "/quit", "/q"): - console.print("[dim]Goodbye![/dim]") state["running"] = False break @@ -997,9 +1001,6 @@ def cmd_interactive( console.print( f"[dim]Memory dir:[/dim] [cyan]{_shorten_path(memory_dir)}[/cyan]" ) - console.print( - f"[dim]UI:[/dim] [cyan]{state['ui_backend']}[/cyan]" - ) console.print() continue @@ -1107,12 +1108,12 @@ def cmd_interactive( _print_separator() except KeyboardInterrupt: - console.print("\n[dim]Goodbye![/dim]") + console.print() state["running"] = False break except EOFError: # Handle Ctrl+D - console.print("\n[dim]Goodbye![/dim]") + console.print() state["running"] = False break except Exception as e: @@ -1135,12 +1136,31 @@ def cmd_interactive( await queue_task except asyncio.CancelledError: pass + # Best-effort: guard so a DB lookup failure here can't + # shadow the original exception exiting _async_main_loop. + current_tid = state.get("thread_id") + if current_tid: + try: + if await thread_exists(current_tid): + state["resume_hint_thread_id"] = current_tid + except Exception: + _channel_logger.debug( + "resume-hint thread_exists lookup failed", + exc_info=True, + ) # Run the async main loop + from .resume_hint import print_resume_hint + try: asyncio.run(_async_main_loop()) except KeyboardInterrupt: - console.print("\n[dim]Goodbye![/dim]") + console.print() + finally: + try: + print_resume_hint(state.get("resume_hint_thread_id"), console=console) + except Exception: + _channel_logger.debug("print_resume_hint failed", exc_info=True) def cmd_run( diff --git a/EvoScientist/cli/resume_hint.py b/EvoScientist/cli/resume_hint.py new file mode 100644 index 0000000..8ec69c5 --- /dev/null +++ b/EvoScientist/cli/resume_hint.py @@ -0,0 +1,19 @@ +"""Helper for printing the session-exit Goodbye message and resume hint.""" + +from __future__ import annotations + +from rich.console import Console +from rich.markup import escape + + +def print_resume_hint( + thread_id: str | None, + console: Console | None = None, +) -> None: + """Print ``Goodbye!`` and, when available, a resume hint for *thread_id*.""" + out = console or Console() + out.print("[dim]Goodbye![/dim]") + if thread_id: + out.print() + out.print("[dim]Resume this session with:[/dim]") + out.print(f"[cyan]EvoSci --resume {escape(thread_id)}[/cyan]") diff --git a/EvoScientist/cli/tui_interactive.py b/EvoScientist/cli/tui_interactive.py index fb8968b..2b0fde9 100644 --- a/EvoScientist/cli/tui_interactive.py +++ b/EvoScientist/cli/tui_interactive.py @@ -24,11 +24,11 @@ from ..commands import CommandContext from ..commands import manager as cmd_manager from ..paths import DATA_DIR from ..sessions import ( - find_similar_threads, generate_thread_id, get_checkpointer, get_thread_messages, get_thread_metadata, + resolve_thread_id_prefix, thread_exists, ) from ..stream.events import stream_agent_events @@ -391,7 +391,7 @@ def run_textual_interactive( title=title, ) await container.mount(picker) - container.scroll_end(animate=False) + self._schedule_scroll_to_bottom(container, delays=()) picker.focus() return await self._wait_for_thread_pick(picker) @@ -408,7 +408,7 @@ def run_textual_interactive( pre_filter_tag=pre_filter_tag, ) await container.mount(browser) - container.scroll_end(animate=False) + self._schedule_scroll_to_bottom(container, delays=()) browser.focus() return await self._wait_for_skill_browse(browser) @@ -425,7 +425,7 @@ def run_textual_interactive( pre_filter_tag=pre_filter_tag, ) await container.mount(browser) - container.scroll_end(animate=False) + self._schedule_scroll_to_bottom(container, delays=()) browser.focus() return await self._wait_for_mcp_browse(browser) @@ -609,6 +609,32 @@ def run_textual_interactive( # ── Widget helpers ───────────────────────────────────── + def _schedule_scroll_to_bottom( + self, + container: VerticalScroll, + *, + delays: tuple[float, ...] = (0.3, 0.8), + immediate: bool = True, + ) -> None: + """Schedule deferred scrolls so the viewport lands at the bottom. + + Markdown- and list-heavy widgets lay out across multiple refresh + cycles, so a single ``scroll_end()`` may fire against a stale + ``virtual_size`` and leave the viewport mid-content. Re-schedule + ``scroll_end`` at each delay to follow subsequent reflows. + """ + if immediate: + self.call_after_refresh( + lambda: container.scroll_end(animate=False), + ) + for delay in delays: + self.set_timer( + delay, + lambda: self.call_after_refresh( + lambda: container.scroll_end(animate=False), + ), + ) + def _append_system(self, text: str, style: str = "dim") -> None: """Mount a SystemMessage widget into #chat.""" container = self.query_one("#chat", VerticalScroll) @@ -1405,17 +1431,11 @@ def run_textual_interactive( clean or state.response_text ) await container.mount(assistant_w) - # Markdown rendering is async and needs multiple - # layout cycles to compute final height. Schedule - # repeated deferred scrolls so long content stays - # visible even when Markdown takes time to lay out. - for delay in (0.15, 0.4, 0.8, 1.5): - self.set_timer( - delay, - lambda: self.call_after_refresh( - lambda: container.scroll_end(animate=False), - ), - ) + self._schedule_scroll_to_bottom( + container, + delays=(0.15, 0.4, 0.8, 1.5), + immediate=False, + ) # Mount token usage stats if state.total_input_tokens or state.total_output_tokens: await container.mount( @@ -1498,18 +1518,7 @@ def run_textual_interactive( ): on_thinking_cb(state.thinking_text.rstrip()) # Final scrolls to ensure last content is visible. - # Markdown layout is async — schedule multiple deferred - # scrolls so long content eventually scrolls into view. - self.call_after_refresh( - lambda: container.scroll_end(animate=False), - ) - for delay in (0.3, 0.8): - self.set_timer( - delay, - lambda: self.call_after_refresh( - lambda: container.scroll_end(animate=False), - ), - ) + self._schedule_scroll_to_bottom(container) # HITL / ask_user: if interrupt was handled, loop back to resume stream if state.pending_interrupt is None and state.pending_ask_user is None: @@ -2163,7 +2172,12 @@ def run_textual_interactive( await container.mount( SystemMessage("── End of history ──", msg_style="dim") ) - container.scroll_end(animate=False) + # History can hold dozens of Markdown-heavy AssistantMessages + # whose async layout keeps growing virtual_size for several + # seconds; schedule enough retries to catch the final reflow. + self._schedule_scroll_to_bottom( + container, delays=(0.1, 0.3, 0.6, 1.0, 1.8, 3.0) + ) # ── Quit handling ────────────────────────────────────── @@ -2414,15 +2428,11 @@ def run_textual_interactive( async def _amain() -> None: async with get_checkpointer() as checkpointer: effective_workspace = workspace_dir - effective_thread_id = thread_id + effective_thread_id: str | None = None resumed = False resume_warning = "" if thread_id: - if await thread_exists(thread_id): - resolved = thread_id - else: - similar = await find_similar_threads(thread_id) - resolved = similar[0] if len(similar) == 1 else None + resolved, matches = await resolve_thread_id_prefix(thread_id) if resolved: meta = await get_thread_metadata(resolved) ws = (meta or {}).get("workspace_dir", "") @@ -2430,6 +2440,11 @@ def run_textual_interactive( effective_workspace = ws effective_thread_id = resolved resumed = True + elif matches: + resume_warning = ( + f"Thread prefix '{thread_id}' is ambiguous " + f"({', '.join(matches)}). Starting new session." + ) else: resume_warning = ( f"Thread '{thread_id}' not found. Starting new session." @@ -2450,7 +2465,29 @@ def run_textual_interactive( resumed=resumed, resume_warning=resume_warning, ) - await app.run_async() + try: + await app.run_async() + finally: + from .resume_hint import print_resume_hint + + # Best-effort resume hint — guarded so failures here (e.g. + # DB teardown race during abnormal shutdown) cannot shadow + # the original run_async traceback. + exit_tid = getattr(app, "_conversation_tid", None) + hint_tid: str | None = None + if exit_tid: + try: + if await thread_exists(exit_tid): + hint_tid = exit_tid + except Exception: + _channel_logger.debug( + "resume-hint thread_exists lookup failed", + exc_info=True, + ) + try: + print_resume_hint(hint_tid) + except Exception: + _channel_logger.debug("print_resume_hint failed", exc_info=True) import nest_asyncio # type: ignore[import-untyped] diff --git a/EvoScientist/commands/implementation/general.py b/EvoScientist/commands/implementation/general.py index 42c6b93..8680fb5 100644 --- a/EvoScientist/commands/implementation/general.py +++ b/EvoScientist/commands/implementation/general.py @@ -56,10 +56,6 @@ class CurrentCommand(Command): ctx.ui.append_system( f"Memory dir: {_shorten_path(str(memory_path))}", style="dim" ) - # How to determine UI type here? - # Maybe ctx.ui has a name or we pass it in ctx. - # For now, let's keep it simple. - ctx.ui.append_system("UI: auto", style="dim") # Register commands diff --git a/EvoScientist/sessions.py b/EvoScientist/sessions.py index 54565ca..29b9b43 100644 --- a/EvoScientist/sessions.py +++ b/EvoScientist/sessions.py @@ -323,19 +323,40 @@ async def find_similar_threads(thread_id: str, limit: int = 5) -> list[str]: async with aiosqlite.connect(db_path, timeout=30.0) as conn: if not await _table_exists(conn, "checkpoints"): return [] - query = """ + # Escape SQL LIKE wildcards so user-supplied prefixes are matched + # literally (e.g. `--resume %` must not match every thread). + escaped = ( + thread_id.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") + ) + query = r""" SELECT DISTINCT thread_id FROM checkpoints - WHERE thread_id LIKE ? + WHERE thread_id LIKE ? ESCAPE '\' AND json_extract(metadata, '$.agent_name') = ? ORDER BY thread_id LIMIT ? """ - async with conn.execute(query, (thread_id + "%", AGENT_NAME, limit)) as cur: + async with conn.execute(query, (escaped + "%", AGENT_NAME, limit)) as cur: rows = await cur.fetchall() return [r[0] for r in rows] +async def resolve_thread_id_prefix(tid: str) -> tuple[str | None, list[str]]: + """Resolve a (possibly partial) thread ID. + + Returns ``(resolved_id, matches)``: + - ``(full_id, [])`` when *tid* is an exact hit or a unique prefix. + - ``(None, [a, b, ...])`` when the prefix is ambiguous (multiple matches). + - ``(None, [])`` when no thread matches. + """ + if await thread_exists(tid): + return tid, [] + similar = await find_similar_threads(tid) + if len(similar) == 1: + return similar[0], [] + return None, similar + + async def delete_thread(thread_id: str) -> bool: """Delete all EvoScientist checkpoints (and writes) for *thread_id*.""" db_path = str(get_db_path()) diff --git a/tests/test_cli_resume_flag.py b/tests/test_cli_resume_flag.py new file mode 100644 index 0000000..096c67e --- /dev/null +++ b/tests/test_cli_resume_flag.py @@ -0,0 +1,35 @@ +"""Tests for the ``--resume`` CLI flag (alias of ``--thread-id``).""" + +from __future__ import annotations + +import re + +from typer.testing import CliRunner + +from EvoScientist.cli._app import app + +runner = CliRunner() + +_ANSI_RE = re.compile(r"\x1b\[[0-9;]*m") + + +def _plain_help() -> str: + """Render ``--help`` and strip ANSI escapes. + + Rich inserts per-character color codes (e.g. ``\\x1b[36m-\\x1b[0m\\x1b[36m-name``), + which break literal substring lookups like ``"--resume" in stdout``. + """ + result = runner.invoke(app, ["--help"]) + assert result.exit_code == 0 + return _ANSI_RE.sub("", result.stdout) + + +def test_resume_flag_listed_in_help(): + plain = _plain_help() + assert "--resume" in plain + assert "--thread-id" in plain + + +def test_thread_id_flag_still_works(): + """Backwards compatibility: --thread-id should remain a valid flag.""" + assert "--thread-id" in _plain_help() diff --git a/tests/test_resume_hint.py b/tests/test_resume_hint.py new file mode 100644 index 0000000..f9419a3 --- /dev/null +++ b/tests/test_resume_hint.py @@ -0,0 +1,35 @@ +"""Tests for the session-exit resume hint helper.""" + +from __future__ import annotations + +import io + +from rich.console import Console + +from EvoScientist.cli.resume_hint import print_resume_hint + + +def _capture(thread_id: str | None) -> str: + buf = io.StringIO() + console = Console(file=buf, force_terminal=False, width=120, color_system=None) + print_resume_hint(thread_id, console=console) + return buf.getvalue() + + +def test_prints_goodbye_and_hint_for_thread_id(): + output = _capture("365cf731") + assert "Goodbye!" in output + assert "Resume this session with:" in output + assert "EvoSci --resume 365cf731" in output + + +def test_none_thread_id_prints_only_goodbye(): + output = _capture(None) + assert "Goodbye!" in output + assert "Resume" not in output + + +def test_empty_thread_id_prints_only_goodbye(): + output = _capture("") + assert "Goodbye!" in output + assert "Resume" not in output diff --git a/tests/test_sessions.py b/tests/test_sessions.py index 25d40b0..02df03a 100644 --- a/tests/test_sessions.py +++ b/tests/test_sessions.py @@ -21,6 +21,7 @@ from EvoScientist.sessions import ( get_thread_messages, get_thread_metadata, list_threads, + resolve_thread_id_prefix, thread_exists, ) from tests.conftest import run_async as _run @@ -210,6 +211,33 @@ class TestThreadFunctions(unittest.TestCase): similar = _run(find_similar_threads("xyz")) assert len(similar) == 0 + def test_resolve_prefix_exact_match(self): + resolved, matches = _run(resolve_thread_id_prefix("abc12345")) + assert resolved == "abc12345" + assert matches == [] + + def test_resolve_prefix_unique_prefix(self): + resolved, matches = _run(resolve_thread_id_prefix("def00")) + assert resolved == "def00001" + assert matches == [] + + def test_resolve_prefix_ambiguous(self): + resolved, matches = _run(resolve_thread_id_prefix("abc1")) + assert resolved is None + assert set(matches) == {"abc12345", "abc12399"} + + def test_resolve_prefix_not_found(self): + resolved, matches = _run(resolve_thread_id_prefix("zzz")) + assert resolved is None + assert matches == [] + + def test_find_similar_escapes_sql_wildcards(self): + # '%' / '_' must be treated as literal characters, not SQL LIKE + # wildcards, so a prefix that doesn't occur verbatim returns nothing + # (prior buggy behavior: '%' matched every thread). + assert _run(find_similar_threads("%")) == [] + assert _run(find_similar_threads("_")) == [] + def test_get_most_recent(self): recent = _run(get_most_recent()) assert recent is not None