File
Blob: tun/client/proxy_forwarded_test.go
| 1 | package client |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "crypto/tls" |
| 6 | "encoding/json" |
| 7 | "net" |
| 8 | "net/http" |
| 9 | "net/http/httptest" |
| 10 | "testing" |
| 11 | "time" |
| 12 | |
| 13 | "go.miragespace.co/specter/gateway" |
| 14 | "go.miragespace.co/specter/spec/protocol" |
| 15 | "go.miragespace.co/specter/spec/tun" |
| 16 | "go.miragespace.co/specter/util/bufconn" |
| 17 | "go.miragespace.co/specter/util/pipe" |
| 18 | |
| 19 | "github.com/quic-go/quic-go" |
| 20 | "github.com/stretchr/testify/require" |
| 21 | "github.com/zhangyunhao116/skipmap" |
| 22 | "go.uber.org/atomic" |
| 23 | "go.uber.org/zap/zaptest" |
| 24 | ) |
| 25 | |
| 26 | type forwardedRequest struct { |
| 27 | Host string |
| 28 | URI string |
| 29 | Headers http.Header |
| 30 | } |
| 31 | |
| 32 | func captureForwardedRequest(w http.ResponseWriter, r *http.Request) { |
| 33 | _ = json.NewEncoder(w).Encode(forwardedRequest{Host: r.Host, URI: r.RequestURI, Headers: r.Header}) |
| 34 | } |
| 35 | |
| 36 | // Supply the same private stream entrypoint used by authenticated server |
| 37 | // delegations, while retaining actual HTTP serialization between both proxies. |
| 38 | type forwardedTunnelServer struct { |
| 39 | client *Client |
| 40 | } |
| 41 | |
| 42 | func (s *forwardedTunnelServer) DialClient(ctx context.Context, link *protocol.Link) (net.Conn, error) { |
| 43 | local, remote := bufconn.BufferedPipe(8192) |
| 44 | if err := s.client.handleIncomingDelegation(ctx, link, forwardedPeerConn{local}); err != nil { |
| 45 | remote.Close() |
| 46 | return nil, err |
| 47 | } |
| 48 | return remote, nil |
| 49 | } |
| 50 | |
| 51 | func (*forwardedTunnelServer) Identity() *protocol.Node { |
| 52 | return &protocol.Node{Address: "192.0.2.20:443"} |
| 53 | } |
| 54 | |
| 55 | func (*forwardedTunnelServer) DialInternal(context.Context, *protocol.Node) (net.Conn, error) { |
| 56 | return nil, tun.ErrLookupFailed |
| 57 | } |
| 58 | |
| 59 | type forwardedPeerConn struct{ net.Conn } |
| 60 | |
| 61 | func (forwardedPeerConn) RemoteAddr() net.Addr { |
| 62 | return &net.TCPAddr{IP: net.ParseIP("192.0.2.20"), Port: 443} |
| 63 | } |
| 64 | |
| 65 | func TestGatewayClientForwardedContext(t *testing.T) { |
| 66 | der, _, key := testMakeRSACert(require.New(t)) |
| 67 | tlsConfig := &tls.Config{ |
| 68 | Certificates: []tls.Certificate{{Certificate: [][]byte{der}, PrivateKey: key}}, |
| 69 | NextProtos: []string{"h2", "http/1.1", "h3"}, |
| 70 | } |
| 71 | |
| 72 | for _, tc := range []struct { |
| 73 | target, mode, proto string |
| 74 | }{ |
| 75 | {"http", "", "http/1.1"}, |
| 76 | {"http", "target", "h2"}, |
| 77 | {"http", "hostname", "http/1.1"}, |
| 78 | {"pipe", "custom", "h2"}, |
| 79 | {"pipe", "legacy-override", "http/1.1"}, |
| 80 | } { |
| 81 | t.Run(tc.target+"/"+tc.mode+"/"+tc.proto, func(t *testing.T) { |
| 82 | pipeTarget := tc.target == "pipe" |
| 83 | mode, proto := tc.mode, tc.proto |
| 84 | as := require.New(t) |
| 85 | var target string |
| 86 | if pipeTarget { |
| 87 | path, address := randomPipeTarget(t, "forwarded-*") |
| 88 | target = address |
| 89 | listener, err := pipe.ListenPipe(path) |
| 90 | as.NoError(err) |
| 91 | server := &http.Server{Handler: http.HandlerFunc(captureForwardedRequest)} |
| 92 | go server.Serve(listener) |
| 93 | t.Cleanup(func() { server.Close() }) |
| 94 | } else { |
| 95 | server := httptest.NewServer(http.HandlerFunc(captureForwardedRequest)) |
| 96 | target = server.URL |
| 97 | t.Cleanup(server.Close) |
| 98 | } |
| 99 | |
| 100 | // HTTP/1 routes by SNI; HTTP/2 routes by authority for coalescing. |
| 101 | hostname := "tls.example.com" |
| 102 | if proto == "h2" { |
| 103 | hostname = "visitor.example.com" |
| 104 | } |
| 105 | tunnel := Tunnel{Hostname: hostname, Target: target, ProxyHeaderMode: mode} |
| 106 | if mode == "custom" || mode == "legacy-override" { |
| 107 | tunnel.ProxyHeaderHost = "backend.example.com:9000" |
| 108 | } |
| 109 | if mode == "legacy-override" { |
| 110 | tunnel.ProxyHeaderMode = "" |
| 111 | } |
| 112 | cfg := &Config{router: skipmap.NewString[route](), Tunnels: []Tunnel{tunnel}} |
| 113 | as.NoError(cfg.validate()) |
| 114 | cfg.buildRouter() |
| 115 | c := &Client{ |
| 116 | ClientConfig: ClientConfig{Logger: zaptest.NewLogger(t), Configuration: cfg}, |
| 117 | forwarder: &forwarder{ |
| 118 | logger: zaptest.NewLogger(t), |
| 119 | rootDomain: atomic.NewString("example.com"), |
| 120 | proxies: skipmap.NewString[*httpProxy](), |
| 121 | }, |
| 122 | } |
| 123 | t.Cleanup(func() { |
| 124 | c.proxies.Range(func(_ string, proxy *httpProxy) bool { |
| 125 | proxy.acceptor.Close() |
| 126 | proxy.forwarder.Close() |
| 127 | return true |
| 128 | }) |
| 129 | }) |
| 130 | |
| 131 | listener, err := tls.Listen("tcp", "127.0.0.1:0", tlsConfig) |
| 132 | as.NoError(err) |
| 133 | t.Cleanup(func() { listener.Close() }) |
| 134 | quicListener, err := quic.ListenAddr("127.0.0.1:0", tlsConfig, nil) |
| 135 | as.NoError(err) |
| 136 | t.Cleanup(func() { quicListener.Close() }) |
| 137 | gw := gateway.New(gateway.GatewayConfig{ |
| 138 | Logger: c.Logger, |
| 139 | TunnelServer: &forwardedTunnelServer{client: c}, |
| 140 | H2Listener: listener, |
| 141 | H3Listener: quicListener, |
| 142 | GatewayPort: 8443, |
| 143 | Options: gateway.Options{TransportBufferSize: 8192, ProxyBufferSize: 8192}, |
| 144 | }) |
| 145 | gw.MustStart(t.Context()) |
| 146 | t.Cleanup(gw.Close) |
| 147 | |
| 148 | transport := &http.Transport{ |
| 149 | ForceAttemptHTTP2: proto == "h2", |
| 150 | DialTLSContext: func(ctx context.Context, _, _ string) (net.Conn, error) { |
| 151 | dialer := &tls.Dialer{Config: &tls.Config{ |
| 152 | InsecureSkipVerify: true, |
| 153 | ServerName: "tls.example.com", |
| 154 | NextProtos: []string{proto}, |
| 155 | }} |
| 156 | return dialer.DialContext(ctx, "tcp", listener.Addr().String()) |
| 157 | }, |
| 158 | } |
| 159 | t.Cleanup(transport.CloseIdleConnections) |
| 160 | visitor := &http.Client{Transport: transport, Timeout: 5 * time.Second} |
| 161 | req, err := http.NewRequestWithContext(t.Context(), "GET", "https://visitor.example.com:8443/a%2Fb?q=one%2Ftwo", nil) |
| 162 | as.NoError(err) |
| 163 | req.Header = http.Header{ |
| 164 | "Forwarded": {"for=198.51.100.1;proto=http;host=attacker.example"}, |
| 165 | "X-Forwarded-For": {"198.51.100.1, 198.51.100.2", "198.51.100.3"}, |
| 166 | "X-Forwarded-Host": {"attacker.example:1234", "second.example"}, |
| 167 | "X-Forwarded-Proto": {"http", "ftp"}, |
| 168 | "X-Real-Ip": {"198.51.100.4"}, |
| 169 | "True-Client-Ip": {"198.51.100.5"}, |
| 170 | } |
| 171 | if proto == "http/1.1" { |
| 172 | req.Header.Set("Connection", "X-Forwarded-For, x-forwarded-host, X-Forwarded-Proto") |
| 173 | } |
| 174 | resp, err := visitor.Do(req) |
| 175 | as.NoError(err) |
| 176 | defer resp.Body.Close() |
| 177 | as.Equal(http.StatusOK, resp.StatusCode) |
| 178 | as.Equal(proto == "h2", resp.ProtoMajor == 2) |
| 179 | var received forwardedRequest |
| 180 | as.NoError(json.NewDecoder(resp.Body).Decode(&received)) |
| 181 | |
| 182 | expectedHost := cfg.Tunnels[0].parsed.Host |
| 183 | if pipeTarget { |
| 184 | expectedHost = "pipe" |
| 185 | } |
| 186 | if mode == "hostname" { |
| 187 | expectedHost = hostname |
| 188 | } else if mode == "custom" || mode == "legacy-override" { |
| 189 | expectedHost = tunnel.ProxyHeaderHost |
| 190 | } |
| 191 | as.Equal(expectedHost, received.Host) |
| 192 | as.Equal("/a%2Fb?q=one%2Ftwo", received.URI) |
| 193 | as.Equal([]string{"127.0.0.1"}, received.Headers.Values("X-Forwarded-For")) |
| 194 | as.Equal([]string{hostname + ":8443"}, received.Headers.Values("X-Forwarded-Host")) |
| 195 | as.Equal([]string{"https"}, received.Headers.Values("X-Forwarded-Proto")) |
| 196 | for _, header := range []string{"Forwarded", "X-Real-IP", "True-Client-IP", "Connection"} { |
| 197 | as.Empty(received.Headers.Values(header), header) |
| 198 | } |
| 199 | }) |
| 200 | } |
| 201 | } |