516 lines
15 KiB
Go
516 lines
15 KiB
Go
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
|
|
}
|