From db1abce8d8d5e64a91b291dce9cd445d3d3e7105 Mon Sep 17 00:00:00 2001 From: dinos Date: Tue, 14 Jul 2026 16:59:41 +0200 Subject: [PATCH] refactor: extract a shared HITL/ask_user interaction engine (#342) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * chore: add pytest-asyncio in auto mode * test: migrate channel and stream tests to native async Convert run_async() wrapper tests to plain 'async def test_*' under pytest-asyncio auto mode. collect_events() in stream_v3_fakes becomes a coroutine awaited at every call site. * test: migrate command and model/middleware tests to native async Convert run_async() wrappers (import, alias, and fixture forms) to plain 'async def test_*'. Multi-call tests merge onto one loop as sequential awaits; none asserted on loop identity. * test: migrate TUI, notifier, gateway, and session tests to native async TUI/notifier/gateway files convert run_async wrappers to plain async tests. test_sessions.py's unittest.TestCase classes move to unittest.IsolatedAsyncioTestCase (pytest-asyncio does not await async methods on plain TestCase; converting blindly would have made ~70 tests silently vacuous). Its setUpClass keeps a one-shot asyncio.run() since IsolatedAsyncioTestCase has no async class-level hook. TestLoadingWidget in test_tui_widgets.py drops its TestCase base for the same reason. * test: replace direct asyncio.run() calls with native async tests Convert tests that called asyncio.run() (directly or via a local _run helper) to plain 'async def test_*'; delete the local helpers. * test: drop undeclared anyio markers and delete run_async helper The @pytest.mark.anyio tests relied on anyio being a transitive dep of httpx; auto-mode pytest-asyncio collects them natively. run_async() and its fixture are unreferenced after the migration, so remove them — pytest-asyncio's per-test loop teardown covers the pending-task cancellation the helper existed for (verified: full suite runs with no 'Event loop is closed' errors or destroyed-task warnings). * test: add autouse fixture for watcher cleanup * refactor: remove redundant hasattr calls * refactor: extract shared HITL/ask_user interaction grammar Extract prompt/question formatting, the reply grammar (approval letters, ask_user choice letters + the 'Other' sub-flow, stop-commands), the ApprovalPolicy (config auto-approve rule + session registry + session-key derivation), per-flow timeout constants, and the bilingual feedback strings into channels/interaction.py. Both drivers now point at the shared functions: this reverses cli/channel.py's imports of consumer privates and closes the /stop drift at the parsing layer (serve-mode ask_user now checks stop-commands before parsing an answer, matching the CLI path). * refactor: add interaction engine + registry; port InboundConsumer Introduce InteractionIO (transport adapter Protocol), PendingReplyRegistry (one asyncio-based reply router per process), and the engine coroutines resolve_ask_user / resolve_approval in channels/interaction.py. Port InboundConsumer onto them: a _ConsumerIO adapter over bus.publish_outbound + the registry, one ApprovalPolicy replacing the config/session auto-approve checks, and a single reply-interception point (registry.try_resolve) replacing the parallel ask_user/HITL pending dicts. _resolve_ask_user and the approval section of _stream_with_hitl are now thin engine calls. Behavior unification (serve mode): an unrecognized HITL reply now declines with the shared 'Unrecognized reply' notice instead of rejecting-and-refeeding as a fresh turn, and /stop mid-approval cancels cleanly — both via the shared parser. * refactor: port CLI channel bridge onto the interaction engine Replace the ~250-line parallel bodies of channel_ask_user_prompt / channel_hitl_prompt with thin bridges that run resolve_ask_user / resolve_approval on the bus loop via run_coroutine_threadsafe(...).result() (outer = engine per-flow timeout + slack, so the engine's own timeout fires first). The 15s send timeout moves into the _BridgeIO adapter. Delete the _pending_hitl / _hitl_lock / _hitl_auto_approve module globals and the _register_hitl_wait / _try_set_hitl_reply / _pop_hitl_reply helpers, absorbed by one bus-loop PendingReplyRegistry + one ApprovalPolicy. The bus consumer feeds the registry via try_resolve ahead of normal enqueue. * refactor: restore serve-mode refeed for unrecognized HITL replies Gate-review fix: the engine no longer decides transport policy for unparseable approval replies. resolve_approval now returns an ApprovalOutcome carrying unrecognized_reply (raw text) when parsing fails, sending no feedback itself; recognized reject keeps the sharedi rejection message. Consumer driver (serve mode) restores the pre-engine semantics: an unrecognized reply rejects the pending action, confirms with the rejection message, and the text is re-dispatched as a NEW agent turn — _stream_with_hitl returns the captured text and _handle_message starts the refeed turn only after the current one has released the chat lock (old fall-through ordering). CLI bridge keeps its old no-refeed path byte-for-byte: 'Unrecognized reply. Action rejected.' and decline. Tests: serve refeed pinned end-to-end (prompt → unrecognized text → rejection feedback → text reaches the stream path as a new turn), CLI no-refeed pinned (notice sent, nothing enqueued), engine test updated to assert the outcome struct with no engine-side feedback. * refactor: polish the interaction engine surface - English feedback strings (Approved / Rejected / auto-approving) - drop the consumer's backwards-compatible re-exports and both modules' private timeout aliases; callers use the canonical interaction names - replace byte-for-byte prompt goldens with structural format tests and assert feedback via the shared constants instead of string literals - strip audit/design shorthand (R1/R2/G3, stage numbers) from comments * fix: propagate pending reply task cancellation * fix: preserve reply context when refeeding HITL replies * fix: honor HITL session grants without bus loop * fix: bound bridge waits by send latency * fix: handle empty ask_user replies explicitly * chore: remove stale interaction helpers * fix: harden interaction engine reply edge cases Review follow-ups on the interaction engine: - normalize ask_user choices before .get(): the tool args come from model JSON and only presence is validated, so plain-string choices must render and parse instead of crashing the turn - treat only None as an approval timeout, so an empty/media-only reply flows through the unrecognized path and serve mode refeeds it with its preserved context - intercept prompt replies before _get_thread_id so a consumed reply cannot create an orphan graph thread or touch the sender-session LRU - close engine coroutines the bridge failed to schedule (no bus loop / scheduling error) to avoid never-awaited warnings - clear pending-response and channel-request state in the bridge test fixture --------- Co-authored-by: Xi Zhang <106144707+X-iZhang@users.noreply.github.com> --- EvoScientist/channels/consumer.py | 509 +++++++----------------- EvoScientist/channels/interaction.py | 540 +++++++++++++++++++++++++ EvoScientist/channels/qq/channel.py | 6 +- EvoScientist/cli/channel.py | 453 ++++++++++----------- tests/test_bus_integration.py | 29 +- tests/test_cli_channel_bridge.py | 426 ++++++++++++++++++++ tests/test_hitl.py | 126 +++--- tests/test_interaction_engine.py | 563 +++++++++++++++++++++++++++ tests/test_interaction_grammar.py | 341 ++++++++++++++++ 9 files changed, 2296 insertions(+), 697 deletions(-) create mode 100644 EvoScientist/channels/interaction.py create mode 100644 tests/test_cli_channel_bridge.py create mode 100644 tests/test_interaction_engine.py create mode 100644 tests/test_interaction_grammar.py diff --git a/EvoScientist/channels/consumer.py b/EvoScientist/channels/consumer.py index f411552..d426094 100644 --- a/EvoScientist/channels/consumer.py +++ b/EvoScientist/channels/consumer.py @@ -20,6 +20,17 @@ from ..gateway import GraphGateway, GraphRunInput, GraphTarget, RunRequest from .base import Channel from .bus import MessageBus from .bus.events import InboundMessage, OutboundMessage +from .capabilities import ChannelCapabilities +from .interaction import ( + ASK_USER_TIMEOUT, + HITL_APPROVAL_TIMEOUT, + REJECTED_FEEDBACK, + ApprovalPolicy, + InteractionIO, + PendingReplyRegistry, + resolve_approval, + resolve_ask_user, +) logger = logging.getLogger(__name__) @@ -28,10 +39,6 @@ T = TypeVar("T") _MAX_CHAT_LOCKS = 10_000 _MAX_SESSIONS = 10_000 _MAX_HITL_ROUNDS = 50 -_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) -) @dataclass @@ -108,120 +115,56 @@ def _join_subagent_text(buffers: dict[str, tuple[str, list[str]]]) -> str: return "\n\n".join(sections) -def _should_auto_approve(action_requests: list[dict]) -> bool: - """Check if all action requests can be auto-approved via config. +class _ConsumerIO(InteractionIO): + """:class:`InteractionIO` over the consumer's bus + reply registry. - Returns True if no manual approval is needed (config auto_approve, - non-execute tools, or shell_allow_list match). + Publishes prompts through ``bus.publish_outbound`` and blocks for + replies on the consumer's shared :class:`PendingReplyRegistry` — both + on the consumer's own event loop, so the engine runs natively async + here with no thread hand-off. """ - if not action_requests: + + def __init__( + self, consumer: InboundConsumer, msg: InboundMessage, session_key: str + ) -> None: + self._consumer = consumer + self._msg = msg + self._session_key = session_key + self._last_reply_message: InboundMessage | None = None + channel = consumer._get_channel(msg.channel) + self.capabilities = ( + channel.capabilities if channel is not None else ChannelCapabilities() + ) + self.base_metadata = msg.metadata + + async def send(self, content: str, *, metadata: dict | None = None) -> bool: + await self._consumer.bus.publish_outbound( + OutboundMessage( + channel=self._msg.channel, + chat_id=self._msg.chat_id, + content=content, + metadata=metadata if metadata is not None else self._msg.metadata, + ) + ) return True - try: - from ..config.settings import HITL_SHELL_TOOLS, load_config + async def wait_reply(self, *, timeout: float) -> str | None: + reply = await self._consumer._reply_registry.wait_event( + self._session_key, timeout + ) + if reply is None: + self._last_reply_message = None + return None + self._last_reply_message = ( + reply.context if isinstance(reply.context, InboundMessage) else None + ) + return reply.content - cfg = load_config() - except Exception: - return False # fail-closed - - if cfg.auto_approve: - return True - - shell_allow_list = ( - [s.strip() for s in cfg.shell_allow_list.split(",") if s.strip()] - if cfg.shell_allow_list - else [] - ) - - for req in action_requests: - name = req.get("name", "") - if name not in HITL_SHELL_TOOLS: - continue - args = req.get("args", {}) - command = args.get("command", "") if isinstance(args, dict) else "" - cmd = command.strip() - if not any(cmd.startswith(prefix) for prefix in shell_allow_list): - return False - return True - - -def _format_approval_prompt( - action_requests: list[dict], *, with_buttons: bool = False -) -> str: - """Format an approval prompt as a text message for channel users. - - When *with_buttons* is True, the trailing "Reply: 1=Approve..." - instruction is dropped — the buttons replace the textual cue. - """ - lines = ["\u26a0\ufe0f Approval Required\n"] - for i, req in enumerate(action_requests, 1): - name = req.get("name", "") - args = req.get("args", {}) - if isinstance(args, dict): - command = args.get("command", args.get("path", "")) - else: - command = "" - if command: - lines.append(f" {i}. {name}: {command}") - else: - lines.append(f" {i}. {name}") - if not with_buttons: - lines.append("") - lines.append("Reply: 1=Approve, 2=Reject, 3=Approve all") - lines.append("(Auto-reject in 2 min if no reply)") - return "\n".join(lines) - - -def _parse_approval_reply(text: str) -> str | None: - """Parse a channel user's reply as an approval decision. - - Returns "approve", "reject", "auto", or None if not recognized. - """ - t = text.strip().lower() - if t in ("1", "y", "yes", "approve", "ok"): - return "approve" - if t in ("2", "n", "no", "reject"): - return "reject" - if t in ("3", "a", "auto", "approve all"): - return "auto" - return None - - -def _approval_prompt_metadata( - base_metadata: dict | None, *, with_buttons: bool -) -> dict: - """Outbound metadata for the HITL approval prompt. - - When *with_buttons* is True, attaches Approve/Reject/Auto buttons whose - values match ``_parse_approval_reply`` so a click flows through the same - path as a typed ``"1"``/``"2"``/``"3"`` reply. - """ - metadata = dict(base_metadata or {}) - if with_buttons: - metadata["buttons"] = [ - {"text": "Approve", "value": "1", "type": "primary"}, - {"text": "Reject", "value": "2", "type": "danger"}, - {"text": "Approve all", "value": "3"}, - ] - return metadata - - -@dataclass -class _PendingInterrupt: - """Stored state for a pending HITL interrupt awaiting channel user reply.""" - - thread_id: str - action_requests: list - event: asyncio.Event # set when user replies - decision: str | None = None # "approve", "reject", "auto" - - -@dataclass -class _PendingAskUserReply: - """Stored state for a pending ask_user question awaiting channel user reply.""" - - event: asyncio.Event # set when user replies - reply: str | None = None # raw reply text + def take_reply_context(self) -> InboundMessage | None: + """Consume the last inbound reply context captured by ``wait_reply``.""" + msg = self._last_reply_message + self._last_reply_message = None + return msg class InboundConsumer: @@ -310,12 +253,12 @@ class InboundConsumer: # Metrics self._metrics = ConsumerMetrics() - # HITL: pending interrupts per session_key, and auto-approve sessions - self._pending_interrupts: dict[str, _PendingInterrupt] = {} - self._auto_approve_sessions: set[str] = set() - - # ask_user: pending reply per session_key - self._pending_ask_user_replies: dict[str, _PendingAskUserReply] = {} + # Interaction engine state: one reply registry (routes the next + # message from a chat into a waiting prompt) and one approval + # policy (config rule + session "Approve all" grants), shared by + # the ask_user and HITL flows via ``channels.interaction``. + self._reply_registry = PendingReplyRegistry() + self._approval_policy = ApprovalPolicy() async def _get_thread_id(self, sender_id: str) -> str: """Get or create a thread ID for the given sender. @@ -428,8 +371,6 @@ class InboundConsumer: except Exception: pass - channel = self._get_channel(msg.channel) - thread_id = await self._get_thread_id(msg.sender_id) session_key = msg.session_key # "channel:chat_id" # Lazily create per-chat lock; evict stale locks when too many @@ -440,29 +381,39 @@ class InboundConsumer: self._metrics.total_processed += 1 - # ask_user: check if this message is a reply to a pending question. - # Must be checked BEFORE HITL approval — any text is a valid answer. - if session_key in self._pending_ask_user_replies: - pending_ask = self._pending_ask_user_replies[session_key] - pending_ask.reply = msg.content - pending_ask.event.set() - return # consumed as ask_user answer + # Reply interception: if a prompt (ask_user question or HITL + # approval) is waiting on this chat, hand it this message instead + # of starting a fresh agent turn. The engine parses it (stop / + # cancel / choice / approval grammar), so the registry only routes + # text plus the original inbound context — one path for both flows. + if self._reply_registry.try_resolve(session_key, msg.content, context=msg): + return - # HITL: check if this message is a reply to a pending approval - if session_key in self._pending_interrupts: - pending = self._pending_interrupts[session_key] - decision = _parse_approval_reply(msg.content) - if decision is not None: - pending.decision = decision - pending.event.set() - return # don't process as a new agent message - # Unrecognized reply — treat as new message, cancel pending - pending.decision = "reject" - pending.event.set() - del self._pending_interrupts[session_key] + # Resolved only for real agent turns — a consumed prompt reply must + # not create a graph thread or touch the sender-session LRU. + channel = self._get_channel(msg.channel) + thread_id = await self._get_thread_id(msg.sender_id) async with self._chat_locks[session_key]: - await self._stream_with_hitl(msg, channel, thread_id, session_key) + refeed = await self._stream_with_hitl(msg, channel, thread_id, session_key) + + # An unrecognized reply to a pending approval rejects the action and + # then becomes a new agent turn. The lock was released above, so the + # previous turn has fully unwound before the refeed turn acquires it. + # Loops in case the refeed turn hits another approval that is again + # answered with unparseable text. + while refeed is not None: + channel = self._get_channel(refeed.channel) + thread_id = await self._get_thread_id(refeed.sender_id) + session_key = refeed.session_key + if session_key not in self._chat_locks: + self._chat_locks[session_key] = asyncio.Lock() + if len(self._chat_locks) > _MAX_CHAT_LOCKS: + self._evict_chat_locks() + async with self._chat_locks[session_key]: + refeed = await self._stream_with_hitl( + refeed, channel, thread_id, session_key + ) async def _stream_with_hitl( self, @@ -470,8 +421,13 @@ class InboundConsumer: channel: Channel | None, thread_id: str, session_key: str, - ) -> None: - """Stream agent events with HITL interrupt handling.""" + ) -> InboundMessage | None: + """Stream agent events with HITL interrupt handling. + + Returns ``None`` normally. When a pending approval is answered + with unrecognized text, returns the intercepted inbound reply so the + caller can refeed it as a new agent turn after this one unwinds. + """ from langgraph.types import Command stream_input: GraphRunInput = msg.content @@ -609,109 +565,39 @@ class InboundConsumer: stream_input = Command(resume=result) continue - # HITL: resolve the interrupt + # HITL: resolve the interrupt through the shared engine. + # ``resolve_approval`` handles session/config auto-approve, + # the approval prompt (with capability-driven buttons), the + # reply wait, parsing (incl. /stop), and feedback strings. action_reqs = interrupt_data.get("action_requests", []) - n = len(action_reqs) or 1 - - # Session auto-approve (user previously chose "Approve all") - if session_key in self._auto_approve_sessions: - stream_input = Command( - resume={"decisions": [{"type": "approve"} for _ in range(n)]} - ) - continue - - # Config auto-approve (auto_approve, non-execute, allow_list) - if _should_auto_approve(action_reqs): - stream_input = Command( - resume={"decisions": [{"type": "approve"} for _ in range(n)]} - ) - continue - - # Needs user approval — send prompt to channel - has_buttons = ( - channel is not None and channel.capabilities.inline_buttons + io = _ConsumerIO(self, msg, session_key) + outcome = await resolve_approval( + action_reqs, + io, + self._approval_policy, + session_key, + timeout=HITL_APPROVAL_TIMEOUT, ) - prompt_text = _format_approval_prompt( - action_reqs, with_buttons=has_buttons - ) - approval_metadata = _approval_prompt_metadata( - msg.metadata, with_buttons=has_buttons - ) - await self.bus.publish_outbound( - OutboundMessage( - channel=msg.channel, - chat_id=msg.chat_id, - content=prompt_text, - metadata=approval_metadata, - ) - ) - - # Wait for user reply - pending = _PendingInterrupt( - thread_id=thread_id, - action_requests=action_reqs, - event=asyncio.Event(), - ) - self._pending_interrupts[session_key] = pending - - timed_out = False - try: - await asyncio.wait_for( - pending.event.wait(), - timeout=_HITL_APPROVAL_TIMEOUT, - ) - except TimeoutError: - timed_out = True - finally: - # Unregister BEFORE any further await so a late reply can't flip - # the decision back to approve during the notification round-trip. - self._pending_interrupts.pop(session_key, None) - - if timed_out: - # Reject on timeout (fail-closed; matches cli/channel.py). Decision - # is a local constant, not pending.decision, so it can't be - # overwritten by a late reply after we unregistered above. - decision = "reject" - await self.bus.publish_outbound( - OutboundMessage( - channel=msg.channel, - chat_id=msg.chat_id, - content="⏰ Approval timed out. Action rejected.", - metadata=msg.metadata, + if outcome.unrecognized_reply is not None: + # Serve-mode policy: an unrecognized reply rejects the + # pending action, confirms with reject feedback, and is + # then processed as a new agent turn. The refeed is + # returned to ``_handle_message`` so chat-lock ordering + # stays serialized. + await io.send(REJECTED_FEEDBACK) + # In this flow, the final wait_reply call is exactly the + # unrecognized approval reply. ask_user does not read this. + refeed_msg = io.take_reply_context() + if refeed_msg is None: + logger.warning( + "Unrecognized approval reply had no inbound context; " + "dropping refeed" ) - ) - else: - decision = pending.decision or "reject" + return refeed_msg + if outcome.decisions is None: + return None # reject / timeout / stop — end the turn - # Visible confirmation so the click/reply registers (QQ has no - # message recall API for C2C). Only fires when the user - # actually responded — silent on timeout to avoid claiming - # the user approved when they just walked away. - if pending.event.is_set(): - feedback_text = { - "approve": "\u2705 已批准", - "auto": "\u2705 已批准(后续自动通过)", - "reject": "\u274c 已拒绝", - }.get(decision) - if feedback_text: - await self.bus.publish_outbound( - OutboundMessage( - channel=msg.channel, - chat_id=msg.chat_id, - content=feedback_text, - metadata=msg.metadata, - ) - ) - - if decision == "reject": - return - - if decision == "auto": - self._auto_approve_sessions.add(session_key) - - stream_input = Command( - resume={"decisions": [{"type": "approve"} for _ in range(n)]} - ) + stream_input = Command(resume={"decisions": outcome.decisions}) # continue to next HITL round except TimeoutError: @@ -773,150 +659,27 @@ class InboundConsumer: # ── ask_user helpers ── - async def _wait_for_ask_user_reply( - self, - session_key: str, - timeout: float, - ) -> str | None: - """Register a pending ask_user slot and wait for the user to reply. - - Returns the raw reply text, or ``None`` on timeout. - """ - pending = _PendingAskUserReply(event=asyncio.Event()) - self._pending_ask_user_replies[session_key] = pending - try: - await asyncio.wait_for(pending.event.wait(), timeout=timeout) - except TimeoutError: - pass - finally: - self._pending_ask_user_replies.pop(session_key, None) - return pending.reply - async def _resolve_ask_user( self, msg: InboundMessage, event_data: dict, session_key: str, ) -> dict: - """Handle an ask_user interrupt: send questions to channel, collect answers. + """Handle an ask_user interrupt via the shared engine. - Mirrors the logic of ``cli.channel.channel_ask_user_prompt`` but runs - fully async inside the consumer event loop. + Delegates the whole question/answer flow (prompt formatting, choice + + "Other" grammar, ``/stop`` handling) to + :func:`channels.interaction.resolve_ask_user` over a + :class:`_ConsumerIO` adapter, so serve mode and the CLI bridge + cannot drift. Returns a dict suitable for ``Command(resume=...)``: ``{"answers": [...], "status": "answered"}`` or ``{"status": "cancelled"}``. """ questions = event_data.get("questions", []) - if not questions: - return {"answers": [], "status": "answered"} - - total = len(questions) - answers: list[str] = [] - - for i, q in enumerate(questions): - q_text = q.get("question", "") - q_type = q.get("type", "text") - required = q.get("required", True) - - # -- Format question header -- - if total == 1: - header = "\u2753 Quick check-in from EvoScientist\n" - else: - header = f"\u2753 Question {i + 1}/{total}\n" - - lines: list[str] = [header, f"{i + 1}. {q_text}"] - if not required: - lines[-1] += " (optional)" - - if q_type == "multiple_choice": - choices = q.get("choices", []) - for j, choice in enumerate(choices): - label = choice.get("value", str(choice)) - letter = chr(ord("A") + j) - lines.append(f" {letter}. {label}") - other_letter = chr(ord("A") + len(choices)) - lines.append(f" {other_letter}. Other") - letters = "/".join(chr(ord("A") + k) for k in range(len(choices) + 1)) - lines.append(f"\nReply with a letter ({letters}), or 'cancel'.") - else: - skip_hint = " Leave empty to skip." if not required else "" - lines.append(f"\nReply with your answer, or 'cancel'.{skip_hint}") - - # -- Send question -- - await self.bus.publish_outbound( - OutboundMessage( - channel=msg.channel, - chat_id=msg.chat_id, - content="\n".join(lines), - metadata=msg.metadata, - ) - ) - - # -- Wait for user reply -- - reply = await self._wait_for_ask_user_reply( - session_key, - _ASK_USER_TIMEOUT, - ) - - if not reply: - await self.bus.publish_outbound( - OutboundMessage( - channel=msg.channel, - chat_id=msg.chat_id, - content="\u23f0 Response timed out.", - metadata=msg.metadata, - ) - ) - return {"status": "cancelled"} - - raw = reply.strip() - if raw.lower() == "cancel": - return {"status": "cancelled"} - - # -- Parse answer -- - if q_type == "multiple_choice": - choices = q.get("choices", []) - other_letter = chr(ord("A") + len(choices)) - if len(raw) == 1 and raw.upper() == other_letter: - # "Other" selected — ask for free-form input - await self.bus.publish_outbound( - OutboundMessage( - channel=msg.channel, - chat_id=msg.chat_id, - content="Please type your answer:", - metadata=msg.metadata, - ) - ) - other_reply = await self._wait_for_ask_user_reply( - session_key, - _ASK_USER_TIMEOUT, - ) - if not other_reply: - await self.bus.publish_outbound( - OutboundMessage( - channel=msg.channel, - chat_id=msg.chat_id, - content="\u23f0 Response timed out.", - metadata=msg.metadata, - ) - ) - return {"status": "cancelled"} - if other_reply.strip().lower() == "cancel": - return {"status": "cancelled"} - answers.append(other_reply.strip()) - elif len(raw) == 1 and raw.upper().isalpha(): - idx = ord(raw.upper()) - ord("A") - if 0 <= idx < len(choices): - answers.append(choices[idx].get("value", raw)) - else: - answers.append(raw) - else: - answers.append(raw) - else: - answers.append(raw) - - return {"answers": answers, "status": "answered"} + io = _ConsumerIO(self, msg, session_key) + return await resolve_ask_user(questions, io, timeout=ASK_USER_TIMEOUT) # ── internal ── diff --git a/EvoScientist/channels/interaction.py b/EvoScientist/channels/interaction.py new file mode 100644 index 0000000..4882c7e --- /dev/null +++ b/EvoScientist/channels/interaction.py @@ -0,0 +1,540 @@ +"""Transport-agnostic HITL and ask_user interaction engine. + +The module defines the channel-side protocol shared by serve mode and the +CLI/TUI bridge: prompt formatting, reply grammar, stop handling, approval +policy, pending-reply routing, and the async engine coroutines for approval +and ask_user flows. Drivers provide transport-specific IO through +:class:`InteractionIO`. +""" + +from __future__ import annotations + +import asyncio +from dataclasses import dataclass +from typing import TYPE_CHECKING, Protocol + +if TYPE_CHECKING: + from .capabilities import ChannelCapabilities + +# ── timeout constants ────────────────────────────────────────────────── +# Per-flow defaults. HITL approval is short (a yes/no gate); ask_user is +# longer because the human may need thinking time. +HITL_APPROVAL_TIMEOUT = 120.0 # seconds to wait for a HITL approval reply +ASK_USER_TIMEOUT = 300.0 # seconds to wait for an ask_user reply + +# ── stop-command grammar ────────────────────────────────────────── +# Checked before reply parsing in *both* flows so a `/stop` mid-prompt +# always cancels instead of being captured as a literal answer. +_STOP_COMMANDS = frozenset(("/stop", "/cancel")) + +# ── feedback strings ───────────────────────────────────────── +# Visible confirmations so a click/reply registers on channels without a +# message-recall API (e.g. QQ C2C). +APPROVED_FEEDBACK = "✅ Approved" +APPROVED_AUTO_FEEDBACK = "✅ Approved (auto-approving future actions)" +REJECTED_FEEDBACK = "❌ Rejected" +UNRECOGNIZED_FEEDBACK = "Unrecognized reply. Action rejected." +APPROVAL_TIMEOUT_FEEDBACK = "⏰ Approval timed out. Action rejected." +ASK_USER_TIMEOUT_FEEDBACK = "⏰ Response timed out." +OTHER_PROMPT = "Please type your answer:" + +# ── stop / cancel helpers ────────────────────────────────────────────── + + +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 is_cancel_reply(content: str | None) -> bool: + """Whether a reply is the literal ``cancel`` sentinel (case-insensitive).""" + return (content or "").strip().lower() == "cancel" + + +# ── approval reply grammar ───────────────────────────────────────────── + + +def parse_approval_reply(text: str) -> str | None: + """Parse a channel user's reply as an approval decision. + + Returns "approve", "reject", "auto", or None if not recognized. + """ + t = text.strip().lower() + if t in ("1", "y", "yes", "approve", "ok"): + return "approve" + if t in ("2", "n", "no", "reject"): + return "reject" + if t in ("3", "a", "auto", "approve all"): + return "auto" + return None + + +def approve_decisions(action_requests: list) -> list[dict]: + """Build the ``decisions`` payload that approves every action request. + + Length matches ``action_requests`` (with a floor of 1, matching the + consumer's historical ``len(...) or 1`` so an empty request list still + yields a single approve — the shape ``Command(resume=...)`` expects). + """ + n = len(action_requests) or 1 + return [{"type": "approve"} for _ in range(n)] + + +# ── approval prompt formatting ───────────────────────────────────────── + + +def format_approval_prompt( + action_requests: list[dict], *, with_buttons: bool = False +) -> str: + """Format an approval prompt as a text message for channel users. + + When *with_buttons* is True, the trailing "Reply: 1=Approve..." + instruction is dropped — the buttons replace the textual cue. + """ + lines = ["⚠️ Approval Required\n"] + for i, req in enumerate(action_requests, 1): + name = req.get("name", "") + args = req.get("args", {}) + if isinstance(args, dict): + command = args.get("command", args.get("path", "")) + else: + command = "" + if command: + lines.append(f" {i}. {name}: {command}") + else: + lines.append(f" {i}. {name}") + if not with_buttons: + lines.append("") + lines.append("Reply: 1=Approve, 2=Reject, 3=Approve all") + lines.append("(Auto-reject in 2 min if no reply)") + return "\n".join(lines) + + +def approval_prompt_metadata(base_metadata: dict | None, *, with_buttons: bool) -> dict: + """Outbound metadata for the HITL approval prompt. + + When *with_buttons* is True, attaches Approve/Reject/Auto buttons whose + values match ``parse_approval_reply`` so a click flows through the same + path as a typed ``"1"``/``"2"``/``"3"`` reply. + """ + metadata = dict(base_metadata or {}) + if with_buttons: + metadata["buttons"] = [ + {"text": "Approve", "value": "1", "type": "primary"}, + {"text": "Reject", "value": "2", "type": "danger"}, + {"text": "Approve all", "value": "3"}, + ] + return metadata + + +# ── ask_user question formatting & answer grammar ────────────────────── + + +def _choice_value(choice: object, fallback: str = "") -> str: + """Normalize one ask_user choice to its display/answer string. + + Choices arrive from model-produced tool args; the schema says dicts with + a ``value`` key, but nothing enforces that at runtime, so plain strings + (or anything else) must not crash the prompt. + """ + if isinstance(choice, dict): + return str(choice.get("value", fallback or choice)) + return str(choice) + + +def format_question_prompt(question: dict, index: int, total: int) -> str: + """Format one ask_user *question* as a channel message. + + *index* is 0-based; *total* is the number of questions in the batch. + """ + q_text = question.get("question", "") + q_type = question.get("type", "text") + required = question.get("required", True) + + if total == 1: + header = "❓ Quick check-in from EvoScientist\n" + else: + header = f"❓ Question {index + 1}/{total}\n" + + lines: list[str] = [header, f"{index + 1}. {q_text}"] + if not required: + lines[-1] += " (optional)" + + if q_type == "multiple_choice": + choices = question.get("choices", []) + for j, choice in enumerate(choices): + label = _choice_value(choice) + letter = chr(ord("A") + j) + lines.append(f" {letter}. {label}") + other_letter = chr(ord("A") + len(choices)) + lines.append(f" {other_letter}. Other") + letters = "/".join(chr(ord("A") + k) for k in range(len(choices) + 1)) + lines.append(f"\nReply with a letter ({letters}), or 'cancel'.") + else: + skip_hint = " Leave empty to skip." if not required else "" + lines.append(f"\nReply with your answer, or 'cancel'.{skip_hint}") + return "\n".join(lines) + + +def parse_choice_answer(raw: str, choices: list) -> tuple[str, str | None]: + """Classify a multiple-choice reply. + + Returns ``(kind, value)``: + + * ``("other", None)`` — the "Other" letter was chosen; the caller must + run the free-form sub-flow (send :data:`OTHER_PROMPT`, wait again). + * ``("answer", value)`` — a resolved answer string (the chosen + choice's ``value``, or the raw text when it isn't a valid letter). + """ + other_letter = chr(ord("A") + len(choices)) + if len(raw) == 1 and raw.upper() == other_letter: + return ("other", None) + if len(raw) == 1 and raw.upper().isalpha(): + idx = ord(raw.upper()) - ord("A") + if 0 <= idx < len(choices): + return ("answer", _choice_value(choices[idx], raw)) + return ("answer", raw) + return ("answer", raw) + + +# ── approval policy ──────────────────────────────────────────────────── + + +def config_auto_approve(action_requests: list[dict]) -> bool: + """Whether config rules alone clear every action request. + + Returns True if no manual approval is needed via config: the global + ``auto_approve`` flag, non-execute tools, or a ``shell_allow_list`` + match on every shell command. Fail-closed on config load errors. + """ + if not action_requests: + return True + + try: + from ..config.settings import HITL_SHELL_TOOLS, load_config + + cfg = load_config() + except Exception: + return False # fail-closed + + if cfg.auto_approve: + return True + + shell_allow_list = ( + [s.strip() for s in cfg.shell_allow_list.split(",") if s.strip()] + if cfg.shell_allow_list + else [] + ) + + for req in action_requests: + name = req.get("name", "") + if name not in HITL_SHELL_TOOLS: + continue + args = req.get("args", {}) + command = args.get("command", "") if isinstance(args, dict) else "" + cmd = command.strip() + if not any(cmd.startswith(prefix) for prefix in shell_allow_list): + return False + return True + + +class ApprovalPolicy: + """Auto-approve policy backed by config rules and session grants. + + One instance is owned per process. The consumer keeps one on its event + loop; the CLI bridge keeps one on the bus loop. + """ + + def __init__(self) -> None: + self._granted_sessions: set[str] = set() + + def is_session_granted(self, session_key: str) -> bool: + """Whether the user previously chose "Approve all" for this session.""" + return session_key in self._granted_sessions + + def grant_session(self, session_key: str) -> None: + """Record an "Approve all" grant for this session.""" + self._granted_sessions.add(session_key) + + def clear_sessions(self) -> None: + """Forget all session grants (test hygiene / session reset).""" + self._granted_sessions.clear() + + def auto_decision( + self, session_key: str, action_requests: list[dict] + ) -> list[dict] | None: + """Return an approve-all ``decisions`` list if this can auto-resolve. + + Auto-resolves when the session was granted "Approve all" or when + config rules clear every request; otherwise returns ``None`` and + the caller must prompt the user. + """ + if self.is_session_granted(session_key) or config_auto_approve(action_requests): + return approve_decisions(action_requests) + return None + + +# ── transport adapter + reply registry ───────────────────────────────── + + +class InteractionIO(Protocol): + """One conversation partner on one channel chat. + + A transport adapter: the engine coroutines below drive a human + interaction entirely through this interface, so the same protocol + logic runs over the consumer's async loop and over the CLI bus loop. + + Attributes + ---------- + capabilities: + The channel's :class:`ChannelCapabilities` — the engine reads + ``inline_buttons`` to decide whether to attach approval buttons. + base_metadata: + The default outbound metadata for this chat (echoed back on each + send unless the engine supplies richer metadata, e.g. buttons). + """ + + capabilities: ChannelCapabilities + base_metadata: dict | None + + async def send(self, content: str, *, metadata: dict | None = None) -> bool: + """Send *content* to the user; return True on success.""" + ... + + async def wait_reply(self, *, timeout: float) -> str | None: + """Wait for the user's next reply; return None on timeout.""" + ... + + +@dataclass(frozen=True) +class PendingReply: + """A pending prompt reply plus optional transport-specific context.""" + + content: str + context: object | None = None + + +class PendingReplyRegistry: + """Route "the next message from this chat" into a waiting coroutine. + + One instance per process (the consumer owns one on its loop; the CLI + bridge owns one on the bus loop). ``register`` / ``wait`` are used by + an :class:`InteractionIO` adapter to block for a reply; the inbound + interception point calls ``try_resolve`` to hand a message to that + waiter instead of enqueuing it as a fresh turn. + + Asyncio-based: register/resolve must happen on the same event loop. + """ + + def __init__(self) -> None: + self._pending: dict[str, asyncio.Future[PendingReply]] = {} + + def register(self, session_key: str) -> asyncio.Future[PendingReply]: + """Create and store a future awaiting the next reply for *session_key*.""" + loop = asyncio.get_running_loop() + # A stale waiter for the same chat should never linger; cancel it + # so its coroutine unwinds instead of hanging until timeout. + stale = self._pending.get(session_key) + if stale is not None and not stale.done(): + stale.cancel() + fut: asyncio.Future[PendingReply] = loop.create_future() + self._pending[session_key] = fut + return fut + + def try_resolve( + self, + session_key: str, + content: str, + *, + context: object | None = None, + ) -> bool: + """Deliver *content* to a pending waiter. Returns True if consumed.""" + fut = self._pending.get(session_key) + if fut is not None and not fut.done(): + fut.set_result(PendingReply(content=content, context=context)) + return True + return False + + def discard(self, session_key: str) -> None: + """Drop any pending waiter for *session_key* (idempotent).""" + self._pending.pop(session_key, None) + + async def wait(self, session_key: str, timeout: float) -> str | None: + """Register, await a reply for *timeout* seconds, then clean up. + + Returns the reply text, or ``None`` on timeout / cancellation. + """ + reply = await self.wait_event(session_key, timeout) + return reply.content if reply is not None else None + + async def wait_event(self, session_key: str, timeout: float) -> PendingReply | None: + """Register, await a reply event, then clean up. + + Returns the full reply envelope, or ``None`` on timeout / registry + cancellation. Cancellation of the task awaiting this method propagates. + """ + fut = self.register(session_key) + try: + done, _pending = await asyncio.wait({fut}, timeout=timeout) + if not done: + fut.cancel() + return None + try: + return fut.result() + except asyncio.CancelledError: + return None + finally: + # Identity-safe: only drop *our* slot, never a newer waiter that + # re-registered on the same chat while we were unwinding. + if self._pending.get(session_key) is fut: + self._pending.pop(session_key, None) + + def clear(self) -> None: + """Cancel and forget every pending waiter (shutdown / test hygiene).""" + for fut in self._pending.values(): + if not fut.done(): + fut.cancel() + self._pending.clear() + + def __contains__(self, session_key: str) -> bool: + return session_key in self._pending + + +# ── engine coroutines ────────────────────────────────────────────────── + + +async def resolve_ask_user( + questions: list[dict], io: InteractionIO, *, timeout: float = ASK_USER_TIMEOUT +) -> dict: + """Drive an ask_user interrupt to a resume payload. + + Sends each question in turn, collects answers, and handles the choice + grammar (letters + the "Other" free-form sub-flow). ``/stop`` and + ``cancel`` are checked *before* parsing every reply. + + Returns a dict suitable for ``Command(resume=...)``: + ``{"answers": [...], "status": "answered"}`` or ``{"status": "cancelled"}``. + """ + if not questions: + return {"answers": [], "status": "answered"} + + total = len(questions) + answers: list[str] = [] + + for i, q in enumerate(questions): + if not await io.send(format_question_prompt(q, i, total)): + return {"status": "cancelled"} + + reply = await io.wait_reply(timeout=timeout) + if reply is None: + await io.send(ASK_USER_TIMEOUT_FEEDBACK) + return {"status": "cancelled"} + + raw = reply.strip() + required = q.get("required", True) is not False + if raw == "": + if required: + return {"status": "cancelled"} + answers.append("") + continue + + if is_stop_command(raw) or is_cancel_reply(raw): + return {"status": "cancelled"} + + if q.get("type", "text") == "multiple_choice": + choices = q.get("choices", []) + kind, value = parse_choice_answer(raw, choices) + if kind == "other": + if not await io.send(OTHER_PROMPT): + return {"status": "cancelled"} + other = await io.wait_reply(timeout=timeout) + if other is None: + await io.send(ASK_USER_TIMEOUT_FEEDBACK) + return {"status": "cancelled"} + other_raw = other.strip() + if other_raw == "": + if required: + return {"status": "cancelled"} + answers.append("") + continue + if is_stop_command(other_raw) or is_cancel_reply(other_raw): + return {"status": "cancelled"} + answers.append(other_raw) + else: + answers.append(value) + else: + answers.append(raw) + + return {"answers": answers, "status": "answered"} + + +@dataclass +class ApprovalOutcome: + """Result of :func:`resolve_approval`. + + ``decisions`` is the approve-all payload on approve/auto, or ``None`` + when the action was declined (reject / timeout / stop / unrecognized). + + ``unrecognized_reply`` carries the raw reply text when parsing failed. + The engine centralizes *parsing* but does not decide the transport + policy for unparseable text — the drivers do: the consumer rejects the + pending action and refeeds the text as a new agent turn (a channel user + who ignores the prompt and types a fresh instruction must not lose it); + the CLI bridge sends :data:`UNRECOGNIZED_FEEDBACK` and declines. + """ + + decisions: list[dict] | None = None + unrecognized_reply: str | None = None + + +async def resolve_approval( + action_requests: list, + io: InteractionIO, + policy: ApprovalPolicy, + session_key: str, + *, + timeout: float = HITL_APPROVAL_TIMEOUT, +) -> ApprovalOutcome: + """Drive a HITL approval interrupt to an :class:`ApprovalOutcome`. + + Auto-resolves via *policy* (session grant or config rule) without + prompting. Otherwise sends the approval prompt (with capability-driven + buttons), waits for a reply, and parses it. ``/stop`` cancels silently + (it already got its own ack from the transport's stop fast-path). An + unrecognized reply declines *without feedback* and hands the raw text + back to the driver via ``unrecognized_reply`` (see + :class:`ApprovalOutcome` for the per-driver policy). + """ + auto = policy.auto_decision(session_key, action_requests) + if auto is not None: + return ApprovalOutcome(decisions=auto) + + has_buttons = bool(io.capabilities.inline_buttons) + prompt = format_approval_prompt(action_requests, with_buttons=has_buttons) + metadata = approval_prompt_metadata(io.base_metadata, with_buttons=has_buttons) + if not await io.send(prompt, metadata=metadata): + return ApprovalOutcome() + + reply = await io.wait_reply(timeout=timeout) + if reply is None: + await io.send(APPROVAL_TIMEOUT_FEEDBACK) + return ApprovalOutcome() + + if is_stop_command(reply): + return ApprovalOutcome() + + decision = parse_approval_reply(reply) + if decision == "auto": + policy.grant_session(session_key) + await io.send(APPROVED_AUTO_FEEDBACK) + return ApprovalOutcome(decisions=approve_decisions(action_requests)) + if decision == "approve": + await io.send(APPROVED_FEEDBACK) + return ApprovalOutcome(decisions=approve_decisions(action_requests)) + if decision == "reject": + await io.send(REJECTED_FEEDBACK) + return ApprovalOutcome() + + # Unrecognized — decline and report the raw text; the driver chooses + # the feedback / refeed policy. + return ApprovalOutcome(unrecognized_reply=reply) diff --git a/EvoScientist/channels/qq/channel.py b/EvoScientist/channels/qq/channel.py index 36554eb..1f13eba 100644 --- a/EvoScientist/channels/qq/channel.py +++ b/EvoScientist/channels/qq/channel.py @@ -235,7 +235,7 @@ class QQChannel(Channel): Surfaces the click as an :class:`InboundMessage` whose ``content`` is the button's ``data`` verbatim — so a "1"/"approve"/… click flows - through ``_parse_approval_reply`` exactly like a typed reply. + through ``parse_approval_reply`` exactly like a typed reply. The click runs through inbound middleware (Dedup suppresses QQ retries) but is published directly to the bus so the per-sender @@ -264,7 +264,7 @@ class QQChannel(Channel): triggering_msg_id = getattr(resolved, "message_id", "") or "" # QQ may serialize non-str values; coerce. Fall back to button id - # when no data — same path as a typed reply via _parse_approval_reply. + # when no data — same path as a typed reply via parse_approval_reply. button_value = str(button_data) if button_data != "" else "" text = button_value or button_id @@ -360,7 +360,7 @@ class QQChannel(Channel): plain_text = self._plain_formatter.format(raw_text) # Plain-text fallback can't carry a keyboard. Append `value=label` # pairs so the user can still type "1"/"approve"/… instead of - # tapping (`_parse_approval_reply` accepts the same values). + # tapping (`parse_approval_reply` accepts the same values). if buttons: pairs = [] for btn in buttons: diff --git a/EvoScientist/cli/channel.py b/EvoScientist/cli/channel.py index 67ae4fc..a8dcc13 100644 --- a/EvoScientist/cli/channel.py +++ b/EvoScientist/cli/channel.py @@ -12,6 +12,7 @@ for the main thread to set a response via ``_set_channel_response()``. from __future__ import annotations import asyncio +import concurrent.futures import logging import queue import threading @@ -24,6 +25,18 @@ from typing import TYPE_CHECKING, Any from rich.panel import Panel from rich.text import Text +from ..channels.capabilities import ChannelCapabilities +from ..channels.interaction import ( + ASK_USER_TIMEOUT, + HITL_APPROVAL_TIMEOUT, + UNRECOGNIZED_FEEDBACK, + ApprovalPolicy, + InteractionIO, + PendingReplyRegistry, + is_stop_command, + resolve_approval, + resolve_ask_user, +) from ..commands.base import ChannelRuntime from ..stream.console import console @@ -448,21 +461,99 @@ async def _dispatch_channel_slash_impl( # --------------------------------------------------------------------------- -# HITL approval intercept: bus thread ⇄ main CLI thread +# HITL / ask_user interaction bridge: bus loop ⇄ main CLI thread # --------------------------------------------------------------------------- -# When the main thread needs HITL approval from a channel user, it registers -# a pending HITL wait for (channel, chat_id). The bus consumer checks this -# BEFORE normal enqueue, so the next reply from that user is intercepted. +# The interaction protocol itself (prompt formatting, reply grammar, +# feedback, auto-approve policy) lives in ``channels.interaction``. Here we +# only bridge it: the whole engine coroutine runs on the bus loop via +# ``run_coroutine_threadsafe`` while the calling (main / TUI) thread blocks +# on the resulting future. Replies are routed by a single asyncio-based +# ``PendingReplyRegistry`` fed from the inbound interception point — the bus +# consumer checks it BEFORE normal enqueue, so the next reply from that chat +# is intercepted. -_pending_hitl: dict[str, dict] = {} # "channel:chat_id" -> {event, reply} -_hitl_lock = threading.Lock() -_hitl_auto_approve: set[str] = set() # "channel:chat_id" keys with auto-approve -_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")) +# Extra head-room on the outer ``.result()`` wait so the engine's own +# per-flow timeout always fires first and returns a clean cancelled/None +# instead of the bridge tearing the coroutine down mid-flight. +_ENGINE_RESULT_SLACK = 30.0 +_ENGINE_CANCEL_SETTLE_TIMEOUT = 1.0 +# Send timeout inside the bridge IO adapter (kept per-flow-independent, as +# the standalone consumer has no send timeout). +_BRIDGE_SEND_TIMEOUT = 15.0 +_ASK_USER_WAITS_PER_QUESTION = 2 +_ASK_USER_SENDS_PER_QUESTION = 3 +_HITL_SENDS_PER_APPROVAL = 2 + +# One reply registry + one approval policy for the whole bridge process, +# both living on the bus loop (replacing the old ``_pending_hitl`` / +# ``_hitl_lock`` / ``_hitl_auto_approve`` module globals). +_reply_registry = PendingReplyRegistry() +_approval_policy = ApprovalPolicy() + + +class _BridgeIO(InteractionIO): + """:class:`InteractionIO` for the CLI bridge, running on the bus loop. + + ``send`` publishes outbound (bounded by :data:`_BRIDGE_SEND_TIMEOUT`); + ``wait_reply`` blocks on the shared :data:`_reply_registry`. Both run on + the bus loop because the engine coroutine is scheduled there via + ``run_coroutine_threadsafe`` — no per-message thread hop. + """ + + def __init__( + self, + bus: Any, + msg: ChannelMessage, + capabilities: ChannelCapabilities, + session_key: str, + ) -> None: + self._bus = bus + self._msg = msg + self.capabilities = capabilities + self.base_metadata = msg.metadata + self._session_key = session_key + + async def send(self, content: str, *, metadata: dict | None = None) -> bool: + from ..channels.bus.events import OutboundMessage + + try: + await asyncio.wait_for( + self._bus.publish_outbound( + OutboundMessage( + channel=self._msg.channel_type, + chat_id=self._msg.chat_id, + content=content, + metadata=metadata + if metadata is not None + else self._msg.metadata or {}, + ) + ), + timeout=_BRIDGE_SEND_TIMEOUT, + ) + return True + except Exception as exc: + _channel_logger.debug("bridge send failed: %s", exc) + return False + + async def wait_reply(self, *, timeout: float) -> str | None: + return await _reply_registry.wait(self._session_key, timeout) + + +def _ask_user_result_timeout(question_count: int) -> float: + per_question = ( + ASK_USER_TIMEOUT * _ASK_USER_WAITS_PER_QUESTION + + _BRIDGE_SEND_TIMEOUT * _ASK_USER_SENDS_PER_QUESTION + ) + return per_question * question_count + _ENGINE_RESULT_SLACK + + +def _hitl_result_timeout() -> float: + return ( + HITL_APPROVAL_TIMEOUT + + _BRIDGE_SEND_TIMEOUT * _HITL_SENDS_PER_APPROVAL + + _ENGINE_RESULT_SLACK + ) # --------------------------------------------------------------------------- @@ -596,265 +687,132 @@ def publish_to_channel_origin(thread_id: str | None, content: str) -> bool: return True -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 _run_engine_on_bus(coro, *, result_timeout: float, on_error): + """Run *coro* (an engine coroutine) on the bus loop and block for it. + Schedules the coroutine on ``_bus_loop`` via ``run_coroutine_threadsafe`` + and waits up to *result_timeout* seconds for it (the outer bound is the + engine's own per-flow timeout plus slack, so the engine's timeout fires + first). Returns *on_error* (a zero-arg factory) on any failure. + """ + bus_loop = _bus_loop + if bus_loop is None: + coro.close() + return on_error() + try: + fut = asyncio.run_coroutine_threadsafe(coro, bus_loop) + except Exception as exc: + coro.close() + _channel_logger.debug("interaction engine bridge failed: %s", exc) + return on_error() -def _register_hitl_wait(channel_type: str, chat_id: str) -> threading.Event: - """Register a pending HITL wait. Returns a threading.Event to block on.""" - key = f"{channel_type}:{chat_id}" - event = threading.Event() - with _hitl_lock: - _pending_hitl[key] = {"event": event, "reply": None} - return event - - -def _pop_hitl_reply(channel_type: str, chat_id: str) -> str | None: - """Pop and return the HITL reply (or None if not set).""" - key = f"{channel_type}:{chat_id}" - with _hitl_lock: - slot = _pending_hitl.pop(key, None) - return slot["reply"] if slot else None - - -def _try_set_hitl_reply(channel_type: str, chat_id: str, content: str) -> bool: - """Try to intercept a message as a HITL reply. Returns True if consumed.""" - key = f"{channel_type}:{chat_id}" - with _hitl_lock: - slot = _pending_hitl.get(key) - if slot: - slot["reply"] = content - slot["event"].set() - return True - return False + try: + return fut.result(timeout=result_timeout) + except concurrent.futures.TimeoutError as exc: + fut.cancel() + try: + asyncio.run_coroutine_threadsafe(asyncio.sleep(0), bus_loop).result( + timeout=_ENGINE_CANCEL_SETTLE_TIMEOUT + ) + except concurrent.futures.TimeoutError: + _channel_logger.debug("interaction engine cancellation did not settle") + except Exception as settle_exc: + _channel_logger.debug( + "interaction engine failed while settling cancellation: %s", + settle_exc, + ) + _channel_logger.debug("interaction engine bridge timed out: %s", exc) + return on_error() + except Exception as exc: + _channel_logger.debug("interaction engine bridge failed: %s", exc) + return on_error() def channel_ask_user_prompt( ask_user_data: dict, msg: ChannelMessage | None = None, ) -> dict: - """Format ask_user questions and collect answers from a channel user. + """Collect answers to ask_user questions from a channel user. - If *msg* is provided, sends questions via the bus and waits for a reply. - Otherwise falls back to returning a cancelled result. + Thin bridge: runs :func:`channels.interaction.resolve_ask_user` on the + bus loop over a :class:`_BridgeIO` and blocks for the result. Signature + and return shape are unchanged (callers in ``interactive.py`` / + ``commands.py`` / ``tui_interactive.py`` are untouched). - Returns: - ``{"answers": [...], "status": "answered"}`` or - ``{"status": "cancelled"}``. + Returns ``{"answers": [...], "status": "answered"}`` or + ``{"status": "cancelled"}``. """ - from ..channels.bus.events import OutboundMessage - questions = ask_user_data.get("questions", []) if not questions: return {"answers": [], "status": "answered"} - - if msg is None or not msg.bus_ref: + if msg is None or not msg.bus_ref or _bus_loop is None: return {"status": "cancelled"} - bus_loop = _bus_loop - if not bus_loop: - return {"status": "cancelled"} - - def _send(content: str) -> bool: - try: - asyncio.run_coroutine_threadsafe( - msg.bus_ref.publish_outbound( - OutboundMessage( - channel=msg.channel_type, - chat_id=msg.chat_id, - content=content, - metadata=msg.metadata or {}, - ) - ), - bus_loop, - ).result(timeout=15) - return True - except Exception as exc: - _channel_logger.debug("ask_user send failed: %s", exc) - return False - - # Ask one question at a time (consistent with Rich CLI / TUI) - total = len(questions) - answers: list[str] = [] - - for i, q in enumerate(questions): - q_text = q.get("question", "") - q_type = q.get("type", "text") - required = q.get("required", True) - - # Format single question - if total == 1: - header = "\u2753 Quick check-in from EvoScientist\n" - else: - header = f"\u2753 Question {i + 1}/{total}\n" - - lines = [header, f"{i + 1}. {q_text}"] - if not required: - lines[-1] += " (optional)" - - if q_type == "multiple_choice": - choices = q.get("choices", []) - for j, choice in enumerate(choices): - label = choice.get("value", str(choice)) - letter = chr(ord("A") + j) - lines.append(f" {letter}. {label}") - other_letter = chr(ord("A") + len(choices)) - lines.append(f" {other_letter}. Other") - lines.append( - f"\nReply with a letter ({'/'.join(chr(ord('A') + k) for k in range(len(choices) + 1))}), or 'cancel'." - ) - else: - skip_hint = " Leave empty to skip." if not required else "" - lines.append(f"\nReply with your answer, or 'cancel'.{skip_hint}") - - if not _send("\n".join(lines)): - return {"status": "cancelled"} - - # Wait for reply - hitl_event = _register_hitl_wait(msg.channel_type, msg.chat_id) - replied = hitl_event.wait(timeout=_ASK_USER_TIMEOUT) - reply_text = _pop_hitl_reply(msg.channel_type, msg.chat_id) - - if not replied or not reply_text: - _send("\u23f0 Response timed out.") - return {"status": "cancelled"} - - raw = reply_text.strip() - if _is_stop_command(raw): - return {"status": "cancelled"} - if raw.lower() == "cancel": - return {"status": "cancelled"} - - # Parse answer - if q_type == "multiple_choice": - choices = q.get("choices", []) - other_letter = chr(ord("A") + len(choices)) - if len(raw) == 1 and raw.upper() == other_letter: - # Other selected — ask for free-form input - if not _send("Please type your answer:"): - return {"status": "cancelled"} - hitl_event = _register_hitl_wait(msg.channel_type, msg.chat_id) - replied = hitl_event.wait(timeout=_ASK_USER_TIMEOUT) - other_text = _pop_hitl_reply(msg.channel_type, msg.chat_id) - 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()) - elif len(raw) == 1 and raw.upper().isalpha(): - idx = ord(raw.upper()) - ord("A") - if 0 <= idx < len(choices): - answers.append(choices[idx].get("value", raw)) - else: - answers.append(raw) - else: - answers.append(raw) - else: - answers.append(raw) - - return {"answers": answers, "status": "answered"} + # ask_user never uses buttons; a plain capability set suffices. + io = _BridgeIO( + msg.bus_ref, msg, ChannelCapabilities(), _channel_message_session_key(msg) + ) + return _run_engine_on_bus( + resolve_ask_user(questions, io, timeout=ASK_USER_TIMEOUT), + result_timeout=_ask_user_result_timeout(len(questions)), + on_error=lambda: {"status": "cancelled"}, + ) def channel_hitl_prompt( action_requests: list, msg: ChannelMessage, ) -> list[dict] | None: - """Send HITL approval prompt to channel user and wait for reply. + """Resolve a HITL approval prompt with a channel user. - Blocking function — uses threading.Event.wait(). Safe to call from a - background thread (CLI channel processing or asyncio.to_thread in TUI). + Thin bridge: runs :func:`channels.interaction.resolve_approval` on the + bus loop over a :class:`_BridgeIO` and blocks for the result. Signature + and return shape are unchanged (callers are untouched). Safe to call + from a background thread (CLI channel processing / TUI ``to_thread``). - Returns approval decisions list on approve/auto, or None on reject/timeout. + Returns the approval decisions list on approve/auto, or None on + reject / unrecognized / timeout / stop. """ - from ..channels.bus.events import OutboundMessage - from ..channels.consumer import ( - _approval_prompt_metadata, - _format_approval_prompt, - _parse_approval_reply, - ) + session_key = _channel_message_session_key(msg) + decisions = _approval_policy.auto_decision(session_key, action_requests) + if decisions is not None: + return decisions - # Check session auto-approve (set by a previous "3" reply) - session_key = f"{msg.channel_type}:{msg.chat_id}" - if session_key in _hitl_auto_approve: - return [{"type": "approve"} for _ in action_requests] - - bus_loop = _bus_loop - if not (bus_loop and msg.bus_ref): + if not (_bus_loop and msg.bus_ref): _channel_logger.debug("HITL: no bus_loop or bus_ref, rejecting") return None - # Look up the channel instance so we can attach buttons when the channel - # supports `inline_buttons` (Feishu cards, QQ keyboards, …). + # Look up the channel instance so the engine can attach buttons when the + # channel supports `inline_buttons` (Feishu cards, QQ keyboards, …). channel_obj = ( _manager.get_channel(msg.channel_type) if _manager is not None else None ) - has_buttons = channel_obj is not None and channel_obj.capabilities.inline_buttons - approval_metadata = _approval_prompt_metadata( - msg.metadata, with_buttons=has_buttons + capabilities = ( + channel_obj.capabilities if channel_obj is not None else ChannelCapabilities() ) + io = _BridgeIO(msg.bus_ref, msg, capabilities, session_key) - def _send(content: str, *, metadata: dict | None = None) -> bool: - """Send a message to the channel user. Returns True on success.""" - try: - asyncio.run_coroutine_threadsafe( - msg.bus_ref.publish_outbound( - OutboundMessage( - channel=msg.channel_type, - chat_id=msg.chat_id, - content=content, - metadata=metadata - if metadata is not None - else msg.metadata or {}, - ) - ), - bus_loop, - ).result(timeout=15) - return True - except Exception as exc: - _channel_logger.debug("HITL send failed: %s", exc) - return False + async def _hitl_flow() -> list[dict] | None: + outcome = await resolve_approval( + action_requests, + io, + _approval_policy, + session_key, + timeout=HITL_APPROVAL_TIMEOUT, + ) + if outcome.unrecognized_reply is not None: + # CLI-bridge policy: an unparseable reply declines with the + # explicit notice. Only the serve-mode consumer refeeds the + # text as a new turn. + await io.send(UNRECOGNIZED_FEEDBACK) + return None + return outcome.decisions - # 1. Send approval prompt - prompt_text = _format_approval_prompt(action_requests, with_buttons=has_buttons) - if not _send(prompt_text, metadata=approval_metadata): - return None - - # 2. Wait for channel user's reply - hitl_event = _register_hitl_wait(msg.channel_type, msg.chat_id) - replied = hitl_event.wait(timeout=_HITL_APPROVAL_TIMEOUT) - reply_text = _pop_hitl_reply(msg.channel_type, msg.chat_id) - - if not replied or not reply_text: - _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": - _hitl_auto_approve.add(session_key) - _send("\u2705 已批准(后续自动通过)") - return [{"type": "approve"} for _ in action_requests] - if decision == "approve": - _send("\u2705 已批准") - return [{"type": "approve"} for _ in action_requests] - - feedback = ( - "\u274c 已拒绝" - if decision == "reject" - else "Unrecognized reply. Action rejected." + return _run_engine_on_bus( + _hitl_flow(), + result_timeout=_hitl_result_timeout(), + on_error=lambda: None, ) - _send(feedback) - return None # --------------------------------------------------------------------------- @@ -1040,13 +998,16 @@ 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): + session_key = _channel_session_key(msg.channel, msg.chat_id) + + # /stop should preempt interaction interception so cancel works + # while waiting for approvals/questions. If a prompt wait is + # pending, still deliver /stop into it so the blocking engine + # unwinds immediately (it treats /stop as a clean cancel). + if is_stop_command(msg.content): + if _reply_registry.try_resolve(session_key, msg.content): _channel_logger.info( - f"[bus] stop request released HITL wait for " + f"[bus] stop request released interaction wait for " f"{msg.channel}:{msg.chat_id}" ) _task = asyncio.create_task(_handle_bus_message(bus, manager, msg)) @@ -1054,10 +1015,12 @@ async def _bus_inbound_consumer(bus, manager) -> None: _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): + # Reply interception sits ahead of normal enqueue — if a prompt + # is waiting on this chat, the next message is its reply and + # must NOT be enqueued as a fresh agent turn. + if _reply_registry.try_resolve(session_key, msg.content): _channel_logger.info( - f"[bus] HITL reply from {msg.channel}:{msg.sender_id}: " + f"[bus] interaction reply from {msg.channel}:{msg.sender_id}: " f"{msg.content[:60]}" ) continue @@ -1085,7 +1048,7 @@ async def _handle_bus_message(bus, manager, msg) -> None: # 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): + if is_stop_command(msg.content): cancelled_count, active_count = _cancel_channel_session( msg.channel, msg.chat_id ) diff --git a/tests/test_bus_integration.py b/tests/test_bus_integration.py index 5f15c9a..f121637 100644 --- a/tests/test_bus_integration.py +++ b/tests/test_bus_integration.py @@ -39,9 +39,8 @@ def clean_channel_state(): 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() + channel_mod._reply_registry.clear() + channel_mod._approval_policy.clear_sessions() with display_mod._stream_cancel_lock: display_mod._stream_cancel_event.clear() display_mod._stream_cancel_events.clear() @@ -365,7 +364,13 @@ class TestBusInboundConsumer: assert queued.msg_id not in channel_mod._pending_responses async def test_stop_during_hitl_wait_releases_wait_and_acks(self): - """`/stop` should wake pending HITL wait and publish immediate ack.""" + """`/stop` should wake a pending interaction wait and publish an ack. + + The bus consumer delivers ``/stop`` into the reply registry (so the + blocking engine unwinds) AND acks with "Stopped." — the registry + interception sits ahead of normal enqueue, so the message never + becomes a fresh agent turn. + """ from EvoScientist.cli import channel as channel_mod from EvoScientist.cli.channel import _bus_inbound_consumer, _message_queue @@ -374,7 +379,12 @@ class TestBusInboundConsumer: ch = FakeChannel() manager.register(ch) - hitl_event = channel_mod._register_hitl_wait("fake", "chat1") + # Simulate a HITL/ask_user prompt waiting for this chat's reply. + reply_fut = asyncio.ensure_future( + channel_mod._reply_registry.wait("fake:chat1", timeout=5.0) + ) + await asyncio.sleep(0.01) # let the wait register + consumer = asyncio.create_task(_bus_inbound_consumer(bus, manager)) await bus.publish_inbound( @@ -387,12 +397,9 @@ class TestBusInboundConsumer: ) ) - 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" + # The pending wait receives "/stop" (engine will treat it as cancel). + released = await asyncio.wait_for(reply_fut, timeout=2.0) + assert released == "/stop" outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=2.0) assert outbound.content == "Stopped." diff --git a/tests/test_cli_channel_bridge.py b/tests/test_cli_channel_bridge.py new file mode 100644 index 0000000..ea57d8d --- /dev/null +++ b/tests/test_cli_channel_bridge.py @@ -0,0 +1,426 @@ +"""Tests for the CLI channel bridge. + +The bridge runs the shared interaction engine on the bus loop via +``run_coroutine_threadsafe`` while the calling thread blocks. Covered here: + +* **Ordering** — the reply-interception point sits *ahead* of normal + enqueue, so a reply to a pending prompt is delivered into the engine's + wait and never becomes a fresh agent turn. +* **Bridge round-trip** — ``channel_hitl_prompt`` drives ``resolve_approval`` + end-to-end over a real bus loop and returns today's decision payloads. +""" + +import asyncio +import threading + +import pytest + +from EvoScientist.channels import interaction as interaction_mod +from EvoScientist.channels.bus.events import InboundMessage +from EvoScientist.channels.bus.message_bus import MessageBus +from EvoScientist.channels.channel_manager import ChannelManager +from EvoScientist.cli import channel as channel_mod +from EvoScientist.cli.channel import ChannelMessage +from tests.fakes import QueueFakeChannel + + +def _reset_channel_state(): + channel_mod._reply_registry.clear() + channel_mod._approval_policy.clear_sessions() + while not channel_mod._message_queue.empty(): + channel_mod._message_queue.get_nowait() + 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() + + +@pytest.fixture(autouse=True) +def _clean_bridge_state(): + _reset_channel_state() + yield + _reset_channel_state() + + +# ═══════════════════════════════════════════════════════════════════════ +# Reply interception sits ahead of normal enqueue +# ═══════════════════════════════════════════════════════════════════════ + + +class TestReplyInterceptionOrdering: + async def test_pending_reply_intercepted_not_enqueued(self): + """A reply to a pending prompt resolves the wait and is NOT enqueued.""" + from EvoScientist.cli.channel import _bus_inbound_consumer, _message_queue + + bus = MessageBus() + manager = ChannelManager(bus) + manager.register(QueueFakeChannel()) + + # A prompt is waiting on fake:chat1 (as the engine's wait_reply would). + reply_fut = asyncio.ensure_future( + channel_mod._reply_registry.wait("fake:chat1", timeout=5.0) + ) + await asyncio.sleep(0.01) + + consumer = asyncio.create_task(_bus_inbound_consumer(bus, manager)) + await bus.publish_inbound( + InboundMessage( + channel="fake", + sender_id="user1", + chat_id="chat1", + content="1", + message_id="m-reply", + ) + ) + + # The reply is delivered into the pending wait... + got = await asyncio.wait_for(reply_fut, timeout=2.0) + assert got == "1" + # ...and did NOT become a queued agent turn. + await asyncio.sleep(0.05) + assert _message_queue.empty() + + consumer.cancel() + try: + await consumer + except asyncio.CancelledError: + pass + + async def test_message_without_pending_wait_is_enqueued(self): + """With no pending prompt, a normal message flows to enqueue as before.""" + from EvoScientist.cli.channel import _bus_inbound_consumer, _message_queue + + bus = MessageBus() + manager = ChannelManager(bus) + manager.register(QueueFakeChannel()) + + consumer = asyncio.create_task(_bus_inbound_consumer(bus, manager)) + await bus.publish_inbound( + InboundMessage( + channel="fake", + sender_id="user1", + chat_id="chat1", + content="hello there", + message_id="m-normal", + ) + ) + + queued = None + for _ in range(40): + 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.02) + assert queued is not None + assert queued.content == "hello there" + + consumer.cancel() + try: + await consumer + except asyncio.CancelledError: + pass + # Drain the pending waiter created by _handle_bus_message. + channel_mod._pop_channel_response(queued.msg_id, cancel_pending=True) + + +# ═══════════════════════════════════════════════════════════════════════ +# Bridge timeout budgeting and cancellation +# ═══════════════════════════════════════════════════════════════════════ + + +class TestBridgeTimeouts: + def test_ask_user_outer_timeout_includes_waits_and_sends(self): + question_count = 2 + per_question_worst_case = ( + channel_mod.ASK_USER_TIMEOUT * channel_mod._ASK_USER_WAITS_PER_QUESTION + + channel_mod._BRIDGE_SEND_TIMEOUT + * channel_mod._ASK_USER_SENDS_PER_QUESTION + ) + + assert channel_mod._ask_user_result_timeout(question_count) == ( + per_question_worst_case * question_count + channel_mod._ENGINE_RESULT_SLACK + ) + + def test_hitl_outer_timeout_exceeds_engine_worst_case(self): + engine_worst_case = ( + channel_mod.HITL_APPROVAL_TIMEOUT + + channel_mod._BRIDGE_SEND_TIMEOUT * channel_mod._HITL_SENDS_PER_APPROVAL + ) + + assert channel_mod._hitl_result_timeout() == ( + engine_worst_case + channel_mod._ENGINE_RESULT_SLACK + ) + assert channel_mod._hitl_result_timeout() > engine_worst_case + + def test_outer_timeout_cancels_engine_and_releases_reply_slot(self, monkeypatch): + session_key = "fake:timeout" + registered = threading.Event() + cancelled = threading.Event() + + async def _wait_forever(): + channel_mod._reply_registry.register(session_key) + registered.set() + try: + await asyncio.Future() + except asyncio.CancelledError: + cancelled.set() + raise + finally: + channel_mod._reply_registry.discard(session_key) + + with _BusLoopThread() as loop: + monkeypatch.setattr(channel_mod, "_bus_loop", loop) + result = channel_mod._run_engine_on_bus( + _wait_forever(), + result_timeout=0.2, + on_error=lambda: "cancelled", + ) + + assert result == "cancelled" + assert registered.wait(timeout=1.0) + assert cancelled.wait(timeout=1.0) + assert channel_mod._reply_registry.try_resolve(session_key, "late") is False + + +# ═══════════════════════════════════════════════════════════════════════ +# Bridge round-trip: channel_hitl_prompt over a real bus loop +# ═══════════════════════════════════════════════════════════════════════ + + +class _BusLoopThread: + """A dedicated event loop running in a background thread (like the bus).""" + + def __init__(self): + self.loop = asyncio.new_event_loop() + self._thread = threading.Thread(target=self.loop.run_forever, daemon=True) + + def __enter__(self): + self._thread.start() + return self.loop + + def __exit__(self, *exc): + self.loop.call_soon_threadsafe(self.loop.stop) + self._thread.join(timeout=2) + self.loop.close() + + +def _feed_reply_when_ready(loop, session_key, reply, *, tries=200): + """Schedule a coroutine on *loop* that resolves the pending wait.""" + + async def _feeder(): + for _ in range(tries): + if session_key in channel_mod._reply_registry: + channel_mod._reply_registry.try_resolve(session_key, reply) + return + await asyncio.sleep(0.01) + + asyncio.run_coroutine_threadsafe(_feeder(), loop) + + +class TestHitlPromptBridge: + def test_no_bus_loop_rejects(self, monkeypatch): + monkeypatch.setattr(interaction_mod, "config_auto_approve", lambda reqs: False) + monkeypatch.setattr(channel_mod, "_bus_loop", None) + msg = ChannelMessage( + msg_id="m1", + content="", + sender="u1", + channel_type="fake", + chat_id="chat1", + bus_ref=object(), + ) + assert channel_mod.channel_hitl_prompt([{"name": "execute"}], msg) is None + + def test_session_grant_approves_when_bus_loop_down(self, monkeypatch): + monkeypatch.setattr(interaction_mod, "config_auto_approve", lambda reqs: False) + monkeypatch.setattr(channel_mod, "_bus_loop", None) + msg = ChannelMessage( + msg_id="m1", + content="", + sender="u1", + channel_type="fake", + chat_id="chat1", + bus_ref=object(), + ) + + channel_mod._approval_policy.grant_session("fake:chat1") + + assert channel_mod.channel_hitl_prompt([{"name": "execute"}], msg) == [ + {"type": "approve"} + ] + + def test_approve_round_trip(self, monkeypatch): + # Force the manual-prompt path (no config auto-approve). + monkeypatch.setattr(interaction_mod, "config_auto_approve", lambda reqs: False) + with _BusLoopThread() as loop: + monkeypatch.setattr(channel_mod, "_bus_loop", loop) + monkeypatch.setattr(channel_mod, "_manager", None) # default caps + bus = MessageBus() + msg = ChannelMessage( + msg_id="m1", + content="", + sender="u1", + channel_type="fake", + chat_id="chat1", + bus_ref=bus, + metadata={}, + ) + _feed_reply_when_ready(loop, "fake:chat1", "1") + result = channel_mod.channel_hitl_prompt( + [{"name": "execute", "args": {"command": "ls"}}], msg + ) + assert result == [{"type": "approve"}] + + def test_reject_round_trip(self, monkeypatch): + monkeypatch.setattr(interaction_mod, "config_auto_approve", lambda reqs: False) + with _BusLoopThread() as loop: + monkeypatch.setattr(channel_mod, "_bus_loop", loop) + monkeypatch.setattr(channel_mod, "_manager", None) + bus = MessageBus() + msg = ChannelMessage( + msg_id="m1", + content="", + sender="u1", + channel_type="fake", + chat_id="chat1", + bus_ref=bus, + metadata={}, + ) + _feed_reply_when_ready(loop, "fake:chat1", "2") + result = channel_mod.channel_hitl_prompt( + [{"name": "execute", "args": {"command": "ls"}}], msg + ) + assert result is None + + def test_unrecognized_reply_declines_without_refeed(self, monkeypatch): + """CLI-bridge policy: unparseable reply → explicit notice, NO refeed. + + The reply is consumed by the interception registry, never enqueued + as a new turn, and the user gets the unrecognized-reply notice. Only + the serve-mode consumer refeeds; see + TestConsumerUnrecognizedRefeed in tests/test_interaction_engine.py. + """ + monkeypatch.setattr(interaction_mod, "config_auto_approve", lambda reqs: False) + with _BusLoopThread() as loop: + monkeypatch.setattr(channel_mod, "_bus_loop", loop) + monkeypatch.setattr(channel_mod, "_manager", None) + bus = MessageBus() + msg = ChannelMessage( + msg_id="m1", + content="", + sender="u1", + channel_type="fake", + chat_id="chat1", + bus_ref=bus, + metadata={}, + ) + _feed_reply_when_ready(loop, "fake:chat1", "do something else instead") + result = channel_mod.channel_hitl_prompt( + [{"name": "execute", "args": {"command": "ls"}}], msg + ) + assert result is None + + # Outbound: prompt, then the exact old unrecognized notice. + async def _drain(): + out = [] + while True: + try: + m = await asyncio.wait_for(bus.consume_outbound(), timeout=0.2) + except TimeoutError: + return out + out.append(m.content) + + contents = asyncio.run_coroutine_threadsafe(_drain(), loop).result( + timeout=5 + ) + assert contents[-1] == interaction_mod.UNRECOGNIZED_FEEDBACK + # No refeed: nothing was enqueued for the main thread. + assert channel_mod._message_queue.empty() + + def test_approve_all_grants_channel_session(self, monkeypatch): + monkeypatch.setattr(interaction_mod, "config_auto_approve", lambda reqs: False) + with _BusLoopThread() as loop: + monkeypatch.setattr(channel_mod, "_bus_loop", loop) + monkeypatch.setattr(channel_mod, "_manager", None) + bus = MessageBus() + msg = ChannelMessage( + msg_id="m1", + content="", + sender="u1", + channel_type="fake", + chat_id="chat1", + bus_ref=bus, + metadata={}, + ) + _feed_reply_when_ready(loop, "fake:chat1", "3") + result = channel_mod.channel_hitl_prompt( + [{"name": "execute", "args": {"command": "ls"}}], msg + ) + assert result == [{"type": "approve"}] + # "Approve all" grant persists: a second prompt auto-approves with + # no reply fed at all. + result2 = channel_mod.channel_hitl_prompt( + [{"name": "execute", "args": {"command": "rm"}}], msg + ) + assert result2 == [{"type": "approve"}] + + +class TestAskUserPromptBridge: + def test_no_msg_cancels(self): + assert channel_mod.channel_ask_user_prompt( + {"questions": [{"question": "Q?", "type": "text"}]}, None + ) == {"status": "cancelled"} + + def test_empty_questions_answered(self): + assert channel_mod.channel_ask_user_prompt({"questions": []}, None) == { + "answers": [], + "status": "answered", + } + + def test_text_answer_round_trip(self, monkeypatch): + with _BusLoopThread() as loop: + monkeypatch.setattr(channel_mod, "_bus_loop", loop) + bus = MessageBus() + msg = ChannelMessage( + msg_id="m1", + content="", + sender="u1", + channel_type="fake", + chat_id="chat1", + bus_ref=bus, + metadata={}, + ) + _feed_reply_when_ready(loop, "fake:chat1", "CIFAR-10") + result = channel_mod.channel_ask_user_prompt( + {"questions": [{"question": "Which dataset?", "type": "text"}]}, msg + ) + assert result == {"answers": ["CIFAR-10"], "status": "answered"} + + +class TestBridgeClosesUnscheduledCoroutine: + def test_no_bus_loop_closes_coro(self, monkeypatch): + """The bridge must close an engine coroutine it never scheduled — + otherwise GC emits a "was never awaited" RuntimeWarning. (cr_frame + is not a reliable observable for close() on unstarted coroutines.)""" + import gc + import warnings + + monkeypatch.setattr(channel_mod, "_bus_loop", None) + + async def _engine(): + return "never" + + coro = _engine() + result = channel_mod._run_engine_on_bus( + coro, result_timeout=1.0, on_error=lambda: "fallback" + ) + assert result == "fallback" + + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + del coro + gc.collect() + assert not [w for w in caught if issubclass(w.category, RuntimeWarning)] diff --git a/tests/test_hitl.py b/tests/test_hitl.py index 61d1f80..400ca71 100644 --- a/tests/test_hitl.py +++ b/tests/test_hitl.py @@ -1,5 +1,6 @@ """Tests for HITL (Human-in-the-Loop) approval mechanism.""" +import asyncio from unittest.mock import MagicMock, patch from langgraph.types import Interrupt @@ -437,34 +438,34 @@ class TestInterruptEventParsing: class TestConsumerHitlHelpers: def test_parse_approval_approve(self): - from EvoScientist.channels.consumer import _parse_approval_reply + from EvoScientist.channels.interaction import parse_approval_reply for text in ("1", "y", "yes", "approve", "ok", " 1 ", " Y "): - assert _parse_approval_reply(text) == "approve", f"Failed for: {text!r}" + assert parse_approval_reply(text) == "approve", f"Failed for: {text!r}" def test_parse_approval_reject(self): - from EvoScientist.channels.consumer import _parse_approval_reply + from EvoScientist.channels.interaction import parse_approval_reply for text in ("2", "n", "no", "reject"): - assert _parse_approval_reply(text) == "reject", f"Failed for: {text!r}" + assert parse_approval_reply(text) == "reject", f"Failed for: {text!r}" def test_parse_approval_auto(self): - from EvoScientist.channels.consumer import _parse_approval_reply + from EvoScientist.channels.interaction import parse_approval_reply for text in ("3", "a", "auto", "approve all"): - assert _parse_approval_reply(text) == "auto", f"Failed for: {text!r}" + assert parse_approval_reply(text) == "auto", f"Failed for: {text!r}" def test_parse_approval_unrecognized(self): - from EvoScientist.channels.consumer import _parse_approval_reply + from EvoScientist.channels.interaction import parse_approval_reply - assert _parse_approval_reply("hello world") is None - assert _parse_approval_reply("") is None - assert _parse_approval_reply("maybe") is None + assert parse_approval_reply("hello world") is None + assert parse_approval_reply("") is None + assert parse_approval_reply("maybe") is None def test_format_approval_prompt(self): - from EvoScientist.channels.consumer import _format_approval_prompt + from EvoScientist.channels.interaction import format_approval_prompt - prompt = _format_approval_prompt( + prompt = format_approval_prompt( [ {"name": "execute", "args": {"command": "ls -la"}}, ] @@ -476,9 +477,9 @@ class TestConsumerHitlHelpers: assert "2=Reject" in prompt def test_format_approval_prompt_multiple(self): - from EvoScientist.channels.consumer import _format_approval_prompt + from EvoScientist.channels.interaction import format_approval_prompt - prompt = _format_approval_prompt( + prompt = format_approval_prompt( [ {"name": "execute", "args": {"command": "ls"}}, {"name": "write_file", "args": {"path": "/out.txt"}}, @@ -488,17 +489,17 @@ class TestConsumerHitlHelpers: assert "2. write_file: /out.txt" in prompt def test_should_auto_approve_non_execute(self): - from EvoScientist.channels.consumer import _should_auto_approve + from EvoScientist.channels.interaction import config_auto_approve - assert _should_auto_approve([{"name": "write_file", "args": {}}]) is True + assert config_auto_approve([{"name": "write_file", "args": {}}]) is True def test_should_auto_approve_empty(self): - from EvoScientist.channels.consumer import _should_auto_approve + from EvoScientist.channels.interaction import config_auto_approve - assert _should_auto_approve([]) is True + assert config_auto_approve([]) is True def test_should_auto_approve_execute_no_allowlist(self): - from EvoScientist.channels.consumer import _should_auto_approve + from EvoScientist.channels.interaction import config_auto_approve # With default config (auto_approve=False, shell_allow_list=""), # execute should NOT auto-approve @@ -506,7 +507,7 @@ class TestConsumerHitlHelpers: mock_cfg.auto_approve = False mock_cfg.shell_allow_list = "" with patch("EvoScientist.config.settings.load_config", return_value=mock_cfg): - result = _should_auto_approve( + result = config_auto_approve( [ {"name": "execute", "args": {"command": "rm -rf /"}}, ] @@ -515,13 +516,13 @@ class TestConsumerHitlHelpers: def test_should_auto_approve_run_in_background_no_allowlist(self): """Channel path must NOT auto-approve run_in_background (same as execute).""" - from EvoScientist.channels.consumer import _should_auto_approve + from EvoScientist.channels.interaction import config_auto_approve mock_cfg = MagicMock() mock_cfg.auto_approve = False mock_cfg.shell_allow_list = "" with patch("EvoScientist.config.settings.load_config", return_value=mock_cfg): - result = _should_auto_approve( + result = config_auto_approve( [ {"name": "run_in_background", "args": {"command": "rm -rf /"}}, ] @@ -529,12 +530,12 @@ class TestConsumerHitlHelpers: assert result is False def test_should_auto_approve_config_true(self): - from EvoScientist.channels.consumer import _should_auto_approve + from EvoScientist.channels.interaction import config_auto_approve mock_cfg = MagicMock() mock_cfg.auto_approve = True with patch("EvoScientist.config.settings.load_config", return_value=mock_cfg): - result = _should_auto_approve( + result = config_auto_approve( [ {"name": "execute", "args": {"command": "rm -rf /"}}, ] @@ -542,13 +543,13 @@ class TestConsumerHitlHelpers: assert result is True def test_should_auto_approve_allowlist_match(self): - from EvoScientist.channels.consumer import _should_auto_approve + from EvoScientist.channels.interaction import config_auto_approve mock_cfg = MagicMock() mock_cfg.auto_approve = False mock_cfg.shell_allow_list = "ls,python" with patch("EvoScientist.config.settings.load_config", return_value=mock_cfg): - result = _should_auto_approve( + result = config_auto_approve( [ {"name": "execute", "args": {"command": "ls -la"}}, ] @@ -557,53 +558,48 @@ class TestConsumerHitlHelpers: # ============================================================================= -# Channel HITL intercept mechanism (channel.py) +# Channel reply-interception mechanism (channel.py PendingReplyRegistry) # ============================================================================= +# The CLI bridge routes prompt replies through the shared asyncio-based +# ``PendingReplyRegistry`` on the bus loop (replacing the old threading.Event +# ``_pending_hitl`` globals). ``_bus_inbound_consumer`` feeds it via +# ``try_resolve`` ahead of normal enqueue. -class TestChannelHitlIntercept: - def test_register_and_set_hitl_reply(self): - from EvoScientist.cli.channel import ( - _pop_hitl_reply, - _register_hitl_wait, - _try_set_hitl_reply, +class TestChannelReplyRegistry: + async def test_register_and_resolve_reply(self): + from EvoScientist.cli import channel as channel_mod + + reg = channel_mod._reply_registry + reg.clear() + + async def _resolver(): + await asyncio.sleep(0.01) # let wait() register first + assert reg.try_resolve("telegram:chat123", "1") is True + + got, _ = await asyncio.gather( + reg.wait("telegram:chat123", timeout=1.0), _resolver() ) + assert got == "1" + assert "telegram:chat123" not in reg - event = _register_hitl_wait("telegram", "chat123") - assert not event.is_set() + def test_try_resolve_no_pending(self): + from EvoScientist.cli import channel as channel_mod - # Simulate reply arriving - intercepted = _try_set_hitl_reply("telegram", "chat123", "1") - assert intercepted is True - assert event.is_set() + channel_mod._reply_registry.clear() + # No pending wait — should not intercept. + resolved = channel_mod._reply_registry.try_resolve("discord:no_pending", "y") + assert resolved is False - reply = _pop_hitl_reply("telegram", "chat123") - assert reply == "1" + async def test_reply_timeout_returns_none(self): + from EvoScientist.cli import channel as channel_mod - def test_try_set_hitl_reply_no_pending(self): - from EvoScientist.cli.channel import _try_set_hitl_reply - - # No pending HITL — should not intercept - assert _try_set_hitl_reply("discord", "no_pending", "y") is False - - def test_pop_hitl_reply_no_pending(self): - from EvoScientist.cli.channel import _pop_hitl_reply - - assert _pop_hitl_reply("discord", "no_pending") is None - - def test_hitl_reply_timeout(self): - from EvoScientist.cli.channel import ( - _pop_hitl_reply, - _register_hitl_wait, - ) - - event = _register_hitl_wait("telegram", "timeout_chat") - # Don't set reply — simulate timeout - replied = event.wait(timeout=0.01) - assert replied is False - # Pop should still return None (reply was never set) - reply = _pop_hitl_reply("telegram", "timeout_chat") - assert reply is None + reg = channel_mod._reply_registry + reg.clear() + # No reply delivered — wait should time out and clean up. + got = await reg.wait("telegram:timeout_chat", timeout=0.02) + assert got is None + assert "telegram:timeout_chat" not in reg # ============================================================================= diff --git a/tests/test_interaction_engine.py b/tests/test_interaction_engine.py new file mode 100644 index 0000000..8ed2f7f --- /dev/null +++ b/tests/test_interaction_engine.py @@ -0,0 +1,563 @@ +"""Tests for the interaction engine coroutines + reply registry. + +The engine (:func:`resolve_ask_user`, :func:`resolve_approval`) is pure +async with an injected :class:`InteractionIO`, so it is exercised here with +a scripted ``FakeIO``: assert the prompts it emits and feed it the replies +a user would send. Covers the whole grammar — single/multi question, +optional, choice letters, "Other", timeout, ``/stop``, approve / reject / +approve-all, and capability-driven button formatting. +""" + +import asyncio + +import pytest + +from EvoScientist.channels import interaction as I +from EvoScientist.channels.capabilities import ChannelCapabilities + +# ═══════════════════════════════════════════════════════════════════════ +# Scripted fake IO +# ═══════════════════════════════════════════════════════════════════════ + + +class FakeIO(I.InteractionIO): + """A scripted :class:`InteractionIO`. + + *replies* is the queue of reply strings ``wait_reply`` hands back in + order; a ``None`` entry (or exhausting the queue) simulates a timeout. + Every ``send`` is recorded as ``(content, metadata)`` in ``sent``. + """ + + def __init__(self, replies=None, *, capabilities=None, base_metadata=None): + self.capabilities = capabilities or ChannelCapabilities() + self.base_metadata = base_metadata + self._replies = list(replies or []) + self.sent: list[tuple[str, dict | None]] = [] + self.send_ok = True + + async def send(self, content, *, metadata=None): + self.sent.append((content, metadata)) + return self.send_ok + + async def wait_reply(self, *, timeout): + if not self._replies: + return None + return self._replies.pop(0) + + @property + def contents(self): + return [c for c, _ in self.sent] + + +QQ_CAPS = ChannelCapabilities(inline_buttons=True) + + +# ═══════════════════════════════════════════════════════════════════════ +# resolve_ask_user +# ═══════════════════════════════════════════════════════════════════════ + + +class TestResolveAskUser: + async def test_empty_questions(self): + io = FakeIO() + result = await I.resolve_ask_user([], io) + assert result == {"answers": [], "status": "answered"} + assert io.sent == [] + + async def test_single_text_answered(self): + io = FakeIO(["CIFAR-10"]) + result = await I.resolve_ask_user( + [{"question": "Which dataset?", "type": "text"}], io + ) + assert result == {"answers": ["CIFAR-10"], "status": "answered"} + assert "Quick check-in" in io.contents[0] + + async def test_multi_question_answered(self): + io = FakeIO(["ans1", "ans2"]) + result = await I.resolve_ask_user( + [ + {"question": "Q1?", "type": "text"}, + {"question": "Q2?", "type": "text"}, + ], + io, + ) + assert result == {"answers": ["ans1", "ans2"], "status": "answered"} + assert io.contents[0].startswith("❓ Question 1/2") + assert io.contents[1].startswith("❓ Question 2/2") + + async def test_choice_letter(self): + io = FakeIO(["B"]) + q = { + "question": "Which?", + "type": "multiple_choice", + "choices": [{"value": "CIFAR-10"}, {"value": "ImageNet"}], + } + result = await I.resolve_ask_user([q], io) + assert result == {"answers": ["ImageNet"], "status": "answered"} + + async def test_choice_other_subflow(self): + # "C" is the Other letter for two choices; then a free-form answer. + io = FakeIO(["C", "my custom dataset"]) + q = { + "question": "Which?", + "type": "multiple_choice", + "choices": [{"value": "CIFAR-10"}, {"value": "ImageNet"}], + } + result = await I.resolve_ask_user([q], io) + assert result == {"answers": ["my custom dataset"], "status": "answered"} + assert io.contents[1] == I.OTHER_PROMPT + + async def test_optional_suffix_in_prompt(self): + io = FakeIO(["ans"]) + await I.resolve_ask_user( + [{"question": "Notes?", "type": "text", "required": False}], io + ) + assert "(optional)" in io.contents[0] + assert "Leave empty to skip." in io.contents[0] + + async def test_optional_empty_reply_skips_and_continues(self): + io = FakeIO(["", "next"]) + result = await I.resolve_ask_user( + [ + {"question": "Notes?", "type": "text", "required": False}, + {"question": "Next?", "type": "text"}, + ], + io, + ) + assert result == {"answers": ["", "next"], "status": "answered"} + assert io.contents[1].startswith("❓ Question 2/2") + assert I.ASK_USER_TIMEOUT_FEEDBACK not in io.contents + + async def test_required_empty_reply_cancels_without_timeout_notice(self): + io = FakeIO([""]) + result = await I.resolve_ask_user( + [{"question": "Required?", "type": "text"}], + io, + ) + assert result == {"status": "cancelled"} + assert I.ASK_USER_TIMEOUT_FEEDBACK not in io.contents + + async def test_timeout_first_question(self): + io = FakeIO([]) # no replies -> timeout + result = await I.resolve_ask_user([{"question": "Q?", "type": "text"}], io) + assert result == {"status": "cancelled"} + assert io.contents[-1] == I.ASK_USER_TIMEOUT_FEEDBACK + + async def test_timeout_in_other_subflow(self): + io = FakeIO(["C"]) # picks Other, then times out on free-form + q = { + "question": "Which?", + "type": "multiple_choice", + "choices": [{"value": "A"}, {"value": "B"}], + } + result = await I.resolve_ask_user([q], io) + assert result == {"status": "cancelled"} + assert io.contents[-1] == I.ASK_USER_TIMEOUT_FEEDBACK + + async def test_stop_command_cancels(self): + io = FakeIO(["/stop"]) + result = await I.resolve_ask_user([{"question": "Q?", "type": "text"}], io) + assert result == {"status": "cancelled"} + # /stop is a pure cancel — no timeout notice sent. + assert I.ASK_USER_TIMEOUT_FEEDBACK not in io.contents + + async def test_stop_command_in_other_subflow(self): + io = FakeIO(["C", "/stop"]) + q = { + "question": "Which?", + "type": "multiple_choice", + "choices": [{"value": "A"}, {"value": "B"}], + } + result = await I.resolve_ask_user([q], io) + assert result == {"status": "cancelled"} + + async def test_cancel_reply(self): + io = FakeIO(["cancel"]) + result = await I.resolve_ask_user([{"question": "Q?", "type": "text"}], io) + assert result == {"status": "cancelled"} + + async def test_send_failure_cancels(self): + io = FakeIO(["ans"]) + io.send_ok = False + result = await I.resolve_ask_user([{"question": "Q?", "type": "text"}], io) + assert result == {"status": "cancelled"} + + +# ═══════════════════════════════════════════════════════════════════════ +# resolve_approval +# ═══════════════════════════════════════════════════════════════════════ + + +REQS = [{"name": "execute", "args": {"command": "rm -rf /tmp/x"}}] + + +class TestResolveApproval: + async def test_session_granted_short_circuits(self, monkeypatch): + monkeypatch.setattr(I, "config_auto_approve", lambda reqs: False) + p = I.ApprovalPolicy() + p.grant_session("tg:c1") + io = FakeIO() + result = await I.resolve_approval(REQS, io, p, "tg:c1") + assert result == I.ApprovalOutcome(decisions=[{"type": "approve"}]) + assert io.sent == [] # no prompt, silent + + async def test_config_auto_approve_short_circuits(self, monkeypatch): + monkeypatch.setattr(I, "config_auto_approve", lambda reqs: True) + io = FakeIO() + result = await I.resolve_approval(REQS, io, I.ApprovalPolicy(), "tg:c1") + assert result == I.ApprovalOutcome(decisions=[{"type": "approve"}]) + assert io.sent == [] + + async def test_approve(self, monkeypatch): + monkeypatch.setattr(I, "config_auto_approve", lambda reqs: False) + io = FakeIO(["1"]) + result = await I.resolve_approval(REQS, io, I.ApprovalPolicy(), "tg:c1") + assert result == I.ApprovalOutcome(decisions=[{"type": "approve"}]) + assert io.contents[0].startswith("⚠️ Approval Required") + assert io.contents[-1] == I.APPROVED_FEEDBACK + + async def test_reject(self, monkeypatch): + monkeypatch.setattr(I, "config_auto_approve", lambda reqs: False) + io = FakeIO(["2"]) + result = await I.resolve_approval(REQS, io, I.ApprovalPolicy(), "tg:c1") + assert result == I.ApprovalOutcome() + assert io.contents[-1] == I.REJECTED_FEEDBACK + + async def test_approve_all_grants_session(self, monkeypatch): + monkeypatch.setattr(I, "config_auto_approve", lambda reqs: False) + io = FakeIO(["3"]) + p = I.ApprovalPolicy() + result = await I.resolve_approval(REQS, io, p, "tg:c1") + assert result == I.ApprovalOutcome(decisions=[{"type": "approve"}]) + assert io.contents[-1] == I.APPROVED_AUTO_FEEDBACK + assert p.is_session_granted("tg:c1") # future prompts auto-approve + + async def test_multi_request_approve_length(self, monkeypatch): + monkeypatch.setattr(I, "config_auto_approve", lambda reqs: False) + reqs = [ + {"name": "execute", "args": {"command": "a"}}, + {"name": "execute", "args": {"command": "b"}}, + ] + io = FakeIO(["1"]) + result = await I.resolve_approval(reqs, io, I.ApprovalPolicy(), "tg:c1") + assert result.decisions == [{"type": "approve"}, {"type": "approve"}] + + async def test_unrecognized_reply_reported_not_judged(self, monkeypatch): + # The engine declines but hands the raw text back — the *driver* + # decides the feedback / refeed policy (consumer refeeds as a new + # turn; CLI bridge sends the unrecognized notice). + monkeypatch.setattr(I, "config_auto_approve", lambda reqs: False) + io = FakeIO(["huh?"]) + result = await I.resolve_approval(REQS, io, I.ApprovalPolicy(), "tg:c1") + assert result.decisions is None + assert result.unrecognized_reply == "huh?" + # No feedback sent by the engine itself on the unrecognized path. + assert I.UNRECOGNIZED_FEEDBACK not in io.contents + assert I.REJECTED_FEEDBACK not in io.contents + + async def test_timeout(self, monkeypatch): + monkeypatch.setattr(I, "config_auto_approve", lambda reqs: False) + io = FakeIO([]) # times out + result = await I.resolve_approval(REQS, io, I.ApprovalPolicy(), "tg:c1") + assert result == I.ApprovalOutcome() + assert io.contents[-1] == I.APPROVAL_TIMEOUT_FEEDBACK + + async def test_stop_command_silent_cancel(self, monkeypatch): + monkeypatch.setattr(I, "config_auto_approve", lambda reqs: False) + io = FakeIO(["/stop"]) + result = await I.resolve_approval(REQS, io, I.ApprovalPolicy(), "tg:c1") + assert result == I.ApprovalOutcome() + # /stop already got its own ack; no reject/unrecognized feedback here. + assert I.REJECTED_FEEDBACK not in io.contents + assert I.UNRECOGNIZED_FEEDBACK not in io.contents + + async def test_send_failure_declines(self, monkeypatch): + monkeypatch.setattr(I, "config_auto_approve", lambda reqs: False) + io = FakeIO(["1"]) + io.send_ok = False + result = await I.resolve_approval(REQS, io, I.ApprovalPolicy(), "tg:c1") + assert result == I.ApprovalOutcome() + + # ── R3: button-capability formatting + payload normalization ── + + async def test_buttons_attached_when_capable(self, monkeypatch): + monkeypatch.setattr(I, "config_auto_approve", lambda reqs: False) + io = FakeIO(["1"], capabilities=QQ_CAPS, base_metadata={"chat": "x"}) + await I.resolve_approval(REQS, io, I.ApprovalPolicy(), "tg:c1") + prompt, metadata = io.sent[0] + # Button channels drop the textual "Reply: 1=..." cue. + assert "Reply: 1=Approve" not in prompt + assert metadata["buttons"] == [ + {"text": "Approve", "value": "1", "type": "primary"}, + {"text": "Reject", "value": "2", "type": "danger"}, + {"text": "Approve all", "value": "3"}, + ] + assert metadata["chat"] == "x" # base metadata preserved + + async def test_no_buttons_when_incapable(self, monkeypatch): + monkeypatch.setattr(I, "config_auto_approve", lambda reqs: False) + io = FakeIO(["1"]) # default caps: no inline_buttons + await I.resolve_approval(REQS, io, I.ApprovalPolicy(), "tg:c1") + prompt, metadata = io.sent[0] + assert "Reply: 1=Approve, 2=Reject, 3=Approve all" in prompt + assert "buttons" not in (metadata or {}) + + async def test_button_press_payload_normalizes(self, monkeypatch): + # A button click delivers its `value` ("3") through the same reply + # path; the engine must treat it exactly like a typed "3". + monkeypatch.setattr(I, "config_auto_approve", lambda reqs: False) + io = FakeIO(["3"], capabilities=QQ_CAPS) + p = I.ApprovalPolicy() + result = await I.resolve_approval(REQS, io, p, "tg:c1") + assert result == I.ApprovalOutcome(decisions=[{"type": "approve"}]) + assert p.is_session_granted("tg:c1") + + +# ═══════════════════════════════════════════════════════════════════════ +# PendingReplyRegistry +# ═══════════════════════════════════════════════════════════════════════ + + +class TestPendingReplyRegistry: + async def test_register_wait_resolve(self): + reg = I.PendingReplyRegistry() + + async def _resolver(): + # Give wait() a tick to register before resolving. + await asyncio.sleep(0.01) + assert reg.try_resolve("s1", "hello") is True + + got, _ = await asyncio.gather(reg.wait("s1", timeout=1.0), _resolver()) + assert got == "hello" + assert "s1" not in reg # cleaned up after wait + + async def test_wait_event_returns_reply_context(self): + reg = I.PendingReplyRegistry() + context = object() + + async def _resolver(): + await asyncio.sleep(0.01) + assert reg.try_resolve("s1", "hello", context=context) is True + + got, _ = await asyncio.gather(reg.wait_event("s1", timeout=1.0), _resolver()) + assert got is not None + assert got.content == "hello" + assert got.context is context + assert "s1" not in reg + + async def test_wait_timeout_returns_none(self): + reg = I.PendingReplyRegistry() + got = await reg.wait("s1", timeout=0.02) + assert got is None + assert "s1" not in reg + + def test_try_resolve_no_pending(self): + reg = I.PendingReplyRegistry() + assert reg.try_resolve("nope", "x") is False + + async def test_reregister_cancels_stale(self): + reg = I.PendingReplyRegistry() + first = asyncio.ensure_future(reg.wait("s1", timeout=1.0)) + await asyncio.sleep(0.01) # let first register + # A second interaction on the same chat re-registers, cancelling the + # first waiter so it unwinds promptly (returns None) instead of + # hanging until its own timeout. + new_fut = reg.register("s1") + assert await first is None + # The stale waiter's cleanup must not evict the newer registration. + assert reg._pending.get("s1") is new_fut + reg.discard("s1") + + async def test_clear_cancels_all(self): + reg = I.PendingReplyRegistry() + fut = asyncio.ensure_future(reg.wait("s1", timeout=1.0)) + await asyncio.sleep(0.01) + reg.clear() + got = await fut + assert got is None + + async def test_task_cancellation_propagates(self): + reg = I.PendingReplyRegistry() + task = asyncio.create_task(reg.wait("s1", timeout=1.0)) + await asyncio.sleep(0.01) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + assert "s1" not in reg + + +# ═══════════════════════════════════════════════════════════════════════ +# Consumer driver: unrecognized-reply refeed (serve-mode policy) +# ═══════════════════════════════════════════════════════════════════════ +# Pre-engine semantics that must survive the extraction: an unrecognized +# reply while a HITL approval is pending REJECTS the pending action, sends +# the rejection feedback, and the user's text is then processed as a NEW +# agent turn — a user who ignores the prompt and types a fresh instruction +# must not lose it. (The CLI bridge deliberately does NOT refeed; see +# tests/test_cli_channel_bridge.py) + + +class TestConsumerUnrecognizedRefeed: + async def test_unrecognized_reply_rejects_and_refeeds(self, monkeypatch): + from unittest.mock import MagicMock + + from EvoScientist.channels.bus.events import ( + InboundMessage as BusInbound, + ) + from EvoScientist.channels.bus.message_bus import MessageBus + from EvoScientist.channels.channel_manager import ChannelManager + from EvoScientist.channels.consumer import InboundConsumer + from tests.fakes import FakeGraphGateway, StubChannel + + # Force the manual-prompt path (no config auto-approve). + monkeypatch.setattr(I, "config_auto_approve", lambda reqs: False) + + bus = MessageBus() + mgr = ChannelManager(bus) + mgr.register(StubChannel()) + + stream_calls = 0 + + async def _fake_stream(request): + nonlocal stream_calls + stream_calls += 1 + if stream_calls == 1: + # First turn hits a HITL interrupt. + yield { + "type": "interrupt", + "interrupt_id": "main", + "action_requests": [ + {"name": "execute", "args": {"command": "rm -rf /x"}} + ], + "review_configs": [], + } + return + # The refeed turn: echo what we were given. + yield {"type": "text", "content": f"handled: {request.message}"} + yield {"type": "done", "content": f"handled: {request.message}"} + + gateway = FakeGraphGateway( + stream=_fake_stream, + generated_thread_ids=["thread-original", "thread-reply"], + ) + consumer = InboundConsumer( + bus=bus, + manager=mgr, + agent=MagicMock(), + thread_id="", + graph_gateway=gateway, + max_concurrent=2, + max_pending=10, + inference_timeout=5.0, + drain_timeout=1.0, + ) + + task = asyncio.create_task(consumer.run()) + try: + await bus.publish_inbound( + BusInbound( + channel="stub", + sender_id="u1", + chat_id="c1", + content="do the thing", + message_id="msg-original", + metadata={"origin": "original"}, + ) + ) + + # 1. Approval prompt goes out. + prompt = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0) + assert prompt.content.startswith("⚠️ Approval Required") + assert prompt.metadata == {"origin": "original"} + + # 2. User ignores the prompt and types a fresh instruction. + await bus.publish_inbound( + BusInbound( + channel="stub", + sender_id="u2", + chat_id="c1", + content="actually, summarize the report", + message_id="msg-reply", + media=["file-report.pdf"], + metadata={"origin": "reply"}, + ) + ) + + # 3. Pending action is rejected with the old serve feedback... + feedback = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0) + assert feedback.content == I.REJECTED_FEEDBACK + assert feedback.metadata == {"origin": "original"} + + # 4. ...and the text is processed as a NEW agent turn. + response = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0) + assert response.content == "handled: actually, summarize the report" + assert response.reply_to == "msg-reply" + assert response.metadata == {"origin": "reply"} + + # The refeed reached the stream path as its own request. + assert stream_calls == 2 + assert gateway.requests[0].thread_id == "thread-original" + assert gateway.requests[-1].message == "actually, summarize the report" + assert gateway.requests[-1].thread_id == "thread-reply" + assert gateway.requests[-1].media == ["file-report.pdf"] + finally: + await consumer.stop() + await task + + +class TestApprovalEmptyReply: + async def test_empty_reply_is_unrecognized_not_timeout(self, monkeypatch): + """A media-only/empty reply must reach the refeed path, not timeout.""" + monkeypatch.setattr(I, "config_auto_approve", lambda reqs: False) + io = FakeIO(replies=[""]) + policy = I.ApprovalPolicy() + outcome = await I.resolve_approval( + [{"name": "execute", "args": {"command": "ls"}}], + io, + policy, + "stub:c1", + timeout=1.0, + ) + assert outcome.decisions is None + assert outcome.unrecognized_reply == "" + sent = [content for content, _ in io.sent] + assert I.APPROVAL_TIMEOUT_FEEDBACK not in sent + + +class TestReplyInterceptionSkipsThreadCreation: + async def test_consumed_reply_creates_no_thread(self): + """A registry-consumed reply must not create a graph thread or touch + the sender-session LRU.""" + from unittest.mock import AsyncMock, MagicMock + + from EvoScientist.channels.bus.events import InboundMessage as BusInbound + from EvoScientist.channels.consumer import InboundConsumer + + gateway = MagicMock() + gateway.create_thread = AsyncMock(return_value="t-should-not-exist") + consumer = InboundConsumer( + bus=MagicMock(), + manager=MagicMock(), + agent=MagicMock(), + thread_id="", + graph_gateway=gateway, + max_concurrent=1, + max_pending=5, + inference_timeout=1.0, + drain_timeout=0.5, + ) + msg = BusInbound( + channel="stub", + sender_id="u1", + chat_id="c1", + content="1", + message_id="m1", + ) + fut = consumer._reply_registry.register(msg.session_key) + + await consumer._handle_message(msg) + + assert fut.done() + assert fut.result().content == "1" + gateway.create_thread.assert_not_awaited() + assert consumer._sessions == {} diff --git a/tests/test_interaction_grammar.py b/tests/test_interaction_grammar.py new file mode 100644 index 0000000..95703a7 --- /dev/null +++ b/tests/test_interaction_grammar.py @@ -0,0 +1,341 @@ +"""Tests the shared interaction grammar for ``channels.interaction``.""" + +from typing import ClassVar +from unittest.mock import MagicMock, patch + +import pytest + +from EvoScientist.channels import interaction as I + +# ═══════════════════════════════════════════════════════════════════════ +# Stop / cancel grammar +# ═══════════════════════════════════════════════════════════════════════ + + +class TestStopCommand: + @pytest.mark.parametrize( + "text", + ["/stop", "/cancel", " /stop ", "/STOP", "/Cancel", "\t/stop\n"], + ) + def test_stop_recognized(self, text): + assert I.is_stop_command(text) is True + + @pytest.mark.parametrize( + "text", + ["stop", "cancel", "/stopp", "1", "", None, "please /stop"], + ) + def test_stop_not_recognized(self, text): + assert I.is_stop_command(text) is False + + @pytest.mark.parametrize("text", ["cancel", "CANCEL", " Cancel ", "\tcancel"]) + def test_cancel_recognized(self, text): + assert I.is_cancel_reply(text) is True + + @pytest.mark.parametrize("text", ["/cancel", "cancelled", "c", "", None]) + def test_cancel_not_recognized(self, text): + assert I.is_cancel_reply(text) is False + + +# ═══════════════════════════════════════════════════════════════════════ +# Approval reply grammar +# ═══════════════════════════════════════════════════════════════════════ + + +class TestParseApprovalReply: + @pytest.mark.parametrize( + ("text", "expected"), + [ + # approve + ("1", "approve"), + ("y", "approve"), + ("yes", "approve"), + ("approve", "approve"), + ("ok", "approve"), + (" 1 ", "approve"), + (" Y ", "approve"), + ("YES", "approve"), + # reject + ("2", "reject"), + ("n", "reject"), + ("no", "reject"), + ("reject", "reject"), + ("REJECT", "reject"), + # auto / approve-all + ("3", "auto"), + ("a", "auto"), + ("auto", "auto"), + ("approve all", "auto"), + ("APPROVE ALL", "auto"), + # unrecognized + ("hello world", None), + ("", None), + ("maybe", None), + ("4", None), + ], + ) + def test_parse(self, text, expected): + assert I.parse_approval_reply(text) == expected + + def test_button_values_normalize_to_decisions(self): + # Feishu/QQ buttons deliver their `value` ("1"/"2"/"3") through + # the same reply path, so the shared parser must map them + # identically to a typed reply. + buttons = I.approval_prompt_metadata(None, with_buttons=True)["buttons"] + values = [b["value"] for b in buttons] + assert values == ["1", "2", "3"] + assert [I.parse_approval_reply(v) for v in values] == [ + "approve", + "reject", + "auto", + ] + + def test_approve_decisions_length(self): + assert I.approve_decisions([{"name": "a"}, {"name": "b"}]) == [ + {"type": "approve"}, + {"type": "approve"}, + ] + # empty request list still yields a single approve (Command shape) + assert I.approve_decisions([]) == [{"type": "approve"}] + + +# ═══════════════════════════════════════════════════════════════════════ +# ask_user choice grammar (letters + "Other") +# ═══════════════════════════════════════════════════════════════════════ + + +class TestParseChoiceAnswer: + CHOICES: ClassVar = [{"value": "CIFAR-10"}, {"value": "ImageNet"}] + + def test_letter_selects_choice(self): + assert I.parse_choice_answer("A", self.CHOICES) == ("answer", "CIFAR-10") + assert I.parse_choice_answer("b", self.CHOICES) == ("answer", "ImageNet") + + def test_other_letter(self): + # Two choices -> "Other" is C. + assert I.parse_choice_answer("C", self.CHOICES) == ("other", None) + assert I.parse_choice_answer("c", self.CHOICES) == ("other", None) + + def test_out_of_range_letter_is_literal(self): + # Z is a single alpha char but past the choice range -> literal answer. + assert I.parse_choice_answer("Z", self.CHOICES) == ("answer", "Z") + + def test_multichar_reply_is_literal(self): + assert I.parse_choice_answer("CIFAR-10", self.CHOICES) == ( + "answer", + "CIFAR-10", + ) + + def test_no_choices_other_is_a(self): + assert I.parse_choice_answer("A", []) == ("other", None) + + +# ═══════════════════════════════════════════════════════════════════════ +# ApprovalPolicy (config rule + session registry + session key) +# ═══════════════════════════════════════════════════════════════════════ + + +class TestApprovalPolicy: + def test_grant_and_is_granted(self): + p = I.ApprovalPolicy() + assert p.is_session_granted("tg:c1") is False + p.grant_session("tg:c1") + assert p.is_session_granted("tg:c1") is True + p.clear_sessions() + assert p.is_session_granted("tg:c1") is False + + def test_auto_decision_session_granted(self): + p = I.ApprovalPolicy() + p.grant_session("tg:c1") + reqs = [{"name": "execute", "args": {"command": "rm -rf /"}}] + # Session grant short-circuits config entirely. + assert p.auto_decision("tg:c1", reqs) == [{"type": "approve"}] + + def test_auto_decision_config_true(self): + p = I.ApprovalPolicy() + cfg = MagicMock() + cfg.auto_approve = True + with patch("EvoScientist.config.settings.load_config", return_value=cfg): + reqs = [{"name": "execute", "args": {"command": "rm -rf /"}}] + assert p.auto_decision("tg:c1", reqs) == [{"type": "approve"}] + + def test_auto_decision_needs_prompt(self): + p = I.ApprovalPolicy() + cfg = MagicMock() + cfg.auto_approve = False + cfg.shell_allow_list = "" + with patch("EvoScientist.config.settings.load_config", return_value=cfg): + reqs = [{"name": "execute", "args": {"command": "rm -rf /"}}] + assert p.auto_decision("tg:c1", reqs) is None + + +class TestConfigAutoApprove: + def test_empty(self): + assert I.config_auto_approve([]) is True + + def test_non_execute(self): + assert I.config_auto_approve([{"name": "write_file", "args": {}}]) is True + + def test_execute_no_allowlist(self): + cfg = MagicMock() + cfg.auto_approve = False + cfg.shell_allow_list = "" + with patch("EvoScientist.config.settings.load_config", return_value=cfg): + assert ( + I.config_auto_approve( + [{"name": "execute", "args": {"command": "rm -rf /"}}] + ) + is False + ) + + def test_execute_allowlist_match(self): + cfg = MagicMock() + cfg.auto_approve = False + cfg.shell_allow_list = "ls,python" + with patch("EvoScientist.config.settings.load_config", return_value=cfg): + assert ( + I.config_auto_approve( + [{"name": "execute", "args": {"command": "ls -la"}}] + ) + is True + ) + + def test_run_in_background_not_allowlisted(self): + cfg = MagicMock() + cfg.auto_approve = False + cfg.shell_allow_list = "ls,cat" + with patch("EvoScientist.config.settings.load_config", return_value=cfg): + assert ( + I.config_auto_approve( + [{"name": "run_in_background", "args": {"command": "rm -rf /"}}] + ) + is False + ) + + def test_fail_closed_on_config_error(self): + with patch( + "EvoScientist.config.settings.load_config", side_effect=RuntimeError("boom") + ): + assert ( + I.config_auto_approve([{"name": "execute", "args": {"command": "ls"}}]) + is False + ) + + +# ═══════════════════════════════════════════════════════════════════════ +# Prompt-format checks. +# ═══════════════════════════════════════════════════════════════════════ + + +class TestApprovalPromptFormat: + def test_lists_each_action_with_reply_options(self): + got = I.format_approval_prompt( + [ + {"name": "execute", "args": {"command": "ls"}}, + {"name": "write_file", "args": {"path": "/out.txt"}}, + ] + ) + assert "execute: ls" in got + assert "write_file: /out.txt" in got + # The offered options must match what parse_approval_reply accepts. + for option in ("1=Approve", "2=Reject", "3=Approve all"): + assert option in got + + def test_with_buttons_drops_text_instruction(self): + got = I.format_approval_prompt( + [{"name": "execute", "args": {"command": "ls -la"}}], + with_buttons=True, + ) + assert "execute: ls -la" in got + assert "1=Approve" not in got # buttons replace the typed-reply hint + + def test_no_command_falls_back_to_name(self): + got = I.format_approval_prompt([{"name": "ask_user", "args": {}}]) + assert "ask_user" in got + + def test_metadata_no_buttons(self): + assert I.approval_prompt_metadata({"k": "v"}, with_buttons=False) == {"k": "v"} + + def test_metadata_with_buttons(self): + md = I.approval_prompt_metadata({"k": "v"}, with_buttons=True) + assert md["k"] == "v" + # Button values must be replies parse_approval_reply understands. + assert [b["value"] for b in md["buttons"]] == ["1", "2", "3"] + + +class TestQuestionPromptFormat: + def test_single_question_offers_cancel(self): + got = I.format_question_prompt( + {"question": "What dataset?", "type": "text"}, 0, 1 + ) + assert "What dataset?" in got + assert "cancel" in got + + def test_optional_question_is_marked_and_skippable(self): + got = I.format_question_prompt( + {"question": "Notes?", "type": "text", "required": False}, 0, 1 + ) + assert "(optional)" in got + assert "skip" in got.lower() + + def test_multi_question_header_shows_position(self): + got = I.format_question_prompt( + { + "question": "Which?", + "type": "multiple_choice", + "choices": [{"value": "A"}, {"value": "B"}], + "required": False, + }, + 1, + 3, + ) + assert "2/3" in got + + def test_choices_get_letters_plus_other(self): + got = I.format_question_prompt( + { + "question": "Pick one", + "type": "multiple_choice", + "choices": [{"value": "CIFAR-10"}, {"value": "ImageNet"}], + }, + 0, + 1, + ) + # Displayed letters must match what the choice parser accepts, with + # the "Other" free-form option appended after the real choices. + assert "A. CIFAR-10" in got + assert "B. ImageNet" in got + assert "C. Other" in got + + +class TestChoiceNormalization: + """Choices arrive from model tool args — plain strings must not crash.""" + + def test_prompt_renders_plain_string_choices(self): + got = I.format_question_prompt( + { + "question": "Pick one", + "type": "multiple_choice", + "choices": ["CIFAR-10", "ImageNet"], + }, + 0, + 1, + ) + assert "A. CIFAR-10" in got + assert "B. ImageNet" in got + assert "C. Other" in got + + def test_parse_returns_plain_string_choice(self): + kind, value = I.parse_choice_answer("b", ["CIFAR-10", "ImageNet"]) + assert (kind, value) == ("answer", "ImageNet") + + def test_mixed_dict_and_string_choices(self): + choices = [{"value": "CIFAR-10"}, "ImageNet"] + got = I.format_question_prompt( + {"question": "Pick", "type": "multiple_choice", "choices": choices}, + 0, + 1, + ) + assert "A. CIFAR-10" in got + assert "B. ImageNet" in got + kind, value = I.parse_choice_answer("a", choices) + assert (kind, value) == ("answer", "CIFAR-10")