soda-sparkdf 4.23.1__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,11 @@
1
+ Metadata-Version: 2.4
2
+ Name: soda-sparkdf
3
+ Version: 4.23.1
4
+ Summary: Soda SparkDF V4
5
+ Author-email: "Soda Data N.V." <info@soda.io>
6
+ License: Proprietary
7
+ Requires-Python: >=3.10
8
+ Requires-Dist: soda-core==4.23.1
9
+ Requires-Dist: freezegun
10
+ Requires-Dist: pyspark[connect]>=3.5.0
11
+ Requires-Dist: soda-databricks==4.23.1
@@ -0,0 +1,29 @@
1
+ [project]
2
+ name = "soda-sparkdf"
3
+ version = "4.23.1"
4
+ description = "Soda SparkDF V4"
5
+ requires-python = ">=3.10"
6
+ license = {text = "Proprietary"}
7
+ authors = [
8
+ {name = "Soda Data N.V.", email = "info@soda.io"}
9
+ ]
10
+ dependencies = [
11
+ "soda-core==4.23.1",
12
+ "freezegun",
13
+ "pyspark[connect]>=3.5.0",
14
+ "soda-databricks==4.23.1",
15
+ ]
16
+
17
+ [project.entry-points."soda.plugins.data_source.sparkdf"]
18
+ SparkDataFrameDataSourceImpl = "soda_sparkdf.common.data_sources.sparkdf_data_source:SparkDataFrameDataSourceImpl"
19
+
20
+ [tool.uv.sources]
21
+ soda-core = { workspace = true }
22
+ soda-databricks = { workspace = true }
23
+
24
+ [build-system]
25
+ requires = ["setuptools>=45", "wheel"]
26
+ build-backend = "setuptools.build_meta"
27
+
28
+ [tool.setuptools]
29
+ package-dir = {"" = "src"}
@@ -0,0 +1,4 @@
1
+ [egg_info]
2
+ tag_build =
3
+ tag_date = 0
4
+
@@ -0,0 +1,7 @@
1
+ from soda_sparkdf.common.data_sources.sparkdf_data_source import (
2
+ SparkDataFrameDataSourceImpl as SparkDataFrameDataSource,
3
+ )
4
+
5
+ __all__ = [
6
+ "SparkDataFrameDataSource",
7
+ ]
@@ -0,0 +1,493 @@
1
+ from datetime import datetime, timezone, tzinfo
2
+ from typing import Any, Optional
3
+
4
+ from freezegun import freeze_time
5
+ from pyspark.sql import DataFrame, SparkSession
6
+ from pyspark.sql.types import Row
7
+ from soda_core.common.data_source_connection import DataSourceConnection, parse_session_timezone
8
+ from soda_core.common.data_source_impl import DataSourceImpl, MetadataTablesQuery
9
+ from soda_core.common.data_source_results import QueryResult
10
+ from soda_core.common.metadata_types import ColumnMetadata, SodaDataTypeName
11
+ from soda_core.common.sql_dialect import SqlDialect
12
+ from soda_core.common.statements.metadata_tables_query import FullyQualifiedTableName
13
+ from soda_core.common.statements.table_types import FullyQualifiedViewName, TableType
14
+ from soda_databricks.common.data_sources.databricks_data_source import DatabricksSqlDialect
15
+ from soda_databricks.common.statements.hive_metadata_tables_query import HiveMetadataTablesQuery
16
+ from soda_sparkdf.common.data_sources.sparkdf_data_source_connection import (
17
+ SparkDataFrameActiveSessionProperties,
18
+ SparkDataFrameConnectionProperties,
19
+ )
20
+ from soda_sparkdf.common.data_sources.sparkdf_data_source_connection import (
21
+ SparkDataFrameDataSource as SparkDataFrameDataSourceModel,
22
+ )
23
+ from soda_sparkdf.common.data_sources.sparkdf_data_source_connection import (
24
+ SparkDataFrameExistingSessionProperties,
25
+ SparkDataFrameNewSessionProperties,
26
+ SparkDataFrameRemoteSessionProperties,
27
+ )
28
+
29
+ _in_memory_connection = None
30
+
31
+
32
+ class SparkDataFrameCursor:
33
+ CACHE_ROW_COUNT = 100
34
+
35
+ def __init__(self, spark_session: SparkSession, test_dir: Optional[str] = None):
36
+ self.spark_session = spark_session
37
+ self.df: DataFrame | None = None
38
+ self.description: tuple[tuple] | None = None
39
+ self.cursor_index: int = -1
40
+ self._cache_index = -1
41
+ self._cached_rows: list[Row] | None = None
42
+
43
+ def execute(self, sql: str):
44
+ self.df = self.spark_session.sql(sqlQuery=sql)
45
+ self.description = self.convert_spark_df_schema_to_dbapi_description(self.df)
46
+ self.cursor_index = 0
47
+ self._cache_index = -1
48
+ self._cached_rows = None
49
+
50
+ @property
51
+ def rowcount(self) -> int:
52
+ if self.df is None:
53
+ return -1
54
+ return self.df.count()
55
+
56
+ def fetchall(self) -> tuple[tuple]:
57
+ rows = []
58
+ with freeze_time(
59
+ datetime.now(timezone.utc)
60
+ ): # We need to freeze the time to UTC at the time of collecting to avoid issues with timestamps
61
+ # Spark stores the timestamps in UTC (verified by reading the parquet file that spark creates), but when querying it converts it to the (python) session-local timezone.
62
+ # By using freeze_time, we can ensure that the timestamps are collected in UTC, regardless of the session-local timezone. This plays nice with the freshness checks.
63
+ spark_rows: list[Row] = self.df.collect()
64
+ # Alternative approach: convert to PyArrow. This will set the timestamps correctly, but introduces memory and time overhead (for the conversion).
65
+ # This also requires more changes regarding the cursor implementation: spark_rows: list[Row] = self.df.toArrow().to_pylist()
66
+ for spark_row in spark_rows:
67
+ row = self.convert_spark_row_to_dbapi_row(spark_row)
68
+ rows.append(row)
69
+ return tuple(rows)
70
+
71
+ def fetchmany(self, size: int) -> tuple[tuple]:
72
+ rows = []
73
+ with freeze_time(
74
+ datetime.now(timezone.utc)
75
+ ): # We need to freeze the time to UTC at the time of collecting to avoid issues with timestamps. See the comment in fetchall() for more details.
76
+ spark_rows: list[Row] = self.df.offset(self.cursor_index).limit(size).collect()
77
+ self.cursor_index += len(spark_rows)
78
+ for spark_row in spark_rows:
79
+ row = self.convert_spark_row_to_dbapi_row(spark_row)
80
+ rows.append(row)
81
+ return tuple(rows)
82
+
83
+ def fetchone(self) -> tuple | None:
84
+ # Fetches have overhead, so we load and cache small pages here as a compromise
85
+ if self._cached_rows is None or self.cursor_index >= self._cache_index + self.CACHE_ROW_COUNT:
86
+ with freeze_time(
87
+ datetime.now(timezone.utc)
88
+ ): # We need to freeze the time to UTC at the time of collecting to avoid issues with timestamps. See the comment in fetchall() for more details.
89
+ self._cached_rows = self.df.offset(self.cursor_index).limit(self.CACHE_ROW_COUNT).collect()
90
+ self._cache_index = self.cursor_index
91
+ access_index = self.cursor_index - self._cache_index
92
+ if not self._cached_rows or access_index >= len(self._cached_rows):
93
+ return None
94
+ spark_row = self._cached_rows[access_index]
95
+ self.cursor_index += 1
96
+ row = self.convert_spark_row_to_dbapi_row(spark_row)
97
+ return tuple(row)
98
+
99
+ @staticmethod
100
+ def convert_spark_row_to_dbapi_row(spark_row):
101
+ return [spark_row[field] for field in spark_row.__fields__]
102
+
103
+ def close(self):
104
+ pass # No-op
105
+
106
+ @staticmethod
107
+ def convert_spark_df_schema_to_dbapi_description(df) -> tuple[tuple]:
108
+ # simpleString() yields Spark SQL type names like "int", "decimal(10,2)", "string" —
109
+ # parseable by the DWH extension and aligned with the dialect's supported-type list.
110
+ return tuple((field.name, field.dataType.simpleString()) for field in df.schema.fields)
111
+
112
+
113
+ class SparkDataFrameDataSourceConnectionWrapper:
114
+ def __init__(self, session: SparkSession):
115
+ self._session = session
116
+
117
+ def __getattr__(self, attr):
118
+ if attr in self.__dict__:
119
+ return getattr(self, attr)
120
+ return getattr(self._session, attr)
121
+
122
+ def commit(self):
123
+ pass # Do nothing, Spark does not have a commit concept
124
+
125
+ def cursor(self):
126
+ return SparkDataFrameCursor(self._session)
127
+
128
+
129
+ class SparkDataFrameSqlDialect(DatabricksSqlDialect, sqlglot_dialect="spark"):
130
+ SODA_DATA_TYPE_SYNONYMS = (
131
+ (SodaDataTypeName.TEXT, SodaDataTypeName.VARCHAR, SodaDataTypeName.CHAR),
132
+ (SodaDataTypeName.NUMERIC, SodaDataTypeName.DECIMAL),
133
+ (SodaDataTypeName.TIMESTAMP_TZ, SodaDataTypeName.TIMESTAMP),
134
+ )
135
+ # Class-level default: legacy local-Spark rejects ``DROP TABLE ... CASCADE``
136
+ # with a ParseException. In catalog mode (Databricks) CASCADE is supported, so
137
+ # __init__ promotes this to True via an instance attribute that shadows the class.
138
+ SUPPORTS_DROP_TABLE_CASCADE: bool = False
139
+
140
+ def __init__(self, use_catalog: bool = False):
141
+ super().__init__()
142
+ # In catalog mode, prefixes are [catalog, schema] (Unity-Catalog style 3-level
143
+ # namespace). In legacy mode, prefixes are [schema] only — local Spark has no
144
+ # catalog concept beyond ``spark_catalog``.
145
+ self.use_catalog = use_catalog
146
+ # Instance override — see class-level comment above.
147
+ self.SUPPORTS_DROP_TABLE_CASCADE = use_catalog
148
+
149
+ def supports_primary_keys(self) -> bool:
150
+ # Open-source Spark SQL (the local, Hive-metastore-backed engine this data source runs
151
+ # against) has no information_schema and cannot declare a PRIMARY KEY in CREATE TABLE,
152
+ # so primary-key introspection stays opt-out — overriding the Unity Catalog default
153
+ # inherited from DatabricksSqlDialect.
154
+ return False
155
+
156
+ def get_database_prefix_index(self) -> int | None:
157
+ return 0 if self.use_catalog else None
158
+
159
+ def get_schema_prefix_index(self) -> int | None:
160
+ return 1 if self.use_catalog else 0
161
+
162
+ def create_schema_if_not_exists_sql(self, prefixes: list[str], add_semicolon: bool = True) -> str:
163
+ if self.use_catalog:
164
+ if len(prefixes) < 2:
165
+ raise ValueError(
166
+ f"Catalog-mode SparkDF requires 2 prefixes [catalog, schema]; got {len(prefixes)}: {prefixes}"
167
+ )
168
+ catalog_name: str = prefixes[0]
169
+ schema_name: str = prefixes[1]
170
+ quoted = f"{self.quote_default(catalog_name)}.{self.quote_default(schema_name)}"
171
+ else:
172
+ if len(prefixes) < 1:
173
+ raise ValueError(f"SparkDF requires at least 1 prefix [schema]; got {len(prefixes)}: {prefixes}")
174
+ quoted = self.quote_default(prefixes[0])
175
+ return f"CREATE SCHEMA IF NOT EXISTS {quoted}" + (";" if add_semicolon else "")
176
+
177
+ def post_schema_create_sql(self, prefixes: list[str]) -> Optional[list[str]]:
178
+ pass # Do nothing, Spark does not have a post-schema-create concept
179
+
180
+ def build_column_metadatas_from_query_result(self, query_result: QueryResult) -> list[ColumnMetadata]:
181
+ # Filter out dataset description rows (first such line starts with #, ignore the rest) or empty
182
+ filtered_rows = []
183
+ for row in query_result.rows:
184
+ if row[0].startswith("#"): # ignore all description rows
185
+ break
186
+ if not row[0] and not row[1]: # empty row
187
+ continue
188
+
189
+ filtered_rows.append(row)
190
+
191
+ return super().build_column_metadatas_from_query_result(
192
+ QueryResult(rows=filtered_rows, columns=query_result.columns)
193
+ )
194
+
195
+ def literal_datetime_with_tz(self, datetime: datetime):
196
+ # Always convert the timestamp to utc when we insert. Spark is not aware of the timezones, so we need to do this conversion so it's ready to be extracted as UTC.
197
+ return f"to_utc_timestamp('{datetime.isoformat()}', 'UTC')"
198
+
199
+ def literal_datetime(self, datetime: datetime):
200
+ # Always convert the timestamp to utc when we insert. Spark is not aware of the timezones, so we need to do this conversion so it's ready to be extracted as UTC.
201
+ return f"to_utc_timestamp('{datetime.isoformat()}', 'UTC')"
202
+
203
+
204
+ class SparkDataFrameDataSourceConnection(DataSourceConnection):
205
+ def __init__(self, name: str, connection_properties: dict, connection: Optional[object] = None):
206
+ # When the caller supplies a pre-built ``connection``, the base ``open_connection`` skips
207
+ # ``_create_connection`` (which is where ``self.session`` would normally be assigned).
208
+ # Pull the session out of ``connection_properties`` up-front so ``self.session`` is
209
+ # always available — DWH calls like ``test_schema_exists`` rely on it.
210
+ self.session: Optional[SparkSession] = None
211
+ if isinstance(connection_properties, dict):
212
+ existing_session = connection_properties.get("spark_session")
213
+ if existing_session is not None:
214
+ self.session = existing_session
215
+ super().__init__(name, connection_properties, connection)
216
+
217
+ def _create_connection(
218
+ self,
219
+ config: SparkDataFrameConnectionProperties,
220
+ ):
221
+ session = None
222
+ if isinstance(config, SparkDataFrameExistingSessionProperties):
223
+ session = config.spark_session
224
+ elif isinstance(config, SparkDataFrameActiveSessionProperties):
225
+ # Pick up the thread-local active SparkSession (the Databricks notebook's
226
+ # ``spark``, or whatever the caller last did ``.getOrCreate()`` on). No URI,
227
+ # no credentials, no global state we own — just pyspark's existing notion
228
+ # of which session is active in this thread.
229
+ session = SparkSession.getActiveSession()
230
+ if session is None:
231
+ raise ValueError(
232
+ "SparkDataFrame is configured with use_active_session=True but no active "
233
+ "SparkSession was found. Build a session (e.g. SparkSession.builder.…"
234
+ "getOrCreate()) before opening this connection, or switch to another "
235
+ "connection mode: pass ``spark_session`` (existing session), "
236
+ "``host`` + ``token`` + ``cluster_id`` (remote Spark Connect), or "
237
+ "``new_session: true`` (local session)."
238
+ )
239
+ elif isinstance(config, SparkDataFrameRemoteSessionProperties):
240
+ # Spark Connect URI. ``token`` becomes a gRPC bearer header (handled by
241
+ # pyspark.sql.connect.ChannelBuilder); ``x-databricks-cluster-id`` is forwarded
242
+ # as gRPC metadata, which is how Databricks routes the session to a cluster.
243
+ # ``getOrCreate`` caches per URI in-process, so multiple data sources pointing
244
+ # at the same workspace+cluster end up sharing one underlying session.
245
+ # NB: ``config.token`` is a SecretStr — unwrap only when assembling the URI
246
+ # (which is kept as a local variable, never logged).
247
+ uri = (
248
+ f"sc://{config.host}:443/"
249
+ f";use_ssl=true"
250
+ f";token={config.token.get_secret_value()}"
251
+ f";x-databricks-cluster-id={config.cluster_id}"
252
+ )
253
+ session = SparkSession.builder.remote(uri).getOrCreate()
254
+ session.sql("SET TIME ZONE 'UTC'")
255
+ elif isinstance(config, SparkDataFrameNewSessionProperties):
256
+ session = (
257
+ SparkSession.builder.master("local")
258
+ .appName(self.name)
259
+ .config("spark.sql.warehouse.dir", config.test_dir)
260
+ .getOrCreate()
261
+ )
262
+ session.sql("SET spark.sql.session.timeZone = +00:00;")
263
+ session.sql("SET TIME ZONE 'UTC';")
264
+ if session is None:
265
+ raise ValueError("No session provided")
266
+ self.session = session
267
+ return SparkDataFrameDataSourceConnectionWrapper(session=session)
268
+
269
+ def close_connection(self) -> None:
270
+ "This is a no-op for SparkDataFrameDataSourceConnection, there is no connection to close."
271
+
272
+ def _fetch_session_timezone(self) -> tzinfo:
273
+ # New sessions created by this adapter explicitly ``SET TIME ZONE 'UTC';``,
274
+ # but ``SparkDataFrameExistingSessionProperties`` lets a caller wrap an existing
275
+ # SparkSession that may have any configured zone. Read the live setting so the
276
+ # value mappers see the same zone the Spark engine will use to interpret
277
+ # naive returns. ``parse_session_timezone`` accepts Spark's reported value
278
+ # ('UTC', '+00:00', 'America/Los_Angeles', etc.) and the connection-level
279
+ # wrapper falls back to UTC if the call raises.
280
+ tz_value = self.session.conf.get("spark.sql.session.timeZone")
281
+ return parse_session_timezone(tz_value)
282
+
283
+ def _execute_query_get_result_row_column_name(self, column) -> str:
284
+ return column[0] # The first element of the tuple is the column name
285
+
286
+ def _cursor_execute_update_and_commit(self, cursor: Any, sql: str) -> int:
287
+ cursor.execute(sql)
288
+ # Skip cursor.rowcount for Spark — it triggers df.count() which runs a full Spark job.
289
+ self.commit()
290
+ return 0
291
+
292
+
293
+ class SparkDataFrameMetadataTablesQuery(HiveMetadataTablesQuery):
294
+ """SparkDF variant of HiveMetadataTablesQuery.
295
+
296
+ Handles two SparkDF-specific quirks:
297
+ 1. SHOW TABLES / SHOW VIEWS throws AnalysisException when the schema doesn't exist
298
+ (the Databricks SQL connector returns empty). Wrap each call in try/except so DWH
299
+ introspection of the not-yet-created diagnostics schema doesn't crash.
300
+ 2. SHOW TABLES in Spark returns both tables and views; collect view names first and
301
+ subtract them so TableType.TABLE results are disjoint from TableType.VIEW.
302
+
303
+ Also emits ``database_name=None`` since SparkDF has no catalog concept (the dialect
304
+ uses ``database_prefix_index=None``).
305
+ """
306
+
307
+ def execute(
308
+ self,
309
+ database_name: Optional[str] = None,
310
+ schema_name: Optional[str] = None,
311
+ include_table_name_like_filters: Optional[list[str]] = None,
312
+ exclude_table_name_like_filters: Optional[list[str]] = None,
313
+ types_to_return: Optional[list[TableType]] = None,
314
+ ) -> list[FullyQualifiedTableName]:
315
+ if types_to_return is None:
316
+ types_to_return = [TableType.TABLE]
317
+ results: list[FullyQualifiedTableName] = []
318
+
319
+ view_names: set[str] = set()
320
+ if TableType.TABLE in types_to_return:
321
+ try:
322
+ view_sql = self.build_sql_statement(
323
+ database_name=database_name, schema_name=schema_name, object_type_to_fetch=TableType.VIEW
324
+ )
325
+ view_result: QueryResult = self.data_source_connection.execute_query(view_sql)
326
+ view_names = {row[1] for row in view_result.rows}
327
+ except Exception:
328
+ pass # Schema may not exist yet
329
+
330
+ if TableType.TABLE in types_to_return:
331
+ try:
332
+ sql = self.build_sql_statement(
333
+ database_name=database_name, schema_name=schema_name, object_type_to_fetch=TableType.TABLE
334
+ )
335
+ query_result: QueryResult = self.data_source_connection.execute_query(sql)
336
+ filtered_rows = [row for row in query_result.rows if row[1] not in view_names]
337
+ filtered_result = QueryResult(rows=filtered_rows, columns=query_result.columns)
338
+ results.extend(
339
+ self.get_results(
340
+ filtered_result,
341
+ object_type_to_fetch=TableType.TABLE,
342
+ include_table_name_like_filters=include_table_name_like_filters,
343
+ exclude_table_name_like_filters=exclude_table_name_like_filters,
344
+ )
345
+ )
346
+ except Exception:
347
+ pass
348
+
349
+ if TableType.VIEW in types_to_return:
350
+ try:
351
+ sql = self.build_sql_statement(
352
+ database_name=database_name, schema_name=schema_name, object_type_to_fetch=TableType.VIEW
353
+ )
354
+ query_result = self.data_source_connection.execute_query(sql)
355
+ results.extend(
356
+ self.get_results(
357
+ query_result,
358
+ object_type_to_fetch=TableType.VIEW,
359
+ include_table_name_like_filters=include_table_name_like_filters,
360
+ exclude_table_name_like_filters=exclude_table_name_like_filters,
361
+ )
362
+ )
363
+ except Exception:
364
+ pass
365
+
366
+ return results
367
+
368
+ def get_results(
369
+ self,
370
+ query_result: QueryResult,
371
+ object_type_to_fetch: TableType,
372
+ include_table_name_like_filters: Optional[list[str]] = None,
373
+ exclude_table_name_like_filters: Optional[list[str]] = None,
374
+ ) -> list[FullyQualifiedTableName]:
375
+ if object_type_to_fetch == TableType.TABLE:
376
+ names_for_filtering = [table_name for _, table_name, _ in query_result.rows]
377
+ elif object_type_to_fetch == TableType.VIEW:
378
+ names_for_filtering = [view_name for _, view_name, *_ in query_result.rows]
379
+ else:
380
+ raise ValueError(f"Invalid object type to fetch: {object_type_to_fetch}")
381
+ filtered_names = self._filter_include_exclude(
382
+ names_for_filtering, include_table_name_like_filters, exclude_table_name_like_filters
383
+ )
384
+
385
+ if object_type_to_fetch == TableType.TABLE:
386
+ return [
387
+ FullyQualifiedTableName(database_name=None, schema_name=schema_name, table_name=table_name)
388
+ for schema_name, table_name, _is_temporary in query_result.rows
389
+ if table_name in filtered_names
390
+ ]
391
+ elif object_type_to_fetch == TableType.VIEW:
392
+ return [
393
+ FullyQualifiedViewName(database_name=None, schema_name=schema_name, view_name=view_name)
394
+ for schema_name, view_name, *_ in query_result.rows
395
+ if view_name in filtered_names
396
+ ]
397
+ else:
398
+ raise ValueError(f"Invalid object type to fetch: {object_type_to_fetch}")
399
+
400
+
401
+ class SparkDataFrameDataSourceImpl(DataSourceImpl, model_class=SparkDataFrameDataSourceModel):
402
+ def _create_sql_dialect(self) -> SqlDialect:
403
+ return SparkDataFrameSqlDialect(use_catalog=self._read_use_catalog_flag())
404
+
405
+ def _read_use_catalog_flag(self) -> bool:
406
+ # Pydantic model when fully parsed, dict during the from_existing_session bootstrap.
407
+ props = self.data_source_model.connection_properties
408
+ if isinstance(props, dict):
409
+ return bool(props.get("use_catalog", False))
410
+ return bool(getattr(props, "use_catalog", False))
411
+
412
+ def _create_data_source_connection(self) -> DataSourceConnection:
413
+ return SparkDataFrameDataSourceConnection(
414
+ name=self.data_source_model.name, connection_properties=self.data_source_model.connection_properties
415
+ )
416
+
417
+ def create_metadata_tables_query(self) -> MetadataTablesQuery:
418
+ return SparkDataFrameMetadataTablesQuery(
419
+ sql_dialect=self.sql_dialect, data_source_connection=self.data_source_connection
420
+ )
421
+
422
+ @classmethod
423
+ def from_existing_session(
424
+ cls,
425
+ session: SparkSession,
426
+ name: str,
427
+ use_catalog: bool = False,
428
+ catalog: Optional[str] = None,
429
+ ) -> DataSourceImpl:
430
+ # Locally-owned dict so we never mutate caller-shared state via the
431
+ # ``connection_properties`` plumbing.
432
+ connection_properties = {
433
+ "spark_session": session,
434
+ "schema_": name,
435
+ "use_catalog": use_catalog,
436
+ "catalog": catalog,
437
+ }
438
+ ds_model = SparkDataFrameDataSourceModel(
439
+ name=name,
440
+ connection_properties=connection_properties,
441
+ )
442
+ soda_connection = SparkDataFrameDataSourceConnection(
443
+ name=name,
444
+ connection_properties=connection_properties,
445
+ connection=SparkDataFrameDataSourceConnectionWrapper(session),
446
+ )
447
+ return cls(data_source_model=ds_model, connection=soda_connection)
448
+
449
+ def build_columns_metadata_query_str(self, dataset_prefixes: list[str], dataset_name: str) -> str:
450
+ if len(dataset_prefixes) == 0:
451
+ return f"DESCRIBE {dataset_name}"
452
+ elif len(dataset_prefixes) == 1:
453
+ schema_name: str = dataset_prefixes[0]
454
+ return f"DESCRIBE {schema_name}.{dataset_name}"
455
+ elif len(dataset_prefixes) == 2:
456
+ database_name: str = dataset_prefixes[0]
457
+ schema_name: str = dataset_prefixes[1]
458
+ return f"DESCRIBE {database_name}.{schema_name}.{dataset_name}"
459
+ else:
460
+ raise ValueError(f"Invalid number of dataset prefixes: {len(dataset_prefixes)}")
461
+
462
+ @property
463
+ def bulk_columns_metadata_available(self) -> bool:
464
+ return False
465
+
466
+ def test_schema_exists(self, prefixes: list[str]) -> bool:
467
+ # Catalog-mode prefixes are [catalog, schema]; we list schemas in the catalog and do
468
+ # exact-match in Python rather than ``LIKE '<schema>'`` so a schema named
469
+ # ``soda_diagnostics`` doesn't false-positive against ``sodaXdiagnostics`` via the
470
+ # ``_`` wildcard. Quote the catalog so names with hyphens (legal in UC) parse cleanly.
471
+ if self.sql_dialect.get_database_prefix_index() is not None and len(prefixes) >= 2:
472
+ catalog_name = prefixes[0]
473
+ schema_name = prefixes[1]
474
+ quoted_catalog = self.sql_dialect.quote_default(catalog_name)
475
+ try:
476
+ result = self.connection.session.sql(f"SHOW SCHEMAS IN {quoted_catalog}").collect()
477
+ except Exception:
478
+ # Catalog doesn't exist (or we can't see into it) — the subsequent
479
+ # ``CREATE SCHEMA IF NOT EXISTS <cat>.<schema>`` will surface the real error.
480
+ return False
481
+ for row in result:
482
+ if row[0] and row[0].lower() == schema_name.lower():
483
+ return True
484
+ return False
485
+ result = self.connection.session.sql("SHOW SCHEMAS").collect()
486
+ for row in result:
487
+ if row[0] and row[0].lower() == prefixes[0].lower():
488
+ return True
489
+ return False
490
+
491
+
492
+ # Alias to make the import and usage cleaner
493
+ SparkDataFrameDataSource = SparkDataFrameDataSourceImpl
@@ -0,0 +1,152 @@
1
+ import re
2
+ from abc import ABC
3
+ from typing import Any, Literal, Optional, Union
4
+
5
+ from pydantic import Field, SecretStr, field_validator
6
+ from soda_core.model.data_source.data_source import DataSourceBase
7
+ from soda_core.model.data_source.data_source_connection_properties import DataSourceConnectionProperties
8
+
9
+
10
+ class SparkDataFrameConnectionProperties(DataSourceConnectionProperties, ABC):
11
+ schema_: Optional[str] = Field(
12
+ "main", description="Optional schema name to use for the SparkDataFrame connection", alias="schema"
13
+ )
14
+ test_dir: Optional[str] = Field(None, description="The directory to use for the test")
15
+ use_catalog: bool = Field(
16
+ False,
17
+ description=(
18
+ "When True, treat DWH prefixes as [catalog, schema] so DWH tables land in a "
19
+ "Unity-Catalog-style 3-level namespace. The dialect qualifies CREATE SCHEMA and "
20
+ "schema-existence checks with the catalog (e.g. ``catalog.schema``); the catalog "
21
+ "itself is assumed to already exist (Soda does not auto-create catalogs). Leave "
22
+ "False for local Spark which has no catalog concept."
23
+ ),
24
+ )
25
+ catalog: Optional[str] = Field(
26
+ None,
27
+ description=(
28
+ "Explicit catalog binding for downstream consumers that need a 3-level "
29
+ "Unity-Catalog identifier (e.g. soda-extensions' diagnostics-warehouse writer "
30
+ "when ``use_catalog=true``). soda-core itself does not interpret this field; "
31
+ "consumers are expected to fall back to ``SELECT current_catalog()`` against "
32
+ "the active SparkSession when it is unset. Ignored when ``use_catalog=false``."
33
+ ),
34
+ )
35
+
36
+
37
+ class SparkDataFrameNewSessionProperties(SparkDataFrameConnectionProperties):
38
+ new_session: bool = Field(True, description="Whether to create a new Spark session")
39
+
40
+
41
+ class SparkDataFrameExistingSessionProperties(SparkDataFrameConnectionProperties, arbitrary_types_allowed=True):
42
+ # We set the type to Any to avoid type errors when the SparkSession is not a SparkSession object
43
+ # This could be the case on Databricks serverless, where the SparkSession is imported as a different object
44
+ spark_session: Any = Field(..., description="The existing Spark session to use")
45
+
46
+
47
+ class SparkDataFrameRemoteSessionProperties(SparkDataFrameConnectionProperties):
48
+ """SparkSession built via Spark Connect against a remote workspace (e.g. Databricks).
49
+
50
+ The connection builds ``SparkSession.builder.remote(<uri>).getOrCreate()`` using a
51
+ Spark Connect URI assembled from ``host``, ``token``, and ``cluster_id``. Because
52
+ pyspark's Spark Connect builder caches sessions by URI per Python process, two
53
+ DataSourceImpls configured against the same workspace+cluster share one underlying
54
+ session — which is what we want for between-source DWH against Databricks.
55
+ """
56
+
57
+ host: str = Field(..., description="Workspace host (e.g. dbc-12345.cloud.databricks.com)")
58
+ # SecretStr so the PAT renders as '**********' in repr/str — soda-core logs
59
+ # connection_properties at DEBUG, which would otherwise leak the token verbatim
60
+ # whenever verbose logging is on.
61
+ token: SecretStr = Field(..., description="Personal access token, sent as gRPC bearer auth")
62
+ cluster_id: str = Field(
63
+ ...,
64
+ description=(
65
+ "All-purpose cluster id, forwarded as ``x-databricks-cluster-id`` gRPC metadata "
66
+ "to route the Spark Connect session to a specific cluster"
67
+ ),
68
+ )
69
+
70
+ @field_validator("host", mode="before")
71
+ @classmethod
72
+ def _strip_host_scheme(cls, value):
73
+ # Users naturally paste workspace URLs like ``https://dbc-12345.cloud.databricks.com``.
74
+ # The Spark Connect URI hard-codes its own ``sc://`` scheme and ``use_ssl=true``, so a
75
+ # scheme in the host would produce ``sc://https://...:443/`` and surface as an opaque
76
+ # gRPC error. Strip the scheme and trailing slash up-front.
77
+ if isinstance(value, str):
78
+ return re.sub(r"^https?://", "", value.strip()).rstrip("/")
79
+ return value
80
+
81
+
82
+ class SparkDataFrameActiveSessionProperties(SparkDataFrameConnectionProperties):
83
+ """SparkSession picked up from ``SparkSession.getActiveSession()`` at connect time.
84
+
85
+ Lets a DWH YAML reuse the SparkSession that's already active in the current thread —
86
+ typically the Databricks notebook's ``spark``, or a session a caller built via
87
+ ``SparkSession.builder.…getOrCreate()`` before invoking contract verification. No
88
+ credentials in YAML, no module-level registry, no monkey-patching: pyspark already
89
+ tracks the active session per thread and we just retrieve it.
90
+
91
+ The connection raises a clear error when no active session is found.
92
+ """
93
+
94
+ use_active_session: Literal[True] = Field(
95
+ ...,
96
+ description=(
97
+ "Must be ``true``. Discriminator for picking this connection mode in YAML; the "
98
+ "session itself is fetched via SparkSession.getActiveSession() at connect time."
99
+ ),
100
+ )
101
+
102
+
103
+ class SparkDataFrameDataSource(DataSourceBase, ABC):
104
+ type: Literal["sparkdf"] = Field("sparkdf")
105
+ connection_properties: Union[
106
+ SparkDataFrameExistingSessionProperties,
107
+ SparkDataFrameRemoteSessionProperties,
108
+ SparkDataFrameActiveSessionProperties,
109
+ SparkDataFrameNewSessionProperties,
110
+ ] = Field(..., alias="connection", description="SparkDataFrame connection configuration")
111
+
112
+ @field_validator("connection_properties", mode="before")
113
+ @classmethod
114
+ def infer_connection_type(cls, value):
115
+ if isinstance(value, SparkDataFrameNewSessionProperties):
116
+ return value
117
+ if isinstance(value, SparkDataFrameExistingSessionProperties):
118
+ return value
119
+ if isinstance(value, SparkDataFrameRemoteSessionProperties):
120
+ return value
121
+ if isinstance(value, SparkDataFrameActiveSessionProperties):
122
+ return value
123
+
124
+ # Reject ambiguous combinations up-front rather than letting the first matching
125
+ # branch win silently. Each mode has a unique discriminator; mixing them is
126
+ # almost always a typo or a half-finished YAML edit.
127
+ modes_present = []
128
+ if "spark_session" in value:
129
+ modes_present.append("spark_session (existing_session mode)")
130
+ if value.get("use_active_session") is True:
131
+ modes_present.append("use_active_session (active_session mode)")
132
+ if "host" in value or "cluster_id" in value:
133
+ modes_present.append("host/cluster_id (remote_session mode)")
134
+ if "new_session" in value:
135
+ modes_present.append("new_session mode")
136
+ if len(modes_present) > 1:
137
+ raise ValueError(
138
+ "Conflicting SparkDataFrame connection config — multiple modes detected: "
139
+ + ", ".join(modes_present)
140
+ + ". Pick exactly one of: spark_session, use_active_session=true, "
141
+ "host+cluster_id, or new_session."
142
+ )
143
+
144
+ if "spark_session" in value:
145
+ return SparkDataFrameExistingSessionProperties(**value)
146
+ elif value.get("use_active_session") is True:
147
+ return SparkDataFrameActiveSessionProperties(**value)
148
+ elif "host" in value and "cluster_id" in value:
149
+ return SparkDataFrameRemoteSessionProperties(**value)
150
+ elif "new_session" in value:
151
+ return SparkDataFrameNewSessionProperties(**value)
152
+ raise ValueError("Could not infer SparkDataFrame connection type from input")
@@ -0,0 +1,55 @@
1
+ from __future__ import annotations
2
+
3
+ import logging
4
+ import tempfile
5
+ from typing import Optional
6
+
7
+ from helpers.data_source_test_helper import DataSourceTestHelper
8
+
9
+ logger = logging.getLogger(__name__)
10
+
11
+
12
+ class SparkDataFrameDataSourceTestHelper(DataSourceTestHelper):
13
+ def __init__(self, name: str):
14
+ super().__init__(name)
15
+ # SparkDF uses in-memory session state that doesn't survive the snapshot
16
+ # connection's per-test reset, causing full-suite record mode to fail.
17
+ # Force snapshot mode off for SparkDF.
18
+ if self._snapshot_mode != "off":
19
+ logger.info(f"SparkDF does not support snapshot mode '{self._snapshot_mode}', forcing 'off'")
20
+ self._snapshot_mode = "off"
21
+
22
+ def _create_data_source_yaml_str(self) -> str:
23
+ """
24
+ Called in _create_data_source_impl to initialized self.data_source_impl
25
+ self.database_name and self.schema_name are available if appropriate for the data source type
26
+ """
27
+ self.test_dir = tempfile.mkdtemp(prefix=f"soda_test_sparkdf_{self.name}_")
28
+ return f"""
29
+ type: sparkdf
30
+ name: {self.name}
31
+ connection:
32
+ new_session: true
33
+ test_dir: {self.test_dir}
34
+ """
35
+
36
+ # We need these methods to comply with the rest of the test helper infrastructure
37
+ def _create_database_name(self) -> Optional[str]:
38
+ return None
39
+
40
+ def _create_schema_name(self) -> Optional[str]:
41
+ # Two helpers (primary + secondary) in the same test session share the
42
+ # JVM-level SparkSession via ``SparkSession.builder.getOrCreate()``, which
43
+ # also means they share the catalog. Use a helper-name-suffixed schema so
44
+ # primary and secondary don't collide on identical TestTableSpecification
45
+ # hashes (e.g. reconciliation tests that build the same fixture on both).
46
+ return f"main_{self.name}"
47
+
48
+ def _create_dataset_prefix(self) -> list[str]:
49
+ schema_name: str = self._create_schema_name()
50
+ return [schema_name]
51
+
52
+ def drop_test_schema_if_exists(self) -> None:
53
+ """
54
+ In-memory SparkDF does not support schemas, so this is a no-op.
55
+ """
@@ -0,0 +1,11 @@
1
+ Metadata-Version: 2.4
2
+ Name: soda-sparkdf
3
+ Version: 4.23.1
4
+ Summary: Soda SparkDF V4
5
+ Author-email: "Soda Data N.V." <info@soda.io>
6
+ License: Proprietary
7
+ Requires-Python: >=3.10
8
+ Requires-Dist: soda-core==4.23.1
9
+ Requires-Dist: freezegun
10
+ Requires-Dist: pyspark[connect]>=3.5.0
11
+ Requires-Dist: soda-databricks==4.23.1
@@ -0,0 +1,11 @@
1
+ pyproject.toml
2
+ src/soda_sparkdf/__init__.py
3
+ src/soda_sparkdf.egg-info/PKG-INFO
4
+ src/soda_sparkdf.egg-info/SOURCES.txt
5
+ src/soda_sparkdf.egg-info/dependency_links.txt
6
+ src/soda_sparkdf.egg-info/entry_points.txt
7
+ src/soda_sparkdf.egg-info/requires.txt
8
+ src/soda_sparkdf.egg-info/top_level.txt
9
+ src/soda_sparkdf/common/data_sources/sparkdf_data_source.py
10
+ src/soda_sparkdf/common/data_sources/sparkdf_data_source_connection.py
11
+ src/soda_sparkdf/test_helpers/sparkdf_data_source_test_helper.py
@@ -0,0 +1,2 @@
1
+ [soda.plugins.data_source.sparkdf]
2
+ SparkDataFrameDataSourceImpl = soda_sparkdf.common.data_sources.sparkdf_data_source:SparkDataFrameDataSourceImpl
@@ -0,0 +1,4 @@
1
+ soda-core==4.23.1
2
+ freezegun
3
+ pyspark[connect]>=3.5.0
4
+ soda-databricks==4.23.1
@@ -0,0 +1 @@
1
+ soda_sparkdf