210 lines
9.0 KiB
Python
210 lines
9.0 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import json
|
|
import unittest
|
|
from datetime import datetime, timedelta, timezone
|
|
from types import SimpleNamespace
|
|
from unittest.mock import MagicMock, patch
|
|
from uuid import uuid4
|
|
|
|
from core.wechat.context_router import (
|
|
RouteContext,
|
|
_fallback,
|
|
_parse_result,
|
|
_recent_two_rounds,
|
|
_serialize_route_input,
|
|
fresh_task_values,
|
|
route_channel_context,
|
|
)
|
|
|
|
|
|
class ContextRouterUnitTests(unittest.TestCase):
|
|
def test_high_confidence_carry_survives_next_day(self) -> None:
|
|
ctx = RouteContext(
|
|
total_messages=8,
|
|
last_user_at=datetime.now(timezone.utc) - timedelta(days=1),
|
|
context_summary="正在修改报告",
|
|
history=(("user", "修改第一章"), ("assistant", "已完成")),
|
|
)
|
|
response = SimpleNamespace(
|
|
choices=[SimpleNamespace(message=SimpleNamespace(
|
|
content='{"decision":"carry","confidence":0.96,"reason":"依赖上一版报告"}'
|
|
))],
|
|
usage=None,
|
|
)
|
|
llm = MagicMock()
|
|
llm.chat.return_value = response
|
|
with (
|
|
patch("core.wechat.context_router._load_context", return_value=ctx),
|
|
patch("core.wechat.context_router.ModelCapabilities.load", return_value=MagicMock()),
|
|
patch("core.wechat.context_router.LLM", return_value=llm),
|
|
patch("core.wechat.context_router.record_chat_usage"),
|
|
patch("core.wechat.context_router._apply_fresh") as apply_fresh,
|
|
):
|
|
result = route_channel_context(
|
|
task_id=uuid4(), user_id=uuid4(), current_message="继续修改第二章"
|
|
)
|
|
self.assertTrue(result.carry)
|
|
apply_fresh.assert_not_called()
|
|
self.assertEqual(llm.chat.call_args.kwargs["timeout_s"], 8.0)
|
|
self.assertEqual(llm.chat.call_args.kwargs["reasoning_effort"], "low")
|
|
|
|
def test_short_gap_independent_question_is_fresh_and_clears_old_window(self) -> None:
|
|
tid = uuid4()
|
|
ctx = RouteContext(
|
|
total_messages=5,
|
|
last_user_at=datetime.now(timezone.utc) - timedelta(minutes=2),
|
|
context_summary="旧摘要",
|
|
history=(("user", "上一题"), ("assistant", "上一题答案")),
|
|
)
|
|
response = SimpleNamespace(
|
|
choices=[SimpleNamespace(message=SimpleNamespace(
|
|
content='{"decision":"fresh","confidence":0.93,"reason":"问题可独立回答"}'
|
|
))],
|
|
usage=None,
|
|
)
|
|
with (
|
|
patch("core.wechat.context_router._load_context", return_value=ctx),
|
|
patch("core.wechat.context_router.ModelCapabilities.load", return_value=MagicMock()),
|
|
patch("core.wechat.context_router.LLM") as llm_cls,
|
|
patch("core.wechat.context_router.record_chat_usage"),
|
|
patch("core.wechat.context_router._apply_fresh") as apply_fresh,
|
|
):
|
|
llm_cls.return_value.chat.return_value = response
|
|
result = route_channel_context(
|
|
task_id=tid, user_id=uuid4(), current_message="今天北京天气如何"
|
|
)
|
|
self.assertEqual(result.decision, "fresh")
|
|
apply_fresh.assert_called_once_with(tid, 5)
|
|
|
|
def test_push_and_tool_rows_do_not_enter_recent_context(self) -> None:
|
|
rows = [
|
|
SimpleNamespace(payload={"role": "assistant", "content": "主动简报"}, kind="push"),
|
|
SimpleNamespace(payload={"role": "tool", "content": "工具结果"}, kind=None),
|
|
SimpleNamespace(payload={"role": "assistant", "content": "答复二"}, kind=None),
|
|
SimpleNamespace(payload={"role": "user", "content": "问题二"}, kind=None),
|
|
SimpleNamespace(payload={"role": "assistant", "content": "答复一"}, kind=None),
|
|
SimpleNamespace(payload={"role": "user", "content": "问题一"}, kind=None),
|
|
SimpleNamespace(payload={"role": "user", "content": "更早问题"}, kind=None),
|
|
]
|
|
self.assertEqual(
|
|
_recent_two_rounds(rows),
|
|
(("user", "问题一"), ("assistant", "答复一"), ("user", "问题二"), ("assistant", "答复二")),
|
|
)
|
|
|
|
def test_failure_falls_back_to_six_hour_user_gap(self) -> None:
|
|
now = datetime.now(timezone.utc)
|
|
self.assertEqual(_fallback(now - timedelta(hours=5), 6), "carry")
|
|
self.assertEqual(_fallback(now - timedelta(hours=7), 6), "fresh")
|
|
|
|
def test_router_exception_applies_fresh_for_stale_user(self) -> None:
|
|
tid = uuid4()
|
|
ctx = RouteContext(
|
|
total_messages=11,
|
|
last_user_at=datetime.now(timezone.utc) - timedelta(hours=7),
|
|
context_summary="旧摘要",
|
|
history=(("user", "旧问题"),),
|
|
)
|
|
with (
|
|
patch("core.wechat.context_router._load_context", return_value=ctx),
|
|
patch("core.wechat.context_router.ModelCapabilities.load", return_value=MagicMock()),
|
|
patch("core.wechat.context_router.LLM", side_effect=RuntimeError("down")),
|
|
patch("core.wechat.context_router._apply_fresh") as apply_fresh,
|
|
):
|
|
result = route_channel_context(
|
|
task_id=tid, user_id=uuid4(), current_message="一个新问题"
|
|
)
|
|
self.assertEqual((result.decision, result.source), ("fresh", "fallback"))
|
|
apply_fresh.assert_called_once_with(tid, 11)
|
|
|
|
def test_fresh_values_advance_base_and_clear_summary(self) -> None:
|
|
self.assertEqual(
|
|
fresh_task_values(17),
|
|
{"context_base_idx": 17, "context_summary": None},
|
|
)
|
|
|
|
def test_serialized_router_input_has_a_hard_character_limit(self) -> None:
|
|
encoded = _serialize_route_input({
|
|
"current_message": "\\\n" * 5000,
|
|
"last_user_at": datetime.now(timezone.utc).isoformat(),
|
|
"context_summary": "\t" * 5000,
|
|
"recent_messages": [
|
|
{"role": "user", "content": '"' * 5000},
|
|
{"role": "assistant", "content": "a" * 5000},
|
|
],
|
|
})
|
|
self.assertLessEqual(len(encoded), 10000)
|
|
self.assertIsInstance(json.loads(encoded), dict)
|
|
|
|
def test_uncertain_and_low_confidence_carry_default_to_fresh(self) -> None:
|
|
self.assertEqual(_parse_result('{"decision":"uncertain","confidence":0.9}')[0], "fresh")
|
|
self.assertEqual(_parse_result('{"decision":"carry","confidence":0.79}')[0], "fresh")
|
|
|
|
def test_exact_continue_is_local_and_does_not_call_model(self) -> None:
|
|
ctx = RouteContext(
|
|
total_messages=3,
|
|
last_user_at=datetime.now(timezone.utc) - timedelta(days=2),
|
|
context_summary="",
|
|
history=(("user", "旧问题"),),
|
|
)
|
|
with (
|
|
patch("core.wechat.context_router._load_context", return_value=ctx),
|
|
patch("core.wechat.context_router.LLM") as llm_cls,
|
|
):
|
|
result = route_channel_context(
|
|
task_id=uuid4(), user_id=uuid4(), current_message="继续上文"
|
|
)
|
|
self.assertEqual((result.decision, result.source), ("carry", "local_continue"))
|
|
llm_cls.assert_not_called()
|
|
|
|
|
|
class SharedChannelEntryTests(unittest.IsolatedAsyncioTestCase):
|
|
async def test_new_topic_is_hard_reset_without_router(self) -> None:
|
|
from web.runs import run_channel_conversation
|
|
|
|
tid = uuid4()
|
|
with (
|
|
patch("core.wechat.service.ensure_channel_chat_task", return_value=tid),
|
|
patch("core.wechat.service.reset_channel_context") as reset,
|
|
patch("core.wechat.context_router.route_channel_context") as router,
|
|
):
|
|
reply = await run_channel_conversation(
|
|
MagicMock(), uuid4(), "新话题", [], channel="wechat"
|
|
)
|
|
self.assertIn("已开启新话题", reply)
|
|
reset.assert_called_once_with(tid, hard=True)
|
|
router.assert_not_called()
|
|
|
|
async def test_wechat_and_wecom_use_same_router_entry(self) -> None:
|
|
from web.runs import run_channel_conversation
|
|
from web.run_lifecycle import RunTaskBusy
|
|
|
|
tid = uuid4()
|
|
uid = uuid4()
|
|
|
|
async def run(channel: str) -> None:
|
|
with (
|
|
patch("core.wechat.service.ensure_channel_chat_task", return_value=tid),
|
|
patch("core.agent_builder.resolve_workspace", return_value=MagicMock()),
|
|
patch("core.shortcuts.expand", return_value=("独立问题", None)),
|
|
patch("core.wechat.context_router.route_channel_context") as router,
|
|
patch(
|
|
"web.run_lifecycle.claim_run_with_message",
|
|
side_effect=RunTaskBusy("running"),
|
|
),
|
|
):
|
|
router.return_value = SimpleNamespace(carry=False)
|
|
reply = await run_channel_conversation(
|
|
MagicMock(), uid, "独立问题", [], channel=channel
|
|
)
|
|
self.assertEqual(reply, "上一条还在处理中,请稍候再发。")
|
|
router.assert_called_once()
|
|
|
|
await run("wechat")
|
|
await run("wecom")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|