isolab 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.
isolab/__init__.py ADDED
@@ -0,0 +1,2 @@
1
+ """isolab - a deterministic lab for PostgreSQL transaction-isolation anomalies."""
2
+ __version__ = "0.1.0"
isolab/__main__.py ADDED
@@ -0,0 +1,3 @@
1
+ from .cli import main
2
+
3
+ raise SystemExit(main())
isolab/cli.py ADDED
@@ -0,0 +1,147 @@
1
+ """Command line interface."""
2
+ from __future__ import annotations
3
+
4
+ import argparse
5
+ import os
6
+ import sys
7
+
8
+ from . import scenarios
9
+ from .engine import normalize, run_cells
10
+ from .model import LEVELS
11
+ from .report import render_json, render_markdown, render_table, render_trace, use_color
12
+
13
+ DEFAULT_DSN = "postgresql://postgres:postgres@localhost:5433/isolab"
14
+
15
+
16
+ def _levels(arg):
17
+ if not arg:
18
+ return list(LEVELS)
19
+ wanted = arg.upper().replace("-", " ").replace("_", " ")
20
+ short = {"RC": "READ COMMITTED", "RR": "REPEATABLE READ", "SER": "SERIALIZABLE", "SI": "REPEATABLE READ"}
21
+ wanted = short.get(wanted, wanted)
22
+ if wanted not in LEVELS:
23
+ raise SystemExit(f"unknown isolation level {arg!r}; choose from {', '.join(LEVELS)} (or rc/rr/ser)")
24
+ return [wanted]
25
+
26
+
27
+ def _pick(keys):
28
+ try:
29
+ return [scenarios.get(k) for k in keys] if keys else list(scenarios.ALL)
30
+ except KeyError as exc:
31
+ raise SystemExit(str(exc.args[0]))
32
+
33
+
34
+ def cmd_list(args):
35
+ for s in scenarios.ALL:
36
+ print(f"{s.key} - {s.title}\n {s.summary}")
37
+ for v in s.variants:
38
+ print(f" - {v.key:<22} {'fix' if v.is_fix else 'bug'} {v.title}")
39
+ print()
40
+
41
+
42
+ def cmd_explain(args):
43
+ s = _pick([args.scenario])[0]
44
+ print(f"{s.title}\n{'=' * len(s.title)}\n{s.summary}\nReference: {s.book_ref}\n")
45
+ for v in s.variants:
46
+ print(f"[{v.key}] {v.title}\n {v.description}")
47
+ print(" schedule: " + " -> ".join(f"{t}.{n}" for t, n in v.schedule))
48
+ print(" documented: " + ", ".join(f"{k.lower()}={e}" for k, e in v.expected.items()) + "\n")
49
+
50
+
51
+ def _execute(args, quiet=False):
52
+ dsn = args.dsn
53
+ chosen = _pick(getattr(args, "scenario", None))
54
+ done = [0]
55
+
56
+ def progress(cell):
57
+ done[0] += 1
58
+ if not quiet and sys.stderr.isatty():
59
+ print(f"\r ran {done[0]} cells...", end="", file=sys.stderr, flush=True)
60
+
61
+ try:
62
+ cells, info = run_cells(dsn, chosen, _levels(getattr(args, "level", None)),
63
+ variant_keys=[args.variant] if getattr(args, "variant", None) else None,
64
+ repeat=args.repeat, retry=not args.no_retry, progress=progress)
65
+ except Exception as exc: # connection problems are the common case
66
+ raise SystemExit(f"could not run against {dsn!r}: {exc}\n"
67
+ f"Start Postgres with `make up` or pass --dsn / set ISOLAB_DSN.")
68
+ if sys.stderr.isatty() and not quiet:
69
+ print("\r" + " " * 30 + "\r", end="", file=sys.stderr)
70
+ if not cells:
71
+ raise SystemExit("nothing matched those filters")
72
+ return cells, info
73
+
74
+
75
+ def cmd_run(args):
76
+ cells, info = _execute(args)
77
+ color = use_color() and not args.output
78
+ if args.format == "markdown":
79
+ text = render_markdown(cells, info, args.repeat, not args.no_retry)
80
+ elif args.format == "json":
81
+ text = render_json(cells, info)
82
+ else:
83
+ text = f"PostgreSQL {info['version']} (repeat={args.repeat}, retry={'off' if args.no_retry else 'on'})\n\n" \
84
+ + render_table(cells, color)
85
+ if args.output:
86
+ with open(args.output, "w", encoding="utf-8") as fh:
87
+ fh.write(text + "\n")
88
+ print(f"wrote {args.output}")
89
+ else:
90
+ print(text)
91
+ if args.trace and args.format == "table":
92
+ print()
93
+ for c in cells:
94
+ if c.skipped:
95
+ continue
96
+ outcome = c.dominant
97
+ print(render_trace(c.sample[outcome],
98
+ f"--- {c.scenario.key} / {c.variant.key} @ {c.level} ---", color))
99
+ print()
100
+
101
+
102
+ def cmd_verify(args):
103
+ args.no_retry = args.no_retry
104
+ cells, info = _execute(args, quiet=True)
105
+ bad = [c for c in cells if not c.skipped and not c.matches]
106
+ skipped = [c for c in cells if c.skipped]
107
+ print(f"PostgreSQL {info['version']}: {len(cells) - len(skipped)} cells checked, "
108
+ f"{len(skipped)} skipped, {len(bad)} mismatched")
109
+ for c in bad:
110
+ got = ", ".join(f"{k}×{n}" for k, n in c.outcomes.items())
111
+ print(f" MISMATCH {c.scenario.key}/{c.variant.key} @ {c.level}: "
112
+ f"expected {c.expected}, got {got}")
113
+ return 1 if bad else 0
114
+
115
+
116
+ def main(argv=None):
117
+ p = argparse.ArgumentParser(prog="isolab", description="Deterministic PostgreSQL isolation-anomaly lab.")
118
+ p.add_argument("--dsn", default=os.environ.get("ISOLAB_DSN", DEFAULT_DSN),
119
+ help="libpq connection string (default: $ISOLAB_DSN or %(default)s)")
120
+ sub = p.add_subparsers(dest="cmd", required=True)
121
+
122
+ sub.add_parser("list", help="list scenarios and variants").set_defaults(fn=cmd_list)
123
+ e = sub.add_parser("explain", help="show a scenario's story, schedules and documented results")
124
+ e.add_argument("scenario")
125
+ e.set_defaults(fn=cmd_explain)
126
+
127
+ def common(sp):
128
+ sp.add_argument("scenario", nargs="*", help="scenario keys (default: all)")
129
+ sp.add_argument("--variant", help="only this variant (e.g. naive, for_update)")
130
+ sp.add_argument("--level", help="only this isolation level (rc | rr | ser or full name)")
131
+ sp.add_argument("--repeat", type=int, default=1, help="run every cell N times to prove determinism")
132
+ sp.add_argument("--no-retry", action="store_true",
133
+ help="do not retry aborted transactions (report them as ABORTED)")
134
+
135
+ r = sub.add_parser("run", help="run scenarios and print the anomaly matrix")
136
+ common(r)
137
+ r.add_argument("--format", choices=["table", "markdown", "json"], default="table")
138
+ r.add_argument("--output", "-o", help="write the report to a file")
139
+ r.add_argument("--trace", action="store_true", help="print the step-by-step timeline of each cell")
140
+ r.set_defaults(fn=cmd_run)
141
+
142
+ v = sub.add_parser("verify", help="exit 1 if any result differs from the documented expectation")
143
+ common(v)
144
+ v.set_defaults(fn=cmd_verify)
145
+
146
+ args = p.parse_args(argv)
147
+ return args.fn(args) or 0
isolab/engine.py ADDED
@@ -0,0 +1,417 @@
1
+ """Execution engine: a deterministic scheduler over real PostgreSQL connections.
2
+
3
+ Every transaction in a variant gets its own connection and worker thread. The
4
+ scheduler releases *one step at a time* in the order given by the variant's
5
+ schedule. If a step blocks on a lock (detected through pg_stat_activity, not a
6
+ sleep), the scheduler moves on to the next step; the blocked step completes
7
+ later, when the lock holder commits. After every step the scheduler waits for
8
+ the system to "settle" (every worker idle or lock-blocked) so the interleaving
9
+ is identical on every run.
10
+ """
11
+ from __future__ import annotations
12
+
13
+ import queue
14
+ import threading
15
+ import time
16
+ from collections import Counter
17
+ from dataclasses import dataclass, field
18
+ from typing import Callable, Optional
19
+
20
+ import psycopg
21
+ from psycopg import errors as pgerr
22
+
23
+ from .model import ABORTED, AppConflict, LEVELS, Scenario, Variant
24
+
25
+ SCHEMA = "isolab"
26
+ CONN_OPTIONS = f"-c search_path={SCHEMA},public"
27
+
28
+
29
+ class EngineError(RuntimeError):
30
+ """The harness itself failed (timeout, stuck worker, ...)."""
31
+
32
+
33
+ # --------------------------------------------------------------------------- errors
34
+
35
+ def classify(exc: BaseException) -> tuple[str, bool]:
36
+ """Map an exception to (kind, retryable)."""
37
+ table = (
38
+ (pgerr.SerializationFailure, "serialization_failure"),
39
+ (pgerr.DeadlockDetected, "deadlock"),
40
+ (pgerr.UniqueViolation, "unique_violation"),
41
+ (pgerr.ExclusionViolation, "exclusion_violation"),
42
+ (pgerr.CheckViolation, "check_violation"),
43
+ (AppConflict, "app_conflict"),
44
+ )
45
+ for cls, kind in table:
46
+ if isinstance(exc, cls):
47
+ return kind, True
48
+ return "error", False
49
+
50
+
51
+ # --------------------------------------------------------------------------- session
52
+
53
+ def _inline(sql: str, params: tuple) -> str:
54
+ out = " ".join(sql.split())
55
+ for p in params:
56
+ out = out.replace("%s", f"'{p}'" if isinstance(p, str) else str(p), 1)
57
+ return out
58
+
59
+
60
+ class Session:
61
+ """What scenario steps talk to: a connection plus a per-transaction ctx dict."""
62
+
63
+ def __init__(self, conn, log: Optional[Callable[[str], None]] = None, ctx: Optional[dict] = None):
64
+ self.conn = conn
65
+ self.ctx = ctx if ctx is not None else {}
66
+ self._log = log or (lambda line: None)
67
+
68
+ def raw(self, sql: str) -> None:
69
+ self._log(sql)
70
+ self.conn.execute(sql)
71
+
72
+ def _go(self, sql: str, params: tuple):
73
+ try:
74
+ return self.conn.execute(sql, params or None)
75
+ except Exception as exc:
76
+ self._log(f"{_inline(sql, params)} -> ERROR")
77
+ raise exc
78
+
79
+ def rows(self, sql: str, *params) -> list:
80
+ data = self._go(sql, params).fetchall()
81
+ self._log(f"{_inline(sql, params)} -> {data}")
82
+ return data
83
+
84
+ def scalar(self, sql: str, *params):
85
+ return self.rows(sql, *params)[0][0]
86
+
87
+ def run(self, sql: str, *params) -> int:
88
+ n = self._go(sql, params).rowcount
89
+ self._log(f"{_inline(sql, params)} -> {n} row(s)")
90
+ return n
91
+
92
+
93
+ # --------------------------------------------------------------------------- trace
94
+
95
+ @dataclass
96
+ class Event:
97
+ seq: int
98
+ txn: str
99
+ step: str
100
+ status: str # ok | blocked | error | skipped
101
+ lines: list
102
+ note: str = ""
103
+ ts: float = 0.0 # ordering key; COMMIT events use the moment COMMIT was *sent*
104
+
105
+
106
+ class Trace:
107
+ def __init__(self):
108
+ self.events: list[Event] = []
109
+ self._lock = threading.Lock()
110
+
111
+ def add(self, txn, step, status, lines=(), note="", ts=None):
112
+ with self._lock:
113
+ self.events.append(Event(len(self.events) + 1, txn, step, status, list(lines), note,
114
+ time.monotonic() if ts is None else ts))
115
+
116
+ def ordered(self) -> list:
117
+ """Events in causal order. A COMMIT releases locks on the server *before* its
118
+ client sees the reply, so a transaction it unblocks can finish first; stamping
119
+ commits at send time keeps the timeline readable."""
120
+ return sorted(self.events, key=lambda e: (e.ts, e.seq))
121
+
122
+
123
+ @dataclass
124
+ class TxnResult:
125
+ label: str
126
+ ctx: dict = field(default_factory=dict)
127
+ committed: bool = False
128
+ first_error: Optional[str] = None # kind of the first failure, if any
129
+ first_detail: str = ""
130
+ first_retryable: bool = False
131
+ attempts: int = 1
132
+
133
+
134
+ def connect(dsn: str):
135
+ return psycopg.connect(dsn, autocommit=True, options=CONN_OPTIONS)
136
+
137
+
138
+ def _waiting_on_lock(obs, pid: int) -> bool:
139
+ row = obs.execute("SELECT wait_event_type FROM pg_stat_activity WHERE pid = %s", (pid,)).fetchone()
140
+ return bool(row) and row[0] == "Lock"
141
+
142
+
143
+ # --------------------------------------------------------------------------- worker
144
+
145
+ class Worker(threading.Thread):
146
+ def __init__(self, label: str, steps: list, dsn: str, level: str, trace: Trace):
147
+ super().__init__(daemon=True, name=f"txn-{label}")
148
+ if level not in LEVELS:
149
+ raise ValueError(level)
150
+ self.label, self.level, self.trace = label, level, trace
151
+ self.steps = {s.name: s for s in steps}
152
+ self.conn = connect(dsn)
153
+ self.pid = self.conn.info.backend_pid
154
+ self.result = TxnResult(label)
155
+ self.q: queue.Queue = queue.Queue()
156
+ self._pending = 0
157
+ self._lock = threading.Lock()
158
+ self._begun = False
159
+ self._failed = False
160
+
161
+ def submit(self, step_name: str) -> threading.Event:
162
+ done = threading.Event()
163
+ with self._lock:
164
+ self._pending += 1
165
+ self.q.put((step_name, done))
166
+ return done
167
+
168
+ def idle(self) -> bool:
169
+ with self._lock:
170
+ return self._pending == 0
171
+
172
+ def run(self):
173
+ while True:
174
+ item = self.q.get()
175
+ if item is None:
176
+ return
177
+ name, done = item
178
+ try:
179
+ self._exec(name)
180
+ finally:
181
+ with self._lock:
182
+ self._pending -= 1
183
+ done.set()
184
+
185
+ def _exec(self, name: str):
186
+ if self._failed:
187
+ self.trace.add(self.label, name, "skipped", note="transaction already aborted")
188
+ return
189
+ lines: list = []
190
+ started = time.monotonic()
191
+ stamp = started if name == "commit" else None
192
+ s = Session(self.conn, lines.append, self.result.ctx)
193
+ try:
194
+ if not self._begun:
195
+ s.raw(f"BEGIN ISOLATION LEVEL {self.level}")
196
+ self._begun = True
197
+ if name == "commit":
198
+ s.raw("COMMIT")
199
+ self.result.committed = True
200
+ else:
201
+ self.steps[name].fn(s)
202
+ self.trace.add(self.label, name, "ok", lines, ts=stamp)
203
+ except Exception as exc: # noqa: BLE001 - we classify everything
204
+ kind, retryable = classify(exc)
205
+ self._failed = True
206
+ r = self.result
207
+ r.first_error, r.first_retryable = kind, retryable
208
+ r.first_detail = (str(exc).strip().splitlines() or [repr(exc)])[0]
209
+ try:
210
+ self.conn.execute("ROLLBACK")
211
+ except Exception: # noqa: BLE001
212
+ pass
213
+ self.trace.add(self.label, name, "error", lines, f"{kind}: {r.first_detail}", ts=stamp)
214
+
215
+
216
+ # --------------------------------------------------------------------------- scheduling
217
+
218
+ def _run_schedule(workers: dict, schedule: list, obs, trace: Trace, timeout: float) -> None:
219
+ def blocked(w: Worker) -> bool:
220
+ return _waiting_on_lock(obs, w.pid)
221
+
222
+ def settle():
223
+ deadline = time.monotonic() + timeout
224
+ while True:
225
+ if all(w.idle() or blocked(w) for w in workers.values()):
226
+ return
227
+ if time.monotonic() > deadline:
228
+ raise EngineError("workers never settled (stuck outside a lock wait?)")
229
+ time.sleep(0.002)
230
+
231
+ for label, name in schedule:
232
+ w = workers[label]
233
+ done = w.submit(name)
234
+ deadline = time.monotonic() + timeout
235
+ while not done.wait(0.004):
236
+ if blocked(w):
237
+ trace.add(label, name, "blocked", note=f"txn {label} is waiting on a lock; schedule continues")
238
+ break
239
+ if time.monotonic() > deadline:
240
+ raise EngineError(f"step {label}.{name} timed out")
241
+ settle()
242
+
243
+
244
+ def _drain(workers: dict, timeout: float) -> None:
245
+ for w in workers.values():
246
+ w.q.put(None)
247
+ deadline = time.monotonic() + timeout
248
+ for w in workers.values():
249
+ w.join(max(0.1, deadline - time.monotonic()))
250
+ if w.is_alive():
251
+ raise EngineError(f"txn {w.label} never finished (deadlock the server did not detect?)")
252
+
253
+
254
+ def _retry(dsn: str, variant: Variant, results: dict, level: str, trace: Trace, max_attempts: int) -> None:
255
+ """Re-run aborted transactions the way a well-behaved application would:
256
+ from the top, in a fresh transaction, one after another."""
257
+ for label, res in results.items():
258
+ if res.committed or not res.first_retryable:
259
+ continue
260
+ for attempt in range(2, max_attempts + 1):
261
+ conn = connect(dsn)
262
+ res.attempts, res.ctx = attempt, {}
263
+ lines: list = []
264
+ s = Session(conn, lines.append, res.ctx)
265
+ try:
266
+ s.raw(f"BEGIN ISOLATION LEVEL {level}")
267
+ for st in variant.txns[label]:
268
+ st.fn(s)
269
+ s.raw("COMMIT")
270
+ res.committed = True
271
+ trace.add(f"{label}-retry", f"attempt {attempt}", "ok", lines)
272
+ break
273
+ except Exception as exc: # noqa: BLE001
274
+ kind, retryable = classify(exc)
275
+ try:
276
+ conn.execute("ROLLBACK")
277
+ except Exception: # noqa: BLE001
278
+ pass
279
+ trace.add(f"{label}-retry", f"attempt {attempt}", "error", lines, kind)
280
+ if not retryable:
281
+ break
282
+ finally:
283
+ conn.close()
284
+
285
+
286
+ # --------------------------------------------------------------------------- cells
287
+
288
+ @dataclass
289
+ class CellResult:
290
+ outcome: str # SAFE | ABORTED | RETRIED | ANOMALY | ERROR
291
+ violations: list
292
+ txns: dict
293
+ trace: Trace
294
+ duration: float
295
+ error: str = ""
296
+
297
+
298
+ def normalize(outcome: str) -> str:
299
+ """RETRIED (aborted, then retried successfully) counts as ABORTED for expectations."""
300
+ return ABORTED if outcome == "RETRIED" else outcome
301
+
302
+
303
+ def _outcome(violations: list, results: dict) -> str:
304
+ if violations:
305
+ return "ANOMALY"
306
+ failed = [r for r in results.values() if r.first_error]
307
+ if any(not r.first_retryable for r in failed):
308
+ return "ERROR"
309
+ if not failed:
310
+ return "SAFE"
311
+ return "RETRIED" if all(r.committed for r in failed) else "ABORTED"
312
+
313
+
314
+ def reset_schema(obs, scenario: Scenario, variant: Variant) -> None:
315
+ obs.execute(f"DROP SCHEMA IF EXISTS {SCHEMA} CASCADE")
316
+ obs.execute(f"CREATE SCHEMA {SCHEMA}")
317
+ for stmt in [*scenario.setup, *variant.setup]:
318
+ obs.execute(stmt)
319
+
320
+
321
+ def run_cell(dsn: str, scenario: Scenario, variant: Variant, level: str, *,
322
+ retry: bool = True, timeout: float = 10.0, max_attempts: int = 3) -> CellResult:
323
+ t0 = time.monotonic()
324
+ trace = Trace()
325
+ obs = connect(dsn)
326
+ workers: dict = {}
327
+ try:
328
+ reset_schema(obs, scenario, variant)
329
+ workers = {lbl: Worker(lbl, steps, dsn, level, trace) for lbl, steps in variant.txns.items()}
330
+ for w in workers.values():
331
+ w.start()
332
+ _run_schedule(workers, variant.schedule, obs, trace, timeout)
333
+ _drain(workers, timeout)
334
+ results = {lbl: w.result for lbl, w in workers.items()}
335
+ if retry:
336
+ _retry(dsn, variant, results, level, trace, max_attempts)
337
+ violations = scenario.invariant(Session(obs), results)
338
+ return CellResult(_outcome(violations, results), violations, results, trace,
339
+ time.monotonic() - t0)
340
+ except EngineError as exc:
341
+ return CellResult("ERROR", [], {l: w.result for l, w in workers.items()}, trace,
342
+ time.monotonic() - t0, error=str(exc))
343
+ finally:
344
+ for w in workers.values():
345
+ try:
346
+ w.conn.cancel()
347
+ except Exception: # noqa: BLE001
348
+ pass
349
+ w.join(1)
350
+ try:
351
+ w.conn.close()
352
+ except Exception: # noqa: BLE001
353
+ pass
354
+ obs.close()
355
+
356
+
357
+ @dataclass
358
+ class Cell:
359
+ scenario: Scenario
360
+ variant: Variant
361
+ level: str
362
+ outcomes: Counter = field(default_factory=Counter)
363
+ sample: dict = field(default_factory=dict) # outcome -> first CellResult with that outcome
364
+ skipped: str = ""
365
+
366
+ @property
367
+ def dominant(self) -> str:
368
+ return self.outcomes.most_common(1)[0][0] if self.outcomes else "SKIPPED"
369
+
370
+ @property
371
+ def deterministic(self) -> bool:
372
+ return len(self.outcomes) <= 1
373
+
374
+ @property
375
+ def expected(self) -> str:
376
+ return self.variant.expected[self.level]
377
+
378
+ @property
379
+ def matches(self) -> bool:
380
+ return (not self.skipped and self.deterministic
381
+ and normalize(self.dominant) == self.expected)
382
+
383
+
384
+ def preflight(dsn: str) -> dict:
385
+ with psycopg.connect(dsn, autocommit=True) as c:
386
+ version = c.execute("SHOW server_version").fetchone()[0]
387
+ extensions = set()
388
+ try:
389
+ c.execute("CREATE EXTENSION IF NOT EXISTS btree_gist")
390
+ extensions.add("btree_gist")
391
+ except Exception: # noqa: BLE001
392
+ pass
393
+ return {"version": version, "extensions": extensions}
394
+
395
+
396
+ def run_cells(dsn: str, scenarios: list, levels: list, *, variant_keys: Optional[list] = None,
397
+ repeat: int = 1, retry: bool = True, progress: Optional[Callable] = None) -> tuple[list, dict]:
398
+ info = preflight(dsn)
399
+ cells: list = []
400
+ for sc in scenarios:
401
+ for v in sc.variants:
402
+ if variant_keys and v.key not in variant_keys:
403
+ continue
404
+ for lvl in levels:
405
+ cell = Cell(sc, v, lvl)
406
+ missing = [e for e in v.requires if e not in info["extensions"]]
407
+ if missing:
408
+ cell.skipped = f"needs extension {', '.join(missing)}"
409
+ else:
410
+ for _ in range(repeat):
411
+ res = run_cell(dsn, sc, v, lvl, retry=retry)
412
+ cell.outcomes[res.outcome] += 1
413
+ cell.sample.setdefault(res.outcome, res)
414
+ cells.append(cell)
415
+ if progress:
416
+ progress(cell)
417
+ return cells, info
isolab/model.py ADDED
@@ -0,0 +1,89 @@
1
+ """Data model: scenarios, variants and steps.
2
+
3
+ A *Scenario* is a business rule that can be broken by concurrency
4
+ ("a doctor must always be on call"). Each scenario has one or more
5
+ *Variants*: the naive implementation plus the fixes that work (or
6
+ deliberately don't). A variant is a set of transactions, each a list of
7
+ named *Steps*, and a *schedule* that fixes the exact interleaving.
8
+ """
9
+ from __future__ import annotations
10
+
11
+ from dataclasses import dataclass, field
12
+ from typing import Callable
13
+
14
+ LEVELS = ("READ COMMITTED", "REPEATABLE READ", "SERIALIZABLE")
15
+
16
+ # Vocabulary used for documented expectations.
17
+ ANOMALY = "ANOMALY" # the invariant was violated
18
+ ABORTED = "ABORTED" # the database/app rejected a transaction; invariant holds
19
+ SAFE = "SAFE" # nothing was rejected and the invariant holds
20
+ EXPECTATIONS = (ANOMALY, ABORTED, SAFE)
21
+
22
+
23
+ class AppConflict(Exception):
24
+ """Raised by application-level optimistic checks (e.g. a version column
25
+ mismatch). Treated like a retryable serialization failure."""
26
+
27
+
28
+ @dataclass(frozen=True)
29
+ class Step:
30
+ name: str
31
+ fn: Callable # fn(session) -> None; see engine.Session
32
+
33
+
34
+ @dataclass
35
+ class Variant:
36
+ key: str
37
+ title: str
38
+ description: str
39
+ txns: dict # label -> list[Step]
40
+ schedule: list # [(label, step_name | "commit"), ...]
41
+ expected: dict # level -> ANOMALY | ABORTED | SAFE
42
+ setup: list = field(default_factory=list) # extra DDL/DML after the scenario setup
43
+ is_fix: bool = True
44
+ requires: tuple = () # PostgreSQL extensions needed
45
+
46
+ def validate(self, scenario_key: str) -> None:
47
+ where = f"{scenario_key}/{self.key}"
48
+ if set(self.expected) != set(LEVELS):
49
+ raise ValueError(f"{where}: expected must cover exactly {LEVELS}")
50
+ for lvl, exp in self.expected.items():
51
+ if exp not in EXPECTATIONS:
52
+ raise ValueError(f"{where}: bad expectation {exp!r} for {lvl}")
53
+ per_txn: dict = {label: [] for label in self.txns}
54
+ for label, step in self.schedule:
55
+ if label not in self.txns:
56
+ raise ValueError(f"{where}: schedule references unknown txn {label!r}")
57
+ names = {s.name for s in self.txns[label]}
58
+ if step != "commit" and step not in names:
59
+ raise ValueError(f"{where}: txn {label} has no step {step!r}")
60
+ per_txn[label].append(step)
61
+ for label, steps in self.txns.items():
62
+ declared = [s.name for s in steps] + ["commit"]
63
+ if per_txn[label] != declared:
64
+ raise ValueError(
65
+ f"{where}: schedule for txn {label} is {per_txn[label]} "
66
+ f"but its steps are {declared} (each step once, in order, then commit)"
67
+ )
68
+
69
+
70
+ @dataclass
71
+ class Scenario:
72
+ key: str
73
+ title: str
74
+ summary: str
75
+ book_ref: str
76
+ setup: list # SQL run once per cell, in a fresh schema
77
+ invariant: Callable # invariant(db: Session, txns: dict[str, TxnResult]) -> list[str]
78
+ variants: list
79
+
80
+ def __post_init__(self):
81
+ for v in self.variants:
82
+ v.validate(self.key)
83
+
84
+ def variant(self, key: str) -> Variant:
85
+ for v in self.variants:
86
+ if v.key == key:
87
+ return v
88
+ raise KeyError(f"{self.key} has no variant {key!r} "
89
+ f"(choose from {[v.key for v in self.variants]})")