diff --git a/codecarbon/cli/cli_utils.py b/codecarbon/cli/cli_utils.py index 04567c4db..14f2a2336 100644 --- a/codecarbon/cli/cli_utils.py +++ b/codecarbon/cli/cli_utils.py @@ -15,8 +15,8 @@ def get_config(path: Optional[Path] = None): config = configparser.ConfigParser() config.read(str(p)) if "codecarbon" in config.sections(): - d = dict(config["codecarbon"]) - return d + return dict(config["codecarbon"]) + return {} else: raise FileNotFoundError( @@ -34,8 +34,9 @@ def get_api_endpoint(path: Optional[Path] = None): if "api_endpoint" in d: return d["api_endpoint"] else: - with p.open("a") as f: - f.write("api_endpoint=https://api.codecarbon.io\n") + config["codecarbon"]["api_endpoint"] = "https://api.codecarbon.io" + with p.open("w") as f: + config.write(f) return "https://api.codecarbon.io" diff --git a/tests/cli/test_cli_utils.py b/tests/cli/test_cli_utils.py index a5840df92..6c5619a1d 100644 --- a/tests/cli/test_cli_utils.py +++ b/tests/cli/test_cli_utils.py @@ -105,3 +105,22 @@ def test_create_new_config_file_expands_home(monkeypatch, tmp_path): assert created_path == target assert target.exists() + + +def test_get_config_returns_empty_dict_without_codecarbon_section(tmp_path): + config_path = tmp_path / ".codecarbon.config" + config_path.write_text("[other]\nkey=value\n") + + assert cli_utils.get_config(config_path) == {} + + +def test_get_api_endpoint_writes_default_under_codecarbon_section(tmp_path): + config_path = tmp_path / ".codecarbon.config" + config_path.write_text("[codecarbon]\nproject_id=abc\n\n[other]\nkey=value\n") + + cli_utils.get_api_endpoint(config_path) + + parser = configparser.ConfigParser() + parser.read(config_path) + assert parser["codecarbon"]["api_endpoint"] == "https://api.codecarbon.io" + assert "api_endpoint" not in parser["other"]