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