"""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)