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:
@@ -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"),
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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():
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
|
||||
@@ -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 = ""
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user