diff --git a/workflows/workflow_baike.py b/workflows/workflow_baike.py index 891318f..080caba 100644 --- a/workflows/workflow_baike.py +++ b/workflows/workflow_baike.py @@ -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"<(?Ppre|code|script|style)\b[^>]*>" + r".*?" + r"", + 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"(? $$x + y$$ + text = re.sub( + r"(?s)(? 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 正文中的完全重复行。 + + 不对以下内容进行去重: + - $$ ... $$ 段落公式 + - \\[ ... \\] 段落公式 + - equation、align、gather 等 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"^]*>$", + + # 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"(? 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'(? str: +# placeholder = "\x00DOUBLE_DOLLAR\x00" +# text = text.replace("$$", placeholder) +# text = re.sub(r'(?