mirror of
https://github.com/NCBM/plyngent.git
synced 2026-07-23 05:55:16 +08:00
core/cli: richer Tab completion for slash commands
Complete /todos actions and todo ids, session ids for /resume|/delete (from cache after /sessions), /rounds hints, and --flags. Pass prior tokens so multi-arg completion works.
This commit is contained in:
@@ -54,8 +54,10 @@ def build_completer(state: ReplState) -> Callable[[str, int], str | None]:
|
|||||||
options = filter_prefix(text, slash_commands())
|
options = filter_prefix(text, slash_commands())
|
||||||
else:
|
else:
|
||||||
head = buffer[:begidx].strip()
|
head = buffer[:begidx].strip()
|
||||||
command = head.split()[0] if head else ""
|
tokens = head.split() if head else []
|
||||||
options = complete_slash_args(state, command, text)
|
command = tokens[0] if tokens else ""
|
||||||
|
prior_args = tokens[1:] if len(tokens) > 1 else []
|
||||||
|
options = complete_slash_args(state, command, text, prior_args=prior_args)
|
||||||
if state_index < len(options):
|
if state_index < len(options):
|
||||||
return options[state_index]
|
return options[state_index]
|
||||||
return None
|
return None
|
||||||
|
|||||||
+136
-17
@@ -33,6 +33,19 @@ _COMPACT_PREVIEW = 400
|
|||||||
_ON_OFF_CHOICES = ("on", "off")
|
_ON_OFF_CHOICES = ("on", "off")
|
||||||
_YOLO_MODE_CHOICES = ("on", "off", "once")
|
_YOLO_MODE_CHOICES = ("on", "off", "once")
|
||||||
_EXPORT_FORMAT_CHOICES = ("md", "json")
|
_EXPORT_FORMAT_CHOICES = ("md", "json")
|
||||||
|
_TODO_ACTION_CHOICES = (
|
||||||
|
"list",
|
||||||
|
"show",
|
||||||
|
"ls",
|
||||||
|
"push",
|
||||||
|
"pop",
|
||||||
|
"done",
|
||||||
|
"cancel",
|
||||||
|
"pending",
|
||||||
|
"in_progress",
|
||||||
|
"clear",
|
||||||
|
)
|
||||||
|
_ROUNDS_CHOICES = ("8", "16", "32", "64", "128")
|
||||||
|
|
||||||
|
|
||||||
class HistoryLimitType(click.ParamType[int]):
|
class HistoryLimitType(click.ParamType[int]):
|
||||||
@@ -69,9 +82,9 @@ HELP_FOOTER = (
|
|||||||
"assistant/tool output is discarded but the user message stays (so /retry\n"
|
"assistant/tool output is discarded but the user message stays (so /retry\n"
|
||||||
"works after resume, not only via readline history). Auto-retry: 10s/20s/30s.\n"
|
"works after resume, not only via readline history). Auto-retry: 10s/20s/30s.\n"
|
||||||
"\n"
|
"\n"
|
||||||
"Tab completes slash commands and some arguments (provider, model, tools,\n"
|
"Tab completes slash commands and arguments (provider, model, session ids,\n"
|
||||||
"stream, verbose, yolo, export). Use --session ID or /resume to continue a prior\n"
|
"todos actions, on/off, yolo, export, flags). Use --session ID or /resume to\n"
|
||||||
"chat after restart.\n"
|
"continue a prior chat after restart.\n"
|
||||||
"\n"
|
"\n"
|
||||||
'Multiline: start a message with """ then end a later line with """.\n'
|
'Multiline: start a message with """ then end a later line with """.\n'
|
||||||
"Long prompts: /edit opens $EDITOR.\n"
|
"Long prompts: /edit opens $EDITOR.\n"
|
||||||
@@ -161,6 +174,80 @@ class ExportFormatParam(click.ParamType[str]):
|
|||||||
EXPORT_FORMAT = ExportFormatParam()
|
EXPORT_FORMAT = ExportFormatParam()
|
||||||
|
|
||||||
|
|
||||||
|
class TodosActionParam(click.ParamType[str]):
|
||||||
|
"""First token of ``/todos``: list|push|pop|done|…"""
|
||||||
|
|
||||||
|
name: str = "todos_action"
|
||||||
|
|
||||||
|
@override
|
||||||
|
def convert(self, value: Any, param: click.Parameter | None, ctx: click.Context | None) -> str:
|
||||||
|
del param, ctx
|
||||||
|
return str(value)
|
||||||
|
|
||||||
|
@override
|
||||||
|
def shell_complete(self, ctx: click.Context, param: click.Parameter, incomplete: str) -> list[CompletionItem]:
|
||||||
|
del ctx, param
|
||||||
|
return _filter_choices(incomplete, _TODO_ACTION_CHOICES)
|
||||||
|
|
||||||
|
|
||||||
|
TODOS_ACTION = TodosActionParam()
|
||||||
|
|
||||||
|
|
||||||
|
class SessionIdParam(click.ParamType[int]):
|
||||||
|
"""Session id for ``/resume`` / ``/delete`` (from ReplState cache + current)."""
|
||||||
|
|
||||||
|
name: str = "session_id"
|
||||||
|
|
||||||
|
@override
|
||||||
|
def convert(self, value: Any, param: click.Parameter | None, ctx: click.Context | None) -> int:
|
||||||
|
if isinstance(value, int):
|
||||||
|
return value
|
||||||
|
text = str(value).strip()
|
||||||
|
try:
|
||||||
|
return int(text, 10)
|
||||||
|
except ValueError:
|
||||||
|
self.fail("expected a session id (integer)", param, ctx)
|
||||||
|
|
||||||
|
@override
|
||||||
|
def shell_complete(self, ctx: click.Context, param: click.Parameter, incomplete: str) -> list[CompletionItem]:
|
||||||
|
del param
|
||||||
|
state = _repl_state(ctx)
|
||||||
|
if state is None:
|
||||||
|
return []
|
||||||
|
# Sync cache only — no async from readline Tab (no awaitlet greenlet).
|
||||||
|
return _filter_choices(incomplete, state.session_ids_for_complete())
|
||||||
|
|
||||||
|
|
||||||
|
SESSION_ID = SessionIdParam()
|
||||||
|
|
||||||
|
|
||||||
|
class RoundsParam(click.ParamType[int]):
|
||||||
|
"""Max tool-loop rounds (int >= 1); Tab suggests common values."""
|
||||||
|
|
||||||
|
name: str = "rounds"
|
||||||
|
|
||||||
|
@override
|
||||||
|
def convert(self, value: Any, param: click.Parameter | None, ctx: click.Context | None) -> int:
|
||||||
|
if isinstance(value, int):
|
||||||
|
n = value
|
||||||
|
else:
|
||||||
|
try:
|
||||||
|
n = int(str(value).strip(), 10)
|
||||||
|
except ValueError:
|
||||||
|
self.fail("expected an integer >= 1", param, ctx)
|
||||||
|
if n < 1:
|
||||||
|
self.fail("max_rounds must be >= 1", param, ctx)
|
||||||
|
return n
|
||||||
|
|
||||||
|
@override
|
||||||
|
def shell_complete(self, ctx: click.Context, param: click.Parameter, incomplete: str) -> list[CompletionItem]:
|
||||||
|
del ctx, param
|
||||||
|
return _filter_choices(incomplete, _ROUNDS_CHOICES)
|
||||||
|
|
||||||
|
|
||||||
|
ROUNDS = RoundsParam()
|
||||||
|
|
||||||
|
|
||||||
class ProviderNameParam(click.ParamType[str]):
|
class ProviderNameParam(click.ParamType[str]):
|
||||||
name: str = "provider"
|
name: str = "provider"
|
||||||
|
|
||||||
@@ -263,21 +350,55 @@ def slash_command_names() -> list[str]:
|
|||||||
return sorted(f"/{name}" for name in slash.list_commands(ctx))
|
return sorted(f"/{name}" for name in slash.list_commands(ctx))
|
||||||
|
|
||||||
|
|
||||||
def complete_slash_args(state: ReplState, command: str, incomplete: str) -> list[str]:
|
def complete_slash_args(
|
||||||
"""Tab-complete arguments for ``command`` (e.g. ``/stream``) from ParamTypes.
|
state: ReplState,
|
||||||
|
command: str,
|
||||||
|
incomplete: str,
|
||||||
|
*,
|
||||||
|
prior_args: Sequence[str] | None = None,
|
||||||
|
) -> list[str]:
|
||||||
|
"""Tab-complete arguments/options for ``command`` from ParamTypes.
|
||||||
|
|
||||||
Uses the first :class:`click.Argument` on the registered command whose type
|
*prior_args* are tokens already typed after the command name (so
|
||||||
implements :meth:`~click.ParamType.shell_complete` with candidates.
|
``/todos done t`` can complete todo ids). Options are completed when
|
||||||
|
*incomplete* starts with ``-``.
|
||||||
"""
|
"""
|
||||||
name = command.lstrip("/").lower()
|
name = command.lstrip("/").lower()
|
||||||
ctx = click.Context(slash, obj=state)
|
ctx = click.Context(slash, obj=state)
|
||||||
cmd = slash.get_command(ctx, name)
|
cmd = slash.get_command(ctx, name)
|
||||||
if cmd is None:
|
if cmd is None:
|
||||||
return []
|
return []
|
||||||
|
prior = list(prior_args or ())
|
||||||
|
|
||||||
with click.Context(cmd, info_name=name, parent=ctx, obj=state) as sub:
|
with click.Context(cmd, info_name=name, parent=ctx, obj=state) as sub:
|
||||||
for param in cmd.params:
|
# Flag / option completion (e.g. --persist, --full).
|
||||||
if not isinstance(param, click.Argument):
|
if incomplete.startswith("-"):
|
||||||
continue
|
flags = [
|
||||||
|
opt
|
||||||
|
for param in cmd.params
|
||||||
|
if isinstance(param, click.Option)
|
||||||
|
for opt in param.opts
|
||||||
|
if opt.startswith(incomplete)
|
||||||
|
]
|
||||||
|
if flags:
|
||||||
|
return sorted(set(flags))
|
||||||
|
|
||||||
|
# Contextual /todos second token: todo ids after done|cancel|pending|in_progress.
|
||||||
|
if name == "todos" and prior:
|
||||||
|
action = prior[0].lower()
|
||||||
|
if action in {"done", "cancel", "pending", "in_progress"} and len(prior) == 1:
|
||||||
|
ids = [item.id for item in state.todo_stack.all_items()]
|
||||||
|
return [i for i in ids if i.startswith(incomplete)]
|
||||||
|
|
||||||
|
# Positional arguments: complete the next unset argument.
|
||||||
|
arg_params = [p for p in cmd.params if isinstance(p, click.Argument)]
|
||||||
|
arg_index = len(prior)
|
||||||
|
if arg_index < len(arg_params):
|
||||||
|
param = arg_params[arg_index]
|
||||||
|
items = param.type.shell_complete(sub, param, incomplete)
|
||||||
|
if items:
|
||||||
|
return [item.value for item in items]
|
||||||
|
for param in arg_params:
|
||||||
items = param.type.shell_complete(sub, param, incomplete)
|
items = param.type.shell_complete(sub, param, incomplete)
|
||||||
if items:
|
if items:
|
||||||
return [item.value for item in items]
|
return [item.value for item in items]
|
||||||
@@ -429,6 +550,7 @@ def status_cmd(state: ReplState) -> None:
|
|||||||
def sessions_cmd(state: ReplState) -> None:
|
def sessions_cmd(state: ReplState) -> None:
|
||||||
"""List sessions for this workspace (newest first)."""
|
"""List sessions for this workspace (newest first)."""
|
||||||
sessions = _await(state.memory.list_sessions(workspace=state.workspace))
|
sessions = _await(state.memory.list_sessions(workspace=state.workspace))
|
||||||
|
state.remember_session_ids([s.sid for s in sessions])
|
||||||
if not sessions:
|
if not sessions:
|
||||||
click.echo(f"(no sessions for workspace {state.workspace})")
|
click.echo(f"(no sessions for workspace {state.workspace})")
|
||||||
return
|
return
|
||||||
@@ -466,7 +588,7 @@ def rename_cmd(state: ReplState, name: tuple[str, ...]) -> None:
|
|||||||
|
|
||||||
|
|
||||||
@slash.command("delete")
|
@slash.command("delete")
|
||||||
@click.argument("session_id", type=int, required=False)
|
@click.argument("session_id", type=SESSION_ID, required=False)
|
||||||
@click.pass_obj
|
@click.pass_obj
|
||||||
def delete_cmd(state: ReplState, session_id: int | None) -> None:
|
def delete_cmd(state: ReplState, session_id: int | None) -> None:
|
||||||
"""Hard-delete a session (confirm; current → new empty)."""
|
"""Hard-delete a session (confirm; current → new empty)."""
|
||||||
@@ -569,7 +691,7 @@ def export_cmd(state: ReplState, parts: tuple[str, ...]) -> None:
|
|||||||
|
|
||||||
|
|
||||||
@slash.command("resume")
|
@slash.command("resume")
|
||||||
@click.argument("session_id", type=int, required=False)
|
@click.argument("session_id", type=SESSION_ID, required=False)
|
||||||
@click.pass_obj
|
@click.pass_obj
|
||||||
def resume_cmd(state: ReplState, session_id: int | None) -> None:
|
def resume_cmd(state: ReplState, session_id: int | None) -> None:
|
||||||
"""Resume session id, or latest for this workspace if omitted."""
|
"""Resume session id, or latest for this workspace if omitted."""
|
||||||
@@ -868,16 +990,13 @@ def markdown_cmd(state: ReplState, enabled: bool | None) -> None: # noqa: FBT00
|
|||||||
|
|
||||||
|
|
||||||
@slash.command("rounds")
|
@slash.command("rounds")
|
||||||
@click.argument("n", required=False, type=int)
|
@click.argument("n", required=False, type=ROUNDS)
|
||||||
@click.pass_obj
|
@click.pass_obj
|
||||||
def rounds_cmd(state: ReplState, n: int | None) -> None:
|
def rounds_cmd(state: ReplState, n: int | None) -> None:
|
||||||
"""Show or set max tool-loop rounds."""
|
"""Show or set max tool-loop rounds."""
|
||||||
if n is None:
|
if n is None:
|
||||||
click.echo(f"max_rounds={state.max_rounds}")
|
click.echo(f"max_rounds={state.max_rounds}")
|
||||||
return
|
return
|
||||||
if n < 1:
|
|
||||||
msg = "max_rounds must be >= 1"
|
|
||||||
raise click.UsageError(msg)
|
|
||||||
state.max_rounds = n
|
state.max_rounds = n
|
||||||
state.agent.max_rounds = n
|
state.agent.max_rounds = n
|
||||||
click.echo(f"max_rounds={state.max_rounds}")
|
click.echo(f"max_rounds={state.max_rounds}")
|
||||||
@@ -946,7 +1065,7 @@ def history_cmd(
|
|||||||
|
|
||||||
|
|
||||||
@slash.command("todos")
|
@slash.command("todos")
|
||||||
@click.argument("action", required=False, type=str)
|
@click.argument("action", required=False, type=TODOS_ACTION)
|
||||||
@click.argument("rest", required=False, nargs=-1, type=str)
|
@click.argument("rest", required=False, nargs=-1, type=str)
|
||||||
@click.pass_obj
|
@click.pass_obj
|
||||||
def todos_cmd( # noqa: C901, PLR0911, PLR0912, PLR0915
|
def todos_cmd( # noqa: C901, PLR0911, PLR0912, PLR0915
|
||||||
|
|||||||
@@ -59,6 +59,8 @@ class ReplState:
|
|||||||
session_id: int | None = None
|
session_id: int | None = None
|
||||||
todo_stack: TodoStack = field(default_factory=TodoStack)
|
todo_stack: TodoStack = field(default_factory=TodoStack)
|
||||||
_todo_persist_tasks: set[object] = field(default_factory=set, init=False, repr=False)
|
_todo_persist_tasks: set[object] = field(default_factory=set, init=False, repr=False)
|
||||||
|
# Session ids for Tab complete (updated when listing/creating/resuming).
|
||||||
|
_session_id_cache: list[int] = field(default_factory=list, init=False, repr=False)
|
||||||
# Remote GET /models cache (per provider base).
|
# Remote GET /models cache (per provider base).
|
||||||
_remote_models: list[str] | None = field(default=None, init=False, repr=False)
|
_remote_models: list[str] | None = field(default=None, init=False, repr=False)
|
||||||
_remote_models_fetched_at: float | None = field(default=None, init=False, repr=False)
|
_remote_models_fetched_at: float | None = field(default=None, init=False, repr=False)
|
||||||
@@ -222,6 +224,22 @@ class ReplState:
|
|||||||
self._remote_models_key = self._models_cache_key()
|
self._remote_models_key = self._models_cache_key()
|
||||||
self._remote_models_error = None
|
self._remote_models_error = None
|
||||||
|
|
||||||
|
def remember_session_ids(self, ids: Sequence[int]) -> None:
|
||||||
|
"""Cache session ids for Tab completion (``/resume`` / ``/delete``)."""
|
||||||
|
self._session_id_cache = [int(i) for i in ids]
|
||||||
|
|
||||||
|
def session_ids_for_complete(self) -> list[str]:
|
||||||
|
"""String session ids for completers (cache + current session if any)."""
|
||||||
|
seen: set[int] = set()
|
||||||
|
out: list[str] = []
|
||||||
|
for sid in self._session_id_cache:
|
||||||
|
if sid not in seen:
|
||||||
|
seen.add(sid)
|
||||||
|
out.append(str(sid))
|
||||||
|
if self.session_id is not None and self.session_id not in seen:
|
||||||
|
out.insert(0, str(self.session_id))
|
||||||
|
return out
|
||||||
|
|
||||||
def cached_remote_models(self) -> list[str] | None:
|
def cached_remote_models(self) -> list[str] | None:
|
||||||
"""Return cached remote ids if still valid for the current provider."""
|
"""Return cached remote ids if still valid for the current provider."""
|
||||||
if self._remote_models is None or self._remote_models_fetched_at is None:
|
if self._remote_models is None or self._remote_models_fetched_at is None:
|
||||||
@@ -423,6 +441,8 @@ class ReplState:
|
|||||||
model=self.model,
|
model=self.model,
|
||||||
)
|
)
|
||||||
self.session_id = session.sid
|
self.session_id = session.sid
|
||||||
|
if session.sid not in self._session_id_cache:
|
||||||
|
self._session_id_cache.insert(0, session.sid)
|
||||||
self.todo_stack = TodoStack()
|
self.todo_stack = TodoStack()
|
||||||
self.agent = self._make_agent()
|
self.agent = self._make_agent()
|
||||||
self._bind_todo_tools()
|
self._bind_todo_tools()
|
||||||
@@ -476,6 +496,8 @@ class ReplState:
|
|||||||
raise ValueError(msg) from exc
|
raise ValueError(msg) from exc
|
||||||
|
|
||||||
self.session_id = session_id
|
self.session_id = session_id
|
||||||
|
if session_id not in self._session_id_cache:
|
||||||
|
self._session_id_cache.insert(0, session_id)
|
||||||
if self.apply_session_llm(row):
|
if self.apply_session_llm(row):
|
||||||
self.rebuild_client()
|
self.rebuild_client()
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -183,3 +183,20 @@ def test_complete_slash_args_from_registry(tmp_path: object) -> None:
|
|||||||
assert complete_slash_args(state, "/help", "st") == ["status", "stream"]
|
assert complete_slash_args(state, "/help", "st") == ["status", "stream"]
|
||||||
assert complete_slash_args(state, "/yolo", "") == ["on", "off", "once"]
|
assert complete_slash_args(state, "/yolo", "") == ["on", "off", "once"]
|
||||||
assert complete_slash_args(state, "/yolo", "o") == ["on", "off", "once"]
|
assert complete_slash_args(state, "/yolo", "o") == ["on", "off", "once"]
|
||||||
|
assert "push" in complete_slash_args(state, "/todos", "p")
|
||||||
|
assert "pop" in complete_slash_args(state, "/todos", "p")
|
||||||
|
assert "--persist" in complete_slash_args(state, "/model", "--p")
|
||||||
|
assert "--full" in complete_slash_args(state, "/history", "--f")
|
||||||
|
assert "32" in complete_slash_args(state, "/rounds", "3")
|
||||||
|
|
||||||
|
from plyngent.agent.todo_stack import TodoStack
|
||||||
|
|
||||||
|
state.todo_stack = TodoStack()
|
||||||
|
_ = state.todo_stack.push_group(["alpha", "beta"])
|
||||||
|
ids = complete_slash_args(state, "/todos", "t", prior_args=["done"])
|
||||||
|
assert "t1" in ids
|
||||||
|
assert "t2" in ids
|
||||||
|
|
||||||
|
state.remember_session_ids([3, 12, 40])
|
||||||
|
assert complete_slash_args(state, "/resume", "1") == ["12"]
|
||||||
|
assert complete_slash_args(state, "/delete", "4") == ["40"]
|
||||||
|
|||||||
Reference in New Issue
Block a user