proxy-pool/internal/gateway/server/handler_test.go
youfak c267f77eee
Some checks are pending
ci / proto (push) Waiting to run
ci / test (ubuntu-latest) (push) Waiting to run
ci / test (windows-latest) (push) Waiting to run
ci / race (push) Waiting to run
ci / integration (push) Waiting to run
feat: distribute credentials in worker snapshots
2026-07-31 16:46:47 +08:00

662 lines
22 KiB
Go

package server
import (
"bufio"
"context"
"errors"
"io"
"net"
"net/http"
"net/http/httptest"
"net/netip"
"net/url"
"strings"
"sync"
"testing"
"time"
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"
)
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 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 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)
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())
}
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 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)