Skip to content
File

Blob: tun/server/server_test.go

go341 lines
1package server
2 
3import (
4 "bytes"
5 "context"
6 "fmt"
7 "testing"
8 "time"
9 
10 "go.miragespace.co/specter/spec/chord"
11 mocks "go.miragespace.co/specter/spec/mocks"
12 "go.miragespace.co/specter/spec/protocol"
13 "go.miragespace.co/specter/spec/rpc"
14 "go.miragespace.co/specter/spec/transport"
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 "go.uber.org/zap"
21 "go.uber.org/zap/zaptest"
22)
23 
24const (
25 testRootDomain = "hello.com"
26 testAcmeZone = "acme.example.com"
27)
28 
29func assertBytes(got []byte, exp ...[]byte) bool {
30 for _, b := range exp {
31 r := bytes.Compare(got, b)
32 if r == 0 {
33 return true
34 }
35 }
36 return false
37}
38 
39type configOption func(*Config)
40 
41func withResolver(r tun.DNSResolver) configOption {
42 return func(c *Config) {
43 c.Resolver = r
44 }
45}
46 
47func getFixture(t *testing.T, as *require.Assertions, options ...configOption) (*zap.Logger, *mocks.VNode, *mocks.Transport, *mocks.Transport, *Server) {
48 logger := zaptest.NewLogger(t, zaptest.WrapOptions(zap.AddCaller()))
49 
50 n := new(mocks.VNode)
51 cht := new(mocks.Transport)
52 clt := new(mocks.Transport)
53 
54 as.True(true) // placeholder
55 
56 cfg := Config{
57 ParentContext: context.Background(),
58 Logger: logger,
59 Chord: chord.WrapRetryKV(n, time.Millisecond*100, 3),
60 TunnelTransport: clt,
61 ChordTransport: cht,
62 Apex: testRootDomain,
63 Acme: testAcmeZone,
64 }
65 
66 for _, opt := range options {
67 opt(&cfg)
68 }
69 
70 s := New(cfg)
71 
72 return logger, n, clt, cht, s
73}
74 
75func getExpected(link *protocol.Link) [][]byte {
76 return [][]byte{
77 []byte(tun.RoutingKey(link.GetHostname(), 1)),
78 []byte(tun.RoutingKey(link.GetHostname(), 2)),
79 []byte(tun.RoutingKey(link.GetHostname(), 3)),
80 }
81}
82 
83func getIdentities() (*protocol.Node, *protocol.Node, *protocol.Node) {
84 cl := &protocol.Node{
85 Address: "client:123",
86 Id: chord.Random(),
87 Rendezvous: true,
88 }
89 ch := &protocol.Node{
90 Address: "chord:123",
91 Id: chord.Random(),
92 }
93 tn := &protocol.Node{
94 Address: "tunnel:123",
95 Id: chord.Random(),
96 }
97 return cl, ch, tn
98}
99 
100func TestContinueLookupOnError(t *testing.T) {
101 as := require.New(t)
102 
103 _, node, _, _, serv := getFixture(t, as)
104 
105 link := &protocol.Link{
106 Alpn: protocol.Link_HTTP,
107 Hostname: "test",
108 }
109 expected := getExpected(link)
110 
111 node.On("Get", mock.Anything, mock.MatchedBy(func(k []byte) bool {
112 return assertBytes(k, expected...)
113 })).Return(nil, fmt.Errorf("panic"))
114 
115 c, err := serv.DialClient(context.Background(), link)
116 as.ErrorIs(err, tun.ErrLookupFailed)
117 as.Nil(c)
118 
119 node.AssertExpectations(t)
120}
121 
122func TestLookupSuccessDirect(t *testing.T) {
123 as := require.New(t)
124 
125 _, node, clientT, _, serv := getFixture(t, as)
126 
127 link := &protocol.Link{
128 Alpn: protocol.Link_HTTP,
129 Hostname: "test",
130 }
131 
132 cli, cht, tn := getIdentities()
133 bundle := &protocol.TunnelRoute{
134 ClientDestination: cli,
135 ChordDestination: cht,
136 TunnelDestination: tn,
137 Hostname: link.GetHostname(),
138 }
139 bundleBuf, err := bundle.MarshalVT()
140 as.NoError(err)
141 
142 expected := getExpected(link)
143 
144 // 1. first query the chord network
145 node.On("Get", mock.Anything, mock.MatchedBy(func(k []byte) bool {
146 return assertBytes(k, expected...)
147 })).Return(bundleBuf, nil)
148 
149 // 2. then it should compare the content in bundle
150 clientT.On("Identity").Return(tn)
151 
152 // 3. once we figure out that it is connected to us,
153 // attempt to dial
154 c1, c2 := bufconn.BufferedPipe(8192)
155 go func() {
156 l := &protocol.Link{}
157 err := rpc.BoundedReceive(c2, l, 1024)
158 as.NoError(err)
159 as.Equal(link.GetAlpn(), l.GetAlpn())
160 as.Equal(link.GetHostname(), l.GetHostname())
161 }()
162 clientT.On("DialStream", mock.Anything, mock.MatchedBy(func(n *protocol.Node) bool {
163 return n.GetId() == cli.GetId()
164 }), protocol.Stream_DIRECT).Return(c1, nil)
165 
166 _, err = serv.DialClient(context.Background(), link)
167 as.NoError(err)
168 
169 node.AssertExpectations(t)
170 clientT.AssertExpectations(t)
171}
172 
173func TestLookupSuccessRemote(t *testing.T) {
174 as := require.New(t)
175 
176 _, node, clientT, chordT, serv := getFixture(t, as)
177 
178 link := &protocol.Link{
179 Alpn: protocol.Link_HTTP,
180 Hostname: "test",
181 }
182 
183 cli, cht, tn := getIdentities()
184 bundle := &protocol.TunnelRoute{
185 ClientDestination: cli,
186 ChordDestination: cht,
187 TunnelDestination: tn,
188 Hostname: link.GetHostname(),
189 }
190 bundleBuf, err := bundle.MarshalVT()
191 as.NoError(err)
192 
193 expected := getExpected(link)
194 
195 // 1. first query the chord network
196 node.On("Get", mock.Anything, mock.MatchedBy(func(k []byte) bool {
197 return assertBytes(k, expected...)
198 })).Return(bundleBuf, nil)
199 
200 // 2. then it should compare the content in bundle
201 clientT.On("Identity").Return(nil)
202 
203 // 3. once we figure out that it is NOT connected to us,
204 // attempt to dial via chord
205 c1, c2 := bufconn.BufferedPipe(8192)
206 go func() {
207 // the remote node should receive the bundle
208 bundle := &protocol.TunnelRoute{}
209 err := rpc.BoundedReceive(c2, bundle, 2048)
210 as.NoError(err)
211 
212 // remote node need to send feedback
213 tun.SendStatusProto(c2, nil)
214 
215 // then receive the link information
216 l := &protocol.Link{}
217 err = rpc.BoundedReceive(c2, l, 1024)
218 as.NoError(err)
219 as.Equal(link.GetAlpn(), l.GetAlpn())
220 as.Equal(link.GetHostname(), l.GetHostname())
221 }()
222 chordT.On("DialStream", mock.Anything, mock.MatchedBy(func(n *protocol.Node) bool {
223 return n.GetId() == cht.GetId()
224 }), protocol.Stream_PROXY).Return(c1, nil)
225 
226 _, err = serv.DialClient(context.Background(), link)
227 as.NoError(err)
228 
229 node.AssertExpectations(t)
230 clientT.AssertExpectations(t)
231 chordT.AssertExpectations(t)
232}
233 
234func TestHandleRemoteConnection(t *testing.T) {
235 as := require.New(t)
236 
237 logger, node, clientT, chordT, serv := getFixture(t, as)
238 cli, cht, tn := getIdentities()
239 bundle := &protocol.TunnelRoute{
240 ClientDestination: cli,
241 ChordDestination: cht,
242 TunnelDestination: tn,
243 Hostname: "test",
244 }
245 
246 ctx := t.Context()
247 
248 syncA := make(chan struct{})
249 syncB := make(chan struct{})
250 
251 chordChan := make(chan *transport.StreamDelegate)
252 chordT.On("AcceptStream").Return(chordChan)
253 chordT.On("Identity").Return(cht)
254 
255 clientChan := make(chan *transport.StreamDelegate)
256 clientT.On("AcceptStream").Return(clientChan)
257 clientT.On("Identity").Return(tn)
258 
259 // on start up (Accept), identities should get published
260 node.On("Put", mock.Anything, mock.MatchedBy(func(k []byte) bool {
261 exp := [][]byte{
262 []byte(tun.DestinationByChordKey(cht)),
263 []byte(tun.DestinationByTunnelKey(tn)),
264 }
265 return assertBytes(k, exp...)
266 }), mock.MatchedBy(func(v []byte) bool {
267 pair := &protocol.TunnelDestination{
268 Chord: cht,
269 Tunnel: tn,
270 }
271 buf, err := pair.MarshalVT()
272 if err != nil {
273 return false
274 }
275 return assertBytes(v, buf)
276 })).Return(nil)
277 
278 streamRouter := transport.NewStreamRouter(logger, chordT, clientT)
279 go streamRouter.Accept(ctx)
280 
281 serv.AttachRouter(ctx, streamRouter)
282 
283 // since the "client" is connected to us, we should expect a DialDirect
284 // to the client
285 c1, c2 := bufconn.BufferedPipe(8192)
286 c3, c4 := bufconn.BufferedPipe(8192)
287 clientT.On("DialStream", mock.Anything, mock.MatchedBy(func(n *protocol.Node) bool {
288 return n.GetId() == cli.GetId()
289 }), protocol.Stream_DIRECT).Return(c3, nil)
290 
291 buf := []byte{1, 2, 3}
292 
293 go func() {
294 // the remote should be sending the bundle over
295 err := rpc.Send(c2, bundle)
296 as.NoError(err)
297 
298 // getConn should check the status
299 x := &protocol.TunnelStatus{}
300 err = rpc.BoundedReceive(c2, x, 1024)
301 as.NoError(err)
302 as.Equal(protocol.TunnelStatusCode_STATUS_OK, x.GetStatus())
303 
304 // now the remote gateway is sending data to us
305 _, err = c2.Write(buf)
306 as.NoError(err)
307 
308 close(syncA)
309 }()
310 
311 go func() {
312 <-syncA
313 
314 // we should receive data from the remote side
315 b := make([]byte, len(buf))
316 _, err := c4.Read(b)
317 as.NoError(err)
318 as.EqualValues(buf, b)
319 
320 close(syncB)
321 }()
322 
323 serv.MustRegister(ctx)
324 
325 chordChan <- &transport.StreamDelegate{
326 Conn: c1,
327 Identity: &protocol.Node{Id: chord.Random()},
328 Kind: protocol.Stream_PROXY,
329 }
330 
331 select {
332 case <-syncB:
333 case <-time.After(time.Second * 5):
334 as.FailNow("timeout")
335 }
336 
337 clientT.AssertExpectations(t)
338 chordT.AssertExpectations(t)
339 node.AssertExpectations(t)
340}