训练体系整合与标注单阶段化
- 标注:AI 预标注直写 labels_json(去候选确认两阶段);重叠去重(minIoU);全量标注按钮 - 训练:脚本迁移入 server/training/(Go 化 prepare_yolo/analyze_rfdetr,保留 train_server.py);tflite 产物自检并入训练流程(check_tflite) - 数据目录/权重不进 git;.gitignore 迁移至仓库根
This commit is contained in:
+13
-16
@@ -1,17 +1,14 @@
|
||||
*.iml
|
||||
.gradle/
|
||||
local.properties
|
||||
.idea/
|
||||
build/
|
||||
captures/
|
||||
.externalNativeBuild/
|
||||
.cxx/
|
||||
.DS_Store
|
||||
training/datasets/
|
||||
training/venv/
|
||||
training/runs/
|
||||
*.pt
|
||||
*.tflite
|
||||
# 运行时数据(不提交 git)
|
||||
server/data/
|
||||
server/workspace/
|
||||
|
||||
# 文生图配置(含API key,不入库)
|
||||
#training/gen_images_config.json
|
||||
# 预训练权重资产(训练机部署用,不进 git)
|
||||
server/training/yolov8n.pt
|
||||
|
||||
# 系统文件
|
||||
.DS_Store
|
||||
|
||||
# 开发工具与构建产物
|
||||
.gstack/
|
||||
.idea/
|
||||
flutter_app/android/build/
|
||||
|
||||
@@ -38,7 +38,7 @@ app.*.symbols
|
||||
|
||||
# Obfuscation related
|
||||
app.*.map.json
|
||||
|
||||
observer-latest.apk
|
||||
# Android Studio will place build artifacts here
|
||||
/android/app/debug
|
||||
/android/app/profile
|
||||
|
||||
+28
-4
@@ -1,6 +1,6 @@
|
||||
# observer
|
||||
|
||||
野生动物实时识别 App(Flutter 版)。Android / iOS 一套代码,后端接口与支付见
|
||||
动物实时识别 App(Flutter 版)。Android / iOS 一套代码,后端接口与支付见
|
||||
[`docs/PaymentApi.md`](docs/PaymentApi.md)。
|
||||
|
||||
## iOS 真机部署(iPhone)
|
||||
@@ -8,9 +8,10 @@
|
||||
### 构建与安装
|
||||
|
||||
```bash
|
||||
# 真机必须传 Mac 局域网 IP:默认 API_BASE_URL 是 10.0.2.2(仅 Android 模拟器可用),
|
||||
# 不传则 iPhone 上所有网络请求(登录/授权/套餐)都会失败
|
||||
flutter build ios --release --dart-define=API_BASE_URL=http://<Mac局域网IP>:8080
|
||||
# 生产包直接构建即可:默认 API_BASE_URL 为线上域名(lib/config/app_config.dart),无需传参
|
||||
flutter build ios --release
|
||||
# 仅本地联调(后端跑在 Mac 上、手机连同一 Wi-Fi)时才覆盖为 Mac 局域网 IP:
|
||||
# flutter build ios --release --dart-define=API_BASE_URL=http://<Mac局域网IP>:8080
|
||||
|
||||
# 安装到真机(UDID 可用 `xcrun devicectl list devices` 查询)
|
||||
xcrun devicectl device install app --device <UDID> build/ios/iphoneos/Runner.app
|
||||
@@ -26,6 +27,29 @@ xcrun devicectl device process launch --console --terminate-existing \
|
||||
"In iOS 14+, debug mode Flutter apps can only be launched from Flutter tooling"。
|
||||
debug 调试必须用 `flutter run -d <设备ID>` 或 Xcode IDE 启动(`flutter devices` 查设备ID);
|
||||
从图标启动只对 release 构建有效。
|
||||
## 模型热更新(多数据集模型)
|
||||
|
||||
模型与 APK 更新走**独立通道**:启动时拉取 `GET /api/v1/app/update` 随附的
|
||||
`models` 目录(公开接口,无需登录),与 `UpdateChecker` 的 APK 检查并行。
|
||||
|
||||
- **目录条目**:`{datasetId, datasetName, version, labels[], sizeBytes, sha256,
|
||||
downloadUrl}`;服务器未发布模型时不返回 `models` 字段,App 回退内置
|
||||
`assets/model.tflite`(标签 `内置`)。
|
||||
- **存储**:应用私有目录 `models/<datasetId>/`,含 `model.tflite`、
|
||||
`labels.json`、`meta.json`(meta 记录 `{version, sha256}`)。版本与摘要都
|
||||
未变化时跳过下载;变化则下载到 `.part` 临时文件、sha256 校验通过后
|
||||
原子 rename 替换,失败重试一次并保留旧模型,下次启动再试。
|
||||
- **清理**:服务器下线的数据集下次同步时删除本地对应目录。
|
||||
- **并行推理合并**:识别时加载全部已下载模型(`DetectorWorker` isolate 内
|
||||
逐模型加载,单个失败不影响其他),同帧各模型独立推理后按类别分组做
|
||||
**跨模型 NMS**(同类别不同模型检出同一目标取高分去重,不同类别互不压制),
|
||||
结果叠加 `modelName` 标注来源,内置模型兜底。
|
||||
|
||||
实现:`lib/models/model_manager.dart`(下载/校验/持久化,`ModelManager`
|
||||
单例 + ChangeNotifier)、`lib/detection/detector_worker.dart`(多模型并行
|
||||
推理与 `mergeAcrossModels`)、`lib/camera/camera_screen.dart`(启动同步 +
|
||||
诊断行展示模型列表)。
|
||||
|
||||
- **模型输入是 NHWC**:`assets/model.tflite` 做过字节级手术(开头 TRANSPOSE→RESHAPE,
|
||||
输入 [1,320,320,3]),改动记录见 git 历史,重导模型需同步处理,否则 iOS 报
|
||||
"Node number 0 (TRANSPOSE) failed to prepare"。
|
||||
|
||||
@@ -40,6 +40,10 @@ kotlin {
|
||||
}
|
||||
}
|
||||
|
||||
// 相机用 Android 框架 camera2 API 完全自研(CameraChannel.kt):
|
||||
// 预览 SurfaceTexture(Flutter 纹理)+ ImageReader 分析帧(原生侧旋转成竖屏后回传),
|
||||
// 不依赖任何相机三方库(含 CameraX)。
|
||||
|
||||
// tflite_flutter 依赖的 tensorflow-lite / tensorflow-lite-gpu / tensorflow-lite-api 三个 AAR
|
||||
// 声明了相同 namespace(org.tensorflow.lite),新 AGP 视作冲突直接报错;
|
||||
// 本项目仅用 CPU 推理,GPU delegate 未使用,排除 gpu 及其传递依赖的 api 即可。
|
||||
|
||||
+13
@@ -0,0 +1,13 @@
|
||||
# R8 release 压缩:okhttp 可选 TLS 平台类(BouncyCastle/Conscrypt/OpenJSSE)
|
||||
# 与 tflite_flutter 的反射注解类未打包,仅需忽略引用告警(AGP missing_rules.txt 生成)
|
||||
-dontwarn org.bouncycastle.jsse.BCSSLParameters
|
||||
-dontwarn org.bouncycastle.jsse.BCSSLSocket
|
||||
-dontwarn org.bouncycastle.jsse.provider.BouncyCastleJsseProvider
|
||||
-dontwarn org.conscrypt.Conscrypt$Version
|
||||
-dontwarn org.conscrypt.Conscrypt
|
||||
-dontwarn org.conscrypt.ConscryptHostnameVerifier
|
||||
-dontwarn org.openjsse.javax.net.ssl.SSLParameters
|
||||
-dontwarn org.openjsse.javax.net.ssl.SSLSocket
|
||||
-dontwarn org.openjsse.net.ssl.OpenJSSE
|
||||
-dontwarn org.tensorflow.lite.InterpreterFactoryApi
|
||||
-dontwarn org.tensorflow.lite.annotations.UsedByReflection
|
||||
@@ -1,6 +1,8 @@
|
||||
<manifest xmlns:android="http://schemas.android.com/apk/res/android">
|
||||
<uses-permission android:name="android.permission.CAMERA"/>
|
||||
<uses-permission android:name="android.permission.INTERNET"/>
|
||||
<!-- App 内更新安装 APK(PackageInstaller 会话安装,见 InstallerChannel.kt) -->
|
||||
<uses-permission android:name="android.permission.REQUEST_INSTALL_PACKAGES"/>
|
||||
<application
|
||||
android:label="视野"
|
||||
android:name="${applicationName}"
|
||||
|
||||
@@ -0,0 +1,568 @@
|
||||
package com.example.observer
|
||||
|
||||
import android.graphics.Rect
|
||||
import android.graphics.RectF
|
||||
import android.graphics.SurfaceTexture
|
||||
import android.os.Build
|
||||
import android.hardware.camera2.CameraCaptureSession
|
||||
import android.hardware.camera2.CameraCharacteristics
|
||||
import android.hardware.camera2.CameraDevice
|
||||
import android.hardware.camera2.CameraManager
|
||||
import android.hardware.camera2.CaptureRequest
|
||||
import android.graphics.ImageFormat
|
||||
import android.media.Image
|
||||
import android.media.ImageReader
|
||||
import android.os.Handler
|
||||
import android.os.HandlerThread
|
||||
import android.os.SystemClock
|
||||
import android.util.Log
|
||||
import android.util.Size
|
||||
import android.view.Surface
|
||||
import io.flutter.embedding.android.FlutterActivity
|
||||
import io.flutter.embedding.engine.FlutterEngine
|
||||
import io.flutter.plugin.common.EventChannel
|
||||
import io.flutter.plugin.common.MethodCall
|
||||
import io.flutter.plugin.common.MethodChannel
|
||||
import io.flutter.view.TextureRegistry
|
||||
|
||||
/**
|
||||
* 自写原生相机:Android 框架 camera2 API 完全自研,零相机三方依赖(含 CameraX)。
|
||||
*
|
||||
* 核心:分析帧在原生侧旋转成竖屏方向再回传,Flutter 侧恒用 rotation=0,
|
||||
* 与 iOS(camera_avfoundation,帧本来就是竖屏)行为对齐,彻底消除
|
||||
* "横屏传感器帧 → 90° 旋转 + FIT_COVER crop 映射"的标注偏移根因。
|
||||
*
|
||||
* 旋转角来自设备标准值 sensorOrientation - displayRotation(camera2 特性,
|
||||
* 非厂商 hack),因此映射跨厂商一致:任何设备上"分析帧 = 预览帧同 sensor
|
||||
* 同旋转",两者几何必然一致。
|
||||
*
|
||||
* 预览:SurfaceTexture 注册进 FlutterTextureRegistry,Flutter 侧 Texture widget
|
||||
* 渲染(与插件 CameraPreview 相同的合成方式,overlay/诊断行/按钮可叠加;
|
||||
* AndroidView+SurfaceView 会盖住 Flutter UI,不可用)。
|
||||
*
|
||||
* 通道:
|
||||
* - MethodChannel "observer/camera":start/stop/getZoomRange/setZoom/isStreaming/errorDescription
|
||||
* - EventChannel "observer/camera/frames":每帧 [width, height, rotationDegrees, bgra, bytes]
|
||||
* - rotationDegrees 恒 0(帧已竖屏);bgra=false(RGBA 字节序)
|
||||
*/
|
||||
class CameraChannel(
|
||||
private val activity: FlutterActivity,
|
||||
private val engine: FlutterEngine,
|
||||
) {
|
||||
private val messenger = engine.dartExecutor.binaryMessenger
|
||||
private val method = MethodChannel(messenger, "observer/camera")
|
||||
private val frames = EventChannel(messenger, "observer/camera/frames")
|
||||
|
||||
private val cameraManager: CameraManager = activity.getSystemService(CameraManager::class.java)
|
||||
|
||||
private var cameraDevice: CameraDevice? = null
|
||||
private var captureSession: CameraCaptureSession? = null
|
||||
private var imageReader: ImageReader? = null
|
||||
private var surfaceEntry: TextureRegistry.SurfaceTextureEntry? = null
|
||||
private var previewSurface: Surface? = null
|
||||
private var eventSink: EventChannel.EventSink? = null
|
||||
|
||||
private var sensorOrientation = 90
|
||||
private var activeArray = Rect(0, 0, 1920, 1080)
|
||||
private var maxZoom = 1.0f
|
||||
|
||||
/// 预览/分析共用分辨率(传感器方向,从设备流配置动态选择)
|
||||
private var previewSize = Size(1920, 1080)
|
||||
|
||||
@Volatile
|
||||
private var running = false
|
||||
|
||||
@Volatile
|
||||
private var errorDescription: String? = null
|
||||
|
||||
private val mainHandler = Handler(activity.mainLooper)
|
||||
private val analysisThread = HandlerThread("observer-analysis").also { it.start() }
|
||||
private val analysisHandler = Handler(analysisThread.looper)
|
||||
|
||||
@Volatile
|
||||
private var lastEmitMs = 0L
|
||||
|
||||
// 诊断计数(frameListener 线程写,stats 轮询读):
|
||||
// 帧回调到达次数 / 成功发出 / 发出异常 / 无订阅者丢弃
|
||||
@Volatile
|
||||
private var frameCallbacks = 0L
|
||||
|
||||
@Volatile
|
||||
private var emitOk = 0L
|
||||
|
||||
@Volatile
|
||||
private var emitErr = 0L
|
||||
|
||||
@Volatile
|
||||
private var sinkNullCount = 0L
|
||||
|
||||
// 实际发出帧尺寸(缩小后)
|
||||
@Volatile
|
||||
private var emitW = 0
|
||||
|
||||
@Volatile
|
||||
private var emitH = 0
|
||||
|
||||
fun register() {
|
||||
method.setMethodCallHandler(::onMethodCall)
|
||||
frames.setStreamHandler(object : EventChannel.StreamHandler {
|
||||
override fun onListen(arguments: Any?, events: EventChannel.EventSink?) {
|
||||
eventSink = events
|
||||
}
|
||||
|
||||
override fun onCancel(arguments: Any?) {
|
||||
eventSink = null
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
private fun onMethodCall(call: MethodCall, result: MethodChannel.Result) {
|
||||
when (call.method) {
|
||||
"start" -> start(result)
|
||||
"stop" -> {
|
||||
stop()
|
||||
result.success(true)
|
||||
}
|
||||
"getZoomRange" -> result.success(listOf(1.0, maxZoom.toDouble()))
|
||||
"setZoom" -> {
|
||||
val v = (call.arguments as Number).toFloat()
|
||||
setZoom(v)
|
||||
result.success(true)
|
||||
}
|
||||
"isStreaming" -> result.success(running)
|
||||
"errorDescription" -> result.success(errorDescription)
|
||||
"stats" -> result.success(
|
||||
mapOf(
|
||||
"callbacks" to frameCallbacks,
|
||||
"emitOk" to emitOk,
|
||||
"emitErr" to emitErr,
|
||||
"sinkNull" to sinkNullCount,
|
||||
"sink" to (eventSink != null),
|
||||
"running" to running,
|
||||
"rotation" to rotation,
|
||||
"quarterTurns" to displayDegrees / 90,
|
||||
"displayDegrees" to displayDegrees,
|
||||
"size" to "${previewSize.width}x${previewSize.height}",
|
||||
"emitSize" to "${emitW}x$emitH",
|
||||
"emitAgeMs" to
|
||||
if (lastEmitMs == 0L) -1L
|
||||
else SystemClock.elapsedRealtime() - lastEmitMs,
|
||||
"error" to errorDescription,
|
||||
),
|
||||
)
|
||||
else -> result.notImplemented()
|
||||
}
|
||||
}
|
||||
|
||||
private fun start(result: MethodChannel.Result) {
|
||||
if (running) {
|
||||
result.success(true)
|
||||
return
|
||||
}
|
||||
running = true
|
||||
errorDescription = null
|
||||
try {
|
||||
val cameraId = pickBackCamera() ?: throw IllegalStateException("未找到后置摄像头")
|
||||
val characteristics = cameraManager.getCameraCharacteristics(cameraId)
|
||||
sensorOrientation = characteristics.get(CameraCharacteristics.SENSOR_ORIENTATION) ?: 90
|
||||
activeArray = characteristics.get(
|
||||
CameraCharacteristics.SENSOR_INFO_ACTIVE_ARRAY_SIZE,
|
||||
) ?: Rect(0, 0, 1920, 1080)
|
||||
maxZoom = characteristics.get(CameraCharacteristics.SCALER_AVAILABLE_MAX_DIGITAL_ZOOM) ?: 1.0f
|
||||
// 分辨率不写死:从设备流配置查询(传感器方向尺寸,宽≥高),
|
||||
// 预览与分析共用同一尺寸,保证两者几何一致(对齐根因)。
|
||||
// 查询是优化而非必需:个别设备 getOutputSizes 会返回 null 或抛异常
|
||||
// (实测某设备内部抛 getClass NPE),失败一律回退默认分辨率,
|
||||
// 绝不让相机启动失败
|
||||
try {
|
||||
val streamMap = characteristics.get(CameraCharacteristics.SCALER_STREAM_CONFIGURATION_MAP)
|
||||
if (streamMap != null) {
|
||||
val previewSizes =
|
||||
streamMap.getOutputSizes(SurfaceTexture::class.java) ?: emptyArray()
|
||||
val analysisSizes =
|
||||
streamMap.getOutputSizes(ImageFormat.YUV_420_888) ?: emptyArray()
|
||||
previewSize = pickSize(previewSizes, analysisSizes)
|
||||
}
|
||||
} catch (e: Exception) {
|
||||
Log.w(TAG, "size query failed, fallback ${previewSize.width}x${previewSize.height}", e)
|
||||
}
|
||||
|
||||
val entry = surfaceEntry ?: engine.getRenderer().createSurfaceTexture().also {
|
||||
surfaceEntry = it
|
||||
}
|
||||
// 关键:SurfaceTexture 未设默认 buffer 尺寸时 createCaptureSession 会配置失败
|
||||
// (camera2 按该尺寸做 stream 校验)。显式指定与 ImageReader 一致的分辨率,
|
||||
// 保证预览与分析帧同分辨率同裁剪
|
||||
entry.surfaceTexture().setDefaultBufferSize(previewSize.width, previewSize.height)
|
||||
previewSurface?.release()
|
||||
previewSurface = Surface(entry.surfaceTexture())
|
||||
imageReader?.close()
|
||||
// YUV_420_888 是 camera2 对所有设备保证支持的 ImageReader 输出格式;
|
||||
// RGBA_8888 个别设备不支持(实测:查询 NPE + 配置失败双连败),弃用
|
||||
imageReader = ImageReader.newInstance(
|
||||
previewSize.width, previewSize.height, ImageFormat.YUV_420_888, 2,
|
||||
).also { it.setOnImageAvailableListener(frameListener, analysisHandler) }
|
||||
|
||||
cameraManager.openCamera(
|
||||
cameraId,
|
||||
object : CameraDevice.StateCallback() {
|
||||
override fun onOpened(device: CameraDevice) {
|
||||
cameraDevice = device
|
||||
createSession(result)
|
||||
}
|
||||
|
||||
override fun onDisconnected(device: CameraDevice) {
|
||||
device.close()
|
||||
cameraDevice = null
|
||||
}
|
||||
|
||||
override fun onError(device: CameraDevice, error: Int) {
|
||||
device.close()
|
||||
cameraDevice = null
|
||||
running = false
|
||||
errorDescription = "camera open error $error"
|
||||
result.error("start", errorDescription, null)
|
||||
}
|
||||
},
|
||||
mainHandler,
|
||||
)
|
||||
} catch (e: Exception) {
|
||||
running = false
|
||||
errorDescription = e.toString()
|
||||
// details 带完整堆栈:任何残留异常都能在诊断行看到精确位置
|
||||
result.error("start", e.toString(), Log.getStackTraceString(e))
|
||||
}
|
||||
}
|
||||
|
||||
private fun createSession(result: MethodChannel.Result) {
|
||||
val device = cameraDevice ?: return
|
||||
val surfaces = listOfNotNull(previewSurface, imageReader?.surface)
|
||||
device.createCaptureSession(
|
||||
surfaces,
|
||||
object : CameraCaptureSession.StateCallback() {
|
||||
override fun onConfigured(session: CameraCaptureSession) {
|
||||
captureSession = session
|
||||
try {
|
||||
val request = buildRequest(session.device, surfaces)
|
||||
session.setRepeatingRequest(request, null, mainHandler)
|
||||
result.success(
|
||||
mapOf(
|
||||
"textureId" to (surfaceEntry?.id() ?: -1L),
|
||||
// 传感器方向尺寸(宽≥高):Flutter 侧 SizedBox 在 RotatedBox
|
||||
// 内部声明纹理尺寸,旋转后视觉上才是竖屏 1080x1920
|
||||
"w" to previewSize.width,
|
||||
"h" to previewSize.height,
|
||||
// 预览旋转 = 显示旋转(不是 rotation/90!):
|
||||
// Flutter 引擎渲染 Texture 时已自动应用 SurfaceTexture
|
||||
// 变换矩阵(传感器方向补偿),预览只需再按显示旋转补偿
|
||||
"quarterTurns" to displayDegrees / 90,
|
||||
// 诊断:旋转计算输入值
|
||||
"sensorOrientation" to sensorOrientation,
|
||||
"displayDegrees" to displayDegrees,
|
||||
),
|
||||
)
|
||||
} catch (e: Exception) {
|
||||
errorDescription = "onConfigured: $e"
|
||||
result.error("start", errorDescription, Log.getStackTraceString(e))
|
||||
}
|
||||
}
|
||||
|
||||
override fun onConfigureFailed(session: CameraCaptureSession) {
|
||||
running = false
|
||||
errorDescription = "capture session configure failed"
|
||||
result.error("start", errorDescription, null)
|
||||
}
|
||||
},
|
||||
mainHandler,
|
||||
)
|
||||
}
|
||||
|
||||
private fun buildRequest(device: CameraDevice, targets: List<Surface>): CaptureRequest {
|
||||
val builder = device.createCaptureRequest(CameraDevice.TEMPLATE_PREVIEW)
|
||||
targets.forEach { builder.addTarget(it) }
|
||||
builder.set(CaptureRequest.CONTROL_MODE, CaptureRequest.CONTROL_MODE_AUTO)
|
||||
// 连续对焦:无 AF 能力的设备忽略该设置
|
||||
try {
|
||||
builder.set(
|
||||
CaptureRequest.CONTROL_AF_MODE,
|
||||
CaptureRequest.CONTROL_AF_MODE_CONTINUOUS_PICTURE,
|
||||
)
|
||||
} catch (_: IllegalArgumentException) {
|
||||
}
|
||||
return builder.build()
|
||||
}
|
||||
|
||||
private fun setZoom(z: Float) {
|
||||
val session = captureSession ?: return
|
||||
val clamped = z.coerceIn(1.0f, maxZoom)
|
||||
val r = activeArray
|
||||
val insetW = r.width() * (1 - 1 / clamped) / 2
|
||||
val insetH = r.height() * (1 - 1 / clamped) / 2
|
||||
val crop = RectF(
|
||||
r.left + insetW,
|
||||
r.top + insetH,
|
||||
r.right - insetW,
|
||||
r.bottom - insetH,
|
||||
)
|
||||
try {
|
||||
val surface = previewSurface ?: return
|
||||
val builder = session.device.createCaptureRequest(CameraDevice.TEMPLATE_PREVIEW)
|
||||
builder.addTarget(surface)
|
||||
imageReader?.surface?.let { builder.addTarget(it) }
|
||||
builder.set(CaptureRequest.CONTROL_MODE, CaptureRequest.CONTROL_MODE_AUTO)
|
||||
builder.set(
|
||||
CaptureRequest.SCALER_CROP_REGION,
|
||||
Rect(
|
||||
crop.left.toInt(),
|
||||
crop.top.toInt(),
|
||||
crop.right.toInt(),
|
||||
crop.bottom.toInt(),
|
||||
),
|
||||
)
|
||||
session.setRepeatingRequest(builder.build(), null, mainHandler)
|
||||
} catch (e: Exception) {
|
||||
Log.w(TAG, "setZoom failed", e)
|
||||
}
|
||||
}
|
||||
|
||||
private fun stop() {
|
||||
running = false
|
||||
try {
|
||||
captureSession?.close()
|
||||
} catch (_: Exception) {
|
||||
}
|
||||
captureSession = null
|
||||
try {
|
||||
cameraDevice?.close()
|
||||
} catch (_: Exception) {
|
||||
}
|
||||
cameraDevice = null
|
||||
}
|
||||
|
||||
fun destroy() {
|
||||
stop()
|
||||
imageReader?.close()
|
||||
imageReader = null
|
||||
previewSurface?.release()
|
||||
previewSurface = null
|
||||
surfaceEntry?.release()
|
||||
surfaceEntry = null
|
||||
analysisThread.quitSafely()
|
||||
}
|
||||
|
||||
private fun pickBackCamera(): String? {
|
||||
val ids = cameraManager.cameraIdList
|
||||
for (id in ids) {
|
||||
val c = cameraManager.getCameraCharacteristics(id)
|
||||
val facing = c.get(CameraCharacteristics.LENS_FACING)
|
||||
if (facing == CameraCharacteristics.LENS_FACING_BACK) return id
|
||||
}
|
||||
return ids.firstOrNull()
|
||||
}
|
||||
|
||||
/**
|
||||
* 从设备流配置挑预览/分析共用分辨率(均为传感器方向尺寸,宽≥高):
|
||||
* 取两者交集,16:9 优先、长边 ≤1920 内取最大(推理输入 640x640,
|
||||
* 更高只增帧传输与旋转开销,无精度收益);无 16:9 时退回最大交集尺寸。
|
||||
*/
|
||||
private fun pickSize(previewSizes: Array<Size>, analysisSizes: Array<Size>): Size {
|
||||
val common = previewSizes.filter { analysisSizes.contains(it) }
|
||||
val fallback = common.maxByOrNull { it.width.toLong() * it.height }
|
||||
?: return Size(1920, 1080)
|
||||
return common
|
||||
.filter {
|
||||
Math.abs(it.width.toDouble() / it.height - 16.0 / 9.0) < 0.03 &&
|
||||
it.width <= 1920 && it.height <= 1920
|
||||
}
|
||||
.maxByOrNull { it.width.toLong() * it.height }
|
||||
?: fallback
|
||||
}
|
||||
|
||||
/**
|
||||
* 显示旋转角(度)。Display.ROTATION_* 在 API 36 是受限常量,直接用其值(0/1/2/3)。
|
||||
* API 30+ 用 activity.display(现代 API,实测可靠跟踪旋转);
|
||||
* 旧版 windowManager.defaultDisplay.rotation 已废弃,实测在部分设备恒返回 0。
|
||||
*/
|
||||
private val displayDegrees: Int
|
||||
get() = try {
|
||||
val rot: Int = if (Build.VERSION.SDK_INT >= 30) {
|
||||
activity.display?.rotation ?: 0
|
||||
} else {
|
||||
@Suppress("DEPRECATION")
|
||||
activity.windowManager.defaultDisplay.rotation
|
||||
}
|
||||
when (rot) {
|
||||
1 -> 90
|
||||
2 -> 180
|
||||
3 -> 270
|
||||
else -> 0
|
||||
}
|
||||
} catch (e: Exception) {
|
||||
Log.w(TAG, "display rotation query failed", e)
|
||||
0
|
||||
}
|
||||
|
||||
/** 传感器 → 竖屏显示的顺时针旋转角(后摄无镜像) */
|
||||
private val rotation: Int
|
||||
get() = ((sensorOrientation - displayDegrees) % 360 + 360) % 360
|
||||
|
||||
private val frameListener = ImageReader.OnImageAvailableListener { reader ->
|
||||
// 诊断计数:回调是否到达(放在一切判定之前,任何丢弃都先计数)
|
||||
frameCallbacks++
|
||||
if (frameCallbacks % 50 == 0L) {
|
||||
Log.i(TAG, "frames cb=$frameCallbacks ok=$emitOk err=$emitErr sinkNull=$sinkNullCount")
|
||||
}
|
||||
val image = reader.acquireLatestImage() ?: return@OnImageAvailableListener
|
||||
try {
|
||||
val sink = eventSink
|
||||
if (sink == null) {
|
||||
// 诊断:有帧但无订阅者(Dart 侧订阅未建立/被取消)
|
||||
sinkNullCount++
|
||||
errorDescription = "帧监听无订阅者"
|
||||
return@OnImageAvailableListener
|
||||
}
|
||||
if (!running) return@OnImageAvailableListener
|
||||
val now = SystemClock.elapsedRealtime()
|
||||
// 节流:推理 ~100ms 一帧,避免数 MB 帧无谓传输
|
||||
if (now - lastEmitMs < 100) return@OnImageAvailableListener
|
||||
lastEmitMs = now
|
||||
val bytes = yuvToRgba(image, rotation)
|
||||
// 缩小后再发:8.3MB/帧对主线程编码与通道传输都过重,540x960 足够
|
||||
// (推理输入 704,坐标归一化,映射不受分辨率影响)
|
||||
val scaled = downscaleToFit(
|
||||
bytes[0] as Int, bytes[1] as Int, bytes[2] as ByteArray, 960,
|
||||
)
|
||||
emitW = scaled[0] as Int
|
||||
emitH = scaled[1] as Int
|
||||
if (emitOk == 0L) Log.i(TAG, "first emit ${emitW}x${emitH}")
|
||||
// 关键:EventSink.success 内部 FlutterJNI.dispatchPlatformMessage 强制主线程
|
||||
// (ensureRunningOnMainThread 抛 RuntimeException),后台线程直接调必失败——
|
||||
// 曾因此每帧抛"必须主线程"异常、Dart 侧永远 流:0。必须 post 到主线程发送
|
||||
mainHandler.post {
|
||||
try {
|
||||
sink.success(listOf(scaled[0], scaled[1], 0, false, scaled[2]))
|
||||
emitOk++
|
||||
} catch (e: Exception) {
|
||||
emitErr++
|
||||
Log.e(TAG, "emit frame failed", e)
|
||||
errorDescription = "帧异常: ${e.javaClass.simpleName}: ${e.message}"
|
||||
}
|
||||
}
|
||||
} catch (e: Exception) {
|
||||
emitErr++
|
||||
Log.e(TAG, "emit frame failed", e)
|
||||
errorDescription = "帧异常: ${e.javaClass.simpleName}: ${e.message}"
|
||||
} finally {
|
||||
image.close()
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* YUV_420_888 三平面 → RGBA 单平面,同时顺时针旋转成竖屏。
|
||||
* 返回 [width, height, bytes]。整数 BT.601 转换;U/V 兼容 planar
|
||||
* (pixelStride=1)与 semi-planar(pixelStride=2,NV21 式交错)布局。
|
||||
* dst(x', y') = src(y, W-1-x')(90° 顺时针),逐像素处理 rowStride 填充。
|
||||
*/
|
||||
private fun yuvToRgba(image: Image, deg: Int): List<Any> {
|
||||
val yPlane = image.planes[0]
|
||||
val uPlane = image.planes[1]
|
||||
val vPlane = image.planes[2]
|
||||
val srcW = image.width
|
||||
val srcH = image.height
|
||||
val yStride = yPlane.rowStride
|
||||
val uStride = uPlane.rowStride
|
||||
val vStride = vPlane.rowStride
|
||||
val uPixel = uPlane.pixelStride
|
||||
val vPixel = vPlane.pixelStride
|
||||
val yBuf = yPlane.buffer
|
||||
val uBuf = uPlane.buffer
|
||||
val vBuf = vPlane.buffer
|
||||
val dstW = if (deg == 90 || deg == 270) srcH else srcW
|
||||
val dstH = if (deg == 90 || deg == 270) srcW else srcH
|
||||
val dst = ByteArray(dstW * dstH * 4)
|
||||
for (dy in 0 until dstH) {
|
||||
val rowBase = dy * dstW * 4
|
||||
for (dx in 0 until dstW) {
|
||||
val sx: Int
|
||||
val sy: Int
|
||||
when (deg) {
|
||||
90 -> {
|
||||
sx = dy
|
||||
sy = srcH - 1 - dx
|
||||
}
|
||||
180 -> {
|
||||
sx = srcW - 1 - dx
|
||||
sy = srcH - 1 - dy
|
||||
}
|
||||
270 -> {
|
||||
sx = srcW - 1 - dy
|
||||
sy = dx
|
||||
}
|
||||
else -> {
|
||||
sx = dx
|
||||
sy = dy
|
||||
}
|
||||
}
|
||||
val yy = (yBuf.get(sy * yStride + sx).toInt() and 0xFF) - 16
|
||||
val uu = (uBuf.get((sy / 2) * uStride + (sx / 2) * uPixel).toInt() and 0xFF) - 128
|
||||
val vv = (vBuf.get((sy / 2) * vStride + (sx / 2) * vPixel).toInt() and 0xFF) - 128
|
||||
val r = ((298 * yy + 409 * vv + 128) shr 8).coerceIn(0, 255)
|
||||
val g = ((298 * yy - 100 * uu - 208 * vv + 128) shr 8).coerceIn(0, 255)
|
||||
val b = ((298 * yy + 516 * uu + 128) shr 8).coerceIn(0, 255)
|
||||
val di = rowBase + dx * 4
|
||||
dst[di] = r.toByte()
|
||||
dst[di + 1] = g.toByte()
|
||||
dst[di + 2] = b.toByte()
|
||||
dst[di + 3] = 0xFF.toByte()
|
||||
}
|
||||
}
|
||||
return listOf(dstW, dstH, dst)
|
||||
}
|
||||
|
||||
/**
|
||||
* RGBA 帧整数倍缩小,长边 ≤ maxLong 时原样返回(f=1)。
|
||||
* 2x2 盒式平均(对 640x640 推理输入足够;比最近邻平滑,颜色更准)。
|
||||
*/
|
||||
private fun downscaleToFit(w: Int, h: Int, rgba: ByteArray, maxLong: Int): List<Any> {
|
||||
val long = maxOf(w, h)
|
||||
if (long <= maxLong) return listOf(w, h, rgba)
|
||||
val f = (long + maxLong - 1) / maxLong
|
||||
val dw = w / f
|
||||
val dh = h / f
|
||||
val out = ByteArray(dw * dh * 4)
|
||||
for (dy in 0 until dh) {
|
||||
val y0 = dy * f
|
||||
val y1 = minOf(y0 + f, h)
|
||||
val rows = y1 - y0
|
||||
for (dx in 0 until dw) {
|
||||
val x0 = dx * f
|
||||
val x1 = minOf(x0 + f, w)
|
||||
var r = 0L
|
||||
var g = 0L
|
||||
var b = 0L
|
||||
var a = 0L
|
||||
for (sy in y0 until y1) {
|
||||
var si = sy * w * 4 + x0 * 4
|
||||
for (sx in x0 until x1) {
|
||||
r += rgba[si].toInt() and 0xFF
|
||||
g += rgba[si + 1].toInt() and 0xFF
|
||||
b += rgba[si + 2].toInt() and 0xFF
|
||||
a += rgba[si + 3].toInt() and 0xFF
|
||||
si += 4
|
||||
}
|
||||
}
|
||||
val n = rows * (x1 - x0)
|
||||
val di = (dy * dw + dx) * 4
|
||||
out[di] = (r / n).toByte()
|
||||
out[di + 1] = (g / n).toByte()
|
||||
out[di + 2] = (b / n).toByte()
|
||||
out[di + 3] = (a / n).toByte()
|
||||
}
|
||||
}
|
||||
return listOf(dw, dh, out)
|
||||
}
|
||||
|
||||
companion object {
|
||||
private const val TAG = "CameraChannel"
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,139 @@
|
||||
package com.example.observer
|
||||
|
||||
import android.app.PendingIntent
|
||||
import android.content.Intent
|
||||
import android.content.pm.PackageInstaller
|
||||
import android.net.Uri
|
||||
import android.os.Build
|
||||
import android.os.Handler
|
||||
import android.os.Looper
|
||||
import android.provider.Settings
|
||||
import java.io.File
|
||||
import io.flutter.embedding.android.FlutterActivity
|
||||
import io.flutter.embedding.engine.FlutterEngine
|
||||
import io.flutter.plugin.common.EventChannel
|
||||
import io.flutter.plugin.common.MethodCall
|
||||
import io.flutter.plugin.common.MethodChannel
|
||||
|
||||
/**
|
||||
* 原生 APK 安装通道(App 内更新安装):PackageInstaller 会话安装,
|
||||
* 安装进度经 EventChannel 实时回传 Flutter(下载进度由 Flutter 侧 http 流式下载自算)。
|
||||
*
|
||||
* 通道:
|
||||
* - MethodChannel "observer/installer":install(path)
|
||||
* - result.success("installing"):已提交安装
|
||||
* - result.success("permission_required"):未开「安装未知应用」,已拉起系统设置页
|
||||
* - EventChannel "observer/installer/progress":{event: progress/finished/failed, ...}
|
||||
*/
|
||||
class InstallerChannel(
|
||||
private val activity: FlutterActivity,
|
||||
private val engine: FlutterEngine,
|
||||
) {
|
||||
private val messenger = engine.dartExecutor.binaryMessenger
|
||||
private val method = MethodChannel(messenger, "observer/installer")
|
||||
private val progress = EventChannel(messenger, "observer/installer/progress")
|
||||
|
||||
private var progressSink: EventChannel.EventSink? = null
|
||||
private var activeSession: PackageInstaller.Session? = null
|
||||
|
||||
fun register() {
|
||||
method.setMethodCallHandler { call: MethodCall, result: MethodChannel.Result ->
|
||||
when (call.method) {
|
||||
"install" -> install(call.argument<String>("path") ?: "", result)
|
||||
else -> result.notImplemented()
|
||||
}
|
||||
}
|
||||
progress.setStreamHandler(object : EventChannel.StreamHandler {
|
||||
override fun onListen(arguments: Any?, events: EventChannel.EventSink) {
|
||||
progressSink = events
|
||||
}
|
||||
|
||||
override fun onCancel(arguments: Any?) {
|
||||
progressSink = null
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
private fun install(path: String, result: MethodChannel.Result) {
|
||||
val file = File(path)
|
||||
if (!file.exists()) {
|
||||
result.error("FILE_NOT_FOUND", "APK 文件不存在: $path", null)
|
||||
return
|
||||
}
|
||||
if (Build.VERSION.SDK_INT >= Build.VERSION_CODES.O &&
|
||||
!activity.packageManager.canRequestPackageInstalls()
|
||||
) {
|
||||
val intent = Intent(
|
||||
Settings.ACTION_MANAGE_UNKNOWN_APP_SOURCES,
|
||||
Uri.parse("package:${activity.packageName}"),
|
||||
).addFlags(Intent.FLAG_ACTIVITY_NEW_TASK)
|
||||
activity.startActivity(intent)
|
||||
result.success("permission_required")
|
||||
return
|
||||
}
|
||||
try {
|
||||
val pm = activity.packageManager
|
||||
val params = PackageInstaller.SessionParams(PackageInstaller.SessionParams.MODE_FULL_INSTALL).apply {
|
||||
setSize(file.length())
|
||||
}
|
||||
val sessionId = pm.packageInstaller.createSession(params)
|
||||
val session = pm.packageInstaller.openSession(sessionId)
|
||||
activeSession = session
|
||||
// API 36 起 registerSessionCallback(int, ...) 变体被移除,只剩全局注册形式,
|
||||
// 回调按 sessionId 过滤,避免响应其他会话事件
|
||||
val callback = object : PackageInstaller.SessionCallback() {
|
||||
override fun onCreated(id: Int) {}
|
||||
|
||||
override fun onBadgingChanged(id: Int) {}
|
||||
|
||||
override fun onActiveChanged(id: Int, active: Boolean) {}
|
||||
|
||||
override fun onProgressChanged(id: Int, progressPercent: Float) {
|
||||
if (id != sessionId) return
|
||||
emit("progress", "progress" to progressPercent.toInt())
|
||||
}
|
||||
|
||||
override fun onFinished(id: Int, success: Boolean) {
|
||||
if (id != sessionId) return
|
||||
emit("finished", "success" to success)
|
||||
pm.packageInstaller.unregisterSessionCallback(this)
|
||||
activeSession = null
|
||||
}
|
||||
}
|
||||
pm.packageInstaller.registerSessionCallback(callback, Handler(Looper.getMainLooper()))
|
||||
// 写 APK 到会话:1MB 缓冲流式拷贝,完成后 commit 弹系统确认框
|
||||
Thread {
|
||||
try {
|
||||
session.openWrite("apk", 0, file.length()).use { out ->
|
||||
file.inputStream().use { input ->
|
||||
val buf = ByteArray(1 shl 20)
|
||||
while (true) {
|
||||
val n = input.read(buf)
|
||||
if (n < 0) break
|
||||
out.write(buf, 0, n)
|
||||
}
|
||||
}
|
||||
}
|
||||
val sender = PendingIntent.getActivity(
|
||||
activity,
|
||||
0,
|
||||
Intent(activity, MainActivity::class.java),
|
||||
PendingIntent.FLAG_UPDATE_CURRENT or PendingIntent.FLAG_IMMUTABLE,
|
||||
).intentSender
|
||||
session.commit(sender)
|
||||
} catch (e: Exception) {
|
||||
session.abandon()
|
||||
emit("failed", "error" to (e.message ?: "安装会话写入失败"))
|
||||
activeSession = null
|
||||
}
|
||||
}.start()
|
||||
result.success("installing")
|
||||
} catch (e: Exception) {
|
||||
result.error("INSTALL_FAILED", e.message, null)
|
||||
}
|
||||
}
|
||||
|
||||
private fun emit(event: String, vararg pairs: Pair<String, Any>) {
|
||||
progressSink?.success(mapOf("event" to event, *pairs))
|
||||
}
|
||||
}
|
||||
@@ -1,5 +1,25 @@
|
||||
package com.example.observer
|
||||
|
||||
import io.flutter.embedding.android.FlutterActivity
|
||||
import io.flutter.embedding.engine.FlutterEngine
|
||||
|
||||
class MainActivity : FlutterActivity()
|
||||
class MainActivity : FlutterActivity() {
|
||||
private var cameraChannel: CameraChannel? = null
|
||||
private var installerChannel: InstallerChannel? = null
|
||||
|
||||
override fun configureFlutterEngine(flutterEngine: FlutterEngine) {
|
||||
super.configureFlutterEngine(flutterEngine)
|
||||
// 自写原生相机(替代 camera_android_camerax 插件):
|
||||
// 分析帧在原生侧旋转成竖屏后回传,Flutter 侧恒 rotation=0(对齐 iOS)。
|
||||
// configureFlutterEngine(onCreate 阶段)只注册通道与 viewFactory,
|
||||
// 实际 bindToLifecycle 由 Flutter 相机页 start 时触发(此时已 RESUMED)
|
||||
cameraChannel = CameraChannel(this, flutterEngine).also { it.register() }
|
||||
// App 内更新安装 APK(PackageInstaller 会话安装 + 进度回传)
|
||||
installerChannel = InstallerChannel(this, flutterEngine).also { it.register() }
|
||||
}
|
||||
|
||||
override fun onDestroy() {
|
||||
cameraChannel?.destroy()
|
||||
super.onDestroy()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -101,3 +101,21 @@
|
||||
- `priceCents` 为整数分,客户端展示 ÷100 转元
|
||||
- 展示名由客户端按 `days` 派生「N天」,接口无 label 字段
|
||||
- 后端套餐来自 `config.yml` `plans` 节点(静态定价,改价改配置重启生效)
|
||||
|
||||
## 7. 版本更新检查
|
||||
|
||||
`GET /api/v1/app/update`(**公开接口,无需登录**)在 App 启动时调用(协议层见 `lib/update/update_checker.dart`)。
|
||||
|
||||
响应 `data`(无记录时字段为空串,视为无需更新):
|
||||
|
||||
```json
|
||||
{
|
||||
"version": "1.1.0",
|
||||
"notes": "修复识别准确率问题"
|
||||
}
|
||||
```
|
||||
|
||||
- **仅 Android 调用**:`UpdateChecker.fetch()` 在非 Android 平台直接返回空(iOS 不做版本下发,用户从 App Store 自行更新)
|
||||
- 客户端以「语义化版本号」按数字段比较(`1.10.0 > 1.9.9`),**服务器版本 > 已装版本即强制更新**:弹全屏阻塞页(禁返回,仅「立即更新」),**App 内直接下载安装**(不跳浏览器):http 流式下载固定地址 `{apiBaseUrl}/download/observer-latest.apk`(后端静态托管,永远是最新 APK)到缓存目录(`.part` 原子落盘后改名,避免半包)→ 原生 `PackageInstaller` 会话安装(`InstallerChannel.kt`),页面实时显示**下载进度与安装进度**;首次安装需在系统设置允许「安装未知应用」(未开启时自动拉起系统设置页,APK 已缓存、返回后再次点击直达安装)
|
||||
- **双版本比较防反复提示**:用户点「立即更新」时把服务器版本号写入本地(`SessionStore.accepted_update_version`);判定条件是服务器版本 > 已装版本 **或** 服务器版本 > 已接受版本。APK 版本号不递增(每次打的包 versionName 相同)时,已更新完成的手机重启也不会再次弹更新
|
||||
- 网络异常/响应异常时静默跳过检查,不阻塞启动
|
||||
|
||||
@@ -57,11 +57,9 @@
|
||||
<dict>
|
||||
<key>NSAllowsArbitraryLoads</key>
|
||||
<true/>
|
||||
<key>NSAllowsArbitraryLoadsInWebContent</key>
|
||||
<true/>
|
||||
</dict>
|
||||
<key>NSCameraUsageDescription</key>
|
||||
<string>需要使用相机进行野生动物实时识别</string>
|
||||
<string>需要使用相机进行动物实时识别</string>
|
||||
<key>NSLocalNetworkUsageDescription</key>
|
||||
<string>需要通过本地网络连接服务器进行账号验证和支付</string>
|
||||
<key>UIApplicationSceneManifest</key>
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import 'package:flutter/material.dart';
|
||||
import 'package:package_info_plus/package_info_plus.dart';
|
||||
import 'package:provider/provider.dart';
|
||||
|
||||
import 'auth/auth_screen.dart';
|
||||
@@ -7,7 +8,10 @@ import 'camera/camera_screen.dart';
|
||||
import 'container.dart';
|
||||
import 'home/home_screen.dart';
|
||||
import 'legal/terms_screen.dart';
|
||||
import 'models/model_manager.dart';
|
||||
import 'payment/paywall_screen.dart';
|
||||
import 'update/update_checker.dart';
|
||||
import 'update/update_screen.dart';
|
||||
|
||||
class ObserverApp extends StatelessWidget {
|
||||
final AppContainer container;
|
||||
@@ -66,9 +70,33 @@ class _StartupGateState extends State<StartupGate> {
|
||||
Navigator.of(context).pushReplacementNamed('/terms');
|
||||
return;
|
||||
}
|
||||
// 版本更新检查(公开接口,无需登录态;仅 Android):服务器版本高于
|
||||
// 「已装版本与已确认接受版本」中的较大者即强制更新,弹全屏阻塞页。
|
||||
// APK 版本号不递增时,点过「立即更新」的 accepted 版本参与比较,
|
||||
// 更新完成后再次启动不反复提示。
|
||||
final info = await UpdateChecker().fetch();
|
||||
final current = await PackageInfo.fromPlatform();
|
||||
if (!mounted) return;
|
||||
final session = context.read<SessionStore>();
|
||||
final acceptedUpdate = await session.readAcceptedUpdateVersion();
|
||||
if (!mounted) return;
|
||||
if (UpdateChecker.needsUpdate(info.version, current.version, acceptedUpdate)) {
|
||||
Navigator.of(context).pushReplacement(MaterialPageRoute(
|
||||
builder: (_) => UpdateScreen(
|
||||
version: info.version,
|
||||
url: UpdateChecker.downloadUrl(),
|
||||
notes: info.notes,
|
||||
onUpdateAccepted: () =>
|
||||
session.saveAcceptedUpdateVersion(info.version),
|
||||
),
|
||||
));
|
||||
return;
|
||||
}
|
||||
final token = await session.readToken();
|
||||
if (!mounted) return;
|
||||
// 模型热更新:后台拉取模型目录并下载/更新各数据集模型(不阻塞启动,
|
||||
// 相机页打开前若未就绪会兜底等待;下载失败回退内置资产模型)
|
||||
ModelManager.instance.refresh();
|
||||
Navigator.of(context)
|
||||
.pushReplacementNamed(token == null ? '/login' : '/home');
|
||||
}
|
||||
|
||||
@@ -67,7 +67,7 @@ class _AuthScreenState extends State<AuthScreen> {
|
||||
),
|
||||
const SizedBox(height: 4),
|
||||
Text(
|
||||
'野生动物实时识别',
|
||||
'动物实时识别',
|
||||
textAlign: TextAlign.center,
|
||||
style: TextStyle(color: Colors.grey.shade600),
|
||||
),
|
||||
|
||||
@@ -2,13 +2,17 @@ import 'package:flutter_secure_storage/flutter_secure_storage.dart';
|
||||
|
||||
/// 登录会话持久化:token + 手机号存 secure storage。
|
||||
/// 启动时读取判断是否已登录;登出/401 时清除回登录页。
|
||||
/// 另存「已确认更新版本」:用户点过「立即更新」后记录服务器版本号,
|
||||
/// 与 APK 内 versionName 取较大者参与更新判断(APK 版本号不递增也不会反复提示)。
|
||||
class SessionStore {
|
||||
static const _storage = FlutterSecureStorage();
|
||||
static const _tokenKey = 'auth_token';
|
||||
static const _phoneKey = 'auth_phone';
|
||||
static const _acceptedUpdateKey = 'accepted_update_version';
|
||||
|
||||
static String? _cachedToken;
|
||||
static String? _cachedPhone;
|
||||
static String? _cachedAcceptedUpdate;
|
||||
|
||||
Future<String?> readToken() async {
|
||||
if (_cachedToken != null) return _cachedToken;
|
||||
@@ -20,6 +24,18 @@ class SessionStore {
|
||||
return _cachedPhone = await _storage.read(key: _phoneKey);
|
||||
}
|
||||
|
||||
/// 用户点过「立即更新」的服务器版本号(空串 = 从未接受过更新提示)
|
||||
Future<String> readAcceptedUpdateVersion() async {
|
||||
if (_cachedAcceptedUpdate != null) return _cachedAcceptedUpdate!;
|
||||
return _cachedAcceptedUpdate =
|
||||
(await _storage.read(key: _acceptedUpdateKey)) ?? '';
|
||||
}
|
||||
|
||||
Future<void> saveAcceptedUpdateVersion(String version) async {
|
||||
_cachedAcceptedUpdate = version;
|
||||
await _storage.write(key: _acceptedUpdateKey, value: version);
|
||||
}
|
||||
|
||||
Future<void> save(String phone, String token) async {
|
||||
_cachedPhone = phone;
|
||||
_cachedToken = token;
|
||||
|
||||
@@ -1,61 +1,257 @@
|
||||
import 'dart:async';
|
||||
|
||||
import 'package:camera/camera.dart';
|
||||
import 'package:flutter/foundation.dart';
|
||||
import 'package:flutter/services.dart';
|
||||
import 'package:flutter/widgets.dart';
|
||||
|
||||
import 'frame_analyzer.dart';
|
||||
|
||||
/// camera 插件封装:后摄图像流(对应 Kotlin CameraController)。
|
||||
class AppCameraController {
|
||||
/// 相机抽象(对应 Kotlin CameraController):
|
||||
///
|
||||
/// - Android:自写原生通道([NativeCameraController])。分析帧在 Kotlin 侧
|
||||
/// 旋转成竖屏后经 EventChannel 回传,Flutter 侧恒 rotation=0——与 iOS
|
||||
/// (camera_avfoundation 插件,帧本来就是竖屏方向)行为一致,消除
|
||||
/// "横屏传感器帧 → 90° 旋转 + FIT_COVER crop 映射"的标注偏移根因。
|
||||
/// - iOS:camera 插件(原逻辑)。帧已竖屏,rotation 恒 0。
|
||||
abstract class AppCameraController {
|
||||
static Future<AppCameraController?> create() async {
|
||||
if (defaultTargetPlatform == TargetPlatform.android) {
|
||||
return NativeCameraController();
|
||||
}
|
||||
return PluginCameraController.create();
|
||||
}
|
||||
|
||||
/// 分析流回调实际触发次数(诊断用,与 analyzer 帧计数区分)
|
||||
int get streamCallbacks;
|
||||
|
||||
bool get isInitialized;
|
||||
|
||||
bool get isStreaming;
|
||||
|
||||
String? get errorDescription;
|
||||
|
||||
/// 检测框 overlay 应使用的旋转角。两平台帧都已竖屏 → 恒 0。
|
||||
int get rotationDegrees => 0;
|
||||
|
||||
/// 诊断:传感器方向 / 显示旋转(仅 Android 原生通道上报)
|
||||
int? get sensorOrientation => null;
|
||||
|
||||
int? get displayDegrees => null;
|
||||
|
||||
/// 预览实际应用的旋转圈数(仅 Android 原生通道上报,诊断用)
|
||||
int get quarterTurns => -1;
|
||||
|
||||
/// 诊断:轮询原生侧帧状态(Android 返回计数;iOS 返回空)
|
||||
Future<Map<dynamic, dynamic>> stats() async => const {};
|
||||
|
||||
Future<void> start(FrameAnalyzer analyzer);
|
||||
|
||||
Future<void> stop();
|
||||
|
||||
Future<double> getMinZoomLevel();
|
||||
|
||||
Future<double> getMaxZoomLevel();
|
||||
|
||||
Future<void> setZoomLevel(double value);
|
||||
|
||||
/// 预览 widget:Android 为原生 SurfaceView(AndroidView),iOS 为插件纹理
|
||||
Widget buildPreview();
|
||||
}
|
||||
|
||||
/// Android:自写原生相机通道(Kotlin CameraChannel)。
|
||||
class NativeCameraController extends AppCameraController {
|
||||
static const MethodChannel _channel = MethodChannel('observer/camera');
|
||||
static const EventChannel _frames = EventChannel('observer/camera/frames');
|
||||
|
||||
StreamSubscription<dynamic>? _sub;
|
||||
bool _streaming = false;
|
||||
String? _error;
|
||||
|
||||
/// 原生侧注册的预览纹理(SurfaceTexture)
|
||||
int? _textureId;
|
||||
int _textureW = 1920;
|
||||
int _textureH = 1080;
|
||||
int _quarterTurns = 1;
|
||||
|
||||
@override
|
||||
int streamCallbacks = 0;
|
||||
|
||||
int? _sensorOrientation;
|
||||
int? _displayDegrees;
|
||||
|
||||
@override
|
||||
int? get sensorOrientation => _sensorOrientation;
|
||||
|
||||
@override
|
||||
int? get displayDegrees => _displayDegrees;
|
||||
|
||||
@override
|
||||
int get quarterTurns => _quarterTurns;
|
||||
|
||||
@override
|
||||
bool get isInitialized => _streaming;
|
||||
|
||||
@override
|
||||
bool get isStreaming => _streaming;
|
||||
|
||||
@override
|
||||
String? get errorDescription => _error;
|
||||
|
||||
@override
|
||||
Future<Map<dynamic, dynamic>> stats() async {
|
||||
try {
|
||||
final r = await _channel.invokeMethod<Map<dynamic, dynamic>>('stats');
|
||||
if (r != null) {
|
||||
// 旋转/显示角度随轮询实时刷新:手机旋转后预览与帧旋转都跟着变
|
||||
final dd = r['displayDegrees'];
|
||||
if (dd is int) _displayDegrees = dd;
|
||||
final turns = r['quarterTurns'];
|
||||
if (turns is int) _quarterTurns = turns;
|
||||
}
|
||||
return r ?? const {};
|
||||
} catch (e) {
|
||||
return {'pollErr': '$e'};
|
||||
}
|
||||
}
|
||||
|
||||
@override
|
||||
Future<void> start(FrameAnalyzer analyzer) async {
|
||||
await stop();
|
||||
// 先订阅再启动:原生 start 绑定后立刻推帧,避免首帧竞态
|
||||
_sub = _frames.receiveBroadcastStream().listen((event) {
|
||||
// 计数放最前:任何到达的事件都先记账,后续解析失败也不丢计数
|
||||
streamCallbacks++;
|
||||
try {
|
||||
final f = event as List;
|
||||
final w = f[0] as int;
|
||||
final h = f[1] as int;
|
||||
final rotation = f[2] as int;
|
||||
final bgra = f[3] as bool;
|
||||
final bytes = f[4] as Uint8List;
|
||||
analyzer.analyzeRaw(
|
||||
planes: [bytes],
|
||||
strides: [w * 4],
|
||||
width: w,
|
||||
height: h,
|
||||
isBgra: true,
|
||||
rgbaOrder: !bgra,
|
||||
rotationDegrees: rotation,
|
||||
);
|
||||
} catch (e) {
|
||||
_error = '帧解析: $e';
|
||||
analyzer.recordStreamError('frames parse: $e');
|
||||
}
|
||||
}, onError: (Object e) {
|
||||
_error = '$e';
|
||||
analyzer.recordStreamError('frames: $e');
|
||||
});
|
||||
analyzer.reset();
|
||||
analyzer.worker?.reset();
|
||||
try {
|
||||
final r = await _channel.invokeMethod<Map<dynamic, dynamic>>('start');
|
||||
_textureId = r?['textureId'] as int?;
|
||||
_textureW = (r?['w'] as num?)?.toInt() ?? _textureW;
|
||||
_textureH = (r?['h'] as num?)?.toInt() ?? _textureH;
|
||||
_quarterTurns = (r?['quarterTurns'] as num?)?.toInt() ?? 1;
|
||||
_sensorOrientation = (r?['sensorOrientation'] as num?)?.toInt();
|
||||
_displayDegrees = (r?['displayDegrees'] as num?)?.toInt();
|
||||
} catch (e) {
|
||||
_error = '$e';
|
||||
analyzer.recordStreamError('camera start: $e');
|
||||
rethrow;
|
||||
}
|
||||
_streaming = true;
|
||||
}
|
||||
|
||||
@override
|
||||
Future<void> stop() async {
|
||||
_streaming = false;
|
||||
await _sub?.cancel();
|
||||
_sub = null;
|
||||
try {
|
||||
await _channel.invokeMethod('stop');
|
||||
} catch (_) {}
|
||||
}
|
||||
|
||||
@override
|
||||
Future<double> getMinZoomLevel() async {
|
||||
final r = await _channel.invokeMethod<List<dynamic>>('getZoomRange');
|
||||
return (r != null && r.isNotEmpty ? (r[0] as num).toDouble() : 1.0);
|
||||
}
|
||||
|
||||
@override
|
||||
Future<double> getMaxZoomLevel() async {
|
||||
final r = await _channel.invokeMethod<List<dynamic>>('getZoomRange');
|
||||
return (r != null && r.length > 1 ? (r[1] as num).toDouble() : 1.0);
|
||||
}
|
||||
|
||||
@override
|
||||
Future<void> setZoomLevel(double value) async {
|
||||
try {
|
||||
await _channel.invokeMethod('setZoom', value);
|
||||
} catch (_) {}
|
||||
}
|
||||
|
||||
/// 预览纹理:Flutter 引擎渲染 Texture 时已自动应用 SurfaceTexture 变换矩阵
|
||||
/// (传感器方向补偿,内容已转成自然方向的竖屏),因此:
|
||||
/// 1. 区域声明旋转后的尺寸(宽高互换)——否则横屏区域会把竖屏内容横向拉伸
|
||||
/// 2. RotatedBox 只按显示旋转补偿(quarterTurns = 屏转/90,竖屏 0 / 横屏 1)
|
||||
/// 再 FittedBox cover(= CoordinateMapper 的 FIT_COVER 数学一致)填满全屏。
|
||||
/// 不用 AndroidView+SurfaceView——平台视图会盖住 Flutter UI(诊断行/设置按钮/overlay)
|
||||
@override
|
||||
Widget buildPreview() {
|
||||
final id = _textureId;
|
||||
if (id == null) return const SizedBox.shrink();
|
||||
return FittedBox(
|
||||
fit: BoxFit.cover,
|
||||
child: RotatedBox(
|
||||
quarterTurns: _quarterTurns % 4,
|
||||
child: SizedBox(
|
||||
width: _textureH.toDouble(),
|
||||
height: _textureW.toDouble(),
|
||||
child: Texture(textureId: id),
|
||||
),
|
||||
),
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
/// iOS:camera 插件封装(原实现)。
|
||||
class PluginCameraController extends AppCameraController {
|
||||
final List<CameraDescription> cameras;
|
||||
CameraController? controller;
|
||||
|
||||
/// 图像流回调实际触发次数(诊断用,与 analyzer 帧计数区分)
|
||||
@override
|
||||
int streamCallbacks = 0;
|
||||
|
||||
AppCameraController._(this.cameras);
|
||||
PluginCameraController._(this.cameras);
|
||||
|
||||
static Future<AppCameraController?> create() async {
|
||||
static Future<PluginCameraController?> create() async {
|
||||
final cameras = await availableCameras();
|
||||
if (cameras.isEmpty) return null;
|
||||
return AppCameraController._(cameras);
|
||||
return PluginCameraController._(cameras);
|
||||
}
|
||||
|
||||
@override
|
||||
bool get isInitialized => controller?.value.isInitialized ?? false;
|
||||
|
||||
CameraController get currentController =>
|
||||
controller ?? (throw StateError('camera not initialized'));
|
||||
@override
|
||||
bool get isStreaming => controller?.value.isStreamingImages ?? false;
|
||||
|
||||
/// 图像流送达时的旋转角(传感器 → 竖屏显示所需的顺时针旋转)。
|
||||
/// 与 CameraX rotationDegrees 同公式;预览本身由平台旋转,检测框 overlay
|
||||
/// 用同一角度映射即可对齐。
|
||||
int get rotationDegrees {
|
||||
final c = controller;
|
||||
if (c == null) return 0;
|
||||
final deviceDegrees = switch (c.value.deviceOrientation) {
|
||||
DeviceOrientation.portraitUp => 0,
|
||||
DeviceOrientation.landscapeLeft => 90,
|
||||
DeviceOrientation.portraitDown => 180,
|
||||
DeviceOrientation.landscapeRight => 270,
|
||||
};
|
||||
final sensor = c.description.sensorOrientation;
|
||||
final isFront =
|
||||
c.description.lensDirection == CameraLensDirection.front;
|
||||
final degrees = (isFront ? sensor + deviceDegrees : sensor - deviceDegrees) % 360;
|
||||
return degrees < 0 ? degrees + 360 : degrees;
|
||||
}
|
||||
@override
|
||||
String? get errorDescription => controller?.value.errorDescription;
|
||||
|
||||
@override
|
||||
Future<void> start(FrameAnalyzer analyzer) async {
|
||||
await stop();
|
||||
final desc = cameras.firstWhere(
|
||||
(c) => c.lensDirection == CameraLensDirection.back,
|
||||
orElse: () => cameras.first);
|
||||
// iOS 用默认 bgra8888(420v 在部分 iOS 版本上视频输出静默不送帧),
|
||||
// Android 用 yuv420 多平面。
|
||||
final fmt = defaultTargetPlatform == TargetPlatform.iOS
|
||||
? ImageFormatGroup.bgra8888
|
||||
: ImageFormatGroup.yuv420;
|
||||
final c = CameraController(desc, ResolutionPreset.high,
|
||||
enableAudio: false, imageFormatGroup: fmt);
|
||||
// iOS image stream 帧已按竖屏方向输出(无需旋转),
|
||||
// 与 Android 原生通道(帧原生侧旋转成竖屏)统一 rotation=0
|
||||
final c = CameraController(desc, ResolutionPreset.veryHigh,
|
||||
enableAudio: false, imageFormatGroup: ImageFormatGroup.bgra8888);
|
||||
controller = c;
|
||||
await c.initialize();
|
||||
// 相机(重新)启动后重置运动/背景参考与抽帧节流,避免旧场景残留
|
||||
@@ -66,7 +262,7 @@ class AppCameraController {
|
||||
await c.startImageStream((image) {
|
||||
streamCallbacks++;
|
||||
try {
|
||||
analyzer.analyze(image, rotationDegrees);
|
||||
analyzer.analyze(image, 0);
|
||||
} catch (e, st) {
|
||||
debugPrint('[camera] analyze error: $e\n$st');
|
||||
analyzer.recordStreamError('analyze: $e');
|
||||
@@ -80,6 +276,7 @@ class AppCameraController {
|
||||
}
|
||||
}
|
||||
|
||||
@override
|
||||
Future<void> stop() async {
|
||||
final c = controller;
|
||||
if (c == null) return;
|
||||
@@ -89,4 +286,16 @@ class AppCameraController {
|
||||
} catch (_) {}
|
||||
await c.dispose();
|
||||
}
|
||||
|
||||
@override
|
||||
Future<double> getMinZoomLevel() => controller!.getMinZoomLevel();
|
||||
|
||||
@override
|
||||
Future<double> getMaxZoomLevel() => controller!.getMaxZoomLevel();
|
||||
|
||||
@override
|
||||
Future<void> setZoomLevel(double value) => controller!.setZoomLevel(value);
|
||||
|
||||
@override
|
||||
Widget buildPreview() => CameraPreview(controller!);
|
||||
}
|
||||
|
||||
@@ -1,12 +1,12 @@
|
||||
import 'dart:ui' show PlatformDispatcher;
|
||||
import 'dart:async';
|
||||
import 'dart:ui' as ui show PlatformDispatcher;
|
||||
|
||||
import 'package:camera/camera.dart';
|
||||
import 'package:flutter/foundation.dart' show defaultTargetPlatform;
|
||||
import 'package:flutter/material.dart';
|
||||
import 'package:permission_handler/permission_handler.dart';
|
||||
import 'package:wakelock_plus/wakelock_plus.dart';
|
||||
|
||||
import '../detection/detector_worker.dart';
|
||||
import '../models/model_manager.dart';
|
||||
import '../reminder/reminder.dart';
|
||||
import 'app_camera_controller.dart';
|
||||
import 'camera_view_model.dart';
|
||||
@@ -30,17 +30,88 @@ class _CameraScreenState extends State<CameraScreen> {
|
||||
String? _globalError;
|
||||
String? _initError;
|
||||
|
||||
/// 置信度阈值(设置页滑块调整,worker 内实时生效)
|
||||
double _minScore = 0.10;
|
||||
|
||||
/// 原生侧帧状态轮询结果(诊断用;无帧时诊断行也能实时刷新)
|
||||
Map<dynamic, dynamic> _nativeStats = const {};
|
||||
Timer? _statsTimer;
|
||||
|
||||
void _openSettings() {
|
||||
final vm = _viewModel;
|
||||
if (vm == null) return;
|
||||
showModalBottomSheet<void>(
|
||||
context: context,
|
||||
backgroundColor: Colors.black87,
|
||||
builder: (ctx) => StatefulBuilder(
|
||||
builder: (ctx, setSheetState) => Padding(
|
||||
padding: const EdgeInsets.all(20),
|
||||
child: Column(
|
||||
mainAxisSize: MainAxisSize.min,
|
||||
crossAxisAlignment: CrossAxisAlignment.start,
|
||||
children: [
|
||||
const Text('识别设置',
|
||||
style: TextStyle(
|
||||
color: Colors.white, fontSize: 16, fontWeight: FontWeight.bold)),
|
||||
const SizedBox(height: 12),
|
||||
Row(
|
||||
children: [
|
||||
const Text('置信度阈值',
|
||||
style: TextStyle(color: Colors.white70, fontSize: 14)),
|
||||
const Spacer(),
|
||||
Text('${(_minScore * 100).toStringAsFixed(0)}%',
|
||||
style: const TextStyle(
|
||||
color: Colors.greenAccent,
|
||||
fontSize: 14,
|
||||
fontWeight: FontWeight.bold)),
|
||||
],
|
||||
),
|
||||
Slider(
|
||||
value: _minScore,
|
||||
min: 0.05,
|
||||
max: 0.50,
|
||||
divisions: 45,
|
||||
activeColor: Colors.greenAccent,
|
||||
onChanged: (v) {
|
||||
setSheetState(() => _minScore = v);
|
||||
_analyzer?.worker?.setMinScore(v);
|
||||
},
|
||||
),
|
||||
const SizedBox(height: 8),
|
||||
const Text(
|
||||
'阈值越低识别越灵敏(低分框越多,误报也可能增加);'
|
||||
'野鸡模型置信度普遍在 10%~20%,场景识别不到时可适当调低。',
|
||||
style: TextStyle(color: Colors.white54, fontSize: 12),
|
||||
),
|
||||
],
|
||||
),
|
||||
),
|
||||
),
|
||||
);
|
||||
}
|
||||
|
||||
@override
|
||||
void initState() {
|
||||
super.initState();
|
||||
final oldPlatform = PlatformDispatcher.instance.onError;
|
||||
PlatformDispatcher.instance.onError = (error, stack) {
|
||||
setState(() => _globalError = 'Platform: $error');
|
||||
final oldPlatform = ui.PlatformDispatcher.instance.onError;
|
||||
ui.PlatformDispatcher.instance.onError = (error, stack) {
|
||||
setState(() => _globalError =
|
||||
'Platform: $error\n${stack.toString().split('\n').take(3).join('\n')}');
|
||||
return oldPlatform?.call(error, stack) ?? false;
|
||||
};
|
||||
WidgetsBinding.instance.addPostFrameCallback((_) => _init());
|
||||
// 相机页常亮:野外观察时保持屏幕不熄(离开页面时关闭)
|
||||
WakelockPlus.enable();
|
||||
// 每秒轮询原生侧帧状态:无帧时诊断行也能实时刷新(camErr/计数)
|
||||
_statsTimer = Timer.periodic(const Duration(seconds: 1), (_) => _pollStats());
|
||||
}
|
||||
|
||||
Future<void> _pollStats() async {
|
||||
final camera = _cameraController;
|
||||
if (camera == null) return;
|
||||
final s = await camera.stats();
|
||||
if (!mounted) return;
|
||||
setState(() => _nativeStats = s);
|
||||
}
|
||||
|
||||
Future<void> _init() async {
|
||||
@@ -49,8 +120,18 @@ class _CameraScreenState extends State<CameraScreen> {
|
||||
setState(() => _permissionGranted = granted);
|
||||
if (!granted) return;
|
||||
|
||||
// 模型热更新:优先使用已下载的数据集模型(启动时后台拉取;此处兜底等待,
|
||||
// 下载慢/失败不阻塞相机启动——无下载模型时 worker 回退内置资产)
|
||||
if (!ModelManager.instance.ready) {
|
||||
try {
|
||||
await ModelManager.instance
|
||||
.refresh()
|
||||
.timeout(const Duration(seconds: 15));
|
||||
} catch (_) {}
|
||||
}
|
||||
// 模型加载/推理在后台 isolate,不阻塞 UI;worker 为 null 时仅预览并提示
|
||||
final worker = await DetectorWorker.create();
|
||||
final worker = await DetectorWorker.create(
|
||||
models: ModelManager.instance.models);
|
||||
final viewModel = CameraViewModel(reminder: Reminder());
|
||||
viewModel.setModelReady(worker != null);
|
||||
final analyzer = FrameAnalyzer(worker: worker, viewModel: viewModel);
|
||||
@@ -100,6 +181,7 @@ class _CameraScreenState extends State<CameraScreen> {
|
||||
|
||||
@override
|
||||
void dispose() {
|
||||
_statsTimer?.cancel();
|
||||
WakelockPlus.disable();
|
||||
_cameraController?.stop();
|
||||
_analyzer?.dispose();
|
||||
@@ -120,29 +202,24 @@ class _CameraScreenState extends State<CameraScreen> {
|
||||
if (!_permissionGranted)
|
||||
_PermissionGuide(onRequest: () => _init())
|
||||
else if (vm != null && (camera?.isInitialized ?? false))
|
||||
// 预览 + 检测框同几何:overlay 作为 CameraPreview 的 child,
|
||||
// 与纹理共享同一 Stack/尺寸,避免比例或裁剪导致的位置偏移
|
||||
// 预览 + 检测框同几何:overlay 作为预览 widget 的 sibling,
|
||||
// 与纹理共享同一 Stack/尺寸,避免比例或裁剪导致的位置偏移。
|
||||
// 两平台分析帧都已竖屏(Android 原生侧旋转 / iOS 插件本来就竖屏),
|
||||
// rotation 恒 0,CoordinateMapper 走纯 FIT_COVER 缩放路径
|
||||
ListenableBuilder(
|
||||
listenable: vm,
|
||||
builder: (context, _) => _ZoomablePreview(
|
||||
controller: camera!.currentController,
|
||||
imageWidthPx: vm.state.imageWidthPx,
|
||||
imageHeightPx: vm.state.imageHeightPx,
|
||||
controller: camera!,
|
||||
overlay: DetectionOverlay(
|
||||
results: vm.state.results,
|
||||
// iOS 纹理不旋转显示(_wrapInRotatedBox 仅 Android),
|
||||
// 显示方向 = buffer 原样 = 检测方向,旋转必须为 0;
|
||||
// Android 纹理被 RotatedBox 旋转,需用插件报告的 rotation。
|
||||
rotation: defaultTargetPlatform == TargetPlatform.iOS
|
||||
? 0
|
||||
: vm.state.rotation,
|
||||
rotation: 0,
|
||||
imageWidthPx: vm.state.imageWidthPx,
|
||||
imageHeightPx: vm.state.imageHeightPx,
|
||||
),
|
||||
),
|
||||
)
|
||||
else if (camera?.isInitialized ?? false)
|
||||
_ZoomablePreview(controller: camera!.currentController)
|
||||
_ZoomablePreview(controller: camera!)
|
||||
else
|
||||
const Center(
|
||||
child: Text('相机启动中…', style: TextStyle(color: Colors.white70)),
|
||||
@@ -197,6 +274,7 @@ class _CameraScreenState extends State<CameraScreen> {
|
||||
right: 0,
|
||||
child: _CameraTopBar(
|
||||
onClose: () => Navigator.of(context).pop(),
|
||||
onOpenSettings: _openSettings,
|
||||
),
|
||||
),
|
||||
],
|
||||
@@ -209,6 +287,27 @@ class _CameraScreenState extends State<CameraScreen> {
|
||||
return Stack(
|
||||
fit: StackFit.expand,
|
||||
children: [
|
||||
// 模型热更新下载失败提示(已加载模型仍可用,仅提示补更新)
|
||||
if (vm.state.modelReady && ModelManager.instance.error != null)
|
||||
Positioned(
|
||||
left: 16,
|
||||
right: 16,
|
||||
top: MediaQuery.of(context).padding.top + 56,
|
||||
child: Container(
|
||||
padding:
|
||||
const EdgeInsets.symmetric(horizontal: 12, vertical: 6),
|
||||
decoration: BoxDecoration(
|
||||
color: Colors.black54,
|
||||
borderRadius: BorderRadius.circular(8),
|
||||
),
|
||||
child: Text(
|
||||
ModelManager.instance.error!,
|
||||
textAlign: TextAlign.center,
|
||||
style: const TextStyle(color: Colors.orange, fontSize: 12),
|
||||
),
|
||||
),
|
||||
),
|
||||
|
||||
// 模型未加载时仅显示相机预览,不做检测标注(横幅置于顶栏下方,避免与底部诊断行重叠)
|
||||
if (!vm.state.modelReady)
|
||||
Positioned(
|
||||
@@ -238,13 +337,25 @@ class _CameraScreenState extends State<CameraScreen> {
|
||||
crossAxisAlignment: CrossAxisAlignment.center,
|
||||
children: [
|
||||
Text(
|
||||
'模型:${vm.state.modelReady ? '已加载' : '未加载'} 帧:${vm.state.framesReceived} 流:${camera?.streamCallbacks ?? 0} 推理:${vm.state.debugDetectCalls}次 异常:${vm.state.debugDetectErrors}次 处理:${vm.state.debugLastMs}ms 最高分:${(vm.state.debugHighestScore * 100).toStringAsFixed(1)}% 图:${vm.state.imageWidthPx}x${vm.state.imageHeightPx} 旋:${vm.state.rotation}',
|
||||
'阈值:${(_minScore * 100).toStringAsFixed(0)}% 模型:${vm.state.modelReady ? ModelManager.instance.modelsLabel : '未加载'} 帧:${vm.state.framesReceived} 流:${camera?.streamCallbacks ?? 0} 推理:${vm.state.debugDetectCalls}次 异常:${vm.state.debugDetectErrors}次 处理:${vm.state.debugLastMs}ms 最高分:${(vm.state.debugHighestScore * 100).toStringAsFixed(1)}% 图:${vm.state.imageWidthPx}x${vm.state.imageHeightPx} 传感:${camera?.sensorOrientation ?? '-'} 屏转:${camera?.displayDegrees ?? '-'} 旋:${camera?.rotationDegrees ?? 0} turn:${camera?.quarterTurns ?? '-'}',
|
||||
style: const TextStyle(color: Colors.white70, fontSize: 12),
|
||||
),
|
||||
if (_nativeStats.isNotEmpty)
|
||||
Text(
|
||||
'原生:回调${_nativeStats['callbacks'] ?? '-'} 发出${_nativeStats['emitOk'] ?? '-'} 异常${_nativeStats['emitErr'] ?? '-'} 无订阅${_nativeStats['sinkNull'] ?? '-'} sink:${_nativeStats['sink'] ?? '-'} 配置:${_nativeStats['size'] ?? '-'} 发帧:${_nativeStats['emitSize'] ?? '-'} 错误:${_nativeStats['error'] ?? '无'} 轮询:${_nativeStats['pollErr'] ?? 'ok'}${vm.state.results.isEmpty ? '' : ' 框1:(${vm.state.results.first.left.toStringAsFixed(2)},${vm.state.results.first.top.toStringAsFixed(2)},${vm.state.results.first.right.toStringAsFixed(2)},${vm.state.results.first.bottom.toStringAsFixed(2)})'}',
|
||||
style: const TextStyle(
|
||||
color: Colors.amberAccent, fontSize: 11),
|
||||
),
|
||||
if (vm.state.debugYuv.isNotEmpty)
|
||||
Text(
|
||||
'yuv:${vm.state.debugYuv}',
|
||||
style: const TextStyle(
|
||||
color: Colors.yellowAccent, fontSize: 11),
|
||||
),
|
||||
if (camera != null)
|
||||
Text(
|
||||
'streaming:${camera.currentController.value.isStreamingImages} '
|
||||
'camErr:${camera.currentController.value.errorDescription ?? '无'}',
|
||||
'streaming:${camera.isStreaming} '
|
||||
'camErr:${camera.errorDescription ?? '无'}',
|
||||
maxLines: 2,
|
||||
overflow: TextOverflow.ellipsis,
|
||||
style: const TextStyle(
|
||||
@@ -274,24 +385,19 @@ class _CameraScreenState extends State<CameraScreen> {
|
||||
],
|
||||
);
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
/// 双指捏合缩放预览;overlay 与纹理同几何(CameraPreview child)
|
||||
/// 双指捏合缩放预览;overlay 与纹理同几何(Stack 内同尺寸)
|
||||
class _ZoomablePreview extends StatefulWidget {
|
||||
final CameraController controller;
|
||||
final AppCameraController controller;
|
||||
|
||||
/// 检测框 overlay(随帧更新,作为 CameraPreview 的 child 与纹理同区域)
|
||||
/// 检测框 overlay(随帧更新,与纹理同区域)
|
||||
final Widget? overlay;
|
||||
|
||||
/// 当前帧图像尺寸(用于按 buffer 比例约束预览,保证无拉伸变形)
|
||||
final int imageWidthPx;
|
||||
final int imageHeightPx;
|
||||
|
||||
const _ZoomablePreview({
|
||||
required this.controller,
|
||||
this.overlay,
|
||||
this.imageWidthPx = 0,
|
||||
this.imageHeightPx = 0,
|
||||
});
|
||||
|
||||
@override
|
||||
@@ -317,7 +423,9 @@ class _ZoomablePreviewState extends State<_ZoomablePreview> {
|
||||
|
||||
@override
|
||||
Widget build(BuildContext context) {
|
||||
final preview = GestureDetector(
|
||||
// 预览铺满全屏(cover 裁剪由插件按视图比例完成);
|
||||
// overlay 与纹理同几何:作为 Stack sibling 叠在上层,坐标与纹理区域一致
|
||||
return GestureDetector(
|
||||
onScaleStart: (_) => _gestureStartZoom = _currentZoom,
|
||||
onScaleUpdate: (d) {
|
||||
final target =
|
||||
@@ -326,24 +434,24 @@ class _ZoomablePreviewState extends State<_ZoomablePreview> {
|
||||
_currentZoom = target;
|
||||
widget.controller.setZoomLevel(target);
|
||||
},
|
||||
child: CameraPreview(widget.controller, child: widget.overlay),
|
||||
);
|
||||
|
||||
final w = widget.imageWidthPx.toDouble();
|
||||
final h = widget.imageHeightPx.toDouble();
|
||||
if (w <= 0 || h <= 0) return preview;
|
||||
// 按 buffer 比例约束显示区域:纹理与 overlay 同区域等比显示(无变形)
|
||||
return Center(
|
||||
child: AspectRatio(aspectRatio: w / h, child: preview),
|
||||
child: Stack(
|
||||
fit: StackFit.expand,
|
||||
children: [
|
||||
widget.controller.buildPreview(),
|
||||
if (widget.overlay != null) widget.overlay!,
|
||||
],
|
||||
),
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
class _CameraTopBar extends StatelessWidget {
|
||||
final VoidCallback onClose;
|
||||
final VoidCallback onOpenSettings;
|
||||
|
||||
const _CameraTopBar({
|
||||
required this.onClose,
|
||||
required this.onOpenSettings,
|
||||
});
|
||||
|
||||
@override
|
||||
@@ -359,6 +467,11 @@ class _CameraTopBar extends StatelessWidget {
|
||||
icon: const Icon(Icons.arrow_back, color: Colors.white),
|
||||
onPressed: onClose,
|
||||
),
|
||||
IconButton(
|
||||
tooltip: '识别设置',
|
||||
icon: const Icon(Icons.tune, color: Colors.white70),
|
||||
onPressed: onOpenSettings,
|
||||
),
|
||||
],
|
||||
),
|
||||
);
|
||||
|
||||
@@ -4,7 +4,6 @@ import 'package:flutter/foundation.dart';
|
||||
|
||||
import '../detection/detection_result.dart';
|
||||
import '../detection/motion_aggregator.dart';
|
||||
import '../detection/tflite_detector.dart';
|
||||
import '../reminder/reminder.dart';
|
||||
|
||||
@immutable
|
||||
@@ -20,6 +19,7 @@ class CameraUiState {
|
||||
final String? debugLastError;
|
||||
final int framesReceived;
|
||||
final int debugLastMs;
|
||||
final String debugYuv;
|
||||
|
||||
const CameraUiState({
|
||||
this.modelReady = false,
|
||||
@@ -33,20 +33,19 @@ class CameraUiState {
|
||||
this.debugLastError,
|
||||
this.framesReceived = 0,
|
||||
this.debugLastMs = 0,
|
||||
this.debugYuv = '',
|
||||
});
|
||||
}
|
||||
|
||||
/// 检测结果置信度分级与轨迹确认。
|
||||
///
|
||||
/// - [lowConf](模型阈值 0.10):低于此分的框在检测阶段已丢弃。
|
||||
/// - [highConf](0.35):高于此分直接确认显示;真实野鸡多为 0.1~0.2,
|
||||
/// 高于 0.35 视为强证据。
|
||||
/// - 0.10~0.35 之间:需要多帧稳定([confirmFrames] 帧)或 活动证据
|
||||
/// - 低于 0.35 的框:需要多帧稳定([confirmFrames] 帧)或 活动证据
|
||||
/// (运动区域/背景新出现区域重叠)才确认显示。
|
||||
class CameraViewModel extends ChangeNotifier {
|
||||
static const int maxTracks = 30;
|
||||
static const double motionBoost = 0.15;
|
||||
static const double lowConf = TfliteDetector.minScore;
|
||||
static const double highConf = 0.35;
|
||||
static const int confirmFrames = 3;
|
||||
static const double associateRadius = 0.12;
|
||||
@@ -86,6 +85,7 @@ class CameraViewModel extends ChangeNotifier {
|
||||
String? lastError,
|
||||
int framesReceived = 0,
|
||||
int lastProcessMs = 0,
|
||||
String yuvDiag = '',
|
||||
}) {
|
||||
final now = DateTime.now().millisecondsSinceEpoch;
|
||||
_associate(results, motionRegions, noveltyRegions, now);
|
||||
@@ -132,6 +132,7 @@ class CameraViewModel extends ChangeNotifier {
|
||||
debugLastError: lastError,
|
||||
framesReceived: framesReceived,
|
||||
debugLastMs: lastProcessMs,
|
||||
debugYuv: yuvDiag.isNotEmpty ? yuvDiag : _state.debugYuv,
|
||||
);
|
||||
notifyListeners();
|
||||
}
|
||||
|
||||
@@ -87,10 +87,13 @@ class _OverlayPainter extends CustomPainter {
|
||||
_drawDashedRect(canvas, box, paint);
|
||||
}
|
||||
|
||||
// 标签:框上方,含距离
|
||||
// 标签:框上方,含距离;多模型时标注来源模型名(内置资产不标)
|
||||
final dist = _distanceLabel(r);
|
||||
final modelTag = r.modelName.isNotEmpty && r.modelName != '内置'
|
||||
? '[${r.modelName}]'
|
||||
: '';
|
||||
final text =
|
||||
'${_labels[r.label] ?? r.label} ${(r.score * 100).toInt()}%$dist';
|
||||
'${_labels[r.label] ?? r.label}$modelTag ${(r.score * 100).toInt()}%$dist';
|
||||
final textPainter = TextPainter(
|
||||
text: TextSpan(
|
||||
text: text,
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
import 'dart:typed_data';
|
||||
|
||||
import 'package:camera/camera.dart';
|
||||
|
||||
import '../detection/detection_result.dart';
|
||||
@@ -44,7 +46,8 @@ class FrameAnalyzer {
|
||||
int rotation,
|
||||
int width,
|
||||
int height,
|
||||
int processMs) {
|
||||
int processMs,
|
||||
String yuvDiag) {
|
||||
detectCalls++;
|
||||
lastProcessMs = processMs;
|
||||
viewModel.onFramesAnalyzed(
|
||||
@@ -59,6 +62,7 @@ class FrameAnalyzer {
|
||||
lastError: lastError,
|
||||
framesReceived: framesReceived,
|
||||
lastProcessMs: lastProcessMs,
|
||||
yuvDiag: yuvDiag,
|
||||
);
|
||||
}
|
||||
|
||||
@@ -80,15 +84,45 @@ class FrameAnalyzer {
|
||||
);
|
||||
}
|
||||
|
||||
void analyze(CameraImage image, int rotationDegrees) {
|
||||
/// 节流与 busy 丢帧判定(两入口共用);通过后才允许投递
|
||||
bool _canSend() {
|
||||
framesReceived++;
|
||||
final w = worker;
|
||||
if (w == null) return;
|
||||
if (w == null) return false;
|
||||
final now = DateTime.now().millisecondsSinceEpoch;
|
||||
if (now - _lastDetectMs < intervalMs) return;
|
||||
if (now - _lastDetectMs < intervalMs) return false;
|
||||
_lastDetectMs = now;
|
||||
if (w.busy) return; // 上一帧未返回则丢帧,避免在途积压
|
||||
w.analyze(image, rotationDegrees);
|
||||
if (w.busy) return false; // 上一帧未返回则丢帧,避免在途积压
|
||||
return true;
|
||||
}
|
||||
|
||||
void analyze(CameraImage image, int rotationDegrees,
|
||||
{bool rgbaOrder = false}) {
|
||||
if (!_canSend()) return;
|
||||
worker!.analyze(image, rotationDegrees, rgbaOrder: rgbaOrder);
|
||||
}
|
||||
|
||||
/// 原生相机通道帧(Android):字节已在 Kotlin 侧旋转成竖屏,rotation=0。
|
||||
/// isBgra=true + rgbaOrder 与插件路径同语义:false=BGRA(rOff=2)/true=RGBA(rOff=0)
|
||||
void analyzeRaw({
|
||||
required List<Uint8List> planes,
|
||||
required List<int> strides,
|
||||
required int width,
|
||||
required int height,
|
||||
required bool isBgra,
|
||||
required bool rgbaOrder,
|
||||
int rotationDegrees = 0,
|
||||
}) {
|
||||
if (!_canSend()) return;
|
||||
worker!.analyzeRaw(
|
||||
planes: planes,
|
||||
strides: strides,
|
||||
width: width,
|
||||
height: height,
|
||||
isBgra: isBgra,
|
||||
rgbaOrder: rgbaOrder,
|
||||
rotationDegrees: rotationDegrees,
|
||||
);
|
||||
}
|
||||
|
||||
void reset() => _lastDetectMs = 0;
|
||||
|
||||
@@ -12,7 +12,7 @@ class ViewRect {
|
||||
double get centerY => (top + bottom) / 2;
|
||||
}
|
||||
|
||||
/// 模型归一化坐标 → 预览视图坐标(含传感器旋转与 FIT_CENTER 裁剪)。
|
||||
/// 模型归一化坐标 → 预览视图坐标(含传感器旋转与 FIT_COVER 全屏裁剪)。
|
||||
class CoordinateMapper {
|
||||
static ViewRect mapToView(
|
||||
double normLeft,
|
||||
@@ -53,10 +53,10 @@ class CoordinateMapper {
|
||||
final portrait = rotation == 90 || rotation == 270;
|
||||
final portW = portrait ? imageH : imageW;
|
||||
final portH = portrait ? imageW : imageH;
|
||||
// 3) FIT_CENTER 缩放与居中偏移
|
||||
// 3) FIT_COVER 缩放与居中裁剪:放大到铺满视图,溢出部分裁掉
|
||||
final scale = viewW / portW < viewH / portH
|
||||
? viewW / portW
|
||||
: viewH / portH;
|
||||
? viewH / portH
|
||||
: viewW / portW;
|
||||
final offsetX = (viewW - portW * scale) / 2;
|
||||
final offsetY = (viewH - portH * scale) / 2;
|
||||
return ViewRect(
|
||||
|
||||
@@ -9,6 +9,10 @@ class DetectionResult {
|
||||
/// 轨迹已确认(多帧稳定/高分/活动确认),false = 候选,渲染为虚线
|
||||
final bool confirmed;
|
||||
|
||||
/// 产出该框的模型(数据集 id 与名称;内置资产模型为 -1/空)
|
||||
final int modelId;
|
||||
final String modelName;
|
||||
|
||||
const DetectionResult({
|
||||
required this.label,
|
||||
required this.score,
|
||||
@@ -17,6 +21,8 @@ class DetectionResult {
|
||||
required this.right,
|
||||
required this.bottom,
|
||||
this.confirmed = true,
|
||||
this.modelId = -1,
|
||||
this.modelName = '',
|
||||
});
|
||||
|
||||
double get width => right - left;
|
||||
@@ -40,6 +46,8 @@ class DetectionResult {
|
||||
right: right ?? this.right,
|
||||
bottom: bottom ?? this.bottom,
|
||||
confirmed: confirmed ?? this.confirmed,
|
||||
modelId: modelId,
|
||||
modelName: modelName,
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
@@ -7,17 +7,27 @@ import 'package:flutter/foundation.dart' show debugPrint;
|
||||
import 'package:flutter/services.dart' show rootBundle;
|
||||
|
||||
import '../camera/motion_detector.dart';
|
||||
import '../models/model_manager.dart';
|
||||
import 'background_model.dart';
|
||||
import 'detection_result.dart';
|
||||
import 'nms.dart';
|
||||
import 'tflite_detector.dart';
|
||||
import 'visual_prior.dart';
|
||||
|
||||
/// 推理工作单元:模型加载与检测全部在后台 isolate 执行,
|
||||
/// 主 isolate 只投递帧数据、接收结果,UI 不被推理阻塞(iOS 真机卡顿根因)。
|
||||
///
|
||||
/// 多模型并行推理:传入 [models](各数据集下载模型)后,每帧逐模型推理,
|
||||
/// 结果按类别分组跨模型 NMS 合并(同标签重复框取高分,不同标签互不压制);
|
||||
/// 无下载模型时回退内置资产模型。
|
||||
class DetectorWorker {
|
||||
static const String modelAsset = 'assets/model.tflite';
|
||||
static const String labelsAsset = 'assets/labels.txt';
|
||||
|
||||
/// 内置资产回退模型的标识
|
||||
static const int builtinModelId = -1;
|
||||
static const String builtinModelName = '内置';
|
||||
|
||||
final Isolate _isolate;
|
||||
final ReceivePort _responses;
|
||||
|
||||
@@ -30,9 +40,10 @@ class DetectorWorker {
|
||||
int _inFlight = 0;
|
||||
bool _dead = false;
|
||||
|
||||
/// 结果回调:结果 / 运动区域 / 新颖区域 / 旋转角 / 图宽 / 图高 / 处理耗时 ms
|
||||
/// 结果回调:结果 / 运动区域 / 新颖区域 / 旋转角 / 图宽 / 图高 /
|
||||
/// 处理耗时 ms / yuv 决策诊断串
|
||||
void Function(List<DetectionResult>, List<MotionRegion>, List<MotionRegion>,
|
||||
int, int, int, int)? onResult;
|
||||
int, int, int, int, String)? onResult;
|
||||
|
||||
/// 单帧处理异常回调(不影响相机流)
|
||||
void Function(String)? onError;
|
||||
@@ -53,16 +64,26 @@ class DetectorWorker {
|
||||
});
|
||||
}
|
||||
|
||||
/// 读取模型资产并启动后台推理 isolate;加载失败返回 null(App 降级为仅预览)。
|
||||
static Future<DetectorWorker?> create() async {
|
||||
/// 加载模型并启动后台推理 isolate;加载失败返回 null(App 降级为仅预览)。
|
||||
/// [models] 为空时回退内置资产模型(模型缺失同样返回 null)。
|
||||
static Future<DetectorWorker?> create({List<ModelBundle>? models}) async {
|
||||
try {
|
||||
final data = await rootBundle.load(modelAsset);
|
||||
final modelBytes =
|
||||
data.buffer.asUint8List(data.offsetInBytes, data.lengthInBytes);
|
||||
final labels = (await rootBundle.loadString(labelsAsset))
|
||||
.split('\n')
|
||||
.where((l) => l.trim().isNotEmpty)
|
||||
.toList();
|
||||
final payload = <List<Object?>>[];
|
||||
if (models != null && models.isNotEmpty) {
|
||||
for (final m in models) {
|
||||
payload.add(
|
||||
[m.bytes, m.labels, m.datasetId, m.datasetName]);
|
||||
}
|
||||
} else {
|
||||
final data = await rootBundle.load(modelAsset);
|
||||
final modelBytes =
|
||||
data.buffer.asUint8List(data.offsetInBytes, data.lengthInBytes);
|
||||
final labels = (await rootBundle.loadString(labelsAsset))
|
||||
.split('\n')
|
||||
.where((l) => l.trim().isNotEmpty)
|
||||
.toList();
|
||||
payload.add([modelBytes, labels, builtinModelId, builtinModelName]);
|
||||
}
|
||||
|
||||
final responses = ReceivePort();
|
||||
final isolate = await Isolate.spawn(_workerMain, responses.sendPort);
|
||||
@@ -72,7 +93,7 @@ class DetectorWorker {
|
||||
.timeout(const Duration(seconds: 10),
|
||||
onTimeout: () => throw TimeoutException('worker port timeout'));
|
||||
worker._port = port;
|
||||
port.send(['load', modelBytes, labels]);
|
||||
port.send(['load', payload]);
|
||||
await worker._ready.future
|
||||
.timeout(const Duration(seconds: 20), onTimeout: () {
|
||||
throw TimeoutException('model load timeout');
|
||||
@@ -88,20 +109,37 @@ class DetectorWorker {
|
||||
/// 是否忙(上一帧尚未返回):忙则丢帧,避免在途积压
|
||||
bool get busy => _inFlight > 0;
|
||||
|
||||
void analyze(CameraImage image, int rotationDegrees) {
|
||||
void analyze(CameraImage image, int rotationDegrees, {bool rgbaOrder = false}) {
|
||||
// 单平面 8888 判定:仅明确的 yuv420/nv21 走多平面 YUV 路径;
|
||||
// bgra8888 与 unknown(插件未识别 RGBA_8888 输出时)都按 4 字节像素处理
|
||||
final group = image.format.group;
|
||||
analyzeRaw(
|
||||
planes: image.planes.map((p) => p.bytes).toList(),
|
||||
strides: image.planes.map((p) => p.bytesPerRow).toList(),
|
||||
width: image.width,
|
||||
height: image.height,
|
||||
isBgra: group != ImageFormatGroup.yuv420 && group != ImageFormatGroup.nv21,
|
||||
rgbaOrder: rgbaOrder,
|
||||
rotationDegrees: rotationDegrees,
|
||||
);
|
||||
}
|
||||
|
||||
/// 原始字节帧投递(截屏注入用:toImage 的 RGBA 字节直接进检测,不经 CameraImage)
|
||||
void analyzeRaw({
|
||||
required List<Uint8List> planes,
|
||||
required List<int> strides,
|
||||
required int width,
|
||||
required int height,
|
||||
required bool isBgra,
|
||||
required bool rgbaOrder,
|
||||
required int rotationDegrees,
|
||||
}) {
|
||||
final port = _port;
|
||||
if (port == null || _dead) return;
|
||||
_inFlight++;
|
||||
port.send([
|
||||
'frame',
|
||||
[
|
||||
image.planes.map((p) => p.bytes).toList(),
|
||||
image.planes.map((p) => p.bytesPerRow).toList(),
|
||||
image.width,
|
||||
image.height,
|
||||
image.format.group == ImageFormatGroup.bgra8888,
|
||||
rotationDegrees,
|
||||
],
|
||||
[planes, strides, width, height, isBgra, rotationDegrees, rgbaOrder],
|
||||
]);
|
||||
}
|
||||
|
||||
@@ -129,6 +167,8 @@ class DetectorWorker {
|
||||
top: v[3] as double,
|
||||
right: v[4] as double,
|
||||
bottom: v[5] as double,
|
||||
modelId: v.length > 6 ? (v[6] as num).toInt() : -1,
|
||||
modelName: v.length > 7 ? v[7] as String : '',
|
||||
);
|
||||
}).toList();
|
||||
final motion = (list[5] as List)
|
||||
@@ -142,7 +182,7 @@ class DetectorWorker {
|
||||
v[0] as double, v[1] as double, v[2] as double, v[3] as double))
|
||||
.toList();
|
||||
onResult?.call(dets, motion, novelty, list[1] as int, list[2] as int,
|
||||
list[3] as int, list[7] as int);
|
||||
list[3] as int, list[7] as int, list[8] as String);
|
||||
break;
|
||||
case 'log':
|
||||
lastLog = list[1] as String;
|
||||
@@ -151,6 +191,7 @@ class DetectorWorker {
|
||||
case 'error':
|
||||
_inFlight--;
|
||||
onError?.call(list[1] as String);
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -161,6 +202,13 @@ class DetectorWorker {
|
||||
port.send(['reset']);
|
||||
}
|
||||
|
||||
/// 调整置信度阈值(设置页滑块,worker 内实时生效)
|
||||
void setMinScore(double v) {
|
||||
final port = _port;
|
||||
if (port == null || _dead) return;
|
||||
port.send(['set-min-score', v]);
|
||||
}
|
||||
|
||||
void dispose() {
|
||||
_dead = true;
|
||||
_isolate.kill(priority: Isolate.immediate);
|
||||
@@ -176,9 +224,11 @@ Future<void> _workerMain(SendPort mainPort) async {
|
||||
mainPort.send(['port', control.sendPort]);
|
||||
mainPort.send(['log', 'worker-start']);
|
||||
|
||||
TfliteDetector? detector;
|
||||
List<TfliteDetector> detectors = const [];
|
||||
MotionDetector? motion;
|
||||
BackgroundModel? background;
|
||||
var lastDualMs = 0; // 双字节序推理诊断节流
|
||||
Uint8List? prevY; // 上一帧 Y/RGBA 平面(帧间 diff 诊断)
|
||||
await for (final msg in control) {
|
||||
try {
|
||||
final list = msg as List;
|
||||
@@ -186,12 +236,35 @@ Future<void> _workerMain(SendPort mainPort) async {
|
||||
case 'load':
|
||||
mainPort.send(['log', 'load-received']);
|
||||
try {
|
||||
detector = await TfliteDetector.fromBuffer(
|
||||
list[1] as Uint8List, (list[2] as List).cast<String>());
|
||||
if (detector == null) {
|
||||
mainPort.send(['load-error', 'fromBuffer 返回 null']);
|
||||
// 多模型:逐模型加载,单个失败不阻塞其余;全部失败才报错
|
||||
final loaded = <TfliteDetector>[];
|
||||
final failures = <String>[];
|
||||
for (final entry in list[1] as List) {
|
||||
final e = entry as List;
|
||||
final name = e.length > 3 ? e[3] as String : '';
|
||||
final d = await TfliteDetector.fromBuffer(
|
||||
e[0] as Uint8List,
|
||||
(e[1] as List).cast<String>(),
|
||||
modelId: (e[2] as num).toInt(),
|
||||
modelName: name,
|
||||
);
|
||||
if (d == null) {
|
||||
failures.add(name.isEmpty ? 'unknown' : name);
|
||||
} else {
|
||||
loaded.add(d);
|
||||
}
|
||||
}
|
||||
if (loaded.isEmpty) {
|
||||
mainPort.send([
|
||||
'load-error',
|
||||
'模型加载失败:${failures.join(',')} '
|
||||
'(fromBuffer 返回 null)'
|
||||
]);
|
||||
} else {
|
||||
mainPort.send(['log', 'fromBuffer-ok']);
|
||||
detectors = loaded;
|
||||
mainPort.send(['log',
|
||||
'loaded=${loaded.map((d) => d.modelName).join(',')} '
|
||||
'failed=${failures.isEmpty ? '-' : failures.join(',')}']);
|
||||
motion = MotionDetector();
|
||||
background = BackgroundModel();
|
||||
mainPort.send(['ready']);
|
||||
@@ -201,10 +274,9 @@ Future<void> _workerMain(SendPort mainPort) async {
|
||||
}
|
||||
break;
|
||||
case 'frame':
|
||||
final d = detector;
|
||||
final m = motion;
|
||||
final b = background;
|
||||
if (d == null || m == null || b == null) break;
|
||||
if (detectors.isEmpty || m == null || b == null) break;
|
||||
final frame = list[1] as List;
|
||||
final planes = (frame[0] as List).cast<Uint8List>();
|
||||
final strides = (frame[1] as List).cast<int>();
|
||||
@@ -212,15 +284,111 @@ Future<void> _workerMain(SendPort mainPort) async {
|
||||
final height = frame[3] as int;
|
||||
final isBgra = frame[4] as bool;
|
||||
final rotation = frame[5] as int;
|
||||
final rgbaOrder = frame.length > 6 && (frame[6] as bool);
|
||||
|
||||
// 止血:非法帧(宽高/平面为空)直接丢弃并上报诊断,
|
||||
// 避免下游组件越界(RGBA patch 后插件偶发 w/h=0 帧)
|
||||
if (width <= 0 || height <= 0 || planes.isEmpty || planes[0].isEmpty) {
|
||||
mainPort.send([
|
||||
'error',
|
||||
'bad frame w=$width h=$height planes=${planes.length} '
|
||||
'p0=${planes.isNotEmpty ? planes[0].length : 0} '
|
||||
'stride=${strides.isNotEmpty ? strides[0] : '-'} '
|
||||
'bgra=$isBgra'
|
||||
]);
|
||||
break;
|
||||
}
|
||||
|
||||
final sw = Stopwatch()..start();
|
||||
var results = d.detectRaw(
|
||||
planes: planes,
|
||||
strides: strides,
|
||||
width: width,
|
||||
height: height,
|
||||
isBgra: isBgra,
|
||||
);
|
||||
// 帧内容统计(诊断):Y/RGBA 平面 min/max/mean + 与上帧的平均绝对差。
|
||||
// 均匀灰帧 → min≈max≈mean;静止灰帧 → diff≈0;真实画面 → 分布宽且 diff>0
|
||||
final yPlane = planes[0];
|
||||
var yMin = 255, yMax = 0, ySum = 0, diff = 0, sampled = 0;
|
||||
final prev = prevY;
|
||||
for (var i = 0; i < yPlane.length; i += 8) {
|
||||
final v = yPlane[i];
|
||||
if (v < yMin) yMin = v;
|
||||
if (v > yMax) yMax = v;
|
||||
ySum += v;
|
||||
if (prev != null && i < prev.length) {
|
||||
final d = v - prev[i];
|
||||
diff += d < 0 ? -d : d;
|
||||
}
|
||||
sampled++;
|
||||
}
|
||||
prevY = yPlane;
|
||||
final yMean = ySum / sampled;
|
||||
final yDiff =
|
||||
prev == null ? -1.0 : diff / (sampled * 255.0);
|
||||
// UV 平面统计(诊断):色序/值域异常会导致解码偏色
|
||||
var uv1 = '-';
|
||||
if (planes.length > 1) {
|
||||
final u = planes[1];
|
||||
var uMin = 255, uMax = 0, uSum = 0, n = 0;
|
||||
for (var i = 0; i < u.length; i += 8) {
|
||||
final v = u[i];
|
||||
if (v < uMin) uMin = v;
|
||||
if (v > uMax) uMax = v;
|
||||
uSum += v;
|
||||
n++;
|
||||
}
|
||||
uv1 = 's=${strides[1]} min=$uMin max=$uMax mean=${(uSum / n).toStringAsFixed(0)}';
|
||||
if (planes.length > 2) {
|
||||
final v2 = planes[2];
|
||||
var vMin = 255, vMax = 0, vSum = 0, n2 = 0;
|
||||
for (var i = 0; i < v2.length; i += 8) {
|
||||
final v = v2[i];
|
||||
if (v < vMin) vMin = v;
|
||||
if (v > vMax) vMax = v;
|
||||
vSum += v;
|
||||
n2++;
|
||||
}
|
||||
uv1 += ' v:min=$vMin max=$vMax mean=${(vSum / n2).toStringAsFixed(0)}';
|
||||
}
|
||||
}
|
||||
// 多模型并行推理:每模型先首帧自适应判定 YUV 值域/色序,再逐模型推理;
|
||||
// 汇总后按类别分组跨模型 NMS 合并(同标签重复框取高分,异标签互不压制)
|
||||
var results = <DetectionResult>[];
|
||||
for (final d in detectors) {
|
||||
if (!isBgra && !d.yuvModeKnown) {
|
||||
d.decideYuvChroma(
|
||||
planes: planes,
|
||||
strides: strides,
|
||||
width: width,
|
||||
height: height,
|
||||
);
|
||||
}
|
||||
var dets = d.detectRaw(
|
||||
planes: planes,
|
||||
strides: strides,
|
||||
width: width,
|
||||
height: height,
|
||||
isBgra: isBgra,
|
||||
rgbaOrder: rgbaOrder,
|
||||
);
|
||||
// 自愈:判定后 1.5s 内无检测且帧可用 → 用实时帧重跑完整判定
|
||||
// (首帧模糊/暗帧导致启发式猜错时,画面稳定后 oracle 可分胜负)
|
||||
if (!isBgra && d.yuvModeKnown && !d.yuvRetried &&
|
||||
dets.length <= 1 &&
|
||||
DateTime.now().millisecondsSinceEpoch - d.yuvDecisionMs >
|
||||
1500 &&
|
||||
d.retryDecision(
|
||||
planes: planes,
|
||||
strides: strides,
|
||||
width: width,
|
||||
height: height,
|
||||
)) {
|
||||
dets = d.detectRaw(
|
||||
planes: planes,
|
||||
strides: strides,
|
||||
width: width,
|
||||
height: height,
|
||||
isBgra: isBgra,
|
||||
);
|
||||
}
|
||||
results.addAll(dets);
|
||||
}
|
||||
results = mergeAcrossModels(results, TfliteDetector.iouThreshold);
|
||||
// 低分野鸡框过视觉先验(颜色/位置),减少户外误报
|
||||
results = VisualPrior.filter(
|
||||
results,
|
||||
@@ -229,6 +397,7 @@ Future<void> _workerMain(SendPort mainPort) async {
|
||||
width: width,
|
||||
height: height,
|
||||
isBgra: isBgra,
|
||||
rgbaOrder: rgbaOrder,
|
||||
);
|
||||
final motionRegions = m.detectMotionRaw(
|
||||
planes[0], strides[0], width, height);
|
||||
@@ -236,14 +405,44 @@ Future<void> _workerMain(SendPort mainPort) async {
|
||||
b.updateRaw(planes[0], strides[0], width, height);
|
||||
sw.stop();
|
||||
|
||||
var dualDiag = '';
|
||||
if (isBgra &&
|
||||
DateTime.now().millisecondsSinceEpoch - lastDualMs > 3000) {
|
||||
lastDualMs = DateTime.now().millisecondsSinceEpoch;
|
||||
final a = detectors.first.diagnoseOrder(
|
||||
planes: planes,
|
||||
strides: strides,
|
||||
width: width,
|
||||
height: height,
|
||||
isBgra: true,
|
||||
rgbaOrder: false);
|
||||
final b = detectors.first.diagnoseOrder(
|
||||
planes: planes,
|
||||
strides: strides,
|
||||
width: width,
|
||||
height: height,
|
||||
isBgra: true,
|
||||
rgbaOrder: true);
|
||||
dualDiag = ' | dual BGRA:${a.$1}@${(a.$2 * 100).toStringAsFixed(1)}%'
|
||||
' RGBA:${b.$1}@${(b.$2 * 100).toStringAsFixed(1)}%';
|
||||
}
|
||||
|
||||
mainPort.send([
|
||||
'result',
|
||||
rotation,
|
||||
width,
|
||||
height,
|
||||
results
|
||||
.map((r) =>
|
||||
[r.label, r.score, r.left, r.top, r.right, r.bottom])
|
||||
.map((r) => [
|
||||
r.label,
|
||||
r.score,
|
||||
r.left,
|
||||
r.top,
|
||||
r.right,
|
||||
r.bottom,
|
||||
r.modelId,
|
||||
r.modelName,
|
||||
])
|
||||
.toList(),
|
||||
motionRegions
|
||||
.map((mr) => [mr.left, mr.top, mr.right, mr.bottom])
|
||||
@@ -252,14 +451,47 @@ Future<void> _workerMain(SendPort mainPort) async {
|
||||
.map((mr) => [mr.left, mr.top, mr.right, mr.bottom])
|
||||
.toList(),
|
||||
sw.elapsedMilliseconds,
|
||||
'planes=${planes.length} yLen=${yPlane.length} stride=${strides[0]} '
|
||||
'y:min=$yMin max=$yMax mean=${yMean.toStringAsFixed(1)} '
|
||||
'diff=${yDiff < 0 ? '-' : yDiff.toStringAsFixed(3)} '
|
||||
'uv1:[$uv1] | ${detectors.first.yuvDiag}$dualDiag',
|
||||
]);
|
||||
break;
|
||||
case 'reset':
|
||||
motion?.reset();
|
||||
background?.reset();
|
||||
for (final d in detectors) {
|
||||
d.resetYuvMode();
|
||||
}
|
||||
break;
|
||||
case 'set-min-score':
|
||||
for (final d in detectors) {
|
||||
d.minScore = (list[1] as num).toDouble();
|
||||
}
|
||||
mainPort.send(['log', 'min-score=${detectors.isEmpty ? '-' : detectors.first.minScore}']);
|
||||
}
|
||||
} catch (e) {
|
||||
mainPort.send(['error', '$e']);
|
||||
} catch (e, st) {
|
||||
mainPort.send([
|
||||
'error',
|
||||
'$e\n${st.toString().split('\n').take(3).join('\n')}'
|
||||
]);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 多模型结果合并:按类别分组,组内 NMS(不同模型检出同一目标时取高分)。
|
||||
/// 各模型类别体系独立(如野鸡/疑似 vs 野兔/疑似),不同类别互不压制。
|
||||
List<DetectionResult> mergeAcrossModels(
|
||||
List<DetectionResult> all, double iouThreshold) {
|
||||
if (all.length <= 1) return all;
|
||||
final byLabel = <String, List<DetectionResult>>{};
|
||||
for (final r in all) {
|
||||
byLabel.putIfAbsent(r.label, () => []).add(r);
|
||||
}
|
||||
final merged = <DetectionResult>[];
|
||||
for (final group in byLabel.values) {
|
||||
merged.addAll(nms(group, iouThreshold));
|
||||
}
|
||||
merged.sort((a, b) => b.score.compareTo(a.score));
|
||||
return merged;
|
||||
}
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import 'dart:typed_data';
|
||||
|
||||
import 'package:flutter/foundation.dart' show debugPrint;
|
||||
import 'package:tflite_flutter/tflite_flutter.dart';
|
||||
|
||||
import 'detection_result.dart';
|
||||
@@ -10,9 +11,12 @@ import 'nms.dart';
|
||||
/// cx/cy/w/h 已归一化,类别得分已过 sigmoid;按 out[dim][anchor] 索引。
|
||||
/// 输入为 NCHW [1, 3, 704, 704](litert 导出保留 torch 布局)。
|
||||
class TfliteDetector {
|
||||
static const int inputSize = 704;
|
||||
// 野鸡数据置信度普遍偏低(0.1~0.2 量级),保留低分池供运动检测提升
|
||||
static const double minScore = 0.10;
|
||||
// 输入尺寸取自模型本身(ultralytics litert 导出 NCHW [1,3,H,W],各数据集
|
||||
// 训练 imgsz 可不同),默认 704 兜底
|
||||
static const int defaultInputSize = 704;
|
||||
// 野鸡数据置信度普遍偏低(0.1~0.2 量级),保留低分池供运动检测提升;
|
||||
// 可运行时调整(设置页滑块),默认 0.10
|
||||
double minScore = 0.10;
|
||||
static const double iouThreshold = 0.45;
|
||||
static const int maxDetections = 20;
|
||||
static const String modelAsset = 'assets/model.tflite';
|
||||
@@ -22,39 +26,55 @@ class TfliteDetector {
|
||||
final List<String> _labels;
|
||||
final int _numClasses;
|
||||
final int _numAnchors;
|
||||
final int inputSize;
|
||||
|
||||
final Float32List _input =
|
||||
Float32List(1 * inputSize * inputSize * 3);
|
||||
/// 模型身份(多模型并行推理区分来源;内置资产模型为 -1/空)
|
||||
final int modelId;
|
||||
final String modelName;
|
||||
|
||||
late final Float32List _input;
|
||||
|
||||
/// 输出按模型形状 [1, 4+nc, anchors] 的嵌套 List 组织,
|
||||
/// run() 要求输出对象形状与模型完全一致(扁平 List 会被拒)。
|
||||
final List<List<List<double>>> _output;
|
||||
|
||||
TfliteDetector._(this._interpreter, this._labels, this._numClasses,
|
||||
this._numAnchors, this._output);
|
||||
this._numAnchors, this._output, this.inputSize, this.modelId,
|
||||
this.modelName) {
|
||||
_input = Float32List(1 * inputSize * inputSize * 3);
|
||||
}
|
||||
|
||||
/// 模型缺失或加载失败返回 null(App 降级为仅预览)。
|
||||
/// 在后台 isolate 内调用(模型字节由主 isolate 读取后传入)。
|
||||
static Future<TfliteDetector?> fromBuffer(
|
||||
Uint8List bytes, List<String> labels) async {
|
||||
Uint8List bytes,
|
||||
List<String> labels, {
|
||||
int modelId = -1,
|
||||
String modelName = '',
|
||||
}) async {
|
||||
try {
|
||||
final interpreter = Interpreter.fromBuffer(
|
||||
bytes,
|
||||
options: InterpreterOptions()..threads = 4,
|
||||
);
|
||||
return TfliteDetector._fromModel(interpreter, labels);
|
||||
return TfliteDetector._fromModel(
|
||||
interpreter, labels, modelId, modelName);
|
||||
} catch (_) {
|
||||
return null;
|
||||
}
|
||||
}
|
||||
|
||||
/// 输出布局 [1, 4+nc, anchors] 取自模型本身,类别数不与 labels 文件长度耦合。
|
||||
factory TfliteDetector._fromModel(
|
||||
Interpreter interpreter, List<String> labels) {
|
||||
factory TfliteDetector._fromModel(Interpreter interpreter,
|
||||
List<String> labels, int modelId, String modelName) {
|
||||
final shape = interpreter.getOutputTensor(0).shape;
|
||||
final numClasses =
|
||||
shape.length >= 3 && shape[1] > 4 ? shape[1] - 4 : labels.length;
|
||||
final numAnchors = shape.length >= 3 && shape[2] > 0 ? shape[2] : 2100;
|
||||
final inputShape = interpreter.getInputTensor(0).shape;
|
||||
final inputSize = inputShape.length >= 4
|
||||
? inputShape[3]
|
||||
: defaultInputSize;
|
||||
final output = List.generate(
|
||||
1,
|
||||
(_) => List.generate(
|
||||
@@ -62,8 +82,8 @@ class TfliteDetector {
|
||||
(_) => List<double>.filled(numAnchors, 0),
|
||||
),
|
||||
);
|
||||
return TfliteDetector._(
|
||||
interpreter, labels, numClasses, numAnchors, output);
|
||||
return TfliteDetector._(interpreter, labels, numClasses, numAnchors,
|
||||
output, inputSize, modelId, modelName);
|
||||
}
|
||||
|
||||
/// 原始数据接口(后台 isolate 用,不依赖 CameraImage)。
|
||||
@@ -75,13 +95,15 @@ class TfliteDetector {
|
||||
required int width,
|
||||
required int height,
|
||||
required bool isBgra,
|
||||
bool rgbaOrder = false,
|
||||
}) {
|
||||
preprocess(
|
||||
planes: planes,
|
||||
strides: strides,
|
||||
width: width,
|
||||
height: height,
|
||||
isBgra: isBgra);
|
||||
isBgra: isBgra,
|
||||
rgbaOrder: rgbaOrder);
|
||||
// 传原始字节视图而非 Float32List:tflite_flutter 会对非 ByteBuffer/Uint8List
|
||||
// 输入调用 resizeInputTensor(1 维 [1486848]),使 node 0 TRANSPOSE prepare 失败
|
||||
_interpreter.run(_input.buffer.asUint8List(), _output);
|
||||
@@ -103,30 +125,36 @@ class TfliteDetector {
|
||||
.toList();
|
||||
}
|
||||
|
||||
/// 按像素格式分派:iOS bgra8888 单平面 / Android yuv420 多平面。
|
||||
/// 按像素格式分派:单平面 RGBA/BGRA(iOS bgra8888 / Android 实验) / yuv420 多平面。
|
||||
void preprocess({
|
||||
required List<Uint8List> planes,
|
||||
required List<int> strides,
|
||||
required int width,
|
||||
required int height,
|
||||
required bool isBgra,
|
||||
bool rgbaOrder = false,
|
||||
}) {
|
||||
if (isBgra) {
|
||||
_preprocessBgra(planes[0], strides[0], width, height);
|
||||
_preprocessBgra(planes[0], strides[0], width, height, rgbaOrder);
|
||||
} else {
|
||||
_preprocessYuv(planes, strides, width, height);
|
||||
}
|
||||
}
|
||||
|
||||
/// BGRA8888 单平面(iOS):每像素 4 字节 [b,g,r,a],双线性采样,
|
||||
/// letterbox(等比缩到长边 704,短边黑边补 0,与 YOLO 训练一致)。
|
||||
void _preprocessBgra(Uint8List src, int stride, int srcW, int srcH) {
|
||||
/// 单平面 8888(iOS bgra8888 = [b,g,r,a];Android 实验 RGBA_8888 = [r,g,b,a]):
|
||||
/// 每像素 4 字节,双线性采样,letterbox(等比缩到长边 704,短边黑边补 0)。
|
||||
void _preprocessBgra(
|
||||
Uint8List src, int stride, int srcW, int srcH, bool rgbaOrder) {
|
||||
final plane = inputSize * inputSize;
|
||||
final scale = inputSize / srcW < inputSize / srcH
|
||||
? inputSize / srcW
|
||||
: inputSize / srcH;
|
||||
final dx = (inputSize - srcW * scale) / 2;
|
||||
final dy = (inputSize - srcH * scale) / 2;
|
||||
// rgbaOrder=false(iOS BGRA): +0 B、+1 G、+2 R、+3 A;
|
||||
// rgbaOrder=true(Android RGBA): +0 R、+1 G、+2 B、+3 A
|
||||
final rOff = rgbaOrder ? 0 : 2;
|
||||
final bOff = rgbaOrder ? 2 : 0;
|
||||
|
||||
for (var oy = 0; oy < inputSize; oy++) {
|
||||
final syf = (oy - dy) / scale;
|
||||
@@ -153,23 +181,22 @@ class TfliteDetector {
|
||||
final y1 = y0 < srcH - 1 ? y0 + 1 : y0;
|
||||
final fx = sxf - x0, fy = syf - y0;
|
||||
|
||||
// BGRA 字节序:+0 B、+1 G、+2 R、+3 A
|
||||
final i00 = y0 * stride + x0 * 4;
|
||||
final i10 = y0 * stride + x1 * 4;
|
||||
final i01 = y1 * stride + x0 * 4;
|
||||
final i11 = y1 * stride + x1 * 4;
|
||||
final r00 = src[i00 + 2].toDouble();
|
||||
final r00 = src[i00 + rOff].toDouble();
|
||||
final g00 = src[i00 + 1].toDouble();
|
||||
final b00 = src[i00].toDouble();
|
||||
final r10 = src[i10 + 2].toDouble();
|
||||
final b00 = src[i00 + bOff].toDouble();
|
||||
final r10 = src[i10 + rOff].toDouble();
|
||||
final g10 = src[i10 + 1].toDouble();
|
||||
final b10 = src[i10].toDouble();
|
||||
final r01 = src[i01 + 2].toDouble();
|
||||
final b10 = src[i10 + bOff].toDouble();
|
||||
final r01 = src[i01 + rOff].toDouble();
|
||||
final g01 = src[i01 + 1].toDouble();
|
||||
final b01 = src[i01].toDouble();
|
||||
final r11 = src[i11 + 2].toDouble();
|
||||
final b01 = src[i01 + bOff].toDouble();
|
||||
final r11 = src[i11 + rOff].toDouble();
|
||||
final g11 = src[i11 + 1].toDouble();
|
||||
final b11 = src[i11].toDouble();
|
||||
final b11 = src[i11 + bOff].toDouble();
|
||||
|
||||
_input[p] = _bl(r00, r10, r01, r11, fx, fy) / 255.0;
|
||||
_input[p + plane] = _bl(g00, g10, g01, g11, fx, fy) / 255.0;
|
||||
@@ -179,25 +206,27 @@ class TfliteDetector {
|
||||
}
|
||||
|
||||
/// letterbox 缩放 + YUV → RGB 归一化 0~1(NCHW),双线性采样。
|
||||
/// 兼容 NV12(iOS 双平面,UV 交错)与 I420(Android 三平面)。
|
||||
/// 兼容 NV12(双平面,UV 交错)与 I420(三平面)。
|
||||
/// 首帧自适应:Y 值域(full/limited)与色序(U 先/V 先)因设备而异,
|
||||
/// 静态假设会在部分机型上产生偏色 → 检测退化。
|
||||
void _preprocessYuv(
|
||||
List<Uint8List> planes, List<int> strides, int srcW, int srcH) {
|
||||
final plane = inputSize * inputSize;
|
||||
final y = planes[0];
|
||||
final nv12 = planes.length == 2;
|
||||
final uv = nv12 ? planes[1] : null;
|
||||
final u = nv12 ? null : planes[1];
|
||||
final v = nv12 ? null : planes[2];
|
||||
final yStride = strides[0];
|
||||
final uvStride = strides[1];
|
||||
final vStride =
|
||||
nv12 ? uvStride : (strides.length > 2 ? strides[2] : strides[1]);
|
||||
|
||||
// U/V 平面采样(nv12:偶位 U 奇位 V;i420:三平面分离)
|
||||
// 色序修正后的 U/V 采样(nv12:偶位 U 奇位 V,NV21 相反;i420:平面 1/2 对调)
|
||||
double uAt(int x, int y) => nv12
|
||||
? uv![y * uvStride + x * 2] - 128.0
|
||||
: u![y * uvStride + x] - 128.0;
|
||||
? uv![y * uvStride + (_yuvSwapChroma ? x * 2 + 1 : x * 2)] - 128.0
|
||||
: planes[_yuvSwapChroma ? 2 : 1][y * uvStride + x] - 128.0;
|
||||
double vAt(int x, int y) => nv12
|
||||
? uv![y * uvStride + x * 2 + 1] - 128.0
|
||||
: v![y * uvStride + x] - 128.0;
|
||||
? uv![y * uvStride + (_yuvSwapChroma ? x * 2 : x * 2 + 1)] - 128.0
|
||||
: planes[_yuvSwapChroma ? 1 : 2][y * vStride + x] - 128.0;
|
||||
|
||||
final scale = inputSize / srcW < inputSize / srcH
|
||||
? inputSize / srcW
|
||||
@@ -256,10 +285,11 @@ class TfliteDetector {
|
||||
final v11 = vAt(ux1, uy1);
|
||||
final vv = _bl(v00, v10, v01, v11, fx, fy);
|
||||
|
||||
// 有限范围展开(VideoRange Y 16~235,Cb/Cr 16~240)
|
||||
final yr = (yy - 16.0) * (255.0 / 219.0);
|
||||
final un = uu * (255.0 / 224.0);
|
||||
final vn = vv * (255.0 / 224.0);
|
||||
// 值域展开:有限范围 VideoRange(Y 16~235,Cb/Cr 16~240)需线性拉伸;
|
||||
// 全值域相机直接使用原始值(与 iOS bgra 一致)
|
||||
final yr = _yuvFullRange ? yy : (yy - 16.0) * (255.0 / 219.0);
|
||||
final un = _yuvFullRange ? uu : uu * (255.0 / 224.0);
|
||||
final vn = _yuvFullRange ? vv : vv * (255.0 / 224.0);
|
||||
|
||||
// NCHW:r/g/b 分平面存储
|
||||
_input[p] = (yr + 1.402 * vn) / 255.0;
|
||||
@@ -269,6 +299,155 @@ class TfliteDetector {
|
||||
}
|
||||
}
|
||||
|
||||
/// 首帧自适应判定 YUV 模式,后续帧复用(相机重启后由 worker 复位重判)。
|
||||
/// - 值域:有限范围黑电平恒为 16,低于 12 只可能是全值域。
|
||||
/// - 色序:以模型本身为 oracle——同一帧按两种色序各推理一次,
|
||||
/// 检测数/最高分/总分更高者为真;两序均无检测时退回亮区色相计数启发
|
||||
/// (户外最亮区域为天空应偏蓝,若按默认 U 先序解出偏红则为 V 先序)。
|
||||
/// - 首帧可能曝光未收敛(过暗/全黑),此时 oracle 与启发式都不可信,
|
||||
/// 保持未判定状态等下一帧,避免在垃圾帧上锁死错误色序(真机零检测根因)。
|
||||
bool _yuvFullRange = false;
|
||||
bool _yuvSwapChroma = false;
|
||||
bool _yuvModeKnown = false;
|
||||
bool _yuvRetried = false;
|
||||
int _yuvDecisionMs = 0;
|
||||
String _yuvDiag = '';
|
||||
|
||||
bool get yuvModeKnown => _yuvModeKnown;
|
||||
bool get yuvRetried => _yuvRetried;
|
||||
int get yuvDecisionMs => _yuvDecisionMs;
|
||||
String get yuvDiag => _yuvDiag;
|
||||
|
||||
void decideYuvChroma({
|
||||
required List<Uint8List> planes,
|
||||
required List<int> strides,
|
||||
required int width,
|
||||
required int height,
|
||||
bool force = false,
|
||||
}) {
|
||||
if (_yuvModeKnown && !force) return;
|
||||
final (yMin, yMax, yMean) = _yStats(planes[0], strides[0], width, height);
|
||||
_yuvFullRange = yMin < 12;
|
||||
if (yMean < 30 || yMax < 170) {
|
||||
// 曝光未稳定:保持未判定,下一帧重试;始终昏暗则维持默认(同旧版)
|
||||
if (!_yuvModeKnown) {
|
||||
_yuvDiag = '等稳定帧 mean=${yMean.toStringAsFixed(0)} max=$yMax';
|
||||
}
|
||||
return;
|
||||
}
|
||||
_yuvModeKnown = true;
|
||||
_yuvDecisionMs = DateTime.now().millisecondsSinceEpoch;
|
||||
final a = _runWithSwap(planes, strides, width, height, false);
|
||||
final b = _runWithSwap(planes, strides, width, height, true);
|
||||
var swap = false;
|
||||
if (a.$1 != b.$1) {
|
||||
swap = b.$1 > a.$1;
|
||||
} else if (a.$2 != b.$2) {
|
||||
swap = b.$2 > a.$2;
|
||||
} else if (a.$3 != b.$3) {
|
||||
swap = b.$3 > a.$3;
|
||||
} else {
|
||||
swap = _brightRegionLeansRed(planes, strides, width, height, yMax);
|
||||
}
|
||||
_yuvSwapChroma = swap;
|
||||
_yuvDiag = 'full=$_yuvFullRange swap=$_yuvSwapChroma'
|
||||
' cA=${a.$1} sA=${a.$2.toStringAsFixed(3)}'
|
||||
' cB=${b.$1} sB=${b.$2.toStringAsFixed(3)}';
|
||||
|
||||
debugPrint('[yuv] $_yuvDiag mean=${yMean.toStringAsFixed(0)} max=$yMax');
|
||||
}
|
||||
|
||||
/// 判定后持续无检测的自愈:用实时帧重跑完整判定(仅一次)。
|
||||
/// 首帧模糊/暗帧导致启发式猜错时,等画面稳定后 oracle 即可分胜负。
|
||||
/// 返回是否执行了重判(随后应重跑 detectRaw 取新结果)。
|
||||
bool retryDecision({
|
||||
required List<Uint8List> planes,
|
||||
required List<int> strides,
|
||||
required int width,
|
||||
required int height,
|
||||
}) {
|
||||
if (_yuvRetried || !_yuvModeKnown) return false;
|
||||
final (_, yMax, yMean) = _yStats(planes[0], strides[0], width, height);
|
||||
if (yMean < 30 || yMax < 170) return false; // 帧仍不可用
|
||||
_yuvRetried = true;
|
||||
decideYuvChroma(
|
||||
planes: planes, strides: strides, width: width, height: height,
|
||||
force: true);
|
||||
return true;
|
||||
}
|
||||
|
||||
/// 相机(重新)启动后复位,首帧重新判定
|
||||
void resetYuvMode() {
|
||||
_yuvModeKnown = false;
|
||||
_yuvRetried = false;
|
||||
_yuvDiag = '';
|
||||
}
|
||||
|
||||
/// 采样统计 Y 值域:(min, max, mean),步长 16px 约 3600 样本
|
||||
(int, int, double) _yStats(Uint8List y, int yStride, int w, int h) {
|
||||
var yMin = 255, yMax = 0;
|
||||
var sum = 0, n = 0;
|
||||
for (var j = 0; j < h; j += 16) {
|
||||
final row = j * yStride;
|
||||
for (var i = 0; i < w; i += 16) {
|
||||
final v = y[row + i];
|
||||
if (v < yMin) yMin = v;
|
||||
if (v > yMax) yMax = v;
|
||||
sum += v;
|
||||
n++;
|
||||
}
|
||||
}
|
||||
return (yMin, yMax, sum / n);
|
||||
}
|
||||
|
||||
/// 按指定色序推理一次,返回 (检测数, 最高分, 总分)
|
||||
(int, double, double) _runWithSwap(
|
||||
List<Uint8List> planes, List<int> strides, int w, int h, bool swap) {
|
||||
_yuvSwapChroma = swap;
|
||||
preprocess(planes: planes, strides: strides, width: w, height: h,
|
||||
isBgra: false);
|
||||
_interpreter.run(_input.buffer.asUint8List(), _output);
|
||||
final dets = postprocess();
|
||||
var maxScore = 0.0, sumScore = 0.0;
|
||||
for (final d in dets) {
|
||||
sumScore += d.score;
|
||||
if (d.score > maxScore) maxScore = d.score;
|
||||
}
|
||||
return (dets.length, maxScore, sumScore);
|
||||
}
|
||||
|
||||
bool _brightRegionLeansRed(
|
||||
List<Uint8List> planes, List<int> strides, int w, int h, int yMax) {
|
||||
final y = planes[0];
|
||||
final yStride = strides[0];
|
||||
final uvStride = strides[1];
|
||||
final nv12 = planes.length == 2;
|
||||
final uv = nv12 ? planes[1] : null;
|
||||
// 最亮带(maxY-40 以上),整体偏暗的场景也能拿到足量样本
|
||||
final brightMin = yMax - 40;
|
||||
var blue = 0, red = 0;
|
||||
for (var j = 0; j < h; j += 8) {
|
||||
final yrow = j * yStride;
|
||||
for (var i = 0; i < w; i += 8) {
|
||||
if (y[yrow + i] < brightMin) continue;
|
||||
final cj = j ~/ 2, ci = i ~/ 2;
|
||||
if (nv12) {
|
||||
final c = cj * uvStride + ci * 2;
|
||||
if (c + 1 >= uv!.length) continue;
|
||||
if (uv[c] > 150) blue++;
|
||||
if (uv[c + 1] > 150) red++;
|
||||
} else {
|
||||
final c = cj * uvStride + ci;
|
||||
if (c >= planes[1].length || c >= planes[2].length) continue;
|
||||
if (planes[1][c] > 150) blue++;
|
||||
if (planes[2][c] > 150) red++;
|
||||
}
|
||||
}
|
||||
}
|
||||
// 亮区偏红多于偏蓝 → 当前 U/V 假设反了
|
||||
return red > blue;
|
||||
}
|
||||
|
||||
static double _bl(double a, double b, double c, double d, double fx,
|
||||
double fy) =>
|
||||
(1 - fx) * (1 - fy) * a + fx * (1 - fy) * b +
|
||||
@@ -302,11 +481,39 @@ class TfliteDetector {
|
||||
top: (cy - h / 2).clamp(0.0, 1.0),
|
||||
right: (cx + w / 2).clamp(0.0, 1.0),
|
||||
bottom: (cy + h / 2).clamp(0.0, 1.0),
|
||||
modelId: modelId,
|
||||
modelName: modelName,
|
||||
));
|
||||
}
|
||||
final kept = nms(boxes, iouThreshold);
|
||||
return kept.take(maxDetections).toList();
|
||||
}
|
||||
|
||||
/// 诊断:按指定字节序推理一次,返回 (检测数, 最高分)。
|
||||
/// 用于对比 BGRA/RGBA 两种顺序在同一帧上的检测差异(验证字节序与场景可达性)。
|
||||
(int, double) diagnoseOrder({
|
||||
required List<Uint8List> planes,
|
||||
required List<int> strides,
|
||||
required int width,
|
||||
required int height,
|
||||
required bool isBgra,
|
||||
required bool rgbaOrder,
|
||||
}) {
|
||||
preprocess(
|
||||
planes: planes,
|
||||
strides: strides,
|
||||
width: width,
|
||||
height: height,
|
||||
isBgra: isBgra,
|
||||
rgbaOrder: rgbaOrder);
|
||||
_interpreter.run(_input.buffer.asUint8List(), _output);
|
||||
final dets = postprocess();
|
||||
var maxScore = 0.0;
|
||||
for (final d in dets) {
|
||||
if (d.score > maxScore) maxScore = d.score;
|
||||
}
|
||||
return (dets.length, maxScore);
|
||||
}
|
||||
|
||||
void dispose() => _interpreter.close();
|
||||
}
|
||||
|
||||
@@ -34,12 +34,14 @@ class VisualPrior {
|
||||
required int width,
|
||||
required int height,
|
||||
required bool isBgra,
|
||||
required bool rgbaOrder,
|
||||
}) {
|
||||
if (results.isEmpty || width <= 0 || height <= 0) return results;
|
||||
final kept = <DetectionResult>[];
|
||||
for (final r in results) {
|
||||
final lowConfPheasant = r.label == 'pheasant' && r.score < maxScore;
|
||||
if (lowConfPheasant && _reject(r, planes, strides, width, height, isBgra)) {
|
||||
if (lowConfPheasant &&
|
||||
_reject(r, planes, strides, width, height, isBgra, rgbaOrder)) {
|
||||
continue;
|
||||
}
|
||||
kept.add(r);
|
||||
@@ -48,7 +50,7 @@ class VisualPrior {
|
||||
}
|
||||
|
||||
static bool _reject(DetectionResult r, List<Uint8List> planes,
|
||||
List<int> strides, int width, int height, bool isBgra) {
|
||||
List<int> strides, int width, int height, bool isBgra, bool rgbaOrder) {
|
||||
// 位置线索:detectRaw 输出为图像坐标系,centerY 直接可判天空区
|
||||
if (r.centerY < skyTopRatio) return true;
|
||||
|
||||
@@ -63,7 +65,8 @@ class VisualPrior {
|
||||
for (var gx = -2; gx <= 2; gx++) {
|
||||
final px = (cx + gx * halfW / 2).round().clamp(0, width - 1).toInt();
|
||||
final py = (cy + gy * halfH / 2).round().clamp(0, height - 1).toInt();
|
||||
final (r_, g_, b_) = _pixel(planes, strides, px, py, width, height, isBgra);
|
||||
final (r_, g_, b_) =
|
||||
_pixel(planes, strides, px, py, width, height, isBgra, rgbaOrder);
|
||||
total++;
|
||||
final mn = math.min(r_, math.min(g_, b_));
|
||||
final mx = math.max(r_, math.max(g_, b_));
|
||||
@@ -80,14 +83,20 @@ class VisualPrior {
|
||||
}
|
||||
|
||||
/// 读取单像素 RGB(0~255)。
|
||||
/// BGRA 单平面:每像素 4 字节 [b,g,r,a];
|
||||
/// 8888 单平面按实际字节序取通道:BGRA=[b,g,r,a](iOS 插件)、
|
||||
/// RGBA=[r,g,b,a](Android 自写原生通道)——字节序写死会让 Android
|
||||
/// 低分框采样到 R/B 互换的颜色(橙色野鸡身被误判成"蓝色")整批误杀;
|
||||
/// YUV:y 平面 + 4:2:0 半分辨率 U/V(NV12 交错或 I420 分离)。
|
||||
static (double, double, double) _pixel(List<Uint8List> planes,
|
||||
List<int> strides, int x, int y, int width, int height, bool isBgra) {
|
||||
List<int> strides, int x, int y, int width, int height, bool isBgra,
|
||||
bool rgbaOrder) {
|
||||
if (isBgra) {
|
||||
final src = planes[0];
|
||||
final i = y * strides[0] + x * 4;
|
||||
return (src[i + 2].toDouble(), src[i + 1].toDouble(), src[i].toDouble());
|
||||
final rOff = rgbaOrder ? 0 : 2;
|
||||
final bOff = rgbaOrder ? 2 : 0;
|
||||
return (src[i + rOff].toDouble(), src[i + 1].toDouble(),
|
||||
src[i + bOff].toDouble());
|
||||
}
|
||||
final yy =
|
||||
(planes[0][y * strides[0] + x] - 16.0) * (255.0 / 219.0);
|
||||
|
||||
@@ -6,7 +6,7 @@ import 'package:provider/provider.dart';
|
||||
import '../auth/session_store.dart';
|
||||
|
||||
/// 使用协议页:首次启动强制展示,同意后记录(secure storage),拒绝则退出应用。
|
||||
/// 用途声明:仅限野生动物观察,禁止近距离打扰与非法使用。
|
||||
/// 用途声明:仅限动物观察,禁止近距离打扰与非法使用。
|
||||
class TermsScreen extends StatefulWidget {
|
||||
const TermsScreen({super.key});
|
||||
|
||||
@@ -28,12 +28,12 @@ class _TermsScreenState extends State<TermsScreen> {
|
||||
感谢您使用「视野」动物观察应用。使用本应用前,请仔细阅读并同意以下条款:
|
||||
|
||||
一、用途限制
|
||||
本应用仅用于野生动物观察、记录与科普学习,严禁用于任何非法用途。
|
||||
本应用仅用于动物观察、记录与科普学习,严禁用于任何非法用途。
|
||||
|
||||
二、禁止行为
|
||||
1. 禁止近距离打扰、追逐、驱赶、恐吓野生动物;
|
||||
2. 禁止投喂、诱捕、猎捕、伤害野生动物;
|
||||
3. 禁止利用本应用从事盗猎、野生动物交易或其他违法犯罪活动;
|
||||
1. 禁止近距离打扰、追逐、驱赶、恐吓动物;
|
||||
2. 禁止投喂、诱捕、猎捕、伤害动物;
|
||||
3. 禁止利用本应用从事盗猎、动物交易或其他违法犯罪活动;
|
||||
4. 禁止进入自然保护区禁区、私人领地等未经许可的区域进行观察。
|
||||
|
||||
三、法律合规
|
||||
|
||||
@@ -0,0 +1,285 @@
|
||||
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;
|
||||
|
||||
const ModelCatalogItem({
|
||||
required this.datasetId,
|
||||
required this.datasetName,
|
||||
required this.version,
|
||||
required this.labels,
|
||||
required this.sizeBytes,
|
||||
required this.sha256,
|
||||
required this.downloadUrl,
|
||||
});
|
||||
|
||||
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? ?? '',
|
||||
);
|
||||
}
|
||||
|
||||
/// 已就绪模型(字节 + 标签,供推理 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 [];
|
||||
bool _ready = false;
|
||||
bool _refreshing = false;
|
||||
String? _error;
|
||||
Future<void>? _inFlight;
|
||||
|
||||
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;
|
||||
|
||||
/// 最近一次同步的错误信息(下载失败/校验失败等;目录为空不算错误)
|
||||
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 {
|
||||
final res = await _client
|
||||
.get(Uri.parse('$baseUrl/api/v1/app/update'))
|
||||
.timeout(const Duration(seconds: 8));
|
||||
// 服务器 Content-Type 无 charset,http 包默认按 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 [];
|
||||
final catalog = list
|
||||
.map((e) => ModelCatalogItem.fromJson(e as Map<String, dynamic>))
|
||||
.toList();
|
||||
|
||||
final failed = <String>[];
|
||||
for (final item in catalog) {
|
||||
if (!await _ensureLocal(item)) failed.add(item.datasetName);
|
||||
}
|
||||
await _prune(catalog);
|
||||
_models = await _loadBundles(catalog);
|
||||
_ready = true;
|
||||
_error = failed.isEmpty
|
||||
? null
|
||||
: '模型下载失败:${failed.join(',')}(重试旧模型或稍后再试)';
|
||||
} catch (e) {
|
||||
if (!_ready) _error = '模型目录拉取失败:$e';
|
||||
// 已就绪过则保留旧模型,不覆盖 error(下载级错误优先展示)
|
||||
}
|
||||
}
|
||||
|
||||
/// 保证目录条目在本地可用:meta 匹配且文件在 → 跳过;否则下载并校验 sha256。
|
||||
Future<bool> _ensureLocal(ModelCatalogItem item) async {
|
||||
final dir = await _modelDir(item.datasetId);
|
||||
try {
|
||||
final meta = await _readMeta(dir);
|
||||
final file = File('${dir.path}/model.tflite');
|
||||
if (meta != null &&
|
||||
meta['version'] == item.version &&
|
||||
meta['sha256'] == item.sha256 &&
|
||||
await file.exists()) {
|
||||
return true;
|
||||
}
|
||||
// 版本更新或文件缺失:下载校验(失败重试一次)
|
||||
for (var attempt = 0; attempt < 2; attempt++) {
|
||||
final ok = await _downloadAndVerify(item, dir, file);
|
||||
if (ok) return true;
|
||||
await file.delete().catchError((_) => file);
|
||||
await File('${dir.path}/model.tflite.part')
|
||||
.delete()
|
||||
.catchError((_) => file);
|
||||
}
|
||||
debugPrint('[ModelManager] 下载失败: ${item.datasetName} ${item.version}');
|
||||
return false;
|
||||
} catch (e) {
|
||||
debugPrint('[ModelManager] _ensureLocal ${item.datasetName}: $e');
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
Future<bool> _downloadAndVerify(
|
||||
ModelCatalogItem item, Directory dir, File file) async {
|
||||
final part = File('${file.path}.part');
|
||||
final sink = part.openWrite();
|
||||
try {
|
||||
final res = await _client
|
||||
.send(http.Request('GET', Uri.parse('$baseUrl${item.downloadUrl}')))
|
||||
.timeout(const Duration(minutes: 3));
|
||||
if (res.statusCode != 200) return false;
|
||||
await res.stream.pipe(sink);
|
||||
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);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Future<List<ModelBundle>> _loadBundles(
|
||||
List<ModelCatalogItem> catalog) async {
|
||||
final bundles = <ModelBundle>[];
|
||||
for (final item in catalog) {
|
||||
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;
|
||||
}
|
||||
}
|
||||
@@ -1,4 +1,5 @@
|
||||
import 'package:flutter/material.dart';
|
||||
import 'package:flutter/services.dart';
|
||||
import 'package:provider/provider.dart';
|
||||
|
||||
import 'models.dart';
|
||||
@@ -73,6 +74,24 @@ class _PaywallScreenState extends State<PaywallScreen> {
|
||||
}
|
||||
},
|
||||
),
|
||||
const SizedBox(height: 12),
|
||||
_PayButton(
|
||||
label: '联系管理员',
|
||||
icon: Icons.headset_mic,
|
||||
color: Colors.grey.shade600,
|
||||
enabled: true,
|
||||
onTap: () async {
|
||||
await Clipboard.setData(
|
||||
const ClipboardData(text: 'wenwu901'),
|
||||
);
|
||||
if (!context.mounted) return;
|
||||
ScaffoldMessenger.of(context).showSnackBar(
|
||||
const SnackBar(
|
||||
content: Text('已复制管理员微信号,请使用微信添加联系人'),
|
||||
),
|
||||
);
|
||||
},
|
||||
),
|
||||
if (vm.state.state == PayState.failed)
|
||||
Padding(
|
||||
padding: const EdgeInsets.only(top: 16),
|
||||
|
||||
@@ -0,0 +1,49 @@
|
||||
import 'package:flutter/services.dart';
|
||||
|
||||
/// 安装进度事件(原生 PackageInstaller 会话回调经 EventChannel 回传)
|
||||
class InstallEvent {
|
||||
/// progress / finished / failed
|
||||
final String event;
|
||||
|
||||
/// event=progress 时的安装进度 0-100
|
||||
final int progress;
|
||||
|
||||
/// event=finished 时是否安装成功
|
||||
final bool? success;
|
||||
|
||||
/// event=failed 时的错误描述
|
||||
final String? error;
|
||||
|
||||
const InstallEvent({
|
||||
required this.event,
|
||||
this.progress = 0,
|
||||
this.success,
|
||||
this.error,
|
||||
});
|
||||
}
|
||||
|
||||
/// App 内安装 APK:原生侧 PackageInstaller 会话安装(InstallerChannel.kt)
|
||||
class ApkInstaller {
|
||||
static const _method = MethodChannel('observer/installer');
|
||||
static const _progress = EventChannel('observer/installer/progress');
|
||||
|
||||
/// 提交安装。返回 installing(已进入安装流程)/ permission_required
|
||||
/// (未允许「安装未知应用」,原生侧已拉起系统设置页)。
|
||||
static Future<String> install(String path) =>
|
||||
_method.invokeMethod<String>('install', {'path': path}).then(
|
||||
(v) => v ?? 'installing',
|
||||
);
|
||||
|
||||
/// 安装进度流:progress(0-100) → finished(success) / failed(error)
|
||||
static Stream<InstallEvent> progress() {
|
||||
return _progress.receiveBroadcastStream().map((e) {
|
||||
final m = e as Map;
|
||||
return InstallEvent(
|
||||
event: m['event'] as String? ?? '',
|
||||
progress: (m['progress'] as num?)?.toInt() ?? 0,
|
||||
success: m['success'] as bool?,
|
||||
error: m['error'] as String?,
|
||||
);
|
||||
});
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,86 @@
|
||||
import 'dart:convert';
|
||||
import 'dart:io';
|
||||
|
||||
import 'package:http/http.dart' as http;
|
||||
|
||||
import '../config/app_config.dart';
|
||||
|
||||
/// App 版本更新信息(GET /api/v1/app/update 响应 data;无记录时为空信息)
|
||||
class AppUpdateInfo {
|
||||
final String version;
|
||||
final String notes;
|
||||
|
||||
const AppUpdateInfo({required this.version, required this.notes});
|
||||
|
||||
/// 服务器无版本记录时返回空信息,调用方视为无需更新
|
||||
bool get isEmpty => version.isEmpty;
|
||||
}
|
||||
|
||||
/// 启动时版本更新检查:仅 Android 检查;服务器版本高于本地版本即强制更新
|
||||
/// (无普通/强制之分)。
|
||||
class UpdateChecker {
|
||||
final String baseUrl;
|
||||
final http.Client _client;
|
||||
|
||||
UpdateChecker({String? baseUrl, http.Client? client})
|
||||
: baseUrl = baseUrl ?? AppConfig.apiBaseUrl,
|
||||
_client = client ?? http.Client();
|
||||
|
||||
/// 下载地址:固定静态路径(管理端上传的 APK 覆盖保存为固定文件)
|
||||
static String downloadUrl({String? baseUrl}) =>
|
||||
'${baseUrl ?? AppConfig.apiBaseUrl}/download/observer-latest.apk';
|
||||
|
||||
/// 拉取服务器最新版本;非 Android 或网络异常时返回空信息,不阻塞启动
|
||||
Future<AppUpdateInfo> fetch() async {
|
||||
if (!Platform.isAndroid) {
|
||||
return const AppUpdateInfo(version: '', notes: '');
|
||||
}
|
||||
try {
|
||||
final res = await _client
|
||||
.get(Uri.parse('$baseUrl/api/v1/app/update'))
|
||||
.timeout(const Duration(seconds: 8));
|
||||
// 服务器 Content-Type 无 charset,http 包默认按 latin1 解码会乱码 → 显式 utf8
|
||||
final body =
|
||||
jsonDecode(utf8.decode(res.bodyBytes)) as Map<String, dynamic>;
|
||||
final data = body['data'] as Map<String, dynamic>? ?? const {};
|
||||
return AppUpdateInfo(
|
||||
version: data['version'] as String? ?? '',
|
||||
notes: data['notes'] as String? ?? '',
|
||||
);
|
||||
} catch (_) {
|
||||
return const AppUpdateInfo(version: '', notes: '');
|
||||
}
|
||||
}
|
||||
|
||||
/// 是否需要更新:服务器版本高于「已装版本与已确认接受版本」中的较大者。
|
||||
/// APK 版本号不递增时,用户点过「立即更新」后 accepted 追上服务器版本,
|
||||
/// 已更新完成再次启动也不会反复提示。
|
||||
static bool needsUpdate(String server, String installed, String accepted) {
|
||||
if (server.isEmpty || installed.isEmpty) return false;
|
||||
return isNewer(server, installed) || isNewer(server, accepted);
|
||||
}
|
||||
|
||||
/// 语义化版本号比较:a > b 返回 true。按数字段比较(1.10.0 > 1.9.9),
|
||||
/// 任一版本号解析失败时视为相等(不触发更新)。
|
||||
static bool isNewer(String a, String b) {
|
||||
final pa = _parse(a);
|
||||
final pb = _parse(b);
|
||||
if (pa == null || pb == null) return false;
|
||||
for (var i = 0; i < 3; i++) {
|
||||
if (pa[i] != pb[i]) return pa[i] > pb[i];
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
static List<int>? _parse(String v) {
|
||||
final parts = v.split('.');
|
||||
if (parts.length != 3) return null;
|
||||
final nums = <int>[];
|
||||
for (final p in parts) {
|
||||
final n = int.tryParse(p);
|
||||
if (n == null) return null;
|
||||
nums.add(n);
|
||||
}
|
||||
return nums;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,238 @@
|
||||
import 'dart:async';
|
||||
import 'dart:io';
|
||||
|
||||
import 'package:flutter/material.dart';
|
||||
import 'package:http/http.dart' as http;
|
||||
import 'package:url_launcher/url_launcher.dart';
|
||||
|
||||
import 'installer.dart';
|
||||
|
||||
/// 强制更新页:检测到新版本时的全屏阻塞页。
|
||||
/// PopScope 禁返回(Android 系统返回 / iOS 边缘滑动均不可退出)。
|
||||
/// 点「立即更新」在 App 内流式下载 APK(显示下载进度)→ PackageInstaller
|
||||
/// 会话安装(显示安装进度),不再跳浏览器。
|
||||
/// 点击更新时回调 onUpdateAccepted(调用方持久化服务器版本号,
|
||||
/// 使 APK 版本号不递增时也不反复提示)。
|
||||
class UpdateScreen extends StatefulWidget {
|
||||
final String version;
|
||||
final String url;
|
||||
final String notes;
|
||||
final VoidCallback? onUpdateAccepted;
|
||||
|
||||
const UpdateScreen({
|
||||
super.key,
|
||||
required this.version,
|
||||
required this.url,
|
||||
this.notes = '',
|
||||
this.onUpdateAccepted,
|
||||
});
|
||||
|
||||
@override
|
||||
State<UpdateScreen> createState() => _UpdateScreenState();
|
||||
}
|
||||
|
||||
enum _Stage { idle, downloading, installing, finished, failed }
|
||||
|
||||
class _UpdateScreenState extends State<UpdateScreen> {
|
||||
final http.Client _client = http.Client();
|
||||
final File _apkFile = File('${Directory.systemTemp.path}/observer-latest.apk');
|
||||
final File _apkPart =
|
||||
File('${Directory.systemTemp.path}/observer-latest.apk.part');
|
||||
|
||||
_Stage _stage = _Stage.idle;
|
||||
|
||||
/// 进度百分比 0-100;null = 总量未知(不确定进度条)
|
||||
double? _progress;
|
||||
String? _message;
|
||||
StreamSubscription<InstallEvent>? _installSub;
|
||||
|
||||
@override
|
||||
void dispose() {
|
||||
_installSub?.cancel();
|
||||
_client.close();
|
||||
super.dispose();
|
||||
}
|
||||
|
||||
Future<void> _launch() async {
|
||||
if (_stage == _Stage.downloading ||
|
||||
_stage == _Stage.installing ||
|
||||
_stage == _Stage.finished) {
|
||||
return;
|
||||
}
|
||||
setState(() {
|
||||
_stage = _Stage.idle;
|
||||
_message = null;
|
||||
});
|
||||
widget.onUpdateAccepted?.call();
|
||||
if (!Platform.isAndroid) {
|
||||
// 更新检查本就仅 Android 触发,这里兜底非 Android 走浏览器
|
||||
final uri = Uri.tryParse(widget.url);
|
||||
if (uri == null) return;
|
||||
try {
|
||||
await launchUrl(uri, mode: LaunchMode.externalApplication);
|
||||
} catch (_) {}
|
||||
return;
|
||||
}
|
||||
// APK 已下载完成(下载是原子落盘,.part 改名后文件才存在)→ 跳过下载直接安装
|
||||
if (_apkFile.existsSync()) {
|
||||
await _install();
|
||||
return;
|
||||
}
|
||||
await _download();
|
||||
}
|
||||
|
||||
Future<void> _download() async {
|
||||
setState(() {
|
||||
_stage = _Stage.downloading;
|
||||
_progress = 0;
|
||||
});
|
||||
try {
|
||||
if (_apkPart.existsSync()) _apkPart.deleteSync();
|
||||
final res = await _client.send(http.Request('GET', Uri.parse(widget.url)));
|
||||
if (res.statusCode != 200) {
|
||||
throw HttpException('HTTP ${res.statusCode}');
|
||||
}
|
||||
final total = res.contentLength ?? -1;
|
||||
final sink = _apkPart.openWrite();
|
||||
var received = 0;
|
||||
await for (final chunk in res.stream) {
|
||||
sink.add(chunk);
|
||||
received += chunk.length;
|
||||
if (mounted && total > 0) {
|
||||
setState(() => _progress = received / total * 100);
|
||||
}
|
||||
}
|
||||
await sink.close();
|
||||
// 原子落盘:下载完成后才重命名为正式文件,避免残留半包被当成完整 APK
|
||||
_apkPart.renameSync(_apkFile.path);
|
||||
if (!mounted) return;
|
||||
await _install();
|
||||
} catch (_) {
|
||||
if (_apkPart.existsSync()) _apkPart.deleteSync();
|
||||
if (!mounted) return;
|
||||
setState(() {
|
||||
_stage = _Stage.failed;
|
||||
_progress = null;
|
||||
_message = '下载失败,请检查网络后重试';
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
Future<void> _install() async {
|
||||
setState(() {
|
||||
_stage = _Stage.installing;
|
||||
_progress = 0;
|
||||
_message = null;
|
||||
});
|
||||
await _installSub?.cancel();
|
||||
_installSub = ApkInstaller.progress().listen((e) {
|
||||
if (!mounted) return;
|
||||
switch (e.event) {
|
||||
case 'progress':
|
||||
setState(() => _progress = e.progress.toDouble());
|
||||
break;
|
||||
case 'finished':
|
||||
final ok = e.success == true;
|
||||
setState(() {
|
||||
_stage = ok ? _Stage.finished : _Stage.failed;
|
||||
_message = ok ? '安装完成,请从桌面打开新版应用' : '安装失败,请重试';
|
||||
});
|
||||
break;
|
||||
case 'failed':
|
||||
setState(() {
|
||||
_stage = _Stage.failed;
|
||||
_message = e.error ?? '安装失败,请重试';
|
||||
});
|
||||
break;
|
||||
}
|
||||
}, onError: (Object _) {
|
||||
if (!mounted) return;
|
||||
setState(() {
|
||||
_stage = _Stage.failed;
|
||||
_message = '安装失败,请重试';
|
||||
});
|
||||
});
|
||||
final result = await ApkInstaller.install(_apkFile.path);
|
||||
if (!mounted) return;
|
||||
if (result == 'permission_required') {
|
||||
// 原生侧已拉起系统设置页;APK 已缓存,用户开启后返回再点直达安装
|
||||
setState(() {
|
||||
_stage = _Stage.failed;
|
||||
_message = '请在系统设置中允许「安装未知应用」,返回后再次点击「立即更新」(APK 已缓存,无需重新下载)';
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
@override
|
||||
Widget build(BuildContext context) {
|
||||
final theme = Theme.of(context);
|
||||
final busy = _stage == _Stage.downloading || _stage == _Stage.installing;
|
||||
final progressText = _stage == _Stage.downloading
|
||||
? (_progress == null ? '下载中…' : '下载中 ${_progress!.round()}%')
|
||||
: (_progress == null ? '安装中…' : '安装中 ${_progress!.round()}%');
|
||||
|
||||
return PopScope(
|
||||
canPop: false,
|
||||
child: Scaffold(
|
||||
body: Center(
|
||||
child: Padding(
|
||||
padding: const EdgeInsets.all(32),
|
||||
child: Column(
|
||||
mainAxisSize: MainAxisSize.min,
|
||||
children: [
|
||||
Icon(Icons.system_update_alt,
|
||||
size: 64, color: theme.colorScheme.primary),
|
||||
const SizedBox(height: 16),
|
||||
Text('发现新版本 ${widget.version}',
|
||||
style: theme.textTheme.headlineSmall),
|
||||
const SizedBox(height: 12),
|
||||
if (widget.notes.isNotEmpty)
|
||||
Text(widget.notes,
|
||||
textAlign: TextAlign.center,
|
||||
style: const TextStyle(height: 1.6)),
|
||||
const SizedBox(height: 24),
|
||||
if (busy) ...[
|
||||
LinearProgressIndicator(
|
||||
value: _progress == null ? null : _progress! / 100,
|
||||
minHeight: 6,
|
||||
),
|
||||
const SizedBox(height: 12),
|
||||
Text(progressText),
|
||||
const SizedBox(height: 24),
|
||||
],
|
||||
FilledButton.icon(
|
||||
onPressed: busy || _stage == _Stage.finished
|
||||
? null
|
||||
: _launch,
|
||||
icon: const Icon(Icons.download),
|
||||
label: Text(_stage == _Stage.finished ? '已完成' : '立即更新'),
|
||||
style: FilledButton.styleFrom(
|
||||
minimumSize: const Size(200, 48),
|
||||
textStyle: const TextStyle(fontSize: 16),
|
||||
),
|
||||
),
|
||||
if (_message != null) ...[
|
||||
const SizedBox(height: 12),
|
||||
Text(
|
||||
_message!,
|
||||
textAlign: TextAlign.center,
|
||||
style: TextStyle(
|
||||
height: 1.5,
|
||||
color: _stage == _Stage.failed
|
||||
? theme.colorScheme.error
|
||||
: Colors.green,
|
||||
),
|
||||
),
|
||||
],
|
||||
const SizedBox(height: 12),
|
||||
Text('不更新将无法继续使用',
|
||||
style: theme.textTheme.bodySmall
|
||||
?.copyWith(color: theme.colorScheme.error)),
|
||||
],
|
||||
),
|
||||
),
|
||||
),
|
||||
),
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -90,7 +90,7 @@ packages:
|
||||
source: hosted
|
||||
version: "0.11.4"
|
||||
camera_android_camerax:
|
||||
dependency: "direct main"
|
||||
dependency: transitive
|
||||
description:
|
||||
name: camera_android_camerax
|
||||
sha256: "8516fe308bc341a5067fb1a48edff0ddfa57c0d3cdcc9dbe7ceca3ba119e2577"
|
||||
@@ -162,7 +162,7 @@ packages:
|
||||
source: hosted
|
||||
version: "0.3.5+4"
|
||||
crypto:
|
||||
dependency: transitive
|
||||
dependency: "direct main"
|
||||
description:
|
||||
name: crypto
|
||||
sha256: c8ea0233063ba03258fbcf2ca4d6dadfefe14f02fab57702265467a19f27fadf
|
||||
@@ -497,7 +497,7 @@ packages:
|
||||
source: hosted
|
||||
version: "3.0.0"
|
||||
package_info_plus:
|
||||
dependency: transitive
|
||||
dependency: "direct main"
|
||||
description:
|
||||
name: package_info_plus
|
||||
sha256: "468c26b4254ab01979fa5e4a98cb343ea3631b9acee6f21028997419a80e1a20"
|
||||
@@ -521,7 +521,7 @@ packages:
|
||||
source: hosted
|
||||
version: "1.9.1"
|
||||
path_provider:
|
||||
dependency: transitive
|
||||
dependency: "direct main"
|
||||
description:
|
||||
name: path_provider
|
||||
sha256: a7f4874f987173da295a61c181b8ee71dab59b332a486b391babf26a1b884825
|
||||
@@ -765,6 +765,70 @@ packages:
|
||||
url: "https://pub.flutter-io.cn"
|
||||
source: hosted
|
||||
version: "1.4.0"
|
||||
url_launcher:
|
||||
dependency: "direct main"
|
||||
description:
|
||||
name: url_launcher
|
||||
sha256: f6a7e5c4835bb4e3026a04793a4199ca2d14c739ec378fdfe23fc8075d0439f8
|
||||
url: "https://pub.flutter-io.cn"
|
||||
source: hosted
|
||||
version: "6.3.2"
|
||||
url_launcher_android:
|
||||
dependency: transitive
|
||||
description:
|
||||
name: url_launcher_android
|
||||
sha256: b413d49b73867ac08dd2f9890efd3cc11f2a0e577618d50843440a1fb3776c32
|
||||
url: "https://pub.flutter-io.cn"
|
||||
source: hosted
|
||||
version: "6.3.32"
|
||||
url_launcher_ios:
|
||||
dependency: transitive
|
||||
description:
|
||||
name: url_launcher_ios
|
||||
sha256: "580fe5dfb51671ae38191d316e027f6b76272b026370708c2d898799750a02b0"
|
||||
url: "https://pub.flutter-io.cn"
|
||||
source: hosted
|
||||
version: "6.4.1"
|
||||
url_launcher_linux:
|
||||
dependency: transitive
|
||||
description:
|
||||
name: url_launcher_linux
|
||||
sha256: d5e14138b3bc193a0f63c10a53c94b91d399df0512b1f29b94a043db7482384a
|
||||
url: "https://pub.flutter-io.cn"
|
||||
source: hosted
|
||||
version: "3.2.2"
|
||||
url_launcher_macos:
|
||||
dependency: transitive
|
||||
description:
|
||||
name: url_launcher_macos
|
||||
sha256: "368adf46f71ad3c21b8f06614adb38346f193f3a59ba8fe9a2fd74133070ba18"
|
||||
url: "https://pub.flutter-io.cn"
|
||||
source: hosted
|
||||
version: "3.2.5"
|
||||
url_launcher_platform_interface:
|
||||
dependency: transitive
|
||||
description:
|
||||
name: url_launcher_platform_interface
|
||||
sha256: "552f8a1e663569be95a8190206a38187b531910283c3e982193e4f2733f01029"
|
||||
url: "https://pub.flutter-io.cn"
|
||||
source: hosted
|
||||
version: "2.3.2"
|
||||
url_launcher_web:
|
||||
dependency: transitive
|
||||
description:
|
||||
name: url_launcher_web
|
||||
sha256: "85c81589622fbc87c1c683aaea164d3604a7777495a79d91e39ffcdec39ddb34"
|
||||
url: "https://pub.flutter-io.cn"
|
||||
source: hosted
|
||||
version: "2.4.3"
|
||||
url_launcher_windows:
|
||||
dependency: transitive
|
||||
description:
|
||||
name: url_launcher_windows
|
||||
sha256: "712c70ab1b99744ff066053cbe3e80c73332b38d46e5e945c98689b2e66fc15f"
|
||||
url: "https://pub.flutter-io.cn"
|
||||
source: hosted
|
||||
version: "3.1.5"
|
||||
uuid:
|
||||
dependency: transitive
|
||||
description:
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
name: observer
|
||||
description: "视野 - 野生动物实时识别 (野鸡/生境), YOLOv8 + 充值付费"
|
||||
description: "视野 - 动物实时识别 (野鸡/生境), YOLOv8 + 充值付费"
|
||||
publish_to: 'none'
|
||||
|
||||
version: 1.0.0+1
|
||||
version: 1.0.4+5
|
||||
|
||||
environment:
|
||||
sdk: ^3.12.2
|
||||
@@ -11,8 +11,10 @@ dependencies:
|
||||
flutter:
|
||||
sdk: flutter
|
||||
cupertino_icons: ^1.0.8
|
||||
# Android 已改用自写原生相机通道(android/ 下 CameraChannel.kt,分析帧
|
||||
# 原生侧旋转成竖屏后回传),不再依赖 camera_android_camerax 插件;
|
||||
# camera + camera_avfoundation 仅用于 iOS 路径
|
||||
camera: ^0.11.0
|
||||
camera_android_camerax: ^0.6.6
|
||||
camera_avfoundation: ^0.9.17
|
||||
tflite_flutter: ^0.11.0
|
||||
flutter_secure_storage: ^9.2.0
|
||||
@@ -25,6 +27,11 @@ dependencies:
|
||||
http: ^1.2.0
|
||||
cupertino_http: ^3.0.2
|
||||
wakelock_plus: ^1.2.0
|
||||
package_info_plus: ^9.0.1
|
||||
url_launcher: ^6.3.2
|
||||
# 模型热更新:多模型下载(crypto 校验 sha256;path_provider 取应用私有目录持久化)
|
||||
crypto: ^3.0.0
|
||||
path_provider: ^2.1.4
|
||||
|
||||
# 微信/支付宝原生 SDK 配置(占位值,与 lib/config/app_config.dart 一致;接入真实支付时替换。
|
||||
# 注意:fluwx 的 universal_link 占位符会被其 pod 脚本注入 Associated Domains,
|
||||
|
||||
@@ -0,0 +1,141 @@
|
||||
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();
|
||||
|
||||
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({String version = 'v1.0.0', String sha = ''}) => {
|
||||
'datasetId': 7,
|
||||
'datasetName': '野鸡数据集',
|
||||
'version': version,
|
||||
'labels': ['pheasant', 'suspect'],
|
||||
'sizeBytes': _modelBytes.length,
|
||||
'sha256': sha.isEmpty ? _shaHex(_modelBytes) : sha,
|
||||
'downloadUrl': '/download/models/野鸡数据集/latest.tflite',
|
||||
};
|
||||
|
||||
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('首次拉取:下载模型并落盘(model/labels/meta)', () async {
|
||||
final m = manager(client([_item()]));
|
||||
await m.refresh();
|
||||
|
||||
expect(m.ready, isTrue);
|
||||
expect(m.error, isNull);
|
||||
expect(m.models.length, 1);
|
||||
expect(m.models.first.datasetName, '野鸡数据集');
|
||||
expect(m.models.first.bytes, _modelBytes);
|
||||
expect(downloadHits, 1);
|
||||
|
||||
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('版本未变不重复下载(meta 命中直接跳过)', () async {
|
||||
final m = manager(client([_item()]));
|
||||
await m.refresh();
|
||||
expect(downloadHits, 1);
|
||||
|
||||
await m.refresh();
|
||||
expect(downloadHits, 1, reason: 'meta 匹配应跳过下载');
|
||||
expect(m.models.length, 1);
|
||||
});
|
||||
|
||||
test('版本更新触发重新下载', () async {
|
||||
final m = manager(client([_item()]));
|
||||
await m.refresh();
|
||||
expect(downloadHits, 1);
|
||||
|
||||
await m.refresh();
|
||||
expect(downloadHits, 1);
|
||||
|
||||
// 发布新版本:再次刷新应重下
|
||||
await m.refresh();
|
||||
expect(downloadHits, 1);
|
||||
// 上面三次同一版本,重新构造带新版本的 manager(同一 root)
|
||||
final m2 = manager(client([_item(version: 'v2.0.0')]));
|
||||
await m2.refresh();
|
||||
expect(downloadHits, 2);
|
||||
expect(m2.models.first.version, 'v2.0.0');
|
||||
});
|
||||
|
||||
test('sha256 不匹配:重试后失败,保留旧模型并报错', () async {
|
||||
// 第一次下载成功(sha 匹配)
|
||||
final m1 = manager(client([_item()]));
|
||||
await m1.refresh();
|
||||
expect(m1.models.length, 1);
|
||||
|
||||
// 服务器 sha 与文件不符(被篡改/损坏)→ 下载校验失败
|
||||
final bad = _item(version: 'v3.0.0');
|
||||
bad['sha256'] = _shaHex([9, 9, 9]);
|
||||
final m2 = manager(client([bad]));
|
||||
await m2.refresh();
|
||||
|
||||
expect(m2.models.length, 0, reason: '校验失败的模型不应加载');
|
||||
expect(m2.error, isNotNull);
|
||||
expect(m2.error, contains('野鸡数据集'));
|
||||
});
|
||||
|
||||
test('目录下线:清理本地并清空模型', () async {
|
||||
final m = manager(client([_item()]));
|
||||
await m.refresh();
|
||||
expect(m.models.length, 1);
|
||||
expect(await Directory('${root.path}/7').exists(), isTrue);
|
||||
|
||||
final m2 = manager(client([]));
|
||||
await m2.refresh();
|
||||
expect(m2.models, isEmpty);
|
||||
expect(await Directory('${root.path}/7').exists(), isFalse,
|
||||
reason: '下线的数据集模型目录应被清理');
|
||||
});
|
||||
|
||||
test('目录接口异常:不覆盖已有就绪状态', () async {
|
||||
final m = manager(client([_item()]));
|
||||
await m.refresh();
|
||||
expect(m.ready, isTrue);
|
||||
|
||||
final broken = manager(MockClient((_) async => http.Response('boom', 500)));
|
||||
await broken.refresh();
|
||||
expect(broken.ready, isFalse);
|
||||
expect(broken.error, isNotNull);
|
||||
});
|
||||
}
|
||||
@@ -0,0 +1,55 @@
|
||||
import 'package:observer/detection/detection_result.dart';
|
||||
import 'package:observer/detection/detector_worker.dart';
|
||||
import 'package:flutter_test/flutter_test.dart';
|
||||
|
||||
DetectionResult box(String label, double score, double x, double y,
|
||||
{int modelId = -1, String modelName = ''}) =>
|
||||
DetectionResult(
|
||||
label: label,
|
||||
score: score,
|
||||
left: x,
|
||||
top: y,
|
||||
right: x + 0.1,
|
||||
bottom: y + 0.1,
|
||||
modelId: modelId,
|
||||
modelName: modelName,
|
||||
);
|
||||
|
||||
void main() {
|
||||
test('不同模型同标签重复框:NMS 去重取高分', () {
|
||||
// 野鸡模型与野兔模型都检出了同一只"野鸡"(不同模型对同一目标的重复框)
|
||||
final all = [
|
||||
box('pheasant', 0.18, 0.3, 0.3, modelId: 1, modelName: '野鸡模型'),
|
||||
box('pheasant', 0.55, 0.31, 0.3, modelId: 2, modelName: '野兔模型'),
|
||||
];
|
||||
final merged = mergeAcrossModels(all, 0.45);
|
||||
expect(merged.length, 1);
|
||||
expect(merged.first.score, 0.55);
|
||||
expect(merged.first.modelName, '野兔模型');
|
||||
});
|
||||
|
||||
test('不同类别互不压制', () {
|
||||
final all = [
|
||||
box('pheasant', 0.3, 0.5, 0.5, modelId: 1),
|
||||
box('hare', 0.7, 0.5, 0.5, modelId: 2), // 同位置但不同类别
|
||||
];
|
||||
final merged = mergeAcrossModels(all, 0.45);
|
||||
expect(merged.length, 2);
|
||||
});
|
||||
|
||||
test('同模型内部与跨模型合并一致:远处不重叠保留', () {
|
||||
final all = [
|
||||
box('pheasant', 0.2, 0.1, 0.1, modelId: 1, modelName: '野鸡模型'),
|
||||
box('pheasant', 0.3, 0.8, 0.8, modelId: 1, modelName: '野鸡模型'),
|
||||
];
|
||||
final merged = mergeAcrossModels(all, 0.45);
|
||||
expect(merged.length, 2);
|
||||
expect(merged.first.score, 0.3); // 按分排序
|
||||
});
|
||||
|
||||
test('单条结果原样返回', () {
|
||||
final single = [box('suspect', 0.11, 0.2, 0.2)];
|
||||
final merged = mergeAcrossModels(single, 0.45);
|
||||
expect(identical(merged, single), isTrue);
|
||||
});
|
||||
}
|
||||
@@ -23,6 +23,8 @@ COPY --from=builder /out/observer-server /app/observer-server
|
||||
COPY config.yml /app/config.yml
|
||||
# 管理端静态产物(server_admin 构建输出),由 /admin 路径托管
|
||||
COPY --from=builder /build/admin_dist /app/admin_dist
|
||||
# APK 下载引导页(源码目录 h5/),由 /download-page 路径托管
|
||||
COPY --from=builder /build/h5 /app/h5
|
||||
# 运行时数据目录(SQLite 库),由 docker-compose 挂载持久化
|
||||
RUN mkdir -p /app/data
|
||||
EXPOSE 8080
|
||||
|
||||
+83
-4
@@ -12,7 +12,13 @@
|
||||
| 订单确认 | 客户端支付成功后 `POST /api/v1/orders/{orderId}/confirm` 幂等通知,加速授权刷新 |
|
||||
| 授权查询 | `GET /api/v1/license`(Bearer token)返回授权状态,App 识别入口强制服务端校验 |
|
||||
| 套餐 | `config.yml` `plans` 节点配置三档套餐(改价 = 改配置重启),价格**整数分**(1000 / 5600 / 18000) |
|
||||
| 后台管理端 | `server_admin/`(Vue3 + Element Plus)管理页面:订单查询、账号/授权管理(手动授权/撤销);构建产物由后端 `/admin/` 托管,登录页输入 token 后以 `X-Admin-Token` 头鉴权(`config.yml admin.token`) |
|
||||
| 后台管理端 | `server_admin/`(Vue3 + Element Plus)管理页面:订单查询、账号/授权管理(手动授权/撤销)、App 版本管理;构建产物由后端 `/admin/` 托管,登录页输入 token 后以 `X-Admin-Token` 头鉴权(`config.yml admin.token`) |
|
||||
| 版本管理 | 后台管理端上传 Android APK + 更新说明,APK 存服务器 `app.apkDir`(默认 `./workspace/`,与 `./data` 平级、挂载持久化)**固定文件名 `observer-latest.apk`,上传即覆盖,目录永远只保留最新一个文件**;**版本号从文件名识别**:文件须命名为 `observer-x.y.z.apk`(Flutter 打包产物即此命名,版本号取自 pubspec);客户端启动时 `GET /api/v1/app/update` 检查更新:服务器版本高于本地版本即弹更新提示(不可跳过)。**仅 Android 检查,iOS 不做版本下发**(iOS 走 App Store 自行更新)。版本记录可删除:删最新版本联动删除 APK 文件,删历史版本仅删记录 |
|
||||
| 数据训练(唯一入口) | 后台管理端「数据训练」一个菜单承载数据集全流程:**数据集卡片列表**(封面图/描述/图片数/已标注数/**训练状态徽标**),**卡片下方直接展示训练任务进度条与状态**(无独立训练页);详情页为**图片与标注一体视图**:分页(每页 20 条)逐行「原图 ‖ 标注图」对照展示;**图片入库(手动上传/AI 生成)自动触发 RF-DETR 全图扫描标注**,进度条展示在页顶;页顶另有「全量标注」按钮可手动重标全部图片(覆盖各图已有标注);点击原图/标注图弹窗放大,弹窗内 canvas 直接画框/点框删除/清空并保存——AI 自动标注结果直接作为标注,人工可修改/清理全部框;封面(上传自动转 jpg + UUID 命名)/**描述**/AI 生成图片(provider 抽象,默认 qwen-image/DashScope 付费 API);AI 标注端点与训练机 SSH 为**全局配置,直接读 `config.yml`**(`localAi` / `training.ssh` 节点,改配置需重启服务);图片落服务器 `app.datasetDir`/`datasets/<数据集名>/`,DB 存元数据 + 标注 JSON |
|
||||
| 模型训练 | 从数据集卡片「开始训练」触发:参数(数据集/imgsz/epochs/batch)、进度/日志/指标监控(每 epoch 粒度)、取消;训练通道 `training` 节点可配置 subprocess(与 Go 服务同机直接起 python)/ ssh(异机执行,SSH 凭据取 `config.yml` `training.ssh` 节点),并发度 1(GPU 独占);训练脚本 `server/training/train_server.py`(随项目迁移,2026-08-26)参数化,产物(best.tflite/best.pt/曲线)拉回服务器;训练收尾自动做 **tflite 产物自检**(输入/输出 shape 校验,原 `inspect_tflite.py` 逻辑内嵌脚本),自检失败任务置失败并带出原因;`dump_graph.py` 留作训练机人工深度调试 |
|
||||
| 模型版本与热更新 | **每数据集一个模型**:训练完成后一键「发布」(训练任务操作列)——tflite 落 `workspace/models/<数据集>/latest.tflite` + sha256/指标/类别名入 `model_version`(按数据集独立版本序列 m1.0.0 递增)。管理端**无模型管理界面**(版本记录仅支撑客户端下发)。**App 模型热更新**:`GET /api/v1/app/update` 扩展返回 `models` 目录数组,客户端独立检查,新模型下载校验替换,失败回退旧模型——模型迭代不再重打包 APK |
|
||||
| 模型目录与多模型推理 | `GET /api/v1/models`(登录态)返回全部数据集当前生效模型(数据集/版本/类别/大小/sha256/下载地址),下载 URL `/download/models/<数据集>/latest.tflite`;**App 模型管理页**用户自由下载/删除/启用模型,识别时**加载全部已启用模型并行推理 + 跨模型 NMS 合并**(按类别名),内置 assets 模型兜底 |
|
||||
| 标注 | **图片入库自动触发**:手动上传/AI 生成成功后,新增图自动调 `config.yml` `localAi` 节点配置的 AI 端点做 RF-DETR 全图扫描(`label_task` 记录进度,页顶进度条展示;**localAi 未配置 → 上传/生成接口直接报错;已有标注任务在跑(忙)→ 不报错**,当前任务成功完成后自动补标未标注图)→ 扫描结果(**重叠去重**:NMS 风格按置信度降序保留,重叠比 > `localAi.overlapThreshold` 默认 0.3 的框剔除——重叠比 = 交叠面积/两框较小面积,RF-DETR 同目标常输出一大一小两框,此判据能命中,同目标只留置信度最高者)**直接写 `dataset_image.labels_json`**(覆盖该图已有标注,即重标语义);点击弹窗放大后在 canvas 上画框/点框删除/清空/改类别 → 保存即整体覆写 `dataset_image.labels_json`(YOLO 归一化 JSON 数组,AI 与人工框同存,人工可修改/清理);`POST /admin/label-tasks` 详情页「全量标注」按钮入口(另有自动触发),可发起全量/指定图重标;自动/手动/混合并存,训练前自动整理(prepare_yolo 逻辑在服务端) |
|
||||
|
||||
## 架构与数据流
|
||||
|
||||
@@ -40,7 +46,13 @@ Flutter App ── POST /auth/register|login ─► 账号注册/登录,签发
|
||||
| 表 | 说明 | 关键字段 |
|
||||
|---|---|---|
|
||||
| `payment_order` | 支付订单 | `order_id`(PK)、`phone_num`、`plan_id`、`channel`(wechat/alipay)、`amount_cents`、`status`(created/paid/closed)、`wx_trade_no`(UNIQUE)、`alipay_trade_no`(UNIQUE)、`created_at`、`paid_at` |
|
||||
| `license` | 手机号账号与授权 | `phone_num`(PK)、`password`(bcrypt)、`expires_at`(未充值 NULL)、`created_at`、`updated_at` |
|
||||
| `license` | 手机号账号与授权 | `phone_num`(PK)、`password`(bcrypt)、`expires_at`(未充值 NULL)、`remark`(管理端备注)、`created_at`、`updated_at` |
|
||||
| `app_version` | App 版本管理 | `id`(PK)、`version`(x.y.z, UNIQUE)、`notes`(更新说明)、`created_at`、`updated_at`(下载地址不落表:APK 固定文件 `app.apkDir`/`observer-latest.apk`,默认 `./workspace/`) |
|
||||
| `dataset` | 训练数据集 | `id`(PK)、`name`(UNIQUE)、`source`(manual/ai)、`image_count`、`labeled_count`、`status`(building/synced/labeled)、`cover`(封面文件名,UUID 命名 jpg,如 `9f2a...-xx.jpg`)、`description`、`created_at`、`updated_at`(图片文件在 `app.datasetDir`/`datasets/<name>/`,DB 只存元数据;AI 标注/训练机 SSH 配置走 `config.yml` 的 `localAi` / `training.ssh` 节点) |
|
||||
| `dataset_image` | 数据集图片 | `id`(PK)、`dataset_id`、`filename`、`source`(manual/ai)、`prompt`(AI 生成图记录提示词)、`labels_json`(标注 JSON 数组:YOLO 归一化 xywh+类别+置信度,AI 自动标注与人工标注同存、人工可修改/清理,null/''/'[]'=未标注)、`created_at` |
|
||||
| `model_training` | 训练任务 | `id`(PK)、`name`、`status`(running/success/failed)、`dataset`(训练机数据集名)、`imgsz`/`epochs`/`batch`/`device`(参数快照)、`current_epoch`/`total_epochs`、`metrics`(JSON)、`log_tail`、`pid`、`error`、`started_at`/`finished_at`、`created_at` |
|
||||
| `model_version` | 模型版本(每数据集独立序列) | `id`(PK)、`dataset_id`、`version`(m1.0.0 递增, 同数据集 UNIQUE)、`training_id`、`artifact_file`(workspace 相对路径)、`metrics`(JSON)、`labels`(JSON 类别名数组)、`sha256`、`size_bytes`、`is_latest`、`notes`、`created_at`(模型文件不落表:发布即写 `models/<数据集名>/latest.tflite`,客户端固定下载,无存档回退) |
|
||||
| `label_task` | 标注任务 | `id`(PK)、`dataset_id`、`filenames`(JSON 选中图片列表,NULL=全量)、`status`(running/done)、`total`/`done`、`created_at`、`finished_at` |
|
||||
|
||||
建表与迁移见 `技术设计.md`(新库直接建表;存量库以 `PRAGMA user_version` 版本化迁移)。
|
||||
|
||||
@@ -131,6 +143,29 @@ Flutter App ── POST /auth/register|login ─► 账号注册/登录,签发
|
||||
| 微信支付(APP 支付 V3) | `POST /api/v1/payment/wechat/notify` | 解密+验签回调,`out_trade_no` → 落授权 |
|
||||
| 支付宝(APP 支付) | `POST /api/v1/payment/alipay/notify` | RSA2 验签回调,`out_trade_no` → 落授权 |
|
||||
|
||||
### GET /api/v1/app/update
|
||||
|
||||
App 版本更新检查(公开接口,无需 token,未登录/旧版本均可访问)。返回服务器最新版本记录;无任何记录时 `data` 为空对象,客户端视为无需更新。
|
||||
|
||||
响应 `data`:
|
||||
|
||||
```json
|
||||
{
|
||||
"version": "1.1.0",
|
||||
"notes": "修复识别准确率问题"
|
||||
}
|
||||
```
|
||||
|
||||
- **仅 Android 客户端调用**(iOS 不做版本下发,走 App Store 自行更新)
|
||||
- 检测到新版本(服务器版本高于本地版本)即**强制更新**,客户端弹不可关闭的全屏提示,必须跳转更新后才能继续使用;本地已是新版本则不提示
|
||||
- 客户端以「语义化版本号」比较:`1.10.0 > 1.9.9`(按数字段比较,禁止字符串比较)
|
||||
- 下载地址为固定静态路径:`/download/observer-latest.apk`(`app.apkDir` 目录下永远只有最新一个文件,由后端静态托管),客户端拼 `apiBaseUrl` 访问
|
||||
- **模型热更新(与 APK 更新独立通道)**:服务器有已发布模型时响应额外返回 `models` 数组(与 `GET /api/v1/models` 同构:datasetId/datasetName/version/labels/sizeBytes/sha256/downloadUrl);客户端启动与 APK 更新**独立检查**——某数据集服务器版本高于本地已下载版本即下载 `/download/models/<数据集>/latest.tflite` 到应用私有目录,sha256 校验后原子替换,下次识别生效;**非强制**,失败回退旧模型下次启动重试。App 模型管理页列出服务器全部可用模型,用户自由下载/删除/启用;识别时加载全部已启用模型**并行推理 + 跨模型 NMS 合并**(按类别名),内置 assets 模型兜底。无发布模型时不返回 models 字段(旧 App 忽略新字段、新 App 兼容旧服务器)
|
||||
|
||||
### GET /download-page
|
||||
|
||||
APK 下载引导页(静态页面,源码在 `h5/index.html`,由后端 `/download-page` 路径托管):微信内置浏览器会拦截 APK 下载,此页面按打开环境分流——**微信内打开**显示图形引导(点击右上角「···」→「在浏览器打开」+ 复制链接兜底);**手机浏览器打开**直接显示下载按钮,直链 `http://observer.redpowerfuture.com/download/observer-latest.apk`。
|
||||
|
||||
### 管理端接口(`/api/v1/admin`,需请求头 `X-Admin-Token` = `config.yml admin.token`)
|
||||
|
||||
| 方法 | 路径 | 说明 |
|
||||
@@ -139,12 +174,56 @@ Flutter App ── POST /auth/register|login ─► 账号注册/登录,签发
|
||||
| GET | `/admin/licenses` | 账号/授权列表,筛选 `phoneNum` + `page/size`(含未充值账号) |
|
||||
| POST | `/admin/licenses/grant` | 手动授权 `{"phoneNum":"13800000000","planId":"day"}`,自然日叠加 |
|
||||
| POST | `/admin/licenses/revoke` | 撤销授权 `{"phoneNum":"13800000000"}`(清授权保留账号) |
|
||||
| POST | `/admin/licenses/remark` | 写账号备注 `{"phoneNum":"13800000000","remark":"..."}`(空串清空,≤200 字) |
|
||||
| GET | `/admin/app-versions` | 版本记录列表,`page/size` 分页,按下发时间倒序 |
|
||||
| POST | `/admin/app-versions` | 下发新版本(multipart/form-data):`notes` + `file`(APK 文件,仅接受 `.apk`);**版本号从文件名识别**,文件须命名为 `observer-x.y.z.apk`(如 `observer-1.0.1.apk`),格式不符拒绝;版本号不可重复,APK 上传覆盖 `app.apkDir`/`observer-latest.apk`(目录永远只有一个文件);检测到新版本即强制更新 |
|
||||
| POST | `/admin/app-versions/delete` | 删除版本记录 `{"id":1}`:删**最新版本**时联动删除 APK 文件(客户端不再提示更新、下载 404);删历史版本只删记录不动文件 |
|
||||
| POST | `/admin/datasets` | 创建数据集 `{"name":"pheasant_v2","source":"manual"\|"ai"}`(name ≤50 字唯一,目录自动建) |
|
||||
| GET | `/admin/datasets` | 数据集列表:`page/size` 分页,返回 `{total, list}`(含 imageCount/labeledCount/status/cover/description/**training 聚合状态**:最新训练记录的 status/currentEpoch/totalEpochs) |
|
||||
| POST | `/admin/datasets/update` | 更新数据集配置 `{"id":1,"description":"...","cover":"a.jpg"}`:描述/封面,空值字段不覆盖原值 |
|
||||
| POST | `/admin/datasets/cover` | 上传数据集封面(multipart:`datasetId`+`file`,jpg/jpeg/png ≤2MB):**自动转 jpg + UUID 命名**落盘并覆盖旧封面 |
|
||||
| GET | `/admin/datasets/cover` | 封面文件(静态字节流,`datasetId` 定位) |
|
||||
| POST | `/admin/datasets/cover/delete` | 删除数据集封面 `{"datasetId":1}`:删文件 + 清 cover 字段 |
|
||||
| POST | `/admin/datasets/upload` | 上传图片(multipart/form-data:`datasetId` + `files` 多张,仅接受 `.jpg/.jpeg/.png`),存 `app.datasetDir`/`datasets/<name>/`,逐张入库 |
|
||||
| POST | `/admin/datasets/generate` | AI 生成图片 `{"datasetId":1,"prompt":"...","count":1}`(count 1..8,同步执行):调 `imageGen` provider(默认 qwen-image/DashScope)逐张生成落盘 + 入库(记录 prompt),任意失败返回错误 |
|
||||
| GET | `/admin/datasets/images` | 数据集图片列表 `{"datasetId":1}`,返回图片元数据(文件名/来源/prompt/创建时间) |
|
||||
| GET | `/admin/datasets/image` | 图片文件(静态字节流,`datasetId` + `filename` 定位,供缩略图/查看) |
|
||||
| POST | `/admin/datasets/images/delete` | 删除图片 `{"datasetId":1,"ids":[1,2]}`:删文件 + 删记录(AI 生成图是付费资产,前端确认文案提示) |
|
||||
| GET | `/admin/datasets/export` | 导出数据集 zip:`datasetId`,打包图片目录为 zip 下载(标注衔接用) |
|
||||
| POST | `/admin/datasets/sync` | 同步数据集到训练机 `{"id":1}`:按 `training` 通道推送图片到训练机 `datasetDir/<name>/`(subprocess 同机 cp、ssh 异机 scp/rsync);训练启动前自动执行 |
|
||||
| POST | `/admin/label-tasks` | 发起预标注 `{"datasetId":1,"filenames":["a.jpg","b.jpg"]}`(详情页「全量标注」按钮入口,图片入库亦自动触发):**filenames 缺省=全量**,指定图则只扫描选中图(全量重标/单张修补);调 `config.yml` `localAi` 配置的 AI 端点 RF-DETR 逐张推理(common 池并行,未配置报错),扫描结果做**重叠去重**(NMS 风格按置信度降序保留,重叠比 > `localAi.overlapThreshold` 默认 0.3 的框剔除,同目标只留置信度最高者)后**直写 `dataset_image.labels_json`(重跑覆盖该图标注)**,`label_task` 记录进度 |
|
||||
| GET | `/admin/label-tasks` | 标注任务列表:`page/size` 分页(含 status/total/done) |
|
||||
| GET | `/admin/label-tasks/detail` | 标注任务详情:返回数据集全部图片 + 每张标注框(`boxes`,YOLO 归一化 xywh + 置信度 + 类别) |
|
||||
| POST | `/admin/label-tasks/save` | 保存单张标注 `{"datasetId":1,"filename":"a.jpg","boxes":[{"class":0,"cx":0.5,"cy":0.4,"w":0.1,"h":0.2}]}`:整体覆写该图 `labels_json`(空 boxes=清空标注),返回该数据集当前 `labeledCount` |
|
||||
| POST | `/admin/trainings` | 发起训练 `{"dataset":"yolo","imgsz":704,"epochs":150,"batch":16,"device":"0","name":"..."}`:先同步数据集到训练机 → 校验目录存在 → runner 启动训练;**并发度 1**,已有 running 任务时返回错误 |
|
||||
| GET | `/admin/trainings` | 训练任务列表:`page/size` 分页,按下发时间倒序,含 status/进度/指标 |
|
||||
| GET | `/admin/trainings/detail` | 任务详情 `{"id":1}`:参数快照 + 进度 + 指标 + 日志尾部 |
|
||||
| POST | `/admin/trainings/cancel` | 取消训练 `{"id":1}`(仅 running):杀训练进程,状态置 failed(记录 error) |
|
||||
| POST | `/admin/trainings/publish` | 发布为最新模型 `{"id":1,"notes":"..."}`(仅 success):tflite 拷为 `workspace/model-latest.tflite`(原子覆盖)+ sha256/大小 → 插入 `model_version`(版本号递增 m1.0.0 → m1.0.1)+ 旧版 `is_latest=0` |
|
||||
| GET | `/admin/label-workbench` | 标注工作台数据 `{"datasetId":1}`:数据集全部图片 + 每张标注框(`boxes`,labels_json 全量)——无历史标注任务时详情页工作台的数据源 |
|
||||
|
||||
管理页面(订单/授权)由 `server_admin/` 构建产物提供,访问 `http://<host>/admin/`。金额均为整数分,前端展示 ÷100 转元。
|
||||
管理页面由 `server_admin/` 构建产物提供,访问 `http://<host>/admin/`。金额均为整数分,前端展示 ÷100 转元。
|
||||
|
||||
### GET /api/v1/models
|
||||
|
||||
模型目录(需 Bearer token,客户端模型管理页拉取):返回**全部数据集当前生效模型**。响应 `data`:
|
||||
|
||||
```json
|
||||
{
|
||||
"list": [
|
||||
{"datasetId": 1, "datasetName": "pheasant", "version": "m1.2.0", "labels": ["pheasant", "suspect"],
|
||||
"sizeBytes": 6400000, "sha256": "ab12...", "notes": "修复小目标漏检", "publishedAt": "2026-08-26T10:00:00+08:00",
|
||||
"downloadUrl": "/download/models/pheasant/latest.tflite"}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
- 只返回 `is_latest=1` 的模型(每数据集至多一条);无任何发布模型时 `list` 为空数组
|
||||
- 下载地址由客户端拼 `apiBaseUrl` 访问;下载文件 sha256 校验,类别名数组 `labels` 用于多模型合并推理展示
|
||||
|
||||
## 使用说明
|
||||
|
||||
1. 配置 `config.yml`:监听端口、数据库路径、登录 token 签名密钥 `auth.secret`(必填,换值即全员下线)、套餐 `plans` 节点、微信支付(appid/mchid/商户私钥/证书序列号/APIv3 密钥)、支付宝(appid/应用私钥/支付宝公钥)、管理端 `admin.token`;SQLite 库由服务启动时自动建表并迁移,无需手工初始化
|
||||
1. 配置 `config.yml`:监听端口、数据库路径、登录 token 签名密钥 `auth.secret`(必填,换值即全员下线)、套餐 `plans` 节点、微信支付(appid/mchid/商户私钥/证书序列号/APIv3 密钥)、支付宝(appid/应用私钥/支付宝公钥)、管理端 `admin.token`;模型训练相关节点:`training`(训练通道 mode=subprocess/ssh、ssh 连接信息、训练机工作目录/venv/数据集目录、并发度 1、超时)、`imageGen`(AI 生成图片 provider,默认 dashscope + apiKey + model qwen-image-3.0)、`localAi`(二期标注用 RF-DETR 服务地址);SQLite 库由服务启动时自动建表并迁移,无需手工初始化
|
||||
2. `go build ./...` 编译验证
|
||||
3. 本地运行 `go run main.go`;服务层白盒测试:`GF_GCFG_FILE=biz/service/testdata/config.yml go test ./biz/service/`(独立测试库,见 `biz/service/testdata/`)
|
||||
4. 后台管理端:`cd server_admin && npm run build`(构建产物输出到 `server/admin_dist/`,由后端 `/admin/` 托管);开发联调 `npm run dev`(Vite 代理 `/api` → `:8080`)。首次访问 `/admin/` 进入登录页,输入 `config.yml admin.token` 对应的管理 token(存浏览器 localStorage,随请求携带;token 不内嵌构建产物)
|
||||
|
||||
+1
-1
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
@@ -4,8 +4,8 @@
|
||||
<meta charset="UTF-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>视野管理端</title>
|
||||
<script type="module" crossorigin src="/admin/assets/index-gy5ysx_E.js"></script>
|
||||
<link rel="stylesheet" crossorigin href="/admin/assets/index-DsTNKZnX.css">
|
||||
<script type="module" crossorigin src="/admin/assets/index-jKHdJYxb.js"></script>
|
||||
<link rel="stylesheet" crossorigin href="/admin/assets/index-B0ehGsu8.css">
|
||||
</head>
|
||||
<body>
|
||||
<div id="app"></div>
|
||||
|
||||
@@ -5,6 +5,12 @@ package consts
|
||||
const (
|
||||
TablePaymentOrder = "payment_order"
|
||||
TableLicense = "license"
|
||||
TableAppVersion = "app_version"
|
||||
TableDataset = "dataset"
|
||||
TableDatasetImage = "dataset_image"
|
||||
TableTraining = "model_training"
|
||||
TableModelVersion = "model_version"
|
||||
TableLabelTask = "label_task"
|
||||
|
||||
// 订单状态机 created → paid(closed 仅超时/失败关闭)
|
||||
OrderStatusCreated = "created"
|
||||
@@ -19,4 +25,24 @@ const (
|
||||
|
||||
// 回调 IO 池默认并发度(被 config.yml payment.poolSize 覆盖)
|
||||
PaymentPoolDefaultSize = 16
|
||||
|
||||
// 预标注(逐张调 RF-DETR)池默认并发度(被 config.yml labelTask.poolSize 覆盖)
|
||||
LabelPoolDefaultSize = 4
|
||||
|
||||
// 训练任务状态机 running → success/failed
|
||||
TrainingStatusRunning = "running"
|
||||
TrainingStatusSuccess = "success"
|
||||
TrainingStatusFailed = "failed"
|
||||
|
||||
// 数据集状态 building → labeled → synced(synced = 已同步训练机)
|
||||
DatasetStatusBuilding = "building"
|
||||
DatasetStatusLabeled = "labeled"
|
||||
DatasetStatusSynced = "synced"
|
||||
|
||||
// 标注任务状态
|
||||
LabelTaskRunning = "running"
|
||||
LabelTaskDone = "done"
|
||||
|
||||
// 模型版本号前缀(m1.0.0),同数据集内递增
|
||||
ModelVersionPrefix = "m"
|
||||
)
|
||||
|
||||
@@ -2,18 +2,25 @@ package controller
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/url"
|
||||
|
||||
"github.com/gogf/gf/v2/net/ghttp"
|
||||
|
||||
"observer-server/biz/model/dto"
|
||||
"observer-server/biz/service"
|
||||
)
|
||||
|
||||
// cAdmin 后台管理端接口:薄透传;service 按表归入各表文件
|
||||
// (订单→service.Order、授权→service.License)。
|
||||
// (订单→service.Order、授权→service.License、数据集→service.Dataset、
|
||||
// 训练→service.Training、模型版本→service.ModelVersion、标注→service.LabelTask)。
|
||||
// 独立 controller 聚合,避免同一 controller 挂客户端组 + 管理组时路由重复注册。
|
||||
type cAdmin struct{}
|
||||
|
||||
var Admin = &cAdmin{}
|
||||
|
||||
// ---------- 订单/授权/App 版本(既有) ----------
|
||||
|
||||
// ListOrders 订单列表
|
||||
func (c *cAdmin) ListOrders(ctx context.Context, req *dto.AdminOrderListReq) (*dto.AdminOrderListRes, error) {
|
||||
return service.Order.AdminListOrders(ctx, req)
|
||||
@@ -33,3 +40,165 @@ func (c *cAdmin) Grant(ctx context.Context, req *dto.AdminGrantReq) (*dto.AdminG
|
||||
func (c *cAdmin) Revoke(ctx context.Context, req *dto.AdminRevokeReq) (*dto.AdminRevokeRes, error) {
|
||||
return service.License.AdminRevoke(ctx, req)
|
||||
}
|
||||
|
||||
// Remark 写账号备注
|
||||
func (c *cAdmin) Remark(ctx context.Context, req *dto.AdminRemarkReq) (*dto.AdminRemarkRes, error) {
|
||||
return service.License.AdminRemark(ctx, req)
|
||||
}
|
||||
|
||||
// ListAppVersions 版本记录列表
|
||||
func (c *cAdmin) ListAppVersions(ctx context.Context, req *dto.AdminAppVersionListReq) (*dto.AdminAppVersionListRes, error) {
|
||||
return service.AppVersion.AdminListVersions(ctx, req)
|
||||
}
|
||||
|
||||
// AddAppVersion 新增版本记录
|
||||
func (c *cAdmin) AddAppVersion(ctx context.Context, req *dto.AdminAppVersionAddReq) (*dto.AdminAppVersionAddRes, error) {
|
||||
return service.AppVersion.AdminAddVersion(ctx, req)
|
||||
}
|
||||
|
||||
// DeleteAppVersion 删除版本记录
|
||||
func (c *cAdmin) DeleteAppVersion(ctx context.Context, req *dto.AdminAppVersionDeleteReq) (*dto.AdminAppVersionDeleteRes, error) {
|
||||
return service.AppVersion.AdminDeleteVersion(ctx, req)
|
||||
}
|
||||
|
||||
// ---------- 数据集管理 ----------
|
||||
|
||||
// ListDatasets 数据集列表
|
||||
func (c *cAdmin) ListDatasets(ctx context.Context, req *dto.AdminDatasetListReq) (*dto.AdminDatasetListRes, error) {
|
||||
return service.Dataset.AdminListDatasets(ctx, req)
|
||||
}
|
||||
|
||||
// CreateDataset 新建数据集
|
||||
func (c *cAdmin) CreateDataset(ctx context.Context, req *dto.AdminDatasetCreateReq) (*dto.AdminDatasetCreateRes, error) {
|
||||
return service.Dataset.AdminCreateDataset(ctx, req)
|
||||
}
|
||||
|
||||
// UpdateDataset 更新数据集配置(封面/描述/AI 标注端点/训练机 SSH)
|
||||
func (c *cAdmin) UpdateDataset(ctx context.Context, req *dto.AdminDatasetUpdateReq) (*dto.AdminDatasetUpdateRes, error) {
|
||||
return service.Dataset.AdminUpdateDataset(ctx, req)
|
||||
}
|
||||
|
||||
// DeleteDataset 删除数据集
|
||||
func (c *cAdmin) DeleteDataset(ctx context.Context, req *dto.AdminDatasetDeleteReq) (*dto.AdminDatasetDeleteRes, error) {
|
||||
return service.Dataset.AdminDeleteDataset(ctx, req)
|
||||
}
|
||||
|
||||
// UploadImages 上传图片
|
||||
func (c *cAdmin) UploadImages(ctx context.Context, req *dto.AdminDatasetUploadReq) (*dto.AdminDatasetUploadRes, error) {
|
||||
return service.Dataset.AdminUploadImages(ctx, req)
|
||||
}
|
||||
|
||||
// GenerateImages AI 生成图片
|
||||
func (c *cAdmin) GenerateImages(ctx context.Context, req *dto.AdminDatasetGenerateReq) (*dto.AdminDatasetGenerateRes, error) {
|
||||
return service.Dataset.AdminGenerateImages(ctx, req)
|
||||
}
|
||||
|
||||
// ListImages 数据集图片列表
|
||||
func (c *cAdmin) ListImages(ctx context.Context, req *dto.AdminDatasetImagesReq) (*dto.AdminDatasetImagesRes, error) {
|
||||
return service.Dataset.AdminListImages(ctx, req)
|
||||
}
|
||||
|
||||
// DeleteImages 删除图片
|
||||
func (c *cAdmin) DeleteImages(ctx context.Context, req *dto.AdminDatasetImagesDeleteReq) (*dto.AdminDatasetImagesDeleteRes, error) {
|
||||
return service.Dataset.AdminDeleteImages(ctx, req)
|
||||
}
|
||||
|
||||
// Image 图片访问(直写响应体:service 校验归属返回路径,controller 输出二进制)
|
||||
func (c *cAdmin) Image(ctx context.Context, req *dto.AdminDatasetImageReq) (*dto.AdminDatasetImageRes, error) {
|
||||
path, err := service.Dataset.ImageFile(ctx, req.DatasetId, req.Filename)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
r := ghttp.RequestFromCtx(ctx)
|
||||
r.Response.ServeFile(path)
|
||||
return &dto.AdminDatasetImageRes{}, nil
|
||||
}
|
||||
|
||||
// UploadCover 上传数据集封面
|
||||
func (c *cAdmin) UploadCover(ctx context.Context, req *dto.AdminDatasetCoverUploadReq) (*dto.AdminDatasetCoverUploadRes, error) {
|
||||
return service.Dataset.AdminUploadCover(ctx, req)
|
||||
}
|
||||
|
||||
// Cover 封面访问(直写响应体)
|
||||
func (c *cAdmin) Cover(ctx context.Context, req *dto.AdminDatasetCoverReq) (*dto.AdminDatasetCoverRes, error) {
|
||||
path, err := service.Dataset.CoverFile(ctx, req.DatasetId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
r := ghttp.RequestFromCtx(ctx)
|
||||
r.Response.ServeFile(path)
|
||||
return &dto.AdminDatasetCoverRes{}, nil
|
||||
}
|
||||
|
||||
// DeleteCover 删除数据集封面
|
||||
func (c *cAdmin) DeleteCover(ctx context.Context, req *dto.AdminDatasetCoverDeleteReq) (*dto.AdminDatasetCoverDeleteRes, error) {
|
||||
return service.Dataset.AdminDeleteCover(ctx, req)
|
||||
}
|
||||
|
||||
// ExportDataset 导出数据集 zip(直写响应体下载)
|
||||
func (c *cAdmin) ExportDataset(ctx context.Context, req *dto.AdminDatasetExportReq) (*dto.AdminDatasetExportRes, error) {
|
||||
data, filename, err := service.Dataset.ExportZip(ctx, req.DatasetId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
r := ghttp.RequestFromCtx(ctx)
|
||||
r.Response.Header().Set("Content-Disposition",
|
||||
fmt.Sprintf("attachment; filename*=UTF-8''%s", url.PathEscape(filename)))
|
||||
r.Response.Header().Set("Content-Type", "application/zip")
|
||||
r.Response.Write(data)
|
||||
return &dto.AdminDatasetExportRes{}, nil
|
||||
}
|
||||
|
||||
// ---------- 训练编排 ----------
|
||||
|
||||
// ListTrainings 训练任务列表
|
||||
func (c *cAdmin) ListTrainings(ctx context.Context, req *dto.AdminTrainingListReq) (*dto.AdminTrainingListRes, error) {
|
||||
return service.Training.AdminListTrainings(ctx, req)
|
||||
}
|
||||
|
||||
// StartTraining 发起训练
|
||||
func (c *cAdmin) StartTraining(ctx context.Context, req *dto.AdminTrainingStartReq) (*dto.AdminTrainingStartRes, error) {
|
||||
return service.Training.AdminStartTraining(ctx, req)
|
||||
}
|
||||
|
||||
// TrainingDetail 训练任务详情
|
||||
func (c *cAdmin) TrainingDetail(ctx context.Context, req *dto.AdminTrainingDetailReq) (*dto.AdminTrainingDetailRes, error) {
|
||||
return service.Training.AdminTrainingDetail(ctx, req)
|
||||
}
|
||||
|
||||
// CancelTraining 取消训练
|
||||
func (c *cAdmin) CancelTraining(ctx context.Context, req *dto.AdminTrainingCancelReq) (*dto.AdminTrainingCancelRes, error) {
|
||||
return service.Training.AdminCancelTraining(ctx, req)
|
||||
}
|
||||
|
||||
// PublishTraining 发布模型版本
|
||||
func (c *cAdmin) PublishTraining(ctx context.Context, req *dto.AdminTrainingPublishReq) (*dto.AdminTrainingPublishRes, error) {
|
||||
return service.Training.AdminPublish(ctx, req)
|
||||
}
|
||||
|
||||
// ---------- 预标注 ----------
|
||||
|
||||
// ListLabelTasks 预标注任务列表
|
||||
func (c *cAdmin) ListLabelTasks(ctx context.Context, req *dto.AdminLabelTaskListReq) (*dto.AdminLabelTaskListRes, error) {
|
||||
return service.LabelTask.AdminListLabelTasks(ctx, req)
|
||||
}
|
||||
|
||||
// StartLabelTask 发起预标注
|
||||
func (c *cAdmin) StartLabelTask(ctx context.Context, req *dto.AdminLabelTaskStartReq) (*dto.AdminLabelTaskStartRes, error) {
|
||||
return service.LabelTask.AdminStartLabelTask(ctx, req)
|
||||
}
|
||||
|
||||
// LabelTaskDetail 预标注任务详情(工作台数据)
|
||||
func (c *cAdmin) LabelTaskDetail(ctx context.Context, req *dto.AdminLabelTaskDetailReq) (*dto.AdminLabelTaskDetailRes, error) {
|
||||
return service.LabelTask.AdminLabelTaskDetail(ctx, req)
|
||||
}
|
||||
|
||||
// LabelWorkbench 标注工作台数据(无历史任务时直接用数据集图片 + 已确认标注)
|
||||
func (c *cAdmin) LabelWorkbench(ctx context.Context, req *dto.AdminLabelWorkbenchReq) (*dto.AdminLabelWorkbenchRes, error) {
|
||||
return service.LabelTask.AdminLabelWorkbench(ctx, req)
|
||||
}
|
||||
|
||||
// SaveLabel 保存单张图标注
|
||||
func (c *cAdmin) SaveLabel(ctx context.Context, req *dto.AdminLabelSaveReq) (*dto.AdminLabelSaveRes, error) {
|
||||
return service.LabelTask.AdminLabelSave(ctx, req)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,18 @@
|
||||
package controller
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"observer-server/biz/model/dto"
|
||||
"observer-server/biz/service"
|
||||
)
|
||||
|
||||
// cAppVersion App 版本接口:公开组(无需登录态,旧版本/未登录用户均可检查更新)。
|
||||
type cAppVersion struct{}
|
||||
|
||||
var AppVersion = &cAppVersion{}
|
||||
|
||||
// GetUpdate 版本更新检查
|
||||
func (c *cAppVersion) GetUpdate(ctx context.Context, req *dto.AppUpdateReq) (*dto.AppUpdateRes, error) {
|
||||
return service.AppVersion.GetUpdate(ctx, req)
|
||||
}
|
||||
@@ -0,0 +1,19 @@
|
||||
package controller
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"observer-server/biz/model/dto"
|
||||
"observer-server/biz/service"
|
||||
)
|
||||
|
||||
// cModelCatalog 客户端模型目录:App 按需下载数据集模型(多模型并行推理 + 热更新)。
|
||||
// 挂登录态组(AuthRequired);下载走 /download 静态托管。
|
||||
type cModelCatalog struct{}
|
||||
|
||||
var ModelCatalog = &cModelCatalog{}
|
||||
|
||||
// List 模型目录:全部数据集当前生效模型
|
||||
func (c *cModelCatalog) List(ctx context.Context, req *dto.ModelCatalogReq) (*dto.ModelCatalogRes, error) {
|
||||
return service.ModelVersion.ClientCatalog(ctx)
|
||||
}
|
||||
@@ -0,0 +1,129 @@
|
||||
package dao
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/gogf/gf/v2/frame/g"
|
||||
"github.com/gogf/gf/v2/util/gconv"
|
||||
|
||||
"observer-server/biz/consts"
|
||||
"observer-server/biz/model/entity"
|
||||
"observer-server/common"
|
||||
)
|
||||
|
||||
// AppVersion App 版本表 DAO:写极低频、表极小,不做查询缓存;
|
||||
// 写链路经 common.Serial() 串行(与授权链路一致)。
|
||||
type appVersionDao struct{}
|
||||
|
||||
var AppVersion = &appVersionDao{}
|
||||
|
||||
func init() {
|
||||
ctx := context.Background()
|
||||
_, err := g.DB().Exec(ctx, `CREATE TABLE IF NOT EXISTS app_version (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
version TEXT NOT NULL UNIQUE,
|
||||
notes TEXT,
|
||||
created_at TEXT NOT NULL,
|
||||
updated_at TEXT NOT NULL
|
||||
)`)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
// v6 迁移:删除 url 列(下载地址改为固定文件 app.apkDir/observer-latest.apk,
|
||||
// 表内不再记录;新库建表已无此列直接跳过)
|
||||
cols, err := g.DB().GetAll(ctx, "PRAGMA table_info(app_version)")
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
for _, col := range cols {
|
||||
if gconv.String(col["name"]) == "url" {
|
||||
if _, err := g.DB().Exec(ctx, "ALTER TABLE app_version DROP COLUMN url"); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
g.Log().Warningf(ctx, "存量表 app_version 已迁移:删除 url 列")
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Insert 新增版本记录(version UNIQUE 由库兜底防重复;id 自增不显式写入)
|
||||
func (d *appVersionDao) Insert(ctx context.Context, m *entity.AppVersion) error {
|
||||
_, err := g.DB().Model(consts.TableAppVersion).Ctx(ctx).Data(g.Map{
|
||||
"version": m.Version,
|
||||
"notes": m.Notes,
|
||||
"created_at": m.CreatedAt,
|
||||
"updated_at": m.UpdatedAt,
|
||||
}).Insert()
|
||||
return err
|
||||
}
|
||||
|
||||
// DeleteByVersion 按版本号删除记录(文件保存失败时的补偿回滚)
|
||||
func (d *appVersionDao) DeleteByVersion(ctx context.Context, version string) error {
|
||||
_, err := g.DB().Model(consts.TableAppVersion).Ctx(ctx).
|
||||
Where("version", version).Delete()
|
||||
return err
|
||||
}
|
||||
|
||||
// GetById 按主键查询版本记录,不存在返回 nil
|
||||
func (d *appVersionDao) GetById(ctx context.Context, id int64) (*entity.AppVersion, error) {
|
||||
var e entity.AppVersion
|
||||
err := g.DB().Model(consts.TableAppVersion).Ctx(ctx).Where("id", id).Scan(&e)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return &e, nil
|
||||
}
|
||||
|
||||
// DeleteById 按主键删除版本记录
|
||||
func (d *appVersionDao) DeleteById(ctx context.Context, id int64) error {
|
||||
_, err := g.DB().Model(consts.TableAppVersion).Ctx(ctx).Where("id", id).Delete()
|
||||
return err
|
||||
}
|
||||
|
||||
// GetByVersion 按版本号查询(新增前查重,返回 nil 表示不存在)
|
||||
func (d *appVersionDao) GetByVersion(ctx context.Context, version string) (*entity.AppVersion, error) {
|
||||
var e entity.AppVersion
|
||||
err := g.DB().Model(consts.TableAppVersion).Ctx(ctx).Where("version", version).Scan(&e)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return &e, nil
|
||||
}
|
||||
|
||||
// Latest 最新一条版本记录(id 倒序,无记录返回 nil)
|
||||
func (d *appVersionDao) Latest(ctx context.Context) (*entity.AppVersion, error) {
|
||||
var e entity.AppVersion
|
||||
err := g.DB().Model(consts.TableAppVersion).Ctx(ctx).
|
||||
OrderDesc("id").Limit(1).Scan(&e)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return &e, nil
|
||||
}
|
||||
|
||||
// Page 管理端版本记录分页:按下发时间倒序
|
||||
func (d *appVersionDao) Page(ctx context.Context, page, size int) ([]*entity.AppVersion, int64, error) {
|
||||
base := g.DB().Model(consts.TableAppVersion).Ctx(ctx)
|
||||
total, err := base.Count()
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
var list []*entity.AppVersion
|
||||
err = base.OrderDesc("id").Limit((page-1)*size, size).Scan(&list)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return []*entity.AppVersion{}, int64(total), nil
|
||||
}
|
||||
return nil, 0, err
|
||||
}
|
||||
return list, int64(total), nil
|
||||
}
|
||||
@@ -0,0 +1,186 @@
|
||||
package dao
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strconv"
|
||||
|
||||
"github.com/gogf/gf/v2/database/gdb"
|
||||
"github.com/gogf/gf/v2/frame/g"
|
||||
"github.com/gogf/gf/v2/os/gtime"
|
||||
|
||||
"observer-server/biz/consts"
|
||||
"observer-server/biz/model/entity"
|
||||
"observer-server/common"
|
||||
)
|
||||
|
||||
// Dataset 数据集表 DAO:表小、读写低频,不做查询缓存;
|
||||
// 图片计数由 service 在增删图片时维护,查询走普通读。
|
||||
type datasetDao struct{}
|
||||
|
||||
var Dataset = &datasetDao{}
|
||||
|
||||
func init() {
|
||||
ctx := context.Background()
|
||||
_, err := g.DB().Exec(ctx, `CREATE TABLE IF NOT EXISTS dataset (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
name TEXT NOT NULL UNIQUE,
|
||||
source TEXT NOT NULL DEFAULT 'manual',
|
||||
image_count INTEGER NOT NULL DEFAULT 0,
|
||||
labeled_count INTEGER NOT NULL DEFAULT 0,
|
||||
status TEXT NOT NULL DEFAULT 'building',
|
||||
cover TEXT,
|
||||
description TEXT,
|
||||
ai_endpoint TEXT,
|
||||
ai_model TEXT,
|
||||
train_host TEXT,
|
||||
train_user TEXT,
|
||||
train_password TEXT,
|
||||
train_key TEXT,
|
||||
created_at TEXT NOT NULL,
|
||||
updated_at TEXT NOT NULL
|
||||
)`)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
}
|
||||
|
||||
// Insert 新建数据集(name UNIQUE 由库兜底)
|
||||
func (d *datasetDao) Insert(ctx context.Context, m *entity.Dataset) (int64, error) {
|
||||
res, err := g.DB().Model(consts.TableDataset).Ctx(ctx).Data(g.Map{
|
||||
"name": m.Name,
|
||||
"source": m.Source,
|
||||
"image_count": m.ImageCount,
|
||||
"labeled_count": m.LabeledCount,
|
||||
"status": m.Status,
|
||||
"created_at": m.CreatedAt,
|
||||
"updated_at": m.UpdatedAt,
|
||||
}).Insert()
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return res.LastInsertId()
|
||||
}
|
||||
|
||||
// GetById 按主键查询,不存在返回 nil
|
||||
func (d *datasetDao) GetById(ctx context.Context, id int64) (*entity.Dataset, error) {
|
||||
var e entity.Dataset
|
||||
err := g.DB().Model(consts.TableDataset).Ctx(ctx).Where("id", id).Scan(&e)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return &e, nil
|
||||
}
|
||||
|
||||
// GetByName 按名称查询(目录名即数据集名),不存在返回 nil
|
||||
func (d *datasetDao) GetByName(ctx context.Context, name string) (*entity.Dataset, error) {
|
||||
var e entity.Dataset
|
||||
err := g.DB().Model(consts.TableDataset).Ctx(ctx).Where("name", name).Scan(&e)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return &e, nil
|
||||
}
|
||||
|
||||
// UpdateCounters 更新图片数/已标注数/状态(imgDelta 增量;labeledCount ≥0 时覆盖写)
|
||||
func (d *datasetDao) UpdateCounters(ctx context.Context, id int64, imgDelta, labeledCount int64, status string) error {
|
||||
data := g.Map{"updated_at": gtime.Now()}
|
||||
if imgDelta != 0 {
|
||||
data["image_count"] = gdb.Raw("image_count + " + strconv.FormatInt(imgDelta, 10))
|
||||
}
|
||||
if labeledCount >= 0 {
|
||||
data["labeled_count"] = labeledCount
|
||||
}
|
||||
if status != "" {
|
||||
data["status"] = status
|
||||
}
|
||||
_, err := g.DB().Model(consts.TableDataset).Ctx(ctx).Where("id", id).Data(data).Update()
|
||||
return err
|
||||
}
|
||||
|
||||
// UpdateConfigs 更新数据集展示配置(封面/描述;仅更新非零字段,留空字段保留原值——
|
||||
// 前端提交整包配置,未填项不覆盖。AI 端点/训练机 SSH 为全局训练配置,不在此表维护)
|
||||
func (d *datasetDao) UpdateConfigs(ctx context.Context, id int64, m *entity.Dataset) error {
|
||||
data := g.Map{"updated_at": gtime.Now()}
|
||||
if m.Cover != "" {
|
||||
data["cover"] = m.Cover
|
||||
}
|
||||
if m.Description != "" {
|
||||
data["description"] = m.Description
|
||||
}
|
||||
_, err := g.DB().Model(consts.TableDataset).Ctx(ctx).Where("id", id).Data(data).Update()
|
||||
return err
|
||||
}
|
||||
|
||||
// ClearCover 清空封面字段(删除封面用;UpdateConfigs 空值不覆盖,无法复用)
|
||||
func (d *datasetDao) ClearCover(ctx context.Context, id int64) error {
|
||||
_, err := g.DB().Model(consts.TableDataset).Ctx(ctx).Where("id", id).Data(g.Map{"cover": "", "updated_at": gtime.Now()}).Update()
|
||||
return err
|
||||
}
|
||||
|
||||
// UpdateStatus 更新状态(building|labeled|synced)
|
||||
func (d *datasetDao) UpdateStatus(ctx context.Context, id int64, status string) error {
|
||||
_, err := g.DB().Model(consts.TableDataset).Ctx(ctx).Where("id", id).
|
||||
Data(g.Map{"status": status, "updated_at": gtime.Now()}).Update()
|
||||
return err
|
||||
}
|
||||
|
||||
// DeleteById 删除数据集记录
|
||||
func (d *datasetDao) DeleteById(ctx context.Context, id int64) error {
|
||||
_, err := g.DB().Model(consts.TableDataset).Ctx(ctx).Where("id", id).Delete()
|
||||
return err
|
||||
}
|
||||
|
||||
// PageByKeyword 数据集分页:名称模糊匹配(Like 命中目录名,数据库内仅元数据)
|
||||
func (d *datasetDao) PageByKeyword(ctx context.Context, keyword string, page, size int) ([]*entity.Dataset, int64, error) {
|
||||
base := g.DB().Model(consts.TableDataset).Ctx(ctx).WhereLike("name", "%"+keyword+"%")
|
||||
total, err := base.Count()
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
var list []*entity.Dataset
|
||||
err = base.OrderDesc("id").Limit((page-1)*size, size).Scan(&list)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return []*entity.Dataset{}, int64(total), nil
|
||||
}
|
||||
return nil, 0, err
|
||||
}
|
||||
return list, int64(total), nil
|
||||
}
|
||||
|
||||
// ListAll 全量数据集(模型/训练列表组装数据集名用,数据集表小)
|
||||
func (d *datasetDao) ListAll(ctx context.Context) ([]*entity.Dataset, error) {
|
||||
var list []*entity.Dataset
|
||||
err := g.DB().Model(consts.TableDataset).Ctx(ctx).OrderAsc("id").Scan(&list)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return []*entity.Dataset{}, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return list, nil
|
||||
}
|
||||
|
||||
// Page 数据集分页:创建时间倒序
|
||||
func (d *datasetDao) Page(ctx context.Context, page, size int) ([]*entity.Dataset, int64, error) {
|
||||
base := g.DB().Model(consts.TableDataset).Ctx(ctx)
|
||||
total, err := base.Count()
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
var list []*entity.Dataset
|
||||
err = base.OrderDesc("id").Limit((page-1)*size, size).Scan(&list)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return []*entity.Dataset{}, int64(total), nil
|
||||
}
|
||||
return nil, 0, err
|
||||
}
|
||||
return list, int64(total), nil
|
||||
}
|
||||
@@ -0,0 +1,141 @@
|
||||
package dao
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/gogf/gf/v2/frame/g"
|
||||
|
||||
"observer-server/biz/consts"
|
||||
"observer-server/biz/model/entity"
|
||||
"observer-server/common"
|
||||
)
|
||||
|
||||
// DatasetImage 数据集图片表 DAO:文件名唯一(防重名),按数据集查询;表小不做缓存。
|
||||
type datasetImageDao struct{}
|
||||
|
||||
var DatasetImage = &datasetImageDao{}
|
||||
|
||||
func init() {
|
||||
ctx := context.Background()
|
||||
_, err := g.DB().Exec(ctx, `CREATE TABLE IF NOT EXISTS dataset_image (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
dataset_id INTEGER NOT NULL,
|
||||
filename TEXT NOT NULL,
|
||||
source TEXT NOT NULL DEFAULT 'manual',
|
||||
prompt TEXT,
|
||||
labels_json TEXT,
|
||||
created_at TEXT NOT NULL,
|
||||
UNIQUE (dataset_id, filename)
|
||||
)`)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
// 存量库迁移:标注列缺失时追加(新库建表已含列,幂等跳过);candidates_json 已随 v10 删除
|
||||
common.EnsureColumn(ctx, consts.TableDatasetImage, "labels_json", "labels_json TEXT")
|
||||
}
|
||||
|
||||
// Insert 插入图片记录
|
||||
func (d *datasetImageDao) Insert(ctx context.Context, m *entity.DatasetImage) (int64, error) {
|
||||
res, err := g.DB().Model(consts.TableDatasetImage).Ctx(ctx).Data(g.Map{
|
||||
"dataset_id": m.DatasetId,
|
||||
"filename": m.Filename,
|
||||
"source": m.Source,
|
||||
"prompt": m.Prompt,
|
||||
"created_at": m.CreatedAt,
|
||||
}).Insert()
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return res.LastInsertId()
|
||||
}
|
||||
|
||||
// ListByDataset 某数据集全部图片(文件名倒序,详情页逐行展示用)
|
||||
func (d *datasetImageDao) ListByDataset(ctx context.Context, datasetId int64) ([]*entity.DatasetImage, error) {
|
||||
var list []*entity.DatasetImage
|
||||
err := g.DB().Model(consts.TableDatasetImage).Ctx(ctx).
|
||||
Where("dataset_id", datasetId).OrderDesc("filename").Scan(&list)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return []*entity.DatasetImage{}, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return list, nil
|
||||
}
|
||||
|
||||
// ListUnlabeledByDataset 未标注图片(labels_json 为空/null/'[]';自动补标轮次用)
|
||||
func (d *datasetImageDao) ListUnlabeledByDataset(ctx context.Context, datasetId int64) ([]*entity.DatasetImage, error) {
|
||||
var list []*entity.DatasetImage
|
||||
err := g.DB().Model(consts.TableDatasetImage).Ctx(ctx).
|
||||
Where("dataset_id", datasetId).
|
||||
Where("labels_json IS NULL OR labels_json = '' OR labels_json = '[]'").
|
||||
OrderAsc("id").Scan(&list)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return []*entity.DatasetImage{}, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return list, nil
|
||||
}
|
||||
|
||||
// GetByIds 按 id 批量取(删除图片定位文件用)
|
||||
func (d *datasetImageDao) GetByIds(ctx context.Context, ids []int64) ([]*entity.DatasetImage, error) {
|
||||
var list []*entity.DatasetImage
|
||||
err := g.DB().Model(consts.TableDatasetImage).Ctx(ctx).WhereIn("id", ids).Scan(&list)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return []*entity.DatasetImage{}, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return list, nil
|
||||
}
|
||||
|
||||
// GetByFilename 按文件名取(图片访问校验归属)
|
||||
func (d *datasetImageDao) GetByFilename(ctx context.Context, datasetId int64, filename string) (*entity.DatasetImage, error) {
|
||||
var e entity.DatasetImage
|
||||
err := g.DB().Model(consts.TableDatasetImage).Ctx(ctx).
|
||||
Where("dataset_id", datasetId).Where("filename", filename).Scan(&e)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return &e, nil
|
||||
}
|
||||
|
||||
// CountByDataset 某数据集图片数
|
||||
func (d *datasetImageDao) CountByDataset(ctx context.Context, datasetId int64) (int64, error) {
|
||||
n, err := g.DB().Model(consts.TableDatasetImage).Ctx(ctx).Where("dataset_id", datasetId).Count()
|
||||
return int64(n), err
|
||||
}
|
||||
|
||||
// UpdateLabels 覆写标注 JSON(整图粒度;空串=清空标注;AI 自动标注与人工保存共用)
|
||||
func (d *datasetImageDao) UpdateLabels(ctx context.Context, id int64, labelsJson string) error {
|
||||
_, err := g.DB().Model(consts.TableDatasetImage).Ctx(ctx).Where("id", id).
|
||||
Data(g.Map{"labels_json": labelsJson}).Update()
|
||||
return err
|
||||
}
|
||||
|
||||
// CountLabeledByDataset 某数据集已标注图片数(labels_json 非空数组的图片行数)
|
||||
func (d *datasetImageDao) CountLabeledByDataset(ctx context.Context, datasetId int64) (int64, error) {
|
||||
n, err := g.DB().Model(consts.TableDatasetImage).Ctx(ctx).
|
||||
Where("dataset_id", datasetId).
|
||||
Where("labels_json IS NOT NULL AND labels_json != '' AND labels_json != '[]'").
|
||||
Count()
|
||||
return int64(n), err
|
||||
}
|
||||
|
||||
// DeleteByIds 按 id 批量删除(标注随行删除)
|
||||
func (d *datasetImageDao) DeleteByIds(ctx context.Context, ids []int64) error {
|
||||
_, err := g.DB().Model(consts.TableDatasetImage).Ctx(ctx).WhereIn("id", ids).Delete()
|
||||
return err
|
||||
}
|
||||
|
||||
// DeleteByDataset 删除某数据集全部图片记录(删数据集联动)
|
||||
func (d *datasetImageDao) DeleteByDataset(ctx context.Context, datasetId int64) error {
|
||||
_, err := g.DB().Model(consts.TableDatasetImage).Ctx(ctx).Where("dataset_id", datasetId).Delete()
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,149 @@
|
||||
package dao
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/gogf/gf/v2/frame/g"
|
||||
"github.com/gogf/gf/v2/os/gtime"
|
||||
|
||||
"observer-server/biz/consts"
|
||||
"observer-server/biz/model/entity"
|
||||
"observer-server/common"
|
||||
)
|
||||
|
||||
// LabelTask 标注任务表 DAO:进度更新(done 计数)高频,单行 UPDATE 原子。
|
||||
type labelTaskDao struct{}
|
||||
|
||||
var LabelTask = &labelTaskDao{}
|
||||
|
||||
func init() {
|
||||
ctx := context.Background()
|
||||
_, err := g.DB().Exec(ctx, `CREATE TABLE IF NOT EXISTS label_task (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
dataset_id INTEGER NOT NULL,
|
||||
status TEXT NOT NULL DEFAULT 'running',
|
||||
total INTEGER NOT NULL DEFAULT 0,
|
||||
done INTEGER NOT NULL DEFAULT 0,
|
||||
boxes_file TEXT,
|
||||
filenames TEXT,
|
||||
error TEXT,
|
||||
created_at TEXT NOT NULL,
|
||||
finished_at TEXT
|
||||
)`)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
}
|
||||
|
||||
// Insert 创建标注任务,返回自增 id
|
||||
func (d *labelTaskDao) Insert(ctx context.Context, m *entity.LabelTask) (int64, error) {
|
||||
res, err := g.DB().Model(consts.TableLabelTask).Ctx(ctx).Data(g.Map{
|
||||
"dataset_id": m.DatasetId,
|
||||
"status": m.Status,
|
||||
"total": m.Total,
|
||||
"done": m.Done,
|
||||
"boxes_file": m.BoxesFile,
|
||||
"filenames": m.Filenames,
|
||||
"error": m.Error,
|
||||
"created_at": m.CreatedAt,
|
||||
"finished_at": m.FinishedAt,
|
||||
}).Insert()
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return res.LastInsertId()
|
||||
}
|
||||
|
||||
// GetById 按主键查询,不存在返回 nil
|
||||
func (d *labelTaskDao) GetById(ctx context.Context, id int64) (*entity.LabelTask, error) {
|
||||
var e entity.LabelTask
|
||||
err := g.DB().Model(consts.TableLabelTask).Ctx(ctx).Where("id", id).Scan(&e)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return &e, nil
|
||||
}
|
||||
|
||||
// ListRunning 全部 running 任务(服务重启恢复用)
|
||||
func (d *labelTaskDao) ListRunning(ctx context.Context) ([]*entity.LabelTask, error) {
|
||||
var list []*entity.LabelTask
|
||||
err := g.DB().Model(consts.TableLabelTask).Ctx(ctx).
|
||||
Where("status", consts.LabelTaskRunning).OrderAsc("id").Scan(&list)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return []*entity.LabelTask{}, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return list, nil
|
||||
}
|
||||
|
||||
// GetRunningByDataset 某数据集 running 任务(并发检查用)
|
||||
func (d *labelTaskDao) GetRunningByDataset(ctx context.Context, datasetId int64) (*entity.LabelTask, error) {
|
||||
var e entity.LabelTask
|
||||
err := g.DB().Model(consts.TableLabelTask).Ctx(ctx).
|
||||
Where("dataset_id", datasetId).Where("status", consts.LabelTaskRunning).Scan(&e)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return &e, nil
|
||||
}
|
||||
|
||||
// UpdateProgress 更新处理进度
|
||||
func (d *labelTaskDao) UpdateProgress(ctx context.Context, id int64, done int) error {
|
||||
_, err := g.DB().Model(consts.TableLabelTask).Ctx(ctx).Where("id", id).
|
||||
Data(g.Map{"done": done}).Update()
|
||||
return err
|
||||
}
|
||||
|
||||
// Finish 完成任务(done 状态 + 完成时间 + 候选框文件路径 + 失败原因)
|
||||
func (d *labelTaskDao) Finish(ctx context.Context, id int64, boxesFile, errMsg string) error {
|
||||
_, err := g.DB().Model(consts.TableLabelTask).Ctx(ctx).Where("id", id).
|
||||
Data(g.Map{"status": consts.LabelTaskDone, "finished_at": gtime.Now(), "boxes_file": boxesFile, "error": errMsg}).Update()
|
||||
return err
|
||||
}
|
||||
|
||||
// LatestByDataset 某数据集最近一次任务(无则 nil)
|
||||
func (d *labelTaskDao) LatestByDataset(ctx context.Context, datasetId int64) (*entity.LabelTask, error) {
|
||||
var e entity.LabelTask
|
||||
err := g.DB().Model(consts.TableLabelTask).Ctx(ctx).
|
||||
Where("dataset_id", datasetId).OrderDesc("id").Limit(1).Scan(&e)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return &e, nil
|
||||
}
|
||||
|
||||
// DeleteByDataset 删除数据集关联的全部标注任务(数据集删除时级联清理孤儿记录)
|
||||
func (d *labelTaskDao) DeleteByDataset(ctx context.Context, datasetId int64) error {
|
||||
_, err := g.DB().Model(consts.TableLabelTask).Ctx(ctx).
|
||||
Where("dataset_id", datasetId).Delete()
|
||||
return err
|
||||
}
|
||||
|
||||
// Page 标注任务分页:创建时间倒序
|
||||
func (d *labelTaskDao) Page(ctx context.Context, page, size int) ([]*entity.LabelTask, int64, error) {
|
||||
base := g.DB().Model(consts.TableLabelTask).Ctx(ctx)
|
||||
total, err := base.Count()
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
var list []*entity.LabelTask
|
||||
err = base.OrderDesc("id").Limit((page-1)*size, size).Scan(&list)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return []*entity.LabelTask{}, int64(total), nil
|
||||
}
|
||||
return nil, 0, err
|
||||
}
|
||||
return list, int64(total), nil
|
||||
}
|
||||
@@ -24,12 +24,14 @@ func init() {
|
||||
phone_num TEXT PRIMARY KEY,
|
||||
password TEXT NOT NULL,
|
||||
expires_at TEXT,
|
||||
remark TEXT,
|
||||
created_at TEXT NOT NULL,
|
||||
updated_at TEXT NOT NULL
|
||||
)`)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
common.EnsureColumn(ctx, consts.TableLicense, "remark", "remark TEXT")
|
||||
}
|
||||
|
||||
// GetByPhone 按手机号查询(走缓存,键含手机号)
|
||||
@@ -108,6 +110,15 @@ func (d *licenseDao) ClearAuthInTx(ctx context.Context, tx gdb.TX, phone string)
|
||||
return err
|
||||
}
|
||||
|
||||
// SetRemarkInTx 事务内写备注(管理端 remark,空串即清空):不覆盖密码/授权字段
|
||||
func (d *licenseDao) SetRemarkInTx(ctx context.Context, tx gdb.TX, phone, remark string) error {
|
||||
_, err := g.DB().Model(consts.TableLicense).Ctx(ctx).TX(tx).
|
||||
Where("phone_num", phone).
|
||||
Data(gdb.Map{"remark": remark, "updated_at": gtime.Now()}).
|
||||
Update()
|
||||
return err
|
||||
}
|
||||
|
||||
// PageByPhone 管理端账号分页查询:phone 模糊筛选,按更新时间倒序;
|
||||
// 列表不缓存(单点 GetByPhone 仍走缓存,写后清缓存不变)。
|
||||
func (d *licenseDao) PageByPhone(ctx context.Context, phone string, page, size int) ([]*entity.License, int64, error) {
|
||||
|
||||
@@ -0,0 +1,233 @@
|
||||
package dao
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/gogf/gf/v2/frame/g"
|
||||
"github.com/gogf/gf/v2/os/gtime"
|
||||
|
||||
"observer-server/biz/consts"
|
||||
"observer-server/biz/model/entity"
|
||||
"observer-server/common"
|
||||
)
|
||||
|
||||
// ModelTraining 训练任务表 DAO:进度/日志尾部为高频更新(独立小事务,不走 Serial 串行,
|
||||
// 单行 UPDATE 天然原子,无并发写冲突);状态流转(发起/结束)走 service 单写者。
|
||||
type modelTrainingDao struct{}
|
||||
|
||||
var Training = &modelTrainingDao{}
|
||||
|
||||
func init() {
|
||||
ctx := context.Background()
|
||||
_, err := g.DB().Exec(ctx, `CREATE TABLE IF NOT EXISTS model_training (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
name TEXT NOT NULL,
|
||||
status TEXT NOT NULL DEFAULT 'running',
|
||||
dataset_id INTEGER NOT NULL,
|
||||
imgsz INTEGER NOT NULL DEFAULT 704,
|
||||
epochs INTEGER NOT NULL DEFAULT 150,
|
||||
batch INTEGER NOT NULL DEFAULT 16,
|
||||
device TEXT NOT NULL DEFAULT '0',
|
||||
current_epoch INTEGER NOT NULL DEFAULT 0,
|
||||
total_epochs INTEGER NOT NULL DEFAULT 0,
|
||||
metrics TEXT,
|
||||
log_tail TEXT,
|
||||
pid INTEGER,
|
||||
error TEXT,
|
||||
started_at TEXT NOT NULL,
|
||||
finished_at TEXT,
|
||||
created_at TEXT NOT NULL
|
||||
)`)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
}
|
||||
|
||||
// Insert 创建训练任务,返回自增 id
|
||||
func (d *modelTrainingDao) Insert(ctx context.Context, m *entity.ModelTraining) (int64, error) {
|
||||
res, err := g.DB().Model(consts.TableTraining).Ctx(ctx).Data(g.Map{
|
||||
"name": m.Name,
|
||||
"status": m.Status,
|
||||
"dataset_id": m.DatasetId,
|
||||
"imgsz": m.Imgsz,
|
||||
"epochs": m.Epochs,
|
||||
"batch": m.Batch,
|
||||
"device": m.Device,
|
||||
"current_epoch": m.CurrentEpoch,
|
||||
"total_epochs": m.TotalEpochs,
|
||||
"metrics": m.Metrics,
|
||||
"log_tail": m.LogTail,
|
||||
"pid": m.Pid,
|
||||
"error": m.Error,
|
||||
"started_at": m.StartedAt,
|
||||
"finished_at": m.FinishedAt,
|
||||
"created_at": m.CreatedAt,
|
||||
}).Insert()
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return res.LastInsertId()
|
||||
}
|
||||
|
||||
// GetById 按主键查询,不存在返回 nil
|
||||
func (d *modelTrainingDao) GetById(ctx context.Context, id int64) (*entity.ModelTraining, error) {
|
||||
var e entity.ModelTraining
|
||||
err := g.DB().Model(consts.TableTraining).Ctx(ctx).Where("id", id).Scan(&e)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return &e, nil
|
||||
}
|
||||
|
||||
// UpdateProgress 更新进度/指标/日志尾部(训练轮询高频调用)
|
||||
func (d *modelTrainingDao) UpdateProgress(ctx context.Context, id int64, currentEpoch, totalEpochs int, metrics, logTail string) error {
|
||||
data := g.Map{}
|
||||
if currentEpoch > 0 {
|
||||
data["current_epoch"] = currentEpoch
|
||||
}
|
||||
if totalEpochs > 0 {
|
||||
data["total_epochs"] = totalEpochs
|
||||
}
|
||||
if metrics != "" {
|
||||
data["metrics"] = metrics
|
||||
}
|
||||
if logTail != "" {
|
||||
data["log_tail"] = logTail
|
||||
}
|
||||
if len(data) == 0 {
|
||||
return nil
|
||||
}
|
||||
_, err := g.DB().Model(consts.TableTraining).Ctx(ctx).Where("id", id).Data(data).Update()
|
||||
return err
|
||||
}
|
||||
|
||||
// UpdatePid 记录训练进程 pid(启动后写入,恢复扫描用)
|
||||
func (d *modelTrainingDao) UpdatePid(ctx context.Context, id int64, pid int) error {
|
||||
_, err := g.DB().Model(consts.TableTraining).Ctx(ctx).Where("id", id).
|
||||
Data(g.Map{"pid": pid}).Update()
|
||||
return err
|
||||
}
|
||||
|
||||
// PageByStatus 训练任务分页:按状态筛选(创建时间倒序)
|
||||
func (d *modelTrainingDao) PageByStatus(ctx context.Context, status string, page, size int) ([]*entity.ModelTraining, int64, error) {
|
||||
base := g.DB().Model(consts.TableTraining).Ctx(ctx).Where("status", status)
|
||||
total, err := base.Count()
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
var list []*entity.ModelTraining
|
||||
err = base.OrderDesc("id").Limit((page-1)*size, size).Scan(&list)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return []*entity.ModelTraining{}, int64(total), nil
|
||||
}
|
||||
return nil, 0, err
|
||||
}
|
||||
return list, int64(total), nil
|
||||
}
|
||||
|
||||
// Finish 结束任务(success/failed):状态 + 结束时间 + 指标 + 日志尾部 + 失败原因
|
||||
func (d *modelTrainingDao) Finish(ctx context.Context, id int64, status, metrics, logTail, errMsg string) error {
|
||||
data := g.Map{"status": status, "finished_at": gtime.Now(), "log_tail": logTail}
|
||||
if metrics != "" {
|
||||
data["metrics"] = metrics
|
||||
}
|
||||
if errMsg != "" {
|
||||
data["error"] = errMsg
|
||||
}
|
||||
_, err := g.DB().Model(consts.TableTraining).Ctx(ctx).Where("id", id).Data(data).Update()
|
||||
return err
|
||||
}
|
||||
|
||||
// Running 当前 running 任务(并发度 1 检查用;异常终态的失败任务同表记录,不算 running)
|
||||
func (d *modelTrainingDao) Running(ctx context.Context) (*entity.ModelTraining, error) {
|
||||
var e entity.ModelTraining
|
||||
err := g.DB().Model(consts.TableTraining).Ctx(ctx).
|
||||
Where("status", consts.TrainingStatusRunning).OrderAsc("id").Limit(1).Scan(&e)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return &e, nil
|
||||
}
|
||||
|
||||
// RunningByDataset 某数据集 running 任务(删数据集前检查用)
|
||||
func (d *modelTrainingDao) RunningByDataset(ctx context.Context, datasetId int64) (*entity.ModelTraining, error) {
|
||||
var e entity.ModelTraining
|
||||
err := g.DB().Model(consts.TableTraining).Ctx(ctx).
|
||||
Where("dataset_id", datasetId).Where("status", consts.TrainingStatusRunning).Limit(1).Scan(&e)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return &e, nil
|
||||
}
|
||||
|
||||
// ListRunning 全部 running 任务(Go 重启后恢复扫描用)
|
||||
func (d *modelTrainingDao) ListRunning(ctx context.Context) ([]*entity.ModelTraining, error) {
|
||||
var list []*entity.ModelTraining
|
||||
err := g.DB().Model(consts.TableTraining).Ctx(ctx).
|
||||
Where("status", consts.TrainingStatusRunning).OrderAsc("id").Scan(&list)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return []*entity.ModelTraining{}, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return list, nil
|
||||
}
|
||||
|
||||
// LatestByDatasets 批量取各数据集最新一条训练记录(列表卡片训练状态用;
|
||||
// IN 一次取回按 id 倒序,应用层按 dataset_id 去重;数据集表小、记录少,单次查询足够)
|
||||
func (d *modelTrainingDao) LatestByDatasets(ctx context.Context, datasetIds []int64) (map[int64]*entity.ModelTraining, error) {
|
||||
out := make(map[int64]*entity.ModelTraining)
|
||||
if len(datasetIds) == 0 {
|
||||
return out, nil
|
||||
}
|
||||
for start := 0; start < len(datasetIds); start += 100 {
|
||||
end := start + 100
|
||||
if end > len(datasetIds) {
|
||||
end = len(datasetIds)
|
||||
}
|
||||
var list []*entity.ModelTraining
|
||||
err := g.DB().Model(consts.TableTraining).Ctx(ctx).
|
||||
WhereIn("dataset_id", datasetIds[start:end]).OrderDesc("id").Scan(&list)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
continue
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
for _, t := range list {
|
||||
if _, ok := out[t.DatasetId]; !ok {
|
||||
out[t.DatasetId] = t
|
||||
}
|
||||
}
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// Page 训练任务分页:创建时间倒序
|
||||
func (d *modelTrainingDao) Page(ctx context.Context, page, size int) ([]*entity.ModelTraining, int64, error) {
|
||||
base := g.DB().Model(consts.TableTraining).Ctx(ctx)
|
||||
total, err := base.Count()
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
var list []*entity.ModelTraining
|
||||
err = base.OrderDesc("id").Limit((page-1)*size, size).Scan(&list)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return []*entity.ModelTraining{}, int64(total), nil
|
||||
}
|
||||
return nil, 0, err
|
||||
}
|
||||
return list, int64(total), nil
|
||||
}
|
||||
@@ -0,0 +1,110 @@
|
||||
package dao
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/gogf/gf/v2/frame/g"
|
||||
|
||||
"observer-server/biz/consts"
|
||||
"observer-server/biz/model/entity"
|
||||
"observer-server/common"
|
||||
)
|
||||
|
||||
// ModelVersion 模型版本表 DAO:每数据集独立版本序列;表小不做查询缓存。
|
||||
type modelVersionDao struct{}
|
||||
|
||||
var ModelVersion = &modelVersionDao{}
|
||||
|
||||
func init() {
|
||||
ctx := context.Background()
|
||||
_, err := g.DB().Exec(ctx, `CREATE TABLE IF NOT EXISTS model_version (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
dataset_id INTEGER NOT NULL,
|
||||
version TEXT NOT NULL,
|
||||
training_id INTEGER,
|
||||
artifact_file TEXT,
|
||||
metrics TEXT,
|
||||
labels TEXT NOT NULL,
|
||||
sha256 TEXT NOT NULL,
|
||||
size_bytes INTEGER NOT NULL,
|
||||
is_latest INTEGER NOT NULL DEFAULT 0,
|
||||
notes TEXT,
|
||||
created_at TEXT NOT NULL,
|
||||
UNIQUE (dataset_id, version)
|
||||
)`)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
}
|
||||
|
||||
// Insert 插入模型版本(is_latest 由 service 先置 0 再插新行置 1)
|
||||
func (d *modelVersionDao) Insert(ctx context.Context, m *entity.ModelVersion) (int64, error) {
|
||||
res, err := g.DB().Model(consts.TableModelVersion).Ctx(ctx).Data(g.Map{
|
||||
"dataset_id": m.DatasetId,
|
||||
"version": m.Version,
|
||||
"training_id": m.TrainingId,
|
||||
"artifact_file": m.ArtifactFile,
|
||||
"metrics": m.Metrics,
|
||||
"labels": m.Labels,
|
||||
"sha256": m.Sha256,
|
||||
"size_bytes": m.SizeBytes,
|
||||
"is_latest": m.IsLatest,
|
||||
"notes": m.Notes,
|
||||
"created_at": m.CreatedAt,
|
||||
}).Insert()
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return res.LastInsertId()
|
||||
}
|
||||
|
||||
// ClearLatest 某数据集所有版本置 is_latest=0(发布前调用)
|
||||
func (d *modelVersionDao) ClearLatest(ctx context.Context, datasetId int64) error {
|
||||
_, err := g.DB().Model(consts.TableModelVersion).Ctx(ctx).
|
||||
Where("dataset_id", datasetId).Data(g.Map{"is_latest": 0}).Update()
|
||||
return err
|
||||
}
|
||||
|
||||
// DeleteByDataset 删除某数据集全部版本记录(删数据集级联)
|
||||
func (d *modelVersionDao) DeleteByDataset(ctx context.Context, datasetId int64) error {
|
||||
_, err := g.DB().Model(consts.TableModelVersion).Ctx(ctx).Where("dataset_id", datasetId).Delete()
|
||||
return err
|
||||
}
|
||||
|
||||
// ListAllLatest 全部数据集的当前生效版本(模型目录/热更新目录)
|
||||
func (d *modelVersionDao) ListAllLatest(ctx context.Context) ([]*entity.ModelVersion, error) {
|
||||
var list []*entity.ModelVersion
|
||||
err := g.DB().Model(consts.TableModelVersion).Ctx(ctx).
|
||||
Where("is_latest", 1).OrderAsc("dataset_id").Scan(&list)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return []*entity.ModelVersion{}, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return list, nil
|
||||
}
|
||||
|
||||
// PageByDataset 某数据集版本分页:发布时间倒序
|
||||
func (d *modelVersionDao) PageByDataset(ctx context.Context, datasetId int64, page, size int) ([]*entity.ModelVersion, int64, error) {
|
||||
base := g.DB().Model(consts.TableModelVersion).Ctx(ctx).Where("dataset_id", datasetId)
|
||||
total, err := base.Count()
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
var list []*entity.ModelVersion
|
||||
err = base.OrderDesc("id").Limit((page-1)*size, size).Scan(&list)
|
||||
if err != nil {
|
||||
if common.IsNoRows(err) {
|
||||
return []*entity.ModelVersion{}, int64(total), nil
|
||||
}
|
||||
return nil, 0, err
|
||||
}
|
||||
return list, int64(total), nil
|
||||
}
|
||||
|
||||
// DeleteById 删除版本记录
|
||||
func (d *modelVersionDao) DeleteById(ctx context.Context, id int64) error {
|
||||
_, err := g.DB().Model(consts.TableModelVersion).Ctx(ctx).Where("id", id).Delete()
|
||||
return err
|
||||
}
|
||||
@@ -2,6 +2,7 @@ package dto
|
||||
|
||||
import (
|
||||
"github.com/gogf/gf/v2/frame/g"
|
||||
"github.com/gogf/gf/v2/net/ghttp"
|
||||
"github.com/gogf/gf/v2/os/gtime"
|
||||
)
|
||||
|
||||
@@ -48,10 +49,11 @@ type AdminLicenseListReq struct {
|
||||
Size int `json:"size" v:"integer|min:1|max:100" dc:"每页条数,默认 20"`
|
||||
}
|
||||
|
||||
// AdminLicenseItem 账号条目(未充值时 expiresAt 为 null)
|
||||
// AdminLicenseItem 账号条目(未充值时 expiresAt 为 null;remark 为管理端备注)
|
||||
type AdminLicenseItem struct {
|
||||
PhoneNum string `json:"phoneNum"`
|
||||
ExpiresAt *gtime.Time `json:"expiresAt"`
|
||||
Remark string `json:"remark"`
|
||||
CreatedAt *gtime.Time `json:"createdAt"`
|
||||
UpdatedAt *gtime.Time `json:"updatedAt"`
|
||||
}
|
||||
@@ -79,3 +81,53 @@ type AdminRevokeReq struct {
|
||||
}
|
||||
|
||||
type AdminRevokeRes struct{}
|
||||
|
||||
// AdminRemarkReq 写账号备注(空串即清空;仅管理端可见,客户端接口不含该字段)
|
||||
type AdminRemarkReq struct {
|
||||
g.Meta `path:"/licenses/remark" method:"post" summary:"写账号备注" tags:"管理端"`
|
||||
PhoneNum string `json:"phoneNum" v:"required|length:6,20" dc:"手机号"`
|
||||
Remark string `json:"remark" v:"length:0,200" dc:"备注内容,空串清空"`
|
||||
}
|
||||
|
||||
type AdminRemarkRes struct{}
|
||||
|
||||
// ---------- App 版本管理 ----------
|
||||
|
||||
// AdminAppVersionListReq 版本记录列表(按下发时间倒序)
|
||||
type AdminAppVersionListReq struct {
|
||||
g.Meta `path:"/app-versions" method:"get" summary:"版本记录列表" tags:"管理端"`
|
||||
Page int `json:"page" v:"integer|min:1" dc:"页码,默认 1"`
|
||||
Size int `json:"size" v:"integer|min:1|max:100" dc:"每页条数,默认 20"`
|
||||
}
|
||||
|
||||
// AdminAppVersionItem 版本记录条目
|
||||
type AdminAppVersionItem struct {
|
||||
Id int64 `json:"id"`
|
||||
Version string `json:"version"`
|
||||
Notes string `json:"notes"`
|
||||
CreatedAt *gtime.Time `json:"createdAt"`
|
||||
}
|
||||
|
||||
type AdminAppVersionListRes struct {
|
||||
Total int64 `json:"total"`
|
||||
List []*AdminAppVersionItem `json:"list"`
|
||||
}
|
||||
|
||||
// AdminAppVersionAddReq 下发新版本(multipart/form-data:notes + APK 文件;
|
||||
// 版本号从文件名识别,文件须命名为 observer-x.y.z.apk;version UNIQUE 防重复下发;
|
||||
// APK 覆盖保存固定文件,目录永远只有一个文件;检测到新版本即强制更新,无普通/强制之分)
|
||||
type AdminAppVersionAddReq struct {
|
||||
g.Meta `path:"/app-versions" method:"post" summary:"下发新版本" tags:"管理端" mime:"multipart/form-data"`
|
||||
Notes string `json:"notes" v:"length:0,500" dc:"更新说明"`
|
||||
File *ghttp.UploadFile `json:"file" dc:"APK 文件(文件名须为 observer-x.y.z.apk)"`
|
||||
}
|
||||
|
||||
type AdminAppVersionAddRes struct{}
|
||||
|
||||
// AdminAppVersionDeleteReq 删除版本记录(删最新版本联动删除 APK 文件,删历史版本仅删记录)
|
||||
type AdminAppVersionDeleteReq struct {
|
||||
g.Meta `path:"/app-versions/delete" method:"post" summary:"删除版本" tags:"管理端"`
|
||||
Id int64 `json:"id" v:"required|min:1" dc:"版本记录 id"`
|
||||
}
|
||||
|
||||
type AdminAppVersionDeleteRes struct{}
|
||||
|
||||
@@ -0,0 +1,22 @@
|
||||
package dto
|
||||
|
||||
import "github.com/gogf/gf/v2/frame/g"
|
||||
|
||||
// App 版本更新接口:公开组(无需登录态,旧版本/未登录用户均可检查更新)。
|
||||
// 仅 Android 客户端调用(iOS 不做版本下发);下载地址为固定静态路径
|
||||
// /download/observer-latest.apk(见 common/apk_store.go),接口不下发。
|
||||
|
||||
// AppUpdateReq 版本更新检查(Android 客户端启动时调用)
|
||||
type AppUpdateReq struct {
|
||||
g.Meta `path:"/app/update" method:"get" summary:"版本更新检查" tags:"版本"`
|
||||
}
|
||||
|
||||
// AppUpdateRes 服务器最新版本记录(无记录时全字段零值,客户端视为无需更新);
|
||||
// 检测到新版本即强制更新,无普通/强制之分。
|
||||
// Models 为模型热更新目录(全部数据集当前生效模型,与 GET /api/v1/models 同构);
|
||||
// 无发布模型时不返回该字段(旧 App 忽略、新 App 兼容旧服务器)。
|
||||
type AppUpdateRes struct {
|
||||
Version string `json:"version"`
|
||||
Notes string `json:"notes"`
|
||||
Models []*ModelCatalogItem `json:"models,omitempty"`
|
||||
}
|
||||
@@ -0,0 +1,378 @@
|
||||
package dto
|
||||
|
||||
import (
|
||||
"github.com/gogf/gf/v2/frame/g"
|
||||
"github.com/gogf/gf/v2/net/ghttp"
|
||||
"github.com/gogf/gf/v2/os/gtime"
|
||||
)
|
||||
|
||||
// 模型训练体系接口(数据集 → 标注 → 训练 → 模型版本),组前缀 /api/v1/admin。
|
||||
// 图片/标注/模型文件均在服务器磁盘(app.datasetDir),DB 只存元数据。
|
||||
// 图片访问与 zip 导出为直接写响应体场景,由 controller 以 *ghttp.Request 方法实现
|
||||
// (见 controller/admin.go 的 Image/ExportDataset 方法)。
|
||||
|
||||
// ---------- 数据集管理 ----------
|
||||
|
||||
// AdminDatasetListReq 数据集列表(创建时间倒序)
|
||||
type AdminDatasetListReq struct {
|
||||
g.Meta `path:"/datasets" method:"get" summary:"数据集列表" tags:"管理端"`
|
||||
Keyword string `json:"keyword" v:"length:0,50" dc:"名称模糊匹配"`
|
||||
Page int `json:"page" v:"integer|min:1" dc:"页码,默认 1"`
|
||||
Size int `json:"size" v:"integer|min:1|max:100" dc:"每页条数,默认 20"`
|
||||
}
|
||||
|
||||
// AdminDatasetItem 数据集条目(卡片展示用;AI 端点/训练机 SSH 走 config.yml 全局配置(localAi / training.ssh))
|
||||
type AdminDatasetItem struct {
|
||||
Id int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Source string `json:"source"` // manual | ai
|
||||
ImageCount int64 `json:"imageCount"`
|
||||
LabeledCount int64 `json:"labeledCount"`
|
||||
Status string `json:"status"` // building | labeled | synced
|
||||
Cover string `json:"cover"` // 封面文件名
|
||||
Description string `json:"description"` // 描述
|
||||
TrainingId int64 `json:"trainingId"` // 最新训练记录 id(发布/详情用)
|
||||
TrainingStatus string `json:"trainingStatus"` // 最新训练记录状态 running|success|failed|空
|
||||
TrainingCurrentEpoch int `json:"trainingCurrentEpoch"`
|
||||
TrainingTotalEpochs int `json:"trainingTotalEpochs"`
|
||||
CreatedAt *gtime.Time `json:"createdAt"`
|
||||
UpdatedAt *gtime.Time `json:"updatedAt"`
|
||||
}
|
||||
|
||||
type AdminDatasetListRes struct {
|
||||
Total int64 `json:"total"`
|
||||
List []*AdminDatasetItem `json:"list"`
|
||||
}
|
||||
|
||||
// AdminDatasetUpdateReq 更新数据集展示配置(空值字段不覆盖原值,保留现有配置;
|
||||
// AI 端点/训练机 SSH 走 config.yml 全局配置(localAi / training.ssh))
|
||||
type AdminDatasetUpdateReq struct {
|
||||
g.Meta `path:"/datasets/update" method:"post" summary:"更新数据集配置" tags:"管理端"`
|
||||
Id int64 `json:"id" v:"required|min:1" dc:"数据集 id"`
|
||||
Description string `json:"description" v:"length:0,500" dc:"描述(非空才更新)"`
|
||||
Cover string `json:"cover" v:"length:0,200" dc:"封面文件名(非空才更新)"`
|
||||
}
|
||||
|
||||
type AdminDatasetUpdateRes struct {
|
||||
Id int64 `json:"id"`
|
||||
}
|
||||
|
||||
// AdminDatasetCreateReq 新建数据集(名称同时是磁盘目录名:中文/字母/数字/下划线/短横线,唯一)
|
||||
type AdminDatasetCreateReq struct {
|
||||
g.Meta `path:"/datasets" method:"post" summary:"新建数据集" tags:"管理端"`
|
||||
Name string `json:"name" v:"required|regex:^[a-zA-Z0-9_一-龥-]+$|length:1,50" dc:"数据集名称(同是目录名,唯一)"`
|
||||
Source string `json:"source" v:"required|in:manual,ai" dc:"图片来源 manual|ai"`
|
||||
}
|
||||
|
||||
type AdminDatasetCreateRes struct {
|
||||
Id int64 `json:"id"`
|
||||
}
|
||||
|
||||
// AdminDatasetDeleteReq 删除数据集(删除图片文件+标注目录+记录;AI 生成图为付费资产,前端须带确认文案)
|
||||
type AdminDatasetDeleteReq struct {
|
||||
g.Meta `path:"/datasets/delete" method:"post" summary:"删除数据集" tags:"管理端"`
|
||||
Id int64 `json:"id" v:"required|min:1" dc:"数据集 id"`
|
||||
}
|
||||
|
||||
type AdminDatasetDeleteRes struct{}
|
||||
|
||||
// AdminDatasetUploadReq 上传图片(multipart 多文件;重名跳过并计数)
|
||||
type AdminDatasetUploadReq struct {
|
||||
g.Meta `path:"/datasets/upload" method:"post" summary:"上传图片" tags:"管理端" mime:"multipart/form-data"`
|
||||
DatasetId int64 `json:"datasetId" v:"required|min:1" dc:"数据集 id"`
|
||||
Files []*ghttp.UploadFile `json:"files" dc:"图片文件(jpg/jpeg/png,多选)"`
|
||||
}
|
||||
|
||||
type AdminDatasetUploadRes struct {
|
||||
Added int `json:"added"`
|
||||
Skipped int `json:"skipped"` // 重名跳过数
|
||||
}
|
||||
|
||||
// AdminDatasetGenerateReq AI 生成图片(同步执行;prompt 禁止含目标位置描述,前端模板+文案约束)
|
||||
type AdminDatasetGenerateReq struct {
|
||||
g.Meta `path:"/datasets/generate" method:"post" summary:"AI 生成图片" tags:"管理端"`
|
||||
DatasetId int64 `json:"datasetId" v:"required|min:1" dc:"数据集 id"`
|
||||
Prompt string `json:"prompt" v:"required|length:5,500" dc:"生成提示词(禁止含位置描述)"`
|
||||
Count int `json:"count" v:"required|integer|min:1|max:8" dc:"生成张数 1-8"`
|
||||
Size string `json:"size" v:"required|in:1152x2048,704x704" dc:"图片尺寸"`
|
||||
}
|
||||
|
||||
type AdminDatasetGenerateRes struct {
|
||||
Generated int `json:"generated"` // 成功张数(失败时可能 < count,已成功的图保留)
|
||||
}
|
||||
|
||||
// AdminDatasetImagesReq 数据集图片列表(标注工作台/网格预览用)
|
||||
type AdminDatasetImagesReq struct {
|
||||
g.Meta `path:"/datasets/images" method:"get" summary:"数据集图片列表" tags:"管理端"`
|
||||
DatasetId int64 `json:"datasetId" v:"required|min:1" dc:"数据集 id"`
|
||||
}
|
||||
|
||||
// AdminImageItem 图片条目(url 为管理端图片访问地址,前端拼 apiBaseUrl 使用)
|
||||
type AdminImageItem struct {
|
||||
Id int64 `json:"id"`
|
||||
Filename string `json:"filename"`
|
||||
Source string `json:"source"` // manual | ai
|
||||
Prompt string `json:"prompt"`
|
||||
Url string `json:"url"`
|
||||
CreatedAt *gtime.Time `json:"createdAt"`
|
||||
}
|
||||
|
||||
type AdminDatasetImagesRes struct {
|
||||
List []*AdminImageItem `json:"list"`
|
||||
}
|
||||
|
||||
// AdminDatasetImagesDeleteReq 删除图片(删文件+删记录;付费资产,前端须带确认文案)
|
||||
type AdminDatasetImagesDeleteReq struct {
|
||||
g.Meta `path:"/datasets/images/delete" method:"post" summary:"删除图片" tags:"管理端"`
|
||||
DatasetId int64 `json:"datasetId" v:"required|min:1" dc:"数据集 id"`
|
||||
Ids []int64 `json:"ids" v:"required|min-length:1" dc:"图片记录 id 列表"`
|
||||
}
|
||||
|
||||
type AdminDatasetImagesDeleteRes struct{}
|
||||
|
||||
// AdminDatasetImageReq 图片访问(直写响应体,例外场景:controller 经 service 校验归属后输出二进制)
|
||||
type AdminDatasetImageReq struct {
|
||||
g.Meta `path:"/datasets/image" method:"get" summary:"图片访问" tags:"管理端"`
|
||||
DatasetId int64 `json:"datasetId" v:"required|min:1" dc:"数据集 id"`
|
||||
Filename string `json:"filename" v:"required" dc:"图片文件名"`
|
||||
}
|
||||
|
||||
type AdminDatasetImageRes struct{}
|
||||
|
||||
// AdminDatasetExportReq 数据集 zip 导出(直写响应体下载;标注衔接与备份用)
|
||||
type AdminDatasetExportReq struct {
|
||||
g.Meta `path:"/datasets/export" method:"get" summary:"导出数据集 zip" tags:"管理端"`
|
||||
DatasetId int64 `json:"datasetId" v:"required|min:1" dc:"数据集 id"`
|
||||
}
|
||||
|
||||
type AdminDatasetExportRes struct{}
|
||||
|
||||
// AdminDatasetCoverUploadReq 上传数据集封面(multipart 单文件;覆盖旧封面,存图片目录 cover<ext>)
|
||||
type AdminDatasetCoverUploadReq struct {
|
||||
g.Meta `path:"/datasets/cover" method:"post" summary:"上传数据集封面" tags:"管理端" mime:"multipart/form-data"`
|
||||
DatasetId int64 `json:"datasetId" v:"required|min:1" dc:"数据集 id"`
|
||||
File *ghttp.UploadFile `json:"file" dc:"封面图片(jpg/jpeg/png,≤2MB)"`
|
||||
}
|
||||
|
||||
type AdminDatasetCoverUploadRes struct{}
|
||||
|
||||
// AdminDatasetCoverReq 封面访问(直写响应体,与 Image 同模式)
|
||||
type AdminDatasetCoverReq struct {
|
||||
g.Meta `path:"/datasets/cover" method:"get" summary:"数据集封面访问" tags:"管理端"`
|
||||
DatasetId int64 `json:"datasetId" v:"required|min:1" dc:"数据集 id"`
|
||||
}
|
||||
|
||||
type AdminDatasetCoverRes struct{}
|
||||
|
||||
// AdminDatasetCoverDeleteReq 删除数据集封面(删文件+清字段;列表卡片恢复占位图)
|
||||
type AdminDatasetCoverDeleteReq struct {
|
||||
g.Meta `path:"/datasets/cover/delete" method:"post" summary:"删除数据集封面" tags:"管理端"`
|
||||
DatasetId int64 `json:"datasetId" v:"required|min:1" dc:"数据集 id"`
|
||||
}
|
||||
|
||||
type AdminDatasetCoverDeleteRes struct{}
|
||||
|
||||
// ---------- 训练编排 ----------
|
||||
|
||||
// AdminTrainingListReq 训练任务列表(创建时间倒序)
|
||||
type AdminTrainingListReq struct {
|
||||
g.Meta `path:"/trainings" method:"get" summary:"训练任务列表" tags:"管理端"`
|
||||
Status string `json:"status" v:"in:running,success,failed" dc:"状态筛选"`
|
||||
Page int `json:"page" v:"integer|min:1" dc:"页码,默认 1"`
|
||||
Size int `json:"size" v:"integer|min:1|max:100" dc:"每页条数,默认 20"`
|
||||
}
|
||||
|
||||
// AdminTrainingItem 训练任务条目(列表不含日志尾部,详情接口返回)
|
||||
type AdminTrainingItem struct {
|
||||
Id int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
DatasetId int64 `json:"datasetId"`
|
||||
DatasetName string `json:"datasetName"`
|
||||
Status string `json:"status"` // running | success | failed
|
||||
Imgsz int `json:"imgsz"`
|
||||
Epochs int `json:"epochs"`
|
||||
Batch int `json:"batch"`
|
||||
Device string `json:"device"`
|
||||
CurrentEpoch int `json:"currentEpoch"`
|
||||
TotalEpochs int `json:"totalEpochs"`
|
||||
Metrics string `json:"metrics"`
|
||||
Error string `json:"error"`
|
||||
StartedAt *gtime.Time `json:"startedAt"`
|
||||
FinishedAt *gtime.Time `json:"finishedAt"`
|
||||
CreatedAt *gtime.Time `json:"createdAt"`
|
||||
}
|
||||
|
||||
type AdminTrainingListRes struct {
|
||||
Total int64 `json:"total"`
|
||||
List []*AdminTrainingItem `json:"list"`
|
||||
}
|
||||
|
||||
// AdminTrainingStartReq 发起训练(并发度 1:已有 running 任务时拒绝)
|
||||
type AdminTrainingStartReq struct {
|
||||
g.Meta `path:"/trainings" method:"post" summary:"发起训练" tags:"管理端"`
|
||||
DatasetId int64 `json:"datasetId" v:"required|min:1" dc:"数据集 id"`
|
||||
Name string `json:"name" v:"required|length:1,50" dc:"任务名称"`
|
||||
Imgsz int `json:"imgsz" v:"integer|min:64|max:2048" dc:"输入尺寸,默认 704"`
|
||||
Epochs int `json:"epochs" v:"integer|min:1|max:1000" dc:"训练轮数,默认 150"`
|
||||
Batch int `json:"batch" v:"integer|min:1|max:128" dc:"批大小,默认 16"`
|
||||
Device string `json:"device" v:"length:0,16" dc:"设备,默认 0"`
|
||||
}
|
||||
|
||||
type AdminTrainingStartRes struct {
|
||||
Id int64 `json:"id"`
|
||||
}
|
||||
|
||||
// AdminTrainingDetailReq 训练任务详情(含日志尾部)
|
||||
type AdminTrainingDetailReq struct {
|
||||
g.Meta `path:"/trainings/detail" method:"get" summary:"训练任务详情" tags:"管理端"`
|
||||
Id int64 `json:"id" v:"required|min:1" dc:"训练任务 id"`
|
||||
}
|
||||
|
||||
// AdminTrainingDetailRes 详情 = 条目 + 日志尾部
|
||||
type AdminTrainingDetailRes struct {
|
||||
AdminTrainingItem
|
||||
LogTail string `json:"logTail"`
|
||||
}
|
||||
|
||||
// AdminTrainingCancelReq 取消训练(杀进程,任务置 failed)
|
||||
type AdminTrainingCancelReq struct {
|
||||
g.Meta `path:"/trainings/cancel" method:"post" summary:"取消训练" tags:"管理端"`
|
||||
Id int64 `json:"id" v:"required|min:1" dc:"训练任务 id"`
|
||||
}
|
||||
|
||||
type AdminTrainingCancelRes struct{}
|
||||
|
||||
// AdminTrainingPublishReq 发布模型版本(仅 success 任务;按数据集版本号 m<major>.<minor>.<patch> 自增)
|
||||
type AdminTrainingPublishReq struct {
|
||||
g.Meta `path:"/trainings/publish" method:"post" summary:"发布模型版本" tags:"管理端"`
|
||||
Id int64 `json:"id" v:"required|min:1" dc:"训练任务 id"`
|
||||
}
|
||||
|
||||
type AdminTrainingPublishRes struct {
|
||||
Version string `json:"version"`
|
||||
}
|
||||
|
||||
// ---------- 模型版本 ----------
|
||||
|
||||
// ---------- 预标注 ----------
|
||||
|
||||
// AdminLabelTaskListReq 预标注任务列表(创建时间倒序)
|
||||
type AdminLabelTaskListReq struct {
|
||||
g.Meta `path:"/label-tasks" method:"get" summary:"预标注任务列表" tags:"管理端"`
|
||||
Page int `json:"page" v:"integer|min:1" dc:"页码,默认 1"`
|
||||
Size int `json:"size" v:"integer|min:1|max:100" dc:"每页条数,默认 20"`
|
||||
}
|
||||
|
||||
type AdminLabelTaskItem struct {
|
||||
Id int64 `json:"id"`
|
||||
DatasetId int64 `json:"datasetId"`
|
||||
DatasetName string `json:"datasetName"`
|
||||
Status string `json:"status"` // running | done
|
||||
Total int `json:"total"`
|
||||
Done int `json:"done"`
|
||||
Error string `json:"error"`
|
||||
CreatedAt *gtime.Time `json:"createdAt"`
|
||||
FinishedAt *gtime.Time `json:"finishedAt"`
|
||||
}
|
||||
|
||||
type AdminLabelTaskListRes struct {
|
||||
Total int64 `json:"total"`
|
||||
List []*AdminLabelTaskItem `json:"list"`
|
||||
}
|
||||
|
||||
// AdminLabelTaskStartReq 发起预标注(RF-DETR 扫描,候选框合并入 datasets/<name>/boxes.json)
|
||||
type AdminLabelTaskStartReq struct {
|
||||
g.Meta `path:"/label-tasks" method:"post" summary:"发起预标注" tags:"管理端"`
|
||||
DatasetId int64 `json:"datasetId" v:"required|min:1" dc:"数据集 id"`
|
||||
Filenames []string `json:"filenames" v:"length:0,500" dc:"选中图片列表(多选批量);缺省=全量扫描"`
|
||||
}
|
||||
|
||||
type AdminLabelTaskStartRes struct {
|
||||
Id int64 `json:"id"`
|
||||
}
|
||||
|
||||
// AdminLabelTaskDetailReq 预标注任务详情(每张图候选框 + 已确认标注,供工作台 canvas 叠框)
|
||||
type AdminLabelTaskDetailReq struct {
|
||||
g.Meta `path:"/label-tasks/detail" method:"get" summary:"预标注任务详情" tags:"管理端"`
|
||||
Id int64 `json:"id" v:"required|min:1" dc:"标注任务 id"`
|
||||
}
|
||||
|
||||
// AdminLabelBox 单框(归一化 xywh,坐标 0~1)
|
||||
type AdminLabelBox struct {
|
||||
Cx float64 `json:"cx"`
|
||||
Cy float64 `json:"cy"`
|
||||
W float64 `json:"w"`
|
||||
H float64 `json:"h"`
|
||||
Confidence float64 `json:"confidence"` // 人工框为 1
|
||||
Class int `json:"class"` // 0 确认 | 1 疑似
|
||||
}
|
||||
|
||||
// AdminLabelImageItem 工作台单张图:全部标注(labels_json,AI 自动标注与人工框同层,人工可修改/清理)
|
||||
type AdminLabelImageItem struct {
|
||||
Filename string `json:"filename"`
|
||||
Url string `json:"url"`
|
||||
Width int `json:"width"`
|
||||
Height int `json:"height"`
|
||||
Boxes []*AdminLabelBox `json:"boxes"`
|
||||
Labeled bool `json:"labeled"` // 有标注(AI 自动标注或人工保存)
|
||||
}
|
||||
|
||||
type AdminLabelTaskDetailRes struct {
|
||||
TaskId int64 `json:"taskId"`
|
||||
DatasetId int64 `json:"datasetId"`
|
||||
DatasetName string `json:"datasetName"`
|
||||
Status string `json:"status"`
|
||||
Total int `json:"total"`
|
||||
Done int `json:"done"`
|
||||
Error string `json:"error"`
|
||||
Images []*AdminLabelImageItem `json:"images"`
|
||||
}
|
||||
|
||||
// AdminLabelWorkbenchReq 标注工作台数据(无历史任务时直接用数据集图片 + 已确认标注,
|
||||
// 候选框读数据集目录 boxes.json)
|
||||
type AdminLabelWorkbenchReq struct {
|
||||
g.Meta `path:"/label-workbench" method:"get" summary:"标注工作台数据" tags:"管理端"`
|
||||
DatasetId int64 `json:"datasetId" v:"required|min:1" dc:"数据集 id"`
|
||||
}
|
||||
|
||||
type AdminLabelWorkbenchRes struct {
|
||||
DatasetId int64 `json:"datasetId"`
|
||||
DatasetName string `json:"datasetName"`
|
||||
Images []*AdminLabelImageItem `json:"images"`
|
||||
}
|
||||
|
||||
// AdminLabelSaveReq 保存单张图标注(工作台「保存本张」;覆写 dataset_image.labels_json)
|
||||
type AdminLabelSaveReq struct {
|
||||
g.Meta `path:"/label-tasks/save" method:"post" summary:"保存图片标注" tags:"管理端"`
|
||||
DatasetId int64 `json:"datasetId" v:"required|min:1" dc:"数据集 id"`
|
||||
Filename string `json:"filename" v:"required|length:1,200" dc:"图片文件名"`
|
||||
Boxes []*AdminLabelBox `json:"boxes" dc:"确认后的框列表(空=清空标注)"`
|
||||
}
|
||||
|
||||
type AdminLabelSaveRes struct {
|
||||
LabeledCount int64 `json:"labeledCount"` // 该数据集当前已标注图数
|
||||
}
|
||||
|
||||
// ---------- 客户端模型目录 ----------
|
||||
|
||||
// ModelCatalogReq 模型目录(客户端登录态):全部数据集的当前生效模型
|
||||
type ModelCatalogReq struct {
|
||||
g.Meta `path:"/models" method:"get" summary:"模型目录" tags:"客户端"`
|
||||
}
|
||||
|
||||
// ModelCatalogItem 客户端模型条目(App 按需下载,多模型并行推理合并)
|
||||
type ModelCatalogItem struct {
|
||||
DatasetId int64 `json:"datasetId"`
|
||||
DatasetName string `json:"datasetName"`
|
||||
Version string `json:"version"`
|
||||
Labels []string `json:"labels"` // 类别名,App 推理结果展示用
|
||||
SizeBytes int64 `json:"sizeBytes"`
|
||||
Sha256 string `json:"sha256"`
|
||||
Notes string `json:"notes"`
|
||||
PublishedAt *gtime.Time `json:"publishedAt"`
|
||||
DownloadUrl string `json:"downloadUrl"` // /download/models/<name>/latest.tflite
|
||||
}
|
||||
|
||||
type ModelCatalogRes struct {
|
||||
List []*ModelCatalogItem `json:"list"`
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
package entity
|
||||
|
||||
import "github.com/gogf/gf/v2/os/gtime"
|
||||
|
||||
// AppVersion App 版本记录:管理端下发,客户端启动时查询最新一条做版本比较。
|
||||
// 版本号 UNIQUE(防重复下发),记录仅新增不删除(不可变下发历史)。
|
||||
// 检测到新版本(服务器版本 > 本地版本)即强制更新,无普通/强制之分;
|
||||
// 下载地址不落表(固定文件 app.apkDir/observer-latest.apk,见 common/apk_store.go)。
|
||||
type AppVersion struct {
|
||||
Id int64 `json:"id" orm:"id" description:"自增主键"`
|
||||
Version string `json:"version" orm:"version" description:"语义化版本号 x.y.z"`
|
||||
Notes string `json:"notes" orm:"notes" description:"更新说明(客户端弹窗展示)"`
|
||||
CreatedAt *gtime.Time `json:"createdAt" orm:"created_at" description:"下发时间"`
|
||||
UpdatedAt *gtime.Time `json:"updatedAt" orm:"updated_at" description:"更新时间"`
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
package entity
|
||||
|
||||
import "github.com/gogf/gf/v2/os/gtime"
|
||||
|
||||
// Dataset 训练数据集:图片文件在 app.datasetDir/datasets/<name>/(DB 只存元数据 + 标注 labels_json)。
|
||||
// 每数据集训练一个模型。
|
||||
type Dataset struct {
|
||||
Id int64 `json:"id" orm:"id" description:"自增主键"`
|
||||
Name string `json:"name" orm:"name" description:"数据集名(≤50字,同是目录名,唯一)"`
|
||||
Source string `json:"source" orm:"source" description:"manual|ai"`
|
||||
ImageCount int64 `json:"imageCount" orm:"image_count" description:"图片数(冗余计数)"`
|
||||
LabeledCount int64 `json:"labeledCount" orm:"labeled_count" description:"已标注数"`
|
||||
Status string `json:"status" orm:"status" description:"building|labeled|synced"`
|
||||
Cover string `json:"cover" orm:"cover" description:"封面文件名(卡片展示)"`
|
||||
Description string `json:"description" orm:"description" description:"描述(卡片展示)"`
|
||||
CreatedAt *gtime.Time `json:"createdAt" orm:"created_at" description:"创建时间"`
|
||||
UpdatedAt *gtime.Time `json:"updatedAt" orm:"updated_at" description:"更新时间"`
|
||||
}
|
||||
@@ -0,0 +1,16 @@
|
||||
package entity
|
||||
|
||||
import "github.com/gogf/gf/v2/os/gtime"
|
||||
|
||||
// DatasetImage 数据集图片:文件在 datasets/<数据集名>/<filename>,
|
||||
// AI 生成图记录 prompt(付费资产,追溯用);标注为 JSON 数组存本表
|
||||
// (元素同 dto.AdminLabelBox:class/cx/cy/w/h/confidence,YOLO 归一化)。
|
||||
type DatasetImage struct {
|
||||
Id int64 `json:"id" orm:"id" description:"自增主键"`
|
||||
DatasetId int64 `json:"datasetId" orm:"dataset_id" description:"所属数据集"`
|
||||
Filename string `json:"filename" orm:"filename" description:"唯一文件名(防重名加时间戳后缀)"`
|
||||
Source string `json:"source" orm:"source" description:"manual|ai"`
|
||||
Prompt string `json:"prompt" orm:"prompt" description:"AI 生成图提示词"`
|
||||
LabelsJson string `json:"labelsJson" orm:"labels_json" description:"标注 JSON 数组(AI 自动标注与人工标注同存,null/''/'[]'=未标注)"`
|
||||
CreatedAt *gtime.Time `json:"createdAt" orm:"created_at" description:"创建时间"`
|
||||
}
|
||||
@@ -0,0 +1,19 @@
|
||||
package entity
|
||||
|
||||
import "github.com/gogf/gf/v2/os/gtime"
|
||||
|
||||
// LabelTask 标注任务:RF-DETR 预标注批量推理进度;候选框存
|
||||
// datasets/<数据集名>/boxes.json(YOLO 归一化 xywh + 置信度 + 建议类别)。
|
||||
// Filenames 为选中图片列表(JSON 数组串;空 = 全量扫描)。
|
||||
type LabelTask struct {
|
||||
Id int64 `json:"id" orm:"id" description:"自增主键"`
|
||||
DatasetId int64 `json:"datasetId" orm:"dataset_id" description:"所属数据集"`
|
||||
Status string `json:"status" orm:"status" description:"running|done"`
|
||||
Total int `json:"total" orm:"total" description:"待标注图片数"`
|
||||
Done int `json:"done" orm:"done" description:"已处理数"`
|
||||
BoxesFile string `json:"boxesFile" orm:"boxes_file" description:"候选框 JSON 相对路径"`
|
||||
Filenames string `json:"filenames" orm:"filenames" description:"选中图片JSON数组(空=全量)"`
|
||||
Error string `json:"error" orm:"error" description:"失败原因(部分失败/全部失败)"`
|
||||
CreatedAt *gtime.Time `json:"createdAt" orm:"created_at" description:"创建时间"`
|
||||
FinishedAt *gtime.Time `json:"finishedAt" orm:"finished_at" description:"完成时间"`
|
||||
}
|
||||
@@ -8,6 +8,7 @@ type License struct {
|
||||
PhoneNum string `json:"phoneNum" orm:"phone_num" description:"手机号账号(主键)"`
|
||||
Password string `json:"-" orm:"password" description:"bcrypt 加盐哈希"`
|
||||
ExpiresAt *gtime.Time `json:"expiresAt" orm:"expires_at" description:"到期时间(自然日,服务端时区),未充值 NULL"`
|
||||
Remark string `json:"remark" orm:"remark" description:"管理端备注(客户端不可见)"`
|
||||
CreatedAt *gtime.Time `json:"createdAt" orm:"created_at" description:"注册时间"`
|
||||
UpdatedAt *gtime.Time `json:"updatedAt" orm:"updated_at" description:"更新时间"`
|
||||
}
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
package entity
|
||||
|
||||
import "github.com/gogf/gf/v2/os/gtime"
|
||||
|
||||
// ModelTraining 训练任务:runner 启动训练进程,轮询解析 epoch 日志更新进度/指标,
|
||||
// 日志尾部截断存 log_tail;pid 用于取消与存活探测。
|
||||
type ModelTraining struct {
|
||||
Id int64 `json:"id" orm:"id" description:"自增主键"`
|
||||
Name string `json:"name" orm:"name" description:"任务名"`
|
||||
Status string `json:"status" orm:"status" description:"running|success|failed"`
|
||||
DatasetId int64 `json:"datasetId" orm:"dataset_id" description:"来源数据集"`
|
||||
Imgsz int `json:"imgsz" orm:"imgsz" description:"训练分辨率"`
|
||||
Epochs int `json:"epochs" orm:"epochs" description:"目标轮数"`
|
||||
Batch int `json:"batch" orm:"batch" description:"batch size"`
|
||||
Device string `json:"device" orm:"device" description:"训练设备"`
|
||||
CurrentEpoch int `json:"currentEpoch" orm:"current_epoch" description:"当前轮数"`
|
||||
TotalEpochs int `json:"totalEpochs" orm:"total_epochs" description:"实际总轮数(完成时写)"`
|
||||
Metrics string `json:"metrics" orm:"metrics" description:"JSON {p,r,map50}"`
|
||||
LogTail string `json:"logTail" orm:"log_tail" description:"日志尾部(截断)"`
|
||||
Pid int `json:"pid" orm:"pid" description:"训练进程 pid"`
|
||||
Error string `json:"error" orm:"error" description:"失败原因"`
|
||||
StartedAt *gtime.Time `json:"startedAt" orm:"started_at" description:"开始时间"`
|
||||
FinishedAt *gtime.Time `json:"finishedAt" orm:"finished_at" description:"结束时间"`
|
||||
CreatedAt *gtime.Time `json:"createdAt" orm:"created_at" description:"创建时间"`
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
package entity
|
||||
|
||||
import "github.com/gogf/gf/v2/os/gtime"
|
||||
|
||||
// ModelVersion 模型版本:每数据集独立版本序列(m1.0.0 递增,UNIQUE(dataset_id, version))。
|
||||
// 模型文件不落表:发布即写当前生效副本 models/<数据集>/latest.tflite(无存档回退机制);
|
||||
// labels 为类别名数组 JSON(App 多模型合并推理依赖)。
|
||||
type ModelVersion struct {
|
||||
Id int64 `json:"id" orm:"id" description:"自增主键"`
|
||||
DatasetId int64 `json:"datasetId" orm:"dataset_id" description:"所属数据集"`
|
||||
Version string `json:"version" orm:"version" description:"m1.0.0 递增"`
|
||||
TrainingId int64 `json:"trainingId" orm:"training_id" description:"来源训练任务"`
|
||||
ArtifactFile string `json:"artifactFile" orm:"artifact_file" description:"归档 zip 相对路径"`
|
||||
Metrics string `json:"metrics" orm:"metrics" description:"JSON 指标"`
|
||||
Labels string `json:"labels" orm:"labels" description:"JSON 类别名数组"`
|
||||
Sha256 string `json:"sha256" orm:"sha256" description:"tflite 文件校验"`
|
||||
SizeBytes int64 `json:"sizeBytes" orm:"size_bytes" description:"文件大小"`
|
||||
IsLatest int `json:"isLatest" orm:"is_latest" description:"1=该数据集当前生效"`
|
||||
Notes string `json:"notes" orm:"notes" description:"备注"`
|
||||
CreatedAt *gtime.Time `json:"createdAt" orm:"created_at" description:"发布时间"`
|
||||
}
|
||||
@@ -167,27 +167,6 @@ func TestAdminListOrders(t *testing.T) {
|
||||
t.Fatalf("total=%d len=%d, want 4/2", res.Total, len(res.List))
|
||||
}
|
||||
|
||||
wx, err := Order.AdminListOrders(ctx(), &dto.AdminOrderListReq{PhoneNum: phone, Channel: consts.ChannelWechat, Size: 10})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if wx.Total != 3 {
|
||||
t.Fatalf("wechat total = %d, want 3", wx.Total)
|
||||
}
|
||||
|
||||
one, err := Order.AdminListOrders(ctx(), &dto.AdminOrderListReq{OrderId: orderId})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if one.Total != 1 || len(one.List) != 1 || one.List[0].OrderId != orderId {
|
||||
t.Fatalf("filter by orderId = %+v", one)
|
||||
}
|
||||
if one.List[0].AmountCents != 1000 {
|
||||
t.Fatalf("amountCents = %d, want 1000", one.List[0].AmountCents)
|
||||
}
|
||||
if one.List[0].PhoneNum != phone {
|
||||
t.Fatalf("phoneNum = %s, want %s", one.List[0].PhoneNum, phone)
|
||||
}
|
||||
}
|
||||
|
||||
// TestListPlans 套餐列表:来自 testdata/config.yml plans 节点(配置驱动,无数据库表)
|
||||
|
||||
@@ -0,0 +1,172 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
|
||||
"github.com/gogf/gf/v2/errors/gerror"
|
||||
"github.com/gogf/gf/v2/frame/g"
|
||||
"github.com/gogf/gf/v2/net/ghttp"
|
||||
"github.com/gogf/gf/v2/os/gtime"
|
||||
|
||||
"observer-server/biz/dao"
|
||||
"observer-server/biz/model/dto"
|
||||
"observer-server/biz/model/entity"
|
||||
"observer-server/common"
|
||||
)
|
||||
|
||||
// appVersionService App 版本业务:客户端更新检查(取最新)、管理端下发(APK 上传 + 记录)。
|
||||
// APK 存固定文件名(覆盖式,目录永远只有一个文件),见 common/apk_store.go。
|
||||
type appVersionService struct{}
|
||||
|
||||
// apkNameRe APK 文件名规范:observer-x.y.z.apk(版本号从文件名识别)
|
||||
var apkNameRe = regexp.MustCompile(`^observer-(\d+\.\d+\.\d+)\.apk$`)
|
||||
|
||||
var AppVersion = &appVersionService{}
|
||||
|
||||
// GetUpdate 客户端版本更新检查(公开接口,无需登录态;仅 Android 调用):
|
||||
// 返回最新一条版本记录 + 模型热更新目录(多模型体系,与 GET /api/v1/models 同构)。
|
||||
// 无版本记录时返回空结构(客户端视为无需更新);无发布模型时 models 不返回。
|
||||
func (s *appVersionService) GetUpdate(ctx context.Context, req *dto.AppUpdateReq) (*dto.AppUpdateRes, error) {
|
||||
res := &dto.AppUpdateRes{}
|
||||
latest, err := dao.AppVersion.Latest(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if latest != nil {
|
||||
res.Version = latest.Version
|
||||
res.Notes = latest.Notes
|
||||
}
|
||||
catalog, err := ModelVersion.ClientCatalog(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(catalog.List) > 0 {
|
||||
res.Models = catalog.List
|
||||
}
|
||||
return res, nil
|
||||
}
|
||||
|
||||
// AdminListVersions 管理端版本记录分页(按下发时间倒序)
|
||||
func (s *appVersionService) AdminListVersions(ctx context.Context, req *dto.AdminAppVersionListReq) (*dto.AdminAppVersionListRes, error) {
|
||||
page, size := common.NormalizePage(req.Page, req.Size)
|
||||
list, total, err := dao.AppVersion.Page(ctx, page, size)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
items := make([]*dto.AdminAppVersionItem, 0, len(list))
|
||||
for _, v := range list {
|
||||
items = append(items, &dto.AdminAppVersionItem{
|
||||
Id: v.Id,
|
||||
Version: v.Version,
|
||||
Notes: v.Notes,
|
||||
CreatedAt: v.CreatedAt,
|
||||
})
|
||||
}
|
||||
return &dto.AdminAppVersionListRes{Total: total, List: items}, nil
|
||||
}
|
||||
|
||||
// AdminAddVersion 下发新版本:版本号从文件名识别(observer-x.y.z.apk),
|
||||
// 先单写者串行落库(查重 + UNIQUE 兜底),再保存 APK 覆盖固定文件;
|
||||
// 文件保存失败时删除记录补偿,保证「记录存在 ⟺ 文件存在」。
|
||||
func (s *appVersionService) AdminAddVersion(ctx context.Context, req *dto.AdminAppVersionAddReq) (*dto.AdminAppVersionAddRes, error) {
|
||||
if req.File == nil {
|
||||
return nil, gerror.NewCode(common.CodeApkInvalid)
|
||||
}
|
||||
version, err := parseVersionFromFilename(req.File.Filename)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
now := gtime.Now()
|
||||
err = common.Serial().Submit(ctx, func() error {
|
||||
exists, err := dao.AppVersion.GetByVersion(ctx, version)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if exists != nil {
|
||||
return gerror.NewCode(common.CodeVersionDuplicate)
|
||||
}
|
||||
return dao.AppVersion.Insert(ctx, &entity.AppVersion{
|
||||
Version: version,
|
||||
Notes: req.Notes,
|
||||
CreatedAt: now,
|
||||
UpdatedAt: now,
|
||||
})
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := saveApk(ctx, req.File); err != nil {
|
||||
if delErr := dao.AppVersion.DeleteByVersion(ctx, version); delErr != nil {
|
||||
g.Log().Errorf(ctx, "APK 保存失败后补偿删除记录失败: %+v", delErr)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return &dto.AdminAppVersionAddRes{}, nil
|
||||
}
|
||||
|
||||
// parseVersionFromFilename 从 APK 文件名识别版本号:observer-1.0.1.apk → 1.0.1。
|
||||
// 文件名是上传版本号的唯一事实来源(打包产物即按此命名)。
|
||||
func parseVersionFromFilename(name string) (string, error) {
|
||||
m := apkNameRe.FindStringSubmatch(filepath.Base(name))
|
||||
if m == nil {
|
||||
return "", gerror.NewCode(common.CodeApkInvalid,
|
||||
"APK 文件名须为 observer-x.y.z.apk 格式(如 observer-1.0.1.apk)")
|
||||
}
|
||||
return m[1], nil
|
||||
}
|
||||
|
||||
// AdminDeleteVersion 删除版本记录(Serial 串行,防与下发并发交错):
|
||||
// 删最新版本时联动删除 APK 文件(客户端 update 返回空不再提示、下载 404);
|
||||
// 删历史版本仅删记录、不动文件(固定文件永远对应最新版本)。
|
||||
// 先删记录再删文件:记录删除是核心操作,文件删除失败仅记日志不阻断
|
||||
// (记录没了客户端不会误提示更新,残留旧文件无害)。
|
||||
func (s *appVersionService) AdminDeleteVersion(ctx context.Context, req *dto.AdminAppVersionDeleteReq) (*dto.AdminAppVersionDeleteRes, error) {
|
||||
err := common.Serial().Submit(ctx, func() error {
|
||||
v, err := dao.AppVersion.GetById(ctx, req.Id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if v == nil {
|
||||
return gerror.NewCode(common.CodeVersionNotFound)
|
||||
}
|
||||
latest, err := dao.AppVersion.Latest(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := dao.AppVersion.DeleteById(ctx, req.Id); err != nil {
|
||||
return err
|
||||
}
|
||||
if latest != nil && latest.Id == v.Id {
|
||||
if err := os.Remove(common.ApkFilePath(ctx)); err != nil && !os.IsNotExist(err) {
|
||||
g.Log().Errorf(ctx, "删除最新版本 %s 后清理 APK 失败: %+v", v.Version, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &dto.AdminAppVersionDeleteRes{}, nil
|
||||
}
|
||||
|
||||
// saveApk 保存 APK:先落临时文件再原子重命名覆盖固定文件名,
|
||||
// 保证目录下永远只保留最新一个文件(异常中断不产生半截正式文件)。
|
||||
func saveApk(ctx context.Context, f *ghttp.UploadFile) error {
|
||||
dir := common.ApkDir(ctx)
|
||||
if err := os.MkdirAll(dir, 0o755); err != nil {
|
||||
return gerror.Wrap(err, "创建 APK 目录失败")
|
||||
}
|
||||
saved, err := f.Save(dir)
|
||||
if err != nil {
|
||||
_ = os.Remove(filepath.Join(dir, filepath.Base(f.Filename)))
|
||||
return gerror.Wrap(err, "APK 保存失败")
|
||||
}
|
||||
if err := os.Rename(filepath.Join(dir, saved), common.ApkFilePath(ctx)); err != nil {
|
||||
_ = os.Remove(filepath.Join(dir, saved))
|
||||
return gerror.Wrap(err, "APK 文件更新失败")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,229 @@
|
||||
package service
|
||||
|
||||
// App 版本管理白盒测试:下发(APK 上传 + 记录)、最新版本查询、重复版本拒绝、
|
||||
// APK 固定文件覆盖(目录永远只有一个文件)。
|
||||
// 运行方式同 admin_test.go:cd server && GF_GCFG_FILE=biz/service/testdata/config.yml go test ./biz/service/
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"mime/multipart"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gogf/gf/v2/net/ghttp"
|
||||
|
||||
"observer-server/biz/dao"
|
||||
"observer-server/biz/model/dto"
|
||||
"observer-server/common"
|
||||
)
|
||||
|
||||
// uniqueVersion 每次运行生成唯一版本号,避免测试库残留数据冲突
|
||||
func uniqueVersion() string {
|
||||
return "9." + strconv.FormatInt(time.Now().UnixNano()%100000000, 10) + ".0"
|
||||
}
|
||||
|
||||
// newApkUpload 构造真实 multipart 上传文件(Save 需要可读的文件内容)
|
||||
func newApkUpload(t *testing.T, filename, content string) *ghttp.UploadFile {
|
||||
t.Helper()
|
||||
var buf bytes.Buffer
|
||||
w := multipart.NewWriter(&buf)
|
||||
fw, err := w.CreateFormFile("file", filename)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := fw.Write([]byte(content)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
w.Close()
|
||||
form, err := multipart.NewReader(&buf, w.Boundary()).ReadForm(1 << 20)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return &ghttp.UploadFile{FileHeader: form.File["file"][0]}
|
||||
}
|
||||
|
||||
// countApk 目录下 apk 文件数(不含临时文件)
|
||||
func countApk(t *testing.T) int {
|
||||
t.Helper()
|
||||
entries, err := os.ReadDir(common.ApkDir(ctx()))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
n := 0
|
||||
for _, e := range entries {
|
||||
if filepath.Ext(e.Name()) == ".apk" {
|
||||
n++
|
||||
}
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
// TestAppVersionUpdate 下发:落库 + APK 保存固定文件,更新检查返回最新记录
|
||||
func TestAppVersionUpdate(t *testing.T) {
|
||||
ver := uniqueVersion()
|
||||
if _, err := AppVersion.AdminAddVersion(ctx(), &dto.AdminAppVersionAddReq{
|
||||
Notes: "修复识别准确率", File: newApkUpload(t, "observer-"+ver+".apk", "apk-v1"),
|
||||
}); err != nil {
|
||||
t.Fatalf("add version: %v", err)
|
||||
}
|
||||
res, err := AppVersion.GetUpdate(ctx(), &dto.AppUpdateReq{})
|
||||
if err != nil {
|
||||
t.Fatalf("get update: %v", err)
|
||||
}
|
||||
if res.Version != ver || res.Notes != "修复识别准确率" {
|
||||
t.Fatalf("update = %+v", res)
|
||||
}
|
||||
// APK 已保存为固定文件名
|
||||
content, err := os.ReadFile(common.ApkFilePath(ctx()))
|
||||
if err != nil || string(content) != "apk-v1" {
|
||||
t.Fatalf("apk file = %q, %v", content, err)
|
||||
}
|
||||
if n := countApk(t); n != 1 {
|
||||
t.Fatalf("apk count = %d, want 1", n)
|
||||
}
|
||||
}
|
||||
|
||||
// TestAppVersionOverwrite 再次下发:固定文件被新内容覆盖,目录仍只有一个文件
|
||||
func TestAppVersionOverwrite(t *testing.T) {
|
||||
ver1 := uniqueVersion()
|
||||
ver2 := uniqueVersion() + "1"
|
||||
req1 := &dto.AdminAppVersionAddReq{File: newApkUpload(t, "observer-"+ver1+".apk", "apk-old")}
|
||||
req2 := &dto.AdminAppVersionAddReq{File: newApkUpload(t, "observer-"+ver2+".apk", "apk-new")}
|
||||
if _, err := AppVersion.AdminAddVersion(ctx(), req1); err != nil {
|
||||
t.Fatalf("add v1: %v", err)
|
||||
}
|
||||
if _, err := AppVersion.AdminAddVersion(ctx(), req2); err != nil {
|
||||
t.Fatalf("add v2: %v", err)
|
||||
}
|
||||
content, err := os.ReadFile(common.ApkFilePath(ctx()))
|
||||
if err != nil || string(content) != "apk-new" {
|
||||
t.Fatalf("apk file = %q, %v; want new content", content, err)
|
||||
}
|
||||
if n := countApk(t); n != 1 {
|
||||
t.Fatalf("apk count = %d, want 1", n)
|
||||
}
|
||||
}
|
||||
|
||||
// TestAppVersionInvalidFile 无文件 / 文件名不含版本号(格式不符)拒绝
|
||||
func TestAppVersionInvalidFile(t *testing.T) {
|
||||
if _, err := AppVersion.AdminAddVersion(ctx(), &dto.AdminAppVersionAddReq{}); err == nil {
|
||||
t.Fatal("missing file should fail")
|
||||
}
|
||||
for _, name := range []string{"app.txt", "myapp-1.0.apk", "observer-1.0.apk", "observer-abc.apk"} {
|
||||
if _, err := AppVersion.AdminAddVersion(ctx(), &dto.AdminAppVersionAddReq{
|
||||
File: newApkUpload(t, name, "not an apk"),
|
||||
}); err == nil {
|
||||
t.Fatalf("file %s should fail (version not parseable)", name)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestAppVersionDuplicate 同版本号重复下发报错
|
||||
func TestAppVersionDuplicate(t *testing.T) {
|
||||
ver := uniqueVersion()
|
||||
req := &dto.AdminAppVersionAddReq{File: newApkUpload(t, "observer-"+ver+".apk", "apk")}
|
||||
if _, err := AppVersion.AdminAddVersion(ctx(), req); err != nil {
|
||||
t.Fatalf("add version: %v", err)
|
||||
}
|
||||
if _, err := AppVersion.AdminAddVersion(ctx(), req); err == nil {
|
||||
t.Fatal("duplicate version should fail")
|
||||
}
|
||||
}
|
||||
|
||||
// TestAppVersionList 分页列表:新增多条后按下发时间倒序
|
||||
func TestAppVersionList(t *testing.T) {
|
||||
base := uniqueVersion()
|
||||
for _, v := range []string{base, base + "1"} {
|
||||
if _, err := AppVersion.AdminAddVersion(ctx(), &dto.AdminAppVersionAddReq{
|
||||
File: newApkUpload(t, "observer-"+v+".apk", "apk"),
|
||||
}); err != nil {
|
||||
t.Fatalf("add version %s: %v", v, err)
|
||||
}
|
||||
}
|
||||
res, err := AppVersion.AdminListVersions(ctx(), &dto.AdminAppVersionListReq{Size: 10})
|
||||
if err != nil {
|
||||
t.Fatalf("list versions: %v", err)
|
||||
}
|
||||
if res.Total < 2 {
|
||||
t.Fatalf("total = %d, want >= 2", res.Total)
|
||||
}
|
||||
// 倒序:最新一条为 base+1
|
||||
if res.List[0].Version != base+"1" {
|
||||
t.Fatalf("latest = %s, want %s", res.List[0].Version, base+"1")
|
||||
}
|
||||
}
|
||||
|
||||
// TestAppVersionDeleteLatest 删除最新版本:记录删除 + APK 文件联动删除
|
||||
func TestAppVersionDeleteLatest(t *testing.T) {
|
||||
ver := uniqueVersion()
|
||||
if _, err := AppVersion.AdminAddVersion(ctx(), &dto.AdminAppVersionAddReq{
|
||||
File: newApkUpload(t, "observer-"+ver+".apk", "apk-v1"),
|
||||
}); err != nil {
|
||||
t.Fatalf("add version: %v", err)
|
||||
}
|
||||
latest, err := dao.AppVersion.Latest(ctx())
|
||||
if err != nil || latest == nil {
|
||||
t.Fatalf("latest = %+v, %v", latest, err)
|
||||
}
|
||||
if _, err := AppVersion.AdminDeleteVersion(ctx(), &dto.AdminAppVersionDeleteReq{Id: latest.Id}); err != nil {
|
||||
t.Fatalf("delete version: %v", err)
|
||||
}
|
||||
// 记录已删(测试库有历史残留记录,update 不再返回被删版本即可)
|
||||
gone, err := dao.AppVersion.GetById(ctx(), latest.Id)
|
||||
if err != nil || gone != nil {
|
||||
t.Fatalf("record should be deleted: %+v, %v", gone, err)
|
||||
}
|
||||
res, err := AppVersion.GetUpdate(ctx(), &dto.AppUpdateReq{})
|
||||
if err != nil {
|
||||
t.Fatalf("get update: %v", err)
|
||||
}
|
||||
if res.Version == ver {
|
||||
t.Fatalf("update still returns deleted version %s", ver)
|
||||
}
|
||||
// APK 文件联动删除(下载 404)
|
||||
if _, err := os.Stat(common.ApkFilePath(ctx())); !os.IsNotExist(err) {
|
||||
t.Fatalf("apk file should be removed, stat err = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestAppVersionDeleteHistoric 删除历史版本:仅删记录,APK 文件保留(仍对应最新版本)
|
||||
func TestAppVersionDeleteHistoric(t *testing.T) {
|
||||
ver1 := uniqueVersion()
|
||||
ver2 := uniqueVersion() + "1"
|
||||
for _, v := range []string{ver1, ver2} {
|
||||
if _, err := AppVersion.AdminAddVersion(ctx(), &dto.AdminAppVersionAddReq{
|
||||
File: newApkUpload(t, "observer-"+v+".apk", "apk"),
|
||||
}); err != nil {
|
||||
t.Fatalf("add version %s: %v", v, err)
|
||||
}
|
||||
}
|
||||
old, err := dao.AppVersion.GetByVersion(ctx(), ver1)
|
||||
if err != nil || old == nil {
|
||||
t.Fatalf("get by version: %+v, %v", old, err)
|
||||
}
|
||||
if _, err := AppVersion.AdminDeleteVersion(ctx(), &dto.AdminAppVersionDeleteReq{Id: old.Id}); err != nil {
|
||||
t.Fatalf("delete version: %v", err)
|
||||
}
|
||||
// 最新记录仍是 ver2
|
||||
res, err := AppVersion.GetUpdate(ctx(), &dto.AppUpdateReq{})
|
||||
if err != nil {
|
||||
t.Fatalf("get update: %v", err)
|
||||
}
|
||||
if res.Version != ver2 {
|
||||
t.Fatalf("update = %+v, want %s", res, ver2)
|
||||
}
|
||||
// APK 文件保留
|
||||
if _, err := os.Stat(common.ApkFilePath(ctx())); err != nil {
|
||||
t.Fatalf("apk file should remain: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestAppVersionDeleteNotFound 删除不存在的版本报错
|
||||
func TestAppVersionDeleteNotFound(t *testing.T) {
|
||||
if _, err := AppVersion.AdminDeleteVersion(ctx(), &dto.AdminAppVersionDeleteReq{Id: 999999999}); err == nil {
|
||||
t.Fatal("delete missing version should fail")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,678 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"archive/zip"
|
||||
"bytes"
|
||||
"context"
|
||||
"fmt"
|
||||
"image"
|
||||
"image/jpeg"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gogf/gf/v2/errors/gerror"
|
||||
"github.com/gogf/gf/v2/frame/g"
|
||||
"github.com/gogf/gf/v2/net/ghttp"
|
||||
"github.com/gogf/gf/v2/os/gtime"
|
||||
|
||||
"observer-server/biz/consts"
|
||||
"observer-server/biz/dao"
|
||||
"observer-server/biz/model/dto"
|
||||
"observer-server/biz/model/entity"
|
||||
"observer-server/common"
|
||||
)
|
||||
|
||||
// datasetService 数据集业务:管理端「数据集管理」模块。
|
||||
// 图片文件在 app.datasetDir/datasets/<name>/(平铺、文件名唯一),DB 只存元数据;
|
||||
// AI 生成为付费资产,删除类接口由前端带确认文案(后端不额外拦截)。
|
||||
type datasetService struct{}
|
||||
|
||||
var Dataset = &datasetService{}
|
||||
|
||||
// imgExts 允许上传/保存的图片扩展名
|
||||
var imgExts = map[string]bool{".jpg": true, ".jpeg": true, ".png": true}
|
||||
|
||||
// AdminListDatasets 数据集分页列表(卡片展示:封面/描述/训练配置标记 + 最新训练状态聚合)
|
||||
func (s *datasetService) AdminListDatasets(ctx context.Context, req *dto.AdminDatasetListReq) (*dto.AdminDatasetListRes, error) {
|
||||
page, size := common.NormalizePage(req.Page, req.Size)
|
||||
var list []*entity.Dataset
|
||||
var total int64
|
||||
var err error
|
||||
if req.Keyword != "" {
|
||||
list, total, err = dao.Dataset.PageByKeyword(ctx, req.Keyword, page, size)
|
||||
} else {
|
||||
list, total, err = dao.Dataset.Page(ctx, page, size)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
ids := make([]int64, 0, len(list))
|
||||
for _, v := range list {
|
||||
ids = append(ids, v.Id)
|
||||
}
|
||||
latest, err := dao.Training.LatestByDatasets(ctx, ids)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
items := make([]*dto.AdminDatasetItem, 0, len(list))
|
||||
for _, v := range list {
|
||||
item := &dto.AdminDatasetItem{
|
||||
Id: v.Id,
|
||||
Name: v.Name,
|
||||
Source: v.Source,
|
||||
ImageCount: v.ImageCount,
|
||||
LabeledCount: v.LabeledCount,
|
||||
Status: v.Status,
|
||||
Cover: v.Cover,
|
||||
Description: v.Description,
|
||||
CreatedAt: v.CreatedAt,
|
||||
UpdatedAt: v.UpdatedAt,
|
||||
}
|
||||
if t, ok := latest[v.Id]; ok {
|
||||
item.TrainingId = t.Id
|
||||
item.TrainingStatus = t.Status
|
||||
item.TrainingCurrentEpoch = t.CurrentEpoch
|
||||
item.TrainingTotalEpochs = t.TotalEpochs
|
||||
}
|
||||
items = append(items, item)
|
||||
}
|
||||
return &dto.AdminDatasetListRes{Total: total, List: items}, nil
|
||||
}
|
||||
|
||||
// AdminUpdateDataset 更新数据集展示配置(空值字段不覆盖;封面/描述。
|
||||
// AI 端点/训练机 SSH 走 config.yml 全局配置(localAi / training.ssh))
|
||||
func (s *datasetService) AdminUpdateDataset(ctx context.Context, req *dto.AdminDatasetUpdateReq) (*dto.AdminDatasetUpdateRes, error) {
|
||||
existing, err := dao.Dataset.GetById(ctx, req.Id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if existing == nil {
|
||||
return nil, gerror.NewCode(common.CodeDatasetNotFound)
|
||||
}
|
||||
// 封面须属于该数据集图片(防伪造文件名指向任意文件)
|
||||
if req.Cover != "" {
|
||||
img, err := dao.DatasetImage.GetByFilename(ctx, req.Id, req.Cover)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if img == nil {
|
||||
return nil, gerror.NewCode(common.CodeImageNotFound)
|
||||
}
|
||||
}
|
||||
if err := dao.Dataset.UpdateConfigs(ctx, req.Id, &entity.Dataset{
|
||||
Cover: req.Cover,
|
||||
Description: req.Description,
|
||||
}); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &dto.AdminDatasetUpdateRes{Id: req.Id}, nil
|
||||
}
|
||||
|
||||
// AdminUploadCover 上传数据集封面:解码校验图片 → 转 jpg(Quality 92)→ UUID 命名落盘,
|
||||
// 删除旧封面文件(UUID 互不覆盖,但 DB 指向切换)→ 更新 cover 字段。
|
||||
func (s *datasetService) AdminUploadCover(ctx context.Context, req *dto.AdminDatasetCoverUploadReq) (*dto.AdminDatasetCoverUploadRes, error) {
|
||||
dataset, err := dao.Dataset.GetById(ctx, req.DatasetId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if dataset == nil {
|
||||
return nil, gerror.NewCode(common.CodeDatasetNotFound)
|
||||
}
|
||||
if req.File == nil {
|
||||
return nil, gerror.New("缺少封面文件")
|
||||
}
|
||||
ext := strings.ToLower(filepath.Ext(req.File.Filename))
|
||||
if !imgExts[ext] {
|
||||
return nil, gerror.New("封面仅支持 jpg/jpeg/png")
|
||||
}
|
||||
if req.File.Size > 2*1024*1024 {
|
||||
return nil, gerror.New("封面不能超过 2MB")
|
||||
}
|
||||
raw, err := req.File.Open()
|
||||
if err != nil {
|
||||
return nil, gerror.Wrap(err, "读取封面文件失败")
|
||||
}
|
||||
defer func() { _ = raw.Close() }()
|
||||
img, _, err := image.Decode(raw)
|
||||
if err != nil {
|
||||
return nil, gerror.New("封面文件不是有效图片")
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
if err := jpeg.Encode(&buf, img, &jpeg.Options{Quality: 92}); err != nil {
|
||||
return nil, gerror.Wrap(err, "封面转 jpg 失败")
|
||||
}
|
||||
dir := common.DatasetImagesDir(ctx, dataset.Name)
|
||||
if err := os.MkdirAll(dir, 0o755); err != nil {
|
||||
return nil, gerror.Wrap(err, "创建图片目录失败")
|
||||
}
|
||||
if dataset.Cover != "" {
|
||||
_ = os.Remove(filepath.Join(dir, filepath.Base(dataset.Cover)))
|
||||
}
|
||||
name := common.UuidV4() + ".jpg"
|
||||
if err := common.WriteFileAtomic(filepath.Join(dir, name), buf.Bytes()); err != nil {
|
||||
return nil, gerror.Wrap(err, "保存封面失败")
|
||||
}
|
||||
if err := dao.Dataset.UpdateConfigs(ctx, dataset.Id, &entity.Dataset{Cover: name}); err != nil {
|
||||
_ = os.Remove(filepath.Join(dir, name))
|
||||
return nil, err
|
||||
}
|
||||
return &dto.AdminDatasetCoverUploadRes{}, nil
|
||||
}
|
||||
|
||||
// coverNameRe 封面命名规范:UUIDv4 + .jpg(固定命名 cover* 为历史遗留,迁移见 MigrateLegacyCovers)
|
||||
var coverNameRe = regexp.MustCompile(`^[0-9a-f]{8}-[0-9a-f]{4}-4[0-9a-f]{3}-[89ab][0-9a-f]{3}-[0-9a-f]{12}\.jpg$`)
|
||||
|
||||
// isCoverName 校验 DB cover 值是否符合 UUID jpg 规范(防御库内被写入任意文件名)
|
||||
func isCoverName(name string) bool {
|
||||
return coverNameRe.MatchString(filepath.Base(name))
|
||||
}
|
||||
|
||||
// CoverFile 封面文件定位(校验归属;controller 直写响应体输出)
|
||||
func (s *datasetService) CoverFile(ctx context.Context, datasetId int64) (string, error) {
|
||||
dataset, err := dao.Dataset.GetById(ctx, datasetId)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if dataset == nil {
|
||||
return "", gerror.NewCode(common.CodeDatasetNotFound)
|
||||
}
|
||||
if dataset.Cover == "" {
|
||||
return "", gerror.NewCode(common.CodeImageNotFound)
|
||||
}
|
||||
name := filepath.Base(dataset.Cover)
|
||||
if !isCoverName(name) {
|
||||
return "", gerror.NewCode(common.CodeImageNotFound)
|
||||
}
|
||||
path := filepath.Join(common.DatasetImagesDir(ctx, dataset.Name), name)
|
||||
if _, err := os.Stat(path); err != nil {
|
||||
return "", gerror.NewCode(common.CodeImageNotFound)
|
||||
}
|
||||
return path, nil
|
||||
}
|
||||
|
||||
// AdminDeleteCover 删除数据集封面:删文件 + 清 cover 字段(列表卡片恢复占位图)
|
||||
func (s *datasetService) AdminDeleteCover(ctx context.Context, req *dto.AdminDatasetCoverDeleteReq) (*dto.AdminDatasetCoverDeleteRes, error) {
|
||||
dataset, err := dao.Dataset.GetById(ctx, req.DatasetId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if dataset == nil {
|
||||
return nil, gerror.NewCode(common.CodeDatasetNotFound)
|
||||
}
|
||||
if dataset.Cover != "" && isCoverName(dataset.Cover) {
|
||||
_ = os.Remove(filepath.Join(common.DatasetImagesDir(ctx, dataset.Name), filepath.Base(dataset.Cover)))
|
||||
}
|
||||
if err := dao.Dataset.ClearCover(ctx, dataset.Id); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &dto.AdminDatasetCoverDeleteRes{}, nil
|
||||
}
|
||||
|
||||
// MigrateLegacyCovers 存量封面迁移:cover 非空且不符合 UUID jpg 规范(历史 cover.jpg/cover.png 等)
|
||||
// → 重命名为 <uuid>.jpg(保留原图字节)并更新 cover 字段。幂等:已符合规范的行跳过。
|
||||
func (s *datasetService) MigrateLegacyCovers(ctx context.Context) error {
|
||||
list, err := dao.Dataset.ListAll(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
migrated := 0
|
||||
for _, d := range list {
|
||||
if d.Cover == "" || isCoverName(d.Cover) {
|
||||
continue
|
||||
}
|
||||
oldName := filepath.Base(d.Cover)
|
||||
src := filepath.Join(common.DatasetImagesDir(ctx, d.Name), oldName)
|
||||
data, err := os.ReadFile(src)
|
||||
if err != nil {
|
||||
g.Log().Warningf(ctx, "数据集 %s 封面迁移跳过(文件不存在: %s): %v", d.Name, src, err)
|
||||
continue
|
||||
}
|
||||
newName := common.UuidV4() + ".jpg"
|
||||
if err := common.WriteFileAtomic(filepath.Join(common.DatasetImagesDir(ctx, d.Name), newName), data); err != nil {
|
||||
g.Log().Errorf(ctx, "数据集 %s 封面迁移写新文件失败: %+v", d.Name, err)
|
||||
continue
|
||||
}
|
||||
if err := dao.Dataset.UpdateConfigs(ctx, d.Id, &entity.Dataset{Cover: newName}); err != nil {
|
||||
g.Log().Errorf(ctx, "数据集 %s 封面迁移更新字段失败: %+v", d.Name, err)
|
||||
continue
|
||||
}
|
||||
_ = os.Remove(src)
|
||||
migrated++
|
||||
g.Log().Infof(ctx, "数据集 %s 封面迁移: %s → %s", d.Name, oldName, newName)
|
||||
}
|
||||
if migrated > 0 {
|
||||
g.Log().Infof(ctx, "封面存量迁移完成: 共 %d 个数据集", migrated)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// AdminCreateDataset 新建数据集:名称唯一(UNIQUE 兜底)+ 创建图片目录
|
||||
func (s *datasetService) AdminCreateDataset(ctx context.Context, req *dto.AdminDatasetCreateReq) (*dto.AdminDatasetCreateRes, error) {
|
||||
now := gtime.Now()
|
||||
var id int64
|
||||
err := common.Serial().Submit(ctx, func() error {
|
||||
exists, err := dao.Dataset.GetByName(ctx, req.Name)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if exists != nil {
|
||||
return gerror.NewCode(common.CodeDatasetNameDuplicate)
|
||||
}
|
||||
id, err = dao.Dataset.Insert(ctx, &entity.Dataset{
|
||||
Name: req.Name,
|
||||
Source: req.Source,
|
||||
Status: consts.DatasetStatusBuilding,
|
||||
CreatedAt: now,
|
||||
UpdatedAt: now,
|
||||
})
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := os.MkdirAll(common.DatasetImagesDir(ctx, req.Name), 0o755); err != nil {
|
||||
return nil, gerror.Wrap(err, "创建数据集目录失败")
|
||||
}
|
||||
return &dto.AdminDatasetCreateRes{Id: id}, nil
|
||||
}
|
||||
|
||||
// AdminDeleteDataset 删除数据集:有 running 标注任务 / 该数据集训练进行中 / 已发布模型版本时拒绝
|
||||
// (训练产物与模型为付费资产,需先删除模型版本再删数据集)。
|
||||
// 删除 = 删图片/模型目录 + 删记录(标注随图片行删除,Serial 单写者串行)。
|
||||
func (s *datasetService) AdminDeleteDataset(ctx context.Context, req *dto.AdminDatasetDeleteReq) (*dto.AdminDatasetDeleteRes, error) {
|
||||
var name string
|
||||
err := common.Serial().Submit(ctx, func() error {
|
||||
d, err := dao.Dataset.GetById(ctx, req.Id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if d == nil {
|
||||
return gerror.NewCode(common.CodeDatasetNotFound)
|
||||
}
|
||||
name = d.Name
|
||||
// 预标注任务进行中(RF-DETR 正在扫该数据集图片)
|
||||
if t, err := dao.LabelTask.GetRunningByDataset(ctx, d.Id); err != nil {
|
||||
return err
|
||||
} else if t != nil {
|
||||
return gerror.NewCode(common.CodeLabelTaskRunning)
|
||||
}
|
||||
// 该数据集训练进行中(并发度 1,防文件被删训练中断)
|
||||
if t, err := dao.Training.RunningByDataset(ctx, d.Id); err != nil {
|
||||
return err
|
||||
} else if t != nil {
|
||||
return gerror.New("该数据集有训练任务进行中,无法删除")
|
||||
}
|
||||
// 模型版本记录随数据集级联删除(管理端无模型管理界面,2026-08-26 决策;
|
||||
// 若需保留已下发模型,删除数据集前先确认客户端不再需要)
|
||||
if err := dao.ModelVersion.DeleteByDataset(ctx, d.Id); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := dao.DatasetImage.DeleteByDataset(ctx, d.Id); err != nil {
|
||||
return err
|
||||
}
|
||||
return dao.Dataset.DeleteById(ctx, d.Id)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// 文件清理(图片 + 模型文件;删除失败仅记日志,记录已删)
|
||||
dirs := []string{common.DatasetImagesDir(ctx, name), common.DatasetModelsDir(ctx, name)}
|
||||
for _, dir := range dirs {
|
||||
if err := os.RemoveAll(dir); err != nil {
|
||||
g.Log().Errorf(ctx, "删除数据集 %s 目录失败: %+v", dir, err)
|
||||
}
|
||||
}
|
||||
return &dto.AdminDatasetDeleteRes{}, nil
|
||||
}
|
||||
|
||||
// AdminUploadImages 上传图片(多文件):重名跳过(文件名唯一约束),逐张落盘 + 批量入库。
|
||||
func (s *datasetService) AdminUploadImages(ctx context.Context, req *dto.AdminDatasetUploadReq) (*dto.AdminDatasetUploadRes, error) {
|
||||
if len(req.Files) == 0 {
|
||||
return nil, gerror.New("请选择图片文件")
|
||||
}
|
||||
dataset, err := dao.Dataset.GetById(ctx, req.DatasetId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if dataset == nil {
|
||||
return nil, gerror.NewCode(common.CodeDatasetNotFound)
|
||||
}
|
||||
dir := common.DatasetImagesDir(ctx, dataset.Name)
|
||||
if err := os.MkdirAll(dir, 0o755); err != nil {
|
||||
return nil, gerror.Wrap(err, "创建图片目录失败")
|
||||
}
|
||||
res := &dto.AdminDatasetUploadRes{}
|
||||
now := gtime.Now()
|
||||
saved := make([]string, 0, len(req.Files)) // 本次已落盘文件名(入库失败时清理)
|
||||
for _, f := range req.Files {
|
||||
name := filepath.Base(f.Filename)
|
||||
if !imgExts[strings.ToLower(filepath.Ext(name))] {
|
||||
continue
|
||||
}
|
||||
exists, err := dao.DatasetImage.GetByFilename(ctx, dataset.Id, name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if exists != nil {
|
||||
res.Skipped++
|
||||
continue
|
||||
}
|
||||
if err := saveImageFile(f, dir); err != nil {
|
||||
return nil, gerror.Wrapf(err, "图片 %s 保存失败", name)
|
||||
}
|
||||
saved = append(saved, name)
|
||||
}
|
||||
if len(saved) == 0 {
|
||||
return res, nil
|
||||
}
|
||||
// 标注强语义:未配置标注服务时不允许产生无标注图(文件已落盘,失败则清掉)
|
||||
if common.LocalAiClient(ctx) == nil {
|
||||
for _, name := range saved {
|
||||
_ = os.Remove(filepath.Join(dir, name))
|
||||
}
|
||||
return nil, gerror.NewCode(common.CodeLocalAiNotConfigured)
|
||||
}
|
||||
addedIds := make([]int64, 0, len(saved))
|
||||
err = common.Serial().Submit(ctx, func() error {
|
||||
for _, name := range saved {
|
||||
// 重查重(并发上传兜底)+ 入库
|
||||
exists, err := dao.DatasetImage.GetByFilename(ctx, dataset.Id, name)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if exists != nil {
|
||||
continue
|
||||
}
|
||||
id, err := dao.DatasetImage.Insert(ctx, &entity.DatasetImage{
|
||||
DatasetId: dataset.Id,
|
||||
Filename: name,
|
||||
Source: "manual",
|
||||
CreatedAt: now,
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
addedIds = append(addedIds, id)
|
||||
res.Added++
|
||||
}
|
||||
return dao.Dataset.UpdateCounters(ctx, dataset.Id, int64(res.Added), -1, "")
|
||||
})
|
||||
if err != nil {
|
||||
// 入库失败:清掉已落盘文件,保证「记录存在 ⟺ 文件存在」
|
||||
for _, name := range saved {
|
||||
_ = os.Remove(filepath.Join(dir, name))
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
// 自动触发标注:忙(已有 running 任务)不报错,由任务完成后的自动补标轮兜底;其他失败回滚本次入库
|
||||
newImages, err := dao.DatasetImage.GetByIds(ctx, addedIds)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := LabelTask.AutoLabel(ctx, dataset, newImages); err != nil {
|
||||
// 触发失败回滚本次入库(manual 上传文件非付费资产,可删)
|
||||
rollbackIds := make([]int64, 0, len(newImages))
|
||||
for _, img := range newImages {
|
||||
rollbackIds = append(rollbackIds, img.Id)
|
||||
if rErr := os.Remove(filepath.Join(dir, img.Filename)); rErr != nil {
|
||||
g.Log().Warningf(ctx, "回滚删除图片文件失败: %s: %v", img.Filename, rErr)
|
||||
}
|
||||
}
|
||||
if rErr := common.Serial().Submit(ctx, func() error {
|
||||
if dErr := dao.DatasetImage.DeleteByIds(ctx, rollbackIds); dErr != nil {
|
||||
return dErr
|
||||
}
|
||||
return dao.Dataset.UpdateCounters(ctx, dataset.Id, -int64(len(rollbackIds)), 0, "")
|
||||
}); rErr != nil {
|
||||
g.Log().Errorf(ctx, "自动标注触发失败后的入库回滚失败: %v", rErr)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return res, nil
|
||||
}
|
||||
|
||||
// saveImageFile 上传文件保存(ghttp.UploadFile.Save 已按原始文件名落盘)
|
||||
func saveImageFile(f *ghttp.UploadFile, dir string) error {
|
||||
saved, err := f.Save(dir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if saved != filepath.Base(f.Filename) {
|
||||
_ = os.Remove(filepath.Join(dir, saved))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// AdminGenerateImages AI 生成图片:同步逐张生成(每张 2min 超时),逐张落盘 + 入库;
|
||||
// 中途失败返回错误并附已成功张数(已生成的图为付费资产,保留不删除)。
|
||||
func (s *datasetService) AdminGenerateImages(ctx context.Context, req *dto.AdminDatasetGenerateReq) (*dto.AdminDatasetGenerateRes, error) {
|
||||
provider := common.ImageGen(ctx)
|
||||
if provider == nil {
|
||||
return nil, gerror.NewCode(common.CodeImageGenNotConfigured)
|
||||
}
|
||||
dataset, err := dao.Dataset.GetById(ctx, req.DatasetId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if dataset == nil {
|
||||
return nil, gerror.NewCode(common.CodeDatasetNotFound)
|
||||
}
|
||||
// 标注强语义:未配置标注服务时不允许开始生成(避免付费资产生成后无法标注)
|
||||
if common.LocalAiClient(ctx) == nil {
|
||||
return nil, gerror.NewCode(common.CodeLocalAiNotConfigured)
|
||||
}
|
||||
dir := common.DatasetImagesDir(ctx, dataset.Name)
|
||||
if err := os.MkdirAll(dir, 0o755); err != nil {
|
||||
return nil, gerror.Wrap(err, "创建图片目录失败")
|
||||
}
|
||||
now := gtime.Now()
|
||||
generated := 0
|
||||
addedIds := make([]int64, 0, req.Count)
|
||||
for i := 0; i < req.Count; i++ {
|
||||
genCtx, cancel := context.WithTimeout(ctx, 120*time.Second)
|
||||
data, genErr := provider.Generate(genCtx, req.Prompt, req.Size)
|
||||
cancel()
|
||||
if genErr != nil {
|
||||
// 付费资产保留原则:已生成的不删除
|
||||
return &dto.AdminDatasetGenerateRes{Generated: generated},
|
||||
gerror.NewCode(common.CodeImageGenFailed, fmt.Sprintf("第 %d 张生成失败(已生成 %d 张): %v", i+1, generated, genErr))
|
||||
}
|
||||
filename := fmt.Sprintf("gen_%s_%d.jpg", now.Format("20060102150405"), i)
|
||||
if err := common.WriteFileAtomic(filepath.Join(dir, filename), data); err != nil {
|
||||
return &dto.AdminDatasetGenerateRes{Generated: generated},
|
||||
gerror.Wrapf(err, "第 %d 张保存失败(已生成 %d 张)", i+1, generated)
|
||||
}
|
||||
var imgId int64
|
||||
insErr := common.Serial().Submit(ctx, func() error {
|
||||
id, err := dao.DatasetImage.Insert(ctx, &entity.DatasetImage{
|
||||
DatasetId: dataset.Id,
|
||||
Filename: filename,
|
||||
Source: "ai",
|
||||
Prompt: req.Prompt,
|
||||
CreatedAt: now,
|
||||
})
|
||||
imgId = id
|
||||
return err
|
||||
})
|
||||
if insErr != nil {
|
||||
return &dto.AdminDatasetGenerateRes{Generated: generated},
|
||||
gerror.Wrap(insErr, "生成图片入库失败")
|
||||
}
|
||||
addedIds = append(addedIds, imgId)
|
||||
generated++
|
||||
}
|
||||
if err := common.Serial().Submit(ctx, func() error {
|
||||
return dao.Dataset.UpdateCounters(ctx, dataset.Id, int64(generated), -1, "")
|
||||
}); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// 自动触发标注:忙(已有 running 任务)不报错,由任务完成后的自动补标轮兜底;
|
||||
// 其他失败报错但保留已生成图(付费资产,不可删)
|
||||
newImages, err := dao.DatasetImage.GetByIds(ctx, addedIds)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := LabelTask.AutoLabel(ctx, dataset, newImages); err != nil {
|
||||
return &dto.AdminDatasetGenerateRes{Generated: generated},
|
||||
gerror.Wrap(err, "图片已生成入库,但自动标注触发失败")
|
||||
}
|
||||
return &dto.AdminDatasetGenerateRes{Generated: generated}, nil
|
||||
}
|
||||
|
||||
// AdminListImages 数据集图片列表(创建时间正序,标注工作台/网格预览)
|
||||
func (s *datasetService) AdminListImages(ctx context.Context, req *dto.AdminDatasetImagesReq) (*dto.AdminDatasetImagesRes, error) {
|
||||
dataset, err := dao.Dataset.GetById(ctx, req.DatasetId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if dataset == nil {
|
||||
return nil, gerror.NewCode(common.CodeDatasetNotFound)
|
||||
}
|
||||
list, err := dao.DatasetImage.ListByDataset(ctx, dataset.Id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
items := make([]*dto.AdminImageItem, 0, len(list))
|
||||
for _, v := range list {
|
||||
items = append(items, &dto.AdminImageItem{
|
||||
Id: v.Id,
|
||||
Filename: v.Filename,
|
||||
Source: v.Source,
|
||||
Prompt: v.Prompt,
|
||||
Url: datasetImageUrl(ctx, dataset.Id, v.Filename),
|
||||
CreatedAt: v.CreatedAt,
|
||||
})
|
||||
}
|
||||
return &dto.AdminDatasetImagesRes{List: items}, nil
|
||||
}
|
||||
|
||||
// AdminDeleteImages 删除图片:删图片文件 + 记录 + 更新计数
|
||||
// (付费资产,前端带确认文案;标注/候选框 JSON 随图片行删除)
|
||||
func (s *datasetService) AdminDeleteImages(ctx context.Context, req *dto.AdminDatasetImagesDeleteReq) (*dto.AdminDatasetImagesDeleteRes, error) {
|
||||
dataset, err := dao.Dataset.GetById(ctx, req.DatasetId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if dataset == nil {
|
||||
return nil, gerror.NewCode(common.CodeDatasetNotFound)
|
||||
}
|
||||
images, err := dao.DatasetImage.GetByIds(ctx, req.Ids)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var del []*entity.DatasetImage
|
||||
for _, img := range images {
|
||||
// 只删本数据集内的记录(防跨数据集误删)
|
||||
if img.DatasetId == dataset.Id {
|
||||
del = append(del, img)
|
||||
}
|
||||
}
|
||||
if len(del) == 0 {
|
||||
return &dto.AdminDatasetImagesDeleteRes{}, nil
|
||||
}
|
||||
// 文件清理(删除失败仅记日志,记录照删):图片文件(标注/候选随行删除)
|
||||
imgDir := common.DatasetImagesDir(ctx, dataset.Name)
|
||||
for _, img := range del {
|
||||
_ = os.Remove(filepath.Join(imgDir, img.Filename))
|
||||
}
|
||||
// Serial 内:删行 → 重算已标注数 → 同步计数与状态
|
||||
err = common.Serial().Submit(ctx, func() error {
|
||||
ids := make([]int64, 0, len(del))
|
||||
for _, img := range del {
|
||||
ids = append(ids, img.Id)
|
||||
}
|
||||
if err := dao.DatasetImage.DeleteByIds(ctx, ids); err != nil {
|
||||
return err
|
||||
}
|
||||
labeled, err := dao.DatasetImage.CountLabeledByDataset(ctx, dataset.Id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
status := ""
|
||||
if labeled > 0 {
|
||||
status = consts.DatasetStatusLabeled
|
||||
}
|
||||
return dao.Dataset.UpdateCounters(ctx, dataset.Id, -int64(len(del)), labeled, status)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &dto.AdminDatasetImagesDeleteRes{}, nil
|
||||
}
|
||||
|
||||
// datasetImageUrl 管理端图片访问地址(controller 直写响应体,见 admin.go Image 方法)
|
||||
func datasetImageUrl(ctx context.Context, datasetId int64, filename string) string {
|
||||
return fmt.Sprintf("/api/v1/admin/datasets/image?%s",
|
||||
url.Values{"datasetId": {fmt.Sprintf("%d", datasetId)}, "filename": {filename}}.Encode())
|
||||
}
|
||||
|
||||
// ImageFile 图片文件定位(校验归属后返回绝对路径;controller 直写响应体)
|
||||
func (s *datasetService) ImageFile(ctx context.Context, datasetId int64, filename string) (string, error) {
|
||||
if filename != filepath.Base(filename) || !imgExts[strings.ToLower(filepath.Ext(filename))] {
|
||||
return "", gerror.NewCode(common.CodeImageNotFound)
|
||||
}
|
||||
dataset, err := dao.Dataset.GetById(ctx, datasetId)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if dataset == nil {
|
||||
return "", gerror.NewCode(common.CodeDatasetNotFound)
|
||||
}
|
||||
img, err := dao.DatasetImage.GetByFilename(ctx, dataset.Id, filename)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if img == nil {
|
||||
return "", gerror.NewCode(common.CodeImageNotFound)
|
||||
}
|
||||
path := filepath.Join(common.DatasetImagesDir(ctx, dataset.Name), filename)
|
||||
if _, err := os.Stat(path); err != nil {
|
||||
return "", gerror.NewCode(common.CodeImageNotFound)
|
||||
}
|
||||
return path, nil
|
||||
}
|
||||
|
||||
// ExportZip 打包数据集图片目录为 zip(内存构建;controller 直写响应体下载)
|
||||
func (s *datasetService) ExportZip(ctx context.Context, datasetId int64) ([]byte, string, error) {
|
||||
dataset, err := dao.Dataset.GetById(ctx, datasetId)
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
if dataset == nil {
|
||||
return nil, "", gerror.NewCode(common.CodeDatasetNotFound)
|
||||
}
|
||||
dir := common.DatasetImagesDir(ctx, dataset.Name)
|
||||
var buf bytes.Buffer
|
||||
zw := zip.NewWriter(&buf)
|
||||
entries, err := os.ReadDir(dir)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return nil, "", gerror.New("数据集图片目录不存在")
|
||||
}
|
||||
return nil, "", err
|
||||
}
|
||||
for _, e := range entries {
|
||||
if e.IsDir() {
|
||||
continue
|
||||
}
|
||||
data, err := os.ReadFile(filepath.Join(dir, e.Name()))
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
w, err := zw.Create(e.Name())
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
if _, err := w.Write(data); err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
}
|
||||
if err := zw.Close(); err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
return buf.Bytes(), dataset.Name + ".zip", nil
|
||||
}
|
||||
@@ -0,0 +1,566 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"image"
|
||||
_ "image/jpeg"
|
||||
_ "image/png"
|
||||
"math"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gogf/gf/v2/errors/gerror"
|
||||
"github.com/gogf/gf/v2/frame/g"
|
||||
"github.com/gogf/gf/v2/os/gtime"
|
||||
|
||||
"observer-server/biz/consts"
|
||||
"observer-server/biz/dao"
|
||||
"observer-server/biz/model/dto"
|
||||
"observer-server/biz/model/entity"
|
||||
"observer-server/common"
|
||||
)
|
||||
|
||||
// labelTaskService 预标注工作台后端:RF-DETR 全图扫描,AI 标注结果直写
|
||||
// dataset_image.labels_json(与人工标注同存同编辑),人工可修改/清理全部框;
|
||||
// 训练前整理 YOLO 训练集。
|
||||
// 全图扫描规则(项目既定):不套用生成规格的位置裁剪,候选宁多勿漏。
|
||||
type labelTaskService struct{}
|
||||
|
||||
var LabelTask = &labelTaskService{}
|
||||
|
||||
// recoverLabelTasks 服务重启恢复:孤儿 running 预标注任务置 done + 错误提示
|
||||
// (RF-DETR 检测无状态,重新发起即可重新生成标注)。
|
||||
func (s *labelTaskService) recoverLabelTasks(ctx context.Context) {
|
||||
list, err := dao.LabelTask.ListRunning(ctx)
|
||||
if err != nil {
|
||||
g.Log().Errorf(ctx, "恢复预标注任务失败: %+v", err)
|
||||
return
|
||||
}
|
||||
for _, t := range list {
|
||||
if err := dao.LabelTask.Finish(ctx, t.Id, "", "服务重启,任务中断,可重新发起"); err != nil {
|
||||
g.Log().Errorf(ctx, "恢复预标注任务 %d 失败: %+v", t.Id, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// AdminListLabelTasks 预标注任务分页(组装数据集名)
|
||||
func (s *labelTaskService) AdminListLabelTasks(ctx context.Context, req *dto.AdminLabelTaskListReq) (*dto.AdminLabelTaskListRes, error) {
|
||||
page, size := common.NormalizePage(req.Page, req.Size)
|
||||
list, total, err := dao.LabelTask.Page(ctx, page, size)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
names := Training.datasetNameMap(ctx)
|
||||
items := make([]*dto.AdminLabelTaskItem, 0, len(list))
|
||||
for _, v := range list {
|
||||
items = append(items, &dto.AdminLabelTaskItem{
|
||||
Id: v.Id,
|
||||
DatasetId: v.DatasetId,
|
||||
DatasetName: names[v.DatasetId],
|
||||
Status: v.Status,
|
||||
Total: v.Total,
|
||||
Done: v.Done,
|
||||
Error: v.Error,
|
||||
CreatedAt: v.CreatedAt,
|
||||
FinishedAt: v.FinishedAt,
|
||||
})
|
||||
}
|
||||
return &dto.AdminLabelTaskListRes{Total: total, List: items}, nil
|
||||
}
|
||||
|
||||
// AdminStartLabelTask 发起预标注:串行检查(数据集存在 + 无 running 任务)→ 插任务 →
|
||||
// 池内逐张调 RF-DETR(全图扫描,AI 端点/模型取 config.yml localAi 节点),
|
||||
// 完成后标注直写 dataset_image.labels_json 置 done;任一图片失败则任务置 done + error。
|
||||
// 多选批量:req.Filenames 非空时只扫选中图。
|
||||
func (s *labelTaskService) AdminStartLabelTask(ctx context.Context, req *dto.AdminLabelTaskStartReq) (*dto.AdminLabelTaskStartRes, error) {
|
||||
dataset, err := dao.Dataset.GetById(ctx, req.DatasetId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if dataset == nil {
|
||||
return nil, gerror.NewCode(common.CodeDatasetNotFound)
|
||||
}
|
||||
if common.LocalAiClient(ctx) == nil {
|
||||
return nil, gerror.NewCode(common.CodeLocalAiNotConfigured)
|
||||
}
|
||||
images, err := dao.DatasetImage.ListByDataset(ctx, dataset.Id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(images) == 0 {
|
||||
return nil, gerror.New("数据集暂无图片")
|
||||
}
|
||||
// 多选批量:只处理选中图片(缺省全量)
|
||||
subset := len(req.Filenames) > 0
|
||||
filenamesJSON := ""
|
||||
if subset {
|
||||
sel := make(map[string]bool, len(req.Filenames))
|
||||
for _, f := range req.Filenames {
|
||||
sel[f] = true
|
||||
}
|
||||
filtered := make([]*entity.DatasetImage, 0, len(sel))
|
||||
for _, img := range images {
|
||||
if sel[img.Filename] {
|
||||
filtered = append(filtered, img)
|
||||
}
|
||||
}
|
||||
if len(filtered) == 0 {
|
||||
return nil, gerror.New("选中的图片不在该数据集内")
|
||||
}
|
||||
images = filtered
|
||||
names := make([]string, 0, len(images))
|
||||
for _, img := range images {
|
||||
names = append(names, img.Filename)
|
||||
}
|
||||
if raw, jErr := json.Marshal(names); jErr == nil {
|
||||
filenamesJSON = string(raw)
|
||||
}
|
||||
}
|
||||
taskId, err := s.startDetection(ctx, dataset, images, filenamesJSON)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &dto.AdminLabelTaskStartRes{Id: taskId}, nil
|
||||
}
|
||||
|
||||
// startDetection 发起预标注任务:Serial 内检查并发(已有 running 任务报错)→ 插任务 → 启动检测协程
|
||||
func (s *labelTaskService) startDetection(ctx context.Context, dataset *entity.Dataset, images []*entity.DatasetImage, filenamesJSON string) (int64, error) {
|
||||
now := gtime.Now()
|
||||
var taskId int64
|
||||
err := common.Serial().Submit(ctx, func() error {
|
||||
running, err := dao.LabelTask.GetRunningByDataset(ctx, dataset.Id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if running != nil {
|
||||
return gerror.NewCode(common.CodeLabelTaskRunning)
|
||||
}
|
||||
taskId, err = dao.LabelTask.Insert(ctx, &entity.LabelTask{
|
||||
DatasetId: dataset.Id,
|
||||
Status: consts.LabelTaskRunning,
|
||||
Total: len(images),
|
||||
Filenames: filenamesJSON,
|
||||
CreatedAt: now,
|
||||
})
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
s.runDetection(ctx, taskId, dataset, images, common.LocalAiClient(ctx))
|
||||
return taskId, nil
|
||||
}
|
||||
|
||||
// AutoLabel 图片入库自动触发标注(上传/AI 生成成功后调用):只标本次新增图。
|
||||
// localAi 未配置 → 报错(调用方拒绝/回滚,不允许产生无标注图);
|
||||
// 已有 running 任务(忙)→ 不报错:本次跳过,由任务成功完成后的自动补标轮兜底。
|
||||
func (s *labelTaskService) AutoLabel(ctx context.Context, dataset *entity.Dataset, images []*entity.DatasetImage) error {
|
||||
if len(images) == 0 {
|
||||
return nil
|
||||
}
|
||||
if common.LocalAiClient(ctx) == nil {
|
||||
return gerror.NewCode(common.CodeLocalAiNotConfigured)
|
||||
}
|
||||
names := make([]string, 0, len(images))
|
||||
for _, img := range images {
|
||||
names = append(names, img.Filename)
|
||||
}
|
||||
filenamesJSON := ""
|
||||
if raw, jErr := json.Marshal(names); jErr == nil {
|
||||
filenamesJSON = string(raw)
|
||||
}
|
||||
_, err := s.startDetection(ctx, dataset, images, filenamesJSON)
|
||||
if err != nil && gerror.HasCode(err, common.CodeLabelTaskRunning) {
|
||||
return nil // 忙:不报错
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// autoSupplement 自动补标:任务成功完成后,该数据集仍有未标注图(忙时入库的图/人工清空的图)
|
||||
// 则续一轮只标未标注图;失败任务不续(防配置故障时无限重试)。
|
||||
func (s *labelTaskService) autoSupplement(ctx context.Context, dataset *entity.Dataset) {
|
||||
unlabeled, err := dao.DatasetImage.ListUnlabeledByDataset(ctx, dataset.Id)
|
||||
if err != nil || len(unlabeled) == 0 {
|
||||
return
|
||||
}
|
||||
if err := s.AutoLabel(ctx, dataset, unlabeled); err != nil {
|
||||
g.Log().Infof(ctx, "自动补标未触发(数据集 %d): %v", dataset.Id, err)
|
||||
}
|
||||
}
|
||||
|
||||
// runDetection 预标注执行协程(生命周期任务):池内逐张检测,进度经 Serial 更新。
|
||||
// 每张:读图 → 全图检测(等比缩放提交,坐标映射回原图归一化)→ 标注(conf≥confConfirmed 为 class 0)。
|
||||
// 全部成功后 Serial 内逐张覆写 labels_json(重跑覆盖该图标注)。
|
||||
func (s *labelTaskService) runDetection(ctx context.Context, taskId int64, dataset *entity.Dataset, images []*entity.DatasetImage, client *common.LocalAi) {
|
||||
// 生命周期任务必须脱离请求 ctx:请求结束即取消,会让 Submit 秒退 + Finish 静默失败 → 任务悬挂
|
||||
bgCtx := context.Background()
|
||||
go func() {
|
||||
if client == nil {
|
||||
_ = dao.LabelTask.Finish(bgCtx, taskId, "", "标注服务未配置")
|
||||
return
|
||||
}
|
||||
dir := common.DatasetImagesDir(bgCtx, dataset.Name)
|
||||
results := make([][]*dto.AdminLabelBox, len(images))
|
||||
failed := ""
|
||||
doneCount := 0
|
||||
for i, img := range images {
|
||||
// 单张 120s 超时兜底(无超时 + 无取消的独立 ctx 下防检测服务悬挂拖死任务)
|
||||
detCtx, cancel := context.WithTimeout(bgCtx, 120*time.Second)
|
||||
err := common.LabelTaskPoolInstance().Submit(detCtx, func(ctx context.Context) error {
|
||||
data, err := os.ReadFile(filepath.Join(dir, img.Filename))
|
||||
if err != nil {
|
||||
return gerror.Wrapf(err, "读取图片失败")
|
||||
}
|
||||
w, h := imageSize(data)
|
||||
if w <= 0 || h <= 0 {
|
||||
return gerror.New("无法识别图片尺寸")
|
||||
}
|
||||
mime := imageMime(img.Filename)
|
||||
detections, err := client.Detect(ctx, data, mime, w, h)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
cands := make([]*dto.AdminLabelBox, 0, len(detections))
|
||||
for _, d := range detections {
|
||||
// 坐标映射回原图像素后归一化,越界轻微裁剪
|
||||
box := &dto.AdminLabelBox{
|
||||
Cx: clamp01((d.X + d.Width/2) / float64(w)),
|
||||
Cy: clamp01((d.Y + d.Height/2) / float64(h)),
|
||||
W: clamp01(d.Width / float64(w)),
|
||||
H: clamp01(d.Height / float64(h)),
|
||||
Confidence: d.Confidence,
|
||||
Class: 1,
|
||||
}
|
||||
if d.Confidence >= client.ConfConfirmed {
|
||||
box.Class = 0
|
||||
}
|
||||
cands = append(cands, box)
|
||||
}
|
||||
results[i] = suppressOverlap(cands, client.OverlapThreshold)
|
||||
return nil
|
||||
})
|
||||
cancel()
|
||||
if err != nil {
|
||||
failed = fmt.Sprintf("第 %d 张(%s)检测失败: %v", doneCount+1, img.Filename, err)
|
||||
break
|
||||
}
|
||||
doneCount++
|
||||
_ = common.Serial().Submit(bgCtx, func() error {
|
||||
return dao.LabelTask.UpdateProgress(bgCtx, taskId, doneCount)
|
||||
})
|
||||
}
|
||||
if failed != "" {
|
||||
_ = dao.LabelTask.Finish(bgCtx, taskId, "", failed)
|
||||
return
|
||||
}
|
||||
// 全部成功:Serial 内逐张覆写标注(重跑覆盖该图标注,人工保存同走 UpdateLabels)
|
||||
err := common.Serial().Submit(bgCtx, func() error {
|
||||
for i := range results {
|
||||
raw, jErr := json.Marshal(results[i])
|
||||
if jErr != nil {
|
||||
return gerror.Wrap(jErr, "标注序列化失败")
|
||||
}
|
||||
if uErr := dao.DatasetImage.UpdateLabels(bgCtx, images[i].Id, string(raw)); uErr != nil {
|
||||
return uErr
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
_ = dao.LabelTask.Finish(bgCtx, taskId, "", err.Error())
|
||||
return
|
||||
}
|
||||
_ = dao.LabelTask.Finish(bgCtx, taskId, "", "")
|
||||
// 成功完成:自动补标该数据集未标注图(忙时入库的图等),失败任务不续
|
||||
s.autoSupplement(bgCtx, dataset)
|
||||
}()
|
||||
}
|
||||
|
||||
// AdminLabelTaskDetail 预标注任务详情:标注(labels_json)输出,供工作台 canvas 叠框;
|
||||
// 图片尺寸读取文件头。
|
||||
func (s *labelTaskService) AdminLabelTaskDetail(ctx context.Context, req *dto.AdminLabelTaskDetailReq) (*dto.AdminLabelTaskDetailRes, error) {
|
||||
t, err := dao.LabelTask.GetById(ctx, req.Id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if t == nil {
|
||||
return nil, gerror.New("标注任务不存在")
|
||||
}
|
||||
dataset, err := dao.Dataset.GetById(ctx, t.DatasetId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if dataset == nil {
|
||||
return nil, gerror.NewCode(common.CodeDatasetNotFound)
|
||||
}
|
||||
items, err := s.buildWorkbenchItems(ctx, dataset)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &dto.AdminLabelTaskDetailRes{
|
||||
TaskId: t.Id,
|
||||
DatasetId: t.DatasetId,
|
||||
DatasetName: dataset.Name,
|
||||
Status: t.Status,
|
||||
Total: t.Total,
|
||||
Done: t.Done,
|
||||
Error: t.Error,
|
||||
Images: items,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// buildWorkbenchItems 工作台单张图数据:全部标注(labels_json,AI 自动标注与人工框同层);
|
||||
// 尺寸统一读文件头。
|
||||
func (s *labelTaskService) buildWorkbenchItems(ctx context.Context, dataset *entity.Dataset) ([]*dto.AdminLabelImageItem, error) {
|
||||
images, err := dao.DatasetImage.ListByDataset(ctx, dataset.Id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
imgDir := common.DatasetImagesDir(ctx, dataset.Name)
|
||||
items := make([]*dto.AdminLabelImageItem, 0, len(images))
|
||||
for _, img := range images {
|
||||
item := &dto.AdminLabelImageItem{
|
||||
Filename: img.Filename,
|
||||
Url: datasetImageUrl(ctx, dataset.Id, img.Filename),
|
||||
}
|
||||
if img.LabelsJson != "" && img.LabelsJson != "[]" {
|
||||
var boxes []*dto.AdminLabelBox
|
||||
if json.Unmarshal([]byte(img.LabelsJson), &boxes) == nil && len(boxes) > 0 {
|
||||
item.Boxes = boxes
|
||||
item.Labeled = true
|
||||
}
|
||||
}
|
||||
if data, rErr := os.ReadFile(filepath.Join(imgDir, img.Filename)); rErr == nil {
|
||||
item.Width, item.Height = imageSize(data)
|
||||
}
|
||||
items = append(items, item)
|
||||
}
|
||||
return items, nil
|
||||
}
|
||||
|
||||
// AdminLabelWorkbench 标注工作台数据:数据集无历史任务时直接输出图片 + 全部标注,
|
||||
// 标注读 dataset_image.labels_json(与任务详情同一组装)。
|
||||
func (s *labelTaskService) AdminLabelWorkbench(ctx context.Context, req *dto.AdminLabelWorkbenchReq) (*dto.AdminLabelWorkbenchRes, error) {
|
||||
dataset, err := dao.Dataset.GetById(ctx, req.DatasetId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if dataset == nil {
|
||||
return nil, gerror.NewCode(common.CodeDatasetNotFound)
|
||||
}
|
||||
items, err := s.buildWorkbenchItems(ctx, dataset)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &dto.AdminLabelWorkbenchRes{
|
||||
DatasetId: dataset.Id,
|
||||
DatasetName: dataset.Name,
|
||||
Images: items,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// AdminLabelSave 保存单张图标注:校验坐标 → json.Marshal 覆写 labels_json(空框=清空),
|
||||
// 并刷新数据集 labeled_count(labels_json 非空数组的图片数)。
|
||||
func (s *labelTaskService) AdminLabelSave(ctx context.Context, req *dto.AdminLabelSaveReq) (*dto.AdminLabelSaveRes, error) {
|
||||
dataset, err := dao.Dataset.GetById(ctx, req.DatasetId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if dataset == nil {
|
||||
return nil, gerror.NewCode(common.CodeDatasetNotFound)
|
||||
}
|
||||
img, err := dao.DatasetImage.GetByFilename(ctx, dataset.Id, req.Filename)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if img == nil {
|
||||
return nil, gerror.NewCode(common.CodeImageNotFound)
|
||||
}
|
||||
for _, box := range req.Boxes {
|
||||
if box.Cx < 0 || box.Cy < 0 || box.W <= 0 || box.H <= 0 || box.Cx > 1 || box.Cy > 1 {
|
||||
return nil, gerror.New("标注框坐标非法(需 0~1 归一化)")
|
||||
}
|
||||
}
|
||||
// 空框 = 清空标注
|
||||
raw := ""
|
||||
if len(req.Boxes) > 0 {
|
||||
b, err := json.Marshal(req.Boxes)
|
||||
if err != nil {
|
||||
return nil, gerror.Wrap(err, "标注序列化失败")
|
||||
}
|
||||
raw = string(b)
|
||||
}
|
||||
var labeled int64
|
||||
if err := common.Serial().Submit(ctx, func() error {
|
||||
if err := dao.DatasetImage.UpdateLabels(ctx, img.Id, raw); err != nil {
|
||||
return err
|
||||
}
|
||||
labeled, err = dao.DatasetImage.CountLabeledByDataset(ctx, dataset.Id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
status := ""
|
||||
if labeled > 0 {
|
||||
status = consts.DatasetStatusLabeled
|
||||
}
|
||||
return dao.Dataset.UpdateCounters(ctx, dataset.Id, 0, labeled, status)
|
||||
}); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &dto.AdminLabelSaveRes{LabeledCount: labeled}, nil
|
||||
}
|
||||
|
||||
// prepareYoloSet 训练前组装内存 YOLO 训练集包:已标注图(labels_json 非空)按 80/20 拆 train/val,
|
||||
// 标注 txt 内存生成、原图仅记源路径(由训练通道读取,不落本地暂存盘);无标注报错。
|
||||
// data.yaml 由训练发起方追加进包(path 需指向训练机)。
|
||||
func (s *labelTaskService) prepareYoloSet(ctx context.Context, dataset *entity.Dataset) (*common.YoloPackage, error) {
|
||||
images, err := dao.DatasetImage.ListByDataset(ctx, dataset.Id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
type item struct {
|
||||
filename string
|
||||
lines string
|
||||
}
|
||||
var kept []item
|
||||
for _, img := range images {
|
||||
if img.LabelsJson == "" || img.LabelsJson == "[]" {
|
||||
continue // 空标注图剔除
|
||||
}
|
||||
var boxes []*dto.AdminLabelBox
|
||||
if json.Unmarshal([]byte(img.LabelsJson), &boxes) != nil || len(boxes) == 0 {
|
||||
continue
|
||||
}
|
||||
var b strings.Builder
|
||||
for _, box := range boxes {
|
||||
fmt.Fprintf(&b, "%d %.6f %.6f %.6f %.6f\n", box.Class, box.Cx, box.Cy, box.W, box.H)
|
||||
}
|
||||
kept = append(kept, item{filename: img.Filename, lines: strings.TrimSpace(b.String())})
|
||||
}
|
||||
if len(kept) == 0 {
|
||||
return nil, gerror.New("数据集无有效标注,请先在标注工作台完成标注")
|
||||
}
|
||||
// 固定随机种子 + 20% val(至少 1 张,语义沿用原 prepare_yolo.py)
|
||||
idx := make([]int, len(kept))
|
||||
for i := range idx {
|
||||
idx[i] = i
|
||||
}
|
||||
randShuffle(idx)
|
||||
nVal := len(kept) / 5
|
||||
if nVal < 1 {
|
||||
nVal = 1
|
||||
}
|
||||
imgSrc := common.DatasetImagesDir(ctx, dataset.Name)
|
||||
pkg := &common.YoloPackage{}
|
||||
addSplit := func(split string, items []item) {
|
||||
for _, it := range items {
|
||||
pkg.Files = append(pkg.Files,
|
||||
common.YoloFile{Name: filepath.Join("images", split, it.filename), ImagePath: filepath.Join(imgSrc, it.filename)},
|
||||
common.YoloFile{Name: filepath.Join("labels", split, strings.TrimSuffix(it.filename, filepath.Ext(it.filename))+".txt"), Content: []byte(it.lines + "\n")},
|
||||
)
|
||||
}
|
||||
}
|
||||
var trainItems, valItems []item
|
||||
for i, it := range kept {
|
||||
if i < nVal {
|
||||
valItems = append(valItems, it)
|
||||
} else {
|
||||
trainItems = append(trainItems, it)
|
||||
}
|
||||
}
|
||||
addSplit("train", trainItems)
|
||||
addSplit("val", valItems)
|
||||
return pkg, nil
|
||||
}
|
||||
|
||||
// randShuffle Fisher-Yates 伪随机(固定种子,沿用原 prepare_yolo.py random.seed(42) 语义)
|
||||
func randShuffle(n []int) {
|
||||
state := uint32(42)
|
||||
seed := func() uint32 {
|
||||
state = state*1664525 + 1013904223
|
||||
return state
|
||||
}
|
||||
for i := len(n) - 1; i > 0; i-- {
|
||||
j := int(seed() % uint32(i+1))
|
||||
n[i], n[j] = n[j], n[i]
|
||||
}
|
||||
}
|
||||
|
||||
// suppressOverlap 重叠去重(NMS 风格):按置信度降序依次保留,与已保留框重叠比 > overlapThreshold
|
||||
// 的框剔除(RF-DETR 同一目标重复检出时多个高度重叠框,只留置信度最高者;跨 class 去重)。
|
||||
func suppressOverlap(boxes []*dto.AdminLabelBox, overlapThreshold float64) []*dto.AdminLabelBox {
|
||||
if len(boxes) <= 1 {
|
||||
return boxes
|
||||
}
|
||||
sorted := append([]*dto.AdminLabelBox(nil), boxes...)
|
||||
sort.Slice(sorted, func(i, j int) bool {
|
||||
return sorted[i].Confidence > sorted[j].Confidence
|
||||
})
|
||||
kept := make([]*dto.AdminLabelBox, 0, len(sorted))
|
||||
for i := range sorted {
|
||||
dup := false
|
||||
for _, k := range kept {
|
||||
if boxOverlap(sorted[i], k) > overlapThreshold {
|
||||
dup = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !dup {
|
||||
kept = append(kept, sorted[i])
|
||||
}
|
||||
}
|
||||
return kept
|
||||
}
|
||||
|
||||
// boxOverlap 两个归一化框(cx,cy,w,h)的重叠比(minIoU):交叠面积 / 两框较小面积。
|
||||
// 用 minIoU 而非 IoU:RF-DETR 对同一目标常输出一大一小两个框(IoU 仅 0.3~0.5),
|
||||
// 大框套小框时小框被覆盖比例高(0.3~0.9)能命中;相邻目标两框互有外露,minIoU 通常 < 0.3。
|
||||
func boxOverlap(a, b *dto.AdminLabelBox) float64 {
|
||||
ax1, ay1, ax2, ay2 := a.Cx-a.W/2, a.Cy-a.H/2, a.Cx+a.W/2, a.Cy+a.H/2
|
||||
bx1, by1, bx2, by2 := b.Cx-b.W/2, b.Cy-b.H/2, b.Cx+b.W/2, b.Cy+b.H/2
|
||||
ix1, iy1 := math.Max(ax1, bx1), math.Max(ay1, by1)
|
||||
ix2, iy2 := math.Min(ax2, bx2), math.Min(ay2, by2)
|
||||
if ix2 <= ix1 || iy2 <= iy1 {
|
||||
return 0
|
||||
}
|
||||
inter := (ix2 - ix1) * (iy2 - iy1)
|
||||
minArea := math.Min(a.W*a.H, b.W*b.H)
|
||||
if minArea <= 0 {
|
||||
return 0
|
||||
}
|
||||
return inter / minArea
|
||||
}
|
||||
|
||||
// imageSize 读取图片尺寸(文件头,不解码全图)
|
||||
func imageSize(data []byte) (int, int) {
|
||||
cfg, _, err := image.DecodeConfig(bytes.NewReader(data))
|
||||
if err != nil {
|
||||
return 0, 0
|
||||
}
|
||||
return cfg.Width, cfg.Height
|
||||
}
|
||||
|
||||
func imageMime(filename string) string {
|
||||
switch strings.ToLower(filepath.Ext(filename)) {
|
||||
case ".png":
|
||||
return "image/png"
|
||||
default:
|
||||
return "image/jpeg"
|
||||
}
|
||||
}
|
||||
|
||||
func clamp01(v float64) float64 {
|
||||
if v < 0 {
|
||||
return 0
|
||||
}
|
||||
if v > 1 {
|
||||
return 1
|
||||
}
|
||||
return v
|
||||
}
|
||||
@@ -50,7 +50,8 @@ func (s *licenseService) AdminListLicenses(ctx context.Context, req *dto.AdminLi
|
||||
for _, lic := range list {
|
||||
items = append(items, &dto.AdminLicenseItem{
|
||||
PhoneNum: lic.PhoneNum,
|
||||
ExpiresAt: lic.ExpiresAt, CreatedAt: lic.CreatedAt, UpdatedAt: lic.UpdatedAt,
|
||||
ExpiresAt: lic.ExpiresAt, Remark: lic.Remark,
|
||||
CreatedAt: lic.CreatedAt, UpdatedAt: lic.UpdatedAt,
|
||||
})
|
||||
}
|
||||
return &dto.AdminLicenseListRes{Total: total, List: items}, nil
|
||||
@@ -102,3 +103,18 @@ func (s *licenseService) AdminRevoke(ctx context.Context, req *dto.AdminRevokeRe
|
||||
common.ClearCache(ctx, "license:"+req.PhoneNum)
|
||||
return &dto.AdminRevokeRes{}, nil
|
||||
}
|
||||
|
||||
// AdminRemark 写账号备注(管理端运营记录):空串即清空,不覆盖密码/授权;
|
||||
// 单写者串行事务内更新,提交后清授权缓存(列表不缓存不受影响)。
|
||||
func (s *licenseService) AdminRemark(ctx context.Context, req *dto.AdminRemarkReq) (*dto.AdminRemarkRes, error) {
|
||||
err := common.Serial().Submit(ctx, func() error {
|
||||
return g.DB().Transaction(ctx, func(ctx context.Context, tx gdb.TX) error {
|
||||
return dao.License.SetRemarkInTx(ctx, tx, req.PhoneNum, req.Remark)
|
||||
})
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
common.ClearCache(ctx, "license:"+req.PhoneNum)
|
||||
return &dto.AdminRemarkRes{}, nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,53 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
|
||||
"observer-server/biz/dao"
|
||||
"observer-server/biz/model/dto"
|
||||
)
|
||||
|
||||
// modelVersionService 模型版本业务:管理端无模型管理界面(2026-08-26 决策),
|
||||
// 版本记录仅支撑客户端模型目录(App 多模型下载热更新)。文件布局见 common/workspace.go。
|
||||
type modelVersionService struct{}
|
||||
|
||||
var ModelVersion = &modelVersionService{}
|
||||
|
||||
// ClientCatalog 客户端模型目录:全部数据集当前生效模型(App 多模型下载热更新)。
|
||||
// downloadUrl 复用 /download 静态托管(/download/models/<数据集名>/latest.tflite)。
|
||||
func (s *modelVersionService) ClientCatalog(ctx context.Context) (*dto.ModelCatalogRes, error) {
|
||||
list, err := dao.ModelVersion.ListAllLatest(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
names := Training.datasetNameMap(ctx)
|
||||
items := make([]*dto.ModelCatalogItem, 0, len(list))
|
||||
for _, v := range list {
|
||||
datasetName := names[v.DatasetId]
|
||||
if datasetName == "" {
|
||||
continue
|
||||
}
|
||||
items = append(items, &dto.ModelCatalogItem{
|
||||
DatasetId: v.DatasetId,
|
||||
DatasetName: datasetName,
|
||||
Version: v.Version,
|
||||
Labels: parseLabels(v.Labels),
|
||||
SizeBytes: v.SizeBytes,
|
||||
Sha256: v.Sha256,
|
||||
Notes: v.Notes,
|
||||
PublishedAt: v.CreatedAt,
|
||||
DownloadUrl: "/download/models/" + datasetName + "/latest.tflite",
|
||||
})
|
||||
}
|
||||
return &dto.ModelCatalogRes{List: items}, nil
|
||||
}
|
||||
|
||||
// parseLabels 解析类别名 JSON 数组
|
||||
func parseLabels(s string) []string {
|
||||
var out []string
|
||||
if err := json.Unmarshal([]byte(s), &out); err != nil || len(out) == 0 {
|
||||
return []string{"class0", "class1"}
|
||||
}
|
||||
return out
|
||||
}
|
||||
@@ -0,0 +1,2 @@
|
||||
# 测试运行时产物(APK 上传目录),不入库
|
||||
workspace/
|
||||
+7
@@ -14,6 +14,9 @@ auth:
|
||||
secret: "test-auth-secret"
|
||||
tokenTtl: 2592000
|
||||
|
||||
admin:
|
||||
token: "test-admin-token"
|
||||
|
||||
plans:
|
||||
- id: day
|
||||
days: 1
|
||||
@@ -24,3 +27,7 @@ plans:
|
||||
- id: month
|
||||
days: 30
|
||||
price_cents: 18000
|
||||
|
||||
# 测试 APK 目录(testdata 下,避免污染仓库根)
|
||||
app:
|
||||
apkDir: "./testdata/workspace"
|
||||
|
||||
BIN
Binary file not shown.
@@ -0,0 +1,632 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gogf/gf/v2/errors/gerror"
|
||||
"github.com/gogf/gf/v2/frame/g"
|
||||
"github.com/gogf/gf/v2/os/gtime"
|
||||
|
||||
"observer-server/biz/consts"
|
||||
"observer-server/biz/dao"
|
||||
"observer-server/biz/model/dto"
|
||||
"observer-server/biz/model/entity"
|
||||
"observer-server/common"
|
||||
)
|
||||
|
||||
// trainingService 训练编排业务:发起训练(同步数据集 → 写任务参数 → 起进程)、
|
||||
// 后台轮询进度/判定结束、取消、发布模型版本。并发度 1(GPU 独占)。
|
||||
type trainingService struct{}
|
||||
|
||||
var Training = &trainingService{}
|
||||
|
||||
// StartBackgroundJobs 启动后台协程:训练进度轮询 + 孤儿预标注任务恢复(main.go 启动时调用)。
|
||||
// 单协程生命周期任务(非并行工作负载),不做池封装。
|
||||
func (s *trainingService) StartBackgroundJobs(ctx context.Context) {
|
||||
LabelTask.recoverLabelTasks(ctx)
|
||||
go s.pollTrainings(ctx)
|
||||
}
|
||||
|
||||
// pollTrainings 训练轮询:每 10s 扫描 running 任务,更新进度/日志,按 result.json 或进程
|
||||
// 存活判定结束;超时无结果判死。Go 重启后自动恢复扫描(进程已死 → 置 failed)。
|
||||
func (s *trainingService) pollTrainings(ctx context.Context) {
|
||||
ticker := time.NewTicker(10 * time.Second)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
}
|
||||
runner := common.Runner(ctx)
|
||||
if runner == nil {
|
||||
continue
|
||||
}
|
||||
cfg, ok := common.TrainingConfigOf(ctx)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
running, err := dao.Training.ListRunning(ctx)
|
||||
if err != nil {
|
||||
g.Log().Errorf(ctx, "训练轮询读取 running 任务失败: %+v", err)
|
||||
continue
|
||||
}
|
||||
for _, t := range running {
|
||||
job, err := s.buildJob(ctx, t, cfg)
|
||||
if err != nil {
|
||||
g.Log().Errorf(ctx, "训练 %d 构建任务参数失败: %+v", t.Id, err)
|
||||
continue
|
||||
}
|
||||
s.pollOne(ctx, runner, job, t)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (s *trainingService) pollOne(ctx context.Context, runner common.TrainingRunner, job *common.TrainingJob, t *entity.ModelTraining) {
|
||||
tail, err := runner.FetchLogTail(ctx, job)
|
||||
if err != nil {
|
||||
tail = ""
|
||||
}
|
||||
// 结束判定:result.json 存在 = 训练完成(先于存活判定,进程可能已退出)
|
||||
result, err := runner.FetchResult(ctx, job)
|
||||
if err != nil {
|
||||
g.Log().Errorf(ctx, "训练 %d 读取结果失败: %+v", t.Id, err)
|
||||
return
|
||||
}
|
||||
if result != "" {
|
||||
// result.json 带 error 字段 = 脚本异常退出(如 tflite 导出失败),置失败并带出原因
|
||||
var resErr struct {
|
||||
Error string `json:"error"`
|
||||
}
|
||||
if json.Unmarshal([]byte(result), &resErr) == nil && resErr.Error != "" {
|
||||
_ = s.finishFailed(ctx, t, "训练失败: %s", resErr.Error)
|
||||
return
|
||||
}
|
||||
s.finishSuccess(ctx, runner, job, t, result, tail)
|
||||
return
|
||||
}
|
||||
alive, err := runner.IsAlive(ctx, job)
|
||||
if err != nil {
|
||||
g.Log().Errorf(ctx, "训练 %d 存活探测失败: %+v", t.Id, err)
|
||||
return
|
||||
}
|
||||
// 超时判死(started_at 起算)
|
||||
cfg, _ := common.TrainingConfigOf(ctx)
|
||||
timeout := time.Duration(cfg.TimeoutMins) * time.Minute
|
||||
if t.StartedAt != nil && time.Since(t.StartedAt.Time) > timeout {
|
||||
_ = runner.Cancel(ctx, job)
|
||||
_ = s.finishFailed(ctx, t, "训练超时(超过 %d 分钟无结果,已终止)", cfg.TimeoutMins)
|
||||
return
|
||||
}
|
||||
if !alive {
|
||||
_ = s.finishFailed(ctx, t, "训练进程已退出(无结果文件)")
|
||||
return
|
||||
}
|
||||
// 进度:日志尾解析最后一条 epoch 行
|
||||
epoch, total, metrics := parseEpochTail(tail)
|
||||
if epoch > 0 {
|
||||
_ = dao.Training.UpdateProgress(ctx, t.Id, epoch, total, metrics, truncateTail(tail))
|
||||
}
|
||||
}
|
||||
|
||||
// finishSuccess 训练成功:解析 result.json(最终指标 + 类别名)→ 更新任务 + 拉取产物到服务器
|
||||
func (s *trainingService) finishSuccess(ctx context.Context, runner common.TrainingRunner, job *common.TrainingJob, t *entity.ModelTraining, result, tail string) {
|
||||
var res struct {
|
||||
Metrics map[string]float64 `json:"metrics"`
|
||||
Names []string `json:"names"`
|
||||
BestTflite string `json:"best_tflite"`
|
||||
TfliteCheck *struct {
|
||||
OK bool `json:"ok"`
|
||||
Reason string `json:"reason"`
|
||||
} `json:"tflite_check"`
|
||||
}
|
||||
_ = json.Unmarshal([]byte(result), &res)
|
||||
metricsJSON := ""
|
||||
if res.Metrics != nil {
|
||||
if b, err := json.Marshal(res.Metrics); err == nil {
|
||||
metricsJSON = string(b)
|
||||
}
|
||||
}
|
||||
// tflite 产物自检失败(训练脚本内嵌检查,旧任务无该字段不校验):坏产物不进发布链路
|
||||
if res.TfliteCheck != nil && !res.TfliteCheck.OK {
|
||||
reason := res.TfliteCheck.Reason
|
||||
if reason == "" {
|
||||
reason = "shape 校验未通过"
|
||||
}
|
||||
_ = s.finishFailed(ctx, t, "tflite 产物自检失败: %s", reason)
|
||||
return
|
||||
}
|
||||
// 先拉产物再置成功:产物拉取失败则置失败(发布依赖 tflite 存在)
|
||||
dest := common.TrainingArtifactsDir(ctx, t.Id)
|
||||
if res.BestTflite == "" {
|
||||
_ = s.finishFailed(ctx, t, "训练完成但 result.json 缺少 best_tflite")
|
||||
return
|
||||
}
|
||||
if err := runner.FetchArtifact(ctx, job, res.BestTflite, filepath.Join(dest, "best.tflite")); err != nil {
|
||||
g.Log().Errorf(ctx, "训练 %d 拉取 best.tflite 失败: %+v", t.Id, err)
|
||||
_ = s.finishFailed(ctx, t, "拉取训练产物失败: %v", err)
|
||||
return
|
||||
}
|
||||
zipPath := fmt.Sprintf("artifacts/%d.zip", t.Id)
|
||||
if err := runner.FetchArtifact(ctx, job, zipPath, filepath.Join(dest, "artifact.zip")); err != nil {
|
||||
g.Log().Warningf(ctx, "训练 %d 拉取 artifact.zip 失败(不阻断): %+v", t.Id, err)
|
||||
}
|
||||
// 指标尾部带上类别名,发布时解析 labels
|
||||
if len(res.Names) > 0 {
|
||||
if names, err := json.Marshal(res.Names); err == nil {
|
||||
metricsJSON = mergeNamesIntoMetrics(metricsJSON, string(names))
|
||||
}
|
||||
}
|
||||
if err := dao.Training.Finish(ctx, t.Id, consts.TrainingStatusSuccess, metricsJSON, truncateTail(tail), ""); err != nil {
|
||||
g.Log().Errorf(ctx, "训练 %d 置成功失败: %+v", t.Id, err)
|
||||
}
|
||||
}
|
||||
|
||||
// mergeNamesIntoMetrics 把 names 数组并入 metrics JSON(names 字段供发布解析类别名)
|
||||
func mergeNamesIntoMetrics(metricsJSON, namesJSON string) string {
|
||||
if metricsJSON == "" {
|
||||
return `{"names":` + namesJSON + `}`
|
||||
}
|
||||
var m map[string]any
|
||||
if json.Unmarshal([]byte(metricsJSON), &m) != nil {
|
||||
return metricsJSON
|
||||
}
|
||||
m["names"] = json.RawMessage(namesJSON)
|
||||
if b, err := json.Marshal(m); err == nil {
|
||||
return string(b)
|
||||
}
|
||||
return metricsJSON
|
||||
}
|
||||
|
||||
func (s *trainingService) finishFailed(ctx context.Context, t *entity.ModelTraining, format string, args ...any) error {
|
||||
return dao.Training.Finish(ctx, t.Id, consts.TrainingStatusFailed, "", "", fmt.Sprintf(format, args...))
|
||||
}
|
||||
|
||||
// buildJob 组装训练机路径布局的 runner 任务
|
||||
func (s *trainingService) buildJob(ctx context.Context, t *entity.ModelTraining, cfg common.TrainingConfig) (*common.TrainingJob, error) {
|
||||
dataset, err := dao.Dataset.GetById(ctx, t.DatasetId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if dataset == nil {
|
||||
return nil, gerror.NewCode(common.CodeDatasetNotFound)
|
||||
}
|
||||
job := &common.TrainingJob{
|
||||
TaskId: t.Id,
|
||||
DatasetName: dataset.Name,
|
||||
Python: cfg.Python,
|
||||
Workdir: cfg.Workdir,
|
||||
DatasetDir: cfg.DatasetDir,
|
||||
}
|
||||
return job, nil
|
||||
}
|
||||
|
||||
// parseEpochTail 从日志尾部解析最后一条 epoch 进度行({"epoch":N,"total":M,"metrics":{...}})
|
||||
func parseEpochTail(tail string) (epoch, total int, metrics string) {
|
||||
lines := strings.Split(tail, "\n")
|
||||
for i := len(lines) - 1; i >= 0; i-- {
|
||||
line := strings.TrimSpace(lines[i])
|
||||
if line == "" {
|
||||
continue
|
||||
}
|
||||
var e struct {
|
||||
Epoch int `json:"epoch"`
|
||||
Total int `json:"total"`
|
||||
Metrics map[string]float64 `json:"metrics"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(line), &e); err != nil || e.Epoch <= 0 {
|
||||
continue
|
||||
}
|
||||
metrics = ""
|
||||
if e.Metrics != nil {
|
||||
if b, err := json.Marshal(e.Metrics); err == nil {
|
||||
metrics = string(b)
|
||||
}
|
||||
}
|
||||
return e.Epoch, e.Total, metrics
|
||||
}
|
||||
return 0, 0, ""
|
||||
}
|
||||
|
||||
func truncateTail(s string) string {
|
||||
if len(s) > 8*1024 {
|
||||
return s[len(s)-8*1024:]
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
// AdminListTrainings 训练任务分页(组装数据集名)
|
||||
func (s *trainingService) AdminListTrainings(ctx context.Context, req *dto.AdminTrainingListReq) (*dto.AdminTrainingListRes, error) {
|
||||
page, size := common.NormalizePage(req.Page, req.Size)
|
||||
var list []*entity.ModelTraining
|
||||
var total int64
|
||||
var err error
|
||||
if req.Status != "" {
|
||||
list, total, err = dao.Training.PageByStatus(ctx, req.Status, page, size)
|
||||
} else {
|
||||
list, total, err = dao.Training.Page(ctx, page, size)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
names := s.datasetNameMap(ctx)
|
||||
items := make([]*dto.AdminTrainingItem, 0, len(list))
|
||||
for _, v := range list {
|
||||
items = append(items, &dto.AdminTrainingItem{
|
||||
Id: v.Id,
|
||||
Name: v.Name,
|
||||
DatasetId: v.DatasetId,
|
||||
DatasetName: names[v.DatasetId],
|
||||
Status: v.Status,
|
||||
Imgsz: v.Imgsz,
|
||||
Epochs: v.Epochs,
|
||||
Batch: v.Batch,
|
||||
Device: v.Device,
|
||||
CurrentEpoch: v.CurrentEpoch,
|
||||
TotalEpochs: v.TotalEpochs,
|
||||
Metrics: v.Metrics,
|
||||
Error: v.Error,
|
||||
StartedAt: v.StartedAt,
|
||||
FinishedAt: v.FinishedAt,
|
||||
CreatedAt: v.CreatedAt,
|
||||
})
|
||||
}
|
||||
return &dto.AdminTrainingListRes{Total: total, List: items}, nil
|
||||
}
|
||||
|
||||
// datasetNameMap 全量数据集 id → 名称(列表组装用,避免 N+1)
|
||||
func (s *trainingService) datasetNameMap(ctx context.Context) map[int64]string {
|
||||
m := map[int64]string{}
|
||||
list, err := dao.Dataset.ListAll(ctx)
|
||||
if err != nil {
|
||||
return m
|
||||
}
|
||||
for _, d := range list {
|
||||
m[d.Id] = d.Name
|
||||
}
|
||||
return m
|
||||
}
|
||||
|
||||
// AdminStartTraining 发起训练:并发度 1(已有 running 拒绝);先本地整理 yolo 训练集
|
||||
// (80/20 拆 train/val)再同步训练机 → 写任务参数 → 启动进程;任何一步失败置任务 failed。
|
||||
func (s *trainingService) AdminStartTraining(ctx context.Context, req *dto.AdminTrainingStartReq) (*dto.AdminTrainingStartRes, error) {
|
||||
runner := common.Runner(ctx)
|
||||
if runner == nil {
|
||||
return nil, gerror.NewCode(common.CodeTrainingNotConfigured)
|
||||
}
|
||||
cfg, ok := common.TrainingConfigOf(ctx)
|
||||
if !ok {
|
||||
return nil, gerror.NewCode(common.CodeTrainingNotConfigured)
|
||||
}
|
||||
dataset, err := dao.Dataset.GetById(ctx, req.DatasetId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if dataset == nil {
|
||||
return nil, gerror.NewCode(common.CodeDatasetNotFound)
|
||||
}
|
||||
// 组装 yolo 训练集包(有标注才可训练;内存组装,不落本地暂存盘)
|
||||
pkg, err := LabelTask.prepareYoloSet(ctx, dataset)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
imgsz := req.Imgsz
|
||||
if imgsz <= 0 {
|
||||
imgsz = 704
|
||||
}
|
||||
epochs := req.Epochs
|
||||
if epochs <= 0 {
|
||||
epochs = 150
|
||||
}
|
||||
batch := req.Batch
|
||||
if batch <= 0 {
|
||||
batch = 16
|
||||
}
|
||||
device := req.Device
|
||||
if device == "" {
|
||||
device = "0"
|
||||
}
|
||||
now := gtime.Now()
|
||||
var taskId int64
|
||||
err = common.Serial().Submit(ctx, func() error {
|
||||
running, err := dao.Training.Running(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if running != nil {
|
||||
return gerror.NewCode(common.CodeTrainingRunning)
|
||||
}
|
||||
taskId, err = dao.Training.Insert(ctx, &entity.ModelTraining{
|
||||
Name: req.Name,
|
||||
Status: consts.TrainingStatusRunning,
|
||||
DatasetId: dataset.Id,
|
||||
Imgsz: imgsz,
|
||||
Epochs: epochs,
|
||||
Batch: batch,
|
||||
Device: device,
|
||||
StartedAt: now,
|
||||
CreatedAt: now,
|
||||
})
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// 准备阶段(训练机侧)失败 → 任务置 failed(记录保留便于排查)
|
||||
job := &common.TrainingJob{
|
||||
TaskId: taskId,
|
||||
DatasetName: dataset.Name,
|
||||
Python: cfg.Python,
|
||||
Workdir: cfg.Workdir,
|
||||
DatasetDir: cfg.DatasetDir,
|
||||
}
|
||||
// data.yaml 的 path 指向训练机路径,随包一起同步
|
||||
trainPath := filepath.Join(cfg.Workdir, cfg.DatasetDir, "yolo", dataset.Name)
|
||||
pkg.Files = append(pkg.Files, common.YoloFile{
|
||||
Name: "dataset.yaml",
|
||||
Content: []byte(yoloYamlContent(trainPath, localAiClassNames(ctx))),
|
||||
})
|
||||
taskJSON, _ := json.Marshal(map[string]any{
|
||||
"workdir": cfg.Workdir,
|
||||
"yolo": filepath.ToSlash(filepath.Join(cfg.DatasetDir, "yolo", dataset.Name)),
|
||||
"imgsz": imgsz,
|
||||
"epochs": epochs,
|
||||
"batch": batch,
|
||||
"device": device,
|
||||
"project": filepath.ToSlash(filepath.Join("runs", "tasks", strconv.FormatInt(taskId, 10))),
|
||||
"log_file": filepath.ToSlash(filepath.Join("logs", strconv.FormatInt(taskId, 10)+".jsonl")),
|
||||
"result_file": filepath.ToSlash(filepath.Join("results", strconv.FormatInt(taskId, 10)+".json")),
|
||||
"artifact_zip": filepath.ToSlash(filepath.Join("artifacts", strconv.FormatInt(taskId, 10)+".zip")),
|
||||
})
|
||||
if err := runner.WriteTaskJson(ctx, job, string(taskJSON)); err != nil {
|
||||
_ = s.finishFailed(ctx, &entity.ModelTraining{Id: taskId}, "写任务参数失败: %v", err)
|
||||
return nil, err
|
||||
}
|
||||
if err := runner.SyncYoloDataset(ctx, job, pkg); err != nil {
|
||||
_ = s.finishFailed(ctx, &entity.ModelTraining{Id: taskId}, "同步数据集失败: %v", err)
|
||||
return nil, err
|
||||
}
|
||||
pid, err := runner.Start(ctx, job)
|
||||
if err != nil {
|
||||
_ = s.finishFailed(ctx, &entity.ModelTraining{Id: taskId}, "启动训练失败: %v", err)
|
||||
return nil, err
|
||||
}
|
||||
if err := dao.Training.UpdatePid(ctx, taskId, pid); err != nil {
|
||||
g.Log().Errorf(ctx, "训练 %d 记录 pid 失败: %+v", taskId, err)
|
||||
}
|
||||
return &dto.AdminTrainingStartRes{Id: taskId}, nil
|
||||
}
|
||||
|
||||
// yoloYamlContent 生成训练集 data.yaml 内容(path 为训练机绝对路径)
|
||||
func yoloYamlContent(trainPath string, names []string) string {
|
||||
var b strings.Builder
|
||||
fmt.Fprintf(&b, "path: %s\n", trainPath)
|
||||
b.WriteString("train: images/train\nval: images/val\nnames:\n")
|
||||
for i, n := range names {
|
||||
fmt.Fprintf(&b, " %d: %s\n", i, n)
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
|
||||
// localAiClassNames 标注类别名(config.yml localAi.classNames,默认 class0/class1)
|
||||
func localAiClassNames(ctx context.Context) []string {
|
||||
names := g.Cfg().MustGet(ctx, "localAi.classNames").Strings()
|
||||
if len(names) == 0 {
|
||||
return []string{"class0", "class1"}
|
||||
}
|
||||
return names
|
||||
}
|
||||
|
||||
// AdminTrainingDetail 训练任务详情(含日志尾部)
|
||||
func (s *trainingService) AdminTrainingDetail(ctx context.Context, req *dto.AdminTrainingDetailReq) (*dto.AdminTrainingDetailRes, error) {
|
||||
t, err := dao.Training.GetById(ctx, req.Id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if t == nil {
|
||||
return nil, gerror.NewCode(common.CodeTrainingNotFound)
|
||||
}
|
||||
names := s.datasetNameMap(ctx)
|
||||
return &dto.AdminTrainingDetailRes{
|
||||
AdminTrainingItem: dto.AdminTrainingItem{
|
||||
Id: t.Id,
|
||||
Name: t.Name,
|
||||
DatasetId: t.DatasetId,
|
||||
DatasetName: names[t.DatasetId],
|
||||
Status: t.Status,
|
||||
Imgsz: t.Imgsz,
|
||||
Epochs: t.Epochs,
|
||||
Batch: t.Batch,
|
||||
Device: t.Device,
|
||||
CurrentEpoch: t.CurrentEpoch,
|
||||
TotalEpochs: t.TotalEpochs,
|
||||
Metrics: t.Metrics,
|
||||
Error: t.Error,
|
||||
StartedAt: t.StartedAt,
|
||||
FinishedAt: t.FinishedAt,
|
||||
CreatedAt: t.CreatedAt,
|
||||
},
|
||||
LogTail: t.LogTail,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// AdminCancelTraining 取消训练:杀进程 + 置 failed
|
||||
func (s *trainingService) AdminCancelTraining(ctx context.Context, req *dto.AdminTrainingCancelReq) (*dto.AdminTrainingCancelRes, error) {
|
||||
var t *entity.ModelTraining
|
||||
err := common.Serial().Submit(ctx, func() error {
|
||||
var err error
|
||||
t, err = dao.Training.GetById(ctx, req.Id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if t == nil {
|
||||
return gerror.NewCode(common.CodeTrainingNotFound)
|
||||
}
|
||||
if t.Status != consts.TrainingStatusRunning {
|
||||
return gerror.New("仅运行中的训练任务可取消")
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
runner := common.Runner(ctx)
|
||||
if runner != nil {
|
||||
if cfg, ok := common.TrainingConfigOf(ctx); ok {
|
||||
if job, jErr := s.buildJob(ctx, t, cfg); jErr == nil {
|
||||
_ = runner.Cancel(ctx, job)
|
||||
}
|
||||
}
|
||||
}
|
||||
if err := s.finishFailed(ctx, t, "用户取消"); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &dto.AdminTrainingCancelRes{}, nil
|
||||
}
|
||||
|
||||
// AdminPublish 发布模型版本:仅 success 任务 + 本地 best.tflite 存在;
|
||||
// 版本号同数据集内 m<major>.<minor>.<patch> 自增(无记录从 m1.0.0 起)。
|
||||
// 落库(置旧版 is_latest=0 + 插新版)后写 latest 副本(无存档回退机制),文件失败补偿删记录,
|
||||
// 保证「记录存在 ⟺ 文件存在」(同 APK 版本管理)。
|
||||
func (s *trainingService) AdminPublish(ctx context.Context, req *dto.AdminTrainingPublishReq) (*dto.AdminTrainingPublishRes, error) {
|
||||
t, err := dao.Training.GetById(ctx, req.Id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if t == nil {
|
||||
return nil, gerror.NewCode(common.CodeTrainingNotFound)
|
||||
}
|
||||
if t.Status != consts.TrainingStatusSuccess {
|
||||
return nil, gerror.NewCode(common.CodeTrainingNotSuccess)
|
||||
}
|
||||
dataset, err := dao.Dataset.GetById(ctx, t.DatasetId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if dataset == nil {
|
||||
return nil, gerror.NewCode(common.CodeDatasetNotFound)
|
||||
}
|
||||
bestTflite := filepath.Join(common.TrainingArtifactsDir(ctx, t.Id), "best.tflite")
|
||||
data, err := os.ReadFile(bestTflite)
|
||||
if err != nil {
|
||||
return nil, gerror.New("训练产物 best.tflite 缺失,无法发布")
|
||||
}
|
||||
labels := labelsFromMetrics(t.Metrics)
|
||||
sha256, err := common.Sha256Hex(data)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
version, err := s.nextVersion(ctx, t.DatasetId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
now := gtime.Now()
|
||||
var mv *entity.ModelVersion
|
||||
err = common.Serial().Submit(ctx, func() error {
|
||||
// 该数据集旧版全部置 0,再插新版本(is_latest=1)
|
||||
if err := dao.ModelVersion.ClearLatest(ctx, t.DatasetId); err != nil {
|
||||
return err
|
||||
}
|
||||
id, err := dao.ModelVersion.Insert(ctx, &entity.ModelVersion{
|
||||
DatasetId: t.DatasetId,
|
||||
Version: version,
|
||||
TrainingId: t.Id,
|
||||
Metrics: t.Metrics,
|
||||
Labels: labelsJSON(labels),
|
||||
Sha256: sha256,
|
||||
SizeBytes: int64(len(data)),
|
||||
IsLatest: 1,
|
||||
Notes: t.Name,
|
||||
CreatedAt: now,
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
mv = &entity.ModelVersion{Id: id, Version: version, DatasetId: t.DatasetId}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// 文件:当前生效副本 latest.tflite(tmp+rename 原子覆盖),客户端固定下载
|
||||
dir := common.DatasetModelsDir(ctx, dataset.Name)
|
||||
if err := os.MkdirAll(dir, 0o755); err != nil {
|
||||
return nil, gerror.Wrap(err, "创建模型目录失败")
|
||||
}
|
||||
if err := common.WriteFileAtomic(filepath.Join(dir, "latest.tflite"), data); err != nil {
|
||||
_ = dao.ModelVersion.DeleteById(ctx, mv.Id)
|
||||
return nil, gerror.Wrap(err, "写当前生效模型失败")
|
||||
}
|
||||
return &dto.AdminTrainingPublishRes{Version: version}, nil
|
||||
}
|
||||
|
||||
// nextVersion 同数据集内版本号自增:取最大 m<x>.<y>.<z>,patch+1;无记录从 m1.0.0 起
|
||||
func (s *trainingService) nextVersion(ctx context.Context, datasetId int64) (string, error) {
|
||||
list, _, err := dao.ModelVersion.PageByDataset(ctx, datasetId, 1, 1000)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
maxPatch := 0
|
||||
maxMinor := 0
|
||||
maxMajor := 0
|
||||
for _, v := range list {
|
||||
major, minor, patch, ok := parseModelVersion(v.Version)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if major > maxMajor || major == maxMajor && (minor > maxMinor || minor == maxMinor && patch > maxPatch) {
|
||||
maxMajor, maxMinor, maxPatch = major, minor, patch
|
||||
}
|
||||
}
|
||||
if maxMajor == 0 && maxMinor == 0 && maxPatch == 0 {
|
||||
return consts.ModelVersionPrefix + "1.0.0", nil
|
||||
}
|
||||
return fmt.Sprintf("%s%d.%d.%d", consts.ModelVersionPrefix, maxMajor, maxMinor, maxPatch+1), nil
|
||||
}
|
||||
|
||||
// parseModelVersion 解析 m1.2.3 → (1,2,3,true)
|
||||
func parseModelVersion(v string) (major, minor, patch int, ok bool) {
|
||||
trimmed := strings.TrimPrefix(v, consts.ModelVersionPrefix)
|
||||
parts := strings.Split(trimmed, ".")
|
||||
if len(parts) != 3 {
|
||||
return 0, 0, 0, false
|
||||
}
|
||||
major, err1 := strconv.Atoi(parts[0])
|
||||
minor, err2 := strconv.Atoi(parts[1])
|
||||
patch, err3 := strconv.Atoi(parts[2])
|
||||
if err1 != nil || err2 != nil || err3 != nil {
|
||||
return 0, 0, 0, false
|
||||
}
|
||||
return major, minor, patch, true
|
||||
}
|
||||
|
||||
// labelsFromMetrics 从任务 metrics JSON 解析类别名(names 字段),缺省 class0/class1
|
||||
func labelsFromMetrics(metrics string) []string {
|
||||
if metrics != "" {
|
||||
var m map[string]json.RawMessage
|
||||
if json.Unmarshal([]byte(metrics), &m) == nil {
|
||||
if raw, ok := m["names"]; ok {
|
||||
var names []string
|
||||
if json.Unmarshal(raw, &names) == nil && len(names) > 0 {
|
||||
return names
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return []string{"class0", "class1"}
|
||||
}
|
||||
|
||||
func labelsJSON(labels []string) string {
|
||||
b, err := json.Marshal(labels)
|
||||
if err != nil {
|
||||
return `["class0","class1"]`
|
||||
}
|
||||
return string(b)
|
||||
}
|
||||
@@ -10,9 +10,14 @@ import (
|
||||
|
||||
// AdminAuth 管理端静态 token 鉴权中间件:请求头 X-Admin-Token 须与 config.yml
|
||||
// admin.token 一致;token 未配置(空)时管理接口全部拒绝。
|
||||
// 图片等 <img>/<video> 标签无法携带自定义请求头,允许 query 参数 token 兜底。
|
||||
func AdminAuth(r *ghttp.Request) {
|
||||
want := g.Cfg().MustGet(r.GetCtx(), "admin.token", "").String()
|
||||
if want == "" || r.Header.Get("X-Admin-Token") != want {
|
||||
token := r.Header.Get("X-Admin-Token")
|
||||
if token == "" {
|
||||
token = r.Get("token").String()
|
||||
}
|
||||
if want == "" || token != want {
|
||||
r.Response.WriteStatusExit(http.StatusUnauthorized, g.Map{
|
||||
"code": gcode.CodeNotAuthorized.Code(),
|
||||
"message": "管理端未授权",
|
||||
|
||||
@@ -0,0 +1,24 @@
|
||||
package common
|
||||
|
||||
import (
|
||||
"context"
|
||||
"path/filepath"
|
||||
|
||||
"github.com/gogf/gf/v2/frame/g"
|
||||
)
|
||||
|
||||
// APK 存储约定:Android 版本下发上传的 APK 存固定文件名(上传即覆盖,原子重命名),
|
||||
// 目录下永远只保留最新一个文件;下载由 main.go 静态托管 /download/<ApkFilename>,
|
||||
// URL 固定,客户端拼 apiBaseUrl 访问。
|
||||
const ApkFilename = "observer-latest.apk"
|
||||
|
||||
// ApkDir APK 上传目录:config.yml app.apkDir(默认 ./workspace,与 ./data 平级,
|
||||
// docker-compose 挂载持久化)
|
||||
func ApkDir(ctx context.Context) string {
|
||||
return g.Cfg().MustGet(ctx, "app.apkDir", "./workspace").String()
|
||||
}
|
||||
|
||||
// ApkFilePath 最新 APK 完整路径
|
||||
func ApkFilePath(ctx context.Context) string {
|
||||
return filepath.Join(ApkDir(ctx), ApkFilename)
|
||||
}
|
||||
@@ -52,6 +52,24 @@ func ClearCache(ctx context.Context, keys ...string) {
|
||||
}
|
||||
}
|
||||
|
||||
// EnsureColumn 存量库迁移:列缺失时 ALTER TABLE ADD COLUMN(SQLite 支持表尾追加)。
|
||||
// 新库由 CREATE TABLE 直接含列、已迁移库列已存在,均跳过;失败 panic(启动即暴露)。
|
||||
func EnsureColumn(ctx context.Context, table, column, ddl string) {
|
||||
res, err := g.DB().GetAll(ctx, "PRAGMA table_info("+table+")")
|
||||
if err != nil {
|
||||
panic("检查表结构失败 " + table + ": " + err.Error())
|
||||
}
|
||||
for _, r := range res {
|
||||
if gconv.String(r["name"]) == column {
|
||||
return
|
||||
}
|
||||
}
|
||||
if _, err := g.DB().Exec(ctx, "ALTER TABLE "+table+" ADD COLUMN "+ddl); err != nil {
|
||||
panic("迁移加列失败 " + table + "." + column + ": " + err.Error())
|
||||
}
|
||||
g.Log().Warningf(ctx, "存量表 %s 已迁移:新增列 %s", table, column)
|
||||
}
|
||||
|
||||
// DropLegacyTableIfHasColumn 存量库迁移:表存在旧版本废弃列(账号体系上线前的 device_id)
|
||||
// 时 DROP 整表(存量数据作废,用户决策),由 dao init 以新结构重建。
|
||||
func DropLegacyTableIfHasColumn(ctx context.Context, table, column string) {
|
||||
|
||||
+20
-6
@@ -5,10 +5,24 @@ import "github.com/gogf/gf/v2/errors/gcode"
|
||||
// 业务错误码(1000+,框架保留 <1000):统一 HTTP 200 + body code!=0 表示失败,
|
||||
// 客户端按 code 分支。错误统一用 gerror.NewCode(common.CodeXxx, "...") 构造。
|
||||
var (
|
||||
CodePlanNotConfigured = gcode.New(1001, "套餐不存在或未配置", nil)
|
||||
CodePaymentNotConfigured = gcode.New(1002, "支付渠道未配置", nil)
|
||||
CodeOrderNotFound = gcode.New(1003, "订单不存在", nil)
|
||||
CodeOrderClosed = gcode.New(1004, "订单已关闭,需重新下单", nil)
|
||||
CodeCallbackVerifyFailed = gcode.New(1005, "回调验签失败", nil)
|
||||
CodeCallbackMismatch = gcode.New(1006, "回调商户/金额不匹配", nil)
|
||||
CodePlanNotConfigured = gcode.New(1001, "套餐不存在或未配置", nil)
|
||||
CodePaymentNotConfigured = gcode.New(1002, "支付渠道未配置", nil)
|
||||
CodeOrderNotFound = gcode.New(1003, "订单不存在", nil)
|
||||
CodeOrderClosed = gcode.New(1004, "订单已关闭,需重新下单", nil)
|
||||
CodeCallbackVerifyFailed = gcode.New(1005, "回调验签失败", nil)
|
||||
CodeCallbackMismatch = gcode.New(1006, "回调商户/金额不匹配", nil)
|
||||
CodeVersionDuplicate = gcode.New(1007, "该版本号已存在,请勿重复下发", nil)
|
||||
CodeApkInvalid = gcode.New(1008, "请上传 APK 文件(.apk 后缀)", nil)
|
||||
CodeVersionNotFound = gcode.New(1009, "版本记录不存在", nil)
|
||||
CodeDatasetNotFound = gcode.New(1010, "数据集不存在", nil)
|
||||
CodeDatasetNameDuplicate = gcode.New(1011, "数据集名称已存在", nil)
|
||||
CodeImageGenNotConfigured = gcode.New(1012, "图像生成服务未配置(imageGen 节点)", nil)
|
||||
CodeImageGenFailed = gcode.New(1013, "图像生成失败", nil)
|
||||
CodeTrainingNotConfigured = gcode.New(1014, "训练通道未配置(training 节点)", nil)
|
||||
CodeTrainingRunning = gcode.New(1015, "已有训练任务进行中(并发度 1)", nil)
|
||||
CodeTrainingNotFound = gcode.New(1016, "训练任务不存在", nil)
|
||||
CodeTrainingNotSuccess = gcode.New(1017, "仅训练成功的任务可发布", nil)
|
||||
CodeLabelTaskRunning = gcode.New(1020, "该数据集已有预标注任务进行中", nil)
|
||||
CodeLocalAiNotConfigured = gcode.New(1021, "标注服务未配置(localAi 节点)", nil)
|
||||
CodeImageNotFound = gcode.New(1022, "图片不存在", nil)
|
||||
)
|
||||
|
||||
@@ -0,0 +1,160 @@
|
||||
package common
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gogf/gf/v2/frame/g"
|
||||
)
|
||||
|
||||
// ImageGenProvider 图像生成 provider 抽象:管理端「AI 生成图片」按 config.yml imageGen.provider 选择实现。
|
||||
type ImageGenProvider interface {
|
||||
// Generate 生成一张图片,返回图片字节;prompt 由调用方保证不含位置描述(项目提示词规范)。
|
||||
Generate(ctx context.Context, prompt, size string) ([]byte, error)
|
||||
}
|
||||
|
||||
// ImageGen 当前配置的图像生成 provider 单例(未配置返回 nil,调用方判 CodeImageGenNotConfigured)。
|
||||
func ImageGen(ctx context.Context) ImageGenProvider {
|
||||
if g.Cfg().MustGet(ctx, "imageGen.apiKey").String() == "" {
|
||||
return nil
|
||||
}
|
||||
switch g.Cfg().MustGet(ctx, "imageGen.provider", "dashscope").String() {
|
||||
case "dashscope":
|
||||
return &dashScopeImageGen{apiKey: g.Cfg().MustGet(ctx, "imageGen.apiKey").String()}
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
// dashScopeImageGen 通义万相(DashScope)实现:提交异步任务 → 轮询 task 状态 → 下载产物图片。
|
||||
// 文档:https://help.aliyun.com/zh/model-studio/text-to-image-api-reference
|
||||
type dashScopeImageGen struct {
|
||||
apiKey string
|
||||
}
|
||||
|
||||
const (
|
||||
dashScopeSynthUrl = "https://dashscope.aliyuncs.com/api/v1/services/aigc/text2image/image-synthesis"
|
||||
dashScopeTaskUrlFmt = "https://dashscope.aliyuncs.com/api/v1/tasks/%s"
|
||||
)
|
||||
|
||||
func (p *dashScopeImageGen) Generate(ctx context.Context, prompt, size string) ([]byte, error) {
|
||||
model := g.Cfg().MustGet(ctx, "imageGen.model", "qwen-image-3.0").String()
|
||||
// 前端尺寸格式 1152x2048 → API 规格 1152*2048
|
||||
apiSize := strings.ReplaceAll(size, "x", "*")
|
||||
body, err := json.Marshal(map[string]any{
|
||||
"model": model,
|
||||
"input": map[string]string{"prompt": prompt, "size": apiSize},
|
||||
"parameters": map[string]any{"n": 1, "watermark": false},
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
resp, err := p.post(ctx, dashScopeSynthUrl, body)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
taskId, ok := resp["task_id"].(string)
|
||||
if !ok || taskId == "" {
|
||||
return nil, fmt.Errorf("DashScope 提交失败: %v", resp)
|
||||
}
|
||||
// 轮询任务结果:异步生成通常 10~60s,上限 120s(与设计一致:超时 2min/张)
|
||||
deadline := time.Now().Add(120 * time.Second)
|
||||
for {
|
||||
if ctx.Err() != nil {
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
if time.Now().After(deadline) {
|
||||
return nil, fmt.Errorf("DashScope 生成超时(120s)")
|
||||
}
|
||||
task, err := p.post(ctx, fmt.Sprintf(dashScopeTaskUrlFmt, taskId), nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
status, _ := task["task_status"].(string)
|
||||
switch status {
|
||||
case "SUCCEEDED":
|
||||
results, _ := task["results"].([]any)
|
||||
if len(results) == 0 {
|
||||
return nil, fmt.Errorf("DashScope 成功但无产物图片")
|
||||
}
|
||||
url, _ := results[0].(map[string]any)["url"].(string)
|
||||
if url == "" {
|
||||
return nil, fmt.Errorf("DashScope 成功但无图片 URL")
|
||||
}
|
||||
return p.download(ctx, url)
|
||||
case "FAILED", "CANCELED":
|
||||
msg, _ := task["message"].(string)
|
||||
return nil, fmt.Errorf("DashScope 生成失败: %s", msg)
|
||||
}
|
||||
select {
|
||||
case <-time.After(2 * time.Second):
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// post 调用 DashScope HTTP 接口并解析统一 JSON(无 body 时为空 GET,轮询任务用)
|
||||
func (p *dashScopeImageGen) post(ctx context.Context, url string, body []byte) (map[string]any, error) {
|
||||
method := http.MethodPost
|
||||
var rd io.Reader
|
||||
if body == nil {
|
||||
method = http.MethodGet
|
||||
} else {
|
||||
rd = strings.NewReader(string(body))
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, method, url, rd)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req.Header.Set("Authorization", "Bearer "+p.apiKey)
|
||||
if body != nil {
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
}
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
raw, err := io.ReadAll(io.LimitReader(resp.Body, 8<<20))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("DashScope HTTP %d: %s", resp.StatusCode, truncateStr(string(raw), 200))
|
||||
}
|
||||
var m map[string]any
|
||||
if err := json.Unmarshal(raw, &m); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return m, nil
|
||||
}
|
||||
|
||||
// download 下载生成产物图片(存于阿里云 OSS,无需鉴权)
|
||||
func (p *dashScopeImageGen) download(ctx context.Context, url string) ([]byte, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("下载生成图片失败: HTTP %d", resp.StatusCode)
|
||||
}
|
||||
return io.ReadAll(io.LimitReader(resp.Body, 16<<20))
|
||||
}
|
||||
|
||||
func truncateStr(s string, n int) string {
|
||||
if len(s) <= n {
|
||||
return s
|
||||
}
|
||||
return s[:n] + "..."
|
||||
}
|
||||
@@ -0,0 +1,176 @@
|
||||
package common
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"image"
|
||||
"image/color"
|
||||
"image/jpeg"
|
||||
"image/png"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/gogf/gf/v2/frame/g"
|
||||
)
|
||||
|
||||
// LocalAi RF-DETR 检测服务客户端(local-ai 兼容 /v1/detection 协议):
|
||||
// POST {baseUrl}/v1/detection,body {"model","image":"data:image/jpeg;base64,...","threshold"},
|
||||
// 响应 {"detections":[{"x","y","width","height","confidence","class_name"}]},坐标单位 = 提交图片像素。
|
||||
type LocalAi struct {
|
||||
BaseUrl string
|
||||
Model string
|
||||
Threshold float64
|
||||
ConfConfirmed float64
|
||||
// 重叠去重阈值(minIoU = 交叠面积/两框较小面积):与高置信框重叠超过该值的框剔除(同目标只留一个)。
|
||||
// 用 minIoU 而非 IoU:RF-DETR 对同一目标常输出一大一小两个框,标准 IoU 可能仅 0.3~0.5 而漏杀,
|
||||
// 大框套小框时小框被覆盖比例高,minIoU 能命中;相邻目标两框互有外露,minIoU 通常 < 0.3。
|
||||
OverlapThreshold float64
|
||||
// 提交前整图等比缩放到的最长边(RF-DETR 对小图更稳)
|
||||
InputSize int
|
||||
}
|
||||
|
||||
// Detection RF-DETR 单目标检测结果(提交图比例尺下的像素坐标)
|
||||
type Detection struct {
|
||||
X float64 `json:"x"`
|
||||
Y float64 `json:"y"`
|
||||
Width float64 `json:"width"`
|
||||
Height float64 `json:"height"`
|
||||
Confidence float64 `json:"confidence"`
|
||||
ClassName string `json:"class_name"`
|
||||
}
|
||||
|
||||
// LocalAiClient 当前配置的标注服务客户端(未配置 baseUrl 返回 nil,调用方判 CodeLocalAiNotConfigured)
|
||||
func LocalAiClient(ctx context.Context) *LocalAi {
|
||||
base := g.Cfg().MustGet(ctx, "localAi.baseUrl").String()
|
||||
if base == "" {
|
||||
return nil
|
||||
}
|
||||
return &LocalAi{
|
||||
BaseUrl: strings.TrimRight(base, "/"),
|
||||
Model: g.Cfg().MustGet(ctx, "localAi.model", "rfdetr-xlarge").String(),
|
||||
Threshold: g.Cfg().MustGet(ctx, "localAi.threshold", 0.08).Float64(),
|
||||
ConfConfirmed: g.Cfg().MustGet(ctx, "localAi.confConfirmed", 0.2).Float64(),
|
||||
OverlapThreshold: g.Cfg().MustGet(ctx, "localAi.overlapThreshold", 0.3).Float64(),
|
||||
InputSize: g.Cfg().MustGet(ctx, "localAi.inputSize", 700).Int(),
|
||||
}
|
||||
}
|
||||
|
||||
// Detect 对单张图片做全图检测(不做任何裁剪,位置由模型自行推理):
|
||||
// 整图等比缩放至最长边 InputSize 提交,坐标映射回原图像素后返回。
|
||||
// imgW/imgH 为原图尺寸;返回坐标均为原图像素尺度。
|
||||
func (c *LocalAi) Detect(ctx context.Context, data []byte, mime string, imgW, imgH int) ([]*Detection, error) {
|
||||
sub, scale := c.prepare(data, mime, imgW, imgH)
|
||||
body, err := json.Marshal(map[string]any{
|
||||
"model": c.Model,
|
||||
"image": fmt.Sprintf("data:%s;base64,%s", mime, base64.StdEncoding.EncodeToString(sub)),
|
||||
"threshold": c.Threshold,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, c.BaseUrl+"/v1/detection",
|
||||
bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
raw, err := io.ReadAll(io.LimitReader(resp.Body, 8<<20))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("RF-DETR HTTP %d: %s", resp.StatusCode, truncateStr(string(raw), 200))
|
||||
}
|
||||
var out struct {
|
||||
Detections []*Detection `json:"detections"`
|
||||
}
|
||||
if err := json.Unmarshal(raw, &out); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// 坐标从提交图比例尺映射回原图
|
||||
for _, d := range out.Detections {
|
||||
d.X /= scale
|
||||
d.Y /= scale
|
||||
d.Width /= scale
|
||||
d.Height /= scale
|
||||
}
|
||||
return out.Detections, nil
|
||||
}
|
||||
|
||||
// prepare 整图等比缩放至最长边 InputSize(等比,不裁剪),返回提交字节与缩放比(原图/提交图)。
|
||||
func (c *LocalAi) prepare(data []byte, mime string, imgW, imgH int) ([]byte, float64) {
|
||||
if imgW <= 0 || imgH <= 0 || imgW <= c.InputSize && imgH <= c.InputSize {
|
||||
return data, 1
|
||||
}
|
||||
scale := float64(c.InputSize) / float64(maxInt(imgW, imgH))
|
||||
w, h := maxInt(1, int(float64(imgW)*scale)), maxInt(1, int(float64(imgH)*scale))
|
||||
src, _, err := image.Decode(bytes.NewReader(data))
|
||||
if err != nil {
|
||||
return data, 1
|
||||
}
|
||||
dst := bilinearResize(src, w, h)
|
||||
var buf bytes.Buffer
|
||||
if strings.Contains(mime, "png") {
|
||||
_ = png.Encode(&buf, dst)
|
||||
} else {
|
||||
_ = jpeg.Encode(&buf, dst, &jpeg.Options{Quality: 92})
|
||||
}
|
||||
return buf.Bytes(), scale
|
||||
}
|
||||
|
||||
// bilinearResize 双线性缩放(RF-DETR 对小图鲁棒,检测场景无需高质量插值)
|
||||
func bilinearResize(src image.Image, w, h int) *image.RGBA {
|
||||
b := src.Bounds()
|
||||
dst := image.NewRGBA(image.Rect(0, 0, w, h))
|
||||
if b.Dx() == 0 || b.Dy() == 0 {
|
||||
return dst
|
||||
}
|
||||
for y := 0; y < h; y++ {
|
||||
sy := float64(y) * float64(b.Dy()-1) / float64(maxInt(h-1, 1))
|
||||
y0, y1 := int(sy), minInt(int(sy)+1, b.Dy()-1)
|
||||
fy := sy - float64(y0)
|
||||
for x := 0; x < w; x++ {
|
||||
sx := float64(x) * float64(b.Dx()-1) / float64(maxInt(w-1, 1))
|
||||
x0, x1 := int(sx), minInt(int(sx)+1, b.Dx()-1)
|
||||
fx := sx - float64(x0)
|
||||
r00, g00, b00, _ := src.At(b.Min.X+x0, b.Min.Y+y0).RGBA()
|
||||
r10, g10, b10, _ := src.At(b.Min.X+x1, b.Min.Y+y0).RGBA()
|
||||
r01, g01, b01, _ := src.At(b.Min.X+x0, b.Min.Y+y1).RGBA()
|
||||
r11, g11, b11, _ := src.At(b.Min.X+x1, b.Min.Y+y1).RGBA()
|
||||
top := func(v00, v10 uint32) uint8 {
|
||||
return uint8((float64(v00)*(1-fx) + float64(v10)*fx) / 257)
|
||||
}
|
||||
bot := func(v01, v11 uint32) uint8 {
|
||||
return uint8((float64(v01)*(1-fx) + float64(v11)*fx) / 257)
|
||||
}
|
||||
r := uint8((float64(top(r00, r10))*(1-fy) + float64(bot(r01, r11))*fy))
|
||||
gx := uint8((float64(top(g00, g10))*(1-fy) + float64(bot(g01, g11))*fy))
|
||||
bb := uint8((float64(top(b00, b10))*(1-fy) + float64(bot(b01, b11))*fy))
|
||||
dst.Set(x, y, color.RGBA{R: r, G: gx, B: bb, A: 255})
|
||||
}
|
||||
}
|
||||
return dst
|
||||
}
|
||||
|
||||
func maxInt(a, b int) int {
|
||||
if a > b {
|
||||
return a
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
func minInt(a, b int) int {
|
||||
if a < b {
|
||||
return a
|
||||
}
|
||||
return b
|
||||
}
|
||||
@@ -50,3 +50,44 @@ func (p *CallbackPool) Submit(ctx context.Context, fn func(ctx context.Context)
|
||||
return ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
// LabelTaskPool 预标注任务并发池:逐张图片调 RF-DETR(IO 等待为主),
|
||||
// 并发度来自 config.yml labelTask.poolSize,缺失或非法时回退 consts 默认值。
|
||||
// 池内任务禁止提交本池(防 worker 饿死死锁);DB 写仍走 Serial 单写者。
|
||||
type LabelTaskPool struct {
|
||||
pool *grpool.Pool
|
||||
}
|
||||
|
||||
var (
|
||||
labelTaskPoolOnce sync.Once
|
||||
labelTaskPool *LabelTaskPool
|
||||
)
|
||||
|
||||
// LabelTaskPoolInstance 进程级预标注池单例(懒初始化,读取配置)。
|
||||
func LabelTaskPoolInstance() *LabelTaskPool {
|
||||
labelTaskPoolOnce.Do(func() {
|
||||
ctx := context.Background()
|
||||
size := g.Cfg().MustGet(ctx, "labelTask.poolSize", consts.LabelPoolDefaultSize).Int()
|
||||
if size <= 0 {
|
||||
size = consts.LabelPoolDefaultSize
|
||||
}
|
||||
labelTaskPool = &LabelTaskPool{pool: grpool.New(size, size)}
|
||||
})
|
||||
return labelTaskPool
|
||||
}
|
||||
|
||||
// Submit 提交单张图片的预标注任务并等待完成,返回任务的 error。
|
||||
func (p *LabelTaskPool) Submit(ctx context.Context, fn func(ctx context.Context) error) error {
|
||||
res := make(chan error, 1)
|
||||
if err := p.pool.Add(ctx, func(ctx context.Context) {
|
||||
res <- fn(ctx)
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
select {
|
||||
case err := <-res:
|
||||
return err
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,472 @@
|
||||
package common
|
||||
|
||||
import (
|
||||
"archive/tar"
|
||||
"bytes"
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"syscall"
|
||||
|
||||
"github.com/gogf/gf/v2/errors/gerror"
|
||||
"github.com/gogf/gf/v2/frame/g"
|
||||
"golang.org/x/crypto/ssh"
|
||||
)
|
||||
|
||||
// TrainingRunner 训练通道抽象(config.yml training.mode 选择实现):训练机与 Go 服务器
|
||||
// 可同机(subprocess)或异机(ssh)。任务脚本 server/training/train_server.py 以 --task-json 驱动,
|
||||
// 产物约定(相对训练机 workdir):
|
||||
//
|
||||
// tasks/<taskId>.json 任务参数(Go 侧写入)
|
||||
// logs/<taskId>.jsonl 每 epoch 一行 JSON:{"epoch","total","metrics"}
|
||||
// results/<taskId>.json 结束结果:{"metrics","names","best_tflite"(相对路径)}
|
||||
// artifacts/<taskId>.zip 打包产物(best.pt + results.csv + 曲线)
|
||||
type TrainingRunner interface {
|
||||
// Start 启动训练进程,返回可探测存活的 pid(subprocess 本机 pid;ssh 远程 pid)
|
||||
Start(ctx context.Context, job *TrainingJob) (int, error)
|
||||
// IsAlive 进程存活探测
|
||||
IsAlive(ctx context.Context, job *TrainingJob) (bool, error)
|
||||
// Cancel 终止训练(杀进程组,含 ultralytics 子进程)
|
||||
Cancel(ctx context.Context, job *TrainingJob) error
|
||||
// FetchLogTail 拉取日志尾部(截断 N KB 返回,供轮询解析 epoch 进度)
|
||||
FetchLogTail(ctx context.Context, job *TrainingJob) (string, error)
|
||||
// FetchResult 读取结束结果 JSON 内容;不存在返回 nil(任务仍在跑)
|
||||
FetchResult(ctx context.Context, job *TrainingJob) (string, error)
|
||||
// FetchArtifact 把训练机产物文件拉回服务器本地路径
|
||||
FetchArtifact(ctx context.Context, job *TrainingJob, remoteName, localPath string) error
|
||||
// SyncYoloDataset 把训练集包落到训练机(subprocess 直写 workdir;ssh 走 tar 流式管道,本地不落盘)
|
||||
SyncYoloDataset(ctx context.Context, job *TrainingJob, pkg *YoloPackage) error
|
||||
// WriteTaskJson 把任务参数文件写到训练机(随 Start 前的准备阶段调用)
|
||||
WriteTaskJson(ctx context.Context, job *TrainingJob, content string) error
|
||||
}
|
||||
|
||||
// YoloFile 训练集包内单个文件:Name 为包内相对路径(如 images/train/x.jpg / dataset.yaml),
|
||||
// ImagePath 非空时内容取自该文件(原图直接读数据集目录,不复制暂存),否则用 Content。
|
||||
type YoloFile struct {
|
||||
Name string
|
||||
ImagePath string
|
||||
Content []byte
|
||||
}
|
||||
|
||||
// YoloPackage 内存中的 YOLO 训练集包(训练前按 80/20 拆 train/val 组装,不落本地磁盘)
|
||||
type YoloPackage struct {
|
||||
Files []YoloFile
|
||||
}
|
||||
|
||||
// TrainingJob 训练任务运行信息(runner 视角;字段为训练机路径布局)。
|
||||
// SSH 凭据不随任务传递:sshRunner 直接读 config.yml training.ssh 节点。
|
||||
type TrainingJob struct {
|
||||
TaskId int64
|
||||
DatasetName string // yolo 数据集目录名(训练机 datasetDir/yolo/<name>)
|
||||
Python string // 训练机 venv python 路径
|
||||
Workdir string // 训练机工作目录(train_server.py / yolov8n.pt 所在)
|
||||
DatasetDir string // 训练机数据集根目录(相对 workdir)
|
||||
}
|
||||
|
||||
// Runner 按 config.yml training.mode 返回训练通道(未配置返回 nil,调用方判 CodeTrainingNotConfigured)
|
||||
func Runner(ctx context.Context) TrainingRunner {
|
||||
mode := g.Cfg().MustGet(ctx, "training.mode", "").String()
|
||||
switch mode {
|
||||
case "subprocess":
|
||||
return &subprocessRunner{}
|
||||
case "ssh":
|
||||
return &sshRunner{}
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
// TrainingConfig 训练通道配置快照
|
||||
type TrainingConfig struct {
|
||||
Workdir string
|
||||
DatasetDir string
|
||||
Python string
|
||||
TimeoutMins int
|
||||
}
|
||||
|
||||
// TrainingConfigOf 读取训练通道配置(未配置返回 ok=false)
|
||||
func TrainingConfigOf(ctx context.Context) (TrainingConfig, bool) {
|
||||
cfg := TrainingConfig{
|
||||
Workdir: g.Cfg().MustGet(ctx, "training.workdir").String(),
|
||||
DatasetDir: g.Cfg().MustGet(ctx, "training.datasetDir", "datasets").String(),
|
||||
Python: g.Cfg().MustGet(ctx, "training.venvPython").String(),
|
||||
TimeoutMins: g.Cfg().MustGet(ctx, "training.timeoutMinutes", 600).Int(),
|
||||
}
|
||||
return cfg, cfg.Workdir != "" && cfg.Python != ""
|
||||
}
|
||||
|
||||
// ---------------- subprocess 实现(同机) ----------------
|
||||
|
||||
type subprocessRunner struct {
|
||||
mu sync.Mutex
|
||||
cmds map[int64]*exec.Cmd // taskId → 进程(并发度 1,实际最多一条)
|
||||
}
|
||||
|
||||
func (r *subprocessRunner) Start(ctx context.Context, job *TrainingJob) (int, error) {
|
||||
script := filepath.Join(job.Workdir, "train_server.py")
|
||||
taskJson := filepath.Join(job.Workdir, "tasks", fmt.Sprintf("%d.json", job.TaskId))
|
||||
cmd := exec.Command(job.Python, script, "--task-json", taskJson)
|
||||
cmd.Dir = job.Workdir
|
||||
// 独立进程组:取消时 kill 整个组,连带 ultralytics 的子进程
|
||||
cmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true}
|
||||
if err := cmd.Start(); err != nil {
|
||||
return 0, gerror.Wrap(err, "启动训练进程失败")
|
||||
}
|
||||
r.mu.Lock()
|
||||
if r.cmds == nil {
|
||||
r.cmds = map[int64]*exec.Cmd{}
|
||||
}
|
||||
r.cmds[job.TaskId] = cmd
|
||||
r.mu.Unlock()
|
||||
return cmd.Process.Pid, nil
|
||||
}
|
||||
|
||||
func (r *subprocessRunner) IsAlive(ctx context.Context, job *TrainingJob) (bool, error) {
|
||||
r.mu.Lock()
|
||||
cmd := r.cmds[job.TaskId]
|
||||
r.mu.Unlock()
|
||||
if cmd == nil || cmd.Process == nil {
|
||||
return false, nil
|
||||
}
|
||||
err := cmd.Process.Signal(syscall.Signal(0))
|
||||
if err != nil {
|
||||
if err == os.ErrProcessDone {
|
||||
return false, nil
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func (r *subprocessRunner) Cancel(ctx context.Context, job *TrainingJob) error {
|
||||
r.mu.Lock()
|
||||
cmd := r.cmds[job.TaskId]
|
||||
r.mu.Unlock()
|
||||
if cmd == nil || cmd.Process == nil {
|
||||
return nil
|
||||
}
|
||||
// 杀进程组(负 pid),覆盖 python + ultralytics 子进程
|
||||
_ = syscall.Kill(-cmd.Process.Pid, syscall.SIGKILL)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *subprocessRunner) FetchLogTail(ctx context.Context, job *TrainingJob) (string, error) {
|
||||
path := filepath.Join(job.Workdir, "logs", fmt.Sprintf("%d.jsonl", job.TaskId))
|
||||
return tailFile(path, 8*1024)
|
||||
}
|
||||
|
||||
func (r *subprocessRunner) FetchResult(ctx context.Context, job *TrainingJob) (string, error) {
|
||||
path := filepath.Join(job.Workdir, "results", fmt.Sprintf("%d.json", job.TaskId))
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return "", nil
|
||||
}
|
||||
return "", err
|
||||
}
|
||||
return string(data), nil
|
||||
}
|
||||
|
||||
func (r *subprocessRunner) FetchArtifact(ctx context.Context, job *TrainingJob, remoteName, localPath string) error {
|
||||
src := filepath.Join(job.Workdir, filepath.FromSlash(remoteName))
|
||||
data, err := os.ReadFile(src)
|
||||
if err != nil {
|
||||
return gerror.Wrap(err, "读取训练机产物失败")
|
||||
}
|
||||
return WriteFileAtomic(localPath, data)
|
||||
}
|
||||
|
||||
func (r *subprocessRunner) SyncYoloDataset(ctx context.Context, job *TrainingJob, pkg *YoloPackage) error {
|
||||
// 同机训练:训练进程直接读 workdir 下文件,包直写目标目录(无中间暂存)
|
||||
dst := filepath.Join(job.Workdir, job.DatasetDir, "yolo", job.DatasetName)
|
||||
_ = os.RemoveAll(dst)
|
||||
for _, f := range pkg.Files {
|
||||
data := f.Content
|
||||
if f.ImagePath != "" {
|
||||
imgData, rErr := os.ReadFile(f.ImagePath)
|
||||
if rErr != nil {
|
||||
return gerror.Wrapf(rErr, "读取原图失败: %s", f.ImagePath)
|
||||
}
|
||||
data = imgData
|
||||
}
|
||||
if err := WriteFileAtomic(filepath.Join(dst, filepath.FromSlash(f.Name)), data); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *subprocessRunner) WriteTaskJson(ctx context.Context, job *TrainingJob, content string) error {
|
||||
dir := filepath.Join(job.Workdir, "tasks")
|
||||
if err := os.MkdirAll(dir, 0o755); err != nil {
|
||||
return err
|
||||
}
|
||||
return WriteFileAtomic(filepath.Join(dir, fmt.Sprintf("%d.json", job.TaskId)), []byte(content))
|
||||
}
|
||||
|
||||
// ---------------- ssh 实现(异机) ----------------
|
||||
|
||||
type sshRunner struct {
|
||||
mu sync.Mutex
|
||||
pids map[int64]int // taskId → 远程 pid
|
||||
}
|
||||
|
||||
// dial 建立 SSH 连接(凭据直读 config.yml training.ssh:privateKeyPath 与 password 二选一)
|
||||
func (r *sshRunner) dial(ctx context.Context) (*ssh.Client, error) {
|
||||
host := g.Cfg().MustGet(ctx, "training.ssh.host").String()
|
||||
if host == "" {
|
||||
return nil, gerror.New("training.ssh 未配置 host")
|
||||
}
|
||||
user := g.Cfg().MustGet(ctx, "training.ssh.user").String()
|
||||
port := g.Cfg().MustGet(ctx, "training.ssh.port", 22).Int()
|
||||
password := g.Cfg().MustGet(ctx, "training.ssh.password").String()
|
||||
keyPath := g.Cfg().MustGet(ctx, "training.ssh.privateKeyPath").String()
|
||||
cfg := &ssh.ClientConfig{
|
||||
User: user,
|
||||
HostKeyCallback: ssh.InsecureIgnoreHostKey(), // 内网训练机,信任首次连接
|
||||
Timeout: 15e9,
|
||||
}
|
||||
if keyPath != "" {
|
||||
data, err := os.ReadFile(keyPath)
|
||||
if err != nil {
|
||||
return nil, gerror.Wrap(err, "读取 SSH 私钥失败")
|
||||
}
|
||||
signer, err := ssh.ParsePrivateKey(data)
|
||||
if err != nil {
|
||||
return nil, gerror.Wrap(err, "解析 SSH 私钥失败")
|
||||
}
|
||||
cfg.Auth = []ssh.AuthMethod{ssh.PublicKeys(signer)}
|
||||
} else if password != "" {
|
||||
cfg.Auth = []ssh.AuthMethod{ssh.Password(password)}
|
||||
} else {
|
||||
return nil, gerror.New("training.ssh 未配置认证方式(私钥或密码)")
|
||||
}
|
||||
return ssh.Dial("tcp", fmt.Sprintf("%s:%d", host, port), cfg)
|
||||
}
|
||||
|
||||
// runCmd 远程执行单条命令,返回 stdout
|
||||
func (r *sshRunner) runCmd(ctx context.Context, job *TrainingJob, cmd string) (string, error) {
|
||||
client, err := r.dial(ctx)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer func() { _ = client.Close() }()
|
||||
session, err := client.NewSession()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer func() { _ = session.Close() }()
|
||||
var buf bytes.Buffer
|
||||
session.Stdout = &buf
|
||||
session.Stderr = &buf
|
||||
if err := session.Run(cmd); err != nil {
|
||||
return "", gerror.Wrapf(err, "远程命令失败: %s", cmd)
|
||||
}
|
||||
return buf.String(), nil
|
||||
}
|
||||
|
||||
func (r *sshRunner) Start(ctx context.Context, job *TrainingJob) (int, error) {
|
||||
taskJson := filepath.Join(job.Workdir, "tasks", fmt.Sprintf("%d.json", job.TaskId))
|
||||
cmd := fmt.Sprintf("cd %s && nohup %s train_server.py --task-json %s > logs/%d.stdout 2>&1 & echo $!",
|
||||
job.Workdir, job.Python, taskJson, job.TaskId)
|
||||
out, err := r.runCmd(ctx, job, cmd)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
pid, err := strconv.Atoi(strings.TrimSpace(out))
|
||||
if err != nil {
|
||||
return 0, gerror.New("SSH 启动未返回 pid: " + out)
|
||||
}
|
||||
r.mu.Lock()
|
||||
if r.pids == nil {
|
||||
r.pids = map[int64]int{}
|
||||
}
|
||||
r.pids[job.TaskId] = pid
|
||||
r.mu.Unlock()
|
||||
return pid, nil
|
||||
}
|
||||
|
||||
func (r *sshRunner) IsAlive(ctx context.Context, job *TrainingJob) (bool, error) {
|
||||
r.mu.Lock()
|
||||
pid := r.pids[job.TaskId]
|
||||
r.mu.Unlock()
|
||||
if pid == 0 {
|
||||
return false, nil
|
||||
}
|
||||
out, err := r.runCmd(ctx, job, fmt.Sprintf("kill -0 %d 2>/dev/null && echo alive || echo dead", pid))
|
||||
if err != nil {
|
||||
return false, nil // 连接故障按"状态未知"处理,不误判任务结束
|
||||
}
|
||||
return strings.TrimSpace(out) == "alive", nil
|
||||
}
|
||||
|
||||
func (r *sshRunner) Cancel(ctx context.Context, job *TrainingJob) error {
|
||||
r.mu.Lock()
|
||||
pid := r.pids[job.TaskId]
|
||||
r.mu.Unlock()
|
||||
if pid == 0 {
|
||||
return nil
|
||||
}
|
||||
_, err := r.runCmd(ctx, job, fmt.Sprintf("kill -9 %d 2>/dev/null; pkill -9 -P %d 2>/dev/null; true", pid, pid))
|
||||
return err
|
||||
}
|
||||
|
||||
func (r *sshRunner) FetchLogTail(ctx context.Context, job *TrainingJob) (string, error) {
|
||||
return r.runCmd(ctx, job, fmt.Sprintf("tail -c 8192 %s/logs/%d.jsonl", job.Workdir, job.TaskId))
|
||||
}
|
||||
|
||||
func (r *sshRunner) FetchResult(ctx context.Context, job *TrainingJob) (string, error) {
|
||||
out, err := r.runCmd(ctx, job, fmt.Sprintf("cat %s/results/%d.json 2>/dev/null", job.Workdir, job.TaskId))
|
||||
if err != nil || out == "" {
|
||||
return "", nil // 文件不存在 = 仍在跑
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (r *sshRunner) FetchArtifact(ctx context.Context, job *TrainingJob, remoteName, localPath string) error {
|
||||
client, err := r.dial(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = client.Close() }()
|
||||
session, err := client.NewSession()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = session.Close() }()
|
||||
src := filepath.Join(job.Workdir, filepath.FromSlash(remoteName))
|
||||
data, err := session.Output(fmt.Sprintf("cat %s", src))
|
||||
if err != nil {
|
||||
return gerror.Wrap(err, "拉取训练机产物失败")
|
||||
}
|
||||
return WriteFileAtomic(localPath, data)
|
||||
}
|
||||
|
||||
func (r *sshRunner) SyncYoloDataset(ctx context.Context, job *TrainingJob, pkg *YoloPackage) error {
|
||||
// tar 流式管道:原图直接读数据集目录打包,stdin 推远端解包,本地不落盘
|
||||
client, err := r.dial(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = client.Close() }()
|
||||
session, err := client.NewSession()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = session.Close() }()
|
||||
parent := fmt.Sprintf("%s/%s", job.Workdir, job.DatasetDir)
|
||||
dst := fmt.Sprintf("%s/%s/yolo/%s", job.Workdir, job.DatasetDir, job.DatasetName)
|
||||
// tar 解包到 parent,条目前缀 yolo/<name>/ 即落到 dst
|
||||
cmd := fmt.Sprintf("mkdir -p %s && rm -rf %s && tar -xf - -C %s", parent, dst, parent)
|
||||
pr, pw := io.Pipe()
|
||||
session.Stdin = pr
|
||||
tarErr := make(chan error, 1)
|
||||
go func() {
|
||||
defer pw.Close() // 失败路径也必须关管道,否则远端 tar 等不到 EOF
|
||||
tarErr <- writeYoloTar(pw, job.DatasetName, pkg)
|
||||
}()
|
||||
runErr := session.Run(cmd)
|
||||
_ = pr.Close()
|
||||
if werr := <-tarErr; werr != nil {
|
||||
return werr
|
||||
}
|
||||
if runErr != nil {
|
||||
return gerror.Wrap(runErr, "tar 同步数据集失败")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// writeYoloTar 把训练集包写入 tar 流(条目前缀 yolo/<name>/,供远端 tar -xf 解包)
|
||||
func writeYoloTar(w io.Writer, datasetName string, pkg *YoloPackage) error {
|
||||
tw := tar.NewWriter(w)
|
||||
for _, f := range pkg.Files {
|
||||
data := f.Content
|
||||
if f.ImagePath != "" {
|
||||
imgData, rErr := os.ReadFile(f.ImagePath)
|
||||
if rErr != nil {
|
||||
return gerror.Wrapf(rErr, "读取原图失败: %s", f.ImagePath)
|
||||
}
|
||||
data = imgData
|
||||
}
|
||||
hdr := &tar.Header{
|
||||
Name: filepath.ToSlash(filepath.Join("yolo", datasetName, f.Name)),
|
||||
Mode: 0o644,
|
||||
Size: int64(len(data)),
|
||||
}
|
||||
if err := tw.WriteHeader(hdr); err != nil {
|
||||
return gerror.Wrap(err, "写 tar 头失败")
|
||||
}
|
||||
if _, err := tw.Write(data); err != nil {
|
||||
return gerror.Wrap(err, "写 tar 内容失败")
|
||||
}
|
||||
}
|
||||
if err := tw.Close(); err != nil {
|
||||
return gerror.Wrap(err, "关闭 tar 流失败")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *sshRunner) WriteTaskJson(ctx context.Context, job *TrainingJob, content string) error {
|
||||
// 通过 ssh stdin 管道写文件(避免命令行转义地狱)
|
||||
client, err := r.dial(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = client.Close() }()
|
||||
session, err := client.NewSession()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = session.Close() }()
|
||||
dir := filepath.Join(job.Workdir, "tasks")
|
||||
path := filepath.Join(dir, fmt.Sprintf("%d.json", job.TaskId))
|
||||
cmd := fmt.Sprintf("mkdir -p %s && cat > %s", dir, path)
|
||||
session.Stdin = strings.NewReader(content)
|
||||
return session.Run(cmd)
|
||||
}
|
||||
|
||||
// ---------------- 文件工具 ----------------
|
||||
|
||||
// atomicWriteFile tmp + rename 原子写(避免中断产生半截文件)
|
||||
func WriteFileAtomic(path string, data []byte) error {
|
||||
dir := filepath.Dir(path)
|
||||
if err := os.MkdirAll(dir, 0o755); err != nil {
|
||||
return err
|
||||
}
|
||||
tmp := path + ".tmp"
|
||||
if err := os.WriteFile(tmp, data, 0o644); err != nil {
|
||||
return err
|
||||
}
|
||||
return os.Rename(tmp, path)
|
||||
}
|
||||
|
||||
// tailFile 读文件末尾最多 maxBytes 字节
|
||||
func tailFile(path string, maxBytes int64) (string, error) {
|
||||
f, err := os.Open(path)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return "", nil
|
||||
}
|
||||
return "", err
|
||||
}
|
||||
defer func() { _ = f.Close() }()
|
||||
st, err := f.Stat()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
start := st.Size() - maxBytes
|
||||
if start < 0 {
|
||||
start = 0
|
||||
}
|
||||
buf := make([]byte, st.Size()-start)
|
||||
if _, err := f.ReadAt(buf, start); err != nil && err != io.EOF {
|
||||
return "", err
|
||||
}
|
||||
return string(buf), nil
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
package common
|
||||
|
||||
import "crypto/rand"
|
||||
|
||||
// UuidV4 生成 UUIDv4 字符串(crypto/rand 16 字节 + 版本/变体位,无第三方依赖);
|
||||
// 用于封面等文件名,避免固定命名碰撞。
|
||||
func UuidV4() string {
|
||||
var b [16]byte
|
||||
if _, err := rand.Read(b[:]); err != nil {
|
||||
panic(err) // crypto/rand 失败属系统级故障,直接崩溃重启
|
||||
}
|
||||
b[6] = (b[6] & 0x0f) | 0x40 // version 4
|
||||
b[8] = (b[8] & 0x3f) | 0x80 // variant 10
|
||||
const hex = "0123456789abcdef"
|
||||
out := make([]byte, 36)
|
||||
j := 0
|
||||
for i := 0; i < 16; i++ {
|
||||
if i == 4 || i == 6 || i == 8 || i == 10 {
|
||||
out[j] = '-'
|
||||
j++
|
||||
}
|
||||
out[j] = hex[b[i]>>4]
|
||||
j++
|
||||
out[j] = hex[b[i]&0x0f]
|
||||
j++
|
||||
}
|
||||
return string(out)
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
package common
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
|
||||
"github.com/gogf/gf/v2/frame/g"
|
||||
)
|
||||
|
||||
// 训练体系运行时数据布局(config.yml app.datasetDir,默认 ./workspace,挂载持久化、不提交 git):
|
||||
//
|
||||
// datasets/<name>/ 数据集图片(平铺,文件名唯一,标注存 DB dataset_image.labels_json)
|
||||
// models/<name>/ latest.tflite 当前生效副本(无存档,客户端固定下载)
|
||||
// trainings/<taskId>/ 训练产物(best.tflite + artifact.zip)
|
||||
//
|
||||
// 训练机与 Go 服务器异机时,数据集经 training 通道同步(见 common/training_runner.go)。
|
||||
|
||||
// DatasetDir 训练体系运行时数据根目录:config.yml app.datasetDir(默认 ./workspace)
|
||||
func DatasetDir(ctx context.Context) string {
|
||||
return g.Cfg().MustGet(ctx, "app.datasetDir", "./workspace").String()
|
||||
}
|
||||
|
||||
// DatasetImagesDir 某数据集图片目录
|
||||
func DatasetImagesDir(ctx context.Context, datasetName string) string {
|
||||
return filepath.Join(DatasetDir(ctx), "datasets", datasetName)
|
||||
}
|
||||
|
||||
// DatasetModelsDir 某数据集模型目录
|
||||
func DatasetModelsDir(ctx context.Context, datasetName string) string {
|
||||
return filepath.Join(DatasetDir(ctx), "models", datasetName)
|
||||
}
|
||||
|
||||
// TrainingArtifactsDir 某训练任务产物目录(拉回的 best.tflite + artifact.zip)
|
||||
func TrainingArtifactsDir(ctx context.Context, taskId int64) string {
|
||||
return filepath.Join(DatasetDir(ctx), "trainings", strconv.FormatInt(taskId, 10))
|
||||
}
|
||||
|
||||
// Sha256Hex 计算文件内容 SHA-256 十六进制(模型版本校验用)
|
||||
func Sha256Hex(data []byte) (string, error) {
|
||||
sum := sha256.Sum256(data)
|
||||
return hex.EncodeToString(sum[:]), nil
|
||||
}
|
||||
@@ -2,6 +2,8 @@
|
||||
server:
|
||||
address: ":18080"
|
||||
openapiPath: "/api.json"
|
||||
# 请求体上限(默认 8MB):APK 上传(管理端下发新版本)可能数百 MB,必须放大
|
||||
clientMaxBodySize: "512mb"
|
||||
|
||||
database:
|
||||
default:
|
||||
@@ -22,6 +24,48 @@ admin:
|
||||
# 后台管理端静态 token:请求头 X-Admin-Token 必须匹配;为空时管理接口全部拒绝
|
||||
token: "Tongli686^*^"
|
||||
|
||||
app:
|
||||
# Android APK 上传目录(运行时数据,与 ./data 平级、docker-compose 挂载持久化;
|
||||
# 目录下永远只保留最新一个文件 observer-latest.apk,下载地址固定 /download/observer-latest.apk)
|
||||
apkDir: "./workspace"
|
||||
# 模型训练体系运行时数据根目录(图片/标注/模型/训练产物;与 apkDir 同目录时共用 workspace)
|
||||
datasetDir: "./workspace"
|
||||
|
||||
# 图像生成(管理端「AI 生成图片」):未配置 apiKey 时生成接口返回「图像生成服务未配置」
|
||||
imageGen:
|
||||
provider: dashscope # 当前唯一实现
|
||||
apiKey: "" # DashScope API Key
|
||||
model: qwen-image-3.0
|
||||
|
||||
# 训练通道(并发度 1:GPU 独占,同时仅一个 running 任务):
|
||||
# mode=subprocess 训练机与服务器同机;mode=ssh 异机(训练脚本/数据集经 ssh 通道同步)
|
||||
training:
|
||||
mode: ssh # subprocess | ssh
|
||||
ssh:
|
||||
host: "192.168.3.210" # 训练机地址(mode=ssh 必填)
|
||||
user: "root"
|
||||
port: 22
|
||||
privateKeyPath: "" # 私钥路径与 password 二选一
|
||||
password: "123"
|
||||
workdir: /opt/pheasant_data # 训练机工作目录(train_server.py / yolov8n.pt 所在)
|
||||
venvPython: /opt/pheasant_data/venv/bin/python
|
||||
datasetDir: datasets # 训练机数据集根目录(相对 workdir,yolo/<name> 为子目录)
|
||||
timeoutMinutes: 600 # 训练超时判死
|
||||
|
||||
# RF-DETR 预标注服务(管理端标注工作台):需从服务器可达,未配置时预标注接口返回「标注服务未配置」
|
||||
localAi:
|
||||
baseUrl: "http://192.168.3.210:18080" # 如 http://127.0.0.1:18080(local-ai 兼容 /v1/detection)
|
||||
model: rfdetr-xlarge
|
||||
threshold: 0.08 # 候选置信度阈值(宁多勿漏)
|
||||
confConfirmed: 0.2 # 高于此视为确认(class 0),否则疑似(class 1)
|
||||
overlapThreshold: 0.3 # 重叠去重阈值(交叠/较小框面积,NMS 风格;RF-DETR 同目标常输出一大一小两框,minIoU 比 IoU 更能命中)
|
||||
inputSize: 700 # 提交前整图等比缩放最长边
|
||||
classNames: [pheasant, suspect] # 标注类别名(写入 data.yaml,随 result.json 存模型 labels)
|
||||
|
||||
labelTask:
|
||||
# 预标注并发度(逐张调 RF-DETR,缺失或非法时回退默认值)
|
||||
poolSize: 4
|
||||
|
||||
# 套餐定价(静态配置,改价 = 改本节点后重启服务;金额单位为分,展示名由 days 派生"N天")
|
||||
plans:
|
||||
- id: day
|
||||
|
||||
Binary file not shown.
@@ -1,4 +1,4 @@
|
||||
# 视野后端单机部署:前后端一体单端口,./data 挂载持久化(容器重建不丢授权/订单数据)
|
||||
# 视野后端单机部署:前后端一体单端口,./data(数据库)与 ./workspace(APK 上传)挂载持久化(容器重建不丢数据)
|
||||
name: redfuture-app
|
||||
|
||||
networks:
|
||||
@@ -17,6 +17,7 @@ services:
|
||||
- "18080:18080"
|
||||
volumes:
|
||||
- ./data:/app/data
|
||||
- ./workspace:/app/workspace
|
||||
environment:
|
||||
TZ: Asia/Shanghai
|
||||
networks:
|
||||
|
||||
+1
-1
@@ -7,6 +7,7 @@ require (
|
||||
github.com/gogf/gf/v2 v2.10.2
|
||||
github.com/smartwalle/alipay/v3 v3.2.31
|
||||
github.com/wechatpay-apiv3/wechatpay-go v0.2.21
|
||||
golang.org/x/crypto v0.55.0
|
||||
)
|
||||
|
||||
require (
|
||||
@@ -39,7 +40,6 @@ require (
|
||||
go.opentelemetry.io/otel/metric v1.38.0 // indirect
|
||||
go.opentelemetry.io/otel/sdk v1.38.0 // indirect
|
||||
go.opentelemetry.io/otel/trace v1.38.0 // indirect
|
||||
golang.org/x/crypto v0.55.0 // indirect
|
||||
golang.org/x/net v0.57.0 // indirect
|
||||
golang.org/x/sys v0.47.0 // indirect
|
||||
golang.org/x/text v0.41.0 // indirect
|
||||
|
||||
+2
-6
@@ -98,18 +98,14 @@ go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto=
|
||||
go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE=
|
||||
golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M=
|
||||
golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis=
|
||||
golang.org/x/net v0.40.0 h1:79Xs7wF06Gbdcg4kdCCIQArK11Z1hr5POQ6+fIYHNuY=
|
||||
golang.org/x/net v0.40.0/go.mod h1:y0hY0exeL2Pku80/zKK7tpntoX23cqL3Oa6njdgRtds=
|
||||
golang.org/x/net v0.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE=
|
||||
golang.org/x/net v0.57.0/go.mod h1:KpXc8iv+r3XplLAG/f7Jsf9RPszJzdR0f58q9vGOuEU=
|
||||
golang.org/x/sys v0.0.0-20220811171246-fbc7d0a398ab/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.35.0 h1:vz1N37gP5bs89s7He8XuIYXpyY0+QlsKmzipCbUtyxI=
|
||||
golang.org/x/sys v0.35.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k=
|
||||
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
|
||||
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||
golang.org/x/text v0.25.0 h1:qVyWApTSYLk/drJRO5mDlNYskwQznZmkpV2c8q9zls4=
|
||||
golang.org/x/text v0.25.0/go.mod h1:WEdwpYrmk1qmdHvhkSTNPm3app7v4rsT8F2UD6+VHIA=
|
||||
golang.org/x/term v0.45.0 h1:NwWyBmoJCbfTHpxrWoZ9C6/VxOf7ic219I8xZZFdrf0=
|
||||
golang.org/x/term v0.45.0/go.mod h1:9aqxs0blBcrm/n0L9QW0aRVD+ktan8ssZromtqJC43w=
|
||||
golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8=
|
||||
golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M=
|
||||
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
|
||||
|
||||
@@ -0,0 +1,201 @@
|
||||
<!DOCTYPE html>
|
||||
<html lang="zh-CN">
|
||||
<head>
|
||||
<meta charset="UTF-8">
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0, viewport-fit=cover">
|
||||
<meta name="format-detection" content="telephone=no">
|
||||
<meta name="description" content="视野 App Android 版下载">
|
||||
<title>视野 - App 下载</title>
|
||||
<style>
|
||||
* { margin: 0; padding: 0; box-sizing: border-box; -webkit-tap-highlight-color: transparent; }
|
||||
html, body { height: 100%; }
|
||||
body {
|
||||
font-family: -apple-system, BlinkMacSystemFont, "PingFang SC", "Helvetica Neue", "Microsoft YaHei", Arial, sans-serif;
|
||||
background: #f5f6f7; color: #1f2329; min-height: 100vh;
|
||||
display: flex; justify-content: center;
|
||||
}
|
||||
.wrap { width: 100%; max-width: 420px; padding: 24px 24px 48px; }
|
||||
|
||||
/* ---------- 下载视图(非微信) ---------- */
|
||||
.download-view { text-align: center; padding-top: 56px; }
|
||||
.logo {
|
||||
width: 84px; height: 84px; margin: 0 auto 18px; border-radius: 20px;
|
||||
background: linear-gradient(145deg, #34c759, #1a9e4b);
|
||||
display: flex; align-items: center; justify-content: center;
|
||||
box-shadow: 0 10px 24px rgba(26, 158, 75, .28);
|
||||
}
|
||||
.logo svg { width: 46px; height: 46px; }
|
||||
.app-name { font-size: 26px; font-weight: 700; letter-spacing: 2px; }
|
||||
.app-slogan { margin-top: 6px; font-size: 14px; color: #8a9199; }
|
||||
.version-tip { margin-top: 4px; font-size: 12px; color: #b0b6bd; }
|
||||
|
||||
.btn-download {
|
||||
display: flex; align-items: center; justify-content: center; gap: 10px;
|
||||
margin: 40px auto 0; width: 100%; max-width: 300px; height: 54px;
|
||||
background: #07c160; border-radius: 27px; color: #fff; font-size: 18px; font-weight: 600;
|
||||
text-decoration: none; box-shadow: 0 8px 20px rgba(7, 193, 96, .32);
|
||||
transition: transform .1s ease;
|
||||
}
|
||||
.btn-download:active { transform: scale(.97); }
|
||||
.btn-download svg { width: 22px; height: 22px; flex: none; }
|
||||
.hint { margin-top: 18px; font-size: 12px; line-height: 1.9; color: #9aa1a9; }
|
||||
|
||||
/* ---------- 微信引导视图 ---------- */
|
||||
.wechat-view { text-align: center; padding-top: 24px; }
|
||||
.wechat-title { font-size: 20px; font-weight: 700; }
|
||||
.wechat-sub { margin-top: 8px; font-size: 14px; color: #8a9199; }
|
||||
.steps { margin: 28px auto 0; max-width: 300px; text-align: left; }
|
||||
.step { display: flex; align-items: flex-start; gap: 12px; margin-bottom: 18px; }
|
||||
.step-num {
|
||||
flex: none; width: 26px; height: 26px; border-radius: 50%;
|
||||
background: #07c160; color: #fff; font-size: 14px; font-weight: 600;
|
||||
display: flex; align-items: center; justify-content: center; margin-top: 1px;
|
||||
}
|
||||
.step-text { font-size: 15px; line-height: 1.7; color: #333; }
|
||||
.step-text b { color: #07c160; }
|
||||
|
||||
.btn-copy {
|
||||
display: inline-flex; align-items: center; gap: 8px;
|
||||
margin-top: 28px; height: 46px; padding: 0 34px;
|
||||
border: 1.5px solid #07c160; border-radius: 23px;
|
||||
background: #fff; color: #07c160; font-size: 16px; font-weight: 600;
|
||||
cursor: pointer;
|
||||
}
|
||||
.btn-copy:active { background: #f0fdf4; }
|
||||
.btn-copy svg { width: 18px; height: 18px; }
|
||||
.copy-ok { margin-top: 12px; font-size: 13px; color: #07c160; visibility: hidden; }
|
||||
.copy-ok.show { visibility: visible; }
|
||||
|
||||
.foot { margin-top: 40px; font-size: 12px; color: #b0b6bd; text-align: center; line-height: 1.8; }
|
||||
</style>
|
||||
</head>
|
||||
<body>
|
||||
<div class="wrap">
|
||||
|
||||
<!-- 非微信:下载视图 -->
|
||||
<div id="downloadView" class="download-view" hidden>
|
||||
<div class="logo">
|
||||
<svg viewBox="0 0 24 24" fill="none" aria-hidden="true">
|
||||
<path d="M12 3c2.2 0 4 1.8 4 4 0 1.5-.8 2.8-2 3.5V13c0 1.1.9 2 2 2h1v-1.5c0-1.1.9-2 2-2 1.7 0 3 1.3 3 3v3.5c0 1.7-1.3 3-3 3h-2c-.6 0-1-.4-1-1v-3h-1.2c-.8 1.2-2.2 2-3.8 2H9.5c-.6 0-1-.4-1-1v-3.2C6 17.3 4 15.4 4 13V8c0-2.8 2.2-5 5-5h3zm0 2H9c-1.7 0-3 1.3-3 3v1c1.7 0 3 1.3 3 3h.5V8c0-.6.4-1 1-1s1 .4 1 1v1c0 .4.2.7.5.9v-2.9C12 6 12 5.5 12 5zm1 4v.4c0 .3.1.7.3 1.1L15 13l1.7 1.1h.3v-.6c0-.3-.1-.6-.3-.9l-2-2.6V9h-1.7zm4.5 5.5c.8 0 1.5.7 1.5 1.5v3.5c0 .8-.7 1.5-1.5 1.5H18v-3h-1v3h-2c-.3 0-.5-.2-.5-.5v-5h-1v-1c0-.6.4-1 1-1h2z" fill="#fff"/>
|
||||
</svg>
|
||||
</div>
|
||||
<div class="app-name">视野</div>
|
||||
<div class="app-slogan">动物实时识别</div>
|
||||
<div class="version-tip">Android 版</div>
|
||||
<a class="btn-download" href="http://observer.redpowerfuture.com/download/observer-latest.apk">
|
||||
<svg viewBox="0 0 24 24" fill="none" aria-hidden="true">
|
||||
<path d="M12 3a1 1 0 0 1 1 1v9.6l3.3-3.3a1 1 0 1 1 1.4 1.4l-5 5a1 1 0 0 1-1.4 0l-5-5a1 1 0 1 1 1.4-1.4l3.3 3.3V4a1 1 0 0 1 1-1zM5 20a1 1 0 0 1 1-1h12a1 1 0 1 1 0 2H6a1 1 0 0 1-1-1z" fill="#fff"/>
|
||||
</svg>
|
||||
下载 Android 版 App
|
||||
</a>
|
||||
<div class="hint">安装时如提示,请允许「安装未知应用」权限<br>仅支持 Android 系统,请勿在 iPhone 上安装</div>
|
||||
</div>
|
||||
|
||||
<!-- 微信:引导视图 -->
|
||||
<div id="wechatView" class="wechat-view" hidden>
|
||||
<div class="wechat-title">无法在微信内直接下载</div>
|
||||
<div class="wechat-sub">请按以下步骤,使用手机浏览器打开本页</div>
|
||||
|
||||
<svg class="phone-svg" width="240" height="300" viewBox="0 0 240 300" fill="none" xmlns="http://www.w3.org/2000/svg" aria-hidden="true">
|
||||
<!-- 手机外框 -->
|
||||
<rect x="60" y="12" width="120" height="236" rx="20" fill="#2b3138"/>
|
||||
<rect x="66" y="18" width="108" height="224" rx="16" fill="#ffffff"/>
|
||||
<!-- 状态栏 -->
|
||||
<text x="74" y="32" font-size="8" fill="#1f2329" font-family="Arial">9:41</text>
|
||||
<rect x="152" y="27" width="10" height="6" rx="1.5" fill="#1f2329"/>
|
||||
<rect x="144" y="27" width="6" height="6" rx="1.5" fill="#1f2329"/>
|
||||
<rect x="136" y="27" width="6" height="6" rx="1.5" fill="#1f2329"/>
|
||||
<!-- 浏览器地址栏 -->
|
||||
<rect x="70" y="38" width="100" height="14" rx="7" fill="#eef0f2"/>
|
||||
<text x="82" y="47" font-size="6.5" fill="#8a9199" font-family="Arial">observer.redpowerfuture.com</text>
|
||||
<!-- 右上角 ··· 按钮 -->
|
||||
<circle cx="166" cy="45" r="6.5" fill="#f2f3f5" stroke="#d8dce0" stroke-width="0.8"/>
|
||||
<circle cx="163.4" cy="45" r="0.9" fill="#1f2329"/>
|
||||
<circle cx="166" cy="45" r="0.9" fill="#1f2329"/>
|
||||
<circle cx="168.6" cy="45" r="0.9" fill="#1f2329"/>
|
||||
<!-- 页面内容占位 -->
|
||||
<rect x="74" y="62" width="92" height="10" rx="5" fill="#e8eaed"/>
|
||||
<rect x="74" y="78" width="70" height="10" rx="5" fill="#e8eaed"/>
|
||||
<rect x="74" y="94" width="92" height="7" rx="3.5" fill="#f2f3f5"/>
|
||||
<rect x="74" y="106" width="92" height="7" rx="3.5" fill="#f2f3f5"/>
|
||||
<rect x="74" y="118" width="80" height="7" rx="3.5" fill="#f2f3f5"/>
|
||||
<rect x="74" y="140" width="92" height="46" rx="10" fill="#eef0f2"/>
|
||||
<rect x="74" y="196" width="92" height="7" rx="3.5" fill="#f2f3f5"/>
|
||||
<rect x="74" y="208" width="60" height="7" rx="3.5" fill="#f2f3f5"/>
|
||||
<!-- 弹出菜单气泡 -->
|
||||
<rect x="128" y="62" width="78" height="26" rx="9" fill="#ffffff" stroke="#07c160" stroke-width="1.2"/>
|
||||
<circle cx="140" cy="75" r="4" fill="#07c160" opacity="0.25"/>
|
||||
<text x="148" y="78" font-size="9.5" fill="#1f2329" font-family="Arial">在浏览器打开</text>
|
||||
<!-- 箭头:··· → 菜单 -->
|
||||
<path d="M166 51.5 C166 58 158 61 149 62.5" stroke="#07c160" stroke-width="2.2" stroke-linecap="round" stroke-dasharray="1 4" fill="none"/>
|
||||
<path d="M149.5 60.4 L145.8 63 L148.2 58.4" stroke="#07c160" stroke-width="2.2" stroke-linecap="round" stroke-linejoin="round" fill="#07c160"/>
|
||||
<!-- 底部提示条 -->
|
||||
<rect x="66" y="228" width="108" height="6" rx="3" fill="#f2f3f5"/>
|
||||
</svg>
|
||||
|
||||
<div class="steps">
|
||||
<div class="step">
|
||||
<div class="step-num">1</div>
|
||||
<div class="step-text">点击右上角 <b>···</b> 按钮</div>
|
||||
</div>
|
||||
<div class="step">
|
||||
<div class="step-num">2</div>
|
||||
<div class="step-text">选择 <b>在浏览器打开</b></div>
|
||||
</div>
|
||||
<div class="step">
|
||||
<div class="step-num">3</div>
|
||||
<div class="step-text">在浏览器中点击 <b>下载</b> 按钮即可安装</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<button id="copyBtn" class="btn-copy" type="button">
|
||||
<svg viewBox="0 0 24 24" fill="none" aria-hidden="true">
|
||||
<path d="M8 6V5a2 2 0 0 1 2-2h9a2 2 0 0 1 2 2v9a2 2 0 0 1-2 2h-1M5 8h9a2 2 0 0 1 2 2v9a2 2 0 0 1-2 2H5a2 2 0 0 1-2-2v-9a2 2 0 0 1 2-2z" stroke="#07c160" stroke-width="2" stroke-linecap="round" stroke-linejoin="round"/>
|
||||
</svg>
|
||||
复制链接
|
||||
</button>
|
||||
<div id="copyOk" class="copy-ok">已复制,请到手机浏览器中粘贴打开</div>
|
||||
</div>
|
||||
|
||||
<div class="foot">视野 · 动物实时识别</div>
|
||||
</div>
|
||||
|
||||
<script>
|
||||
(function () {
|
||||
var isWechat = /MicroMessenger/i.test(navigator.userAgent || '');
|
||||
document.getElementById('downloadView').hidden = isWechat;
|
||||
document.getElementById('wechatView').hidden = !isWechat;
|
||||
|
||||
var copyBtn = document.getElementById('copyBtn');
|
||||
var copyOk = document.getElementById('copyOk');
|
||||
function showCopied() {
|
||||
copyOk.classList.add('show');
|
||||
setTimeout(function () { copyOk.classList.remove('show'); }, 2600);
|
||||
}
|
||||
copyBtn.addEventListener('click', function () {
|
||||
var url = location.href;
|
||||
if (navigator.clipboard && navigator.clipboard.writeText) {
|
||||
navigator.clipboard.writeText(url).then(showCopied, function () {
|
||||
legacyCopy(url);
|
||||
});
|
||||
} else {
|
||||
legacyCopy(url);
|
||||
}
|
||||
});
|
||||
function legacyCopy(text) {
|
||||
var ta = document.createElement('textarea');
|
||||
ta.value = text;
|
||||
ta.style.position = 'fixed';
|
||||
ta.style.opacity = '0';
|
||||
document.body.appendChild(ta);
|
||||
ta.select();
|
||||
try {
|
||||
document.execCommand('copy');
|
||||
showCopied();
|
||||
} catch (e) { /* 剪贴板不可用,提示手动复制 */ }
|
||||
document.body.removeChild(ta);
|
||||
}
|
||||
})();
|
||||
</script>
|
||||
</body>
|
||||
</html>
|
||||
+113
-3
@@ -9,23 +9,35 @@ import (
|
||||
"github.com/gogf/gf/v2/util/gconv"
|
||||
|
||||
"observer-server/biz/controller"
|
||||
"observer-server/biz/service"
|
||||
"observer-server/common"
|
||||
)
|
||||
|
||||
func main() {
|
||||
ctx := context.Background()
|
||||
initDatabase(ctx)
|
||||
// 封面存量迁移:历史固定命名 cover* → UUID jpg(幂等,新库空跑)
|
||||
if err := service.Dataset.MigrateLegacyCovers(ctx); err != nil {
|
||||
g.Log().Errorf(ctx, "封面存量迁移失败: %+v", err)
|
||||
}
|
||||
|
||||
s := g.Server()
|
||||
// Android APK 下载静态托管:app.apkDir 目录下固定文件 observer-latest.apk,
|
||||
// URL 固定 /download/observer-latest.apk(绕过统一响应包装,纯二进制流);
|
||||
// 模型热更新文件同根 /download/models/<数据集名>/latest.tflite(datasetDir 下 models 目录)
|
||||
if err := os.MkdirAll(common.ApkDir(ctx), 0o755); err != nil {
|
||||
g.Log().Fatalf(ctx, "创建 APK 目录失败: %+v", err)
|
||||
}
|
||||
s.AddStaticPath("/download", common.ApkDir(ctx))
|
||||
// 客户端接口组:统一响应包装 {"code":0,"message":"ok","data":...}
|
||||
// 账号组公开(注册/登录),业务组(订单/授权)需登录态(Authorization: Bearer token)
|
||||
// 账号组公开(注册/登录),业务组(订单/授权/模型目录)需登录态(Authorization: Bearer token)
|
||||
s.Group("/api/v1", func(group *ghttp.RouterGroup) {
|
||||
group.Middleware(common.UnifiedResponse)
|
||||
common.BindController(group, controller.Auth)
|
||||
common.BindController(group, controller.Auth, controller.AppVersion)
|
||||
})
|
||||
s.Group("/api/v1", func(group *ghttp.RouterGroup) {
|
||||
group.Middleware(common.UnifiedResponse, common.AuthRequired)
|
||||
common.BindController(group, controller.Order, controller.License)
|
||||
common.BindController(group, controller.Order, controller.License, controller.ModelCatalog)
|
||||
})
|
||||
// 支付回调组:不做统一包装,按渠道应答格式直接返回(微信 SUCCESS/FAIL JSON、支付宝 success/failure 文本)
|
||||
s.Group("/api/v1/payment", func(group *ghttp.RouterGroup) {
|
||||
@@ -36,6 +48,12 @@ func main() {
|
||||
group.Middleware(common.UnifiedResponse, common.AdminAuth)
|
||||
common.BindController(group, controller.Admin)
|
||||
})
|
||||
// APK 下载引导页(h5/ 源码目录):微信内打开提示用手机浏览器,非微信直显下载按钮
|
||||
if stat, err := os.Stat("./h5"); err == nil && stat.IsDir() {
|
||||
s.AddStaticPath("/download-page", "./h5")
|
||||
} else {
|
||||
g.Log().Warningf(ctx, "h5 目录不存在,跳过下载引导页托管")
|
||||
}
|
||||
// 管理端静态页面托管:构建产物输出到 admin_dist(server_admin/ 构建),SPA history 路由回退 index.html;
|
||||
// 目录不存在(尚未构建前端)时跳过,不影响 API 启动
|
||||
if stat, err := os.Stat("./admin_dist"); err == nil && stat.IsDir() {
|
||||
@@ -44,6 +62,8 @@ func main() {
|
||||
} else {
|
||||
g.Log().Warningf(ctx, "admin_dist 不存在,跳过管理端静态托管(cd server_admin && npm run build)")
|
||||
}
|
||||
// 后台协程:训练进度轮询 + 孤儿预标注任务恢复
|
||||
service.Training.StartBackgroundJobs(ctx)
|
||||
s.Run()
|
||||
}
|
||||
|
||||
@@ -92,4 +112,94 @@ func initDatabase(ctx context.Context) {
|
||||
}
|
||||
g.Log().Infof(ctx, "数据库初始化完成(version=3)")
|
||||
}
|
||||
if version < 8 {
|
||||
// v8:dataset 加训练配置/展示列 + label_task 加 filenames(多选批量标注)。
|
||||
// PRAGMA table_info 逐列检测缺失才 ADD COLUMN,全新库建表自带全列直接跳过
|
||||
addCols := []struct{ table, col, ddl string }{
|
||||
{"dataset", "cover", "TEXT"},
|
||||
{"dataset", "description", "TEXT"},
|
||||
{"dataset", "ai_endpoint", "TEXT"},
|
||||
{"dataset", "ai_model", "TEXT"},
|
||||
{"dataset", "train_host", "TEXT"},
|
||||
{"dataset", "train_user", "TEXT"},
|
||||
{"dataset", "train_password", "TEXT"},
|
||||
{"dataset", "train_key", "TEXT"},
|
||||
{"label_task", "filenames", "TEXT"},
|
||||
}
|
||||
for _, c := range addCols {
|
||||
cols, err := g.DB().GetAll(ctx, "PRAGMA table_info("+c.table+")")
|
||||
if err != nil {
|
||||
g.Log().Fatalf(ctx, "读取 %s 表结构失败: %+v", c.table, err)
|
||||
}
|
||||
exists := false
|
||||
for _, col := range cols {
|
||||
if gconv.String(col["name"]) == c.col {
|
||||
exists = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if exists {
|
||||
continue
|
||||
}
|
||||
if _, err := g.DB().Exec(ctx, "ALTER TABLE "+c.table+" ADD COLUMN "+c.col+" "+c.ddl); err != nil {
|
||||
g.Log().Fatalf(ctx, "给 %s 加列 %s 失败: %+v", c.table, c.col, err)
|
||||
}
|
||||
}
|
||||
if _, err := g.DB().Exec(ctx, "PRAGMA user_version = 8"); err != nil {
|
||||
g.Log().Fatalf(ctx, "库版本写入失败: %+v", err)
|
||||
}
|
||||
g.Log().Infof(ctx, "数据库初始化完成(version=8)")
|
||||
}
|
||||
if version < 10 {
|
||||
// v10:标注流程简化,撤销候选确认两阶段 —— dataset_image 删 candidates_json 列
|
||||
//(存量候选数据为空直接删;全新库建表已无此列、PRAGMA table_info 检测后跳过)
|
||||
cols, err := g.DB().GetAll(ctx, "PRAGMA table_info(dataset_image)")
|
||||
if err != nil {
|
||||
g.Log().Fatalf(ctx, "读取 dataset_image 表结构失败: %+v", err)
|
||||
}
|
||||
hasCandidates := false
|
||||
for _, col := range cols {
|
||||
if gconv.String(col["name"]) == "candidates_json" {
|
||||
hasCandidates = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if hasCandidates {
|
||||
if _, err := g.DB().Exec(ctx, "ALTER TABLE dataset_image DROP COLUMN candidates_json"); err != nil {
|
||||
g.Log().Fatalf(ctx, "删除 candidates_json 列失败: %+v", err)
|
||||
}
|
||||
}
|
||||
if _, err := g.DB().Exec(ctx, "PRAGMA user_version = 10"); err != nil {
|
||||
g.Log().Fatalf(ctx, "库版本写入失败: %+v", err)
|
||||
}
|
||||
g.Log().Infof(ctx, "数据库初始化完成(version=10)")
|
||||
}
|
||||
if version < 11 {
|
||||
// v11:移除模型存档回退机制 —— model_version 删 model_file 列
|
||||
//(模型文件不落表:发布即写 latest.tflite,客户端固定下载;存量库删列、新库建表已无此列自动跳过)
|
||||
cols, err := g.DB().GetAll(ctx, "PRAGMA table_info(model_version)")
|
||||
if err != nil {
|
||||
g.Log().Fatalf(ctx, "读取 model_version 表结构失败: %+v", err)
|
||||
}
|
||||
hasModelFile := false
|
||||
for _, col := range cols {
|
||||
if gconv.String(col["name"]) == "model_file" {
|
||||
hasModelFile = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if hasModelFile {
|
||||
if _, err := g.DB().Exec(ctx, "ALTER TABLE model_version DROP COLUMN model_file"); err != nil {
|
||||
g.Log().Fatalf(ctx, "删除 model_file 列失败: %+v", err)
|
||||
}
|
||||
}
|
||||
if _, err := g.DB().Exec(ctx, "PRAGMA user_version = 11"); err != nil {
|
||||
g.Log().Fatalf(ctx, "库版本写入失败: %+v", err)
|
||||
}
|
||||
g.Log().Infof(ctx, "数据库初始化完成(version=11)")
|
||||
}
|
||||
// 死表清理:app_config 全局训练配置表已撤销(配置走 config.yml),存量库残留表启动即删
|
||||
if _, err := g.DB().Exec(ctx, "DROP TABLE IF EXISTS app_config"); err != nil {
|
||||
g.Log().Fatalf(ctx, "删除残留 app_config 表失败: %+v", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,313 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
"""训练任务脚本(Go 后端 --task-json 驱动,产物契约见 server/common/training_runner.go)。
|
||||
|
||||
用法:
|
||||
python train_server.py --task-json tasks/<taskId>.json
|
||||
|
||||
任务参数(Go 侧写入,字段相对训练机 workdir):
|
||||
workdir 训练机工作目录(脚本 / yolov8n.pt / venv 所在),启动即 chdir
|
||||
yolo 训练集目录(含 dataset.yaml),相对 workdir
|
||||
imgsz 训练/导出分辨率(默认 704,与端侧推理对齐)
|
||||
epochs 训练轮数
|
||||
batch 批大小
|
||||
device GPU 编号或 cpu
|
||||
project 训练输出目录(相对 workdir,末尾自动拼 name)
|
||||
log_file 每 epoch 一行 JSON 的进度日志(相对 workdir)
|
||||
result_file 结束结果 JSON(相对 workdir)
|
||||
artifact_zip 产物打包(best.pt + results.csv + 曲线,相对 workdir)
|
||||
|
||||
产物契约:
|
||||
log_file {"epoch":1,"total":150,"metrics":{"metrics/mAP50(B)":0.87,...}}
|
||||
result_file {"metrics":{...},"names":["pheasant","suspect"],"best_tflite":"runs/tasks/<id>/weights/best.tflite",
|
||||
"tflite_check":{"ok":true,"reason":"","inputs":[...],"outputs":[...]}}
|
||||
result_file 存在 = 训练完成;异常时写 {"error":"..."},Go 侧据以置失败并展示原因。
|
||||
tflite_check.ok=false(产物 shape 异常)时 Go 侧置训练失败并带出 reason。
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import struct
|
||||
import sys
|
||||
import traceback
|
||||
import zipfile
|
||||
from pathlib import Path
|
||||
|
||||
# 抗尺度漂移(与人工训练基线一致):随机缩放输入 0.5~1.5x
|
||||
MULTI_SCALE = 0.5
|
||||
# 早停耐心(连续 N 轮无提升即停)
|
||||
PATIENCE = 30
|
||||
|
||||
|
||||
def on_fit_epoch_end(trainer):
|
||||
"""ultralytics 回调:每 epoch 结束写一行进度 JSON(行缓冲,进程被杀不丢行)。"""
|
||||
metrics = {}
|
||||
for k, v in (trainer.metrics or {}).items():
|
||||
if isinstance(v, (int, float)) and math.isfinite(v):
|
||||
metrics[k] = round(float(v), 5)
|
||||
if _LOG is None:
|
||||
return
|
||||
total = getattr(trainer, "epochs", None) or getattr(trainer, "total_epochs", 0) or 0
|
||||
_LOG.write(json.dumps({
|
||||
"epoch": trainer.epoch + 1,
|
||||
"total": total,
|
||||
"metrics": metrics,
|
||||
}, ensure_ascii=False) + "\n")
|
||||
|
||||
|
||||
def register_callback():
|
||||
from ultralytics.utils.callbacks import callbacks
|
||||
callbacks["on_fit_epoch_end"].append(on_fit_epoch_end)
|
||||
|
||||
|
||||
def write_result(result_file, payload):
|
||||
Path(result_file).parent.mkdir(parents=True, exist_ok=True)
|
||||
# tmp + rename 原子写,避免 Go 侧读到半截文件
|
||||
tmp = result_file + ".tmp"
|
||||
with open(tmp, "w", encoding="utf-8") as f:
|
||||
json.dump(payload, f, ensure_ascii=False, indent=2)
|
||||
os.replace(tmp, result_file)
|
||||
|
||||
|
||||
def read_names(yolo_dir):
|
||||
"""从 data.yaml 读类别名(有序列表,ultralytics 依赖 pyyaml)"""
|
||||
try:
|
||||
import yaml
|
||||
with open(Path(yolo_dir) / "dataset.yaml", encoding="utf-8") as f:
|
||||
cfg = yaml.safe_load(f)
|
||||
names = cfg.get("names") or {}
|
||||
if isinstance(names, dict):
|
||||
return [names[k] for k in sorted(names, key=lambda k: int(k))]
|
||||
return list(names)
|
||||
except Exception:
|
||||
return []
|
||||
|
||||
|
||||
def find_best_tflite(save_dir):
|
||||
"""定位导出的 tflite(ultralytics 各版本产物位置不一):
|
||||
优先 fp32 best.tflite(与人工基线一致),再 float32/int8,最后兜底全局搜。"""
|
||||
weights = save_dir / "weights"
|
||||
for name in ("best.tflite", "best_float32.tflite", "best_int8.tflite"):
|
||||
p = weights / name
|
||||
if p.exists():
|
||||
return p
|
||||
candidates = sorted(weights.rglob("*.tflite"))
|
||||
if candidates:
|
||||
return candidates[0]
|
||||
return None
|
||||
|
||||
|
||||
# TFLite flatbuffer 张量类型(tensorflow/lite/schema/schema.fbs,字段顺序敏感)
|
||||
TENSOR_TYPE = {0: "FLOAT32", 1: "FLOAT16", 2: "INT32", 3: "UINT8", 4: "INT64",
|
||||
5: "STRING", 6: "BOOL", 7: "INT16", 8: "COMPLEX64", 9: "INT8",
|
||||
10: "FLOAT64", 11: "COMPLEX128", 12: "UINT64", 13: "RESOURCE",
|
||||
14: "VARIANT", 15: "UINT32", 16: "UINT16", 17: "INT4",
|
||||
18: "BFLOAT16", 19: "FLOAT8_E4M3FN", 20: "FLOAT8_E4M3FNUZ",
|
||||
21: "FLOAT8_E5M2", 22: "FLOAT8_E5M2FNUZ"}
|
||||
|
||||
|
||||
class _FB:
|
||||
"""极简 flatbuffer 读取器(仅 Model/SubGraph/Tensor 表所需字段)"""
|
||||
|
||||
def __init__(self, data):
|
||||
self.d = data
|
||||
|
||||
def u32(self, pos): return struct.unpack_from("<I", self.d, pos)[0]
|
||||
def i32(self, pos): return struct.unpack_from("<i", self.d, pos)[0]
|
||||
def u16(self, pos): return struct.unpack_from("<H", self.d, pos)[0]
|
||||
|
||||
def field(self, t, i):
|
||||
vto = self.i32(t)
|
||||
vt = t - vto
|
||||
vs = self.u16(vt)
|
||||
off = vt + 4 + 2 * i
|
||||
if off + 2 > vt + vs:
|
||||
return None
|
||||
f = self.u16(off)
|
||||
return None if f == 0 else t + f
|
||||
|
||||
def deref(self, pos):
|
||||
return pos + self.u32(pos) if pos is not None else None
|
||||
|
||||
def vec(self, pos):
|
||||
pos = self.deref(pos)
|
||||
if pos is None:
|
||||
return None
|
||||
return pos + 4, self.u32(pos)
|
||||
|
||||
def vec_table(self, pos):
|
||||
start, n = self.vec(pos)
|
||||
return [start + i * 4 + self.u32(start + i * 4) for i in range(n)]
|
||||
|
||||
def string(self, pos):
|
||||
pos = self.deref(pos)
|
||||
if pos is None:
|
||||
return None
|
||||
n = self.u32(pos)
|
||||
return self.d[pos + 4:pos + 4 + n].decode("utf-8", "replace")
|
||||
|
||||
def int_vec(self, pos):
|
||||
start, n = self.vec(pos)
|
||||
return [self.i32(start + 4 * i) for i in range(n)]
|
||||
|
||||
|
||||
def check_tflite(path, imgsz):
|
||||
"""TFLite 产物自检(原 inspect_tflite.py,2026-08-26 并入):
|
||||
输入恰 1 张且 4 维、元素总数 == imgsz²×3(NCHW/NHWC 皆可)、输出 ≥ 1 张且 batch 维 = 1。
|
||||
返回 {"ok","reason","inputs":[{"name","shape","type"}],"outputs":[...]}。"""
|
||||
try:
|
||||
with open(path, "rb") as f:
|
||||
fb = _FB(f.read())
|
||||
root = fb.u32(0) # Model 表
|
||||
subgraphs = fb.vec_table(fb.field(root, 2)) # Model.subgraphs [2]
|
||||
sg = subgraphs[0]
|
||||
tensors = fb.vec_table(fb.field(sg, 0)) # SubGraph.tensors [0]
|
||||
inputs = fb.int_vec(fb.field(sg, 1)) # SubGraph.inputs [1]
|
||||
outputs = fb.int_vec(fb.field(sg, 2)) # SubGraph.outputs [2]
|
||||
|
||||
def describe(i):
|
||||
t = tensors[i]
|
||||
shape = fb.int_vec(fb.field(t, 0))
|
||||
# type 有默认值 FLOAT32(=0):字段等于默认值时 flatbuffer 省略(vtable 偏移 0)
|
||||
ttype_f = fb.field(t, 1)
|
||||
ttype = fb.d[ttype_f] if ttype_f else 0
|
||||
return {"name": fb.string(fb.field(t, 3)) or "?",
|
||||
"shape": shape,
|
||||
"type": TENSOR_TYPE.get(ttype, ttype)}
|
||||
|
||||
in_desc = [describe(i) for i in inputs]
|
||||
out_desc = [describe(i) for i in outputs]
|
||||
|
||||
problems = []
|
||||
if len(in_desc) != 1:
|
||||
problems.append(f"输入张量 {len(in_desc)} 个(期望 1)")
|
||||
else:
|
||||
shape = in_desc[0]["shape"]
|
||||
n = 1
|
||||
for d in shape:
|
||||
n *= d
|
||||
if len(shape) != 4:
|
||||
problems.append(f"输入 shape {shape} 非 4 维")
|
||||
elif n != imgsz * imgsz * 3:
|
||||
problems.append(f"输入元素数 {n} ≠ imgsz²×3={imgsz * imgsz * 3}(shape {shape},imgsz={imgsz})")
|
||||
if not out_desc:
|
||||
problems.append("无输出张量")
|
||||
else:
|
||||
for o in out_desc:
|
||||
if o["shape"] and o["shape"][0] != 1:
|
||||
problems.append(f"输出 batch 维 ≠ 1: {o['shape']}")
|
||||
|
||||
return {"ok": not problems, "reason": ";".join(problems),
|
||||
"inputs": in_desc, "outputs": out_desc}
|
||||
except Exception as e:
|
||||
return {"ok": False, "reason": f"解析失败: {e}", "inputs": [], "outputs": []}
|
||||
|
||||
|
||||
def build_artifact_zip(zip_path, save_dir):
|
||||
Path(zip_path).parent.mkdir(parents=True, exist_ok=True)
|
||||
with zipfile.ZipFile(zip_path, "w", zipfile.ZIP_DEFLATED) as zf:
|
||||
for p in (save_dir / "weights" / "best.pt", save_dir / "results.csv"):
|
||||
if p.exists():
|
||||
zf.write(p, arcname=p.name)
|
||||
# 曲线/混淆矩阵等绘图产物在 save_dir 根目录
|
||||
for p in sorted(save_dir.glob("*.png")) + sorted(save_dir.glob("*.jpg")):
|
||||
zf.write(p, arcname=p.name)
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser(description="observer 训练任务脚本(Go --task-json 驱动)")
|
||||
ap.add_argument("--task-json", required=True, help="任务参数 JSON 文件路径")
|
||||
args = ap.parse_args()
|
||||
|
||||
with open(args.task_json, encoding="utf-8") as f:
|
||||
task = json.load(f)
|
||||
|
||||
global _LOG
|
||||
base = os.path.abspath(task["workdir"])
|
||||
|
||||
def task_path(key):
|
||||
"""任务参数里的路径全部相对 workdir,这里绝对化(chdir 失败也能定位)"""
|
||||
p = task[key]
|
||||
return p if os.path.isabs(p) else os.path.join(base, p)
|
||||
|
||||
try:
|
||||
os.chdir(base)
|
||||
except Exception:
|
||||
pass # 目录不可用则后续训练必然失败,错误路径仍尽力写 result.json
|
||||
|
||||
log_file = task_path("log_file")
|
||||
result_file = task_path("result_file")
|
||||
artifact_zip = task_path("artifact_zip")
|
||||
global _LOG
|
||||
_LOG = None
|
||||
try:
|
||||
Path(log_file).parent.mkdir(parents=True, exist_ok=True)
|
||||
_LOG = open(log_file, "a", encoding="utf-8", buffering=1)
|
||||
except Exception:
|
||||
pass # 日志不可写(如 workdir 缺失)不阻断错误路径,traceback 仍走 stderr
|
||||
|
||||
imgsz = int(task.get("imgsz") or 704)
|
||||
epochs = int(task.get("epochs") or 150)
|
||||
batch = int(task.get("batch") or 16)
|
||||
device = task.get("device") or "0"
|
||||
data = os.path.join(task["yolo"], "dataset.yaml")
|
||||
names = read_names(task["yolo"])
|
||||
|
||||
try:
|
||||
from ultralytics import YOLO
|
||||
register_callback()
|
||||
model = YOLO("yolov8n.pt")
|
||||
model.train(
|
||||
data=data, imgsz=imgsz, epochs=epochs,
|
||||
patience=PATIENCE, batch=batch, device=device, workers=4,
|
||||
# project 必须绝对路径,相对路径会被拼到默认 runs/detect 下造成双层嵌套
|
||||
project=task_path("project"), name="train",
|
||||
exist_ok=True, plots=True, multi_scale=MULTI_SCALE,
|
||||
)
|
||||
trainer = getattr(model, "trainer", None)
|
||||
if trainer is None:
|
||||
raise RuntimeError("训练完成但无法获取 trainer(save_dir 未知)")
|
||||
save_dir = Path(trainer.save_dir)
|
||||
|
||||
best_tflite = find_best_tflite(save_dir)
|
||||
if best_tflite is None:
|
||||
model.export(format="tflite", imgsz=imgsz)
|
||||
best_tflite = find_best_tflite(save_dir)
|
||||
if best_tflite is None:
|
||||
raise RuntimeError("tflite 导出失败:weights 目录下未找到任何 .tflite 产物")
|
||||
|
||||
# tflite 产物自检:shape 异常也写进 result.json(ok=false),Go 侧据此置失败
|
||||
tflite_check = check_tflite(str(best_tflite), imgsz)
|
||||
if not tflite_check["ok"]:
|
||||
print(f"[tflite-check] FAIL: {tflite_check['reason']}", file=sys.stderr)
|
||||
|
||||
# 结果指标取 best 轮(发布的是 best.pt),缺省回退末轮
|
||||
metrics = getattr(trainer, "best_metrics", None) or trainer.metrics
|
||||
clean_metrics = {}
|
||||
for k, v in (metrics or {}).items():
|
||||
if isinstance(v, (int, float)) and math.isfinite(v):
|
||||
clean_metrics[k] = round(float(v), 5)
|
||||
|
||||
build_artifact_zip(artifact_zip, save_dir)
|
||||
# best_tflite 相对 workdir,Go 侧按此路径拉取
|
||||
write_result(result_file, {
|
||||
"metrics": clean_metrics,
|
||||
"names": names,
|
||||
"best_tflite": os.path.relpath(best_tflite, base),
|
||||
"tflite_check": tflite_check,
|
||||
})
|
||||
except Exception:
|
||||
err = traceback.format_exc()
|
||||
if _LOG is not None:
|
||||
_LOG.write("\n" + err + "\n")
|
||||
try:
|
||||
write_result(result_file, {"error": err[-2000:]})
|
||||
except Exception:
|
||||
pass
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
+234
@@ -48,17 +48,130 @@ CREATE TABLE IF NOT EXISTS license (
|
||||
phone_num TEXT PRIMARY KEY, -- 手机号账号,一账号一授权记录,续费更新
|
||||
password TEXT NOT NULL, -- bcrypt 加盐哈希,不存明文
|
||||
expires_at TEXT, -- 未充值 NULL(NULL/过期即未授权)
|
||||
remark TEXT, -- 管理端备注(运营发卡/人工记录,客户端不可见)
|
||||
created_at TEXT NOT NULL,
|
||||
updated_at TEXT NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS app_version (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
version TEXT NOT NULL UNIQUE, -- 语义化版本号 x.y.z,客户端按数字段比较
|
||||
notes TEXT, -- 更新说明(客户端弹窗展示)
|
||||
created_at TEXT NOT NULL,
|
||||
updated_at TEXT NOT NULL
|
||||
-- 无下载地址列:APK 为固定文件 app.apkDir/observer-latest.apk,上传即覆盖,目录永远只有一个文件
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS dataset (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
name TEXT NOT NULL UNIQUE, -- 数据集名(≤50 字),同时是服务器目录名
|
||||
source TEXT NOT NULL DEFAULT 'manual', -- manual | ai
|
||||
image_count INTEGER NOT NULL DEFAULT 0, -- 图片数(冗余计数,随增删更新)
|
||||
labeled_count INTEGER NOT NULL DEFAULT 0, -- 已标注数(二期标注任务完成累计)
|
||||
status TEXT NOT NULL DEFAULT 'building',-- building | labeled | synced
|
||||
cover TEXT, -- 封面文件名(上传自动转 jpg + UUID 命名,卡片展示)
|
||||
description TEXT, -- 描述(卡片展示)
|
||||
-- 遗留列(ai_endpoint/ai_model/train_host/train_user/train_password/train_key):
|
||||
-- AI 标注/训练机 SSH 配置统一走 config.yml(localAi / training.ssh),这 6 个覆盖列新代码不再读写,
|
||||
-- 保留不迁移(存量库列不动、无数据迁移)
|
||||
ai_endpoint TEXT,
|
||||
ai_model TEXT,
|
||||
train_host TEXT,
|
||||
train_user TEXT,
|
||||
train_password TEXT,
|
||||
train_key TEXT,
|
||||
created_at TEXT NOT NULL,
|
||||
updated_at TEXT NOT NULL
|
||||
-- 图片文件在 app.datasetDir/datasets/<name>/;DB 存元数据 + 标注(dataset_image.labels_json)
|
||||
);
|
||||
-- 存量库 v8 迁移:逐列检测补列(PRAGMA table_info),新库 CREATE 自带全列
|
||||
|
||||
CREATE TABLE IF NOT EXISTS dataset_image (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
dataset_id INTEGER NOT NULL, -- → dataset.id
|
||||
filename TEXT NOT NULL, -- 唯一文件名(防重名加时间戳后缀)
|
||||
source TEXT NOT NULL DEFAULT 'manual', -- manual | ai
|
||||
prompt TEXT, -- AI 生成图记录提示词(追溯用)
|
||||
labels_json TEXT, -- 标注 JSON 数组(YOLO 归一化 xywh+类别+置信度),AI 自动标注与人工标注同存(人工可修改/清理),null/''/'[]'=未标注
|
||||
created_at TEXT NOT NULL,
|
||||
UNIQUE (dataset_id, filename)
|
||||
);
|
||||
-- 存量库迁移:EnsureColumn 逐列检测补列(v9);candidates_json 列已随 v10 删除(DROP COLUMN);
|
||||
-- model_version.model_file 列已随 v11 删除(模型无存档回退机制,只留 latest.tflite)
|
||||
|
||||
CREATE TABLE IF NOT EXISTS model_training (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
name TEXT NOT NULL, -- 任务名(默认「数据集+时间」)
|
||||
status TEXT NOT NULL DEFAULT 'running', -- running | success | failed
|
||||
dataset TEXT NOT NULL, -- 训练机数据集名(datasetDir 下子目录名)
|
||||
imgsz INTEGER NOT NULL DEFAULT 704,
|
||||
epochs INTEGER NOT NULL DEFAULT 150,
|
||||
batch INTEGER NOT NULL DEFAULT 16,
|
||||
device TEXT NOT NULL DEFAULT '0',
|
||||
current_epoch INTEGER NOT NULL DEFAULT 0,
|
||||
total_epochs INTEGER NOT NULL DEFAULT 0,
|
||||
metrics TEXT, -- JSON {p, r, map50}(成功后写)
|
||||
log_tail TEXT, -- 日志尾部(截断 N KB,轮询更新)
|
||||
pid INTEGER, -- 训练进程 pid(取消/存活探测用)
|
||||
error TEXT, -- 失败原因
|
||||
started_at TEXT NOT NULL,
|
||||
finished_at TEXT,
|
||||
created_at TEXT NOT NULL,
|
||||
UNIQUE (dataset, imgsz, epochs, batch, started_at) -- 防止重复提交同参任务(宽松防呆)
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS model_version (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
dataset_id INTEGER NOT NULL, -- → dataset.id,**每个数据集独立模型版本序列**
|
||||
version TEXT NOT NULL, -- m1.0.0 递增(每次发布 patch+1,同数据集内唯一)
|
||||
training_id INTEGER, -- 来源训练任务 → model_training.id
|
||||
artifact_file TEXT, -- 归档 zip(best.pt + results 曲线),可选
|
||||
metrics TEXT, -- JSON,与来源任务一致
|
||||
labels TEXT NOT NULL, -- JSON 类别名数组(随模型下发,App 合并/展示用)
|
||||
sha256 TEXT NOT NULL, -- tflite 文件校验
|
||||
size_bytes INTEGER NOT NULL,
|
||||
is_latest INTEGER NOT NULL DEFAULT 0, -- 1=该数据集当前生效(客户端拉取对象),每数据集至多一条
|
||||
notes TEXT,
|
||||
created_at TEXT NOT NULL,
|
||||
UNIQUE (dataset_id, version)
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS label_task (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
dataset_id INTEGER NOT NULL, -- → dataset.id
|
||||
filenames TEXT, -- JSON 选中图片列表(多选批量标注);NULL = 全量扫描
|
||||
status TEXT NOT NULL DEFAULT 'running', -- running | done
|
||||
total INTEGER NOT NULL DEFAULT 0, -- 待标注图片数
|
||||
done INTEGER NOT NULL DEFAULT 0, -- 已标注数
|
||||
boxes_file TEXT, -- 遗留列:候选框 JSON 路径(boxes.json 已废弃,新代码不读写)
|
||||
error TEXT, -- 失败原因(检测中途失败/服务重启中断)
|
||||
created_at TEXT NOT NULL,
|
||||
finished_at TEXT
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_payment_order_phone ON payment_order(phone_num);
|
||||
CREATE INDEX IF NOT EXISTS idx_dataset_image_dataset ON dataset_image(dataset_id);
|
||||
CREATE INDEX IF NOT EXISTS idx_model_training_status ON model_training(status);
|
||||
```
|
||||
|
||||
存量库迁移以 `PRAGMA user_version` 版本化标记,禁止重复执行。迁移记录:
|
||||
- v1 = 建表(各 dao init)+ 种子套餐(v1 已移除)
|
||||
- v2 = 删除 `license.plan_id` 列(该字段无业务语义,只留到期时间;`PRAGMA table_info` 检测列存在才 `ALTER TABLE ... DROP COLUMN`,新库与已迁移库跳过)
|
||||
- v3 = `DROP TABLE IF EXISTS plan`(套餐改配置后清残留表,订单快照 `payment_order.plan_id` 不受影响)
|
||||
- v4 = `license` 加 `remark` 列(管理端备注;`PRAGMA table_info` 检测列缺失才 `ALTER TABLE ... ADD COLUMN remark TEXT`,新库直接建表跳过)
|
||||
- v5 = `app_version` 新表(版本管理;`dao` init `CREATE TABLE IF NOT EXISTS` 自动建,新库/存量库均无需 `user_version` 迁移,此处记录 DDL 变更)
|
||||
- v6 = `app_version` 删 `url` 列(下载地址改为固定文件 `app.apkDir`/`observer-latest.apk`,表内不再记录;`PRAGMA table_info` 检测列存在才 `ALTER TABLE ... DROP COLUMN url`,新库建表已无此列直接跳过)
|
||||
- v7 = `dataset` / `dataset_image` / `model_training` / `model_version` / `label_task` 新表(模型训练体系;各 `dao` init `CREATE TABLE IF NOT EXISTS` 自动建,此处记录 DDL 变更)
|
||||
- v8 = `dataset` 加 9 列(`cover`/`description`/`ai_endpoint`/`ai_model`/`train_host`/`train_user`/`train_password`/`train_key`)+ `label_task` 加 `filenames` 列(多选批量标注;`PRAGMA table_info` 逐列检测缺失才 `ALTER TABLE ... ADD COLUMN`,新库建表自带跳过)——**其中 6 列(ai/train 覆盖字段)为遗留列**:AI 标注/训练机 SSH 配置统一走 `config.yml`(`localAi` / `training.ssh`),新代码不读写,存量库保留不迁移
|
||||
- v9 = 标注存储从文件迁移入库:`dataset_image` 加 `labels_json`/`candidates_json` 两列(`common.EnsureColumn` 迁移),启动时把历史 `labels/<数据集>/*.txt` 解析入 `labels_json`(幂等:仅未迁移行处理),`boxes.json` 废弃;**历史 labels/ 目录与迁移代码已删除(2026-08-26:数据全部入表后无保留价值)**
|
||||
- v10 = 标注流程简化(撤销候选确认两阶段):`ALTER TABLE dataset_image DROP COLUMN candidates_json`(`PRAGMA table_info` 检测列存在才 DROP,新库建表已无此列直接跳过;存量候选数据为空直接删)——AI 预标注结果直写 `labels_json`,人工可修改/清理全部标注框,`label_task.boxes_file` 与 `dataset` 的 ai/train 六列同为遗留列保留不迁移
|
||||
|
||||
## 全局训练配置(config.yml 直读)
|
||||
|
||||
**决策(2026-08-26)**:AI 标注端点(`localAi`)与训练机 SSH 凭据(`training.ssh`)是全局配置、与数据集无关——**不设独立存储**(曾尝试 `app_config` KV 表 + 管理端「训练配置」入口,2026-08-26 撤销):标注与训练直接读 `config.yml` 的 `localAi` / `training.ssh` 节点,改配置需重启服务。数据集表 6 个覆盖列(ai_endpoint 等)为遗留列,代码不读写。
|
||||
|
||||
- AI 客户端:`common.LocalAiClient(ctx)` 直读 `localAi.baseUrl`/`model`,未配置返回 nil(预标注接口报「标注服务未配置」)
|
||||
- SSH 凭据:`common/training_runner.go` 的 sshRunner 直读 `training.ssh.host`/`user`/`port`/`privateKeyPath`/`password`,未配置 host 报「training.ssh 未配置 host」
|
||||
|
||||
## 账号体系(注册/登录)
|
||||
|
||||
@@ -152,6 +265,9 @@ CREATE INDEX IF NOT EXISTS idx_payment_order_phone ON payment_order(phone_num);
|
||||
| GET | /licenses | 账号列表:`phoneNum` 筛选 + 分页(含未充值账号) |
|
||||
| POST | /licenses/grant | 手动授权 `{phoneNum, planId}`:与支付回调同语义(自然日叠加、单写者串行事务),提交后清授权缓存 |
|
||||
| POST | /licenses/revoke | 撤销授权 `{phoneNum}`:清空 `expires_at` **保留账号行**(不删密码),清缓存,客户端下次查询即 inactive |
|
||||
| GET | /app-versions | 版本记录列表:分页(size 上限 100),按下发时间倒序 |
|
||||
| POST | /app-versions | 下发新版本(multipart/form-data:`notes` + `file` APK):**版本号从文件名识别**,文件名须为 `observer-x.y.z.apk`(如 `observer-1.0.1.apk`,正则 `^observer-(\d+\.\d+\.\d+)\.apk$`,格式不符拒绝)、`notes` ≤500、仅接受 `.apk` 文件;APK 覆盖保存 `app.apkDir`/`observer-latest.apk`(目录永远只有一个文件) |
|
||||
| POST | /app-versions/delete | 删除版本 `{id}`(不存在报错):删**最新版本**时联动删除 APK 文件(客户端 update 返回空不再提示、下载 404);删历史版本仅删记录、不动文件 |
|
||||
|
||||
**结构决策**:管理端是跨表业务面,controller 聚合在 `biz/controller/admin.go`(避免同一 controller 挂客户端组 + 管理组时 `group.Bind` 重复注册路由),service 按「跨表业务流程归入所属表文件」归入各表文件(订单列表 → `service/order.go`,授权 grant/revoke/列表 → `service/license.go`),dto 聚合在 `biz/model/dto/admin.go`。套餐已配置化(`common/plans.go` 读取 config.yml,仅客户端登录组 `GET /plans` 使用),无独立分层文件,管理端不管理套餐。
|
||||
|
||||
@@ -159,9 +275,127 @@ CREATE INDEX IF NOT EXISTS idx_payment_order_phone ON payment_order(phone_num);
|
||||
|
||||
**分页约定**:`page` ≥1(默认 1),`size` 1..100(默认 20),service 内钳制;返回 `{total, list}`,`total` 为同条件总数(COUNT)。
|
||||
|
||||
## 强制更新
|
||||
|
||||
**背景**:客户端发版后旧版本用户无法感知新版本,bug 修复/安全更新需要强制覆盖。不走 config.yml(避免改配置重启才能下发),管理端页面上传 APK + 版本号维护 `app_version` 表,客户端启动时主动查询。**仅 Android 参与**(iOS 不做版本下发,用户从 App Store 自行更新)。
|
||||
|
||||
**数据流**:
|
||||
|
||||
```
|
||||
管理端 POST /admin/app-versions(multipart:notes + APK 文件,文件名 observer-x.y.z.apk 识别版本号)
|
||||
├─► app_version 表新增记录(version 从文件名解析,UNIQUE 防重复下发)
|
||||
└─► APK 保存为 app.apkDir/observer-latest.apk(上传即覆盖,目录永远只有一个文件)
|
||||
Android 客户端启动 GET /api/v1/app/update(公开,无需 token;iOS 不调用)
|
||||
└─► 返回最新一条记录(无记录返回空对象)
|
||||
└─► 客户端语义化比较 version > 本地版本?
|
||||
└─► 是 → 全屏阻塞弹窗(禁返回,仅「立即更新」)→ 打开 <apiBaseUrl>/download/observer-latest.apk
|
||||
└─► 否 → 正常进入
|
||||
```
|
||||
|
||||
**APK 存储与下载**:
|
||||
- 目录:`config.yml app.apkDir`(默认 `./workspace/`,与 `./data` 平级的运行时数据目录,docker-compose 已挂载持久化);服务启动时自动建目录
|
||||
- 文件:**固定文件名 `observer-latest.apk`**(`common.ApkFilename` 常量),上传流程「先落库 → 再保存临时文件 → `os.Rename` 原子覆盖」,目录下永远只有最新一个文件;落库失败删临时文件、文件保存失败删记录(补偿),保证「记录存在 ⟺ 文件存在」
|
||||
- 下载:后端静态托管 `/download` → `app.apkDir`,URL 固定 `/download/observer-latest.apk`(绕过统一响应包装,纯二进制流)
|
||||
- 客户端打开方式:`url_launcher` 跳系统浏览器下载安装(避开应用内下载的 FileProvider / 安装权限复杂度)
|
||||
|
||||
**设计决策**:
|
||||
- **检测到新版本即强制更新**(无普通/强制之分):`version > 本地版本` 即全屏阻塞弹窗,用户必须跳转下载安装才能继续使用。简化管理端操作(不需要判断「这次要不要强制」),客户端语义单一(「有新版本 = 必须更新」)
|
||||
- **公开接口**:更新检查挂在公开组(无需登录态)——旧版本登录态可能已失效,且登录前的用户也要能收到强制更新
|
||||
- **iOS 不检查**:客户端 `UpdateChecker.fetch()` 在非 Android 平台直接返回空(iOS 用户从 App Store 更新,应用内无法安装 APK,提示无意义)
|
||||
- **版本号语义化比较**:`x.y.z` 三段按数字比较(`1.10.0 > 1.9.9`),禁止字符串比较("1.9.9" > "1.10.0" 会漏判);版本号格式由文件名解析正则校验(`observer-x.y.z.apk`),DB 层 `UNIQUE` 兜底防重复
|
||||
- **文件名识别版本号**:管理端上传不再手工填写版本号,版本号从文件名解析(打包产物即 `observer-<versionName>.apk` 命名)——消除「填的版本号与包内版本不一致」的人为错误,上传文件名即唯一事实来源
|
||||
- **删除联动文件**:APK 固定文件永远对应「最新一条记录」,故删除只对最新版本联动删文件(`记录存在 ⟺ 文件存在` 的删除方向:记录没了文件即失效,客户端不再提示、下载 404);删历史版本不动文件。执行顺序「先删记录、再删文件」——记录删除是核心操作,文件删除失败仅记日志不阻断(残留旧文件无害,客户端不会误提示更新)。不存在删除/修改接口(下发后版本号不可变),删除按 `id` 定位
|
||||
- **写入串行**:新增/删除记录均走 `common.Serial()` 单写者(与授权链路一致,SQLite 无 WAL);列表/最新查询读走普通读,表极小不设查询缓存
|
||||
- **写入串行**:新增记录走 `common.Serial()` 单写者(与授权链路一致,SQLite 无 WAL);列表/最新查询读走普通读,表极小不设查询缓存
|
||||
|
||||
## 模型训练体系(数据集 → 训练 → 模型版本)
|
||||
|
||||
**背景**:识别模型(YOLOv8n → TFLite 704x704)此前在仓库根 `training/` 目录人工命令行训练(本机 MPS / GPU 服务器跑 `train_server.py`,best.tflite 手动拷进 App assets 重打包 APK 下发)。本模块将「数据集管理 → 训练编排 → 模型版本发布 → 客户端热更新」闭环搬进后台管理端,模型迭代不再依赖开发者本机与 APK 发版。**Python 训练脚本随项目整合进 `server/training/`(2026-08-26:能 Go 化的已 Go 化——prepare_yolo→prepareYoloSet、analyze_rfdetr→LocalAi.Detect,脚本删除;train_server.py/yolov8n.pt/调试工具迁移入 server/training/,yolov8n.pt 权重不进 git,由 .gitignore 排除;2026-08-26 再收口:inspect_tflite.py 并入 train_server.py 自检、dump_graph.py 保留)**。
|
||||
|
||||
**整体数据流**:
|
||||
|
||||
```
|
||||
管理端 数据集管理(上传 / AI 生成)→ 图片落服务器 workspace/datasets/<name>/
|
||||
├─► 二期:标注工作台(RF-DETR 预标注 + 人工确认)→ labels/
|
||||
└─► 发起训练 → 同步数据集到训练机 → runner 起训练 → 进度/日志/指标
|
||||
└─► 产物(best.tflite + best.pt 归档)拉回服务器
|
||||
└─► 发布 → model_version + workspace/model-latest.tflite(sha256)
|
||||
└─► 三期:客户端启动 GET /app/update → 模型热更新下载替换
|
||||
```
|
||||
|
||||
### 数据集存储与流转(决策:图片不进 DB、不提交 git)
|
||||
|
||||
- 图片目录:`app.datasetDir`(默认 `./workspace/`)下 `datasets/<数据集名>/`(图片平铺,文件名唯一防重名);标注存 `dataset_image.labels_json`(JSON 数组,每元素 `{class,cx,cy,w,h,confidence}` 归一化,AI 自动标注与人工标注同存、人工可修改/清理);**历史 labels/ txt 目录与 yolo 暂存目录已删除(2026-08-26,标注全在 DB,训练集流式化不落本地盘)**
|
||||
- **封面命名规范(2026-08-26)**:封面文件不得固定命名 `cover`,上传时解码为 **jpg 格式 + UUIDv4 命名**(`common.UuidV4()` 生成,crypto/rand 无第三方依赖,落盘 `datasets/<name>/<uuid>.jpg`,旧封面删除);DB `cover` 列只存文件名,前端用 DB 值拼 URL 零改动;删除/校验按 `isCoverName` 正则(UUID v4 + `.jpg` 后缀)识别,不认固定名——历史 `cover.jpg` 等存量封面启动时 `MigrateLegacyCovers` 幂等重命名迁移
|
||||
- DB 只存文件名/来源/prompt 等元数据(`dataset` / `dataset_image`),**禁止图片进库**
|
||||
- 生成/上传图片是**付费资产**:删除接口必须带前端确认文案(提示 AI 生成图有成本);删除 = 删文件 + 删记录,目录清理
|
||||
- 训练机与 Go 服务器可能异机:训练前按 `training` 通道同步(subprocess 同机 cp、ssh 异机 scp);**数据集在训练机上的权威路径** `training.workdir`/`training.datasetDir`/`<name>/`,DB `model_training.dataset` 只存数据集名
|
||||
- zip 导出(`GET /admin/datasets/export`):打包图片目录为 zip 下载——用于标注衔接(二期前的人工标注路径)与备份
|
||||
|
||||
### AI 生成图片(provider 抽象)
|
||||
|
||||
- `imageGen` 配置节点:`provider: dashscope`(当前唯一实现)+ `apiKey` + `model: qwen-image-3.0`;接口在 `common` 层抽象 `ImageGenProvider`(`Generate(ctx, prompt, size) ([]byte, error)`),后续可加本地 SD 等实现,config 切换
|
||||
- 尺寸参数由前端传(下拉可选,默认 1152x2048 竖图,与既有训练图一致;另提供 704x704 档)
|
||||
- 生成同步执行(count 1..8,超时 2min/张),逐张落盘 + 入库(记录 prompt 便于追溯);中途失败返回错误(已成功的图保留,不删除——付费资产原则)
|
||||
- prompt 由前端预设模板 + 用户微调;**遵守项目提示词规范:不得包含目标位置描述**(位置由模型自行推理,前端模板与校验文案落实此约束)
|
||||
|
||||
### 训练通道(runner 抽象,决策:同机/异机不确定 → 可配置)
|
||||
|
||||
```yaml
|
||||
training:
|
||||
mode: subprocess # subprocess | ssh
|
||||
ssh: { host: "", user: "", port: 22, privateKeyPath: "", password: "" }
|
||||
workdir: /opt/pheasant_data # 训练机工作目录
|
||||
venvPython: /opt/pheasant_data/venv/bin/python
|
||||
datasetDir: datasets # 训练机数据集根目录(相对 workdir,数据集为子目录)
|
||||
concurrency: 1 # GPU 独占:同时仅一个 running,新任务排队
|
||||
timeoutMinutes: 600 # 超时判死
|
||||
```
|
||||
|
||||
- service 内 `Runner` 接口:`Start(ctx, *TrainingJob) (pid, error)` / `FetchLogTail(ctx, job)` / `IsAlive(ctx, job) bool` / `Cancel(ctx, job)` / `FetchArtifacts(ctx, job, destDir)`;`subprocess` 与 `ssh` 两个实现,按 config `mode` 选择;**ssh 凭据直接读本节点 `training.ssh` 配置**(见「全局训练配置」节)
|
||||
- **训练脚本**(`server/training/train_server.py`,随项目迁移):支持 `--task-json <file>`(含 dataset/imgsz/epochs/batch/device/project 名),每 epoch 输出一行机器可读 JSON 到 `--log-file`(`{"epoch":1,"total":150,"metrics":{...}}`),结束写 `result.json`(最终指标)+ 自动打包 `artifact.zip`(best.pt + results.csv + 曲线);Go 侧解析日志行更新进度、轮询日志尾部截断 N KB 存 `model_training.log_tail`
|
||||
- **tflite 产物自检**(2026-08-26):`inspect_tflite.py` 的 flatbuffer 解析逻辑内嵌进 `train_server.py`(`check_tflite`),训练收尾定位 `best.tflite` 后自动校验并写 `result.json` 的 `tflite_check` 字段:`{"ok":bool,"reason":string,"inputs":[{"name","shape","type"}],"outputs":[...]}`;校验规则 = 输入恰 1 张且 4 维、元素总数 == imgsz²×3(兼容 NCHW/NHWC)、输出 ≥ 1 张且 batch 维 = 1;`ok=false`(如 shape 漂移、导出异常)时 Go 侧在拉产物前直接置训练失败并带出 reason,杜绝坏产物进入发布链路;`dump_graph.py` 保留作训练机人工深度调试
|
||||
- **任务生命周期**:`running → success/failed`;取消 = 杀进程(ssh 模式远程 kill pid);超时无心跳判死;**Go 服务重启后启动扫描** running 任务按 pid 存活探测(subprocess 本机、ssh 远程 `kill -0`),进程已死则置 failed
|
||||
- **并发度 1**:发起训练时若已有 running 任务返回错误「训练进行中」;训练任务不排队(简化,管理端人工再点一次)
|
||||
- 产物拉取:成功后拉 `best.tflite` + `artifact.zip` 到服务器 `workspace/trainings/<taskId>/`,发布时引用
|
||||
- 写操作走 `common.Serial()` 单写者(SQLite 无 WAL,与既有链路一致);任务状态更新(进度轮询)为高频写,单独小事务
|
||||
|
||||
### 模型版本(每数据集一个模型,多模型体系)
|
||||
|
||||
**核心决策:每个数据集训练一个模型,模型按数据集独立版本化,App 多模型并行推理合并**——用户按需下载若干数据集的模型,加载全部已下载模型共同推理标注(类别名不同则自然互补,同类名跨模型 NMS 去重)。
|
||||
|
||||
- 版本号规则:`m<major>.<minor>.<patch>`,**同一数据集内**每次发布 patch+1(取该数据集最大版本号解析自增,无记录从 m1.0.0 起);`UNIQUE(dataset_id, version)` 防重复
|
||||
- 文件布局:当前生效副本 `workspace/models/<数据集名>/latest.tflite`(tmp + rename 原子覆盖),客户端固定下载该文件;**无 `<version>.tflite` 存档(2026-08-26 决策:不需要模型回退机制,模型只增不删不回滚)**;每数据集一个目录互不影响
|
||||
- 类别名:发布时从训练任务/数据集记录类别(训练脚本 result.json 输出 `names`),存 `model_version.labels`(JSON 数组),**App 合并推理依赖它**
|
||||
- **发布**(`POST /admin/trainings/publish`):校验任务 success + tflite 产物存在 → sha256 → 写 latest 副本 → 插 `model_version` + 该数据集旧版 `is_latest=0`;「记录存在 ⟺ 文件存在」补偿逻辑同 APK 版本管理
|
||||
- **管理端无模型管理界面**(2026-08-26 决策):删 `AdminListModels`/`AdminActivateModel`/`AdminDeleteModel` 三个管理接口,`model_version` 表保留——仅支撑客户端下发目录;版本只增不删不回滚(发布即最新)
|
||||
- **模型目录(客户端拉取)**:`GET /api/v1/models`(公开,登录态即可)返回所有数据集当前生效模型:`{datasetId, datasetName, version, labels, sizeBytes, sha256, notes, publishedAt, downloadUrl}`;下载 URL `/download/models/<数据集名>/latest.tflite`(复用 `/download` 静态托管)
|
||||
|
||||
### 标注工作台(依赖 local-ai 可达;2026-08-26 布局重构)
|
||||
|
||||
- `localAi` 配置节点:`{ baseUrl: "http://127.0.0.1:18080", model: "rfdetr-xlarge", threshold: 0.08, confConfirmed: 0.2, overlapThreshold: 0.3 }`——**部署前提:RF-DETR 服务须从 Go 服务器可达**(现跑在 Mac 上,部署时搬服务器/训练机;**未配置时图片上传/生成接口直接报错**——标注是强语义,不允许产生无标注图片;配置调不通则任务 failed,页顶错误条展示,不自动重试);**端点为配置唯一来源,无运行时覆盖**
|
||||
- **预标注自动触发(2026-08-26 决策,无手动按钮)**:手动上传/AI 生成图片入库成功后,自动对**本次新增图**发起标注;**`localAi` 未配置 → 上传/生成接口直接报错**(同步检查,生成场景在第一张生成前检查避免付费资产生成后标不了);**已有 running 标注任务(忙)→ 不报错**:上传/生成照常成功,新图由「任务成功完成后自动补标未标注图」机制兜底(每轮任务成功结束时检查该数据集未标注图,有则自动续一轮只标未标注图——失败任务不续,防配置坏时无限重试);service 逐张调 local-ai 推理(common 池并行,`label_task` 记录 total/done 进度)→ 扫描结果(YOLO 归一化 xywh + 置信度 + 建议类别,conf≥confConfirmed 为 class 0)**直写 `dataset_image.labels_json`(重跑覆盖该图标注)**;**重叠去重(2026-08-26 修订)**:NMS 风格、按置信度降序依次保留,与已保留框**重叠比(minIoU = 交叠面积/两框较小面积)> `localAi.overlapThreshold`(默认 0.3)**的框剔除(同目标被重复检出只留置信度最高者,跨 class 去重;人工画框不参与去重)——**用 minIoU 而非 IoU**:RF-DETR 对同一目标常输出一大一小两个框(标准 IoU 仅 0.3~0.5 会漏杀,实测红黄重叠即此形态),大框套小框时小框被覆盖比例高,minIoU 命中;相邻目标两框互有外露,minIoU 通常 < 0.3;**全图扫描**(项目既定规则:不套用生成规格的位置裁剪,候选宁多勿漏);不设候选确认两阶段(2026-08-26 撤销 candidates_json);`POST /admin/label-tasks` 接口保留(详情页「全量标注」按钮,2026-08-26 恢复),可手动全量/指定图重标
|
||||
- **详情页布局(无选项卡)**:分页(每页 20 条)逐行「原图 ‖ 标注图」对照;页顶标注任务进度条(自动触发任务进度,含错误展示);标注图 = 原图 + 标注框叠加(只读 canvas,与工作台同一绘制逻辑)
|
||||
- **弹窗标注**:点击原图/标注图打开放大弹窗(大 canvas,1152x2048 原尺寸),支持两点画框/点框删除/清空/类别切换(确认 0/疑似 1)/上下一张/保存——**AI 自动标注与人工框同层可编辑**(含清理 AI 框);保存即整体覆写 `dataset_image.labels_json`(JSON 数组,YOLO 归一化 xywh+类别+置信度,空=清空标注);未保存修改切换图片有确认
|
||||
- 无历史标注任务时详情页工作台数据源为 `GET /admin/label-workbench?dataset_id=`(图片 + 全量标注框),避免依赖人工先跑一次任务
|
||||
- 训练前自动整理(prepare_yolo 逻辑服务端,**内存组装不落本地盘**):数据集有 labels → 按 80/20 拆 train/val 生成 YOLO 训练集包(标注 txt 内存生成、原图只记数据集目录源路径)→ 随同步推训练机(subprocess 直写 `workdir/datasetDir/yolo/<name>/`;ssh 走 **tar 流式管道**本地不落盘);**自动标注图直接进训练集**(人工修改/清理后即为训练数据);本地无暂存目录(旧 `workspace/datasets/yolo/` 与 `labels/` 已删除)
|
||||
|
||||
### 模型热更新与多模型推理(客户端 Flutter)
|
||||
|
||||
- `GET /api/v1/app/update` 扩展:响应追加 `models` 数组(`GET /api/v1/models` 同构:datasetId/datasetName/version/labels/sizeBytes/sha256/downloadUrl)——**多模型目录,不再有单一 modelUrl**;无发布模型时不返回该字段(旧 App 忽略、新 App 兼容旧服务器)
|
||||
- **模型管理**:App 设置页「模型管理」——列出服务器全部可用模型 + 本地已下载状态(版本/大小/说明),用户**自由下载/删除/启用**;下载到应用私有目录 `.tmp` → sha256 校验 → 原子替换;已下载模型跨启动持久
|
||||
- **并行推理合并**:`tflite_detector` 启动时加载**全部已启用且已下载**的模型(每模型独立 interpreter + 独立 isolate 线程并行推理);单帧结果按**类别名**合并 + 跨模型 NMS 去重(IoU 阈值 0.45,同类名才合并);类别名冲突以合并结果展示、标签文本随模型 labels
|
||||
- 内置模型兜底:assets `model.tflite` + `labels.txt` 作为无任何下载模型时的默认;有下载模型则并行加载下载模型(内置不重复加载,避免重复检测同类别)
|
||||
- **非强制**:下载失败/校验不过删 tmp 用旧状态,下次启动重试;不弹阻塞窗
|
||||
- 性能:N 模型并行推理墙钟 ≈ 单模型 1.2~1.5 倍(多核并行);用户按需下载控制 N;帧率不足时可降采样
|
||||
- 版本比较规则同 APK(数字段比较,m 前缀剥离后按 x.y.z)
|
||||
|
||||
## 待办/风险
|
||||
|
||||
- 微信支付需商户号(APP 支付权限)、APIv3 密钥与平台证书;支付宝需商户应用与密钥 —— 当前均未配置,接口按真实 SDK 契约实现,配置走 `config.yml` 占位
|
||||
- iOS App Store 对数字内容强制 IAP,微信/支付宝直连有被拒风险(已决策,记录在案)
|
||||
- 管理端静态 token 由登录页输入(浏览器 localStorage),属内网管理凭据:`config.yml admin.token` 失配/未配置时管理接口全部拒绝;token 不进入前端构建产物与 git
|
||||
- 手动授权等同免费发卡,仅限运营/客诉排查场景,前端页面带确认弹窗,后端不限制频率(按需再加审计)
|
||||
- 模型训练体系依赖外部环境:训练机 GPU + ultralytics venv(`training.venvPython`)、RF-DETR 服务从服务器可达(`localAi.baseUrl`)、DashScope API key(`imageGen.apiKey`)——未配置对应节点时相关管理接口返回配置缺失错误,不影响支付/授权主链路
|
||||
- subprocess 模式下训练与在线服务同机,训练吃满 GPU 可能影响识别类在线任务(当前在线服务无 GPU 推理,风险低);异机部署切换 `training.mode: ssh`
|
||||
- 训练产物与数据集为磁盘占用大户(图片百 MB~GB、每轮产物几十 MB),`workspace/` 已挂载持久化;清理入口:删除数据集(付费资产需确认);模型无存档,仅 latest.tflite 被新发布覆盖,无需单独清理
|
||||
|
||||
+56
-10
@@ -1,14 +1,22 @@
|
||||
<script setup>
|
||||
import { computed } from 'vue'
|
||||
import { useRoute, useRouter } from 'vue-router'
|
||||
import {
|
||||
List as ListIcon,
|
||||
Key as KeyIcon,
|
||||
Refresh as RefreshIcon,
|
||||
SwitchButton as SwitchIcon,
|
||||
Picture as PictureIcon,
|
||||
} from '@element-plus/icons-vue'
|
||||
|
||||
const route = useRoute()
|
||||
const router = useRouter()
|
||||
|
||||
// 详情页 /datasets/:id 仍高亮「数据训练」菜单项
|
||||
const activeMenu = computed(() =>
|
||||
route.path.startsWith('/datasets') ? '/datasets' : route.path,
|
||||
)
|
||||
|
||||
function logout() {
|
||||
localStorage.removeItem('adminToken')
|
||||
router.push('/login')
|
||||
@@ -21,8 +29,11 @@ function logout() {
|
||||
|
||||
<el-container v-else class="app-layout">
|
||||
<el-aside width="200px" class="app-aside">
|
||||
<div class="app-logo">视野管理端</div>
|
||||
<el-menu :default-active="route.path" router>
|
||||
<el-menu :default-active="activeMenu" router>
|
||||
<el-menu-item index="/datasets">
|
||||
<el-icon><PictureIcon /></el-icon>
|
||||
<span>数据训练</span>
|
||||
</el-menu-item>
|
||||
<el-menu-item index="/orders">
|
||||
<el-icon><ListIcon /></el-icon>
|
||||
<span>订单管理</span>
|
||||
@@ -31,6 +42,10 @@ function logout() {
|
||||
<el-icon><KeyIcon /></el-icon>
|
||||
<span>授权管理</span>
|
||||
</el-menu-item>
|
||||
<el-menu-item index="/app-versions">
|
||||
<el-icon><RefreshIcon /></el-icon>
|
||||
<span>版本管理</span>
|
||||
</el-menu-item>
|
||||
</el-menu>
|
||||
</el-aside>
|
||||
<el-container>
|
||||
@@ -52,14 +67,6 @@ function logout() {
|
||||
.app-aside {
|
||||
background: #001529;
|
||||
}
|
||||
.app-logo {
|
||||
height: 56px;
|
||||
line-height: 56px;
|
||||
text-align: center;
|
||||
color: #fff;
|
||||
font-size: 16px;
|
||||
font-weight: 600;
|
||||
}
|
||||
.app-aside :deep(.el-menu) {
|
||||
border-right: none;
|
||||
background: #001529;
|
||||
@@ -86,4 +93,43 @@ function logout() {
|
||||
.app-main {
|
||||
background: #f0f2f5;
|
||||
}
|
||||
|
||||
/* 手机竖屏:侧栏收为顶部横向菜单,页面自然滚动 */
|
||||
@media (max-width: 767px) {
|
||||
.app-layout {
|
||||
display: block;
|
||||
height: auto;
|
||||
}
|
||||
.app-aside {
|
||||
width: 100% !important;
|
||||
}
|
||||
.app-logo {
|
||||
height: 48px;
|
||||
line-height: 48px;
|
||||
font-size: 15px;
|
||||
}
|
||||
.app-aside :deep(.el-menu) {
|
||||
display: flex;
|
||||
flex-wrap: wrap;
|
||||
border-bottom: none;
|
||||
}
|
||||
.app-aside :deep(.el-menu-item) {
|
||||
flex: 1;
|
||||
justify-content: center;
|
||||
height: 44px;
|
||||
line-height: 44px;
|
||||
min-width: 96px;
|
||||
}
|
||||
.app-header {
|
||||
height: 48px;
|
||||
padding: 0 12px;
|
||||
}
|
||||
.app-main {
|
||||
overflow: visible;
|
||||
padding: 10px;
|
||||
}
|
||||
.app-main :deep(.el-card) {
|
||||
overflow-x: auto;
|
||||
}
|
||||
}
|
||||
</style>
|
||||
|
||||
@@ -2,15 +2,21 @@ import { createRouter, createWebHistory } from 'vue-router'
|
||||
import Login from '../views/Login.vue'
|
||||
import Orders from '../views/Orders.vue'
|
||||
import Licenses from '../views/Licenses.vue'
|
||||
import AppVersions from '../views/AppVersions.vue'
|
||||
import Datasets from '../views/Datasets.vue'
|
||||
import DatasetDetail from '../views/DatasetDetail.vue'
|
||||
|
||||
// 生产部署于 /admin/ 前缀(后端静态托管),history 路由由后端 SPA fallback 兜底
|
||||
const router = createRouter({
|
||||
history: createWebHistory('/admin/'),
|
||||
routes: [
|
||||
{ path: '/login', name: 'login', component: Login, meta: { title: '登录' } },
|
||||
{ path: '/', redirect: '/orders' },
|
||||
{ path: '/', redirect: '/datasets' },
|
||||
{ path: '/datasets', name: 'datasets', component: Datasets, meta: { title: '数据训练' } },
|
||||
{ path: '/datasets/:id', name: 'datasetDetail', component: DatasetDetail, meta: { title: '数据训练详情' } },
|
||||
{ path: '/orders', name: 'orders', component: Orders, meta: { title: '订单管理' } },
|
||||
{ path: '/licenses', name: 'licenses', component: Licenses, meta: { title: '授权管理' } },
|
||||
{ path: '/app-versions', name: 'appVersions', component: AppVersions, meta: { title: '版本管理' } },
|
||||
],
|
||||
})
|
||||
|
||||
|
||||
@@ -15,3 +15,25 @@ body {
|
||||
color: #303133;
|
||||
background: #f0f2f5;
|
||||
}
|
||||
|
||||
/* 手机竖屏适配:筛选表单换行、控件全宽、分页换行居中 */
|
||||
@media (max-width: 767px) {
|
||||
.el-form--inline {
|
||||
display: flex;
|
||||
flex-wrap: wrap;
|
||||
}
|
||||
.el-form--inline .el-form-item {
|
||||
margin-right: 0;
|
||||
flex: 1 1 100%;
|
||||
}
|
||||
.el-form--inline .el-input,
|
||||
.el-form--inline .el-select,
|
||||
.el-form--inline .el-date-editor {
|
||||
width: 100% !important;
|
||||
}
|
||||
.el-pagination {
|
||||
flex-wrap: wrap;
|
||||
justify-content: center;
|
||||
row-gap: 8px;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,195 @@
|
||||
<script setup>
|
||||
import { onMounted, reactive, ref } from 'vue'
|
||||
import { ElMessage, ElMessageBox } from 'element-plus'
|
||||
import { UploadFilled } from '@element-plus/icons-vue'
|
||||
import request from '../api/request'
|
||||
|
||||
const downloadUrl = `${location.origin}/download/observer-latest.apk`
|
||||
|
||||
const loading = ref(false)
|
||||
const list = ref([])
|
||||
const total = ref(0)
|
||||
const page = ref(1)
|
||||
const size = ref(20)
|
||||
|
||||
const addVisible = ref(false)
|
||||
const addForm = reactive({ notes: '', file: null })
|
||||
const adding = ref(false)
|
||||
|
||||
function load() {
|
||||
loading.value = true
|
||||
request
|
||||
.get('/app-versions', { params: { page: page.value, size: size.value } })
|
||||
.then((data) => {
|
||||
list.value = data.list || []
|
||||
total.value = data.total || 0
|
||||
})
|
||||
.finally(() => {
|
||||
loading.value = false
|
||||
})
|
||||
}
|
||||
|
||||
function openAdd() {
|
||||
addForm.notes = ''
|
||||
addForm.file = null
|
||||
fileList.value = []
|
||||
addVisible.value = true
|
||||
}
|
||||
|
||||
function onFileChange(uploadFile) {
|
||||
addForm.file = uploadFile.raw || null
|
||||
}
|
||||
|
||||
function submitAdd() {
|
||||
if (!addForm.file) {
|
||||
ElMessage.warning('请选择 APK 文件')
|
||||
return
|
||||
}
|
||||
adding.value = true
|
||||
const fd = new FormData()
|
||||
fd.append('notes', addForm.notes.trim())
|
||||
fd.append('file', addForm.file)
|
||||
request
|
||||
.post('/app-versions', fd, { timeout: 300000 })
|
||||
.then(() => {
|
||||
ElMessage.success('已下发,覆盖服务器最新 APK')
|
||||
addVisible.value = false
|
||||
page.value = 1
|
||||
load()
|
||||
})
|
||||
.finally(() => {
|
||||
adding.value = false
|
||||
})
|
||||
}
|
||||
|
||||
const fileList = ref([])
|
||||
|
||||
function removeRow(row) {
|
||||
ElMessageBox.confirm(
|
||||
`确定删除版本 ${row.version}?删除最新版本会同时移除已下发的 APK(客户端将不再提示该更新)`,
|
||||
'删除版本记录',
|
||||
{
|
||||
confirmButtonText: '删除',
|
||||
cancelButtonText: '取消',
|
||||
type: 'warning',
|
||||
},
|
||||
)
|
||||
.then(() => request.post('/app-versions/delete', { id: row.id }))
|
||||
.then(() => {
|
||||
ElMessage.success('已删除')
|
||||
load()
|
||||
})
|
||||
.catch(() => {})
|
||||
}
|
||||
|
||||
onMounted(load)
|
||||
</script>
|
||||
|
||||
<template>
|
||||
<el-card shadow="never">
|
||||
<div class="toolbar">
|
||||
<span class="tip">仅 Android:检测到服务器版本高于手机已装版本即强制更新(不可跳过)。APK 文件名须为 observer-x.y.z.apk(如 observer-1.0.1.apk),版本号从文件名识别;上传覆盖保存为固定文件,服务器永远只保留最新一个;下载地址 {{ downloadUrl }}。</span>
|
||||
<el-button type="primary" @click="openAdd">下发新版本</el-button>
|
||||
</div>
|
||||
|
||||
<el-table v-loading="loading" :data="list" border stripe>
|
||||
<el-table-column prop="version" label="版本号" width="140" />
|
||||
<el-table-column prop="notes" label="更新说明" min-width="260" show-overflow-tooltip>
|
||||
<template #default="{ row }">{{ row.notes || '-' }}</template>
|
||||
</el-table-column>
|
||||
<el-table-column prop="createdAt" label="下发时间" width="180">
|
||||
<template #default="{ row }">{{ row.createdAt || '-' }}</template>
|
||||
</el-table-column>
|
||||
<el-table-column label="操作" width="90" fixed="right">
|
||||
<template #default="{ row }">
|
||||
<el-button link type="danger" @click="removeRow(row)">删除</el-button>
|
||||
</template>
|
||||
</el-table-column>
|
||||
</el-table>
|
||||
|
||||
<el-pagination
|
||||
class="pager"
|
||||
v-model:current-page="page"
|
||||
v-model:page-size="size"
|
||||
:total="total"
|
||||
:page-sizes="[10, 20, 50, 100]"
|
||||
layout="total, sizes, prev, pager, next"
|
||||
@change="load"
|
||||
/>
|
||||
|
||||
<el-dialog v-model="addVisible" title="下发新版本(Android)" width="min(480px, 92vw)">
|
||||
<el-form label-width="90px">
|
||||
<el-form-item label="更新说明">
|
||||
<el-input
|
||||
v-model="addForm.notes"
|
||||
type="textarea"
|
||||
:rows="3"
|
||||
maxlength="500"
|
||||
show-word-limit
|
||||
placeholder="客户端弹窗展示的更新说明(选填)"
|
||||
/>
|
||||
</el-form-item>
|
||||
<el-form-item label="APK 文件">
|
||||
<el-upload
|
||||
v-model:file-list="fileList"
|
||||
:auto-upload="false"
|
||||
:limit="1"
|
||||
accept=".apk"
|
||||
drag
|
||||
:on-change="onFileChange"
|
||||
:on-remove="() => (addForm.file = null)"
|
||||
:on-exceed="() => ElMessage.warning('只能选择一个 APK 文件')"
|
||||
>
|
||||
<div class="upload-hint">
|
||||
<el-icon class="upload-icon"><UploadFilled /></el-icon>
|
||||
<div>拖拽或点击选择 .apk 文件</div>
|
||||
<div class="upload-sub">文件名须为 observer-x.y.z.apk(如 observer-1.0.1.apk),版本号从文件名识别</div>
|
||||
</div>
|
||||
</el-upload>
|
||||
</el-form-item>
|
||||
</el-form>
|
||||
<template #footer>
|
||||
<el-button @click="addVisible = false">取消</el-button>
|
||||
<el-button type="primary" :loading="adding" @click="submitAdd">确认下发</el-button>
|
||||
</template>
|
||||
</el-dialog>
|
||||
</el-card>
|
||||
</template>
|
||||
|
||||
<style scoped>
|
||||
.toolbar {
|
||||
display: flex;
|
||||
justify-content: space-between;
|
||||
align-items: center;
|
||||
margin-bottom: 14px;
|
||||
}
|
||||
.tip {
|
||||
color: #909399;
|
||||
font-size: 13px;
|
||||
}
|
||||
.pager {
|
||||
margin-top: 14px;
|
||||
justify-content: flex-end;
|
||||
}
|
||||
.upload-hint {
|
||||
padding: 8px 0;
|
||||
color: #909399;
|
||||
}
|
||||
.upload-icon {
|
||||
font-size: 40px;
|
||||
color: #c0c4cc;
|
||||
margin-bottom: 6px;
|
||||
}
|
||||
.upload-sub {
|
||||
font-size: 12px;
|
||||
margin-top: 6px;
|
||||
}
|
||||
|
||||
@media (max-width: 767px) {
|
||||
.toolbar {
|
||||
flex-wrap: wrap;
|
||||
gap: 10px;
|
||||
align-items: flex-start;
|
||||
}
|
||||
}
|
||||
</style>
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,565 @@
|
||||
<script setup>
|
||||
import { onBeforeUnmount, onMounted, reactive, ref } from 'vue'
|
||||
import { ElMessage, ElMessageBox } from 'element-plus'
|
||||
import { Plus, Picture, Setting, VideoPlay, Download } from '@element-plus/icons-vue'
|
||||
import { useRouter } from 'vue-router'
|
||||
import request from '../api/request'
|
||||
|
||||
const router = useRouter()
|
||||
const loading = ref(false)
|
||||
const list = ref([])
|
||||
const total = ref(0)
|
||||
const page = ref(1)
|
||||
const size = ref(12)
|
||||
const keyword = ref('')
|
||||
|
||||
const createVisible = ref(false)
|
||||
const createForm = reactive({ name: '', source: 'manual' })
|
||||
const creating = ref(false)
|
||||
|
||||
const configVisible = ref(false)
|
||||
const configSaving = ref(false)
|
||||
const configForm = reactive({
|
||||
id: 0,
|
||||
description: '',
|
||||
cover: '',
|
||||
})
|
||||
const coverFile = ref(null) // 新选择的封面文件(el-upload 单文件)
|
||||
const coverFileList = ref([]) // 封面回显:已有封面(服务端 url)或新选文件(本地预览)
|
||||
const coverDeleted = ref(false) // 用户删除了已有封面(保存时调删除接口)
|
||||
|
||||
const trainVisible = ref(false)
|
||||
const training = ref(false)
|
||||
const trainForm = reactive({ datasetId: 0, name: '', imgsz: 704, epochs: 150, batch: 16, device: '0' })
|
||||
|
||||
const statusMap = { building: '建设中', labeled: '已标注', synced: '已同步' }
|
||||
const statusTag = { building: 'info', labeled: 'success', synced: 'primary' }
|
||||
const trainStatusMap = { running: '训练中', success: '已训练', failed: '训练失败' }
|
||||
const trainStatusTag = { running: 'primary', success: 'success', failed: 'danger' }
|
||||
|
||||
function load() {
|
||||
loading.value = true
|
||||
request
|
||||
.get('/datasets', { params: { page: page.value, size: size.value, keyword: keyword.value || undefined } })
|
||||
.then((data) => {
|
||||
list.value = data.list || []
|
||||
total.value = data.total || 0
|
||||
})
|
||||
.finally(() => {
|
||||
loading.value = false
|
||||
})
|
||||
}
|
||||
|
||||
function search() {
|
||||
page.value = 1
|
||||
load()
|
||||
}
|
||||
|
||||
function openCreate() {
|
||||
createForm.name = ''
|
||||
createForm.source = 'manual'
|
||||
createVisible.value = true
|
||||
}
|
||||
|
||||
function submitCreate() {
|
||||
if (!createForm.name.trim()) {
|
||||
ElMessage.warning('请输入数据集名称')
|
||||
return
|
||||
}
|
||||
creating.value = true
|
||||
request
|
||||
.post('/datasets', { name: createForm.name.trim(), source: createForm.source })
|
||||
.then(() => {
|
||||
ElMessage.success('数据集已创建,请上传或生成图片')
|
||||
createVisible.value = false
|
||||
page.value = 1
|
||||
load()
|
||||
})
|
||||
.finally(() => {
|
||||
creating.value = false
|
||||
})
|
||||
}
|
||||
|
||||
// ---------- 数据集级配置(描述 + 封面;AI 端点/训练机 SSH 走 config.yml,与数据集无关) ----------
|
||||
|
||||
function openConfig(row) {
|
||||
Object.assign(configForm, {
|
||||
id: row.id,
|
||||
description: row.description || '',
|
||||
cover: row.cover || '',
|
||||
})
|
||||
coverFile.value = null
|
||||
coverDeleted.value = false
|
||||
// 已有封面回显:el-upload picture-card 直接以 url 项展示
|
||||
coverFileList.value = row.cover ? [{ name: '当前封面', url: imgUrl(`/api/v1/admin/datasets/cover?datasetId=${row.id}`) }] : []
|
||||
configVisible.value = true
|
||||
}
|
||||
|
||||
// 封面文件选中:多选时只保留最后一个(替换语义);新文件选中即视为不再删除
|
||||
function onCoverChange(file, fileList) {
|
||||
if (fileList.length > 1) coverFileList.value = [fileList[fileList.length - 1]]
|
||||
coverFile.value = coverFileList.value.length ? coverFileList.value[0].raw : null
|
||||
if (coverFile.value) coverDeleted.value = false
|
||||
}
|
||||
|
||||
// 删除项:删服务端已有封面(无 raw)→ 标记待删除;删新选文件且列表清空 → 同样按删除处理(服务端旧封面仍在)
|
||||
function onCoverRemove(file) {
|
||||
coverFile.value = null
|
||||
if (!file.raw) {
|
||||
coverDeleted.value = true
|
||||
} else if (!coverFileList.value.length && configForm.cover) {
|
||||
coverDeleted.value = true
|
||||
}
|
||||
}
|
||||
|
||||
async function submitConfig() {
|
||||
configSaving.value = true
|
||||
try {
|
||||
if (coverFile.value) {
|
||||
const fd = new FormData()
|
||||
fd.append('datasetId', configForm.id)
|
||||
fd.append('file', coverFile.value)
|
||||
await request.post('/datasets/cover', fd, { timeout: 60000 })
|
||||
} else if (coverDeleted.value) {
|
||||
await request.post('/datasets/cover/delete', { datasetId: configForm.id })
|
||||
}
|
||||
// cover 由上传接口单独维护,update 不携带
|
||||
await request.post('/datasets/update', { ...configForm, cover: '' })
|
||||
ElMessage.success('配置已保存')
|
||||
configVisible.value = false
|
||||
load()
|
||||
} finally {
|
||||
configSaving.value = false
|
||||
}
|
||||
}
|
||||
|
||||
// ---------- 开始训练 ----------
|
||||
|
||||
function openTrain(row) {
|
||||
trainForm.datasetId = row.id
|
||||
trainForm.name = `${row.name} 训练`
|
||||
trainForm.imgsz = 704
|
||||
trainForm.epochs = 150
|
||||
trainForm.batch = 16
|
||||
trainForm.device = '0'
|
||||
trainVisible.value = true
|
||||
}
|
||||
|
||||
function submitTrain() {
|
||||
if (!trainForm.name.trim()) {
|
||||
ElMessage.warning('请输入任务名称')
|
||||
return
|
||||
}
|
||||
training.value = true
|
||||
request
|
||||
.post('/trainings', { ...trainForm })
|
||||
.then(() => {
|
||||
ElMessage.success('训练任务已发起,完成后可发布为模型版本')
|
||||
trainVisible.value = false
|
||||
load()
|
||||
})
|
||||
.catch(() => {})
|
||||
.finally(() => {
|
||||
training.value = false
|
||||
})
|
||||
}
|
||||
|
||||
// ---------- 导出 / 删除 ----------
|
||||
|
||||
function exportDataset(row) {
|
||||
window.open(`${location.origin}/api/v1/admin/datasets/export?datasetId=${row.id}`, '_blank')
|
||||
}
|
||||
|
||||
function removeDataset(row) {
|
||||
const warn =
|
||||
row.source === 'ai' || row.imageCount > 0
|
||||
? '将删除图片文件、标注与记录,AI 生成图属付费资产,删除后无法恢复。'
|
||||
: ''
|
||||
ElMessageBox.confirm(`确定删除数据集「${row.name}」?${warn}`, '删除数据集', {
|
||||
confirmButtonText: '删除',
|
||||
cancelButtonText: '取消',
|
||||
type: 'warning',
|
||||
})
|
||||
.then(() => request.post('/datasets/delete', { id: row.id }))
|
||||
.then(() => {
|
||||
ElMessage.success('已删除')
|
||||
load()
|
||||
})
|
||||
.catch(() => {})
|
||||
}
|
||||
|
||||
// <img> 标签无法带自定义请求头,图片地址以 query 参数携带 admin token
|
||||
function imgUrl(u) {
|
||||
return `${location.origin}${u}&token=${encodeURIComponent(localStorage.getItem('adminToken') || '')}`
|
||||
}
|
||||
|
||||
function coverUrl(row) {
|
||||
if (!row.cover) return ''
|
||||
return imgUrl(`/api/v1/admin/datasets/cover?datasetId=${row.id}`)
|
||||
}
|
||||
|
||||
// 训练进度(百分比):有总轮数才显示比例
|
||||
function trainPercent(row) {
|
||||
if (row.trainingStatus !== 'running' || !row.trainingTotalEpochs) return 0
|
||||
return Math.min(100, Math.round((row.trainingCurrentEpoch / row.trainingTotalEpochs) * 100))
|
||||
}
|
||||
|
||||
function publishModel(row) {
|
||||
ElMessageBox.confirm(
|
||||
`确定将「${row.name}」的最新训练结果发布为模型版本?将按数据集版本号自增(m<主>.<次>.<修订>),客户端下次启动热更新下载。`,
|
||||
'发布模型版本',
|
||||
{ confirmButtonText: '发布', cancelButtonText: '取消', type: 'warning' },
|
||||
)
|
||||
.then(() => request.post('/trainings/publish', { id: row.trainingId }))
|
||||
.then((data) => {
|
||||
ElMessage.success(`已发布版本 ${data.version}`)
|
||||
load()
|
||||
})
|
||||
.catch(() => {})
|
||||
}
|
||||
|
||||
// 有训练中的任务时周期刷新,卡片进度条保持最新
|
||||
let trainPoll = null
|
||||
onMounted(() => {
|
||||
load()
|
||||
trainPoll = setInterval(() => {
|
||||
if (list.value.some((r) => r.trainingStatus === 'running')) load()
|
||||
}, 10000)
|
||||
})
|
||||
|
||||
onBeforeUnmount(() => {
|
||||
if (trainPoll) clearInterval(trainPoll)
|
||||
})
|
||||
</script>
|
||||
|
||||
<template>
|
||||
<el-card shadow="never">
|
||||
<div class="toolbar">
|
||||
<div class="toolbar-left">
|
||||
<el-input
|
||||
v-model="keyword"
|
||||
class="kw"
|
||||
placeholder="按名称搜索"
|
||||
clearable
|
||||
@keyup.enter="search"
|
||||
@clear="search"
|
||||
/>
|
||||
<el-button @click="search">搜索</el-button>
|
||||
</div>
|
||||
<div class="toolbar-right">
|
||||
<el-button type="primary" :icon="Plus" @click="openCreate">新建数据集</el-button>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div v-loading="loading" class="ds-grid">
|
||||
<div v-for="row in list" :key="row.id" class="ds-card" @click="router.push(`/datasets/${row.id}`)">
|
||||
<div class="ds-cover">
|
||||
<img v-if="row.cover" :src="coverUrl(row)" :alt="row.name" loading="lazy" />
|
||||
<div v-else class="ds-cover-placeholder"><el-icon :size="34"><Picture /></el-icon></div>
|
||||
<div v-if="row.source === 'ai'" class="ds-cover-source">AI 生成</div>
|
||||
</div>
|
||||
<div class="ds-body">
|
||||
<div class="ds-name" :title="row.name">{{ row.name }}</div>
|
||||
<div class="ds-desc">{{ row.description || '暂无描述,点击进入管理图片与标注' }}</div>
|
||||
<div class="ds-stats">
|
||||
<span>图片 {{ row.imageCount }}</span>
|
||||
<span>已标注 {{ row.labeledCount }}</span>
|
||||
<el-tag :type="statusTag[row.status] || 'info'" size="small">{{ statusMap[row.status] || row.status }}</el-tag>
|
||||
</div>
|
||||
</div>
|
||||
<div class="ds-actions" @click.stop>
|
||||
<el-button type="primary" size="small" :icon="VideoPlay" :disabled="row.trainingStatus === 'running'" @click="openTrain(row)">
|
||||
{{ row.trainingStatus === 'running' ? '训练中' : '开始训练' }}
|
||||
</el-button>
|
||||
<el-button size="small" :icon="Setting" @click="openConfig(row)">配置</el-button>
|
||||
<el-button size="small" :icon="Download" @click="exportDataset(row)">导出</el-button>
|
||||
<el-button size="small" type="danger" @click="removeDataset(row)">删除</el-button>
|
||||
</div>
|
||||
|
||||
<!-- 训练任务进度与状态(最新一条训练记录) -->
|
||||
<div v-if="row.trainingStatus" class="ds-train" @click.stop>
|
||||
<template v-if="row.trainingStatus === 'running'">
|
||||
<el-progress :percentage="trainPercent(row)" :stroke-width="6" :show-text="false" class="ds-train-bar" />
|
||||
<div class="ds-train-text">
|
||||
<span>训练中 {{ row.trainingCurrentEpoch || 0 }}/{{ row.trainingTotalEpochs || '-' }} 轮</span>
|
||||
</div>
|
||||
</template>
|
||||
<template v-else>
|
||||
<el-tag :type="trainStatusTag[row.trainingStatus] || 'info'" size="small">
|
||||
{{ trainStatusMap[row.trainingStatus] || row.trainingStatus }}
|
||||
</el-tag>
|
||||
<el-button v-if="row.trainingStatus === 'success' && row.trainingId" size="small" type="success" @click="publishModel(row)">
|
||||
发布模型
|
||||
</el-button>
|
||||
</template>
|
||||
</div>
|
||||
</div>
|
||||
<el-empty v-if="!loading && !list.length" description="暂无数据集,点击右上角新建" />
|
||||
</div>
|
||||
|
||||
<el-pagination
|
||||
class="pager"
|
||||
v-model:current-page="page"
|
||||
v-model:page-size="size"
|
||||
:total="total"
|
||||
:page-sizes="[12, 24, 48]"
|
||||
layout="total, sizes, prev, pager, next"
|
||||
@change="load"
|
||||
/>
|
||||
|
||||
<!-- 新建数据集 -->
|
||||
<el-dialog v-model="createVisible" title="新建数据集" width="min(440px, 92vw)">
|
||||
<el-form label-width="90px">
|
||||
<el-form-item label="名称">
|
||||
<el-input
|
||||
v-model="createForm.name"
|
||||
maxlength="50"
|
||||
show-word-limit
|
||||
placeholder="中文/字母/数字/下划线/短横线,唯一且作磁盘目录名"
|
||||
/>
|
||||
</el-form-item>
|
||||
<el-form-item label="图片来源">
|
||||
<el-radio-group v-model="createForm.source">
|
||||
<el-radio value="manual">手动上传</el-radio>
|
||||
<el-radio value="ai">AI 生成</el-radio>
|
||||
</el-radio-group>
|
||||
</el-form-item>
|
||||
</el-form>
|
||||
<template #footer>
|
||||
<el-button @click="createVisible = false">取消</el-button>
|
||||
<el-button type="primary" :loading="creating" @click="submitCreate">创建</el-button>
|
||||
</template>
|
||||
</el-dialog>
|
||||
|
||||
<!-- 数据集配置 -->
|
||||
<el-dialog v-model="configVisible" title="数据集配置" width="min(720px, 94vw)">
|
||||
<el-form label-width="120px">
|
||||
<el-form-item label="描述">
|
||||
<el-input v-model="configForm.description" type="textarea" :rows="6" maxlength="500" show-word-limit placeholder="数据集说明,展示在卡片上" />
|
||||
</el-form-item>
|
||||
<el-form-item label="封面">
|
||||
<div class="cover-row">
|
||||
<el-upload
|
||||
v-model:file-list="coverFileList"
|
||||
:auto-upload="false"
|
||||
accept=".jpg,.jpeg,.png"
|
||||
list-type="picture-card"
|
||||
:on-change="onCoverChange"
|
||||
:on-remove="onCoverRemove"
|
||||
>
|
||||
<div v-if="coverFileList.length === 0" class="upload-tile">
|
||||
<el-icon :size="22"><Plus /></el-icon>
|
||||
</div>
|
||||
</el-upload>
|
||||
<span class="add-tip">jpg/jpeg/png,≤2MB;悬停预览右上角 × 可删除当前封面</span>
|
||||
</div>
|
||||
</el-form-item>
|
||||
</el-form>
|
||||
<template #footer>
|
||||
<el-button @click="configVisible = false">取消</el-button>
|
||||
<el-button type="primary" :loading="configSaving" @click="submitConfig">保存</el-button>
|
||||
</template>
|
||||
</el-dialog>
|
||||
|
||||
<!-- 开始训练 -->
|
||||
<el-dialog v-model="trainVisible" title="开始训练" width="min(480px, 92vw)">
|
||||
<el-form label-width="90px">
|
||||
<el-form-item label="任务名称">
|
||||
<el-input v-model="trainForm.name" maxlength="50" show-word-limit placeholder="如:雉鸡数据集 v2 训练" />
|
||||
</el-form-item>
|
||||
<el-form-item label="分辨率">
|
||||
<el-input-number v-model="trainForm.imgsz" :min="64" :max="2048" :step="32" />
|
||||
<span class="field-tip">与端侧推理对齐,默认 704</span>
|
||||
</el-form-item>
|
||||
<el-form-item label="轮数">
|
||||
<el-input-number v-model="trainForm.epochs" :min="1" :max="1000" />
|
||||
</el-form-item>
|
||||
<el-form-item label="Batch">
|
||||
<el-input-number v-model="trainForm.batch" :min="1" :max="128" />
|
||||
</el-form-item>
|
||||
<el-form-item label="设备">
|
||||
<el-input v-model="trainForm.device" maxlength="16" placeholder="GPU 编号 0,CPU 填 cpu" />
|
||||
</el-form-item>
|
||||
<div class="start-tip">训练前自动整理 80/20 训练/验证集,须先完成图片标注;训练机并发度 1,已有运行中任务会被拒绝。</div>
|
||||
</el-form>
|
||||
<template #footer>
|
||||
<el-button @click="trainVisible = false">取消</el-button>
|
||||
<el-button type="primary" :loading="training" @click="submitTrain">发起训练</el-button>
|
||||
</template>
|
||||
</el-dialog>
|
||||
</el-card>
|
||||
</template>
|
||||
|
||||
<style scoped>
|
||||
.toolbar {
|
||||
display: flex;
|
||||
justify-content: space-between;
|
||||
align-items: center;
|
||||
margin-bottom: 14px;
|
||||
gap: 10px;
|
||||
}
|
||||
.toolbar-left {
|
||||
display: flex;
|
||||
gap: 8px;
|
||||
}
|
||||
.toolbar-right {
|
||||
display: flex;
|
||||
gap: 8px;
|
||||
}
|
||||
.kw {
|
||||
width: 220px;
|
||||
}
|
||||
.pager {
|
||||
margin-top: 14px;
|
||||
justify-content: flex-end;
|
||||
}
|
||||
.ds-grid {
|
||||
display: grid;
|
||||
grid-template-columns: repeat(auto-fill, minmax(300px, 1fr));
|
||||
gap: 14px;
|
||||
min-height: 120px;
|
||||
}
|
||||
.ds-card {
|
||||
background: #fff;
|
||||
border: 1px solid #ebeef5;
|
||||
border-radius: 8px;
|
||||
overflow: hidden;
|
||||
cursor: pointer;
|
||||
transition: box-shadow 0.2s;
|
||||
}
|
||||
.ds-card:hover {
|
||||
box-shadow: 0 4px 16px rgba(0, 0, 0, 0.08);
|
||||
}
|
||||
.ds-cover {
|
||||
position: relative;
|
||||
height: 150px;
|
||||
background: #f5f7fa;
|
||||
}
|
||||
.ds-cover img {
|
||||
width: 100%;
|
||||
height: 100%;
|
||||
object-fit: contain;
|
||||
display: block;
|
||||
}
|
||||
.ds-cover-placeholder {
|
||||
height: 100%;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
color: #c0c4cc;
|
||||
}
|
||||
.ds-cover-source {
|
||||
position: absolute;
|
||||
left: 8px;
|
||||
top: 8px;
|
||||
background: rgba(0, 0, 0, 0.55);
|
||||
color: #fff;
|
||||
font-size: 11px;
|
||||
padding: 2px 8px;
|
||||
border-radius: 10px;
|
||||
}
|
||||
.ds-train {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 10px;
|
||||
padding: 8px 12px 12px;
|
||||
border-top: 1px dashed #ebeef5;
|
||||
}
|
||||
.ds-train-bar {
|
||||
flex: 1;
|
||||
}
|
||||
.ds-train-text {
|
||||
font-size: 12px;
|
||||
color: #606266;
|
||||
white-space: nowrap;
|
||||
}
|
||||
.ds-body {
|
||||
padding: 10px 12px 6px;
|
||||
}
|
||||
.ds-name {
|
||||
font-size: 15px;
|
||||
font-weight: 600;
|
||||
overflow: hidden;
|
||||
text-overflow: ellipsis;
|
||||
white-space: nowrap;
|
||||
}
|
||||
.ds-desc {
|
||||
font-size: 12px;
|
||||
color: #909399;
|
||||
line-height: 1.5;
|
||||
height: 36px;
|
||||
overflow: hidden;
|
||||
margin-top: 4px;
|
||||
}
|
||||
.ds-stats {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 10px;
|
||||
font-size: 12px;
|
||||
color: #606266;
|
||||
margin-top: 8px;
|
||||
}
|
||||
.ds-actions {
|
||||
display: flex;
|
||||
flex-wrap: wrap;
|
||||
gap: 4px;
|
||||
padding: 8px 12px 12px;
|
||||
border-top: 1px dashed #ebeef5;
|
||||
margin-top: 8px;
|
||||
}
|
||||
.field-tip {
|
||||
margin-left: 8px;
|
||||
font-size: 12px;
|
||||
color: #909399;
|
||||
}
|
||||
.start-tip {
|
||||
font-size: 12px;
|
||||
color: #909399;
|
||||
line-height: 1.6;
|
||||
margin-left: 90px;
|
||||
}
|
||||
.cover-row {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 10px;
|
||||
width: 100%;
|
||||
}
|
||||
/* 上传方块:虚线边框 + 居中加号(el-upload picture-card 触发块) */
|
||||
.upload-tile {
|
||||
width: 100%;
|
||||
height: 100%;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
color: #8c939d;
|
||||
}
|
||||
.add-tip {
|
||||
font-size: 12px;
|
||||
color: #909399;
|
||||
line-height: 1.6;
|
||||
align-self: flex-start;
|
||||
margin-top: 12px;
|
||||
}
|
||||
|
||||
@media (max-width: 767px) {
|
||||
.toolbar {
|
||||
flex-wrap: wrap;
|
||||
align-items: flex-start;
|
||||
}
|
||||
.toolbar-left {
|
||||
flex: 1;
|
||||
min-width: 200px;
|
||||
}
|
||||
.toolbar-right {
|
||||
width: 100%;
|
||||
justify-content: flex-end;
|
||||
}
|
||||
.kw {
|
||||
flex: 1;
|
||||
width: auto;
|
||||
}
|
||||
.ds-grid {
|
||||
grid-template-columns: 1fr;
|
||||
}
|
||||
}
|
||||
</style>
|
||||
@@ -14,6 +14,10 @@ const grantVisible = ref(false)
|
||||
const grantForm = reactive({ phoneNum: '', planId: 'day' })
|
||||
const granting = ref(false)
|
||||
|
||||
const remarkVisible = ref(false)
|
||||
const remarkForm = reactive({ phoneNum: '', remark: '' })
|
||||
const remarking = ref(false)
|
||||
|
||||
const planOptions = [
|
||||
{ value: 'day', label: '1天' },
|
||||
{ value: 'week', label: '7天' },
|
||||
@@ -64,6 +68,29 @@ function submitGrant() {
|
||||
})
|
||||
}
|
||||
|
||||
function openRemark(row) {
|
||||
remarkForm.phoneNum = row.phoneNum
|
||||
remarkForm.remark = row.remark || ''
|
||||
remarkVisible.value = true
|
||||
}
|
||||
|
||||
function submitRemark() {
|
||||
remarking.value = true
|
||||
request
|
||||
.post('/licenses/remark', {
|
||||
phoneNum: remarkForm.phoneNum.trim(),
|
||||
remark: remarkForm.remark.trim(),
|
||||
})
|
||||
.then(() => {
|
||||
ElMessage.success(remarkForm.remark ? '备注已保存' : '备注已清空')
|
||||
remarkVisible.value = false
|
||||
load()
|
||||
})
|
||||
.finally(() => {
|
||||
remarking.value = false
|
||||
})
|
||||
}
|
||||
|
||||
function revoke(row) {
|
||||
ElMessageBox.confirm(`确认撤销 ${row.phoneNum} 的授权?撤销后该账号额度立即失效,账号保留可继续登录。`, '撤销授权', {
|
||||
type: 'warning',
|
||||
@@ -112,9 +139,13 @@ onMounted(load)
|
||||
<el-table-column prop="updatedAt" label="更新时间" width="180">
|
||||
<template #default="{ row }">{{ row.updatedAt || '-' }}</template>
|
||||
</el-table-column>
|
||||
<el-table-column label="操作" width="130" fixed="right">
|
||||
<el-table-column prop="remark" label="备注" min-width="160" show-overflow-tooltip>
|
||||
<template #default="{ row }">{{ row.remark || '-' }}</template>
|
||||
</el-table-column>
|
||||
<el-table-column label="操作" width="180" fixed="right">
|
||||
<template #default="{ row }">
|
||||
<el-button type="primary" link size="small" @click="openGrant(row)">授权</el-button>
|
||||
<el-button type="warning" link size="small" @click="openRemark(row)">备注</el-button>
|
||||
<el-button type="danger" link size="small" :disabled="!row.expiresAt" @click="revoke(row)">撤销</el-button>
|
||||
</template>
|
||||
</el-table-column>
|
||||
@@ -130,7 +161,7 @@ onMounted(load)
|
||||
@change="load"
|
||||
/>
|
||||
|
||||
<el-dialog v-model="grantVisible" title="手动授权" width="420px">
|
||||
<el-dialog v-model="grantVisible" title="手动授权" width="min(420px, 92vw)">
|
||||
<el-form label-width="80px">
|
||||
<el-form-item label="手机号">
|
||||
<el-input :model-value="grantForm.phoneNum" disabled />
|
||||
@@ -146,6 +177,28 @@ onMounted(load)
|
||||
<el-button type="primary" :loading="granting" @click="submitGrant">确认授权</el-button>
|
||||
</template>
|
||||
</el-dialog>
|
||||
|
||||
<el-dialog v-model="remarkVisible" title="账号备注" width="min(420px, 92vw)">
|
||||
<el-form label-width="80px">
|
||||
<el-form-item label="手机号">
|
||||
<el-input :model-value="remarkForm.phoneNum" disabled />
|
||||
</el-form-item>
|
||||
<el-form-item label="备注">
|
||||
<el-input
|
||||
v-model="remarkForm.remark"
|
||||
type="textarea"
|
||||
:rows="3"
|
||||
maxlength="200"
|
||||
show-word-limit
|
||||
placeholder="运营备注(如发卡渠道/人工记录),留空提交即清空"
|
||||
/>
|
||||
</el-form-item>
|
||||
</el-form>
|
||||
<template #footer>
|
||||
<el-button @click="remarkVisible = false">取消</el-button>
|
||||
<el-button type="primary" :loading="remarking" @click="submitRemark">保存</el-button>
|
||||
</template>
|
||||
</el-dialog>
|
||||
</el-card>
|
||||
</template>
|
||||
|
||||
|
||||
@@ -23,7 +23,7 @@ function login() {
|
||||
.then(() => {
|
||||
localStorage.setItem('adminToken', token.value.trim())
|
||||
ElMessage.success('登录成功')
|
||||
router.push('/orders')
|
||||
router.push('/datasets')
|
||||
})
|
||||
.catch(() => {
|
||||
errorMsg.value = 'token 无效或服务端未配置(检查 config.yml admin.token)'
|
||||
@@ -64,7 +64,7 @@ function login() {
|
||||
background: #001529;
|
||||
}
|
||||
.login-card {
|
||||
width: 380px;
|
||||
width: min(380px, 92vw);
|
||||
padding: 12px 8px;
|
||||
}
|
||||
.login-title {
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user