ritvik360/nl2sql-bench
0
1import os2import sys3from fastapi import FastAPI, Request4import uvicorn5 6sys.path.insert(0, "./server")7from environment import NL2SQLEnvironment8from models import NL2SQLAction9 10app = FastAPI()11env = NL2SQLEnvironment()12 13@app.post("/reset")14async def reset(request: Request):15 data = await request.json()16 # Now we take task_name directly from the API call17 task_name = data.get("task_name", "simple-filter")18 print(f"๐ Environment Resetting for Task: {task_name}")19 obs = env.reset(task_name=task_name)20 return {"observation": obs.__dict__}21 22@app.post("/step")23async def step(request: Request):24 data = await request.json()25 query = data.get("query", "")26 print(f"โฉ Executing SQL: {query[:60]}...")27 28 action = NL2SQLAction(query=query)29 obs = env.step(action)30 return {"observation": obs.__dict__}31 32if __name__ == "__main__":33 uvicorn.run(app, host="0.0.0.0", port=8000)