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 +1 -0
- dbreduce/__main__.py +3 -0
- dbreduce/cache/__init__.py +0 -0
- dbreduce/cache/store.py +32 -0
- dbreduce/cli/__init__.py +193 -0
- dbreduce/graph/__init__.py +0 -0
- dbreduce/graph/dependencies.py +33 -0
- dbreduce/models/__init__.py +0 -0
- dbreduce/models/schema.py +33 -0
- dbreduce/oracle/__init__.py +0 -0
- dbreduce/oracle/runner.py +52 -0
- dbreduce/postgres/__init__.py +0 -0
- dbreduce/postgres/backend.py +64 -0
- dbreduce/postgres/database.py +238 -0
- dbreduce/postgres/dump.py +85 -0
- dbreduce/postgres/introspection.py +113 -0
- dbreduce/postgres/rows.py +112 -0
- dbreduce/reducer/__init__.py +0 -0
- dbreduce/reducer/ddmin.py +17 -0
- dbreduce/reducer/engine.py +35 -0
- dbreduce-0.1.0.dist-info/METADATA +211 -0
- dbreduce-0.1.0.dist-info/RECORD +25 -0
- dbreduce-0.1.0.dist-info/WHEEL +4 -0
- dbreduce-0.1.0.dist-info/entry_points.txt +2 -0
- dbreduce-0.1.0.dist-info/licenses/LICENSE +201 -0
dbreduce/__init__.py
ADDED
|
@@ -0,0 +1 @@
|
|
|
1
|
+
"""PostgreSQL reproducer reduction."""
|
dbreduce/__main__.py
ADDED
|
File without changes
|
dbreduce/cache/store.py
ADDED
|
@@ -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
|
dbreduce/cli/__init__.py
ADDED
|
@@ -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
|