File
Blob: gateway/internal_proxy_test.go
| 1 | package gateway |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "crypto/tls" |
| 6 | "fmt" |
| 7 | "io" |
| 8 | "net" |
| 9 | "net/http" |
| 10 | "os" |
| 11 | "testing" |
| 12 | |
| 13 | "go.miragespace.co/specter/overlay" |
| 14 | "go.miragespace.co/specter/spec/cipher" |
| 15 | mocks "go.miragespace.co/specter/spec/mocks" |
| 16 | "go.miragespace.co/specter/spec/protocol" |
| 17 | "go.miragespace.co/specter/spec/transport" |
| 18 | "go.miragespace.co/specter/spec/tun" |
| 19 | "go.miragespace.co/specter/util/acceptor" |
| 20 | "go.miragespace.co/specter/util/bufconn" |
| 21 | |
| 22 | "github.com/go-chi/chi/v5" |
| 23 | "github.com/quic-go/quic-go" |
| 24 | "github.com/stretchr/testify/mock" |
| 25 | "github.com/stretchr/testify/require" |
| 26 | "go.uber.org/zap" |
| 27 | "go.uber.org/zap/zaptest" |
| 28 | ) |
| 29 | |
| 30 | func setupGatewayWithRouter(t *testing.T, as *require.Assertions, logger *zap.Logger, streamRouter *transport.StreamRouter) (udpPort int, tcpPort int, mockS *mocks.TunnelServer, done func()) { |
| 31 | var q net.PacketConn |
| 32 | var h2 net.Listener |
| 33 | |
| 34 | q, udpPort = getUDPListener(as) |
| 35 | |
| 36 | h2, tcpPort = getH2Listener(as) |
| 37 | |
| 38 | ss := generateTLSConfig([]string{}) |
| 39 | alpnMux, err := overlay.NewMux(&quic.Transport{Conn: q}) |
| 40 | as.NoError(err) |
| 41 | |
| 42 | h3 := alpnMux.With(cipher.GetGatewayTLSConfig(func(chi *tls.ClientHelloInfo) (*tls.Certificate, error) { |
| 43 | return &ss.Certificates[0], nil |
| 44 | }, nil), append(cipher.H3Protos, tun.ALPN(protocol.Link_TCP))...) |
| 45 | |
| 46 | mockS = new(mocks.TunnelServer) |
| 47 | |
| 48 | ctx, cancel := context.WithCancel(context.Background()) |
| 49 | |
| 50 | fakeStats := chi.NewRouter() |
| 51 | fakeStats.Get("/stats", func(w http.ResponseWriter, r *http.Request) { |
| 52 | w.WriteHeader(http.StatusOK) |
| 53 | }) |
| 54 | |
| 55 | conf := GatewayConfig{ |
| 56 | Logger: logger, |
| 57 | TunnelServer: mockS, |
| 58 | HTTPListener: nil, |
| 59 | H2Listener: h2, |
| 60 | H3Listener: h3, |
| 61 | RootDomains: []string{testDomain}, |
| 62 | GatewayPort: udpPort, |
| 63 | AdminUser: os.Getenv("INTERNAL_USER"), |
| 64 | AdminPass: os.Getenv("INTERNAL_PASS"), |
| 65 | Handlers: InternalHandlers{ |
| 66 | Chord: fakeStats, |
| 67 | }, |
| 68 | Options: Options{ |
| 69 | TransportBufferSize: 1024 * 8, |
| 70 | ProxyBufferSize: 1024 * 8, |
| 71 | }, |
| 72 | } |
| 73 | g := New(conf) |
| 74 | g.AttachRouter(ctx, streamRouter) |
| 75 | g.MustStart(ctx) |
| 76 | |
| 77 | go alpnMux.Accept(ctx) |
| 78 | |
| 79 | return udpPort, tcpPort, mockS, func() { |
| 80 | cancel() |
| 81 | h2.Close() |
| 82 | h3.Close() |
| 83 | alpnMux.Close() |
| 84 | g.Close() |
| 85 | } |
| 86 | } |
| 87 | |
| 88 | func TestInternalProxy(t *testing.T) { |
| 89 | os.Setenv("INTERNAL_USER", testUser) |
| 90 | os.Setenv("INTERNAL_PASS", testPass) |
| 91 | defer func() { |
| 92 | os.Setenv("INTERNAL_USER", "") |
| 93 | os.Setenv("INTERNAL_PASS", "") |
| 94 | }() |
| 95 | |
| 96 | as := require.New(t) |
| 97 | logger := zaptest.NewLogger(t, zaptest.WrapOptions(zap.AddCaller())) |
| 98 | tp := mocks.SelfTransport() |
| 99 | streamRouter := transport.NewStreamRouter(logger, tp, nil) |
| 100 | |
| 101 | _, tcpPort, mockS, done := setupGatewayWithRouter(t, as, logger, streamRouter) |
| 102 | defer done() |
| 103 | |
| 104 | ctx := t.Context() |
| 105 | |
| 106 | go streamRouter.Accept(ctx) |
| 107 | |
| 108 | fakeNode := &protocol.Node{ |
| 109 | Address: "127.0.0.1:4444", |
| 110 | } |
| 111 | |
| 112 | respString := "proxied" |
| 113 | fakeServer := &http.Server{ |
| 114 | Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { |
| 115 | fmt.Fprint(w, respString) |
| 116 | }), |
| 117 | } |
| 118 | acc := acceptor.NewH2Acceptor(nil) |
| 119 | go fakeServer.Serve(acc) |
| 120 | |
| 121 | c1, c2 := bufconn.BufferedPipe(8192) |
| 122 | |
| 123 | mockS.On("DialInternal", mock.Anything, mock.MatchedBy(func(node *protocol.Node) bool { |
| 124 | return node.GetAddress() == fakeNode.GetAddress() && node.GetId() == fakeNode.GetId() |
| 125 | })).Return(c1, nil).Run(func(args mock.Arguments) { |
| 126 | acc.Handle(c2) |
| 127 | }) |
| 128 | |
| 129 | req, err := http.NewRequest("GET", fmt.Sprintf("https://%s/_internal/chord/stats", testDomain), nil) |
| 130 | as.NoError(err) |
| 131 | |
| 132 | req.SetBasicAuth(testUser, testPass) |
| 133 | req.Header.Set(internalProxyNodeAddress, fakeNode.GetAddress()) |
| 134 | |
| 135 | c := getH2Client("", tcpPort) |
| 136 | |
| 137 | resp, err := c.Do(req) |
| 138 | as.NoError(err) |
| 139 | defer resp.Body.Close() |
| 140 | |
| 141 | b, err := io.ReadAll(resp.Body) |
| 142 | as.NoError(err) |
| 143 | |
| 144 | as.Contains(string(b), respString) |
| 145 | as.NotEmpty(resp.Header.Get("alt-svc")) |
| 146 | as.Equal("false", resp.Header.Get("http3")) |
| 147 | |
| 148 | mockS.AssertExpectations(t) |
| 149 | } |