feat: add ask_user middleware and interactive widget for user prompts

- Implemented `ask_user` middleware to facilitate agent-initiated questions during research workflows.
- Created `AskUserWidget` for interactive user prompts, supporting both text and multiple choice questions.
- Enhanced `StreamEventEmitter` to handle `ask_user` interrupts and updated event handling in `stream_agent_events`.
- Added state management for pending `ask_user` events in `StreamState`.
- Developed validation for question structures and parsing for responses.
- Introduced unit tests covering middleware functionality, event handling, and widget behavior.
This commit is contained in:
X-iZhang
2026-03-10 18:17:52 +00:00
parent bc7a0b753c
commit c15c0dfdf0
18 changed files with 1693 additions and 10 deletions
+9 -1
View File
@@ -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)
+121
View File
@@ -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",
+6
View File
@@ -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)
+5
View File
@@ -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}"
+3
View File
@@ -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,
)
+104 -4
View File
@@ -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
+3
View File
@@ -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
+2
View File
@@ -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",
]
+375
View File
@@ -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()
+3
View File
@@ -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"
+12
View File
@@ -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",
]
+405
View File
@@ -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))
+12 -2
View File
@@ -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)
+92 -2
View File
@@ -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" <b><style fg='#1565c0'>Choice [{letters}]:</style></b> ")).strip()
if raw.upper() == other_letter:
raw = pt_prompt(HTML(" <b><style fg='#42a5f5'>&gt; Your answer:</style></b> ")).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(" <b><style fg='#42a5f5'>&gt; Answer:</style></b> ")).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,
+12
View File
@@ -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."""
+25 -1
View File
@@ -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", [])
+5
View File
@@ -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", "")
+499
View File
@@ -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