Files
wyndham-ARR/tests/test_monthly_reports_repository.py
2026-07-31 15:11:42 +08:00

379 lines
14 KiB
Python

from __future__ import annotations
import os
import unittest
from contextlib import contextmanager
from datetime import date
from decimal import Decimal
from typing import Any, Callable, List, Optional, Sequence, Tuple
from unittest.mock import patch
from monthly_reports.contracts import DailyVersionPin, ErrorCode
from monthly_reports.core import build_monthly_report
from monthly_reports.repository import (
DatabaseConfig,
FileMetadata,
PostgresReportRepository,
RepositoryError,
ReservedReport,
)
Handler = Callable[[str, Optional[Sequence[Any]]], Tuple[List[Tuple[Any, ...]], int]]
class FakeCursor:
def __init__(self, handler: Handler):
self.handler = handler
self.calls: List[Tuple[str, Optional[Sequence[Any]]]] = []
self.many_calls: List[Tuple[str, Sequence[Sequence[Any]]]] = []
self._rows: List[Tuple[Any, ...]] = []
self.rowcount = 0
def __enter__(self) -> "FakeCursor":
return self
def __exit__(self, *_args: Any) -> None:
return None
def execute(self, sql: str, params: Optional[Sequence[Any]] = None) -> None:
normalized = " ".join(sql.split())
self.calls.append((normalized, params))
self._rows, self.rowcount = self.handler(normalized, params)
def executemany(self, sql: str, params: Sequence[Sequence[Any]]) -> None:
normalized = " ".join(sql.split())
self.many_calls.append((normalized, params))
self._rows, self.rowcount = self.handler(normalized, params)
def fetchall(self) -> List[Tuple[Any, ...]]:
rows = list(self._rows)
self._rows = []
return rows
def fetchone(self) -> Optional[Tuple[Any, ...]]:
if not self._rows:
return None
row = self._rows[0]
self._rows = self._rows[1:]
return row
class FakeConnection:
def __init__(self, handler: Handler):
self.cursor_instance = FakeCursor(handler)
self.closed = False
@contextmanager
def transaction(self):
yield
def cursor(self) -> FakeCursor:
return self.cursor_instance
def close(self) -> None:
self.closed = True
def source_row() -> Tuple[Any, ...]:
return (
501,
102,
date(2026, 7, 2),
"QBD",
date(2026, 7, 2),
date(2026, 7, 4),
2,
2,
0,
"SYN-BLOCK",
1,
"SYN-COMPANY",
"SYN-CONFIRMATION",
"SYN-ROOM",
Decimal("450.00"),
"SYN-GUEST",
"SYN-COMMENT",
"",
"",
"SYN-RATE",
"SYN-TYPE",
Decimal("450.00"),
Decimal("900.00"),
None,
)
class MonthlyReportsRepositoryTests(unittest.TestCase):
def test_database_config_uses_specific_then_arr_fallback(self):
with patch.dict(os.environ, {"ARR_DATABASE_URL": "postgresql://arr"}, clear=True):
self.assertEqual(DatabaseConfig.from_environment().dsn, "postgresql://arr")
with patch.dict(
os.environ,
{
"ARR_DATABASE_URL": "postgresql://arr",
"MONTHLY_REPORT_DATABASE_URL": "postgresql://monthly",
},
clear=True,
):
self.assertEqual(
DatabaseConfig.from_environment().dsn,
"postgresql://monthly",
)
with patch.dict(os.environ, {}, clear=True):
with self.assertRaises(RepositoryError) as caught:
DatabaseConfig.from_environment()
self.assertEqual(caught.exception.code, ErrorCode.REQUEST_INVALID)
def test_snapshot_is_repeatable_and_maps_current_manifest(self):
def handler(sql: str, _params: Optional[Sequence[Any]]):
if sql == "SELECT current_database()":
return [("booking_test",)], 1
if "FROM finance.v_active_daily_facts AS facts" in sql:
return [source_row()], 1
if "FROM finance.current_daily_versions WHERE" in sql:
return [(date(2026, 7, 2), 102)], 1
if (
"JOIN finance.daily_channel_metrics AS metrics" in sql
and "current_version.business_date," in sql
):
return [(date(2026, 7, 2), 102, "QBD", 1, 1)], 1
if "SELECT metrics.channel_key" in sql:
return [("LIANTAI-GROUP",), ("QBD",), ("HANA",)], 3
return [], 0
connection = FakeConnection(handler)
repository = PostgresReportRepository(
DatabaseConfig("postgresql://synthetic"),
connect=lambda _dsn: connection,
)
snapshot = repository.load_snapshot(2026, 7, date(2026, 7, 2))
self.assertTrue(connection.closed)
self.assertEqual(snapshot.facts[0].daily_record_id, 501)
self.assertEqual(snapshot.facts[0].total_price, Decimal("900.00"))
self.assertEqual(snapshot.daily_versions[0].daily_version_id, 102)
self.assertEqual(snapshot.channel_observations[0].row_count, 1)
self.assertEqual(snapshot.preferred_channel_order[-1], "HANA")
executed = " ".join(sql for sql, _params in connection.cursor_instance.calls)
self.assertIn("REPEATABLE READ READ ONLY", executed)
self.assertIn("finance.v_active_daily_facts", executed)
self.assertIn("facts.arrival BETWEEN", executed)
def test_daily_event_scope_uses_arrival_and_month_maximum(self):
def handler(sql: str, _params: Optional[Sequence[Any]]):
if sql == "SELECT current_database()":
return [("booking_test",)], 1
if "SELECT min(record.arrival), max(record.arrival)" in sql:
return [(date(2026, 7, 21), date(2026, 7, 21))], 1
if "SELECT max(fact.arrival)" in sql:
return [(date(2026, 7, 27),)], 1
return [], 0
connection = FakeConnection(handler)
repository = PostgresReportRepository(
DatabaseConfig("postgresql://synthetic"),
connect=lambda _dsn: connection,
)
request = repository.derive_monthly_request(4)
self.assertIsNotNone(request)
self.assertEqual(request.report_year, 2026)
self.assertEqual(request.report_month, 7)
self.assertEqual(request.as_of_date, date(2026, 7, 27))
executed = " ".join(sql for sql, _params in connection.cursor_instance.calls)
self.assertIn("record.arrival", executed)
self.assertIn("fact.arrival", executed)
def test_reservation_freezes_daily_and_channel_manifests(self):
snapshot_connection = FakeConnection(
lambda sql, _params: (
[("booking_test",)],
1,
)
if sql == "SELECT current_database()"
else (
[source_row()],
1,
)
if "FROM finance.v_active_daily_facts AS facts" in sql
else (
[(date(2026, 7, 2), 102)],
1,
)
if "FROM finance.current_daily_versions WHERE" in sql
else (
[(date(2026, 7, 2), 102, "QBD", None, 1)],
1,
)
if (
"JOIN finance.daily_channel_metrics AS metrics" in sql
and "current_version.business_date," in sql
)
else ([], 0)
)
repository = PostgresReportRepository(
DatabaseConfig("postgresql://synthetic"),
connect=lambda _dsn: snapshot_connection,
)
report = build_monthly_report(
2026,
7,
date(2026, 7, 2),
repository.load_snapshot(2026, 7, date(2026, 7, 2)),
)
def reserve_handler(sql: str, _params: Optional[Sequence[Any]]):
if sql == "SELECT current_database()":
return [("booking_test",)], 1
if "FROM finance.current_daily_versions" in sql:
return [(date(2026, 7, 2), 102)], 1
if "FROM reporting.monthly_runs AS run" in sql:
return [], 0
if "SELECT COALESCE(max(version_no), 0) + 1" in sql:
return [(2,)], 1
if "INSERT INTO reporting.monthly_runs" in sql:
return [(44,)], 1
return [], 1
reserve_connection = FakeConnection(reserve_handler)
repository = PostgresReportRepository(
DatabaseConfig("postgresql://synthetic"),
connect=lambda _dsn: reserve_connection,
)
reservation = repository.reserve_report(report)
self.assertEqual(reservation.report_version_id, 44)
self.assertEqual(reservation.version_no, 2)
self.assertEqual(reservation.period_start, date(2026, 7, 1))
self.assertEqual(reservation.as_of_date, date(2026, 7, 2))
self.assertEqual(reservation.daily_versions, report.daily_versions)
self.assertRegex(reservation.source_snapshot_sha256 or "", r"^[0-9a-f]{64}$")
self.assertTrue(reserve_connection.closed)
self.assertEqual(len(reserve_connection.cursor_instance.many_calls), 0)
executed = " ".join(
sql for sql, _params in reserve_connection.cursor_instance.calls
)
self.assertNotIn("finance.report_versions", executed)
self.assertIn("INSERT INTO reporting.monthly_runs", executed)
self.assertIn("INSERT INTO reporting.monthly_run_daily_versions", executed)
self.assertIn("INSERT INTO reporting.monthly_channel_manifest", executed)
def test_stale_snapshot_is_rejected_before_version_insert(self):
from monthly_reports.contracts import (
ChannelObservation,
DailyVersionPin,
MonthlySnapshot,
)
row = source_row()
from monthly_reports.repository import _decimal
from monthly_reports.contracts import MonthlyFact
source = MonthlyFact(
daily_record_id=row[0], daily_version_id=row[1], business_date=row[2],
channel_key=row[3], arrival=row[4], departure=row[5], nights=row[6],
adults=row[7], children=row[8], block_code=row[9], no_of_rooms=row[10],
company_name=row[11], confirmation_no=row[12], disp_room_no=row[13],
effective_rate_amount=_decimal(row[14]), full_name=row[15],
res_comment=row[16], trace_text=row[17], products=row[18], rate_code=row[19],
room_category_label=row[20], real_price=_decimal(row[21]),
total_price=_decimal(row[22]), kb_amount=row[23],
)
report = build_monthly_report(
2026,
7,
date(2026, 7, 2),
MonthlySnapshot(
(source,),
(DailyVersionPin(date(2026, 7, 2), 102),),
(ChannelObservation(date(2026, 7, 2), 102, "QBD", 1, 1),),
),
)
def handler(sql: str, _params: Optional[Sequence[Any]]):
if sql == "SELECT current_database()":
return [("booking_test",)], 1
if "FROM finance.current_daily_versions" in sql:
return [(date(2026, 7, 2), 999)], 1
return [], 0
connection = FakeConnection(handler)
repository = PostgresReportRepository(
DatabaseConfig("postgresql://synthetic"),
connect=lambda _dsn: connection,
)
with self.assertRaises(RepositoryError) as caught:
repository.reserve_report(report)
self.assertEqual(caught.exception.code, ErrorCode.SOURCE_SNAPSHOT_STALE)
self.assertFalse(
any(
"INSERT INTO finance.report_versions" in sql
for sql, _params in connection.cursor_instance.calls
)
)
def test_activation_registers_artifacts_and_supersedes_atomically(self):
def handler(sql: str, params: Optional[Sequence[Any]]):
if sql == "SELECT current_database()":
return [("booking_test",)], 1
if "FROM finance.current_daily_versions" in sql:
return [(date(2026, 7, 2), 102)], 1
if "SELECT report_status, source_snapshot_sha256" in sql:
return [("reserved", "c" * 64)], 1
if "FROM ingestion.artifacts" in sql:
return [], 0
if "INSERT INTO ingestion.artifacts" in sql:
artifact_id = 701 if params and params[0] == "monthly_xlsx" else 702
return [(artifact_id,)], 1
return [], 1
connection = FakeConnection(handler)
repository = PostgresReportRepository(
DatabaseConfig("postgresql://synthetic"),
connect=lambda _dsn: connection,
)
repository.activate_report(
ReservedReport(
44,
44,
date(2026, 7, 1),
date(2026, 7, 2),
(DailyVersionPin(date(2026, 7, 2), 102),),
"c" * 64,
),
FileMetadata(
"monthly_xlsx",
"report.xlsx",
"outputs/monthly_reports/2026/07/archive/v0002/report.xlsx",
"a" * 64,
100,
"application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
),
FileMetadata(
"result_json",
"v0002.result.json",
"outputs/monthly_reports/2026/07/archive/v0002/v0002.result.json",
"b" * 64,
50,
"application/json",
),
"d" * 64,
)
self.assertTrue(connection.closed)
executed = " ".join(sql for sql, _params in connection.cursor_instance.calls)
self.assertIn("finance.current_daily_versions", executed)
self.assertEqual(executed.count("INSERT INTO ingestion.artifacts"), 2)
self.assertIn("SET report_status = 'superseded'", executed)
self.assertIn("SET report_status = 'active'", executed)
self.assertIn("workbook_artifact_id", executed)
self.assertNotIn("finance.report_versions", executed)
self.assertNotIn("booking.file_objects", executed)
if __name__ == "__main__":
unittest.main()