From bfd281bff7d2b13305dbefd8503afe330f40bd3d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?T=C3=B5nis=20Ormisson?= Date: Wed, 9 Sep 2026 15:24:01 +0300 Subject: [PATCH] perf: reuse canonical import rows --- src/openstatspec/sql/wide.py | 14 ++--- tests/test_atomic_import.py | 117 +++++++++++++++++++++++++++++++++++ 2 files changed, 123 insertions(+), 8 deletions(-) diff --git a/src/openstatspec/sql/wide.py b/src/openstatspec/sql/wide.py index ac8ee64..1f0a72a 100644 --- a/src/openstatspec/sql/wide.py +++ b/src/openstatspec/sql/wide.py @@ -879,10 +879,8 @@ def create_wide_dataset( for item in variables ), ) - materialized = [ - {"__case_ordinal": ordinal, **row} - for ordinal, row in enumerate(source_rows, start=1) - ] + for ordinal, row in enumerate(source_rows, start=1): + row.setdefault("__case_ordinal", ordinal) docs_rows = document_rows(normative_dataset_id, documents) labels_rows = value_label_rows(normative_dataset_id, variables) missing_rows = missing_rule_rows(normative_dataset_id, variables) @@ -940,7 +938,7 @@ def create_wide_dataset( connection, normative, dataset_name=dataset_id, source_format=source_format, physical_table_name=data_table.name, dataset_label=file_label, source_encoding=source_encoding, - source_hash=source_sha256, source_case_count=len(materialized), + source_hash=source_sha256, source_case_count=len(source_rows), imported_at=imported_at or None, variables=variables, documents=docs_rows, value_labels=labels_rows, missing_rules=missing_rows, attributes=attributes_rows, @@ -963,9 +961,9 @@ def create_wide_dataset( },) if operation_details else ()), ), ) - if materialized: + if source_rows: for batch in _bounded_batches( - materialized, variables, profile.max_statement_bytes, + source_rows, variables, profile.max_statement_bytes, ): connection.execute(insert(data_table), batch) finish_normative_operation( @@ -1038,7 +1036,7 @@ def create_wide_dataset( "dataset_id": normative_dataset_id, "dataset_name": dataset_id, "data_table": data_table.name, - "case_count": len(materialized), + "case_count": len(source_rows), "operation_id": operation_id, } diff --git a/tests/test_atomic_import.py b/tests/test_atomic_import.py index 95bdb5a..f58a701 100755 --- a/tests/test_atomic_import.py +++ b/tests/test_atomic_import.py @@ -1,6 +1,8 @@ import sqlite3 from contextlib import contextmanager from dataclasses import replace +from decimal import Decimal +import struct import pytest @@ -9,6 +11,121 @@ from openstatspec.sql.wide import create_wide_dataset +def _reuse_variables() -> list[dict[str, object]]: + return [ + { + "ordinal": 1, "source_name": "score", "physical_name": "score", + "storage_kind": "numeric", "string_width": None, "label": "Score", + "format": "F8.2", "measure": "scale", "alignment": "right", + "display_width": 8, "value_labels": "{}", "missing_ranges": "[]", + }, + { + "ordinal": 2, "source_name": "label", "physical_name": "label", + "storage_kind": "string", "string_width": 16, "label": "Label", + "format": "A16", "measure": "nominal", "alignment": "left", + "display_width": 16, "value_labels": "{}", "missing_ranges": "[]", + }, + ] + + +def test_import_reuses_canonical_rows_after_preflight(tmp_path, monkeypatch) -> None: + database_path = tmp_path / "reuse.sqlite" + database = f"sqlite:///{database_path}" + input_rows = [ + {"score": Decimal("1.25"), "label": " Ä "}, + {"score": None, "label": ""}, + {"score": 0.1, "label": "trailing "}, + ] + original_rows = [dict(row) for row in input_rows] + canonical_rows = None + batches = [] + real_canonicalize = wide._canonicalize_database_numeric_rows + real_batches = wide._bounded_batches + + def capture_canonicalize(rows, variables): + nonlocal canonical_rows + canonical_rows = real_canonicalize(rows, variables) + return canonical_rows + + def capture_batches(rows, variables, maximum_statement_bytes): + assert rows is canonical_rows + assert all(actual is expected for actual, expected in zip(rows, canonical_rows, strict=True)) + for batch in real_batches(rows, variables, maximum_statement_bytes): + batches.append(batch) + yield batch + + monkeypatch.setattr( + wide, "_canonicalize_database_numeric_rows", capture_canonicalize, + ) + monkeypatch.setattr(wide, "_bounded_batches", capture_batches) + monkeypatch.setattr( + wide, "effective_profile", + lambda _url: (replace(SQLITE, max_statement_bytes=80), {}), + ) + + result = create_wide_dataset( + database_url=database, dataset_id="reuse", source_name="fixture.sav", + source_format="SAV", rows=(row for row in input_rows), + variables=_reuse_variables(), + ) + + assert input_rows == original_rows + assert result["case_count"] == len(input_rows) + assert len(batches) == len(input_rows) + connection = sqlite3.connect(database_path) + try: + rows = connection.execute( + "select __case_ordinal, score, label from data_reuse " + "order by __case_ordinal" + ).fetchall() + finally: + connection.close() + assert [row[0] for row in rows] == [1, 2, 3] + assert rows[0][1] == 1.25 and rows[0][2] == " Ä " + assert rows[1][1] is None and rows[1][2] == "" + assert struct.pack(">d", rows[2][1]) == struct.pack(">d", 0.1) + assert rows[2][2] == "trailing " + + +@pytest.mark.parametrize( + ("rows", "expected"), + [ + ( + [ + {"score": 0.1, "label": "ten", "__case_ordinal": 10}, + {"score": 2.0, "label": "two", "__case_ordinal": 2}, + ], + [(2, 2.0, "two"), (10, 0.1, "ten")], + ), + ([], []), + ], +) +def test_import_preserves_supplied_ordinals_and_empty_input( + tmp_path, rows, expected, +) -> None: + database_path = tmp_path / "ordinals.sqlite" + result = create_wide_dataset( + database_url=f"sqlite:///{database_path}", dataset_id="ordinals", + source_name="fixture.sav", source_format="SAV", rows=rows, + variables=_reuse_variables(), + ) + + connection = sqlite3.connect(database_path) + try: + actual = connection.execute( + "select __case_ordinal, score, label from data_ordinals " + "order by __case_ordinal" + ).fetchall() + case_count = connection.execute( + "select source_case_count from dataset where dataset_id = ?", + (result["dataset_id"],), + ).fetchone()[0] + finally: + connection.close() + assert actual == expected + assert case_count == len(rows) + + def test_invalid_string_row_leaves_no_dataset_or_data_table(tmp_path) -> None: database_path = tmp_path / "dataset.sqlite" database = f"sqlite:///{database_path}"