diff --git a/README.md b/README.md index 36f3e01..d7bb340 100644 --- a/README.md +++ b/README.md @@ -82,6 +82,30 @@ if country != "ALL": df.write.mode("overwrite").saveAsTable("silver.customers_enriched") ``` +## Spark write modes and schema options + +Table writes (`saveAsTable` / `insertInto`) stay CSV-backed. File writes (`csv` / `parquet` / `json` / `save`) use native Spark under `base_path`. `format("delta").save(path)` is stored as parquet (no Delta log). + +| Write | Missing table | Existing table | +|---|---|---| +| default / `overwrite` | create | replace rows | +| `append` | create | append rows (exact column set, unless `mergeSchema`) | +| `error` / `errorIfExists` | create | raise `AnalysisException` | +| `ignore` | create | no-op (file and temp view unchanged) | +| `insertInto` | raise `AnalysisException` | append, or replace when `overwrite=True` / `mode("overwrite")` | + +Schema flags: + +| Option | Effect | +|---|---| +| `overwriteSchema=true` + overwrite | replace the CSV even when columns change | +| `overwriteSchema=false` (default) + overwrite | raise `SchemaMismatchError` on incompatible schema change | +| `mergeSchema=true` + append | union missing columns with nulls | +| `mergeSchema=false` (default) + append | raise `SchemaMismatchError` if columns differ | +| same column, incompatible types | always raise `SchemaMismatchError` (merge only adds columns) | + +`partitionBy` columns must exist on the DataFrame. `option("replaceWhere", "")` with `mode("overwrite")` deletes matching stored rows then appends the new frame. `bucketBy` / `sortBy` are accepted no-ops (bucketing is not simulated). `df.writeTo(table).using(...).create()` / `.replace()` / `.append()` maps onto the same table writer; `createOrReplace` and `overwritePartitions` raise `NotImplementedError` with a `saveAsTable` hint. + ## Key Modules 1. `SparkProxy` - A Spark proxy that manipulates incoming Delta table reads and writes and redirects them to interactions with CSV files stored locally 2. `LocalWorkflowRunner` - A notebook orchestrator that takes the notebook .py files as defined in a Databricks Workflow JSON file and executes them as per the DAG definition. Databricks comment magics `%run` and `%sh` (`# %sh` / `# MAGIC %sh`, including `%sh -e`) work in those notebooks. diff --git a/docs/superpowers/specs/2026-09-03-spark-write-lakeflow-roadmap-design.md b/docs/superpowers/specs/2026-09-03-spark-write-lakeflow-roadmap-design.md index 5c91adf..3748935 100644 --- a/docs/superpowers/specs/2026-09-03-spark-write-lakeflow-roadmap-design.md +++ b/docs/superpowers/specs/2026-09-03-spark-write-lakeflow-roadmap-design.md @@ -42,14 +42,14 @@ Storage stays `{base_path}/{schema}/{table}.csv` + temp views `{schema}_{table}` ### Cluster 2 — `spark.write` fidelity (in order) -- [ ] **W1. Mode fidelity: `error` / `errorIfExists` + `ignore`** — `saveAsTable` with `error`/`errorIfExists` raises when CSV exists; `ignore` silently skips the write (no file touch, no view refresh). Few lines in `save_dataframe`; completes the Spark save-mode truth table alongside existing overwrite/append/default-overwrite. -- [ ] **W2. `insertInto(table, overwrite=False)`** — DataFrame API twin of append/overwrite `saveAsTable`; honor writer `mode` (`append` vs `overwrite`); require the table (CSV) to exist and raise a Spark-like AnalysisException message when missing; accept the `overwrite=True` kwarg for full-refresh semantics. -- [ ] **W3. `partitionBy` validation + `replaceWhere` dynamic overwrite** — validate `partitionBy` cols exist in the DataFrame (raise, don't silently ignore typos); support `.option("replaceWhere", "")` with `mode("overwrite")` as overwrite-where-predicate: delete matching rows from stored CSV via pandas query, append the new frame. Full partition-overwrite layout stays out of scope (still one CSV). -- [ ] **W4. CSV write options** — honor `delimiter`/`sep`, `quote`, `escape`, `nullValue`, `dateFormat`, `timestampFormat` on both `saveAsTable` (persist options per table for round-trip reads) and `csv(path)` passthrough; store the effective options so `read.table` round-trips without callers repeating them. -- [ ] **W5. File-write dispatch: `parquet` / `json` / `save` + `format().save()`** — route `df.write.parquet(path)` / `.json(path)` / `.save(path)` and `.format("parquet"|"json"|"csv"|"delta").save(path)` under `base_path` using native Spark writers; `format("delta").save(path)` maps to parquet-on-disk (documented, no Delta log). Table APIs stay CSV; file APIs use real formats. -- [ ] **W6. `overwriteSchema` / `mergeSchema` truth table** — `overwriteSchema=true` + overwrite = replace file even on schema change (already true, add tests); `mergeSchema=true` + append = union missing columns with nulls via pandas instead of raising `SchemaMismatchError`; `overwriteSchema=false` + incompatible change = raise. Document the matrix in README. -- [ ] **W7. `bucketBy` / `sortBy` accepted no-ops** — accept and ignore (like `partitionBy` today) with an explicit log/docstring that bucketing/sorting is not simulated; prevents AttributeError on production chains that call `.bucketBy(n, col).sortBy(col)`. -- [ ] **W8. `writeTo` (DataFrameWriterV2) decision** — either a minimal `writeTo(table).using(...).partitionedBy(...).option(...).create()/replace()/append()` façade over `save_dataframe`, or an explicit `NotImplementedError` with a migration hint to `saveAsTable`. Decide once real notebooks show which V2 verbs appear; do not build full V2 (overwritePartitions, createOrReplace) preemptively. +- [x] **W1. Mode fidelity: `error` / `errorIfExists` + `ignore`** — `saveAsTable` with `error`/`errorIfExists` raises when CSV exists; `ignore` silently skips the write (no file touch, no view refresh). Few lines in `save_dataframe`; completes the Spark save-mode truth table alongside existing overwrite/append/default-overwrite. +- [x] **W2. `insertInto(table, overwrite=False)`** — DataFrame API twin of append/overwrite `saveAsTable`; honor writer `mode` (`append` vs `overwrite`); require the table (CSV) to exist and raise a Spark-like AnalysisException message when missing; accept the `overwrite=True` kwarg for full-refresh semantics. +- [x] **W3. `partitionBy` validation + `replaceWhere` dynamic overwrite** — validate `partitionBy` cols exist in the DataFrame (raise, don't silently ignore typos); support `.option("replaceWhere", "")` with `mode("overwrite")` as overwrite-where-predicate: delete matching rows from stored CSV via pandas query, append the new frame. Full partition-overwrite layout stays out of scope (still one CSV). +- [x] **W4. CSV write options** — honor `delimiter`/`sep`, `quote`, `escape`, `nullValue`, `dateFormat`, `timestampFormat` on both `saveAsTable` (persist options per table for round-trip reads) and `csv(path)` passthrough; store the effective options so `read.table` round-trips without callers repeating them. +- [x] **W5. File-write dispatch: `parquet` / `json` / `save` + `format().save()`** — route `df.write.parquet(path)` / `.json(path)` / `.save(path)` and `.format("parquet"|"json"|"csv"|"delta").save(path)` under `base_path` using native Spark writers; `format("delta").save(path)` maps to parquet-on-disk (documented, no Delta log). Table APIs stay CSV; file APIs use real formats. +- [x] **W6. `overwriteSchema` / `mergeSchema` truth table** — `overwriteSchema=true` + overwrite = replace file even on schema change (already true, add tests); `mergeSchema=true` + append = union missing columns with nulls via pandas instead of raising `SchemaMismatchError`; `overwriteSchema=false` + incompatible change = raise. Document the matrix in README. +- [x] **W7. `bucketBy` / `sortBy` accepted no-ops** — accept and ignore (like `partitionBy` today) with an explicit log/docstring that bucketing/sorting is not simulated; prevents AttributeError on production chains that call `.bucketBy(n, col).sortBy(col)`. +- [x] **W8. `writeTo` (DataFrameWriterV2) decision** — either a minimal `writeTo(table).using(...).partitionedBy(...).option(...).create()/replace()/append()` façade over `save_dataframe`, or an explicit `NotImplementedError` with a migration hint to `saveAsTable`. Decide once real notebooks show which V2 verbs appear; do not build full V2 (overwritePartitions, createOrReplace) preemptively. ### Cluster 3 — missing `dbutils` APIs (in order) diff --git a/src/testbricks/catalog/table_catalog.py b/src/testbricks/catalog/table_catalog.py index ce21410..04963fa 100644 --- a/src/testbricks/catalog/table_catalog.py +++ b/src/testbricks/catalog/table_catalog.py @@ -2,18 +2,38 @@ from __future__ import annotations +import json import os +import re import tempfile from pathlib import Path from typing import Mapping, Optional import pandas as pd from pyspark.sql import DataFrame, SparkSession +from pyspark.sql.utils import AnalysisException from .errors import SchemaMismatchError from .identifier import TableIdentifier +_ERROR_MODES = frozenset({"error", "errorifexists"}) +_APPEND_MODES = frozenset({"append"}) +_OVERWRITE_MODES = frozenset({"overwrite"}) +_IGNORE_MODES = frozenset({"ignore"}) + _DEFAULT_READ_OPTIONS = {"header": "true", "inferSchema": "true"} +_CSV_OPTION_KEYS = frozenset( + { + "delimiter", + "sep", + "quote", + "escape", + "nullvalue", + "dateformat", + "timestampformat", + "header", + } +) class TableCatalog: @@ -23,6 +43,7 @@ def __init__(self, spark_session: SparkSession, base_path: str): self._spark = spark_session self._base_path = base_path self._root = Path(base_path) + self._csv_options: dict[str, dict[str, str]] = {} @property def base_path(self) -> str: @@ -31,6 +52,9 @@ def base_path(self) -> str: def path_for(self, ident: TableIdentifier) -> str: return str(self._root / ident.relative_csv_path) + def options_path_for(self, ident: TableIdentifier) -> str: + return str(self._root / ident.schema / f"{ident.table}.options.json") + def full_path(self, relative_path: str) -> str: return str(self._root / relative_path) @@ -42,6 +66,12 @@ def ensure_schema_dir(self, ident: TableIdentifier) -> str: def exists(self, ident: TableIdentifier) -> bool: return os.path.exists(self.path_for(ident)) + def csv_options_for(self, ident: TableIdentifier) -> dict[str, str]: + key = str(ident) + if key not in self._csv_options: + self._csv_options[key] = _load_options_file(self.options_path_for(ident)) + return dict(self._csv_options[key]) + def iter_schema_names(self) -> list[str]: if not self._root.exists(): return [] @@ -56,17 +86,20 @@ def iter_identifiers(self) -> list[TableIdentifier]: def load_all(self) -> None: for ident in self.iter_identifiers(): - self.read_csv(ident, _DEFAULT_READ_OPTIONS).createOrReplaceTempView( - ident.view_name - ) + self.read_csv(ident).createOrReplaceTempView(ident.view_name) def read_csv( self, ident: TableIdentifier, options: Optional[Mapping[str, str]] = None, ) -> DataFrame: + merged = { + **_DEFAULT_READ_OPTIONS, + **self.csv_options_for(ident), + **dict(options or {}), + } reader = self._spark.read - for key, value in (options or {}).items(): + for key, value in merged.items(): reader = reader.option(key, value) return reader.csv(self.path_for(ident)) @@ -76,36 +109,320 @@ def save_dataframe( dataframe: DataFrame, mode: Optional[str] = None, header: bool = True, + replace_where: Optional[str] = None, + csv_options: Optional[Mapping[str, str]] = None, + overwrite_schema: bool = False, + merge_schema: bool = False, ) -> None: self.ensure_schema_dir(ident) csv_path = self.path_for(ident) + exists = os.path.exists(csv_path) + save_mode = _normalize_save_mode(mode) + + if exists and save_mode in _ERROR_MODES: + raise AnalysisException( + f"[TABLE_OR_VIEW_ALREADY_EXISTS] Cannot create table or view " + f"{ident} because it already exists." + ) + if exists and save_mode in _IGNORE_MODES: + return + + stored = self.csv_options_for(ident) if exists else {} + incoming = _normalize_csv_options(csv_options) + if save_mode in _APPEND_MODES and exists: + effective_options = {**stored, **incoming} + else: + effective_options = {**incoming} + if header: + effective_options.setdefault("header", "true") + else: + effective_options["header"] = "false" + new_pdf = dataframe.toPandas() + new_pdf = _format_temporal_columns(new_pdf, effective_options) - if mode == "append" and os.path.exists(csv_path): - existing_pdf = pd.read_csv(csv_path) - if set(existing_pdf.columns) != set(new_pdf.columns): + if replace_where: + if save_mode != "overwrite": + raise ValueError( + "option('replaceWhere') requires mode('overwrite'); " + f"got mode '{mode}'." + ) + if exists: + existing_pdf = _pandas_read_csv(csv_path, stored or effective_options) + remaining = _apply_replace_where(existing_pdf, replace_where) + aligned = _align_columns_for_concat(remaining, new_pdf) + new_pdf = pd.concat(aligned, ignore_index=True) + elif exists and save_mode == "overwrite": + existing_pdf = _pandas_read_csv(csv_path, stored or effective_options) + if _schema_incompatible(existing_pdf, new_pdf) and not overwrite_schema: + raise SchemaMismatchError( + f"Cannot overwrite '{ident}' with an incompatible schema unless " + "overwriteSchema=true. " + f"Existing columns={list(existing_pdf.columns)}, " + f"new columns={list(new_pdf.columns)}" + ) + elif save_mode in _APPEND_MODES and exists: + existing_pdf = _pandas_read_csv(csv_path, stored or effective_options) + type_conflict = _overlapping_type_conflicts(existing_pdf, new_pdf) + column_mismatch = set(existing_pdf.columns) != set(new_pdf.columns) + if type_conflict: raise SchemaMismatchError( f"Cannot append to '{ident}': schema mismatch. " f"Existing columns={list(existing_pdf.columns)}, " f"new columns={list(new_pdf.columns)}" ) - new_pdf = pd.concat( - [existing_pdf, new_pdf[existing_pdf.columns]], - ignore_index=True, - ) + if column_mismatch and merge_schema: + aligned = _align_columns_for_concat(existing_pdf, new_pdf) + new_pdf = pd.concat(aligned, ignore_index=True) + elif column_mismatch: + raise SchemaMismatchError( + f"Cannot append to '{ident}': schema mismatch. " + f"Existing columns={list(existing_pdf.columns)}, " + f"new columns={list(new_pdf.columns)}" + ) + else: + new_pdf = pd.concat( + [existing_pdf, new_pdf[existing_pdf.columns]], + ignore_index=True, + ) + + self._write_csv_atomic(new_pdf, csv_path, header=header, options=effective_options) + self._persist_csv_options(ident, effective_options) + self._spark.createDataFrame(_nulls_for_spark(new_pdf)).createOrReplaceTempView( + ident.view_name + ) - self._write_csv_atomic(new_pdf, csv_path, header=header) - self._spark.createDataFrame(new_pdf).createOrReplaceTempView(ident.view_name) + def _persist_csv_options(self, ident: TableIdentifier, options: Mapping[str, str]) -> None: + payload = {key: str(value) for key, value in options.items()} + self._csv_options[str(ident)] = dict(payload) + options_path = self.options_path_for(ident) + with open(options_path, "w", encoding="utf-8") as handle: + json.dump(payload, handle, indent=2, sort_keys=True) @staticmethod - def _write_csv_atomic(pandas_df: pd.DataFrame, csv_path: str, header: bool = True) -> None: + def _write_csv_atomic( + pandas_df: pd.DataFrame, + csv_path: str, + header: bool = True, + options: Optional[Mapping[str, str]] = None, + ) -> None: directory = os.path.dirname(csv_path) fd, temp_path = tempfile.mkstemp(suffix=".csv", dir=directory) os.close(fd) try: - pandas_df.to_csv(temp_path, index=False, header=header) + pandas_df.to_csv(temp_path, **_pandas_write_kwargs(options, header=header)) os.replace(temp_path, csv_path) except Exception: if os.path.exists(temp_path): os.remove(temp_path) raise + + +def _normalize_save_mode(mode: Optional[str]) -> str: + """Map Spark save-mode aliases; default (None) overwrites, matching SparkProxy today.""" + if mode is None: + return "overwrite" + normalized = str(mode).strip().lower() + if normalized in _ERROR_MODES: + return "error" + if normalized in _APPEND_MODES: + return "append" + if normalized in _OVERWRITE_MODES: + return "overwrite" + if normalized in _IGNORE_MODES: + return "ignore" + raise ValueError( + f"Unknown save mode '{mode}'. Expected overwrite, append, ignore, " + "error, or errorIfExists." + ) + + +def _normalize_csv_options(options: Optional[Mapping[str, str]]) -> dict[str, str]: + if not options: + return {} + aliases = { + "sep": "delimiter", + "delimiter": "delimiter", + "quote": "quote", + "escape": "escape", + "nullvalue": "nullValue", + "dateformat": "dateFormat", + "timestampformat": "timestampFormat", + "header": "header", + } + normalized: dict[str, str] = {} + for key, value in options.items(): + canonical = aliases.get(key.lower()) + if canonical is None: + continue + normalized[canonical] = str(value) + return normalized + + +def _option_lookup(options: Optional[Mapping[str, str]], *names: str) -> Optional[str]: + if not options: + return None + lowered = {key.lower(): value for key, value in options.items()} + for name in names: + if name.lower() in lowered: + return str(lowered[name.lower()]) + return None + + +def _pandas_write_kwargs(options: Optional[Mapping[str, str]], header: bool) -> dict: + kwargs: dict = {"index": False, "header": header} + delimiter = _option_lookup(options, "delimiter", "sep") + if delimiter: + kwargs["sep"] = delimiter + quote = _option_lookup(options, "quote") + if quote: + kwargs["quotechar"] = quote + escape = _option_lookup(options, "escape") + if escape: + kwargs["escapechar"] = escape + kwargs["doublequote"] = False + null_value = _option_lookup(options, "nullValue") + if null_value is not None: + kwargs["na_rep"] = null_value + date_format = _option_lookup(options, "dateFormat", "timestampFormat") + if date_format: + kwargs["date_format"] = java_date_format_to_strftime(date_format) + return kwargs + + +def _pandas_read_csv(csv_path: str, options: Optional[Mapping[str, str]]) -> pd.DataFrame: + kwargs: dict = {} + delimiter = _option_lookup(options, "delimiter", "sep") + if delimiter: + kwargs["sep"] = delimiter + quote = _option_lookup(options, "quote") + if quote: + kwargs["quotechar"] = quote + escape = _option_lookup(options, "escape") + if escape: + kwargs["escapechar"] = escape + null_value = _option_lookup(options, "nullValue") + if null_value is not None: + kwargs["na_values"] = [null_value] + header = _option_lookup(options, "header") + if header and header.lower() == "false": + kwargs["header"] = None + return pd.read_csv(csv_path, **kwargs) + + +def _load_options_file(path: str) -> dict[str, str]: + if not os.path.exists(path): + return {} + with open(path, encoding="utf-8") as handle: + payload = json.load(handle) + if not isinstance(payload, dict): + return {} + return {str(key): str(value) for key, value in payload.items()} + + +def java_date_format_to_strftime(fmt: str) -> str: + result = fmt + for java_token, python_token in ( + ("yyyy", "%Y"), + ("SSS", "%f"), + ("yy", "%y"), + ("MM", "%m"), + ("dd", "%d"), + ("HH", "%H"), + ("mm", "%M"), + ("ss", "%S"), + ): + result = result.replace(java_token, python_token) + return result + + +def _format_value(value, strftime_fmt: str): + if value is None or (isinstance(value, float) and pd.isna(value)): + return value + if hasattr(value, "strftime"): + return value.strftime(strftime_fmt) + return value + + +def _format_temporal_columns(pdf: pd.DataFrame, options: Mapping[str, str]) -> pd.DataFrame: + date_fmt = _option_lookup(options, "dateFormat") + ts_fmt = _option_lookup(options, "timestampFormat") + if not date_fmt and not ts_fmt: + return pdf + formatted = pdf.copy() + for column in formatted.columns: + series = formatted[column] + if pd.api.types.is_datetime64_any_dtype(series): + has_time = bool((series.dt.hour.fillna(0) != 0).any() or (series.dt.minute.fillna(0) != 0).any()) + chosen = ts_fmt if (has_time and ts_fmt) else (date_fmt or ts_fmt) + formatted[column] = series.dt.strftime(java_date_format_to_strftime(chosen)) + continue + sample = series.dropna() + if sample.empty or not hasattr(sample.iloc[0], "strftime"): + continue + chosen = ts_fmt if ts_fmt and hasattr(sample.iloc[0], "hour") else (date_fmt or ts_fmt) + if not chosen: + continue + strftime_fmt = java_date_format_to_strftime(chosen) + formatted[column] = series.map(lambda value: _format_value(value, strftime_fmt)) + return formatted + + +def _spark_predicate_to_pandas_query(predicate: str) -> str: + query = predicate.strip() + query = query.replace("`", "") + query = re.sub(r"\bAND\b", "and", query, flags=re.IGNORECASE) + query = re.sub(r"\bOR\b", "or", query, flags=re.IGNORECASE) + query = re.sub(r"\bNOT\b", "not", query, flags=re.IGNORECASE) + query = re.sub(r"(?!=])=(?!=)", "==", query) + return query + + +def _apply_replace_where(existing_pdf: pd.DataFrame, predicate: str) -> pd.DataFrame: + query = _spark_predicate_to_pandas_query(predicate) + try: + matching = existing_pdf.query(query) + except Exception as exc: + raise ValueError( + f"Unparsable replaceWhere predicate: {predicate!r}" + ) from exc + return existing_pdf.drop(matching.index) + + +def _nulls_for_spark(pdf: pd.DataFrame) -> pd.DataFrame: + cleaned = pdf.copy() + for column in cleaned.columns: + cleaned[column] = cleaned[column].where(pd.notna(cleaned[column]), None) + return cleaned + + +def _dtype_family(series: pd.Series) -> str: + if pd.api.types.is_bool_dtype(series): + return "bool" + if pd.api.types.is_numeric_dtype(series): + return "numeric" + if pd.api.types.is_datetime64_any_dtype(series): + return "datetime" + return "string" + + +def _overlapping_type_conflicts(left: pd.DataFrame, right: pd.DataFrame) -> bool: + for column in set(left.columns) & set(right.columns): + left_series = left[column].dropna() + right_series = right[column].dropna() + if left_series.empty or right_series.empty: + continue + if _dtype_family(left_series) != _dtype_family(right_series): + return True + return False + + +def _schema_incompatible(left: pd.DataFrame, right: pd.DataFrame) -> bool: + if set(left.columns) != set(right.columns): + return True + return _overlapping_type_conflicts(left, right) + + +def _align_columns_for_concat(left: pd.DataFrame, right: pd.DataFrame) -> list[pd.DataFrame]: + columns = list(dict.fromkeys(list(left.columns) + list(right.columns))) + return [left.reindex(columns=columns), right.reindex(columns=columns)] diff --git a/src/testbricks/data_frame_wrapper.py b/src/testbricks/data_frame_wrapper.py index 216c330..076eeb9 100644 --- a/src/testbricks/data_frame_wrapper.py +++ b/src/testbricks/data_frame_wrapper.py @@ -1,8 +1,33 @@ +import logging + from pyspark.sql import DataFrame from pyspark.sql.group import GroupedData +from pyspark.sql.utils import AnalysisException from .catalog import TableIdentifier +logger = logging.getLogger(__name__) + + +_FILE_FORMATS = { + "parquet": "parquet", + "delta": "parquet", + "json": "json", + "csv": "csv", +} + + +def _resolve_file_format(fmt) -> str: + if fmt is None: + return "parquet" + resolved = _FILE_FORMATS.get(str(fmt).strip().lower()) + if resolved is None: + raise ValueError( + f"Unknown format '{fmt}' for file save(). " + "Supported formats: parquet, json, csv, delta (delta is stored as parquet)." + ) + return resolved + class _IoBuilder: """Shared format/option chaining used by both reader and writer.""" @@ -43,9 +68,40 @@ def __init__(self, spark_proxy, dataframe): self._dataframe = dataframe self._mode = None self._partition_by = () + self._bucket_by = None + self._sort_by = () def partitionBy(self, *cols): - self._partition_by = cols + flattened = [] + for col in cols: + if isinstance(col, (list, tuple)): + flattened.extend(col) + else: + flattened.append(col) + self._partition_by = tuple(flattened) + return self + + def bucketBy(self, numBuckets, *cols): + """Accepted no-op: Hive-style bucketing is not simulated locally.""" + flattened = [] + for col in cols: + if isinstance(col, (list, tuple)): + flattened.extend(col) + else: + flattened.append(col) + self._bucket_by = (numBuckets, tuple(flattened)) + logger.info( + "bucketBy(%s, %s) is accepted but not simulated", + numBuckets, + flattened, + ) + return self + + def sortBy(self, col, *cols): + """Accepted no-op: sortBy is not simulated locally.""" + flattened = [col, *cols] + self._sort_by = tuple(flattened) + logger.info("sortBy(%s) is accepted but not simulated", flattened) return self def mode(self, save_mode): @@ -53,21 +109,191 @@ def mode(self, save_mode): return self def csv(self, path): + self._native_writer().csv(self._spark._get_full_path(path)) + + def parquet(self, path): + self._native_writer().parquet(self._spark._get_full_path(path)) + + def json(self, path): + self._native_writer().json(self._spark._get_full_path(path)) + + def save(self, path, format=None, **options): + if options: + self._options.update(options) + fmt = format or self._format or "parquet" + resolved = _resolve_file_format(fmt) + full_path = self._spark._get_full_path(path) + writer = self._native_writer() + if resolved == "csv": + writer.csv(full_path) + elif resolved == "json": + writer.json(full_path) + else: + writer.parquet(full_path) + + def _native_writer(self): writer = self._dataframe.write if self._mode: writer = writer.mode(self._mode) + if self._partition_by: + writer = writer.partitionBy(*self._partition_by) for key, value in self._options.items(): writer = writer.option(key, value) - writer.csv(self._spark._get_full_path(path)) + return writer + + def _validate_partition_columns(self): + if not self._partition_by: + return + available = list(self._dataframe.columns) + missing = [col for col in self._partition_by if col not in available] + if missing: + raise AnalysisException( + f"partitionBy columns {missing} do not exist in the DataFrame. " + f"Available columns: {available}" + ) + + def _header_flag(self): + return str(self._options.get("header", "true")).lower() == "true" + + def _replace_where(self): + return self._options.get("replaceWhere") or self._options.get("replacewhere") + + def _csv_options(self): + return { + key: value + for key, value in self._options.items() + if key.lower() + in { + "delimiter", + "sep", + "quote", + "escape", + "nullvalue", + "dateformat", + "timestampformat", + "header", + } + } + + def _option_flag(self, *names): + wanted = {name.lower() for name in names} + for key, value in self._options.items(): + if key.lower() in wanted: + return str(value).lower() in {"true", "1", "yes"} + return False def saveAsTable(self, table_name): ident = TableIdentifier.parse(table_name) - header = self._options.get("header", "true").lower() == "true" + self._validate_partition_columns() self._spark._catalog.save_dataframe( ident, self._dataframe, mode=self._mode, - header=header, + header=self._header_flag(), + replace_where=self._replace_where(), + csv_options=self._csv_options(), + overwrite_schema=self._option_flag("overwriteSchema"), + merge_schema=self._option_flag("mergeSchema"), + ) + + def insertInto(self, table_name, overwrite=False): + """Append or overwrite rows in an existing table (Spark DataFrameWriter.insertInto).""" + ident = TableIdentifier.parse(table_name) + self._validate_partition_columns() + if not self._spark._catalog.exists(ident): + raise AnalysisException( + f"[TABLE_OR_VIEW_NOT_FOUND] The table or view {ident} cannot be found. " + "Verify the table exists before calling insertInto." + ) + writer_mode = str(self._mode).strip().lower() if self._mode else None + if overwrite or writer_mode == "overwrite": + mode = "overwrite" + else: + mode = "append" + self._spark._catalog.save_dataframe( + ident, + self._dataframe, + mode=mode, + header=self._header_flag(), + replace_where=self._replace_where(), + csv_options=self._csv_options(), + overwrite_schema=self._option_flag("overwriteSchema"), + merge_schema=self._option_flag("mergeSchema"), + ) + + +class DataFrameWriterV2: + """Minimal Spark DataFrameWriterV2 façade over ``saveAsTable``. + + Implements ``create`` / ``replace`` / ``append``. Full V2 verbs such as + ``createOrReplace`` and ``overwritePartitions`` raise ``NotImplementedError`` + with a migration hint to ``saveAsTable``. + """ + + def __init__(self, spark_proxy, dataframe, table_name): + self._spark = spark_proxy + self._dataframe = dataframe + self._table_name = table_name + self._options = {} + self._partitioned_by = () + self._using = None + + def using(self, provider): + self._using = provider + return self + + def option(self, key, value): + self._options[key] = value + return self + + def options(self, **kwargs): + self._options.update(kwargs) + return self + + def tableProperty(self, property, value): + return self + + def partitionedBy(self, *cols): + flattened = [] + for col in cols: + if isinstance(col, (list, tuple)): + flattened.extend(col) + else: + flattened.append(col) + self._partitioned_by = tuple(flattened) + return self + + def _writer(self, mode): + writer = DataFrameWriter(self._spark, self._dataframe) + writer._mode = mode + writer._format = self._using + writer._partition_by = self._partitioned_by + writer._options.update(self._options) + return writer + + def create(self): + self._writer("error").saveAsTable(self._table_name) + + def replace(self): + writer = self._writer("overwrite") + writer._options.setdefault("overwriteSchema", "true") + writer.saveAsTable(self._table_name) + + def append(self): + self._writer("append").saveAsTable(self._table_name) + + def createOrReplace(self): + raise NotImplementedError( + "writeTo(...).createOrReplace() is not implemented. " + "Use df.write.mode('overwrite').option('overwriteSchema', 'true')" + ".saveAsTable(...) instead." + ) + + def overwritePartitions(self): + raise NotImplementedError( + "writeTo(...).overwritePartitions() is not implemented. " + "Use df.write.mode('overwrite').option('replaceWhere', predicate)" + ".saveAsTable(...) instead." ) @@ -113,3 +339,6 @@ def write(self): if self._write is None: self._write = DataFrameWriter(self._spark, self._dataframe) return self._write + + def writeTo(self, table): + return DataFrameWriterV2(self._spark, self._dataframe, table) diff --git a/tests/test_basic.py b/tests/test_basic.py index ada0d0a..2030927 100644 --- a/tests/test_basic.py +++ b/tests/test_basic.py @@ -227,6 +227,287 @@ def test_save_as_table_append_schema_mismatch_raises(self, temp_spark): with pytest.raises(ValueError, match="schema mismatch"): second.write.mode("append").saveAsTable("default.people") + def test_save_as_table_error_mode_raises_when_table_exists(self, temp_spark): + from pyspark.sql.utils import AnalysisException + + first = _make_df(temp_spark, [("Alice", 30)], ["Name", "Age"]) + first.write.mode("overwrite").saveAsTable("default.people") + + second = _make_df(temp_spark, [("Bob", 25)], ["Name", "Age"]) + with pytest.raises(AnalysisException, match="default.people"): + second.write.mode("error").saveAsTable("default.people") + with pytest.raises(AnalysisException, match="already exists"): + second.write.mode("errorIfExists").saveAsTable("default.people") + + result = temp_spark.sql("SELECT * FROM default.people") + assert result.count() == 1 + assert result.collect()[0].Name == "Alice" + + def test_save_as_table_ignore_mode_skips_existing_table(self, temp_spark): + first = _make_df(temp_spark, [("Alice", 30)], ["Name", "Age"]) + first.write.mode("overwrite").saveAsTable("default.people") + csv_path = os.path.join(temp_spark._base_path, "default", "people.csv") + mtime_before = os.path.getmtime(csv_path) + + second = _make_df(temp_spark, [("Bob", 25)], ["Name", "Age"]) + second.write.mode("ignore").saveAsTable("default.people") + + assert os.path.getmtime(csv_path) == mtime_before + result = temp_spark.sql("SELECT * FROM default.people") + assert result.count() == 1 + assert result.collect()[0].Name == "Alice" + + def test_save_as_table_error_and_ignore_create_missing_table(self, temp_spark): + error_df = _make_df(temp_spark, [("Alice", 30)], ["Name", "Age"]) + error_df.write.mode("error").saveAsTable("default.from_error") + assert temp_spark.sql("SELECT * FROM default.from_error").count() == 1 + + ignore_df = _make_df(temp_spark, [("Bob", 25)], ["Name", "Age"]) + ignore_df.write.mode("IGNORE").saveAsTable("default.from_ignore") + assert temp_spark.sql("SELECT * FROM default.from_ignore").count() == 1 + + def test_save_as_table_unknown_mode_raises(self, temp_spark): + df = _make_df(temp_spark, [("Alice", 30)], ["Name", "Age"]) + with pytest.raises(ValueError, match="Unknown save mode"): + df.write.mode("upsert").saveAsTable("default.people") + + +class TestInsertInto: + def test_insert_into_appends_to_existing_table(self, temp_spark): + first = _make_df(temp_spark, [("Alice", 30)], ["Name", "Age"]) + first.write.mode("overwrite").saveAsTable("default.people") + + second = _make_df(temp_spark, [("Bob", 25)], ["Name", "Age"]) + second.write.insertInto("default.people") + + result = temp_spark.sql("SELECT * FROM default.people") + assert result.count() == 2 + assert {row.Name for row in result.collect()} == {"Alice", "Bob"} + + def test_insert_into_overwrite_kwarg_replaces_rows(self, temp_spark): + first = _make_df(temp_spark, [("Alice", 30), ("Bob", 25)], ["Name", "Age"]) + first.write.mode("overwrite").saveAsTable("default.people") + + second = _make_df(temp_spark, [("Charlie", 35)], ["Name", "Age"]) + second.write.insertInto("default.people", overwrite=True) + + result = temp_spark.sql("SELECT * FROM default.people") + assert result.count() == 1 + assert result.collect()[0].Name == "Charlie" + + def test_insert_into_honors_writer_overwrite_mode(self, temp_spark): + first = _make_df(temp_spark, [("Alice", 30)], ["Name", "Age"]) + first.write.mode("overwrite").saveAsTable("default.people") + + second = _make_df(temp_spark, [("Dana", 40)], ["Name", "Age"]) + second.write.mode("overwrite").insertInto("default.people") + + result = temp_spark.sql("SELECT * FROM default.people") + assert result.count() == 1 + assert result.collect()[0].Name == "Dana" + + def test_insert_into_missing_table_raises(self, temp_spark): + from pyspark.sql.utils import AnalysisException + + df = _make_df(temp_spark, [("Alice", 30)], ["Name", "Age"]) + with pytest.raises(AnalysisException, match="TABLE_OR_VIEW_NOT_FOUND"): + df.write.insertInto("default.missing") + + +class TestPartitionByAndReplaceWhere: + def test_partition_by_unknown_column_raises(self, temp_spark): + from pyspark.sql.utils import AnalysisException + + df = _make_df(temp_spark, [("Alice", 30, "2024-01-01")], ["Name", "Age", "dt"]) + with pytest.raises(AnalysisException, match="partitionBy columns"): + df.write.mode("overwrite").partitionBy("missing").saveAsTable("silver.people") + + def test_replace_where_overwrites_matching_rows_only(self, temp_spark): + first = _make_df( + temp_spark, + [("Alice", "2026-09-01"), ("Bob", "2026-09-02")], + ["Name", "dt"], + ) + first.write.mode("overwrite").partitionBy("dt").saveAsTable("silver.people") + + replacement = _make_df( + temp_spark, + [("Carol", "2026-09-01")], + ["Name", "dt"], + ) + replacement.write.format("delta").mode("overwrite").option( + "replaceWhere", "dt = '2026-09-01'" + ).partitionBy("dt").saveAsTable("silver.people") + + result = temp_spark.sql("SELECT * FROM silver.people") + rows = {(row.Name, row.dt) for row in result.collect()} + assert rows == {("Carol", "2026-09-01"), ("Bob", "2026-09-02")} + + def test_replace_where_unparsable_predicate_raises(self, temp_spark): + first = _make_df(temp_spark, [("Alice", "2026-09-01")], ["Name", "dt"]) + first.write.mode("overwrite").saveAsTable("silver.people") + + replacement = _make_df(temp_spark, [("Carol", "2026-09-01")], ["Name", "dt"]) + with pytest.raises(ValueError, match="Unparsable replaceWhere predicate"): + replacement.write.mode("overwrite").option( + "replaceWhere", "not a valid predicate !!!" + ).saveAsTable("silver.people") + + assert temp_spark.sql("SELECT * FROM silver.people").collect()[0].Name == "Alice" + + def test_replace_where_requires_overwrite_mode(self, temp_spark): + first = _make_df(temp_spark, [("Alice", "2026-09-01")], ["Name", "dt"]) + first.write.mode("overwrite").saveAsTable("silver.people") + replacement = _make_df(temp_spark, [("Carol", "2026-09-01")], ["Name", "dt"]) + with pytest.raises(ValueError, match="replaceWhere"): + replacement.write.mode("append").option( + "replaceWhere", "dt = '2026-09-01'" + ).saveAsTable("silver.people") + + +class TestCsvWriteOptions: + def test_save_as_table_pipe_delimiter_round_trips_on_read(self, temp_spark): + df = _make_df(temp_spark, [("Alice", 30), ("Bob", 25)], ["Name", "Age"]) + df.write.mode("overwrite").option("delimiter", "|").saveAsTable("default.people") + + csv_path = os.path.join(temp_spark._base_path, "default", "people.csv") + with open(csv_path, encoding="utf-8") as handle: + contents = handle.read() + assert "Alice|30" in contents + + result = temp_spark.read.table("default.people") + assert result.count() == 2 + assert {row.Name for row in result.collect()} == {"Alice", "Bob"} + + def test_save_as_table_null_value_round_trips(self, temp_spark): + df = temp_spark.createDataFrame([(None, 1), ("Alice", 2)], ["Name", "Age"]) + df.write.mode("overwrite").option("nullValue", "NA").saveAsTable("default.people") + + csv_path = os.path.join(temp_spark._base_path, "default", "people.csv") + with open(csv_path, encoding="utf-8") as handle: + contents = handle.read() + assert "NA" in contents + + result = temp_spark.read.table("default.people") + names = [row.Name for row in result.collect()] + assert None in names + assert "Alice" in names + + def test_save_as_table_date_format(self, temp_spark): + from datetime import date + + df = temp_spark.createDataFrame([(date(2026, 9, 1),)], ["dt"]) + df.write.mode("overwrite").option("dateFormat", "dd/MM/yyyy").saveAsTable( + "default.dates" + ) + csv_path = os.path.join(temp_spark._base_path, "default", "dates.csv") + with open(csv_path, encoding="utf-8") as handle: + contents = handle.read() + assert "01/09/2026" in contents + + result = temp_spark.read.table("default.dates") + assert result.count() == 1 + + @pytest.mark.skipif(sys.platform == "win32", reason="Native Spark CSV writer requires Hadoop winutils on Windows") + def test_csv_path_write_honors_delimiter(self, temp_spark): + df = _make_df(temp_spark, [(1, "a")], ["id", "name"]) + df.write.mode("overwrite").option("delimiter", "|").option("header", "true").csv( + "output/pipe" + ) + output_dir = os.path.join(temp_spark._base_path, "output", "pipe") + csv_files = [f for f in os.listdir(output_dir) if f.endswith(".csv")] + with open(os.path.join(output_dir, csv_files[0]), encoding="utf-8") as handle: + body = handle.read() + assert "|" in body + + +class TestFileWriteDispatch: + @pytest.mark.skipif(sys.platform == "win32", reason="Native Spark file writers require Hadoop winutils on Windows") + def test_parquet_write_is_readable(self, temp_spark): + df = _make_df(temp_spark, [("Alice", 30)], ["Name", "Age"]) + df.write.mode("overwrite").parquet("output/people_parquet") + path = os.path.join(temp_spark._base_path, "output", "people_parquet") + result = temp_spark._spark_session.read.parquet(path) + assert result.count() == 1 + assert result.collect()[0].Name == "Alice" + + @pytest.mark.skipif(sys.platform == "win32", reason="Native Spark file writers require Hadoop winutils on Windows") + def test_json_write_is_readable(self, temp_spark): + df = _make_df(temp_spark, [("Alice", 30)], ["Name", "Age"]) + df.write.mode("overwrite").json("output/people_json") + path = os.path.join(temp_spark._base_path, "output", "people_json") + result = temp_spark._spark_session.read.json(path) + assert result.count() == 1 + assert result.collect()[0].Name == "Alice" + + @pytest.mark.skipif(sys.platform == "win32", reason="Native Spark file writers require Hadoop winutils on Windows") + def test_format_delta_save_writes_parquet(self, temp_spark): + df = _make_df(temp_spark, [("Alice", 30)], ["Name", "Age"]) + df.write.format("delta").mode("overwrite").save("output/people_delta") + path = os.path.join(temp_spark._base_path, "output", "people_delta") + result = temp_spark._spark_session.read.parquet(path) + assert result.count() == 1 + parquet_files = [ + name for name in os.listdir(path) if name.endswith(".parquet") + ] + assert parquet_files + + def test_unknown_file_format_raises(self, temp_spark): + df = _make_df(temp_spark, [("Alice", 30)], ["Name", "Age"]) + with pytest.raises(ValueError, match="Unknown format"): + df.write.format("avro").save("output/people_avro") + + +class TestSchemaOptions: + def test_overwrite_schema_true_replaces_columns(self, temp_spark): + first = _make_df(temp_spark, [("Alice", 30)], ["Name", "Age"]) + first.write.mode("overwrite").saveAsTable("default.people") + + second = _make_df(temp_spark, [("Bob",)], ["Name"]) + second.write.mode("overwrite").option("overwriteSchema", "true").saveAsTable( + "default.people" + ) + result = temp_spark.sql("SELECT * FROM default.people") + assert result.columns == ["Name"] + assert result.collect()[0].Name == "Bob" + + def test_overwrite_schema_false_raises_on_column_change(self, temp_spark): + from testbricks.catalog import SchemaMismatchError + + first = _make_df(temp_spark, [("Alice", 30)], ["Name", "Age"]) + first.write.mode("overwrite").saveAsTable("default.people") + + second = _make_df(temp_spark, [("Bob",)], ["Name"]) + with pytest.raises(SchemaMismatchError, match="overwriteSchema"): + second.write.mode("overwrite").saveAsTable("default.people") + + def test_merge_schema_appends_missing_columns_as_nulls(self, temp_spark): + first = _make_df(temp_spark, [("Alice", 30)], ["Name", "Age"]) + first.write.mode("overwrite").saveAsTable("default.people") + + second = _make_df(temp_spark, [("Bob", 25, "UK")], ["Name", "Age", "Country"]) + second.write.mode("append").option("mergeSchema", "true").saveAsTable( + "default.people" + ) + result = temp_spark.sql("SELECT * FROM default.people") + assert result.count() == 2 + assert "Country" in result.columns + rows = {row.Name: row.Country for row in result.collect()} + assert rows["Alice"] is None + assert rows["Bob"] == "UK" + + def test_merge_schema_type_conflict_raises(self, temp_spark): + from testbricks.catalog import SchemaMismatchError + + first = _make_df(temp_spark, [("Alice", 30)], ["Name", "Age"]) + first.write.mode("overwrite").saveAsTable("default.people") + + second = temp_spark.createDataFrame([("Bob", "thirty")], ["Name", "Age"]) + with pytest.raises(SchemaMismatchError, match="schema mismatch"): + second.write.mode("append").option("mergeSchema", "true").saveAsTable( + "default.people" + ) + class TestWriteTransformedTable: def test_write_transformed_table_creates_expected_csv(self, spark): @@ -274,6 +555,41 @@ def test_writer_mode_option_chain(self, temp_spark): assert writer._partition_by == ("id",) assert writer._options == {"header": "true"} + def test_bucket_by_sort_by_chain_is_noop(self, temp_spark): + df = _make_df(temp_spark, [(1, "a")], ["id", "name"]) + writer = df.write.mode("overwrite").bucketBy(4, "id").sortBy("name") + assert writer._bucket_by == (4, ("id",)) + assert writer._sort_by == ("name",) + writer.saveAsTable("default.bucketed") + assert temp_spark.sql("SELECT * FROM default.bucketed").count() == 1 + + +class TestWriteTo: + def test_write_to_create_replace_and_append(self, temp_spark): + first = _make_df(temp_spark, [("Alice", 30)], ["Name", "Age"]) + first.writeTo("default.people").using("delta").partitionedBy("Age").create() + assert temp_spark.sql("SELECT * FROM default.people").count() == 1 + + from pyspark.sql.utils import AnalysisException + + with pytest.raises(AnalysisException, match="already exists"): + first.writeTo("default.people").create() + + replacement = _make_df(temp_spark, [("Bob",)], ["Name"]) + replacement.writeTo("default.people").option("header", "true").replace() + result = temp_spark.sql("SELECT * FROM default.people") + assert result.columns == ["Name"] + assert result.collect()[0].Name == "Bob" + + extra = _make_df(temp_spark, [("Carol",)], ["Name"]) + extra.writeTo("default.people").append() + assert temp_spark.sql("SELECT * FROM default.people").count() == 2 + + def test_write_to_create_or_replace_raises_with_hint(self, temp_spark): + df = _make_df(temp_spark, [("Alice", 30)], ["Name", "Age"]) + with pytest.raises(NotImplementedError, match="saveAsTable"): + df.writeTo("default.people").createOrReplace() + class TestSparkProxyLifecycle: def test_base_path_unchanged(self, spark):