143 lines
5.5 KiB
Python
143 lines
5.5 KiB
Python
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_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>中控tag:complex_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()
|