merchant

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

test_order_sequence_migrations.py (19355B)


      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_reset_many_instance_schemas(self):
    196         db = self.database_before(47)
    197         for number in range(1, 25):
    198             db.add_instance(number)
    199         db.sql("CREATE SCHEMA merchant_instance_backup; "
    200                "CREATE SCHEMA merchantxinstancey1; "
    201                "SELECT _v.register_patch('other-0001', NULL, NULL)")
    202         # Simulate resuming a reset interrupted after one instance was dropped.
    203         db.sql("DROP SCHEMA merchant_instance_1 CASCADE")
    204 
    205         # The cluster uses PostgreSQL's default lock budget. A transaction
    206         # holding locks on every instance's tables cannot reset this database.
    207         db.apply_file(self.sql_dir / 'drop.sql')
    208         self.assertEqual(db.sql(
    209             "SELECT count(*) FROM pg_namespace WHERE nspname='merchant' "
    210             "OR nspname ~ '^merchant_instance_[0-9]+$'"), "0")
    211         self.assertEqual(db.sql(
    212             "SELECT count(*) FROM pg_namespace WHERE nspname IN "
    213             "('merchant_instance_backup', 'merchantxinstancey1')"), "2")
    214         self.assertEqual(db.sql(
    215             "SELECT patch_name FROM _v.patches ORDER BY patch_name"), "other-0001")
    216         db.apply_file(self.sql_dir / 'drop.sql')
    217 
    218     def test_reset_empty_database(self):
    219         self.admin.sql(f"CREATE DATABASE {self._testMethodName}")
    220         db = Database(self._testMethodName, self.bindir, self.cluster_env,
    221                       self.sql_dir)
    222         db.apply_file(self.sql_dir / 'drop.sql')
    223         db.apply_file(self.sql_dir / 'drop.sql')
    224 
    225     def test_0036_keeps_ids_from_both_tables(self):
    226         db = self.database_before(36)
    227         paid = db.add_instance(1, legacy=True)
    228         unpaid = db.add_instance(2, legacy=True)
    229         db.add_instance(3, legacy=True)
    230         paid.add_paid_contract(78)  # The corresponding order has expired.
    231         unpaid.add_order(90)
    232 
    233         db.apply_migration(36)
    234 
    235         self.assertEqual(OrderFixture(db, 1).add_order(), 79)
    236         self.assertEqual(OrderFixture(db, 2).add_order(), 91)
    237         self.assertEqual(OrderFixture(db, 3).add_order(), 1)
    238 
    239     def check_0036_preserves_sequence(self, *, is_called, expected_next):
    240         db = self.database_before(36)
    241         paid = db.add_instance(1, legacy=True)
    242         db.add_instance(2, legacy=True)
    243         paid.add_paid_contract(78)
    244         paid.set_sequence(last_value=200, is_called=is_called)
    245 
    246         db.apply_migration(36)
    247 
    248         # Both new sequences inherit the shared sequence's higher position.
    249         self.assertEqual(OrderFixture(db, 1).add_order(), expected_next)
    250         self.assertEqual(OrderFixture(db, 2).add_order(), expected_next)
    251 
    252     def test_0036_preserves_called_sequence(self):
    253         self.check_0036_preserves_sequence(is_called=True, expected_next=201)
    254 
    255     def test_0036_preserves_uncalled_sequence(self):
    256         self.check_0036_preserves_sequence(is_called=False, expected_next=200)
    257 
    258     def test_0036_rejects_exhausted_sequence(self):
    259         db = self.database_before(36)
    260         orders = db.add_instance(1, legacy=True)
    261         orders.set_sequence(last_value=MAX_SERIAL, is_called=True)
    262 
    263         with self.assertRaisesRegex(RuntimeError, "Order serial sequence exhausted"):
    264             db.apply_migration(36)
    265 
    266         self.assertEqual(db.sql("SELECT count(*) FROM _v.patches "
    267                                 "WHERE patch_name='merchant-0036'"), "0")
    268 
    269     def test_0047_repairs_restarted_ids_without_changing_orders(self):
    270         db = self.database_before(47)
    271         paid = db.add_instance(1)
    272         paid.add_paid_contract(78)
    273         # Reproduce the reported history: new IDs 1-4 follow historical ID 78.
    274         for serial in range(1, 5):
    275             paid.add_order(serial)
    276             paid.add_paid_contract(serial)
    277         paid.set_sequence(last_value=4, is_called=True)
    278         unpaid = db.add_instance(2)
    279         unpaid.add_paid_contract(80)
    280         unpaid.add_order(90)
    281         paid_before, unpaid_before = paid.snapshot(), unpaid.snapshot()
    282 
    283         db.apply_migration(47)
    284 
    285         self.assert_sequence(paid, last_value=79, is_called=False)
    286         self.assert_sequence(unpaid, last_value=91, is_called=False)
    287         self.assertEqual(paid.snapshot(), paid_before)
    288         self.assertEqual(unpaid.snapshot(), unpaid_before)
    289 
    290         # Reapplying the registered fixup must not consume or rewind IDs.
    291         db.sql("CALL merchant.fixup_instance_schema(47::INT8)")
    292         self.assert_sequence(paid, last_value=79, is_called=False)
    293         self.assert_sequence(unpaid, last_value=91, is_called=False)
    294         self.assertEqual(paid.add_order(), 79)
    295         self.assertEqual(unpaid.add_order(), 91)
    296 
    297         # At the database level, claiming can copy the new serial without a
    298         # primary-key collision. This is not an HTTP/backend claim test.
    299         db.sql(f"""
    300             INSERT INTO {paid.schema}.merchant_contract_terms
    301               (order_serial, order_id, contract_terms, h_contract_terms,
    302                creation_time, pay_deadline, refund_deadline, claim_token)
    303             SELECT order_serial, order_id, contract_terms,
    304                    decode(repeat('ff',64),'hex'), creation_time,
    305                    pay_deadline, pay_deadline, claim_token
    306               FROM {paid.schema}.merchant_orders WHERE order_serial=79
    307         """)
    308         self.assertEqual(db.sql(
    309             f"SELECT order_serial FROM {paid.schema}.merchant_contract_terms "
    310             "WHERE order_id='new-order'"), "79")
    311 
    312     def test_0047_preserves_safe_sequence_states(self):
    313         db = self.database_before(47)
    314         called = db.add_instance(1)
    315         called.set_sequence(last_value=200, is_called=True)
    316         uncalled = db.add_instance(2)
    317         uncalled.set_sequence(last_value=200, is_called=False)
    318         just_above_history = db.add_instance(3)
    319         just_above_history.add_paid_contract(78)
    320         just_above_history.set_sequence(last_value=79, is_called=False)
    321 
    322         db.apply_migration(47)
    323 
    324         self.assert_sequence(called, last_value=200, is_called=True)
    325         self.assert_sequence(uncalled, last_value=200, is_called=False)
    326         self.assert_sequence(just_above_history, last_value=79, is_called=False)
    327 
    328     def test_0047_keeps_empty_and_new_instances_starting_at_one(self):
    329         db = self.database_before(47)
    330         empty = db.add_instance(1)
    331 
    332         db.apply_migration(47)
    333         new = db.add_instance(2)
    334 
    335         self.assert_sequence(empty, last_value=1, is_called=False)
    336         self.assertEqual(empty.add_order(), 1)
    337         self.assertEqual(new.add_order(), 1)
    338 
    339     def test_0048_adds_dd97_and_dd98_to_existing_and_new_instances(self):
    340         db = self.database_before(47)
    341         existing = db.add_instance(1)
    342         db.apply_migration(47)
    343         db.apply_migration(48)
    344         new = db.add_instance(2)
    345         for instance in (existing, new):
    346             schema = instance.schema
    347             self.assertEqual(db.sql(f"""
    348                 SELECT count(*) FROM information_schema.columns
    349                  WHERE table_schema='{schema}'
    350                    AND ((table_name='merchant_otp_devices'
    351                          AND column_name='otp_device_pub')
    352                         OR (table_name IN ('merchant_orders', 'merchant_contract_terms')
    353                             AND column_name='pos_challenge'))
    354             """), "3")
    355             self.assertEqual(db.sql(f"""
    356                 SELECT count(*) FROM information_schema.tables
    357                  WHERE table_schema='{schema}' AND table_name IN
    358                    ('merchant_fountains', 'merchant_fountain_grants',
    359                     'merchant_fountain_withdrawals', 'merchant_fountain_withdraw_sigs')
    360             """), "4")
    361             self.assertEqual(db.sql(f"""
    362                 SELECT is_nullable FROM information_schema.columns
    363                  WHERE table_schema='{schema}' AND table_name='merchant_issued_tokens'
    364                    AND column_name='h_contract_terms'
    365             """), "YES")
    366 
    367     def test_0047_sequence_restart_rolls_back(self):
    368         db = self.database_before(47)
    369         orders = db.add_instance(1)
    370         orders.add_paid_contract(78)
    371         db.apply_migration(47)
    372         orders.set_sequence(last_value=4, is_called=True)
    373 
    374         db.sql(f"BEGIN; {orders.repair_statement()} ROLLBACK;")
    375 
    376         self.assert_sequence(orders, last_value=4, is_called=True)
    377 
    378     def test_0047_later_failure_rolls_back_earlier_repair(self):
    379         db = self.database_before(47)
    380         repairable = db.add_instance(1)
    381         exhausted = db.add_instance(2)
    382         repairable.add_paid_contract(78)
    383         db.apply_migration(47)
    384         repairable.set_sequence(last_value=4, is_called=True)
    385         exhausted.set_sequence(last_value=MAX_SERIAL, is_called=True)
    386 
    387         with self.assertRaisesRegex(RuntimeError, "Order serial sequence exhausted"):
    388             db.sql(f"BEGIN; {repairable.repair_statement()} "
    389                    f"{exhausted.repair_statement()} COMMIT;")
    390 
    391         self.assert_sequence(repairable, last_value=4, is_called=True)
    392 
    393     def test_0047_rejects_exhausted_stored_ids(self):
    394         db = self.database_before(47)
    395         orders = db.add_instance(1)
    396         orders.add_order(MAX_SERIAL)
    397 
    398         with self.assertRaisesRegex(RuntimeError, "Order serial sequence exhausted"):
    399             db.apply_migration(47)
    400 
    401         self.assert_sequence(orders, last_value=1, is_called=False)
    402 
    403 
    404 def main():
    405     source, build = (Path(arg).resolve() for arg in sys.argv[1:])
    406     # Missing CI prerequisites must not silently disable migration coverage.
    407     unavailable_status = 1 if os.geteuid() == 0 else 77
    408     if not shutil.which("pg_config"):
    409         print("PostgreSQL server tools unavailable")
    410         return unavailable_status
    411     bindir = Path(run(["pg_config", "--bindir"]))
    412     if not all((bindir / tool).exists() for tool in ("initdb", "pg_ctl", "psql")):
    413         print("PostgreSQL server tools unavailable")
    414         return unavailable_status
    415 
    416     with postgres_cluster(bindir) as env:
    417         admin = Database("template1", bindir, env, build)
    418         admin.sql("CREATE DATABASE before_36")
    419         before_36 = Database("before_36", bindir, env, build)
    420         before_36.apply_file(source / "versioning.sql")
    421         for version in range(1, 36):
    422             before_36.apply_migration(version)
    423         admin.sql("CREATE DATABASE before_47 TEMPLATE before_36")
    424         before_47 = Database("before_47", bindir, env, build)
    425         for version in range(36, 47):
    426             before_47.apply_migration(version)
    427 
    428         OrderSequenceMigrations.admin = admin
    429         OrderSequenceMigrations.bindir = bindir
    430         OrderSequenceMigrations.cluster_env = env
    431         OrderSequenceMigrations.sql_dir = build
    432         suite = unittest.defaultTestLoader.loadTestsFromTestCase(OrderSequenceMigrations)
    433         result = unittest.TextTestRunner(verbosity=2).run(suite)
    434         return 0 if result.wasSuccessful() else 1
    435 
    436 
    437 if __name__ == "__main__":
    438     sys.exit(main())