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