Team Ai
Apppublic

MauryaVivek/sql-data-quality-agent

sourceHugging Facemitupdated 6mo agoView on Hugging Face
0likes
environment.py268 linesDownload Raw Back to root
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        )