Files
2026-09-21 17:37:53 +08:00

129 lines
3.9 KiB
Dart

import 'package:flutter_riverpod/flutter_riverpod.dart';
import '../api/bill_api.dart';
import '../api/category_api.dart';
import '../api/ocr_api.dart';
import '../core/network/api_exception.dart';
import '../models/bill_request.dart';
import '../models/ocr_result.dart';
import '../utils/date_util.dart';
import 'bill_refresh.dart';
/// OCR 识别流程状态
class OcrState {
final String? imagePath;
final List<OcrResult> results;
final Set<int> expandedIndexes;
final bool recognizing;
final bool saving;
final String? error;
const OcrState({
this.imagePath,
this.results = const [],
this.expandedIndexes = const {},
this.recognizing = false,
this.saving = false,
this.error,
});
OcrState copyWith({
Object? imagePath = _sentinel,
List<OcrResult>? results,
Set<int>? expandedIndexes,
bool? recognizing,
bool? saving,
Object? error = _sentinel,
}) {
return OcrState(
imagePath: imagePath == _sentinel ? this.imagePath : imagePath as String?,
results: results ?? this.results,
expandedIndexes: expandedIndexes ?? this.expandedIndexes,
recognizing: recognizing ?? this.recognizing,
saving: saving ?? this.saving,
error: error == _sentinel ? this.error : error as String?,
);
}
static const _sentinel = Object();
}
class OcrNotifier extends Notifier<OcrState> {
@override
OcrState build() => const OcrState();
void setImagePath(String path) {
state = OcrState(imagePath: path);
}
void clearImage() {
state = const OcrState();
}
void toggleExpanded(int index) {
final indexes = {...state.expandedIndexes};
if (!indexes.remove(index)) indexes.add(index);
state = state.copyWith(expandedIndexes: indexes);
}
/// 上传识别
Future<void> recognize() async {
final path = state.imagePath;
if (path == null) throw ApiException(-1, '请先选择图片');
state = state.copyWith(recognizing: true, error: null);
try {
final results = await ref.read(ocrApiProvider).recognize(path);
state = state.copyWith(
results: results,
// 只有一条时自动展开
expandedIndexes: results.length == 1 ? {0} : <int>{},
);
} on ApiException catch (e) {
state = state.copyWith(error: e.message);
rethrow;
} finally {
state = state.copyWith(recognizing: false);
}
}
/// 直接批量存储:按金额正负推断类型,分类取该类型下第一项
Future<void> saveDirectly() async {
if (state.results.isEmpty) throw ApiException(-1, '没有可存储的识别结果');
state = state.copyWith(saving: true, error: null);
try {
final categoryApi = ref.read(categoryApiProvider);
final bills = <BillRequest>[];
final cache = <int, int?>{}; // type -> 默认分类id
for (final r in state.results) {
final type = r.isExpense ? 1 : 2;
if (!cache.containsKey(type)) {
final categories = await categoryApi.getCategories(type);
cache[type] = categories.isNotEmpty ? categories.first.id : null;
}
final categoryId = cache[type];
if (categoryId == null) {
throw ApiException(-1, '缺少${type == 1 ? '支出' : '收入'}分类,请先创建分类');
}
bills.add(BillRequest(
categoryId: categoryId,
amount: r.absAmount,
description: r.merchant ?? '',
billDate: parseBillDate(r.date),
type: type,
));
}
await ref.read(billApiProvider).createBillsBatch(bills);
// 让首页列表与统计页重新取数
invalidateBillData(ref);
state = const OcrState();
} on ApiException catch (e) {
state = state.copyWith(error: e.message);
rethrow;
} finally {
state = state.copyWith(saving: false);
}
}
}
final ocrProvider = NotifierProvider<OcrNotifier, OcrState>(OcrNotifier.new);