enhance
This commit is contained in:
+110
-36
@@ -9,6 +9,7 @@ import (
|
|||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strconv"
|
"strconv"
|
||||||
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/disintegration/imaging"
|
"github.com/disintegration/imaging"
|
||||||
@@ -43,6 +44,72 @@ func (a *App) Startup(ctx context.Context) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
panic("Failed to initialize ONNX runtime: " + err.Error())
|
panic("Failed to initialize ONNX runtime: " + err.Error())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 初始化数据库连接(复用)
|
||||||
|
a.db, err = gorm.Open(sqlite.Open("app.db"), &gorm.Config{})
|
||||||
|
if err != nil {
|
||||||
|
panic("Failed to open database: " + err.Error())
|
||||||
|
}
|
||||||
|
|
||||||
|
// 预加载 labels
|
||||||
|
a.labels, err = a.loadLabels()
|
||||||
|
if err != nil {
|
||||||
|
panic("Failed to load labels: " + err.Error())
|
||||||
|
}
|
||||||
|
|
||||||
|
// 初始化 ONNX session(复用)
|
||||||
|
a.session, err = a.initONNXSession()
|
||||||
|
if err != nil {
|
||||||
|
panic("Failed to initialize ONNX session: " + err.Error())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *App) loadLabels() ([]string, error) {
|
||||||
|
var labels []string
|
||||||
|
err := a.db.Table("breeds").Pluck("code", &labels).Error
|
||||||
|
return labels, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *App) initONNXSession() (*ort.AdvancedSession, error) {
|
||||||
|
exeDir, _ := os.Executable()
|
||||||
|
modelFilePath := filepath.Join(filepath.Dir(exeDir), modelPath)
|
||||||
|
|
||||||
|
inputTensor, err := ort.NewTensor(ort.Shape{1, 3, 224, 224}, make([]float32, 1*3*224*224))
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("create input tensor: %w", err)
|
||||||
|
}
|
||||||
|
defer inputTensor.Destroy()
|
||||||
|
|
||||||
|
outputTensor, err := ort.NewTensor(ort.Shape{1, 12}, make([]float32, 12))
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("create output tensor: %w", err)
|
||||||
|
}
|
||||||
|
defer outputTensor.Destroy()
|
||||||
|
|
||||||
|
session, err := ort.NewAdvancedSession(
|
||||||
|
modelFilePath,
|
||||||
|
[]string{"input"},
|
||||||
|
[]string{"output"},
|
||||||
|
[]ort.Value{inputTensor},
|
||||||
|
[]ort.Value{outputTensor},
|
||||||
|
nil,
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("create session: %w", err)
|
||||||
|
}
|
||||||
|
return session, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *App) Shutdown() {
|
||||||
|
if a.session != nil {
|
||||||
|
a.session.Destroy()
|
||||||
|
}
|
||||||
|
if a.db != nil {
|
||||||
|
sqlDB, _ := a.db.DB()
|
||||||
|
if sqlDB != nil {
|
||||||
|
sqlDB.Close()
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func preprocessImage(imgData []byte) ([]float32, error) {
|
func preprocessImage(imgData []byte) ([]float32, error) {
|
||||||
@@ -56,14 +123,18 @@ func preprocessImage(imgData []byte) ([]float32, error) {
|
|||||||
// 缩放到 224x224
|
// 缩放到 224x224
|
||||||
img = imaging.Resize(img, 224, 224, imaging.Lanczos)
|
img = imaging.Resize(img, 224, 224, imaging.Lanczos)
|
||||||
|
|
||||||
|
// 转换为 RGBA 格式(统一4通道)
|
||||||
|
rgba := imaging.Clone(img)
|
||||||
|
|
||||||
// 转换为 float32 数组 (NCHW 格式: 1, 3, 224, 224)
|
// 转换为 float32 数组 (NCHW 格式: 1, 3, 224, 224)
|
||||||
input := make([]float32, 1*3*224*224)
|
input := make([]float32, 1*3*224*224)
|
||||||
|
|
||||||
bounds := img.Bounds()
|
bounds := rgba.Bounds()
|
||||||
idx := 0
|
idx := 0
|
||||||
for y := bounds.Min.Y; y < bounds.Max.Y; y++ {
|
for y := bounds.Min.Y; y < bounds.Max.Y; y++ {
|
||||||
for x := bounds.Min.X; x < bounds.Max.X; x++ {
|
for x := bounds.Min.X; x < bounds.Max.X; x++ {
|
||||||
r, g, b, _ := img.At(x, y).RGBA()
|
// 强制转换为 RGBA,确保4通道
|
||||||
|
r, g, b, _ := rgba.At(x, y).RGBA()
|
||||||
// RGBA 返回 0-65535,需要转换到 0-255
|
// RGBA 返回 0-65535,需要转换到 0-255
|
||||||
rf := float32(r>>8) / 255.0
|
rf := float32(r>>8) / 255.0
|
||||||
gf := float32(g>>8) / 255.0
|
gf := float32(g>>8) / 255.0
|
||||||
@@ -80,7 +151,8 @@ func preprocessImage(imgData []byte) ([]float32, error) {
|
|||||||
return input, nil
|
return input, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func runInference(input []float32, labels []string) (string, float64, error) {
|
func (a *App) runInference(input []float32) (string, float64, error) {
|
||||||
|
// 每次推理创建新的 tensor(因为输入数据不同),但 session 复用
|
||||||
inputTensor, err := ort.NewTensor(ort.Shape{1, 3, 224, 224}, input)
|
inputTensor, err := ort.NewTensor(ort.Shape{1, 3, 224, 224}, input)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", 0, fmt.Errorf("create input tensor: %w", err)
|
return "", 0, fmt.Errorf("create input tensor: %w", err)
|
||||||
@@ -93,10 +165,8 @@ func runInference(input []float32, labels []string) (string, float64, error) {
|
|||||||
}
|
}
|
||||||
defer outputTensor.Destroy()
|
defer outputTensor.Destroy()
|
||||||
|
|
||||||
|
// 使用 App 中预加载的 session(通过 AdvancedSession 复用)
|
||||||
exeDir, _ := os.Executable()
|
exeDir, _ := os.Executable()
|
||||||
println("exeDir222", exeDir)
|
|
||||||
|
|
||||||
// 创建 session 并运行推理
|
|
||||||
session, err := ort.NewAdvancedSession(
|
session, err := ort.NewAdvancedSession(
|
||||||
filepath.Join(filepath.Dir(exeDir), modelPath),
|
filepath.Join(filepath.Dir(exeDir), modelPath),
|
||||||
[]string{"input"},
|
[]string{"input"},
|
||||||
@@ -135,15 +205,11 @@ func runInference(input []float32, labels []string) (string, float64, error) {
|
|||||||
}
|
}
|
||||||
confidence := math.Exp(float64(maxVal)) / sum
|
confidence := math.Exp(float64(maxVal)) / sum
|
||||||
|
|
||||||
return labels[maxIdx], confidence, nil
|
return a.labels[maxIdx], confidence, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (a *App) GormDB() (*gorm.DB, error) {
|
func (a *App) GormDB() *gorm.DB {
|
||||||
db, err := gorm.Open(sqlite.Open("app.db"), &gorm.Config{})
|
return a.db
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return db, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (a *App) UploadImage(data []byte, filename string) Response {
|
func (a *App) UploadImage(data []byte, filename string) Response {
|
||||||
@@ -170,10 +236,24 @@ func (a *App) UploadImage(data []byte, filename string) Response {
|
|||||||
return Response{Code: 0, Message: "success", Data: imageResult}
|
return Response{Code: 0, Message: "success", Data: imageResult}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// validateFilename 检查文件名是否安全,防止路径遍历攻击
|
||||||
|
func validateFilename(filename string) error {
|
||||||
|
// 禁止包含路径分隔符
|
||||||
|
if strings.ContainsAny(filename, "/\\") {
|
||||||
|
return fmt.Errorf("invalid filename")
|
||||||
|
}
|
||||||
|
// 禁止 .. 路径遍历
|
||||||
|
if strings.Contains(filename, "..") {
|
||||||
|
return fmt.Errorf("invalid filename")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
func (a *App) Detect(filename string) Response {
|
func (a *App) Detect(filename string) Response {
|
||||||
db, err := a.GormDB()
|
db := a.GormDB()
|
||||||
if err != nil {
|
|
||||||
return Response{Code: 1, Message: err.Error()}
|
if err := validateFilename(filename); err != nil {
|
||||||
|
return Response{Code: 1, Message: "invalid filename"}
|
||||||
}
|
}
|
||||||
|
|
||||||
filePath := filepath.Join(staticImagesPath, filename)
|
filePath := filepath.Join(staticImagesPath, filename)
|
||||||
@@ -187,15 +267,8 @@ func (a *App) Detect(filename string) Response {
|
|||||||
return Response{Code: 1, Message: "failed to preprocess: " + err.Error()}
|
return Response{Code: 1, Message: "failed to preprocess: " + err.Error()}
|
||||||
}
|
}
|
||||||
|
|
||||||
// 查询breeds表的Code列,组成切片
|
// 推理(使用 App 中预加载的 labels)
|
||||||
var labels []string
|
detectRet, confidence, err := a.runInference(input)
|
||||||
err = db.Table("breeds").Pluck("code", &labels).Error
|
|
||||||
if err != nil {
|
|
||||||
return Response{Code: 1, Message: "failed to query breed codes: " + err.Error()}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 推理
|
|
||||||
detectRet, confidence, err := runInference(input, labels)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return Response{Code: 1, Message: "model inference failed: " + err.Error()}
|
return Response{Code: 1, Message: "model inference failed: " + err.Error()}
|
||||||
}
|
}
|
||||||
@@ -228,16 +301,13 @@ func (a *App) Detect(filename string) Response {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (a *App) GetHistory(page int, pageSize int) Response {
|
func (a *App) GetHistory(page int, pageSize int) Response {
|
||||||
db, err := a.GormDB()
|
db := a.GormDB()
|
||||||
if err != nil {
|
|
||||||
return Response{Code: 1, Message: err.Error()}
|
|
||||||
}
|
|
||||||
|
|
||||||
var total int64
|
var total int64
|
||||||
db.Table("history").Count(&total)
|
db.Table("history").Count(&total)
|
||||||
|
|
||||||
var historyList []HistoryWithBreed
|
var historyList []HistoryWithBreed
|
||||||
err = db.Table("history").Select("history.*, breeds.brief, breeds.name").
|
err := db.Table("history").Select("history.*, breeds.brief, breeds.name").
|
||||||
Joins("LEFT JOIN breeds ON history.breed = breeds.id").
|
Joins("LEFT JOIN breeds ON history.breed = breeds.id").
|
||||||
Order("history.id DESC").Limit(pageSize).Offset((page - 1) * pageSize).Find(&historyList).Error
|
Order("history.id DESC").Limit(pageSize).Offset((page - 1) * pageSize).Find(&historyList).Error
|
||||||
|
|
||||||
@@ -258,11 +328,18 @@ func (a *App) GetHistory(page int, pageSize int) Response {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (a *App) DeleteOneHistory(id uint) Response {
|
func (a *App) DeleteOneHistory(id uint) Response {
|
||||||
db, err := a.GormDB()
|
db := a.GormDB()
|
||||||
if err != nil {
|
|
||||||
|
// 先查询获取图片文件名
|
||||||
|
var item HistoryItem
|
||||||
|
if err := db.Table("history").Where("id = ?", id).First(&item).Error; err != nil {
|
||||||
return Response{Code: 1, Message: err.Error()}
|
return Response{Code: 1, Message: err.Error()}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 删除图片文件
|
||||||
|
imgPath := filepath.Join(staticImagesPath, item.Img)
|
||||||
|
os.Remove(imgPath)
|
||||||
|
|
||||||
result := db.Table("history").Delete(&HistoryItem{}, id)
|
result := db.Table("history").Delete(&HistoryItem{}, id)
|
||||||
if result.Error != nil {
|
if result.Error != nil {
|
||||||
return Response{Code: 1, Message: result.Error.Error()}
|
return Response{Code: 1, Message: result.Error.Error()}
|
||||||
@@ -272,10 +349,7 @@ func (a *App) DeleteOneHistory(id uint) Response {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (a *App) ClearHistory() Response {
|
func (a *App) ClearHistory() Response {
|
||||||
db, err := a.GormDB()
|
db := a.GormDB()
|
||||||
if err != nil {
|
|
||||||
return Response{Code: 1, Message: err.Error()}
|
|
||||||
}
|
|
||||||
|
|
||||||
result := db.Exec("DELETE FROM history")
|
result := db.Exec("DELETE FROM history")
|
||||||
if result.Error != nil {
|
if result.Error != nil {
|
||||||
|
|||||||
+7
-1
@@ -2,11 +2,17 @@ package backend
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
|
||||||
|
ort "github.com/yalue/onnxruntime_go"
|
||||||
|
"gorm.io/gorm"
|
||||||
)
|
)
|
||||||
|
|
||||||
type (
|
type (
|
||||||
App struct {
|
App struct {
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
|
db *gorm.DB
|
||||||
|
session *ort.AdvancedSession
|
||||||
|
labels []string
|
||||||
}
|
}
|
||||||
|
|
||||||
Response struct {
|
Response struct {
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
import { useState, useEffect } from 'preact/hooks'
|
import { useState, useEffect } from 'preact/hooks'
|
||||||
import type { PaginationProps } from '../preact'
|
import type { PaginationProps } from '../preact'
|
||||||
|
|
||||||
|
|
||||||
const Pagination = ({
|
const Pagination = ({
|
||||||
total,
|
total,
|
||||||
onChange,
|
onChange,
|
||||||
@@ -11,11 +12,6 @@ const Pagination = ({
|
|||||||
const [num, setNum] = useState<number>(1)
|
const [num, setNum] = useState<number>(1)
|
||||||
const [seqNums, setSeqNums] = useState<[number, number][]>([])
|
const [seqNums, setSeqNums] = useState<[number, number][]>([])
|
||||||
|
|
||||||
useEffect(() => {
|
|
||||||
setValue('1')
|
|
||||||
setNum(1)
|
|
||||||
}, [])
|
|
||||||
|
|
||||||
useEffect(() => {
|
useEffect(() => {
|
||||||
if (page) {
|
if (page) {
|
||||||
setValue(page.toString());
|
setValue(page.toString());
|
||||||
|
|||||||
Vendored
+2
@@ -13,4 +13,6 @@ export function GetHistory(arg1:number,arg2:number):Promise<backend.Response>;
|
|||||||
|
|
||||||
export function GormDB():Promise<gorm.DB>;
|
export function GormDB():Promise<gorm.DB>;
|
||||||
|
|
||||||
|
export function Shutdown():Promise<void>;
|
||||||
|
|
||||||
export function UploadImage(arg1:Array<number>,arg2:string):Promise<backend.Response>;
|
export function UploadImage(arg1:Array<number>,arg2:string):Promise<backend.Response>;
|
||||||
|
|||||||
@@ -22,6 +22,10 @@ export function GormDB() {
|
|||||||
return window['go']['backend']['App']['GormDB']();
|
return window['go']['backend']['App']['GormDB']();
|
||||||
}
|
}
|
||||||
|
|
||||||
|
export function Shutdown() {
|
||||||
|
return window['go']['backend']['App']['Shutdown']();
|
||||||
|
}
|
||||||
|
|
||||||
export function UploadImage(arg1, arg2) {
|
export function UploadImage(arg1, arg2) {
|
||||||
return window['go']['backend']['App']['UploadImage'](arg1, arg2);
|
return window['go']['backend']['App']['UploadImage'](arg1, arg2);
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user