package transport import ( "bufio" "context" "encoding/base64" "errors" "fmt" "io" "net" "net/http" "net/http/httptest" "net/url" "strings" "testing" "time" proxyDomain "proxy-pool/internal/domain/proxy" ) func TestRoundTripForwardsHTTPViaSelectedProxy(t *testing.T) { t.Parallel() requestSeen := make(chan *http.Request, 1) upstream := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { requestSeen <- request.Clone(request.Context()) writer.Header().Set("X-Upstream", "selected") writer.WriteHeader(http.StatusCreated) _, _ = writer.Write([]byte("forwarded")) })) defer upstream.Close() selected := proxyFromURL(t, upstream.URL) selected.ID = "proxy-a" selected.Username = "alice" selected.CredentialVersion = "v1" selected.SecretRef = "secret://proxy-a" client := New(Config{}, CredentialResolverFunc(func(context.Context, proxyDomain.Proxy) (Credentials, error) { return Credentials{Username: "alice", Password: "s3cret"}, nil })) request := httptest.NewRequest(http.MethodGet, "http://TARGET/resource?q=1", nil) response, err := client.RoundTrip(request.Context(), selected, request) if err != nil { t.Fatalf("RoundTrip() error = %v", err) } defer response.Body.Close() body, err := io.ReadAll(response.Body) if err != nil { t.Fatalf("read response body: %v", err) } if response.StatusCode != http.StatusCreated || string(body) != "forwarded" { t.Fatalf("response = (%d, %q), want (201, forwarded)", response.StatusCode, body) } seen := <-requestSeen if seen.RequestURI != "http://TARGET/resource?q=1" { t.Fatalf("proxy request URI = %q", seen.RequestURI) } wantAuth := "Basic " + base64.StdEncoding.EncodeToString([]byte("alice:s3cret")) if got := seen.Header.Get("Proxy-Authorization"); got != wantAuth { t.Fatalf("Proxy-Authorization = %q, want %q", got, wantAuth) } } func TestRoundTripDirectForwardsWithoutProxyAuthorization(t *testing.T) { t.Parallel() upstream := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { if request.Header.Get("Proxy-Authorization") != "" { t.Fatalf("Proxy-Authorization = %q, want empty", request.Header.Get("Proxy-Authorization")) } writer.WriteHeader(http.StatusNoContent) })) defer upstream.Close() request := httptest.NewRequest(http.MethodGet, upstream.URL+"/resource", nil) request.Header.Set("Proxy-Authorization", "Basic should-not-forward") response, err := New(Config{}, nil).RoundTripDirect(request.Context(), request) if err != nil { t.Fatalf("RoundTripDirect(): %v", err) } defer response.Body.Close() if response.StatusCode != http.StatusNoContent { t.Fatalf("status = %d, want 204", response.StatusCode) } } func TestNewAppliesConfiguredConnectionPoolLimits(t *testing.T) { t.Parallel() client := New(Config{ MaxIdleConns: 256, MaxIdleConnsPerHost: 24, MaxConnsPerHost: 12, }, nil) t.Cleanup(client.CloseIdleConnections) for name, item := range map[string]*http.Transport{ "proxy": client.client, "direct": client.direct, } { if item.MaxIdleConns != 256 { t.Fatalf("%s MaxIdleConns = %d, want 256", name, item.MaxIdleConns) } if item.MaxIdleConnsPerHost != 24 { t.Fatalf("%s MaxIdleConnsPerHost = %d, want 24", name, item.MaxIdleConnsPerHost) } if item.MaxConnsPerHost != 12 { t.Fatalf("%s MaxConnsPerHost = %d, want 12", name, item.MaxConnsPerHost) } } } func TestOpenDirectTunnelDialsTarget(t *testing.T) { t.Parallel() listener, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("Listen(): %v", err) } defer listener.Close() accepted := make(chan struct{}) go func() { connection, acceptErr := listener.Accept() if acceptErr == nil { _ = connection.Close() close(accepted) } }() connection, err := New(Config{DialTimeout: time.Second}, nil).OpenDirectTunnel(context.Background(), listener.Addr().String()) if err != nil { t.Fatalf("OpenDirectTunnel(): %v", err) } defer connection.Close() select { case <-accepted: case <-time.After(time.Second): t.Fatal("target listener did not accept direct tunnel") } } func TestRoundTripCommitsReservationAfterConnectionAcquisition(t *testing.T) { t.Parallel() allowResponse := make(chan struct{}) upstream := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { <-allowResponse writer.WriteHeader(http.StatusNoContent) })) defer upstream.Close() client := New(Config{}, nil) request := httptest.NewRequest(http.MethodGet, "http://TARGET/resource", nil) committed := make(chan struct{}, 1) done := make(chan error, 1) go func() { response, err := client.RoundTrip(request.Context(), proxyFromURL(t, upstream.URL), request, func() error { committed <- struct{}{} return nil }) if response != nil { _ = response.Body.Close() } done <- err }() select { case <-committed: case <-time.After(time.Second): t.Fatal("commit hook was not called after acquiring the proxy connection") } close(allowResponse) if err := <-done; err != nil { t.Fatalf("RoundTrip() error = %v", err) } } func TestOpenTunnelPreservesBytesBufferedAfterSuccessfulHandshake(t *testing.T) { t.Parallel() listener, requests := startConnectProxy(t, func(connection net.Conn) { _, _ = io.WriteString(connection, "HTTP/1.1 200 Connection Established\r\n\r\nREADY") }) selected := proxyFromAddress(listener.Addr().String()) selected.Username = "alice" client := New(Config{}, CredentialResolverFunc(func(context.Context, proxyDomain.Proxy) (Credentials, error) { return Credentials{Username: "alice", Password: "s3cret"}, nil })) connection, err := client.OpenTunnel(context.Background(), selected, "example.test:443") if err != nil { t.Fatalf("OpenTunnel() error = %v", err) } defer connection.Close() preface := make([]byte, len("READY")) if _, err := io.ReadFull(connection, preface); err != nil { t.Fatalf("read buffered tunnel bytes: %v", err) } if string(preface) != "READY" { t.Fatalf("tunnel preface = %q", preface) } request := <-requests if request.Method != http.MethodConnect || request.Host != "example.test:443" { t.Fatalf("CONNECT request = %s %s", request.Method, request.Host) } wantAuth := "Basic " + base64.StdEncoding.EncodeToString([]byte("alice:s3cret")) if got := request.Header.Get("Proxy-Authorization"); got != wantAuth { t.Fatalf("Proxy-Authorization = %q, want %q", got, wantAuth) } } func TestOpenTunnelReturnsBoundedProxyResponseError(t *testing.T) { t.Parallel() listener, _ := startConnectProxy(t, func(connection net.Conn) { _, _ = io.WriteString(connection, "HTTP/1.1 407 Proxy Authentication Required\r\nContent-Length: 8\r\nProxy-Authenticate: Basic\r\n\r\ndenied!!") }) client := New(Config{MaxErrorResponseBytes: 4}, nil) connection, err := client.OpenTunnel(context.Background(), proxyFromAddress(listener.Addr().String()), "example.test:443") if connection != nil { _ = connection.Close() t.Fatal("OpenTunnel() returned a connection for 407") } var responseError *ProxyResponseError if !errors.As(err, &responseError) { t.Fatalf("OpenTunnel() error = %T %v, want *ProxyResponseError", err, err) } if responseError.StatusCode != http.StatusProxyAuthRequired { t.Fatalf("status = %d, want 407", responseError.StatusCode) } if string(responseError.Body) != "deni" { t.Fatalf("bounded body = %q, want deni", responseError.Body) } } func TestProxyResponseErrorRetryable(t *testing.T) { t.Parallel() tests := []struct { statusCode int want bool }{ {statusCode: http.StatusProxyAuthRequired, want: false}, {statusCode: http.StatusRequestTimeout, want: true}, {statusCode: http.StatusTooEarly, want: true}, {statusCode: http.StatusTooManyRequests, want: true}, {statusCode: http.StatusInternalServerError, want: true}, {statusCode: http.StatusBadGateway, want: true}, {statusCode: http.StatusServiceUnavailable, want: true}, {statusCode: http.StatusGatewayTimeout, want: true}, {statusCode: http.StatusBadRequest, want: false}, } for _, tt := range tests { t.Run(fmt.Sprint(tt.statusCode), func(t *testing.T) { err := &ProxyResponseError{StatusCode: tt.statusCode} if got := err.Retryable(); got != tt.want { t.Fatalf("Retryable() = %t, want %t", got, tt.want) } }) } } func TestOpenTunnelRejectsOversizedHandshakeResponse(t *testing.T) { t.Parallel() listener, _ := startConnectProxy(t, func(connection net.Conn) { _, _ = io.WriteString(connection, "HTTP/1.1 200 Connection Established\r\nX-Large: "+strings.Repeat("x", 4096)+"\r\n\r\n") }) client := New(Config{MaxResponseHeaderBytes: 256}, nil) connection, err := client.OpenTunnel(context.Background(), proxyFromAddress(listener.Addr().String()), "example.test:443") if connection != nil { _ = connection.Close() t.Fatal("OpenTunnel() returned a connection for an oversized handshake") } if !errors.Is(err, ErrProxyResponseTooLarge) { t.Fatalf("OpenTunnel() error = %v, want ErrProxyResponseTooLarge", err) } } func TestOpenTunnelHonorsContextCancellationDuringHandshake(t *testing.T) { t.Parallel() listener, _ := startConnectProxy(t, func(connection net.Conn) { <-time.After(time.Second) }) client := New(Config{HandshakeTimeout: time.Second}, nil) ctx, cancel := context.WithTimeout(context.Background(), 30*time.Millisecond) defer cancel() started := time.Now() connection, err := client.OpenTunnel(ctx, proxyFromAddress(listener.Addr().String()), "example.test:443") if connection != nil { _ = connection.Close() } if !errors.Is(err, context.DeadlineExceeded) { t.Fatalf("OpenTunnel() error = %v, want context deadline exceeded", err) } if elapsed := time.Since(started); elapsed > 300*time.Millisecond { t.Fatalf("cancellation took %s", elapsed) } } func TestRelayPreservesTCPHalfCloseInBothDirections(t *testing.T) { t.Parallel() client, gatewayClient := tcpPair(t) gatewayUpstream, upstream := tcpPair(t) defer client.Close() defer gatewayClient.Close() defer gatewayUpstream.Close() defer upstream.Close() relay := New(Config{TunnelBufferBytes: 1024}, nil) relayDone := make(chan error, 1) go func() { relayDone <- relay.Relay(context.Background(), gatewayClient, gatewayUpstream) }() if _, err := io.WriteString(client, "request"); err != nil { t.Fatalf("write client request: %v", err) } if err := client.CloseWrite(); err != nil { t.Fatalf("half-close client: %v", err) } request, err := io.ReadAll(upstream) if err != nil { t.Fatalf("read upstream request: %v", err) } if string(request) != "request" { t.Fatalf("upstream request = %q", request) } if _, err := io.WriteString(upstream, "response"); err != nil { t.Fatalf("write upstream response: %v", err) } if err := upstream.CloseWrite(); err != nil { t.Fatalf("half-close upstream: %v", err) } response, err := io.ReadAll(client) if err != nil { t.Fatalf("read client response: %v", err) } if string(response) != "response" { t.Fatalf("client response = %q", response) } select { case err := <-relayDone: if err != nil { t.Fatalf("Relay() error = %v", err) } case <-time.After(time.Second): t.Fatal("Relay() did not finish after both half-closes") } } func TestRelayKeepsTunnelAliveWhileTrafficFlowsInOneDirection(t *testing.T) { t.Parallel() client, gatewayClient := tcpPair(t) gatewayUpstream, upstream := tcpPair(t) defer client.Close() defer gatewayClient.Close() defer gatewayUpstream.Close() defer upstream.Close() ctx, cancel := context.WithCancel(context.Background()) defer cancel() relay := New(Config{TunnelIdleTimeout: 100 * time.Millisecond}, nil) done := make(chan error, 1) go func() { done <- relay.Relay(ctx, gatewayClient, gatewayUpstream) }() if err := client.SetReadDeadline(time.Now().Add(time.Second)); err != nil { t.Fatalf("set client read deadline: %v", err) } for range 6 { time.Sleep(30 * time.Millisecond) if _, err := upstream.Write([]byte("x")); err != nil { t.Fatalf("write one-way tunnel traffic: %v", err) } buffer := make([]byte, 1) if _, err := io.ReadFull(client, buffer); err != nil { t.Fatalf("read one-way tunnel traffic: %v", err) } } select { case err := <-done: t.Fatalf("Relay() ended during one-way activity: %v", err) default: } cancel() select { case err := <-done: if !errors.Is(err, context.Canceled) { t.Fatalf("Relay() error = %v, want context canceled", err) } case <-time.After(time.Second): t.Fatal("Relay() did not stop after cancellation") } } func TestRelayStopsAnIdleTunnelAtConfiguredDeadline(t *testing.T) { t.Parallel() left, leftPeer := net.Pipe() right, rightPeer := net.Pipe() defer leftPeer.Close() defer rightPeer.Close() relay := New(Config{TunnelIdleTimeout: 30 * time.Millisecond}, nil) done := make(chan error, 1) go func() { done <- relay.Relay(context.Background(), left, right) }() select { case err := <-done: if err == nil { t.Fatal("Relay() error = nil, want idle timeout") } case <-time.After(300 * time.Millisecond): t.Fatal("Relay() did not enforce tunnel idle timeout") } } func startConnectProxy(t *testing.T, respond func(net.Conn)) (net.Listener, <-chan *http.Request) { t.Helper() listener, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("listen: %v", err) } t.Cleanup(func() { _ = listener.Close() }) requests := make(chan *http.Request, 1) go func() { connection, acceptErr := listener.Accept() if acceptErr != nil { return } defer connection.Close() request, readErr := http.ReadRequest(bufio.NewReader(connection)) if readErr != nil { return } requests <- request respond(connection) }() return listener, requests } func proxyFromURL(t *testing.T, rawURL string) proxyDomain.Proxy { t.Helper() parsed, err := url.Parse(rawURL) if err != nil { t.Fatalf("parse proxy URL: %v", err) } return proxyFromAddress(parsed.Host) } func proxyFromAddress(address string) proxyDomain.Proxy { host, portText, err := net.SplitHostPort(address) if err != nil { panic(fmt.Sprintf("split proxy address %q: %v", address, err)) } var port uint16 if _, err := fmt.Sscanf(portText, "%d", &port); err != nil { panic(fmt.Sprintf("parse proxy port %q: %v", portText, err)) } return proxyDomain.Proxy{ ID: strings.ReplaceAll(address, ":", "-"), Scheme: proxyDomain.SchemeHTTP, Host: host, Port: port, MaxConcurrency: 1, } } func tcpPair(t *testing.T) (*net.TCPConn, *net.TCPConn) { t.Helper() listener, err := net.ListenTCP("tcp", &net.TCPAddr{IP: net.ParseIP("127.0.0.1")}) if err != nil { t.Fatalf("listen TCP pair: %v", err) } defer listener.Close() accepted := make(chan *net.TCPConn, 1) acceptErrors := make(chan error, 1) go func() { connection, acceptErr := listener.AcceptTCP() if acceptErr != nil { acceptErrors <- acceptErr return } accepted <- connection }() client, err := net.DialTCP("tcp", nil, listener.Addr().(*net.TCPAddr)) if err != nil { t.Fatalf("dial TCP pair: %v", err) } select { case server := <-accepted: return client, server case acceptErr := <-acceptErrors: _ = client.Close() t.Fatalf("accept TCP pair: %v", acceptErr) } return nil, nil }