Skip to content

Commit 7ec8d4b

Browse files
Merge pull request #28 from ArshVermaGit/main
feat: implement SQLite persistence and database management
2 parents f3c6396 + 4df824f commit 7ec8d4b

7 files changed

Lines changed: 426 additions & 46 deletions

File tree

.gitignore

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -11,3 +11,8 @@ dist/
1111
build/
1212
.idea/
1313
.vscode/
14+
15+
# Persistence
16+
data/*.db
17+
data/*.db-shm
18+
data/*.db-wal

app.py

Lines changed: 205 additions & 44 deletions
Original file line numberDiff line numberDiff line change
@@ -1,60 +1,109 @@
11
import 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
512
from 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

718
from codereview_env.models import (
8-
TaskId, Action, ResetResult, StepResult, EpisodeResult
19+
TaskId, Action, ResetResult, StepResult, EpisodeResult, ActionRecord
920
)
1021
from 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 ────────────────────────────────────────────────────────
1238
app = 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 ────────────────────────────────────────────────────────────────────
2688
class ResetRequest(BaseModel):
2789
task_id: TaskId
2890
seed: int = 42
2991

30-
3192
class 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-
4496
class SubmitScore(BaseModel):
4597
agent_name: str
4698
task_id: TaskId
4799
score: float
48100
seed: int
49101

50-
51102
# ── WebSocket clients ─────────────────────────────────────────────────────────
52103
clients = set()
53104

54-
55105
async 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")
71142
def 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")
127289
async def websocket_endpoint(websocket: WebSocket):
@@ -135,7 +297,6 @@ async def websocket_endpoint(websocket: WebSocket):
135297
finally:
136298
clients.discard(websocket)
137299

138-
139300
if __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)

codereview_env/config.py

Lines changed: 27 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,27 @@
1+
from functools import lru_cache
2+
from pydantic_settings import BaseSettings, SettingsConfigDict
3+
4+
class Settings(BaseSettings):
5+
model_config = SettingsConfigDict(env_file=".env", env_file_encoding="utf-8", extra="ignore")
6+
7+
app_host: str = "0.0.0.0"
8+
app_port: int = 7860
9+
app_env: str = "development"
10+
11+
api_key: str = "changeme"
12+
api_key_enabled: bool = False
13+
14+
leaderboard_max_entries: int = 10
15+
16+
log_level: str = "INFO"
17+
18+
episode_ttl_seconds: int = 3600 # episodes expire after 1 hour
19+
rate_limit_per_minute: int = 60 # requests per minute per IP
20+
21+
# Persistence
22+
db_path: str = "./data/codereview.db"
23+
db_echo: bool = False # Set True to log all SQL queries
24+
25+
@lru_cache
26+
def get_settings() -> Settings:
27+
return Settings()

0 commit comments

Comments
 (0)