feat: 启动时自动下载缺失的模型文件
This commit is contained in:
+3
-4
@@ -33,13 +33,12 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# 复制二进制文件、前端和模型
|
||||
# 复制二进制文件和前端
|
||||
COPY --from=builder /app/anticaptcha .
|
||||
COPY --from=builder /app/web/dist ./web/dist
|
||||
COPY --from=builder /app/models ./models
|
||||
|
||||
# 创建数据目录
|
||||
RUN mkdir -p /app/data
|
||||
# 创建数据和模型目录(模型启动时自动下载)
|
||||
RUN mkdir -p /app/data /app/models
|
||||
|
||||
ENV GIN_MODE=release
|
||||
ENV TZ=Asia/Shanghai
|
||||
|
||||
@@ -1,6 +1,10 @@
|
||||
package captcha
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
|
||||
"anticaptcha/pkg/opencv"
|
||||
@@ -11,10 +15,69 @@ type Handler struct {
|
||||
mu sync.RWMutex
|
||||
}
|
||||
|
||||
// 模型配置
|
||||
var modelConfigs = map[string]string{
|
||||
"math.onnx": "[AntiCAP]-CRNN_Math.onnx",
|
||||
"ocr.onnx": "[Dddd]-OCR.onnx",
|
||||
"rotate.onnx": "[AntiCAP]-Rotation-RotNetR.onnx",
|
||||
"siamese.onnx": "[AntiCAP]-Siamese-ResNet18.onnx",
|
||||
"charsets.txt": "[Dddd]-CharSets.txt",
|
||||
}
|
||||
|
||||
const modelBaseURL = "https://raw.githubusercontent.com/81NewArk/AntiCAP/main/AntiCAP/AntiCAP-Models"
|
||||
|
||||
func NewHandler(modelPath string) *Handler {
|
||||
return &Handler{
|
||||
h := &Handler{
|
||||
modelPath: modelPath,
|
||||
}
|
||||
// 确保模型目录存在并下载缺失的模型
|
||||
h.ensureModels()
|
||||
return h
|
||||
}
|
||||
|
||||
// ensureModels 检查并下载缺失的模型
|
||||
func (h *Handler) ensureModels() {
|
||||
// 确保目录存在
|
||||
if err := os.MkdirAll(h.modelPath, 0755); err != nil {
|
||||
fmt.Printf("警告: 创建模型目录失败: %v\n", err)
|
||||
return
|
||||
}
|
||||
|
||||
for localName, remoteName := range modelConfigs {
|
||||
localPath := filepath.Join(h.modelPath, localName)
|
||||
if _, err := os.Stat(localPath); os.IsNotExist(err) {
|
||||
fmt.Printf("下载模型: %s -> %s\n", remoteName, localName)
|
||||
if err := h.downloadModel(remoteName, localPath); err != nil {
|
||||
fmt.Printf("警告: 下载模型 %s 失败: %v\n", remoteName, err)
|
||||
} else {
|
||||
fmt.Printf("模型下载完成: %s\n", localName)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// downloadModel 下载模型文件
|
||||
func (h *Handler) downloadModel(remoteName, localPath string) error {
|
||||
url := fmt.Sprintf("%s/%s", modelBaseURL, remoteName)
|
||||
|
||||
resp, err := http.Get(url)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return fmt.Errorf("HTTP %d", resp.StatusCode)
|
||||
}
|
||||
|
||||
out, err := os.Create(localPath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer out.Close()
|
||||
|
||||
_, err = out.ReadFrom(resp.Body)
|
||||
return err
|
||||
}
|
||||
|
||||
// OCR 文字识别(需要 ONNX 模型)
|
||||
@@ -143,4 +206,4 @@ func (h *Handler) DetectionIconOrder(orderImgBase64, targetImgBase64 string) ([]
|
||||
func (h *Handler) DetectionTextOrder(orderImgBase64, targetImgBase64 string) ([]map[string]int, error) {
|
||||
// 暂时返回空结果
|
||||
return []map[string]int{}, nil
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user