pdf-quiz-generator/backend/app/services/ai_service.py
Daniel fdebda993c
Some checks failed
Mobile Android Release / android-release (push) Failing after 1s
Improve quiz TTS and Android release build
2026-05-12 01:01:01 +02:00

515 lines
21 KiB
Python
Raw Permalink 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.

import json
import logging
import litellm
from app.config import settings
logger = logging.getLogger(__name__)
def _proxy_model(model_id: str) -> str:
"""Always prefix model with openai/ when using a proxy. The proxy (Bifrost)
expects provider/model format; litellm strips ONE openai/ prefix before
sending to OpenAI-compatible endpoints, so we wrap with an extra openai/."""
if settings.LITELLM_API_BASE:
return f"openai/{model_id}"
return model_id
EXTRACTION_PROMPT = """You are extracting questions from a PREP (Pediatric Review and Education Program) exam PDF.
These PDFs follow a strict format:
1. A numbered question with a clinical vignette (patient scenario)
2. Five answer options labeled A, B, C, D, E
3. A line "Correct Answer: X" or "Preferred Response: X" where X is the letter of the correct option
4. An explanation paragraph
5. A "Critique:" section with detailed reasoning
6. A "Content Specifications:" section listing the learning objectives
Your task: extract every question and return ONLY a JSON object in this exact format:
{{"questions": [
{{
"item_number": "<item number digits only, e.g. '193' — null if not found>",
"question_text": "<full question stem including any patient vignette>",
"question_type": "mcq",
"options": ["<option A text>", "<option B text>", "<option C text>", "<option D text>", "<option E text>"],
"correct_answer": "<full text of the correct option, NOT the letter>",
"explanation": "<ALL text verbatim from after the correct answer line to before the next question — explanation, Critique, Learning Points, Content Specifications, Suggested Reading, everything>",
"has_figure": false,
"page_reference": {page_ref}
}}
]}}
CRITICAL RULES — follow exactly:
1. CORRECT ANSWER: Find "Correct Answer: X" or "Preferred Response: X" after each question.
X is a letter (AE). Look up the full text of that option and store it as correct_answer.
NEVER store just the letter. NEVER guess. If not found, set correct_answer to null.
Example: options=["alpha","beta","gamma","delta","epsilon"] and "Correct Answer: C" → correct_answer="gamma"
2. EXPLANATION: Copy EVERYTHING that appears after the correct answer line and before the next question,
verbatim and in full — including the explanation paragraph, any Critique section, Learning Points,
Content Specifications, American Board of Pediatrics specifications, Suggested Readings, and any
other content. Do NOT summarize, shorten, skip, or paraphrase a single word.
Keep all section headers (e.g. "Critique:", "Content Specifications:", "Suggested Reading:").
The explanation field should be the complete, unaltered text block.
3. QUESTIONS: Extract ALL questions, even if partially cut off.
A new question starts with "Item NNN" or "ltem NNN" (OCR artifact — lowercase l instead of I).
Content after an image reference (e.g. "Item C123A", figure captions) is NOT a new question.
4. OPTIONS: Extract just the text for each option, without the letter prefix (A. B. C. D. E.).
5. HAS_FIGURE: Set to true ONLY if the question text explicitly references an image, figure, photograph,
radiograph, X-ray, ECG, rash photo, growth chart, or similar visual that is essential to answering.
Do NOT set true for decorative logos, page headers, publisher icons, or repeated branding images.
If the question can be answered from text alone, set false.
6. Return ONLY the JSON — no markdown fences, no explanation, no preamble.
Content from page(s) {page_info}:
{content}"""
ANSWER_KEY_PROMPT = """Extract the answer key from this PREP exam content.
The answer key lists items with their correct answer letters, like:
"Item 193 Preferred Response: D"
"Item 194 Preferred Response: A"
Return ONLY a JSON object mapping item numbers to correct letters:
{{"answers": {{"193": "D", "194": "A", "211": "C"}}}}
Rules:
- "Preferred Response: X" or "Correct Answer: X" — X is the letter.
- Item numbers may appear as "ltemXXX" (OCR artifact — l is actually I).
- Only include items where you find a clear correct answer letter.
- Return ONLY the JSON — no markdown, no preamble.
Content from page(s) {page_info}:
{content}"""
def get_model_for_task(db, task: str = "extraction") -> tuple[str, str | None]:
"""Get the configured model for a specific task from DB, or fall back to settings."""
try:
from app.models.ai_model_config import AIModelConfig
config = db.query(AIModelConfig).filter(
AIModelConfig.task == task,
AIModelConfig.is_active == True,
AIModelConfig.is_default == True,
).first()
if config:
return config.model_id, config.api_key
except Exception as e:
logger.warning(f"Failed to load AI model config for task '{task}', using fallback: {e}")
return settings.LITELLM_MODEL, settings.LITELLM_API_KEY or None
def _truncate_content(content: str, max_chars: int = 100000) -> str:
if len(content) <= max_chars:
return content
half = max_chars // 2
return content[:half] + "\n\n... [content truncated] ...\n\n" + content[-half:]
def extract_questions(
content: str,
page_info: str = "unknown",
page_ref: int | None = None,
model_id: str | None = None,
api_key: str | None = None,
) -> list[dict]:
"""Extract quiz questions from PDF content using LiteLLM."""
content = _truncate_content(content)
prompt = EXTRACTION_PROMPT.format(
content=content,
page_info=page_info,
page_ref=page_ref if page_ref else "null",
)
use_model = _proxy_model(model_id or settings.LITELLM_MODEL)
use_key = api_key or settings.LITELLM_API_KEY
last_error = None
for attempt in range(3):
try:
kwargs = {
"model": use_model,
"messages": [{"role": "user", "content": prompt}],
"temperature": 0.1, # low temp for faithful extraction
}
if use_key:
kwargs["api_key"] = use_key
if settings.LITELLM_API_BASE:
kwargs["api_base"] = settings.LITELLM_API_BASE
# Don't force JSON mode — let the model respond naturally and we parse it
response = litellm.completion(**kwargs)
response_text = response.choices[0].message.content
logger.info(f"Model raw response (first 500 chars): {response_text[:500]!r}")
# Try to parse JSON, handle markdown code blocks
text = response_text.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)
# Handle all common response shapes
if isinstance(data, list):
questions = data
elif isinstance(data, dict):
# Try common keys
for key in ("questions", "items", "results", "data"):
if isinstance(data.get(key), list):
questions = data[key]
break
else:
# Maybe the dict itself is a single question
if "question_text" in data:
questions = [data]
else:
raise ValueError(f"Unexpected response shape: {list(data.keys())}")
else:
raise ValueError("Response is not a list of questions")
validated = []
skipped = []
for q in questions:
if "question_text" not in q:
continue
correct = q.get("correct_answer")
if not correct:
skipped.append(q.get("question_text", "")[:120])
continue
qtype = q.get("question_type", "mcq")
if qtype not in ("mcq", "true_false", "fill_blank"):
qtype = "mcq"
validated.append({
"question_text": q["question_text"],
"question_type": qtype,
"options": q.get("options"),
"correct_answer": correct,
"explanation": q.get("explanation", ""),
"page_reference": q.get("page_reference"),
"skipped": [],
})
# Attach skipped list to first question so caller can surface it
if validated and skipped:
validated[0]["skipped"] = skipped
if validated:
return validated
raise ValueError("No valid questions extracted from content")
except Exception as e:
last_error = e
logger.warning(f"Extraction attempt {attempt + 1} failed: {e!r}")
raise RuntimeError(f"Failed to extract questions after 3 attempts: {last_error}")
def _call_model(prompt: str, model_id: str | None, api_key: str | None) -> str:
"""Call the configured LLM and return raw text response."""
use_model = _proxy_model(model_id or settings.LITELLM_MODEL)
use_key = api_key or settings.LITELLM_API_KEY
kwargs = {
"model": use_model,
"messages": [{"role": "user", "content": prompt}],
"temperature": 0.1,
}
if use_key:
kwargs["api_key"] = use_key
if settings.LITELLM_API_BASE:
kwargs["api_base"] = settings.LITELLM_API_BASE
response = litellm.completion(**kwargs)
return response.choices[0].message.content
def _parse_json_response(text: str) -> dict:
text = text.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()
return json.loads(text)
def transcribe_audio(
audio: bytes,
filename: str = "audio.webm",
content_type: str = "audio/webm",
model_id: str | None = None,
api_key: str | None = None,
) -> str | None:
"""Transcribe uploaded audio through LiteLLM's OpenAI-compatible audio endpoint."""
import httpx
if not audio:
return None
base = (settings.LITELLM_API_BASE or "").rstrip("/").removesuffix("/v1")
if not base:
logger.error("LiteLLM API base not configured for STT")
return None
key = api_key or settings.LITELLM_API_KEY
headers = {"Authorization": f"Bearer {key}"} if key else {}
try:
resp = httpx.post(
f"{base}/v1/audio/transcriptions",
headers=headers,
data={"model": model_id or "local-parakeet-v3", "response_format": "json"},
files={"file": (filename or "audio.webm", audio, content_type or "application/octet-stream")},
timeout=120,
)
resp.raise_for_status()
try:
data = resp.json()
if isinstance(data, dict):
return str(data.get("text") or data.get("transcript") or data.get("transcription") or "").strip()
except ValueError:
return resp.text.strip()
except Exception as e:
logger.error(f"LiteLLM STT failed: {e}")
return None
def extract_questions_no_answers(
content: str,
page_info: str = "unknown",
page_ref: int | None = None,
model_id: str | None = None,
api_key: str | None = None,
) -> list[dict]:
"""Extract questions allowing null correct_answer — for PDFs where answers are at the end."""
content = _truncate_content(content)
prompt = EXTRACTION_PROMPT.format(
content=content,
page_info=page_info,
page_ref=page_ref if page_ref else "null",
)
last_error = None
for attempt in range(3):
try:
text = _call_model(prompt, model_id, api_key)
text = text.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)
questions = []
if isinstance(data, list):
questions = data
elif isinstance(data, dict):
for key in ("questions", "items", "results", "data"):
if isinstance(data.get(key), list):
questions = data[key]
break
else:
if "question_text" in data:
questions = [data]
result = []
for q in questions:
if "question_text" not in q:
continue
qtype = q.get("question_type", "mcq")
if qtype not in ("mcq", "true_false", "fill_blank"):
qtype = "mcq"
result.append({
"item_number": str(q.get("item_number") or "").strip().lstrip("0") or None,
"question_text": q["question_text"],
"question_type": qtype,
"options": q.get("options"),
"correct_answer": q.get("correct_answer"), # may be null
"explanation": q.get("explanation", ""),
"page_reference": q.get("page_reference"),
})
if result:
return result
raise ValueError("No questions found in content")
except Exception as e:
last_error = e
logger.warning(f"No-answer extraction attempt {attempt + 1} failed: {e!r}")
raise RuntimeError(f"Failed after 3 attempts: {last_error}")
def extract_answer_key(
content: str,
page_info: str = "unknown",
model_id: str | None = None,
api_key: str | None = None,
) -> dict[str, str]:
"""Extract answer key from end-of-document content. Returns {item_number: letter}."""
content = _truncate_content(content, max_chars=80000)
prompt = ANSWER_KEY_PROMPT.format(content=content, page_info=page_info)
last_error = None
for attempt in range(3):
try:
text = _call_model(prompt, model_id, api_key)
text = text.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)
answers = data.get("answers", data) if isinstance(data, dict) else {}
# Normalize: strip leading zeros, uppercase letters
return {str(k).strip().lstrip("0"): str(v).strip().upper() for k, v in answers.items() if v}
except Exception as e:
last_error = e
logger.warning(f"Answer key extraction attempt {attempt + 1} failed: {e!r}")
logger.warning(f"Answer key extraction failed: {last_error}")
return {}
def generate_tts_audio(
text: str,
model_id: str | None = None,
api_key: str | None = None,
) -> bytes | None:
"""Generate TTS audio. Supports local LiteLLM, local Sherpa, OpenAI, ElevenLabs, and Google Cloud TTS.
model_id conventions:
tts-1:alloy → OpenAI TTS (voice after colon)
tts-1-hd:nova → OpenAI TTS HD
elevenlabs/<voice> → ElevenLabs
google/<voice_name> → Google Cloud TTS (e.g. google/en-US-Wavenet-D)
sherpa/<profile>:<voice> → Local speech gateway (e.g. sherpa/kokoro:am_adam)
local-<model>:<voice> → Local LiteLLM TTS model + voice (e.g. local-kokoro-tts:am_adam)
"""
import httpx, base64
use_model = model_id or "tts-1:alloy"
# ── Local LiteLLM TTS models ─────────────────────────────────
if use_model.startswith("local-"):
local_model = use_model
local_voice = "am_adam" if use_model.startswith("local-kokoro-tts") else "alloy"
if ":" in use_model:
local_model, local_voice = use_model.split(":", 1)
key = api_key or settings.LITELLM_API_KEY
base = (settings.LITELLM_API_BASE or "").rstrip("/").removesuffix("/v1")
if not base:
logger.error("LiteLLM API base not configured for local TTS")
return None
try:
resp = httpx.post(
f"{base}/v1/audio/speech",
headers={"Authorization": f"Bearer {key}", "Content-Type": "application/json"} if key else {"Content-Type": "application/json"},
json={"model": local_model, "input": text, "voice": local_voice, "response_format": "mp3"},
timeout=90,
)
resp.raise_for_status()
return resp.content
except Exception as e:
logger.error(f"Local LiteLLM TTS failed: {e}")
return None
# ── Local Sherpa gateway ───────────────────────────────────
if use_model.startswith("sherpa/"):
payload = use_model[len("sherpa/"):]
profile = payload
voice = "0"
if ":" in payload:
profile, voice = payload.split(":", 1)
try:
resp = httpx.post(
f"{settings.LOCAL_SPEECH_GATEWAY_URL.rstrip('/')}/v1/audio/speech",
json={"model": f"sherpa/{profile}", "input": text, "voice": voice, "response_format": "mp3"},
timeout=90,
)
resp.raise_for_status()
return resp.content
except Exception as e:
logger.error(f"Local Sherpa TTS failed: {e}")
return None
# ── ElevenLabs ─────────────────────────────────────────────
if use_model.startswith("elevenlabs/") or use_model.startswith("eleven_labs/"):
voice = use_model.split("/", 1)[1]
key = api_key or settings.ELEVENLABS_API_KEY
if not key:
logger.error("ElevenLabs API key not configured")
return None
try:
resp = httpx.post(
f"https://api.elevenlabs.io/v1/text-to-speech/{voice}",
headers={"xi-api-key": key, "Content-Type": "application/json"},
json={"text": text, "model_id": "eleven_turbo_v2_5"},
timeout=30,
)
resp.raise_for_status()
return resp.content
except Exception as e:
logger.error(f"ElevenLabs TTS failed: {e}")
return None
# ── Google Cloud TTS ────────────────────────────────────────
if use_model.startswith("google/"):
voice_name = use_model[len("google/"):]
key = api_key or settings.GOOGLE_TTS_API_KEY
if not key:
logger.error("Google TTS API key not configured (GOOGLE_TTS_API_KEY)")
return None
# Parse language code from voice name (e.g. "en-US-Wavenet-D" → "en-US")
parts = voice_name.split("-")
lang_code = f"{parts[0]}-{parts[1]}" if len(parts) >= 2 else "en-US"
try:
resp = httpx.post(
f"https://texttospeech.googleapis.com/v1/text:synthesize?key={key}",
json={
"input": {"text": text},
"voice": {"languageCode": lang_code, "name": voice_name},
"audioConfig": {"audioEncoding": "MP3"},
},
timeout=30,
)
resp.raise_for_status()
return base64.b64decode(resp.json()["audioContent"])
except Exception as e:
logger.error(f"Google TTS failed: {e}")
return None
# ── OpenAI (default) ────────────────────────────────────────
# model_id may encode voice as "tts-1:nova", "tts-1-hd:alloy", etc.
clean_model = use_model.replace("openai/", "")
oai_voice = "alloy"
if ":" in clean_model:
clean_model, oai_voice = clean_model.split(":", 1)
# Per-model key > OPENAI_API_KEY (direct) > LITELLM_API_KEY (proxy)
if api_key:
key = api_key
base = (settings.LITELLM_API_BASE or "https://api.openai.com").rstrip("/").removesuffix("/v1")
elif settings.OPENAI_API_KEY:
key = settings.OPENAI_API_KEY
base = "https://api.openai.com"
else:
key = settings.LITELLM_API_KEY
base = (settings.LITELLM_API_BASE or "https://api.openai.com").rstrip("/").removesuffix("/v1")
try:
resp = httpx.post(
f"{base}/v1/audio/speech",
headers={"Authorization": f"Bearer {key}", "Content-Type": "application/json"},
json={"model": clean_model, "input": text, "voice": oai_voice},
timeout=60,
)
resp.raise_for_status()
return resp.content
except Exception as e:
logger.error(f"OpenAI TTS failed: {e}")
return None