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)