Team Ai
Apppublic

MauryaVivek/sql-data-quality-agent

sourceHugging Facemitupdated 6mo agoView on Hugging Face
0likes
tasks.py366 linesDownload Raw Back to root
1"""2tasks.py3========4Task registry for the SQL Data Quality Agent.5 6Each task defines:7  - task_id, name, difficulty, description8  - max_steps9  - db_generator function reference10  - grader: computes a quality score 0.0–1.0 from a live SQLite connection11"""12 13import math14import sqlite315from typing import Dict, Any, Callable, Optional16from pydantic import BaseModel, field_validator, model_validator17 18 19# ---------------------------------------------------------------------------20# Universal clamping — THE single source of truth for score safety21# ---------------------------------------------------------------------------22 23_SCORE_LOW = 0.0124_SCORE_HIGH = 0.9925 26 27def _safe_float(v: Any) -> float:28    """Convert any value to a safe float. Returns 0.5 for any invalid input."""29    try:30        f = float(v)31    except (TypeError, ValueError):32        return 0.533    if math.isnan(f) or math.isinf(f):34        return 0.535    return f36 37 38def clamp_score(score: Any, low: float = _SCORE_LOW, high: float = _SCORE_HIGH) -> float:39    """Clamp score strictly into (0, 1) exclusive — required by OpenEnv Phase 2 validator.40 41    Default range: [0.01, 0.99] — safe margin from both boundaries.42    Handles NaN, Inf, None, strings, and floating-point edge cases.43    44    CRITICAL: The validator checks that no score equals exactly 0.0 or 1.0.45    This function guarantees the output is always in [low, high].46    """47    f = _safe_float(score)48    # Clamp to specified range49    result = max(low, min(high, f))50    # Final paranoid guard against floating-point weirdness51    if result <= 0.0 or result >= 1.0:52        return 0.553    return round(result, 4)54 55 56def clamp_ratio(ratio: Any) -> float:57    """Clamp a ratio value strictly into (0, 1) exclusive.58    Uses [0.001, 0.999] to preserve more precision for ratios.59    """60    f = _safe_float(ratio)61    if f <= 0.0:62        return 0.00163    if f >= 1.0:64        return 0.99965    result = max(0.001, min(0.999, f))66    if result <= 0.0 or result >= 1.0:67        return 0.568    return round(result, 4)69 70 71class QualityReport(BaseModel):72    null_ratio: float = 0.0173    duplicate_ratio: float = 0.0174    type_error_ratio: float = 0.0175    constraint_violation_ratio: float = 0.0176    value_error_ratio: float = 0.0177    overall_score: float = 0.578    details: Dict[str, Any] = {}79 80    @field_validator("overall_score", mode="before")81    @classmethod82    def _clamp_overall_score(cls, v: Any) -> float:83        """Clamp overall_score strictly into (0, 1) — OpenEnv Phase 2 requirement."""84        return clamp_score(v)85 86    @field_validator(87        "null_ratio",88        "duplicate_ratio",89        "type_error_ratio",90        "constraint_violation_ratio",91        "value_error_ratio",92        mode="before",93    )94    @classmethod95    def _validate_ratios(cls, v: Any) -> float:96        """Clamp ratio fields — keep in (0, 1) exclusive for safety."""97        return clamp_ratio(v)98 99    @model_validator(mode="after")100    def _final_check(self):101        """Absolute final guard: overall_score and all ratios strictly in (0, 1)."""102        # Guard overall_score103        s = self.overall_score104        if s is None or (isinstance(s, float) and (math.isnan(s) or math.isinf(s))) or s <= 0.0 or s >= 1.0:105            self.overall_score = 0.5106        # Guard all ratios107        for field_name in ["null_ratio", "duplicate_ratio", "type_error_ratio",108                           "constraint_violation_ratio", "value_error_ratio"]:109            val = getattr(self, field_name)110            if val is None or (isinstance(val, float) and (math.isnan(val) or math.isinf(val))) or val <= 0.0 or val >= 1.0:111                setattr(self, field_name, 0.01)112        return self113 114    def model_dump(self, **kwargs) -> Dict[str, Any]:115        """Override model_dump to ensure all numeric fields are clamped on serialization."""116        d = super().model_dump(**kwargs)117        # Final safety net before JSON serialization118        d["overall_score"] = clamp_score(d.get("overall_score", 0.5))119        for ratio_field in ["null_ratio", "duplicate_ratio", "type_error_ratio",120                            "constraint_violation_ratio", "value_error_ratio"]:121            d[ratio_field] = clamp_ratio(d.get(ratio_field, 0.01))122        return d123 124 125# ---------------------------------------------------------------------------126# Task 1 Grader — Null Patrol127# ---------------------------------------------------------------------------128 129def grade_null_patrol(conn: sqlite3.Connection) -> QualityReport:130    """Score based on how many email/phone nulls remain."""131    cur = conn.cursor()132    cur.execute("SELECT COUNT(*) FROM customers")133    total = cur.fetchone()[0]134    if total == 0:135        return QualityReport(overall_score=0.5)136 137    cur.execute("SELECT COUNT(*) FROM customers WHERE email IS NULL")138    null_emails = cur.fetchone()[0]139 140    cur.execute("SELECT COUNT(*) FROM customers WHERE phone IS NULL")141    null_phones = cur.fetchone()[0]142 143    total_nullable_fields = total * 2  # email + phone144    total_nulls = null_emails + null_phones145    null_ratio_raw = total_nulls / total_nullable_fields if total_nullable_fields > 0 else 0.0146 147    # Scale score: 0% nulls -> 0.85, 100% nulls -> 0.15 (safe range)148    quality = 1.0 - null_ratio_raw149    score = 0.15 + quality * 0.70  # Maps quality [0, 1] -> [0.15, 0.85]150 151    return QualityReport(152        null_ratio=null_ratio_raw,153        overall_score=score,154        details={155            "total_rows": total,156            "null_emails": null_emails,157            "null_phones": null_phones,158        },159    )160 161 162# ---------------------------------------------------------------------------163# Task 2 Grader — Duplicate Destroyer164# ---------------------------------------------------------------------------165 166def grade_duplicate_destroyer(conn: sqlite3.Connection) -> QualityReport:167    """Score based on how many duplicate order_ids remain."""168    cur = conn.cursor()169    cur.execute("SELECT COUNT(*) FROM orders")170    total = cur.fetchone()[0]171    if total == 0:172        return QualityReport(overall_score=0.5)173 174    # Count rows that are NOT the earliest row for their order_id175    cur.execute("""176        SELECT COUNT(*) FROM orders177        WHERE row_id NOT IN (178            SELECT MIN(row_id) FROM orders GROUP BY order_id179        )180    """)181    duplicate_count = cur.fetchone()[0]182 183    duplicate_ratio_raw = duplicate_count / total if total > 0 else 0.0184    # Scale score: 0% dupes -> 0.85, 100% dupes -> 0.15 (safe range)185    quality = 1.0 - duplicate_ratio_raw186    score = 0.15 + quality * 0.70  # Maps quality [0, 1] -> [0.15, 0.85]187 188    return QualityReport(189        duplicate_ratio=duplicate_ratio_raw,190        overall_score=score,191        details={192            "total_rows": total,193            "duplicate_rows": duplicate_count,194        },195    )196 197 198# ---------------------------------------------------------------------------199# Task 3 Grader — Constraint Cascade200# ---------------------------------------------------------------------------201 202def grade_constraint_cascade(conn: sqlite3.Connection) -> QualityReport:203    """204    Composite score (weighted):205      30% — type errors (price not numeric, e.g. '$12.99')206      30% — FK violations (inventory rows referencing missing product_ids)207      20% — value errors (negative quantities)208      20% — category normalisation (mixed-case like 'electronics', 'BOOKS')209    """210    cur = conn.cursor()211 212    # --- Type errors in products.price ---213    cur.execute("SELECT COUNT(*) FROM products")214    total_products = cur.fetchone()[0] or 1215 216    cur.execute("SELECT price FROM products")217    prices = [row[0] for row in cur.fetchall()]218    type_errors = 0219    for p in prices:220        try:221            float(str(p).strip())222        except (ValueError, TypeError):223            type_errors += 1224 225    type_error_ratio_raw = type_errors / total_products if total_products > 0 else 0.0226 227    # --- FK violations in inventory ---228    cur.execute("SELECT COUNT(*) FROM inventory")229    total_inv = cur.fetchone()[0] or 1230 231    cur.execute("""232        SELECT COUNT(*) FROM inventory233        WHERE product_id NOT IN (SELECT product_id FROM products)234    """)235    fk_violations = cur.fetchone()[0]236    fk_ratio_raw = fk_violations / total_inv if total_inv > 0 else 0.0237 238    # --- Negative quantities ---239    cur.execute("SELECT COUNT(*) FROM inventory WHERE quantity < 0")240    neg_qty = cur.fetchone()[0]241    neg_ratio_raw = neg_qty / total_inv if total_inv > 0 else 0.0242 243    # --- Mixed-case category names ---244    VALID_CATEGORIES = {"Electronics", "Clothing", "Food", "Books", "Tools"}245    cur.execute("SELECT category FROM products")246    categories = [row[0] for row in cur.fetchall()]247    case_errors = sum(1 for c in categories if c not in VALID_CATEGORIES)248    case_error_ratio_raw = case_errors / total_products if total_products > 0 else 0.0249 250    # Weighted composite (4 dimensions) scaled to (0.15, 0.85)251    type_score  = 1.0 - type_error_ratio_raw252    fk_score    = 1.0 - fk_ratio_raw253    neg_score   = 1.0 - neg_ratio_raw254    case_score  = 1.0 - case_error_ratio_raw255    quality = (0.30 * type_score) + (0.30 * fk_score) + (0.20 * neg_score) + (0.20 * case_score)256    overall = 0.15 + quality * 0.70  # Maps quality [0, 1] -> [0.15, 0.85]257 258    return QualityReport(259        type_error_ratio=type_error_ratio_raw,260        constraint_violation_ratio=fk_ratio_raw,261        value_error_ratio=neg_ratio_raw,262        overall_score=overall,263        details={264            "total_products": total_products,265            "type_errors": type_errors,266            "total_inventory": total_inv,267            "fk_violations": fk_violations,268            "negative_quantity_rows": neg_qty,269            "category_case_errors": case_errors,270        },271    )272 273 274# ---------------------------------------------------------------------------275# Task Registry276# ---------------------------------------------------------------------------277 278class TaskDefinition(BaseModel):279    task_id: str280    name: str281    difficulty: str  # easy | medium | hard282    description: str283    max_steps: int284    success_threshold: float  # score needed to consider done285 286    class Config:287        arbitrary_types_allowed = True288 289 290TASK_REGISTRY: Dict[str, Dict[str, Any]] = {291    "null_patrol": {292        "meta": TaskDefinition(293            task_id="null_patrol",294            name="Null Patrol",295            difficulty="easy",296            description=(297                "A customers table has ~20% of email and phone fields set to NULL. "298                "Your job is to fill them with sensible placeholder values so the "299                "null ratio drops below 5%."300            ),301            max_steps=15,302            success_threshold=0.85,303        ),304        "grader": grade_null_patrol,305        "db_tables": ["customers"],306        "hints": [307            "Inspect NULLs: SELECT COUNT(*) FROM customers WHERE email IS NULL OR phone IS NULL",308            "UPDATE rows where email is NULL — replace with a placeholder string",309            "UPDATE rows where phone is NULL — replace with a placeholder string",310        ],311    },312    "duplicate_destroyer": {313        "meta": TaskDefinition(314            task_id="duplicate_destroyer",315            name="Duplicate Destroyer",316            difficulty="medium",317            description=(318                "An orders table has ~15% duplicate rows (same order_id with different "319                "inserted_at timestamps). Identify and remove the duplicates, keeping "320                "the earliest entry for each order_id."321            ),322            max_steps=20,323            success_threshold=0.85,324        ),325        "grader": grade_duplicate_destroyer,326        "db_tables": ["orders"],327        "hints": [328            "First inspect: SELECT order_id, COUNT(*) as cnt FROM orders GROUP BY order_id HAVING cnt > 1",329            "You need to keep the earliest row per order_id and delete the rest",330            "Hint: use MIN(row_id) grouped by order_id to identify which rows to keep",331        ],332    },333    "constraint_cascade": {334        "meta": TaskDefinition(335            task_id="constraint_cascade",336            name="Constraint Cascade",337            difficulty="hard",338            description=(339                "A products + inventory database has FOUR categories of issues: "340                "(1) ~20% of product prices are stored as strings like '$12.99' instead of numeric values, "341                "(2) ~15% of inventory rows reference non-existent product_ids (FK violations), "342                "(3) ~10% of inventory quantities are negative, "343                "(4) ~25% of product category names have inconsistent casing (e.g. 'electronics', 'BOOKS'). "344                "Fix all four categories to reach a composite quality score >= 0.80."345            ),346            max_steps=30,347            success_threshold=0.80,348        ),349        "grader": grade_constraint_cascade,350        "db_tables": ["products", "inventory"],351        "hints": [352            "Inspect issues: SELECT DISTINCT category FROM products; SELECT price FROM products WHERE price LIKE '$%'",353            "Check inventory: SELECT COUNT(*) FROM inventory WHERE product_id NOT IN (SELECT product_id FROM products)",354            "Valid categories are: Electronics, Clothing, Food, Books, Tools (title-case exactly)",355            "Quantities should be non-negative; prices should be castable to float",356        ],357    },358}359 360 361def get_task(task_id: str) -> Optional[Dict[str, Any]]:362    return TASK_REGISTRY.get(task_id)363 364 365def list_tasks():366    return [v["meta"].model_dump() for v in TASK_REGISTRY.values()]