package server import ( "bufio" "context" "errors" "io" "net" "net/http" "net/http/httptest" "net/netip" "net/url" "strings" "sync" "testing" "time" "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 TestHandlerRejectsGatewayRoutingOutsideCredentialPolicy(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", Upstreams: []string{"provider-a"}}, 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 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 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 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)