Skip to content
File

Blob: tun/server/revalidate_test.go

go268 lines
1package server
2 
3import (
4 "context"
5 "errors"
6 "testing"
7 "testing/synctest"
8 "time"
9 
10 "go.miragespace.co/specter/spec/mocks"
11 "go.miragespace.co/specter/spec/protocol"
12 "go.miragespace.co/specter/spec/transport"
13 "go.miragespace.co/specter/spec/tun"
14 
15 "github.com/stretchr/testify/mock"
16 "github.com/stretchr/testify/require"
17)
18 
19func revalidationFixture(t *testing.T, count int) (*Server, *mocks.VNode, []*session) {
20 t.Helper()
21 n := new(mocks.VNode)
22 s := &Server{
23 Config: Config{
24 ParentContext: t.Context(),
25 Chord: n,
26 },
27 sessions: newSessionRegistry(),
28 }
29 sessions := make([]*session, 0, count)
30 for i := 0; i < count; i++ {
31 pc := mocks.NewPhysicalConn(func(*transport.StreamDelegate) {})
32 t.Cleanup(func() { pc.Close("test cleanup") })
33 sess := &session{
34 alias: tun.SessionAlias([16]byte{byte(i + 1)}),
35 hostname: "test",
36 conn: pc,
37 mode: delegated,
38 grantID: tun.DelegationID([32]byte{2}),
39 owner: &protocol.ClientToken{Token: []byte("owner")},
40 certNotAfter: time.Now().Add(time.Hour),
41 }
42 require.NoError(t, s.sessions.reserve(sess))
43 sessions = append(sessions, sess)
44 }
45 return s, n, sessions
46}
47 
48func revalidationRecord(t *testing.T, sess *session) []byte {
49 t.Helper()
50 rec := &protocol.DelegationRecord{
51 Version: 1,
52 Id: sess.grantID,
53 Hostname: sess.hostname,
54 Owner: sess.owner,
55 }
56 if !sess.grantExpiry.IsZero() {
57 rec.ExpiresAt = sess.grantExpiry.Unix()
58 }
59 data, err := rec.MarshalVT()
60 require.NoError(t, err)
61 return data
62}
63 
64func TestRevalidationIndependentSessions(t *testing.T) {
65 synctest.Test(t, func(t *testing.T) {
66 s, n, sessions := revalidationFixture(t, 3)
67 // All three connections share a token; each must maintain its own authority.
68 n.On("Get", mock.Anything, []byte(tun.DelegationKey(sessions[0].grantID))).
69 Run(func(mock.Arguments) { time.Sleep(2 * time.Second) }).
70 Return(revalidationRecord(t, sessions[0]), nil)
71 n.On("PrefixContains", mock.Anything, mock.Anything, mock.Anything).Return(true, nil)
72 until := time.Now().Add(revalidateGrace)
73 for _, sess := range sessions {
74 require.True(t, sess.activate(t.Context(), until))
75 go s.maintainSession(s.ParentContext, sess)
76 }
77 time.Sleep(revalidateInterval + lookupTimeout)
78 synctest.Wait()
79 for i, sess := range sessions {
80 sess.mu.Lock()
81 extended := sess.authorizedUntil.After(until)
82 sess.mu.Unlock()
83 require.True(t, extended, "session %d was starved", i)
84 }
85 sessions[0].conn.Close("one connection ended")
86 // Successful checks must keep resetting expiry without affecting siblings.
87 time.Sleep(2 * revalidateGrace)
88 synctest.Wait()
89 for _, sess := range sessions[1:] {
90 require.False(t, connectionDone(sess.conn))
91 sess.mu.Lock()
92 authorized := time.Now().Before(sess.authorizedUntil)
93 sess.mu.Unlock()
94 require.True(t, authorized)
95 }
96 n.AssertExpectations(t)
97 })
98}
99 
100func TestRevalidationExpiryDuringBlockedRead(t *testing.T) {
101 for _, limit := range []string{"grace", "grant", "certificate"} {
102 t.Run(limit, func(t *testing.T) {
103 synctest.Test(t, func(t *testing.T) {
104 s, n, sessions := revalidationFixture(t, 1)
105 sess := sessions[0]
106 switch limit {
107 case "grant":
108 sess.grantExpiry = time.Now().Add(time.Minute)
109 case "certificate":
110 sess.certNotAfter = time.Now().Add(time.Minute)
111 }
112 until := sess.authorizationDeadline(time.Now(), sess.grantExpiry)
113 require.True(t, sess.activate(t.Context(), until))
114 stream, err := sess.conn.OpenStream(protocol.Stream_DIRECT)
115 require.NoError(t, err)
116 defer stream.Close()
117 started := make(chan context.Context, 1)
118 n.On("Get", mock.Anything, mock.Anything).Run(func(args mock.Arguments) {
119 started <- args.Get(0).(context.Context)
120 time.Sleep(revalidateGrace) // Deliberately ignore the lookup deadline.
121 }).Return(revalidationRecord(t, sess), nil).Once()
122 done := make(chan struct{})
123 go func() {
124 defer close(done)
125 s.maintainSession(s.ParentContext, sess)
126 }()
127 readCtx := <-started
128 time.Sleep(time.Until(until))
129 synctest.Wait()
130 require.True(t, connectionDone(sess.conn))
131 require.Error(t, readCtx.Err())
132 s.sessions.mu.Lock()
133 pending := len(s.sessions.byConn)
134 s.sessions.mu.Unlock()
135 require.Equal(t, 1, pending, "unfinished reads must retain their admission slot")
136 _, err = stream.Read(make([]byte, 1))
137 require.Error(t, err, "expiry must close existing streams")
138 <-done // Let the late successful read return; it must not revive authority.
139 synctest.Wait()
140 s.sessions.mu.Lock()
141 pending = len(s.sessions.byConn)
142 s.sessions.mu.Unlock()
143 require.Zero(t, pending)
144 sess.mu.Lock()
145 state, finalDeadline := sess.state, sess.authorizedUntil
146 sess.mu.Unlock()
147 require.Equal(t, terminal, state)
148 require.Equal(t, until, finalDeadline)
149 n.AssertNotCalled(t, "PrefixContains", mock.Anything, mock.Anything, mock.Anything)
150 n.AssertExpectations(t)
151 })
152 })
153 }
154}
155 
156func TestRevalidationStopsWithSession(t *testing.T) {
157 for _, reason := range []string{"disconnect", "shutdown"} {
158 for _, pendingRead := range []bool{false, true} {
159 name := reason + "/idle"
160 if pendingRead {
161 name = reason + "/lookup"
162 }
163 t.Run(name, func(t *testing.T) {
164 synctest.Test(t, func(t *testing.T) {
165 s, n, sessions := revalidationFixture(t, 1)
166 sess := sessions[0]
167 require.True(t, sess.activate(t.Context(), time.Now().Add(revalidateGrace)))
168 ctx, cancel := context.WithCancel(s.ParentContext)
169 defer cancel()
170 started := make(chan context.Context, 1)
171 if pendingRead {
172 n.On("Get", mock.Anything, mock.Anything).Run(func(args mock.Arguments) {
173 readCtx := args.Get(0).(context.Context)
174 started <- readCtx
175 <-readCtx.Done()
176 }).Return(nil, context.Canceled).Once()
177 }
178 done := make(chan struct{})
179 go func() {
180 defer close(done)
181 s.maintainSession(ctx, sess)
182 }()
183 synctest.Wait()
184 var readCtx context.Context
185 if pendingRead {
186 readCtx = <-started
187 }
188 if reason == "disconnect" {
189 sess.conn.Close("disconnect")
190 } else {
191 cancel()
192 }
193 synctest.Wait()
194 select {
195 case <-done:
196 default:
197 t.Fatal("session worker did not stop")
198 }
199 require.True(t, connectionDone(sess.conn))
200 if pendingRead {
201 require.ErrorIs(t, readCtx.Err(), context.Canceled)
202 }
203 s.sessions.mu.Lock()
204 remaining := len(s.sessions.byConn)
205 s.sessions.mu.Unlock()
206 require.Zero(t, remaining)
207 n.AssertExpectations(t)
208 })
209 })
210 }
211 }
212}
213 
214func TestRevalidation(t *testing.T) {
215 for _, scenario := range []string{"transient", "revoked", "delayed success"} {
216 t.Run(scenario, func(t *testing.T) {
217 s, n, _, cert := sessionFixture(t)
218 ctx, pc := sessionContext(t, cert)
219 sess := &session{
220 alias: tun.SessionAlias([16]byte{1}),
221 hostname: "test",
222 conn: pc,
223 mode: delegated,
224 grantID: tun.DelegationID([32]byte{2}),
225 owner: &protocol.ClientToken{Token: []byte("owner")},
226 certNotAfter: time.Now().Add(time.Hour),
227 }
228 require.NoError(t, s.sessions.reserve(sess))
229 until := time.Now().Add(40 * time.Millisecond)
230 require.True(t, sess.activate(ctx, until))
231 switch scenario {
232 case "transient":
233 n.On("Get", mock.Anything, mock.Anything).Return(nil, errors.New("temporary failure")).Once()
234 s.revalidateSession(ctx, sess)
235 require.Equal(t, until, sess.authorizedUntil)
236 require.False(t, connectionDone(pc))
237 time.Sleep(time.Until(until) + time.Millisecond)
238 _, err := s.sessions.dial(sess.alias, sess.hostname)
239 require.ErrorIs(t, err, transport.ErrNoDirect)
240 s.revalidateSession(ctx, sess)
241 case "revoked":
242 n.On("Get", mock.Anything, mock.Anything).Return(nil, nil).Once()
243 s.revalidateSession(ctx, sess)
244 case "delayed success":
245 rec := &protocol.DelegationRecord{
246 Version: 1,
247 Id: sess.grantID,
248 Hostname: sess.hostname,
249 Owner: sess.owner,
250 }
251 data, _ := rec.MarshalVT()
252 n.On("Get", mock.Anything, mock.Anything).Run(func(mock.Arguments) { time.Sleep(10 * time.Millisecond) }).Return(data, nil).Once()
253 n.On("PrefixContains", mock.Anything, mock.Anything, mock.Anything).Return(true, nil).Once()
254 start := time.Now()
255 s.revalidateSession(ctx, sess)
256 require.True(t, sess.authorizedUntil.Before(start.Add(revalidateGrace+5*time.Millisecond)))
257 sess.close("test close")
258 previous := sess.authorizedUntil
259 sess.extend(time.Now().Add(time.Hour), time.Time{})
260 require.Equal(t, previous, sess.authorizedUntil)
261 }
262 require.True(t, connectionDone(pc))
263 n.AssertNotCalled(t, "Delete", mock.Anything, mock.Anything)
264 n.AssertExpectations(t)
265 })
266 }
267}