Files
Auto-Finance/app/ocr.py

454 lines
16 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
OCR 刷卡信息提取
从支付截图中识别刷卡记录(姓名、日期、金额),回填到发票数据中。
匹配策略:
1. 先按文件名匹配PDF 和图片同名)
2. 未匹配的通过金额近邻匹配
对外接口:
enrich_with_ocr(rows, directory) -> list[dict] 用 OCR 识别结果丰富发票数据
"""
import csv
import os
import re
from pathlib import Path
from . import get_logger
log = get_logger("ocr")
# ------------------------------------------------------------------
# 懒加载 OCR
# ------------------------------------------------------------------
_ocr_instance = None
def _get_ocr():
"""懒加载 PaddleOCR 实例(兼容 2.x / 3.x"""
global _ocr_instance
if _ocr_instance is not None:
return _ocr_instance
os.environ.setdefault("FLAGS_use_mkldnn", "0")
os.environ.setdefault("FLAGS_mkldnn_cache_enabled", "0")
from paddleocr import PaddleOCR
try:
_ocr_instance = PaddleOCR(use_textline_orientation=True, lang="ch")
except TypeError:
try:
_ocr_instance = PaddleOCR(lang="ch")
except TypeError:
_ocr_instance = PaddleOCR()
return _ocr_instance
# ------------------------------------------------------------------
# OCR 识别
# ------------------------------------------------------------------
def ocr_image(image_path: Path) -> list[dict]:
"""对单张图片执行 OCR返回 [{"text": str, "confidence": float}, ...]"""
ocr = _get_ocr()
texts = []
try:
results = ocr.ocr(str(image_path), cls=True)
if results and isinstance(results, list):
for page_result in results:
if not page_result:
continue
for line in page_result:
if isinstance(line, (list, tuple)) and len(line) >= 2:
_, text_info = line[0], line[1]
if isinstance(text_info, (list, tuple)) and len(text_info) >= 2:
texts.append({
"text": str(text_info[0]),
"confidence": float(text_info[1]),
})
except Exception:
try:
if hasattr(ocr, "predict"):
results = ocr.predict(str(image_path))
if results:
for result in results:
if hasattr(result, "rec_result_list"):
for line in result.rec_result_list:
t = getattr(line, "text", "") or ""
s = getattr(line, "score", 0.0) or 0.0
texts.append({"text": str(t), "confidence": float(s)})
elif isinstance(result, list):
for line in result:
if isinstance(line, (list, tuple)) and len(line) >= 2:
t = line[1][0] if isinstance(line[1], (list, tuple)) else str(line[1])
s = line[1][1] if isinstance(line[1], (list, tuple)) and len(line[1]) > 1 else 0.0
texts.append({"text": str(t), "confidence": float(s)})
except Exception as e:
log.error(f"OCR 识别失败: {e}")
return texts
def extract_card_info(texts: list[dict]) -> dict:
"""从 OCR 文本中提取刷卡信息(日期 / 金额 / 姓名)"""
info = {"刷卡日期": "", "刷卡金额": "", "人员姓名": ""}
valid = [t for t in texts if t["confidence"] > 0.5]
full_text = " ".join(t["text"] for t in valid)
if not full_text:
return info
# 日期(优先级匹配,避免误抓发票开票日期)
date_candidates = []
for pattern, priority in [
(r"记账时间[:\s]*(\d{4}[-/]\d{1,2}[-/]\d{1,2})", 10),
(r"交易时间[:\s]*(\d{4}[-/]\d{1,2}[-/]\d{1,2})", 9),
(r"刷卡日期[:\s]*(\d{4}[-/]\d{1,2}[-/]\d{1,2})", 9),
(r"日期[:\s]*(\d{4}[-/]\d{1,2}[-/]\d{1,2})", 5),
]:
for m in re.finditer(pattern, full_text):
date_candidates.append((priority, m.start(), m.group(1).replace("-", "/")))
if date_candidates:
date_candidates.sort(key=lambda x: (-x[0], x[1]))
info["刷卡日期"] = date_candidates[0][2]
# 金额
amount_candidates = []
for pattern, priority in [
(r"交易金额[:\s]*([+-]?[\d,]+\.?\d*)", 10),
(r"刷卡金额[:\s]*([+-]?[\d,]+\.?\d*)", 10),
(r"金额[:\s]*([+-]?[\d,]+\.?\d*)", 5),
]:
for m in re.finditer(pattern, full_text):
amt_str = m.group(1).replace(",", "").replace("+", "")
try:
val = float(amt_str)
if 0 < val < 999999:
amount_candidates.append((priority, m.start(), amt_str))
except ValueError:
continue
if amount_candidates:
amount_candidates.sort(key=lambda x: (-x[0], x[1]))
info["刷卡金额"] = amount_candidates[0][2]
# 姓名(排除公司/机构后缀)
EXCLUDE_SUFFIXES = ("公司", "银行", "中心", "支行", "商户", "网点", "有限", "责任")
name_candidates = []
for pattern, priority in [
(r"交易户名[:\s]*([\u4e00-\u9fff]{2,6})", 10),
(r"户名[:\s]*([\u4e00-\u9fff]{2,6})", 8),
(r"持卡人[:\s]*([\u4e00-\u9fff]{2,6})", 8),
(r"姓名[:\s]*([\u4e00-\u9fff]{2,6})", 8),
]:
for m in re.finditer(pattern, full_text):
name = m.group(1)
if not any(name.endswith(s) for s in EXCLUDE_SUFFIXES):
name_candidates.append((priority, m.start(), name))
if name_candidates:
name_candidates.sort(key=lambda x: (-x[0], x[1]))
info["人员姓名"] = name_candidates[0][2]
return info
# ------------------------------------------------------------------
# PDF 发票号提取
# ------------------------------------------------------------------
def extract_invoice_number(pdf_path: Path) -> str:
"""从 PDF 中提取发票号码"""
try:
import pdfplumber
except ImportError:
log.warning("缺少 pdfplumber跳过发票号提取")
return ""
try:
with pdfplumber.open(str(pdf_path)) as pdf_file:
page_text = ""
for page in pdf_file.pages:
page_text += page.extract_text() or ""
for pattern in [
r"发票号码[:\s]*([A-Za-z0-9]{8,20})",
r"发票代码[:\s]*([A-Za-z0-9]{10,12})",
r"号码[:\s]*([A-Za-z0-9]{8,20})",
]:
m = re.search(pattern, page_text)
if m:
return m.group(1)
except Exception as e:
log.warning(f"PDF 读取失败 ({pdf_path.name}): {e}")
return ""
# ------------------------------------------------------------------
# 图片配对
# ------------------------------------------------------------------
def _extract_amount_from_pdf(pdf_path: Path) -> float | None:
"""从 PDF 中提取价税合计金额"""
try:
import pdfplumber
with pdfplumber.open(str(pdf_path)) as pdf:
text = ""
for page in pdf.pages:
t = page.extract_text()
if t:
text += t + "\n"
m = re.search(r"价税合计.*?(小写)[¥¥]?\s*(\d+\.?\d*)", text)
if m:
return float(m.group(1))
except Exception:
pass
return None
def _extract_amount_from_image(img_path: Path) -> float | None:
"""从图片 OCR 中提取刷卡金额"""
texts = ocr_image(img_path)
if not texts:
return None
info = extract_card_info(texts)
amt_str = info.get("刷卡金额", "")
if amt_str:
try:
return float(amt_str)
except ValueError:
pass
return None
def find_image_pairs(directory: str = ".") -> list[tuple[Path, Path]]:
"""查找 PDF 和对应图片的配对
1. 先按文件名匹配PDF 和图片同名)
2. 未匹配的通过金额近邻匹配
"""
base = Path(directory)
pdfs = sorted(base.glob("*.pdf"))
image_exts = {".png", ".jpg", ".jpeg", ".bmp", ".webp"}
all_images = sorted(
f for ext in image_exts for f in base.glob(f"*{ext}")
)
# ---- Phase 1: 文件名匹配 ----
pairs: list[tuple[Path, Path]] = []
matched_pdfs: set[Path] = set()
matched_imgs: set[Path] = set()
for pdf in pdfs:
for ext in image_exts:
img = base / f"{pdf.stem}{ext}"
if img.exists():
pairs.append((pdf, img))
matched_pdfs.add(pdf)
matched_imgs.add(img)
break
unmatched_pdfs = [p for p in pdfs if p not in matched_pdfs]
unmatched_imgs = [i for i in all_images if i not in matched_imgs]
if not unmatched_pdfs or not unmatched_imgs:
return pairs
# ---- Phase 2: 金额近邻匹配 ----
if len(unmatched_pdfs) > 0 and len(unmatched_imgs) > 0:
log.info(f"文件名匹配 {len(pairs)} 组,剩余 {len(unmatched_pdfs)} 个 PDF、{len(unmatched_imgs)} 张图片,尝试金额匹配...")
pdf_amounts: dict[Path, float] = {}
for pdf in unmatched_pdfs:
amt = _extract_amount_from_pdf(pdf)
if amt is not None:
pdf_amounts[pdf] = amt
img_amounts: dict[Path, float] = {}
for img in unmatched_imgs:
amt = _extract_amount_from_image(img)
if amt is not None:
img_amounts[img] = amt
# 贪婪匹配:每张图片找金额差最小的 PDF
used_pdfs: set[Path] = set()
for img, img_amt in sorted(img_amounts.items(), key=lambda x: x[0].name):
best_pdf: Path | None = None
best_diff: float = float("inf")
for pdf, pdf_amt in pdf_amounts.items():
if pdf in used_pdfs:
continue
diff = abs(pdf_amt - img_amt)
if diff < best_diff:
best_diff = diff
best_pdf = pdf
if best_pdf is not None:
pairs.append((best_pdf, img))
used_pdfs.add(best_pdf)
log.info(f"金额匹配完成,共 {len(pairs)} 组配对")
return pairs
# ------------------------------------------------------------------
# CSV 读写
# ------------------------------------------------------------------
from .extractor import CSV_COLUMNS
def _load_csv(csv_path: Path) -> list[dict] | None:
"""读取现有 CSV 为 dict 列表,失败返回 None"""
try:
with open(csv_path, encoding="utf-8", newline="") as f:
reader = csv.DictReader(f)
fieldnames = reader.fieldnames or []
missing = [c for c in CSV_COLUMNS if c not in fieldnames]
if missing:
log.error(f"CSV 缺少必要列: {missing}")
return None
return [row for row in reader]
except FileNotFoundError:
log.error(f"CSV 文件不存在: {csv_path.name}")
return None
except Exception as e:
log.error(f"CSV 读取失败: {e}")
return None
def _save_csv(csv_path: Path, rows: list[dict]):
"""保存 CSV"""
with open(csv_path, "w", encoding="utf-8", newline="") as f:
writer = csv.DictWriter(f, fieldnames=CSV_COLUMNS)
writer.writeheader()
writer.writerows(rows)
# ------------------------------------------------------------------
# Markdown 同步
# ------------------------------------------------------------------
def save_markdown_from_csv(csv_path: Path, rows: list[dict]):
"""根据最新 CSV 数据生成 Markdown 汇总表"""
md_path = csv_path.with_suffix(".md")
columns = [
("序号", "序号"), ("发票号码", "发票号码"), ("开票日期", "开票日期"),
("项目名称", "项目名称"), ("规格型号", "规格型号"), ("价税合计", "价税合计"),
("销售方名称", "销售方名称"), ("人员姓名", "人员姓名"),
("刷卡日期", "刷卡日期"), ("公务卡号", "公务卡号"),
("刷卡金额", "刷卡金额"), ("备注", "备注"), ("工号", "工号"),
]
lines = ["# 发票信息汇总表", ""]
header = " | ".join(col[1] for col in columns)
separator = "|".join(["------" for _ in columns])
lines.append(f"| {header} |")
lines.append(f"|{separator}|")
total_price = 0.0
total_card = 0.0
for row in rows:
cells = []
for key, _ in columns:
value = row.get(key, "").strip()
if key == "价税合计" and value:
try:
total_price += float(value.replace(",", ""))
cells.append(f"¥{float(value.replace(',', '')):,.2f}")
except (ValueError, TypeError):
cells.append(value)
elif key == "刷卡金额" and value:
try:
total_card += float(value.replace(",", ""))
cells.append(f"¥{float(value.replace(',', '')):,.2f}")
except (ValueError, TypeError):
cells.append(value)
else:
cells.append(value if value else "")
lines.append("| " + " | ".join(cells) + " |")
lines.append("")
lines.append(f"**价税合计总计: ¥{total_price:,.2f}**")
lines.append(f"**刷卡金额总计: ¥{total_card:,.2f}**")
lines.append("")
with open(md_path, "w", encoding="utf-8") as f:
f.write("\n".join(lines))
log.info(f"Markdown 已同步: {md_path.name}")
# ------------------------------------------------------------------
# 主入口
# ------------------------------------------------------------------
def enrich_with_ocr(rows: list[dict], directory: str = ".") -> list[dict]:
"""用 OCR 识别结果丰富发票数据,返回更新后的行列表
rows 应包含「发票号码」列,已存在的字段不会覆盖。
"""
pairs = find_image_pairs(directory)
if not pairs:
log.warning("未找到 PDF-图片配对文件,跳过 OCR")
return rows
log.info(f"找到 {len(pairs)} 组 PDF-图片配对")
ocr_by_invoice: dict[str, dict] = {}
for idx, (pdf, img) in enumerate(pairs, 1):
inv_num = extract_invoice_number(pdf)
if not inv_num:
inv_num = pdf.stem
texts = ocr_image(img)
if not texts:
log.warning(f"OCR 未识别到文本: {img.name}")
continue
info = extract_card_info(texts)
ocr_by_invoice[inv_num] = info
# 更新行数据
updated = 0
matched = 0
for i, row in enumerate(rows):
inv_num = row.get("发票号码", "").strip()
ocr_info = ocr_by_invoice.get(inv_num)
if not ocr_info:
for key, val in ocr_by_invoice.items():
if inv_num in key or key in inv_num:
ocr_info = val
break
if ocr_info:
matched += 1
if not row.get("人员姓名", "").strip() and ocr_info["人员姓名"]:
row["人员姓名"] = ocr_info["人员姓名"]
updated += 1
if not row.get("刷卡日期", "").strip() and ocr_info["刷卡日期"]:
row["刷卡日期"] = ocr_info["刷卡日期"]
updated += 1
if not row.get("刷卡金额", "").strip() and ocr_info["刷卡金额"]:
row["刷卡金额"] = ocr_info["刷卡金额"]
updated += 1
log.info(f"OCR 完成: 匹配 {matched}/{len(rows)} 行,更新 {updated} 个字段")
return rows