refactor: extract a shared HITL/ask_user interaction engine (#342)

* 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>
This commit is contained in:
dinos
2026-07-14 16:59:41 +02:00
committed by GitHub
parent 753c745405
commit db1abce8d8
9 changed files with 2296 additions and 697 deletions
+136 -373
View File
@@ -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 ──
+540
View File
@@ -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)
+3 -3
View File
@@ -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:
+208 -245
View File
@@ -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
)
+18 -11
View File
@@ -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."
+426
View File
@@ -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)]
+61 -65
View File
@@ -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
# =============================================================================
+563
View File
@@ -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 == {}
+341
View File
@@ -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")