优化切片,增加父节点章节、图片描述

This commit is contained in:
Defeng 2026-07-20 20:40:39 +08:00
parent 626139226d
commit 129846dc17
2 changed files with 456 additions and 200 deletions

View File

@ -24,6 +24,7 @@ import threading
import aiofiles import aiofiles
import httpx import httpx
import logging import logging
from resetlevel import reset_textlevel
from rapidocr_onnxruntime import RapidOCR from rapidocr_onnxruntime import RapidOCR
from PIL import Image, ImageOps, UnidentifiedImageError from PIL import Image, ImageOps, UnidentifiedImageError
from config import ( from config import (
@ -51,6 +52,8 @@ _image_ocr_cache: Dict[str, str] = {}
_image_ocr_cache_lock = threading.Lock() _image_ocr_cache_lock = threading.Lock()
_api_url_inflight = [0] * len(API_URLS) _api_url_inflight = [0] * len(API_URLS)
_api_url_lock = threading.Lock() _api_url_lock = threading.Lock()
IMAGE_OCR_PREFIX = "图片中文本内容为:"
IMAGE_OCR_PREFIXES = (IMAGE_OCR_PREFIX, "图片包含的文字内容为:")
def _image_mime_type(suffix: str) -> str: def _image_mime_type(suffix: str) -> str:
@ -250,7 +253,7 @@ def extract_image_text_sync(image_path: Union[str, PathLib]) -> dict:
ocr_result, elapse = ocr_engine(resolved_path) ocr_result, elapse = ocr_engine(resolved_path)
texts = [item[1] for item in (ocr_result or [])] texts = [item[1] for item in (ocr_result or [])]
full_text = "图片包含的文字内容为:" + ",".join(texts) if texts else "" full_text = IMAGE_OCR_PREFIX + ",".join(texts) if texts else ""
with _image_ocr_cache_lock: with _image_ocr_cache_lock:
_image_ocr_cache[resolved_path] = full_text _image_ocr_cache[resolved_path] = full_text
return { return {
@ -261,7 +264,7 @@ def extract_image_text_sync(image_path: Union[str, PathLib]) -> dict:
def build_image_markdown_with_ocr(markdown_image: str, local_image_path: Optional[PathLib]) -> str: def build_image_markdown_with_ocr(markdown_image: str, local_image_path: Optional[PathLib]) -> str:
if "图片包含的文字内容为:" in markdown_image: if any(prefix in markdown_image for prefix in IMAGE_OCR_PREFIXES):
return markdown_image return markdown_image
if not local_image_path: if not local_image_path:
return markdown_image return markdown_image
@ -278,6 +281,179 @@ def build_image_markdown_with_ocr(markdown_image: str, local_image_path: Optiona
return f"{markdown_image}\n{ocr_full_text}" return f"{markdown_image}\n{ocr_full_text}"
async def get_image_ocr_text(local_image_path: Optional[PathLib]) -> str:
if not local_image_path:
return ""
try:
ocr_data = await extract_image_text(local_image_path)
except Exception as exc:
logger.warning(f"Failed to OCR local image {local_image_path}: {exc}")
return ""
return (ocr_data.get("full_text") or "").strip()
def stringify_content_field(value: Any) -> str:
if value is None:
return ""
if isinstance(value, list):
return "\n".join(str(item) for item in value if str(item).strip())
if isinstance(value, dict):
return json.dumps(value, ensure_ascii=False)
return str(value)
def clean_table_body(table_body: Any) -> str:
table_body = stringify_content_field(table_body)
for pattern in [' colspan="1"', ' rowspan="1"', " colspan='1'", " rowspan='1'", " colspan=1", " rowspan=1"]:
table_body = table_body.replace(pattern, "")
return table_body
def join_nonempty_parts(*parts: Any) -> str:
return "\n".join(part for part in (stringify_content_field(part).strip() for part in parts) if part)
def record_to_chunk_text(ins: Dict[str, Any]) -> str:
record_type = ins.get("type")
ocr_text = ins.get("ocr_text", "")
if record_type == "text":
text = stringify_content_field(ins.get("text", ""))
level = ins.get("text_level")
if level is not None and isinstance(level, int) and level > 0:
text = "#" * min(level, 6) + " " + text
return join_nonempty_parts(ocr_text, text)
if record_type == "table":
return join_nonempty_parts(
ocr_text,
ins.get("table_caption"),
clean_table_body(ins.get("table_body", "")),
ins.get("table_footnote"),
)
if record_type == "image":
return join_nonempty_parts(
ocr_text,
ins.get("image_caption"),
ins.get("img_path"),
ins.get("image_footnote"),
)
if record_type == "chart":
return join_nonempty_parts(
ocr_text,
ins.get("chart_caption"),
ins.get("content"),
ins.get("img_path"),
ins.get("chart_footnote"),
)
if record_type == "equation":
return join_nonempty_parts(
ocr_text,
ins.get("text"),
ins.get("img_path"),
)
if record_type == "code":
return join_nonempty_parts(
ocr_text,
ins.get("code_caption"),
ins.get("code_body"),
ins.get("code_footnote"),
)
if record_type == "list":
return join_nonempty_parts(ocr_text, ins.get("list_items"))
if record_type in {"discarded", "header", "footer", "page_number"}:
return join_nonempty_parts(ocr_text, ins.get("text"))
return join_nonempty_parts(ocr_text, ins.get("text"))
def is_supported_chunk_record(ins: Dict[str, Any]) -> bool:
return ins.get("type") in {
"text",
"table",
"image",
"chart",
"equation",
"code",
"list",
"discarded",
"header",
"footer",
"page_number",
}
def append_caption_ocr(ins: Dict[str, Any], ocr_texts: List[str]) -> None:
caption_field = {
"image": "image_caption",
"table": "table_caption",
"chart": "chart_caption",
"code": "code_caption",
}.get(ins.get("type"))
if not caption_field:
caption_field = "ocr_text"
caption = ins.get(caption_field, [])
existing_text = stringify_content_field(caption)
if caption_field == "ocr_text":
existing_text = join_nonempty_parts(existing_text, ins.get("text"))
additions = []
for text in ocr_texts:
text = (text or "").strip()
text_without_prefix = text
for prefix in IMAGE_OCR_PREFIXES:
if text_without_prefix.startswith(prefix):
text_without_prefix = text_without_prefix[len(prefix):].strip()
break
if (
text
and text not in existing_text
and text_without_prefix not in existing_text
and text not in additions
):
additions.append(text)
if not additions:
return
if caption_field == "ocr_text":
ins[caption_field] = join_nonempty_parts(additions, caption)
return
if isinstance(caption, list):
ins[caption_field] = additions + caption
elif caption is None or not str(caption).strip():
ins[caption_field] = additions
else:
ins[caption_field] = "\n".join(additions) + "\n" + str(caption).lstrip()
def build_image_markdown(markdown_or_path: str, prefix: str) -> Tuple[str, Optional[PathLib]]:
img_path = str(markdown_or_path).strip()
if not img_path:
return "", None
if img_path.strip().startswith("!["):
return img_path, resolve_local_image_path(img_path)
p = PurePosixPath(img_path)
if p.parts and p.parts[0] == "images":
relative_path = str(PurePosixPath(*p.parts[1:]))
return f"![图片]({prefix}{relative_path})", resolve_local_image_path(relative_path)
return f"![图片]({prefix}{img_path})", resolve_local_image_path(img_path)
def image_base64_to_data_url(image_base64: Union[str, bytes], suffix: str) -> str: def image_base64_to_data_url(image_base64: Union[str, bytes], suffix: str) -> str:
if isinstance(image_base64, bytes): if isinstance(image_base64, bytes):
raw = base64.b64encode(image_base64).decode("ascii") raw = base64.b64encode(image_base64).decode("ascii")
@ -602,76 +778,14 @@ def split_content_bbox(data, max_length=8000):
bbox = ins.get("bbox") bbox = ins.get("bbox")
page_idx = ins.get("page_idx", -1) page_idx = ins.get("page_idx", -1)
new_text = "" if not is_supported_chunk_record(ins):
should_record_highlight = False # 标记是否需要记录 highlight即使 new_text 为空) continue
if ins["type"] == "text": new_text = record_to_chunk_text(ins)
text = ins.get("text", "") if new_text.strip():
level = ins.get("text_level") new_text += "\n"
if level is not None and isinstance(level, int) and level > 0:
text = "#" * min(level, 6) + " " + text # 安全限制标题级别
new_text = text + "\n"
should_record_highlight = True should_record_highlight = True
elif ins["type"] == "table":
table_caption = ins.get("table_caption")
table_body = ins.get("table_body", "")
# 清理冗余属性
for pattern in [' colspan="1"', ' rowspan="1"', " colspan='1'", " rowspan='1'", " colspan=1", " rowspan=1"]:
table_body = table_body.replace(pattern, "")
caption_str = ""
if table_caption:
if isinstance(table_caption, list):
caption_str = "\n".join(str(c) for c in table_caption)
else:
caption_str = str(table_caption)
# 判断是否有实质内容
has_caption = bool(caption_str.strip())
has_body = bool(table_body.strip())
if has_caption or has_body:
new_text = caption_str + ("\n" if has_caption else "") + table_body + "\n"
else:
new_text = "" # 空内容,不拼接
# 即使内容为空,只要是 table 类型,就记录 highlight满足你的需求
should_record_highlight = True
elif ins["type"] == "image":
image_caption = ins.get("image_caption", "")
img_path = ins.get("img_path", "")
if isinstance(image_caption, list):
image_caption = "\n".join(str(item) for item in image_caption)
elif not isinstance(image_caption, str):
image_caption = str(image_caption)
if isinstance(img_path, list):
img_path = "\n".join(str(p) for p in img_path)
elif not isinstance(img_path, str):
img_path = str(img_path)
if image_caption.strip() or img_path.strip():
new_text = image_caption + "\n" + img_path + "\n"
else:
new_text = ""
should_record_highlight = True
elif ins["type"] == "discarded":
text = ins.get("text", "")
if text.strip():
new_text = text + "\n"
else:
new_text = ""
should_record_highlight = True
else:
continue # 不支持的 type跳过且不记录 highlight
# --- 关键:即使 new_text 为空,只要 should_record_highlight 为 True就要处理 chunk 切分和 highlight 记录 --- # --- 关键:即使 new_text 为空,只要 should_record_highlight 为 True就要处理 chunk 切分和 highlight 记录 ---
# 构造 highlight 项(即使 new_text 为空) # 构造 highlight 项(即使 new_text 为空)
@ -707,7 +821,7 @@ def split_content_bbox(data, max_length=8000):
def group_records(records): def _group_records_without_parent_context(records):
result = [] result = []
current_group = [] current_group = []
@ -728,10 +842,50 @@ def group_records(records):
return result return result
def _get_group_level(group):
if not group:
return None
level = group[0].get("text_level")
if isinstance(level, int) and level > 0:
return level
return None
def group_records(records):
"""
Split records by text_level, then prepend ancestor section content.
For example, a text_level=2 section will include the nearest text_level=1
group content before its own content. A text_level=3 section will include
the nearest level 1 and level 2 group content before its own content.
"""
base_groups = _group_records_without_parent_context(records)
result = []
section_stack = [] # [(level, merged_group)]
for group in base_groups:
level = _get_group_level(group)
if level is None:
result.append(group)
continue
while section_stack and section_stack[-1][0] >= level:
section_stack.pop()
parent_group = section_stack[-1][1] if section_stack else []
merged_group = parent_group + group
result.append(merged_group)
section_stack.append((level, merged_group))
return result
def get_chunk_bbox_otherfile(data, max_length=8000): def get_chunk_bbox_otherfile(data, max_length=8000):
""" """
处理没有 bbox 信息的文件数据主要面向纯表格类内容 处理没有 bbox 信息的文件数据
每个 table 作为一个独立切片不与其他记录拼接 条支持的记录作为一个独立切片不与其他记录拼接
- 单个 table 内容超过 max_length 调用 chunk_reset </tr> 处安全切分 - 单个 table 内容超过 max_length 调用 chunk_reset </tr> 处安全切分
切分后的多个 chunk 共享同一组 positions 切分后的多个 chunk 共享同一组 positions
@ -743,38 +897,14 @@ def get_chunk_bbox_otherfile(data, max_length=8000):
slices = [] slices = []
for ins in data: for ins in data:
if ins.get("type") != "table": if not is_supported_chunk_record(ins):
continue continue
page_idx = ins.get("page_idx", -1) page_idx = ins.get("page_idx", -1)
bbox = ins.get("bbox", []) bbox = ins.get("bbox", [])
content = record_to_chunk_text(ins)
table_caption = ins.get("table_caption") if not content.strip():
table_body = ins.get("table_body", "") or "" continue
# 清理冗余属性
for pattern in [
' colspan="1"', ' rowspan="1"',
" colspan='1'", " rowspan='1'",
" colspan=1", " rowspan=1",
]:
table_body = table_body.replace(pattern, "")
# 拼接 caption
caption_str = ""
if table_caption:
if isinstance(table_caption, list):
caption_str = "\n".join(str(c) for c in table_caption)
else:
caption_str = str(table_caption)
has_caption = bool(caption_str.strip())
has_body = bool(table_body.strip())
if not (has_caption or has_body):
continue # 空表格跳过
content = caption_str + ("\n" if has_caption else "") + table_body
positions = [{"page_idx": page_idx, "bbox": bbox}] positions = [{"page_idx": page_idx, "bbox": bbox}]
# 单表超长 → 在 </tr> 处切分,多个 chunk 共享 positions # 单表超长 → 在 </tr> 处切分,多个 chunk 共享 positions
@ -875,38 +1005,11 @@ def get_chunk_bbox(data):
bbox = ins["bbox"] bbox = ins["bbox"]
positions = [{"page_idx": page_idx, "bbox": bbox}] positions = [{"page_idx": page_idx, "bbox": bbox}]
if ins["type"] == "text": if not is_supported_chunk_record(ins):
if ins.get("text_level"): continue
text = ins.get("text", "")
level = ins.get("text_level")
content = "#" * level + text
else:
content = ins.get("text", "")
elif ins["type"] == "table":
content = (ins.get("table_caption") or "") + "\n" + (ins.get("table_body") or "")
elif ins["type"] == "image":
image_caption = ins.get("image_caption", [])
img_path = ins.get("img_path", [])
# 如果 image_caption 是列表,则将其转换为字符串 content = record_to_chunk_text(ins)
if isinstance(image_caption, list): if not content.strip():
image_caption = "\n".join(str(item) for item in image_caption)
elif not isinstance(image_caption, str):
image_caption = str(image_caption)
# 如果 img_path 是列表,则将其转换为字符串
if isinstance(img_path, list):
img_path = "\n".join(str(p) for p in img_path)
elif not isinstance(img_path, str):
img_path = str(img_path)
if image_caption.strip() or img_path.strip():
content = (image_caption or "") + "\n" + (img_path or "")
else:
continue # 跳过空图像项
elif ins["type"] == "discarded":
content = ins.get("text", "")
else:
continue continue
slices.append({"content": content, "positions": positions}) slices.append({"content": content, "positions": positions})
@ -977,7 +1080,7 @@ def split_text_preserve_sentences(text, max_length=7500):
return chunks return chunks
def data_replace(data, prefix): async def data_replace(data, prefix):
""" """
遍历 data 列表若元素包含 'img_path' 字段则移除前缀 'images/' 并拼接新前缀 遍历 data 列表若元素包含 'img_path' 字段则移除前缀 'images/' 并拼接新前缀
使用 pathlib 安全处理路径 使用 pathlib 安全处理路径
@ -985,32 +1088,27 @@ def data_replace(data, prefix):
for ins in data: for ins in data:
if isinstance(ins, dict) and "img_path" in ins: if isinstance(ins, dict) and "img_path" in ins:
img_path = ins["img_path"] img_path = ins["img_path"]
local_image_paths = []
if isinstance(img_path, list): if isinstance(img_path, list):
ins["img_path"] = "\n".join( markdown_images = []
data_replace([{"img_path": str(item)}], prefix)[0]["img_path"] for item in img_path:
for item in img_path if not str(item).strip():
if str(item).strip()
)
continue continue
markdown_image, local_image_path = build_image_markdown(str(item), prefix)
img_path = str(img_path) markdown_images.append(markdown_image)
if img_path.strip().startswith("!["): if local_image_path:
local_image_path = resolve_local_image_path(img_path) local_image_paths.append(local_image_path)
ins["img_path"] = build_image_markdown_with_ocr(img_path, local_image_path) ins["img_path"] = "\n".join(markdown_images)
continue
p = PurePosixPath(img_path)
# 如果以 images/ 开头,去掉第一级目录
if p.parts and p.parts[0] == "images":
relative_path = str(PurePosixPath(*p.parts[1:]))
markdown_image = f"![图片]({prefix}{relative_path})"
local_image_path = resolve_local_image_path(relative_path)
ins["img_path"] = build_image_markdown_with_ocr(markdown_image, local_image_path)
else: else:
# 否则保留原路径(或按需处理) markdown_image, local_image_path = build_image_markdown(str(img_path), prefix)
markdown_image = f"![图片]({prefix}{img_path})" ins["img_path"] = markdown_image
local_image_path = resolve_local_image_path(img_path) if local_image_path:
ins["img_path"] = build_image_markdown_with_ocr(markdown_image, local_image_path) local_image_paths.append(local_image_path)
ocr_texts = await asyncio.gather(
*(get_image_ocr_text(path) for path in local_image_paths)
)
append_caption_ocr(ins, ocr_texts)
return data return data
def filter_content_bbox(data, max_length=3000): def filter_content_bbox(data, max_length=3000):
@ -1034,45 +1132,15 @@ def filter_content_bbox(data, max_length=3000):
if "type" not in ins: if "type" not in ins:
continue continue
added_content = ""
bbox = ins.get("bbox", []) bbox = ins.get("bbox", [])
page_idx = ins.get("page_idx", -1) page_idx = ins.get("page_idx", -1)
if ins["type"] == "text": if not is_supported_chunk_record(ins):
text_content = ins.get("text", "") continue
added_content = text_content + "\n"
elif ins["type"] == "table": added_content = record_to_chunk_text(ins)
table_caption = ins.get("table_caption") if added_content.strip():
table_body = ins.get("table_body") added_content += "\n"
caption_str = ""
if table_caption:
if isinstance(table_caption, list):
caption_str = "\n".join(str(c) for c in table_caption)
else:
caption_str = str(table_caption)
added_content = caption_str + "\n" + (table_body or "") + "\n"
elif ins["type"] == "image":
image_caption = ins.get("image_caption", "")
img_path = ins.get("img_path", "")
if isinstance(image_caption, list):
image_caption = "\n".join(str(item) for item in image_caption)
elif not isinstance(image_caption, str):
image_caption = str(image_caption)
if isinstance(img_path, list):
img_path = "\n".join(str(p) for p in img_path)
elif not isinstance(img_path, str):
img_path = str(img_path)
if image_caption or img_path:
added_content = image_caption + "\n" + img_path + "\n"
else:
continue # 跳过空图像项
else: else:
continue continue
@ -1319,7 +1387,8 @@ async def process_pdf_file(pdf_path: Path, image_prefix: str,filename:str) -> Op
if doc_result is not None: if doc_result is not None:
content_list, images = doc_result content_list, images = doc_result
if content_list and images: if content_list and images:
replaced_content_list = data_replace(content_list, image_prefix) replaced_content_textlevel = reset_textlevel(content_list)
replaced_content_list = await data_replace(replaced_content_textlevel, image_prefix)
slices = get_chunk_bbox(replaced_content_list) slices = get_chunk_bbox(replaced_content_list)
slices_check = chunk_check(slices, 8000) slices_check = chunk_check(slices, 8000)
slices_check = merge_short_slices(slices_check, min_length=30,filename=filename) slices_check = merge_short_slices(slices_check, min_length=30,filename=filename)
@ -1342,7 +1411,8 @@ async def process_pdf_file(pdf_path: Path, image_prefix: str,filename:str) -> Op
data = result.get("data", {}) data = result.get("data", {})
content_list = data.get("content_list", []) content_list = data.get("content_list", [])
images = data.get("images", {}) images = data.get("images", {})
replaced_content_list = data_replace(content_list, image_prefix) replaced_content_textlevel = reset_textlevel(content_list)
replaced_content_list = await data_replace(replaced_content_textlevel, image_prefix)
slices = get_chunk_bbox(replaced_content_list) slices = get_chunk_bbox(replaced_content_list)
slices_check = chunk_check(slices, 8000) slices_check = chunk_check(slices, 8000)
slices_check = merge_short_slices(slices_check, min_length=30,filename=filename) slices_check = merge_short_slices(slices_check, min_length=30,filename=filename)
@ -1367,7 +1437,8 @@ async def process_other_file(pdf_path: Path, image_prefix: str, filename: str) -
if doc_result is not None: if doc_result is not None:
content_list, images = doc_result content_list, images = doc_result
if content_list and images: if content_list and images:
replaced_content_list = data_replace(content_list, image_prefix) replaced_content_textlevel = reset_textlevel(content_list)
replaced_content_list = await data_replace(replaced_content_textlevel, image_prefix)
slices = get_chunk_bbox(replaced_content_list) slices = get_chunk_bbox(replaced_content_list)
slices_check = chunk_check(slices, 8000) slices_check = chunk_check(slices, 8000)
return {"slices": slices_check, "images": images} return {"slices": slices_check, "images": images}
@ -1397,7 +1468,8 @@ async def process_other_file(pdf_path: Path, image_prefix: str, filename: str) -
images = data.get("images", {}) images = data.get("images", {})
# 后续的数据替换与切片处理逻辑 # 后续的数据替换与切片处理逻辑
replaced_content_list = data_replace(content_list, image_prefix) replaced_content_textlevel = reset_textlevel(content_list)
replaced_content_list = await data_replace(replaced_content_textlevel, image_prefix)
slices = get_chunk_bbox(replaced_content_list) slices = get_chunk_bbox(replaced_content_list)
slices_check = chunk_check(slices, 8000) slices_check = chunk_check(slices, 8000)
@ -1415,7 +1487,7 @@ async def process_other_file(pdf_path: Path, image_prefix: str, filename: str) -
if __name__ == "__main__": if __name__ == "__main__":
filepath = r"E:\ZKYNLP\Hjunproject\project0506\kgrag\122-06A0014-B01001_发动机-维修手册_content_list.json" filepath = r"E:\ZKYNLP\Hjunproject\project0506\kgrag\舰船抗沉损管训练仿真系统研究_content_list.json"
try: try:
with open(filepath, 'r', encoding='utf-8') as file: with open(filepath, 'r', encoding='utf-8') as file:
data = json.load(file) data = json.load(file)
@ -1425,12 +1497,12 @@ if __name__ == "__main__":
# print(len(data)) # print(len(data))
start_chars = ["1", "2", "3", "4", "5", "6", "7", "8", "9", "0", "", "", start_chars = ["1", "2", "3", "4", "5", "6", "7", "8", "9", "0", "", "",
"", "", "", "", "", "", "", "", ""] "", "", "", "", "", "", "", "", ""]
replaced_content_textlevel = reset_textlevel(data)
data = data_replace(data=data,prefix="/api/v1/knowledge/files/images/") data = asyncio.run(data_replace(data=replaced_content_textlevel,prefix="/api/v1/knowledge/files/images/"))
slices = get_chunk_bbox(data=data) slices = get_chunk_bbox(data=data)
slices_1 =chunk_check(slices, 8000) slices_1 =chunk_check(slices, 8000)
slices_2 = merge_short_slices(slices_1, min_length=30) # ← 新增这一行 slices_2 = merge_short_slices(slices_1, min_length=30) # ← 新增这一行
for ins in slices_2: for ins in slices_2[:30]:
print(ins) print(ins)

184
resetlevel.py Normal file
View File

@ -0,0 +1,184 @@
import copy
import json
import re
import sys
from typing import Any
CN_NUM = "\u4e00\u4e8c\u4e09\u56db\u4e94\u516d\u4e03\u516b\u4e5d\u5341\u767e\u5343\u4e07\u3007\u96f6\u4e24"
CHAPTER_PATTERN = re.compile(rf"^\s*\u7b2c\s*([{CN_NUM}\d]+)\s*\u7ae0")
CN_LEVEL1_PATTERN = re.compile(rf"^\s*([{CN_NUM}]+)(?=\s|[\u3001\u3002\uff0e.])")
CASE_LIKE_PATTERN = re.compile(r"^\s*\d+(?:\.|\uff0e)(?!\d|\s|$)")
ARABIC_PATTERN = re.compile(
r"^\s*(\d+(?:\.\d+)*)(?=\s|[\u4e00-\u9fff\u3001\u3002\uff0e]|$|\.(?:\s|$))"
)
def has_valid_text_level(item: dict) -> bool:
return "text_level" in item and item.get("text_level") not in (None, "")
def parse_heading(text: str):
text = text.strip()
if not text:
return None, None
if CASE_LIKE_PATTERN.match(text):
return "case_like", None
if CHAPTER_PATTERN.match(text):
return "level1", None
if CN_LEVEL1_PATTERN.match(text):
return "level1", None
match = ARABIC_PATTERN.match(text)
if match:
return "arabic", match.group(1)
return None, None
def arabic_depth(number: str) -> int:
return number.count(".") + 1
def parent_number(number: str) -> str | None:
parts = number.split(".")
if len(parts) <= 1:
return None
return ".".join(parts[:-1])
def collect_headings(data: list) -> list:
headings = []
for item in data:
if not isinstance(item, dict):
continue
if item.get("type") != "text":
continue
if not has_valid_text_level(item):
continue
kind, number = parse_heading(item.get("text", ""))
if kind in {"level1", "arabic", "case_like"}:
headings.append({"kind": kind, "number": number})
return headings
def is_valid_arabic_hierarchy(numbers: list[str], allow_missing_depth1_parent: bool) -> bool:
seen = set()
has_depth2 = False
has_depth3 = False
for number in numbers:
depth = arabic_depth(number)
if depth >= 2:
has_depth2 = True
if depth >= 3:
has_depth3 = True
parent = parent_number(number)
if parent is not None:
parent_depth = arabic_depth(parent)
parent_is_missing_depth1 = allow_missing_depth1_parent and parent_depth == 1
if parent not in seen and not parent_is_missing_depth1:
return False
seen.add(number)
return has_depth2 and has_depth3
def detect_valid_mode(headings: list):
if not headings:
return None
if any(heading["kind"] == "case_like" for heading in headings):
return None
first = headings[0]
if first["kind"] == "arabic":
numbers = [heading["number"] for heading in headings if heading["kind"] == "arabic"]
if numbers and arabic_depth(numbers[0]) == 1:
if is_valid_arabic_hierarchy(numbers, allow_missing_depth1_parent=False):
return "A"
return None
if first["kind"] == "level1":
numbers = [heading["number"] for heading in headings[1:] if heading["kind"] == "arabic"]
if not numbers:
return None
first_depth = arabic_depth(numbers[0])
if first_depth == 2 and is_valid_arabic_hierarchy(numbers, allow_missing_depth1_parent=True):
return "B"
if first_depth == 1 and is_valid_arabic_hierarchy(numbers, allow_missing_depth1_parent=False):
return "C"
return None
JsonContent = Any
def reset_textlevel(json_content: JsonContent) -> JsonContent:
"""
输入 JSON 内容返回更新 text_level 后的新 JSON 内容
json_content 可以是已经 json.load/json.loads 后的 Python 对象
也可以是 JSON 字符串函数不会修改原始入参
"""
if isinstance(json_content, str):
data = json.loads(json_content)
else:
data = copy.deepcopy(json_content)
if not isinstance(data, list):
return data
headings = collect_headings(data)
mode = detect_valid_mode(headings)
if mode is None:
return data
for item in data:
if not isinstance(item, dict):
continue
if item.get("type") != "text":
continue
if not has_valid_text_level(item):
continue
kind, number = parse_heading(item.get("text", ""))
if kind == "level1":
item["text_level"] = 1
continue
if kind != "arabic":
continue
depth = arabic_depth(number)
item["text_level"] = depth if mode in {"A", "B"} else depth + 1
return data
if __name__ == "__main__":
json_file_path = r'E:\ZKYNLP\Hjunproject\project0506\kgrag\163-06A0014-B01001_发动机-维修手册_content_list.json'
with open(json_file_path, "r", encoding="utf-8") as f:
json_content = json.load(f)
result= reset_textlevel(json_content)
for ins in result[:30]:
print(ins)
# print(json.dumps(result, ensure_ascii=False, indent=2))