mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-04 12:39:06 +02:00
Restrict the HTTP/1.1 fallback retry and tighten the h2 failure classification
This commit is contained in:
@@ -10,10 +10,12 @@ import (
|
||||
"crypto/x509"
|
||||
"crypto/x509/pkix"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"math/big"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -87,6 +89,146 @@ func TestUpstreamTransport_ExplicitHTTP2NeverDowngrades(t *testing.T) {
|
||||
assert.False(t, mt.insecure.isDowngraded(srv.addr), "an explicit version must never pin an upstream")
|
||||
}
|
||||
|
||||
// TestUpstreamTransport_AutoDoesNotReplayUnsafeRequests covers the
|
||||
// other half of the fallback: an h2 failure says nothing about whether
|
||||
// the upstream already applied the request, so a state-changing one is
|
||||
// not replayed. The host is still pinned, so the next request rides
|
||||
// HTTP/1.1 without a second h2 attempt.
|
||||
func TestUpstreamTransport_AutoDoesNotReplayUnsafeRequests(t *testing.T) {
|
||||
t.Setenv(EnvUpstreamHTTPVersion, string(upstreamHTTPAuto))
|
||||
srv := startBrokenHTTP2Server(t)
|
||||
// A stream error means the stream was open, so the upstream had the
|
||||
// request in hand — unlike the GOAWAY at stream 0 the fake server
|
||||
// sends, which states it processed nothing.
|
||||
streamErr := http2.StreamError{StreamID: 1, Code: http2.ErrCodeProtocol}
|
||||
|
||||
mt := NewDirectOnly(nil)
|
||||
transport := mt.insecure
|
||||
ctx := WithSkipTLSVerify(WithDirectUpstream(context.Background()))
|
||||
|
||||
assert.False(t, safeToRetry(newTestRequest(t, ctx, http.MethodPost, srv.addr), streamErr),
|
||||
"a POST must not be replayed after a failure that may have been applied")
|
||||
assert.True(t, safeToRetry(newTestRequest(t, ctx, http.MethodGet, srv.addr), streamErr),
|
||||
"a GET is safe to replay whatever the failure was")
|
||||
|
||||
// The fake server's GOAWAY names stream 0, so even a POST is safe
|
||||
// there and the request must succeed over HTTP/1.1.
|
||||
resp, err := transport.RoundTrip(newTestRequest(t, ctx, http.MethodPost, srv.addr))
|
||||
require.NoError(t, err, "a GOAWAY at stream 0 means the upstream applied nothing, so the POST may be replayed")
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
_ = resp.Body.Close()
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "http/1.1", string(body), "the retry must reach the upstream over http/1.1")
|
||||
|
||||
h2Attempts := srv.http2Handshakes()
|
||||
resp, err = transport.RoundTrip(newTestRequest(t, ctx, http.MethodPost, srv.addr))
|
||||
require.NoError(t, err)
|
||||
_ = resp.Body.Close()
|
||||
assert.Equal(t, h2Attempts, srv.http2Handshakes(),
|
||||
"the pin must carry later requests without another h2 attempt")
|
||||
}
|
||||
|
||||
func newTestRequest(t *testing.T, ctx context.Context, method, addr string) *http.Request {
|
||||
t.Helper()
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, method, "https://"+addr, strings.NewReader("payload"))
|
||||
require.NoError(t, err)
|
||||
|
||||
return req
|
||||
}
|
||||
|
||||
func TestUpstreamKey(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
url string
|
||||
want string
|
||||
}{
|
||||
{name: "host", url: "https://backend.invalid/path", want: "backend.invalid"},
|
||||
// DNS is case-insensitive, so these are one upstream and must
|
||||
// share one pin.
|
||||
{name: "mixed case", url: "https://Backend.INVALID/path", want: "backend.invalid"},
|
||||
// The default port is implied on every path that can downgrade.
|
||||
{name: "explicit default port", url: "https://backend.invalid:443/", want: "backend.invalid"},
|
||||
{name: "non-default port", url: "https://backend.invalid:8443/", want: "backend.invalid:8443"},
|
||||
// An IPv6 literal needs its brackets back after Hostname strips
|
||||
// them, or the key is not a dialable authority.
|
||||
{name: "ipv6 default port", url: "https://[2001:db8::1]/", want: "2001:db8::1"},
|
||||
{name: "ipv6 with port", url: "https://[2001:db8::1]:8443/", want: "[2001:db8::1]:8443"},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
parsed, err := url.Parse(tc.url)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, tc.want, upstreamKey(parsed))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSafeToRetry(t *testing.T) {
|
||||
streamErr := http2.StreamError{StreamID: 1, Code: http2.ErrCodeProtocol}
|
||||
goAwayAtZero := errors.New(`http2: server sent GOAWAY and closed the connection; LastStreamID=0, ErrCode=HTTP_1_1_REQUIRED, debug=""`)
|
||||
goAwayLater := errors.New(`http2: server sent GOAWAY and closed the connection; LastStreamID=11, ErrCode=PROTOCOL_ERROR, debug=""`)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
method string
|
||||
headers map[string]string
|
||||
err error
|
||||
want bool
|
||||
}{
|
||||
{name: "get", method: http.MethodGet, err: streamErr, want: true},
|
||||
{name: "head", method: http.MethodHead, err: streamErr, want: true},
|
||||
{name: "options", method: http.MethodOptions, err: streamErr, want: true},
|
||||
{name: "trace", method: http.MethodTrace, err: streamErr, want: true},
|
||||
{name: "post", method: http.MethodPost, err: streamErr, want: false},
|
||||
{name: "put", method: http.MethodPut, err: streamErr, want: false},
|
||||
{name: "patch", method: http.MethodPatch, err: streamErr, want: false},
|
||||
{name: "delete", method: http.MethodDelete, err: streamErr, want: false},
|
||||
// The upstream reported it handled nothing, so repeating the
|
||||
// request cannot duplicate anything.
|
||||
{name: "post with goaway at stream 0", method: http.MethodPost, err: goAwayAtZero, want: true},
|
||||
// What the bundled transport actually raises for the first
|
||||
// stream on a connection the upstream GOAWAYs with a real error
|
||||
// code, which is the shape a real IIS site produces.
|
||||
{
|
||||
name: "post with first-stream goaway abort",
|
||||
method: http.MethodPost,
|
||||
err: errors.New("http2: Transport received GOAWAY from server ErrCode:HTTP_1_1_REQUIRED"),
|
||||
want: true,
|
||||
},
|
||||
// It handled earlier streams, so this one may have been applied.
|
||||
{name: "post with goaway after other streams", method: http.MethodPost, err: goAwayLater, want: false},
|
||||
{
|
||||
name: "post with idempotency key",
|
||||
method: http.MethodPost,
|
||||
headers: map[string]string{"Idempotency-Key": "abc"},
|
||||
err: streamErr,
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "post with prefixed idempotency key",
|
||||
method: http.MethodPost,
|
||||
headers: map[string]string{"X-Idempotency-Key": "abc"},
|
||||
err: streamErr,
|
||||
want: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
req, err := http.NewRequest(tc.method, "https://backend.invalid", nil)
|
||||
require.NoError(t, err)
|
||||
for k, v := range tc.headers {
|
||||
req.Header.Set(k, v)
|
||||
}
|
||||
|
||||
assert.Equal(t, tc.want, safeToRetry(req, tc.err))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpstreamTransport_MayDowngrade(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
@@ -227,6 +369,30 @@ func TestIsHTTP2ProtocolError(t *testing.T) {
|
||||
err: errors.New("Get \"https://backend.invalid\": http2: client connection lost"),
|
||||
want: true,
|
||||
},
|
||||
// A GOAWAY with NO_ERROR is a server draining a connection —
|
||||
// recycling an application pool, capping requests per
|
||||
// connection, shutting down gracefully. It speaks h2 fine.
|
||||
{
|
||||
name: "graceful goaway",
|
||||
err: errors.New(`http2: server sent GOAWAY and closed the connection; LastStreamID=9, ErrCode=NO_ERROR, debug=""`),
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
name: "graceful shutdown abort",
|
||||
err: errors.New("http2: Transport received Server's graceful shutdown GOAWAY"),
|
||||
want: false,
|
||||
},
|
||||
// Markers are substrings, so a URL carried by a *url.Error must
|
||||
// not be able to classify a plain failure as an h2 one.
|
||||
{
|
||||
name: "url error whose path looks like a marker",
|
||||
err: &url.Error{
|
||||
Op: "Get",
|
||||
URL: "https://backend.invalid/http2:/connection error: x",
|
||||
Err: errors.New("dial tcp 10.0.0.1:443: connect: connection refused"),
|
||||
},
|
||||
want: false,
|
||||
},
|
||||
// Retrying these on HTTP/1.1 fixes nothing, so they must never
|
||||
// pin an upstream.
|
||||
{name: "dial failure", err: errors.New("dial tcp 10.0.0.1:443: connect: connection refused"), want: false},
|
||||
@@ -332,7 +498,8 @@ func (s *brokenHTTP2Server) handle(conn net.Conn) {
|
||||
return
|
||||
}
|
||||
|
||||
if tlsConn.ConnectionState().NegotiatedProtocol == "h2" {
|
||||
proto := tlsConn.ConnectionState().NegotiatedProtocol
|
||||
if proto == "h2" {
|
||||
select {
|
||||
case s.handshakes <- struct{}{}:
|
||||
default:
|
||||
@@ -341,7 +508,7 @@ func (s *brokenHTTP2Server) handle(conn net.Conn) {
|
||||
return
|
||||
}
|
||||
|
||||
s.serveHTTP1(tlsConn)
|
||||
s.serveHTTP1(tlsConn, proto)
|
||||
}
|
||||
|
||||
// refuseHTTP2 completes just enough of the h2 handshake for the client
|
||||
@@ -370,16 +537,18 @@ func (s *brokenHTTP2Server) refuseHTTP2(conn net.Conn) {
|
||||
_, _ = io.Copy(io.Discard, conn)
|
||||
}
|
||||
|
||||
// serveHTTP1 answers a single request with the protocol the upstream
|
||||
// saw, so the test can tell which transport carried it.
|
||||
func (s *brokenHTTP2Server) serveHTTP1(conn net.Conn) {
|
||||
// serveHTTP1 answers a single request with the ALPN protocol the
|
||||
// upstream actually settled on, so a test asserting on the body is
|
||||
// checking what the upstream saw rather than a constant.
|
||||
func (s *brokenHTTP2Server) serveHTTP1(conn net.Conn, alpn string) {
|
||||
reader := bufio.NewReader(conn)
|
||||
if _, err := http.ReadRequest(reader); err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
const body = "http/1.1"
|
||||
_, _ = io.WriteString(conn, "HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nContent-Length: 8\r\nConnection: close\r\n\r\n"+body)
|
||||
_, _ = fmt.Fprintf(conn,
|
||||
"HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nContent-Length: %d\r\nConnection: close\r\n\r\n%s",
|
||||
len(alpn), alpn)
|
||||
}
|
||||
|
||||
func selfSignedCert(t *testing.T) tls.Certificate {
|
||||
|
||||
Reference in New Issue
Block a user