RcmEmailAutomation / app /lib /utils /classifier.py
cmoss3's picture
Replace "tie" logic with "ambiguous" logic so that daily brief manual reviews happen when there is more than 1 doc type that meets confidence threshold
189d6c8
Raw
History Blame Contribute Delete
5.51 kB
from __future__ import annotations
import logging
from dataclasses import dataclass, field
logger = logging.getLogger(__name__)
# Each doc type maps keywords to their individual signal weight (0.0–1.0).
# Higher weight = stronger evidence for that doc type when matched.
_KEYWORD_MAP: dict[str, dict[str, float]] = {
"mtr": {
"material test": 1.0,
"mill test": 1.0,
" mtr ": 1.0,
"material remarks": 0.7,
"inspection certificate": 0.7,
"inspection document": 0.7,
"certificate no.": 0.7,
"certificate number": 0.7,
"test specimen": 0.6,
"test certificate": 0.6,
"specimen": 0.4,
"product test": 0.4,
"heat test": 0.4,
"hardness test": 0.4,
"heat treatment": 0.4,
"chemical composition": 0.4,
"chemical analysis": 0.4,
"flang test": 0.4,
"flattening test": 0.4,
"flaring test": 0.4,
"material": 0.2,
# "certificate": 0.2,
},
"po": {
"purchase order": 1.0,
"sales order": 1.0,
"total sales order amount": 1.0,
"p.o.": 0.6,
" po ": 0.6,
"po date": 0.5,
"po number": 0.5,
" po#": 0.5,
" po.": 0.5,
" so#": 0.5,
" so.": 0.5,
"total due": 0.5,
},
"invoice": {
"invoice": 0.6,
"customer statement": 1.0,
"invoice #": 1.0,
"invoice date": 0.6,
"paid to": 0.5,
"payment type": 0.5,
"bill payment": 0.5,
},
"quote": {
"quotation": 1.0,
"request for quote": 1.0,
"rfq": 0.7,
"quote": 0.8,
"bid": 0.5,
},
}
# quote is classified but intentionally not routed to Dropbox
_ROUTABLE: frozenset[str] = frozenset({"po", "invoice", "mtr"})
@dataclass
class KeywordMatch:
keyword: str
weight: float
filename_hits: int
ocr_hits: int
filename_contrib: float
ocr_contrib: float
@dataclass
class ClassifyResult:
doc_type: str
reason: str
scores: dict[str, float] = field(default_factory=dict)
# Only doc types with at least one keyword hit are present.
breakdown: dict[str, list[KeywordMatch]] = field(default_factory=dict)
@dataclass
class ScoringWeights:
filename: float
ocr: float
freq_multiplier: float
min_threshold: float
def classify(
filename: str,
file_bytes: bytes,
content_type: str,
weights: ScoringWeights,
) -> ClassifyResult:
if not content_type.startswith("application/pdf"):
return ClassifyResult(doc_type="skipped", reason="non_pdf", scores={})
filename_lower = filename.lower() if filename else ""
ocr_text = _ocr_pdf(file_bytes)
scores: dict[str, float] = {}
breakdown: dict[str, list[KeywordMatch]] = {}
for doc_type, keywords in _KEYWORD_MAP.items():
score = 0.0
matches: list[KeywordMatch] = []
for keyword, kw_weight in keywords.items():
fn_hits = filename_lower.count(keyword)
ocr_hits = ocr_text.count(keyword) if ocr_text else 0
fn_contrib = 0.0
ocr_contrib = 0.0
if fn_hits > 0:
fn_contrib = weights.filename * kw_weight * (1 + (fn_hits - 1) * weights.freq_multiplier)
score += fn_contrib
if ocr_hits > 0:
ocr_contrib = weights.ocr * kw_weight * (1 + (ocr_hits - 1) * weights.freq_multiplier)
score += ocr_contrib
if fn_hits > 0 or ocr_hits > 0:
matches.append(KeywordMatch(
keyword=keyword,
weight=kw_weight,
filename_hits=fn_hits,
ocr_hits=ocr_hits,
filename_contrib=round(fn_contrib, 4),
ocr_contrib=round(ocr_contrib, 4),
))
scores[doc_type] = round(score, 4)
if matches:
breakdown[doc_type] = matches
max_score = max(scores.values(), default=0.0)
if max_score < weights.min_threshold:
return ClassifyResult(doc_type="unknown", reason="below_threshold", scores=scores, breakdown=breakdown)
above_threshold = [t for t, s in scores.items() if s >= weights.min_threshold]
if len(above_threshold) > 1:
return ClassifyResult(doc_type="ambiguous", reason="ambiguous", scores=scores, breakdown=breakdown)
winner = above_threshold[0]
if winner not in _ROUTABLE:
return ClassifyResult(doc_type=winner, reason="not_routable", scores=scores, breakdown=breakdown)
return ClassifyResult(doc_type=winner, reason="routed", scores=scores, breakdown=breakdown)
def _ocr_pdf(file_bytes: bytes) -> str:
from pdf2image import convert_from_bytes
import pytesseract
images = convert_from_bytes(file_bytes, first_page=1, last_page=2)
return "\n".join(pytesseract.image_to_string(img).lower() for img in images)