ArshVerma commited on
Commit
4df824f
·
1 Parent(s): f3c6396

feat: add SQLite persistence for episodes, results, and leaderboard via SQLModel

Browse files
.gitignore CHANGED
@@ -11,3 +11,8 @@ dist/
11
  build/
12
  .idea/
13
  .vscode/
 
 
 
 
 
 
11
  build/
12
  .idea/
13
  .vscode/
14
+
15
+ # Persistence
16
+ data/*.db
17
+ data/*.db-shm
18
+ data/*.db-wal
app.py CHANGED
@@ -1,60 +1,109 @@
1
  import uuid
2
- from typing import Dict
3
-
4
- from fastapi import FastAPI, HTTPException, WebSocket, WebSocketDisconnect
 
 
 
 
 
 
 
5
  from pydantic import BaseModel
 
 
 
 
6
 
7
  from codereview_env.models import (
8
- TaskId, Action, ResetResult, StepResult, EpisodeResult
9
  )
10
  from codereview_env.env import CodeReviewEnv
 
 
 
 
 
 
 
 
 
 
 
 
 
 
11
 
 
12
  app = FastAPI(
13
  title="AgentOrg CodeReview OpenEnv API",
14
  description=(
15
  "AI Senior Code Reviewer evaluation environment. "
16
  "Trains agents to detect bugs, security vulnerabilities, and architectural issues "
17
- "in realistic Python PRs grounded in real-world incident patterns."
18
  ),
19
  version="1.0.0",
20
  )
21
 
22
- # Simple in-memory storage for active episodes
23
- episodes: Dict[str, CodeReviewEnv] = {}
 
 
 
 
 
24
 
 
 
 
 
 
25
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
26
  class ResetRequest(BaseModel):
27
  task_id: TaskId
28
  seed: int = 42
29
 
30
-
31
  class ResetResponse(BaseModel):
32
  episode_id: str
33
  result: ResetResult
34
 
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
-
44
  class SubmitScore(BaseModel):
45
  agent_name: str
46
  task_id: TaskId
47
  score: float
48
  seed: int
49
 
50
-
51
  # ── WebSocket clients ─────────��───────────────────────────────────────────────
52
  clients = set()
53
 
54
-
55
  async def broadcast_event(data: dict):
56
  from fastapi.encoders import jsonable_encoder
57
- import json
58
  message = json.dumps(jsonable_encoder(data))
59
  dead = set()
60
  for client in clients:
@@ -64,30 +113,55 @@ async def broadcast_event(data: dict):
64
  dead.add(client)
65
  clients.difference_update(dead)
66
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
67
 
68
  # ── Endpoints ─────────────────────────────────────────────────────────────────
69
 
70
  @app.get("/health")
71
  def health_check():
72
  return {
73
- "status": "ok",
74
- "version": "1.0.0",
75
  "env_ready": True,
 
76
  "active_episodes": len(episodes),
 
77
  }
78
 
79
-
80
  @app.post("/reset", response_model=ResetResponse)
81
- def reset_env(req: ResetRequest):
 
82
  episode_id = str(uuid.uuid4())
83
  env = CodeReviewEnv()
84
  result = env.reset(req.task_id, req.seed)
85
  episodes[episode_id] = env
 
86
  return ResetResponse(episode_id=episode_id, result=result)
87
 
88
-
89
  @app.post("/step/{episode_id}", response_model=StepResult)
90
- async def step_env(episode_id: str, action: Action):
 
91
  if episode_id not in episodes:
92
  raise HTTPException(status_code=404, detail="Episode not found")
93
 
@@ -99,29 +173,117 @@ async def step_env(episode_id: str, action: Action):
99
  except RuntimeError as e:
100
  raise HTTPException(status_code=400, detail=str(e))
101
 
102
-
103
  @app.get("/result/{episode_id}", response_model=EpisodeResult)
104
- def get_result(episode_id: str):
105
- if episode_id not in episodes:
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
106
  raise HTTPException(status_code=404, detail="Episode not found")
107
- return episodes[episode_id].get_final_result()
108
-
 
 
 
 
 
 
 
 
 
 
 
 
109
 
110
  @app.get("/leaderboard")
111
- def get_leaderboard():
112
- return leaderboard
113
-
 
 
 
 
 
 
 
 
 
 
 
 
 
 
114
 
115
  @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}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
124
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
125
 
126
  @app.websocket("/ws/events")
127
  async def websocket_endpoint(websocket: WebSocket):
@@ -135,7 +297,6 @@ async def websocket_endpoint(websocket: WebSocket):
135
  finally:
136
  clients.discard(websocket)
137
 
138
-
139
  if __name__ == "__main__":
140
  import uvicorn
141
- uvicorn.run(app, host="0.0.0.0", port=7860)
 
1
  import uuid
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
12
  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
17
 
18
  from codereview_env.models import (
19
+ TaskId, Action, ResetResult, StepResult, EpisodeResult, ActionRecord
20
  )
21
  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")
36
 
37
+ # ── App Initialization ────────────────────────────────────────────────────────
38
  app = FastAPI(
39
  title="AgentOrg CodeReview OpenEnv API",
40
  description=(
41
  "AI Senior Code Reviewer evaluation environment. "
42
  "Trains agents to detect bugs, security vulnerabilities, and architectural issues "
43
+ "in realistic Python PRs."
44
  ),
45
  version="1.0.0",
46
  )
47
 
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)
55
 
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")
61
 
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 ────────────────────────────────────────────────────────────────────
88
  class ResetRequest(BaseModel):
89
  task_id: TaskId
90
  seed: int = 42
91
 
 
92
  class ResetResponse(BaseModel):
93
  episode_id: str
94
  result: ResetResult
95
 
 
 
 
 
 
 
 
 
 
96
  class SubmitScore(BaseModel):
97
  agent_name: str
98
  task_id: TaskId
99
  score: float
100
  seed: int
101
 
 
102
  # ── WebSocket clients ─────────��───────────────────────────────────────────────
103
  clients = set()
104
 
 
105
  async def broadcast_event(data: dict):
106
  from fastapi.encoders import jsonable_encoder
 
107
  message = json.dumps(jsonable_encoder(data))
108
  dead = set()
109
  for client in clients:
 
113
  dead.add(client)
114
  clients.difference_update(dead)
115
 
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
+ )
138
 
139
  # ── Endpoints ─────────────────────────────────────────────────────────────────
140
 
141
  @app.get("/health")
142
  def health_check():
143
  return {
144
+ "status": "ok",
145
+ "version": "1.0.0",
146
  "env_ready": True,
147
+ "env": settings.app_env,
148
  "active_episodes": len(episodes),
149
+ "auth_enabled": settings.api_key_enabled
150
  }
151
 
 
152
  @app.post("/reset", response_model=ResetResponse)
153
+ @limiter.limit(f"{settings.rate_limit_per_minute}/minute")
154
+ def reset_env(request: Request, req: ResetRequest, _: None = Depends(verify_api_key)):
155
  episode_id = str(uuid.uuid4())
156
  env = CodeReviewEnv()
157
  result = env.reset(req.task_id, req.seed)
158
  episodes[episode_id] = env
159
+ episode_timestamps[episode_id] = datetime.now(timezone.utc)
160
  return ResetResponse(episode_id=episode_id, result=result)
161
 
 
162
  @app.post("/step/{episode_id}", response_model=StepResult)
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)):
165
  if episode_id not in episodes:
166
  raise HTTPException(status_code=404, detail="Episode not found")
167
 
 
173
  except RuntimeError as e:
174
  raise HTTPException(status_code=400, detail=str(e))
175
 
 
176
  @app.get("/result/{episode_id}", response_model=EpisodeResult)
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:
197
  raise HTTPException(status_code=404, detail="Episode not found")
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
+ )
212
 
213
  @app.get("/leaderboard")
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
231
 
232
  @app.post("/submit")
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
+ }
270
 
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)}
287
 
288
  @app.websocket("/ws/events")
289
  async def websocket_endpoint(websocket: WebSocket):
 
297
  finally:
298
  clients.discard(websocket)
299
 
 
300
  if __name__ == "__main__":
301
  import uvicorn
302
+ uvicorn.run(app, host=settings.app_host, port=settings.app_port)
codereview_env/config.py ADDED
@@ -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()
codereview_env/database.py ADDED
@@ -0,0 +1,142 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from pathlib import Path
2
+ from sqlmodel import SQLModel, Field, Session, create_engine, select
3
+ from typing import Optional, List, Tuple
4
+ import json
5
+ from codereview_env.config import get_settings
6
+ from codereview_env.models import EpisodeResult, TaskId
7
+
8
+ def get_engine():
9
+ settings = get_settings()
10
+ Path(settings.db_path).parent.mkdir(parents=True, exist_ok=True)
11
+ return create_engine(
12
+ f"sqlite:///{settings.db_path}",
13
+ echo=settings.db_echo,
14
+ connect_args={"check_same_thread": False}
15
+ )
16
+
17
+ def create_db_and_tables():
18
+ engine = get_engine()
19
+ SQLModel.metadata.create_all(engine)
20
+
21
+ def get_session():
22
+ engine = get_engine()
23
+ with Session(engine) as session:
24
+ yield session
25
+
26
+ class EpisodeRecord(SQLModel, table=True):
27
+ __tablename__ = "episodes"
28
+
29
+ id: Optional[int] = Field(default=None, primary_key=True)
30
+ episode_id: str = Field(index=True, unique=True)
31
+ task_id: str
32
+ scenario_hash: str
33
+ seed: int
34
+ final_score: float
35
+ steps_taken: int
36
+ issues_found: int
37
+ issues_total: int
38
+ noise_penalties: int
39
+ terminated_reason: str
40
+ history_json: str = "" # JSON-serialized list of ActionRecord dicts
41
+ created_at: str = "" # ISO datetime
42
+ agent_name: str = "" # optional, set via /submit
43
+
44
+ class LeaderboardRecord(SQLModel, table=True):
45
+ __tablename__ = "leaderboard"
46
+
47
+ id: Optional[int] = Field(default=None, primary_key=True)
48
+ agent_name: str
49
+ task_id: str = Field(index=True)
50
+ score: float
51
+ seed: int
52
+ episode_id: str = ""
53
+ submitted_at: str = "" # ISO datetime
54
+
55
+ def save_episode(session: Session, result: EpisodeResult) -> EpisodeRecord:
56
+ """Persist a completed episode result."""
57
+ from datetime import datetime, timezone
58
+ record = EpisodeRecord(
59
+ episode_id=result.episode_id,
60
+ task_id=result.task_id.value,
61
+ scenario_hash=result.scenario_hash,
62
+ seed=result.seed,
63
+ final_score=result.final_score,
64
+ steps_taken=result.steps_taken,
65
+ issues_found=result.issues_found,
66
+ issues_total=result.issues_total,
67
+ noise_penalties=result.noise_penalties,
68
+ terminated_reason=result.terminated_reason,
69
+ history_json=json.dumps([r.model_dump() for r in result.history]),
70
+ created_at=datetime.now(timezone.utc).isoformat()
71
+ )
72
+ session.add(record)
73
+ session.commit()
74
+ session.refresh(record)
75
+ return record
76
+
77
+ def get_episode(session: Session, episode_id: str) -> Optional[EpisodeRecord]:
78
+ return session.exec(select(EpisodeRecord).where(EpisodeRecord.episode_id == episode_id)).first()
79
+
80
+ def get_leaderboard_db(session: Session, task_id: str, limit: int = 10, offset: int = 0) -> Tuple[List[LeaderboardRecord], int]:
81
+ results = session.exec(
82
+ select(LeaderboardRecord)
83
+ .where(LeaderboardRecord.task_id == task_id)
84
+ .order_by(LeaderboardRecord.score.desc())
85
+ .offset(offset)
86
+ .limit(limit)
87
+ ).all()
88
+ # Fixed: session.exec(select(LeaderboardRecord).where(LeaderboardRecord.task_id == task_id)).all() is not efficient but the snippet used it.
89
+ # To get count specifically:
90
+ from sqlmodel import func
91
+ total = session.exec(
92
+ select(func.count()).select_from(LeaderboardRecord).where(LeaderboardRecord.task_id == task_id)
93
+ ).one()
94
+ return list(results), total
95
+
96
+ def submit_leaderboard(session: Session, agent_name: str, task_id: str, score: float, seed: int, episode_id: str = "") -> int:
97
+ """Add entry to leaderboard. Returns rank (1-indexed)."""
98
+ from datetime import datetime, timezone
99
+ record = LeaderboardRecord(
100
+ agent_name=agent_name,
101
+ task_id=task_id,
102
+ score=score,
103
+ seed=seed,
104
+ episode_id=episode_id,
105
+ submitted_at=datetime.now(timezone.utc).isoformat()
106
+ )
107
+ session.add(record)
108
+ session.commit()
109
+ # Calculate rank
110
+ rank_result = session.exec(
111
+ select(LeaderboardRecord)
112
+ .where(LeaderboardRecord.task_id == task_id)
113
+ .order_by(LeaderboardRecord.score.desc())
114
+ ).all()
115
+ for i, r in enumerate(rank_result):
116
+ if r.id == record.id:
117
+ return i + 1
118
+ return -1
119
+
120
+ def get_stats(session: Session) -> dict:
121
+ """Return aggregate statistics about all recorded episodes."""
122
+ all_episodes = session.exec(select(EpisodeRecord)).all()
123
+ if not all_episodes:
124
+ return {"total_episodes": 0, "avg_score": 0.0, "by_task": {}}
125
+
126
+ from collections import defaultdict
127
+ by_task = defaultdict(list)
128
+ for ep in all_episodes:
129
+ by_task[ep.task_id].append(ep.final_score)
130
+
131
+ return {
132
+ "total_episodes": len(all_episodes),
133
+ "avg_score": sum(ep.final_score for ep in all_episodes) / len(all_episodes),
134
+ "by_task": {
135
+ task: {
136
+ "count": len(scores),
137
+ "avg_score": sum(scores) / len(scores) if scores else 0,
138
+ "best_score": max(scores) if scores else 0,
139
+ }
140
+ for task, scores in by_task.items()
141
+ }
142
+ }
requirements.txt CHANGED
@@ -6,3 +6,8 @@ requests>=2.31.0
6
  websockets>=12.0
7
  httpx<0.28.0
8
  openai>=1.0.0
 
 
 
 
 
 
6
  websockets>=12.0
7
  httpx<0.28.0
8
  openai>=1.0.0
9
+ pydantic-settings==2.2.1
10
+ slowapi==0.1.9
11
+ python-dotenv==1.0.1
12
+ sqlmodel==0.0.16
13
+ aiosqlite==0.20.0
scripts/migrate.py ADDED
@@ -0,0 +1,38 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """Initialize or reset the CodeReview database."""
3
+ import sys
4
+ import os
5
+ # Ensure PYTHONPATH is set so we can import codereview_env
6
+ sys.path.append(os.getcwd())
7
+
8
+ from codereview_env.database import create_db_and_tables, get_engine
9
+ from codereview_env.config import get_settings
10
+ from sqlmodel import SQLModel
11
+
12
+ def init():
13
+ settings = get_settings()
14
+ print(f"Initializing database at: {settings.db_path}")
15
+ create_db_and_tables()
16
+ print("Database initialized successfully.")
17
+
18
+ def reset():
19
+ settings = get_settings()
20
+ engine = get_engine()
21
+ print(f"Dropping all tables in: {settings.db_path}")
22
+ SQLModel.metadata.drop_all(engine)
23
+ SQLModel.metadata.create_all(engine)
24
+ print("Database reset successfully.")
25
+
26
+ if __name__ == "__main__":
27
+ if len(sys.argv) < 2:
28
+ print("Usage: python3 scripts/migrate.py [init|reset]")
29
+ sys.exit(1)
30
+
31
+ cmd = sys.argv[1]
32
+ if cmd == "init":
33
+ init()
34
+ elif cmd == "reset":
35
+ reset()
36
+ else:
37
+ print(f"Unknown command: {cmd}. Use 'init' or 'reset'.")
38
+ sys.exit(1)
tests/test_api.py CHANGED
@@ -50,8 +50,10 @@ def test_api_leaderboard():
50
  # Check leaderboard
51
  lb_resp = client.get("/leaderboard")
52
  assert lb_resp.status_code == 200
53
- assert len(lb_resp.json()["bug_detection"]) > 0
54
- assert lb_resp.json()["bug_detection"][0]["agent_name"] == "test_agent"
 
 
55
 
56
  def test_api_invalid_episode():
57
  client = TestClient(app)
 
50
  # Check leaderboard
51
  lb_resp = client.get("/leaderboard")
52
  assert lb_resp.status_code == 200
53
+ lb_data = lb_resp.json()
54
+ bug_entries = lb_data["bug_detection"]["entries"]
55
+ assert len(bug_entries) > 0
56
+ assert bug_entries[0]["agent_name"] == "test_agent"
57
 
58
  def test_api_invalid_episode():
59
  client = TestClient(app)