ritvik360/nl2sql-bench
0
1"""2nl2sql-bench/server/grader.py3==============================4Deterministic, programmatic reward grader.5 6No LLM-as-judge. Every reward is computed by comparing the agent's SQL7execution results against a ground-truth result set.8 9Reward decomposition (sums to 1.0 for a perfect first-attempt answer):10 +0.10 syntax_ok — query runs without SQLite error11 +0.20 columns_match — returned column names match ground truth exactly12 +0.20 row_count_match — number of returned rows matches13 +0.50 exact_match — full result set equals ground truth (order-aware14 for ORDER BY queries, order-agnostic otherwise)15 16Step penalty:17 -0.05 per step beyond the first (encourages solving in fewer steps),18 clamped so the minimum is always 0.0.19 20All rewards are floats in [0.0, 1.0].21"""22 23from __future__ import annotations24 25import sqlite326from typing import Any, Dict, List, Optional, Tuple27 28 29# ── Result normalisation ───────────────────────────────────────────────────30 31def _normalise_value(v: Any) -> Any:32 """Round floats for comparison so 1.2300000001 == 1.23."""33 if isinstance(v, float):34 return round(v, 4)35 if isinstance(v, str):36 return v.strip()37 return v38 39 40def _normalise_row(row: Dict[str, Any]) -> Dict[str, Any]:41 return {k: _normalise_value(v) for k, v in row.items()}42 43 44def _normalise_rows(rows: List[Dict[str, Any]]) -> List[Dict[str, Any]]:45 return [_normalise_row(r) for r in rows]46 47 48# ── SQL execution ──────────────────────────────────────────────────────────49 50def execute_query(51 conn: sqlite3.Connection,52 query: str,53 max_rows: int = 200,54) -> Tuple[Optional[List[Dict[str, Any]]], Optional[str]]:55 """56 Execute a SQL query safely.57 58 Returns (rows, error_string).59 rows is None on error.60 """61 query = query.strip().rstrip(";")62 if not query:63 return None, "Empty query."64 65 # Block write operations — the environment is read-only from the agent's view.66 forbidden = ("insert", "update", "delete", "drop", "alter",67 "create", "replace", "truncate", "pragma")68 first_word = query.split()[0].lower() if query.split() else ""69 if first_word in forbidden:70 return None, (71 f"Operation '{first_word.upper()}' is not allowed. "72 "Only SELECT queries are permitted."73 )74 75 try:76 cur = conn.execute(query)77 cols = [d[0] for d in cur.description] if cur.description else []78 rows = [dict(zip(cols, row)) for row in cur.fetchmany(max_rows)]79 return rows, None80 except sqlite3.Error as exc:81 return None, str(exc)82 83 84# ── Grading logic ──────────────────────────────────────────────────────────85 86class GradeResult:87 __slots__ = (88 "reward", "syntax_ok", "columns_match",89 "row_count_match", "exact_match", "step_penalty",90 "breakdown",91 )92 93 def __init__(94 self,95 reward: float,96 syntax_ok: bool,97 columns_match: bool,98 row_count_match: bool,99 exact_match: bool,100 step_penalty: float,101 ) -> None:102 self.reward = reward103 self.syntax_ok = syntax_ok104 self.columns_match = columns_match105 self.row_count_match = row_count_match106 self.exact_match = exact_match107 self.step_penalty = step_penalty108 self.breakdown = {109 "syntax_ok": 0.10 if syntax_ok else 0.0,110 "columns_match": 0.20 if (syntax_ok and columns_match) else 0.0,111 "row_count_match": 0.20 if (syntax_ok and row_count_match) else 0.0,112 "exact_match": 0.50 if (syntax_ok and exact_match) else 0.0,113 "step_penalty": -step_penalty,114 }115 116 def __repr__(self) -> str: # pragma: no cover117 return (118 f"GradeResult(reward={self.reward:.3f}, "119 f"exact={self.exact_match}, cols={self.columns_match}, "120 f"rows={self.row_count_match}, syntax={self.syntax_ok})"121 )122 123 124def grade(125 actual_rows: Optional[List[Dict[str, Any]]],126 ground_truth_rows: List[Dict[str, Any]],127 error: Optional[str],128 step: int,129 order_sensitive: bool = False,130) -> GradeResult:131 """132 Grade the agent's query result against ground truth.133 134 Parameters135 ----------136 actual_rows : Rows returned by the agent's query (None on error).137 ground_truth_rows : Expected rows (pre-computed at task load time).138 error : SQLite error string (None if query ran successfully).139 step : Current step number (1-indexed) for penalty calculation.140 order_sensitive : If True, row order matters (queries with ORDER BY).141 """142 # ── Syntax ──────────────────────────────────────────────────────────143 syntax_ok = error is None and actual_rows is not None144 145 if not syntax_ok:146 return GradeResult(147 reward=0.0,148 syntax_ok=False,149 columns_match=False,150 row_count_match=False,151 exact_match=False,152 step_penalty=0.0,153 )154 155 gt_norm = _normalise_rows(ground_truth_rows)156 act_norm = _normalise_rows(actual_rows)157 158 gt_cols = set(gt_norm[0].keys()) if gt_norm else set()159 act_cols = set(act_norm[0].keys()) if act_norm else set()160 columns_match = act_cols == gt_cols161 row_count_match = len(act_norm) == len(gt_norm)162 163 # Exact match: if order matters, compare list; otherwise compare sorted sets164 if columns_match and row_count_match:165 if order_sensitive:166 exact_match = act_norm == gt_norm167 else:168 # Sort rows by their string representation for order-agnostic compare169 def _sort_key(r: Dict) -> str:170 return str(sorted(r.items()))171 exact_match = (172 sorted(act_norm, key=_sort_key) == sorted(gt_norm, key=_sort_key)173 )174 else:175 exact_match = False176 177 # ── Score assembly ────────────────────────────────────────────────178 raw = (179 0.10 # syntax180 + (0.20 if columns_match else 0.0)181 + (0.20 if row_count_match else 0.0)182 + (0.50 if exact_match else 0.0)183 )184 185 penalty = max(0.0, step - 1) * 0.05186 reward = float(max(0.0, min(1.0, raw - penalty)))187 188 return GradeResult(189 reward=reward,190 syntax_ok=syntax_ok,191 columns_match=columns_match,192 row_count_match=row_count_match,193 exact_match=exact_match,194 step_penalty=penalty,195 )196 197 198# ── Convenience: pre-compute ground truth rows ─────────────────────────────199 200def compute_ground_truth(201 conn: sqlite3.Connection,202 sql: str,203) -> List[Dict[str, Any]]:204 """Execute the ground-truth SQL and return normalised rows."""205 rows, error = execute_query(conn, sql)206 if error or rows is None:207 raise ValueError(f"Ground-truth SQL failed: {error}\nSQL: {sql}")208 return _normalise_rows(rows)209 210 211def has_order_by(sql: str) -> bool:212 """Heuristic: does the top-level query have an ORDER BY?"""213 # Simple check sufficient for our controlled task SQL214 return "ORDER BY" in sql.upper()215 