Explorar el Código

pkg/transport: fix a data race in TestReadWriteTimeoutDialer

Accessing test.T async will cause data race.

Change to use select to coordinate the access of test.T.
Yicheng Qin hace 11 años
padre
commit
de1a16e0f1
Se han modificado 1 ficheros con 21 adiciones y 7 borrados
  1. 21 7
      pkg/transport/timeout_dialer_test.go

+ 21 - 7
pkg/transport/timeout_dialer_test.go

@@ -42,18 +42,22 @@ func TestReadWriteTimeoutDialer(t *testing.T) {
 
 
 	// fill the socket buffer
 	// fill the socket buffer
 	data := make([]byte, 5*1024*1024)
 	data := make([]byte, 5*1024*1024)
-	timer := time.AfterFunc(d.wtimeoutd*5, func() {
+	done := make(chan struct{})
+	go func() {
+		_, err = conn.Write(data)
+		done <- struct{}{}
+	}()
+
+	select {
+	case <-done:
+	case <-time.After(d.wtimeoutd * 5):
 		t.Fatal("wait timeout")
 		t.Fatal("wait timeout")
-	})
-	defer timer.Stop()
+	}
 
 
-	_, err = conn.Write(data)
 	if operr, ok := err.(*net.OpError); !ok || operr.Op != "write" || !operr.Timeout() {
 	if operr, ok := err.(*net.OpError); !ok || operr.Op != "write" || !operr.Timeout() {
 		t.Errorf("err = %v, want write i/o timeout error", err)
 		t.Errorf("err = %v, want write i/o timeout error", err)
 	}
 	}
 
 
-	timer.Reset(d.rdtimeoutd * 5)
-
 	conn, err = d.Dial("tcp", ln.Addr().String())
 	conn, err = d.Dial("tcp", ln.Addr().String())
 	if err != nil {
 	if err != nil {
 		t.Fatalf("unexpected dial error: %v", err)
 		t.Fatalf("unexpected dial error: %v", err)
@@ -61,7 +65,17 @@ func TestReadWriteTimeoutDialer(t *testing.T) {
 	defer conn.Close()
 	defer conn.Close()
 
 
 	buf := make([]byte, 10)
 	buf := make([]byte, 10)
-	_, err = conn.Read(buf)
+	go func() {
+		_, err = conn.Read(buf)
+		done <- struct{}{}
+	}()
+
+	select {
+	case <-done:
+	case <-time.After(d.rdtimeoutd * 5):
+		t.Fatal("wait timeout")
+	}
+
 	if operr, ok := err.(*net.OpError); !ok || operr.Op != "read" || !operr.Timeout() {
 	if operr, ok := err.(*net.OpError); !ok || operr.Op != "read" || !operr.Timeout() {
 		t.Errorf("err = %v, want write i/o timeout error", err)
 		t.Errorf("err = %v, want write i/o timeout error", err)
 	}
 	}