fix(channel): scope stop and restore resume history (#186)

* fix(channel): scope stop and restore resume history

* refactor(channel): simplify stop and resume patch

* Delete PR_MESSAGE.md

* fix(channel): address review feedback

* fix(channel): address remaining review bugs

* fix(channel): clean up stopped request handling

* fix(channel): preserve resolved replies and sync tui commands

---------

Co-authored-by: Xi Zhang <106144707+X-iZhang@users.noreply.github.com>
This commit is contained in:
Ziheng Zhang
2026-04-25 23:12:49 +08:00
committed by GitHub
parent 27fee3256e
commit da74c325d6
14 changed files with 1616 additions and 408 deletions
+192
View File
@@ -59,6 +59,10 @@ _response_lock = threading.Lock()
_RESPONSE_TIMEOUT = 600.0
_LATE_RESPONSE_TIMEOUT = 86400.0
_LATE_RESPONSE_NOTICE = "Still working on it. I'll send the result when it's ready."
_channel_request_lock = threading.Lock()
_channel_requests: dict[str, dict[str, str]] = {}
_session_requests: dict[str, list[str]] = {}
_cancelled_channel_messages: set[str] = set()
def _enqueue_channel_message(msg: ChannelMessage) -> asyncio.Future[str]:
@@ -71,6 +75,7 @@ def _enqueue_channel_message(msg: ChannelMessage) -> asyncio.Future[str]:
"loop": loop,
"response": None,
}
_register_channel_request(msg)
_message_queue.put(msg)
return future
@@ -106,6 +111,128 @@ def _pop_channel_response(msg_id: str, *, cancel_pending: bool = False) -> str |
return slot["response"]
def _channel_session_key(channel_type: str, chat_id: str) -> str:
return f"{channel_type}:{chat_id}"
def _channel_message_session_key(msg: ChannelMessage) -> str:
return _channel_session_key(msg.channel_type, msg.chat_id)
def _channel_message_cancel_scope(msg: ChannelMessage) -> str:
return f"channel:{msg.channel_type}:{msg.chat_id}:{msg.msg_id}"
def _register_channel_request(msg: ChannelMessage) -> None:
"""Track a queued channel request so `/stop` can find it later."""
session_key = _channel_message_session_key(msg)
with _channel_request_lock:
_channel_requests[msg.msg_id] = {
"session_key": session_key,
"cancel_scope": _channel_message_cancel_scope(msg),
"state": "queued",
}
_session_requests.setdefault(session_key, []).append(msg.msg_id)
def _claim_channel_request(msg: ChannelMessage) -> bool:
"""Mark a queued request active. Returns False if it was cancelled first."""
with _channel_request_lock:
slot = _channel_requests.get(msg.msg_id)
if slot is None or msg.msg_id in _cancelled_channel_messages:
return False
slot["state"] = "active"
return True
def _claim_or_complete_channel_request(msg: ChannelMessage) -> bool:
"""Claim a request, or clean it up if `/stop` cancelled it while queued."""
if _claim_channel_request(msg):
return True
_complete_channel_request(msg.msg_id)
return False
def _channel_request_state(msg_id: str) -> str | None:
with _channel_request_lock:
slot = _channel_requests.get(msg_id)
return slot.get("state") if slot is not None else None
def _complete_channel_request(
msg_id: str,
*,
discard_cancel_scope: bool = True,
) -> None:
"""Forget a request once its waiter is resolved or cancelled."""
with _channel_request_lock:
slot = _channel_requests.pop(msg_id, None)
_cancelled_channel_messages.discard(msg_id)
if slot is not None:
request_ids = _session_requests.get(slot["session_key"])
if request_ids:
try:
request_ids.remove(msg_id)
except ValueError:
pass
if not request_ids:
_session_requests.pop(slot["session_key"], None)
if slot is not None and discard_cancel_scope:
from ..stream.display import discard_stream_cancel
discard_stream_cancel(slot["cancel_scope"])
def _cancel_channel_session(channel_type: str, chat_id: str) -> tuple[int, int]:
"""Cancel queued and active work for one channel chat session."""
session_key = _channel_session_key(channel_type, chat_id)
with _channel_request_lock:
request_ids: list[str] = []
cancelled_ids: list[str] = []
active_scopes: list[str] = []
with _response_lock:
for msg_id in tuple(_session_requests.get(session_key, ())):
request_slot = _channel_requests.get(msg_id)
if request_slot is None:
continue
response_slot = _pending_responses.get(msg_id)
response_resolved = False
if response_slot is not None:
future = response_slot["future"]
# Once a response is already resolved, leave the slot alone
# so the bus waiter can still publish it instead of falling
# back to "No response".
response_resolved = (
response_slot.get("response") is not None or future.done()
)
if not response_resolved:
request_ids.append(msg_id)
should_cancel = False
if response_slot is None:
should_cancel = request_slot.get("state") == "active"
else:
should_cancel = not response_resolved
if should_cancel:
cancelled_ids.append(msg_id)
if request_slot.get("state") == "active" and should_cancel:
active_scopes.append(request_slot["cancel_scope"])
_cancelled_channel_messages.update(cancelled_ids)
for msg_id in request_ids:
_pop_channel_response(msg_id, cancel_pending=True)
if active_scopes:
from ..stream.display import request_stream_cancel
for cancel_scope in active_scopes:
request_stream_cancel(cancel_scope)
return len(request_ids), len(active_scopes)
# ---------------------------------------------------------------------------
# Slash command dispatch for channel messages
# ---------------------------------------------------------------------------
@@ -312,6 +439,12 @@ _HITL_APPROVAL_TIMEOUT = 120.0 # seconds to wait for HITL approval reply
_ASK_USER_TIMEOUT = (
300.0 # seconds to wait for ask_user reply (longer for thinking time)
)
_STOP_COMMANDS = frozenset(("/stop", "/cancel"))
def _is_stop_command(content: str | None) -> bool:
"""Whether incoming content is a stop/cancel slash command."""
return (content or "").strip().lower() in _STOP_COMMANDS
def _register_hitl_wait(channel_type: str, chat_id: str) -> threading.Event:
@@ -434,6 +567,8 @@ def channel_ask_user_prompt(
return {"status": "cancelled"}
raw = reply_text.strip()
if _is_stop_command(raw):
return {"status": "cancelled"}
if raw.lower() == "cancel":
return {"status": "cancelled"}
@@ -451,6 +586,8 @@ def channel_ask_user_prompt(
if not replied or not other_text:
_send("\u23f0 Response timed out.")
return {"status": "cancelled"}
if _is_stop_command(other_text):
return {"status": "cancelled"}
if other_text.strip().lower() == "cancel":
return {"status": "cancelled"}
answers.append(other_text.strip())
@@ -528,6 +665,12 @@ def channel_hitl_prompt(
_send("\u23f0 Approval timed out. Action rejected.")
return None
if _is_stop_command(reply_text):
# `/stop` already got its own immediate ack from the bus fast-path.
# Treat it as a pure cancel signal here so we don't send a second,
# contradictory "Unrecognized reply" message.
return None
# 3. Parse decision
decision = _parse_approval_reply(reply_text)
if decision == "auto":
@@ -724,6 +867,20 @@ async def _bus_inbound_consumer(bus, manager) -> None:
except asyncio.CancelledError:
break
# /stop should preempt HITL interception so cancel works while
# waiting for approvals/questions. If a HITL wait is pending,
# still release it so the blocking prompt can unwind immediately.
if _is_stop_command(msg.content):
if _try_set_hitl_reply(msg.channel, msg.chat_id, msg.content):
_channel_logger.info(
f"[bus] stop request released HITL wait for "
f"{msg.channel}:{msg.chat_id}"
)
_task = asyncio.create_task(_handle_bus_message(bus, manager, msg))
_tasks.add(_task)
_task.add_done_callback(_tasks.discard)
continue
# Check if this message is a HITL approval reply
if _try_set_hitl_reply(msg.channel, msg.chat_id, msg.content):
_channel_logger.info(
@@ -752,6 +909,37 @@ async def _handle_bus_message(bus, manager, msg) -> None:
)
manager.record_message(msg.channel, "received")
# Fast-path: /stop intercept. Handle on the bus task itself so we
# don't deadlock behind the main-thread stream we're trying to
# interrupt. No typing indicator, no queue entry.
if _is_stop_command(msg.content):
cancelled_count, active_count = _cancel_channel_session(
msg.channel, msg.chat_id
)
try:
await bus.publish_outbound(
OutboundMessage(
channel=msg.channel,
chat_id=msg.chat_id,
content="Stopped.",
reply_to=msg.message_id or None,
metadata=msg.metadata,
)
)
manager.record_message(msg.channel, "sent")
except Exception as e:
_channel_logger.error(f"[bus] /stop ack send error: {e}")
else:
if cancelled_count or active_count:
_channel_logger.info(
"[bus] /stop cancelled %d request(s) (%d active) for %s:%s",
cancelled_count,
active_count,
msg.channel,
msg.chat_id,
)
return
channel = manager.get_channel(msg.channel)
typing_active = False
if channel:
@@ -821,6 +1009,8 @@ async def _handle_bus_message(bus, manager, msg) -> None:
f"for {cm.msg_id}"
)
_pop_channel_response(cm.msg_id, cancel_pending=True)
if _channel_request_state(cm.msg_id) != "active":
_complete_channel_request(cm.msg_id)
return
response = _pop_channel_response(cm.msg_id) or "No response"
@@ -836,6 +1026,8 @@ async def _handle_bus_message(bus, manager, msg) -> None:
manager.record_message(msg.channel, "sent")
except asyncio.CancelledError:
_pop_channel_response(cm.msg_id, cancel_pending=True)
if _channel_request_state(cm.msg_id) != "active":
_complete_channel_request(cm.msg_id)
raise
except Exception as e:
_channel_logger.error(f"[bus] Outbound error: {e}")
+93 -70
View File
@@ -27,7 +27,10 @@ from .agent import (
)
from .channel import (
ChannelMessage,
_channel_message_cancel_scope,
_channels_stop,
_claim_or_complete_channel_request,
_complete_channel_request,
_message_queue,
_set_channel_response,
_start_channels_bus_mode,
@@ -530,9 +533,9 @@ def _make_serve_start_new_session_cb(agent_holder: dict[str, Any]):
def _make_serve_cmd_completed_hook(agent_holder: dict[str, Any]):
"""Build the ``on_cmd_completed`` hook used by serve mode.
Adopts ``/model`` agent swaps and ``/resume`` thread swaps back
into ``agent_holder`` so the outer poll loop picks up the new
handles on subsequent messages. Also keeps
Adopts ``/model`` agent swaps and ``/resume`` thread/workspace
swaps back into ``agent_holder`` so the outer poll loop picks up
the new handles on subsequent messages. Also keeps
``EvoScientist.cli.channel`` globals in sync so other readers
(e.g. the bus) see the new values.
@@ -574,6 +577,10 @@ def _make_serve_cmd_completed_hook(agent_holder: dict[str, Any]):
except Exception: # pragma: no cover — defensive
pass
new_workspace = getattr(ctx, "workspace_dir", None)
if new_workspace and new_workspace != agent_holder.get("workspace_dir"):
agent_holder["workspace_dir"] = new_workspace
# Surface the in-memory-state limitation to the channel user
# for ``/resume`` so the missing history isn't silent. Flush
# is required because ``cmd_manager.execute`` already flushed
@@ -607,19 +614,25 @@ def _serve_process_message(
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``)
shared with the outer ``serve()`` loop. ``on_cmd_completed`` (the
agent-swap / thread-swap adoption hook) and ``start_new_session_cb``
(thread rotation for ``/new``) are constructed once in ``serve()``
— if omitted, they're rebuilt per message (backward compat for
existing tests). ``/resume`` lands via the ``on_cmd_completed``
hook because the command mutates ``ctx.thread_id`` directly.
``agent_holder`` is a mutable dict (keys: ``agent``, ``thread_id``,
``workspace_dir``) 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
message (backward compat for existing tests). ``/resume`` lands
via the ``on_cmd_completed`` hook because the command mutates
``ctx.thread_id`` / ``ctx.workspace_dir`` directly.
"""
import asyncio
from .channel import _bus_loop
from .tui_runtime import run_streaming
if not _claim_or_complete_channel_request(msg):
return
runtime_workspace = agent_holder.get("workspace_dir") or workspace_dir
console.print(
f"[dim][{msg.channel_type}] {msg.sender}: {escape(msg.content[:80])}[/dim]"
)
@@ -694,70 +707,76 @@ def _serve_process_message(
# the ``finally`` below so subsequent messages start from a clean
# slate. Loop creation lives inside the try so an exception between
# creation and ``set_event_loop`` still closes the loop.
_prev_loop: asyncio.AbstractEventLoop | None
try:
_prev_loop = asyncio.get_event_loop_policy().get_event_loop()
except RuntimeError:
_prev_loop = None
_slash_loop: asyncio.AbstractEventLoop | None = None
_slash_handled = False
_slash_error: Exception | None = None
try:
_slash_loop = asyncio.new_event_loop()
asyncio.set_event_loop(_slash_loop)
_slash_handled = _slash_loop.run_until_complete(
dispatch_channel_slash_command(
msg,
agent=agent_holder["agent"],
thread_id=agent_holder["thread_id"],
workspace_dir=workspace_dir,
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),
on_cmd_completed=on_cmd_completed
or _make_serve_cmd_completed_hook(agent_holder),
_prev_loop: asyncio.AbstractEventLoop | None
try:
_prev_loop = asyncio.get_event_loop_policy().get_event_loop()
except RuntimeError:
_prev_loop = None
_slash_loop: asyncio.AbstractEventLoop | None = None
_slash_handled = False
_slash_error: Exception | None = None
try:
_slash_loop = asyncio.new_event_loop()
asyncio.set_event_loop(_slash_loop)
_slash_handled = _slash_loop.run_until_complete(
dispatch_channel_slash_command(
msg,
agent=agent_holder["agent"],
thread_id=agent_holder["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),
on_cmd_completed=on_cmd_completed
or _make_serve_cmd_completed_hook(agent_holder),
)
)
)
except Exception as exc:
_slash_error = exc
_serve_logger.exception("Slash dispatch failed for %s", msg.channel_type)
finally:
if _slash_loop is not None:
_slash_loop.close()
asyncio.set_event_loop(_prev_loop)
except Exception as exc:
_slash_error = exc
_serve_logger.exception("Slash dispatch failed for %s", msg.channel_type)
finally:
if _slash_loop is not None:
_slash_loop.close()
asyncio.set_event_loop(_prev_loop)
if _slash_error is not None:
_set_channel_response(msg.msg_id, f"Command error: {_slash_error}")
console.print(f"[red]Slash command error: {escape(str(_slash_error))}[/red]")
return
if _slash_error is not None:
_set_channel_response(msg.msg_id, f"Command error: {_slash_error}")
console.print(
f"[red]Slash command error: {escape(str(_slash_error))}[/red]"
)
return
if _slash_handled:
if _slash_handled:
console.print(f"[dim][{msg.channel_type}] Replied to {msg.sender}[/dim]")
return
meta = build_metadata(runtime_workspace, model)
try:
response = run_streaming(
ui_backend="cli",
agent=agent_holder["agent"],
message=msg.content,
thread_id=agent_holder["thread_id"],
show_thinking=show_thinking,
interactive=True,
metadata=meta,
on_thinking=_send_thinking,
on_todo=_send_todo,
on_file_write=_send_media,
hitl_prompt_fn=_hitl_prompt,
ask_user_prompt_fn=_ask_user_prompt,
cancel_scope=_channel_message_cancel_scope(msg),
)
except Exception as e:
response = f"Error: {e}"
console.print(f"[red]Serve error: {e}[/red]")
_set_channel_response(msg.msg_id, response)
console.print(f"[dim][{msg.channel_type}] Replied to {msg.sender}[/dim]")
return
meta = build_metadata(workspace_dir, model)
try:
response = run_streaming(
ui_backend="cli",
agent=agent_holder["agent"],
message=msg.content,
thread_id=agent_holder["thread_id"],
show_thinking=show_thinking,
interactive=True,
metadata=meta,
on_thinking=_send_thinking,
on_todo=_send_todo,
on_file_write=_send_media,
hitl_prompt_fn=_hitl_prompt,
ask_user_prompt_fn=_ask_user_prompt,
)
except Exception as e:
response = f"Error: {e}"
console.print(f"[red]Serve error: {e}[/red]")
_set_channel_response(msg.msg_id, response)
console.print(f"[dim][{msg.channel_type}] Replied to {msg.sender}[/dim]")
finally:
_complete_channel_request(msg.msg_id)
# =============================================================================
@@ -862,7 +881,11 @@ def serve(
# 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}
agent_holder: dict[str, Any] = {
"agent": agent,
"thread_id": tid,
"workspace_dir": ws,
}
# Build the slash-dispatch callbacks once; the poll loop reuses
# them for every inbound message. Without this hoist each message
+170 -161
View File
@@ -729,175 +729,184 @@ def cmd_interactive(
[channel: Replied to sender]
─────────────────
"""
# Clear the waiting ❯ prompt line
sys.stdout.write("\r\033[2K")
sys.stdout.flush()
# Reprint as if user typed it after ❯
prompt_line = Text()
prompt_line.append("\u276f ", style="bold blue")
prompt_line.append(msg.content)
console.print(prompt_line)
rx = Text()
rx.append(f"[{msg.channel_type}: Received from ", style="dim")
rx.append(msg.sender, style="cyan")
rx.append("]", style="dim")
console.print(rx)
_print_separator()
def _send_to_channel(coro, label: str, timeout: int = 15) -> None:
"""Schedule an async channel send on the bus loop."""
loop = _ch_mod._bus_loop
if not loop:
return
try:
asyncio.run_coroutine_threadsafe(coro, loop).result(
timeout=timeout
)
except Exception as e:
_channel_logger.debug(f"{label} send failed: {e}")
def _send_thinking_to_channel(thinking: str) -> None:
ch = msg.channel_ref
if ch and ch.send_thinking:
_send_to_channel(
ch.send_thinking_message(
sender=msg.chat_id,
thinking=thinking,
metadata=msg.metadata,
),
"Thinking",
)
def _send_todo_to_channel(items: list[dict]) -> None:
from ..channels.consumer import _format_todo_list
if msg.channel_ref:
_send_to_channel(
msg.channel_ref.send_todo_message(
sender=msg.chat_id,
content=_format_todo_list(items),
metadata=msg.metadata,
),
"Todo",
)
def _send_media_to_channel(file_path: str) -> None:
if msg.channel_ref:
_send_to_channel(
msg.channel_ref.send_media(
recipient=msg.chat_id,
file_path=file_path,
metadata=msg.metadata,
),
"Media",
timeout=30,
)
def _channel_hitl_prompt(action_requests: list) -> list[dict] | None:
"""Send HITL approval prompt to channel user and wait for reply."""
return _ch_mod.channel_hitl_prompt(action_requests, msg)
def _channel_ask_user(ask_user_data: dict) -> dict:
"""Send ask_user questions to channel user and wait for reply."""
return _ch_mod.channel_ask_user_prompt(ask_user_data, msg)
# ---- Slash command dispatch (cmd_manager, not the agent) ----
# Mirrors the TUI's behavior so ``/evoskills``, ``/mcp list``
# etc. sent via iMessage actually execute instead of being
# fed to the LLM as a plain prompt.
async def _on_channel_cmd_completed(
ctx: Any, original_agent: Any, cmd: Any
) -> None:
"""Mirror the REPL adoption block at
``interactive.py:1005-1030`` so ``/model`` and similar
state-mutating commands invoked via a channel actually
rebind the running session and keep the status bar
in sync."""
nonlocal model
agent_swapped = (
ctx.agent is not None and ctx.agent is not original_agent
)
if agent_swapped:
from ..EvoScientist import _ensure_config
agent_loader.adopt(ctx.agent)
cfg = _ensure_config()
model = cfg.model
state["status_base_snapshot"] = make_empty_status_snapshot(
model
)
if _channels_is_running():
_ch_mod._cli_agent = ctx.agent
_ch_mod._cli_thread_id = state["thread_id"]
# ``/new`` rotates ``state["thread_id"]`` / workspace,
# ``/compact`` reduces token usage — both need the
# status snapshot re-rendered even when the agent
# didn't swap. ``/resume`` refreshes inline in its
# own async callback.
if agent_swapped or getattr(cmd, "name", None) in (
"/compact",
"/new",
):
await _refresh_status_snapshot(reset_streaming_text=True)
_slash_handled = await dispatch_channel_slash_command(
msg,
agent=agent_loader.agent,
thread_id=state["thread_id"],
workspace_dir=state["workspace_dir"],
checkpointer=checkpointer,
append_system=lambda t, s="dim": console.print(t, style=s),
start_new_session_cb=_on_start_new_session,
handle_session_resume_cb=_on_handle_session_resume,
await_agent_ready=_await_agent_ready,
on_cmd_completed=_on_channel_cmd_completed,
)
if _slash_handled:
_print_separator()
sys.stdout.write("\033[34;1m❯\033[0m ")
sys.stdout.flush()
if not _ch_mod._claim_or_complete_channel_request(msg):
return
try:
ready_agent = await _await_agent_ready()
meta = build_metadata(state["workspace_dir"], model)
await _refresh_status_snapshot(
msg.content, reset_streaming_text=True
)
response = run_streaming(
ui_backend=state["ui_backend"],
agent=ready_agent,
message=msg.content,
# Clear the waiting ❯ prompt line
sys.stdout.write("\r\033[2K")
sys.stdout.flush()
# Reprint as if user typed it after ❯
prompt_line = Text()
prompt_line.append("\u276f ", style="bold blue")
prompt_line.append(msg.content)
console.print(prompt_line)
rx = Text()
rx.append(f"[{msg.channel_type}: Received from ", style="dim")
rx.append(msg.sender, style="cyan")
rx.append("]", style="dim")
console.print(rx)
_print_separator()
def _send_to_channel(coro, label: str, timeout: int = 15) -> None:
"""Schedule an async channel send on the bus loop."""
loop = _ch_mod._bus_loop
if not loop:
return
try:
asyncio.run_coroutine_threadsafe(coro, loop).result(
timeout=timeout
)
except Exception as e:
_channel_logger.debug(f"{label} send failed: {e}")
def _send_thinking_to_channel(thinking: str) -> None:
ch = msg.channel_ref
if ch and ch.send_thinking:
_send_to_channel(
ch.send_thinking_message(
sender=msg.chat_id,
thinking=thinking,
metadata=msg.metadata,
),
"Thinking",
)
def _send_todo_to_channel(items: list[dict]) -> None:
from ..channels.consumer import _format_todo_list
if msg.channel_ref:
_send_to_channel(
msg.channel_ref.send_todo_message(
sender=msg.chat_id,
content=_format_todo_list(items),
metadata=msg.metadata,
),
"Todo",
)
def _send_media_to_channel(file_path: str) -> None:
if msg.channel_ref:
_send_to_channel(
msg.channel_ref.send_media(
recipient=msg.chat_id,
file_path=file_path,
metadata=msg.metadata,
),
"Media",
timeout=30,
)
def _channel_hitl_prompt(
action_requests: list,
) -> list[dict] | None:
"""Send HITL approval prompt to channel user and wait for reply."""
return _ch_mod.channel_hitl_prompt(action_requests, msg)
def _channel_ask_user(ask_user_data: dict) -> dict:
"""Send ask_user questions to channel user and wait for reply."""
return _ch_mod.channel_ask_user_prompt(ask_user_data, msg)
# ---- Slash command dispatch (cmd_manager, not the agent) ----
# Mirrors the TUI's behavior so ``/evoskills``, ``/mcp list``
# etc. sent via iMessage actually execute instead of being
# fed to the LLM as a plain prompt.
async def _on_channel_cmd_completed(
ctx: Any, original_agent: Any, cmd: Any
) -> None:
"""Mirror the REPL adoption block at
``interactive.py:1005-1030`` so ``/model`` and similar
state-mutating commands invoked via a channel actually
rebind the running session and keep the status bar
in sync."""
nonlocal model
agent_swapped = (
ctx.agent is not None and ctx.agent is not original_agent
)
if agent_swapped:
from ..EvoScientist import _ensure_config
agent_loader.adopt(ctx.agent)
cfg = _ensure_config()
model = cfg.model
state["status_base_snapshot"] = make_empty_status_snapshot(
model
)
if _channels_is_running():
_ch_mod._cli_agent = ctx.agent
_ch_mod._cli_thread_id = state["thread_id"]
# ``/new`` rotates ``state["thread_id"]`` / workspace,
# ``/compact`` reduces token usage — both need the
# status snapshot re-rendered even when the agent
# didn't swap. ``/resume`` refreshes inline in its
# own async callback.
if agent_swapped or getattr(cmd, "name", None) in (
"/compact",
"/new",
):
await _refresh_status_snapshot(reset_streaming_text=True)
_slash_handled = await dispatch_channel_slash_command(
msg,
agent=agent_loader.agent,
thread_id=state["thread_id"],
show_thinking=show_thinking,
interactive=True,
metadata=meta,
on_thinking=_send_thinking_to_channel,
on_todo=_send_todo_to_channel,
on_file_write=_send_media_to_channel,
hitl_prompt_fn=_channel_hitl_prompt,
ask_user_prompt_fn=_channel_ask_user,
on_stream_event=_handle_stream_status_event,
status_footer_builder=_stream_status_footer,
workspace_dir=state["workspace_dir"],
checkpointer=checkpointer,
append_system=lambda t, s="dim": console.print(t, style=s),
start_new_session_cb=_on_start_new_session,
handle_session_resume_cb=_on_handle_session_resume,
await_agent_ready=_await_agent_ready,
on_cmd_completed=_on_channel_cmd_completed,
)
except Exception as e:
response = f"Error: {e}"
console.print(f"[red]Channel error: {e}[/red]")
if _slash_handled:
_print_separator()
sys.stdout.write("\033[34;1m❯\033[0m ")
sys.stdout.flush()
return
_set_channel_response(msg.msg_id, response)
await _refresh_status_snapshot(reset_streaming_text=True)
try:
ready_agent = await _await_agent_ready()
meta = build_metadata(state["workspace_dir"], model)
await _refresh_status_snapshot(
msg.content, reset_streaming_text=True
)
response = run_streaming(
ui_backend=state["ui_backend"],
agent=ready_agent,
message=msg.content,
thread_id=state["thread_id"],
show_thinking=show_thinking,
interactive=True,
metadata=meta,
on_thinking=_send_thinking_to_channel,
on_todo=_send_todo_to_channel,
on_file_write=_send_media_to_channel,
hitl_prompt_fn=_channel_hitl_prompt,
ask_user_prompt_fn=_channel_ask_user,
on_stream_event=_handle_stream_status_event,
status_footer_builder=_stream_status_footer,
cancel_scope=_ch_mod._channel_message_cancel_scope(msg),
)
except Exception as e:
response = f"Error: {e}"
console.print(f"[red]Channel error: {e}[/red]")
tx = Text()
tx.append(f"[{msg.channel_type}: Replied to ", style="dim")
tx.append(msg.sender, style="cyan")
tx.append("]", style="dim")
console.print(tx)
_print_separator()
_set_channel_response(msg.msg_id, response)
await _refresh_status_snapshot(reset_streaming_text=True)
# Redraw the ❯ prompt on a new line after separator
sys.stdout.write("\033[34;1m\u276f\033[0m ")
sys.stdout.flush()
tx = Text()
tx.append(f"[{msg.channel_type}: Replied to ", style="dim")
tx.append(msg.sender, style="cyan")
tx.append("]", style="dim")
console.print(tx)
_print_separator()
# Redraw the ❯ prompt on a new line after separator
sys.stdout.write("\033[34;1m\u276f\033[0m ")
sys.stdout.flush()
finally:
_ch_mod._complete_channel_request(msg.msg_id)
async def _check_channel_queue() -> None:
"""Poll the channel message queue and dispatch to the agent."""
+3
View File
@@ -30,6 +30,7 @@ class StreamingTUIBackend(Protocol):
metadata: dict | 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,
) -> str:
"""Run streaming and return final response text."""
@@ -56,6 +57,7 @@ class RichStreamingBackend:
metadata: dict | 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,
) -> str:
return _run_streaming(
agent=agent,
@@ -71,4 +73,5 @@ class RichStreamingBackend:
metadata=metadata,
hitl_prompt_fn=hitl_prompt_fn,
ask_user_prompt_fn=ask_user_prompt_fn,
cancel_scope=cancel_scope,
)
+88 -21
View File
@@ -191,6 +191,29 @@ _SUMMARY_CONTINUATION_EVENTS = {
}
async def _sync_tui_command_completion(
app: Any,
ctx: CommandContext,
original_agent: Any,
cmd: Any,
) -> None:
"""Adopt successful command-side state changes back into the TUI app."""
agent_swapped = ctx.agent is not None and ctx.agent is not original_agent
if agent_swapped:
from ..EvoScientist import _ensure_config
app._agent_loader.adopt(ctx.agent)
cfg = _ensure_config()
update_model = getattr(app, "update_status_after_model_change", None)
if callable(update_model):
update_model(cfg.model, cfg.provider)
if _channels_is_running():
_ch_mod._cli_agent = ctx.agent
_ch_mod._cli_thread_id = app._conversation_tid
await app._refresh_status_snapshot(reset_streaming_text=True)
def _should_finalize_active_summarization(event_type: str) -> bool:
"""Return whether an active summary panel should stop for this event."""
return bool(event_type) and event_type not in _SUMMARY_CONTINUATION_EVENTS
@@ -730,6 +753,14 @@ def run_textual_interactive(
lambda m=msg: asyncio.ensure_future(self._process_channel_message(m))
)
async def _on_channel_cmd_completed(
self,
ctx: CommandContext,
original_agent: Any,
cmd: Any,
) -> None:
await _sync_tui_command_completion(self, ctx, original_agent, cmd)
# ── Widget helpers ─────────────────────────────────────
def _schedule_scroll_to_bottom(
@@ -975,6 +1006,7 @@ def run_textual_interactive(
file_warnings: list[str] | None = None,
channel_hitl_fn: Callable[[list], list[dict] | None] | None = None,
channel_ask_user_fn: Callable[[dict], dict] | None = None,
cancel_scope: str | None = None,
) -> str:
"""Stream agent events and mount widgets. Returns response text.
@@ -996,6 +1028,11 @@ def run_textual_interactive(
When provided (channel messages), this is called instead
of mounting the AskUserWidget.
"""
from ..stream.display import (
build_stopped_response_text,
is_stream_cancel_requested,
)
container = self.query_one("#chat", VerticalScroll)
# 1. Mount user message + loading spinner
@@ -1065,6 +1102,30 @@ def run_textual_interactive(
except Exception:
pass
async def _mark_cancelled_response() -> str:
nonlocal assistant_w
previous_text = state.response_text or ""
current, final_text = build_stopped_response_text(previous_text)
state.response_text = final_text
self._set_status_streaming_text(final_text)
if assistant_w is None:
if final_text:
assistant_w = AssistantMessage(final_text)
await container.mount(assistant_w)
else:
if previous_text != current:
assistant_w._content = final_text
await assistant_w.stop_stream()
else:
suffix = final_text[len(current) :]
if suffix:
await assistant_w.append_content(suffix)
_schedule_scroll()
return final_text
def _finalize_active_summarization() -> None:
"""Stop the active summary timer once the stream moves on."""
if summarization_w is not None and summarization_w._is_active:
@@ -1140,6 +1201,9 @@ def run_textual_interactive(
_stream_input: Any = user_text # str or Command for HITL resume
for _hitl_round in range(_MAX_HITL_ROUNDS):
if is_stream_cancel_requested(cancel_scope):
response = await _mark_cancelled_response()
break
state.pending_interrupt = None
state.pending_ask_user = None
_hitl_resuming = False
@@ -1154,6 +1218,9 @@ def run_textual_interactive(
self._conversation_tid,
metadata=metadata,
):
if is_stream_cancel_requested(cancel_scope):
response = await _mark_cancelled_response()
break
event_type = state.handle_event(event)
if event_type == "usage_stats":
@@ -1457,6 +1524,10 @@ def run_textual_interactive(
result = await asyncio.to_thread(
lambda f=_ask_fn, e=event: f(e),
)
if is_stream_cancel_requested(cancel_scope):
state.pending_ask_user = None
response = await _mark_cancelled_response()
break
else:
# Interactive TUI: display widget, collect via arrow keys
from .widgets.ask_user_widget import AskUserWidget
@@ -1511,6 +1582,10 @@ def run_textual_interactive(
channel_hitl_fn,
action_reqs,
)
if is_stream_cancel_requested(cancel_scope):
state.pending_interrupt = None
response = await _mark_cancelled_response()
break
if decisions is not None:
from langgraph.types import (
Command, # type: ignore[import-untyped]
@@ -1674,6 +1749,9 @@ def run_textual_interactive(
self._schedule_scroll_to_bottom(container)
# HITL / ask_user: if interrupt was handled, loop back to resume stream
if is_stream_cancel_requested(cancel_scope):
response = await _mark_cancelled_response()
break
if state.pending_interrupt is None and state.pending_ask_user is None:
break # normal completion or rejection — exit HITL loop
# Otherwise _stream_input was set to Command(resume=...)
@@ -1735,6 +1813,8 @@ def run_textual_interactive(
[channel: Replied to sender]
"""
prompt_widget = None
if not _ch_mod._claim_or_complete_channel_request(msg):
return
try:
self._busy = True
await self._refresh_status_snapshot(msg.content)
@@ -1834,6 +1914,7 @@ def run_textual_interactive(
start_new_session_cb=self.start_new_session,
handle_session_resume_cb=self.handle_session_resume,
await_agent_ready=self._await_agent_ready,
on_cmd_completed=self._on_channel_cmd_completed,
)
if _slash_handled:
return # outer finally handles _busy / widget cleanup
@@ -1858,6 +1939,7 @@ def run_textual_interactive(
skip_user_message=True,
channel_hitl_fn=_channel_hitl_prompt,
channel_ask_user_fn=_channel_ask_user,
cancel_scope=_ch_mod._channel_message_cancel_scope(msg),
)
except Exception as exc:
response = f"Error: {exc}"
@@ -1876,6 +1958,7 @@ def run_textual_interactive(
if prompt_widget is not None:
prompt_widget.disabled = False
prompt_widget.focus()
_ch_mod._complete_channel_request(msg.msg_id)
# ── Clipboard (copy on mouse select) ─────────────────
@@ -2292,27 +2375,11 @@ def run_textual_interactive(
)
if await cmd_manager.execute(command, ctx):
# Sync agent back if command replaced it (e.g. /model).
# ``is not None`` guard: non-agent commands (ctx.agent
# starts None) must not clobber a valid loaded agent.
if (
ctx.agent is not None
and ctx.agent is not self._agent_loader.agent
):
# ``adopt`` also cancels/supersedes any in-flight
# load so a late completion can't overwrite the
# replacement agent (/model on a broken provider).
self._agent_loader.adopt(ctx.agent)
if _channels_is_running():
_ch_mod._cli_agent = ctx.agent
_ch_mod._cli_thread_id = self._conversation_tid
# Do NOT invalidate the usage baseline after /compact.
# build_session_status_snapshot() only counts raw checkpoint
# messages (~46 tokens) and misses system prompt + tool
# definitions (~50K overhead). The stale pre-compact count
# is far more accurate; the next LLM call will correct it.
await self._refresh_status_snapshot(
reset_streaming_text=True,
await _sync_tui_command_completion(
self,
ctx,
self._agent_loader.agent,
cmd,
)
return
+3
View File
@@ -74,6 +74,7 @@ def run_streaming(
metadata: dict | 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,
) -> str:
"""Run streaming with the selected backend."""
backend = get_backend(ui_backend, warn_fallback=True)
@@ -92,6 +93,7 @@ def run_streaming(
metadata=metadata,
hitl_prompt_fn=hitl_prompt_fn,
ask_user_prompt_fn=ask_user_prompt_fn,
cancel_scope=cancel_scope,
)
except RuntimeError:
requested = normalize_ui_backend(ui_backend)
@@ -113,5 +115,6 @@ def run_streaming(
metadata=metadata,
hitl_prompt_fn=hitl_prompt_fn,
ask_user_prompt_fn=ask_user_prompt_fn,
cancel_scope=cancel_scope,
)
raise
+90 -3
View File
@@ -1,14 +1,19 @@
from __future__ import annotations
import asyncio
import logging
from typing import Any
from .base import CommandUI
_logger = logging.getLogger(__name__)
class ChannelCommandUI(CommandUI):
"""CommandUI implementation for messaging channels with output buffering."""
_TEXT_CHUNK_LIMIT = 3500
@property
def supports_interactive(self) -> bool:
return False
@@ -26,14 +31,54 @@ class ChannelCommandUI(CommandUI):
self.handle_session_resume_callback = handle_session_resume_callback
self._system_buffer: list[str] = []
def append_system(self, text: str, style: str = "dim") -> None:
if self.append_system_callback:
def _queue_system(
self,
text: str,
style: str = "dim",
*,
mirror_local: bool = True,
) -> None:
if mirror_local and self.append_system_callback:
self.append_system_callback(text, style)
# Buffer the text for grouped delivery to the channel
# We ignore style for grouping but keep it for individual lines if needed
self._system_buffer.append(text)
def append_system(self, text: str, style: str = "dim") -> None:
self._queue_system(text, style)
@staticmethod
def _extract_message_text(message: Any) -> str:
content = getattr(message, "content", "") or ""
if isinstance(content, list):
parts = [
block.get("text", "")
for block in content
if isinstance(block, dict) and block.get("type") == "text"
]
content = " ".join(parts) if parts else ""
return str(content).strip()
async def _send_text_chunks(self, text: str, *, mirror_local: bool = True) -> None:
"""Flush long plain-text payloads in channel-safe chunks."""
text = (text or "").strip()
if not text:
return
pending = text
while pending:
chunk = pending[: self._TEXT_CHUNK_LIMIT]
if len(pending) > self._TEXT_CHUNK_LIMIT:
split_at = chunk.rfind("\n")
if split_at > 0:
chunk = chunk[:split_at]
chunk = chunk.rstrip()
if not chunk:
chunk = pending[: self._TEXT_CHUNK_LIMIT]
self._queue_system(chunk, mirror_local=mirror_local)
await self.flush()
pending = pending[len(chunk) :].lstrip("\n")
async def flush(self) -> None:
"""Send all buffered system messages as a single grouped message."""
if not self._system_buffer:
@@ -140,5 +185,47 @@ class ChannelCommandUI(CommandUI):
async def handle_session_resume(
self, thread_id: str, workspace_dir: str | None = None
) -> None:
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)
except Exception as exc:
_logger.exception(
"Failed to load saved history for resumed thread %s",
thread_id,
)
lines.append(f"(history unavailable: {exc})")
await self._send_text_chunks("\n".join(lines), mirror_local=mirror_local)
return
display = [m for m in messages if getattr(m, "type", None) in ("human", "ai")]
if not display:
if messages:
lines.append("No displayable messages in this session.")
else:
lines.append("No saved messages in this session.")
await self._send_text_chunks("\n".join(lines), mirror_local=mirror_local)
return
HISTORY_WINDOW = 10
if len(display) > HISTORY_WINDOW:
display = display[-HISTORY_WINDOW:]
lines.append(f"Conversation history (last {HISTORY_WINDOW} messages):")
else:
lines.append("Conversation history:")
for message in display:
text = self._extract_message_text(message)
if not text:
continue
if getattr(message, "type", None) == "human":
lines.append(f"User: {text}")
else:
lines.append(f"EvoScientist: {text}")
await self._send_text_chunks("\n".join(lines), mirror_local=mirror_local)
+253 -146
View File
@@ -9,6 +9,7 @@ import asyncio
import inspect
import logging
import os
import threading
from collections.abc import Callable
from typing import Any
@@ -49,6 +50,79 @@ _MEDIA_EXTENSIONS = {".png", ".jpg", ".jpeg", ".gif", ".webp", ".bmp", ".svg", "
formatter = ToolResultFormatter()
# Stream-cancel events keyed by logical stream scope. Channel messages pass a
# per-message scope so `/stop` only affects that message's run; scope-less
# callers retain the legacy process-wide default event.
_DEFAULT_STREAM_CANCEL_SCOPE = "__default__"
_stream_cancel_lock = threading.Lock()
_stream_cancel_events: dict[str, threading.Event] = {
_DEFAULT_STREAM_CANCEL_SCOPE: threading.Event()
}
# Backward-compat alias used by older tests and direct imports.
_stream_cancel_event = _stream_cancel_events[_DEFAULT_STREAM_CANCEL_SCOPE]
def _stream_cancel_scope_key(cancel_scope: str | None) -> str:
return cancel_scope or _DEFAULT_STREAM_CANCEL_SCOPE
def _get_stream_cancel_event(
cancel_scope: str | None,
*,
create: bool = False,
) -> threading.Event | None:
scope_key = _stream_cancel_scope_key(cancel_scope)
with _stream_cancel_lock:
event = _stream_cancel_events.get(scope_key)
if event is None and create:
event = threading.Event()
_stream_cancel_events[scope_key] = event
return event
def request_stream_cancel(cancel_scope: str | None = None) -> bool:
"""Signal a specific in-flight stream to terminate."""
event = _get_stream_cancel_event(cancel_scope, create=True)
already_requested = event.is_set()
event.set()
return not already_requested
def is_stream_cancel_requested(cancel_scope: str | None = None) -> bool:
event = _get_stream_cancel_event(cancel_scope)
return event.is_set() if event is not None else False
def clear_stream_cancel(cancel_scope: str | None = None) -> None:
"""Clear a scope's stop signal without dropping the scope entry."""
event = _get_stream_cancel_event(cancel_scope)
if event is not None:
event.clear()
def discard_stream_cancel(cancel_scope: str | None = None) -> None:
"""Drop a scope's stop signal after the owning request is fully done."""
scope_key = _stream_cancel_scope_key(cancel_scope)
with _stream_cancel_lock:
if scope_key == _DEFAULT_STREAM_CANCEL_SCOPE:
_stream_cancel_events[scope_key].clear()
else:
_stream_cancel_events.pop(scope_key, None)
def build_stopped_response_text(previous_text: str | None) -> tuple[str, str]:
"""Normalize a cancelled response and return `(trimmed_previous, final_text)`."""
marker = "[Stopped.]"
current = (previous_text or "").rstrip()
if not current:
final_text = marker
elif current.endswith(marker):
final_text = current
else:
final_text = f"{current}\n{marker}"
return current, final_text
# ---------------------------------------------------------------------------
# Todo formatting
# ---------------------------------------------------------------------------
@@ -1092,6 +1166,7 @@ def _run_streaming(
metadata: dict | 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,
*,
_state: StreamState | None = None,
_hitl_depth: int = 0,
@@ -1123,17 +1198,31 @@ def _run_streaming(
Returns:
The final response text.
"""
# Scope-less callers keep the legacy single-event semantics. Scoped
# callers use unique per-request scopes, so pre-start `/stop` must
# remain armed until this run consumes it.
if _state is None and cancel_scope is None:
clear_stream_cancel()
state = _state if _state is not None else StreamState()
_todo_sent = False
if _media_sent is None:
_media_sent = set()
_MIN_THINKING_LEN = 200
def _stopped_response() -> str:
_, final_text = build_stopped_response_text(state.response_text)
state.response_text = final_text
return final_text
async def _consume() -> None:
nonlocal _sent_thinking_text, _todo_sent
async for event in stream_agent_events(
agent, message, thread_id, metadata=metadata
):
if is_stream_cancel_requested(cancel_scope):
_stopped_response()
return
event_type = state.handle_event(event)
# Relay thinking to channel when transitioning away from
@@ -1226,159 +1315,136 @@ def _run_streaming(
)
)
with Live(
console=console,
auto_refresh=False,
transient=False,
vertical_overflow="visible",
) as live:
live.update(
create_streaming_display(
is_waiting=True,
status_footer=(
status_footer_builder() if status_footer_builder else None
),
try:
if is_stream_cancel_requested(cancel_scope):
return _stopped_response()
with Live(
console=console,
auto_refresh=False,
transient=False,
vertical_overflow="visible",
) as live:
live.update(
create_streaming_display(
is_waiting=True,
status_footer=(
status_footer_builder() if status_footer_builder else None
),
)
)
)
# Determine how to run the async streaming coroutine.
# - In TUI mode (Textual), there's already a running event loop;
# nest_asyncio is needed to allow run_until_complete inside it.
# - In serve/CLI mode, the main thread has no running loop;
# use a fresh event loop directly (no nest_asyncio needed or wanted,
# since nest_asyncio.apply() patches globally and breaks the bus
# thread's event loop Task-context detection).
try:
running_loop = asyncio.get_running_loop()
except RuntimeError:
running_loop = None
if running_loop is not None:
# Already inside a running loop (TUI) — must use nest_asyncio.
# NOTE: nest_asyncio.apply() is global and irreversible within
# the process; avoid mixing TUI and serve modes in one process.
import nest_asyncio # type: ignore[import-untyped]
nest_asyncio.apply()
loop = running_loop
else:
# No running loop (serve/CLI) — create a fresh one
# Determine how to run the async streaming coroutine.
# - In TUI mode (Textual), there's already a running event loop;
# nest_asyncio is needed to allow run_until_complete inside it.
# - In serve/CLI mode, the main thread has no running loop;
# use a fresh event loop directly (no nest_asyncio needed or wanted,
# since nest_asyncio.apply() patches globally and breaks the bus
# thread's event loop Task-context detection).
try:
loop = _get_event_loop()
running_loop = asyncio.get_running_loop()
except RuntimeError:
loop = _create_event_loop()
running_loop = None
async def _run_with_refresh() -> None:
async def _periodic_refresh() -> None:
if running_loop is not None:
# Already inside a running loop (TUI) — must use nest_asyncio.
# NOTE: nest_asyncio.apply() is global and irreversible within
# the process; avoid mixing TUI and serve modes in one process.
import nest_asyncio # type: ignore[import-untyped]
nest_asyncio.apply()
loop = running_loop
else:
# No running loop (serve/CLI) — create a fresh one
try:
while True:
await asyncio.sleep(0.05)
live.refresh()
except asyncio.CancelledError:
pass
loop = _get_event_loop()
except RuntimeError:
loop = _create_event_loop()
refresh_task = asyncio.ensure_future(_periodic_refresh())
try:
await _consume()
finally:
refresh_task.cancel()
async def _run_with_refresh() -> None:
async def _periodic_refresh() -> None:
try:
while True:
await asyncio.sleep(0.05)
live.refresh()
except asyncio.CancelledError:
pass
refresh_task = asyncio.ensure_future(_periodic_refresh())
try:
await refresh_task
except asyncio.CancelledError:
pass
# Render clean final frame before Live exits (no spinners, expanded tools)
if (
state.pending_interrupt is not None
or state.pending_ask_user is not None
):
# Interrupted: render current state (not final) so it
# looks continuous when prompt appears.
final_display = create_streaming_display(
**state.get_display_args(),
show_thinking=show_thinking,
response_markdown=state.get_response_markdown(),
status_footer=resolve_final_status_footer(
interactive, status_footer_builder
),
)
elif interactive:
final_display = create_streaming_display(
**state.get_display_args(),
show_thinking=show_thinking,
is_final=True,
final_show_thinking=False,
response_markdown=state.get_response_markdown(),
status_footer=resolve_final_status_footer(
interactive, status_footer_builder
),
)
else:
final_display = create_streaming_display(
**state.get_display_args(),
show_thinking=show_thinking,
is_final=True,
final_show_thinking=True,
final_thinking_max_length=DisplayLimits.THINKING_FINAL,
response_markdown=state.get_response_markdown(),
status_footer=resolve_final_status_footer(
interactive, status_footer_builder
),
)
live.update(final_display)
live.refresh()
await _consume()
finally:
refresh_task.cancel()
try:
await refresh_task
except asyncio.CancelledError:
pass
# Render clean final frame before Live exits (no spinners, expanded tools)
if (
state.pending_interrupt is not None
or state.pending_ask_user is not None
):
# Interrupted: render current state (not final) so it
# looks continuous when prompt appears.
final_display = create_streaming_display(
**state.get_display_args(),
show_thinking=show_thinking,
response_markdown=state.get_response_markdown(),
status_footer=resolve_final_status_footer(
interactive, status_footer_builder
),
)
elif interactive:
final_display = create_streaming_display(
**state.get_display_args(),
show_thinking=show_thinking,
is_final=True,
final_show_thinking=False,
response_markdown=state.get_response_markdown(),
status_footer=resolve_final_status_footer(
interactive, status_footer_builder
),
)
else:
final_display = create_streaming_display(
**state.get_display_args(),
show_thinking=show_thinking,
is_final=True,
final_show_thinking=True,
final_thinking_max_length=DisplayLimits.THINKING_FINAL,
response_markdown=state.get_response_markdown(),
status_footer=resolve_final_status_footer(
interactive, status_footer_builder
),
)
live.update(final_display)
live.refresh()
loop.run_until_complete(_run_with_refresh())
loop.run_until_complete(_run_with_refresh())
# Flush any remaining thinking that wasn't sent during streaming.
if on_thinking and state.thinking_text:
current = state.thinking_text.rstrip()
if len(current) >= _MIN_THINKING_LEN and current != _sent_thinking_text:
on_thinking(current)
_sent_thinking_text = current
# Flush any remaining thinking that wasn't sent during streaming.
if on_thinking and state.thinking_text:
current = state.thinking_text.rstrip()
if len(current) >= _MIN_THINKING_LEN and current != _sent_thinking_text:
on_thinking(current)
_sent_thinking_text = current
# ask_user: check before HITL (ask_user uses the same resume loop)
if state.pending_ask_user is not None and _hitl_depth < _MAX_HITL_ITERATIONS:
if ask_user_prompt_fn is not None:
result = ask_user_prompt_fn(state.pending_ask_user)
else:
result = _resolve_ask_user_prompt(state.pending_ask_user)
from langgraph.types import Command # type: ignore[import-untyped]
state.pending_ask_user = None
state.thinking_text = "" # reset accumulation for fresh round
return _run_streaming(
agent=agent,
message=Command(resume=result),
thread_id=thread_id,
show_thinking=show_thinking,
interactive=interactive,
on_thinking=on_thinking,
on_todo=on_todo,
on_file_write=on_file_write,
on_stream_event=on_stream_event,
status_footer_builder=status_footer_builder,
metadata=metadata,
hitl_prompt_fn=hitl_prompt_fn,
ask_user_prompt_fn=ask_user_prompt_fn,
_state=state,
_hitl_depth=_hitl_depth + 1,
_media_sent=_media_sent,
_sent_thinking_text=_sent_thinking_text,
)
# HITL: check for pending interrupt and handle approval
if state.pending_interrupt is not None and _hitl_depth < _MAX_HITL_ITERATIONS:
decisions = _resolve_hitl_approval(
state.pending_interrupt,
prompt_fn=hitl_prompt_fn,
)
if decisions is not None:
# ask_user: check before HITL (ask_user uses the same resume loop)
if state.pending_ask_user is not None and _hitl_depth < _MAX_HITL_ITERATIONS:
if is_stream_cancel_requested(cancel_scope):
return _stopped_response()
if ask_user_prompt_fn is not None:
result = ask_user_prompt_fn(state.pending_ask_user)
else:
result = _resolve_ask_user_prompt(state.pending_ask_user)
from langgraph.types import Command # type: ignore[import-untyped]
state.pending_interrupt = None
state.pending_ask_user = None
state.thinking_text = "" # reset accumulation for fresh round
if is_stream_cancel_requested(cancel_scope):
return _stopped_response()
return _run_streaming(
agent=agent,
message=Command(resume={"decisions": decisions}),
message=Command(resume=result),
thread_id=thread_id,
show_thinking=show_thinking,
interactive=interactive,
@@ -1390,21 +1456,62 @@ def _run_streaming(
metadata=metadata,
hitl_prompt_fn=hitl_prompt_fn,
ask_user_prompt_fn=ask_user_prompt_fn,
cancel_scope=cancel_scope,
_state=state,
_hitl_depth=_hitl_depth + 1,
_media_sent=_media_sent,
_sent_thinking_text=_sent_thinking_text,
)
elif state.pending_interrupt is not None:
_logger.warning(
"HITL loop reached max iterations (%d), stopping",
_MAX_HITL_ITERATIONS,
)
# Everything (tools, thinking, todos, response) is already on screen
# from Live's final frame (transient=False). No need to re-print.
# HITL: check for pending interrupt and handle approval
if state.pending_interrupt is not None and _hitl_depth < _MAX_HITL_ITERATIONS:
if is_stream_cancel_requested(cancel_scope):
return _stopped_response()
decisions = _resolve_hitl_approval(
state.pending_interrupt,
prompt_fn=hitl_prompt_fn,
)
if is_stream_cancel_requested(cancel_scope):
return _stopped_response()
if decisions is not None:
from langgraph.types import Command # type: ignore[import-untyped]
return (state.response_text or "").strip()
state.pending_interrupt = None
state.thinking_text = "" # reset accumulation for fresh round
if is_stream_cancel_requested(cancel_scope):
return _stopped_response()
return _run_streaming(
agent=agent,
message=Command(resume={"decisions": decisions}),
thread_id=thread_id,
show_thinking=show_thinking,
interactive=interactive,
on_thinking=on_thinking,
on_todo=on_todo,
on_file_write=on_file_write,
on_stream_event=on_stream_event,
status_footer_builder=status_footer_builder,
metadata=metadata,
hitl_prompt_fn=hitl_prompt_fn,
ask_user_prompt_fn=ask_user_prompt_fn,
cancel_scope=cancel_scope,
_state=state,
_hitl_depth=_hitl_depth + 1,
_media_sent=_media_sent,
_sent_thinking_text=_sent_thinking_text,
)
elif state.pending_interrupt is not None:
_logger.warning(
"HITL loop reached max iterations (%d), stopping",
_MAX_HITL_ITERATIONS,
)
# Everything (tools, thinking, todos, response) is already on screen
# from Live's final frame (transient=False). No need to re-print.
return (state.response_text or "").strip()
finally:
discard_stream_cancel(cancel_scope)
# ---------------------------------------------------------------------------
+262 -6
View File
@@ -30,14 +30,29 @@ def clean_channel_state():
"""Reset shared channel bridge state before and after each test."""
from EvoScientist.cli import channel as channel_mod
from EvoScientist.cli.channel import _message_queue
from EvoScientist.stream import display as display_mod
_drain_queue(_message_queue)
with channel_mod._response_lock:
channel_mod._pending_responses.clear()
def _reset() -> None:
_drain_queue(_message_queue)
with channel_mod._response_lock:
channel_mod._pending_responses.clear()
with channel_mod._channel_request_lock:
channel_mod._channel_requests.clear()
channel_mod._session_requests.clear()
channel_mod._cancelled_channel_messages.clear()
with channel_mod._hitl_lock:
channel_mod._pending_hitl.clear()
channel_mod._hitl_auto_approve.clear()
with display_mod._stream_cancel_lock:
display_mod._stream_cancel_event.clear()
display_mod._stream_cancel_events.clear()
display_mod._stream_cancel_events[
display_mod._DEFAULT_STREAM_CANCEL_SCOPE
] = display_mod._stream_cancel_event
_reset()
yield
_drain_queue(_message_queue)
with channel_mod._response_lock:
channel_mod._pending_responses.clear()
_reset()
class _FakeConfig:
@@ -251,6 +266,80 @@ class TestBusInboundConsumer:
_run(_test())
def test_late_timeout_keeps_active_request_cancellable(self, monkeypatch):
"""Late timeout must not discard an active request's cancel scope."""
from EvoScientist.cli import channel as channel_mod
from EvoScientist.cli.channel import (
_channel_message_cancel_scope,
_channel_request_state,
_claim_channel_request,
_handle_bus_message,
_message_queue,
)
from EvoScientist.stream import display as display_mod
monkeypatch.setattr(channel_mod, "_RESPONSE_TIMEOUT", 0.05)
monkeypatch.setattr(channel_mod, "_LATE_RESPONSE_TIMEOUT", 0.05)
async def _test():
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
task = asyncio.create_task(
_handle_bus_message(
bus,
manager,
InboundMessage(
channel="fake",
sender_id="user1",
chat_id="chat1",
content="still running",
message_id="msg-active",
),
)
)
queued = None
for _ in range(20):
with _message_queue.mutex:
queued = _message_queue.queue[0] if _message_queue.queue else None
if queued is not None:
break
await asyncio.sleep(0.05)
assert queued is not None
assert _claim_channel_request(queued) is True
notice = await asyncio.wait_for(bus.consume_outbound(), timeout=1.0)
assert "Still working on it" in notice.content
await task
assert _channel_request_state(queued.msg_id) == "active"
cancel_scope = _channel_message_cancel_scope(queued)
assert not display_mod.is_stream_cancel_requested(cancel_scope)
await _handle_bus_message(
bus,
manager,
InboundMessage(
channel="fake",
sender_id="user1",
chat_id="chat1",
content="/stop",
message_id="msg-stop-active",
),
)
ack = await asyncio.wait_for(bus.consume_outbound(), timeout=1.0)
assert ack.content == "Stopped."
assert ack.reply_to == "msg-stop-active"
assert display_mod.is_stream_cancel_requested(cancel_scope)
_run(_test())
def test_cancelled_wait_cleans_pending_response(self):
"""Cancelling a pending bus message should not leak its response slot."""
from EvoScientist.cli import channel as channel_mod
@@ -332,6 +421,173 @@ class TestBusInboundConsumer:
_run(_test())
def test_stop_during_hitl_wait_releases_wait_and_acks(self):
"""`/stop` should wake pending HITL wait and publish immediate ack."""
from EvoScientist.cli import channel as channel_mod
from EvoScientist.cli.channel import _bus_inbound_consumer, _message_queue
async def _test():
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
hitl_event = channel_mod._register_hitl_wait("fake", "chat1")
consumer = asyncio.create_task(_bus_inbound_consumer(bus, manager))
await bus.publish_inbound(
InboundMessage(
channel="fake",
sender_id="user1",
chat_id="chat1",
content="/stop",
message_id="m-stop-1",
)
)
for _ in range(20):
if hitl_event.is_set():
break
await asyncio.sleep(0.05)
assert hitl_event.is_set()
assert channel_mod._pop_hitl_reply("fake", "chat1") == "/stop"
outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=2.0)
assert outbound.content == "Stopped."
assert outbound.reply_to == "m-stop-1"
assert _message_queue.empty()
consumer.cancel()
try:
await consumer
except asyncio.CancelledError:
pass
_run(_test())
def test_stop_cancels_queued_request_before_main_thread_processes_it(self):
"""`/stop` should cancel a queued request instead of only acking."""
from EvoScientist.cli import channel as channel_mod
from EvoScientist.cli.channel import (
_claim_or_complete_channel_request,
_handle_bus_message,
_message_queue,
)
async def _test():
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
task = asyncio.create_task(
_handle_bus_message(
bus,
manager,
InboundMessage(
channel="fake",
sender_id="user1",
chat_id="chat1",
content="please work",
message_id="m-work-1",
),
)
)
queued = None
for _ in range(20):
with _message_queue.mutex:
queued = _message_queue.queue[0] if _message_queue.queue else None
if queued is not None:
break
await asyncio.sleep(0.05)
assert queued is not None
with channel_mod._response_lock:
assert queued.msg_id in channel_mod._pending_responses
await _handle_bus_message(
bus,
manager,
InboundMessage(
channel="fake",
sender_id="user1",
chat_id="chat1",
content="/stop",
message_id="m-stop-2",
),
)
with pytest.raises(asyncio.CancelledError):
await task
skipped = _message_queue.get_nowait()
assert skipped.msg_id == queued.msg_id
assert _claim_or_complete_channel_request(skipped) is False
with channel_mod._response_lock:
assert queued.msg_id not in channel_mod._pending_responses
with channel_mod._channel_request_lock:
assert queued.msg_id not in channel_mod._channel_requests
assert queued.msg_id not in channel_mod._cancelled_channel_messages
assert "fake:chat1" not in channel_mod._session_requests
outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=2.0)
assert outbound.content == "Stopped."
assert outbound.reply_to == "m-stop-2"
with pytest.raises(asyncio.TimeoutError):
await asyncio.wait_for(bus.consume_outbound(), timeout=0.2)
_run(_test())
def test_stop_leaves_resolved_response_available_for_delivery(self):
"""`/stop` must not steal a response whose waiter already resolved."""
from EvoScientist.cli import channel as channel_mod
from EvoScientist.cli.channel import (
ChannelMessage,
_cancel_channel_session,
_claim_channel_request,
_complete_channel_request,
_enqueue_channel_message,
_pop_channel_response,
_set_channel_response,
)
async def _test():
msg = ChannelMessage(
msg_id="msg-resolved",
content="already answered",
sender="user1",
channel_type="fake",
metadata={},
channel_ref=None,
bus_ref=None,
chat_id="chat1",
message_id="m-resolved",
)
waiter = _enqueue_channel_message(msg)
assert _claim_channel_request(msg) is True
_set_channel_response(msg.msg_id, "final answer")
assert await asyncio.wait_for(asyncio.shield(waiter), timeout=1.0) == (
"final answer"
)
cancelled_count, active_count = _cancel_channel_session("fake", "chat1")
assert cancelled_count == 0
assert active_count == 0
with channel_mod._response_lock:
assert msg.msg_id in channel_mod._pending_responses
with channel_mod._channel_request_lock:
assert msg.msg_id not in channel_mod._cancelled_channel_messages
assert _pop_channel_response(msg.msg_id) == "final answer"
_complete_channel_request(msg.msg_id)
_run(_test())
def test_message_counting(self):
"""Messages are counted via record_message."""
from EvoScientist.cli.channel import (
+112
View File
@@ -0,0 +1,112 @@
from __future__ import annotations
import asyncio
from types import SimpleNamespace
from unittest.mock import AsyncMock, patch
from EvoScientist.commands.channel_ui import ChannelCommandUI
from tests.conftest import run_async as _run
def _make_ui(callback=None, bus_ref=None):
captured: list[str] = []
ui = ChannelCommandUI(
SimpleNamespace(
channel_type="fake",
chat_id="chat-1",
message_id="msg-1",
metadata={},
bus_ref=bus_ref,
channel_ref=None,
),
append_system_callback=lambda text, style="dim": captured.append(text),
handle_session_resume_callback=callback,
)
return ui, captured
async def _run_resume(ui, thread_id: str, workspace_dir: str):
loop = asyncio.get_running_loop()
scheduled: list[asyncio.Task] = []
def _schedule(coro, _loop):
task = loop.create_task(coro)
scheduled.append(task)
return task
with (
patch("EvoScientist.cli.channel._bus_loop", new=loop),
patch(
"EvoScientist.commands.channel_ui.asyncio.run_coroutine_threadsafe",
side_effect=_schedule,
),
):
await ui.handle_session_resume(thread_id, workspace_dir)
if scheduled:
await asyncio.gather(*scheduled)
def _sent_text(bus_ref) -> str:
return "\n".join(
call.args[0].content for call in bus_ref.publish_outbound.await_args_list
)
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."),
]
with patch(
"EvoScientist.sessions.get_thread_messages",
new=AsyncMock(return_value=messages),
):
_run(_run_resume(ui, "thread-42", "/workspace"))
callback.assert_awaited_once_with("thread-42", "/workspace")
assert captured == []
text = _sent_text(bus_ref)
assert "Resumed session: thread-42" in text
assert "Conversation history:" in text
assert "User: How does this work?" in text
assert "EvoScientist: Here is the saved answer." in text
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)
with patch(
"EvoScientist.sessions.get_thread_messages",
new=AsyncMock(side_effect=RuntimeError("db locked")),
):
_run(_run_resume(ui, "thread-42", "/workspace"))
callback.assert_awaited_once_with("thread-42", "/workspace")
assert captured == []
text = _sent_text(bus_ref)
assert "Resumed session: thread-42" in text
assert "history unavailable: db locked" in text
def test_handle_session_resume_distinguishes_non_displayable_messages():
bus_ref = SimpleNamespace(publish_outbound=AsyncMock())
ui, captured = _make_ui(bus_ref=bus_ref)
with patch(
"EvoScientist.sessions.get_thread_messages",
new=AsyncMock(return_value=[SimpleNamespace(type="tool", content="hidden")]),
):
_run(_run_resume(ui, "thread-42", "/workspace"))
assert captured == [
"Resumed session: thread-42\nNo displayable messages in this session."
]
text = _sent_text(bus_ref)
assert "No displayable messages in this session." in text
+74 -1
View File
@@ -10,7 +10,10 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from EvoScientist.cli.channel import ChannelMessage
from EvoScientist.cli.channel import (
ChannelMessage,
_register_channel_request,
)
from EvoScientist.cli.commands import (
_make_serve_cmd_completed_hook,
_make_serve_start_new_session_cb,
@@ -117,6 +120,23 @@ def test_hook_updates_thread_id_on_resume():
assert holder["thread_id"] == "new-tid"
def test_hook_updates_workspace_dir_on_resume():
"""`/resume` can restore a different workspace; serve must adopt it."""
holder = {"agent": "a", "thread_id": "original-tid", "workspace_dir": "/old-ws"}
hook = _make_serve_cmd_completed_hook(holder)
ctx = MagicMock()
ctx.agent = "a"
ctx.thread_id = "new-tid"
ctx.workspace_dir = "/restored-ws"
cmd = MagicMock()
cmd.name = "/resume"
_run(hook(ctx, "a", cmd))
assert holder["workspace_dir"] == "/restored-ws"
def test_hook_syncs_channel_module_thread_id():
"""The bus reads ``cli.channel._cli_thread_id``; hook must sync it
alongside the holder update."""
@@ -274,6 +294,7 @@ def test_serve_process_message_reports_slash_dispatch_error_without_fallback():
patch("EvoScientist.cli.commands._set_channel_response") as mock_set_resp,
patch("EvoScientist.cli.tui_runtime.run_streaming") as mock_run_streaming,
):
_register_channel_request(msg)
_serve_process_message(
msg,
agent_holder=holder,
@@ -284,3 +305,55 @@ def test_serve_process_message_reports_slash_dispatch_error_without_fallback():
mock_set_resp.assert_called_once_with("msg-1", "Command error: slash broke")
mock_run_streaming.assert_not_called()
def test_serve_process_message_uses_runtime_workspace_from_holder():
"""After `/resume`, serve should use the adopted workspace, not startup ws."""
msg = ChannelMessage(
msg_id="msg-2",
content="hello",
sender="channel-user",
channel_type="imessage",
metadata={},
channel_ref=None,
bus_ref=None,
chat_id="channel-user",
message_id="ts-2",
)
holder = {
"agent": "agent",
"thread_id": "tid",
"workspace_dir": "/restored-workspace",
}
captured: dict[str, str] = {}
async def _fake_dispatch(*args, **kwargs):
captured["slash_workspace"] = kwargs["workspace_dir"]
return False
def _fake_build_metadata(workspace_dir: str, _model: str | None):
captured["meta_workspace"] = workspace_dir
return {}
with (
patch(
"EvoScientist.cli.commands.dispatch_channel_slash_command",
new=AsyncMock(side_effect=_fake_dispatch),
),
patch(
"EvoScientist.cli.commands.build_metadata",
side_effect=_fake_build_metadata,
),
patch("EvoScientist.cli.tui_runtime.run_streaming", return_value="ok"),
):
_register_channel_request(msg)
_serve_process_message(
msg,
agent_holder=holder,
model="model",
workspace_dir="/startup-workspace",
show_thinking=False,
)
assert captured["slash_workspace"] == "/restored-workspace"
assert captured["meta_workspace"] == "/restored-workspace"
+166
View File
@@ -0,0 +1,166 @@
"""Tests for channel-initiated stream cancellation."""
from __future__ import annotations
from unittest.mock import MagicMock, patch
import pytest
from EvoScientist.stream import display as display_mod
@pytest.fixture(autouse=True)
def _clean_cancel_event():
"""Ensure all stream-cancel scopes start clear for every test."""
with display_mod._stream_cancel_lock:
display_mod._stream_cancel_event.clear()
display_mod._stream_cancel_events.clear()
display_mod._stream_cancel_events[display_mod._DEFAULT_STREAM_CANCEL_SCOPE] = (
display_mod._stream_cancel_event
)
yield
with display_mod._stream_cancel_lock:
display_mod._stream_cancel_event.clear()
display_mod._stream_cancel_events.clear()
display_mod._stream_cancel_events[display_mod._DEFAULT_STREAM_CANCEL_SCOPE] = (
display_mod._stream_cancel_event
)
# ---------------------------------------------------------------------------
# 1. _consume breaks on cancel event
# ---------------------------------------------------------------------------
def test_consume_breaks_on_cancel_event():
"""After set(), ``_consume`` should stop pulling events and mark
``state.response_text`` with the ``[Stopped.]`` suffix."""
seen_events: list[int] = []
cancel_scope = "scope:consume"
async def _fake_stream(agent, message, thread_id, **kwargs):
for i in range(100):
if i == 3:
# Set during iteration — next loop iter should bail.
display_mod.request_stream_cancel(cancel_scope)
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,
)
# 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
# the generator before exit.
assert len(seen_events) <= 5
assert "[Stopped.]" in result
# ---------------------------------------------------------------------------
# 2. fresh _run_streaming clears stale set event
# ---------------------------------------------------------------------------
def test_run_streaming_short_circuits_when_scope_already_cancelled():
"""A queued request that is cancelled before start should stop immediately."""
seen_event = False
cancel_scope = "scope:queued"
async def _fake_stream(agent, message, thread_id, **kwargs):
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,
)
assert result == "[Stopped.]"
assert seen_event is False
assert not display_mod.is_stream_cancel_requested(cancel_scope)
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):
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",
)
assert "[Stopped.]" not in result
# ---------------------------------------------------------------------------
# 3. pending HITL/ask_user branches short-circuit when stop is requested
# ---------------------------------------------------------------------------
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):
if False:
yield {}
state = display_mod.StreamState()
state.response_text = "Partial answer"
state.pending_interrupt = {
"action_requests": [{"name": "execute", "args": {"command": "echo hi"}}]
}
display_mod.request_stream_cancel("scope:hitl")
prompt_called = False
def _prompt(_requests):
nonlocal prompt_called
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,
)
assert result == "Partial answer\n[Stopped.]"
assert prompt_called is False
+90
View File
@@ -0,0 +1,90 @@
"""Tests for TUI command-completion state sync."""
from types import SimpleNamespace
import pytest
from EvoScientist.commands.base import CommandContext
from tests.conftest import run_async as _run
pytest.importorskip("textual")
class _Loader:
def __init__(self) -> None:
self.adopt_calls: list[object] = []
def adopt(self, agent: object) -> None:
self.adopt_calls.append(agent)
class _StubApp:
def __init__(self) -> None:
self._agent_loader = _Loader()
self._conversation_tid = "thread-1"
self.model_updates: list[tuple[str, str | None]] = []
self.refresh_calls: list[bool] = []
def update_status_after_model_change(
self,
new_model: str,
new_provider: str | None = None,
) -> None:
self.model_updates.append((new_model, new_provider))
async def _refresh_status_snapshot(
self,
*,
reset_streaming_text: bool = True,
) -> None:
self.refresh_calls.append(reset_streaming_text)
def test_sync_tui_command_completion_adopts_agent_swap(monkeypatch):
import EvoScientist.cli.tui_interactive as tui_mod
from EvoScientist import EvoScientist as evosci_mod
app = _StubApp()
ctx = CommandContext(
agent="new-agent",
thread_id="thread-1",
ui=SimpleNamespace(),
)
cmd = SimpleNamespace(name="/model")
monkeypatch.setattr(
evosci_mod,
"_ensure_config",
lambda: SimpleNamespace(model="gpt-5.5", provider="openai"),
)
monkeypatch.setattr(tui_mod, "_channels_is_running", lambda: True)
monkeypatch.setattr(tui_mod._ch_mod, "_cli_agent", "old-agent", raising=False)
monkeypatch.setattr(tui_mod._ch_mod, "_cli_thread_id", "old-thread", raising=False)
_run(tui_mod._sync_tui_command_completion(app, ctx, "old-agent", cmd))
assert app._agent_loader.adopt_calls == ["new-agent"]
assert app.model_updates == [("gpt-5.5", "openai")]
assert app.refresh_calls == [True]
assert tui_mod._ch_mod._cli_agent == "new-agent"
assert tui_mod._ch_mod._cli_thread_id == "thread-1"
def test_sync_tui_command_completion_refreshes_without_agent_swap(monkeypatch):
import EvoScientist.cli.tui_interactive as tui_mod
app = _StubApp()
ctx = CommandContext(
agent="same-agent",
thread_id="thread-1",
ui=SimpleNamespace(),
)
cmd = SimpleNamespace(name="/compact")
monkeypatch.setattr(tui_mod, "_channels_is_running", lambda: False)
_run(tui_mod._sync_tui_command_completion(app, ctx, "same-agent", cmd))
assert app._agent_loader.adopt_calls == []
assert app.model_updates == []
assert app.refresh_calls == [True]
+20
View File
@@ -151,6 +151,26 @@ class TestSummarizationStateMachine(unittest.TestCase):
assert _should_finalize_active_summarization(event_type) is True
class TestStoppedResponseText(unittest.TestCase):
"""Stopped-response text normalization."""
def test_trims_before_appending_marker(self):
from EvoScientist.stream.display import build_stopped_response_text
current, final_text = build_stopped_response_text("partial answer \n")
assert current == "partial answer"
assert final_text == "partial answer\n[Stopped.]"
def test_does_not_duplicate_marker(self):
from EvoScientist.stream.display import build_stopped_response_text
current, final_text = build_stopped_response_text("partial\n[Stopped.]")
assert current == "partial\n[Stopped.]"
assert final_text == "partial\n[Stopped.]"
@unittest.skipUnless(_has_textual, "textual not installed")
class TestAssistantMessage(unittest.TestCase):
"""AssistantMessage construction."""