226 lines
12 KiB
Python
226 lines
12 KiB
Python
"""Native replay is isolated from OHIP and retains actual processing/SQL boundaries."""
|
|
|
|
from datetime import date
|
|
import hashlib
|
|
import json
|
|
import os
|
|
from pathlib import Path
|
|
import tempfile
|
|
import time
|
|
import unittest
|
|
import sys
|
|
from unittest.mock import patch
|
|
|
|
from arr_ingestion.repository import InMemoryIngestionRepository
|
|
from arr_ingestion.service import IngestionService
|
|
from arr_ingestion.validation import DeliveryValidator
|
|
from arr_processing.policy import load_processor_policy
|
|
from arr_storage.filesystem import FilesystemObjectBackend
|
|
from arr_storage.store import ManagedObjectStore
|
|
from arr_web.app import PortalApplication
|
|
from arr_web.auth import LoginCredentials
|
|
from arr_web.local_replay import create_instance, open_instance
|
|
from arr_web.local_replay_database import ReplayDatabase
|
|
from arr_web.local_xml_replay import (COOKIE, LocalReplayPortal, NativeXMLSnapshot,
|
|
NativeXMLReplayExecutor, document)
|
|
from tests.test_arr_opera_daily_ingest import reservation, xml_document
|
|
from tests.test_ohip_processing_handoff import CountingProcessor
|
|
from integrations.ohip.collect_arr_source import CollectionError
|
|
|
|
|
|
DAY = date(2026, 7, 27)
|
|
REQUEST = "a" * 32
|
|
PORT = 18873
|
|
HOST = {"Host": f"127.0.0.1:{PORT}", "Content-Type": "application/json"}
|
|
|
|
|
|
def xml():
|
|
row = reservation(1).replace("<G_RESERVATION>",
|
|
"<G_RESERVATION><RESORT>LOCAL</RESORT><RESV_NAME_ID>1</RESV_NAME_ID>")
|
|
return xml_document(row).encode()
|
|
|
|
|
|
def login(app, credentials):
|
|
response = app.handle("POST", "/api/login", HOST, json.dumps(credentials).encode())
|
|
if response.status != 200:
|
|
raise AssertionError(response.status)
|
|
headers = {**HOST, "Cookie": response.headers["Set-Cookie"].split(";", 1)[0]}
|
|
session = app.handle("GET", "/api/session", headers)
|
|
headers["X-ARR-CSRF"] = json.loads(session.body)["data"]["csrf_token"]
|
|
return headers
|
|
|
|
|
|
class NativeReplayTests(unittest.TestCase):
|
|
@classmethod
|
|
def setUpClass(cls):
|
|
cls.policy = load_processor_policy(Path(__file__).resolve().parents[1])
|
|
|
|
def setUp(self):
|
|
temp = tempfile.TemporaryDirectory()
|
|
self.addCleanup(temp.cleanup)
|
|
self.root = Path(temp.name)
|
|
self.source = self.root / "input.xml"
|
|
self.raw = xml()
|
|
self.source.write_bytes(self.raw)
|
|
self.digest = hashlib.sha256(self.raw).hexdigest()
|
|
self.snapshot = NativeXMLSnapshot.create(self.root / "fixture", self.source, self.digest, DAY.isoformat())
|
|
self.repository = InMemoryIngestionRepository()
|
|
self.store = ManagedObjectStore(FilesystemObjectBackend(self.root / "objects", create=True))
|
|
self.service = IngestionService(DeliveryValidator(self.store, self.policy), self.repository)
|
|
self.processor = CountingProcessor(self.policy)
|
|
self.executor = NativeXMLReplayExecutor(self.root / "handoffs", self.snapshot, self.policy,
|
|
self.store, self.repository, self.service, self.processor)
|
|
|
|
def execute(self, **kwargs):
|
|
args = dict(request_id=REQUEST, from_date=DAY, to_date=DAY, report_stage=lambda _: None)
|
|
args.update(kwargs)
|
|
return self.executor.execute(**args)
|
|
|
|
def test_exact_copy_and_source_date_hash_required(self):
|
|
self.assertEqual(self.snapshot.payload(), self.raw)
|
|
self.assertEqual((self.root / "fixture/source.xml").stat().st_mode & 0o777, 0o600)
|
|
self.source.write_bytes(b"changed external file")
|
|
self.assertEqual(self.snapshot.payload(), self.raw)
|
|
for pin, day in (("0" * 64, DAY.isoformat()), (self.digest, "2026-07-28")):
|
|
self.source.write_bytes(self.raw)
|
|
with self.assertRaises(ValueError):
|
|
NativeXMLSnapshot.create(self.root / "bad", self.source, pin, day)
|
|
self.assertFalse((self.root / "bad").exists())
|
|
|
|
def test_source_tampering_stops_before_processing_or_database(self):
|
|
(self.root / "fixture/source.xml").write_bytes(self.raw + b" ")
|
|
with self.assertRaises(ValueError):
|
|
self.execute()
|
|
self.assertEqual(self.processor.calls, 0)
|
|
self.assertFalse(self.repository._jobs)
|
|
|
|
def test_unowned_database_directory_is_not_reused_or_erased(self):
|
|
candidate = self.root / "postgres"
|
|
candidate.mkdir(mode=0o700)
|
|
sentinel = candidate / "sentinel"
|
|
sentinel.write_text("untouched")
|
|
with (
|
|
patch.dict(sys.modules, {"psycopg": object()}),
|
|
patch("shutil.which", return_value="/unused"),
|
|
self.assertRaisesRegex(RuntimeError, "refusing_unowned_postgres_directory"),
|
|
):
|
|
ReplayDatabase(self.root).start()
|
|
self.assertEqual(sentinel.read_text(), "untouched")
|
|
self.assertFalse((self.root / "database-owner.json").exists())
|
|
|
|
def test_manifest_numeric_alias_is_not_accepted(self):
|
|
changed = dict(self.snapshot.manifest, source_records=True)
|
|
(self.snapshot.root / "source.json").write_text(json.dumps(changed))
|
|
with self.assertRaisesRegex(ValueError, "identity_changed"):
|
|
self.snapshot.payload()
|
|
|
|
def test_wrong_day_and_invalid_id_cannot_queue_or_process(self):
|
|
for args in ({"from_date": date(2026, 7, 28)}, {"request_id": "../../escape"}):
|
|
with self.assertRaises((ValueError, RuntimeError)):
|
|
self.execute(**args)
|
|
self.assertEqual(self.processor.calls, 0)
|
|
|
|
def test_unknown_commit_reuses_frozen_package_and_one_finance_version(self):
|
|
original = self.service.ingest
|
|
def lose_ack(*args, **kwargs):
|
|
original(*args, **kwargs)
|
|
raise ConnectionError("synthetic lost acknowledgement")
|
|
with patch.object(self.service, "ingest", side_effect=lose_ack), self.assertRaises(ConnectionError):
|
|
self.execute()
|
|
self.assertEqual(len(self.repository._versions), 1)
|
|
self.assertEqual(self.execute().status, "succeeded")
|
|
self.assertEqual(len(self.repository._versions), 1)
|
|
self.assertEqual(self.processor.calls, 1)
|
|
|
|
def test_replay_portal_auth_csrf_origin_cookie_and_labels(self):
|
|
credentials = {"username": "replay-test", "password": "synthetic-password-only"}
|
|
app = LocalReplayPortal(PortalApplication(credentials=LoginCredentials(**credentials)), self.snapshot, PORT)
|
|
self.assertEqual(app.handle("GET", "/api/session", HOST).status, 401)
|
|
self.assertEqual(app.handle("GET", "/api/session", {"Host": "untrusted.test"}).status, 403)
|
|
self.assertEqual(app.handle("POST", "/api/login", {**HOST, "Origin": "https://untrusted.test"}, b"{}").status, 403)
|
|
headers = login(app, credentials)
|
|
self.assertTrue(headers["Cookie"].startswith(COOKIE + "="))
|
|
wrong_cookie = {**headers, "Cookie": headers["Cookie"].replace(COOKIE, "arr_session")}
|
|
self.assertEqual(app.handle("GET", "/api/session", wrong_cookie).status, 401)
|
|
page = app.handle("GET", "/", headers)
|
|
self.assertIn("本机 XML 重放", page.body.decode())
|
|
self.assertIn("未连接 Oracle", page.body.decode())
|
|
status = json.loads(app.handle("GET", "/api/arr-downloads", headers).body)["data"]
|
|
self.assertFalse(status["oracle_connected"])
|
|
self.assertEqual(status["default_date"], DAY.isoformat())
|
|
no_csrf = {k: v for k, v in headers.items() if k != "X-ARR-CSRF"}
|
|
self.assertEqual(app.handle("POST", "/api/arr-downloads", no_csrf, b"{}").status, 403)
|
|
for route in ("/api/jobs", "/api/company-reports/source", "/api/monthly-runs"):
|
|
self.assertEqual(app.handle("POST", route, headers, b"{}").status, 405)
|
|
self.assertEqual(app.handle("GET", "/api/public/h5/months", HOST).status, 404)
|
|
|
|
|
|
@unittest.skipUnless(os.environ.get("ARR_TEST_LOCAL_POSTGRES") == "1", "owned local PostgreSQL opt-in required")
|
|
class ReplaySQLTests(unittest.TestCase):
|
|
def setUp(self):
|
|
temporary = tempfile.TemporaryDirectory()
|
|
self.addCleanup(temporary.cleanup)
|
|
parent = Path(temporary.name)
|
|
source = parent / "synthetic.xml"
|
|
source.write_bytes(xml())
|
|
self.root = create_instance(parent, source, hashlib.sha256(xml()).hexdigest(), DAY.isoformat())
|
|
|
|
@staticmethod
|
|
def wait(queue):
|
|
deadline = time.monotonic() + 30
|
|
while time.monotonic() < deadline:
|
|
value = queue.get(REQUEST)
|
|
if value["status"] not in {"queued", "downloading", "processing"}:
|
|
return value
|
|
time.sleep(.02)
|
|
raise AssertionError("replay_task_timeout")
|
|
|
|
def test_owned_sql_http_download_monthly_and_restart(self):
|
|
with patch.dict(os.environ, {"PGHOST": "invalid.example", "ARR_WEB_PASSWORD": "not-used"}):
|
|
with open_instance(self.root, PORT) as runtime:
|
|
headers = login(runtime.app, document(self.root / "login.json"))
|
|
response = runtime.app.handle("POST", "/api/arr-downloads", headers,
|
|
json.dumps({"request_id": REQUEST, "report_date": DAY.isoformat()}).encode())
|
|
self.assertEqual(response.status, 202)
|
|
result = self.wait(runtime.queue)
|
|
self.assertEqual(result["status"], "succeeded")
|
|
daily = runtime.app.handle("GET", "/api/download/daily?job_id=" + result["job_id"], headers)
|
|
self.assertEqual(daily.status, 200)
|
|
self.assertIn("LOCAL-REPLAY-", daily.headers["Content-Disposition"])
|
|
self.assertTrue(daily.body.startswith(b"PK"))
|
|
daily_hash = hashlib.sha256(daily.body).hexdigest()
|
|
self.assertEqual(runtime.worker.process_next().status, "published")
|
|
reports = runtime.portal_repository.list_monthly_runs("2026-07")[0]
|
|
self.assertEqual(len(reports), 1)
|
|
with runtime.database.connect() as connection:
|
|
report_id = connection.execute("SELECT id FROM reporting.monthly_runs").fetchone()[0]
|
|
self.assertEqual(connection.execute("SELECT count(*) FROM finance.daily_versions").fetchone()[0], 1)
|
|
monthly = runtime.app.handle("GET", f"/api/download/monthly?report_id={report_id}", headers)
|
|
self.assertEqual(monthly.status, 200)
|
|
self.assertTrue(monthly.body.startswith(b"PK"))
|
|
self.assertEqual(os.environ["PGHOST"], "invalid.example")
|
|
with open_instance(self.root, PORT) as runtime:
|
|
self.assertEqual(runtime.queue.get(REQUEST)["status"], "succeeded")
|
|
runtime.queue.create(DAY.isoformat(), REQUEST)
|
|
headers = login(runtime.app, document(self.root / "login.json"))
|
|
daily = runtime.app.handle("GET", "/api/download/daily?job_id=" + result["job_id"], headers)
|
|
self.assertEqual(hashlib.sha256(daily.body).hexdigest(), daily_hash)
|
|
with runtime.database.connect() as connection:
|
|
self.assertEqual(connection.execute("SELECT count(*) FROM finance.daily_versions").fetchone()[0], 1)
|
|
self.assertEqual(runtime.worker.process_next().status, "idle")
|
|
|
|
def test_wrong_date_fails_before_task_creation_and_concurrent_owner_rejected(self):
|
|
with open_instance(self.root, PORT) as runtime:
|
|
headers = login(runtime.app, document(self.root / "login.json"))
|
|
response = runtime.app.handle("POST", "/api/arr-downloads", headers,
|
|
json.dumps({"request_id": REQUEST, "report_date": "2026-07-28"}).encode())
|
|
self.assertEqual(response.status, 409)
|
|
self.assertIsNone(runtime.queue.latest())
|
|
with self.assertRaisesRegex(CollectionError, "batch_busy"):
|
|
with open_instance(self.root, PORT):
|
|
self.fail("second owner acquired the same instance")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|