Skip to content
File

Blob: gateway/http_test.go

go114 lines
1package gateway
2 
3import (
4 "bufio"
5 "crypto/rand"
6 "fmt"
7 "io"
8 "net"
9 "net/http"
10 "net/url"
11 "testing"
12 "time"
13 
14 "go.miragespace.co/specter/spec/protocol"
15 "go.miragespace.co/specter/spec/tun"
16 "go.miragespace.co/specter/util/bufconn"
17 
18 "github.com/stretchr/testify/mock"
19 "github.com/stretchr/testify/require"
20)
21 
22func TestHTTPRedirect(t *testing.T) {
23 as := require.New(t)
24 
25 testHost := "http"
26 testPath := "/sup"
27 
28 httpListener, httpPort := getTCPListener(as)
29 port, _, _, done := setupGateway(t, as, httpListener)
30 defer done()
31 
32 c := getHTTPClient(httpPort)
33 c.CheckRedirect = func(req *http.Request, via []*http.Request) error {
34 return http.ErrUseLastResponse
35 }
36 
37 req, err := http.NewRequest("GET", fmt.Sprintf("http://%s.%s%s", testHost, testDomain, testPath), nil)
38 as.NoError(err)
39 
40 resp, err := c.Do(req)
41 as.NoError(err)
42 defer resp.Body.Close()
43 
44 re := resp.Header.Get("location")
45 as.Equal(fmt.Sprintf("https://%s.%s:%d%s", testHost, testDomain, port, testPath), re)
46 as.Equal(http.StatusMovedPermanently, resp.StatusCode)
47}
48 
49func TestHTTPConnectProxy(t *testing.T) {
50 as := require.New(t)
51 
52 httpListener, httpPort := getTCPListener(as)
53 _, _, mockS, done := setupGateway(t, as, httpListener)
54 defer done()
55 
56 testHost := "hello"
57 bufLength := 16
58 
59 c1, c2 := bufconn.BufferedPipe(8192)
60 
61 go func() {
62 tun.SendStatusProto(c2, nil)
63 
64 buf := make([]byte, bufLength)
65 n, err := io.ReadFull(c2, buf)
66 as.NoError(err)
67 as.Equal(bufLength, n)
68 
69 rand.Read(buf)
70 n, err = c2.Write(buf)
71 as.NoError(err)
72 as.Equal(bufLength, n)
73 }()
74 
75 mockS.On("DialClient", mock.Anything, mock.MatchedBy(func(l *protocol.Link) bool {
76 return l.GetAlpn() == protocol.Link_TCP && l.GetHostname() == testHost
77 })).Return(c1, nil)
78 
79 // HTTP Connect start
80 dialer := &net.Dialer{
81 Timeout: time.Second,
82 }
83 conn, err := dialer.Dial("tcp", fmt.Sprintf("127.0.0.1:%d", httpPort))
84 as.NoError(err)
85 
86 proxyAddr := fmt.Sprintf("%s.%s:1234", testHost, testDomain) // port doesn't matter
87 req := &http.Request{
88 Method: http.MethodConnect,
89 URL: &url.URL{
90 Opaque: proxyAddr,
91 },
92 Host: proxyAddr,
93 Header: make(http.Header),
94 }
95 as.NoError(req.Write(conn))
96 resp, err := http.ReadResponse(bufio.NewReader(conn), req)
97 as.NoError(err)
98 as.Equal(http.StatusOK, resp.StatusCode)
99 // HTTP Connect end
100 
101 buf := make([]byte, bufLength)
102 
103 rand.Read(buf)
104 n, err := conn.Write(buf)
105 as.NoError(err)
106 as.Equal(bufLength, n)
107 
108 n, err = io.ReadFull(conn, buf)
109 as.NoError(err)
110 as.Equal(bufLength, n)
111 
112 mockS.AssertExpectations(t)
113}