Skip to content
File

Blob: gateway/internal_proxy_test.go

go150 lines
1package gateway
2 
3import (
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 
30func 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 
88func 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}