Skip to content
File

Blob: tun/server/session_rpc_test.go

go269 lines
1package server
2 
3import (
4 "context"
5 "crypto/x509"
6 "errors"
7 "net"
8 "net/http"
9 "testing"
10 "time"
11 
12 "go.miragespace.co/specter/spec/mocks"
13 "go.miragespace.co/specter/spec/protocol"
14 "go.miragespace.co/specter/spec/rpc"
15 "go.miragespace.co/specter/spec/transport"
16 "go.miragespace.co/specter/spec/tun"
17 
18 "github.com/stretchr/testify/mock"
19 "github.com/stretchr/testify/require"
20 "github.com/twitchtv/twirp"
21)
22 
23func sessionFixture(t *testing.T) (*Server, *mocks.VNode, *mocks.Transport, *x509.Certificate) {
24 logger, n, clientT, chordT, s := getFixture(t, require.New(t))
25 s.Chord, s.ParentContext = n, t.Context()
26 cli, ch, tn := getIdentities()
27 n.On("ID").Return(ch.Id).Maybe()
28 clientT.On("Identity").Return(tn)
29 chordT.On("Identity").Return(ch)
30 cert := toCertificate(require.New(t), logger, cli, &protocol.ClientToken{Token: mustGenerateToken()})
31 t.Cleanup(func() { s.rpcAcceptor.Close(); s.routeCache.Close(); s.keylessCache.Close() })
32 return s, n, clientT, cert
33}
34 
35func sessionContext(t *testing.T, cert *x509.Certificate) (context.Context, *mocks.PhysicalConn) {
36 pc := mocks.NewPhysicalConn(func(d *transport.StreamDelegate) { d.Close() })
37 t.Cleanup(func() { pc.Close("test cleanup") })
38 conn, err := pc.OpenStream(protocol.Stream_RPC)
39 require.NoError(t, err)
40 return rpc.WithDelegation(t.Context(), &transport.StreamDelegate{
41 Conn: conn,
42 Certificate: cert,
43 }), pc
44}
45 
46func TestOpenEphemeralSession(t *testing.T) {
47 s, n, clientT, cert := sessionFixture(t)
48 tp := mocks.SelfTransport()
49 tp.WithCertificate(cert)
50 delivered := make(chan *transport.StreamDelegate, 1)
51 pc := mocks.NewPhysicalConn(func(d *transport.StreamDelegate) { delivered <- d })
52 t.Cleanup(func() { pc.Close("cleanup") })
53 tp.Physical = pc
54 router := transport.NewStreamRouter(s.Logger, nil, tp)
55 s.AttachRouter(t.Context(), router)
56 go router.Accept(t.Context())
57 ctx := t.Context()
58 httpClient := &http.Client{Transport: &http.Transport{
59 DisableKeepAlives: true,
60 DialContext: func(context.Context, string, string) (net.Conn, error) {
61 return tp.DialStream(ctx, &protocol.Node{Address: "test"}, protocol.Stream_RPC)
62 },
63 }}
64 cli := protocol.NewTunnelServiceProtobufClient("http://tunnel", httpClient)
65 // Open methods authenticate their attachment without owner registration.
66 tp.WithCertificate(nil)
67 _, err := cli.OpenEphemeralSession(ctx, &protocol.OpenEphemeralSessionRequest{})
68 require.Equal(t, twirp.Unauthenticated, err.(twirp.Error).Code())
69 _, err = cli.OpenDelegatedSession(ctx, &protocol.OpenDelegatedSessionRequest{})
70 require.Equal(t, twirp.Unauthenticated, err.(twirp.Error).Code())
71 tp.WithCertificate(cert)
72 resp, err := cli.OpenEphemeralSession(ctx, &protocol.OpenEphemeralSessionRequest{})
73 require.NoError(t, err)
74 expected, _, err := tun.EphemeralLabel(s.Chord.ID(), cert.RawSubjectPublicKeyInfo)
75 require.NoError(t, err)
76 require.Equal(t, expected, resp.Hostname)
77 require.Equal(t, s.TunnelTransport.Identity(), resp.Node)
78 require.Equal(t, s.Apex, resp.Apex)
79 conn, err := s.DialClient(t.Context(), &protocol.Link{
80 Hostname: expected,
81 Alpn: protocol.Link_HTTP,
82 })
83 require.NoError(t, err)
84 conn.Close()
85 (<-delivered).Close()
86 _, err = cli.OpenEphemeralSession(ctx, &protocol.OpenEphemeralSessionRequest{})
87 require.Equal(t, twirp.FailedPrecondition, err.(twirp.Error).Code())
88 require.NoError(t, pc.Err())
89 pc.Close("disconnected")
90 _, err = s.DialClient(t.Context(), &protocol.Link{Hostname: expected})
91 require.ErrorIs(t, err, tun.ErrTunnelClientNotConnected)
92 clientT.AssertNotCalled(t, "DialStream", mock.Anything, mock.Anything, mock.Anything)
93 nextCtx, _ := sessionContext(t, cert)
94 reopened, err := s.OpenEphemeralSession(nextCtx, &protocol.OpenEphemeralSessionRequest{})
95 require.NoError(t, err)
96 require.Equal(t, expected, reopened.Hostname)
97 _, _, _, another := sessionFixture(t)
98 nextCtx, _ = sessionContext(t, another)
99 fresh, err := s.OpenEphemeralSession(nextCtx, &protocol.OpenEphemeralSessionRequest{})
100 require.NoError(t, err)
101 require.NotEqual(t, expected, fresh.Hostname)
102 for _, method := range []string{"Put", "PrefixAppend", "Delete", "Acquire"} {
103 n.AssertNotCalled(t, method, mock.Anything, mock.Anything, mock.Anything)
104 }
105}
106 
107func TestOpenDelegatedSession(t *testing.T) {
108 for _, scenario := range []string{"slot 1", "slot 2", "slot 3", "custom valid", "malformed", "unknown", "version", "expired", "not owned", "custom mismatch", "publication failed"} {
109 t.Run(scenario, func(t *testing.T) {
110 s, n, _, cert := sessionFixture(t)
111 ctx, pc := sessionContext(t, cert)
112 secret := [32]byte{7}
113 token := tun.EncodeDelegationToken(secret)
114 rec := &protocol.DelegationRecord{
115 Version: 1,
116 Id: tun.DelegationID(secret),
117 Owner: &protocol.ClientToken{Token: []byte("owner")},
118 Hostname: "test-host",
119 }
120 code := twirp.NoError
121 slot := uint32(1)
122 switch scenario {
123 case "slot 2":
124 slot = 2
125 rec.ExpiresAt = time.Now().Add(time.Hour).Unix()
126 case "slot 3":
127 slot = 3
128 case "malformed":
129 token = "bad"
130 code = twirp.InvalidArgument
131 case "unknown":
132 code = twirp.Unauthenticated
133 case "version":
134 rec.Version = 2
135 code = twirp.Unauthenticated
136 case "expired":
137 rec.ExpiresAt = time.Now().Add(-time.Hour).Unix()
138 code = twirp.PermissionDenied
139 case "not owned", "custom mismatch":
140 code = twirp.PermissionDenied
141 case "publication failed":
142 code = twirp.Unavailable
143 }
144 if scenario == "custom mismatch" || scenario == "custom valid" {
145 rec.Hostname = "test.example.com"
146 }
147 if scenario != "malformed" {
148 data, err := rec.MarshalVT()
149 require.NoError(t, err)
150 if scenario == "unknown" {
151 data = nil
152 }
153 n.On("Get", mock.Anything, []byte(tun.DelegationKey(rec.Id))).Return(data, nil).Once()
154 }
155 if code == twirp.NoError || scenario == "publication failed" || scenario == "not owned" || scenario == "custom mismatch" {
156 acquired := n.On("Acquire", mock.Anything, []byte(tun.ClientLeaseKey(rec.Owner)), 30*time.Second).Return(uint64(1), nil).Once()
157 n.On("Release", mock.Anything, []byte(tun.ClientLeaseKey(rec.Owner)), uint64(1)).Return(nil).Once().NotBefore(acquired)
158 n.On("PrefixContains", mock.Anything, []byte(tun.ClientHostnamesPrefix(rec.Owner)), []byte(rec.Hostname)).Return(scenario != "not owned", nil).Once()
159 if scenario == "custom mismatch" || scenario == "custom valid" {
160 owner := rec.Owner
161 if scenario == "custom mismatch" {
162 owner = &protocol.ClientToken{Token: []byte("someone else")}
163 }
164 data, _ := (&protocol.CustomHostname{ClientToken: owner}).MarshalVT()
165 n.On("Get", mock.Anything, []byte(tun.CustomHostnameKey(rec.Hostname))).Return(data, nil).Once()
166 }
167 }
168 var written []*protocol.TunnelRoute
169 if code == twirp.NoError || scenario == "publication failed" {
170 var putErr error
171 if scenario == "publication failed" {
172 putErr = errors.New("storage unavailable")
173 }
174 n.On("Put", mock.Anything, []byte(tun.RoutingKey(rec.Hostname, int(slot))), mock.Anything).Run(func(args mock.Arguments) {
175 route := &protocol.TunnelRoute{}
176 require.NoError(t, route.UnmarshalVT(args.Get(2).([]byte)))
177 written = append(written, route)
178 }).Return(putErr).Once()
179 }
180 resp, err := s.OpenDelegatedSession(ctx, &protocol.OpenDelegatedSessionRequest{
181 Token: token,
182 RouteSlot: slot,
183 })
184 if code == twirp.NoError {
185 require.NoError(t, err)
186 require.Equal(t, rec.Id, resp.GrantId)
187 require.Equal(t, rec.ExpiresAt, resp.ExpiresAt)
188 require.Equal(t, rec.Hostname, resp.Hostname)
189 require.Equal(t, s.TunnelTransport.Identity(), resp.Node)
190 require.Equal(t, s.Apex, resp.Apex)
191 require.Len(t, written, 1)
192 require.True(t, tun.IsSessionAlias(written[0].ClientDestination.Address))
193 require.True(t, written[0].ClientDestination.Rendezvous)
194 require.Equal(t, s.ChordTransport.Identity(), written[0].ChordDestination)
195 require.Equal(t, s.TunnelTransport.Identity(), written[0].TunnelDestination)
196 conn, err := s.sessions.dial(written[0].ClientDestination.Address, rec.Hostname)
197 require.NoError(t, err)
198 require.Equal(t, pc, conn.(transport.PhysicalConnProvider).PhysicalConn())
199 conn.Close()
200 } else {
201 require.Equal(t, code, err.(twirp.Error).Code())
202 if len(written) != 0 {
203 _, err := s.sessions.dial(written[0].ClientDestination.Address, rec.Hostname)
204 require.ErrorIs(t, err, transport.ErrNoDirect)
205 }
206 require.Eventually(t, func() bool { return connectionDone(pc) }, 2*time.Second, 10*time.Millisecond)
207 }
208 n.AssertNotCalled(t, "Delete", mock.Anything, mock.Anything)
209 n.AssertExpectations(t)
210 })
211 }
212}
213 
214func TestOpenDelegatedSessionRejectsInvalidSlot(t *testing.T) {
215 s, n, _, cert := sessionFixture(t)
216 ctx, pc := sessionContext(t, cert)
217 for _, slot := range []uint32{0, tun.NumRedundantLinks + 1} {
218 _, err := s.OpenDelegatedSession(ctx, &protocol.OpenDelegatedSessionRequest{RouteSlot: slot})
219 require.Equal(t, twirp.InvalidArgument, err.(twirp.Error).Code())
220 }
221 s.sessions.mu.Lock()
222 reserved := len(s.sessions.byConn)
223 s.sessions.mu.Unlock()
224 require.Zero(t, reserved)
225 require.NoError(t, pc.Err())
226 require.Empty(t, n.Calls)
227}
228 
229func TestSessionCleanupPreservesReboundAlias(t *testing.T) {
230 registry := newSessionRegistry()
231 physical := func() *mocks.PhysicalConn {
232 conn := mocks.NewPhysicalConn(func(d *transport.StreamDelegate) { d.Close() })
233 t.Cleanup(func() { conn.Close("test cleanup") })
234 return conn
235 }
236 old := &session{
237 alias: tun.SessionAlias([16]byte{1}),
238 hostname: "test",
239 conn: physical(),
240 mode: ephemeral,
241 spki: []byte("same key"),
242 }
243 require.NoError(t, registry.reserve(old))
244 // Model failed setup awaiting physical closure while its error response flushes.
245 old.mu.Lock()
246 old.state = terminal
247 old.mu.Unlock()
248 replacement := &session{
249 alias: old.alias,
250 hostname: old.hostname,
251 conn: physical(),
252 mode: ephemeral,
253 spki: old.spki,
254 }
255 require.NoError(t, registry.reserve(replacement))
256 require.True(t, replacement.activate(t.Context(), time.Time{}))
257 old.conn.Close("old attachment ended")
258 require.Eventually(t, func() bool {
259 registry.mu.Lock()
260 defer registry.mu.Unlock()
261 return registry.byConn[old.conn] == nil
262 }, time.Second, time.Millisecond)
263 
264 conn, err := registry.dial(replacement.alias, replacement.hostname)
265 require.NoError(t, err)
266 defer conn.Close()
267 require.Equal(t, replacement.conn, conn.(transport.PhysicalConnProvider).PhysicalConn())
268}