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

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()