Skip to content
File

Blob: tun/client/proxy_test.go

go397 lines
1package client
2 
3import (
4 "context"
5 "fmt"
6 "io"
7 "net"
8 "net/http"
9 "net/http/httptest"
10 "net/url"
11 "os"
12 "path/filepath"
13 "runtime"
14 "testing"
15 
16 "go.miragespace.co/specter/spec/chord"
17 "go.miragespace.co/specter/spec/mocks"
18 "go.miragespace.co/specter/spec/protocol"
19 "go.miragespace.co/specter/spec/rpc"
20 "go.miragespace.co/specter/util/pipe"
21 
22 "github.com/stretchr/testify/mock"
23 "github.com/stretchr/testify/require"
24 "github.com/zhangyunhao116/skipmap"
25 "go.uber.org/atomic"
26 "go.uber.org/zap/zaptest"
27)
28 
29func randomPipeTarget(t *testing.T, pattern string) (string, string) {
30 t.Helper()
31 
32 file, err := os.CreateTemp("", pattern)
33 require.NoError(t, err)
34 
35 name := file.Name()
36 require.NoError(t, file.Close())
37 require.NoError(t, os.Remove(name))
38 
39 if runtime.GOOS == "windows" {
40 path := `\\.\pipe\` + filepath.Base(name)
41 return path, path
42 }
43 
44 t.Cleanup(func() {
45 _ = os.Remove(name)
46 })
47 
48 return name, "unix://" + name
49}
50 
51func TestJustHTTPProxy(t *testing.T) {
52 as := require.New(t)
53 logger := zaptest.NewLogger(t)
54 
55 file, err := os.CreateTemp("", "client")
56 as.NoError(err)
57 defer os.Remove(file.Name())
58 
59 ctx := t.Context()
60 
61 ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
62 fmt.Fprintf(w, "ok")
63 }))
64 defer ts.Close()
65 
66 token := &protocol.ClientToken{
67 Token: []byte("test"),
68 }
69 cl := &protocol.Node{
70 Id: chord.Random(),
71 }
72 
73 der, cert, key := makeCertificate(as, logger, cl, token, nil)
74 cfg := &Config{
75 path: file.Name(),
76 router: skipmap.NewString[route](),
77 Apex: testApex,
78 Certificate: cert,
79 PrivKey: key,
80 Tunnels: []Tunnel{
81 {
82 Target: ts.URL,
83 },
84 },
85 }
86 as.NoError(cfg.validate())
87 
88 m := func(s *mocks.TunnelService, t1 *mocks.MemoryTransport, publishCall *mock.Call) {
89 defaultNoHostnames(s)
90 transportHelper(t1, der)
91 }
92 
93 client, t2, assertion := setupClient(t, as, ctx, logger, nil, cfg, nil, m, false, 1)
94 defer assertion()
95 defer client.Close()
96 
97 client.Start(ctx)
98 
99 httpCl := &http.Client{
100 Transport: &http.Transport{
101 MaxIdleConnsPerHost: -1,
102 DisableKeepAlives: true,
103 DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
104 c1, err := t2.DialStream(ctx, cl, protocol.Stream_DIRECT)
105 as.NoError(err)
106 as.NoError(rpc.Send(c1, &protocol.Link{
107 Alpn: protocol.Link_HTTP,
108 Hostname: testHostname,
109 }))
110 return c1, nil
111 },
112 },
113 }
114 resp, err := httpCl.Get("http://test/")
115 as.NoError(err)
116 defer resp.Body.Close()
117 
118 buf, err := io.ReadAll(resp.Body)
119 as.NoError(err)
120 as.Equal("ok", string(buf))
121}
122 
123func TestPipeHTTP(t *testing.T) {
124 as := require.New(t)
125 logger := zaptest.NewLogger(t)
126 
127 file, err := os.CreateTemp("", "client")
128 as.NoError(err)
129 defer os.Remove(file.Name())
130 
131 ctx := t.Context()
132 
133 path, target := randomPipeTarget(t, "specterhttp-*")
134 
135 pipeListener, err := pipe.ListenPipe(path)
136 as.NoError(err)
137 defer pipeListener.Close()
138 
139 svc := &http.Server{
140 Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
141 fmt.Fprintf(w, "ok")
142 }),
143 }
144 go svc.Serve(pipeListener)
145 defer svc.Shutdown(ctx)
146 
147 token := &protocol.ClientToken{
148 Token: []byte("test"),
149 }
150 cl := &protocol.Node{
151 Id: chord.Random(),
152 }
153 
154 der, cert, key := makeCertificate(as, logger, cl, token, nil)
155 cfg := &Config{
156 path: file.Name(),
157 router: skipmap.NewString[route](),
158 Apex: testApex,
159 Certificate: cert,
160 PrivKey: key,
161 Tunnels: []Tunnel{
162 {
163 Target: target,
164 },
165 },
166 }
167 as.NoError(cfg.validate())
168 
169 m := func(s *mocks.TunnelService, t1 *mocks.MemoryTransport, publishCall *mock.Call) {
170 defaultNoHostnames(s)
171 transportHelper(t1, der)
172 }
173 
174 client, t2, assertion := setupClient(t, as, ctx, logger, nil, cfg, nil, m, false, 1)
175 defer assertion()
176 defer client.Close()
177 
178 client.Start(ctx)
179 
180 httpCl := &http.Client{
181 Transport: &http.Transport{
182 MaxIdleConnsPerHost: -1,
183 DisableKeepAlives: true,
184 DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
185 c1, err := t2.DialStream(ctx, cl, protocol.Stream_DIRECT)
186 as.NoError(err)
187 as.NoError(rpc.Send(c1, &protocol.Link{
188 Alpn: protocol.Link_HTTP,
189 Hostname: testHostname,
190 }))
191 return c1, nil
192 },
193 },
194 }
195 resp, err := httpCl.Get("http://test/")
196 as.NoError(err)
197 defer resp.Body.Close()
198 
199 buf, err := io.ReadAll(resp.Body)
200 as.NoError(err)
201 as.Equal("ok", string(buf))
202}
203 
204func TestPipeTCP(t *testing.T) {
205 as := require.New(t)
206 logger := zaptest.NewLogger(t)
207 
208 file, err := os.CreateTemp("", "client")
209 as.NoError(err)
210 defer os.Remove(file.Name())
211 
212 ctx := t.Context()
213 
214 path, target := randomPipeTarget(t, "spectertcp-*")
215 
216 pipeListener, err := pipe.ListenPipe(path)
217 as.NoError(err)
218 defer pipeListener.Close()
219 
220 go func() {
221 conn, err := pipeListener.Accept()
222 as.NoError(err)
223 conn.Write([]byte("hi"))
224 }()
225 
226 token := &protocol.ClientToken{
227 Token: []byte("test"),
228 }
229 cl := &protocol.Node{
230 Id: chord.Random(),
231 }
232 
233 der, cert, key := makeCertificate(as, logger, cl, token, nil)
234 cfg := &Config{
235 path: file.Name(),
236 router: skipmap.NewString[route](),
237 Apex: testApex,
238 Certificate: cert,
239 PrivKey: key,
240 Tunnels: []Tunnel{
241 {
242 Target: target,
243 },
244 },
245 }
246 as.NoError(cfg.validate())
247 
248 m := func(s *mocks.TunnelService, t1 *mocks.MemoryTransport, publishCall *mock.Call) {
249 defaultNoHostnames(s)
250 transportHelper(t1, der)
251 }
252 
253 client, t2, assertion := setupClient(t, as, ctx, logger, nil, cfg, nil, m, false, 1)
254 defer assertion()
255 defer client.Close()
256 
257 client.Start(ctx)
258 
259 conn, err := t2.DialStream(ctx, cl, protocol.Stream_DIRECT)
260 as.NoError(err)
261 as.NoError(rpc.Send(conn, &protocol.Link{
262 Alpn: protocol.Link_TCP,
263 Hostname: testHostname,
264 }))
265 
266 status := &protocol.TunnelStatus{}
267 rpc.BoundedReceive(conn, status, 1024)
268 as.Equal(protocol.TunnelStatusCode_STATUS_OK, status.GetStatus())
269 
270 buf := make([]byte, 2)
271 n, err := io.ReadFull(conn, buf)
272 as.NoError(err)
273 as.Equal(2, n)
274 as.Equal("hi", string(buf))
275}
276 
277func TestHTTPProxyHostHeaderDefault(t *testing.T) {
278 as := require.New(t)
279 logger := zaptest.NewLogger(t)
280 
281 ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
282 fmt.Fprintf(w, "%s", r.Host)
283 }))
284 defer ts.Close()
285 
286 u, err := url.Parse(ts.URL)
287 as.NoError(err)
288 
289 c := &Client{
290 ClientConfig: ClientConfig{
291 Logger: logger,
292 },
293 forwarder: &forwarder{
294 logger: zaptest.NewLogger(t),
295 rootDomain: atomic.NewString(testApex),
296 proxies: skipmap.NewString[*httpProxy](),
297 },
298 }
299 
300 r := route{
301 parsed: u,
302 }
303 
304 hp := c.getHTTPProxy(context.Background(), testHostname, r)
305 
306 req := httptest.NewRequest("GET", "http://example/", nil)
307 rr := httptest.NewRecorder()
308 hp.forwarder.Handler.ServeHTTP(rr, req)
309 resp := rr.Result()
310 defer resp.Body.Close()
311 body, err := io.ReadAll(resp.Body)
312 as.NoError(err)
313 as.Equal(u.Host, string(body))
314}
315 
316func TestHTTPProxyHostHeaderHostnameMode(t *testing.T) {
317 as := require.New(t)
318 logger := zaptest.NewLogger(t)
319 
320 ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
321 fmt.Fprintf(w, "%s", r.Host)
322 }))
323 defer ts.Close()
324 
325 u, err := url.Parse(ts.URL)
326 as.NoError(err)
327 
328 c := &Client{
329 ClientConfig: ClientConfig{
330 Logger: logger,
331 },
332 forwarder: &forwarder{
333 logger: zaptest.NewLogger(t),
334 rootDomain: atomic.NewString(testApex),
335 proxies: skipmap.NewString[*httpProxy](),
336 },
337 }
338 
339 r := route{
340 parsed: u,
341 proxyHeaderMode: "hostname",
342 }
343 
344 hp := c.getHTTPProxy(context.Background(), testHostname, r)
345 
346 req := httptest.NewRequest("GET", "http://example/", nil)
347 rr := httptest.NewRecorder()
348 hp.forwarder.Handler.ServeHTTP(rr, req)
349 resp := rr.Result()
350 defer resp.Body.Close()
351 body, err := io.ReadAll(resp.Body)
352 as.NoError(err)
353 as.Equal(fmt.Sprintf("%s.%s", testHostname, testApex), string(body))
354}
355 
356func TestHTTPProxyHostHeaderCustomMode(t *testing.T) {
357 as := require.New(t)
358 logger := zaptest.NewLogger(t)
359 
360 ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
361 fmt.Fprintf(w, "%s", r.Host)
362 }))
363 defer ts.Close()
364 
365 u, err := url.Parse(ts.URL)
366 as.NoError(err)
367 
368 c := &Client{
369 ClientConfig: ClientConfig{
370 Logger: logger,
371 },
372 forwarder: &forwarder{
373 logger: zaptest.NewLogger(t),
374 rootDomain: atomic.NewString(testApex),
375 proxies: skipmap.NewString[*httpProxy](),
376 },
377 }
378 
379 customHost := "custom.example.com"
380 r := route{
381 parsed: u,
382 proxyHeaderMode: "custom",
383 proxyHeaderHost: customHost,
384 }
385 
386 hp := c.getHTTPProxy(context.Background(), testHostname, r)
387 
388 req := httptest.NewRequest("GET", "http://example/", nil)
389 rr := httptest.NewRecorder()
390 hp.forwarder.Handler.ServeHTTP(rr, req)
391 resp := rr.Result()
392 defer resp.Body.Close()
393 body, err := io.ReadAll(resp.Body)
394 as.NoError(err)
395 as.Equal(customHost, string(body))
396}