mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-09-21 15:42:01 +00:00
fix(quickproxy): ensure request body is preserved on fallback
This commit is contained in:
@@ -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)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user