Skip to content

Commit bd23ce5

Browse files
authored
FIX: CLI bug fixes and minor updates (#1559)
1 parent 30dc0b0 commit bd23ce5

8 files changed

Lines changed: 736 additions & 129 deletions

File tree

pyrit/cli/_banner.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -255,7 +255,7 @@ def add(line: str, role: ColorRole, segments: Optional[list[tuple[int, int, Colo
255255
"Commands:",
256256
" • list-scenarios - See all available scenarios",
257257
" • list-initializers - See all available initializers",
258-
" • list-targets - See all available targets in the registry",
258+
" • list-targets [opts] - See all available targets in the registry",
259259
" • run <scenario> [opts] - Execute a security scenario",
260260
" • scenario-history - View your session history",
261261
" • print-scenario [N] - Display detailed results",

pyrit/cli/_cli_args.py

Lines changed: 188 additions & 83 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,7 @@
1414
from __future__ import annotations
1515

1616
import argparse
17+
import dataclasses
1718
import inspect
1819
import json
1920
import logging
@@ -342,6 +343,172 @@ def _parse_initializer_arg(arg: str) -> str | dict[str, Any]:
342343
return name
343344

344345

346+
# ---------------------------------------------------------------------------
347+
# Shell argument specification
348+
# ---------------------------------------------------------------------------
349+
350+
351+
@dataclasses.dataclass(frozen=True)
352+
class _ArgSpec:
353+
"""
354+
Declarative specification for a single shell-mode CLI argument.
355+
356+
Each instance describes one CLI flag (or set of aliases) and how its
357+
value(s) should be collected and validated. A list of ``_ArgSpec`` objects
358+
is passed to ``_parse_shell_arguments`` which handles the actual parsing
359+
loop. Adding a new flag only requires defining a new ``_ArgSpec``
360+
constant, not editing any parsing logic.
361+
362+
Attributes:
363+
flags: CLI flag strings that trigger this argument (e.g., ``["--strategies", "-s"]``).
364+
result_key: Key name in the returned dict (e.g., ``"scenario_strategies"``).
365+
multi_value: If True, collect values until the next flag.
366+
If False, consume exactly one value.
367+
parser: Optional callable to transform each raw string value.
368+
Applied per-item for multi-value args, or to the single value otherwise.
369+
"""
370+
371+
flags: list[str]
372+
result_key: str
373+
multi_value: bool = False
374+
parser: Callable[[str], Any] | None = None
375+
376+
377+
_INITIALIZERS_ARG = _ArgSpec(
378+
flags=["--initializers"],
379+
result_key="initializers",
380+
multi_value=True,
381+
parser=_parse_initializer_arg,
382+
)
383+
_INIT_SCRIPTS_ARG = _ArgSpec(
384+
flags=["--initialization-scripts"],
385+
result_key="initialization_scripts",
386+
multi_value=True,
387+
)
388+
389+
_STRATEGIES_ARG = _ArgSpec(
390+
flags=["--strategies", "-s"],
391+
result_key="scenario_strategies",
392+
multi_value=True,
393+
)
394+
_MAX_CONCURRENCY_ARG = _ArgSpec(
395+
flags=["--max-concurrency"],
396+
result_key="max_concurrency",
397+
parser=lambda v: validate_integer(v, name="--max-concurrency", min_value=1),
398+
)
399+
_MAX_RETRIES_ARG = _ArgSpec(
400+
flags=["--max-retries"],
401+
result_key="max_retries",
402+
parser=lambda v: validate_integer(v, name="--max-retries", min_value=0),
403+
)
404+
_MEMORY_LABELS_ARG = _ArgSpec(
405+
flags=["--memory-labels"],
406+
result_key="memory_labels",
407+
parser=parse_memory_labels,
408+
)
409+
_LOG_LEVEL_ARG = _ArgSpec(
410+
flags=["--log-level"],
411+
result_key="log_level",
412+
parser=lambda v: validate_log_level(log_level=v),
413+
)
414+
_DATASET_NAMES_ARG = _ArgSpec(
415+
flags=["--dataset-names"],
416+
result_key="dataset_names",
417+
multi_value=True,
418+
)
419+
_MAX_DATASET_SIZE_ARG = _ArgSpec(
420+
flags=["--max-dataset-size"],
421+
result_key="max_dataset_size",
422+
parser=lambda v: validate_integer(v, name="--max-dataset-size", min_value=1),
423+
)
424+
_TARGET_ARG = _ArgSpec(
425+
flags=["--target"],
426+
result_key="target",
427+
)
428+
429+
_RUN_ARG_SPECS: list[_ArgSpec] = [
430+
_INITIALIZERS_ARG,
431+
_INIT_SCRIPTS_ARG,
432+
_STRATEGIES_ARG,
433+
_MAX_CONCURRENCY_ARG,
434+
_MAX_RETRIES_ARG,
435+
_MEMORY_LABELS_ARG,
436+
_LOG_LEVEL_ARG,
437+
_DATASET_NAMES_ARG,
438+
_MAX_DATASET_SIZE_ARG,
439+
_TARGET_ARG,
440+
]
441+
442+
_LIST_TARGETS_ARG_SPECS: list[_ArgSpec] = [
443+
_INITIALIZERS_ARG,
444+
_INIT_SCRIPTS_ARG,
445+
]
446+
447+
448+
# ---------------------------------------------------------------------------
449+
# Generic shell argument parser
450+
# ---------------------------------------------------------------------------
451+
452+
453+
def _parse_shell_arguments(*, parts: list[str], arg_specs: list[_ArgSpec]) -> dict[str, Any]:
454+
"""
455+
Parse a list of shell tokens against a set of argument specifications.
456+
457+
Each ``_ArgSpec`` in *arg_specs* declares how its flag(s) should be handled
458+
(multi-value collection vs. single-value consumption) and what validation
459+
or transformation to apply.
460+
461+
Args:
462+
parts: Token list (already split on whitespace, positional args removed).
463+
arg_specs: Argument specifications that this command accepts.
464+
465+
Returns:
466+
Dictionary mapping each spec's ``result_key`` to its parsed value,
467+
defaulting to ``None`` for arguments not present in *parts*.
468+
469+
Raises:
470+
ValueError: On unknown flags or missing values.
471+
"""
472+
# Build lookup: flag string → spec
473+
flag_to_spec: dict[str, _ArgSpec] = {}
474+
for spec in arg_specs:
475+
for flag in spec.flags:
476+
flag_to_spec[flag] = spec
477+
478+
# Initialise result with None defaults
479+
result: dict[str, Any] = {spec.result_key: None for spec in arg_specs}
480+
481+
i = 0
482+
while i < len(parts):
483+
token = parts[i]
484+
spec = flag_to_spec.get(token)
485+
486+
if spec is None:
487+
valid = sorted(flag_to_spec.keys())
488+
raise ValueError(f"Unknown argument: {token}. Valid arguments: {', '.join(valid)}")
489+
490+
i += 1
491+
492+
if spec.multi_value:
493+
values: list[Any] = []
494+
# Collect values until the next flag (whether valid or invalid)
495+
while i < len(parts) and not (parts[i].startswith("--") or parts[i] in flag_to_spec):
496+
item = spec.parser(parts[i]) if spec.parser else parts[i]
497+
values.append(item)
498+
i += 1
499+
if len(values) == 0:
500+
raise ValueError(f"{spec.flags[0]} requires at least one value")
501+
result[spec.result_key] = values
502+
else:
503+
if i >= len(parts):
504+
raise ValueError(f"{spec.flags[0]} requires a value")
505+
raw = parts[i]
506+
result[spec.result_key] = spec.parser(raw) if spec.parser else raw
507+
i += 1
508+
509+
return result
510+
511+
345512
def parse_run_arguments(*, args_string: str) -> dict[str, Any]:
346513
"""
347514
Parse run command arguments from a string (for shell mode).
@@ -371,92 +538,30 @@ def parse_run_arguments(*, args_string: str) -> dict[str, Any]:
371538
if not parts:
372539
raise ValueError("No scenario name provided")
373540

374-
result: dict[str, Any] = {
375-
"scenario_name": parts[0],
376-
"initializers": None,
377-
"initialization_scripts": None,
378-
"scenario_strategies": None,
379-
"max_concurrency": None,
380-
"max_retries": None,
381-
"memory_labels": None,
382-
"log_level": None,
383-
"dataset_names": None,
384-
"max_dataset_size": None,
385-
"target": None,
386-
}
387-
388-
i = 1
389-
while i < len(parts):
390-
if parts[i] == "--initializers":
391-
# Collect initializers until next flag, parsing name:key=val syntax
392-
result["initializers"] = []
393-
i += 1
394-
while i < len(parts) and not parts[i].startswith("--"):
395-
result["initializers"].append(_parse_initializer_arg(parts[i]))
396-
i += 1
397-
elif parts[i] == "--initialization-scripts":
398-
# Collect script paths until next flag
399-
result["initialization_scripts"] = []
400-
i += 1
401-
while i < len(parts) and not parts[i].startswith("--"):
402-
result["initialization_scripts"].append(parts[i])
403-
i += 1
404-
elif parts[i] in ("--strategies", "-s"):
405-
# Collect strategies until next flag
406-
result["scenario_strategies"] = []
407-
i += 1
408-
while i < len(parts) and not parts[i].startswith("--") and parts[i] != "-s":
409-
result["scenario_strategies"].append(parts[i])
410-
i += 1
411-
elif parts[i] == "--max-concurrency":
412-
i += 1
413-
if i >= len(parts):
414-
raise ValueError("--max-concurrency requires a value")
415-
result["max_concurrency"] = validate_integer(parts[i], name="--max-concurrency", min_value=1)
416-
i += 1
417-
elif parts[i] == "--max-retries":
418-
i += 1
419-
if i >= len(parts):
420-
raise ValueError("--max-retries requires a value")
421-
result["max_retries"] = validate_integer(parts[i], name="--max-retries", min_value=0)
422-
i += 1
423-
elif parts[i] == "--memory-labels":
424-
i += 1
425-
if i >= len(parts):
426-
raise ValueError("--memory-labels requires a value")
427-
result["memory_labels"] = parse_memory_labels(parts[i])
428-
i += 1
429-
elif parts[i] == "--log-level":
430-
i += 1
431-
if i >= len(parts):
432-
raise ValueError("--log-level requires a value")
433-
result["log_level"] = validate_log_level(log_level=parts[i])
434-
i += 1
435-
elif parts[i] == "--dataset-names":
436-
# Collect dataset names until next flag
437-
result["dataset_names"] = []
438-
i += 1
439-
while i < len(parts) and not parts[i].startswith("--"):
440-
result["dataset_names"].append(parts[i])
441-
i += 1
442-
elif parts[i] == "--max-dataset-size":
443-
i += 1
444-
if i >= len(parts):
445-
raise ValueError("--max-dataset-size requires a value")
446-
result["max_dataset_size"] = validate_integer(parts[i], name="--max-dataset-size", min_value=1)
447-
i += 1
448-
elif parts[i] == "--target":
449-
i += 1
450-
if i >= len(parts):
451-
raise ValueError("--target requires a value")
452-
result["target"] = parts[i]
453-
i += 1
454-
else:
455-
raise ValueError(f"Unknown argument: {parts[i]}")
456-
541+
result = _parse_shell_arguments(parts=parts[1:], arg_specs=_RUN_ARG_SPECS)
542+
result["scenario_name"] = parts[0]
457543
return result
458544

459545

546+
def parse_list_targets_arguments(*, args_string: str) -> dict[str, Any]:
547+
"""
548+
Parse list-targets command arguments from a string (for shell mode).
549+
550+
Args:
551+
args_string: Space-separated argument string (e.g., "--initializers target").
552+
553+
Returns:
554+
Dictionary with parsed arguments:
555+
- initializers: Optional[list[str | dict[str, Any]]]
556+
- initialization_scripts: Optional[list[str]]
557+
558+
Raises:
559+
ValueError: If parsing or validation fails.
560+
"""
561+
parts = args_string.split()
562+
return _parse_shell_arguments(parts=parts, arg_specs=_LIST_TARGETS_ARG_SPECS)
563+
564+
460565
# ---------------------------------------------------------------------------
461566
# Shared argparse builder
462567
# ---------------------------------------------------------------------------

0 commit comments

Comments
 (0)