test/tests/test_hotword_prompt.py

40 lines
1.2 KiB
Python

# -*- coding: utf-8 -*-
import unittest
from app.core.hotword_resolver import format_hotword_prompt_context
from app.core.hotword_resolver import strip_hotword_prompt_leakage
class HotwordPromptTest(unittest.TestCase):
def test_format_prompt_context_removes_weights(self) -> None:
hotwords = "KPI考核考核 0.2 华智 0.2 合川 0.2 商客市场拓展 0.2"
self.assertEqual(
format_hotword_prompt_context(hotwords),
"热词列表:[KPI考核考核, 华智, 合川, 商客市场拓展]",
)
def test_strip_leaked_hotword_prompt(self) -> None:
text = "热词列表:[KPI考核考核, 华智, 合川, 商客市场拓展] 今天我们看一下执行力。"
self.assertEqual(
strip_hotword_prompt_leakage(text),
"今天我们看一下执行力。",
)
def test_strip_legacy_context_leak_prefix(self) -> None:
text = (
"Use this context when resolving named entities: "
"热词列表:[KPI考核考核, 华智] 今天继续看KPI考核。"
)
self.assertEqual(
strip_hotword_prompt_leakage(text),
"今天继续看KPI考核。",
)
if __name__ == "__main__":
unittest.main()