feat: add support for session resumption with --resume flag and enhan… (#170)
* 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
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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]")
|
||||
@@ -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]
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user