Skip to content

Commit efd6737

Browse files
committed
Add test for PSK and Identity
* Assert that ServerKeyExchange is only sent with PSKIdentityHint is set on the server side. * Assert that a empty PSKIdentityHint can be used for clients. Resolves #389
1 parent cb62aac commit efd6737

4 files changed

Lines changed: 115 additions & 10 deletions

File tree

conn.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -969,7 +969,7 @@ func (c *Conn) isHandshakeCompletedSuccessfully() bool {
969969
return boolean.bool
970970
}
971971

972-
func (c *Conn) handshake(ctx context.Context, cfg *handshakeConfig, initialFlight flightVal, initialState handshakeState) error { //nolint:gocognit
972+
func (c *Conn) handshake(ctx context.Context, cfg *handshakeConfig, initialFlight flightVal, initialState handshakeState) error { //nolint:gocognit,contextcheck
973973
c.fsm = newHandshakeFSM(&c.state, c.handshakeCache, cfg, initialFlight)
974974

975975
done := make(chan struct{})

conn_test.go

Lines changed: 112 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -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+
716822
func TestClientTimeout(t *testing.T) {
717823
// Limit runtime in case of deadlocks
718824
lim := test.TimeOut(time.Second * 20)

examples/dial/psk/main.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -28,7 +28,7 @@ func main() {
2828
fmt.Printf("Server's hint: %s \n", hint)
2929
return []byte{0xAB, 0xC1, 0x23}, nil
3030
},
31-
PSKIdentityHint: []byte("Pion DTLS Client"),
31+
PSKIdentityHint: []byte{},
3232
CipherSuites: []dtls.CipherSuiteID{dtls.TLS_PSK_WITH_AES_128_CCM_8},
3333
ExtendedMasterSecret: dtls.RequireExtendedMasterSecret,
3434
}

examples/util/util.go

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,6 @@ import (
1111
"errors"
1212
"fmt"
1313
"io"
14-
"io/ioutil"
1514
"net"
1615
"os"
1716
"path/filepath"
@@ -70,7 +69,7 @@ func LoadKeyAndCertificate(keyPath string, certificatePath string) (tls.Certific
7069

7170
// LoadCertificate Load/read certificate(s) from file
7271
func LoadCertificate(path string) (*tls.Certificate, error) {
73-
rawData, err := ioutil.ReadFile(filepath.Clean(path))
72+
rawData, err := os.ReadFile(filepath.Clean(path))
7473
if err != nil {
7574
return nil, err
7675
}

0 commit comments

Comments
 (0)