mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
fix(gateway): 解压 zstd 上游响应体
This commit is contained in:
@@ -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()
|
||||
}
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user