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