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