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.
This commit is contained in:
X-iZhang
2026-02-11 23:34:25 +00:00
parent 63517a2a04
commit f29b254d9a
9 changed files with 957 additions and 132 deletions
+5 -1
View File
@@ -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"),
}
+5 -3
View File
@@ -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)
+20 -8
View File
@@ -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)
+379 -117
View File
@@ -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 <id>[/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 <id> — delete a saved session."""
if not arg:
console.print("[red]Usage: /delete <thread-id>[/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('<ansiblue><b>\u276f</b></ansiblue> ')
)
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('<ansiblue><b>\u276f</b></ansiblue> ')
)
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():
+315
View File
@@ -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)
+4 -1
View File
@@ -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
+11 -2
View File
@@ -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 = ""
+1
View File
@@ -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",
+217
View File
@@ -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()