diff --git a/CLAUDE.md b/CLAUDE.md index 6814d11..702da11 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -48,7 +48,7 @@ TOML load/store (`ConfigStore`): `[providers]` tagged union presets, `[database] ### Memory (`memory/`) -Async SQLAlchemy + aiosqlite. `MemoryStore`: schema init (+ lightweight SQLite `ALTER` for new columns), default local user, sessions (bound to `workspace` path), messages stored as msgspec chat message JSON. +Async SQLAlchemy + aiosqlite. `MemoryStore`: schema init (+ lightweight SQLite `ALTER` for new columns), default local user, sessions (bound to `workspace` path; optional `provider_name`/`model`), messages stored as msgspec chat message JSON. ### Agent (`agent/`) @@ -83,7 +83,7 @@ Shared interactive I/O: `ask` / `choose` / `form` / `confirm` with pluggable bac Click app + readline REPL. Entry: `plyngent` / `python -m plyngent`. -- **`plyngent chat`**: provider/model (flags or interactive); SQLite sessions via `[database]` (file DB under user data if unset/`:memory:`); sessions bound to workspace; resume latest for cwd/`--workspace` by default (`--new` / `--session`). One-shot: `-p/--prompt` and non-TTY stdin; exit codes 0/1/2/3; `--yes`, `--stream/--no-stream`, `--quiet`. Root `--log-level`. +- **`plyngent chat`**: provider/model (flags or interactive; Tab via readline in `prompting`); sessions store `provider_name`/`model` and restore on resume; SQLite via `[database]` (file DB under user data if unset/`:memory:`); workspace-bound; resume latest for cwd/`--workspace` by default (`--new` / `--session`). One-shot: `-p/--prompt` and non-TTY stdin; exit codes 0/1/2/3; `--yes`, `--stream/--no-stream`, `--quiet`. Root `--log-level`. - Slash: Click group in `cli/slash.py` + `awaitlet` for async work; Tab completer from registry + ParamType `shell_complete`. Multiline `"""` … `"""`; `/edit` via `$EDITOR`. - Explicit `/resume` or `--session` from another workspace prompts: **keep** / **update** / **abort**. - Failed/cancelled turns: user message kept; partial assistant/tool rolled back; Ctrl+C cancels turn; TTY confirms off-loop; auto-retry 10s/20s/30s then `/retry`. diff --git a/README.md b/README.md index ce050a0..8cfeda0 100644 --- a/README.md +++ b/README.md @@ -85,7 +85,7 @@ plyngent chat --session 3 | `--yes` | Allow destructive tools without confirm (also for one-shot) | | `--log-level` | On the root CLI: `DEBUG`, `INFO`, `WARNING`, … | -Sessions resume the **most recently updated** session for the current workspace unless you pass `--new` or `--session`. +Sessions resume the **most recently updated** session for the current workspace unless you pass `--new` or `--session`. Each session remembers the last **provider** and **model** (restored on resume so you are not re-prompted). ### One-shot (scripts / CI) diff --git a/src/plyngent/cli/app.py b/src/plyngent/cli/app.py index c638a84..fef5af6 100644 --- a/src/plyngent/cli/app.py +++ b/src/plyngent/cli/app.py @@ -140,7 +140,7 @@ async def _run_oneshot(state: ReplState, prompt_text: str) -> int: return EXIT_OK if ok else EXIT_TURN_FAILED -async def _run_chat( +async def _run_chat( # noqa: C901 — chat orchestration *, config_path: Path | None, provider_name: str | None, @@ -166,22 +166,36 @@ async def _run_chat( names = ", ".join(sorted(store.bad_providers.keys())) click.secho(f"warning: ignored bad providers: {names}", fg="yellow", err=True) - try: - pname, provider = select_provider( - store.providers, - preferred=provider_name, - interactive=interactive, - ) - model_id = select_model(provider, preferred=model, interactive=interactive) - _ = create_client(provider) - except ProviderNotSupportedError as exc: - raise click.ClickException(str(exc)) from exc - _setup_workspace_and_hooks(store, workspace, interactive=interactive) confirm_destructive: bool | None = False if yes else None memory = await MemoryStore.open(_database_config(store, quiet=quiet or oneshot)) try: + # Prefer session-remembered LLM when resuming (unless flags override). + preferred_provider = provider_name + preferred_model = model + if not oneshot and not new_session and preferred_provider is None: + if session_id is not None: + row = await memory.get_session(session_id) + else: + row = await memory.get_latest_session(workspace=workspace) + if row is not None: + if preferred_provider is None and row.provider_name: + preferred_provider = row.provider_name + if preferred_model is None and row.model: + preferred_model = row.model + + try: + pname, provider = select_provider( + store.providers, + preferred=preferred_provider, + interactive=interactive, + ) + model_id = select_model(provider, preferred=preferred_model, interactive=interactive) + _ = create_client(provider) + except ProviderNotSupportedError as exc: + raise click.ClickException(str(exc)) from exc + state = ReplState( config=store, memory=memory, @@ -205,6 +219,8 @@ async def _run_chat( oneshot=oneshot, quiet=quiet, ) + # Ensure new sessions / flag overrides are stored for next resume. + await state.persist_llm_selection() if oneshot: assert prompt_text is not None diff --git a/src/plyngent/cli/export.py b/src/plyngent/cli/export.py index 9665e49..21bd750 100644 --- a/src/plyngent/cli/export.py +++ b/src/plyngent/cli/export.py @@ -34,6 +34,8 @@ def session_export_payload( created_at: datetime | None, updated_at: datetime | None, messages: Sequence[AnyChatMessage], + provider_name: str | None = None, + model: str | None = None, ) -> dict[str, object]: """Build a JSON-serializable dict for a session transcript. @@ -44,6 +46,8 @@ def session_export_payload( "session_id": sid, "name": name, "workspace": workspace, + "provider": provider_name, + "model": model, "created_at": _iso(created_at), "updated_at": _iso(updated_at), "messages": [msgspec.to_builtins(m) for m in messages], diff --git a/src/plyngent/cli/selection.py b/src/plyngent/cli/selection.py index b8555af..4d8ac2b 100644 --- a/src/plyngent/cli/selection.py +++ b/src/plyngent/cli/selection.py @@ -4,6 +4,8 @@ from typing import TYPE_CHECKING import click +from plyngent.prompting import ChoiceOption, choose + if TYPE_CHECKING: from collections.abc import Mapping @@ -16,7 +18,7 @@ def select_provider( preferred: str | None = None, interactive: bool = True, ) -> tuple[str, Provider]: - """Pick a provider by name or interactive prompt.""" + """Pick a provider by name or interactive prompt (readline + Tab).""" if not providers: msg = "no providers configured; edit your plyngent.toml" raise click.ClickException(msg) @@ -37,11 +39,15 @@ def select_provider( msg = f"multiple providers; pass --provider ({', '.join(names)})" raise click.ClickException(msg) - click.echo("Available providers:") - for index, name in enumerate(names, start=1): - preset = type(providers[name]).__struct_config__.tag - click.echo(f" {index}. {name} ({preset})") - choice = click.prompt("Select provider", type=click.Choice(names), show_choices=True) + options = [ + ChoiceOption( + label=name, + description=str(type(providers[name]).__struct_config__.tag), + value=name, + ) + for name in names + ] + choice = choose("Select provider", options, allow_custom=False) return choice, providers[choice] @@ -51,7 +57,7 @@ def select_model( preferred: str | None = None, interactive: bool = True, ) -> str: - """Pick a model id from provider.models or free-form prompt.""" + """Pick a model id from provider.models or free-form prompt (readline + Tab).""" model_names = sorted(provider.models.keys()) if preferred is not None: if model_names and preferred not in provider.models: @@ -72,9 +78,12 @@ def select_model( raise click.ClickException(msg) if model_names: - click.echo("Available models:") - for index, name in enumerate(model_names, start=1): - click.echo(f" {index}. {name}") - return click.prompt("Select model", type=click.Choice(model_names), show_choices=True) + return choose( + "Select model", + model_names, + allow_custom=False, + ) - return click.prompt("Model id (not listed in config)", type=str) + from plyngent.prompting import ask + + return ask("Model id (not listed in config)") diff --git a/src/plyngent/cli/slash.py b/src/plyngent/cli/slash.py index e0977db..d94fdf6 100644 --- a/src/plyngent/cli/slash.py +++ b/src/plyngent/cli/slash.py @@ -455,6 +455,8 @@ def export_cmd(state: ReplState, parts: tuple[str, ...]) -> None: created_at=row.created_at, updated_at=row.updated_at, messages=messages, + provider_name=row.provider_name, + model=row.model, ) ) else: @@ -526,8 +528,12 @@ def provider_cmd(state: ReplState, name: str | None) -> None: pname, provider = select_provider(state.config.providers, preferred=name.strip()) state.provider_name = pname state.provider = provider + # Keep model if still listed; else first model on the new provider. + if state.model not in provider.models and provider.models: + state.model = next(iter(sorted(provider.models.keys()))) state.rebuild_client() - click.echo(f"switched provider to {pname}") + _await(state.persist_llm_selection()) + click.echo(f"switched provider to {pname} model={state.model}") except (click.ClickException, ProviderNotSupportedError) as exc: click.echo(f"error: {exc}") @@ -543,6 +549,7 @@ def model_cmd(state: ReplState, model_id: str | None) -> None: try: state.model = select_model(state.provider, preferred=model_id.strip()) state.rebuild_client() + _await(state.persist_llm_selection()) click.echo(f"switched model to {state.model}") except click.ClickException as exc: click.echo(f"error: {exc}") diff --git a/src/plyngent/cli/state.py b/src/plyngent/cli/state.py index 8904506..5a9a7c7 100644 --- a/src/plyngent/cli/state.py +++ b/src/plyngent/cli/state.py @@ -126,8 +126,77 @@ class ReplState: return self._set_workspace(Path(row.workspace)) + async def persist_llm_selection(self) -> None: + """Write current provider/model onto the active session row (if any).""" + if self.session_id is None: + return + _ = await self.memory.update_session_llm( + self.session_id, + provider_name=self.provider_name, + model=self.model, + ) + + def _try_set_provider(self, pname: str) -> bool: + import click + + from plyngent.cli.selection import select_provider + from plyngent.runtime import ProviderNotSupportedError + + if pname not in self.config.providers: + return False + try: + name, provider = select_provider( + self.config.providers, + preferred=pname, + interactive=False, + ) + except click.ClickException, ProviderNotSupportedError: + return False + if name != self.provider_name or provider is not self.provider: + self.provider_name = name + self.provider = provider + return True + return False + + def _try_set_model(self, model_id: str) -> bool: + import click + + from plyngent.cli.selection import select_model + + try: + resolved = select_model(self.provider, preferred=model_id, interactive=False) + except click.ClickException: + return False + if resolved != self.model: + self.model = resolved + return True + return False + + def apply_session_llm(self, row: SessionRow) -> bool: + """Apply stored provider/model from ``row`` when still valid in config. + + Returns True when selection changed (caller should rebuild agent). + """ + changed = False + pname = row.provider_name + if pname: + changed = self._try_set_provider(pname) or changed + if row.model and self._try_set_model(row.model): + changed = True + elif changed and self.provider.models and self.model not in self.provider.models: + self.model = next(iter(sorted(self.provider.models.keys()))) + return changed + if row.model: + return self._try_set_model(row.model) + return False + async def new_session(self, name: str = "chat") -> None: - session = await self.memory.create_session(name=name, workspace=self.workspace) + session = await self.memory.create_session( + name=name, + workspace=self.workspace, + provider_name=self.provider_name, + model=self.model, + ) self.session_id = session.sid self.agent = self._make_agent() @@ -180,7 +249,10 @@ class ReplState: raise ValueError(msg) from exc self.session_id = session_id - self.agent = self._make_agent() + if self.apply_session_llm(row): + self.rebuild_client() + else: + self.agent = self._make_agent() await self.agent.load_history() async def resume_latest_or_new(self, name: str = "chat") -> str: @@ -191,7 +263,10 @@ class ReplState: return "new" # Same-workspace list: no mismatch prompt expected. self.session_id = latest.sid - self.agent = self._make_agent() + if self.apply_session_llm(latest): + self.rebuild_client() + else: + self.agent = self._make_agent() await self.agent.load_history() _ = await self.memory.touch_session(latest.sid) return "resume" diff --git a/src/plyngent/memory/database/schema.py b/src/plyngent/memory/database/schema.py index a315ded..867f64d 100644 --- a/src/plyngent/memory/database/schema.py +++ b/src/plyngent/memory/database/schema.py @@ -29,6 +29,9 @@ class Session(PlyngentBase): name: Mapped[str] = mapped_column(String(64)) # Absolute workspace path this chat is bound to (tools root); null = legacy unbound. workspace: Mapped[str | None] = mapped_column(String(1024), nullable=True, index=True) + # Last selected provider/model for this session (config provider key + model id). + provider_name: Mapped[str | None] = mapped_column(String(128), nullable=True) + model: Mapped[str | None] = mapped_column(String(256), nullable=True) created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), server_default=func.now()) updated_at: Mapped[datetime] = mapped_column( DateTime(timezone=True), server_default=func.now(), onupdate=func.now() diff --git a/src/plyngent/memory/database/store.py b/src/plyngent/memory/database/store.py index 3471814..e70f3db 100644 --- a/src/plyngent/memory/database/store.py +++ b/src/plyngent/memory/database/store.py @@ -67,6 +67,7 @@ class MemoryStore: async with self._engine.begin() as conn: _ = await conn.run_sync(PlyngentBase.metadata.create_all) await conn.run_sync(_migrate_session_workspace) + await conn.run_sync(_migrate_session_llm) async def close(self) -> None: """Dispose the underlying engine.""" @@ -100,17 +101,26 @@ class MemoryStore: uid: int | None = None, name: str = "default", workspace: str | Path | None = None, + provider_name: str | None = None, + model: str | None = None, ) -> Session: """Create a chat session for ``uid`` (default local user when omitted). ``workspace`` is stored as a resolved absolute path when provided. + Optional ``provider_name`` / ``model`` remember LLM selection for resume. """ if uid is None: user = await self.ensure_default_user() uid = user.uid ws = normalize_workspace(workspace) async with self._session_factory() as session: - row = Session(uid=uid, name=name, workspace=ws) + row = Session( + uid=uid, + name=name, + workspace=ws, + provider_name=provider_name, + model=model, + ) session.add(row) await session.commit() await session.refresh(row) @@ -176,6 +186,28 @@ class MemoryStore: await session.refresh(row) return row + async def update_session_llm( + self, + sid: int, + *, + provider_name: str | None = None, + model: str | None = None, + ) -> Session: + """Update remembered provider/model for a session (omit a field to leave it).""" + async with self._session_factory() as session: + row = await session.get(Session, sid) + if row is None: + msg = f"session not found: {sid}" + raise ValueError(msg) + if provider_name is not None: + row.provider_name = provider_name + if model is not None: + row.model = model + row.updated_at = datetime.now(UTC) + await session.commit() + await session.refresh(row) + return row + async def rename_session(self, sid: int, name: str) -> Session: """Rename a session (max 64 characters, non-empty after strip).""" cleaned = name.strip() @@ -244,14 +276,35 @@ class MemoryStore: return result.scalars().all() +def _session_columns(sync_conn: object) -> set[str]: + from sqlalchemy.engine import Connection + + if not isinstance(sync_conn, Connection): + return set() + rows = sync_conn.execute(text("PRAGMA table_info(session)")).fetchall() + return {str(row[1]) for row in rows} + + def _migrate_session_workspace(sync_conn: object) -> None: """Add ``session.workspace`` on existing SQLite DBs created before the column existed.""" from sqlalchemy.engine import Connection if not isinstance(sync_conn, Connection): return - rows = sync_conn.execute(text("PRAGMA table_info(session)")).fetchall() - columns = {str(row[1]) for row in rows} + columns = _session_columns(sync_conn) if "workspace" in columns: return _ = sync_conn.execute(text("ALTER TABLE session ADD COLUMN workspace VARCHAR(1024)")) + + +def _migrate_session_llm(sync_conn: object) -> None: + """Add ``session.provider_name`` / ``session.model`` for remembered LLM selection.""" + from sqlalchemy.engine import Connection + + if not isinstance(sync_conn, Connection): + return + columns = _session_columns(sync_conn) + if "provider_name" not in columns: + _ = sync_conn.execute(text("ALTER TABLE session ADD COLUMN provider_name VARCHAR(128)")) + if "model" not in columns: + _ = sync_conn.execute(text("ALTER TABLE session ADD COLUMN model VARCHAR(256)")) diff --git a/src/plyngent/prompting.py b/src/plyngent/prompting.py index 676c290..de53fe4 100644 --- a/src/plyngent/prompting.py +++ b/src/plyngent/prompting.py @@ -47,7 +47,13 @@ class PromptBackend(Protocol): def is_interactive(self) -> bool: ... - def read_line(self, prompt: str, *, default: str | None = None) -> str: ... + def read_line( + self, + prompt: str, + *, + default: str | None = None, + completions: Sequence[str] | None = None, + ) -> str: ... def confirm(self, prompt: str, *, default: bool = False) -> bool: ... @@ -56,20 +62,60 @@ class PromptBackend(Protocol): def secho(self, message: str, *, fg: str | None = None, err: bool = False) -> None: ... +def _readline_input(prompt: str, *, completions: Sequence[str] | None = None) -> str: + """``input()`` with optional Tab completion via readline when available.""" + try: + import readline + except ImportError: + return input(prompt) + + previous_completer = readline.get_completer() + previous_delims = readline.get_completer_delims() + options = list(completions or ()) + + def completer(text: str, state: int) -> str | None: + matches = [c for c in options if c.startswith(text)] if text else list(options) + if state < len(matches): + return matches[state] + return None + + try: + readline.set_completer_delims(" \t\n") + readline.set_completer(completer if options else None) + # GNU readline + libedit bindings (best-effort). + with contextlib.suppress(Exception): + _ = readline.parse_and_bind("tab: complete") + with contextlib.suppress(Exception): + _ = readline.parse_and_bind("bind ^I rl_complete") + return input(prompt) + finally: + readline.set_completer(previous_completer) + with contextlib.suppress(Exception): + readline.set_completer_delims(previous_delims) + + class ClickPromptBackend: - """Click/TTY backend for interactive prompts.""" + """Click/TTY backend for interactive prompts (readline + Tab when available).""" def is_interactive(self) -> bool: return bool(sys.stdin.isatty() and sys.stdout.isatty()) - def read_line(self, prompt: str, *, default: str | None = None) -> str: + def read_line( + self, + prompt: str, + *, + default: str | None = None, + completions: Sequence[str] | None = None, + ) -> str: try: - if default is None: - return str(click.prompt(prompt, prompt_suffix=": ")) - return str(click.prompt(prompt, default=default, show_default=True, prompt_suffix=": ")) - except (click.Abort, KeyboardInterrupt, EOFError) as exc: + display = f"{prompt} [{default}]: " if default is not None else f"{prompt}: " + raw = _readline_input(display, completions=completions) + except (KeyboardInterrupt, EOFError) as exc: msg = "prompt cancelled" raise NonInteractiveError(msg) from exc + if not raw.strip() and default is not None: + return default + return raw def confirm(self, prompt: str, *, default: bool = False) -> bool: try: @@ -91,7 +137,14 @@ class NonInteractiveBackend: def is_interactive(self) -> bool: return False - def read_line(self, prompt: str, *, default: str | None = None) -> str: + def read_line( + self, + prompt: str, + *, + default: str | None = None, + completions: Sequence[str] | None = None, + ) -> str: + del completions if default is not None: return default msg = f"non-interactive: cannot prompt for {prompt!r}" @@ -184,14 +237,22 @@ def _show_choices( backend.echo(" (or type a custom answer)") -def ask(prompt: str, *, default: str | None = None) -> str: - """Free-form question; always allows arbitrary user text.""" +def ask( + prompt: str, + *, + default: str | None = None, + completions: Sequence[str] | None = None, +) -> str: + """Free-form question; always allows arbitrary user text. + + Optional ``completions`` enable Tab completion when the backend supports it. + """ backend = get_prompt_backend() if not backend.is_interactive() and default is None: msg = f"non-interactive: cannot prompt for {prompt!r}" raise NonInteractiveError(msg) backend.secho(prompt, fg="yellow") - return backend.read_line("Answer", default=default).strip() + return backend.read_line("Answer", default=default, completions=completions).strip() def choose( @@ -210,13 +271,24 @@ def choose( _show_choices(backend, prompt, choices, allow_custom=allow_custom) default_display = _default_display(choices, default) + # Tab: indices, labels, and resolved values. + completions: list[str] = [] + for index, option in enumerate(choices, start=1): + completions.append(str(index)) + completions.append(option.label) + if option.resolved_value != option.label: + completions.append(option.resolved_value) if not backend.is_interactive() and default is None: msg = f"non-interactive: cannot prompt for {prompt!r}" raise NonInteractiveError(msg) while True: - raw = backend.read_line("Choice", default=default_display).strip() + raw = backend.read_line( + "Choice", + default=default_display, + completions=completions, + ).strip() if not raw: if default is not None: return default diff --git a/tests/test_cli/test_export.py b/tests/test_cli/test_export.py index 060c68e..229cce3 100644 --- a/tests/test_cli/test_export.py +++ b/tests/test_cli/test_export.py @@ -38,7 +38,11 @@ def test_session_export_json_roundtrip() -> None: created_at=now, updated_at=now, messages=messages, + provider_name="local", + model="tiny", ) + assert payload.get("provider") == "local" + assert payload.get("model") == "tiny" assert isinstance(payload, dict) raw = encode_session_export_json(payload) assert "7" in raw diff --git a/tests/test_cli/test_repl_commands.py b/tests/test_cli/test_repl_commands.py index 5f233f1..188f946 100644 --- a/tests/test_cli/test_repl_commands.py +++ b/tests/test_cli/test_repl_commands.py @@ -173,6 +173,23 @@ async def test_rename_slash(state: ReplState) -> None: assert row.name == "my-chat" +async def test_model_switch_persists(state: ReplState) -> None: + from plyngent.config.models import ModelConfig + + state.provider.models = { + "m1": ModelConfig(), + "m2": ModelConfig(), + } + state.model = "m1" + await state.persist_llm_selection() + assert await handle_slash(state, "/model m2") is True + assert state.model == "m2" + assert state.session_id is not None + row = await state.memory.get_session(state.session_id) + assert row is not None + assert row.model == "m2" + + async def test_delete_slash_confirm( state: ReplState, capsys: pytest.CaptureFixture[str], diff --git a/tests/test_cli/test_selection.py b/tests/test_cli/test_selection.py index 2fe85aa..30f9d44 100644 --- a/tests/test_cli/test_selection.py +++ b/tests/test_cli/test_selection.py @@ -38,11 +38,23 @@ def test_select_model_from_list() -> None: assert select_model(provider, preferred="m1") == "m1" -def test_select_model_prompt_when_empty(monkeypatch: pytest.MonkeyPatch) -> None: +def test_select_model_prompt_when_empty() -> None: + from plyngent.prompting import temporary_backend + from tests.test_prompting import ScriptedBackend + provider = OpenAIProvider(access_key_or_token="sk") + with temporary_backend(ScriptedBackend(["gpt-test"])): + assert select_model(provider) == "gpt-test" - def _prompt(*_args: object, **_kwargs: object) -> str: - return "gpt-test" - monkeypatch.setattr("click.prompt", _prompt) - assert select_model(provider) == "gpt-test" +def test_select_provider_interactive_choose() -> None: + from plyngent.prompting import temporary_backend + from tests.test_prompting import ScriptedBackend + + providers = { + "a": OpenAIProvider(access_key_or_token="sk"), + "b": OpenAICompatibleProvider(access_key_or_token="sk", url="https://x/v1"), + } + with temporary_backend(ScriptedBackend(["2"])): + name, _ = select_provider(providers) + assert name == "b" diff --git a/tests/test_memory/test_store.py b/tests/test_memory/test_store.py index 6ebb8b4..f984466 100644 --- a/tests/test_memory/test_store.py +++ b/tests/test_memory/test_store.py @@ -129,6 +129,24 @@ async def test_delete_session_cascades_messages(store: MemoryStore) -> None: assert await store.delete_session(session.sid) is False +async def test_session_llm_remembered(store: MemoryStore) -> None: + session = await store.create_session( + name="llm", + provider_name="deepseek", + model="deepseek-v4-flash", + ) + row = await store.get_session(session.sid) + assert row is not None + assert row.provider_name == "deepseek" + assert row.model == "deepseek-v4-flash" + updated = await store.update_session_llm(session.sid, model="deepseek-v4-pro") + assert updated.provider_name == "deepseek" + assert updated.model == "deepseek-v4-pro" + again = await store.update_session_llm(session.sid, provider_name="other") + assert again.provider_name == "other" + assert again.model == "deepseek-v4-pro" + + async def test_update_session_workspace(store: MemoryStore, tmp_path: object) -> None: from pathlib import Path diff --git a/tests/test_prompting.py b/tests/test_prompting.py index 5d38981..14e3d00 100644 --- a/tests/test_prompting.py +++ b/tests/test_prompting.py @@ -27,8 +27,14 @@ class ScriptedBackend: def is_interactive(self) -> bool: return self.interactive - def read_line(self, prompt: str, *, default: str | None = None) -> str: - del prompt + def read_line( + self, + prompt: str, + *, + default: str | None = None, + completions: object = None, + ) -> str: + del prompt, completions if self.lines: return self.lines.pop(0) if default is not None: