# -*- 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