diff --git a/src/openstatspec/spss/sav.py b/src/openstatspec/spss/sav.py index b914dda..5a9af41 100644 --- a/src/openstatspec/spss/sav.py +++ b/src/openstatspec/spss/sav.py @@ -54,8 +54,9 @@ from ..sql.wide import ( create_wide_dataset, physical_name, - read_fidelity_events, - read_wide_dataset, + _read_fidelity_events, + _read_snapshot, + _read_wide_dataset, validate_spss_catalog, ) @@ -556,20 +557,21 @@ def export_sav_dataset( destination_path = Path(destination) if destination_path.suffix.lower() not in {".sav", ".zsav"}: raise UnsupportedOperationError("Export destinations must use the .sav or .zsav extension.") - dataset, variables, rows = read_wide_dataset( - database_url=database_url, dataset_id=dataset_id, + with _read_snapshot( + database_url=database_url, dolt_conformance_source=dolt_conformance_source, - ) + ) as (connection, profile): + dataset, variables, rows = _read_wide_dataset( + connection, dataset_id=dataset_id, profile=profile, + ) + persisted_events = _read_fidelity_events( + connection, dataset_id=dataset["dataset_id"], direction="import", + ) validate_spss_catalog( variables, case_weight_variable=dataset.get("case_weight_variable"), multiple_response_sets=dataset.get("multiple_response_sets"), ) - persisted_events = read_fidelity_events( - database_url=database_url, dataset_id=dataset_id, - direction="import", - dolt_conformance_source=dolt_conformance_source, - ) if legacy_locale is not None: persisted_events = tuple( event for event in persisted_events diff --git a/src/openstatspec/sql/wide.py b/src/openstatspec/sql/wide.py index 7f0d128..99bfa79 100644 --- a/src/openstatspec/sql/wide.py +++ b/src/openstatspec/sql/wide.py @@ -1045,131 +1045,158 @@ def create_wide_dataset( -def read_wide_dataset( - *, database_url: str, dataset_id: str, profile: Any | None = None, +@contextmanager +def _read_snapshot( + *, database_url: str, profile: Any | None = None, dolt_conformance_source: Any | None = None, -) -> tuple[dict[str, Any], list[dict[str, Any]], list[dict[str, Any]]]: - """Read an export descriptor from the normative catalog without mutation.""" +): + """Own one verified native read snapshot, rolling back on every exit.""" require_existing_database_url(database_url) if profile is None: profile, _active = read_profile( database_url, dolt_conformance_source=dolt_conformance_source, ) engine = create_engine(database_url) + try: + with engine.connect() as connection: + if profile.name in {"postgresql", "mysql", "mariadb", "dolt"}: + connection = connection.execution_options(isolation_level="REPEATABLE READ") + if profile.name in {"sqlite", "dolt"}: + # sqlite3 SELECTs and Dolt's retained autocommit need native BEGIN. + connection.exec_driver_sql("BEGIN") + require_verified_catalog(connection) + yield connection, profile + finally: + engine.dispose() + + +def read_wide_dataset( + *, database_url: str, dataset_id: str, profile: Any | None = None, + dolt_conformance_source: Any | None = None, +) -> tuple[dict[str, Any], list[dict[str, Any]], list[dict[str, Any]]]: + """Read an export descriptor from the normative catalog without mutation.""" + with _read_snapshot( + database_url=database_url, profile=profile, + dolt_conformance_source=dolt_conformance_source, + ) as (connection, profile): + return _read_wide_dataset(connection, dataset_id=dataset_id, profile=profile) + + +def _read_wide_dataset( + connection: Any, *, dataset_id: str, profile: Any, +) -> tuple[dict[str, Any], list[dict[str, Any]], list[dict[str, Any]]]: normative = normative_catalog(MetaData()) - with engine.connect() as connection: - require_verified_catalog(connection) - dataset_row = _resolve_normative_dataset(connection, normative, dataset_id) - core_id = str(dataset_row["dataset_id"]) - data_table = Table( - str(dataset_row["physical_table_name"]), MetaData(), - schema=dataset_row["physical_table_schema"], - autoload_with=connection, + dataset_row = _resolve_normative_dataset(connection, normative, dataset_id) + core_id = str(dataset_row["dataset_id"]) + data_table = Table( + str(dataset_row["physical_table_name"]), MetaData(), + schema=dataset_row["physical_table_schema"], + autoload_with=connection, + ) + source_variables = connection.execute( + select(normative.variable) + .where(normative.variable.c.dataset_id == core_id) + .order_by(normative.variable.c.source_ordinal) + ).mappings().all() + variables = [_export_variable(row) for row in source_variables] + variables_by_id = { + str(row["variable_id"]): variable + for row, variable in zip(source_variables, variables, strict=True) + } + variable_ids = tuple(variables_by_id) + rows = [dict(row) for row in connection.execute( + select(data_table).order_by(data_table.c.__case_ordinal) + ).mappings()] + rows = _canonicalize_database_numeric_rows(rows, variables) + documents = connection.execute( + select(normative.document) + .where(normative.document.c.dataset_id == core_id) + .order_by(normative.document.c.source_ordinal) + ).mappings().all() + dataset_attributes = connection.execute( + select(normative.dataset_attribute) + .where(normative.dataset_attribute.c.dataset_id == core_id) + .order_by( + normative.dataset_attribute.c.attribute_name, + normative.dataset_attribute.c.array_ordinal, ) - source_variables = connection.execute( - select(normative.variable) - .where(normative.variable.c.dataset_id == core_id) - .order_by(normative.variable.c.source_ordinal) - ).mappings().all() - variables = [_export_variable(row) for row in source_variables] - variables_by_id = { - str(row["variable_id"]): variable - for row, variable in zip(source_variables, variables, strict=True) - } - variable_ids = tuple(variables_by_id) - rows = [dict(row) for row in connection.execute( - select(data_table).order_by(data_table.c.__case_ordinal) - ).mappings()] - rows = _canonicalize_database_numeric_rows(rows, variables) - documents = connection.execute( - select(normative.document) - .where(normative.document.c.dataset_id == core_id) - .order_by(normative.document.c.source_ordinal) - ).mappings().all() - dataset_attributes = connection.execute( - select(normative.dataset_attribute) - .where(normative.dataset_attribute.c.dataset_id == core_id) - .order_by( - normative.dataset_attribute.c.attribute_name, - normative.dataset_attribute.c.array_ordinal, - ) - ).mappings().all() - variable_attributes = connection.execute( - select(normative.variable_attribute) - .where(normative.variable_attribute.c.variable_id.in_(variable_ids)) - .order_by( - normative.variable_attribute.c.variable_id, - normative.variable_attribute.c.attribute_name, - normative.variable_attribute.c.array_ordinal, - ) - ).mappings().all() - labels = connection.execute( - select( - normative.variable_value_label_set.c.variable_id, - normative.value_label, - ) - .join( - normative.value_label, - normative.value_label.c.value_label_set_id - == normative.variable_value_label_set.c.value_label_set_id, - ) - .where( - normative.variable_value_label_set.c.variable_id.in_(variable_ids) - ) - .order_by( - normative.variable_value_label_set.c.variable_id, - normative.value_label.c.ordinal, - ) - ).mappings().all() - missing_rules = connection.execute( - select(normative.missing_rule) - .where(normative.missing_rule.c.variable_id.in_(variable_ids)) - .order_by( - normative.missing_rule.c.variable_id, - normative.missing_rule.c.ordinal, - ) - ).mappings().all() - variable_sets = connection.execute( - select(normative.variable_set) - .where(normative.variable_set.c.dataset_id == core_id) - .order_by(normative.variable_set.c.source_ordinal) - ).mappings().all() - variable_set_ids = tuple( - str(row["variable_set_id"]) for row in variable_sets + ).mappings().all() + variable_attributes = connection.execute( + select(normative.variable_attribute) + .where(normative.variable_attribute.c.variable_id.in_(variable_ids)) + .order_by( + normative.variable_attribute.c.variable_id, + normative.variable_attribute.c.attribute_name, + normative.variable_attribute.c.array_ordinal, ) - variable_set_members = connection.execute( - select(normative.variable_set_member) - .where(normative.variable_set_member.c.variable_set_id.in_(variable_set_ids)) - .order_by( - normative.variable_set_member.c.variable_set_id, - normative.variable_set_member.c.source_ordinal, - ) - ).mappings().all() - response_sets = connection.execute( - select(normative.multiple_response_set) - .where(normative.multiple_response_set.c.dataset_id == core_id) - .order_by(normative.multiple_response_set.c.source_ordinal) - ).mappings().all() - response_set_ids = tuple( - str(row["multiple_response_set_id"]) for row in response_sets + ).mappings().all() + labels = connection.execute( + select( + normative.variable_value_label_set.c.variable_id, + normative.value_label, ) - response_members = connection.execute( - select(normative.multiple_response_member) - .where( - normative.multiple_response_member.c.multiple_response_set_id.in_( - response_set_ids - ) - ) - .order_by( - normative.multiple_response_member.c.multiple_response_set_id, - normative.multiple_response_member.c.source_ordinal, - ) - ).mappings().all() - weight_id = connection.execute( - select(normative.dataset_weight_variable.c.variable_id).where( - normative.dataset_weight_variable.c.dataset_id == core_id + .join( + normative.value_label, + normative.value_label.c.value_label_set_id + == normative.variable_value_label_set.c.value_label_set_id, + ) + .where( + normative.variable_value_label_set.c.variable_id.in_(variable_ids) + ) + .order_by( + normative.variable_value_label_set.c.variable_id, + normative.value_label.c.ordinal, + ) + ).mappings().all() + missing_rules = connection.execute( + select(normative.missing_rule) + .where(normative.missing_rule.c.variable_id.in_(variable_ids)) + .order_by( + normative.missing_rule.c.variable_id, + normative.missing_rule.c.ordinal, + ) + ).mappings().all() + variable_sets = connection.execute( + select(normative.variable_set) + .where(normative.variable_set.c.dataset_id == core_id) + .order_by(normative.variable_set.c.source_ordinal) + ).mappings().all() + variable_set_ids = tuple( + str(row["variable_set_id"]) for row in variable_sets + ) + variable_set_members = connection.execute( + select(normative.variable_set_member) + .where(normative.variable_set_member.c.variable_set_id.in_(variable_set_ids)) + .order_by( + normative.variable_set_member.c.variable_set_id, + normative.variable_set_member.c.source_ordinal, + ) + ).mappings().all() + response_sets = connection.execute( + select(normative.multiple_response_set) + .where(normative.multiple_response_set.c.dataset_id == core_id) + .order_by(normative.multiple_response_set.c.source_ordinal) + ).mappings().all() + response_set_ids = tuple( + str(row["multiple_response_set_id"]) for row in response_sets + ) + response_members = connection.execute( + select(normative.multiple_response_member) + .where( + normative.multiple_response_member.c.multiple_response_set_id.in_( + response_set_ids ) - ).scalar_one_or_none() + ) + .order_by( + normative.multiple_response_member.c.multiple_response_set_id, + normative.multiple_response_member.c.source_ordinal, + ) + ).mappings().all() + weight_id = connection.execute( + select(normative.dataset_weight_variable.c.variable_id).where( + normative.dataset_weight_variable.c.dataset_id == core_id + ) + ).scalar_one_or_none() dataset = { "dataset_id": core_id, @@ -1386,26 +1413,33 @@ def read_fidelity_events( dolt_conformance_source: Any | None = None, ) -> tuple[dict[str, Any], ...]: """Read fidelity diagnostics, optionally limited to one lifecycle direction.""" - read_profile( - database_url, dolt_conformance_source=dolt_conformance_source, - ) - engine = create_engine(database_url) - normative = normative_catalog(MetaData()) - with engine.connect() as connection: - require_verified_catalog(connection) + with _read_snapshot( + database_url=database_url, + dolt_conformance_source=dolt_conformance_source, + ) as (connection, _profile): + normative = normative_catalog(MetaData()) dataset = _resolve_normative_dataset(connection, normative, dataset_id) - statement = ( - select(normative.fidelity_event) - .where(normative.fidelity_event.c.dataset_id == dataset["dataset_id"]) - .where(normative.fidelity_event.c.severity != "info") + return _read_fidelity_events( + connection, dataset_id=str(dataset["dataset_id"]), direction=direction, ) - if direction is not None: - statement = statement.where( - normative.fidelity_event.c.direction == direction - ) - events = connection.execute(statement.order_by( - normative.fidelity_event.c.event_code - )).mappings().all() + + +def _read_fidelity_events( + connection: Any, *, dataset_id: str, direction: str | None = None, +) -> tuple[dict[str, Any], ...]: + normative = normative_catalog(MetaData()) + statement = ( + select(normative.fidelity_event) + .where(normative.fidelity_event.c.dataset_id == dataset_id) + .where(normative.fidelity_event.c.severity != "info") + ) + if direction is not None: + statement = statement.where( + normative.fidelity_event.c.direction == direction + ) + events = connection.execute(statement.order_by( + normative.fidelity_event.c.event_code + )).mappings().all() result = [] for item in events: details = json.loads(item["detail_json"] or "{}") @@ -1421,13 +1455,18 @@ def validate_wide_dataset( *, database_url: str, dataset_id: str, dolt_conformance_source: Any | None = None, ) -> dict[str, Any]: - profile, _active = read_profile( - database_url, dolt_conformance_source=dolt_conformance_source, - ) - dataset, variables, rows = read_wide_dataset( - database_url=database_url, dataset_id=dataset_id, profile=profile, + with _read_snapshot( + database_url=database_url, dolt_conformance_source=dolt_conformance_source, - ) + ) as (connection, profile): + dataset, variables, rows = _read_wide_dataset( + connection, dataset_id=dataset_id, profile=profile, + ) + reflected_table = Table( + dataset["data_table"], MetaData(), + schema=dataset["physical_table_schema"], + autoload_with=connection, + ) preflight( profile, variables, rows=rows, require_canonical_mapping=False, ) @@ -1439,12 +1478,6 @@ def validate_wide_dataset( if not variables: raise ValueError("A conforming dataset needs at least one source variable.") expected_columns = {"__case_ordinal", *(item["physical_name"] for item in variables)} - engine = create_engine(database_url) - reflected_table = Table( - dataset["data_table"], MetaData(), - schema=dataset["physical_table_schema"], - autoload_with=engine, - ) reflected_columns = {column.name: column for column in reflected_table.columns} actual_columns = set(reflected_columns) if actual_columns != expected_columns: diff --git a/tests/test_dolt_conformance.py b/tests/test_dolt_conformance.py index 7474a0c..9a68c83 100644 --- a/tests/test_dolt_conformance.py +++ b/tests/test_dolt_conformance.py @@ -255,27 +255,20 @@ def test_validate_wide_dataset_propagates_explicit_source( ) -> None: sentinel = object() calls: list[object] = [] - profile_calls: list[object] = [] - def capture_effective_profile( + def stop_after_read_preflight( _database_url: str, *, dolt_conformance_source: object, - ) -> tuple[object, dict[str, object]]: - profile_calls.append(dolt_conformance_source) - return object(), {} - - def stop_after_read_preflight(**kwargs: object) -> tuple[object, object, object]: - calls.append(kwargs["dolt_conformance_source"]) + ) -> None: + calls.append(dolt_conformance_source) raise UnsupportedOperationError("stop after propagation check") - monkeypatch.setattr(wide, "read_profile", capture_effective_profile) - monkeypatch.setattr(wide, "read_wide_dataset", stop_after_read_preflight) + monkeypatch.setattr(wide, "read_profile", stop_after_read_preflight) with pytest.raises(UnsupportedOperationError, match="propagation check"): wide.validate_wide_dataset( database_url="mysql+pymysql://example.invalid/catalog", dataset_id="synthetic", dolt_conformance_source=sentinel, ) - assert profile_calls == [sentinel] assert calls == [sentinel] @@ -362,11 +355,13 @@ def test_sav_export_propagates_explicit_source_to_read_gate( sentinel = object() calls: list[object] = [] - def stop_after_read(**kwargs: object) -> tuple[object, object, object]: - calls.append(kwargs["dolt_conformance_source"]) + def stop_after_read( + _database_url: str, *, dolt_conformance_source: object, + ) -> None: + calls.append(dolt_conformance_source) raise UnsupportedOperationError("stop after SAV read propagation") - monkeypatch.setattr(sav_module, "read_wide_dataset", stop_after_read) + monkeypatch.setattr(wide, "read_profile", stop_after_read) with pytest.raises(UnsupportedOperationError, match="SAV read propagation"): sav_module.export_sav_dataset( database_url="sqlite://", dataset_id="synthetic", diff --git a/tests/test_read_only_export.py b/tests/test_read_only_export.py index 95d8184..7ecdcda 100644 --- a/tests/test_read_only_export.py +++ b/tests/test_read_only_export.py @@ -5,6 +5,7 @@ python -m pytest tests/test_read_only_export.py The fixture creates and removes its own database and SELECT-only user. """ +import json import os import sqlite3 from uuid import uuid4 @@ -43,6 +44,9 @@ def seeded_catalog(request, tmp_path): with source_engine.connect() as connection: dataset = connection.execute(select(logical.dataset)).mappings().one() if request.param == "sqlite": + with sqlite3.connect(path) as connection: + connection.execute("PRAGMA journal_mode = WAL") + def snapshot(): with sqlite3.connect(path) as connection: return tuple(connection.iterdump()) @@ -119,7 +123,8 @@ def read_catalog(seeded_catalog): forbidden = [] def check(_connection, _cursor, statement, _parameters, _context, _many): - if statement.lstrip().split()[0].upper() not in {"SELECT", "SHOW", "PRAGMA", "DESCRIBE"}: + normalized = " ".join(statement.split()).upper() + if normalized != "BEGIN" and normalized.split()[0] not in {"SELECT", "SHOW", "PRAGMA", "DESCRIBE"}: forbidden.append(statement) raise AssertionError("Read operation attempted a non-read SQL statement") @@ -151,7 +156,7 @@ def test_export_and_reads_leave_no_database_trace(read_catalog, tmp_path, suffix assert not list(tmp_path.glob(".*.staging.*")) -@pytest.mark.parametrize("failure", ["loss", "writer", "publish", "restore", "staging", "backup"]) +@pytest.mark.parametrize("failure", ["read", "loss", "writer", "publish", "restore", "staging", "backup"]) def test_failed_export_never_writes_database(read_catalog, tmp_path, monkeypatch, failure): url, snapshot, _source = read_catalog before = snapshot() @@ -161,7 +166,16 @@ def test_failed_export_never_writes_database(read_catalog, tmp_path, monkeypatch def fail(*_args, **_kwargs): raise OSError("injected failure") - if failure == "loss": + readers = [] + + def fail_read(connection, _cursor, statement, _parameters, _context, _many): + if "FROM variable " in " ".join(statement.split()): + readers.append(connection) + raise OSError("injected read failure") + + if failure == "read": + event.listen(Engine, "after_cursor_execute", fail_read) + elif failure == "loss": monkeypatch.setattr(sav, "_export_loss_report", lambda *_a, **_k: ( {"code": "test-loss", "detail": "Requires consent", "details": {}}, )) @@ -188,8 +202,14 @@ def fail_backup(self, *args, **kwargs): return unlink(self, *args, **kwargs) monkeypatch.setattr(type(output), "unlink", fail_backup) - with pytest.raises((OSError, UnsupportedOperationError)): - openstatspec.export_sav(database_url=url, dataset_id="sample", destination=output) + try: + with pytest.raises((OSError, UnsupportedOperationError), match="injected read failure" if failure == "read" else None): + openstatspec.export_sav(database_url=url, dataset_id="sample", destination=output) + finally: + if failure == "read": + event.remove(Engine, "after_cursor_execute", fail_read) + if failure == "read": + assert readers and all(connection.closed for connection in readers) assert snapshot() == before if failure in {"restore", "backup"}: # If filesystem recovery itself fails, preserve the previous bytes in @@ -205,6 +225,117 @@ def fail_backup(self, *args, **kwargs): assert snapshot() == before +@pytest.mark.parametrize("seeded_catalog", ["sqlite"], indirect=True) +@pytest.mark.parametrize("operation", ["descriptor", "export"]) +def test_reads_use_one_snapshot(read_catalog, tmp_path, monkeypatch, operation): + url, snapshot, source = read_catalog + path = make_url(url).database + readers = [] + catalog_transactions = [] + committed = False + + def observe_catalog(connection, _cursor, statement, _parameters, _context, _many): + if "catalog_identity" in statement and connection not in readers: + readers.append(connection) + catalog_transactions.append(connection.connection.driver_connection.in_transaction) + + def interleave(connection, _cursor, statement, _parameters, _context, _many): + nonlocal committed + sql = " ".join(statement.split()) + if committed or "FROM variable " not in sql or "ORDER BY variable.source_ordinal" not in sql: + return + assert not connection.closed + # Raw connection deliberately bypasses SQLAlchemy's query-only listener. + with sqlite3.connect(path, timeout=1) as writer: + table, dataset_id, operation_id = writer.execute( + "SELECT physical_table_name, dataset_id, (SELECT operation_id FROM operation) FROM dataset" + ).fetchone() + writer.execute(f'UPDATE "{table}" SET answer = 3 WHERE __case_ordinal = 1') + writer.execute("UPDATE variable SET variable_label = 'New answer' WHERE source_name = 'answer'") + writer.execute("UPDATE value_label SET label = 'New code' WHERE numeric_code = 1") + writer.execute("UPDATE dataset SET dataset_label = 'New dataset'") + writer.execute( + "INSERT INTO fidelity_event VALUES (?, ?, ?, 'import', 'warning', 'snapshot-loss', NULL, '{}', CURRENT_TIMESTAMP)", + (str(uuid4()), operation_id, dataset_id), + ) + committed = True + assert not connection.closed # writer committed without waiting for reader release + + real_writer = sav._write_with_dictionary_bridge + + def write_after_read(*args, **kwargs): + assert readers and all(connection.closed for connection in readers) + return real_writer(*args, **kwargs) + + monkeypatch.setattr(sav, "_write_with_dictionary_bridge", write_after_read) + event.listen(Engine, "before_cursor_execute", observe_catalog) + event.listen(Engine, "after_cursor_execute", interleave) + try: + if operation == "descriptor": + dataset, variables, rows = wide.read_wide_dataset(database_url=url, dataset_id="sample") + assert rows[0]["answer"] == 1 + assert variables[0]["label"] == "Answer" + assert json.loads(variables[0]["value_labels"])["1.0"] == "Yes" + assert dataset["file_label"] == "" + else: + output = tmp_path / "snapshot.sav" + result = openstatspec.export_sav(database_url=url, dataset_id="sample", destination=output) + assert not result.diagnostics + expected, expected_meta = pyspssio.read_sav(str(source), include_user_missing=True) + actual, actual_meta = pyspssio.read_sav(str(output), include_user_missing=True) + pd.testing.assert_frame_equal(actual, expected) + for key in ("var_labels", "var_value_labels", "file_label"): + assert actual_meta[key] == expected_meta[key] + finally: + event.remove(Engine, "after_cursor_execute", interleave) + event.remove(Engine, "before_cursor_execute", observe_catalog) + assert committed + assert catalog_transactions and all(catalog_transactions), "first catalog read needs native BEGIN" + assert all(connection.closed for connection in readers) + after_writer = snapshot() + dataset, variables, rows = wide.read_wide_dataset(database_url=url, dataset_id="sample") + assert rows[0]["answer"] == 3 + assert variables[0]["label"] == "New answer" + assert json.loads(variables[0]["value_labels"])["1.0"] == "New code" + assert dataset["file_label"] == "New dataset" + if operation == "export": + with pytest.raises(UnsupportedOperationError, match="snapshot-loss"): + openstatspec.export_sav(database_url=url, dataset_id="sample", destination=output) + result = openstatspec.export_sav( + database_url=url, dataset_id="sample", destination=output, allow_loss=["snapshot-loss"], + ) + assert {d.code for d in result.diagnostics} == {"snapshot-loss"} + assert snapshot() == after_writer + + +@pytest.mark.parametrize("seeded_catalog", ["sqlite"], indirect=True) +def test_validation_reflects_the_same_schema_snapshot(read_catalog): + url, snapshot, _source = read_catalog + with sqlite3.connect(make_url(url).database) as connection: + table = connection.execute("SELECT physical_table_name FROM dataset").fetchone()[0] + readers = [] + + def rename_after_values(connection, _cursor, statement, _parameters, _context, _many): + if readers or f"FROM {table} " not in " ".join(statement.split()): + return + readers.append(connection) + with sqlite3.connect(make_url(url).database, timeout=1) as writer: + writer.execute("BEGIN") + writer.execute(f'ALTER TABLE "{table}" RENAME COLUMN answer TO renamed_answer') + writer.execute("UPDATE variable SET physical_name = 'renamed_answer' WHERE source_name = 'answer'") + assert not connection.closed + + event.listen(Engine, "after_cursor_execute", rename_after_values) + try: + assert openstatspec.validate(database_url=url, dataset_id="sample")["valid"] + finally: + event.remove(Engine, "after_cursor_execute", rename_after_values) + assert readers and all(connection.closed for connection in readers) + after_writer = snapshot() + assert openstatspec.validate(database_url=url, dataset_id="sample")["valid"] + assert snapshot() == after_writer + + def test_dolt_reads_do_not_enable_undeclared_writes(monkeypatch): url = "mysql+pymysql://reader@host/dataset" active = { diff --git a/tests/test_sql_services.py b/tests/test_sql_services.py index 1537435..b945899 100755 --- a/tests/test_sql_services.py +++ b/tests/test_sql_services.py @@ -1,12 +1,15 @@ """Real-service writes, round trips, failure boundaries, and candidate probes.""" +import json import os +from datetime import datetime, timezone from uuid import uuid4 import pandas as pd import pyspssio import pytest -from sqlalchemy import MetaData, create_engine, inspect as inspect_database, text +from sqlalchemy import MetaData, Table, create_engine, event, inspect as inspect_database, select, text +from sqlalchemy.engine import Engine from sqlalchemy.exc import DBAPIError import openstatspec @@ -17,6 +20,7 @@ delete_dataset_representation, ) from openstatspec.sql.wide import create_wide_dataset +from openstatspec.spss import sav from conformance import compare_sav_semantics, write_supported_semantics_fixture @@ -111,6 +115,160 @@ def test_live_profile_preserves_supported_sav_semantics(environment_name, datase ) assert compare_sav_semantics(source, destination) == {"equivalent": True, "differences": []} +@pytest.mark.parametrize("environment_name", [ + "OPENSTATSPEC_POSTGRES_URL", "OPENSTATSPEC_MYSQL_URL", + "OPENSTATSPEC_MARIADB_URL", "OPENSTATSPEC_DOLT_URL", +]) +def test_live_export_uses_one_snapshot(environment_name, source_sav, tmp_path, monkeypatch, record_property): + database_url = os.environ.get(environment_name) + if not database_url: + pytest.skip(f"{environment_name} is not configured") + imported = openstatspec.import_sav( + source_sav, database_url=database_url, dataset_id="snapshot_" + uuid4().hex, + ) + engine = create_engine(database_url) + tables = normative_catalog(MetaData()) + readers, isolations, transaction_checks = [], [], [] + committed = False + + def weaker_default(connection): + if connection.engine is not engine: + # Test-only DBAPI setup, before the reader can request its isolation. + # SQLAlchemy's MySQL setter itself emits SET SESSION ... and COMMIT + # outside cursor events; the event SQL allowlist must not allow SET. + connection.dialect.set_isolation_level( + connection.connection.driver_connection, "READ COMMITTED", + ) + assert connection.get_isolation_level() == "READ COMMITTED" + + def only_reads(connection, _cursor, statement, _parameters, _context, _many): + if connection.engine is engine: + return + sql = " ".join(statement.split()).upper() + assert sql == "BEGIN" or sql.split()[0] in {"SELECT", "SHOW", "DESCRIBE"}, statement + + def owned_contents(): + with engine.connect() as connection: + return ( + tuple(connection.execute(select(physical).order_by(physical.c.__case_ordinal))), + tuple(connection.execute(select(tables.variable).where( + tables.variable.c.dataset_id == imported["dataset_id"], + ).order_by(tables.variable.c.source_ordinal))), + tuple(connection.execute(select(tables.value_label).where( + tables.value_label.c.value_label_set_id.in_(select( + tables.value_label_set.c.value_label_set_id, + ).where(tables.value_label_set.c.dataset_id == imported["dataset_id"])), + ).order_by(tables.value_label.c.value_label_set_id, tables.value_label.c.ordinal))), + tuple(connection.execute(select(tables.fidelity_event).where( + tables.fidelity_event.c.dataset_id == imported["dataset_id"], + ).order_by(tables.fidelity_event.c.fidelity_event_id))), + ) + + after_writer = None + + def interleave(connection, _cursor, statement, _parameters, _context, _many): + nonlocal committed, after_writer + if connection.engine is engine: # never recurse through the writer + return + sql = " ".join(statement.split()) + if "FROM catalog_identity" in sql: + readers.append(connection) + raw = connection.connection.driver_connection + isolations.append(connection.dialect.get_isolation_level(raw)) + if connection.dialect.name == "postgresql": + transaction_checks.append(("native_transaction_active", raw.info.transaction_status != 0)) + elif environment_name == "OPENSTATSPEC_DOLT_URL": + # Dolt's explicit BEGIN supplies an OK packet with a valid IN_TRANS bit. + transaction_checks.append(("begin_in_trans", bool(raw.server_status & 1))) + else: + # PyMySQL updates server_status from OK, not SELECT EOF packets: + # IN_TRANS may still describe the isolation setter's COMMIT. + # Autocommit off implies a native txn at this real table SELECT; + # the interleaving below proves snapshot safety, not a measured bit. + transaction_checks.append(("autocommit_disabled", not raw.get_autocommit())) + if committed or "FROM variable " not in sql or "ORDER BY variable.source_ordinal" not in sql: + return + with engine.begin() as writer: + if writer.dialect.name in {"mysql", "mariadb"}: + writer.exec_driver_sql("BEGIN") # Dolt may retain DBAPI autocommit=1. + writer.execute(physical.update().where(physical.c.__case_ordinal == 1).values(age=35)) + writer.execute(tables.variable.update().where( + tables.variable.c.dataset_id == imported["dataset_id"], + tables.variable.c.source_name == "age", + ).values(variable_label="New age")) + writer.execute(tables.value_label.update().where( + tables.value_label.c.value_label_set_id.in_(select( + tables.value_label_set.c.value_label_set_id, + ).where(tables.value_label_set.c.dataset_id == imported["dataset_id"])), + ).values(label="New code")) + writer.execute(tables.fidelity_event.insert().values( + fidelity_event_id=str(uuid4()), operation_id=imported["operation_id"], + dataset_id=imported["dataset_id"], direction="import", severity="warning", + event_code="snapshot-loss", detail_json="{}", created_at=datetime.now(timezone.utc), + )) + committed = True + assert not connection.closed + after_writer = owned_contents() + + real_writer = sav._write_with_dictionary_bridge + + def write_after_read(*args, **kwargs): + assert readers and all(connection.closed for connection in readers) + return real_writer(*args, **kwargs) + + try: + physical = Table(imported["data_table"], MetaData(), autoload_with=engine) + with engine.connect() as connection: + record_property("server_version", connection.exec_driver_sql("SELECT VERSION()").scalar_one()) + monkeypatch.setattr(sav, "_write_with_dictionary_bridge", write_after_read) + event.listen(Engine, "engine_connect", weaker_default) + event.listen(Engine, "before_cursor_execute", only_reads) + event.listen(Engine, "after_cursor_execute", interleave) + try: + destination = tmp_path / "snapshot.sav" + result = openstatspec.export_sav( + database_url=database_url, dataset_id=imported["dataset_id"], destination=destination, + ) + finally: + event.remove(Engine, "after_cursor_execute", interleave) + event.remove(Engine, "before_cursor_execute", only_reads) + event.remove(Engine, "engine_connect", weaker_default) + record_property("catalog_isolations", isolations) + record_property("catalog_transaction_checks", transaction_checks) + record_property("writer_committed", committed) + assert committed + assert not result.diagnostics + expected, expected_meta = pyspssio.read_sav(str(source_sav), include_user_missing=True) + actual, actual_meta = pyspssio.read_sav(str(destination), include_user_missing=True) + pd.testing.assert_frame_equal(actual, expected) + for key in ("var_labels", "var_value_labels"): + assert actual_meta[key] == expected_meta[key] + assert isolations and set(isolations) == {"REPEATABLE READ"} + assert transaction_checks and all(passed for _, passed in transaction_checks), transaction_checks + assert all(connection.closed for connection in readers) + _, variables, rows = wide.read_wide_dataset( + database_url=database_url, dataset_id=imported["dataset_id"], + ) + assert rows[0]["age"] == 35 + assert variables[0]["label"] == "New age" + assert json.loads(variables[0]["value_labels"])["34.0"] == "New code" + with pytest.raises(UnsupportedOperationError, match="snapshot-loss"): + openstatspec.export_sav( + database_url=database_url, dataset_id=imported["dataset_id"], destination=destination, + ) + assert owned_contents() == after_writer + finally: + # Never reset a service/catalog or touch another dataset's objects. + with engine.begin() as connection: + delete_dataset_representation(connection, tables, imported["dataset_id"]) + quote = connection.dialect.identifier_preparer.quote + connection.exec_driver_sql(f"DROP TABLE {quote(imported['data_table'])}") + connection.execute(tables.operation.delete().where( + tables.operation.c.operation_id == imported["operation_id"], + )) + engine.dispose() + + def test_live_unknown_dolt_is_read_only_and_rejects_default_writes( tmp_path, ) -> None: