feat: add /model command for changing models inside TUI/CLI (#162)

* rebase main

* fix: apply pr comments

* fix: fix critical issue

* fix: ordering /model in cli mode

* feat: refactor /model command handling and add Rich CLI support

---------

Co-authored-by: X-iZhang <zacharyzhang2022@gmail.com>
This commit is contained in:
Wiktor Cupiał
2026-04-22 15:43:21 +02:00
committed by GitHub
parent aa3dd00409
commit 28f3e81b4c
11 changed files with 1263 additions and 17 deletions
+15
View File
@@ -86,6 +86,21 @@ def _ensure_chat_model():
return _chat_model
def set_chat_model(model: str, provider: str | None = None):
"""Replace the cached chat model with a new one.
Called by ``/model`` to switch the LLM mid-session.
Returns the new chat model instance.
"""
global _chat_model, _EvoScientist_agent
from .llm import get_chat_model
_chat_model = get_chat_model(model=model, provider=provider)
# Invalidate the cached default agent so it gets rebuilt with the new model.
_EvoScientist_agent = None
return _chat_model
# =============================================================================
# MCP caching
# =============================================================================
+33
View File
@@ -162,6 +162,7 @@ _SLASH_COMMANDS = [
("/mcp", "Manage MCP servers"),
("/channel", "Configure messaging channels"),
("/compact", "Compact conversation to free context"),
("/model", "Switch model (--save to persist)"),
("/exit", "Quit EvoScientist"),
]
@@ -671,6 +672,7 @@ def cmd_interactive(
async def _async_main_loop():
"""Async main loop with prompt_async and channel queue checking."""
nonlocal model
async with get_checkpointer() as checkpointer:
# Handle --thread-id resume
if thread_id:
@@ -1077,6 +1079,37 @@ def cmd_interactive(
)
continue
if user_input.lower().startswith("/model"):
from ..commands.base import CommandContext
from ..commands.manager import manager as cmd_manager
from ..EvoScientist import _ensure_config
from .rich_command_ui import RichCLICommandUI
ctx = CommandContext(
agent=state["agent"],
thread_id=state["thread_id"],
ui=RichCLICommandUI(console),
workspace_dir=state["workspace_dir"],
checkpointer=checkpointer,
)
await cmd_manager.execute(user_input, ctx)
# Sync agent back if command replaced it (e.g. /model)
if ctx.agent is not state["agent"]:
state["agent"] = ctx.agent
cfg = _ensure_config()
model = cfg.model
state["status_base_snapshot"] = (
make_empty_status_snapshot(model)
)
await _refresh_status_snapshot(
reset_streaming_text=True,
)
if _channels_is_running():
_ch_mod._cli_agent = state["agent"]
_ch_mod._cli_thread_id = state["thread_id"]
continue
# Resolve @file mentions — inject file contents inline
_, message_to_send, file_warnings = resolve_file_mentions(
user_input, state["workspace_dir"]
+126
View File
@@ -0,0 +1,126 @@
"""CommandUI Protocol adapter for the Rich CLI surface.
Methods not exercised by the currently-migrated commands raise
``NotImplementedError`` rather than silently returning ``None``, so
future callers fail loudly instead of pretending the command ran.
"""
from __future__ import annotations
from typing import Any
from rich.console import Console
from rich.table import Table
from ..commands.base import CommandUI
class RichCLICommandUI(CommandUI):
"""CommandUI implementation that prints to a Rich ``Console``."""
def __init__(self, console: Console) -> None:
self.console = console
# ── Core I/O ─────────────────────────────────────────────
@property
def supports_interactive(self) -> bool:
# Rich CLI has no picker widget, but wait_for_* fall back to
# printing a table and returning None (see wait_for_model_pick).
return True
def append_system(self, text: str, style: str = "dim") -> None:
self.console.print(text, style=style)
def mount_renderable(self, renderable: Any) -> None:
self.console.print(renderable)
async def flush(self) -> None:
# Rich console flushes synchronously; nothing to await.
return
# ── /model interactive picker fallback ──────────────────
async def wait_for_model_pick(
self,
entries: list[tuple[str, str, str]],
current_model: str | None,
current_provider: str | None,
) -> tuple[str, str] | None:
"""Print the model table and return ``None``; user re-runs with
``/model <name>`` since the CLI has no interactive picker."""
table = Table(
title="Available Models",
show_header=True,
header_style="bold cyan",
)
table.add_column("Name", style="bold")
table.add_column("Provider", style="dim")
for name, _mid, prov in entries:
marker = " *" if name == current_model and prov == current_provider else ""
table.add_row(f"{name}{marker}", prov)
self.console.print(table)
self.console.print(
"[dim]Usage: /model <name> [provider] [--save] — "
"provider is optional, auto-detected from model name[/dim]"
)
return None
def update_status_after_model_change(
self, new_model: str, new_provider: str | None = None
) -> None:
"""No-op; the CLI REPL refreshes status itself after detecting an
``ctx.agent`` change post-``cmd_manager.execute``."""
return
# ── Not yet migrated ────────────────────────────────────
async def wait_for_thread_pick(
self, threads: list[dict], current_thread: str, title: str
) -> str | None:
raise NotImplementedError(
"RichCLICommandUI.wait_for_thread_pick — implement when "
"migrating /threads / /resume"
)
async def wait_for_skill_browse(
self, index: list[dict], installed_names: set[str], pre_filter_tag: str
) -> list[str] | None:
raise NotImplementedError(
"RichCLICommandUI.wait_for_skill_browse — implement when "
"migrating /skills / /evoskills"
)
async def wait_for_mcp_browse(
self, servers: list, installed_names: set[str], pre_filter_tag: str
) -> list | None:
raise NotImplementedError(
"RichCLICommandUI.wait_for_mcp_browse — implement when migrating /mcp"
)
def clear_chat(self) -> None:
raise NotImplementedError(
"RichCLICommandUI.clear_chat — implement when migrating /clear"
)
def request_quit(self) -> None:
raise NotImplementedError(
"RichCLICommandUI.request_quit — implement when migrating /exit"
)
def force_quit(self) -> None:
raise NotImplementedError(
"RichCLICommandUI.force_quit — implement when migrating /exit"
)
def start_new_session(self) -> None:
raise NotImplementedError(
"RichCLICommandUI.start_new_session — implement when migrating /new"
)
async def handle_session_resume(
self, thread_id: str, workspace_dir: str | None = None
) -> None:
raise NotImplementedError(
"RichCLICommandUI.handle_session_resume — implement when migrating /resume"
)
+103 -15
View File
@@ -19,6 +19,7 @@ from rich.console import Group
from rich.text import Text
import EvoScientist.cli.channel as _ch_mod
from EvoScientist.cli.widgets.thread_selector import ThreadPickerWidget
from ..commands import CommandContext
from ..commands import manager as cmd_manager
@@ -353,13 +354,16 @@ def run_textual_interactive(
self._picker_future: asyncio.Future | None = None
self._browser_future: asyncio.Future | None = None
self._mcp_browser_future: asyncio.Future | None = None
self._model_picker_future: asyncio.Future | None = None
self._history_suggester = HistorySuggester(DATA_DIR / "history")
self._history_index: int = -1 # -1 = not browsing history
self._history_saved_input: str = "" # saved current input before browsing
self._background_tasks: set[asyncio.Task] = set()
self._quit_pending: bool = False
self._current_model: str | None = model
self._current_provider: str | None = provider
self._status_started_at = datetime.now()
self._status_base_snapshot = make_empty_status_snapshot(model)
self._status_base_snapshot = make_empty_status_snapshot(self._current_model)
self._status_snapshot = self._status_base_snapshot
self._status_streaming_text = ""
self._status_last_input_tokens: int | None = None
@@ -430,6 +434,26 @@ def run_textual_interactive(
return await self._wait_for_mcp_browse(browser)
async def wait_for_model_pick(
self,
entries: list[tuple[str, str, str]],
current_model: str | None,
current_provider: str | None,
) -> tuple[str, str] | None:
from .widgets.model_picker import ModelPickerWidget
container = self.query_one("#chat", VerticalScroll)
picker = ModelPickerWidget(
entries,
current_model=current_model,
current_provider=current_provider,
)
await container.mount(picker)
self._schedule_scroll_to_bottom(container, delays=())
picker.focus()
return await self._wait_for_model_pick(picker)
def clear_chat(self) -> None:
container = self.query_one("#chat", VerticalScroll)
welcome = self.query_one("#welcome", Static)
@@ -452,7 +476,7 @@ def run_textual_interactive(
checkpointer=self._checkpointer,
)
self._status_started_at = datetime.now()
self._status_base_snapshot = make_empty_status_snapshot(model)
self._status_base_snapshot = make_empty_status_snapshot(self._current_model)
self._status_snapshot = self._status_base_snapshot
self._status_streaming_text = ""
self._status_last_input_tokens = None
@@ -478,7 +502,7 @@ def run_textual_interactive(
checkpointer=self._checkpointer,
)
self._status_started_at = datetime.now()
self._status_base_snapshot = make_empty_status_snapshot(model)
self._status_base_snapshot = make_empty_status_snapshot(self._current_model)
self._status_snapshot = self._status_base_snapshot
self._status_streaming_text = ""
self._status_last_input_tokens = None
@@ -725,7 +749,9 @@ def run_textual_interactive(
return {"answers": result.get("answers", []), "status": "answered"}
return {"status": "cancelled"}
async def _wait_for_thread_pick(self, picker_widget) -> str | None:
async def _wait_for_thread_pick(
self, picker_widget: ThreadPickerWidget
) -> str | None:
"""Wait for user to pick a thread from ThreadPickerWidget.
Returns the selected thread_id, or ``None`` on cancel/timeout.
@@ -808,6 +834,34 @@ def run_textual_interactive(
if self._mcp_browser_future and not self._mcp_browser_future.done():
self._mcp_browser_future.set_result(None)
async def _wait_for_model_pick(self, picker_widget) -> tuple[str, str] | None:
"""Wait for user to pick a model from ModelPickerWidget.
Returns ``(name, provider)`` or ``None`` on cancel/timeout.
"""
self._model_picker_future = asyncio.get_event_loop().create_future()
try:
return await asyncio.wait_for(self._model_picker_future, timeout=120)
except (TimeoutError, asyncio.CancelledError):
return None
finally:
self._model_picker_future = None
try:
picker_widget.remove()
except Exception:
_channel_logger.debug("model picker cleanup failed", exc_info=True)
self.query_one("#prompt", ChatTextArea).focus()
def on_model_picker_widget_picked(self, event) -> None: # type: ignore[override]
"""Handle ModelPickerWidget.Picked message."""
if self._model_picker_future and not self._model_picker_future.done():
self._model_picker_future.set_result((event.name, event.provider))
def on_model_picker_widget_cancelled(self, event) -> None: # type: ignore[override]
"""Handle ModelPickerWidget.Cancelled message."""
if self._model_picker_future and not self._model_picker_future.done():
self._model_picker_future.set_result(None)
# ── Streaming core ─────────────────────────────────────
async def _stream_with_widgets(
@@ -901,7 +955,7 @@ def run_textual_interactive(
lambda: container.scroll_end(animate=False),
)
metadata = build_metadata(self._workspace_dir, model)
metadata = build_metadata(self._workspace_dir, self._current_model)
response = ""
async def _remove_w(w: Static | None) -> None:
@@ -1832,11 +1886,12 @@ def run_textual_interactive(
# Force-resolve the future
self._ask_user_future.set_result({"type": "cancelled"})
return
# Delegate to ApprovalWidget, ThreadPickerWidget, or SkillBrowserWidget if focused
# Delegate to focused interactive widget
focused = self.focused
if focused is not None:
from .widgets.approval_widget import ApprovalWidget
from .widgets.mcp_browser import MCPBrowserWidget
from .widgets.model_picker import ModelPickerWidget
from .widgets.skill_browser import SkillBrowserWidget
from .widgets.thread_selector import ThreadPickerWidget
@@ -1852,6 +1907,14 @@ def run_textual_interactive(
if isinstance(focused, MCPBrowserWidget):
focused.action_cancel()
return
if isinstance(focused, ModelPickerWidget):
focused.action_cancel()
if (
self._model_picker_future
and not self._model_picker_future.done()
):
self._model_picker_future.set_result(None)
return
if self._queued_messages:
self._queued_messages.pop()
self._render_queue_indicator()
@@ -1865,12 +1928,13 @@ def run_textual_interactive(
self._render_completions()
return
# Skip if an ApprovalWidget, AskUserWidget, ThreadPickerWidget, or SkillBrowserWidget has focus
# Skip if an interactive picker widget 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.mcp_browser import MCPBrowserWidget
from .widgets.model_picker import ModelPickerWidget
from .widgets.skill_browser import SkillBrowserWidget
from .widgets.thread_selector import ThreadPickerWidget
@@ -1889,6 +1953,9 @@ def run_textual_interactive(
if isinstance(focused, MCPBrowserWidget):
focused.action_move_up()
return
if isinstance(focused, ModelPickerWidget):
focused.action_move_up()
return
if self._queued_messages:
last = self._queued_messages.pop()
prompt = self.query_one("#prompt", ChatTextArea)
@@ -1924,6 +1991,7 @@ def run_textual_interactive(
from .widgets.approval_widget import ApprovalWidget
from .widgets.ask_user_widget import AskUserWidget
from .widgets.mcp_browser import MCPBrowserWidget
from .widgets.model_picker import ModelPickerWidget
from .widgets.skill_browser import SkillBrowserWidget
from .widgets.thread_selector import ThreadPickerWidget
@@ -1942,6 +2010,9 @@ def run_textual_interactive(
if isinstance(focused, MCPBrowserWidget):
focused.action_move_down()
return
if isinstance(focused, ModelPickerWidget):
focused.action_move_down()
return
# History browsing (down key)
if self._history_index >= 0:
@@ -2076,6 +2147,12 @@ def run_textual_interactive(
try:
if await cmd_manager.execute(command, ctx):
# Sync agent back if command replaced it (e.g. /model)
if ctx.agent is not self._agent:
self._agent = ctx.agent
if _channels_is_running():
_ch_mod._cli_agent = self._agent
_ch_mod._cli_thread_id = self._conversation_tid
# Do NOT invalidate the usage baseline after /compact.
# build_session_status_snapshot() only counts raw checkpoint
# messages (~46 tokens) and misses system prompt + tool
@@ -2247,25 +2324,25 @@ def run_textual_interactive(
self._status_base_snapshot = apply_user_text_to_snapshot(
make_usage_status_snapshot(
self._status_last_input_tokens,
model_name=model,
model_name=self._current_model,
),
pending,
)
else:
self._status_base_snapshot = await build_session_status_snapshot(
self._conversation_tid,
model_name=model,
model_name=self._current_model,
pending_user_text=pending,
)
elif self._status_last_input_tokens is not None:
self._status_base_snapshot = make_usage_status_snapshot(
self._status_last_input_tokens,
model_name=model,
model_name=self._current_model,
)
else:
self._status_base_snapshot = await build_session_status_snapshot(
self._conversation_tid,
model_name=model,
model_name=self._current_model,
)
if reset_streaming_text:
self._status_streaming_text = ""
@@ -2278,7 +2355,7 @@ def run_textual_interactive(
self._status_last_input_tokens = input_tokens
self._status_base_snapshot = make_usage_status_snapshot(
input_tokens,
model_name=model,
model_name=self._current_model,
)
self._rebuild_status_snapshot()
@@ -2293,10 +2370,21 @@ def run_textual_interactive(
self._status_last_input_tokens = tokens_after
self._status_base_snapshot = make_usage_status_snapshot(
tokens_after,
model_name=model,
model_name=self._current_model,
)
self._rebuild_status_snapshot()
def update_status_after_model_change(
self, new_model: str, new_provider: str | None = None
) -> None:
"""Update the status bar and welcome banner after /model switches the LLM."""
self._current_model = new_model
if new_provider is not None:
self._current_provider = new_provider
self._status_base_snapshot = make_empty_status_snapshot(new_model)
self._rebuild_status_snapshot()
self._render_welcome()
def _set_status_streaming_text(self, text: str | None) -> None:
"""Update in-flight assistant text shown in the context bar."""
new_text = text or ""
@@ -2342,8 +2430,8 @@ def run_textual_interactive(
thread_id=self._conversation_tid,
workspace_dir=self._workspace_dir,
mode=mode,
model=model,
provider=provider,
model=self._current_model,
provider=self._current_provider,
ui_backend="tui",
channels=channels_info,
)
+294
View File
@@ -0,0 +1,294 @@
"""Inline model picker widget for /model command in TUI.
Keyboard-driven widget mounted directly into the chat container.
Models are grouped by provider with a search/filter input.
"""
from __future__ import annotations
from typing import TYPE_CHECKING, Any, ClassVar
from rich.text import Text
from textual.binding import Binding, BindingType
from textual.containers import Container
from textual.message import Message
from textual.widget import Widget
from textual.widgets import Static
if TYPE_CHECKING:
from textual import events
from textual.app import ComposeResult
def _build_items(
entries: list[tuple[str, str, str]],
current_model: str | None = None,
current_provider: str | None = None,
filter_text: str = "",
) -> list[dict]:
"""Build the flat item list rendered by ModelPickerWidget.
Returns a list of::
{"type": "header", "label": str}
{"type": "model", "name": str, "model_id": str, "provider": str, "current": bool}
"""
# Apply filter
if filter_text:
ft = filter_text.lower()
entries = [
(n, mid, p) for n, mid, p in entries if ft in n.lower() or ft in p.lower()
]
# Group by provider preserving order
groups: dict[str, list[tuple[str, str, str]]] = {}
for name, model_id, provider in entries:
if provider not in groups:
groups[provider] = []
groups[provider].append((name, model_id, provider))
items: list[dict] = []
for provider, models in groups.items():
items.append({"type": "header", "label": provider})
for name, model_id, prov in models:
is_current = name == current_model and prov == current_provider
items.append(
{
"type": "model",
"name": name,
"model_id": model_id,
"provider": prov,
"current": is_current,
}
)
return items
class ModelPickerWidget(Widget):
"""Inline model picker -- mounts in chat, keyboard-driven.
Posts ``Picked(name, provider)`` on Enter, ``Cancelled()`` on Esc.
Type to filter models.
"""
can_focus = True
can_focus_children = False
DEFAULT_CSS = """
ModelPickerWidget {
height: auto;
max-height: 30;
margin: 1 0;
padding: 0 1;
background: $surface;
border: solid $primary;
}
ModelPickerWidget .picker-title {
height: 1;
text-style: bold;
color: $primary;
}
ModelPickerWidget .picker-filter {
height: 1;
padding: 0 1;
color: $text;
}
ModelPickerWidget .picker-rows {
height: auto;
max-height: 22;
overflow-y: auto;
}
ModelPickerWidget .picker-header {
height: 1;
padding: 0 1;
margin-top: 1;
}
ModelPickerWidget .picker-row {
height: 1;
padding: 0 1;
}
ModelPickerWidget .picker-row-selected {
background: $primary;
text-style: bold;
}
ModelPickerWidget .picker-help {
height: 1;
color: $text-muted;
text-style: italic;
}
"""
BINDINGS: ClassVar[list[BindingType]] = [
Binding("up", "move_up", "Up", show=False),
Binding("down", "move_down", "Down", show=False),
Binding("enter", "select", "Select", show=False),
Binding("escape", "cancel", "Cancel", show=False),
Binding("backspace", "backspace", "Backspace", show=False),
]
class Picked(Message):
def __init__(self, name: str, provider: str) -> None:
super().__init__()
self.name = name
self.provider = provider
class Cancelled(Message):
"""Posted when user cancels selection."""
def __init__(
self,
entries: list[tuple[str, str, str]],
*,
current_model: str | None = None,
current_provider: str | None = None,
title: str = ">>> Select model <<<",
**kwargs: Any,
) -> None:
super().__init__(**kwargs)
self._entries = entries
self._current_model = current_model
self._current_provider = current_provider
self._title = title
self._filter_text = ""
self._items = _build_items(
entries,
current_model=current_model,
current_provider=current_provider,
)
self._selected = self._first_model_index()
self._row_widgets: list[Static] = []
self._filter_widget: Static | None = None
def _first_model_index(self) -> int:
for i, item in enumerate(self._items):
if item["type"] == "model":
return i
return 0
def _move(self, direction: int) -> None:
if not self._items:
return
i = (self._selected + direction) % len(self._items)
steps = 0
while self._items[i]["type"] != "model" and steps < len(self._items):
i = (i + direction) % len(self._items)
steps += 1
if self._items[i]["type"] == "model":
self._selected = i
self._update_rows()
def _rebuild(self) -> None:
"""Rebuild items from filter and re-render."""
self._items = _build_items(
self._entries,
current_model=self._current_model,
current_provider=self._current_provider,
filter_text=self._filter_text,
)
self._selected = self._first_model_index()
# Re-mount rows
rows_container = self.query_one(".picker-rows", Container)
for w in list(rows_container.children):
w.remove()
self._row_widgets.clear()
for item in self._items:
css = "picker-header" if item["type"] == "header" else "picker-row"
widget = Static("", classes=css)
self._row_widgets.append(widget)
rows_container.mount(widget)
self._update_rows()
self._update_filter()
def compose(self) -> ComposeResult:
yield Static(self._title, classes="picker-title")
self._filter_widget = Static("", classes="picker-filter")
yield self._filter_widget
with Container(classes="picker-rows"):
for item in self._items:
css = "picker-header" if item["type"] == "header" else "picker-row"
widget = Static("", classes=css)
self._row_widgets.append(widget)
yield widget
yield Static(
"\u2191/\u2193 navigate \u00b7 Enter select \u00b7 Type to filter \u00b7 Esc cancel",
classes="picker-help",
)
def on_mount(self) -> None:
self._update_rows()
self._update_filter()
self.call_later(self.focus)
def _update_filter(self) -> None:
if self._filter_widget is not None:
if self._filter_text:
t = Text()
t.append(" Filter: ", style="dim")
t.append(self._filter_text, style="bold")
t.append("\u2588", style="blink")
self._filter_widget.update(t)
else:
self._filter_widget.update(
Text(" Type to filter...", style="dim italic")
)
def _update_rows(self) -> None:
for i, (item, widget) in enumerate(
zip(self._items, self._row_widgets, strict=False)
):
widget.remove_class("picker-row-selected")
if item["type"] == "header":
t = Text()
t.append("\u2500\u2500 ", style="bold cyan")
t.append(item["label"], style="bold cyan")
widget.update(t)
else:
is_selected = i == self._selected
t = Text()
cursor = "\u25b8 " if is_selected else " "
t.append(cursor, style="bold cyan" if is_selected else "dim")
t.append(item["name"], style="bold" if is_selected else "")
if item["current"]:
t.append(" *", style="bold green")
t.append(f" ({item['provider']})", style="dim italic")
widget.update(t)
if is_selected:
widget.add_class("picker-row-selected")
widget.scroll_visible()
def on_key(self, event: events.Key) -> None:
# Let bindings handle special keys
if event.key in ("up", "down", "enter", "escape", "backspace"):
return
# Printable characters -> filter
if event.character and event.character.isprintable():
self._filter_text += event.character
self._rebuild()
event.prevent_default()
def action_backspace(self) -> None:
if self._filter_text:
self._filter_text = self._filter_text[:-1]
self._rebuild()
def action_move_up(self) -> None:
self._move(-1)
def action_move_down(self) -> None:
self._move(1)
def action_select(self) -> None:
if not self._items or self._selected >= len(self._items):
self.post_message(self.Cancelled())
return
item = self._items[self._selected]
if item["type"] == "model":
self.post_message(self.Picked(item["name"], item["provider"]))
else:
self.post_message(self.Cancelled())
def action_cancel(self) -> None:
self.post_message(self.Cancelled())
def on_blur(self, event: events.Blur) -> None:
self.call_after_refresh(self.focus)
+6
View File
@@ -35,6 +35,12 @@ class CommandUI(Protocol):
async def wait_for_mcp_browse(
self, servers: list, installed_names: set[str], pre_filter_tag: str
) -> list | None: ...
async def wait_for_model_pick(
self,
entries: list[tuple[str, str, str]],
current_model: str | None,
current_provider: str | None,
) -> tuple[str, str] | None: ...
def clear_chat(self) -> None: ...
def request_quit(self) -> None: ...
def force_quit(self) -> None: ...
@@ -1,5 +1,5 @@
from __future__ import annotations
from . import channel, general, mcp, session, skills
from . import channel, general, mcp, model, session, skills
__all__ = ["channel", "general", "mcp", "session", "skills"]
__all__ = ["channel", "general", "mcp", "model", "session", "skills"]
@@ -0,0 +1,173 @@
from __future__ import annotations
from typing import ClassVar
from ..base import Argument, Command, CommandContext
from ..manager import manager
def extract_model_and_provider(args: list[str]) -> tuple[str, str]:
"""Parse model name and provider from argument list.
Args:
args: Non-empty argument list (model_name [provider]).
Returns:
``(model_name, provider)`` tuple.
Raises:
ValueError: If the model is not in the registry.
"""
from ...llm.models import MODELS
model_name = args[0]
provider_override = args[1] if len(args) > 1 else None
if model_name not in MODELS:
raise ValueError(f"Unknown model '{model_name}'")
if provider_override:
provider = provider_override
else:
_, provider = MODELS[model_name]
return model_name, provider
class ModelCommand(Command):
"""Switch the LLM model for the current session."""
name = "/model"
description = "Switch model (--save to persist)"
# ``--save`` is parsed manually in ``execute`` via ``"--save" in args``;
# ``type=bool`` below is declarative metadata, not enforced by the manager.
arguments: ClassVar[list[Argument]] = [
Argument(
name="model_name",
type=str,
description="Model short name (e.g. claude-sonnet-4-6). Opens picker if omitted.",
required=False,
),
Argument(
name="--save",
type=bool,
description="Save the choice to config file",
required=False,
),
]
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
from ...EvoScientist import _ensure_config
from ...llm.models import list_models_by_provider
cfg = _ensure_config()
current_model = cfg.model
current_provider = cfg.provider
# Parse --save flag
save = "--save" in args
args = [a for a in args if a != "--save"]
if args:
try:
model_name, provider = extract_model_and_provider(args)
except ValueError:
ctx.ui.append_system(
f"Unknown model '{args[0]}'. Use /model to browse available models.",
style="red",
)
return
await self._apply_model(ctx, model_name, provider, save=save)
return
# Interactive picker
if not ctx.ui.supports_interactive:
ctx.ui.append_system(
"Usage: /model <name> [provider] [--save]",
style="yellow",
)
return
entries = list_models_by_provider()
result = await ctx.ui.wait_for_model_pick(
entries,
current_model=current_model,
current_provider=current_provider,
)
if result is None:
return
name, provider = result
await self._apply_model(ctx, name, provider, save=save)
async def _apply_model(
self,
ctx: CommandContext,
model_name: str,
provider: str,
*,
save: bool = False,
) -> None:
import copy
from ...cli.agent import _load_agent
from ...EvoScientist import _ensure_config, set_chat_model
cfg = _ensure_config()
# Build a temporary config to verify the agent can be created
# before mutating any global state.
temp_cfg = copy.copy(cfg)
temp_cfg.model = model_name
temp_cfg.provider = provider
try:
new_agent = _load_agent(
workspace_dir=ctx.workspace_dir,
checkpointer=ctx.checkpointer,
config=temp_cfg,
)
except Exception as e:
ctx.ui.append_system(f"Failed to switch model: {e}", style="red")
return
# Agent built successfully — now commit the change globally.
try:
set_chat_model(model_name, provider=provider)
except Exception as e:
ctx.ui.append_system(f"Failed to switch model: {e}", style="red")
return
cfg.model = model_name
cfg.provider = provider
ctx.agent = new_agent
# Persist to config file if --save was given
if save:
from ...config.settings import set_config_value
set_config_value("model", model_name)
set_config_value("provider", provider)
# Propagate to channel module if channels are running
try:
import EvoScientist.cli.channel as _ch_mod
if getattr(_ch_mod, "_cli_agent", None) is not None:
_ch_mod._cli_agent = new_agent
except Exception:
pass
# Update status bar if available
update_model_fn = getattr(ctx.ui, "update_status_after_model_change", None)
if callable(update_model_fn):
update_model_fn(model_name, provider)
saved_note = " (saved to config)" if save else ""
ctx.ui.append_system(
f"Switched to {model_name} ({provider}){saved_note}", style="green"
)
manager.register(ModelCommand())
+16
View File
@@ -461,6 +461,22 @@ def list_models() -> list[str]:
return result
def list_models_by_provider() -> list[tuple[str, str, str]]:
"""List all unique (short_name, model_id, provider) entries.
Returns:
De-duplicated list of model entries preserving registry order.
"""
seen: set[tuple[str, str]] = set()
result: list[tuple[str, str, str]] = []
for name, model_id, provider in _MODEL_ENTRIES:
key = (name, provider)
if key not in seen:
seen.add(key)
result.append((name, model_id, provider))
return result
def get_model_info(model: str) -> tuple[str, str] | None:
"""Get the (model_id, provider) tuple for a short name.
+321
View File
@@ -0,0 +1,321 @@
"""Tests for the /model command and extract_model_and_provider helper."""
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from tests.conftest import run_async as _run
class TestExtractModelAndProvider:
"""Unit tests for the argument parser helper."""
def test_known_model_no_provider(self):
from EvoScientist.commands.implementation.model import (
extract_model_and_provider,
)
name, prov = extract_model_and_provider(["claude-sonnet-4-6"])
assert name == "claude-sonnet-4-6"
assert prov == "anthropic"
def test_known_model_with_provider_override(self):
from EvoScientist.commands.implementation.model import (
extract_model_and_provider,
)
name, prov = extract_model_and_provider(["claude-sonnet-4-6", "openrouter"])
assert name == "claude-sonnet-4-6"
assert prov == "openrouter"
def test_unknown_model_no_provider_raises(self):
from EvoScientist.commands.implementation.model import (
extract_model_and_provider,
)
with pytest.raises(ValueError, match="Unknown model"):
extract_model_and_provider(["nonexistent-model-xyz"])
def test_unknown_model_with_provider_still_raises(self):
from EvoScientist.commands.implementation.model import (
extract_model_and_provider,
)
# Unknown models are always rejected, even with an explicit provider
with pytest.raises(ValueError, match="Unknown model"):
extract_model_and_provider(["my-custom-model", "custom-openai"])
def test_provider_override_on_known_model(self):
from EvoScientist.commands.implementation.model import (
extract_model_and_provider,
)
# Known model with explicit provider override uses the override
name, prov = extract_model_and_provider(["claude-sonnet-4-6", "openrouter"])
assert name == "claude-sonnet-4-6"
assert prov == "openrouter"
class TestModelCommandUnknownModel:
"""Verify error message for unknown models."""
def test_unknown_model_shows_error(self):
from EvoScientist.commands.implementation.model import ModelCommand
cmd = ModelCommand()
ui = MagicMock()
ui.supports_interactive = True
cfg = SimpleNamespace(model="claude-sonnet-4-6", provider="anthropic")
ctx = MagicMock()
ctx.ui = ui
with patch(
"EvoScientist.EvoScientist._ensure_config",
return_value=cfg,
):
_run(cmd.execute(ctx, ["nonexistent-model-xyz"]))
ui.append_system.assert_called_once()
call_args = ui.append_system.call_args
assert "Unknown model" in call_args[0][0]
assert call_args[1]["style"] == "red"
class TestModelCommandPickerCancelled:
"""Verify no-op when the interactive picker is cancelled."""
def test_picker_returns_none(self):
from EvoScientist.commands.implementation.model import ModelCommand
cmd = ModelCommand()
ui = MagicMock()
ui.supports_interactive = True
ui.wait_for_model_pick = AsyncMock(return_value=None)
cfg = SimpleNamespace(model="claude-sonnet-4-6", provider="anthropic")
ctx = MagicMock()
ctx.ui = ui
with patch(
"EvoScientist.EvoScientist._ensure_config",
return_value=cfg,
):
_run(cmd.execute(ctx, []))
# No model switch should have happened
ui.append_system.assert_not_called()
class TestModelCommandSwitch:
"""Verify a successful model switch updates config and rebuilds agent."""
def test_switch_known_model(self):
from EvoScientist.commands.implementation.model import ModelCommand
cmd = ModelCommand()
ui = MagicMock()
ui.supports_interactive = True
cfg = SimpleNamespace(model="claude-sonnet-4-6", provider="anthropic")
new_agent = MagicMock()
ctx = MagicMock()
ctx.ui = ui
ctx.workspace_dir = "/tmp/test"
ctx.checkpointer = MagicMock()
with (
patch(
"EvoScientist.EvoScientist._ensure_config",
return_value=cfg,
),
patch(
"EvoScientist.EvoScientist.set_chat_model",
),
patch(
"EvoScientist.cli.agent._load_agent",
return_value=new_agent,
),
):
_run(cmd.execute(ctx, ["claude-opus-4-6"]))
# Config should be updated
assert cfg.model == "claude-opus-4-6"
assert cfg.provider == "anthropic"
# Agent should be replaced on context
assert ctx.agent == new_agent
# Success message shown
ui.append_system.assert_called_once()
msg = ui.append_system.call_args[0][0]
assert "claude-opus-4-6" in msg
assert "anthropic" in msg
def test_switch_with_save_flag(self):
from EvoScientist.commands.implementation.model import ModelCommand
cmd = ModelCommand()
ui = MagicMock()
ui.supports_interactive = True
cfg = SimpleNamespace(model="claude-sonnet-4-6", provider="anthropic")
ctx = MagicMock()
ctx.ui = ui
ctx.workspace_dir = "/tmp/test"
ctx.checkpointer = MagicMock()
with (
patch(
"EvoScientist.EvoScientist._ensure_config",
return_value=cfg,
),
patch("EvoScientist.EvoScientist.set_chat_model"),
patch(
"EvoScientist.cli.agent._load_agent",
return_value=MagicMock(),
),
patch("EvoScientist.config.settings.set_config_value") as mock_save,
):
_run(cmd.execute(ctx, ["claude-opus-4-6", "--save"]))
# Config file should be updated
mock_save.assert_any_call("model", "claude-opus-4-6")
mock_save.assert_any_call("provider", "anthropic")
# Success message should mention save
msg = ui.append_system.call_args[0][0]
assert "saved to config" in msg
def test_switch_without_save_flag_does_not_persist(self):
from EvoScientist.commands.implementation.model import ModelCommand
cmd = ModelCommand()
ui = MagicMock()
ui.supports_interactive = True
cfg = SimpleNamespace(model="claude-sonnet-4-6", provider="anthropic")
ctx = MagicMock()
ctx.ui = ui
ctx.workspace_dir = "/tmp/test"
ctx.checkpointer = MagicMock()
with (
patch(
"EvoScientist.EvoScientist._ensure_config",
return_value=cfg,
),
patch("EvoScientist.EvoScientist.set_chat_model"),
patch(
"EvoScientist.cli.agent._load_agent",
return_value=MagicMock(),
),
patch("EvoScientist.config.settings.set_config_value") as mock_save,
):
_run(cmd.execute(ctx, ["claude-opus-4-6"]))
# Config file should NOT be updated
mock_save.assert_not_called()
# Message should not mention save
msg = ui.append_system.call_args[0][0]
assert "saved to config" not in msg
class TestModelCommandFailure:
"""Verify error handling when set_chat_model raises."""
def test_set_chat_model_error(self):
from EvoScientist.commands.implementation.model import ModelCommand
cmd = ModelCommand()
ui = MagicMock()
ui.supports_interactive = True
cfg = SimpleNamespace(model="claude-sonnet-4-6", provider="anthropic")
ctx = MagicMock()
ctx.ui = ui
ctx.workspace_dir = "/tmp/test"
ctx.checkpointer = MagicMock()
with (
patch(
"EvoScientist.EvoScientist._ensure_config",
return_value=cfg,
),
patch(
"EvoScientist.cli.agent._load_agent",
return_value=MagicMock(),
),
patch(
"EvoScientist.EvoScientist.set_chat_model",
side_effect=RuntimeError("API key missing"),
) as mock_set,
):
_run(cmd.execute(ctx, ["claude-opus-4-6"]))
mock_set.assert_called_once()
ui.append_system.assert_called_once()
call_args = ui.append_system.call_args
assert "Failed to switch model" in call_args[0][0]
assert call_args[1]["style"] == "red"
class TestModelCommandLoadAgentFailure:
"""Verify the transactional ordering: when ``_load_agent`` raises,
nothing downstream (``set_chat_model``, ``cfg`` mutation,
``set_config_value``) should happen.
This is the core guarantee of the refactor that established
"build agent first, commit state only on success". Without this test
the ordering could silently regress (e.g. if ``_apply_model`` were
reordered to call ``set_chat_model`` first)."""
def test_load_agent_error_is_transactional(self):
from EvoScientist.commands.implementation.model import ModelCommand
cmd = ModelCommand()
ui = MagicMock()
ui.supports_interactive = True
cfg = SimpleNamespace(model="claude-sonnet-4-6", provider="anthropic")
ctx = MagicMock()
ctx.ui = ui
ctx.workspace_dir = "/tmp/test"
ctx.checkpointer = MagicMock()
with (
patch(
"EvoScientist.EvoScientist._ensure_config",
return_value=cfg,
),
patch(
"EvoScientist.cli.agent._load_agent",
side_effect=RuntimeError("agent build failed"),
) as mock_load,
patch(
"EvoScientist.EvoScientist.set_chat_model",
) as mock_set,
patch(
"EvoScientist.config.settings.set_config_value",
) as mock_save,
):
# Pass ``--save`` to strengthen the assertion: if the ordering
# ever regresses, ``set_config_value`` would be called with
# stale data.
_run(cmd.execute(ctx, ["claude-opus-4-6", "--save"]))
# _load_agent was attempted (transactional first step).
mock_load.assert_called_once()
# Downstream side-effects must NOT have happened.
mock_set.assert_not_called()
mock_save.assert_not_called()
# Config must be untouched.
assert cfg.model == "claude-sonnet-4-6"
assert cfg.provider == "anthropic"
# User sees a red error message.
ui.append_system.assert_called_once()
call_args = ui.append_system.call_args
assert "Failed to switch model" in call_args[0][0]
assert call_args[1]["style"] == "red"
+174
View File
@@ -0,0 +1,174 @@
"""Tests for the Rich CLI CommandUI adapter."""
from unittest.mock import MagicMock
import pytest
from rich.console import Console
from rich.table import Table
from tests.conftest import run_async as _run
def _make_ui():
"""Build a RichCLICommandUI backed by a MagicMock console."""
from EvoScientist.cli.rich_command_ui import RichCLICommandUI
console = MagicMock(spec=Console)
ui = RichCLICommandUI(console)
return ui, console
class TestBasicIO:
"""Core CommandUI methods used by /model path."""
def test_supports_interactive_true(self):
ui, _ = _make_ui()
assert ui.supports_interactive is True
def test_append_system_forwards_style(self):
ui, console = _make_ui()
ui.append_system("hello", style="green")
console.print.assert_called_once_with("hello", style="green")
def test_append_system_default_style(self):
ui, console = _make_ui()
ui.append_system("info")
console.print.assert_called_once_with("info", style="dim")
def test_mount_renderable_preserves_type(self):
ui, console = _make_ui()
table = Table(title="demo")
ui.mount_renderable(table)
console.print.assert_called_once_with(table)
def test_flush_is_async_noop(self):
ui, console = _make_ui()
_run(ui.flush())
# flush should not print anything
console.print.assert_not_called()
class TestWaitForModelPick:
"""CLI model picker fallback: print table + return None."""
def test_returns_none(self):
ui, _ = _make_ui()
entries = [
("claude-sonnet-4-6", "anthropic/claude-sonnet", "anthropic"),
("gpt-4o", "openai/gpt-4o", "openai"),
]
result = _run(
ui.wait_for_model_pick(
entries,
current_model="claude-sonnet-4-6",
current_provider="anthropic",
)
)
assert result is None
def test_prints_table_with_current_model_marker(self):
ui, console = _make_ui()
entries = [
("claude-sonnet-4-6", "anthropic/claude-sonnet", "anthropic"),
("gpt-4o", "openai/gpt-4o", "openai"),
]
_run(
ui.wait_for_model_pick(
entries,
current_model="claude-sonnet-4-6",
current_provider="anthropic",
)
)
# First call renders the Table (Rich renderable), second prints usage.
assert console.print.call_count == 2
first_arg = console.print.call_args_list[0].args[0]
assert isinstance(first_arg, Table)
usage_arg = console.print.call_args_list[1].args[0]
assert "Usage: /model" in usage_arg
assert "--save" in usage_arg
def test_no_current_model_no_marker(self):
ui, console = _make_ui()
entries = [("claude-sonnet-4-6", "anthropic/claude-sonnet", "anthropic")]
_run(
ui.wait_for_model_pick(
entries,
current_model=None,
current_provider=None,
)
)
# Just asserts the coroutine runs without marker-branch issues.
assert console.print.call_count == 2
def test_empty_entries_still_prints_header_and_usage(self):
ui, console = _make_ui()
result = _run(
ui.wait_for_model_pick(
[],
current_model=None,
current_provider=None,
)
)
assert result is None
# Header table + usage hint should still be printed even with
# no entries.
assert console.print.call_count == 2
class TestUpdateStatusHook:
"""update_status_after_model_change is a deliberate no-op on CLI."""
def test_no_op(self):
ui, console = _make_ui()
ui.update_status_after_model_change("claude-opus-4-6", "anthropic")
console.print.assert_not_called()
class TestUnmigratedMethodsStub:
"""Protocol methods not yet wired for CLI must raise NotImplementedError.
These stubs are signposts for the A1 migration (see
cli-commandmanager-migration.md). Each one is replaced by a real
implementation when its corresponding command is migrated.
"""
def test_wait_for_thread_pick(self):
ui, _ = _make_ui()
with pytest.raises(NotImplementedError, match="/threads"):
_run(ui.wait_for_thread_pick([], "tid", "title"))
def test_wait_for_skill_browse(self):
ui, _ = _make_ui()
with pytest.raises(NotImplementedError, match="/skills"):
_run(ui.wait_for_skill_browse([], set(), ""))
def test_wait_for_mcp_browse(self):
ui, _ = _make_ui()
with pytest.raises(NotImplementedError, match="/mcp"):
_run(ui.wait_for_mcp_browse([], set(), ""))
def test_clear_chat(self):
ui, _ = _make_ui()
with pytest.raises(NotImplementedError, match="/clear"):
ui.clear_chat()
def test_request_quit(self):
ui, _ = _make_ui()
with pytest.raises(NotImplementedError, match="/exit"):
ui.request_quit()
def test_force_quit(self):
ui, _ = _make_ui()
with pytest.raises(NotImplementedError, match="/exit"):
ui.force_quit()
def test_start_new_session(self):
ui, _ = _make_ui()
with pytest.raises(NotImplementedError, match="/new"):
ui.start_new_session()
def test_handle_session_resume(self):
ui, _ = _make_ui()
with pytest.raises(NotImplementedError, match="/resume"):
_run(ui.handle_session_resume("tid"))