Spaces:
Sleeping
Sleeping
| """Repair-aware randomized human annotation application.""" | |
| from __future__ import annotations | |
| import hmac | |
| import json | |
| import mimetypes | |
| import os | |
| from contextlib import asynccontextmanager | |
| from pathlib import Path | |
| from fastapi import Depends, FastAPI, Form, Request | |
| from fastapi.responses import FileResponse, HTMLResponse, JSONResponse, RedirectResponse | |
| from fastapi.staticfiles import StaticFiles | |
| from fastapi.templating import Jinja2Templates | |
| from sqlalchemy import func, select | |
| from . import db, study | |
| from .auth import ( | |
| COOKIE_NAME, | |
| COOKIE_SECURE, | |
| create_session_cookie, | |
| current_email, | |
| get_csrf_token, | |
| set_csrf_cookie, | |
| verify_csrf, | |
| ) | |
| mimetypes.add_type("audio/mpeg", ".mp3") | |
| APP_DIR = Path(__file__).resolve().parent | |
| PACKAGE_DIR = (APP_DIR / ".." / "data").resolve() | |
| STUDY_MANIFEST = PACKAGE_DIR / "study_manifest.json" | |
| STUDY_ASSIGNMENTS = PACKAGE_DIR / "study_assignments.json" | |
| MEDIA_DIR = Path(os.environ.get("ANNOTATOR_MEDIA_DIR", "/data/study_clips")) | |
| STUDY_PHASE = os.environ.get("STUDY_PHASE", "development").strip().lower() | |
| REQUIRE_ACCESS_CODE = os.environ.get("REQUIRE_ACCESS_CODE", "true").strip().lower() != "false" | |
| STUDY_OPEN = os.environ.get("STUDY_OPEN", "false").strip().lower() == "true" | |
| async def lifespan(_app: FastAPI): | |
| db.init_db() | |
| session = db.get_session() | |
| try: | |
| result = db.seed_study(session, STUDY_MANIFEST, STUDY_ASSIGNMENTS) | |
| print(f"[startup] study package: {result}") | |
| finally: | |
| session.close() | |
| MEDIA_DIR.mkdir(parents=True, exist_ok=True) | |
| yield | |
| app = FastAPI(title="Listening and Responsiveness Study", lifespan=lifespan) | |
| app.mount("/static", StaticFiles(directory=str(APP_DIR / "static")), name="static") | |
| templates = Jinja2Templates(directory=str(APP_DIR / "templates")) | |
| def _get_db(): | |
| session = db.get_session() | |
| try: | |
| yield session | |
| finally: | |
| session.close() | |
| def _valid_access_code(candidate: str) -> bool: | |
| expected = os.environ.get("ANNOTATOR_ACCESS_CODE", "") | |
| if not REQUIRE_ACCESS_CODE: | |
| return True | |
| return bool(expected) and hmac.compare_digest(candidate.strip(), expected) | |
| def _admin_ok(token: str) -> bool: | |
| expected = os.environ.get("ANNOTATOR_ADMIN_TOKEN", "") | |
| return bool(expected) and hmac.compare_digest(token, expected) | |
| def index(request: Request, session=Depends(_get_db)): | |
| identity = current_email(request) | |
| if identity: | |
| annotator = db.get_annotator(session, identity) | |
| if annotator and annotator.consented_at: | |
| destination = "/label" if annotator.training_passed_at else "/instructions" | |
| return RedirectResponse(destination, status_code=303) | |
| csrf_token = get_csrf_token(request) | |
| response = templates.TemplateResponse( | |
| request, | |
| "login.html", | |
| { | |
| "csrf_token": csrf_token, | |
| "access_configured": bool(os.environ.get("ANNOTATOR_ACCESS_CODE")) or not REQUIRE_ACCESS_CODE, | |
| }, | |
| ) | |
| set_csrf_cookie(response, csrf_token) | |
| return response | |
| def login( | |
| request: Request, | |
| participant_id: str = Form(...), | |
| access_code: str = Form(""), | |
| finance_familiarity: str = Form(...), | |
| english_proficiency: str = Form(...), | |
| consent: str | None = Form(None), | |
| csrf_token: str = Form(...), | |
| session=Depends(_get_db), | |
| ): | |
| if not verify_csrf(request, csrf_token): | |
| return RedirectResponse("/?error=invalid_csrf", status_code=303) | |
| identity = participant_id.strip().lower() | |
| if not (3 <= len(identity) <= 80) or any(character.isspace() for character in identity): | |
| return RedirectResponse("/?error=invalid_id", status_code=303) | |
| if not _valid_access_code(access_code): | |
| return RedirectResponse("/?error=invalid_access", status_code=303) | |
| if finance_familiarity not in {"none", "some", "professional"}: | |
| return RedirectResponse("/?error=profile", status_code=303) | |
| if english_proficiency not in {"fluent", "native"} or consent != "yes": | |
| return RedirectResponse("/?error=consent", status_code=303) | |
| db.enroll_annotator(session, identity, finance_familiarity, english_proficiency) | |
| response = RedirectResponse("/instructions", status_code=303) | |
| response.set_cookie( | |
| COOKIE_NAME, | |
| create_session_cookie(identity), | |
| max_age=60 * 60 * 24 * 365, | |
| httponly=True, | |
| samesite="lax", | |
| secure=COOKIE_SECURE, | |
| ) | |
| return response | |
| def instructions(request: Request, session=Depends(_get_db)): | |
| identity = current_email(request) | |
| annotator = db.get_annotator(session, identity) if identity else None | |
| if annotator is None or not annotator.consented_at: | |
| return RedirectResponse("/", status_code=303) | |
| csrf_token = get_csrf_token(request) | |
| response = templates.TemplateResponse( | |
| request, | |
| "instructions.html", | |
| { | |
| "participant_id": identity, | |
| "csrf_token": csrf_token, | |
| "passed": bool(annotator.training_passed_at), | |
| "error": request.query_params.get("error"), | |
| }, | |
| ) | |
| set_csrf_cookie(response, csrf_token) | |
| return response | |
| def complete_instructions( | |
| request: Request, | |
| practice_1: str = Form(""), | |
| practice_2: str = Form(""), | |
| practice_3: str = Form(""), | |
| csrf_token: str = Form(...), | |
| session=Depends(_get_db), | |
| ): | |
| identity = current_email(request) | |
| annotator = db.get_annotator(session, identity) if identity else None | |
| if annotator is None or not verify_csrf(request, csrf_token): | |
| return RedirectResponse("/", status_code=303) | |
| passed = ( | |
| practice_1 == "substantive_answer_attempt" | |
| and practice_2 == "clarification_repair" | |
| and practice_3 == "explicit_disclosure_boundary" | |
| ) | |
| db.record_training_attempt(session, annotator, passed) | |
| if not passed: | |
| return RedirectResponse("/instructions?error=practice", status_code=303) | |
| return RedirectResponse("/label", status_code=303) | |
| def logout(request: Request, csrf_token: str = Form(...)): | |
| if not verify_csrf(request, csrf_token): | |
| return RedirectResponse("/", status_code=303) | |
| response = RedirectResponse("/", status_code=303) | |
| response.delete_cookie(COOKIE_NAME) | |
| return response | |
| def label_view(request: Request, session=Depends(_get_db)): | |
| identity = current_email(request) | |
| if not identity: | |
| return RedirectResponse("/", status_code=303) | |
| annotator = db.get_annotator(session, identity) | |
| if annotator is None or not annotator.consented_at: | |
| return RedirectResponse("/", status_code=303) | |
| if not annotator.training_passed_at: | |
| return RedirectResponse("/instructions", status_code=303) | |
| if not STUDY_OPEN: | |
| progress = db.study_progress(session, annotator.id, STUDY_PHASE) | |
| csrf_token = get_csrf_token(request) | |
| response = templates.TemplateResponse( | |
| request, | |
| "done.html", | |
| { | |
| "participant_id": identity, | |
| "progress": progress, | |
| "csrf_token": csrf_token, | |
| "phase": "prelaunch — collection is paused", | |
| }, | |
| ) | |
| set_csrf_cookie(response, csrf_token) | |
| return response | |
| claimed = db.claim_task(session, annotator, STUDY_PHASE) | |
| progress = db.study_progress(session, annotator.id, STUDY_PHASE) | |
| csrf_token = get_csrf_token(request) | |
| if claimed is None: | |
| response = templates.TemplateResponse( | |
| request, | |
| "done.html", | |
| { | |
| "participant_id": identity, | |
| "progress": progress, | |
| "csrf_token": csrf_token, | |
| "phase": STUDY_PHASE, | |
| }, | |
| ) | |
| else: | |
| task, item = claimed | |
| media_exists = bool(item.audio_filename and (MEDIA_DIR / item.audio_filename).exists()) | |
| response = templates.TemplateResponse( | |
| request, | |
| "label.html", | |
| { | |
| "participant_id": identity, | |
| "task": task, | |
| "item": item, | |
| "progress": progress, | |
| "csrf_token": csrf_token, | |
| "gate_options": study.GATE_OPTIONS, | |
| "rasiah_options": study.RASIAH_OPTIONS, | |
| "supplied_options": study.SUPPLIED_OPTIONS, | |
| "descriptors": study.DESCRIPTORS, | |
| "media_exists": media_exists, | |
| "phase": STUDY_PHASE, | |
| "error": request.query_params.get("error"), | |
| }, | |
| ) | |
| set_csrf_cookie(response, csrf_token) | |
| return response | |
| def protected_media(item_id: str, request: Request, session=Depends(_get_db)): | |
| """Serve audio only to the participant currently assigned this item.""" | |
| identity = current_email(request) | |
| annotator = db.get_annotator(session, identity) if identity else None | |
| if annotator is None or not annotator.current_task_id: | |
| return JSONResponse({"error": "not found"}, status_code=404) | |
| task = session.get(db.StudyTask, annotator.current_task_id) | |
| item = session.get(db.StudyItem, item_id) | |
| if ( | |
| task is None | |
| or item is None | |
| or task.item_id != item_id | |
| or task.condition != "text_audio" | |
| or task.assigned_annotator_id != annotator.id | |
| or not item.audio_filename | |
| ): | |
| return JSONResponse({"error": "not found"}, status_code=404) | |
| path = MEDIA_DIR / item.audio_filename | |
| if not path.exists() or path.parent.resolve() != MEDIA_DIR.resolve(): | |
| return JSONResponse({"error": "not found"}, status_code=404) | |
| return FileResponse(path, media_type="audio/mpeg", filename=None) | |
| async def submit(request: Request, session=Depends(_get_db)): | |
| identity = current_email(request) | |
| annotator = db.get_annotator(session, identity) if identity else None | |
| if annotator is None: | |
| return RedirectResponse("/", status_code=303) | |
| form_data = await request.form() | |
| form = dict(form_data) | |
| if not verify_csrf(request, str(form.get("csrf_token") or "")): | |
| return RedirectResponse("/label?error=session", status_code=303) | |
| try: | |
| task_id = int(str(form.get("task_id") or "")) | |
| except ValueError: | |
| return RedirectResponse("/label?error=assignment", status_code=303) | |
| task = session.get(db.StudyTask, task_id) | |
| if task is None or task.assigned_annotator_id != annotator.id: | |
| return RedirectResponse("/label?error=assignment", status_code=303) | |
| action = str(form.get("action") or "label") | |
| if action == "flag": | |
| reason = str(form.get("flag_reason") or "").strip() | |
| if len(reason) < 4: | |
| return RedirectResponse("/label?error=flag_reason", status_code=303) | |
| db.flag_task(session, annotator, task_id, reason) | |
| return RedirectResponse("/label", status_code=303) | |
| item = session.get(db.StudyItem, task.item_id) | |
| clean, error = study.validate_submission( | |
| form, | |
| task.condition, | |
| audio_duration_s=item.duration_s if item is not None else None, | |
| ) | |
| if error: | |
| return RedirectResponse(f"/label?error={error}", status_code=303) | |
| if not db.complete_task(session, annotator, task_id, clean or {}): | |
| return RedirectResponse("/label?error=assignment", status_code=303) | |
| return RedirectResponse("/label", status_code=303) | |
| def export(token: str = "", session=Depends(_get_db)): | |
| if not _admin_ok(token): | |
| return JSONResponse({"error": "not found"}, status_code=404) | |
| return JSONResponse(db.export_payload(session)) | |
| def admin_status(token: str = "", session=Depends(_get_db)): | |
| if not _admin_ok(token): | |
| return JSONResponse({"error": "not found"}, status_code=404) | |
| by_condition = dict( | |
| session.execute( | |
| select(db.StudyTask.condition, func.count(db.StudyTask.id)) | |
| .where(db.StudyTask.completed.is_(True)) | |
| .group_by(db.StudyTask.condition) | |
| ).all() | |
| ) | |
| return { | |
| "phase": STUDY_PHASE, | |
| "items": session.scalar(select(func.count(db.StudyItem.id))) or 0, | |
| "tasks": session.scalar(select(func.count(db.StudyTask.id))) or 0, | |
| "completed": session.scalar(select(func.count(db.StudyTask.id)).where(db.StudyTask.completed.is_(True))) or 0, | |
| "annotators": session.scalar(select(func.count(db.Annotator.id))) or 0, | |
| "media_files": len(list(MEDIA_DIR.glob("*.mp3"))), | |
| "completed_by_condition": by_condition, | |
| } | |
| async def upload_media(item_id: str, request: Request, token: str = "", session=Depends(_get_db)): | |
| """Upload one authorized MP3 to private persistent storage.""" | |
| if not _admin_ok(token): | |
| return JSONResponse({"error": "not found"}, status_code=404) | |
| item = session.get(db.StudyItem, item_id) | |
| if item is None or not item.audio_filename: | |
| return JSONResponse({"error": "unknown item"}, status_code=404) | |
| body = await request.body() | |
| if not body or len(body) > 2_000_000: | |
| return JSONResponse({"error": "invalid media size"}, status_code=400) | |
| target = MEDIA_DIR / item.audio_filename | |
| target.write_bytes(body) | |
| return {"ok": True, "bytes": len(body)} | |
| def healthz(): | |
| return { | |
| "status": "ok", | |
| "study_id": study.STUDY_ID, | |
| "phase": STUDY_PHASE, | |
| "db_path": str(Path(db.DB_PATH).parent), | |
| "access_code_configured": bool(os.environ.get("ANNOTATOR_ACCESS_CODE")) or not REQUIRE_ACCESS_CODE, | |
| "study_open": STUDY_OPEN, | |
| } | |