Files
netbird/proxy/internal/roundtrip/multi_test.go
T
Viktor Liu 582c8b1ae8 Fall back to HTTP/1.1 when an upstream cannot serve the h2 it negotiated
ForceAttemptHTTP2 only puts h2 in the ALPN offer — the upstream still
picks — so "auto" already meant "whatever the upstream chose". What ALPN
cannot express is an upstream that selects h2 and then fails to speak it,
which is the case the setting was added for: today that leaves the
operator pinning every upstream to 1.1 to work around one broken backend.

Auto now completes itself. The first h2-level failure for a host pins
that host to an HTTP/1.1-only clone of its transport for 10 minutes and
retries the request there when the body can be replayed, so a broken
backend costs one failed h2 attempt instead of a configuration change.
The pin is per upstream host, so one broken backend does not drop the
others, and it expires so a fixed backend returns to h2 on its own.

Only h2 framing errors trigger it: a dial, TLS or context error says
nothing about the protocol and retrying it over HTTP/1.1 would fix
nothing. The explicit values stay absolute — "2" never downgrades.

Pinning HTTP/1.1 now also strips h2 from the ALPN offer. Configuring h2
makes net/http append it to the transport's TLSClientConfig, so a clone
taken from a transport that already served a request would otherwise
advertise a protocol the clone refuses to speak, and the reply would come
back as h2 frames parsed as an HTTP/1.1 message.
2026-09-03 18:18:00 +02:00

191 lines
7.0 KiB
Go

package roundtrip
import (
"context"
"io"
"net/http"
"net/http/httptest"
"os"
"strings"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// stubRoundTripper records whether RoundTrip was called and returns a
// canned response so tests can assert the dispatch decision without
// running a real network.
type stubRoundTripper struct {
called bool
body string
}
func (s *stubRoundTripper) RoundTrip(_ *http.Request) (*http.Response, error) {
s.called = true
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(s.body)),
Header: http.Header{},
}, nil
}
func TestMultiTransport_DispatchesByContextFlag(t *testing.T) {
embedded := &stubRoundTripper{body: "embedded"}
mt := NewMultiTransport(embedded, nil)
t.Run("default routes to embedded", func(t *testing.T) {
embedded.called = false
req := httptest.NewRequest(http.MethodGet, "http://example.invalid", nil)
resp, err := mt.RoundTrip(req)
require.NoError(t, err, "embedded path must not error on stubbed transport")
require.NotNil(t, resp)
_ = resp.Body.Close()
assert.True(t, embedded.called, "request without WithDirectUpstream must hit the embedded transport")
})
t.Run("WithDirectUpstream skips embedded", func(t *testing.T) {
embedded.called = false
// Hit a server we control to verify the stdlib transport is used.
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
_, _ = io.WriteString(w, "direct")
}))
defer srv.Close()
req, err := http.NewRequestWithContext(WithDirectUpstream(context.Background()), http.MethodGet, srv.URL, nil)
require.NoError(t, err)
resp, err := mt.RoundTrip(req)
require.NoError(t, err, "direct path must dial via stdlib transport")
body, err := io.ReadAll(resp.Body)
_ = resp.Body.Close()
require.NoError(t, err)
assert.Equal(t, "direct", string(body), "stdlib transport must reach the test server")
assert.False(t, embedded.called, "WithDirectUpstream must bypass the embedded transport")
})
}
// TestMultiTransport_AppliesEnvOverridesToDirect verifies that the
// NB_PROXY_* env vars consumed by loadTransportConfig flow into the
// direct branches (previously they only applied to the embedded
// roundtripper, so direct-upstream traffic ignored operator tuning).
func TestMultiTransport_AppliesEnvOverridesToDirect(t *testing.T) {
t.Setenv(EnvMaxIdleConns, "42")
t.Setenv(EnvIdleConnTimeout, "11s")
t.Setenv(EnvTLSHandshakeTimeout, "7s")
mt := NewMultiTransport(&stubRoundTripper{body: "embedded"}, nil)
assert.Equal(t, 42, mt.direct.primary.MaxIdleConns,
"NB_PROXY_MAX_IDLE_CONNS must propagate to the direct transport")
assert.Equal(t, 11*time.Second, mt.direct.primary.IdleConnTimeout,
"NB_PROXY_IDLE_CONN_TIMEOUT must propagate to the direct transport")
assert.Equal(t, 7*time.Second, mt.direct.primary.TLSHandshakeTimeout,
"NB_PROXY_TLS_HANDSHAKE_TIMEOUT must propagate to the direct transport")
assert.Equal(t, 42, mt.insecure.primary.MaxIdleConns,
"env tuning must also apply to the insecure-skip-verify direct transport")
}
// TestMultiTransport_UpstreamHTTPVersion pins the protocol actually
// negotiated with an HTTPS upstream that offers both h2 and http/1.1.
// The request rides the insecure clone, so this also covers the version
// surviving http.Transport.Clone.
func TestMultiTransport_UpstreamHTTPVersion(t *testing.T) {
tests := []struct {
name string
env string
wantProto string
}{
{name: "unset negotiates h2", env: "", wantProto: "HTTP/2.0"},
{name: "auto negotiates h2", env: "auto", wantProto: "HTTP/2.0"},
{name: "1.1 pins http/1.1", env: "1.1", wantProto: "HTTP/1.1"},
{name: "2 negotiates h2", env: "2", wantProto: "HTTP/2.0"},
{name: "unsupported value keeps the default", env: "http3", wantProto: "HTTP/2.0"},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
// t.Setenv registers the restore for whatever the process
// inherited; unsetting afterwards lets the default row
// exercise a genuinely absent variable.
t.Setenv(EnvUpstreamHTTPVersion, tc.env)
if tc.env == "" {
require.NoError(t, os.Unsetenv(EnvUpstreamHTTPVersion))
}
// The test server's certificate isn't in any root pool, so the
// request rides the insecure branch via WithSkipTLSVerify.
srv := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_, _ = io.WriteString(w, r.Proto)
}))
srv.EnableHTTP2 = true
srv.StartTLS()
defer srv.Close()
mt := NewDirectOnly(nil)
ctx := WithSkipTLSVerify(WithDirectUpstream(context.Background()))
req, err := http.NewRequestWithContext(ctx, http.MethodGet, srv.URL, nil)
require.NoError(t, err)
resp, err := mt.RoundTrip(req)
require.NoError(t, err)
body, err := io.ReadAll(resp.Body)
_ = resp.Body.Close()
require.NoError(t, err)
assert.Equal(t, tc.wantProto, resp.Proto,
"client-side protocol must follow %s=%q", EnvUpstreamHTTPVersion, tc.env)
assert.Equal(t, tc.wantProto, string(body),
"the upstream must see the same protocol the client negotiated")
})
}
}
// TestMultiTransport_NilEmbeddedErrorsWhenWGPathRequested guards
// against the previous silent fallback: a MultiTransport constructed
// without an embedded transport must reject requests that don't
// explicitly opt into the direct branch, rather than routing them
// over the host stack and bypassing WireGuard.
func TestMultiTransport_NilEmbeddedErrorsWhenWGPathRequested(t *testing.T) {
mt := NewMultiTransport(nil, nil)
req := httptest.NewRequest(http.MethodGet, "http://example.invalid", nil)
resp, err := mt.RoundTrip(req)
if resp != nil {
_ = resp.Body.Close()
}
require.Error(t, err, "nil embedded must surface as an explicit error, not a silent direct dispatch")
assert.Nil(t, resp)
assert.ErrorIs(t, err, errNoEmbeddedTransport,
"the error must be the sentinel so callers can distinguish misconfiguration from network failures")
}
// TestMultiTransport_DirectOnlyServesDirectBranch verifies NewDirectOnly
// constructs a MultiTransport whose direct branch handles requests with
// the direct-upstream flag set, and surfaces the explicit sentinel
// when the embedded path is reached.
func TestMultiTransport_DirectOnlyServesDirectBranch(t *testing.T) {
mt := NewDirectOnly(nil)
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
_, _ = io.WriteString(w, "ok")
}))
defer srv.Close()
req, err := http.NewRequestWithContext(WithDirectUpstream(context.Background()), http.MethodGet, srv.URL, nil)
require.NoError(t, err)
resp, err := mt.RoundTrip(req)
require.NoError(t, err, "direct-only must serve requests that opt into the direct branch")
_ = resp.Body.Close()
assert.Equal(t, http.StatusOK, resp.StatusCode)
wgReq := httptest.NewRequest(http.MethodGet, "http://example.invalid", nil)
resp, err = mt.RoundTrip(wgReq)
if resp != nil {
_ = resp.Body.Close()
}
require.Error(t, err, "direct-only must refuse requests that didn't opt into the direct branch")
assert.Nil(t, resp)
assert.ErrorIs(t, err, errNoEmbeddedTransport)
}