332 lines
12 KiB
Python
332 lines
12 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, 26))
|
|
|
|
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)
|
|
|
|
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, 26),
|
|
repository.load_snapshot(2026, 7, date(2026, 7, 26)),
|
|
)
|
|
|
|
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
|
|
return [], 1
|
|
|
|
reserve_connection = FakeConnection(reserve_handler)
|
|
repository = PostgresReportRepository(
|
|
DatabaseConfig("postgresql://synthetic"),
|
|
connect=lambda _dsn: reserve_connection,
|
|
)
|
|
reservation = repository.reserve_report(report)
|
|
|
|
self.assertGreater(reservation.report_version_id, 0)
|
|
self.assertEqual(
|
|
reservation.report_version_id,
|
|
reservation.version_no,
|
|
)
|
|
self.assertEqual(reservation.period_start, date(2026, 7, 1))
|
|
self.assertEqual(reservation.as_of_date, date(2026, 7, 26))
|
|
self.assertEqual(reservation.daily_versions, report.daily_versions)
|
|
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)
|
|
|
|
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, 26),
|
|
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_rechecks_pins_without_persisting_monthly_rows(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
|
|
return [], 0
|
|
|
|
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, 26),
|
|
(DailyVersionPin(date(2026, 7, 2), 102),),
|
|
),
|
|
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",
|
|
),
|
|
)
|
|
|
|
self.assertTrue(connection.closed)
|
|
executed = " ".join(sql for sql, _params in connection.cursor_instance.calls)
|
|
self.assertIn("finance.current_daily_versions", executed)
|
|
self.assertNotIn("finance.report_versions", executed)
|
|
self.assertNotIn("booking.file_objects", executed)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|