dbreduce 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
dbreduce/__init__.py ADDED
@@ -0,0 +1 @@
1
+ """PostgreSQL reproducer reduction."""
dbreduce/__main__.py ADDED
@@ -0,0 +1,3 @@
1
+ from dbreduce.cli import app
2
+
3
+ app()
File without changes
@@ -0,0 +1,32 @@
1
+ import hashlib
2
+ from pathlib import Path
3
+
4
+
5
+ def fingerprint(snapshot: Path) -> str:
6
+ """Hash schema/sequence statements and COPY rows as a multiset per table."""
7
+ digest = hashlib.sha256()
8
+ with snapshot.open("rb") as stream:
9
+ for line in stream:
10
+ digest.update(line)
11
+ if line.startswith(b"COPY ") and line.rstrip(b"\r\n").endswith(b" FROM stdin;"):
12
+ rows = []
13
+ for line in stream:
14
+ if line.rstrip(b"\r\n") == b"\\.":
15
+ break
16
+ rows.append(line)
17
+ for row in sorted(rows):
18
+ digest.update(row)
19
+ digest.update(line)
20
+ return digest.hexdigest()
21
+
22
+
23
+ class Cache:
24
+ def __init__(self) -> None:
25
+ self.results: dict[str, bool] = {}
26
+ self.hits = 0
27
+
28
+ def get(self, key: str) -> bool | None:
29
+ result = self.results.get(key)
30
+ if result is not None:
31
+ self.hits += 1
32
+ return result
@@ -0,0 +1,193 @@
1
+ import json
2
+ import os
3
+ import shutil
4
+ import tempfile
5
+ import time
6
+ from pathlib import Path
7
+ from typing import Annotated
8
+
9
+ import psycopg
10
+ import typer
11
+
12
+ from dbreduce.cache.store import Cache
13
+ from dbreduce.graph.dependencies import components, dependencies
14
+ from dbreduce.oracle.runner import Oracle
15
+ from dbreduce.postgres.backend import PostgresBackend
16
+ from dbreduce.postgres.database import Workspace, read_archive_settings, read_settings
17
+ from dbreduce.postgres.dump import dump
18
+ from dbreduce.postgres.introspection import check_extension_tables, inspect_database
19
+ from dbreduce.reducer.engine import reduce as reduce_state
20
+
21
+ app = typer.Typer(
22
+ no_args_is_help=True, help="Minimize PostgreSQL bug reproducers in a disposable copy."
23
+ )
24
+
25
+
26
+ def publish_result(exported: Path, output: Path, report: Path, result: dict[str, object]) -> None:
27
+ """Prepare both files, then publish without overwriting existing paths."""
28
+ pending: list[Path] = []
29
+ created: list[Path] = []
30
+ try:
31
+ with tempfile.NamedTemporaryFile(
32
+ dir=output.parent, prefix=".dbreduce-", delete=False
33
+ ) as file:
34
+ pending.append(Path(file.name))
35
+ with exported.open("rb") as source:
36
+ shutil.copyfileobj(source, file)
37
+ with tempfile.NamedTemporaryFile(
38
+ mode="w",
39
+ encoding="utf-8",
40
+ dir=report.parent,
41
+ prefix=".dbreduce-",
42
+ delete=False,
43
+ ) as file:
44
+ pending.append(Path(file.name))
45
+ json.dump(result, file, indent=2)
46
+ file.write("\n")
47
+ for prepared_path, destination in zip(pending, (output, report), strict=True):
48
+ os.link(prepared_path, destination)
49
+ created.append(destination)
50
+ except BaseException:
51
+ for path in created:
52
+ path.unlink(missing_ok=True)
53
+ raise
54
+ finally:
55
+ for path in pending:
56
+ path.unlink(missing_ok=True)
57
+
58
+
59
+ def require_clients() -> None:
60
+ for name in ("pg_dump", "pg_restore"):
61
+ if not shutil.which(name):
62
+ raise ValueError(f"Required PostgreSQL client tool is missing: {name}")
63
+
64
+
65
+ @app.command("inspect")
66
+ def inspect_command(database: Annotated[str, typer.Option(help="PostgreSQL source DSN")]) -> None:
67
+ """Show tables, exact row counts, keys and dependency graph (read-only)."""
68
+ try:
69
+ with psycopg.connect(database, connect_timeout=10) as conn:
70
+ conn.execute("SET TRANSACTION ISOLATION LEVEL REPEATABLE READ READ ONLY")
71
+ schema = inspect_database(conn)
72
+ for table in schema.tables:
73
+ typer.echo(
74
+ f"{table.label}\n rows: {table.rows}\n pk: {', '.join(table.primary_key) or '-'}"
75
+ )
76
+ for fk in schema.foreign_keys:
77
+ if fk.child == table.key:
78
+ typer.echo(
79
+ f" FK {', '.join(fk.columns)} -> {'.'.join(fk.parent)}"
80
+ f"({', '.join(fk.target_columns)})"
81
+ )
82
+ typer.echo("Dependencies (child -> parent):")
83
+ for child, parents in dependencies(schema).items():
84
+ typer.echo(
85
+ f" {'.'.join(child)} -> {', '.join('.'.join(p) for p in sorted(parents)) or '-'}"
86
+ )
87
+ typer.echo(f"Strongly connected components: {components(schema)}")
88
+ except (ValueError, psycopg.Error) as error:
89
+ fail(error)
90
+
91
+
92
+ def fail(error: Exception) -> None:
93
+ message = (
94
+ "PostgreSQL operation failed; check connection and permissions"
95
+ if isinstance(error, psycopg.Error)
96
+ else str(error)
97
+ )
98
+ typer.echo(f"Error: {message}", err=True)
99
+ raise typer.Exit(1)
100
+
101
+
102
+ @app.command("reduce")
103
+ def reduce_command(
104
+ oracle: Annotated[str, typer.Option(help="Shell command; nonzero exit reproduces the bug")],
105
+ database: Annotated[str | None, typer.Option(help="Source PostgreSQL DSN")] = None,
106
+ input_dump: Annotated[
107
+ Path | None, typer.Option("--dump", help="Trusted pg_dump custom-format archive")
108
+ ] = None,
109
+ admin_database: Annotated[str | None, typer.Option(help="DSN with CREATEDB privilege")] = None,
110
+ confirm: Annotated[int, typer.Option(min=1)] = 1,
111
+ timeout: Annotated[float, typer.Option(min=0.01)] = 60,
112
+ output: Annotated[Path, typer.Option()] = Path("dbreduce.min.sql"),
113
+ report: Annotated[Path, typer.Option()] = Path("dbreduce-report.json"),
114
+ ) -> None:
115
+ """Reduce a source database or dump; export SQL and a JSON report."""
116
+ started = time.monotonic()
117
+ try:
118
+ if (database is None) == (input_dump is None) or database == "":
119
+ raise ValueError("Provide exactly one of --database or --dump")
120
+ admin = admin_database or database
121
+ if not admin:
122
+ raise ValueError("--admin-database is required with --dump")
123
+ if output.resolve() == report.resolve() or output.exists() or report.exists():
124
+ raise ValueError("Output and report must be distinct paths that do not already exist")
125
+ require_clients()
126
+ with tempfile.TemporaryDirectory(prefix="dbreduce-") as temporary:
127
+ snapshot = Path(temporary) / "accepted.dump"
128
+ if database:
129
+ with psycopg.connect(database, connect_timeout=10) as conn:
130
+ conn.execute("SET TRANSACTION READ ONLY")
131
+ check_extension_tables(conn)
132
+ dump(database, snapshot)
133
+ else:
134
+ assert input_dump is not None
135
+ shutil.copyfile(input_dump, snapshot)
136
+ settings = read_settings(database) if database else read_archive_settings(snapshot)
137
+ with Workspace(admin, settings=settings) as workspace:
138
+ workspace.reset(snapshot)
139
+ with psycopg.connect(workspace.dsn, connect_timeout=10) as conn:
140
+ check_extension_tables(conn)
141
+ schema = inspect_database(conn)
142
+ # Keep one normalized snapshot for all probes.
143
+ dump(workspace.dsn, snapshot)
144
+ initial_rows = sum(table.rows for table in schema.tables)
145
+ typer.echo(f"Initial database: {len(schema.tables)} tables, {initial_rows} rows")
146
+ runner = Oracle(oracle, confirm=confirm, timeout=timeout)
147
+ cache = Cache()
148
+ if not runner.fails(workspace.url, lambda: workspace.reset(snapshot)):
149
+ raise ValueError(
150
+ "Oracle passed on the initial copy; failure does not reproduce"
151
+ )
152
+ typer.echo("Oracle: FAIL\nReducing tables, row groups and individual rows...")
153
+ workspace.reset(snapshot)
154
+ backend = PostgresBackend(workspace, schema, snapshot, runner, cache)
155
+ final = reduce_state(backend, typer.echo)
156
+ # A fresh uncached final confirmation catches some flaky-oracle failures.
157
+ if not runner.fails(workspace.url, lambda: workspace.reset(snapshot)):
158
+ raise ValueError(
159
+ "Final oracle confirmation passed; no verified result exported"
160
+ )
161
+ workspace.reset(snapshot)
162
+ exported = Path(temporary) / "result.sql"
163
+ dump(workspace.dsn, exported, archive=False, create_database=True)
164
+ final_rows = sum(map(len, final.values()))
165
+ result = {
166
+ "initial_tables": len(schema.tables),
167
+ "initial_rows": initial_rows,
168
+ "final_tables": len(final),
169
+ "final_rows": final_rows,
170
+ "oracle_executions": runner.executions,
171
+ "cache_hits": cache.hits,
172
+ "duration_seconds": time.monotonic() - started,
173
+ "rows_by_table": {
174
+ schema_name: {
175
+ name: len(final[(schema_name, name)])
176
+ for schema, name in final
177
+ if schema == schema_name
178
+ }
179
+ for schema_name, _ in final
180
+ },
181
+ "constraint_rejections": backend.constraint_rejections,
182
+ "raise_exception_rejections": backend.raise_rejections,
183
+ "confirm": confirm,
184
+ "oracle": "FAIL",
185
+ "restore_database": workspace.name,
186
+ }
187
+ publish_result(exported, output, report, result)
188
+ typer.echo(
189
+ f"Final: {final_rows} rows, oracle: FAIL, "
190
+ f"executions: {runner.executions}, cache hits: {cache.hits}"
191
+ )
192
+ except (OSError, ValueError, RuntimeError, psycopg.Error) as error:
193
+ fail(error)
File without changes
@@ -0,0 +1,33 @@
1
+ from dbreduce.models.schema import Schema, TableKey
2
+
3
+
4
+ def dependencies(schema: Schema) -> dict[TableKey, set[TableKey]]:
5
+ graph: dict[TableKey, set[TableKey]] = {table.key: set() for table in schema.tables}
6
+ for fk in schema.foreign_keys:
7
+ graph[fk.child].add(fk.parent)
8
+ return graph
9
+
10
+
11
+ def components(schema: Schema) -> list[tuple[TableKey, ...]]:
12
+ """Iterative reachability SCCs; small table graphs favor a simple implementation."""
13
+ graph = dependencies(schema)
14
+ reach = {}
15
+ for node in graph:
16
+ seen: set[TableKey] = set()
17
+ pending = [node]
18
+ while pending:
19
+ current = pending.pop()
20
+ if current not in seen:
21
+ seen.add(current)
22
+ pending.extend(graph[current] - seen)
23
+ reach[node] = seen
24
+ remaining = set(graph)
25
+ result = []
26
+ while remaining:
27
+ node = min(remaining)
28
+ group = tuple(
29
+ sorted(other for other in remaining if other in reach[node] and node in reach[other])
30
+ )
31
+ result.append(group)
32
+ remaining.difference_update(group)
33
+ return result
File without changes
@@ -0,0 +1,33 @@
1
+ from dataclasses import dataclass
2
+
3
+ type TableKey = tuple[str, str]
4
+ type RowKey = tuple[str, int]
5
+ type State = dict[TableKey, list[RowKey]]
6
+
7
+
8
+ @dataclass(frozen=True)
9
+ class Table:
10
+ key: TableKey
11
+ primary_key: tuple[str, ...]
12
+ rows: int
13
+
14
+ @property
15
+ def label(self) -> str:
16
+ return ".".join(self.key)
17
+
18
+
19
+ @dataclass(frozen=True)
20
+ class ForeignKey:
21
+ name: str
22
+ child: TableKey
23
+ parent: TableKey
24
+ columns: tuple[str, ...]
25
+ target_columns: tuple[str, ...]
26
+ operators: tuple[tuple[str, str], ...]
27
+ collations: tuple[tuple[str, str] | None, ...] = ()
28
+
29
+
30
+ @dataclass(frozen=True)
31
+ class Schema:
32
+ tables: tuple[Table, ...]
33
+ foreign_keys: tuple[ForeignKey, ...]
File without changes
@@ -0,0 +1,52 @@
1
+ import os
2
+ import signal
3
+ import subprocess
4
+ from collections.abc import Callable
5
+
6
+
7
+ class OracleError(RuntimeError):
8
+ pass
9
+
10
+
11
+ class Oracle:
12
+ def __init__(self, command: str, *, confirm: int = 1, timeout: float = 60) -> None:
13
+ if confirm < 1 or timeout <= 0:
14
+ raise ValueError("confirm and timeout must be positive")
15
+ self.command = command
16
+ self.confirm = confirm
17
+ self.timeout = timeout
18
+ self.executions = 0
19
+
20
+ def fails(self, url: str, prepare: Callable[[], None]) -> bool:
21
+ env = os.environ.copy()
22
+ env.update(DATABASE_URL=url, DBREDUCE_DATABASE_URL=url)
23
+ for _ in range(self.confirm):
24
+ prepare()
25
+ self.executions += 1
26
+ with subprocess.Popen(
27
+ self.command,
28
+ shell=True,
29
+ env=env,
30
+ start_new_session=True,
31
+ stdout=subprocess.DEVNULL,
32
+ stderr=subprocess.DEVNULL,
33
+ ) as process:
34
+ try:
35
+ code = process.wait(timeout=self.timeout)
36
+ except BaseException as error:
37
+ os.killpg(process.pid, signal.SIGKILL)
38
+ process.wait()
39
+ if isinstance(error, subprocess.TimeoutExpired):
40
+ raise OracleError(
41
+ "Oracle timed out; this is not a reproduced failure"
42
+ ) from error
43
+ raise
44
+ finally:
45
+ # Do not leave background children connected to the disposable database.
46
+ try:
47
+ os.killpg(process.pid, signal.SIGKILL)
48
+ except ProcessLookupError:
49
+ pass
50
+ if code == 0:
51
+ return False
52
+ return True
File without changes
@@ -0,0 +1,64 @@
1
+ import uuid
2
+ from pathlib import Path
3
+
4
+ import psycopg
5
+
6
+ from dbreduce.cache.store import Cache, fingerprint
7
+ from dbreduce.models.schema import RowKey, Schema, State, TableKey
8
+ from dbreduce.oracle.runner import Oracle
9
+ from dbreduce.postgres.database import Workspace
10
+ from dbreduce.postgres.dump import dump
11
+ from dbreduce.postgres.rows import CandidateRejected, delete_rows, read_rows
12
+
13
+
14
+ class PostgresBackend:
15
+ def __init__(
16
+ self, workspace: Workspace, schema: Schema, snapshot: Path, oracle: Oracle, cache: Cache
17
+ ) -> None:
18
+ self.workspace = workspace
19
+ self.schema = schema
20
+ self.snapshot = snapshot
21
+ self.oracle = oracle
22
+ self.cache = cache
23
+ self.constraint_rejections = 0
24
+ self.raise_rejections = 0
25
+ self.restrict_key = uuid.uuid4().hex
26
+ with psycopg.connect(workspace.dsn, connect_timeout=10) as conn:
27
+ self.current, _ = read_rows(conn, schema)
28
+
29
+ def state(self) -> State:
30
+ return self.current
31
+
32
+ def attempt(self, table: TableKey, rows: list[RowKey]) -> bool:
33
+ self.workspace.reset(self.snapshot)
34
+ try:
35
+ with psycopg.connect(self.workspace.dsn, connect_timeout=10) as conn:
36
+ conn.execute("SET statement_timeout = '600s'")
37
+ delete_rows(conn, self.schema, table, rows)
38
+ except psycopg.errors.IntegrityConstraintViolation:
39
+ self.constraint_rejections += 1
40
+ return False
41
+ except CandidateRejected:
42
+ self.raise_rejections += 1
43
+ return False
44
+ candidate_dump = self.snapshot.with_name("candidate.dump")
45
+ dump(self.workspace.dsn, candidate_dump)
46
+ # Compare and fingerprint the state the oracle will actually see after restore.
47
+ self.workspace.reset(candidate_dump)
48
+ with psycopg.connect(self.workspace.dsn, connect_timeout=10) as conn:
49
+ candidate, _ = read_rows(conn, self.schema)
50
+ if sum(map(len, candidate.values())) >= sum(map(len, self.current.values())):
51
+ return False
52
+ fingerprint_dump = self.snapshot.with_name("fingerprint.sql")
53
+ dump(self.workspace.dsn, fingerprint_dump, archive=False, restrict_key=self.restrict_key)
54
+ key = fingerprint(fingerprint_dump)
55
+ accepted = self.cache.get(key)
56
+ if accepted is None:
57
+ accepted = self.oracle.fails(
58
+ self.workspace.url, lambda: self.workspace.reset(candidate_dump)
59
+ )
60
+ self.cache.results[key] = accepted
61
+ if accepted:
62
+ candidate_dump.replace(self.snapshot)
63
+ self.current = candidate
64
+ return accepted
@@ -0,0 +1,238 @@
1
+ import hashlib
2
+ import os
3
+ import re
4
+ import subprocess
5
+ import uuid
6
+ from dataclasses import dataclass
7
+ from pathlib import Path
8
+ from types import TracebackType
9
+ from urllib.parse import quote, urlencode
10
+
11
+ import psycopg
12
+ from psycopg import sql
13
+ from psycopg.conninfo import conninfo_to_dict, make_conninfo
14
+
15
+ from dbreduce.postgres.dump import restore
16
+
17
+
18
+ @dataclass(frozen=True)
19
+ class DatabaseSettings:
20
+ encoding: str
21
+ lc_collate: str
22
+ lc_ctype: str
23
+ provider: str
24
+ locale: str | None
25
+ icu_rules: str | None
26
+
27
+
28
+ def read_settings(dsn: str) -> DatabaseSettings:
29
+ with psycopg.connect(dsn, connect_timeout=10) as conn:
30
+ row = conn.execute("""
31
+ SELECT pg_encoding_to_char(d.encoding), d.datcollate, d.datctype,
32
+ coalesce(to_jsonb(d)->>'datlocprovider', 'c'),
33
+ coalesce(to_jsonb(d)->>'datlocale',
34
+ to_jsonb(d)->>'daticulocale'),
35
+ to_jsonb(d)->>'daticurules'
36
+ FROM pg_database d WHERE d.datname = current_database()
37
+ """).fetchone()
38
+ if row is None:
39
+ raise ValueError("Cannot read source database settings")
40
+ return DatabaseSettings(*row)
41
+
42
+
43
+ def read_archive_settings(path: Path) -> DatabaseSettings:
44
+ """Use pg_restore's own CREATE DATABASE statement as archive metadata."""
45
+ with path.open("rb") as archive:
46
+ try:
47
+ result = subprocess.run(
48
+ [
49
+ "pg_restore",
50
+ "--create",
51
+ "--schema-only",
52
+ "--use-list",
53
+ os.devnull,
54
+ "--file",
55
+ "-",
56
+ ],
57
+ stdin=archive,
58
+ capture_output=True,
59
+ timeout=600,
60
+ env={**os.environ, "LC_ALL": "C"},
61
+ check=False,
62
+ )
63
+ except subprocess.TimeoutExpired as error:
64
+ raise ValueError("Timed out reading archive database settings") from error
65
+ if result.returncode:
66
+ if b"database name contains a newline or carriage return" in result.stderr:
67
+ raise ValueError(
68
+ "Custom archive has a database name containing a newline or carriage return; "
69
+ "pg_restore cannot safely read its settings"
70
+ )
71
+ raise ValueError("Cannot read database settings from custom archive")
72
+ statement = next(
73
+ (
74
+ line
75
+ for line in result.stdout.splitlines()
76
+ if line.startswith(b"CREATE DATABASE ") and b" WITH TEMPLATE = " in line
77
+ ),
78
+ None,
79
+ )
80
+ if statement is None:
81
+ raise ValueError("Archive does not expose source database settings")
82
+
83
+ encoding_match = re.search(rb"\bENCODING\s*=\s*'([^']+)'", statement)
84
+ if encoding_match is None:
85
+ raise ValueError("Archive database encoding is missing")
86
+ encoding = encoding_match.group(1).decode("ascii")
87
+ codec = (
88
+ "cp" + encoding[3:]
89
+ if encoding.startswith("WIN")
90
+ else "ascii"
91
+ if encoding == "SQL_ASCII"
92
+ else encoding
93
+ )
94
+
95
+ def setting(name: str) -> str | None:
96
+ match = re.search(rb"\b" + name.encode() + rb"\s*=\s*'((?:''|[^'])*)'", statement)
97
+ if not match:
98
+ return None
99
+ try:
100
+ return match.group(1).replace(b"''", b"'").decode(codec)
101
+ except (UnicodeError, LookupError) as error:
102
+ raise ValueError("Archive locale encoding cannot be decoded") from error
103
+
104
+ collate = setting("LC_COLLATE") or setting("LOCALE")
105
+ ctype = setting("LC_CTYPE") or setting("LOCALE")
106
+ provider_match = re.search(rb"\bLOCALE_PROVIDER\s*=\s*(\w+)", statement)
107
+ provider = {"libc": "c", "icu": "i", "builtin": "b"}.get(
108
+ provider_match.group(1).decode("ascii") if provider_match else "libc"
109
+ )
110
+ if not encoding or not collate or not ctype or provider is None:
111
+ raise ValueError("Archive database encoding/locale is incomplete")
112
+ locale = setting("ICU_LOCALE") or setting("BUILTIN_LOCALE") or setting("LOCALE")
113
+ return DatabaseSettings(encoding, collate, ctype, provider, locale, setting("ICU_RULES"))
114
+
115
+
116
+ def database_url(dsn: str, name: str) -> str:
117
+ params = conninfo_to_dict(dsn)
118
+ params.pop("dbname", None)
119
+ user = str(params.pop("user", ""))
120
+ password = params.pop("password", None)
121
+ authority = quote(user, safe="")
122
+ if password is not None:
123
+ authority += ":" + quote(str(password), safe="")
124
+ if authority:
125
+ authority += "@"
126
+ host = str(params.get("host", ""))
127
+ # Keep socket paths and multi-host libpq settings in the query string.
128
+ if host and "/" not in host and "," not in host:
129
+ params.pop("host")
130
+ authority += f"[{host}]" if ":" in host else host
131
+ port = params.pop("port", None)
132
+ if port is not None:
133
+ authority += ":" + str(port)
134
+ query = "?" + urlencode(params) if params else ""
135
+ return f"postgresql://{authority}/{quote(name, safe='')}{query}"
136
+
137
+
138
+ class Workspace:
139
+ """Own exactly one random database; no destructive method accepts a user database name."""
140
+
141
+ def __init__(self, admin_dsn: str, settings: DatabaseSettings | None = None) -> None:
142
+ self.admin_dsn = admin_dsn
143
+ self.name = f"dbreduce_{uuid.uuid4().hex}"
144
+ self.lock_key = int.from_bytes(
145
+ hashlib.sha256(self.name.encode()).digest()[:8], "big", signed=True
146
+ )
147
+ self.dsn = make_conninfo(admin_dsn, dbname=self.name)
148
+ self.url = database_url(admin_dsn, self.name)
149
+ self.settings = settings
150
+ self.created = False
151
+
152
+ def __enter__(self) -> "Workspace":
153
+ return self
154
+
155
+ def create(self) -> None:
156
+ self.close()
157
+ with psycopg.connect(self.admin_dsn, autocommit=True, connect_timeout=10) as conn:
158
+ conn.execute("SET statement_timeout = '600s'")
159
+ # Keep cleanup behind a CREATE that continues after its client loses the reply.
160
+ conn.execute("SELECT pg_advisory_lock(%s)", (self.lock_key,))
161
+ statement = sql.SQL("CREATE DATABASE {} TEMPLATE template0").format(
162
+ sql.Identifier(self.name)
163
+ )
164
+ if self.settings is not None:
165
+ settings = self.settings
166
+ statement += sql.SQL(" ENCODING {} LC_COLLATE {} LC_CTYPE {}").format(
167
+ sql.Literal(settings.encoding),
168
+ sql.Literal(settings.lc_collate),
169
+ sql.Literal(settings.lc_ctype),
170
+ )
171
+ if settings.provider == "i":
172
+ if settings.locale is None:
173
+ raise ValueError("Source ICU locale is missing")
174
+ statement += sql.SQL(" LOCALE_PROVIDER icu ICU_LOCALE {}").format(
175
+ sql.Literal(settings.locale)
176
+ )
177
+ if settings.icu_rules:
178
+ statement += sql.SQL(" ICU_RULES {}").format(
179
+ sql.Literal(settings.icu_rules)
180
+ )
181
+ elif settings.provider == "b":
182
+ if settings.locale is None:
183
+ raise ValueError("Source builtin locale is missing")
184
+ statement += sql.SQL(" LOCALE_PROVIDER builtin BUILTIN_LOCALE {}").format(
185
+ sql.Literal(settings.locale)
186
+ )
187
+ elif settings.provider != "c":
188
+ raise ValueError(f"Unsupported locale provider: {settings.provider}")
189
+ # CREATE may commit even if its reply is lost; cleanup must still try the name.
190
+ self.created = True
191
+ try:
192
+ conn.execute(statement)
193
+ except psycopg.errors.DuplicateDatabase:
194
+ # This name existed before our CREATE; it is not ours to drop.
195
+ self.created = False
196
+ raise
197
+
198
+ def reset(self, snapshot: Path) -> None:
199
+ self.create()
200
+ restore(self.dsn, snapshot)
201
+
202
+ def close(self) -> None:
203
+ if self.created:
204
+ try:
205
+ with psycopg.connect(self.admin_dsn, autocommit=True, connect_timeout=10) as conn:
206
+ conn.execute("SET statement_timeout = '600s'")
207
+ conn.execute("SELECT pg_advisory_lock(%s)", (self.lock_key,))
208
+ conn.execute(
209
+ sql.SQL("DROP DATABASE IF EXISTS {} WITH (FORCE)").format(
210
+ sql.Identifier(self.name)
211
+ )
212
+ )
213
+ except psycopg.Error as error:
214
+ raise RuntimeError(
215
+ f"Could not confirm cleanup of workspace database {self.name}; "
216
+ "inspect and remove it if present"
217
+ ) from error
218
+ self.created = False
219
+
220
+ def __exit__(
221
+ self,
222
+ exc_type: type[BaseException] | None,
223
+ exc: BaseException | None,
224
+ tb: TracebackType | None,
225
+ ) -> None:
226
+ try:
227
+ self.close()
228
+ except RuntimeError as cleanup_error:
229
+ if isinstance(exc, Exception):
230
+ original = (
231
+ "PostgreSQL operation failed; check connection and permissions"
232
+ if isinstance(exc, psycopg.Error)
233
+ else str(exc) or type(exc).__name__
234
+ )
235
+ raise RuntimeError(
236
+ f"Original failure: {original}; {cleanup_error}"
237
+ ) from cleanup_error
238
+ raise