878 lines
37 KiB
Python
878 lines
37 KiB
Python
from __future__ import annotations
|
|
|
|
import csv
|
|
import io
|
|
import json
|
|
import tempfile
|
|
import unittest
|
|
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,
|
|
export_planning_eval_csv,
|
|
export_training_jsonl,
|
|
get_generation_plan,
|
|
normalize_dataset_draft,
|
|
prepare_generation_goal,
|
|
prepare_generation_plan,
|
|
render_dataset_records_table_csv,
|
|
render_planning_eval_csv,
|
|
update_generation_plan,
|
|
validate_dataset_records,
|
|
)
|
|
|
|
|
|
class DataAgentRecordTests(unittest.TestCase):
|
|
def test_normalize_dataset_draft_builds_canonical_records(self) -> None:
|
|
draft = '''
|
|
# dataset_label: 地图和生活边界数据
|
|
|
|
### case: 附近餐饮查询
|
|
用户: 帮我看看附近有什么好吃的
|
|
target: Agent(tag="life_service")
|
|
notes: 附近生活服务查询
|
|
|
|
### case: 多轮承接附近餐饮
|
|
用户: 我想出门逛逛
|
|
小爱: 好的
|
|
用户: 看看附近有什么好吃的
|
|
target: Agent(tag="life_service")
|
|
notes: 多轮承接附近生活服务查询
|
|
'''.strip()
|
|
|
|
payload = normalize_dataset_draft(
|
|
draft,
|
|
batch_id='demo',
|
|
base_timestamp=1_755_567_930_500,
|
|
timestamp_step_ms=60_000,
|
|
)
|
|
|
|
self.assertEqual(payload['warnings'], [])
|
|
records = payload['records']
|
|
self.assertEqual(len(records), 2)
|
|
self.assertEqual(records[0]['record_id'], 'gen_demo_000001')
|
|
self.assertEqual(records[0]['turn']['query'], '帮我看看附近有什么好吃的')
|
|
self.assertEqual(records[0]['prev_session'], [])
|
|
self.assertEqual(records[0]['label']['target_type'], 'agent')
|
|
self.assertEqual(records[0]['dimensions'], {'complex': False})
|
|
self.assertEqual(records[1]['turn']['query'], '看看附近有什么好吃的')
|
|
self.assertEqual(
|
|
records[1]['prev_session'],
|
|
[
|
|
{
|
|
'query': '我想出门逛逛',
|
|
'tts': '好的',
|
|
'timestamp': 1_755_569_070_500,
|
|
}
|
|
],
|
|
)
|
|
|
|
def test_normalize_dataset_draft_canonicalizes_single_quote_agent_target(self) -> None:
|
|
payload = normalize_dataset_draft(
|
|
'''
|
|
# dataset_label: 地图导航边界数据
|
|
### case: 顺路停车场
|
|
用户: 帮我找个顺路的停车场
|
|
target: Agent(tag='地图导航')
|
|
'''.strip(),
|
|
batch_id='demo',
|
|
base_timestamp=1_755_567_930_500,
|
|
)
|
|
|
|
record = payload['records'][0]
|
|
self.assertEqual(record['label']['target'], 'Agent(tag="地图导航")')
|
|
self.assertEqual(record['label']['target_type'], 'agent')
|
|
|
|
def test_normalize_dataset_draft_keeps_complex_dimension_independent(self) -> None:
|
|
payload = normalize_dataset_draft(
|
|
'''
|
|
# dataset_label: 地图导航复杂样例
|
|
### case: 复杂导航
|
|
用户: 帮我规划一条先去加油站再去公司的路线
|
|
complex: true
|
|
target: Agent(tag='地图导航')
|
|
'''.strip(),
|
|
batch_id='demo',
|
|
base_timestamp=1_755_567_930_500,
|
|
)
|
|
|
|
record = payload['records'][0]
|
|
self.assertEqual(record['label']['target'], 'Agent(tag="地图导航")')
|
|
self.assertEqual(record['dimensions']['complex'], True)
|
|
|
|
def test_validate_dataset_records_reports_valid_payload(self) -> None:
|
|
records = normalize_dataset_draft(
|
|
'''
|
|
# dataset_label: 地图和生活边界数据
|
|
### case: 附近餐饮查询
|
|
用户: 附近有什么好吃的
|
|
target: Agent(tag="life_service")
|
|
'''.strip(),
|
|
base_timestamp=1_755_567_930_500,
|
|
)['records']
|
|
|
|
result = validate_dataset_records(records)
|
|
|
|
self.assertTrue(result['ok'])
|
|
self.assertEqual(result['error_count'], 0)
|
|
|
|
def test_validate_dataset_records_rejects_current_turn_tts(self) -> None:
|
|
records = normalize_dataset_draft(
|
|
'''
|
|
# dataset_label: 地图和生活边界数据
|
|
### case: 附近餐饮查询
|
|
用户: 附近有什么好吃的
|
|
target: Agent(tag="life_service")
|
|
'''.strip(),
|
|
base_timestamp=1_755_567_930_500,
|
|
)['records']
|
|
records[0]['turn']['tts'] = '不应该出现'
|
|
|
|
result = validate_dataset_records(records)
|
|
|
|
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()
|
|
table_path = Path(tmp_dir) / payload['table_output_path']
|
|
table_rows = list(csv.DictReader(io.StringIO(table_path.read_text(encoding='utf-8-sig'))))
|
|
|
|
self.assertEqual(payload['output_format'], 'jsonl')
|
|
self.assertEqual(payload['table_output_format'], 'csv')
|
|
self.assertEqual(payload['record_count'], 1)
|
|
self.assertEqual(len(lines), 1)
|
|
self.assertNotIn(': ', lines[0])
|
|
self.assertEqual(json.loads(lines[0])['turn']['query'], '附近有什么好吃的')
|
|
self.assertEqual(len(table_rows), 1)
|
|
self.assertEqual(
|
|
list(table_rows[0].keys()),
|
|
['request_id', 'timestamp', 'query', 'prev_session', 'context', 'label', '是否迁移Function', 'function'],
|
|
)
|
|
self.assertEqual(table_rows[0]['request_id'], 'aabbccdd')
|
|
self.assertEqual(table_rows[0]['query'], '附近有什么好吃的')
|
|
self.assertEqual(table_rows[0]['prev_session'], '[]')
|
|
self.assertEqual(table_rows[0]['context'], '{}')
|
|
self.assertEqual(table_rows[0]['label'], '地图和生活边界数据')
|
|
self.assertEqual(table_rows[0]['是否迁移Function'], '')
|
|
self.assertEqual(table_rows[0]['function'], 'complex=false\nAgent(tag="life_service")')
|
|
|
|
def test_render_dataset_records_table_csv_keeps_prev_session_json(self) -> None:
|
|
records = normalize_dataset_draft(
|
|
'''
|
|
# dataset_label: 地图导航边界数据
|
|
### case: 多轮顺路停车场
|
|
用户: 查一下附近停车场
|
|
小爱: 找到了附近停车场
|
|
用户: 帮我找个最顺路的
|
|
target: Agent(tag='地图导航')
|
|
'''.strip(),
|
|
batch_id='demo',
|
|
base_timestamp=1_755_567_930_500,
|
|
timestamp_step_ms=60_000,
|
|
)['records']
|
|
|
|
rows = list(csv.DictReader(io.StringIO(render_dataset_records_table_csv(records))))
|
|
|
|
prev_session = json.loads(rows[0]['prev_session'])
|
|
self.assertEqual(prev_session[0]['query'], '查一下附近停车场')
|
|
self.assertEqual(prev_session[0]['tts'], '找到了附近停车场')
|
|
self.assertEqual(prev_session[0]['timestamp'], '1755567870500')
|
|
self.assertEqual(rows[0]['function'], 'complex=false\nAgent(tag="地图导航")')
|
|
|
|
def test_export_training_jsonl_uses_prompt_template_and_history(self) -> None:
|
|
records = normalize_dataset_draft(
|
|
'''
|
|
# dataset_label: 地图导航边界数据
|
|
### case: 多轮顺路停车场
|
|
用户: 查一下附近停车场
|
|
小爱: 找到了附近停车场
|
|
用户: 帮我找个顺路的
|
|
target: Agent(tag='地图导航')
|
|
'''.strip(),
|
|
batch_id='demo',
|
|
base_timestamp=1_755_567_930_500,
|
|
timestamp_step_ms=60_000,
|
|
)['records']
|
|
records[0]['context'] = {'location': '北京', 'rag': '地图服务可用', 'other': '忽略'}
|
|
with tempfile.TemporaryDirectory() as tmp_dir:
|
|
payload = export_training_jsonl(
|
|
records,
|
|
root=tmp_dir,
|
|
output_path='output/training.jsonl',
|
|
)
|
|
output_path = Path(tmp_dir) / payload['output_path']
|
|
line = json.loads(output_path.read_text(encoding='utf-8').splitlines()[0])
|
|
|
|
self.assertEqual(payload['output_format'], 'jsonl')
|
|
self.assertEqual(line['system'], '你是小爱同学,中文智能语音助手。')
|
|
self.assertEqual(line['output'], 'complex=false\nAgent(tag="地图导航")')
|
|
self.assertIn('[知识注入]\n{\n"location": "北京",\n"rag": "地图服务可用"\n}', line['instruction'])
|
|
self.assertIn('用户: 查一下附近停车场\n小爱: 找到了附近停车场', line['instruction'])
|
|
self.assertIn('[当前query]\n用户: 帮我找个顺路的', line['instruction'])
|
|
|
|
def test_render_planning_eval_csv_uses_expected_columns(self) -> None:
|
|
records = normalize_dataset_draft(
|
|
'''
|
|
# dataset_label: 地图导航边界数据
|
|
### case: 多轮顺路停车场
|
|
用户: 查一下附近停车场
|
|
小爱: 找到了附近停车场
|
|
用户: 帮我找个顺路的
|
|
target: Agent(tag='地图导航')
|
|
'''.strip(),
|
|
batch_id='demo',
|
|
base_timestamp=1_755_567_930_500,
|
|
)['records']
|
|
|
|
rows = list(csv.DictReader(io.StringIO(render_planning_eval_csv(records))))
|
|
|
|
self.assertEqual(
|
|
list(rows[0].keys()),
|
|
['request_id', 'newPrompt', 'query', '类别真实标签', 'code标签', 'complex'],
|
|
)
|
|
self.assertEqual(rows[0]['request_id'], 'aabbccdd')
|
|
self.assertEqual(rows[0]['query'], '帮我找个顺路的')
|
|
self.assertEqual(rows[0]['类别真实标签'], '地图导航')
|
|
self.assertEqual(rows[0]['code标签'], 'Agent(tag="地图导航")')
|
|
self.assertEqual(rows[0]['complex'], 'FALSE')
|
|
self.assertTrue(rows[0]['newPrompt'].startswith('<|im_start|>system\n你是小爱同学,中文智能语音助手。<|im_end|>'))
|
|
self.assertIn('<|im_start|>user\n请参考用户的[当前query]', rows[0]['newPrompt'])
|
|
self.assertIn('[function]', rows[0]['newPrompt'])
|
|
self.assertTrue(rows[0]['newPrompt'].endswith('<|im_start|>assistant\n'))
|
|
|
|
def test_export_planning_eval_csv_writes_file(self) -> None:
|
|
records = normalize_dataset_draft(
|
|
'''
|
|
# dataset_label: 时间工具数据
|
|
### case: 几点
|
|
用户: 现在几点
|
|
complex: true
|
|
target: CalendarQA(type="TIME")
|
|
'''.strip(),
|
|
batch_id='demo',
|
|
base_timestamp=1_755_567_930_500,
|
|
)['records']
|
|
with tempfile.TemporaryDirectory() as tmp_dir:
|
|
payload = export_planning_eval_csv(
|
|
records,
|
|
root=tmp_dir,
|
|
output_path='output/eval_planning.csv',
|
|
)
|
|
output_path = Path(tmp_dir) / payload['output_path']
|
|
rows = list(csv.DictReader(io.StringIO(output_path.read_text(encoding='utf-8-sig'))))
|
|
|
|
self.assertEqual(payload['output_format'], 'csv')
|
|
self.assertEqual(rows[0]['类别真实标签'], 'CalendarQA')
|
|
self.assertEqual(rows[0]['code标签'], 'CalendarQA(type="TIME")')
|
|
self.assertEqual(rows[0]['complex'], 'TRUE')
|
|
|
|
def test_normalize_online_draft_uses_real_request_metadata(self) -> None:
|
|
records = normalize_dataset_draft(
|
|
'''
|
|
# dataset_label: 线上误召回badcase专项
|
|
### case: 线上样例
|
|
request_id: rid-123
|
|
timestamp: 1755567930500
|
|
用户: 附近有什么好吃的
|
|
target: Agent(tag="life_service")
|
|
'''.strip(),
|
|
source_type='online',
|
|
)['records']
|
|
|
|
self.assertEqual(records[0]['record_id'], 'online_rid-123_000001')
|
|
self.assertEqual(records[0]['source']['request_id'], 'rid-123')
|
|
self.assertEqual(records[0]['source']['timestamp'], 1_755_567_930_500)
|
|
|
|
def test_normalize_requires_confirmed_plan_when_state_root_is_provided(self) -> None:
|
|
with tempfile.TemporaryDirectory() as tmp_dir:
|
|
with self.assertRaisesRegex(Exception, 'confirmed_plan_id is required'):
|
|
normalize_dataset_draft(
|
|
'''
|
|
# dataset_label: 地图和生活边界数据
|
|
### case: 附近餐饮查询
|
|
用户: 附近有什么好吃的
|
|
target: Agent(tag="life_service")
|
|
'''.strip(),
|
|
plan_state_root=tmp_dir,
|
|
)
|
|
|
|
def test_generation_plan_can_be_confirmed_then_used(self) -> None:
|
|
with tempfile.TemporaryDirectory() as tmp_dir:
|
|
plan_payload = prepare_generation_plan(
|
|
root=tmp_dir,
|
|
dataset_label='地图和生活边界数据',
|
|
target='Agent(tag="life_service")',
|
|
total_count=2,
|
|
turn_mix='1 条单轮,1 条多轮',
|
|
coverage='附近吃喝玩乐',
|
|
exclusions='不要生成导航路线类 query',
|
|
output_path='tasks/demo/artifacts',
|
|
)
|
|
plan_id = plan_payload['plan']['plan_id']
|
|
confirmed = confirm_generation_plan(root=tmp_dir, plan_id=plan_id, confirmation='确认,开始生成')
|
|
records = normalize_dataset_draft(
|
|
'''
|
|
# dataset_label: 地图和生活边界数据
|
|
### case: 附近餐饮查询
|
|
用户: 附近有什么好吃的
|
|
target: Agent(tag="life_service")
|
|
'''.strip(),
|
|
confirmed_plan_id=confirmed['confirmed_plan_id'],
|
|
plan_state_root=tmp_dir,
|
|
)['records']
|
|
|
|
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_goal_requires_declared_targets_before_review(self) -> None:
|
|
with tempfile.TemporaryDirectory() as tmp_dir:
|
|
with self.assertRaisesRegex(Exception, 'must include target or target_definitions'):
|
|
prepare_generation_goal(
|
|
root=tmp_dir,
|
|
dataset_label='地图和生活边界数据',
|
|
goal_summary='生成附近生活服务边界数据',
|
|
coverage='附近吃喝玩乐',
|
|
exclusions='不要生成导航路线类 query',
|
|
)
|
|
|
|
def test_generation_plan_inherits_targets_from_confirmed_goal(self) -> None:
|
|
with tempfile.TemporaryDirectory() as tmp_dir:
|
|
goal_payload = prepare_generation_goal(
|
|
root=tmp_dir,
|
|
dataset_label='餐饮和导航边界数据',
|
|
goal_summary='生成餐饮服务和地图导航边界数据',
|
|
target_definitions=[
|
|
{'name': '餐饮服务', 'target': 'Agent(tag="餐饮服务")', 'rule': '找附近餐饮'},
|
|
{'name': '地图导航', 'target': 'Agent(tag="地图导航")', 'rule': '明确导航'},
|
|
],
|
|
coverage='餐饮服务和地图导航边界',
|
|
exclusions='不要生成无关闲聊',
|
|
)
|
|
goal_id = goal_payload['goal']['goal_id']
|
|
confirm_generation_goal(root=tmp_dir, goal_id=goal_id, confirmation='确认目标')
|
|
|
|
plan_payload = prepare_generation_plan(
|
|
root=tmp_dir,
|
|
confirmed_goal_id=goal_id,
|
|
dataset_label='餐饮和导航边界数据',
|
|
target='',
|
|
total_count=2,
|
|
turn_mix='2 条单轮',
|
|
coverage='餐饮服务和地图导航边界',
|
|
exclusions='不要生成无关闲聊',
|
|
output_path='tasks/demo/artifacts',
|
|
)
|
|
|
|
self.assertEqual(
|
|
[item['target'] for item in plan_payload['plan']['target_definitions']],
|
|
['Agent(tag="餐饮服务")', 'Agent(tag="地图导航")'],
|
|
)
|
|
|
|
def test_generation_plan_dedupes_existing_task_output_path(self) -> None:
|
|
with tempfile.TemporaryDirectory() as tmp_dir:
|
|
existing = Path(tmp_dir) / 'tasks' / 'demo' / 'artifacts'
|
|
existing.mkdir(parents=True)
|
|
plan_payload = prepare_generation_plan(
|
|
root=tmp_dir,
|
|
dataset_label='地图和生活边界数据',
|
|
target='Agent(tag="life_service")',
|
|
total_count=2,
|
|
turn_mix='1 条单轮,1 条多轮',
|
|
coverage='附近吃喝玩乐',
|
|
exclusions='不要生成导航路线类 query',
|
|
output_path='tasks/demo/artifacts',
|
|
)
|
|
|
|
self.assertEqual(
|
|
plan_payload['plan']['output_path'],
|
|
'tasks/demo-data_plan_000001/artifacts',
|
|
)
|
|
self.assertEqual(plan_payload['plan']['requested_output_path'], 'tasks/demo/artifacts')
|
|
|
|
def test_generation_plan_review_update_and_revision_confirmation(self) -> None:
|
|
with tempfile.TemporaryDirectory() as tmp_dir:
|
|
plan_payload = prepare_generation_plan(
|
|
root=tmp_dir,
|
|
dataset_label='地图和生活边界数据',
|
|
target='Agent(tag="life_service")',
|
|
total_count=2,
|
|
turn_mix='1 条单轮,1 条多轮',
|
|
coverage='附近吃喝玩乐',
|
|
exclusions='不要生成导航路线类 query',
|
|
output_path='tasks/demo/artifacts',
|
|
)
|
|
plan_id = plan_payload['plan']['plan_id']
|
|
updated = update_generation_plan(
|
|
root=tmp_dir,
|
|
plan_id=plan_id,
|
|
review_feedback='多轮数据多一点',
|
|
updates={'turn_mix': '1 条单轮,3 条多轮', 'total_count': 4},
|
|
)
|
|
|
|
self.assertEqual(updated['plan']['revision'], 2)
|
|
self.assertEqual(updated['plan']['turn_mix'], '1 条单轮,3 条多轮')
|
|
self.assertEqual(len(updated['plan']['review_history']), 1)
|
|
shown = get_generation_plan(root=tmp_dir, plan_id=plan_id)
|
|
self.assertEqual(shown['plan']['revision'], 2)
|
|
with self.assertRaisesRegex(Exception, 'reviewed_revision must match'):
|
|
confirm_generation_plan(
|
|
root=tmp_dir,
|
|
plan_id=plan_id,
|
|
confirmation='确认,开始生成',
|
|
reviewed_revision=1,
|
|
)
|
|
confirmed = confirm_generation_plan(
|
|
root=tmp_dir,
|
|
plan_id=plan_id,
|
|
confirmation='确认,开始生成',
|
|
reviewed_revision=2,
|
|
)
|
|
|
|
self.assertEqual(confirmed['plan']['status'], 'confirmed')
|
|
self.assertEqual(confirmed['plan']['confirmed_revision'], 2)
|
|
|
|
def test_multi_target_plan_allows_only_declared_targets(self) -> None:
|
|
with tempfile.TemporaryDirectory() as tmp_dir:
|
|
plan_payload = prepare_generation_plan(
|
|
root=tmp_dir,
|
|
dataset_label='餐饮和导航边界数据',
|
|
target='',
|
|
target_definitions=[
|
|
{
|
|
'name': '餐饮服务',
|
|
'target': 'Agent(tag="餐饮服务")',
|
|
'rule': '找附近美食但不导航',
|
|
},
|
|
{
|
|
'name': '地图导航',
|
|
'target': 'Agent(tag="地图导航")',
|
|
'rule': '明确要求导航去某地',
|
|
},
|
|
],
|
|
total_count=2,
|
|
turn_mix='2 条单轮',
|
|
coverage='餐饮服务和地图导航边界',
|
|
exclusions='不要生成无关闲聊',
|
|
output_path='tasks/demo/artifacts',
|
|
)
|
|
plan_id = plan_payload['plan']['plan_id']
|
|
confirm_generation_plan(root=tmp_dir, plan_id=plan_id, confirmation='确认,开始生成')
|
|
records = normalize_dataset_draft(
|
|
'''
|
|
# dataset_label: 餐饮和导航边界数据
|
|
### case: 找奶茶
|
|
用户: 附近有没有奶茶店
|
|
target: Agent(tag="餐饮服务")
|
|
### case: 导航去奶茶店
|
|
用户: 导航去最近的奶茶店
|
|
target: Agent(tag="地图导航")
|
|
'''.strip(),
|
|
confirmed_plan_id=plan_id,
|
|
plan_state_root=tmp_dir,
|
|
)['records']
|
|
|
|
self.assertEqual(records[0]['label']['target'], 'Agent(tag="餐饮服务")')
|
|
self.assertEqual(records[1]['label']['target'], 'Agent(tag="地图导航")')
|
|
|
|
def test_multi_target_plan_rejects_undeclared_targets(self) -> None:
|
|
with tempfile.TemporaryDirectory() as tmp_dir:
|
|
plan_payload = prepare_generation_plan(
|
|
root=tmp_dir,
|
|
dataset_label='餐饮和导航边界数据',
|
|
target='',
|
|
target_definitions=[
|
|
{'target': 'Agent(tag="餐饮服务")'},
|
|
{'target': 'Agent(tag="地图导航")'},
|
|
],
|
|
total_count=1,
|
|
turn_mix='1 条单轮',
|
|
coverage='餐饮服务和地图导航边界',
|
|
exclusions='不要生成无关闲聊',
|
|
output_path='tasks/demo/artifacts',
|
|
)
|
|
plan_id = plan_payload['plan']['plan_id']
|
|
confirm_generation_plan(root=tmp_dir, plan_id=plan_id, confirmation='确认,开始生成')
|
|
with self.assertRaisesRegex(Exception, 'record targets must match'):
|
|
normalize_dataset_draft(
|
|
'''
|
|
# dataset_label: 餐饮和导航边界数据
|
|
### case: 天气
|
|
用户: 明天天气怎么样
|
|
target: Agent(tag="天气")
|
|
'''.strip(),
|
|
confirmed_plan_id=plan_id,
|
|
plan_state_root=tmp_dir,
|
|
)
|
|
|
|
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,
|
|
'turn_mix': '1 条单轮',
|
|
'coverage': '附近吃喝玩乐',
|
|
'exclusions': '不要生成导航路线类 query',
|
|
'output_path': 'tasks/demo/artifacts',
|
|
},
|
|
context,
|
|
)
|
|
self.assertTrue(plan_result.ok)
|
|
plan_id = json.loads(plan_result.content)['plan']['plan_id']
|
|
update_result = execute_tool(
|
|
default_tool_registry(),
|
|
'data_agent_update_generation_plan',
|
|
{
|
|
'plan_id': plan_id,
|
|
'review_feedback': '多轮不需要,先只测单轮',
|
|
'updates': {'turn_mix': '1 条单轮'},
|
|
},
|
|
context,
|
|
)
|
|
self.assertTrue(update_result.ok)
|
|
revision = json.loads(update_result.content)['plan']['revision']
|
|
confirm_result = execute_tool(
|
|
default_tool_registry(),
|
|
'data_agent_confirm_generation_plan',
|
|
{'plan_id': plan_id, 'confirmation': '确认,开始生成', 'reviewed_revision': revision},
|
|
context,
|
|
)
|
|
self.assertTrue(confirm_result.ok)
|
|
normalize_result = execute_tool(
|
|
default_tool_registry(),
|
|
'data_agent_normalize_dataset_draft',
|
|
{
|
|
'draft_text': '''
|
|
# dataset_label: 地图和生活边界数据
|
|
### case: 附近餐饮查询
|
|
用户: 附近有什么好吃的
|
|
target: Agent(tag="life_service")
|
|
'''.strip(),
|
|
'batch_id': 'demo',
|
|
'base_timestamp': 1_755_567_930_500,
|
|
'confirmed_plan_id': plan_id,
|
|
},
|
|
context,
|
|
)
|
|
self.assertTrue(normalize_result.ok)
|
|
records = json.loads(normalize_result.content)['records']
|
|
|
|
validate_result = execute_tool(
|
|
default_tool_registry(),
|
|
'data_agent_validate_dataset_records',
|
|
{'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)
|
|
|
|
def test_data_agent_tools_accept_draft_and_records_paths(self) -> None:
|
|
with tempfile.TemporaryDirectory() as tmp_dir:
|
|
root = Path(tmp_dir)
|
|
context = build_tool_context(AgentRuntimeConfig(cwd=root))
|
|
draft_path = root / 'draft.txt'
|
|
draft_path.write_text(
|
|
'''
|
|
# dataset_label: 地图导航边界数据
|
|
### case: 顺路停车场
|
|
用户: 帮我找个顺路的停车场
|
|
target: Agent(tag='地图导航')
|
|
'''.strip(),
|
|
encoding='utf-8',
|
|
)
|
|
|
|
normalize_result = execute_tool(
|
|
default_tool_registry(),
|
|
'data_agent_normalize_dataset_draft',
|
|
{
|
|
'draft_path': str(draft_path),
|
|
'batch_id': 'demo',
|
|
'source_type': 'manual',
|
|
'base_timestamp': 1_755_567_930_500,
|
|
},
|
|
context,
|
|
)
|
|
self.assertTrue(normalize_result.ok, normalize_result.content)
|
|
records = json.loads(normalize_result.content)['records']
|
|
self.assertEqual(records[0]['label']['target'], 'Agent(tag="地图导航")')
|
|
|
|
records_path = root / 'records.jsonl'
|
|
records_path.write_text(
|
|
'\n'.join(json.dumps(record, ensure_ascii=False) for record in records),
|
|
encoding='utf-8',
|
|
)
|
|
validate_result = execute_tool(
|
|
default_tool_registry(),
|
|
'data_agent_validate_dataset_records',
|
|
{'records_path': str(records_path)},
|
|
context,
|
|
)
|
|
export_result = execute_tool(
|
|
default_tool_registry(),
|
|
'data_agent_export_dataset_records',
|
|
{
|
|
'records_path': str(records_path),
|
|
'output_path': 'output/records.jsonl',
|
|
},
|
|
context,
|
|
)
|
|
self.assertTrue(export_result.ok, export_result.content)
|
|
export_payload = json.loads(export_result.content)
|
|
table_exists = (root / export_payload['table_output_path']).exists()
|
|
training_result = execute_tool(
|
|
default_tool_registry(),
|
|
'data_agent_export_training_jsonl',
|
|
{'records_path': str(records_path)},
|
|
context,
|
|
)
|
|
self.assertTrue(training_result.ok, training_result.content)
|
|
training_payload = json.loads(training_result.content)
|
|
training_exists = (root / training_payload['output_path']).exists()
|
|
eval_result = execute_tool(
|
|
default_tool_registry(),
|
|
'data_agent_export_planning_eval_csv',
|
|
{'records_path': str(records_path)},
|
|
context,
|
|
)
|
|
self.assertTrue(eval_result.ok, eval_result.content)
|
|
eval_payload = json.loads(eval_result.content)
|
|
eval_exists = (root / eval_payload['output_path']).exists()
|
|
|
|
self.assertTrue(validate_result.ok, validate_result.content)
|
|
self.assertTrue(json.loads(validate_result.content)['ok'])
|
|
self.assertEqual(export_payload['record_count'], 1)
|
|
self.assertTrue(table_exists)
|
|
self.assertEqual(training_payload['output_path'], '.port_sessions/data_agent_output/training.jsonl')
|
|
self.assertTrue(training_exists)
|
|
self.assertEqual(eval_payload['output_path'], '.port_sessions/data_agent_output/eval_planning.csv')
|
|
self.assertTrue(eval_exists)
|
|
|
|
def test_tool_schemas_avoid_top_level_composition_keywords(self) -> None:
|
|
# 部分 Bedrock/Anthropic 兼容后端不接受顶层 anyOf/oneOf/allOf。
|
|
blocked_keywords = {'anyOf', 'oneOf', 'allOf'}
|
|
violations = [
|
|
name
|
|
for name, tool in default_tool_registry().items()
|
|
if blocked_keywords.intersection(tool.parameters)
|
|
]
|
|
|
|
self.assertEqual(violations, [])
|
|
|
|
def test_data_agent_tools_route_relative_outputs_to_session_output(self) -> None:
|
|
with tempfile.TemporaryDirectory() as tmp_dir:
|
|
root = Path(tmp_dir)
|
|
scratchpad = root / '.port_sessions' / 'accounts' / 'alice' / 'sessions' / 's1' / 'scratchpad'
|
|
scratchpad.mkdir(parents=True)
|
|
context = build_tool_context(AgentRuntimeConfig(cwd=root), scratchpad_directory=scratchpad)
|
|
registry = default_tool_registry()
|
|
|
|
plan_result = execute_tool(
|
|
registry,
|
|
'data_agent_prepare_generation_plan',
|
|
{
|
|
'direct_review': True,
|
|
'dataset_label': '地图和生活边界数据',
|
|
'target': 'Agent(tag="life_service")',
|
|
'total_count': 1,
|
|
'turn_mix': '1 条单轮',
|
|
'coverage': '附近吃喝玩乐',
|
|
'exclusions': '不要生成导航路线类 query',
|
|
'output_path': 'output/demo/records.jsonl',
|
|
},
|
|
context,
|
|
)
|
|
self.assertTrue(plan_result.ok, plan_result.content)
|
|
plan_payload = json.loads(plan_result.content)
|
|
|
|
records = normalize_dataset_draft(
|
|
'''
|
|
# dataset_label: 地图和生活边界数据
|
|
### case: 附近餐饮查询
|
|
用户: 附近有什么好吃的
|
|
target: Agent(tag="life_service")
|
|
'''.strip(),
|
|
batch_id='demo',
|
|
base_timestamp=1_755_567_930_500,
|
|
)['records']
|
|
export_result = execute_tool(
|
|
registry,
|
|
'data_agent_export_dataset_records',
|
|
{'records': records, 'output_path': 'output/地图和生活边界数据/自定义名字.jsonl'},
|
|
context,
|
|
)
|
|
self.assertTrue(export_result.ok, export_result.content)
|
|
export_payload = json.loads(export_result.content)
|
|
|
|
expected = '.port_sessions/accounts/alice/sessions/s1/output/records.jsonl'
|
|
expected_table = '.port_sessions/accounts/alice/sessions/s1/output/records.csv'
|
|
self.assertEqual(plan_payload['plan']['output_path'], expected)
|
|
self.assertEqual(export_payload['output_path'], expected)
|
|
self.assertEqual(export_payload['table_output_path'], expected_table)
|
|
|
|
def test_data_agent_session_records_path_stays_stable_when_file_exists(self) -> None:
|
|
with tempfile.TemporaryDirectory() as tmp_dir:
|
|
root = Path(tmp_dir)
|
|
scratchpad = root / '.port_sessions' / 'accounts' / 'alice' / 'sessions' / 's1' / 'scratchpad'
|
|
output_root = scratchpad.parent / 'output'
|
|
output_root.mkdir(parents=True)
|
|
(output_root / 'records.jsonl').write_text('{}\n', encoding='utf-8')
|
|
context = build_tool_context(AgentRuntimeConfig(cwd=root), scratchpad_directory=scratchpad)
|
|
|
|
plan_result = execute_tool(
|
|
default_tool_registry(),
|
|
'data_agent_prepare_generation_plan',
|
|
{
|
|
'direct_review': True,
|
|
'dataset_label': '地图和生活边界数据',
|
|
'target': 'Agent(tag="life_service")',
|
|
'total_count': 1,
|
|
'turn_mix': '1 条单轮',
|
|
'coverage': '附近吃喝玩乐',
|
|
'exclusions': '不要生成导航路线类 query',
|
|
'output_path': 'output/任意子目录/任意名字.jsonl',
|
|
},
|
|
context,
|
|
)
|
|
self.assertTrue(plan_result.ok, plan_result.content)
|
|
plan_payload = json.loads(plan_result.content)
|
|
|
|
self.assertEqual(
|
|
plan_payload['plan']['output_path'],
|
|
'.port_sessions/accounts/alice/sessions/s1/output/records.jsonl',
|
|
)
|
|
|
|
|
|
if __name__ == '__main__':
|
|
unittest.main()
|