From 438f9e56d5b15a6949649285b27660477be99db9 Mon Sep 17 00:00:00 2001 From: Ziheng Zhang <142805986+MuXinCG2004@users.noreply.github.com> Date: Thu, 12 Mar 2026 19:28:37 +0800 Subject: [PATCH] Fix: channel can't ask user (#23) * small fix * small fix * remove local file * remove local file --- EvoScientist/channels/consumer.py | 172 ++++++++++++++++++++++++++++ EvoScientist/cli/commands.py | 5 + EvoScientist/cli/tui_interactive.py | 18 ++- 3 files changed, 193 insertions(+), 2 deletions(-) diff --git a/EvoScientist/channels/consumer.py b/EvoScientist/channels/consumer.py index 23126fe..7aa58d6 100644 --- a/EvoScientist/channels/consumer.py +++ b/EvoScientist/channels/consumer.py @@ -153,6 +153,13 @@ class _PendingInterrupt: 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 + + class InboundConsumer: """Consume inbound messages from the bus, process via agent, publish outbound. @@ -237,6 +244,9 @@ class InboundConsumer: 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] = {} + def _get_thread_id(self, sender_id: str) -> str: """Get or create a thread ID for the given sender. @@ -357,6 +367,14 @@ 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 + # 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] @@ -447,6 +465,10 @@ class InboundConsumer: interrupt_data = event break # exit async for to handle interrupt + elif event_type == "ask_user": + interrupt_data = event + break # exit async for to handle ask_user + # Flush thinking if thinking_buffer and not thinking_sent and channel: full_thinking = "".join(thinking_buffer) @@ -473,6 +495,15 @@ class InboundConsumer: pass return # done + # ask_user: send questions to channel user, collect answers + if interrupt_data.get("type") == "ask_user": + result = await self._resolve_ask_user( + msg, interrupt_data, session_key, + ) + from langgraph.types import Command # type: ignore[import-untyped] + stream_input = Command(resume=result) + continue + # HITL: resolve the interrupt action_reqs = interrupt_data.get("action_requests", []) n = len(action_reqs) or 1 @@ -588,6 +619,147 @@ class InboundConsumer: "sessions": len(self._sessions), } + # ── 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 asyncio.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. + + Mirrors the logic of ``cli.channel.channel_ask_user_prompt`` but runs + fully async inside the consumer event loop. + + 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, _HITL_APPROVAL_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, _HITL_APPROVAL_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"} + # ── internal ── def _evict_chat_locks(self) -> None: diff --git a/EvoScientist/cli/commands.py b/EvoScientist/cli/commands.py index 167e08c..b94210b 100644 --- a/EvoScientist/cli/commands.py +++ b/EvoScientist/cli/commands.py @@ -24,6 +24,7 @@ from .channel import ( _message_queue, _set_channel_response, _start_channels_bus_mode, + channel_ask_user_prompt, channel_hitl_prompt, ) from .tui_runtime import run_streaming @@ -170,6 +171,9 @@ def _serve_process_message( def _hitl_prompt(action_requests: list) -> list[dict] | None: return channel_hitl_prompt(action_requests, msg) + def _ask_user_prompt(ask_user_data: dict) -> dict: + return channel_ask_user_prompt(ask_user_data, msg) + meta = build_metadata(workspace_dir, model) try: response = run_streaming( @@ -184,6 +188,7 @@ def _serve_process_message( on_todo=_send_todo, on_file_write=_send_media, hitl_prompt_fn=_hitl_prompt, + ask_user_prompt_fn=_ask_user_prompt, ) except Exception as e: response = f"Error: {e}" diff --git a/EvoScientist/cli/tui_interactive.py b/EvoScientist/cli/tui_interactive.py index 1db34f6..aba101c 100644 --- a/EvoScientist/cli/tui_interactive.py +++ b/EvoScientist/cli/tui_interactive.py @@ -496,6 +496,7 @@ def run_textual_interactive( on_media_cb: Callable[[str], None] | None = None, skip_user_message: bool = False, channel_hitl_fn: Callable[[list], list[dict] | None] | None = None, + channel_ask_user_fn: Callable[[dict], dict] | None = None, ) -> str: """Stream agent events and mount widgets. Returns response text. @@ -508,6 +509,9 @@ def run_textual_interactive( channel_hitl_fn: Optional channel-based HITL approval function. When provided (channel messages), this is called instead of mounting the ApprovalWidget. + channel_ask_user_fn: Optional channel-based ask_user function. + When provided (channel messages), this is called instead + of mounting the AskUserWidget. """ container = self.query_one("#chat", VerticalScroll) @@ -896,13 +900,14 @@ def run_textual_interactive( questions = event.get("questions", []) if questions: # Channel messages: use channel-based text prompt - if channel_hitl_fn is not None: + if channel_ask_user_fn is not None: self._append_system( "Waiting for channel user input...", style="dim italic", ) + _ask_fn = channel_ask_user_fn result = await asyncio.to_thread( - lambda: _channel_ask_user_prompt_from_event(event, channel_hitl_fn), + lambda: _ask_fn(event), ) else: # Interactive TUI: display widget, collect via arrow keys @@ -1209,6 +1214,14 @@ def run_textual_interactive( """ return _ch_mod.channel_hitl_prompt(action_requests, msg) + def _channel_ask_user(ask_user_data: dict) -> dict: + """Send ask_user questions to channel user and wait for reply. + + This runs in a thread (called via asyncio.to_thread) so it can + block without freezing the Textual event loop. + """ + return _ch_mod.channel_ask_user_prompt(ask_user_data, msg) + response = "" try: response = await self._stream_with_widgets( @@ -1218,6 +1231,7 @@ def run_textual_interactive( on_media_cb=_send_media, skip_user_message=True, channel_hitl_fn=_channel_hitl_prompt, + channel_ask_user_fn=_channel_ask_user, ) except Exception as exc: response = f"Error: {exc}"