MauryaVivek/sql-data-quality-agent
0
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()]