Files
Cloud-Tour-to-Libo/app/api/doc_restructure.py

321 lines
13 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.

"""文档 MD 的 ②结构化重排 + ③事实级校验(供 8765 上传转换后调用)。
来自 doc-eval 压测验证过的逻辑,纯函数、无外部副作用:
· restructure_md(md) 把"忠实转换"的宽表 MD 折叠成「每产品一节 + 干净明细表」的结构化 MD
· validate(src, out) 事实级校验:数字保真(幻觉)、按产品价格不串位、内容不被删减、业务规则
"""
from __future__ import annotations
import re
from collections import Counter, defaultdict
NUM = re.compile(r"\d+(?:\.\d+)?")
def _clean(s: str) -> str:
s = (s or "").replace("<br>", " ")
s = re.sub(r"\s+", " ", s).strip()
n = len(s)
if n > 20: # 折叠"整串精确对折重复"(源 Excel 合并单元格常存双份)
for mid in range(n // 2 - 2, n // 2 + 3):
if 0 < mid < n:
left, right = s[:mid].strip(), s[mid:].strip()
if left and len(left) > 10 and left == right:
return left
return s
def parse_blocks(md: str):
lines = md.splitlines()
blocks, i, n, buf = [], 0, len(lines), []
def flush():
if buf:
for para in re.split(r"\n\s*\n", "\n".join(buf)):
if para.strip():
blocks.append(("text", para.strip()))
buf.clear()
def is_row(s):
return s.strip().startswith("|") and s.count("|") >= 2
def is_sep(s):
return bool(re.match(r"^\s*\|[\s:|-]+\|\s*$", s)) and "-" in s
while i < n:
if is_row(lines[i]) and i + 1 < n and is_sep(lines[i + 1]):
flush()
header = [c.strip() for c in lines[i].strip().strip("|").split("|")]
j, rows = i + 2, []
while j < n and is_row(lines[j]) and not is_sep(lines[j]):
row = [c.strip() for c in lines[j].strip().strip("|").split("|")]
rows.append((row + [""] * len(header))[: len(header)])
j += 1
blocks.append(("table", header, rows))
i = j
else:
buf.append(lines[i])
i += 1
flush()
return blocks
def classify_columns(header, rows):
n, doc, keys, data = len(rows), [], [], []
for c in range(len(header)):
vals = [r[c] for r in rows]
nonempty = [v for v in vals if v]
if not nonempty:
continue
distinct = len(set(nonempty))
changes = sum(1 for k in range(1, n) if vals[k] != vals[k - 1])
if distinct <= 1:
doc.append(c)
elif n >= 4 and changes <= max(1, int(n * 0.45)):
keys.append(c)
else:
data.append(c)
return doc, keys, data
def _restructure_table(header, rows):
doc, keys, data = classify_columns(header, rows)
docmeta = {}
for c in doc:
v = next((r[c] for r in rows if r[c]), "")
if v:
docmeta[_clean(header[c])] = _clean(v)
if not keys or not data:
keep = [c for c in range(len(header)) if any(r[c] for r in rows)]
out = ["| " + " | ".join(_clean(header[c]) for c in keep) + " |",
"| " + " | ".join("---" for _ in keep) + " |"]
out += ["| " + " | ".join(r[c] for c in keep) + " |" for r in rows]
return docmeta, ["\n".join(out)]
title_col = max(keys, key=lambda c: len({r[c] for r in rows if r[c]}))
attr_keys = [c for c in keys if c != title_col]
sections, cur_key, grp = [], object(), []
def emit(group):
if not group:
return
seg = [f"## {_clean(group[0][title_col]) or '(未命名)'}"]
for c in attr_keys:
v = _clean(group[0][c])
if v:
seg.append(f"- {_clean(header[c])}: {v}")
seg.append("")
seg.append("| " + " | ".join(_clean(header[c]) for c in data) + " |")
seg.append("| " + " | ".join("---" for _ in data) + " |")
for r in group:
seg.append("| " + " | ".join(_clean(r[c]) for c in data) + " |")
sections.append("\n".join(seg))
for r in rows:
kt = tuple(r[c] for c in keys)
if kt != cur_key:
emit(grp)
grp, cur_key = [], kt
grp.append(r)
emit(grp)
return docmeta, sections
def restructure_md(md: str, source: str = "") -> str:
blocks = parse_blocks(md)
docmeta, body, seen, title = {}, [], set(), ""
for blk in blocks:
if blk[0] == "text":
para = blk[1]
key = re.sub(r"\s+", "", para)
if key in seen:
continue
seen.add(key)
if not title and not para.startswith(("#", "|")):
title = para.split("\n")[0]
if "\n" not in para:
continue
if re.fullmatch(r"##\s*Sheet\d+\s*", para):
continue
body.append(("text", para))
else:
dm, secs = _restructure_table(blk[1], blk[2])
docmeta.update(dm)
body.append(("secs", secs))
out = ["## 基本信息", ""]
if source:
out.append(f"- 来源文件: {source}")
for k, v in docmeta.items():
out.append(f"- {k}: {v}")
out.append("")
if title:
out.append(f"# {title}\n")
for kind, val in body:
if kind == "text" and not val.startswith("#"):
out.append(val + "\n")
elif kind == "text":
out.append(val + "\n")
else:
out.extend(s + "\n" for s in val)
return "\n".join(out).strip() + "\n"
# ---------- ②-LLM 语义结构化(类型自适应,用于行程单/复杂件;价格表走免费确定性)----------
STRUCTURE_PROMPT = """你是旅游 ToB 文档结构化助手。把下面从源文件转出的原始 Markdown重排成"结构化科学 MD",供后期 LLM 检索入库。
【先判断文档类型,套用对应结构】
· 报价单/价目表:顶部「## 基本信息」表(供应商/有效期/出发地/币种);每个产品一个 `## 产品:<名称>` 小节(拼团类型/起订人数/景区小交通费/退费政策等);价格用「房型 | 成人价 | 儿童价 | 单房差」表;通用规则、参考酒店各自独立小节。
· 行程单:顶部「## 基本信息」表(产品名/天数/车型/出发地);按天分 `## D1 <当天路线>`、`## D2 …` 小节,每天完整保留 行程详情/餐食(早中晚)/住宿/交通;费用包含、费用不含、购物场所、预订须知、温馨提示 各自独立小节。
· 资源表(酒店/车辆/餐厅):每个资源一个小节 + 属性。
· 其它:用标题分层、表格如实保留。
【铁律,违反即事故】
1. 完整保留原文每一条信息、逐条照搬——严禁摘要、严禁概括、严禁删减合并、严禁省略。这是入库用的全量数据;输出篇幅应与原文相当。
2. 所有数值/价格/日期/时间/电话严格来自原文,一字不改、不新增、不计算、不脑补;看不清就原样抄。
3. 可纠正明显的"表头误标"(如"车型"列内容其实是"N人拼小团"→改标「拼团类型」),但只改标签、不改数值。
4. 元数据用可见的「## 基本信息」表,不要 YAML frontmatter。
只输出 MD不要任何解释。"""
def resolve_llm_key() -> str | None:
"""LLM key 解析顺序:环境变量 LLM_API_KEY / DEEPSEEK_API_KEY → ~/.ark_key 文件。"""
import os
k = os.environ.get("LLM_API_KEY") or os.environ.get("DEEPSEEK_API_KEY")
if k:
return k.strip()
try:
with open(os.path.expanduser("~/.ark_key")) as f:
return f.read().strip() or None
except Exception:
return None
def structure_llm(md: str, api_key: str) -> str:
"""调 LLM(默认火山豆包,可用 LLM_BASE_URL/LLM_MODEL 覆盖)做类型自适应结构化,带重试退避。"""
import os
import time
import requests
base_url = os.environ.get("LLM_BASE_URL", "https://ark.cn-beijing.volces.com/api/v3")
model = os.environ.get("LLM_MODEL", "doubao-seed-2-0-lite-260428")
payload = {"model": model, "temperature": 0,
"messages": [{"role": "system", "content": STRUCTURE_PROMPT},
{"role": "user", "content": md}]}
last = None
for attempt in range(4):
try:
r = requests.post(f"{base_url}/chat/completions",
headers={"Authorization": f"Bearer {api_key}"},
json=payload, timeout=180)
r.raise_for_status()
return r.json()["choices"][0]["message"]["content"]
except requests.exceptions.RequestException as e:
last = e
if attempt < 3:
time.sleep(2 * (attempt + 1))
raise last
# ---------- ③ 事实级校验 ----------
def _wide_product_numbers(md: str) -> dict:
out = defaultdict(set)
for blk in parse_blocks(md):
if blk[0] != "table":
continue
header, rows = blk[1], blk[2]
_doc, keys, _data = classify_columns(header, rows)
if not keys:
continue
tcol = max(keys, key=lambda c: len({r[c] for r in rows if r[c]}))
for r in rows:
for cell in r:
out[_clean(r[tcol])].update(NUM.findall(cell))
return out
def _section_numbers(md: str) -> dict:
out, cur = defaultdict(set), None
for line in md.splitlines():
m = re.match(r"^##\s+(?:产品[:])?\s*(.+)$", line)
if m:
cur = _clean(m.group(1))
elif cur and line.count("|") >= 2:
out[cur].update(NUM.findall(line))
return out
def _has_price_table(md: str) -> bool:
for blk in parse_blocks(md):
if blk[0] != "table":
continue
_doc, keys, data = classify_columns(blk[1], blk[2])
if keys and len(data) >= 2:
vals = [r[c] for r in blk[2] for c in data]
nums = [float(v) for v in vals if re.fullmatch(r"\d+(?:\.\d+)?", v)]
if nums and sum(x >= 100 for x in nums) >= len(nums) * 0.5:
return True
return False
def _strip_fm(md: str) -> str:
m = re.match(r"^---\n.*?\n---\n", md, re.S)
return md[m.end():] if m else md
def validate(src_md: str, out_md: str) -> dict:
out_body = _strip_fm(out_md)
# 去掉"来源文件"行(文件名里的日期是元数据、不是业务数据,不参与数字校验)
out_body = "\n".join(l for l in out_body.splitlines() if not l.lstrip().startswith("- 来源文件"))
src_nums, out_nums = Counter(NUM.findall(src_md)), Counter(NUM.findall(out_body))
invented = {n: c for n, c in (out_nums - src_nums).items() if n not in src_nums}
misassigned: dict = {}
if _has_price_table(src_md):
src_by = _wide_product_numbers(src_md)
for prod, nums in _section_numbers(out_md).items():
match = next((src_by[k] for k in src_by if k and (k in prod or prod in k)), None)
if match is None:
continue
bad = {n for n in (nums - match) if float(n) >= 100 and not 1900 <= float(n) <= 2099}
if bad:
misassigned[prod] = sorted(bad)
shared = {n for n, c in Counter(n for ns in misassigned.values() for n in ns).items()
if c >= max(3, int(len(misassigned) * 0.5))}
misassigned = {p: [n for n in ns if n not in shared]
for p, ns in misassigned.items() if [n for n in ns if n not in shared]}
def _prose(md):
return sum(len(re.sub(r"\s", "", l)) for l in md.splitlines() if l.count("|") < 2)
sp, op = _prose(src_md), _prose(out_body)
content_loss = f"②正文{op}字 / ①{sp}字,疑似删减" if (sp > 300 and op < sp * 0.5) else None
warn = []
for line in out_md.splitlines():
if line.count("|") >= 3:
cells = [c.strip() for c in line.strip("|").split("|")]
nums = [c for c in cells if re.fullmatch(r"\d+(?:\.\d+)?", c)]
if len(nums) >= 2 and float(nums[0]) < float(nums[1]):
warn.append(f"成人价<儿童价? {line.strip()[:50]}")
return {
"passed": not invented and not misassigned and not content_loss,
"invented": invented,
"misassigned": misassigned,
"content_loss": content_loss,
"business_warn": warn[:5],
}
def build_structured(md: str, source: str = "") -> tuple[str, dict, str]:
"""8765 用:自动路由结构化。价格表且确定性已过 → 免费确定性重排;
行程单/复杂件/确定性没过 → 若有 key 则上 LLM(类型自适应)。返回 (结构化MD, 校验, 模式)。"""
deterministic = restructure_md(md, source=source)
rep = validate(md, deterministic)
key = resolve_llm_key()
if key and not (_has_price_table(md) and rep["passed"]):
try:
llm_md = structure_llm(md, key)
if llm_md and llm_md.strip():
return llm_md, validate(md, llm_md), "llm"
except Exception: # noqa: BLE001
pass
return deterministic, rep, "deterministic"