saharak's picture
Initial DeepSite commit
15621ed verified
Raw
History Blame Contribute Delete
4.8 kB
import time
import uuid
from typing import List, Dict, Any
from concurrent.futures import ThreadPoolExecutor
import asyncio
from modules.ocr.receipt_detector import detect_receipt_type
from modules.ocr.models_donut import donut_model
from modules.ocr.models_paddle import paddle_model
from modules.ocr.models_trocr import trocr_model
from core.telemetry import log_telemetry
# Job storage (in production, use Redis or persistent DB)
jobs = {}
def get_model_priority(receipt_type: str) -> List[str]:
"""Get model priority based on receipt type"""
priorities = {
"structured": ["donut", "paddle", "trocr"],
"printed": ["trocr", "donut", "paddle"],
"handwritten": ["paddle", "trocr", "donut"],
"blurry": ["paddle", "trocr", "donut"]
}
return priorities.get(receipt_type, ["donut", "paddle", "trocr"])
def process_single_receipt(image_path: str, filename: str) -> Dict[str, Any]:
"""Process one receipt through the cascade"""
start_time = time.time()
receipt_type = detect_receipt_type(image_path)
models = get_model_priority(receipt_type)
models_tried = []
best_result = None
best_confidence = 0.0
best_model = ""
for model_name in models:
models_tried.append(model_name)
if model_name == "donut":
result, confidence = donut_model.load(), 0.0
result, confidence = donut_model.process(image_path)
elif model_name == "paddle":
result, confidence = paddle_model.process(image_path)
elif model_name == "trocr":
result, confidence = trocr_model.process(image_path)
else:
continue
if confidence >= best_confidence:
best_confidence = confidence
best_result = result
best_model = model_name
if confidence >= 0.70:
break # Stop cascade if good enough
processing_time = int((time.time() - start_time) * 1000)
# Log telemetry
log_telemetry({
"model_used": best_model,
"models_tried": models_tried,
"confidence": best_confidence,
"receipt_type": receipt_type,
"file_type": image_path.split('.')[-1],
"field_count": len(best_result.keys()) if best_result else 0,
"processing_ms": processing_time
})
warning = None
if best_confidence < 0.70:
warning = "No model reached 70% confidence. Best result shown."
return {
"filename": filename,
"receiptType": receipt_type,
"modelUsed": best_model,
"modelsTried": models_tried,
"overallConfidence": round(best_confidence, 2),
"fields": [
{"label": k.replace("_", " ").title(), "value": str(v), "confidence": best_confidence}
for k, v in (best_result or {}).items() if k != "line_items"
],
"lineItems": best_result.get("line_items", []) if best_result else [],
"rawText": "", # Would contain actual OCR text
"processingMs": processing_time,
"warning": warning
}
async def process_batch_async(file_paths: List[tuple], job_id: str):
"""Process batch in background"""
jobs[job_id] = {
"status": "processing",
"progress": 0,
"total": len(file_paths),
"results": []
}
loop = asyncio.get_event_loop()
with ThreadPoolExecutor() as pool:
for idx, (path, filename) in enumerate(file_paths):
try:
result = await loop.run_in_executor(
pool, process_single_receipt, path, filename
)
jobs[job_id]["results"].append(result)
jobs[job_id]["progress"] = idx + 1
except Exception as e:
jobs[job_id]["results"].append({
"filename": filename,
"error": str(e),
"receiptType": "error",
"modelUsed": "none",
"modelsTried": [],
"overallConfidence": 0.0,
"fields": [],
"lineItems": [],
"rawText": "",
"processingMs": 0,
"warning": str(e)
})
jobs[job_id]["status"] = "completed"
jobs[job_id]["progress"] = len(file_paths)
def create_job() -> str:
job_id = str(uuid.uuid4())
jobs[job_id] = {"status": "pending", "progress": 0, "total": 0, "results": []}
return job_id
def get_job_status(job_id: str) -> dict:
return jobs.get(job_id, {"status": "not_found"})
def get_job_results(job_id: str) -> List[Dict]:
job = jobs.get(job_id)
if job and job["status"] == "completed":
return job["results"]
return []