167 lines
6.4 KiB
Python
167 lines
6.4 KiB
Python
# -*- coding: utf-8 -*-
|
||
"""腾讯云 OCR 适配层 — 封装通用票据识别 & 通用文字识别"""
|
||
import base64
|
||
import logging
|
||
|
||
_logger = logging.getLogger(__name__)
|
||
|
||
# tencentcloud-sdk-python 为可选依赖,导入失败时给出明确提示
|
||
try:
|
||
from tencentcloud.common import credential
|
||
from tencentcloud.common.exception.tencent_cloud_sdk_exception import TencentCloudSDKException
|
||
from tencentcloud.ocr.v20181119 import ocr_client, models
|
||
HAS_TENCENT_SDK = True
|
||
except ImportError:
|
||
HAS_TENCENT_SDK = False
|
||
_logger.warning(
|
||
'yuthon_ai_agent: tencentcloud-sdk-python 未安装,OCR 功能不可用。'
|
||
'请执行: pip install tencentcloud-sdk-python'
|
||
)
|
||
|
||
# 支持的图片 MIME 类型
|
||
SUPPORTED_IMAGE_TYPES = {
|
||
'image/png', 'image/jpeg', 'image/jpg', 'image/bmp', 'image/gif',
|
||
}
|
||
SUPPORTED_PDF_TYPE = 'application/pdf'
|
||
MAX_IMAGE_SIZE = 10 * 1024 * 1024 # 10MB
|
||
|
||
|
||
class TencentOCR:
|
||
"""腾讯云 OCR 适配器"""
|
||
|
||
def __init__(self, secret_id, secret_key, region='ap-guangzhou'):
|
||
if not HAS_TENCENT_SDK:
|
||
raise RuntimeError(
|
||
'tencentcloud-sdk-python 未安装,请执行: pip install tencentcloud-sdk-python'
|
||
)
|
||
if not secret_id or not secret_key:
|
||
raise ValueError('OCR SecretId / SecretKey 未配置')
|
||
self.cred = credential.Credential(secret_id, secret_key)
|
||
self.region = region
|
||
self._client = None
|
||
|
||
@property
|
||
def client(self):
|
||
if self._client is None:
|
||
self._client = ocr_client.OcrClient(self.cred, self.region)
|
||
return self._client
|
||
|
||
def _encode_image(self, raw_bytes):
|
||
"""将图片/PDF 二进制数据转为 Base64 字符串"""
|
||
return base64.b64encode(raw_bytes).decode('utf-8')
|
||
|
||
def _validate_input(self, raw_bytes, mimetype):
|
||
"""校验输入大小 & 格式"""
|
||
if len(raw_bytes) > MAX_IMAGE_SIZE:
|
||
raise ValueError(f'文件大小超过限制(最大 10MB),当前 {len(raw_bytes) / 1024 / 1024:.1f}MB')
|
||
if mimetype not in SUPPORTED_IMAGE_TYPES and mimetype != SUPPORTED_PDF_TYPE:
|
||
raise ValueError(f'不支持的文件格式:{mimetype},仅支持 PNG/JPG/BMP/GIF/PDF')
|
||
|
||
def recognize_general(self, raw_bytes, mimetype='image/png'):
|
||
"""通用文字识别 — 适用于任意图片/PDF 的文字提取。
|
||
|
||
:param raw_bytes: 文件二进制数据
|
||
:param mimetype: MIME 类型
|
||
:return: dict {success, text, error}
|
||
"""
|
||
self._validate_input(raw_bytes, mimetype)
|
||
try:
|
||
image_base64 = self._encode_image(raw_bytes)
|
||
req = models.GeneralBasicOCRRequest()
|
||
req.ImageBase64 = image_base64
|
||
# PDF 需要额外参数
|
||
if mimetype == SUPPORTED_PDF_TYPE:
|
||
req.IsPdf = True
|
||
|
||
resp = self.client.GeneralBasicOCR(req)
|
||
texts = []
|
||
for item in resp.TextDetections:
|
||
texts.append(item.DetectedText)
|
||
return {
|
||
'success': True,
|
||
'text': '\n'.join(texts),
|
||
}
|
||
except TencentCloudSDKException as e:
|
||
_logger.exception('Tencent OCR GeneralBasicOCR 调用失败')
|
||
return {'success': False, 'error': f'OCR 识别失败:{e.message}'}
|
||
except Exception as e:
|
||
_logger.exception('OCR GeneralBasicOCR 异常')
|
||
return {'success': False, 'error': f'OCR 异常:{str(e)}'}
|
||
|
||
def recognize_invoice(self, raw_bytes, mimetype='image/png'):
|
||
"""通用票据识别(高级版)— 自动识别增值税发票、火车票、机票等。
|
||
|
||
:param raw_bytes: 文件二进制数据
|
||
:param mimetype: MIME 类型
|
||
:return: dict {success, text, error}
|
||
"""
|
||
self._validate_input(raw_bytes, mimetype)
|
||
try:
|
||
image_base64 = self._encode_image(raw_bytes)
|
||
req = models.RecognizeGeneralInvoiceRequest()
|
||
req.ImageBase64 = image_base64
|
||
if mimetype == SUPPORTED_PDF_TYPE:
|
||
req.IsPdf = True
|
||
|
||
resp = self.client.RecognizeGeneralInvoice(req)
|
||
items = resp.MixedInvoiceItems or []
|
||
|
||
if not items:
|
||
return {'success': True, 'text': '(未识别到票据信息)'}
|
||
|
||
texts = []
|
||
for item in items:
|
||
item_type = getattr(item, 'Type', '未知票据')
|
||
single_invoice = item.SingleInvoiceInfos
|
||
if single_invoice:
|
||
texts.append(f'—— {item_type} ——')
|
||
for key, value in single_invoice.items():
|
||
if value:
|
||
texts.append(f'{key}:{value}')
|
||
return {
|
||
'success': True,
|
||
'text': '\n'.join(texts) if texts else '(未识别到票据内容)',
|
||
}
|
||
except TencentCloudSDKException as e:
|
||
_logger.exception('Tencent OCR RecognizeGeneralInvoice 调用失败')
|
||
return {'success': False, 'error': f'OCR 识别失败:{e.message}'}
|
||
except Exception as e:
|
||
_logger.exception('OCR RecognizeGeneralInvoice 异常')
|
||
return {'success': False, 'error': f'OCR 异常:{str(e)}'}
|
||
|
||
def smart_recognize(self, raw_bytes, mimetype='image/png'):
|
||
"""智能识别:先尝试票据识别,失败则回退到通用文字识别。
|
||
|
||
:return: dict {success, text, error}
|
||
"""
|
||
# 先尝试票据识别
|
||
result = self.recognize_invoice(raw_bytes, mimetype)
|
||
if result['success'] and result.get('text') and '(未识别到票据' not in result['text']:
|
||
return result
|
||
# 回退到通用文字
|
||
return self.recognize_general(raw_bytes, mimetype)
|
||
|
||
|
||
def get_ocr_instance(provider):
|
||
"""从 AiProvider 记录中创建 OCR 实例。
|
||
|
||
:param provider: ai.provider 记录
|
||
:return: TencentOCR 实例或 None
|
||
"""
|
||
if not provider.ocr_enabled:
|
||
return None
|
||
if not provider.ocr_secret_id or not provider.ocr_secret_key:
|
||
_logger.warning('OCR 已启用但 SecretId/SecretKey 未配置 (provider=%s)', provider.name)
|
||
return None
|
||
if not HAS_TENCENT_SDK:
|
||
return None
|
||
try:
|
||
return TencentOCR(
|
||
secret_id=provider.ocr_secret_id,
|
||
secret_key=provider.ocr_secret_key,
|
||
region=provider.ocr_region or 'ap-guangzhou',
|
||
)
|
||
except Exception as e:
|
||
_logger.exception('创建 OCR 实例失败')
|
||
return None
|