fix: restore reviewable migration evidence

This commit is contained in:
MiMoCode
2026-07-10 18:26:48 +08:00
commit d1bbb5370c
42 changed files with 6419 additions and 0 deletions
+336
View File
@@ -0,0 +1,336 @@
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")
}
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 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 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())
}
}