test_order_sequence_migrations.py (16457B)
1 #!/usr/bin/env python3 2 3 # This file is part of TALER 4 # Copyright (C) 2026 Taler Systems SA 5 # 6 # TALER is free software; you can redistribute it and/or modify it under the 7 # terms of the GNU General Public License as published by the Free Software 8 # Foundation; either version 3, or (at your option) any later version. 9 # 10 # TALER is distributed in the hope that it will be useful, but WITHOUT ANY 11 # WARRANTY; without even the implied warranty of MERCHANTABILITY or FITNESS FOR 12 # A PARTICULAR PURPOSE. See the GNU General Public License for more details. 13 # 14 # You should have received a copy of the GNU General Public License along with 15 # TALER; see the file COPYING. If not, see <http://www.gnu.org/licenses/> 16 17 """Regression tests for the order ID reset introduced by migration 0036. 18 19 Unclaimed orders live in merchant_orders. Paid contracts can remain in 20 merchant_contract_terms after the corresponding order rows have expired. 21 The next order ID must therefore exceed IDs in both tables. 22 23 Each test clones an empty database at version 35 or 46, seeds the relevant 24 history, and runs the real migration SQL. These are database migration tests; 25 the contract insertion check does not exercise the HTTP claim endpoint. 26 """ 27 28 from contextlib import contextmanager 29 import os 30 from pathlib import Path 31 import pwd 32 import shutil 33 import subprocess 34 import sys 35 import tempfile 36 import unittest 37 38 39 MAX_SERIAL = 9223372036854775807 40 41 42 def run(command, **kwargs): 43 result = subprocess.run(command, text=True, capture_output=True, **kwargs) 44 if result.returncode: 45 raise RuntimeError(f"{command[0]} failed:\n{result.stdout}\n{result.stderr}") 46 return result.stdout.strip() 47 48 49 @contextmanager 50 def postgres_cluster(bindir, *, max_locks=64): 51 """Keep all test data in a disposable server, accessible only by Unix socket.""" 52 with tempfile.TemporaryDirectory(prefix="merchant-seq-", dir="/tmp") as tmp: 53 data = Path(tmp) / "data" 54 server_options = {"cwd": tmp} 55 if os.geteuid() == 0: 56 # CI runs as root, but PostgreSQL requires an unprivileged server. 57 # Keep SQL clients as the caller so they can read the checkout. 58 try: 59 account = pwd.getpwnam("postgres") 60 except KeyError: 61 raise RuntimeError("Root execution requires the postgres account") from None 62 if account.pw_uid == 0: 63 raise RuntimeError("The postgres account must be unprivileged") 64 os.chown(tmp, account.pw_uid, account.pw_gid) 65 server_options.update(user=account.pw_uid, group=account.pw_gid, 66 extra_groups=[]) 67 env = {key: value for key, value in os.environ.items() 68 if not key.startswith("PG")} 69 env.update(PGHOST=tmp, PGPORT="5432", PGUSER="postgres", 70 PGOPTIONS="-c client_min_messages=warning") 71 run([str(bindir / "initdb"), "-D", str(data), "-A", "trust", 72 "-U", "postgres", "--no-locale"], env=env, **server_options) 73 try: 74 run([str(bindir / "pg_ctl"), "-D", str(data), 75 "-l", str(Path(tmp) / "server.log"), 76 "-o", f"-F -k {tmp} -c listen_addresses='' " 77 f"-c max_locks_per_transaction={max_locks}", 78 "-w", "start"], 79 env=env, **server_options) 80 yield env 81 finally: 82 if (data / "postmaster.pid").exists(): 83 run([str(bindir / "pg_ctl"), "-D", str(data), "-m", "immediate", 84 "-w", "stop"], env=env, **server_options) 85 86 87 class Database: 88 """A fixed database connection; choosing another database creates a new handle.""" 89 90 def __init__(self, name, bindir, env, sql_dir): 91 self.name = name 92 self.command = [str(bindir / "psql"), "-X", "-qAt", "-v", "ON_ERROR_STOP=1"] 93 self.env = dict(env, PGDATABASE=name) 94 self.sql_dir = sql_dir 95 96 def sql(self, statement): 97 return run(self.command, input=statement, env=self.env) 98 99 def apply_migration(self, version): 100 self.apply_file(self.sql_dir / f"merchant-{version:04}.sql") 101 102 def apply_file(self, path): 103 return run(self.command + ["-f", str(path)], env=self.env) 104 105 def add_instance(self, number, *, legacy=False): 106 self.sql(f""" 107 INSERT INTO merchant.merchant_instances 108 (merchant_serial, merchant_id, merchant_name, merchant_pub, 109 address, jurisdiction, default_wire_transfer_delay, default_pay_delay) 110 VALUES ({number}, 'test-{number}', 'Test', 111 decode(lpad(to_hex({number}),64,'0'),'hex'), '{{}}', '{{}}', 1, 1) 112 """) 113 # Runtime procedure bundles are not loaded in this migration-only fixture. 114 # Invoke the schema constructor explicitly instead of its runtime trigger. 115 if not legacy: 116 self.sql(f"SELECT merchant.create_instance_schema({number})") 117 return OrderFixture(self, number, legacy=legacy) 118 119 120 class OrderFixture: 121 """Seed only the order columns needed for the sequence migration scenarios.""" 122 123 def __init__(self, db, instance, *, legacy=False): 124 self.db = db 125 self.schema = "merchant" if legacy else f"merchant_instance_{instance}" 126 self.sequence = f"{self.schema}.merchant_orders_order_serial_seq" 127 self.instance_column = "merchant_serial," if legacy else "" 128 self.instance_value = f"{instance}," if legacy else "" 129 # Old statistics triggers need runtime procedures absent from the fixture. 130 self.seed_setup = "SET session_replication_role=replica;" if legacy else "" 131 132 def add_order(self, serial=None): 133 """An explicit serial seeds history; omitting it exercises ID allocation.""" 134 serial_value = "DEFAULT" if serial is None else str(serial) 135 order_id = "new-order" if serial is None else f"order-{serial}" 136 return int(self.db.sql(f""" 137 {self.seed_setup} 138 INSERT INTO {self.schema}.merchant_orders 139 ({self.instance_column}order_serial, order_id, claim_token, 140 h_post_data, pay_deadline, creation_time, contract_terms) 141 VALUES ({self.instance_value}{serial_value}, '{order_id}', 142 decode(repeat('01',16),'hex'), decode(repeat('02',64),'hex'), 143 2000000000000000, 1788800253000000, '{{}}') 144 RETURNING order_serial 145 """)) 146 147 def add_paid_contract(self, serial): 148 self.db.sql(f""" 149 {self.seed_setup} 150 INSERT INTO {self.schema}.merchant_contract_terms 151 ({self.instance_column}order_serial, order_id, contract_terms, 152 h_contract_terms, creation_time, pay_deadline, refund_deadline, 153 claim_token, paid) 154 VALUES ({self.instance_value}{serial}, 'order-{serial}', '{{}}', 155 decode(lpad(to_hex({serial}),128,'0'),'hex'), 156 1788719948000000, 2000000000000000, 2000000000000000, 157 decode(repeat('01',16),'hex'), true) 158 """) 159 160 def set_sequence(self, *, last_value, is_called): 161 self.db.sql(f"SELECT setval('{self.sequence}', {last_value}, " 162 f"{str(is_called).lower()})") 163 164 def sequence_state(self): 165 last_value, is_called = self.db.sql( 166 f"SELECT last_value, is_called FROM {self.sequence}" 167 ).split("|") 168 return int(last_value), is_called == "t" 169 170 def snapshot(self): 171 return self.db.sql(f""" 172 SELECT jsonb_agg(to_jsonb(t) ORDER BY order_serial) 173 FROM {self.schema}.merchant_orders t; 174 SELECT jsonb_agg(to_jsonb(t) ORDER BY order_serial) 175 FROM {self.schema}.merchant_contract_terms t; 176 """) 177 178 def repair_statement(self): 179 return f"CALL merchant.merchant_0047_init('{self.schema}');" 180 181 182 class OrderSequenceMigrations(unittest.TestCase): 183 """Each test gets its own clone; no scenario depends on an earlier test.""" 184 185 def database_before(self, version): 186 self.admin.sql(f"CREATE DATABASE {self._testMethodName} " 187 f"TEMPLATE before_{version}") 188 return Database(self._testMethodName, self.bindir, self.cluster_env, 189 self.sql_dir) 190 191 def assert_sequence(self, orders, *, last_value, is_called): 192 self.assertEqual(orders.sequence_state(), (last_value, is_called), 193 f"Unexpected sequence state in {orders.schema}") 194 195 def test_0036_keeps_ids_from_both_tables(self): 196 db = self.database_before(36) 197 paid = db.add_instance(1, legacy=True) 198 unpaid = db.add_instance(2, legacy=True) 199 db.add_instance(3, legacy=True) 200 paid.add_paid_contract(78) # The corresponding order has expired. 201 unpaid.add_order(90) 202 203 db.apply_migration(36) 204 205 self.assertEqual(OrderFixture(db, 1).add_order(), 79) 206 self.assertEqual(OrderFixture(db, 2).add_order(), 91) 207 self.assertEqual(OrderFixture(db, 3).add_order(), 1) 208 209 def check_0036_preserves_sequence(self, *, is_called, expected_next): 210 db = self.database_before(36) 211 paid = db.add_instance(1, legacy=True) 212 db.add_instance(2, legacy=True) 213 paid.add_paid_contract(78) 214 paid.set_sequence(last_value=200, is_called=is_called) 215 216 db.apply_migration(36) 217 218 # Both new sequences inherit the shared sequence's higher position. 219 self.assertEqual(OrderFixture(db, 1).add_order(), expected_next) 220 self.assertEqual(OrderFixture(db, 2).add_order(), expected_next) 221 222 def test_0036_preserves_called_sequence(self): 223 self.check_0036_preserves_sequence(is_called=True, expected_next=201) 224 225 def test_0036_preserves_uncalled_sequence(self): 226 self.check_0036_preserves_sequence(is_called=False, expected_next=200) 227 228 def test_0036_rejects_exhausted_sequence(self): 229 db = self.database_before(36) 230 orders = db.add_instance(1, legacy=True) 231 orders.set_sequence(last_value=MAX_SERIAL, is_called=True) 232 233 with self.assertRaisesRegex(RuntimeError, "Order serial sequence exhausted"): 234 db.apply_migration(36) 235 236 self.assertEqual(db.sql("SELECT count(*) FROM _v.patches " 237 "WHERE patch_name='merchant-0036'"), "0") 238 239 def test_0047_repairs_restarted_ids_without_changing_orders(self): 240 db = self.database_before(47) 241 paid = db.add_instance(1) 242 paid.add_paid_contract(78) 243 # Reproduce the reported history: new IDs 1-4 follow historical ID 78. 244 for serial in range(1, 5): 245 paid.add_order(serial) 246 paid.add_paid_contract(serial) 247 paid.set_sequence(last_value=4, is_called=True) 248 unpaid = db.add_instance(2) 249 unpaid.add_paid_contract(80) 250 unpaid.add_order(90) 251 paid_before, unpaid_before = paid.snapshot(), unpaid.snapshot() 252 253 db.apply_migration(47) 254 255 self.assert_sequence(paid, last_value=79, is_called=False) 256 self.assert_sequence(unpaid, last_value=91, is_called=False) 257 self.assertEqual(paid.snapshot(), paid_before) 258 self.assertEqual(unpaid.snapshot(), unpaid_before) 259 260 # Reapplying the registered fixup must not consume or rewind IDs. 261 db.sql("CALL merchant.fixup_instance_schema(47::INT8)") 262 self.assert_sequence(paid, last_value=79, is_called=False) 263 self.assert_sequence(unpaid, last_value=91, is_called=False) 264 self.assertEqual(paid.add_order(), 79) 265 self.assertEqual(unpaid.add_order(), 91) 266 267 # At the database level, claiming can copy the new serial without a 268 # primary-key collision. This is not an HTTP/backend claim test. 269 db.sql(f""" 270 INSERT INTO {paid.schema}.merchant_contract_terms 271 (order_serial, order_id, contract_terms, h_contract_terms, 272 creation_time, pay_deadline, refund_deadline, claim_token) 273 SELECT order_serial, order_id, contract_terms, 274 decode(repeat('ff',64),'hex'), creation_time, 275 pay_deadline, pay_deadline, claim_token 276 FROM {paid.schema}.merchant_orders WHERE order_serial=79 277 """) 278 self.assertEqual(db.sql( 279 f"SELECT order_serial FROM {paid.schema}.merchant_contract_terms " 280 "WHERE order_id='new-order'"), "79") 281 282 def test_0047_preserves_safe_sequence_states(self): 283 db = self.database_before(47) 284 called = db.add_instance(1) 285 called.set_sequence(last_value=200, is_called=True) 286 uncalled = db.add_instance(2) 287 uncalled.set_sequence(last_value=200, is_called=False) 288 just_above_history = db.add_instance(3) 289 just_above_history.add_paid_contract(78) 290 just_above_history.set_sequence(last_value=79, is_called=False) 291 292 db.apply_migration(47) 293 294 self.assert_sequence(called, last_value=200, is_called=True) 295 self.assert_sequence(uncalled, last_value=200, is_called=False) 296 self.assert_sequence(just_above_history, last_value=79, is_called=False) 297 298 def test_0047_keeps_empty_and_new_instances_starting_at_one(self): 299 db = self.database_before(47) 300 empty = db.add_instance(1) 301 302 db.apply_migration(47) 303 new = db.add_instance(2) 304 305 self.assert_sequence(empty, last_value=1, is_called=False) 306 self.assertEqual(empty.add_order(), 1) 307 self.assertEqual(new.add_order(), 1) 308 309 def test_0047_sequence_restart_rolls_back(self): 310 db = self.database_before(47) 311 orders = db.add_instance(1) 312 orders.add_paid_contract(78) 313 db.apply_migration(47) 314 orders.set_sequence(last_value=4, is_called=True) 315 316 db.sql(f"BEGIN; {orders.repair_statement()} ROLLBACK;") 317 318 self.assert_sequence(orders, last_value=4, is_called=True) 319 320 def test_0047_later_failure_rolls_back_earlier_repair(self): 321 db = self.database_before(47) 322 repairable = db.add_instance(1) 323 exhausted = db.add_instance(2) 324 repairable.add_paid_contract(78) 325 db.apply_migration(47) 326 repairable.set_sequence(last_value=4, is_called=True) 327 exhausted.set_sequence(last_value=MAX_SERIAL, is_called=True) 328 329 with self.assertRaisesRegex(RuntimeError, "Order serial sequence exhausted"): 330 db.sql(f"BEGIN; {repairable.repair_statement()} " 331 f"{exhausted.repair_statement()} COMMIT;") 332 333 self.assert_sequence(repairable, last_value=4, is_called=True) 334 335 def test_0047_rejects_exhausted_stored_ids(self): 336 db = self.database_before(47) 337 orders = db.add_instance(1) 338 orders.add_order(MAX_SERIAL) 339 340 with self.assertRaisesRegex(RuntimeError, "Order serial sequence exhausted"): 341 db.apply_migration(47) 342 343 self.assert_sequence(orders, last_value=1, is_called=False) 344 345 346 def main(): 347 source, build = (Path(arg).resolve() for arg in sys.argv[1:]) 348 # Missing CI prerequisites must not silently disable migration coverage. 349 unavailable_status = 1 if os.geteuid() == 0 else 77 350 if not shutil.which("pg_config"): 351 print("PostgreSQL server tools unavailable") 352 return unavailable_status 353 bindir = Path(run(["pg_config", "--bindir"])) 354 if not all((bindir / tool).exists() for tool in ("initdb", "pg_ctl", "psql")): 355 print("PostgreSQL server tools unavailable") 356 return unavailable_status 357 358 with postgres_cluster(bindir) as env: 359 admin = Database("template1", bindir, env, build) 360 admin.sql("CREATE DATABASE before_36") 361 before_36 = Database("before_36", bindir, env, build) 362 before_36.apply_file(source / "versioning.sql") 363 for version in range(1, 36): 364 before_36.apply_migration(version) 365 admin.sql("CREATE DATABASE before_47 TEMPLATE before_36") 366 before_47 = Database("before_47", bindir, env, build) 367 for version in range(36, 47): 368 before_47.apply_migration(version) 369 370 OrderSequenceMigrations.admin = admin 371 OrderSequenceMigrations.bindir = bindir 372 OrderSequenceMigrations.cluster_env = env 373 OrderSequenceMigrations.sql_dir = build 374 suite = unittest.defaultTestLoader.loadTestsFromTestCase(OrderSequenceMigrations) 375 result = unittest.TextTestRunner(verbosity=2).run(suite) 376 return 0 if result.wasSuccessful() else 1 377 378 379 if __name__ == "__main__": 380 sys.exit(main())