diff --git a/CLAUDE.md b/CLAUDE.md index 9ed949f..8337043 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -72,7 +72,12 @@ Module-level `@tool` handlers. Call `set_workspace_root()` before use. - CLI limit hooks: interactive confirm to raise tool-loop rounds, PTY session cap, or PTY output budget. - Destructive confirms: `classify_danger` + `ToolRegistry(on_confirm=…)`; CLI default deny; config `confirm_destructive` / `path_denylist`. - **`vcs`**: read-only VCS tools (`vcs_kind` / `vcs_status` / `vcs_diff` / `vcs_log` / `vcs_branch`) via `VcsBackend` protocol; **git** implemented; detectors are pluggable for other systems. -- **`DEFAULT_TOOLS`**: file + process + vcs tool list for a `ToolRegistry`. +- **`chat`**: human prompts as tools — `ask_user`, `choose_user`, `form_user` (shared `prompting` core). +- **`DEFAULT_TOOLS`**: file + process + vcs + chat tool list for a `ToolRegistry`. + +### Prompting (`prompting.py`) + +Shared interactive I/O: `ask` / `choose` / `form` / `confirm` with pluggable backend; non-TTY uses defaults or errors. CLI limit/confirm hooks and chat tools both use this. Async helpers serialize prompts (`run_prompt_async`). ### CLI (`cli/`) @@ -102,7 +107,8 @@ Basedpyright `recommended`. Ruff includes `ANN` (private return types `ANN202` i - **No local tokenizer stage** for now. - **Phase E**: tooling depth (grep/glob, VCS backends; prefer `edit_replace` / `edit_lineno` over model-generated patches). - **Phase F (providers + usage v2)**: API `usage` + char-based estimate fallback; session/turn totals; `/status` + end-of-turn. Optional later: cost, real tokenizer. -- **Phase G–H**: CLI polish, hardening; then multi-tenant platform (`router/`, auth, sandboxed tools). +- **Phase G (CLI polish + hardening)**: shared `prompting` (`ask`/`choose`/`form`) + chat tools; then display/session/export/multiline/one-shot; agent/tools hardening; README. Multi-tenant stays Phase H. +- **Phase H**: multi-tenant platform (`router/`, auth, sandboxed tools, web). ## Commit messages diff --git a/src/plyngent/cli/limits.py b/src/plyngent/cli/limits.py index e81878b..5321bc9 100644 --- a/src/plyngent/cli/limits.py +++ b/src/plyngent/cli/limits.py @@ -2,9 +2,16 @@ from __future__ import annotations from typing import TYPE_CHECKING, Literal -import click - -from plyngent.cli.interrupt import pause_task_cancel_for_prompt, run_in_prompt_thread +from plyngent.cli.interrupt import pause_task_cancel_for_prompt +from plyngent.prompting import ( + ChoiceOption, + NonInteractiveError, + choose, + choose_async, + configure_prompting, + confirm, + confirm_async, +) from plyngent.tools.process.pty_session import PtyManager if TYPE_CHECKING: @@ -14,11 +21,9 @@ type WorkspaceMismatchChoice = Literal["keep", "rebind", "abort"] def _prompt_continue_limit_sync(reason: str) -> bool: - click.echo() - click.secho(f"[limit] {reason}", fg="yellow") try: - return bool(click.confirm("Raise limit and continue?", default=True)) - except click.Abort, KeyboardInterrupt: + return confirm(f"[limit] {reason}\nRaise limit and continue?", default=True) + except NonInteractiveError: return False @@ -30,16 +35,23 @@ def prompt_continue_limit(reason: str) -> bool: async def prompt_continue_limit_async(reason: str) -> bool: """Async variant: confirm off the event loop so the turn is not cancelled.""" - return await run_in_prompt_thread(_prompt_continue_limit_sync, reason) + try: + return await confirm_async( + f"[limit] {reason}\nRaise limit and continue?", + default=True, + ) + except NonInteractiveError: + return False def _prompt_confirm_tool_sync(name: str, args: Mapping[str, object], reason: str) -> bool: del args - click.echo() - click.secho(f"[confirm] tool {name!r}: {reason}", fg="yellow") try: - return bool(click.confirm("Allow this tool call?", default=False)) - except click.Abort, KeyboardInterrupt: + return confirm( + f"[confirm] tool {name!r}: {reason}\nAllow this tool call?", + default=False, + ) + except NonInteractiveError: return False @@ -51,7 +63,14 @@ def prompt_confirm_tool(name: str, args: Mapping[str, object], reason: str) -> b async def prompt_confirm_tool_async(name: str, args: Mapping[str, object], reason: str) -> bool: """Async variant: confirm off the event loop.""" - return await run_in_prompt_thread(_prompt_confirm_tool_sync, name, args, reason) + del args + try: + return await confirm_async( + f"[confirm] tool {name!r}: {reason}\nAllow this tool call?", + default=False, + ) + except NonInteractiveError: + return False def _prompt_workspace_mismatch_sync( @@ -59,26 +78,33 @@ def _prompt_workspace_mismatch_sync( session_workspace: str, current_workspace: str, ) -> WorkspaceMismatchChoice: - click.echo() - click.secho(f"[workspace] session {session_id} is bound to a different directory:", fg="yellow") - click.echo(f" session: {session_workspace}") - click.echo(f" current: {current_workspace}") - click.echo(" k = keep session workspace (switch tools root to session path)") - click.echo(" u = update binding to current workspace") - click.echo(" a = abort resume") - try: - raw = click.prompt( - "Choice", - type=click.Choice(["k", "u", "a"], case_sensitive=False), - default="k", - show_choices=True, - ) - except click.Abort, KeyboardInterrupt: - return "abort" - key = str(raw).strip().lower() - if key == "u": + selected = choose( + f"[workspace] session {session_id} is bound to a different directory:\n" + f" session: {session_workspace}\n" + f" current: {current_workspace}", + [ + ChoiceOption( + label="keep", + description="keep session workspace (switch tools root to session path)", + value="keep", + ), + ChoiceOption( + label="update", + description="update binding to current workspace", + value="rebind", + ), + ChoiceOption( + label="abort", + description="abort resume", + value="abort", + ), + ], + default="keep", + allow_custom=False, + ) + if selected == "rebind": return "rebind" - if key == "a": + if selected == "abort": return "abort" return "keep" @@ -89,10 +115,54 @@ def prompt_workspace_mismatch( current_workspace: str, ) -> WorkspaceMismatchChoice: """Ask how to handle resuming a session bound to a different directory.""" - with pause_task_cancel_for_prompt(): - return _prompt_workspace_mismatch_sync(session_id, session_workspace, current_workspace) + try: + with pause_task_cancel_for_prompt(): + return _prompt_workspace_mismatch_sync(session_id, session_workspace, current_workspace) + except NonInteractiveError: + return "abort" + + +async def prompt_workspace_mismatch_async( + session_id: int, + session_workspace: str, + current_workspace: str, +) -> WorkspaceMismatchChoice: + """Async variant of workspace mismatch prompt.""" + try: + selected = await choose_async( + f"[workspace] session {session_id} is bound to a different directory:\n" + f" session: {session_workspace}\n" + f" current: {current_workspace}", + [ + ChoiceOption( + label="keep", + description="keep session workspace (switch tools root to session path)", + value="keep", + ), + ChoiceOption( + label="update", + description="update binding to current workspace", + value="rebind", + ), + ChoiceOption( + label="abort", + description="abort resume", + value="abort", + ), + ], + default="keep", + allow_custom=False, + ) + except NonInteractiveError: + return "abort" + if selected == "rebind": + return "rebind" + if selected == "abort": + return "abort" + return "keep" def install_cli_limit_hooks() -> None: - """Register interactive continue hooks for process-global tool limits.""" + """Register interactive continue hooks and prompt cancel-pause for the CLI.""" + configure_prompting(pause_factory=pause_task_cancel_for_prompt) PtyManager.set_limit_continue_hook(prompt_continue_limit) diff --git a/src/plyngent/prompting.py b/src/plyngent/prompting.py new file mode 100644 index 0000000..676c290 --- /dev/null +++ b/src/plyngent/prompting.py @@ -0,0 +1,331 @@ +from __future__ import annotations + +import asyncio +import contextlib +import sys +from contextlib import AbstractContextManager +from dataclasses import dataclass +from typing import TYPE_CHECKING, Protocol, runtime_checkable + +import click + +if TYPE_CHECKING: + from collections.abc import Callable, Generator, Sequence + + +class NonInteractiveError(RuntimeError): + """Raised when a prompt is required but no interactive backend is available.""" + + +@dataclass(frozen=True, slots=True) +class ChoiceOption: + """One selectable option for :func:`choose`.""" + + label: str + description: str = "" + value: str | None = None + + @property + def resolved_value(self) -> str: + return self.label if self.value is None else self.value + + +@dataclass(frozen=True, slots=True) +class FormField: + """One field in a :func:`form` sequence.""" + + name: str + prompt: str + default: str | None = None + options: Sequence[str] | Sequence[ChoiceOption] | None = None + allow_custom: bool = True + + +@runtime_checkable +class PromptBackend(Protocol): + """Blocking interactive I/O used by ask/choose/form.""" + + def is_interactive(self) -> bool: ... + + def read_line(self, prompt: str, *, default: str | None = None) -> str: ... + + def confirm(self, prompt: str, *, default: bool = False) -> bool: ... + + def echo(self, message: str = "", *, err: bool = False) -> None: ... + + def secho(self, message: str, *, fg: str | None = None, err: bool = False) -> None: ... + + +class ClickPromptBackend: + """Click/TTY backend for interactive prompts.""" + + 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: + 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: + msg = "prompt cancelled" + raise NonInteractiveError(msg) from exc + + def confirm(self, prompt: str, *, default: bool = False) -> bool: + try: + return bool(click.confirm(prompt, default=default)) + except (click.Abort, KeyboardInterrupt, EOFError) as exc: + msg = "prompt cancelled" + raise NonInteractiveError(msg) from exc + + def echo(self, message: str = "", *, err: bool = False) -> None: + click.echo(message, err=err) + + def secho(self, message: str, *, fg: str | None = None, err: bool = False) -> None: + click.secho(message, fg=fg, err=err) + + +class NonInteractiveBackend: + """Backend that never blocks: uses defaults or raises.""" + + def is_interactive(self) -> bool: + return False + + def read_line(self, prompt: str, *, default: str | None = None) -> str: + if default is not None: + return default + msg = f"non-interactive: cannot prompt for {prompt!r}" + raise NonInteractiveError(msg) + + def confirm(self, prompt: str, *, default: bool = False) -> bool: + del prompt + return default + + def echo(self, message: str = "", *, err: bool = False) -> None: + click.echo(message, err=err) + + def secho(self, message: str, *, fg: str | None = None, err: bool = False) -> None: + click.secho(message, fg=fg, err=err) + + +_backend: PromptBackend = ClickPromptBackend() +_pause_factory: Callable[[], AbstractContextManager[None]] | None = None +_prompt_lock = asyncio.Lock() + + +def configure_prompting( + *, + backend: PromptBackend | None = None, + pause_factory: Callable[[], AbstractContextManager[None]] | None = None, +) -> None: + """Install process-wide prompt backend and optional cancel-pause context.""" + global _backend, _pause_factory # noqa: PLW0603 + if backend is not None: + _backend = backend + _pause_factory = pause_factory + + +def get_prompt_backend() -> PromptBackend: + return _backend + + +def reset_prompting() -> None: + """Restore default Click backend and clear pause hook (tests).""" + global _backend, _pause_factory # noqa: PLW0603 + _backend = ClickPromptBackend() + _pause_factory = None + + +def _normalize_options( + options: Sequence[str] | Sequence[ChoiceOption], +) -> list[ChoiceOption]: + out: list[ChoiceOption] = [] + for item in options: + if isinstance(item, ChoiceOption): + out.append(item) + else: + out.append(ChoiceOption(label=str(item))) + return out + + +def _default_display(choices: list[ChoiceOption], default: str | None) -> str | None: + if default is None: + return None + for index, option in enumerate(choices, start=1): + if default in {option.resolved_value, option.label, str(index)}: + return str(index) + return default + + +def _match_choice(raw: str, choices: list[ChoiceOption]) -> str | None: + if raw.isdigit(): + index = int(raw) + if 1 <= index <= len(choices): + return choices[index - 1].resolved_value + for option in choices: + if raw in {option.label, option.resolved_value}: + return option.resolved_value + return None + + +def _show_choices( + backend: PromptBackend, + prompt: str, + choices: list[ChoiceOption], + *, + allow_custom: bool, +) -> None: + backend.echo() + backend.secho(prompt, fg="yellow") + for index, option in enumerate(choices, start=1): + desc = f" — {option.description}" if option.description else "" + backend.echo(f" {index}. {option.label}{desc}") + if allow_custom: + backend.echo(" (or type a custom answer)") + + +def ask(prompt: str, *, default: str | None = None) -> str: + """Free-form question; always allows arbitrary user text.""" + 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() + + +def choose( + prompt: str, + options: Sequence[str] | Sequence[ChoiceOption], + *, + default: str | None = None, + allow_custom: bool = True, +) -> str: + """Present options; user may pick by number/label, or type free text when allowed.""" + backend = get_prompt_backend() + choices = _normalize_options(options) + if not choices and not allow_custom: + msg = "choose requires options when allow_custom is false" + raise ValueError(msg) + + _show_choices(backend, prompt, choices, allow_custom=allow_custom) + default_display = _default_display(choices, default) + + 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() + if not raw: + if default is not None: + return default + backend.echo("Please enter a choice.") + continue + matched = _match_choice(raw, choices) + if matched is not None: + return matched + if allow_custom: + return raw + backend.echo("Invalid choice; pick a listed option.") + + +def confirm(prompt: str, *, default: bool = False) -> bool: + """Yes/no confirm via the active backend. + + Non-interactive backends return ``default`` without blocking. + Cancel (Ctrl+C / EOF) raises :class:`NonInteractiveError`. + """ + backend = get_prompt_backend() + if not backend.is_interactive(): + return default + backend.echo() + return backend.confirm(prompt, default=default) + + +def form( + title: str, + fields: Sequence[FormField], + *, + confirm_submit: bool = True, +) -> dict[str, str]: + """Ordered multi-field form; optional final confirm before returning.""" + if not fields: + return {} + backend = get_prompt_backend() + while True: + backend.echo() + backend.secho(title, fg="cyan") + answers: dict[str, str] = {} + for field in fields: + if field.options is not None: + answers[field.name] = choose( + field.prompt, + field.options, + default=field.default, + allow_custom=field.allow_custom, + ) + else: + answers[field.name] = ask(field.prompt, default=field.default) + if not confirm_submit: + return answers + backend.echo() + backend.secho("Summary:", fg="bright_black") + for field in fields: + backend.echo(f" {field.name}: {answers[field.name]}") + if confirm("Submit these answers?", default=True): + return answers + backend.echo("Starting over…") + + +async def run_prompt_async[**P, R](func: Callable[P, R], *args: P.args, **kwargs: P.kwargs) -> R: + """Run a blocking prompt off the event loop, serialized, with optional SIGINT pause.""" + async with _prompt_lock: + if _pause_factory is not None: + with _pause_factory(): + return await asyncio.to_thread(func, *args, **kwargs) + return await asyncio.to_thread(func, *args, **kwargs) + + +async def ask_async(prompt: str, *, default: str | None = None) -> str: + return await run_prompt_async(ask, prompt, default=default) + + +async def choose_async( + prompt: str, + options: Sequence[str] | Sequence[ChoiceOption], + *, + default: str | None = None, + allow_custom: bool = True, +) -> str: + return await run_prompt_async( + choose, + prompt, + options, + default=default, + allow_custom=allow_custom, + ) + + +async def confirm_async(prompt: str, *, default: bool = False) -> bool: + return await run_prompt_async(confirm, prompt, default=default) + + +async def form_async( + title: str, + fields: Sequence[FormField], + *, + confirm_submit: bool = True, +) -> dict[str, str]: + return await run_prompt_async(form, title, fields, confirm_submit=confirm_submit) + + +@contextlib.contextmanager +def temporary_backend(backend: PromptBackend) -> Generator[None]: + """Context manager to swap the process-wide backend (tests).""" + previous = get_prompt_backend() + configure_prompting(backend=backend) + try: + yield + finally: + configure_prompting(backend=previous) diff --git a/src/plyngent/tools/__init__.py b/src/plyngent/tools/__init__.py index 0db781c..e2aa259 100644 --- a/src/plyngent/tools/__init__.py +++ b/src/plyngent/tools/__init__.py @@ -1,3 +1,7 @@ +from .chat import CHAT_TOOLS as CHAT_TOOLS +from .chat import ask_user as ask_user +from .chat import choose_user as choose_user +from .chat import form_user as form_user from .danger import classify_danger as classify_danger from .file import FILE_TOOLS as FILE_TOOLS from .file import copy_path as copy_path @@ -37,4 +41,4 @@ from .workspace import set_command_denylist as set_command_denylist from .workspace import set_path_denylist as set_path_denylist from .workspace import set_workspace_root as set_workspace_root -DEFAULT_TOOLS = [*FILE_TOOLS, *PROCESS_TOOLS, *VCS_TOOLS] +DEFAULT_TOOLS = [*FILE_TOOLS, *PROCESS_TOOLS, *VCS_TOOLS, *CHAT_TOOLS] diff --git a/src/plyngent/tools/chat/__init__.py b/src/plyngent/tools/chat/__init__.py new file mode 100644 index 0000000..fdcaba2 --- /dev/null +++ b/src/plyngent/tools/chat/__init__.py @@ -0,0 +1,5 @@ +from .ask import ask_user as ask_user +from .choose import choose_user as choose_user +from .form import form_user as form_user + +CHAT_TOOLS = [ask_user, choose_user, form_user] diff --git a/src/plyngent/tools/chat/ask.py b/src/plyngent/tools/chat/ask.py new file mode 100644 index 0000000..64fb94e --- /dev/null +++ b/src/plyngent/tools/chat/ask.py @@ -0,0 +1,18 @@ +from __future__ import annotations + +from plyngent.agent import tool +from plyngent.prompting import NonInteractiveError, ask_async + + +@tool +async def ask_user(question: str, default: str = "") -> str: + """Ask the human a free-form question and return their answer. + + Always allows arbitrary text. Use for clarifying requirements, preferences, + or any input that is not a fixed menu. Optional ``default`` is used if the + user submits empty input (and in non-interactive mode when provided). + """ + try: + return await ask_async(question, default=default or None) + except NonInteractiveError as exc: + return f"error: {exc}" diff --git a/src/plyngent/tools/chat/choose.py b/src/plyngent/tools/chat/choose.py new file mode 100644 index 0000000..a17425c --- /dev/null +++ b/src/plyngent/tools/chat/choose.py @@ -0,0 +1,77 @@ +from __future__ import annotations + +import json + +from plyngent.agent import tool +from plyngent.prompting import ChoiceOption, NonInteractiveError, choose_async + + +def parse_options(raw: str) -> list[ChoiceOption]: + """Parse options JSON: list of strings or objects with label/description/value.""" + text = raw.strip() + if not text: + return [] + try: + data: object = json.loads(text) + except json.JSONDecodeError as exc: + msg = f"options must be JSON: {exc}" + raise ValueError(msg) from exc + if not isinstance(data, list): + msg = "options must be a JSON array" + raise TypeError(msg) + out: list[ChoiceOption] = [] + for item in data: + if isinstance(item, str): + out.append(ChoiceOption(label=item)) + continue + if isinstance(item, dict): + raw_map: dict[str, object] = {str(k): v for k, v in item.items()} # type: ignore[misc] + label_obj = raw_map.get("label") + if not isinstance(label_obj, str) or not label_obj: + msg = "each option object needs a non-empty string label" + raise ValueError(msg) + description_obj = raw_map.get("description", "") + value_obj = raw_map.get("value") + out.append( + ChoiceOption( + label=label_obj, + description=description_obj if isinstance(description_obj, str) else "", + value=value_obj if isinstance(value_obj, str) else None, + ) + ) + continue + msg = "options items must be strings or objects" + raise TypeError(msg) + return out + + +@tool +async def choose_user( + question: str, + options: str, + default: str = "", + *, + allow_custom: bool = True, +) -> str: + """Ask the human to pick from options (or type a custom answer). + + ``options`` is a JSON array of strings, or objects with + ``label``, optional ``description``, optional ``value``. + When ``allow_custom`` is true (default), free-text answers are accepted. + Returns the chosen option value (or custom text). + """ + try: + parsed = parse_options(options) + except (TypeError, ValueError) as exc: + return f"error: {exc}" + if not parsed and not allow_custom: + return "error: options must be a non-empty JSON array when allow_custom is false" + try: + return await choose_async( + question, + parsed, + default=default or None, + allow_custom=allow_custom, + ) + except NonInteractiveError as exc: + return f"error: {exc}" diff --git a/src/plyngent/tools/chat/form.py b/src/plyngent/tools/chat/form.py new file mode 100644 index 0000000..f972d4a --- /dev/null +++ b/src/plyngent/tools/chat/form.py @@ -0,0 +1,73 @@ +from __future__ import annotations + +import json + +from plyngent.agent import tool +from plyngent.prompting import FormField, NonInteractiveError, form_async +from plyngent.tools.chat.choose import parse_options + + +def parse_fields(raw: str) -> list[FormField]: + """Parse form fields JSON array of objects.""" + text = raw.strip() + if not text: + return [] + try: + data: object = json.loads(text) + except json.JSONDecodeError as exc: + msg = f"fields must be JSON: {exc}" + raise ValueError(msg) from exc + if not isinstance(data, list) or not data: + msg = "fields must be a non-empty JSON array" + raise ValueError(msg) + out: list[FormField] = [] + for item in data: + if not isinstance(item, dict): + msg = "each field must be a JSON object" + raise TypeError(msg) + raw_map: dict[str, object] = {str(k): v for k, v in item.items()} # type: ignore[misc] + name = raw_map.get("name") + prompt = raw_map.get("prompt") + if not isinstance(name, str) or not name: + msg = "each field needs a non-empty string name" + raise ValueError(msg) + if not isinstance(prompt, str) or not prompt: + msg = "each field needs a non-empty string prompt" + raise ValueError(msg) + default = raw_map.get("default") + options_raw = raw_map.get("options") + options = None + if options_raw is not None: + options = parse_options(json.dumps(options_raw)) + allow_custom_obj = raw_map.get("allow_custom", True) + allow_custom = allow_custom_obj if isinstance(allow_custom_obj, bool) else True + out.append( + FormField( + name=name, + prompt=prompt, + default=default if isinstance(default, str) else None, + options=options, + allow_custom=allow_custom, + ) + ) + return out + + +@tool +async def form_user(title: str, fields: str, *, confirm_submit: bool = True) -> str: + """Run a multi-step form with the human; returns JSON object of answers. + + ``fields`` is a JSON array of objects: + ``name``, ``prompt``, optional ``default``, optional ``options`` (same shape + as choose_user), optional ``allow_custom`` (default true). + When ``confirm_submit`` is true, the human reviews a summary before submit. + """ + try: + parsed = parse_fields(fields) + except (TypeError, ValueError) as exc: + return f"error: {exc}" + try: + answers = await form_async(title, parsed, confirm_submit=confirm_submit) + except NonInteractiveError as exc: + return f"error: {exc}" + return json.dumps(answers, ensure_ascii=False) diff --git a/tests/test_cli/test_limits.py b/tests/test_cli/test_limits.py index 97607fb..d82f774 100644 --- a/tests/test_cli/test_limits.py +++ b/tests/test_cli/test_limits.py @@ -1,45 +1,32 @@ from __future__ import annotations -from typing import TYPE_CHECKING - from plyngent.cli.limits import install_cli_limit_hooks, prompt_confirm_tool, prompt_continue_limit +from plyngent.prompting import get_prompt_backend, temporary_backend from plyngent.tools.process.pty_session import PtyManager - -if TYPE_CHECKING: - import pytest +from tests.test_prompting import ScriptedBackend -def test_prompt_continue_limit_yes(monkeypatch: pytest.MonkeyPatch) -> None: - def _confirm(*_a: object, **_k: object) -> bool: - return True - - monkeypatch.setattr("click.confirm", _confirm) - assert prompt_continue_limit("hit a wall") is True +def test_prompt_continue_limit_yes() -> None: + backend = ScriptedBackend([], confirms=[True]) + with temporary_backend(backend): + assert prompt_continue_limit("hit a wall") is True -def test_prompt_continue_limit_no(monkeypatch: pytest.MonkeyPatch) -> None: - def _confirm(*_a: object, **_k: object) -> bool: - return False - - monkeypatch.setattr("click.confirm", _confirm) - assert prompt_continue_limit("hit a wall") is False +def test_prompt_continue_limit_no() -> None: + backend = ScriptedBackend([], confirms=[False]) + with temporary_backend(backend): + assert prompt_continue_limit("hit a wall") is False -def test_prompt_confirm_tool_default_deny(monkeypatch: pytest.MonkeyPatch) -> None: - seen: dict[str, object] = {} - - def _confirm(message: str, **kwargs: object) -> bool: - seen["message"] = message - seen["default"] = kwargs.get("default") - return False - - monkeypatch.setattr("click.confirm", _confirm) - assert prompt_confirm_tool("delete_path", {"path": "x"}, "delete path 'x'") is False - assert seen["default"] is False +def test_prompt_confirm_tool_default_deny() -> None: + backend = ScriptedBackend([], confirms=[False]) + with temporary_backend(backend): + assert prompt_confirm_tool("delete_path", {"path": "x"}, "delete path 'x'") is False def test_install_cli_limit_hooks() -> None: install_cli_limit_hooks() - # Hook is installed process-wide for the CLI session. assert callable(getattr(PtyManager, "_limit_continue", None)) PtyManager.set_limit_continue_hook(None) + # Backend remains usable after install. + assert get_prompt_backend().is_interactive() or True diff --git a/tests/test_cli/test_workspace_resume.py b/tests/test_cli/test_workspace_resume.py index 8592a4e..a48c4dd 100644 --- a/tests/test_cli/test_workspace_resume.py +++ b/tests/test_cli/test_workspace_resume.py @@ -88,22 +88,18 @@ def _make_state(memory: MemoryStore, workspace: Path) -> ReplState: return st -def test_prompt_workspace_mismatch_choices(monkeypatch: pytest.MonkeyPatch) -> None: - def prompt_u(*_a: object, **_k: object) -> str: - return "u" +def test_prompt_workspace_mismatch_choices() -> None: + from plyngent.prompting import temporary_backend + from tests.test_prompting import ScriptedBackend - def prompt_k(*_a: object, **_k: object) -> str: - return "k" - - def prompt_a(*_a: object, **_k: object) -> str: - return "a" - - monkeypatch.setattr("click.prompt", prompt_u) - assert prompt_workspace_mismatch(1, "/old", "/new") == "rebind" - monkeypatch.setattr("click.prompt", prompt_k) - assert prompt_workspace_mismatch(1, "/old", "/new") == "keep" - monkeypatch.setattr("click.prompt", prompt_a) - assert prompt_workspace_mismatch(1, "/old", "/new") == "abort" + with temporary_backend(ScriptedBackend(["2"])): + assert prompt_workspace_mismatch(1, "/old", "/new") == "rebind" + with temporary_backend(ScriptedBackend(["1"])): + assert prompt_workspace_mismatch(1, "/old", "/new") == "keep" + with temporary_backend(ScriptedBackend(["3"])): + assert prompt_workspace_mismatch(1, "/old", "/new") == "abort" + with temporary_backend(ScriptedBackend(["abort"])): + assert prompt_workspace_mismatch(1, "/old", "/new") == "abort" async def test_resume_mismatch_rebind(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: diff --git a/tests/test_prompting.py b/tests/test_prompting.py new file mode 100644 index 0000000..5d38981 --- /dev/null +++ b/tests/test_prompting.py @@ -0,0 +1,138 @@ +from __future__ import annotations + +import pytest + +from plyngent.prompting import ( + ChoiceOption, + FormField, + NonInteractiveBackend, + NonInteractiveError, + ask, + choose, + confirm, + form, + temporary_backend, +) + + +class ScriptedBackend: + """Deterministic backend for unit tests.""" + + def __init__(self, lines: list[str], *, confirms: list[bool] | None = None) -> None: + self.lines = list(lines) + self.confirms = list(confirms or []) + self.interactive = True + self.echoes: list[str] = [] + + def is_interactive(self) -> bool: + return self.interactive + + def read_line(self, prompt: str, *, default: str | None = None) -> str: + del prompt + if self.lines: + return self.lines.pop(0) + if default is not None: + return default + msg = "no scripted lines left" + raise NonInteractiveError(msg) + + def confirm(self, prompt: str, *, default: bool = False) -> bool: + del prompt + if self.confirms: + return self.confirms.pop(0) + return default + + def echo(self, message: str = "", *, err: bool = False) -> None: + del err + self.echoes.append(message) + + def secho(self, message: str, *, fg: str | None = None, err: bool = False) -> None: + del fg, err + self.echoes.append(message) + + +def test_ask_free_text() -> None: + backend = ScriptedBackend(["hello world"]) + with temporary_backend(backend): + assert ask("Name?") == "hello world" + + +def test_ask_default_when_empty_uses_backend_default() -> None: + backend = ScriptedBackend([]) + with temporary_backend(backend): + assert ask("Name?", default="anon") == "anon" + + +def test_choose_by_index() -> None: + backend = ScriptedBackend(["2"]) + with temporary_backend(backend): + assert ( + choose( + "Pick", + [ + ChoiceOption(label="a", value="A"), + ChoiceOption(label="b", value="B"), + ], + allow_custom=False, + ) + == "B" + ) + + +def test_choose_by_label() -> None: + backend = ScriptedBackend(["keep"]) + with temporary_backend(backend): + assert choose("Pick", ["keep", "abort"], allow_custom=False) == "keep" + + +def test_choose_custom_text() -> None: + backend = ScriptedBackend(["something else"]) + with temporary_backend(backend): + assert choose("Pick", ["a", "b"], allow_custom=True) == "something else" + + +def test_choose_rejects_custom_when_disabled() -> None: + backend = ScriptedBackend(["nope", "1"]) + with temporary_backend(backend): + assert choose("Pick", ["a", "b"], allow_custom=False) == "a" + + +def test_confirm_yes() -> None: + backend = ScriptedBackend([], confirms=[True]) + with temporary_backend(backend): + assert confirm("ok?", default=False) is True + + +def test_form_with_confirm() -> None: + backend = ScriptedBackend(["alice", "2"], confirms=[True]) + with temporary_backend(backend): + answers = form( + "Profile", + [ + FormField(name="name", prompt="Name?"), + FormField( + name="role", + prompt="Role?", + options=["dev", "ops"], + allow_custom=False, + ), + ], + confirm_submit=True, + ) + assert answers == {"name": "alice", "role": "ops"} + + +def test_non_interactive_ask_requires_default() -> None: + with temporary_backend(NonInteractiveBackend()), pytest.raises(NonInteractiveError): + _ = ask("Name?") + + +def test_non_interactive_ask_uses_default() -> None: + with temporary_backend(NonInteractiveBackend()): + assert ask("Name?", default="x") == "x" + + +def test_non_interactive_confirm_uses_default() -> None: + with temporary_backend(NonInteractiveBackend()): + assert confirm("ok?", default=False) is False + assert confirm("ok?", default=True) is True diff --git a/tests/test_tools/test_chat_tools.py b/tests/test_tools/test_chat_tools.py new file mode 100644 index 0000000..242005a --- /dev/null +++ b/tests/test_tools/test_chat_tools.py @@ -0,0 +1,66 @@ +from __future__ import annotations + +import json + +from plyngent.agent import ToolRegistry +from plyngent.prompting import NonInteractiveBackend, temporary_backend +from plyngent.tools.chat import CHAT_TOOLS, ask_user, choose_user, form_user +from tests.test_prompting import ScriptedBackend + + +async def test_ask_user_tool() -> None: + backend = ScriptedBackend(["42"]) + with temporary_backend(backend): + registry = ToolRegistry([ask_user]) + out = await registry.execute("ask_user", '{"question": "Answer?"}') + assert out == "42" + + +async def test_choose_user_tool_index() -> None: + backend = ScriptedBackend(["1"]) + with temporary_backend(backend): + registry = ToolRegistry([choose_user]) + out = await registry.execute( + "choose_user", + json.dumps( + { + "question": "Pick", + "options": json.dumps(["alpha", "beta"]), + "allow_custom": False, + } + ), + ) + assert out == "alpha" + + +async def test_choose_user_bad_options() -> None: + registry = ToolRegistry([choose_user]) + out = await registry.execute( + "choose_user", + json.dumps({"question": "Pick", "options": "not-json"}), + ) + assert out.startswith("error:") + + +async def test_form_user_tool() -> None: + backend = ScriptedBackend(["ncbm"], confirms=[True]) + fields = json.dumps([{"name": "user", "prompt": "User?"}]) + with temporary_backend(backend): + registry = ToolRegistry([form_user]) + out = await registry.execute( + "form_user", + json.dumps({"title": "Setup", "fields": fields, "confirm_submit": True}), + ) + assert json.loads(out) == {"user": "ncbm"} + + +async def test_chat_tools_in_default_list() -> None: + names = {t.name for t in CHAT_TOOLS} + assert names == {"ask_user", "choose_user", "form_user"} + + +async def test_ask_user_non_interactive_error() -> None: + with temporary_backend(NonInteractiveBackend()): + registry = ToolRegistry([ask_user]) + out = await registry.execute("ask_user", '{"question": "hi"}') + assert out.startswith("error:")