package captcha import ( "bufio" "bytes" "encoding/base64" "fmt" "image" "math" "net/http" "os" "path/filepath" "strings" "sync" "anticaptcha/pkg/onnx" "anticaptcha/pkg/opencv" "github.com/disintegration/imaging" ) type Handler struct { modelPath string mu sync.RWMutex } // 模型配置:本地文件名 -> 远程文件名 var modelConfigs = map[string]string{ "CRNN_Math.onnx": "[AntiCAP]-CRNN_Math.onnx", "OCR.onnx": "[Dddd]-OCR.onnx", "Rotation-RotNetR.onnx": "[AntiCAP]-Rotation-RotNetR.onnx", "Siamese-ResNet18.onnx": "[AntiCAP]-Siamese-ResNet18.onnx", "CharSets.txt": "[Dddd]-CharSets.txt", } // 从 Gitea 仓库下载(公开仓库,无需认证) const modelBaseURL = "https://git.viaeon.com/admin/anticaptcha/raw/branch/main/models" func NewHandler(modelPath string) *Handler { h := &Handler{ modelPath: modelPath, } // 确保模型目录存在并下载缺失的模型 h.ensureModels() // 加载模型 h.loadModels() return h } // loadModels 加载所有模型 func (h *Handler) loadModels() { // 加载 OCR 模型 ocrPath := filepath.Join(h.modelPath, "OCR.onnx") if _, err := os.Stat(ocrPath); err == nil { if err := onnx.LoadModel("ocr", ocrPath); err != nil { fmt.Printf("警告: 加载 OCR 模型失败: %v\n", err) } else { fmt.Println("OCR 模型加载成功") } } // 加载 Math 模型 mathPath := filepath.Join(h.modelPath, "CRNN_Math.onnx") if _, err := os.Stat(mathPath); err == nil { if err := onnx.LoadModel("math", mathPath); err != nil { fmt.Printf("警告: 加载 Math 模型失败: %v\n", err) } else { fmt.Println("Math 模型加载成功") } } // 加载 Rotation 模型 rotatePath := filepath.Join(h.modelPath, "Rotation-RotNetR.onnx") if _, err := os.Stat(rotatePath); err == nil { if err := onnx.LoadModel("rotate", rotatePath); err != nil { fmt.Printf("警告: 加载 Rotation 模型失败: %v\n", err) } else { fmt.Println("Rotation 模型加载成功") } } // 加载 Siamese 模型 siamesePath := filepath.Join(h.modelPath, "Siamese-ResNet18.onnx") if _, err := os.Stat(siamesePath); err == nil { if err := onnx.LoadModel("siamese", siamesePath); err != nil { fmt.Printf("警告: 加载 Siamese 模型失败: %v\n", err) } else { fmt.Println("Siamese 模型加载成功") } } } // 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 文字识别 ===================== func (h *Handler) OCR(imageBase64 string) (string, error) { sess, ok := onnx.GetSession("ocr") if !ok { return "", fmt.Errorf("OCR 模型未加载") } // 解码图片 img, err := decodeBase64ToImage(imageBase64) if err != nil { return "", err } // 加载字符集 charset, err := h.loadCharset() if err != nil { return "", err } // 预处理 input, width, err := preprocessOCR(img) if err != nil { return "", err } // 推理 dims := []int64{1, 1, 64, int64(width)} output, err := sess.Run(input, dims) if err != nil { return "", err } // CTC 解码 return ctcDecode(output, charset), nil } // preprocessOCR OCR 预处理 func preprocessOCR(img image.Image) ([]float32, int, error) { // 调整高度为 64,保持宽高比 bounds := img.Bounds() width := bounds.Dx() height := bounds.Dy() newHeight := 64 newWidth := width * newHeight / height resized := imaging.Resize(img, newWidth, newHeight, imaging.Lanczos) gray := imaging.Grayscale(resized) // 转换为模型输入 pixels := make([]float32, newWidth*newHeight) for y := 0; y < newHeight; y++ { for x := 0; x < newWidth; x++ { c := gray.At(x, y) r, _, _, _ := c.RGBA() val := float32(r) / 65535.0 pixels[y*newWidth+x] = (val - 0.5) / 0.5 } } return pixels, newWidth, nil } // ctcDecode CTC 解码 func ctcDecode(output []float32, charset []string) string { if len(charset) == 0 { return "" } result := "" lastIdx := 0 numClasses := len(charset) timesteps := len(output) / numClasses for t := 0; t < timesteps; t++ { maxIdx := 0 maxProb := float32(-math.MaxFloat32) for c := 0; c < numClasses; c++ { idx := t * numClasses + c if idx < len(output) && output[idx] > maxProb { maxProb = output[idx] maxIdx = c } } if maxIdx != 0 && maxIdx != lastIdx && maxIdx < len(charset) { result += charset[maxIdx] } lastIdx = maxIdx } return result } // loadCharset 加载字符集 func (h *Handler) loadCharset() ([]string, error) { charsetPath := filepath.Join(h.modelPath, "CharSets.txt") file, err := os.Open(charsetPath) if err != nil { return nil, err } defer file.Close() charset := make([]string, 0, 6000) scanner := bufio.NewScanner(file) for scanner.Scan() { line := strings.TrimSpace(scanner.Text()) if line != "" { charset = append(charset, line) } } // 添加空白符作为第一个字符 result := make([]string, len(charset)+1) result[0] = "" copy(result[1:], charset) return result, scanner.Err() } // ===================== Math 数学计算 ===================== const mathChars = "0123456789+-*/÷×=?" func (h *Handler) Math(imageBase64 string) (string, error) { sess, ok := onnx.GetSession("math") if !ok { return "", fmt.Errorf("Math 模型未加载") } // 解码图片 img, err := decodeBase64ToImage(imageBase64) if err != nil { return "", err } // 预处理 input, err := preprocessMath(img) if err != nil { return "", err } // 推理 dims := []int64{1, 3, 70, 200} output, err := sess.Run(input, dims) if err != nil { return "", err } // 解码表达式 expr := decodeMath(output) if expr == "" { return "", fmt.Errorf("无法识别表达式") } // 计算结果 result, err := evalMathExpression(expr) if err != nil { return "", err } return fmt.Sprintf("%v", result), nil } // preprocessMath Math 预处理 func preprocessMath(img image.Image) ([]float32, error) { // 调整大小为 200x70,保持比例 resized := imaging.Resize(img, 200, 70, imaging.Lanczos) // 转换为 RGB rgb := imaging.Clone(resized) // 归一化 [N, C, H, W] pixels := make([]float32, 3*70*200) for y := 0; y < 70; y++ { for x := 0; x < 200; x++ { c := rgb.At(x, y) r, g, b, _ := c.RGBA() // CHW 格式,归一化 pixels[0*70*200+y*200+x] = (float32(r)/65535.0 - 0.5) / 0.5 pixels[1*70*200+y*200+x] = (float32(g)/65535.0 - 0.5) / 0.5 pixels[2*70*200+y*200+x] = (float32(b)/65535.0 - 0.5) / 0.5 } } return pixels, nil } // decodeMath 解码数学表达式 func decodeMath(output []float32) string { numChars := len(mathChars) + 1 timesteps := len(output) / numChars result := "" lastIdx := 0 for t := 0; t < timesteps; t++ { maxIdx := 0 maxProb := float32(-math.MaxFloat32) for c := 0; c < numChars; c++ { idx := t * numChars + c if idx < len(output) && output[idx] > maxProb { maxProb = output[idx] maxIdx = c } } if maxIdx != 0 && maxIdx != lastIdx { if maxIdx-1 < len(mathChars) { result += string(mathChars[maxIdx-1]) } } lastIdx = maxIdx } return result } // evalMathExpression 计算数学表达式 func evalMathExpression(expr string) (interface{}, error) { // 替换特殊符号 expr = strings.ReplaceAll(expr, "×", "*") expr = strings.ReplaceAll(expr, "÷", "/") expr = strings.ReplaceAll(expr, "?", "") expr = strings.ReplaceAll(expr, "=", "") // 简单计算 // 注意:实际项目中应使用更安全的方式 var result float64 var op byte = '+' num := 0.0 for i := 0; i < len(expr); i++ { c := expr[i] if c >= '0' && c <= '9' { num = num*10 + float64(c-'0') } else if c == '+' || c == '-' || c == '*' || c == '/' { switch op { case '+': result += num case '-': result -= num case '*': result *= num case '/': if num != 0 { result /= num } } op = c num = 0 } } // 处理最后一个数字 switch op { case '+': result += num case '-': result -= num case '*': result *= num case '/': if num != 0 { result /= num } } // 返回整数或浮点数 if result == float64(int(result)) { return int(result), nil } return result, nil } // ===================== 滑块匹配 ===================== func (h *Handler) SliderMatch(targetBase64, backgroundBase64 string) (int, error) { target, err := opencv.DecodeFromBase64(targetBase64) if err != nil { return 0, err } defer target.Free() background, err := opencv.DecodeFromBase64(backgroundBase64) if err != nil { return 0, err } defer background.Free() return opencv.SliderMatch(target, background) } func (h *Handler) SliderComparison(targetBase64, backgroundBase64 string) (int, error) { target, err := opencv.DecodeFromBase64(targetBase64) if err != nil { return 0, err } defer target.Free() background, err := opencv.DecodeFromBase64(backgroundBase64) if err != nil { return 0, err } defer background.Free() return opencv.SliderComparison(target, background) } // ===================== 图像相似度 ===================== func (h *Handler) CompareSimilarity(img1Base64, img2Base64 string) (float32, error) { // 使用 ONNX 模型 sess, ok := onnx.GetSession("siamese") if ok { return h.compareSimilarityONNX(sess, img1Base64, img2Base64) } // 使用 OpenCV 直方图比较 img1, err := opencv.DecodeFromBase64(img1Base64) if err != nil { return 0, err } defer img1.Free() img2, err := opencv.DecodeFromBase64(img2Base64) if err != nil { return 0, err } defer img2.Free() return opencv.CompareSimilarity(img1, img2) } func (h *Handler) compareSimilarityONNX(sess *onnx.Session, img1Base64, img2Base64 string) (float32, error) { img1, err := decodeBase64ToImage(img1Base64) if err != nil { return 0, err } img2, err := decodeBase64ToImage(img2Base64) if err != nil { return 0, err } // 预处理 input1, err := preprocessSiamese(img1) if err != nil { return 0, err } input2, err := preprocessSiamese(img2) if err != nil { return 0, err } // 推理 dims := []int64{1, 3, 105, 105} output, err := sess.RunDualInput(input1, dims, input2, dims) if err != nil { return 0, err } // 计算相似度 if len(output) >= 2 { emb1 := output[:len(output)/2] emb2 := output[len(output)/2:] // 欧氏距离 var dist float32 for i := 0; i < len(emb1); i++ { d := emb1[i] - emb2[i] dist += d * d } dist = float32(math.Sqrt(float64(dist))) // 相似度 similarity := 1.0 / (1.0 + dist) return similarity, nil } return 0, fmt.Errorf("输出格式错误") } func preprocessSiamese(img image.Image) ([]float32, error) { // 调整大小为 105x105 resized := imaging.Resize(img, 105, 105, imaging.Lanczos) rgb := imaging.Clone(resized) // ImageNet 归一化 mean := [3]float32{0.485, 0.456, 0.406} std := [3]float32{0.229, 0.224, 0.225} pixels := make([]float32, 3*105*105) for y := 0; y < 105; y++ { for x := 0; x < 105; x++ { c := rgb.At(x, y) r, g, b, _ := c.RGBA() pixels[0*105*105+y*105+x] = (float32(r)/65535.0 - mean[0]) / std[0] pixels[1*105*105+y*105+x] = (float32(g)/65535.0 - mean[1]) / std[1] pixels[2*105*105+y*105+x] = (float32(b)/65535.0 - mean[2]) / std[2] } } return pixels, nil } // ===================== 旋转检测 ===================== func (h *Handler) SingleRotate(imageBase64 string) (float32, error) { // 使用 ONNX 模型 sess, ok := onnx.GetSession("rotate") if ok { return h.singleRotateONNX(sess, imageBase64) } // 使用 OpenCV img, err := opencv.DecodeFromBase64(imageBase64) if err != nil { return 0, err } defer img.Free() return opencv.DetectRotation(img) } func (h *Handler) singleRotateONNX(sess *onnx.Session, imageBase64 string) (float32, error) { img, err := decodeBase64ToImage(imageBase64) if err != nil { return 0, err } // 预处理 input, err := preprocessRotation(img) if err != nil { return 0, err } // 推理 dims := []int64{1, 3, 224, 224} output, err := sess.Run(input, dims) if err != nil { return 0, err } // 找到最大概率的角度 maxIdx := 0 maxProb := float32(-math.MaxFloat32) for i := 0; i < len(output); i++ { if output[i] > maxProb { maxProb = output[i] maxIdx = i } } return float32(maxIdx), nil } func preprocessRotation(img image.Image) ([]float32, error) { // 调整大小为 224x224 resized := imaging.Resize(img, 224, 224, imaging.Lanczos) rgb := imaging.Clone(resized) // ImageNet 归一化 mean := [3]float32{0.485, 0.456, 0.406} std := [3]float32{0.229, 0.224, 0.225} pixels := make([]float32, 3*224*224) for y := 0; y < 224; y++ { for x := 0; x < 224; x++ { c := rgb.At(x, y) r, g, b, _ := c.RGBA() pixels[0*224*224+y*224+x] = (float32(r)/65535.0 - mean[0]) / std[0] pixels[1*224*224+y*224+x] = (float32(g)/65535.0 - mean[1]) / std[1] pixels[2*224*224+y*224+x] = (float32(b)/65535.0 - mean[2]) / std[2] } } return pixels, nil } func (h *Handler) DoubleRotate(insideBase64, outsideBase64 string) (float32, error) { inside, err := opencv.DecodeFromBase64(insideBase64) if err != nil { return 0, err } defer inside.Free() outside, err := opencv.DecodeFromBase64(outsideBase64) if err != nil { return 0, err } defer outside.Free() angleInside, err := opencv.DetectRotation(inside) if err != nil { return 0, err } angleOutside, err := opencv.DetectRotation(outside) if err != nil { return 0, err } return angleInside - angleOutside, nil } // ===================== 图标/文字检测 ===================== func (h *Handler) DetectionIcon(imageBase64 string) ([]map[string]int, error) { // 暂时返回空结果 return []map[string]int{}, nil } func (h *Handler) DetectionText(imageBase64 string) ([]map[string]int, error) { return []map[string]int{}, nil } func (h *Handler) DetectionIconOrder(orderImgBase64, targetImgBase64 string) ([]map[string]int, error) { return []map[string]int{}, nil } func (h *Handler) DetectionTextOrder(orderImgBase64, targetImgBase64 string) ([]map[string]int, error) { return []map[string]int{}, nil } // ===================== 工具函数 ===================== func decodeBase64ToImage(base64Str string) (image.Image, error) { data, err := base64.StdEncoding.DecodeString(base64Str) if err != nil { // 尝试去掉 data URL 前缀 if strings.Contains(base64Str, ",") { parts := strings.SplitN(base64Str, ",", 2) if len(parts) == 2 { data, err = base64.StdEncoding.DecodeString(parts[1]) if err != nil { return nil, err } } } else { return nil, err } } return imaging.Decode(bytes.NewReader(data)) }