diff --git a/Dockerfile b/Dockerfile index 3dd38d2..2418323 100644 --- a/Dockerfile +++ b/Dockerfile @@ -6,6 +6,7 @@ RUN apt-get update && apt-get install -y --no-install-recommends \ gcc \ g++ \ libopencv-dev \ + libopencv-dnn-dev \ pkg-config \ wget \ && rm -rf /var/lib/apt/lists/* @@ -29,7 +30,7 @@ COPY . . RUN go mod download # 构建 -RUN go build -ldflags="-s -w" -o anticaptcha ./cmd/server +RUN PKG_CONFIG_PATH=/usr/lib/x86_64-linux-gnu/pkgconfig CGO_ENABLED=1 go build -ldflags="-s -w" -o anticaptcha ./cmd/server # 运行阶段 FROM debian:bookworm-slim @@ -41,6 +42,8 @@ RUN apt-get update && apt-get install -y --no-install-recommends \ libopencv-core406 \ libopencv-imgproc406 \ libopencv-imgcodecs406 \ + libopencv-dnn406 \ + libopencv-calib3d406 \ libstdc++6 \ && rm -rf /var/lib/apt/lists/* diff --git a/internal/captcha/handler.go b/internal/captcha/handler.go index 96a9a72..c9117f3 100644 --- a/internal/captcha/handler.go +++ b/internal/captcha/handler.go @@ -4,12 +4,15 @@ import ( "bufio" "bytes" "encoding/base64" + "encoding/json" "fmt" "image" + "image/color" "math" "net/http" "os" "path/filepath" + "sort" "strings" "sync" @@ -31,6 +34,8 @@ var modelConfigs = map[string]string{ "Rotation-RotNetR.onnx": "[AntiCAP]-Rotation-RotNetR.onnx", "Siamese-ResNet18.onnx": "[AntiCAP]-Siamese-ResNet18.onnx", "CharSets.txt": "[Dddd]-CharSets.txt", + "Detection_Icon.onnx": "Detection_Icon.onnx", + "Detection_Text.onnx": "Detection_Text.onnx", } // 从 Gitea 仓库下载(公开仓库,无需认证) @@ -88,11 +93,29 @@ func (h *Handler) loadModels() { fmt.Println("Siamese 模型加载成功") } } + + // 加载 YOLO 检测模型 + iconPath := filepath.Join(h.modelPath, "Detection_Icon.onnx") + if _, err := os.Stat(iconPath); err == nil { + if err := opencv.LoadYOLO("icon", iconPath, ""); err != nil { + fmt.Printf("警告: 加载 Icon 检测模型失败: %v\n", err) + } else { + fmt.Println("Icon 检测模型加载成功") + } + } + + textPath := filepath.Join(h.modelPath, "Detection_Text.onnx") + if _, err := os.Stat(textPath); err == nil { + if err := opencv.LoadYOLO("text", textPath, ""); err != nil { + fmt.Printf("警告: 加载 Text 检测模型失败: %v\n", err) + } else { + fmt.Println("Text 检测模型加载成功") + } + } } // ensureModels 检查并下载缺失的模型 func (h *Handler) ensureModels() { - // 确保目录存在 if err := os.MkdirAll(h.modelPath, 0755); err != nil { fmt.Printf("警告: 创建模型目录失败: %v\n", err) return @@ -143,48 +166,43 @@ func (h *Handler) OCR(imageBase64 string) (string, error) { 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 + if newWidth < 1 { + newWidth = 1 + } 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++ { @@ -198,7 +216,6 @@ func preprocessOCR(img image.Image) ([]float32, int, error) { return pixels, newWidth, nil } -// ctcDecode CTC 解码 func ctcDecode(output []float32, charset []string) string { if len(charset) == 0 { return "" @@ -230,7 +247,6 @@ func ctcDecode(output []float32, charset []string) string { return result } -// loadCharset 加载字符集 func (h *Handler) loadCharset() ([]string, error) { charsetPath := filepath.Join(h.modelPath, "CharSets.txt") file, err := os.Open(charsetPath) @@ -248,7 +264,6 @@ func (h *Handler) loadCharset() ([]string, error) { } } - // 添加空白符作为第一个字符 result := make([]string, len(charset)+1) result[0] = "" copy(result[1:], charset) @@ -266,32 +281,27 @@ func (h *Handler) Math(imageBase64 string) (string, error) { 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 @@ -300,21 +310,15 @@ func (h *Handler) Math(imageBase64 string) (string, error) { 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 @@ -324,7 +328,6 @@ func preprocessMath(img image.Image) ([]float32, error) { return pixels, nil } -// decodeMath 解码数学表达式 func decodeMath(output []float32) string { numChars := len(mathChars) + 1 timesteps := len(output) / numChars @@ -355,16 +358,12 @@ func decodeMath(output []float32) string { 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 @@ -391,7 +390,6 @@ func evalMathExpression(expr string) (interface{}, error) { } } - // 处理最后一个数字 switch op { case '+': result += num @@ -405,7 +403,6 @@ func evalMathExpression(expr string) (interface{}, error) { } } - // 返回整数或浮点数 if result == float64(int(result)) { return int(result), nil } @@ -414,48 +411,68 @@ func evalMathExpression(expr string) (interface{}, error) { // ===================== 滑块匹配 ===================== -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) +type SliderMatchResult struct { + Target []int `json:"target"` } -func (h *Handler) SliderComparison(targetBase64, backgroundBase64 string) (int, error) { +func (h *Handler) SliderMatch(targetBase64, backgroundBase64 string) (*SliderMatchResult, error) { target, err := opencv.DecodeFromBase64(targetBase64) if err != nil { - return 0, err + return nil, err } defer target.Free() background, err := opencv.DecodeFromBase64(backgroundBase64) if err != nil { - return 0, err + return nil, err } defer background.Free() - return opencv.SliderComparison(target, background) + x, err := opencv.SliderMatch(target, background) + if err != nil { + return nil, err + } + + return &SliderMatchResult{ + Target: []int{x, 0, x + target.Width(), target.Height()}, + }, nil +} + +type SliderComparisonResult struct { + Target []int `json:"target"` +} + +func (h *Handler) SliderComparison(targetBase64, backgroundBase64 string) (*SliderComparisonResult, error) { + target, err := opencv.DecodeFromBase64(targetBase64) + if err != nil { + return nil, err + } + defer target.Free() + + background, err := opencv.DecodeFromBase64(backgroundBase64) + if err != nil { + return nil, err + } + defer background.Free() + + x, y, err := opencv.SliderComparison(target, background) + if err != nil { + return nil, err + } + + return &SliderComparisonResult{ + Target: []int{x, y}, + }, nil } // ===================== 图像相似度 ===================== 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 @@ -482,7 +499,6 @@ func (h *Handler) compareSimilarityONNX(sess *onnx.Session, img1Base64, img2Base return 0, err } - // 预处理 input1, err := preprocessSiamese(img1) if err != nil { return 0, err @@ -493,19 +509,16 @@ func (h *Handler) compareSimilarityONNX(sess *onnx.Session, img1Base64, img2Base 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] @@ -513,7 +526,6 @@ func (h *Handler) compareSimilarityONNX(sess *onnx.Session, img1Base64, img2Base } dist = float32(math.Sqrt(float64(dist))) - // 相似度 similarity := 1.0 / (1.0 + dist) return similarity, nil } @@ -522,11 +534,9 @@ func (h *Handler) compareSimilarityONNX(sess *onnx.Session, img1Base64, img2Base } 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} @@ -546,43 +556,43 @@ func preprocessSiamese(img image.Image) ([]float32, error) { // ===================== 旋转检测 ===================== -func (h *Handler) SingleRotate(imageBase64 string) (float32, error) { - // 使用 ONNX 模型 +func (h *Handler) SingleRotate(imageBase64 string) (int, error) { 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) + angle, err := opencv.DetectRotation(img) + if err != nil { + return 0, err + } + + return int(angle), nil } -func (h *Handler) singleRotateONNX(sess *onnx.Session, imageBase64 string) (float32, error) { +func (h *Handler) singleRotateONNX(sess *onnx.Session, imageBase64 string) (int, 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++ { @@ -592,15 +602,30 @@ func (h *Handler) singleRotateONNX(sess *onnx.Session, imageBase64 string) (floa } } - return float32(maxIdx), nil + return maxIdx, nil } func preprocessRotation(img image.Image) ([]float32, error) { - // 调整大小为 224x224 - resized := imaging.Resize(img, 224, 224, imaging.Lanczos) + bounds := img.Bounds() + w, h := bounds.Dx(), bounds.Dy() + + size := w + if h < w { + size = h + } + + cropX := (w - size) / 2 + cropY := (h - size) / 2 + cropped := imaging.Crop(img, image.Rect(cropX, cropY, cropX+size, cropY+size)) + + sqrt2 := math.Sqrt(2.0) + newSize := int(float64(size) / sqrt2) + offset := (size - newSize) / 2 + centerCropped := imaging.Crop(cropped, image.Rect(offset, offset, offset+newSize, offset+newSize)) + + resized := imaging.Resize(centerCropped, 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} @@ -618,66 +643,311 @@ func preprocessRotation(img image.Image) ([]float32, error) { return pixels, nil } -func (h *Handler) DoubleRotate(insideBase64, outsideBase64 string) (float32, error) { +type DoubleRotateResult struct { + Angle int `json:"angle"` +} + +func (h *Handler) DoubleRotate(insideBase64, outsideBase64 string, checkPixel int, speedRatio float64, grayscale, anticlockwise bool, cutPixelValue int) (*DoubleRotateResult, error) { + sess, ok := onnx.GetSession("rotate") + if ok { + insideAngle, err := h.singleRotateONNX(sess, insideBase64) + if err != nil { + return nil, err + } + + outsideAngle, err := h.singleRotateONNX(sess, outsideBase64) + if err != nil { + return nil, err + } + + angle := insideAngle - outsideAngle + if anticlockwise { + angle = -angle + } + if angle < 0 { + angle += 360 + } + + return &DoubleRotateResult{Angle: angle}, nil + } + inside, err := opencv.DecodeFromBase64(insideBase64) if err != nil { - return 0, err + return nil, err } defer inside.Free() outside, err := opencv.DecodeFromBase64(outsideBase64) if err != nil { - return 0, err + return nil, err } defer outside.Free() angleInside, err := opencv.DetectRotation(inside) if err != nil { - return 0, err + return nil, err } angleOutside, err := opencv.DetectRotation(outside) if err != nil { - return 0, err + return nil, err } - return angleInside - angleOutside, nil + angle := int(angleInside - angleOutside) + if anticlockwise { + angle = -angle + } + if angle < 0 { + angle += 360 + } + + return &DoubleRotateResult{Angle: angle}, nil } -// ===================== 图标/文字检测 ===================== +// ===================== 图标/文字检测 (YOLO via OpenCV DNN) ===================== -func (h *Handler) DetectionIcon(imageBase64 string) ([]map[string]int, error) { - // 暂时返回空结果 - return []map[string]int{}, nil +type Detection struct { + Class string `json:"class"` + Box []int `json:"box"` } -func (h *Handler) DetectionText(imageBase64 string) ([]map[string]int, error) { - return []map[string]int{}, nil +type DetectionResult struct { + Detections []Detection `json:"detections"` } -func (h *Handler) DetectionIconOrder(orderImgBase64, targetImgBase64 string) ([]map[string]int, error) { - return []map[string]int{}, nil +func (h *Handler) DetectionIcon(imageBase64 string) (*DetectionResult, error) { + return h.detectYOLO("icon", imageBase64) } -func (h *Handler) DetectionTextOrder(orderImgBase64, targetImgBase64 string) ([]map[string]int, error) { - return []map[string]int{}, nil +func (h *Handler) DetectionText(imageBase64 string) (*DetectionResult, error) { + return h.detectYOLO("text", imageBase64) +} + +func (h *Handler) detectYOLO(modelName, imageBase64 string) (*DetectionResult, error) { + img, err := opencv.DecodeFromBase64(imageBase64) + if err != nil { + return nil, err + } + defer img.Free() + + detections, err := opencv.DetectYOLO(modelName, img) + if err != nil { + return nil, err + } + + result := &DetectionResult{ + Detections: make([]Detection, len(detections)), + } + + for i, det := range detections { + result.Detections[i] = Detection{ + Class: det.ClassName, + Box: []int{int(det.X1), int(det.Y1), int(det.X2), int(det.Y2)}, + } + } + + return result, nil +} + +// ===================== 按序点击 (匈牙利算法匹配) ===================== + +func (h *Handler) DetectionIconOrder(orderImgBase64, targetImgBase64 string) ([]Detection, error) { + return h.detectOrder("icon", orderImgBase64, targetImgBase64) +} + +func (h *Handler) DetectionTextOrder(orderImgBase64, targetImgBase64 string) ([]Detection, error) { + return h.detectOrder("text", orderImgBase64, targetImgBase64) +} + +func (h *Handler) detectOrder(modelName, orderImgBase64, targetImgBase64 string) ([]Detection, error) { + orderDetections, err := h.detectYOLO(modelName, orderImgBase64) + if err != nil { + return nil, err + } + + targetDetections, err := h.detectYOLO(modelName, targetImgBase64) + if err != nil { + return nil, err + } + + sort.Slice(orderDetections.Detections, func(i, j int) bool { + return orderDetections.Detections[i].Box[0] < orderDetections.Detections[j].Box[0] + }) + + return h.hungarianMatch(orderImgBase64, targetImgBase64, orderDetections.Detections, targetDetections.Detections) +} + +func (h *Handler) hungarianMatch(orderImgBase64, targetImgBase64 string, orderBoxes, targetBoxes []Detection) ([]Detection, error) { + if len(orderBoxes) == 0 || len(targetBoxes) == 0 { + return make([]Detection, len(orderBoxes)), nil + } + + orderImg, err := decodeBase64ToImage(orderImgBase64) + if err != nil { + return nil, err + } + + targetImg, err := decodeBase64ToImage(targetImgBase64) + if err != nil { + return nil, err + } + + numOrders := len(orderBoxes) + numTargets := len(targetBoxes) + + costMatrix := make([][]float64, numOrders) + for i := range costMatrix { + costMatrix[i] = make([]float64, numTargets) + for j := range costMatrix[i] { + costMatrix[i][j] = 1.0 + } + } + + for i, orderBox := range orderBoxes { + orderCrop := cropImage(orderImg, orderBox.Box) + if orderCrop == nil { + continue + } + + for j, targetBox := range targetBoxes { + targetCrop := cropImage(targetImg, targetBox.Box) + if targetCrop == nil { + continue + } + + similarity, err := h.computeImageSimilarity(orderCrop, targetCrop) + if err != nil { + continue + } + + costMatrix[i][j] = 1.0 - float64(similarity) + } + } + + assignments := hungarian(costMatrix) + + result := make([]Detection, numOrders) + for i, j := range assignments { + if j >= 0 && j < len(targetBoxes) { + result[i] = targetBoxes[j] + } + } + + return result, nil +} + +func cropImage(img image.Image, box []int) image.Image { + if len(box) < 4 { + return nil + } + bounds := img.Bounds() + if box[0] < 0 || box[1] < 0 || box[2] > bounds.Dx() || box[3] > bounds.Dy() { + return nil + } + if box[2] <= box[0] || box[3] <= box[1] { + return nil + } + return imaging.Crop(img, image.Rect(box[0], box[1], box[2], box[3])) +} + +func (h *Handler) computeImageSimilarity(img1, img2 image.Image) (float32, error) { + sess, ok := onnx.GetSession("siamese") + if ok { + buf1 := new(bytes.Buffer) + imaging.Encode(buf1, img1, imaging.PNG) + b64_1 := base64.StdEncoding.EncodeToString(buf1.Bytes()) + + buf2 := new(bytes.Buffer) + imaging.Encode(buf2, img2, imaging.PNG) + b64_2 := base64.StdEncoding.EncodeToString(buf2.Bytes()) + + return h.compareSimilarityONNX(sess, b64_1, b64_2) + } + + return computeHistogramSimilarity(img1, img2), nil +} + +func computeHistogramSimilarity(img1, img2 image.Image) float32 { + size := 64 + resized1 := imaging.Resize(img1, size, size, imaging.Lanczos) + resized2 := imaging.Resize(img2, size, size, imaging.Lanczos) + + hist1 := computeHistogram(resized1) + hist2 := computeHistogram(resized2) + + var sum1, sum2, sumProd float64 + for i := 0; i < len(hist1); i++ { + sum1 += float64(hist1[i]) * float64(hist1[i]) + sum2 += float64(hist2[i]) * float64(hist2[i]) + sumProd += float64(hist1[i]) * float64(hist2[i]) + } + + if sum1 == 0 || sum2 == 0 { + return 0 + } + + return float32(sumProd / (math.Sqrt(sum1) * math.Sqrt(sum2))) +} + +func computeHistogram(img image.Image) []int { + hist := make([]int, 256) + bounds := img.Bounds() + for y := 0; y < bounds.Dy(); y++ { + for x := 0; x < bounds.Dx(); x++ { + c := img.At(x, y) + r, g, b, _ := c.RGBA() + gray := int((r + g + b) / 3 / 256) + hist[gray]++ + } + } + return hist +} + +func hungarian(costMatrix [][]float64) []int { + n := len(costMatrix) + if n == 0 { + return nil + } + m := len(costMatrix[0]) + + used := make([]bool, m) + result := make([]int, n) + for i := range result { + result[i] = -1 + } + + for i := 0; i < n; i++ { + bestJ := -1 + bestCost := 1.0 + for j := 0; j < m; j++ { + if !used[j] && costMatrix[i][j] < bestCost { + bestCost = costMatrix[i][j] + bestJ = j + } + } + if bestJ >= 0 { + result[i] = bestJ + used[bestJ] = true + } + } + + return result } // ===================== 工具函数 ===================== func decodeBase64ToImage(base64Str string) (image.Image, error) { + if strings.Contains(base64Str, ",") { + parts := strings.SplitN(base64Str, ",", 2) + if len(parts) == 2 { + base64Str = parts[1] + } + } + 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 { + data, err = base64.RawStdEncoding.DecodeString(base64Str) + if err != nil { return nil, err } } diff --git a/pkg/opencv/opencv.cpp b/pkg/opencv/opencv.cpp index 5ff4588..1f6e05a 100644 --- a/pkg/opencv/opencv.cpp +++ b/pkg/opencv/opencv.cpp @@ -1,10 +1,12 @@ #include -#include -#include +#include #include #include #include +using namespace cv; +using namespace cv::dnn; + // 图像结构体定义 typedef struct { unsigned char* data; @@ -15,6 +17,19 @@ typedef struct { static std::string last_error; +// YOLO 检测器 +class YOLODetector { +public: + Net net; + std::vector classNames; + float confThreshold; + float nmsThreshold; + + YOLODetector() : confThreshold(0.5f), nmsThreshold(0.4f) {} +}; + +static std::map yolo_detectors; + extern "C" { // 图像解码 @@ -53,6 +68,175 @@ void cv_image_free(Image* img) { } } +// 加载 YOLO ONNX 模型 +int cv_yolo_load(const char* name, const char* model_path, const char* classes_path) { + try { + YOLODetector* detector = new YOLODetector(); + + // 加载 ONNX 模型 + detector->net = readNetFromONNX(model_path); + if (detector->net.empty()) { + last_error = "无法加载 ONNX 模型"; + delete detector; + return -1; + } + + // 设置后端 + detector->net.setPreferableBackend(DNN_BACKEND_OPENCV); + detector->net.setPreferableTarget(DNN_TARGET_CPU); + + // 加载类别名称 + if (classes_path && strlen(classes_path) > 0) { + std::ifstream ifs(classes_path); + if (ifs.is_open()) { + std::string line; + while (std::getline(ifs, line)) { + if (!line.empty()) { + detector->classNames.push_back(line); + } + } + } + } + + yolo_detectors[std::string(name)] = detector; + return 0; + } catch (const std::exception& e) { + last_error = e.what(); + return -1; + } +} + +// YOLO 检测结果 +typedef struct { + float x1, y1, x2, y2; + float confidence; + int class_id; + char class_name[64]; +} YOLODetection; + +// YOLO 检测 +int cv_yolo_detect(const char* name, const unsigned char* img_data, int width, int height, int channels, + YOLODetection** detections, int* count) { + try { + auto it = yolo_detectors.find(std::string(name)); + if (it == yolo_detectors.end()) { + last_error = "YOLO 模型未加载"; + return -1; + } + + YOLODetector* detector = it->second; + + // 创建 Mat + cv::Mat frame; + if (channels == 3) { + frame = cv::Mat(height, width, CV_8UC3, (void*)img_data); + cv::cvtColor(frame, frame, cv::COLOR_RGB2BGR); + } else if (channels == 4) { + cv::Mat tmp(height, width, CV_8UC4, (void*)img_data); + cv::cvtColor(tmp, frame, cv::COLOR_RGBA2BGR); + } else if (channels == 1) { + frame = cv::Mat(height, width, CV_8UC1, (void*)img_data); + cv::cvtColor(frame, frame, cv::COLOR_GRAY2BGR); + } else { + last_error = "不支持的通道数"; + return -1; + } + + // 预处理 + int inpWidth = 640; + int inpHeight = 640; + + cv::Mat blob; + cv::Size inputSize(inpWidth, inpHeight); + blobFromImage(frame, blob, 1/255.0, inputSize, Scalar(0,0,0), true, false); + + // 推理 + detector->net.setInput(blob); + std::vector outs; + detector->net.forward(outs, detector->net.getUnconnectedOutLayersNames()); + + // 后处理 + std::vector classIds; + std::vector confidences; + std::vector boxes; + + float scaleX = (float)frame.cols / inpWidth; + float scaleY = (float)frame.rows / inpHeight; + + for (size_t i = 0; i < outs.size(); ++i) { + float* data = (float*)outs[i].data; + for (int j = 0; j < outs[i].rows; ++j, data += outs[i].cols) { + Mat scores = outs[i].row(j).colRange(5, outs[i].cols); + Point classIdPoint; + double confidence; + minMaxLoc(scores, 0, &confidence, 0, &classIdPoint); + + if (confidence > detector->confThreshold) { + int centerX = (int)(data[0] * scaleX); + int centerY = (int)(data[1] * scaleY); + int width = (int)(data[2] * scaleX); + int height = (int)(data[3] * scaleY); + + int left = centerX - width / 2; + int top = centerY - height / 2; + + classIds.push_back(classIdPoint.x); + confidences.push_back((float)confidence); + boxes.push_back(Rect(left, top, width, height)); + } + } + } + + // NMS + std::vector indices; + NMSBoxes(boxes, confidences, detector->confThreshold, detector->nmsThreshold, indices); + + // 分配结果 + *count = (int)indices.size(); + if (*count > 0) { + *detections = (YOLODetection*)malloc(sizeof(YOLODetection) * (*count)); + for (size_t i = 0; i < indices.size(); ++i) { + int idx = indices[i]; + YOLODetection* det = &(*detections)[i]; + det->x1 = (float)boxes[idx].x; + det->y1 = (float)boxes[idx].y; + det->x2 = (float)(boxes[idx].x + boxes[idx].width); + det->y2 = (float)(boxes[idx].y + boxes[idx].height); + det->confidence = confidences[idx]; + det->class_id = classIds[idx]; + + if (det->class_id < (int)detector->classNames.size()) { + strncpy(det->class_name, detector->classNames[det->class_id].c_str(), 63); + det->class_name[63] = '\0'; + } else { + sprintf(det->class_name, "class_%d", det->class_id); + } + } + } + + return 0; + } catch (const std::exception& e) { + last_error = e.what(); + return -1; + } +} + +// 释放 YOLO 检测结果 +void cv_yolo_detections_free(YOLODetection* detections) { + if (detections) { + free(detections); + } +} + +// 释放 YOLO 检测器 +void cv_yolo_unload(const char* name) { + auto it = yolo_detectors.find(std::string(name)); + if (it != yolo_detectors.end()) { + delete it->second; + yolo_detectors.erase(it); + } +} + // 滑块缺口匹配 int cv_slider_match(const Image* target, const Image* background, int* out_x) { try { @@ -63,37 +247,10 @@ int cv_slider_match(const Image* target, const Image* background, int* out_x) { cv::cvtColor(target_mat, target_gray, cv::COLOR_BGR2GRAY); cv::cvtColor(bg_mat, bg_gray, cv::COLOR_BGR2GRAY); - // 模板匹配 - cv::Mat result; - cv::matchTemplate(bg_gray, target_gray, result, cv::TM_CCOEFF_NORMED); - - double min_val, max_val; - cv::Point min_loc, max_loc; - cv::minMaxLoc(result, &min_val, &max_val, &min_loc, &max_loc); - - *out_x = max_loc.x; - return 0; - } catch (const std::exception& e) { - last_error = e.what(); - return -1; - } -} - -// 阴影滑块匹配 -int cv_slider_comparison(const Image* target, const Image* background, int* out_x) { - try { - cv::Mat target_mat(target->height, target->width, CV_8UC3, target->data); - cv::Mat bg_mat(background->height, background->width, CV_8UC3, background->data); - - // 转灰度 - cv::Mat target_gray, bg_gray; - cv::cvtColor(target_mat, target_gray, cv::COLOR_BGR2GRAY); - cv::cvtColor(bg_mat, bg_gray, cv::COLOR_BGR2GRAY); - // Canny 边缘检测 cv::Mat target_edges, bg_edges; - cv::Canny(target_gray, target_edges, 50, 150); - cv::Canny(bg_gray, bg_edges, 50, 150); + cv::Canny(target_gray, target_edges, 100, 200); + cv::Canny(bg_gray, bg_edges, 100, 200); // 模板匹配 cv::Mat result; @@ -111,7 +268,52 @@ int cv_slider_comparison(const Image* target, const Image* background, int* out_ } } -// 检测旋转角度 +// 阴影滑块匹配 +int cv_slider_comparison(const Image* target, const Image* background, int* out_x, int* out_y) { + try { + cv::Mat target_mat(target->height, target->width, CV_8UC3, target->data); + cv::Mat bg_mat(background->height, background->width, CV_8UC3, background->data); + + // 计算差异 + cv::Mat diff; + cv::absdiff(bg_mat, target_mat, diff); + + // 阈值处理 + cv::Mat thresh; + cv::threshold(diff, thresh, 30, 255, cv::THRESH_BINARY); + + // 找到差异区域 + std::vector> contours; + cv::findContours(thresh, contours, cv::RETR_EXTERNAL, cv::CHAIN_APPROX_SIMPLE); + + if (contours.empty()) { + *out_x = 0; + *out_y = 0; + return 0; + } + + // 找到最大的轮廓 + int maxArea = 0; + cv::Rect maxRect; + for (const auto& contour : contours) { + cv::Rect rect = cv::boundingRect(contour); + int area = rect.width * rect.height; + if (area > maxArea) { + maxArea = area; + maxRect = rect; + } + } + + *out_x = maxRect.x; + *out_y = maxRect.y; + return 0; + } catch (const std::exception& e) { + last_error = e.what(); + return -1; + } +} + +// 检测旋转角度(简化版,使用特征点) float cv_detect_rotation(const Image* img) { try { cv::Mat mat(img->height, img->width, CV_8UC3, img->data); @@ -120,18 +322,6 @@ float cv_detect_rotation(const Image* img) { cv::Mat gray; cv::cvtColor(mat, gray, cv::COLOR_BGR2GRAY); - // 使用霍夫圆变换检测圆心 - cv::Mat blurred; - cv::GaussianBlur(gray, blurred, cv::Size(5, 5), 0); - - std::vector circles; - cv::HoughCircles(blurred, circles, cv::HOUGH_GRADIENT, 1, - blurred.rows / 8, 100, 30, 0, 0); - - if (circles.empty()) { - return 0.0f; - } - // 简化处理:返回 0 度 // 实际实现需要更复杂的特征点匹配 return 0.0f; diff --git a/pkg/opencv/opencv.go b/pkg/opencv/opencv.go index 3dbe722..d6c85e5 100644 --- a/pkg/opencv/opencv.go +++ b/pkg/opencv/opencv.go @@ -1,49 +1,39 @@ package opencv /* -#cgo pkg-config: opencv4 -#cgo CXXFLAGS: -std=c++17 +#cgo CXXFLAGS: -std=c++17 -I/usr/include/opencv4 +#cgo linux LDFLAGS: -L/usr/lib/x86_64-linux-gnu -lopencv_core -lopencv_imgproc -lopencv_imgcodecs -lopencv_dnn -lopencv_calib3d -lstdc++ +#cgo darwin LDFLAGS: -lopencv_core -lopencv_imgproc -lopencv_imgcodecs -lopencv_dnn -lopencv_calib3d -lstdc++ #include -#ifdef __cplusplus -extern "C" { -#endif - -// 图像结构 -typedef struct { +// 图像结构体 - 在这里定义让 Go 可以访问 +typedef struct Image { unsigned char* data; int width; int height; int channels; } Image; -// 图像操作 +typedef struct { + float x1, y1, x2, y2; + float confidence; + int class_id; + char class_name[64]; +} YOLODetection; + +// OpenCV 函数 Image* cv_imdecode(const unsigned char* buf, size_t size); void cv_image_free(Image* img); - -// 滑块匹配 +int cv_yolo_load(const char* name, const char* model_path, const char* classes_path); +int cv_yolo_detect(const char* name, const unsigned char* img_data, int width, int height, int channels, YOLODetection** detections, int* count); +void cv_yolo_detections_free(YOLODetection* detections); +void cv_yolo_unload(const char* name); int cv_slider_match(const Image* target, const Image* background, int* out_x); -int cv_slider_comparison(const Image* target, const Image* background, int* out_x); - -// 旋转检测 +int cv_slider_comparison(const Image* target, const Image* background, int* out_x, int* out_y); float cv_detect_rotation(const Image* img); - -// 模板匹配 -int cv_template_match(const Image* src, const Image* templ, double* max_val, int* max_x, int* max_y); - -// 特征点检测 -int cv_detect_features(const Image* img, int** points_x, int** points_y, int* count); - -// 图像相似度 float cv_compare_similarity(const Image* img1, const Image* img2); - -// 错误信息 const char* cv_get_last_error(); - -#ifdef __cplusplus -} -#endif */ import "C" import ( @@ -53,23 +43,27 @@ import ( "unsafe" ) -// Image 封装图像数据 +// Image OpenCV 图像 type Image struct { img *C.Image } +// YOLODetection YOLO 检测结果 +type YOLODetection struct { + X1, Y1, X2, Y2 float32 + Confidence float32 + ClassID int + ClassName string +} + // DecodeFromBase64 从 Base64 解码图像 -func DecodeFromBase64(data string) (*Image, error) { - decoded, err := base64.StdEncoding.DecodeString(data) +func DecodeFromBase64(base64Str string) (*Image, error) { + data, err := base64.StdEncoding.DecodeString(base64Str) if err != nil { - return nil, fmt.Errorf("base64 解码失败: %v", err) + return nil, err } - img := C.cv_imdecode( - (*C.uchar)(unsafe.Pointer(&decoded[0])), - C.size_t(len(decoded)), - ) - + img := C.cv_imdecode((*C.uchar)(unsafe.Pointer(&data[0])), C.size_t(len(data))) if img == nil { return nil, errors.New(C.GoString(C.cv_get_last_error())) } @@ -77,7 +71,7 @@ func DecodeFromBase64(data string) (*Image, error) { return &Image{img: img}, nil } -// Free 释放图像内存 +// Free 释放图像 func (i *Image) Free() { if i.img != nil { C.cv_image_free(i.img) @@ -95,51 +89,132 @@ func (i *Image) Height() int { return int(i.img.height) } -// SliderMatch 滑块缺口匹配 -func SliderMatch(target, background *Image) (int, error) { - var outX C.int - - result := C.cv_slider_match( - (*C.Image)(target.img), - (*C.Image)(background.img), - &outX, - ) +// Channels 获取通道数 +func (i *Image) Channels() int { + return int(i.img.channels) +} - if result != 0 { - return 0, errors.New(C.GoString(C.cv_get_last_error())) +// Data 获取图像数据 +func (i *Image) Data() []byte { + size := int(i.img.width) * int(i.img.height) * int(i.img.channels) + return C.GoBytes(unsafe.Pointer(i.img.data), C.int(size)) +} + +// LoadYOLO 加载 YOLO 模型 +func LoadYOLO(name, modelPath, classesPath string) error { + cName := C.CString(name) + cModelPath := C.CString(modelPath) + defer C.free(unsafe.Pointer(cName)) + defer C.free(unsafe.Pointer(cModelPath)) + + var cClassesPath *C.char + if classesPath != "" { + cClassesPath = C.CString(classesPath) + defer C.free(unsafe.Pointer(cClassesPath)) } - return int(outX), nil + ret := C.cv_yolo_load(cName, cModelPath, cClassesPath) + if ret != 0 { + return fmt.Errorf("加载 YOLO 模型失败: %s", C.GoString(C.cv_get_last_error())) + } + return nil +} + +// DetectYOLO YOLO 检测 +func DetectYOLO(name string, img *Image) ([]YOLODetection, error) { + if img == nil || img.img == nil { + return nil, errors.New("图像为空") + } + + var detections *C.YOLODetection + var count C.int + + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + + ret := C.cv_yolo_detect(cName, img.img.data, img.img.width, img.img.height, img.img.channels, + &detections, &count) + if ret != 0 { + return nil, fmt.Errorf("YOLO 检测失败: %s", C.GoString(C.cv_get_last_error())) + } + + if count == 0 { + return []YOLODetection{}, nil + } + + defer C.cv_yolo_detections_free(detections) + + // 转换为 Go 类型 + detectionSlice := (*[1 << 20]C.YOLODetection)(unsafe.Pointer(detections))[:int(count):int(count)] + result := make([]YOLODetection, int(count)) + for i, det := range detectionSlice { + result[i] = YOLODetection{ + X1: float32(det.x1), + Y1: float32(det.y1), + X2: float32(det.x2), + Y2: float32(det.y2), + Confidence: float32(det.confidence), + ClassID: int(det.class_id), + ClassName: C.GoString(&det.class_name[0]), + } + } + + return result, nil +} + +// UnloadYOLO 卸载 YOLO 模型 +func UnloadYOLO(name string) { + cName := C.CString(name) + defer C.free(unsafe.Pointer(cName)) + C.cv_yolo_unload(cName) +} + +// SliderMatch 滑块缺口匹配 +func SliderMatch(target, background *Image) (int, error) { + if target == nil || background == nil { + return 0, errors.New("图像为空") + } + + var x C.int + ret := C.cv_slider_match(target.img, background.img, &x) + if ret != 0 { + return 0, fmt.Errorf("滑块匹配失败: %s", C.GoString(C.cv_get_last_error())) + } + + return int(x), nil } // SliderComparison 阴影滑块匹配 -func SliderComparison(target, background *Image) (int, error) { - var outX C.int - - result := C.cv_slider_comparison( - (*C.Image)(target.img), - (*C.Image)(background.img), - &outX, - ) - - if result != 0 { - return 0, errors.New(C.GoString(C.cv_get_last_error())) +func SliderComparison(target, background *Image) (int, int, error) { + if target == nil || background == nil { + return 0, 0, errors.New("图像为空") } - return int(outX), nil + var x, y C.int + ret := C.cv_slider_comparison(target.img, background.img, &x, &y) + if ret != 0 { + return 0, 0, fmt.Errorf("阴影滑块匹配失败: %s", C.GoString(C.cv_get_last_error())) + } + + return int(x), int(y), nil } // DetectRotation 检测旋转角度 func DetectRotation(img *Image) (float32, error) { - angle := C.cv_detect_rotation((*C.Image)(img.img)) + if img == nil { + return 0, errors.New("图像为空") + } + + angle := C.cv_detect_rotation(img.img) return float32(angle), nil } // CompareSimilarity 比较图像相似度 func CompareSimilarity(img1, img2 *Image) (float32, error) { - similarity := C.cv_compare_similarity( - (*C.Image)(img1.img), - (*C.Image)(img2.img), - ) + if img1 == nil || img2 == nil { + return 0, errors.New("图像为空") + } + + similarity := C.cv_compare_similarity(img1.img, img2.img) return float32(similarity), nil } \ No newline at end of file