test_qa_question_context.py 9.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183
  1. """Offline semantic context routing and state propagation regression tests."""
  2. from __future__ import annotations
  3. import copy
  4. import hashlib
  5. import json
  6. import tempfile
  7. import unittest
  8. from pathlib import Path
  9. from unittest.mock import patch
  10. from langgraph.graph import END, START, StateGraph
  11. from step3_qa_agent.agent import nodes
  12. from step3_qa_agent.agent.state import AgentState
  13. from step3_qa_agent.agent import question_context as qc
  14. from step3_qa_agent.agent.schema_context import build_schema_text
  15. class QuestionContextTests(unittest.TestCase):
  16. def setUp(self):
  17. self.tmp = tempfile.TemporaryDirectory()
  18. self.root = Path(self.tmp.name)
  19. self.schema_path = self.root / 'schema.json'
  20. self.assessment_path = self.root / 'qa_data_context.json'
  21. self.schema = {'meta': {'schema_version': 2, 'build_id': 'version-one'},
  22. 'nodes': [{'id': '人员信息', 'name': '人员信息', 'attributes': ['工号', '岗位名称'],
  23. 'count': 2, 'source_rows': 3, 'identity_fields': ['工号'],
  24. 'multivalue_fields': {'岗位名称': 1}, 'active': True},
  25. {'id': '项目信息', 'name': '项目信息', 'attributes': ['项目编号', '项目名称'], 'active': True}],
  26. 'relations': [{'id': 'service-rule', 'source': '人员信息', 'target': '项目信息',
  27. 'type': '服务于', 'key': '服务项目-项目名称', 'method': '包含', 'active': True}]}
  28. reports = []
  29. for n in self.schema['nodes']:
  30. content = ('private-source-value-' + n['id']).encode()
  31. (self.root / (n['id'] + '.xlsx')).write_bytes(content)
  32. reports.append({'dataset': n['id'], 'source_sha256': hashlib.sha256(content).hexdigest(),
  33. 'row_count': 3, 'completeness_rate': .9, 'llm_quality_score': 80,
  34. 'summary': n['id'] + '-ASSESSMENT-BODY',
  35. 'routing_terms': [n['id']],
  36. 'record_granularity': {'one_row_represents': '源记录', 'candidate_business_key': [n['attributes'][0]]},
  37. 'fields': [{'field': f, 'missing_rate': .1, 'business_meaning': f,
  38. 'qa_usage': {'cautions': ['不要把缺失当成零']}} for f in n['attributes']],
  39. 'qa_guidance': {'interpretation_notes': ['源行数不等于节点数']}})
  40. self.payload = {'raw_values_sent_to_llm': False, 'datasets': reports}
  41. self.write()
  42. self.patches = [patch.object(qc, 'SCHEMA_PATH', self.schema_path),
  43. patch.object(qc, 'DEFAULT_CONTEXT_PATH', self.assessment_path),
  44. patch.object(qc, 'PRODUCTION_DIR', self.root)]
  45. for p in self.patches:
  46. p.start()
  47. def tearDown(self):
  48. for p in reversed(self.patches):
  49. p.stop()
  50. self.tmp.cleanup()
  51. def write(self):
  52. self.schema_path.write_text(json.dumps(self.schema, ensure_ascii=False), encoding='utf-8')
  53. self.assessment_path.write_text(json.dumps(self.payload, ensure_ascii=False), encoding='utf-8')
  54. def select(self, needs):
  55. return qc.select_assessments(qc.load_question_context(), needs)
  56. def test_semantic_selection_after_understanding_reaches_plan(self):
  57. captured = []
  58. def llm(system, user):
  59. captured.append(system)
  60. if '问题理解器' in system:
  61. return {'category': '图谱检索', 'slots': {'intent': '其他', 'anchors': []},
  62. 'data_needs': {'datasets': ['人员信息'], 'relations': []}}
  63. return {'steps': [{'step_id': 's1', 'tool': '人员', 'params': {}, 'fields': ['工号'], 'depends': []}]}
  64. graph = StateGraph(AgentState)
  65. graph.add_node('understand', nodes.understand)
  66. graph.add_node('prepare_plan', lambda state: {'grounded': state['slots']})
  67. graph.add_node('plan', nodes.plan)
  68. graph.add_edge(START, 'understand')
  69. graph.add_edge('understand', 'prepare_plan')
  70. graph.add_edge('prepare_plan', 'plan')
  71. graph.add_edge('plan', END)
  72. with patch.object(nodes, 'chat_json', side_effect=llm):
  73. state = graph.compile().invoke({'question': '他们的岗位分别是什么?',
  74. 'messages': [{'role': 'assistant', 'content': '上轮正在讨论员工。'}]})
  75. self.assertIn('service-rule', captured[0])
  76. self.assertIn('multivalue_fields', captured[0])
  77. self.assertNotIn('ASSESSMENT-BODY', captured[0])
  78. self.assertIn('人员信息-ASSESSMENT-BODY', captured[1])
  79. self.assertNotIn('项目信息-ASSESSMENT-BODY', captured[1])
  80. self.assertNotIn('private-source-value', captured[1])
  81. self.assertEqual(state['qa_context']['selected_datasets'], ['人员信息'])
  82. self.assertNotIn('_reports', state['qa_context'])
  83. self.assertEqual(state['trace'][0]['selection_source'], 'understanding')
  84. self.assertEqual(state['trace'][0]['schema_build_id'], 'version-one')
  85. def test_relation_selection_includes_both_endpoints(self):
  86. context = self.select({'datasets': [], 'relations': ['service-rule']})
  87. self.assertEqual(context['loaded_datasets'], ['人员信息', '项目信息'])
  88. self.assertIn('人员信息-ASSESSMENT-BODY', context['assessment_text'])
  89. self.assertIn('项目信息-ASSESSMENT-BODY', context['assessment_text'])
  90. def test_new_turn_reload_and_within_turn_snapshot(self):
  91. answers = [{'category': '图谱检索', 'slots': {}, 'data_needs': {'datasets': ['人员信息'], 'relations': []}},
  92. {'category': '图谱检索', 'slots': {}, 'data_needs': {'datasets': ['项目信息'], 'relations': []}}]
  93. with patch.object(nodes, 'chat_json', side_effect=answers):
  94. first = nodes.understand({'question': '同一个问题'})
  95. self.schema['meta']['build_id'] = 'version-two'
  96. self.payload['datasets'][1]['summary'] = 'NEW-PROJECT-ASSESSMENT'
  97. self.write()
  98. second = nodes.understand({**first, 'question': '同一个问题'})
  99. self.assertEqual(first['qa_context']['build_id'], 'version-one')
  100. self.assertEqual(second['qa_context']['build_id'], 'version-two')
  101. frozen = build_schema_text(state=first)
  102. fresh = build_schema_text(state=second)
  103. self.assertIn('version-one', frozen)
  104. self.assertNotIn('version-two', frozen)
  105. self.assertIn('NEW-PROJECT-ASSESSMENT', fresh)
  106. self.assertNotIn('人员信息-ASSESSMENT-BODY', fresh)
  107. def test_stale_report_omitted_but_schema_kept(self):
  108. (self.root / '人员信息.xlsx').write_bytes(b'new content')
  109. context = self.select({'datasets': ['人员信息'], 'relations': []})
  110. self.assertEqual(context['loaded_datasets'], [])
  111. self.assertIn('不一致', ' '.join(context['issues']))
  112. self.assertIn('version-one', build_schema_text(state={'qa_context': context}))
  113. self.assertNotIn('ASSESSMENT-BODY', build_schema_text(state={'qa_context': context}))
  114. def test_invalid_selection_never_injects_all_reports(self):
  115. context = self.select({'datasets': ['不存在'], 'relations': ['unknown-relation']})
  116. self.assertEqual(context['selected_datasets'], [])
  117. self.assertEqual(context['assessment_text'], '')
  118. self.assertTrue(context['issues'])
  119. context = self.select({'datasets': [], 'relations': []})
  120. self.assertEqual(context['assessment_text'], '')
  121. def test_missing_malformed_and_mismatched_reports(self):
  122. self.assessment_path.unlink()
  123. context = self.select({'datasets': ['人员信息'], 'relations': []})
  124. self.assertEqual(context['assessment_text'], '')
  125. self.assertTrue(context['issues'])
  126. self.write()
  127. self.payload['datasets'][0]['fields'] = None
  128. self.write()
  129. context = self.select({'datasets': ['人员信息'], 'relations': []})
  130. self.assertEqual(context['loaded_datasets'], [])
  131. self.payload['raw_values_sent_to_llm'] = True
  132. self.write()
  133. context = self.select({'datasets': ['项目信息'], 'relations': []})
  134. self.assertEqual(context['loaded_datasets'], [])
  135. def test_legacy_llm_reply_uses_current_question_fallback(self):
  136. context = qc.select_assessments(qc.load_question_context(), None,
  137. question='查询人员信息', slots={})
  138. self.assertEqual(context['loaded_datasets'], ['人员信息'])
  139. self.assertEqual(context['selection_source'], 'lexical_fallback')
  140. def test_missing_schema_stops_before_llm(self):
  141. self.schema_path.unlink()
  142. with patch.object(nodes, 'chat_json') as llm:
  143. with self.assertRaisesRegex(ValueError, '元知识图谱'):
  144. nodes.understand({'question': '查询'})
  145. llm.assert_not_called()
  146. def test_old_graph_answers_not_reused(self):
  147. state = {"question": "继续统计", "qa_context": {"build_id": "new"},
  148. "rounds": [{"question": "查询", "answer": "旧数据", "schema_build_id": "old"}]}
  149. with patch.object(nodes, "chat_json") as llm:
  150. result = nodes.reuse_check(state)
  151. self.assertTrue(result["need_full_query"])
  152. llm.assert_not_called()
  153. def test_chat_receives_schema_without_prior_assessment(self):
  154. with patch.object(nodes, 'chat_json', return_value={'category': '闲聊', 'slots': None}):
  155. state = nodes.understand({'question': '你能做什么', 'qa_context': {'assessment_text': 'OLD'}})
  156. captured = []
  157. with patch.object(nodes, 'chat_text', side_effect=lambda system, user, **kwargs: captured.append(system) or 'OK'):
  158. nodes.chat({**state, 'question': '你能做什么'})
  159. self.assertIn('service-rule', captured[0])
  160. self.assertNotIn('ASSESSMENT-BODY', captured[0])
  161. self.assertNotIn('OLD', captured[0])
  162. if __name__ == '__main__':
  163. unittest.main(verbosity=2)