core/memory+cli: remember session LLM; readline prompts

Store provider_name/model on sessions (SQLite migrate); restore on
resume and persist on create//model//provider. Startup prefers session
selection over re-asking. ask/choose and provider/model pick use
readline with Tab completion when available.
This commit is contained in:
2026-07-15 14:12:10 +08:00
parent cc298c6cd2
commit 64d288a9b1
15 changed files with 349 additions and 53 deletions
+2 -2
View File
@@ -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`.
+1 -1
View File
@@ -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)
+28 -12
View File
@@ -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
+4
View File
@@ -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],
+21 -12
View File
@@ -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)")
+8 -1
View File
@@ -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}")
+78 -3
View File
@@ -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"
+3
View File
@@ -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()
+56 -3
View File
@@ -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)"))
+84 -12
View File
@@ -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
+4
View File
@@ -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
+17
View File
@@ -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],
+17 -5
View File
@@ -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"
+18
View File
@@ -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
+8 -2
View File
@@ -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: