refactor: LangGraph gateway layer for UI-agnostic graph and thread access (#295)

* feat(gateway): graph gateway protocol

* refactor(cli): wire gateway in cli/tui

* refactor(gateway): centralize runtime gateway init

* chore(gateway): restrict RunRequest message type

* feat(gateway): add langgraph server gateway

* chore(cli): tighten serve runtime state typing

* refactor(cli): route async task state reads through graph gateway

* refactor(gateway): support graph targets in server gateway

* refactor(cli): route session commands through graph gateway

* refactor(cli): fold thread store under graph gateway

* refactor(gateway): route graph state access through gateway

* refactor(channels): wire graph gateway

* refactor(memory): preserve graph threads for cloning

* feat(gateway): add thread cloning

* fix(tui): pass effective workspace for thread creation

* chore(memory): add workspare dir to memory worker metadata

* fix(sessions): filter preloaded UUID registy entries by the current scope

* test(fakes): use https

* refactor(consumer): consolidate imports

* fix(stream): optional summarization event

* fix(gateway): resolve abbreviated thread IDs by search

* fix(gateway): page server thread listings

* fix(gateway): emit pending interrupt events

* style: fmt

* feat(gateway): persist workspace_dir & model in thread metadata

* fix(gateway): page server thread prefix resolution

* fix(gateway): expose server thread list metadata

* refactor: add back type def

* refactor: tighten types

* revert: add back worker thread deletion

The worker thread forking changes are out of scope for now, so to
maintain parity with the existing behavior we'll leave this intact.

* fix(gateway): apply compaction to server thread history

* refactor(stream): restore direct summary replay suppression

* fix(gateway): preserve compaction state and server stream output

* fix(gateway): close local stream generator on cancellation

---------

Co-authored-by: Xi Zhang <106144707+X-iZhang@users.noreply.github.com>
This commit is contained in:
dinos
2026-06-22 15:54:14 +02:00
committed by GitHub
parent 6eb467e70b
commit bd307f3a11
49 changed files with 5071 additions and 1582 deletions
+5 -1
View File
@@ -20,6 +20,7 @@ import json
import logging
import os
from pathlib import Path
from typing import TYPE_CHECKING
from langchain.agents.middleware import AgentMiddleware, HumanInTheLoopMiddleware
@@ -37,6 +38,9 @@ from .prompts import get_system_prompt
# Suppress noisy warnings from deepagents skill loader (non-string frontmatter fields, etc.)
logging.getLogger("deepagents.middleware.skills").setLevel(logging.ERROR)
if TYPE_CHECKING:
from langgraph.graph.state import CompiledStateGraph
# =============================================================================
# Constants
# =============================================================================
@@ -848,7 +852,7 @@ def create_cli_agent(
chat_model=None,
*,
on_mcp_progress=None,
):
) -> "CompiledStateGraph":
"""Create agent with checkpointer for CLI multi-turn support.
A fresh backend is constructed on every call using the current
+22 -21
View File
@@ -11,12 +11,12 @@ from __future__ import annotations
import asyncio
import logging
import uuid
from collections import OrderedDict
from collections.abc import AsyncIterator, Callable
from dataclasses import dataclass
from typing import Any, TypeVar
from ..gateway import GraphGateway, GraphRunInput, GraphTarget, RunRequest
from .base import Channel
from .bus import MessageBus
from .bus.events import InboundMessage, OutboundMessage
@@ -234,9 +234,11 @@ class InboundConsumer:
manager:
The ChannelManager (used to look up channel instances).
agent:
The agent object (must support ``stream_agent_events``).
The local agent object used by local graph gateway targets.
thread_id:
Default thread ID for agent conversations.
graph_gateway:
Gateway used for thread creation and graph streaming.
send_thinking:
Whether to forward thinking messages to the channel.
on_message_received:
@@ -267,6 +269,7 @@ class InboundConsumer:
agent: Any,
thread_id: str,
*,
graph_gateway: GraphGateway,
send_thinking: bool = False,
on_message_received: Callable[[InboundMessage], None] | None = None,
on_streaming_event: Callable[[dict], None] | None = None,
@@ -280,6 +283,7 @@ class InboundConsumer:
self.manager = manager
self.agent = agent
self.thread_id = thread_id
self.graph_gateway = graph_gateway
self.send_thinking = send_thinking
self._on_message_received = on_message_received
self._on_streaming_event = on_streaming_event
@@ -313,7 +317,7 @@ class InboundConsumer:
# ask_user: pending reply per session_key
self._pending_ask_user_replies: dict[str, _PendingAskUserReply] = {}
def _get_thread_id(self, sender_id: str) -> str:
async def _get_thread_id(self, sender_id: str) -> str:
"""Get or create a thread ID for the given sender.
Uses LRU ordering: recently accessed senders are moved to the
@@ -329,7 +333,9 @@ class InboundConsumer:
if self.thread_id:
self._sessions[sender_id] = f"{self.thread_id}:{sender_id}"
else:
self._sessions[sender_id] = str(uuid.uuid4())
self._sessions[sender_id] = await self.graph_gateway.create_thread(
GraphTarget(local_graph=self.agent)
)
return self._sessions[sender_id]
def _get_channel(self, channel_name: str) -> Channel | None:
@@ -423,7 +429,7 @@ class InboundConsumer:
pass
channel = self._get_channel(msg.channel)
thread_id = self._get_thread_id(msg.sender_id)
thread_id = await self._get_thread_id(msg.sender_id)
session_key = msg.session_key # "channel:chat_id"
# Lazily create per-chat lock; evict stale locks when too many
@@ -466,9 +472,9 @@ class InboundConsumer:
session_key: str,
) -> None:
"""Stream agent events with HITL interrupt handling."""
from ..stream.events import stream_agent_events
from langgraph.types import Command
stream_input: Any = msg.content
stream_input: GraphRunInput = msg.content
try:
if channel:
@@ -507,13 +513,15 @@ class InboundConsumer:
return True
async for event in _timeout_aiter(
stream_agent_events(
self.agent,
stream_input,
thread_id,
media=msg.media or None
if isinstance(stream_input, str)
else None,
self.graph_gateway.stream_events(
RunRequest(
message=stream_input,
thread_id=thread_id,
media=msg.media or None
if isinstance(stream_input, str)
else None,
target=GraphTarget(local_graph=self.agent),
)
),
self._inference_timeout,
):
@@ -597,7 +605,6 @@ class InboundConsumer:
interrupt_data,
session_key,
)
from langgraph.types import Command # type: ignore[import-untyped]
stream_input = Command(resume=result)
continue
@@ -608,8 +615,6 @@ class InboundConsumer:
# Session auto-approve (user previously chose "Approve all")
if session_key in self._auto_approve_sessions:
from langgraph.types import Command # type: ignore[import-untyped]
stream_input = Command(
resume={"decisions": [{"type": "approve"} for _ in range(n)]}
)
@@ -617,8 +622,6 @@ class InboundConsumer:
# Config auto-approve (auto_approve, non-execute, allow_list)
if _should_auto_approve(action_reqs):
from langgraph.types import Command # type: ignore[import-untyped]
stream_input = Command(
resume={"decisions": [{"type": "approve"} for _ in range(n)]}
)
@@ -706,8 +709,6 @@ class InboundConsumer:
if decision == "auto":
self._auto_approve_sessions.add(session_key)
from langgraph.types import Command # type: ignore[import-untyped]
stream_input = Command(
resume={"decisions": [{"type": "approve"} for _ in range(n)]}
)
+3
View File
@@ -108,8 +108,10 @@ async def _async_main(
if use_agent:
logger.info("Loading EvoScientist agent...")
from ..EvoScientist import create_cli_agent
from ..gateway import create_runtime_gateways
agent = create_cli_agent()
runtime_gateways = create_runtime_gateways()
logger.info("Agent loaded")
consumer = InboundConsumer(
@@ -117,6 +119,7 @@ async def _async_main(
manager=manager,
agent=agent,
thread_id="",
graph_gateway=runtime_gateways.graph_gateway,
send_thinking=send_thinking,
)
manager.register_health_provider("consumer", lambda: consumer.metrics)
+13 -10
View File
@@ -9,15 +9,16 @@ from __future__ import annotations
import asyncio
import logging
from collections.abc import Callable
from typing import Any
from typing import Any, Generic, TypeVar
_logger = logging.getLogger(__name__)
ProgressEvent = str # "start" | "success" | "error"
ProgressState = str # "pending" | "ok" | "error"
AgentT = TypeVar("AgentT")
ProgressCallback = Callable[[ProgressEvent, str, str], None]
SuccessCallback = Callable[[Any], None]
SuccessCallback = Callable[[AgentT], None]
FailureCallback = Callable[[BaseException], None]
@@ -73,7 +74,7 @@ class MCPProgressTracker:
return done, total
class BackgroundAgentLoader:
class BackgroundAgentLoader(Generic[AgentT]):
"""Owns the background ``_load_agent`` task and its generation token.
Each :meth:`start` bumps an internal id; callbacks from a superseded
@@ -88,7 +89,7 @@ class BackgroundAgentLoader:
def __init__(
self,
loader_fn: Callable[..., Any],
loader_fn: Callable[..., AgentT],
*,
on_progress: ProgressCallback | None = None,
on_success: SuccessCallback | None = None,
@@ -98,12 +99,12 @@ class BackgroundAgentLoader:
self._on_progress = on_progress
self._on_success = on_success
self._on_failure = on_failure
self.agent: Any = None
self._task: asyncio.Task | None = None
self.agent: AgentT | None = None
self._task: asyncio.Task[AgentT] | None = None
self._load_id: int = 0
@property
def task(self) -> asyncio.Task | None:
def task(self) -> asyncio.Task[AgentT] | None:
return self._task
@property
@@ -146,7 +147,7 @@ class BackgroundAgentLoader:
)
self._task.add_done_callback(lambda task, lid=load_id: self._on_done(task, lid))
def adopt(self, agent: Any) -> None:
def adopt(self, agent: AgentT) -> None:
"""Install an externally-built agent and supersede any in-flight load.
Used by ``/model`` (and any other caller that constructs a
@@ -162,7 +163,7 @@ class BackgroundAgentLoader:
self._task = None
self.agent = agent
async def await_ready(self) -> Any:
async def await_ready(self) -> AgentT:
"""Return the loaded agent; re-raises on load failure.
Idempotent. State transitions (setting ``self.agent``, calling
@@ -177,9 +178,11 @@ class BackgroundAgentLoader:
"BackgroundAgentLoader.await_ready called before start()"
)
await self._task
if self.agent is None:
raise RuntimeError("BackgroundAgentLoader completed without an agent")
return self.agent
def _on_done(self, task: asyncio.Task, load_id: int) -> None:
def _on_done(self, task: asyncio.Task[AgentT], load_id: int) -> None:
if load_id != self._load_id:
return
if task.cancelled():
+5 -1
View File
@@ -3,9 +3,13 @@
import os
from datetime import datetime
from pathlib import Path
from typing import TYPE_CHECKING
from ..paths import new_run_dir
if TYPE_CHECKING:
from langgraph.graph.state import CompiledStateGraph
def _shorten_path(path: str) -> str:
"""Shorten absolute path to relative path from current directory."""
@@ -65,7 +69,7 @@ def _load_agent(
chat_model=None,
*,
on_mcp_progress=None,
):
) -> "CompiledStateGraph":
"""Load the CLI agent with optional persistent checkpointer.
Args:
+33 -10
View File
@@ -16,7 +16,10 @@ import threading
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from datetime import UTC, datetime
from typing import Final
from typing import TYPE_CHECKING, Final, TypeAlias, TypedDict
if TYPE_CHECKING:
from ..gateway import GraphGateway, GraphTarget
TERMINAL_STATUSES: Final = frozenset({"success", "error", "timeout", "interrupted"})
"""Aligned with langgraph_sdk.schema.RunStatus terminal values.
@@ -31,6 +34,15 @@ Cancel operations transition runs into ``interrupted`` (not ``cancelled``).
_MAX_RECONNECT_ATTEMPTS: Final = 10
class AsyncTaskState(TypedDict, total=False):
status: str
last_checked_at: str
last_updated_at: str
AsyncTasksState: TypeAlias = dict[str, AsyncTaskState]
@dataclass(frozen=True)
class AsyncTaskNotification:
"""A completed-async-task signal pushed by a watcher."""
@@ -66,12 +78,10 @@ _notification_queue = _unrouted_queue
# dict[handle, origin_cli_thread_id] so the consumer's batching grace loop
# can filter for watchers tied to the current CLI thread (or unrouted)
# without being delayed by sibling-thread watchers.
_active_watchers: dict = {}
_active_watchers: dict[object, str | None] = {}
# Map thread_id (sub-agent thread) → current watcher handle (supports
# replacement on update_async_task). Value type widens from asyncio.Task
# to "anything with .cancel()/.done()/.add_done_callback()" so we can
# move watcher scheduling onto a background loop in a follow-up fix.
_watcher_by_thread: dict[str, object] = {}
# replacement on update_async_task).
_watcher_by_thread: dict[str, asyncio.Task[None]] = {}
def _has_relevant_active_watchers(current_thread_id: str | None) -> bool:
@@ -129,6 +139,19 @@ def pending_thread_ids() -> set[str]:
return {tid for tid, q in _notifications_by_thread.items() if not q.empty()}
async def read_async_tasks_from_gateway(
gateway: GraphGateway,
target: GraphTarget,
thread_id: str,
) -> AsyncTasksState:
"""Read async_tasks state through the active graph gateway."""
try:
values = await gateway.get_state_values(target, thread_id)
except Exception:
return {}
return values.get("async_tasks", {})
async def watch_run_and_notify(
client,
thread_id: str,
@@ -285,7 +308,7 @@ def spawn_watcher(
agent_name: str,
prompt: str = "",
origin_cli_thread_id: str | None = None,
) -> asyncio.Task:
) -> asyncio.Task[None]:
"""Spawn a watcher on the caller's asyncio loop.
Replacement semantics support ``update_async_task`` which creates a new
@@ -319,7 +342,7 @@ def spawn_watcher(
_watcher_by_thread[thread_id] = task
_active_watchers[task] = origin_cli_thread_id
def _cleanup(t: asyncio.Task) -> None:
def _cleanup(t: asyncio.Task[None]) -> None:
_active_watchers.pop(t, None)
# Only remove if THIS task is still the registered one — could
# have been replaced by a newer spawn_watcher call already.
@@ -366,7 +389,7 @@ def drain_notifications(
def dedup_notifications(
notifs: list[AsyncTaskNotification],
async_tasks: dict[str, dict] | None,
async_tasks: AsyncTasksState | None,
) -> list[AsyncTaskNotification]:
"""Filter notifications the agent has already 'seen' via prior check.
@@ -513,7 +536,7 @@ NOTIFICATION_ACTIVE_WATCHER_WAIT_SECONDS = 3.0
async def consume_notifications(
run_message: Callable[[str, list[AsyncTaskNotification]], Awaitable[None]],
read_async_tasks_state: Callable[[], Awaitable[dict[str, dict]]],
read_async_tasks_state: Callable[[], Awaitable[AsyncTasksState]],
current_thread_id: str | None = None,
) -> None:
"""Drain queue, dedup, batch, and inject as a synthetic user message.
+22 -7
View File
@@ -9,6 +9,8 @@ enqueues a ``ChannelMessage`` on a thread-safe ``queue.Queue`` and waits
for the main thread to set a response via ``_set_channel_response()``.
"""
from __future__ import annotations
import asyncio
import logging
import queue
@@ -17,7 +19,7 @@ import time
import uuid
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from typing import Any
from typing import TYPE_CHECKING, Any
from rich.panel import Panel
from rich.text import Text
@@ -25,6 +27,9 @@ from rich.text import Text
from ..commands.base import ChannelRuntime
from ..stream.console import console
if TYPE_CHECKING:
from ..gateway import GraphGateway
_channel_logger = logging.getLogger(__name__)
@@ -253,7 +258,8 @@ async def dispatch_channel_slash_command(
workspace_dir: str | None,
checkpointer: Any,
append_system: Callable[[str, str], None],
start_new_session_cb: Callable[[], None] | None = None,
graph_gateway: GraphGateway,
start_new_session_cb: Callable[[], Awaitable[None]] | None = None,
handle_session_resume_cb: Callable[..., Awaitable[None]] | None = None,
await_agent_ready: Callable[[], Awaitable[Any]] | None = None,
on_cmd_completed: Callable[..., Awaitable[None]] | None = None,
@@ -284,6 +290,9 @@ async def dispatch_channel_slash_command(
Optional lifecycle callbacks forwarded to ``ChannelCommandUI``.
Headless serve passes ``None`` — ``/new`` and ``/resume`` degrade
gracefully via the default ``ChannelCommandUI`` messages.
graph_gateway:
Graph gateway forwarded to slash commands and channel resume-history
rendering.
await_agent_ready:
Optional async resolver that blocks until the background agent
load finishes. Called only when ``cmd.needs_agent(args)`` is
@@ -320,6 +329,7 @@ async def dispatch_channel_slash_command(
await_agent_ready=await_agent_ready,
on_cmd_completed=on_cmd_completed,
channel_runtime=channel_runtime,
graph_gateway=graph_gateway,
)
except Exception as exc:
# Last-ditch safety: any uncaught exception from inside the
@@ -350,7 +360,8 @@ async def _dispatch_channel_slash_impl(
workspace_dir: str | None,
checkpointer: Any,
append_system: Callable[[str, str], None],
start_new_session_cb: Callable[[], None] | None,
graph_gateway: GraphGateway,
start_new_session_cb: Callable[[], Awaitable[None]] | None,
handle_session_resume_cb: Callable[..., Awaitable[None]] | None,
await_agent_ready: Callable[[], Awaitable[Any]] | None,
on_cmd_completed: Callable[..., Awaitable[None]] | None,
@@ -386,6 +397,7 @@ async def _dispatch_channel_slash_impl(
append_system_callback=append_system,
start_new_session_callback=start_new_session_cb,
handle_session_resume_callback=handle_session_resume_cb,
graph_gateway=graph_gateway,
)
ctx = CommandContext(
agent=agent_for_ctx,
@@ -394,6 +406,7 @@ async def _dispatch_channel_slash_impl(
workspace_dir=workspace_dir,
checkpointer=checkpointer,
channel_runtime=channel_runtime,
graph_gateway=graph_gateway,
)
try:
@@ -619,7 +632,7 @@ def _try_set_hitl_reply(channel_type: str, chat_id: str, content: str) -> bool:
def channel_ask_user_prompt(
ask_user_data: dict,
msg: "ChannelMessage | None" = None,
msg: ChannelMessage | None = None,
) -> dict:
"""Format ask_user questions and collect answers from a channel user.
@@ -651,7 +664,7 @@ def channel_ask_user_prompt(
channel=msg.channel_type,
chat_id=msg.chat_id,
content=content,
metadata=msg.metadata,
metadata=msg.metadata or {},
)
),
bus_loop,
@@ -748,7 +761,7 @@ def channel_ask_user_prompt(
def channel_hitl_prompt(
action_requests: list,
msg: "ChannelMessage",
msg: ChannelMessage,
) -> list[dict] | None:
"""Send HITL approval prompt to channel user and wait for reply.
@@ -793,7 +806,9 @@ def channel_hitl_prompt(
channel=msg.channel_type,
chat_id=msg.chat_id,
content=content,
metadata=metadata if metadata is not None else msg.metadata,
metadata=metadata
if metadata is not None
else msg.metadata or {},
)
),
bus_loop,
+186 -139
View File
@@ -1,23 +1,32 @@
"""Typer command registrations — onboard, config, mcp, main callback."""
import asyncio
import logging
import os
import queue
import re
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from datetime import datetime
from importlib.metadata import version as _pkg_version
from pathlib import Path
from typing import Annotated, Any, cast
from typing import TYPE_CHECKING, Annotated, Any, cast
import typer
from rich.markup import escape
from rich.table import Table
from ..commands.base import Command, CommandContext
from ..commands.base import ChannelRuntime, Command, CommandContext
from ..gateway import (
GraphGateway,
GraphTarget,
RuntimeGateways,
create_runtime_gateways,
)
from ..llm.context_window import DEFAULT_CONTEXT_WINDOW_FALLBACK, resolve_context_window
from ..paths import ensure_dirs, set_active_workspace, set_workspace_root
from ..stream.console import console
from . import async_notifier
from ._app import app, channel_app, config_app, configure_app, mcp_app, sessions_app
from ._constants import build_metadata
from .agent import (
@@ -51,6 +60,11 @@ from .mcp_ui import (
_show_mcp_config,
)
if TYPE_CHECKING:
from langgraph.graph.state import CompiledStateGraph
from ..config import EvoScientistConfig
# =============================================================================
# Onboard command
# =============================================================================
@@ -596,16 +610,17 @@ def build_compact_summary_renderable(
async def compact_conversation(
agent: Any,
thread_id: str | None,
graph_gateway: GraphGateway,
thread_id: str,
target: GraphTarget,
*,
input_tokens_hint: int | None = None,
) -> CompactResult:
"""Compact the conversation by summarizing old messages.
Reads the agent's checkpointed state, creates a temporary
Reads the graph's checkpointed state, creates a temporary
``SummarizationMiddleware``, generates a summary, and writes
the compacted state back via ``aupdate_state``.
the compacted state back through ``GraphGateway``.
``input_tokens_hint`` is the real LLM input token count from the last
``usage_metadata`` (includes system prompt + tool schemas). When
@@ -615,20 +630,17 @@ async def compact_conversation(
Returns a structured ``CompactResult``.
"""
if not agent or not thread_id:
return CompactResult("noop", "Nothing to compact — start a conversation first.")
from langchain_core.messages.utils import count_tokens_approximately
from langchain_core.runnables import RunnableConfig
config: RunnableConfig = {"configurable": {"thread_id": thread_id}}
try:
state_snapshot = await agent.aget_state(config)
state_values = await graph_gateway.get_state_values(target, thread_id)
except Exception as exc:
return CompactResult("error", f"Failed to read state: {exc}")
messages = state_snapshot.values.get("messages", [])
messages = state_values.get("messages", [])
if not messages:
return CompactResult(
"noop", "Nothing to compact — no messages in conversation."
@@ -661,7 +673,7 @@ async def compact_conversation(
)
# Rebuild effective message list accounting for prior compaction
event = state_snapshot.values.get("_summarization_event")
event = state_values.get("_summarization_event")
effective = middleware._apply_event_to_messages(messages, event)
effective_tokens = count_tokens_approximately(effective)
@@ -783,7 +795,11 @@ async def compact_conversation(
"file_path": file_path,
}
await agent.aupdate_state(config, {"_summarization_event": new_event})
await graph_gateway.update_state_values(
target,
thread_id,
{"_summarization_event": new_event},
)
return CompactResult(
"ok",
@@ -809,9 +825,44 @@ async def compact_conversation(
_serve_logger = logging.getLogger(__name__)
@dataclass(slots=True)
class ServeRuntimeState:
"""Mutable serve-mode runtime shared by the poll loop and slash callbacks."""
agent: "CompiledStateGraph"
thread_id: str
workspace_dir: str | None
config: "EvoScientistConfig | None"
runtime_gateways: RuntimeGateways
resume_warning_thread_id: str | None = None
def set_agent(
self,
agent: "CompiledStateGraph",
channel_runtime: ChannelRuntime | None,
) -> None:
self.agent = agent
if channel_runtime is not None:
channel_runtime.agent = agent
def set_thread_id(
self,
thread_id: str,
channel_runtime: ChannelRuntime | None,
*,
forget_previous_origin: bool = True,
) -> None:
old_thread_id = self.thread_id
if forget_previous_origin:
forget_channel_origin(old_thread_id)
self.thread_id = thread_id
if channel_runtime is not None:
channel_runtime.thread_id = thread_id
def _make_serve_start_new_session_cb(
agent_holder: dict[str, Any],
channel_runtime: Any | None = None,
runtime_state: ServeRuntimeState,
channel_runtime: ChannelRuntime | None = None,
):
"""Build the ``start_new_session_cb`` used by serve mode.
@@ -821,55 +872,52 @@ def _make_serve_start_new_session_cb(
fresh thread id. Without a wired callback the channel user gets
``ChannelCommandUI``'s fallback "restart the channel link" message
and nothing actually rotates. This helper generates a new thread
id, updates the shared holder, and syncs the channel runtime so
id, updates the shared runtime state, and syncs the channel runtime so
subsequent messages land on the new thread.
"""
def _cb() -> None:
from ..sessions import generate_thread_id
new_tid = generate_thread_id()
forget_channel_origin(agent_holder.get("thread_id"))
agent_holder["thread_id"] = new_tid
if channel_runtime is not None:
channel_runtime.thread_id = new_tid
async def _cb() -> None:
new_tid = await runtime_state.runtime_gateways.graph_gateway.create_thread(
GraphTarget(workspace_dir=runtime_state.workspace_dir)
)
runtime_state.set_thread_id(new_tid, channel_runtime)
console.print(f"[dim][serve] New thread: {new_tid}[/dim]")
return _cb
def _serve_resume_config(
agent_holder: dict[str, Any],
config: Any | None,
) -> Any | None:
runtime_state: ServeRuntimeState,
config: "EvoScientistConfig | None",
) -> "EvoScientistConfig | None":
"""Return the effective config to use for serve-mode resume sync."""
return config if config is not None else agent_holder.get("config")
return config if config is not None else runtime_state.config
async def _apply_serve_resume_state(
agent_holder: dict[str, Any],
channel_runtime: Any | None,
runtime_state: ServeRuntimeState,
channel_runtime: ChannelRuntime | None,
*,
thread_id: str,
workspace_dir: str | None,
config: Any | None = None,
config: "EvoScientistConfig | None" = None,
) -> None:
"""Adopt a resumed thread/workspace into serve-mode runtime state.
Workspace-bound resources are rebuilt and synced before mutating the shared
holder. The agent is loaded before syncing the external server so a load
state. The agent is loaded before syncing the external server so a load
failure cannot move the server away from the currently active session.
"""
import asyncio
old_workspace = agent_holder.get("workspace_dir")
old_workspace = runtime_state.workspace_dir
new_workspace = (
workspace_dir if workspace_dir and workspace_dir != old_workspace else None
)
new_agent: Any | None = None
workspace_update: tuple[str, CompiledStateGraph] | None = None
if new_workspace is not None:
effective_config = _serve_resume_config(agent_holder, config)
effective_config = _serve_resume_config(runtime_state, config)
if effective_config is None:
raise RuntimeError(
"Cannot resume into a different workspace in serve mode without "
@@ -885,59 +933,56 @@ async def _apply_serve_resume_state(
effective_config,
workspace_dir=new_workspace,
)
workspace_update = (new_workspace, new_agent)
except Exception:
if old_workspace:
set_active_workspace(old_workspace)
raise
old_thread_id = agent_holder.get("thread_id")
thread_changed = bool(thread_id) and thread_id != old_thread_id
old_thread_id = runtime_state.thread_id
thread_changed = thread_id != old_thread_id
if thread_changed:
forget_channel_origin(old_thread_id)
agent_holder["thread_id"] = thread_id
if channel_runtime is not None:
channel_runtime.thread_id = thread_id
runtime_state.set_thread_id(thread_id, channel_runtime)
if new_workspace is not None:
agent_holder["workspace_dir"] = new_workspace
agent_holder["agent"] = new_agent
if channel_runtime is not None:
channel_runtime.agent = new_agent
if workspace_update is not None:
updated_workspace, updated_agent = workspace_update
runtime_state.workspace_dir = updated_workspace
runtime_state.set_agent(updated_agent, channel_runtime)
def _make_serve_handle_session_resume_cb(
agent_holder: dict[str, Any],
channel_runtime: Any | None = None,
runtime_state: ServeRuntimeState,
channel_runtime: ChannelRuntime | None = None,
*,
config: Any | None = None,
config: "EvoScientistConfig | None" = None,
):
"""Build the ChannelCommandUI resume callback for serve mode."""
async def _cb(thread_id: str, workspace_dir: str | None = None) -> None:
old_thread_id = agent_holder.get("thread_id")
old_thread_id = runtime_state.thread_id
await _apply_serve_resume_state(
agent_holder,
runtime_state,
channel_runtime,
thread_id=thread_id,
workspace_dir=workspace_dir,
config=config,
)
if thread_id and thread_id != old_thread_id:
agent_holder["_resume_warning_thread_id"] = thread_id
if thread_id != old_thread_id:
runtime_state.resume_warning_thread_id = thread_id
return _cb
def _make_serve_cmd_completed_hook(
agent_holder: dict[str, Any],
channel_runtime: Any | None = None,
runtime_state: ServeRuntimeState,
channel_runtime: ChannelRuntime | None = None,
*,
config: Any | None = None,
config: "EvoScientistConfig | None" = None,
):
"""Build the ``on_cmd_completed`` hook used by serve mode.
Adopts ``/model`` agent swaps and ``/resume`` thread/workspace
swaps back into ``agent_holder`` so the outer poll loop picks up
swaps back into ``runtime_state`` so the outer poll loop picks up
the new handles on subsequent messages. Also keeps
``channel_runtime`` in sync so the bus sees the new values.
@@ -952,17 +997,17 @@ def _make_serve_cmd_completed_hook(
without spinning up the whole serve loop.
"""
async def _hook(ctx: CommandContext, original_agent: Any, cmd: Command) -> None:
async def _hook(
ctx: CommandContext,
original_agent: "CompiledStateGraph",
cmd: Command,
) -> None:
if ctx.agent is not None and ctx.agent is not original_agent:
agent_holder["agent"] = ctx.agent
if channel_runtime is not None:
channel_runtime.agent = ctx.agent
runtime_state.set_agent(ctx.agent, channel_runtime)
old_thread_id = agent_holder.get("thread_id")
resume_warning_thread_id = agent_holder.pop(
"_resume_warning_thread_id",
None,
)
old_thread_id = runtime_state.thread_id
resume_warning_thread_id = runtime_state.resume_warning_thread_id
runtime_state.resume_warning_thread_id = None
# ``/resume`` mutates ``ctx.thread_id`` directly (its UI callback
# is a no-op in serve mode since there's no REPL to reset). Pick
@@ -975,21 +1020,18 @@ def _make_serve_cmd_completed_hook(
new_tid = ctx.thread_id
if cmd.name == "/resume":
await _apply_serve_resume_state(
agent_holder,
runtime_state,
channel_runtime,
thread_id=new_tid,
workspace_dir=ctx.workspace_dir,
config=config,
)
else:
thread_changed = bool(new_tid) and new_tid != old_thread_id
thread_changed = new_tid != old_thread_id
if thread_changed:
forget_channel_origin(old_thread_id)
agent_holder["thread_id"] = new_tid
if channel_runtime is not None:
channel_runtime.thread_id = new_tid
runtime_state.set_thread_id(new_tid, channel_runtime)
thread_changed = bool(new_tid) and new_tid != old_thread_id
thread_changed = new_tid != old_thread_id
# Surface the in-memory-state limitation to the channel user
# for ``/resume`` so the missing history isn't silent. Flush
@@ -1014,22 +1056,21 @@ def _make_serve_cmd_completed_hook(
def _serve_process_message(
msg: ChannelMessage,
*,
agent_holder: dict[str, Any],
runtime_state: ServeRuntimeState,
model: str | None,
workspace_dir: str,
show_thinking: bool,
on_cmd_completed: Callable[..., Awaitable[None]] | None = None,
handle_session_resume_cb: Callable[..., Awaitable[None]] | None = None,
start_new_session_cb: Callable[[], None] | None = None,
channel_runtime: Any | None = None,
start_new_session_cb: Callable[[], Awaitable[None]] | None = None,
channel_runtime: ChannelRuntime | None = None,
) -> None:
"""Process a single channel message in headless serve mode.
Headless equivalent of interactive.py's ``_process_channel_message``.
No CLI prompt manipulation — just log lines for monitoring.
``agent_holder`` is a mutable dict (keys: ``agent``, ``thread_id``,
``workspace_dir``) shared with the outer ``serve()`` loop.
``runtime_state`` is shared with the outer ``serve()`` loop.
``on_cmd_completed`` (the agent-swap / session-adoption hook) and
``start_new_session_cb`` (thread rotation for ``/new``) are
constructed once in ``serve()`` — if omitted, they're rebuilt per
@@ -1042,12 +1083,14 @@ def _serve_process_message(
from .channel import _bus_loop
from .tui_runtime import run_streaming
runtime_gateways = runtime_state.runtime_gateways
if not _claim_or_complete_channel_request(msg):
return
remember_channel_origin(agent_holder.get("thread_id"), msg)
remember_channel_origin(runtime_state.thread_id, msg)
runtime_workspace = agent_holder.get("workspace_dir") or workspace_dir
runtime_workspace = runtime_state.workspace_dir or workspace_dir
console.print(
f"[dim][{msg.channel_type}] {msg.sender}: {escape(msg.content[:80])}[/dim]"
@@ -1138,25 +1181,29 @@ def _serve_process_message(
_slash_handled = _slash_loop.run_until_complete(
dispatch_channel_slash_command(
msg,
agent=agent_holder["agent"],
thread_id=agent_holder["thread_id"],
agent=runtime_state.agent,
thread_id=runtime_state.thread_id,
workspace_dir=runtime_workspace,
checkpointer=None,
append_system=lambda t, s="dim": console.print(t, style=s),
start_new_session_cb=start_new_session_cb
or _make_serve_start_new_session_cb(agent_holder, channel_runtime),
or _make_serve_start_new_session_cb(
runtime_state,
channel_runtime,
),
handle_session_resume_cb=handle_session_resume_cb
or _make_serve_handle_session_resume_cb(
agent_holder,
runtime_state,
channel_runtime,
),
on_cmd_completed=on_cmd_completed
or _make_serve_cmd_completed_hook(
agent_holder,
runtime_state,
channel_runtime,
config=agent_holder.get("config"),
config=runtime_state.config,
),
channel_runtime=channel_runtime,
graph_gateway=runtime_gateways.graph_gateway,
)
)
except Exception as exc:
@@ -1178,7 +1225,7 @@ def _serve_process_message(
# A channel-issued /new or /resume rotates the thread inside the
# dispatch above; re-bind the now-current thread to this channel
# so async-notifier turns on it still forward back here.
remember_channel_origin(agent_holder["thread_id"], msg)
remember_channel_origin(runtime_state.thread_id, msg)
console.print(f"[dim][{msg.channel_type}] Replied to {msg.sender}[/dim]")
return
@@ -1186,9 +1233,9 @@ def _serve_process_message(
try:
response = run_streaming(
ui_backend="cli",
agent=agent_holder["agent"],
agent=runtime_state.agent,
message=msg.content,
thread_id=agent_holder["thread_id"],
thread_id=runtime_state.thread_id,
show_thinking=show_thinking,
interactive=True,
metadata=meta,
@@ -1198,6 +1245,7 @@ def _serve_process_message(
hitl_prompt_fn=_hitl_prompt,
ask_user_prompt_fn=_ask_user_prompt,
cancel_scope=_channel_message_cancel_scope(msg),
gateway=runtime_gateways.graph_gateway,
)
except Exception as e:
response = f"Error: {e}"
@@ -1216,7 +1264,7 @@ def _serve_process_message(
def _serve_drain_notifications(
*,
agent_holder: dict,
runtime_state: ServeRuntimeState,
model: str | None,
workspace_dir: str,
show_thinking: bool,
@@ -1228,8 +1276,6 @@ def _serve_drain_notifications(
"""
import asyncio as _aio
from EvoScientist.cli import async_notifier
from .tui_runtime import run_streaming
def _run_notification_message(text: str, notifs: list) -> None:
@@ -1239,20 +1285,21 @@ def _serve_drain_notifications(
for line_text, line_style in format_notification_lines(notifs):
console.print(line_text, style=line_style, markup=False)
# Use the current workspace from agent_holder (updated by /resume's
# Use the current workspace from runtime_state (updated by /resume's
# session-rebind callback), falling back to the startup value.
runtime_workspace = agent_holder.get("workspace_dir") or workspace_dir
runtime_workspace = runtime_state.workspace_dir or workspace_dir
meta = build_metadata(runtime_workspace, model)
tid = agent_holder["thread_id"]
tid = runtime_state.thread_id
try:
response = run_streaming(
ui_backend="cli",
agent=agent_holder["agent"],
agent=runtime_state.agent,
message=text,
thread_id=tid,
show_thinking=show_thinking,
interactive=True,
metadata=meta,
gateway=runtime_state.runtime_gateways.graph_gateway,
)
except Exception as exc:
_serve_logger.warning("Notification agent turn failed: %s", exc)
@@ -1270,22 +1317,24 @@ def _serve_drain_notifications(
async def _run_notification_message_async(text: str, notifs: list) -> None:
await _aio.to_thread(_run_notification_message, text, notifs)
async def _read_async_tasks() -> dict:
agent = agent_holder.get("agent")
thread_id = agent_holder.get("thread_id")
if agent is None or not thread_id:
return {}
try:
snap = await agent.aget_state({"configurable": {"thread_id": thread_id}})
return (snap.values or {}).get("async_tasks") or {}
except Exception:
async def _read_async_tasks() -> async_notifier.AsyncTasksState:
thread_id = runtime_state.thread_id
if not thread_id:
return {}
return await async_notifier.read_async_tasks_from_gateway(
runtime_state.runtime_gateways.graph_gateway,
GraphTarget(
local_graph=runtime_state.agent,
workspace_dir=runtime_state.workspace_dir,
),
thread_id,
)
async def _consume() -> None:
await async_notifier.consume_notifications(
run_message=_run_notification_message_async,
read_async_tasks_state=_read_async_tasks,
current_thread_id=agent_holder.get("thread_id"),
current_thread_id=runtime_state.thread_id,
)
_notif_loop: _aio.AbstractEventLoop | None = None
@@ -1406,22 +1455,22 @@ def serve(
)
console.print("[dim]Loading agent...[/dim]")
agent = _load_agent(workspace_dir=ws, config=config)
from ..sessions import generate_thread_id
tid = generate_thread_id()
runtime_gateways = create_runtime_gateways()
tid = asyncio.run(
runtime_gateways.graph_gateway.create_thread(GraphTarget(workspace_dir=ws))
)
# Mutable holder shared with _serve_process_message so ``/model``
# invoked over a channel can hot-swap the agent for subsequent
# messages. A pass-by-value parameter gets captured once at startup
# and never updated.
agent_holder: dict[str, Any] = {
"agent": agent,
"thread_id": tid,
"workspace_dir": ws,
"config": config,
}
from ..commands.base import ChannelRuntime
# Mutable runtime shared with _serve_process_message so channel slash
# commands can update the active agent/thread/workspace for subsequent
# messages.
runtime_state = ServeRuntimeState(
agent=agent,
thread_id=tid,
workspace_dir=ws,
config=config,
runtime_gateways=runtime_gateways,
)
channel_runtime = ChannelRuntime(agent=agent, thread_id=tid)
@@ -1429,13 +1478,13 @@ def serve(
# them for every inbound message. Without this hoist each message
# would allocate a fresh closure pair.
_serve_on_cmd_completed = _make_serve_cmd_completed_hook(
agent_holder, channel_runtime, config=config
runtime_state, channel_runtime, config=config
)
_serve_handle_session_resume_cb = _make_serve_handle_session_resume_cb(
agent_holder, channel_runtime, config=config
runtime_state, channel_runtime, config=config
)
_serve_start_new_session_cb = _make_serve_start_new_session_cb(
agent_holder, channel_runtime
runtime_state, channel_runtime
)
_start_channels_bus_mode(
@@ -1486,7 +1535,7 @@ def serve(
try:
_serve_process_message(
msg,
agent_holder=agent_holder,
runtime_state=runtime_state,
model=config.model,
workspace_dir=ws,
show_thinking=effective_channel_thinking,
@@ -1500,11 +1549,9 @@ def serve(
break
# Poll notification queue when idle (no channel message was pending).
from EvoScientist.cli import async_notifier
if async_notifier.has_pending_notifications(agent_holder.get("thread_id")):
if async_notifier.has_pending_notifications(runtime_state.thread_id):
_serve_drain_notifications(
agent_holder=agent_holder,
runtime_state=runtime_state,
model=config.model,
workspace_dir=ws,
show_thinking=effective_channel_thinking,
@@ -2173,27 +2220,26 @@ def _main_callback(
# Single-shot mode: wrap in persistent checkpointer
import asyncio
from ..sessions import (
generate_thread_id,
get_checkpointer,
resolve_thread_id_prefix,
)
from ..sessions import get_checkpointer
from .interactive import cmd_run
from .resume_hint import print_resume_hint
runtime_gateways = create_runtime_gateways()
graph_gateway = runtime_gateways.graph_gateway
async def _single_shot():
async with get_checkpointer() as checkpointer:
# Resolve resume target first so a bad --resume/--thread-id
# exits before the slow _load_agent() provider setup.
if thread_id:
resolved, matches = await resolve_thread_id_prefix(thread_id)
if resolved:
tid = resolved
elif matches:
resolution = await graph_gateway.resolve_thread(thread_id)
if resolution.thread_id:
tid = resolution.thread_id
elif resolution.matches:
console.print(
f"[yellow]Ambiguous thread ID '{escape(thread_id)}'. Matches:[/yellow]"
)
for s in matches:
for s in resolution.matches:
console.print(f" [cyan]{escape(s)}[/cyan]")
raise typer.Exit(1)
else:
@@ -2202,7 +2248,7 @@ def _main_callback(
)
raise typer.Exit(1)
else:
tid = generate_thread_id()
tid = await graph_gateway.create_thread()
console.print("[dim]Loading agent...[/dim]")
agent = _load_agent(
workspace_dir=workspace_dir,
@@ -2218,6 +2264,7 @@ def _main_callback(
workspace_dir=workspace_dir,
model=config.model,
ui_backend=config.ui_backend,
runtime_gateways=runtime_gateways,
)
finally:
try:
+153 -99
View File
@@ -7,8 +7,9 @@ import random
import sys
import time
from collections.abc import Callable
from dataclasses import dataclass
from datetime import datetime
from typing import Any
from typing import TYPE_CHECKING, Any
import typer # type: ignore[import-untyped]
from prompt_toolkit import PromptSession # type: ignore[import-untyped]
@@ -34,17 +35,16 @@ import EvoScientist.cli.channel as _ch_mod
from ..commands.base import Command, CommandContext
from ..commands.manager import manager as cmd_manager
from ..sessions import (
generate_thread_id,
get_checkpointer,
get_thread_messages,
get_thread_metadata,
resolve_thread_id_prefix,
short_thread_id,
thread_exists,
from ..gateway import (
GraphGateway,
GraphTarget,
RuntimeGateways,
create_runtime_gateways,
)
from ..sessions import get_checkpointer, short_thread_id
from ..stream.console import console
from ..stream.display import _fix_markdown_heading_spacing
from . import async_notifier
from ._agent_loader import BackgroundAgentLoader, MCPProgressTracker
from ._constants import (
DANGEROUS_BANNER_LABEL,
@@ -95,6 +95,19 @@ _channel_logger = logging.getLogger(__name__)
# Keeps references to fire-and-forget coroutines so they aren't GC'd mid-flight.
_background_tasks: set[asyncio.Task] = set()
if TYPE_CHECKING:
from langgraph.graph.state import CompiledStateGraph
@dataclass(frozen=True, slots=True)
class _StartupSession:
"""Resolved interactive startup session with a concrete active thread."""
thread_id: str
workspace_dir: str | None
resumed: bool
# =============================================================================
# Banner
# =============================================================================
@@ -257,6 +270,67 @@ class SlashCommandCompleter(Completer):
)
async def _resolve_startup_session(
requested_thread_id: str | None,
*,
workspace_dir: str | None,
graph_gateway: GraphGateway,
config: Any,
) -> _StartupSession:
"""Resolve/create the initial CLI session before shared REPL state exists."""
if not requested_thread_id:
return _StartupSession(
thread_id=await graph_gateway.create_thread(
GraphTarget(workspace_dir=workspace_dir)
),
workspace_dir=workspace_dir,
resumed=False,
)
resolution = await graph_gateway.resolve_thread(requested_thread_id)
if resolution.thread_id is None:
if resolution.matches:
console.print(
f"[yellow]Ambiguous thread ID '{escape(requested_thread_id)}'. "
"Matches:[/yellow]"
)
for match in resolution.matches:
console.print(f" [cyan]{match}[/cyan]")
else:
console.print(
f"[red]Thread '{escape(requested_thread_id)}' not found.[/red]"
)
return _StartupSession(
thread_id=await graph_gateway.create_thread(
GraphTarget(workspace_dir=workspace_dir)
),
workspace_dir=workspace_dir,
resumed=False,
)
resolved_thread_id = resolution.thread_id
metadata = await graph_gateway.get_thread_metadata(resolved_thread_id)
resolved_workspace = (metadata or {}).get("workspace_dir") or workspace_dir
if resolved_workspace:
from ..langgraph_dev.manager import WorkspaceMismatchError
from .commands import _sync_background_agent_server_workspace
try:
await _sync_background_agent_server_workspace(
config,
workspace_dir=resolved_workspace,
)
except WorkspaceMismatchError as exc:
console.print(f"[red]{exc}[/red]")
raise typer.Exit(1) from exc
return _StartupSession(
thread_id=resolved_thread_id,
workspace_dir=resolved_workspace,
resumed=True,
)
# =============================================================================
# Interactive & single-shot modes
# =============================================================================
@@ -353,20 +427,6 @@ def cmd_interactive(
width = console.size.width
console.print(Text("\u2500" * width, style="dim"))
# Mutable state for async loop
state: dict[str, Any] = {
"thread_id": thread_id or generate_thread_id(),
"workspace_dir": workspace_dir,
"running": True,
"resumed": False,
"ui_backend": resolved_ui_backend,
"status_started_at": datetime.now(),
"status_base_snapshot": make_empty_status_snapshot(model),
"status_snapshot": make_empty_status_snapshot(model),
"status_streaming_text": "",
"status_last_input_tokens": None,
}
from ..commands.base import ChannelRuntime
channel_runtime = ChannelRuntime()
@@ -396,6 +456,23 @@ def cmd_interactive(
on_progress=_on_mcp_progress,
)
runtime_gateways = create_runtime_gateways()
graph_gateway = runtime_gateways.graph_gateway
requested_thread_id = thread_id
# Mutable state for async loop
state: dict[str, Any] = {
"workspace_dir": workspace_dir,
"running": True,
"resumed": False,
"ui_backend": resolved_ui_backend,
"status_started_at": datetime.now(),
"status_base_snapshot": make_empty_status_snapshot(model),
"status_snapshot": make_empty_status_snapshot(model),
"status_streaming_text": "",
"status_last_input_tokens": None,
}
def _on_status_after_compact(input_tokens: int) -> None:
"""Mirror inline /compact post-update: refresh both fields so the
next status render reflects the reduced context immediately.
@@ -419,7 +496,7 @@ def cmd_interactive(
config=config,
)
async def _await_agent_ready() -> Any:
async def _await_agent_ready() -> "CompiledStateGraph":
"""Await the agent load and apply CLI-side post-load side effects.
Raises when called before ``_start_agent_load``: reloading here
@@ -475,6 +552,7 @@ def cmd_interactive(
state["thread_id"],
model_name=model,
pending_user_text=pending,
graph_gateway=graph_gateway,
)
elif state["status_last_input_tokens"] is not None:
state["status_base_snapshot"] = make_usage_status_snapshot(
@@ -485,6 +563,7 @@ def cmd_interactive(
state["status_base_snapshot"] = await build_session_status_snapshot(
state["thread_id"],
model_name=model,
graph_gateway=graph_gateway,
)
if reset_streaming_text:
state["status_streaming_text"] = ""
@@ -546,24 +625,9 @@ def cmd_interactive(
elif event_type in ("done", "error"):
_set_status_streaming_text("")
async def _resolve_thread_id(tid: str) -> str | None:
"""Resolve a (possibly partial) thread ID. Returns full ID or None."""
resolved, matches = await resolve_thread_id_prefix(tid)
if resolved:
return resolved
if matches:
console.print(
f"[yellow]Ambiguous thread ID '{escape(tid)}'. Matches:[/yellow]"
)
for s in matches:
console.print(f" [cyan]{s}[/cyan]")
return None
console.print(f"[red]Thread '{escape(tid)}' not found.[/red]")
return None
async def _render_history(thread_id: str):
"""Display conversation history for a resumed session."""
messages = await get_thread_messages(thread_id)
messages = await graph_gateway.get_thread_messages(thread_id)
if not messages:
return
@@ -639,11 +703,23 @@ def cmd_interactive(
"""Async main loop with prompt_async and channel queue checking."""
nonlocal model
async with get_checkpointer() as checkpointer:
startup = await _resolve_startup_session(
requested_thread_id,
workspace_dir=state["workspace_dir"],
graph_gateway=graph_gateway,
config=config,
)
state["thread_id"] = startup.thread_id
state["workspace_dir"] = startup.workspace_dir
state["resumed"] = startup.resumed
if startup.resumed:
state["status_started_at"] = datetime.now()
state["status_last_input_tokens"] = None
# Lifecycle callbacks (new / resume) need ``checkpointer``
# in scope — define the ``rich_ui`` adapter here rather than
# at the outer function level.
def _on_start_new_session() -> None:
async def _on_start_new_session() -> None:
"""NewCommand callback — rotate workspace (if not fixed),
issue a new thread id, reset session-scoped status fields,
and kick off background agent reload. The dispatch block
@@ -652,7 +728,9 @@ def cmd_interactive(
_ch_mod.forget_channel_origin(state.get("thread_id"))
if not workspace_fixed:
state["workspace_dir"] = _create_session_workspace(run_name)
state["thread_id"] = generate_thread_id()
state["thread_id"] = await graph_gateway.create_thread(
GraphTarget(workspace_dir=state["workspace_dir"])
)
state["resumed"] = False
state["status_started_at"] = datetime.now()
state["status_last_input_tokens"] = None
@@ -742,45 +820,6 @@ def cmd_interactive(
on_handle_session_resume=_on_handle_session_resume,
)
# 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
state["status_started_at"] = datetime.now()
state["status_last_input_tokens"] = None
if ws:
state["workspace_dir"] = ws
# CLI-startup --resume path: sync langgraph dev
# subprocess to the thread's saved workspace if it
# differs from the one we initially launched it with.
# Show a spinner during the 10-15s restart, and run
# the sync call in a worker thread so the asyncio
# event loop stays responsive.
from ..langgraph_dev.manager import WorkspaceMismatchError
from .commands import _sync_background_agent_server_workspace
try:
await _sync_background_agent_server_workspace(
config,
workspace_dir=ws,
)
except WorkspaceMismatchError as exc:
# Startup --resume into a workspace owned by
# a different EvoSci process: refuse to start
# the CLI so the user can resolve the conflict.
console.print(f"[red]{exc}[/red]")
raise typer.Exit(1) from exc
else:
# Resolution failed (ambiguous/not-found); the user's raw
# input is still seeded in state["thread_id"] from init.
# Replace with a fresh ID so a new session isn't
# checkpointed under the bad prefix.
state["thread_id"] = generate_thread_id()
# Kick off agent construction (MCP tool enumeration is the
# slow part) in the background so the banner and prompt can
# appear immediately. The status bar shows a spinner while
@@ -976,6 +1015,7 @@ def cmd_interactive(
await_agent_ready=_await_agent_ready,
on_cmd_completed=_on_channel_cmd_completed,
channel_runtime=channel_runtime,
graph_gateway=runtime_gateways.graph_gateway,
)
if _slash_handled:
# A channel-issued /new or /resume rotates the thread
@@ -1010,6 +1050,7 @@ def cmd_interactive(
on_stream_event=_handle_stream_status_event,
status_footer_builder=_stream_status_footer,
cancel_scope=_ch_mod._channel_message_cancel_scope(msg),
gateway=runtime_gateways.graph_gateway,
)
except Exception as e:
response = f"Error: {e}"
@@ -1050,9 +1091,10 @@ def cmd_interactive(
console.print(line_text, style=line_style, markup=False)
meta = build_metadata(state["workspace_dir"], model)
await _refresh_status_snapshot(text, reset_streaming_text=True)
ready_agent = await _await_agent_ready()
response = run_streaming(
ui_backend=state["ui_backend"],
agent=await _await_agent_ready(),
agent=ready_agent,
message=text,
# Falls back to live state["thread_id"] if no override is
# passed (legacy / direct-call paths). Dedup reader has no
@@ -1065,6 +1107,7 @@ def cmd_interactive(
metadata=meta,
on_stream_event=_handle_stream_status_event,
status_footer_builder=_stream_status_footer,
gateway=runtime_gateways.graph_gateway,
)
_notif_tid = target_thread_id or state["thread_id"]
if _ch_mod.publish_to_channel_origin(_notif_tid, response):
@@ -1084,9 +1127,12 @@ def cmd_interactive(
sys.stdout.write("\033[34;1m❯\033[0m ")
sys.stdout.flush()
async def _empty_async_tasks() -> async_notifier.AsyncTasksState:
return {}
async def _read_current_async_tasks(
target_thread_id: str | None,
) -> dict[str, dict]:
target_thread_id: str,
) -> async_notifier.AsyncTasksState:
"""Snapshot async_tasks from the active agent state for dedup.
Uses ``agent_loader.agent`` (the currently loaded agent) and
@@ -1095,20 +1141,22 @@ def cmd_interactive(
cannot make us read the wrong thread's state).
"""
agent = agent_loader.agent
if agent is None or not target_thread_id:
if agent is None:
return {}
try:
snap = await agent.aget_state(
{"configurable": {"thread_id": target_thread_id}}
return await async_notifier.read_async_tasks_from_gateway(
runtime_gateways.graph_gateway,
GraphTarget(
local_graph=agent,
workspace_dir=state["workspace_dir"],
),
target_thread_id,
)
return (snap.values or {}).get("async_tasks") or {}
except Exception:
return {}
async def _check_channel_queue() -> None:
"""Poll the channel + notification queues and dispatch."""
from EvoScientist.cli import async_notifier
while True:
try:
msg = _message_queue.get_nowait()
@@ -1124,6 +1172,11 @@ def cmd_interactive(
# would silently die otherwise (Fix #4).
current_tid = state.get("thread_id")
if async_notifier.has_pending_notifications(current_tid):
read_async_tasks_state = (
(lambda _tid=current_tid: _read_current_async_tasks(_tid))
if current_tid
else _empty_async_tasks
)
try:
await async_notifier.consume_notifications(
run_message=lambda text, notifs, _tid=current_tid: (
@@ -1131,9 +1184,7 @@ def cmd_interactive(
text, notifs, target_thread_id=_tid
)
),
read_async_tasks_state=lambda _tid=current_tid: (
_read_current_async_tasks(_tid)
),
read_async_tasks_state=read_async_tasks_state,
current_thread_id=current_tid,
)
except Exception:
@@ -1264,6 +1315,7 @@ def cmd_interactive(
config=config,
input_tokens_hint=state.get("status_last_input_tokens"),
channel_runtime=channel_runtime,
graph_gateway=runtime_gateways.graph_gateway,
)
await cmd_manager.execute(user_input, ctx)
@@ -1356,6 +1408,7 @@ def cmd_interactive(
metadata=meta,
on_stream_event=_handle_stream_status_event,
status_footer_builder=_stream_status_footer,
gateway=runtime_gateways.graph_gateway,
)
await _refresh_status_snapshot(reset_streaming_text=True)
console.print()
@@ -1395,7 +1448,7 @@ def cmd_interactive(
current_tid = state.get("thread_id")
if current_tid:
try:
if await thread_exists(current_tid):
if await graph_gateway.thread_exists(current_tid):
state["resume_hint_thread_id"] = current_tid
except Exception:
_channel_logger.debug(
@@ -1418,27 +1471,27 @@ def cmd_interactive(
def cmd_run(
agent: Any,
agent: "CompiledStateGraph",
prompt: str,
thread_id: str | None = None,
thread_id: str,
show_thinking: bool = True,
workspace_dir: str | None = None,
model: str | None = None,
ui_backend: str = "cli",
*,
runtime_gateways: RuntimeGateways,
) -> None:
"""Single-shot execution with streaming display.
Args:
agent: Compiled agent graph
prompt: User prompt
thread_id: Optional thread ID (generates new one if None)
thread_id: Thread ID for conversation persistence.
show_thinking: Whether to display thinking panels
workspace_dir: Per-session workspace directory path
model: Model name for checkpoint metadata
ui_backend: UI backend ('cli' or 'tui')
"""
thread_id = thread_id or generate_thread_id()
width = console.size.width
sep = Text("\u2500" * width, style="dim")
console.print(sep)
@@ -1459,6 +1512,7 @@ def cmd_run(
show_thinking=show_thinking,
interactive=False,
metadata=meta,
gateway=runtime_gateways.graph_gateway,
)
_wait_for_memory_workers_before_exit()
except Exception as e:
+3 -3
View File
@@ -41,7 +41,7 @@ class RichCLICommandUI(CommandUI):
on_force_quit: Callable[[], None] | None = None,
on_clear_chat: Callable[[], None] | None = None,
on_status_after_compact: Callable[[int], None] | None = None,
on_start_new_session: Callable[[], None] | None = None,
on_start_new_session: Callable[[], Awaitable[None]] | None = None,
on_handle_session_resume: (
Callable[[str, str | None], Awaitable[None]] | None
) = None,
@@ -188,9 +188,9 @@ class RichCLICommandUI(CommandUI):
if self._on_force_quit is not None:
self._on_force_quit()
def start_new_session(self) -> None:
async def start_new_session(self) -> None:
if self._on_start_new_session is not None:
self._on_start_new_session()
await self._on_start_new_session()
async def handle_session_resume(
self, thread_id: str, workspace_dir: str | None = None
+6 -3
View File
@@ -4,7 +4,7 @@ from __future__ import annotations
from dataclasses import dataclass, replace
from datetime import datetime
from typing import Any
from typing import TYPE_CHECKING, Any
from langchain_core.messages import AIMessage, HumanMessage
from langchain_core.messages.utils import count_tokens_approximately
@@ -14,7 +14,9 @@ from ..llm.context_window import (
resolve_context_window,
)
from ..memory.worker_activity import MemoryWorkerStatusSnapshot, memory_worker_status
from ..sessions import get_thread_messages
if TYPE_CHECKING:
from ..gateway import GraphGateway
_FALLBACK_CONTEXT_WINDOW = DEFAULT_CONTEXT_WINDOW_FALLBACK
STATUS_BAR_BG = "#171a20"
@@ -409,11 +411,12 @@ async def build_session_status_snapshot(
model_name: str | None = None,
model_obj: Any | None = None,
pending_user_text: str | None = None,
graph_gateway: GraphGateway,
) -> SessionStatusSnapshot:
"""Count current thread context and return a display snapshot."""
resolved_name = _resolve_model_name(model_name, model_obj)
window = _resolve_context_window(model_obj)
messages = list(await get_thread_messages(thread_id))
messages = list(await graph_gateway.get_thread_messages(thread_id))
pending = (pending_user_text or "").strip()
if pending:
+4
View File
@@ -6,6 +6,7 @@ from collections.abc import Callable
from dataclasses import dataclass
from typing import Any, Protocol
from ..gateway import GraphGateway
from ..stream.display import _run_streaming
@@ -31,6 +32,7 @@ class StreamingTUIBackend(Protocol):
hitl_prompt_fn: Callable[[list], list[dict] | None] | None = None,
ask_user_prompt_fn: Callable[[dict], dict] | None = None,
cancel_scope: str | None = None,
gateway: GraphGateway,
) -> str:
"""Run streaming and return final response text."""
@@ -58,6 +60,7 @@ class RichStreamingBackend:
hitl_prompt_fn: Callable[[list], list[dict] | None] | None = None,
ask_user_prompt_fn: Callable[[dict], dict] | None = None,
cancel_scope: str | None = None,
gateway: GraphGateway,
) -> str:
return _run_streaming(
agent=agent,
@@ -74,4 +77,5 @@ class RichStreamingBackend:
hitl_prompt_fn=hitl_prompt_fn,
ask_user_prompt_fn=ask_user_prompt_fn,
cancel_scope=cancel_scope,
gateway=gateway,
)
+81 -41
View File
@@ -14,7 +14,7 @@ import sys
from collections.abc import Callable
from dataclasses import dataclass
from datetime import datetime
from typing import Any, ClassVar
from typing import TYPE_CHECKING, Any, ClassVar
from rich.console import Group
from rich.text import Text
@@ -24,16 +24,15 @@ from EvoScientist.cli.widgets.thread_selector import ThreadPickerWidget
from ..commands import Command, CommandContext
from ..commands import manager as cmd_manager
from ..paths import DATA_DIR
from ..sessions import (
generate_thread_id,
get_checkpointer,
get_thread_messages,
get_thread_metadata,
resolve_thread_id_prefix,
thread_exists,
from ..gateway import (
GraphGateway,
GraphTarget,
RunRequest,
RuntimeGateways,
create_runtime_gateways,
)
from ..stream.events import stream_agent_events
from ..paths import DATA_DIR
from ..sessions import get_checkpointer
from ..stream.state import ResearchPhase, StreamState
from ._agent_loader import BackgroundAgentLoader, MCPProgressTracker
from ._constants import (
@@ -44,6 +43,12 @@ from ._constants import (
WELCOME_SLOGANS,
build_metadata,
)
from .async_notifier import (
AsyncTasksState,
consume_notifications,
has_pending_notifications,
read_async_tasks_from_gateway,
)
from .channel import (
ChannelMessage,
_auto_start_channel,
@@ -74,6 +79,9 @@ from .status_bar import (
_channel_logger = logging.getLogger(__name__)
if TYPE_CHECKING:
from langgraph.graph.state import CompiledStateGraph
def _shorten_path(path: str) -> str:
"""Shorten absolute path to a cwd-relative form (consistent with Rich CLI)."""
@@ -286,6 +294,9 @@ def run_textual_interactive(
config = get_effective_config()
runtime_gateways = create_runtime_gateways()
graph_gateway = runtime_gateways.graph_gateway
try:
from textual.app import App, ComposeResult
from textual.binding import Binding
@@ -405,6 +416,7 @@ def run_textual_interactive(
thread_id_value: str,
workspace: str | None,
checkpointer: Any,
runtime_gateways: RuntimeGateways,
channel_send_thinking_value: bool = True,
resumed: bool = False,
resume_warning: str = "",
@@ -421,6 +433,7 @@ def run_textual_interactive(
self._conversation_tid = thread_id_value
self._workspace_dir = workspace
self._checkpointer = checkpointer
self._runtime_gateways = runtime_gateways
self._channel_send_thinking = channel_send_thinking_value
self._resumed = resumed
self._resume_warning = resume_warning
@@ -539,12 +552,15 @@ def run_textual_interactive(
if widget.dismissed:
self._mcp_loader_widget = None
async def _await_agent_ready(self) -> Any:
async def _await_agent_ready(self) -> CompiledStateGraph:
"""Await the agent load, auto-retrying on cold-start or failure."""
if self._agent_loader.needs_restart:
self._start_background_agent_load(self._workspace_dir)
return await self._agent_loader.await_ready()
def _graph_gateway(self) -> GraphGateway:
return self._runtime_gateways.graph_gateway
# ── CommandUI implementation ─────────────────────────
def append_system(self, text: str, style: str = "dim") -> None:
@@ -640,14 +656,18 @@ def run_textual_interactive(
def request_quit(self) -> None:
self.action_request_quit()
def start_new_session(self) -> None:
async def start_new_session(self) -> None:
# Clear all widgets except #welcome
self.clear_chat()
_ch_mod.forget_channel_origin(self._conversation_tid)
if not workspace_fixed:
self._workspace_dir = create_session_workspace(run_name)
self._conversation_tid = generate_thread_id()
self._conversation_tid = (
await self._runtime_gateways.graph_gateway.create_thread(
GraphTarget(workspace_dir=self._workspace_dir)
)
)
# Background reload: next user message awaits it.
self._start_background_agent_load(self._workspace_dir)
self._status_started_at = datetime.now()
@@ -870,8 +890,6 @@ def run_textual_interactive(
def _poll_channel_queue(self) -> None:
"""Poll the channel + notification queues (every 100ms)."""
from EvoScientist.cli import async_notifier
try:
msg = _message_queue.get_nowait()
except queue.Empty:
@@ -892,7 +910,7 @@ def run_textual_interactive(
# so that the next poll tick cannot schedule a second consumer before
# the first one has a chance to run (fixes overlapping-turn bug).
if (
async_notifier.has_pending_notifications(self._conversation_tid)
has_pending_notifications(self._conversation_tid)
and not self._busy
and not self._notification_consuming
):
@@ -909,12 +927,10 @@ def run_textual_interactive(
``asyncio.ensure_future(...)`` scheduled by ``_poll_channel_queue``
and silently kill notification + channel dispatch.
"""
from EvoScientist.cli import async_notifier
target_tid = self._conversation_tid
try:
try:
await async_notifier.consume_notifications(
await consume_notifications(
run_message=lambda text, notifs: self._inject_notification_tui(
text, notifs, target_thread_id=target_tid
),
@@ -987,9 +1003,7 @@ def run_textual_interactive(
self._run_task = asyncio.ensure_future(_run_and_publish())
async def _read_async_tasks_tui(
self, target_thread_id: str | None
) -> dict[str, dict]:
async def _read_async_tasks_tui(self, target_thread_id: str) -> AsyncTasksState:
"""Read async_tasks from agent state for dedup, against a frozen tid.
``target_thread_id`` is captured by ``_consume_notifications_tui`` at
@@ -997,13 +1011,17 @@ def run_textual_interactive(
make us read the wrong thread's state.
"""
agent = self._agent_loader.agent
if agent is None or not target_thread_id:
if agent is None:
return {}
try:
snap = await agent.aget_state(
{"configurable": {"thread_id": target_thread_id}}
return await read_async_tasks_from_gateway(
self._graph_gateway(),
GraphTarget(
local_graph=agent,
workspace_dir=self._workspace_dir,
),
target_thread_id,
)
return (snap.values or {}).get("async_tasks") or {}
except Exception:
return {}
@@ -1312,6 +1330,7 @@ def run_textual_interactive(
metadata = build_metadata(self._workspace_dir, self._current_model)
response = ""
agent = await self._await_agent_ready()
async def _remove_w(w: Widget | None) -> None:
"""Safely remove a transient indicator widget."""
@@ -1472,6 +1491,7 @@ def run_textual_interactive(
_MAX_HITL_ROUNDS = 50
_stream_input: Any = user_text # str or Command for HITL resume
graph_gateway = self._graph_gateway()
for _hitl_round in range(_MAX_HITL_ROUNDS):
if is_stream_cancel_requested(cancel_scope):
@@ -1486,11 +1506,16 @@ def run_textual_interactive(
summarization_w = None
try:
_anchor_engaged = False
async for event in stream_agent_events(
self._agent_loader.agent,
_stream_input,
thread_id_override or self._conversation_tid,
metadata=metadata,
async for event in graph_gateway.stream_events(
RunRequest(
message=_stream_input,
thread_id=thread_id_override or self._conversation_tid,
metadata=metadata,
target=GraphTarget(
local_graph=agent,
workspace_dir=self._workspace_dir,
),
)
):
if is_stream_cancel_requested(cancel_scope):
response = await _mark_cancelled_response()
@@ -2241,6 +2266,7 @@ def run_textual_interactive(
await_agent_ready=self._await_agent_ready,
on_cmd_completed=self._on_channel_cmd_completed,
channel_runtime=self._channel_runtime,
graph_gateway=self._runtime_gateways.graph_gateway,
)
if _slash_handled:
# A channel-issued /new or /resume rotates the thread in
@@ -2715,8 +2741,12 @@ def run_textual_interactive(
# Only gate on agent readiness for commands that need it —
# recovery commands like ``/mcp add`` must run even when
# ``_await_agent_ready`` would hang on a broken MCP load.
cmd, cmd_args = cmd_manager.resolve(command) or (None, [])
parsed = cmd_manager.resolve(command)
cmd = None
cmd_args: list[str] = []
agent = None
if parsed is not None:
cmd, cmd_args = parsed
if cmd is not None and cmd.needs_agent(cmd_args):
try:
agent = await self._await_agent_ready()
@@ -2731,9 +2761,12 @@ def run_textual_interactive(
checkpointer=self._checkpointer,
input_tokens_hint=self._status_last_input_tokens,
channel_runtime=self._channel_runtime,
graph_gateway=self._runtime_gateways.graph_gateway,
)
if await cmd_manager.execute(command, ctx):
if cmd is None:
return
await _sync_tui_command_completion(
self,
ctx,
@@ -2757,7 +2790,9 @@ def run_textual_interactive(
skipped — they are difficult to faithfully reproduce from
checkpoint data.
"""
messages = await get_thread_messages(thread_id_value)
messages = await self._runtime_gateways.graph_gateway.get_thread_messages(
thread_id_value
)
if not messages:
return
@@ -2912,6 +2947,7 @@ def run_textual_interactive(
self._conversation_tid,
model_name=self._current_model,
pending_user_text=pending,
graph_gateway=self._runtime_gateways.graph_gateway,
)
elif self._status_last_input_tokens is not None:
self._status_base_snapshot = make_usage_status_snapshot(
@@ -2922,6 +2958,7 @@ def run_textual_interactive(
self._status_base_snapshot = await build_session_status_snapshot(
self._conversation_tid,
model_name=self._current_model,
graph_gateway=self._runtime_gateways.graph_gateway,
)
if reset_streaming_text:
self._status_streaming_text = ""
@@ -3132,9 +3169,9 @@ def run_textual_interactive(
resumed = False
resume_warning = ""
if thread_id:
resolved, matches = await resolve_thread_id_prefix(thread_id)
if resolved:
meta = await get_thread_metadata(resolved)
resolution = await graph_gateway.resolve_thread(thread_id)
if resolution.thread_id:
meta = await graph_gateway.get_thread_metadata(resolution.thread_id)
ws = (meta or {}).get("workspace_dir", "")
mismatch_aborted = False
if ws:
@@ -3193,19 +3230,21 @@ def run_textual_interactive(
"workspace conflict. Starting new session."
)
else:
effective_thread_id = resolved
effective_thread_id = resolution.thread_id
resumed = True
elif matches:
elif resolution.matches:
resume_warning = (
f"Thread prefix '{thread_id}' is ambiguous "
f"({', '.join(matches)}). Starting new session."
f"({', '.join(resolution.matches)}). Starting new session."
)
else:
resume_warning = (
f"Thread '{thread_id}' not found. Starting new session."
)
if not effective_thread_id:
effective_thread_id = generate_thread_id()
effective_thread_id = await graph_gateway.create_thread(
GraphTarget(workspace_dir=effective_workspace)
)
# The TUI opens instantly and starts MCP loading in the
# background; ``on_mount`` in the app kicks off the real
@@ -3214,6 +3253,7 @@ def run_textual_interactive(
thread_id_value=effective_thread_id,
workspace=effective_workspace,
checkpointer=checkpointer,
runtime_gateways=runtime_gateways,
channel_send_thinking_value=channel_send_thinking,
resumed=resumed,
resume_warning=resume_warning,
@@ -3230,7 +3270,7 @@ def run_textual_interactive(
hint_tid: str | None = None
if exit_tid:
try:
if await thread_exists(exit_tid):
if await graph_gateway.thread_exists(exit_tid):
hint_tid = exit_tid
except Exception:
_channel_logger.debug(
+4
View File
@@ -5,6 +5,7 @@ from __future__ import annotations
from collections.abc import Callable
from typing import Any
from ..gateway import GraphGateway
from ..stream.console import console
from .tui_backends import RichStreamingBackend, StreamingTUIBackend
@@ -79,6 +80,7 @@ def run_streaming(
hitl_prompt_fn: Callable[[list], list[dict] | None] | None = None,
ask_user_prompt_fn: Callable[[dict], dict] | None = None,
cancel_scope: str | None = None,
gateway: GraphGateway,
) -> str:
"""Run streaming with the selected backend."""
backend = get_backend(ui_backend, warn_fallback=True)
@@ -98,6 +100,7 @@ def run_streaming(
hitl_prompt_fn=hitl_prompt_fn,
ask_user_prompt_fn=ask_user_prompt_fn,
cancel_scope=cancel_scope,
gateway=gateway,
)
except RuntimeError:
requested = normalize_ui_backend(ui_backend)
@@ -120,5 +123,6 @@ def run_streaming(
hitl_prompt_fn=hitl_prompt_fn,
ask_user_prompt_fn=ask_user_prompt_fn,
cancel_scope=cancel_scope,
gateway=gateway,
)
raise
+7 -3
View File
@@ -2,7 +2,10 @@ from __future__ import annotations
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from typing import Any, ClassVar, Protocol, runtime_checkable
from typing import TYPE_CHECKING, Any, ClassVar, Protocol, runtime_checkable
if TYPE_CHECKING:
from ..gateway import GraphGateway
@dataclass
@@ -53,7 +56,7 @@ class CommandUI(Protocol):
def clear_chat(self) -> None: ...
def request_quit(self) -> None: ...
def force_quit(self) -> None: ...
def start_new_session(self) -> None: ...
async def start_new_session(self) -> None: ...
async def handle_session_resume(
self, thread_id: str, workspace_dir: str | None = None
) -> None: ...
@@ -67,7 +70,7 @@ class ChannelRuntime:
agent: Any = None
thread_id: str | None = None
def bind(self, agent: Any, thread_id: str | None) -> None:
def bind(self, agent: Any, thread_id: str) -> None:
self.agent = agent
self.thread_id = thread_id
@@ -87,6 +90,7 @@ class CommandContext:
checkpointer: Any = None
config: Any = None
channel_runtime: ChannelRuntime | None = None
graph_gateway: GraphGateway | None = None
command_error: str | None = None
# Real LLM input token count from last usage_metadata (includes system
# prompt + tool schemas). Used by /compact for accurate display.
+12 -7
View File
@@ -2,10 +2,14 @@ from __future__ import annotations
import asyncio
import logging
from typing import Any
from collections.abc import Awaitable, Callable
from typing import TYPE_CHECKING, Any
from .base import CommandUI
if TYPE_CHECKING:
from ..gateway import GraphGateway
_logger = logging.getLogger(__name__)
@@ -21,14 +25,17 @@ class ChannelCommandUI(CommandUI):
def __init__(
self,
channel_msg: Any,
*,
graph_gateway: GraphGateway,
append_system_callback: Any = None,
start_new_session_callback: Any = None,
start_new_session_callback: Callable[[], Awaitable[None]] | None = None,
handle_session_resume_callback: Any = None,
):
self.msg = channel_msg
self.append_system_callback = append_system_callback
self.start_new_session_callback = start_new_session_callback
self.handle_session_resume_callback = handle_session_resume_callback
self.graph_gateway = graph_gateway
self._system_buffer: list[str] = []
def _queue_system(
@@ -174,9 +181,9 @@ class ChannelCommandUI(CommandUI):
def force_quit(self) -> None:
self.request_quit()
def start_new_session(self) -> None:
async def start_new_session(self) -> None:
if self.start_new_session_callback:
self.start_new_session_callback()
await self.start_new_session_callback()
else:
self.append_system(
"New session requested. Please restart the channel link or use /new if supported."
@@ -188,11 +195,9 @@ class ChannelCommandUI(CommandUI):
mirror_local = self.handle_session_resume_callback is None
if self.handle_session_resume_callback:
await self.handle_session_resume_callback(thread_id, workspace_dir)
from ..sessions import get_thread_messages
lines = [f"Resumed session: {thread_id}"]
try:
messages = await get_thread_messages(thread_id)
messages = await self.graph_gateway.get_thread_messages(thread_id)
except Exception as exc:
_logger.exception(
"Failed to load saved history for resumed thread %s",
+37 -47
View File
@@ -5,10 +5,17 @@ from typing import ClassVar
from rich.table import Table
from ...gateway import GraphGateway, GraphTarget
from ..base import Argument, Command, CommandContext
from ..manager import manager
def _graph_gateway(ctx: CommandContext) -> GraphGateway:
if ctx.graph_gateway is None:
raise RuntimeError("Session commands require a graph_gateway")
return ctx.graph_gateway
class CompactCommand(Command):
"""Compact conversation to free context."""
@@ -36,8 +43,12 @@ class CompactCommand(Command):
try:
result = await compact_conversation(
agent=ctx.agent,
graph_gateway=_graph_gateway(ctx),
thread_id=ctx.thread_id,
target=GraphTarget(
local_graph=ctx.agent,
workspace_dir=ctx.workspace_dir,
),
input_tokens_hint=ctx.input_tokens_hint,
)
finally:
@@ -73,9 +84,10 @@ class ThreadsCommand(Command):
description = "List recent sessions"
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
from ...sessions import _format_relative_time, list_threads
from ...sessions import _format_relative_time, short_thread_id
threads = await list_threads(
gateway = _graph_gateway(ctx)
threads = await gateway.list_threads(
limit=0,
include_message_count=True,
include_preview=True,
@@ -98,8 +110,6 @@ class ThreadsCommand(Command):
table.add_column("Model", style="dim")
table.add_column("Last Used", style="dim")
from ...sessions import short_thread_id
for thread in threads:
thread_id_value = thread["thread_id"]
marker = " *" if thread_id_value == ctx.thread_id else ""
@@ -137,14 +147,10 @@ class ResumeCommand(Command):
]
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
from ...sessions import (
get_thread_metadata,
list_threads,
)
gateway = _graph_gateway(ctx)
arg = args[0] if args else ""
if not arg:
threads = await list_threads(
threads = await gateway.list_threads(
limit=0,
include_message_count=True,
include_preview=True,
@@ -168,7 +174,7 @@ class ResumeCommand(Command):
if not resolved:
return
metadata = await get_thread_metadata(resolved)
metadata = await gateway.get_thread_metadata(resolved)
restored_workspace = (metadata or {}).get("workspace_dir", "")
if restored_workspace:
ctx.workspace_dir = restored_workspace
@@ -180,21 +186,16 @@ class ResumeCommand(Command):
await ctx.ui.handle_session_resume(resolved, restored_workspace)
async def _resolve_thread_id(self, prefix: str, ctx: CommandContext) -> str | None:
from ...sessions import find_similar_threads, thread_exists
resolution = await _graph_gateway(ctx).resolve_thread(prefix)
if resolution.thread_id:
return resolution.thread_id
if await thread_exists(prefix):
return prefix
similar = await find_similar_threads(prefix)
if len(similar) == 1:
return similar[0]
if len(similar) > 1:
if resolution.matches:
ctx.ui.append_system(
f"Ambiguous thread ID '{prefix}'. Use a longer prefix.",
style="yellow",
)
for thread in similar:
for thread in resolution.matches:
ctx.ui.append_system(f" - {thread}", style="dim")
return None
@@ -209,7 +210,7 @@ class NewCommand(Command):
description = "Start a new session"
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
ctx.ui.start_new_session()
await ctx.ui.start_new_session()
class ClearCommand(Command):
@@ -237,16 +238,10 @@ class DeleteCommand(Command):
]
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
from ...sessions import (
delete_thread,
find_similar_threads,
list_threads,
thread_exists,
)
gateway = _graph_gateway(ctx)
arg = args[0] if args else ""
if not arg:
threads = await list_threads(
threads = await gateway.list_threads(
limit=0,
include_message_count=True,
include_preview=True,
@@ -266,22 +261,17 @@ class DeleteCommand(Command):
arg = selected
# Resolve thread_id
resolved = None
if await thread_exists(arg):
resolved = arg
else:
similar = await find_similar_threads(arg)
if len(similar) == 1:
resolved = similar[0]
elif len(similar) > 1:
ctx.ui.append_system(
f"Ambiguous thread ID '{arg}'. Use a longer prefix.",
style="yellow",
)
for thread in similar:
ctx.ui.append_system(f" - {thread}", style="dim")
return
resolution = await gateway.resolve_thread(arg)
if resolution.matches:
ctx.ui.append_system(
f"Ambiguous thread ID '{arg}'. Use a longer prefix.",
style="yellow",
)
for thread in resolution.matches:
ctx.ui.append_system(f" - {thread}", style="dim")
return
resolved = resolution.thread_id
if not resolved:
ctx.ui.append_system(f"Session '{arg}' not found.", style="red")
return
@@ -293,7 +283,7 @@ class DeleteCommand(Command):
)
return
deleted = await delete_thread(resolved)
deleted = await gateway.delete_thread(resolved)
if deleted:
ctx.ui.append_system(f"Deleted session {resolved}.", style="green")
else:
+48
View File
@@ -0,0 +1,48 @@
"""Graph/thread gateway abstractions.
The gateway package is the migration seam between UI surfaces and graph
execution. CLI, TUI, channels, and future frontends should depend on this
package for thread/run operations instead of reaching directly into
``sessions.py``, ``stream.events``, or the LangGraph SDK.
"""
from .local import LocalGraphGateway, LocalThreadStore
from .runtime import (
RuntimeGatewayBackend,
RuntimeGateways,
create_runtime_gateways,
)
from .server import (
LangGraphServerGateway,
LangGraphServerThreadStore,
)
from .types import (
DEFAULT_GRAPH_ID,
GraphEvent,
GraphGateway,
GraphRunInput,
GraphStateValues,
GraphTarget,
RunRequest,
ThreadResolution,
ThreadStore,
)
__all__ = [
"DEFAULT_GRAPH_ID",
"GraphEvent",
"GraphGateway",
"GraphRunInput",
"GraphStateValues",
"GraphTarget",
"LangGraphServerGateway",
"LangGraphServerThreadStore",
"LocalGraphGateway",
"LocalThreadStore",
"RunRequest",
"RuntimeGatewayBackend",
"RuntimeGateways",
"ThreadResolution",
"ThreadStore",
"create_runtime_gateways",
]
+194
View File
@@ -0,0 +1,194 @@
"""Local in-process gateway backend preserving current behavior."""
from __future__ import annotations
from collections.abc import AsyncIterator
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, Any
from .. import sessions as session_store
from .types import (
GraphEvent,
GraphStateValues,
GraphTarget,
RunRequest,
ThreadResolution,
ThreadStore,
)
if TYPE_CHECKING:
from langgraph.graph.state import CompiledStateGraph
@dataclass(frozen=True, slots=True)
class LocalThreadStore:
"""Thread store backed by the current ``sessions.py`` module."""
def generate_thread_id(self) -> str:
return session_store.generate_thread_id()
async def list_threads(
self,
*,
limit: int = 20,
include_message_count: bool = False,
include_preview: bool = False,
) -> list[dict[str, Any]]:
return await session_store.list_threads(
limit=limit,
include_message_count=include_message_count,
include_preview=include_preview,
)
async def resolve_thread_id_prefix(
self,
thread_id_or_prefix: str,
) -> tuple[str | None, list[str]]:
return await session_store.resolve_thread_id_prefix(thread_id_or_prefix)
async def get_thread_metadata(self, thread_id: str) -> dict[str, Any] | None:
return await session_store.get_thread_metadata(thread_id)
async def get_thread_messages(self, thread_id: str) -> list[Any]:
return await session_store.get_thread_messages(thread_id)
async def thread_exists(self, thread_id: str) -> bool:
return await session_store.thread_exists(thread_id)
async def delete_thread(self, thread_id: str) -> bool:
return await session_store.delete_thread(thread_id)
@dataclass(slots=True)
class LocalGraphGateway:
"""Gateway backed by the current in-process graph and session helpers."""
thread_store: ThreadStore = field(default_factory=LocalThreadStore)
async def create_thread(
self,
target: GraphTarget | None = None,
*,
metadata: dict[str, Any] | None = None,
) -> str:
return self.thread_store.generate_thread_id()
async def list_threads(
self,
*,
limit: int = 20,
include_message_count: bool = False,
include_preview: bool = False,
target: GraphTarget | None = None,
) -> list[dict[str, Any]]:
return await self.thread_store.list_threads(
limit=limit,
include_message_count=include_message_count,
include_preview=include_preview,
)
async def resolve_thread(
self,
thread_id_or_prefix: str,
target: GraphTarget | None = None,
) -> ThreadResolution:
resolved, matches = await self.thread_store.resolve_thread_id_prefix(
thread_id_or_prefix
)
return ThreadResolution(resolved, tuple(matches))
async def get_thread_metadata(
self,
thread_id: str,
target: GraphTarget | None = None,
) -> dict[str, Any] | None:
return await self.thread_store.get_thread_metadata(thread_id)
async def get_thread_messages(
self,
thread_id: str,
target: GraphTarget | None = None,
) -> list[Any]:
return await self.thread_store.get_thread_messages(thread_id)
async def thread_exists(
self,
thread_id: str,
target: GraphTarget | None = None,
) -> bool:
return await self.thread_store.thread_exists(thread_id)
async def delete_thread(
self,
thread_id: str,
target: GraphTarget | None = None,
) -> bool:
return await self.thread_store.delete_thread(thread_id)
async def clone_thread(
self,
source_thread_id: str,
*,
metadata: dict[str, Any] | None = None,
target: GraphTarget | None = None,
) -> str:
raise NotImplementedError("LocalGraphGateway does not support thread cloning.")
def stream_events(self, request: RunRequest) -> AsyncIterator[GraphEvent]:
target = request.target
local_graph = self._require_local_graph(target)
if target is None:
raise RuntimeError("LocalGraphGateway requires GraphTarget.local_graph")
return self._stream_events(local_graph, target, request)
async def _stream_events(
self,
local_graph: CompiledStateGraph,
target: GraphTarget,
request: RunRequest,
) -> AsyncIterator[GraphEvent]:
from ..stream.events import stream_agent_events
inner = stream_agent_events(
local_graph,
request.message,
request.thread_id,
metadata=request.metadata,
media=request.media,
)
try:
async for event in inner:
yield event
finally:
await inner.aclose()
async def get_state_values(
self,
target: GraphTarget,
thread_id: str,
) -> GraphStateValues:
local_graph = self._require_local_graph(target)
snapshot = await local_graph.aget_state(
{"configurable": {"thread_id": thread_id}}
)
values: GraphStateValues = snapshot.values
return values
async def update_state_values(
self,
target: GraphTarget,
thread_id: str,
values: GraphStateValues,
) -> None:
local_graph = self._require_local_graph(target)
as_node = "model" if "_summarization_event" in values else None
await local_graph.aupdate_state(
{"configurable": {"thread_id": thread_id}},
values,
as_node=as_node,
)
def _require_local_graph(self, target: GraphTarget | None) -> CompiledStateGraph:
if target is None or target.local_graph is None:
raise RuntimeError("LocalGraphGateway requires GraphTarget.local_graph")
return target.local_graph
+68
View File
@@ -0,0 +1,68 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Literal
from .local import LocalGraphGateway, LocalThreadStore
from .server import (
DEFAULT_GRAPH_ID,
LangGraphClientFactory,
LangGraphServerGateway,
LangGraphServerThreadStore,
)
from .types import GraphGateway, ThreadStore
RuntimeGatewayBackend = Literal["local", "langgraph_server"]
@dataclass(frozen=True, slots=True)
class RuntimeGateways:
"""Gateway handles for one CLI/TUI/serve runtime."""
thread_store: ThreadStore
graph_gateway: GraphGateway
def create_runtime_gateways(
*,
backend: RuntimeGatewayBackend = "local",
base_url: str | None = None,
graph_id: str = DEFAULT_GRAPH_ID,
headers: dict[str, str] | None = None,
client_factory: LangGraphClientFactory | None = None,
) -> RuntimeGateways:
"""Create gateway handles for CLI/TUI/serve execution."""
if backend == "langgraph_server":
if base_url is None:
raise ValueError("base_url is required for langgraph_server gateways")
if client_factory is not None:
server_thread_store = LangGraphServerThreadStore(
base_url=base_url,
graph_id=graph_id,
headers=headers,
client_factory=client_factory,
)
else:
server_thread_store = LangGraphServerThreadStore(
base_url=base_url,
graph_id=graph_id,
headers=headers,
)
return RuntimeGateways(
thread_store=server_thread_store,
graph_gateway=LangGraphServerGateway(
server_thread_store,
graph_id=graph_id,
),
)
if backend != "local":
raise ValueError(f"Unsupported runtime gateway backend: {backend}")
local_thread_store = LocalThreadStore()
return RuntimeGateways(
thread_store=local_thread_store,
graph_gateway=LocalGraphGateway(thread_store=local_thread_store),
)
+737
View File
@@ -0,0 +1,737 @@
"""LangGraph server-backed gateway implementation."""
from __future__ import annotations
import asyncio
import uuid
from collections.abc import AsyncIterator, Callable, Mapping
from dataclasses import dataclass, field
from datetime import UTC, datetime
from typing import Any
from langchain_core.messages import BaseMessage, convert_to_messages, messages_from_dict
from langgraph.types import Command
from langgraph_sdk import get_client
from langgraph_sdk._async.stream import AsyncThreadStream
from langgraph_sdk.client import LangGraphClient
from langgraph_sdk.errors import NotFoundError
from langgraph_sdk.schema import Thread, ThreadState
from ..sessions import _apply_summarization_event
from ..stream.emitter import StreamEventEmitter
from ..stream.events import (
_SubagentRegistry,
_V3EventProcessor,
build_agent_stream_input,
)
from ..stream.summarization import _find_summarization_event_payload
from ..stream.v3_payloads import _as_raw_map, _event_namespace
from .types import (
DEFAULT_GRAPH_ID,
GraphEvent,
GraphStateValues,
GraphTarget,
RunRequest,
ThreadResolution,
ThreadStore,
)
_THREAD_SEARCH_LIMIT = 1000
_RUN_SUBSCRIBE_CHANNELS = [
"messages",
"tools",
"updates",
"values",
"tasks",
"lifecycle",
"input",
]
LangGraphClientFactory = Callable[
[str, Mapping[str, str] | None],
LangGraphClient,
]
def _default_client_factory(
base_url: str,
headers: Mapping[str, str] | None,
) -> LangGraphClient:
return get_client(url=base_url, headers=headers)
def _thread_metadata(thread: Thread) -> dict[str, Any]:
metadata = thread.get("metadata")
return dict(metadata) if isinstance(metadata, dict) else {}
def _build_thread_metadata(
*,
graph_id: str,
workspace_dir: str | None,
metadata: Mapping[str, Any] | None = None,
) -> dict[str, Any]:
merged = dict(metadata or {})
merged["graph_id"] = graph_id
if graph_id == DEFAULT_GRAPH_ID:
merged["agent_name"] = DEFAULT_GRAPH_ID
else:
merged.pop("agent_name", None)
if workspace_dir is not None:
merged["workspace_dir"] = workspace_dir
merged.setdefault("updated_at", datetime.now(UTC).isoformat())
return merged
def _thread_preview(messages: list[BaseMessage]) -> str:
for message in reversed(messages):
if getattr(message, "type", None) != "human":
continue
content = message.content
if isinstance(content, str):
return content.strip().replace("\n", " ")[:120]
if isinstance(content, list):
text_parts = [
str(block.get("text", ""))
for block in content
if isinstance(block, dict) and block.get("type") == "text"
]
if text := " ".join(part for part in text_parts if part).strip():
return text.replace("\n", " ")[:120]
return ""
def _is_uuid(value: str) -> bool:
try:
uuid.UUID(value)
except ValueError:
return False
return True
def _input_requested_event_from_interrupt(
interrupt: Mapping[str, object],
) -> dict[str, Any]:
return {
"type": "event",
"method": "input.requested",
"params": {
"namespace": interrupt.get("namespace") or [],
"data": {
"interrupt_id": interrupt.get("interrupt_id")
or interrupt.get("id")
or "default",
"value": interrupt.get("value"),
},
},
}
def _state_interrupts(state: ThreadState) -> list[Mapping[str, object]]:
interrupts = state.get("interrupts")
if not isinstance(interrupts, list):
return []
return [interrupt for interrupt in interrupts if isinstance(interrupt, Mapping)]
def _is_interrupt_event(event: Mapping[str, object]) -> bool:
return event.get("type") in {"interrupt", "ask_user"}
def _messages_from_state(state: ThreadState) -> list[BaseMessage]:
values = state.get("values")
if not isinstance(values, dict):
return []
raw_messages = values.get("messages")
if not isinstance(raw_messages, list):
return []
event = values.get("_summarization_event")
summarization_event = dict(event) if isinstance(event, Mapping) else None
effective_messages = _apply_summarization_event(
raw_messages,
summarization_event,
)
try:
return list(convert_to_messages(effective_messages))
except ValueError:
return messages_from_dict(
[message for message in effective_messages if isinstance(message, dict)]
)
@dataclass(frozen=True, slots=True)
class LangGraphServerThreadStore(ThreadStore):
"""Thread store backed by the LangGraph server Threads API."""
base_url: str
graph_id: str = DEFAULT_GRAPH_ID
headers: Mapping[str, str] | None = None
client_factory: LangGraphClientFactory = _default_client_factory
_client: LangGraphClient = field(init=False, repr=False)
def __post_init__(self) -> None:
object.__setattr__(
self,
"_client",
self.client_factory(self.base_url, self.headers),
)
@property
def client(self) -> LangGraphClient:
return self._client
def generate_thread_id(self) -> str:
return str(uuid.uuid4())
def _target_graph_id(self, graph_id: str | None = None) -> str:
return graph_id or self.graph_id
async def create_thread(
self,
graph_id: str | None = None,
*,
metadata: Mapping[str, Any] | None = None,
workspace_dir: str | None = None,
) -> str:
target_graph_id = self._target_graph_id(graph_id)
thread = await self.client.threads.create(
graph_id=target_graph_id,
metadata=_build_thread_metadata(
graph_id=target_graph_id,
workspace_dir=workspace_dir,
metadata=metadata,
),
)
return thread["thread_id"]
async def ensure_thread_exists(
self,
thread_id: str,
graph_id: str | None = None,
*,
metadata: Mapping[str, Any] | None = None,
workspace_dir: str | None = None,
) -> None:
target_graph_id = self._target_graph_id(graph_id)
await self.client.threads.create(
thread_id=thread_id,
graph_id=target_graph_id,
metadata=_build_thread_metadata(
graph_id=target_graph_id,
workspace_dir=workspace_dir,
metadata=metadata,
),
if_exists="do_nothing",
)
async def list_threads(
self,
*,
limit: int = 20,
include_message_count: bool = False,
include_preview: bool = False,
graph_id: str | None = None,
) -> list[dict[str, Any]]:
target_graph_id = self._target_graph_id(graph_id)
threads = await self._search_threads(
target_graph_id=target_graph_id,
limit=limit,
)
rows: list[dict[str, Any]] = []
for thread in threads:
thread_id = thread["thread_id"]
metadata = _thread_metadata(thread)
row: dict[str, Any] = {
"thread_id": thread_id,
"created_at": thread.get("created_at"),
"updated_at": thread.get("updated_at"),
"workspace_dir": metadata.get("workspace_dir"),
"model": metadata.get("model"),
"metadata": metadata,
}
if include_message_count or include_preview:
messages = await self.get_thread_messages(thread_id)
if include_message_count:
row["message_count"] = len(messages)
if include_preview:
row["preview"] = _thread_preview(messages)
rows.append(row)
return rows
async def resolve_thread_id_prefix(
self,
thread_id_or_prefix: str,
graph_id: str | None = None,
) -> tuple[str | None, list[str]]:
target_graph_id = self._target_graph_id(graph_id)
if _is_uuid(thread_id_or_prefix):
try:
thread = await self.client.threads.get(thread_id_or_prefix)
if _thread_metadata(thread).get("graph_id") == target_graph_id:
return thread["thread_id"], []
except NotFoundError:
pass
threads = await self._search_threads(target_graph_id=target_graph_id)
matches = sorted(
thread["thread_id"]
for thread in threads
if thread["thread_id"].startswith(thread_id_or_prefix)
)
if len(matches) == 1:
return matches[0], []
return None, matches
async def _search_threads(
self,
*,
target_graph_id: str,
limit: int | None = None,
) -> list[Thread]:
if limit is not None and limit > 0:
return await self._search_thread_page(
target_graph_id=target_graph_id,
limit=limit,
)
return await self._search_all_threads(target_graph_id=target_graph_id)
async def _search_thread_page(
self,
*,
target_graph_id: str,
limit: int,
offset: int = 0,
) -> list[Thread]:
return await self.client.threads.search(
metadata={"graph_id": target_graph_id},
limit=limit,
offset=offset,
sort_by="updated_at",
sort_order="desc",
)
async def _search_all_threads(self, *, target_graph_id: str) -> list[Thread]:
threads: list[Thread] = []
offset = 0
while True:
page = await self._search_thread_page(
target_graph_id=target_graph_id,
limit=_THREAD_SEARCH_LIMIT,
offset=offset,
)
threads.extend(page)
if len(page) < _THREAD_SEARCH_LIMIT:
break
offset += _THREAD_SEARCH_LIMIT
return threads
async def get_thread_metadata(self, thread_id: str) -> dict[str, Any] | None:
try:
thread = await self.client.threads.get(thread_id)
except NotFoundError:
return None
return _thread_metadata(thread)
async def get_thread_messages(self, thread_id: str) -> list[BaseMessage]:
try:
state = await self.client.threads.get_state(thread_id)
except NotFoundError:
return []
return _messages_from_state(state)
async def thread_exists(self, thread_id: str) -> bool:
try:
await self.client.threads.get(thread_id)
except NotFoundError:
return False
return True
async def delete_thread(self, thread_id: str) -> bool:
try:
await self.client.threads.delete(thread_id)
except NotFoundError:
return False
return True
async def clone_thread(
self,
source_thread_id: str,
*,
metadata: dict[str, Any] | None = None,
) -> str:
copy_response: object = await self.client.threads.copy(source_thread_id)
if not isinstance(copy_response, Mapping):
raise RuntimeError(
"LangGraph thread copy did not return a cloned thread id"
)
cloned_thread_id = copy_response.get("thread_id")
if not isinstance(cloned_thread_id, str) or not cloned_thread_id:
raise RuntimeError(
"LangGraph thread copy did not return a cloned thread id"
)
if metadata:
await self.client.threads.update(
cloned_thread_id,
metadata=metadata,
)
return cloned_thread_id
@dataclass(slots=True)
class _ServerSubagentTracker:
"""Infer subagent start/end events from LangGraph server namespaces."""
emitter: StreamEventEmitter
registry: _SubagentRegistry
_active: dict[tuple[str, ...], tuple[str, str | None]] = field(default_factory=dict)
def process(self, event: Mapping[str, Any]) -> list[dict[str, Any]]:
events: list[dict[str, Any]] = []
namespace = tuple(_event_namespace(event))
if namespace:
events.extend(self._ensure_registered(namespace[:1], tool_call_id=None))
method = event.get("method")
params = _as_raw_map(event.get("params"))
data = _as_raw_map(params.get("data")) if params is not None else None
if data is None:
return events
if method == "lifecycle":
phase = data.get("event")
if phase == "started" and namespace:
events.extend(self._ensure_registered(namespace, tool_call_id=None))
elif phase in ("completed", "failed") and namespace:
events.extend(self._end(namespace))
elif method == "tasks":
if "result" in data:
events.extend(self._end_triggered_child(namespace, data.get("id")))
elif namespace:
events.extend(self._ensure_registered(namespace, tool_call_id=None))
return events
def finish(self) -> list[dict[str, Any]]:
events: list[dict[str, Any]] = []
for path in sorted(
self._active.keys(), key=lambda item: len(item), reverse=True
):
events.extend(self._end(path))
self.registry.close()
return events
def _ensure_registered(
self,
path: tuple[str, ...],
*,
tool_call_id: str | None,
) -> list[dict[str, Any]]:
if not path or path in self._active:
return []
name, parsed_tool_call_id = self._parse_namespace_segment(path[-1])
trigger_call_id = tool_call_id or parsed_tool_call_id
instance_id = ":".join(path)
self._active[path] = (name, trigger_call_id)
self.registry.register(path, name)
return [
self.emitter.subagent_start(
name,
"",
instance_id=instance_id,
tool_call_id=trigger_call_id or "",
).data
]
def _end(self, path: tuple[str, ...]) -> list[dict[str, Any]]:
active = self._active.pop(path, None)
if active is None:
return []
name, _tool_call_id = active
return [self.emitter.subagent_end(name, instance_id=":".join(path)).data]
def _end_triggered_child(
self,
namespace: tuple[str, ...],
result_id: object,
) -> list[dict[str, Any]]:
if not result_id:
return []
events: list[dict[str, Any]] = []
for path, (_name, tool_call_id) in list(self._active.items()):
if path[:-1] == namespace and tool_call_id == result_id:
events.extend(self._end(path))
return events
@staticmethod
def _parse_namespace_segment(segment: str) -> tuple[str, str | None]:
name, sep, task_id = segment.partition(":")
return name, task_id if sep else None
@dataclass(slots=True)
class LangGraphServerGateway:
"""Gateway backed by a running LangGraph server."""
thread_store: LangGraphServerThreadStore
graph_id: str = DEFAULT_GRAPH_ID
interrupt_wait_seconds: float = 5.0
def _target_graph_id(self, target: GraphTarget | None = None) -> str:
return target.graph_id if target is not None else self.graph_id
async def create_thread(
self,
target: GraphTarget | None = None,
*,
metadata: dict[str, Any] | None = None,
) -> str:
return await self.thread_store.create_thread(
graph_id=self._target_graph_id(target),
metadata=metadata,
workspace_dir=target.workspace_dir if target is not None else None,
)
async def list_threads(
self,
*,
limit: int = 20,
include_message_count: bool = False,
include_preview: bool = False,
target: GraphTarget | None = None,
) -> list[dict[str, Any]]:
return await self.thread_store.list_threads(
limit=limit,
include_message_count=include_message_count,
include_preview=include_preview,
graph_id=self._target_graph_id(target),
)
async def resolve_thread(
self,
thread_id_or_prefix: str,
target: GraphTarget | None = None,
) -> ThreadResolution:
resolved, matches = await self.thread_store.resolve_thread_id_prefix(
thread_id_or_prefix,
graph_id=self._target_graph_id(target),
)
return ThreadResolution(resolved, tuple(matches))
async def get_thread_metadata(
self,
thread_id: str,
target: GraphTarget | None = None,
) -> dict[str, Any] | None:
return await self.thread_store.get_thread_metadata(thread_id)
async def get_thread_messages(
self,
thread_id: str,
target: GraphTarget | None = None,
) -> list[BaseMessage]:
return await self.thread_store.get_thread_messages(thread_id)
async def thread_exists(
self,
thread_id: str,
target: GraphTarget | None = None,
) -> bool:
return await self.thread_store.thread_exists(thread_id)
async def delete_thread(
self,
thread_id: str,
target: GraphTarget | None = None,
) -> bool:
return await self.thread_store.delete_thread(thread_id)
async def clone_thread(
self,
source_thread_id: str,
*,
metadata: dict[str, Any] | None = None,
target: GraphTarget | None = None,
) -> str:
return await self.thread_store.clone_thread(
source_thread_id,
metadata=metadata,
)
async def _start_or_resume(
self,
stream: AsyncThreadStream,
request: RunRequest,
) -> None:
config: dict[str, Any] = {"configurable": {"thread_id": request.thread_id}}
await self.thread_store.ensure_thread_exists(
request.thread_id,
graph_id=self._target_graph_id(request.target),
metadata=request.metadata,
workspace_dir=(
request.target.workspace_dir if request.target is not None else None
),
)
request_workspace = (
request.target.workspace_dir if request.target is not None else None
)
if request.metadata or request_workspace is not None:
await self.thread_store.client.threads.update(
request.thread_id,
metadata=_build_thread_metadata(
graph_id=self._target_graph_id(request.target),
workspace_dir=request_workspace,
metadata=request.metadata,
),
)
if isinstance(request.message, Command):
if request.message.resume is not None:
await self._respond_to_interrupt(stream, request.message.resume)
return
raise RuntimeError(
"LangGraph server gateway only supports Command(resume=...) messages."
)
run_input = await build_agent_stream_input(
request.message,
media=request.media,
)
await stream.run.start(
input=run_input,
config=config,
metadata=request.metadata,
)
async def _respond_to_interrupt(
self,
stream: AsyncThreadStream,
response: object,
) -> None:
loop = asyncio.get_running_loop()
deadline = loop.time() + self.interrupt_wait_seconds
while not stream.interrupts and loop.time() < deadline:
await asyncio.sleep(0.05)
interrupt_id = None
if len(stream.interrupts) == 1:
interrupt_id = str(stream.interrupts[0].get("interrupt_id") or "")
await stream.run.respond(response, interrupt_id=interrupt_id or None)
def stream_events(self, request: RunRequest) -> AsyncIterator[GraphEvent]:
return self._stream_events(request)
async def get_state_values(
self,
target: GraphTarget,
thread_id: str,
) -> GraphStateValues:
return await self._get_state_values(thread_id)
async def update_state_values(
self,
target: GraphTarget,
thread_id: str,
values: GraphStateValues,
) -> None:
as_node = "model" if "_summarization_event" in values else None
await self.thread_store.client.threads.update_state(
thread_id,
values,
as_node=as_node,
)
async def _get_state_values(self, thread_id: str) -> GraphStateValues:
state = await self.thread_store.client.threads.get_state(thread_id)
values = state.get("values")
if not isinstance(values, dict):
return {}
return {str(key): value for key, value in values.items()}
async def _pending_interrupt_events(
self,
stream: AsyncThreadStream,
thread_id: str,
processor: _V3EventProcessor,
) -> list[GraphEvent]:
events: list[GraphEvent] = []
for interrupt in stream.interrupts:
events.extend(
await processor.process(
_input_requested_event_from_interrupt(interrupt)
)
)
if events or not stream.interrupted:
return events
try:
state = await self.thread_store.client.threads.get_state(thread_id)
except NotFoundError:
return events
for interrupt in _state_interrupts(state):
events.extend(
await processor.process(
_input_requested_event_from_interrupt(interrupt)
)
)
return events
async def _stream_events(self, request: RunRequest) -> AsyncIterator[GraphEvent]:
emitter = StreamEventEmitter()
state_values: GraphStateValues = {}
existing_summarization_event: Mapping[str, object] | None = None
process_value_messages = True
try:
state_values = await self._get_state_values(request.thread_id)
existing_summarization_event = _find_summarization_event_payload(
state_values
)
except NotFoundError:
pass
except Exception:
process_value_messages = False
subagents = _SubagentRegistry()
processor = _V3EventProcessor(
emitter,
subagents,
existing_summarization_event,
state_values.get("messages"),
process_value_messages=process_value_messages,
)
tracker = _ServerSubagentTracker(emitter, subagents)
stream = self.thread_store.client.threads.stream(
request.thread_id,
assistant_id=self._target_graph_id(request.target),
)
try:
async with stream:
await self._start_or_resume(stream, request)
emitted_interrupt = False
async for event in stream.subscribe(_RUN_SUBSCRIBE_CHANNELS):
raw_event = _as_raw_map(event)
if raw_event is None:
continue
event_map: dict[str, Any] = dict(raw_event)
for subagent_event in tracker.process(event_map):
yield subagent_event
for normalized in await processor.process(event_map):
emitted_interrupt = emitted_interrupt or _is_interrupt_event(
normalized
)
yield normalized
if not emitted_interrupt:
for event in await self._pending_interrupt_events(
stream,
request.thread_id,
processor,
):
yield event
except Exception as exc:
yield emitter.error(str(exc)).data
raise
finally:
for event in tracker.finish():
yield event
yield emitter.done(processor.full_response).data
+175
View File
@@ -0,0 +1,175 @@
"""Shared types for graph/thread gateway implementations."""
from __future__ import annotations
from collections.abc import AsyncIterator
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Protocol, TypeAlias
from langgraph.types import Command
if TYPE_CHECKING:
from langgraph.graph.state import CompiledStateGraph
GraphEvent: TypeAlias = dict[str, Any]
GraphRunInput: TypeAlias = str | Command
GraphStateValues: TypeAlias = dict[str, Any]
DEFAULT_GRAPH_ID = "EvoScientist"
@dataclass(frozen=True, slots=True)
class GraphTarget:
"""Identifies the graph/workspace a thread operation targets.
``local_graph`` is the in-process execution handle required only by the
local backend. Server backends select execution via ``graph_id``.
"""
graph_id: str = DEFAULT_GRAPH_ID
workspace_dir: str | None = None
local_graph: CompiledStateGraph | None = None
@dataclass(frozen=True, slots=True)
class RunRequest:
"""A graph turn request, independent of the UI that initiated it."""
message: GraphRunInput
thread_id: str
metadata: dict[str, Any] | None = None
media: list[str] | None = None
target: GraphTarget | None = None
@dataclass(frozen=True, slots=True)
class ThreadResolution:
"""Result of resolving an exact or prefix thread id."""
thread_id: str | None
matches: tuple[str, ...] = ()
@property
def found(self) -> bool:
return self.thread_id is not None
@property
def ambiguous(self) -> bool:
return self.thread_id is None and bool(self.matches)
class ThreadStore(Protocol):
"""Thread persistence operations used by graph gateways."""
def generate_thread_id(self) -> str:
"""Generate a new thread id."""
async def list_threads(
self,
*,
limit: int = 20,
include_message_count: bool = False,
include_preview: bool = False,
) -> list[dict[str, Any]]:
"""Return persisted threads."""
async def resolve_thread_id_prefix(
self,
thread_id_or_prefix: str,
) -> tuple[str | None, list[str]]:
"""Resolve an exact or prefix thread id."""
async def get_thread_metadata(self, thread_id: str) -> dict[str, Any] | None:
"""Return persisted metadata for a thread, if available."""
async def get_thread_messages(self, thread_id: str) -> list[Any]:
"""Return persisted messages for a thread."""
async def thread_exists(self, thread_id: str) -> bool:
"""Return whether a thread exists."""
async def delete_thread(self, thread_id: str) -> bool:
"""Delete a thread and its persisted state."""
class GraphGateway(Protocol):
"""One authority for graph runs and thread lifecycle operations."""
async def create_thread(
self,
target: GraphTarget | None = None,
*,
metadata: dict[str, Any] | None = None,
) -> str:
"""Create or reserve a new thread id."""
async def list_threads(
self,
*,
limit: int = 20,
include_message_count: bool = False,
include_preview: bool = False,
target: GraphTarget | None = None,
) -> list[dict[str, Any]]:
"""Return user-facing threads for the active backend."""
async def resolve_thread(
self,
thread_id_or_prefix: str,
target: GraphTarget | None = None,
) -> ThreadResolution:
"""Resolve a thread id or prefix."""
async def get_thread_metadata(
self,
thread_id: str,
target: GraphTarget | None = None,
) -> dict[str, Any] | None:
"""Return persisted metadata for a thread, if available."""
async def get_thread_messages(
self,
thread_id: str,
target: GraphTarget | None = None,
) -> list[Any]:
"""Return persisted messages for a thread."""
async def thread_exists(
self,
thread_id: str,
target: GraphTarget | None = None,
) -> bool:
"""Return whether a thread exists in the active backend."""
async def delete_thread(
self,
thread_id: str,
target: GraphTarget | None = None,
) -> bool:
"""Delete a thread and its persisted state."""
async def clone_thread(
self,
source_thread_id: str,
*,
metadata: dict[str, Any] | None = None,
target: GraphTarget | None = None,
) -> str:
"""Clone a thread and return the cloned thread id."""
def stream_events(self, request: RunRequest) -> AsyncIterator[GraphEvent]:
"""Stream normalized graph events for the request target."""
async def get_state_values(
self,
target: GraphTarget,
thread_id: str,
) -> GraphStateValues:
"""Return the graph state values for a thread."""
async def update_state_values(
self,
target: GraphTarget,
thread_id: str,
values: GraphStateValues,
) -> None:
"""Update graph state values for a thread."""
+130 -37
View File
@@ -18,7 +18,7 @@ from dataclasses import dataclass
from datetime import UTC, datetime
from enum import StrEnum
from pathlib import Path
from typing import TYPE_CHECKING, Any, NotRequired, TypedDict, TypeVar, cast
from typing import TYPE_CHECKING, Any, NotRequired, Protocol, TypedDict, TypeVar, cast
from langchain.agents.middleware.types import AgentMiddleware, AgentState
from langchain_core.messages import AIMessage, BaseMessage, ToolMessage, filter_messages
@@ -44,7 +44,7 @@ from ..memory.worker_activity import (
)
if TYPE_CHECKING:
from langgraph_sdk.schema import Config, Input
from langgraph_sdk.schema import Config, Input, Run, Thread
logger = logging.getLogger(__name__)
@@ -139,6 +139,7 @@ class MemoryWorkerLaunchArgs(TypedDict):
role: MemoryLifecycleRole
memory_dir: str | Path
workspace_dir: str | Path
project_id: str
source_agent: str
session_id: str
@@ -154,6 +155,62 @@ class MemoryWorkerRunPayload(TypedDict):
config: Config
class _SyncMemoryWorkerThreads(Protocol):
def create(
self,
*,
graph_id: str,
metadata: dict[str, str],
) -> Thread: ...
class _SyncMemoryWorkerRuns(Protocol):
def create(
self,
thread_id: str,
assistant_id: str,
*,
input: Input,
metadata: dict[str, str],
config: Config,
) -> Run: ...
def get(self, thread_id: str, run_id: str) -> Run: ...
class _SyncMemoryWorkerClient(Protocol):
threads: _SyncMemoryWorkerThreads
runs: _SyncMemoryWorkerRuns
class _AsyncMemoryWorkerThreads(Protocol):
async def create(
self,
*,
graph_id: str,
metadata: dict[str, str],
) -> Thread: ...
class _AsyncMemoryWorkerRuns(Protocol):
async def create(
self,
thread_id: str,
assistant_id: str,
*,
input: Input,
metadata: dict[str, str],
config: Config,
) -> Run: ...
async def get(self, thread_id: str, run_id: str) -> Run: ...
class _AsyncMemoryWorkerClient(Protocol):
threads: _AsyncMemoryWorkerThreads
runs: _AsyncMemoryWorkerRuns
@dataclass(frozen=True)
class _SummaryWriteArgs:
"""Concrete metadata needed to write a subagent execution summary."""
@@ -617,20 +674,6 @@ def _safe_segment(value: str) -> str:
return safe.strip("-") or "unknown"
def _worker_thread_id(
*,
role: MemoryLifecycleRole,
session_id: str,
source_agent: str,
trajectory: list[CompactMessage],
) -> str:
"""Return a deterministic thread id for a background worker run."""
key = "\n".join(
[role.value, session_id, source_agent, _trajectory_digest(trajectory)]
)
return f"evomemory-{role.value}:{_short_hash(key)}"
def _agent_result_model(result: Mapping[str, object], model_type: type[T]) -> T | None:
"""Extract a DeepAgents/LangChain structured response from agent state."""
value = result.get("structured_response")
@@ -952,30 +995,49 @@ def _runs_create_kwargs(kwargs: MemoryWorkerRunPayload) -> MemoryWorkerRunPayloa
return cast("MemoryWorkerRunPayload", _merge_runs_config_kwargs(dict(kwargs)))
def _worker_workspace_dir(workspace_dir: str | Path) -> str:
return str(Path(workspace_dir).expanduser().resolve())
def _memory_worker_metadata(
*,
role: MemoryLifecycleRole,
workspace_dir: str | Path,
project_id: str,
source_agent: str,
session_id: str,
trajectory_digest: str,
) -> dict[str, str]:
return {
"run_kind": f"evomemory_{role.value}_worker",
"source_session_id": session_id,
"source_agent": source_agent,
"project_id": project_id,
"trajectory_digest": trajectory_digest,
"workspace_dir": _worker_workspace_dir(workspace_dir),
}
def _memory_worker_run_kwargs(
*,
role: MemoryLifecycleRole,
thread_id: str,
workspace_dir: str | Path,
project_id: str,
source_agent: str,
session_id: str,
trajectory: list[CompactMessage],
) -> MemoryWorkerRunPayload:
"""Build the LangGraph SDK run payload for a memory worker."""
worker_thread_id = _worker_thread_id(
role=role,
session_id=session_id,
source_agent=source_agent,
trajectory=trajectory,
)
trajectory_digest = _trajectory_digest(trajectory)
metadata = {
"agent_name": "EvoScientist",
"run_kind": f"evomemory_{role.value}_worker",
"source_session_id": session_id,
"source_agent": source_agent,
"project_id": project_id,
"trajectory_digest": trajectory_digest,
}
metadata = _memory_worker_metadata(
role=role,
workspace_dir=workspace_dir,
project_id=project_id,
source_agent=source_agent,
session_id=session_id,
trajectory_digest=trajectory_digest,
)
payload: MemoryWorkerRunPayload = {
"assistant_id": role.graph_id,
"input": {
@@ -993,7 +1055,7 @@ def _memory_worker_run_kwargs(
"metadata": metadata,
"config": {
"configurable": {
"thread_id": worker_thread_id,
"thread_id": thread_id,
"evomemory_source_session_id": session_id,
"evomemory_source_agent": source_agent,
"evomemory_project_id": project_id,
@@ -1126,7 +1188,7 @@ def _watch_memory_worker_run_sync(
def _spawn_memory_worker_status_task(
client: Any,
client: _AsyncMemoryWorkerClient,
*,
thread_id: str,
run_id: str,
@@ -1140,7 +1202,7 @@ def _spawn_memory_worker_status_task(
async def _watch_memory_worker_run_async(
client: Any,
client: _AsyncMemoryWorkerClient,
*,
thread_id: str,
run_id: str,
@@ -1189,6 +1251,7 @@ def _launch_memory_worker(
*,
role: MemoryLifecycleRole,
memory_dir: str | Path,
workspace_dir: str | Path,
project_id: str,
source_agent: str,
session_id: str,
@@ -1204,12 +1267,24 @@ def _launch_memory_worker(
logger.info("Skipping EvoMemory worker launch; LangGraph dev is unavailable")
return
client = get_sync_client(url=url, headers={"x-auth-scheme": "langsmith"})
thread = client.threads.create(graph_id=role.graph_id)
client: _SyncMemoryWorkerClient = get_sync_client(
url=url, headers={"x-auth-scheme": "langsmith"}
)
metadata = _memory_worker_metadata(
role=role,
workspace_dir=workspace_dir,
project_id=project_id,
source_agent=source_agent,
session_id=session_id,
trajectory_digest=_trajectory_digest(trajectory),
)
thread = client.threads.create(graph_id=role.graph_id, metadata=metadata)
worker_thread_id = str(thread["thread_id"])
before_outputs = snapshot_memory_outputs(memory_dir)
payload = _memory_worker_run_kwargs(
role=role,
thread_id=worker_thread_id,
workspace_dir=workspace_dir,
project_id=project_id,
source_agent=source_agent,
session_id=session_id,
@@ -1244,6 +1319,7 @@ async def _alaunch_memory_worker(
*,
role: MemoryLifecycleRole,
memory_dir: str | Path,
workspace_dir: str | Path,
project_id: str,
source_agent: str,
session_id: str,
@@ -1259,12 +1335,24 @@ async def _alaunch_memory_worker(
logger.info("Skipping EvoMemory worker launch; LangGraph dev is unavailable")
return
client = get_client(url=url, headers={"x-auth-scheme": "langsmith"})
thread = await client.threads.create(graph_id=role.graph_id)
client: _AsyncMemoryWorkerClient = get_client(
url=url, headers={"x-auth-scheme": "langsmith"}
)
metadata = _memory_worker_metadata(
role=role,
workspace_dir=workspace_dir,
project_id=project_id,
source_agent=source_agent,
session_id=session_id,
trajectory_digest=_trajectory_digest(trajectory),
)
thread = await client.threads.create(graph_id=role.graph_id, metadata=metadata)
worker_thread_id = str(thread["thread_id"])
before_outputs = await asyncio.to_thread(snapshot_memory_outputs, memory_dir)
payload = _memory_worker_run_kwargs(
role=role,
thread_id=worker_thread_id,
workspace_dir=workspace_dir,
project_id=project_id,
source_agent=source_agent,
session_id=session_id,
@@ -1310,6 +1398,9 @@ class EvoMemoryLifecycleMiddleware(AgentMiddleware):
source_agent: str,
) -> None:
self._memory_dir = Path(memory_dir).expanduser()
self._workspace_dir = Path(
_paths.WORKSPACE_ROOT if workspace_dir is None else workspace_dir
).expanduser()
self._project_id = project_id
self._role = role
self._source_agent = source_agent
@@ -1329,6 +1420,7 @@ class EvoMemoryLifecycleMiddleware(AgentMiddleware):
return {
"role": MemoryLifecycleRole.TURN,
"memory_dir": self._memory_dir,
"workspace_dir": self._workspace_dir,
"project_id": self._project_id,
"source_agent": self._source_agent,
"session_id": session_id,
@@ -1341,6 +1433,7 @@ class EvoMemoryLifecycleMiddleware(AgentMiddleware):
return {
"role": MemoryLifecycleRole.SUBAGENT,
"memory_dir": self._memory_dir,
"workspace_dir": self._workspace_dir,
"project_id": self._project_id,
"source_agent": self._source_agent,
"session_id": session_id,
+129 -90
View File
@@ -34,6 +34,7 @@ import math
import uuid
from collections.abc import AsyncIterator, Awaitable, Callable
from contextlib import asynccontextmanager
from dataclasses import dataclass
from datetime import UTC, datetime
from pathlib import Path
from typing import Any, cast
@@ -45,6 +46,7 @@ from langchain_core.messages import (
RemoveMessage,
convert_to_messages,
)
from langchain_core.runnables import RunnableConfig
from langgraph.checkpoint.serde.jsonplus import JsonPlusSerializer
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
from langgraph.graph.message import REMOVE_ALL_MESSAGES
@@ -68,6 +70,12 @@ if not hasattr(aiosqlite.Connection, "is_alive"):
# ---------------------------------------------------------------------------
AGENT_NAME = "EvoScientist"
MAIN_THREAD_FILTER_SQL = (
"json_extract(metadata, '$.agent_name') = ? "
"AND (json_extract(metadata, '$.graph_id') IS NULL "
" OR json_extract(metadata, '$.graph_id') = ?)"
)
MAIN_THREAD_FILTER_PARAMS = (AGENT_NAME, AGENT_NAME)
# ---------------------------------------------------------------------------
@@ -637,8 +645,8 @@ async def _load_checkpoint_messages(
Returns a list of LangChain message objects, or an empty list on failure.
"""
# Pre-resolve the latest EvoScientist checkpoint_id with an
# ``agent_name`` filter, then pin it into the config so
# Pre-resolve the latest main EvoScientist checkpoint_id, then pin it
# into the config so
# ``aget_tuple`` fetches THAT specific row. Without the pin,
# ``aget_tuple`` returns the latest by ``checkpoint_id`` alone — in
# a multi-agent DB where a third-party tool shares the same
@@ -650,14 +658,16 @@ async def _load_checkpoint_messages(
head_query = (
"SELECT checkpoint_id FROM checkpoints "
"WHERE thread_id = ? AND checkpoint_ns = '' "
" AND json_extract(metadata, '$.agent_name') = ? "
f" AND {MAIN_THREAD_FILTER_SQL} "
"ORDER BY checkpoint_id DESC LIMIT 1"
)
async with saver.conn.execute(head_query, (thread_id, AGENT_NAME)) as cur:
async with saver.conn.execute(
head_query, (thread_id, *MAIN_THREAD_FILTER_PARAMS)
) as cur:
head_row = await cur.fetchone()
if head_row is None:
return []
config = {
config: RunnableConfig = {
"configurable": {
"thread_id": thread_id,
"checkpoint_ns": "",
@@ -880,20 +890,20 @@ async def list_threads(
if not await _table_exists(conn, "checkpoints"):
return []
query = """
query = f"""
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') = ?
WHERE {MAIN_THREAD_FILTER_SQL}
GROUP BY thread_id
ORDER BY updated_at DESC
"""
params: tuple = (AGENT_NAME,)
params: tuple = MAIN_THREAD_FILTER_PARAMS
if limit > 0:
query += " LIMIT ?\n"
params = (AGENT_NAME, limit)
params = (*MAIN_THREAD_FILTER_PARAMS, limit)
async with conn.execute(query, params) as cur:
rows = await cur.fetchall()
@@ -927,13 +937,13 @@ async def get_most_recent() -> str | None:
async with aiosqlite.connect(db_path, timeout=30.0) as conn:
if not await _table_exists(conn, "checkpoints"):
return None
query = """
query = f"""
SELECT thread_id FROM checkpoints
WHERE json_extract(metadata, '$.agent_name') = ?
WHERE {MAIN_THREAD_FILTER_SQL}
ORDER BY checkpoint_id DESC
LIMIT 1
"""
async with conn.execute(query, (AGENT_NAME,)) as cur:
async with conn.execute(query, MAIN_THREAD_FILTER_PARAMS) as cur:
row = await cur.fetchone()
return row[0] if row else None
@@ -944,12 +954,12 @@ async def thread_exists(thread_id: str) -> bool:
async with aiosqlite.connect(db_path, timeout=30.0) as conn:
if not await _table_exists(conn, "checkpoints"):
return False
query = """
query = f"""
SELECT 1 FROM checkpoints
WHERE thread_id = ? AND json_extract(metadata, '$.agent_name') = ?
WHERE thread_id = ? AND {MAIN_THREAD_FILTER_SQL}
LIMIT 1
"""
async with conn.execute(query, (thread_id, AGENT_NAME)) as cur:
async with conn.execute(query, (thread_id, *MAIN_THREAD_FILTER_PARAMS)) as cur:
return (await cur.fetchone()) is not None
@@ -964,15 +974,17 @@ async def find_similar_threads(thread_id: str, limit: int = 5) -> list[str]:
escaped = (
thread_id.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
)
query = r"""
query = f"""
SELECT DISTINCT thread_id
FROM checkpoints
WHERE thread_id LIKE ? ESCAPE '\'
AND json_extract(metadata, '$.agent_name') = ?
WHERE thread_id LIKE ? ESCAPE '\\'
AND {MAIN_THREAD_FILTER_SQL}
ORDER BY thread_id
LIMIT ?
"""
async with conn.execute(query, (escaped + "%", AGENT_NAME, limit)) as cur:
async with conn.execute(
query, (escaped + "%", *MAIN_THREAD_FILTER_PARAMS, limit)
) as cur:
rows = await cur.fetchall()
return [r[0] for r in rows]
@@ -1002,18 +1014,18 @@ async def delete_thread(thread_id: str) -> bool:
# Delete writes FIRST — the subquery needs checkpoints to still exist
if await _table_exists(conn, "writes"):
await conn.execute(
"""DELETE FROM writes
f"""DELETE FROM writes
WHERE thread_id = ?
AND checkpoint_id IN (
SELECT checkpoint_id FROM checkpoints
WHERE thread_id = ?
AND json_extract(metadata, '$.agent_name') = ?
AND {MAIN_THREAD_FILTER_SQL}
)""",
(thread_id, thread_id, AGENT_NAME),
(thread_id, thread_id, *MAIN_THREAD_FILTER_PARAMS),
)
cur = await conn.execute(
"DELETE FROM checkpoints WHERE thread_id = ? AND json_extract(metadata, '$.agent_name') = ?",
(thread_id, AGENT_NAME),
f"DELETE FROM checkpoints WHERE thread_id = ? AND {MAIN_THREAD_FILTER_SQL}",
(thread_id, *MAIN_THREAD_FILTER_PARAMS),
)
deleted = cur.rowcount > 0
await conn.commit()
@@ -1029,17 +1041,17 @@ async def get_thread_metadata(thread_id: str) -> dict | None:
async with aiosqlite.connect(db_path, timeout=30.0) as conn:
if not await _table_exists(conn, "checkpoints"):
return None
query = """
query = f"""
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 = ?
AND json_extract(metadata, '$.agent_name') = ?
AND {MAIN_THREAD_FILTER_SQL}
ORDER BY checkpoint_id DESC
LIMIT 1
"""
async with conn.execute(query, (thread_id, AGENT_NAME)) as cur:
async with conn.execute(query, (thread_id, *MAIN_THREAD_FILTER_PARAMS)) as cur:
row = await cur.fetchone()
if not row:
return None
@@ -1066,12 +1078,12 @@ async def get_thread_messages(thread_id: str) -> list:
if not await _table_exists(conn, "checkpoints"):
return []
# Verify this thread belongs to EvoScientist before loading messages
check = """
check = f"""
SELECT 1 FROM checkpoints
WHERE thread_id = ? AND json_extract(metadata, '$.agent_name') = ?
WHERE thread_id = ? AND {MAIN_THREAD_FILTER_SQL}
LIMIT 1
"""
async with conn.execute(check, (thread_id, AGENT_NAME)) as cur:
async with conn.execute(check, (thread_id, *MAIN_THREAD_FILTER_PARAMS)) as cur:
if not await cur.fetchone():
return []
serde = JsonPlusSerializer()
@@ -1237,7 +1249,7 @@ async def _run_migration_sweep(
"WHERE json_extract(metadata, '$.agent_name') = ?",
(AGENT_NAME,),
) as cur:
pairs = await cur.fetchall()
pairs = list(await cur.fetchall())
# Reuse the DeltaChannel-aware prune logic from PruningCheckpointer
# instead of running naive keep_latest SQL: legacy DBs almost always
@@ -1437,13 +1449,15 @@ class _ApiPruningCheckpointer(PruningCheckpointer):
"""``PruningCheckpointer`` that stamps CLI-compatible ownership metadata.
langgraph-api run metadata carries ``graph_id``/``assistant_id`` but not
the ``agent_name`` / ``workspace_dir`` / ``updated_at`` keys that the CLI
session surface (``list_threads``, ``/resume``, ``/delete``,
``_prune_after_put``) filters and sorts on. Stamping them at write time
— for main-graph runs only — makes WebUI threads first-class CLI
sessions in the same workspace, and brings them under the existing
pruning/retention machinery. Worker and async-subagent graphs are left
unstamped on purpose: they must not surface in CLI listings.
always the ``workspace_dir`` / ``updated_at`` keys needed to safely
rebuild the in-memory thread registry after server restart. Stamping graph
rows with the current workspace keeps main and async-subagent threads
restorable without exposing other workspaces. Memory-worker rows still get
workspace metadata, but remain disposable until worker cloning lands.
Only the main graph receives ``agent_name``. The local CLI session
surface still uses that ownership key, so worker/subagent graph rows must
remain outside ordinary ``/threads``, ``/resume``, and ``/delete``.
"""
async def aput(
@@ -1453,9 +1467,8 @@ class _ApiPruningCheckpointer(PruningCheckpointer):
metadata: Any,
new_versions: Any,
) -> Any:
if isinstance(metadata, dict) and metadata.get("graph_id") == AGENT_NAME:
if isinstance(metadata, dict) and isinstance(metadata.get("graph_id"), str):
metadata = dict(metadata)
metadata.setdefault("agent_name", AGENT_NAME)
# _api_workspace_dir() calls Path.resolve()/Path.cwd() -> os.getcwd(),
# a blocking syscall flagged by the dev runtime's blockbuster guard.
# Run it in a thread, and only when actually needed — ``setdefault``
@@ -1464,6 +1477,8 @@ class _ApiPruningCheckpointer(PruningCheckpointer):
if "workspace_dir" not in metadata:
metadata["workspace_dir"] = await _api_workspace_dir_async()
metadata["updated_at"] = datetime.now(UTC).isoformat()
if metadata.get("graph_id") == AGENT_NAME:
metadata.setdefault("agent_name", AGENT_NAME)
return await super().aput(config, checkpoint, metadata, new_versions)
@@ -1508,6 +1523,15 @@ async def _purge_internal_worker_threads() -> None:
)
@dataclass(frozen=True, slots=True)
class _RestoredThreadInfo:
updated_at: str | None
assistant_id: str | None
graph_id: str
workspace_dir: str
model: str | None
async def _restore_webui_threads_to_global_store() -> None:
"""Re-populate ``GlobalStore["threads"]`` from SQLite on server startup.
@@ -1519,20 +1543,21 @@ async def _restore_webui_threads_to_global_store() -> None:
are normalized in place, and missing threads are appended as stub
dicts that satisfy ``POST /threads/search``.
Restore scope — only threads that are BOTH main-graph
(``metadata.graph_id == AGENT_NAME``) and owned by this server's
workspace (``metadata.workspace_dir`` matches): sessions.db is
machine-global, and an unscoped restore would expose every workspace's
history (and internal worker threads) on the unauthenticated API —
worst case ``--tunnel``. CLI/TUI threads (8-char hex IDs, managed by
``list_threads()``) and pre-stamping rows without ``workspace_dir``
are excluded.
Restore scope — UUID-format graph threads owned by this server's
workspace (``metadata.workspace_dir`` matches). This includes the main
graph and async-subagent graphs. Memory-worker graphs are excluded for now:
they are still treated as disposable residue until worker cloning lands.
The workspace filter is required because sessions.db is machine-global,
and an unscoped restore would expose every workspace's history on the
unauthenticated API — worst case ``--tunnel``. CLI/TUI threads (8-char
hex IDs, managed by ``list_threads()``) and pre-stamping rows without
``workspace_dir`` are excluded.
Best-effort: any exception is logged and swallowed so a broken restore
never prevents the ``langgraph dev`` server from starting.
"""
try:
from langgraph_runtime_inmem.database import ( # type: ignore[import-untyped]
from langgraph_runtime_inmem.database import (
GLOBAL_STORE,
)
except ImportError:
@@ -1547,24 +1572,15 @@ async def _restore_webui_threads_to_global_store() -> None:
try:
rows: list[Any] = []
# All UUID threads that have ANY checkpoint rows — the existence
# check for ghost removal (deliberately unscoped: a thread whose
# checkpoints exist but fall outside the restore scope is not a
# ghost, its state still loads when opened).
uuid_threads_in_db: set[uuid.UUID] = set()
# Restore scope: ONLY main-graph threads belonging to THIS server's
# workspace. sessions.db is machine-global, so an unscoped restore
# would resurrect every workspace's history (and internal
# worker/subagent threads) into this server's thread registry — and
# expose it over the unauthenticated API / --tunnel. Main-graph =
# metadata.graph_id == AGENT_NAME (langgraph-api rows, stamped by
# _ApiPruningCheckpointer) OR no graph_id but agent_name ==
# AGENT_NAME (CLI rows via build_metadata). Worker residue carries
# graph_id='evomemory-*' and is excluded by the first clause even
# though it also stamps agent_name. Rows predating stamping have no
# workspace_dir and are deliberately excluded.
# Restore scope: graph threads belonging to THIS server's workspace.
# sessions.db is machine-global, so an unscoped restore would
# resurrect every workspace's history into this server's thread
# registry — and expose it over the unauthenticated API / --tunnel.
# Legacy WebUI/CLI interop rows without graph_id are restored as the
# main graph only when they carry agent_name == AGENT_NAME. Rows
# predating workspace stamping remain deliberately excluded.
current_workspace = await _api_workspace_dir_async()
sqlite_data: dict[uuid.UUID, tuple[str | None, str | None, str]] = {}
sqlite_data: dict[uuid.UUID, _RestoredThreadInfo] = {}
titles: dict[uuid.UUID, str] = {}
db_path = str(get_db_path())
async with aiosqlite.connect(db_path, timeout=30.0) as conn:
@@ -1584,14 +1600,19 @@ async def _restore_webui_threads_to_global_store() -> None:
MAX(json_extract(metadata, '$.assistant_id')) as assistant_id,
MAX(json_extract(metadata, '$.graph_id')) as graph_id,
MAX(json_extract(metadata, '$.workspace_dir')) as workspace_dir,
MAX(json_extract(metadata, '$.model')) as model,
MAX(json_extract(metadata, '$.agent_name')) as agent_name
FROM checkpoints
WHERE thread_id LIKE '________-____-____-____-____________'
AND (
json_extract(metadata, '$.graph_id') IS NULL
OR json_extract(metadata, '$.graph_id') NOT LIKE 'evomemory-%'
)
GROUP BY thread_id
ORDER BY updated_at DESC
"""
async with conn.execute(query) as cur:
rows = await cur.fetchall()
rows = list(await cur.fetchall())
for row in rows:
(
@@ -1600,20 +1621,26 @@ async def _restore_webui_threads_to_global_store() -> None:
assistant_id,
graph_id,
workspace_dir,
model,
agent_name,
) = row
thread_uuid = _to_uuid_safe(thread_id_str)
if thread_uuid is None:
continue
uuid_threads_in_db.add(thread_uuid)
is_main_graph = graph_id == AGENT_NAME or (
graph_id is None and agent_name == AGENT_NAME
)
if not is_main_graph:
restored_graph_id = graph_id
if restored_graph_id is None and agent_name == AGENT_NAME:
restored_graph_id = AGENT_NAME
if restored_graph_id is None:
continue
if not workspace_dir or workspace_dir != current_workspace:
continue
sqlite_data[thread_uuid] = (updated_at, assistant_id, AGENT_NAME)
sqlite_data[thread_uuid] = _RestoredThreadInfo(
updated_at=updated_at,
assistant_id=assistant_id,
graph_id=restored_graph_id,
workspace_dir=workspace_dir,
model=model,
)
# Derive a sidebar title from each scoped thread's first human
# message (stubs carry values=None, so the WebUI would otherwise
@@ -1638,16 +1665,16 @@ async def _restore_webui_threads_to_global_store() -> None:
pass
return datetime.now(UTC)
# Drop ghost entries: a .pckl-loaded UUID entry with no checkpoint
# rows opens as an empty session (the #277 symptom). Slice
# assignment mutates the live registry list.
# Drop stale registry entries: UUID entries outside the scoped restore
# set either point at missing state or another workspace's state.
# Slice assignment mutates the live registry list.
store_threads: list[dict[str, Any]] = GLOBAL_STORE.get("threads", [])
before = len(store_threads)
store_threads[:] = [
entry
for entry in store_threads
if (tid := _to_uuid_safe(entry.get("thread_id"))) is None
or tid in uuid_threads_in_db
or tid in sqlite_data
]
removed = before - len(store_threads)
@@ -1665,15 +1692,21 @@ async def _restore_webui_threads_to_global_store() -> None:
entry["thread_id"] = tid_uuid
changed = True
if tid_uuid in sqlite_data:
_updated_at, asst_id_str, gid = sqlite_data[tid_uuid]
info = sqlite_data[tid_uuid]
meta: dict[str, Any] = entry.setdefault("metadata", {})
if asst_id_str and "assistant_id" not in meta:
if info.assistant_id and "assistant_id" not in meta:
# str, not uuid.UUID: the runtime stores str and search
# filters compare with raw == against JSON strings.
meta["assistant_id"] = str(asst_id_str)
meta["assistant_id"] = str(info.assistant_id)
changed = True
if gid and "graph_id" not in meta:
meta["graph_id"] = gid
if info.graph_id and "graph_id" not in meta:
meta["graph_id"] = info.graph_id
changed = True
if meta.get("workspace_dir") != info.workspace_dir:
meta["workspace_dir"] = info.workspace_dir
changed = True
if info.model and meta.get("model") != info.model:
meta["model"] = info.model
changed = True
if "title" not in meta and tid_uuid in titles:
meta["title"] = titles[tid_uuid]
@@ -1692,16 +1725,21 @@ async def _restore_webui_threads_to_global_store() -> None:
# Append threads present in SQLite but absent from the registry.
restored = 0
for thread_uuid, (updated_at, assistant_id, graph_id) in sqlite_data.items():
for thread_uuid, info in sqlite_data.items():
if thread_uuid in existing_uuids:
continue
stub_metadata: dict[str, Any] = {"graph_id": graph_id}
if assistant_id:
stub_metadata: dict[str, Any] = {
"graph_id": info.graph_id,
"workspace_dir": info.workspace_dir,
}
if info.assistant_id:
# str, not uuid.UUID — same convention as above.
stub_metadata["assistant_id"] = str(assistant_id)
stub_metadata["assistant_id"] = str(info.assistant_id)
if info.model:
stub_metadata["model"] = info.model
if thread_uuid in titles:
stub_metadata["title"] = titles[thread_uuid]
ts = _parse_dt(updated_at)
ts = _parse_dt(info.updated_at)
stub: dict[str, Any] = {
"thread_id": thread_uuid,
"created_at": ts,
@@ -1746,9 +1784,10 @@ async def create_checkpointer_for_langgraph_api() -> AsyncIterator[PruningCheckp
bad row only loses that row. The langgraph-api adapter detects async
context managers and enters them automatically.
The yielded ``_ApiPruningCheckpointer`` stamps main-graph rows with
``agent_name`` / ``workspace_dir`` / ``updated_at`` so WebUI threads
surface in the CLI session commands and participate in
The yielded ``_ApiPruningCheckpointer`` stamps graph rows with
``workspace_dir`` / ``updated_at`` so they can be restored into the
LangGraph server registry. Main-graph rows also get ``agent_name`` so
WebUI threads surface in the CLI session commands and participate in
``_prune_after_put`` retention.
Capability note: ``adelete_thread`` is real, but ``aprune`` /
+42 -9
View File
@@ -12,7 +12,7 @@ import os
import re
import threading
from collections.abc import Callable
from typing import Any
from typing import TYPE_CHECKING, Any
from rich.console import Group # type: ignore[import-untyped]
from rich.live import Live # type: ignore[import-untyped]
@@ -21,10 +21,10 @@ from rich.panel import Panel # type: ignore[import-untyped]
from rich.spinner import Spinner # type: ignore[import-untyped]
from rich.text import Text # type: ignore[import-untyped]
from ..gateway import GraphGateway, GraphRunInput, GraphTarget, RunRequest
from ..paths import resolve_virtual_path
from .console import console
from .diff_format import build_edit_diff
from .events import stream_agent_events
from .formatter import ToolResultFormatter
from .state import (
StreamState,
@@ -40,6 +40,9 @@ from .utils import (
is_success,
)
if TYPE_CHECKING:
from langgraph.graph.state import CompiledStateGraph
# ---------------------------------------------------------------------------
# Shared globals
# ---------------------------------------------------------------------------
@@ -47,6 +50,19 @@ from .utils import (
# Media file extensions that should trigger on_file_write callback
_MEDIA_EXTENSIONS = {".png", ".jpg", ".jpeg", ".gif", ".webp", ".bmp", ".svg", ".pdf"}
def _graph_target_for_local_agent(
agent: "CompiledStateGraph",
metadata: dict[str, object] | None = None,
) -> GraphTarget:
workspace = None
if metadata is not None:
raw_workspace = metadata.get("workspace_dir")
if isinstance(raw_workspace, str) and raw_workspace:
workspace = raw_workspace
return GraphTarget(local_graph=agent, workspace_dir=workspace)
# LLM output sometimes omits the CommonMark-required space after `#` (e.g.
# "###文件系统"), which makes Rich render the line as raw text. The lookahead
# `(?=[^ \t#\r\n])` requires a real non-excluded next char, so the helper is
@@ -1264,8 +1280,8 @@ def _resolve_ask_user_prompt(ask_user_data: dict) -> dict:
def _run_streaming(
agent: Any,
message: Any,
agent: "CompiledStateGraph",
message: GraphRunInput,
thread_id: str,
show_thinking: bool,
interactive: bool,
@@ -1274,11 +1290,12 @@ def _run_streaming(
on_file_write: Callable[[str], None] | None = None,
on_stream_event: Callable[[str, Any], Any] | None = None,
status_footer_builder: Callable[[], Any] | None = None,
metadata: dict | None = None,
metadata: dict[str, object] | None = None,
hitl_prompt_fn: Callable[[list], list[dict] | None] | None = None,
ask_user_prompt_fn: Callable[[dict], dict] | None = None,
cancel_scope: str | None = None,
*,
gateway: GraphGateway,
_state: StreamState | None = None,
_hitl_depth: int = 0,
_media_sent: set[str] | None = None,
@@ -1305,6 +1322,7 @@ def _run_streaming(
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.
gateway: Graph/thread gateway supplied by the active runtime.
Returns:
The final response text.
@@ -1328,8 +1346,13 @@ def _run_streaming(
async def _consume() -> None:
nonlocal _sent_thinking_text, _todo_sent
async for event in stream_agent_events(
agent, message, thread_id, metadata=metadata
async for event in gateway.stream_events(
RunRequest(
message=message,
thread_id=thread_id,
metadata=metadata,
target=_graph_target_for_local_agent(agent, metadata),
)
):
if is_stream_cancel_requested(cancel_scope):
_stopped_response()
@@ -1568,6 +1591,7 @@ def _run_streaming(
hitl_prompt_fn=hitl_prompt_fn,
ask_user_prompt_fn=ask_user_prompt_fn,
cancel_scope=cancel_scope,
gateway=gateway,
_state=state,
_hitl_depth=_hitl_depth + 1,
_media_sent=_media_sent,
@@ -1606,6 +1630,7 @@ def _run_streaming(
hitl_prompt_fn=hitl_prompt_fn,
ask_user_prompt_fn=ask_user_prompt_fn,
cancel_scope=cancel_scope,
gateway=gateway,
_state=state,
_hitl_depth=_hitl_depth + 1,
_media_sent=_media_sent,
@@ -1631,10 +1656,12 @@ def _run_streaming(
async def _astream_to_console(
agent: Any,
agent: "CompiledStateGraph",
message: str,
thread_id: str,
show_thinking: bool = True,
*,
gateway: GraphGateway,
) -> str:
"""Stream agent events to console using static prints (thread-safe, no Live).
@@ -1654,7 +1681,13 @@ async def _astream_to_console(
"""
state = StreamState()
async for event in stream_agent_events(agent, message, thread_id):
async for event in gateway.stream_events(
RunRequest(
message=message,
thread_id=thread_id,
target=_graph_target_for_local_agent(agent),
)
):
etype = state.handle_event(event)
# Only show subagent starts as real-time progress.
+198 -98
View File
@@ -8,9 +8,9 @@ import base64
import inspect
import mimetypes
import os
from collections.abc import AsyncGenerator, AsyncIterator
from collections.abc import AsyncGenerator, AsyncIterator, Mapping
from dataclasses import dataclass
from typing import Any
from typing import Any, TypeAlias
from langchain_core.messages import AIMessage, AIMessageChunk, BaseMessage, ToolMessage
from langgraph.graph import END
@@ -38,6 +38,17 @@ from .v3_payloads import (
_usage_counts,
)
UserMessageContent: TypeAlias = str | list[dict[str, object]]
GraphRunInput: TypeAlias = str | Command
LangGraphStreamInput: TypeAlias = dict[str, list[dict[str, object]]] | Command
_ValueMessageKey: TypeAlias = tuple[str, ...]
@dataclass(frozen=True, slots=True)
class _AssistantValueMessage:
key: _ValueMessageKey
content: object
def _is_interrupt_error_message(message: object) -> bool:
if not isinstance(message, str):
@@ -181,11 +192,17 @@ class _V3EventProcessor:
self,
emitter: StreamEventEmitter,
subagents: _SubagentRegistry,
baseline_summarization_signature: tuple[object, ...] | None,
existing_summarization_event: Mapping[str, object] | None,
existing_messages: object = None,
process_value_messages: bool = False,
) -> None:
self.emitter = emitter
self.subagents = subagents
self.baseline_summarization_signature = baseline_summarization_signature
self._suppressed_summarization_signature = _summarization_event_signature(
existing_summarization_event
)
self._seen_value_message_keys = self._message_keys(existing_messages)
self._process_value_message_snapshots = process_value_messages
self.full_response = ""
self._summarization_in_progress = False
self._tool_inputs: dict[
@@ -211,18 +228,82 @@ class _V3EventProcessor:
return []
if method == "messages":
return self._process_message_event(_event_data(event), subagent, namespace)
events = self._process_message_event(
_event_data(event), subagent, namespace
)
if not namespace and any(item.get("type") == "text" for item in events):
self._process_value_message_snapshots = False
return events
if method == "tools":
return self._process_tool_event(namespace, _event_data(event), subagent)
if method == "updates":
return self._process_update_event(_event_data(event))
if method == "values":
events: list[dict[str, Any]] = []
params = event.get("params") or {}
interrupts = params.get("interrupts") or ()
if interrupts:
return self._process_update_event({"__interrupt__": interrupts})
events.extend(self._process_update_event({"__interrupt__": interrupts}))
if self._process_value_message_snapshots and not namespace:
events.extend(self._process_value_messages(_event_data(event)))
return events
if method == "input.requested":
return self._process_input_requested(event.get("params"))
return []
@classmethod
def _message_keys(cls, messages: object) -> set[_ValueMessageKey]:
if not isinstance(messages, list):
return set()
keys: set[_ValueMessageKey] = set()
for message in messages:
if parsed := cls._assistant_value_message(message):
keys.add(parsed.key)
return keys
@staticmethod
def _assistant_value_message(message: object) -> _AssistantValueMessage | None:
message_map = _as_raw_map(message)
if message_map is not None:
raw_id = message_map.get("id")
raw_role = message_map.get("type") or message_map.get("role")
content = message_map.get("content")
elif isinstance(message, BaseMessage):
raw_id = message.id
raw_role = message.type
content = message.content
else:
return None
if raw_role not in ("ai", "assistant"):
return None
if raw_id:
key = ("id", str(raw_id))
else:
key = ("body", str(raw_role), repr(content))
return _AssistantValueMessage(key=key, content=content)
def _process_value_messages(self, data: object) -> list[dict[str, Any]]:
data_map = _as_raw_map(data)
if data_map is None:
return []
messages = data_map.get("messages")
if not isinstance(messages, list):
return []
events: list[dict[str, Any]] = []
for message in messages:
parsed = self._assistant_value_message(message)
if parsed is None:
continue
if parsed.key in self._seen_value_message_keys:
continue
text = _text_from_content(parsed.content)
if text:
self._seen_value_message_keys.add(parsed.key)
events.extend(self._emit_text(text, subagent=None))
return events
def _process_message_event(
self,
data: object,
@@ -523,7 +604,7 @@ class _V3EventProcessor:
signature = _summarization_event_signature(summarization_event)
if (
signature is not None
and signature == self.baseline_summarization_signature
and signature == self._suppressed_summarization_signature
):
return events
summary_message = summarization_event.get("summary_message")
@@ -543,35 +624,56 @@ class _V3EventProcessor:
continue
interrupt_value = interrupt_obj.value
if not isinstance(interrupt_value, dict):
continue
iv_type = interrupt_value.get("type")
interrupt_id = interrupt_obj.id or "default"
if iv_type == "ask_user":
questions = interrupt_value.get("questions", [])
tc_id = str(interrupt_value.get("tool_call_id", ""))
events.extend(
self._dedupe_interrupt_event(
self.emitter.ask_user_interrupt(
interrupt_id, questions, tc_id
).data
)
)
continue
action_reqs = interrupt_value.get("action_requests", [])
review_cfgs = interrupt_value.get("review_configs", [])
if action_reqs:
events.extend(
self._dedupe_interrupt_event(
self.emitter.interrupt(
interrupt_id, action_reqs, review_cfgs
).data
)
)
events.extend(self._process_interrupt_value(interrupt_id, interrupt_value))
return events
def _process_input_requested(self, params: object) -> list[dict[str, Any]]:
params_map = _as_raw_map(params)
if params_map is None:
return []
data = _as_raw_map(params_map.get("data"))
if data is None:
return []
interrupt_id = str(data.get("interrupt_id") or "default")
return self._process_interrupt_value(interrupt_id, data.get("value"))
def _process_interrupt_value(
self,
interrupt_id: str,
interrupt_value: object,
) -> list[dict[str, Any]]:
interrupt_map = _as_raw_map(interrupt_value)
if interrupt_map is None:
return []
iv_type = interrupt_map.get("type")
if iv_type == "ask_user":
raw_questions = interrupt_map.get("questions")
questions = raw_questions if isinstance(raw_questions, list) else []
tc_id = str(interrupt_map.get("tool_call_id", ""))
return self._dedupe_interrupt_event(
self.emitter.ask_user_interrupt(
interrupt_id,
questions,
tc_id,
).data
)
raw_action_reqs = interrupt_map.get("action_requests")
action_reqs = raw_action_reqs if isinstance(raw_action_reqs, list) else []
raw_review_cfgs = interrupt_map.get("review_configs")
review_cfgs = raw_review_cfgs if isinstance(raw_review_cfgs, list) else None
if action_reqs:
return self._dedupe_interrupt_event(
self.emitter.interrupt(
interrupt_id,
action_reqs,
review_cfgs,
).data
)
return []
def _dedupe_interrupt_event(self, event: dict[str, Any]) -> list[dict[str, Any]]:
signature = repr(event)
if signature in self._emitted_interrupts:
@@ -631,9 +733,63 @@ class _V3EventProcessor:
return _text_from_content(payload.content)
async def build_agent_stream_input(
message: GraphRunInput,
*,
media: list[str] | None = None,
) -> LangGraphStreamInput:
"""Build the LangGraph run input shared by local and server gateways."""
if not isinstance(message, str):
return message
user_content: UserMessageContent = message
if media:
image_exts = frozenset({".jpg", ".jpeg", ".png", ".gif", ".webp", ".bmp"})
max_inline_size = 5 * 1024 * 1024
content_blocks: list[dict[str, object]] = []
if message:
content_blocks.append({"type": "text", "text": message})
def _read_file_b64(path: str) -> str:
with open(path, "rb") as fh:
return base64.b64encode(fh.read()).decode("ascii")
file_refs: list[str] = []
for path in media:
ext = os.path.splitext(path)[1].lower()
is_image = ext in image_exts and await asyncio.to_thread(
os.path.isfile, path
)
if is_image:
fsize = await asyncio.to_thread(os.path.getsize, path)
if fsize <= max_inline_size:
mime = mimetypes.guess_type(path)[0] or "image/png"
b64 = await asyncio.to_thread(_read_file_b64, path)
content_blocks.append(
{
"type": "image_url",
"image_url": {
"url": f"data:{mime};base64,{b64}",
},
}
)
else:
file_refs.append(path)
else:
file_refs.append(path)
if file_refs:
ref_text = "\n".join(
f"[attached file: {os.path.basename(p)}] path: {p}" for p in file_refs
)
content_blocks.append({"type": "text", "text": ref_text})
if content_blocks:
user_content = content_blocks
return {"messages": [{"role": "user", "content": user_content}]}
async def stream_agent_events(
agent: Any,
message: str | Command,
message: GraphRunInput,
thread_id: str,
metadata: dict[str, Any] | None = None,
media: list[str] | None = None,
@@ -662,74 +818,18 @@ async def stream_agent_events(
if metadata:
config["metadata"] = metadata
emitter = StreamEventEmitter()
clear_memory_worker_saved_counts()
# Build input for agent.astream_events()
if isinstance(message, str):
# Build user message content: text + inline images + file path references
user_content: str | list[dict[str, object]] = message
if media:
_IMAGE_EXTS = frozenset({".jpg", ".jpeg", ".png", ".gif", ".webp", ".bmp"})
_MAX_INLINE_SIZE = 5 * 1024 * 1024 # 5 MB
content_blocks: list[dict[str, object]] = []
if message:
content_blocks.append({"type": "text", "text": message})
def _read_file_b64(path: str) -> str:
with open(path, "rb") as fh:
return base64.b64encode(fh.read()).decode("ascii")
file_refs: list[str] = []
for path in media:
ext = os.path.splitext(path)[1].lower()
is_image = ext in _IMAGE_EXTS and await asyncio.to_thread(
os.path.isfile, path
)
if is_image:
fsize = await asyncio.to_thread(os.path.getsize, path)
if fsize <= _MAX_INLINE_SIZE:
mime = mimetypes.guess_type(path)[0] or "image/png"
b64 = await asyncio.to_thread(_read_file_b64, path)
content_blocks.append(
{
"type": "image_url",
"image_url": {
"url": f"data:{mime};base64,{b64}",
},
}
)
else:
file_refs.append(path)
else:
file_refs.append(path)
if file_refs:
ref_text = "\n".join(
f"[attached file: {os.path.basename(p)}] path: {p}"
for p in file_refs
)
content_blocks.append({"type": "text", "text": ref_text})
if content_blocks:
user_content = content_blocks
astream_input: dict[str, list[dict[str, object]]] | Command = {
"messages": [{"role": "user", "content": user_content}]
}
else:
# HITL resume: Command object passed directly to agent
astream_input = message
_baseline_summarization_signature: tuple[object, ...] | None = None
existing_summarization_event: Mapping[str, object] | None = None
try:
snapshot = await agent.aget_state(config)
values = snapshot.values
if isinstance(values, dict):
baseline_event = _find_summarization_event_payload(values)
_baseline_summarization_signature = _summarization_event_signature(
baseline_event
)
existing_summarization_event = _find_summarization_event_payload(
getattr(snapshot, "values", None)
)
except Exception:
pass
clear_memory_worker_saved_counts()
astream_input = await build_agent_stream_input(message, media=media)
stream: Any | None = None
producers: list[asyncio.Task[Any]] = []
_run_raised: bool = False
@@ -753,7 +853,7 @@ async def stream_agent_events(
processor = _V3EventProcessor(
emitter,
subagents,
_baseline_summarization_signature,
existing_summarization_event,
)
queue: asyncio.Queue[Any] = asyncio.Queue()
producer_done = object()
+653
View File
@@ -0,0 +1,653 @@
"""Shared test doubles for gateway/runtime boundaries."""
from __future__ import annotations
import asyncio
from collections.abc import AsyncIterator, Callable, Iterable
from dataclasses import dataclass
from typing import Any
import httpx
from langgraph_sdk.client import LangGraphClient
from EvoScientist.channels.base import Channel
from EvoScientist.channels.bus.events import InboundMessage, OutboundMessage
from EvoScientist.commands.base import CommandUI
from EvoScientist.gateway import (
GraphEvent,
GraphGateway,
GraphStateValues,
GraphTarget,
RunRequest,
ThreadResolution,
ThreadStore,
)
_DEFAULT_COPY_RESPONSE = object()
class FakeCommandUI(CommandUI):
"""Command UI test double with recorded calls and inert pickers."""
def __init__(self, *, supports_interactive: bool = True) -> None:
self._supports_interactive = supports_interactive
self.system_messages: list[str] = []
self.renderables: list[object] = []
self.started = 0
self.stopped = 0
self.updated_tokens: list[int] = []
self.chat_cleared = False
self.quit_requested = False
self.force_quit_requested = False
self.started_sessions = 0
self.resumed_sessions: list[tuple[str, str | None]] = []
self.flushes = 0
@property
def supports_interactive(self) -> bool:
return self._supports_interactive
def append_system(self, text: str, style: str = "dim") -> None:
self.system_messages.append(text)
def mount_renderable(self, renderable: object) -> None:
self.renderables.append(renderable)
async def wait_for_thread_pick(
self,
threads: list[dict],
current_thread: str,
title: str,
) -> str | None:
return None
async def wait_for_skill_browse(
self,
index: list[dict],
installed_names: set[str],
pre_filter_tag: str,
) -> list[str] | None:
return None
async def wait_for_mcp_browse(
self,
servers: list,
installed_names: set[str],
pre_filter_tag: str,
) -> list | None:
return None
async def wait_for_model_pick(
self,
entries: list[tuple[str, str, str]],
current_model: str | None,
current_provider: str | None,
) -> tuple[str, str] | None:
return None
def clear_chat(self) -> None:
self.chat_cleared = True
def request_quit(self) -> None:
self.quit_requested = True
def force_quit(self) -> None:
self.force_quit_requested = True
async def start_new_session(self) -> None:
self.started_sessions += 1
async def handle_session_resume(
self,
thread_id: str,
workspace_dir: str | None = None,
) -> None:
self.resumed_sessions.append((thread_id, workspace_dir))
async def flush(self) -> None:
self.flushes += 1
async def start_compacting_indicator(self) -> None:
self.started += 1
async def stop_compacting_indicator(self) -> None:
self.stopped += 1
def update_status_after_compact(self, tokens_after: int) -> None:
self.updated_tokens.append(tokens_after)
@dataclass
class FakeChannelConfig:
"""Minimal config surface consumed by channel base tests."""
text_chunk_limit: int = 4096
allowed_senders: list | None = None
allowed_channels: list | None = None
proxy: str | None = None
require_mention: str = "group"
dm_policy: str = "allowlist"
class StubChannel(Channel):
"""Minimal concrete channel for unit tests of channel base behavior."""
name = "stub"
def __init__(self, config: Any | None = None) -> None:
super().__init__(config or FakeChannelConfig())
self._sent_chunks: list[tuple] = []
self._typing_started: list[str] = []
self._typing_stopped: list[str] = []
self._started = False
async def start(self) -> None:
self._started = True
self._running = True
async def _send_chunk(
self,
chat_id: str,
formatted_text: str,
raw_text: str,
reply_to: str | None,
metadata: dict,
) -> None:
self._sent_chunks.append(
(chat_id, formatted_text, raw_text, reply_to, metadata)
)
async def _send_typing_action(self, chat_id: str) -> None:
self._typing_started.append(chat_id)
class QueueFakeChannel(Channel):
"""Concrete channel with queue receive and captured outbound messages."""
name = "fake"
def __init__(self, config: Any | None = None) -> None:
super().__init__(config or FakeChannelConfig())
self._started = False
self._stopped = False
self._sent: list[OutboundMessage] = []
async def start(self) -> None:
self._started = True
async def stop(self) -> None:
self._stopped = True
async def receive(self) -> AsyncIterator[InboundMessage]:
while True:
try:
msg = await asyncio.wait_for(self._queue.get(), timeout=0.5)
yield msg
except TimeoutError:
return
async def send(self, message: OutboundMessage) -> bool:
self._sent.append(message)
return True
async def _send_chunk(
self,
chat_id: str,
formatted_text: str,
raw_text: str,
reply_to: str | None,
metadata: dict,
) -> None:
pass
class FakeThreadStore(ThreadStore):
"""Configurable ``ThreadStore`` test double with call recording."""
def __init__(
self,
*,
generated_thread_id: str = "unused",
threads: list[dict[str, Any]] | None = None,
resolved_thread_id: str | None = None,
matches: list[str] | None = None,
metadata: dict[str, Any] | None = None,
messages: list[Any] | None = None,
exists: bool = False,
deleted: bool = False,
errors: dict[str, BaseException] | None = None,
) -> None:
self.generated_thread_id = generated_thread_id
self.threads = threads or []
self.resolved_thread_id = resolved_thread_id
self.matches = matches or []
self.metadata = metadata
self.messages = messages or []
self.exists = exists
self.deleted = deleted
self.errors = errors or {}
self.calls: list[tuple[str, Any]] = []
def _maybe_raise(self, method: str) -> None:
error = self.errors.get(method)
if error is not None:
raise error
def generate_thread_id(self) -> str:
self.calls.append(("generate_thread_id", None))
self._maybe_raise("generate_thread_id")
return self.generated_thread_id
async def list_threads(
self,
*,
limit: int = 20,
include_message_count: bool = False,
include_preview: bool = False,
) -> list[dict[str, Any]]:
self.calls.append(
(
"list_threads",
{
"limit": limit,
"include_message_count": include_message_count,
"include_preview": include_preview,
},
)
)
self._maybe_raise("list_threads")
return self.threads
async def resolve_thread_id_prefix(
self,
thread_id_or_prefix: str,
) -> tuple[str | None, list[str]]:
self.calls.append(("resolve_thread_id_prefix", thread_id_or_prefix))
self._maybe_raise("resolve_thread_id_prefix")
return self.resolved_thread_id, self.matches
async def get_thread_metadata(self, thread_id: str) -> dict[str, Any] | None:
self.calls.append(("get_thread_metadata", thread_id))
self._maybe_raise("get_thread_metadata")
return self.metadata
async def get_thread_messages(self, thread_id: str) -> list[Any]:
self.calls.append(("get_thread_messages", thread_id))
self._maybe_raise("get_thread_messages")
return self.messages
async def thread_exists(self, thread_id: str) -> bool:
self.calls.append(("thread_exists", thread_id))
self._maybe_raise("thread_exists")
return self.exists
async def delete_thread(self, thread_id: str) -> bool:
self.calls.append(("delete_thread", thread_id))
self._maybe_raise("delete_thread")
return self.deleted
FakeStreamFactory = Callable[[RunRequest], AsyncIterator[GraphEvent]]
class FakeGraphGateway(GraphGateway):
"""Configurable graph gateway test double with request recording."""
def __init__(
self,
events: Iterable[GraphEvent] | None = None,
*,
stream: FakeStreamFactory | None = None,
state_values: GraphStateValues | None = None,
state_error: BaseException | None = None,
generated_thread_ids: Iterable[str] | None = None,
thread_store: ThreadStore | None = None,
) -> None:
self.events = list(events or [])
self.stream = stream
self.state_values = state_values or {}
self.state_error = state_error
self.generated_thread_ids = list(generated_thread_ids or [])
self.thread_store = thread_store or FakeThreadStore()
self.requests: list[RunRequest] = []
self.clone_calls: list[
tuple[str, dict[str, Any] | None, GraphTarget | None]
] = []
self.updated_states: list[tuple[GraphTarget, str, GraphStateValues]] = []
async def create_thread(
self,
target: GraphTarget | None = None,
*,
metadata: dict[str, Any] | None = None,
) -> str:
if self.generated_thread_ids:
return self.generated_thread_ids.pop(0)
return self.thread_store.generate_thread_id()
async def list_threads(
self,
*,
limit: int = 20,
include_message_count: bool = False,
include_preview: bool = False,
target: GraphTarget | None = None,
) -> list[dict[str, Any]]:
return await self.thread_store.list_threads(
limit=limit,
include_message_count=include_message_count,
include_preview=include_preview,
)
async def resolve_thread(
self,
thread_id_or_prefix: str,
target: GraphTarget | None = None,
) -> ThreadResolution:
resolved, matches = await self.thread_store.resolve_thread_id_prefix(
thread_id_or_prefix
)
return ThreadResolution(resolved, tuple(matches))
async def get_thread_metadata(
self,
thread_id: str,
target: GraphTarget | None = None,
) -> dict[str, Any] | None:
return await self.thread_store.get_thread_metadata(thread_id)
async def get_thread_messages(
self,
thread_id: str,
target: GraphTarget | None = None,
) -> list[Any]:
return await self.thread_store.get_thread_messages(thread_id)
async def thread_exists(
self,
thread_id: str,
target: GraphTarget | None = None,
) -> bool:
return await self.thread_store.thread_exists(thread_id)
async def delete_thread(
self,
thread_id: str,
target: GraphTarget | None = None,
) -> bool:
return await self.thread_store.delete_thread(thread_id)
async def clone_thread(
self,
source_thread_id: str,
*,
metadata: dict[str, Any] | None = None,
target: GraphTarget | None = None,
) -> str:
self.clone_calls.append((source_thread_id, metadata, target))
if self.generated_thread_ids:
return self.generated_thread_ids.pop(0)
return f"{source_thread_id}-clone"
def stream_events(self, request: RunRequest) -> AsyncIterator[GraphEvent]:
self.requests.append(request)
if self.stream is not None:
return self.stream(request)
async def _events() -> AsyncIterator[GraphEvent]:
for event in self.events:
yield event
return _events()
async def get_state_values(
self,
target: GraphTarget,
thread_id: str,
) -> GraphStateValues:
if self.state_error is not None:
raise self.state_error
return self.state_values
async def update_state_values(
self,
target: GraphTarget,
thread_id: str,
values: GraphStateValues,
) -> None:
self.updated_states.append((target, thread_id, values))
class FakeLangGraphRunModule:
"""Fake thread-stream run controller for server gateway tests."""
def __init__(self) -> None:
self.starts: list[dict[str, Any]] = []
self.responses: list[dict[str, Any]] = []
async def start(
self,
*,
input: object = None,
config: dict[str, Any] | None = None,
metadata: dict[str, Any] | None = None,
) -> dict[str, Any]:
self.starts.append(
{
"input": input,
"config": config,
"metadata": metadata,
}
)
return {"run_id": "run-1"}
async def respond(
self,
response: object,
*,
interrupt_id: str | None = None,
) -> dict[str, Any]:
self.responses.append(
{
"response": response,
"interrupt_id": interrupt_id,
}
)
return {"run_id": "run-1"}
class FakeLangGraphThreadStream:
"""Finite fake of the LangGraph SDK thread stream."""
def __init__(
self,
thread_id: str,
events: Iterable[dict[str, Any]] | None = None,
*,
interrupts: list[dict[str, Any]] | None = None,
interrupted: bool = False,
) -> None:
self.thread_id = thread_id
self.events = list(events or [])
self.interrupts = interrupts or []
self.interrupted = interrupted
self.run = FakeLangGraphRunModule()
self.subscribed_channels: list[list[str]] = []
self.entered = False
self.exited = False
async def __aenter__(self) -> FakeLangGraphThreadStream:
self.entered = True
return self
async def __aexit__(self, exc_type: object, exc: object, tb: object) -> None:
self.exited = True
async def _iter_events(self) -> AsyncIterator[dict[str, Any]]:
for event in self.events:
yield event
def subscribe(self, channels: list[str]) -> AsyncIterator[dict[str, Any]]:
self.subscribed_channels.append(channels)
return self._iter_events()
class FakeLangGraphThreadsClient:
"""Fake LangGraph ``client.threads`` surface."""
def __init__(
self,
*,
threads: list[dict[str, Any]] | None = None,
states: dict[str, dict[str, Any]] | None = None,
streams: dict[str, FakeLangGraphThreadStream] | None = None,
copy_response: object = _DEFAULT_COPY_RESPONSE,
) -> None:
self.threads = threads or []
self.states = states or {}
self.streams = streams or {}
self.copy_response = copy_response
self.created: list[dict[str, Any]] = []
self.copied: list[str] = []
self.metadata_updates: list[tuple[str, dict[str, Any]]] = []
self.deleted: list[str] = []
self.gets: list[str] = []
self.searches: list[dict[str, Any]] = []
self.stream_calls: list[tuple[str, str]] = []
self.state_updates: list[tuple[str, GraphStateValues, str | None]] = []
async def create(
self,
*,
metadata: dict[str, Any] | None = None,
thread_id: str | None = None,
if_exists: str | None = None,
graph_id: str | None = None,
) -> dict[str, Any]:
if thread_id is not None and if_exists == "do_nothing":
for thread in self.threads:
if thread.get("thread_id") == thread_id:
return thread
created = {
"thread_id": thread_id or "server-thread",
"metadata": {
**(metadata or {}),
**({"graph_id": graph_id} if graph_id else {}),
},
}
self.created.append(created)
self.threads.append(created)
return created
async def search(
self,
*,
metadata: dict[str, Any] | None = None,
limit: int = 10,
offset: int = 0,
sort_by: str | None = None,
sort_order: str | None = None,
) -> list[dict[str, Any]]:
self.searches.append(
{
"metadata": metadata,
"limit": limit,
"offset": offset,
"sort_by": sort_by,
"sort_order": sort_order,
}
)
rows = self.threads
if metadata:
rows = [
thread
for thread in rows
if all(
(thread.get("metadata") or {}).get(key) == value
for key, value in metadata.items()
)
]
return rows[offset : offset + limit]
async def get(self, thread_id: str) -> dict[str, Any]:
from langgraph_sdk.errors import NotFoundError
self.gets.append(thread_id)
for thread in self.threads:
if thread.get("thread_id") == thread_id:
return thread
raise NotFoundError("not found", response=_not_found_response(), body=None)
async def copy(self, thread_id: str) -> object:
source = await self.get(thread_id)
self.copied.append(thread_id)
if self.copy_response is not _DEFAULT_COPY_RESPONSE:
return self.copy_response
copied = {
"thread_id": f"{thread_id}-copy",
"metadata": dict(source.get("metadata") or {}),
}
self.threads.append(copied)
return copied
async def update(
self,
thread_id: str,
*,
metadata: dict[str, Any],
) -> dict[str, Any]:
thread = await self.get(thread_id)
existing_metadata = thread.get("metadata")
merged = {
**(existing_metadata if isinstance(existing_metadata, dict) else {}),
**metadata,
}
thread["metadata"] = merged
self.metadata_updates.append((thread_id, metadata))
return thread
async def get_state(self, thread_id: str) -> dict[str, Any]:
from langgraph_sdk.errors import NotFoundError
if thread_id in self.states:
return self.states[thread_id]
raise NotFoundError("not found", response=_not_found_response(), body=None)
async def update_state(
self,
thread_id: str,
values: GraphStateValues,
*,
as_node: str | None = None,
) -> dict[str, Any]:
self.state_updates.append((thread_id, values, as_node))
return {"checkpoint": {"thread_id": thread_id}}
async def delete(self, thread_id: str) -> None:
await self.get(thread_id)
self.deleted.append(thread_id)
self.threads = [
thread for thread in self.threads if thread.get("thread_id") != thread_id
]
def stream(
self,
thread_id: str | None = None,
*,
assistant_id: str,
) -> FakeLangGraphThreadStream:
assert thread_id is not None
self.stream_calls.append((thread_id, assistant_id))
return self.streams[thread_id]
class FakeLangGraphClient(LangGraphClient):
"""Fake LangGraph SDK async client."""
def __init__(self, threads: FakeLangGraphThreadsClient) -> None:
self.threads = threads
def _not_found_response() -> httpx.Response:
request = httpx.Request("GET", "https://test.local/not-found")
return httpx.Response(404, request=request)
+10 -2
View File
@@ -17,12 +17,20 @@ async def async_iter(items: Iterable[Any]) -> AsyncIterator[Any]:
yield item
def collect_events(agent, message: str = "hi", thread_id: str = "t1"):
def collect_events(
agent,
message: str = "hi",
thread_id: str = "t1",
):
"""Collect stream_agent_events output for synchronous tests."""
async def _run():
events = []
async for ev in stream_agent_events(agent, message, thread_id):
async for ev in stream_agent_events(
agent,
message,
thread_id,
):
events.append(ev)
return events
+23 -1
View File
@@ -12,6 +12,8 @@ from EvoScientist.cli.async_notifier import (
format_batch_message,
format_notification_lines,
)
from EvoScientist.gateway import GraphTarget
from tests.fakes import FakeGraphGateway
def test_notification_dataclass_fields():
@@ -49,6 +51,26 @@ def _drain_queue(q):
return items
def test_read_async_tasks_from_gateway_reads_state_values(run_async):
gateway = FakeGraphGateway(
state_values={
"async_tasks": {
"task-1": {"status": "success"},
}
}
)
tasks = run_async(
async_notifier.read_async_tasks_from_gateway(
gateway,
GraphTarget(local_graph=MagicMock()),
"tid",
)
)
assert tasks == {"task-1": {"status": "success"}}
def test_watcher_pushes_notification_on_stream_end(run_async):
# Stream yields one "values" chunk with the final state, then closes
final_state = {
@@ -322,7 +344,7 @@ def test_drain_returns_all_pending_and_empties_queue():
def test_dedup_skips_tasks_already_checked_after_terminal():
"""dedup_notifications skips tasks with terminal status and last_checked_at >= last_updated_at."""
async_tasks = {
async_tasks: async_notifier.AsyncTasksState = {
"a": {
"status": "success",
"last_checked_at": "2026-05-06T12:01:00Z",
+1 -39
View File
@@ -9,11 +9,11 @@ import asyncio
import pytest
from EvoScientist.channels.base import Channel, OutgoingMessage
from EvoScientist.channels.bus.events import InboundMessage
from EvoScientist.channels.bus.message_bus import MessageBus
from EvoScientist.channels.channel_manager import ChannelManager
from tests.conftest import run_async as _run
from tests.fakes import QueueFakeChannel as FakeChannel
def _drain_queue(q):
@@ -55,44 +55,6 @@ def clean_channel_state():
_reset()
class _FakeConfig:
text_chunk_limit = 4096
allowed_senders = None
class FakeChannel(Channel):
"""Minimal channel for bus integration testing."""
name = "fake"
def __init__(self):
super().__init__(_FakeConfig())
self._started = False
self._stopped = False
self._sent: list[OutgoingMessage] = []
async def start(self):
self._started = True
async def stop(self):
self._stopped = True
async def receive(self):
while True:
try:
msg = await asyncio.wait_for(self._queue.get(), timeout=0.5)
yield msg
except TimeoutError:
return
async def send(self, message: OutgoingMessage) -> bool:
self._sent.append(message)
return True
async def _send_chunk(self, chat_id, formatted_text, raw_text, reply_to, metadata):
pass
class TestBusInboundConsumer:
"""Test the _bus_inbound_consumer queue bridge."""
+36 -27
View File
@@ -7,10 +7,12 @@ from unittest.mock import AsyncMock, patch
import pytest
from EvoScientist.commands.channel_ui import ChannelCommandUI
from EvoScientist.gateway import ThreadStore
from tests.conftest import run_async as _run
from tests.fakes import FakeGraphGateway, FakeThreadStore
def _make_ui(callback=None, bus_ref=None):
def _make_ui(*, thread_store: ThreadStore, callback=None, bus_ref=None):
captured: list[str] = []
ui = ChannelCommandUI(
SimpleNamespace(
@@ -23,6 +25,7 @@ def _make_ui(callback=None, bus_ref=None):
),
append_system_callback=lambda text, style="dim": captured.append(text),
handle_session_resume_callback=callback,
graph_gateway=FakeGraphGateway(thread_store=thread_store),
)
return ui, captured
@@ -57,20 +60,22 @@ def _sent_text(bus_ref) -> str:
def test_handle_session_resume_sends_history_back_to_channel_without_local_duplicate():
callback = AsyncMock()
bus_ref = SimpleNamespace(publish_outbound=AsyncMock())
ui, captured = _make_ui(callback=callback, bus_ref=bus_ref)
messages = [
SimpleNamespace(type="human", content="How does this work?"),
SimpleNamespace(type="ai", content="Here is the saved answer."),
]
thread_store = FakeThreadStore(messages=messages)
ui, captured = _make_ui(
callback=callback,
bus_ref=bus_ref,
thread_store=thread_store,
)
with patch(
"EvoScientist.sessions.get_thread_messages",
new=AsyncMock(return_value=messages),
):
_run(_run_resume(ui, "thread-42", "/workspace"))
_run(_run_resume(ui, "thread-42", "/workspace"))
callback.assert_awaited_once_with("thread-42", "/workspace")
assert thread_store.calls == [("get_thread_messages", "thread-42")]
assert captured == []
text = _sent_text(bus_ref)
assert "Resumed session: thread-42" in text
@@ -82,17 +87,18 @@ def test_handle_session_resume_sends_history_back_to_channel_without_local_dupli
def test_handle_session_resume_propagates_callback_abort_without_history():
callback = AsyncMock(side_effect=RuntimeError("workspace conflict"))
bus_ref = SimpleNamespace(publish_outbound=AsyncMock())
ui, captured = _make_ui(callback=callback, bus_ref=bus_ref)
thread_store = FakeThreadStore()
ui, captured = _make_ui(
callback=callback,
bus_ref=bus_ref,
thread_store=thread_store,
)
with patch(
"EvoScientist.sessions.get_thread_messages",
new=AsyncMock(),
) as get_messages:
with pytest.raises(RuntimeError, match="workspace conflict"):
_run(_run_resume(ui, "thread-42", "/workspace"))
with pytest.raises(RuntimeError, match="workspace conflict"):
_run(_run_resume(ui, "thread-42", "/workspace"))
callback.assert_awaited_once_with("thread-42", "/workspace")
get_messages.assert_not_awaited()
assert thread_store.calls == []
bus_ref.publish_outbound.assert_not_awaited()
assert captured == []
@@ -100,13 +106,15 @@ def test_handle_session_resume_propagates_callback_abort_without_history():
def test_handle_session_resume_reports_history_load_error():
callback = AsyncMock()
bus_ref = SimpleNamespace(publish_outbound=AsyncMock())
ui, captured = _make_ui(callback=callback, bus_ref=bus_ref)
ui, captured = _make_ui(
callback=callback,
bus_ref=bus_ref,
thread_store=FakeThreadStore(
errors={"get_thread_messages": RuntimeError("db locked")}
),
)
with patch(
"EvoScientist.sessions.get_thread_messages",
new=AsyncMock(side_effect=RuntimeError("db locked")),
):
_run(_run_resume(ui, "thread-42", "/workspace"))
_run(_run_resume(ui, "thread-42", "/workspace"))
callback.assert_awaited_once_with("thread-42", "/workspace")
assert captured == []
@@ -117,13 +125,14 @@ def test_handle_session_resume_reports_history_load_error():
def test_handle_session_resume_distinguishes_non_displayable_messages():
bus_ref = SimpleNamespace(publish_outbound=AsyncMock())
ui, captured = _make_ui(bus_ref=bus_ref)
ui, captured = _make_ui(
bus_ref=bus_ref,
thread_store=FakeThreadStore(
messages=[SimpleNamespace(type="tool", content="hidden")]
),
)
with patch(
"EvoScientist.sessions.get_thread_messages",
new=AsyncMock(return_value=[SimpleNamespace(type="tool", content="hidden")]),
):
_run(_run_resume(ui, "thread-42", "/workspace"))
_run(_run_resume(ui, "thread-42", "/workspace"))
assert captured == [
"Resumed session: thread-42\nNo displayable messages in this session."
+20 -45
View File
@@ -16,14 +16,12 @@ from __future__ import annotations
import asyncio
import time
from dataclasses import dataclass
from datetime import datetime
from unittest.mock import AsyncMock, MagicMock
import pytest
from EvoScientist.channels.base import (
Channel,
ChannelError,
InboundMessage,
OutboundMessage,
@@ -47,40 +45,8 @@ from EvoScientist.channels.retry import RetryConfig, RetryInfo, retry_async
# Helpers
# ═══════════════════════════════════════════════════════════════════
from tests.conftest import run_async as _run
@dataclass
class _FakeConfig:
text_chunk_limit: int = 4096
allowed_senders: list | None = None
allowed_channels: list | None = None
proxy: str | None = None
require_mention: str = "group"
dm_policy: str = "allowlist"
class StubChannel(Channel):
"""Minimal concrete channel for unit testing."""
name = "stub"
def __init__(self, config=None):
super().__init__(config or _FakeConfig())
self._sent_chunks: list[tuple] = []
self._typing_started: list[str] = []
self._typing_stopped: list[str] = []
self._started = False
async def start(self):
self._started = True
self._running = True
async def _send_chunk(self, chat_id, formatted, raw, reply_to, metadata):
self._sent_chunks.append((chat_id, formatted, raw, reply_to, metadata))
async def _send_typing_action(self, chat_id):
self._typing_started.append(chat_id)
from tests.fakes import FakeChannelConfig as _FakeConfig
from tests.fakes import FakeGraphGateway, StubChannel
# ═══════════════════════════════════════════════════════════════════
# 1. DedupCache
@@ -1395,6 +1361,7 @@ class TestInboundConsumer:
mgr.register(StubChannel())
if agent is None:
agent = MagicMock()
kw.setdefault("graph_gateway", FakeGraphGateway())
return InboundConsumer(
bus=bus,
manager=mgr,
@@ -1412,15 +1379,21 @@ class TestInboundConsumer:
assert msg.session_key == "tg:c1"
def test_get_thread_id_creates_unique(self):
consumer = self._make_consumer()
tid1 = consumer._get_thread_id("user_a")
tid2 = consumer._get_thread_id("user_b")
consumer = self._make_consumer(
graph_gateway=FakeGraphGateway(
generated_thread_ids=["thread-a", "thread-b"]
)
)
tid1 = _run(consumer._get_thread_id("user_a"))
tid2 = _run(consumer._get_thread_id("user_b"))
assert tid1 != tid2
def test_get_thread_id_returns_same_for_same_sender(self):
consumer = self._make_consumer()
tid1 = consumer._get_thread_id("user_a")
tid2 = consumer._get_thread_id("user_a")
consumer = self._make_consumer(
graph_gateway=FakeGraphGateway(generated_thread_ids=["thread-a"])
)
tid1 = _run(consumer._get_thread_id("user_a"))
tid2 = _run(consumer._get_thread_id("user_a"))
assert tid1 == tid2
def test_shared_thread_id_bug(self):
@@ -1433,9 +1406,10 @@ class TestInboundConsumer:
manager=mgr,
agent=MagicMock(),
thread_id="shared_thread", # Non-empty!
graph_gateway=FakeGraphGateway(),
)
tid1 = consumer._get_thread_id("alice")
tid2 = consumer._get_thread_id("bob")
tid1 = _run(consumer._get_thread_id("alice"))
tid2 = _run(consumer._get_thread_id("bob"))
# Fixed: Each sender gets a unique thread_id using thread_id as prefix
assert tid1 != tid2
assert tid1 == "shared_thread:alice"
@@ -1451,7 +1425,7 @@ class TestInboundConsumer:
consumer._sessions[f"user_{i}"] = f"thread_{i}"
# Access "user_0" via _get_thread_id (triggers LRU move_to_end)
consumer._get_thread_id("user_0")
_run(consumer._get_thread_id("user_0"))
# "user_0" should now be at the end (most recently used)
oldest = next(iter(consumer._sessions))
@@ -1494,6 +1468,7 @@ class TestInboundConsumerErrorHandling:
manager=mgr,
agent=MagicMock(),
thread_id="",
graph_gateway=FakeGraphGateway(),
)
# The error message format includes the raw exception
+55 -2
View File
@@ -10,9 +10,21 @@ from unittest.mock import AsyncMock, MagicMock, patch
from EvoScientist.cli.channel import (
ChannelMessage,
dispatch_channel_slash_command,
)
from EvoScientist.cli.channel import (
dispatch_channel_slash_command as _dispatch_channel_slash_command,
)
from tests.conftest import run_async as _run
from tests.fakes import FakeGraphGateway, FakeThreadStore
def _thread_store() -> FakeThreadStore:
return FakeThreadStore()
def dispatch_channel_slash_command(*args, **kwargs):
kwargs.setdefault("graph_gateway", FakeGraphGateway())
return _dispatch_channel_slash_command(*args, **kwargs)
def _make_msg(
@@ -107,6 +119,45 @@ def test_successful_slash_execution_sets_response_and_breadcrumb():
assert any("Executed command from" in t for t in breadcrumbs)
def test_slash_dispatch_passes_graph_gateway_to_command_context():
msg = _make_msg()
fake_cmd = MagicMock()
fake_cmd.needs_agent.return_value = False
append = MagicMock()
graph_gateway = FakeGraphGateway(thread_store=_thread_store())
captured = {}
async def _execute(_content, ctx):
captured["graph_gateway"] = ctx.graph_gateway
return True
with (
patch(
"EvoScientist.commands.manager.manager.resolve",
return_value=(fake_cmd, ["core"]),
),
patch(
"EvoScientist.commands.manager.manager.execute",
new=AsyncMock(side_effect=_execute),
),
patch("EvoScientist.cli.channel._set_channel_response"),
):
handled = _run(
dispatch_channel_slash_command(
msg,
agent="fake-agent",
thread_id="t1",
workspace_dir="/tmp",
checkpointer=None,
append_system=append,
graph_gateway=graph_gateway,
)
)
assert handled is True
assert captured["graph_gateway"] is graph_gateway
def test_needs_agent_awaits_loader_and_passes_result():
"""Commands with needs_agent=True must await the loader and the
resulting agent must flow through the CommandContext."""
@@ -144,7 +195,9 @@ def test_needs_agent_awaits_loader_and_passes_result():
)
assert handled is True
await_called.assert_called_once()
ctx_arg = mock_execute.await_args.args[1]
await_args = mock_execute.await_args
assert await_args is not None
ctx_arg = await_args.args[1]
assert ctx_arg.agent == "ready-agent"
+55 -103
View File
@@ -3,44 +3,45 @@
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
from EvoScientist.gateway import GraphTarget
from tests.conftest import run_async as _run
from tests.fakes import FakeCommandUI, FakeGraphGateway
_TARGET = GraphTarget()
def _compact(
graph_gateway: FakeGraphGateway,
*,
thread_id: str = "tid-1",
input_tokens_hint: int | None = None,
):
from EvoScientist.cli.commands import compact_conversation
return _run(
compact_conversation(
graph_gateway=graph_gateway,
thread_id=thread_id,
target=_TARGET,
input_tokens_hint=input_tokens_hint,
)
)
class TestCompactGuards:
"""Guard conditions that return early without touching the middleware."""
def test_no_agent(self):
from EvoScientist.cli.commands import compact_conversation
result = _run(compact_conversation(agent=None, thread_id="abc"))
assert result.status == "noop"
assert "Nothing to compact" in result.message
def test_no_thread_id(self):
from EvoScientist.cli.commands import compact_conversation
result = _run(compact_conversation(agent=MagicMock(), thread_id=None))
assert result.status == "noop"
assert "Nothing to compact" in result.message
def test_empty_messages(self):
from EvoScientist.cli.commands import compact_conversation
graph_gateway = FakeGraphGateway(state_values={"messages": []})
agent = MagicMock()
snapshot = SimpleNamespace(values={"messages": []})
agent.aget_state = AsyncMock(return_value=snapshot)
result = _run(compact_conversation(agent=agent, thread_id="tid-1"))
result = _compact(graph_gateway)
assert result.status == "noop"
assert "no messages" in result.message
def test_state_read_failure(self):
from EvoScientist.cli.commands import compact_conversation
graph_gateway = FakeGraphGateway(state_error=RuntimeError("DB gone"))
agent = MagicMock()
agent.aget_state = AsyncMock(side_effect=RuntimeError("DB gone"))
result = _run(compact_conversation(agent=agent, thread_id="tid-1"))
result = _compact(graph_gateway)
assert result.status == "error"
assert "Failed to read state" in result.message
@@ -49,12 +50,8 @@ class TestCompactCutoffZero:
"""When cutoff == 0, conversation is within retention budget."""
def test_nothing_to_compact_short_conversation(self):
from EvoScientist.cli.commands import compact_conversation
agent = MagicMock()
msgs = [MagicMock() for _ in range(3)]
snapshot = SimpleNamespace(values={"messages": msgs})
agent.aget_state = AsyncMock(return_value=snapshot)
graph_gateway = FakeGraphGateway(state_values={"messages": msgs})
mock_middleware_inst = MagicMock()
mock_middleware_inst._apply_event_to_messages.return_value = msgs
@@ -82,7 +79,7 @@ class TestCompactCutoffZero:
return_value=500,
),
):
result = _run(compact_conversation(agent=agent, thread_id="tid-1"))
result = _compact(graph_gateway)
assert result.status == "noop"
assert "within the retention budget" in result.message
@@ -93,14 +90,10 @@ class TestCompactNegligibleSavings:
"""When cutoff > 0 but savings are too small to be worth it."""
def test_skip_when_few_messages_and_low_tokens(self):
from EvoScientist.cli.commands import compact_conversation
agent = MagicMock()
msgs = [MagicMock() for _ in range(15)]
snapshot = SimpleNamespace(
values={"messages": msgs, "_summarization_event": None}
graph_gateway = FakeGraphGateway(
state_values={"messages": msgs, "_summarization_event": None}
)
agent.aget_state = AsyncMock(return_value=snapshot)
mock_middleware_inst = MagicMock()
mock_middleware_inst._apply_event_to_messages.return_value = msgs
@@ -133,7 +126,7 @@ class TestCompactNegligibleSavings:
side_effect=lambda x: next(token_values),
),
):
result = _run(compact_conversation(agent=agent, thread_id="tid-1"))
result = _compact(graph_gateway)
assert result.status == "noop"
assert "not worth" in result.message
@@ -144,15 +137,10 @@ class TestCompactNegligibleSavings:
"""2 messages but they account for >2% of tokens — should compact."""
from langchain_core.messages import HumanMessage
from EvoScientist.cli.commands import compact_conversation
agent = MagicMock()
msgs = [MagicMock() for _ in range(10)]
snapshot = SimpleNamespace(
values={"messages": msgs, "_summarization_event": None}
graph_gateway = FakeGraphGateway(
state_values={"messages": msgs, "_summarization_event": None}
)
agent.aget_state = AsyncMock(return_value=snapshot)
agent.aupdate_state = AsyncMock()
summary_msg = HumanMessage(content="Summary")
@@ -190,24 +178,20 @@ class TestCompactNegligibleSavings:
side_effect=lambda x: next(token_values),
),
):
result = _run(compact_conversation(agent=agent, thread_id="tid-1"))
result = _compact(graph_gateway)
assert result.status == "ok"
agent.aupdate_state.assert_awaited_once()
assert len(graph_gateway.updated_states) == 1
class TestCompactSuccess:
"""Normal compaction flow."""
def test_manual_threshold_blocks_low_context_compaction(self):
from EvoScientist.cli.commands import compact_conversation
agent = MagicMock()
msgs = [MagicMock() for _ in range(20)]
snapshot = SimpleNamespace(
values={"messages": msgs, "_summarization_event": None}
graph_gateway = FakeGraphGateway(
state_values={"messages": msgs, "_summarization_event": None}
)
agent.aget_state = AsyncMock(return_value=snapshot)
mock_middleware_inst = MagicMock()
mock_middleware_inst._apply_event_to_messages.return_value = msgs
@@ -233,7 +217,7 @@ class TestCompactSuccess:
return_value=30_000,
),
):
result = _run(compact_conversation(agent=agent, thread_id="tid-1"))
result = _compact(graph_gateway)
assert result.status == "noop"
assert "40%" in result.message
@@ -244,15 +228,10 @@ class TestCompactSuccess:
def test_successful_compaction(self):
from langchain_core.messages import HumanMessage
from EvoScientist.cli.commands import compact_conversation
agent = MagicMock()
msgs = [MagicMock() for _ in range(20)]
snapshot = SimpleNamespace(
values={"messages": msgs, "_summarization_event": None}
graph_gateway = FakeGraphGateway(
state_values={"messages": msgs, "_summarization_event": None}
)
agent.aget_state = AsyncMock(return_value=snapshot)
agent.aupdate_state = AsyncMock()
summary_msg = HumanMessage(content="Summary of conversation")
to_summarize = msgs[:15]
@@ -294,7 +273,7 @@ class TestCompactSuccess:
side_effect=lambda x: next(token_values),
),
):
result = _run(compact_conversation(agent=agent, thread_id="tid-1"))
result = _compact(graph_gateway)
assert result.status == "ok"
assert result.messages_compacted == 15
@@ -305,11 +284,10 @@ class TestCompactSuccess:
# context_percent reflects usage AFTER compact (12%), not before (60%)
assert result.context_percent == 12
assert result.summary_text == "Summary text"
agent.aupdate_state.assert_awaited_once()
assert len(graph_gateway.updated_states) == 1
# Verify the event structure passed to aupdate_state
call_args = agent.aupdate_state.call_args
event_data = call_args[0][1]
# Verify the event structure passed through the graph gateway.
event_data = graph_gateway.updated_states[0][2]
assert "_summarization_event" in event_data
assert event_data["_summarization_event"]["cutoff_index"] == 15
@@ -317,15 +295,10 @@ class TestCompactSuccess:
"""Offload failure should not prevent compaction."""
from langchain_core.messages import HumanMessage
from EvoScientist.cli.commands import compact_conversation
agent = MagicMock()
msgs = [MagicMock() for _ in range(10)]
snapshot = SimpleNamespace(
values={"messages": msgs, "_summarization_event": None}
graph_gateway = FakeGraphGateway(
state_values={"messages": msgs, "_summarization_event": None}
)
agent.aget_state = AsyncMock(return_value=snapshot)
agent.aupdate_state = AsyncMock()
summary_msg = HumanMessage(content="Summary")
@@ -362,13 +335,13 @@ class TestCompactSuccess:
return_value=1000,
),
):
result = _run(compact_conversation(agent=agent, thread_id="tid-1"))
result = _compact(graph_gateway)
assert result.status == "ok"
agent.aupdate_state.assert_awaited_once()
assert len(graph_gateway.updated_states) == 1
# file_path should be None in the event
event_data = agent.aupdate_state.call_args[0][1]
event_data = graph_gateway.updated_states[0][2]
assert event_data["_summarization_event"]["file_path"] is None
@@ -409,36 +382,15 @@ class TestCompactCommandUI:
from EvoScientist.commands.base import CommandContext
from EvoScientist.commands.implementation.session import CompactCommand
class _UI:
supports_interactive = True
def __init__(self) -> None:
self.system_messages: list[str] = []
self.renderables: list[object] = []
self.started = 0
self.stopped = 0
self.updated_tokens: list[int] = []
def append_system(self, text: str, style: str = "dim") -> None:
self.system_messages.append(text)
def mount_renderable(self, renderable):
self.renderables.append(renderable)
async def start_compacting_indicator(self) -> None:
self.started += 1
async def stop_compacting_indicator(self) -> None:
self.stopped += 1
def update_status_after_compact(self, tokens_after: int) -> None:
self.updated_tokens.append(tokens_after)
ui = _UI()
ui = FakeCommandUI()
# input_tokens_hint must be set for update_status_after_compact to fire
# (without it, tokens_after is message-level and the unit would be wrong)
ctx = CommandContext(
agent=MagicMock(), thread_id="tid-1", ui=ui, input_tokens_hint=5000
agent=MagicMock(),
thread_id="tid-1",
ui=ui,
graph_gateway=FakeGraphGateway(),
input_tokens_hint=5000,
)
result = CompactResult(
"ok",
+33 -74
View File
@@ -1,87 +1,47 @@
"""Tests for the /delete command."""
from unittest.mock import AsyncMock, MagicMock, patch
from unittest.mock import AsyncMock, MagicMock
from tests.conftest import run_async as _run
from tests.fakes import FakeGraphGateway, FakeThreadStore
def _ctx(thread_id="current"):
def _ctx(thread_id="current", thread_store=None):
from EvoScientist.commands.base import CommandContext
store = thread_store or FakeThreadStore()
ui = MagicMock()
ui.supports_interactive = True
return CommandContext(agent=None, thread_id=thread_id, ui=ui), ui
def _patches(thread_exists=False, similar=None, deleted=True, threads=None):
"""Return a context manager stack patching the sessions module."""
from contextlib import ExitStack
stack = ExitStack()
stack.enter_context(
patch(
"EvoScientist.sessions.thread_exists",
new=AsyncMock(return_value=thread_exists),
)
)
stack.enter_context(
patch(
"EvoScientist.sessions.find_similar_threads",
new=AsyncMock(return_value=similar or []),
)
)
stack.enter_context(
patch(
"EvoScientist.sessions.delete_thread",
new=AsyncMock(return_value=deleted),
)
)
stack.enter_context(
patch(
"EvoScientist.sessions.list_threads",
new=AsyncMock(return_value=threads or []),
)
)
return stack
return CommandContext(
agent=None,
thread_id=thread_id,
ui=ui,
graph_gateway=FakeGraphGateway(thread_store=store),
), ui
class TestDeleteCommand:
def test_refuses_to_delete_current(self):
from EvoScientist.commands.implementation.session import DeleteCommand
ctx, ui = _ctx(thread_id="current")
# Inline the patches here (rather than using ``_patches``) so we
# can keep a direct handle on the ``delete_thread`` mock and
# assert on it *inside* the context. Asserting after the
# context exits hits the real function (no ``await_count``
# attr), which silently degrades into ``assert True``.
mock_delete = AsyncMock(return_value=True)
with (
patch(
"EvoScientist.sessions.thread_exists",
new=AsyncMock(return_value=True),
),
patch("EvoScientist.sessions.delete_thread", new=mock_delete),
patch(
"EvoScientist.sessions.find_similar_threads",
new=AsyncMock(return_value=[]),
),
patch(
"EvoScientist.sessions.list_threads",
new=AsyncMock(return_value=[]),
),
):
_run(DeleteCommand().execute(ctx, ["current"]))
assert mock_delete.await_count == 0
thread_store = FakeThreadStore(resolved_thread_id="current", deleted=True)
ctx, ui = _ctx(thread_id="current", thread_store=thread_store)
_run(DeleteCommand().execute(ctx, ["current"]))
assert ("delete_thread", "current") not in thread_store.calls
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Cannot delete the current session" in m for m in msgs)
def test_happy_path_success(self):
from EvoScientist.commands.implementation.session import DeleteCommand
ctx, ui = _ctx(thread_id="current")
with _patches(thread_exists=True, deleted=True):
_run(DeleteCommand().execute(ctx, ["other-thread"]))
ctx, ui = _ctx(
thread_id="current",
thread_store=FakeThreadStore(
resolved_thread_id="other-thread",
deleted=True,
),
)
_run(DeleteCommand().execute(ctx, ["other-thread"]))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Deleted session other-thread" in m for m in msgs)
@@ -89,26 +49,25 @@ class TestDeleteCommand:
from EvoScientist.commands.implementation.session import DeleteCommand
ctx, ui = _ctx()
with _patches(thread_exists=False, similar=[]):
_run(DeleteCommand().execute(ctx, ["missing"]))
_run(DeleteCommand().execute(ctx, ["missing"]))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("not found" in m for m in msgs)
def test_ambiguous_prefix(self):
from EvoScientist.commands.implementation.session import DeleteCommand
ctx, ui = _ctx()
with _patches(thread_exists=False, similar=["abc-one", "abc-two"]):
_run(DeleteCommand().execute(ctx, ["abc"]))
ctx, ui = _ctx(thread_store=FakeThreadStore(matches=["abc-one", "abc-two"]))
_run(DeleteCommand().execute(ctx, ["abc"]))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Ambiguous" in m for m in msgs)
def test_prefix_resolves_to_unique_match(self):
from EvoScientist.commands.implementation.session import DeleteCommand
ctx, ui = _ctx()
with _patches(thread_exists=False, similar=["abc-one"], deleted=True):
_run(DeleteCommand().execute(ctx, ["abc"]))
ctx, ui = _ctx(
thread_store=FakeThreadStore(resolved_thread_id="abc-one", deleted=True)
)
_run(DeleteCommand().execute(ctx, ["abc"]))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Deleted session abc-one" in m for m in msgs)
@@ -116,8 +75,7 @@ class TestDeleteCommand:
from EvoScientist.commands.implementation.session import DeleteCommand
ctx, ui = _ctx()
with _patches(threads=[]):
_run(DeleteCommand().execute(ctx, []))
_run(DeleteCommand().execute(ctx, []))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("No sessions to delete" in m for m in msgs)
@@ -136,6 +94,7 @@ class TestDeleteCommand:
"updated_at": None,
}
]
with _patches(threads=threads):
_run(DeleteCommand().execute(ctx, []))
store = FakeThreadStore(threads=threads)
ctx.graph_gateway = FakeGraphGateway(thread_store=store)
_run(DeleteCommand().execute(ctx, []))
ui.wait_for_thread_pick.assert_awaited_once()
+63 -67
View File
@@ -6,6 +6,7 @@ from unittest.mock import Mock, patch
import pytest
from EvoScientist.stream.display import _create_event_loop, _get_event_loop
from tests.fakes import FakeGraphGateway
class _TrackingEventLoopPolicy(asyncio.DefaultEventLoopPolicy):
@@ -128,7 +129,7 @@ class TestMultipleStreamingCalls:
# Mock agent that returns simple events
mock_agent = Mock()
async def mock_stream(*args, **kwargs):
async def mock_stream(_request):
"""Mock event stream."""
yield {"type": "text", "content": "test response"}
yield {"type": "done", "response": "test response"}
@@ -141,38 +142,39 @@ class TestMultipleStreamingCalls:
except RuntimeError:
pass
# Patch the stream_agent_events function
with patch(
"EvoScientist.stream.display.stream_agent_events", side_effect=mock_stream
):
# Patch Live to avoid terminal output during tests
with patch("EvoScientist.stream.display.Live"):
# First call
_run_streaming(
agent=mock_agent,
message="test message 1",
thread_id="thread1",
show_thinking=False,
interactive=True,
)
gateway = FakeGraphGateway(stream=mock_stream)
# Second call - this would fail with "Event loop is closed" before the fix
_run_streaming(
agent=mock_agent,
message="test message 2",
thread_id="thread1",
show_thinking=False,
interactive=True,
)
# Patch Live to avoid terminal output during tests
with patch("EvoScientist.stream.display.Live"):
# First call
_run_streaming(
agent=mock_agent,
message="test message 1",
thread_id="thread1",
show_thinking=False,
interactive=True,
gateway=gateway,
)
# Third call for good measure
_run_streaming(
agent=mock_agent,
message="test message 3",
thread_id="thread1",
show_thinking=False,
interactive=True,
)
# Second call - this would fail with "Event loop is closed" before the fix
_run_streaming(
agent=mock_agent,
message="test message 2",
thread_id="thread1",
show_thinking=False,
interactive=True,
gateway=gateway,
)
# Third call for good measure
_run_streaming(
agent=mock_agent,
message="test message 3",
thread_id="thread1",
show_thinking=False,
interactive=True,
gateway=gateway,
)
def test_loop_reused_across_calls(self):
"""Event loop should be reused across multiple calls."""
@@ -227,7 +229,7 @@ class TestMultipleStreamingCalls:
thinking = "Initial plan. " * 20
stream_calls = 0
async def mock_stream(*args, **kwargs):
async def mock_stream(_request):
nonlocal stream_calls
stream_calls += 1
if stream_calls == 1:
@@ -245,23 +247,20 @@ class TestMultipleStreamingCalls:
sent_thinking: list[str] = []
with patch(
"EvoScientist.stream.display.stream_agent_events",
side_effect=mock_stream,
):
with patch("EvoScientist.stream.display.Live"):
result = _run_streaming(
agent=mock_agent,
message="test message",
thread_id="thread1",
show_thinking=False,
interactive=True,
on_thinking=sent_thinking.append,
ask_user_prompt_fn=lambda _data: {
"answers": ["yes"],
"status": "answered",
},
)
with patch("EvoScientist.stream.display.Live"):
result = _run_streaming(
agent=mock_agent,
message="test message",
thread_id="thread1",
show_thinking=False,
interactive=True,
on_thinking=sent_thinking.append,
ask_user_prompt_fn=lambda _data: {
"answers": ["yes"],
"status": "answered",
},
gateway=FakeGraphGateway(stream=mock_stream),
)
assert result == "final answer"
assert sent_thinking == [thinking.rstrip()]
@@ -275,7 +274,7 @@ class TestMultipleStreamingCalls:
thinking_r2 = "Revised plan. " * 20
stream_calls = 0
async def mock_stream(*args, **kwargs):
async def mock_stream(_request):
nonlocal stream_calls
stream_calls += 1
if stream_calls == 1:
@@ -294,23 +293,20 @@ class TestMultipleStreamingCalls:
sent_thinking: list[str] = []
with patch(
"EvoScientist.stream.display.stream_agent_events",
side_effect=mock_stream,
):
with patch("EvoScientist.stream.display.Live"):
result = _run_streaming(
agent=mock_agent,
message="test message",
thread_id="thread1",
show_thinking=False,
interactive=True,
on_thinking=sent_thinking.append,
ask_user_prompt_fn=lambda _data: {
"answers": ["yes"],
"status": "answered",
},
)
with patch("EvoScientist.stream.display.Live"):
result = _run_streaming(
agent=mock_agent,
message="test message",
thread_id="thread1",
show_thinking=False,
interactive=True,
on_thinking=sent_thinking.append,
ask_user_prompt_fn=lambda _data: {
"answers": ["yes"],
"status": "answered",
},
gateway=FakeGraphGateway(stream=mock_stream),
)
assert result == "final answer"
assert sent_thinking == [thinking_r1.rstrip(), thinking_r2.rstrip()]
File diff suppressed because it is too large Load Diff
+4 -2
View File
@@ -1,6 +1,6 @@
"""Tests for the /new command."""
from unittest.mock import MagicMock
from unittest.mock import AsyncMock, MagicMock
from tests.conftest import run_async as _run
@@ -11,6 +11,7 @@ class TestNewCommand:
from EvoScientist.commands.implementation.session import NewCommand
ui = MagicMock()
ui.start_new_session = AsyncMock()
ctx = CommandContext(
agent=None,
thread_id="old-tid",
@@ -18,7 +19,7 @@ class TestNewCommand:
workspace_dir="/old/ws",
)
_run(NewCommand().execute(ctx, []))
ui.start_new_session.assert_called_once()
ui.start_new_session.assert_awaited_once()
def test_requires_agent_false(self):
from EvoScientist.commands.implementation.session import NewCommand
@@ -31,6 +32,7 @@ class TestNewCommand:
from EvoScientist.commands.implementation.session import NewCommand
ui = MagicMock()
ui.start_new_session = AsyncMock()
ctx = CommandContext(agent=None, thread_id="tid", ui=ui)
# No AttributeError even though ctx.agent is None
_run(NewCommand().execute(ctx, []))
+40 -8
View File
@@ -777,13 +777,20 @@ def test_subagent_summary_writer_uses_worker_metadata(tmp_path, monkeypatch):
assert _markdown_sections(body) == {"Summary": summary}
def test_memory_worker_run_kwargs_use_graph_id_and_source_metadata_only():
def test_memory_worker_run_kwargs_use_server_thread_id_and_source_metadata(monkeypatch):
monkeypatch.setattr(
memory_lifecycle,
"_worker_workspace_dir",
lambda _workspace_dir: "/tmp/ws",
)
trajectory: list[memory_lifecycle.CompactMessage] = [
{"role": "human", "content": "hi"}
]
kwargs = memory_lifecycle._memory_worker_run_kwargs(
role=memory_lifecycle.MemoryLifecycleRole.SUBAGENT,
thread_id="worker-thread",
workspace_dir="/active/workspace",
project_id="P-project",
source_agent="writing-agent",
session_id="thread-1",
@@ -792,20 +799,15 @@ def test_memory_worker_run_kwargs_use_graph_id_and_source_metadata_only():
assert kwargs["assistant_id"] == memory_lifecycle.SUBAGENT_MEMORY_WORKER_GRAPH_ID
assert kwargs["metadata"] == {
"agent_name": "EvoScientist",
"run_kind": "evomemory_subagent_worker",
"source_session_id": "thread-1",
"source_agent": "writing-agent",
"project_id": "P-project",
"trajectory_digest": memory_lifecycle._trajectory_digest(trajectory),
"workspace_dir": "/tmp/ws",
}
configurable = kwargs["config"]["configurable"]
assert configurable["thread_id"] == memory_lifecycle._worker_thread_id(
role=memory_lifecycle.MemoryLifecycleRole.SUBAGENT,
session_id="thread-1",
source_agent="writing-agent",
trajectory=trajectory,
)
assert configurable["thread_id"] == "worker-thread"
assert {
key: value
for key, value in configurable.items()
@@ -1285,6 +1287,7 @@ def test_memory_worker_skips_when_langgraph_dev_unavailable(tmp_path, monkeypatc
memory_lifecycle._launch_memory_worker(
role=memory_lifecycle.MemoryLifecycleRole.TURN,
memory_dir=tmp_path / "memories",
workspace_dir=tmp_path / "workspace",
project_id="P-project",
source_agent="EvoScientist",
session_id="thread-1",
@@ -1295,6 +1298,11 @@ def test_memory_worker_skips_when_langgraph_dev_unavailable(tmp_path, monkeypatc
def test_memory_worker_launch_marks_active_status(tmp_path, monkeypatch):
worker_activity.reset_memory_worker_status_for_tests()
monkeypatch.setattr(memory_lifecycle, "_memory_worker_url", lambda: "http://x")
monkeypatch.setattr(
memory_lifecycle,
"_worker_workspace_dir",
lambda _workspace_dir: "/tmp/ws",
)
monkeypatch.setattr(
"EvoScientist.langgraph_dev.manager.is_langgraph_dev_running",
lambda **_kwargs: True,
@@ -1320,6 +1328,7 @@ def test_memory_worker_launch_marks_active_status(tmp_path, monkeypatch):
memory_lifecycle._launch_memory_worker(
role=memory_lifecycle.MemoryLifecycleRole.TURN,
memory_dir=memory_dir,
workspace_dir=tmp_path / "workspace",
project_id="P-project",
source_agent="EvoScientist",
session_id="thread-1",
@@ -1328,6 +1337,23 @@ def test_memory_worker_launch_marks_active_status(tmp_path, monkeypatch):
try:
assert worker_activity.memory_worker_status().is_running is True
expected_metadata = {
"run_kind": "evomemory_turn_worker",
"source_session_id": "thread-1",
"source_agent": "EvoScientist",
"project_id": "P-project",
"trajectory_digest": memory_lifecycle._trajectory_digest(trajectory),
"workspace_dir": "/tmp/ws",
}
fake_client.threads.create.assert_called_once_with(
graph_id=memory_lifecycle.TURN_MEMORY_WORKER_GRAPH_ID,
metadata=expected_metadata,
)
fake_client.runs.create.assert_called_once()
run_kwargs = fake_client.runs.create.call_args.kwargs
assert run_kwargs["thread_id"] == "worker-thread"
assert run_kwargs["metadata"] == expected_metadata
assert run_kwargs["config"]["configurable"]["thread_id"] == "worker-thread"
assert spawned == [
{"url": "http://x", "thread_id": "worker-thread", "run_id": "run-1"}
]
@@ -1351,6 +1377,11 @@ def test_async_memory_worker_launch_offloads_blocking_work(
):
worker_activity.reset_memory_worker_status_for_tests()
monkeypatch.setattr(memory_lifecycle, "_memory_worker_url", lambda: "http://x")
monkeypatch.setattr(
memory_lifecycle,
"_worker_workspace_dir",
lambda _workspace_dir: "/tmp/ws",
)
call_threads: list[tuple[str, int]] = []
@@ -1394,6 +1425,7 @@ def test_async_memory_worker_launch_offloads_blocking_work(
await memory_lifecycle._alaunch_memory_worker(
role=memory_lifecycle.MemoryLifecycleRole.TURN,
memory_dir=tmp_path / "memories",
workspace_dir=tmp_path / "workspace",
project_id="P-project",
source_agent="EvoScientist",
session_id="thread-1",
+41 -63
View File
@@ -1,65 +1,39 @@
"""Tests for the /resume command."""
from contextlib import ExitStack
from unittest.mock import AsyncMock, MagicMock, patch
from unittest.mock import AsyncMock, MagicMock
from tests.conftest import run_async as _run
from tests.fakes import FakeGraphGateway, FakeThreadStore
def _ctx(thread_id="current", workspace_dir="/ws"):
def _ctx(thread_id="current", workspace_dir="/ws", thread_store=None):
from EvoScientist.commands.base import CommandContext
store = thread_store or FakeThreadStore()
ui = MagicMock()
ui.supports_interactive = True
ui.wait_for_thread_pick = AsyncMock()
ui.handle_session_resume = AsyncMock()
return CommandContext(
agent=None, thread_id=thread_id, ui=ui, workspace_dir=workspace_dir
agent=None,
thread_id=thread_id,
ui=ui,
workspace_dir=workspace_dir,
graph_gateway=FakeGraphGateway(thread_store=store),
), ui
def _patches(
*,
thread_exists=False,
similar=None,
threads=None,
metadata=None,
):
stack = ExitStack()
stack.enter_context(
patch(
"EvoScientist.sessions.thread_exists",
new=AsyncMock(return_value=thread_exists),
)
)
stack.enter_context(
patch(
"EvoScientist.sessions.find_similar_threads",
new=AsyncMock(return_value=similar or []),
)
)
stack.enter_context(
patch(
"EvoScientist.sessions.list_threads",
new=AsyncMock(return_value=threads or []),
)
)
stack.enter_context(
patch(
"EvoScientist.sessions.get_thread_metadata",
new=AsyncMock(return_value=metadata or {}),
)
)
return stack
class TestResumeCommand:
def test_with_arg_resolves_and_calls_ui(self):
from EvoScientist.commands.implementation.session import ResumeCommand
ctx, ui = _ctx()
with _patches(thread_exists=True, metadata={"workspace_dir": "/restored"}):
_run(ResumeCommand().execute(ctx, ["target-tid"]))
ctx, ui = _ctx(
thread_store=FakeThreadStore(
resolved_thread_id="target-tid",
metadata={"workspace_dir": "/restored"},
)
)
_run(ResumeCommand().execute(ctx, ["target-tid"]))
ui.handle_session_resume.assert_awaited_once_with("target-tid", "/restored")
# ctx mutations
assert ctx.thread_id == "target-tid"
@@ -69,8 +43,7 @@ class TestResumeCommand:
from EvoScientist.commands.implementation.session import ResumeCommand
ctx, ui = _ctx()
with _patches(threads=[]):
_run(ResumeCommand().execute(ctx, []))
_run(ResumeCommand().execute(ctx, []))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("No sessions to resume" in m for m in msgs)
ui.wait_for_thread_pick.assert_not_called()
@@ -82,8 +55,12 @@ class TestResumeCommand:
ctx, ui = _ctx()
ui.wait_for_thread_pick.return_value = "picked-tid"
threads = [{"thread_id": "picked-tid", "preview": "p", "message_count": 1}]
with _patches(thread_exists=True, threads=threads):
_run(ResumeCommand().execute(ctx, []))
store = FakeThreadStore(
threads=threads,
resolved_thread_id="picked-tid",
)
ctx.graph_gateway = FakeGraphGateway(thread_store=store)
_run(ResumeCommand().execute(ctx, []))
ui.wait_for_thread_pick.assert_awaited_once()
ui.handle_session_resume.assert_awaited_once()
@@ -93,16 +70,16 @@ class TestResumeCommand:
ctx, ui = _ctx()
ui.wait_for_thread_pick.return_value = None
threads = [{"thread_id": "t1", "preview": "", "message_count": 0}]
with _patches(threads=threads):
_run(ResumeCommand().execute(ctx, []))
store = FakeThreadStore(threads=threads)
ctx.graph_gateway = FakeGraphGateway(thread_store=store)
_run(ResumeCommand().execute(ctx, []))
ui.handle_session_resume.assert_not_called()
def test_ambiguous_prefix(self):
from EvoScientist.commands.implementation.session import ResumeCommand
ctx, ui = _ctx()
with _patches(thread_exists=False, similar=["abc-one", "abc-two"]):
_run(ResumeCommand().execute(ctx, ["abc"]))
ctx, ui = _ctx(thread_store=FakeThreadStore(matches=["abc-one", "abc-two"]))
_run(ResumeCommand().execute(ctx, ["abc"]))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Ambiguous" in m for m in msgs)
ui.handle_session_resume.assert_not_called()
@@ -111,8 +88,7 @@ class TestResumeCommand:
from EvoScientist.commands.implementation.session import ResumeCommand
ctx, ui = _ctx()
with _patches(thread_exists=False, similar=[]):
_run(ResumeCommand().execute(ctx, ["missing"]))
_run(ResumeCommand().execute(ctx, ["missing"]))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("not found" in m for m in msgs)
ui.handle_session_resume.assert_not_called()
@@ -120,22 +96,24 @@ class TestResumeCommand:
def test_prefix_resolves_to_unique_match(self):
from EvoScientist.commands.implementation.session import ResumeCommand
ctx, ui = _ctx()
with _patches(
thread_exists=False,
similar=["abc-one"],
metadata={"workspace_dir": "/ws1"},
):
_run(ResumeCommand().execute(ctx, ["abc"]))
ctx, ui = _ctx(
thread_store=FakeThreadStore(
resolved_thread_id="abc-one",
metadata={"workspace_dir": "/ws1"},
)
)
_run(ResumeCommand().execute(ctx, ["abc"]))
ui.handle_session_resume.assert_awaited_once_with("abc-one", "/ws1")
assert ctx.thread_id == "abc-one"
def test_empty_workspace_metadata_preserves_ctx_workspace(self):
from EvoScientist.commands.implementation.session import ResumeCommand
ctx, ui = _ctx(workspace_dir="/keep")
with _patches(thread_exists=True, metadata={}):
_run(ResumeCommand().execute(ctx, ["tid"]))
ctx, ui = _ctx(
workspace_dir="/keep",
thread_store=FakeThreadStore(resolved_thread_id="tid", metadata={}),
)
_run(ResumeCommand().execute(ctx, ["tid"]))
# ResumeCommand only overwrites ctx.workspace_dir if metadata has one
assert ctx.workspace_dir == "/keep"
# Callback still fires with the metadata value (empty string)
+7 -5
View File
@@ -283,14 +283,16 @@ class TestPhaseBMigrated:
"""Session lifecycle callbacks (start/resume) filled in Phase B."""
def test_start_new_session_fires_callback(self):
called = []
ui, _ = _make_ui(on_start_new_session=lambda: called.append("new"))
ui.start_new_session()
assert called == ["new"]
from unittest.mock import AsyncMock
cb = AsyncMock()
ui, _ = _make_ui(on_start_new_session=cb)
_run(ui.start_new_session())
cb.assert_awaited_once()
def test_start_new_session_without_callback_is_noop(self):
ui, console = _make_ui()
ui.start_new_session()
_run(ui.start_new_session())
console.print.assert_not_called()
def test_handle_session_resume_awaits_callback(self):
+242 -161
View File
@@ -6,50 +6,102 @@ subsequent messages, not silently keep the stale one the while-loop
captured at startup.
"""
from __future__ import annotations
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from langgraph.graph.state import CompiledStateGraph
from EvoScientist.cli.channel import (
ChannelMessage,
_register_channel_request,
)
from EvoScientist.cli.commands import (
ServeRuntimeState,
_make_serve_cmd_completed_hook,
_make_serve_handle_session_resume_cb,
_make_serve_start_new_session_cb,
_serve_process_message,
)
from EvoScientist.commands.base import ChannelRuntime
from EvoScientist.config import EvoScientistConfig
from EvoScientist.gateway import RuntimeGateways, ThreadStore
from tests.conftest import run_async as _run
from tests.fakes import FakeGraphGateway, FakeThreadStore
def test_hook_updates_holder_on_agent_swap():
def _agent(name: str = "agent") -> CompiledStateGraph:
return MagicMock(name=name, spec=CompiledStateGraph)
def _config() -> EvoScientistConfig:
return EvoScientistConfig()
def _thread_store(thread_id: str = "unused") -> ThreadStore:
return FakeThreadStore(generated_thread_id=thread_id)
def _runtime_gateways(thread_store: ThreadStore | None = None) -> RuntimeGateways:
store = thread_store or _thread_store()
return RuntimeGateways(
thread_store=store,
graph_gateway=FakeGraphGateway(thread_store=store),
)
def _runtime_state(
*,
agent: CompiledStateGraph | None = None,
thread_id: str = "tid",
workspace_dir: str | None = None,
config: EvoScientistConfig | None = None,
thread_store: ThreadStore | None = None,
runtime_gateways: RuntimeGateways | None = None,
) -> ServeRuntimeState:
store = thread_store or _thread_store()
return ServeRuntimeState(
agent=agent if agent is not None else _agent(),
thread_id=thread_id,
workspace_dir=workspace_dir,
config=config,
runtime_gateways=runtime_gateways or _runtime_gateways(store),
)
def test_hook_updates_runtime_state_on_agent_swap():
"""``/model`` mutates ``ctx.agent`` to a new handle — the hook must
push that handle into the shared holder so the outer poll loop sees
push that handle into the shared runtime state so the outer poll loop sees
it on the next message."""
holder = {"agent": "original-agent"}
hook = _make_serve_cmd_completed_hook(holder)
original_agent = _agent("original-agent")
new_agent = _agent("new-agent")
state = _runtime_state(agent=original_agent)
hook = _make_serve_cmd_completed_hook(state)
ctx = MagicMock()
ctx.agent = "new-agent"
ctx.agent = new_agent
ctx.thread_id = state.thread_id
cmd = MagicMock()
cmd.name = "/model"
_run(hook(ctx, "original-agent", cmd))
_run(hook(ctx, original_agent, cmd))
assert holder["agent"] == "new-agent"
assert state.agent is new_agent
def test_hook_syncs_channel_runtime():
"""Other readers (the bus) look at ``ChannelRuntime.agent``; the
hook keeps the runtime in sync with the holder update."""
holder = {"agent": "original-agent", "thread_id": "t"}
runtime = ChannelRuntime(agent="original-agent", thread_id="t")
hook = _make_serve_cmd_completed_hook(holder, runtime)
hook keeps the runtime in sync with the runtime state update."""
original_agent = _agent("original-agent")
new_agent = _agent("new-agent")
state = _runtime_state(agent=original_agent, thread_id="t")
runtime = ChannelRuntime(agent=original_agent, thread_id="t")
hook = _make_serve_cmd_completed_hook(state, runtime)
ctx = MagicMock()
ctx.agent = "new-agent"
ctx.agent = new_agent
# Pin ctx.thread_id explicitly — a bare MagicMock would let the
# hook's getattr fall through to a fresh MagicMock attribute and
# silently mutate runtime.thread_id, hiding regressions.
@@ -57,76 +109,83 @@ def test_hook_syncs_channel_runtime():
cmd = MagicMock()
cmd.name = "/model"
_run(hook(ctx, "original-agent", cmd))
_run(hook(ctx, original_agent, cmd))
assert runtime.agent == "new-agent"
assert runtime.agent is new_agent
assert runtime.thread_id == "t"
def test_hook_noop_when_agent_unchanged():
"""Commands like ``/evoskills`` don't touch ``ctx.agent`` — the
holder must stay put."""
holder = {"agent": "original-agent"}
hook = _make_serve_cmd_completed_hook(holder)
runtime state must stay put."""
original_agent = _agent("original-agent")
state = _runtime_state(agent=original_agent)
hook = _make_serve_cmd_completed_hook(state)
ctx = MagicMock()
ctx.agent = "original-agent" # no swap
ctx.agent = original_agent # no swap
ctx.thread_id = state.thread_id
cmd = MagicMock()
cmd.name = "/evoskills"
_run(hook(ctx, "original-agent", cmd))
_run(hook(ctx, original_agent, cmd))
assert holder["agent"] == "original-agent"
assert state.agent is original_agent
def test_hook_noop_when_ctx_agent_is_none():
"""Guard against commands that reset ``ctx.agent`` to ``None`` —
we never want to write ``None`` into the holder."""
holder = {"agent": "original-agent"}
hook = _make_serve_cmd_completed_hook(holder)
we never want to write ``None`` into runtime state."""
original_agent = _agent("original-agent")
state = _runtime_state(agent=original_agent)
hook = _make_serve_cmd_completed_hook(state)
ctx = MagicMock()
ctx.agent = None
ctx.thread_id = state.thread_id
cmd = MagicMock()
cmd.name = "/whatever"
_run(hook(ctx, "original-agent", cmd))
_run(hook(ctx, original_agent, cmd))
assert holder["agent"] == "original-agent"
assert state.agent is original_agent
def test_hook_updates_thread_id_on_resume():
"""``/resume`` mutates ``ctx.thread_id`` — the hook must push the
new id into the holder so the outer poll loop runs subsequent
new id into runtime state so the outer poll loop runs subsequent
messages on the resumed thread."""
holder = {"agent": "a", "thread_id": "original-tid"}
hook = _make_serve_cmd_completed_hook(holder)
agent = _agent("a")
state = _runtime_state(agent=agent, thread_id="original-tid")
hook = _make_serve_cmd_completed_hook(state)
ctx = MagicMock()
ctx.agent = "a" # no agent swap
ctx.agent = agent # no agent swap
ctx.thread_id = "new-tid"
ctx.workspace_dir = None
cmd = MagicMock()
cmd.name = "/resume"
_run(hook(ctx, "a", cmd))
_run(hook(ctx, agent, cmd))
assert holder["thread_id"] == "new-tid"
assert state.thread_id == "new-tid"
def test_hook_updates_workspace_dir_on_resume():
"""`/resume` can restore a different workspace; serve must reload for it."""
cfg = object()
holder = {
"agent": "old-agent",
"thread_id": "original-tid",
"workspace_dir": "/old-ws",
"config": cfg,
}
hook = _make_serve_cmd_completed_hook(holder, config=cfg)
cfg = _config()
old_agent = _agent("old-agent")
reloaded_agent = _agent("reloaded-agent")
state = _runtime_state(
agent=old_agent,
thread_id="original-tid",
workspace_dir="/old-ws",
config=cfg,
)
hook = _make_serve_cmd_completed_hook(state, config=cfg)
ctx = MagicMock()
ctx.agent = "old-agent"
ctx.agent = old_agent
ctx.thread_id = "new-tid"
ctx.workspace_dir = "/restored-ws"
cmd = MagicMock()
@@ -139,67 +198,70 @@ def test_hook_updates_workspace_dir_on_resume():
) as sync_server,
patch(
"EvoScientist.cli.commands._load_agent",
return_value="reloaded-agent",
return_value=reloaded_agent,
) as load_agent,
):
_run(hook(ctx, "old-agent", cmd))
_run(hook(ctx, old_agent, cmd))
sync_server.assert_awaited_once_with(cfg, workspace_dir="/restored-ws")
load_agent.assert_called_once_with(workspace_dir="/restored-ws", config=cfg)
assert holder["workspace_dir"] == "/restored-ws"
assert holder["agent"] == "reloaded-agent"
assert state.workspace_dir == "/restored-ws"
assert state.agent is reloaded_agent
def test_hook_syncs_channel_runtime_thread_id():
"""The bus reads ``ChannelRuntime.thread_id``; hook must sync it
alongside the holder update."""
holder = {"agent": "a", "thread_id": "original-tid"}
runtime = ChannelRuntime(agent="a", thread_id="original-tid")
hook = _make_serve_cmd_completed_hook(holder, runtime)
alongside the runtime state update."""
agent = _agent("a")
state = _runtime_state(agent=agent, thread_id="original-tid")
runtime = ChannelRuntime(agent=agent, thread_id="original-tid")
hook = _make_serve_cmd_completed_hook(state, runtime)
ctx = MagicMock()
ctx.agent = "a"
ctx.agent = agent
ctx.thread_id = "new-tid"
ctx.workspace_dir = None
cmd = MagicMock()
cmd.name = "/resume"
_run(hook(ctx, "a", cmd))
_run(hook(ctx, agent, cmd))
assert runtime.thread_id == "new-tid"
def test_hook_noop_when_thread_id_unchanged():
"""Most commands don't touch thread_id — holder stays put."""
holder = {"agent": "a", "thread_id": "same-tid"}
hook = _make_serve_cmd_completed_hook(holder)
"""Most commands don't touch thread_id — runtime state stays put."""
agent = _agent("a")
state = _runtime_state(agent=agent, thread_id="same-tid")
hook = _make_serve_cmd_completed_hook(state)
ctx = MagicMock()
ctx.agent = "a"
ctx.agent = agent
ctx.thread_id = "same-tid"
cmd = MagicMock()
cmd.name = "/evoskills"
_run(hook(ctx, "a", cmd))
_run(hook(ctx, agent, cmd))
assert holder["thread_id"] == "same-tid"
assert state.thread_id == "same-tid"
def test_hook_skips_resume_warning_when_thread_unchanged():
"""Bare ``/resume`` with no argument prints usage but leaves
``ctx.thread_id`` unchanged — the in-memory-state warning must NOT
fire because no resume actually happened."""
holder = {"agent": "a", "thread_id": "original-tid"}
hook = _make_serve_cmd_completed_hook(holder)
agent = _agent("a")
state = _runtime_state(agent=agent, thread_id="original-tid")
hook = _make_serve_cmd_completed_hook(state)
ctx = MagicMock()
ctx.agent = "a"
ctx.agent = agent
ctx.thread_id = "original-tid" # unchanged — bare /resume case
ctx.workspace_dir = None
cmd = MagicMock()
cmd.name = "/resume"
_run(hook(ctx, "a", cmd))
_run(hook(ctx, agent, cmd))
ctx.ui.append_system.assert_not_called()
ctx.ui.flush.assert_not_called()
@@ -208,19 +270,20 @@ def test_hook_skips_resume_warning_when_thread_unchanged():
def test_hook_emits_resume_warning_when_thread_changed():
"""``/resume <tid>`` that actually changes thread_id must surface
the in-memory-state warning via ``ctx.ui``."""
holder = {"agent": "a", "thread_id": "original-tid"}
hook = _make_serve_cmd_completed_hook(holder)
agent = _agent("a")
state = _runtime_state(agent=agent, thread_id="original-tid")
hook = _make_serve_cmd_completed_hook(state)
ctx = MagicMock()
# Mock out async flush so the test can synchronously run the hook.
ctx.ui.flush = AsyncMock()
ctx.agent = "a"
ctx.agent = agent
ctx.thread_id = "abc12345-resumed-tid"
ctx.workspace_dir = None
cmd = MagicMock()
cmd.name = "/resume"
_run(hook(ctx, "a", cmd))
_run(hook(ctx, agent, cmd))
ctx.ui.append_system.assert_called_once()
warn_text, warn_kwargs = (
@@ -235,51 +298,58 @@ def test_hook_emits_resume_warning_when_thread_changed():
def test_start_new_session_cb_rotates_thread_id():
"""``/new`` via channel calls this callback — must generate a new
thread id, push into holder, and sync the channel runtime."""
holder = {"agent": "a", "thread_id": "old-tid"}
runtime = ChannelRuntime(agent="a", thread_id="old-tid")
thread id, push into runtime state, and sync the channel runtime."""
agent = _agent("a")
state = _runtime_state(
agent=agent,
thread_id="old-tid",
thread_store=_thread_store("freshly-generated-tid"),
)
runtime = ChannelRuntime(agent=agent, thread_id="old-tid")
with patch(
"EvoScientist.sessions.generate_thread_id",
return_value="freshly-generated-tid",
):
cb = _make_serve_start_new_session_cb(holder, runtime)
cb()
cb = _make_serve_start_new_session_cb(
state,
runtime,
)
_run(cb())
assert holder["thread_id"] == "freshly-generated-tid"
assert state.thread_id == "freshly-generated-tid"
assert runtime.thread_id == "freshly-generated-tid"
def test_start_new_session_cb_leaves_agent_alone():
"""``/new`` rotates thread only — agent handle must stay put
(serve's agent is a single pre-loaded instance, not per-thread)."""
holder = {"agent": "a", "thread_id": "old-tid"}
agent = _agent("a")
state = _runtime_state(
agent=agent,
thread_id="old-tid",
thread_store=_thread_store("new-tid"),
)
with patch(
"EvoScientist.sessions.generate_thread_id",
return_value="new-tid",
):
cb = _make_serve_start_new_session_cb(holder)
cb()
cb = _make_serve_start_new_session_cb(state)
_run(cb())
assert holder["agent"] == "a"
assert state.agent is agent
def test_serve_resume_callback_syncs_reloads_and_adopts_workspace():
cfg = object()
holder = {
"agent": "old-agent",
"thread_id": "old-tid",
"workspace_dir": "/old-ws",
"config": cfg,
}
runtime = ChannelRuntime(agent="old-agent", thread_id="old-tid")
cb = _make_serve_handle_session_resume_cb(holder, runtime, config=cfg)
cfg = _config()
old_agent = _agent("old-agent")
reloaded_agent = _agent("reloaded-agent")
state = _runtime_state(
agent=old_agent,
thread_id="old-tid",
workspace_dir="/old-ws",
config=cfg,
)
runtime = ChannelRuntime(agent=old_agent, thread_id="old-tid")
cb = _make_serve_handle_session_resume_cb(state, runtime, config=cfg)
call_order: list[str] = []
def _load_agent(**_kwargs):
call_order.append("load")
return "reloaded-agent"
return reloaded_agent
async def _sync_server(*_args, **_kwargs):
call_order.append("sync")
@@ -299,23 +369,25 @@ def test_serve_resume_callback_syncs_reloads_and_adopts_workspace():
sync_server.assert_awaited_once_with(cfg, workspace_dir="/new-ws")
load_agent.assert_called_once_with(workspace_dir="/new-ws", config=cfg)
assert call_order == ["load", "sync"]
assert holder["thread_id"] == "new-tid"
assert holder["workspace_dir"] == "/new-ws"
assert holder["agent"] == "reloaded-agent"
assert state.thread_id == "new-tid"
assert state.workspace_dir == "/new-ws"
assert state.agent is reloaded_agent
assert runtime.thread_id == "new-tid"
assert runtime.agent == "reloaded-agent"
assert runtime.agent is reloaded_agent
def test_hook_emits_resume_warning_after_resume_callback_adopts_thread():
cfg = object()
holder = {
"agent": "old-agent",
"thread_id": "old-tid",
"workspace_dir": "/old-ws",
"config": cfg,
}
runtime = ChannelRuntime(agent="old-agent", thread_id="old-tid")
cb = _make_serve_handle_session_resume_cb(holder, runtime, config=cfg)
cfg = _config()
old_agent = _agent("old-agent")
reloaded_agent = _agent("reloaded-agent")
state = _runtime_state(
agent=old_agent,
thread_id="old-tid",
workspace_dir="/old-ws",
config=cfg,
)
runtime = ChannelRuntime(agent=old_agent, thread_id="old-tid")
cb = _make_serve_handle_session_resume_cb(state, runtime, config=cfg)
with (
patch(
@@ -324,21 +396,21 @@ def test_hook_emits_resume_warning_after_resume_callback_adopts_thread():
),
patch(
"EvoScientist.cli.commands._load_agent",
return_value="reloaded-agent",
return_value=reloaded_agent,
),
):
_run(cb("abc12345-resumed-tid", "/new-ws"))
hook = _make_serve_cmd_completed_hook(holder, runtime, config=cfg)
hook = _make_serve_cmd_completed_hook(state, runtime, config=cfg)
ctx = MagicMock()
ctx.ui.flush = AsyncMock()
ctx.agent = "reloaded-agent"
ctx.agent = reloaded_agent
ctx.thread_id = "abc12345-resumed-tid"
ctx.workspace_dir = "/new-ws"
cmd = MagicMock()
cmd.name = "/resume"
_run(hook(ctx, "reloaded-agent", cmd))
_run(hook(ctx, reloaded_agent, cmd))
ctx.ui.append_system.assert_called_once()
assert "in-memory state" in ctx.ui.append_system.call_args.args[0]
@@ -346,15 +418,17 @@ def test_hook_emits_resume_warning_after_resume_callback_adopts_thread():
def test_serve_resume_callback_preserves_state_when_sync_fails():
cfg = object()
holder = {
"agent": "old-agent",
"thread_id": "old-tid",
"workspace_dir": "/old-ws",
"config": cfg,
}
runtime = ChannelRuntime(agent="old-agent", thread_id="old-tid")
cb = _make_serve_handle_session_resume_cb(holder, runtime, config=cfg)
cfg = _config()
old_agent = _agent("old-agent")
loaded_but_not_adopted = _agent("loaded-but-not-adopted")
state = _runtime_state(
agent=old_agent,
thread_id="old-tid",
workspace_dir="/old-ws",
config=cfg,
)
runtime = ChannelRuntime(agent=old_agent, thread_id="old-tid")
cb = _make_serve_handle_session_resume_cb(state, runtime, config=cfg)
with (
patch(
@@ -363,7 +437,7 @@ def test_serve_resume_callback_preserves_state_when_sync_fails():
),
patch(
"EvoScientist.cli.commands._load_agent",
return_value="loaded-but-not-adopted",
return_value=loaded_but_not_adopted,
) as load_agent,
patch("EvoScientist.cli.commands.set_active_workspace") as set_active,
pytest.raises(RuntimeError, match="workspace conflict"),
@@ -372,28 +446,26 @@ def test_serve_resume_callback_preserves_state_when_sync_fails():
load_agent.assert_called_once_with(workspace_dir="/new-ws", config=cfg)
set_active.assert_called_once_with("/old-ws")
assert "loaded-but-not-adopted" not in holder.values()
assert "_resume_warning_thread_id" not in holder
assert holder == {
"agent": "old-agent",
"thread_id": "old-tid",
"workspace_dir": "/old-ws",
"config": cfg,
}
assert runtime.agent == "old-agent"
assert state.agent is old_agent
assert state.resume_warning_thread_id is None
assert state.thread_id == "old-tid"
assert state.workspace_dir == "/old-ws"
assert state.config is cfg
assert runtime.agent is old_agent
assert runtime.thread_id == "old-tid"
def test_serve_resume_callback_load_failure_does_not_sync_or_adopt():
cfg = object()
holder = {
"agent": "old-agent",
"thread_id": "old-tid",
"workspace_dir": "/old-ws",
"config": cfg,
}
runtime = ChannelRuntime(agent="old-agent", thread_id="old-tid")
cb = _make_serve_handle_session_resume_cb(holder, runtime, config=cfg)
cfg = _config()
old_agent = _agent("old-agent")
state = _runtime_state(
agent=old_agent,
thread_id="old-tid",
workspace_dir="/old-ws",
config=cfg,
)
runtime = ChannelRuntime(agent=old_agent, thread_id="old-tid")
cb = _make_serve_handle_session_resume_cb(state, runtime, config=cfg)
with (
patch(
@@ -412,32 +484,32 @@ def test_serve_resume_callback_load_failure_does_not_sync_or_adopt():
load_agent.assert_called_once_with(workspace_dir="/new-ws", config=cfg)
set_active.assert_called_once_with("/old-ws")
sync_server.assert_not_awaited()
assert "_resume_warning_thread_id" not in holder
assert holder == {
"agent": "old-agent",
"thread_id": "old-tid",
"workspace_dir": "/old-ws",
"config": cfg,
}
assert runtime.agent == "old-agent"
assert state.resume_warning_thread_id is None
assert state.agent is old_agent
assert state.thread_id == "old-tid"
assert state.workspace_dir == "/old-ws"
assert state.config is cfg
assert runtime.agent is old_agent
assert runtime.thread_id == "old-tid"
def test_hook_handles_both_agent_and_thread_swap():
"""Edge case: a command that changes both (hypothetical). Both
updates must land in the holder."""
holder = {"agent": "old-agent", "thread_id": "old-tid"}
hook = _make_serve_cmd_completed_hook(holder)
updates must land in runtime state."""
old_agent = _agent("old-agent")
new_agent = _agent("new-agent")
state = _runtime_state(agent=old_agent, thread_id="old-tid")
hook = _make_serve_cmd_completed_hook(state)
ctx = MagicMock()
ctx.agent = "new-agent"
ctx.agent = new_agent
ctx.thread_id = "new-tid"
cmd = MagicMock()
_run(hook(ctx, "old-agent", cmd))
_run(hook(ctx, old_agent, cmd))
assert holder["agent"] == "new-agent"
assert holder["thread_id"] == "new-tid"
assert state.agent is new_agent
assert state.thread_id == "new-tid"
def test_serve_process_message_reports_slash_dispatch_error_without_fallback():
@@ -456,7 +528,13 @@ def test_serve_process_message_reports_slash_dispatch_error_without_fallback():
chat_id="channel-user",
message_id="ts-1",
)
holder = {"agent": "agent", "thread_id": "tid"}
thread_store = _thread_store()
state = _runtime_state(
agent=_agent(),
thread_id="tid",
thread_store=thread_store,
runtime_gateways=_runtime_gateways(thread_store),
)
with (
patch(
@@ -469,7 +547,7 @@ def test_serve_process_message_reports_slash_dispatch_error_without_fallback():
_register_channel_request(msg)
_serve_process_message(
msg,
agent_holder=holder,
runtime_state=state,
model="model",
workspace_dir="/tmp",
show_thinking=False,
@@ -479,7 +557,7 @@ def test_serve_process_message_reports_slash_dispatch_error_without_fallback():
mock_run_streaming.assert_not_called()
def test_serve_process_message_uses_runtime_workspace_from_holder():
def test_serve_process_message_uses_runtime_workspace_from_state():
"""After `/resume`, serve should use the adopted workspace, not startup ws."""
msg = ChannelMessage(
msg_id="msg-2",
@@ -492,11 +570,14 @@ def test_serve_process_message_uses_runtime_workspace_from_holder():
chat_id="channel-user",
message_id="ts-2",
)
holder = {
"agent": "agent",
"thread_id": "tid",
"workspace_dir": "/restored-workspace",
}
thread_store = _thread_store()
state = _runtime_state(
agent=_agent(),
thread_id="tid",
workspace_dir="/restored-workspace",
thread_store=thread_store,
runtime_gateways=_runtime_gateways(thread_store),
)
captured: dict[str, str] = {}
async def _fake_dispatch(*args, **kwargs):
@@ -521,7 +602,7 @@ def test_serve_process_message_uses_runtime_workspace_from_holder():
_register_channel_request(msg)
_serve_process_message(
msg,
agent_holder=holder,
runtime_state=state,
model="model",
workspace_dir="/startup-workspace",
show_thinking=False,
+85 -30
View File
@@ -918,7 +918,7 @@ class TestPruningCheckpointer(unittest.TestCase):
async def _boom(*args, **kwargs):
raise RuntimeError("simulated prune failure")
wrapper._prune_after_put = _boom # type: ignore[assignment]
wrapper._prune_after_put = _boom
return await wrapper.aput(
self._config(tid),
self._checkpoint("cpf_0001", step=0),
@@ -1040,7 +1040,7 @@ class TestPruningCheckpointer(unittest.TestCase):
await release_prune.wait()
await orig_prune(thread_id, checkpoint_ns)
saver._prune_after_put = _gated_prune # type: ignore[method-assign]
saver._prune_after_put = _gated_prune
cfg_a = self._config(tid)
cfg_b = self._config(tid)
@@ -2149,10 +2149,8 @@ class TestCreateCheckpointerForLanggraphApi(unittest.TestCase):
"docstring in create_checkpointer_for_langgraph_api"
)
def test_aput_stamps_cli_metadata_for_main_graph_rows(self):
"""Main-graph (graph_id == AGENT_NAME) rows get agent_name /
workspace_dir / updated_at so they surface in CLI listings and
participate in pruning."""
def test_aput_stamps_workspace_metadata_for_graph_rows(self):
"""Graph rows get workspace metadata; only main rows get agent_name."""
import json
import aiosqlite
@@ -2212,7 +2210,8 @@ class TestCreateCheckpointerForLanggraphApi(unittest.TestCase):
assert main.get("workspace_dir") == "/tmp/test-workspace"
assert main.get("updated_at"), "updated_at drives /threads ordering"
assert "agent_name" not in worker
assert "workspace_dir" not in worker
assert worker.get("workspace_dir") == "/tmp/test-workspace"
assert worker.get("updated_at")
with tempfile.TemporaryDirectory() as td:
db = os.path.join(td, "sessions.db")
@@ -2250,6 +2249,7 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase):
assistant_id: str | None = "aaaa-bbbb",
graph_id: str | None = "EvoScientist",
workspace_dir: str | None = _WS,
model: str | None = "test-model",
agent_name: str | None = "EvoScientist",
ckpt_prefix: str = "ckpt",
) -> None:
@@ -2273,6 +2273,8 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase):
meta_dict["graph_id"] = graph_id
if workspace_dir is not None:
meta_dict["workspace_dir"] = workspace_dir
if model is not None:
meta_dict["model"] = model
meta = json.dumps(meta_dict)
con.execute(
"INSERT INTO checkpoints VALUES (?,?,?,?,?,?,?)",
@@ -2345,6 +2347,8 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase):
assert added[0]["metadata"].get("assistant_id") == asst_uuid_id
assert isinstance(added[0]["metadata"].get("assistant_id"), str)
assert added[0]["metadata"].get("graph_id") == "EvoScientist"
assert added[0]["metadata"].get("workspace_dir") == self._WS
assert added[0]["metadata"].get("model") == "test-model"
# created_at / updated_at must be datetime objects, not ISO strings.
# Threads.search() sorts by these fields using sorted(); mixing
# datetime and str raises TypeError: '<' not supported.
@@ -2413,14 +2417,18 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase):
assert t["metadata"].get("assistant_id") == asst_uuid_id
assert isinstance(t["metadata"].get("assistant_id"), str)
assert t["metadata"].get("graph_id") == "EvoScientist"
assert t["metadata"].get("workspace_dir") == self._WS
assert t["metadata"].get("model") == "test-model"
def test_restore_excludes_other_workspaces_and_internal_graphs(self):
"""The restore scope is graph_id==AGENT_NAME AND current workspace.
def test_restore_includes_current_workspace_graph_threads_only(self):
"""Restore includes current-workspace graph threads only.
Threads from other workspaces, internal worker/subagent graphs, and
pre-stamping rows without workspace_dir must NOT be resurrected —
sessions.db is machine-global and an unscoped restore would expose
them on the unauthenticated API (worst case --tunnel).
Threads from other workspaces and pre-stamping rows without
workspace_dir must NOT be resurrected — sessions.db is machine-global
and an unscoped restore would expose them on the unauthenticated API
(worst case --tunnel). Current-workspace async-subagent graph threads
are restored; memory-worker graph threads remain disposable until
worker cloning lands.
"""
import sys
import uuid as _uuid_mod
@@ -2463,11 +2471,24 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase):
_run(_restore_webui_threads_to_global_store())
added = mock_store["threads"]
assert len(added) == 1, (
f"Only the current-workspace main-graph thread may be restored, "
f"got {len(added)}: {[t['thread_id'] for t in added]}"
restored = {entry["thread_id"]: entry for entry in added}
assert set(restored) == {
_uuid_mod.UUID(mine),
_uuid_mod.UUID(subagent),
}
assert restored[_uuid_mod.UUID(mine)]["metadata"].get("graph_id") == (
"EvoScientist"
)
assert restored[_uuid_mod.UUID(mine)]["metadata"].get("workspace_dir") == (
self._WS
)
assert restored[_uuid_mod.UUID(mine)]["metadata"].get("model") == ("test-model")
assert restored[_uuid_mod.UUID(subagent)]["metadata"].get("graph_id") == (
"writing-agent"
)
assert restored[_uuid_mod.UUID(subagent)]["metadata"].get("workspace_dir") == (
self._WS
)
assert added[0]["thread_id"] == _uuid_mod.UUID(mine)
def test_purge_removes_only_evomemory_rows(self):
"""Startup purge drops evomemory-* residue, leaves everything else."""
@@ -2479,6 +2500,7 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase):
keep_main = "11111111-1111-1111-1111-111111111111"
keep_cli = "abcd1234"
drop_worker = "33333333-3333-3333-3333-333333333333"
keep_subagent = "44444444-4444-4444-4444-444444444444"
with tempfile.TemporaryDirectory() as td:
db = os.path.join(td, "sessions.db")
@@ -2486,6 +2508,7 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase):
self._make_db_with_threads(
db, [drop_worker], graph_id="evomemory-turn-worker"
)
self._make_db_with_threads(db, [keep_subagent], graph_id="writing-agent")
with patch(
"EvoScientist.sessions.get_db_path",
return_value=_mock_path(db),
@@ -2499,7 +2522,37 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase):
r[0] for r in con.execute("SELECT DISTINCT thread_id FROM checkpoints")
}
con.close()
assert remaining == {keep_main, keep_cli}
assert remaining == {keep_main, keep_cli, keep_subagent}
def test_cli_session_filters_exclude_non_main_graph_rows(self):
from unittest.mock import patch
from EvoScientist.sessions import (
list_threads,
resolve_thread_id_prefix,
thread_exists,
)
main_thread = "11111111-1111-1111-1111-111111111111"
worker_thread = "33333333-3333-3333-3333-333333333333"
with tempfile.TemporaryDirectory() as td:
db = os.path.join(td, "sessions.db")
self._make_db_with_threads(db, [main_thread])
self._make_db_with_threads(
db, [worker_thread], graph_id="evomemory-turn-worker"
)
with patch(
"EvoScientist.sessions.get_db_path",
return_value=_mock_path(db),
):
assert [row["thread_id"] for row in _run(list_threads())] == [
main_thread
]
assert _run(thread_exists(main_thread))
assert not _run(thread_exists(worker_thread))
assert _run(resolve_thread_id_prefix(worker_thread[:8])) == (None, [])
def test_restores_cli_rows_and_excludes_worker_residue(self):
"""CLI rows (agent_name, no graph_id) are restored with graph_id
@@ -2547,10 +2600,12 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase):
_run(_restore_webui_threads_to_global_store())
added = mock_store["threads"]
assert len(added) == 1, f"expected only the CLI thread, got {added}"
assert len(added) == 1
assert added[0]["thread_id"] == _uuid_mod.UUID(cli_thread)
# graph_id backfilled so Threads.State.get works on the stub.
# graph_id backfilled so Threads.State.get works on the CLI stub.
assert added[0]["metadata"].get("graph_id") == "EvoScientist"
assert added[0]["metadata"].get("workspace_dir") == self._WS
assert added[0]["metadata"].get("model") == "test-model"
def test_mixed_cli_webui_rows_keep_assistant_and_graph_id(self):
"""Interop thread (CLI rows + WebUI rows under one UUID): bare
@@ -2601,6 +2656,8 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase):
assert added[0]["thread_id"] == _uuid_mod.UUID(tid)
assert added[0]["metadata"].get("assistant_id") == asst
assert added[0]["metadata"].get("graph_id") == "EvoScientist"
assert added[0]["metadata"].get("workspace_dir") == self._WS
assert added[0]["metadata"].get("model") == "test-model"
def test_restored_stub_gets_title_from_first_human_message(self):
"""Stubs carry metadata.title derived from the thread's first human
@@ -2668,15 +2725,13 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase):
assert len(added) == 1, f"expected 1 restored thread, got {added}"
assert added[0]["metadata"].get("title") == "hello title test"
def test_removes_ghost_entries_absent_from_sqlite(self):
"""Stale .pckl UUID entries with no checkpoint rows are dropped.
def test_removes_preloaded_uuid_entries_outside_restore_scope(self):
"""Stale and out-of-scope .pckl UUID entries are dropped.
Ghost entries point at deleted/lost state and render as empty
sessions (the #277 symptom). Existence is checked against ALL UUID
threads in the DB, not the scoped restore set: a thread whose
checkpoints exist but fall outside the restore scope still opens
fine, so it must NOT be treated as a ghost. CLI-style non-UUID
entries are never touched.
Stale UUID entries point at deleted/lost state and render as empty
sessions (the #277 symptom). Out-of-scope UUID entries point at another
workspace's state and must not remain in this server's unauthenticated
thread registry. CLI-style non-UUID entries are never touched.
"""
import sys
import uuid as _uuid_mod
@@ -2724,8 +2779,8 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase):
assert _uuid_mod.UUID(ghost) not in ids, f"ghost must be removed, got {ids}"
assert ghost not in ids, f"ghost must be removed (str form), got {ids}"
assert "notauuid" in ids, "CLI-style entries must never be touched"
# Out-of-scope but existing in DB: kept (state still loads when opened).
assert _uuid_mod.UUID(out_of_scope) in ids
assert _uuid_mod.UUID(out_of_scope) not in ids
assert out_of_scope not in ids
# In-scope thread restored as usual.
assert _uuid_mod.UUID(in_scope) in ids
+6 -8
View File
@@ -21,6 +21,7 @@ from EvoScientist.cli.status_bar import (
status_style_name,
trim_status_text,
)
from tests.fakes import FakeGraphGateway, FakeThreadStore
def _render_fragments(fragments: list[tuple[str, str]]) -> str:
@@ -208,10 +209,6 @@ def test_build_status_text_uses_rich_styles():
def test_build_session_status_snapshot_uses_fallback_window(monkeypatch):
async def _fake_messages(thread_id: str):
assert thread_id == "thread-1"
return [HumanMessage(content="existing")]
class _FakeModel:
model_name: ClassVar[str] = "provider/demo-model"
profile: ClassVar[dict[str, object]] = {}
@@ -221,10 +218,6 @@ def test_build_session_status_snapshot_uses_fallback_window(monkeypatch):
assert messages[-1].content == "pending"
return 42_000
monkeypatch.setattr(
"EvoScientist.cli.status_bar.get_thread_messages",
_fake_messages,
)
monkeypatch.setattr(
"EvoScientist.cli.status_bar._get_default_chat_model",
lambda: _FakeModel(),
@@ -238,6 +231,11 @@ def test_build_session_status_snapshot_uses_fallback_window(monkeypatch):
build_session_status_snapshot(
"thread-1",
pending_user_text="pending",
graph_gateway=FakeGraphGateway(
thread_store=FakeThreadStore(
messages=[HumanMessage(content="existing")]
)
),
)
)
+44 -52
View File
@@ -2,11 +2,12 @@
from __future__ import annotations
from unittest.mock import MagicMock, patch
from unittest.mock import MagicMock
import pytest
from EvoScientist.stream import display as display_mod
from tests.fakes import FakeGraphGateway
@pytest.fixture(autouse=True)
@@ -38,7 +39,7 @@ def test_consume_breaks_on_cancel_event():
seen_events: list[int] = []
cancel_scope = "scope:consume"
async def _fake_stream(agent, message, thread_id, **kwargs):
async def _fake_stream(_request):
for i in range(100):
if i == 3:
# Set during iteration — next loop iter should bail.
@@ -46,18 +47,15 @@ def test_consume_breaks_on_cancel_event():
seen_events.append(i)
yield {"type": "text", "content": f"chunk-{i}"}
with patch(
"EvoScientist.stream.display.stream_agent_events",
new=_fake_stream,
):
result = display_mod._run_streaming(
agent=MagicMock(),
message="hello",
thread_id="t1",
show_thinking=False,
interactive=True,
cancel_scope=cancel_scope,
)
result = display_mod._run_streaming(
agent=MagicMock(),
message="hello",
thread_id="t1",
show_thinking=False,
interactive=True,
cancel_scope=cancel_scope,
gateway=FakeGraphGateway(stream=_fake_stream),
)
# We set the flag during event index 3; the cancel check runs at the
# top of the NEXT iteration (index 4), so indices 0-3 are pulled from
@@ -76,25 +74,22 @@ def test_run_streaming_short_circuits_when_scope_already_cancelled():
seen_event = False
cancel_scope = "scope:queued"
async def _fake_stream(agent, message, thread_id, **kwargs):
async def _fake_stream(_request):
nonlocal seen_event
seen_event = True
yield {"type": "text", "content": "ok"}
display_mod.request_stream_cancel(cancel_scope)
with patch(
"EvoScientist.stream.display.stream_agent_events",
new=_fake_stream,
):
result = display_mod._run_streaming(
agent=MagicMock(),
message="hello",
thread_id="t1",
show_thinking=False,
interactive=True,
cancel_scope=cancel_scope,
)
result = display_mod._run_streaming(
agent=MagicMock(),
message="hello",
thread_id="t1",
show_thinking=False,
interactive=True,
cancel_scope=cancel_scope,
gateway=FakeGraphGateway(stream=_fake_stream),
)
assert result == "[Stopped.]"
assert seen_event is False
@@ -105,21 +100,18 @@ def test_run_streaming_ignores_other_scope_cancel():
"""Cancelling one scope must not bleed into a different stream."""
display_mod.request_stream_cancel("scope:other")
async def _fake_stream(agent, message, thread_id, **kwargs):
async def _fake_stream(_request):
yield {"type": "text", "content": "ok"}
with patch(
"EvoScientist.stream.display.stream_agent_events",
new=_fake_stream,
):
result = display_mod._run_streaming(
agent=MagicMock(),
message="hello",
thread_id="t1",
show_thinking=False,
interactive=True,
cancel_scope="scope:self",
)
result = display_mod._run_streaming(
agent=MagicMock(),
message="hello",
thread_id="t1",
show_thinking=False,
interactive=True,
cancel_scope="scope:self",
gateway=FakeGraphGateway(stream=_fake_stream),
)
assert "[Stopped.]" not in result
@@ -132,7 +124,7 @@ def test_run_streaming_ignores_other_scope_cancel():
def test_run_streaming_pending_interrupt_short_circuits_on_cancel():
"""If cancel is already set, pending HITL prompt should not run."""
async def _empty_stream(agent, message, thread_id, **kwargs):
async def _empty_stream(_request):
if False:
yield {}
@@ -150,17 +142,17 @@ def test_run_streaming_pending_interrupt_short_circuits_on_cancel():
prompt_called = True
return None
with patch("EvoScientist.stream.display.stream_agent_events", new=_empty_stream):
result = display_mod._run_streaming(
agent=MagicMock(),
message="hello",
thread_id="t1",
show_thinking=False,
interactive=True,
hitl_prompt_fn=_prompt,
cancel_scope="scope:hitl",
_state=state,
)
result = display_mod._run_streaming(
agent=MagicMock(),
message="hello",
thread_id="t1",
show_thinking=False,
interactive=True,
hitl_prompt_fn=_prompt,
cancel_scope="scope:hitl",
_state=state,
gateway=FakeGraphGateway(stream=_empty_stream),
)
assert result == "Partial answer\n[Stopped.]"
assert prompt_called is False
+52 -5
View File
@@ -360,6 +360,41 @@ class TestV3ProtocolStreaming:
summary_events = [e for e in events if e.get("type") == "summarization"]
assert summary_events == []
def test_direct_stream_loads_existing_summarization_event_when_omitted(self):
"""Public stream_agent_events() suppresses persisted summary replays."""
summary_message = HumanMessage(
content="Here is a summary of the conversation to date:\n\nKey facts",
)
summary_event = {
"_summarization_event": {
"summary_message": summary_message,
"cutoff_index": 12,
"file_path": None,
}
}
agent = FakeV3Agent(
[
protocol_event("updates", summary_event),
message_delta("real content"),
],
state_values=summary_event,
)
async def _collect():
events = []
async for event in stream_agent_events(agent, "hi", "t1"):
events.append(event)
return events
events = run_async(_collect())
summary_start_events = [
e for e in events if e.get("type") == "summarization_start"
]
assert summary_start_events == []
summary_events = [e for e in events if e.get("type") == "summarization"]
assert summary_events == []
def test_whole_message_reasoning_is_not_duplicated(self):
"""Providers can expose the same reasoning in kwargs and content blocks."""
message = AIMessage(
@@ -432,7 +467,9 @@ class TestV3ProtocolStreaming:
return [
event
async for event in stream_agent_events(
agent, "run probe", "live-deepagents-tool-id"
agent,
"run probe",
"live-deepagents-tool-id",
)
]
@@ -493,7 +530,9 @@ class TestV3ProtocolStreaming:
return [
event
async for event in stream_agent_events(
agent, "run echo", "live-deepagents-hitl"
agent,
"run echo",
"live-deepagents-hitl",
)
]
@@ -554,7 +593,9 @@ class TestV3ProtocolStreaming:
return [
event
async for event in stream_agent_events(
agent, message, "live-deepagents-ask-user"
agent,
message,
"live-deepagents-ask-user",
)
]
@@ -625,7 +666,9 @@ class TestV3ProtocolStreaming:
return [
event
async for event in stream_agent_events(
agent, "delegate", "live-deepagents-subagent"
agent,
"delegate",
"live-deepagents-subagent",
)
]
@@ -983,7 +1026,11 @@ class TestV3ProtocolStreaming:
async def consume_one_and_close():
agent = HangingV3Agent([message_delta("hi")])
stream = stream_agent_events(agent, "hi", "t1")
stream = stream_agent_events(
agent,
"hi",
"t1",
)
first = await stream.__anext__()
await stream.aclose()
return first, agent.aborted
+173 -241
View File
@@ -10,16 +10,16 @@ Covers:
from __future__ import annotations
import asyncio
from dataclasses import dataclass
from unittest.mock import AsyncMock, MagicMock, patch
from unittest.mock import AsyncMock, MagicMock
from EvoScientist.channels.base import Channel
from EvoScientist.channels.bus.events import InboundMessage as BusInbound
from EvoScientist.channels.bus.message_bus import MessageBus
from EvoScientist.channels.channel_manager import ChannelManager
from EvoScientist.channels.consumer import InboundConsumer, _join_subagent_text
from EvoScientist.stream.emitter import StreamEvent, StreamEventEmitter
from tests.conftest import run_async as _run
from tests.fakes import FakeGraphGateway
from tests.fakes import StubChannel as _StubChannel
from tests.stream_v3_fakes import (
FakeSubagent,
FakeV3Agent,
@@ -32,16 +32,6 @@ from tests.stream_v3_fakes import (
# ═══════════════════════════════════════════════════════════════════
@dataclass
class _FakeConfig:
text_chunk_limit: int = 4096
allowed_senders: list | None = None
allowed_channels: list | None = None
proxy: str | None = None
require_mention: str = "group"
dm_policy: str = "allowlist"
# ═══════════════════════════════════════════════════════════════════
# 1. StreamEventEmitter.subagent_text
# ═══════════════════════════════════════════════════════════════════
@@ -192,24 +182,6 @@ class TestStreamAgentEventsSubagentText:
# ═══════════════════════════════════════════════════════════════════
class _StubChannel(Channel):
"""Minimal concrete channel for consumer tests."""
name = "stub"
def __init__(self, config=None):
super().__init__(config or _FakeConfig())
async def start(self):
self._running = True
async def _send_chunk(self, chat_id, formatted, raw, reply_to, metadata):
pass
async def _send_typing_action(self, chat_id):
pass
def _make_consumer(stream_events: list[dict], **kw):
"""Create an InboundConsumer whose agent streams the given event dicts.
@@ -220,24 +192,20 @@ def _make_consumer(stream_events: list[dict], **kw):
mgr = ChannelManager(bus)
mgr.register(_StubChannel())
# Patch stream_agent_events to yield pre-built events
async def _fake_stream(agent, message, thread_id, **kwargs):
for ev in stream_events:
yield ev
agent = MagicMock()
consumer = InboundConsumer(
bus=bus,
manager=mgr,
agent=agent,
thread_id="",
graph_gateway=FakeGraphGateway(events=stream_events),
max_concurrent=2,
max_pending=10,
inference_timeout=5.0,
drain_timeout=1.0,
**kw,
)
return consumer, bus, _fake_stream
return consumer, bus
class TestConsumerSubagentTextFallback:
@@ -260,31 +228,25 @@ class TestConsumerSubagentTextFallback:
},
{"type": "done", "content": ""},
]
consumer, bus, fake_stream = _make_consumer(events)
consumer, bus = _make_consumer(events)
async def _test():
with patch(
"EvoScientist.stream.events.stream_agent_events",
new=fake_stream,
):
msg = BusInbound(
channel="stub",
sender_id="u1",
chat_id="c1",
content="analyze papers",
)
await bus.publish_inbound(msg)
msg = BusInbound(
channel="stub",
sender_id="u1",
chat_id="c1",
content="analyze papers",
)
await bus.publish_inbound(msg)
task = asyncio.create_task(consumer.run())
outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0)
task = asyncio.create_task(consumer.run())
outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0)
assert (
outbound.content == "Found 3 relevant papers. Key insight: X is Y."
)
assert outbound.channel == "stub"
assert outbound.content == "Found 3 relevant papers. Key insight: X is Y."
assert outbound.channel == "stub"
await consumer.stop()
await task
await consumer.stop()
await task
_run(_test())
@@ -300,28 +262,24 @@ class TestConsumerSubagentTextFallback:
{"type": "text", "content": "Here is my summary."},
{"type": "done", "content": ""},
]
consumer, bus, fake_stream = _make_consumer(events)
consumer, bus = _make_consumer(events)
async def _test():
with patch(
"EvoScientist.stream.events.stream_agent_events",
new=fake_stream,
):
msg = BusInbound(
channel="stub",
sender_id="u1",
chat_id="c1",
content="test",
)
await bus.publish_inbound(msg)
msg = BusInbound(
channel="stub",
sender_id="u1",
chat_id="c1",
content="test",
)
await bus.publish_inbound(msg)
task = asyncio.create_task(consumer.run())
outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0)
task = asyncio.create_task(consumer.run())
outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0)
assert outbound.content == "Here is my summary."
assert outbound.content == "Here is my summary."
await consumer.stop()
await task
await consumer.stop()
await task
_run(_test())
@@ -335,25 +293,10 @@ class TestConsumerSubagentTextFallback:
assert channel is not None
channel.send_thinking_message = AsyncMock()
consumer = InboundConsumer(
bus=bus,
manager=mgr,
agent=MagicMock(),
thread_id="",
max_concurrent=2,
max_pending=10,
inference_timeout=5.0,
drain_timeout=1.0,
send_thinking=True,
)
consumer._resolve_ask_user = AsyncMock( # type: ignore[method-assign]
return_value={"answers": ["yes"], "status": "answered"}
)
thinking = "Initial plan. " * 20
stream_calls = 0
async def _fake_stream(agent, message, thread_id, **kwargs):
async def _fake_stream(_request):
nonlocal stream_calls
stream_calls += 1
if stream_calls == 1:
@@ -370,30 +313,42 @@ class TestConsumerSubagentTextFallback:
yield {"type": "text", "content": "final answer"}
yield {"type": "done", "content": "final answer"}
consumer = InboundConsumer(
bus=bus,
manager=mgr,
agent=MagicMock(),
thread_id="",
graph_gateway=FakeGraphGateway(stream=_fake_stream),
max_concurrent=2,
max_pending=10,
inference_timeout=5.0,
drain_timeout=1.0,
send_thinking=True,
)
consumer._resolve_ask_user = AsyncMock( # type: ignore[method-assign]
return_value={"answers": ["yes"], "status": "answered"}
)
async def _test():
with patch(
"EvoScientist.stream.events.stream_agent_events",
new=_fake_stream,
):
await bus.publish_inbound(
BusInbound(
channel="stub",
sender_id="u1",
chat_id="c1",
content="analyze papers",
)
await bus.publish_inbound(
BusInbound(
channel="stub",
sender_id="u1",
chat_id="c1",
content="analyze papers",
)
)
task = asyncio.create_task(consumer.run())
outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0)
task = asyncio.create_task(consumer.run())
outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0)
assert outbound.content == "final answer"
assert channel.send_thinking_message.await_count == 1
call = channel.send_thinking_message.await_args_list[0]
assert call.args[1] == thinking.rstrip()
assert outbound.content == "final answer"
assert channel.send_thinking_message.await_count == 1
call = channel.send_thinking_message.await_args_list[0]
assert call.args[1] == thinking.rstrip()
await consumer.stop()
await task
await consumer.stop()
await task
_run(_test())
@@ -407,26 +362,11 @@ class TestConsumerSubagentTextFallback:
assert channel is not None
channel.send_thinking_message = AsyncMock()
consumer = InboundConsumer(
bus=bus,
manager=mgr,
agent=MagicMock(),
thread_id="",
max_concurrent=2,
max_pending=10,
inference_timeout=5.0,
drain_timeout=1.0,
send_thinking=True,
)
consumer._resolve_ask_user = AsyncMock( # type: ignore[method-assign]
return_value={"answers": ["yes"], "status": "answered"}
)
thinking_r1 = "Initial plan. " * 20
thinking_r2 = "Revised plan. " * 20
stream_calls = 0
async def _fake_stream(agent, message, thread_id, **kwargs):
async def _fake_stream(_request):
nonlocal stream_calls
stream_calls += 1
if stream_calls == 1:
@@ -443,32 +383,44 @@ class TestConsumerSubagentTextFallback:
yield {"type": "text", "content": "final answer"}
yield {"type": "done", "content": "final answer"}
consumer = InboundConsumer(
bus=bus,
manager=mgr,
agent=MagicMock(),
thread_id="",
graph_gateway=FakeGraphGateway(stream=_fake_stream),
max_concurrent=2,
max_pending=10,
inference_timeout=5.0,
drain_timeout=1.0,
send_thinking=True,
)
consumer._resolve_ask_user = AsyncMock( # type: ignore[method-assign]
return_value={"answers": ["yes"], "status": "answered"}
)
async def _test():
with patch(
"EvoScientist.stream.events.stream_agent_events",
new=_fake_stream,
):
await bus.publish_inbound(
BusInbound(
channel="stub",
sender_id="u1",
chat_id="c1",
content="analyze papers",
)
await bus.publish_inbound(
BusInbound(
channel="stub",
sender_id="u1",
chat_id="c1",
content="analyze papers",
)
)
task = asyncio.create_task(consumer.run())
outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0)
task = asyncio.create_task(consumer.run())
outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0)
assert outbound.content == "final answer"
assert channel.send_thinking_message.await_count == 2
call1 = channel.send_thinking_message.await_args_list[0]
call2 = channel.send_thinking_message.await_args_list[1]
assert call1.args[1] == thinking_r1.rstrip()
assert call2.args[1] == thinking_r2.rstrip()
assert outbound.content == "final answer"
assert channel.send_thinking_message.await_count == 2
call1 = channel.send_thinking_message.await_args_list[0]
call2 = channel.send_thinking_message.await_args_list[1]
assert call1.args[1] == thinking_r1.rstrip()
assert call2.args[1] == thinking_r2.rstrip()
await consumer.stop()
await task
await consumer.stop()
await task
_run(_test())
@@ -477,28 +429,24 @@ class TestConsumerSubagentTextFallback:
events = [
{"type": "done", "content": ""},
]
consumer, bus, fake_stream = _make_consumer(events)
consumer, bus = _make_consumer(events)
async def _test():
with patch(
"EvoScientist.stream.events.stream_agent_events",
new=fake_stream,
):
msg = BusInbound(
channel="stub",
sender_id="u1",
chat_id="c1",
content="test",
)
await bus.publish_inbound(msg)
msg = BusInbound(
channel="stub",
sender_id="u1",
chat_id="c1",
content="test",
)
await bus.publish_inbound(msg)
task = asyncio.create_task(consumer.run())
outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0)
task = asyncio.create_task(consumer.run())
outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0)
assert outbound.content == "No response"
assert outbound.content == "No response"
await consumer.stop()
await task
await consumer.stop()
await task
_run(_test())
@@ -513,28 +461,24 @@ class TestConsumerSubagentTextFallback:
},
{"type": "done", "content": "Final summary from done event."},
]
consumer, bus, fake_stream = _make_consumer(events)
consumer, bus = _make_consumer(events)
async def _test():
with patch(
"EvoScientist.stream.events.stream_agent_events",
new=fake_stream,
):
msg = BusInbound(
channel="stub",
sender_id="u1",
chat_id="c1",
content="test",
)
await bus.publish_inbound(msg)
msg = BusInbound(
channel="stub",
sender_id="u1",
chat_id="c1",
content="test",
)
await bus.publish_inbound(msg)
task = asyncio.create_task(consumer.run())
outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0)
task = asyncio.create_task(consumer.run())
outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0)
assert outbound.content == "Final summary from done event."
assert outbound.content == "Final summary from done event."
await consumer.stop()
await task
await consumer.stop()
await task
_run(_test())
@@ -639,29 +583,25 @@ class TestConsumerParallelSubagentFallback:
},
{"type": "done", "content": ""},
]
consumer, bus, fake_stream = _make_consumer(events)
consumer, bus = _make_consumer(events)
async def _test():
with patch(
"EvoScientist.stream.events.stream_agent_events",
new=fake_stream,
):
msg = BusInbound(
channel="stub",
sender_id="u1",
chat_id="c1",
content="test",
)
await bus.publish_inbound(msg)
msg = BusInbound(
channel="stub",
sender_id="u1",
chat_id="c1",
content="test",
)
await bus.publish_inbound(msg)
task = asyncio.create_task(consumer.run())
outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0)
task = asyncio.create_task(consumer.run())
outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0)
assert "[research]: Found papers. Key insight." in outbound.content
assert "[analysis]: Metric is high." in outbound.content
assert "[research]: Found papers. Key insight." in outbound.content
assert "[analysis]: Metric is high." in outbound.content
await consumer.stop()
await task
await consumer.stop()
await task
_run(_test())
@@ -676,29 +616,25 @@ class TestConsumerParallelSubagentFallback:
},
{"type": "done", "content": ""},
]
consumer, bus, fake_stream = _make_consumer(events)
consumer, bus = _make_consumer(events)
async def _test():
with patch(
"EvoScientist.stream.events.stream_agent_events",
new=fake_stream,
):
msg = BusInbound(
channel="stub",
sender_id="u1",
chat_id="c1",
content="test",
)
await bus.publish_inbound(msg)
msg = BusInbound(
channel="stub",
sender_id="u1",
chat_id="c1",
content="test",
)
await bus.publish_inbound(msg)
task = asyncio.create_task(consumer.run())
outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0)
task = asyncio.create_task(consumer.run())
outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0)
assert outbound.content == "Only agent."
assert "[research]" not in outbound.content
assert outbound.content == "Only agent."
assert "[research]" not in outbound.content
await consumer.stop()
await task
await consumer.stop()
await task
_run(_test())
@@ -740,36 +676,32 @@ class TestConsumerSameNameInterleaved:
},
{"type": "done", "content": ""},
]
consumer, bus, fake_stream = _make_consumer(events)
consumer, bus = _make_consumer(events)
async def _test():
with patch(
"EvoScientist.stream.events.stream_agent_events",
new=fake_stream,
):
msg = BusInbound(
channel="stub",
sender_id="u1",
chat_id="c1",
content="test",
)
await bus.publish_inbound(msg)
msg = BusInbound(
channel="stub",
sender_id="u1",
chat_id="c1",
content="test",
)
await bus.publish_inbound(msg)
task = asyncio.create_task(consumer.run())
outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0)
task = asyncio.create_task(consumer.run())
outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0)
# Fixed: instances are now properly separated with numbered labels
assert (
"[research-agent #1]: Instance-1 sentence A. Instance-1 sentence B."
in outbound.content
)
assert (
"[research-agent #2]: Instance-2 sentence X. Instance-2 sentence Y."
in outbound.content
)
# Fixed: instances are now properly separated with numbered labels
assert (
"[research-agent #1]: Instance-1 sentence A. Instance-1 sentence B."
in outbound.content
)
assert (
"[research-agent #2]: Instance-2 sentence X. Instance-2 sentence Y."
in outbound.content
)
await consumer.stop()
await task
await consumer.stop()
await task
_run(_test())
+14 -21
View File
@@ -1,10 +1,11 @@
"""Tests for the /threads command."""
from unittest.mock import AsyncMock, MagicMock, patch
from unittest.mock import MagicMock
from rich.table import Table
from tests.conftest import run_async as _run
from tests.fakes import FakeGraphGateway, FakeThreadStore
def _ctx(**overrides):
@@ -12,11 +13,13 @@ def _ctx(**overrides):
ui = MagicMock()
ui.supports_interactive = overrides.pop("supports_interactive", True)
store = overrides.pop("thread_store", FakeThreadStore())
return CommandContext(
agent=None,
thread_id=overrides.pop("thread_id", "tid-1"),
ui=ui,
workspace_dir=overrides.pop("workspace_dir", "/ws"),
graph_gateway=FakeGraphGateway(thread_store=store),
), ui
@@ -25,11 +28,7 @@ class TestThreadsCommand:
from EvoScientist.commands.implementation.session import ThreadsCommand
ctx, ui = _ctx()
with patch(
"EvoScientist.sessions.list_threads",
new=AsyncMock(return_value=[]),
):
_run(ThreadsCommand().execute(ctx, []))
_run(ThreadsCommand().execute(ctx, []))
ui.append_system.assert_called_once()
assert "No saved sessions" in ui.append_system.call_args.args[0]
@@ -53,11 +52,9 @@ class TestThreadsCommand:
"updated_at": None,
},
]
with patch(
"EvoScientist.sessions.list_threads",
new=AsyncMock(return_value=threads),
):
_run(ThreadsCommand().execute(ctx, []))
store = FakeThreadStore(threads=threads)
ctx.graph_gateway = FakeGraphGateway(thread_store=store)
_run(ThreadsCommand().execute(ctx, []))
ui.mount_renderable.assert_called_once()
table = ui.mount_renderable.call_args.args[0]
assert isinstance(table, Table)
@@ -81,11 +78,9 @@ class TestThreadsCommand:
"updated_at": None,
}
]
with patch(
"EvoScientist.sessions.list_threads",
new=AsyncMock(return_value=threads),
):
_run(ThreadsCommand().execute(ctx, []))
store = FakeThreadStore(threads=threads)
ctx.graph_gateway = FakeGraphGateway(thread_store=store)
_run(ThreadsCommand().execute(ctx, []))
ui.append_system.assert_not_called()
def test_channel_mode_drops_model_column(self):
@@ -102,11 +97,9 @@ class TestThreadsCommand:
"updated_at": None,
}
]
with patch(
"EvoScientist.sessions.list_threads",
new=AsyncMock(return_value=threads),
):
_run(ThreadsCommand().execute(ctx, []))
store = FakeThreadStore(threads=threads)
ctx.graph_gateway = FakeGraphGateway(thread_store=store)
_run(ThreadsCommand().execute(ctx, []))
# Channel mode: no Model column. 4 columns: ID, Preview, Msgs, Last Used.
table = ui.mount_renderable.call_args.args[0]
column_headers = [col.header for col in table.columns]
+2
View File
@@ -7,6 +7,7 @@ from EvoScientist.cli.tui_runtime import (
resolve_ui_backend,
run_streaming,
)
from tests.fakes import FakeGraphGateway
def test_normalize_ui_backend_defaults_to_cli():
@@ -73,5 +74,6 @@ def test_run_streaming_falls_back_to_cli_on_runtime_error(monkeypatch):
thread_id="t1",
show_thinking=False,
interactive=True,
gateway=FakeGraphGateway(),
)
assert result == "fallback-ok"