diff --git a/codecarbon/cli/main.py b/codecarbon/cli/main.py index c10b32338..3626d896a 100644 --- a/codecarbon/cli/main.py +++ b/codecarbon/cli/main.py @@ -2,8 +2,9 @@ import signal import sys import time +from inspect import signature from pathlib import Path -from typing import Optional +from typing import Optional, get_args import typer from rich import print @@ -27,6 +28,78 @@ codecarbon = typer.Typer(no_args_is_help=True) +def _tracker_option_types() -> dict[str, object]: + """Return CLI-forwardable options from the tracker constructor.""" + from codecarbon.emissions_tracker import BaseEmissionsTracker + + return { + name: parameter.annotation + for name, parameter in signature( + BaseEmissionsTracker.__init__ + ).parameters.items() + if name != "self" and parameter.kind.name != "VAR_KEYWORD" + } + + +def _tracker_option_value_type(annotation: object) -> object: + if isinstance(annotation, str): + for value_type in (bool, int, float, str): + if value_type.__name__ in annotation: + return value_type + return str + + candidates = [ + candidate for candidate in get_args(annotation) if candidate is not type(None) + ] + return candidates[0] if candidates else annotation + + +def _parse_tracker_option(value: str, annotation: object) -> object: + """Convert a CLI value using the scalar type declared by the tracker.""" + value_type = _tracker_option_value_type(annotation) + if value_type is bool: + return value.lower() == "true" + if value_type in (int, float, str): + return value_type(value) + return value + + +def _extract_tracker_options(args: list[str]) -> tuple[dict[str, object], list[str]]: + """Remove tracker options that Typer does not declare from command arguments.""" + option_types = _tracker_option_types() + tracker_options: dict[str, object] = {} + remaining = list(args) + + while remaining and remaining[0].startswith("--"): + raw_option = remaining[0][2:] + option_name, separator, raw_value = raw_option.partition("=") + negated = option_name.startswith("no-") + normalized_name = option_name[3:] if negated else option_name + parameter_name = normalized_name.replace("-", "_") + annotation = option_types.get(parameter_name) + if annotation is None: + break + + remaining.pop(0) + value_type = _tracker_option_value_type(annotation) + if value_type is bool: + if separator: + tracker_options[parameter_name] = _parse_tracker_option( + raw_value, annotation + ) + else: + tracker_options[parameter_name] = not negated + continue + + if not separator: + if not remaining: + raise typer.BadParameter(f"Option '--{option_name}' requires a value") + raw_value = remaining.pop(0) + tracker_options[parameter_name] = _parse_tracker_option(raw_value, annotation) + + return tracker_options, remaining + + def main(): """ Main entry point for the CodeCarbon CLI application. @@ -362,11 +435,17 @@ def monitor( ): """Monitor your machine's carbon emissions.""" + extra_tracker_args, command_args = _extract_tracker_options( + list(getattr(ctx, "args", None) or []) + ) + ctx.args = command_args + # Shared tracker args so monitor and run_and_monitor behave the same tracker_args = { "measure_power_secs": measure_power_secs, "api_call_interval": api_call_interval, "log_level": log_level, + **extra_tracker_args, } # Set up the tracker arguments based on mode (offline vs online) and validate required args for each mode if offline: diff --git a/tests/cli/test_cli_main.py b/tests/cli/test_cli_main.py index 84f42493d..4c05d9484 100644 --- a/tests/cli/test_cli_main.py +++ b/tests/cli/test_cli_main.py @@ -305,6 +305,45 @@ def stop(self): assert calls["kwargs"]["region"] == "IDF" +def test_monitor_forwards_tracker_options_not_declared_by_typer(monkeypatch): + calls = {"kwargs": None, "started": 0} + + class FakeOfflineTracker: + def __init__(self, **kwargs): + calls["kwargs"] = kwargs + self._another_instance_already_running = True + + def start(self): + calls["started"] += 1 + + def stop(self): + return None + + monkeypatch.setattr( + "codecarbon.emissions_tracker.OfflineEmissionsTracker", FakeOfflineTracker + ) + monkeypatch.setattr(cli_main.signal, "signal", lambda *args, **kwargs: None) + + runner = CliRunner() + result = runner.invoke( + cli_main.codecarbon, + [ + "monitor", + "--offline", + "--country-iso-code", + "FRA", + "--pue", + "1.25", + "--allow-multiple-runs", + ], + ) + + assert result.exit_code == 0 + assert calls["started"] == 1 + assert calls["kwargs"]["pue"] == 1.25 + assert calls["kwargs"]["allow_multiple_runs"] is True + + def test_monitor_delegates_offline_flag_to_run_and_monitor(monkeypatch): captured = {}