diff --git a/README.md b/README.md index 2be9778..dffb9a4 100644 --- a/README.md +++ b/README.md @@ -130,6 +130,25 @@ rowLevelResult_df.show() Each check produces a Boolean column (named after the check description) indicating pass/fail per row. When a single Check contains multiple constraints, they are ANDed together into one Boolean column — the row passes only if all constraints in that Check pass. Only checks with row-level-capable constraints (e.g., `isComplete`, `isContainedIn`, `hasPattern`, `isUnique`) will produce output columns. +### DQDL Rules + +Rules can also be written in [DQDL](https://docs.aws.amazon.com/glue/latest/dg/dqdl.html) (Data Quality Definition Language). See the [Deequ README](https://github.com/awslabs/deequ#supported-dqdl-rules) for the supported rules, such as `DataFreshness` for checking how recent the data is: + +```python +from pydeequ.dqdl import EvaluateDataQuality + +ruleset = """Rules=[ + RowCount >= 3, + IsComplete "a", + DataFreshness "updated_at" <= 24 hours +]""" + +outcomes_df = EvaluateDataQuality.process(spark, df, ruleset) +outcomes_df.show() +``` + +Use `processRows()` to see which rows passed or failed each rule. It returns a dict with the `originalData`, `ruleOutcomes` and `rowLevelOutcomes` DataFrames. Rules without row-level support in Deequ (such as `RowCount` and `DataFreshness`) are listed under `DataQualityRulesSkip`. Dataset comparison rules such as `RowCountMatch "reference" >= 0.9` take their reference DataFrames through `additionalDataSources={"reference": reference_df}`. + ### Repository Save to a Metrics Repository by adding the `useRepository()` and `saveOrAppendResult()` calls to your Analysis Runner. diff --git a/docs/source/pydeequ.rst b/docs/source/pydeequ.rst index 68c253b..68b0403 100644 --- a/docs/source/pydeequ.rst +++ b/docs/source/pydeequ.rst @@ -25,6 +25,14 @@ Checks :undoc-members: :show-inheritance: +DQDL +------------------- + +.. automodule:: pydeequ.dqdl + :members: + :undoc-members: + :show-inheritance: + Profiles ----------------------- diff --git a/pydeequ/dqdl.py b/pydeequ/dqdl.py new file mode 100644 index 0000000..71ab5d2 --- /dev/null +++ b/pydeequ/dqdl.py @@ -0,0 +1,95 @@ +# -*- coding: utf-8 -*- +"""Evaluate data quality rules written in DQDL (Data Quality Definition Language). + +See https://docs.aws.amazon.com/glue/latest/dg/dqdl.html for the DQDL syntax. +""" +from typing import Dict, Optional + +from pyspark.sql import DataFrame, SparkSession + +from pydeequ.pandas_utils import ensure_pyspark_df + + +class EvaluateDataQuality: + """Validates a DataFrame against a ruleset defined in DQDL. + + Example:: + + ruleset = '''Rules=[ + IsComplete "id", + DataFreshness "updated_at" <= 24 hours + ]''' + outcomes = EvaluateDataQuality.process(spark, df, ruleset) + """ + + ORIGINAL_DATA_KEY = "originalData" + RULE_OUTCOMES_KEY = "ruleOutcomes" + ROW_LEVEL_OUTCOMES_KEY = "rowLevelOutcomes" + + @classmethod + def process( + cls, + spark_session: SparkSession, + data: DataFrame, + rulesetDefinition: str, + additionalDataSources: Optional[Dict[str, DataFrame]] = None, + pandas: bool = False, + ): + """ + Evaluates a DQDL ruleset and returns one row per rule. + + :param SparkSession spark_session: SparkSession + :param DataFrame data: DataFrame to validate + :param str rulesetDefinition: DQDL ruleset, e.g. 'Rules=[DataFreshness "ts" <= 24 hours]' + :param dict additionalDataSources: alias -> DataFrame for dataset comparison rules + (e.g. RowCountMatch, ReferentialIntegrity) + :param bool pandas: If True, return a Pandas DataFrame instead of PySpark + :return: DataFrame with columns Rule, Outcome, FailureReason, EvaluatedMetrics, EvaluatedRule + """ + jdf = cls._evaluator(spark_session).process( + *cls._arguments(spark_session, data, rulesetDefinition, additionalDataSources) + ) + df = DataFrame(jdf, spark_session) + return df.toPandas() if pandas else df + + @classmethod + def processRows( + cls, + spark_session: SparkSession, + data: DataFrame, + rulesetDefinition: str, + additionalDataSources: Optional[Dict[str, DataFrame]] = None, + pandas: bool = False, + ) -> Dict[str, DataFrame]: + """ + Evaluates a DQDL ruleset and returns both rule-level and row-level outcomes. + + :param SparkSession spark_session: SparkSession + :param DataFrame data: DataFrame to validate + :param str rulesetDefinition: DQDL ruleset + :param dict additionalDataSources: alias -> DataFrame for dataset comparison rules + :param bool pandas: If True, the returned DataFrames are Pandas DataFrames + :return: dict with keys "originalData" (the input data), "ruleOutcomes" (one row per rule) + and "rowLevelOutcomes" (input rows with per-row passed/failed/skipped rule arrays) + """ + results = cls._evaluator(spark_session).processRows( + *cls._arguments(spark_session, data, rulesetDefinition, additionalDataSources) + ) + keys = (cls.ORIGINAL_DATA_KEY, cls.RULE_OUTCOMES_KEY, cls.ROW_LEVEL_OUTCOMES_KEY) + dfs = {key: DataFrame(results.apply(key), spark_session) for key in keys} + return {key: df.toPandas() for key, df in dfs.items()} if pandas else dfs + + @staticmethod + def _evaluator(spark_session: SparkSession): + return spark_session._jvm.com.amazon.deequ.dqdl.EvaluateDataQuality + + @staticmethod + def _arguments(spark_session, data, rulesetDefinition, additionalDataSources): + if not isinstance(rulesetDefinition, str): + raise TypeError(f"Expected str for rulesetDefinition, not {type(rulesetDefinition)}") + data = ensure_pyspark_df(spark_session, data) + sources = { + alias: ensure_pyspark_df(spark_session, df)._jdf + for alias, df in (additionalDataSources or {}).items() + } + return data._jdf, rulesetDefinition, spark_session._jvm.PythonUtils.toScalaMap(sources) diff --git a/tests/test_dqdl.py b/tests/test_dqdl.py new file mode 100644 index 0000000..f9859e6 --- /dev/null +++ b/tests/test_dqdl.py @@ -0,0 +1,85 @@ +# -*- coding: utf-8 -*- +import unittest +from datetime import datetime, timedelta + +from pyspark.sql import Row + +from pydeequ.dqdl import EvaluateDataQuality +from tests.conftest import setup_pyspark + + +class TestDQDL(unittest.TestCase): + @classmethod + def setUpClass(cls): + cls.spark = setup_pyspark().appName("test-dqdl-local").getOrCreate() + cls.sc = cls.spark.sparkContext + now = datetime.now() + cls.df = cls.sc.parallelize( + [ + Row(id="1", name="foo", updated_at=now - timedelta(hours=1)), + Row(id="2", name="bar", updated_at=now - timedelta(hours=2)), + Row(id="3", name=None, updated_at=now - timedelta(hours=50)), + ] + ).toDF() + + @classmethod + def tearDownClass(cls): + cls.spark.sparkContext._gateway.shutdown_callback_server() + cls.spark.stop() + + def outcomes(self, ruleset, **kwargs): + result = EvaluateDataQuality.process(self.spark, self.df, ruleset, **kwargs) + return {row.Rule: row for row in result.collect()} + + def test_process_returns_one_row_per_rule(self): + result = EvaluateDataQuality.process(self.spark, self.df, 'Rules=[RowCount = 3, IsComplete "name"]') + self.assertEqual( + result.columns, ["Rule", "Outcome", "FailureReason", "EvaluatedMetrics", "EvaluatedRule"] + ) + outcomes = {row.Rule: row.Outcome for row in result.collect()} + self.assertEqual(outcomes, {"RowCount = 3": "Passed", 'IsComplete "name"': "Failed"}) + + def test_data_freshness(self): + outcomes = self.outcomes( + 'Rules=[DataFreshness "updated_at" <= 72 hours, DataFreshness "updated_at" <= 24 hours]' + ) + fresh = outcomes['DataFreshness "updated_at" <= 72 hours'] + stale = outcomes['DataFreshness "updated_at" <= 24 hours'] + self.assertEqual(fresh.Outcome, "Passed") + self.assertEqual(stale.Outcome, "Failed") + self.assertAlmostEqual(stale.EvaluatedMetrics["Column.updated_at.DataFreshness.Compliance"], 2 / 3) + + def test_data_freshness_units(self): + outcomes = self.outcomes( + 'Rules=[DataFreshness "updated_at" <= 3 days, DataFreshness "updated_at" > 30 minutes]' + ) + self.assertEqual({row.Outcome for row in outcomes.values()}, {"Passed"}) + + def test_additional_data_sources(self): + reference = self.sc.parallelize([Row(id="1"), Row(id="2"), Row(id="3")]).toDF() + outcomes = self.outcomes( + 'Rules=[RowCountMatch "reference" = 1.0]', additionalDataSources={"reference": reference} + ) + self.assertEqual(outcomes['RowCountMatch "reference" = 1.0'].Outcome, "Passed") + + def test_process_pandas(self): + result = EvaluateDataQuality.process(self.spark, self.df, "Rules=[RowCount > 0]", pandas=True) + self.assertEqual(result["Outcome"].tolist(), ["Passed"]) + + def test_process_rows(self): + results = EvaluateDataQuality.processRows(self.spark, self.df, 'Rules=[IsComplete "name"]') + self.assertEqual(set(results), {"originalData", "ruleOutcomes", "rowLevelOutcomes"}) + self.assertEqual(results["originalData"].count(), 3) + self.assertEqual(results["ruleOutcomes"].first().Outcome, "Failed") + row_results = { + row.id: row.DataQualityEvaluationResult for row in results["rowLevelOutcomes"].collect() + } + self.assertEqual(row_results, {"1": "Passed", "2": "Passed", "3": "Failed"}) + + def test_process_rows_pandas(self): + results = EvaluateDataQuality.processRows(self.spark, self.df, 'Rules=[IsComplete "id"]', pandas=True) + self.assertEqual(len(results["rowLevelOutcomes"]), 3) + + def test_ruleset_must_be_str(self): + with self.assertRaises(TypeError): + EvaluateDataQuality.process(self.spark, self.df, ["RowCount > 0"])