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:
+136
-373
@@ -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 ──
|
||||
|
||||
|
||||
@@ -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)
|
||||
@@ -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
@@ -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
|
||||
)
|
||||
|
||||
@@ -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."
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
# =============================================================================
|
||||
|
||||
@@ -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 == {}
|
||||
@@ -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")
|
||||
Reference in New Issue
Block a user