AI_Development_Automation_Manager / tests /test_external_tools.py
SyntheticMDProductions's picture
Some of Adams structure
e0265b9 verified
Raw
History Blame Contribute Delete
4.31 kB
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")