Skip to content
File

Blob: tun/server/route_cache.go

go243 lines
1package server
2 
3import (
4 "context"
5 "io/fs"
6 "sort"
7 "sync/atomic"
8 "time"
9 
10 "go.miragespace.co/specter/spec/chord"
11 "go.miragespace.co/specter/spec/protocol"
12 "go.miragespace.co/specter/spec/tun"
13 "go.miragespace.co/specter/util/promise"
14 
15 "github.com/Yiling-J/theine-go"
16 "go.uber.org/zap"
17)
18 
19type routesResult struct {
20 err error
21 routes []*protocol.TunnelRoute
22 refreshAfter time.Time
23}
24 
25const (
26 routeCacheBytes = 1 << 20 // 1MiB
27 routePositiveTTL = time.Minute * 5
28 routeNegativeTTL = time.Second * 15
29 routeFailedTTL = time.Second * 5
30)
31 
32func (s *Server) initRouteCache() {
33 routeCache, err := theine.NewBuilder[string, *routesResult](routeCacheBytes).Build()
34 
35 if err != nil {
36 panic("BUG: " + err.Error())
37 }
38 
39 s.routeCache = routeCache
40}
41 
42func (s *Server) RoutesPreload(hostname string) {
43 s.lookupRoutes(s.ParentContext, hostname, nil)
44}
45 
46// A stale result requests one refresh, unless another caller already replaced it
47// or the most recent refresh is still in its cooldown. Keep results immutable so
48// slow callers can identify the exact generation they exhausted.
49func (s *Server) lookupRoutes(ctx context.Context, hostname string, stale *routesResult) (*routesResult, error) {
50 if err := ctx.Err(); err != nil {
51 return nil, err
52 }
53 cached := func() (*routesResult, bool) {
54 ret, ok := s.routeCache.Get(hostname)
55 return ret, ok && (ret != stale || time.Now().Before(ret.refreshAfter))
56 }
57 if ret, ok := cached(); ok {
58 return ret, nil
59 }
60 
61 result := s.routeLoads.DoChan(hostname, func() (any, error) {
62 // Recheck after joining the flight: a delayed caller must not refresh a
63 // newer result just because its own connection attempts took longer.
64 if ret, ok := cached(); ok {
65 return ret, nil
66 }
67 // The loader has its own timeout. A cancelled waiter must not cancel
68 // the shared lookup or poison the cache for subsequent requests.
69 loaded := s.routeCacheLoader(s.ParentContext, hostname)
70 if stale != nil {
71 loaded.Value.refreshAfter = time.Now().Add(routeFailedTTL)
72 }
73 ret := &loaded.Value
74 s.routeCache.SetWithTTL(hostname, ret, loaded.Cost, loaded.TTL)
75 return ret, nil
76 })
77 select {
78 case <-ctx.Done():
79 return nil, ctx.Err()
80 case <-s.ParentContext.Done():
81 return nil, s.ParentContext.Err()
82 case result := <-result:
83 return result.Val.(*routesResult), nil
84 }
85}
86 
87func (s *Server) routeCacheLoader(ctx context.Context, hostname string) (ret theine.Loaded[routesResult]) {
88 start := time.Now()
89 defer func() {
90 s.Logger.Debug("Route cache loader invoked",
91 zap.String("hostname", hostname),
92 zap.Duration("duration", time.Since(start)),
93 zap.Bool("error", ret.Value.err != nil),
94 zap.Int64("cost", ret.Cost),
95 zap.Duration("ttl", ret.TTL),
96 )
97 }()
98 
99 if home, inv, ok := tun.ParseEphemeralLabel(hostname); ok {
100 return s.ephemeralRouteLoader(ctx, hostname, home, inv)
101 }
102 
103 var (
104 numNotFound = 0
105 numError = 0
106 numLookup = tun.NumRedundantLinks
107 lookupJobs = make([]func(context.Context) (*protocol.TunnelRoute, error), tun.NumRedundantLinks)
108 )
109 
110 for i := range lookupJobs {
111 k := i + 1
112 lookupJobs[i] = func(ctx context.Context) (*protocol.TunnelRoute, error) {
113 key := tun.RoutingKey(hostname, k)
114 val, err := s.Chord.Get(ctx, []byte(key))
115 if err != nil {
116 return nil, err
117 }
118 if len(val) == 0 {
119 return nil, fs.ErrNotExist
120 }
121 route := &protocol.TunnelRoute{}
122 if err := route.UnmarshalVT(val); err != nil {
123 return nil, err
124 }
125 atomic.AddInt64(&ret.Cost, int64(len(val)))
126 return route, nil
127 }
128 }
129 
130 lookupCtx, lookupCancel := context.WithTimeout(ctx, lookupTimeout)
131 defer lookupCancel()
132 
133 routes, errors := promise.All(lookupCtx, lookupJobs...)
134 for _, err := range errors {
135 switch err {
136 case nil:
137 case fs.ErrNotExist:
138 numNotFound++
139 default:
140 numError++
141 }
142 }
143 
144 if numLookup == numNotFound {
145 ret.Value.err = tun.ErrDestinationNotFound
146 ret.TTL = routeNegativeTTL // cache negative result with shorter ttl
147 ret.Cost = 8 // use a (1 pointer) cost for negative result
148 return
149 }
150 
151 if numLookup == numNotFound+numError {
152 ret.Value.err = tun.ErrLookupFailed
153 ret.TTL = routeFailedTTL // also cache failed result with an even shorter ttl
154 ret.Cost = 16 // use a (2 pointers) cost for failed result
155 return
156 }
157 
158 // we don't know which one error'd, need to filter nil routes
159 filtered := routes[:0]
160 for _, route := range routes {
161 if route != nil {
162 filtered = append(filtered, route)
163 }
164 }
165 // need to nil the elements for gc, if any
166 // see https://github.com/golang/go/wiki/SliceTricks#filtering-without-allocating
167 for i := len(filtered); i < len(routes); i++ {
168 routes[i] = nil
169 }
170 
171 // prioritize directly connected route
172 localAddress := s.TunnelTransport.Identity().GetAddress()
173 sort.SliceStable(filtered, func(i, j int) bool {
174 return filtered[i].GetTunnelDestination().GetAddress() == localAddress &&
175 filtered[j].GetTunnelDestination().GetAddress() != localAddress
176 })
177 
178 // now we can store the routes on a longer ttl
179 ret.Value.routes = filtered
180 ret.TTL = routePositiveTTL
181 // cost is added atomically during lookup
182 
183 return
184}
185 
186func (s *Server) ephemeralRouteLoader(ctx context.Context, label string, home uint64, inv [16]byte) (ret theine.Loaded[routesResult]) {
187 ret.Cost, ret.TTL, ret.Value.err = 256, routeFailedTTL, tun.ErrLookupFailed
188 ctx, cancel := context.WithTimeout(ctx, lookupTimeout)
189 defer cancel()
190 var dst *protocol.TunnelDestination
191 if home == s.Chord.ID() {
192 dst = &protocol.TunnelDestination{
193 Chord: s.ChordTransport.Identity(),
194 Tunnel: s.TunnelTransport.Identity(),
195 }
196 } else {
197 select {
198 case s.ephemeralLoads <- struct{}{}:
199 default:
200 return
201 }
202 type result struct {
203 node chord.VNode
204 err error
205 }
206 done := make(chan result, 1)
207 go func() {
208 defer func() { <-s.ephemeralLoads }()
209 node, err := s.Chord.FindSuccessor(home)
210 done <- result{node, err}
211 }()
212 select {
213 case <-ctx.Done():
214 return
215 case found := <-done:
216 if found.err != nil {
217 return
218 }
219 if found.node == nil || found.node.ID() != home {
220 ret.TTL, ret.Value.err = routeNegativeTTL, tun.ErrDestinationNotFound
221 return
222 }
223 var err error
224 dst, err = s.lookupDestination(ctx, tun.DestinationByChordKey(found.node.Identity()))
225 if err != nil {
226 return
227 }
228 }
229 }
230 route := &protocol.TunnelRoute{
231 ClientDestination: &protocol.Node{
232 Address: tun.SessionAlias(inv),
233 Rendezvous: true,
234 },
235 ChordDestination: dst.GetChord(),
236 TunnelDestination: dst.GetTunnel(),
237 Hostname: label,
238 }
239 ret.Cost += int64(route.SizeVT())
240 ret.Value.routes, ret.Value.err, ret.TTL = []*protocol.TunnelRoute{route}, nil, routePositiveTTL
241 return
242}