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