Fix startup and Gemini provider handling
This commit is contained in:
@@ -7,6 +7,7 @@ backups/
|
|||||||
__pycache__/
|
__pycache__/
|
||||||
*.pyc
|
*.pyc
|
||||||
.venv/
|
.venv/
|
||||||
|
.venv312/
|
||||||
venv/
|
venv/
|
||||||
*.egg-info/
|
*.egg-info/
|
||||||
.claude/
|
.claude/
|
||||||
@@ -14,3 +15,4 @@ venv/
|
|||||||
*.jpeg
|
*.jpeg
|
||||||
*.jpg
|
*.jpg
|
||||||
*.png
|
*.png
|
||||||
|
tmp_samples/
|
||||||
|
|||||||
@@ -12,6 +12,10 @@ from routers import documents, review, search, upload, auth
|
|||||||
from services.auth_service import get_session, cleanup_expired_sessions
|
from services.auth_service import get_session, cleanup_expired_sessions
|
||||||
|
|
||||||
|
|
||||||
|
Path(UPLOAD_DIR).mkdir(parents=True, exist_ok=True)
|
||||||
|
Path("data").mkdir(exist_ok=True)
|
||||||
|
|
||||||
|
|
||||||
class CachedStaticFiles(StaticFiles):
|
class CachedStaticFiles(StaticFiles):
|
||||||
"""StaticFiles that sets a long Cache-Control header so browsers
|
"""StaticFiles that sets a long Cache-Control header so browsers
|
||||||
don't re-download the same image on every refresh."""
|
don't re-download the same image on every refresh."""
|
||||||
@@ -29,8 +33,6 @@ class CachedStaticFiles(StaticFiles):
|
|||||||
|
|
||||||
@asynccontextmanager
|
@asynccontextmanager
|
||||||
async def lifespan(app: FastAPI):
|
async def lifespan(app: FastAPI):
|
||||||
Path(UPLOAD_DIR).mkdir(parents=True, exist_ok=True)
|
|
||||||
Path("data").mkdir(exist_ok=True)
|
|
||||||
create_tables()
|
create_tables()
|
||||||
try:
|
try:
|
||||||
cleanup_expired_sessions()
|
cleanup_expired_sessions()
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ google-genai>=1.70.0
|
|||||||
easyocr>=1.7.0
|
easyocr>=1.7.0
|
||||||
pymupdf>=1.24.0
|
pymupdf>=1.24.0
|
||||||
jinja2>=3.1.4
|
jinja2>=3.1.4
|
||||||
|
python-dotenv>=1.0.1
|
||||||
python-multipart>=0.0.9
|
python-multipart>=0.0.9
|
||||||
aiofiles>=23.0.0
|
aiofiles>=23.0.0
|
||||||
pillow>=10.0.0
|
pillow>=10.0.0
|
||||||
|
|||||||
+20
-28
@@ -45,7 +45,7 @@ def require_admin(request: Request):
|
|||||||
|
|
||||||
@router.get("/login", response_class=HTMLResponse)
|
@router.get("/login", response_class=HTMLResponse)
|
||||||
async def login_get(request: Request):
|
async def login_get(request: Request):
|
||||||
return templates.TemplateResponse(request=request, name="login.html", context={"error": None})
|
return templates.TemplateResponse("login.html", {"request": request, "error": None})
|
||||||
|
|
||||||
@router.post("/login", response_class=HTMLResponse)
|
@router.post("/login", response_class=HTMLResponse)
|
||||||
async def login_post(request: Request, username: str = Form(...), password: str = Form(...)):
|
async def login_post(request: Request, username: str = Form(...), password: str = Form(...)):
|
||||||
@@ -54,9 +54,8 @@ async def login_post(request: Request, username: str = Form(...), password: str
|
|||||||
|
|
||||||
if is_login_blocked(username, ip):
|
if is_login_blocked(username, ip):
|
||||||
return templates.TemplateResponse(
|
return templates.TemplateResponse(
|
||||||
request=request,
|
"login.html",
|
||||||
name="login.html",
|
{"request": request, "error": "تم حجب محاولات تسجيل الدخول مؤقتاً. حاول بعد 15 دقيقة."},
|
||||||
context={"error": "تم حجب محاولات تسجيل الدخول مؤقتاً. حاول بعد 15 دقيقة."},
|
|
||||||
status_code=429,
|
status_code=429,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -64,9 +63,8 @@ async def login_post(request: Request, username: str = Form(...), password: str
|
|||||||
if not user or not verify_password(user["password_hash"], password):
|
if not user or not verify_password(user["password_hash"], password):
|
||||||
record_login_attempt(username, ip, success=False)
|
record_login_attempt(username, ip, success=False)
|
||||||
return templates.TemplateResponse(
|
return templates.TemplateResponse(
|
||||||
request=request,
|
"login.html",
|
||||||
name="login.html",
|
{"request": request, "error": "Invalid username or password"},
|
||||||
context={"error": "Invalid username or password"},
|
|
||||||
status_code=401,
|
status_code=401,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -109,9 +107,8 @@ async def logout(request: Request):
|
|||||||
async def users_list(request: Request, current=Depends(require_admin)):
|
async def users_list(request: Request, current=Depends(require_admin)):
|
||||||
users = get_all_users()
|
users = get_all_users()
|
||||||
return templates.TemplateResponse(
|
return templates.TemplateResponse(
|
||||||
request=request,
|
"users.html",
|
||||||
name="users.html",
|
{"request": request, "users": users, "current_user": current},
|
||||||
context={"users": users, "current_user": current},
|
|
||||||
)
|
)
|
||||||
|
|
||||||
@router.post("/users/create")
|
@router.post("/users/create")
|
||||||
@@ -128,16 +125,14 @@ async def add_user(
|
|||||||
if len(password) < 8:
|
if len(password) < 8:
|
||||||
users = get_all_users()
|
users = get_all_users()
|
||||||
return templates.TemplateResponse(
|
return templates.TemplateResponse(
|
||||||
request=request,
|
"users.html",
|
||||||
name="users.html",
|
{"request": request, "users": users, "current_user": current, "error": "Password must be at least 8 characters."},
|
||||||
context={"users": users, "current_user": current, "error": "Password must be at least 8 characters."},
|
|
||||||
)
|
)
|
||||||
if get_user_by_username(username):
|
if get_user_by_username(username):
|
||||||
users = get_all_users()
|
users = get_all_users()
|
||||||
return templates.TemplateResponse(
|
return templates.TemplateResponse(
|
||||||
request=request,
|
"users.html",
|
||||||
name="users.html",
|
{"request": request, "users": users, "current_user": current, "error": f"User '{username}' already exists."},
|
||||||
context={"users": users, "current_user": current, "error": f"User '{username}' already exists."},
|
|
||||||
)
|
)
|
||||||
create_user(username, password, role=role)
|
create_user(username, password, role=role)
|
||||||
return RedirectResponse(url="/auth/users", status_code=status.HTTP_303_SEE_OTHER)
|
return RedirectResponse(url="/auth/users", status_code=status.HTTP_303_SEE_OTHER)
|
||||||
@@ -147,9 +142,8 @@ async def remove_user(request: Request, user_id: int, current=Depends(require_ad
|
|||||||
if user_id == current["user_id"]:
|
if user_id == current["user_id"]:
|
||||||
users = get_all_users()
|
users = get_all_users()
|
||||||
return templates.TemplateResponse(
|
return templates.TemplateResponse(
|
||||||
request=request,
|
"users.html",
|
||||||
name="users.html",
|
{"request": request, "users": users, "current_user": current, "error": "لا يمكنك حذف حسابك الخاص."},
|
||||||
context={"users": users, "current_user": current, "error": "لا يمكنك حذف حسابك الخاص."},
|
|
||||||
status_code=400,
|
status_code=400,
|
||||||
)
|
)
|
||||||
# Prevent removing the last admin
|
# Prevent removing the last admin
|
||||||
@@ -157,9 +151,8 @@ async def remove_user(request: Request, user_id: int, current=Depends(require_ad
|
|||||||
if target and target["role"] == "admin" and count_admins() <= 1:
|
if target and target["role"] == "admin" and count_admins() <= 1:
|
||||||
users = get_all_users()
|
users = get_all_users()
|
||||||
return templates.TemplateResponse(
|
return templates.TemplateResponse(
|
||||||
request=request,
|
"users.html",
|
||||||
name="users.html",
|
{"request": request, "users": users, "current_user": current, "error": "لا يمكن حذف آخر مسؤول في النظام."},
|
||||||
context={"users": users, "current_user": current, "error": "لا يمكن حذف آخر مسؤول في النظام."},
|
|
||||||
status_code=400,
|
status_code=400,
|
||||||
)
|
)
|
||||||
delete_user(user_id)
|
delete_user(user_id)
|
||||||
@@ -180,9 +173,8 @@ async def change_role(
|
|||||||
if target and target["role"] == "admin" and role != "admin" and count_admins() <= 1:
|
if target and target["role"] == "admin" and role != "admin" and count_admins() <= 1:
|
||||||
users = get_all_users()
|
users = get_all_users()
|
||||||
return templates.TemplateResponse(
|
return templates.TemplateResponse(
|
||||||
request=request,
|
"users.html",
|
||||||
name="users.html",
|
{"request": request, "users": users, "current_user": current, "error": "لا يمكن تخفيض رتبة آخر مسؤول."},
|
||||||
context={"users": users, "current_user": current, "error": "لا يمكن تخفيض رتبة آخر مسؤول."},
|
|
||||||
status_code=400,
|
status_code=400,
|
||||||
)
|
)
|
||||||
set_user_role(user_id, role)
|
set_user_role(user_id, role)
|
||||||
@@ -196,9 +188,9 @@ async def backup_db(request: Request, current=Depends(require_admin)):
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
msg = f"Backup failed: {str(e)}"
|
msg = f"Backup failed: {str(e)}"
|
||||||
return templates.TemplateResponse(
|
return templates.TemplateResponse(
|
||||||
request=request,
|
"users.html",
|
||||||
name="users.html",
|
{
|
||||||
context={
|
"request": request,
|
||||||
"users": get_all_users(),
|
"users": get_all_users(),
|
||||||
"current_user": current,
|
"current_user": current,
|
||||||
"backup_msg": msg,
|
"backup_msg": msg,
|
||||||
|
|||||||
+3
-1
@@ -12,6 +12,7 @@ from services.extractor import (
|
|||||||
extract_document,
|
extract_document,
|
||||||
get_available_providers,
|
get_available_providers,
|
||||||
get_default_provider,
|
get_default_provider,
|
||||||
|
provider_supports_ai_verification,
|
||||||
verify_page_correlation,
|
verify_page_correlation,
|
||||||
)
|
)
|
||||||
from services.search_service import normalize_arabic, _normalize_scope
|
from services.search_service import normalize_arabic, _normalize_scope
|
||||||
@@ -336,7 +337,8 @@ async def review_document(request: Request, doc_id: int, wait: int = 0):
|
|||||||
"providers": get_available_providers(),
|
"providers": get_available_providers(),
|
||||||
"current_provider": doc.get("provider") or get_default_provider(),
|
"current_provider": doc.get("provider") or get_default_provider(),
|
||||||
"ai_verification_available": any(
|
"ai_verification_available": any(
|
||||||
provider["id"] in {"claude", "gemini"} for provider in get_available_providers()
|
provider_supports_ai_verification(provider["id"])
|
||||||
|
for provider in get_available_providers()
|
||||||
),
|
),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|||||||
+12
-1
@@ -9,7 +9,7 @@ from fastapi import APIRouter, File, Form, HTTPException, Request, UploadFile
|
|||||||
from fastapi.responses import JSONResponse, RedirectResponse
|
from fastapi.responses import JSONResponse, RedirectResponse
|
||||||
from fastapi.templating import Jinja2Templates
|
from fastapi.templating import Jinja2Templates
|
||||||
|
|
||||||
from config import UPLOAD_DIR
|
from config import UPLOAD_DIR, MAX_CONCURRENT_EXTRACTIONS
|
||||||
from database.connection import get_db
|
from database.connection import get_db
|
||||||
from services.extractor import extract_document, get_available_providers, get_default_provider
|
from services.extractor import extract_document, get_available_providers, get_default_provider
|
||||||
from services.pdf_handler import pdf_to_images
|
from services.pdf_handler import pdf_to_images
|
||||||
@@ -17,6 +17,16 @@ from services.pdf_handler import pdf_to_images
|
|||||||
router = APIRouter()
|
router = APIRouter()
|
||||||
templates = Jinja2Templates(directory="templates")
|
templates = Jinja2Templates(directory="templates")
|
||||||
|
|
||||||
|
# Semaphore: limits simultaneous AI API calls to avoid quota exhaustion on bulk uploads
|
||||||
|
_extraction_semaphore: asyncio.Semaphore | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def _get_semaphore() -> asyncio.Semaphore:
|
||||||
|
global _extraction_semaphore
|
||||||
|
if _extraction_semaphore is None:
|
||||||
|
_extraction_semaphore = asyncio.Semaphore(MAX_CONCURRENT_EXTRACTIONS)
|
||||||
|
return _extraction_semaphore
|
||||||
|
|
||||||
ALLOWED_IMAGE_EXTS = {".jpg", ".jpeg", ".png", ".webp"}
|
ALLOWED_IMAGE_EXTS = {".jpg", ".jpeg", ".png", ".webp"}
|
||||||
ALLOWED_PDF_EXTS = {".pdf"}
|
ALLOWED_PDF_EXTS = {".pdf"}
|
||||||
ALLOWED_EXTENSIONS = ALLOWED_IMAGE_EXTS | ALLOWED_PDF_EXTS
|
ALLOWED_EXTENSIONS = ALLOWED_IMAGE_EXTS | ALLOWED_PDF_EXTS
|
||||||
@@ -91,6 +101,7 @@ def _save_image(file_bytes: bytes, original_name: str) -> str:
|
|||||||
|
|
||||||
async def _extract_and_save(doc_id: int, image_path: str, provider: str = ""):
|
async def _extract_and_save(doc_id: int, image_path: str, provider: str = ""):
|
||||||
"""Background task: call extractor, parse result, update DB."""
|
"""Background task: call extractor, parse result, update DB."""
|
||||||
|
async with _get_semaphore():
|
||||||
try:
|
try:
|
||||||
data = await extract_document(image_path, provider=provider)
|
data = await extract_document(image_path, provider=provider)
|
||||||
raw_json = json.dumps(data, ensure_ascii=False)
|
raw_json = json.dumps(data, ensure_ascii=False)
|
||||||
|
|||||||
+50
-14
@@ -1,4 +1,5 @@
|
|||||||
import base64
|
import base64
|
||||||
|
import importlib.util
|
||||||
import json
|
import json
|
||||||
import re
|
import re
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
@@ -61,7 +62,7 @@ Rules:
|
|||||||
|
|
||||||
def _get_ai_verification_provider(preferred_provider: str = "") -> str:
|
def _get_ai_verification_provider(preferred_provider: str = "") -> str:
|
||||||
"""Return a provider capable of vision reasoning for verification."""
|
"""Return a provider capable of vision reasoning for verification."""
|
||||||
providers = [p["id"] for p in get_available_providers() if p["id"] in {"claude", "gemini"}]
|
providers = [p["id"] for p in get_available_providers() if provider_supports_ai_verification(p["id"])]
|
||||||
if preferred_provider in providers:
|
if preferred_provider in providers:
|
||||||
return preferred_provider
|
return preferred_provider
|
||||||
if DEFAULT_PROVIDER in providers:
|
if DEFAULT_PROVIDER in providers:
|
||||||
@@ -71,6 +72,10 @@ def _get_ai_verification_provider(preferred_provider: str = "") -> str:
|
|||||||
raise ValueError("No AI verification provider configured")
|
raise ValueError("No AI verification provider configured")
|
||||||
|
|
||||||
|
|
||||||
|
def provider_supports_ai_verification(provider_id: str) -> bool:
|
||||||
|
return provider_id == "claude" or provider_id == "gemini" or provider_id.startswith("gemini-")
|
||||||
|
|
||||||
|
|
||||||
def _build_correlation_user_prompt(current_context: dict, candidate_context: dict) -> str:
|
def _build_correlation_user_prompt(current_context: dict, candidate_context: dict) -> str:
|
||||||
return (
|
return (
|
||||||
"Determine whether PAGE_A and PAGE_B belong to the same multi-page request/document for the same person. "
|
"Determine whether PAGE_A and PAGE_B belong to the same multi-page request/document for the same person. "
|
||||||
@@ -249,15 +254,29 @@ def _parse_json_lenient(text: str) -> dict:
|
|||||||
raise ValueError(f"Could not parse model response as JSON: {text[:200]}")
|
raise ValueError(f"Could not parse model response as JSON: {text[:200]}")
|
||||||
|
|
||||||
|
|
||||||
|
# Individual Gemini models exposed to the UI
|
||||||
|
_GEMINI_MODEL_OPTIONS = [
|
||||||
|
("gemini-2.5-pro", "Gemini 2.5 Pro (أفضل دقة)"),
|
||||||
|
("gemini-2.5-flash", "Gemini 2.5 Flash (سريع)"),
|
||||||
|
("gemini-2.0-flash", "Gemini 2.0 Flash"),
|
||||||
|
("gemini-2.0-flash-lite", "Gemini 2.0 Flash Lite (احتياطي)"),
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def _easyocr_is_available() -> bool:
|
||||||
|
return importlib.util.find_spec("easyocr") is not None
|
||||||
|
|
||||||
|
|
||||||
def get_available_providers() -> list[dict]:
|
def get_available_providers() -> list[dict]:
|
||||||
"""Return list of all available providers (EasyOCR is always available)."""
|
"""Return only providers that are actually usable in the current runtime."""
|
||||||
providers = [
|
providers = []
|
||||||
{"id": "easyocr", "name": "EasyOCR (مجاني)", "model": "local"},
|
if _easyocr_is_available():
|
||||||
]
|
providers.append({"id": "easyocr", "name": "EasyOCR (مجاني)", "model": "local"})
|
||||||
if ANTHROPIC_API_KEY:
|
if ANTHROPIC_API_KEY:
|
||||||
providers.append({"id": "claude", "name": "Claude (Anthropic)", "model": CLAUDE_MODEL})
|
providers.append({"id": "claude", "name": "Claude (Anthropic)", "model": CLAUDE_MODEL})
|
||||||
if GEMINI_API_KEY:
|
if GEMINI_API_KEY:
|
||||||
providers.append({"id": "gemini", "name": "Gemini (Google)", "model": GEMINI_MODEL})
|
for model_id, label in _GEMINI_MODEL_OPTIONS:
|
||||||
|
providers.append({"id": model_id, "name": label, "model": model_id})
|
||||||
return providers
|
return providers
|
||||||
|
|
||||||
|
|
||||||
@@ -265,8 +284,12 @@ def get_default_provider() -> str:
|
|||||||
"""Return the default provider, falling back to whichever is available."""
|
"""Return the default provider, falling back to whichever is available."""
|
||||||
providers = get_available_providers()
|
providers = get_available_providers()
|
||||||
if not providers:
|
if not providers:
|
||||||
return "easyocr"
|
return GEMINI_MODEL if DEFAULT_PROVIDER == "gemini" else DEFAULT_PROVIDER
|
||||||
ids = [p["id"] for p in providers]
|
ids = [p["id"] for p in providers]
|
||||||
|
# Map legacy "gemini" default to the configured GEMINI_MODEL
|
||||||
|
resolved = GEMINI_MODEL if DEFAULT_PROVIDER == "gemini" else DEFAULT_PROVIDER
|
||||||
|
if resolved in ids:
|
||||||
|
return resolved
|
||||||
if DEFAULT_PROVIDER in ids:
|
if DEFAULT_PROVIDER in ids:
|
||||||
return DEFAULT_PROVIDER
|
return DEFAULT_PROVIDER
|
||||||
return ids[0]
|
return ids[0]
|
||||||
@@ -329,7 +352,7 @@ async def _extract_with_claude(image_path: str) -> dict:
|
|||||||
|
|
||||||
# ─── Gemini extraction ───────────────────────────────────────────
|
# ─── Gemini extraction ───────────────────────────────────────────
|
||||||
|
|
||||||
async def _extract_with_gemini(image_path: str) -> dict:
|
async def _extract_with_gemini(image_path: str, model: str = "") -> dict:
|
||||||
import asyncio
|
import asyncio
|
||||||
from google import genai
|
from google import genai
|
||||||
from google.genai import types
|
from google.genai import types
|
||||||
@@ -345,8 +368,17 @@ async def _extract_with_gemini(image_path: str) -> dict:
|
|||||||
".png": "image/png", ".webp": "image/webp"}
|
".png": "image/png", ".webp": "image/webp"}
|
||||||
mime_type = mime_map.get(suffix, "image/jpeg")
|
mime_type = mime_map.get(suffix, "image/jpeg")
|
||||||
|
|
||||||
# Try primary model first; only fall back on 503 (overloaded), NOT on 429 (quota)
|
# Full fallback chain: start from requested model, cascade through cheaper/available ones
|
||||||
models_to_try = [GEMINI_MODEL, "gemini-2.0-flash-lite"]
|
_ALL_GEMINI_FALLBACK = [
|
||||||
|
"gemini-2.5-pro",
|
||||||
|
"gemini-2.5-flash",
|
||||||
|
"gemini-2.0-flash",
|
||||||
|
"gemini-2.0-flash-lite",
|
||||||
|
]
|
||||||
|
# If a specific model was requested, start from it; otherwise start from the configured default
|
||||||
|
start_model = model if model else GEMINI_MODEL
|
||||||
|
# Build ordered list: requested model first, then remaining fallbacks in order
|
||||||
|
models_to_try = [start_model] + [m for m in _ALL_GEMINI_FALLBACK if m != start_model]
|
||||||
|
|
||||||
def _call(model_name: str):
|
def _call(model_name: str):
|
||||||
response = client.models.generate_content(
|
response = client.models.generate_content(
|
||||||
@@ -731,6 +763,8 @@ async def extract_document(image_path: str, provider: str = "") -> dict:
|
|||||||
"""
|
"""
|
||||||
Extract structured data from a document image using the specified provider.
|
Extract structured data from a document image using the specified provider.
|
||||||
Falls back to the default provider if none specified.
|
Falls back to the default provider if none specified.
|
||||||
|
provider can be: "claude", "easyocr", "gemini" (uses GEMINI_MODEL),
|
||||||
|
or a specific Gemini model ID like "gemini-2.5-flash".
|
||||||
"""
|
"""
|
||||||
if not provider:
|
if not provider:
|
||||||
provider = get_default_provider()
|
provider = get_default_provider()
|
||||||
@@ -739,11 +773,13 @@ async def extract_document(image_path: str, provider: str = "") -> dict:
|
|||||||
if not ANTHROPIC_API_KEY:
|
if not ANTHROPIC_API_KEY:
|
||||||
raise ValueError("ANTHROPIC_API_KEY not set")
|
raise ValueError("ANTHROPIC_API_KEY not set")
|
||||||
return await _extract_with_claude(image_path)
|
return await _extract_with_claude(image_path)
|
||||||
elif provider == "gemini":
|
|
||||||
if not GEMINI_API_KEY:
|
|
||||||
raise ValueError("GEMINI_API_KEY not set")
|
|
||||||
return await _extract_with_gemini(image_path)
|
|
||||||
elif provider == "easyocr":
|
elif provider == "easyocr":
|
||||||
return await _extract_with_easyocr(image_path)
|
return await _extract_with_easyocr(image_path)
|
||||||
|
elif provider == "gemini" or provider.startswith("gemini-"):
|
||||||
|
if not GEMINI_API_KEY:
|
||||||
|
raise ValueError("GEMINI_API_KEY not set")
|
||||||
|
# Pass the specific model if the provider ID encodes one
|
||||||
|
model_override = provider if provider.startswith("gemini-") else ""
|
||||||
|
return await _extract_with_gemini(image_path, model=model_override)
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Unknown provider: {provider}")
|
raise ValueError(f"Unknown provider: {provider}")
|
||||||
|
|||||||
+15
-4
@@ -6,25 +6,35 @@ from pathlib import Path
|
|||||||
|
|
||||||
from config import UPLOAD_DIR
|
from config import UPLOAD_DIR
|
||||||
|
|
||||||
|
# Maximum pages we will extract from a single PDF upload.
|
||||||
|
# Beyond this limit pages are silently dropped to prevent runaway queue growth.
|
||||||
|
MAX_PDF_PAGES = 200
|
||||||
|
|
||||||
|
|
||||||
def pdf_to_images(pdf_bytes: bytes, original_name: str) -> list[dict]:
|
def pdf_to_images(pdf_bytes: bytes, original_name: str) -> list[dict]:
|
||||||
"""
|
"""
|
||||||
Convert a PDF to individual page images.
|
Convert a PDF to individual page images.
|
||||||
Returns list of dicts: [{"image_path": "relative/path.png", "page_number": 1}, ...]
|
|
||||||
|
Each page is rendered at 250 DPI — sufficient for Arabic OCR while keeping
|
||||||
|
file sizes manageable. For scanned/image-only PDFs PyMuPDF simply renders
|
||||||
|
the embedded raster; no additional image extraction step is needed.
|
||||||
|
|
||||||
|
Returns list of dicts: [{"image_path": "relative/path.png", "page_number": 1, ...}, ...]
|
||||||
"""
|
"""
|
||||||
import fitz # PyMuPDF
|
import fitz # PyMuPDF
|
||||||
|
|
||||||
doc = fitz.open(stream=pdf_bytes, filetype="pdf")
|
doc = fitz.open(stream=pdf_bytes, filetype="pdf")
|
||||||
|
total_pages = len(doc)
|
||||||
group_id = uuid.uuid4().hex
|
group_id = uuid.uuid4().hex
|
||||||
today = date.today().isoformat()
|
today = date.today().isoformat()
|
||||||
dest_dir = Path(UPLOAD_DIR) / today
|
dest_dir = Path(UPLOAD_DIR) / today
|
||||||
dest_dir.mkdir(parents=True, exist_ok=True)
|
dest_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
pages = []
|
pages = []
|
||||||
for page_num in range(len(doc)):
|
for page_num in range(min(total_pages, MAX_PDF_PAGES)):
|
||||||
page = doc[page_num]
|
page = doc[page_num]
|
||||||
# Render at 200 DPI for good OCR quality
|
# 250 DPI gives sharper Arabic text while staying under ~2 MB per PNG
|
||||||
pix = page.get_pixmap(dpi=200)
|
pix = page.get_pixmap(dpi=250)
|
||||||
|
|
||||||
filename = f"{group_id}_p{page_num + 1}.png"
|
filename = f"{group_id}_p{page_num + 1}.png"
|
||||||
dest = dest_dir / filename
|
dest = dest_dir / filename
|
||||||
@@ -34,6 +44,7 @@ def pdf_to_images(pdf_bytes: bytes, original_name: str) -> list[dict]:
|
|||||||
"image_path": str(Path(today) / filename),
|
"image_path": str(Path(today) / filename),
|
||||||
"page_number": page_num + 1,
|
"page_number": page_num + 1,
|
||||||
"pdf_group_id": group_id,
|
"pdf_group_id": group_id,
|
||||||
|
"total_pages": total_pages,
|
||||||
})
|
})
|
||||||
|
|
||||||
doc.close()
|
doc.close()
|
||||||
|
|||||||
Reference in New Issue
Block a user