257 lines
8.9 KiB
Dart
257 lines
8.9 KiB
Dart
import 'dart:convert';
|
||
import 'dart:io';
|
||
|
||
import 'package:crypto/crypto.dart' show sha256;
|
||
import 'package:flutter_test/flutter_test.dart';
|
||
import 'package:http/http.dart' as http;
|
||
import 'package:http/testing.dart';
|
||
import 'package:observer/models/model_manager.dart';
|
||
|
||
const _modelBytes = [1, 2, 3, 4, 5, 6, 7, 8];
|
||
|
||
String _shaHex(List<int> bytes) => sha256.convert(bytes).toString();
|
||
|
||
/// 分两段到达的下载流(中间 50ms 停顿,供取消测试在下载中触发)
|
||
Stream<List<int>> _delayedChunks() async* {
|
||
yield [1, 2, 3, 4];
|
||
await Future<void>.delayed(const Duration(milliseconds: 50));
|
||
yield [5, 6, 7, 8];
|
||
}
|
||
|
||
Map<String, dynamic> _catalog(List<Map<String, dynamic>> models) => {
|
||
'code': 0,
|
||
'message': 'ok',
|
||
'data': {'version': '0.0.2', 'notes': '', 'models': models}
|
||
};
|
||
|
||
Map<String, dynamic> _item({
|
||
int datasetId = 7,
|
||
String name = '环颈雉鸡数据集',
|
||
String version = 'v1.0.0',
|
||
String sha = '',
|
||
}) =>
|
||
{
|
||
'datasetId': datasetId,
|
||
'datasetName': name,
|
||
'version': version,
|
||
'labels': ['pheasant', 'suspect'],
|
||
'sizeBytes': _modelBytes.length,
|
||
'sha256': sha.isEmpty ? _shaHex(_modelBytes) : sha,
|
||
'downloadUrl': '/download/models/$datasetId/latest.tflite',
|
||
'coverUrl': '/api/v1/app/cover?namePrefix=RNPHE',
|
||
};
|
||
|
||
void main() {
|
||
late Directory root;
|
||
late int downloadHits;
|
||
|
||
setUp(() async {
|
||
root = await Directory.systemTemp.createTemp('model-manager-test');
|
||
downloadHits = 0;
|
||
});
|
||
|
||
tearDown(() => root.delete(recursive: true));
|
||
|
||
ModelManager manager(MockClient client) => ModelManager(
|
||
baseUrl: 'http://test.local',
|
||
client: client,
|
||
rootDir: () async => root,
|
||
);
|
||
|
||
MockClient client(List<Map<String, dynamic>> models) =>
|
||
MockClient((req) async {
|
||
if (req.url.path == '/api/v1/app/update') {
|
||
// Response(String) 默认 latin1 编码,中文数据集名会抛异常 → 用 bytes
|
||
return http.Response.bytes(
|
||
utf8.encode(jsonEncode(_catalog(models))), 200);
|
||
}
|
||
if (req.url.path.startsWith('/download/models/')) {
|
||
downloadHits++;
|
||
return http.Response.bytes(_modelBytes, 200);
|
||
}
|
||
return http.Response('not found', 404);
|
||
});
|
||
|
||
test('refresh 只拉目录:不下载任何模型,coverUrl 解析正确', () async {
|
||
final m = manager(client([_item()]));
|
||
await m.refresh();
|
||
|
||
expect(m.ready, isTrue);
|
||
expect(m.error, isNull);
|
||
expect(downloadHits, 0, reason: 'refresh 不应触发下载');
|
||
expect(m.models, isEmpty, reason: '未激活的模型不应出现在 models');
|
||
expect(m.catalog.length, 1);
|
||
expect(m.catalog.first.coverUrl, '/api/v1/app/cover?namePrefix=RNPHE');
|
||
});
|
||
|
||
test('downloadModel:下载+校验+落盘+自动激活+进度回调', () async {
|
||
final m = manager(client([_item()]));
|
||
await m.refresh();
|
||
|
||
final progresses = <double>[];
|
||
final ok = await m.downloadModel(m.catalog.first, onProgress: (r, t) {
|
||
progresses.add(t == 0 ? 0 : r / t);
|
||
});
|
||
|
||
expect(ok, isTrue);
|
||
expect(downloadHits, 1);
|
||
expect(m.isDownloaded(7), isTrue);
|
||
expect(m.isActive(7), isTrue, reason: '下载完成应自动使用');
|
||
expect(progresses.last, 1.0);
|
||
expect(m.models.length, 1);
|
||
expect(m.models.first.datasetName, '环颈雉鸡数据集');
|
||
|
||
final dir = Directory('${root.path}/7');
|
||
expect(await File('${dir.path}/model.tflite').exists(), isTrue);
|
||
expect(await File('${dir.path}/labels.json').exists(), isTrue);
|
||
expect(await File('${dir.path}/meta.json').exists(), isTrue);
|
||
});
|
||
|
||
test('已下载未激活:setActive 直接使用,不触发下载', () async {
|
||
final m = manager(client([_item()]));
|
||
await m.refresh();
|
||
await m.downloadModel(m.catalog.first);
|
||
await m.setActive(7, false);
|
||
expect(m.isActive(7), isFalse);
|
||
expect(m.models, isEmpty);
|
||
|
||
final before = downloadHits;
|
||
await m.setActive(7, true);
|
||
expect(m.isActive(7), isTrue);
|
||
expect(downloadHits, before, reason: '已下载直接使用不应重新下载');
|
||
expect(m.models.length, 1);
|
||
});
|
||
|
||
test('版本更新:downloadModel 重新下载', () async {
|
||
final m = manager(client([_item()]));
|
||
await m.refresh();
|
||
await m.downloadModel(m.catalog.first);
|
||
expect(downloadHits, 1);
|
||
|
||
final m2 = manager(client([_item(version: 'v2.0.0')]));
|
||
await m2.refresh();
|
||
await m2.downloadModel(m2.catalog.first);
|
||
expect(downloadHits, 2);
|
||
expect(m2.models.first.version, 'v2.0.0');
|
||
});
|
||
|
||
test('autoUpdate:已激活模型出新版本,refresh 后自动重下并立即生效', () async {
|
||
final m = manager(client([_item()]));
|
||
await m.refresh();
|
||
await m.downloadModel(m.catalog.first);
|
||
expect(downloadHits, 1);
|
||
|
||
// 重启 + 服务器目录出 v2:refresh 内部 autoUpdate 自动重下(无需手动)
|
||
final m2 = manager(client([_item(version: 'v2.0.0')]));
|
||
await m2.refresh();
|
||
for (var i = 0; i < 50 && downloadHits < 2; i++) {
|
||
await Future<void>.delayed(const Duration(milliseconds: 20));
|
||
}
|
||
|
||
expect(downloadHits, 2, reason: '新版本应自动重下');
|
||
expect(m2.isActive(7), isTrue, reason: '自动更新应保持激活');
|
||
expect(m2.models.first.version, 'v2.0.0',
|
||
reason: '自动更新后立即生效新版本字节');
|
||
expect(m2.models.first.bytes, _modelBytes);
|
||
});
|
||
|
||
test('sha256 不匹配:重试后失败、不激活、错误可见', () async {
|
||
final m = manager(client([_item()]));
|
||
await m.refresh();
|
||
final bad = _item(version: 'v3.0.0');
|
||
bad['sha256'] = _shaHex([9, 9, 9]);
|
||
final m2 = manager(client([bad]));
|
||
await m2.refresh();
|
||
|
||
final ok = await m2.downloadModel(m2.catalog.first);
|
||
expect(ok, isFalse);
|
||
expect(downloadHits, 2, reason: '校验失败应重试一次');
|
||
expect(m2.isDownloaded(7), isFalse);
|
||
expect(m2.isActive(7), isFalse);
|
||
expect(m2.errorOf(7), isNotNull);
|
||
});
|
||
|
||
test('激活集持久化:重启后恢复激活且已下载的模型', () async {
|
||
final m = manager(client([_item()]));
|
||
await m.refresh();
|
||
await m.downloadModel(m.catalog.first);
|
||
|
||
// 同一 root 新建 manager 模拟重启
|
||
final m2 = manager(client([_item()]));
|
||
await m2.refresh();
|
||
expect(m2.isActive(7), isTrue, reason: '激活集应持久化');
|
||
expect(m2.models.length, 1);
|
||
expect(downloadHits, 1, reason: '重启不应触发下载');
|
||
});
|
||
|
||
test('目录下线:清理本地并移除激活', () async {
|
||
final m = manager(client([_item()]));
|
||
await m.refresh();
|
||
await m.downloadModel(m.catalog.first);
|
||
expect(await Directory('${root.path}/7').exists(), isTrue);
|
||
|
||
final m2 = manager(client([]));
|
||
await m2.refresh();
|
||
expect(m2.models, isEmpty);
|
||
expect(m2.isActive(7), isFalse, reason: '下线的模型应移出激活集');
|
||
expect(await Directory('${root.path}/7').exists(), isFalse,
|
||
reason: '下线的数据集模型目录应被清理');
|
||
});
|
||
|
||
test('多模型:只激活其一则只加载其一', () async {
|
||
final m = manager(
|
||
client([_item(datasetId: 7, name: '环颈雉'), _item(datasetId: 8, name: '斑鸠')]));
|
||
await m.refresh();
|
||
await m.downloadModel(m.catalog.first); // 只下载并激活 7
|
||
expect(m.models.length, 1);
|
||
expect(m.models.first.datasetId, 7);
|
||
|
||
await m.downloadModel(m.catalog.last); // 下载 8 自动激活
|
||
expect(m.models.length, 2);
|
||
});
|
||
|
||
test('目录接口异常:不覆盖已有就绪状态', () async {
|
||
final m = manager(client([_item()]));
|
||
await m.refresh();
|
||
await m.downloadModel(m.catalog.first);
|
||
expect(m.ready, isTrue);
|
||
|
||
final broken = manager(MockClient((_) async => http.Response('boom', 500)));
|
||
await broken.refresh();
|
||
expect(broken.ready, isFalse);
|
||
expect(broken.error, isNotNull);
|
||
});
|
||
|
||
test('取消下载:中断后恢复未下载、不记错误', () async {
|
||
final m = manager(MockClient((req) async {
|
||
if (req.url.path == '/api/v1/app/update') {
|
||
return http.Response.bytes(
|
||
utf8.encode(jsonEncode(_catalog([_item()]))), 200);
|
||
}
|
||
if (req.url.path.startsWith('/download/models/')) {
|
||
downloadHits++;
|
||
return http.Response.fromStream(http.StreamedResponse(
|
||
_delayedChunks(), 200,
|
||
contentLength: _modelBytes.length));
|
||
}
|
||
return http.Response('not found', 404);
|
||
}));
|
||
await m.refresh();
|
||
|
||
final fut = m.downloadModel(m.catalog.first);
|
||
// 首个分块到达后取消(模拟用户在下载中点取消)
|
||
await Future<void>.delayed(const Duration(milliseconds: 20));
|
||
m.cancelDownload(7);
|
||
final ok = await fut;
|
||
|
||
expect(ok, isFalse);
|
||
expect(downloadHits, 1);
|
||
expect(m.isDownloaded(7), isFalse);
|
||
expect(m.isActive(7), isFalse);
|
||
expect(m.progressOf(7), isNull);
|
||
expect(m.errorOf(7), isNull, reason: '取消不记错误');
|
||
expect(await File('${root.path}/7/model.tflite.part').exists(), isFalse,
|
||
reason: '取消后 .part 残留应被清理');
|
||
});
|
||
}
|