merchant

Merchant backend to process payments, run by merchants
Log | Files | Refs | Submodules | README | LICENSE

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())