table-validator 0.1.0__py3-none-any.whl

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,1230 @@
1
+ """
2
+ Databricks SQL Warehouse Connector
3
+
4
+ Responsible solely for establishing connectivity to a Databricks SQL Warehouse
5
+ and retrieving data / schema information.
6
+
7
+ Contains no comparison logic. Every method here answers a factual question
8
+ ("does this catalog exist", "what are the null counts for these columns")
9
+ - it never decides PASS/FAIL. That decision lives in validators/catalog_validator.py.
10
+ """
11
+
12
+ from __future__ import annotations
13
+
14
+ import datetime
15
+ import logging
16
+ import numbers
17
+ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Sequence, Tuple
18
+
19
+ import pandas as pd
20
+ from databricks import sql
21
+ from databricks.sql.client import Connection
22
+
23
+ if TYPE_CHECKING:
24
+ from table_validator.models import HashCanonicalizationSpec
25
+
26
+ logger = logging.getLogger(__name__)
27
+
28
+ # numbers.Number covers int/float/Decimal AND numpy's numeric scalar types
29
+ # (np.int64, np.float64, ...) in one check - values returned by pandas/
30
+ # numpy-backed connectors (Databricks) need to compare cleanly against
31
+ # plain Decimal/int values from a DB-API driver (e.g. pyodbc), and
32
+ # Decimal.__eq__ raises TypeError rather than returning False when handed
33
+ # a numpy scalar it doesn't recognize.
34
+ _NUMERIC_TYPES = numbers.Number
35
+
36
+
37
+ def values_differ(a: Any, b: Any) -> bool:
38
+ """
39
+ Robust value comparison for two independently-fetched cells (e.g. one
40
+ row read from a source and one from a target, via separate round-trips
41
+ - a SQL query pair, or a CSV row vs a SQL row).
42
+
43
+ Raw `!=` is too strict here: two sides can return numerically-equal
44
+ values with different Python representations for the SAME underlying
45
+ type family (e.g. Decimal vs float precision, or date vs datetime),
46
+ which would otherwise register as a false mismatch against a row
47
+ already flagged as changed by a whole-row hash comparison.
48
+
49
+ Deliberately NOT applied across a numeric-vs-string type change (e.g.
50
+ a column migrated from double to string) - that IS a real, reportable
51
+ difference even if the string happens to parse to the same number.
52
+ """
53
+ if a is None or b is None:
54
+ return a is not b
55
+
56
+ if isinstance(a, _NUMERIC_TYPES) and isinstance(b, _NUMERIC_TYPES):
57
+ return abs(float(a) - float(b)) > 1e-9
58
+
59
+ if isinstance(a, (datetime.date, datetime.datetime)) and isinstance(
60
+ b, (datetime.date, datetime.datetime)
61
+ ):
62
+ a_cmp = a.date() if isinstance(a, datetime.datetime) else a
63
+ b_cmp = b.date() if isinstance(b, datetime.datetime) else b
64
+ return a_cmp != b_cmp
65
+
66
+ # A type change between the two sides (e.g. one side returned a number,
67
+ # the other a string) is itself a real difference, even if their
68
+ # string forms happen to look identical.
69
+ if type(a) is not type(b):
70
+ return True
71
+
72
+ return str(a) != str(b)
73
+
74
+
75
+ # Data types for which MIN/MAX is meaningful. Kept as a prefix match against
76
+ # the raw Databricks type string (e.g. "decimal(10,2)" -> "decimal").
77
+ _MIN_MAX_ELIGIBLE_TYPE_PREFIXES = (
78
+ "tinyint",
79
+ "smallint",
80
+ "int",
81
+ "bigint",
82
+ "float",
83
+ "double",
84
+ "decimal",
85
+ "date",
86
+ "timestamp",
87
+ )
88
+
89
+
90
+ class DatabricksConnector:
91
+ """
92
+ Lightweight reusable connector for Databricks SQL Warehouse.
93
+ """
94
+
95
+ def __init__(
96
+ self,
97
+ host: Optional[str] = None,
98
+ token: Optional[str] = None,
99
+ http_path: Optional[str] = None,
100
+ ) -> None:
101
+ """
102
+ host/token/http_path must be resolved by the caller before
103
+ construction - e.g. via table_validator.auth.databricks_auth.
104
+ get_databricks_token() for the token, and config.databricks.
105
+ workspace_url/http_path for the rest. This connector does not
106
+ read credentials from the environment itself.
107
+ """
108
+
109
+ self._host = host
110
+ self._token = token
111
+ self._http_path = http_path
112
+
113
+ if not self._host or not self._token:
114
+ raise ValueError(
115
+ "Databricks host and token are required. "
116
+ "Provide them via constructor arguments."
117
+ )
118
+
119
+ if not self._http_path:
120
+ raise ValueError(
121
+ "Databricks HTTP path is required."
122
+ )
123
+
124
+ self._connection: Optional[Connection] = None
125
+
126
+ logger.debug(
127
+ "DatabricksConnector initialized for host=%s",
128
+ self._host,
129
+ )
130
+
131
+ # ------------------------------------------------------------------
132
+ # Connection Lifecycle
133
+ # ------------------------------------------------------------------
134
+ def connect(self) -> None:
135
+
136
+ if self._connection is not None:
137
+ return
138
+
139
+ try:
140
+
141
+ self._connection = sql.connect(
142
+ server_hostname=self._host,
143
+ http_path=self._http_path,
144
+ access_token=self._token,
145
+ )
146
+
147
+ with self._connection.cursor() as cursor:
148
+ cursor.execute("SELECT 1")
149
+ cursor.fetchall()
150
+
151
+ logger.info(
152
+ "Successfully connected to Databricks SQL Warehouse"
153
+ )
154
+
155
+ except Exception as exc:
156
+
157
+ self._connection = None
158
+
159
+ logger.exception(
160
+ "Failed to connect to Databricks SQL Warehouse"
161
+ )
162
+
163
+ raise ConnectionError(
164
+ f"Unable to connect to Databricks: {exc}"
165
+ ) from exc
166
+
167
+ def disconnect(self) -> None:
168
+
169
+ if self._connection is None:
170
+ return
171
+
172
+ try:
173
+
174
+ self._connection.close()
175
+
176
+ logger.info(
177
+ "Disconnected from Databricks SQL Warehouse"
178
+ )
179
+
180
+ except Exception as exc:
181
+
182
+ logger.warning(
183
+ "Error while closing Databricks connection: %s",
184
+ exc,
185
+ )
186
+
187
+ finally:
188
+ self._connection = None
189
+
190
+ def test_connection(self) -> bool:
191
+
192
+ try:
193
+
194
+ self.connect()
195
+
196
+ with self._connection.cursor() as cursor:
197
+ cursor.execute("SELECT 1")
198
+ cursor.fetchall()
199
+
200
+ logger.info(
201
+ "Databricks connection test succeeded"
202
+ )
203
+
204
+ return True
205
+
206
+ except Exception as exc:
207
+
208
+ logger.error(
209
+ "Databricks connection test failed: %s",
210
+ exc,
211
+ )
212
+
213
+ return False
214
+
215
+ # ------------------------------------------------------------------
216
+ # Internal Helpers
217
+ # ------------------------------------------------------------------
218
+ def _ensure_connected(self) -> Connection:
219
+
220
+ if self._connection is None:
221
+ self.connect()
222
+
223
+ if self._connection is None:
224
+ raise ConnectionError(
225
+ "Databricks connection is not available"
226
+ )
227
+
228
+ return self._connection
229
+
230
+ def _execute_to_dataframe(
231
+ self,
232
+ query: str,
233
+ ) -> pd.DataFrame:
234
+
235
+ connection = self._ensure_connected()
236
+
237
+ try:
238
+
239
+ with connection.cursor() as cursor:
240
+
241
+ cursor.execute(query)
242
+
243
+ if cursor.description is None:
244
+ return pd.DataFrame()
245
+
246
+ columns = [
247
+ desc[0]
248
+ for desc in cursor.description
249
+ ]
250
+
251
+ rows = cursor.fetchall()
252
+
253
+ return pd.DataFrame(
254
+ rows,
255
+ columns=columns,
256
+ )
257
+
258
+ except Exception as exc:
259
+
260
+ logger.exception(
261
+ "Failed to execute query against Databricks"
262
+ )
263
+
264
+ raise RuntimeError(
265
+ f"Unable to execute query: {exc}"
266
+ ) from exc
267
+
268
+ @staticmethod
269
+ def _quote_ident(identifier: str) -> str:
270
+ """Backtick-quote a single identifier part, escaping embedded backticks."""
271
+ escaped = identifier.replace("`", "``")
272
+ return f"`{escaped}`"
273
+
274
+ @classmethod
275
+ def _qualify(cls, *parts: str) -> str:
276
+ """Build a fully-qualified, backtick-quoted `a`.`b`.`c` identifier."""
277
+ return ".".join(cls._quote_ident(p) for p in parts if p is not None)
278
+
279
+ # ------------------------------------------------------------------
280
+ # Data Retrieval (existing - unchanged)
281
+ # ------------------------------------------------------------------
282
+ def read_table(
283
+ self,
284
+ table_name: str,
285
+ ) -> pd.DataFrame:
286
+
287
+ if not table_name or not table_name.strip():
288
+ raise ValueError(
289
+ "table_name must be a non-empty string"
290
+ )
291
+
292
+ safe_name = ".".join(
293
+ f"`{part}`"
294
+ for part in table_name.split(".")
295
+ )
296
+
297
+ query = f"SELECT * FROM {safe_name}"
298
+
299
+ logger.info(
300
+ "Reading table '%s' from Databricks",
301
+ table_name,
302
+ )
303
+
304
+ df = self._execute_to_dataframe(query)
305
+
306
+ logger.info(
307
+ "Successfully read table '%s' - shape=%s",
308
+ table_name,
309
+ df.shape,
310
+ )
311
+
312
+ return df
313
+
314
+ def read_query(
315
+ self,
316
+ query: str,
317
+ ) -> pd.DataFrame:
318
+
319
+ if not query or not query.strip():
320
+ raise ValueError(
321
+ "query must be a non-empty string"
322
+ )
323
+
324
+ logger.info(
325
+ "Executing custom query against Databricks"
326
+ )
327
+
328
+ df = self._execute_to_dataframe(query)
329
+
330
+ logger.info(
331
+ "Query returned shape=%s",
332
+ df.shape,
333
+ )
334
+
335
+ return df
336
+
337
+ def get_schema(
338
+ self,
339
+ table_name: str,
340
+ ) -> pd.DataFrame:
341
+
342
+ if not table_name or not table_name.strip():
343
+ raise ValueError(
344
+ "table_name must be a non-empty string"
345
+ )
346
+
347
+ safe_name = ".".join(
348
+ f"`{part}`"
349
+ for part in table_name.split(".")
350
+ )
351
+
352
+ describe_query = (
353
+ f"DESCRIBE TABLE {safe_name}"
354
+ )
355
+
356
+ try:
357
+
358
+ raw_df = self._execute_to_dataframe(
359
+ describe_query
360
+ )
361
+
362
+ if raw_df.empty:
363
+
364
+ return pd.DataFrame(
365
+ columns=[
366
+ "column_name",
367
+ "data_type",
368
+ "is_nullable",
369
+ "character_maximum_length",
370
+ ]
371
+ )
372
+
373
+ raw_df = raw_df[
374
+ raw_df["col_name"].notna()
375
+ ]
376
+
377
+ raw_df = raw_df[
378
+ raw_df["col_name"] != ""
379
+ ]
380
+
381
+ raw_df = raw_df[
382
+ ~raw_df["col_name"].astype(str).str.startswith("#")
383
+ ]
384
+
385
+ schema_df = pd.DataFrame(
386
+ {
387
+ "column_name": raw_df["col_name"],
388
+ "data_type": raw_df["data_type"],
389
+ "is_nullable": "YES",
390
+ "character_maximum_length": None,
391
+ }
392
+ )
393
+
394
+ logger.info(
395
+ "Schema for '%s' retrieved - %d columns",
396
+ table_name,
397
+ len(schema_df),
398
+ )
399
+
400
+ return schema_df.reset_index(drop=True)
401
+
402
+ except Exception as exc:
403
+
404
+ logger.exception(
405
+ "Failed to retrieve schema for '%s'",
406
+ table_name,
407
+ )
408
+
409
+ raise RuntimeError(
410
+ f"Unable to retrieve schema for '{table_name}': {exc}"
411
+ ) from exc
412
+
413
+ # ------------------------------------------------------------------
414
+ # NEW: generic passthrough (public alias, used by CatalogValidator
415
+ # for anything not covered by a dedicated method below)
416
+ # ------------------------------------------------------------------
417
+ def execute_query(self, query: str) -> pd.DataFrame:
418
+ """Public entry point for executing an arbitrary read-only query."""
419
+ return self._execute_to_dataframe(query)
420
+
421
+ # ------------------------------------------------------------------
422
+ # NEW: Catalog / schema / table metadata methods
423
+ # ------------------------------------------------------------------
424
+ def catalog_exists(self, catalog: str) -> bool:
425
+ try:
426
+ df = self._execute_to_dataframe("SHOW CATALOGS")
427
+ except Exception as exc:
428
+ logger.exception("Failed to list catalogs")
429
+ raise RuntimeError(f"Unable to list catalogs: {exc}") from exc
430
+
431
+ if df.empty:
432
+ return False
433
+
434
+ col = df.columns[0]
435
+ return catalog.lower() in {str(v).lower() for v in df[col]}
436
+
437
+ def get_schemas(self, catalog: str) -> List[str]:
438
+ try:
439
+ df = self._execute_to_dataframe(
440
+ f"SHOW SCHEMAS IN {self._quote_ident(catalog)}"
441
+ )
442
+ except Exception as exc:
443
+ logger.exception("Failed to list schemas for catalog '%s'", catalog)
444
+ raise RuntimeError(
445
+ f"Unable to list schemas for catalog '{catalog}': {exc}"
446
+ ) from exc
447
+
448
+ if df.empty:
449
+ return []
450
+
451
+ col = "databaseName" if "databaseName" in df.columns else df.columns[0]
452
+ return sorted(str(v) for v in df[col])
453
+
454
+ def schema_exists(self, catalog: str, schema: str) -> bool:
455
+ return schema.lower() in {s.lower() for s in self.get_schemas(catalog)}
456
+
457
+ def get_tables(self, catalog: str, schema: str) -> List[str]:
458
+ try:
459
+ df = self._execute_to_dataframe(
460
+ f"SHOW TABLES IN {self._qualify(catalog, schema)}"
461
+ )
462
+ except Exception as exc:
463
+ logger.exception(
464
+ "Failed to list tables for '%s.%s'", catalog, schema
465
+ )
466
+ raise RuntimeError(
467
+ f"Unable to list tables for '{catalog}.{schema}': {exc}"
468
+ ) from exc
469
+
470
+ if df.empty:
471
+ return []
472
+
473
+ col = "tableName" if "tableName" in df.columns else df.columns[0]
474
+ return sorted(str(v) for v in df[col])
475
+
476
+ def table_exists(self, catalog: str, schema: str, table: str) -> bool:
477
+ return table.lower() in {t.lower() for t in self.get_tables(catalog, schema)}
478
+
479
+ def get_table_schema(
480
+ self,
481
+ catalog: str,
482
+ schema: str,
483
+ table: str,
484
+ ) -> pd.DataFrame:
485
+ """
486
+ Returns columns: column_name, data_type, is_nullable (bool),
487
+ ordinal_position - sourced from information_schema, which (unlike
488
+ DESCRIBE TABLE) gives real nullability and a reliable column order.
489
+ """
490
+ query = f"""
491
+ SELECT column_name, full_data_type AS data_type,
492
+ is_nullable, ordinal_position
493
+ FROM {self._quote_ident(catalog)}.information_schema.columns
494
+ WHERE table_schema = '{schema}' AND table_name = '{table}'
495
+ ORDER BY ordinal_position
496
+ """
497
+
498
+ try:
499
+ df = self._execute_to_dataframe(query)
500
+ except Exception as exc:
501
+ logger.exception(
502
+ "Failed to retrieve column metadata for '%s.%s.%s'",
503
+ catalog, schema, table,
504
+ )
505
+ raise RuntimeError(
506
+ f"Unable to retrieve column metadata for "
507
+ f"'{catalog}.{schema}.{table}': {exc}"
508
+ ) from exc
509
+
510
+ if df.empty:
511
+ return pd.DataFrame(
512
+ columns=["column_name", "data_type", "is_nullable", "ordinal_position"]
513
+ )
514
+
515
+ df["is_nullable"] = df["is_nullable"].astype(str).str.upper().eq("YES")
516
+ return df.reset_index(drop=True)
517
+
518
+ def get_row_count(self, catalog: str, schema: str, table: str) -> int:
519
+ query = f"SELECT COUNT(*) AS row_count FROM {self._qualify(catalog, schema, table)}"
520
+ try:
521
+ df = self._execute_to_dataframe(query)
522
+ except Exception as exc:
523
+ logger.exception(
524
+ "Failed to get row count for '%s.%s.%s'", catalog, schema, table
525
+ )
526
+ raise RuntimeError(
527
+ f"Unable to get row count for '{catalog}.{schema}.{table}': {exc}"
528
+ ) from exc
529
+
530
+ if df.empty:
531
+ return 0
532
+ return int(df.iloc[0]["row_count"])
533
+
534
+ def get_column_statistics(
535
+ self,
536
+ catalog: str,
537
+ schema: str,
538
+ table: str,
539
+ columns: Sequence[str],
540
+ min_max_columns: Optional[Sequence[str]] = None,
541
+ ) -> Dict[str, Dict[str, Any]]:
542
+ """
543
+ Single aggregate query returning null count, distinct count, and
544
+ (for min_max_columns) MIN/MAX for every requested column - avoids
545
+ one round-trip per column.
546
+
547
+ Returns: {column_name: {"null_count": int, "distinct_count": int,
548
+ "min": Any | None, "max": Any | None}}
549
+ """
550
+ if not columns:
551
+ return {}
552
+
553
+ min_max_set = {c.lower() for c in (min_max_columns or [])}
554
+
555
+ select_parts = []
556
+ for col in columns:
557
+ q = self._quote_ident(col)
558
+ select_parts.append(f"SUM(CASE WHEN {q} IS NULL THEN 1 ELSE 0 END) AS `{col}__nulls`")
559
+ select_parts.append(f"COUNT(DISTINCT {q}) AS `{col}__distinct`")
560
+ if col.lower() in min_max_set:
561
+ select_parts.append(f"MIN({q}) AS `{col}__min`")
562
+ select_parts.append(f"MAX({q}) AS `{col}__max`")
563
+
564
+ query = (
565
+ f"SELECT {', '.join(select_parts)} "
566
+ f"FROM {self._qualify(catalog, schema, table)}"
567
+ )
568
+
569
+ try:
570
+ df = self._execute_to_dataframe(query)
571
+ except Exception as exc:
572
+ logger.exception(
573
+ "Failed to compute column statistics for '%s.%s.%s'",
574
+ catalog, schema, table,
575
+ )
576
+ raise RuntimeError(
577
+ f"Unable to compute column statistics for "
578
+ f"'{catalog}.{schema}.{table}': {exc}"
579
+ ) from exc
580
+
581
+ result: Dict[str, Dict[str, Any]] = {}
582
+
583
+ if df.empty:
584
+ return {col: {"null_count": None, "distinct_count": None,
585
+ "min": None, "max": None} for col in columns}
586
+
587
+ row = df.iloc[0]
588
+
589
+ for col in columns:
590
+ entry: Dict[str, Any] = {
591
+ "null_count": int(row.get(f"{col}__nulls"))
592
+ if row.get(f"{col}__nulls") is not None else None,
593
+ "distinct_count": int(row.get(f"{col}__distinct"))
594
+ if row.get(f"{col}__distinct") is not None else None,
595
+ "min": None,
596
+ "max": None,
597
+ }
598
+ if col.lower() in min_max_set:
599
+ entry["min"] = row.get(f"{col}__min")
600
+ entry["max"] = row.get(f"{col}__max")
601
+ result[col] = entry
602
+
603
+ return result
604
+
605
+ @staticmethod
606
+ def is_min_max_eligible(data_type: str) -> bool:
607
+ dt = (data_type or "").strip().lower()
608
+ return any(dt.startswith(prefix) for prefix in _MIN_MAX_ELIGIBLE_TYPE_PREFIXES)
609
+
610
+ def _row_hash_expr(
611
+ self,
612
+ columns: Sequence[str],
613
+ spec: Optional["HashCanonicalizationSpec"] = None,
614
+ ) -> str:
615
+ """
616
+ Build the canonical per-row hash SQL expression shared by
617
+ get_row_hashes, get_row_hashes_by_row_number, and
618
+ get_table_fingerprint, so all three tiers hash identically.
619
+
620
+ Defaults reproduce the original inline expression byte-for-byte:
621
+ sha2(concat_ws('||', COALESCE(CAST(col AS STRING), sentinel)...), 256).
622
+ `spec` is accepted for forward compatibility with
623
+ HashCanonicalizationSpec but is not yet applied here - see
624
+ models.HashCanonicalizationSpec docstring.
625
+ """
626
+ null_sentinel = (spec.null_sentinel if spec else None) or "\x01NULL\x01"
627
+
628
+ hashed_exprs = [
629
+ f"COALESCE(CAST({self._quote_ident(c)} AS STRING), '{null_sentinel}')"
630
+ for c in columns
631
+ ]
632
+
633
+ if hashed_exprs:
634
+ return f"sha2(concat_ws('||', {', '.join(hashed_exprs)}), 256)"
635
+ return f"sha2('{null_sentinel}', 256)"
636
+
637
+ def get_table_fingerprint(
638
+ self,
639
+ catalog: str,
640
+ schema: str,
641
+ table: str,
642
+ columns: Sequence[str],
643
+ spec: Optional["HashCanonicalizationSpec"] = None,
644
+ ) -> Dict[str, Any]:
645
+ """
646
+ Tier 2: single order-independent whole-table fingerprint, computed
647
+ entirely server-side - no row data ever leaves the warehouse.
648
+
649
+ Combines three aggregates that are individually weak but strong
650
+ together: COUNT(*) alone misses swapped/altered values; SUM alone
651
+ is collision-prone/overflow-prone; XOR alone is blind to
652
+ duplicated rows. A 15-hex-char (60-bit) prefix of the per-row hash
653
+ is converted to a numeric value via conv(hex, 16, 10) - Databricks
654
+ SQL has no native hex-to-numeric cast - which keeps the XOR
655
+ argument within BIGINT range, while the SUM accumulates as
656
+ DECIMAL(38,0) to stay overflow-safe across an entire table.
657
+
658
+ Returns {"row_count": int, "hash_sum": Decimal|None, "hash_xor": int|None}.
659
+ """
660
+ row_hash_expr = self._row_hash_expr(columns, spec)
661
+ hash_prefix = f"substr({row_hash_expr}, 1, 15)"
662
+
663
+ query = f"""
664
+ SELECT
665
+ COUNT(*) AS row_count,
666
+ SUM(CAST(conv({hash_prefix}, 16, 10) AS DECIMAL(38,0))) AS hash_sum,
667
+ BIT_XOR(CAST(conv({hash_prefix}, 16, 10) AS BIGINT)) AS hash_xor
668
+ FROM {self._qualify(catalog, schema, table)}
669
+ """
670
+
671
+ try:
672
+ df = self._execute_to_dataframe(query)
673
+ except Exception as exc:
674
+ logger.exception(
675
+ "Failed to compute table fingerprint for '%s.%s.%s'",
676
+ catalog, schema, table,
677
+ )
678
+ raise RuntimeError(
679
+ f"Unable to compute table fingerprint for "
680
+ f"'{catalog}.{schema}.{table}': {exc}"
681
+ ) from exc
682
+
683
+ if df.empty:
684
+ return {"row_count": 0, "hash_sum": None, "hash_xor": None}
685
+
686
+ row = df.iloc[0]
687
+ return {
688
+ "row_count": int(row.get("row_count") or 0),
689
+ "hash_sum": row.get("hash_sum"),
690
+ "hash_xor": row.get("hash_xor"),
691
+ }
692
+
693
+ def get_table_fingerprint_by_bucket(
694
+ self,
695
+ catalog: str,
696
+ schema: str,
697
+ table: str,
698
+ columns: Sequence[str],
699
+ bucket_column: str,
700
+ spec: Optional["HashCanonicalizationSpec"] = None,
701
+ ) -> pd.DataFrame:
702
+ """
703
+ Tier 3: the same triple fingerprint as get_table_fingerprint, but
704
+ GROUP BY a chosen bucket column - one row per distinct bucket
705
+ value, computed entirely server-side. Used to narrow a confirmed
706
+ table-level mismatch down to the specific bucket(s) that actually
707
+ differ, so Tier 4's row-hash diff only needs to scan those buckets
708
+ instead of the whole table.
709
+
710
+ Returns a DataFrame with columns: bucket_value, row_count,
711
+ hash_sum, hash_xor - one row per distinct value of bucket_column
712
+ present on this side.
713
+ """
714
+ row_hash_expr = self._row_hash_expr(columns, spec)
715
+ hash_prefix = f"substr({row_hash_expr}, 1, 15)"
716
+ bucket_ident = self._quote_ident(bucket_column)
717
+
718
+ query = f"""
719
+ SELECT
720
+ {bucket_ident} AS bucket_value,
721
+ COUNT(*) AS row_count,
722
+ SUM(CAST(conv({hash_prefix}, 16, 10) AS DECIMAL(38,0))) AS hash_sum,
723
+ BIT_XOR(CAST(conv({hash_prefix}, 16, 10) AS BIGINT)) AS hash_xor
724
+ FROM {self._qualify(catalog, schema, table)}
725
+ GROUP BY {bucket_ident}
726
+ """
727
+
728
+ try:
729
+ df = self._execute_to_dataframe(query)
730
+ except Exception as exc:
731
+ logger.exception(
732
+ "Failed to compute bucketed table fingerprint for '%s.%s.%s' "
733
+ "(bucket_column='%s')",
734
+ catalog, schema, table, bucket_column,
735
+ )
736
+ raise RuntimeError(
737
+ f"Unable to compute bucketed table fingerprint for "
738
+ f"'{catalog}.{schema}.{table}' (bucket_column='{bucket_column}'): {exc}"
739
+ ) from exc
740
+
741
+ if df.empty:
742
+ return pd.DataFrame(columns=["bucket_value", "row_count", "hash_sum", "hash_xor"])
743
+
744
+ return df
745
+
746
+ def key_based_row_diff(
747
+ self,
748
+ source_fqtn: str,
749
+ target_fqtn: str,
750
+ key_columns: Sequence[str],
751
+ value_columns: Sequence[str],
752
+ limit_samples: int = 50,
753
+ ) -> Dict[str, Any]:
754
+ """
755
+ Push-down key-based row comparison between two fully-qualified
756
+ tables (e.g. 'cat_a.schema.table' vs 'cat_b.schema.table').
757
+
758
+ Returns counts of source-only / target-only rows (via SQL EXCEPT
759
+ on key columns) and changed-row count (matching key, differing
760
+ row hash), plus a small bounded sample of each - never a full
761
+ collect() of either table.
762
+ """
763
+ key_idents = [self._quote_ident(k) for k in key_columns]
764
+ key_list = ", ".join(key_idents)
765
+
766
+ src = ".".join(self._quote_ident(p) for p in source_fqtn.split("."))
767
+ tgt = ".".join(self._quote_ident(p) for p in target_fqtn.split("."))
768
+
769
+ # Source-only / target-only keys via EXCEPT (fully pushed down)
770
+ source_only_query = f"""
771
+ SELECT {key_list} FROM {src}
772
+ EXCEPT
773
+ SELECT {key_list} FROM {tgt}
774
+ """
775
+ target_only_query = f"""
776
+ SELECT {key_list} FROM {tgt}
777
+ EXCEPT
778
+ SELECT {key_list} FROM {src}
779
+ """
780
+
781
+ source_only_df = self._execute_to_dataframe(
782
+ f"SELECT COUNT(*) AS c FROM ({source_only_query}) x"
783
+ )
784
+ target_only_df = self._execute_to_dataframe(
785
+ f"SELECT COUNT(*) AS c FROM ({target_only_query}) x"
786
+ )
787
+
788
+ source_only_count = int(source_only_df.iloc[0]["c"]) if not source_only_df.empty else 0
789
+ target_only_count = int(target_only_df.iloc[0]["c"]) if not target_only_df.empty else 0
790
+
791
+ sample_source_only = self._execute_to_dataframe(
792
+ f"{source_only_query} LIMIT {int(limit_samples)}"
793
+ ).to_dict(orient="records")
794
+ sample_target_only = self._execute_to_dataframe(
795
+ f"{target_only_query} LIMIT {int(limit_samples)}"
796
+ ).to_dict(orient="records")
797
+
798
+ # Changed rows: matching key, differing hash of value columns
799
+ changed_count = 0
800
+ sample_changed: List[Dict[str, Any]] = []
801
+ sample_changed_detail: List[Dict[str, Any]] = []
802
+
803
+ if value_columns:
804
+ value_concat = ", ".join(self._quote_ident(c) for c in value_columns)
805
+ changed_query = f"""
806
+ SELECT {key_list} FROM (
807
+ SELECT {key_list}, hash({value_concat}) AS __row_hash
808
+ FROM {src}
809
+ ) s
810
+ JOIN (
811
+ SELECT {key_list}, hash({value_concat}) AS __row_hash
812
+ FROM {tgt}
813
+ ) t
814
+ USING ({key_list})
815
+ WHERE s.__row_hash != t.__row_hash
816
+ """
817
+ changed_df = self._execute_to_dataframe(
818
+ f"SELECT COUNT(*) AS c FROM ({changed_query}) x"
819
+ )
820
+ changed_count = int(changed_df.iloc[0]["c"]) if not changed_df.empty else 0
821
+
822
+ sample_changed = self._execute_to_dataframe(
823
+ f"{changed_query} LIMIT {int(limit_samples)}"
824
+ ).to_dict(orient="records")
825
+
826
+ sample_changed_detail = self._changed_row_detail(
827
+ src=src,
828
+ tgt=tgt,
829
+ key_columns=key_columns,
830
+ key_idents=key_idents,
831
+ key_list=key_list,
832
+ value_columns=value_columns,
833
+ changed_query=changed_query,
834
+ limit_samples=limit_samples,
835
+ )
836
+
837
+ return {
838
+ "source_only_rows": source_only_count,
839
+ "target_only_rows": target_only_count,
840
+ "changed_rows": changed_count,
841
+ "sample_source_only": sample_source_only,
842
+ "sample_target_only": sample_target_only,
843
+ "sample_changed": sample_changed,
844
+ "sample_changed_detail": sample_changed_detail,
845
+ }
846
+
847
+ def get_row_detail_for_keys(
848
+ self,
849
+ source_catalog: str,
850
+ target_catalog: str,
851
+ schema: str,
852
+ table: str,
853
+ key_column: str,
854
+ key_values: Sequence[str],
855
+ value_columns: Sequence[str],
856
+ limit_samples: int = 500,
857
+ ) -> List[Dict[str, Any]]:
858
+ """
859
+ Tier 5: column-level diff for a bounded, already-known set of
860
+ mismatched keys (single-column key only). Fetches key + value
861
+ columns plus a whole-row hash() from both sides for exactly
862
+ those keys - never a full-table pull - and diffs them
863
+ column-by-column so callers can report exactly which column(s)
864
+ differ per row. `key_values` are treated as opaque string literals
865
+ (matching Tier 4's compare_row_hashes display-key convention).
866
+ """
867
+ if not key_values:
868
+ return []
869
+
870
+ src = self._qualify(source_catalog, schema, table)
871
+ tgt = self._qualify(target_catalog, schema, table)
872
+ key_ident = self._quote_ident(key_column)
873
+ quoted_values = ", ".join(f"'{str(v).replace(chr(39), chr(39) * 2)}'" for v in key_values)
874
+ changed_query = (
875
+ f"SELECT {key_ident} FROM {src} "
876
+ f"WHERE CAST({key_ident} AS STRING) IN ({quoted_values})"
877
+ )
878
+
879
+ return self._changed_row_detail(
880
+ src=src,
881
+ tgt=tgt,
882
+ key_columns=[key_column],
883
+ key_idents=[key_ident],
884
+ key_list=key_ident,
885
+ value_columns=value_columns,
886
+ changed_query=changed_query,
887
+ limit_samples=limit_samples,
888
+ )
889
+
890
+ def _changed_row_detail(
891
+ self,
892
+ src: str,
893
+ tgt: str,
894
+ key_columns: Sequence[str],
895
+ key_idents: List[str],
896
+ key_list: str,
897
+ value_columns: Sequence[str],
898
+ changed_query: str,
899
+ limit_samples: int,
900
+ ) -> List[Dict[str, Any]]:
901
+ """
902
+ For a bounded sample of changed keys (from `changed_query`), fetch
903
+ the full source and target rows (key + value columns) plus a
904
+ whole-row hash for each side, so callers can report exactly which
905
+ column(s) differ per row without ever collecting a full table.
906
+ """
907
+ value_idents = [self._quote_ident(c) for c in value_columns]
908
+ all_idents = key_idents + value_idents
909
+ select_list = ", ".join(all_idents)
910
+ value_concat = ", ".join(value_idents)
911
+
912
+ source_rows = self._execute_to_dataframe(f"""
913
+ SELECT {select_list}, hash({value_concat}) AS __row_hash
914
+ FROM {src}
915
+ WHERE ({key_list}) IN (SELECT {key_list} FROM ({changed_query} LIMIT {int(limit_samples)}) __k)
916
+ """).to_dict(orient="records")
917
+
918
+ target_rows = self._execute_to_dataframe(f"""
919
+ SELECT {select_list}, hash({value_concat}) AS __row_hash
920
+ FROM {tgt}
921
+ WHERE ({key_list}) IN (SELECT {key_list} FROM ({changed_query} LIMIT {int(limit_samples)}) __k)
922
+ """).to_dict(orient="records")
923
+
924
+ def _key_tuple(row: Dict[str, Any]) -> tuple:
925
+ return tuple(row.get(k) for k in key_columns)
926
+
927
+ target_by_key = {_key_tuple(r): r for r in target_rows}
928
+
929
+ detail: List[Dict[str, Any]] = []
930
+ for src_row in source_rows:
931
+ tgt_row = target_by_key.get(_key_tuple(src_row))
932
+ if tgt_row is None:
933
+ continue
934
+
935
+ mismatched_columns = [
936
+ col for col in value_columns
937
+ if values_differ(src_row.get(col), tgt_row.get(col))
938
+ ]
939
+ if not mismatched_columns:
940
+ # SQL-side hash() flagged this row as changed, but our
941
+ # tolerant per-column comparison found nothing - the two
942
+ # comparisons disagree (e.g. a difference the hash catches
943
+ # that our value comparison normalizes away). Report the
944
+ # row anyway rather than silently dropping a row the
945
+ # mismatch count already accounts for.
946
+ mismatched_columns = list(value_columns)
947
+
948
+ detail.append(
949
+ {
950
+ "key": {k: src_row.get(k) for k in key_columns},
951
+ "mismatched_columns": mismatched_columns,
952
+ "source_values": {c: src_row.get(c) for c in mismatched_columns},
953
+ "target_values": {c: tgt_row.get(c) for c in mismatched_columns},
954
+ "source_row_hash": src_row.get("__row_hash"),
955
+ "target_row_hash": tgt_row.get("__row_hash"),
956
+ }
957
+ )
958
+
959
+ return detail
960
+
961
+ def _bucket_where_clause(
962
+ self,
963
+ bucket_predicate: Optional[Tuple[str, Any]],
964
+ ) -> str:
965
+ """
966
+ Build a `WHERE {col} = {value}` clause (or `WHERE {col} IS NULL`
967
+ for a null bucket value) scoping a query to exactly one partition
968
+ bucket, or an empty string when no predicate is given (whole-table
969
+ query, today's default behavior). The value is treated as an
970
+ opaque literal from a bucket-fingerprint query's own result - it
971
+ did not come from user input, but is still quoted defensively.
972
+ """
973
+ if bucket_predicate is None:
974
+ return ""
975
+ column, value = bucket_predicate
976
+ ident = self._quote_ident(column)
977
+ if value is None:
978
+ return f"WHERE {ident} IS NULL"
979
+ escaped = str(value).replace("'", "''")
980
+ return f"WHERE CAST({ident} AS STRING) = '{escaped}'"
981
+
982
+ def get_row_hashes(
983
+ self,
984
+ catalog: str,
985
+ schema: str,
986
+ table: str,
987
+ columns: Sequence[str],
988
+ primary_key_cols: Sequence[str],
989
+ bucket_predicate: Optional[Tuple[str, Any]] = None,
990
+ ) -> pd.DataFrame:
991
+ """
992
+ Single push-down query returning one deterministic row hash per
993
+ primary key value(s). `columns` is the fixed, already-sorted list
994
+ of business columns (PK excluded) to hash - callers must pass the
995
+ SAME order for both source and target so the hashes are directly
996
+ comparable. Never pulls row data into pandas beyond the key(s) and
997
+ the resulting hash column.
998
+
999
+ Each column is COALESCE(CAST(col AS STRING), sentinel)'d before
1000
+ concatenation so NULLs hash consistently and never collapse into
1001
+ an empty-string collision with a genuinely empty string value.
1002
+
1003
+ `bucket_predicate`, when given as (column, value), scopes the
1004
+ query to exactly one partition bucket (Tier 3) instead of the
1005
+ whole table - this is what makes a partitioned Tier 4 cheaper
1006
+ than an unpartitioned one.
1007
+
1008
+ Returns a DataFrame with one row per primary key: the key column(s)
1009
+ plus `row_hash`.
1010
+ """
1011
+ if not primary_key_cols:
1012
+ raise ValueError("primary_key_cols must be non-empty")
1013
+
1014
+ key_list = ", ".join(self._quote_ident(k) for k in primary_key_cols)
1015
+
1016
+ row_hash_expr = self._row_hash_expr(columns)
1017
+ where_clause = self._bucket_where_clause(bucket_predicate)
1018
+
1019
+ query = f"""
1020
+ SELECT {key_list}, {row_hash_expr} AS row_hash
1021
+ FROM {self._qualify(catalog, schema, table)}
1022
+ {where_clause}
1023
+ """
1024
+
1025
+ try:
1026
+ df = self._execute_to_dataframe(query)
1027
+ except Exception as exc:
1028
+ logger.exception(
1029
+ "Failed to compute row hashes for '%s.%s.%s'", catalog, schema, table
1030
+ )
1031
+ raise RuntimeError(
1032
+ f"Unable to compute row hashes for '{catalog}.{schema}.{table}': {exc}"
1033
+ ) from exc
1034
+
1035
+ if df.empty:
1036
+ return pd.DataFrame(columns=list(primary_key_cols) + ["row_hash"])
1037
+
1038
+ return df
1039
+
1040
+ def get_row_hashes_by_row_number(
1041
+ self,
1042
+ catalog: str,
1043
+ schema: str,
1044
+ table: str,
1045
+ columns: Sequence[str],
1046
+ bucket_predicate: Optional[Tuple[str, Any]] = None,
1047
+ ) -> pd.DataFrame:
1048
+ """
1049
+ Fallback for tables with no configured primary key: assigns a
1050
+ synthetic row number via ROW_NUMBER() OVER (ORDER BY <every
1051
+ requested column>) on both sides, then hashes each row the same
1052
+ way get_row_hashes does. Ordering by every column (not insertion
1053
+ order, which SQL never guarantees) means two logically-identical
1054
+ rows always sort next to each other and get matching numbers
1055
+ regardless of physical storage order - but this is NOT a
1056
+ substitute for a real key: if the two sides don't contain the
1057
+ same *set* of rows, row N on one side is not necessarily the same
1058
+ logical record as row N on the other, and comparisons will be
1059
+ misleading. Only use when no real shared key exists.
1060
+
1061
+ `bucket_predicate`, when given as (column, value), scopes both the
1062
+ ROW_NUMBER() sort and the hash computation to exactly one
1063
+ partition bucket (Tier 3) rather than the whole table - row
1064
+ numbers are still only comparable within the same bucket on both
1065
+ sides, which is exactly the intended scope here.
1066
+
1067
+ Returns a DataFrame with columns: row_number, row_hash.
1068
+ """
1069
+ if not columns:
1070
+ raise ValueError("columns must be non-empty for row-number based hashing")
1071
+
1072
+ order_by = ", ".join(self._quote_ident(c) for c in columns)
1073
+ row_hash_expr = self._row_hash_expr(columns)
1074
+ where_clause = self._bucket_where_clause(bucket_predicate)
1075
+
1076
+ query = f"""
1077
+ SELECT
1078
+ ROW_NUMBER() OVER (ORDER BY {order_by}) AS row_number,
1079
+ {row_hash_expr} AS row_hash
1080
+ FROM {self._qualify(catalog, schema, table)}
1081
+ {where_clause}
1082
+ """
1083
+
1084
+ try:
1085
+ df = self._execute_to_dataframe(query)
1086
+ except Exception as exc:
1087
+ logger.exception(
1088
+ "Failed to compute row-number-based hashes for '%s.%s.%s'", catalog, schema, table
1089
+ )
1090
+ raise RuntimeError(
1091
+ f"Unable to compute row-number-based hashes for '{catalog}.{schema}.{table}': {exc}"
1092
+ ) from exc
1093
+
1094
+ if df.empty:
1095
+ return pd.DataFrame(columns=["row_number", "row_hash"])
1096
+
1097
+ return df
1098
+
1099
+ def get_row_detail_for_row_numbers(
1100
+ self,
1101
+ source_catalog: str,
1102
+ target_catalog: str,
1103
+ schema: str,
1104
+ table: str,
1105
+ order_by_columns: Sequence[str],
1106
+ row_numbers: Sequence[int],
1107
+ value_columns: Sequence[str],
1108
+ limit_samples: int = 500,
1109
+ bucket_predicate: Optional[Tuple[str, Any]] = None,
1110
+ ) -> List[Dict[str, Any]]:
1111
+ """
1112
+ Best-effort Tier 5 column-level diff for the ROW_NUMBER() fallback
1113
+ (no primary key configured). Re-executes the SAME
1114
+ ROW_NUMBER() OVER (ORDER BY <order_by_columns>) window per side
1115
+ used by get_row_hashes_by_row_number, filtered down to the given
1116
+ row numbers, then diffs the fetched rows column-by-column.
1117
+
1118
+ `order_by_columns` MUST be the exact same column list (same order)
1119
+ passed to get_row_hashes_by_row_number for this table - that is
1120
+ what keeps row numbers consistent between the hash-computation
1121
+ pass (Tier 4) and this re-fetch (Tier 5). `value_columns` is
1122
+ normally the same list too, since the row-number fallback has no
1123
+ key to exclude.
1124
+
1125
+ This is inherently best-effort, not a substitute for a real key:
1126
+ "row N" on the source and target are only the same logical record
1127
+ if both sides otherwise contain the same row set in the same
1128
+ relative order. Callers must mark any result derived from this
1129
+ method as unverified (see RowMismatchDetail.verified).
1130
+
1131
+ A window-function output column (ROW_NUMBER() here) can only be
1132
+ filtered in a query that reads it as a plain column from a
1133
+ subquery - it cannot be filtered in the same SELECT that computes
1134
+ it via OVER(). Hence the subquery/CTE shape below, rather than a
1135
+ flat SELECT ... WHERE row_number IN (...).
1136
+
1137
+ Returns the same shape as _changed_row_detail: one dict per
1138
+ row with "key" (here always {"row_number": N}),
1139
+ "mismatched_columns", "source_values", "target_values",
1140
+ "source_row_hash", "target_row_hash".
1141
+ """
1142
+ if not row_numbers:
1143
+ return []
1144
+
1145
+ src = self._qualify(source_catalog, schema, table)
1146
+ tgt = self._qualify(target_catalog, schema, table)
1147
+ order_by = ", ".join(self._quote_ident(c) for c in order_by_columns)
1148
+ value_idents = [self._quote_ident(c) for c in value_columns]
1149
+ select_list = ", ".join(value_idents)
1150
+ where_clause = self._bucket_where_clause(bucket_predicate)
1151
+ row_hash_expr = self._row_hash_expr(value_columns)
1152
+
1153
+ # row_numbers are Python ints derived from our own prior
1154
+ # ROW_NUMBER() output (never user input) - safe to inline.
1155
+ unique_row_numbers = sorted(set(int(n) for n in row_numbers))[: int(limit_samples)]
1156
+ row_numbers_csv = ", ".join(str(n) for n in unique_row_numbers)
1157
+
1158
+ def _numbered_query(fqtn: str) -> str:
1159
+ return f"""
1160
+ SELECT row_number, {select_list}, {row_hash_expr} AS __row_hash
1161
+ FROM (
1162
+ SELECT
1163
+ ROW_NUMBER() OVER (ORDER BY {order_by}) AS row_number,
1164
+ {select_list}
1165
+ FROM {fqtn}
1166
+ {where_clause}
1167
+ ) numbered
1168
+ WHERE row_number IN ({row_numbers_csv})
1169
+ """
1170
+
1171
+ try:
1172
+ source_rows = self._execute_to_dataframe(
1173
+ _numbered_query(src)
1174
+ ).to_dict(orient="records")
1175
+ target_rows = self._execute_to_dataframe(
1176
+ _numbered_query(tgt)
1177
+ ).to_dict(orient="records")
1178
+ except Exception as exc:
1179
+ logger.exception(
1180
+ "Failed to fetch row-number-based row detail for '%s.%s'", schema, table
1181
+ )
1182
+ raise RuntimeError(
1183
+ f"Unable to fetch row-number-based row detail for '{schema}.{table}': {exc}"
1184
+ ) from exc
1185
+
1186
+ target_by_row_number = {r["row_number"]: r for r in target_rows}
1187
+
1188
+ detail: List[Dict[str, Any]] = []
1189
+ for src_row in source_rows:
1190
+ tgt_row = target_by_row_number.get(src_row["row_number"])
1191
+ if tgt_row is None:
1192
+ continue
1193
+
1194
+ mismatched_columns = [
1195
+ col for col in value_columns
1196
+ if values_differ(src_row.get(col), tgt_row.get(col))
1197
+ ]
1198
+ if not mismatched_columns:
1199
+ mismatched_columns = list(value_columns)
1200
+
1201
+ detail.append(
1202
+ {
1203
+ "key": {"row_number": src_row["row_number"]},
1204
+ "mismatched_columns": mismatched_columns,
1205
+ "source_values": {c: src_row.get(c) for c in mismatched_columns},
1206
+ "target_values": {c: tgt_row.get(c) for c in mismatched_columns},
1207
+ "source_row_hash": src_row.get("__row_hash"),
1208
+ "target_row_hash": tgt_row.get("__row_hash"),
1209
+ }
1210
+ )
1211
+
1212
+ return detail
1213
+
1214
+ # ------------------------------------------------------------------
1215
+ # Context Manager Support
1216
+ # ------------------------------------------------------------------
1217
+ def __enter__(self) -> "DatabricksConnector":
1218
+
1219
+ self.connect()
1220
+
1221
+ return self
1222
+
1223
+ def __exit__(
1224
+ self,
1225
+ exc_type,
1226
+ exc_val,
1227
+ exc_tb,
1228
+ ) -> None:
1229
+
1230
+ self.disconnect()