Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 15 additions & 4 deletions src/openstatspec/sql/inplace_transform.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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",
Expand All @@ -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,
Expand Down
73 changes: 32 additions & 41 deletions src/openstatspec/transform/validation.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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":
Expand All @@ -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.",
Expand Down Expand Up @@ -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
Expand All @@ -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:
Expand Down Expand Up @@ -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":
Expand All @@ -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":
Expand All @@ -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)
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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:
Expand Down
68 changes: 58 additions & 10 deletions tests/test_inplace_transform.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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()
Expand All @@ -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:
Expand Down Expand Up @@ -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"
Expand All @@ -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(
Expand Down
Loading