import base64 import io import os import re import traceback from datetime import date from pathlib import Path import requests from fastapi import FastAPI, HTTPException from fastapi.middleware.cors import CORSMiddleware from pydantic import BaseModel from PIL import Image import numpy as np from ultralytics import YOLO from paddleocr import PaddleOCR app = FastAPI(title="PFM Product Classifier and OCR API") CLASSIFIER_WEIGHTS_GLOB = "produk-pfm-classifier-26n-*e-*.pt" CLASSIFIER_DATE_IN_NAME = re.compile( r"produk-pfm-classifier-26n-\d+e-(\d{4}-\d{2}-\d{2})\.pt$" ) def _classifier_date_from_name(path: Path) -> date | None: match = CLASSIFIER_DATE_IN_NAME.match(path.name) if not match: return None year, month, day = (int(part) for part in match.group(1).split("-")) return date(year, month, day) def latest_classifier_weights(models_dir: Path) -> Path | None: """Return the newest produk-pfm-classifier .pt weights in models/.""" if not models_dir.is_dir(): return None candidates = list(models_dir.glob(CLASSIFIER_WEIGHTS_GLOB)) if not candidates: return None def sort_key(path: Path) -> tuple[date, float]: name_date = _classifier_date_from_name(path) or date.min return (name_date, path.stat().st_mtime) return max(candidates, key=sort_key) def resolve_classifier_models_dir() -> Path | None: """Locate produk-pfm/models (env override, repo path, or Docker mount).""" env_dir = os.environ.get("CLASSIFIER_MODELS_DIR") if env_dir: path = Path(env_dir) if path.is_dir(): return path repo_root = Path(__file__).resolve().parent.parent for candidate in ( repo_root / "pfm-web-app/public/produk-pfm/models", Path("/app/pfm-web-app/public/produk-pfm/models"), Path(__file__).resolve().parent, ): if candidate.is_dir(): return candidate return None def resolve_classifier_weights_path() -> Path | None: """Resolve YOLO weights: CLASSIFIER_MODEL_PATH or newest file in models/.""" explicit = os.environ.get("CLASSIFIER_MODEL_PATH") if explicit: path = Path(explicit) if path.is_file(): return path print(f"CLASSIFIER_MODEL_PATH not found: {path}") models_dir = resolve_classifier_models_dir() if models_dir is None: return None weights = latest_classifier_weights(models_dir) if weights is None: print(f"No classifier weights matching {CLASSIFIER_WEIGHTS_GLOB} in {models_dir}") return weights # Enable CORS app.add_middleware( CORSMiddleware, allow_origins=["*"], allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) # Load models on startup print("Loading YOLO model...") yolo_model = None try: classifier_weights = resolve_classifier_weights_path() if classifier_weights is None: raise FileNotFoundError( "No produk-pfm classifier weights found. Train with train_classifier.py " "or set CLASSIFIER_MODEL_PATH / CLASSIFIER_MODELS_DIR." ) print(f"Using classifier weights: {classifier_weights}") yolo_model = YOLO(str(classifier_weights)) print("YOLO model loaded successfully.") except Exception as e: print(f"Error loading YOLO model: {e}") yolo_model = None print("Loading PaddleOCR...") try: # Use standard textline orientation detection for PaddleOCR 3.x ocr = PaddleOCR(use_textline_orientation=True, lang='en') print("PaddleOCR loaded successfully.") except Exception as e: print(f"Error loading PaddleOCR: {e}") ocr = None class ScanRequest(BaseModel): image_base64: str def clean_ocr_text(text: str) -> str: return re.sub(r'^[^\w\s./-]+|[^\w\s./-]+$', '', text).strip() EXP_KEYWORD_RE = re.compile( r'(?:exp(?:\.|ired)?|tgl(?:\s*exp)?|expiry|bbd|best\s*before|before|best|baik\s*digunakan|\bbb\b)', re.IGNORECASE, ) DD_MM_YYYY_RE = re.compile( r'(? bool: if len(val) != 8 or not val.isdigit(): return False day, month, year = int(val[0:2]), int(val[2:4]), int(val[4:8]) return 1 <= day <= 31 and 1 <= month <= 12 and 2000 <= year <= 2099 def format_ddmmyyyy(val: str) -> str: if is_valid_ddmmyyyy_digits(val): return f"{val[0:2]}/{val[2:4]}/{val[4:8]}" return val.upper() def format_ddmmyy(val: str) -> str: if len(val) == 6 and val.isdigit(): day, month = int(val[0:2]), int(val[2:4]) if 1 <= day <= 31 and 1 <= month <= 12: return f"{val[0:2]}/{val[2:4]}/{val[4:6]}" return val.upper() def line_has_exp_keyword(line: str) -> bool: if EXP_KEYWORD_RE.search(line): return True # BB05032027 — keyword directly followed by digits return bool(re.search(r'(?i)\b(?:bb|bestbefore)(?:\s*[:.-]?\s*)?\d', line)) def clean_date_line(line: str) -> str: # 1) Replace "1)" with "0" cleaned = line.replace("1)", "0") # 2) Replace "()" with "0" cleaned = cleaned.replace("()", "0") # Clean BB misrecognitions (convert B8, 8B, 88 to BB when followed by digits) cleaned = re.sub(r'\b(?:B8|8B|88)(?=\d)', 'BB', cleaned) cleaned = re.sub(r'^(?:B8|8B|88)(?=\d)', 'BB', cleaned) # Clean 012 month misrecognition (e.g. 020122027 -> 02022027) cleaned = re.sub(r'(?02\g<2>', cleaned) cleaned = re.sub(r'(?\g<2>02\g<3>\g<4>', cleaned) # Clean 112 month misrecognition (e.g. 021122027 -> 02022027) cleaned = re.sub(r'(?02\g<2>', cleaned) cleaned = re.sub(r'(?\g<2>02\g<3>\g<4>', cleaned) # Run contextual replacements for _ in range(3): # letter o/O flanked by digits or boundary -> 0 cleaned = re.compile(r'(\d)[oO](\d|\b)').sub(r'\g<1>0\g<2>', cleaned) cleaned = re.compile(r'(\b|\d)[oO](\d)').sub(r'\g<1>0\g<2>', cleaned) # letter I/i/l/| flanked by digits -> 1 cleaned = re.compile(r'(\d)[Ii|l](\d|\b)').sub(r'\g<1>1\g<2>', cleaned) cleaned = re.compile(r'(\b|\d)[Ii|l](\d)').sub(r'\g<1>1\g<2>', cleaned) # letter S/s flanked by digits -> 5 cleaned = re.compile(r'(\d)[Ss](\d|\b)').sub(r'\g<1>5\g<2>', cleaned) cleaned = re.compile(r'(\b|\d)[Ss](\d)').sub(r'\g<1>5\g<2>', cleaned) # letter Z/z flanked by digits -> 2 cleaned = re.compile(r'(\d)[Zz](\d|\b)').sub(r'\g<1>2\g<2>', cleaned) cleaned = re.compile(r'(\b|\d)[Zz](\d)').sub(r'\g<1>2\g<2>', cleaned) # letter B flanked by digits -> 8 cleaned = re.compile(r'(\d)B(\d|\b)').sub(r'\g<1>8\g<2>', cleaned) cleaned = re.compile(r'(\b|\d)B(\d)').sub(r'\g<1>8\g<2>', cleaned) return cleaned def extract_expired_date(text_lines): """Return (formatted_date, line_index, source_line). Prioritises BB/EXP + DDMMYYYY or DD MM YYYY.""" if not text_lines: return None, None, None cleaned_lines = [clean_date_line(line) for line in text_lines] def pick(match, idx, cleaned_line, formatter=None): raw = match.group(0) original_line = text_lines[idx].strip() if match.lastindex and match.lastindex >= 3: formatted = f"{match.group(1)}/{match.group(2)}/{match.group(3)}" elif match.lastindex and match.lastindex >= 1 and match.group(1).isdigit(): digits = match.group(1) if len(digits) == 8: formatted = format_ddmmyyyy(digits) elif len(digits) == 6: formatted = format_ddmmyy(digits) else: formatted = digits elif formatter: formatted = formatter(raw) else: formatted = raw.strip().upper() return formatted, idx, original_line # 1) BB/EXP keyword lines — compact DDMMYYYY (e.g. BB05032027, EXP 05032027) for idx, line in enumerate(cleaned_lines): if not line_has_exp_keyword(line): continue match = BB_ATTACHED_DATE_RE.search(line) or DDMMYYYY_RE.search(line) if match: return pick(match, idx, line) # 2) BB/EXP keyword lines — spaced DD MM YYYY (e.g. BB 05 03 2027) for idx, line in enumerate(cleaned_lines): if not line_has_exp_keyword(line): continue match = DD_MM_YYYY_RE.search(line) if match: return pick(match, idx, line) # 3) Keyword + 6–8 digit run (BB05032027 via keyword_digits) for idx, line in enumerate(cleaned_lines): match = KEYWORD_DIGITS_RE.search(line) if match: digits = match.group(1) if len(digits) == 8 and is_valid_ddmmyyyy_digits(digits): return format_ddmmyyyy(digits), idx, text_lines[idx].strip() if len(digits) == 6: return format_ddmmyy(digits), idx, text_lines[idx].strip() # 3.5) BB/EXP keyword lines — lenient check for unclear/noisy date formats (e.g. BB 02J 132027) for idx, line in enumerate(cleaned_lines): if not line_has_exp_keyword(line): continue match = LENIENT_DATE_RE.search(line) if match: return pick(match, idx, line) # 4) Any line — spaced DD MM YYYY for idx, line in enumerate(cleaned_lines): match = DD_MM_YYYY_RE.search(line) if match: return pick(match, idx, line) # 5) Any line — compact DDMMYYYY (skip likely SKU: same line has 8-digit product code context) for idx, line in enumerate(cleaned_lines): for match in DDMMYYYY_RE.finditer(line): digits = f"{match.group(1)}{match.group(2)}{match.group(3)}" if is_valid_ddmmyyyy_digits(digits): # Skip if this 8-digit block is the only digits and looks like SKU on label top if re.search(r'\b\d{8}\b', line) and not line_has_exp_keyword(line): if re.search(r'(?:nugget|chicken|fiesta|champ|okey|akumo|frozen|gr)', line, re.I): continue return format_ddmmyyyy(digits), idx, text_lines[idx].strip() # 6) Legacy patterns (slashes, month names, etc.) date_patterns = [ r'\b\d{2}[-./]\d{2}[-./]\d{2,4}\b', r'\b\d{4}[-./]\d{2}[-./]\d{2}\b', r'\b\d{2}\s+(?:JAN|FEB|MAR|APR|MAY|JUN|JUL|AUG|SEP|OCT|NOV|DEC)[a-zA-Z]*\s+\d{2,4}\b', ] for idx, line in enumerate(cleaned_lines): if not line_has_exp_keyword(line): continue for pat in date_patterns: match = re.search(pat, line, re.IGNORECASE) if match: return match.group(0).upper(), idx, text_lines[idx].strip() return None, None, None def line_contains_expired_date(line: str, expired_date: str) -> bool: if not line or not expired_date: return False digits_only = re.sub(r"\D", "", expired_date) line_digits = re.sub(r"\D", "", line) if len(digits_only) >= 6 and digits_only in line_digits: return True compact = expired_date.replace("/", "") return compact in line.replace(" ", "") or expired_date in line def find_expired_crop_index(text_lines, expired_idx, expired_date, polys_len): """Pick OCR box index for cropping; prefer the line that actually contains the date.""" if not expired_date or polys_len <= 0: return None if ( expired_idx is not None and expired_idx < polys_len and expired_idx < len(text_lines) and line_contains_expired_date(text_lines[expired_idx], expired_date) ): return expired_idx keyword_match = None for idx, line in enumerate(text_lines): if idx >= polys_len: break if not line_contains_expired_date(line, expired_date): continue if line_has_exp_keyword(line): return idx if keyword_match is None: keyword_match = idx if keyword_match is not None: return keyword_match if expired_idx is not None and expired_idx < polys_len: return expired_idx return None def ocr_coordinate_image(res_entry, fallback_image: Image.Image) -> Image.Image: """Image in the same pixel space as rec_polys (after doc orientation + unwarping).""" dpr = res_entry.get("doc_preprocessor_res") or {} output_arr = dpr.get("output_img") if output_arr is not None: return Image.fromarray(np.asarray(output_arr)).convert("RGB") return fallback_image def ocr_text_polys(res_entry): """Recognition polygons — 1:1 aligned with rec_texts.""" return res_entry.get("rec_polys") or res_entry.get("dt_polys") or [] def extract_sku(text_lines): # SKU is usually an 8-digit number (e.g. 12010111) for line in text_lines: match = re.search(r'\b(\d{8})\b', line) if match: return match.group(1) # Try finding 7-9 digit numbers for line in text_lines: match = re.search(r'\b(\d{7,9})\b', line) if match: return match.group(1) return None def crop_poly_region(image, poly, padding=8): x_coords = [float(p[0]) for p in poly] y_coords = [float(p[1]) for p in poly] x_min = max(0, int(min(x_coords))) y_min = max(0, int(min(y_coords))) x_max = min(image.width, int(max(x_coords))) y_max = min(image.height, int(max(y_coords))) crop_left = max(0, x_min - padding) crop_top = max(0, y_min - padding) crop_right = min(image.width, x_max + padding) crop_bottom = min(image.height, y_max + padding) if crop_right <= crop_left or crop_bottom <= crop_top: return None cropped = image.crop((crop_left, crop_top, crop_right, crop_bottom)) buffered = io.BytesIO() cropped.save(buffered, format="JPEG") return "data:image/jpeg;base64," + base64.b64encode(buffered.getvalue()).decode("utf-8") def extract_product_name(text_lines, classified_name=None): keywords = ['NUGGET', 'CHICKEN', 'CHAMP', 'FIESTA', 'AKUMO', 'ASIMO', 'OKEY', 'FRIES', 'BURGER', 'SAUSAGE', 'SOSIS', 'KARAGE', 'SIOMAY', 'BUMBU', 'RACIK'] matches = [] for line in text_lines: line_upper = line.upper() if any(kw in line_upper for kw in keywords): cleaned = re.sub(r'\b\d{8}\b', '', line) cleaned = re.sub(r'(?:exp|expired|tgl|expiry|bbd|before)[^ \n]*', '', cleaned, flags=re.IGNORECASE) cleaned = re.sub(r'\b\d{2}[-./]\d{2}[-./]\d{2,4}\b', '', cleaned) cleaned = cleaned.strip() if len(cleaned) > 3: matches.append(cleaned) if matches: return max(matches, key=len) if classified_name: return classified_name candidate_lines = [l for l in text_lines if not re.match(r'^\d+$', l) and len(l) > 3] if candidate_lines: return max(candidate_lines, key=len) return "Unknown Product" @app.post("/classify-ocr") async def classify_ocr(payload: ScanRequest): try: # Decode image from PIL import ImageOps img_data = base64.b64decode(payload.image_base64.split(",")[-1]) raw_image = Image.open(io.BytesIO(img_data)) image = ImageOps.exif_transpose(raw_image).convert("RGB") # 1. Run YOLO Classification classification_result = {} top1_name = None if yolo_model: results = yolo_model(image) probs = results[0].probs top1_idx = probs.top1 top1_conf = float(probs.top1conf) top1_name = results[0].names[top1_idx] all_probs = [] for idx, val in enumerate(probs.data): all_probs.append({ "name": results[0].names[idx], "confidence": float(val) }) all_probs.sort(key=lambda x: x["confidence"], reverse=True) classification_result = { "top1_name": top1_name, "top1_confidence": top1_conf, "all_probabilities": all_probs } else: classification_result = { "error": "YOLO model not loaded" } # 2. Run PaddleOCR ocr_result = {} if ocr: img_arr = np.array(image) # Use predict method and convert generator to list res_list = list(ocr.predict(img_arr)) text_lines = [] if res_list and len(res_list) > 0: text_lines = res_list[0].get("rec_texts", []) sku = extract_sku(text_lines) expired_date, expired_idx, expired_source_line = extract_expired_date(text_lines) product_name = extract_product_name(text_lines, top1_name) res_entry = res_list[0] if res_list else {} coord_image = ocr_coordinate_image(res_entry, image) text_polys = ocr_text_polys(res_entry) crop_idx = find_expired_crop_index( text_lines, expired_idx, expired_date, len(text_polys) ) # Create visual OCR image with bounding boxes vis_image_b64 = None try: vis_image = coord_image.copy() from PIL import ImageDraw, ImageFont draw = ImageDraw.Draw(vis_image) try: font = ImageFont.load_default() except: font = None for idx, poly in enumerate(text_polys): is_expired = (crop_idx is not None and idx == crop_idx) pts = [(float(p[0]), float(p[1])) for p in poly] if is_expired: color = (245, 158, 11) # Amber label = "EXP" else: color = (13, 148, 136) # Teal label = "TEXT" draw.polygon(pts, outline=color, width=3) x0, y0 = pts[0] label_w = 32 if label == "EXP" else 38 draw.rectangle([x0, y0 - 15, x0 + label_w, y0], fill=color) if font: draw.text((x0 + 4, y0 - 14), label, fill=(255, 255, 255), font=font) else: draw.text((x0 + 4, y0 - 14), label, fill=(255, 255, 255)) buffered = io.BytesIO() vis_image.save(buffered, format="JPEG") vis_image_b64 = "data:image/jpeg;base64," + base64.b64encode(buffered.getvalue()).decode("utf-8") except Exception as draw_err: print(f"Error drawing visual OCR: {draw_err}") traceback.print_exc() # Crop expired date OCR region for summary verification expired_date_crop_b64 = None try: if crop_idx is not None and crop_idx < len(text_polys): expired_date_crop_b64 = crop_poly_region(coord_image, text_polys[crop_idx]) except Exception as crop_err: print(f"Error cropping expired date image: {crop_err}") traceback.print_exc() # 3. Call Spotting API spotting_image_b64 = None try: img_b64_only = payload.image_base64.split(",")[-1] spotting_payload = { "file": img_b64_only, "matchHistoryJob": False, "useLayoutDetection": False, "fileType": 1, "useDocUnwarping": False, "useDocOrientationClassify": False, "promptLabel": "spotting" } spotting_url = "http://localhost:8090/layout-parsing" spotting_resp = requests.post(spotting_url, json=spotting_payload, timeout=60) if spotting_resp.status_code == 200: spotting_data = spotting_resp.json() if spotting_data.get("errorCode") == 0: layout_results = spotting_data.get("result", {}).get("layoutParsingResults", []) if layout_results: page0 = layout_results[0] out_imgs = page0.get("outputImages", {}) spotting_img = out_imgs.get("spotting_res_img") if spotting_img: spotting_image_b64 = "data:image/jpeg;base64," + spotting_img else: print(f"Spotting API error: {spotting_resp.text}") except Exception as spotting_err: print(f"Error calling spotting API: {spotting_err}") traceback.print_exc() ocr_result = { "text_lines": text_lines, "extracted_product_name": product_name, "extracted_sku": sku, "extracted_expired_date": expired_date, "expired_line_index": crop_idx, "expired_source_line": expired_source_line, "expired_date_crop_base64": expired_date_crop_b64, "vis_image_base64": vis_image_b64, "spotting_image_base64": spotting_image_b64 } else: ocr_result = { "error": "PaddleOCR not loaded" } return { "classification": classification_result, "ocr": ocr_result } except Exception as e: traceback.print_exc() raise HTTPException(status_code=500, detail=str(e)) if __name__ == "__main__": import uvicorn uvicorn.run(app, host="0.0.0.0", port=8120)