proxy-pool/internal/gateway/server/e2e_test.go

305 lines
9.3 KiB
Go

package server
import (
"bufio"
"context"
"errors"
"io"
"net"
"net/http"
"net/http/httptest"
"net/netip"
"net/url"
"strconv"
"strings"
"sync/atomic"
"testing"
"time"
proxyDomain "github.com/proxy-pool/proxy-pool/internal/domain/proxy"
"github.com/proxy-pool/proxy-pool/internal/gateway/dispatch"
"github.com/proxy-pool/proxy-pool/internal/gateway/policy"
"github.com/proxy-pool/proxy-pool/internal/gateway/snapshot"
transportDomain "github.com/proxy-pool/proxy-pool/internal/gateway/transport"
)
func TestHTTPProxyEndToEndRetriesDialFailure(t *testing.T) {
t.Parallel()
requestURI := make(chan string, 1)
goodProxy := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
requestURI <- request.RequestURI
writer.WriteHeader(http.StatusOK)
_, _ = writer.Write([]byte("through-good-proxy"))
}))
defer goodProxy.Close()
closedAddress := reserveClosedAddress(t)
dispatcher, view := dispatcherWithDescriptors(t,
proxyDescriptor(t, "proxy-a", "http://"+closedAddress),
proxyDescriptor(t, "proxy-b", goodProxy.URL),
)
proxyTransport := transportDomain.New(transportDomain.Config{
DialTimeout: 100 * time.Millisecond,
ResponseHeaderTimeout: time.Second,
}, nil)
defer proxyTransport.CloseIdleConnections()
handler := newTestHandler(t, Config{
MaxAttempts: 2,
RetryMethods: []string{http.MethodGet},
}, dispatcher, proxyTransport)
response := httptest.NewRecorder()
handler.ServeHTTP(response, httptest.NewRequest(http.MethodGet, "http://TARGET/resource?q=1", nil))
if response.Code != http.StatusOK || response.Body.String() != "through-good-proxy" {
t.Fatalf("response = (%d, %q)", response.Code, response.Body.String())
}
if got := <-requestURI; got != "http://TARGET/resource?q=1" {
t.Fatalf("upstream request URI = %q", got)
}
assertNoLeakedCapacity(t, view)
}
func TestHTTPProxyEndToEndDoesNotRetry407(t *testing.T) {
t.Parallel()
var deniedCalls atomic.Int64
deniedProxy := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
deniedCalls.Add(1)
writer.Header().Set("Proxy-Authenticate", "Basic")
writer.WriteHeader(http.StatusProxyAuthRequired)
_, _ = writer.Write([]byte("denied"))
}))
defer deniedProxy.Close()
var fallbackCalls atomic.Int64
fallbackProxy := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
fallbackCalls.Add(1)
writer.WriteHeader(http.StatusOK)
}))
defer fallbackProxy.Close()
dispatcher, view := dispatcherWithDescriptors(t,
proxyDescriptor(t, "proxy-a", deniedProxy.URL),
proxyDescriptor(t, "proxy-b", fallbackProxy.URL),
)
proxyTransport := transportDomain.New(transportDomain.Config{}, nil)
defer proxyTransport.CloseIdleConnections()
handler := newTestHandler(t, Config{MaxAttempts: 2, RetryMethods: []string{http.MethodGet}}, dispatcher, proxyTransport)
response := httptest.NewRecorder()
handler.ServeHTTP(response, httptest.NewRequest(http.MethodGet, "http://TARGET/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 deniedCalls.Load() != 1 || fallbackCalls.Load() != 0 {
t.Fatalf("proxy calls = denied:%d fallback:%d", deniedCalls.Load(), fallbackCalls.Load())
}
assertNoLeakedCapacity(t, view)
}
func TestHTTPProxyEndToEndBlocksPrivateTargetBeforeDispatch(t *testing.T) {
t.Parallel()
targets, err := policy.NewTargetPolicy(policy.Config{Resolver: staticResolver{
addresses: []netip.Addr{netip.MustParseAddr("127.0.0.1")},
}})
if err != nil {
t.Fatalf("NewTargetPolicy() error = %v", err)
}
var dispatchCalls atomic.Int64
handler, err := New(Config{}, Dependencies{
Targets: targets,
Router: RouteFunc(func(*http.Request) (dispatch.Request, error) {
t.Fatal("routing must not run for a blocked target")
return dispatch.Request{}, nil
}),
Dispatcher: DispatcherFunc(func(dispatch.Request) (*dispatch.Lease, error) {
dispatchCalls.Add(1)
return nil, dispatch.ErrNoCandidate
}),
Transport: &fakeTransport{},
})
if err != nil {
t.Fatalf("New() error = %v", err)
}
response := httptest.NewRecorder()
handler.ServeHTTP(response, httptest.NewRequest(http.MethodGet, "http://127.0.0.1/admin", nil))
if response.Code != http.StatusForbidden {
t.Fatalf("status = %d, want 403", response.Code)
}
if dispatchCalls.Load() != 0 {
t.Fatalf("dispatch calls = %d, want 0", dispatchCalls.Load())
}
}
func TestCONNECTEndToEndRelaysDataAndHalfClose(t *testing.T) {
t.Parallel()
upstream, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("listen fake CONNECT proxy: %v", err)
}
defer upstream.Close()
upstreamDone := make(chan error, 1)
go func() {
connection, acceptErr := upstream.Accept()
if acceptErr != nil {
upstreamDone <- acceptErr
return
}
defer connection.Close()
request, readErr := http.ReadRequest(bufio.NewReader(connection))
if readErr != nil {
upstreamDone <- readErr
return
}
if request.Method != http.MethodConnect || request.Host != "example.test:443" {
upstreamDone <- errors.New("unexpected upstream CONNECT target: " + request.Host)
return
}
if _, writeErr := io.WriteString(connection, "HTTP/1.1 200 Connection Established\r\n\r\n"); writeErr != nil {
upstreamDone <- writeErr
return
}
buffer := make([]byte, 1024)
for {
count, readErr := connection.Read(buffer)
if count > 0 {
if _, writeErr := connection.Write(buffer[:count]); writeErr != nil {
upstreamDone <- writeErr
return
}
}
if readErr != nil {
if readErr == io.EOF {
upstreamDone <- nil
} else {
upstreamDone <- readErr
}
return
}
}
}()
selected := proxyDescriptor(t, "proxy-a", "http://"+upstream.Addr().String())
dispatcher, view := dispatcherWithDescriptors(t, selected)
proxyTransport := transportDomain.New(transportDomain.Config{TunnelIdleTimeout: time.Second}, nil)
handler := newTestHandler(t, Config{MaxAttempts: 1}, dispatcher, proxyTransport)
gateway := httptest.NewServer(handler)
defer gateway.Close()
parsedGateway, _ := url.Parse(gateway.URL)
client, err := net.DialTCP("tcp", nil, tcpAddress(t, parsedGateway.Host))
if err != nil {
t.Fatalf("dial gateway: %v", err)
}
defer client.Close()
if _, err := io.WriteString(client, "CONNECT example.test:443 HTTP/1.1\r\nHost: example.test:443\r\n\r\n"); err != nil {
t.Fatalf("write client CONNECT: %v", err)
}
response, err := http.ReadResponse(bufio.NewReader(client), &http.Request{Method: http.MethodConnect})
if err != nil {
t.Fatalf("read gateway CONNECT response: %v", err)
}
defer response.Body.Close()
if response.StatusCode != http.StatusOK {
t.Fatalf("CONNECT status = %d, want 200", response.StatusCode)
}
if _, err := io.WriteString(client, "ping"); err != nil {
t.Fatalf("write tunnel data: %v", err)
}
echo := make([]byte, 4)
if _, err := io.ReadFull(response.Body, echo); err != nil {
t.Fatalf("read tunnel echo: %v", err)
}
if string(echo) != "ping" {
t.Fatalf("tunnel echo = %q", echo)
}
if err := client.CloseWrite(); err != nil {
t.Fatalf("half-close client tunnel: %v", err)
}
if _, err := io.Copy(io.Discard, response.Body); err != nil {
t.Fatalf("drain tunnel after half-close: %v", err)
}
if err := <-upstreamDone; err != nil {
t.Fatalf("fake upstream: %v", err)
}
assertNoLeakedCapacity(t, view)
}
type staticResolver struct{ addresses []netip.Addr }
func (resolver staticResolver) LookupNetIP(context.Context, string) ([]netip.Addr, error) {
return append([]netip.Addr(nil), resolver.addresses...), nil
}
func dispatcherWithDescriptors(t *testing.T, proxies ...proxyDomain.Proxy) (*dispatch.Dispatcher, *snapshot.View) {
t.Helper()
store := snapshot.NewStore("cluster-e2e", "worker-e2e")
envelope := snapshot.Envelope{
ClusterID: "cluster-e2e",
WorkerID: "worker-e2e",
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 proxyDescriptor(t *testing.T, id, rawURL string) proxyDomain.Proxy {
t.Helper()
parsed, err := url.Parse(rawURL)
if err != nil {
t.Fatalf("parse proxy URL: %v", err)
}
host, portText, err := net.SplitHostPort(parsed.Host)
if err != nil {
t.Fatalf("split proxy address: %v", err)
}
port, err := strconv.ParseUint(portText, 10, 16)
if err != nil {
t.Fatalf("parse proxy port: %v", err)
}
return proxyDomain.Proxy{
ID: id,
Scheme: proxyDomain.Scheme(strings.ToLower(parsed.Scheme)),
Host: host,
Port: uint16(port),
SourceUpstream: "provider-a",
MaxConcurrency: 2,
State: proxyDomain.StateAvailable,
CredentialVersion: "v1",
}
}
func reserveClosedAddress(t *testing.T) string {
t.Helper()
listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("reserve address: %v", err)
}
address := listener.Addr().String()
if err := listener.Close(); err != nil {
t.Fatalf("close reserved address: %v", err)
}
return address
}
func tcpAddress(t *testing.T, address string) *net.TCPAddr {
t.Helper()
resolved, err := net.ResolveTCPAddr("tcp", address)
if err != nil {
t.Fatalf("resolve TCP address: %v", err)
}
return resolved
}