Kalletlamadhav/sql-optimization-env
0
1# server/main.py2import sqlite33from contextlib import asynccontextmanager4from pathlib import Path5from typing import Any6 7from fastapi import FastAPI, HTTPException, Query8from fastapi.responses import HTMLResponse9from fastapi.staticfiles import StaticFiles10 11from .environment import SQLOptEnvironment12from .models import SQLOptAction, SQLOptObservation, EnvironmentState13 14env: SQLOptEnvironment = None15 16ROOT = Path(__file__).resolve().parent.parent17FIXTURE_DB = ROOT / "data" / "fixtures" / "benchmark_seed42.db"18SCHEMA_DIR = ROOT / "data" / "schemas"19 20DOMAIN_SPECS: list[dict[str, Any]] = [21 {22 "id": "gst",23 "name": "GST",24 "description": "Goods & Services Tax — B2B invoices, HSN line items, intra/inter-state tax splits.",25 "schema_file": "gst_schema.sql",26 "tables": ["gst_invoice_records", "gst_invoice_items"],27 "real_world_scale": "At 100k invoice scale: ~100k header rows and ~300k line items (typical seed ratios).",28 "sample_queries": [29 "SELECT state_code, COUNT(*) AS invoices FROM gst_invoice_records GROUP BY state_code ORDER BY invoices DESC LIMIT 10;",30 "SELECT invoice_id, taxable_value, cgst_amount, sgst_amount, igst_amount FROM gst_invoice_records WHERE igst_amount = 0 AND cgst_amount > 0 AND ABS(cgst_amount - sgst_amount) > 1 LIMIT 20;",31 "SELECT gstin_supplier, COUNT(*) FROM gst_invoice_records WHERE invoice_date >= '2025-01-01' GROUP BY gstin_supplier ORDER BY COUNT(*) DESC LIMIT 5;",32 ],33 },34 {35 "id": "pds",36 "name": "PDS",37 "description": "Public Distribution System — ration cards, monthly allotments, fair-price shops.",38 "schema_file": "pds_schema.sql",39 "tables": ["ration_card_beneficiaries", "pds_allotments"],40 "real_world_scale": "At 100k invoice-scale seed: ~20k cardholders and ~80k allotment rows.",41 "sample_queries": [42 "SELECT state_code, COUNT(*) FROM ration_card_beneficiaries GROUP BY state_code;",43 "SELECT r.card_id, COUNT(*) FROM ration_card_beneficiaries r JOIN pds_allotments a ON r.card_id = a.card_id GROUP BY r.card_id LIMIT 10;",44 "SELECT commodity, SUM(offtake_qty_kg) FROM pds_allotments GROUP BY commodity;",45 ],46 },47 {48 "id": "railway",49 "name": "Railway (IRCTC-style)",50 "description": "Train master and PNR bookings — Tatkal / availability style workloads.",51 "schema_file": "railway_schema.sql",52 "tables": ["railway_trains", "railway_pnr_bookings"],53 "real_world_scale": "Fixed train catalog plus booking volume tied to seed row count; journey dates skew to festival months.",54 "sample_queries": [55 "SELECT strftime('%m', journey_date) AS m, COUNT(*) FROM railway_pnr_bookings GROUP BY m ORDER BY m;",56 "SELECT train_no, journey_date, COUNT(*) FROM railway_pnr_bookings GROUP BY train_no, journey_date ORDER BY COUNT(*) DESC LIMIT 15;",57 "SELECT booking_status, COUNT(*) FROM railway_pnr_bookings WHERE booked_via = 'TATKAL' GROUP BY booking_status;",58 ],59 },60 {61 "id": "mgnrega",62 "name": "MGNREGA",63 "description": "Rural employment guarantee — workers, muster attendance, wage payments.",64 "schema_file": "mgnrega_schema.sql",65 "tables": ["mgnrega_workers", "mgnrega_attendance", "mgnrega_payments"],66 "real_world_scale": "Worker count ~ n/3 of GST invoices in seed; attendance and payment rows scale with workers.",67 "sample_queries": [68 "SELECT state_code, COUNT(*) FROM mgnrega_workers GROUP BY state_code;",69 "SELECT w.worker_id, SUM(a.days_worked) FROM mgnrega_workers w JOIN mgnrega_attendance a ON w.worker_id = a.worker_id GROUP BY w.worker_id LIMIT 10;",70 "SELECT payment_month, SUM(amount_due), SUM(amount_paid) FROM mgnrega_payments GROUP BY payment_month ORDER BY payment_month LIMIT 12;",71 ],72 },73]74 75 76def _fixture_db_path() -> Path:77 if env is not None and getattr(env, "_db_path", None):78 return env._db_path79 return FIXTURE_DB80 81 82def _table_row_counts() -> dict[str, int]:83 db = _fixture_db_path()84 if not db.is_file():85 return {}86 out: dict[str, int] = {}87 conn = sqlite3.connect(str(db))88 try:89 names = [90 r[0]91 for r in conn.execute(92 "SELECT name FROM sqlite_master WHERE type='table' AND name NOT LIKE 'sqlite_%' ORDER BY name"93 ).fetchall()94 ]95 for name in names:96 try:97 out[name] = int(98 conn.execute(f"SELECT COUNT(*) FROM {name}").fetchone()[0]99 )100 except sqlite3.Error:101 out[name] = -1102 finally:103 conn.close()104 return out105 106 107def _state_counts(conn: sqlite3.Connection, table: str) -> dict[str, int]:108 try:109 rows = conn.execute(110 f"SELECT state_code, COUNT(*) FROM {table} GROUP BY state_code"111 ).fetchall()112 return {str(k): int(v) for k, v in rows if k is not None}113 except sqlite3.Error:114 return {}115 116 117@asynccontextmanager118async def lifespan(app: FastAPI):119 global env120 env = SQLOptEnvironment()121 122 # Mount static files for dashboard123 static_dir = Path(__file__).parent.parent / "static"124 if static_dir.exists():125 app.mount("/static", StaticFiles(directory=str(static_dir)), name="static")126 127 print("Environment initialized", flush=True)128 yield129 130 131app = FastAPI(132 title="SQL Optimization RL Environment",133 description="OpenEnv-compliant SQL optimization environment with Indian data domains",134 version="1.0.0",135 lifespan=lifespan136)137 138 139@app.get("/", response_class=HTMLResponse, include_in_schema=False)140async def dashboard():141 """Serve the main dashboard UI."""142 html_file = Path(__file__).parent.parent / "static" / "index.html"143 if html_file.exists():144 return HTMLResponse(content=html_file.read_text(encoding="utf-8"))145 return HTMLResponse(content="Dashboard not found. Add static/index.html")146 147 148@app.get("/health")149async def health():150 row_counts = _table_row_counts()151 return {152 "status": "ok",153 "environment": "sql-optimization-env",154 "version": "1.0.0",155 "database_path": str(_fixture_db_path()),156 "table_row_counts": row_counts,157 "total_rows": sum(v for v in row_counts.values() if v >= 0),158 }159 160 161@app.get("/domains")162async def domains():163 """Per-domain metadata, table counts, sample SQL, DDL, and state-level heatmap inputs."""164 row_counts = _table_row_counts()165 db = _fixture_db_path()166 gst_by_state: dict[str, int] = {}167 pds_by_state: dict[str, int] = {}168 if db.is_file():169 conn = sqlite3.connect(str(db))170 try:171 gst_by_state = _state_counts(conn, "gst_invoice_records")172 pds_by_state = _state_counts(conn, "ration_card_beneficiaries")173 finally:174 conn.close()175 176 payload: list[dict[str, Any]] = []177 for spec in DOMAIN_SPECS:178 schema_path = SCHEMA_DIR / spec["schema_file"]179 ddl = (180 schema_path.read_text(encoding="utf-8")181 if schema_path.is_file()182 else f"-- Schema file not found: {schema_path.name}\n"183 )184 tables_out = []185 for t in spec["tables"]:186 tables_out.append({"table": t, "row_count": row_counts.get(t, 0)})187 entry: dict[str, Any] = {188 "domain": spec["name"],189 "id": spec["id"],190 "description": spec["description"],191 "real_world_scale": spec["real_world_scale"],192 "tables": tables_out,193 "sample_queries": spec["sample_queries"],194 "schema_ddl": ddl,195 }196 if spec["id"] == "gst":197 entry["state_row_counts"] = gst_by_state198 elif spec["id"] == "pds":199 entry["state_row_counts"] = pds_by_state200 payload.append(entry)201 202 return {203 "domains": payload,204 "map": {205 "gst_by_state": gst_by_state,206 "pds_by_state": pds_by_state,207 "state_labels": {208 "01": "JK",209 "02": "HP",210 "03": "PB",211 "06": "HR",212 "07": "DL",213 "08": "RJ",214 "09": "UP",215 "10": "BR",216 "11": "SK",217 "12": "AR",218 "13": "NL",219 "14": "MN",220 "18": "AS",221 "19": "WB",222 "21": "OR",223 "22": "CG",224 "23": "MP",225 "24": "GJ",226 "27": "MH",227 "29": "KA",228 "32": "KL",229 "33": "TN",230 "36": "TG",231 "37": "AP",232 },233 },234 }235 236 237@app.get("/reset")238async def reset(task_id: str = Query(default=None)) -> SQLOptObservation:239 try:240 return env.reset(task_id=task_id)241 except Exception as e:242 raise HTTPException(status_code=500, detail=str(e))243 244 245@app.post("/step")246async def step(action: SQLOptAction):247 try:248 return env.step(action)249 except Exception as e:250 raise HTTPException(status_code=500, detail=str(e))251 252 253@app.get("/state")254async def state() -> EnvironmentState:255 try:256 return env.state()257 except Exception as e:258 raise HTTPException(status_code=500, detail=str(e))259 260 261@app.get("/tasks")262async def list_tasks():263 """Return all available tasks with metadata."""264 try:265 tasks = env.task_registry.all_tasks()266 return [267 {268 "task_id": t.task_id,269 "difficulty": t.difficulty,270 "curriculum_level": t.curriculum_level,271 "expected_pattern": t.expected_pattern,272 "tables": t.tables273 }274 for t in tasks275 ]276 except Exception as e:277 raise HTTPException(status_code=500, detail=str(e))