kumar6591/data-quality-env
0
1from __future__ import annotations2 3import itertools4import re5from dataclasses import dataclass6from typing import Iterable7 8 9@dataclass(frozen=True)10class AlgoConfig:11 w_coverage: float12 w_stat: float13 w_risk: float14 w_novelty: float15 limit_bonus: float16 repeat_penalty: float17 18 19def _query_features(sql: str) -> dict[str, float]:20 s = (sql or "").lower()21 return {22 "coverage": float(any(k in s for k in ["count(", "sum(", "avg(", "group by", "distinct"])),23 "stat": float(any(k in s for k in ["avg(", "stddev", "variance", "percentile", "try_cast", "strptime"])),24 "risk": float(any(k in s for k in ["drop", "truncate", "delete", "insert", "update", "alter", "create"])),25 "novelty": float(any(k in s for k in ["left join", "except", "not in", "having", "case when"])),26 "has_limit": float("limit" in s),27 }28 29 30def _task_keywords(task_id: int) -> list[str]:31 if task_id == 1:32 return ["null", "email", "customer_id", "duplicate", "group by"]33 if task_id == 2:34 return ["quantity", "amount", "n/a", "try_cast", "order_date"]35 return ["transactions_baseline", "transactions_current", "category", "user_id", "avg(amount)"]36 37 38def _task_relevance(task_id: int, sql: str) -> float:39 s = (sql or "").lower()40 keys = _task_keywords(task_id)41 hits = sum(1 for k in keys if k in s)42 return hits / max(1, len(keys))43 44 45def _sql_shape_penalty(sql: str) -> float:46 # Penalize very long and likely redundant SQL in a constrained step budget.47 length = len(sql or "")48 if length < 120:49 return 0.050 if length < 300:51 return 0.0252 return 0.0553 54 55def algorithm_config_stream() -> Iterable[AlgoConfig]:56 # 11^4 * 7^2 = 717,409 total algorithm configurations.57 grid_a = [i / 10 for i in range(0, 11)]58 grid_b = [i / 20 for i in range(0, 7)]59 for a, b, c, d, e, f in itertools.product(grid_a, grid_a, grid_a, grid_a, grid_b, grid_b):60 yield AlgoConfig(61 w_coverage=a,62 w_stat=b,63 w_risk=c,64 w_novelty=d,65 limit_bonus=e,66 repeat_penalty=f,67 )68 69 70def _config_query_score(task_id: int, sql: str, cfg: AlgoConfig, q_prior: float) -> float:71 f = _query_features(sql)72 relevance = _task_relevance(task_id, sql)73 penalty_len = _sql_shape_penalty(sql)74 score = (75 cfg.w_coverage * f["coverage"]76 + cfg.w_stat * f["stat"]77 + cfg.w_novelty * f["novelty"]78 + cfg.limit_bonus * f["has_limit"]79 + 0.6 * relevance80 + 0.4 * q_prior81 - cfg.w_risk * f["risk"]82 - penalty_len83 )84 return score85 86 87def _ranking_for_config(task_id: int, queries: list[str], cfg: AlgoConfig, priors: list[float]) -> list[int]:88 pairs = []89 for i, q in enumerate(queries):90 pairs.append((i, _config_query_score(task_id, q, cfg, priors[i])))91 pairs.sort(key=lambda x: x[1], reverse=True)92 return [i for i, _ in pairs]93 94 95def select_best_config(task_id: int, queries: list[str], priors: list[float], max_configs: int = 100_000) -> AlgoConfig:96 best_cfg = None97 best_obj = -10**998 99 for idx, cfg in enumerate(algorithm_config_stream()):100 if idx >= max_configs:101 break102 ranking = _ranking_for_config(task_id, queries, cfg, priors)103 104 # Objective: prioritize top-2 quality and diversity in SQL intent.105 top = ranking[:2]106 top_score = sum(_config_query_score(task_id, queries[i], cfg, priors[i]) for i in top)107 108 intents = set()109 for i in top:110 s = queries[i].lower()111 intent = "join" if any(k in s for k in ["join", "except", "not in"]) else "agg"112 intents.add(intent)113 diversity_bonus = 0.05 if len(intents) > 1 else 0.0114 115 obj = top_score + diversity_bonus116 if obj > best_obj:117 best_obj = obj118 best_cfg = cfg119 120 return best_cfg if best_cfg is not None else AlgoConfig(0.5, 0.5, 1.0, 0.5, 0.0, 0.0)121 122 123def ensemble_order(task_id: int, queries: list[str], priors: list[float], max_configs: int = 100_000) -> list[str]:124 cfg = select_best_config(task_id, queries, priors, max_configs=max_configs)125 ranking = _ranking_for_config(task_id, queries, cfg, priors)126 127 # De-prioritize unsafe SQL just in case external user-provided probes are included.128 safe = []129 unsafe = []130 for i in ranking:131 if re.search(r"\b(drop|truncate|delete|insert|update|alter|create)\b", queries[i], re.IGNORECASE):132 unsafe.append(queries[i])133 else:134 safe.append(queries[i])135 return safe + unsafe136 