diff --git a/EvoScientist/EvoScientist.py b/EvoScientist/EvoScientist.py index b00b739..99b9af7 100644 --- a/EvoScientist/EvoScientist.py +++ b/EvoScientist/EvoScientist.py @@ -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 # ============================================================================= diff --git a/EvoScientist/cli/interactive.py b/EvoScientist/cli/interactive.py index 28b725d..c66123b 100644 --- a/EvoScientist/cli/interactive.py +++ b/EvoScientist/cli/interactive.py @@ -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"] diff --git a/EvoScientist/cli/rich_command_ui.py b/EvoScientist/cli/rich_command_ui.py new file mode 100644 index 0000000..122a472 --- /dev/null +++ b/EvoScientist/cli/rich_command_ui.py @@ -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 `` 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 [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" + ) diff --git a/EvoScientist/cli/tui_interactive.py b/EvoScientist/cli/tui_interactive.py index 2b0fde9..aa18622 100644 --- a/EvoScientist/cli/tui_interactive.py +++ b/EvoScientist/cli/tui_interactive.py @@ -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, ) diff --git a/EvoScientist/cli/widgets/model_picker.py b/EvoScientist/cli/widgets/model_picker.py new file mode 100644 index 0000000..590e6ad --- /dev/null +++ b/EvoScientist/cli/widgets/model_picker.py @@ -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) diff --git a/EvoScientist/commands/base.py b/EvoScientist/commands/base.py index f3e6198..7bf41c9 100644 --- a/EvoScientist/commands/base.py +++ b/EvoScientist/commands/base.py @@ -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: ... diff --git a/EvoScientist/commands/implementation/__init__.py b/EvoScientist/commands/implementation/__init__.py index 087d1bd..24e3cab 100644 --- a/EvoScientist/commands/implementation/__init__.py +++ b/EvoScientist/commands/implementation/__init__.py @@ -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"] diff --git a/EvoScientist/commands/implementation/model.py b/EvoScientist/commands/implementation/model.py new file mode 100644 index 0000000..ccc4209 --- /dev/null +++ b/EvoScientist/commands/implementation/model.py @@ -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 [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()) diff --git a/EvoScientist/llm/models.py b/EvoScientist/llm/models.py index 3d440f0..2cb77a7 100644 --- a/EvoScientist/llm/models.py +++ b/EvoScientist/llm/models.py @@ -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. diff --git a/tests/test_model_command.py b/tests/test_model_command.py new file mode 100644 index 0000000..51bce76 --- /dev/null +++ b/tests/test_model_command.py @@ -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" diff --git a/tests/test_rich_command_ui.py b/tests/test_rich_command_ui.py new file mode 100644 index 0000000..e6aac92 --- /dev/null +++ b/tests/test_rich_command_ui.py @@ -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"))