From c1c28ac7bb0b5569d2dabe99c5fe6e3e22068f9f Mon Sep 17 00:00:00 2001 From: wucm667 Date: Fri, 12 Jun 2026 15:06:01 +0800 Subject: [PATCH] =?UTF-8?q?fix(gateway):=20=E8=A7=A3=E5=8E=8B=20zstd=20?= =?UTF-8?q?=E4=B8=8A=E6=B8=B8=E5=93=8D=E5=BA=94=E4=BD=93?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../repository/decompress_response_test.go | 190 ++++++++++++++++++ backend/internal/repository/http_upstream.go | 36 +++- 2 files changed, 225 insertions(+), 1 deletion(-) create mode 100644 backend/internal/repository/decompress_response_test.go diff --git a/backend/internal/repository/decompress_response_test.go b/backend/internal/repository/decompress_response_test.go new file mode 100644 index 0000000000..8de64ef0f0 --- /dev/null +++ b/backend/internal/repository/decompress_response_test.go @@ -0,0 +1,190 @@ +package repository + +import ( + "bytes" + "compress/flate" + "compress/gzip" + "io" + "log/slog" + "net/http" + "testing" + + "github.com/andybalholm/brotli" + "github.com/klauspost/compress/zstd" + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +func TestDecompressResponseBodyZstdUsage(t *testing.T) { + payload := []byte(`{"usage":{"input_tokens":123,"output_tokens":45,"cache_read_input_tokens":67}}`) + compressed := compressZstd(t, payload) + resp := newEncodedResponse("zstd", compressed) + + decompressResponseBody(resp) + + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + require.Equal(t, payload, body) + require.Equal(t, int64(123), gjson.GetBytes(body, "usage.input_tokens").Int()) + require.Equal(t, int64(45), gjson.GetBytes(body, "usage.output_tokens").Int()) + require.Equal(t, int64(67), gjson.GetBytes(body, "usage.cache_read_input_tokens").Int()) + require.Empty(t, resp.Header.Get("Content-Encoding")) + require.Empty(t, resp.Header.Get("Content-Length")) + require.Equal(t, int64(-1), resp.ContentLength) + require.NoError(t, resp.Body.Close()) +} + +func TestDecompressResponseBodyExistingEncodings(t *testing.T) { + payload := []byte(`{"ok":true}`) + tests := []struct { + name string + encoding string + compress func(*testing.T, []byte) []byte + }{ + {name: "gzip", encoding: "gzip", compress: compressGzip}, + {name: "brotli", encoding: "br", compress: compressBrotli}, + {name: "deflate", encoding: "deflate", compress: compressDeflate}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + resp := newEncodedResponse(tt.encoding, tt.compress(t, payload)) + + decompressResponseBody(resp) + + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + require.Equal(t, payload, body) + require.Empty(t, resp.Header.Get("Content-Encoding")) + require.Empty(t, resp.Header.Get("Content-Length")) + require.Equal(t, int64(-1), resp.ContentLength) + require.NoError(t, resp.Body.Close()) + }) + } +} + +func TestDecompressResponseBodyWithoutEncodingLeavesBodyUntouched(t *testing.T) { + originalBody := &responseTestBody{Reader: bytes.NewReader([]byte("plain"))} + resp := &http.Response{ + Header: make(http.Header), + Body: originalBody, + ContentLength: 5, + } + + decompressResponseBody(resp) + + require.Same(t, originalBody, resp.Body) + require.Equal(t, int64(5), resp.ContentLength) + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + require.Equal(t, "plain", string(body)) + require.NoError(t, resp.Body.Close()) +} + +func TestDecompressResponseBodyInvalidZstdWarnsAndPreservesBody(t *testing.T) { + previousLogger := slog.Default() + var logOutput bytes.Buffer + slog.SetDefault(slog.New(slog.NewTextHandler(&logOutput, nil))) + t.Cleanup(func() { + slog.SetDefault(previousLogger) + }) + + payload := []byte("not a zstd response") + resp := newEncodedResponse("zstd", payload) + + require.NotPanics(t, func() { + decompressResponseBody(resp) + }) + + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + require.Equal(t, payload, body) + require.Equal(t, "zstd", resp.Header.Get("Content-Encoding")) + require.Equal(t, int64(len(payload)), resp.ContentLength) + require.Contains(t, logOutput.String(), "msg=zstd_decompress_failed") + require.NoError(t, resp.Body.Close()) +} + +func TestDecompressResponseBodyEmptyZstdWarnsAndPreservesBody(t *testing.T) { + previousLogger := slog.Default() + var logOutput bytes.Buffer + slog.SetDefault(slog.New(slog.NewTextHandler(&logOutput, nil))) + t.Cleanup(func() { + slog.SetDefault(previousLogger) + }) + + resp := newEncodedResponse("zstd", nil) + + require.NotPanics(t, func() { + decompressResponseBody(resp) + }) + + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + require.Empty(t, body) + require.Equal(t, "zstd", resp.Header.Get("Content-Encoding")) + require.Equal(t, int64(0), resp.ContentLength) + require.Contains(t, logOutput.String(), "msg=zstd_decompress_failed") + require.NoError(t, resp.Body.Close()) +} + +type responseTestBody struct { + io.Reader +} + +func (b *responseTestBody) Close() error { + return nil +} + +func newEncodedResponse(encoding string, body []byte) *http.Response { + header := make(http.Header) + header.Set("Content-Encoding", encoding) + header.Set("Content-Length", "123") + return &http.Response{ + Header: header, + Body: io.NopCloser(bytes.NewReader(body)), + ContentLength: int64(len(body)), + } +} + +func compressZstd(t *testing.T, payload []byte) []byte { + t.Helper() + var buf bytes.Buffer + zw, err := zstd.NewWriter(&buf) + require.NoError(t, err) + _, err = zw.Write(payload) + require.NoError(t, err) + require.NoError(t, zw.Close()) + return buf.Bytes() +} + +func compressGzip(t *testing.T, payload []byte) []byte { + t.Helper() + var buf bytes.Buffer + zw := gzip.NewWriter(&buf) + _, err := zw.Write(payload) + require.NoError(t, err) + require.NoError(t, zw.Close()) + return buf.Bytes() +} + +func compressBrotli(t *testing.T, payload []byte) []byte { + t.Helper() + var buf bytes.Buffer + zw := brotli.NewWriter(&buf) + _, err := zw.Write(payload) + require.NoError(t, err) + require.NoError(t, zw.Close()) + return buf.Bytes() +} + +func compressDeflate(t *testing.T, payload []byte) []byte { + t.Helper() + var buf bytes.Buffer + zw, err := flate.NewWriter(&buf, flate.DefaultCompression) + require.NoError(t, err) + _, err = zw.Write(payload) + require.NoError(t, err) + require.NoError(t, zw.Close()) + return buf.Bytes() +} diff --git a/backend/internal/repository/http_upstream.go b/backend/internal/repository/http_upstream.go index 476da3aed6..eac60d6268 100644 --- a/backend/internal/repository/http_upstream.go +++ b/backend/internal/repository/http_upstream.go @@ -1,6 +1,7 @@ package repository import ( + "bufio" "compress/flate" "compress/gzip" "context" @@ -18,6 +19,7 @@ import ( "time" "github.com/andybalholm/brotli" + "github.com/klauspost/compress/zstd" "github.com/Wei-Shaw/sub2api/internal/config" "github.com/Wei-Shaw/sub2api/internal/pkg/proxyurl" @@ -1178,6 +1180,7 @@ func decompressResponseBody(resp *http.Response) { return } + originalBody := resp.Body var reader io.Reader switch ce { case "gzip": @@ -1190,17 +1193,48 @@ func decompressResponseBody(resp *http.Response) { reader = brotli.NewReader(resp.Body) case "deflate": reader = flate.NewReader(resp.Body) + case "zstd": + bufferedBody := bufio.NewReader(resp.Body) + resp.Body = &decompressedBody{reader: bufferedBody, closer: originalBody} + + headerBytes, _ := bufferedBody.Peek(zstd.HeaderMaxSize) + var header zstd.Header + if err := header.Decode(headerBytes); err != nil { + slog.Warn("zstd_decompress_failed", "error", err) + return + } + + zr, err := zstd.NewReader(bufferedBody) + if err != nil { + slog.Warn("zstd_decompress_failed", "error", err) + return + } + reader = &zstdResponseReader{ReadCloser: zr.IOReadCloser()} default: return } - originalBody := resp.Body resp.Body = &decompressedBody{reader: reader, closer: originalBody} resp.Header.Del("Content-Encoding") resp.Header.Del("Content-Length") // 解压后长度不确定 resp.ContentLength = -1 } +type zstdResponseReader struct { + io.ReadCloser + warnOnce sync.Once +} + +func (r *zstdResponseReader) Read(p []byte) (int, error) { + n, err := r.ReadCloser.Read(p) + if err != nil && !errors.Is(err, io.EOF) { + r.warnOnce.Do(func() { + slog.Warn("zstd_decompress_failed", "error", err) + }) + } + return n, err +} + // decompressedBody 组合解压 reader 和原始 body 的 close。 type decompressedBody struct { reader io.Reader