105 lines
3.4 KiB
Python
105 lines
3.4 KiB
Python
from __future__ import annotations
|
|
|
|
import unittest
|
|
|
|
from booking_ingestion.excel_postgres import BookingExcelRepositoryError, DatabaseConfig
|
|
from booking_ingestion.excel_review_postgres import PostgresBookingReviewRepository
|
|
|
|
|
|
DRAFT_ID = "bookingdraft-" + "a" * 32
|
|
|
|
|
|
class DeleteCursor:
|
|
def __init__(self, returned_ids: list[int]) -> None:
|
|
self.returned_ids = returned_ids
|
|
self.fetchone_value: object = None
|
|
self.executed: list[tuple[str, object]] = []
|
|
|
|
def __enter__(self) -> "DeleteCursor":
|
|
return self
|
|
|
|
def __exit__(self, *_args: object) -> None:
|
|
return None
|
|
|
|
def execute(self, query: str, params: object = None) -> None:
|
|
normalized = " ".join(query.split())
|
|
self.executed.append((normalized, params))
|
|
if "SELECT current_database()" in normalized:
|
|
self.fetchone_value = ("booking_test", "booking.current_source_batch")
|
|
elif "to_regclass('booking.extraction_drafts')" in normalized:
|
|
self.fetchone_value = (
|
|
"booking.extraction_drafts",
|
|
"booking.extraction_draft_items",
|
|
)
|
|
|
|
def fetchone(self) -> object:
|
|
return self.fetchone_value
|
|
|
|
def fetchall(self) -> list[tuple[int]]:
|
|
return [(item_id,) for item_id in self.returned_ids]
|
|
|
|
|
|
class DeleteConnection:
|
|
def __init__(self, returned_ids: list[int]) -> None:
|
|
self.cursor_instance = DeleteCursor(returned_ids)
|
|
self.committed = False
|
|
self.rolled_back = False
|
|
self.closed = False
|
|
|
|
def cursor(self) -> DeleteCursor:
|
|
return self.cursor_instance
|
|
|
|
def commit(self) -> None:
|
|
self.committed = True
|
|
|
|
def rollback(self) -> None:
|
|
self.rolled_back = True
|
|
|
|
def close(self) -> None:
|
|
self.closed = True
|
|
|
|
|
|
class PostgresBookingReviewDeleteTests(unittest.TestCase):
|
|
def repository(self, connection: DeleteConnection) -> PostgresBookingReviewRepository:
|
|
return PostgresBookingReviewRepository(
|
|
DatabaseConfig("postgresql://unused"),
|
|
connect=lambda _dsn: connection,
|
|
)
|
|
|
|
def test_bulk_delete_commits_all_selected_items_once(self) -> None:
|
|
connection = DeleteConnection([8, 7])
|
|
|
|
self.repository(connection).delete_items(DRAFT_ID, [7, 8])
|
|
|
|
self.assertTrue(connection.committed)
|
|
self.assertFalse(connection.rolled_back)
|
|
self.assertTrue(connection.closed)
|
|
update = next(
|
|
entry
|
|
for entry in connection.cursor_instance.executed
|
|
if "UPDATE booking.extraction_draft_items" in entry[0]
|
|
)
|
|
self.assertIn("item.id = ANY(%s::bigint[])", update[0])
|
|
self.assertEqual(update[1], ([7, 8], DRAFT_ID))
|
|
|
|
def test_partial_match_rolls_back_without_touching_draft_timestamp(self) -> None:
|
|
connection = DeleteConnection([7])
|
|
|
|
with self.assertRaises(BookingExcelRepositoryError) as raised:
|
|
self.repository(connection).delete_items(DRAFT_ID, [7, 8])
|
|
|
|
self.assertEqual(raised.exception.code, "BOOKING_EXCEL_REVIEW_ITEM_NOT_FOUND")
|
|
self.assertFalse(connection.committed)
|
|
self.assertTrue(connection.rolled_back)
|
|
self.assertTrue(connection.closed)
|
|
self.assertFalse(
|
|
any(
|
|
"UPDATE booking.extraction_drafts SET updated_at" in query
|
|
for query, _params in connection.cursor_instance.executed
|
|
)
|
|
)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main(verbosity=2)
|