diff --git a/src/openstatspec/sql/inplace_transform.py b/src/openstatspec/sql/inplace_transform.py index 022d928..437286b 100644 --- a/src/openstatspec/sql/inplace_transform.py +++ b/src/openstatspec/sql/inplace_transform.py @@ -2,6 +2,8 @@ from __future__ import annotations +import hashlib + from collections.abc import Callable, Mapping from dataclasses import dataclass from datetime import UTC, datetime @@ -471,6 +473,7 @@ def _apply_plan_on_connection( dolt_branch: str | None, dolt_head: str | None, mutation_journal: dict[str, Any] | None = None, + encoded_plan: tuple[bytes, str] | None = None, ) -> dict[str, Any]: before_identity = _target_identity_state(connection, dataset_id, lock_dataset=True) if before_identity[3] != 1: @@ -792,6 +795,10 @@ def _apply_plan_on_connection( started = _now() if mutation_journal is not None: mutation_journal["apply_id"] = apply_id + if encoded_plan is None: + plan_bytes = plan.canonical_bytes() + encoded_plan = plan_bytes, hashlib.sha256(plan_bytes).hexdigest() + plan_bytes, plan_hash = encoded_plan connection.execute(insert(audit).values( apply_id=apply_id, contract_id=APPLY_CONTRACT, @@ -802,8 +809,8 @@ def _apply_plan_on_connection( source_kind=submission.source_kind, source_hash=submission.source_hash, frontend_contract=submission.frontend_contract, - plan_hash=plan.sha256(), - canonical_plan_json=plan.canonical_json(), + plan_hash=plan_hash, + canonical_plan_json=plan_bytes.decode("utf-8"), actor=actor, status="succeeded", dolt_branch=dolt_branch, @@ -837,7 +844,7 @@ def _apply_plan_on_connection( "source_kind": submission.source_kind, "source_hash": submission.source_hash, "frontend_contract": submission.frontend_contract, - "plan_hash": plan.sha256(), + "plan_hash": plan_hash, "dolt_branch": dolt_branch, "dolt_head_before": dolt_head, "dolt_head_after": dolt_head, @@ -978,6 +985,7 @@ def _run_in_place_submission( dataset_id: str, actor: str, prepare: Callable[[Any, str], InPlacePlanSubmission], + encoded_plan: tuple[bytes, str] | None = None, expected_branch: str | None = None, expected_head: str | None = None, dolt_conformance_source: DoltConformanceSource | None = None, @@ -1062,6 +1070,7 @@ def enable_sqlite_foreign_keys(dbapi_connection, _connection_record): dolt_branch=branch, dolt_head=head, mutation_journal=journal, + encoded_plan=encoded_plan, ) if profile.name == "dolt": after_branch, after_head, dirty_after = _dolt_state(connection) @@ -1116,7 +1125,8 @@ def apply_transformation_plan_in_place( if isinstance(plan, TransformationPlan) else transformation_plan_from_dict(plan) ) - plan_hash = normalized.sha256() + plan_bytes = normalized.canonical_bytes() + plan_hash = hashlib.sha256(plan_bytes).hexdigest() submission = InPlacePlanSubmission( plan=normalized, source_kind="canonical_plan", @@ -1127,6 +1137,7 @@ def apply_transformation_plan_in_place( dataset_id=dataset_id, actor=actor, prepare=lambda _connection, _dataset_id: submission, + encoded_plan=(plan_bytes, plan_hash), expected_branch=expected_branch, expected_head=expected_head, dolt_conformance_source=dolt_conformance_source, diff --git a/src/openstatspec/transform/validation.py b/src/openstatspec/transform/validation.py index 7a29000..93feea9 100644 --- a/src/openstatspec/transform/validation.py +++ b/src/openstatspec/transform/validation.py @@ -40,20 +40,16 @@ def _expected_type(storage_kind: StorageKind) -> ValueType: def _resolve( - variables: list[VariableDefinition], name: str -) -> tuple[int, VariableDefinition]: - matches = [ - (index, variable) - for index, variable in enumerate(variables) - if variable.name.casefold() == name.casefold() - ] - if len(matches) != 1: + variables: dict[str, VariableDefinition], name: str +) -> tuple[str, VariableDefinition]: + key = name.casefold() + if key not in variables: raise frontend_error( "unknown_variable", f"Variable {name!r} is not present in the current schema.", variable=name, ) - return matches[0] + return key, variables[key] def _validate_match(match: RecodeMatch, source: VariableDefinition) -> None: @@ -133,7 +129,7 @@ def _validate_recode_string_width( def _bind_recode( - operation: RecodeOperation, variables: list[VariableDefinition] + operation: RecodeOperation, variables: dict[str, VariableDefinition] ) -> None: _, source = _resolve(variables, operation.source) if operation.target_mode == "create": @@ -143,10 +139,8 @@ def _bind_recode( f"Target name {operation.target!r} is reserved.", target=operation.target, ) - if any( - variable.name.casefold() == operation.target.casefold() - for variable in variables - ): + target_key = operation.target.casefold() + if target_key in variables: raise frontend_error( "target_already_exists", f"Target name {operation.target!r} already exists.", @@ -183,16 +177,14 @@ def _bind_recode( variable=source.name, ) if operation.target_mode == "create": - variables.append( - VariableDefinition( - operation.target, - "numeric" if output_type == "binary64" else "string", - ) + variables[target_key] = VariableDefinition( + operation.target, + "numeric" if output_type == "binary64" else "string", ) def _operand_type( - operand: Operand, variables: list[VariableDefinition] + operand: Operand, variables: dict[str, VariableDefinition] ) -> ValueType: if operand.kind == "literal": assert operand.value is not None @@ -204,7 +196,7 @@ def _operand_type( def _validate_predicate( predicate: ComparisonExpression | BooleanExpression, - variables: list[VariableDefinition], + variables: dict[str, VariableDefinition], ) -> None: if isinstance(predicate, BooleanExpression): for item in predicate.operands: @@ -236,7 +228,7 @@ def _validate_predicate( def _bind_assign( - operation: AssignOperation, variables: list[VariableDefinition] + operation: AssignOperation, variables: dict[str, VariableDefinition] ) -> None: output_type = _operand_type(operation.value, variables) if output_type == "string": @@ -253,19 +245,17 @@ def _bind_assign( f"Target name {operation.target!r} is reserved.", target=operation.target, ) - if any( - variable.name.casefold() == operation.target.casefold() - for variable in variables - ): + target_key = operation.target.casefold() + if target_key in variables: raise frontend_error( "target_already_exists", f"Target name {operation.target!r} already exists.", target=operation.target, ) - variables.append(VariableDefinition( + variables[target_key] = VariableDefinition( operation.target, "numeric" if output_type == "binary64" else "string", - )) + ) return _, target = _resolve(variables, operation.target) if target.storage_kind == "string": @@ -284,7 +274,7 @@ def _bind_assign( def _bind_conditional_assign( operation: ConditionalAssignOperation, - variables: list[VariableDefinition], + variables: dict[str, VariableDefinition], ) -> None: _validate_predicate(operation.condition, variables) _, target = _resolve(variables, operation.target) @@ -328,17 +318,17 @@ def bind_transformation_plan( raise TypeError("plan must be a TransformationPlan.") if not isinstance(schema, VariableSchema): raise TypeError("schema must be a VariableSchema.") - variables = list(schema.variables) + variables = {variable.name.casefold(): variable for variable in schema.variables} + last_create = max( + (index for index, operation in enumerate(plan.operations) if _creates_variable(operation)), + default=-1, + ) for operation_index, operation in enumerate(plan.operations): - later_create = any( - _creates_variable(later_operation) - for later_operation in plan.operations[operation_index + 1:] - ) if isinstance(operation, CreateVariableOperation): _bind_create(operation, variables) continue if isinstance(operation, DeleteVariableOperation): - _bind_delete(operation, variables, allow_empty=later_create) + _bind_delete(operation, variables, allow_empty=operation_index < last_create) continue if isinstance(operation, RecodeOperation): _bind_recode(operation, variables) @@ -397,27 +387,28 @@ def bind_transformation_plan( if isinstance(operation, ExecuteOperation): continue raise AssertionError(f"Unknown plan operation: {type(operation)!r}") - return BoundTransformation(plan, VariableSchema(tuple(variables))) + return BoundTransformation(plan, VariableSchema(tuple(variables.values()))) def _bind_create( operation: CreateVariableOperation, - variables: list[VariableDefinition], + variables: dict[str, VariableDefinition], ) -> None: - if any(variable.name.casefold() == operation.variable.casefold() for variable in variables): + key = operation.variable.casefold() + if key in variables: raise frontend_error( "target_already_exists", f"Target name {operation.variable!r} already exists.", target=operation.variable, ) - variables.append(VariableDefinition( + variables[key] = VariableDefinition( operation.variable, operation.storage_kind, declared_string_width=operation.declared_string_width, - )) + ) def _bind_delete( operation: DeleteVariableOperation, - variables: list[VariableDefinition], + variables: dict[str, VariableDefinition], *, allow_empty: bool, ) -> None: diff --git a/tests/test_inplace_transform.py b/tests/test_inplace_transform.py index 11e0b32..fbecf6a 100644 --- a/tests/test_inplace_transform.py +++ b/tests/test_inplace_transform.py @@ -1,9 +1,11 @@ from __future__ import annotations +import hashlib import json import sqlite3 from dataclasses import replace from types import SimpleNamespace +from unittest.mock import patch import pytest from sqlalchemy import create_engine, inspect, text @@ -632,12 +634,25 @@ def test_temporary_target_type_is_not_taken_from_same_name_recreation(catalog) - def test_public_apply_supports_non_dolt_without_building_undo(catalog) -> None: url, path, dataset_id, table_name = catalog plan = _plan("RECODE score (1 = 0).") - result = openstatspec.apply_spss_in_place( - database_url=url, - dataset_id=dataset_id, - source_text="RECODE score (1 = 0).", - actor="test-agent", + expected_json = ( + '{"contract":"openstatspec-transformation-plan-v0.1","input_alias":"parent",' + '"operations":[{"op":"recode","rules":[{"match":{"kind":"values","values":' + '[{"bits":"3ff0000000000000","type":"binary64"}]},"result":{"kind":"literal",' + '"value":{"bits":"0000000000000000","type":"binary64"}}}],"source":"score",' + '"target":"score","target_mode":"replace","unmatched":{"kind":"copy"}}]}' ) + expected_hash = "c1a0025028816c7b424312612af4e0a95d00da6b1bcff25e9b742e4be25dbb42" + source_hash = "7e86162055308158fa2cf1781eae7d997c07bb870fc66becdabdfdd00580dc7b" + with patch.object( + openstatspec.TransformationPlan, "canonical_bytes", autospec=True, + side_effect=openstatspec.TransformationPlan.canonical_bytes, + ) as encodes, patch("hashlib.sha256", wraps=hashlib.sha256) as hashes: + result = openstatspec.apply_spss_in_place( + database_url=url, + dataset_id=dataset_id, + source_text="RECODE score (1 = 0).\r\n", + actor="test-agent", + ) assert result["dolt_branch"] is None assert result["dolt_commit_performed"] is False assert result["plan_hash"] == plan.sha256() @@ -651,6 +666,18 @@ def test_public_apply_supports_non_dolt_without_building_undo(catalog) -> None: "spss_syntax", "openstatspec-spss-syntax-frontend-v0.2", ) + assert (result["dataset_id"], result["physical_table_name"]) == (dataset_id, table_name) + assert (result["plan_hash"], result["source_hash"]) == (expected_hash, source_hash) + with sqlite3.connect(path) as connection: + audit_json, audit_hash, audit_source = connection.execute( + "SELECT canonical_plan_json, plan_hash, source_hash FROM transformation_apply" + ).fetchone() + assert (audit_json.encode("utf-8"), audit_hash, audit_source) == ( + expected_json.encode("utf-8"), expected_hash, source_hash, + ) + assert (encodes.call_count, sum( + call.args == (expected_json.encode("utf-8"),) for call in hashes.call_args_list + )) == (1, 1), "public SPSS apply must encode/hash the canonical plan once" def test_schema_commands_record_the_v03_frontend_contract(catalog) -> None: @@ -679,12 +706,24 @@ def test_public_generic_plan_apply_accepts_object_and_mapping( url, path, dataset_id, table_name = catalog plan = _plan("RECODE score (1 = 7).") supplied = plan.as_dict() if as_mapping else plan - result = openstatspec.apply_transformation_plan_in_place( - database_url=url, - dataset_id=dataset_id, - plan=supplied, - actor="test-agent", + expected_json = ( + '{"contract":"openstatspec-transformation-plan-v0.1","input_alias":"parent",' + '"operations":[{"op":"recode","rules":[{"match":{"kind":"values","values":' + '[{"bits":"3ff0000000000000","type":"binary64"}]},"result":{"kind":"literal",' + '"value":{"bits":"401c000000000000","type":"binary64"}}}],"source":"score",' + '"target":"score","target_mode":"replace","unmatched":{"kind":"copy"}}]}' ) + expected_hash = "089de9933b97157e9cb8b595d6f00afc11044907ca69e4ab747f6d0aac0bf60f" + with patch.object( + openstatspec.TransformationPlan, "canonical_bytes", autospec=True, + side_effect=openstatspec.TransformationPlan.canonical_bytes, + ) as encodes, patch("hashlib.sha256", wraps=hashlib.sha256) as hashes: + result = openstatspec.apply_transformation_plan_in_place( + database_url=url, + dataset_id=dataset_id, + plan=supplied, + actor="test-agent", + ) assert result["dataset_id"] == dataset_id assert result["physical_table_name"] == table_name assert result["source_kind"] == "canonical_plan" @@ -702,6 +741,15 @@ def test_public_generic_plan_apply_accepts_object_and_mapping( None, plan.sha256(), ) + audit_json, audit_hash = connection.execute( + "SELECT canonical_plan_json, plan_hash FROM transformation_apply" + ).fetchone() + assert (audit_json.encode("utf-8"), audit_hash, result["plan_hash"]) == ( + expected_json.encode("utf-8"), expected_hash, expected_hash, + ) + assert (encodes.call_count, sum( + call.args == (expected_json.encode("utf-8"),) for call in hashes.call_args_list + )) == (1, 1), "public canonical apply must encode/hash the plan once" def test_generic_string_width_is_rejected_before_ddl( diff --git a/tests/test_transform_frontend.py b/tests/test_transform_frontend.py index 1c73862..ed01ee2 100644 --- a/tests/test_transform_frontend.py +++ b/tests/test_transform_frontend.py @@ -14,7 +14,10 @@ spss_source_hash, ) from openstatspec.transform import ( + AssignOperation, CreateVariableOperation, + ExecuteOperation, + Operand, DeleteVariableOperation, RecodeMatch, RecodeOperation, @@ -680,6 +683,112 @@ def test_recode_copy_preserves_source_declared_width() -> None: assert bound.output_schema == schema +@pytest.mark.parametrize("workload", ["execute", "labels"]) +def test_generic_binding_10k_has_linear_bookkeeping(workload) -> None: + size = 10_000 + budget = 10 * (size + size) + probes = 0 + + class Name(str): + def casefold(self): + nonlocal probes + probes += 1 + assert probes <= budget, f"variable-name casefold probes: {probes} > {budget}" + return super().casefold() + + class Operations(tuple): + def __getitem__(self, key): + nonlocal probes + if isinstance(key, slice): + probes += len(range(*key.indices(len(self)))) + assert probes <= budget, f"operation suffix entries: {probes} > {budget}" + return super().__getitem__(key) + + schema = _schema(*(VariableDefinition(Name(f"V{i}"), "numeric") for i in range(size))) + operations = ( + Operations((ExecuteOperation(),) * size) if workload == "execute" + else tuple(SetVariableLabelOperation(f"v{i}", "L") for i in range(size)) + ) + plan = TransformationPlan(operations) + probes = 0 # Count only the public bind, not input construction. + bound = bind_transformation_plan(plan, schema) + assert bound.plan == plan + assert bound.output_schema == _schema(*( + VariableDefinition(f"V{i}", "numeric", variable_label="L" if workload == "labels" else None) + for i in range(size) + )) + + +@pytest.mark.parametrize("create", [ + CreateVariableOperation("replacement", "numeric"), + AssignOperation("replacement", "create", Operand("literal", value=TypedValue.binary64(1))), +], ids=["explicit", "assign"]) +def test_generic_distant_create_preserves_first_error_and_empty_schema_exception(create) -> None: + schema = _schema(VariableDefinition("Only", "numeric")) + prefix = (DeleteVariableOperation("oNLY"),) + (ExecuteOperation(),) * 100 + plan = TransformationPlan(prefix + (create,), contract="openstatspec-transformation-plan-v0.3") + assert bind_transformation_plan(plan, schema).output_schema == _schema( + VariableDefinition("replacement", "numeric"), + ) + plan = TransformationPlan( + prefix + (create, SetVariableLabelOperation("missing", "L"), create), + contract=plan.contract, + ) + with pytest.raises(TransformationFrontendError) as caught: + bind_transformation_plan(plan, schema) + assert caught.value.as_dict() == { + "code": "unknown_variable", + "detail": "Variable 'missing' is not present in the current schema.", + "details": {"variable": "missing"}, + } + + +@pytest.mark.parametrize("later", [(), (RecodeOperation( + source="missing", target="replacement", target_mode="create", + rules=(RecodeRule(RecodeMatch("values", (TypedValue.binary64(1),)), + RecodeResult("literal", TypedValue.binary64(0))),), + unmatched=RecodeResult("copy"), +),)], ids=["no-create", "recode-create"]) +def test_generic_last_delete_precedes_later_lookup_error(later) -> None: + plan = TransformationPlan( + (DeleteVariableOperation("oNLY"), SetVariableLabelOperation("missing", "L")) + + (ExecuteOperation(),) * 100 + later, + contract="openstatspec-transformation-plan-v0.3", + ) + with pytest.raises(TransformationFrontendError) as caught: + bind_transformation_plan(plan, _schema(VariableDefinition("Only", "numeric"))) + assert caught.value.as_dict() == { + "code": "cannot_delete_last_variable", + "detail": "A dataset must retain at least one variable.", + "details": {"variable": "Only"}, + } + + +def test_generic_delete_recreate_appends_with_new_spelling_and_metadata() -> None: + plan = TransformationPlan(( + DeleteVariableOperation("oNLY"), + CreateVariableOperation("ONLY", "string", 3), + SetVariableLabelOperation("only", "New"), + ), contract="openstatspec-transformation-plan-v0.3") + bound = bind_transformation_plan(plan, _schema( + VariableDefinition("Only", "numeric", variable_label="Old"), + VariableDefinition("Keep", "numeric"), + )) + assert bound.output_schema == _schema( + VariableDefinition("Keep", "numeric"), + VariableDefinition("ONLY", "string", variable_label="New", declared_string_width=3), + ) + + +def test_literal_unicode_canonical_plan_vector() -> None: + plan = TransformationPlan((SetVariableLabelOperation("Only", "Täpselt"),), input_alias="survey") + assert plan.canonical_bytes() == ( + '{"contract":"openstatspec-transformation-plan-v0.2","input_alias":"survey",' + '"operations":[{"label":"Täpselt","op":"set_variable_label","variable":"Only"}]}' + ).encode("utf-8") + assert plan.sha256() == "2ab3bf3c6dd2c21ce0fed076d43c49d93b6cc38e3e38529a18b7563286af9e03" + + def test_generic_plan_binding_validates_sequential_schema_state() -> None: plan = TransformationPlan(( RecodeOperation(