Files
sub2api/backend/internal/repository/decompress_response_test.go
T
李建琦 6d655c9903
Release / update-version (push) Has been cancelled
Release / build-frontend (push) Has been cancelled
Release / release (push) Has been cancelled
Release / sync-version-file (push) Has been cancelled
CI / shell (push) Canceled after 0s
CI / test (push) Canceled after 0s
CI / frontend (push) Canceled after 0s
CI / golangci-lint (push) Canceled after 0s
Security Scan / backend-security (push) Canceled after 0s
Security Scan / frontend-security (push) Canceled after 0s
Sub2API v1.0 - AI API 网关(二开初始版本,基于上游 Wei-Shaw/sub2api)
2026-08-21 18:30:13 +08:00

191 lines
5.2 KiB
Go

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()
}