Spaces:
Sleeping
Sleeping
| import os | |
| import time | |
| import asyncio | |
| from typing import Optional | |
| from contextlib import asynccontextmanager | |
| from fastapi import FastAPI, HTTPException, Request | |
| from fastapi.middleware.cors import CORSMiddleware | |
| from fastapi.responses import JSONResponse, Response | |
| from pydantic import BaseModel, ValidationError | |
| from env.environment import environment | |
| from env.models import ( | |
| Action, Observation, EpisodeState, | |
| DifficultyLevel, ActionType, | |
| StepResponse, ResetResponse, TaskListResponse, | |
| BaselineResponse, BaselineResult, | |
| GraderRequest, GraderResponse, | |
| HealthResponse, TaskInfo | |
| ) | |
| from env.tasks import task_manager, ACTION_SCHEMA | |
| from env.graders import grade | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| # STARTUP / SHUTDOWN | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| _startup_time = time.time() | |
| async def lifespan(app: FastAPI): | |
| environment.reset(difficulty="easy") | |
| yield | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| # APP DEFINITION | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| app = FastAPI( | |
| title = "SQL Query Debugger β OpenEnv Environment", | |
| description = ( | |
| "An OpenEnv-compliant reinforcement learning environment where AI agents " | |
| "learn to debug SQL queries across syntax errors, logic bugs, and performance issues. " | |
| "Built for the META x PyTorch x SST OpenEnv Hackathon." | |
| ), | |
| version = "1.0.0", | |
| lifespan = lifespan, | |
| docs_url = "/docs", | |
| redoc_url = "/redoc", | |
| ) | |
| app.add_middleware( | |
| CORSMiddleware, | |
| allow_origins = ["*"], | |
| allow_credentials = True, | |
| allow_methods = ["*"], | |
| allow_headers = ["*"], | |
| ) | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| # GLOBAL EXCEPTION HANDLER | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| async def global_exception_handler(request: Request, exc: Exception): | |
| return JSONResponse( | |
| status_code = 500, | |
| content = {"error": str(exc), "type": type(exc).__name__} | |
| ) | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| # FAVICON β fix 404 | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| async def favicon(): | |
| """Returns 204 No Content instead of 404 for favicon requests.""" | |
| return Response(status_code=204) | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| # 1. /health β GET | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| async def health(): | |
| """Liveness check. Always returns 200. Used by HF Space health monitoring.""" | |
| return HealthResponse( | |
| status = "ok", | |
| version = "1.0.0", | |
| uptime = round(time.time() - _startup_time, 2) | |
| ) | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| # 2. /reset β POST | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| class ResetBody(BaseModel): | |
| difficulty: Optional[str] = None | |
| task_id: Optional[str] = None | |
| async def reset(body: ResetBody = ResetBody()): | |
| """ | |
| Starts a fresh episode. Returns the initial Observation the agent sees. | |
| Edge case: always returns valid Observation even if dataset issues occur. | |
| """ | |
| try: | |
| obs = environment.reset( | |
| difficulty = body.difficulty, | |
| task_id = body.task_id | |
| ) | |
| return obs | |
| except ValueError as e: | |
| raise HTTPException(status_code=400, detail=str(e)) | |
| except Exception as e: | |
| raise HTTPException(status_code=500, detail=f"Reset failed: {str(e)}") | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| # 3. /step β POST | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| async def step(action: Action): | |
| """ | |
| Submits an action to the environment. | |
| Returns (observation, reward, done, info). | |
| Edge cases: null action, malformed payload, episode already done. | |
| """ | |
| try: | |
| response = environment.step(action) | |
| return response | |
| except ValidationError as e: | |
| from env.models import Reward | |
| return StepResponse( | |
| observation = environment._build_observation(), | |
| reward = Reward( | |
| score = -0.1, | |
| breakdown = {"validation_error": -0.1}, | |
| feedback = f"Malformed action: {str(e)}" | |
| ), | |
| done = False, | |
| info = {"error": "validation_error", "detail": str(e)} | |
| ) | |
| except Exception as e: | |
| raise HTTPException(status_code=500, detail=f"Step failed: {str(e)}") | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| # 4. /state β GET | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| async def state(): | |
| """ | |
| Returns full current environment state. | |
| Works before reset() is called β returns default empty state. | |
| Always JSON-serializable. Never crashes. | |
| """ | |
| return environment.state() | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| # 5. /tasks β GET | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| async def tasks(): | |
| """ | |
| Lists all 15 tasks with full action schema definitions. | |
| Validator checks for action field definitions, not just task names. | |
| """ | |
| all_tasks = task_manager.list_all_tasks() | |
| return TaskListResponse( | |
| tasks = all_tasks, | |
| total = len(all_tasks), | |
| action_types = [a.value for a in ActionType] | |
| ) | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| # 6. /grader β POST | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| async def grader(request: GraderRequest): | |
| """ | |
| Grades a completed episode action. | |
| Returns float score 0.0-1.0. Never crashes. | |
| Edge cases: null action β 0.0, unknown task β 0.0. | |
| """ | |
| try: | |
| if request.action is None: | |
| return GraderResponse( | |
| score = 0.0, | |
| feedback = "No action provided for grading.", | |
| breakdown = {"error": "null_action"} | |
| ) | |
| score, breakdown, feedback = grade(request.action, request.task_id) | |
| return GraderResponse( | |
| score = score, | |
| feedback = feedback, | |
| breakdown = breakdown | |
| ) | |
| except Exception as e: | |
| return GraderResponse( | |
| score = 0.0, | |
| feedback = f"Grader error: {str(e)}", | |
| breakdown = {"error": str(e)} | |
| ) | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| # 7. /baseline β POST | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| async def baseline(): | |
| """ | |
| Runs the baseline agent against all 3 difficulty levels. | |
| Returns scores JSON. Must complete within 60 seconds. | |
| Edge case: OPENAI_API_KEY not set β continues with rule-based agent. | |
| """ | |
| try: | |
| import baseline as baseline_module | |
| results = await asyncio.wait_for( | |
| asyncio.to_thread(baseline_module.run_baseline), | |
| timeout=55.0 | |
| ) | |
| return results | |
| except asyncio.TimeoutError: | |
| return BaselineResponse( | |
| results=[ | |
| BaselineResult( | |
| task_id = "timeout", | |
| difficulty = DifficultyLevel.EASY, | |
| score = 0.0, | |
| steps = 0, | |
| feedback = "Baseline timed out after 55 seconds." | |
| ) | |
| ], | |
| average_score=0.0 | |
| ) | |
| except Exception as e: | |
| return BaselineResponse( | |
| results=[ | |
| BaselineResult( | |
| task_id = "error", | |
| difficulty = DifficultyLevel.EASY, | |
| score = 0.0, | |
| steps = 0, | |
| feedback = f"Baseline error: {str(e)}" | |
| ) | |
| ], | |
| average_score=0.0 | |
| ) | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| # ROOT β project info | |
| # βββββββββββββββββββββββββββββββββββββββββββββ | |
| async def root(): | |
| return { | |
| "name": "SQL Query Debugger β OpenEnv Environment", | |
| "version": "1.0.0", | |
| "docs": "/docs", | |
| "health": "/health", | |
| "endpoints": ["/reset", "/step", "/state", "/tasks", "/grader", "/baseline", "/health"], | |
| "hackathon": "META x PyTorch x SST OpenEnv Hackathon", | |
| "domain": "SQL Query Debugging", | |
| "tasks_count": 15, | |
| } |