proxy-pool/internal/gateway/transport/transport_test.go

438 lines
13 KiB
Go

package transport
import (
"bufio"
"context"
"encoding/base64"
"errors"
"fmt"
"io"
"net"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"testing"
"time"
proxyDomain "github.com/proxy-pool/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 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
}