1192 lines
40 KiB
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)
|