|
1 | 1 | """Step implementations in Spark.""" |
2 | 2 |
|
3 | 3 | # pylint: disable=R0801 |
4 | | -from collections.abc import Callable, Iterator |
| 4 | +from collections.abc import Callable, Iterable, Iterator |
5 | 5 | from typing import Optional |
6 | 6 | from uuid import uuid4 |
7 | 7 |
|
|
13 | 13 | from dve.core_engine.backends.exceptions import ConstraintError |
14 | 14 | from dve.core_engine.backends.implementations.spark.spark_helpers import ( |
15 | 15 | create_udf, |
| 16 | + df_is_empty, |
16 | 17 | get_all_registered_udfs, |
17 | 18 | object_to_spark_literal, |
18 | 19 | spark_filter_contract_errors, |
|
48 | 49 | SemiJoin, |
49 | 50 | TableUnion, |
50 | 51 | ) |
| 52 | +from dve.core_engine.constants import RECORD_INDEX_COLUMN_NAME |
51 | 53 | from dve.core_engine.functions import implementations as functions |
52 | 54 | from dve.core_engine.message import FeedbackMessage |
53 | 55 | from dve.core_engine.templating import template_object |
@@ -341,16 +343,107 @@ def union(self, entities: SparkEntities, *, config: TableUnion) -> Messages: |
341 | 343 | return [] |
342 | 344 |
|
343 | 345 | def identify_orphans( |
344 | | - self, entities: SparkEntities, *, config: OrphanIdentification |
345 | | - ) -> tuple[Messages, int]: |
346 | | - # TODO - adjust this to new setup of identify and remove orphans |
347 | | - raise NotImplementedError |
| 346 | + self, |
| 347 | + entities: SparkEntities, |
| 348 | + *, |
| 349 | + config: OrphanIdentification, |
| 350 | + ) -> Iterable: |
| 351 | + """Identify records in an entity which don't have at least one corresponding |
| 352 | + match in the target. A new boolean column will be added to `entity` ('IsOrphaned') |
| 353 | + indicating whether the condition matched. |
| 354 | +
|
| 355 | + If there is already an 'IsOrphaned' column in the entity, this will be set to the |
| 356 | + logical OR of its current value and the value it would have been set to otherwise. |
| 357 | +
|
| 358 | + """ |
| 359 | + source_df: DataFrame = entities[config.entity_name] |
| 360 | + source_df = source_df.alias(config.entity_name) |
| 361 | + |
| 362 | + if df_is_empty(source_df): |
| 363 | + self.logger.info(f"{config.entity_name} is empty. Skipping orphan check.") |
| 364 | + return |
| 365 | + |
| 366 | + target_df: DataFrame = entities[config.target_name] |
| 367 | + match_name = f"matched_{uuid4().hex}" |
| 368 | + target_df = target_df.select("*", sf.lit(1).alias(match_name)).alias(config.target_name) |
| 369 | + |
| 370 | + orphaned_df: DataFrame = ( |
| 371 | + source_df.join(target_df, on=sf.expr(config.join_condition), how="left") |
| 372 | + .groupBy(f"{config.entity_name}.{config.id}") |
| 373 | + .agg( |
| 374 | + sf.first(f"{config.entity_name}.{RECORD_INDEX_COLUMN_NAME}").alias( |
| 375 | + RECORD_INDEX_COLUMN_NAME |
| 376 | + ), # pylint: disable=C0301 |
| 377 | + (sf.coalesce(sf.count(match_name), sf.lit(0)) == sf.lit(0)).alias("IsOrphaned"), |
| 378 | + ) |
| 379 | + .filter(sf.col("IsOrphaned")) |
| 380 | + .select(RECORD_INDEX_COLUMN_NAME) |
| 381 | + .alias("orphan") |
| 382 | + ) |
| 383 | + |
| 384 | + message_df = ( |
| 385 | + entities[config.entity_name] |
| 386 | + .alias(config.entity_name) |
| 387 | + .join( |
| 388 | + orphaned_df, |
| 389 | + sf.expr( |
| 390 | + f"{config.entity_name}.{RECORD_INDEX_COLUMN_NAME} = orphan.{RECORD_INDEX_COLUMN_NAME}", # pylint: disable=C0301 |
| 391 | + ), |
| 392 | + "semi", |
| 393 | + ) |
| 394 | + ) |
| 395 | + filtered_rel = ( |
| 396 | + entities[config.entity_name] |
| 397 | + .alias(config.entity_name) |
| 398 | + .join( |
| 399 | + orphaned_df, |
| 400 | + sf.expr( |
| 401 | + f"{config.entity_name}.{RECORD_INDEX_COLUMN_NAME} = orphan.{RECORD_INDEX_COLUMN_NAME}", # pylint: disable=C0301 |
| 402 | + ), |
| 403 | + "anti", |
| 404 | + ) |
| 405 | + ) |
| 406 | + |
| 407 | + entities[config.entity_name] = filtered_rel |
| 408 | + |
| 409 | + for r in message_df.toLocalIterator(): |
| 410 | + yield r.asDict() |
348 | 411 |
|
349 | 412 | def check_mandatory_group( |
350 | 413 | self, entities: SparkEntities, *, config: GroupIdentification |
351 | 414 | ) -> Iterator: |
352 | | - # TODO - implement for spark |
353 | | - raise NotImplementedError |
| 415 | + """ |
| 416 | + Check that a mandatory key in an entity has at least one valid entry in the all the |
| 417 | + child entities. |
| 418 | + """ |
| 419 | + source_df: DataFrame = entities[config.entity_name] |
| 420 | + source_df = source_df.alias(config.entity_name) |
| 421 | + target_df: DataFrame = entities[config.target_name] |
| 422 | + target_df = target_df.alias(config.target_name) |
| 423 | + |
| 424 | + source_columns = [f"{config.entity_name}.{c.strip()}" for c in source_df.columns] |
| 425 | + _pk, fk = config.join_condition.split("=") |
| 426 | + |
| 427 | + joined_df = source_df.join(target_df, sf.expr(config.join_condition), "left").select( |
| 428 | + *source_columns, |
| 429 | + sf.col(fk.strip()).alias("fk"), |
| 430 | + ) |
| 431 | + |
| 432 | + missing_children_df = joined_df.filter("fk IS NULL") |
| 433 | + filtered_df = joined_df.filter("fk IS NOT NULL").select("*").drop(sf.col("fk")) |
| 434 | + |
| 435 | + if not df_is_empty(missing_children_df): |
| 436 | + _no_valid_children = missing_children_df.count() |
| 437 | + else: |
| 438 | + _no_valid_children = 0 |
| 439 | + self.logger.info( |
| 440 | + f"Found {_no_valid_children} records with no valid children in {config.entity_name}." |
| 441 | + ) |
| 442 | + |
| 443 | + entities[config.entity_name] = filtered_df |
| 444 | + |
| 445 | + for r in missing_children_df.toLocalIterator(): |
| 446 | + yield r.asDict() |
354 | 447 |
|
355 | 448 | def filter(self, entities: SparkEntities, *, config: ImmediateFilter) -> Messages: |
356 | 449 | """Filter an entity immediately, and do not emit any messages. |
|
0 commit comments