File
Blob: overlay/reuse.go
| 1 | package overlay |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "fmt" |
| 6 | "time" |
| 7 | |
| 8 | "go.miragespace.co/specter/spec" |
| 9 | "go.miragespace.co/specter/spec/pki" |
| 10 | "go.miragespace.co/specter/spec/protocol" |
| 11 | "go.miragespace.co/specter/spec/rpc" |
| 12 | "go.miragespace.co/specter/spec/transport" |
| 13 | |
| 14 | "github.com/quic-go/quic-go" |
| 15 | "go.uber.org/zap" |
| 16 | ) |
| 17 | |
| 18 | const ( |
| 19 | reuseErrorState = "invalid state" |
| 20 | ) |
| 21 | |
| 22 | func wrapReuseError(msg string) error { |
| 23 | return fmt.Errorf("%s: %s", reuseErrorState, msg) |
| 24 | } |
| 25 | |
| 26 | func (t *QUIC) reuseConnection(_ context.Context, q *quic.Conn, s *quic.Stream, dir direction) (*nodeConnection, bool, error) { |
| 27 | negotiation := &protocol.Connection{ |
| 28 | Identity: t.Endpoint, |
| 29 | Version: spec.BuildVersion, |
| 30 | } |
| 31 | |
| 32 | err := rpc.Send(s, negotiation) |
| 33 | if err != nil { |
| 34 | return nil, false, fmt.Errorf("error sending identity: %w", err) |
| 35 | } |
| 36 | negotiation.Reset() |
| 37 | |
| 38 | s.SetReadDeadline(time.Now().Add(quicConfig.HandshakeIdleTimeout)) |
| 39 | err = rpc.BoundedReceive(s, negotiation, 256) |
| 40 | if err != nil { |
| 41 | return nil, false, fmt.Errorf("error receiving identity: %w", err) |
| 42 | } |
| 43 | s.SetReadDeadline(time.Time{}) |
| 44 | |
| 45 | if t.UseCertificateIdentity { |
| 46 | chain := q.ConnectionState().TLS.VerifiedChains |
| 47 | if len(chain) == 0 || len(chain[0]) == 0 { |
| 48 | return nil, false, transport.ErrNoCertificate |
| 49 | } |
| 50 | cert := chain[0][0] |
| 51 | identity, err := pki.ExtractCertificateIdentity(cert) |
| 52 | if err != nil { |
| 53 | return nil, false, fmt.Errorf("extracting certificate identity: %w", err) |
| 54 | } |
| 55 | t.Logger.Debug("Using identify from certificate", zap.String("remote", q.RemoteAddr().String()), zap.Object("identity", identity)) |
| 56 | negotiation.Identity = identity.NodeIdentity() |
| 57 | } |
| 58 | |
| 59 | qKey := t.makeCachedKey(negotiation.GetIdentity()) |
| 60 | fresh := &nodeConnection{ |
| 61 | peer: negotiation.GetIdentity(), |
| 62 | quic: q, |
| 63 | direction: dir, |
| 64 | version: negotiation.GetVersion(), |
| 65 | } |
| 66 | |
| 67 | negotiation.Reset() |
| 68 | |
| 69 | rUnlock := t.cachedMutex.RLock(qKey) |
| 70 | cache, cached := t.cachedConnections.Load(qKey) |
| 71 | if cached { |
| 72 | negotiation.CacheState = protocol.Connection_CACHED |
| 73 | if cache.direction == directionIncoming { |
| 74 | negotiation.CacheDirection = protocol.Connection_INCOMING |
| 75 | } else { |
| 76 | negotiation.CacheDirection = protocol.Connection_OUTGOING |
| 77 | } |
| 78 | } else { |
| 79 | negotiation.CacheState = protocol.Connection_FRESH |
| 80 | if dir == directionIncoming { |
| 81 | negotiation.CacheDirection = protocol.Connection_INCOMING |
| 82 | } else { |
| 83 | negotiation.CacheDirection = protocol.Connection_OUTGOING |
| 84 | } |
| 85 | } |
| 86 | rUnlock() |
| 87 | |
| 88 | err = rpc.Send(s, negotiation) |
| 89 | if err != nil { |
| 90 | return nil, false, fmt.Errorf("error sending cache status: %w", err) |
| 91 | } |
| 92 | negotiation.Reset() |
| 93 | |
| 94 | s.SetReadDeadline(time.Now().Add(quicConfig.HandshakeIdleTimeout)) |
| 95 | err = rpc.BoundedReceive(s, negotiation, 8) |
| 96 | if err != nil { |
| 97 | return nil, false, fmt.Errorf("error receiving cache status: %w", err) |
| 98 | } |
| 99 | s.SetReadDeadline(time.Time{}) |
| 100 | |
| 101 | unlock := t.cachedMutex.Lock(qKey) |
| 102 | defer unlock() |
| 103 | |
| 104 | // I really should make a state machine for this |
| 105 | switch negotiation.CacheState { |
| 106 | case protocol.Connection_CACHED: |
| 107 | switch negotiation.CacheDirection { |
| 108 | case protocol.Connection_INCOMING: |
| 109 | if cached { |
| 110 | if cache.direction == directionIncoming { |
| 111 | // other: cached incoming |
| 112 | // us: cached incoming |
| 113 | return nil, false, wrapReuseError("both peers have cached incoming connections") |
| 114 | } else { |
| 115 | // other: cached incoming |
| 116 | // us: cached outgoing |
| 117 | fresh.quic.CloseWithError(508, wrapReuseError("previously cached connection was reused").Error()) |
| 118 | return cache, true, nil |
| 119 | } |
| 120 | } else { |
| 121 | if dir == directionIncoming { |
| 122 | // other: cached incoming |
| 123 | // us: new incoming |
| 124 | return nil, false, wrapReuseError("both peers have incoming connections") |
| 125 | } else { |
| 126 | // other: cached incoming |
| 127 | // us: new outgoing |
| 128 | return nil, false, wrapReuseError("other peer has cached connection while we are establishing a new outgoing connection") |
| 129 | } |
| 130 | } |
| 131 | case protocol.Connection_OUTGOING: |
| 132 | if cached { |
| 133 | if cache.direction == directionIncoming { |
| 134 | // other: cached outgoing |
| 135 | // us: cached incoming |
| 136 | // we will let the receiver side close the connection |
| 137 | return cache, true, nil |
| 138 | } else { |
| 139 | // other: cached outgoing |
| 140 | // us: cached outgoing |
| 141 | return nil, false, wrapReuseError("both peers have cached outgoing connections") |
| 142 | } |
| 143 | } else { |
| 144 | if dir == directionIncoming { |
| 145 | // other: cached outgoing |
| 146 | // us: new incoming |
| 147 | return nil, false, wrapReuseError("other peer has cached connection while we are handling a new incoming connection") |
| 148 | } else { |
| 149 | // other: cached outgoing |
| 150 | // us: new outgoing |
| 151 | return nil, false, wrapReuseError("both peers have outgoing connections") |
| 152 | } |
| 153 | } |
| 154 | default: |
| 155 | return nil, false, fmt.Errorf("unknown transport cache direction") |
| 156 | } |
| 157 | case protocol.Connection_FRESH: |
| 158 | switch negotiation.CacheDirection { |
| 159 | case protocol.Connection_INCOMING: |
| 160 | if cached { |
| 161 | if cache.direction == directionIncoming { |
| 162 | // other: new incoming |
| 163 | // us: cached incoming |
| 164 | return nil, false, wrapReuseError("both peers have cached incoming connections") |
| 165 | } else { |
| 166 | // other: new incoming |
| 167 | // us: cached outgoing |
| 168 | return nil, false, wrapReuseError("we have cached connection while peer is handling a new incoming connection") |
| 169 | } |
| 170 | } else { |
| 171 | if dir == directionIncoming { |
| 172 | // other: new incoming |
| 173 | // us: new incoming |
| 174 | return nil, false, wrapReuseError("both peers have incoming connections") |
| 175 | } else { |
| 176 | // need to check again because we released the lock |
| 177 | cache, cached = t.cachedConnections.Load(qKey) |
| 178 | if cached { |
| 179 | // other: new incoming |
| 180 | // us: cached |
| 181 | fresh.quic.CloseWithError(508, wrapReuseError("previously cached connection was reused").Error()) |
| 182 | return cache, true, nil |
| 183 | } else { |
| 184 | // other: new incoming |
| 185 | // us: new outgoing |
| 186 | t.cachedConnections.Store(qKey, fresh) |
| 187 | return fresh, false, nil |
| 188 | } |
| 189 | } |
| 190 | } |
| 191 | case protocol.Connection_OUTGOING: |
| 192 | if cached { |
| 193 | if cache.direction == directionIncoming { |
| 194 | // other: new outgoing |
| 195 | // us: cached incoming |
| 196 | return nil, false, wrapReuseError("we have cached connection while peer is establishing a new outgoing connection") |
| 197 | } else { |
| 198 | // other: new outgoing |
| 199 | // us: cached outgoing |
| 200 | return nil, false, wrapReuseError("both peers have outgoing connections") |
| 201 | } |
| 202 | } else { |
| 203 | if dir == directionIncoming { |
| 204 | // need to check again because we released the lock |
| 205 | cache, cached = t.cachedConnections.Load(qKey) |
| 206 | if cached { |
| 207 | // other: new outgoing |
| 208 | // us: cached |
| 209 | fresh.quic.CloseWithError(508, wrapReuseError("previously cached connection was reused").Error()) |
| 210 | return cache, true, nil |
| 211 | } else { |
| 212 | // other: new outgoing |
| 213 | // us: new incoming |
| 214 | t.cachedConnections.Store(qKey, fresh) |
| 215 | return fresh, false, nil |
| 216 | } |
| 217 | } else { |
| 218 | // other: new outgoing |
| 219 | // us: new outgoing |
| 220 | return nil, false, wrapReuseError("both peers have outgoing connections") |
| 221 | } |
| 222 | } |
| 223 | default: |
| 224 | return nil, false, fmt.Errorf("unknown transport cache direction") |
| 225 | } |
| 226 | default: |
| 227 | return nil, false, fmt.Errorf("unknown transport cache state") |
| 228 | } |
| 229 | } |