Skip to content

Commit c86e09f

Browse files
committed
Resolve comments from wei
1 parent 01f4f67 commit c86e09f

11 files changed

Lines changed: 218 additions & 164 deletions

File tree

providers/common/dataquality/docs/rules.rst

Lines changed: 4 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -26,26 +26,10 @@ their own. They describe *what* to check; :class:`~airflow.providers.common.data
2626
(see :doc:`operators`) is what actually runs them. Rules are usually written by hand, but they
2727
can also be proposed by an LLM -- see :doc:`agents`.
2828

29-
.. code-block:: python
30-
31-
from airflow.providers.common.dataquality.rules import DQRule, RuleSet
32-
33-
orders_ruleset = RuleSet(
34-
name="orders_quality",
35-
rules=(
36-
DQRule(
37-
name="order_id_not_null",
38-
check="null_count",
39-
column="order_id",
40-
condition={"equal_to": 0},
41-
),
42-
DQRule(
43-
name="row_count_present",
44-
check="row_count",
45-
condition={"greater_than": 0},
46-
),
47-
),
48-
)
29+
.. exampleinclude:: /../src/airflow/providers/common/dataquality/example_dags/example_dq_in_taskflow.py
30+
:language: python
31+
:start-after: [START howto_rules_dq_ruleset]
32+
:end-before: [END howto_rules_dq_ruleset]
4933

5034
``DQRule`` fields
5135
------------------

providers/common/dataquality/src/airflow/providers/common/dataquality/backends/object_storage.py

Lines changed: 43 additions & 34 deletions
Original file line numberDiff line numberDiff line change
@@ -37,32 +37,39 @@
3737
import json
3838
import logging
3939
from datetime import datetime, timezone
40-
from typing import Any
40+
from typing import Any, TypedDict
4141

4242
from airflow.providers.common.dataquality.results import DQRun, RuleResult, build_summary
4343
from airflow.sdk import ObjectStoragePath
4444

4545
log = logging.getLogger(__name__)
4646

4747

48+
class DQPageResult(TypedDict):
49+
"""Paginated DQ read result."""
50+
51+
items: list[dict[str, Any]]
52+
next_cursor: str | None
53+
54+
4855
class ObjectStorageResultsBackend:
4956
"""Persist DQ results as JSON files via ``ObjectStoragePath``."""
5057

51-
def __init__(self, results_path: str, conn_id: str | None = None) -> None:
58+
def __init__(self, *, results_path: str, conn_id: str | None = None) -> None:
5259
self.root = ObjectStoragePath(results_path, conn_id=conn_id)
5360

54-
def write_run(self, run: DQRun, results: list[RuleResult]) -> None:
61+
def write_run(self, *, run: DQRun, results: list[RuleResult]) -> None:
5562
timestamp = run.started_at or datetime.now(tz=timezone.utc).isoformat()
5663
compact_ts = self._get_safe_key(timestamp)
57-
payload = self._build_run_payload(run, results)
64+
payload = self._build_run_payload(run=run, results=results)
5865

59-
self._write_run_file(run, timestamp[:10], compact_ts, payload)
60-
self._write_task_instance_index(run, payload)
61-
self._write_rule_indexes(run, results, timestamp)
66+
self._write_run_file(run=run, date_part=timestamp[:10], compact_ts=compact_ts, payload=payload)
67+
self._write_task_instance_index(run=run, payload=payload)
68+
self._write_rule_indexes(run=run, results=results, timestamp=timestamp)
6269

6370
def read_task_rule_history(
64-
self, dag_id: str, task_id: str, rule_uid: str, limit: int = 100, before: str | None = None
65-
) -> dict[str, Any]:
71+
self, *, dag_id: str, task_id: str, rule_uid: str, limit: int = 100, before: str | None = None
72+
) -> DQPageResult:
6673
"""Return recent results for one rule produced by one task, newest first."""
6774
rule_dir = (
6875
self.root
@@ -72,11 +79,11 @@ def read_task_rule_history(
7279
/ f"task_id={task_id}"
7380
/ f"rule_uid={rule_uid}"
7481
)
75-
return self._read_rule_history_dir(rule_dir, limit, before)
82+
return self._read_rule_history_dir(rule_dir=rule_dir, limit=limit, before=before)
7683

7784
def read_task_runs(
78-
self, dag_id: str, task_id: str, limit: int = 50, before: str | None = None
79-
) -> dict[str, Any]:
85+
self, *, dag_id: str, task_id: str, limit: int = 50, before: str | None = None
86+
) -> DQPageResult:
8087
"""
8188
Return recent data quality runs for one task, newest first.
8289
@@ -103,11 +110,9 @@ def read_task_runs(
103110
for path in sorted(date_dir.iterdir(), key=lambda p: p.name, reverse=True):
104111
if not path.name.endswith(".json"):
105112
continue
106-
payload = self._read_json(path)
107-
if payload is None:
113+
if (payload := self._read_json(path)) is None:
108114
continue
109-
cursor = self._get_run_payload_cursor(payload)
110-
if before is not None and cursor >= before:
115+
if before is not None and self._get_run_payload_cursor(payload) >= before:
111116
continue
112117
runs.append(payload)
113118
if len(runs) > limit:
@@ -121,7 +126,7 @@ def read_task_runs(
121126
return {"items": page, "next_cursor": next_cursor}
122127

123128
def read_by_task_instance(
124-
self, dag_id: str, task_id: str, run_id: str, map_index: int = -1
129+
self, *, dag_id: str, task_id: str, run_id: str, map_index: int = -1
125130
) -> dict[str, Any]:
126131
"""Read the latest run for one task instance as ``{"run": ..., "results": ..., "summary": ...}``."""
127132
path = (
@@ -134,7 +139,9 @@ def read_by_task_instance(
134139
)
135140
return self._read_json_or_raise(path)
136141

137-
def _write_run_file(self, run: DQRun, date_part: str, compact_ts: str, payload: dict[str, Any]) -> None:
142+
def _write_run_file(
143+
self, *, run: DQRun, date_part: str, compact_ts: str, payload: dict[str, Any]
144+
) -> None:
138145
run_dir = (
139146
self.root
140147
/ "runs"
@@ -146,39 +153,41 @@ def _write_run_file(self, run: DQRun, date_part: str, compact_ts: str, payload:
146153
run_dir.mkdir(parents=True, exist_ok=True)
147154
(run_dir / f"{compact_ts}__{run.run_uid}.json").write_text(json.dumps(payload, default=str))
148155

149-
def _write_task_instance_index(self, run: DQRun, payload: dict[str, Any]) -> None:
156+
def _write_task_instance_index(self, *, run: DQRun, payload: dict[str, Any]) -> None:
150157
ti_dir = self.root / "runs" / "by_task_instance" / f"dag_id={run.dag_id}" / f"task_id={run.task_id}"
151158
ti_dir.mkdir(parents=True, exist_ok=True)
152159
(ti_dir / f"{self._get_safe_key(run.run_id)}__{run.map_index}.json").write_text(
153160
json.dumps(payload, default=str)
154161
)
155162

156-
def _write_rule_indexes(self, run: DQRun, results: list[RuleResult], timestamp: str) -> None:
163+
def _write_rule_indexes(self, *, run: DQRun, results: list[RuleResult], timestamp: str) -> None:
157164
run_context = self._build_run_context(run)
158165
compact_ts = self._get_safe_key(timestamp)
159166
for result in results:
160167
payload = {"run": run_context, "result": result.to_dict()}
161168
self._write_rule_index(
162-
self.root
163-
/ "rules"
164-
/ "by_task_rule"
165-
/ f"dag_id={run.dag_id}"
166-
/ f"task_id={run.task_id}"
167-
/ f"rule_uid={result.rule_uid}",
168-
compact_ts,
169-
run.run_uid,
170-
payload,
169+
rule_dir=(
170+
self.root
171+
/ "rules"
172+
/ "by_task_rule"
173+
/ f"dag_id={run.dag_id}"
174+
/ f"task_id={run.task_id}"
175+
/ f"rule_uid={result.rule_uid}"
176+
),
177+
compact_ts=compact_ts,
178+
run_uid=run.run_uid,
179+
payload=payload,
171180
)
172181

173182
def _write_rule_index(
174-
self, rule_dir: ObjectStoragePath, compact_ts: str, run_uid: str, payload: dict[str, Any]
183+
self, *, rule_dir: ObjectStoragePath, compact_ts: str, run_uid: str, payload: dict[str, Any]
175184
) -> None:
176185
rule_dir.mkdir(parents=True, exist_ok=True)
177186
(rule_dir / f"{compact_ts}__{run_uid}.json").write_text(json.dumps(payload, default=str))
178187

179188
def _read_rule_history_dir(
180-
self, rule_dir: ObjectStoragePath, limit: int, before: str | None = None
181-
) -> dict[str, Any]:
189+
self, *, rule_dir: ObjectStoragePath, limit: int, before: str | None = None
190+
) -> DQPageResult:
182191
"""
183192
Read rule-result records newest-first, as ``{"items": [...], "next_cursor": ...}``.
184193
@@ -211,12 +220,12 @@ def _get_safe_key(value: str) -> str:
211220
return value.replace("/", "_").replace(":", "_").replace("+", "_")
212221

213222
@staticmethod
214-
def _build_run_payload(run: DQRun, results: list[RuleResult]) -> dict[str, Any]:
223+
def _build_run_payload(*, run: DQRun, results: list[RuleResult]) -> dict[str, Any]:
215224
result_records = [result.to_dict() for result in results]
216225
return {
217226
"run": run.to_dict(),
218227
"results": result_records,
219-
"summary": build_summary(run, results),
228+
"summary": build_summary(run=run, results=results),
220229
}
221230

222231
@staticmethod

providers/common/dataquality/src/airflow/providers/common/dataquality/engines/sql.py

Lines changed: 20 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -57,21 +57,27 @@ class SQLDQEngine:
5757
def __init__(self, hook: DbApiHook) -> None:
5858
self.hook = hook
5959

60-
def measure(self, ruleset: RuleSet, table: str, partition_clause: str | None = None) -> list[Observation]:
60+
def measure(
61+
self, *, ruleset: RuleSet, table: str, partition_clause: str | None = None
62+
) -> list[Observation]:
6163
builtin_rules = [rule for rule in ruleset.rules if rule.check != CUSTOM_SQL_CHECK]
6264
custom_rules = [rule for rule in ruleset.rules if rule.check == CUSTOM_SQL_CHECK]
6365

6466
observations = []
6567
if builtin_rules:
66-
observations.extend(self._measure_builtin(builtin_rules, table, partition_clause))
68+
observations.extend(
69+
self._measure_builtin(rules=builtin_rules, table=table, partition_clause=partition_clause)
70+
)
6771
for rule in custom_rules:
68-
observations.append(self._measure_custom(rule, table))
72+
observations.append(self._measure_custom(rule=rule, table=table))
6973
return observations
7074

71-
def build_batch_sql(self, rules: list[DQRule], table: str, partition_clause: str | None) -> str:
72-
return " UNION ALL ".join(self.build_rule_sql(rule, table, partition_clause) for rule in rules)
75+
def build_batch_sql(self, *, rules: list[DQRule], table: str, partition_clause: str | None) -> str:
76+
return " UNION ALL ".join(
77+
self.build_rule_sql(rule=rule, table=table, partition_clause=partition_clause) for rule in rules
78+
)
7379

74-
def build_rule_sql(self, rule: DQRule, table: str, partition_clause: str | None = None) -> str:
80+
def build_rule_sql(self, *, rule: DQRule, table: str, partition_clause: str | None = None) -> str:
7581
"""Build the SQL used to measure one built-in rule."""
7682
expression = CHECK_SPECS[rule.check].expression.format(column=rule.column)
7783
predicates = [p for p in (partition_clause, rule.partition_clause) if p]
@@ -81,11 +87,14 @@ def build_rule_sql(self, rule: DQRule, table: str, partition_clause: str | None
8187
)
8288

8389
def _measure_builtin(
84-
self, rules: list[DQRule], table: str, partition_clause: str | None
90+
self, *, rules: list[DQRule], table: str, partition_clause: str | None
8591
) -> list[Observation]:
8692
# Built once per rule and reused below, instead of re-deriving each rule's SQL for the
8793
# batch, for the returned Observation, and again in the per-rule fallback.
88-
rule_sql = {rule.rule_uid: self.build_rule_sql(rule, table, partition_clause) for rule in rules}
94+
rule_sql = {
95+
rule.rule_uid: self.build_rule_sql(rule=rule, table=table, partition_clause=partition_clause)
96+
for rule in rules
97+
}
8998
sql = " UNION ALL ".join(rule_sql.values())
9099
log.info("Running %d built-in checks against %s", len(rules), table)
91100
started = time.monotonic()
@@ -97,7 +106,7 @@ def _measure_builtin(
97106
table,
98107
)
99108

100-
return [self._measure_builtin_single(rule, rule_sql[rule.rule_uid]) for rule in rules]
109+
return [self._measure_builtin_single(rule=rule, sql=rule_sql[rule.rule_uid]) for rule in rules]
101110
elapsed_ms = (time.monotonic() - started) * 1000
102111
observed_by_uid = {str(row[0]): row[1] for row in records or []}
103112
return [
@@ -111,7 +120,7 @@ def _measure_builtin(
111120
for rule in rules
112121
]
113122

114-
def _measure_builtin_single(self, rule: DQRule, sql: str) -> Observation:
123+
def _measure_builtin_single(self, *, rule: DQRule, sql: str) -> Observation:
115124
started = time.monotonic()
116125
try:
117126
row = self.hook.get_first(sql)
@@ -126,7 +135,7 @@ def _measure_builtin_single(self, rule: DQRule, sql: str) -> Observation:
126135
)
127136
return Observation(rule=rule, observed_value=row[1], duration_ms=elapsed_ms, sql=sql)
128137

129-
def _measure_custom(self, rule: DQRule, table: str) -> Observation:
138+
def _measure_custom(self, *, rule: DQRule, table: str) -> Observation:
130139
if rule.sql is None:
131140
raise ValueError(f"Rule {rule.name!r} has no SQL to execute")
132141
sql = rule.sql.replace("{table}", table)

providers/common/dataquality/src/airflow/providers/common/dataquality/example_dags/example_dq_in_taskflow.py

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -39,8 +39,7 @@
3939
READY_TABLE = "dq_taskflow_ready_orders"
4040
RESULTS_PATH = Path("/tmp/airflow_dq_example/results")
4141

42-
os.environ.setdefault("AIRFLOW__COMMON_DATAQUALITY__RESULTS_PATH", f"file://{RESULTS_PATH}")
43-
42+
# [START howto_rules_dq_ruleset]
4443
orders_ruleset = RuleSet(
4544
name="orders_taskflow_quality",
4645
rules=(
@@ -63,6 +62,9 @@
6362
),
6463
),
6564
)
65+
# [END howto_rules_dq_ruleset]
66+
67+
os.environ.setdefault("AIRFLOW__COMMON_DATAQUALITY__RESULTS_PATH", f"file://{RESULTS_PATH}")
6668

6769
with DAG(
6870
dag_id=DAG_ID,

providers/common/dataquality/src/airflow/providers/common/dataquality/execution.py

Lines changed: 7 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -90,7 +90,11 @@ def run_quality_checks(
9090
resolved_ruleset = _resolve_ruleset(ruleset)
9191
resolved_hook = hook if hook is not None else _get_hook(conn_id, hook_params)
9292
started_at = datetime.now(tz=timezone.utc).isoformat()
93-
observations = SQLDQEngine(resolved_hook).measure(resolved_ruleset, table, partition_clause)
93+
observations = SQLDQEngine(resolved_hook).measure(
94+
ruleset=resolved_ruleset,
95+
table=table,
96+
partition_clause=partition_clause,
97+
)
9498
finished_at = datetime.now(tz=timezone.utc).isoformat()
9599
return DataQualityResult(
96100
ruleset=resolved_ruleset,
@@ -123,7 +127,7 @@ def persist_quality_results(
123127
finished_at=result.finished_at,
124128
)
125129
results = list(result.results)
126-
summary = build_summary(run, results)
130+
summary = build_summary(run=run, results=results)
127131

128132
backend = get_backend_from_config()
129133
if backend is None:
@@ -132,7 +136,7 @@ def persist_quality_results(
132136
# Persistence is best-effort: an unreachable results store leaves a gap in
133137
# history but must not change the outcome of the check itself.
134138
try:
135-
backend.write_run(run, results)
139+
backend.write_run(run=run, results=results)
136140
except Exception:
137141
log.exception("Failed to persist data quality results; continuing")
138142

providers/common/dataquality/src/airflow/providers/common/dataquality/results.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -100,7 +100,7 @@ def compute_score(results: list[RuleResult]) -> float | None:
100100
return round(1.0 - penalty / len(results), 4)
101101

102102

103-
def build_summary(run: DQRun, results: list[RuleResult]) -> dict[str, Any]:
103+
def build_summary(*, run: DQRun, results: list[RuleResult]) -> dict[str, Any]:
104104
"""Compact run summary attached to XCom and outlet asset events."""
105105
return {
106106
"run_uid": run.run_uid,

0 commit comments

Comments
 (0)