"""
基于耗时数据的用例拆分模块。
根据 time_data.json 中的用例文件执行耗时数据和类级执行耗时数据,
将用例拆分成 world_size 份,目标是使拆分后的每一个机器上的总耗时尽量相等。
"""
import ast
import json
import os
from pathlib import Path
from access_control import TEST_DIR, NETWORK_OPS_DIR
DEFAULT_FILE_TIME = 10.0
DEFAULT_CLASS_TIME = 60.0
SLOW_FILE_THRESHOLD = 300.0
def get_op_name(ut_file):
"""从用例文件路径获取操作名(与 exec_ut 中逻辑一致)"""
op_name = str(Path(ut_file).name).split('.')[0]
return op_name[5:] if op_name.startswith("test_") else op_name
def get_test_key(ut_file, ut_type):
"""
生成与 time_data.json 中 timedata 键匹配的标识。
与 exec_ut 中 ut_info 的生成逻辑保持一致:
- op_ut_files: "test_ops _" + op_name
- ut_files (op-plugin): 相对 NETWORK_OPS_DIR 的路径(去 .py)
- ut_files (其他): 相对 TEST_DIR 的路径(去 .py)
"""
if ut_type == "op_ut_files":
return "test_ops _" + get_op_name(ut_file)
if 'op-plugin' in str(Path(ut_file)):
return Path(ut_file).relative_to(NETWORK_OPS_DIR).as_posix()[:-3]
return Path(ut_file).relative_to(TEST_DIR).as_posix()[:-3]
def load_and_validate_time_data(file_path):
"""
加载并验证 time_data.json 文件。
验证条件:
1. 文件存在
2. 内容为有效 JSON 且为字典
3. 包含 timedata 字段且为字典
4. timedata 中每项为字典且含 total_time 数值
5. 若存在 classes 字段,则必须为字典且每个 class_time 为数值
Returns:
timedata 字典(成功)或 None(失败,触发回退)
"""
if not os.path.exists(file_path):
return None
try:
with open(file_path, 'r', encoding='utf-8') as f:
content = json.loads(f.read())
if not isinstance(content, dict):
return None
timedata = content.get('timedata')
if not isinstance(timedata, dict):
return None
for key, value in timedata.items():
if not isinstance(value, dict):
return None
total_time = value.get('total_time')
if total_time is None or not isinstance(total_time, (int, float)):
return None
classes = value.get('classes')
if classes is not None:
if not isinstance(classes, dict):
return None
for class_time in classes.values():
if not isinstance(class_time, (int, float)):
return None
return timedata
except (json.JSONDecodeError, OSError, ValueError):
return None
def discover_test_classes(file_path):
"""
用 AST 解析发现测试文件中以 Test 开头的类名列表。
用于慢文件(>= 300秒)无类级耗时数据时,发现类名进行类级拆分。
"""
try:
with open(file_path, 'r', encoding='utf-8') as f:
tree = ast.parse(f.read())
classes = []
for node in ast.walk(tree):
if isinstance(node, ast.ClassDef) and node.name.startswith('Test'):
classes.append(node.name)
return classes
except Exception:
return []
def build_split_units(test_files, timedata):
"""
构建拆分单元列表。
拆分规则:
- op_ut_files 始终按文件级拆分(通过 -k 过滤运行,类级拆分不适用)
- 无耗时数据 → 文件级单元,时间=10秒
- total_time < 300 → 文件级单元,时间=total_time
- total_time >= 300 且有 classes → 类级单元,每个类单独一个单元
- total_time >= 300 但无 classes → AST发现类,类级单元,每个类时间=60秒
Returns:
list of (ut_type, ut_file, class_name_or_None, estimated_time)
"""
units = []
for ut_type, ut_files in test_files.items():
for ut_file in ut_files:
test_key = get_test_key(ut_file, ut_type)
time_data = timedata.get(test_key)
if ut_type == "op_ut_files":
if time_data:
est_time = time_data.get('total_time', DEFAULT_FILE_TIME)
else:
est_time = DEFAULT_FILE_TIME
units.append((ut_type, ut_file, None, est_time))
continue
if time_data is None:
units.append((ut_type, ut_file, None, DEFAULT_FILE_TIME))
else:
total_time = time_data.get('total_time', DEFAULT_FILE_TIME)
classes = time_data.get('classes')
if total_time < SLOW_FILE_THRESHOLD:
units.append((ut_type, ut_file, None, total_time))
elif classes and isinstance(classes, dict):
for class_name, class_time in classes.items():
if not isinstance(class_time, (int, float)):
continue
units.append((ut_type, ut_file, class_name, class_time))
else:
discovered = discover_test_classes(ut_file)
if discovered:
for class_name in discovered:
units.append((ut_type, ut_file, class_name, DEFAULT_CLASS_TIME))
else:
units.append((ut_type, ut_file, None, total_time))
return units
def split_by_time(test_files, timedata, rank, world_size):
"""
基于耗时数据拆分用例,使每个机器上的总耗时尽量相等。
使用 LPT(Longest Processing Time first)贪心算法:
1. 将所有拆分单元按预估时间降序排序
2. 每个单元分配给当前总时间最小的机器
Args:
test_files: {'ut_files': [...], 'op_ut_files': [...]}
timedata: time_data.json 中的 timedata 字典
rank: 当前机器的序号(从1开始)
world_size: 机器总数
Returns:
{
'ut_files': [file_path, ...],
'op_ut_files': [file_path, ...],
'ut_classes': {file_path: [class_name, ...]},
'op_ut_classes': {file_path: [class_name, ...]}
}
"""
if rank > world_size:
raise Exception(f'rank {rank} is greater than world_size {world_size}')
units = build_split_units(test_files, timedata)
units.sort(key=lambda x: x[3], reverse=True)
machine_loads = [0.0] * world_size
machine_units = [[] for _ in range(world_size)]
for unit in units:
min_machine = machine_loads.index(min(machine_loads))
machine_units[min_machine].append(unit)
machine_loads[min_machine] += unit[3]
print("***** Time-based split result:")
for i, (load, mus) in enumerate(zip(machine_loads, machine_units)):
print(f" Machine {i + 1}: estimated_time={load:.1f}s, units={len(mus)}")
my_units = machine_units[rank - 1]
result = {
'ut_files': [],
'op_ut_files': [],
'ut_classes': {},
'op_ut_classes': {}
}
for ut_type, ut_file, class_name, _ in my_units:
if class_name is None:
if ut_file not in result[ut_type]:
result[ut_type].append(ut_file)
else:
if ut_file not in result[ut_type]:
result[ut_type].append(ut_file)
classes_key = 'ut_classes' if ut_type == 'ut_files' else 'op_ut_classes'
if ut_file not in result[classes_key]:
result[classes_key][ut_file] = []
result[classes_key][ut_file].append(class_name)
return result