Fix: channel can't ask user (#23)

* small fix

* small fix

* remove local file

* remove local file
This commit is contained in:
Ziheng Zhang
2026-03-12 19:28:37 +08:00
committed by GitHub
parent 9bbde69f7f
commit 438f9e56d5
3 changed files with 193 additions and 2 deletions
+172
View File
@@ -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:
+5
View File
@@ -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}"
+16 -2
View File
@@ -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}"