dkAmulet/sql-query-optimizer
0
1"""2Typed Pydantic models for the SQL Query Optimizer environment.3Scores are strictly in (0.0, 1.0) exclusive as required by OpenEnv validator.4"""5from __future__ import annotations6from pydantic import BaseModel, Field, field_validator7from typing import Optional, Dict, Any, List8 9 10class ExecutionMetrics(BaseModel):11 uses_index: bool = False12 full_scan_count: int = 013 query_plan: List[str] = Field(default_factory=list)14 15 16class SQLObservation(BaseModel):17 task_id: str18 task_name: str19 difficulty: str20 description: str21 schema_ddl: str22 slow_query: str23 current_query: str24 step_number: int25 max_steps: int26 slow_metrics: ExecutionMetrics27 last_feedback: str = ""28 last_reward: float = 0.00129 30 31class SQLAction(BaseModel):32 optimized_query: str = Field(33 ...,34 description="The rewritten SQL query to evaluate.",35 )36 37 38class RewardBreakdown(BaseModel):39 validity: float = Field(0.001, gt=0.0, lt=1.0)40 correctness: float = Field(0.001, gt=0.0, lt=1.0)41 performance: float = Field(0.001, gt=0.0, lt=1.0)42 style: float = Field(0.001, gt=0.0, lt=1.0)43 44 @field_validator('validity', 'correctness', 'performance', 'style', mode='before')45 @classmethod46 def clamp_strict(cls, v):47 return max(0.001, min(0.999, float(v)))48 49 50class SQLReward(BaseModel):51 value: float = Field(..., gt=0.0, lt=1.0)52 breakdown: RewardBreakdown53 feedback: str = Field(...,)54 55 @field_validator('value', mode='before')56 @classmethod57 def clamp_value(cls, v):58 return max(0.001, min(0.999, float(v)))59 60 61class StepResult(BaseModel):62 observation: SQLObservation63 reward: SQLReward64 done: bool65 info: Dict[str, Any] = Field(default_factory=dict)66 67 68class EnvironmentState(BaseModel):69 task_id: str70 step_number: int71 max_steps: int72 best_reward: float = 0.00173 done: bool74 current_query: str = ""75 