orena-procedure-runtime / test_runtime.py
gwd200's picture
Publish PROCEDURE ep15 merged3225 self-contained runtime v1
f6588d9 verified
Raw History Blame Contribute Delete
7.66 kB
from __future__ import annotations
import hashlib
import json
import os
from fractions import Fraction
from pathlib import Path
from types import SimpleNamespace
import pytest
from bundle import infer
from bundle import verify_model
from bundle.verify_model import (
EXPECTED_FILE_COUNT,
EXPECTED_MANIFEST_SHA256,
EXPECTED_TOTAL_BYTES,
)
ROOT = Path(__file__).resolve().parent
def test_shipped_model_manifest_identity() -> None:
manifest = ROOT / "model.manifest.tsv"
assert hashlib.sha256(manifest.read_bytes()).hexdigest() == EXPECTED_MANIFEST_SHA256
rows = [line.split("\t") for line in manifest.read_text().splitlines()]
assert len(rows) == EXPECTED_FILE_COUNT
assert [row[0] for row in rows] == sorted(row[0] for row in rows)
assert sum(int(row[1]) for row in rows) == EXPECTED_TOTAL_BYTES
@pytest.mark.parametrize(
("value", "expected"),
[
(Fraction(0), 0),
(Fraction(1, 2), 500),
(Fraction(1, 2000), 1),
(Fraction(1, 3000), 0),
],
)
def test_fraction_milliseconds_rounds_half_up(value: Fraction, expected: int) -> None:
assert infer._fraction_milliseconds(value) == expected
def test_preprocess_parser_requires_persistent_work_dir() -> None:
parser = infer.build_parser()
with pytest.raises(SystemExit):
parser.parse_args(
[
"preprocess",
"--video",
"video.mp4",
"--question",
"When?",
"--output",
"out.json",
]
)
def test_llm_input_preserves_video_metadata_tuple() -> None:
video = ("pixels", {"fps": 1000.0, "frames_indices": [100, 300]})
value = infer._llm_input(
{
"input_ids": [1, 2],
"videos": [video],
"mm_processor_kwargs": {"do_sample_frames": False},
}
)
assert value == {
"prompt_token_ids": [1, 2],
"multi_modal_data": {"video": [video]},
"mm_processor_kwargs": {"do_sample_frames": False},
}
def test_generate_uses_frozen_engine_and_sampling(
monkeypatch: pytest.MonkeyPatch,
) -> None:
observed: dict[str, object] = {}
class FakeCompletion:
text = "answer"
token_ids = [7]
finish_reason = "stop"
class FakeRequest:
request_id = "0"
prompt_token_ids = [1, 2, 3]
outputs = [FakeCompletion()]
class FakeLLM:
def __init__(self, **kwargs):
observed["llm"] = kwargs
def generate(self, inputs, sampling, *, use_tqdm):
observed["inputs"] = inputs
observed["sampling"] = sampling.kwargs
observed["use_tqdm"] = use_tqdm
return [FakeRequest()]
class FakeSamplingParams:
def __init__(self, **kwargs):
self.kwargs = kwargs
fake_vllm = SimpleNamespace(LLM=FakeLLM, SamplingParams=FakeSamplingParams)
real_import = infer.importlib.import_module
def fake_import(name: str):
if name == "vllm":
return fake_vllm
return real_import(name)
monkeypatch.setattr(infer.importlib, "import_module", fake_import)
encoded = {
"input_ids": [11, 12],
"videos": ["video"],
"mm_processor_kwargs": {},
}
report = infer._generate(
encoded,
Path("/model"),
)
assert report["answer"] == "answer"
assert observed["llm"] == {
"model": "/model",
"dtype": "bfloat16",
"max_model_len": 32768,
"max_num_seqs": 1,
"limit_mm_per_prompt": {"image": 0, "video": 1},
"enforce_eager": True,
"gpu_memory_utilization": 0.9,
"seed": 43,
"gdn_prefill_backend": "triton",
}
assert observed["sampling"] == {
"max_tokens": 128,
"temperature": 0.0,
"seed": 43,
}
assert observed["use_tqdm"] is False
assert json.loads(json.dumps(report))["finish_reason"] == "stop"
def test_infer_parser_rejects_science_overrides() -> None:
parser = infer.build_parser()
with pytest.raises(SystemExit):
parser.parse_args(
[
"infer",
"--video",
"video.mp4",
"--question",
"When?",
"--output",
"out.json",
"--model",
"model",
"--model-manifest",
"model.manifest.tsv",
"--seed",
"99",
]
)
def test_atomic_json_is_no_clobber(tmp_path: Path) -> None:
output = tmp_path / "nested" / "result.json"
infer._atomic_json(output, {"status": "first"})
assert json.loads(output.read_text()) == {"status": "first"}
assert os.stat(output).st_nlink == 1
with pytest.raises(FileExistsError):
infer._atomic_json(output, {"status": "second"})
assert json.loads(output.read_text()) == {"status": "first"}
def test_model_verifier_accepts_only_exact_inventory(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
model = tmp_path / "model"
model.mkdir()
(model / "a.bin").write_bytes(b"alpha")
manifest = tmp_path / "model.manifest.tsv"
manifest.write_text(
f"a.bin\t5\t{hashlib.sha256(b'alpha').hexdigest()}\n",
encoding="utf-8",
)
monkeypatch.setattr(
verify_model,
"EXPECTED_MANIFEST_SHA256",
hashlib.sha256(manifest.read_bytes()).hexdigest(),
)
monkeypatch.setattr(verify_model, "EXPECTED_FILE_COUNT", 1)
monkeypatch.setattr(verify_model, "EXPECTED_TOTAL_BYTES", 5)
report = verify_model.verify_model(model, manifest)
assert report["status"] == "PASS"
assert report["file_count"] == 1
(model / "extra.bin").write_bytes(b"extra")
with pytest.raises(RuntimeError, match="regular-file set"):
verify_model.verify_model(model, manifest)
def test_model_verifier_rejects_top_level_symlink(
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
model = tmp_path / "model"
model.mkdir()
(model / "a.bin").write_bytes(b"alpha")
alias = tmp_path / "alias"
alias.symlink_to(model, target_is_directory=True)
manifest = tmp_path / "model.manifest.tsv"
manifest.write_text(
f"a.bin\t5\t{hashlib.sha256(b'alpha').hexdigest()}\n",
encoding="utf-8",
)
monkeypatch.setattr(
verify_model,
"EXPECTED_MANIFEST_SHA256",
hashlib.sha256(manifest.read_bytes()).hexdigest(),
)
monkeypatch.setattr(verify_model, "EXPECTED_FILE_COUNT", 1)
monkeypatch.setattr(verify_model, "EXPECTED_TOTAL_BYTES", 5)
with pytest.raises(RuntimeError, match="must not be a symlink"):
verify_model.verify_model(alias, manifest)
def test_docker_and_readme_ship_the_full_runtime_contract() -> None:
dockerfile = (ROOT / "Dockerfile").read_text()
readme = (ROOT / "README.md").read_text()
normalized_readme = " ".join(readme.lower().split())
lock = (ROOT / "requirements.lock.txt").read_text()
assert "cuda12-4@sha256:d29902d885459b" in dockerfile
assert "--requirement /opt/procedure/requirements.lock.txt" in dockerfile
assert "--force-reinstall --no-deps opencv-python-headless==5.0.0.93" in dockerfile
assert "python -m bundle.verify_runtime" in dockerfile
assert "torch==2.11.0+cu130" in lock
for required_text in (
"PyAV",
"Decord",
"OpenCV",
"256",
"stride 5",
"fps=1000",
"old/uniform temporal",
"not checkpoint430",
):
assert required_text.lower() in normalized_readme