@@ -539,6 +539,7 @@ func TestPSK(t *testing.T) {
539539
540540 for _ , test := range []struct {
541541 Name string
542+ ClientIdentity []byte
542543 ServerIdentity []byte
543544 CipherSuites []CipherSuiteID
544545 ClientVerifyConnection func (* State ) error
@@ -550,11 +551,13 @@ func TestPSK(t *testing.T) {
550551 {
551552 Name : "Server identity specified" ,
552553 ServerIdentity : []byte ("Test Identity" ),
554+ ClientIdentity : []byte ("Client Identity" ),
553555 CipherSuites : []CipherSuiteID {TLS_PSK_WITH_AES_128_CCM_8 },
554556 },
555557 {
556558 Name : "Server identity specified - Server verify connection fails" ,
557559 ServerIdentity : []byte ("Test Identity" ),
560+ ClientIdentity : []byte ("Client Identity" ),
558561 CipherSuites : []CipherSuiteID {TLS_PSK_WITH_AES_128_CCM_8 },
559562 ServerVerifyConnection : func (* State ) error {
560563 return errExample
@@ -566,6 +569,7 @@ func TestPSK(t *testing.T) {
566569 {
567570 Name : "Server identity specified - Client verify connection fails" ,
568571 ServerIdentity : []byte ("Test Identity" ),
572+ ClientIdentity : []byte ("Client Identity" ),
569573 CipherSuites : []CipherSuiteID {TLS_PSK_WITH_AES_128_CCM_8 },
570574 ClientVerifyConnection : func (* State ) error {
571575 return errExample
@@ -577,25 +581,33 @@ func TestPSK(t *testing.T) {
577581 {
578582 Name : "Server identity nil" ,
579583 ServerIdentity : nil ,
584+ ClientIdentity : []byte ("Client Identity" ),
580585 CipherSuites : []CipherSuiteID {TLS_PSK_WITH_AES_128_CCM_8 },
581586 },
582587 {
583588 Name : "TLS_PSK_WITH_AES_128_CBC_SHA256" ,
584589 ServerIdentity : nil ,
590+ ClientIdentity : []byte ("Client Identity" ),
585591 CipherSuites : []CipherSuiteID {TLS_PSK_WITH_AES_128_CBC_SHA256 },
586592 },
587593 {
588594 Name : "TLS_ECDHE_PSK_WITH_AES_128_CBC_SHA256" ,
589595 ServerIdentity : nil ,
596+ ClientIdentity : []byte ("Client Identity" ),
590597 CipherSuites : []CipherSuiteID {TLS_ECDHE_PSK_WITH_AES_128_CBC_SHA256 },
591598 },
599+ {
600+ Name : "Client identity empty" ,
601+ ServerIdentity : nil ,
602+ ClientIdentity : []byte {},
603+ CipherSuites : []CipherSuiteID {TLS_PSK_WITH_AES_128_CCM_8 },
604+ },
592605 } {
593606 test := test
594607 t .Run (test .Name , func (t * testing.T ) {
595608 ctx , cancel := context .WithTimeout (context .Background (), 10 * time .Second )
596609 defer cancel ()
597610
598- clientIdentity := []byte ("Client Identity" )
599611 type result struct {
600612 c * Conn
601613 err error
@@ -612,7 +624,7 @@ func TestPSK(t *testing.T) {
612624
613625 return []byte {0xAB , 0xC1 , 0x23 }, nil
614626 },
615- PSKIdentityHint : clientIdentity ,
627+ PSKIdentityHint : test . ClientIdentity ,
616628 CipherSuites : test .CipherSuites ,
617629 VerifyConnection : test .ClientVerifyConnection ,
618630 }
@@ -623,8 +635,9 @@ func TestPSK(t *testing.T) {
623635
624636 config := & Config {
625637 PSK : func (hint []byte ) ([]byte , error ) {
626- if ! bytes .Equal (clientIdentity , hint ) {
627- return nil , fmt .Errorf ("%w: expected(% 02x) actual(% 02x)" , errTestPSKInvalidIdentity , clientIdentity , hint )
638+ fmt .Println (hint )
639+ if ! bytes .Equal (test .ClientIdentity , hint ) {
640+ return nil , fmt .Errorf ("%w: expected(% 02x) actual(% 02x)" , errTestPSKInvalidIdentity , test .ClientIdentity , hint )
628641 }
629642 return []byte {0xAB , 0xC1 , 0x23 }, nil
630643 },
@@ -649,8 +662,8 @@ func TestPSK(t *testing.T) {
649662 }
650663
651664 actualPSKIdentityHint := server .ConnectionState ().IdentityHint
652- if ! bytes .Equal (actualPSKIdentityHint , clientIdentity ) {
653- t .Errorf ("TestPSK: Server ClientPSKIdentity Mismatch '%s': expected(%v) actual(%v)" , test .Name , clientIdentity , actualPSKIdentityHint )
665+ if ! bytes .Equal (actualPSKIdentityHint , test . ClientIdentity ) {
666+ t .Errorf ("TestPSK: Server ClientPSKIdentity Mismatch '%s': expected(%v) actual(%v)" , test .Name , test . ClientIdentity , actualPSKIdentityHint )
654667 }
655668
656669 defer func () {
@@ -713,6 +726,99 @@ func TestPSKHintFail(t *testing.T) {
713726 }
714727}
715728
729+ // Assert that ServerKeyExchange is only sent if Identity is set on server side
730+ func TestPSKServerKeyExchange (t * testing.T ) {
731+ // Limit runtime in case of deadlocks
732+ lim := test .TimeOut (time .Second * 20 )
733+ defer lim .Stop ()
734+
735+ // Check for leaking routines
736+ report := test .CheckRoutines (t )
737+ defer report ()
738+
739+ for _ , test := range []struct {
740+ Name string
741+ SetIdentity bool
742+ }{
743+ {
744+ Name : "Server Identity Set" ,
745+ SetIdentity : true ,
746+ },
747+ {
748+ Name : "Server Not Identity Set" ,
749+ SetIdentity : false ,
750+ },
751+ } {
752+ test := test
753+ t .Run (test .Name , func (t * testing.T ) {
754+ ctx , cancel := context .WithTimeout (context .Background (), 10 * time .Second )
755+ defer cancel ()
756+ gotServerKeyExchange := false
757+
758+ clientErr := make (chan error , 1 )
759+ ca , cb := dpipe .Pipe ()
760+ cbAnalyzer := & connWithCallback {Conn : cb }
761+ cbAnalyzer .onWrite = func (in []byte ) {
762+ messages , err := recordlayer .UnpackDatagram (in )
763+ if err != nil {
764+ t .Fatal (err )
765+ }
766+
767+ for i := range messages {
768+ h := & handshake.Handshake {}
769+ _ = h .Unmarshal (messages [i ][recordlayer .FixedHeaderSize :])
770+
771+ if h .Header .Type == handshake .TypeServerKeyExchange {
772+ gotServerKeyExchange = true
773+ }
774+ }
775+ }
776+
777+ go func () {
778+ conf := & Config {
779+ PSK : func ([]byte ) ([]byte , error ) {
780+ return []byte {0xAB , 0xC1 , 0x23 }, nil
781+ },
782+ PSKIdentityHint : []byte {0xAB , 0xC1 , 0x23 },
783+ CipherSuites : []CipherSuiteID {TLS_PSK_WITH_AES_128_CCM_8 },
784+ }
785+
786+ if client , err := testClient (ctx , dtlsnet .PacketConnFromConn (ca ), ca .RemoteAddr (), conf , false ); err != nil {
787+ clientErr <- err
788+ } else {
789+ clientErr <- client .Close () //nolint
790+ }
791+ }()
792+
793+ config := & Config {
794+ PSK : func ([]byte ) ([]byte , error ) {
795+ return []byte {0xAB , 0xC1 , 0x23 }, nil
796+ },
797+ CipherSuites : []CipherSuiteID {TLS_PSK_WITH_AES_128_CCM_8 },
798+ }
799+ if test .SetIdentity {
800+ config .PSKIdentityHint = []byte {0xAB , 0xC1 , 0x23 }
801+ }
802+
803+ if server , err := testServer (ctx , dtlsnet .PacketConnFromConn (cbAnalyzer ), cbAnalyzer .RemoteAddr (), config , false ); err != nil {
804+ t .Fatalf ("TestPSK: Server error %v" , err )
805+ } else {
806+ if err = server .Close (); err != nil {
807+ t .Fatal (err )
808+ }
809+ }
810+
811+ if err := <- clientErr ; err != nil {
812+ t .Fatalf ("TestPSK: Client error %v" , err )
813+ }
814+
815+ if gotServerKeyExchange != test .SetIdentity {
816+ t .Fatalf ("Mismatch between setting Identity and getting a ServerKeyExchange exp(%t) actual(%t)" , test .SetIdentity , gotServerKeyExchange )
817+ }
818+ })
819+ }
820+ }
821+
716822func TestClientTimeout (t * testing.T ) {
717823 // Limit runtime in case of deadlocks
718824 lim := test .TimeOut (time .Second * 20 )
0 commit comments