Files
token_thief/tests/proxy/robustness_test.go
T

493 lines
16 KiB
Go

package proxy_test
import (
"context"
"errors"
"io"
"net/http"
"net/http/httptest"
"net/netip"
"net/url"
"strings"
"sync/atomic"
"testing"
"time"
"git.misaka.ren/M1saka/token_thief/config"
"git.misaka.ren/M1saka/token_thief/proxy"
)
type failingBody struct{ err error }
func (b failingBody) Read([]byte) (int, error) { return 0, b.err }
func (failingBody) Close() error { return nil }
type lateFailingBody struct {
remaining int
err error
}
func (b *lateFailingBody) Read(p []byte) (int, error) {
if b.remaining == 0 {
return 0, b.err
}
n := min(len(p), b.remaining)
for i := range p[:n] {
p[i] = 'x'
}
b.remaining -= n
return n, nil
}
func (*lateFailingBody) Close() error { return nil }
type failingResponseWriter struct {
header http.Header
short bool
}
func (w *failingResponseWriter) Header() http.Header { return w.header }
func (*failingResponseWriter) WriteHeader(int) {}
func (w *failingResponseWriter) Write(p []byte) (int, error) {
if w.short {
return len(p) - 1, nil
}
return 0, errors.New("write failed")
}
type failingFlushResponseWriter struct{ header http.Header }
func (w *failingFlushResponseWriter) Header() http.Header { return w.header }
func (*failingFlushResponseWriter) WriteHeader(int) {}
func (*failingFlushResponseWriter) Write(p []byte) (int, error) {
return len(p), nil
}
func (*failingFlushResponseWriter) Flush() {}
func (*failingFlushResponseWriter) FlushError() error {
return errors.New("flush failed")
}
type unflushableResponseWriter struct{ header http.Header }
func (w *unflushableResponseWriter) Header() http.Header { return w.header }
func (*unflushableResponseWriter) WriteHeader(int) {}
func (*unflushableResponseWriter) Write(p []byte) (int, error) {
return len(p), nil
}
func newTestHandler(t *testing.T, upstream *url.URL, sub *captureSubmitter, opts proxy.Options) *proxy.Handler {
t.Helper()
filter, err := config.NewFilter(config.FilterDisabled, nil)
if err != nil {
t.Fatal(err)
}
return proxy.NewWithOptions(upstream, filter, sub, 1024*1024, opts)
}
func TestRequestBodyReadFailureReturnsFixed400WithoutUpstream(t *testing.T) {
var upstreamCalls atomic.Int32
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
upstreamCalls.Add(1)
w.WriteHeader(http.StatusOK)
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
h := newTestHandler(t, u, &captureSubmitter{}, proxy.Options{})
req := httptest.NewRequest(http.MethodPost, "http://proxy.test/v1/chat", nil)
req.Body = failingBody{err: errors.New("secret read failure")}
req.ContentLength = 1
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
if rec.Code != http.StatusBadRequest || rec.Body.String() != "bad request\n" {
t.Fatalf("response=(%d, %q), want fixed 400 bad request", rec.Code, rec.Body.String())
}
if upstreamCalls.Load() != 0 {
t.Fatalf("upstream called %d times", upstreamCalls.Load())
}
}
func TestRequestBodyOverLimitReturns413WithoutUpstream(t *testing.T) {
var upstreamCalls atomic.Int32
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
upstreamCalls.Add(1)
w.WriteHeader(http.StatusOK)
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
h := newTestHandler(t, u, &captureSubmitter{}, proxy.Options{MaxRequestBytes: 4})
for _, tc := range []struct {
name string
contentLength int64
}{
{name: "known length", contentLength: 5},
{name: "chunked", contentLength: -1},
} {
t.Run(tc.name, func(t *testing.T) {
req := httptest.NewRequest(http.MethodPost, "http://proxy.test/v1/chat", strings.NewReader("12345"))
req.ContentLength = tc.contentLength
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
if rec.Code != http.StatusRequestEntityTooLarge || rec.Body.String() != "request body too large\n" {
t.Fatalf("response=(%d, %q), want fixed 413", rec.Code, rec.Body.String())
}
})
}
if upstreamCalls.Load() != 0 {
t.Fatalf("upstream called %d times", upstreamCalls.Load())
}
}
func TestRequestBodyReadFailureWithZeroContentLengthDoesNotReachUpstream(t *testing.T) {
var upstreamCalls atomic.Int32
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
upstreamCalls.Add(1)
w.WriteHeader(http.StatusOK)
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
h := newTestHandler(t, u, &captureSubmitter{}, proxy.Options{})
req := httptest.NewRequest(http.MethodPost, "http://proxy.test/v1/chat", nil)
req.Body = failingBody{err: errors.New("secret read failure")}
req.ContentLength = 0
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
if rec.Code != http.StatusBadRequest || rec.Body.String() != "bad request\n" {
t.Fatalf("response=(%d, %q), want fixed 400 bad request", rec.Code, rec.Body.String())
}
if upstreamCalls.Load() != 0 {
t.Fatalf("upstream called %d times", upstreamCalls.Load())
}
}
func TestRequestBodyReadFailureAfterCaptureLimitDoesNotReachUpstream(t *testing.T) {
var upstreamCalls atomic.Int32
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
upstreamCalls.Add(1)
_, _ = io.Copy(io.Discard, r.Body)
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
h := newTestHandler(t, u, &captureSubmitter{}, proxy.Options{})
req := httptest.NewRequest(http.MethodPost, "http://proxy.test/v1/chat", nil)
req.Body = &lateFailingBody{remaining: 1024*1024 + 1, err: errors.New("late read failure")}
req.ContentLength = 1024*1024 + 2
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
if rec.Code != http.StatusBadRequest || rec.Body.String() != "bad request\n" {
t.Fatalf("response=(%d, %q), want fixed 400 bad request", rec.Code, rec.Body.String())
}
if upstreamCalls.Load() != 0 {
t.Fatalf("upstream called %d times", upstreamCalls.Load())
}
}
func TestFilteredRequestBodyReadFailureDoesNotReachUpstream(t *testing.T) {
var upstreamCalls atomic.Int32
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
upstreamCalls.Add(1)
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
filter, err := config.NewFilter(config.FilterBlacklist, []string{"/ignored"})
if err != nil {
t.Fatal(err)
}
h := proxy.NewWithOptions(u, filter, &captureSubmitter{}, 1024, proxy.Options{})
req := httptest.NewRequest(http.MethodPost, "http://proxy.test/ignored", nil)
req.Body = failingBody{err: errors.New("read failure")}
req.ContentLength = 1
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
if rec.Code != http.StatusBadRequest || rec.Body.String() != "bad request\n" {
t.Fatalf("response=(%d, %q), want fixed 400 bad request", rec.Code, rec.Body.String())
}
if upstreamCalls.Load() != 0 {
t.Fatalf("upstream called %d times", upstreamCalls.Load())
}
}
func TestBadGatewayResponseDoesNotLeakUpstreamError(t *testing.T) {
badURL, _ := url.Parse("http://127.0.0.1:1")
h := newTestHandler(t, badURL, &captureSubmitter{}, proxy.Options{})
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "http://proxy.test/v1/test", nil))
if rec.Code != http.StatusBadGateway || rec.Body.String() != "bad gateway\n" {
t.Fatalf("response=(%d, %q), want fixed 502 bad gateway", rec.Code, rec.Body.String())
}
}
func TestResponseWriteFailurePreventsCommit(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_, _ = io.WriteString(w, "response")
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
for _, short := range []bool{false, true} {
sub := &captureSubmitter{}
h := newTestHandler(t, u, sub, proxy.Options{})
w := &failingResponseWriter{header: make(http.Header), short: short}
h.ServeHTTP(w, httptest.NewRequest(http.MethodGet, "http://proxy.test/v1/test", nil))
if sub.Len() != 0 {
t.Fatalf("short=%v: failed response write was committed", short)
}
}
}
func TestResponseFlushFailurePreventsCommit(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_, _ = io.WriteString(w, "response")
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
sub := &captureSubmitter{}
h := newTestHandler(t, u, sub, proxy.Options{})
w := &failingFlushResponseWriter{header: make(http.Header)}
h.ServeHTTP(w, httptest.NewRequest(http.MethodGet, "http://proxy.test/v1/test", nil))
if sub.Len() != 0 {
t.Fatal("failed response flush was committed")
}
}
func TestHeaderOnlyResponseFlushFailurePreventsCommit(t *testing.T) {
tests := []struct {
name string
method string
status int
}{
{name: "head", method: http.MethodHead, status: http.StatusOK},
{name: "no content", method: http.MethodGet, status: http.StatusNoContent},
{name: "not modified", method: http.MethodGet, status: http.StatusNotModified},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(tc.status)
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
sub := &captureSubmitter{}
h := newTestHandler(t, u, sub, proxy.Options{})
w := &failingFlushResponseWriter{header: make(http.Header)}
h.ServeHTTP(w, httptest.NewRequest(tc.method, "http://proxy.test/v1/test", nil))
if sub.Len() != 0 {
t.Fatal("response with undelivered headers was committed")
}
})
}
}
func TestHeaderOnlyResponseSuccessfulFlushCommits(t *testing.T) {
for _, status := range []int{http.StatusNoContent, http.StatusNotModified} {
t.Run(http.StatusText(status), func(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(status)
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
sub := &captureSubmitter{}
h := newTestHandler(t, u, sub, proxy.Options{})
h.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "http://proxy.test/v1/test", nil))
if sub.Len() != 1 {
t.Fatalf("entries=%d, want successfully delivered response committed", sub.Len())
}
})
}
}
func TestHeaderOnlyResponseWithoutDeliveryBoundaryDoesNotCommit(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusNoContent)
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
sub := &captureSubmitter{}
h := newTestHandler(t, u, sub, proxy.Options{})
w := &unflushableResponseWriter{header: make(http.Header)}
h.ServeHTTP(w, httptest.NewRequest(http.MethodGet, "http://proxy.test/v1/test", nil))
if sub.Len() != 0 {
t.Fatal("header-only response without a delivery boundary was committed")
}
}
func TestTrustedProxyClientIP(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_, _ = io.WriteString(w, "ok")
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
trusted := netip.MustParsePrefix("10.0.0.0/8")
tests := []struct {
name string
peer string
xff string
wantIP string
}{
{name: "untrusted peer ignores xff", peer: "203.0.113.9:1234", xff: "198.51.100.1", wantIP: "203.0.113.9"},
{name: "strip trusted from right", peer: "10.0.0.2:1234", xff: "198.51.100.7, 10.0.0.3", wantIP: "198.51.100.7"},
{name: "ipv6", peer: "[2001:db8::2]:1234", xff: "198.51.100.7", wantIP: "2001:db8::2"},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
sub := &captureSubmitter{}
h := newTestHandler(t, u, sub, proxy.Options{TrustedProxies: []netip.Prefix{trusted}})
req := httptest.NewRequest(http.MethodGet, "http://proxy.test/v1/test", nil)
req.RemoteAddr = tc.peer
req.Header.Set("X-Forwarded-For", tc.xff)
h.ServeHTTP(httptest.NewRecorder(), req)
if sub.Len() != 1 || sub.Entry(0).ClientIP != tc.wantIP {
t.Fatalf("ClientIP=%q, want %q", sub.Entry(0).ClientIP, tc.wantIP)
}
})
}
}
func TestTrustedProxyUsesValidXRealIPWithoutXFF(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
sub := &captureSubmitter{}
trusted := netip.MustParsePrefix("192.0.2.0/24")
h := newTestHandler(t, u, sub, proxy.Options{TrustedProxies: []netip.Prefix{trusted}})
req := httptest.NewRequest(http.MethodGet, "http://proxy.test/x", nil)
req.RemoteAddr = "192.0.2.10:1234"
req.Header.Set("X-Real-IP", "198.51.100.20")
h.ServeHTTP(httptest.NewRecorder(), req)
if sub.Len() != 1 || sub.Entry(0).ClientIP != "198.51.100.20" {
t.Fatalf("entries=%d", sub.Len())
}
}
func TestResponseBodyTimeoutPreventsCommit(t *testing.T) {
release := make(chan struct{})
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusOK)
w.(http.Flusher).Flush()
<-release
}))
u, _ := url.Parse(upstream.URL)
sub := &captureSubmitter{}
h := newTestHandler(t, u, sub, proxy.Options{ResponseTimeout: 50 * time.Millisecond})
srv := httptest.NewServer(h)
resp, err := http.Get(srv.URL + "/v1/test")
if err == nil {
_, _ = io.ReadAll(resp.Body)
_ = resp.Body.Close()
}
time.Sleep(50 * time.Millisecond)
if sub.Len() != 0 {
t.Fatal("timed out response must not be committed")
}
close(release)
srv.Close()
upstream.Close()
}
func TestSSEIdleTimeoutResetsAfterReads(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "text/event-stream")
f := w.(http.Flusher)
for _, event := range []string{"data: one\n\n", "data: two\n\n", "data: [DONE]\n\n"} {
_, _ = io.WriteString(w, event)
f.Flush()
time.Sleep(30 * time.Millisecond)
}
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
sub := &captureSubmitter{}
h := newTestHandler(t, u, sub, proxy.Options{SSEIdleTimeout: 60 * time.Millisecond})
srv := httptest.NewServer(h)
defer srv.Close()
resp, err := http.Get(srv.URL + "/v1/test")
if err != nil {
t.Fatal(err)
}
_, _ = io.ReadAll(resp.Body)
_ = resp.Body.Close()
if sub.Len() != 1 {
t.Fatalf("entries=%d, want completed SSE commit", sub.Len())
}
}
func TestIncompleteSSEEventPreventsCommit(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "text/event-stream")
_, _ = io.WriteString(w, "data: partial")
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
sub := &captureSubmitter{}
h := newTestHandler(t, u, sub, proxy.Options{})
h.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "http://proxy.test/v1/test", nil))
if sub.Len() != 0 {
t.Fatal("SSE ending mid-event must not be committed as complete")
}
}
func TestIncompleteSSEEventAfterTerminalEventPreventsCommit(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "text/event-stream")
_, _ = io.WriteString(w, "data: [DONE]\n\n")
w.(http.Flusher).Flush()
_, _ = io.WriteString(w, "data: partial")
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
sub := &captureSubmitter{}
h := newTestHandler(t, u, sub, proxy.Options{})
h.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "http://proxy.test/v1/test", nil))
if sub.Len() != 0 {
t.Fatal("SSE ending mid-event after a terminal event must not be committed")
}
}
func TestShutdownReturnsWithoutHijackedConnections(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) }))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
h := newTestHandler(t, u, &captureSubmitter{}, proxy.Options{})
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := h.Shutdown(ctx); err != nil {
t.Fatal(err)
}
}
func TestTerminalTextInNonSSEBodyDoesNotAffectCommit(t *testing.T) {
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
_, _ = io.WriteString(w, `{"text":"data: [DONE]"}`)
}))
defer upstream.Close()
u, _ := url.Parse(upstream.URL)
sub := &captureSubmitter{}
h := newTestHandler(t, u, sub, proxy.Options{})
h.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "http://proxy.test/v1/test", strings.NewReader("")))
if sub.Len() != 1 {
t.Fatalf("entries=%d, want 1", sub.Len())
}
}