Team Ai
Apppublic

Kalletlamadhav/sql-optimization-env

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
main.py277 linesDownload Raw Back to server
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))