CoolFace
Apppublic

kumar6591/data-quality-env

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
sql_brain.py81 linesDownload Raw Back to env
1from __future__ import annotations2 3from dataclasses import dataclass4 5 6@dataclass(frozen=True)7class SQLProbe:8    name: str9    purpose: str10    sql_template: str11 12 13TASK1_PROBES = [14    SQLProbe("sample_rows", "Quick table sanity sample", "SELECT * FROM {table} LIMIT 5"),15    SQLProbe("null_email", "Count null emails", "SELECT SUM(CASE WHEN email IS NULL THEN 1 ELSE 0 END) AS null_email FROM {table}"),16    SQLProbe("null_customer_id", "Count null customer IDs", "SELECT SUM(CASE WHEN customer_id IS NULL THEN 1 ELSE 0 END) AS null_customer_id FROM {table}"),17    SQLProbe(18        "duplicate_rows",19        "Estimate exact duplicate row count",20        "SELECT COALESCE(SUM(c-1),0) AS duplicate_rows FROM ("21        "SELECT customer_id, email, name, signup_date, country, COUNT(*) AS c "22        "FROM {table} GROUP BY 1,2,3,4,5 HAVING COUNT(*) > 1) t",23    ),24    SQLProbe("country_dist", "Distribution by country", "SELECT country, COUNT(*) AS n FROM {table} GROUP BY country ORDER BY n DESC"),25]26 27TASK2_PROBES = [28    SQLProbe("sample_rows", "Quick table sanity sample", "SELECT * FROM {table} LIMIT 5"),29    SQLProbe(30        "negative_quantity_rows",31        "Count negative quantity violations",32        "SELECT SUM(CASE WHEN quantity < 0 THEN 1 ELSE 0 END) AS negative_quantity_rows FROM {table}",33    ),34    SQLProbe(35        "unparseable_amount_rows",36        "Count unparseable amount values",37        "SELECT SUM(CASE WHEN try_cast(replace(amount, '$', '') AS DOUBLE) IS NULL THEN 1 ELSE 0 END) AS unparseable_amount_rows FROM {table}",38    ),39    SQLProbe(40        "amount_parse_preview",41        "Preview parsed amounts",42        "SELECT amount, try_cast(replace(amount, '$', '') AS DOUBLE) AS amount_num FROM {table} LIMIT 20",43    ),44    SQLProbe("status_dist", "Distribution by status", "SELECT status, COUNT(*) AS n FROM {table} GROUP BY status ORDER BY n DESC"),45]46 47TASK3_PROBES = [48    SQLProbe(49        "mean_shift",50        "Compare baseline/current amount means",51        "SELECT (SELECT AVG(amount) FROM transactions_baseline) AS baseline_mean, "52        "(SELECT AVG(amount) FROM transactions_current) AS current_mean",53    ),54    SQLProbe(55        "new_categories",56        "Find categories present only in current snapshot",57        "SELECT DISTINCT c.category FROM transactions_current c "58        "LEFT JOIN (SELECT DISTINCT category FROM transactions_baseline) b "59        "ON c.category=b.category WHERE b.category IS NULL ORDER BY c.category",60    ),61    SQLProbe(62        "new_user_row_pct",63        "Estimate referential drift on user_id",64        "SELECT AVG(CASE WHEN user_id >= 1000 THEN 1.0 ELSE 0.0 END) AS new_user_row_pct "65        "FROM transactions_current",66    ),67    SQLProbe(68        "mean_by_category",69        "Amount mean by category in current snapshot",70        "SELECT category, AVG(amount) AS avg_amount FROM transactions_current GROUP BY category ORDER BY avg_amount DESC",71    ),72]73 74 75def probes_for_task(task_id: int, table_name: str) -> list[str]:76    if task_id == 1:77        return [p.sql_template.format(table=table_name) for p in TASK1_PROBES]78    if task_id == 2:79        return [p.sql_template.format(table=table_name) for p in TASK2_PROBES]80    return [p.sql_template.format(table=table_name) for p in TASK3_PROBES]81