proxy-pool/internal/gateway/server/handler_test.go

1192 lines
40 KiB
Go

package server
import (
"bufio"
"context"
"errors"
"io"
"net"
"net/http"
"net/http/httptest"
"net/netip"
"net/url"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
"proxy-pool/internal/config"
"proxy-pool/internal/domain/clientpolicy"
outcomeDomain "proxy-pool/internal/domain/outcome"
proxyDomain "proxy-pool/internal/domain/proxy"
"proxy-pool/internal/domain/routing"
"proxy-pool/internal/gateway/dispatch"
"proxy-pool/internal/gateway/policy"
"proxy-pool/internal/gateway/snapshot"
transportDomain "proxy-pool/internal/gateway/transport"
"proxy-pool/internal/platform/httpsecurity"
)
func TestHandlerRunsProtectionAndTargetPolicyBeforeRouting(t *testing.T) {
t.Parallel()
var calls []string
var mu sync.Mutex
record := func(name string) {
mu.Lock()
calls = append(calls, name)
mu.Unlock()
}
handler, err := New(Config{MaxAttempts: 1}, Dependencies{
Auth: GuardFunc(func(context.Context, *http.Request) error { record("auth"); return nil }),
Access: GuardFunc(func(context.Context, *http.Request) error { record("access"); return nil }),
Admission: GuardFunc(func(context.Context, *http.Request) error { record("admission"); return nil }),
Targets: fakeTargets{evaluateURL: func(context.Context, string) (policy.Authority, error) {
record("target")
return policy.Authority{Host: "example.test", Port: 80}, nil
}},
Router: RouteFunc(func(*http.Request) (dispatch.Request, error) {
record("route")
return dispatch.Request{}, errors.New("route stopped")
}),
Dispatcher: DispatcherFunc(func(dispatch.Request) (*dispatch.Lease, error) {
t.Fatal("dispatcher must not run after route error")
return nil, nil
}),
Transport: &fakeTransport{},
})
if err != nil {
t.Fatalf("New() error = %v", err)
}
response := httptest.NewRecorder()
handler.ServeHTTP(response, httptest.NewRequest(http.MethodGet, "http://example.test/resource", nil))
if got, want := strings.Join(calls, ","), "auth,access,admission,target,route"; got != want {
t.Fatalf("pipeline order = %q, want %q", got, want)
}
if response.Code != http.StatusBadGateway {
t.Fatalf("status = %d, want 502", response.Code)
}
}
func TestHandlerReportsLocalRequestAndTunnelLifecycles(t *testing.T) {
t.Parallel()
metrics := &recordingRequestMetrics{}
handler, err := New(Config{}, Dependencies{
Targets: fakeTargets{},
Router: RouteFunc(func(*http.Request) (dispatch.Request, error) {
return dispatch.Request{}, errors.New("route stopped")
}),
Dispatcher: DispatcherFunc(func(dispatch.Request) (*dispatch.Lease, error) {
t.Fatal("dispatcher must not run after route error")
return nil, nil
}),
Transport: &fakeTransport{},
Metrics: metrics,
})
if err != nil {
t.Fatalf("New() error = %v", err)
}
handler.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "http://example.test/resource", nil))
connect := httptest.NewRequest(http.MethodConnect, "http://example.test", nil)
connect.Host = "example.test:443"
handler.ServeHTTP(httptest.NewRecorder(), connect)
tunnel := &activeTunnel{}
if !handler.registerTunnel(tunnel) {
t.Fatal("registerTunnel() = false")
}
handler.unregisterTunnel(tunnel)
if got := metrics.started("HTTP"); got != 1 {
t.Fatalf("HTTP request starts = %d, want 1", got)
}
if got := metrics.finished("HTTP"); got != 1 {
t.Fatalf("HTTP request finishes = %d, want 1", got)
}
if got := metrics.started("CONNECT"); got != 1 {
t.Fatalf("CONNECT request starts = %d, want 1", got)
}
if got := metrics.finished("CONNECT"); got != 1 {
t.Fatalf("CONNECT request finishes = %d, want 1", got)
}
if metrics.opened != 1 || metrics.closed != 1 {
t.Fatalf("tunnel lifecycle = opened:%d closed:%d, want 1:1", metrics.opened, metrics.closed)
}
}
func TestHandlerRejectsStaticDirectRouteOutsideCredentialPolicy(t *testing.T) {
t.Parallel()
authentication, err := httpsecurity.New(httpsecurity.Config{
Authentication: httpsecurity.Authentication{
Mode: httpsecurity.ModeBearer, Token: "gateway-token",
ClientPolicy: clientpolicy.Policy{AllowedRoutings: []string{"checkout"}},
},
ClientIdentification: httpsecurity.ClientAuthenticated,
Semantics: httpsecurity.ProxySemantics,
}, nil)
if err != nil {
t.Fatalf("New HTTP protection: %v", err)
}
handler, err := New(Config{}, Dependencies{
Auth: authentication,
Targets: fakeTargets{evaluateURL: func(context.Context, string) (policy.Authority, error) {
return policy.Authority{Host: "example.test", Port: 80}, nil
}},
Router: RouteFunc(func(*http.Request) (dispatch.Request, error) {
return dispatch.Request{RoutingName: "catalog", Action: routing.ActionDirect}, nil
}),
Dispatcher: DispatcherFunc(func(dispatch.Request) (*dispatch.Lease, error) {
t.Fatal("dispatcher must not run for a route outside the credential policy")
return nil, nil
}),
Transport: &fakeTransport{},
})
if err != nil {
t.Fatalf("New gateway handler: %v", err)
}
request := httptest.NewRequest(http.MethodGet, "http://example.test/catalog", nil)
request.Header.Set("Proxy-Authorization", "Bearer gateway-token")
response := httptest.NewRecorder()
handler.ServeHTTP(response, request)
if response.Code != http.StatusForbidden {
t.Fatalf("status = %d, want 403", response.Code)
}
}
func TestHandlerEvaluatesTargetPolicyBeforeStaticDirectRoute(t *testing.T) {
t.Parallel()
handler, err := New(Config{}, Dependencies{
Targets: fakeTargets{evaluateURL: func(context.Context, string) (policy.Authority, error) {
return policy.Authority{}, policy.ErrTargetDenied
}},
Router: RouteFunc(func(*http.Request) (dispatch.Request, error) {
t.Fatal("router must not run before target policy")
return dispatch.Request{}, nil
}),
Dispatcher: DispatcherFunc(func(dispatch.Request) (*dispatch.Lease, error) {
t.Fatal("dispatcher must not run for a denied target")
return nil, nil
}),
Transport: &fakeTransport{directRoundTrip: func(context.Context, *http.Request) (*http.Response, error) {
t.Fatal("direct transport must not run for a denied target")
return nil, nil
}},
})
if err != nil {
t.Fatalf("New(): %v", err)
}
response := httptest.NewRecorder()
handler.ServeHTTP(response, httptest.NewRequest(http.MethodGet, "http://blocked.example/resource", nil))
if response.Code != http.StatusForbidden {
t.Fatalf("status = %d, want 403", response.Code)
}
}
func TestHandlerReleasesCredentialConcurrencyAfterRequest(t *testing.T) {
t.Parallel()
protection, err := BuildProtection(config.Listener{
Auth: config.Auth{Mode: "bearer", Token: "gateway-token",
ClientPolicy: clientpolicy.Policy{MaxConcurrentConnections: 1}},
})
if err != nil {
t.Fatalf("BuildProtection() error = %v", err)
}
dispatcher, view := dispatcherWithProxies(t, "proxy-a")
started := make(chan struct{})
release := make(chan struct{})
transport := &fakeTransport{roundTrip: func(
_ context.Context, _ proxyDomain.Proxy, _ *http.Request, commit ...func() error,
) (*http.Response, error) {
if err := commit[0](); err != nil {
return nil, err
}
select {
case started <- struct{}{}:
default:
}
<-release
return &http.Response{StatusCode: http.StatusNoContent, Header: make(http.Header), Body: http.NoBody}, nil
}}
handler, err := New(Config{}, Dependencies{
Auth: protection.Auth, Admission: protection.Admission, Targets: fakeTargets{},
Router: RouteFunc(func(*http.Request) (dispatch.Request, error) {
return dispatch.Request{Upstreams: []string{"provider-a"}}, nil
}),
Dispatcher: dispatcher, Transport: transport,
})
if err != nil {
t.Fatalf("New() error = %v", err)
}
request := func() *http.Request {
result := httptest.NewRequest(http.MethodGet, "http://example.test/resource", nil)
result.Header.Set("Proxy-Authorization", "Bearer gateway-token")
return result
}
first := httptest.NewRecorder()
done := make(chan struct{})
go func() {
handler.ServeHTTP(first, request())
close(done)
}()
<-started
second := httptest.NewRecorder()
handler.ServeHTTP(second, request())
if second.Code != http.StatusTooManyRequests {
t.Fatalf("second status = %d, want 429", second.Code)
}
close(release)
<-done
third := httptest.NewRecorder()
handler.ServeHTTP(third, request())
if third.Code != http.StatusNoContent {
t.Fatalf("third status = %d, want 204", third.Code)
}
assertNoLeakedCapacity(t, view)
}
func TestHandlerPinsAuthenticatedSessionToCommittedProxyAndStripsHeader(t *testing.T) {
t.Parallel()
dispatcher, view := dispatcherWithProxies(t, "proxy-a", "proxy-b")
auth, err := httpsecurity.New(httpsecurity.Config{
Authentication: httpsecurity.Authentication{Mode: httpsecurity.ModeBearer, Token: "gateway-token"},
ClientIdentification: httpsecurity.ClientAuthenticated,
Semantics: httpsecurity.ProxySemantics,
}, nil)
if err != nil {
t.Fatalf("New HTTP protection: %v", err)
}
transport := &fakeTransport{roundTrip: func(
_ context.Context,
_ proxyDomain.Proxy,
request *http.Request,
commit ...func() error,
) (*http.Response, error) {
if request.Header.Get("X-Proxy-Session") != "" {
t.Fatal("sticky session header was forwarded upstream")
}
if err := commit[0](); err != nil {
return nil, err
}
return &http.Response{StatusCode: http.StatusNoContent, Header: make(http.Header), Body: http.NoBody}, nil
}}
handler, err := New(Config{
StickySession: StickySessionConfig{Header: "X-Proxy-Session", TTL: time.Minute, MaxEntries: 100},
}, Dependencies{
Auth: auth,
Targets: fakeTargets{evaluateURL: func(context.Context, string) (policy.Authority, error) {
return policy.Authority{Host: "example.test", Port: 80}, nil
}},
Router: RouteFunc(func(*http.Request) (dispatch.Request, error) {
return dispatch.Request{RoutingName: "catalog", Upstreams: []string{"provider-a"}}, nil
}),
Dispatcher: dispatcher,
Transport: transport,
})
if err != nil {
t.Fatalf("New gateway handler: %v", err)
}
for range 2 {
request := httptest.NewRequest(http.MethodGet, "http://example.test/catalog", nil)
request.Header.Set("Proxy-Authorization", "Bearer gateway-token")
request.Header.Set("X-Proxy-Session", "order-182736")
response := httptest.NewRecorder()
handler.ServeHTTP(response, request)
if response.Code != http.StatusNoContent {
t.Fatalf("status = %d, want 204", response.Code)
}
}
if got := strings.Join(transport.attempts(), ","); got != "proxy-a,proxy-a" {
t.Fatalf("selected proxies = %q, want proxy-a,proxy-a", got)
}
assertNoLeakedCapacity(t, view)
}
func TestHandlerRebindsStickySessionWhenSnapshotDropsBoundProxy(t *testing.T) {
t.Parallel()
store := snapshot.NewStore("cluster-a", "worker-a")
expiresAt := time.Now().Add(time.Minute).UTC()
first := []proxyDomain.Proxy{
{ID: "proxy-a", Scheme: proxyDomain.SchemeHTTP, SourceUpstream: "provider-a", State: proxyDomain.StateAvailable, MaxConcurrency: 1, ExpiresAt: &expiresAt},
{ID: "proxy-b", Scheme: proxyDomain.SchemeHTTP, SourceUpstream: "provider-a", State: proxyDomain.StateAvailable, MaxConcurrency: 1, ExpiresAt: &expiresAt},
}
applySnapshot := func(version uint64, proxies []proxyDomain.Proxy) {
t.Helper()
envelope := snapshot.Envelope{
ClusterID: "cluster-a", WorkerID: "worker-a", Epoch: 1, Version: version, Full: true, Proxies: proxies,
}
envelope.Checksum = snapshot.Checksum(proxies)
if err := store.Apply(envelope); err != nil {
t.Fatalf("Apply(%d): %v", version, err)
}
}
applySnapshot(1, first)
auth, err := httpsecurity.New(httpsecurity.Config{
Authentication: httpsecurity.Authentication{Mode: httpsecurity.ModeBearer, Token: "gateway-token"},
ClientIdentification: httpsecurity.ClientAuthenticated,
Semantics: httpsecurity.ProxySemantics,
}, nil)
if err != nil {
t.Fatalf("New HTTP protection: %v", err)
}
transport := &fakeTransport{roundTrip: func(
_ context.Context,
_ proxyDomain.Proxy,
_ *http.Request,
commit ...func() error,
) (*http.Response, error) {
if err := commit[0](); err != nil {
return nil, err
}
return &http.Response{StatusCode: http.StatusNoContent, Header: make(http.Header), Body: http.NoBody}, nil
}}
handler, err := New(Config{
StickySession: StickySessionConfig{Header: "X-Proxy-Session", TTL: time.Minute, MaxEntries: 100},
}, Dependencies{
Auth: auth,
Targets: fakeTargets{evaluateURL: func(context.Context, string) (policy.Authority, error) {
return policy.Authority{Host: "example.test", Port: 80}, nil
}},
Router: RouteFunc(func(*http.Request) (dispatch.Request, error) {
return dispatch.Request{RoutingName: "catalog", Upstreams: []string{"provider-a"}}, nil
}),
Dispatcher: dispatch.New(store),
Transport: transport,
})
if err != nil {
t.Fatalf("New gateway handler: %v", err)
}
serve := func() {
t.Helper()
request := httptest.NewRequest(http.MethodGet, "http://example.test/catalog", nil)
request.Header.Set("Proxy-Authorization", "Bearer gateway-token")
request.Header.Set("X-Proxy-Session", "order-182736")
response := httptest.NewRecorder()
handler.ServeHTTP(response, request)
if response.Code != http.StatusNoContent {
t.Fatalf("status = %d, want 204", response.Code)
}
}
serve()
applySnapshot(2, first[1:])
serve()
if got := strings.Join(transport.attempts(), ","); got != "proxy-a,proxy-b" {
t.Fatalf("selected proxies = %q, want proxy-a,proxy-b", got)
}
assertNoLeakedCapacity(t, store.Current())
}
func TestHandlerRetriesGETWithAnotherProxyBeforeResponseCommit(t *testing.T) {
t.Parallel()
dispatcher, view := dispatcherWithProxies(t, "proxy-a", "proxy-b")
transport := &fakeTransport{}
transport.roundTrip = func(
_ context.Context,
selected proxyDomain.Proxy,
_ *http.Request,
commit ...func() error,
) (*http.Response, error) {
if selected.ID == "proxy-a" {
return nil, errors.New("dial failed")
}
if err := commit[0](); err != nil {
return nil, err
}
return &http.Response{
StatusCode: http.StatusOK,
Header: make(http.Header),
Body: io.NopCloser(strings.NewReader("ok")),
}, nil
}
handler := newTestHandler(t, Config{MaxAttempts: 2, RetryMethods: []string{http.MethodGet}}, dispatcher, transport)
response := httptest.NewRecorder()
handler.ServeHTTP(response, httptest.NewRequest(http.MethodGet, "http://example.test/resource", nil))
if response.Code != http.StatusOK || response.Body.String() != "ok" {
t.Fatalf("response = (%d, %q), want (200, ok)", response.Code, response.Body.String())
}
if got := strings.Join(transport.attempts(), ","); got != "proxy-a,proxy-b" {
t.Fatalf("proxy attempts = %q", got)
}
assertNoLeakedCapacity(t, view)
}
func TestHandlerRecordsOutcomeForEachProxyAttempt(t *testing.T) {
dispatcher, _ := dispatcherWithProxies(t, "proxy-a", "proxy-b")
transport := &fakeTransport{roundTrip: func(
_ context.Context,
selected proxyDomain.Proxy,
_ *http.Request,
commit ...func() error,
) (*http.Response, error) {
if selected.ID == "proxy-a" {
return nil, errors.New("dial failed")
}
if err := commit[0](); err != nil {
return nil, err
}
return &http.Response{StatusCode: http.StatusNoContent, Header: make(http.Header), Body: http.NoBody}, nil
}}
handler := newTestHandler(t, Config{MaxAttempts: 2, RetryMethods: []string{http.MethodGet}}, dispatcher, transport)
recorder := &outcomeRecorder{}
handler.outcomes = recorder
response := httptest.NewRecorder()
handler.ServeHTTP(response, httptest.NewRequest(http.MethodGet, "http://example.test/resource", nil))
events := recorder.Events()
if response.Code != http.StatusNoContent || len(events) != 2 {
t.Fatalf("status=%d events=%+v", response.Code, events)
}
if events[0].ProxyID != "proxy-a" || events[0].Stage != outcomeDomain.StageDial || events[0].Success ||
events[0].ErrorClass != outcomeDomain.ErrorClassDial || events[1].ProxyID != "proxy-b" ||
events[1].Stage != outcomeDomain.StageResponseHeaders || !events[1].Success {
t.Fatalf("outcomes = %+v", events)
}
}
func TestHandlerClampsNegativeOutcomeLatency(t *testing.T) {
recorder := &outcomeRecorder{}
handler := &Handler{outcomes: recorder}
handler.recordOutcome("proxy-a", "route-a", outcomeDomain.StageDial, true, nil, time.Now().Add(time.Second))
events := recorder.Events()
if len(events) != 1 || events[0].Latency != 0 {
t.Fatalf("outcomes = %+v, want zero latency", events)
}
}
func TestHandlerDoesNotRetryPOST(t *testing.T) {
t.Parallel()
dispatcher, view := dispatcherWithProxies(t, "proxy-a", "proxy-b")
transport := &fakeTransport{roundTrip: func(
_ context.Context,
selected proxyDomain.Proxy,
_ *http.Request,
_ ...func() error,
) (*http.Response, error) {
return nil, errors.New("dial failed through " + selected.ID)
}}
handler := newTestHandler(t, Config{MaxAttempts: 2, RetryMethods: []string{http.MethodGet}}, dispatcher, transport)
response := httptest.NewRecorder()
handler.ServeHTTP(response, httptest.NewRequest(http.MethodPost, "http://example.test/resource", strings.NewReader("body")))
if response.Code != http.StatusBadGateway {
t.Fatalf("status = %d, want 502", response.Code)
}
if len(transport.attempts()) != 1 {
t.Fatalf("attempts = %v, want one POST attempt", transport.attempts())
}
assertNoLeakedCapacity(t, view)
}
func TestHandlerReturns407WithoutRetry(t *testing.T) {
t.Parallel()
dispatcher, view := dispatcherWithProxies(t, "proxy-a", "proxy-b")
transport := &fakeTransport{roundTrip: func(
_ context.Context,
_ proxyDomain.Proxy,
_ *http.Request,
commit ...func() error,
) (*http.Response, error) {
if err := commit[0](); err != nil {
return nil, err
}
return &http.Response{
StatusCode: http.StatusProxyAuthRequired,
Header: http.Header{"Proxy-Authenticate": []string{"Basic"}},
Body: io.NopCloser(strings.NewReader("denied")),
}, nil
}}
handler := newTestHandler(t, Config{MaxAttempts: 2, RetryMethods: []string{http.MethodGet}}, dispatcher, transport)
recorder := &outcomeRecorder{}
handler.outcomes = recorder
response := httptest.NewRecorder()
handler.ServeHTTP(response, httptest.NewRequest(http.MethodGet, "http://example.test/resource", nil))
if response.Code != http.StatusProxyAuthRequired || response.Body.String() != "denied" {
t.Fatalf("response = (%d, %q), want (407, denied)", response.Code, response.Body.String())
}
if len(transport.attempts()) != 1 {
t.Fatalf("attempts = %v, want no retry for 407", transport.attempts())
}
events := recorder.Events()
if len(events) != 1 || events[0].Success || events[0].Stage != outcomeDomain.StageResponseHeaders ||
events[0].ErrorClass != outcomeDomain.ErrorClassProxyResponse {
t.Fatalf("outcomes = %+v, want failed proxy response", events)
}
assertNoLeakedCapacity(t, view)
}
func TestHandlerPinsValidatedHTTPDestinationWithoutChangingHost(t *testing.T) {
t.Parallel()
dispatcher, view := dispatcherWithProxies(t, "proxy-a")
transport := &fakeTransport{roundTrip: func(
_ context.Context,
_ proxyDomain.Proxy,
request *http.Request,
commit ...func() error,
) (*http.Response, error) {
if request.URL.Host != "198.51.100.10:80" {
t.Fatalf("pinned URL host = %q", request.URL.Host)
}
if request.Host != "example.test" {
t.Fatalf("original Host = %q", request.Host)
}
if err := commit[0](); err != nil {
return nil, err
}
return &http.Response{StatusCode: http.StatusNoContent, Header: make(http.Header), Body: http.NoBody}, nil
}}
handler, err := New(Config{}, Dependencies{
Targets: fakeTargets{evaluateURL: func(context.Context, string) (policy.Authority, error) {
return policy.Authority{
Host: "example.test", Port: 80,
ResolvedIP: netip.MustParseAddr("198.51.100.10"),
}, nil
}},
Router: RouteFunc(func(*http.Request) (dispatch.Request, error) {
return dispatch.Request{Upstreams: []string{"provider-a"}}, nil
}),
Dispatcher: dispatcher,
Transport: transport,
})
if err != nil {
t.Fatalf("New() error = %v", err)
}
response := httptest.NewRecorder()
handler.ServeHTTP(response, httptest.NewRequest(http.MethodGet, "http://example.test/resource", nil))
if response.Code != http.StatusNoContent {
t.Fatalf("status = %d, want 204", response.Code)
}
assertNoLeakedCapacity(t, view)
}
func TestHandlerUsesDirectFallbackAfterTargetPolicy(t *testing.T) {
var directCalls int
transport := &fakeTransport{directRoundTrip: func(_ context.Context, request *http.Request) (*http.Response, error) {
directCalls++
if request.URL.Host != "198.51.100.10:80" || request.Host != "example.test" {
t.Fatalf("direct request = URL %q Host %q", request.URL.Host, request.Host)
}
return &http.Response{StatusCode: http.StatusNoContent, Header: make(http.Header), Body: http.NoBody}, nil
}}
handler, err := New(Config{}, Dependencies{
Targets: fakeTargets{evaluateURL: func(context.Context, string) (policy.Authority, error) {
return policy.Authority{Host: "example.test", Port: 80, ResolvedIP: netip.MustParseAddr("198.51.100.10")}, nil
}},
Router: RouteFunc(func(*http.Request) (dispatch.Request, error) {
return dispatch.Request{OnUnavailable: routing.OnUnavailableDirect}, nil
}),
Dispatcher: DispatcherFunc(func(dispatch.Request) (*dispatch.Lease, error) { return nil, dispatch.ErrNoCandidate }),
Transport: transport,
})
if err != nil {
t.Fatalf("New(): %v", err)
}
response := httptest.NewRecorder()
handler.ServeHTTP(response, httptest.NewRequest(http.MethodGet, "http://example.test/resource", nil))
if response.Code != http.StatusNoContent || directCalls != 1 {
t.Fatalf("response = %d, direct calls = %d", response.Code, directCalls)
}
}
func TestHandlerUsesStaticDirectRouteForHTTP(t *testing.T) {
t.Parallel()
var dispatcherCalls atomic.Int64
transport := &fakeTransport{directRoundTrip: func(_ context.Context, request *http.Request) (*http.Response, error) {
if request.URL.Host != "198.51.100.10:80" || request.Host != "example.test" {
t.Fatalf("direct request = URL %q Host %q", request.URL.Host, request.Host)
}
return &http.Response{StatusCode: http.StatusNoContent, Header: make(http.Header), Body: http.NoBody}, nil
}}
handler, err := New(Config{StickySession: StickySessionConfig{
Header: "X-Proxy-Session", TTL: time.Minute, MaxEntries: 10,
}}, Dependencies{
Targets: fakeTargets{evaluateURL: func(context.Context, string) (policy.Authority, error) {
return policy.Authority{Host: "example.test", Port: 80, ResolvedIP: netip.MustParseAddr("198.51.100.10")}, nil
}},
Router: RouteFunc(func(*http.Request) (dispatch.Request, error) {
return dispatch.Request{RoutingName: "direct-api", Action: routing.ActionDirect}, nil
}),
Dispatcher: DispatcherFunc(func(dispatch.Request) (*dispatch.Lease, error) {
dispatcherCalls.Add(1)
return nil, dispatch.ErrNoCandidate
}),
Transport: transport,
})
if err != nil {
t.Fatalf("New(): %v", err)
}
recorder := &outcomeRecorder{}
handler.outcomes = recorder
response := httptest.NewRecorder()
handler.ServeHTTP(response, httptest.NewRequest(http.MethodGet, "http://example.test/resource", nil))
if response.Code != http.StatusNoContent || dispatcherCalls.Load() != 0 || len(recorder.Events()) != 0 {
t.Fatalf("response = %d, dispatcher calls = %d, outcomes = %+v", response.Code, dispatcherCalls.Load(), recorder.Events())
}
}
func TestHandlerUsesStaticDirectRouteForCONNECT(t *testing.T) {
t.Parallel()
var dispatcherCalls atomic.Int64
var directCalls atomic.Int64
transport := &fakeTransport{
directTunnel: func(_ context.Context, target string) (net.Conn, error) {
if target != "198.51.100.10:443" {
t.Fatalf("direct target = %q", target)
}
directCalls.Add(1)
upstream, peer := net.Pipe()
_ = peer.Close()
return upstream, nil
},
relay: func(context.Context, net.Conn, net.Conn) error { return nil },
}
handler, err := New(Config{}, Dependencies{
Targets: fakeTargets{evaluateConnect: func(context.Context, string) (policy.Authority, error) {
return policy.Authority{Host: "example.test", Port: 443, ResolvedIP: netip.MustParseAddr("198.51.100.10")}, nil
}},
Router: RouteFunc(func(*http.Request) (dispatch.Request, error) {
return dispatch.Request{RoutingName: "direct-connect", Action: routing.ActionDirect}, nil
}),
Dispatcher: DispatcherFunc(func(dispatch.Request) (*dispatch.Lease, error) {
dispatcherCalls.Add(1)
return nil, dispatch.ErrNoCandidate
}),
Transport: transport,
})
if err != nil {
t.Fatalf("New(): %v", err)
}
recorder := &outcomeRecorder{}
handler.outcomes = recorder
gateway := httptest.NewServer(handler)
defer gateway.Close()
response := sendConnect(t, gateway.URL, "example.test:443")
defer response.Body.Close()
if response.StatusCode != http.StatusOK || directCalls.Load() != 1 || dispatcherCalls.Load() != 0 || len(recorder.Events()) != 0 {
t.Fatalf("response = %d, direct calls = %d, dispatcher calls = %d, outcomes = %+v",
response.StatusCode, directCalls.Load(), dispatcherCalls.Load(), recorder.Events())
}
}
func TestHandlerWaitsOnlyForWaitRoutingAction(t *testing.T) {
dispatcher := &waitRecordingDispatcher{}
handler := &Handler{dispatcher: dispatcher}
_, err := handler.acquireRoute(context.Background(), dispatch.Request{
OnUnavailable: routing.OnUnavailableWait, WaitTimeout: 25 * time.Millisecond,
})
if !errors.Is(err, dispatch.ErrNoCandidate) || dispatcher.waitTimeout != 25*time.Millisecond {
t.Fatalf("acquireRoute() error = %v, wait timeout = %s", err, dispatcher.waitTimeout)
}
}
func TestHandlerEnforcesConcurrentRequestLimit(t *testing.T) {
t.Parallel()
dispatcher, view := dispatcherWithProxies(t, "proxy-a")
started := make(chan struct{})
release := make(chan struct{})
transport := &fakeTransport{roundTrip: func(
_ context.Context,
_ proxyDomain.Proxy,
_ *http.Request,
commit ...func() error,
) (*http.Response, error) {
if err := commit[0](); err != nil {
return nil, err
}
close(started)
<-release
return &http.Response{StatusCode: http.StatusNoContent, Header: make(http.Header), Body: http.NoBody}, nil
}}
handler := newTestHandler(t, Config{MaxAttempts: 1, MaxConcurrentRequests: 1}, dispatcher, transport)
first := httptest.NewRecorder()
firstDone := make(chan struct{})
go func() {
handler.ServeHTTP(first, httptest.NewRequest(http.MethodGet, "http://example.test/first", nil))
close(firstDone)
}()
<-started
second := httptest.NewRecorder()
handler.ServeHTTP(second, httptest.NewRequest(http.MethodGet, "http://example.test/second", nil))
if second.Code != http.StatusServiceUnavailable {
t.Fatalf("second status = %d, want 503", second.Code)
}
close(release)
<-firstDone
assertNoLeakedCapacity(t, view)
}
func TestHandlerRetriesCONNECTOnlyBeforeClientSuccess(t *testing.T) {
t.Parallel()
dispatcher, view := dispatcherWithProxies(t, "proxy-a", "proxy-b")
transport := &fakeTransport{}
transport.openTunnel = func(_ context.Context, selected proxyDomain.Proxy, target string) (net.Conn, error) {
if target != "example.test:443" {
t.Fatalf("CONNECT target = %q", target)
}
if selected.ID == "proxy-a" {
return nil, errors.New("handshake failed")
}
gateway, peer := net.Pipe()
_ = peer.Close()
return gateway, nil
}
transport.relay = func(context.Context, net.Conn, net.Conn) error {
return errors.New("tunnel broke after client commit")
}
handler := newTestHandler(t, Config{MaxAttempts: 2, RetryMethods: []string{http.MethodConnect}}, dispatcher, transport)
gateway := httptest.NewServer(handler)
defer gateway.Close()
response := sendConnect(t, gateway.URL, "example.test:443")
defer response.Body.Close()
if response.StatusCode != http.StatusOK {
t.Fatalf("CONNECT status = %d, want 200", response.StatusCode)
}
if got := strings.Join(transport.attempts(), ","); got != "proxy-a,proxy-b" {
t.Fatalf("CONNECT proxy attempts = %q", got)
}
assertNoLeakedCapacity(t, view)
}
func TestHandlerReturnsCONNECT407WithoutRetry(t *testing.T) {
t.Parallel()
dispatcher, view := dispatcherWithProxies(t, "proxy-a", "proxy-b")
transport := &fakeTransport{openTunnel: func(context.Context, proxyDomain.Proxy, string) (net.Conn, error) {
return nil, &transportDomain.ProxyResponseError{
StatusCode: http.StatusProxyAuthRequired,
Status: "407 Proxy Authentication Required",
Header: http.Header{"Proxy-Authenticate": []string{"Basic"}},
Body: []byte("denied"),
}
}}
handler := newTestHandler(t, Config{MaxAttempts: 2, RetryMethods: []string{http.MethodConnect}}, dispatcher, transport)
response := httptest.NewRecorder()
request := httptest.NewRequest(http.MethodConnect, "http://example.test:443", nil)
request.Host = "example.test:443"
handler.ServeHTTP(response, request)
if response.Code != http.StatusProxyAuthRequired || strings.TrimSpace(response.Body.String()) != "denied" {
t.Fatalf("CONNECT response = (%d, %q), want (407, denied)", response.Code, response.Body.String())
}
if len(transport.attempts()) != 1 {
t.Fatalf("CONNECT attempts = %v, want no retry for 407", transport.attempts())
}
assertNoLeakedCapacity(t, view)
}
func TestHandlerRetriesRetryableCONNECTResponseBeforeClientSuccess(t *testing.T) {
t.Parallel()
dispatcher, view := dispatcherWithProxies(t, "proxy-a", "proxy-b")
transport := &fakeTransport{openTunnel: func(_ context.Context, selected proxyDomain.Proxy, _ string) (net.Conn, error) {
if selected.ID == "proxy-a" {
return nil, &transportDomain.ProxyResponseError{
StatusCode: http.StatusServiceUnavailable,
Status: "503 Service Unavailable",
Header: make(http.Header),
Body: []byte("retry"),
}
}
gateway, peer := net.Pipe()
_ = peer.Close()
return gateway, nil
}}
transport.relay = func(context.Context, net.Conn, net.Conn) error { return nil }
handler := newTestHandler(t, Config{MaxAttempts: 2, RetryMethods: []string{http.MethodConnect}}, dispatcher, transport)
gateway := httptest.NewServer(handler)
defer gateway.Close()
response := sendConnect(t, gateway.URL, "example.test:443")
defer response.Body.Close()
if response.StatusCode != http.StatusOK {
t.Fatalf("CONNECT status = %d, want 200", response.StatusCode)
}
if got := strings.Join(transport.attempts(), ","); got != "proxy-a,proxy-b" {
t.Fatalf("CONNECT attempts = %q, want proxy-a,proxy-b", got)
}
assertNoLeakedCapacity(t, view)
}
func TestHandlerShutdownForcesHijackedTunnelsAtDeadlineAndRejectsNewRequests(t *testing.T) {
t.Parallel()
dispatcher, view := dispatcherWithProxies(t, "proxy-a")
peerConnections := make(chan net.Conn, 1)
relayStarted := make(chan struct{})
transport := &fakeTransport{}
transport.openTunnel = func(context.Context, proxyDomain.Proxy, string) (net.Conn, error) {
gateway, peer := net.Pipe()
peerConnections <- peer
return gateway, nil
}
transport.relay = func(_ context.Context, client, _ net.Conn) error {
close(relayStarted)
_, err := io.Copy(io.Discard, client)
return err
}
handler := newTestHandler(t, Config{MaxAttempts: 1}, dispatcher, transport)
gateway := httptest.NewServer(handler)
defer gateway.Close()
response := sendConnect(t, gateway.URL, "example.test:443")
defer response.Body.Close()
peer := <-peerConnections
defer peer.Close()
<-relayStarted
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Millisecond)
defer cancel()
if err := handler.Shutdown(ctx); !errors.Is(err, context.DeadlineExceeded) {
t.Fatalf("Shutdown() error = %v, want context deadline exceeded", err)
}
rejected := httptest.NewRecorder()
handler.ServeHTTP(rejected, httptest.NewRequest(http.MethodGet, "http://example.test/new", nil))
if rejected.Code != http.StatusServiceUnavailable {
t.Fatalf("request after shutdown status = %d, want 503", rejected.Code)
}
assertNoLeakedCapacity(t, view)
}
func TestHandlerShutdownWaitsForInFlightHTTP(t *testing.T) {
t.Parallel()
dispatcher, view := dispatcherWithProxies(t, "proxy-a")
started := make(chan struct{})
release := make(chan struct{})
transport := &fakeTransport{roundTrip: func(
_ context.Context,
_ proxyDomain.Proxy,
_ *http.Request,
commit ...func() error,
) (*http.Response, error) {
if err := commit[0](); err != nil {
return nil, err
}
close(started)
<-release
return &http.Response{StatusCode: http.StatusNoContent, Header: make(http.Header), Body: http.NoBody}, nil
}}
handler := newTestHandler(t, Config{}, dispatcher, transport)
requestDone := make(chan struct{})
go func() {
handler.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "http://example.test/resource", nil))
close(requestDone)
}()
<-started
shutdownDone := make(chan error, 1)
go func() { shutdownDone <- handler.Shutdown(context.Background()) }()
select {
case err := <-shutdownDone:
t.Fatalf("Shutdown() returned before HTTP completed: %v", err)
case <-time.After(20 * time.Millisecond):
}
close(release)
<-requestDone
if err := <-shutdownDone; err != nil {
t.Fatalf("Shutdown() error = %v", err)
}
assertNoLeakedCapacity(t, view)
}
type fakeTargets struct {
evaluateURL func(context.Context, string) (policy.Authority, error)
evaluateConnect func(context.Context, string) (policy.Authority, error)
}
func (targets fakeTargets) EvaluateURL(ctx context.Context, raw string) (policy.Authority, error) {
if targets.evaluateURL != nil {
return targets.evaluateURL(ctx, raw)
}
return policy.Authority{Host: "example.test", Port: 80}, nil
}
func (targets fakeTargets) EvaluateConnectAuthority(ctx context.Context, raw string) (policy.Authority, error) {
if targets.evaluateConnect != nil {
return targets.evaluateConnect(ctx, raw)
}
return policy.Authority{Host: "example.test", Port: 443}, nil
}
type fakeTransport struct {
mu sync.Mutex
seen []string
roundTrip func(context.Context, proxyDomain.Proxy, *http.Request, ...func() error) (*http.Response, error)
openTunnel func(context.Context, proxyDomain.Proxy, string) (net.Conn, error)
directRoundTrip func(context.Context, *http.Request) (*http.Response, error)
directTunnel func(context.Context, string) (net.Conn, error)
relay func(context.Context, net.Conn, net.Conn) error
}
type recordingRequestMetrics struct {
mu sync.Mutex
starts map[string]int
finishes map[string]int
opened int
closed int
}
func (metrics *recordingRequestMetrics) ObserveRequestStarted(protocol string) {
metrics.mu.Lock()
defer metrics.mu.Unlock()
if metrics.starts == nil {
metrics.starts = make(map[string]int)
}
metrics.starts[protocol]++
}
func (metrics *recordingRequestMetrics) ObserveRequestFinished(protocol string) {
metrics.mu.Lock()
defer metrics.mu.Unlock()
if metrics.finishes == nil {
metrics.finishes = make(map[string]int)
}
metrics.finishes[protocol]++
}
func (metrics *recordingRequestMetrics) ObserveTunnelOpened() {
metrics.mu.Lock()
defer metrics.mu.Unlock()
metrics.opened++
}
func (metrics *recordingRequestMetrics) ObserveTunnelClosed() {
metrics.mu.Lock()
defer metrics.mu.Unlock()
metrics.closed++
}
func (metrics *recordingRequestMetrics) started(protocol string) int {
metrics.mu.Lock()
defer metrics.mu.Unlock()
return metrics.starts[protocol]
}
func (metrics *recordingRequestMetrics) finished(protocol string) int {
metrics.mu.Lock()
defer metrics.mu.Unlock()
return metrics.finishes[protocol]
}
type outcomeRecorder struct {
mu sync.Mutex
events []outcomeDomain.Event
}
func (recorder *outcomeRecorder) Record(event outcomeDomain.Event) {
recorder.mu.Lock()
recorder.events = append(recorder.events, event)
recorder.mu.Unlock()
}
func (recorder *outcomeRecorder) Events() []outcomeDomain.Event {
recorder.mu.Lock()
defer recorder.mu.Unlock()
return append([]outcomeDomain.Event(nil), recorder.events...)
}
type waitRecordingDispatcher struct{ waitTimeout time.Duration }
func (dispatcher *waitRecordingDispatcher) Acquire(dispatch.Request) (*dispatch.Lease, error) {
return nil, dispatch.ErrNoCandidate
}
func (dispatcher *waitRecordingDispatcher) AcquireWait(_ context.Context, _ dispatch.Request, timeout time.Duration) (*dispatch.Lease, error) {
dispatcher.waitTimeout = timeout
return nil, dispatch.ErrNoCandidate
}
func (transport *fakeTransport) RoundTripDirect(ctx context.Context, request *http.Request) (*http.Response, error) {
if transport.directRoundTrip == nil {
return nil, errors.New("direct round trip not configured")
}
return transport.directRoundTrip(ctx, request)
}
func (transport *fakeTransport) OpenDirectTunnel(ctx context.Context, target string) (net.Conn, error) {
if transport.directTunnel == nil {
return nil, errors.New("direct tunnel not configured")
}
return transport.directTunnel(ctx, target)
}
func (transport *fakeTransport) RoundTrip(
ctx context.Context,
selected proxyDomain.Proxy,
request *http.Request,
commit ...func() error,
) (*http.Response, error) {
transport.record(selected.ID)
if transport.roundTrip == nil {
return nil, errors.New("round trip not configured")
}
return transport.roundTrip(ctx, selected, request, commit...)
}
func (transport *fakeTransport) OpenTunnel(ctx context.Context, selected proxyDomain.Proxy, target string) (net.Conn, error) {
transport.record(selected.ID)
if transport.openTunnel == nil {
return nil, errors.New("tunnel not configured")
}
return transport.openTunnel(ctx, selected, target)
}
func (transport *fakeTransport) Relay(ctx context.Context, left, right net.Conn) error {
if transport.relay == nil {
return errors.New("relay not configured")
}
return transport.relay(ctx, left, right)
}
func (transport *fakeTransport) record(id string) {
transport.mu.Lock()
defer transport.mu.Unlock()
transport.seen = append(transport.seen, id)
}
func (transport *fakeTransport) attempts() []string {
transport.mu.Lock()
defer transport.mu.Unlock()
return append([]string(nil), transport.seen...)
}
func newTestHandler(
t *testing.T,
config Config,
dispatcher Dispatcher,
transport ProxyTransport,
) *Handler {
t.Helper()
handler, err := New(config, Dependencies{
Targets: fakeTargets{},
Router: RouteFunc(func(*http.Request) (dispatch.Request, error) {
return dispatch.Request{Upstreams: []string{"provider-a"}}, nil
}),
Dispatcher: dispatcher,
Transport: transport,
})
if err != nil {
t.Fatalf("New() error = %v", err)
}
return handler
}
func dispatcherWithProxies(t *testing.T, ids ...string) (*dispatch.Dispatcher, *snapshot.View) {
t.Helper()
store := snapshot.NewStore("cluster-a", "worker-a")
proxies := make([]proxyDomain.Proxy, 0, len(ids))
for index, id := range ids {
proxies = append(proxies, proxyDomain.Proxy{
ID: id,
Scheme: proxyDomain.SchemeHTTP,
Host: "127.0.0.1",
Port: uint16(20000 + index),
SourceUpstream: "provider-a",
MaxConcurrency: 2,
State: proxyDomain.StateAvailable,
})
}
envelope := snapshot.Envelope{
ClusterID: "cluster-a",
WorkerID: "worker-a",
Epoch: 1,
Version: 1,
Full: true,
Proxies: proxies,
}
envelope.Checksum = snapshot.Checksum(proxies)
if err := store.Apply(envelope); err != nil {
t.Fatalf("apply snapshot: %v", err)
}
return dispatch.New(store), store.Current()
}
func assertNoLeakedCapacity(t *testing.T, view *snapshot.View) {
t.Helper()
deadline := time.Now().Add(500 * time.Millisecond)
for {
allReleased := true
for _, entry := range view.Entries {
if entry.Runtime.Active() != 0 || entry.Runtime.Reserved() != 0 {
allReleased = false
break
}
}
if allReleased {
return
}
if time.Now().After(deadline) {
for _, entry := range view.Entries {
if active, reserved := entry.Runtime.Active(), entry.Runtime.Reserved(); active != 0 || reserved != 0 {
t.Fatalf("proxy %s capacity leaked: active=%d reserved=%d", entry.Proxy.ID, active, reserved)
}
}
}
time.Sleep(time.Millisecond)
}
}
func sendConnect(t *testing.T, gatewayURL, target string) *http.Response {
t.Helper()
parsed, err := url.Parse(gatewayURL)
if err != nil {
t.Fatalf("parse gateway URL: %v", err)
}
connection, err := net.Dial("tcp", parsed.Host)
if err != nil {
t.Fatalf("dial gateway: %v", err)
}
t.Cleanup(func() { _ = connection.Close() })
if _, err := io.WriteString(connection, "CONNECT "+target+" HTTP/1.1\r\nHost: "+target+"\r\n\r\n"); err != nil {
t.Fatalf("write CONNECT request: %v", err)
}
response, err := http.ReadResponse(bufio.NewReader(connection), &http.Request{Method: http.MethodConnect})
if err != nil {
t.Fatalf("read CONNECT response: %v", err)
}
return response
}
var _ ProxyTransport = (*transportDomain.Transport)(nil)