Initial commit: Auto-Finance 财务报销自动化系统
This commit is contained in:
454
app/ocr.py
Normal file
454
app/ocr.py
Normal file
@@ -0,0 +1,454 @@
|
||||
"""
|
||||
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
|
||||
Reference in New Issue
Block a user