package handshake
import (
"crypto"
"crypto/cipher"
"fmt"
tls "github.com/tmc/go-iroh/internal/itls/tls"
"sync/atomic"
"github.com/tmc/go-iroh/internal/qng/internal/monotime"
"github.com/tmc/go-iroh/internal/qng/internal/protocol"
"github.com/tmc/go-iroh/internal/qng/internal/qerr"
"github.com/tmc/go-iroh/internal/qng/internal/utils"
"github.com/tmc/go-iroh/internal/qng/qlog"
"github.com/tmc/go-iroh/internal/qng/qlogwriter"
)
var keyUpdateInterval atomic .Uint64
func init() {
keyUpdateInterval .Store (protocol .KeyUpdateInterval )
}
func SetKeyUpdateInterval (v uint64 ) (reset func ()) {
old := keyUpdateInterval .Swap (v )
return func () { keyUpdateInterval .Store (old ) }
}
var FirstKeyUpdateInterval uint64 = 100
type updatableAEAD struct {
suite cipherSuite
keyPhase protocol .KeyPhase
largestAcked protocol .PacketNumber
firstPacketNumber protocol .PacketNumber
handshakeConfirmed bool
invalidPacketLimit uint64
invalidPacketCount uint64
prevRcvAEADExpiry monotime .Time
prevRcvAEAD cipher .AEAD
firstRcvdWithCurrentKey protocol .PacketNumber
firstSentWithCurrentKey protocol .PacketNumber
highestRcvdPN map [protocol .PathID ]protocol .PacketNumber
numRcvdWithCurrentKey uint64
numSentWithCurrentKey uint64
rcvAEAD cipher .AEAD
sendAEAD cipher .AEAD
aeadOverhead int
nextRcvAEAD cipher .AEAD
nextSendAEAD cipher .AEAD
nextRcvTrafficSecret []byte
nextSendTrafficSecret []byte
headerDecrypter headerProtector
headerEncrypter headerProtector
rttStats *utils .RTTStats
qlogger qlogwriter .Recorder
logger utils .Logger
version protocol .Version
nonceBuf []byte
}
var (
_ ShortHeaderOpener = &updatableAEAD {}
_ ShortHeaderSealer = &updatableAEAD {}
)
func newUpdatableAEAD(rttStats *utils .RTTStats , qlogger qlogwriter .Recorder , logger utils .Logger , version protocol .Version ) *updatableAEAD {
return &updatableAEAD {
firstPacketNumber : protocol .InvalidPacketNumber ,
largestAcked : protocol .InvalidPacketNumber ,
firstRcvdWithCurrentKey : protocol .InvalidPacketNumber ,
firstSentWithCurrentKey : protocol .InvalidPacketNumber ,
highestRcvdPN : map [protocol .PathID ]protocol .PacketNumber {},
rttStats : rttStats ,
qlogger : qlogger ,
logger : logger ,
version : version ,
}
}
func (a *updatableAEAD ) rollKeys () {
if a .prevRcvAEAD != nil {
a .logger .Debugf ("Dropping key phase %d ahead of scheduled time. Drop time was: %s" , a .keyPhase -1 , a .prevRcvAEADExpiry )
if a .qlogger != nil {
a .qlogger .RecordEvent (qlog .KeyDiscarded {
KeyType : qlog .KeyTypeClient1RTT ,
KeyPhase : a .keyPhase - 1 ,
})
a .qlogger .RecordEvent (qlog .KeyDiscarded {
KeyType : qlog .KeyTypeServer1RTT ,
KeyPhase : a .keyPhase - 1 ,
})
}
a .prevRcvAEADExpiry = 0
}
a .keyPhase ++
a .firstRcvdWithCurrentKey = protocol .InvalidPacketNumber
a .firstSentWithCurrentKey = protocol .InvalidPacketNumber
a .numRcvdWithCurrentKey = 0
a .numSentWithCurrentKey = 0
a .prevRcvAEAD = a .rcvAEAD
a .rcvAEAD = a .nextRcvAEAD
a .sendAEAD = a .nextSendAEAD
a .nextRcvTrafficSecret = a .getNextTrafficSecret (a .suite .Hash , a .nextRcvTrafficSecret )
a .nextSendTrafficSecret = a .getNextTrafficSecret (a .suite .Hash , a .nextSendTrafficSecret )
a .nextRcvAEAD = createAEAD (a .suite , a .nextRcvTrafficSecret , a .version )
a .nextSendAEAD = createAEAD (a .suite , a .nextSendTrafficSecret , a .version )
}
func (a *updatableAEAD ) startKeyDropTimer (now monotime .Time ) {
d := 3 * a .rttStats .PTO (true )
a .logger .Debugf ("Starting key drop timer to drop key phase %d (in %s)" , a .keyPhase -1 , d )
a .prevRcvAEADExpiry = now .Add (d )
}
func (a *updatableAEAD ) getNextTrafficSecret (hash crypto .Hash , ts []byte ) []byte {
return hkdfExpandLabel (hash , ts , []byte {}, "quic ku" , hash .Size ())
}
func (a *updatableAEAD ) SetReadKey (suite cipherSuite , trafficSecret []byte ) {
a .rcvAEAD = createAEAD (suite , trafficSecret , a .version )
a .headerDecrypter = newHeaderProtector (suite , trafficSecret , false , a .version )
if a .suite .ID == 0 {
a .setAEADParameters (a .rcvAEAD , suite )
}
a .nextRcvTrafficSecret = a .getNextTrafficSecret (suite .Hash , trafficSecret )
a .nextRcvAEAD = createAEAD (suite , a .nextRcvTrafficSecret , a .version )
}
func (a *updatableAEAD ) SetWriteKey (suite cipherSuite , trafficSecret []byte ) {
a .sendAEAD = createAEAD (suite , trafficSecret , a .version )
a .headerEncrypter = newHeaderProtector (suite , trafficSecret , false , a .version )
if a .suite .ID == 0 {
a .setAEADParameters (a .sendAEAD , suite )
}
a .nextSendTrafficSecret = a .getNextTrafficSecret (suite .Hash , trafficSecret )
a .nextSendAEAD = createAEAD (suite , a .nextSendTrafficSecret , a .version )
}
func (a *updatableAEAD ) setAEADParameters (aead cipher .AEAD , suite cipherSuite ) {
a .nonceBuf = make ([]byte , aeadNonceLength )
a .aeadOverhead = aead .Overhead ()
a .suite = suite
switch suite .ID {
case tls .TLS_AES_128_GCM_SHA256 , tls .TLS_AES_256_GCM_SHA384 :
a .invalidPacketLimit = protocol .InvalidPacketLimitAES
case tls .TLS_CHACHA20_POLY1305_SHA256 :
a .invalidPacketLimit = protocol .InvalidPacketLimitChaCha
default :
panic (fmt .Sprintf ("unknown cipher suite %d" , suite .ID ))
}
}
func (a *updatableAEAD ) DecodePacketNumber (pid protocol .PathID , wirePN protocol .PacketNumber , wirePNLen protocol .PacketNumberLen ) protocol .PacketNumber {
return protocol .DecodePacketNumber (wirePNLen , a .highestRcvdPN [pid ], wirePN )
}
func (a *updatableAEAD ) Open (dst , src []byte , rcvTime monotime .Time , pid protocol .PathID , pn protocol .PacketNumber , kp protocol .KeyPhaseBit , ad []byte ) ([]byte , error ) {
dec , err := a .open (dst , src , rcvTime , pid , pn , kp , ad )
if err == ErrDecryptionFailed {
a .invalidPacketCount ++
if a .invalidPacketCount >= a .invalidPacketLimit {
return nil , &qerr .TransportError {ErrorCode : qerr .AEADLimitReached }
}
}
if err == nil {
a .highestRcvdPN [pid ] = max (a .highestRcvdPN [pid ], pn )
}
return dec , err
}
func (a *updatableAEAD ) open (dst , src []byte , rcvTime monotime .Time , pid protocol .PathID , pn protocol .PacketNumber , kp protocol .KeyPhaseBit , ad []byte ) ([]byte , error ) {
if a .prevRcvAEAD != nil && !a .prevRcvAEADExpiry .IsZero () && rcvTime .After (a .prevRcvAEADExpiry ) {
a .prevRcvAEAD = nil
a .logger .Debugf ("Dropping key phase %d" , a .keyPhase -1 )
a .prevRcvAEADExpiry = 0
if a .qlogger != nil {
a .qlogger .RecordEvent (qlog .KeyDiscarded {
KeyType : qlog .KeyTypeClient1RTT ,
KeyPhase : a .keyPhase - 1 ,
})
a .qlogger .RecordEvent (qlog .KeyDiscarded {
KeyType : qlog .KeyTypeServer1RTT ,
KeyPhase : a .keyPhase - 1 ,
})
}
}
nonce := putPathNonce (a .nonceBuf , pid , pn )
if kp != a .keyPhase .Bit () {
if a .keyPhase > 0 && a .firstRcvdWithCurrentKey == protocol .InvalidPacketNumber || pn < a .firstRcvdWithCurrentKey {
if a .prevRcvAEAD == nil {
return nil , ErrKeysDropped
}
dec , err := a .prevRcvAEAD .Open (dst , nonce , src , ad )
if err != nil {
err = ErrDecryptionFailed
}
return dec , err
}
dec , err := a .nextRcvAEAD .Open (dst , nonce , src , ad )
if err != nil {
return nil , ErrDecryptionFailed
}
if a .keyPhase > 0 && a .firstSentWithCurrentKey == protocol .InvalidPacketNumber {
return nil , &qerr .TransportError {
ErrorCode : qerr .KeyUpdateError ,
ErrorMessage : "keys updated too quickly" ,
}
}
a .rollKeys ()
a .logger .Debugf ("Peer updated keys to %d" , a .keyPhase )
a .startKeyDropTimer (rcvTime )
if a .qlogger != nil {
a .qlogger .RecordEvent (qlog .KeyUpdated {
Trigger : qlog .KeyUpdateRemote ,
KeyType : qlog .KeyTypeClient1RTT ,
KeyPhase : a .keyPhase ,
})
a .qlogger .RecordEvent (qlog .KeyUpdated {
Trigger : qlog .KeyUpdateRemote ,
KeyType : qlog .KeyTypeServer1RTT ,
KeyPhase : a .keyPhase ,
})
}
a .firstRcvdWithCurrentKey = pn
return dec , err
}
dec , err := a .rcvAEAD .Open (dst , nonce , src , ad )
if err != nil {
return dec , ErrDecryptionFailed
}
a .numRcvdWithCurrentKey ++
if a .firstRcvdWithCurrentKey == protocol .InvalidPacketNumber {
if a .keyPhase > 0 {
a .logger .Debugf ("Peer confirmed key update to phase %d" , a .keyPhase )
a .startKeyDropTimer (rcvTime )
}
a .firstRcvdWithCurrentKey = pn
}
return dec , err
}
func (a *updatableAEAD ) Seal (dst , src []byte , pid protocol .PathID , pn protocol .PacketNumber , ad []byte ) []byte {
if a .firstSentWithCurrentKey == protocol .InvalidPacketNumber {
a .firstSentWithCurrentKey = pn
}
if a .firstPacketNumber == protocol .InvalidPacketNumber {
a .firstPacketNumber = pn
}
a .numSentWithCurrentKey ++
return a .sendAEAD .Seal (dst , putPathNonce (a .nonceBuf , pid , pn ), src , ad )
}
func (a *updatableAEAD ) SetLargestAcked (pn protocol .PacketNumber ) error {
if a .firstSentWithCurrentKey != protocol .InvalidPacketNumber &&
pn >= a .firstSentWithCurrentKey && a .numRcvdWithCurrentKey == 0 {
return &qerr .TransportError {
ErrorCode : qerr .KeyUpdateError ,
ErrorMessage : fmt .Sprintf ("received ACK for key phase %d, but peer didn't update keys" , a .keyPhase ),
}
}
a .largestAcked = pn
return nil
}
func (a *updatableAEAD ) SetHandshakeConfirmed () {
a .handshakeConfirmed = true
}
func (a *updatableAEAD ) updateAllowed () bool {
if !a .handshakeConfirmed {
return false
}
return a .keyPhase == 0 ||
(a .firstSentWithCurrentKey != protocol .InvalidPacketNumber &&
a .largestAcked != protocol .InvalidPacketNumber &&
a .largestAcked >= a .firstSentWithCurrentKey )
}
func (a *updatableAEAD ) shouldInitiateKeyUpdate () bool {
if !a .updateAllowed () {
return false
}
if a .keyPhase == 0 {
if a .numRcvdWithCurrentKey >= FirstKeyUpdateInterval || a .numSentWithCurrentKey >= FirstKeyUpdateInterval {
return true
}
}
if a .numRcvdWithCurrentKey >= keyUpdateInterval .Load () {
a .logger .Debugf ("Received %d packets with current key phase. Initiating key update to the next key phase: %d" , a .numRcvdWithCurrentKey , a .keyPhase +1 )
return true
}
if a .numSentWithCurrentKey >= keyUpdateInterval .Load () {
a .logger .Debugf ("Sent %d packets with current key phase. Initiating key update to the next key phase: %d" , a .numSentWithCurrentKey , a .keyPhase +1 )
return true
}
return false
}
func (a *updatableAEAD ) KeyPhase () protocol .KeyPhaseBit {
if a .shouldInitiateKeyUpdate () {
a .rollKeys ()
if a .qlogger != nil {
a .qlogger .RecordEvent (qlog .KeyUpdated {
Trigger : qlog .KeyUpdateLocal ,
KeyType : qlog .KeyTypeClient1RTT ,
KeyPhase : a .keyPhase ,
})
a .qlogger .RecordEvent (qlog .KeyUpdated {
Trigger : qlog .KeyUpdateLocal ,
KeyType : qlog .KeyTypeServer1RTT ,
KeyPhase : a .keyPhase ,
})
}
}
return a .keyPhase .Bit ()
}
func (a *updatableAEAD ) Overhead () int {
return a .aeadOverhead
}
func (a *updatableAEAD ) EncryptHeader (sample []byte , firstByte *byte , hdrBytes []byte ) {
a .headerEncrypter .EncryptHeader (sample , firstByte , hdrBytes )
}
func (a *updatableAEAD ) DecryptHeader (sample []byte , firstByte *byte , hdrBytes []byte ) {
a .headerDecrypter .DecryptHeader (sample , firstByte , hdrBytes )
}
func (a *updatableAEAD ) FirstPacketNumber () protocol .PacketNumber {
return a .firstPacketNumber
}
The pages are generated with Golds v0.8.4 . (GOOS=linux GOARCH=amd64)
Golds is a Go 101 project developed by Tapir Liu .
PR and bug reports are welcome and can be submitted to the issue list .
Please follow @zigo_101 (reachable from the left QR code) to get the latest news of Golds .