diff --git a/EvoScientist/EvoScientist.py b/EvoScientist/EvoScientist.py
index 7cdca10..9320764 100644
--- a/EvoScientist/EvoScientist.py
+++ b/EvoScientist/EvoScientist.py
@@ -275,11 +275,16 @@ def _get_default_middleware():
"""Build the default middleware list."""
from .middleware import create_memory_middleware, ToolErrorHandlerMiddleware
+ cfg = _ensure_config()
memory_dir = str(_paths_mod.MEMORY_DIR)
- return [
+ mw = [
ToolErrorHandlerMiddleware(),
create_memory_middleware(memory_dir, extraction_model=_ensure_chat_model()),
]
+ if cfg.enable_ask_user and not cfg.auto_approve:
+ from .middleware.ask_user import AskUserMiddleware
+ mw.insert(0, AskUserMiddleware())
+ return mw
def _get_default_agent():
@@ -389,6 +394,9 @@ def create_cli_agent(workspace_dir: str | None = None, checkpointer=None, config
ToolErrorHandlerMiddleware(),
create_memory_middleware(_mem_dir, extraction_model=_ensure_chat_model()),
]
+ if cfg.enable_ask_user and not cfg.auto_approve:
+ from .middleware.ask_user import AskUserMiddleware
+ mw.insert(0, AskUserMiddleware())
# Re-load MCP tools from current config (picks up /mcp add changes)
kwargs = load_mcp_and_build_kwargs(be, mw)
diff --git a/EvoScientist/cli/channel.py b/EvoScientist/cli/channel.py
index 38679b9..4a15fee 100644
--- a/EvoScientist/cli/channel.py
+++ b/EvoScientist/cli/channel.py
@@ -122,6 +122,127 @@ def _try_set_hitl_reply(channel_type: str, chat_id: str, content: str) -> bool:
return False
+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.
+
+ If *msg* is provided, sends questions via the bus and waits for a reply.
+ Otherwise falls back to returning a cancelled result.
+
+ 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:
+ 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,
+ )),
+ 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=_HITL_APPROVAL_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 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=_HITL_APPROVAL_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 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"}
+
+
def channel_hitl_prompt(
action_requests: list,
msg: "ChannelMessage",
diff --git a/EvoScientist/cli/commands.py b/EvoScientist/cli/commands.py
index 415f669..167e08c 100644
--- a/EvoScientist/cli/commands.py
+++ b/EvoScientist/cli/commands.py
@@ -202,6 +202,7 @@ def serve(
no_thinking: bool = typer.Option(False, "--no-thinking", help="Disable thinking relay to channels"),
workdir: Optional[str] = typer.Option(None, "--workdir", help="Override workspace directory"),
auto_approve: bool = typer.Option(False, "--auto-approve", help="Auto-approve all tool executions without prompting"),
+ ask_user: bool = typer.Option(False, "--ask-user", help="Enable agent to ask clarifying questions about your research preferences"),
):
"""Run EvoScientist in headless mode -- channels only, no interactive prompt.
@@ -219,6 +220,8 @@ def serve(
cli_overrides = {}
if auto_approve:
cli_overrides["auto_approve"] = True
+ if ask_user:
+ cli_overrides["enable_ask_user"] = True
config = get_effective_config(cli_overrides)
apply_config_to_env(config)
@@ -539,6 +542,7 @@ def _main_callback(
use_cwd: bool = typer.Option(False, "--use-cwd", help="Use current working directory as workspace"),
no_thinking: bool = typer.Option(False, "--no-thinking", help="Disable thinking display"),
auto_approve: bool = typer.Option(False, "--auto-approve", help="Auto-approve all tool executions without prompting"),
+ ask_user: bool = typer.Option(False, "--ask-user", help="Enable agent to ask clarifying questions about your research preferences"),
ui: Optional[str] = typer.Option(
None,
"--ui",
@@ -569,6 +573,8 @@ def _main_callback(
cli_overrides["ui_backend"] = ui
if auto_approve:
cli_overrides["auto_approve"] = True
+ if ask_user:
+ cli_overrides["enable_ask_user"] = True
config = get_effective_config(cli_overrides)
apply_config_to_env(config)
diff --git a/EvoScientist/cli/interactive.py b/EvoScientist/cli/interactive.py
index f272531..9328dae 100644
--- a/EvoScientist/cli/interactive.py
+++ b/EvoScientist/cli/interactive.py
@@ -534,6 +534,10 @@ def cmd_interactive(
"""Send HITL approval prompt to channel user and wait for reply."""
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."""
+ return _ch_mod.channel_ask_user_prompt(ask_user_data, msg)
+
meta = build_metadata(state["workspace_dir"], model)
try:
response = run_streaming(
@@ -548,6 +552,7 @@ def cmd_interactive(
on_todo=_send_todo_to_channel,
on_file_write=_send_media_to_channel,
hitl_prompt_fn=_channel_hitl_prompt,
+ ask_user_prompt_fn=_channel_ask_user,
)
except Exception as e:
response = f"Error: {e}"
diff --git a/EvoScientist/cli/tui_backends.py b/EvoScientist/cli/tui_backends.py
index 1d8bc87..e5c4d9b 100644
--- a/EvoScientist/cli/tui_backends.py
+++ b/EvoScientist/cli/tui_backends.py
@@ -26,6 +26,7 @@ class StreamingTUIBackend(Protocol):
on_file_write: Callable[[str], None] | None = None,
metadata: dict | None = None,
hitl_prompt_fn: Callable[[list], list[dict] | None] | None = None,
+ ask_user_prompt_fn: Callable[[dict], dict] | None = None,
) -> str:
"""Run streaming and return final response text."""
@@ -49,6 +50,7 @@ class RichStreamingBackend:
on_file_write: Callable[[str], None] | None = None,
metadata: dict | None = None,
hitl_prompt_fn: Callable[[list], list[dict] | None] | None = None,
+ ask_user_prompt_fn: Callable[[dict], dict] | None = None,
) -> str:
return _run_streaming(
agent=agent,
@@ -61,4 +63,5 @@ class RichStreamingBackend:
on_file_write=on_file_write,
metadata=metadata,
hitl_prompt_fn=hitl_prompt_fn,
+ ask_user_prompt_fn=ask_user_prompt_fn,
)
diff --git a/EvoScientist/cli/tui_interactive.py b/EvoScientist/cli/tui_interactive.py
index a94b8b7..1db34f6 100644
--- a/EvoScientist/cli/tui_interactive.py
+++ b/EvoScientist/cli/tui_interactive.py
@@ -72,6 +72,18 @@ def _shorten_path(path: str) -> str:
return _sp(path)
+def _channel_ask_user_prompt_from_event(event: dict, channel_hitl_fn: Any) -> dict:
+ """Bridge ask_user event to channel prompt (called via asyncio.to_thread).
+
+ This is a thin wrapper used by TUI channel mode — it delegates to the
+ channel module's ``channel_ask_user_prompt`` if available, otherwise
+ returns a cancelled result.
+ """
+ from .channel import channel_ask_user_prompt as _ch_ask
+
+ return _ch_ask(event)
+
+
def _build_welcome_banner(
*,
thread_id: str,
@@ -308,6 +320,7 @@ def run_textual_interactive(
self._comp_index: int = -1
self._hitl_auto_approve: bool = False
self._approval_future: asyncio.Future | None = None
+ self._ask_user_future: asyncio.Future | None = None
self._picker_future: asyncio.Future | None = None
self._history_suggester = HistorySuggester(get_config_dir() / "history")
@@ -419,6 +432,32 @@ def run_textual_interactive(
if self._approval_future and not self._approval_future.done():
self._approval_future.set_result(event)
+ async def _wait_for_ask_user(self, ask_w) -> dict:
+ """Wait for the interactive ask_user widget to resolve via Future.
+
+ Returns ``{"answers": [...], "status": "answered"}``
+ or ``{"status": "cancelled"}``.
+ """
+ loop = asyncio.get_running_loop()
+ self._ask_user_future = loop.create_future()
+ ask_w.set_future(self._ask_user_future)
+
+ try:
+ result = await asyncio.wait_for(self._ask_user_future, timeout=300)
+ except (asyncio.TimeoutError, asyncio.CancelledError):
+ ask_w.action_cancel()
+ return {"status": "cancelled"}
+ finally:
+ self._ask_user_future = None
+
+ if not isinstance(result, dict):
+ return {"status": "cancelled"}
+
+ result_type = result.get("type", "")
+ if result_type == "answered":
+ return {"answers": result.get("answers", []), "status": "answered"}
+ return {"status": "cancelled"}
+
async def _wait_for_thread_pick(self, picker_widget) -> str | None:
"""Wait for user to pick a thread from ThreadPickerWidget.
@@ -605,7 +644,12 @@ def run_textual_interactive(
for _hitl_round in range(_MAX_HITL_ROUNDS):
state.pending_interrupt = None
+ state.pending_ask_user = None
_hitl_resuming = False
+ # Reset per-round widgets so resumed streams get fresh ones
+ if _hitl_round > 0:
+ thinking_w = None
+ summarization_w = None
try:
async for event in stream_agent_events(
self._agent,
@@ -848,6 +892,38 @@ def run_textual_interactive(
if sa_w is not None:
sa_w.finalize()
+ elif event_type == "ask_user":
+ questions = event.get("questions", [])
+ if questions:
+ # Channel messages: use channel-based text prompt
+ if channel_hitl_fn is not None:
+ self._append_system(
+ "Waiting for channel user input...",
+ style="dim italic",
+ )
+ result = await asyncio.to_thread(
+ lambda: _channel_ask_user_prompt_from_event(event, channel_hitl_fn),
+ )
+ else:
+ # Interactive TUI: display widget, collect via arrow keys
+ from .widgets.ask_user_widget import AskUserWidget
+ _prompt = self.query_one("#prompt", Input)
+ _prompt.disabled = True
+ ask_w = AskUserWidget(questions)
+ await container.mount(ask_w)
+ _schedule_scroll()
+ self.call_after_refresh(ask_w.focus_active)
+ result = await self._wait_for_ask_user(ask_w)
+ try:
+ await ask_w.remove()
+ except Exception:
+ pass
+ _prompt.disabled = False
+ from langgraph.types import Command # type: ignore[import-untyped]
+ _stream_input = Command(resume=result)
+ _hitl_resuming = True
+ break # re-enter outer HITL loop
+
elif event_type == "interrupt":
action_reqs = event.get("action_requests", [])
n = len(action_reqs) or 1
@@ -885,12 +961,16 @@ def run_textual_interactive(
continue
# Interactive TUI: mount approval widget
+ # Disable main prompt so it can't steal focus
+ _prompt = self.query_one("#prompt", Input)
+ _prompt.disabled = True
from .widgets.approval_widget import ApprovalWidget
approval_w = ApprovalWidget(action_reqs)
await container.mount(approval_w)
_schedule_scroll()
decided_event = await self._wait_for_approval(approval_w)
await approval_w.remove()
+ _prompt.disabled = False
if decided_event and decided_event.decisions is not None:
if decided_event.auto_approve_session:
self._hitl_auto_approve = True
@@ -1019,8 +1099,8 @@ def run_textual_interactive(
),
)
- # HITL: if interrupt was handled, loop back to resume stream
- if state.pending_interrupt is None:
+ # HITL / ask_user: if interrupt was handled, loop back to resume stream
+ if state.pending_interrupt is None and state.pending_ask_user is None:
break # normal completion or rejection — exit HITL loop
# Otherwise _stream_input was set to Command(resume=...)
# by the interrupt handler above; loop continues.
@@ -1166,6 +1246,7 @@ def run_textual_interactive(
text = event.value.strip()
prompt = self.query_one("#prompt", Input)
prompt.value = ""
+
if not text:
return
@@ -1225,6 +1306,17 @@ def run_textual_interactive(
def action_cancel_queued(self) -> None:
"""Cancel the last queued message on Esc."""
+ # Cancel ask_user if active (widget handles Escape internally,
+ # but this is a safety fallback)
+ if self._ask_user_future and not self._ask_user_future.done():
+ try:
+ from .widgets.ask_user_widget import AskUserWidget
+ ask_w = self.query_one(AskUserWidget)
+ ask_w.action_cancel()
+ except Exception:
+ # Force-resolve the future
+ self._ask_user_future.set_result({"type": "cancelled"})
+ return
# Delegate to ApprovalWidget or ThreadPickerWidget if focused
focused = self.focused
if focused is not None:
@@ -1242,14 +1334,18 @@ def run_textual_interactive(
def action_edit_queued(self) -> None:
"""Pop the last queued message back into input for editing."""
- # Skip if an ApprovalWidget or ThreadPickerWidget has focus
+ # Skip if an ApprovalWidget, AskUserWidget, or ThreadPickerWidget has focus
focused = self.focused
if focused is not None:
from .widgets.approval_widget import ApprovalWidget
+ from .widgets.ask_user_widget import AskUserWidget
from .widgets.thread_selector import ThreadPickerWidget
if isinstance(focused, ApprovalWidget):
focused.action_move_up()
return
+ if isinstance(focused, AskUserWidget):
+ focused.action_move_up()
+ return
if isinstance(focused, ThreadPickerWidget):
focused.action_move_up()
return
@@ -1262,14 +1358,18 @@ def run_textual_interactive(
self._render_queue_indicator()
def action_down_delegate(self) -> None:
- """Delegate down key to focused ApprovalWidget or ThreadPickerWidget."""
+ """Delegate down key to focused ApprovalWidget, AskUserWidget, or ThreadPickerWidget."""
focused = self.focused
if focused is not None:
from .widgets.approval_widget import ApprovalWidget
+ from .widgets.ask_user_widget import AskUserWidget
from .widgets.thread_selector import ThreadPickerWidget
if isinstance(focused, ApprovalWidget):
focused.action_move_down()
return
+ if isinstance(focused, AskUserWidget):
+ focused.action_move_down()
+ return
if isinstance(focused, ThreadPickerWidget):
focused.action_move_down()
return
diff --git a/EvoScientist/cli/tui_runtime.py b/EvoScientist/cli/tui_runtime.py
index 3a9b1d2..244d24a 100644
--- a/EvoScientist/cli/tui_runtime.py
+++ b/EvoScientist/cli/tui_runtime.py
@@ -67,6 +67,7 @@ def run_streaming(
on_file_write: Callable[[str], None] | None = None,
metadata: dict | None = None,
hitl_prompt_fn: Callable[[list], list[dict] | None] | None = None,
+ ask_user_prompt_fn: Callable[[dict], dict] | None = None,
) -> str:
"""Run streaming with the selected backend."""
backend = get_backend(ui_backend, warn_fallback=True)
@@ -82,6 +83,7 @@ def run_streaming(
on_file_write=on_file_write,
metadata=metadata,
hitl_prompt_fn=hitl_prompt_fn,
+ ask_user_prompt_fn=ask_user_prompt_fn,
)
except RuntimeError:
requested = normalize_ui_backend(ui_backend)
@@ -100,5 +102,6 @@ def run_streaming(
on_file_write=on_file_write,
metadata=metadata,
hitl_prompt_fn=hitl_prompt_fn,
+ ask_user_prompt_fn=ask_user_prompt_fn,
)
raise
diff --git a/EvoScientist/cli/widgets/__init__.py b/EvoScientist/cli/widgets/__init__.py
index 668ed75..a0e3f7f 100644
--- a/EvoScientist/cli/widgets/__init__.py
+++ b/EvoScientist/cli/widgets/__init__.py
@@ -11,6 +11,7 @@ from .user_message import UserMessage
from .system_message import SystemMessage
from .usage_widget import UsageWidget
from .approval_widget import ApprovalWidget
+from .ask_user_widget import AskUserWidget
from .thread_selector import ThreadPickerWidget
__all__ = [
@@ -25,5 +26,6 @@ __all__ = [
"SystemMessage",
"UsageWidget",
"ApprovalWidget",
+ "AskUserWidget",
"ThreadPickerWidget",
]
diff --git a/EvoScientist/cli/widgets/ask_user_widget.py b/EvoScientist/cli/widgets/ask_user_widget.py
new file mode 100644
index 0000000..a287b2d
--- /dev/null
+++ b/EvoScientist/cli/widgets/ask_user_widget.py
@@ -0,0 +1,375 @@
+"""Interactive ask_user widget for Textual TUI.
+
+Shows one question at a time with a progress indicator. Keyboard-driven
+like ApprovalWidget: all bindings on the top-level widget, compact layout.
+
+The widget is self-contained: the main ``#prompt`` Input should be
+**disabled** while this widget is mounted so it cannot steal focus.
+"""
+
+from __future__ import annotations
+
+import asyncio
+import logging
+from typing import TYPE_CHECKING, Any, ClassVar, Literal
+
+from rich.markup import escape as escape_markup
+from textual.binding import Binding, BindingType
+from textual.message import Message
+from textual.widget import Widget
+from textual.widgets import Input, Static
+
+if TYPE_CHECKING:
+ from textual import events
+ from textual.app import ComposeResult
+
+logger = logging.getLogger(__name__)
+
+OTHER_CHOICE_LABEL = "Other (type your answer)"
+_CURSOR = "▸"
+
+
+class AskUserWidget(Widget):
+ """Interactive widget for asking the user questions one at a time.
+
+ Extends Widget (not Container) to avoid built-in scroll behavior
+ that captures arrow keys. Same pattern as ApprovalWidget.
+ """
+
+ can_focus = True
+ can_focus_children = True # needed for text Input
+
+ BINDINGS: ClassVar[list[BindingType]] = [
+ Binding("up", "move_up", "Up", show=False),
+ Binding("k", "move_up", "Up", show=False),
+ Binding("down", "move_down", "Down", show=False),
+ Binding("j", "move_down", "Down", show=False),
+ Binding("enter", "confirm", "Confirm", show=False),
+ Binding("escape", "cancel", "Cancel", show=False),
+ ]
+
+ DEFAULT_CSS = """
+ AskUserWidget {
+ height: auto;
+ margin: 1 0;
+ padding: 0 1;
+ background: $surface;
+ border: solid $success;
+ }
+ AskUserWidget .ask-title {
+ height: 1;
+ text-style: bold;
+ color: $success;
+ }
+ AskUserWidget .ask-question-text {
+ height: auto;
+ margin: 0 0 0 1;
+ text-style: bold;
+ }
+ AskUserWidget .ask-choice {
+ height: 1;
+ padding: 0 2;
+ margin: 0;
+ }
+ AskUserWidget .ask-choice-selected {
+ background: $primary;
+ text-style: bold;
+ }
+ AskUserWidget .ask-text-input {
+ height: auto;
+ margin: 0 2;
+ }
+ AskUserWidget .ask-help {
+ height: 1;
+ color: $text-muted;
+ text-style: italic;
+ margin: 0;
+ }
+ """
+
+ class Answered(Message):
+ """Posted when the user submits all answers."""
+
+ def __init__(self, answers: list[str]) -> None:
+ super().__init__()
+ self.answers = answers
+
+ class Cancelled(Message):
+ """Posted when the user cancels the ask_user prompt."""
+
+ def __init__(self) -> None:
+ super().__init__()
+
+ def __init__(
+ self,
+ questions: list[dict],
+ id: str | None = None,
+ **kwargs: Any,
+ ) -> None:
+ super().__init__(id=id or "ask-user-widget", **kwargs)
+ self._questions = questions
+ self._answers: list[str] = []
+ self._current_index = 0
+ self._future: asyncio.Future | None = None
+ self._submitted = False
+
+ # Current question state
+ self._q_type: Literal["text", "multiple_choice"] = "text"
+ self._choices: list[dict] = []
+ self._required: bool = True
+ self._selected_choice: int = 0
+ self._is_other: bool = False
+
+ # Widgets (composed once, updated per question)
+ self._title_w: Static | None = None
+ self._question_w: Static | None = None
+ self._choice_widgets: list[Static] = []
+ self._text_input: Input | None = None
+ self._other_input: Input | None = None
+ self._help_w: Static | None = None
+
+ def set_future(self, future: asyncio.Future) -> None:
+ """Set the future to resolve when user answers."""
+ self._future = future
+
+ def compose(self) -> ComposeResult:
+ total = len(self._questions)
+ if total == 1:
+ title = ">>> Quick check-in from EvoScientist <<<"
+ else:
+ title = ">>> Question 1/{} — Quick check-in from EvoScientist <<<".format(total)
+
+ self._title_w = Static(title, classes="ask-title")
+ yield self._title_w
+
+ self._question_w = Static("", classes="ask-question-text")
+ yield self._question_w
+
+ # Pre-create max choice slots (choices + Other = up to ~10)
+ # Unused slots stay hidden.
+ self._choice_widgets = []
+ for _ in range(12):
+ cw = Static("", classes="ask-choice")
+ cw.display = False
+ self._choice_widgets.append(cw)
+ yield cw
+
+ self._text_input = Input(
+ placeholder="Type your answer...",
+ classes="ask-text-input",
+ )
+ self._text_input.display = False
+ yield self._text_input
+
+ self._other_input = Input(
+ placeholder="Type your answer...",
+ classes="ask-text-input",
+ )
+ self._other_input.display = False
+ yield self._other_input
+
+ self._help_w = Static("", classes="ask-help")
+ yield self._help_w
+
+ async def on_mount(self) -> None:
+ self._show_question(0)
+
+ def focus_active(self) -> None:
+ """Focus the appropriate element for the current question."""
+ if self._q_type == "text":
+ if self._text_input:
+ self._text_input.focus()
+ elif self._is_other:
+ if self._other_input:
+ self._other_input.focus()
+ else:
+ self.focus()
+
+ # ------------------------------------------------------------------
+ # Show a question
+ # ------------------------------------------------------------------
+
+ def _show_question(self, index: int) -> None:
+ """Populate widgets for question at *index*."""
+ q = self._questions[index]
+ q_text = q.get("question", "")
+ q_type = q.get("type", "text")
+ self._choices = q.get("choices", [])
+ self._required = q.get("required", True)
+ self._q_type = "multiple_choice" if q_type == "multiple_choice" else "text"
+ self._selected_choice = 0
+ self._is_other = False
+
+ # Title
+ total = len(self._questions)
+ if self._title_w:
+ if total == 1:
+ self._title_w.update(">>> Quick check-in from EvoScientist <<<")
+ else:
+ self._title_w.update(
+ f">>> Question {index + 1}/{total}"
+ " — Quick check-in from EvoScientist <<<"
+ )
+
+ # Question text
+ suffix = " [dim](required)[/dim]" if self._required else " [dim](optional)[/dim]"
+ if self._question_w:
+ self._question_w.update(
+ f"[bold]{index + 1}. {escape_markup(q_text)}[/bold]{suffix}"
+ )
+
+ # Reset all choice slots
+ for cw in self._choice_widgets:
+ cw.display = False
+ cw.remove_class("ask-choice-selected")
+
+ if self._text_input:
+ self._text_input.display = False
+ self._text_input.value = ""
+ if self._other_input:
+ self._other_input.display = False
+ self._other_input.value = ""
+
+ if self._q_type == "multiple_choice" and self._choices:
+ # Show choice options + Other
+ for i, choice in enumerate(self._choices):
+ if i < len(self._choice_widgets):
+ label = escape_markup(choice.get("value", str(choice)))
+ cursor = f"{_CURSOR} " if i == 0 else " "
+ self._choice_widgets[i].update(f"{cursor}{label}")
+ self._choice_widgets[i].display = True
+ if i == 0:
+ self._choice_widgets[i].add_class("ask-choice-selected")
+
+ other_idx = len(self._choices)
+ if other_idx < len(self._choice_widgets):
+ self._choice_widgets[other_idx].update(f" {OTHER_CHOICE_LABEL}")
+ self._choice_widgets[other_idx].display = True
+
+ # Help text
+ if self._help_w:
+ self._help_w.update("↑/↓ select · Enter confirm · Esc cancel")
+ self.focus()
+ else:
+ # Text input
+ if self._text_input:
+ self._text_input.display = True
+ self._text_input.focus()
+ if self._help_w:
+ self._help_w.update("Enter confirm · Esc cancel")
+
+ # ------------------------------------------------------------------
+ # Key bindings
+ # ------------------------------------------------------------------
+
+ def action_move_up(self) -> None:
+ if self._q_type != "multiple_choice":
+ return
+ if self._is_other and self._other_input and self._other_input.has_focus:
+ # Jump back from Other input to choice list
+ self._is_other = False
+ if self._other_input:
+ self._other_input.display = False
+ self._selected_choice = len(self._choices) # stay on Other option
+ self._update_choices()
+ self.focus()
+ return
+ total_opts = len(self._choices) + 1 # choices + Other
+ self._selected_choice = (self._selected_choice - 1) % total_opts
+ self._update_choices()
+
+ def action_move_down(self) -> None:
+ if self._q_type != "multiple_choice":
+ return
+ total_opts = len(self._choices) + 1
+ self._selected_choice = (self._selected_choice + 1) % total_opts
+ self._update_choices()
+
+ def action_confirm(self) -> None:
+ if self._q_type == "multiple_choice":
+ is_other = self._selected_choice == len(self._choices)
+ if is_other and not self._is_other:
+ # Show Other text input
+ self._is_other = True
+ if self._other_input:
+ self._other_input.display = True
+ self._other_input.focus()
+ return
+ if is_other and self._is_other:
+ # Submit Other answer
+ answer = self._other_input.value if self._other_input else ""
+ if answer.strip() or not self._required:
+ self._advance(answer)
+ return
+ # Regular choice
+ if self._selected_choice < len(self._choices):
+ answer = self._choices[self._selected_choice].get("value", "")
+ self._advance(answer)
+ else:
+ # Text question — Enter on widget (not Input) acts as confirm
+ answer = self._text_input.value if self._text_input else ""
+ if answer.strip() or not self._required:
+ self._advance(answer)
+
+ def on_input_submitted(self, event: Input.Submitted) -> None:
+ """Handle Enter in text Input widgets."""
+ event.stop()
+ if event.input is self._text_input:
+ answer = self._text_input.value if self._text_input else ""
+ if answer.strip() or not self._required:
+ self._advance(answer)
+ elif event.input is self._other_input:
+ answer = self._other_input.value if self._other_input else ""
+ if answer.strip() or not self._required:
+ self._advance(answer)
+
+ def action_cancel(self) -> None:
+ if self._submitted:
+ return
+ self._submitted = True
+ if self._future and not self._future.done():
+ self._future.set_result({"type": "cancelled"})
+ self.post_message(self.Cancelled())
+
+ # ------------------------------------------------------------------
+ # Helpers
+ # ------------------------------------------------------------------
+
+ def _update_choices(self) -> None:
+ """Update choice display to reflect current selection."""
+ total_opts = len(self._choices) + 1
+ for i in range(total_opts):
+ if i >= len(self._choice_widgets):
+ break
+ if i < len(self._choices):
+ label = escape_markup(self._choices[i].get("value", ""))
+ else:
+ label = OTHER_CHOICE_LABEL
+ cursor = f"{_CURSOR} " if i == self._selected_choice else " "
+ self._choice_widgets[i].update(f"{cursor}{label}")
+ self._choice_widgets[i].remove_class("ask-choice-selected")
+ if i == self._selected_choice:
+ self._choice_widgets[i].add_class("ask-choice-selected")
+
+ def _advance(self, answer: str) -> None:
+ """Record answer, show next question or submit."""
+ self._answers.append(answer)
+ self._current_index += 1
+
+ if self._current_index >= len(self._questions):
+ self._submit()
+ return
+
+ self._show_question(self._current_index)
+
+ def _submit(self) -> None:
+ if self._submitted:
+ return
+ self._submitted = True
+ if self._future and not self._future.done():
+ self._future.set_result({"type": "answered", "answers": self._answers})
+ self.post_message(self.Answered(self._answers))
+
+ def on_blur(self, event: events.Blur) -> None:
+ """Prevent blur from propagating and dismissing the widget."""
+ event.stop()
diff --git a/EvoScientist/config/settings.py b/EvoScientist/config/settings.py
index 07d1854..d204968 100644
--- a/EvoScientist/config/settings.py
+++ b/EvoScientist/config/settings.py
@@ -181,6 +181,9 @@ class EvoScientistConfig:
auto_approve: bool = False # Auto-approve all tool executions without prompting
shell_allow_list: str = "" # Comma-separated shell command prefixes to auto-approve
+ # Agent features
+ enable_ask_user: bool = True # Enable ask_user tool for agent-initiated questions
+
# DM access control policy
dm_policy: str = "allowlist"
diff --git a/EvoScientist/middleware/__init__.py b/EvoScientist/middleware/__init__.py
index 93e4dbb..ddbcc38 100644
--- a/EvoScientist/middleware/__init__.py
+++ b/EvoScientist/middleware/__init__.py
@@ -4,6 +4,13 @@ Re-exports middleware classes and factory functions so that existing
``from EvoScientist.middleware import X`` imports continue to work.
"""
+from .ask_user import (
+ AskUserMiddleware,
+ AskUserRequest,
+ AskUserWidgetResult,
+ Choice,
+ Question,
+)
from .memory import (
EvoMemoryMiddleware,
EvoMemoryState,
@@ -13,9 +20,14 @@ from .memory import (
from .tool_error_handler import ToolErrorHandlerMiddleware
__all__ = [
+ "AskUserMiddleware",
+ "AskUserRequest",
+ "AskUserWidgetResult",
+ "Choice",
"EvoMemoryMiddleware",
"EvoMemoryState",
"ExtractedMemory",
+ "Question",
"ToolErrorHandlerMiddleware",
"create_memory_middleware",
]
diff --git a/EvoScientist/middleware/ask_user.py b/EvoScientist/middleware/ask_user.py
new file mode 100644
index 0000000..0b99924
--- /dev/null
+++ b/EvoScientist/middleware/ask_user.py
@@ -0,0 +1,405 @@
+"""Ask user middleware for agent-initiated interactive questions.
+
+Enables the agent to proactively ask the user for clarification during
+research workflows. The middleware registers an ``ask_user`` tool that
+uses LangGraph ``interrupt()`` to pause the graph; the UI layer collects
+answers and resumes with ``Command(resume={...})``.
+
+Ported from upstream DeepAgents ``AskUserMiddleware`` with a system
+prompt tailored for scientific research contexts.
+"""
+
+from __future__ import annotations
+
+import json
+import logging
+from typing import TYPE_CHECKING, Annotated, Any, Literal, cast
+
+if TYPE_CHECKING:
+ from collections.abc import Awaitable, Callable
+
+from typing import NotRequired
+
+from langchain.agents.middleware.types import (
+ AgentMiddleware,
+ ContextT,
+ ModelRequest,
+ ModelResponse,
+ ResponseT,
+)
+from langchain.tools import InjectedToolCallId
+from langchain_core.messages import AIMessage, SystemMessage, ToolMessage
+from langchain_core.tools import tool
+from langgraph.types import Command, interrupt
+from pydantic import BeforeValidator, Field
+from typing_extensions import TypedDict
+
+logger = logging.getLogger(__name__)
+
+
+# ---------------------------------------------------------------------------
+# Data types
+# ---------------------------------------------------------------------------
+
+
+class Choice(TypedDict):
+ """A single choice option for a multiple choice question."""
+
+ value: Annotated[str, Field(description="The display label for this choice.")]
+
+
+class Question(TypedDict):
+ """A question to ask the user."""
+
+ question: Annotated[str, Field(description="The question text to display.")]
+
+ type: Annotated[
+ Literal["text", "multiple_choice"],
+ Field(
+ description=(
+ "Question type. 'text' for free-form input, 'multiple_choice' for "
+ "predefined options."
+ )
+ ),
+ ]
+
+ choices: NotRequired[
+ Annotated[
+ list[Choice],
+ Field(
+ description=(
+ "Options for multiple_choice questions. An 'Other' free-form "
+ "option is always appended automatically."
+ )
+ ),
+ ]
+ ]
+
+ required: NotRequired[
+ Annotated[
+ bool,
+ Field(
+ description="Whether the user must answer. Defaults to true if omitted."
+ ),
+ ]
+ ]
+
+
+def _coerce_questions_list(v: Any) -> Any:
+ """Accept a JSON string and parse it to a list.
+
+ LLMs sometimes serialize the ``questions`` argument as a JSON string
+ instead of a native list. This before-validator transparently handles
+ that case so the tool invocation succeeds on the first try.
+ """
+ if isinstance(v, str):
+ try:
+ parsed = json.loads(v)
+ if isinstance(parsed, list):
+ return parsed
+ except (json.JSONDecodeError, TypeError):
+ pass
+ return v
+
+
+# Annotated type for the tool parameter — schema stays ``list[Question]``
+# for the LLM, but runtime accepts JSON strings too.
+QuestionsList = Annotated[list[Question], BeforeValidator(_coerce_questions_list)]
+
+
+class AskUserRequest(TypedDict):
+ """Request payload sent via interrupt when asking the user questions."""
+
+ type: Literal["ask_user"]
+
+ questions: list[Question]
+
+ tool_call_id: str
+
+
+class AskUserAnswered(TypedDict):
+ """Widget result when the user submits answers."""
+
+ type: Literal["answered"]
+ """Discriminator tag, always ``'answered'``."""
+
+ answers: list[str]
+ """User-provided answers, one per question."""
+
+
+class AskUserCancelled(TypedDict):
+ """Widget result when the user cancels the prompt."""
+
+ type: Literal["cancelled"]
+ """Discriminator tag, always ``'cancelled'``."""
+
+
+# Discriminated union for the ask_user widget Future result.
+AskUserWidgetResult = AskUserAnswered | AskUserCancelled
+
+
+# ---------------------------------------------------------------------------
+# Prompts (research-tailored)
+# ---------------------------------------------------------------------------
+
+
+ASK_USER_TOOL_DESCRIPTION = """\
+Ask the user one or more questions when you need clarification or specific input
+before proceeding with a research task.
+
+Each question can be:
+- "text": Free-form text response
+- "multiple_choice": User selects from predefined options (an "Other" option is always appended)
+
+Use when: dataset/model/framework selection is ambiguous, experiment parameters unclear,
+research scope needs clarification, or significant plan needs confirmation.
+Do NOT use for trivial decisions or questions answerable from context/memory/web search."""
+
+
+ASK_USER_SYSTEM_PROMPT = """\
+## `ask_user` — Interactive Clarification Tool
+
+You have access to the `ask_user` tool to ask the user questions when you need
+information that cannot be determined from the conversation, loaded skills,
+or available tools.
+
+### When to use `ask_user`:
+- **Dataset or benchmark selection**: "Which dataset should I use: CIFAR-10, ImageNet, or a custom dataset?"
+- **Experiment parameters**: "What GPU memory budget should I target? What batch size range is acceptable?"
+- **Research scope**: "Should I focus on accuracy improvements or inference speed?"
+- **Methodology choice**: "For the baseline comparison, should I reimplement from the paper or use the official repo?"
+- **Paper or report preferences**: "Which venue format should I target: NeurIPS, ICML, or ICLR?"
+- **Ambiguous instructions**: When the user's request has multiple valid interpretations
+- **Resource constraints**: When the approach depends on available compute, time, or data
+
+### When NOT to use `ask_user`:
+- Simple yes/no decisions — proceed with your best judgment
+- Information already provided in the conversation or memory
+- Trivial choices that don't meaningfully affect outcomes
+- Questions you can answer by searching with `tavily_search` or reading files
+- During sub-agent execution (only the main agent should ask the user)
+
+### Guidelines:
+- Be concise and specific — avoid vague, open-ended questions
+- Use `multiple_choice` when there are 2-5 clear options
+- Use `text` for open-ended input (preferred language, custom parameters, etc.)
+- Group related questions into a **single** `ask_user` call (max 5 questions)
+- Never ask more than once per decision point — respect the user's time
+- After receiving answers, summarize what you understood before proceeding"""
+
+
+# ---------------------------------------------------------------------------
+# Validation & parsing
+# ---------------------------------------------------------------------------
+
+
+def _validate_questions(questions: list[Question]) -> None:
+ """Validate ask_user question structure before interrupting.
+
+ Raises:
+ ValueError: If the questions list or an individual question is invalid.
+ """
+ if not questions:
+ msg = "ask_user requires at least one question"
+ raise ValueError(msg)
+
+ for q in questions:
+ question_text = q.get("question")
+ if not isinstance(question_text, str) or not question_text.strip():
+ msg = "ask_user questions must have non-empty 'question' text"
+ raise ValueError(msg)
+
+ question_type = q.get("type")
+ if question_type not in {"text", "multiple_choice"}:
+ msg = f"unsupported ask_user question type: {question_type!r}"
+ raise ValueError(msg)
+
+ if question_type == "multiple_choice" and not q.get("choices"):
+ msg = (
+ f"multiple_choice question "
+ f"{q.get('question')!r} requires a "
+ f"non-empty 'choices' list"
+ )
+ raise ValueError(msg)
+
+ if question_type == "text" and q.get("choices"):
+ msg = f"text question {q.get('question')!r} must not define 'choices'"
+ raise ValueError(msg)
+
+
+def _parse_answers(
+ response: object,
+ questions: list[Question],
+ tool_call_id: str,
+) -> Command[Any]:
+ """Parse an interrupt response into a ``Command`` with a ``ToolMessage``.
+
+ Supports explicit status signaling from the adapter:
+
+ - ``answered`` (default): consume provided ``answers``
+ - ``cancelled``: synthesize ``(cancelled)`` answers
+ - ``error``: synthesize ``(error: ...)`` answers
+
+ Malformed payloads are converted into explicit error answers instead of
+ silently defaulting to ``(no answer)``.
+ """
+ status: str = "answered"
+ error_text: str | None = None
+ answers: list[str]
+ if not isinstance(response, dict):
+ logger.error(
+ "ask_user received malformed resume payload "
+ "(expected dict, got %s); returning explicit error answers",
+ type(response).__name__,
+ )
+ answers = []
+ status = "error"
+ error_text = "invalid ask_user response payload"
+ else:
+ response_dict = cast("dict[str, Any]", response)
+ response_status = response_dict.get("status")
+ if isinstance(response_status, str):
+ status = response_status
+
+ if "answers" not in response_dict:
+ if status == "answered":
+ logger.error(
+ "ask_user received resume payload without 'answers'; "
+ "returning explicit error answers"
+ )
+ answers = []
+ status = "error"
+ error_text = "missing ask_user answers payload"
+ else:
+ answers = []
+ else:
+ raw_answers = response_dict["answers"]
+ if isinstance(raw_answers, list):
+ answers = [str(answer) for answer in raw_answers]
+ else:
+ logger.error(
+ "ask_user received non-list 'answers' payload (%s); "
+ "returning explicit error answers",
+ type(raw_answers).__name__,
+ )
+ answers = []
+ status = "error"
+ error_text = "invalid ask_user answers payload"
+
+ if status == "error":
+ response_error = response_dict.get("error")
+ if isinstance(response_error, str) and response_error:
+ error_text = response_error
+ elif status == "cancelled":
+ answers = ["(cancelled)" for _ in questions]
+ elif status == "answered":
+ if len(answers) != len(questions):
+ logger.warning(
+ "ask_user answer count mismatch: expected %d, got %d",
+ len(questions),
+ len(answers),
+ )
+ else:
+ logger.error(
+ "ask_user received unknown status %r; returning explicit error answers",
+ status,
+ )
+ answers = []
+ status = "error"
+ error_text = "invalid ask_user response status"
+
+ if status == "error":
+ detail = error_text or "ask_user interaction failed"
+ answers = [f"(error: {detail})" for _ in questions]
+
+ formatted_answers = []
+ for i, q in enumerate(questions):
+ answer = answers[i] if i < len(answers) else "(no answer)"
+ formatted_answers.append(f"Q: {q['question']}\nA: {answer}")
+ result_text = "\n\n".join(formatted_answers)
+ return Command(
+ update={
+ "messages": [ToolMessage(result_text, tool_call_id=tool_call_id)],
+ }
+ )
+
+
+# ---------------------------------------------------------------------------
+# Middleware class
+# ---------------------------------------------------------------------------
+
+
+class AskUserMiddleware(AgentMiddleware[Any, ContextT, ResponseT]):
+ """Middleware that provides an ``ask_user`` tool for interactive questioning.
+
+ Adds an ``ask_user`` tool that allows the main agent to ask the user
+ questions during execution. Questions can be free-form text or multiple
+ choice. The tool uses LangGraph ``interrupt()`` to pause execution and
+ wait for user input.
+ """
+
+ def __init__(
+ self,
+ *,
+ system_prompt: str = ASK_USER_SYSTEM_PROMPT,
+ tool_description: str = ASK_USER_TOOL_DESCRIPTION,
+ ) -> None:
+ super().__init__()
+ self.system_prompt = system_prompt
+ self.tool_description = tool_description
+
+ @tool(description=self.tool_description)
+ def _ask_user(
+ questions: QuestionsList,
+ tool_call_id: Annotated[str, InjectedToolCallId],
+ ) -> Command[Any]:
+ """Ask the user one or more questions."""
+ _validate_questions(questions)
+ ask_request = AskUserRequest(
+ type="ask_user",
+ questions=questions,
+ tool_call_id=tool_call_id,
+ )
+ response = interrupt(ask_request)
+ return _parse_answers(response, questions, tool_call_id)
+
+ _ask_user.name = "ask_user"
+ self.tools = [_ask_user]
+
+ def wrap_model_call(
+ self,
+ request: ModelRequest[ContextT],
+ handler: Callable[[ModelRequest[ContextT]], ModelResponse[ResponseT]],
+ ) -> ModelResponse[ResponseT] | AIMessage:
+ """Inject the ask_user system prompt."""
+ if request.system_message is not None:
+ new_system_content = [
+ *request.system_message.content_blocks,
+ {"type": "text", "text": f"\n\n{self.system_prompt}"},
+ ]
+ else:
+ new_system_content = [{"type": "text", "text": self.system_prompt}]
+ new_system_message = SystemMessage(
+ content=cast("list[str | dict[str, str]]", new_system_content)
+ )
+ return handler(request.override(system_message=new_system_message))
+
+ async def awrap_model_call(
+ self,
+ request: ModelRequest[ContextT],
+ handler: Callable[
+ [ModelRequest[ContextT]], Awaitable[ModelResponse[ResponseT]]
+ ],
+ ) -> ModelResponse[ResponseT] | AIMessage:
+ """Inject the ask_user system prompt (async)."""
+ if request.system_message is not None:
+ new_system_content = [
+ *request.system_message.content_blocks,
+ {"type": "text", "text": f"\n\n{self.system_prompt}"},
+ ]
+ else:
+ new_system_content = [{"type": "text", "text": self.system_prompt}]
+ new_system_message = SystemMessage(
+ content=cast("list[str | dict[str, str]]", new_system_content)
+ )
+ return await handler(request.override(system_message=new_system_message))
diff --git a/EvoScientist/middleware/tool_error_handler.py b/EvoScientist/middleware/tool_error_handler.py
index 63d90df..ac052a8 100644
--- a/EvoScientist/middleware/tool_error_handler.py
+++ b/EvoScientist/middleware/tool_error_handler.py
@@ -21,6 +21,12 @@ from langchain.agents.middleware.types import AgentMiddleware
from langchain_core.messages import ToolMessage
from langgraph.types import Command
+# GraphInterrupt must propagate — never catch it as a tool error.
+try:
+ from langgraph.errors import GraphInterrupt as _GraphInterrupt
+except ImportError: # older langgraph versions
+ _GraphInterrupt = None # type: ignore[assignment,misc]
+
if TYPE_CHECKING:
from langchain.agents.middleware.types import ToolCallRequest
@@ -39,7 +45,9 @@ class ToolErrorHandlerMiddleware(AgentMiddleware):
) -> ToolMessage | Command[Any]:
try:
return handler(request)
- except Exception:
+ except Exception as exc:
+ if _GraphInterrupt is not None and isinstance(exc, _GraphInterrupt):
+ raise
return _build_error_message(request)
async def awrap_tool_call(
@@ -49,7 +57,9 @@ class ToolErrorHandlerMiddleware(AgentMiddleware):
) -> ToolMessage | Command[Any]:
try:
return await handler(request)
- except Exception:
+ except Exception as exc:
+ if _GraphInterrupt is not None and isinstance(exc, _GraphInterrupt):
+ raise
return _build_error_message(request)
diff --git a/EvoScientist/stream/display.py b/EvoScientist/stream/display.py
index d2d58c2..569fb7c 100644
--- a/EvoScientist/stream/display.py
+++ b/EvoScientist/stream/display.py
@@ -864,6 +864,69 @@ def _get_event_loop() -> asyncio.AbstractEventLoop:
loop = _create_event_loop()
return loop
+def _resolve_ask_user_prompt(ask_user_data: dict) -> dict:
+ """Interactive console Q&A for ask_user events.
+
+ Presents questions via ``prompt_toolkit.prompt()`` (not ``input()``)
+ for proper CJK IME support and styled prompts without cursor drift.
+ """
+ from prompt_toolkit import prompt as pt_prompt # type: ignore[import-untyped]
+ from prompt_toolkit.formatted_text import HTML # type: ignore[import-untyped]
+
+ questions = ask_user_data.get("questions", [])
+ if not questions:
+ return {"answers": [], "status": "answered"}
+
+ console.print()
+ console.print(Panel(
+ Text("Quick check-in from EvoScientist", style="bold"),
+ border_style="cyan",
+ padding=(0, 1),
+ ))
+ console.print()
+
+ answers: list[str] = []
+ try:
+ for i, q in enumerate(questions):
+ q_text = q.get("question", "")
+ q_type = q.get("type", "text")
+ required = q.get("required", True)
+ tag = " [dim](optional)[/dim]" if not required else ""
+ console.print(f" [bold]{i + 1}. {q_text}[/bold]{tag}")
+
+ 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)
+ console.print(Text(f" {letter}. {label}", style="dim"))
+ other_letter = chr(ord("A") + len(choices))
+ console.print(Text(f" {other_letter}. Other (type your answer)", style="dim"))
+
+ letters = "/".join(chr(ord("A") + k) for k in range(len(choices) + 1))
+ raw = pt_prompt(HTML(f" ")).strip()
+ if raw.upper() == other_letter:
+ raw = pt_prompt(HTML(" ")).strip()
+ answers.append(raw)
+ 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:
+ raw = pt_prompt(HTML(" ")).strip()
+ answers.append(raw)
+ console.print()
+ except (EOFError, KeyboardInterrupt):
+ console.print("[dim] Cancelled.[/dim]")
+ return {"status": "cancelled"}
+
+ return {"answers": answers, "status": "answered"}
+
+
def _run_streaming(
agent: Any,
message: Any,
@@ -875,6 +938,7 @@ def _run_streaming(
on_file_write: Callable[[str], None] | None = None,
metadata: dict | None = None,
hitl_prompt_fn: Callable[[list], list[dict] | None] | None = None,
+ ask_user_prompt_fn: Callable[[dict], dict] | None = None,
*,
_state: StreamState | None = None,
_hitl_depth: int = 0,
@@ -1015,9 +1079,9 @@ def _run_streaming(
except asyncio.CancelledError:
pass
# Render clean final frame before Live exits (no spinners, expanded tools)
- if state.pending_interrupt is not None:
+ if state.pending_interrupt is not None or state.pending_ask_user is not None:
# Interrupted: render current state (not final) so it
- # looks continuous when approval prompt appears.
+ # looks continuous when prompt appears.
final_display = create_streaming_display(
**state.get_display_args(),
show_thinking=show_thinking,
@@ -1050,6 +1114,31 @@ def _run_streaming(
if len(state.thinking_text) >= _MIN_THINKING_LEN:
on_thinking(state.thinking_text.rstrip())
+ # ask_user: check before HITL (ask_user uses the same resume loop)
+ if state.pending_ask_user is not None and _hitl_depth < _MAX_HITL_ITERATIONS:
+ if ask_user_prompt_fn is not None:
+ result = ask_user_prompt_fn(state.pending_ask_user)
+ else:
+ result = _resolve_ask_user_prompt(state.pending_ask_user)
+ from langgraph.types import Command # type: ignore[import-untyped]
+ state.pending_ask_user = None
+ return _run_streaming(
+ agent=agent,
+ message=Command(resume=result),
+ thread_id=thread_id,
+ show_thinking=show_thinking,
+ interactive=interactive,
+ on_thinking=on_thinking,
+ on_todo=on_todo,
+ on_file_write=on_file_write,
+ metadata=metadata,
+ hitl_prompt_fn=hitl_prompt_fn,
+ ask_user_prompt_fn=ask_user_prompt_fn,
+ _state=state,
+ _hitl_depth=_hitl_depth + 1,
+ _media_sent=_media_sent,
+ )
+
# HITL: check for pending interrupt and handle approval
if state.pending_interrupt is not None and _hitl_depth < _MAX_HITL_ITERATIONS:
decisions = _resolve_hitl_approval(
@@ -1069,6 +1158,7 @@ def _run_streaming(
on_file_write=on_file_write,
metadata=metadata,
hitl_prompt_fn=hitl_prompt_fn,
+ ask_user_prompt_fn=ask_user_prompt_fn,
_state=state,
_hitl_depth=_hitl_depth + 1,
_media_sent=_media_sent,
diff --git a/EvoScientist/stream/emitter.py b/EvoScientist/stream/emitter.py
index d5d6093..565bf13 100644
--- a/EvoScientist/stream/emitter.py
+++ b/EvoScientist/stream/emitter.py
@@ -109,6 +109,18 @@ class StreamEventEmitter:
"review_configs": review_configs or [],
})
+ @staticmethod
+ def ask_user_interrupt(
+ interrupt_id: str, questions: list, tool_call_id: str = "",
+ ) -> StreamEvent:
+ """Agent-initiated ask_user interrupt event."""
+ return StreamEvent("ask_user", {
+ "type": "ask_user",
+ "interrupt_id": interrupt_id,
+ "questions": questions,
+ "tool_call_id": tool_call_id,
+ })
+
@staticmethod
def summarization(content: str) -> StreamEvent:
"""Context summarization event."""
diff --git a/EvoScientist/stream/events.py b/EvoScientist/stream/events.py
index ada32e7..c1286c4 100644
--- a/EvoScientist/stream/events.py
+++ b/EvoScientist/stream/events.py
@@ -326,7 +326,7 @@ async def stream_agent_events(
else:
continue
- # Parse HITL interrupts from updates mode
+ # Parse HITL / ask_user interrupts from updates mode
if mode_str == "updates":
if isinstance(data, dict) and "__interrupt__" in data:
for interrupt_obj in data["__interrupt__"]:
@@ -334,6 +334,30 @@ async def stream_agent_events(
interrupt_value = interrupt_obj.get("value", {})
else:
interrupt_value = getattr(interrupt_obj, "value", {})
+
+ # Discriminate ask_user vs HITL interrupts
+ iv_type = (
+ interrupt_value.get("type")
+ if isinstance(interrupt_value, dict)
+ else getattr(interrupt_value, "type", None)
+ )
+ if iv_type == "ask_user":
+ questions = (
+ interrupt_value.get("questions", [])
+ if isinstance(interrupt_value, dict)
+ else getattr(interrupt_value, "questions", [])
+ )
+ tc_id = (
+ interrupt_value.get("tool_call_id", "")
+ if isinstance(interrupt_value, dict)
+ else getattr(interrupt_value, "tool_call_id", "")
+ )
+ ns_parts = interrupt_obj.get("ns", [""]) if isinstance(interrupt_obj, dict) else getattr(interrupt_obj, "ns", [""])
+ interrupt_id = str(ns_parts[0]) if ns_parts else "default"
+ yield emitter.ask_user_interrupt(interrupt_id, questions, tc_id).data
+ continue
+
+ # Standard HITL approval interrupt
if isinstance(interrupt_value, dict):
action_reqs = interrupt_value.get("action_requests", [])
review_cfgs = interrupt_value.get("review_configs", [])
diff --git a/EvoScientist/stream/state.py b/EvoScientist/stream/state.py
index 95203a6..69a38fe 100644
--- a/EvoScientist/stream/state.py
+++ b/EvoScientist/stream/state.py
@@ -98,6 +98,8 @@ class StreamState:
self.total_output_tokens = 0
# HITL interrupt tracking
self.pending_interrupt: dict | None = None
+ # ask_user interrupt tracking
+ self.pending_ask_user: dict | None = None
# Cached Markdown object for Rich CLI display (avoids O(n²) re-parsing)
self._cached_md_text: str = ""
self._cached_md: object | None = None
@@ -260,6 +262,9 @@ class StreamState:
elif event_type == "interrupt":
self.pending_interrupt = event
+ elif event_type == "ask_user":
+ self.pending_ask_user = event
+
elif event_type == "summarization":
self.summarization_text = event.get("content", "")
diff --git a/tests/test_ask_user.py b/tests/test_ask_user.py
new file mode 100644
index 0000000..8047b31
--- /dev/null
+++ b/tests/test_ask_user.py
@@ -0,0 +1,499 @@
+"""Tests for the ask_user middleware, stream events, state, and UI helpers."""
+
+import pytest
+from unittest.mock import patch
+
+
+# ---------------------------------------------------------------------------
+# Middleware data types
+# ---------------------------------------------------------------------------
+
+
+class TestDataTypes:
+ """Test Question, Choice, AskUserRequest construction."""
+
+ def test_choice_construction(self):
+ from EvoScientist.middleware.ask_user import Choice
+
+ choice: Choice = {"value": "CIFAR-10"}
+ assert choice["value"] == "CIFAR-10"
+
+ def test_question_text_construction(self):
+ from EvoScientist.middleware.ask_user import Question
+
+ q: Question = {"question": "Which dataset?", "type": "text"}
+ assert q["question"] == "Which dataset?"
+ assert q["type"] == "text"
+
+ def test_question_multiple_choice_construction(self):
+ from EvoScientist.middleware.ask_user import Question
+
+ q: Question = {
+ "question": "Which dataset?",
+ "type": "multiple_choice",
+ "choices": [{"value": "CIFAR-10"}, {"value": "ImageNet"}],
+ }
+ assert q["type"] == "multiple_choice"
+ assert len(q["choices"]) == 2
+
+ def test_ask_user_request_construction(self):
+ from EvoScientist.middleware.ask_user import AskUserRequest
+
+ req: AskUserRequest = {
+ "type": "ask_user",
+ "questions": [{"question": "test?", "type": "text"}],
+ "tool_call_id": "tc_123",
+ }
+ assert req["type"] == "ask_user"
+ assert req["tool_call_id"] == "tc_123"
+
+ def test_ask_user_answered_construction(self):
+ from EvoScientist.middleware.ask_user import AskUserAnswered
+
+ result: AskUserAnswered = {"type": "answered", "answers": ["CIFAR-10"]}
+ assert result["type"] == "answered"
+
+ def test_ask_user_cancelled_construction(self):
+ from EvoScientist.middleware.ask_user import AskUserCancelled
+
+ result: AskUserCancelled = {"type": "cancelled"}
+ assert result["type"] == "cancelled"
+
+
+# ---------------------------------------------------------------------------
+# BeforeValidator: _coerce_questions_list
+# ---------------------------------------------------------------------------
+
+
+class TestCoerceQuestionsList:
+ """Test _coerce_questions_list() — the BeforeValidator that handles
+ LLMs serializing the questions param as a JSON string instead of a list."""
+
+ def test_json_string_parsed_to_list(self):
+ from EvoScientist.middleware.ask_user import _coerce_questions_list
+
+ raw = '[{"question": "Which dataset?", "type": "text"}]'
+ result = _coerce_questions_list(raw)
+ assert isinstance(result, list)
+ assert len(result) == 1
+ assert result[0]["question"] == "Which dataset?"
+
+ def test_list_passthrough(self):
+ from EvoScientist.middleware.ask_user import _coerce_questions_list
+
+ data = [{"question": "Q?", "type": "text"}]
+ result = _coerce_questions_list(data)
+ assert result is data # same object, no copy
+
+ def test_non_list_json_string_passthrough(self):
+ from EvoScientist.middleware.ask_user import _coerce_questions_list
+
+ # JSON object string should NOT be parsed (not a list)
+ raw = '{"question": "Q?", "type": "text"}'
+ result = _coerce_questions_list(raw)
+ assert result == raw # returned as-is
+
+ def test_invalid_json_string_passthrough(self):
+ from EvoScientist.middleware.ask_user import _coerce_questions_list
+
+ raw = "not json at all"
+ result = _coerce_questions_list(raw)
+ assert result == raw
+
+ def test_empty_list_json_string(self):
+ from EvoScientist.middleware.ask_user import _coerce_questions_list
+
+ result = _coerce_questions_list("[]")
+ assert result == []
+
+ def test_non_string_non_list_passthrough(self):
+ from EvoScientist.middleware.ask_user import _coerce_questions_list
+
+ assert _coerce_questions_list(42) == 42
+ assert _coerce_questions_list(None) is None
+
+
+# ---------------------------------------------------------------------------
+# Validation
+# ---------------------------------------------------------------------------
+
+
+class TestValidateQuestions:
+ """Test _validate_questions()."""
+
+ def test_empty_list_raises(self):
+ from EvoScientist.middleware.ask_user import _validate_questions
+
+ with pytest.raises(ValueError, match="at least one question"):
+ _validate_questions([])
+
+ def test_missing_question_text_raises(self):
+ from EvoScientist.middleware.ask_user import _validate_questions
+
+ with pytest.raises(ValueError, match="non-empty 'question' text"):
+ _validate_questions([{"question": "", "type": "text"}])
+
+ def test_wrong_type_raises(self):
+ from EvoScientist.middleware.ask_user import _validate_questions
+
+ with pytest.raises(ValueError, match="unsupported"):
+ _validate_questions([{"question": "Q?", "type": "radio"}])
+
+ def test_multiple_choice_no_choices_raises(self):
+ from EvoScientist.middleware.ask_user import _validate_questions
+
+ with pytest.raises(ValueError, match="non-empty 'choices' list"):
+ _validate_questions(
+ [{"question": "Q?", "type": "multiple_choice"}]
+ )
+
+ def test_text_with_choices_raises(self):
+ from EvoScientist.middleware.ask_user import _validate_questions
+
+ with pytest.raises(ValueError, match="must not define 'choices'"):
+ _validate_questions(
+ [
+ {
+ "question": "Q?",
+ "type": "text",
+ "choices": [{"value": "A"}],
+ }
+ ]
+ )
+
+ def test_valid_text_question(self):
+ from EvoScientist.middleware.ask_user import _validate_questions
+
+ # Should not raise
+ _validate_questions([{"question": "What dataset?", "type": "text"}])
+
+ def test_valid_multiple_choice_question(self):
+ from EvoScientist.middleware.ask_user import _validate_questions
+
+ _validate_questions(
+ [
+ {
+ "question": "Which?",
+ "type": "multiple_choice",
+ "choices": [{"value": "A"}, {"value": "B"}],
+ }
+ ]
+ )
+
+
+# ---------------------------------------------------------------------------
+# _parse_answers
+# ---------------------------------------------------------------------------
+
+
+class TestParseAnswers:
+ """Test _parse_answers()."""
+
+ def test_answered_status(self):
+ from EvoScientist.middleware.ask_user import _parse_answers
+
+ questions = [{"question": "Q1?", "type": "text"}]
+ result = _parse_answers(
+ {"answers": ["my answer"], "status": "answered"},
+ questions,
+ "tc_1",
+ )
+ assert hasattr(result, "update")
+ msgs = result.update["messages"]
+ assert len(msgs) == 1
+ assert "Q: Q1?" in msgs[0].content
+ assert "A: my answer" in msgs[0].content
+
+ def test_cancelled_status(self):
+ from EvoScientist.middleware.ask_user import _parse_answers
+
+ questions = [{"question": "Q1?", "type": "text"}]
+ result = _parse_answers(
+ {"status": "cancelled"},
+ questions,
+ "tc_1",
+ )
+ msgs = result.update["messages"]
+ assert "(cancelled)" in msgs[0].content
+
+ def test_malformed_payload_non_dict(self):
+ from EvoScientist.middleware.ask_user import _parse_answers
+
+ questions = [{"question": "Q1?", "type": "text"}]
+ result = _parse_answers("not a dict", questions, "tc_1")
+ msgs = result.update["messages"]
+ assert "(error:" in msgs[0].content
+
+ def test_missing_answers_key(self):
+ from EvoScientist.middleware.ask_user import _parse_answers
+
+ questions = [{"question": "Q1?", "type": "text"}]
+ result = _parse_answers(
+ {"status": "answered"},
+ questions,
+ "tc_1",
+ )
+ msgs = result.update["messages"]
+ assert "(error:" in msgs[0].content
+
+ def test_unknown_status(self):
+ from EvoScientist.middleware.ask_user import _parse_answers
+
+ questions = [{"question": "Q1?", "type": "text"}]
+ result = _parse_answers(
+ {"answers": ["x"], "status": "unknown_status"},
+ questions,
+ "tc_1",
+ )
+ msgs = result.update["messages"]
+ assert "(error:" in msgs[0].content
+
+
+# ---------------------------------------------------------------------------
+# Middleware class
+# ---------------------------------------------------------------------------
+
+
+class TestAskUserMiddleware:
+ """Test AskUserMiddleware initialization and tool creation."""
+
+ def test_init_creates_tool(self):
+ from EvoScientist.middleware.ask_user import AskUserMiddleware
+
+ mw = AskUserMiddleware()
+ assert len(mw.tools) == 1
+ assert mw.tools[0].name == "ask_user"
+
+ def test_system_prompt_set(self):
+ from EvoScientist.middleware.ask_user import (
+ ASK_USER_SYSTEM_PROMPT,
+ AskUserMiddleware,
+ )
+
+ mw = AskUserMiddleware()
+ assert mw.system_prompt == ASK_USER_SYSTEM_PROMPT
+
+ def test_custom_prompt(self):
+ from EvoScientist.middleware.ask_user import AskUserMiddleware
+
+ mw = AskUserMiddleware(system_prompt="custom prompt")
+ assert mw.system_prompt == "custom prompt"
+
+
+# ---------------------------------------------------------------------------
+# Stream event emitter
+# ---------------------------------------------------------------------------
+
+
+class TestStreamEmitter:
+ """Test ask_user_interrupt event creation."""
+
+ def test_ask_user_interrupt_event_structure(self):
+ from EvoScientist.stream.emitter import StreamEventEmitter
+
+ emitter = StreamEventEmitter()
+ event = emitter.ask_user_interrupt(
+ interrupt_id="default",
+ questions=[{"question": "Q?", "type": "text"}],
+ tool_call_id="tc_1",
+ )
+ assert event.type == "ask_user"
+ assert event.data["type"] == "ask_user"
+ assert event.data["interrupt_id"] == "default"
+ assert event.data["tool_call_id"] == "tc_1"
+ assert len(event.data["questions"]) == 1
+
+ def test_ask_user_interrupt_default_tool_call_id(self):
+ from EvoScientist.stream.emitter import StreamEventEmitter
+
+ emitter = StreamEventEmitter()
+ event = emitter.ask_user_interrupt("ns1", [])
+ assert event.data["tool_call_id"] == ""
+
+
+# ---------------------------------------------------------------------------
+# Stream state
+# ---------------------------------------------------------------------------
+
+
+class TestStreamState:
+ """Test StreamState handling of ask_user events."""
+
+ def test_pending_ask_user_starts_none(self):
+ from EvoScientist.stream.state import StreamState
+
+ state = StreamState()
+ assert state.pending_ask_user is None
+
+ def test_handle_ask_user_sets_pending(self):
+ from EvoScientist.stream.state import StreamState
+
+ state = StreamState()
+ event = {
+ "type": "ask_user",
+ "interrupt_id": "default",
+ "questions": [{"question": "Q?", "type": "text"}],
+ "tool_call_id": "tc_1",
+ }
+ result = state.handle_event(event)
+ assert result == "ask_user"
+ assert state.pending_ask_user is not None
+ assert state.pending_ask_user["tool_call_id"] == "tc_1"
+
+ def test_ask_user_does_not_affect_pending_interrupt(self):
+ from EvoScientist.stream.state import StreamState
+
+ state = StreamState()
+ event = {
+ "type": "ask_user",
+ "interrupt_id": "default",
+ "questions": [],
+ "tool_call_id": "tc_1",
+ }
+ state.handle_event(event)
+ assert state.pending_interrupt is None
+ assert state.pending_ask_user is not None
+
+
+# ---------------------------------------------------------------------------
+# Config
+# ---------------------------------------------------------------------------
+
+
+class TestConfig:
+ """Test enable_ask_user config field."""
+
+ def test_default_is_true(self):
+ from EvoScientist.config.settings import EvoScientistConfig
+
+ cfg = EvoScientistConfig()
+ assert cfg.enable_ask_user is True
+
+ def test_set_to_false(self):
+ from EvoScientist.config.settings import EvoScientistConfig
+
+ cfg = EvoScientistConfig(enable_ask_user=False)
+ assert cfg.enable_ask_user is False
+
+
+# ---------------------------------------------------------------------------
+# Rich CLI prompt (mocking input)
+# ---------------------------------------------------------------------------
+
+
+class TestRichCLIPrompt:
+ """Test _resolve_ask_user_prompt with mocked prompt_toolkit.prompt()."""
+
+ _PT_PROMPT = "prompt_toolkit.prompt"
+
+ def test_text_question_returns_answered(self):
+ from EvoScientist.stream.display import _resolve_ask_user_prompt
+
+ data = {
+ "questions": [{"question": "What dataset?", "type": "text"}],
+ "tool_call_id": "tc_1",
+ }
+ with patch(self._PT_PROMPT, side_effect=["CIFAR-10"]):
+ result = _resolve_ask_user_prompt(data)
+ assert result["status"] == "answered"
+ assert result["answers"] == ["CIFAR-10"]
+
+ def test_keyboard_interrupt_returns_cancelled(self):
+ from EvoScientist.stream.display import _resolve_ask_user_prompt
+
+ data = {
+ "questions": [{"question": "What?", "type": "text"}],
+ "tool_call_id": "tc_1",
+ }
+ with patch(self._PT_PROMPT, side_effect=KeyboardInterrupt):
+ result = _resolve_ask_user_prompt(data)
+ assert result["status"] == "cancelled"
+
+ def test_empty_questions_returns_empty(self):
+ from EvoScientist.stream.display import _resolve_ask_user_prompt
+
+ data = {"questions": [], "tool_call_id": "tc_1"}
+ result = _resolve_ask_user_prompt(data)
+ assert result["status"] == "answered"
+ assert result["answers"] == []
+
+ def test_multiple_choice_letter_mapping(self):
+ from EvoScientist.stream.display import _resolve_ask_user_prompt
+
+ data = {
+ "questions": [
+ {
+ "question": "Which?",
+ "type": "multiple_choice",
+ "choices": [{"value": "CIFAR-10"}, {"value": "ImageNet"}],
+ }
+ ],
+ "tool_call_id": "tc_1",
+ }
+ with patch(self._PT_PROMPT, side_effect=["B"]):
+ result = _resolve_ask_user_prompt(data)
+ assert result["status"] == "answered"
+ assert result["answers"] == ["ImageNet"]
+
+
+# ---------------------------------------------------------------------------
+# TUI widget (basic construction)
+# ---------------------------------------------------------------------------
+
+
+class TestAskUserWidget:
+ """Test AskUserWidget basic construction."""
+
+ def test_widget_instantiation(self):
+ from EvoScientist.cli.widgets.ask_user_widget import AskUserWidget
+
+ questions = [{"question": "Q?", "type": "text"}]
+ w = AskUserWidget(questions)
+ assert w._questions == questions
+ assert w._answers == []
+
+ def test_answered_message_class_exists(self):
+ from EvoScientist.cli.widgets.ask_user_widget import AskUserWidget
+
+ msg = AskUserWidget.Answered(["answer1"])
+ assert msg.answers == ["answer1"]
+
+ def test_cancelled_message_class_exists(self):
+ from EvoScientist.cli.widgets.ask_user_widget import AskUserWidget
+
+ msg = AskUserWidget.Cancelled()
+ assert isinstance(msg, AskUserWidget.Cancelled)
+
+
+# ---------------------------------------------------------------------------
+# Middleware __init__ exports
+# ---------------------------------------------------------------------------
+
+
+class TestMiddlewareExports:
+ """Test that ask_user types are exported from middleware package."""
+
+ def test_ask_user_middleware_exported(self):
+ from EvoScientist.middleware import AskUserMiddleware
+
+ assert AskUserMiddleware is not None
+
+ def test_ask_user_request_exported(self):
+ from EvoScientist.middleware import AskUserRequest
+
+ assert AskUserRequest is not None
+
+ def test_question_exported(self):
+ from EvoScientist.middleware import Question
+
+ assert Question is not None
+
+ def test_choice_exported(self):
+ from EvoScientist.middleware import Choice
+
+ assert Choice is not None
+
+ def test_widget_result_exported(self):
+ from EvoScientist.middleware import AskUserWidgetResult
+
+ assert AskUserWidgetResult is not None