feat: prepare ARR for controlled public deployment
This commit is contained in:
221
tests/test_channel_analytics_postgres.py
Normal file
221
tests/test_channel_analytics_postgres.py
Normal file
@@ -0,0 +1,221 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import unittest
|
||||
from contextlib import contextmanager
|
||||
from datetime import date, datetime, timezone
|
||||
from decimal import Decimal
|
||||
from typing import Any, Callable, List, Optional, Sequence, Tuple
|
||||
from unittest.mock import patch
|
||||
|
||||
from channel_analytics.contracts import AnalyticsError
|
||||
from channel_analytics.postgres import (
|
||||
CHANNEL_DETAIL_SQL,
|
||||
MONTHS_SQL,
|
||||
ROOM_AGGREGATES_SQL,
|
||||
AnalyticsRepositoryError,
|
||||
DatabaseConfig,
|
||||
PostgresAnalyticsRepository,
|
||||
)
|
||||
|
||||
|
||||
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._rows: List[Tuple[Any, ...]] = []
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *_args):
|
||||
return None
|
||||
|
||||
def execute(self, sql: str, params: Optional[Sequence[Any]] = None) -> None:
|
||||
normalized = " ".join(sql.split())
|
||||
self.calls.append((normalized, params))
|
||||
self._rows, _rowcount = self.handler(normalized, params)
|
||||
|
||||
def fetchall(self):
|
||||
rows = list(self._rows)
|
||||
self._rows = []
|
||||
return rows
|
||||
|
||||
def fetchone(self):
|
||||
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):
|
||||
return self.cursor_instance
|
||||
|
||||
def close(self):
|
||||
self.closed = True
|
||||
|
||||
|
||||
def handler(sql: str, _params: Optional[Sequence[Any]]):
|
||||
if sql.startswith("SELECT current_database()"):
|
||||
return [("booking_test", "on")], 1
|
||||
if "max(current_version.business_date) AS as_of_date" in sql:
|
||||
return [(
|
||||
date(2026, 7, 26),
|
||||
datetime(2026, 7, 28, 15, 54, tzinfo=timezone.utc),
|
||||
6,
|
||||
)], 1
|
||||
if "JOIN ingestion.artifacts AS source" in sql:
|
||||
return [
|
||||
(date(2026, 7, day), 100 + day, "a" * 64)
|
||||
for day in range(21, 27)
|
||||
], 6
|
||||
if "GROUP BY metrics.channel_key" in sql:
|
||||
return [("B", 2), ("A", 1), ("EMPTY", 0)], 3
|
||||
if "FROM finance.v_active_daily_facts AS facts" in sql and "GROUP BY" in sql:
|
||||
return [
|
||||
("A", "RM2", 1, 2, 6, Decimal("100")),
|
||||
("B", "RM3", 2, 3, 4, Decimal("90")),
|
||||
], 2
|
||||
if "FROM finance.v_active_daily_facts AS facts" in sql and "LIMIT" in sql:
|
||||
return [(
|
||||
date(2026, 7, 2),
|
||||
date(2026, 7, 4),
|
||||
2,
|
||||
1,
|
||||
"SYN-COMPANY",
|
||||
"SYN-RATE",
|
||||
"RM3",
|
||||
Decimal("45"),
|
||||
Decimal("90"),
|
||||
)], 1
|
||||
if "SELECT DISTINCT date_trunc('month', business_date)::date" in sql:
|
||||
return [(date(2026, 7, 1),)], 1
|
||||
return [], 0
|
||||
|
||||
|
||||
class ChannelAnalyticsPostgresTests(unittest.TestCase):
|
||||
def repository(self):
|
||||
connection = FakeConnection(handler)
|
||||
return (
|
||||
PostgresAnalyticsRepository(
|
||||
DatabaseConfig("postgresql://synthetic"),
|
||||
connect=lambda _dsn: connection,
|
||||
),
|
||||
connection,
|
||||
)
|
||||
|
||||
def test_database_config_uses_dashboard_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",
|
||||
"DASHBOARD_DATABASE_URL": "postgresql://dashboard",
|
||||
},
|
||||
clear=True,
|
||||
):
|
||||
self.assertEqual(DatabaseConfig.from_environment().dsn, "postgresql://dashboard")
|
||||
with patch.dict(os.environ, {}, clear=True):
|
||||
with self.assertRaises(AnalyticsRepositoryError):
|
||||
DatabaseConfig.from_environment()
|
||||
|
||||
def test_dashboard_uses_current_daily_versions_and_channel_order(self):
|
||||
repository, connection = self.repository()
|
||||
payload = repository.read_dashboard("2026-07")
|
||||
|
||||
self.assertTrue(connection.closed)
|
||||
self.assertEqual(
|
||||
[channel["worksheet"] for channel in payload["channels"]],
|
||||
["B", "A", "EMPTY"],
|
||||
)
|
||||
self.assertEqual(payload["overall"]["totals"]["rooms_sold"], 5)
|
||||
self.assertEqual(payload["overall"]["totals"]["total_price"], 190)
|
||||
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("finance.current_daily_versions", executed)
|
||||
self.assertIn("ingestion.artifacts", executed)
|
||||
|
||||
def test_channel_detail_is_paginated_and_contains_no_guest_identity(self):
|
||||
repository, connection = self.repository()
|
||||
payload = repository.read_channel_detail("2026-07", "B", limit=20, offset=0)
|
||||
|
||||
self.assertEqual(payload["version"], "1.0")
|
||||
self.assertEqual(payload["total_rows"], 2)
|
||||
self.assertEqual(payload["rows"][0]["total_price"], 90)
|
||||
self.assertEqual(
|
||||
set(payload["rows"][0]),
|
||||
{
|
||||
"arrival",
|
||||
"departure",
|
||||
"nights",
|
||||
"no_of_rooms",
|
||||
"company_name",
|
||||
"rate_code",
|
||||
"room_type",
|
||||
"real_price",
|
||||
"total_price",
|
||||
},
|
||||
)
|
||||
query = " ".join(sql for sql, _params in connection.cursor_instance.calls)
|
||||
for forbidden in ("full_name", "confirmation_no", "disp_room_no", "res_comment"):
|
||||
self.assertNotIn(forbidden, query.lower())
|
||||
|
||||
def test_manifest_mismatch_and_unknown_channel_fail_closed(self):
|
||||
def mismatched(sql: str, params: Optional[Sequence[Any]]):
|
||||
rows, count = handler(sql, params)
|
||||
if "GROUP BY metrics.channel_key" in sql:
|
||||
return [("B", 99), ("A", 1)], 2
|
||||
return rows, count
|
||||
|
||||
repository = PostgresAnalyticsRepository(
|
||||
DatabaseConfig("postgresql://synthetic"),
|
||||
connect=lambda _dsn: FakeConnection(mismatched),
|
||||
)
|
||||
with self.assertRaises(AnalyticsError) as caught:
|
||||
repository.read_dashboard("2026-07")
|
||||
self.assertEqual(caught.exception.code, "ANALYTICS_CHANNEL_MANIFEST_INVALID")
|
||||
|
||||
repository, _connection = self.repository()
|
||||
with self.assertRaises(AnalyticsError) as caught:
|
||||
repository.read_channel_detail("2026-07", "UNKNOWN")
|
||||
self.assertEqual(caught.exception.code, "ANALYTICS_CHANNEL_NOT_FOUND")
|
||||
|
||||
def test_month_list_is_current_and_privacy_minimized(self):
|
||||
repository, connection = self.repository()
|
||||
months = repository.list_months()
|
||||
self.assertEqual(months[0]["month_key"], "2026-07")
|
||||
self.assertEqual(months[0]["row_count"], 3)
|
||||
self.assertTrue(connection.closed)
|
||||
|
||||
def test_sql_surface_is_read_only_and_avoids_private_columns(self):
|
||||
sql = " ".join((ROOM_AGGREGATES_SQL, CHANNEL_DETAIL_SQL, MONTHS_SQL)).lower()
|
||||
for forbidden in (
|
||||
"full_name",
|
||||
"confirmation_no",
|
||||
"disp_room_no",
|
||||
"res_comment",
|
||||
"trace_text",
|
||||
" insert ",
|
||||
" update ",
|
||||
" delete ",
|
||||
):
|
||||
self.assertNotIn(forbidden, f" {sql} ")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user