File
Blob: tun/server/keyless_cache.go
| 1 | package server |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "crypto/tls" |
| 6 | "crypto/x509" |
| 7 | "errors" |
| 8 | "time" |
| 9 | |
| 10 | "go.miragespace.co/specter/spec/rpc" |
| 11 | |
| 12 | "github.com/Yiling-J/theine-go" |
| 13 | "github.com/twitchtv/twirp" |
| 14 | "go.uber.org/zap" |
| 15 | ) |
| 16 | |
| 17 | type keylessCertResult struct { |
| 18 | err error |
| 19 | cert *tls.Certificate |
| 20 | } |
| 21 | |
| 22 | const ( |
| 23 | keylessCacheBytes = 1 << 24 // 16MiB |
| 24 | keylessPositiveTTL = 5 * time.Minute |
| 25 | keylessFailedTTL = 10 * time.Second |
| 26 | keylessExpirySkew = time.Minute |
| 27 | ) |
| 28 | |
| 29 | func (s *Server) initKeylessCache() { |
| 30 | cache, err := theine.NewBuilder[string, keylessCertResult](keylessCacheBytes). |
| 31 | BuildWithLoader(s.keylessCertLoader) |
| 32 | if err != nil { |
| 33 | panic("BUG: " + err.Error()) |
| 34 | } |
| 35 | |
| 36 | s.keylessCache = cache |
| 37 | } |
| 38 | |
| 39 | func computeKeylessTTL(cert *tls.Certificate, now time.Time) time.Duration { |
| 40 | if cert == nil { |
| 41 | return keylessPositiveTTL |
| 42 | } |
| 43 | |
| 44 | leaf := cert.Leaf |
| 45 | if leaf == nil && len(cert.Certificate) > 0 { |
| 46 | if parsed, err := x509.ParseCertificate(cert.Certificate[0]); err == nil { |
| 47 | leaf = parsed |
| 48 | } |
| 49 | } |
| 50 | if leaf == nil { |
| 51 | return keylessPositiveTTL |
| 52 | } |
| 53 | |
| 54 | expiry := leaf.NotAfter.Add(-keylessExpirySkew) |
| 55 | remaining := expiry.Sub(now) |
| 56 | if remaining <= 0 { |
| 57 | // certificate is effectively expired (or about to), use a very short ttl |
| 58 | return time.Second |
| 59 | } |
| 60 | |
| 61 | if remaining < keylessPositiveTTL { |
| 62 | return remaining |
| 63 | } |
| 64 | |
| 65 | return keylessPositiveTTL |
| 66 | } |
| 67 | |
| 68 | func (s *Server) keylessCertLoader(ctx context.Context, hostname string) (ret theine.Loaded[keylessCertResult], loadErr error) { |
| 69 | start := time.Now() |
| 70 | defer func() { |
| 71 | if s.Logger != nil { |
| 72 | s.Logger.Debug("Keyless certificate loader invoked", |
| 73 | zap.String("hostname", hostname), |
| 74 | zap.Duration("duration", time.Since(start)), |
| 75 | zap.Bool("err", ret.Value.err != nil), |
| 76 | zap.Int64("cost", ret.Cost), |
| 77 | zap.Duration("ttl", ret.TTL), |
| 78 | ) |
| 79 | } |
| 80 | }() |
| 81 | |
| 82 | delegation := rpc.GetDelegation(ctx) |
| 83 | if delegation == nil { |
| 84 | ret.Value.err = twirp.Internal.Error("delegation missing in context") |
| 85 | ret.TTL = keylessFailedTTL |
| 86 | return |
| 87 | } |
| 88 | |
| 89 | chi := &tls.ClientHelloInfo{ |
| 90 | ServerName: hostname, |
| 91 | Conn: delegation, |
| 92 | } |
| 93 | |
| 94 | cert, err := s.CertProvider.GetCertificateWithContext(ctx, chi) |
| 95 | if err != nil { |
| 96 | ret.Value.err = err |
| 97 | ret.TTL = keylessFailedTTL |
| 98 | return |
| 99 | } |
| 100 | if cert == nil || len(cert.Certificate) == 0 { |
| 101 | ret.Value.err = errors.New("no certificate returned from cert provider") |
| 102 | ret.TTL = keylessFailedTTL |
| 103 | return |
| 104 | } |
| 105 | |
| 106 | ret.Value.cert = cert |
| 107 | |
| 108 | now := time.Now() |
| 109 | ret.TTL = computeKeylessTTL(cert, now) |
| 110 | if ret.TTL <= 0 { |
| 111 | ret.TTL = keylessFailedTTL |
| 112 | } |
| 113 | |
| 114 | if len(cert.Certificate) > 0 { |
| 115 | ret.Cost = int64(len(cert.Certificate[0])) |
| 116 | if cert.Leaf != nil { |
| 117 | ret.Cost += int64(len(cert.Leaf.Raw)) |
| 118 | } |
| 119 | } else { |
| 120 | ret.Cost = 1 |
| 121 | } |
| 122 | |
| 123 | return |
| 124 | } |