40 lines
1.2 KiB
Python
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()
|