From f29b254d9a0178e62100e802dd384f0440e430ec Mon Sep 17 00:00:00 2001 From: X-iZhang Date: Wed, 11 Feb 2026 23:34:25 +0000 Subject: [PATCH] Implement session persistence and management features - Introduced a new `sessions.py` module for handling session persistence using SQLite. - Added CRUD operations for threads, including listing, checking existence, finding similar threads, and deleting threads. - Enhanced the interactive CLI to support commands for managing sessions: `/current`, `/threads`, `/resume`, and `/delete`. - Updated the `cmd_interactive` function to handle session metadata and improve user experience with session history rendering. - Modified the streaming functions to include metadata for checkpoint persistence. - Added unit tests for session management functionalities to ensure reliability and correctness. --- EvoScientist/__init__.py | 6 +- EvoScientist/cli/agent.py | 8 +- EvoScientist/cli/commands.py | 28 +- EvoScientist/cli/interactive.py | 496 ++++++++++++++++++++++++-------- EvoScientist/sessions.py | 315 ++++++++++++++++++++ EvoScientist/stream/display.py | 5 +- EvoScientist/stream/events.py | 13 +- pyproject.toml | 1 + tests/test_sessions.py | 217 ++++++++++++++ 9 files changed, 957 insertions(+), 132 deletions(-) create mode 100644 EvoScientist/sessions.py create mode 100644 tests/test_sessions.py diff --git a/EvoScientist/__init__.py b/EvoScientist/__init__.py index 5b64d59..c581660 100644 --- a/EvoScientist/__init__.py +++ b/EvoScientist/__init__.py @@ -36,7 +36,11 @@ _EXPORTS: dict[str, tuple[str, str]] = { # Tools "tavily_search": (".tools", "tavily_search"), "think_tool": (".tools", "think_tool"), - + # Sessions + "get_checkpointer": (".sessions", "get_checkpointer"), + "generate_thread_id": (".sessions", "generate_thread_id"), + "list_threads": (".sessions", "list_threads"), + "delete_thread": (".sessions", "delete_thread"), } diff --git a/EvoScientist/cli/agent.py b/EvoScientist/cli/agent.py index 6cb076a..9029194 100644 --- a/EvoScientist/cli/agent.py +++ b/EvoScientist/cli/agent.py @@ -48,11 +48,13 @@ def _create_session_workspace(name: str | None = None) -> str: return workspace_dir -def _load_agent(workspace_dir: str | None = None): - """Load the CLI agent (with InMemorySaver checkpointer for multi-turn). +def _load_agent(workspace_dir: str | None = None, checkpointer=None): + """Load the CLI agent with optional persistent checkpointer. Args: workspace_dir: Optional per-session workspace directory. + checkpointer: Optional LangGraph checkpointer (e.g. ``AsyncSqliteSaver``). + Falls back to ``InMemorySaver`` when ``None``. """ from ..EvoScientist import create_cli_agent - return create_cli_agent(workspace_dir=workspace_dir) + return create_cli_agent(workspace_dir=workspace_dir, checkpointer=checkpointer) diff --git a/EvoScientist/cli/commands.py b/EvoScientist/cli/commands.py index 4e207f8..b1c40cb 100644 --- a/EvoScientist/cli/commands.py +++ b/EvoScientist/cli/commands.py @@ -406,17 +406,24 @@ def _main_callback( os.makedirs(workspace_dir, exist_ok=True) workspace_fixed = True - # Load agent with session workspace - console.print("[dim]Loading agent...[/dim]") - agent = _load_agent(workspace_dir=workspace_dir) - if prompt: - # Single-shot mode: execute query and exit - cmd_run(agent, prompt, thread_id=thread_id, show_thinking=show_thinking, workspace_dir=workspace_dir) + # Single-shot mode: wrap in persistent checkpointer + import asyncio + from ..sessions import get_checkpointer, generate_thread_id + + async def _single_shot(): + async with get_checkpointer() as checkpointer: + console.print("[dim]Loading agent...[/dim]") + agent = _load_agent(workspace_dir=workspace_dir, checkpointer=checkpointer) + tid = thread_id or generate_thread_id() + cmd_run(agent, prompt, thread_id=tid, show_thinking=show_thinking, workspace_dir=workspace_dir) + + import nest_asyncio # type: ignore[import-untyped] + nest_asyncio.apply() + asyncio.get_event_loop().run_until_complete(_single_shot()) else: - # Interactive mode (default) + # Interactive mode (default) — checkpointer managed inside cmd_interactive cmd_interactive( - agent, show_thinking=show_thinking, workspace_dir=workspace_dir, workspace_fixed=workspace_fixed, @@ -427,6 +434,7 @@ def _main_callback( imessage_allowed_senders=config.imessage_allowed_senders, imessage_send_thinking=config.imessage_send_thinking, run_name=name, + thread_id=thread_id, ) @@ -456,3 +464,7 @@ def _configure_logging(): root_logger.removeHandler(h) root_logger.addHandler(handler) root_logger.setLevel(logging.WARNING) + + # Suppress noisy schema warnings from langchain_google_genai + # (e.g. "Key '$schema' is not supported in schema, ignoring") + logging.getLogger("langchain_google_genai._function_utils").setLevel(logging.ERROR) diff --git a/EvoScientist/cli/interactive.py b/EvoScientist/cli/interactive.py index 15a7523..f99a304 100644 --- a/EvoScientist/cli/interactive.py +++ b/EvoScientist/cli/interactive.py @@ -4,7 +4,7 @@ import asyncio import os import queue import sys -import uuid +from datetime import datetime, timezone from typing import Any import typer # type: ignore[import-untyped] @@ -15,8 +15,21 @@ from prompt_toolkit.auto_suggest import AutoSuggestFromHistory # type: ignore[i from prompt_toolkit.formatted_text import HTML # type: ignore[import-untyped] from prompt_toolkit.shortcuts import CompleteStyle # type: ignore[import-untyped] from prompt_toolkit.styles import Style as PtStyle # type: ignore[import-untyped] +from rich.table import Table from rich.text import Text +from ..sessions import ( + generate_thread_id, + get_checkpointer, + list_threads, + thread_exists, + find_similar_threads, + delete_thread, + get_thread_metadata, + get_thread_messages, + _format_relative_time, + AGENT_NAME, +) from ..stream.display import console, _run_streaming from .agent import _shorten_path, _create_session_workspace, _load_agent from .channel import ( @@ -90,7 +103,10 @@ def print_banner( # ============================================================================= _SLASH_COMMANDS = [ - ("/thread", "Show thread ID, workspace & memory dir"), + ("/current", "Show current session info"), + ("/threads", "List recent sessions"), + ("/resume", "Resume a previous session (prefix match)"), + ("/delete", "Delete a saved session"), ("/new", "Start a new session"), ("/skills", "List installed skills"), ("/install-skill", "Add a skill from path or GitHub"), @@ -110,6 +126,17 @@ _COMPLETION_STYLE = PtStyle.from_dict({ "scrollbar.button": "bg:default", }) +# Style for questionary pickers — matches _COMPLETION_STYLE visual language: +# gray (#888888) for non-selected, bold for selected, no background changes. +_PICKER_STYLE = PtStyle.from_dict({ + "questionmark": "#888888", + "question": "", + "pointer": "bold", + "highlighted": "bold", + "text": "#888888", + "answer": "bold", +}) + class SlashCommandCompleter(Completer): """Autocomplete for slash commands — triggers when input starts with '/'.""" @@ -133,8 +160,17 @@ class SlashCommandCompleter(Completer): # ============================================================================= +def _build_metadata(workspace_dir: str | None, model: str | None) -> dict: + """Build metadata dict for LangGraph checkpoint persistence.""" + return { + "agent_name": AGENT_NAME, + "updated_at": datetime.now(timezone.utc).isoformat(), + "workspace_dir": workspace_dir or "", + "model": model or "", + } + + def cmd_interactive( - agent: Any, show_thinking: bool = True, workspace_dir: str | None = None, workspace_fixed: bool = False, @@ -145,11 +181,14 @@ def cmd_interactive( imessage_allowed_senders: str = "", imessage_send_thinking: bool = True, run_name: str | None = None, + thread_id: str | None = None, ) -> None: """Interactive conversation mode with streaming output. + The persistent ``AsyncSqliteSaver`` checkpointer is opened here and + shared for the entire interactive session lifetime. + Args: - agent: Compiled agent graph show_thinking: Whether to display thinking panels workspace_dir: Per-session workspace directory path workspace_fixed: If True, /new keeps the same workspace directory @@ -160,14 +199,13 @@ def cmd_interactive( imessage_allowed_senders: Comma-separated allowed senders imessage_send_thinking: Whether to forward thinking to channel run_name: Optional run name for /new session deduplication + thread_id: Optional thread ID to resume a previous session """ import nest_asyncio nest_asyncio.apply() - thread_id = str(uuid.uuid4()) from ..EvoScientist import MEMORY_DIR memory_dir = MEMORY_DIR - print_banner(thread_id, workspace_dir, memory_dir, mode, model, provider) history_file = str(os.path.expanduser("~/.EvoScientist_history")) session = PromptSession( @@ -185,11 +223,12 @@ def cmd_interactive( console.print(Text("\u2500" * width, style="dim")) # Mutable state for async loop - state = { - "agent": agent, - "thread_id": thread_id, + state: dict[str, Any] = { + "agent": None, + "thread_id": thread_id or generate_thread_id(), "workspace_dir": workspace_dir, "running": True, + "resumed": False, } def _process_channel_message(msg: ChannelMessage) -> None: @@ -259,11 +298,12 @@ def cmd_interactive( on_file_write = _send_file try: + meta = _build_metadata(state["workspace_dir"], model) # Use SAME _run_streaming as CLI input — full Live experience response_text = _run_streaming( state["agent"], msg.content, state["thread_id"], show_thinking, interactive=True, on_thinking=on_thinking, on_todo=on_todo, - on_file_write=on_file_write, + on_file_write=on_file_write, metadata=meta, ) # Set response for channel handler to retrieve @@ -291,118 +331,331 @@ def cmd_interactive( pass await asyncio.sleep(0.1) # Check every 100ms + 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: + console.print(f"[yellow]Ambiguous thread ID '{tid}'. Matches:[/yellow]") + for s in similar: + console.print(f" [cyan]{s}[/cyan]") + return None + console.print(f"[red]Thread '{tid}' not found.[/red]") + return None + + async def _cmd_threads(): + """Handle /threads command — show recent sessions.""" + threads = await list_threads( + limit=20, include_message_count=True, include_preview=True, + ) + if not threads: + console.print("[yellow]No saved sessions.[/yellow]") + return + table = Table(title="Sessions", show_header=True, header_style="bold cyan") + table.add_column("ID", style="bold") + table.add_column("Preview", style="dim", max_width=50, no_wrap=True) + table.add_column("Messages", justify="right") + table.add_column("Model", style="dim") + table.add_column("Last Used", style="dim") + for t in threads: + tid = t["thread_id"] + marker = " *" if tid == state["thread_id"] else "" + table.add_row( + f"{tid}{marker}", + t.get("preview", "") or "", + str(t.get("message_count", 0)), + t.get("model", "") or "", + _format_relative_time(t.get("updated_at")), + ) + console.print() + console.print(table) + console.print("[dim] /resume[/dim] to continue a session [dim]/delete [/dim] to remove [dim]/new[/dim] to start fresh") + console.print() + + async def _render_history(thread_id: str): + """Display a compact conversation history for a resumed session.""" + messages = await get_thread_messages(thread_id) + if not messages: + return + + MAX_CONTENT_LEN = 200 # truncate long messages + + def _truncate(text: str) -> str: + text = text.strip() + if len(text) <= MAX_CONTENT_LEN: + return text + return text[:MAX_CONTENT_LEN] + "..." + + console.print("[dim]── Conversation history ──[/dim]") + for msg in messages: + msg_type = getattr(msg, "type", None) + content = getattr(msg, "content", "") or "" + # content can be a list of blocks (multimodal) — extract text + if isinstance(content, list): + parts = [b.get("text", "") for b in content if isinstance(b, dict) and b.get("type") == "text"] + content = " ".join(parts) if parts else "" + + if msg_type == "human": + console.print(Text.assemble( + ("\u276f ", "bold blue"), + (_truncate(content), ""), + )) + elif msg_type == "ai": + tool_calls = getattr(msg, "tool_calls", None) or [] + if content: + console.print(Text(_truncate(content), style="dim")) + if tool_calls: + names = [tc.get("name", "?") for tc in tool_calls] + console.print(Text( + f" \u25b6 {', '.join(names)}", + style="dim italic", + )) + # Skip tool messages — they are verbose and not useful in replay + + console.print("[dim]── End of history ──[/dim]") + console.print() + + async def _cmd_resume(arg: str, checkpointer): + """Handle /resume [id] — resume a previous session.""" + if not arg: + # Show interactive session picker with conversation previews + threads = await list_threads( + limit=10, include_message_count=True, include_preview=True, + ) + if not threads: + console.print("[yellow]No sessions to resume.[/yellow]") + return + + import questionary + + choices = [] + # Display-width-aware padding (CJK chars take 2 columns) + import unicodedata + + def _display_width(s: str) -> int: + w = 0 + for ch in s: + w += 2 if unicodedata.east_asian_width(ch) in ("W", "F") else 1 + return w + + def _pad_to_width(s: str, target: int) -> str: + pad = target - _display_width(s) + return s + " " * max(pad, 2) + + lefts = [t.get("preview", "") or t["thread_id"] for t in threads] + col_width = max(_display_width(s) for s in lefts) + 4 + for t, left_text in zip(threads, lefts): + tid = t["thread_id"] + when = _format_relative_time(t.get("updated_at")) + label = f"{_pad_to_width(left_text, col_width)}({tid} {when})" + choices.append(questionary.Choice(title=label, value=tid)) + + selected = questionary.select( + "Select session to resume:", + choices=choices, + style=_PICKER_STYLE, + ).ask() + + if selected is None: + return + arg = selected + + resolved = await _resolve_thread_id(arg) + if not resolved: + return + + meta = await get_thread_metadata(resolved) + ws = (meta or {}).get("workspace_dir", "") or state["workspace_dir"] + + state["thread_id"] = resolved + state["resumed"] = True + if ws: + state["workspace_dir"] = ws + console.print("[dim]Loading session...[/dim]") + state["agent"] = _load_agent(workspace_dir=state["workspace_dir"], checkpointer=checkpointer) + # Sync shared refs if channel is running + if _ChannelState.is_running(): + _ChannelState.agent = state["agent"] + _ChannelState.thread_id = state["thread_id"] + console.print(f"[green]Resumed session:[/green] [yellow]{resolved}[/yellow]") + if state["workspace_dir"]: + console.print(f"[dim]Workspace:[/dim] [cyan]{_shorten_path(state['workspace_dir'])}[/cyan]") + console.print() + await _render_history(resolved) + + async def _cmd_delete(arg: str): + """Handle /delete — delete a saved session.""" + if not arg: + console.print("[red]Usage: /delete [/red]") + return + resolved = await _resolve_thread_id(arg) + if not resolved: + return + if resolved == state["thread_id"]: + console.print("[red]Cannot delete the current session.[/red]") + return + deleted = await delete_thread(resolved) + if deleted: + console.print(f"[green]Deleted session {resolved}.[/green]") + else: + console.print(f"[red]Session {resolved} not found.[/red]") + async def _async_main_loop(): """Async main loop with prompt_async and channel queue checking.""" - # Start background queue checker - queue_task = asyncio.create_task(_check_channel_queue()) + async with get_checkpointer() as checkpointer: + # Handle --thread-id resume + if thread_id: + resolved = await _resolve_thread_id(thread_id) + if resolved: + meta = await get_thread_metadata(resolved) + ws = (meta or {}).get("workspace_dir", "") or state["workspace_dir"] + state["thread_id"] = resolved + state["resumed"] = True + if ws: + state["workspace_dir"] = ws - # Auto-start iMessage channel if enabled in config - if imessage_enabled and not _ChannelState.is_running(): - _auto_start_channel(state["agent"], state["thread_id"], imessage_allowed_senders, imessage_send_thinking) + console.print("[dim]Loading agent...[/dim]") + state["agent"] = _load_agent(workspace_dir=state["workspace_dir"], checkpointer=checkpointer) - try: - _print_separator() - while state["running"]: - try: - user_input = await session.prompt_async( - HTML('\u276f ') - ) - user_input = user_input.strip() + # Print banner + if state["resumed"]: + print_banner(state["thread_id"], state["workspace_dir"], memory_dir, mode, model, provider) + console.print(f"[green]Resumed session [yellow]{state['thread_id']}[/yellow][/green]\n") + else: + print_banner(state["thread_id"], state["workspace_dir"], memory_dir, mode, model, provider) - if not user_input: - # Erase the empty prompt line so it looks like nothing happened - sys.stdout.write("\033[A\033[2K\r") - sys.stdout.flush() - continue + # Start background queue checker + queue_task = asyncio.create_task(_check_channel_queue()) - _print_separator() + # Auto-start iMessage channel if enabled in config + if imessage_enabled and not _ChannelState.is_running(): + _auto_start_channel(state["agent"], state["thread_id"], imessage_allowed_senders, imessage_send_thinking) - # Special commands - if user_input.lower() in ("/exit", "/quit", "/q"): - console.print("[dim]Goodbye![/dim]") - state["running"] = False - break - - if user_input.lower() == "/new": - # New session: new thread; workspace only changes if not fixed - if not workspace_fixed: - state["workspace_dir"] = _create_session_workspace(run_name) - console.print("[dim]Loading new session...[/dim]") - state["agent"] = _load_agent(workspace_dir=state["workspace_dir"]) - state["thread_id"] = str(uuid.uuid4()) - # Sync shared refs if channel is running - if _ChannelState.is_running(): - _ChannelState.agent = state["agent"] - _ChannelState.thread_id = state["thread_id"] - console.print(f"[green]New session:[/green] [yellow]{state['thread_id']}[/yellow]") - if state["workspace_dir"]: - console.print(f"[dim]Workspace:[/dim] [cyan]{_shorten_path(state['workspace_dir'])}[/cyan]\n") - continue - - if user_input.lower() == "/thread": - console.print(f"[dim]Thread:[/dim] [yellow]{state['thread_id']}[/yellow]") - if state["workspace_dir"]: - console.print(f"[dim]Workspace:[/dim] [cyan]{_shorten_path(state['workspace_dir'])}[/cyan]") - if memory_dir: - console.print(f"[dim]Memory dir:[/dim] [cyan]{_shorten_path(memory_dir)}[/cyan]") - console.print() - continue - - if user_input.lower() == "/skills": - _cmd_list_skills() - continue - - if user_input.lower().startswith("/install-skill"): - source = user_input[len("/install-skill"):].strip() - _cmd_install_skill(source) - continue - - if user_input.lower().startswith("/uninstall-skill"): - name = user_input[len("/uninstall-skill"):].strip() - _cmd_uninstall_skill(name) - continue - - if user_input.lower().startswith("/mcp"): - _cmd_mcp(user_input[4:]) - continue - - if user_input.lower().startswith("/channel"): - args = user_input[len("/channel"):].strip() - if args.lower() == "stop": - _cmd_channel_stop() - else: - _cmd_channel(args, state["agent"], state["thread_id"]) - continue - - # Stream agent response - console.print() - _run_streaming(state["agent"], user_input, state["thread_id"], show_thinking, interactive=True) - _print_separator() - - except KeyboardInterrupt: - console.print("\n[dim]Goodbye![/dim]") - state["running"] = False - break - except EOFError: - # Handle Ctrl+D - console.print("\n[dim]Goodbye![/dim]") - state["running"] = False - break - except Exception as e: - error_msg = str(e) - if "authentication" in error_msg.lower() or "api_key" in error_msg.lower(): - console.print("[red]Error: API key not configured.[/red]") - console.print("[dim]Run [bold]EvoSci onboard[/bold] to set up your API key.[/dim]") - state["running"] = False - break - else: - console.print(f"[red]Error: {e}[/red]") - finally: - queue_task.cancel() try: - await queue_task - except asyncio.CancelledError: - pass + _print_separator() + while state["running"]: + try: + user_input = await session.prompt_async( + HTML('\u276f ') + ) + user_input = user_input.strip() + + if not user_input: + # Erase the empty prompt line so it looks like nothing happened + sys.stdout.write("\033[A\033[2K\r") + sys.stdout.flush() + continue + + _print_separator() + + # Special commands + if user_input.lower() in ("/exit", "/quit", "/q"): + console.print("[dim]Goodbye![/dim]") + state["running"] = False + break + + if user_input.lower() == "/threads": + await _cmd_threads() + continue + + if user_input.lower().startswith("/resume"): + arg = user_input[len("/resume"):].strip() + await _cmd_resume(arg, checkpointer) + continue + + if user_input.lower().startswith("/delete"): + arg = user_input[len("/delete"):].strip() + await _cmd_delete(arg) + continue + + if user_input.lower() == "/new": + # New session: new thread; workspace only changes if not fixed + if not workspace_fixed: + state["workspace_dir"] = _create_session_workspace(run_name) + console.print("[dim]Loading new session...[/dim]") + state["agent"] = _load_agent(workspace_dir=state["workspace_dir"], checkpointer=checkpointer) + state["thread_id"] = generate_thread_id() + state["resumed"] = False + # Sync shared refs if channel is running + if _ChannelState.is_running(): + _ChannelState.agent = state["agent"] + _ChannelState.thread_id = state["thread_id"] + console.print(f"[green]New session:[/green] [yellow]{state['thread_id']}[/yellow]") + if state["workspace_dir"]: + console.print(f"[dim]Workspace:[/dim] [cyan]{_shorten_path(state['workspace_dir'])}[/cyan]\n") + continue + + if user_input.lower() == "/current": + console.print(f"[dim]Thread:[/dim] [yellow]{state['thread_id']}[/yellow]") + if state["workspace_dir"]: + console.print(f"[dim]Workspace:[/dim] [cyan]{_shorten_path(state['workspace_dir'])}[/cyan]") + if memory_dir: + console.print(f"[dim]Memory dir:[/dim] [cyan]{_shorten_path(memory_dir)}[/cyan]") + console.print() + continue + + if user_input.lower() == "/skills": + _cmd_list_skills() + continue + + if user_input.lower().startswith("/install-skill"): + source = user_input[len("/install-skill"):].strip() + _cmd_install_skill(source) + continue + + if user_input.lower().startswith("/uninstall-skill"): + name = user_input[len("/uninstall-skill"):].strip() + _cmd_uninstall_skill(name) + continue + + if user_input.lower().startswith("/mcp"): + _cmd_mcp(user_input[4:]) + continue + + if user_input.lower().startswith("/channel"): + args = user_input[len("/channel"):].strip() + if args.lower() == "stop": + _cmd_channel_stop() + else: + _cmd_channel(args, state["agent"], state["thread_id"]) + continue + + # Stream agent response with metadata for persistence + console.print() + meta = _build_metadata(state["workspace_dir"], model) + _run_streaming( + state["agent"], user_input, state["thread_id"], + show_thinking, interactive=True, metadata=meta, + ) + _print_separator() + + except KeyboardInterrupt: + console.print("\n[dim]Goodbye![/dim]") + state["running"] = False + break + except EOFError: + # Handle Ctrl+D + console.print("\n[dim]Goodbye![/dim]") + state["running"] = False + break + except Exception as e: + error_msg = str(e) + if "authentication" in error_msg.lower() or "api_key" in error_msg.lower(): + console.print("[red]Error: API key not configured.[/red]") + console.print("[dim]Run [bold]EvoSci onboard[/bold] to set up your API key.[/dim]") + state["running"] = False + break + else: + console.print(f"[red]Error: {e}[/red]") + finally: + queue_task.cancel() + try: + await queue_task + except asyncio.CancelledError: + pass # Run the async main loop try: @@ -411,7 +664,14 @@ def cmd_interactive( console.print("\n[dim]Goodbye![/dim]") -def cmd_run(agent: Any, prompt: str, thread_id: str | None = None, show_thinking: bool = True, workspace_dir: str | None = None) -> None: +def cmd_run( + agent: Any, + prompt: str, + thread_id: str | None = None, + show_thinking: bool = True, + workspace_dir: str | None = None, + model: str | None = None, +) -> None: """Single-shot execution with streaming display. Args: @@ -420,8 +680,9 @@ def cmd_run(agent: Any, prompt: str, thread_id: str | None = None, show_thinking thread_id: Optional thread ID (generates new one if None) show_thinking: Whether to display thinking panels workspace_dir: Per-session workspace directory path + model: Model name for checkpoint metadata """ - thread_id = thread_id or str(uuid.uuid4()) + thread_id = thread_id or generate_thread_id() width = console.size.width sep = Text("\u2500" * width, style="dim") @@ -433,8 +694,9 @@ def cmd_run(agent: Any, prompt: str, thread_id: str | None = None, show_thinking console.print(f"[dim]Workspace: {_shorten_path(workspace_dir)}[/dim]") console.print() + meta = _build_metadata(workspace_dir, model) try: - _run_streaming(agent, prompt, thread_id, show_thinking, interactive=False) + _run_streaming(agent, prompt, thread_id, show_thinking, interactive=False, metadata=meta) except Exception as e: error_msg = str(e) if "authentication" in error_msg.lower() or "api_key" in error_msg.lower(): diff --git a/EvoScientist/sessions.py b/EvoScientist/sessions.py new file mode 100644 index 0000000..55006c6 --- /dev/null +++ b/EvoScientist/sessions.py @@ -0,0 +1,315 @@ +"""Session persistence using LangGraph's SQLite checkpoint storage. + +Provides thread CRUD operations, prefix-matched resume, and an async +context manager for the shared ``AsyncSqliteSaver`` checkpointer. + +Adapted from upstream ``deepagents_cli/sessions.py``. +""" + +import uuid +from collections.abc import AsyncIterator +from contextlib import asynccontextmanager +from datetime import datetime, timezone +from pathlib import Path + +import aiosqlite +from langgraph.checkpoint.serde.jsonplus import JsonPlusSerializer +from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver + +# --------------------------------------------------------------------------- +# Monkey-patch aiosqlite for langgraph-checkpoint >= 2.1.0 compatibility +# --------------------------------------------------------------------------- +if not hasattr(aiosqlite.Connection, "is_alive"): + + def _is_alive(self: aiosqlite.Connection) -> bool: + return self._connection is not None + + aiosqlite.Connection.is_alive = _is_alive # type: ignore[attr-defined] + + +# --------------------------------------------------------------------------- +# Constants +# --------------------------------------------------------------------------- + +AGENT_NAME = "EvoScientist" + + +# --------------------------------------------------------------------------- +# Paths & ID generation +# --------------------------------------------------------------------------- + +def get_db_path() -> Path: + """Return ``~/.config/evoscientist/sessions.db``, creating parents.""" + db_dir = Path.home() / ".config" / "evoscientist" + db_dir.mkdir(parents=True, exist_ok=True) + return db_dir / "sessions.db" + + +def generate_thread_id() -> str: + """Generate an 8-char hex thread ID.""" + return uuid.uuid4().hex[:8] + + +# --------------------------------------------------------------------------- +# Checkpointer context manager +# --------------------------------------------------------------------------- + +@asynccontextmanager +async def get_checkpointer() -> AsyncIterator[AsyncSqliteSaver]: + """Yield an ``AsyncSqliteSaver`` connected to the global sessions DB.""" + async with AsyncSqliteSaver.from_conn_string(str(get_db_path())) as cp: + yield cp + + +# --------------------------------------------------------------------------- +# Internal helpers +# --------------------------------------------------------------------------- + +async def _table_exists(conn: aiosqlite.Connection, table: str) -> bool: + query = "SELECT 1 FROM sqlite_master WHERE type='table' AND name=?" + async with conn.execute(query, (table,)) as cur: + return await cur.fetchone() is not None + + +async def _load_checkpoint_messages( + conn: aiosqlite.Connection, + thread_id: str, + serde: JsonPlusSerializer, +) -> list: + """Load messages from the most recent checkpoint for *thread_id*. + + Returns a list of LangChain message objects, or an empty list on failure. + """ + query = """ + SELECT type, checkpoint + FROM checkpoints + WHERE thread_id = ? + ORDER BY checkpoint_id DESC + LIMIT 1 + """ + async with conn.execute(query, (thread_id,)) as cur: + row = await cur.fetchone() + if not row or not row[0] or not row[1]: + return [] + try: + data = serde.loads_typed((row[0], row[1])) + return data.get("channel_values", {}).get("messages", []) + except (ValueError, TypeError, KeyError): + return [] + + +async def _count_messages( + conn: aiosqlite.Connection, + thread_id: str, + serde: JsonPlusSerializer, +) -> int: + """Count messages in the most recent checkpoint for *thread_id*.""" + msgs = await _load_checkpoint_messages(conn, thread_id, serde) + return len(msgs) + + +def _extract_preview(messages: list, max_len: int = 50) -> str: + """Extract the first human message as a preview string.""" + for msg in messages: + if getattr(msg, "type", None) != "human": + continue + content = getattr(msg, "content", "") or "" + if isinstance(content, list): + parts = [ + b.get("text", "") + for b in content + if isinstance(b, dict) and b.get("type") == "text" + ] + content = " ".join(parts) + content = content.strip() + if content: + return content[:max_len] + "..." if len(content) > max_len else content + return "" + + +def _format_relative_time(iso_ts: str | None) -> str: + """Convert ISO timestamp to a human-readable relative string.""" + if not iso_ts: + return "" + try: + dt = datetime.fromisoformat(iso_ts) + if dt.tzinfo is None: + dt = dt.replace(tzinfo=timezone.utc) + now = datetime.now(timezone.utc) + delta = now - dt + seconds = int(delta.total_seconds()) + if seconds < 60: + return "just now" + minutes = seconds // 60 + if minutes < 60: + return f"{minutes} min ago" + hours = minutes // 60 + if hours < 24: + return f"{hours} hour{'s' if hours != 1 else ''} ago" + days = hours // 24 + if days < 30: + return f"{days} day{'s' if days != 1 else ''} ago" + months = days // 30 + return f"{months} month{'s' if months != 1 else ''} ago" + except (ValueError, TypeError): + return "" + + +# --------------------------------------------------------------------------- +# Thread CRUD +# --------------------------------------------------------------------------- + +async def list_threads( + limit: int = 20, + include_message_count: bool = False, + include_preview: bool = False, +) -> list[dict]: + """List EvoScientist threads, most-recent first. + + Returns list of dicts with keys: ``thread_id``, ``updated_at``, + ``workspace_dir``, ``model``, and optionally ``message_count`` + and ``preview``. + """ + db_path = str(get_db_path()) + async with aiosqlite.connect(db_path, timeout=30.0) as conn: + if not await _table_exists(conn, "checkpoints"): + return [] + + query = """ + SELECT thread_id, + MAX(json_extract(metadata, '$.updated_at')) as updated_at, + json_extract(metadata, '$.workspace_dir') as workspace_dir, + json_extract(metadata, '$.model') as model + FROM checkpoints + WHERE json_extract(metadata, '$.agent_name') = ? + GROUP BY thread_id + ORDER BY updated_at DESC + LIMIT ? + """ + async with conn.execute(query, (AGENT_NAME, limit)) as cur: + rows = await cur.fetchall() + + threads = [ + { + "thread_id": r[0], + "updated_at": r[1], + "workspace_dir": r[2], + "model": r[3], + } + for r in rows + ] + + if (include_message_count or include_preview) and threads: + serde = JsonPlusSerializer() + for t in threads: + msgs = await _load_checkpoint_messages(conn, t["thread_id"], serde) + if include_message_count: + t["message_count"] = len(msgs) + if include_preview: + t["preview"] = _extract_preview(msgs) + + return threads + + +async def get_most_recent() -> str | None: + """Return the most recent EvoScientist thread ID, or ``None``.""" + db_path = str(get_db_path()) + async with aiosqlite.connect(db_path, timeout=30.0) as conn: + if not await _table_exists(conn, "checkpoints"): + return None + query = """ + SELECT thread_id FROM checkpoints + WHERE json_extract(metadata, '$.agent_name') = ? + ORDER BY checkpoint_id DESC + LIMIT 1 + """ + async with conn.execute(query, (AGENT_NAME,)) as cur: + row = await cur.fetchone() + return row[0] if row else None + + +async def thread_exists(thread_id: str) -> bool: + """Return ``True`` if *thread_id* has at least one checkpoint.""" + db_path = str(get_db_path()) + async with aiosqlite.connect(db_path, timeout=30.0) as conn: + if not await _table_exists(conn, "checkpoints"): + return False + query = "SELECT 1 FROM checkpoints WHERE thread_id = ? LIMIT 1" + async with conn.execute(query, (thread_id,)) as cur: + return (await cur.fetchone()) is not None + + +async def find_similar_threads(thread_id: str, limit: int = 5) -> list[str]: + """Find thread IDs that start with *thread_id* (prefix match).""" + db_path = str(get_db_path()) + async with aiosqlite.connect(db_path, timeout=30.0) as conn: + if not await _table_exists(conn, "checkpoints"): + return [] + query = """ + SELECT DISTINCT thread_id + FROM checkpoints + WHERE thread_id LIKE ? + ORDER BY thread_id + LIMIT ? + """ + async with conn.execute(query, (thread_id + "%", limit)) as cur: + rows = await cur.fetchall() + return [r[0] for r in rows] + + +async def delete_thread(thread_id: str) -> bool: + """Delete all checkpoints (and writes) for *thread_id*.""" + db_path = str(get_db_path()) + async with aiosqlite.connect(db_path, timeout=30.0) as conn: + if not await _table_exists(conn, "checkpoints"): + return False + cur = await conn.execute( + "DELETE FROM checkpoints WHERE thread_id = ?", (thread_id,) + ) + deleted = cur.rowcount > 0 + if await _table_exists(conn, "writes"): + await conn.execute("DELETE FROM writes WHERE thread_id = ?", (thread_id,)) + await conn.commit() + return deleted + + +async def get_thread_metadata(thread_id: str) -> dict | None: + """Return metadata dict for *thread_id*, or ``None`` if not found. + + Keys: ``workspace_dir``, ``model``, ``updated_at``. + """ + db_path = str(get_db_path()) + async with aiosqlite.connect(db_path, timeout=30.0) as conn: + if not await _table_exists(conn, "checkpoints"): + return None + query = """ + SELECT json_extract(metadata, '$.workspace_dir') as workspace_dir, + json_extract(metadata, '$.model') as model, + json_extract(metadata, '$.updated_at') as updated_at + FROM checkpoints + WHERE thread_id = ? + ORDER BY checkpoint_id DESC + LIMIT 1 + """ + async with conn.execute(query, (thread_id,)) as cur: + row = await cur.fetchone() + if not row: + return None + return { + "workspace_dir": row[0], + "model": row[1], + "updated_at": row[2], + } + + +async def get_thread_messages(thread_id: str) -> list: + """Return the list of LangChain message objects for *thread_id*. + + Returns an empty list if the thread has no checkpoints. + """ + db_path = str(get_db_path()) + serde = JsonPlusSerializer() + async with aiosqlite.connect(db_path, timeout=30.0) as conn: + if not await _table_exists(conn, "checkpoints"): + return [] + return await _load_checkpoint_messages(conn, thread_id, serde) diff --git a/EvoScientist/stream/display.py b/EvoScientist/stream/display.py index 50d6394..af9cf85 100644 --- a/EvoScientist/stream/display.py +++ b/EvoScientist/stream/display.py @@ -586,6 +586,7 @@ def _run_streaming( on_thinking: Callable[[str], None] | None = None, on_todo: Callable[[list[dict]], None] | None = None, on_file_write: Callable[[str], None] | None = None, + metadata: dict | None = None, ) -> str: """Run async streaming and render with Rich Live display. @@ -604,6 +605,8 @@ def _run_streaming( Called once when write_todos tool_call is detected. on_file_write: Optional sync callback receiving the real filesystem path when the agent writes a media file (image/pdf) via write_file. + metadata: Optional metadata dict forwarded to ``stream_agent_events`` + for LangGraph checkpoint persistence. Returns: The final response text. @@ -618,7 +621,7 @@ def _run_streaming( async def _consume() -> None: nonlocal _thinking_sent, _todo_sent - async for event in stream_agent_events(agent, message, thread_id): + async for event in stream_agent_events(agent, message, thread_id, metadata=metadata): event_type = state.handle_event(event) # Send thinking to channel when transitioning away from thinking diff --git a/EvoScientist/stream/events.py b/EvoScientist/stream/events.py index 97f570d..97a6d9a 100644 --- a/EvoScientist/stream/events.py +++ b/EvoScientist/stream/events.py @@ -59,7 +59,12 @@ def _extract_tool_content(msg) -> tuple[str, bool]: return str(content), False -async def stream_agent_events(agent: Any, message: str, thread_id: str) -> AsyncIterator[dict]: +async def stream_agent_events( + agent: Any, + message: str, + thread_id: str, + metadata: dict | None = None, +) -> AsyncIterator[dict]: """Stream events from the agent graph using async iteration. Uses agent.astream() with subgraphs=True to see sub-agent activity. @@ -68,13 +73,17 @@ async def stream_agent_events(agent: Any, message: str, thread_id: str) -> Async agent: Compiled state graph from create_deep_agent() message: User message thread_id: Thread ID for conversation persistence + metadata: Optional metadata dict merged into the LangGraph config + (e.g. agent_name, updated_at for checkpoint persistence). Yields: Event dicts: thinking, text, tool_call, tool_result, subagent_start, subagent_tool_call, subagent_tool_result, subagent_end, done, error """ - config = {"configurable": {"thread_id": thread_id}} + config: dict[str, Any] = {"configurable": {"thread_id": thread_id}} + if metadata: + config["metadata"] = metadata emitter = StreamEventEmitter() main_tracker = ToolCallTracker() full_response = "" diff --git a/pyproject.toml b/pyproject.toml index fc610b0..14aeff1 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -30,6 +30,7 @@ dependencies = [ "typer>=0.12", "python-dotenv>=1.0", "langgraph-cli[inmem]>=0.4", + "langgraph-checkpoint-sqlite>=3.0.0", "httpx>=0.27", "markdownify>=0.14", "nest-asyncio>=1.6", diff --git a/tests/test_sessions.py b/tests/test_sessions.py new file mode 100644 index 0000000..7d6db51 --- /dev/null +++ b/tests/test_sessions.py @@ -0,0 +1,217 @@ +"""Tests for EvoScientist.sessions — thread CRUD, ID generation, helpers.""" + +import asyncio +import json +import os +import tempfile +import unittest +from unittest.mock import patch + +from EvoScientist.sessions import ( + AGENT_NAME, + _format_relative_time, + delete_thread, + find_similar_threads, + generate_thread_id, + get_db_path, + get_most_recent, + get_thread_metadata, + list_threads, + thread_exists, +) + + +def _run(coro): + """Run an async coroutine synchronously (resilient to closed loops).""" + try: + loop = asyncio.get_event_loop() + if loop.is_closed(): + raise RuntimeError("closed") + except RuntimeError: + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + return loop.run_until_complete(coro) + + +class TestGenerateThreadId(unittest.TestCase): + def test_length(self): + tid = generate_thread_id() + self.assertEqual(len(tid), 8) + + def test_hex(self): + tid = generate_thread_id() + int(tid, 16) # Should not raise + + def test_uniqueness(self): + ids = {generate_thread_id() for _ in range(100)} + self.assertEqual(len(ids), 100) + + +class TestGetDbPath(unittest.TestCase): + def test_uses_config_dir(self): + path = get_db_path() + self.assertTrue(str(path).endswith("sessions.db")) + self.assertIn(".config", str(path)) + self.assertIn("evoscientist", str(path)) + + +class TestFormatRelativeTime(unittest.TestCase): + def test_none(self): + self.assertEqual(_format_relative_time(None), "") + + def test_invalid(self): + self.assertEqual(_format_relative_time("not-a-date"), "") + + def test_recent(self): + from datetime import datetime, timezone + now = datetime.now(timezone.utc).isoformat() + result = _format_relative_time(now) + self.assertIn("just now", result) + + +class TestThreadFunctions(unittest.TestCase): + """Tests using a real temporary SQLite database.""" + + @classmethod + def setUpClass(cls): + """Create a temp DB and populate with test data.""" + cls._tmpdir = tempfile.mkdtemp() + cls._db_path = os.path.join(cls._tmpdir, "test_sessions.db") + + async def _setup(): + import aiosqlite + async with aiosqlite.connect(cls._db_path) as conn: + # Create tables matching LangGraph checkpoint schema + await conn.execute(""" + CREATE TABLE IF NOT EXISTS checkpoints ( + thread_id TEXT NOT NULL, + checkpoint_ns TEXT NOT NULL DEFAULT '', + checkpoint_id TEXT NOT NULL, + parent_checkpoint_id TEXT, + type TEXT, + checkpoint BLOB, + metadata TEXT NOT NULL DEFAULT '{}', + PRIMARY KEY (thread_id, checkpoint_ns, checkpoint_id) + ) + """) + await conn.execute(""" + CREATE TABLE IF NOT EXISTS writes ( + thread_id TEXT NOT NULL, + checkpoint_ns TEXT NOT NULL DEFAULT '', + checkpoint_id TEXT NOT NULL, + task_id TEXT NOT NULL, + idx INTEGER NOT NULL, + channel TEXT NOT NULL, + type TEXT, + blob BLOB, + PRIMARY KEY (thread_id, checkpoint_ns, checkpoint_id, task_id, idx) + ) + """) + + # Insert test checkpoints + for i, tid in enumerate(["abc12345", "abc12399", "def00001"]): + meta = json.dumps({ + "agent_name": AGENT_NAME, + "updated_at": f"2025-01-{15 + i}T10:00:00+00:00", + "workspace_dir": f"/tmp/ws_{tid}", + "model": "claude-sonnet-4-5", + }) + await conn.execute( + "INSERT INTO checkpoints (thread_id, checkpoint_ns, checkpoint_id, metadata) VALUES (?, '', ?, ?)", + (tid, f"cp_{i}", meta), + ) + + # Insert a non-EvoScientist checkpoint (should be filtered) + other_meta = json.dumps({ + "agent_name": "OtherAgent", + "updated_at": "2025-01-20T10:00:00+00:00", + }) + await conn.execute( + "INSERT INTO checkpoints (thread_id, checkpoint_ns, checkpoint_id, metadata) VALUES (?, '', ?, ?)", + ("zzz99999", "cp_other", other_meta), + ) + await conn.commit() + + _run(_setup()) + + # Patch get_db_path to point to our temp DB + cls._patcher = patch( + "EvoScientist.sessions.get_db_path", + return_value=type("P", (), {"__str__": lambda s: cls._db_path, "__fspath__": lambda s: cls._db_path})(), + ) + cls._patcher.start() + + @classmethod + def tearDownClass(cls): + cls._patcher.stop() + try: + os.unlink(cls._db_path) + os.rmdir(cls._tmpdir) + except OSError: + pass + + def test_list_threads(self): + threads = _run(list_threads(limit=10)) + # Should only contain EvoScientist threads + self.assertEqual(len(threads), 3) + # Most recent first + self.assertEqual(threads[0]["thread_id"], "def00001") + + def test_list_threads_with_message_count(self): + threads = _run(list_threads(limit=10, include_message_count=True)) + self.assertIn("message_count", threads[0]) + + def test_thread_exists_true(self): + self.assertTrue(_run(thread_exists("abc12345"))) + + def test_thread_exists_false(self): + self.assertFalse(_run(thread_exists("nonexist"))) + + def test_find_similar(self): + similar = _run(find_similar_threads("abc1")) + self.assertEqual(len(similar), 2) + self.assertIn("abc12345", similar) + self.assertIn("abc12399", similar) + + def test_find_similar_no_match(self): + similar = _run(find_similar_threads("xyz")) + self.assertEqual(len(similar), 0) + + def test_get_most_recent(self): + recent = _run(get_most_recent()) + self.assertIsNotNone(recent) + self.assertEqual(recent, "def00001") + + def test_get_thread_metadata(self): + meta = _run(get_thread_metadata("abc12345")) + self.assertIsNotNone(meta) + self.assertEqual(meta["workspace_dir"], "/tmp/ws_abc12345") + self.assertEqual(meta["model"], "claude-sonnet-4-5") + + def test_get_thread_metadata_missing(self): + meta = _run(get_thread_metadata("nonexist")) + self.assertIsNone(meta) + + def test_delete_thread(self): + # Insert a thread to delete + async def _insert(): + import aiosqlite + async with aiosqlite.connect(self._db_path) as conn: + meta = json.dumps({"agent_name": AGENT_NAME, "updated_at": "2025-01-01T00:00:00+00:00"}) + await conn.execute( + "INSERT INTO checkpoints (thread_id, checkpoint_ns, checkpoint_id, metadata) VALUES (?, '', ?, ?)", + ("todelete", "cp_del", meta), + ) + await conn.commit() + _run(_insert()) + + self.assertTrue(_run(thread_exists("todelete"))) + self.assertTrue(_run(delete_thread("todelete"))) + self.assertFalse(_run(thread_exists("todelete"))) + + def test_delete_nonexistent(self): + self.assertFalse(_run(delete_thread("nope1234"))) + + +if __name__ == "__main__": + unittest.main()