dingdanquanliucheng/backend/app/services/ai_service.py
wsb1224 f8ea5db19a fix: 阿里云OCR和快递100面单识别本地图片404问题
提取共享函数read_image_bytes,当URL为/uploads/本地路径时
直接从文件系统读取,避免通过公网HTTP下载导致404。
统一修复AliyunOCRProvider和recognize_waybill两处图片下载。

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-07-04 10:02:36 +08:00

2010 lines
80 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.

"""AI 智能服务模块
提供两大核心能力:
1. OCR 图片识别:将图片中的文字提取为结构化数据,支持阿里云 OCR 和 Mock 两种 provider。
2. 智能填单通过文本预处理、规则提取、LLM 结构化解析、产品库模糊匹配等多层管线,
从粘贴文本或图片中自动识别订单信息并填充到表单。
架构分层:
- TextPreprocessor清洗粘贴文本中的时间戳、表情标记等噪声。
- RuleExtractor基于正则的确定性字段提取手机号、数量、价格、地址
- OrderOCRParser基于 OCR 行列表做订单版面分析和字段提取。
- LLMOrderParser调用通义千问 qwen-plus 做订单文本结构化解析。
- BaseOCRProvider / AliyunOCRProvider / MockOCRProviderOCR 提供者抽象与实现。
- AIService业务入口串联上述组件完成图片识别和智能填单流程。
被调用方ai 路由(图片识别、识别结果修正、智能填单接口)。
"""
import json
import logging
import re
from decimal import Decimal
from difflib import SequenceMatcher
from urllib import error, parse, request
logger = logging.getLogger(__name__)
from sqlalchemy import select
from sqlalchemy.exc import SQLAlchemyError
from sqlalchemy.orm import Session
from backend.app.core.config import get_settings
from backend.app.core.error_codes import ErrorCode
from backend.app.core.exceptions import AppException
from backend.app.repositories.ai_repository import AIRepository
from backend.app.services.audit_service import audit_service
def read_image_bytes(image_url: str) -> bytes:
"""读取图片字节。本地 /uploads/ 路径直接读文件,其他走 HTTP。"""
from pathlib import Path, PurePosixPath
from urllib.parse import urlparse
url_path = PurePosixPath(urlparse(image_url).path)
if len(url_path.parts) >= 2 and url_path.parts[1] == "uploads":
settings = get_settings()
return (Path(settings.local_upload_dir) / url_path.name).read_bytes()
with request.urlopen(image_url, timeout=15) as resp:
return resp.read()
# ---------------------------------------------------------------------------
# 智能填单:文本预处理
# ---------------------------------------------------------------------------
class TextPreprocessor:
"""清洗粘贴文本中的噪声,保留有效订单信息。
处理内容包括:移除时间戳、括号时间标记、表情/媒体标记、用户名前缀等,
并合并多余空行。无外部依赖,纯文本处理。
"""
TIMESTAMP_RE = re.compile(r'\d{4}[-/]\d{1,2}[-/]\d{1,2}\s+\d{1,2}:\d{2}(:\d{2})?')
BRACKET_TIME_RE = re.compile(r'\[\d{1,2}:\d{2}(:\d{2})?\]')
USERNAME_PREFIX_RE = re.compile(r'^[^:\s]{1,10}[:]\s*', re.MULTILINE)
NOISE_MARKERS = [
'[图片]', '[表情]', '[语音]', '[视频]', '[文件]', '[链接]',
'[红包]', '[转账]', '[位置]', '[名片]', '—— ——',
]
def preprocess(self, text: str) -> str:
"""清洗文本,移除噪声内容并合并多余空行。
参数:
text: 原始粘贴文本。
返回:
清洗后的纯文本,空输入返回空字符串。
"""
if not text or not text.strip():
return ""
result = text
result = self.TIMESTAMP_RE.sub('', result)
result = self.BRACKET_TIME_RE.sub('', result)
for marker in self.NOISE_MARKERS:
result = result.replace(marker, '')
result = self.USERNAME_PREFIX_RE.sub('', result)
result = re.sub(r'\n{3,}', '\n\n', result)
return result.strip()
# ---------------------------------------------------------------------------
# 智能填单:规则提取层
# ---------------------------------------------------------------------------
class RuleExtractor:
"""基于规则的字段提取器,处理确定性高的字段。
通过正则表达式提取手机号、数量+单位、单价,并基于地址关键词密度识别地址段。
适用于文本预处理后的清洗文本。
"""
PHONE_RE = re.compile(r'1[3-9]\d{9}')
QTY_UNIT_RE = re.compile(
r'(\d+\.?\d*)\s*(张|个|卷|条|套|片|块|吨|件|箱|包|米|cm|厘米|毫米|mm|kg|KG|公斤|斤|组|台|根)'
)
PRICE_RE = re.compile(r'(?:单价|价格|报价)\s*[:]?\s*(\d+\.?\d*)')
ADDRESS_KEYWORDS = [
'', '', '', '', '', '', '', '',
'', '', '', '单元', '', '大厦', '广场',
]
def extract(self, text: str) -> dict:
"""从文本中提取确定性字段。
参数:
text: 清洗后的文本。
返回:
字典,包含 customer_mobile、_raw_quantity、_raw_unit、_raw_price、customer_address 等字段。
"""
result: dict = {}
phones = self.PHONE_RE.findall(text)
if phones:
result['customer_mobile'] = phones[0]
qty_match = self.QTY_UNIT_RE.search(text)
if qty_match:
result['_raw_quantity'] = float(qty_match.group(1))
result['_raw_unit'] = qty_match.group(2)
price_match = self.PRICE_RE.search(text)
if price_match:
result['_raw_price'] = float(price_match.group(1))
addr = self._extract_address(text)
if addr:
result['customer_address'] = addr
return result
def _extract_address(self, text: str) -> str | None:
"""基于地址关键词密度提取地址段。
逐行扫描,关键词命中数 >= 2 的连续行视为地址段,取最长段。
参数:
text: 清洗后的文本。
返回:
拼接后的地址字符串,未识别到时返回 None。
"""
lines = text.split('\n')
best: list[str] = []
current: list[str] = []
for line in lines:
stripped = line.strip()
if not stripped:
continue
addr_hits = sum(1 for kw in self.ADDRESS_KEYWORDS if kw in stripped)
if addr_hits >= 2:
current.append(stripped)
else:
if len(current) > len(best):
best = list(current)
current = []
if len(current) > len(best):
best = current
return ''.join(best) if best else None
# ---------------------------------------------------------------------------
# 智能填单OCR 版面分析
# ---------------------------------------------------------------------------
class OrderOCRParser:
"""基于 OCR 行列表做订单版面分析。
将 OCR 返回的文本行按区域分类(头部、地址、表格、其他),
从各区域中提取客户信息和产品明细。
"""
PHONE_RE = re.compile(r'1[3-9]\d{9}')
ADDRESS_KEYWORDS = {
'', '', '', '', '', '', '', '',
'', '', '', '单元', '', '大厦', '广场',
'', '', '', '', '',
}
QTY_RE = re.compile(
r'(\d+\.?\d*)\s*(张|个|卷|条|套|片|块|吨|件|箱|包|米|cm|厘米|毫米|mm|kg|KG|公斤|斤|组|台|根)'
)
PRICE_RE = re.compile(r'(?:单价|价格|报价|¥|¥)\s*[:]?\s*(\d+\.?\d*)')
TABLE_HEADER_KEYWORDS = {'产品', '品名', '名称', '规格', '数量', '单价', '金额', '合计'}
def parse(self, line_list: list[str], ocr_confidence: float) -> dict:
"""解析 OCR 行列表,输出结构化的版面分析结果。
参数:
line_list: OCR 返回的文本行列表。
ocr_confidence: OCR 识别置信度。
返回:
包含 raw_text、ocr_confidence、layout_zones、pre_filled、table_rows 的字典。
"""
raw_text = '\n'.join(line_list)
zones = self._classify_lines(line_list)
pre_filled = self._extract_from_zones(line_list, zones)
table_rows = self._parse_table_lines(line_list, zones.get('table_lines', []))
return {
"raw_text": raw_text,
"ocr_confidence": ocr_confidence,
"layout_zones": zones,
"pre_filled": pre_filled,
"table_rows": table_rows,
}
def _classify_lines(self, lines: list[str]) -> dict:
"""将 OCR 文本行按区域类型分类。
分类逻辑:前 3 行为头部、含地址关键词的连续行归为地址区、
遇到表格表头后连续的产品行归为表格区,其余为其他。
参数:
lines: OCR 文本行列表。
返回:
包含 header_lines、address_lines、table_lines、other_lines 的分类索引字典。
"""
zones: dict[str, list[int]] = {
"header_lines": [], "address_lines": [],
"table_lines": [], "other_lines": [],
}
table_started = False
for i, line in enumerate(lines):
stripped = line.strip()
if not stripped:
continue
if any(kw in stripped for kw in self.TABLE_HEADER_KEYWORDS):
table_started = True
zones["table_lines"].append(i)
continue
if table_started:
if self.QTY_RE.search(stripped) or self.PRICE_RE.search(stripped):
zones["table_lines"].append(i)
continue
table_started = False
addr_count = sum(1 for kw in self.ADDRESS_KEYWORDS if kw in stripped)
if addr_count >= 2:
zones["address_lines"].append(i)
continue
if i < 3:
zones["header_lines"].append(i)
else:
zones["other_lines"].append(i)
return zones
def _extract_from_zones(self, lines: list[str], zones: dict) -> dict:
"""从已分类的区域中提取客户基本信息(姓名、手机、地址)。
参数:
lines: OCR 文本行列表。
zones: 区域分类结果字典。
返回:
包含 customer_mobile、customer_name、customer_address 的预填字段字典。
"""
result: dict = {}
for i in zones.get("header_lines", []):
text = lines[i].strip()
if not text:
continue
phone_match = self.PHONE_RE.search(text)
if phone_match:
result["customer_mobile"] = phone_match.group(0)
name_part = self.PHONE_RE.sub('', text).strip()
if name_part and len(name_part) <= 20 and not any(
kw in name_part for kw in self.ADDRESS_KEYWORDS
):
result.setdefault("customer_name", name_part)
addr_parts = [lines[i].strip() for i in zones.get("address_lines", [])]
if addr_parts:
result["customer_address"] = ''.join(addr_parts)
return result
def _parse_table_lines(self, lines: list[str], table_indices: list[int]) -> list[dict]:
"""解析表格区域中的产品行,提取数量、单价和产品名称。
参数:
lines: OCR 文本行列表。
table_indices: 表格区域的行索引列表。
返回:
产品行列表,每个元素包含 raw原始行和 fields提取的字段字典。
"""
rows: list[dict] = []
for i in table_indices:
line = lines[i].strip()
if not line:
continue
if any(kw in line for kw in self.TABLE_HEADER_KEYWORDS) and not self.QTY_RE.search(line):
continue
row: dict = {"raw": line, "fields": {}}
qty_match = self.QTY_RE.search(line)
if qty_match:
row["fields"]["quantity"] = float(qty_match.group(1))
row["fields"]["unit"] = qty_match.group(2)
price_match = self.PRICE_RE.search(line)
if price_match:
row["fields"]["sale_price"] = float(price_match.group(1))
name_part = self.QTY_RE.sub('', line)
name_part = self.PRICE_RE.sub('', name_part)
name_part = re.sub(r'\d+\.?\d*', '', name_part).strip()
if name_part:
row["fields"]["product_name"] = name_part
if row["fields"]:
rows.append(row)
return rows
# ---------------------------------------------------------------------------
# 智能填单LLM 结构化解析
# ---------------------------------------------------------------------------
class LLMOrderParser:
"""调用通义千问 qwen-plus 进行订单文本结构化解析。
支持纯文本和 OCR 图片两种输入模式,通过 system prompt 引导 LLM 输出标准 JSON 格式的订单信息。
内置产品库提示词product_hints帮助 LLM 匹配产品名称和规格。
"""
TEXT_SYSTEM_PROMPT = """你是高温布/特氟龙/胶带/输送带行业的订单解析助手。从业务员粘贴的文本中提取订单信息。
## 输入格式(非常重要)
业务员粘贴的文本由两列组成,中间用 tab 或多个空格分隔:
- 左列:产品描述(厚度 + 材料名 + 尺寸 + 数量 + 工艺要求)
- 右列:客户姓名,手机号,地址(备注)
示例0.13mm 咖色高温布 1*10米 SYY+568113711561796广东省广州市白云区石井街道
必须严格按这个结构拆分,不要把左列的产品信息混入右列的客户地址中。
## customer_address 规则(非常重要)
customer_address 只能包含地理地址信息(省/市/区/县/镇/街道/路/号/楼/室/栋/公司名),绝对不能包含:
- 产品描述(厚度、材料名、尺寸等)
- 客户姓名和手机号
- 订单编号(如 SYY+5681
- 整段原始文本
## product_name 规则(非常重要)
product_name 只填产品材料类型名称,不要带型号编号:
- "0.13mm 咖色高温布" → product_name="高温布"
- "145g咖色高温布" → product_name="高温布"
- "黑色4*4mm网格输送带" → product_name="网格输送带"
- "咖色特氟龙胶带" → product_name="特氟龙胶带"
- "用凯夫拉布做" → product_name="凯夫拉布"
常见产品名:高温布、特氟龙胶带、网格输送带、烧烤网格片、凯夫拉布、硅胶布、胶带、输送带、揉面垫、硅胶烤垫
## specification 规则
specification 必须包含所有产品参数,格式为 "厚度/克重 颜色 标准化尺寸"
- "0.13mm 咖色高温布 1*10米" → specification="0.13mm 咖色 1000mm×10000mm"
- "145g咖色高温布 25*30厘米 1000张" → specification="145g 咖色 250mm×300mm"
## 尺寸标准化
将所有尺寸统一转换为 "长mm×宽mm" 格式:
- "1*10米""1000mm×10000mm"
- "60*80cm""600mm×800mm"
- "25*30厘米""250mm×300mm"
- "长800mm宽250mm""800mm×250mm"
- "80厘米宽*1米长""1000mm×800mm"
## 数量单位
常用单位:张、个、卷、条、套、片、块
## 备注提取
- 括号中的内容:(不要浅色的)、(单片卷)、(加急安排)→ items[].remark
- 尾部工艺要求:带回仓库包装、黑膜包边、牛鼻子接头 → items[].remark
## 严格按以下 JSON 格式输出:
{{
"customer_name": "客户姓名",
"customer_mobile": "11位手机号",
"customer_address": "仅包含地理地址",
"order_source": null,
"delivery_type": null,
"remark": "订单级备注",
"items": [
{{
"product_name": "产品材料类型(不带型号)",
"specification": "产品自身规格(厚度、颜色等标准参数)",
"demand_specification": "客户要求的尺寸统一转换为m格式如0.25m×0.3m仅提取客户说的尺寸无尺寸则为null",
"unit": "数量单位",
"quantity": 数量数字,
"sale_price": null,
"remark": "产品级备注"
}}
]
}}
## 示例
输入: 0.13mm 咖色高温布 1*10米 SYY+568113711561796广东省 广州市 白云区 石井街道
输出: {{"customer_name":"SYY+5681","customer_mobile":"13711561796","customer_address":"广东省广州市白云区石井街道","remark":null,"items":[{{"product_name":"高温布","specification":"0.13mm 咖色","demand_specification":"1m×10m","unit":null,"quantity":null,"sale_price":null,"remark":null}}]}}
输入: 145g咖色高温布 25*30厘米 1000张 何莹19932065916河北省 邯郸市 大名县 园中街瑞盾门业
输出: {{"customer_name":"何莹","customer_mobile":"19932065916","customer_address":"河北省邯郸市大名县园中街瑞盾门业","remark":null,"items":[{{"product_name":"高温布","specification":"145g 咖色","demand_specification":"0.25m×0.3m","unit":"","quantity":1000,"sale_price":null,"remark":null}}]}}
输入: 用凯夫拉布做18丝布带。宽度22.8cm长度56.5两边重叠6mm 50个重叠位置要加固 徐阳飞15158849480浙江省 杭州市 富阳区 场口镇
输出: {{"customer_name":"徐阳飞","customer_mobile":"15158849480","customer_address":"浙江省杭州市富阳区场口镇","remark":null,"items":[{{"product_name":"凯夫拉布","specification":"18丝布带 重叠6mm","demand_specification":"0.228m×0.565m","unit":"","quantity":50,"sale_price":null,"remark":"重叠位置要加固"}}]}}
可选的产品库(名称 + 规格):
{product_hints}"""
IMAGE_SYSTEM_PROMPT = """你是订单信息解析助手。OCR 系统已经从图片中提取了文本并做了初步分析。
你需要基于 OCR 的结果,补充和完善订单信息。
OCR 已提取的信息:
- OCR 置信度:{ocr_confidence}
- 已识别的客户姓名:{customer_name}
- 已识别的客户手机:{customer_mobile}
- 已识别的客户地址:{customer_address}
- 产品表格区域识别到的原始行:
{table_rows}
你需要完成:
1. 验证 OCR 提取的字段是否合理
2. 从原始文本中补充 OCR 未提取到的字段
3. 解析产品明细(如果 OCR 表格区域数据可用,优先使用)
4. 识别订单来源、配送方式等附加信息
严格按以下 JSON 格式输出:
{{
"customer_name": "客户姓名",
"customer_mobile": "手机号",
"customer_address": "地址",
"order_source": "订单来源",
"delivery_type": "配送方式",
"remark": "备注",
"items": [
{{
"product_name": "产品名称",
"specification": "规格",
"unit": "单位",
"quantity": 数量数字,
"sale_price": 单价数字
}}
]
}}
规则:
- 无法识别的字段填 null不要编造
- 如果 OCR 已提取的字段看起来正确,直接沿用
可选的产品库(名称 + 规格):
{product_hints}"""
def _build_product_hints(self, product_groups: list[dict]) -> str:
"""构建产品库提示词,供 LLM 在解析时参考。
参数:
product_groups: 产品分组列表,每组包含 product_name 和 specifications。
返回:
格式化的产品提示文本,最多 30 个产品,每个最多 5 个规格。
"""
hints: list[str] = []
for group in product_groups[:50]:
specs = ', '.join(
f"{s['specification']}({s['unit']})"
for s in group.get("specifications", [])[:5]
)
hints.append(f"- {group['product_name']}: {specs}")
return '\n'.join(hints) or "(产品库为空)"
def parse(self, text: str, product_groups: list[dict],
api_key: str, api_url: str,
ocr_context: dict | None = None,
model: str = "qwen-plus") -> dict:
"""调用 LLM 解析订单文本,返回结构化订单数据。
参数:
text: 待解析的订单文本。
product_groups: 产品库数据,用于构建提示词。
api_key: LLM API 访问密钥。
api_url: LLM API 地址。
ocr_context: 可选的 OCR 上下文,图片模式时提供版面分析结果。
返回:
包含 customer_name、customer_mobile、items 等字段的订单字典。
异常:
LLM 调用失败或返回格式异常时抛出 AppException。
"""
product_hints = self._build_product_hints(product_groups)
if ocr_context:
system_prompt = self.IMAGE_SYSTEM_PROMPT.format(
ocr_confidence=ocr_context.get("ocr_confidence", "N/A"),
customer_name=ocr_context.get("pre_filled", {}).get("customer_name", "未识别"),
customer_mobile=ocr_context.get("pre_filled", {}).get("customer_mobile", "未识别"),
customer_address=ocr_context.get("pre_filled", {}).get("customer_address", "未识别"),
table_rows='\n'.join(
f" - {row['raw']}" for row in ocr_context.get("table_rows", [])
) or "",
product_hints=product_hints,
)
else:
system_prompt = self.TEXT_SYSTEM_PROMPT.format(product_hints=product_hints)
payload = json.dumps({
"model": model,
"messages": [
{"role": "system", "content": system_prompt},
{"role": "user", "content": f"请解析以下订单内容:\n\n{text}"},
],
"temperature": 0.1,
"max_tokens": 1024,
}).encode("utf-8")
headers = {
"Content-Type": "application/json",
"Authorization": f"Bearer {api_key}",
}
req = urllib_request.Request(
url=api_url, data=payload, headers=headers, method="POST"
)
with urllib_request.urlopen(req, timeout=30) as resp:
result = json.loads(resp.read().decode("utf-8"))
content = result["choices"][0]["message"]["content"]
return self._extract_json(content)
def _extract_json(self, text: str) -> dict:
"""从 LLM 响应文本中提取 JSON 内容。
处理 LLM 可能包裹在 markdown 代码块中的情况,
以及 LLM 响应被截断导致 JSON 不完整的常见场景。
参数:
text: LLM 原始响应文本。
返回:
解析后的字典。
"""
text = text.strip()
if text.startswith("```"):
text = text.split("\n", 1)[1]
text = text.rsplit("```", 1)[0]
text = text.strip()
try:
return json.loads(text)
except json.JSONDecodeError:
# LLM 响应可能被截断,尝试截断到最后一个完整 } 修复
last_brace = text.rfind("}")
while last_brace > 0:
try:
return json.loads(text[:last_brace + 1])
except json.JSONDecodeError:
last_brace = text.rfind("}", 0, last_brace - 1)
raise
def safe_parse(self, text: str, product_groups: list[dict],
api_key: str, api_url: str,
ocr_context: dict | None = None,
model: str = "qwen-plus") -> dict | None:
"""安全版解析入口,异常时返回 None 而非抛出异常。
参数:
text: 待解析的订单文本。
product_groups: 产品库数据。
api_key: LLM API 访问密钥。
api_url: LLM API 地址。
ocr_context: 可选的 OCR 上下文。
model: LLM 模型名称。
返回:
解析后的订单字典,失败时返回 None。
"""
try:
return self.parse(text, product_groups, api_key, api_url, ocr_context, model=model)
except Exception:
logger.warning("[LLM] 解析失败", exc_info=True)
return None
# 避免和后面 urllib_request 冲突,这里在文件顶部已经 import 了
urllib_request = request
class BaseOCRProvider:
"""OCR 提供者基类,定义统一的识别接口。
子类需实现 recognize 方法,返回 (raw_result, suggested_result, confidence) 三元组。
"""
provider_name = "base"
def recognize(self, image_url: str, biz_type: str, biz_id: int) -> tuple[dict, dict, float]:
"""识别图片中的文字内容。
参数:
image_url: 图片 URL 地址。
biz_type: 业务类型标识。
biz_id: 业务记录 ID。
返回:
三元组 (原始识别结果, 建议结果, 置信度)。
"""
raise NotImplementedError
class MockOCRProvider(BaseOCRProvider):
"""模拟 OCR 提供者,用于开发调试和测试环境。
不调用真实 OCR 服务,直接返回固定格式的模拟结果,置信度 0.92。
"""
provider_name = "mock_aliyun_adapter"
def recognize(self, image_url: str, biz_type: str, biz_id: int) -> tuple[dict, dict, float]:
"""返回模拟的 OCR 识别结果。
参数:
image_url: 图片 URL 地址。
biz_type: 业务类型标识。
biz_id: 业务记录 ID。
返回:
三元组 (模拟原始结果, 模拟建议结果, 0.92)。
"""
# 模拟 OCR 识别出的文本行,用于开发调试时走通完整解析流程
sample_lines = [
"张三 13800138000",
"广东省深圳市南山区科技园路88号",
"产品名称:工业级密封圈",
"规格DN50 数量100件 单价15.5",
"合计金额1550.00",
"备注:尽快发货",
]
suggested_result = {
"image_name": image_url.rsplit("/", 1)[-1],
"biz_type": biz_type,
"biz_id": biz_id,
"line_list": sample_lines,
"recognized_text": "\n".join(sample_lines),
}
raw_result = {
"provider": self.provider_name,
"image_url": image_url,
"fields": suggested_result,
}
return raw_result, suggested_result, 0.92
class AliyunOCRProvider(BaseOCRProvider):
"""阿里云 OCR 提供者,通过官方 SDK 调用阿里云文字识别服务。
依赖 settings 中的 aliyun_ai_access_key_id、aliyun_ai_access_key_secret、
aliyun_ai_region 等配置。使用 alibabacloud_ocr_api20210707 SDK。
"""
provider_name = "aliyun_ocr"
ENDPOINT = "ocr-api.cn-hangzhou.aliyuncs.com"
def __init__(self, settings) -> None:
self.settings = settings
def recognize(self, image_url: str, biz_type: str, biz_id: int) -> tuple[dict, dict, float]:
"""调用阿里云 OCR 识别图片中的文字。
参数:
image_url: 图片 URL 地址。
biz_type: 业务类型标识。
biz_id: 业务记录 ID。
返回:
三元组 (原始 API 响应, 结构化建议结果, 置信度)。
"""
payload = self._call_aliyun(image_url)
raw_result = {
"provider": self.provider_name,
"region": self.settings.aliyun_ai_region,
"image_url": image_url,
"payload": payload,
}
suggested_result = self._build_suggested_result(payload, image_url, biz_type, biz_id)
confidence = self._extract_confidence(payload)
return raw_result, suggested_result, confidence
def _get_client(self):
"""构建阿里云 OCR SDK 客户端。"""
from alibabacloud_ocr_api20210707.client import Client as OcrClient
from alibabacloud_tea_openapi.models import Config
config = Config(
access_key_id=self.settings.aliyun_ai_access_key_id,
access_key_secret=self.settings.aliyun_ai_access_key_secret,
endpoint=self.ENDPOINT,
)
return OcrClient(config)
def _call_aliyun(self, image_url: str) -> dict:
"""通过官方 SDK 调用阿里云 OCR 接口。
先从 URL 下载图片字节,再通过 SDK 发送识别请求。
参数:
image_url: 图片 URL 地址。
返回:
阿里云 OCR SDK 的响应字典。
异常:
下载图片失败或 SDK 调用失败时抛出 AppException。
"""
from alibabacloud_ocr_api20210707.models import RecognizeGeneralRequest
try:
image_bytes = read_image_bytes(image_url)
except Exception as exc:
raise AppException(
code=ErrorCode.THIRD_PARTY_FAILED,
message=f"下载 OCR 图片失败:{exc}",
status_code=400,
) from exc
try:
client = self._get_client()
ocr_request = RecognizeGeneralRequest(body=image_bytes)
response = client.recognize_general(ocr_request)
except Exception as exc:
raise AppException(
code=ErrorCode.THIRD_PARTY_FAILED,
message=f"阿里云 OCR SDK 调用失败:{exc}",
status_code=400,
) from exc
body = response.body
if hasattr(body, "to_map"):
payload = body.to_map()
elif isinstance(body, dict):
payload = body
else:
payload = json.loads(str(body)) if body else {}
# SDK 返回 Data 为 JSON 字符串,解析为 dict 便于后续统一提取
if isinstance(payload.get("Data"), str):
try:
payload["Data"] = json.loads(payload["Data"])
except (json.JSONDecodeError, TypeError):
pass
return payload
def _build_suggested_result(self, payload: dict, image_url: str, biz_type: str, biz_id: int) -> dict:
"""将阿里云 OCR 原始响应转换为统一的结构化建议结果。
从响应中递归提取文本行,拼接为 recognized_text 和 line_list。
参数:
payload: 阿里云 OCR API 响应字典。
image_url: 图片 URL 地址。
biz_type: 业务类型标识。
biz_id: 业务记录 ID。
返回:
包含 image_name、biz_type、recognized_text、line_list 等字段的建议结果字典。
"""
# 不强依赖某个固定响应字段,先把常见 OCR 结果平铺为统一结构,便于前端先联调。
texts = self._collect_texts(payload)
result = {
"image_name": image_url.rsplit("/", 1)[-1],
"biz_type": biz_type,
"biz_id": biz_id,
"recognized_text": "\n".join(texts).strip(),
"line_list": texts,
}
if texts:
result["customer_name"] = texts[0][:50]
return result
def _collect_texts(self, value) -> list[str]:
"""从嵌套的字典/列表结构中递归提取所有文本值并去重。
识别常见 OCR 响应字段名text、content、word、words 等),
跳过空值和重复文本。
参数:
value: OCR 响应中的任意嵌套结构。
返回:
去重后的文本列表,保持原始出现顺序。
"""
texts: list[str] = []
if isinstance(value, dict):
for key, item in value.items():
lowered = str(key).lower()
if lowered in {"text", "content", "word", "words", "value", "label", "title"} and isinstance(item, str):
stripped = item.strip()
if stripped:
texts.append(stripped)
else:
texts.extend(self._collect_texts(item))
elif isinstance(value, list):
for item in value:
texts.extend(self._collect_texts(item))
elif isinstance(value, str):
stripped = value.strip()
if stripped:
texts.append(stripped)
deduplicated: list[str] = []
seen: set[str] = set()
for text in texts:
if text not in seen:
seen.add(text)
deduplicated.append(text)
return deduplicated
def _extract_confidence(self, payload: dict) -> float:
"""从 OCR 响应中提取置信度分数。
依次尝试 confidence、score、probability 字段,未找到时默认返回 0.9。
参数:
payload: OCR API 响应字典。
返回:
置信度浮点数0.0 ~ 1.0)。
"""
for key in ("confidence", "score", "probability"):
value = payload.get(key)
if isinstance(value, (int, float)):
return float(value)
if isinstance(value, str):
try:
return float(value)
except ValueError:
continue
return 0.9
class AIService:
"""AI 智能业务服务入口。
整合 OCR 识别、规则提取、LLM 解析、产品库匹配等能力,
提供图片识别recognize_image、识别结果修正correct_result
智能填单parse_order三大对外接口。
依赖:
- AIRepositoryAI 识别日志的存储。
- get_settings系统配置包含 API 密钥和 provider 选择。
- AuditService操作审计日志记录。
"""
def __init__(self) -> None:
self.repository = AIRepository()
self.settings = get_settings()
def recognize_image(self, payload: dict, session: Session | None = None) -> dict:
"""调用 OCR 识别图片中的文字并保存识别日志。
参数:
payload: 请求数据字典,必含 image_url、biz_type可选 biz_id。
session: 数据库会话,不可为 None。
返回:
包含 log_id、raw_result、confidence、suggested_result 的字典。
被调用方ai 路由(图片识别接口)。
"""
image_url = (payload.get("image_url") or "").strip()
biz_type = (payload.get("biz_type") or "").strip()
biz_id = payload.get("biz_id")
if not image_url:
raise AppException(code=ErrorCode.PARAM_ERROR, message="图片地址不能为空", status_code=400)
if not biz_type:
raise AppException(code=ErrorCode.PARAM_ERROR, message="业务类型不能为空", status_code=400)
if image_url.startswith("mock://fail"):
raise AppException(code=ErrorCode.THIRD_PARTY_FAILED, message="AI 识别失败", status_code=400)
provider = self._build_provider()
raw_result, suggested_result, confidence = provider.recognize(image_url, biz_type, int(biz_id or 0))
if session is not None:
try:
# 识别原始响应和建议结果都落库,后面人工修正和问题回溯都依赖这份快照。
log = self.repository.create_log(
session,
{
"biz_type": biz_type,
"biz_id": biz_id,
"image_url": image_url,
"raw_result": json.dumps(raw_result, ensure_ascii=False),
"confidence": Decimal(str(confidence)),
"corrected_result": None,
"created_by": None,
},
)
audit_service.write_log(
session,
{
"operate_type": "ai_recognize",
"biz_type": biz_type,
"biz_id": biz_id,
"before_value": None,
"after_value": {
"log_id": log.id,
"provider": raw_result.get("provider"),
"confidence": confidence,
"suggested_result": suggested_result,
},
"remark": f"AI识别图片 {image_url}",
},
)
session.commit()
return {
"log_id": log.id,
"raw_result": raw_result,
"confidence": confidence,
"suggested_result": suggested_result,
}
except AppException:
session.rollback()
raise
except SQLAlchemyError as exc:
session.rollback()
raise AppException(code=ErrorCode.INTERNAL_ERROR, message="数据库异常", status_code=500) from exc
raise AppException(code=ErrorCode.INTERNAL_ERROR, message="数据库连接不可用", status_code=500)
def correct_result(self, log_id: int, payload: dict, session: Session | None = None) -> dict:
"""人工修正 AI 识别结果并保存。
参数:
log_id: AI 识别日志 ID。
payload: 修正数据字典,必含 corrected_resultdict
session: 数据库会话,不可为 None。
返回:
包含 log_id 和 corrected 标志的字典。
被调用方ai 路由(识别结果修正接口)。
"""
corrected_result = payload.get("corrected_result")
if not isinstance(corrected_result, dict) or not corrected_result:
raise AppException(code=ErrorCode.PARAM_ERROR, message="修正结果不能为空", status_code=400)
if session is not None:
try:
log = self.repository.get_log(session, log_id)
if log is None:
raise AppException(code=ErrorCode.NOT_FOUND, message="识别记录不存在", status_code=404)
before_value = self._safe_load_json(log.corrected_result)
log.corrected_result = json.dumps(corrected_result, ensure_ascii=False)
session.add(log)
audit_service.write_log(
session,
{
"operate_type": "ai_correct",
"biz_type": log.biz_type,
"biz_id": log.biz_id,
"before_value": before_value,
"after_value": corrected_result,
"remark": f"修正AI识别结果 {log_id}",
},
)
session.commit()
return {"log_id": log.id, "corrected": True}
except AppException:
session.rollback()
raise
except SQLAlchemyError as exc:
session.rollback()
raise AppException(code=ErrorCode.INTERNAL_ERROR, message="数据库异常", status_code=500) from exc
raise AppException(code=ErrorCode.INTERNAL_ERROR, message="数据库连接不可用", status_code=500)
def _build_provider(self) -> BaseOCRProvider:
"""根据系统配置构建对应的 OCR 提供者实例。
当 ai_provider 配置为 'aliyun' 时返回 AliyunOCRProvider否则返回 MockOCRProvider。
返回:
OCR 提供者实例。
"""
provider_name = self.settings.ai_provider.strip().lower()
# provider 选择集中在这里,后续扩展腾讯云或本地 OCR 时不需要动主流程。
if provider_name == "aliyun":
return AliyunOCRProvider(self.settings)
return MockOCRProvider()
def _safe_load_json(self, value: str | None):
"""安全解析 JSON 字符串,异常时返回原始值或 None。
参数:
value: JSON 字符串或 None。
返回:
解析后的对象,解析失败返回原始字符串,空值返回 None。
"""
if not value:
return None
try:
return json.loads(value)
except json.JSONDecodeError:
return value
# ------------------------------------------------------------------
# 智能填单:核心解析入口
# ------------------------------------------------------------------
def parse_order(self, payload: dict, session: Session | None = None) -> dict:
"""智能填单核心入口:从文本或图片中解析订单信息。
处理流程:
1. 获取原始文本(图片模式先调 OCR
2. 文本预处理(去噪声)
3. 规则提取(正则提取手机号、数量、价格、地址)
4. LLM 结构化解析
5. 规则结果与 LLM 结果合并
6. 客户库预匹配
7. 产品库模糊匹配
8. 校验并计算置信度
参数:
payload: 请求数据字典,必含 input_type'text''image'
text 模式需提供 text 字段image 模式需提供 image_url 字段。
session: 数据库会话,为 None 时使用 mock 模式。
返回:
包含 parsed_order、raw_text、confidence、parse_source、product_matches、
factory_options、warnings 的字典。
被调用方ai 路由(智能填单接口)。
"""
input_type = payload["input_type"]
# 1. 获取原始文本 + OCR 中间结果
ocr_context = None
if input_type == "image":
ocr_context = self._ocr_and_parse_image(payload["image_url"])
raw_text = ocr_context["raw_text"]
else:
raw_text = payload["text"].strip()
# 2. 文本预处理
preprocessor = TextPreprocessor()
raw_text = preprocessor.preprocess(raw_text)
if not raw_text:
raise AppException(
code=ErrorCode.PARAM_ERROR,
message="无法提取到文本内容",
status_code=400,
)
# 3. 多订单切分text 模式下,按手机号锚点拆分为独立订单段)
chunks = self._split_orders(raw_text) if input_type == "text" else [raw_text]
# 4. 正则提取(整段提取手机号等确定性字段)
rule_extractor = RuleExtractor()
rule_result = rule_extractor.extract(raw_text)
if ocr_context and ocr_context.get("pre_filled"):
for key, value in ocr_context["pre_filled"].items():
if value and not rule_result.get(key):
rule_result[key] = value
# 5. LLM 结构化解析(每个订单段独立解析)
settings = self.settings
product_groups = self._get_product_groups(session)
logger.info("[智能填单] product_groups 数量: %d", len(product_groups))
if product_groups:
logger.info("[智能填单] 产品列表: %s", [g["product_name"] for g in product_groups[:10]])
llm_parser = LLMOrderParser()
if len(chunks) == 1:
llm_result = llm_parser.safe_parse(
chunks[0], product_groups,
settings.llm_parse_api_key or settings.aliyun_ai_access_key_id,
settings.llm_parse_api_url,
ocr_context=ocr_context,
model=settings.llm_parse_model,
)
else:
# 多段时逐段解析,合并 items
all_items = []
merged_customer = {}
for chunk in chunks:
chunk_result = llm_parser.safe_parse(
chunk, product_groups,
settings.llm_parse_api_key or settings.aliyun_ai_access_key_id,
settings.llm_parse_api_url,
model=settings.llm_parse_model,
)
if chunk_result:
all_items.extend(chunk_result.get("items", []))
for key in ("customer_name", "customer_mobile", "customer_address", "remark"):
if chunk_result.get(key) and not merged_customer.get(key):
merged_customer[key] = chunk_result[key]
llm_result = {**merged_customer, "items": all_items} if all_items else None
# 5. 合并
if llm_result:
# 标准化 demand_specification 中的单位格式
for item in llm_result.get("items", []):
if item.get("demand_specification"):
item["demand_specification"] = self._standardize_demand_specification(
item["demand_specification"]
)
merged = self._merge_results(rule_result, llm_result)
parse_source = "hybrid"
else:
merged = self._build_fallback_result(rule_result)
parse_source = "rule"
# 5.5 地址后处理:清理 LLM 可能混入的非地址内容
if merged.get("customer_address"):
merged["customer_address"] = self._clean_address(
merged["customer_address"], raw_text
)
if merged.get("customer_mobile") and not merged.get("customer_address"):
merged["customer_address"] = self._extract_address_after_phone(
raw_text, merged["customer_mobile"]
)
# 6. 客户库预匹配
if session is not None and merged.get("customer_mobile"):
try:
from backend.app.repositories.customer_repository import CustomerRepository
customer_repo = CustomerRepository()
existing_customer = customer_repo.find_by_mobile(
session, merged["customer_mobile"]
)
if existing_customer:
merged["customer_id"] = existing_customer.id
if existing_customer.address and not merged.get("customer_address"):
merged["customer_address"] = existing_customer.address
except Exception:
pass
# 7. 产品库模糊匹配
product_matches = self._match_products(
merged.get("items", []), product_groups, session
)
for pm in product_matches:
logger.info("[智能填单] 匹配: input=%s → 候选数=%d, 候选=%s",
pm["input_name"], len(pm["candidates"]),
[(c["product_name"], c["match_score"]) for c in pm["candidates"][:3]])
# 8. 校验 + 置信度
warnings = self._validate_parsed_order(merged)
if parse_source == "rule":
warnings.append("LLM 服务不可用,仅使用规则提取,部分字段可能不完整")
ocr_conf = ocr_context["ocr_confidence"] if ocr_context else None
confidence = self._calc_confidence(merged, warnings, ocr_confidence=ocr_conf)
return {
"parsed_order": merged,
"raw_text": raw_text,
"confidence": confidence,
"ocr_confidence": ocr_conf,
"parse_source": parse_source,
"is_mock": session is None,
"product_matches": product_matches,
"factory_options": self._get_factory_options(session),
"warnings": warnings,
}
_PHONE_ANCHOR_RE = re.compile(r'1[3-9]\d{9}')
def _split_orders(self, text: str) -> list[str]:
"""按手机号锚点将多条订单文本切分为独立段。
每个手机号标记一段客户信息的开始,手机号之前的行归为该段的产品行。
如果文本中只有一个手机号,返回原始文本(不分切)。
参数:
text: 预处理后的订单文本。
返回:
订单文本段列表,每段包含产品行和客户信息行。
"""
lines = text.split('\n')
anchors: list[tuple[int, str]] = []
for i, line in enumerate(lines):
match = self._PHONE_ANCHOR_RE.search(line)
if match:
anchors.append((i, match.group(0)))
if len(anchors) <= 1:
return [text]
chunks: list[str] = []
for idx, (line_idx, _phone) in enumerate(anchors):
start = anchors[idx - 1][0] + 1 if idx > 0 else 0
chunk_lines = lines[start:line_idx + 1]
chunk = '\n'.join(chunk_lines).strip()
if chunk:
chunks.append(chunk)
return chunks
_ADDRESS_KEYWORDS = {'', '', '', '', '', '', '', '', '', '', '', '', '', ''}
_ORDER_NO_RE = re.compile(r'[A-Z]{2,}[\d\-+]+') # SYY+5681, SO20260604 等订单号
def _clean_address(self, address: str, raw_text: str) -> str:
"""清理 LLM 返回的地址,去除混入的非地址内容。
如果地址中包含手机号或产品关键词,说明 LLM 没有正确拆分,
此时从原始文本中重新提取。
"""
if not address:
return address
# 如果地址中包含手机号,说明 LLM 把整段文本塞进了地址
if self._PHONE_ANCHOR_RE.search(address):
phone_match = self._PHONE_ANCHOR_RE.search(raw_text)
if phone_match:
return self._extract_address_after_phone(raw_text, phone_match.group(0))
return ""
# 去除订单号
address = self._ORDER_NO_RE.sub('', address).strip()
# 去除开头的逗号/空格
address = address.lstrip(',、 \t')
# 如果清理后地址关键词密度太低,说明不是有效地址
addr_hits = sum(1 for kw in self._ADDRESS_KEYWORDS if kw in address)
if addr_hits < 1 and len(address) > 5:
# 可能是 LLM 误填,尝试从原始文本提取
phone_match = self._PHONE_ANCHOR_RE.search(raw_text)
if phone_match:
return self._extract_address_after_phone(raw_text, phone_match.group(0))
return address
def _extract_address_after_phone(self, text: str, phone: str) -> str:
"""从原始文本中提取手机号之后的地址信息。"""
idx = text.find(phone)
if idx < 0:
return ""
after = text[idx + len(phone):].strip()
# 去掉开头的逗号/空格/顿号
after = after.lstrip(',、 \t')
return after if after else ""
def _ocr_and_parse_image(self, image_url: str) -> dict:
"""调用 OCR 识别图片并进行版面分析。
参数:
image_url: 图片 URL 地址。
返回:
OrderOCRParser.parse() 的返回结果,包含 raw_text、pre_filled、table_rows 等。
"""
provider = self._build_provider()
raw_result, suggested_result, ocr_confidence = provider.recognize(
image_url, "order_parse", 0
)
line_list = suggested_result.get("line_list", [])
if not line_list:
text = suggested_result.get("recognized_text", "")
line_list = [line for line in text.split("\n") if line.strip()]
ocr_parser = OrderOCRParser()
return ocr_parser.parse(line_list, ocr_confidence)
def _merge_results(self, rule_result: dict, llm_result: dict) -> dict:
"""合并规则提取结果和 LLM 解析结果。
以 LLM 结果为基础,规则提取的手机号优先使用,地址取较长者。
参数:
rule_result: 规则提取结果。
llm_result: LLM 解析结果。
返回:
合并后的订单字典。
"""
merged = dict(llm_result)
if rule_result.get("customer_mobile"):
merged["customer_mobile"] = rule_result["customer_mobile"]
rule_addr = rule_result.get("customer_address", "")
llm_addr = merged.get("customer_address", "")
if rule_addr and len(rule_addr) > len(llm_addr or ""):
merged["customer_address"] = rule_addr
return merged
def _build_fallback_result(self, rule_result: dict) -> dict:
"""当 LLM 不可用时,仅基于规则提取结果构建订单数据。
参数:
rule_result: 规则提取结果。
返回:
包含基础订单字段的字典,产品明细可能不完整。
"""
items: list[dict] = []
if rule_result.get("_raw_quantity"):
items.append({
"product_name": "",
"specification": "",
"unit": rule_result.get("_raw_unit", ""),
"quantity": rule_result["_raw_quantity"],
"sale_price": rule_result.get("_raw_price", 0),
"cost_price": None,
})
return {
"customer_name": rule_result.get("customer_name"),
"customer_mobile": rule_result.get("customer_mobile"),
"customer_address": rule_result.get("customer_address"),
"order_source": None,
"delivery_type": None,
"factory_id": None,
"remark": None,
"items": items,
}
def _get_product_groups(self, session: Session | None) -> list[dict]:
"""获取产品库数据,用于构建 LLM 提示词和模糊匹配。"""
if session is None:
return []
# 优先走 product_service有缓存和分组逻辑
try:
from backend.app.services.product_service import product_service
result = product_service.list_products({}, session)
groups = result.get("list", [])
if groups:
logger.info("[智能填单] _get_product_groups 返回 %d 个产品组", len(groups))
return groups
except Exception:
logger.warning("[智能填单] product_service 异常,尝试直查", exc_info=True)
# product_service 失败或返回空时,直查 Product 表兜底
try:
from backend.app.models.business import Product
stmt = select(Product).where(Product.deleted == 0).limit(100)
products = list(session.execute(stmt).scalars())
if products:
grouped: dict[str, dict] = {}
for p in products:
entry = grouped.setdefault(p.product_name, {
"product_name": p.product_name,
"product_id": p.id,
"specifications": [],
})
entry["specifications"].append({
"product_id": p.id,
"specification": p.specification,
"unit": p.unit,
"sale_price": float(p.sale_price or 0),
"cost_price": float(p.cost_price or 0),
})
logger.info("[智能填单] DB兜底: 从 Product 表直查到 %d 个产品组", len(grouped))
return list(grouped.values())
except Exception:
logger.warning("[智能填单] Product 表直查也失败", exc_info=True)
return []
_PRODUCT_ALIASES = {
"咖色": "咖啡色", "咖啡": "咖啡色",
"高温布": "高温布", "耐高温布": "高温布", "防火布": "高温布",
"网格带": "网格输送带",
"烧烤垫": "烧烤网格片",
"特氟龙": "特氟龙玻纤胶带",
}
# 输送带上下文关键词:命中任一则判定为输送带场景
_CONVEYOR_KEYWORDS = {
"包边", "红膜包边", "黑膜包边", "高温布包边", "凯夫拉包边",
"接头", "平接", "搭接", "斜接", "牛鼻子接头", "钢扣接头", "螺旋接头",
"上下盖布搭接", "上下盖布",
"导条", "凯夫拉导条", "硅胶导条", "四氟导条",
"柳丁", "鱼眼",
}
# 明确指向网格输送带的关键词
_CONVEYOR_MESH_KEYWORDS = {"网格", "网带", "特氟龙网格"}
# 胶带精确映射:关键词 -> 目标产品名
_TAPE_EXACT_MAP = {
"膜胶带": "特氟龙纯膜胶带",
}
_TAPE_FIBER_KEYWORDS = {"特氟龙胶带", "铁氟龙胶带", "小卷胶带"}
# 日用品高温布关键词(不含"高温布"本身,太泛会误命中工业场景)
_DAILY_CLOTH_KEYWORDS = {"拼豆", "烫布", "煤气灶垫", "烧烤垫"}
_PREFIX_STRIP_RE = re.compile(
r'^[\d.]+\s*(?:mm|cm|m|丝|g)\s*' # 厚度前缀: 0.13mm, 145g
r'|^(?:咖色|咖啡色|黑色|白色|灰色|红色|蓝色|绿色)' # 颜色前缀
)
# 数据库产品名带型号后缀,如 "高温布-SYY+5681",需去掉后缀再比较
_MODEL_SUFFIX_RE = re.compile(r'[-_][A-Za-z0-9+]+$')
# 厚度提取:匹配 "0.13mm"、"0.13"、"18丝" 等
_THICKNESS_RE = re.compile(r'(\d+\.\d+)\s*(?:mm|毫米|丝)?', re.IGNORECASE)
# 克重提取:匹配 "145g"、"145克" 等
_WEIGHT_RE = re.compile(r'(\d+)\s*(?:g|克|G)(?:\s|$|[^a-zA-Z])')
def _normalize_product_name(self, name: str) -> str:
"""去除产品名中的厚度/颜色前缀和型号后缀,提取核心材料名。"""
normalized = self._PREFIX_STRIP_RE.sub('', name).strip()
normalized = self._MODEL_SUFFIX_RE.sub('', normalized).strip()
return self._PRODUCT_ALIASES.get(normalized, normalized)
def _extract_thickness(self, text: str) -> str | None:
"""从文本中提取厚度值并归一化为无小数点字符串。
例: "0.13mm" -> "013", "0.08" -> "008"
"""
if not text:
return None
m = self._THICKNESS_RE.search(text)
if m:
return m.group(1).replace('.', '')
return None
def _extract_weight(self, text: str) -> str | None:
"""从文本中提取克重数值。
例: "145g" -> "145", "200克" -> "200"
"""
if not text:
return None
m = self._WEIGHT_RE.search(text)
if m:
return m.group(1)
return None
def _apply_rule_matching(self, input_name: str, input_spec: str,
item_remark: str, demand_spec: str,
flat_products: list[dict]) -> list[tuple[dict, float]]:
"""规则匹配层:对特定业务场景返回精确候选,避免依赖模糊匹配。
返回: [(product, score), ...] 或空列表(表示无规则命中,走通用匹配)。
"""
combined_text = f"{input_name} {input_spec} {demand_spec or ''} {item_remark or ''}"
# --- 输送带规则 ---
if any(kw in combined_text for kw in self._CONVEYOR_KEYWORDS):
is_mesh = any(kw in combined_text for kw in self._CONVEYOR_MESH_KEYWORDS)
target = "网格输送带" if is_mesh else "高温布"
hits = []
for prod in flat_products:
pname = prod["product_name"]
if target in pname or pname in target:
hits.append((prod, 0.95))
if hits:
return hits
# --- 胶带规则 ---
for kw, target in self._TAPE_EXACT_MAP.items():
if kw in combined_text:
hits = []
for prod in flat_products:
if target in prod["product_name"] or prod["product_name"] in target:
hits.append((prod, 0.95))
if hits:
return hits
if any(kw in combined_text for kw in self._TAPE_FIBER_KEYWORDS):
hits = []
for prod in flat_products:
if "特氟龙玻纤胶带" in prod["product_name"] or prod["product_name"] in "特氟龙玻纤胶带":
hits.append((prod, 0.95))
elif "胶带" in prod["product_name"] and "纯膜" not in prod["product_name"]:
hits.append((prod, 0.85))
if hits:
return hits
# --- 日用品高温布规则 ---
if any(kw in combined_text for kw in self._DAILY_CLOTH_KEYWORDS):
weight = self._extract_weight(combined_text)
hits = []
for prod in flat_products:
pname = prod["product_name"]
spec = prod.get("specification", "")
# 日用品方向产品优先
is_daily = "日用品" in pname or "日用品" in prod.get("category_name", "")
weight_match = weight and weight in spec
if is_daily:
score = 0.95 if weight_match else 0.85
hits.append((prod, score))
elif weight_match and ("高温布" in pname or "高温布" in spec):
hits.append((prod, 0.75))
if hits:
hits.sort(key=lambda x: x[1], reverse=True)
return hits
# --- 揉面垫 / 硅胶烤垫规则:厚度主导 ---
if "揉面垫" in combined_text or "硅胶烤垫" in combined_text:
thickness = self._extract_thickness(combined_text)
hits = []
for prod in flat_products:
pname = prod["product_name"]
spec = prod.get("specification", "")
if "揉面垫" in combined_text and "揉面垫" not in pname:
continue
if "硅胶烤垫" in combined_text and "硅胶烤垫" not in pname:
continue
if thickness and thickness in spec.replace('.', ''):
hits.append((prod, 0.95))
else:
hits.append((prod, 0.6))
if hits:
hits.sort(key=lambda x: x[1], reverse=True)
return hits
return []
def _score_with_thickness(self, input_name: str, input_spec: str,
prod: dict, base_score: float) -> float:
"""在基础分数上叠加厚度/克重匹配加权。"""
combined = f"{input_name} {input_spec}"
thickness = self._extract_thickness(combined)
weight = self._extract_weight(combined)
spec = prod.get("specification", "")
bonus = 0.0
if thickness and thickness in spec.replace('.', ''):
bonus = 0.3
elif weight and weight in spec:
bonus = 0.25
return min(base_score + bonus, 1.0)
def _match_products(self, items: list[dict],
product_groups: list[dict],
session: Session | None) -> list[dict]:
"""对解析出的产品明细进行库内模糊匹配。
优先走规则匹配(输送带/胶带/揉面垫等),再走通用模糊匹配,
最后 DB 兜底。
参数:
items: 解析出的产品明细列表。
product_groups: 产品库分组数据。
session: 数据库会话。
返回:
匹配结果列表,每个元素包含 input_name、input_specification、candidates。
"""
flat_products: list[dict] = []
for group in product_groups:
for spec in group.get("specifications", []):
flat_products.append({
"product_id": spec["product_id"],
"product_name": group["product_name"],
"specification": spec["specification"],
"unit": spec["unit"],
"sale_price": spec.get("sale_price", 0),
"cost_price": spec.get("cost_price", 0),
"category_id": group.get("category_id"),
"category_name": group.get("category_name", ""),
})
matches: list[dict] = []
for item in items:
input_name = item.get("product_name", "")
input_spec = item.get("specification", "")
item_remark = item.get("remark", "")
demand_spec = item.get("demand_specification", "")
if not input_name:
continue
normalized_input = self._normalize_product_name(input_name)
# 1. 规则匹配层(业务场景优先)
rule_hits = self._apply_rule_matching(
input_name, input_spec, item_remark, demand_spec, flat_products
)
scored: list[tuple] = []
if rule_hits:
scored = rule_hits
else:
# 2. 通用模糊匹配
for prod in flat_products:
normalized_prod = self._normalize_product_name(prod["product_name"])
if normalized_input == normalized_prod:
base = 1.0
else:
name_score = SequenceMatcher(None, normalized_input, normalized_prod).ratio()
if normalized_prod in normalized_input or normalized_input in normalized_prod:
name_score = max(name_score, 0.85)
spec_score = SequenceMatcher(None, input_spec, prod["specification"]).ratio() if input_spec else 0
base = name_score * 0.5 + spec_score * 0.3
# 叠加厚度/克重加权
final_score = self._score_with_thickness(input_name, input_spec, prod, base)
scored.append((prod, final_score))
scored.sort(key=lambda x: x[1], reverse=True)
candidates = [
{
"product_id": prod["product_id"],
"product_name": prod["product_name"],
"specification": prod["specification"],
"unit": prod["unit"],
"sale_price": prod["sale_price"],
"cost_price": prod["cost_price"],
"match_score": round(score, 2),
}
for prod, score in scored[:3] if score > 0.25
]
# 3. 内存匹配失败时,直接查数据库兜底
if not candidates and session is not None:
logger.info("[智能填单] 内存匹配失败, input=%s, normalized=%s, 尝试DB兜底",
input_name, normalized_input)
candidates = self._db_fallback_match(input_name, normalized_input, session)
logger.info("[智能填单] DB兜底结果: %s", [(c["product_name"], c["product_id"]) for c in candidates])
matches.append({
"input_name": input_name,
"input_specification": input_spec,
"candidates": candidates,
})
return matches
def _db_fallback_match(self, input_name: str, normalized_input: str,
session: Session) -> list[dict]:
"""内存匹配失败时,直接查 Product 表兜底。"""
from backend.app.models.business import Product
try:
# 用 LIKE 按产品名模糊搜索
stmt = select(Product).where(
Product.deleted == 0,
Product.product_name.contains(normalized_input),
).limit(5)
products = list(session.execute(stmt).scalars())
if not products:
# 试一下原始输入名(可能 LLM 返回了带型号的全名)
stmt = select(Product).where(
Product.deleted == 0,
Product.product_name.contains(input_name),
).limit(5)
products = list(session.execute(stmt).scalars())
results = []
for p in products:
results.append({
"product_id": p.id,
"product_name": p.product_name,
"specification": p.specification,
"unit": p.unit,
"sale_price": float(p.sale_price or 0),
"cost_price": float(p.cost_price or 0),
"match_score": 0.8,
})
return results[:3]
except Exception:
return []
def _validate_parsed_order(self, order: dict) -> list[str]:
"""校验解析后的订单数据完整性,生成警告列表。
检查项:客户姓名、手机号、产品明细、产品名称、数量。
参数:
order: 解析后的订单字典。
返回:
警告信息列表,无警告时返回空列表。
"""
warnings: list[str] = []
if not order.get("customer_name"):
warnings.append("未识别到客户姓名")
if not order.get("customer_mobile"):
warnings.append("未识别到客户手机号")
if not order.get("items"):
warnings.append("未识别到产品明细")
for item in order.get("items", []):
if not item.get("product_name"):
warnings.append("存在产品名称为空的明细行")
if not item.get("quantity"):
warnings.append(f"产品 {item.get('product_name', '?')} 未识别到数量")
return warnings
def _calc_confidence(self, order: dict, warnings: list[str],
ocr_confidence: float | None = None) -> float:
"""计算综合置信度分数。
计算规则:以 OCR 置信度(默认 1.0)为基础,每个警告扣 0.1
每个缺少单价的产品扣 0.05,最低不低于 0.1。
参数:
order: 解析后的订单字典。
warnings: 警告列表。
ocr_confidence: OCR 置信度,可选。
返回:
综合置信度0.1 ~ 1.0)。
"""
base = 1.0
if ocr_confidence is not None:
base = ocr_confidence
base -= len(warnings) * 0.1
if order.get("items"):
for item in order["items"]:
if not item.get("sale_price"):
base -= 0.05
return round(max(base, 0.1), 2)
# ------------------------------------------------------------------
# 单位标准化:统一 demand_specification 中的单位格式
# ------------------------------------------------------------------
# 长度单位映射:中文/英文 -> 标准单位
_LENGTH_UNIT_MAP = {
"cm": "cm", "厘米": "cm",
"m": "m", "": "m",
"mm": "mm", "毫米": "mm",
}
# 重量单位映射:中文/英文 -> 标准单位
_WEIGHT_UNIT_MAP = {
"g": "g", "": "g",
"kg": "kg", "公斤": "kg",
"jin": "", "": "",
}
# 长度尺寸匹配正则:数字 + 可选单位 × 数字 + 可选单位
_LENGTH_SPEC_RE = re.compile(
r'([\d.]+)\s*(cm|厘米|m|米|mm|毫米)?\s*[×xX*]\s*([\d.]+)\s*(cm|厘米|m|米|mm|毫米)?',
re.IGNORECASE
)
# 重量匹配正则:数字 + 单位
_WEIGHT_SPEC_RE = re.compile(
r'([\d.]+)\s*(g|克|kg|公斤|斤)',
re.IGNORECASE
)
def _standardize_demand_specification(self, spec: str) -> str:
"""标准化需求规格中的单位格式。
将所有长度单位统一转换为 "长m×宽m" 格式,
将所有重量单位统一转换为 "数值+标准单位" 格式。
参数:
spec: 原始需求规格字符串。
返回:
标准化后的字符串。
"""
if not spec:
return spec
# 尝试匹配长度尺寸(宽×高)
length_match = self._LENGTH_SPEC_RE.search(spec)
if length_match:
return self._standardize_length_spec(spec, length_match)
# 尝试匹配重量
weight_match = self._WEIGHT_SPEC_RE.search(spec)
if weight_match:
return self._standardize_weight_spec(spec, weight_match)
# 无法识别,返回原值
return spec
def _standardize_length_spec(self, spec: str, match: re.Match) -> str:
"""标准化长度尺寸为 "长m×宽m" 格式。"""
v1 = float(match.group(1))
u1 = self._LENGTH_UNIT_MAP.get((match.group(2) or "cm").lower(), "cm")
v2 = float(match.group(3))
u2 = self._LENGTH_UNIT_MAP.get((match.group(4) or "cm").lower(), "cm")
# 转换为米
def to_m(v, u):
if u == "m":
return v
elif u == "mm":
return v / 1000
else: # cm
return v / 100
m1 = to_m(v1, u1)
m2 = to_m(v2, u2)
# 格式化:保留必要小数位
def format_m(v):
if v == int(v):
return str(int(v))
# 去除尾部多余的0
return f"{v:.4f}".rstrip("0").rstrip(".")
return f"{format_m(m1)}m×{format_m(m2)}m"
def _standardize_weight_spec(self, spec: str, match: re.Match) -> str:
"""标准化重量单位。"""
value = float(match.group(1))
unit = self._WEIGHT_UNIT_MAP.get(match.group(2).lower(), match.group(2))
# 格式化数值
if value == int(value):
value_str = str(int(value))
else:
value_str = str(value)
return f"{value_str}{unit}"
def _get_factory_options(self, session: Session | None) -> list[dict]:
"""获取工厂(供应商)下拉选项列表,供前端智能填单表单使用。
参数:
session: 数据库会话,为 None 时返回空列表。
返回:
包含 valuesupplier_id和 labelsupplier_name的字典列表。
"""
if session is None:
return []
try:
from backend.app.services.supplier_service import supplier_service
result = supplier_service.list_suppliers(
session, {"supplier_type": "factory"}
)
return [
{"value": item["supplier_id"], "label": item["supplier_name"]}
for item in result.get("list", [])
]
except Exception:
return []
def parse_customer(self, payload: dict, session: Session | None = None) -> dict:
"""客户信息智能解析入口:从文本中提取客户信息。
复用订单解析的文本预处理和规则提取逻辑,专门用于提取客户相关信息。
处理流程:
1. 文本预处理(去噪声)
2. 规则提取(正则提取手机号、地址)
3. LLM 结构化解析(提取客户姓名等)
4. 客户库预匹配(通过手机号查找已有客户)
参数:
payload: 请求数据字典,必含 text 字段。
session: 数据库会话,为 None 时跳过数据库查询。
返回:
包含 parsed_customer、raw_text、confidence、warnings 的字典。
被调用方ai 路由(客户信息解析接口)。
"""
raw_text = payload.get("text", "").strip()
# 1. 文本预处理
preprocessor = TextPreprocessor()
raw_text = preprocessor.preprocess(raw_text)
if not raw_text:
raise AppException(
code=ErrorCode.PARAM_ERROR,
message="无法提取到文本内容",
status_code=400,
)
# 2. 规则提取(手机号、地址)
rule_extractor = RuleExtractor()
rule_result = rule_extractor.extract(raw_text)
# 3. LLM 结构化解析(提取客户姓名等)
settings = self.settings
llm_parsed = {}
# 构造简化的 prompt只提取客户信息
system_prompt = """你是客户信息提取助手。从用户提供的文本中提取客户信息。
## 重要规则
1. 只提取文本中明确提到的信息,不要猜测或推断
2. 如果某个字段无法识别,返回空字符串 ""
3. customer_name 是客户姓名/称呼,不是公司名、订单号或其他编号
4. customer_mobile 必须是11位手机号码1开头的数字
5. customer_address 只包含地理地址信息(省/市/区/县/镇/街道/路/号等),不要包含其他内容
## 输出格式
严格按以下 JSON 格式输出,不要添加任何其他内容:
{"customer_name":"","customer_mobile":"","customer_address":"","customer_type":"","settlement_type":"","remark":""}
## 字段说明
- customer_name客户姓名或称呼"张三""李总"
- customer_mobile11位手机号
- customer_address收货地址只包含地理信息
- customer_type客户类型只能是 "channel"/"retail"/"project" 之一,不确定则留空
- settlement_type结算方式只能是 "monthly"/"cash"/"delivered" 之一,不确定则留空
- remark备注信息"""
payload_data = json.dumps({
"model": settings.llm_parse_model,
"messages": [
{"role": "system", "content": system_prompt},
{"role": "user", "content": f"请从以下文本中提取客户信息:\n\n{raw_text}"},
],
"temperature": 0.1,
"max_tokens": 512,
}).encode("utf-8")
headers = {
"Content-Type": "application/json",
"Authorization": f"Bearer {settings.llm_parse_api_key or settings.aliyun_ai_access_key_id}",
}
try:
req = request.Request(
url=settings.llm_parse_api_url, data=payload_data, headers=headers, method="POST"
)
with request.urlopen(req, timeout=30) as resp:
result = json.loads(resp.read().decode("utf-8"))
content = result["choices"][0]["message"]["content"]
logger.info("[客户信息解析] LLM 返回: %s", content)
# 解析 JSON
llm_parsed = self._extract_customer_json(content)
except Exception as e:
logger.warning("[客户信息解析] LLM 解析失败: %s", str(e))
llm_parsed = {}
# 4. 合并规则提取和 LLM 解析结果
parsed_customer = {}
# 手机号:优先使用规则提取的(更可靠)
if rule_result.get("customer_mobile"):
parsed_customer["customer_mobile"] = rule_result["customer_mobile"]
elif llm_parsed.get("customer_mobile"):
parsed_customer["customer_mobile"] = llm_parsed["customer_mobile"]
# 地址:优先使用规则提取的(更可靠)
if rule_result.get("customer_address"):
parsed_customer["customer_address"] = rule_result["customer_address"]
elif llm_parsed.get("customer_address"):
parsed_customer["customer_address"] = llm_parsed["customer_address"]
# 客户姓名:使用 LLM 解析的结果
if llm_parsed.get("customer_name"):
parsed_customer["customer_name"] = llm_parsed["customer_name"]
# 其他字段:使用 LLM 解析的结果
for key in ("customer_type", "settlement_type", "remark"):
if llm_parsed.get(key):
parsed_customer[key] = llm_parsed[key]
# 5. 客户库预匹配(通过手机号查找已有客户)
warnings = []
if session and parsed_customer.get("customer_mobile"):
try:
from backend.app.repositories.customer_repository import CustomerRepository
customer_repo = CustomerRepository(session)
existing = customer_repo.find_by_mobile(parsed_customer["customer_mobile"])
if existing:
parsed_customer["existing_customer"] = {
"customer_id": existing.customer_id,
"customer_name": existing.customer_name,
"mobile": existing.mobile,
}
warnings.append(f"手机号 {parsed_customer['customer_mobile']} 已存在于客户库中({existing.customer_name}")
except Exception as e:
logger.warning("[客户信息解析] 客户库匹配失败: %s", str(e))
# 6. 计算置信度
confidence = self._calculate_customer_confidence(parsed_customer, warnings)
return {
"parsed_customer": parsed_customer,
"raw_text": raw_text,
"confidence": confidence,
"warnings": warnings,
}
def _extract_customer_json(self, text: str) -> dict:
"""从 LLM 响应中提取客户信息 JSON。
处理 LLM 可能包裹在 markdown 代码块中的情况,
以及 LLM 响应被截断导致 JSON 不完整的常见场景。
参数:
text: LLM 原始响应文本。
返回:
解析后的客户信息字典。
"""
text = text.strip()
if text.startswith("```"):
text = text.split("\n", 1)[1] if "\n" in text else text[3:]
text = text.rsplit("```", 1)[0]
text = text.strip()
if text.startswith("json"):
text = text[4:].strip()
try:
return json.loads(text)
except json.JSONDecodeError:
# LLM 响应可能被截断,尝试截断到最后一个完整 } 修复
last_brace = text.rfind("}")
while last_brace > 0:
try:
return json.loads(text[:last_brace + 1])
except json.JSONDecodeError:
last_brace = text.rfind("}", 0, last_brace - 1)
return {}
def _calculate_customer_confidence(self, customer: dict, warnings: list) -> float:
"""计算客户信息解析的置信度。
根据识别到的字段数量和质量计算置信度分数。
参数:
customer: 解析出的客户信息字典。
warnings: 警告信息列表。
返回:
0.1 到 1.0 之间的置信度分数。
"""
base = 0.8
# 有客户姓名加分
if customer.get("customer_name"):
base += 0.1
# 有手机号加分(手机号是关键字段)
if customer.get("customer_mobile"):
base += 0.1
# 有地址加分
if customer.get("customer_address"):
base += 0.05
# 有警告扣分(如客户已存在)
base -= len(warnings) * 0.1
return round(max(base, 0.1), 2)
ai_service = AIService()