Team Ai
Apppublic

lablab-ai-amd-developer-hackathon/gpu-goblin

sourceHugging Facemitupdated 5mo agoView on Hugging Face
0likes
test_parse_config.py378 linesDownload Raw Back to tests
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