pdf-quiz-generator/backend/app/tasks/quiz_tasks.py
Daniel 2cbbfe00c3 Tag filtering, multi-category, bug fixes, image validation, docs
- Fix tag filtering (sa_text import shadowing caused UnboundLocalError)
- Add TagBrowser component with per-section search
- Multi-category selection (OR within categories, AND with tags)
- AI image validation: has_figure field in extraction prompt
- Skip known branding images by MD5 hash + dimension filters
- Fix quiz timer auto-submit (wrong useEffect dependency)
- Fix QuizResponse schema: section_id nullable
- Fix Question.quiz_id → source_quiz_id attribute name
- Fix SQL injection in quizzes.py vector search
- Add PDF processing progress steps via Redis
- Add delete user from admin panel
- Admin page: no spinner flash on data refresh
- Upload progress: axios 1.x e.progress, remove manual Content-Type
- Duplicate model error: 409 with clear message
- Backend startup: retry DDL migration on lock timeout
- Replace all silent except:pass with warning logs
- Comprehensive multi-page documentation (docs/)

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-04-04 22:48:26 +02:00

499 lines
22 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""Async quiz extraction task with step-by-step progress reporting via Redis."""
import json
import logging
import time
import os
from app.tasks import celery_app
from app.database import SessionLocal
logger = logging.getLogger(__name__)
EXPIRE_SECONDS = 3600
CHUNK_PAGES = 50
def _redis():
import redis
from app.config import settings
return redis.from_url(settings.REDIS_URL, decode_responses=True)
def _push_step(r, job_id: str, step: str, message: str):
key = f"extraction:steps:{job_id}"
entry = json.dumps({"step": step, "message": message, "ts": time.time()})
r.rpush(key, entry)
r.expire(key, EXPIRE_SECONDS)
def _normalize_ocr(text: str) -> str:
"""Fix common OCR artifacts in PREP PDFs."""
return (text
.replace("Pref erred", "Preferred")
.replace("Pre ferred", "Preferred")
.replace("Prefer red", "Preferred")
.replace("ltem", "Item")
.replace("ltcm", "Item"))
@celery_app.task(name="extract_quiz", bind=True)
def extract_quiz(
self,
job_id: str,
user_id: int,
section_id: int,
title: str,
mode: str,
time_limit_minutes: int | None,
model_id: str | None,
question_category_id: int | None,
extraction_mode: str = "standard",
):
r = _redis()
r.set(f"extraction:status:{job_id}", "running", ex=EXPIRE_SECONDS)
db = SessionLocal()
try:
from app.models.section import Section
from app.models.pdf_document import PDFDocument
from app.services import ai_service, vector_service, pdf_service, embedding_service
from app.models.quiz import Quiz
from app.models.question import Question
from app.models.quiz_question_link import QuizQuestionLink
from app.config import settings
_push_step(r, job_id, "start", "Starting extraction…")
section = db.query(Section).filter(Section.id == section_id).first()
if not section:
raise ValueError("Section not found")
document = db.query(PDFDocument).filter(PDFDocument.id == section.document_id).first()
if not document:
raise ValueError("Document not found")
total_pages = section.end_page - section.start_page + 1
_push_step(r, job_id, "text", f"Loading text from pages {section.start_page}{section.end_page} ({total_pages} pages)…")
# Determine model
if model_id:
from app.models.ai_model_config import AIModelConfig
config = db.query(AIModelConfig).filter(AIModelConfig.model_id == model_id).first()
api_key = config.api_key if config and config.api_key else None
model_name = config.name if config else model_id
else:
model_id, api_key = ai_service.get_model_for_task(db, "extraction")
model_name = model_id
# Split into 50-page chunks
if total_pages <= CHUNK_PAGES:
chunks = [(section.start_page, section.end_page)]
else:
chunks = []
p = section.start_page
while p <= section.end_page:
end = min(p + CHUNK_PAGES - 1, section.end_page)
chunks.append((p, end))
p = end + 1
n_chunks = len(chunks)
if n_chunks > 1:
_push_step(r, job_id, "text", f"Large section: splitting into {n_chunks} chunks of up to {CHUNK_PAGES} pages each.")
all_valid_questions = []
all_skipped = []
# ── Non-standard extraction modes ─────────────────────────────────────
if extraction_mode != "standard":
from app.services.extraction_modes import (
extract_questions_only, extract_two_step,
extract_with_regex, ai_decide_strategy, generate_from_text,
)
resolved_mode = extraction_mode
if extraction_mode == "ai_decide":
_push_step(r, job_id, "ai", "AI is analysing the document to choose the best strategy…")
resolved_mode, reasoning = ai_decide_strategy(
section.document_id, section.start_page, section.end_page,
model_id, api_key,
)
_push_step(r, job_id, "ai", f"AI chose: {resolved_mode}{reasoning}")
if resolved_mode == "questions_only":
_push_step(r, job_id, "ai", "Mode: Questions Only — extracting questions without answers.")
for chunk_idx, (start_p, end_p) in enumerate(chunks, 1):
if r.get(f"extraction:status:{job_id}") == "cancelled":
_push_step(r, job_id, "cancelled", "Job cancelled.")
return
_push_step(r, job_id, "ai", f"Chunk {chunk_idx}/{n_chunks}: pages {start_p}{end_p}")
chunk_content = vector_service.get_pages_text(
document_id=section.document_id, start_page=start_p, end_page=end_p)
if not chunk_content:
continue
try:
qs = extract_questions_only(_normalize_ocr(chunk_content),
f"{start_p}-{end_p}", start_p, model_id, api_key)
all_valid_questions.extend(qs)
_push_step(r, job_id, "ai", f" Pages {start_p}{end_p}: {len(qs)} questions. Total: {len(all_valid_questions)}.")
except Exception as e:
_push_step(r, job_id, "ai", f" Pages {start_p}{end_p} failed: {e}")
elif resolved_mode == "two_step":
_push_step(r, job_id, "ai", "Mode: Two-Step (separate answer key section).")
all_valid_questions, all_skipped = extract_two_step(
section.document_id, section.start_page, section.end_page,
model_id, api_key,
push_step=lambda step, msg: _push_step(r, job_id, step, msg),
chunk_pages=CHUNK_PAGES,
)
elif resolved_mode == "regex":
_push_step(r, job_id, "ai", "Mode: AI+Regex — analysing format then applying regex.")
all_valid_questions, all_skipped = extract_with_regex(
section.document_id, section.start_page, section.end_page,
model_id, api_key,
push_step=lambda step, msg: _push_step(r, job_id, step, msg),
chunk_pages=CHUNK_PAGES,
)
elif resolved_mode == "generate":
_push_step(r, job_id, "ai", "Mode: Generate — AI creates questions from the text.")
for chunk_idx, (start_p, end_p) in enumerate(chunks, 1):
if r.get(f"extraction:status:{job_id}") == "cancelled":
_push_step(r, job_id, "cancelled", "Job cancelled.")
return
_push_step(r, job_id, "ai", f"Chunk {chunk_idx}/{n_chunks}: pages {start_p}{end_p}")
chunk_content = vector_service.get_pages_text(
document_id=section.document_id, start_page=start_p, end_page=end_p)
if not chunk_content:
continue
try:
qs = generate_from_text(_normalize_ocr(chunk_content),
f"{start_p}-{end_p}", start_p, model_id, api_key)
all_valid_questions.extend(qs)
_push_step(r, job_id, "ai", f" Pages {start_p}{end_p}: {len(qs)} questions generated. Total: {len(all_valid_questions)}.")
except Exception as e:
_push_step(r, job_id, "ai", f" Pages {start_p}{end_p} failed: {e}")
else:
# ai_decide resolved to standard — fall through to standard loop below
extraction_mode = "standard"
if extraction_mode == "standard":
# ── STANDARD: existing working extraction (unchanged) ──────────────
for chunk_idx, (start_p, end_p) in enumerate(chunks, 1):
if r.get(f"extraction:status:{job_id}") == "cancelled":
_push_step(r, job_id, "cancelled", "Job cancelled.")
return
if n_chunks > 1:
_push_step(r, job_id, "ai", f"Chunk {chunk_idx}/{n_chunks}: pages {start_p}{end_p}{model_name}")
else:
_push_step(r, job_id, "ai", f"Sending pages {start_p}{end_p} to {model_name}")
chunk_content = vector_service.get_pages_text(
document_id=section.document_id, start_page=start_p, end_page=end_p,
)
if not chunk_content:
_push_step(r, job_id, "ai", f" No text found for pages {start_p}{end_p}, skipping.")
continue
try:
chunk_data = ai_service.extract_questions(
_normalize_ocr(chunk_content),
page_info=f"{start_p}-{end_p}",
page_ref=start_p,
model_id=model_id,
api_key=api_key,
)
chunk_skipped = chunk_data[0].pop("skipped", []) if chunk_data else []
chunk_valid = [q for q in chunk_data if q.get("correct_answer")]
all_valid_questions.extend(chunk_valid)
all_skipped.extend(chunk_skipped)
_push_step(r, job_id, "ai",
f" Pages {start_p}{end_p}: {len(chunk_valid)} questions"
f"{f', {len(chunk_skipped)} skipped' if chunk_skipped else ''}. "
f"Total: {len(all_valid_questions)}.")
except Exception as e:
_push_step(r, job_id, "ai", f" Pages {start_p}{end_p} failed: {e}. Continuing…")
valid_questions = all_valid_questions
skipped = all_skipped
_push_step(r, job_id, "ai", f"Extraction complete: {len(valid_questions)} valid questions{f', {len(skipped)} skipped' if skipped else ''}.")
if not valid_questions:
raise ValueError("No valid questions extracted. The AI could not find questions with correct answers in this page range.")
# Refresh DB connection — it may have gone stale during long LLM extraction
from sqlalchemy import text as _text
try:
db.execute(_text("SELECT 1"))
except Exception:
db.rollback()
db.close()
db = SessionLocal()
# Re-fetch objects that were bound to the old session
section = db.query(Section).filter(Section.id == section_id).first()
document = db.query(PDFDocument).filter(PDFDocument.id == section.document_id).first()
# Extract images
_push_step(r, job_id, "images", "Extracting question images…")
file_path = os.path.join(settings.UPLOAD_DIR, document.filename)
page_images = {}
if os.path.exists(file_path):
try:
page_images = pdf_service.extract_all_images(
file_path, document.id, section.start_page, section.end_page
)
except Exception as e:
_push_step(r, job_id, "images", f"Image extraction skipped: {e}")
# Create quiz
quiz = Quiz(
section_id=section_id,
user_id=user_id,
title=title,
questions_count=len(valid_questions),
mode=mode,
time_limit_minutes=time_limit_minutes,
skipped_questions=json.dumps(skipped) if skipped else None,
)
db.add(quiz)
db.flush()
_push_step(r, job_id, "save", f"Saving {len(valid_questions)} questions and generating embeddings…")
for pos, q in enumerate(valid_questions):
page_ref = q.get("page_reference")
image_path = None
# Only link an image if the AI flagged the question as having a figure
if q.get("has_figure") and page_ref and page_ref in page_images and page_images[page_ref]:
image_path = page_images[page_ref].pop(0)
if not page_images[page_ref]:
del page_images[page_ref]
question = Question(
source_quiz_id=quiz.id,
question_category_id=question_category_id,
question_text=q["question_text"],
question_type=q["question_type"],
options=q.get("options"),
correct_answer=q["correct_answer"],
explanation=q.get("explanation", ""),
page_reference=page_ref,
image_path=image_path,
)
db.add(question)
db.flush()
db.add(QuizQuestionLink(quiz_id=quiz.id, question_id=question.id, position=pos))
try:
embedding_service.embed_question(question)
except Exception as e:
logger.warning(f"Embedding failed for question {question.id}: {e}")
db.commit()
db.refresh(quiz)
_push_step(r, job_id, "done", f"Quiz ready! {len(valid_questions)} questions extracted and saved.")
r.set(f"extraction:status:{job_id}", "completed", ex=EXPIRE_SECONDS)
r.set(f"extraction:quiz_id:{job_id}", str(quiz.id), ex=EXPIRE_SECONDS)
return quiz.id
except Exception as e:
logger.exception(f"Quiz extraction failed for job {job_id}")
_push_step(r, job_id, "error", f"Extraction failed: {e}")
r.set(f"extraction:status:{job_id}", "failed", ex=EXPIRE_SECONDS)
r.set(f"extraction:error:{job_id}", str(e)[:500], ex=EXPIRE_SECONDS)
raise
finally:
db.close()
CLASSIFY_EXPIRE = 3600
def _push_classify_step(r, job_id: str, step: str, message: str):
key = f"classify:steps:{job_id}"
entry = json.dumps({"step": step, "message": message, "ts": time.time()})
r.rpush(key, entry)
r.expire(key, CLASSIFY_EXPIRE)
@celery_app.task(name="classify_questions", bind=True)
def classify_questions(self, job_id: str, user_id: int):
"""Classify untagged questions using AI — subjects, diseases, keywords."""
r = _redis()
r.set(f"classify:status:{job_id}", "running", ex=CLASSIFY_EXPIRE)
db = SessionLocal()
try:
from app.models.question import Question
from app.services import ai_service
from sqlalchemy import text as sa_text
_push_classify_step(r, job_id, "start", "Finding untagged questions...")
# Get IDs of questions that already have tags
tagged_ids_rows = db.execute(sa_text(
"SELECT DISTINCT question_id FROM question_tag_links"
)).fetchall()
tagged_ids = {row[0] for row in tagged_ids_rows}
# Get all questions not yet tagged
all_questions = db.query(Question).all()
untagged = [q for q in all_questions if q.id not in tagged_ids]
if not untagged:
_push_classify_step(r, job_id, "done", "All questions are already tagged.")
r.set(f"classify:status:{job_id}", "completed", ex=CLASSIFY_EXPIRE)
return
total = len(untagged)
_push_classify_step(r, job_id, "start", f"Found {total} untagged questions. Starting classification...")
# Get AI model for keyword task
model_id, api_key = ai_service.get_model_for_task(db, "keyword")
batch_size = 10
classified = 0
for i in range(0, total, batch_size):
if r.get(f"classify:status:{job_id}") == "cancelled":
_push_classify_step(r, job_id, "cancelled", "Job cancelled.")
return
batch = untagged[i:i + batch_size]
batch_num = (i // batch_size) + 1
total_batches = (total + batch_size - 1) // batch_size
_push_classify_step(r, job_id, "progress",
f"Batch {batch_num}/{total_batches}: classifying {len(batch)} questions...")
# Build questions JSON for prompt
questions_json = json.dumps([
{"id": q.id, "question_text": q.question_text[:500]}
for q in batch
], indent=2)
prompt = f"""Classify each medical question below. For each question, provide:
- subjects: 1-3 medical subjects/specialties (e.g., "Cardiology", "Infectious Disease", "Neonatology")
- diseases: 1-3 specific diseases/conditions mentioned (e.g., "Kawasaki Disease", "Pneumonia", "Type 1 Diabetes")
- keywords: 2-4 key clinical concepts (e.g., "fever workup", "antibiotic resistance", "fluid management")
Return ONLY JSON:
{{"classifications": [
{{"id": <question_id>, "subjects": [...], "diseases": [...], "keywords": [...]}}
]}}
Questions:
{questions_json}"""
try:
raw = ai_service._call_model(prompt, model_id, api_key)
# Parse JSON from response
text = raw.strip()
if text.startswith("```"):
text = text.split("\n", 1)[1] if "\n" in text else text[3:]
if text.endswith("```"):
text = text[:-3]
text = text.strip()
data = json.loads(text)
classifications = data.get("classifications", [])
# Store tags
for cls in classifications:
q_id = cls.get("id")
if not q_id:
continue
for tag_type, tag_list in [("subject", cls.get("subjects", [])),
("disease", cls.get("diseases", [])),
("keyword", cls.get("keywords", []))]:
for tag_name in tag_list:
if not tag_name or not isinstance(tag_name, str):
continue
normalized = tag_name.strip().title()
if not normalized:
continue
# Insert tag (ON CONFLICT DO NOTHING for case-insensitive uniqueness)
db.execute(sa_text("""
INSERT INTO question_tags (name, type)
VALUES (:name, :type)
ON CONFLICT (LOWER(name), type) DO NOTHING
"""), {"name": normalized, "type": tag_type})
db.flush()
# Get the tag ID
tag_row = db.execute(sa_text("""
SELECT id FROM question_tags
WHERE LOWER(name) = LOWER(:name) AND type = :type
"""), {"name": normalized, "type": tag_type}).fetchone()
if tag_row:
db.execute(sa_text("""
INSERT INTO question_tag_links (question_id, tag_id)
VALUES (:qid, :tid)
ON CONFLICT DO NOTHING
"""), {"qid": q_id, "tid": tag_row[0]})
db.commit()
classified += len(batch)
_push_classify_step(r, job_id, "progress",
f"Batch {batch_num}/{total_batches} done. {classified}/{total} classified.")
except Exception as e:
logger.warning(f"Classification batch {batch_num} failed: {e}")
_push_classify_step(r, job_id, "progress",
f"Batch {batch_num} failed: {e}. Continuing...")
continue
_push_classify_step(r, job_id, "done", f"Classification complete. {classified}/{total} questions classified.")
r.set(f"classify:status:{job_id}", "completed", ex=CLASSIFY_EXPIRE)
except Exception as e:
logger.exception(f"Classification failed for job {job_id}")
_push_classify_step(r, job_id, "error", f"Failed: {e}")
r.set(f"classify:status:{job_id}", "failed", ex=CLASSIFY_EXPIRE)
r.set(f"classify:error:{job_id}", str(e)[:500], ex=CLASSIFY_EXPIRE)
raise
finally:
db.close()
@celery_app.task(name="regenerate_embeddings", bind=True)
def regenerate_embeddings(self, job_id: str, user_id: int):
"""Regenerate embeddings for all questions using the current embedding model."""
r = _redis()
r.set(f"extraction:status:{job_id}", "running", ex=EXPIRE_SECONDS)
r.set(f"extraction:job_title:{job_id}", "Regenerate Embeddings", ex=EXPIRE_SECONDS)
r.lpush(f"extraction:user_jobs:{user_id}", job_id)
r.expire(f"extraction:user_jobs:{user_id}", 86400)
db = SessionLocal()
try:
from app.models.question import Question
from app.services import embedding_service
questions = db.query(Question).all()
total = len(questions)
_push_step(r, job_id, "start", f"Regenerating embeddings for {total} questions…")
ok = 0
for i, q in enumerate(questions):
try:
if embedding_service.embed_question(q):
ok += 1
if (i + 1) % 50 == 0:
db.commit()
_push_step(r, job_id, "progress", f"{i + 1}/{total} processed ({ok} embedded)")
except Exception as e:
logger.warning(f"Embedding failed for question {q.id}: {e}")
db.commit()
_push_step(r, job_id, "done", f"Done — {ok}/{total} questions re-embedded.")
r.set(f"extraction:status:{job_id}", "completed", ex=EXPIRE_SECONDS)
except Exception as e:
logger.exception(f"Embedding regeneration failed for job {job_id}")
_push_step(r, job_id, "error", f"Failed: {e}")
r.set(f"extraction:status:{job_id}", "failed", ex=EXPIRE_SECONDS)
raise
finally:
db.close()