test_order_sequence_migrations.py (16355B)
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): 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=''", "-w", "start"], 77 env=env, **server_options) 78 yield env 79 finally: 80 if (data / "postmaster.pid").exists(): 81 run([str(bindir / "pg_ctl"), "-D", str(data), "-m", "immediate", 82 "-w", "stop"], env=env, **server_options) 83 84 85 class Database: 86 """A fixed database connection; choosing another database creates a new handle.""" 87 88 def __init__(self, name, bindir, env, sql_dir): 89 self.name = name 90 self.command = [str(bindir / "psql"), "-X", "-qAt", "-v", "ON_ERROR_STOP=1"] 91 self.env = dict(env, PGDATABASE=name) 92 self.sql_dir = sql_dir 93 94 def sql(self, statement): 95 return run(self.command, input=statement, env=self.env) 96 97 def apply_migration(self, version): 98 self.apply_file(self.sql_dir / f"merchant-{version:04}.sql") 99 100 def apply_file(self, path): 101 return run(self.command + ["-f", str(path)], env=self.env) 102 103 def add_instance(self, number, *, legacy=False): 104 self.sql(f""" 105 INSERT INTO merchant.merchant_instances 106 (merchant_serial, merchant_id, merchant_name, merchant_pub, 107 address, jurisdiction, default_wire_transfer_delay, default_pay_delay) 108 VALUES ({number}, 'test-{number}', 'Test', 109 decode(lpad(to_hex({number}),64,'0'),'hex'), '{{}}', '{{}}', 1, 1) 110 """) 111 # Runtime procedure bundles are not loaded in this migration-only fixture. 112 # Invoke the schema constructor explicitly instead of its runtime trigger. 113 if not legacy: 114 self.sql(f"SELECT merchant.create_instance_schema({number})") 115 return OrderFixture(self, number, legacy=legacy) 116 117 118 class OrderFixture: 119 """Seed only the order columns needed for the sequence migration scenarios.""" 120 121 def __init__(self, db, instance, *, legacy=False): 122 self.db = db 123 self.schema = "merchant" if legacy else f"merchant_instance_{instance}" 124 self.sequence = f"{self.schema}.merchant_orders_order_serial_seq" 125 self.instance_column = "merchant_serial," if legacy else "" 126 self.instance_value = f"{instance}," if legacy else "" 127 # Old statistics triggers need runtime procedures absent from the fixture. 128 self.seed_setup = "SET session_replication_role=replica;" if legacy else "" 129 130 def add_order(self, serial=None): 131 """An explicit serial seeds history; omitting it exercises ID allocation.""" 132 serial_value = "DEFAULT" if serial is None else str(serial) 133 order_id = "new-order" if serial is None else f"order-{serial}" 134 return int(self.db.sql(f""" 135 {self.seed_setup} 136 INSERT INTO {self.schema}.merchant_orders 137 ({self.instance_column}order_serial, order_id, claim_token, 138 h_post_data, pay_deadline, creation_time, contract_terms) 139 VALUES ({self.instance_value}{serial_value}, '{order_id}', 140 decode(repeat('01',16),'hex'), decode(repeat('02',64),'hex'), 141 2000000000000000, 1788800253000000, '{{}}') 142 RETURNING order_serial 143 """)) 144 145 def add_paid_contract(self, serial): 146 self.db.sql(f""" 147 {self.seed_setup} 148 INSERT INTO {self.schema}.merchant_contract_terms 149 ({self.instance_column}order_serial, order_id, contract_terms, 150 h_contract_terms, creation_time, pay_deadline, refund_deadline, 151 claim_token, paid) 152 VALUES ({self.instance_value}{serial}, 'order-{serial}', '{{}}', 153 decode(lpad(to_hex({serial}),128,'0'),'hex'), 154 1788719948000000, 2000000000000000, 2000000000000000, 155 decode(repeat('01',16),'hex'), true) 156 """) 157 158 def set_sequence(self, *, last_value, is_called): 159 self.db.sql(f"SELECT setval('{self.sequence}', {last_value}, " 160 f"{str(is_called).lower()})") 161 162 def sequence_state(self): 163 last_value, is_called = self.db.sql( 164 f"SELECT last_value, is_called FROM {self.sequence}" 165 ).split("|") 166 return int(last_value), is_called == "t" 167 168 def snapshot(self): 169 return self.db.sql(f""" 170 SELECT jsonb_agg(to_jsonb(t) ORDER BY order_serial) 171 FROM {self.schema}.merchant_orders t; 172 SELECT jsonb_agg(to_jsonb(t) ORDER BY order_serial) 173 FROM {self.schema}.merchant_contract_terms t; 174 """) 175 176 def repair_statement(self): 177 return f"CALL merchant.merchant_0047_init('{self.schema}');" 178 179 180 class OrderSequenceMigrations(unittest.TestCase): 181 """Each test gets its own clone; no scenario depends on an earlier test.""" 182 183 def database_before(self, version): 184 self.admin.sql(f"CREATE DATABASE {self._testMethodName} " 185 f"TEMPLATE before_{version}") 186 return Database(self._testMethodName, self.bindir, self.cluster_env, 187 self.sql_dir) 188 189 def assert_sequence(self, orders, *, last_value, is_called): 190 self.assertEqual(orders.sequence_state(), (last_value, is_called), 191 f"Unexpected sequence state in {orders.schema}") 192 193 def test_0036_keeps_ids_from_both_tables(self): 194 db = self.database_before(36) 195 paid = db.add_instance(1, legacy=True) 196 unpaid = db.add_instance(2, legacy=True) 197 db.add_instance(3, legacy=True) 198 paid.add_paid_contract(78) # The corresponding order has expired. 199 unpaid.add_order(90) 200 201 db.apply_migration(36) 202 203 self.assertEqual(OrderFixture(db, 1).add_order(), 79) 204 self.assertEqual(OrderFixture(db, 2).add_order(), 91) 205 self.assertEqual(OrderFixture(db, 3).add_order(), 1) 206 207 def check_0036_preserves_sequence(self, *, is_called, expected_next): 208 db = self.database_before(36) 209 paid = db.add_instance(1, legacy=True) 210 db.add_instance(2, legacy=True) 211 paid.add_paid_contract(78) 212 paid.set_sequence(last_value=200, is_called=is_called) 213 214 db.apply_migration(36) 215 216 # Both new sequences inherit the shared sequence's higher position. 217 self.assertEqual(OrderFixture(db, 1).add_order(), expected_next) 218 self.assertEqual(OrderFixture(db, 2).add_order(), expected_next) 219 220 def test_0036_preserves_called_sequence(self): 221 self.check_0036_preserves_sequence(is_called=True, expected_next=201) 222 223 def test_0036_preserves_uncalled_sequence(self): 224 self.check_0036_preserves_sequence(is_called=False, expected_next=200) 225 226 def test_0036_rejects_exhausted_sequence(self): 227 db = self.database_before(36) 228 orders = db.add_instance(1, legacy=True) 229 orders.set_sequence(last_value=MAX_SERIAL, is_called=True) 230 231 with self.assertRaisesRegex(RuntimeError, "Order serial sequence exhausted"): 232 db.apply_migration(36) 233 234 self.assertEqual(db.sql("SELECT count(*) FROM _v.patches " 235 "WHERE patch_name='merchant-0036'"), "0") 236 237 def test_0047_repairs_restarted_ids_without_changing_orders(self): 238 db = self.database_before(47) 239 paid = db.add_instance(1) 240 paid.add_paid_contract(78) 241 # Reproduce the reported history: new IDs 1-4 follow historical ID 78. 242 for serial in range(1, 5): 243 paid.add_order(serial) 244 paid.add_paid_contract(serial) 245 paid.set_sequence(last_value=4, is_called=True) 246 unpaid = db.add_instance(2) 247 unpaid.add_paid_contract(80) 248 unpaid.add_order(90) 249 paid_before, unpaid_before = paid.snapshot(), unpaid.snapshot() 250 251 db.apply_migration(47) 252 253 self.assert_sequence(paid, last_value=79, is_called=False) 254 self.assert_sequence(unpaid, last_value=91, is_called=False) 255 self.assertEqual(paid.snapshot(), paid_before) 256 self.assertEqual(unpaid.snapshot(), unpaid_before) 257 258 # Reapplying the registered fixup must not consume or rewind IDs. 259 db.sql("CALL merchant.fixup_instance_schema(47::INT8)") 260 self.assert_sequence(paid, last_value=79, is_called=False) 261 self.assert_sequence(unpaid, last_value=91, is_called=False) 262 self.assertEqual(paid.add_order(), 79) 263 self.assertEqual(unpaid.add_order(), 91) 264 265 # At the database level, claiming can copy the new serial without a 266 # primary-key collision. This is not an HTTP/backend claim test. 267 db.sql(f""" 268 INSERT INTO {paid.schema}.merchant_contract_terms 269 (order_serial, order_id, contract_terms, h_contract_terms, 270 creation_time, pay_deadline, refund_deadline, claim_token) 271 SELECT order_serial, order_id, contract_terms, 272 decode(repeat('ff',64),'hex'), creation_time, 273 pay_deadline, pay_deadline, claim_token 274 FROM {paid.schema}.merchant_orders WHERE order_serial=79 275 """) 276 self.assertEqual(db.sql( 277 f"SELECT order_serial FROM {paid.schema}.merchant_contract_terms " 278 "WHERE order_id='new-order'"), "79") 279 280 def test_0047_preserves_safe_sequence_states(self): 281 db = self.database_before(47) 282 called = db.add_instance(1) 283 called.set_sequence(last_value=200, is_called=True) 284 uncalled = db.add_instance(2) 285 uncalled.set_sequence(last_value=200, is_called=False) 286 just_above_history = db.add_instance(3) 287 just_above_history.add_paid_contract(78) 288 just_above_history.set_sequence(last_value=79, is_called=False) 289 290 db.apply_migration(47) 291 292 self.assert_sequence(called, last_value=200, is_called=True) 293 self.assert_sequence(uncalled, last_value=200, is_called=False) 294 self.assert_sequence(just_above_history, last_value=79, is_called=False) 295 296 def test_0047_keeps_empty_and_new_instances_starting_at_one(self): 297 db = self.database_before(47) 298 empty = db.add_instance(1) 299 300 db.apply_migration(47) 301 new = db.add_instance(2) 302 303 self.assert_sequence(empty, last_value=1, is_called=False) 304 self.assertEqual(empty.add_order(), 1) 305 self.assertEqual(new.add_order(), 1) 306 307 def test_0047_sequence_restart_rolls_back(self): 308 db = self.database_before(47) 309 orders = db.add_instance(1) 310 orders.add_paid_contract(78) 311 db.apply_migration(47) 312 orders.set_sequence(last_value=4, is_called=True) 313 314 db.sql(f"BEGIN; {orders.repair_statement()} ROLLBACK;") 315 316 self.assert_sequence(orders, last_value=4, is_called=True) 317 318 def test_0047_later_failure_rolls_back_earlier_repair(self): 319 db = self.database_before(47) 320 repairable = db.add_instance(1) 321 exhausted = db.add_instance(2) 322 repairable.add_paid_contract(78) 323 db.apply_migration(47) 324 repairable.set_sequence(last_value=4, is_called=True) 325 exhausted.set_sequence(last_value=MAX_SERIAL, is_called=True) 326 327 with self.assertRaisesRegex(RuntimeError, "Order serial sequence exhausted"): 328 db.sql(f"BEGIN; {repairable.repair_statement()} " 329 f"{exhausted.repair_statement()} COMMIT;") 330 331 self.assert_sequence(repairable, last_value=4, is_called=True) 332 333 def test_0047_rejects_exhausted_stored_ids(self): 334 db = self.database_before(47) 335 orders = db.add_instance(1) 336 orders.add_order(MAX_SERIAL) 337 338 with self.assertRaisesRegex(RuntimeError, "Order serial sequence exhausted"): 339 db.apply_migration(47) 340 341 self.assert_sequence(orders, last_value=1, is_called=False) 342 343 344 def main(): 345 source, build = (Path(arg).resolve() for arg in sys.argv[1:]) 346 # Missing CI prerequisites must not silently disable migration coverage. 347 unavailable_status = 1 if os.geteuid() == 0 else 77 348 if not shutil.which("pg_config"): 349 print("PostgreSQL server tools unavailable") 350 return unavailable_status 351 bindir = Path(run(["pg_config", "--bindir"])) 352 if not all((bindir / tool).exists() for tool in ("initdb", "pg_ctl", "psql")): 353 print("PostgreSQL server tools unavailable") 354 return unavailable_status 355 356 with postgres_cluster(bindir) as env: 357 admin = Database("template1", bindir, env, build) 358 admin.sql("CREATE DATABASE before_36") 359 before_36 = Database("before_36", bindir, env, build) 360 before_36.apply_file(source / "versioning.sql") 361 for version in range(1, 36): 362 before_36.apply_migration(version) 363 admin.sql("CREATE DATABASE before_47 TEMPLATE before_36") 364 before_47 = Database("before_47", bindir, env, build) 365 for version in range(36, 47): 366 before_47.apply_migration(version) 367 368 OrderSequenceMigrations.admin = admin 369 OrderSequenceMigrations.bindir = bindir 370 OrderSequenceMigrations.cluster_env = env 371 OrderSequenceMigrations.sql_dir = build 372 suite = unittest.defaultTestLoader.loadTestsFromTestCase(OrderSequenceMigrations) 373 result = unittest.TextTestRunner(verbosity=2).run(suite) 374 return 0 if result.wasSuccessful() else 1 375 376 377 if __name__ == "__main__": 378 sys.exit(main())