merchant

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

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