Update: Refactor to absolute imports, fix QPS limits, and enhance Gitea sync with SQLite support
This commit is contained in:
+127
-4
@@ -9,7 +9,7 @@ import base64
|
||||
import requests
|
||||
from typing import Dict, Optional, Union, List
|
||||
|
||||
from ..utils.log_utils import get_logger
|
||||
from app.core.utils.log_utils import get_logger
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
@@ -229,10 +229,19 @@ class BaiduOCRClient:
|
||||
logger.debug(f"百度OCR API返回结果: {result}")
|
||||
|
||||
if 'error_code' in result:
|
||||
error_code = result.get('error_code')
|
||||
error_msg = result.get('error_msg', '未知错误')
|
||||
|
||||
# 如果是 QPS 限制 (18),增加延迟并重试
|
||||
if error_code == 18:
|
||||
wait_time = self.retry_delay * (attempt + 1)
|
||||
logger.warning(f"触发 QPS 限制,将在 {wait_time} 秒后重试 (尝试 {attempt+1}/{self.max_retries})")
|
||||
time.sleep(wait_time)
|
||||
continue
|
||||
|
||||
logger.error(f"百度OCR API错误: {error_msg}")
|
||||
# 如果是授权错误,尝试刷新令牌
|
||||
if result.get('error_code') in [110, 111]: # 授权相关错误码
|
||||
if error_code in [110, 111]: # 授权相关错误码
|
||||
logger.info("尝试刷新访问令牌...")
|
||||
self.token_manager.refresh_token()
|
||||
return None
|
||||
@@ -292,8 +301,18 @@ class BaiduOCRClient:
|
||||
if response.status_code == 200:
|
||||
result = response.json()
|
||||
if 'error_code' in result:
|
||||
logger.warning(f"通用识别错误: {result.get('error_msg')}")
|
||||
if result.get('error_code') in (110, 111):
|
||||
error_code = result.get('error_code')
|
||||
error_msg = result.get('error_msg', '未知错误')
|
||||
|
||||
# 如果是 QPS 限制 (18),增加延迟并重试
|
||||
if error_code == 18:
|
||||
wait_time = self.retry_delay * (attempt + 1)
|
||||
logger.warning(f"触发 QPS 限制 (General),将在 {wait_time} 秒后重试 (尝试 {attempt+1}/{self.max_retries})")
|
||||
time.sleep(wait_time)
|
||||
continue
|
||||
|
||||
logger.warning(f"通用识别错误: {error_msg}")
|
||||
if error_code in (110, 111):
|
||||
self.token_manager.refresh_token()
|
||||
return None
|
||||
words_list = result.get('words_result') or []
|
||||
@@ -306,6 +325,110 @@ class BaiduOCRClient:
|
||||
time.sleep(self.retry_delay * (2 ** attempt))
|
||||
logger.error("通用识别失败")
|
||||
return None
|
||||
|
||||
def recognize_ticket(self, image_data: Union[str, bytes]) -> Optional[List[Dict]]:
|
||||
"""通用票务识别:用于精准捕获表头抬头、日期等关键信息。
|
||||
使用用户指定的 URL: https://aip.baidubce.com/rest/2.0/ocr/v1/general_ocr
|
||||
|
||||
Returns:
|
||||
[{'words': '...'}, ...]
|
||||
"""
|
||||
access_token = self.token_manager.get_token()
|
||||
if not access_token:
|
||||
logger.error("无法获取访问令牌,无法进行票务识别")
|
||||
return None
|
||||
|
||||
if isinstance(image_data, str):
|
||||
image_data = self.read_image(image_data)
|
||||
if image_data is None:
|
||||
return None
|
||||
|
||||
# 使用用户指定的通用票务识别接口
|
||||
url = self.config.get('API', 'ticket_ocr_url',
|
||||
fallback='https://aip.baidubce.com/rest/2.0/ocr/v1/general_ocr')
|
||||
url = f"{url}?access_token={access_token}"
|
||||
image_base64 = base64.b64encode(image_data).decode('utf-8')
|
||||
|
||||
payload = {
|
||||
'image': image_base64,
|
||||
'detect_direction': 'true', # 开启方向检测,应对旋转的图片
|
||||
}
|
||||
headers = {'Content-Type': 'application/x-www-form-urlencoded',
|
||||
'Accept': 'application/json'}
|
||||
|
||||
for attempt in range(self.max_retries):
|
||||
try:
|
||||
response = requests.post(url, data=payload, headers=headers,
|
||||
timeout=self.timeout)
|
||||
if response.status_code == 200:
|
||||
result = response.json()
|
||||
if 'error_code' in result:
|
||||
error_code = result.get('error_code')
|
||||
error_msg = result.get('error_msg', '未知错误')
|
||||
|
||||
# 如果是 QPS 限制 (18),增加延迟并重试
|
||||
if error_code == 18:
|
||||
wait_time = self.retry_delay * (attempt + 1)
|
||||
logger.warning(f"触发 QPS 限制 (Ticket),将在 {wait_time} 秒后重试 (尝试 {attempt+1}/{self.max_retries})")
|
||||
time.sleep(wait_time)
|
||||
continue
|
||||
|
||||
# 如果是权限错误 (6),记录详细日志告知用户
|
||||
if error_code == 6:
|
||||
logger.error("权限错误 (6): 您的百度 API Key 未开启【通用卡证票据识别】(General OCR) 服务。请前往百度云控制台手动开启该服务,否则无法精准识别表头。")
|
||||
|
||||
logger.warning(f"票务识别错误: {error_msg}")
|
||||
if error_code in (110, 111):
|
||||
self.token_manager.refresh_token()
|
||||
return None
|
||||
|
||||
# 1. 优先处理 results 结构 (通用卡证票据识别的新结构)
|
||||
if 'results' in result and isinstance(result['results'], dict):
|
||||
extracted = []
|
||||
# 通常只有一个结果 "0"
|
||||
for res_id, res_content in result['results'].items():
|
||||
if not isinstance(res_content, dict): continue
|
||||
for k, v_list in res_content.items():
|
||||
if isinstance(v_list, list):
|
||||
for item in v_list:
|
||||
if isinstance(item, dict) and 'words' in item:
|
||||
words_val = item['words']
|
||||
if isinstance(words_val, list):
|
||||
for w in words_val:
|
||||
extracted.append({'words': f"{k}: {w}"})
|
||||
else:
|
||||
extracted.append({'words': f"{k}: {words_val}"})
|
||||
return extracted
|
||||
|
||||
# 2. 通用票据/票务识别通常也返回 words_result
|
||||
words_list = result.get('words_result')
|
||||
|
||||
# 如果返回的是结构化字典,提取其中的文字行
|
||||
if isinstance(words_list, dict):
|
||||
extracted = []
|
||||
for k, v in words_list.items():
|
||||
if isinstance(v, dict) and 'words' in v:
|
||||
words_val = v['words']
|
||||
if isinstance(words_val, list):
|
||||
for w in words_val:
|
||||
extracted.append({'words': f"{k}: {w}"})
|
||||
else:
|
||||
extracted.append({'words': f"{k}: {words_val}"})
|
||||
elif isinstance(v, str):
|
||||
extracted.append({'words': f"{k}: {v}"})
|
||||
return extracted
|
||||
|
||||
return words_list or []
|
||||
|
||||
logger.warning(f"票务识别请求失败 (尝试 {attempt+1}): {response.text[:200]}")
|
||||
except Exception as e:
|
||||
logger.warning(f"票务识别异常 (尝试 {attempt+1}): {e}")
|
||||
|
||||
if attempt < self.max_retries - 1:
|
||||
time.sleep(self.retry_delay * (2 ** attempt))
|
||||
|
||||
logger.error("票务识别失败")
|
||||
return None
|
||||
|
||||
def get_excel_result(self, request_id_or_result: Union[str, Dict]) -> Optional[bytes]:
|
||||
"""
|
||||
|
||||
@@ -16,6 +16,14 @@ from typing import List, Optional, Tuple
|
||||
SUPPLIER_KEYWORDS = (
|
||||
"供货单", "供应商", "供货方", "批发", "酒行", "商行",
|
||||
"经销", "专卖店", "配送单", "送货单", "采购单", "订单",
|
||||
"销售单", "出库单", "经营部", "商贸", "有限公司", "发证机构",
|
||||
)
|
||||
|
||||
# 供应商名称中需要剔除的冗余噪声词
|
||||
SUPPLIER_NOISE_WORDS = (
|
||||
"标题", "采购单", "销售单", "入库单", "出库单", "送货单",
|
||||
"单据", "订单", "配送单", "清单", "供货单", "预览", "详情",
|
||||
"打印", "副本", "记账", "联", "存根",
|
||||
)
|
||||
|
||||
# 总金额关键词(命中后取该行最近一个金额数字)
|
||||
@@ -24,6 +32,11 @@ AMOUNT_KEYWORDS = (
|
||||
"合计", "总计", "总计金额",
|
||||
)
|
||||
|
||||
# 日期关键词(用于辅助定位日期)
|
||||
DATE_KEYWORDS = (
|
||||
"单据日期", "下单时间", "日期", "时间", "开单日期", "制单日期", "打印时间", "业务日期", "下单日期",
|
||||
)
|
||||
|
||||
# 日期正则:4 种格式
|
||||
DATE_PATTERNS = [
|
||||
# 2026年07月17日 / 2026年7月17日
|
||||
@@ -52,6 +65,16 @@ class OrderMetadata:
|
||||
return asdict(self)
|
||||
|
||||
|
||||
# 排除关键词:包含这些词的行绝对不是供应商抬头
|
||||
EXCLUDE_KEYWORDS = (
|
||||
"购货单位", "客户名称", "收货地址", "联系电话", "经手人",
|
||||
"地址", "电话", "传真", "邮编", "网址", "开户行", "账号",
|
||||
"税号", "业务员", "联系人", "单据编号", "流水号", "页码",
|
||||
"四川省", "成都市", "武侯区", "双流区", "高新区", "金牛区",
|
||||
"成华区", "锦江区", "龙泉驿区", "青羊区", "新都区", "温江区",
|
||||
"街道", "社区", "楼", "号", "室", "路", "巷", "段",
|
||||
)
|
||||
|
||||
class OrderMetadataExtractor:
|
||||
"""单据元信息识别器。"""
|
||||
|
||||
@@ -61,20 +84,35 @@ class OrderMetadataExtractor:
|
||||
# 供应商名最大长度
|
||||
MAX_SUPPLIER_LEN = 50
|
||||
|
||||
def extract(self, ocr_text: str, ocr_rows: Optional[List[List[str]]] = None) -> OrderMetadata:
|
||||
def extract(self, ocr_text: str, ocr_rows: Optional[List[List[str]]] = None, general_text: Optional[str] = None) -> OrderMetadata:
|
||||
"""从 OCR 原始文本和/或解析后的二维数组提取三字段。
|
||||
|
||||
Args:
|
||||
ocr_text: OCR 原始字符串全文(带换行)
|
||||
ocr_rows: 解析后的二维数组(可选),用于 row-level 精确匹配
|
||||
ocr_text: OCR 原始字符串全文(从表格 OCR 提取)
|
||||
ocr_rows: 解析后的二维数组(可选)
|
||||
general_text: 从通用票据识别接口 (/v1/general_ocr) 提取的文本,优先级最高
|
||||
|
||||
Returns:
|
||||
OrderMetadata
|
||||
"""
|
||||
text = ocr_text or ''
|
||||
supplier, raw_supplier = self._extract_supplier(text, ocr_rows)
|
||||
bill_date = self._extract_bill_date(text)
|
||||
total_amount = self._extract_total_amount(text, ocr_rows)
|
||||
# 优先使用通用票据识别的文本进行供应商和日期提取
|
||||
supplier_source = general_text if general_text else ocr_text
|
||||
date_source = general_text if general_text else ocr_text
|
||||
|
||||
# 金额通常在表格内,优先使用 ocr_text
|
||||
amount_source = ocr_text or general_text or ''
|
||||
|
||||
supplier, raw_supplier = self._extract_supplier(supplier_source, ocr_rows)
|
||||
bill_date = self._extract_bill_date(date_source)
|
||||
total_amount = self._extract_total_amount(amount_source, ocr_rows)
|
||||
|
||||
# 如果通用票据识别没抓到供应商,尝试用表格 OCR 的文本补位
|
||||
if not supplier and general_text and ocr_text:
|
||||
supplier, raw_supplier = self._extract_supplier(ocr_text, ocr_rows)
|
||||
|
||||
# 如果通用票据识别没抓到日期,尝试用表格 OCR 的文本补位
|
||||
if not bill_date and general_text and ocr_text:
|
||||
bill_date = self._extract_bill_date(ocr_text)
|
||||
|
||||
return OrderMetadata(
|
||||
supplier=supplier,
|
||||
@@ -95,9 +133,15 @@ class OrderMetadataExtractor:
|
||||
cleaned = self._clean_supplier_line(line)
|
||||
if not cleaned:
|
||||
continue
|
||||
|
||||
# 严格排除:购货单位、地址、电话等干扰项
|
||||
if any(k in cleaned for k in EXCLUDE_KEYWORDS):
|
||||
continue
|
||||
|
||||
for kw in SUPPLIER_KEYWORDS:
|
||||
if kw in cleaned:
|
||||
return self._truncate(cleaned), cleaned
|
||||
final_name = self._final_cleanup_supplier(cleaned)
|
||||
return self._truncate(final_name), cleaned
|
||||
|
||||
# 2) 兜底:取顶部第一个"含中文且无数字行号"且长度 ≥ 4 的非空行
|
||||
# 但排除"纯日期行"(避免把日期当供应商)
|
||||
@@ -105,6 +149,11 @@ class OrderMetadataExtractor:
|
||||
cleaned = self._clean_supplier_line(line)
|
||||
if not cleaned:
|
||||
continue
|
||||
|
||||
# 兜底也要严格排除地址和电话行
|
||||
if any(k in cleaned for k in EXCLUDE_KEYWORDS):
|
||||
continue
|
||||
|
||||
if not (re.search(r'[\u4e00-\u9fa5]', cleaned) and len(cleaned) >= 4):
|
||||
continue
|
||||
# 排除:纯日期(YYYY-MM-DD / YYYY/MM/DD / YYYY年MM月DD日 / YYYYMMDD)
|
||||
@@ -117,16 +166,63 @@ class OrderMetadataExtractor:
|
||||
# 排除:以"单据"开头的行
|
||||
if re.match(r'^\s*单据[::]?', cleaned):
|
||||
continue
|
||||
return self._truncate(cleaned), cleaned
|
||||
|
||||
final_name = self._final_cleanup_supplier(cleaned)
|
||||
if final_name:
|
||||
return self._truncate(final_name), cleaned
|
||||
|
||||
return '', ''
|
||||
|
||||
@staticmethod
|
||||
def _final_cleanup_supplier(s: str) -> str:
|
||||
"""最后的供应商名称深度清理:去除“标题”、“采购单”等噪声。"""
|
||||
if not s:
|
||||
return ''
|
||||
|
||||
# 1. 统一处理全角和常见符号
|
||||
s = s.replace(':', ':').replace('(', '(').replace(')', ')')
|
||||
|
||||
# 2. 去除“标题:”或“名称:”这类前缀
|
||||
s = re.sub(r'^(标题|名称|供应商|供货方|单位|商户)[:\s]*', '', s)
|
||||
|
||||
# 3. 循环去除噪声词
|
||||
noise_sorted = sorted(SUPPLIER_NOISE_WORDS, key=len, reverse=True)
|
||||
|
||||
changed = True
|
||||
while changed:
|
||||
original = s
|
||||
for noise in noise_sorted:
|
||||
# 简单替换所有匹配的噪声词
|
||||
s = s.replace(noise, '')
|
||||
|
||||
# 去除括号及其中间的噪声词(如 (采购单) )
|
||||
s = re.sub(r'\(\s*\)', '', s)
|
||||
|
||||
# 去除首尾残留标点和空白
|
||||
s = re.sub(r'^[::\s\-_\|\.\(\)]+', '', s)
|
||||
s = re.sub(r'[::\s\-_\|\.\(\)]+$', '', s)
|
||||
|
||||
changed = (s != original)
|
||||
|
||||
return s.strip()
|
||||
|
||||
@staticmethod
|
||||
def _clean_supplier_line(line: str) -> str:
|
||||
"""清理一行文本:去前后空白、去首尾日期/编号/电话/标点。"""
|
||||
s = line.strip()
|
||||
if not s:
|
||||
return ''
|
||||
|
||||
# 统一全角冒号
|
||||
s = s.replace(':', ':')
|
||||
|
||||
# 如果是“标题: 新双利采购单”,保留冒号后面的部分进行初步处理
|
||||
if ':' in s:
|
||||
parts = s.split(':', 1)
|
||||
# 如果冒号前面是“标题”、“名称”等词,则取后面
|
||||
if any(k in parts[0] for k in ("标题", "名称", "供应商", "商户")):
|
||||
s = parts[1].strip()
|
||||
|
||||
# 去行首日期/编号前缀
|
||||
s = re.sub(r'^[\s\d\-\.\/年月日::]+', '', s)
|
||||
# 去行尾标点
|
||||
@@ -142,13 +238,54 @@ class OrderMetadataExtractor:
|
||||
|
||||
def _extract_bill_date(self, text: str) -> str:
|
||||
"""提取单据日期,标准化为 YYYYMMDD。"""
|
||||
if not text:
|
||||
return ''
|
||||
|
||||
lines = text.splitlines()
|
||||
|
||||
# 1. 优先尝试关键词定位逻辑(如用户建议:搜索“时间”或“日期”)
|
||||
for i, line in enumerate(lines):
|
||||
line_clean = line.strip().replace(':', ':')
|
||||
|
||||
# 检查是否包含日期相关关键词
|
||||
for kw in DATE_KEYWORDS:
|
||||
if kw in line_clean:
|
||||
# a) 尝试在当前行找日期
|
||||
for pat in DATE_PATTERNS:
|
||||
m = pat.search(line_clean)
|
||||
if m:
|
||||
y, mo, d = m.group(1), m.group(2), m.group(3)
|
||||
if self._is_valid_date(y, mo, d):
|
||||
return f"{y}{int(mo):02d}{int(d):02d}"
|
||||
|
||||
# b) 如果当前行没找到,尝试在下一行找(针对表格 OCR 错位情况)
|
||||
if i + 1 < len(lines):
|
||||
next_line = lines[i+1].strip()
|
||||
for pat in DATE_PATTERNS:
|
||||
m = pat.search(next_line)
|
||||
if m:
|
||||
y, mo, d = m.group(1), m.group(2), m.group(3)
|
||||
if self._is_valid_date(y, mo, d):
|
||||
return f"{y}{int(mo):02d}{int(d):02d}"
|
||||
|
||||
# 2. 兜底策略:全文正则匹配
|
||||
for pat in DATE_PATTERNS:
|
||||
m = pat.search(text)
|
||||
if m:
|
||||
# 优先找包含 202 开头的年份(更像当前日期)
|
||||
matches = pat.finditer(text)
|
||||
for m in matches:
|
||||
y, mo, d = m.group(1), m.group(2), m.group(3)
|
||||
if self._is_valid_date(y, mo, d):
|
||||
return f"{y}{int(mo):02d}{int(d):02d}"
|
||||
return ''
|
||||
# 如果年份以 202 开头,优先返回
|
||||
if y.startswith('202'):
|
||||
return f"{y}{int(mo):02d}{int(d):02d}"
|
||||
# 否则记录下来作为候选
|
||||
last_valid = f"{y}{int(mo):02d}{int(d):02d}"
|
||||
|
||||
# 如果没有 202 开头的,返回最后一个有效的
|
||||
try:
|
||||
return last_valid
|
||||
except NameError:
|
||||
return ''
|
||||
|
||||
@staticmethod
|
||||
def _is_valid_date(y: str, m: str, d: str) -> bool:
|
||||
|
||||
@@ -10,8 +10,8 @@ import base64
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Dict, List, Optional, Tuple, Callable
|
||||
|
||||
from ..utils.log_utils import get_logger
|
||||
from ..utils.file_utils import (
|
||||
from app.core.utils.log_utils import get_logger
|
||||
from app.core.utils.file_utils import (
|
||||
ensure_dir,
|
||||
get_file_extension,
|
||||
get_files_by_extensions,
|
||||
@@ -20,6 +20,7 @@ from ..utils.file_utils import (
|
||||
load_json,
|
||||
save_json
|
||||
)
|
||||
from app.config.settings import ConfigManager
|
||||
from .baidu_ocr import BaiduOCRClient
|
||||
|
||||
logger = get_logger(__name__)
|
||||
@@ -102,14 +103,16 @@ class OCRProcessor:
|
||||
OCR处理器,负责协调OCR识别和结果处理
|
||||
"""
|
||||
|
||||
def __init__(self, config):
|
||||
def __init__(self, config: Optional[ConfigManager] = None):
|
||||
"""
|
||||
初始化OCR处理器
|
||||
|
||||
Args:
|
||||
config: 配置信息
|
||||
config: 配置管理器
|
||||
"""
|
||||
self.config = config
|
||||
self.config = config or ConfigManager()
|
||||
self.ocr_client = None
|
||||
self._ensure_ocr_client()
|
||||
|
||||
# 修复ConfigParser对象没有get_path方法的问题
|
||||
try:
|
||||
@@ -348,9 +351,10 @@ class OCRProcessor:
|
||||
|
||||
if max_workers is None:
|
||||
try:
|
||||
max_workers = self.config.getint('Performance', 'max_workers', fallback=4)
|
||||
# 强制设为 1 以严格遵守百度 QPS=2 的限制(一张图调两个接口就满了)
|
||||
max_workers = self.config.getint('Performance', 'max_workers', fallback=1)
|
||||
except Exception:
|
||||
max_workers = 4
|
||||
max_workers = 1
|
||||
|
||||
# 获取未处理的图片
|
||||
unprocessed_images = self.get_unprocessed_images()
|
||||
|
||||
Reference in New Issue
Block a user