diff --git a/app.db b/app.db index 2c7ca50..4a5f25c 100644 Binary files a/app.db and b/app.db differ diff --git a/backend/app.go b/backend/app.go index ae3d349..d9e8c2d 100644 --- a/backend/app.go +++ b/backend/app.go @@ -9,6 +9,7 @@ import ( "os" "path/filepath" "strconv" + "strings" "time" "github.com/disintegration/imaging" @@ -43,6 +44,72 @@ func (a *App) Startup(ctx context.Context) { if err != nil { 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) { @@ -56,14 +123,18 @@ func preprocessImage(imgData []byte) ([]float32, error) { // 缩放到 224x224 img = imaging.Resize(img, 224, 224, imaging.Lanczos) + // 转换为 RGBA 格式(统一4通道) + rgba := imaging.Clone(img) + // 转换为 float32 数组 (NCHW 格式: 1, 3, 224, 224) input := make([]float32, 1*3*224*224) - bounds := img.Bounds() + bounds := rgba.Bounds() idx := 0 for y := bounds.Min.Y; y < bounds.Max.Y; y++ { 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 rf := float32(r>>8) / 255.0 gf := float32(g>>8) / 255.0 @@ -80,7 +151,8 @@ func preprocessImage(imgData []byte) ([]float32, error) { 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) if err != nil { 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() + // 使用 App 中预加载的 session(通过 AdvancedSession 复用) exeDir, _ := os.Executable() - println("exeDir222", exeDir) - - // 创建 session 并运行推理 session, err := ort.NewAdvancedSession( filepath.Join(filepath.Dir(exeDir), modelPath), []string{"input"}, @@ -135,15 +205,11 @@ func runInference(input []float32, labels []string) (string, float64, error) { } confidence := math.Exp(float64(maxVal)) / sum - return labels[maxIdx], confidence, nil + return a.labels[maxIdx], confidence, nil } -func (a *App) GormDB() (*gorm.DB, error) { - db, err := gorm.Open(sqlite.Open("app.db"), &gorm.Config{}) - if err != nil { - return nil, err - } - return db, nil +func (a *App) GormDB() *gorm.DB { + return a.db } 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} } +// 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 { - db, err := a.GormDB() - if err != nil { - return Response{Code: 1, Message: err.Error()} + db := a.GormDB() + + if err := validateFilename(filename); err != nil { + return Response{Code: 1, Message: "invalid 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()} } - // 查询breeds表的Code列,组成切片 - var labels []string - 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) + // 推理(使用 App 中预加载的 labels) + detectRet, confidence, err := a.runInference(input) if err != nil { 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 { - db, err := a.GormDB() - if err != nil { - return Response{Code: 1, Message: err.Error()} - } + db := a.GormDB() var total int64 db.Table("history").Count(&total) 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"). 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 { - db, err := a.GormDB() - if err != nil { + db := a.GormDB() + + // 先查询获取图片文件名 + var item HistoryItem + if err := db.Table("history").Where("id = ?", id).First(&item).Error; err != nil { return Response{Code: 1, Message: err.Error()} } + // 删除图片文件 + imgPath := filepath.Join(staticImagesPath, item.Img) + os.Remove(imgPath) + result := db.Table("history").Delete(&HistoryItem{}, id) if result.Error != nil { return Response{Code: 1, Message: result.Error.Error()} @@ -272,10 +349,7 @@ func (a *App) DeleteOneHistory(id uint) Response { } func (a *App) ClearHistory() Response { - db, err := a.GormDB() - if err != nil { - return Response{Code: 1, Message: err.Error()} - } + db := a.GormDB() result := db.Exec("DELETE FROM history") if result.Error != nil { diff --git a/backend/types.go b/backend/types.go index adad6fe..ba76996 100644 --- a/backend/types.go +++ b/backend/types.go @@ -2,11 +2,17 @@ package backend import ( "context" + + ort "github.com/yalue/onnxruntime_go" + "gorm.io/gorm" ) type ( App struct { - ctx context.Context + ctx context.Context + db *gorm.DB + session *ort.AdvancedSession + labels []string } Response struct { diff --git a/frontend/src/components/Pagination.tsx b/frontend/src/components/Pagination.tsx index 302376f..bff897e 100644 --- a/frontend/src/components/Pagination.tsx +++ b/frontend/src/components/Pagination.tsx @@ -1,6 +1,7 @@ import { useState, useEffect } from 'preact/hooks' import type { PaginationProps } from '../preact' + const Pagination = ({ total, onChange, @@ -11,11 +12,6 @@ const Pagination = ({ const [num, setNum] = useState(1) const [seqNums, setSeqNums] = useState<[number, number][]>([]) - useEffect(() => { - setValue('1') - setNum(1) - }, []) - useEffect(() => { if (page) { setValue(page.toString()); diff --git a/frontend/wailsjs/go/backend/App.d.ts b/frontend/wailsjs/go/backend/App.d.ts index 0658577..7041bfa 100644 --- a/frontend/wailsjs/go/backend/App.d.ts +++ b/frontend/wailsjs/go/backend/App.d.ts @@ -13,4 +13,6 @@ export function GetHistory(arg1:number,arg2:number):Promise; export function GormDB():Promise; +export function Shutdown():Promise; + export function UploadImage(arg1:Array,arg2:string):Promise; diff --git a/frontend/wailsjs/go/backend/App.js b/frontend/wailsjs/go/backend/App.js index 620bc9b..db82292 100644 --- a/frontend/wailsjs/go/backend/App.js +++ b/frontend/wailsjs/go/backend/App.js @@ -22,6 +22,10 @@ export function GormDB() { return window['go']['backend']['App']['GormDB'](); } +export function Shutdown() { + return window['go']['backend']['App']['Shutdown'](); +} + export function UploadImage(arg1, arg2) { return window['go']['backend']['App']['UploadImage'](arg1, arg2); }