Hztech's picture
Update app.py
93cf99b verified
Raw History Blame Contribute Delete
7.81 kB
import os
import sys
import tempfile
from pathlib import Path
import pandas as pd
# Keep Gradio on client-side rendering. In hosted Gradio 6 runs, SSR starts a
# Node proxy that can emit noisy asyncio cleanup tracebacks at startup/shutdown.
os.environ.setdefault("GRADIO_SSR_MODE", "false")
import gradio as gr
import time
# Ensure project root is in python path
sys.path.append(os.path.dirname(os.path.abspath(__file__)))
try:
import spaces
except ImportError:
spaces = None
from extractor import (
PassportExtractor,
format_fly_baghdad,
format_fly_dubai,
format_iraqi,
)
def gpu_task(duration=120):
if spaces is None:
return lambda fn: fn
return spaces.GPU(duration=duration)
# Initialize Extractor
print("Initializing Passport Extractor...")
EXTRACTOR = PassportExtractor(use_gpu=False)
print("Extractor initialized.")
MAX_FILE_SIZE_MB = 100
ALLOWED_FILE_TYPES = [
".png",
".jpg",
".jpeg",
".pdf",
".avif",
".webp",
".bmp",
".tiff",
]
def _format_data(good_results, airline):
if airline == "Fly Dubai":
return format_fly_dubai(good_results)
if airline == "Iraqi":
return format_iraqi(good_results)
if airline == "Fly Baghdad":
return format_fly_baghdad(good_results)
# Default view
df = pd.DataFrame(good_results)
if not df.empty:
from extractor import format_date
if "date_of_birth" in df.columns:
df["date_of_birth"] = df["date_of_birth"].apply(lambda x: format_date(x, "%d/%m/%Y", is_dob=True))
if "expiration_date" in df.columns:
df["expiration_date"] = df["expiration_date"].apply(lambda x: format_date(x, "%d/%m/%Y", is_dob=False))
return df
@gpu_task(duration=300)
def process_files(files, airline, export_all_files, progress=gr.Progress()):
if not files:
return None, "No files uploaded.", None, None
start_time = time.time()
all_results, problematic_files = [], []
total_files = len(files)
from concurrent.futures import ThreadPoolExecutor, as_completed
def _uploaded_path(file_obj):
if isinstance(file_obj, (str, os.PathLike)):
return os.fspath(file_obj)
if hasattr(file_obj, "name"):
return file_obj.name
if isinstance(file_obj, dict):
return file_obj.get("path") or file_obj.get("name")
return None
def process_single_file(file_obj):
file_path = _uploaded_path(file_obj)
if not file_path:
return [], {
"file_name": "unknown",
"reason": f"Unsupported upload object: {type(file_obj).__name__}",
}
file_name = Path(file_path).name
ext = os.path.splitext(file_name)[1].lower()
file_results, file_problem = [], None
if os.path.getsize(file_path) > MAX_FILE_SIZE_MB * 1024 * 1024:
return [], {
"file_name": file_name,
"reason": f"File too large. Max {MAX_FILE_SIZE_MB} MB allowed.",
}
try:
if ext == ".pdf":
def pdf_progress(p):
progress(
completed_files / total_files + (p * (1 / total_files)),
desc=f"Processing {file_name}...",
)
file_results = (
EXTRACTOR.process_pdf(
file_path,
progress_callback=pdf_progress,
airline=(airline or "Default").lower(),
)
or []
)
else:
result = EXTRACTOR.get_data(
file_path, airline=(airline or "Default").lower()
)
if result:
file_results = [result]
except Exception as e:
file_problem = {"file_name": file_name, "reason": f"Error: {e}"}
mrz_found = False
for res in file_results:
res["source_file"] = file_name
if res.get("mrz_found"):
mrz_found = True
if not file_problem:
if not file_results:
file_problem = {
"file_name": file_name,
"reason": "No passport data detected",
}
elif not mrz_found:
file_problem = {"file_name": file_name, "reason": "MRZ data not found"}
return file_results, file_problem
completed_files = 0
with ThreadPoolExecutor(max_workers=2) as executor:
future_to_file = {executor.submit(process_single_file, f): f for f in files}
for future in as_completed(future_to_file):
completed_files += 1
try:
file_results, file_problem = future.result()
except Exception as e:
original_path = _uploaded_path(future_to_file[future]) or "unknown"
file_results = []
file_problem = {
"file_name": Path(original_path).name,
"reason": f"Error: {e}",
}
all_results.extend(file_results)
if file_problem:
problematic_files.append(file_problem)
good_results = [r for r in all_results if r.get("mrz_found")]
end_time = time.time()
duration = round(end_time - start_time, 2)
if not good_results and problematic_files:
status_msg = f"No data could be extracted. (Time: {duration}s)"
status_msg += "\n\nProblems:\n" + "\n".join(
[f"- {p['file_name']}: {p['reason']}" for p in problematic_files]
)
return None, status_msg, None, None
if not good_results:
return (
None,
f"⚠️ No data could be extracted. (Time: {duration}s)",
None,
None,
)
df = _format_data(good_results, airline)
status_msg = f"Extracted {len(good_results)} passport(s) from {total_files} file(s) in {duration} seconds."
if problematic_files:
status_msg += "\n\nProblems:\n" + "\n".join(
[f"- {p['file_name']}: {p['reason']}" for p in problematic_files]
)
csv_path = tempfile.mktemp(suffix=".csv")
df.to_csv(csv_path, index=False)
excel_path = tempfile.mktemp(suffix=".xlsx")
df.to_excel(excel_path, index=False)
return df, status_msg, csv_path, excel_path
def main():
with gr.Blocks(title="Passport OCR Extractor") as demo:
gr.Markdown("#Passport OCR Extractor")
with gr.Row():
with gr.Column(scale=1):
airline = gr.Dropdown(
choices=["Default", "Fly Dubai", "Iraqi", "Fly Baghdad"],
value="Default",
label="Airline Format",
)
file_input = gr.File(
file_count="multiple",
label="Upload Passports",
file_types=ALLOWED_FILE_TYPES,
)
extract_btn = gr.Button("Extract Data", variant="primary")
with gr.Column(scale=2):
status_output = gr.Markdown()
table_output = gr.Dataframe(label="Extracted Data")
with gr.Row():
csv_download = gr.File(label="Download CSV")
excel_download = gr.File(label="Download Excel")
extract_btn.click(
fn=process_files,
inputs=[file_input, airline, gr.State(False)],
outputs=[table_output, status_output, csv_download, excel_download],
)
port = int(os.environ.get("PORT", "7860"))
demo.launch(server_name="0.0.0.0", server_port=port, show_error=True, ssr_mode=False)
if __name__ == "__main__":
main()