From ac53df57b8cc33c5d8654af05feb0cb92176e7fd Mon Sep 17 00:00:00 2001 From: atulikumwenayo Date: Tue, 29 Sep 2026 17:37:20 -0400 Subject: [PATCH] streaming local test --- src/datacustomcode/__init__.py | 10 + src/datacustomcode/io/cdf.py | 37 ++++ src/datacustomcode/io/reader/local_deltas.py | 184 ++++++++++++++++++ .../io/reader/streaming_seeder.py | 113 +++++++++++ src/datacustomcode/io/writer/local_deltas.py | 118 +++++++++++ src/datacustomcode/run.py | 30 ++- src/datacustomcode/template.py | 22 +++ tests/io/reader/test_local_deltas.py | 86 ++++++++ tests/io/reader/test_streaming_seeder.py | 83 ++++++++ tests/io/writer/test_local_deltas.py | 79 ++++++++ tests/test_template.py | 25 +++ 11 files changed, 786 insertions(+), 1 deletion(-) create mode 100644 src/datacustomcode/io/cdf.py create mode 100644 src/datacustomcode/io/reader/local_deltas.py create mode 100644 src/datacustomcode/io/reader/streaming_seeder.py create mode 100644 src/datacustomcode/io/writer/local_deltas.py create mode 100644 tests/io/reader/test_local_deltas.py create mode 100644 tests/io/reader/test_streaming_seeder.py create mode 100644 tests/io/writer/test_local_deltas.py diff --git a/src/datacustomcode/__init__.py b/src/datacustomcode/__init__.py index 4cd56f5..746a694 100644 --- a/src/datacustomcode/__init__.py +++ b/src/datacustomcode/__init__.py @@ -19,6 +19,8 @@ "Credentials", "DefaultSparkEinsteinPredictions", "DefaultSparkLLMGateway", + "LocalDeltasReader", + "LocalDeltasWriter", "PrintDataCloudWriter", "QueryAPIDataCloudReader", "SparkEinsteinPredictions", @@ -55,6 +57,14 @@ def __getattr__(name: str): from datacustomcode.io.reader.query_api import QueryAPIDataCloudReader return QueryAPIDataCloudReader + elif name == "LocalDeltasReader": + from datacustomcode.io.reader.local_deltas import LocalDeltasReader + + return LocalDeltasReader + elif name == "LocalDeltasWriter": + from datacustomcode.io.writer.local_deltas import LocalDeltasWriter + + return LocalDeltasWriter elif name == "SparkLLMGateway": from datacustomcode.llm_gateway import SparkLLMGateway diff --git a/src/datacustomcode/io/cdf.py b/src/datacustomcode/io/cdf.py new file mode 100644 index 0000000..15951e2 --- /dev/null +++ b/src/datacustomcode/io/cdf.py @@ -0,0 +1,37 @@ +# Copyright (c) 2025, Salesforce, Inc. +# SPDX-License-Identifier: Apache-2 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from __future__ import annotations + +from typing import Final + +from datacustomcode.io.writer.base import ( + MERGE_RECORD_TYPE_COLUMN, + MergeRecordType, +) + +COMMIT_VERSION: Final = "_commit_version" +COMMIT_TIMESTAMP: Final = "_commit_timestamp" +MERGE_RECORD_TYPE: Final = MERGE_RECORD_TYPE_COLUMN + +CDF_METADATA_COLUMNS: Final = (COMMIT_VERSION, COMMIT_TIMESTAMP, MERGE_RECORD_TYPE) + +__all__ = [ + "COMMIT_VERSION", + "COMMIT_TIMESTAMP", + "MERGE_RECORD_TYPE", + "CDF_METADATA_COLUMNS", + "MergeRecordType", +] diff --git a/src/datacustomcode/io/reader/local_deltas.py b/src/datacustomcode/io/reader/local_deltas.py new file mode 100644 index 0000000..a79b822 --- /dev/null +++ b/src/datacustomcode/io/reader/local_deltas.py @@ -0,0 +1,184 @@ +# Copyright (c) 2025, Salesforce, Inc. +# SPDX-License-Identifier: Apache-2 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from __future__ import annotations + +import json +import logging +import sys +from pathlib import Path +from typing import TYPE_CHECKING, Optional, Union + +from datacustomcode.config import config +from datacustomcode.io import cdf +from datacustomcode.io.reader.base import BaseDataCloudReader +from datacustomcode.io.reader.query_api import QueryAPIDataCloudReader + +if TYPE_CHECKING: + from pyspark.sql import DataFrame as PySparkDataFrame, SparkSession + from pyspark.sql.types import AtomicType, StructType + +logger = logging.getLogger(__name__) + + +class LocalDeltasReader(BaseDataCloudReader): + """Local reader for streaming transforms""" + + CONFIG_NAME = "LocalDeltasReader" + + def __init__( + self, + spark: SparkSession, + credentials_profile: str = "default", + dataspace: Optional[str] = None, + sf_cli_org: Optional[str] = None, + default_row_limit: Optional[int] = None, + fixtures_root: str = "payload/streaming_fixtures", + ) -> None: + super().__init__(spark) + self._fixtures_root = Path(fixtures_root) + # Credentials captured here and forwarded to StreamingSourceSeeder + # on the first `run` against a source (see _try_seed). + self._credentials_profile = credentials_profile + self._dataspace = dataspace + self._sf_cli_org = sf_cli_org + # Reuse the batch reader for INITIAL_SYNC / REBUILD reads. + self._batch = QueryAPIDataCloudReader( + spark=spark, + credentials_profile=credentials_profile, + dataspace=dataspace, + sf_cli_org=sf_cli_org, + default_row_limit=default_row_limit, + ) + # Which delta method was called; the seeder's LIMIT 1 probe + # needs to target the right layer. + self._current_layer = "dlo" + + def read_dlo( + self, + name: str, + schema: Union[AtomicType, StructType, str, None] = None, + ) -> PySparkDataFrame: + return self._batch.read_dlo(name, schema) + + def read_dmo( + self, + name: str, + schema: Union[AtomicType, StructType, str, None] = None, + ) -> PySparkDataFrame: + return self._batch.read_dmo(name, schema) + + def read_dlo_deltas(self) -> PySparkDataFrame: + self._current_layer = "dlo" + return self._open_stream(self._streaming_source()) + + def read_dmo_deltas(self) -> PySparkDataFrame: + self._current_layer = "dmo" + return self._open_stream(self._streaming_source()) + + def _streaming_source(self) -> str: + source = config.streaming_source + if not source: + raise RuntimeError( + "No streaming source configured. Set streamingSource.name in " + "config.json, or run `datacustomcode init --use-in-feature " + "StreamingTransform` to get started." + ) + return source + + def _open_stream(self, name: str) -> PySparkDataFrame: + drop_dir = self._fixtures_root / name + + if not (drop_dir / "_schema.json").exists(): + seeded = self._try_seed(name) + if not seeded: + print( + f"\nStreaming source {self._current_layer}='{name}' is empty.\n" + f" Populate the stream source with at least one row in Data Cloud, " + f"then re-run `datacustomcode run`.\n" + ) + sys.exit(0) + + print( + f"\nWrote sample streaming fixture in: {drop_dir} " + f"based on the contents of the streaming source {name}.\n" + f" You may add more JSON files alongside it to simulate additional " + f"changes.\n" + ) + + schema = self._build_stream_schema(name) + return ( + self.spark.readStream + .format("json") + .schema(schema) + .option("maxFilesPerTrigger", 1) # one file = one batch + .option("latestFirst", "false") # oldest mtime first + .load(str(drop_dir)) + ) + + def _build_stream_schema(self, name: str) -> "StructType": + """Compose (source schema + CDF metadata columns) for readStream.""" + from pyspark.sql.types import ( + LongType, + StringType, + StructField, + StructType, + TimestampType, + ) + + source_schema = self._resolve_source_schema(name) + cdf_fields = [ + StructField(cdf.COMMIT_VERSION, LongType(), True), + StructField(cdf.COMMIT_TIMESTAMP, TimestampType(), True), + StructField(cdf.MERGE_RECORD_TYPE, StringType(), True), + ] + return StructType(list(source_schema.fields) + cdf_fields) + + def _resolve_source_schema(self, name: str) -> "StructType": + from pyspark.sql.types import StructType + + schema_file = self._fixtures_root / name / "_schema.json" + return StructType.fromJson(json.loads(schema_file.read_text())) + + def _try_seed(self, name: str) -> bool: + """Fetch schema + snapshot from the tenant for this source. + + Called by _open_stream when _schema.json is missing (first run + against `name`, or the customer deleted the cache to force a + fresh seed). The `_` prefix is a Hadoop hidden-file convention; + Spark's FileStreamSource skips it so the schema cache doesn't + get delivered as a micro-batch alongside 000_seed.json. + + Returns True when fixtures were written, False when the source + is empty (caller sys.exit(0)s with a friendly message). Query + API errors propagate as RuntimeError. + """ + # Deferred import: batch-only runs never load the seeder module. + from datacustomcode.io.reader.streaming_seeder import ( + StreamingSourceSeeder, + ) + + seeder = StreamingSourceSeeder( + spark=self.spark, + credentials_profile=self._credentials_profile, + dataspace=self._dataspace, + sf_cli_org=self._sf_cli_org, + ) + # self._current_layer was set by read_dlo_deltas / read_dmo_deltas + # right before _open_stream — pass it through so the seeder's + # LIMIT 1 probe targets the correct read method. + return seeder.seed_source( + name, self._current_layer, str(self._fixtures_root) + ) diff --git a/src/datacustomcode/io/reader/streaming_seeder.py b/src/datacustomcode/io/reader/streaming_seeder.py new file mode 100644 index 0000000..3d9a6ff --- /dev/null +++ b/src/datacustomcode/io/reader/streaming_seeder.py @@ -0,0 +1,113 @@ +# Copyright (c) 2025, Salesforce, Inc. +# SPDX-License-Identifier: Apache-2 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from __future__ import annotations + +import json +import math +from datetime import datetime, timezone +from pathlib import Path +from typing import Any, TYPE_CHECKING, Optional + +from datacustomcode.io import cdf +from datacustomcode.io.reader.query_api import QueryAPIDataCloudReader + +if TYPE_CHECKING: + from pyspark.sql import SparkSession + +SEED_LIMIT = 10 + + +def _clean_for_json(value: Any) -> Any: + if isinstance(value, float): + if math.isnan(value): + return None + if value.is_integer(): + return int(value) + return value + + +class StreamingSourceSeeder: + def __init__( + self, + spark: SparkSession, + credentials_profile: str = "default", + dataspace: Optional[str] = None, + sf_cli_org: Optional[str] = None, + ) -> None: + self.spark = spark + self.reader = QueryAPIDataCloudReader( + spark=spark, + credentials_profile=credentials_profile, + dataspace=dataspace, + sf_cli_org=sf_cli_org, + ) + + def seed_source( + self, name: str, layer: str, fixtures_root: str + ) -> bool: + """Seed streaming fixtures and source schema. + + Returns: + ``True`` when schema + fixtures are written + ``False`` when the source is empty. + Query API errors (missing creds, missing + source) propagate as ``RuntimeError``. + """ + read_fn = self.reader.read_dlo if layer == "dlo" else self.reader.read_dmo + + try: + head_df = read_fn(name).limit(1) + except Exception as exc: + raise RuntimeError( + f"Failed to read {layer}='{name}': {exc}." + ) from exc + + head_pandas = head_df.toPandas() + if len(head_pandas) == 0: + return False + + schema = head_df.schema + out_dir = Path(fixtures_root) / name + out_dir.mkdir(parents=True, exist_ok=True) + schema_file_path = out_dir / "_schema.json" + schema_file_path.write_text(json.dumps(schema.jsonValue())) + print( + f"\nStreaming source schema written to {schema_file_path} " + f"({len(schema.fields)} fields)" + ) + + # Mix UPSERT and DELETE in the seed so the customer sees both + # operation types in the starter fixture + snapshot = read_fn(name).limit(SEED_LIMIT).toPandas() + seed_rows = [] + for i, record in enumerate(snapshot.to_dict("records")): + op = ( + cdf.MergeRecordType.DELETE + if i % 2 + else cdf.MergeRecordType.UPSERT + ) + + cleaned = {k: _clean_for_json(v) for k, v in record.items()} + seed_rows.append({ + **cleaned, + cdf.COMMIT_VERSION: i + 1, + cdf.COMMIT_TIMESTAMP: datetime.now(timezone.utc).isoformat(), + cdf.MERGE_RECORD_TYPE: op.value, + }) + (out_dir / "000_seed.json").write_text( + "\n".join(json.dumps(r, default=str) for r in seed_rows) + ) + return True diff --git a/src/datacustomcode/io/writer/local_deltas.py b/src/datacustomcode/io/writer/local_deltas.py new file mode 100644 index 0000000..086d65b --- /dev/null +++ b/src/datacustomcode/io/writer/local_deltas.py @@ -0,0 +1,118 @@ +# Copyright (c) 2025, Salesforce, Inc. +# SPDX-License-Identifier: Apache-2 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from __future__ import annotations + +import tempfile +from typing import TYPE_CHECKING, Optional + +from datacustomcode.io import cdf +from datacustomcode.io.writer.base import BaseDataCloudWriter, WriteMode + +if TYPE_CHECKING: + from pyspark.sql import DataFrame as PySparkDataFrame, SparkSession + from pyspark.sql.streaming import StreamingQuery + +_HEADER_WARNING = ( + " NOTE: Once this code extension is deployed, Data Cloud will persist " + " only the final state per primary key.\n" + " This local preview shows every emitted row uncollapsed." +) + + +class LocalDeltasWriter(BaseDataCloudWriter): + """Local writer for streaming transforms.""" + + CONFIG_NAME = "LocalDeltasWriter" + + def __init__( + self, + spark: SparkSession, + credentials_profile: str = "default", + dataspace: Optional[str] = None, + sf_cli_org: Optional[str] = None, + ) -> None: + super().__init__(spark) + + from datacustomcode.io.writer.print import PrintDataCloudWriter + self._batch_writer = PrintDataCloudWriter( + spark=spark, + credentials_profile=credentials_profile, + dataspace=dataspace, + sf_cli_org=sf_cli_org, + ) + + def write_to_dlo( + self, name: str, dataframe: PySparkDataFrame, write_mode: WriteMode + ) -> None: + return self._batch_writer.write_to_dlo(name, dataframe, write_mode) + + def write_to_dmo( + self, name: str, dataframe: PySparkDataFrame, write_mode: WriteMode + ) -> None: + return self._batch_writer.write_to_dmo(name, dataframe, write_mode) + + def auto_write_to_dlo( + self, name: str, dataframe: PySparkDataFrame + ) -> None: + return self._batch_writer.auto_write_to_dlo(name, dataframe) + + def auto_write_to_dmo( + self, name: str, dataframe: PySparkDataFrame + ) -> None: + return self._batch_writer.auto_write_to_dmo(name, dataframe) + + def write_dlo_deltas( + self, name: str, dataframe: PySparkDataFrame, **kwargs + ) -> StreamingQuery: + if not name: + raise ValueError("DLO name must be provided.") + + def _preview(batch_df, batch_id): + from pyspark.sql import functions as F + + output_batch = batch_df.withColumn( + "_operation", + F.when( + F.col(cdf.MERGE_RECORD_TYPE) + == cdf.MergeRecordType.DELETE.value, + F.lit("DELETE"), + ).otherwise(F.lit("UPSERT")), + ).drop( + cdf.COMMIT_VERSION, + cdf.COMMIT_TIMESTAMP, + cdf.MERGE_RECORD_TYPE, + ) + + row_count = output_batch.count() + print( + f"\nTarget={name} batch_id={batch_id} " + f"rows={row_count}" + ) + print(_HEADER_WARNING) + output_batch.show(truncate=False) + + checkpoint_dir = tempfile.mkdtemp(prefix=f"local-deltas-ckpt-{name}-") + return ( + dataframe.writeStream + .foreachBatch(_preview) + .option("checkpointLocation", checkpoint_dir) + # AvailableNow: drain every fixture file then + # terminate. The user's `query.awaitTermination()` + # returns without needing a timeout, so + # `datacustomcode run` exits deterministically. + .trigger(availableNow=True) + .start() + ) diff --git a/src/datacustomcode/run.py b/src/datacustomcode/run.py index f615af5..1bc5696 100644 --- a/src/datacustomcode/run.py +++ b/src/datacustomcode/run.py @@ -52,6 +52,27 @@ def _read_streaming_source(config_json: dict) -> Optional[str]: return str(name) if name else None +def _project_config_yaml(entrypoint: str) -> Optional[str]: + """Locate a project-local ``config.yaml`` next to ``payload/``. + + A streaming project scaffolded by + ``datacustomcode init --use-in-feature StreamingTransform`` has + ``config.yaml`` at the project root. + + Layout:: + + / + config.yaml <-- this file (streaming profile) + payload/ + entrypoint.py <-- passed as `entrypoint` + config.json + """ + entrypoint_dir = os.path.dirname(os.path.abspath(entrypoint)) + project_root = os.path.dirname(entrypoint_dir) + candidate = os.path.join(project_root, "config.yaml") + return candidate if os.path.exists(candidate) else None + + def _update_config_options(profile: Optional[str], sf_cli_org: Optional[str]): if sf_cli_org: config_key = "sf_cli_org" @@ -133,9 +154,16 @@ def run_entrypoint( f"Please ensure config.json contains a 'dataspace' field." ) - # Load config file first + # Load config file first. Precedence: + # 1. Explicit --config-file argument (highest) + # 2. Project-local /config.yaml + # 3. SDK-shipped default if config_file: config.load(config_file) + else: + project_config = _project_config_yaml(entrypoint) + if project_config: + config.load(project_config) # Add dataspace to reader and writer config options _set_config_option(config.reader_config, "dataspace", dataspace) diff --git a/src/datacustomcode/template.py b/src/datacustomcode/template.py index a543575..db72bd7 100644 --- a/src/datacustomcode/template.py +++ b/src/datacustomcode/template.py @@ -27,6 +27,10 @@ script_template_dir, "examples", "streaming_deltas", "entrypoint.py" ) +_SDK_CONFIG_YAML = os.path.join( + os.path.dirname(__file__), "config.yaml" +) + def copy_script_template(target_dir: str, streaming: bool = False) -> None: """Copy the template to the target directory.""" @@ -51,6 +55,24 @@ def copy_script_template(target_dir: str, streaming: bool = False) -> None: ) shutil.copy2(STREAMING_EXAMPLE_ENTRYPOINT, destination) + _write_streaming_config_yaml(target_dir) + + +def _write_streaming_config_yaml(target_dir: str) -> None: + """Writes the sdk config file with Streaming overrdies""" + import yaml + + with open(_SDK_CONFIG_YAML) as f: + config_data = yaml.safe_load(f) + + config_data["reader_config"]["type_config_name"] = "LocalDeltasReader" + config_data["writer_config"]["type_config_name"] = "LocalDeltasWriter" + + destination = os.path.join(target_dir, "config.yaml") + logger.debug(f"Writing streaming config.yaml to {destination}...") + with open(destination, "w") as f: + yaml.safe_dump(config_data, f, sort_keys=False) + def copy_function_template(target_dir: str, use_in_feature: Optional[str]) -> None: os.makedirs(target_dir, exist_ok=True) diff --git a/tests/io/reader/test_local_deltas.py b/tests/io/reader/test_local_deltas.py new file mode 100644 index 0000000..7ecaa24 --- /dev/null +++ b/tests/io/reader/test_local_deltas.py @@ -0,0 +1,86 @@ +from __future__ import annotations + +import json +from unittest.mock import MagicMock, patch + +import pytest +from pyspark.sql.types import ( + LongType, + StringType, + StructField, + StructType, +) + +from datacustomcode.io import cdf +from datacustomcode.io.reader.local_deltas import LocalDeltasReader + + +SOURCE_SCHEMA = StructType( + [ + StructField("id__c", StringType(), True), + StructField("age__c", LongType(), True), + ] +) + +@pytest.fixture +def reader(tmp_path): + """Reader with __init__ short-circuited (avoids Query API construction).""" + with patch.object(LocalDeltasReader, "__init__", return_value=None): + r = LocalDeltasReader(spark=None) + r.spark = MagicMock() + r._fixtures_root = tmp_path + r._credentials_profile = "default" + r._dataspace = None + r._sf_cli_org = None + r._batch = MagicMock() + r._current_layer = "dlo" + yield r + + +def _write_schema_file(fixtures_root, name): + out_dir = fixtures_root / name + out_dir.mkdir(parents=True, exist_ok=True) + (out_dir / "_schema.json").write_text(json.dumps(SOURCE_SCHEMA.jsonValue())) + + +def test_stream_schema_is_source_columns_plus_three_cdf_columns(reader, tmp_path): + _write_schema_file(tmp_path, "Foo__dll") + + schema = reader._build_stream_schema("Foo__dll") + + names = [f.name for f in schema.fields] + assert names == [ + "id__c", + "age__c", + cdf.COMMIT_VERSION, + cdf.COMMIT_TIMESTAMP, + cdf.MERGE_RECORD_TYPE, + ] + + +def test_open_stream_invokes_seeder_when_schema_file_missing(reader, tmp_path): + reader.spark.readStream.format.return_value = reader.spark.readStream + reader.spark.readStream.schema.return_value = reader.spark.readStream + reader.spark.readStream.option.return_value = reader.spark.readStream + reader.spark.readStream.load.return_value = MagicMock() + + def fake_seed(name): + _write_schema_file(tmp_path, name) + return True + + with patch.object(reader, "_try_seed", side_effect=fake_seed) as try_seed: + reader._open_stream("Foo__dll") + try_seed.assert_called_once_with("Foo__dll") + + +def test_open_stream_exits_cleanly_on_empty_source(reader): + with patch.object(reader, "_try_seed", return_value=False): + with pytest.raises(SystemExit) as exc: + reader._open_stream("Foo__dll") + assert exc.value.code == 0 + reader.spark.readStream.format.assert_not_called() + + +def test_read_dlo_delegates_to_batch_reader(reader): + reader.read_dlo("Foo__dll", None) + reader._batch.read_dlo.assert_called_once_with("Foo__dll", None) diff --git a/tests/io/reader/test_streaming_seeder.py b/tests/io/reader/test_streaming_seeder.py new file mode 100644 index 0000000..f725895 --- /dev/null +++ b/tests/io/reader/test_streaming_seeder.py @@ -0,0 +1,83 @@ +from __future__ import annotations + +import json +from unittest.mock import MagicMock, patch + +import pandas as pd +import pytest +from pyspark.sql.types import ( + LongType, + StringType, + StructField, + StructType, +) + +from datacustomcode.io import cdf +from datacustomcode.io.reader.streaming_seeder import StreamingSourceSeeder + + +SOURCE_SCHEMA = StructType( + [ + StructField("id__c", StringType(), True), + StructField("age__c", LongType(), True), + ] +) + + +def _stub_read_result(rows_df: pd.DataFrame): + limited = MagicMock() + limited.toPandas.return_value = rows_df + limited.schema = SOURCE_SCHEMA + + df = MagicMock() + df.limit.return_value = limited + return df + + +@pytest.fixture +def seeder(): + """Seeder with __init__ short-circuited (avoids real Query API creds).""" + with patch.object(StreamingSourceSeeder, "__init__", return_value=None): + s = StreamingSourceSeeder(spark=None) + s.spark = MagicMock() + s.reader = MagicMock() + yield s + + +def test_non_empty_source_writes_schema_and_seed(seeder, tmp_path): + seeder.reader.read_dlo.return_value = _stub_read_result( + pd.DataFrame( + [{"id__c": str(i), "age__c": i} for i in range(4)] + ) + ) + + wrote = seeder.seed_source("Foo__dll", "dlo", str(tmp_path)) + + assert wrote is True + out_dir = tmp_path / "Foo__dll" + assert (out_dir / "_schema.json").exists() + assert (out_dir / "000_seed.json").exists() + + rows = [ + json.loads(line) + for line in (out_dir / "000_seed.json").read_text().splitlines() + ] + for row in rows: + assert cdf.COMMIT_VERSION in row + assert cdf.COMMIT_TIMESTAMP in row + assert cdf.MERGE_RECORD_TYPE in row + # Both operation types appear in the seed. + ops = {row[cdf.MERGE_RECORD_TYPE] for row in rows} + assert ops == { + cdf.MergeRecordType.UPSERT.value, + cdf.MergeRecordType.DELETE.value, + } + + +def test_empty_source_returns_false_and_writes_nothing(seeder, tmp_path): + seeder.reader.read_dlo.return_value = _stub_read_result(pd.DataFrame([])) + + wrote = seeder.seed_source("Foo__dll", "dlo", str(tmp_path)) + + assert wrote is False + assert not (tmp_path / "Foo__dll").exists() diff --git a/tests/io/writer/test_local_deltas.py b/tests/io/writer/test_local_deltas.py new file mode 100644 index 0000000..7f48358 --- /dev/null +++ b/tests/io/writer/test_local_deltas.py @@ -0,0 +1,79 @@ +from __future__ import annotations + +from unittest.mock import MagicMock, patch + +import pytest + +from datacustomcode.io import cdf +from datacustomcode.io.writer.base import WriteMode +from datacustomcode.io.writer.local_deltas import LocalDeltasWriter + + +@pytest.fixture +def writer(): + with patch.object(LocalDeltasWriter, "__init__", return_value=None): + w = LocalDeltasWriter(spark=None) + w.spark = MagicMock() + w._batch_writer = MagicMock() + yield w + + +def _configure_writestream_chain(df): + ws = MagicMock() + df.writeStream = ws + ws.foreachBatch.return_value = ws + ws.option.return_value = ws + ws.trigger.return_value = ws + return ws + + +def test_write_dlo_deltas_uses_trigger_available_now(writer): + df = MagicMock() + ws = _configure_writestream_chain(df) + + writer.write_dlo_deltas("Foo__dll", df) + + ws.trigger.assert_called_once_with(availableNow=True) + + +def test_preview_drops_cdf_metadata_columns(writer): + df = MagicMock() + ws = _configure_writestream_chain(df) + writer.write_dlo_deltas("Foo__dll", df) + callback = ws.foreachBatch.call_args.args[0] + + # Set up the batch DF so .withColumn().drop() → a mock we can inspect. + annotated_after_with = MagicMock() + annotated_after_drop = MagicMock() + annotated_after_drop.count.return_value = 2 + annotated_after_with.drop.return_value = annotated_after_drop + batch_df = MagicMock() + batch_df.withColumn.return_value = annotated_after_with + + # pyspark.sql.functions.{col,when,lit} need a live SparkContext to + # produce real Columns; stub them out for the pure-Python paths. + with patch("pyspark.sql.functions.col", MagicMock()), \ + patch("pyspark.sql.functions.when", MagicMock()), \ + patch("pyspark.sql.functions.lit", MagicMock()): + callback(batch_df, 0) + + # drop() removed all three CDF metadata columns. + assert set(annotated_after_with.drop.call_args.args) == { + cdf.COMMIT_VERSION, + cdf.COMMIT_TIMESTAMP, + cdf.MERGE_RECORD_TYPE, + } + + +def test_batch_writes_delegate_to_print_writer(writer): + df = MagicMock() + + writer.write_to_dlo("Foo__dll", df, WriteMode.APPEND) + writer.auto_write_to_dlo("Foo__dll", df) + + writer._batch_writer.write_to_dlo.assert_called_once_with( + "Foo__dll", df, WriteMode.APPEND + ) + writer._batch_writer.auto_write_to_dlo.assert_called_once_with( + "Foo__dll", df + ) diff --git a/tests/test_template.py b/tests/test_template.py index 7475a97..89c0c3e 100644 --- a/tests/test_template.py +++ b/tests/test_template.py @@ -114,3 +114,28 @@ def test_copy_template_with_existing_content(self, temp_dir): template_items = os.listdir(script_template_dir) for item in template_items: assert os.path.exists(os.path.join(temp_dir, item)) + + # --- streaming=True behavior --- + + def test_copy_template_streaming_writes_config_yaml(self, temp_dir): + """Design: streaming init drops a config.yaml at the project root + that wires LocalDeltasReader + LocalDeltasWriter, so + `datacustomcode run` picks the streaming reader/writer via + run.py's auto-discovery. + """ + copy_script_template(temp_dir, streaming=True) + + config_yaml = os.path.join(temp_dir, "config.yaml") + assert os.path.isfile(config_yaml) + content = open(config_yaml).read() + assert "LocalDeltasReader" in content + assert "LocalDeltasWriter" in content + + def test_copy_template_default_batch_omits_streaming_config(self, temp_dir): + """Design: batch init MUST NOT write config.yaml — batch projects + keep using the SDK-shipped default (QueryAPIDataCloudReader + + PrintDataCloudWriter). Regression guard: dropping a project-local + config.yaml would silently override the batch profile. + """ + copy_script_template(temp_dir) + assert not os.path.exists(os.path.join(temp_dir, "config.yaml"))