test_graph.py 3.1 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879808182838485868788
  1. import json
  2. from unittest.mock import patch, MagicMock
  3. from src.schemas.questionnaire import GenerateRequest, Questionnaire, Question, Option
  4. def _mock_llm_response():
  5. """返回合规 JSON 的 mock LLM 响应"""
  6. mock_msg = MagicMock()
  7. mock_msg.content = json.dumps({
  8. "version": 1,
  9. "questions": [
  10. {
  11. "id": "q1",
  12. "dimension": "trust",
  13. "direction": "positive",
  14. "weight": 1.0,
  15. "text": "您对小明的信任程度如何?",
  16. "options": [
  17. {"id": "a", "score": 0},
  18. {"id": "b", "score": 1},
  19. {"id": "c", "score": 2},
  20. ],
  21. }
  22. ]
  23. }, ensure_ascii=False)
  24. return mock_msg
  25. def _mock_llm_invalid_response():
  26. mock_msg = MagicMock()
  27. mock_msg.content = "这不是JSON"
  28. return mock_msg
  29. class TestQuestionnaireSchema:
  30. def test_valid_questionnaire(self):
  31. q = Questionnaire(version=1, questions=[
  32. Question(id="q1", dimension="trust", text="test",
  33. options=[Option(id="a", score=0)])
  34. ])
  35. assert q.validate_structure() is True
  36. def test_empty_questions_fails_validation(self):
  37. q = Questionnaire(version=1, questions=[])
  38. assert q.validate_structure() is False
  39. def test_question_without_options_or_scale_fails(self):
  40. q = Questionnaire(version=1, questions=[
  41. Question(id="q1", dimension="trust", text="test")
  42. ])
  43. assert q.validate_structure() is False
  44. class TestGraph:
  45. def test_graph_valid_output(self):
  46. with patch("src.graphs.questionnaire.get_llm") as mock_get_llm:
  47. mock_llm = MagicMock()
  48. mock_llm.invoke.return_value = _mock_llm_response()
  49. mock_get_llm.return_value = mock_llm
  50. from src.graphs.questionnaire import build_questionnaire_graph
  51. graph = build_questionnaire_graph()
  52. result = graph.invoke({
  53. "request": GenerateRequest(member_name="小明", relationship_type="child"),
  54. "raw_response": "",
  55. "questionnaire": None,
  56. "error": None,
  57. })
  58. assert result["error"] is None
  59. assert result["questionnaire"] is not None
  60. assert len(result["questionnaire"]["questions"]) == 1
  61. def test_graph_invalid_llm_response_rejected(self):
  62. with patch("src.graphs.questionnaire.get_llm") as mock_get_llm:
  63. mock_llm = MagicMock()
  64. mock_llm.invoke.return_value = _mock_llm_invalid_response()
  65. mock_get_llm.return_value = mock_llm
  66. from src.graphs.questionnaire import build_questionnaire_graph
  67. graph = build_questionnaire_graph()
  68. result = graph.invoke({
  69. "request": GenerateRequest(member_name="小明", relationship_type="child"),
  70. "raw_response": "",
  71. "questionnaire": None,
  72. "error": None,
  73. })
  74. assert result["error"] is not None
  75. assert result["questionnaire"] is None