Files
ARR-2.0-0918/tests/test_arr_local_xml_replay.py
T

235 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 Mock, 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, SourceDateDownloads, 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_date_wrapper_preserves_pending_review_list(self):
pending = [{"request_id": REQUEST, "report_date": DAY.isoformat(), "status": "needs_data_review"}]
queue = Mock()
queue.pending_data_reviews.return_value = pending
downloads = SourceDateDownloads(queue, self.snapshot)
self.assertEqual(downloads.pending_data_reviews(), pending)
queue.pending_data_reviews.assert_called_once_with()
queue.create.assert_not_called()
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()