421a664336
- Remove legacy provider profiles, admin-token auth, /model command, model picker widget, and config.yaml LLM fields (design doc section 10) - Wire CLI/channels/cron and async sub-agents through the local snapshot entry; run creation rejects model config outside runtime_snapshot_id - Add periodic run-snapshot TTL cleanup to the config service lifespan - Isolate tests from the real config dir and activate the registry where run/model paths fail closed in bootstrap Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
789 lines
27 KiB
Python
789 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_auxiliary_chat_model")
|
|
@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_aux_model, mock_tool_selector
|
|
):
|
|
cfg = MagicMock()
|
|
cfg.enable_ask_user = True
|
|
cfg.auto_approve = True
|
|
cfg.auto_mode = False
|
|
cfg.auxiliary_model = ""
|
|
cfg.auxiliary_provider = ""
|
|
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_auxiliary_chat_model")
|
|
@patch("EvoScientist.EvoScientist._ensure_chat_model")
|
|
@patch("EvoScientist.EvoScientist._ensure_config")
|
|
def test_auto_mode_disables_ask_user_middleware(
|
|
mock_config, mock_model, mock_aux_model, mock_tool_selector
|
|
):
|
|
cfg = MagicMock()
|
|
cfg.enable_ask_user = True
|
|
cfg.auto_approve = True
|
|
cfg.auto_mode = True
|
|
cfg.auxiliary_model = ""
|
|
cfg.auxiliary_provider = ""
|
|
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_auxiliary_chat_model")
|
|
@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_aux_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
|
|
cfg.auxiliary_model = ""
|
|
cfg.auxiliary_provider = ""
|
|
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
|