11import uuid
2- from typing import Dict
3-
4- from fastapi import FastAPI , HTTPException , WebSocket , WebSocketDisconnect
2+ import logging
3+ import asyncio
4+ import json
5+ from typing import Dict , List , Optional
6+ from datetime import datetime , timezone
7+
8+ from fastapi import FastAPI , HTTPException , WebSocket , WebSocketDisconnect , Depends , Security , Query , BackgroundTasks , Request
9+ from fastapi .responses import JSONResponse
10+ from fastapi .exceptions import RequestValidationError
11+ from fastapi .security .api_key import APIKeyHeader
512from pydantic import BaseModel
13+ from slowapi import Limiter , _rate_limit_exceeded_handler
14+ from slowapi .util import get_remote_address
15+ from slowapi .errors import RateLimitExceeded
16+ from sqlmodel import Session
617
718from codereview_env .models import (
8- TaskId , Action , ResetResult , StepResult , EpisodeResult
19+ TaskId , Action , ResetResult , StepResult , EpisodeResult , ActionRecord
920)
1021from codereview_env .env import CodeReviewEnv
22+ from codereview_env .config import get_settings
23+ from codereview_env .database import (
24+ create_db_and_tables , get_session , save_episode ,
25+ get_episode , get_leaderboard_db , submit_leaderboard , get_stats ,
26+ LeaderboardRecord
27+ )
28+
29+ # ── Logging ───────────────────────────────────────────────────────────────────
30+ settings = get_settings ()
31+ logging .basicConfig (
32+ level = getattr (logging , settings .log_level ),
33+ format = "%(asctime)s [%(levelname)s] %(name)s: %(message)s"
34+ )
35+ logger = logging .getLogger ("codereview_env" )
1136
37+ # ── App Initialization ────────────────────────────────────────────────────────
1238app = FastAPI (
1339 title = "AgentOrg CodeReview OpenEnv API" ,
1440 description = (
1541 "AI Senior Code Reviewer evaluation environment. "
1642 "Trains agents to detect bugs, security vulnerabilities, and architectural issues "
17- "in realistic Python PRs grounded in real-world incident patterns ."
43+ "in realistic Python PRs."
1844 ),
1945 version = "1.0.0" ,
2046)
2147
22- # Simple in-memory storage for active episodes
23- episodes : Dict [str , CodeReviewEnv ] = {}
48+ # ── Rate Limiting ─────────────────────────────────────────────────────────────
49+ limiter = Limiter (key_func = get_remote_address , default_limits = [f"{ settings .rate_limit_per_minute } /minute" ])
50+ app .state .limiter = limiter
51+ app .add_exception_handler (RateLimitExceeded , _rate_limit_exceeded_handler )
52+
53+ # ── API Key Authentication ────────────────────────────────────────────────────
54+ API_KEY_HEADER = APIKeyHeader (name = "X-API-Key" , auto_error = False )
2455
56+ async def verify_api_key (api_key : str = Security (API_KEY_HEADER )):
57+ if not settings .api_key_enabled :
58+ return # Auth disabled in development
59+ if api_key != settings .api_key :
60+ raise HTTPException (status_code = 403 , detail = "Invalid or missing API key" )
2561
62+ # ── Storage & TTL ─────────────────────────────────────────────────────────────
63+ episodes : Dict [str , CodeReviewEnv ] = {}
64+ episode_timestamps : Dict [str , datetime ] = {}
65+
66+ async def cleanup_expired_episodes ():
67+ """Remove episodes older than TTL."""
68+ while True :
69+ await asyncio .sleep (300 ) # run every 5 minutes
70+ cutoff = datetime .now (timezone .utc ).timestamp () - settings .episode_ttl_seconds
71+ expired = [
72+ eid for eid , ts in episode_timestamps .items ()
73+ if ts .timestamp () < cutoff
74+ ]
75+ for eid in expired :
76+ episodes .pop (eid , None )
77+ episode_timestamps .pop (eid , None )
78+ if expired :
79+ logger .info (f"Cleaned up { len (expired )} expired episodes" )
80+
81+ @app .on_event ("startup" )
82+ async def startup_event ():
83+ create_db_and_tables ()
84+ asyncio .create_task (cleanup_expired_episodes ())
85+ logger .info (f"CodeReview API started \u2014 DB at { settings .db_path } " )
86+
87+ # ── Models ────────────────────────────────────────────────────────────────────
2688class ResetRequest (BaseModel ):
2789 task_id : TaskId
2890 seed : int = 42
2991
30-
3192class ResetResponse (BaseModel ):
3293 episode_id : str
3394 result : ResetResult
3495
35-
36- # In-memory leaderboard
37- leaderboard : Dict [TaskId , list ] = {
38- TaskId .BUG_DETECTION : [],
39- TaskId .SECURITY_AUDIT : [],
40- TaskId .ARCHITECTURAL_REVIEW : []
41- }
42-
43-
4496class SubmitScore (BaseModel ):
4597 agent_name : str
4698 task_id : TaskId
4799 score : float
48100 seed : int
49101
50-
51102# ── WebSocket clients ─────────────────────────────────────────────────────────
52103clients = set ()
53104
54-
55105async def broadcast_event (data : dict ):
56106 from fastapi .encoders import jsonable_encoder
57- import json
58107 message = json .dumps (jsonable_encoder (data ))
59108 dead = set ()
60109 for client in clients :
@@ -64,30 +113,55 @@ async def broadcast_event(data: dict):
64113 dead .add (client )
65114 clients .difference_update (dead )
66115
116+ # ── Error Handlers ────────────────────────────────────────────────────────────
117+ @app .exception_handler (RequestValidationError )
118+ async def validation_exception_handler (request , exc ):
119+ return JSONResponse (
120+ status_code = 422 ,
121+ content = {
122+ "error" : "validation_error" ,
123+ "detail" : str (exc ),
124+ "status_code" : 422
125+ }
126+ )
127+
128+ @app .exception_handler (HTTPException )
129+ async def http_exception_handler (request , exc ):
130+ logger .warning (f"HTTP { exc .status_code } : { exc .detail } \u2014 { request .url } " )
131+ return JSONResponse (
132+ status_code = exc .status_code ,
133+ content = {
134+ "error" : exc .detail ,
135+ "status_code" : exc .status_code
136+ }
137+ )
67138
68139# ── Endpoints ─────────────────────────────────────────────────────────────────
69140
70141@app .get ("/health" )
71142def health_check ():
72143 return {
73- "status" : "ok" ,
74- "version" : "1.0.0" ,
144+ "status" : "ok" ,
145+ "version" : "1.0.0" ,
75146 "env_ready" : True ,
147+ "env" : settings .app_env ,
76148 "active_episodes" : len (episodes ),
149+ "auth_enabled" : settings .api_key_enabled
77150 }
78151
79-
80152@app .post ("/reset" , response_model = ResetResponse )
81- def reset_env (req : ResetRequest ):
153+ @limiter .limit (f"{ settings .rate_limit_per_minute } /minute" )
154+ def reset_env (request : Request , req : ResetRequest , _ : None = Depends (verify_api_key )):
82155 episode_id = str (uuid .uuid4 ())
83156 env = CodeReviewEnv ()
84157 result = env .reset (req .task_id , req .seed )
85158 episodes [episode_id ] = env
159+ episode_timestamps [episode_id ] = datetime .now (timezone .utc )
86160 return ResetResponse (episode_id = episode_id , result = result )
87161
88-
89162@app .post ("/step/{episode_id}" , response_model = StepResult )
90- async def step_env (episode_id : str , action : Action ):
163+ @limiter .limit (f"{ settings .rate_limit_per_minute } /minute" )
164+ async def step_env (request : Request , episode_id : str , action : Action , _ : None = Depends (verify_api_key )):
91165 if episode_id not in episodes :
92166 raise HTTPException (status_code = 404 , detail = "Episode not found" )
93167
@@ -99,29 +173,117 @@ async def step_env(episode_id: str, action: Action):
99173 except RuntimeError as e :
100174 raise HTTPException (status_code = 400 , detail = str (e ))
101175
102-
103176@app .get ("/result/{episode_id}" , response_model = EpisodeResult )
104- def get_result (episode_id : str ):
105- if episode_id not in episodes :
177+ def get_result (
178+ episode_id : str ,
179+ session : Session = Depends (get_session ),
180+ _ : None = Depends (verify_api_key )
181+ ):
182+ # Try in-memory (active episode)
183+ if episode_id in episodes :
184+ env = episodes [episode_id ]
185+ result = env .get_final_result ()
186+ result .episode_id = episode_id
187+ # If done, persist and remove from memory
188+ if env .done :
189+ save_episode (session , result )
190+ del episodes [episode_id ]
191+ episode_timestamps .pop (episode_id , None )
192+ return result
193+
194+ # Fall back to DB (completed episode)
195+ record = get_episode (session , episode_id )
196+ if not record :
106197 raise HTTPException (status_code = 404 , detail = "Episode not found" )
107- return episodes [episode_id ].get_final_result ()
108-
198+
199+ return EpisodeResult (
200+ episode_id = record .episode_id ,
201+ task_id = TaskId (record .task_id ),
202+ scenario_hash = record .scenario_hash ,
203+ seed = record .seed ,
204+ final_score = record .final_score ,
205+ steps_taken = record .steps_taken ,
206+ issues_found = record .issues_found ,
207+ issues_total = record .issues_total ,
208+ noise_penalties = record .noise_penalties ,
209+ terminated_reason = record .terminated_reason ,
210+ history = [ActionRecord (** r ) for r in json .loads (record .history_json or "[]" )]
211+ )
109212
110213@app .get ("/leaderboard" )
111- def get_leaderboard ():
112- return leaderboard
113-
214+ def get_leaderboard (
215+ task_id : Optional [TaskId ] = None ,
216+ limit : int = Query (default = 10 , ge = 1 , le = 50 ),
217+ offset : int = Query (default = 0 , ge = 0 ),
218+ session : Session = Depends (get_session )
219+ ):
220+ tasks_to_query = [task_id ] if task_id else list (TaskId )
221+ result = {}
222+ for t in tasks_to_query :
223+ entries , total = get_leaderboard_db (session , t .value , limit , offset )
224+ result [t .value ] = {
225+ "entries" : [e .model_dump () for e in entries ],
226+ "total" : total
227+ }
228+ if task_id :
229+ return result [task_id .value ]
230+ return result
114231
115232@app .post ("/submit" )
116- def submit_to_leaderboard (submission : SubmitScore ):
117- entries = leaderboard .get (submission .task_id , [])
118- new_entry = submission .model_dump ()
119- entries .append (new_entry )
120- entries .sort (key = lambda x : x ["score" ], reverse = True )
121- rank = entries .index (new_entry ) + 1 # capture rank before slicing
122- leaderboard [submission .task_id ] = entries [:5 ]
123- return {"status" : "submitted" , "rank" : rank if rank <= 5 else None }
233+ @limiter .limit (f"{ settings .rate_limit_per_minute } /minute" )
234+ def submit_to_leaderboard (
235+ request : Request ,
236+ submission : SubmitScore ,
237+ session : Session = Depends (get_session ),
238+ _ : None = Depends (verify_api_key )
239+ ):
240+ rank = submit_leaderboard (
241+ session ,
242+ agent_name = submission .agent_name ,
243+ task_id = submission .task_id .value ,
244+ score = submission .score ,
245+ seed = submission .seed
246+ )
247+ return {"status" : "submitted" , "rank" : rank if rank > 0 else None }
248+
249+ @app .get ("/stats" )
250+ def get_aggregate_stats (session : Session = Depends (get_session )):
251+ return get_stats (session )
252+
253+ @app .get ("/episodes/{episode_id}/replay" )
254+ def get_episode_replay (
255+ episode_id : str ,
256+ session : Session = Depends (get_session ),
257+ _ : None = Depends (verify_api_key )
258+ ):
259+ record = get_episode (session , episode_id )
260+ if not record :
261+ raise HTTPException (status_code = 404 , detail = "Episode not found or not yet completed" )
262+ return {
263+ "episode_id" : record .episode_id ,
264+ "task_id" : record .task_id ,
265+ "scenario_hash" : record .scenario_hash ,
266+ "final_score" : record .final_score ,
267+ "history" : json .loads (record .history_json or "[]" ),
268+ "created_at" : record .created_at
269+ }
124270
271+ @app .get ("/episodes" )
272+ def list_episodes (
273+ _ : None = Depends (verify_api_key ),
274+ limit : int = Query (default = 20 , ge = 1 , le = 100 )
275+ ):
276+ episode_list = [
277+ {
278+ "episode_id" : eid ,
279+ "task_id" : env .task_id ,
280+ "step_count" : env .observation .step_count ,
281+ "done" : env .done ,
282+ "created_at" : episode_timestamps .get (eid , "" ).isoformat () if episode_timestamps .get (eid ) else ""
283+ }
284+ for eid , env in list (episodes .items ())[:limit ]
285+ ]
286+ return {"episodes" : episode_list , "total" : len (episodes )}
125287
126288@app .websocket ("/ws/events" )
127289async def websocket_endpoint (websocket : WebSocket ):
@@ -135,7 +297,6 @@ async def websocket_endpoint(websocket: WebSocket):
135297 finally :
136298 clients .discard (websocket )
137299
138-
139300if __name__ == "__main__" :
140301 import uvicorn
141- uvicorn .run (app , host = "0.0.0.0" , port = 7860 )
302+ uvicorn .run (app , host = settings . app_host , port = settings . app_port )
0 commit comments