Passport-Extractor / extractor.py
Hztech's picture
Update extractor.py
185dade verified
Raw History Blame Contribute Delete
23.9 kB
import os
import cv2
import numpy as np
# On Hugging Face GPU Spaces, the `spaces` package must be imported before
# EasyOCR/Torch touches CUDA, otherwise Gradio's Spaces watcher can crash.
try:
import spaces # noqa: F401
except ImportError:
pass
import easyocr
import ssl
import re
import string as st
import logging
import pandas as pd
import tempfile
import fitz # PyMuPDF
import shutil
import pycountry
from datetime import datetime, date
from passporteye import read_mrz
# ---------------------------------------------------
# CONFIGURATION & SETTINGS
# ---------------------------------------------------
OCR_LANGUAGES = ["en"]
USE_GPU = False
# ---------------------------------------------------
# LOGGING & UTILS
# ---------------------------------------------------
logging.basicConfig(
level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
)
logger = logging.getLogger(__name__)
def clean_string(text):
if not text:
return ""
return "".join(i for i in str(text) if i.isalnum()).upper()
def clean_name_field(text):
if not text:
return ""
text = str(text).upper()
text = re.sub(r"([<KG]){6,}$", "", text)
text = text.replace("<", " ")
text = "".join(c for c in text if not c.isdigit())
tokens = [token for token in text.split() if token]
while tokens and (
re.fullmatch(r"([A-Z])\1{2,}", tokens[-1])
or re.fullmatch(r"[KG]{2,}", tokens[-1])
):
tokens.pop()
return " ".join(tokens).strip()
def parse_mrz_line1(line1):
if not line1:
return "", "", ""
line1 = clean_mrz_line(line1)
issuing_country = normalize_country_code(line1[2:5])
name_field = line1[5:44] if len(line1) >= 6 else line1
name_field = re.sub(r"([<KG]){6,}$", "", name_field)
parts = name_field.split("<<", 1)
surname = clean_name_field(parts[0])
given = clean_name_field(parts[1] if len(parts) > 1 else "")
return issuing_country, surname, given
def clean_mrz_line(line):
if not line:
return ""
line = line.upper().replace(" ", "")
allowed = set(st.ascii_uppercase + st.digits + "<")
line = "".join([c for c in line if c in allowed])
p_start = line.find("P<")
if 0 <= p_start <= 3:
line = line[p_start:]
elif len(line) > 44:
line = line[-44:]
if len(line) < 44:
line += "<" * (44 - len(line))
return line[:44]
def mrz_check_digit(value):
weights = [7, 3, 1]
total = 0
for i, char in enumerate(value):
if char.isdigit():
number = int(char)
elif "A" <= char <= "Z":
number = ord(char) - ord("A") + 10
elif char == "<":
number = 0
else:
number = 0
total += number * weights[i % 3]
return str(total % 10)
def normalize_mrz_digits(text):
table = str.maketrans(
{"O": "0", "Q": "0", "D": "0", "I": "1", "L": "1", "S": "5", "B": "8"}
)
return (text or "").upper().translate(table)
def valid_yymmdd(text):
text = normalize_mrz_digits(text)
if len(text) != 6 or not text.isdigit():
return False
try:
month = int(text[2:4])
day = int(text[4:6])
return 1 <= month <= 12 and 1 <= day <= 31
except:
return False
def normalize_country_code(code):
if not code:
return ""
code = code.upper().translate(
str.maketrans({"0": "O", "1": "I", "2": "Z", "4": "A", "5": "S", "8": "B"})
)
code = "".join(c for c in code if c.isalpha())
if len(code) != 3:
return ""
if pycountry.countries.get(alpha_3=code):
return code
return ""
def clean_mrz_line2(line):
if not line:
return ""
line = line.upper().replace(" ", "")
allowed = set(st.ascii_uppercase + st.digits + "<")
line = "".join([c for c in line if c in allowed])
if len(line) <= 44:
return (line + "<" * (44 - len(line)))[:44]
def score(candidate):
doc = candidate[0:9]
doc_check = candidate[9]
nationality = normalize_country_code(candidate[10:13])
dob = normalize_mrz_digits(candidate[13:19])
dob_check = candidate[19]
expiry = normalize_mrz_digits(candidate[21:27])
expiry_check = candidate[27]
value = 0
value += 12 if doc_check.isdigit() and mrz_check_digit(doc) == doc_check else 0
value += 10 if nationality else 0
value += 14 if valid_yymmdd(dob) else 0
value += 18 if dob_check.isdigit() and mrz_check_digit(dob) == dob_check else 0
# Sex field score with OCR error tolerance
sex_char = candidate[20]
if sex_char in "MFX<":
value += 8
elif sex_char in "HN04KW0": # Likely M
value += 6
elif sex_char in "EP735S": # Likely F
value += 6
value += 14 if valid_yymmdd(expiry) else 0
value += 18 if expiry_check.isdigit() and mrz_check_digit(expiry) == expiry_check else 0
value += sum(1 for c in candidate[13:19] + candidate[21:27] if c in st.digits + "OQDISBL")
value -= sum(1 for c in candidate[:9] if c == "<")
return value
windows = [line[i : i + 44] for i in range(0, len(line) - 43)]
return max(windows, key=score)
def clean_mrz_date(text):
return normalize_mrz_digits(text)
def parse_mrz_line2(line):
line = clean_mrz_line2(line)
if not line:
return {}
passport_number = clean_string(line[0:9]).replace("O", "0")
nationality = normalize_country_code(line[10:13])
dob = normalize_mrz_digits(line[13:19])
sex = get_sex(line[20])
expiry = normalize_mrz_digits(line[21:27])
dob_ok = valid_yymmdd(dob) and line[19].isdigit() and mrz_check_digit(dob) == line[19]
expiry_ok = valid_yymmdd(expiry) and line[27].isdigit() and mrz_check_digit(expiry) == line[27]
return {
"passport_number": passport_number,
"nationality": nationality,
"date_of_birth": dob if dob_ok else "",
"sex": sex,
"expiration_date": expiry if expiry_ok else "",
}
def get_country_name(code):
code = normalize_country_code(str(code))
country = pycountry.countries.get(alpha_3=code) if code else None
return country.name.upper() if country else code
def get_sex(code):
if not code:
return ""
code = str(code).upper().strip()
# Standard values
if code in ["M", "F"]:
return code
# Common OCR misreadings for Male (M)
if code in ["H", "N", "0", "4", "K", "W"]:
return "M"
# Common OCR misreadings for Female (F)
if code in ["E", "P", "7", "3", "5", "S"]:
return "F"
return ""
# ---------------------------------------------------
# FORMATTING
# ---------------------------------------------------
def parse_any_date(date_val, is_dob=True):
"""
Robustly parses a date from various formats and types.
Handles century correction for Date of Birth.
"""
if not date_val:
return None
if isinstance(date_val, (date, datetime)):
if isinstance(date_val, datetime):
date_val = date_val.date()
# Still apply century correction if it's a DOB and in the future
if is_dob and date_val > date.today():
date_val = date_val.replace(year=date_val.year - 100)
return date_val
date_str = str(date_val).strip()
if not date_str or "•" in date_str:
return None
# Remove common OCR noise but keep alphanumeric and basic separators
clean_str = re.sub(r"[^0-9A-Z/\-\s]", "", date_str.upper())
# Try parsing with various formats
parsed_date = None
# 1. Try MRZ format (YYMMDD) if it's 6 digits
digits_only = "".join(c for c in clean_str if c.isdigit())
if len(digits_only) == 6:
try:
parsed_date = datetime.strptime(digits_only, "%y%m%d").date()
# For 2-digit years, strptime uses a 69-99 -> 19xx, 00-68 -> 20xx rule.
# We override this for DOB to ensure it's not in the future.
if is_dob and parsed_date > date.today():
parsed_date = parsed_date.replace(year=parsed_date.year - 100)
return parsed_date
except:
pass
# 2. Try common human-readable formats
for fmt in ("%d/%m/%Y", "%d/%m/%y", "%d-%m-%Y", "%d-%m-%y", "%d%b%y", "%Y%m%d", "%d %b %Y"):
try:
parsed_date = datetime.strptime(date_str, fmt).date()
# If we used a 2-digit year format, apply century correction
if "%y" in fmt and is_dob and parsed_date > date.today():
parsed_date = parsed_date.replace(year=parsed_date.year - 100)
return parsed_date
except:
continue
return None
def calculate_passenger_type(dob):
"""
Calculates the Passenger Type Code (PTC) based on the date of birth.
- Infant (INF): Age < 2 years
- Child (CHD): 2 <= Age < 12 years
- Adult (ADT): Age >= 12 years
"""
parsed_dob = parse_any_date(dob, is_dob=True)
if not parsed_dob:
return "Adult", "ADT"
try:
today = date.today()
age = today.year - parsed_dob.year - ((today.month, today.day) < (parsed_dob.month, parsed_dob.day))
if age < 0:
return "Adult", "ADT"
if age < 2:
return "Infant", "INF"
if age < 12:
return "Child", "CHD"
return "Adult", "ADT"
except:
return "Adult", "ADT"
def calculate_title(sex, ptc):
"""
Calculates the Title based on gender and passenger type (PTC).
- MR: Adult Male
- MRS: Adult Female
- MS: Child/Infant Female
- MSTR: Child/Infant Male (Master)
"""
sex = str(sex).upper().strip()
ptc = str(ptc).upper().strip()
if sex == "M":
if ptc == "ADT":
return "MR"
return "MSTR"
elif sex == "F":
if ptc == "ADT":
return "MRS"
return "MS"
return "MR" # Default fallback
def format_date(raw_date, fmt="%d/%m/%Y", is_dob=True):
parsed = parse_any_date(raw_date, is_dob=is_dob)
if not parsed:
return str(raw_date) if raw_date else ""
return parsed.strftime(fmt).upper()
def format_fly_dubai(results):
rows = []
for res in results:
sex = res.get("sex", "").upper()
raw_dob = res.get("date_of_birth", "")
raw_exp = res.get("expiration_date", "")
dob_formatted = format_date(raw_dob, "%d%b%y", is_dob=True)
exp_formatted = format_date(raw_exp, "%d%b%y", is_dob=False)
_, ptc = calculate_passenger_type(raw_dob)
title = calculate_title(sex, ptc)
rows.append(
{
"Last Name": res.get("surname", ""),
"First Name and Middle Name": res.get("name", ""),
"Title": title,
"PTC": ptc,
"Gender": sex,
"Date of Birth": dob_formatted,
"Passport Number": res.get("passport_number", ""),
"Passport Nationality": (res.get("nationality") or "")[:3],
"Passport Issue Country": (res.get("nationality") or "")[:3],
"Passport Expiry Date": exp_formatted,
}
)
return pd.DataFrame(rows)
def format_iraqi(results):
rows = []
for res in results:
sex = res.get("sex", "").upper()
raw_dob = res.get("date_of_birth", "")
raw_exp = res.get("expiration_date", "")
dob_formatted = format_date(raw_dob, "%d/%m/%y", is_dob=True)
exp_formatted = format_date(raw_exp, "%d/%m/%y", is_dob=False)
ptype, ptc = calculate_passenger_type(raw_dob)
title = calculate_title(sex, ptc)
nat = (res.get("nationality") or "")[:3]
rows.append(
{
"TYPE": ptype,
"TITLE": title,
"FIRST NAME": res.get("name", ""),
"LAST NAME": res.get("surname", ""),
"DOB (DD/MM/YYYY)": dob_formatted,
"GENDER": "Male" if sex == "M" else "Female"
}
)
return pd.DataFrame(rows)
def format_fly_baghdad(results):
rows = []
for i, res in enumerate(results):
sex = res.get("sex", "").upper()
raw_dob = res.get("date_of_birth", "")
raw_exp = res.get("expiration_date", "")
dob_formatted = format_date(raw_dob, "%d/%m/%Y", is_dob=True)
exp_formatted = format_date(raw_exp, "%d/%m/%Y", is_dob=False)
_, ptc = calculate_passenger_type(raw_dob)
title = calculate_title(sex, ptc)
nat = (res.get("nationality") or "")[:3]
issuing = (res.get("issuing_country") or "")[:3]
rows.append(
{
"Sequence": i + 1,
"Pax Type": ptc,
"Title": title,
"First Name": res.get("name", ""),
"Last Name": res.get("surname", ""),
"Gender": "MALE" if sex == "M" else "FEMALE",
"DOB (dd/mm/yyyy)": dob_formatted,
"Nationality": get_country_name(nat),
"Passport Number": res.get("passport_number", ""),
"Passport Expiry (dd/mm/yyyy)": exp_formatted,
"Passport Issued Country": get_country_name(issuing)
}
)
return pd.DataFrame(rows)
# ---------------------------------------------------
# CORE EXTRACTOR
# ---------------------------------------------------
class PassportExtractor:
def __init__(self, use_gpu=USE_GPU):
# Fix SSL for model downloads
try:
ssl._create_default_https_context = ssl._create_unverified_context
except:
pass
# Initialize EasyOCR with default system paths
self.reader = easyocr.Reader(OCR_LANGUAGES, gpu=use_gpu, verbose=False)
self.has_tesseract = shutil.which("tesseract") is not None
if not self.has_tesseract:
logger.info("Tesseract not found. Skipping PassportEye and using EasyOCR.")
def crop_mrz_section(self, img_path, top_crop_ratio=0.62):
img = cv2.imread(img_path)
if img is None:
return img_path, False
height = img.shape[0]
crop_start = int(height * top_crop_ratio)
crop_start = min(max(crop_start, 0), height - 1)
mrz_img = img[crop_start:height, :]
with tempfile.NamedTemporaryFile(suffix=".jpg", delete=False) as tmp:
cropped_path = tmp.name
cv2.imwrite(cropped_path, mrz_img)
return cropped_path, True
def preprocess_for_ocr(self, img_path):
img = cv2.imread(img_path)
if img is None:
return []
variants = [img]
gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)
clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))
enhanced = clahe.apply(gray)
variants.append(enhanced)
sharp = cv2.GaussianBlur(enhanced, (0, 0), 1.0)
sharp = cv2.addWeighted(enhanced, 1.6, sharp, -0.6, 0)
variants.append(sharp)
return variants
def _find_mrz_lines(self, candidates):
for i, c in enumerate(candidates):
p_start = c.find("P<")
if 0 <= p_start <= 3:
c = c[p_start:]
for next_line in candidates[i + 1 : i + 4]:
if len(next_line) >= 35:
return c, next_line
for i in range(len(candidates) - 1):
first, second = candidates[i], candidates[i + 1]
if len(first) >= 35 and len(second) >= 35:
return first, second
return None, None
def _mrz_line_boxes(self, image):
if len(image.shape) == 3:
gray = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY)
else:
gray = image
_, bw = cv2.threshold(gray, 0, 255, cv2.THRESH_BINARY_INV + cv2.THRESH_OTSU)
kernel = cv2.getStructuringElement(cv2.MORPH_RECT, (35, 5))
dilated = cv2.dilate(bw, kernel, iterations=2)
contours, _ = cv2.findContours(
dilated, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE
)
h, w = gray.shape
boxes = []
for contour in contours:
x, y, box_w, box_h = cv2.boundingRect(contour)
if box_w > w * 0.25 and box_h > 8:
pad = 8
boxes.append(
[
max(0, x - pad),
min(w, x + box_w + pad),
max(0, y - pad),
min(h, y + box_h + pad),
]
)
boxes = sorted(boxes, key=lambda box: box[2])[-2:]
if len(boxes) == 2:
return sorted(boxes, key=lambda box: box[2])
return []
def extract_mrz_easyocr(self, img_path):
allow = st.ascii_uppercase + st.digits + "<"
l1, l2 = None, None
for image in self.preprocess_for_ocr(img_path):
boxes = self._mrz_line_boxes(image)
if not boxes:
continue
gray = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY) if len(image.shape) == 3 else image
results = self.reader.recognize(
gray,
horizontal_list=boxes,
free_list=[],
detail=0,
allowlist=allow,
batch_size=2,
)
candidates = [
"".join([c for c in r.upper().replace(" ", "") if c in allow])
for r in results
if len(r) >= 25
]
l1, l2 = self._find_mrz_lines(candidates)
if l1 and l2:
return l1, l2
return l1, l2
def extract_mrz_easyocr_slow(self, img_path):
allow = st.ascii_uppercase + st.digits + "<"
best = (None, None)
for image in self.preprocess_for_ocr(img_path):
results = self.reader.readtext(
image,
detail=0,
allowlist=allow,
paragraph=False,
width_ths=0.9,
ycenter_ths=0.7,
contrast_ths=0.05,
adjust_contrast=0.7,
)
candidates = [
"".join([c for c in r.upper().replace(" ", "") if c in allow])
for r in results
if len(r) >= 20
]
l1, l2 = self._find_mrz_lines(candidates)
if l1 and l2:
return l1, l2
if candidates:
best = self._find_mrz_lines(
[c for c in candidates if len(c) >= 30]
)
return best
def get_data(self, img_path, airline="iraqi"):
ocr_path, should_cleanup = self.crop_mrz_section(img_path)
# 1. Try PassportEye (Tesseract)
mrz_data = None
if self.has_tesseract:
try:
mrz = read_mrz(ocr_path)
if mrz and getattr(mrz, "valid_score", 0) > 20:
mrz_data = {
"surname": clean_name_field(getattr(mrz, "surname", "")),
"name": clean_name_field(getattr(mrz, "names", "")),
"passport_number": clean_string(getattr(mrz, "number", "")),
"nationality": clean_string(getattr(mrz, "nationality", "")),
"issuing_country": clean_string(getattr(mrz, "issuing_state", "")),
"date_of_birth": clean_string(
getattr(mrz, "date_of_birth", "")
),
"sex": get_sex(getattr(mrz, "sex", "")),
"expiration_date": clean_string(
getattr(mrz, "expiration_date", "")
),
"mrz_found": True,
}
except Exception as e:
logger.info(f"PassportEye failed, using EasyOCR fallback: {e}")
# 2. Try EasyOCR direct search if Tesseract failed
if not mrz_data:
l1, l2 = self.extract_mrz_easyocr(ocr_path)
if not (l1 and l2):
l1, l2 = self.extract_mrz_easyocr_slow(ocr_path)
if not (l1 and l2) and should_cleanup:
l1, l2 = self.extract_mrz_easyocr_slow(img_path)
if l1 and l2:
try:
l1, l2 = clean_mrz_line(l1), clean_mrz_line2(l2)
issuing_country, surname, given_names = parse_mrz_line1(l1)
line2_data = parse_mrz_line2(l2)
mrz_data = {
"surname": surname,
"name": given_names,
"issuing_country": issuing_country,
"passport_number": line2_data.get("passport_number", ""),
"nationality": line2_data.get("nationality", ""),
"date_of_birth": line2_data.get("date_of_birth", ""),
"sex": line2_data.get("sex", ""),
"expiration_date": line2_data.get("expiration_date", ""),
"mrz_found": True,
}
except:
pass
try:
if mrz_data:
# Final fallback for sex if still missing
if not mrz_data.get("sex"):
try:
# Quick check for gender keywords in the whole image
results = self.reader.readtext(img_path, detail=0)
text_blob = " ".join(results).upper()
if "FEMALE" in text_blob:
mrz_data["sex"] = "F"
elif "MALE" in text_blob:
mrz_data["sex"] = "M"
except:
pass
mrz_data["country"] = get_country_name(mrz_data["nationality"])
return mrz_data
return None
finally:
if should_cleanup and os.path.exists(ocr_path):
try:
os.remove(ocr_path)
except Exception as delete_error:
logger.warning(
f"Could not remove temp cropped file {ocr_path}: {delete_error}"
)
def process_pdf(self, pdf_path, progress_callback=None, airline="iraqi"):
results = []
try:
doc = fitz.open(pdf_path)
total_pages = len(doc)
for i in range(total_pages):
if progress_callback:
progress_callback((i + 1) / total_pages)
page = doc.load_page(i)
pix = page.get_pixmap(
matrix=fitz.Matrix(2, 2)
) # Better resolution for OCR
with tempfile.NamedTemporaryFile(suffix=".jpg", delete=False) as tmp:
tmp_path = tmp.name
try:
pix.save(tmp_path)
res = self.get_data(tmp_path, airline=airline)
if res:
results.append(res)
finally:
if os.path.exists(tmp_path):
try:
os.remove(tmp_path)
except Exception as delete_error:
logger.warning(
f"Could not remove temp file {tmp_path}: {delete_error}"
)
doc.close()
except Exception as e:
logger.error(f"PDF processing error: {e}")
raise e
return results