Download tests/test_remote_phase3a.py from SyntheticMDProductions/AI_Development_Automation_Manager: direct link, hf CLI and curl.
- Browser
- Download file 20.1 kB
-
https://huggingface.co/SyntheticMDProductions/AI_Development_Automation_Manager/resolve/refs%2Fpr%2F1/tests/test_remote_phase3a.py
- Command line
-
hf download hf://SyntheticMDProductions/AI_Development_Automation_Manager@refs/pr/1/tests/test_remote_phase3a.py
-
curl -L -o test_remote_phase3a.py https://huggingface.co/SyntheticMDProductions/AI_Development_Automation_Manager/resolve/refs%2Fpr%2F1/tests/test_remote_phase3a.py
20.1 kB
| from __future__ import annotations | |
| import json | |
| import logging | |
| import threading | |
| import time | |
| from pathlib import Path | |
| from urllib.error import HTTPError | |
| from urllib.request import Request, urlopen | |
| from PIL import Image | |
| import pytest | |
| from adam.assets import AssetRegistry | |
| from adam.config import ConfigManager | |
| from adam.dataset_registry import DatasetRegistry | |
| from adam.executor import ToolContext | |
| from adam.generations import build_generation_plan | |
| from adam.job_manager import JobManager | |
| from adam.models import ExecutionPlan, Job, JobStatus, PlanStep | |
| from adam.planner import Planner | |
| from adam.registry import ToolRegistry | |
| from adam.remote_access import RemoteAccessService | |
| from adam.remote_dispatcher import RemoteCommandDispatcher | |
| from adam.remote_media import OpaqueIdCodec, RemoteMediaStore | |
| from adam.remote_v1 import RemoteV1Service | |
| from adam.studio import caption_path | |
| from adam.tools.lora_adapter import train_lora | |
| def _image(path: Path, color: tuple[int, int, int] = (40, 120, 210)) -> None: | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| Image.new("RGB", (16, 12), color).save(path) | |
| def test_remote_training_reviews_before_auto_approval( | |
| tmp_path: Path, monkeypatch, structured: bool, epochs: int, needs_review: bool, | |
| ) -> None: | |
| dataset = tmp_path / "dataset" | |
| _image(dataset / "image.png") | |
| trainer = tmp_path / "trainer" | |
| trainer.mkdir() | |
| config = _config(tmp_path, { | |
| "tool_folders": {"ddpm_trainer": str(trainer)}, | |
| "remote_access": {"auto_approve_training": True, "token": "test-token"}, | |
| }) | |
| planner = Planner(tmp_path, ToolRegistry(Path.cwd()), config) | |
| asset = planner.assets.register(kind="dataset", name="Test Dataset", path=str(dataset)) | |
| jobs = JobManager(tmp_path, None, logging.getLogger("test.remote.review"), config) | |
| monkeypatch.setattr(jobs, "_start_next", lambda: None) | |
| service = RemoteAccessService(config, jobs, None, planner) | |
| try: | |
| if structured: | |
| payload = { | |
| "trainer": "ddpm", "model_name": "Test Model", "epochs": epochs, | |
| "dataset_id": service.codec.encode({"kind": "dataset", "asset_id": asset.id}), | |
| } | |
| preview = service.api_v1.training_plan(payload) | |
| assert "Pre-flight:" in preview["summary"] | |
| assert "ORION —" in preview["summary"] | |
| response = service.api_v1.start_training(payload) | |
| else: | |
| response = service.submit_prompt( | |
| f"From the Test Dataset dataset, train a DDPM model for {epochs} epochs. " | |
| "Name the model Test Model." | |
| ) | |
| assert response["ok"] is True | |
| job = jobs.get(response["job_id"]) | |
| assert response["requires_approval"] is needs_review | |
| assert job.status == (JobStatus.AWAITING_CONFIRMATION if needs_review else JobStatus.QUEUED) | |
| assert job.plan.orion_review["level"] == ("warning" if needs_review else "ready") | |
| assert job.plan.summary.count("Pre-flight:") == 1 | |
| assert job.plan.summary.count("ORION —") == 1 | |
| assert job.plan.steps[0].arguments["epochs"] == epochs | |
| assert config.get("remote_access")["auto_approve_training"] is True | |
| finally: | |
| service.shutdown() | |
| jobs.shutdown() | |
| def _config(root: Path, values: dict | None = None) -> ConfigManager: | |
| config = ConfigManager(root) | |
| if values: | |
| config.update(values) | |
| return config | |
| def _remote_v1(root: Path, *, jobs=None, planner=None) -> RemoteV1Service: | |
| config = _config(root) | |
| planner = planner or Planner(root, ToolRegistry(Path.cwd()), config) | |
| return RemoteV1Service( | |
| root=root, | |
| config=config, | |
| jobs=jobs, | |
| planner=planner, | |
| dispatcher=RemoteCommandDispatcher(), | |
| codec=OpaqueIdCodec("test-secret"), | |
| media=RemoteMediaStore(root, OpaqueIdCodec("test-secret")), | |
| auto_approve_training=lambda _plan: False, | |
| ) | |
| def test_remote_dispatcher_uses_invoker_from_worker_thread() -> None: | |
| dispatcher = RemoteCommandDispatcher() | |
| calls: list[str] = [] | |
| class Invoker: | |
| def invoke(self, payload): | |
| calls.append("invoked") | |
| payload["result"] = payload["fn"]() | |
| payload["event"].set() | |
| dispatcher._invoker = Invoker() | |
| result: list[str] = [] | |
| thread = threading.Thread(target=lambda: result.append(dispatcher.call_ui(lambda: "done"))) | |
| thread.start() | |
| thread.join(timeout=3) | |
| dispatcher.shutdown() | |
| assert calls == ["invoked"] | |
| assert result == ["done"] | |
| def test_remote_v1_datasets_are_paginated_redacted_and_editable(tmp_path: Path) -> None: | |
| dataset = tmp_path / "datasets" / "Minecraft Steve" | |
| for index in range(3): | |
| image = dataset / f"image_{index}.png" | |
| _image(image, (index * 40, 100, 200)) | |
| caption_path(image).write_text(f"caption {index}\n", encoding="utf-8") | |
| assets = AssetRegistry(tmp_path) | |
| dataset_asset = assets.register(kind="dataset", name="Minecraft Steve", path=str(dataset)) | |
| api = _remote_v1(tmp_path) | |
| public_id = api.codec.encode({"kind": "dataset", "asset_id": dataset_asset.id}) | |
| listed = json.loads(api.route("GET", "/api/v1/datasets").body.decode("utf-8")) | |
| page = json.loads(api.route("GET", f"/api/v1/datasets/{public_id}/items", "page=1&page_size=2").body.decode("utf-8")) | |
| assert listed["datasets"][0]["name"] == "Minecraft Steve" | |
| assert "path" not in listed["datasets"][0] | |
| assert page["pagination"]["total"] == 3 | |
| assert len(page["items"]) == 2 | |
| item = page["items"][0] | |
| assert "path" not in item | |
| assert item["caption"] == "caption 0\n" | |
| caption = json.loads(api.route( | |
| "POST", | |
| f"/api/v1/datasets/{public_id}/items/{item['id']}/caption", | |
| payload={"caption": "new caption"}, | |
| ).body.decode("utf-8")) | |
| decision = json.loads(api.route( | |
| "POST", | |
| f"/api/v1/datasets/{public_id}/items/{item['id']}/decision", | |
| payload={"decision": "reject"}, | |
| ).body.decode("utf-8")) | |
| assert caption["item"]["caption"] == "new caption" | |
| assert (dataset / "image_0.txt").read_text(encoding="utf-8") == "new caption\n" | |
| assert decision["item"]["decision"] == "reject" | |
| def test_dataset_registry_discovers_registered_locations_into_remote(tmp_path: Path) -> None: | |
| location = tmp_path / "Remembered" | |
| dataset = location / "Minecraft Oasis V3" | |
| _image(dataset / "frame_0001.png") | |
| registry = DatasetRegistry(tmp_path, _config(tmp_path)) | |
| registry.register_location(location, name="Oasis datasets") | |
| api = _remote_v1(tmp_path) | |
| listed = json.loads(api.route("GET", "/api/v1/datasets").body.decode("utf-8")) | |
| locations = json.loads(api.route("GET", "/api/v1/datasets/locations").body.decode("utf-8")) | |
| assert listed["datasets"][0]["name"] == "Minecraft Oasis V3" | |
| assert listed["datasets"][0]["available"] is True | |
| assert listed["datasets"][0]["thumbnail_url"] | |
| assert "path" not in listed["datasets"][0] | |
| assert locations["locations"][0]["name"] == "Oasis datasets" | |
| assert "path" not in locations["locations"][0] | |
| def test_remote_caption_cannot_escape_dataset(tmp_path: Path, monkeypatch) -> None: | |
| dataset = tmp_path / 'dataset' | |
| _image(dataset / 'image.png') | |
| assets = AssetRegistry(tmp_path) | |
| asset = assets.register(kind='dataset', name='Test', path=str(dataset)) | |
| api = _remote_v1(tmp_path) | |
| dataset_id = api.codec.encode({'kind': 'dataset', 'asset_id': asset.id}) | |
| item_id = api.media.media_id(kind='dataset_image', asset_id=asset.id, index=0) | |
| outside = tmp_path / 'private.txt' | |
| outside.write_text('private', encoding='utf-8') | |
| monkeypatch.setattr('adam.remote_v1.caption_path', lambda _path: outside) | |
| response = api.route('POST', f'/api/v1/datasets/{dataset_id}/items/{item_id}/caption', payload={'caption': 'overwritten'}) | |
| assert response.status == 403 | |
| assert outside.read_text(encoding='utf-8') == 'private' | |
| assert api.route('GET', f'/api/v1/datasets/{dataset_id}/items').status == 403 | |
| api.dispatcher.shutdown() | |
| def test_remote_caption_replaces_hard_link_without_overwriting_target(tmp_path: Path) -> None: | |
| import os | |
| dataset = tmp_path / 'dataset' | |
| _image(dataset / 'image.png') | |
| outside = tmp_path / 'private.txt' | |
| outside.write_text('private', encoding='utf-8') | |
| os.link(outside, dataset / 'image.txt') | |
| assets = AssetRegistry(tmp_path) | |
| asset = assets.register(kind='dataset', name='Test', path=str(dataset)) | |
| api = _remote_v1(tmp_path) | |
| dataset_id = api.codec.encode({'kind': 'dataset', 'asset_id': asset.id}) | |
| item_id = api.media.media_id(kind='dataset_image', asset_id=asset.id, index=0) | |
| response = api.route('POST', f'/api/v1/datasets/{dataset_id}/items/{item_id}/caption', payload={'caption': 'new caption'}) | |
| assert response.status == 200 | |
| assert outside.read_text(encoding='utf-8') == 'private' | |
| assert (dataset / 'image.txt').read_text(encoding='utf-8') == 'new caption\n' | |
| api.dispatcher.shutdown() | |
| def test_remote_dataset_favorite_and_use_are_persistent_without_paths(tmp_path: Path) -> None: | |
| dataset = tmp_path / "datasets" / "Minecraft Oasis V3" | |
| _image(dataset / "frame_0001.png") | |
| assets = AssetRegistry(tmp_path) | |
| asset = assets.register(kind="dataset", name="Minecraft Oasis V3", path=str(dataset)) | |
| api = _remote_v1(tmp_path) | |
| dataset_id = api.codec.encode({"kind": "dataset", "asset_id": asset.id}) | |
| favorite = json.loads(api.route( | |
| "POST", | |
| f"/api/v1/datasets/{dataset_id}/favorite", | |
| payload={"favorite": True}, | |
| ).body.decode("utf-8")) | |
| used = json.loads(api.route( | |
| "POST", | |
| f"/api/v1/datasets/{dataset_id}/use", | |
| payload={}, | |
| ).body.decode("utf-8")) | |
| assert favorite["dataset"]["favorite"] is True | |
| assert used["dataset"]["last_used_at"] | |
| registry = DatasetRegistry(tmp_path, _config(tmp_path)) | |
| record = registry.record_for_path(dataset) | |
| assert record.favorite is True | |
| assert record.last_used_at | |
| def test_remote_v1_opaque_item_id_cannot_cross_datasets(tmp_path: Path) -> None: | |
| first = tmp_path / "first" | |
| second = tmp_path / "second" | |
| _image(first / "a.png") | |
| _image(second / "b.png") | |
| assets = AssetRegistry(tmp_path) | |
| one = assets.register(kind="dataset", name="One", path=str(first)) | |
| two = assets.register(kind="dataset", name="Two", path=str(second)) | |
| api = _remote_v1(tmp_path) | |
| first_id = api.codec.encode({"kind": "dataset", "asset_id": one.id}) | |
| wrong_dataset = api.codec.encode({"kind": "dataset", "asset_id": two.id}) | |
| item_id = api.media.media_id(kind="dataset_image", asset_id=one.id, index=0) | |
| response = api.route( | |
| "POST", | |
| f"/api/v1/datasets/{wrong_dataset}/items/{item_id}/decision", | |
| payload={"decision": "keep"}, | |
| ) | |
| assert response.status == 403 | |
| assert first_id | |
| def test_remote_thumbnail_cache_reuses_and_invalidates_changed_source(tmp_path: Path) -> None: | |
| source = tmp_path / "image.png" | |
| _image(source, (10, 20, 30)) | |
| media = RemoteMediaStore(tmp_path, OpaqueIdCodec("cache-test")) | |
| first = media.thumbnail(source, size=180) | |
| second = media.thumbnail(source, size=180) | |
| time.sleep(0.02) | |
| _image(source, (200, 40, 30)) | |
| third = media.thumbnail(source, size=180) | |
| assert first.path == second.path | |
| assert second.cache_hit is True | |
| assert third.path != first.path | |
| assert third.cache_hit is False | |
| def test_remote_v1_models_include_lora_trigger_word_without_paths(tmp_path: Path) -> None: | |
| model = tmp_path / "models" / "Adam_OC_LoRA_v2" | |
| model.mkdir(parents=True) | |
| checkpoint = model / "adam.safetensors" | |
| checkpoint.write_bytes(b"weights") | |
| assets = AssetRegistry(tmp_path) | |
| assets.register( | |
| kind="model", | |
| name="Adam_OC_LoRA_v2", | |
| path=str(model), | |
| trainer="lora", | |
| checkpoint=str(checkpoint), | |
| metadata={"trigger_word": "adam_oc"}, | |
| ) | |
| api = _remote_v1(tmp_path) | |
| payload = json.loads(api.route("GET", "/api/v1/models").body.decode("utf-8")) | |
| assert payload["models"][0]["trigger_word"] == "adam_oc" | |
| assert "path" not in payload["models"][0] | |
| assert payload["models"][0]["checkpoint_name"] == "adam.safetensors" | |
| def test_structured_generation_queues_existing_generation_plan(tmp_path: Path) -> None: | |
| model = tmp_path / "ddpm" / "Model" | |
| model.mkdir(parents=True) | |
| (model / "model_index.json").write_text("{}", encoding="utf-8") | |
| assets = AssetRegistry(tmp_path) | |
| model_asset = assets.register(kind="model", name="Minecraft", path=str(model), trainer="ddpm") | |
| class Jobs: | |
| def __init__(self) -> None: | |
| self.jobs = [] | |
| self.active_job = None | |
| def submit(self, plan: ExecutionPlan) -> Job: | |
| job = Job(plan=plan, status=JobStatus.QUEUED) | |
| self.jobs.insert(0, job) | |
| return job | |
| api = _remote_v1(tmp_path, jobs=Jobs()) | |
| model_id = api.codec.encode({"kind": "model", "asset_id": model_asset.id}) | |
| response = json.loads(api.route( | |
| "POST", | |
| "/api/v1/generation/start", | |
| payload={ | |
| "provider_id": "ddpm_generator", | |
| "model_id": model_id, | |
| "prompt": "Minecraft", | |
| "image_count": 1, | |
| "steps": 20, | |
| "seed": 5, | |
| "sampler": "DDIM", | |
| "aspect_ratio": "1:1 (Square)", | |
| }, | |
| ).body.decode("utf-8")) | |
| assert response["job_id"] | |
| assert response["plan"]["steps"][0]["tool_id"] == "ddpm_generator" | |
| def test_lora_trigger_word_survives_planner_command_job_and_experiment(tmp_path: Path) -> None: | |
| dataset = tmp_path / "dataset" | |
| dataset.mkdir() | |
| for index in range(2): | |
| _image(dataset / f"{index}.png") | |
| caption_path(dataset / f"{index}.png").write_text("adam_oc\n", encoding="utf-8") | |
| lora_root = tmp_path / "lora" | |
| (lora_root / "output").mkdir(parents=True) | |
| base = tmp_path / "base.safetensors" | |
| base.write_bytes(b"base") | |
| config = _config(tmp_path, {"tool_folders": {"lora_trainer": str(lora_root)}}) | |
| registry = ToolRegistry(Path.cwd()) | |
| planner = Planner(tmp_path, registry, config) | |
| planner.assets.register(kind="dataset", name="Adam Dataset", path=str(dataset)) | |
| request = ( | |
| "From the Adam Dataset dataset, train a LoRA model for 3 epochs. " | |
| "Name the model Adam_OC_LoRA_v2. " | |
| f"[ADAM_TRAINING_OPTIONS:{{\"base_model\":{json.dumps(str(base))},\"trigger_word\":\"adam_oc\"}}] " | |
| "[ADAM_TRAINER:lora]" | |
| ) | |
| plan = planner.plan(request) | |
| args = plan.steps[0].arguments | |
| job = Job(plan=plan, status=JobStatus.FINISHED, output_folder=args["output_dir"]) | |
| run = planner.assets | |
| experiment = __import__("adam.experiment_tracker", fromlist=["ExperimentStore"]).ExperimentStore(tmp_path).record_job(job) | |
| assert args["model_name"] == "Adam_OC_LoRA_v2" | |
| assert args["trigger_word"] == "adam_oc" | |
| assert experiment is not None | |
| assert experiment.trigger_word == "adam_oc" | |
| assert run | |
| def test_lora_adapter_passes_explicit_trigger_word_to_native_payload(tmp_path: Path) -> None: | |
| trainer = tmp_path / "trainer" | |
| backend = trainer / "src" / "loratrainer" / "trainer" | |
| model_pkg = trainer / "src" / "loratrainer" / "models" | |
| backend.mkdir(parents=True) | |
| model_pkg.mkdir(parents=True) | |
| for package in (trainer / "src" / "loratrainer", backend, model_pkg): | |
| (package / "__init__.py").write_text("", encoding="utf-8") | |
| (model_pkg / "training_config.py").write_text( | |
| "from dataclasses import dataclass\n" | |
| "from pathlib import Path\n" | |
| "@dataclass\n" | |
| "class TrainingConfig:\n" | |
| " dataset_dir: Path\n" | |
| " base_model_path: Path\n" | |
| " output_dir: Path\n" | |
| " resume_checkpoint: Path | None = None\n" | |
| " trigger_word: str = ''\n" | |
| " epochs: int = 1\n", | |
| encoding="utf-8", | |
| ) | |
| (backend / "diffusers_sdxl_lora_backend.py").write_text( | |
| "import json\n" | |
| "class DiffusersSDXLLoRABackend:\n" | |
| " def train(self, config, control, progress):\n" | |
| " config.output_dir.mkdir(parents=True, exist_ok=True)\n" | |
| " (config.output_dir / 'payload.json').write_text(json.dumps({'trigger_word': config.trigger_word, 'epochs': config.epochs}), encoding='utf-8')\n" | |
| " final = config.output_dir / 'final.safetensors'\n" | |
| " final.write_bytes(b'weights')\n" | |
| " return final\n", | |
| encoding="utf-8", | |
| ) | |
| dataset = tmp_path / "dataset" | |
| for index in range(2): | |
| _image(dataset / f"{index}.png") | |
| caption_path(dataset / f"{index}.png").write_text("adam_oc\n", encoding="utf-8") | |
| base = tmp_path / "base.safetensors" | |
| base.write_bytes(b"base") | |
| _config(tmp_path, {"tool_folders": {"lora_trainer": str(trainer)}}) | |
| context = ToolContext( | |
| root=tmp_path, | |
| job_id="LORA1", | |
| tool="lora_trainer", | |
| cancel_event=threading.Event(), | |
| run_event=threading.Event(), | |
| progress_callback=lambda *_args, **_kwargs: None, | |
| log_callback=lambda _message: None, | |
| preview_callback=lambda _payload: None, | |
| ) | |
| context.run_event.set() | |
| result = train_lora( | |
| context, | |
| dataset_dir=str(dataset), | |
| model_name="Adam_OC_LoRA_v2", | |
| trigger_word="adam_oc", | |
| epochs=1, | |
| output_dir=str(trainer / "output" / "Adam_OC_LoRA_v2"), | |
| base_model=str(base), | |
| ) | |
| payload = json.loads((Path(result["output_folder"]) / "payload.json").read_text(encoding="utf-8")) | |
| assert payload["trigger_word"] == "adam_oc" | |
| assert result["trigger_word"] == "adam_oc" | |
| assert result["assets"][0]["metadata"]["trigger_word"] == "adam_oc" | |
| def test_remote_v1_routes_are_authenticated_and_legacy_status_is_redacted(tmp_path: Path) -> None: | |
| assets = AssetRegistry(tmp_path) | |
| dataset = tmp_path / "dataset" | |
| _image(dataset / "a.png") | |
| assets.register(kind="dataset", name="Dataset", path=str(dataset)) | |
| class Config: | |
| root = tmp_path | |
| values = {} | |
| def get(self, key, default=None): | |
| return self.values.get(key, default) | |
| def update(self, values): | |
| self.values.update(values) | |
| asset_registry = assets | |
| class PlannerStub: | |
| root = tmp_path | |
| registry = ToolRegistry(Path.cwd()) | |
| assets = asset_registry | |
| class Jobs: | |
| def __init__(self) -> None: | |
| self.active_job = None | |
| self.jobs = [ | |
| Job( | |
| plan=ExecutionPlan("run", "Run", [PlanStep("preview_generator", "Preview", "Preview")]), | |
| status=JobStatus.QUEUED, | |
| output_folder=str(tmp_path / "secret" / "output"), | |
| ) | |
| ] | |
| service = RemoteAccessService(Config(), Jobs(), monitor=None, planner=PlannerStub()) | |
| token = service.settings()["token"] | |
| service.save_settings({"enabled": True, "port": 0, "token": token}) | |
| import socket | |
| with socket.socket() as sock: | |
| sock.bind(("127.0.0.1", 0)) | |
| port = sock.getsockname()[1] | |
| service.save_settings({"enabled": True, "port": port, "token": token}) | |
| try: | |
| service.start() | |
| try: | |
| urlopen(f"http://127.0.0.1:{port}/api/v1/datasets", timeout=3) | |
| except HTTPError as exc: | |
| assert exc.code == 401 | |
| else: | |
| raise AssertionError("v1 route should require authentication") | |
| payload = json.loads(urlopen(f"http://127.0.0.1:{port}/api/status?token={token}", timeout=3).read().decode("utf-8")) | |
| datasets = json.loads(urlopen(f"http://127.0.0.1:{port}/api/v1/datasets?token={token}", timeout=3).read().decode("utf-8")) | |
| finally: | |
| service.stop() | |
| assert payload["queue"][0]["output_folder"] == "output" | |
| assert str(tmp_path) not in json.dumps(payload) | |
| assert datasets["datasets"][0]["name"] == "Dataset" | |