refactor: LangGraph gateway layer for UI-agnostic graph and thread access (#295)
* feat(gateway): graph gateway protocol * refactor(cli): wire gateway in cli/tui * refactor(gateway): centralize runtime gateway init * chore(gateway): restrict RunRequest message type * feat(gateway): add langgraph server gateway * chore(cli): tighten serve runtime state typing * refactor(cli): route async task state reads through graph gateway * refactor(gateway): support graph targets in server gateway * refactor(cli): route session commands through graph gateway * refactor(cli): fold thread store under graph gateway * refactor(gateway): route graph state access through gateway * refactor(channels): wire graph gateway * refactor(memory): preserve graph threads for cloning * feat(gateway): add thread cloning * fix(tui): pass effective workspace for thread creation * chore(memory): add workspare dir to memory worker metadata * fix(sessions): filter preloaded UUID registy entries by the current scope * test(fakes): use https * refactor(consumer): consolidate imports * fix(stream): optional summarization event * fix(gateway): resolve abbreviated thread IDs by search * fix(gateway): page server thread listings * fix(gateway): emit pending interrupt events * style: fmt * feat(gateway): persist workspace_dir & model in thread metadata * fix(gateway): page server thread prefix resolution * fix(gateway): expose server thread list metadata * refactor: add back type def * refactor: tighten types * revert: add back worker thread deletion The worker thread forking changes are out of scope for now, so to maintain parity with the existing behavior we'll leave this intact. * fix(gateway): apply compaction to server thread history * refactor(stream): restore direct summary replay suppression * fix(gateway): preserve compaction state and server stream output * fix(gateway): close local stream generator on cancellation --------- Co-authored-by: Xi Zhang <106144707+X-iZhang@users.noreply.github.com>
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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)]}
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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():
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -9,6 +9,8 @@ enqueues a ``ChannelMessage`` on a thread-safe ``queue.Queue`` and waits
|
||||
for the main thread to set a response via ``_set_channel_response()``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import queue
|
||||
@@ -17,7 +19,7 @@ import time
|
||||
import uuid
|
||||
from collections.abc import Awaitable, Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from rich.panel import Panel
|
||||
from rich.text import Text
|
||||
@@ -25,6 +27,9 @@ from rich.text import Text
|
||||
from ..commands.base import ChannelRuntime
|
||||
from ..stream.console import console
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..gateway import GraphGateway
|
||||
|
||||
_channel_logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -253,7 +258,8 @@ async def dispatch_channel_slash_command(
|
||||
workspace_dir: str | None,
|
||||
checkpointer: Any,
|
||||
append_system: Callable[[str, str], None],
|
||||
start_new_session_cb: Callable[[], None] | None = None,
|
||||
graph_gateway: GraphGateway,
|
||||
start_new_session_cb: Callable[[], Awaitable[None]] | None = None,
|
||||
handle_session_resume_cb: Callable[..., Awaitable[None]] | None = None,
|
||||
await_agent_ready: Callable[[], Awaitable[Any]] | None = None,
|
||||
on_cmd_completed: Callable[..., Awaitable[None]] | None = None,
|
||||
@@ -284,6 +290,9 @@ async def dispatch_channel_slash_command(
|
||||
Optional lifecycle callbacks forwarded to ``ChannelCommandUI``.
|
||||
Headless serve passes ``None`` — ``/new`` and ``/resume`` degrade
|
||||
gracefully via the default ``ChannelCommandUI`` messages.
|
||||
graph_gateway:
|
||||
Graph gateway forwarded to slash commands and channel resume-history
|
||||
rendering.
|
||||
await_agent_ready:
|
||||
Optional async resolver that blocks until the background agent
|
||||
load finishes. Called only when ``cmd.needs_agent(args)`` is
|
||||
@@ -320,6 +329,7 @@ async def dispatch_channel_slash_command(
|
||||
await_agent_ready=await_agent_ready,
|
||||
on_cmd_completed=on_cmd_completed,
|
||||
channel_runtime=channel_runtime,
|
||||
graph_gateway=graph_gateway,
|
||||
)
|
||||
except Exception as exc:
|
||||
# Last-ditch safety: any uncaught exception from inside the
|
||||
@@ -350,7 +360,8 @@ async def _dispatch_channel_slash_impl(
|
||||
workspace_dir: str | None,
|
||||
checkpointer: Any,
|
||||
append_system: Callable[[str, str], None],
|
||||
start_new_session_cb: Callable[[], None] | None,
|
||||
graph_gateway: GraphGateway,
|
||||
start_new_session_cb: Callable[[], Awaitable[None]] | None,
|
||||
handle_session_resume_cb: Callable[..., Awaitable[None]] | None,
|
||||
await_agent_ready: Callable[[], Awaitable[Any]] | None,
|
||||
on_cmd_completed: Callable[..., Awaitable[None]] | None,
|
||||
@@ -386,6 +397,7 @@ async def _dispatch_channel_slash_impl(
|
||||
append_system_callback=append_system,
|
||||
start_new_session_callback=start_new_session_cb,
|
||||
handle_session_resume_callback=handle_session_resume_cb,
|
||||
graph_gateway=graph_gateway,
|
||||
)
|
||||
ctx = CommandContext(
|
||||
agent=agent_for_ctx,
|
||||
@@ -394,6 +406,7 @@ async def _dispatch_channel_slash_impl(
|
||||
workspace_dir=workspace_dir,
|
||||
checkpointer=checkpointer,
|
||||
channel_runtime=channel_runtime,
|
||||
graph_gateway=graph_gateway,
|
||||
)
|
||||
|
||||
try:
|
||||
@@ -619,7 +632,7 @@ def _try_set_hitl_reply(channel_type: str, chat_id: str, content: str) -> bool:
|
||||
|
||||
def channel_ask_user_prompt(
|
||||
ask_user_data: dict,
|
||||
msg: "ChannelMessage | None" = None,
|
||||
msg: ChannelMessage | None = None,
|
||||
) -> dict:
|
||||
"""Format ask_user questions and collect answers from a channel user.
|
||||
|
||||
@@ -651,7 +664,7 @@ def channel_ask_user_prompt(
|
||||
channel=msg.channel_type,
|
||||
chat_id=msg.chat_id,
|
||||
content=content,
|
||||
metadata=msg.metadata,
|
||||
metadata=msg.metadata or {},
|
||||
)
|
||||
),
|
||||
bus_loop,
|
||||
@@ -748,7 +761,7 @@ def channel_ask_user_prompt(
|
||||
|
||||
def channel_hitl_prompt(
|
||||
action_requests: list,
|
||||
msg: "ChannelMessage",
|
||||
msg: ChannelMessage,
|
||||
) -> list[dict] | None:
|
||||
"""Send HITL approval prompt to channel user and wait for reply.
|
||||
|
||||
@@ -793,7 +806,9 @@ def channel_hitl_prompt(
|
||||
channel=msg.channel_type,
|
||||
chat_id=msg.chat_id,
|
||||
content=content,
|
||||
metadata=metadata if metadata is not None else msg.metadata,
|
||||
metadata=metadata
|
||||
if metadata is not None
|
||||
else msg.metadata or {},
|
||||
)
|
||||
),
|
||||
bus_loop,
|
||||
|
||||
+186
-139
@@ -1,23 +1,32 @@
|
||||
"""Typer command registrations — onboard, config, mcp, main callback."""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
import queue
|
||||
import re
|
||||
from collections.abc import Awaitable, Callable
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from importlib.metadata import version as _pkg_version
|
||||
from pathlib import Path
|
||||
from typing import Annotated, Any, cast
|
||||
from typing import TYPE_CHECKING, Annotated, Any, cast
|
||||
|
||||
import typer
|
||||
from rich.markup import escape
|
||||
from rich.table import Table
|
||||
|
||||
from ..commands.base import Command, CommandContext
|
||||
from ..commands.base import ChannelRuntime, Command, CommandContext
|
||||
from ..gateway import (
|
||||
GraphGateway,
|
||||
GraphTarget,
|
||||
RuntimeGateways,
|
||||
create_runtime_gateways,
|
||||
)
|
||||
from ..llm.context_window import DEFAULT_CONTEXT_WINDOW_FALLBACK, resolve_context_window
|
||||
from ..paths import ensure_dirs, set_active_workspace, set_workspace_root
|
||||
from ..stream.console import console
|
||||
from . import async_notifier
|
||||
from ._app import app, channel_app, config_app, configure_app, mcp_app, sessions_app
|
||||
from ._constants import build_metadata
|
||||
from .agent import (
|
||||
@@ -51,6 +60,11 @@ from .mcp_ui import (
|
||||
_show_mcp_config,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from langgraph.graph.state import CompiledStateGraph
|
||||
|
||||
from ..config import EvoScientistConfig
|
||||
|
||||
# =============================================================================
|
||||
# Onboard command
|
||||
# =============================================================================
|
||||
@@ -596,16 +610,17 @@ def build_compact_summary_renderable(
|
||||
|
||||
|
||||
async def compact_conversation(
|
||||
agent: Any,
|
||||
thread_id: str | None,
|
||||
graph_gateway: GraphGateway,
|
||||
thread_id: str,
|
||||
target: GraphTarget,
|
||||
*,
|
||||
input_tokens_hint: int | None = None,
|
||||
) -> CompactResult:
|
||||
"""Compact the conversation by summarizing old messages.
|
||||
|
||||
Reads the agent's checkpointed state, creates a temporary
|
||||
Reads the graph's checkpointed state, creates a temporary
|
||||
``SummarizationMiddleware``, generates a summary, and writes
|
||||
the compacted state back via ``aupdate_state``.
|
||||
the compacted state back through ``GraphGateway``.
|
||||
|
||||
``input_tokens_hint`` is the real LLM input token count from the last
|
||||
``usage_metadata`` (includes system prompt + tool schemas). When
|
||||
@@ -615,20 +630,17 @@ async def compact_conversation(
|
||||
|
||||
Returns a structured ``CompactResult``.
|
||||
"""
|
||||
if not agent or not thread_id:
|
||||
return CompactResult("noop", "Nothing to compact — start a conversation first.")
|
||||
|
||||
from langchain_core.messages.utils import count_tokens_approximately
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
|
||||
config: RunnableConfig = {"configurable": {"thread_id": thread_id}}
|
||||
|
||||
try:
|
||||
state_snapshot = await agent.aget_state(config)
|
||||
state_values = await graph_gateway.get_state_values(target, thread_id)
|
||||
except Exception as exc:
|
||||
return CompactResult("error", f"Failed to read state: {exc}")
|
||||
|
||||
messages = state_snapshot.values.get("messages", [])
|
||||
messages = state_values.get("messages", [])
|
||||
if not messages:
|
||||
return CompactResult(
|
||||
"noop", "Nothing to compact — no messages in conversation."
|
||||
@@ -661,7 +673,7 @@ async def compact_conversation(
|
||||
)
|
||||
|
||||
# Rebuild effective message list accounting for prior compaction
|
||||
event = state_snapshot.values.get("_summarization_event")
|
||||
event = state_values.get("_summarization_event")
|
||||
effective = middleware._apply_event_to_messages(messages, event)
|
||||
effective_tokens = count_tokens_approximately(effective)
|
||||
|
||||
@@ -783,7 +795,11 @@ async def compact_conversation(
|
||||
"file_path": file_path,
|
||||
}
|
||||
|
||||
await agent.aupdate_state(config, {"_summarization_event": new_event})
|
||||
await graph_gateway.update_state_values(
|
||||
target,
|
||||
thread_id,
|
||||
{"_summarization_event": new_event},
|
||||
)
|
||||
|
||||
return CompactResult(
|
||||
"ok",
|
||||
@@ -809,9 +825,44 @@ async def compact_conversation(
|
||||
_serve_logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ServeRuntimeState:
|
||||
"""Mutable serve-mode runtime shared by the poll loop and slash callbacks."""
|
||||
|
||||
agent: "CompiledStateGraph"
|
||||
thread_id: str
|
||||
workspace_dir: str | None
|
||||
config: "EvoScientistConfig | None"
|
||||
runtime_gateways: RuntimeGateways
|
||||
resume_warning_thread_id: str | None = None
|
||||
|
||||
def set_agent(
|
||||
self,
|
||||
agent: "CompiledStateGraph",
|
||||
channel_runtime: ChannelRuntime | None,
|
||||
) -> None:
|
||||
self.agent = agent
|
||||
if channel_runtime is not None:
|
||||
channel_runtime.agent = agent
|
||||
|
||||
def set_thread_id(
|
||||
self,
|
||||
thread_id: str,
|
||||
channel_runtime: ChannelRuntime | None,
|
||||
*,
|
||||
forget_previous_origin: bool = True,
|
||||
) -> None:
|
||||
old_thread_id = self.thread_id
|
||||
if forget_previous_origin:
|
||||
forget_channel_origin(old_thread_id)
|
||||
self.thread_id = thread_id
|
||||
if channel_runtime is not None:
|
||||
channel_runtime.thread_id = thread_id
|
||||
|
||||
|
||||
def _make_serve_start_new_session_cb(
|
||||
agent_holder: dict[str, Any],
|
||||
channel_runtime: Any | None = None,
|
||||
runtime_state: ServeRuntimeState,
|
||||
channel_runtime: ChannelRuntime | None = None,
|
||||
):
|
||||
"""Build the ``start_new_session_cb`` used by serve mode.
|
||||
|
||||
@@ -821,55 +872,52 @@ def _make_serve_start_new_session_cb(
|
||||
fresh thread id. Without a wired callback the channel user gets
|
||||
``ChannelCommandUI``'s fallback "restart the channel link" message
|
||||
and nothing actually rotates. This helper generates a new thread
|
||||
id, updates the shared holder, and syncs the channel runtime so
|
||||
id, updates the shared runtime state, and syncs the channel runtime so
|
||||
subsequent messages land on the new thread.
|
||||
"""
|
||||
|
||||
def _cb() -> None:
|
||||
from ..sessions import generate_thread_id
|
||||
|
||||
new_tid = generate_thread_id()
|
||||
forget_channel_origin(agent_holder.get("thread_id"))
|
||||
agent_holder["thread_id"] = new_tid
|
||||
if channel_runtime is not None:
|
||||
channel_runtime.thread_id = new_tid
|
||||
async def _cb() -> None:
|
||||
new_tid = await runtime_state.runtime_gateways.graph_gateway.create_thread(
|
||||
GraphTarget(workspace_dir=runtime_state.workspace_dir)
|
||||
)
|
||||
runtime_state.set_thread_id(new_tid, channel_runtime)
|
||||
console.print(f"[dim][serve] New thread: {new_tid}[/dim]")
|
||||
|
||||
return _cb
|
||||
|
||||
|
||||
def _serve_resume_config(
|
||||
agent_holder: dict[str, Any],
|
||||
config: Any | None,
|
||||
) -> Any | None:
|
||||
runtime_state: ServeRuntimeState,
|
||||
config: "EvoScientistConfig | None",
|
||||
) -> "EvoScientistConfig | None":
|
||||
"""Return the effective config to use for serve-mode resume sync."""
|
||||
return config if config is not None else agent_holder.get("config")
|
||||
return config if config is not None else runtime_state.config
|
||||
|
||||
|
||||
async def _apply_serve_resume_state(
|
||||
agent_holder: dict[str, Any],
|
||||
channel_runtime: Any | None,
|
||||
runtime_state: ServeRuntimeState,
|
||||
channel_runtime: ChannelRuntime | None,
|
||||
*,
|
||||
thread_id: str,
|
||||
workspace_dir: str | None,
|
||||
config: Any | None = None,
|
||||
config: "EvoScientistConfig | None" = None,
|
||||
) -> None:
|
||||
"""Adopt a resumed thread/workspace into serve-mode runtime state.
|
||||
|
||||
Workspace-bound resources are rebuilt and synced before mutating the shared
|
||||
holder. The agent is loaded before syncing the external server so a load
|
||||
state. The agent is loaded before syncing the external server so a load
|
||||
failure cannot move the server away from the currently active session.
|
||||
"""
|
||||
import asyncio
|
||||
|
||||
old_workspace = agent_holder.get("workspace_dir")
|
||||
old_workspace = runtime_state.workspace_dir
|
||||
new_workspace = (
|
||||
workspace_dir if workspace_dir and workspace_dir != old_workspace else None
|
||||
)
|
||||
new_agent: Any | None = None
|
||||
workspace_update: tuple[str, CompiledStateGraph] | None = None
|
||||
|
||||
if new_workspace is not None:
|
||||
effective_config = _serve_resume_config(agent_holder, config)
|
||||
effective_config = _serve_resume_config(runtime_state, config)
|
||||
if effective_config is None:
|
||||
raise RuntimeError(
|
||||
"Cannot resume into a different workspace in serve mode without "
|
||||
@@ -885,59 +933,56 @@ async def _apply_serve_resume_state(
|
||||
effective_config,
|
||||
workspace_dir=new_workspace,
|
||||
)
|
||||
workspace_update = (new_workspace, new_agent)
|
||||
except Exception:
|
||||
if old_workspace:
|
||||
set_active_workspace(old_workspace)
|
||||
raise
|
||||
|
||||
old_thread_id = agent_holder.get("thread_id")
|
||||
thread_changed = bool(thread_id) and thread_id != old_thread_id
|
||||
old_thread_id = runtime_state.thread_id
|
||||
thread_changed = thread_id != old_thread_id
|
||||
if thread_changed:
|
||||
forget_channel_origin(old_thread_id)
|
||||
agent_holder["thread_id"] = thread_id
|
||||
if channel_runtime is not None:
|
||||
channel_runtime.thread_id = thread_id
|
||||
runtime_state.set_thread_id(thread_id, channel_runtime)
|
||||
|
||||
if new_workspace is not None:
|
||||
agent_holder["workspace_dir"] = new_workspace
|
||||
agent_holder["agent"] = new_agent
|
||||
if channel_runtime is not None:
|
||||
channel_runtime.agent = new_agent
|
||||
if workspace_update is not None:
|
||||
updated_workspace, updated_agent = workspace_update
|
||||
runtime_state.workspace_dir = updated_workspace
|
||||
runtime_state.set_agent(updated_agent, channel_runtime)
|
||||
|
||||
|
||||
def _make_serve_handle_session_resume_cb(
|
||||
agent_holder: dict[str, Any],
|
||||
channel_runtime: Any | None = None,
|
||||
runtime_state: ServeRuntimeState,
|
||||
channel_runtime: ChannelRuntime | None = None,
|
||||
*,
|
||||
config: Any | None = None,
|
||||
config: "EvoScientistConfig | None" = None,
|
||||
):
|
||||
"""Build the ChannelCommandUI resume callback for serve mode."""
|
||||
|
||||
async def _cb(thread_id: str, workspace_dir: str | None = None) -> None:
|
||||
old_thread_id = agent_holder.get("thread_id")
|
||||
old_thread_id = runtime_state.thread_id
|
||||
await _apply_serve_resume_state(
|
||||
agent_holder,
|
||||
runtime_state,
|
||||
channel_runtime,
|
||||
thread_id=thread_id,
|
||||
workspace_dir=workspace_dir,
|
||||
config=config,
|
||||
)
|
||||
if thread_id and thread_id != old_thread_id:
|
||||
agent_holder["_resume_warning_thread_id"] = thread_id
|
||||
if thread_id != old_thread_id:
|
||||
runtime_state.resume_warning_thread_id = thread_id
|
||||
|
||||
return _cb
|
||||
|
||||
|
||||
def _make_serve_cmd_completed_hook(
|
||||
agent_holder: dict[str, Any],
|
||||
channel_runtime: Any | None = None,
|
||||
runtime_state: ServeRuntimeState,
|
||||
channel_runtime: ChannelRuntime | None = None,
|
||||
*,
|
||||
config: Any | None = None,
|
||||
config: "EvoScientistConfig | None" = None,
|
||||
):
|
||||
"""Build the ``on_cmd_completed`` hook used by serve mode.
|
||||
|
||||
Adopts ``/model`` agent swaps and ``/resume`` thread/workspace
|
||||
swaps back into ``agent_holder`` so the outer poll loop picks up
|
||||
swaps back into ``runtime_state`` so the outer poll loop picks up
|
||||
the new handles on subsequent messages. Also keeps
|
||||
``channel_runtime`` in sync so the bus sees the new values.
|
||||
|
||||
@@ -952,17 +997,17 @@ def _make_serve_cmd_completed_hook(
|
||||
without spinning up the whole serve loop.
|
||||
"""
|
||||
|
||||
async def _hook(ctx: CommandContext, original_agent: Any, cmd: Command) -> None:
|
||||
async def _hook(
|
||||
ctx: CommandContext,
|
||||
original_agent: "CompiledStateGraph",
|
||||
cmd: Command,
|
||||
) -> None:
|
||||
if ctx.agent is not None and ctx.agent is not original_agent:
|
||||
agent_holder["agent"] = ctx.agent
|
||||
if channel_runtime is not None:
|
||||
channel_runtime.agent = ctx.agent
|
||||
runtime_state.set_agent(ctx.agent, channel_runtime)
|
||||
|
||||
old_thread_id = agent_holder.get("thread_id")
|
||||
resume_warning_thread_id = agent_holder.pop(
|
||||
"_resume_warning_thread_id",
|
||||
None,
|
||||
)
|
||||
old_thread_id = runtime_state.thread_id
|
||||
resume_warning_thread_id = runtime_state.resume_warning_thread_id
|
||||
runtime_state.resume_warning_thread_id = None
|
||||
|
||||
# ``/resume`` mutates ``ctx.thread_id`` directly (its UI callback
|
||||
# is a no-op in serve mode since there's no REPL to reset). Pick
|
||||
@@ -975,21 +1020,18 @@ def _make_serve_cmd_completed_hook(
|
||||
new_tid = ctx.thread_id
|
||||
if cmd.name == "/resume":
|
||||
await _apply_serve_resume_state(
|
||||
agent_holder,
|
||||
runtime_state,
|
||||
channel_runtime,
|
||||
thread_id=new_tid,
|
||||
workspace_dir=ctx.workspace_dir,
|
||||
config=config,
|
||||
)
|
||||
else:
|
||||
thread_changed = bool(new_tid) and new_tid != old_thread_id
|
||||
thread_changed = new_tid != old_thread_id
|
||||
if thread_changed:
|
||||
forget_channel_origin(old_thread_id)
|
||||
agent_holder["thread_id"] = new_tid
|
||||
if channel_runtime is not None:
|
||||
channel_runtime.thread_id = new_tid
|
||||
runtime_state.set_thread_id(new_tid, channel_runtime)
|
||||
|
||||
thread_changed = bool(new_tid) and new_tid != old_thread_id
|
||||
thread_changed = new_tid != old_thread_id
|
||||
|
||||
# Surface the in-memory-state limitation to the channel user
|
||||
# for ``/resume`` so the missing history isn't silent. Flush
|
||||
@@ -1014,22 +1056,21 @@ def _make_serve_cmd_completed_hook(
|
||||
def _serve_process_message(
|
||||
msg: ChannelMessage,
|
||||
*,
|
||||
agent_holder: dict[str, Any],
|
||||
runtime_state: ServeRuntimeState,
|
||||
model: str | None,
|
||||
workspace_dir: str,
|
||||
show_thinking: bool,
|
||||
on_cmd_completed: Callable[..., Awaitable[None]] | None = None,
|
||||
handle_session_resume_cb: Callable[..., Awaitable[None]] | None = None,
|
||||
start_new_session_cb: Callable[[], None] | None = None,
|
||||
channel_runtime: Any | None = None,
|
||||
start_new_session_cb: Callable[[], Awaitable[None]] | None = None,
|
||||
channel_runtime: ChannelRuntime | None = None,
|
||||
) -> None:
|
||||
"""Process a single channel message in headless serve mode.
|
||||
|
||||
Headless equivalent of interactive.py's ``_process_channel_message``.
|
||||
No CLI prompt manipulation — just log lines for monitoring.
|
||||
|
||||
``agent_holder`` is a mutable dict (keys: ``agent``, ``thread_id``,
|
||||
``workspace_dir``) shared with the outer ``serve()`` loop.
|
||||
``runtime_state`` is shared with the outer ``serve()`` loop.
|
||||
``on_cmd_completed`` (the agent-swap / session-adoption hook) and
|
||||
``start_new_session_cb`` (thread rotation for ``/new``) are
|
||||
constructed once in ``serve()`` — if omitted, they're rebuilt per
|
||||
@@ -1042,12 +1083,14 @@ def _serve_process_message(
|
||||
from .channel import _bus_loop
|
||||
from .tui_runtime import run_streaming
|
||||
|
||||
runtime_gateways = runtime_state.runtime_gateways
|
||||
|
||||
if not _claim_or_complete_channel_request(msg):
|
||||
return
|
||||
|
||||
remember_channel_origin(agent_holder.get("thread_id"), msg)
|
||||
remember_channel_origin(runtime_state.thread_id, msg)
|
||||
|
||||
runtime_workspace = agent_holder.get("workspace_dir") or workspace_dir
|
||||
runtime_workspace = runtime_state.workspace_dir or workspace_dir
|
||||
|
||||
console.print(
|
||||
f"[dim][{msg.channel_type}] {msg.sender}: {escape(msg.content[:80])}[/dim]"
|
||||
@@ -1138,25 +1181,29 @@ def _serve_process_message(
|
||||
_slash_handled = _slash_loop.run_until_complete(
|
||||
dispatch_channel_slash_command(
|
||||
msg,
|
||||
agent=agent_holder["agent"],
|
||||
thread_id=agent_holder["thread_id"],
|
||||
agent=runtime_state.agent,
|
||||
thread_id=runtime_state.thread_id,
|
||||
workspace_dir=runtime_workspace,
|
||||
checkpointer=None,
|
||||
append_system=lambda t, s="dim": console.print(t, style=s),
|
||||
start_new_session_cb=start_new_session_cb
|
||||
or _make_serve_start_new_session_cb(agent_holder, channel_runtime),
|
||||
or _make_serve_start_new_session_cb(
|
||||
runtime_state,
|
||||
channel_runtime,
|
||||
),
|
||||
handle_session_resume_cb=handle_session_resume_cb
|
||||
or _make_serve_handle_session_resume_cb(
|
||||
agent_holder,
|
||||
runtime_state,
|
||||
channel_runtime,
|
||||
),
|
||||
on_cmd_completed=on_cmd_completed
|
||||
or _make_serve_cmd_completed_hook(
|
||||
agent_holder,
|
||||
runtime_state,
|
||||
channel_runtime,
|
||||
config=agent_holder.get("config"),
|
||||
config=runtime_state.config,
|
||||
),
|
||||
channel_runtime=channel_runtime,
|
||||
graph_gateway=runtime_gateways.graph_gateway,
|
||||
)
|
||||
)
|
||||
except Exception as exc:
|
||||
@@ -1178,7 +1225,7 @@ def _serve_process_message(
|
||||
# A channel-issued /new or /resume rotates the thread inside the
|
||||
# dispatch above; re-bind the now-current thread to this channel
|
||||
# so async-notifier turns on it still forward back here.
|
||||
remember_channel_origin(agent_holder["thread_id"], msg)
|
||||
remember_channel_origin(runtime_state.thread_id, msg)
|
||||
console.print(f"[dim][{msg.channel_type}] Replied to {msg.sender}[/dim]")
|
||||
return
|
||||
|
||||
@@ -1186,9 +1233,9 @@ def _serve_process_message(
|
||||
try:
|
||||
response = run_streaming(
|
||||
ui_backend="cli",
|
||||
agent=agent_holder["agent"],
|
||||
agent=runtime_state.agent,
|
||||
message=msg.content,
|
||||
thread_id=agent_holder["thread_id"],
|
||||
thread_id=runtime_state.thread_id,
|
||||
show_thinking=show_thinking,
|
||||
interactive=True,
|
||||
metadata=meta,
|
||||
@@ -1198,6 +1245,7 @@ def _serve_process_message(
|
||||
hitl_prompt_fn=_hitl_prompt,
|
||||
ask_user_prompt_fn=_ask_user_prompt,
|
||||
cancel_scope=_channel_message_cancel_scope(msg),
|
||||
gateway=runtime_gateways.graph_gateway,
|
||||
)
|
||||
except Exception as e:
|
||||
response = f"Error: {e}"
|
||||
@@ -1216,7 +1264,7 @@ def _serve_process_message(
|
||||
|
||||
def _serve_drain_notifications(
|
||||
*,
|
||||
agent_holder: dict,
|
||||
runtime_state: ServeRuntimeState,
|
||||
model: str | None,
|
||||
workspace_dir: str,
|
||||
show_thinking: bool,
|
||||
@@ -1228,8 +1276,6 @@ def _serve_drain_notifications(
|
||||
"""
|
||||
import asyncio as _aio
|
||||
|
||||
from EvoScientist.cli import async_notifier
|
||||
|
||||
from .tui_runtime import run_streaming
|
||||
|
||||
def _run_notification_message(text: str, notifs: list) -> None:
|
||||
@@ -1239,20 +1285,21 @@ def _serve_drain_notifications(
|
||||
|
||||
for line_text, line_style in format_notification_lines(notifs):
|
||||
console.print(line_text, style=line_style, markup=False)
|
||||
# Use the current workspace from agent_holder (updated by /resume's
|
||||
# Use the current workspace from runtime_state (updated by /resume's
|
||||
# session-rebind callback), falling back to the startup value.
|
||||
runtime_workspace = agent_holder.get("workspace_dir") or workspace_dir
|
||||
runtime_workspace = runtime_state.workspace_dir or workspace_dir
|
||||
meta = build_metadata(runtime_workspace, model)
|
||||
tid = agent_holder["thread_id"]
|
||||
tid = runtime_state.thread_id
|
||||
try:
|
||||
response = run_streaming(
|
||||
ui_backend="cli",
|
||||
agent=agent_holder["agent"],
|
||||
agent=runtime_state.agent,
|
||||
message=text,
|
||||
thread_id=tid,
|
||||
show_thinking=show_thinking,
|
||||
interactive=True,
|
||||
metadata=meta,
|
||||
gateway=runtime_state.runtime_gateways.graph_gateway,
|
||||
)
|
||||
except Exception as exc:
|
||||
_serve_logger.warning("Notification agent turn failed: %s", exc)
|
||||
@@ -1270,22 +1317,24 @@ def _serve_drain_notifications(
|
||||
async def _run_notification_message_async(text: str, notifs: list) -> None:
|
||||
await _aio.to_thread(_run_notification_message, text, notifs)
|
||||
|
||||
async def _read_async_tasks() -> dict:
|
||||
agent = agent_holder.get("agent")
|
||||
thread_id = agent_holder.get("thread_id")
|
||||
if agent is None or not thread_id:
|
||||
return {}
|
||||
try:
|
||||
snap = await agent.aget_state({"configurable": {"thread_id": thread_id}})
|
||||
return (snap.values or {}).get("async_tasks") or {}
|
||||
except Exception:
|
||||
async def _read_async_tasks() -> async_notifier.AsyncTasksState:
|
||||
thread_id = runtime_state.thread_id
|
||||
if not thread_id:
|
||||
return {}
|
||||
return await async_notifier.read_async_tasks_from_gateway(
|
||||
runtime_state.runtime_gateways.graph_gateway,
|
||||
GraphTarget(
|
||||
local_graph=runtime_state.agent,
|
||||
workspace_dir=runtime_state.workspace_dir,
|
||||
),
|
||||
thread_id,
|
||||
)
|
||||
|
||||
async def _consume() -> None:
|
||||
await async_notifier.consume_notifications(
|
||||
run_message=_run_notification_message_async,
|
||||
read_async_tasks_state=_read_async_tasks,
|
||||
current_thread_id=agent_holder.get("thread_id"),
|
||||
current_thread_id=runtime_state.thread_id,
|
||||
)
|
||||
|
||||
_notif_loop: _aio.AbstractEventLoop | None = None
|
||||
@@ -1406,22 +1455,22 @@ def serve(
|
||||
)
|
||||
console.print("[dim]Loading agent...[/dim]")
|
||||
agent = _load_agent(workspace_dir=ws, config=config)
|
||||
from ..sessions import generate_thread_id
|
||||
|
||||
tid = generate_thread_id()
|
||||
runtime_gateways = create_runtime_gateways()
|
||||
tid = asyncio.run(
|
||||
runtime_gateways.graph_gateway.create_thread(GraphTarget(workspace_dir=ws))
|
||||
)
|
||||
|
||||
# Mutable holder shared with _serve_process_message so ``/model``
|
||||
# invoked over a channel can hot-swap the agent for subsequent
|
||||
# messages. A pass-by-value parameter gets captured once at startup
|
||||
# and never updated.
|
||||
agent_holder: dict[str, Any] = {
|
||||
"agent": agent,
|
||||
"thread_id": tid,
|
||||
"workspace_dir": ws,
|
||||
"config": config,
|
||||
}
|
||||
|
||||
from ..commands.base import ChannelRuntime
|
||||
# Mutable runtime shared with _serve_process_message so channel slash
|
||||
# commands can update the active agent/thread/workspace for subsequent
|
||||
# messages.
|
||||
runtime_state = ServeRuntimeState(
|
||||
agent=agent,
|
||||
thread_id=tid,
|
||||
workspace_dir=ws,
|
||||
config=config,
|
||||
runtime_gateways=runtime_gateways,
|
||||
)
|
||||
|
||||
channel_runtime = ChannelRuntime(agent=agent, thread_id=tid)
|
||||
|
||||
@@ -1429,13 +1478,13 @@ def serve(
|
||||
# them for every inbound message. Without this hoist each message
|
||||
# would allocate a fresh closure pair.
|
||||
_serve_on_cmd_completed = _make_serve_cmd_completed_hook(
|
||||
agent_holder, channel_runtime, config=config
|
||||
runtime_state, channel_runtime, config=config
|
||||
)
|
||||
_serve_handle_session_resume_cb = _make_serve_handle_session_resume_cb(
|
||||
agent_holder, channel_runtime, config=config
|
||||
runtime_state, channel_runtime, config=config
|
||||
)
|
||||
_serve_start_new_session_cb = _make_serve_start_new_session_cb(
|
||||
agent_holder, channel_runtime
|
||||
runtime_state, channel_runtime
|
||||
)
|
||||
|
||||
_start_channels_bus_mode(
|
||||
@@ -1486,7 +1535,7 @@ def serve(
|
||||
try:
|
||||
_serve_process_message(
|
||||
msg,
|
||||
agent_holder=agent_holder,
|
||||
runtime_state=runtime_state,
|
||||
model=config.model,
|
||||
workspace_dir=ws,
|
||||
show_thinking=effective_channel_thinking,
|
||||
@@ -1500,11 +1549,9 @@ def serve(
|
||||
break
|
||||
|
||||
# Poll notification queue when idle (no channel message was pending).
|
||||
from EvoScientist.cli import async_notifier
|
||||
|
||||
if async_notifier.has_pending_notifications(agent_holder.get("thread_id")):
|
||||
if async_notifier.has_pending_notifications(runtime_state.thread_id):
|
||||
_serve_drain_notifications(
|
||||
agent_holder=agent_holder,
|
||||
runtime_state=runtime_state,
|
||||
model=config.model,
|
||||
workspace_dir=ws,
|
||||
show_thinking=effective_channel_thinking,
|
||||
@@ -2173,27 +2220,26 @@ def _main_callback(
|
||||
# Single-shot mode: wrap in persistent checkpointer
|
||||
import asyncio
|
||||
|
||||
from ..sessions import (
|
||||
generate_thread_id,
|
||||
get_checkpointer,
|
||||
resolve_thread_id_prefix,
|
||||
)
|
||||
from ..sessions import get_checkpointer
|
||||
from .interactive import cmd_run
|
||||
from .resume_hint import print_resume_hint
|
||||
|
||||
runtime_gateways = create_runtime_gateways()
|
||||
graph_gateway = runtime_gateways.graph_gateway
|
||||
|
||||
async def _single_shot():
|
||||
async with get_checkpointer() as checkpointer:
|
||||
# Resolve resume target first so a bad --resume/--thread-id
|
||||
# exits before the slow _load_agent() provider setup.
|
||||
if thread_id:
|
||||
resolved, matches = await resolve_thread_id_prefix(thread_id)
|
||||
if resolved:
|
||||
tid = resolved
|
||||
elif matches:
|
||||
resolution = await graph_gateway.resolve_thread(thread_id)
|
||||
if resolution.thread_id:
|
||||
tid = resolution.thread_id
|
||||
elif resolution.matches:
|
||||
console.print(
|
||||
f"[yellow]Ambiguous thread ID '{escape(thread_id)}'. Matches:[/yellow]"
|
||||
)
|
||||
for s in matches:
|
||||
for s in resolution.matches:
|
||||
console.print(f" [cyan]{escape(s)}[/cyan]")
|
||||
raise typer.Exit(1)
|
||||
else:
|
||||
@@ -2202,7 +2248,7 @@ def _main_callback(
|
||||
)
|
||||
raise typer.Exit(1)
|
||||
else:
|
||||
tid = generate_thread_id()
|
||||
tid = await graph_gateway.create_thread()
|
||||
console.print("[dim]Loading agent...[/dim]")
|
||||
agent = _load_agent(
|
||||
workspace_dir=workspace_dir,
|
||||
@@ -2218,6 +2264,7 @@ def _main_callback(
|
||||
workspace_dir=workspace_dir,
|
||||
model=config.model,
|
||||
ui_backend=config.ui_backend,
|
||||
runtime_gateways=runtime_gateways,
|
||||
)
|
||||
finally:
|
||||
try:
|
||||
|
||||
+153
-99
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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
|
||||
@@ -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),
|
||||
)
|
||||
@@ -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
|
||||
@@ -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."""
|
||||
@@ -18,7 +18,7 @@ from dataclasses import dataclass
|
||||
from datetime import UTC, datetime
|
||||
from enum import StrEnum
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any, NotRequired, TypedDict, TypeVar, cast
|
||||
from typing import TYPE_CHECKING, Any, NotRequired, Protocol, TypedDict, TypeVar, cast
|
||||
|
||||
from langchain.agents.middleware.types import AgentMiddleware, AgentState
|
||||
from langchain_core.messages import AIMessage, BaseMessage, ToolMessage, filter_messages
|
||||
@@ -44,7 +44,7 @@ from ..memory.worker_activity import (
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from langgraph_sdk.schema import Config, Input
|
||||
from langgraph_sdk.schema import Config, Input, Run, Thread
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -139,6 +139,7 @@ class MemoryWorkerLaunchArgs(TypedDict):
|
||||
|
||||
role: MemoryLifecycleRole
|
||||
memory_dir: str | Path
|
||||
workspace_dir: str | Path
|
||||
project_id: str
|
||||
source_agent: str
|
||||
session_id: str
|
||||
@@ -154,6 +155,62 @@ class MemoryWorkerRunPayload(TypedDict):
|
||||
config: Config
|
||||
|
||||
|
||||
class _SyncMemoryWorkerThreads(Protocol):
|
||||
def create(
|
||||
self,
|
||||
*,
|
||||
graph_id: str,
|
||||
metadata: dict[str, str],
|
||||
) -> Thread: ...
|
||||
|
||||
|
||||
class _SyncMemoryWorkerRuns(Protocol):
|
||||
def create(
|
||||
self,
|
||||
thread_id: str,
|
||||
assistant_id: str,
|
||||
*,
|
||||
input: Input,
|
||||
metadata: dict[str, str],
|
||||
config: Config,
|
||||
) -> Run: ...
|
||||
|
||||
def get(self, thread_id: str, run_id: str) -> Run: ...
|
||||
|
||||
|
||||
class _SyncMemoryWorkerClient(Protocol):
|
||||
threads: _SyncMemoryWorkerThreads
|
||||
runs: _SyncMemoryWorkerRuns
|
||||
|
||||
|
||||
class _AsyncMemoryWorkerThreads(Protocol):
|
||||
async def create(
|
||||
self,
|
||||
*,
|
||||
graph_id: str,
|
||||
metadata: dict[str, str],
|
||||
) -> Thread: ...
|
||||
|
||||
|
||||
class _AsyncMemoryWorkerRuns(Protocol):
|
||||
async def create(
|
||||
self,
|
||||
thread_id: str,
|
||||
assistant_id: str,
|
||||
*,
|
||||
input: Input,
|
||||
metadata: dict[str, str],
|
||||
config: Config,
|
||||
) -> Run: ...
|
||||
|
||||
async def get(self, thread_id: str, run_id: str) -> Run: ...
|
||||
|
||||
|
||||
class _AsyncMemoryWorkerClient(Protocol):
|
||||
threads: _AsyncMemoryWorkerThreads
|
||||
runs: _AsyncMemoryWorkerRuns
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _SummaryWriteArgs:
|
||||
"""Concrete metadata needed to write a subagent execution summary."""
|
||||
@@ -617,20 +674,6 @@ def _safe_segment(value: str) -> str:
|
||||
return safe.strip("-") or "unknown"
|
||||
|
||||
|
||||
def _worker_thread_id(
|
||||
*,
|
||||
role: MemoryLifecycleRole,
|
||||
session_id: str,
|
||||
source_agent: str,
|
||||
trajectory: list[CompactMessage],
|
||||
) -> str:
|
||||
"""Return a deterministic thread id for a background worker run."""
|
||||
key = "\n".join(
|
||||
[role.value, session_id, source_agent, _trajectory_digest(trajectory)]
|
||||
)
|
||||
return f"evomemory-{role.value}:{_short_hash(key)}"
|
||||
|
||||
|
||||
def _agent_result_model(result: Mapping[str, object], model_type: type[T]) -> T | None:
|
||||
"""Extract a DeepAgents/LangChain structured response from agent state."""
|
||||
value = result.get("structured_response")
|
||||
@@ -952,30 +995,49 @@ def _runs_create_kwargs(kwargs: MemoryWorkerRunPayload) -> MemoryWorkerRunPayloa
|
||||
return cast("MemoryWorkerRunPayload", _merge_runs_config_kwargs(dict(kwargs)))
|
||||
|
||||
|
||||
def _worker_workspace_dir(workspace_dir: str | Path) -> str:
|
||||
return str(Path(workspace_dir).expanduser().resolve())
|
||||
|
||||
|
||||
def _memory_worker_metadata(
|
||||
*,
|
||||
role: MemoryLifecycleRole,
|
||||
workspace_dir: str | Path,
|
||||
project_id: str,
|
||||
source_agent: str,
|
||||
session_id: str,
|
||||
trajectory_digest: str,
|
||||
) -> dict[str, str]:
|
||||
return {
|
||||
"run_kind": f"evomemory_{role.value}_worker",
|
||||
"source_session_id": session_id,
|
||||
"source_agent": source_agent,
|
||||
"project_id": project_id,
|
||||
"trajectory_digest": trajectory_digest,
|
||||
"workspace_dir": _worker_workspace_dir(workspace_dir),
|
||||
}
|
||||
|
||||
|
||||
def _memory_worker_run_kwargs(
|
||||
*,
|
||||
role: MemoryLifecycleRole,
|
||||
thread_id: str,
|
||||
workspace_dir: str | Path,
|
||||
project_id: str,
|
||||
source_agent: str,
|
||||
session_id: str,
|
||||
trajectory: list[CompactMessage],
|
||||
) -> MemoryWorkerRunPayload:
|
||||
"""Build the LangGraph SDK run payload for a memory worker."""
|
||||
worker_thread_id = _worker_thread_id(
|
||||
role=role,
|
||||
session_id=session_id,
|
||||
source_agent=source_agent,
|
||||
trajectory=trajectory,
|
||||
)
|
||||
trajectory_digest = _trajectory_digest(trajectory)
|
||||
metadata = {
|
||||
"agent_name": "EvoScientist",
|
||||
"run_kind": f"evomemory_{role.value}_worker",
|
||||
"source_session_id": session_id,
|
||||
"source_agent": source_agent,
|
||||
"project_id": project_id,
|
||||
"trajectory_digest": trajectory_digest,
|
||||
}
|
||||
metadata = _memory_worker_metadata(
|
||||
role=role,
|
||||
workspace_dir=workspace_dir,
|
||||
project_id=project_id,
|
||||
source_agent=source_agent,
|
||||
session_id=session_id,
|
||||
trajectory_digest=trajectory_digest,
|
||||
)
|
||||
payload: MemoryWorkerRunPayload = {
|
||||
"assistant_id": role.graph_id,
|
||||
"input": {
|
||||
@@ -993,7 +1055,7 @@ def _memory_worker_run_kwargs(
|
||||
"metadata": metadata,
|
||||
"config": {
|
||||
"configurable": {
|
||||
"thread_id": worker_thread_id,
|
||||
"thread_id": thread_id,
|
||||
"evomemory_source_session_id": session_id,
|
||||
"evomemory_source_agent": source_agent,
|
||||
"evomemory_project_id": project_id,
|
||||
@@ -1126,7 +1188,7 @@ def _watch_memory_worker_run_sync(
|
||||
|
||||
|
||||
def _spawn_memory_worker_status_task(
|
||||
client: Any,
|
||||
client: _AsyncMemoryWorkerClient,
|
||||
*,
|
||||
thread_id: str,
|
||||
run_id: str,
|
||||
@@ -1140,7 +1202,7 @@ def _spawn_memory_worker_status_task(
|
||||
|
||||
|
||||
async def _watch_memory_worker_run_async(
|
||||
client: Any,
|
||||
client: _AsyncMemoryWorkerClient,
|
||||
*,
|
||||
thread_id: str,
|
||||
run_id: str,
|
||||
@@ -1189,6 +1251,7 @@ def _launch_memory_worker(
|
||||
*,
|
||||
role: MemoryLifecycleRole,
|
||||
memory_dir: str | Path,
|
||||
workspace_dir: str | Path,
|
||||
project_id: str,
|
||||
source_agent: str,
|
||||
session_id: str,
|
||||
@@ -1204,12 +1267,24 @@ def _launch_memory_worker(
|
||||
logger.info("Skipping EvoMemory worker launch; LangGraph dev is unavailable")
|
||||
return
|
||||
|
||||
client = get_sync_client(url=url, headers={"x-auth-scheme": "langsmith"})
|
||||
thread = client.threads.create(graph_id=role.graph_id)
|
||||
client: _SyncMemoryWorkerClient = get_sync_client(
|
||||
url=url, headers={"x-auth-scheme": "langsmith"}
|
||||
)
|
||||
metadata = _memory_worker_metadata(
|
||||
role=role,
|
||||
workspace_dir=workspace_dir,
|
||||
project_id=project_id,
|
||||
source_agent=source_agent,
|
||||
session_id=session_id,
|
||||
trajectory_digest=_trajectory_digest(trajectory),
|
||||
)
|
||||
thread = client.threads.create(graph_id=role.graph_id, metadata=metadata)
|
||||
worker_thread_id = str(thread["thread_id"])
|
||||
before_outputs = snapshot_memory_outputs(memory_dir)
|
||||
payload = _memory_worker_run_kwargs(
|
||||
role=role,
|
||||
thread_id=worker_thread_id,
|
||||
workspace_dir=workspace_dir,
|
||||
project_id=project_id,
|
||||
source_agent=source_agent,
|
||||
session_id=session_id,
|
||||
@@ -1244,6 +1319,7 @@ async def _alaunch_memory_worker(
|
||||
*,
|
||||
role: MemoryLifecycleRole,
|
||||
memory_dir: str | Path,
|
||||
workspace_dir: str | Path,
|
||||
project_id: str,
|
||||
source_agent: str,
|
||||
session_id: str,
|
||||
@@ -1259,12 +1335,24 @@ async def _alaunch_memory_worker(
|
||||
logger.info("Skipping EvoMemory worker launch; LangGraph dev is unavailable")
|
||||
return
|
||||
|
||||
client = get_client(url=url, headers={"x-auth-scheme": "langsmith"})
|
||||
thread = await client.threads.create(graph_id=role.graph_id)
|
||||
client: _AsyncMemoryWorkerClient = get_client(
|
||||
url=url, headers={"x-auth-scheme": "langsmith"}
|
||||
)
|
||||
metadata = _memory_worker_metadata(
|
||||
role=role,
|
||||
workspace_dir=workspace_dir,
|
||||
project_id=project_id,
|
||||
source_agent=source_agent,
|
||||
session_id=session_id,
|
||||
trajectory_digest=_trajectory_digest(trajectory),
|
||||
)
|
||||
thread = await client.threads.create(graph_id=role.graph_id, metadata=metadata)
|
||||
worker_thread_id = str(thread["thread_id"])
|
||||
before_outputs = await asyncio.to_thread(snapshot_memory_outputs, memory_dir)
|
||||
payload = _memory_worker_run_kwargs(
|
||||
role=role,
|
||||
thread_id=worker_thread_id,
|
||||
workspace_dir=workspace_dir,
|
||||
project_id=project_id,
|
||||
source_agent=source_agent,
|
||||
session_id=session_id,
|
||||
@@ -1310,6 +1398,9 @@ class EvoMemoryLifecycleMiddleware(AgentMiddleware):
|
||||
source_agent: str,
|
||||
) -> None:
|
||||
self._memory_dir = Path(memory_dir).expanduser()
|
||||
self._workspace_dir = Path(
|
||||
_paths.WORKSPACE_ROOT if workspace_dir is None else workspace_dir
|
||||
).expanduser()
|
||||
self._project_id = project_id
|
||||
self._role = role
|
||||
self._source_agent = source_agent
|
||||
@@ -1329,6 +1420,7 @@ class EvoMemoryLifecycleMiddleware(AgentMiddleware):
|
||||
return {
|
||||
"role": MemoryLifecycleRole.TURN,
|
||||
"memory_dir": self._memory_dir,
|
||||
"workspace_dir": self._workspace_dir,
|
||||
"project_id": self._project_id,
|
||||
"source_agent": self._source_agent,
|
||||
"session_id": session_id,
|
||||
@@ -1341,6 +1433,7 @@ class EvoMemoryLifecycleMiddleware(AgentMiddleware):
|
||||
return {
|
||||
"role": MemoryLifecycleRole.SUBAGENT,
|
||||
"memory_dir": self._memory_dir,
|
||||
"workspace_dir": self._workspace_dir,
|
||||
"project_id": self._project_id,
|
||||
"source_agent": self._source_agent,
|
||||
"session_id": session_id,
|
||||
|
||||
+129
-90
@@ -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`` /
|
||||
|
||||
@@ -12,7 +12,7 @@ import os
|
||||
import re
|
||||
import threading
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from rich.console import Group # type: ignore[import-untyped]
|
||||
from rich.live import Live # type: ignore[import-untyped]
|
||||
@@ -21,10 +21,10 @@ from rich.panel import Panel # type: ignore[import-untyped]
|
||||
from rich.spinner import Spinner # type: ignore[import-untyped]
|
||||
from rich.text import Text # type: ignore[import-untyped]
|
||||
|
||||
from ..gateway import GraphGateway, GraphRunInput, GraphTarget, RunRequest
|
||||
from ..paths import resolve_virtual_path
|
||||
from .console import console
|
||||
from .diff_format import build_edit_diff
|
||||
from .events import stream_agent_events
|
||||
from .formatter import ToolResultFormatter
|
||||
from .state import (
|
||||
StreamState,
|
||||
@@ -40,6 +40,9 @@ from .utils import (
|
||||
is_success,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from langgraph.graph.state import CompiledStateGraph
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Shared globals
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -47,6 +50,19 @@ from .utils import (
|
||||
# Media file extensions that should trigger on_file_write callback
|
||||
_MEDIA_EXTENSIONS = {".png", ".jpg", ".jpeg", ".gif", ".webp", ".bmp", ".svg", ".pdf"}
|
||||
|
||||
|
||||
def _graph_target_for_local_agent(
|
||||
agent: "CompiledStateGraph",
|
||||
metadata: dict[str, object] | None = None,
|
||||
) -> GraphTarget:
|
||||
workspace = None
|
||||
if metadata is not None:
|
||||
raw_workspace = metadata.get("workspace_dir")
|
||||
if isinstance(raw_workspace, str) and raw_workspace:
|
||||
workspace = raw_workspace
|
||||
return GraphTarget(local_graph=agent, workspace_dir=workspace)
|
||||
|
||||
|
||||
# LLM output sometimes omits the CommonMark-required space after `#` (e.g.
|
||||
# "###文件系统"), which makes Rich render the line as raw text. The lookahead
|
||||
# `(?=[^ \t#\r\n])` requires a real non-excluded next char, so the helper is
|
||||
@@ -1264,8 +1280,8 @@ def _resolve_ask_user_prompt(ask_user_data: dict) -> dict:
|
||||
|
||||
|
||||
def _run_streaming(
|
||||
agent: Any,
|
||||
message: Any,
|
||||
agent: "CompiledStateGraph",
|
||||
message: GraphRunInput,
|
||||
thread_id: str,
|
||||
show_thinking: bool,
|
||||
interactive: bool,
|
||||
@@ -1274,11 +1290,12 @@ def _run_streaming(
|
||||
on_file_write: Callable[[str], None] | None = None,
|
||||
on_stream_event: Callable[[str, Any], Any] | None = None,
|
||||
status_footer_builder: Callable[[], Any] | None = None,
|
||||
metadata: dict | None = None,
|
||||
metadata: dict[str, object] | None = None,
|
||||
hitl_prompt_fn: Callable[[list], list[dict] | None] | None = None,
|
||||
ask_user_prompt_fn: Callable[[dict], dict] | None = None,
|
||||
cancel_scope: str | None = None,
|
||||
*,
|
||||
gateway: GraphGateway,
|
||||
_state: StreamState | None = None,
|
||||
_hitl_depth: int = 0,
|
||||
_media_sent: set[str] | None = None,
|
||||
@@ -1305,6 +1322,7 @@ def _run_streaming(
|
||||
when the agent writes a media file (image/pdf) via write_file.
|
||||
metadata: Optional metadata dict forwarded to ``stream_agent_events``
|
||||
for LangGraph checkpoint persistence.
|
||||
gateway: Graph/thread gateway supplied by the active runtime.
|
||||
|
||||
Returns:
|
||||
The final response text.
|
||||
@@ -1328,8 +1346,13 @@ def _run_streaming(
|
||||
|
||||
async def _consume() -> None:
|
||||
nonlocal _sent_thinking_text, _todo_sent
|
||||
async for event in stream_agent_events(
|
||||
agent, message, thread_id, metadata=metadata
|
||||
async for event in gateway.stream_events(
|
||||
RunRequest(
|
||||
message=message,
|
||||
thread_id=thread_id,
|
||||
metadata=metadata,
|
||||
target=_graph_target_for_local_agent(agent, metadata),
|
||||
)
|
||||
):
|
||||
if is_stream_cancel_requested(cancel_scope):
|
||||
_stopped_response()
|
||||
@@ -1568,6 +1591,7 @@ def _run_streaming(
|
||||
hitl_prompt_fn=hitl_prompt_fn,
|
||||
ask_user_prompt_fn=ask_user_prompt_fn,
|
||||
cancel_scope=cancel_scope,
|
||||
gateway=gateway,
|
||||
_state=state,
|
||||
_hitl_depth=_hitl_depth + 1,
|
||||
_media_sent=_media_sent,
|
||||
@@ -1606,6 +1630,7 @@ def _run_streaming(
|
||||
hitl_prompt_fn=hitl_prompt_fn,
|
||||
ask_user_prompt_fn=ask_user_prompt_fn,
|
||||
cancel_scope=cancel_scope,
|
||||
gateway=gateway,
|
||||
_state=state,
|
||||
_hitl_depth=_hitl_depth + 1,
|
||||
_media_sent=_media_sent,
|
||||
@@ -1631,10 +1656,12 @@ def _run_streaming(
|
||||
|
||||
|
||||
async def _astream_to_console(
|
||||
agent: Any,
|
||||
agent: "CompiledStateGraph",
|
||||
message: str,
|
||||
thread_id: str,
|
||||
show_thinking: bool = True,
|
||||
*,
|
||||
gateway: GraphGateway,
|
||||
) -> str:
|
||||
"""Stream agent events to console using static prints (thread-safe, no Live).
|
||||
|
||||
@@ -1654,7 +1681,13 @@ async def _astream_to_console(
|
||||
"""
|
||||
state = StreamState()
|
||||
|
||||
async for event in stream_agent_events(agent, message, thread_id):
|
||||
async for event in gateway.stream_events(
|
||||
RunRequest(
|
||||
message=message,
|
||||
thread_id=thread_id,
|
||||
target=_graph_target_for_local_agent(agent),
|
||||
)
|
||||
):
|
||||
etype = state.handle_event(event)
|
||||
|
||||
# Only show subagent starts as real-time progress.
|
||||
|
||||
+198
-98
@@ -8,9 +8,9 @@ import base64
|
||||
import inspect
|
||||
import mimetypes
|
||||
import os
|
||||
from collections.abc import AsyncGenerator, AsyncIterator
|
||||
from collections.abc import AsyncGenerator, AsyncIterator, Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
from typing import Any, TypeAlias
|
||||
|
||||
from langchain_core.messages import AIMessage, AIMessageChunk, BaseMessage, ToolMessage
|
||||
from langgraph.graph import END
|
||||
@@ -38,6 +38,17 @@ from .v3_payloads import (
|
||||
_usage_counts,
|
||||
)
|
||||
|
||||
UserMessageContent: TypeAlias = str | list[dict[str, object]]
|
||||
GraphRunInput: TypeAlias = str | Command
|
||||
LangGraphStreamInput: TypeAlias = dict[str, list[dict[str, object]]] | Command
|
||||
_ValueMessageKey: TypeAlias = tuple[str, ...]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _AssistantValueMessage:
|
||||
key: _ValueMessageKey
|
||||
content: object
|
||||
|
||||
|
||||
def _is_interrupt_error_message(message: object) -> bool:
|
||||
if not isinstance(message, str):
|
||||
@@ -181,11 +192,17 @@ class _V3EventProcessor:
|
||||
self,
|
||||
emitter: StreamEventEmitter,
|
||||
subagents: _SubagentRegistry,
|
||||
baseline_summarization_signature: tuple[object, ...] | None,
|
||||
existing_summarization_event: Mapping[str, object] | None,
|
||||
existing_messages: object = None,
|
||||
process_value_messages: bool = False,
|
||||
) -> None:
|
||||
self.emitter = emitter
|
||||
self.subagents = subagents
|
||||
self.baseline_summarization_signature = baseline_summarization_signature
|
||||
self._suppressed_summarization_signature = _summarization_event_signature(
|
||||
existing_summarization_event
|
||||
)
|
||||
self._seen_value_message_keys = self._message_keys(existing_messages)
|
||||
self._process_value_message_snapshots = process_value_messages
|
||||
self.full_response = ""
|
||||
self._summarization_in_progress = False
|
||||
self._tool_inputs: dict[
|
||||
@@ -211,18 +228,82 @@ class _V3EventProcessor:
|
||||
return []
|
||||
|
||||
if method == "messages":
|
||||
return self._process_message_event(_event_data(event), subagent, namespace)
|
||||
events = self._process_message_event(
|
||||
_event_data(event), subagent, namespace
|
||||
)
|
||||
if not namespace and any(item.get("type") == "text" for item in events):
|
||||
self._process_value_message_snapshots = False
|
||||
return events
|
||||
if method == "tools":
|
||||
return self._process_tool_event(namespace, _event_data(event), subagent)
|
||||
if method == "updates":
|
||||
return self._process_update_event(_event_data(event))
|
||||
if method == "values":
|
||||
events: list[dict[str, Any]] = []
|
||||
params = event.get("params") or {}
|
||||
interrupts = params.get("interrupts") or ()
|
||||
if interrupts:
|
||||
return self._process_update_event({"__interrupt__": interrupts})
|
||||
events.extend(self._process_update_event({"__interrupt__": interrupts}))
|
||||
if self._process_value_message_snapshots and not namespace:
|
||||
events.extend(self._process_value_messages(_event_data(event)))
|
||||
return events
|
||||
if method == "input.requested":
|
||||
return self._process_input_requested(event.get("params"))
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
def _message_keys(cls, messages: object) -> set[_ValueMessageKey]:
|
||||
if not isinstance(messages, list):
|
||||
return set()
|
||||
keys: set[_ValueMessageKey] = set()
|
||||
for message in messages:
|
||||
if parsed := cls._assistant_value_message(message):
|
||||
keys.add(parsed.key)
|
||||
return keys
|
||||
|
||||
@staticmethod
|
||||
def _assistant_value_message(message: object) -> _AssistantValueMessage | None:
|
||||
message_map = _as_raw_map(message)
|
||||
if message_map is not None:
|
||||
raw_id = message_map.get("id")
|
||||
raw_role = message_map.get("type") or message_map.get("role")
|
||||
content = message_map.get("content")
|
||||
elif isinstance(message, BaseMessage):
|
||||
raw_id = message.id
|
||||
raw_role = message.type
|
||||
content = message.content
|
||||
else:
|
||||
return None
|
||||
|
||||
if raw_role not in ("ai", "assistant"):
|
||||
return None
|
||||
if raw_id:
|
||||
key = ("id", str(raw_id))
|
||||
else:
|
||||
key = ("body", str(raw_role), repr(content))
|
||||
return _AssistantValueMessage(key=key, content=content)
|
||||
|
||||
def _process_value_messages(self, data: object) -> list[dict[str, Any]]:
|
||||
data_map = _as_raw_map(data)
|
||||
if data_map is None:
|
||||
return []
|
||||
messages = data_map.get("messages")
|
||||
if not isinstance(messages, list):
|
||||
return []
|
||||
|
||||
events: list[dict[str, Any]] = []
|
||||
for message in messages:
|
||||
parsed = self._assistant_value_message(message)
|
||||
if parsed is None:
|
||||
continue
|
||||
if parsed.key in self._seen_value_message_keys:
|
||||
continue
|
||||
text = _text_from_content(parsed.content)
|
||||
if text:
|
||||
self._seen_value_message_keys.add(parsed.key)
|
||||
events.extend(self._emit_text(text, subagent=None))
|
||||
return events
|
||||
|
||||
def _process_message_event(
|
||||
self,
|
||||
data: object,
|
||||
@@ -523,7 +604,7 @@ class _V3EventProcessor:
|
||||
signature = _summarization_event_signature(summarization_event)
|
||||
if (
|
||||
signature is not None
|
||||
and signature == self.baseline_summarization_signature
|
||||
and signature == self._suppressed_summarization_signature
|
||||
):
|
||||
return events
|
||||
summary_message = summarization_event.get("summary_message")
|
||||
@@ -543,35 +624,56 @@ class _V3EventProcessor:
|
||||
continue
|
||||
|
||||
interrupt_value = interrupt_obj.value
|
||||
if not isinstance(interrupt_value, dict):
|
||||
continue
|
||||
|
||||
iv_type = interrupt_value.get("type")
|
||||
interrupt_id = interrupt_obj.id or "default"
|
||||
if iv_type == "ask_user":
|
||||
questions = interrupt_value.get("questions", [])
|
||||
tc_id = str(interrupt_value.get("tool_call_id", ""))
|
||||
events.extend(
|
||||
self._dedupe_interrupt_event(
|
||||
self.emitter.ask_user_interrupt(
|
||||
interrupt_id, questions, tc_id
|
||||
).data
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
action_reqs = interrupt_value.get("action_requests", [])
|
||||
review_cfgs = interrupt_value.get("review_configs", [])
|
||||
if action_reqs:
|
||||
events.extend(
|
||||
self._dedupe_interrupt_event(
|
||||
self.emitter.interrupt(
|
||||
interrupt_id, action_reqs, review_cfgs
|
||||
).data
|
||||
)
|
||||
)
|
||||
events.extend(self._process_interrupt_value(interrupt_id, interrupt_value))
|
||||
return events
|
||||
|
||||
def _process_input_requested(self, params: object) -> list[dict[str, Any]]:
|
||||
params_map = _as_raw_map(params)
|
||||
if params_map is None:
|
||||
return []
|
||||
data = _as_raw_map(params_map.get("data"))
|
||||
if data is None:
|
||||
return []
|
||||
interrupt_id = str(data.get("interrupt_id") or "default")
|
||||
return self._process_interrupt_value(interrupt_id, data.get("value"))
|
||||
|
||||
def _process_interrupt_value(
|
||||
self,
|
||||
interrupt_id: str,
|
||||
interrupt_value: object,
|
||||
) -> list[dict[str, Any]]:
|
||||
interrupt_map = _as_raw_map(interrupt_value)
|
||||
if interrupt_map is None:
|
||||
return []
|
||||
|
||||
iv_type = interrupt_map.get("type")
|
||||
if iv_type == "ask_user":
|
||||
raw_questions = interrupt_map.get("questions")
|
||||
questions = raw_questions if isinstance(raw_questions, list) else []
|
||||
tc_id = str(interrupt_map.get("tool_call_id", ""))
|
||||
return self._dedupe_interrupt_event(
|
||||
self.emitter.ask_user_interrupt(
|
||||
interrupt_id,
|
||||
questions,
|
||||
tc_id,
|
||||
).data
|
||||
)
|
||||
|
||||
raw_action_reqs = interrupt_map.get("action_requests")
|
||||
action_reqs = raw_action_reqs if isinstance(raw_action_reqs, list) else []
|
||||
raw_review_cfgs = interrupt_map.get("review_configs")
|
||||
review_cfgs = raw_review_cfgs if isinstance(raw_review_cfgs, list) else None
|
||||
if action_reqs:
|
||||
return self._dedupe_interrupt_event(
|
||||
self.emitter.interrupt(
|
||||
interrupt_id,
|
||||
action_reqs,
|
||||
review_cfgs,
|
||||
).data
|
||||
)
|
||||
return []
|
||||
|
||||
def _dedupe_interrupt_event(self, event: dict[str, Any]) -> list[dict[str, Any]]:
|
||||
signature = repr(event)
|
||||
if signature in self._emitted_interrupts:
|
||||
@@ -631,9 +733,63 @@ class _V3EventProcessor:
|
||||
return _text_from_content(payload.content)
|
||||
|
||||
|
||||
async def build_agent_stream_input(
|
||||
message: GraphRunInput,
|
||||
*,
|
||||
media: list[str] | None = None,
|
||||
) -> LangGraphStreamInput:
|
||||
"""Build the LangGraph run input shared by local and server gateways."""
|
||||
if not isinstance(message, str):
|
||||
return message
|
||||
|
||||
user_content: UserMessageContent = message
|
||||
if media:
|
||||
image_exts = frozenset({".jpg", ".jpeg", ".png", ".gif", ".webp", ".bmp"})
|
||||
max_inline_size = 5 * 1024 * 1024
|
||||
content_blocks: list[dict[str, object]] = []
|
||||
if message:
|
||||
content_blocks.append({"type": "text", "text": message})
|
||||
|
||||
def _read_file_b64(path: str) -> str:
|
||||
with open(path, "rb") as fh:
|
||||
return base64.b64encode(fh.read()).decode("ascii")
|
||||
|
||||
file_refs: list[str] = []
|
||||
for path in media:
|
||||
ext = os.path.splitext(path)[1].lower()
|
||||
is_image = ext in image_exts and await asyncio.to_thread(
|
||||
os.path.isfile, path
|
||||
)
|
||||
if is_image:
|
||||
fsize = await asyncio.to_thread(os.path.getsize, path)
|
||||
if fsize <= max_inline_size:
|
||||
mime = mimetypes.guess_type(path)[0] or "image/png"
|
||||
b64 = await asyncio.to_thread(_read_file_b64, path)
|
||||
content_blocks.append(
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": f"data:{mime};base64,{b64}",
|
||||
},
|
||||
}
|
||||
)
|
||||
else:
|
||||
file_refs.append(path)
|
||||
else:
|
||||
file_refs.append(path)
|
||||
if file_refs:
|
||||
ref_text = "\n".join(
|
||||
f"[attached file: {os.path.basename(p)}] path: {p}" for p in file_refs
|
||||
)
|
||||
content_blocks.append({"type": "text", "text": ref_text})
|
||||
if content_blocks:
|
||||
user_content = content_blocks
|
||||
return {"messages": [{"role": "user", "content": user_content}]}
|
||||
|
||||
|
||||
async def stream_agent_events(
|
||||
agent: Any,
|
||||
message: str | Command,
|
||||
message: GraphRunInput,
|
||||
thread_id: str,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
media: list[str] | None = None,
|
||||
@@ -662,74 +818,18 @@ async def stream_agent_events(
|
||||
if metadata:
|
||||
config["metadata"] = metadata
|
||||
emitter = StreamEventEmitter()
|
||||
|
||||
clear_memory_worker_saved_counts()
|
||||
# Build input for agent.astream_events()
|
||||
if isinstance(message, str):
|
||||
# Build user message content: text + inline images + file path references
|
||||
user_content: str | list[dict[str, object]] = message
|
||||
if media:
|
||||
_IMAGE_EXTS = frozenset({".jpg", ".jpeg", ".png", ".gif", ".webp", ".bmp"})
|
||||
_MAX_INLINE_SIZE = 5 * 1024 * 1024 # 5 MB
|
||||
content_blocks: list[dict[str, object]] = []
|
||||
if message:
|
||||
content_blocks.append({"type": "text", "text": message})
|
||||
|
||||
def _read_file_b64(path: str) -> str:
|
||||
with open(path, "rb") as fh:
|
||||
return base64.b64encode(fh.read()).decode("ascii")
|
||||
|
||||
file_refs: list[str] = []
|
||||
for path in media:
|
||||
ext = os.path.splitext(path)[1].lower()
|
||||
is_image = ext in _IMAGE_EXTS and await asyncio.to_thread(
|
||||
os.path.isfile, path
|
||||
)
|
||||
if is_image:
|
||||
fsize = await asyncio.to_thread(os.path.getsize, path)
|
||||
if fsize <= _MAX_INLINE_SIZE:
|
||||
mime = mimetypes.guess_type(path)[0] or "image/png"
|
||||
b64 = await asyncio.to_thread(_read_file_b64, path)
|
||||
content_blocks.append(
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": f"data:{mime};base64,{b64}",
|
||||
},
|
||||
}
|
||||
)
|
||||
else:
|
||||
file_refs.append(path)
|
||||
else:
|
||||
file_refs.append(path)
|
||||
if file_refs:
|
||||
ref_text = "\n".join(
|
||||
f"[attached file: {os.path.basename(p)}] path: {p}"
|
||||
for p in file_refs
|
||||
)
|
||||
content_blocks.append({"type": "text", "text": ref_text})
|
||||
if content_blocks:
|
||||
user_content = content_blocks
|
||||
astream_input: dict[str, list[dict[str, object]]] | Command = {
|
||||
"messages": [{"role": "user", "content": user_content}]
|
||||
}
|
||||
else:
|
||||
# HITL resume: Command object passed directly to agent
|
||||
astream_input = message
|
||||
|
||||
_baseline_summarization_signature: tuple[object, ...] | None = None
|
||||
|
||||
existing_summarization_event: Mapping[str, object] | None = None
|
||||
try:
|
||||
snapshot = await agent.aget_state(config)
|
||||
values = snapshot.values
|
||||
if isinstance(values, dict):
|
||||
baseline_event = _find_summarization_event_payload(values)
|
||||
_baseline_summarization_signature = _summarization_event_signature(
|
||||
baseline_event
|
||||
)
|
||||
existing_summarization_event = _find_summarization_event_payload(
|
||||
getattr(snapshot, "values", None)
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
clear_memory_worker_saved_counts()
|
||||
astream_input = await build_agent_stream_input(message, media=media)
|
||||
|
||||
stream: Any | None = None
|
||||
producers: list[asyncio.Task[Any]] = []
|
||||
_run_raised: bool = False
|
||||
@@ -753,7 +853,7 @@ async def stream_agent_events(
|
||||
processor = _V3EventProcessor(
|
||||
emitter,
|
||||
subagents,
|
||||
_baseline_summarization_signature,
|
||||
existing_summarization_event,
|
||||
)
|
||||
queue: asyncio.Queue[Any] = asyncio.Queue()
|
||||
producer_done = object()
|
||||
|
||||
+653
@@ -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)
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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."""
|
||||
|
||||
|
||||
@@ -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."
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -10,9 +10,21 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from EvoScientist.cli.channel import (
|
||||
ChannelMessage,
|
||||
dispatch_channel_slash_command,
|
||||
)
|
||||
from EvoScientist.cli.channel import (
|
||||
dispatch_channel_slash_command as _dispatch_channel_slash_command,
|
||||
)
|
||||
from tests.conftest import run_async as _run
|
||||
from tests.fakes import FakeGraphGateway, FakeThreadStore
|
||||
|
||||
|
||||
def _thread_store() -> FakeThreadStore:
|
||||
return FakeThreadStore()
|
||||
|
||||
|
||||
def dispatch_channel_slash_command(*args, **kwargs):
|
||||
kwargs.setdefault("graph_gateway", FakeGraphGateway())
|
||||
return _dispatch_channel_slash_command(*args, **kwargs)
|
||||
|
||||
|
||||
def _make_msg(
|
||||
@@ -107,6 +119,45 @@ def test_successful_slash_execution_sets_response_and_breadcrumb():
|
||||
assert any("Executed command from" in t for t in breadcrumbs)
|
||||
|
||||
|
||||
def test_slash_dispatch_passes_graph_gateway_to_command_context():
|
||||
msg = _make_msg()
|
||||
fake_cmd = MagicMock()
|
||||
fake_cmd.needs_agent.return_value = False
|
||||
append = MagicMock()
|
||||
graph_gateway = FakeGraphGateway(thread_store=_thread_store())
|
||||
captured = {}
|
||||
|
||||
async def _execute(_content, ctx):
|
||||
captured["graph_gateway"] = ctx.graph_gateway
|
||||
return True
|
||||
|
||||
with (
|
||||
patch(
|
||||
"EvoScientist.commands.manager.manager.resolve",
|
||||
return_value=(fake_cmd, ["core"]),
|
||||
),
|
||||
patch(
|
||||
"EvoScientist.commands.manager.manager.execute",
|
||||
new=AsyncMock(side_effect=_execute),
|
||||
),
|
||||
patch("EvoScientist.cli.channel._set_channel_response"),
|
||||
):
|
||||
handled = _run(
|
||||
dispatch_channel_slash_command(
|
||||
msg,
|
||||
agent="fake-agent",
|
||||
thread_id="t1",
|
||||
workspace_dir="/tmp",
|
||||
checkpointer=None,
|
||||
append_system=append,
|
||||
graph_gateway=graph_gateway,
|
||||
)
|
||||
)
|
||||
|
||||
assert handled is True
|
||||
assert captured["graph_gateway"] is graph_gateway
|
||||
|
||||
|
||||
def test_needs_agent_awaits_loader_and_passes_result():
|
||||
"""Commands with needs_agent=True must await the loader and the
|
||||
resulting agent must flow through the CommandContext."""
|
||||
@@ -144,7 +195,9 @@ def test_needs_agent_awaits_loader_and_passes_result():
|
||||
)
|
||||
assert handled is True
|
||||
await_called.assert_called_once()
|
||||
ctx_arg = mock_execute.await_args.args[1]
|
||||
await_args = mock_execute.await_args
|
||||
assert await_args is not None
|
||||
ctx_arg = await_args.args[1]
|
||||
assert ctx_arg.agent == "ready-agent"
|
||||
|
||||
|
||||
|
||||
+55
-103
@@ -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",
|
||||
|
||||
@@ -1,87 +1,47 @@
|
||||
"""Tests for the /delete command."""
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from tests.conftest import run_async as _run
|
||||
from tests.fakes import FakeGraphGateway, FakeThreadStore
|
||||
|
||||
|
||||
def _ctx(thread_id="current"):
|
||||
def _ctx(thread_id="current", thread_store=None):
|
||||
from EvoScientist.commands.base import CommandContext
|
||||
|
||||
store = thread_store or FakeThreadStore()
|
||||
ui = MagicMock()
|
||||
ui.supports_interactive = True
|
||||
return CommandContext(agent=None, thread_id=thread_id, ui=ui), ui
|
||||
|
||||
|
||||
def _patches(thread_exists=False, similar=None, deleted=True, threads=None):
|
||||
"""Return a context manager stack patching the sessions module."""
|
||||
from contextlib import ExitStack
|
||||
|
||||
stack = ExitStack()
|
||||
stack.enter_context(
|
||||
patch(
|
||||
"EvoScientist.sessions.thread_exists",
|
||||
new=AsyncMock(return_value=thread_exists),
|
||||
)
|
||||
)
|
||||
stack.enter_context(
|
||||
patch(
|
||||
"EvoScientist.sessions.find_similar_threads",
|
||||
new=AsyncMock(return_value=similar or []),
|
||||
)
|
||||
)
|
||||
stack.enter_context(
|
||||
patch(
|
||||
"EvoScientist.sessions.delete_thread",
|
||||
new=AsyncMock(return_value=deleted),
|
||||
)
|
||||
)
|
||||
stack.enter_context(
|
||||
patch(
|
||||
"EvoScientist.sessions.list_threads",
|
||||
new=AsyncMock(return_value=threads or []),
|
||||
)
|
||||
)
|
||||
return stack
|
||||
return CommandContext(
|
||||
agent=None,
|
||||
thread_id=thread_id,
|
||||
ui=ui,
|
||||
graph_gateway=FakeGraphGateway(thread_store=store),
|
||||
), ui
|
||||
|
||||
|
||||
class TestDeleteCommand:
|
||||
def test_refuses_to_delete_current(self):
|
||||
from EvoScientist.commands.implementation.session import DeleteCommand
|
||||
|
||||
ctx, ui = _ctx(thread_id="current")
|
||||
# Inline the patches here (rather than using ``_patches``) so we
|
||||
# can keep a direct handle on the ``delete_thread`` mock and
|
||||
# assert on it *inside* the context. Asserting after the
|
||||
# context exits hits the real function (no ``await_count``
|
||||
# attr), which silently degrades into ``assert True``.
|
||||
mock_delete = AsyncMock(return_value=True)
|
||||
with (
|
||||
patch(
|
||||
"EvoScientist.sessions.thread_exists",
|
||||
new=AsyncMock(return_value=True),
|
||||
),
|
||||
patch("EvoScientist.sessions.delete_thread", new=mock_delete),
|
||||
patch(
|
||||
"EvoScientist.sessions.find_similar_threads",
|
||||
new=AsyncMock(return_value=[]),
|
||||
),
|
||||
patch(
|
||||
"EvoScientist.sessions.list_threads",
|
||||
new=AsyncMock(return_value=[]),
|
||||
),
|
||||
):
|
||||
_run(DeleteCommand().execute(ctx, ["current"]))
|
||||
assert mock_delete.await_count == 0
|
||||
thread_store = FakeThreadStore(resolved_thread_id="current", deleted=True)
|
||||
ctx, ui = _ctx(thread_id="current", thread_store=thread_store)
|
||||
_run(DeleteCommand().execute(ctx, ["current"]))
|
||||
assert ("delete_thread", "current") not in thread_store.calls
|
||||
msgs = [c.args[0] for c in ui.append_system.call_args_list]
|
||||
assert any("Cannot delete the current session" in m for m in msgs)
|
||||
|
||||
def test_happy_path_success(self):
|
||||
from EvoScientist.commands.implementation.session import DeleteCommand
|
||||
|
||||
ctx, ui = _ctx(thread_id="current")
|
||||
with _patches(thread_exists=True, deleted=True):
|
||||
_run(DeleteCommand().execute(ctx, ["other-thread"]))
|
||||
ctx, ui = _ctx(
|
||||
thread_id="current",
|
||||
thread_store=FakeThreadStore(
|
||||
resolved_thread_id="other-thread",
|
||||
deleted=True,
|
||||
),
|
||||
)
|
||||
_run(DeleteCommand().execute(ctx, ["other-thread"]))
|
||||
msgs = [c.args[0] for c in ui.append_system.call_args_list]
|
||||
assert any("Deleted session other-thread" in m for m in msgs)
|
||||
|
||||
@@ -89,26 +49,25 @@ class TestDeleteCommand:
|
||||
from EvoScientist.commands.implementation.session import DeleteCommand
|
||||
|
||||
ctx, ui = _ctx()
|
||||
with _patches(thread_exists=False, similar=[]):
|
||||
_run(DeleteCommand().execute(ctx, ["missing"]))
|
||||
_run(DeleteCommand().execute(ctx, ["missing"]))
|
||||
msgs = [c.args[0] for c in ui.append_system.call_args_list]
|
||||
assert any("not found" in m for m in msgs)
|
||||
|
||||
def test_ambiguous_prefix(self):
|
||||
from EvoScientist.commands.implementation.session import DeleteCommand
|
||||
|
||||
ctx, ui = _ctx()
|
||||
with _patches(thread_exists=False, similar=["abc-one", "abc-two"]):
|
||||
_run(DeleteCommand().execute(ctx, ["abc"]))
|
||||
ctx, ui = _ctx(thread_store=FakeThreadStore(matches=["abc-one", "abc-two"]))
|
||||
_run(DeleteCommand().execute(ctx, ["abc"]))
|
||||
msgs = [c.args[0] for c in ui.append_system.call_args_list]
|
||||
assert any("Ambiguous" in m for m in msgs)
|
||||
|
||||
def test_prefix_resolves_to_unique_match(self):
|
||||
from EvoScientist.commands.implementation.session import DeleteCommand
|
||||
|
||||
ctx, ui = _ctx()
|
||||
with _patches(thread_exists=False, similar=["abc-one"], deleted=True):
|
||||
_run(DeleteCommand().execute(ctx, ["abc"]))
|
||||
ctx, ui = _ctx(
|
||||
thread_store=FakeThreadStore(resolved_thread_id="abc-one", deleted=True)
|
||||
)
|
||||
_run(DeleteCommand().execute(ctx, ["abc"]))
|
||||
msgs = [c.args[0] for c in ui.append_system.call_args_list]
|
||||
assert any("Deleted session abc-one" in m for m in msgs)
|
||||
|
||||
@@ -116,8 +75,7 @@ class TestDeleteCommand:
|
||||
from EvoScientist.commands.implementation.session import DeleteCommand
|
||||
|
||||
ctx, ui = _ctx()
|
||||
with _patches(threads=[]):
|
||||
_run(DeleteCommand().execute(ctx, []))
|
||||
_run(DeleteCommand().execute(ctx, []))
|
||||
msgs = [c.args[0] for c in ui.append_system.call_args_list]
|
||||
assert any("No sessions to delete" in m for m in msgs)
|
||||
|
||||
@@ -136,6 +94,7 @@ class TestDeleteCommand:
|
||||
"updated_at": None,
|
||||
}
|
||||
]
|
||||
with _patches(threads=threads):
|
||||
_run(DeleteCommand().execute(ctx, []))
|
||||
store = FakeThreadStore(threads=threads)
|
||||
ctx.graph_gateway = FakeGraphGateway(thread_store=store)
|
||||
_run(DeleteCommand().execute(ctx, []))
|
||||
ui.wait_for_thread_pick.assert_awaited_once()
|
||||
|
||||
+63
-67
@@ -6,6 +6,7 @@ from unittest.mock import Mock, patch
|
||||
import pytest
|
||||
|
||||
from EvoScientist.stream.display import _create_event_loop, _get_event_loop
|
||||
from tests.fakes import FakeGraphGateway
|
||||
|
||||
|
||||
class _TrackingEventLoopPolicy(asyncio.DefaultEventLoopPolicy):
|
||||
@@ -128,7 +129,7 @@ class TestMultipleStreamingCalls:
|
||||
# Mock agent that returns simple events
|
||||
mock_agent = Mock()
|
||||
|
||||
async def mock_stream(*args, **kwargs):
|
||||
async def mock_stream(_request):
|
||||
"""Mock event stream."""
|
||||
yield {"type": "text", "content": "test response"}
|
||||
yield {"type": "done", "response": "test response"}
|
||||
@@ -141,38 +142,39 @@ class TestMultipleStreamingCalls:
|
||||
except RuntimeError:
|
||||
pass
|
||||
|
||||
# Patch the stream_agent_events function
|
||||
with patch(
|
||||
"EvoScientist.stream.display.stream_agent_events", side_effect=mock_stream
|
||||
):
|
||||
# Patch Live to avoid terminal output during tests
|
||||
with patch("EvoScientist.stream.display.Live"):
|
||||
# First call
|
||||
_run_streaming(
|
||||
agent=mock_agent,
|
||||
message="test message 1",
|
||||
thread_id="thread1",
|
||||
show_thinking=False,
|
||||
interactive=True,
|
||||
)
|
||||
gateway = FakeGraphGateway(stream=mock_stream)
|
||||
|
||||
# Second call - this would fail with "Event loop is closed" before the fix
|
||||
_run_streaming(
|
||||
agent=mock_agent,
|
||||
message="test message 2",
|
||||
thread_id="thread1",
|
||||
show_thinking=False,
|
||||
interactive=True,
|
||||
)
|
||||
# Patch Live to avoid terminal output during tests
|
||||
with patch("EvoScientist.stream.display.Live"):
|
||||
# First call
|
||||
_run_streaming(
|
||||
agent=mock_agent,
|
||||
message="test message 1",
|
||||
thread_id="thread1",
|
||||
show_thinking=False,
|
||||
interactive=True,
|
||||
gateway=gateway,
|
||||
)
|
||||
|
||||
# Third call for good measure
|
||||
_run_streaming(
|
||||
agent=mock_agent,
|
||||
message="test message 3",
|
||||
thread_id="thread1",
|
||||
show_thinking=False,
|
||||
interactive=True,
|
||||
)
|
||||
# Second call - this would fail with "Event loop is closed" before the fix
|
||||
_run_streaming(
|
||||
agent=mock_agent,
|
||||
message="test message 2",
|
||||
thread_id="thread1",
|
||||
show_thinking=False,
|
||||
interactive=True,
|
||||
gateway=gateway,
|
||||
)
|
||||
|
||||
# Third call for good measure
|
||||
_run_streaming(
|
||||
agent=mock_agent,
|
||||
message="test message 3",
|
||||
thread_id="thread1",
|
||||
show_thinking=False,
|
||||
interactive=True,
|
||||
gateway=gateway,
|
||||
)
|
||||
|
||||
def test_loop_reused_across_calls(self):
|
||||
"""Event loop should be reused across multiple calls."""
|
||||
@@ -227,7 +229,7 @@ class TestMultipleStreamingCalls:
|
||||
thinking = "Initial plan. " * 20
|
||||
stream_calls = 0
|
||||
|
||||
async def mock_stream(*args, **kwargs):
|
||||
async def mock_stream(_request):
|
||||
nonlocal stream_calls
|
||||
stream_calls += 1
|
||||
if stream_calls == 1:
|
||||
@@ -245,23 +247,20 @@ class TestMultipleStreamingCalls:
|
||||
|
||||
sent_thinking: list[str] = []
|
||||
|
||||
with patch(
|
||||
"EvoScientist.stream.display.stream_agent_events",
|
||||
side_effect=mock_stream,
|
||||
):
|
||||
with patch("EvoScientist.stream.display.Live"):
|
||||
result = _run_streaming(
|
||||
agent=mock_agent,
|
||||
message="test message",
|
||||
thread_id="thread1",
|
||||
show_thinking=False,
|
||||
interactive=True,
|
||||
on_thinking=sent_thinking.append,
|
||||
ask_user_prompt_fn=lambda _data: {
|
||||
"answers": ["yes"],
|
||||
"status": "answered",
|
||||
},
|
||||
)
|
||||
with patch("EvoScientist.stream.display.Live"):
|
||||
result = _run_streaming(
|
||||
agent=mock_agent,
|
||||
message="test message",
|
||||
thread_id="thread1",
|
||||
show_thinking=False,
|
||||
interactive=True,
|
||||
on_thinking=sent_thinking.append,
|
||||
ask_user_prompt_fn=lambda _data: {
|
||||
"answers": ["yes"],
|
||||
"status": "answered",
|
||||
},
|
||||
gateway=FakeGraphGateway(stream=mock_stream),
|
||||
)
|
||||
|
||||
assert result == "final answer"
|
||||
assert sent_thinking == [thinking.rstrip()]
|
||||
@@ -275,7 +274,7 @@ class TestMultipleStreamingCalls:
|
||||
thinking_r2 = "Revised plan. " * 20
|
||||
stream_calls = 0
|
||||
|
||||
async def mock_stream(*args, **kwargs):
|
||||
async def mock_stream(_request):
|
||||
nonlocal stream_calls
|
||||
stream_calls += 1
|
||||
if stream_calls == 1:
|
||||
@@ -294,23 +293,20 @@ class TestMultipleStreamingCalls:
|
||||
|
||||
sent_thinking: list[str] = []
|
||||
|
||||
with patch(
|
||||
"EvoScientist.stream.display.stream_agent_events",
|
||||
side_effect=mock_stream,
|
||||
):
|
||||
with patch("EvoScientist.stream.display.Live"):
|
||||
result = _run_streaming(
|
||||
agent=mock_agent,
|
||||
message="test message",
|
||||
thread_id="thread1",
|
||||
show_thinking=False,
|
||||
interactive=True,
|
||||
on_thinking=sent_thinking.append,
|
||||
ask_user_prompt_fn=lambda _data: {
|
||||
"answers": ["yes"],
|
||||
"status": "answered",
|
||||
},
|
||||
)
|
||||
with patch("EvoScientist.stream.display.Live"):
|
||||
result = _run_streaming(
|
||||
agent=mock_agent,
|
||||
message="test message",
|
||||
thread_id="thread1",
|
||||
show_thinking=False,
|
||||
interactive=True,
|
||||
on_thinking=sent_thinking.append,
|
||||
ask_user_prompt_fn=lambda _data: {
|
||||
"answers": ["yes"],
|
||||
"status": "answered",
|
||||
},
|
||||
gateway=FakeGraphGateway(stream=mock_stream),
|
||||
)
|
||||
|
||||
assert result == "final answer"
|
||||
assert sent_thinking == [thinking_r1.rstrip(), thinking_r2.rstrip()]
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -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, []))
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -283,14 +283,16 @@ class TestPhaseBMigrated:
|
||||
"""Session lifecycle callbacks (start/resume) filled in Phase B."""
|
||||
|
||||
def test_start_new_session_fires_callback(self):
|
||||
called = []
|
||||
ui, _ = _make_ui(on_start_new_session=lambda: called.append("new"))
|
||||
ui.start_new_session()
|
||||
assert called == ["new"]
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
cb = AsyncMock()
|
||||
ui, _ = _make_ui(on_start_new_session=cb)
|
||||
_run(ui.start_new_session())
|
||||
cb.assert_awaited_once()
|
||||
|
||||
def test_start_new_session_without_callback_is_noop(self):
|
||||
ui, console = _make_ui()
|
||||
ui.start_new_session()
|
||||
_run(ui.start_new_session())
|
||||
console.print.assert_not_called()
|
||||
|
||||
def test_handle_session_resume_awaits_callback(self):
|
||||
|
||||
+242
-161
@@ -6,50 +6,102 @@ subsequent messages, not silently keep the stale one the while-loop
|
||||
captured at startup.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from langgraph.graph.state import CompiledStateGraph
|
||||
|
||||
from EvoScientist.cli.channel import (
|
||||
ChannelMessage,
|
||||
_register_channel_request,
|
||||
)
|
||||
from EvoScientist.cli.commands import (
|
||||
ServeRuntimeState,
|
||||
_make_serve_cmd_completed_hook,
|
||||
_make_serve_handle_session_resume_cb,
|
||||
_make_serve_start_new_session_cb,
|
||||
_serve_process_message,
|
||||
)
|
||||
from EvoScientist.commands.base import ChannelRuntime
|
||||
from EvoScientist.config import EvoScientistConfig
|
||||
from EvoScientist.gateway import RuntimeGateways, ThreadStore
|
||||
from tests.conftest import run_async as _run
|
||||
from tests.fakes import FakeGraphGateway, FakeThreadStore
|
||||
|
||||
|
||||
def test_hook_updates_holder_on_agent_swap():
|
||||
def _agent(name: str = "agent") -> CompiledStateGraph:
|
||||
return MagicMock(name=name, spec=CompiledStateGraph)
|
||||
|
||||
|
||||
def _config() -> EvoScientistConfig:
|
||||
return EvoScientistConfig()
|
||||
|
||||
|
||||
def _thread_store(thread_id: str = "unused") -> ThreadStore:
|
||||
return FakeThreadStore(generated_thread_id=thread_id)
|
||||
|
||||
|
||||
def _runtime_gateways(thread_store: ThreadStore | None = None) -> RuntimeGateways:
|
||||
store = thread_store or _thread_store()
|
||||
|
||||
return RuntimeGateways(
|
||||
thread_store=store,
|
||||
graph_gateway=FakeGraphGateway(thread_store=store),
|
||||
)
|
||||
|
||||
|
||||
def _runtime_state(
|
||||
*,
|
||||
agent: CompiledStateGraph | None = None,
|
||||
thread_id: str = "tid",
|
||||
workspace_dir: str | None = None,
|
||||
config: EvoScientistConfig | None = None,
|
||||
thread_store: ThreadStore | None = None,
|
||||
runtime_gateways: RuntimeGateways | None = None,
|
||||
) -> ServeRuntimeState:
|
||||
store = thread_store or _thread_store()
|
||||
return ServeRuntimeState(
|
||||
agent=agent if agent is not None else _agent(),
|
||||
thread_id=thread_id,
|
||||
workspace_dir=workspace_dir,
|
||||
config=config,
|
||||
runtime_gateways=runtime_gateways or _runtime_gateways(store),
|
||||
)
|
||||
|
||||
|
||||
def test_hook_updates_runtime_state_on_agent_swap():
|
||||
"""``/model`` mutates ``ctx.agent`` to a new handle — the hook must
|
||||
push that handle into the shared holder so the outer poll loop sees
|
||||
push that handle into the shared runtime state so the outer poll loop sees
|
||||
it on the next message."""
|
||||
holder = {"agent": "original-agent"}
|
||||
hook = _make_serve_cmd_completed_hook(holder)
|
||||
original_agent = _agent("original-agent")
|
||||
new_agent = _agent("new-agent")
|
||||
state = _runtime_state(agent=original_agent)
|
||||
hook = _make_serve_cmd_completed_hook(state)
|
||||
|
||||
ctx = MagicMock()
|
||||
ctx.agent = "new-agent"
|
||||
ctx.agent = new_agent
|
||||
ctx.thread_id = state.thread_id
|
||||
cmd = MagicMock()
|
||||
cmd.name = "/model"
|
||||
|
||||
_run(hook(ctx, "original-agent", cmd))
|
||||
_run(hook(ctx, original_agent, cmd))
|
||||
|
||||
assert holder["agent"] == "new-agent"
|
||||
assert state.agent is new_agent
|
||||
|
||||
|
||||
def test_hook_syncs_channel_runtime():
|
||||
"""Other readers (the bus) look at ``ChannelRuntime.agent``; the
|
||||
hook keeps the runtime in sync with the holder update."""
|
||||
holder = {"agent": "original-agent", "thread_id": "t"}
|
||||
runtime = ChannelRuntime(agent="original-agent", thread_id="t")
|
||||
hook = _make_serve_cmd_completed_hook(holder, runtime)
|
||||
hook keeps the runtime in sync with the runtime state update."""
|
||||
original_agent = _agent("original-agent")
|
||||
new_agent = _agent("new-agent")
|
||||
state = _runtime_state(agent=original_agent, thread_id="t")
|
||||
runtime = ChannelRuntime(agent=original_agent, thread_id="t")
|
||||
hook = _make_serve_cmd_completed_hook(state, runtime)
|
||||
|
||||
ctx = MagicMock()
|
||||
ctx.agent = "new-agent"
|
||||
ctx.agent = new_agent
|
||||
# Pin ctx.thread_id explicitly — a bare MagicMock would let the
|
||||
# hook's getattr fall through to a fresh MagicMock attribute and
|
||||
# silently mutate runtime.thread_id, hiding regressions.
|
||||
@@ -57,76 +109,83 @@ def test_hook_syncs_channel_runtime():
|
||||
cmd = MagicMock()
|
||||
cmd.name = "/model"
|
||||
|
||||
_run(hook(ctx, "original-agent", cmd))
|
||||
_run(hook(ctx, original_agent, cmd))
|
||||
|
||||
assert runtime.agent == "new-agent"
|
||||
assert runtime.agent is new_agent
|
||||
assert runtime.thread_id == "t"
|
||||
|
||||
|
||||
def test_hook_noop_when_agent_unchanged():
|
||||
"""Commands like ``/evoskills`` don't touch ``ctx.agent`` — the
|
||||
holder must stay put."""
|
||||
holder = {"agent": "original-agent"}
|
||||
hook = _make_serve_cmd_completed_hook(holder)
|
||||
runtime state must stay put."""
|
||||
original_agent = _agent("original-agent")
|
||||
state = _runtime_state(agent=original_agent)
|
||||
hook = _make_serve_cmd_completed_hook(state)
|
||||
|
||||
ctx = MagicMock()
|
||||
ctx.agent = "original-agent" # no swap
|
||||
ctx.agent = original_agent # no swap
|
||||
ctx.thread_id = state.thread_id
|
||||
cmd = MagicMock()
|
||||
cmd.name = "/evoskills"
|
||||
|
||||
_run(hook(ctx, "original-agent", cmd))
|
||||
_run(hook(ctx, original_agent, cmd))
|
||||
|
||||
assert holder["agent"] == "original-agent"
|
||||
assert state.agent is original_agent
|
||||
|
||||
|
||||
def test_hook_noop_when_ctx_agent_is_none():
|
||||
"""Guard against commands that reset ``ctx.agent`` to ``None`` —
|
||||
we never want to write ``None`` into the holder."""
|
||||
holder = {"agent": "original-agent"}
|
||||
hook = _make_serve_cmd_completed_hook(holder)
|
||||
we never want to write ``None`` into runtime state."""
|
||||
original_agent = _agent("original-agent")
|
||||
state = _runtime_state(agent=original_agent)
|
||||
hook = _make_serve_cmd_completed_hook(state)
|
||||
|
||||
ctx = MagicMock()
|
||||
ctx.agent = None
|
||||
ctx.thread_id = state.thread_id
|
||||
cmd = MagicMock()
|
||||
cmd.name = "/whatever"
|
||||
|
||||
_run(hook(ctx, "original-agent", cmd))
|
||||
_run(hook(ctx, original_agent, cmd))
|
||||
|
||||
assert holder["agent"] == "original-agent"
|
||||
assert state.agent is original_agent
|
||||
|
||||
|
||||
def test_hook_updates_thread_id_on_resume():
|
||||
"""``/resume`` mutates ``ctx.thread_id`` — the hook must push the
|
||||
new id into the holder so the outer poll loop runs subsequent
|
||||
new id into runtime state so the outer poll loop runs subsequent
|
||||
messages on the resumed thread."""
|
||||
holder = {"agent": "a", "thread_id": "original-tid"}
|
||||
hook = _make_serve_cmd_completed_hook(holder)
|
||||
agent = _agent("a")
|
||||
state = _runtime_state(agent=agent, thread_id="original-tid")
|
||||
hook = _make_serve_cmd_completed_hook(state)
|
||||
|
||||
ctx = MagicMock()
|
||||
ctx.agent = "a" # no agent swap
|
||||
ctx.agent = agent # no agent swap
|
||||
ctx.thread_id = "new-tid"
|
||||
ctx.workspace_dir = None
|
||||
cmd = MagicMock()
|
||||
cmd.name = "/resume"
|
||||
|
||||
_run(hook(ctx, "a", cmd))
|
||||
_run(hook(ctx, agent, cmd))
|
||||
|
||||
assert holder["thread_id"] == "new-tid"
|
||||
assert state.thread_id == "new-tid"
|
||||
|
||||
|
||||
def test_hook_updates_workspace_dir_on_resume():
|
||||
"""`/resume` can restore a different workspace; serve must reload for it."""
|
||||
cfg = object()
|
||||
holder = {
|
||||
"agent": "old-agent",
|
||||
"thread_id": "original-tid",
|
||||
"workspace_dir": "/old-ws",
|
||||
"config": cfg,
|
||||
}
|
||||
hook = _make_serve_cmd_completed_hook(holder, config=cfg)
|
||||
cfg = _config()
|
||||
old_agent = _agent("old-agent")
|
||||
reloaded_agent = _agent("reloaded-agent")
|
||||
state = _runtime_state(
|
||||
agent=old_agent,
|
||||
thread_id="original-tid",
|
||||
workspace_dir="/old-ws",
|
||||
config=cfg,
|
||||
)
|
||||
hook = _make_serve_cmd_completed_hook(state, config=cfg)
|
||||
|
||||
ctx = MagicMock()
|
||||
ctx.agent = "old-agent"
|
||||
ctx.agent = old_agent
|
||||
ctx.thread_id = "new-tid"
|
||||
ctx.workspace_dir = "/restored-ws"
|
||||
cmd = MagicMock()
|
||||
@@ -139,67 +198,70 @@ def test_hook_updates_workspace_dir_on_resume():
|
||||
) as sync_server,
|
||||
patch(
|
||||
"EvoScientist.cli.commands._load_agent",
|
||||
return_value="reloaded-agent",
|
||||
return_value=reloaded_agent,
|
||||
) as load_agent,
|
||||
):
|
||||
_run(hook(ctx, "old-agent", cmd))
|
||||
_run(hook(ctx, old_agent, cmd))
|
||||
|
||||
sync_server.assert_awaited_once_with(cfg, workspace_dir="/restored-ws")
|
||||
load_agent.assert_called_once_with(workspace_dir="/restored-ws", config=cfg)
|
||||
assert holder["workspace_dir"] == "/restored-ws"
|
||||
assert holder["agent"] == "reloaded-agent"
|
||||
assert state.workspace_dir == "/restored-ws"
|
||||
assert state.agent is reloaded_agent
|
||||
|
||||
|
||||
def test_hook_syncs_channel_runtime_thread_id():
|
||||
"""The bus reads ``ChannelRuntime.thread_id``; hook must sync it
|
||||
alongside the holder update."""
|
||||
holder = {"agent": "a", "thread_id": "original-tid"}
|
||||
runtime = ChannelRuntime(agent="a", thread_id="original-tid")
|
||||
hook = _make_serve_cmd_completed_hook(holder, runtime)
|
||||
alongside the runtime state update."""
|
||||
agent = _agent("a")
|
||||
state = _runtime_state(agent=agent, thread_id="original-tid")
|
||||
runtime = ChannelRuntime(agent=agent, thread_id="original-tid")
|
||||
hook = _make_serve_cmd_completed_hook(state, runtime)
|
||||
|
||||
ctx = MagicMock()
|
||||
ctx.agent = "a"
|
||||
ctx.agent = agent
|
||||
ctx.thread_id = "new-tid"
|
||||
ctx.workspace_dir = None
|
||||
cmd = MagicMock()
|
||||
cmd.name = "/resume"
|
||||
|
||||
_run(hook(ctx, "a", cmd))
|
||||
_run(hook(ctx, agent, cmd))
|
||||
|
||||
assert runtime.thread_id == "new-tid"
|
||||
|
||||
|
||||
def test_hook_noop_when_thread_id_unchanged():
|
||||
"""Most commands don't touch thread_id — holder stays put."""
|
||||
holder = {"agent": "a", "thread_id": "same-tid"}
|
||||
hook = _make_serve_cmd_completed_hook(holder)
|
||||
"""Most commands don't touch thread_id — runtime state stays put."""
|
||||
agent = _agent("a")
|
||||
state = _runtime_state(agent=agent, thread_id="same-tid")
|
||||
hook = _make_serve_cmd_completed_hook(state)
|
||||
|
||||
ctx = MagicMock()
|
||||
ctx.agent = "a"
|
||||
ctx.agent = agent
|
||||
ctx.thread_id = "same-tid"
|
||||
cmd = MagicMock()
|
||||
cmd.name = "/evoskills"
|
||||
|
||||
_run(hook(ctx, "a", cmd))
|
||||
_run(hook(ctx, agent, cmd))
|
||||
|
||||
assert holder["thread_id"] == "same-tid"
|
||||
assert state.thread_id == "same-tid"
|
||||
|
||||
|
||||
def test_hook_skips_resume_warning_when_thread_unchanged():
|
||||
"""Bare ``/resume`` with no argument prints usage but leaves
|
||||
``ctx.thread_id`` unchanged — the in-memory-state warning must NOT
|
||||
fire because no resume actually happened."""
|
||||
holder = {"agent": "a", "thread_id": "original-tid"}
|
||||
hook = _make_serve_cmd_completed_hook(holder)
|
||||
agent = _agent("a")
|
||||
state = _runtime_state(agent=agent, thread_id="original-tid")
|
||||
hook = _make_serve_cmd_completed_hook(state)
|
||||
|
||||
ctx = MagicMock()
|
||||
ctx.agent = "a"
|
||||
ctx.agent = agent
|
||||
ctx.thread_id = "original-tid" # unchanged — bare /resume case
|
||||
ctx.workspace_dir = None
|
||||
cmd = MagicMock()
|
||||
cmd.name = "/resume"
|
||||
|
||||
_run(hook(ctx, "a", cmd))
|
||||
_run(hook(ctx, agent, cmd))
|
||||
|
||||
ctx.ui.append_system.assert_not_called()
|
||||
ctx.ui.flush.assert_not_called()
|
||||
@@ -208,19 +270,20 @@ def test_hook_skips_resume_warning_when_thread_unchanged():
|
||||
def test_hook_emits_resume_warning_when_thread_changed():
|
||||
"""``/resume <tid>`` that actually changes thread_id must surface
|
||||
the in-memory-state warning via ``ctx.ui``."""
|
||||
holder = {"agent": "a", "thread_id": "original-tid"}
|
||||
hook = _make_serve_cmd_completed_hook(holder)
|
||||
agent = _agent("a")
|
||||
state = _runtime_state(agent=agent, thread_id="original-tid")
|
||||
hook = _make_serve_cmd_completed_hook(state)
|
||||
|
||||
ctx = MagicMock()
|
||||
# Mock out async flush so the test can synchronously run the hook.
|
||||
ctx.ui.flush = AsyncMock()
|
||||
ctx.agent = "a"
|
||||
ctx.agent = agent
|
||||
ctx.thread_id = "abc12345-resumed-tid"
|
||||
ctx.workspace_dir = None
|
||||
cmd = MagicMock()
|
||||
cmd.name = "/resume"
|
||||
|
||||
_run(hook(ctx, "a", cmd))
|
||||
_run(hook(ctx, agent, cmd))
|
||||
|
||||
ctx.ui.append_system.assert_called_once()
|
||||
warn_text, warn_kwargs = (
|
||||
@@ -235,51 +298,58 @@ def test_hook_emits_resume_warning_when_thread_changed():
|
||||
|
||||
def test_start_new_session_cb_rotates_thread_id():
|
||||
"""``/new`` via channel calls this callback — must generate a new
|
||||
thread id, push into holder, and sync the channel runtime."""
|
||||
holder = {"agent": "a", "thread_id": "old-tid"}
|
||||
runtime = ChannelRuntime(agent="a", thread_id="old-tid")
|
||||
thread id, push into runtime state, and sync the channel runtime."""
|
||||
agent = _agent("a")
|
||||
state = _runtime_state(
|
||||
agent=agent,
|
||||
thread_id="old-tid",
|
||||
thread_store=_thread_store("freshly-generated-tid"),
|
||||
)
|
||||
runtime = ChannelRuntime(agent=agent, thread_id="old-tid")
|
||||
|
||||
with patch(
|
||||
"EvoScientist.sessions.generate_thread_id",
|
||||
return_value="freshly-generated-tid",
|
||||
):
|
||||
cb = _make_serve_start_new_session_cb(holder, runtime)
|
||||
cb()
|
||||
cb = _make_serve_start_new_session_cb(
|
||||
state,
|
||||
runtime,
|
||||
)
|
||||
_run(cb())
|
||||
|
||||
assert holder["thread_id"] == "freshly-generated-tid"
|
||||
assert state.thread_id == "freshly-generated-tid"
|
||||
assert runtime.thread_id == "freshly-generated-tid"
|
||||
|
||||
|
||||
def test_start_new_session_cb_leaves_agent_alone():
|
||||
"""``/new`` rotates thread only — agent handle must stay put
|
||||
(serve's agent is a single pre-loaded instance, not per-thread)."""
|
||||
holder = {"agent": "a", "thread_id": "old-tid"}
|
||||
agent = _agent("a")
|
||||
state = _runtime_state(
|
||||
agent=agent,
|
||||
thread_id="old-tid",
|
||||
thread_store=_thread_store("new-tid"),
|
||||
)
|
||||
|
||||
with patch(
|
||||
"EvoScientist.sessions.generate_thread_id",
|
||||
return_value="new-tid",
|
||||
):
|
||||
cb = _make_serve_start_new_session_cb(holder)
|
||||
cb()
|
||||
cb = _make_serve_start_new_session_cb(state)
|
||||
_run(cb())
|
||||
|
||||
assert holder["agent"] == "a"
|
||||
assert state.agent is agent
|
||||
|
||||
|
||||
def test_serve_resume_callback_syncs_reloads_and_adopts_workspace():
|
||||
cfg = object()
|
||||
holder = {
|
||||
"agent": "old-agent",
|
||||
"thread_id": "old-tid",
|
||||
"workspace_dir": "/old-ws",
|
||||
"config": cfg,
|
||||
}
|
||||
runtime = ChannelRuntime(agent="old-agent", thread_id="old-tid")
|
||||
cb = _make_serve_handle_session_resume_cb(holder, runtime, config=cfg)
|
||||
cfg = _config()
|
||||
old_agent = _agent("old-agent")
|
||||
reloaded_agent = _agent("reloaded-agent")
|
||||
state = _runtime_state(
|
||||
agent=old_agent,
|
||||
thread_id="old-tid",
|
||||
workspace_dir="/old-ws",
|
||||
config=cfg,
|
||||
)
|
||||
runtime = ChannelRuntime(agent=old_agent, thread_id="old-tid")
|
||||
cb = _make_serve_handle_session_resume_cb(state, runtime, config=cfg)
|
||||
call_order: list[str] = []
|
||||
|
||||
def _load_agent(**_kwargs):
|
||||
call_order.append("load")
|
||||
return "reloaded-agent"
|
||||
return reloaded_agent
|
||||
|
||||
async def _sync_server(*_args, **_kwargs):
|
||||
call_order.append("sync")
|
||||
@@ -299,23 +369,25 @@ def test_serve_resume_callback_syncs_reloads_and_adopts_workspace():
|
||||
sync_server.assert_awaited_once_with(cfg, workspace_dir="/new-ws")
|
||||
load_agent.assert_called_once_with(workspace_dir="/new-ws", config=cfg)
|
||||
assert call_order == ["load", "sync"]
|
||||
assert holder["thread_id"] == "new-tid"
|
||||
assert holder["workspace_dir"] == "/new-ws"
|
||||
assert holder["agent"] == "reloaded-agent"
|
||||
assert state.thread_id == "new-tid"
|
||||
assert state.workspace_dir == "/new-ws"
|
||||
assert state.agent is reloaded_agent
|
||||
assert runtime.thread_id == "new-tid"
|
||||
assert runtime.agent == "reloaded-agent"
|
||||
assert runtime.agent is reloaded_agent
|
||||
|
||||
|
||||
def test_hook_emits_resume_warning_after_resume_callback_adopts_thread():
|
||||
cfg = object()
|
||||
holder = {
|
||||
"agent": "old-agent",
|
||||
"thread_id": "old-tid",
|
||||
"workspace_dir": "/old-ws",
|
||||
"config": cfg,
|
||||
}
|
||||
runtime = ChannelRuntime(agent="old-agent", thread_id="old-tid")
|
||||
cb = _make_serve_handle_session_resume_cb(holder, runtime, config=cfg)
|
||||
cfg = _config()
|
||||
old_agent = _agent("old-agent")
|
||||
reloaded_agent = _agent("reloaded-agent")
|
||||
state = _runtime_state(
|
||||
agent=old_agent,
|
||||
thread_id="old-tid",
|
||||
workspace_dir="/old-ws",
|
||||
config=cfg,
|
||||
)
|
||||
runtime = ChannelRuntime(agent=old_agent, thread_id="old-tid")
|
||||
cb = _make_serve_handle_session_resume_cb(state, runtime, config=cfg)
|
||||
|
||||
with (
|
||||
patch(
|
||||
@@ -324,21 +396,21 @@ def test_hook_emits_resume_warning_after_resume_callback_adopts_thread():
|
||||
),
|
||||
patch(
|
||||
"EvoScientist.cli.commands._load_agent",
|
||||
return_value="reloaded-agent",
|
||||
return_value=reloaded_agent,
|
||||
),
|
||||
):
|
||||
_run(cb("abc12345-resumed-tid", "/new-ws"))
|
||||
|
||||
hook = _make_serve_cmd_completed_hook(holder, runtime, config=cfg)
|
||||
hook = _make_serve_cmd_completed_hook(state, runtime, config=cfg)
|
||||
ctx = MagicMock()
|
||||
ctx.ui.flush = AsyncMock()
|
||||
ctx.agent = "reloaded-agent"
|
||||
ctx.agent = reloaded_agent
|
||||
ctx.thread_id = "abc12345-resumed-tid"
|
||||
ctx.workspace_dir = "/new-ws"
|
||||
cmd = MagicMock()
|
||||
cmd.name = "/resume"
|
||||
|
||||
_run(hook(ctx, "reloaded-agent", cmd))
|
||||
_run(hook(ctx, reloaded_agent, cmd))
|
||||
|
||||
ctx.ui.append_system.assert_called_once()
|
||||
assert "in-memory state" in ctx.ui.append_system.call_args.args[0]
|
||||
@@ -346,15 +418,17 @@ def test_hook_emits_resume_warning_after_resume_callback_adopts_thread():
|
||||
|
||||
|
||||
def test_serve_resume_callback_preserves_state_when_sync_fails():
|
||||
cfg = object()
|
||||
holder = {
|
||||
"agent": "old-agent",
|
||||
"thread_id": "old-tid",
|
||||
"workspace_dir": "/old-ws",
|
||||
"config": cfg,
|
||||
}
|
||||
runtime = ChannelRuntime(agent="old-agent", thread_id="old-tid")
|
||||
cb = _make_serve_handle_session_resume_cb(holder, runtime, config=cfg)
|
||||
cfg = _config()
|
||||
old_agent = _agent("old-agent")
|
||||
loaded_but_not_adopted = _agent("loaded-but-not-adopted")
|
||||
state = _runtime_state(
|
||||
agent=old_agent,
|
||||
thread_id="old-tid",
|
||||
workspace_dir="/old-ws",
|
||||
config=cfg,
|
||||
)
|
||||
runtime = ChannelRuntime(agent=old_agent, thread_id="old-tid")
|
||||
cb = _make_serve_handle_session_resume_cb(state, runtime, config=cfg)
|
||||
|
||||
with (
|
||||
patch(
|
||||
@@ -363,7 +437,7 @@ def test_serve_resume_callback_preserves_state_when_sync_fails():
|
||||
),
|
||||
patch(
|
||||
"EvoScientist.cli.commands._load_agent",
|
||||
return_value="loaded-but-not-adopted",
|
||||
return_value=loaded_but_not_adopted,
|
||||
) as load_agent,
|
||||
patch("EvoScientist.cli.commands.set_active_workspace") as set_active,
|
||||
pytest.raises(RuntimeError, match="workspace conflict"),
|
||||
@@ -372,28 +446,26 @@ def test_serve_resume_callback_preserves_state_when_sync_fails():
|
||||
|
||||
load_agent.assert_called_once_with(workspace_dir="/new-ws", config=cfg)
|
||||
set_active.assert_called_once_with("/old-ws")
|
||||
assert "loaded-but-not-adopted" not in holder.values()
|
||||
assert "_resume_warning_thread_id" not in holder
|
||||
assert holder == {
|
||||
"agent": "old-agent",
|
||||
"thread_id": "old-tid",
|
||||
"workspace_dir": "/old-ws",
|
||||
"config": cfg,
|
||||
}
|
||||
assert runtime.agent == "old-agent"
|
||||
assert state.agent is old_agent
|
||||
assert state.resume_warning_thread_id is None
|
||||
assert state.thread_id == "old-tid"
|
||||
assert state.workspace_dir == "/old-ws"
|
||||
assert state.config is cfg
|
||||
assert runtime.agent is old_agent
|
||||
assert runtime.thread_id == "old-tid"
|
||||
|
||||
|
||||
def test_serve_resume_callback_load_failure_does_not_sync_or_adopt():
|
||||
cfg = object()
|
||||
holder = {
|
||||
"agent": "old-agent",
|
||||
"thread_id": "old-tid",
|
||||
"workspace_dir": "/old-ws",
|
||||
"config": cfg,
|
||||
}
|
||||
runtime = ChannelRuntime(agent="old-agent", thread_id="old-tid")
|
||||
cb = _make_serve_handle_session_resume_cb(holder, runtime, config=cfg)
|
||||
cfg = _config()
|
||||
old_agent = _agent("old-agent")
|
||||
state = _runtime_state(
|
||||
agent=old_agent,
|
||||
thread_id="old-tid",
|
||||
workspace_dir="/old-ws",
|
||||
config=cfg,
|
||||
)
|
||||
runtime = ChannelRuntime(agent=old_agent, thread_id="old-tid")
|
||||
cb = _make_serve_handle_session_resume_cb(state, runtime, config=cfg)
|
||||
|
||||
with (
|
||||
patch(
|
||||
@@ -412,32 +484,32 @@ def test_serve_resume_callback_load_failure_does_not_sync_or_adopt():
|
||||
load_agent.assert_called_once_with(workspace_dir="/new-ws", config=cfg)
|
||||
set_active.assert_called_once_with("/old-ws")
|
||||
sync_server.assert_not_awaited()
|
||||
assert "_resume_warning_thread_id" not in holder
|
||||
assert holder == {
|
||||
"agent": "old-agent",
|
||||
"thread_id": "old-tid",
|
||||
"workspace_dir": "/old-ws",
|
||||
"config": cfg,
|
||||
}
|
||||
assert runtime.agent == "old-agent"
|
||||
assert state.resume_warning_thread_id is None
|
||||
assert state.agent is old_agent
|
||||
assert state.thread_id == "old-tid"
|
||||
assert state.workspace_dir == "/old-ws"
|
||||
assert state.config is cfg
|
||||
assert runtime.agent is old_agent
|
||||
assert runtime.thread_id == "old-tid"
|
||||
|
||||
|
||||
def test_hook_handles_both_agent_and_thread_swap():
|
||||
"""Edge case: a command that changes both (hypothetical). Both
|
||||
updates must land in the holder."""
|
||||
holder = {"agent": "old-agent", "thread_id": "old-tid"}
|
||||
hook = _make_serve_cmd_completed_hook(holder)
|
||||
updates must land in runtime state."""
|
||||
old_agent = _agent("old-agent")
|
||||
new_agent = _agent("new-agent")
|
||||
state = _runtime_state(agent=old_agent, thread_id="old-tid")
|
||||
hook = _make_serve_cmd_completed_hook(state)
|
||||
|
||||
ctx = MagicMock()
|
||||
ctx.agent = "new-agent"
|
||||
ctx.agent = new_agent
|
||||
ctx.thread_id = "new-tid"
|
||||
cmd = MagicMock()
|
||||
|
||||
_run(hook(ctx, "old-agent", cmd))
|
||||
_run(hook(ctx, old_agent, cmd))
|
||||
|
||||
assert holder["agent"] == "new-agent"
|
||||
assert holder["thread_id"] == "new-tid"
|
||||
assert state.agent is new_agent
|
||||
assert state.thread_id == "new-tid"
|
||||
|
||||
|
||||
def test_serve_process_message_reports_slash_dispatch_error_without_fallback():
|
||||
@@ -456,7 +528,13 @@ def test_serve_process_message_reports_slash_dispatch_error_without_fallback():
|
||||
chat_id="channel-user",
|
||||
message_id="ts-1",
|
||||
)
|
||||
holder = {"agent": "agent", "thread_id": "tid"}
|
||||
thread_store = _thread_store()
|
||||
state = _runtime_state(
|
||||
agent=_agent(),
|
||||
thread_id="tid",
|
||||
thread_store=thread_store,
|
||||
runtime_gateways=_runtime_gateways(thread_store),
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
@@ -469,7 +547,7 @@ def test_serve_process_message_reports_slash_dispatch_error_without_fallback():
|
||||
_register_channel_request(msg)
|
||||
_serve_process_message(
|
||||
msg,
|
||||
agent_holder=holder,
|
||||
runtime_state=state,
|
||||
model="model",
|
||||
workspace_dir="/tmp",
|
||||
show_thinking=False,
|
||||
@@ -479,7 +557,7 @@ def test_serve_process_message_reports_slash_dispatch_error_without_fallback():
|
||||
mock_run_streaming.assert_not_called()
|
||||
|
||||
|
||||
def test_serve_process_message_uses_runtime_workspace_from_holder():
|
||||
def test_serve_process_message_uses_runtime_workspace_from_state():
|
||||
"""After `/resume`, serve should use the adopted workspace, not startup ws."""
|
||||
msg = ChannelMessage(
|
||||
msg_id="msg-2",
|
||||
@@ -492,11 +570,14 @@ def test_serve_process_message_uses_runtime_workspace_from_holder():
|
||||
chat_id="channel-user",
|
||||
message_id="ts-2",
|
||||
)
|
||||
holder = {
|
||||
"agent": "agent",
|
||||
"thread_id": "tid",
|
||||
"workspace_dir": "/restored-workspace",
|
||||
}
|
||||
thread_store = _thread_store()
|
||||
state = _runtime_state(
|
||||
agent=_agent(),
|
||||
thread_id="tid",
|
||||
workspace_dir="/restored-workspace",
|
||||
thread_store=thread_store,
|
||||
runtime_gateways=_runtime_gateways(thread_store),
|
||||
)
|
||||
captured: dict[str, str] = {}
|
||||
|
||||
async def _fake_dispatch(*args, **kwargs):
|
||||
@@ -521,7 +602,7 @@ def test_serve_process_message_uses_runtime_workspace_from_holder():
|
||||
_register_channel_request(msg)
|
||||
_serve_process_message(
|
||||
msg,
|
||||
agent_holder=holder,
|
||||
runtime_state=state,
|
||||
model="model",
|
||||
workspace_dir="/startup-workspace",
|
||||
show_thinking=False,
|
||||
|
||||
+85
-30
@@ -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
|
||||
|
||||
|
||||
@@ -21,6 +21,7 @@ from EvoScientist.cli.status_bar import (
|
||||
status_style_name,
|
||||
trim_status_text,
|
||||
)
|
||||
from tests.fakes import FakeGraphGateway, FakeThreadStore
|
||||
|
||||
|
||||
def _render_fragments(fragments: list[tuple[str, str]]) -> str:
|
||||
@@ -208,10 +209,6 @@ def test_build_status_text_uses_rich_styles():
|
||||
|
||||
|
||||
def test_build_session_status_snapshot_uses_fallback_window(monkeypatch):
|
||||
async def _fake_messages(thread_id: str):
|
||||
assert thread_id == "thread-1"
|
||||
return [HumanMessage(content="existing")]
|
||||
|
||||
class _FakeModel:
|
||||
model_name: ClassVar[str] = "provider/demo-model"
|
||||
profile: ClassVar[dict[str, object]] = {}
|
||||
@@ -221,10 +218,6 @@ def test_build_session_status_snapshot_uses_fallback_window(monkeypatch):
|
||||
assert messages[-1].content == "pending"
|
||||
return 42_000
|
||||
|
||||
monkeypatch.setattr(
|
||||
"EvoScientist.cli.status_bar.get_thread_messages",
|
||||
_fake_messages,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"EvoScientist.cli.status_bar._get_default_chat_model",
|
||||
lambda: _FakeModel(),
|
||||
@@ -238,6 +231,11 @@ def test_build_session_status_snapshot_uses_fallback_window(monkeypatch):
|
||||
build_session_status_snapshot(
|
||||
"thread-1",
|
||||
pending_user_text="pending",
|
||||
graph_gateway=FakeGraphGateway(
|
||||
thread_store=FakeThreadStore(
|
||||
messages=[HumanMessage(content="existing")]
|
||||
)
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
+44
-52
@@ -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
|
||||
|
||||
@@ -360,6 +360,41 @@ class TestV3ProtocolStreaming:
|
||||
summary_events = [e for e in events if e.get("type") == "summarization"]
|
||||
assert summary_events == []
|
||||
|
||||
def test_direct_stream_loads_existing_summarization_event_when_omitted(self):
|
||||
"""Public stream_agent_events() suppresses persisted summary replays."""
|
||||
summary_message = HumanMessage(
|
||||
content="Here is a summary of the conversation to date:\n\nKey facts",
|
||||
)
|
||||
summary_event = {
|
||||
"_summarization_event": {
|
||||
"summary_message": summary_message,
|
||||
"cutoff_index": 12,
|
||||
"file_path": None,
|
||||
}
|
||||
}
|
||||
agent = FakeV3Agent(
|
||||
[
|
||||
protocol_event("updates", summary_event),
|
||||
message_delta("real content"),
|
||||
],
|
||||
state_values=summary_event,
|
||||
)
|
||||
|
||||
async def _collect():
|
||||
events = []
|
||||
async for event in stream_agent_events(agent, "hi", "t1"):
|
||||
events.append(event)
|
||||
return events
|
||||
|
||||
events = run_async(_collect())
|
||||
|
||||
summary_start_events = [
|
||||
e for e in events if e.get("type") == "summarization_start"
|
||||
]
|
||||
assert summary_start_events == []
|
||||
summary_events = [e for e in events if e.get("type") == "summarization"]
|
||||
assert summary_events == []
|
||||
|
||||
def test_whole_message_reasoning_is_not_duplicated(self):
|
||||
"""Providers can expose the same reasoning in kwargs and content blocks."""
|
||||
message = AIMessage(
|
||||
@@ -432,7 +467,9 @@ class TestV3ProtocolStreaming:
|
||||
return [
|
||||
event
|
||||
async for event in stream_agent_events(
|
||||
agent, "run probe", "live-deepagents-tool-id"
|
||||
agent,
|
||||
"run probe",
|
||||
"live-deepagents-tool-id",
|
||||
)
|
||||
]
|
||||
|
||||
@@ -493,7 +530,9 @@ class TestV3ProtocolStreaming:
|
||||
return [
|
||||
event
|
||||
async for event in stream_agent_events(
|
||||
agent, "run echo", "live-deepagents-hitl"
|
||||
agent,
|
||||
"run echo",
|
||||
"live-deepagents-hitl",
|
||||
)
|
||||
]
|
||||
|
||||
@@ -554,7 +593,9 @@ class TestV3ProtocolStreaming:
|
||||
return [
|
||||
event
|
||||
async for event in stream_agent_events(
|
||||
agent, message, "live-deepagents-ask-user"
|
||||
agent,
|
||||
message,
|
||||
"live-deepagents-ask-user",
|
||||
)
|
||||
]
|
||||
|
||||
@@ -625,7 +666,9 @@ class TestV3ProtocolStreaming:
|
||||
return [
|
||||
event
|
||||
async for event in stream_agent_events(
|
||||
agent, "delegate", "live-deepagents-subagent"
|
||||
agent,
|
||||
"delegate",
|
||||
"live-deepagents-subagent",
|
||||
)
|
||||
]
|
||||
|
||||
@@ -983,7 +1026,11 @@ class TestV3ProtocolStreaming:
|
||||
|
||||
async def consume_one_and_close():
|
||||
agent = HangingV3Agent([message_delta("hi")])
|
||||
stream = stream_agent_events(agent, "hi", "t1")
|
||||
stream = stream_agent_events(
|
||||
agent,
|
||||
"hi",
|
||||
"t1",
|
||||
)
|
||||
first = await stream.__anext__()
|
||||
await stream.aclose()
|
||||
return first, agent.aborted
|
||||
|
||||
+173
-241
@@ -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())
|
||||
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user