jw-search / backend /tests /test_runpod_scripts.py
jw-tools's picture
deploy: slim data bundle (drop face-crops/thumbnails-small) + latest main (UI redesign, LLM deps)
09d78f2 verified
Raw
History Blame Contribute Delete
59.6 kB
"""Regression tests for the RunPod helper scripts."""
from __future__ import annotations
import importlib.util
import json
import os
import subprocess
import sys
import time
from pathlib import Path
REPO_ROOT = Path(__file__).resolve().parents[2]
def _load_module(module_name: str, relative_path: str):
module_path = REPO_ROOT / relative_path
script_dir = str(module_path.parent)
added_path = False
if script_dir not in sys.path:
sys.path.insert(0, script_dir)
added_path = True
spec = importlib.util.spec_from_file_location(module_name, module_path)
module = importlib.util.module_from_spec(spec)
assert spec.loader is not None
try:
spec.loader.exec_module(module)
finally:
if added_path:
sys.path.remove(script_dir)
return module
def test_process_local_videos_times_out_and_continues(tmp_path, monkeypatch):
module = _load_module("test_process_local_videos_script", "scripts/process-local-videos.py")
class _FakeImageSearch:
def get_processed_video_keys(self):
return set()
def is_video_processed(self, natural_key):
return False
def _fake_process_video(natural_key, label, delete_video_after):
if natural_key == "video-1":
time.sleep(2)
return {"status": "success", "message": f"processed {natural_key}"}
monkeypatch.setattr(module, "ImageSearch", lambda: _FakeImageSearch())
monkeypatch.setattr(module, "get_all_local_video_keys", lambda: ["video-1", "video-2"])
monkeypatch.setattr(module, "process_video", _fake_process_video)
report_path = tmp_path / "timeout-report.json"
monkeypatch.setattr(
sys,
"argv",
[
"process-local-videos.py",
"--timeout-seconds",
"1",
"--report",
str(report_path),
],
)
exit_code = module.main()
assert exit_code == 1
report = json.loads(report_path.read_text())
assert report["timeout_seconds"] == 1
assert report["processed"] == 1
assert report["failed"] == 1
assert report["results"][0]["natural_key"] == "video-1"
assert report["results"][0]["status"] == "error"
assert "Timed out after 1s" in report["results"][0]["message"]
assert report["results"][1]["natural_key"] == "video-2"
assert report["results"][1]["status"] == "success"
def test_process_local_videos_can_pin_gpu(monkeypatch):
module = _load_module("test_process_local_videos_gpu_pin", "scripts/process-local-videos.py")
monkeypatch.delenv("CUDA_VISIBLE_DEVICES", raising=False)
monkeypatch.setattr(
sys,
"argv",
[
"process-local-videos.py",
"--gpu-id",
"1",
"--limit",
"0",
],
)
captured = {}
def _fake_load_backend_dependencies():
captured["gpu_id"] = os.environ.get("CUDA_VISIBLE_DEVICES")
module.ImageSearch = lambda: type("_Fake", (), {"get_processed_video_keys": lambda self: set()})()
module.get_all_local_video_keys = lambda: []
module.process_video = lambda *args, **kwargs: {"status": "success", "message": "ok"}
module.atomic_write_json = lambda *args, **kwargs: None
monkeypatch.setattr(module, "_load_backend_dependencies", _fake_load_backend_dependencies)
exit_code = module.main()
assert exit_code == 0
assert captured["gpu_id"] == "1"
assert os.environ["CUDA_VISIBLE_DEVICES"] == "1"
def test_process_local_videos_report_write_is_best_effort(tmp_path):
module = _load_module("test_process_local_videos_report_best_effort", "scripts/process-local-videos.py")
calls = []
def _fake_atomic_write(path, payload, indent=None):
calls.append(path)
if "search-ui-report-shadow" in path:
Path(path).parent.mkdir(parents=True, exist_ok=True)
Path(path).write_text(json.dumps(payload))
return
raise OSError(5, "Input/output error")
module.atomic_write_json = _fake_atomic_write
module._save_report(str(tmp_path / "report.json"), {"status": "ok"})
assert any("search-ui-report-shadow" in path for path in calls)
def test_process_local_videos_filters_partition_before_sharding(monkeypatch):
module = _load_module("test_process_local_videos_partition", "scripts/process-local-videos.py")
partition_module = _load_module("test_batch_partition_helper", "backend/batch_partition.py")
class _FakeImageSearch:
def get_processed_video_keys(self):
return set()
def is_video_processed(self, natural_key):
return False
processed = []
monkeypatch.setattr(module, "ImageSearch", lambda: _FakeImageSearch())
monkeypatch.setattr(
module,
"get_all_local_video_keys",
lambda: [f"video-{idx}" for idx in range(12)],
)
monkeypatch.setattr(
module,
"process_video",
lambda natural_key, label, delete_video_after: processed.append(natural_key) or {"status": "success", "message": natural_key},
)
module.atomic_write_json = lambda *args, **kwargs: None
monkeypatch.setattr(
sys,
"argv",
[
"process-local-videos.py",
"--partition-total",
"2",
"--partition-index",
"1",
"--shard-count",
"2",
"--shard-index",
"0",
],
)
exit_code = module.main()
assert exit_code == 0
expected_partition = partition_module.filter_natural_keys_for_partition(
[f"video-{idx}" for idx in range(12)],
partition_total=2,
partition_index=1,
)
expected_shard = [
natural_key
for offset, natural_key in enumerate(expected_partition)
if offset % 2 == 0
]
assert processed == expected_shard
def test_process_local_videos_global_limit_is_exact_across_shards(monkeypatch):
module = _load_module("test_process_local_videos_exact_limit", "scripts/process-local-videos.py")
class _FakeImageSearch:
def get_processed_video_keys(self):
return set()
def is_video_processed(self, natural_key):
return False
processed = []
monkeypatch.setattr(module, "ImageSearch", lambda: _FakeImageSearch())
monkeypatch.setattr(module, "get_all_local_video_keys", lambda: [f"video-{idx}" for idx in range(10)])
monkeypatch.setattr(
module,
"process_video",
lambda natural_key, label, delete_video_after: processed.append(natural_key) or {"status": "success", "message": natural_key},
)
module.atomic_write_json = lambda *args, **kwargs: None
for shard_index in (0, 1):
monkeypatch.setattr(
sys,
"argv",
[
"process-local-videos.py",
"--limit",
"5",
"--shard-count",
"2",
"--shard-index",
str(shard_index),
],
)
exit_code = module.main()
assert exit_code == 0
assert sorted(processed) == [f"video-{idx}" for idx in range(5)]
def test_process_local_videos_target_file_filters_and_preserves_order(tmp_path, monkeypatch):
module = _load_module("test_process_local_videos_target_file", "scripts/process-local-videos.py")
class _FakeImageSearch:
def get_processed_video_keys(self):
return {"video-2"}
def is_video_processed(self, natural_key):
return natural_key == "video-2"
target_file = tmp_path / "targets.txt"
target_file.write_text("video-3\nvideo-1\nvideo-2\nvideo-99\n", encoding="utf-8")
processed = []
monkeypatch.setattr(module, "ImageSearch", lambda: _FakeImageSearch())
monkeypatch.setattr(module, "get_all_local_video_keys", lambda: ["video-1", "video-2", "video-3", "video-4"])
monkeypatch.setattr(
module,
"process_video",
lambda natural_key, label, delete_video_after: processed.append(natural_key) or {"status": "success", "message": natural_key},
)
module.atomic_write_json = lambda *args, **kwargs: None
monkeypatch.setattr(
sys,
"argv",
[
"process-local-videos.py",
"--natural-key-file",
str(target_file),
],
)
exit_code = module.main()
assert exit_code == 0
assert processed == ["video-3", "video-1"]
def test_download_all_videos_removes_stale_part_files(tmp_path, monkeypatch):
module = _load_module("test_download_all_videos_script", "scripts/download-all-videos.py")
stale_dir = tmp_path / "video-1" / "720p"
stale_dir.mkdir(parents=True)
stale_part = stale_dir / "video.mp4.123.part"
stale_part.write_bytes(b"partial")
monkeypatch.setattr(module, "get_videos_dir", lambda: str(tmp_path))
monkeypatch.setattr(module, "refresh_catalog", lambda language, redownload=False: {})
monkeypatch.setattr(module, "get_existing_video_keys", lambda label: set())
report_path = tmp_path / "download-report.json"
monkeypatch.setattr(
sys,
"argv",
[
"download-all-videos.py",
"--report",
str(report_path),
],
)
exit_code = module.main()
assert exit_code == 0
assert not stale_part.exists()
report = json.loads(report_path.read_text())
assert report["partial_files_removed"] == 1
assert report["total"] == 0
def test_download_all_videos_report_write_is_best_effort(tmp_path):
module = _load_module("test_download_all_videos_report_best_effort", "scripts/download-all-videos.py")
calls = []
def _fake_atomic_write(path, payload, indent=None):
calls.append(path)
if "search-ui-report-shadow" in path:
Path(path).parent.mkdir(parents=True, exist_ok=True)
Path(path).write_text(json.dumps(payload))
return
raise OSError(5, "Input/output error")
module.atomic_write_json = _fake_atomic_write
module._save_report(str(tmp_path / "download-report.json"), {"status": "ok"})
assert any("search-ui-report-shadow" in path for path in calls)
def test_download_all_videos_skips_processed_keys_by_default(tmp_path, monkeypatch):
module = _load_module("test_download_all_videos_skip_processed", "scripts/download-all-videos.py")
monkeypatch.setattr(
module,
"refresh_catalog",
lambda language, redownload=False: {
"video-1": {"languageAgnosticNaturalKey": "video-1", "title": "Processed"},
"video-2": {"languageAgnosticNaturalKey": "video-2", "title": "Pending"},
},
)
monkeypatch.setattr(module, "get_existing_video_keys", lambda label: set())
monkeypatch.setattr(module, "find_video_file", lambda media_item, label, subtitled=False: {"label": label})
monkeypatch.setattr(module, "download_video", lambda media_item, label: f"/tmp/{media_item['languageAgnosticNaturalKey']}.mp4")
class _FakeImageSearch:
def get_processed_video_keys(self):
return {"video-1"}
monkeypatch.setattr(module, "ImageSearch", lambda: _FakeImageSearch())
report_path = tmp_path / "download-report.json"
monkeypatch.setattr(
sys,
"argv",
[
"download-all-videos.py",
"--report",
str(report_path),
],
)
exit_code = module.main()
assert exit_code == 0
report = json.loads(report_path.read_text())
assert report["total"] == 1
assert report["processed"] == 1
assert report["downloaded"] == 1
assert report["skipped_already_processed"] == 1
assert report["results"][0]["natural_key"] == "video-2"
def test_download_all_videos_applies_partition_before_limit(tmp_path, monkeypatch):
module = _load_module("test_download_all_videos_partition", "scripts/download-all-videos.py")
catalog = {
f"video-{idx}": {"languageAgnosticNaturalKey": f"video-{idx}", "title": f"Video {idx}"}
for idx in range(8)
}
monkeypatch.setattr(module, "refresh_catalog", lambda language, redownload=False: catalog)
monkeypatch.setattr(module, "get_existing_video_keys", lambda label: set())
monkeypatch.setattr(module, "find_video_file", lambda media_item, label, subtitled=False: {"label": label})
monkeypatch.setattr(module, "download_video", lambda media_item, label: f"/tmp/{media_item['languageAgnosticNaturalKey']}.mp4")
class _FakeImageSearch:
def get_processed_video_keys(self):
return set()
monkeypatch.setattr(module, "ImageSearch", lambda: _FakeImageSearch())
report_path = tmp_path / "download-partition-report.json"
monkeypatch.setattr(
sys,
"argv",
[
"download-all-videos.py",
"--partition-total",
"2",
"--partition-index",
"0",
"--limit",
"2",
"--report",
str(report_path),
],
)
exit_code = module.main()
assert exit_code == 0
report = json.loads(report_path.read_text())
assert report["partition_total"] == 2
assert report["partition_index"] == 0
assert report["processed"] == 2
def test_download_all_videos_pushes_monthly_programs_to_back(tmp_path, monkeypatch):
module = _load_module("test_download_all_videos_priority", "scripts/download-all-videos.py")
monkeypatch.setattr(
module,
"refresh_catalog",
lambda language, redownload=False: {
"pub-jwb-135_1_VIDEO": {
"languageAgnosticNaturalKey": "pub-jwb-135_1_VIDEO",
"title": "JW Broadcasting—April 2026",
"_category_key": "VODStudio",
"_subcategory": "Monthly Programs",
},
"talk-1": {
"languageAgnosticNaturalKey": "talk-1",
"title": "Talk 1",
"_category_key": "VODStudio",
"_subcategory": "Apply Bible Principles",
},
"insert-1": {
"languageAgnosticNaturalKey": "insert-1",
"title": "Insert 1",
"_category_key": "VODStudio",
"_subcategory": "Become Jehovah's Friend",
},
},
)
monkeypatch.setattr(module, "get_existing_video_keys", lambda label: set())
monkeypatch.setattr(module, "find_video_file", lambda media_item, label, subtitled=False: {"label": label})
monkeypatch.setattr(module, "download_video", lambda media_item, label: f"/tmp/{media_item['languageAgnosticNaturalKey']}.mp4")
class _FakeImageSearch:
def get_processed_video_keys(self):
return set()
monkeypatch.setattr(module, "ImageSearch", lambda: _FakeImageSearch())
report_path = tmp_path / "download-priority-report.json"
monkeypatch.setattr(
sys,
"argv",
[
"download-all-videos.py",
"--report",
str(report_path),
],
)
exit_code = module.main()
assert exit_code == 0
report = json.loads(report_path.read_text())
assert [entry["natural_key"] for entry in report["results"]] == [
"insert-1",
"talk-1",
"pub-jwb-135_1_VIDEO",
]
def test_download_all_videos_plan_only_writes_fixed_target_file(tmp_path, monkeypatch):
module = _load_module("test_download_all_videos_plan_only", "scripts/download-all-videos.py")
catalog = {
f"video-{idx}": {"languageAgnosticNaturalKey": f"video-{idx}", "title": f"Video {idx}"}
for idx in range(6)
}
monkeypatch.setattr(module, "refresh_catalog", lambda language, redownload=False: catalog)
monkeypatch.setattr(module, "get_existing_video_keys", lambda label: set())
class _FakeImageSearch:
def get_processed_video_keys(self):
return {"video-0"}
monkeypatch.setattr(module, "ImageSearch", lambda: _FakeImageSearch())
monkeypatch.setattr(module, "_cleanup_partial_downloads", lambda: 99)
target_path = tmp_path / "targets.txt"
report_path = tmp_path / "download-report.json"
monkeypatch.setattr(
sys,
"argv",
[
"download-all-videos.py",
"--partition-total",
"2",
"--partition-index",
"1",
"--limit",
"2",
"--write-target-file",
str(target_path),
"--plan-only",
"--report",
str(report_path),
],
)
exit_code = module.main()
assert exit_code == 0
planned = [line.strip() for line in target_path.read_text().splitlines() if line.strip()]
assert len(planned) == 2
report = json.loads(report_path.read_text())
assert report["total"] == 2
assert report["processed"] == 0
assert report["partial_files_removed"] == 0
def test_download_all_videos_natural_key_file_preserves_fixed_chunk(tmp_path, monkeypatch):
module = _load_module("test_download_all_videos_target_file", "scripts/download-all-videos.py")
catalog = {
f"video-{idx}": {"languageAgnosticNaturalKey": f"video-{idx}", "title": f"Video {idx}"}
for idx in range(4)
}
monkeypatch.setattr(module, "refresh_catalog", lambda language, redownload=False: catalog)
monkeypatch.setattr(module, "get_existing_video_keys", lambda label: set())
monkeypatch.setattr(module, "find_video_file", lambda media_item, label, subtitled=False: {"label": label})
monkeypatch.setattr(module, "download_video", lambda media_item, label: f"/tmp/{media_item['languageAgnosticNaturalKey']}.mp4")
class _FakeImageSearch:
def get_processed_video_keys(self):
return {"video-2"}
monkeypatch.setattr(module, "ImageSearch", lambda: _FakeImageSearch())
target_path = tmp_path / "targets.txt"
target_path.write_text("video-3\nvideo-2\nvideo-1\n", encoding="utf-8")
report_path = tmp_path / "download-report.json"
monkeypatch.setattr(
sys,
"argv",
[
"download-all-videos.py",
"--natural-key-file",
str(target_path),
"--limit",
"1",
"--partition-total",
"2",
"--partition-index",
"0",
"--report",
str(report_path),
],
)
exit_code = module.main()
assert exit_code == 0
report = json.loads(report_path.read_text())
assert report["total"] == 2
assert {result["natural_key"] for result in report["results"]} == {"video-3", "video-1"}
def test_download_all_videos_plan_only_refilters_target_file_against_processed_db(tmp_path, monkeypatch):
module = _load_module("test_download_all_videos_target_refresh", "scripts/download-all-videos.py")
catalog = {
f"video-{idx}": {"languageAgnosticNaturalKey": f"video-{idx}", "title": f"Video {idx}"}
for idx in range(4)
}
monkeypatch.setattr(module, "refresh_catalog", lambda language, redownload=False: catalog)
monkeypatch.setattr(module, "get_existing_video_keys", lambda label: set())
class _FakeImageSearch:
def get_processed_video_keys(self):
return {"video-2"}
monkeypatch.setattr(module, "ImageSearch", lambda: _FakeImageSearch())
monkeypatch.setattr(module, "_cleanup_partial_downloads", lambda: 0)
source_path = tmp_path / "targets.txt"
source_path.write_text("video-3\nvideo-2\nvideo-1\n", encoding="utf-8")
refreshed_path = tmp_path / "targets-refreshed.txt"
report_path = tmp_path / "download-plan-refresh-report.json"
monkeypatch.setattr(
sys,
"argv",
[
"download-all-videos.py",
"--natural-key-file",
str(source_path),
"--write-target-file",
str(refreshed_path),
"--plan-only",
"--report",
str(report_path),
],
)
exit_code = module.main()
assert exit_code == 0
assert refreshed_path.read_text(encoding="utf-8").splitlines() == ["video-3", "video-1"]
report = json.loads(report_path.read_text())
assert report["total"] == 2
assert report["processed"] == 0
def test_download_all_videos_retries_transient_failures(tmp_path, monkeypatch):
module = _load_module("test_download_all_videos_retry", "scripts/download-all-videos.py")
monkeypatch.setattr(
module,
"refresh_catalog",
lambda language, redownload=False: {
"video-1": {"languageAgnosticNaturalKey": "video-1", "title": "Retry me"},
},
)
monkeypatch.setattr(module, "get_existing_video_keys", lambda label: set())
monkeypatch.setattr(module, "find_video_file", lambda media_item, label, subtitled=False: {"label": label})
attempts = {"count": 0}
def _fake_download_video(media_item, label):
attempts["count"] += 1
if attempts["count"] == 1:
return None
return f"/tmp/{media_item['languageAgnosticNaturalKey']}.mp4"
class _FakeImageSearch:
def get_processed_video_keys(self):
return set()
monkeypatch.setattr(module, "download_video", _fake_download_video)
monkeypatch.setattr(module, "ImageSearch", lambda: _FakeImageSearch())
monkeypatch.setattr(module.time, "sleep", lambda *_args, **_kwargs: None)
report_path = tmp_path / "download-report.json"
monkeypatch.setattr(
sys,
"argv",
[
"download-all-videos.py",
"--report",
str(report_path),
"--download-retries",
"3",
"--retry-backoff-seconds",
"0",
],
)
exit_code = module.main()
assert exit_code == 0
assert attempts["count"] == 2
report = json.loads(report_path.read_text())
assert report["downloaded"] == 1
assert report["errors"] == 0
assert report["results"][0]["status"] == "downloaded"
def test_download_all_videos_records_final_item_error_without_crashing(tmp_path, monkeypatch):
module = _load_module("test_download_all_videos_item_error", "scripts/download-all-videos.py")
monkeypatch.setattr(
module,
"refresh_catalog",
lambda language, redownload=False: {
"video-1": {"languageAgnosticNaturalKey": "video-1", "title": "Broken"},
},
)
monkeypatch.setattr(module, "get_existing_video_keys", lambda label: set())
monkeypatch.setattr(module, "find_video_file", lambda media_item, label, subtitled=False: {"label": label})
monkeypatch.setattr(module, "download_video", lambda media_item, label: None)
class _FakeImageSearch:
def get_processed_video_keys(self):
return set()
monkeypatch.setattr(module, "ImageSearch", lambda: _FakeImageSearch())
monkeypatch.setattr(module.time, "sleep", lambda *_args, **_kwargs: None)
report_path = tmp_path / "download-report.json"
monkeypatch.setattr(
sys,
"argv",
[
"download-all-videos.py",
"--report",
str(report_path),
"--download-retries",
"2",
"--retry-backoff-seconds",
"0",
],
)
exit_code = module.main()
assert exit_code == 1
report = json.loads(report_path.read_text())
assert report["processed"] == 1
assert report["errors"] == 1
assert report["results"][0]["status"] == "error"
assert "after 2 attempts" in report["results"][0]["message"]
def test_evaluate_runpod_output_returns_nonzero_when_checks_fail(monkeypatch):
module = _load_module("test_evaluate_runpod_output_script", "scripts/evaluate-runpod-output.py")
monkeypatch.setattr(module, "_collect_phase_reports", lambda report_dir: {})
monkeypatch.setattr(
module,
"_collect_filesystem",
lambda language: {
"subtitles_root": "/tmp/subtitles",
"videos_root": "/tmp/videos",
"vtt_files": 0,
"txt_files": 0,
"mp4_files": 0,
"thumbnails": 0,
"small_thumbnails": 0,
"face_crops": 0,
},
)
monkeypatch.setattr(module, "_gather_search_samples", lambda language: {"stats": {"total_subtitles": 0}})
monkeypatch.setattr(module, "_gather_image_samples", lambda: {"stats": {"unique_videos": 0}})
monkeypatch.setattr(module, "_gather_face_samples", lambda: {"stats": {"total_faces": 0}})
monkeypatch.setattr(sys, "argv", ["evaluate-runpod-output.py"])
exit_code = module.main()
assert exit_code == 1
def test_evaluate_runpod_output_counts_face_crops_directory(tmp_path, monkeypatch):
module = _load_module("test_evaluate_runpod_output_face_crops", "scripts/evaluate-runpod-output.py")
subtitles_root = tmp_path / "subtitles"
videos_root = tmp_path / "videos"
(subtitles_root / "E").mkdir(parents=True)
face_crop_dir = videos_root / "video-1" / "720p" / "face-crops"
face_crop_dir.mkdir(parents=True)
(face_crop_dir / "frame1_face0.jpg").write_bytes(b"jpeg")
monkeypatch.setattr(module, "get_subtitles_dir", lambda: str(subtitles_root))
monkeypatch.setattr(module, "get_videos_dir", lambda: str(videos_root))
filesystem = module._collect_filesystem("E")
assert filesystem["face_crops"] == 1
def test_runpod_batch_worker_splits_process_limit_across_shards(tmp_path, monkeypatch):
module = _load_module("test_runpod_batch_worker_script", "scripts/runpod-batch-worker.py")
class _FakeProc:
def __init__(self, cmd, **kwargs):
self.cmd = cmd
def poll(self):
return 0
def wait(self):
return 0
def terminate(self):
return None
def kill(self):
return None
monkeypatch.setattr(module.subprocess, "Popen", lambda cmd, **kwargs: _FakeProc(cmd, **kwargs))
monkeypatch.setattr(module.time, "sleep", lambda *_args, **_kwargs: None)
manifest_path = tmp_path / "manifest.json"
manifest = {"manifest_path": str(manifest_path), "phases": []}
report_dir = tmp_path / "reports"
report_dir.mkdir()
module._run_processing_shards(
label="720p",
process_shards=2,
process_limit=5,
partition_total=1,
partition_index=0,
delete_after=False,
report_dir=str(report_dir),
manifest=manifest,
shard_timeout_seconds=0,
)
commands = [shard["command"] for shard in manifest["phases"][0]["shards"]]
limits = [cmd[cmd.index("--limit") + 1] for cmd in commands]
assert limits == ["5", "5"]
def test_runpod_batch_worker_passes_gpu_ids_to_shards(tmp_path, monkeypatch):
module = _load_module("test_runpod_batch_worker_gpu_ids", "scripts/runpod-batch-worker.py")
class _FakeProc:
def __init__(self, cmd, **kwargs):
self.cmd = cmd
def poll(self):
return 0
def wait(self):
return 0
def terminate(self):
return None
def kill(self):
return None
monkeypatch.setattr(module.subprocess, "Popen", lambda cmd, **kwargs: _FakeProc(cmd, **kwargs))
monkeypatch.setattr(module.time, "sleep", lambda *_args, **_kwargs: None)
monkeypatch.setenv("RUNPOD_WORKER_GPU_IDS", "3,5")
manifest_path = tmp_path / "manifest.json"
manifest = {"manifest_path": str(manifest_path), "phases": []}
report_dir = tmp_path / "reports"
report_dir.mkdir()
module._run_processing_shards(
label="720p",
process_shards=2,
process_limit=0,
partition_total=1,
partition_index=0,
delete_after=False,
report_dir=str(report_dir),
manifest=manifest,
shard_timeout_seconds=0,
)
commands = [shard["command"] for shard in manifest["phases"][0]["shards"]]
gpu_ids = [cmd[cmd.index("--gpu-id") + 1] for cmd in commands]
assert gpu_ids == ["3", "5"]
def test_runpod_batch_worker_respects_single_visible_cuda_device(tmp_path, monkeypatch):
module = _load_module("test_runpod_batch_worker_single_visible", "scripts/runpod-batch-worker.py")
class _FakeProc:
def __init__(self, cmd, **kwargs):
self.cmd = cmd
def poll(self):
return 0
def wait(self):
return 0
def terminate(self):
return None
def kill(self):
return None
monkeypatch.setattr(module.subprocess, "Popen", lambda cmd, **kwargs: _FakeProc(cmd, **kwargs))
monkeypatch.setattr(module.time, "sleep", lambda *_args, **_kwargs: None)
monkeypatch.delenv("RUNPOD_WORKER_GPU_IDS", raising=False)
monkeypatch.setenv("CUDA_VISIBLE_DEVICES", "7")
manifest_path = tmp_path / "manifest.json"
manifest = {"manifest_path": str(manifest_path), "phases": []}
report_dir = tmp_path / "reports"
report_dir.mkdir()
module._run_processing_shards(
label="720p",
process_shards=1,
process_limit=0,
partition_total=1,
partition_index=0,
delete_after=False,
report_dir=str(report_dir),
manifest=manifest,
shard_timeout_seconds=0,
)
command = manifest["phases"][0]["shards"][0]["command"]
assert command[command.index("--gpu-id") + 1] == "7"
def test_runpod_batch_worker_passes_partition_args_to_shards(tmp_path, monkeypatch):
module = _load_module("test_runpod_batch_worker_partition", "scripts/runpod-batch-worker.py")
class _FakeProc:
def __init__(self, cmd, **kwargs):
self.cmd = cmd
def poll(self):
return 0
def wait(self):
return 0
def terminate(self):
return None
def kill(self):
return None
monkeypatch.setattr(module.subprocess, "Popen", lambda cmd, **kwargs: _FakeProc(cmd, **kwargs))
monkeypatch.setattr(module.time, "sleep", lambda *_args, **_kwargs: None)
manifest_path = tmp_path / "manifest.json"
manifest = {"manifest_path": str(manifest_path), "phases": []}
report_dir = tmp_path / "reports"
report_dir.mkdir()
module._run_processing_shards(
label="720p",
process_shards=2,
process_limit=0,
partition_total=3,
partition_index=1,
delete_after=False,
report_dir=str(report_dir),
manifest=manifest,
shard_timeout_seconds=0,
)
commands = [shard["command"] for shard in manifest["phases"][0]["shards"]]
assert all("--partition-total" in cmd for cmd in commands)
assert all("--partition-index" in cmd for cmd in commands)
def test_run_batch_script_wires_gpu_ids_into_process_workers():
script = (REPO_ROOT / "deploy/runpod/run-batch.sh").read_text()
multi_gpu_proof = (REPO_ROOT / "deploy/runpod/run-multi-gpu-proof.sh").read_text()
assert "RUNPOD_GPU_IDS" in script
assert "get_shard_gpu_id" in script
assert 'shard_cmd+=(--gpu-id "$gpu_id")' in script
assert "RUNPOD_PARTITION_TOTAL" in script
assert "--partition-total \"$RUNPOD_PARTITION_TOTAL\"" in script
assert "count_natural_key_file_unprocessed" in script
assert "DOWNLOAD_TARGET_FILE" in script
assert "DOWNLOAD_PLAN_FILE" in script
assert "DOWNLOAD_PLANNED_TARGET_COUNT" in script
assert "restart_download_if_needed" in script
assert '--natural-key-file "$DOWNLOAD_TARGET_FILE"' in script
assert "restart_dead_shards_if_needed" in script
assert "WATCHDOG_STATE_DIR" in script
assert "/tmp/search-ui-watchdog" in script
assert "SHARD_DISABLED" in script
assert "backup_runtime_databases_best_effort" in script
assert "RunPod Multi-GPU Proof" in multi_gpu_proof
assert "export RUNPOD_GPU_IDS" in multi_gpu_proof
assert "duplicate GPU ID" in multi_gpu_proof
assert "RUNPOD_PARTITION_TOTAL and RUNPOD_PARTITION_INDEX must be set explicitly" in multi_gpu_proof
assert 'bash "$SCRIPT_DIR/run-batch.sh"' in multi_gpu_proof
def test_run_batch_archives_shard_progress_before_restart_cap_failure():
script = (REPO_ROOT / "deploy/runpod/run-batch.sh").read_text()
archive_idx = script.index(' archive_shard_progress "$shard_index"')
cap_idx = script.index(' if [ "${SHARD_RESTART_COUNTS[$shard_index]:-0}" -ge "$WATCHDOG_MAX_SHARD_RESTARTS" ]; then')
assert archive_idx < cap_idx
def test_run_proof_validation_uses_vec_connections():
script = (REPO_ROOT / "deploy/runpod/run-proof.sh").read_text()
assert "from db_utils import IMAGE_DB_PRAGMAS, connect_vec_db" in script
assert "image_conn = connect_vec_db(get_image_db_path(), pragmas=IMAGE_DB_PRAGMAS)" in script
assert "connect_vec_db(face_db_path, pragmas=IMAGE_DB_PRAGMAS)" in script
def test_runpod_scripts_fail_fast_when_system_tools_are_missing():
run_proof = (REPO_ROOT / "deploy/runpod/run-proof.sh").read_text()
run_batch = (REPO_ROOT / "deploy/runpod/run-batch.sh").read_text()
export_results = (REPO_ROOT / "deploy/runpod/export-results.sh").read_text()
setup_pod = (REPO_ROOT / "deploy/runpod/setup-pod.sh").read_text()
watchdog = (REPO_ROOT / "deploy/runpod/watchdog.py").read_text()
check_gpu_runtime = (REPO_ROOT / "deploy/runpod/check_gpu_runtime.py").read_text()
backup_runtime_dbs = (REPO_ROOT / "deploy/runpod/backup_runtime_dbs.py").read_text()
restore_runtime_dbs = (REPO_ROOT / "deploy/runpod/restore_runtime_dbs.py").read_text()
assert "require_system_commands ffmpeg ffprobe sqlite3" in run_proof
assert "Recovery: bash deploy/runpod/setup-pod.sh" in run_proof
assert "require_system_commands ffmpeg ffprobe sqlite3 setsid" in run_batch
assert "check_gpu_runtime" in run_batch
assert 'check_gpu_runtime.py" --context "run-proof.sh"' in run_proof
assert 'check_gpu_runtime.py" --context "run-batch.sh"' in run_batch
assert 'check_gpu_runtime.py" --context "setup-pod.sh"' in setup_pod
assert 'backup_runtime_dbs.py' in run_proof
assert 'backup_runtime_dbs.py' in run_batch
assert "backup_runtime_databases_or_die" in run_batch
assert "count_natural_key_file_processed_in_db" in run_batch
assert "All downloadable targets are already processed in the DB." in run_batch
assert "Recovery: bash deploy/runpod/setup-pod.sh" in run_batch
assert "sqlite3 unavailable" in export_results
assert "Recovery: bash deploy/runpod/setup-pod.sh" in export_results
assert "db-backup (runtime backup failed)" in export_results
assert "--include-large-thumbnails" in export_results
assert "excluding large thumbnails (default)" in export_results
assert 'restore_runtime_dbs.py' in setup_pod
assert 'VENV_DIR=/tmp/search-ui-venv' in setup_pod
assert 'run_pip_with_retries' in setup_pod
assert "--auto-export" in run_batch
assert "Allowing active shards to finish the downloaded backlog" in run_batch
assert "downloader exited too many times" in run_batch
assert "CHUNK COMPLETE WITH DOWNLOAD WARNINGS" in run_batch
assert "all shards are disabled while downloader is still running" in run_batch
assert "watchdog.py" in run_batch
assert "write_watchdog_state" in run_batch
assert "cleanup_on_exit" in run_batch
assert "trap 'cleanup_on_exit $?'" in run_batch
assert 'setsid "${download_cmd[@]}"' in run_batch
assert "os.killpg(pid, signal.SIGTERM)" in watchdog
assert "--dead-shard-checks" in watchdog
assert "GPU capability" in check_gpu_runtime
assert "CUDA smoke test failed" in check_gpu_runtime
assert "source_conn.backup(target_conn)" in backup_runtime_dbs
assert "backup newer than runtime" in restore_runtime_dbs
assert "backup-manifest.json" in backup_runtime_dbs
assert "manifest_sha_mismatch" in restore_runtime_dbs
assert "export-status.json" in export_results
assert "_collect_export_status" in watchdog
assert "RUNPOD_ENABLE_DEEPFACE_GPU" in setup_pod
assert "tensorflow[and-cuda]" in setup_pod
assert "Linked ptxas" in setup_pod
assert "SEARCH_UI_TF_LIBRARY_DIRS" in setup_pod
assert "LD_LIBRARY_PATH" in setup_pod
assert "RUNPOD_PARTITION_TOTAL" in run_batch
def test_check_gpu_runtime_rejects_unsupported_arch():
module = _load_module("test_runpod_check_gpu_runtime", "deploy/runpod/check_gpu_runtime.py")
probe = {
"torch_version": "2.4.1+cu124",
"cuda_version": "12.4",
"cuda_available": True,
"device_count": 1,
"device_name": "NVIDIA RTX PRO 4500 Blackwell",
"total_memory_gb": 32.0,
"capability_tag": "sm_120",
"supported_arches": ["sm_50", "sm_60", "sm_70", "sm_75", "sm_80", "sm_86", "sm_90"],
"smoke_error": None,
}
ok, message = module.evaluate_runtime(gpu_visible=True, probe=probe)
assert ok is False
assert "sm_120" in message
assert "sm_90" in message
def test_watchdog_terminates_orphaned_children(tmp_path):
module = _load_module("test_runpod_watchdog", "deploy/runpod/watchdog.py")
download_proc = subprocess.Popen(["sleep", "300"])
shard_proc = subprocess.Popen(["sleep", "300"])
try:
state_path = tmp_path / "watchdog-state.json"
status_path = tmp_path / "watchdog-status.json"
state_path.write_text(
json.dumps(
{
"phase": "processing",
"heartbeat_epoch": int(time.time()),
"supervisor_pid": 999999,
"download_pid": download_proc.pid,
"shard_pids": [shard_proc.pid],
"report_dir": str(tmp_path),
"workspace_dir": str(tmp_path),
"download_finished": False,
"download_exit_code": 0,
}
)
)
exit_code = module.main(
[
"--state-file",
str(state_path),
"--status-file",
str(status_path),
"--once",
]
)
assert exit_code == 2
download_proc.wait(timeout=10)
shard_proc.wait(timeout=10)
status = json.loads(status_path.read_text())
assert status["action"] == "terminated"
assert status["reason"] == "supervisor_dead"
finally:
for proc in (download_proc, shard_proc):
if proc.poll() is None:
proc.terminate()
try:
proc.wait(timeout=5)
except subprocess.TimeoutExpired:
proc.kill()
def test_watchdog_requests_restart_on_partial_shard_death(tmp_path):
module = _load_module("test_runpod_watchdog_partial", "deploy/runpod/watchdog.py")
alive_shard = subprocess.Popen(["sleep", "300"])
dead_shard = subprocess.Popen(["sleep", "300"])
dead_shard.terminate()
dead_shard.wait(timeout=5)
terminated = {}
try:
state_path = tmp_path / "watchdog-state.json"
status_path = tmp_path / "watchdog-status.json"
state_path.write_text(
json.dumps(
{
"phase": "processing",
"heartbeat_epoch": int(time.time()),
"supervisor_pid": os.getpid(),
"download_pid": None,
"shard_pids": [alive_shard.pid, dead_shard.pid],
"report_dir": str(tmp_path),
"workspace_dir": str(tmp_path),
"download_finished": True,
"download_exit_code": 0,
}
)
)
module._terminate_pids = lambda pids, grace_seconds=5.0: terminated.setdefault("pids", list(pids))
exit_code = module.main(
[
"--state-file",
str(state_path),
"--status-file",
str(status_path),
"--dead-shard-checks",
"1",
"--once",
]
)
assert exit_code == 0
status = json.loads(status_path.read_text())
assert status["action"] == "restart_requested"
assert status["reason"] == "partial_shard_death"
assert dead_shard.pid in status["dead_shard_pids"]
assert "pids" not in terminated
finally:
if alive_shard.poll() is None:
alive_shard.terminate()
try:
alive_shard.wait(timeout=5)
except subprocess.TimeoutExpired:
alive_shard.kill()
def test_watchdog_uses_in_memory_state_when_status_writes_fail(tmp_path, monkeypatch):
module = _load_module("test_runpod_watchdog_memory_state", "deploy/runpod/watchdog.py")
alive_shard = subprocess.Popen(["sleep", "300"])
dead_shard = subprocess.Popen(["sleep", "300"])
dead_shard.terminate()
dead_shard.wait(timeout=5)
try:
state_path = tmp_path / "watchdog-state.json"
status_path = tmp_path / "watchdog-status.json"
state_path.write_text(
json.dumps(
{
"phase": "processing",
"heartbeat_epoch": int(time.time()),
"supervisor_pid": os.getpid(),
"download_pid": None,
"shard_pids": [alive_shard.pid, dead_shard.pid],
"report_dir": str(tmp_path),
"workspace_dir": str(tmp_path),
"download_finished": True,
"download_exit_code": 0,
}
)
)
monkeypatch.setattr(module, "_save_status", lambda path, payload: None)
exit_code, previous = module._evaluate(
module.parse_args(
[
"--state-file",
str(state_path),
"--status-file",
str(status_path),
"--dead-shard-checks",
"2",
"--once",
]
),
str(status_path),
{},
)
assert exit_code == 0
assert previous["dead_shard_streak"] == 1
exit_code, previous = module._evaluate(
module.parse_args(
[
"--state-file",
str(state_path),
"--status-file",
str(status_path),
"--dead-shard-checks",
"2",
"--once",
]
),
str(status_path),
previous,
)
assert exit_code == 0
assert previous["action"] == "restart_requested"
assert previous["dead_shard_streak"] == 2
finally:
if alive_shard.poll() is None:
alive_shard.terminate()
try:
alive_shard.wait(timeout=5)
except subprocess.TimeoutExpired:
alive_shard.kill()
def test_watchdog_terminates_stale_pipeline_when_state_file_disappears(tmp_path):
module = _load_module("test_runpod_watchdog_missing_state", "deploy/runpod/watchdog.py")
status_path = tmp_path / "watchdog-status.json"
previous = {
"action": "healthy",
"phase": "processing",
"supervisor_pid": 11111,
"download_pid": None,
"shard_pids": [22222],
}
terminated = {}
module._pid_exists = lambda pid: pid in {11111, 22222}
module._terminate_pids = lambda pids, grace_seconds=5.0: terminated.setdefault("pids", list(pids))
exit_code, status = module._evaluate(
module.parse_args(
[
"--state-file",
str(tmp_path / "missing-state.json"),
"--status-file",
str(status_path),
"--once",
]
),
str(status_path),
previous,
)
assert exit_code == 6
assert terminated["pids"] == [11111, 22222]
assert status["action"] == "terminated"
assert status["reason"] == "state_file_missing"
def test_watchdog_state_file_missing_scans_for_live_pipeline_candidates(tmp_path):
module = _load_module("test_runpod_watchdog_missing_state_scan", "deploy/runpod/watchdog.py")
status_path = tmp_path / "watchdog-status.json"
terminated = {}
module._collect_pipeline_candidate_pids = lambda: [33333, 44444]
module._pid_exists = lambda pid: pid in {33333, 44444}
module._terminate_pids = lambda pids, grace_seconds=5.0: terminated.setdefault("pids", list(pids))
exit_code, status = module._evaluate(
module.parse_args(
[
"--state-file",
str(tmp_path / "missing-state.json"),
"--status-file",
str(status_path),
"--once",
]
),
str(status_path),
{},
)
assert exit_code == 6
assert terminated["pids"] == [33333, 44444]
assert status["action"] == "terminated"
assert status["reason"] == "state_file_missing"
def test_watchdog_collects_pipeline_candidate_pids_from_pgrep(monkeypatch):
module = _load_module("test_runpod_watchdog_candidate_scan", "deploy/runpod/watchdog.py")
def _fake_run(cmd, capture_output, text, check):
pattern = cmd[-1]
if "run-batch.sh" in pattern:
return type("Result", (), {"returncode": 0, "stdout": "111\n222\n"})()
if "process-local-videos.py" in pattern:
return type("Result", (), {"returncode": 0, "stdout": "222\n333\n"})()
return type("Result", (), {"returncode": 1, "stdout": ""})()
monkeypatch.setattr(module.subprocess, "run", _fake_run)
monkeypatch.setattr(module.os, "getpid", lambda: 333)
assert module._collect_pipeline_candidate_pids() == [111, 222]
def test_watchdog_pid_exists_treats_zombies_as_dead(monkeypatch):
module = _load_module("test_runpod_watchdog_pid_exists", "deploy/runpod/watchdog.py")
monkeypatch.setattr(module.Path, "read_text", lambda self, encoding="utf-8": "123 (python) Z 1 2 3")
monkeypatch.setattr(module.os, "kill", lambda pid, sig: None)
assert module._pid_exists(12345) is False
def test_watchdog_terminates_when_all_children_are_dead(tmp_path):
module = _load_module("test_runpod_watchdog_all_children_dead", "deploy/runpod/watchdog.py")
state_path = tmp_path / "watchdog-state.json"
status_path = tmp_path / "watchdog-status.json"
state_path.write_text(
json.dumps(
{
"phase": "processing",
"heartbeat_epoch": int(time.time()),
"supervisor_pid": os.getpid(),
"download_pid": 999991,
"shard_pids": [999992, 999993],
"report_dir": str(tmp_path),
"workspace_dir": str(tmp_path),
"download_finished": False,
"download_exit_code": 0,
}
)
)
terminated = {}
module._terminate_pids = lambda pids, grace_seconds=5.0: terminated.setdefault("pids", list(pids))
exit_code = module.main(
[
"--state-file",
str(state_path),
"--status-file",
str(status_path),
"--once",
]
)
assert exit_code == 7
assert terminated["pids"] == [os.getpid()]
status = json.loads(status_path.read_text())
assert status["action"] == "terminated"
assert status["reason"] == "all_children_dead"
def test_watchdog_detects_export_stall_before_status_file_exists(tmp_path):
module = _load_module("test_runpod_watchdog_export_stall", "deploy/runpod/watchdog.py")
original_export_status = module.EXPORT_STATUS_FILE
module.EXPORT_STATUS_FILE = str(tmp_path / "missing-export-status.json")
try:
state_path = tmp_path / "watchdog-state.json"
status_path = tmp_path / "watchdog-status.json"
state_path.write_text(
json.dumps(
{
"phase": "auto-export",
"heartbeat_epoch": int(time.time()),
"supervisor_pid": os.getpid(),
"download_pid": None,
"shard_pids": [],
"report_dir": str(tmp_path),
"workspace_dir": str(tmp_path),
"download_finished": True,
"download_exit_code": 0,
}
)
)
args = module.parse_args(
[
"--state-file",
str(state_path),
"--status-file",
str(status_path),
"--stall-checks",
"2",
"--once",
]
)
terminated = {}
module._terminate_pids = lambda pids, grace_seconds=5.0: terminated.setdefault("pids", list(pids))
exit_code, previous = module._evaluate(args, str(status_path), {})
assert exit_code == 0
assert previous["stall_streak"] == 1
exit_code, previous = module._evaluate(args, str(status_path), previous)
assert exit_code == 5
assert terminated["pids"] == [os.getpid()]
assert previous["action"] == "terminated"
assert previous["reason"] == "export_stall"
finally:
module.EXPORT_STATUS_FILE = original_export_status
def test_watchdog_detects_export_stall_when_status_and_signature_freeze(tmp_path):
module = _load_module("test_runpod_watchdog_export_frozen", "deploy/runpod/watchdog.py")
state_path = tmp_path / "watchdog-state.json"
status_path = tmp_path / "watchdog-status.json"
state_path.write_text(
json.dumps(
{
"phase": "auto-export",
"heartbeat_epoch": int(time.time()),
"supervisor_pid": os.getpid(),
"download_pid": None,
"shard_pids": [],
"report_dir": str(tmp_path),
"workspace_dir": str(tmp_path),
"download_finished": True,
"download_exit_code": 0,
}
)
)
module._collect_export_status = lambda: {"stage": "tar", "_mtime_ns": 1, "_size": 10}
module._collect_export_signature = lambda: ["tarball.tgz", 123, 456]
terminated = {}
module._terminate_pids = lambda pids, grace_seconds=5.0: terminated.setdefault("pids", list(pids))
args = module.parse_args(
[
"--state-file",
str(state_path),
"--status-file",
str(status_path),
"--stall-checks",
"2",
"--once",
]
)
exit_code, previous = module._evaluate(args, str(status_path), {})
assert exit_code == 0
assert previous["stall_streak"] == 0
exit_code, previous = module._evaluate(args, str(status_path), previous)
assert exit_code == 0
assert previous["stall_streak"] == 1
exit_code, previous = module._evaluate(args, str(status_path), previous)
assert exit_code == 5
assert terminated["pids"] == [os.getpid()]
assert previous["action"] == "terminated"
assert previous["reason"] == "export_stall"
def test_restore_runtime_dbs_prefers_newer_valid_backup(tmp_path):
module = _load_module("test_restore_runtime_dbs", "deploy/runpod/restore_runtime_dbs.py")
db_utils = _load_module("test_restore_runtime_dbs_utils", "backend/db_utils.py")
runtime_dir = tmp_path / "runtime"
backup_dir = tmp_path / "backup"
runtime_dir.mkdir()
backup_dir.mkdir()
runtime_db = runtime_dir / "images_database.db"
backup_db = backup_dir / "images_database.db"
runtime_conn = db_utils.connect_vec_db(str(runtime_db))
runtime_conn.execute("CREATE TABLE sample(value TEXT)")
runtime_conn.execute("INSERT INTO sample(value) VALUES ('runtime')")
runtime_conn.commit()
runtime_conn.close()
time.sleep(0.01)
backup_conn = db_utils.connect_vec_db(str(backup_db))
backup_conn.execute("CREATE TABLE sample(value TEXT)")
backup_conn.execute("INSERT INTO sample(value) VALUES ('backup')")
backup_conn.commit()
backup_conn.close()
exit_code = module.main(
[
"--runtime-dir",
str(runtime_dir),
"--backup-dir",
str(backup_dir),
]
)
assert exit_code == 0
check_conn = db_utils.connect_vec_db(str(runtime_db))
value = check_conn.execute("SELECT value FROM sample").fetchone()[0]
check_conn.close()
assert value == "backup"
def test_restore_runtime_dbs_keeps_newer_valid_runtime(tmp_path):
module = _load_module("test_restore_runtime_dbs_keep", "deploy/runpod/restore_runtime_dbs.py")
db_utils = _load_module("test_restore_runtime_dbs_keep_utils", "backend/db_utils.py")
runtime_dir = tmp_path / "runtime"
backup_dir = tmp_path / "backup"
runtime_dir.mkdir()
backup_dir.mkdir()
runtime_db = runtime_dir / "images_database.db"
backup_db = backup_dir / "images_database.db"
backup_conn = db_utils.connect_vec_db(str(backup_db))
backup_conn.execute("CREATE TABLE sample(value TEXT)")
backup_conn.execute("INSERT INTO sample(value) VALUES ('backup')")
backup_conn.commit()
backup_conn.close()
time.sleep(0.01)
runtime_conn = db_utils.connect_vec_db(str(runtime_db))
runtime_conn.execute("CREATE TABLE sample(value TEXT)")
runtime_conn.execute("INSERT INTO sample(value) VALUES ('runtime')")
runtime_conn.commit()
runtime_conn.close()
exit_code = module.main(
[
"--runtime-dir",
str(runtime_dir),
"--backup-dir",
str(backup_dir),
]
)
assert exit_code == 0
check_conn = db_utils.connect_vec_db(str(runtime_db))
value = check_conn.execute("SELECT value FROM sample").fetchone()[0]
check_conn.close()
assert value == "runtime"
def test_backup_runtime_dbs_writes_manifest(tmp_path):
module = _load_module("test_backup_runtime_dbs_manifest", "deploy/runpod/backup_runtime_dbs.py")
db_utils = _load_module("test_backup_runtime_dbs_manifest_utils", "backend/db_utils.py")
source_dir = tmp_path / "runtime"
target_dir = tmp_path / "backup"
source_dir.mkdir()
target_dir.mkdir()
source_db = source_dir / "images_database.db"
conn = db_utils.connect_vec_db(str(source_db))
conn.execute("CREATE TABLE sample(value TEXT)")
conn.execute("INSERT INTO sample(value) VALUES ('runtime')")
conn.commit()
conn.close()
exit_code = module.main(
[
"--source-dir",
str(source_dir),
"--target-dir",
str(target_dir),
]
)
assert exit_code == 0
manifest = json.loads((target_dir / "backup-manifest.json").read_text())
entry = manifest["databases"]["images_database.db"]
assert entry["quick_check"] == "ok"
assert entry["size_bytes"] > 0
assert entry["sha256"]
def test_restore_runtime_dbs_rejects_manifest_checksum_mismatch(tmp_path):
backup_module = _load_module("test_backup_runtime_dbs_manifest_restore", "deploy/runpod/backup_runtime_dbs.py")
restore_module = _load_module("test_restore_runtime_dbs_manifest_restore", "deploy/runpod/restore_runtime_dbs.py")
db_utils = _load_module("test_restore_runtime_dbs_manifest_restore_utils", "backend/db_utils.py")
runtime_dir = tmp_path / "runtime"
backup_dir = tmp_path / "backup"
runtime_dir.mkdir()
backup_dir.mkdir()
source_db = runtime_dir / "images_database.db"
conn = db_utils.connect_vec_db(str(source_db))
conn.execute("CREATE TABLE sample(value TEXT)")
conn.execute("INSERT INTO sample(value) VALUES ('runtime')")
conn.commit()
conn.close()
backup_module.main(
[
"--source-dir",
str(runtime_dir),
"--target-dir",
str(backup_dir),
]
)
backup_db = backup_dir / "images_database.db"
backup_db.write_bytes(backup_db.read_bytes() + b"corrupt")
source_db.unlink()
exit_code = restore_module.main(
[
"--runtime-dir",
str(runtime_dir),
"--backup-dir",
str(backup_dir),
]
)
assert exit_code == 1
assert not source_db.exists()
# ---------------------------------------------------------------------------
# _chat_with_retry — backoff on transient API errors (429/5xx/timeout)
# ---------------------------------------------------------------------------
class _FakeStatusError(Exception):
"""Mimics an openai SDK error carrying a .status_code attribute."""
def __init__(self, status_code, message=""):
super().__init__(message or f"status {status_code}")
self.status_code = status_code
class _FakeClient:
"""A client whose chat.completions.create raises a queued sequence of
exceptions, then returns a sentinel once the queue is exhausted."""
def __init__(self, errors):
self._errors = list(errors)
self.calls = 0
outer = self
class _Completions:
def create(self, **kwargs):
outer.calls += 1
if outer._errors:
raise outer._errors.pop(0)
return "OK_RESPONSE"
class _Chat:
completions = _Completions()
self.chat = _Chat()
def _load_runpod_processor(monkeypatch):
module = _load_module("test_runpod_scene_processor", "scripts/runpod_scene_processor.py")
# Make backoff instant so the test doesn't actually sleep.
monkeypatch.setattr(module.time, "sleep", lambda *_a, **_k: None)
return module
def test_chat_with_retry_recovers_after_429(monkeypatch):
module = _load_runpod_processor(monkeypatch)
client = _FakeClient([_FakeStatusError(429), _FakeStatusError(429)])
result = module._chat_with_retry(client, model="m", messages=[])
assert result == "OK_RESPONSE"
assert client.calls == 3 # two failures + one success
def test_chat_with_retry_reraises_fatal_400_immediately(monkeypatch):
module = _load_runpod_processor(monkeypatch)
client = _FakeClient([_FakeStatusError(400, "bad request")])
try:
module._chat_with_retry(client, model="m", messages=[])
raise AssertionError("expected the 400 to be re-raised")
except _FakeStatusError as exc:
assert exc.status_code == 400
assert client.calls == 1 # no retries on a non-transient error
def test_chat_with_retry_gives_up_after_max_retries(monkeypatch):
module = _load_runpod_processor(monkeypatch)
# Always 429 — should exhaust max_retries (8) and then raise.
client = _FakeClient([_FakeStatusError(429) for _ in range(20)])
try:
module._chat_with_retry(client, model="m", messages=[], max_retries=8)
raise AssertionError("expected to give up and raise after max retries")
except _FakeStatusError as exc:
assert exc.status_code == 429
assert client.calls == 8 # exactly max_retries attempts, no infinite loop
def test_chat_with_retry_retries_on_timeout_message(monkeypatch):
module = _load_runpod_processor(monkeypatch)
# No status_code, but message indicates a timeout → should retry.
client = _FakeClient([TimeoutError("request timeout")])
result = module._chat_with_retry(client, model="m", messages=[])
assert result == "OK_RESPONSE"
assert client.calls == 2