@@ -389,7 +389,7 @@ func TestHandshakeWithAlert(t *testing.T) {
389389 clientErr <- err
390390 }()
391391
392- _ , errServer := testServer (ctx , dtlsnet .PacketConnFromConn (cb ), ca .RemoteAddr (), testCase .configServer , true )
392+ _ , errServer := testServer (ctx , dtlsnet .PacketConnFromConn (cb ), cb .RemoteAddr (), testCase .configServer , true )
393393 if ! errors .Is (errServer , testCase .errServer ) {
394394 t .Fatalf ("Server error exp(%v) failed(%v)" , testCase .errServer , errServer )
395395 }
@@ -402,6 +402,71 @@ func TestHandshakeWithAlert(t *testing.T) {
402402 }
403403}
404404
405+ func TestHandshakeWithInvalidRecord (t * testing.T ) {
406+ // Limit runtime in case of deadlocks
407+ lim := test .TimeOut (time .Second * 20 )
408+ defer lim .Stop ()
409+
410+ // Check for leaking routines
411+ report := test .CheckRoutines (t )
412+ defer report ()
413+
414+ ctx , cancel := context .WithTimeout (context .Background (), 10 * time .Second )
415+ defer cancel ()
416+
417+ type result struct {
418+ c * Conn
419+ err error
420+ }
421+ clientErr := make (chan result , 1 )
422+ ca , cb := dpipe .Pipe ()
423+ caWithInvalidRecord := & connWithCallback {Conn : ca }
424+
425+ var msgSeq atomic.Int32
426+ // Send invalid record after first message
427+ caWithInvalidRecord .onWrite = func (b []byte ) {
428+ if msgSeq .Add (1 ) == 2 {
429+ if _ , err := ca .Write ([]byte {0x01 , 0x02 }); err != nil {
430+ t .Fatal (err )
431+ }
432+ }
433+ }
434+ go func () {
435+ client , err := testClient (ctx , dtlsnet .PacketConnFromConn (caWithInvalidRecord ), caWithInvalidRecord .RemoteAddr (), & Config {
436+ CipherSuites : []CipherSuiteID {TLS_ECDHE_ECDSA_WITH_AES_128_GCM_SHA256 },
437+ }, true )
438+ clientErr <- result {client , err }
439+ }()
440+
441+ server , errServer := testServer (ctx , dtlsnet .PacketConnFromConn (cb ), cb .RemoteAddr (), & Config {
442+ CipherSuites : []CipherSuiteID {TLS_ECDHE_ECDSA_WITH_AES_128_GCM_SHA256 },
443+ }, true )
444+
445+ errClient := <- clientErr
446+
447+ defer func () {
448+ if server != nil {
449+ if err := server .Close (); err != nil {
450+ t .Fatal (err )
451+ }
452+ }
453+
454+ if errClient .c != nil {
455+ if err := errClient .c .Close (); err != nil {
456+ t .Fatal (err )
457+ }
458+ }
459+ }()
460+
461+ if errServer != nil {
462+ t .Fatalf ("Server failed(%v)" , errServer )
463+ }
464+
465+ if errClient .err != nil {
466+ t .Fatalf ("Client failed(%v)" , errClient .err )
467+ }
468+ }
469+
405470func TestExportKeyingMaterial (t * testing.T ) {
406471 // Check for leaking routines
407472 report := test .CheckRoutines (t )
@@ -3096,3 +3161,15 @@ func TestSkipHelloVerify(t *testing.T) {
30963161 t .Error (err )
30973162 }
30983163}
3164+
3165+ type connWithCallback struct {
3166+ net.Conn
3167+ onWrite func ([]byte )
3168+ }
3169+
3170+ func (c * connWithCallback ) Write (b []byte ) (int , error ) {
3171+ if c .onWrite != nil {
3172+ c .onWrite (b )
3173+ }
3174+ return c .Conn .Write (b )
3175+ }
0 commit comments