|
10 | 10 | from pathlib import Path |
11 | 11 | from contextlib import contextmanager |
12 | 12 |
|
13 | | -from fastapi import FastAPI, File, Form, UploadFile, HTTPException, Header, Depends |
14 | | -from fastapi.responses import JSONResponse |
| 13 | +from fastapi import FastAPI, File, Form, UploadFile, HTTPException, Header, Depends, Request |
| 14 | +from fastapi.responses import JSONResponse, HTMLResponse, RedirectResponse |
| 15 | +from fastapi.templating import Jinja2Templates |
15 | 16 | from pydantic import BaseModel |
16 | 17 |
|
17 | 18 | from server.evaluate import evaluate_submission, EvalResult |
18 | 19 |
|
| 20 | +# Templates directory |
| 21 | +TEMPLATES_DIR = Path(__file__).parent / "templates" |
| 22 | +templates = Jinja2Templates(directory=TEMPLATES_DIR) |
| 23 | + |
19 | 24 | app = FastAPI(title="RL Policy Evaluation Server") |
20 | 25 |
|
21 | 26 | # Use DATA_DIR env var for Railway volume, fallback to local directory |
@@ -355,12 +360,12 @@ async def get_leaderboard(env_name: str, limit: int = 50): |
355 | 360 | ] |
356 | 361 |
|
357 | 362 |
|
358 | | -@app.get("/my-submissions") |
| 363 | +@app.get("/api/my-submissions") |
359 | 364 | async def get_my_submissions( |
360 | 365 | student: dict = Depends(verify_student_key), |
361 | 366 | env_name: str | None = None |
362 | 367 | ): |
363 | | - """Get all submissions for the authenticated student.""" |
| 368 | + """Get all submissions for the authenticated student (API endpoint).""" |
364 | 369 | student_id = student["email"] |
365 | 370 |
|
366 | 371 | with get_db() as conn: |
@@ -399,3 +404,272 @@ async def list_environments(): |
399 | 404 | @app.get("/health") |
400 | 405 | async def health(): |
401 | 406 | return {"status": "ok"} |
| 407 | + |
| 408 | + |
| 409 | +# ============ Web Pages ============ |
| 410 | + |
| 411 | +def get_environments_list(): |
| 412 | + """Get list of available environments.""" |
| 413 | + return [ |
| 414 | + {"name": "cartpole", "description": "Classic CartPole balancing task"}, |
| 415 | + ] |
| 416 | + |
| 417 | + |
| 418 | +# Simple flash message support (stored in query params for simplicity) |
| 419 | +def flash_context(): |
| 420 | + """Return empty flash messages context (messages passed via template).""" |
| 421 | + return {"get_flashed_messages": lambda with_categories=False: []} |
| 422 | + |
| 423 | + |
| 424 | +@app.get("/", response_class=HTMLResponse) |
| 425 | +async def web_leaderboard(request: Request, env: str = "cartpole"): |
| 426 | + """Web page showing the leaderboard.""" |
| 427 | + environments = get_environments_list() |
| 428 | + |
| 429 | + with get_db() as conn: |
| 430 | + rows = conn.execute( |
| 431 | + """ |
| 432 | + SELECT student_id, MAX(mean_reward) as mean_reward, std_reward, submitted_at |
| 433 | + FROM submissions |
| 434 | + WHERE env_name = ? |
| 435 | + GROUP BY student_id |
| 436 | + ORDER BY mean_reward DESC |
| 437 | + LIMIT 50 |
| 438 | + """, |
| 439 | + (env,), |
| 440 | + ).fetchall() |
| 441 | + |
| 442 | + entries = [ |
| 443 | + { |
| 444 | + "rank": i + 1, |
| 445 | + "student_id": row["student_id"], |
| 446 | + "mean_reward": row["mean_reward"], |
| 447 | + "std_reward": row["std_reward"], |
| 448 | + "submitted_at": row["submitted_at"], |
| 449 | + } |
| 450 | + for i, row in enumerate(rows) |
| 451 | + ] |
| 452 | + |
| 453 | + return templates.TemplateResponse( |
| 454 | + "leaderboard.html", |
| 455 | + { |
| 456 | + "request": request, |
| 457 | + "entries": entries, |
| 458 | + "environments": environments, |
| 459 | + "current_env": env, |
| 460 | + **flash_context(), |
| 461 | + }, |
| 462 | + ) |
| 463 | + |
| 464 | + |
| 465 | +@app.get("/submit", response_class=HTMLResponse) |
| 466 | +async def web_submit_form(request: Request): |
| 467 | + """Web page for submitting policies.""" |
| 468 | + return templates.TemplateResponse( |
| 469 | + "submit.html", |
| 470 | + { |
| 471 | + "request": request, |
| 472 | + "environments": get_environments_list(), |
| 473 | + "error": None, |
| 474 | + "result": None, |
| 475 | + **flash_context(), |
| 476 | + }, |
| 477 | + ) |
| 478 | + |
| 479 | + |
| 480 | +@app.post("/submit", response_class=HTMLResponse) |
| 481 | +async def web_submit_policy( |
| 482 | + request: Request, |
| 483 | + api_key: str = Form(...), |
| 484 | + env_name: str = Form(default="cartpole"), |
| 485 | + policy_file: UploadFile = File(...), |
| 486 | + checkpoint_file: UploadFile = File(...), |
| 487 | + num_episodes: int = Form(default=100), |
| 488 | +): |
| 489 | + """Handle web form submission.""" |
| 490 | + # Verify API key |
| 491 | + student = get_student_by_api_key(api_key) |
| 492 | + if not student: |
| 493 | + return templates.TemplateResponse( |
| 494 | + "submit.html", |
| 495 | + { |
| 496 | + "request": request, |
| 497 | + "environments": get_environments_list(), |
| 498 | + "error": "Invalid API key", |
| 499 | + "result": None, |
| 500 | + **flash_context(), |
| 501 | + }, |
| 502 | + ) |
| 503 | + |
| 504 | + student_id = student["email"] |
| 505 | + |
| 506 | + # Validate file extensions |
| 507 | + if not policy_file.filename.endswith(".py"): |
| 508 | + return templates.TemplateResponse( |
| 509 | + "submit.html", |
| 510 | + { |
| 511 | + "request": request, |
| 512 | + "environments": get_environments_list(), |
| 513 | + "error": "Policy file must be a .py file", |
| 514 | + "result": None, |
| 515 | + **flash_context(), |
| 516 | + }, |
| 517 | + ) |
| 518 | + |
| 519 | + if not checkpoint_file.filename.endswith(".pt"): |
| 520 | + return templates.TemplateResponse( |
| 521 | + "submit.html", |
| 522 | + { |
| 523 | + "request": request, |
| 524 | + "environments": get_environments_list(), |
| 525 | + "error": "Checkpoint file must be a .pt file", |
| 526 | + "result": None, |
| 527 | + **flash_context(), |
| 528 | + }, |
| 529 | + ) |
| 530 | + |
| 531 | + # Create temporary directory for submission |
| 532 | + with tempfile.TemporaryDirectory() as tmpdir: |
| 533 | + submission_dir = Path(tmpdir) |
| 534 | + |
| 535 | + # Save uploaded files |
| 536 | + policy_path = submission_dir / "policy.py" |
| 537 | + checkpoint_path = submission_dir / "checkpoint.pt" |
| 538 | + |
| 539 | + with open(policy_path, "wb") as f: |
| 540 | + shutil.copyfileobj(policy_file.file, f) |
| 541 | + with open(checkpoint_path, "wb") as f: |
| 542 | + shutil.copyfileobj(checkpoint_file.file, f) |
| 543 | + |
| 544 | + # Run evaluation |
| 545 | + try: |
| 546 | + result = evaluate_submission( |
| 547 | + submission_dir, |
| 548 | + env_name=env_name, |
| 549 | + num_episodes=num_episodes, |
| 550 | + seed=42, |
| 551 | + timeout=60.0, |
| 552 | + ) |
| 553 | + except Exception as e: |
| 554 | + return templates.TemplateResponse( |
| 555 | + "submit.html", |
| 556 | + { |
| 557 | + "request": request, |
| 558 | + "environments": get_environments_list(), |
| 559 | + "error": f"Evaluation failed: {str(e)}", |
| 560 | + "result": None, |
| 561 | + **flash_context(), |
| 562 | + }, |
| 563 | + ) |
| 564 | + |
| 565 | + # Store in database |
| 566 | + submission_id = str(uuid.uuid4())[:8] |
| 567 | + submitted_at = datetime.utcnow().isoformat() |
| 568 | + |
| 569 | + with get_db() as conn: |
| 570 | + conn.execute( |
| 571 | + """ |
| 572 | + INSERT INTO submissions |
| 573 | + (id, student_id, env_name, mean_reward, std_reward, mean_length, episodes, eval_time, submitted_at) |
| 574 | + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) |
| 575 | + """, |
| 576 | + ( |
| 577 | + submission_id, |
| 578 | + student_id, |
| 579 | + env_name, |
| 580 | + result.mean_reward, |
| 581 | + result.std_reward, |
| 582 | + result.mean_length, |
| 583 | + result.episodes, |
| 584 | + result.eval_time, |
| 585 | + submitted_at, |
| 586 | + ), |
| 587 | + ) |
| 588 | + conn.commit() |
| 589 | + |
| 590 | + # Get rank |
| 591 | + rank = conn.execute( |
| 592 | + """ |
| 593 | + SELECT COUNT(*) + 1 FROM ( |
| 594 | + SELECT student_id, MAX(mean_reward) as best |
| 595 | + FROM submissions |
| 596 | + WHERE env_name = ? |
| 597 | + GROUP BY student_id |
| 598 | + HAVING best > ? |
| 599 | + ) |
| 600 | + """, |
| 601 | + (env_name, result.mean_reward), |
| 602 | + ).fetchone()[0] |
| 603 | + |
| 604 | + return templates.TemplateResponse( |
| 605 | + "submit.html", |
| 606 | + { |
| 607 | + "request": request, |
| 608 | + "environments": get_environments_list(), |
| 609 | + "error": None, |
| 610 | + "result": { |
| 611 | + "mean_reward": result.mean_reward, |
| 612 | + "std_reward": result.std_reward, |
| 613 | + "episodes": result.episodes, |
| 614 | + "rank": rank, |
| 615 | + }, |
| 616 | + **flash_context(), |
| 617 | + }, |
| 618 | + ) |
| 619 | + |
| 620 | + |
| 621 | +@app.get("/my-submissions", response_class=HTMLResponse) |
| 622 | +async def web_my_submissions(request: Request, api_key: str | None = None): |
| 623 | + """Web page showing user's submissions.""" |
| 624 | + if not api_key: |
| 625 | + return templates.TemplateResponse( |
| 626 | + "my_submissions.html", |
| 627 | + { |
| 628 | + "request": request, |
| 629 | + "api_key": None, |
| 630 | + "submissions": [], |
| 631 | + "student_email": None, |
| 632 | + "error": None, |
| 633 | + **flash_context(), |
| 634 | + }, |
| 635 | + ) |
| 636 | + |
| 637 | + student = get_student_by_api_key(api_key) |
| 638 | + if not student: |
| 639 | + return templates.TemplateResponse( |
| 640 | + "my_submissions.html", |
| 641 | + { |
| 642 | + "request": request, |
| 643 | + "api_key": None, |
| 644 | + "submissions": [], |
| 645 | + "student_email": None, |
| 646 | + "error": "Invalid API key", |
| 647 | + **flash_context(), |
| 648 | + }, |
| 649 | + ) |
| 650 | + |
| 651 | + student_id = student["email"] |
| 652 | + |
| 653 | + with get_db() as conn: |
| 654 | + rows = conn.execute( |
| 655 | + """ |
| 656 | + SELECT * FROM submissions |
| 657 | + WHERE student_id = ? |
| 658 | + ORDER BY submitted_at DESC |
| 659 | + """, |
| 660 | + (student_id,), |
| 661 | + ).fetchall() |
| 662 | + |
| 663 | + submissions = [dict(row) for row in rows] |
| 664 | + |
| 665 | + return templates.TemplateResponse( |
| 666 | + "my_submissions.html", |
| 667 | + { |
| 668 | + "request": request, |
| 669 | + "api_key": api_key, |
| 670 | + "submissions": submissions, |
| 671 | + "student_email": student_id, |
| 672 | + "error": None, |
| 673 | + **flash_context(), |
| 674 | + }, |
| 675 | + ) |
0 commit comments