diff --git a/pkg/server/quickproxy/quickproxy.go b/pkg/server/quickproxy/quickproxy.go index 036258a..872e899 100644 --- a/pkg/server/quickproxy/quickproxy.go +++ b/pkg/server/quickproxy/quickproxy.go @@ -5,8 +5,10 @@ package quickproxy import ( + "bytes" "errors" "fmt" + "io" "net" "net/http" "net/http/httputil" @@ -191,6 +193,15 @@ func (s *Service) Handler(fallback http.Handler) http.Handler { for i := range s.rules { s.rules[i].proxy.ErrorHandler = func(w http.ResponseWriter, r *http.Request, _ error) { + // ReverseProxy consumes and closes r.Body while attempting the + // upstream request, even when that attempt fails (per the + // http.RoundTripper contract). Restore a fresh copy from + // r.GetBody, set below, before handing the request to fallback. + if r.GetBody != nil { + if body, err := r.GetBody(); err == nil { + r.Body = body + } + } fallback.ServeHTTP(w, r) } } @@ -201,6 +212,22 @@ func (s *Service) Handler(fallback http.Handler) http.Handler { fallback.ServeHTTP(w, r) return } + + // Buffer the body so it can be replayed to fallback if the upstream + // attempt fails; see ErrorHandler above. + if r.Body != nil && r.Body != http.NoBody { + bodyBytes, err := io.ReadAll(r.Body) + r.Body.Close() + if err != nil { + http.Error(w, "failed to read request body", http.StatusInternalServerError) + return + } + r.Body = io.NopCloser(bytes.NewReader(bodyBytes)) + r.GetBody = func() (io.ReadCloser, error) { + return io.NopCloser(bytes.NewReader(bodyBytes)), nil + } + } + rule.proxy.ServeHTTP(w, r) }) } diff --git a/pkg/server/quickproxy/quickproxy_test.go b/pkg/server/quickproxy/quickproxy_test.go index 7a4ced7..43df888 100644 --- a/pkg/server/quickproxy/quickproxy_test.go +++ b/pkg/server/quickproxy/quickproxy_test.go @@ -4,6 +4,7 @@ import ( "io" "net/http" "net/http/httptest" + "strings" "testing" "time" ) @@ -130,6 +131,75 @@ func TestHandler_UnreachableUpstreamFallsBack(t *testing.T) { } } +func TestHandler_UnreachableUpstreamFallsBackWithBody(t *testing.T) { + // A closed listener address: nothing is listening, so dialing fails and + // ReverseProxy invokes ErrorHandler. The fallback handler must still see + // the original request body, even though ReverseProxy consumed and + // closed it while attempting (and failing) the upstream request. + unreachable := "http://127.0.0.1:1" + + svc, err := NewService([]Rule{{URLPrefix: "/", Target: unreachable}}, WithTimeout(500*time.Millisecond)) + if err != nil { + t.Fatalf("NewService: %v", err) + } + + echoBody := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, err := io.ReadAll(r.Body) + if err != nil { + t.Fatalf("fallback reading body: %v", err) + } + w.WriteHeader(http.StatusOK) + _, _ = w.Write(body) + }) + + handler := svc.Handler(echoBody) + + req := httptest.NewRequest(http.MethodPost, "/submit", strings.NewReader("payload=1")) + rr := httptest.NewRecorder() + handler.ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Fatalf("status = %d, want 200", rr.Code) + } + if got := rr.Body.String(); got != "payload=1" { + t.Fatalf("body = %q, want payload=1", got) + } +} + +func TestHandler_404FallsBackWithBody(t *testing.T) { + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusNotFound) + })) + defer upstream.Close() + + svc, err := NewService([]Rule{{URLPrefix: "/", Target: upstream.URL}}) + if err != nil { + t.Fatalf("NewService: %v", err) + } + + echoBody := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, err := io.ReadAll(r.Body) + if err != nil { + t.Fatalf("fallback reading body: %v", err) + } + w.WriteHeader(http.StatusOK) + _, _ = w.Write(body) + }) + + handler := svc.Handler(echoBody) + + req := httptest.NewRequest(http.MethodPut, "/missing", strings.NewReader("payload=2")) + rr := httptest.NewRecorder() + handler.ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Fatalf("status = %d, want 200", rr.Code) + } + if got := rr.Body.String(); got != "payload=2" { + t.Fatalf("body = %q, want payload=2", got) + } +} + func TestHandler_NonNotFoundErrorsPassThrough(t *testing.T) { codes := []int{http.StatusOK, http.StatusForbidden, http.StatusBadRequest, http.StatusInternalServerError}