Files
zk-data-agent/tests/test_data_agent_inputs.py
2026-05-08 17:09:07 +08:00

154 lines
6.1 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
from __future__ import annotations
import json
import tempfile
import unittest
import zipfile
from pathlib import Path
import openpyxl
from src.agent_tools import build_tool_context, default_tool_registry, execute_tool
from src.agent_types import AgentRuntimeConfig
from src.data_agent_inputs import (
extract_case_evidence,
load_input_sources,
render_source_context,
)
class DataAgentInputTests(unittest.TestCase):
def test_load_input_sources_reads_xlsx_tables(self) -> None:
with tempfile.TemporaryDirectory() as tmp_dir:
path = Path(tmp_dir) / 'cases.xlsx'
_write_xlsx(path)
payload = load_input_sources(tmp_dir, ['cases.xlsx'])
self.assertEqual(payload['source_count'], 1)
table = payload['sources'][0]['tables'][0]
self.assertEqual(table['title'], 'Sheet')
self.assertEqual(table['rows'][0], ['query', '预期domain', '0106-prev-domain', 'type', '备注'])
def test_load_input_sources_reads_explicit_external_file(self) -> None:
with tempfile.TemporaryDirectory() as root_dir, tempfile.TemporaryDirectory() as external_dir:
path = Path(external_dir) / 'cases.xlsx'
_write_xlsx(path)
payload = load_input_sources(root_dir, [str(path)])
self.assertEqual(payload['source_count'], 1)
self.assertEqual(payload['sources'][0]['path'], path.resolve().as_posix())
self.assertEqual(payload['sources'][0]['tables'][0]['rows'][1][0], '怎么开启查找设备')
def test_extract_case_evidence_profiles_and_extracts_rows(self) -> None:
with tempfile.TemporaryDirectory() as tmp_dir:
path = Path(tmp_dir) / 'cases.xlsx'
_write_xlsx(path)
payload = extract_case_evidence(tmp_dir, paths=['cases.xlsx'])
self.assertEqual(payload['required_questions'], [])
self.assertEqual(payload['evidence_count'], 2)
self.assertEqual(payload['evidence'][0]['query'], '怎么开启查找设备')
self.assertEqual(payload['evidence'][0]['expected_label'], 'productAgent')
self.assertEqual(payload['evidence'][0]['predicted_label'], 'QA')
def test_extract_case_evidence_asks_when_expected_label_is_missing(self) -> None:
with tempfile.TemporaryDirectory() as tmp_dir:
path = Path(tmp_dir) / 'ambiguous.xlsx'
workbook = openpyxl.Workbook()
sheet = workbook.active
sheet.append(['query', '备注'])
sheet.append(['怎么开启查找设备', '需要产品问答'])
workbook.save(path)
payload = extract_case_evidence(tmp_dir, paths=['ambiguous.xlsx'])
self.assertEqual(payload['evidence_count'], 0)
self.assertTrue(any('预期标签列' in question for question in payload['required_questions']))
def test_render_source_context_reads_docx_tables(self) -> None:
with tempfile.TemporaryDirectory() as tmp_dir:
path = Path(tmp_dir) / 'prd.docx'
_write_minimal_docx(path)
payload = render_source_context(tmp_dir, paths=['prd.docx'])
self.assertFalse(payload['truncated'])
self.assertIn('complex_task(tag="设备控制")', payload['context_text'])
self.assertIn('query示例: 我到家了布置下客厅灯光', payload['context_text'])
def test_tools_execute_against_registry(self) -> None:
with tempfile.TemporaryDirectory() as tmp_dir:
path = Path(tmp_dir) / 'cases.xlsx'
_write_xlsx(path)
context = build_tool_context(AgentRuntimeConfig(cwd=Path(tmp_dir)))
registry = default_tool_registry()
load_result = execute_tool(
registry,
'data_agent_load_input_sources',
{'paths': ['cases.xlsx'], 'max_rows_per_table': 5},
context,
)
self.assertTrue(load_result.ok, load_result.content)
loaded = json.loads(load_result.content)
evidence_result = execute_tool(
registry,
'data_agent_extract_case_evidence',
{'loaded_sources': loaded},
context,
)
render_result = execute_tool(
registry,
'data_agent_render_source_context',
{'loaded_sources': loaded, 'max_chars': 2000},
context,
)
self.assertTrue(evidence_result.ok, evidence_result.content)
evidence = json.loads(evidence_result.content)
self.assertEqual(evidence['evidence_count'], 2)
self.assertTrue(render_result.ok, render_result.content)
rendered = json.loads(render_result.content)
self.assertIn('怎么开启查找设备', rendered['context_text'])
def _write_xlsx(path: Path) -> None:
workbook = openpyxl.Workbook()
sheet = workbook.active
sheet.append(['query', '预期domain', '0106-prev-domain', 'type', '备注'])
sheet.append(['怎么开启查找设备', 'productAgent', 'QA', '设置问答', '预期落产品问答'])
sheet.append(['查找设备打开了吗', 'productAgent', 'smartApp', '状态问答', ''])
workbook.save(path)
def _write_minimal_docx(path: Path) -> None:
document_xml = '''<?xml version="1.0" encoding="UTF-8" standalone="yes"?>
<w:document xmlns:w="http://schemas.openxmlformats.org/wordprocessingml/2006/main">
<w:body>
<w:p><w:r><w:t>中控tagcomplex_task(tag="设备控制")</w:t></w:r></w:p>
<w:tbl>
<w:tr>
<w:tc><w:p><w:r><w:t>意图类别</w:t></w:r></w:p></w:tc>
<w:tc><w:p><w:r><w:t>query示例</w:t></w:r></w:p></w:tc>
<w:tc><w:p><w:r><w:t>高置信品类</w:t></w:r></w:p></w:tc>
</w:tr>
<w:tr>
<w:tc><w:p><w:r><w:t>设备控制</w:t></w:r></w:p></w:tc>
<w:tc><w:p><w:r><w:t>我到家了布置下客厅灯光</w:t></w:r></w:p></w:tc>
<w:tc><w:p><w:r><w:t>light</w:t></w:r></w:p></w:tc>
</w:tr>
</w:tbl>
</w:body>
</w:document>
'''
with zipfile.ZipFile(path, 'w') as archive:
archive.writestr('word/document.xml', document_xml)
if __name__ == '__main__':
unittest.main()