100 lines
3.5 KiB
Python
100 lines
3.5 KiB
Python
from __future__ import annotations
|
|
|
|
import unittest
|
|
|
|
from owner_voice_pet.continuation import (
|
|
HybridContinuationDecisionProvider,
|
|
LlmContinuationDecisionProvider,
|
|
RuleContinuationDecisionProvider,
|
|
)
|
|
from owner_voice_pet.models import Message, ReplyDelta
|
|
|
|
|
|
class FakeClassifierLlm:
|
|
def __init__(self, text: str) -> None:
|
|
self.text = text
|
|
self.calls: list[list[Message]] = []
|
|
|
|
def stream_reply(self, messages):
|
|
self.calls.append(list(messages))
|
|
yield ReplyDelta(self.text, finish_reason="stop")
|
|
|
|
|
|
class ContinuationDecisionTests(unittest.TestCase):
|
|
def test_rule_continue_for_assistant_question(self) -> None:
|
|
decision = RuleContinuationDecisionProvider().decide(
|
|
user_text="讲讲天气",
|
|
assistant_text="你想继续听哪一部分?",
|
|
history=[],
|
|
)
|
|
|
|
self.assertEqual(decision.action, "continue")
|
|
self.assertGreaterEqual(decision.confidence, 0.65)
|
|
|
|
def test_rule_standby_for_completed_answer(self) -> None:
|
|
decision = RuleContinuationDecisionProvider().decide(
|
|
user_text="今天天气怎么样",
|
|
assistant_text="这是今天的天气。",
|
|
history=[],
|
|
)
|
|
|
|
self.assertEqual(decision.action, "standby")
|
|
self.assertGreaterEqual(decision.confidence, 0.65)
|
|
|
|
def test_hybrid_low_confidence_llm_defaults_to_standby(self) -> None:
|
|
llm = FakeClassifierLlm('{"action":"continue","confidence":0.3,"reason":"弱"}')
|
|
provider = HybridContinuationDecisionProvider(
|
|
RuleContinuationDecisionProvider(),
|
|
LlmContinuationDecisionProvider(llm, threshold=0.65),
|
|
threshold=0.65,
|
|
)
|
|
|
|
decision = provider.decide(
|
|
user_text="继续",
|
|
assistant_text="我还可以从背景原因和下一步影响两个方向继续展开",
|
|
history=[],
|
|
)
|
|
|
|
self.assertEqual(decision.action, "standby")
|
|
self.assertEqual(len(llm.calls), 1)
|
|
|
|
def test_hybrid_generic_completed_reply_does_not_call_llm_classifier(self) -> None:
|
|
llm = FakeClassifierLlm('{"action":"continue","confidence":0.99,"reason":"不应调用"}')
|
|
provider = HybridContinuationDecisionProvider(
|
|
RuleContinuationDecisionProvider(),
|
|
LlmContinuationDecisionProvider(llm, threshold=0.65),
|
|
threshold=0.65,
|
|
)
|
|
|
|
decision = provider.decide(
|
|
user_text="没有呢",
|
|
assistant_text="明白了,我先保持待机。有需要再叫我就行。",
|
|
history=[],
|
|
)
|
|
|
|
self.assertEqual(decision.action, "standby")
|
|
self.assertEqual(decision.provider, "rule")
|
|
self.assertEqual(llm.calls, [])
|
|
|
|
def test_hybrid_high_confidence_llm_can_continue(self) -> None:
|
|
llm = FakeClassifierLlm('{"action":"continue","confidence":0.91,"reason":"等待用户选择"}')
|
|
provider = HybridContinuationDecisionProvider(
|
|
RuleContinuationDecisionProvider(),
|
|
LlmContinuationDecisionProvider(llm, threshold=0.65),
|
|
threshold=0.65,
|
|
)
|
|
|
|
decision = provider.decide(
|
|
user_text="继续",
|
|
assistant_text="我还可以从背景原因和下一步影响两个方向继续展开",
|
|
history=[],
|
|
)
|
|
|
|
self.assertEqual(decision.action, "continue")
|
|
self.assertEqual(decision.provider, "hybrid")
|
|
self.assertEqual(len(llm.calls), 1)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|