225 lines
7.5 KiB
Go
225 lines
7.5 KiB
Go
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()
|
|
}
|