test_direct_reply.py 2.3 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071
  1. from types import SimpleNamespace
  2. import pytest
  3. from data_query_agent.models import IncomingMessage, QueryDecision, SkillParameters
  4. from data_query_agent.service import DataQueryService
  5. def _parameters() -> SkillParameters:
  6. return SkillParameters(**{name: None for name in SkillParameters.model_fields})
  7. class FakeState:
  8. def __init__(self) -> None:
  9. self.saved_thread: tuple[str, str] | None = None
  10. def get_conversation(self, key: str) -> SimpleNamespace:
  11. return SimpleNamespace(thread_id="old-thread", session_id="session-1")
  12. def set_thread(self, key: str, thread_id: str) -> None:
  13. self.saved_thread = (key, thread_id)
  14. def create_run(self, *args: object, **kwargs: object) -> None:
  15. raise AssertionError("direct_reply must not create a query run")
  16. class FakeCodex:
  17. async def plan(self, thread_id: str | None, question: str) -> tuple[str, QueryDecision]:
  18. assert thread_id == "old-thread"
  19. assert question == "分析下刚才结果"
  20. return "new-thread", QueryDecision(
  21. status="ready",
  22. reply="实验组整体效率小幅走弱,主要来自推荐场景回流下降。",
  23. title="数据助手回复",
  24. selected_skill="odps-product-efficiency-report",
  25. execution_mode="direct_reply",
  26. sql=None,
  27. parameters=_parameters(),
  28. assumptions=["仅基于本线程已有查询结果"],
  29. )
  30. class FakeFeishu:
  31. def __init__(self) -> None:
  32. self.replies: list[tuple[str, str]] = []
  33. async def reply_text(self, message_id: str, text: str) -> None:
  34. self.replies.append((message_id, text))
  35. @pytest.mark.asyncio
  36. async def test_direct_reply_uses_thread_context_without_starting_query() -> None:
  37. service = object.__new__(DataQueryService)
  38. service.state = FakeState()
  39. service.codex = FakeCodex()
  40. service.feishu = FakeFeishu()
  41. message = IncomingMessage(
  42. message_id="m1",
  43. chat_id="c1",
  44. chat_type="group",
  45. sender_open_id="u1",
  46. text="分析下刚才结果",
  47. mentioned_bot=True,
  48. )
  49. await service._handle_query(message)
  50. assert service.state.saved_thread == (message.conversation_key, "new-thread")
  51. assert service.feishu.replies == [
  52. ("m1", "实验组整体效率小幅走弱,主要来自推荐场景回流下降。")
  53. ]