Files
ARR-2.0-0918/tests/test_ohip_capture_jobs.py

321 lines
15 KiB
Python

from __future__ import annotations
import copy
import hashlib
import json
import multiprocessing
import os
from pathlib import Path
import stat
import subprocess
import sys
import tempfile
import unittest
from unittest.mock import patch
from integrations.ohip import capture_job as jobs
from integrations.ohip import collect_arr_source as collector
from tests.test_ohip_arr_collection import DAY, HOTEL, FakeService
BATCH = "arr-day-001"
SCRIPT = Path(__file__).resolve().parents[1] / "integrations/ohip/capture_job.py"
def assert_strict_persisted_identity(test, run_job, directory, fields):
"""A completed job must match JSON types as well as Python numeric values."""
test.assertTrue(run_job()["candidate_capture_complete"])
path = directory / "job.json"
original = json.loads(path.read_bytes())
# Object key order is immaterial; numeric type changes are not.
path.write_bytes(collector.json_bytes(dict(reversed(list(original.items())))))
test.assertTrue(run_job(factory=lambda *_: test.fail("unexpected reader"))["reused"])
for field in fields:
with test.subTest(field=field):
changed = copy.deepcopy(original)
changed["options"][field] = float(changed["options"][field])
path.write_bytes(collector.json_bytes(changed))
before = {str(p): p.read_bytes() for p in directory.rglob("*") if p.is_file()}
result = run_job(factory=lambda *_: test.fail("unexpected reader"))
test.assertEqual(result.get("error"), "batch_request_conflict")
test.assertEqual(result["status"], "capture_job_refused")
test.assertEqual(before, {str(p): p.read_bytes() for p in directory.rglob("*") if p.is_file()})
def child_hold_batch(root, entered, release):
def factory(archive, hotel):
entered.set()
if not release.wait(10):
raise RuntimeError("test release timed out")
return collector.Reader(archive, hotel, FakeService(count=1), sleep=lambda _: None)
result = jobs.run_batch(Path(root), BATCH, collector.Options(DAY, HOTEL, page_size=2), factory)
if not result.get("candidate_capture_complete"):
raise RuntimeError("child capture failed")
def child_crash_batch(root, phase):
original_atomic = jobs.atomic_json
def atomic(path, doc, *, replace):
if phase == "after_capture" and path.name == "state.json" and doc["attempts"] and doc["attempts"][-1]["status"] == "complete":
os._exit(24)
return original_atomic(path, doc, replace=replace)
jobs.atomic_json = atomic
def factory(archive, hotel):
if phase == "before_capture":
os._exit(23)
return collector.Reader(archive, hotel, FakeService(count=1), sleep=lambda _: None)
jobs.run_batch(Path(root), BATCH, collector.Options(DAY, HOTEL, page_size=2), factory)
class CaptureJobTests(unittest.TestCase):
def setUp(self):
self.temp = tempfile.TemporaryDirectory()
self.addCleanup(self.temp.cleanup)
self.root = Path(self.temp.name) / "jobs"
self.options = collector.Options(DAY, HOTEL, page_size=2)
self.service = FakeService(count=3)
self.factory_calls = 0
def factory(self, archive, hotel):
self.factory_calls += 1
return collector.Reader(archive, hotel, self.service, sleep=lambda _: None)
def run_job(self, **kwargs):
return jobs.run_batch(kwargs.get("root", self.root), kwargs.get("batch_id", BATCH),
kwargs.get("options", self.options), kwargs.get("factory", self.factory))
def state(self):
return json.loads((self.root / BATCH / "state.json").read_bytes())
def write_state(self, value):
(self.root / BATCH / "state.json").write_bytes(collector.json_bytes(value))
def digests(self, path):
return {str(p.relative_to(path)): hashlib.sha256(p.read_bytes()).hexdigest()
for p in path.rglob("*") if p.is_file()}
def test_success_reused_without_credentials_or_new_reads(self):
first = self.run_job()
self.assertTrue(first["candidate_capture_complete"])
self.assertFalse(first["reused"])
self.assertFalse(first["finance_ready"])
before = self.digests(self.root)
second = self.run_job(factory=lambda *args: self.fail("factory must not be called for completed batch"))
self.assertTrue(second["reused"])
self.assertEqual(second["manifest_sha256"], first["manifest_sha256"])
self.assertEqual(second["attempt_no"], 1)
self.assertEqual(self.factory_calls, 1)
self.assertEqual(len(self.service.calls), 7)
self.assertEqual(self.digests(self.root), before)
def test_persisted_job_identity_rejects_numeric_type_changes(self):
assert_strict_persisted_identity(self, self.run_job, self.root / BATCH, ("page_size", "max_pages"))
self.assertEqual(self.factory_calls, 1)
def test_same_day_new_batch_is_explicit_fresh_capture(self):
first, second = self.run_job(), self.run_job(batch_id="arr-day-002")
self.assertTrue(first["candidate_capture_complete"] and second["candidate_capture_complete"])
self.assertFalse(second["reused"])
self.assertNotEqual(first["capture_dir"], second["capture_dir"])
self.assertEqual(self.factory_calls, 2)
def test_batch_id_cannot_change_day_hotel_or_limits(self):
self.run_job()
for options in (collector.Options("2026-09-14", HOTEL, page_size=2),
collector.Options(DAY, "different-hotel", page_size=2), collector.Options(DAY, HOTEL)):
with self.subTest(options=options):
result = self.run_job(options=options)
self.assertEqual(result["error"], "batch_request_conflict")
self.assertEqual(self.factory_calls, 1)
def test_partial_failure_retries_a_whole_new_capture(self):
underlying = self.service
count = 0
def failing(method, path, raw):
nonlocal count
count += 1
if count == 3:
return 403, {}, b"PRIVATE_DENIAL"
return underlying(method, path, raw)
first = self.run_job(factory=lambda archive, hotel: collector.Reader(archive, hotel, failing, sleep=lambda _: None))
self.assertEqual(first["status"], "capture_failed")
old = self.root / BATCH / "attempt-0001"
before = self.digests(old)
self.service = FakeService(count=3)
second = self.run_job()
self.assertEqual(second["attempt_no"], 2)
self.assertTrue(second["candidate_capture_complete"])
self.assertEqual(json.loads(self.service.calls[0][2])["offset"], 0)
self.assertEqual(len(self.service.calls), 7)
self.assertEqual(self.digests(old), before)
self.assertEqual([item["status"] for item in self.state()["attempts"]], ["failed", "complete"])
self.assertNotIn("PRIVATE_DENIAL", json.dumps(first))
def test_empty_source_has_no_complete_pointer(self):
self.service = FakeService(count=0)
result = self.run_job()
self.assertEqual(result["error"], "empty_source_requires_review")
self.assertEqual(self.state()["attempts"][0]["status"], "failed")
def test_tampered_completed_source_refuses_without_recollection(self):
first = self.run_job()
path = Path(first["capture_dir"]) / "request-000001.response.bin"
path.write_bytes(path.read_bytes().replace(b"PRIVATE_NOTE", b"CHANGED_NOTE"))
result = self.run_job()
self.assertEqual(result["error"], "archive_hash_mismatch")
self.assertEqual(self.factory_calls, 1)
self.assertEqual(len(self.state()["attempts"]), 1)
def test_changed_manifest_refuses_even_when_files_unchanged(self):
first = self.run_job()
path = Path(first["capture_dir"]) / "result.json"
path.write_bytes(path.read_bytes() + b"\n")
result = self.run_job()
self.assertEqual(result["error"], "manifest_hash_mismatch")
self.assertEqual(self.factory_calls, 1)
def test_missing_pin_is_not_reconstructed_from_completed_capture(self):
self.run_job()
state = self.state()
del state["attempts"][0]["manifest_sha256"]
self.write_state(state)
self.assertEqual(self.run_job()["error"], "missing_completed_capture_pin")
self.assertEqual(self.factory_calls, 1)
def test_private_modes_for_all_job_and_capture_files(self):
self.run_job()
for path in [self.root, *self.root.rglob("*")]:
self.assertEqual(stat.S_IMODE(path.stat().st_mode), 0o700 if path.is_dir() else 0o600)
self.assertFalse(list(self.root.rglob(".metadata-*")))
def test_unsafe_directory_and_symlink_refused(self):
self.root.mkdir(mode=0o755)
self.root.chmod(0o755)
self.assertEqual(self.run_job()["error"], "unsafe_job_directory")
self.root.chmod(0o700)
linked = self.root.parent / "linked"
linked.symlink_to(self.root, target_is_directory=True)
self.assertEqual(self.run_job(root=linked)["error"], "unsafe_job_directory")
self.assertEqual(self.factory_calls, 0)
def test_unsafe_lock_symlink_not_followed(self):
self.root.mkdir(mode=0o700)
(self.root / BATCH).mkdir(mode=0o700)
outside = self.root.parent / "outside"
outside.write_text("unchanged")
(self.root / BATCH / ".lock").symlink_to(outside)
result = self.run_job()
self.assertEqual(result["status"], "capture_job_refused")
self.assertEqual(outside.read_text(), "unchanged")
self.assertEqual(self.factory_calls, 0)
def test_invalid_batch_path_and_date_do_not_create_job(self):
for value in ("../escape", "/absolute", "UPPER", "", "a" * 65):
self.assertEqual(self.run_job(batch_id=value)["error"], "invalid_batch_id")
self.assertEqual(self.run_job(options=collector.Options("2026-99-99", HOTEL))["error"], "invalid_arrival_date")
self.assertFalse(self.root.exists())
def test_repository_store_refused(self):
repo_root = Path(__file__).resolve().parents[1]
result = self.run_job(root=repo_root / "never-created-capture-store")
self.assertEqual(result["error"], "job_store_must_be_outside_repository")
self.assertFalse((repo_root / "never-created-capture-store").exists())
def test_orphan_attempt_does_not_get_overwritten(self):
self.run_job()
(self.root / BATCH / "attempt-0002").mkdir(mode=0o700)
result = self.run_job()
self.assertEqual(result["error"], "orphan_capture_attempt")
self.assertEqual(self.factory_calls, 1)
def test_missing_state_with_existing_attempt_is_not_recreated(self):
self.run_job()
(self.root / BATCH / "state.json").unlink()
self.assertEqual(self.run_job()["error"], "missing_job_state")
self.assertEqual(self.factory_calls, 1)
def test_state_identity_and_attempt_sequence_are_checked(self):
self.run_job()
original = self.state()
changed = copy.deepcopy(original)
changed["request_sha256"] = "0" * 64
self.write_state(changed)
self.assertEqual(self.run_job()["error"], "job_state_identity_mismatch")
changed = copy.deepcopy(original)
changed["attempts"][0]["attempt_no"] = True
self.write_state(changed)
self.assertEqual(self.run_job()["error"], "invalid_job_attempt_number")
def test_setup_failure_safe_and_attempts_bounded(self):
def failing(*args):
raise RuntimeError("PRIVATE_KEY_SHOULD_NOT_APPEAR")
for number in range(1, jobs.MAX_ATTEMPTS + 1):
result = self.run_job(factory=failing)
self.assertEqual(result["attempt_no"], number)
self.assertEqual(result["error"], "capture_setup_or_storage_failure")
self.assertNotIn("PRIVATE_KEY", json.dumps(result))
result = self.run_job(factory=failing)
self.assertEqual(result["error"], "batch_attempt_limit_reached")
self.assertEqual(len(self.state()["attempts"]), jobs.MAX_ATTEMPTS)
def test_cli_reuses_complete_capture_without_opening_credential(self):
self.options = collector.Options(DAY, HOTEL)
first = self.run_job()
result = subprocess.run([sys.executable, str(SCRIPT), "--job-store", str(self.root), "--batch-id", BATCH,
"--arrival-date", DAY, "--hotel-id", HOTEL, "--credential-file", str(self.root.parent / "does-not-exist")],
capture_output=True, text=True)
self.assertEqual(result.returncode, 0, result.stderr)
output = json.loads(result.stdout)
self.assertTrue(output["reused"])
self.assertEqual(output["manifest_sha256"], first["manifest_sha256"])
self.assertNotIn("PRIVATE_", result.stdout + result.stderr)
def test_real_second_process_cannot_collect_same_batch_concurrently(self):
context = multiprocessing.get_context("spawn")
entered, release = context.Event(), context.Event()
process = context.Process(target=child_hold_batch, args=(str(self.root), entered, release))
process.start()
try:
self.assertTrue(entered.wait(5))
result = self.run_job()
self.assertEqual(result["error"], "batch_busy")
self.assertEqual(self.factory_calls, 0)
finally:
release.set()
process.join(10)
if process.is_alive():
process.terminate()
process.join(2)
self.assertEqual(process.exitcode, 0)
result = self.run_job()
self.assertTrue(result["reused"])
self.assertEqual(self.factory_calls, 0)
def test_process_death_releases_lock_and_restarts_complete_attempt(self):
for phase, expected_code in (("before_capture", 23), ("after_capture", 24)):
with self.subTest(phase=phase):
self.root = Path(self.temp.name) / phase
process = multiprocessing.get_context("spawn").Process(target=child_crash_batch, args=(str(self.root), phase))
process.start()
process.join(10)
if process.is_alive():
process.terminate()
process.join(2)
self.assertEqual(process.exitcode, expected_code)
self.assertEqual(self.state()["attempts"][0]["status"], "running")
old = self.root / BATCH / "attempt-0001"
before = self.digests(old)
self.assertEqual((old / "result.json").exists(), phase == "after_capture")
self.service = FakeService(count=3)
result = self.run_job()
self.assertTrue(result["candidate_capture_complete"])
self.assertFalse(result["reused"])
self.assertEqual(result["attempt_no"], 2)
self.assertEqual(self.digests(old), before)
self.assertEqual([item["status"] for item in self.state()["attempts"]], ["interrupted", "complete"])
self.assertEqual(json.loads(self.service.calls[0][2])["offset"], 0)
if __name__ == "__main__":
unittest.main()