lablab-ai-amd-developer-hackathon/gpu-goblin
0
1"""Tests for ``agent.tools.parse_config._parse_config``.2 3Covers all three input shapes (Python AST, JSON, YAML), redaction of every4secret pattern we ship, error paths for missing/malformed inputs, and the5WorkloadConfig field mapping for HF TrainingArguments + DataLoader kwargs.6"""7 8from __future__ import annotations9 10import json11from pathlib import Path12 13import pytest14 15from agent.schemas import WorkloadConfig16from agent.tools.parse_config import PARSE_CONFIG, _parse_config, _parse_config_full17 18FIXTURES = Path(__file__).parent / "fixtures"19 20 21# ---------------------------------------------------------------------------22# Python script path23# ---------------------------------------------------------------------------24 25 26class TestPythonScript:27 def test_returns_ok(self) -> None:28 result = _parse_config(str(FIXTURES / "sample_train.py"))29 assert result.ok, result.error30 assert result.result is not None31 32 def test_extracts_training_arguments_kwargs(self) -> None:33 result = _parse_config(str(FIXTURES / "sample_train.py"))34 cfg = result.result35 assert cfg["batch_size"] == 436 assert cfg["grad_accum_steps"] == 837 assert cfg["lr"] == pytest.approx(2e-4)38 assert cfg["warmup_steps"] == 10039 assert cfg["optimizer"] == "adamw_torch"40 # fp16=True in TrainingArguments → precision should resolve to fp16.41 assert cfg["precision"] == "fp16"42 43 def test_dataloader_kwargs_captured(self) -> None:44 result = _parse_config(str(FIXTURES / "sample_train.py"))45 cfg = result.result46 assert cfg["dataloader_workers"] == 047 assert cfg["dataloader_pin_memory"] is False48 assert cfg["dataloader_prefetch_factor"] == 249 assert cfg["dataloader_persistent_workers"] is False50 51 def test_torch_compile_call_flips_flag(self) -> None:52 result = _parse_config(str(FIXTURES / "sample_train.py"))53 # The script calls torch.compile(model, ...), even though the54 # TrainingArguments has torch_compile=False — the explicit call wins.55 assert result.result["torch_compile"] is True56 57 def test_gradient_checkpointing_enable_call(self) -> None:58 result = _parse_config(str(FIXTURES / "sample_train.py"))59 assert result.result["gradient_checkpointing"] is True60 61 def test_env_vars_captured(self) -> None:62 result = _parse_config(str(FIXTURES / "sample_train.py"))63 env = result.result["env_vars"]64 assert env["HSA_FORCE_FINE_GRAIN_PCIE"] == "1"65 assert env["MIOPEN_FIND_MODE"] == "3"66 assert env["NCCL_MIN_NCHANNELS"] == "112"67 68 def test_lora_rank_extracted(self) -> None:69 result = _parse_config(str(FIXTURES / "sample_train.py"))70 assert result.result["lora_rank"] == 1671 72 def test_attention_impl_from_from_pretrained(self) -> None:73 result = _parse_config(str(FIXTURES / "sample_train.py"))74 assert result.result["attention_impl"] == "eager"75 76 def test_model_name_resolved(self) -> None:77 result = _parse_config(str(FIXTURES / "sample_train.py"))78 assert result.result["model_name"] == "Qwen/Qwen2.5-7B-Instruct"79 80 81# ---------------------------------------------------------------------------82# JSON path83# ---------------------------------------------------------------------------84 85 86class TestJsonConfig:87 def test_returns_ok(self) -> None:88 result = _parse_config(str(FIXTURES / "sample_train.json"))89 assert result.ok, result.error90 91 def test_field_mapping(self) -> None:92 cfg = _parse_config(str(FIXTURES / "sample_train.json")).result93 assert cfg["model_name"] == "Qwen/Qwen2.5-7B-Instruct"94 assert cfg["batch_size"] == 895 assert cfg["grad_accum_steps"] == 496 assert cfg["seq_len"] == 409697 assert cfg["precision"] == "bf16"98 assert cfg["optimizer"] == "adamw_torch_fused"99 assert cfg["torch_compile"] is True100 assert cfg["gradient_checkpointing"] is True101 assert cfg["dataloader_workers"] == 4102 assert cfg["dataloader_pin_memory"] is True103 assert cfg["dataloader_persistent_workers"] is True104 assert cfg["attention_impl"] == "flash"105 106 def test_env_vars_dict(self) -> None:107 cfg = _parse_config(str(FIXTURES / "sample_train.json")).result108 assert cfg["env_vars"]["HSA_FORCE_FINE_GRAIN_PCIE"] == "1"109 110 def test_extras_collects_unmapped_fields(self) -> None:111 cfg = _parse_config(str(FIXTURES / "sample_train.json")).result112 # num_train_epochs has no slot in WorkloadConfig — must land in extras.113 assert cfg["extras"]["num_train_epochs"] == 3114 assert cfg["extras"]["save_steps"] == 500115 116 117# ---------------------------------------------------------------------------118# YAML path119# ---------------------------------------------------------------------------120 121 122class TestYamlConfig:123 def test_returns_ok(self) -> None:124 result = _parse_config(str(FIXTURES / "sample_train.yaml"))125 assert result.ok, result.error126 127 def test_field_mapping(self) -> None:128 cfg = _parse_config(str(FIXTURES / "sample_train.yaml")).result129 assert cfg["batch_size"] == 2130 assert cfg["grad_accum_steps"] == 16131 assert cfg["seq_len"] == 8192132 assert cfg["precision"] == "bf16"133 assert cfg["torch_compile"] is False134 assert cfg["gradient_checkpointing"] is True135 assert cfg["attention_impl"] == "sdpa"136 assert cfg["dataloader_workers"] == 8137 assert cfg["dataloader_persistent_workers"] is True138 139 140# ---------------------------------------------------------------------------141# Redaction142# ---------------------------------------------------------------------------143 144 145class TestRedaction:146 def test_python_script_redactions(self) -> None:147 cfg = _parse_config(str(FIXTURES / "sample_train.py")).result148 labels = set(cfg["redactions"])149 # Every secret pattern in sample_train.py should fire.150 assert "hf_token" in labels151 assert "openai_key" in labels152 assert "github_token" in labels153 assert "bearer_token" in labels154 assert "home_path" in labels155 assert "s3_uri" in labels156 assert "ws_uri" in labels157 158 def test_raw_source_is_scrubbed(self) -> None:159 # `raw_source` is intentionally stripped from the tool result envelope160 # (keeps the LLM conversation small) — use the `_full` helper to read161 # it. The redaction labels list still proves which patterns fired.162 cfg = _parse_config_full(str(FIXTURES / "sample_train.py"))163 assert isinstance(cfg, WorkloadConfig)164 raw = cfg.raw_source165 assert "hf_abcdefghijklmnopqrstuvwxyz123456" not in raw166 assert "sk-abcdefghijklmnopqrstuvwxyz1234567890" not in raw167 assert "gho_abcdefghijklmnopqrstuvwxyz123456" not in raw168 assert "/home/researcher/datasets/alpaca" not in raw169 assert "s3://my-team/checkpoints/qwen-lora/" not in raw170 assert "wss://logs.internal.example.com/stream" not in raw171 assert "<REDACTED:hf_token>" in raw172 assert "<REDACTED:openai_key>" in raw173 174 def test_raw_source_excluded_from_tool_result(self) -> None:175 # The tool result MUST NOT carry raw_source — it bloated the audit176 # conversation past 8K on Qwen2.5-7B during the live AMD GPU run.177 cfg = _parse_config(str(FIXTURES / "sample_train.py")).result178 assert "raw_source" not in cfg179 180 def test_json_redactions(self) -> None:181 cfg = _parse_config(str(FIXTURES / "sample_train.json")).result182 labels = set(cfg["redactions"])183 assert "hf_token" in labels184 assert "s3_uri" in labels185 # raw_source is no longer in the result; verify scrubbing via the186 # full-config helper.187 full = _parse_config_full(str(FIXTURES / "sample_train.json"))188 assert isinstance(full, WorkloadConfig)189 assert "hf_jsonsamplehfabcdefghijklmnopqrs" not in full.raw_source190 191 def test_extras_values_are_scrubbed(self) -> None:192 # Secret-shaped values that landed in extras must also be redacted —193 # otherwise the leak just moves from raw_source into extras.194 cfg = _parse_config(str(FIXTURES / "sample_train.json")).result195 extras = cfg["extras"]196 assert "hf_jsonsamplehfabcdefghijklmnopqrs" not in extras.get("hub_token", "")197 assert extras.get("hub_token", "").startswith("<REDACTED:")198 assert extras.get("checkpoint_uri", "").startswith("<REDACTED:")199 200 def test_yaml_redactions(self) -> None:201 cfg = _parse_config(str(FIXTURES / "sample_train.yaml")).result202 labels = set(cfg["redactions"])203 assert "hf_token" in labels204 assert "bearer_token" in labels205 assert "home_path" in labels206 207 208# ---------------------------------------------------------------------------209# Failure modes210# ---------------------------------------------------------------------------211 212 213class TestErrors:214 def test_missing_file(self) -> None:215 result = _parse_config("/tmp/definitely-does-not-exist-xyz.py")216 assert result.ok is False217 assert "not found" in (result.error or "").lower()218 219 def test_unsupported_extension(self, tmp_path: Path) -> None:220 bad = tmp_path / "config.toml"221 bad.write_text("model_name = 'foo'\n")222 result = _parse_config(str(bad))223 assert result.ok is False224 assert "unsupported" in (result.error or "").lower()225 226 def test_malformed_python(self, tmp_path: Path) -> None:227 bad = tmp_path / "broken.py"228 bad.write_text("def oops(:\n pass\n")229 result = _parse_config(str(bad))230 assert result.ok is False231 assert "parse error" in (result.error or "").lower()232 233 def test_malformed_json(self, tmp_path: Path) -> None:234 bad = tmp_path / "broken.json"235 bad.write_text("{not really json")236 result = _parse_config(str(bad))237 assert result.ok is False238 assert "json" in (result.error or "").lower()239 240 def test_json_top_level_must_be_dict(self, tmp_path: Path) -> None:241 bad = tmp_path / "list.json"242 bad.write_text(json.dumps([{"foo": 1}]))243 result = _parse_config(str(bad))244 assert result.ok is False245 246 def test_yaml_top_level_must_be_mapping(self, tmp_path: Path) -> None:247 bad = tmp_path / "scalar.yaml"248 bad.write_text("- 1\n- 2\n")249 result = _parse_config(str(bad))250 assert result.ok is False251 252 253# ---------------------------------------------------------------------------254# Schema invariants255# ---------------------------------------------------------------------------256 257 258class TestSchema:259 def test_result_round_trips_through_workload_config(self) -> None:260 result = _parse_config(str(FIXTURES / "sample_train.py"))261 # Must be reconstructible — guards against extras-vs-fields collisions.262 cfg = WorkloadConfig(**result.result)263 assert cfg.model_name == "Qwen/Qwen2.5-7B-Instruct"264 265 def test_defaults_when_field_absent(self, tmp_path: Path) -> None:266 # Minimal config — only model_name. Everything else should fall back to schema defaults.267 path = tmp_path / "tiny.json"268 path.write_text(json.dumps({"model_name": "test/tiny"}))269 result = _parse_config(str(path))270 assert result.ok271 cfg = result.result272 assert cfg["batch_size"] == 1273 assert cfg["precision"] == "fp16"274 assert cfg["optimizer"] == "adamw_torch"275 assert cfg["gradient_checkpointing"] is False276 assert cfg["redactions"] == []277 278 def test_tool_definition_unchanged_in_shape(self) -> None:279 # The Tool definition should still expose name/description/input_schema/fn.280 assert PARSE_CONFIG.name == "parse_config"281 assert PARSE_CONFIG.fn is _parse_config282 assert "file_path" in PARSE_CONFIG.input_schema["properties"]283 284 285# ---------------------------------------------------------------------------286# Regression: canonical + scenario workloads must parse with all the right287# audit-relevant fields. These are what the live agent actually sees, so a288# regression here directly degrades audit quality (the agent reasons over289# HF defaults instead of the script's settings).290# ---------------------------------------------------------------------------291 292 293REPO_ROOT = Path(__file__).resolve().parent.parent294 295 296class TestCanonicalWorkload:297 def test_canonical_workload_extracts_full_config(self) -> None:298 """The canonical demo workload must yield batch_size=4, lr=2e-4, etc.299 — not HF defaults. Catches the `**dict_var` splat regression where300 every TrainingArguments kwarg disappears.301 """302 result = _parse_config(str(REPO_ROOT / "workloads" / "train_qwen_lora.py"))303 assert result.ok, result.error304 cfg = result.result305 assert cfg["model_name"] == "Qwen/Qwen2.5-7B-Instruct"306 assert cfg["batch_size"] == 4, (307 "expected batch_size=4 from per_device_train_batch_size; "308 "did `**_ta_kwargs` splat hide the kwargs?"309 )310 assert cfg["grad_accum_steps"] == 8311 assert cfg["lr"] == 2e-4312 assert cfg["warmup_steps"] == 100313 assert cfg["precision"] == "fp16"314 assert cfg["attention_impl"] == "eager"315 assert cfg["dataloader_workers"] == 0316 assert cfg["dataloader_pin_memory"] is False317 assert cfg["lora_rank"] == 16318 assert cfg["torch_compile"] is False319 assert cfg["env_vars"]["HSA_FORCE_FINE_GRAIN_PCIE"] == "1"320 321 322class TestSplatKwargsResolution:323 """`_ta = dict(k=v); Foo(**_ta)` must resolve back through the dict324 constant. Defensive — the canonical workload no longer uses this325 pattern, but third-party scripts often do.326 """327 328 def test_dict_function_call_splat(self, tmp_path) -> None:329 src = """330from transformers import TrainingArguments331 332_ta = dict(333 per_device_train_batch_size=8,334 gradient_accumulation_steps=2,335 fp16=True,336 optim=\"adamw_torch_fused\",337)338training_args = TrainingArguments(output_dir=\"./out\", **_ta)339"""340 p = tmp_path / "splat.py"341 p.write_text(src)342 cfg = _parse_config(str(p)).result343 assert cfg["batch_size"] == 8344 assert cfg["grad_accum_steps"] == 2345 assert cfg["precision"] == "fp16"346 assert cfg["optimizer"] == "adamw_torch_fused"347 348 def test_dict_literal_splat(self, tmp_path) -> None:349 src = """350from transformers import TrainingArguments351 352_ta = {353 "per_device_train_batch_size": 16,354 "bf16": True,355}356training_args = TrainingArguments(output_dir=\"./out\", **_ta)357"""358 p = tmp_path / "splat_literal.py"359 p.write_text(src)360 cfg = _parse_config(str(p)).result361 assert cfg["batch_size"] == 16362 assert cfg["precision"] == "bf16"363 364 def test_explicit_kwarg_overrides_splat(self, tmp_path) -> None:365 src = """366from transformers import TrainingArguments367 368_ta = dict(per_device_train_batch_size=8)369training_args = TrainingArguments(per_device_train_batch_size=32, **_ta)370"""371 p = tmp_path / "splat_override.py"372 p.write_text(src)373 cfg = _parse_config(str(p)).result374 # Explicit kwarg wins over splat (both occur in the kwargs list,375 # explicit comes first in the AST → setdefault keeps it).376 assert cfg["batch_size"] == 32377 378 