diff --git a/strix/interface/completions.py b/strix/interface/completions.py new file mode 100644 index 00000000..013533e7 --- /dev/null +++ b/strix/interface/completions.py @@ -0,0 +1,158 @@ +"""Shell completion scripts and candidates for the Strix CLI.""" + +from __future__ import annotations + +import sys +from typing import Any + +from strix.interface.cloud.spec import SPEC, Cmd + + +_ROOT_COMMANDS = ("cloud", "auth", "view", "completions", "completion") +_SESSION_COMMANDS = ("login", "logout", "whoami", "credits") +_COMMON_FLAGS = ("--json", "--token", "--app-url", "--timeout", "--help") + + +def run_completions(argv: list[str]) -> int: + """Print a shell integration script or hidden completion candidates.""" + if argv and argv[0] == "--candidates": + for candidate in completion_candidates(argv[1:]): + sys.stdout.write(candidate + "\n") + return 0 + if not argv or argv[0] in ("-h", "--help", "help"): + sys.stdout.write( + "Usage: strix completions \n\n" + "Enable tab completion for the current shell:\n" + " zsh: source <(strix completions zsh)\n" + " bash: source <(strix completions bash)\n" + " fish: strix completions fish | source\n" + ) + return 0 + shell = argv[0].lower() + scripts = {"zsh": _zsh_script, "bash": _bash_script, "fish": _fish_script} + generator = scripts.get(shell) + if generator is None: + sys.stderr.write(f"Unknown shell: {shell}. Choose zsh, bash, or fish.\n") + return 2 + sys.stdout.write(generator()) + return 0 + + +def completion_candidates(words: list[str]) -> list[str]: + """Return candidates for words after the ``strix`` executable.""" + prior, current = _split_cursor(words) + if not prior: + return _matching(_ROOT_COMMANDS, current) + if prior[0] != "cloud": + return [] + return _cloud_candidates(prior[1:], current) + + +def _split_cursor(words: list[str]) -> tuple[list[str], str]: + if not words: + return [], "" + return words[:-1], words[-1] + + +def _cloud_candidates(prior: list[str], current: str) -> list[str]: + groups = (*_SESSION_COMMANDS, *SPEC, "workspace") + if not prior: + return _matching(groups, current) + group = "workspaces" if prior[0] == "workspace" else prior[0] + rest = prior[1:] + if group in _SESSION_COMMANDS: + return _matching(_session_flags(group), current) + commands = SPEC.get(group) + if commands is None: + return _matching(groups, current) + + verb_paths = [verb.split() for verb in commands] + if group == "workspaces": + verb_paths.append(["use"]) + matching_paths = [path for path in verb_paths if path[: len(rest)] == rest] + if not matching_paths: + return [] + next_words = sorted({path[len(rest)] for path in matching_paths if len(path) > len(rest)}) + exact_verbs = [" ".join(path) for path in matching_paths if len(path) == len(rest)] + candidates: list[str] = list(next_words) + for verb in exact_verbs: + if verb == "use" and group == "workspaces": + candidates.extend(_COMMON_FLAGS) + else: + candidates.extend(_command_flags(commands[verb])) + return _matching(candidates, current) + + +def _session_flags(group: str) -> tuple[str, ...]: + if group == "login": + return ("--no-browser", "--scopes", "--workspace", "--help") + if group == "whoami": + return ("--json", "--help") + return ("--help",) + + +def _command_flags(cmd: Cmd) -> tuple[str, ...]: + flags: list[str] = list(_COMMON_FLAGS) + for param in cmd.query + cmd.body: + flag = "--" + (param.flag or _kebab(param.name)) + flags.append(flag) + if param.kind == "bool": + flags.append("--no-" + flag.removeprefix("--")) + if cmd.method in ("POST", "PUT", "PATCH"): + flags.append("--data") + if cmd.binary: + flags.append("--output") + if cmd.link: + flags.append("--no-browser") + if cmd.wait_path or cmd.wait_self: + flags.append("--wait") + if cmd.path == "/billing/topup": + flags.extend(("--yes", "--no-pay", "--payment-method")) + if cmd.path == "/billing/auto-topup" and cmd.method == "PUT": + flags.append("--no-monthly-cap") + return tuple(dict.fromkeys(flags)) + + +def _kebab(value: str) -> str: + output: list[str] = [] + for char in value: + if char.isupper(): + output.extend(("-", char.lower())) + else: + output.append("-" if char == "_" else char) + return "".join(output) + + +def _matching(candidates: Any, prefix: str) -> list[str]: + return sorted({str(candidate) for candidate in candidates if str(candidate).startswith(prefix)}) + + +def _zsh_script() -> str: + return r"""#compdef strix +_strix() { + local -a candidates + candidates=("${(@f)$($words[1] completions --candidates "${words[@]:2}")}") + _describe 'strix' candidates +} +compdef _strix strix +""" + + +def _bash_script() -> str: + return r"""_strix_completion() { + local -a candidates + mapfile -t candidates < <(strix completions --candidates "${COMP_WORDS[@]:1:$COMP_CWORD}") + COMPREPLY=( $(compgen -W "${candidates[*]}" -- "${COMP_WORDS[$COMP_CWORD]}") ) +} +complete -F _strix_completion strix +""" + + +def _fish_script() -> str: + return r"""function __strix_candidates + set -l words (commandline -opc) + set -e words[1] + command strix completions --candidates $words (commandline -ct) +end +complete -c strix -f -a '(__strix_candidates)' +""" diff --git a/strix/interface/main.py b/strix/interface/main.py index 69a74b3b..e7dfc4ab 100644 --- a/strix/interface/main.py +++ b/strix/interface/main.py @@ -431,6 +431,12 @@ def main() -> None: sys.exit(run_auth(sys.argv[2:])) + # Generate native shell completion scripts before scan argument parsing. + if len(sys.argv) > 1 and sys.argv[1] in ("completion", "completions"): + from strix.interface.completions import run_completions + + sys.exit(run_completions(sys.argv[2:])) + # `strix cloud …` drives the managed platform (app.strix.ai) and exits; # it needs no target, Docker, or scan setup. if len(sys.argv) > 1 and sys.argv[1] == "cloud": diff --git a/tests/test_completions.py b/tests/test_completions.py new file mode 100644 index 00000000..373f4634 --- /dev/null +++ b/tests/test_completions.py @@ -0,0 +1,49 @@ +from __future__ import annotations + +from typing import Any + +from strix.interface.completions import completion_candidates, run_completions + + +def test_root_completion_candidates() -> None: + assert completion_candidates(["cl"]) == ["cloud"] + assert "completions" in completion_candidates([""]) + + +def test_cloud_group_and_alias_candidates() -> None: + candidates = completion_candidates(["cloud", "work"]) + assert candidates == ["workspace", "workspaces"] + + +def test_cloud_verb_candidates_include_multiword_prefixes() -> None: + assert "test-users" in completion_candidates(["cloud", "domains", ""]) + assert completion_candidates(["cloud", "domains", "test-users", "in"]) == [ + "inbox", + "inbox-message", + ] + + +def test_cloud_leaf_flag_candidates_come_from_command_spec() -> None: + candidates = completion_candidates(["cloud", "scans", "start", "--"]) + assert "--domain-ids" in candidates + assert "--json" in candidates + assert "--wait" in candidates + + +def test_boolean_completion_includes_positive_and_negative_flags() -> None: + candidates = completion_candidates(["cloud", "billing", "auto-topup", "update", "--"]) + assert "--enabled" in candidates + assert "--no-enabled" in candidates + assert "--no-monthly-cap" in candidates + + +def test_completion_scripts_cover_supported_shells(capsys: Any) -> None: + for shell in ("zsh", "bash", "fish"): + assert run_completions([shell]) == 0 + output = capsys.readouterr().out + assert "completions --candidates" in output + + +def test_completion_rejects_unknown_shell(capsys: Any) -> None: + assert run_completions(["powershell"]) == 2 + assert "Choose zsh, bash, or fish" in capsys.readouterr().err