@@ -85,10 +85,12 @@ type comm struct {
8585 messageRecvCount * uint64 // Counter to make sure both sides got a message
8686 clientMutex * sync.Mutex
8787 clientConn net.Conn
88+ clientDone chan error
8889 serverMutex * sync.Mutex
8990 serverConn net.Conn
9091 serverListener net.Listener
9192 serverReady chan struct {}
93+ serverDone chan error
9294 errChan chan error
9395 clientChan chan string
9496 serverChan chan string
@@ -107,6 +109,8 @@ func newComm(ctx context.Context, clientConfig, serverConfig *dtls.Config, serve
107109 clientMutex : & sync.Mutex {},
108110 serverMutex : & sync.Mutex {},
109111 serverReady : make (chan struct {}),
112+ serverDone : make (chan error ),
113+ clientDone : make (chan error ),
110114 errChan : make (chan error ),
111115 clientChan : make (chan string ),
112116 serverChan : make (chan string ),
@@ -172,6 +176,32 @@ func (c *comm) assert(t *testing.T) {
172176 }()
173177}
174178
179+ func (c * comm ) cleanup (t * testing.T ) {
180+ clientDone , serverDone := false , false
181+ for {
182+ select {
183+ case err := <- c .clientDone :
184+ if err != nil {
185+ t .Fatal (err )
186+ }
187+ clientDone = true
188+ if clientDone && serverDone {
189+ return
190+ }
191+ case err := <- c .serverDone :
192+ if err != nil {
193+ t .Fatal (err )
194+ }
195+ serverDone = true
196+ if clientDone && serverDone {
197+ return
198+ }
199+ case <- time .After (testTimeLimit ):
200+ t .Fatalf ("Test timeout waiting for server shutdown" )
201+ }
202+ }
203+ }
204+
175205func clientPion (c * comm ) {
176206 select {
177207 case <- c .serverReady :
@@ -194,6 +224,8 @@ func clientPion(c *comm) {
194224 }
195225
196226 simpleReadWrite (c .errChan , c .clientChan , c .clientConn , c .messageRecvCount )
227+ c .clientDone <- nil
228+ close (c .clientDone )
197229}
198230
199231func serverPion (c * comm ) {
@@ -217,6 +249,8 @@ func serverPion(c *comm) {
217249 }
218250
219251 simpleReadWrite (c .errChan , c .serverChan , c .serverConn , c .messageRecvCount )
252+ c .serverDone <- nil
253+ close (c .serverDone )
220254}
221255
222256/*
@@ -254,6 +288,7 @@ func testPionE2ESimple(t *testing.T, server, client func(*comm)) {
254288 }
255289 serverPort := randomPort (t )
256290 comm := newComm (ctx , cfg , cfg , serverPort , server , client )
291+ defer comm .cleanup (t )
257292 comm .assert (t )
258293 })
259294 }
@@ -287,6 +322,7 @@ func testPionE2ESimplePSK(t *testing.T, server, client func(*comm)) {
287322 }
288323 serverPort := randomPort (t )
289324 comm := newComm (ctx , cfg , cfg , serverPort , server , client )
325+ defer comm .cleanup (t )
290326 comm .assert (t )
291327 })
292328 }
@@ -322,6 +358,7 @@ func testPionE2EMTUs(t *testing.T, server, client func(*comm)) {
322358 }
323359 serverPort := randomPort (t )
324360 comm := newComm (ctx , cfg , cfg , serverPort , server , client )
361+ defer comm .cleanup (t )
325362 comm .assert (t )
326363 })
327364 }
@@ -362,6 +399,7 @@ func testPionE2ESimpleED25519(t *testing.T, server, client func(*comm)) {
362399 }
363400 serverPort := randomPort (t )
364401 comm := newComm (ctx , cfg , cfg , serverPort , server , client )
402+ defer comm .cleanup (t )
365403 comm .assert (t )
366404 })
367405 }
@@ -407,6 +445,7 @@ func testPionE2ESimpleED25519ClientCert(t *testing.T, server, client func(*comm)
407445 }
408446 serverPort := randomPort (t )
409447 comm := newComm (ctx , ccfg , scfg , serverPort , server , client )
448+ defer comm .cleanup (t )
410449 comm .assert (t )
411450}
412451
@@ -450,6 +489,7 @@ func testPionE2ESimpleECDSAClientCert(t *testing.T, server, client func(*comm))
450489 }
451490 serverPort := randomPort (t )
452491 comm := newComm (ctx , ccfg , scfg , serverPort , server , client )
492+ defer comm .cleanup (t )
453493 comm .assert (t )
454494}
455495
@@ -493,6 +533,7 @@ func testPionE2ESimpleRSAClientCert(t *testing.T, server, client func(*comm)) {
493533 }
494534 serverPort := randomPort (t )
495535 comm := newComm (ctx , ccfg , scfg , serverPort , server , client )
536+ defer comm .cleanup (t )
496537 comm .assert (t )
497538}
498539
0 commit comments