Add data agent input and export workflow
This commit is contained in:
@@ -0,0 +1,142 @@
|
||||
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()
|
||||
@@ -8,9 +8,12 @@ from pathlib import Path
|
||||
from src.agent_tools import build_tool_context, default_tool_registry, execute_tool
|
||||
from src.agent_types import AgentRuntimeConfig
|
||||
from src.data_agent_records import (
|
||||
confirm_generation_goal,
|
||||
confirm_generation_plan,
|
||||
export_dataset_records,
|
||||
get_generation_plan,
|
||||
normalize_dataset_draft,
|
||||
prepare_generation_goal,
|
||||
prepare_generation_plan,
|
||||
update_generation_plan,
|
||||
validate_dataset_records,
|
||||
@@ -94,6 +97,32 @@ target: Agent(tag="life_service")
|
||||
self.assertFalse(result['ok'])
|
||||
self.assertEqual(result['errors'][0]['path'], 'records[0].turn.tts')
|
||||
|
||||
def test_export_dataset_records_writes_compact_jsonl(self) -> None:
|
||||
records = normalize_dataset_draft(
|
||||
'''
|
||||
# dataset_label: 地图和生活边界数据
|
||||
### case: 附近餐饮查询
|
||||
用户: 附近有什么好吃的
|
||||
target: Agent(tag="life_service")
|
||||
'''.strip(),
|
||||
batch_id='demo',
|
||||
base_timestamp=1_755_567_930_500,
|
||||
)['records']
|
||||
with tempfile.TemporaryDirectory() as tmp_dir:
|
||||
payload = export_dataset_records(
|
||||
records,
|
||||
root=tmp_dir,
|
||||
output_path='tasks/demo/records.jsonl',
|
||||
)
|
||||
output_path = Path(tmp_dir) / payload['output_path']
|
||||
lines = output_path.read_text(encoding='utf-8').splitlines()
|
||||
|
||||
self.assertEqual(payload['output_format'], 'jsonl')
|
||||
self.assertEqual(payload['record_count'], 1)
|
||||
self.assertEqual(len(lines), 1)
|
||||
self.assertNotIn(': ', lines[0])
|
||||
self.assertEqual(json.loads(lines[0])['turn']['query'], '附近有什么好吃的')
|
||||
|
||||
def test_normalize_online_draft_uses_real_request_metadata(self) -> None:
|
||||
records = normalize_dataset_draft(
|
||||
'''
|
||||
@@ -151,6 +180,55 @@ target: Agent(tag="life_service")
|
||||
|
||||
self.assertEqual(records[0]['label']['dataset_label'], '地图和生活边界数据')
|
||||
|
||||
def test_generation_goal_can_gate_plan_creation(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as tmp_dir:
|
||||
goal_payload = prepare_generation_goal(
|
||||
root=tmp_dir,
|
||||
dataset_label='地图和生活边界数据',
|
||||
goal_summary='生成附近生活服务边界数据',
|
||||
target='Agent(tag="life_service")',
|
||||
plan_hint='建议先生成 2 条单轮,输出到 tasks/demo/records.jsonl',
|
||||
coverage='附近吃喝玩乐',
|
||||
exclusions='不要生成导航路线类 query',
|
||||
source_refs=['manual:user'],
|
||||
)
|
||||
goal_id = goal_payload['goal']['goal_id']
|
||||
with self.assertRaisesRegex(Exception, 'generation goal is not confirmed'):
|
||||
prepare_generation_plan(
|
||||
root=tmp_dir,
|
||||
confirmed_goal_id=goal_id,
|
||||
dataset_label='地图和生活边界数据',
|
||||
target='Agent(tag="life_service")',
|
||||
total_count=2,
|
||||
turn_mix='2 条单轮',
|
||||
coverage='附近吃喝玩乐',
|
||||
exclusions='不要生成导航路线类 query',
|
||||
output_path='tasks/demo/artifacts',
|
||||
)
|
||||
confirmed_goal = confirm_generation_goal(
|
||||
root=tmp_dir,
|
||||
goal_id=goal_id,
|
||||
confirmation='确认 goal',
|
||||
reviewed_revision=1,
|
||||
)
|
||||
plan_payload = prepare_generation_plan(
|
||||
root=tmp_dir,
|
||||
confirmed_goal_id=confirmed_goal['confirmed_goal_id'],
|
||||
dataset_label='地图和生活边界数据',
|
||||
target='Agent(tag="life_service")',
|
||||
total_count=2,
|
||||
turn_mix='2 条单轮',
|
||||
coverage='附近吃喝玩乐',
|
||||
exclusions='不要生成导航路线类 query',
|
||||
output_path='tasks/demo/artifacts',
|
||||
)
|
||||
|
||||
self.assertEqual(plan_payload['plan']['confirmed_goal_id'], goal_id)
|
||||
self.assertEqual(
|
||||
goal_payload['goal']['plan_hint'],
|
||||
'建议先生成 2 条单轮,输出到 tasks/demo/records.jsonl',
|
||||
)
|
||||
|
||||
def test_generation_plan_dedupes_existing_task_output_path(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as tmp_dir:
|
||||
existing = Path(tmp_dir) / 'tasks' / 'demo' / 'artifacts'
|
||||
@@ -290,10 +368,48 @@ target: Agent(tag="天气")
|
||||
def test_tools_execute_against_registry(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as tmp_dir:
|
||||
context = build_tool_context(AgentRuntimeConfig(cwd=Path(tmp_dir)))
|
||||
goal_result = execute_tool(
|
||||
default_tool_registry(),
|
||||
'data_agent_prepare_generation_goal',
|
||||
{
|
||||
'dataset_label': '地图和生活边界数据',
|
||||
'goal_summary': '生成附近生活服务边界数据',
|
||||
'target': 'Agent(tag="life_service")',
|
||||
'coverage': '附近吃喝玩乐',
|
||||
'exclusions': '不要生成导航路线类 query',
|
||||
'source_refs': ['manual:user'],
|
||||
},
|
||||
context,
|
||||
)
|
||||
self.assertTrue(goal_result.ok, goal_result.content)
|
||||
goal_id = json.loads(goal_result.content)['goal']['goal_id']
|
||||
blocked_plan_result = execute_tool(
|
||||
default_tool_registry(),
|
||||
'data_agent_prepare_generation_plan',
|
||||
{
|
||||
'dataset_label': '地图和生活边界数据',
|
||||
'target': 'Agent(tag="life_service")',
|
||||
'total_count': 1,
|
||||
'turn_mix': '1 条单轮',
|
||||
'coverage': '附近吃喝玩乐',
|
||||
'exclusions': '不要生成导航路线类 query',
|
||||
'output_path': 'tasks/demo/artifacts',
|
||||
},
|
||||
context,
|
||||
)
|
||||
self.assertFalse(blocked_plan_result.ok)
|
||||
confirmed_goal_result = execute_tool(
|
||||
default_tool_registry(),
|
||||
'data_agent_confirm_generation_goal',
|
||||
{'goal_id': goal_id, 'confirmation': '确认 goal', 'reviewed_revision': 1},
|
||||
context,
|
||||
)
|
||||
self.assertTrue(confirmed_goal_result.ok, confirmed_goal_result.content)
|
||||
plan_result = execute_tool(
|
||||
default_tool_registry(),
|
||||
'data_agent_prepare_generation_plan',
|
||||
{
|
||||
'confirmed_goal_id': goal_id,
|
||||
'dataset_label': '地图和生活边界数据',
|
||||
'target': 'Agent(tag="life_service")',
|
||||
'total_count': 1,
|
||||
@@ -350,9 +466,27 @@ target: Agent(tag="life_service")
|
||||
{'records': records},
|
||||
context,
|
||||
)
|
||||
export_result = execute_tool(
|
||||
default_tool_registry(),
|
||||
'data_agent_export_dataset_records',
|
||||
{
|
||||
'records': records,
|
||||
'output_path': 'tasks/demo/records.jsonl',
|
||||
},
|
||||
context,
|
||||
)
|
||||
export_payload = json.loads(export_result.content) if export_result.ok else {}
|
||||
exported_lines = (
|
||||
(Path(tmp_dir) / export_payload['output_path']).read_text(encoding='utf-8').splitlines()
|
||||
if export_result.ok
|
||||
else []
|
||||
)
|
||||
|
||||
self.assertTrue(validate_result.ok)
|
||||
self.assertTrue(json.loads(validate_result.content)['ok'])
|
||||
self.assertTrue(export_result.ok, export_result.content)
|
||||
self.assertEqual(export_payload['record_count'], 1)
|
||||
self.assertEqual(len(exported_lines), 1)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
|
||||
Reference in New Issue
Block a user