ai4mats-mcp-tools/test/test_dataset.py

148 lines
5.2 KiB
Python
Raw 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.

import asyncio
import json
import sys
import os
import time
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from tools.dataset_mcp_server import (
login,
add_dataset,
get_asset_icon,
get_data_types,
upload_chunk,
add_version,
update_dataset,
update_desc,
publish_dataset,
download_all_files,
download_single_file,
delete_version,
delete_dataset,
query_datasets,
get_dataset_detail
)
async def run_test(test_name, func, params, retries=1):
"""运行单个测试用例"""
print(f"\n{'='*60}")
print(f"🔧 测试: {test_name}")
print(f"{'='*60}")
print(f"参数: {json.dumps(params, indent=2, ensure_ascii=False)}")
try:
result = await func(**params)
print(f"\n✅ 成功!")
print(f"返回结果: {result[:500]}{'...' if len(result) > 500 else ''}")
return result, True
except Exception as e:
if retries > 0:
print(f"\n⚠️ 失败,重试中 ({retries}次剩余): {e}")
return await run_test(test_name, func, params, retries - 1)
print(f"\n❌ 失败: {e}")
return None, False
async def main():
print("🚀 开始测试数据集 MCP 服务...")
try:
with open("../dataset/test_dataset.json", "r", encoding="utf-8") as f:
test_data = json.load(f)
except FileNotFoundError:
print("❌ 错误:找不到 test_dataset.json 文件")
return
except json.JSONDecodeError as e:
print(f"❌ 错误test_dataset.json 格式不正确: {e}")
return
token = None
login_result, success = await run_test("登录获取Token", login, {})
if success:
try:
login_data = json.loads(login_result)
token = login_data.get("data", {}).get("access_token")
print(f"\n📋 获取到Token: {token[:20]}...")
except:
print("❌ 解析Token失败")
return
else:
print("❌ 登录失败,无法继续测试")
return
test_cases = [
("查询数据集分类", get_asset_icon, {"token": token, **test_data["get_asset_icon"]}),
("获取可用数据类型", get_data_types, {"token": token}),
("查询数据集列表", query_datasets, {"token": token, **test_data["query_datasets"]}),
]
passed = 1
failed = 0
for test_name, func, params in test_cases:
_, success = await run_test(test_name, func, params)
if success:
passed += 1
else:
failed += 1
new_dataset_params = {"token": token, **test_data["add_dataset"]}
dataset_name = test_data["add_dataset"]["name"]
add_result, success = await run_test("新增数据集", add_dataset, new_dataset_params)
while not success or (add_result and "项目名称已被使用" in add_result):
if add_result and "项目名称已被使用" in add_result:
timestamp = int(time.time())
new_name = f"{dataset_name}_{timestamp}"
print(f"\n⚠️ 项目名称已被使用,尝试新名称: {new_name}")
new_dataset_params["name"] = new_name
add_result, success = await run_test("新增数据集", add_dataset, new_dataset_params)
else:
break
if success:
passed += 1
try:
add_result_data = json.loads(add_result)
dataset_id = add_result_data.get("data", {}).get("id", test_data["add_version"]["id"])
print(f"\n📋 创建的数据集ID: {dataset_id}")
except:
dataset_id = test_data["add_version"]["id"]
else:
failed += 1
dataset_id = test_data["add_version"]["id"]
remaining_tests = [
("上传文件分片", upload_chunk, {"token": token, **test_data["upload_chunk"]}),
("新增版本", add_version, {"token": token, **{"id": dataset_id, **test_data["add_version"]}}),
("修改数据集", update_dataset, {"token": token, **{"id": dataset_id, **test_data["update_dataset"]}}),
("编辑数据集简介", update_desc, {"token": token, **test_data["update_desc"]}),
("发布数据集", publish_dataset, {"token": token, **{"id": dataset_id, **test_data["publish_dataset"]}}),
("下载全部文件", download_all_files, {"token": token, **test_data["download_all_files"]}),
("下载单个文件", download_single_file, {"token": token, **test_data["download_single_file"]}),
("删除版本", delete_version, {"token": token, **test_data["delete_version"]}),
("删除数据集", delete_dataset, {"token": token, **{"id": dataset_id}}),
("查询数据集详情", get_dataset_detail, {"token": token, **test_data["get_dataset_detail"]})
]
for test_name, func, params in remaining_tests:
_, success = await run_test(test_name, func, params)
if success:
passed += 1
else:
failed += 1
print(f"\n{'='*60}")
print("📊 测试结果汇总")
print(f"{'='*60}")
print(f"通过: {passed}")
print(f"失败: {failed}")
print(f"成功率: {passed / (passed + failed) * 100:.1f}%")
if __name__ == "__main__":
asyncio.run(main())