test/app/core/hotword_resolver.py

217 lines
8.3 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

# -*- coding: utf-8 -*-
"""Rule-based hotword correction for ASR post-processing."""
from __future__ import annotations
import re
from difflib import SequenceMatcher
from typing import Any
try:
from pypinyin import lazy_pinyin
except Exception: # pragma: no cover - optional runtime enhancement.
lazy_pinyin = None # type: ignore[assignment]
_HOTWORD_PROMPT_LEAK_PATTERNS = (
re.compile(r"^\s*Use this context when resolving named entities:\s*", re.IGNORECASE),
re.compile(r"^\s*(?:上下文信息[::]?\s*)?热词列表[::]\s*[[\[][^\]]]{0,500}[]\]]\s*[,。,::;;\s]*"),
)
def parse_hotwords(hotwords: str | list[dict[str, object]] | None) -> list[dict[str, Any]]:
if isinstance(hotwords, list):
return _extract_hotword_entries(hotwords)
raw_text = str(hotwords or "").strip()
if not raw_text:
return []
tokens = [part for part in re.split(r"[\s,,;;]+", raw_text) if part]
entries: list[dict[str, Any]] = []
index = 0
order = 0
while index < len(tokens):
word = tokens[index].strip()
if not word:
index += 1
continue
weight = 1.0
if index + 1 < len(tokens):
try:
weight = float(tokens[index + 1])
index += 2
except ValueError:
index += 1
else:
index += 1
entries.append({"text": word, "weight": weight, "order": order})
order += 1
return entries
def format_hotword_prompt_context(hotwords: str | list[dict[str, object]] | None) -> str:
entries = parse_hotwords(hotwords)
if not entries:
return ""
unique_words: list[str] = []
for entry in entries:
word = str(entry["text"]).strip()
if word and word not in unique_words:
unique_words.append(word)
if not unique_words:
return ""
return f"热词列表:[{', '.join(unique_words)}]"
def strip_hotword_prompt_leakage(text: str) -> str:
cleaned = str(text or "")
changed = True
while changed and cleaned:
changed = False
for pattern in _HOTWORD_PROMPT_LEAK_PATTERNS:
updated, count = pattern.subn("", cleaned, count=1)
if count:
cleaned = updated
changed = True
return cleaned.strip()
def _extract_hotword_entries(hotwords: list[dict[str, object]] | None) -> list[dict[str, Any]]:
entries: list[dict[str, Any]] = []
for index, item in enumerate(hotwords or []):
hotword_text = str(item.get("hotword", "")).strip()
if not hotword_text:
continue
try:
weight = float(item.get("weight", 1.0))
except Exception:
weight = 1.0
entries.append({"text": hotword_text, "weight": weight, "order": index})
return entries
def _normalize_ascii_token(text: str) -> str:
return re.sub(r"[^a-z0-9]+", "", str(text or "").lower())
def _common_prefix_length(left: str, right: str) -> int:
matched = 0
for left_char, right_char in zip(left, right):
if left_char != right_char:
break
matched += 1
return matched
def _common_suffix_length(left: str, right: str) -> int:
return _common_prefix_length(left[::-1], right[::-1])
def _normalize_cjk_pinyin(text: str) -> tuple[str, ...]:
if lazy_pinyin is None:
return ()
return tuple(part.strip().lower() for part in lazy_pinyin(str(text or ""), errors="ignore") if str(part).strip())
def _replace_ascii_hotwords(text: str, hotword_entries: list[dict[str, Any]]) -> tuple[str, bool, list[str]]:
updated_text = str(text or "")
changed = False
matched_hotwords: list[str] = []
ascii_pattern = re.compile(r"[A-Za-z][A-Za-z0-9\s._-]{0,40}")
for entry in hotword_entries:
hotword = str(entry["text"])
if not re.search(r"[A-Za-z]", hotword):
continue
normalized_hotword = _normalize_ascii_token(hotword)
if not normalized_hotword:
continue
def _replace_match(match: re.Match[str]) -> str:
nonlocal changed
candidate = match.group(0)
if _normalize_ascii_token(candidate) == normalized_hotword and candidate != hotword:
changed = True
if hotword not in matched_hotwords:
matched_hotwords.append(hotword)
return hotword
return candidate
updated_text = ascii_pattern.sub(_replace_match, updated_text)
return updated_text, changed, matched_hotwords
def _score_cjk_candidate(candidate: str, hotword_entry: dict[str, Any]) -> tuple[int, float, int, float, int] | None:
hotword = str(hotword_entry["text"])
if candidate == hotword or len(candidate) != len(hotword):
return None
candidate_pinyin = _normalize_cjk_pinyin(candidate)
hotword_pinyin = _normalize_cjk_pinyin(hotword)
if candidate_pinyin and candidate_pinyin == hotword_pinyin:
return (3, float(hotword_entry["weight"]), len(hotword), 1.0, -int(hotword_entry["order"]))
if len(hotword) <= 2:
return None
prefix_length = _common_prefix_length(candidate, hotword)
suffix_length = _common_suffix_length(candidate, hotword)
similarity = SequenceMatcher(None, candidate, hotword).ratio()
if prefix_length >= len(hotword) - 1 or suffix_length >= len(hotword) - 1:
return (2, float(hotword_entry["weight"]), prefix_length + suffix_length, similarity, -int(hotword_entry["order"]))
if similarity >= 0.67 and prefix_length >= 1 and suffix_length >= 1:
return (2, float(hotword_entry["weight"]), prefix_length + suffix_length, similarity, -int(hotword_entry["order"]))
return None
def _replace_cjk_hotwords(text: str, hotword_entries: list[dict[str, Any]]) -> tuple[str, bool, list[str]]:
updated_text = str(text or "")
changed = False
matched_hotwords: list[str] = []
chinese_entries = [entry for entry in hotword_entries if re.fullmatch(r"[\u4e00-\u9fff]{2,12}", str(entry["text"]))]
if not chinese_entries:
return updated_text, False, matched_hotwords
candidate_lengths = sorted({len(str(entry["text"])) for entry in chinese_entries})
token_pattern = re.compile(r"[\u4e00-\u9fff]{2,24}")
def _replace_match(match: re.Match[str]) -> str:
nonlocal changed
token = match.group(0)
best_choice: tuple[tuple[int, float, int, float, int], int, int, str] | None = None
for hotword_length in candidate_lengths:
if hotword_length > len(token):
continue
scoped_entries = [entry for entry in chinese_entries if len(str(entry["text"])) == hotword_length]
for index in range(0, len(token) - hotword_length + 1):
candidate = token[index:index + hotword_length]
for entry in scoped_entries:
score = _score_cjk_candidate(candidate, entry)
if score is None:
continue
current_choice = (score, index, hotword_length, str(entry["text"]))
if best_choice is None or current_choice > best_choice:
best_choice = current_choice
if best_choice is None:
return token
_, begin_index, hotword_length, hotword = best_choice
token_chars = list(token)
token_chars[begin_index:begin_index + hotword_length] = list(hotword)
changed = True
if hotword not in matched_hotwords:
matched_hotwords.append(hotword)
return "".join(token_chars)
updated_text = token_pattern.sub(_replace_match, updated_text)
return updated_text, changed, matched_hotwords
def apply_hotword_rules(text: str, hotwords: str | list[dict[str, object]] | None = None) -> tuple[str, bool, list[str]]:
hotword_entries = parse_hotwords(hotwords)
if not hotword_entries:
return str(text or ""), False, []
updated_text, ascii_changed, ascii_matches = _replace_ascii_hotwords(str(text or ""), hotword_entries)
updated_text, cjk_changed, cjk_matches = _replace_cjk_hotwords(updated_text, hotword_entries)
matched_hotwords: list[str] = []
for hotword_text in ascii_matches + cjk_matches:
if hotword_text not in matched_hotwords:
matched_hotwords.append(hotword_text)
return updated_text, ascii_changed or cjk_changed, matched_hotwords