MauryaVivek/sql-data-quality-agent
0
1"""2environment.py3==============4Core OpenEnv environment class for the SQL Data Quality Agent.5 6Implements the standard OpenEnv interface:7 - reset(task_id, seed) -> DataQualityObservation8 - step(action) -> (DataQualityObservation, reward, done, info)9 - state() -> DataQualityState10 11All models are typed Pydantic BaseModels for spec compliance.12"""13 14import math15import sqlite316import uuid17from typing import Any, Dict, List, Optional, Tuple18 19from pydantic import BaseModel, Field, model_validator20 21from tasks import TASK_REGISTRY, QualityReport, get_task, clamp_score, clamp_ratio22from data_generator import TASK_DB_GENERATORS23from reward import compute_reward24 25 26# ---------------------------------------------------------------------------27# Typed Pydantic Models (OpenEnv spec requirement)28# ---------------------------------------------------------------------------29 30class DataQualityAction(BaseModel):31 """Action: a single SQL statement + optional agent rationale."""32 sql: str = Field(..., description="SQL statement to execute (SELECT/UPDATE/DELETE/INSERT)")33 rationale: str = Field(default="", description="Agent's reasoning for this action")34 35 36class DataQualityObservation(BaseModel):37 """Observation returned after reset() or step()."""38 task_id: str39 task_description: str40 table_schema: Dict[str, Dict[str, str]] = Field(41 description="Table name -> {column: type}"42 )43 sample_rows: Dict[str, List[Dict[str, Any]]] = Field(44 description="Table name -> list of up to 20 sample rows"45 )46 quality_report: QualityReport47 last_action_result: str = Field(default="", description="'success' or error message")48 step: int = 049 done: bool = False50 hints: List[str] = Field(default_factory=list)51 52 53class DataQualityState(BaseModel):54 """Full internal state (for the /state endpoint)."""55 episode_id: str56 task_id: str57 step: int58 max_steps: int59 current_score: float60 cumulative_reward: float61 tables: List[str]62 db_row_counts: Dict[str, int]63 64 @model_validator(mode="after")65 def _ensure_scores_valid(self):66 """Ensure current_score and cumulative_reward are never exactly 0.0 or 1.0."""67 self.current_score = clamp_score(self.current_score)68 self.cumulative_reward = clamp_score(self.cumulative_reward)69 return self70 71 72# ---------------------------------------------------------------------------73# Main Environment Class74# ---------------------------------------------------------------------------75 76class DataQualityEnv:77 """78 SQL Data Quality Agent Environment.79 80 The agent receives a dirty SQLite database and must issue SQL statements81 to improve its data quality score (0.0 -> 1.0).82 """83 84 def __init__(self):85 self._conn: Optional[sqlite3.Connection] = None86 self._task_id: Optional[str] = None87 self._episode_id: Optional[str] = None88 self._step: int = 089 self._max_steps: int = 090 self._prev_score: float = 0.591 self._cumulative_reward: float = 0.0192 self._done: bool = False93 self._seed: int = 4294 95 # ------------------------------------------------------------------96 # Public API97 # ------------------------------------------------------------------98 99 def reset(self, task_id: str = "null_patrol", seed: int = 42) -> DataQualityObservation:100 """101 Initialise a fresh episode for the given task.102 Returns the initial observation.103 """104 task_data = get_task(task_id)105 if task_data is None:106 raise ValueError(f"Unknown task_id: '{task_id}'. Valid tasks: {list(TASK_REGISTRY.keys())}")107 108 # Build a fresh dirty database109 generator = TASK_DB_GENERATORS[task_id]110 self._conn = generator(seed=seed)111 self._task_id = task_id112 self._episode_id = str(uuid.uuid4())113 self._step = 0114 self._max_steps = task_data["meta"].max_steps115 self._done = False116 self._cumulative_reward = 0.01 # Start at safe non-zero117 self._seed = seed118 119 # Initial quality score120 report = task_data["grader"](self._conn)121 # Ensure clamped — grader already clamps but be safe122 report.overall_score = clamp_score(report.overall_score)123 self._prev_score = report.overall_score124 125 return self._build_observation(report, last_result="", task_data=task_data)126 127 def step(self, action: DataQualityAction) -> Tuple[DataQualityObservation, float, bool, Dict]:128 """129 Execute one SQL action, compute reward, and return (obs, reward, done, info).130 """131 if self._conn is None or self._done:132 raise RuntimeError("Call reset() before step(), or episode is already done.")133 134 task_data = get_task(self._task_id)135 self._step += 1136 137 # Execute SQL138 result, last_result = self._execute_sql(action.sql)139 140 # Compute new quality score141 report = task_data["grader"](self._conn)142 curr_score = clamp_score(report.overall_score)143 report.overall_score = curr_score # ensure observation also has clamped value144 145 # Compute reward — already returns value in (0.01, 0.99)146 reward = compute_reward(147 prev_score=self._prev_score,148 curr_score=curr_score,149 sql=action.sql,150 action_result=last_result,151 step=self._step,152 max_steps=self._max_steps,153 success_threshold=task_data["meta"].success_threshold,154 )155 # Belt-and-suspenders: clamp once more156 reward = clamp_score(reward)157 self._cumulative_reward = clamp_score(self._cumulative_reward + reward)158 self._prev_score = curr_score159 160 # Check done conditions161 threshold = task_data["meta"].success_threshold162 self._done = (curr_score >= threshold) or (self._step >= self._max_steps)163 164 obs = self._build_observation(report, last_result=last_result, task_data=task_data)165 info = {166 "episode_id": self._episode_id,167 "cumulative_reward": clamp_score(self._cumulative_reward),168 "step": self._step,169 "success": curr_score >= threshold,170 "sql_result": result,171 }172 return obs, reward, self._done, info173 174 def state(self) -> DataQualityState:175 """Return the current non-observation state (episode metadata)."""176 if self._conn is None:177 raise RuntimeError("Call reset() first.")178 179 cur = self._conn.cursor()180 cur.execute("SELECT name FROM sqlite_master WHERE type='table'")181 tables = [row[0] for row in cur.fetchall()]182 183 row_counts = {}184 for table in tables:185 cur.execute(f"SELECT COUNT(*) FROM '{table}'")186 row_counts[table] = cur.fetchone()[0]187 188 return DataQualityState(189 episode_id=self._episode_id,190 task_id=self._task_id,191 step=self._step,192 max_steps=self._max_steps,193 current_score=clamp_score(self._prev_score),194 cumulative_reward=clamp_score(self._cumulative_reward),195 tables=tables,196 db_row_counts=row_counts,197 )198 199 # ------------------------------------------------------------------200 # Internal helpers201 # ------------------------------------------------------------------202 203 def _execute_sql(self, sql: str) -> Tuple[Any, str]:204 """Execute a SQL statement. Returns (result, status_string)."""205 cur = self._conn.cursor()206 try:207 cur.execute(sql)208 self._conn.commit()209 rows = cur.fetchall()210 if rows:211 result = [dict(row) for row in rows]212 else:213 result = f"OK ({cur.rowcount} rows affected)"214 return result, "success"215 except sqlite3.Error as e:216 return None, f"error: {str(e)}"217 except Exception as e:218 return None, f"error: {str(e)}"219 220 def _get_schema(self) -> Dict[str, Dict[str, str]]:221 """Return table schemas as {table: {col: type}}."""222 cur = self._conn.cursor()223 cur.execute("SELECT name FROM sqlite_master WHERE type='table'")224 tables = [row[0] for row in cur.fetchall()]225 226 schema = {}227 for table in tables:228 cur.execute(f"PRAGMA table_info('{table}')")229 schema[table] = {row[1]: row[2] for row in cur.fetchall()}230 return schema231 232 def _get_samples(self) -> Dict[str, List[Dict[str, Any]]]:233 """Return up to 20 sample rows per table."""234 cur = self._conn.cursor()235 cur.execute("SELECT name FROM sqlite_master WHERE type='table'")236 tables = [row[0] for row in cur.fetchall()]237 238 samples = {}239 for table in tables:240 cur.execute(f"SELECT * FROM '{table}' LIMIT 20")241 samples[table] = [dict(row) for row in cur.fetchall()]242 return samples243 244 def _build_observation(245 self,246 report: QualityReport,247 last_result: str,248 task_data: Dict,249 ) -> DataQualityObservation:250 # Safety: ensure overall_score is clamped before building observation251 report.overall_score = clamp_score(report.overall_score)252 # Also ensure all ratio fields are clamped253 report.null_ratio = clamp_ratio(report.null_ratio)254 report.duplicate_ratio = clamp_ratio(report.duplicate_ratio)255 report.type_error_ratio = clamp_ratio(report.type_error_ratio)256 report.constraint_violation_ratio = clamp_ratio(report.constraint_violation_ratio)257 report.value_error_ratio = clamp_ratio(report.value_error_ratio)258 return DataQualityObservation(259 task_id=self._task_id,260 task_description=task_data["meta"].description,261 table_schema=self._get_schema(),262 sample_rows=self._get_samples(),263 quality_report=report,264 last_action_result=last_result,265 step=self._step,266 done=self._done,267 hints=task_data.get("hints", []),268 )