File
Blob: tun/server/revalidate_test.go
| 1 | package server |
| 2 | |
| 3 | import ( |
| 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 | |
| 19 | func 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 | |
| 48 | func 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 | |
| 64 | func 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 | |
| 100 | func 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 | |
| 156 | func 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 | |
| 214 | func 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 | } |