|
14 | 14 | from __future__ import annotations |
15 | 15 |
|
16 | 16 | import argparse |
| 17 | +import dataclasses |
17 | 18 | import inspect |
18 | 19 | import json |
19 | 20 | import logging |
@@ -342,6 +343,172 @@ def _parse_initializer_arg(arg: str) -> str | dict[str, Any]: |
342 | 343 | return name |
343 | 344 |
|
344 | 345 |
|
| 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 | + |
345 | 512 | def parse_run_arguments(*, args_string: str) -> dict[str, Any]: |
346 | 513 | """ |
347 | 514 | Parse run command arguments from a string (for shell mode). |
@@ -371,92 +538,30 @@ def parse_run_arguments(*, args_string: str) -> dict[str, Any]: |
371 | 538 | if not parts: |
372 | 539 | raise ValueError("No scenario name provided") |
373 | 540 |
|
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] |
457 | 543 | return result |
458 | 544 |
|
459 | 545 |
|
| 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 | + |
460 | 565 | # --------------------------------------------------------------------------- |
461 | 566 | # Shared argparse builder |
462 | 567 | # --------------------------------------------------------------------------- |
|
0 commit comments