@@ -15,6 +15,7 @@ import ( //nolint:gci
1515 "github.com/pion/dtls/v2/pkg/crypto/prf"
1616 "github.com/pion/dtls/v2/pkg/protocol"
1717 "github.com/pion/dtls/v2/pkg/protocol/recordlayer"
18+ "golang.org/x/crypto/cryptobyte"
1819)
1920
2021// block ciphers using cipher block chaining.
@@ -64,18 +65,24 @@ func NewCBC(localKey, localWriteIV, localMac, remoteKey, remoteWriteIV, remoteMa
6465
6566// Encrypt encrypt a DTLS RecordLayer message
6667func (c * CBC ) Encrypt (pkt * recordlayer.RecordLayer , raw []byte ) ([]byte , error ) {
67- payload := raw [recordlayer . HeaderSize :]
68- raw = raw [:recordlayer . HeaderSize ]
68+ payload := raw [pkt . Header . Size () :]
69+ raw = raw [:pkt . Header . Size () ]
6970 blockSize := c .writeCBC .BlockSize ()
7071
7172 // Generate + Append MAC
7273 h := pkt .Header
7374
74- MAC , err := c .hmac (h .Epoch , h .SequenceNumber , h .ContentType , h .Version , payload , c .writeMac , c .h )
75+ var err error
76+ var mac []byte
77+ if h .ContentType == protocol .ContentTypeConnectionID {
78+ mac , err = c .hmacCID (h .Epoch , h .SequenceNumber , h .Version , payload , c .writeMac , c .h , h .ConnectionID )
79+ } else {
80+ mac , err = c .hmac (h .Epoch , h .SequenceNumber , h .ContentType , h .Version , payload , c .writeMac , c .h )
81+ }
7582 if err != nil {
7683 return nil , err
7784 }
78- payload = append (payload , MAC ... )
85+ payload = append (payload , mac ... )
7986
8087 // Generate + Append padding
8188 padding := make ([]byte , blockSize - len (payload )% blockSize )
@@ -96,26 +103,26 @@ func (c *CBC) Encrypt(pkt *recordlayer.RecordLayer, raw []byte) ([]byte, error)
96103 c .writeCBC .CryptBlocks (payload , payload )
97104 payload = append (iv , payload ... )
98105
99- // Prepend unencrypte header with encrypted payload
106+ // Prepend unencrypted header with encrypted payload
100107 raw = append (raw , payload ... )
101108
102109 // Update recordLayer size to include IV+MAC+Padding
103- binary .BigEndian .PutUint16 (raw [recordlayer . HeaderSize - 2 :], uint16 (len (raw )- recordlayer . HeaderSize ))
110+ binary .BigEndian .PutUint16 (raw [pkt . Header . Size () - 2 :], uint16 (len (raw )- pkt . Header . Size () ))
104111
105112 return raw , nil
106113}
107114
108115// Decrypt decrypts a DTLS RecordLayer message
109- func (c * CBC ) Decrypt (in []byte ) ([]byte , error ) {
110- body := in [recordlayer .HeaderSize :]
116+ func (c * CBC ) Decrypt (h recordlayer.Header , in []byte ) ([]byte , error ) {
111117 blockSize := c .readCBC .BlockSize ()
112118 mac := c .h ()
113119
114- var h recordlayer.Header
115- err := h .Unmarshal (in )
116- switch {
117- case err != nil :
120+ if err := h .Unmarshal (in ); err != nil {
118121 return nil , err
122+ }
123+ body := in [h .Size ():]
124+
125+ switch {
119126 case h .ContentType == protocol .ContentTypeChangeCipherSpec :
120127 // Nothing to encrypt with ChangeCipherSpec
121128 return in , nil
@@ -145,14 +152,19 @@ func (c *CBC) Decrypt(in []byte) ([]byte, error) {
145152 dataEnd := len (body ) - macSize - paddingLen
146153
147154 expectedMAC := body [dataEnd : dataEnd + macSize ]
148- actualMAC , err := c .hmac (h .Epoch , h .SequenceNumber , h .ContentType , h .Version , body [:dataEnd ], c .readMac , c .h )
149-
155+ var err error
156+ var actualMAC []byte
157+ if h .ContentType == protocol .ContentTypeConnectionID {
158+ actualMAC , err = c .hmacCID (h .Epoch , h .SequenceNumber , h .Version , body [:dataEnd ], c .readMac , c .h , h .ConnectionID )
159+ } else {
160+ actualMAC , err = c .hmac (h .Epoch , h .SequenceNumber , h .ContentType , h .Version , body [:dataEnd ], c .readMac , c .h )
161+ }
150162 // Compute Local MAC and compare
151163 if err != nil || ! hmac .Equal (actualMAC , expectedMAC ) {
152164 return nil , errInvalidMAC
153165 }
154166
155- return append (in [:recordlayer . HeaderSize ], body [:dataEnd ]... ), nil
167+ return append (in [:h . Size () ], body [:dataEnd ]... ), nil
156168}
157169
158170func (c * CBC ) hmac (epoch uint16 , sequenceNumber uint64 , contentType protocol.ContentType , protocolVersion protocol.Version , payload []byte , key []byte , hf func () hash.Hash ) ([]byte , error ) {
@@ -169,7 +181,45 @@ func (c *CBC) hmac(epoch uint16, sequenceNumber uint64, contentType protocol.Con
169181
170182 if _ , err := h .Write (msg ); err != nil {
171183 return nil , err
172- } else if _ , err := h .Write (payload ); err != nil {
184+ }
185+ if _ , err := h .Write (payload ); err != nil {
186+ return nil , err
187+ }
188+
189+ return h .Sum (nil ), nil
190+ }
191+
192+ // hmacCID calculates a MAC according to
193+ // https://datatracker.ietf.org/doc/html/rfc9146#section-5.1
194+ func (c * CBC ) hmacCID (epoch uint16 , sequenceNumber uint64 , protocolVersion protocol.Version , payload []byte , key []byte , hf func () hash.Hash , cid []byte ) ([]byte , error ) {
195+ // Must unmarshal inner plaintext in orde to perform MAC.
196+ ip := & recordlayer.InnerPlaintext {}
197+ if err := ip .Unmarshal (payload ); err != nil {
198+ return nil , err
199+ }
200+
201+ h := hmac .New (hf , key )
202+
203+ var msg cryptobyte.Builder
204+
205+ msg .AddUint64 (seqNumPlaceholder )
206+ msg .AddUint8 (uint8 (protocol .ContentTypeConnectionID ))
207+ msg .AddUint8 (uint8 (len (cid )))
208+ msg .AddUint8 (uint8 (protocol .ContentTypeConnectionID ))
209+ msg .AddUint8 (protocolVersion .Major )
210+ msg .AddUint8 (protocolVersion .Minor )
211+ msg .AddUint16 (epoch )
212+ util .AddUint48 (& msg , sequenceNumber )
213+ msg .AddBytes (cid )
214+ msg .AddUint16 (uint16 (len (payload )))
215+ msg .AddBytes (ip .Content )
216+ msg .AddUint8 (uint8 (ip .RealType ))
217+ msg .AddBytes (make ([]byte , ip .Zeros ))
218+
219+ if _ , err := h .Write (msg .BytesOrPanic ()); err != nil {
220+ return nil , err
221+ }
222+ if _ , err := h .Write (payload ); err != nil {
173223 return nil , err
174224 }
175225
0 commit comments