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:
@@ -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)
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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}"
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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()
|
||||
@@ -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"
|
||||
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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))
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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'>> 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'>> 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,
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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", [])
|
||||
|
||||
@@ -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", "")
|
||||
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user