273 lines
14 KiB
Python
273 lines
14 KiB
Python
"""Source composition/lifetime checks; isolated synthetic sources, no live APIs."""
|
|
|
|
from dataclasses import replace
|
|
import json
|
|
from pathlib import Path
|
|
import signal
|
|
import sqlite3
|
|
import tempfile
|
|
import threading
|
|
import time
|
|
from types import SimpleNamespace
|
|
import unittest
|
|
from unittest.mock import Mock, patch
|
|
|
|
from arr_ingestion.repository import InMemoryIngestionRepository
|
|
from arr_processing.policy import load_processor_policy
|
|
from arr_storage.filesystem import FilesystemObjectBackend
|
|
from arr_storage.store import ManagedObjectStore
|
|
from arr_web import run, processing_runtime
|
|
from arr_web.arr_download_runtime import CapturedARRSource, compose_arr_downloads
|
|
from arr_web.arr_downloads import DownloadOutcome, PersistentARRDownloads
|
|
from arr_web.local_replay import ReplayRuntime
|
|
from integrations.ohip import collect_arr_source as source
|
|
from integrations.ohip.rate_info import RateInfoReader
|
|
from tests.test_arr_web import FakeRepository, TEST_CREDENTIALS, login
|
|
from tests.test_arr_web_capture_executor import FixtureAdapter, FixtureValidator
|
|
from tests.test_arr_web_download_handoff import REQUEST_ID
|
|
from tests.test_ohip_day_capture import DayService, DAY, HOTEL
|
|
from tests.test_ohip_processing_handoff import CountingProcessor
|
|
|
|
|
|
class RuntimeTests(unittest.TestCase):
|
|
@classmethod
|
|
def setUpClass(cls):
|
|
cls.policy = load_processor_policy(Path(__file__).resolve().parents[1])
|
|
|
|
def setUp(self):
|
|
temporary = tempfile.TemporaryDirectory()
|
|
self.addCleanup(temporary.cleanup)
|
|
self.root = Path(temporary.name)
|
|
self.store = ManagedObjectStore(FilesystemObjectBackend(self.root / 'objects', create=True))
|
|
self.repository = InMemoryIngestionRepository()
|
|
self.processor = CountingProcessor(self.policy)
|
|
self.client = Mock()
|
|
with patch.object(processing_runtime, 'compose_object_store', return_value=
|
|
processing_runtime.ObjectStoreRuntime(self.client, self.store)), \
|
|
patch.object(processing_runtime, 'PostgresIngestionRepository', return_value=self.repository), \
|
|
patch.object(processing_runtime, 'load_processor_policy', return_value=self.policy), \
|
|
patch.object(processing_runtime, 'LocalDailyProcessor', return_value=self.processor):
|
|
self.processing = processing_runtime.compose_programmatic_processing(
|
|
project_root=Path(__file__).resolve().parents[1], connect=Mock())
|
|
self.transport = DayService()
|
|
self.factory_calls = 0
|
|
self.source = CapturedARRSource(root=self.root / 'downloads', hotel_id=HOTEL,
|
|
adapter_contract='synthetic-runtime-only/v1', adapter=FixtureAdapter(),
|
|
mapping_validator=FixtureValidator(), reader_factory=self.readers,
|
|
page_size=2, max_pages=10, max_records=10)
|
|
|
|
def readers(self, archive, hotel):
|
|
self.factory_calls += 1
|
|
return (source.Reader(archive, hotel, self.transport, sleep=lambda _: None),
|
|
RateInfoReader(archive, hotel, self.transport, sleep=lambda _: None))
|
|
|
|
def queue(self, configured=None):
|
|
queue = compose_arr_downloads(configured or self.source, self.processing)
|
|
self.addCleanup(queue.close, wait=True)
|
|
return queue
|
|
|
|
def test_shared_processing_and_pinned_day_flow_with_restart_reuse(self):
|
|
queue = self.queue()
|
|
self.assertEqual(self.factory_calls, 0)
|
|
self.assertEqual(self.transport.calls, [])
|
|
queue.create(DAY, REQUEST_ID)
|
|
deadline = time.monotonic() + 15
|
|
while queue.get(REQUEST_ID)['status'] in {'queued', 'downloading', 'processing'}:
|
|
self.assertLess(time.monotonic(), deadline)
|
|
time.sleep(.01)
|
|
result = queue.get(REQUEST_ID)
|
|
self.assertEqual(result['status'], 'succeeded')
|
|
self.assertEqual(len(self.repository._versions), 1)
|
|
self.assertEqual(self.processor.calls, 1)
|
|
pin = json.loads((self.source.root / 'source.json').read_bytes())
|
|
self.assertEqual(pin['rule_set_sha256'], self.policy.rule_set_sha256)
|
|
request = json.loads((self.source.root / 'acquisition/requests' / REQUEST_ID / 'request.json').read_bytes())
|
|
options = request['capture']['options']
|
|
self.assertEqual([options[k] for k in ['from_date', 'to_date', 'rate_date']], [DAY] * 3)
|
|
queue.close(wait=True)
|
|
restarted = self.queue()
|
|
self.assertEqual(restarted.create(DAY, REQUEST_ID), result)
|
|
self.assertEqual(self.factory_calls, 1)
|
|
self.client.close.assert_not_called() # queue borrows processing dependencies
|
|
|
|
def test_invalid_source_does_not_create_state_or_read_credentials(self):
|
|
cases = [dict(adapter=object()), dict(mapping_validator=object()), dict(page_size=0),
|
|
dict(capture_version='v3'), dict(max_profiles=3)]
|
|
for case in cases:
|
|
with self.subTest(case=case), self.assertRaises((ValueError, source.CollectionError)):
|
|
self.queue(replace(self.source, **case))
|
|
self.assertFalse(self.source.root.exists())
|
|
self.assertEqual(self.factory_calls, 0)
|
|
|
|
def test_runtime_pin_refuses_rebinding_and_unbound_existing_state(self):
|
|
queue = self.queue()
|
|
queue.close(wait=True)
|
|
with sqlite3.connect(self.source.root / 'queue/arr-downloads.sqlite3') as db:
|
|
db.execute("INSERT INTO downloads(request_id, report_date, status, created_at, updated_at) VALUES(?,?,?,?,?)",
|
|
(REQUEST_ID, DAY, 'queued', '2026-09-17T00:00:00+00:00', '2026-09-17T00:00:00+00:00'))
|
|
for change in [dict(hotel_id='OTHER'), dict(adapter_contract='synthetic-runtime-only/v2'),
|
|
dict(max_records=11), dict(capture_version='v3', max_profiles=3)]:
|
|
with self.subTest(change=change), self.assertRaisesRegex(source.CollectionError, 'source_conflict'):
|
|
self.queue(replace(self.source, **change))
|
|
with patch.object(self.processing, 'policy', replace(self.policy, rule_set_sha256='a' * 64)), \
|
|
self.assertRaisesRegex(source.CollectionError, 'source_conflict'):
|
|
self.queue()
|
|
pin = self.source.root / 'source.json'
|
|
raw = pin.read_bytes()
|
|
data = json.loads(raw)
|
|
data['page_size'] = float(data['page_size'])
|
|
pin.write_text(json.dumps(data))
|
|
with self.assertRaisesRegex(source.CollectionError, 'source_conflict'):
|
|
self.queue()
|
|
pin.write_bytes(raw)
|
|
pin.unlink()
|
|
with self.assertRaisesRegex(source.CollectionError, 'unbound_download_runtime_state'):
|
|
self.queue()
|
|
self.assertEqual(self.factory_calls, 0)
|
|
self.assertEqual(queue.get(REQUEST_ID)['status'], 'queued')
|
|
|
|
def test_v3_bounds_are_bound_before_dispatch(self):
|
|
self.queue(replace(self.source, capture_version='v3', max_profiles=3))
|
|
data = json.loads((self.source.root / 'source.json').read_bytes())
|
|
self.assertEqual((data['capture_version'], data['max_profiles']), ('v3', 3))
|
|
self.assertEqual(self.factory_calls, 0)
|
|
|
|
def launch(self, serve, *, configured=True, args=None, repository=None):
|
|
with patch.object(run.LoginCredentials, 'from_environment', return_value=TEST_CREDENTIALS), \
|
|
patch.object(run.PostgresPortalRepository, 'from_environment', return_value=repository or FakeRepository()), \
|
|
patch.object(run, 'compose_programmatic_processing', return_value=self.processing), \
|
|
patch.object(run, 'serve', side_effect=serve):
|
|
return run.main(['--enable-processing'] if args is None else args,
|
|
arr_source=self.source if configured else None)
|
|
|
|
def test_default_entry_has_no_download_queue(self):
|
|
def check(app, *_):
|
|
_, headers = login(app)
|
|
config = json.loads(app.handle('GET', '/api/arr-downloads', headers).body)['data']
|
|
self.assertFalse(config['ready'])
|
|
self.assertIsNone(config['latest_task'])
|
|
self.assertEqual(self.launch(check, configured=False), 0)
|
|
self.assertFalse(self.source.root.exists())
|
|
self.client.close.assert_called_once()
|
|
|
|
def test_injection_requires_processing_opt_in_and_ready_dependencies(self):
|
|
with self.assertRaisesRegex(SystemExit, '--enable-processing'):
|
|
self.launch(Mock(), args=[])
|
|
self.client.close.assert_not_called()
|
|
unavailable = Mock()
|
|
unavailable.list_months.side_effect = RuntimeError('private connection details')
|
|
with self.assertRaisesRegex(SystemExit, 'dependencies are unavailable'):
|
|
self.launch(Mock(), repository=unavailable)
|
|
self.assertFalse(self.source.root.exists())
|
|
self.client.close.assert_not_called()
|
|
|
|
def test_bad_source_fails_startup_without_leaking_details_and_closes_storage(self):
|
|
serve = Mock()
|
|
with patch.object(run, 'compose_arr_downloads', side_effect=ValueError('private source details')):
|
|
with self.assertRaises(SystemExit) as raised:
|
|
self.launch(serve)
|
|
self.assertEqual(str(raised.exception), 'ARR download source initialization failed')
|
|
serve.assert_not_called()
|
|
self.client.close.assert_called_once()
|
|
|
|
def test_application_constructor_failure_closes_queue_before_shared_storage(self):
|
|
order = []
|
|
queue = Mock()
|
|
queue.close.side_effect = lambda **kwargs: order.append(('queue', kwargs))
|
|
self.client.close.side_effect = lambda: order.append(('storage', {}))
|
|
with patch.object(run, 'compose_arr_downloads', return_value=queue), \
|
|
patch.object(run, 'PortalApplication', side_effect=RuntimeError('startup')):
|
|
with self.assertRaisesRegex(RuntimeError, 'startup'):
|
|
self.launch(Mock())
|
|
self.assertEqual(order, [('queue', {'wait': True}), ('storage', {})])
|
|
|
|
def test_configured_sigterm_restores_handler_after_owned_cleanup(self):
|
|
previous = signal.getsignal(signal.SIGTERM)
|
|
queue = Mock()
|
|
def terminate(*_):
|
|
handler = signal.getsignal(signal.SIGTERM)
|
|
self.assertIsNot(handler, previous)
|
|
handler(signal.SIGTERM, None)
|
|
with patch.object(run, 'compose_arr_downloads', return_value=queue):
|
|
with self.assertRaises(SystemExit) as raised:
|
|
self.launch(terminate)
|
|
self.assertEqual(raised.exception.code, 128 + signal.SIGTERM)
|
|
queue.close.assert_called_once_with(wait=True)
|
|
self.client.close.assert_called_once()
|
|
self.assertIs(signal.getsignal(signal.SIGTERM), previous)
|
|
|
|
def test_local_runtime_waits_before_allowing_owned_database_to_close(self):
|
|
entered, release, stopped = threading.Event(), threading.Event(), threading.Event()
|
|
def execute(**kwargs):
|
|
entered.set()
|
|
if not release.wait(10):
|
|
raise RuntimeError('test release timed out')
|
|
return DownloadOutcome('succeeded', 'arrjob-local-lifetime-test')
|
|
queue = PersistentARRDownloads(self.root / 'local-queue', SimpleNamespace(execute=execute))
|
|
self.addCleanup(queue.close, wait=True)
|
|
self.addCleanup(release.set)
|
|
runtime = object.__new__(ReplayRuntime)
|
|
runtime.queue, runtime.stop, runtime.thread = queue, threading.Event(), None
|
|
queue.create(DAY, REQUEST_ID)
|
|
self.assertTrue(entered.wait(2))
|
|
def shutdown():
|
|
runtime.close() # open_instance closes the owned DB immediately after this
|
|
stopped.set()
|
|
thread = threading.Thread(target=shutdown)
|
|
thread.start()
|
|
try:
|
|
self.assertFalse(stopped.wait(2.2))
|
|
finally:
|
|
release.set()
|
|
thread.join(5)
|
|
self.assertFalse(thread.is_alive())
|
|
self.assertTrue(stopped.is_set())
|
|
self.assertEqual(queue.get(REQUEST_ID)['status'], 'succeeded')
|
|
|
|
def test_serving_error_drains_active_delivery_before_storage_close(self):
|
|
entered, release, finished = threading.Event(), threading.Event(), threading.Event()
|
|
order = []
|
|
def execute(**kwargs):
|
|
entered.set()
|
|
if not release.wait(10):
|
|
raise RuntimeError('test release timed out')
|
|
self.client.close.assert_not_called()
|
|
order.append('delivery')
|
|
return DownloadOutcome('succeeded', 'arrjob-lifetime-test')
|
|
queue = PersistentARRDownloads(self.root / 'blocking-queue', SimpleNamespace(execute=execute))
|
|
self.addCleanup(queue.close, wait=True)
|
|
self.addCleanup(release.set)
|
|
def serve(app, *_):
|
|
_, headers = login(app)
|
|
response = app.handle('POST', '/api/arr-downloads', headers,
|
|
json.dumps({'request_id': REQUEST_ID, 'report_date': DAY}).encode())
|
|
self.assertEqual(response.status, 202)
|
|
self.assertTrue(entered.wait(2))
|
|
raise RuntimeError('test server stop')
|
|
self.client.close.side_effect = lambda: order.append('storage')
|
|
errors = []
|
|
def launch():
|
|
try:
|
|
with patch.object(run, 'compose_arr_downloads', return_value=queue):
|
|
self.launch(serve)
|
|
except BaseException as error:
|
|
errors.append(error)
|
|
finally:
|
|
finished.set()
|
|
thread = threading.Thread(target=launch)
|
|
thread.start()
|
|
self.assertTrue(entered.wait(2))
|
|
# Beyond the legacy2s close timeout: dependencies must still be alive.
|
|
self.assertFalse(finished.wait(2.2))
|
|
self.client.close.assert_not_called()
|
|
release.set()
|
|
thread.join(5)
|
|
self.assertFalse(thread.is_alive())
|
|
self.assertEqual(order, ['delivery', 'storage'])
|
|
self.assertEqual([str(error) for error in errors], ['test server stop'])
|
|
self.assertEqual(queue.get(REQUEST_ID)['status'], 'succeeded')
|
|
|
|
|
|
if __name__ == '__main__':
|
|
unittest.main()
|