305 lines
9.3 KiB
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
|
|
}
|