package chat import ( "bytes" "context" "encoding/binary" "errors" "fmt" "hash/crc32" "image" "image/color" "image/draw" "image/png" "net/http" "net/http/httptest" "strings" "sync/atomic" "testing" "time" "go.uber.org/zap" "go.uber.org/zap/zaptest/observer" ) func newOCRTestRepository(client *http.Client, logger *zap.Logger, sleep func(context.Context, time.Duration) error) *Repository { return NewRepository(nil, nil, nil, withOCRHTTPClient(client), withOCRGate(newOCRRequestGate(ocrRequestConcurrency, 0)), WithLogger(logger), withOCRSleep(sleep), ) } func TestDoPaddleRequestRetries429AndHonorsRetryAfter(t *testing.T) { var requests atomic.Int32 server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { if requests.Add(1) == 1 { w.Header().Set("Retry-After", "2") w.WriteHeader(http.StatusTooManyRequests) _, _ = w.Write([]byte(`{"code":"rate_limit","message":"busy"}`)) return } _, _ = w.Write([]byte("ok")) })) t.Cleanup(server.Close) var delays []time.Duration repo := newOCRTestRepository(server.Client(), zap.NewNop(), func(_ context.Context, delay time.Duration) error { delays = append(delays, delay) return nil }) body, err := repo.doPaddleRequest(t.Context(), "submit", func(ctx context.Context) (*http.Request, error) { return http.NewRequestWithContext(ctx, http.MethodPost, server.URL, nil) }) if err != nil { t.Fatalf("Paddle 请求失败: %v", err) } if string(body) != "ok" || requests.Load() != 2 { t.Fatalf("body = %q, requests = %d", body, requests.Load()) } if len(delays) != 1 || delays[0] != 2*time.Second { t.Fatalf("retry delays = %v, want [2s]", delays) } } func TestDoPaddleRequestDoesNotRetryDeterministic4xx(t *testing.T) { var requests atomic.Int32 server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { requests.Add(1) w.WriteHeader(http.StatusUnsupportedMediaType) _, _ = w.Write([]byte(`{"code":"invalid_format","message":"unsupported image"}`)) })) t.Cleanup(server.Close) repo := newOCRTestRepository(server.Client(), zap.NewNop(), func(context.Context, time.Duration) error { return nil }) _, err := repo.doPaddleRequest(t.Context(), "submit", func(ctx context.Context) (*http.Request, error) { return http.NewRequestWithContext(ctx, http.MethodPost, server.URL, nil) }) var upstreamErr *ocrUpstreamError if !errors.As(err, &upstreamErr) { t.Fatalf("error = %v, want ocrUpstreamError", err) } if requests.Load() != 1 || upstreamErr.Retriable { t.Fatalf("requests = %d, retriable = %v", requests.Load(), upstreamErr.Retriable) } if upstreamErr.StatusCode != http.StatusUnsupportedMediaType || upstreamErr.Code != "invalid_format" { t.Fatalf("upstream error = %+v", upstreamErr) } } func TestDoPaddleRequestLogsStructured502Failure(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.Header().Set("X-Trace-ID", "trace-502") w.WriteHeader(http.StatusBadGateway) _, _ = w.Write([]byte(`{"code":"capacity","message":"upstream busy"}`)) })) t.Cleanup(server.Close) core, observed := observer.New(zap.DebugLevel) repo := newOCRTestRepository(server.Client(), zap.New(core), func(context.Context, time.Duration) error { return nil }) _, err := repo.doPaddleRequest(t.Context(), "submit", func(ctx context.Context) (*http.Request, error) { req, requestErr := http.NewRequestWithContext(ctx, http.MethodPost, server.URL, nil) if requestErr == nil { req.Header.Set("Authorization", "bearer secret-token") } return req, requestErr }) var upstreamErr *ocrUpstreamError if !errors.As(err, &upstreamErr) { t.Fatalf("error = %v, want ocrUpstreamError", err) } if upstreamErr.StatusCode != http.StatusBadGateway || upstreamErr.TraceID != "trace-502" || upstreamErr.Code != "capacity" { t.Fatalf("upstream error = %+v", upstreamErr) } entries := observed.AllUntimed() if len(entries) != ocrRequestMaxAttempts { t.Fatalf("log entries = %d, want %d", len(entries), ocrRequestMaxAttempts) } fields := entries[len(entries)-1].ContextMap() if fields["final"] != true || fields["upstream_status"] != int64(http.StatusBadGateway) || fields["upstream_trace_id"] != "trace-502" { t.Fatalf("final log fields = %#v", fields) } for _, entry := range entries { if strings.Contains(entry.Message, "secret-token") || strings.Contains(fmt.Sprint(entry.Context), "secret-token") { t.Fatal("OCR 日志泄露了 Authorization Token") } } } func TestReadAndValidateOCRFileNormalizesJPEGAndSize(t *testing.T) { source := image.NewNRGBA(image.Rect(0, 0, 3000, 1200)) draw.Draw(source, source.Bounds(), &image.Uniform{C: color.NRGBA{R: 25, G: 80, B: 160, A: 255}}, image.Point{}, draw.Src) var input bytes.Buffer if err := png.Encode(&input, source); err != nil { t.Fatalf("编码测试图片失败: %v", err) } data, contentType, err := readAndValidateOCRFile(&input, "image/png") if err != nil { t.Fatalf("归一化图片失败: %v", err) } if contentType != "image/jpeg" || len(data) < 2 || data[0] != 0xff || data[1] != 0xd8 { t.Fatalf("content_type = %q, signature = %x", contentType, data[:min(2, len(data))]) } config, format, err := image.DecodeConfig(bytes.NewReader(data)) if err != nil { t.Fatalf("读取归一化图片失败: %v", err) } if format != "jpeg" || config.Width != 2560 || config.Height != 1024 { t.Fatalf("format = %q, size = %dx%d", format, config.Width, config.Height) } } func TestReadAndValidateOCRFileRejectsExcessiveDecodedPixels(t *testing.T) { data := pngHeader(8001, 5000) config, _, decodeErr := image.DecodeConfig(bytes.NewReader(data)) if decodeErr != nil || config.Width != 8001 || config.Height != 5000 { t.Fatalf("测试 PNG 头无效: config=%+v, error=%v", config, decodeErr) } _, _, err := readAndValidateOCRFile(bytes.NewReader(data), "image/png") if !errors.Is(err, ErrQrCodeOCRInvalidFile) { t.Fatalf("error = %v, want ErrQrCodeOCRInvalidFile", err) } } func TestOCRRequestGateLimitsConcurrency(t *testing.T) { gate := newOCRRequestGate(2, 0) sleep := func(context.Context, time.Duration) error { return nil } releaseFirst, err := gate.acquire(t.Context(), sleep) if err != nil { t.Fatal(err) } releaseSecond, err := gate.acquire(t.Context(), sleep) if err != nil { t.Fatal(err) } third := make(chan func(), 1) go func() { release, acquireErr := gate.acquire(t.Context(), sleep) if acquireErr == nil { third <- release } }() select { case release := <-third: release() t.Fatal("第三个请求在并发槽释放前进入") case <-time.After(30 * time.Millisecond): } releaseFirst() select { case release := <-third: release() case <-time.After(time.Second): t.Fatal("并发槽释放后第三个请求仍未进入") } releaseSecond() } func TestNormalizeOCRFilenameUsesJPEGExtension(t *testing.T) { if got := normalizeOCRFilename("group.qrcode.webp"); got != "group.qrcode.jpg" { t.Fatalf("filename = %q", got) } if got := normalizeOCRFilename(""); got != "qrcode.jpg" { t.Fatalf("empty filename = %q", got) } } func pngHeader(width, height uint32) []byte { var result bytes.Buffer result.Write([]byte{137, 80, 78, 71, 13, 10, 26, 10}) data := make([]byte, 13) binary.BigEndian.PutUint32(data[0:4], width) binary.BigEndian.PutUint32(data[4:8], height) data[8] = 8 data[9] = 2 binary.Write(&result, binary.BigEndian, uint32(len(data))) result.WriteString("IHDR") result.Write(data) checksum := crc32.ChecksumIEEE(append([]byte("IHDR"), data...)) binary.Write(&result, binary.BigEndian, checksum) return result.Bytes() }