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
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
[project]
name = "lakebench"
version = "2.0.1"
version = "2.1.1"
authors = [
{ name="Miles Cole" },
]
Expand Down
12 changes: 12 additions & 0 deletions src/lakebench/engines/fabric_spark.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@ class FabricSpark(Spark):
Fabric Spark Engine
"""

_FAST_OPTIMIZE_CONFIG = "spark.microsoft.delta.optimize.fast.enabled"
_WRITE_STATS_CONFIGS = (
"spark.microsoft.delta.stats.collect.extended",
"spark.microsoft.delta.stats.injection.enabled",
Expand Down Expand Up @@ -139,3 +140,14 @@ def _configure_write_stats_collection(self) -> None:
for config_name in self._WRITE_STATS_CONFIGS:
self.spark.conf.set(config_name, config_value)
self.spark_configs[config_name] = config_value

def optimize_table(self, table_name: str):
fast_optimize_value = self.spark.conf.get(self._FAST_OPTIMIZE_CONFIG, None)
if str(fast_optimize_value).lower() != "true":
return super().optimize_table(table_name)

self.spark.conf.set(self._FAST_OPTIMIZE_CONFIG, "false")
try:
return super().optimize_table(table_name)
finally:
self.spark.conf.set(self._FAST_OPTIMIZE_CONFIG, fast_optimize_value)
64 changes: 62 additions & 2 deletions tests/test_fabric_spark.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,18 +4,46 @@


class _SparkConf:
def __init__(self):
def __init__(self, values=None):
self.initial_values = dict(values or {})
self.values = {}

def set(self, name, value):
self.values[name] = value

def get(self, name, default=None):
return self.initial_values.get(name, default)


class _Spark:
def __init__(self, conf_values=None, sql_error=None):
self.conf = _SparkConf(conf_values)
self.executed_statements = []
self.config_values_at_execution = []
self.sql_error = sql_error

def sql(self, statement):
self.executed_statements.append(statement)
self.config_values_at_execution.append(dict(self.conf.values))
if self.sql_error is not None:
raise self.sql_error


def _make_engine(collect_stats_on_write):
engine = object.__new__(FabricSpark)
engine.collect_stats_on_write = collect_stats_on_write
engine.spark_configs = {}
engine.spark = type("SparkStub", (), {"conf": _SparkConf()})()
engine.spark = _Spark()
return engine


def _make_optimize_engine(fast_optimize_value=None, sql_error=None):
engine = object.__new__(FabricSpark)
config_values = (
{FabricSpark._FAST_OPTIMIZE_CONFIG: fast_optimize_value} if fast_optimize_value is not None else None
)
engine.spark = _Spark(config_values, sql_error)
engine.full_catalog_schema_reference = "`lakehouse`.`schema`"
return engine


Expand All @@ -38,3 +66,35 @@ def test_write_stats_configs_are_set_and_logged(enabled, expected):

assert engine.spark.conf.values == {config_name: expected for config_name in FabricSpark._WRITE_STATS_CONFIGS}
assert engine.spark_configs == {config_name: expected for config_name in FabricSpark._WRITE_STATS_CONFIGS}


@pytest.mark.parametrize("fast_optimize_value", [None, "false", "FALSE"])
def test_optimize_table_leaves_disabled_fast_optimize_unchanged(fast_optimize_value):
engine = _make_optimize_engine(fast_optimize_value)

engine.optimize_table("store_sales")

assert engine.spark.executed_statements == ["OPTIMIZE `lakehouse`.`schema`.store_sales"]
assert engine.spark.config_values_at_execution == [{}]
assert engine.spark.conf.values == {}


def test_optimize_table_temporarily_disables_and_restores_fast_optimize():
engine = _make_optimize_engine("TRUE")

engine.optimize_table("store_sales")

assert engine.spark.executed_statements == ["OPTIMIZE `lakehouse`.`schema`.store_sales"]
assert engine.spark.config_values_at_execution == [{FabricSpark._FAST_OPTIMIZE_CONFIG: "false"}]
assert engine.spark.conf.values == {FabricSpark._FAST_OPTIMIZE_CONFIG: "TRUE"}


def test_optimize_table_restores_fast_optimize_when_optimize_fails():
error = RuntimeError("optimize failed")
engine = _make_optimize_engine("true", sql_error=error)

with pytest.raises(RuntimeError, match="optimize failed"):
engine.optimize_table("store_sales")

assert engine.spark.config_values_at_execution == [{FabricSpark._FAST_OPTIMIZE_CONFIG: "false"}]
assert engine.spark.conf.values == {FabricSpark._FAST_OPTIMIZE_CONFIG: "true"}
Loading
Loading