| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183 |
- """Offline semantic context routing and state propagation regression tests."""
- from __future__ import annotations
- import copy
- import hashlib
- import json
- import tempfile
- import unittest
- from pathlib import Path
- from unittest.mock import patch
- from langgraph.graph import END, START, StateGraph
- from step3_qa_agent.agent import nodes
- from step3_qa_agent.agent.state import AgentState
- from step3_qa_agent.agent import question_context as qc
- from step3_qa_agent.agent.schema_context import build_schema_text
- class QuestionContextTests(unittest.TestCase):
- def setUp(self):
- self.tmp = tempfile.TemporaryDirectory()
- self.root = Path(self.tmp.name)
- self.schema_path = self.root / 'schema.json'
- self.assessment_path = self.root / 'qa_data_context.json'
- self.schema = {'meta': {'schema_version': 2, 'build_id': 'version-one'},
- 'nodes': [{'id': '人员信息', 'name': '人员信息', 'attributes': ['工号', '岗位名称'],
- 'count': 2, 'source_rows': 3, 'identity_fields': ['工号'],
- 'multivalue_fields': {'岗位名称': 1}, 'active': True},
- {'id': '项目信息', 'name': '项目信息', 'attributes': ['项目编号', '项目名称'], 'active': True}],
- 'relations': [{'id': 'service-rule', 'source': '人员信息', 'target': '项目信息',
- 'type': '服务于', 'key': '服务项目-项目名称', 'method': '包含', 'active': True}]}
- reports = []
- for n in self.schema['nodes']:
- content = ('private-source-value-' + n['id']).encode()
- (self.root / (n['id'] + '.xlsx')).write_bytes(content)
- reports.append({'dataset': n['id'], 'source_sha256': hashlib.sha256(content).hexdigest(),
- 'row_count': 3, 'completeness_rate': .9, 'llm_quality_score': 80,
- 'summary': n['id'] + '-ASSESSMENT-BODY',
- 'routing_terms': [n['id']],
- 'record_granularity': {'one_row_represents': '源记录', 'candidate_business_key': [n['attributes'][0]]},
- 'fields': [{'field': f, 'missing_rate': .1, 'business_meaning': f,
- 'qa_usage': {'cautions': ['不要把缺失当成零']}} for f in n['attributes']],
- 'qa_guidance': {'interpretation_notes': ['源行数不等于节点数']}})
- self.payload = {'raw_values_sent_to_llm': False, 'datasets': reports}
- self.write()
- self.patches = [patch.object(qc, 'SCHEMA_PATH', self.schema_path),
- patch.object(qc, 'DEFAULT_CONTEXT_PATH', self.assessment_path),
- patch.object(qc, 'PRODUCTION_DIR', self.root)]
- for p in self.patches:
- p.start()
- def tearDown(self):
- for p in reversed(self.patches):
- p.stop()
- self.tmp.cleanup()
- def write(self):
- self.schema_path.write_text(json.dumps(self.schema, ensure_ascii=False), encoding='utf-8')
- self.assessment_path.write_text(json.dumps(self.payload, ensure_ascii=False), encoding='utf-8')
- def select(self, needs):
- return qc.select_assessments(qc.load_question_context(), needs)
- def test_semantic_selection_after_understanding_reaches_plan(self):
- captured = []
- def llm(system, user):
- captured.append(system)
- if '问题理解器' in system:
- return {'category': '图谱检索', 'slots': {'intent': '其他', 'anchors': []},
- 'data_needs': {'datasets': ['人员信息'], 'relations': []}}
- return {'steps': [{'step_id': 's1', 'tool': '人员', 'params': {}, 'fields': ['工号'], 'depends': []}]}
- graph = StateGraph(AgentState)
- graph.add_node('understand', nodes.understand)
- graph.add_node('prepare_plan', lambda state: {'grounded': state['slots']})
- graph.add_node('plan', nodes.plan)
- graph.add_edge(START, 'understand')
- graph.add_edge('understand', 'prepare_plan')
- graph.add_edge('prepare_plan', 'plan')
- graph.add_edge('plan', END)
- with patch.object(nodes, 'chat_json', side_effect=llm):
- state = graph.compile().invoke({'question': '他们的岗位分别是什么?',
- 'messages': [{'role': 'assistant', 'content': '上轮正在讨论员工。'}]})
- self.assertIn('service-rule', captured[0])
- self.assertIn('multivalue_fields', captured[0])
- self.assertNotIn('ASSESSMENT-BODY', captured[0])
- self.assertIn('人员信息-ASSESSMENT-BODY', captured[1])
- self.assertNotIn('项目信息-ASSESSMENT-BODY', captured[1])
- self.assertNotIn('private-source-value', captured[1])
- self.assertEqual(state['qa_context']['selected_datasets'], ['人员信息'])
- self.assertNotIn('_reports', state['qa_context'])
- self.assertEqual(state['trace'][0]['selection_source'], 'understanding')
- self.assertEqual(state['trace'][0]['schema_build_id'], 'version-one')
- def test_relation_selection_includes_both_endpoints(self):
- context = self.select({'datasets': [], 'relations': ['service-rule']})
- self.assertEqual(context['loaded_datasets'], ['人员信息', '项目信息'])
- self.assertIn('人员信息-ASSESSMENT-BODY', context['assessment_text'])
- self.assertIn('项目信息-ASSESSMENT-BODY', context['assessment_text'])
- def test_new_turn_reload_and_within_turn_snapshot(self):
- answers = [{'category': '图谱检索', 'slots': {}, 'data_needs': {'datasets': ['人员信息'], 'relations': []}},
- {'category': '图谱检索', 'slots': {}, 'data_needs': {'datasets': ['项目信息'], 'relations': []}}]
- with patch.object(nodes, 'chat_json', side_effect=answers):
- first = nodes.understand({'question': '同一个问题'})
- self.schema['meta']['build_id'] = 'version-two'
- self.payload['datasets'][1]['summary'] = 'NEW-PROJECT-ASSESSMENT'
- self.write()
- second = nodes.understand({**first, 'question': '同一个问题'})
- self.assertEqual(first['qa_context']['build_id'], 'version-one')
- self.assertEqual(second['qa_context']['build_id'], 'version-two')
- frozen = build_schema_text(state=first)
- fresh = build_schema_text(state=second)
- self.assertIn('version-one', frozen)
- self.assertNotIn('version-two', frozen)
- self.assertIn('NEW-PROJECT-ASSESSMENT', fresh)
- self.assertNotIn('人员信息-ASSESSMENT-BODY', fresh)
- def test_stale_report_omitted_but_schema_kept(self):
- (self.root / '人员信息.xlsx').write_bytes(b'new content')
- context = self.select({'datasets': ['人员信息'], 'relations': []})
- self.assertEqual(context['loaded_datasets'], [])
- self.assertIn('不一致', ' '.join(context['issues']))
- self.assertIn('version-one', build_schema_text(state={'qa_context': context}))
- self.assertNotIn('ASSESSMENT-BODY', build_schema_text(state={'qa_context': context}))
- def test_invalid_selection_never_injects_all_reports(self):
- context = self.select({'datasets': ['不存在'], 'relations': ['unknown-relation']})
- self.assertEqual(context['selected_datasets'], [])
- self.assertEqual(context['assessment_text'], '')
- self.assertTrue(context['issues'])
- context = self.select({'datasets': [], 'relations': []})
- self.assertEqual(context['assessment_text'], '')
- def test_missing_malformed_and_mismatched_reports(self):
- self.assessment_path.unlink()
- context = self.select({'datasets': ['人员信息'], 'relations': []})
- self.assertEqual(context['assessment_text'], '')
- self.assertTrue(context['issues'])
- self.write()
- self.payload['datasets'][0]['fields'] = None
- self.write()
- context = self.select({'datasets': ['人员信息'], 'relations': []})
- self.assertEqual(context['loaded_datasets'], [])
- self.payload['raw_values_sent_to_llm'] = True
- self.write()
- context = self.select({'datasets': ['项目信息'], 'relations': []})
- self.assertEqual(context['loaded_datasets'], [])
- def test_legacy_llm_reply_uses_current_question_fallback(self):
- context = qc.select_assessments(qc.load_question_context(), None,
- question='查询人员信息', slots={})
- self.assertEqual(context['loaded_datasets'], ['人员信息'])
- self.assertEqual(context['selection_source'], 'lexical_fallback')
- def test_missing_schema_stops_before_llm(self):
- self.schema_path.unlink()
- with patch.object(nodes, 'chat_json') as llm:
- with self.assertRaisesRegex(ValueError, '元知识图谱'):
- nodes.understand({'question': '查询'})
- llm.assert_not_called()
- def test_old_graph_answers_not_reused(self):
- state = {"question": "继续统计", "qa_context": {"build_id": "new"},
- "rounds": [{"question": "查询", "answer": "旧数据", "schema_build_id": "old"}]}
- with patch.object(nodes, "chat_json") as llm:
- result = nodes.reuse_check(state)
- self.assertTrue(result["need_full_query"])
- llm.assert_not_called()
- def test_chat_receives_schema_without_prior_assessment(self):
- with patch.object(nodes, 'chat_json', return_value={'category': '闲聊', 'slots': None}):
- state = nodes.understand({'question': '你能做什么', 'qa_context': {'assessment_text': 'OLD'}})
- captured = []
- with patch.object(nodes, 'chat_text', side_effect=lambda system, user, **kwargs: captured.append(system) or 'OK'):
- nodes.chat({**state, 'question': '你能做什么'})
- self.assertIn('service-rule', captured[0])
- self.assertNotIn('ASSESSMENT-BODY', captured[0])
- self.assertNotIn('OLD', captured[0])
- if __name__ == '__main__':
- unittest.main(verbosity=2)
|