Download code/tt_diffusion_planner/tests/test_bundle_host.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 22 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tests/test_bundle_host.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/tests/test_bundle_host.py
-
curl -L -o test_bundle_host.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/tests/test_bundle_host.py
22 kB
| # SPDX-License-Identifier: Apache-2.0 | |
| """Host tests of the bundle wiring (no device, no weights): the vendored ttaw copy is intact and every module | |
| imports without ttnn / torch; the identity in api.py / device.py and the numerics knobs of tt/config.py match | |
| tt-model.yaml; the device configuration resolves like the server's; the stdlib client and smoke test work against | |
| this bundle's app served over real HTTP (ETH dispatch and 12x10 asserted, serve-profile pins, agreement with a stored | |
| reference output); the pip project is the repo-root pyproject.toml; code/scripts/container_smoke.sh (run against fake | |
| tt-model / docker tools) keeps its evidence and always stops the container. | |
| TT_VISIBLE_DEVICES=none python -m pytest -q code/tt_diffusion_planner/tests/test_bundle_host.py | |
| """ | |
| from __future__ import annotations | |
| import contextlib | |
| import fnmatch | |
| import hashlib | |
| import json | |
| import os | |
| import shutil | |
| import socket | |
| import subprocess | |
| import sys | |
| import threading | |
| from pathlib import Path | |
| from typing import Iterator | |
| import numpy as np | |
| import pytest | |
| import tt_diffusion_planner | |
| from tt_diffusion_planner import device as tdev | |
| from tt_diffusion_planner.api import DiffusionPlanner | |
| from tt_diffusion_planner.tests.stubs import ETH_P150, StubModel, sample_inputs | |
| PKG_DIR = Path(tt_diffusion_planner.__file__).resolve().parent | |
| CODE_DIR, BUNDLE_DIR, TTAW_DIR = PKG_DIR.parent, PKG_DIR.parents[1], PKG_DIR / "ttaw" | |
| PATCH_MARKER = "single_chip_arch_1cq_no_dispatch_s" # added to topology.cpp by patches/tt-metal-eth-dispatch.patch | |
| VENDOR_IGNORE = ("__pycache__", "*.pyc", "*.pyo", ".pytest_cache", "*.bin", "*.log", ".DS_Store", "VENDORED.json") | |
| BLOCKER = """ | |
| import importlib, importlib.abc, sys | |
| BLOCKED = set(sys.argv[2].split(",")) | |
| class Block(importlib.abc.MetaPathFinder): | |
| def find_spec(self, name, path, target=None): | |
| if name.split(".")[0] in BLOCKED: | |
| raise ImportError("blocked import of " + name) | |
| sys.meta_path.insert(0, Block()) | |
| sys.path.insert(0, sys.argv[1]) | |
| for target in sys.argv[3:]: | |
| module, _, attr = target.partition(":") | |
| mod = importlib.import_module(module) | |
| if attr: | |
| getattr(mod, attr) | |
| leaked = sorted(k for k in sys.modules if k.split(".")[0] in BLOCKED) | |
| print("ok" if not leaked else "leaked: " + " ".join(leaked)) | |
| """ | |
| # ------------------------------------------------------------------------------------------ vendored ttaw | |
| def test_vendored_ttaw_matches_its_manifest(): | |
| """The copy under code/tt_diffusion_planner/ttaw is exactly what common/tools/vendor.py wrote (never edit it | |
| here).""" | |
| manifest = json.loads((TTAW_DIR / "VENDORED.json").read_text()) | |
| recorded = manifest["files"] | |
| actual = {} | |
| for p in sorted(TTAW_DIR.rglob("*")): | |
| rel = p.relative_to(TTAW_DIR) | |
| if p.is_file() and not any(fnmatch.fnmatch(part, pat) for part in rel.parts for pat in VENDOR_IGNORE): | |
| actual[rel.as_posix()] = hashlib.sha256(p.read_bytes()).hexdigest() | |
| problems = ([f"missing: {k}" for k in sorted(set(recorded) - set(actual))] | |
| + [f"extra: {k}" for k in sorted(set(actual) - set(recorded))] | |
| + [f"modified: {k}" for k in sorted(set(recorded) & set(actual)) if recorded[k] != actual[k]]) | |
| assert problems == [], "re-vendor with common/tools/vendor.py instead of editing the copy" | |
| from tt_diffusion_planner import ttaw | |
| assert manifest["package"] == "ttaw" and manifest["version"] == ttaw.__version__ | |
| _P = "tt_diffusion_planner" | |
| def test_imports_have_no_side_effects(blocked, targets): | |
| if blocked == "ttnn,torch": | |
| pytest.importorskip("fastapi") | |
| targets = targets + ["tt_diffusion_planner.server.app:app"] | |
| out = subprocess.run([sys.executable, "-c", BLOCKER, str(CODE_DIR), blocked] + targets, capture_output=True, | |
| text=True) | |
| assert out.returncode == 0 and out.stdout.strip() == "ok", out.stderr[-2000:] | |
| # ----------------------------------------------------------------------------------- identity and config | |
| def manifest(): | |
| yaml = pytest.importorskip("yaml") | |
| path = BUNDLE_DIR / "tt-model.yaml" | |
| if not path.is_file(): | |
| pytest.skip("no tt-model.yaml next to code/ (installed package or container image)") | |
| return yaml.safe_load(path.read_text()) | |
| def test_identity_matches_manifest(manifest): | |
| cls = DiffusionPlanner | |
| assert manifest["name"] == cls.MODEL_NAME and manifest["repo"].endswith("/" + cls.MODEL_NAME) | |
| w = manifest["weights"] | |
| assert (w["repo"], w["revision"]) == (cls.DEFAULT_REPO, cls.DEFAULT_REVISION) | |
| assert list(w.get("allow_patterns") or []) == list(cls.ALLOW_PATTERNS or []) | |
| assert manifest["runtime"]["app"] == "tt_diffusion_planner.server.app:app" | |
| env = manifest["serve"]["env"] | |
| assert env["TT_WEIGHTS_REVISION"] == cls.DEFAULT_REVISION | |
| assert int(env["DIFFUSION_PLANNER_NUM_CQS"]) == tdev.DEVICE_DEFAULTS["num_command_queues"] | |
| for profile in [{"name": "serve", "env": {}}] + list(manifest.get("serve_profiles") or []): | |
| merged = {**env, **(profile.get("env") or {})} | |
| assert merged.get("DIFFUSION_PLANNER_DISPATCH") == "eth", f"{profile['name']}: the p150 target is ETH dispatch" | |
| assert merged.get("DIFFUSION_PLANNER_VARIANT") in cls.VARIANTS, profile["name"] | |
| verify = "\n".join(manifest["verify"]) | |
| assert PATCH_MARKER in verify and cls.DEFAULT_REVISION in verify | |
| def test_serve_env_pins_the_numerics(manifest): | |
| """The image serves the gated graph: ``serve.env`` pins every numerics knob of ``tt/config.py`` at its published | |
| default (``KNOBS.serve_env()``, the ttaw.knobs convention) and an empty ``DIFFUSION_PLANNER_PRECISION`` (the | |
| ``DEFAULT_PRECISION`` policy, no override rule); no serve profile overrides them (VERIFY_PORT L2).""" | |
| from tt_diffusion_planner.tt.config import KNOBS | |
| env = manifest["serve"]["env"] | |
| pins = KNOBS.serve_env() | |
| assert set(pins) == {f"DIFFUSION_PLANNER_{k}" for k in | |
| ("LN_FP32", "HIDDEN_FP32", "SPLIT_MATMUL", "ATTN_FP32_ACC", "ATTN_MATMUL", | |
| "ENC_CH2D", "ATTN_FAST", | |
| "DEC_MMCFG", "LN_KERNEL", "LN_RESID", "SPLIT_KCAT", "LN_SFPU_BCAST", "LN_SPLIT", "KCAT_ACT", "ATTN_SMASK", | |
| "ATTN_SMSM", "ENC_KCAT", "LIN_ACT", "ATTN_FUSED", "KCAT_EMIT", "LN_TR", "KCAT_L1", "KCAT_ACT_ONCE", "ATTN_L1", "DEC_L1", "ENC_L1", "FUS_L1", | |
| "COMPACT", "AGENT_BUCKETS", "HOST_FAST", "INPUT_TRIM")} | |
| assert {k: env.get(k) for k in pins} == pins | |
| assert env.get("DIFFUSION_PLANNER_PRECISION") == "" | |
| for profile in manifest.get("serve_profiles") or []: | |
| assert not set(profile.get("env") or {}) & (set(pins) | {"DIFFUSION_PLANNER_PRECISION"}), profile["name"] | |
| def test_device_config_resolution(monkeypatch): | |
| for name in ("DISPATCH", "NUM_CQS", "L1_SMALL", "TRACE_REGION", "WORKER_L1_SIZE"): | |
| monkeypatch.delenv(f"{tdev.ENV_PREFIX}_{name}", raising=False) | |
| monkeypatch.delenv("TT_DEVICE_ID", raising=False) | |
| cfg = tdev.device_config() | |
| assert (cfg.dispatch, cfg.device_id, cfg.allow_fallback) == ("eth", 0, True) | |
| assert {k: getattr(cfg, k) for k in tdev.DEVICE_DEFAULTS} == tdev.DEVICE_DEFAULTS | |
| assert cfg == DiffusionPlanner.device_config() | |
| monkeypatch.setenv("DIFFUSION_PLANNER_DISPATCH", "worker") | |
| monkeypatch.setenv("DIFFUSION_PLANNER_TRACE_REGION", str(1 << 20)) | |
| monkeypatch.setenv("TT_DEVICE_ID", "1") | |
| cfg = tdev.device_config(num_command_queues=2, worker_l1_size=None) | |
| assert (cfg.dispatch, cfg.trace_region_size, cfg.device_id, cfg.num_command_queues) == ("worker", 1 << 20, 1, 2) | |
| assert cfg == DiffusionPlanner.device_config(num_command_queues=2) | |
| assert tdev.device_config(dispatch="eth").dispatch == "eth" # an explicit argument beats the environment | |
| # ------------------------------------------------------------------------------- client and smoke test | |
| def test_client_script_runs_standalone(tmp_path): | |
| """server/client.py runs with any python3 in isolated mode: no package install, no numpy needed.""" | |
| np.savez(tmp_path / "scene.npz", **sample_inputs()) | |
| subprocess.run([sys.executable, "-I", str(PKG_DIR / "server" / "client.py"), | |
| "--inputs", str(tmp_path / "scene.npz"), | |
| "--param", "stopping_threshold=0.5", "--out", str(tmp_path / "req.json")], check=True, | |
| capture_output=True, text=True) | |
| req = json.loads((tmp_path / "req.json").read_text()) | |
| assert req["inputs"]["format"] == "npz" and req["params"] == {"stopping_threshold": 0.5} | |
| class WorkerStub(StubModel): | |
| """A server whose ETH open fell back to WORKER dispatch.""" | |
| device_info = dict(ETH_P150, dispatch="worker", fallback="RuntimeError: no ETH", grid="11x10", grid_x=11, | |
| cores=110) | |
| def _free_port() -> int: | |
| with socket.socket() as s: | |
| s.bind(("127.0.0.1", 0)) | |
| return s.getsockname()[1] | |
| def _serve(stub_cls) -> Iterator[str]: | |
| """This bundle's real app under uvicorn on a free port with ``stub_cls`` as the model; yields the base URL.""" | |
| uvicorn = pytest.importorskip("uvicorn") | |
| pytest.importorskip("fastapi") | |
| from tt_diffusion_planner.server import app as server | |
| from tt_diffusion_planner.server import smoke_test | |
| state = server.app.state.ttaw | |
| saved, state.model_factory = state.model_factory, stub_cls | |
| srv = uvicorn.Server(uvicorn.Config(server.app, host="127.0.0.1", port=_free_port(), log_level="warning", | |
| lifespan="on")) | |
| thread = threading.Thread(target=srv.run, daemon=True) | |
| thread.start() | |
| base = f"http://127.0.0.1:{srv.config.port}" | |
| try: | |
| assert smoke_test.ttaw_client.wait_ready(base, wait_s=30, poll_s=0.1)["status"] == "ok" | |
| yield base | |
| finally: | |
| srv.should_exit = True | |
| thread.join(timeout=30) | |
| state.model_factory = saved | |
| def _staged_manifest(tmp_path: Path) -> Path: | |
| """A staged tt_kernel_manifest.json (wire format) with a default profile and one whose pins differ.""" | |
| env = {"TT_WEIGHTS_REVISION": DiffusionPlanner.DEFAULT_REVISION, "DIFFUSION_PLANNER_DISPATCH": "eth", | |
| "DIFFUSION_PLANNER_NUM_CQS": str(tdev.DEVICE_DEFAULTS["num_command_queues"]), | |
| "DIFFUSION_PLANNER_VARIANT": DiffusionPlanner.DEFAULT_VARIANT} | |
| wire = {"name": DiffusionPlanner.MODEL_NAME, | |
| "weights": {"repo_id": DiffusionPlanner.DEFAULT_REPO, "revision": DiffusionPlanner.DEFAULT_REVISION}, | |
| "container": {"serve": {"env": env}, "default_profile": None, | |
| "serve_profiles": [{"name": "default", "env": {}}, | |
| {"name": "other", "env": {"DIFFUSION_PLANNER_VARIANT": "not-a-variant"}}]}} | |
| path = tmp_path / "tt_kernel_manifest.json" | |
| path.write_text(json.dumps(wire)) | |
| return path | |
| def test_smoke_test_against_live_server(tmp_path, monkeypatch, capsys): | |
| from tt_diffusion_planner.server import smoke_test | |
| monkeypatch.setenv("TT_MESH_SHAPE", "1x1") | |
| monkeypatch.setenv("TT_MODEL_WEIGHTS_REVISION", DiffusionPlanner.DEFAULT_REVISION) | |
| for name in ("HF_MODEL", "DIFFUSION_PLANNER_VARIANT", "DIFFUSION_PLANNER_DISPATCH", "DIFFUSION_PLANNER_NUM_CQS"): | |
| monkeypatch.delenv(name, raising=False) | |
| sample = tmp_path / "sample.npz" | |
| np.savez(sample, **sample_inputs()) | |
| reference = StubModel()(inputs=str(sample)).to_dict() # the stub plays the CPU reference of the sample | |
| (tmp_path / "sample.reference.json").write_text(json.dumps(reference)) | |
| shifted = json.loads(json.dumps(reference)) | |
| for row in shifted["trajectory"]: | |
| row[0] += 2.0 | |
| (tmp_path / "shifted.json").write_text(json.dumps(shifted)) | |
| args = ["--input", str(sample), "--out", str(tmp_path / "out.json")] | |
| pinned = args + ["--manifest", str(_staged_manifest(tmp_path))] | |
| with _serve(StubModel) as url: | |
| assert smoke_test.main(["--url", url] + pinned) == 0 | |
| out = capsys.readouterr().out | |
| assert out.startswith("PASS diffusion-planner-p150: profile=default") and "dispatch=eth grid=12x10" in out | |
| assert "reference=sample.reference.json (trajectory ade=0.0000 fde=0.0000" in out | |
| assert json.loads((tmp_path / "out.json").read_text())["num_poses"] == 80 | |
| assert smoke_test.main(["--url", url, "--profile", "other"] + pinned) == 1 | |
| assert "profile 'other' pins DIFFUSION_PLANNER_VARIANT=not-a-variant" in capsys.readouterr().out | |
| assert smoke_test.main(["--url", url, "--reference", str(tmp_path / "shifted.json")] + args) == 1 | |
| assert "trajectory ADE 2.0000 > 0.3" in capsys.readouterr().out | |
| assert smoke_test.main(["--url", url, "--reference", str(tmp_path / "missing.json")] + args) == 1 | |
| assert "reference " in capsys.readouterr().out | |
| assert smoke_test.main(["--url", url, "--input", str(tmp_path / "missing.npz")]) == 1 | |
| assert "/predict: FileNotFoundError" in capsys.readouterr().out # a FAIL line, not a traceback | |
| with _serve(WorkerStub) as url: | |
| assert smoke_test.main(["--url", url] + args) == 1 | |
| out = capsys.readouterr().out | |
| assert out.startswith("FAIL diffusion-planner-p150:") | |
| assert "dispatch is 'worker', expected 'eth'" in out and "grid is '11x10', expected '12x10'" in out | |
| # ------------------------------------------------------------------------- pip project and container smoke | |
| def test_pip_project_at_repo_root(): | |
| """The pip project is the repo-root pyproject.toml over code/ (never code/pyproject.toml: tt-model copies code/ | |
| over the tt-metal tree before building ttnn, PACKAGING_PILOT.md problem 1).""" | |
| tomllib = pytest.importorskip("tomllib") | |
| path = BUNDLE_DIR / "pyproject.toml" | |
| if not (BUNDLE_DIR / "tt-model.yaml").is_file(): | |
| pytest.skip("not a bundle source checkout (installed package or container image)") | |
| assert not (CODE_DIR / "pyproject.toml").exists(), "the pip project lives at the repo root, never in code/" | |
| project = tomllib.loads(path.read_text()) | |
| find = project["tool"]["setuptools"]["packages"]["find"] | |
| assert find["where"] == ["code"] and any(fnmatch.fnmatch("tt_diffusion_planner", pat) for pat in find["include"]) | |
| data = project["tool"]["setuptools"]["package-data"] | |
| assert {"API.md", "VENDORED.json"} <= set(data["tt_diffusion_planner.ttaw"]) | |
| assert {"kernels/*.cpp", "kernels/*.hpp", "kernels/*.h"} <= set(data["tt_diffusion_planner.ttaw.ops"]) | |
| assert (BUNDLE_DIR / project["project"]["readme"]).is_file() | |
| assert project["tool"]["setuptools"]["dynamic"]["version"]["attr"] == "tt_diffusion_planner.__version__" | |
| yaml = pytest.importorskip("yaml") | |
| shipped = [p for e in yaml.safe_load((BUNDLE_DIR / "tt-model.yaml").read_text())["source"]["extra_code"] | |
| for p in e["paths"]] | |
| assert "pyproject.toml" not in shipped and "setup.py" not in shipped | |
| setuptools = pytest.importorskip("setuptools") | |
| found = set(setuptools.find_packages(where=str(CODE_DIR), include=find["include"], exclude=find.get("exclude", []))) | |
| assert {"tt_diffusion_planner", "tt_diffusion_planner.ttaw", "tt_diffusion_planner.ttaw.server", | |
| "tt_diffusion_planner.ttaw.ops", "tt_diffusion_planner.host", "tt_diffusion_planner.reference"} <= found | |
| assert not any(p.startswith("tt_diffusion_planner.tests") for p in found) | |
| FAKE_TT_MODEL = r''' | |
| import json, os, signal, subprocess, sys, time, urllib.request | |
| from pathlib import Path | |
| state = Path(os.environ["FAKE_STATE"]) | |
| args = sys.argv[1:] | |
| with open(state / "calls.txt", "a") as f: | |
| f.write(" ".join(args) + "\n") | |
| opt = lambda name: args[args.index(name) + 1] if name in args else None | |
| if args[0] == "serve": | |
| port, manifest = int(opt("--port")), json.load(open(args[-1])) | |
| (state / "container").write_text("tt-model-" + manifest["name"] + "-" + (opt("--profile") or "default")) | |
| (state / "container.log").write_text("boot: Opening device 0 (dispatch eth, 1 CQ)\n") | |
| if os.environ.get("FAKE_SERVE_FAIL"): | |
| sys.exit("fake serve: the boot failed") | |
| srv = subprocess.Popen([sys.executable, str(state / "server.py"), str(port)], start_new_session=True) | |
| (state / "server.pid").write_text(str(srv.pid)) | |
| for _ in range(200): | |
| try: | |
| urllib.request.urlopen("http://127.0.0.1:%d/health" % port, timeout=1) | |
| break | |
| except OSError: | |
| time.sleep(0.05) | |
| elif args[0] == "stop": | |
| if (state / "server.pid").exists(): | |
| os.kill(int((state / "server.pid").read_text()), signal.SIGTERM) | |
| with open(state / "container.log", "a") as f: | |
| f.write("shutdown: traces released, device closed\n") | |
| (state / "stopped").write_text("1") | |
| elif args[0] == "logs": | |
| print("fallback: tt-model logs") | |
| else: | |
| sys.exit(2) | |
| ''' | |
| FAKE_DOCKER = r''' | |
| import os, sys, time | |
| from pathlib import Path | |
| state = Path(os.environ["FAKE_STATE"]) | |
| args = sys.argv[1:] | |
| with open(state / "docker_calls.txt", "a") as f: | |
| f.write(" ".join(args) + "\n") | |
| known = (state / "container").read_text() if (state / "container").exists() else None | |
| if args[0] == "inspect": | |
| sys.exit(0 if args[-1] == known else 1) | |
| if args[0] == "logs" and args[-1] == known: | |
| deadline = time.time() + 20 | |
| while "--follow" in args and not (state / "stopped").exists() and time.time() < deadline: | |
| time.sleep(0.05) | |
| sys.stdout.write((state / "container.log").read_text()) | |
| sys.exit(0) | |
| sys.exit(1) | |
| ''' | |
| FAKE_SERVER = r''' | |
| import json, sys | |
| from http.server import BaseHTTPRequestHandler, HTTPServer | |
| INFO = {"model": "fake", "device": {"dispatch": "eth", "grid": "12x10", "cores": 120}} | |
| class Handler(BaseHTTPRequestHandler): | |
| def _send(self, code, body): | |
| data = json.dumps(body).encode() | |
| self.send_response(code) | |
| self.send_header("Content-Type", "application/json") | |
| self.send_header("Content-Length", str(len(data))) | |
| self.end_headers() | |
| self.wfile.write(data) | |
| def do_GET(self): | |
| self._send(200, {"status": "ok"} if self.path.startswith("/health") else INFO if self.path == "/info" else {}) | |
| def do_POST(self): | |
| self._send(500, {"detail": "fake server"}) | |
| def log_message(self, *args): | |
| pass | |
| HTTPServer(("127.0.0.1", int(sys.argv[1])), Handler).serve_forever() | |
| ''' | |
| def _fake_tools(tmp_path: Path) -> dict: | |
| """PATH with fake `tt-model` / `docker` / `python3` (the test interpreter) and a stub server: no docker, no chip.""" | |
| state, bin_dir = tmp_path / "state", tmp_path / "bin" | |
| state.mkdir() | |
| bin_dir.mkdir() | |
| (state / "server.py").write_text(FAKE_SERVER) | |
| for name, body in (("tt-model", FAKE_TT_MODEL), ("docker", FAKE_DOCKER)): | |
| (bin_dir / name).write_text(f"#!{sys.executable}\n{body}") | |
| (bin_dir / name).chmod(0o755) | |
| (bin_dir / "python3").write_text(f'#!/bin/sh\nexec "{sys.executable}" "$@"\n') | |
| (bin_dir / "python3").chmod(0o755) | |
| env = {k: v for k, v in os.environ.items() if k not in ("SMOKE_OUT", "HF_TOKEN", "HUGGING_FACE_HUB_TOKEN")} | |
| env.update(PATH=f"{bin_dir}:{env.get('PATH', '/usr/bin:/bin')}", FAKE_STATE=str(state)) | |
| for tool in ("tt-model", "docker", "python3"): # never the real tools: a real serve would claim the chip | |
| assert shutil.which(tool, path=env["PATH"]) == str(bin_dir / tool) | |
| return env | |
| def test_container_smoke_keeps_evidence_and_stops(tmp_path, serve_fails): | |
| """code/scripts/container_smoke.sh passes --profile through, saves /info once READY and the container log through | |
| the shutdown (whatever the outcome), and always stops the container.""" | |
| staged = tmp_path / "staged" | |
| staged.mkdir() | |
| manifest = _staged_manifest(staged) | |
| env, logs, port = _fake_tools(tmp_path), tmp_path / "evidence", _free_port() | |
| profile = [] if serve_fails else ["--profile", "other"] | |
| if serve_fails: | |
| env["FAKE_SERVE_FAIL"] = "1" | |
| r = subprocess.run(["bash", str(CODE_DIR / "scripts" / "container_smoke.sh"), str(staged), str(port), *profile, | |
| f"--log-dir={logs}"], env=env, capture_output=True, text=True, timeout=180) | |
| assert r.returncode == 1, r.stdout + r.stderr # the stub serves no valid /predict: the smoke FAILS | |
| calls = (tmp_path / "state" / "calls.txt").read_text().splitlines() | |
| assert calls[0] == " ".join(["serve", "--port", str(port), *profile, str(manifest)]) | |
| assert calls[-1] == " ".join(["stop", *profile, str(manifest)]) and len(calls) == 2 | |
| stem = "diffusion-planner-p150" + ("" if serve_fails else "-other") + "-" | |
| files = {p.name.split(".", 1)[1]: p for p in logs.iterdir()} | |
| assert all(p.name.startswith(stem) for p in files.values()), sorted(files) | |
| log = files["container.log"].read_text() | |
| assert "boot: Opening device 0" in log and "shutdown: traces released" in log # followed through the stop | |
| result = json.loads(files["result.json"].read_text()) | |
| assert (result["rc"], result["result"], result["profile"]) == (1, "FAIL", "default" if serve_fails else "other") | |
| if serve_fails: | |
| assert "info.json" not in files | |
| else: | |
| assert json.loads(files["info.json"].read_text())["device"] == {"dispatch": "eth", "grid": "12x10", | |
| "cores": 120} | |
| assert "FAIL diffusion-planner-p150:" in r.stdout | |