| 1 | #!/usr/bin/env python3 |
| 2 | import argparse |
| 3 | import hashlib |
| 4 | import json |
| 5 | import os |
| 6 | from pathlib import Path |
| 7 | import re |
| 8 | import subprocess |
| 9 | import tempfile |
| 10 | import urllib.error |
| 11 | import urllib.request |
| 12 | |
| 13 | |
| 14 | def identifier(value): |
| 15 | return '"' + value.replace('"', '""') + '"' |
| 16 | |
| 17 | |
| 18 | def digest(path): |
| 19 | result = hashlib.sha256() |
| 20 | with path.open('rb') as stream: |
| 21 | for block in iter(lambda: stream.read(1024 * 1024), b''): |
| 22 | result.update(block) |
| 23 | return result.hexdigest() |
| 24 | |
| 25 | |
| 26 | def fingerprint(sql, database): |
| 27 | relations = json.loads(sql(database, """ |
| 28 | SELECT coalesce(json_agg(x ORDER BY n, c), '[]') FROM ( |
| 29 | SELECT n.nspname n, c.relname c, c.relkind k |
| 30 | FROM pg_class c JOIN pg_namespace n ON n.oid=c.relnamespace |
| 31 | WHERE n.nspname NOT LIKE 'pg_%' AND n.nspname <> 'information_schema' |
| 32 | AND c.relkind IN ('r','m','S') |
| 33 | ) x |
| 34 | """)) |
| 35 | result = {} |
| 36 | for relation in relations: |
| 37 | table = identifier(relation['n']) + '.' + identifier(relation['c']) |
| 38 | query = (f'SELECT last_value,is_called FROM {table}' if relation['k'] == 'S' else |
| 39 | f'COPY (SELECT row_to_json(t)::text FROM {table} t ' |
| 40 | 'ORDER BY row_to_json(t)::text COLLATE "C") TO STDOUT') |
| 41 | content = sql(database, query) |
| 42 | result[table] = {'rows': len(content.splitlines()), |
| 43 | 'sha256': hashlib.sha256(content).hexdigest()} |
| 44 | content = sql(database, "COPY (SELECT loid,pageno,encode(data,'hex') " |
| 45 | 'FROM pg_largeobject ORDER BY loid,pageno) TO STDOUT') |
| 46 | result['large_objects'] = {'rows': len(content.splitlines()), |
| 47 | 'sha256': hashlib.sha256(content).hexdigest()} |
| 48 | return result |
| 49 | |
| 50 | |
| 51 | def main(): |
| 52 | parser = argparse.ArgumentParser(description='Restore a personal database before starting its service.') |
| 53 | parser.add_argument('service', choices=('keycloak', 'dawarich')) |
| 54 | parser.add_argument('handoff', type=Path) |
| 55 | parser.add_argument('--container', help='Restore inside the isolated initializer before deploying Postgres') |
| 56 | args = parser.parse_args() |
| 57 | if os.geteuid() != 0: |
| 58 | parser.error('Run as root on Zenith') |
| 59 | os.umask(0o077) |
| 60 | handoff = args.handoff.resolve() |
| 61 | cluster = json.loads((handoff / 'manifest.json').read_text())['clusters']['18'] |
| 62 | if not cluster.get('restoreVerified'): |
| 63 | raise ValueError('PostgreSQL 18 backup is unverified. Use the verified handoff') |
| 64 | entry = cluster['databases'][args.service] |
| 65 | dump = (handoff / entry['file']).resolve() |
| 66 | if not dump.is_relative_to(handoff / 'pg18') or dump.suffix != '.dump': |
| 67 | raise ValueError('Dump is outside the PostgreSQL 18 handoff') |
| 68 | if dump.stat().st_size != entry['bytes'] or digest(dump) != entry['sha256']: |
| 69 | raise ValueError('Backup checksum differs. Recover the verified handoff before importing') |
| 70 | |
| 71 | token = os.environ.get('NOMAD_TOKEN') or Path('/var/lib/studio/nomad.token').read_text().strip() |
| 72 | |
| 73 | def nomad(path): |
| 74 | request = urllib.request.Request('http://127.0.0.1:4646/v1/' + path, |
| 75 | headers={'X-Nomad-Token': token}) |
| 76 | with urllib.request.urlopen(request, timeout=10) as response: |
| 77 | return json.load(response) |
| 78 | |
| 79 | def stopped(service): |
| 80 | try: |
| 81 | job = nomad('job/' + service) |
| 82 | except urllib.error.HTTPError as error: |
| 83 | if error.code == 404: |
| 84 | return |
| 85 | raise |
| 86 | if not job.get('Stop') or any( |
| 87 | item['ClientStatus'] in ('pending', 'running') |
| 88 | for item in nomad('job/' + service + '/allocations') |
| 89 | ): |
| 90 | raise ValueError(f'Stop {service} and wait for its allocations before importing') |
| 91 | |
| 92 | stopped(args.service) |
| 93 | connection = nomad('var/nomad/jobs/' + args.service + '/inputs/database')['Items'] |
| 94 | database, owner = connection['name'], connection['username'] |
| 95 | expected_database = 'keycloak_next' if args.service == 'keycloak' else 'dawarich' |
| 96 | if database != expected_database or owner != 'svc_' + database: |
| 97 | raise ValueError('Database variable differs from the production service. Allocate its database first') |
| 98 | if args.container: |
| 99 | stopped('postgres') |
| 100 | else: |
| 101 | allocations = [item['ID'] for item in nomad('job/postgres/allocations') |
| 102 | if item['ClientStatus'] == 'running' and item['DesiredStatus'] == 'run'] |
| 103 | if len(allocations) != 1: |
| 104 | raise ValueError('Postgres needs exactly one running allocation') |
| 105 | folder = Path(tempfile.mkdtemp(prefix=args.service + '-import-', dir='/var/lib/studio')) |
| 106 | diagnostics = folder / 'diagnostics.log' |
| 107 | podman = ['podman'] |
| 108 | |
| 109 | def run(command, *, source=None, output=None): |
| 110 | with diagnostics.open('ab') as errors: |
| 111 | result = subprocess.run(command, input=source if isinstance(source, bytes) else None, |
| 112 | stdin=source if source is not None and not isinstance(source, bytes) else None, |
| 113 | stdout=output or subprocess.PIPE, stderr=errors) |
| 114 | if result.returncode: |
| 115 | raise RuntimeError(f'{Path(command[0]).name} failed. See {diagnostics}') |
| 116 | return result.stdout |
| 117 | |
| 118 | if args.container: |
| 119 | container = json.loads(run(podman + ['inspect', args.container]))[0] |
| 120 | if not container['State']['Running'] or container['HostConfig']['NetworkMode'] != 'none': |
| 121 | raise ValueError('Initializer must be running with --network none') |
| 122 | container_id = container['Id'] |
| 123 | else: |
| 124 | containers = [parts[0] for line in run(podman + ['ps', '--format', '{{.ID}} {{.Names}}']).decode().splitlines() |
| 125 | if len(parts := line.split()) == 2 and parts[1].endswith(allocations[0])] |
| 126 | if len(containers) != 1: |
| 127 | raise ValueError('Postgres allocation container is unavailable') |
| 128 | container_id = containers[0] |
| 129 | execute = podman + ['exec', '-i', container_id] |
| 130 | |
| 131 | def sql(db, query): |
| 132 | return run(execute + ['psql', '-X', '-U', 'postgres', '-d', db, |
| 133 | '-At', '-v', 'ON_ERROR_STOP=1'], source=(query + '\n').encode()) |
| 134 | |
| 135 | current_owner = sql('postgres', f"SELECT pg_get_userbyid(datdba) FROM pg_database WHERE datname='{database}'").decode().strip() |
| 136 | if current_owner != owner: |
| 137 | raise ValueError('Production database is missing or has another owner. Allocate it with the Postgres provider') |
| 138 | with dump.open('rb') as source: |
| 139 | listing = run(execute + ['pg_restore', '--list'], source=source).decode() |
| 140 | extensions = re.findall(r'^\d+; \d+ \d+ EXTENSION - (\S+) ', listing, re.MULTILINE) |
| 141 | if set(extensions) - {'plpgsql', 'postgis'} or ('postgis' in extensions) != (args.service == 'dawarich'): |
| 142 | raise ValueError('Backup extensions differ from the service. Review the personal database before importing') |
| 143 | extension_data, application = [], [] |
| 144 | for line in listing.splitlines(): |
| 145 | if re.match(r'^\d+; \d+ \d+ (EXTENSION |COMMENT - EXTENSION |SCHEMA - public )', line): |
| 146 | continue |
| 147 | if re.match(r'^\d+; \d+ \d+ TABLE DATA public spatial_ref_sys ', line): |
| 148 | extension_data.append(line) |
| 149 | else: |
| 150 | application.append(line) |
| 151 | contents = {name: value for name, value in entry['contents'].items() if name != 'large_object_owners'} |
| 152 | restore_list = '/tmp/' + folder.name + '.list' |
| 153 | run(execute + ['tee', restore_list], source=('\n'.join(application) + '\n').encode()) |
| 154 | try: |
| 155 | stopped(args.service) |
| 156 | if args.container: |
| 157 | stopped('postgres') |
| 158 | if sql('postgres', f"SELECT count(*) FROM pg_stat_activity WHERE datname='{database}'").strip() != b'0': |
| 159 | raise ValueError('Database has connected clients. Stop them before importing') |
| 160 | backup = folder / 'before.dump' |
| 161 | with backup.open('wb') as output: |
| 162 | run(execute + ['pg_dump', '-U', 'postgres', '-Fc', database], output=output) |
| 163 | if not backup.stat().st_size: |
| 164 | raise RuntimeError('Target backup is empty. Import stopped') |
| 165 | sql('postgres', f'DROP DATABASE {database};') |
| 166 | sql('postgres', f'CREATE DATABASE {database} OWNER {owner};') |
| 167 | sql(database, f'ALTER SCHEMA public OWNER TO {owner};') |
| 168 | if args.service == 'dawarich': |
| 169 | sql(database, 'CREATE EXTENSION postgis;') |
| 170 | with dump.open('rb') as source: |
| 171 | run(execute + ['pg_restore', '-U', 'postgres', '-d', database, '--no-owner', '--no-acl', |
| 172 | '--role=' + owner, '--single-transaction', '--exit-on-error', |
| 173 | '--use-list=' + restore_list], source=source) |
| 174 | if extension_data: |
| 175 | run(execute + ['tee', restore_list], source=('\n'.join(extension_data) + '\n').encode()) |
| 176 | with dump.open('rb') as source: |
| 177 | extension_sql = run(execute + ['pg_restore', '-U', 'postgres', '--data-only', '--no-owner', |
| 178 | '--no-acl', '--use-list=' + restore_list, '--file=-'], source=source) |
| 179 | srids = [] |
| 180 | copying = False |
| 181 | for line in extension_sql.splitlines(): |
| 182 | if line.startswith(b'COPY public.spatial_ref_sys ('): |
| 183 | copying = True |
| 184 | elif copying and line == b'\\.': |
| 185 | copying = False |
| 186 | elif copying: |
| 187 | srids.append(int(line.split(b'\t', 1)[0])) |
| 188 | # PostGIS dumps include custom SRIDs; the extension supplies built-in definitions. |
| 189 | if srids: |
| 190 | sql(database, 'DELETE FROM public.spatial_ref_sys WHERE srid IN (' + ','.join(map(str, srids)) + ');') |
| 191 | with dump.open('rb') as source: |
| 192 | run(execute + ['pg_restore', '-U', 'postgres', '-d', database, '--no-owner', '--no-acl', |
| 193 | '--data-only', '--single-transaction', '--exit-on-error', |
| 194 | '--use-list=' + restore_list], source=source) |
| 195 | actual = fingerprint(sql, database) |
| 196 | if actual != contents: |
| 197 | changed = sorted(name for name in actual.keys() | contents.keys() if actual.get(name) != contents.get(name)) |
| 198 | (folder / 'differences.json').write_text(json.dumps(changed, indent=2) + '\n') |
| 199 | raise RuntimeError(f'Restored data differs. Keep {args.service} stopped and review {folder}') |
| 200 | if dump.stat().st_size != entry['bytes'] or digest(dump) != entry['sha256']: |
| 201 | raise RuntimeError('Backup changed during import. Keep the service stopped') |
| 202 | (folder / 'verified.json').write_text(json.dumps({ |
| 203 | 'service': args.service, 'database': database, 'source': str(dump), |
| 204 | 'sha256': entry['sha256'], 'contents': actual, |
| 205 | }, indent=2) + '\n') |
| 206 | print(f'Imported and verified {args.service}. Previous database: {backup}') |
| 207 | finally: |
| 208 | run(execute + ['rm', '-f', restore_list]) |
| 209 | |
| 210 | |
| 211 | if __name__ == '__main__': |
| 212 | try: |
| 213 | main() |
| 214 | except (ValueError, RuntimeError, OSError, KeyError, urllib.error.HTTPError) as error: |
| 215 | raise SystemExit(str(error)) from None |