Files
observer/flutter_app/lib/models/model_manager.dart
T

403 lines
14 KiB
Dart
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 'dart:convert';
import 'dart:io';
import 'package:crypto/crypto.dart' show sha256;
import 'package:flutter/foundation.dart';
import 'package:http/http.dart' as http;
import 'package:path_provider/path_provider.dart';
import '../config/app_config.dart';
/// 模型目录条目(GET /api/v1/app/update 响应 data.models[])。
/// 服务器发布模型后随版本检查一同下发,App 按目录逐数据集下载/更新。
class ModelCatalogItem {
final int datasetId;
final String datasetName;
final String version;
final List<String> labels;
final int sizeBytes;
final String sha256;
final String downloadUrl;
final String coverUrl;
const ModelCatalogItem({
required this.datasetId,
required this.datasetName,
required this.version,
required this.labels,
required this.sizeBytes,
required this.sha256,
required this.downloadUrl,
this.coverUrl = '',
});
factory ModelCatalogItem.fromJson(Map<String, dynamic> j) =>
ModelCatalogItem(
datasetId: (j['datasetId'] as num?)?.toInt() ?? 0,
datasetName: j['datasetName'] as String? ?? '',
version: j['version'] as String? ?? '',
labels: (j['labels'] as List? ?? const [])
.map((e) => e.toString())
.toList(),
sizeBytes: (j['sizeBytes'] as num?)?.toInt() ?? 0,
sha256: j['sha256'] as String? ?? '',
downloadUrl: j['downloadUrl'] as String? ?? '',
coverUrl: j['coverUrl'] as String? ?? '',
);
}
/// 已就绪模型(字节 + 标签,供推理 worker 加载)
class ModelBundle {
final int datasetId;
final String datasetName;
final String version;
final List<String> labels;
final Uint8List bytes;
const ModelBundle({
required this.datasetId,
required this.datasetName,
required this.version,
required this.labels,
required this.bytes,
});
}
/// 模型热更新管理:启动时拉取模型目录(随 /app/update 公开接口下发,无需登录态),
/// 按需下载/校验/持久化各数据集模型,供相机页多模型并行推理。
///
/// 存储:应用私有目录 `models/<datasetId>/`model.tflite + labels.json + meta.json),
/// meta 记录 {version, sha256},服务器发布新版本时按版本+摘要重下,不重复下载旧模型。
class ModelManager extends ChangeNotifier {
static final ModelManager instance = ModelManager._();
final String baseUrl;
final http.Client _client;
final Future<Directory> Function()? _rootDirOverride;
List<ModelBundle> _models = const [];
List<ModelCatalogItem> _catalog = const [];
final Set<int> _activeIds = {};
final Set<int> _downloadedIds = {};
final Map<int, double> _progress = {};
final Map<int, String> _errors = {};
bool _activeLoaded = false;
bool _ready = false;
bool _refreshing = false;
String? _error;
Future<void>? _inFlight;
/// 服务器目录(弹层模型清单展示用)
List<ModelCatalogItem> get catalog => _catalog;
/// 激活模型 id 集合(多选叠加)
Set<int> get activeDatasetIds => Set.unmodifiable(_activeIds);
bool isActive(int datasetId) => _activeIds.contains(datasetId);
/// 该数据集模型文件是否已下载到本地(同步判断,内存态)
bool isDownloaded(int datasetId) => _downloadedIds.contains(datasetId);
/// 下载进度 0..1(无下载/已完成为 null)
double? progressOf(int datasetId) => _progress[datasetId];
/// 下载失败原因(失败后可重试)
String? errorOf(int datasetId) => _errors[datasetId];
ModelManager._({String? baseUrl, http.Client? client})
: this(baseUrl: baseUrl, client: client);
/// 可注入 baseUrl / client / 存储根目录(单测用)
@visibleForTesting
ModelManager({
String? baseUrl,
http.Client? client,
Future<Directory> Function()? rootDir,
}) : baseUrl = baseUrl ?? AppConfig.apiBaseUrl,
_client = client ?? http.Client(),
_rootDirOverride = rootDir;
/// 已激活且已下载的模型列表(空 = 未下载任何模型,相机页仅预览)
List<ModelBundle> get models => _models;
/// 是否成功拉取过目录(即使下载失败也为 true,用于区分"从未联网"与"目录为空"
bool get ready => _ready;
/// 模型目录拉取失败的错误信息(仅目录级;下载/校验失败见 errorOf)
String? get error => _error;
bool get refreshing => _refreshing;
/// 模型名摘要(诊断行展示):未下载 / 数据集名×n
String get modelsLabel {
if (_models.isEmpty) return '未下载';
return _models.map((m) => m.datasetName).join(',');
}
/// 拉取目录并同步本地模型;并发调用共享同一进行中的刷新。
Future<void> refresh() {
if (_refreshing) return _inFlight ?? Future.value();
_refreshing = true;
_inFlight = _doRefresh().whenComplete(() {
_refreshing = false;
_inFlight = null;
notifyListeners();
});
return _inFlight!;
}
Future<void> _doRefresh() async {
try {
await _loadActive();
final res = await _client
.get(Uri.parse('$baseUrl/api/v1/app/update'))
.timeout(const Duration(seconds: 8));
// 服务器 Content-Type 无 charsethttp 包默认按 latin1 解码会乱码 → 显式 utf8
final body =
jsonDecode(utf8.decode(res.bodyBytes)) as Map<String, dynamic>;
final data = body['data'] as Map<String, dynamic>? ?? const {};
final list = data['models'] as List? ?? const [];
_catalog = list
.map((e) => ModelCatalogItem.fromJson(e as Map<String, dynamic>))
.toList();
// 只拉目录不下载;扫描本地已下载(meta+文件齐备)供清单展示
final downloaded = <int>{};
for (final item in _catalog) {
if (await _isLocal(item)) downloaded.add(item.datasetId);
}
_downloadedIds
..clear()
..addAll(downloaded);
await _prune(_catalog);
// 服务器已下线的数据集移出激活集
final catalogIds = _catalog.map((c) => c.datasetId).toSet();
if (_activeIds.any((id) => !catalogIds.contains(id))) {
_activeIds.removeWhere((id) => !catalogIds.contains(id));
await _saveActive();
}
_models = await _loadBundles(_catalog);
_ready = true;
_error = null;
} catch (e) {
if (!_ready) _error = '模型目录拉取失败:$e';
// 已就绪过则保留旧目录/旧模型,不覆盖 error(下载级错误优先展示)
}
}
/// 本地是否已有匹配版本的文件(meta 版本+sha256 相符且文件存在)
Future<bool> _isLocal(ModelCatalogItem item) async {
final dir = await _modelDir(item.datasetId);
try {
final meta = await _readMeta(dir);
final file = File('${dir.path}/model.tflite');
return meta != null &&
meta['version'] == item.version &&
meta['sha256'] == item.sha256 &&
await file.exists();
} catch (e) {
return false;
}
}
/// 按需下载并激活:流式下载 + sha256 校验 + 落盘(labels/meta);
/// 成功自动加入激活集(下载完成即使用)。失败重试一次并记录错误。
Future<bool> downloadModel(ModelCatalogItem item,
{void Function(int received, int total)? onProgress}) async {
// 并发保护:同一数据集已有进行中的下载则直接短路(预置 0 先占位,
// 使 onProgress 首次回调前的双击/refresh 交错也被 containsKey 拦下)
if (_progress.containsKey(item.datasetId)) return false;
_progress[item.datasetId] = 0;
final dir = await _modelDir(item.datasetId);
final file = File('${dir.path}/model.tflite');
try {
for (var attempt = 0; attempt < 2; attempt++) {
final ok = await _downloadAndVerify(item, dir, file,
onProgress: (r, t) {
_progress[item.datasetId] = t == 0 ? 0 : r / t;
onProgress?.call(r, t);
notifyListeners();
});
if (ok) {
_progress.remove(item.datasetId);
_errors.remove(item.datasetId);
_downloadedIds.add(item.datasetId);
await setActive(item.datasetId, true);
return true;
}
await file.delete().catchError((_) => file);
await File('${dir.path}/model.tflite.part')
.delete()
.catchError((_) => file);
}
_progress.remove(item.datasetId);
_errors[item.datasetId] = '下载失败,请重试';
notifyListeners();
debugPrint('[ModelManager] 下载失败: ${item.datasetName} ${item.version}');
return false;
} catch (e) {
_progress.remove(item.datasetId);
_errors[item.datasetId] = '下载异常:$e';
notifyListeners();
debugPrint('[ModelManager] 下载异常 ${item.datasetName}: $e');
return false;
}
}
Future<bool> _downloadAndVerify(
ModelCatalogItem item, Directory dir, File file,
{void Function(int received, int total)? onProgress}) async {
final part = File('${file.path}.part');
final sink = part.openWrite();
var received = 0;
try {
final res = await _client
.send(http.Request('GET', Uri.parse('$baseUrl${item.downloadUrl}')))
.timeout(const Duration(minutes: 3));
if (res.statusCode != 200) {
await sink.close();
return false;
}
final total = res.contentLength ?? item.sizeBytes;
await for (final chunk in res.stream) {
sink.add(chunk);
received += chunk.length;
onProgress?.call(received, total);
}
await sink.close();
final bytes = await part.readAsBytes();
final hex = sha256.convert(bytes).toString();
if (item.sha256.isNotEmpty && hex != item.sha256) {
debugPrint('[ModelManager] sha256 不匹配: ${item.datasetName} '
'want=${item.sha256} got=$hex');
return false;
}
await part.rename(file.path);
await dir.create(recursive: true);
await File('${dir.path}/labels.json')
.writeAsString(jsonEncode(item.labels));
await File('${dir.path}/meta.json').writeAsString(jsonEncode({
'version': item.version,
'sha256': item.sha256,
}));
debugPrint('[ModelManager] 已下载 ${item.datasetName} '
'${bytes.length}B -> ${file.path}');
return true;
} catch (e) {
await sink.close().catchError((_) {});
debugPrint('[ModelManager] 下载异常 ${item.datasetName}: $e');
return false;
}
}
/// 清理服务器目录中已下线的数据集模型(不再发布则删本地)
Future<void> _prune(List<ModelCatalogItem> catalog) async {
final root = await _rootDir();
if (!await root.exists()) return;
final keep = catalog.map((c) => '${c.datasetId}').toSet();
await for (final e in root.list()) {
if (e is Directory) {
// 目录 URI 末尾带 '/'pathSegments 末位为空串 → 过滤后取目录名
final name = e.uri.pathSegments.where((s) => s.isNotEmpty).last;
if (!keep.contains(name)) {
await e.delete(recursive: true).catchError((_) => e);
}
}
}
}
/// 设置激活状态(true=使用,false=取消);持久化到 `root/active.json`。
/// 未下载的模型不可激活(下载完成由 downloadModel 自动激活)。
Future<void> setActive(int datasetId, bool active) async {
final changed =
active ? _activeIds.add(datasetId) : _activeIds.remove(datasetId);
if (!changed) return;
_models = await _loadBundles(_catalog);
await _saveActive();
notifyListeners();
}
Future<void> _saveActive() async {
try {
final root = await _rootDir();
await root.create(recursive: true);
await File('${root.path}/active.json')
.writeAsString(jsonEncode({'active': _activeIds.toList()}));
} catch (e) {
debugPrint('[ModelManager] 激活集持久化失败: $e');
}
}
Future<void> _loadActive() async {
if (_activeLoaded) return;
_activeLoaded = true;
try {
final root = await _rootDir();
final f = File('${root.path}/active.json');
if (!await f.exists()) return;
final data = jsonDecode(await f.readAsString()) as Map<String, dynamic>;
_activeIds
..clear()
..addAll((data['active'] as List? ?? const [])
.map((e) => (e as num).toInt()));
} catch (e) {
debugPrint('[ModelManager] 激活集读取失败: $e');
}
}
Future<List<ModelBundle>> _loadBundles(
List<ModelCatalogItem> catalog) async {
final bundles = <ModelBundle>[];
for (final item in catalog) {
if (!_activeIds.contains(item.datasetId)) continue;
try {
final dir = await _modelDir(item.datasetId);
final file = File('${dir.path}/model.tflite');
if (!await file.exists()) continue;
final labels = await File('${dir.path}/labels.json').exists()
? (jsonDecode(
await File('${dir.path}/labels.json').readAsString())
as List)
.map((e) => e.toString())
.toList()
: item.labels;
bundles.add(ModelBundle(
datasetId: item.datasetId,
datasetName: item.datasetName,
version: item.version,
labels: labels,
bytes: await file.readAsBytes(),
));
} catch (e) {
debugPrint('[ModelManager] 读取 ${item.datasetName} 失败: $e');
}
}
return bundles;
}
Future<Map<String, dynamic>?> _readMeta(Directory dir) async {
final f = File('${dir.path}/meta.json');
if (!await f.exists()) return null;
try {
return jsonDecode(await f.readAsString()) as Map<String, dynamic>;
} catch (_) {
return null;
}
}
Future<Directory> _rootDir() async {
if (_rootDirOverride != null) return _rootDirOverride();
final support = await getApplicationSupportDirectory();
return Directory('${support.path}/models');
}
Future<Directory> _modelDir(int datasetId) async {
final root = await _rootDir();
final dir = Directory('${root.path}/$datasetId');
await dir.create(recursive: true);
return dir;
}
}