1223 lines
36 KiB
Python
1223 lines
36 KiB
Python
"""
|
||
工作流公共工具函数
|
||
提取自 workflow_fault_diagnosis.py 和 workflow_operate.py 的共享代码
|
||
"""
|
||
|
||
from typing import TypedDict, Optional, Dict, Any, List
|
||
from difflib import SequenceMatcher
|
||
from modelsAPI.model_api import OpenaiAPI
|
||
from utils.function_tracker import get_all_callbacks, get_current_stream_handler
|
||
from utils.image_utils import extract_image_paths, normalize_markdown_images
|
||
import json
|
||
import re
|
||
import random
|
||
import inspect
|
||
import unicodedata
|
||
|
||
|
||
RAG_DEDUP_SIMILARITY_THRESHOLD = 0.86
|
||
RAG_DEDUP_CONTAINMENT_THRESHOLD = 0.92
|
||
RAG_DEDUP_MIN_TEXT_LENGTH = 20
|
||
|
||
|
||
def safe_json_extract(text: str) -> Optional[Any]:
|
||
if not text:
|
||
return None
|
||
cleaned = text.strip()
|
||
cleaned = re.sub(r"```json\s*", "", cleaned, flags=re.IGNORECASE)
|
||
cleaned = re.sub(r"```\s*", "", cleaned)
|
||
m = re.search(r"\[[\s\S]*\]", cleaned)
|
||
if m:
|
||
try:
|
||
return json.loads(m.group(0))
|
||
except Exception:
|
||
pass
|
||
m = re.search(r"\{[\s\S]*\}", cleaned)
|
||
if m:
|
||
try:
|
||
return json.loads(m.group(0))
|
||
except Exception:
|
||
pass
|
||
try:
|
||
return json.loads(cleaned)
|
||
except Exception:
|
||
return None
|
||
|
||
|
||
from workflows.history_manager import parse_history, build_history_str, filter_image_urls, init_history # noqa: F401
|
||
|
||
|
||
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 正文中的完全重复行。
|
||
|
||
不对以下内容进行去重:
|
||
- $$ ... $$ 段落公式
|
||
- \\[ ... \\] 段落公式
|
||
- 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"^</?[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)
|
||
|
||
#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 remove_duplicate_content(text: str) -> str:
|
||
# if not text or len(text) < 10:
|
||
# return text
|
||
#
|
||
# lines = text.split('\n')
|
||
# result = []
|
||
# seen_segments = set()
|
||
#
|
||
# for line in lines:
|
||
# line = line.rstrip()
|
||
# if not line:
|
||
# result.append(line)
|
||
# continue
|
||
#
|
||
# normalized_line = line.strip().lower()
|
||
# if normalized_line in seen_segments:
|
||
# continue
|
||
#
|
||
# seen_segments.add(normalized_line)
|
||
# result.append(line)
|
||
#
|
||
# cleaned_text = '\n'.join(result)
|
||
# return cleaned_text
|
||
|
||
|
||
def _extract_rag_text(item: Any) -> str:
|
||
if isinstance(item, dict):
|
||
return str(item.get("text") or item.get("content") or "")
|
||
if isinstance(item, str):
|
||
return item
|
||
return ""
|
||
|
||
|
||
def _extract_rag_score(item: Any) -> float:
|
||
if not isinstance(item, dict):
|
||
return 0.0
|
||
try:
|
||
return float(item.get("score") or item.get("socre") or 0.0)
|
||
except (TypeError, ValueError):
|
||
return 0.0
|
||
|
||
|
||
def _normalize_rag_text_for_dedupe(text: Any) -> str:
|
||
value = unicodedata.normalize("NFKC", str(text or ""))
|
||
value = re.sub(r"!\[[^\]]*\]\([^)]+\)", "", value)
|
||
value = re.sub(r"\[([^\]]*)\]\([^)]+\)", r"\1", value)
|
||
value = re.sub(r"https?://\S+", "", value, flags=re.IGNORECASE)
|
||
value = re.sub(r"[/\\]?picture[/\\]\S+", "", value, flags=re.IGNORECASE)
|
||
value = re.sub(r"\s+", "", value).lower()
|
||
return "".join(ch for ch in value if ch.isalnum())
|
||
|
||
|
||
def _char_ngrams(text: str, n: int) -> set:
|
||
if len(text) <= n:
|
||
return {text} if text else set()
|
||
return {text[i:i + n] for i in range(len(text) - n + 1)}
|
||
|
||
|
||
def _rag_texts_are_similar(text_a: str, text_b: str, threshold: float) -> bool:
|
||
if not text_a or not text_b:
|
||
return False
|
||
if text_a == text_b:
|
||
return True
|
||
|
||
min_len = min(len(text_a), len(text_b))
|
||
max_len = max(len(text_a), len(text_b))
|
||
if min_len < 8:
|
||
return False
|
||
|
||
if min_len >= RAG_DEDUP_MIN_TEXT_LENGTH and (text_a in text_b or text_b in text_a):
|
||
return (min_len / max_len) >= RAG_DEDUP_CONTAINMENT_THRESHOLD
|
||
|
||
if min_len < RAG_DEDUP_MIN_TEXT_LENGTH:
|
||
return SequenceMatcher(None, text_a, text_b).ratio() >= 0.95
|
||
|
||
ngram_size = 3 if min_len >= 30 else 2
|
||
grams_a = _char_ngrams(text_a, ngram_size)
|
||
grams_b = _char_ngrams(text_b, ngram_size)
|
||
if not grams_a or not grams_b:
|
||
return False
|
||
|
||
jaccard = len(grams_a & grams_b) / len(grams_a | grams_b)
|
||
if jaccard >= threshold:
|
||
return True
|
||
|
||
return SequenceMatcher(None, text_a, text_b).ratio() >= threshold
|
||
|
||
|
||
def dedupe_rag_results(
|
||
items: Any,
|
||
limit: Optional[int] = None,
|
||
similarity_threshold: float = RAG_DEDUP_SIMILARITY_THRESHOLD,
|
||
) -> Any:
|
||
"""
|
||
Deduplicate RAG result chunks while preserving the original item shape.
|
||
|
||
Higher-score items are kept first. Graph results are intentionally not handled
|
||
here; callers should pass only text RAG chunks.
|
||
"""
|
||
if not isinstance(items, list):
|
||
return items
|
||
if len(items) < 2:
|
||
return items[:limit] if limit else items
|
||
|
||
prepared = []
|
||
for idx, item in enumerate(items):
|
||
prepared.append({
|
||
"item": item,
|
||
"index": idx,
|
||
"score": _extract_rag_score(item),
|
||
"normalized_text": _normalize_rag_text_for_dedupe(_extract_rag_text(item)),
|
||
})
|
||
|
||
prepared.sort(key=lambda entry: (-entry["score"], entry["index"]))
|
||
|
||
kept = []
|
||
for entry in prepared:
|
||
current_text = entry["normalized_text"]
|
||
is_duplicate = False
|
||
if current_text:
|
||
for kept_entry in kept:
|
||
if _rag_texts_are_similar(
|
||
current_text,
|
||
kept_entry["normalized_text"],
|
||
similarity_threshold,
|
||
):
|
||
is_duplicate = True
|
||
break
|
||
if not is_duplicate:
|
||
kept.append(entry)
|
||
|
||
deduped = [entry["item"] for entry in kept]
|
||
return deduped[:limit] if limit else deduped
|
||
|
||
|
||
def convert_rag_result(rag_result):
|
||
if not rag_result or not rag_result.get("success", False):
|
||
return []
|
||
source_citation = rag_result.get("sourceCitation", {})
|
||
flat_results = []
|
||
for resource, items in source_citation.items():
|
||
for item in items:
|
||
flat_results.append(item)
|
||
return flat_results
|
||
|
||
|
||
def calculate_info_completeness(weights: Dict[str, float], values: Dict[str, str]) -> float:
|
||
score = 0.0
|
||
for field, weight in weights.items():
|
||
value = values.get(field, "") or ""
|
||
if value and str(value).strip():
|
||
score += weight
|
||
return round(score, 2)
|
||
|
||
|
||
def emit_callback_event(title_options: List[str], details: str):
|
||
start_event = {"type": "function_execution", "title": random.choice(title_options), "details": details}
|
||
for callback in get_all_callbacks():
|
||
try:
|
||
callback(start_event)
|
||
except Exception:
|
||
pass
|
||
|
||
|
||
async def resolve_ship_number_for_workflow(
|
||
ship_number: str,
|
||
context: str = "workflow",
|
||
) -> tuple[str, Dict[str, Any]]:
|
||
"""
|
||
工作流内舷号/舰名标准化。
|
||
|
||
能映射到具体舷号时返回映射后的舷号;失败或未匹配时返回原值。
|
||
"""
|
||
original_ship_number = str(ship_number or "").strip()
|
||
if not original_ship_number:
|
||
return original_ship_number, {}
|
||
|
||
try:
|
||
from utils.ship_number_search import smart_ship_number_mapping
|
||
|
||
matched_ship_number, matched_model, all_numbers, match_type = await smart_ship_number_mapping(
|
||
original_ship_number
|
||
)
|
||
result = {
|
||
"matched_ship_number": matched_ship_number,
|
||
"matched_model": matched_model,
|
||
"all_numbers": all_numbers,
|
||
"match_type": match_type,
|
||
}
|
||
if match_type in ["exact_number", "exact_name", "embedding"] and matched_ship_number:
|
||
mapped_ship_number = str(matched_ship_number).strip()
|
||
if mapped_ship_number and mapped_ship_number != original_ship_number:
|
||
print(
|
||
f"[{context}] 舷号映射: '{original_ship_number}' -> "
|
||
f"'{mapped_ship_number}',匹配类型={match_type}"
|
||
)
|
||
return mapped_ship_number or original_ship_number, result
|
||
|
||
print(f"[{context}] 舷号未映射,继续使用原值: {original_ship_number}")
|
||
return original_ship_number, result
|
||
except Exception as e:
|
||
print(f"[{context}] 舷号映射异常,继续使用原值: {original_ship_number}; 错误: {str(e)}")
|
||
return original_ship_number, {}
|
||
|
||
|
||
async def resolve_ship_scope_for_workflow(
|
||
ship_number: str,
|
||
context: str = "workflow",
|
||
) -> tuple[str, List[str], Dict[str, Any]]:
|
||
"""
|
||
Resolve user ship input into a display hull number and a hull-number scope.
|
||
|
||
Hull number and ship-name matches collapse to one hull. Model matches keep the
|
||
user's original value for display but return every hull number under that model
|
||
so shared device normalization and system lookup can search the whole scope.
|
||
"""
|
||
original_ship_number = str(ship_number or "").strip()
|
||
if not original_ship_number:
|
||
return original_ship_number, [], {}
|
||
|
||
try:
|
||
from utils.ship_number_search import smart_ship_number_mapping
|
||
|
||
matched_ship_number, matched_model, all_numbers, match_type = await smart_ship_number_mapping(
|
||
original_ship_number
|
||
)
|
||
scope_numbers = [str(num).strip() for num in (all_numbers or []) if str(num or "").strip()]
|
||
result = {
|
||
"matched_ship_number": matched_ship_number,
|
||
"matched_model": matched_model,
|
||
"all_numbers": scope_numbers,
|
||
"match_type": match_type,
|
||
}
|
||
|
||
if match_type in ["exact_number", "exact_name", "embedding"] and matched_ship_number:
|
||
mapped_ship_number = str(matched_ship_number).strip()
|
||
return mapped_ship_number or original_ship_number, [mapped_ship_number], result
|
||
|
||
if match_type == "exact_model" and scope_numbers:
|
||
print(
|
||
f"[{context}] ship model scope resolved: '{original_ship_number}' -> "
|
||
f"{matched_model}, hull_numbers={scope_numbers}"
|
||
)
|
||
return original_ship_number, scope_numbers, result
|
||
|
||
print(f"[{context}] ship scope not resolved, continuing with original value: {original_ship_number}")
|
||
return original_ship_number, [], result
|
||
except Exception as e:
|
||
print(f"[{context}] ship scope resolve failed, continuing with original value: {original_ship_number}; error: {str(e)}")
|
||
return original_ship_number, [], {}
|
||
|
||
|
||
async def normalize_device_name_for_workflow(
|
||
device_name: str,
|
||
ship_number: str = "",
|
||
ship_numbers: Optional[List[str]] = None,
|
||
context: str = "workflow",
|
||
) -> tuple[str, Dict[str, Any]]:
|
||
"""
|
||
工作流内设备名称标准化。
|
||
|
||
标准化失败时返回原设备名,避免检索流程被 Neo4j 或 embedding 服务异常阻断。
|
||
"""
|
||
original_device_name = str(device_name or "").strip()
|
||
if not original_device_name:
|
||
return original_device_name, {}
|
||
|
||
try:
|
||
from tools.graph_tools import normalize_device_name_by_ship_graph
|
||
|
||
print(
|
||
f"[{context}] 开始设备名称标准化: device='{original_device_name}', "
|
||
f"ship_number='{ship_number or '未提供'}'"
|
||
)
|
||
call_kwargs = {}
|
||
if "ship_number" in inspect.signature(normalize_device_name_by_ship_graph).parameters:
|
||
call_kwargs["ship_number"] = str(ship_number or "").strip() or None
|
||
if "ship_numbers" in inspect.signature(normalize_device_name_by_ship_graph).parameters:
|
||
scoped_numbers = [str(num).strip() for num in (ship_numbers or []) if str(num or "").strip()]
|
||
call_kwargs["ship_numbers"] = scoped_numbers or None
|
||
elif ship_number:
|
||
print(f"[{context}] 当前进程加载的设备标准化函数不支持 ship_number,请重启服务加载最新 graph_tools.py")
|
||
|
||
result = await normalize_device_name_by_ship_graph(original_device_name, **call_kwargs)
|
||
if result.get("success", False):
|
||
normalized_device_name = str(result.get("device_name") or original_device_name).strip()
|
||
print(
|
||
f"[{context}] 设备名称标准化完成: '{original_device_name}' -> "
|
||
f"'{normalized_device_name or original_device_name}',舷号={ship_number or '未提供'}"
|
||
)
|
||
emit_callback_event(
|
||
["🔎 设备名称标准化", "🔧 设备匹配"],
|
||
f"{original_device_name} -> {normalized_device_name or original_device_name}"
|
||
)
|
||
return normalized_device_name or original_device_name, result
|
||
|
||
print(
|
||
f"[{context}] 设备名称标准化失败,继续使用原设备名: {original_device_name}; "
|
||
f"原因: {result.get('error', '未知错误')}"
|
||
)
|
||
return original_device_name, result
|
||
except Exception as e:
|
||
print(f"[{context}] 设备名称标准化异常,继续使用原设备名: {original_device_name}; 错误: {str(e)}")
|
||
return original_device_name, {}
|
||
|
||
|
||
async def resolve_device_system_for_workflow(
|
||
device_name: str,
|
||
ship_number: str = "",
|
||
ship_numbers: Optional[List[str]] = None,
|
||
context: str = "workflow",
|
||
) -> tuple[str, Dict[str, Any]]:
|
||
"""
|
||
工作流内根据设备名称反推所属系统。
|
||
|
||
失败时返回空字符串,调用方可继续执行主流程。
|
||
"""
|
||
device_name = str(device_name or "").strip()
|
||
if not device_name:
|
||
return "", {}
|
||
|
||
try:
|
||
print(
|
||
f"[{context}] 开始设备反推系统: device='{device_name}', "
|
||
f"ship_number='{ship_number or '未提供'}'"
|
||
)
|
||
from tools.graph_tools import find_system_by_device
|
||
|
||
call_kwargs = {"device_name": device_name}
|
||
if "ship_number" in inspect.signature(find_system_by_device).parameters:
|
||
call_kwargs["ship_number"] = str(ship_number or "").strip() or None
|
||
if "ship_numbers" in inspect.signature(find_system_by_device).parameters:
|
||
scoped_numbers = [str(num).strip() for num in (ship_numbers or []) if str(num or "").strip()]
|
||
call_kwargs["ship_numbers"] = scoped_numbers or None
|
||
elif ship_number:
|
||
print(f"[{context}] 当前进程加载的设备反推系统函数不支持 ship_number,请重启服务加载最新 graph_tools.py")
|
||
|
||
result = await find_system_by_device(**call_kwargs)
|
||
if not result.get("success", False):
|
||
print(f"[{context}] 设备反推系统失败: {result.get('error', '未知错误')}")
|
||
return "", result
|
||
|
||
systems = result.get("systems") or []
|
||
first_system = systems[0] if systems else {}
|
||
system_name = ""
|
||
if isinstance(first_system, dict):
|
||
system_name = str(first_system.get("name") or first_system.get("系统名称") or first_system.get("名称") or "").strip()
|
||
|
||
if system_name:
|
||
print(f"[{context}] 设备反推系统完成: {device_name} -> {system_name}")
|
||
emit_callback_event(
|
||
["🔎 设备所属系统识别", "🧭 系统反查"],
|
||
f"{device_name} -> {system_name}"
|
||
)
|
||
else:
|
||
print(f"[{context}] 设备反推系统无结果: {device_name}")
|
||
return system_name, result
|
||
except Exception as e:
|
||
print(f"[{context}] 设备反推系统异常: {str(e)}")
|
||
return "", {}
|
||
|
||
|
||
def build_detail_info(fields: Dict[str, str]) -> str:
|
||
detail_info = ""
|
||
for label, value in fields.items():
|
||
if value:
|
||
detail_info += f"- {label}:{value}\n"
|
||
return detail_info
|
||
|
||
|
||
def append_atlas_section(answer: str, atlas_results: Dict[str, Any]) -> str:
|
||
if not atlas_results:
|
||
return answer
|
||
atlas_section = "\n\n---\n\n**参考图册:**\n"
|
||
for node_name, node_data in atlas_results.items():
|
||
if isinstance(node_data, dict):
|
||
for title, urls in node_data.items():
|
||
if isinstance(urls, list) and urls:
|
||
for url in urls:
|
||
atlas_section += f"- {title}: [ `{title}` ]({url})\n"
|
||
return answer + atlas_section
|
||
|
||
|
||
|
||
|
||
async def stream_generate_and_postprocess(
|
||
system_prompt: str,
|
||
messages: List[Dict[str, Any]],
|
||
original_image_paths: List[str],
|
||
temperature: float = 0.3,
|
||
fallback: str = "抱歉,无法生成回答。",
|
||
) -> str:
|
||
"""单次流式生成最终回复,并在本地做稳定的格式与图片校验。"""
|
||
stream_handler = get_current_stream_handler()
|
||
answer = ""
|
||
accumulated_content = ""
|
||
|
||
async for chunk in OpenaiAPI.open_api_chat_stream(
|
||
model=None,
|
||
system_prompt=system_prompt,
|
||
messages=messages,
|
||
temperature=temperature,
|
||
):
|
||
answer += chunk
|
||
accumulated_content += chunk
|
||
if stream_handler:
|
||
stream_handler.send_stream_content(accumulated_content)
|
||
|
||
answer = answer.strip() if answer else fallback
|
||
answer = re.sub(r'ynchroneg>.*?ost switching>', '', answer, flags=re.DOTALL)
|
||
answer = remove_duplicate_content(answer)
|
||
answer = normalize_markdown_images(answer, original_paths=original_image_paths)
|
||
answer = normalize_latex_spacing(answer)
|
||
|
||
return answer.strip() if answer else fallback
|
||
|