Files
EvoScientist/tests/test_ask_user.py
T
m4 bb9bed82e1 feat(runtime)!: remove auxiliary model role, resolve all roles from run snapshot
ModelRole collapses to "primary": every role (main, tool selector, memory
agents, subagents, summarizer) resolves to the snapshot's frozen primary
model, per design 6.1/8.3 — users typically configure a single usable LLM,
so compile-time auxiliary bindings were bypassing run snapshots and
mis-attributing usage. Legacy auxiliary keys in stored snapshots, registry
JSON, and thread metadata are tolerated on read and dropped.

BREAKING CHANGE: ThreadModelSelection no longer carries an auxiliary ref;
snapshot selection_hash is computed over {primary, reasoning_effort} only;
ConfigurableModelMiddleware(role="auxiliary") is rejected.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-07-23 10:18:44 +08:00

780 lines
27 KiB
Python

"""Tests for the ask_user middleware, stream events, state, and UI helpers."""
from contextlib import contextmanager
from unittest.mock import MagicMock, patch
import pytest
@contextmanager
def _patched_runtime_config(config):
import langgraph.config as langgraph_config
if isinstance(config, Exception):
with patch.object(langgraph_config, "get_config", side_effect=config):
yield
else:
with patch.object(langgraph_config, "get_config", return_value=config):
yield
# ---------------------------------------------------------------------------
# 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"
def test_system_prompt_mentions_resource_estimation(self):
from EvoScientist.middleware.ask_user import ASK_USER_SYSTEM_PROMPT
assert "estimation" in ASK_USER_SYSTEM_PROMPT.lower()
def test_system_prompt_mentions_timeout(self):
from EvoScientist.middleware.ask_user import ASK_USER_SYSTEM_PROMPT
assert "timeout" in ASK_USER_SYSTEM_PROMPT.lower()
def test_tool_description_mentions_resource(self):
from EvoScientist.middleware.ask_user import ASK_USER_TOOL_DESCRIPTION
assert "resource" in ASK_USER_TOOL_DESCRIPTION.lower()
@pytest.mark.parametrize("mode", ["manual", "auto"])
def test_manual_and_auto_keep_ask_user(self, mode):
from langchain.agents.middleware.types import ModelRequest
from EvoScientist.middleware.ask_user import (
ASK_USER_SYSTEM_PROMPT,
AskUserMiddleware,
)
middleware = AskUserMiddleware()
request = ModelRequest(
model=MagicMock(),
messages=[],
tools=[middleware.tools[0], {"name": "execute"}],
)
handler = MagicMock(return_value="response")
with _patched_runtime_config({"configurable": {"review_mode": mode}}):
result = middleware.wrap_model_call(request, handler)
assert result == "response"
forwarded = handler.call_args.args[0]
assert [
tool.name if hasattr(tool, "name") else tool["name"]
for tool in forwarded.tools
] == [
"ask_user",
"execute",
]
assert ASK_USER_SYSTEM_PROMPT in forwarded.system_message.text
def test_full_removes_ask_user_and_adds_unattended_prompt(self):
from langchain.agents.middleware.types import ModelRequest
from EvoScientist.middleware.ask_user import (
FULL_APPROVE_SYSTEM_PROMPT,
AskUserMiddleware,
)
middleware = AskUserMiddleware()
request = ModelRequest(
model=MagicMock(),
messages=[],
tools=[
middleware.tools[0],
{"name": "execute"},
{"type": "function", "function": {"name": "ask_user"}},
],
)
handler = MagicMock(return_value="response")
with _patched_runtime_config({"configurable": {"review_mode": "full"}}):
result = middleware.wrap_model_call(request, handler)
assert result == "response"
forwarded = handler.call_args.args[0]
assert forwarded.tools == [{"name": "execute"}]
assert FULL_APPROVE_SYSTEM_PROMPT in forwarded.system_message.text
async def test_full_async_wrapper_matches_sync_behavior(self):
from langchain.agents.middleware.types import ModelRequest
from EvoScientist.middleware.ask_user import AskUserMiddleware
middleware = AskUserMiddleware()
request = ModelRequest(
model=MagicMock(),
messages=[],
tools=[middleware.tools[0], {"name": "execute"}],
)
async def handler(forwarded):
assert forwarded.tools == [{"name": "execute"}]
return "response"
with _patched_runtime_config({"configurable": {"review_mode": "full"}}):
result = await middleware.awrap_model_call(request, handler)
assert result == "response"
def test_full_tool_defense_does_not_interrupt(self):
from EvoScientist.middleware.ask_user import (
FULL_APPROVE_TOOL_MESSAGE,
AskUserMiddleware,
)
middleware = AskUserMiddleware()
with (
_patched_runtime_config({"configurable": {"review_mode": "full"}}),
patch("EvoScientist.middleware.ask_user.interrupt") as mock_interrupt,
):
result = middleware.tools[0].func(
questions=[],
tool_call_id="tool-call-1",
)
mock_interrupt.assert_not_called()
message = result.update["messages"][0]
assert message.content == FULL_APPROVE_TOOL_MESSAGE
assert message.tool_call_id == "tool-call-1"
class TestReviewMode:
@pytest.mark.parametrize(
("config", "expected"),
[
({"configurable": {"review_mode": "manual"}}, "manual"),
({"configurable": {"review_mode": "auto"}}, "auto"),
({"configurable": {"review_mode": "full"}}, "full"),
({"configurable": {"review_mode": "invalid"}}, "manual"),
({"configurable": "invalid"}, "manual"),
("invalid", "manual"),
],
)
def test_review_mode_parsing(self, config, expected):
from EvoScientist.middleware.ask_user import _review_mode
with _patched_runtime_config(config):
assert _review_mode() == expected
def test_outside_runtime_defaults_to_manual(self):
from EvoScientist.middleware.ask_user import _review_mode
with _patched_runtime_config(RuntimeError("outside runtime")):
assert _review_mode() == "manual"
# ---------------------------------------------------------------------------
# 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
def test_auto_mode_default_is_false(self):
from EvoScientist.config.settings import EvoScientistConfig
cfg = EvoScientistConfig()
assert cfg.auto_mode is False
def test_auto_mode_set_to_true(self):
from EvoScientist.config.settings import EvoScientistConfig
cfg = EvoScientistConfig(auto_mode=True)
assert cfg.auto_mode is True
@patch("EvoScientist.middleware.create_tool_selector_middleware", return_value=[])
@patch("EvoScientist.EvoScientist._ensure_chat_model")
@patch("EvoScientist.EvoScientist._ensure_config")
def test_auto_approve_still_includes_ask_user_middleware(
mock_config, mock_model, mock_tool_selector
):
cfg = MagicMock()
cfg.enable_ask_user = True
cfg.auto_approve = True
cfg.auto_mode = False
mock_config.return_value = cfg
mock_model.return_value = MagicMock(profile={"max_input_tokens": 200_000})
from EvoScientist.EvoScientist import _get_default_middleware
type_names = [type(m).__name__ for m in _get_default_middleware()]
assert "AskUserMiddleware" in type_names
@patch("EvoScientist.middleware.create_tool_selector_middleware", return_value=[])
@patch("EvoScientist.EvoScientist._ensure_chat_model")
@patch("EvoScientist.EvoScientist._ensure_config")
def test_auto_mode_disables_ask_user_middleware(
mock_config, mock_model, mock_tool_selector
):
cfg = MagicMock()
cfg.enable_ask_user = True
cfg.auto_approve = True
cfg.auto_mode = True
mock_config.return_value = cfg
mock_model.return_value = MagicMock(profile={"max_input_tokens": 200_000})
from EvoScientist.EvoScientist import _get_default_middleware
type_names = [type(m).__name__ for m in _get_default_middleware()]
assert "AskUserMiddleware" not in type_names
@patch("EvoScientist.middleware.create_tool_selector_middleware", return_value=[])
@patch("EvoScientist.EvoScientist._ensure_chat_model")
@patch("EvoScientist.EvoScientist._ensure_config")
def test_for_async_subagent_omits_ask_user_middleware(
mock_config, mock_model, mock_tool_selector
):
"""``AskUserMiddleware`` uses ``interrupt()`` to wait on user input.
Async sub-agents run in the langgraph dev subprocess where the parent
only holds a ``task_id`` and has no UI path to surface or resume an
interrupt. Including ``AskUserMiddleware`` would deadlock the
sub-agent the first time the LLM calls ``ask_user``. The
``for_async_subagent=True`` flag must suppress it even when the user
has globally enabled ``ask_user``.
"""
cfg = MagicMock()
cfg.enable_ask_user = True
cfg.auto_approve = False
cfg.auto_mode = False
mock_config.return_value = cfg
mock_model.return_value = MagicMock(profile={"max_input_tokens": 200_000})
from EvoScientist.EvoScientist import _get_default_middleware
# Sanity: with the default flag, ask_user IS present.
default_names = [type(m).__name__ for m in _get_default_middleware()]
assert "AskUserMiddleware" in default_names
# With for_async_subagent=True, ask_user is suppressed.
async_names = [
type(m).__name__ for m in _get_default_middleware(for_async_subagent=True)
]
assert "AskUserMiddleware" not in async_names
# Other middleware must remain — only ask_user is filtered.
assert "ConfigurableModelMiddleware" in async_names
assert "ContextEditingMiddleware" in async_names
# The fallback chain was removed (design doc 8.3): it must appear in
# neither the default nor the async sub-agent middleware stack.
assert "ModelFallbackMiddleware" not in async_names
assert "ModelFallbackMiddleware" not in default_names
# ---------------------------------------------------------------------------
# Rich CLI prompt (mocking input)
# ---------------------------------------------------------------------------
class TestRichCLIPrompt:
"""Test _resolve_ask_user_prompt with mocked questionary."""
def test_text_question_returns_answered(self):
from unittest.mock import MagicMock
from EvoScientist.stream.display import _resolve_ask_user_prompt
data = {
"questions": [{"question": "What dataset?", "type": "text"}],
"tool_call_id": "tc_1",
}
mock_text = MagicMock()
mock_text.return_value.ask.return_value = "CIFAR-10"
with patch("questionary.text", mock_text):
result = _resolve_ask_user_prompt(data)
assert result["status"] == "answered"
assert result["answers"] == ["CIFAR-10"]
def test_keyboard_interrupt_returns_cancelled(self):
from unittest.mock import MagicMock
from EvoScientist.stream.display import _resolve_ask_user_prompt
data = {
"questions": [{"question": "What?", "type": "text"}],
"tool_call_id": "tc_1",
}
# questionary returns None when user presses Ctrl+C
mock_text = MagicMock()
mock_text.return_value.ask.return_value = None
with patch("questionary.text", mock_text):
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_selection(self):
from unittest.mock import MagicMock
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",
}
mock_select = MagicMock()
mock_select.return_value.ask.return_value = "ImageNet"
with patch("questionary.select", mock_select):
result = _resolve_ask_user_prompt(data)
assert result["status"] == "answered"
assert result["answers"] == ["ImageNet"]
def test_multiple_choice_other_option(self):
from unittest.mock import MagicMock
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",
}
mock_select = MagicMock()
mock_select.return_value.ask.return_value = "Other (type your answer)"
mock_text = MagicMock()
mock_text.return_value.ask.return_value = "custom dataset"
with (
patch("questionary.select", mock_select),
patch("questionary.text", mock_text),
):
result = _resolve_ask_user_prompt(data)
assert result["status"] == "answered"
assert result["answers"] == ["custom dataset"]
# ---------------------------------------------------------------------------
# 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