package server import ( "bufio" "context" "errors" "io" "net" "net/http" "net/http/httptest" "net/netip" "net/url" "strconv" "strings" "sync/atomic" "testing" "time" proxyDomain "proxy-pool/internal/domain/proxy" "proxy-pool/internal/gateway/dispatch" "proxy-pool/internal/gateway/policy" "proxy-pool/internal/gateway/snapshot" transportDomain "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 }