diff --git a/kt-kernel/python/cli/commands/config.py b/kt-kernel/python/cli/commands/config.py index 84f24759f..551767b34 100644 --- a/kt-kernel/python/cli/commands/config.py +++ b/kt-kernel/python/cli/commands/config.py @@ -137,10 +137,10 @@ def model_path_remove( def _parse_value(value: str): """Parse a string value into appropriate Python type.""" - # Try boolean - if value.lower() in ("true", "yes", "on", "1"): + # Try boolean ("1"/"0" fall through to int so e.g. CUDA_VISIBLE_DEVICES=1 stays "1") + if value.lower() in ("true", "yes", "on"): return True - if value.lower() in ("false", "no", "off", "0"): + if value.lower() in ("false", "no", "off"): return False # Try integer diff --git a/kt-kernel/test/per_commit/test_config_parse_value.py b/kt-kernel/test/per_commit/test_config_parse_value.py new file mode 100644 index 000000000..7c672eacd --- /dev/null +++ b/kt-kernel/test/per_commit/test_config_parse_value.py @@ -0,0 +1,44 @@ +import tempfile +import unittest +from pathlib import Path + +import yaml + +from ci.ci_register import register_cpu_ci +from kt_kernel.cli.commands.config import _parse_value +from kt_kernel.cli.config.settings import Settings + + +register_cpu_ci(est_time=0.1, suite="default") + + +class TestConfigParseValue(unittest.TestCase): + def test_numeric_strings_stay_integers(self): + for raw, expected in (("0", 0), ("1", 1), ("2", 2), ("30000", 30000)): + value = _parse_value(raw) + self.assertIs(type(value), int, raw) + self.assertEqual(value, expected) + + def test_boolean_words_still_parse_as_booleans(self): + for raw in ("true", "True", "yes", "on"): + self.assertIs(_parse_value(raw), True, raw) + for raw in ("false", "False", "no", "off"): + self.assertIs(_parse_value(raw), False, raw) + + def test_env_var_set_to_one_is_exported_as_one(self): + with tempfile.TemporaryDirectory() as tmp: + tmp = Path(tmp) + config_path = tmp / "config.yaml" + # Keep every directory Settings creates inside the temp dir. + config_path.write_text( + yaml.safe_dump({"paths": {"models": str(tmp / "models"), "cache": str(tmp / "cache")}}), + encoding="utf-8", + ) + + Settings(config_path=config_path).set("advanced.env.CUDA_VISIBLE_DEVICES", _parse_value("1")) + + self.assertEqual(Settings(config_path=config_path).get_env_vars(), {"CUDA_VISIBLE_DEVICES": "1"}) + + +if __name__ == "__main__": + unittest.main()