File size: 4,308 Bytes
e0265b9 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 | from __future__ import annotations
import json
import shutil
from pathlib import Path
from adam.config import ConfigManager
from adam.external_tools import ExternalToolStore, scan_folder
from adam.planner import Planner
from adam.registry import ToolRegistry
ROOT = Path(__file__).resolve().parents[1]
def make_project(tmp_path: Path) -> Path:
(tmp_path / "config").mkdir()
shutil.copy2(ROOT / "config" / "tools.json", tmp_path / "config" / "tools.json")
return tmp_path
def test_analyzer_detects_training_contract(tmp_path: Path) -> None:
script = tmp_path / "train.py"
script.write_text(
"""
import argparse
from tqdm import tqdm
parser = argparse.ArgumentParser()
parser.add_argument("--dataset-dir", required=True)
parser.add_argument("--epochs", type=int, default=10)
parser.add_argument("--output-dir", required=True)
parser.add_argument("--resume-from")
if __name__ == "__main__":
args = parser.parse_args()
for epoch in tqdm(range(args.epochs)):
print("loss", epoch)
checkpoint = "checkpoint.pt"
""",
encoding="utf-8",
)
analysis = scan_folder(str(tmp_path))
assert analysis.selected_entry == "train.py"
assert analysis.score >= 8
assert analysis.required_arguments == ["dataset_dir", "output_dir"]
assert "resume_from" in analysis.resume_behavior
assert "tqdm" in analysis.progress_behavior
def test_analyzer_lowers_rating_for_risky_calls(tmp_path: Path) -> None:
(tmp_path / "train.py").write_text(
"""
import os
import shutil
if __name__ == "__main__":
os.system("unknown command")
shutil.rmtree("output")
""",
encoding="utf-8",
)
analysis = scan_folder(str(tmp_path))
assert analysis.score <= 3
assert any("delete" in warning or "os.system" in warning for warning in analysis.warnings)
def test_saved_external_tool_is_confirmation_gated_and_plannable(tmp_path: Path) -> None:
project = make_project(tmp_path)
tool_folder = tmp_path / "apvd"
tool_folder.mkdir()
(tool_folder / "train.py").write_text(
"""
import argparse
parser = argparse.ArgumentParser()
parser.add_argument("--dataset", required=True)
parser.add_argument("--epochs", type=int, required=True)
parser.add_argument("--output")
if __name__ == "__main__":
args = parser.parse_args()
print("training progress")
""",
encoding="utf-8",
)
analysis = scan_folder(str(tool_folder))
ExternalToolStore(project).save_connector(
name="APVD Model Trainer",
description="Train the APVD model.",
analysis=analysis,
arguments=analysis.arguments,
required_arguments=analysis.required_arguments,
)
registry = ToolRegistry(project)
spec = registry.get("external_apvd_model_trainer")
assert spec.requires_confirmation is True
assert spec.backend["type"] == "script"
config = ConfigManager(project)
config.settings["provider"] = "manual"
planner = Planner(project, registry, config)
plan = planner.plan(
"Run APVD Model Trainer with dataset=D:/DreamData, epochs=20, output=D:/Runs"
)
assert plan.requires_confirmation is True
assert plan.steps[0].tool_id == "external_apvd_model_trainer"
assert plan.steps[0].arguments["epochs"] == 20
def test_external_registry_cannot_override_builtin_tool(tmp_path: Path) -> None:
project = make_project(tmp_path)
(project / "config" / "external_tools.json").write_text(
json.dumps(
{
"tools": [
{
"id": "ddpm_trainer",
"name": "Replacement",
"description": "Not allowed",
"category": "External",
"backend": {
"type": "script",
"path": str((project / "train.py").resolve()),
"root": str(project.resolve()),
},
}
]
}
),
encoding="utf-8",
)
try:
ToolRegistry(project)
except Exception as exc:
assert "external_" in str(exc) or "replace" in str(exc)
else:
raise AssertionError("External registry override was accepted")
|