更新公式输出

This commit is contained in:
ChenRui 2026-07-21 18:23:41 +08:00
parent 0ddd86860e
commit d1ba8905cc

View File

@ -51,6 +51,752 @@ class QAState(TypedDict):
suggestedReplies: List[Dict[str, Any]] # 建议回复问题列表
error_message: str # 错误信息
# =============== 辅助函数 ===============
import re
def normalize_latex_spacing(text: str) -> str:
"""
规范 Markdown 中的 LaTeX 公式
处理
- $...$行内公式清理内侧空格并补充正文间空格
- $$...$$段落公式完整保留开始和结束定界符
- \\(...\\)转换为 $...$
- \\[...\\]转换为 $$...$$
不修改
- Markdown 围栏代码块
- Markdown 行内代码
- HTML code/pre/script/style
- HTML 注释
- 转义美元符号 \\$
"""
if not isinstance(text, str):
raise TypeError(
f"text 必须是 str实际类型为 {type(text).__name__}"
)
if not text:
return text
text = text.replace("\r\n", "\n").replace("\r", "\n")
protected_parts: list[str] = []
def protect(content: str) -> str:
token = (
f"\uE000LATEX_PROTECTED_"
f"{len(protected_parts)}"
f"\uE001"
)
protected_parts.append(content)
return token
def restore(content: str) -> str:
for index, original in enumerate(protected_parts):
token = f"\uE000LATEX_PROTECTED_{index}\uE001"
content = content.replace(token, original)
return content
def is_escaped(content: str, position: int) -> bool:
slash_count = 0
position -= 1
while position >= 0 and content[position] == "\\":
slash_count += 1
position -= 1
return slash_count % 2 == 1
# ================================================================
# 保护 Markdown 围栏代码块
# ================================================================
lines = text.splitlines(keepends=True)
output_lines: list[str] = []
line_index = 0
while line_index < len(lines):
line = lines[line_index]
fence_match = re.match(
r"^[ \t]{0,3}(`{3,}|~{3,})",
line,
)
if not fence_match:
output_lines.append(line)
line_index += 1
continue
opening_fence = fence_match.group(1)
fence_character = opening_fence[0]
fence_length = len(opening_fence)
closing_pattern = re.compile(
rf"^[ \t]{{0,3}}"
rf"{re.escape(fence_character)}"
rf"{{{fence_length},}}"
rf"[ \t]*(?:\n|$)"
)
block_lines = [line]
line_index += 1
while line_index < len(lines):
current_line = lines[line_index]
block_lines.append(current_line)
line_index += 1
if closing_pattern.match(current_line):
break
output_lines.append(
protect("".join(block_lines))
)
text = "".join(output_lines)
# ================================================================
# 保护 HTML 代码区域和注释
# ================================================================
text = re.sub(
r"(?is)"
r"<(?P<tag>pre|code|script|style)\b[^>]*>"
r".*?"
r"</(?P=tag)\s*>",
lambda match: protect(match.group(0)),
text,
)
text = re.sub(
r"(?s)<!--.*?-->",
lambda match: protect(match.group(0)),
text,
)
# ================================================================
# 保护 Markdown 链接地址
# ================================================================
text = re.sub(
r"(!?\[[^\]\n]*\]\()([^)\n]*)(\))",
lambda match: (
match.group(1)
+ protect(match.group(2))
+ match.group(3)
),
text,
)
# ================================================================
# 保护 Markdown 行内代码
# ================================================================
inline_code_output: list[str] = []
position = 0
while position < len(text):
if text[position] != "`":
inline_code_output.append(text[position])
position += 1
continue
delimiter_end = position
while (
delimiter_end < len(text)
and text[delimiter_end] == "`"
):
delimiter_end += 1
delimiter = text[position:delimiter_end]
search_position = delimiter_end
closing_position = -1
while search_position < len(text):
candidate = text.find(
delimiter,
search_position,
)
if candidate < 0:
break
previous_character = (
text[candidate - 1]
if candidate > 0
else ""
)
next_position = candidate + len(delimiter)
next_character = (
text[next_position]
if next_position < len(text)
else ""
)
if (
previous_character != "`"
and next_character != "`"
):
closing_position = candidate
break
search_position = candidate + 1
if closing_position < 0:
inline_code_output.append(delimiter)
position = delimiter_end
continue
end_position = closing_position + len(delimiter)
inline_code_output.append(
protect(text[position:end_position])
)
position = end_position
text = "".join(inline_code_output)
# ================================================================
# 转换其他 LaTeX 定界符
# ================================================================
# \( x + y \) -> $x + y$
text = re.sub(
r"(?<!\\)\\\((.*?)(?<!\\)\\\)",
lambda match: (
"$" + match.group(1).strip() + "$"
),
text,
)
# \[ x + y \] -> $$x + y$$
text = re.sub(
r"(?s)(?<!\\)\\\[(.*?)(?<!\\)\\\]",
lambda match: (
"$$" + match.group(1).strip() + "$$"
),
text,
)
# ================================================================
# 扫描公式
# ================================================================
opening_punctuation = set(
"([{(【《〈“‘"
)
closing_punctuation = set(
")]})】》〉”’、,。;:!?,.!?;:"
)
result: list[str] = []
position = 0
while position < len(text):
# ------------------------------------------------------------
# 段落公式 $$...$$
# ------------------------------------------------------------
if (
text.startswith("$$", position)
and not is_escaped(text, position)
):
search_position = position + 2
closing_position = -1
while search_position < len(text):
candidate = text.find(
"$$",
search_position,
)
if candidate < 0:
break
if not is_escaped(text, candidate):
closing_position = candidate
break
search_position = candidate + 2
# 未闭合时,原样保留剩余内容。
if closing_position < 0:
result.append(text[position:])
break
complete_formula = text[
position:closing_position + 2
]
formula_body = text[
position + 2:closing_position
]
# 多行段落公式完整原样保留。
if "\n" in complete_formula:
result.append(complete_formula)
else:
# 单行段落公式只清理内侧空格,
# 不删除或重建结束的 $$。
result.append(
"$$" + formula_body.strip() + "$$"
)
position = closing_position + 2
continue
# ------------------------------------------------------------
# 行内公式 $...$
# ------------------------------------------------------------
if (
text[position] == "$"
and not is_escaped(text, position)
and not text.startswith("$$", position)
and not (
position > 0
and text[position - 1] == "$"
)
):
opening_position = position
search_position = position + 1
closing_position = -1
while search_position < len(text):
candidate = text.find(
"$",
search_position,
)
if candidate < 0:
break
# 行内公式不允许跨行。
if "\n" in text[
opening_position + 1:candidate
]:
break
if is_escaped(text, candidate):
search_position = candidate + 1
continue
if text.startswith("$$", candidate):
search_position = candidate + 2
continue
next_character = (
text[candidate + 1]
if candidate + 1 < len(text)
else ""
)
# 避免把 $100 和 $200 配对成公式。
if next_character.isdigit():
search_position = candidate + 1
continue
closing_position = candidate
break
if closing_position < 0:
result.append("$")
position += 1
continue
formula_body = text[
opening_position + 1:closing_position
].strip()
if not formula_body:
result.append(
text[
opening_position:
closing_position + 1
]
)
position = closing_position + 1
continue
formula_body = re.sub(
r"[ \t]+",
" ",
formula_body,
)
previous_character = ""
if result and result[-1]:
previous_character = result[-1][-1]
next_character = (
text[closing_position + 1]
if closing_position + 1 < len(text)
else ""
)
if (
previous_character
and not previous_character.isspace()
and previous_character
not in opening_punctuation
):
result.append(" ")
result.append(
"$" + formula_body + "$"
)
if (
next_character
and not next_character.isspace()
and next_character
not in closing_punctuation
):
result.append(" ")
position = closing_position + 1
continue
result.append(text[position])
position += 1
return restore("".join(result))
def remove_duplicate_content(text: str) -> str:
"""
删除普通 Markdown 正文中的完全重复行
不对以下内容进行去重
- $$ ... $$ 段落公式
- \\[ ... \\] 段落公式
- equationaligngather LaTeX 环境
- Markdown 围栏代码块
- Markdown 缩进代码块
- 标题列表引用表格等 Markdown 结构行
这样不会删除段落公式末尾的第二个 $$
"""
if not text or len(text) < 10:
return text
text = text.replace("\r\n", "\n").replace("\r", "\n")
lines = text.split("\n")
result: list[str] = []
seen_segments: set[str] = set()
in_code_block = False
code_fence_character = ""
code_fence_length = 0
in_dollar_math = False
in_bracket_math = False
latex_environment_stack: list[str] = []
display_environments = {
"equation",
"equation*",
"align",
"align*",
"alignat",
"alignat*",
"gather",
"gather*",
"multline",
"multline*",
"flalign",
"flalign*",
"eqnarray",
"eqnarray*",
"displaymath",
"math",
}
def is_escaped(value: str, position: int) -> bool:
slash_count = 0
position -= 1
while position >= 0 and value[position] == "\\":
slash_count += 1
position -= 1
return slash_count % 2 == 1
def count_unescaped_double_dollars(value: str) -> int:
count = 0
position = 0
while position < len(value) - 1:
if (
value[position:position + 2] == "$$"
and not is_escaped(value, position)
):
count += 1
position += 2
else:
position += 1
return count
def is_markdown_structure_line(value: str) -> bool:
stripped = value.strip()
if not stripped:
return True
patterns = (
# Markdown 标题
r"^#{1,6}\s+",
# 无序列表
r"^[-+*]\s+",
# 有序列表
r"^\d+[.)]\s+",
# 引用
r"^>\s*",
# 任务列表
r"^[-+*]\s+\[[ xX]\]\s+",
# 水平分隔线
r"^(?:-{3,}|\*{3,}|_{3,})$",
# Markdown 表格分隔行
r"^\|?\s*:?-{3,}:?"
r"(?:\s*\|\s*:?-{3,}:?)+\s*\|?$",
# 普通表格行
r"^\|.*\|$",
# HTML 标签
r"^</?[A-Za-z][^>]*>$",
# LaTeX 环境边界
r"^\\(?:begin|end)\{[^}]+\}",
# 公式定界符
r"^(?:\$\$|\\\[|\\\])$",
# Markdown 链接引用定义
r"^\[[^\]]+\]:\s*\S+",
)
return any(
re.match(pattern, stripped)
for pattern in patterns
)
for line in lines:
line = line.rstrip()
stripped = line.strip()
# ============================================================
# Markdown 围栏代码块
# ============================================================
fence_match = re.match(
r"^[ \t]{0,3}(`{3,}|~{3,})",
line,
)
if not in_code_block and fence_match:
fence = fence_match.group(1)
in_code_block = True
code_fence_character = fence[0]
code_fence_length = len(fence)
result.append(line)
continue
if in_code_block:
result.append(line)
closing_pattern = re.compile(
rf"^[ \t]{{0,3}}"
rf"{re.escape(code_fence_character)}"
rf"{{{code_fence_length},}}"
rf"[ \t]*$"
)
if closing_pattern.match(line):
in_code_block = False
code_fence_character = ""
code_fence_length = 0
continue
# Markdown 缩进代码块。
if re.match(r"^(?:\t| {4})", line):
result.append(line)
continue
# ============================================================
# \[ ... \] 段落公式
# ============================================================
if in_bracket_math:
result.append(line)
if re.search(r"(?<!\\)\\\]", line):
in_bracket_math = False
continue
bracket_open_count = len(
re.findall(r"(?<!\\)\\\[", line)
)
bracket_close_count = len(
re.findall(r"(?<!\\)\\\]", line)
)
if bracket_open_count:
result.append(line)
if bracket_open_count > bracket_close_count:
in_bracket_math = True
continue
# ============================================================
# LaTeX display 环境
# ============================================================
begin_matches = re.findall(
r"\\begin\{([^}]+)\}",
line,
)
end_matches = re.findall(
r"\\end\{([^}]+)\}",
line,
)
relevant_begins = [
environment
for environment in begin_matches
if environment in display_environments
]
relevant_ends = [
environment
for environment in end_matches
if environment in display_environments
]
if latex_environment_stack:
result.append(line)
for environment in relevant_begins:
latex_environment_stack.append(environment)
for environment in relevant_ends:
if (
latex_environment_stack
and latex_environment_stack[-1]
== environment
):
latex_environment_stack.pop()
elif environment in latex_environment_stack:
latex_environment_stack.remove(environment)
continue
if relevant_begins:
result.append(line)
for environment in relevant_begins:
latex_environment_stack.append(environment)
for environment in relevant_ends:
if (
latex_environment_stack
and latex_environment_stack[-1]
== environment
):
latex_environment_stack.pop()
elif environment in latex_environment_stack:
latex_environment_stack.remove(environment)
continue
# ============================================================
# $$ ... $$ 段落公式
# ============================================================
double_dollar_count = (
count_unescaped_double_dollars(line)
)
if in_dollar_math:
# 公式内部和结束的 $$ 全部原样保留。
result.append(line)
if double_dollar_count % 2 == 1:
in_dollar_math = False
continue
if double_dollar_count > 0:
# 包含 $$ 的行不参与去重。
result.append(line)
# 奇数个 $$ 表示开启了跨行公式。
if double_dollar_count % 2 == 1:
in_dollar_math = True
continue
# ============================================================
# 空行和 Markdown 结构行
# ============================================================
if not stripped:
result.append(line)
continue
if is_markdown_structure_line(line):
result.append(line)
continue
# ============================================================
# 仅对普通正文行去重
# ============================================================
normalized_line = re.sub(
r"\s+",
" ",
stripped.lower(),
)
if normalized_line in seen_segments:
continue
seen_segments.add(normalized_line)
result.append(line)
return "\n".join(result)
# 推荐调用顺序:
#
# text = remove_duplicate_content(text)
# text = normalize_latex_spacing(text)
# =============== 节点函数 ===============
@ -418,18 +1164,20 @@ async def generate_answer(state: QAState) -> Dict[str, Any]:
answer = answer.strip() if answer else raw_content
# LaTeX 公式空格规范化
def normalize_latex_spacing(text: str) -> str:
placeholder = "\x00DOUBLE_DOLLAR\x00"
text = text.replace("$$", placeholder)
text = re.sub(r'(?<!\s)\$', r' $', text)
text = re.sub(r'\$(?!\s)', r'$ ', text)
text = text.replace(placeholder, "$$")
text = re.sub(r'(?<!\s)\$\$', r' $$', text)
text = re.sub(r'\$\$(?!\s)', r'$$ ', text)
return text
# def normalize_latex_spacing(text: str) -> str:
# placeholder = "\x00DOUBLE_DOLLAR\x00"
# text = text.replace("$$", placeholder)
# text = re.sub(r'(?<!\s)\$', r' $', text)
# text = re.sub(r'\$(?!\s)', r'$ ', text)
# text = text.replace(placeholder, "$$")
# text = re.sub(r'(?<!\s)\$\$', r' $$', text)
# text = re.sub(r'\$\$(?!\s)', r'$$ ', text)
# return text
answer = remove_duplicate_content(answer)
answer = normalize_latex_spacing(answer)
# answer = normalize_latex_spacing(answer)
# 使用增强的图片路径匹配与规范化(仅在有检索结果时)
if need_retrieval and original_image_paths:
answer = normalize_markdown_images(answer, original_paths=original_image_paths)