diff --git a/EvoScientist/EvoScientist.py b/EvoScientist/EvoScientist.py index d8402eb..9bb7059 100644 --- a/EvoScientist/EvoScientist.py +++ b/EvoScientist/EvoScientist.py @@ -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 diff --git a/EvoScientist/channels/consumer.py b/EvoScientist/channels/consumer.py index 9b83d0d..f411552 100644 --- a/EvoScientist/channels/consumer.py +++ b/EvoScientist/channels/consumer.py @@ -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)]} ) diff --git a/EvoScientist/channels/standalone.py b/EvoScientist/channels/standalone.py index cdd8d10..902bb69 100644 --- a/EvoScientist/channels/standalone.py +++ b/EvoScientist/channels/standalone.py @@ -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) diff --git a/EvoScientist/cli/_agent_loader.py b/EvoScientist/cli/_agent_loader.py index e7534f1..d7256df 100644 --- a/EvoScientist/cli/_agent_loader.py +++ b/EvoScientist/cli/_agent_loader.py @@ -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(): diff --git a/EvoScientist/cli/agent.py b/EvoScientist/cli/agent.py index 2e8caa8..ce12fd4 100644 --- a/EvoScientist/cli/agent.py +++ b/EvoScientist/cli/agent.py @@ -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: diff --git a/EvoScientist/cli/async_notifier.py b/EvoScientist/cli/async_notifier.py index d6cebff..5389c99 100644 --- a/EvoScientist/cli/async_notifier.py +++ b/EvoScientist/cli/async_notifier.py @@ -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. diff --git a/EvoScientist/cli/channel.py b/EvoScientist/cli/channel.py index e904476..67ae4fc 100644 --- a/EvoScientist/cli/channel.py +++ b/EvoScientist/cli/channel.py @@ -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, diff --git a/EvoScientist/cli/commands.py b/EvoScientist/cli/commands.py index 3dae5c8..a2e846a 100644 --- a/EvoScientist/cli/commands.py +++ b/EvoScientist/cli/commands.py @@ -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: diff --git a/EvoScientist/cli/interactive.py b/EvoScientist/cli/interactive.py index 61e1853..ae0a4d5 100644 --- a/EvoScientist/cli/interactive.py +++ b/EvoScientist/cli/interactive.py @@ -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: diff --git a/EvoScientist/cli/rich_command_ui.py b/EvoScientist/cli/rich_command_ui.py index e710665..26c5056 100644 --- a/EvoScientist/cli/rich_command_ui.py +++ b/EvoScientist/cli/rich_command_ui.py @@ -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 diff --git a/EvoScientist/cli/status_bar.py b/EvoScientist/cli/status_bar.py index 88d148b..1f2f54a 100644 --- a/EvoScientist/cli/status_bar.py +++ b/EvoScientist/cli/status_bar.py @@ -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: diff --git a/EvoScientist/cli/tui_backends.py b/EvoScientist/cli/tui_backends.py index a28178d..e196a69 100644 --- a/EvoScientist/cli/tui_backends.py +++ b/EvoScientist/cli/tui_backends.py @@ -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, ) diff --git a/EvoScientist/cli/tui_interactive.py b/EvoScientist/cli/tui_interactive.py index 64b9e5c..86fd88f 100644 --- a/EvoScientist/cli/tui_interactive.py +++ b/EvoScientist/cli/tui_interactive.py @@ -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( diff --git a/EvoScientist/cli/tui_runtime.py b/EvoScientist/cli/tui_runtime.py index 8a95d05..f028506 100644 --- a/EvoScientist/cli/tui_runtime.py +++ b/EvoScientist/cli/tui_runtime.py @@ -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 diff --git a/EvoScientist/commands/base.py b/EvoScientist/commands/base.py index 27ff2cb..f52863a 100644 --- a/EvoScientist/commands/base.py +++ b/EvoScientist/commands/base.py @@ -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. diff --git a/EvoScientist/commands/channel_ui.py b/EvoScientist/commands/channel_ui.py index 8c288ff..0d2ba36 100644 --- a/EvoScientist/commands/channel_ui.py +++ b/EvoScientist/commands/channel_ui.py @@ -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", diff --git a/EvoScientist/commands/implementation/session.py b/EvoScientist/commands/implementation/session.py index 47492cd..fcb27d0 100644 --- a/EvoScientist/commands/implementation/session.py +++ b/EvoScientist/commands/implementation/session.py @@ -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: diff --git a/EvoScientist/gateway/__init__.py b/EvoScientist/gateway/__init__.py new file mode 100644 index 0000000..e146632 --- /dev/null +++ b/EvoScientist/gateway/__init__.py @@ -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", +] diff --git a/EvoScientist/gateway/local.py b/EvoScientist/gateway/local.py new file mode 100644 index 0000000..aac0096 --- /dev/null +++ b/EvoScientist/gateway/local.py @@ -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 diff --git a/EvoScientist/gateway/runtime.py b/EvoScientist/gateway/runtime.py new file mode 100644 index 0000000..bf93a72 --- /dev/null +++ b/EvoScientist/gateway/runtime.py @@ -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), + ) diff --git a/EvoScientist/gateway/server.py b/EvoScientist/gateway/server.py new file mode 100644 index 0000000..b792e33 --- /dev/null +++ b/EvoScientist/gateway/server.py @@ -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 diff --git a/EvoScientist/gateway/types.py b/EvoScientist/gateway/types.py new file mode 100644 index 0000000..65c30fc --- /dev/null +++ b/EvoScientist/gateway/types.py @@ -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.""" diff --git a/EvoScientist/middleware/memory_lifecycle.py b/EvoScientist/middleware/memory_lifecycle.py index 8d27a4c..cd2c02c 100644 --- a/EvoScientist/middleware/memory_lifecycle.py +++ b/EvoScientist/middleware/memory_lifecycle.py @@ -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, diff --git a/EvoScientist/sessions.py b/EvoScientist/sessions.py index 11cefc9..b968e58 100644 --- a/EvoScientist/sessions.py +++ b/EvoScientist/sessions.py @@ -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`` / diff --git a/EvoScientist/stream/display.py b/EvoScientist/stream/display.py index c205e60..164afeb 100644 --- a/EvoScientist/stream/display.py +++ b/EvoScientist/stream/display.py @@ -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. diff --git a/EvoScientist/stream/events.py b/EvoScientist/stream/events.py index 6c31ffc..9ac9a95 100644 --- a/EvoScientist/stream/events.py +++ b/EvoScientist/stream/events.py @@ -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() diff --git a/tests/fakes.py b/tests/fakes.py new file mode 100644 index 0000000..aeabb4c --- /dev/null +++ b/tests/fakes.py @@ -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) diff --git a/tests/stream_v3_fakes.py b/tests/stream_v3_fakes.py index b8316ab..d518caa 100644 --- a/tests/stream_v3_fakes.py +++ b/tests/stream_v3_fakes.py @@ -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 diff --git a/tests/test_async_notifier.py b/tests/test_async_notifier.py index b470d21..309b746 100644 --- a/tests/test_async_notifier.py +++ b/tests/test_async_notifier.py @@ -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", diff --git a/tests/test_bus_integration.py b/tests/test_bus_integration.py index 543b806..fe70165 100644 --- a/tests/test_bus_integration.py +++ b/tests/test_bus_integration.py @@ -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.""" diff --git a/tests/test_channel_command_ui.py b/tests/test_channel_command_ui.py index 1b04030..2d65a17 100644 --- a/tests/test_channel_command_ui.py +++ b/tests/test_channel_command_ui.py @@ -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." diff --git a/tests/test_channel_comprehensive.py b/tests/test_channel_comprehensive.py index 4af92c6..c22d941 100644 --- a/tests/test_channel_comprehensive.py +++ b/tests/test_channel_comprehensive.py @@ -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 diff --git a/tests/test_cli_channel_slash.py b/tests/test_cli_channel_slash.py index 9df17be..aa41728 100644 --- a/tests/test_cli_channel_slash.py +++ b/tests/test_cli_channel_slash.py @@ -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" diff --git a/tests/test_compact_command.py b/tests/test_compact_command.py index 027b809..74444be 100644 --- a/tests/test_compact_command.py +++ b/tests/test_compact_command.py @@ -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", diff --git a/tests/test_delete_command.py b/tests/test_delete_command.py index e38532c..082a8ea 100644 --- a/tests/test_delete_command.py +++ b/tests/test_delete_command.py @@ -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() diff --git a/tests/test_event_loop.py b/tests/test_event_loop.py index 4f06e5f..186e66f 100644 --- a/tests/test_event_loop.py +++ b/tests/test_event_loop.py @@ -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()] diff --git a/tests/test_graph_gateway.py b/tests/test_graph_gateway.py new file mode 100644 index 0000000..7fc9055 --- /dev/null +++ b/tests/test_graph_gateway.py @@ -0,0 +1,1095 @@ +"""Tests for the graph/thread gateway abstraction.""" + +from __future__ import annotations + +from types import SimpleNamespace +from typing import Any +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from langchain_core.messages import AIMessage, HumanMessage + +from EvoScientist.gateway import ( + GraphTarget, + LangGraphServerGateway, + LangGraphServerThreadStore, + LocalGraphGateway, + RunRequest, + RuntimeGateways, + create_runtime_gateways, +) +from EvoScientist.gateway.server import _THREAD_SEARCH_LIMIT +from EvoScientist.stream import display as display_mod +from tests.conftest import run_async +from tests.fakes import ( + FakeGraphGateway, + FakeLangGraphClient, + FakeLangGraphThreadsClient, + FakeLangGraphThreadStream, + FakeThreadStore, +) + + +def test_local_gateway_streams_from_injected_streamer(): + seen: dict[str, Any] = {} + + async def _streamer(agent, message, thread_id, **kwargs): + seen.update( + { + "agent": agent, + "message": message, + "thread_id": thread_id, + "metadata": kwargs.get("metadata"), + "media": kwargs.get("media"), + } + ) + yield {"type": "text", "content": "hi"} + yield {"type": "done", "response": "hi"} + + agent = MagicMock() + gateway = LocalGraphGateway() + + async def _collect(): + request = RunRequest( + message="hello", + thread_id="t1", + metadata={"workspace_dir": "/tmp/ws"}, + media=["plot.png"], + target=GraphTarget(local_graph=agent, workspace_dir="/tmp/ws"), + ) + return [event async for event in gateway.stream_events(request)] + + with patch("EvoScientist.stream.events.stream_agent_events", new=_streamer): + events = run_async(_collect()) + + assert events == [ + {"type": "text", "content": "hi"}, + {"type": "done", "response": "hi"}, + ] + assert seen == { + "agent": agent, + "message": "hello", + "thread_id": "t1", + "metadata": {"workspace_dir": "/tmp/ws"}, + "media": ["plot.png"], + } + + +def test_local_graph_gateway_delegates_thread_operations(): + thread_store = FakeThreadStore( + generated_thread_id="new12345", + threads=[{"thread_id": "abc12345"}], + resolved_thread_id="abc12345", + metadata={"workspace_dir": "/tmp/ws"}, + messages=["message"], + exists=True, + deleted=True, + ) + + async def _run(): + gateway = LocalGraphGateway(thread_store=thread_store) + resolution = await gateway.resolve_thread("abc") + return { + "created": await gateway.create_thread(), + "threads": await gateway.list_threads( + limit=3, + include_message_count=True, + ), + "resolution": resolution, + "metadata": await gateway.get_thread_metadata("abc12345"), + "messages": await gateway.get_thread_messages("abc12345"), + "exists": await gateway.thread_exists("abc12345"), + "deleted": await gateway.delete_thread("abc12345"), + } + + result = run_async(_run()) + + assert result["created"] == "new12345" + assert result["threads"] == [{"thread_id": "abc12345"}] + assert result["resolution"].thread_id == "abc12345" + assert result["resolution"].matches == () + assert result["resolution"].found + assert not result["resolution"].ambiguous + assert result["metadata"] == {"workspace_dir": "/tmp/ws"} + assert result["messages"] == ["message"] + assert result["exists"] is True + assert result["deleted"] is True + assert thread_store.calls == [ + ("resolve_thread_id_prefix", "abc"), + ("generate_thread_id", None), + ( + "list_threads", + { + "limit": 3, + "include_message_count": True, + "include_preview": False, + }, + ), + ("get_thread_metadata", "abc12345"), + ("get_thread_messages", "abc12345"), + ("thread_exists", "abc12345"), + ("delete_thread", "abc12345"), + ] + + +def test_local_graph_gateway_reads_state_values(): + agent = MagicMock() + agent.aget_state = AsyncMock( + return_value=SimpleNamespace(values={"async_tasks": {"task-1": {}}}) + ) + gateway = LocalGraphGateway() + + values = run_async( + gateway.get_state_values(GraphTarget(local_graph=agent), "abc12345") + ) + + assert values == {"async_tasks": {"task-1": {}}} + agent.aget_state.assert_awaited_once_with( + {"configurable": {"thread_id": "abc12345"}} + ) + + +def test_local_graph_gateway_updates_state_values(): + agent = MagicMock() + agent.aupdate_state = AsyncMock() + gateway = LocalGraphGateway() + + run_async( + gateway.update_state_values( + GraphTarget(local_graph=agent), + "abc12345", + {"_summarization_event": {"cutoff_index": 2}}, + ) + ) + + agent.aupdate_state.assert_awaited_once_with( + {"configurable": {"thread_id": "abc12345"}}, + {"_summarization_event": {"cutoff_index": 2}}, + as_node="model", + ) + + +def test_local_stream_events_delegates_aclose_to_inner(): + cleanup_ran = False + + async def _streamer(_agent, _message, _thread_id, **_kwargs): + nonlocal cleanup_ran + try: + while True: + yield {"type": "event"} + finally: + cleanup_ran = True + + async def _run(): + gateway = LocalGraphGateway() + stream = gateway.stream_events( + RunRequest( + message="hi", + thread_id="t1", + target=GraphTarget(local_graph=object()), + ) + ) + await stream.__anext__() + await stream.aclose() + assert cleanup_ran is True + + with patch("EvoScientist.stream.events.stream_agent_events", new=_streamer): + run_async(_run()) + + +def test_run_streaming_can_consume_injected_gateway(): + agent = MagicMock() + gateway = FakeGraphGateway( + events=[ + {"type": "text", "content": "gateway-ok"}, + {"type": "done", "response": "gateway-ok"}, + ] + ) + + with patch("EvoScientist.stream.display.Live"): + result = display_mod._run_streaming( + agent=agent, + message="hello", + thread_id="t1", + show_thinking=False, + interactive=True, + metadata={"workspace_dir": "/tmp/ws"}, + gateway=gateway, + ) + + assert result == "gateway-ok" + assert gateway.requests == [ + RunRequest( + message="hello", + thread_id="t1", + metadata={"workspace_dir": "/tmp/ws"}, + target=GraphTarget(local_graph=agent, workspace_dir="/tmp/ws"), + ) + ] + + +def test_resume_command_consumes_context_gateway(): + from EvoScientist.commands.base import CommandContext + from EvoScientist.commands.implementation.session import ResumeCommand + + ui = MagicMock() + ui.handle_session_resume = AsyncMock() + thread_store = FakeThreadStore( + resolved_thread_id="abc12345", + metadata={"workspace_dir": "/restored"}, + ) + ctx = CommandContext( + agent=None, + thread_id="current", + ui=ui, + workspace_dir="/old", + graph_gateway=FakeGraphGateway(thread_store=thread_store), + ) + + run_async(ResumeCommand().execute(ctx, ["abc"])) + + assert ctx.thread_id == "abc12345" + assert ctx.workspace_dir == "/restored" + ui.handle_session_resume.assert_awaited_once_with("abc12345", "/restored") + + +def test_cmd_run_passes_local_graph_gateway(monkeypatch): + from EvoScientist.cli import interactive + + thread_store = FakeThreadStore(generated_thread_id="generated-thread") + + runtime_gateways = RuntimeGateways( + thread_store=thread_store, + graph_gateway=LocalGraphGateway(thread_store=thread_store), + ) + seen: dict[str, Any] = {} + + def _run_streaming(**kwargs): + seen.update(kwargs) + return "ok" + + monkeypatch.setattr(interactive, "run_streaming", _run_streaming) + + agent = MagicMock() + interactive.cmd_run( + agent, + "hello", + thread_id="generated-thread", + show_thinking=False, + workspace_dir="/tmp/ws", + model="test-model", + runtime_gateways=runtime_gateways, + ) + + assert seen["agent"] is agent + assert seen["thread_id"] == "generated-thread" + assert isinstance(seen["gateway"], LocalGraphGateway) + assert seen["gateway"].thread_store is thread_store + + +def test_langgraph_server_thread_store_delegates_to_sdk_threads(): + threads = FakeLangGraphThreadsClient( + threads=[ + { + "thread_id": "abc12345", + "created_at": "2026-01-01T00:00:00Z", + "updated_at": "2026-01-02T00:00:00Z", + "metadata": {"graph_id": "EvoScientist", "workspace_dir": "/tmp/ws"}, + }, + { + "thread_id": "worker123", + "metadata": {"graph_id": "evomemory-turn-worker"}, + }, + ], + states={ + "abc12345": { + "values": { + "messages": [ + {"role": "user", "content": "hello from server"}, + {"role": "assistant", "content": "hi"}, + ] + } + } + }, + ) + client = FakeLangGraphClient(threads) + + def _client_factory(_base_url, _headers): + return client + + store = LangGraphServerThreadStore( + base_url="http://localhost:2024", + client_factory=_client_factory, + ) + + async def _run(): + return { + "created": await store.create_thread( + metadata={"model": "test-model"}, + workspace_dir="/tmp/new-ws", + ), + "threads": await store.list_threads( + include_message_count=True, + include_preview=True, + ), + "resolution": await store.resolve_thread_id_prefix("abc"), + "metadata": await store.get_thread_metadata("abc12345"), + "messages": await store.get_thread_messages("abc12345"), + "exists": await store.thread_exists("abc12345"), + "deleted": await store.delete_thread("abc12345"), + } + + result = run_async(_run()) + + assert result["created"] == "server-thread" + assert len(threads.created) == 1 + assert threads.created[0]["thread_id"] == "server-thread" + created_metadata = threads.created[0]["metadata"] + assert created_metadata["graph_id"] == "EvoScientist" + assert created_metadata["agent_name"] == "EvoScientist" + assert created_metadata["workspace_dir"] == "/tmp/new-ws" + assert created_metadata["model"] == "test-model" + assert isinstance(created_metadata["updated_at"], str) + assert result["threads"] == [ + { + "thread_id": "abc12345", + "created_at": "2026-01-01T00:00:00Z", + "updated_at": "2026-01-02T00:00:00Z", + "workspace_dir": "/tmp/ws", + "model": None, + "metadata": {"graph_id": "EvoScientist", "workspace_dir": "/tmp/ws"}, + "message_count": 2, + "preview": "hello from server", + }, + { + "thread_id": "server-thread", + "created_at": None, + "updated_at": None, + "workspace_dir": "/tmp/new-ws", + "model": "test-model", + "metadata": created_metadata, + "message_count": 0, + "preview": "", + }, + ] + assert result["resolution"] == ("abc12345", []) + assert result["metadata"] == { + "graph_id": "EvoScientist", + "workspace_dir": "/tmp/ws", + } + assert [message.type for message in result["messages"]] == ["human", "ai"] + assert result["exists"] is True + assert result["deleted"] is True + assert threads.deleted == ["abc12345"] + + +def test_langgraph_server_thread_store_limit_zero_pages_all_threads(): + rows = [ + { + "thread_id": f"thread-{index}", + "metadata": {"graph_id": "EvoScientist"}, + } + for index in range(_THREAD_SEARCH_LIMIT + 1) + ] + threads = FakeLangGraphThreadsClient(threads=rows) + store = LangGraphServerThreadStore( + base_url="http://localhost:2024", + client_factory=lambda _base_url, _headers: FakeLangGraphClient(threads), + ) + + result = run_async(store.list_threads(limit=0)) + + assert [row["thread_id"] for row in result] == [ + f"thread-{index}" for index in range(_THREAD_SEARCH_LIMIT + 1) + ] + assert [(search["limit"], search["offset"]) for search in threads.searches] == [ + (_THREAD_SEARCH_LIMIT, 0), + (_THREAD_SEARCH_LIMIT, _THREAD_SEARCH_LIMIT), + ] + + +def test_langgraph_server_thread_store_positive_limit_uses_single_search(): + threads = FakeLangGraphThreadsClient( + threads=[ + { + "thread_id": f"thread-{index}", + "metadata": {"graph_id": "EvoScientist"}, + } + for index in range(3) + ] + ) + store = LangGraphServerThreadStore( + base_url="http://localhost:2024", + client_factory=lambda _base_url, _headers: FakeLangGraphClient(threads), + ) + + result = run_async(store.list_threads(limit=2)) + + assert [row["thread_id"] for row in result] == ["thread-0", "thread-1"] + assert [(search["limit"], search["offset"]) for search in threads.searches] == [ + (2, 0) + ] + + +def test_langgraph_server_thread_store_prefix_resolution_skips_exact_lookup(): + threads = FakeLangGraphThreadsClient( + threads=[ + { + "thread_id": "abc12345", + "metadata": {"graph_id": "EvoScientist"}, + } + ] + ) + store = LangGraphServerThreadStore( + base_url="http://localhost:2024", + client_factory=lambda _base_url, _headers: FakeLangGraphClient(threads), + ) + + result = run_async(store.resolve_thread_id_prefix("abc")) + + assert result == ("abc12345", []) + assert threads.gets == [] + assert len(threads.searches) == 1 + + +def test_langgraph_server_thread_store_prefix_resolution_pages_all_threads(): + rows = [ + { + "thread_id": f"thread-{index}", + "metadata": {"graph_id": "EvoScientist"}, + } + for index in range(_THREAD_SEARCH_LIMIT) + ] + rows.append( + { + "thread_id": "older-thread-match", + "metadata": {"graph_id": "EvoScientist"}, + } + ) + threads = FakeLangGraphThreadsClient(threads=rows) + store = LangGraphServerThreadStore( + base_url="http://localhost:2024", + client_factory=lambda _base_url, _headers: FakeLangGraphClient(threads), + ) + + result = run_async(store.resolve_thread_id_prefix("older-thread")) + + assert result == ("older-thread-match", []) + assert [(search["limit"], search["offset"]) for search in threads.searches] == [ + (_THREAD_SEARCH_LIMIT, 0), + (_THREAD_SEARCH_LIMIT, _THREAD_SEARCH_LIMIT), + ] + + +def test_langgraph_server_thread_store_uuid_resolution_uses_exact_lookup(): + thread_id = "019ed9e4-4253-7f62-b50f-f0470a4b3c9f" + threads = FakeLangGraphThreadsClient( + threads=[ + { + "thread_id": thread_id, + "metadata": {"graph_id": "EvoScientist"}, + } + ] + ) + store = LangGraphServerThreadStore( + base_url="http://localhost:2024", + client_factory=lambda _base_url, _headers: FakeLangGraphClient(threads), + ) + + result = run_async(store.resolve_thread_id_prefix(thread_id)) + + assert result == (thread_id, []) + assert threads.gets == [thread_id] + assert threads.searches == [] + + +def test_langgraph_server_thread_store_uuid_resolution_filters_graph_id(): + thread_id = "019ed9e4-4253-7f62-b50f-f0470a4b3c9f" + threads = FakeLangGraphThreadsClient( + threads=[ + { + "thread_id": thread_id, + "metadata": {"graph_id": "other-agent"}, + } + ] + ) + store = LangGraphServerThreadStore( + base_url="http://localhost:2024", + client_factory=lambda _base_url, _headers: FakeLangGraphClient(threads), + ) + + result = run_async(store.resolve_thread_id_prefix(thread_id)) + + assert result == (None, []) + assert threads.gets == [thread_id] + assert [(search["limit"], search["offset"]) for search in threads.searches] == [ + (_THREAD_SEARCH_LIMIT, 0) + ] + + +def test_langgraph_server_thread_store_clones_thread_with_metadata(): + clone_metadata = { + "clone_purpose": "memory_extraction", + "source_thread_id": "source-thread", + } + threads = FakeLangGraphThreadsClient( + threads=[ + { + "thread_id": "source-thread", + "metadata": {"graph_id": "writing-agent", "workspace_dir": "/tmp/ws"}, + } + ] + ) + store = LangGraphServerThreadStore( + base_url="http://localhost:2024", + client_factory=lambda _base_url, _headers: FakeLangGraphClient(threads), + ) + + cloned_thread_id = run_async( + store.clone_thread("source-thread", metadata=clone_metadata) + ) + + assert cloned_thread_id == "source-thread-copy" + assert threads.copied == ["source-thread"] + assert threads.metadata_updates == [("source-thread-copy", clone_metadata)] + assert threads.threads[-1] == { + "thread_id": "source-thread-copy", + "metadata": { + "graph_id": "writing-agent", + "workspace_dir": "/tmp/ws", + "clone_purpose": "memory_extraction", + "source_thread_id": "source-thread", + }, + } + + +def test_langgraph_server_thread_store_rejects_copy_without_thread_id(): + threads = FakeLangGraphThreadsClient( + threads=[{"thread_id": "source-thread", "metadata": {"graph_id": "agent"}}], + copy_response=None, + ) + store = LangGraphServerThreadStore( + base_url="http://localhost:2024", + client_factory=lambda _base_url, _headers: FakeLangGraphClient(threads), + ) + + async def _run(): + await store.clone_thread("source-thread") + + with pytest.raises(RuntimeError, match="did not return a cloned thread id"): + run_async(_run()) + + +def test_langgraph_server_gateway_clones_thread(): + threads = FakeLangGraphThreadsClient( + threads=[{"thread_id": "source-thread", "metadata": {"graph_id": "agent"}}] + ) + gateway = LangGraphServerGateway( + LangGraphServerThreadStore( + base_url="http://localhost:2024", + client_factory=lambda _base_url, _headers: FakeLangGraphClient(threads), + ) + ) + + cloned_thread_id = run_async( + gateway.clone_thread( + "source-thread", + metadata={"clone_purpose": "manual"}, + target=GraphTarget(graph_id="agent"), + ) + ) + + assert cloned_thread_id == "source-thread-copy" + assert threads.metadata_updates == [ + ("source-thread-copy", {"clone_purpose": "manual"}) + ] + + +def test_local_graph_gateway_clone_thread_is_explicitly_unsupported(): + async def _run(): + await LocalGraphGateway().clone_thread("source-thread") + + with pytest.raises(NotImplementedError, match="does not support thread cloning"): + run_async(_run()) + + +def test_runtime_gateways_can_use_langgraph_server_backend(): + threads = FakeLangGraphThreadsClient() + client = FakeLangGraphClient(threads) + + def _client_factory(_base_url, _headers): + return client + + runtime_gateways = create_runtime_gateways( + backend="langgraph_server", + base_url="http://localhost:2024", + client_factory=_client_factory, + ) + + gateway = runtime_gateways.graph_gateway + + assert isinstance(runtime_gateways.thread_store, LangGraphServerThreadStore) + assert isinstance(gateway, LangGraphServerGateway) + assert gateway.thread_store is runtime_gateways.thread_store + + +def test_langgraph_server_gateway_reads_state_values(): + threads = FakeLangGraphThreadsClient( + threads=[{"thread_id": "abc12345", "metadata": {"graph_id": "EvoScientist"}}], + states={"abc12345": {"values": {"async_tasks": {"task-1": {}}}}}, + ) + gateway = LangGraphServerGateway( + LangGraphServerThreadStore( + base_url="http://localhost:2024", + client_factory=lambda _base_url, _headers: FakeLangGraphClient(threads), + ) + ) + + values = run_async(gateway.get_state_values(GraphTarget(), "abc12345")) + + assert values == {"async_tasks": {"task-1": {}}} + + +def test_langgraph_server_gateway_messages_apply_summarization_event(): + threads = FakeLangGraphThreadsClient( + threads=[{"thread_id": "abc12345", "metadata": {"graph_id": "EvoScientist"}}], + states={ + "abc12345": { + "values": { + "messages": [ + HumanMessage(content="first"), + AIMessage(content="second"), + HumanMessage(content="third"), + ], + "_summarization_event": { + "cutoff_index": 2, + "summary_message": AIMessage(content="summary"), + "file_path": None, + }, + } + } + }, + ) + gateway = LangGraphServerGateway( + LangGraphServerThreadStore( + base_url="http://localhost:2024", + client_factory=lambda _base_url, _headers: FakeLangGraphClient(threads), + ) + ) + + messages = run_async(gateway.get_thread_messages("abc12345")) + + assert len(messages) == 2 + assert isinstance(messages[0], AIMessage) + assert messages[0].content == "summary" + assert isinstance(messages[1], HumanMessage) + assert messages[1].content == "third" + + +def test_langgraph_server_gateway_updates_state_values(): + threads = FakeLangGraphThreadsClient( + threads=[{"thread_id": "abc12345", "metadata": {"graph_id": "EvoScientist"}}], + ) + gateway = LangGraphServerGateway( + LangGraphServerThreadStore( + base_url="http://localhost:2024", + client_factory=lambda _base_url, _headers: FakeLangGraphClient(threads), + ) + ) + + run_async( + gateway.update_state_values( + GraphTarget(), + "abc12345", + {"_summarization_event": {"cutoff_index": 2}}, + ) + ) + + assert threads.state_updates == [ + ("abc12345", {"_summarization_event": {"cutoff_index": 2}}, "model") + ] + + +def test_langgraph_server_gateway_streams_root_protocol_events(): + stream = FakeLangGraphThreadStream( + "abc12345", + events=[ + { + "method": "messages", + "params": { + "namespace": [], + "data": { + "event": "content-block-delta", + "delta": {"type": "text-delta", "text": "hello"}, + }, + }, + }, + { + "method": "messages", + "params": { + "namespace": [], + "data": {"event": "message-finish"}, + }, + }, + ], + ) + threads = FakeLangGraphThreadsClient( + threads=[], + states={"abc12345": {"values": {}}}, + streams={"abc12345": stream}, + ) + gateway = LangGraphServerGateway( + LangGraphServerThreadStore( + base_url="http://localhost:2024", + client_factory=lambda _base_url, _headers: FakeLangGraphClient(threads), + ) + ) + + async def _collect(): + return [ + event + async for event in gateway.stream_events( + RunRequest( + message="hi", + thread_id="abc12345", + metadata={"workspace_dir": "/tmp/ws"}, + target=GraphTarget(graph_id="writing-agent"), + ) + ) + ] + + events = run_async(_collect()) + + assert len(threads.created) == 1 + assert threads.created[0]["thread_id"] == "abc12345" + created_metadata = threads.created[0]["metadata"] + assert created_metadata["graph_id"] == "writing-agent" + assert created_metadata["workspace_dir"] == "/tmp/ws" + assert isinstance(created_metadata["updated_at"], str) + assert len(threads.metadata_updates) == 1 + update_thread_id, update_metadata = threads.metadata_updates[0] + assert update_thread_id == "abc12345" + assert update_metadata["graph_id"] == "writing-agent" + assert update_metadata["workspace_dir"] == "/tmp/ws" + assert isinstance(update_metadata["updated_at"], str) + assert threads.stream_calls == [("abc12345", "writing-agent")] + assert stream.run.starts == [ + { + "input": {"messages": [{"role": "user", "content": "hi"}]}, + "config": {"configurable": {"thread_id": "abc12345"}}, + "metadata": {"workspace_dir": "/tmp/ws"}, + } + ] + assert events == [ + {"type": "text", "content": "hello"}, + {"type": "done", "content": "hello", "response": "hello"}, + ] + + +_OLD_AI = {"type": "ai", "content": "old", "id": "old-ai"} +_HUMAN = {"type": "human", "content": "hi", "id": "human-1"} +_NEW_AI = {"type": "ai", "content": "new", "id": "new-ai"} + + +def _value_snapshot( + messages: list[dict[str, object]], + *, + namespace: list[str] | None = None, +) -> dict[str, object]: + return { + "method": "values", + "params": { + "namespace": namespace or [], + "data": {"messages": messages}, + }, + } + + +def _root_text_delta(text: str) -> dict[str, object]: + return { + "method": "messages", + "params": { + "namespace": [], + "data": { + "event": "content-block-delta", + "delta": {"type": "text-delta", "text": text}, + }, + }, + } + + +def _root_message_finish() -> dict[str, object]: + return { + "method": "messages", + "params": {"namespace": [], "data": {"event": "message-finish"}}, + } + + +def _collect_server_gateway_stream( + events: list[dict[str, object]], + *, + state_messages: list[dict[str, object]] | None = None, +) -> list[dict[str, Any]]: + stream = FakeLangGraphThreadStream("abc12345", events=events) + state_values: dict[str, object] = {} + if state_messages is not None: + state_values["messages"] = state_messages + threads = FakeLangGraphThreadsClient( + threads=[], + states={"abc12345": {"values": state_values}}, + streams={"abc12345": stream}, + ) + gateway = LangGraphServerGateway( + LangGraphServerThreadStore( + base_url="http://localhost:2024", + client_factory=lambda _base_url, _headers: FakeLangGraphClient(threads), + ) + ) + + async def _collect(): + return [ + event + async for event in gateway.stream_events( + RunRequest(message="hi", thread_id="abc12345") + ) + ] + + return run_async(_collect()) + + +def test_langgraph_server_gateway_streams_value_message_snapshots(): + events = _collect_server_gateway_stream( + [ + _value_snapshot([_OLD_AI, _HUMAN]), + _value_snapshot([_OLD_AI, _HUMAN, _NEW_AI]), + ], + state_messages=[_OLD_AI], + ) + + assert events == [ + {"type": "text", "content": "new"}, + {"type": "done", "content": "new", "response": "new"}, + ] + + +def test_langgraph_server_gateway_values_do_not_duplicate_message_stream(): + events = _collect_server_gateway_stream( + [ + _root_text_delta("new"), + _root_message_finish(), + _value_snapshot([_OLD_AI, _HUMAN, _NEW_AI]), + ], + state_messages=[_OLD_AI], + ) + + assert events == [ + {"type": "text", "content": "new"}, + {"type": "done", "content": "new", "response": "new"}, + ] + + +def test_langgraph_server_gateway_ignores_non_root_value_messages(): + events = _collect_server_gateway_stream( + [ + _value_snapshot( + [{"type": "ai", "content": "subagent text", "id": "subagent-ai"}], + namespace=["research:task-1"], + ) + ], + ) + + assert not any(event.get("type") == "text" for event in events) + assert events[-1] == {"type": "done", "content": "", "response": ""} + + +def test_langgraph_server_gateway_emits_state_interrupt_before_done(): + stream = FakeLangGraphThreadStream( + "abc12345", + events=[], + interrupts=[{"interrupt_id": "interrupt-1", "value": None}], + interrupted=True, + ) + threads = FakeLangGraphThreadsClient( + threads=[], + states={ + "abc12345": { + "values": {}, + "interrupts": [ + { + "id": "interrupt-1", + "value": { + "action_requests": [ + { + "name": "execute", + "args": {"command": "echo hello"}, + "id": "tool-1", + } + ], + "review_configs": [ + { + "action_name": "execute", + "allowed_decisions": ["approve", "reject"], + } + ], + }, + } + ], + } + }, + streams={"abc12345": stream}, + ) + gateway = LangGraphServerGateway( + LangGraphServerThreadStore( + base_url="http://localhost:2024", + client_factory=lambda _base_url, _headers: FakeLangGraphClient(threads), + ) + ) + + async def _collect(): + return [ + event + async for event in gateway.stream_events( + RunRequest(message="hi", thread_id="abc12345") + ) + ] + + events = run_async(_collect()) + + assert events == [ + { + "type": "interrupt", + "interrupt_id": "interrupt-1", + "action_requests": [ + { + "name": "execute", + "args": {"command": "echo hello"}, + "id": "tool-1", + } + ], + "review_configs": [ + { + "action_name": "execute", + "allowed_decisions": ["approve", "reject"], + } + ], + }, + {"type": "done", "content": "", "response": ""}, + ] + + +def test_langgraph_server_gateway_streams_subagent_protocol_events(): + stream = FakeLangGraphThreadStream( + "abc12345", + events=[ + { + "method": "lifecycle", + "params": { + "namespace": ["data-analysis-agent:tool-1"], + "data": {"event": "started"}, + }, + }, + { + "method": "messages", + "params": { + "namespace": ["data-analysis-agent:tool-1"], + "data": { + "event": "content-block-delta", + "delta": {"type": "text-delta", "text": "sub text"}, + }, + }, + }, + { + "method": "lifecycle", + "params": { + "namespace": ["data-analysis-agent:tool-1"], + "data": {"event": "completed"}, + }, + }, + ], + ) + threads = FakeLangGraphThreadsClient( + threads=[{"thread_id": "abc12345", "metadata": {"graph_id": "EvoScientist"}}], + states={"abc12345": {"values": {}}}, + streams={"abc12345": stream}, + ) + gateway = LangGraphServerGateway( + LangGraphServerThreadStore( + base_url="http://localhost:2024", + client_factory=lambda _base_url, _headers: FakeLangGraphClient(threads), + ) + ) + + async def _collect(): + return [ + event + async for event in gateway.stream_events( + RunRequest(message="hi", thread_id="abc12345") + ) + ] + + events = run_async(_collect()) + + assert events == [ + { + "type": "subagent_start", + "name": "data-analysis-agent", + "description": "", + "instance_id": "data-analysis-agent:tool-1", + "tool_call_id": "tool-1", + }, + { + "type": "subagent_text", + "subagent": "data-analysis-agent", + "content": "sub text", + "instance_id": "data-analysis-agent:tool-1", + }, + { + "type": "subagent_end", + "name": "data-analysis-agent", + "instance_id": "data-analysis-agent:tool-1", + }, + {"type": "done", "content": "", "response": ""}, + ] + + +def test_langgraph_server_gateway_resumes_interrupt_with_thread_stream(): + from langgraph.types import Command + + stream = FakeLangGraphThreadStream( + "abc12345", + events=[], + interrupts=[{"interrupt_id": "interrupt-1"}], + ) + threads = FakeLangGraphThreadsClient( + threads=[{"thread_id": "abc12345", "metadata": {"graph_id": "EvoScientist"}}], + states={"abc12345": {"values": {}}}, + streams={"abc12345": stream}, + ) + gateway = LangGraphServerGateway( + LangGraphServerThreadStore( + base_url="http://localhost:2024", + client_factory=lambda _base_url, _headers: FakeLangGraphClient(threads), + ) + ) + + async def _collect(): + return [ + event + async for event in gateway.stream_events( + RunRequest( + message=Command(resume={"decisions": [{"allowed": True}]}), + thread_id="abc12345", + ) + ) + ] + + events = run_async(_collect()) + + assert stream.run.starts == [] + assert stream.run.responses == [ + { + "response": {"decisions": [{"allowed": True}]}, + "interrupt_id": "interrupt-1", + } + ] + assert events == [{"type": "done", "content": "", "response": ""}] diff --git a/tests/test_new_command.py b/tests/test_new_command.py index c924481..c3620f1 100644 --- a/tests/test_new_command.py +++ b/tests/test_new_command.py @@ -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, [])) diff --git a/tests/test_observation_memory.py b/tests/test_observation_memory.py index 2373f5c..56d4a68 100644 --- a/tests/test_observation_memory.py +++ b/tests/test_observation_memory.py @@ -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", diff --git a/tests/test_resume_command.py b/tests/test_resume_command.py index 5919ed1..6dde6df 100644 --- a/tests/test_resume_command.py +++ b/tests/test_resume_command.py @@ -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) diff --git a/tests/test_rich_command_ui.py b/tests/test_rich_command_ui.py index d41afab..f540286 100644 --- a/tests/test_rich_command_ui.py +++ b/tests/test_rich_command_ui.py @@ -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): diff --git a/tests/test_serve_agent_holder.py b/tests/test_serve_agent_holder.py index 5170c8f..b94a273 100644 --- a/tests/test_serve_agent_holder.py +++ b/tests/test_serve_agent_holder.py @@ -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 `` 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, diff --git a/tests/test_sessions.py b/tests/test_sessions.py index b0e15f9..a430cb2 100644 --- a/tests/test_sessions.py +++ b/tests/test_sessions.py @@ -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 diff --git a/tests/test_status_bar.py b/tests/test_status_bar.py index 6e81d2b..405b142 100644 --- a/tests/test_status_bar.py +++ b/tests/test_status_bar.py @@ -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")] + ) + ), ) ) diff --git a/tests/test_stream_cancel.py b/tests/test_stream_cancel.py index 0b0db8c..16e20c1 100644 --- a/tests/test_stream_cancel.py +++ b/tests/test_stream_cancel.py @@ -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 diff --git a/tests/test_stream_events.py b/tests/test_stream_events.py index 3ce871a..627ef28 100644 --- a/tests/test_stream_events.py +++ b/tests/test_stream_events.py @@ -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 diff --git a/tests/test_subagent_summarize.py b/tests/test_subagent_summarize.py index 474ac8a..331c0d5 100644 --- a/tests/test_subagent_summarize.py +++ b/tests/test_subagent_summarize.py @@ -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()) diff --git a/tests/test_threads_command.py b/tests/test_threads_command.py index 3baa60e..d86cbbe 100644 --- a/tests/test_threads_command.py +++ b/tests/test_threads_command.py @@ -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] diff --git a/tests/test_ui_runtime.py b/tests/test_ui_runtime.py index 31addee..19b84ae 100644 --- a/tests/test_ui_runtime.py +++ b/tests/test_ui_runtime.py @@ -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"