129 lines
3.9 KiB
Dart
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);
|