@@ -573,7 +573,7 @@ func (c *Conn) Write(payload []byte) (int, error) {
573573 return 0 , err
574574 }
575575
576- ctx , cancel := c .contextWithClose (c . writeDeadline )
576+ ctx , cancel := c .contextWithClose ()
577577 defer cancel ()
578578
579579 err := c .writeApplicationData (ctx , []* dtlsflight.Outbound {
@@ -829,62 +829,94 @@ func (c *Conn) cacheHandshake(outbound *dtlsflight.Outbound, dtlsHandshake *hand
829829 return nil
830830}
831831
832- func (c * Conn ) contextWithClose (ctx context.Context ) (context.Context , context.CancelFunc ) {
833- closeCtx , cancel := context .WithCancelCause (context .WithoutCancel (ctx ))
834- go func () {
835- select {
836- case <- c .closed .Done ():
837- cancel (context .Canceled )
838- case <- ctx .Done ():
839- err := ctx .Err ()
840- if err == nil {
841- err = context .DeadlineExceeded
842- }
843- cancel (err )
844- case <- closeCtx .Done ():
845- }
846- }()
832+ func (c * Conn ) contextWithClose () (context.Context , context.CancelFunc ) {
833+ ctx := context .Background ()
834+
835+ var cancelDeadline context.CancelFunc = func () {}
836+ if deadline , ok := c .readDeadline .Deadline (); ok {
837+ ctx , cancelDeadline = context .WithDeadline (ctx , deadline )
838+ }
839+
840+ closeCtx , cancelClose := context .WithCancelCause (ctx )
841+ detachLifetime := context .AfterFunc (c .closed , func () {
842+ cancelClose (c .closed .Err ())
843+ })
847844
848845 return closeCtx , func () {
849- cancel (context .Canceled )
846+ detachLifetime ()
847+ cancelDeadline ()
848+ cancelClose (context .Canceled )
850849 }
851850}
852851
853852func (c * Conn ) contextWithCloseAndWriteDeadline (ctx context.Context ) (context.Context , context.CancelFunc ) {
854- operationCtx , cancel := context .WithCancelCause (context .Background ())
855- go func () {
856- select {
857- case <- c .closed .Done ():
858- cancel (context .Canceled )
859- case <- c .writeDeadline .Done ():
860- cancel (context .DeadlineExceeded )
861- case <- ctx .Done ():
862- cancel (ctx .Err ())
863- case <- operationCtx .Done ():
853+ if ctx == nil {
854+ ctx = context .Background ()
855+ }
856+
857+ var cancelDeadline context.CancelFunc = func () {}
858+ if deadline , ok := c .writeDeadline .Deadline (); ok {
859+ ctx , cancelDeadline = context .WithDeadline (ctx , deadline )
860+ }
861+
862+ operationCtx , cancelClose := context .WithCancelCause (ctx )
863+ detachLifetime := context .AfterFunc (c .closed , func () {
864+ err := c .closed .Err ()
865+ if err == nil {
866+ err = context .Canceled
864867 }
865- }()
868+ cancelClose (err )
869+ })
866870
867871 return operationCtx , func () {
868- cancel (context .Canceled )
872+ detachLifetime ()
873+ cancelDeadline ()
874+ cancelClose (context .Canceled )
869875 }
870876}
871877
872878func (c * Conn ) compactPreparedRecords (records []preparedRecord ) []preparedDatagram {
873- datagrams := make ([]preparedDatagram , 0 , len (records ))
874- current := preparedDatagram {}
879+ if len (records ) == 0 {
880+ return []preparedDatagram {}
881+ }
882+
883+ totalSize := 0
875884 for _ , record := range records {
876- if len (current .raw ) > 0 && len (current .raw )+ len (record .raw ) >= c .maximumTransmissionUnit {
877- datagrams = append (datagrams , current )
878- current = preparedDatagram {}
885+ totalSize += len (record .raw )
886+ }
887+
888+ datagrams := make ([]preparedDatagram , len (records ))
889+ flatRaw := make ([]byte , 0 , totalSize )
890+
891+ datagramIndex := 0
892+ currentSize := 0
893+ offset := 0
894+
895+ for _ , record := range records {
896+ recordSize := len (record .raw )
897+
898+ flatRaw = append (flatRaw , record .raw ... )
899+
900+ if currentSize > 0 && currentSize + recordSize >= c .maximumTransmissionUnit {
901+ datagrams [datagramIndex ].raw = flatRaw [offset : offset + currentSize ]
902+ datagramIndex ++
903+ offset += currentSize
904+ currentSize = 0
879905 }
880- current .raw = append (current .raw , record .raw ... )
906+
907+ currentSize += recordSize
908+
881909 if record .tracked != nil {
882- current .tracked = append (current .tracked , * record .tracked )
910+ datagrams [datagramIndex ].tracked = append (
911+ datagrams [datagramIndex ].tracked ,
912+ * record .tracked ,
913+ )
883914 }
884915 }
885- datagrams = append (datagrams , current )
886916
887- return datagrams
917+ datagrams [datagramIndex ].raw = flatRaw [offset : offset + currentSize ]
918+
919+ return datagrams [:datagramIndex + 1 ]
888920}
889921
890922func (c * Conn ) prepareRecord (outbound * dtlsflight.Outbound ) ([]byte , error ) {
@@ -1214,42 +1246,76 @@ func selectHandshakeFragment(offsets map[uint32]uint32, raw []byte) (bool, error
12141246 return ok && length == header .FragmentLength , nil
12151247}
12161248
1249+ var noContentFragments = [][]byte { //nolint:gochecknoglobals
1250+ {},
1251+ }
1252+
12171253func (c * Conn ) fragmentHandshake (dtlsHandshake * handshake.Handshake ) ([][]byte , error ) {
1254+ messageSize := dtlsHandshake .Message .MarshalSize ()
1255+ numFragments := (messageSize - 1 )/ c .maximumTransmissionUnit + 1
1256+
1257+ fragmentedHandshakes := make ([][]byte , numFragments )
1258+
1259+ if numFragments == 1 {
1260+ fragmentedHandshake := make ([]byte , handshake .HeaderLength + messageSize )
1261+
1262+ headerFragment := handshake.Header {
1263+ Type : dtlsHandshake .Header .Type ,
1264+ Length : dtlsHandshake .Header .Length ,
1265+ MessageSequence : dtlsHandshake .Header .MessageSequence ,
1266+ FragmentOffset : uint32 (0 ),
1267+ FragmentLength : uint32 (messageSize ), //nolint:gosec // G115
1268+ }
1269+
1270+ _ , err := headerFragment .MarshalTo (fragmentedHandshake )
1271+ if err != nil {
1272+ return nil , err
1273+ }
1274+
1275+ _ , err = dtlsHandshake .Message .MarshalTo (fragmentedHandshake [handshake .HeaderLength :])
1276+ if err != nil {
1277+ return nil , err
1278+ }
1279+
1280+ fragmentedHandshakes [0 ] = fragmentedHandshake
1281+
1282+ return fragmentedHandshakes , nil
1283+ }
1284+
12181285 content , err := dtlsHandshake .Message .Marshal ()
12191286 if err != nil {
12201287 return nil , err
12211288 }
12221289
1223- fragmentedHandshakes := make ([][]byte , 0 )
1224-
12251290 contentFragments := util .SplitBytes (content , c .maximumTransmissionUnit )
12261291 if len (contentFragments ) == 0 {
1227- contentFragments = [][]byte {
1228- {},
1229- }
1292+ contentFragments = noContentFragments
12301293 }
12311294
12321295 offset := 0
1233- for _ , contentFragment := range contentFragments {
1296+ for i , contentFragment := range contentFragments {
12341297 contentFragmentLen := len (contentFragment )
12351298
1236- headerFragment := & handshake.Header {
1299+ headerFragment := handshake.Header {
12371300 Type : dtlsHandshake .Header .Type ,
12381301 Length : dtlsHandshake .Header .Length ,
12391302 MessageSequence : dtlsHandshake .Header .MessageSequence ,
12401303 FragmentOffset : uint32 (offset ),
12411304 FragmentLength : uint32 (contentFragmentLen ), //nolint:gosec // G115
12421305 }
12431306
1307+ fragmentedHandshake := make ([]byte , handshake .HeaderLength + contentFragmentLen )
1308+
12441309 offset += contentFragmentLen
12451310
1246- fragmentedHandshake , err := headerFragment .Marshal ( )
1311+ _ , err := headerFragment .MarshalTo ( fragmentedHandshake )
12471312 if err != nil {
12481313 return nil , err
12491314 }
12501315
1251- fragmentedHandshake = append (fragmentedHandshake , contentFragment ... )
1252- fragmentedHandshakes = append (fragmentedHandshakes , fragmentedHandshake )
1316+ copy (fragmentedHandshake [handshake .HeaderLength :], contentFragment )
1317+
1318+ fragmentedHandshakes [i ] = fragmentedHandshake
12531319 }
12541320
12551321 return fragmentedHandshakes , nil
@@ -2756,12 +2822,7 @@ func (c *Conn) close(byUser bool) error {
27562822}
27572823
27582824func (c * Conn ) isConnectionClosed () bool {
2759- select {
2760- case <- c .closed .Done ():
2761- return true
2762- default :
2763- return false
2764- }
2825+ return c .closed .Err () != nil
27652826}
27662827
27672828func (c * Conn ) setLocalEpoch (epoch uint16 ) {
0 commit comments