Download tests/test_registry_executor.py from SyntheticMDProductions/AI_Development_Automation_Manager: direct link, hf CLI and curl.
- Browser
- Download file 4.09 kB
-
https://huggingface.co/SyntheticMDProductions/AI_Development_Automation_Manager/resolve/refs%2Fpr%2F1/tests/test_registry_executor.py
- Command line
-
hf download hf://SyntheticMDProductions/AI_Development_Automation_Manager@refs/pr/1/tests/test_registry_executor.py
-
curl -L -o test_registry_executor.py https://huggingface.co/SyntheticMDProductions/AI_Development_Automation_Manager/resolve/refs%2Fpr%2F1/tests/test_registry_executor.py
4.09 kB
| from __future__ import annotations | |
| import json | |
| import shutil | |
| import threading | |
| from pathlib import Path | |
| import pytest | |
| from adam.executor import ToolExecutionError, ToolExecutor | |
| from adam.logging_setup import configure_logging | |
| from adam.registry import RegistryError, ToolRegistry | |
| ROOT = Path(__file__).resolve().parents[1] | |
| def make_root(tmp_path: Path) -> Path: | |
| (tmp_path / "config").mkdir() | |
| shutil.copy2(ROOT / "config" / "tools.json", tmp_path / "config" / "tools.json") | |
| return tmp_path | |
| def execute( | |
| executor: ToolExecutor, | |
| tool_id: str, | |
| arguments: dict, | |
| ) -> dict: | |
| run_event = threading.Event() | |
| run_event.set() | |
| return executor.execute( | |
| tool_id, | |
| arguments, | |
| job_id="TEST0001", | |
| cancel_event=threading.Event(), | |
| run_event=run_event, | |
| progress_callback=lambda _percent, _message: None, | |
| log_callback=lambda _message: None, | |
| ) | |
| def test_registry_exposes_enabled_trainers(tmp_path: Path) -> None: | |
| registry = ToolRegistry(make_root(tmp_path)) | |
| assert registry.get("lora_trainer").demo is False | |
| assert "resume_training" in registry.get("lora_trainer").capabilities | |
| assert registry.get("flow_trainer").demo is False | |
| assert "fresh_training" in registry.get("flow_trainer").capabilities | |
| def test_executor_rejects_unregistered_arguments(tmp_path: Path) -> None: | |
| project = make_root(tmp_path) | |
| executor = ToolExecutor( | |
| project, | |
| ToolRegistry(project), | |
| configure_logging(project), | |
| step_delay=0, | |
| ) | |
| with pytest.raises(ToolExecutionError, match="unsupported arguments"): | |
| execute( | |
| executor, | |
| "preview_generator", | |
| { | |
| "subject": "test", | |
| "project_name": "test", | |
| "preview_count": 1, | |
| "shell_command": "dangerous", | |
| }, | |
| ) | |
| def test_demo_pipeline_creates_truthful_reviewable_artifacts(tmp_path: Path) -> None: | |
| project = make_root(tmp_path) | |
| registry_path = project / "config" / "tools.json" | |
| registry_payload = json.loads(registry_path.read_text(encoding="utf-8")) | |
| collector = next( | |
| tool | |
| for tool in registry_payload["tools"] | |
| if tool["id"] == "dataset_collector" | |
| ) | |
| collector["demo"] = True | |
| collector["backend"] = { | |
| "type": "python", | |
| "module": "adam.tools.demo_backends", | |
| "function": "collect_dataset", | |
| } | |
| lora = next( | |
| tool for tool in registry_payload["tools"] if tool["id"] == "lora_trainer" | |
| ) | |
| lora["demo"] = True | |
| lora["arguments"] = ["subject", "project_name", "epochs"] | |
| lora["required_arguments"] = ["subject", "project_name", "epochs"] | |
| lora["backend"] = { | |
| "type": "python", | |
| "module": "adam.tools.demo_backends", | |
| "function": "train_lora", | |
| } | |
| registry_path.write_text( | |
| json.dumps(registry_payload), | |
| encoding="utf-8", | |
| ) | |
| executor = ToolExecutor( | |
| project, | |
| ToolRegistry(project), | |
| configure_logging(project), | |
| step_delay=0, | |
| ) | |
| common = {"subject": "Test Subject", "project_name": "Test Subject LoRA"} | |
| steps = [ | |
| ("dataset_collector", {**common, "image_count": 12}), | |
| ("dataset_preparer", {"project_name": common["project_name"]}), | |
| ("caption_generator", common), | |
| ("lora_trainer", {**common, "epochs": 2}), | |
| ("preview_generator", {**common, "preview_count": 2}), | |
| ("completion_notifier", {"project_name": common["project_name"]}), | |
| ] | |
| result = {} | |
| for tool_id, arguments in steps: | |
| result = execute(executor, tool_id, arguments) | |
| output = Path(result["output_folder"]) | |
| assert output.is_relative_to(project / "data" / "projects") | |
| summary = json.loads( | |
| (output / "training" / "training_summary.json").read_text(encoding="utf-8") | |
| ) | |
| assert summary["mode"] == "demo" | |
| assert summary["model_created"] is False | |
| assert len(list((output / "previews").glob("preview_*.svg"))) == 2 | |
| assert (output / "completion.json").exists() | |