493 lines
16 KiB
Go
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())
|
|
}
|
|
}
|