Source file src/internal/nettest/conn_test.go

     1  // Copyright 2026 The Go Authors. All rights reserved.
     2  // Use of this source code is governed by a BSD-style
     3  // license that can be found in the LICENSE file.
     4  
     5  package nettest_test
     6  
     7  import (
     8  	"bytes"
     9  	"errors"
    10  	"internal/nettest"
    11  	"io"
    12  	"net"
    13  	"os"
    14  	"testing"
    15  	"testing/synctest"
    16  	"time"
    17  )
    18  
    19  func TestConnReadWrite(t *testing.T) {
    20  	synctest.Test(t, func(t *testing.T) {
    21  		cliConn, srvConn := nettest.NewConnPair()
    22  
    23  		cliData := []byte("hello")
    24  		srvData := []byte("HELLO")
    25  		if n, err := cliConn.Write(cliData); n != len(cliData) || err != nil {
    26  			t.Fatalf("cliConn.Write(%q) = %v, %v; want %v, nil", cliData, n, err, len(cliData))
    27  		}
    28  		if err := cliConn.CloseWrite(); err != nil {
    29  			t.Fatalf("cliConn.CloseWrite() = %v, want nil", err)
    30  		}
    31  		if n, err := srvConn.Write(srvData); n != len(srvData) || err != nil {
    32  			t.Fatalf("srvConn.Write(%q) = %v, %v; want %v, nil", srvData, n, err, len(srvData))
    33  		}
    34  		if err := srvConn.CloseWrite(); err != nil {
    35  			t.Fatalf("cliConn.CloseWrite() = %v, want nil", err)
    36  		}
    37  		gotCli, err := io.ReadAll(cliConn)
    38  		if !bytes.Equal(gotCli, srvData) || err != nil {
    39  			t.Fatalf("io.ReadAll(cliConn) = %q, %v; want %v, nil", gotCli, err, srvData)
    40  		}
    41  		gotSrv, err := io.ReadAll(srvConn)
    42  		if !bytes.Equal(gotSrv, cliData) || err != nil {
    43  			t.Fatalf("io.ReadAll(srvConn) = %q, %v; want %v, nil", gotSrv, err, cliData)
    44  		}
    45  	})
    46  }
    47  
    48  func TestConnZeroBuffer(t *testing.T) {
    49  	// Exercise the case where one side of the conn is blocked writing and the
    50  	// other side is blocked reading.
    51  	// This can only happen when the read buffer has been set to 0, blocking all writes.
    52  	synctest.Test(t, func(t *testing.T) {
    53  		rconn, wconn := nettest.NewConnPair()
    54  		rconn.SetReadBufferSize(0)
    55  		var readDone, writeDone bool
    56  		go func() {
    57  			rconn.Read(make([]byte, 100))
    58  			readDone = true
    59  		}()
    60  		go func() {
    61  			wconn.Write([]byte("a"))
    62  			writeDone = true
    63  		}()
    64  		synctest.Wait()
    65  		if readDone || writeDone {
    66  			t.Errorf("before unblocking: readDone=%v, writeDone=%v; want false", readDone, writeDone)
    67  		}
    68  		wconn.Close()
    69  		synctest.Wait()
    70  		if !readDone || !writeDone {
    71  			t.Errorf("after unblocking: readDone=%v, writeDone=%v; want true", readDone, writeDone)
    72  		}
    73  	})
    74  }
    75  
    76  func TestConnPartialWrite(t *testing.T) {
    77  	// A blocking write to a conn successfully writes some, but not all data.
    78  	synctest.Test(t, func(t *testing.T) {
    79  		const readSize = 5
    80  		data := []byte("0123456789")
    81  		rconn, wconn := nettest.NewConnPair()
    82  		rconn.SetReadBufferSize(1)
    83  		go func() {
    84  			got := make([]byte, readSize)
    85  			if n, err := io.ReadFull(rconn, got); n != readSize || err != nil {
    86  				t.Errorf("io.ReadFull() = %v, %v; want %v, nil", n, err, readSize)
    87  			}
    88  			if want := data[:readSize]; !bytes.Equal(got, want) {
    89  				t.Errorf("read %q, want %q", got, want)
    90  			}
    91  			synctest.Wait()
    92  			rconn.Close()
    93  		}()
    94  		n, err := wconn.Write(data)
    95  		if n != readSize+1 || err == nil {
    96  			t.Errorf("Write() = %v, %v; want %v, error", n, err, readSize+1)
    97  		}
    98  	})
    99  }
   100  
   101  func TestConnReadDeadline(t *testing.T) {
   102  	for _, unblock := range []struct {
   103  		name string
   104  		f    func(*nettest.Conn)
   105  	}{{
   106  		name: "Write",
   107  		f: func(c *nettest.Conn) {
   108  			c.Write([]byte("x"))
   109  		},
   110  	}, {
   111  		name: "Close",
   112  		f: func(c *nettest.Conn) {
   113  			c.Close()
   114  		},
   115  	}, {
   116  		name: "CloseWrite",
   117  		f: func(c *nettest.Conn) {
   118  			c.CloseWrite()
   119  		},
   120  	}} {
   121  		for _, setDeadline := range []struct {
   122  			name string
   123  			f    func(*nettest.Conn, time.Time) error
   124  		}{{
   125  			name: "SetDeadline",
   126  			f:    (*nettest.Conn).SetDeadline,
   127  		}, {
   128  			name: "SetReadDeadline",
   129  			f:    (*nettest.Conn).SetReadDeadline,
   130  		}} {
   131  			t.Run(unblock.name+"/"+setDeadline.name, func(t *testing.T) {
   132  				testDeadline(t, func() deadlineTest {
   133  					rconn, wconn := nettest.NewConnPair()
   134  					return deadlineTest{
   135  						what: "Read()",
   136  						block: func() error {
   137  							_, err := rconn.Read(make([]byte, 1))
   138  							return err
   139  						},
   140  						unblock: func() {
   141  							unblock.f(wconn)
   142  						},
   143  						setDeadline: func(d time.Duration) {
   144  							setDeadline.f(rconn, time.Now().Add(d))
   145  						},
   146  					}
   147  				})
   148  			})
   149  		}
   150  	}
   151  }
   152  
   153  func TestConnWriteDeadline(t *testing.T) {
   154  	for _, unblock := range []struct {
   155  		name string
   156  		f    func(*nettest.Conn)
   157  	}{{
   158  		name: "Read",
   159  		f: func(c *nettest.Conn) {
   160  			io.Copy(io.Discard, c)
   161  		},
   162  	}, {
   163  		name: "Close",
   164  		f: func(c *nettest.Conn) {
   165  			c.Close()
   166  		},
   167  	}, {
   168  		name: "CloseRead",
   169  		f: func(c *nettest.Conn) {
   170  			c.CloseRead()
   171  		},
   172  	}} {
   173  		for _, setDeadline := range []struct {
   174  			name string
   175  			f    func(*nettest.Conn, time.Time) error
   176  		}{{
   177  			name: "SetDeadline",
   178  			f:    (*nettest.Conn).SetDeadline,
   179  		}, {
   180  			name: "SetWriteDeadline",
   181  			f:    (*nettest.Conn).SetWriteDeadline,
   182  		}} {
   183  			t.Run(unblock.name+"/"+setDeadline.name, func(t *testing.T) {
   184  				testDeadline(t, func() deadlineTest {
   185  					rconn, wconn := nettest.NewConnPair()
   186  					rconn.SetReadBufferSize(1)
   187  					return deadlineTest{
   188  						what: "Write()",
   189  						block: func() error {
   190  							_, err := wconn.Write([]byte("1234"))
   191  							wconn.Close()
   192  							return err
   193  						},
   194  						unblock: func() {
   195  							go unblock.f(rconn)
   196  						},
   197  						setDeadline: func(d time.Duration) {
   198  							setDeadline.f(wconn, time.Now().Add(d))
   199  						},
   200  					}
   201  				})
   202  			})
   203  		}
   204  	}
   205  }
   206  
   207  func TestConnCanRead(t *testing.T) {
   208  	synctest.Test(t, func(t *testing.T) {
   209  		rconn, wconn := nettest.NewConnPair()
   210  		if got, want := rconn.CanRead(), false; got != want {
   211  			t.Fatalf("before writing data: rconn.CanRead() = %v, want %v", got, want)
   212  		}
   213  		wconn.Write([]byte("a"))
   214  		if got, want := rconn.CanRead(), true; got != want {
   215  			t.Fatalf("after writing data: rconn.CanRead() = %v, want %v", got, want)
   216  		}
   217  		rconn.Read(make([]byte, 1))
   218  		if got, want := rconn.CanRead(), false; got != want {
   219  			t.Fatalf("after reading data: rconn.CanRead() = %v, want %v", got, want)
   220  		}
   221  		wconn.Close()
   222  		if got, want := rconn.CanRead(), true; got != want {
   223  			t.Fatalf("after closing: rconn.CanRead() = %v, want %v", got, want)
   224  		}
   225  	})
   226  }
   227  
   228  func TestConnIsClosed(t *testing.T) {
   229  	for _, test := range []struct {
   230  		name string
   231  		f    func() *nettest.Conn
   232  		want bool
   233  	}{{
   234  		name: "unclosed",
   235  		f: func() *nettest.Conn {
   236  			conn, _ := nettest.NewConnPair()
   237  			return conn
   238  		},
   239  		want: false,
   240  	}, {
   241  		name: "closed",
   242  		f: func() *nettest.Conn {
   243  			conn, _ := nettest.NewConnPair()
   244  			conn.Close()
   245  			return conn
   246  		},
   247  		want: true,
   248  	}, {
   249  		name: "read-closed",
   250  		f: func() *nettest.Conn {
   251  			conn, _ := nettest.NewConnPair()
   252  			conn.CloseRead()
   253  			return conn
   254  		},
   255  		want: false,
   256  	}, {
   257  		name: "write-closed",
   258  		f: func() *nettest.Conn {
   259  			conn, _ := nettest.NewConnPair()
   260  			conn.CloseWrite()
   261  			return conn
   262  		},
   263  		want: false,
   264  	}, {
   265  		name: "read-write-closed",
   266  		f: func() *nettest.Conn {
   267  			conn, _ := nettest.NewConnPair()
   268  			conn.CloseRead()
   269  			conn.CloseWrite()
   270  			return conn
   271  		},
   272  		want: true,
   273  	}} {
   274  		synctestSubtest(t, test.name, func(t *testing.T) {
   275  			conn := test.f()
   276  			if got, want := conn.IsClosed(), test.want; got != want {
   277  				t.Fatalf("conn.IsClosed() = %v, want %v", got, want)
   278  			}
   279  			if got, want := conn.Peer().IsClosed(), false; got != want {
   280  				t.Fatalf("conn.Peer().IsClosed() = %v, want %v", got, want)
   281  			}
   282  		})
   283  	}
   284  }
   285  
   286  var anyError = errors.New("any") // anyError is passed to isOpError to match any error
   287  
   288  func isOpError(err, want error) bool {
   289  	oe, ok := err.(*net.OpError)
   290  	return ok && (oe.Err == want || want == anyError)
   291  }
   292  
   293  func wantConnReadBytes(t *testing.T, c *nettest.Conn, want []byte) {
   294  	t.Helper()
   295  	got := make([]byte, len(want))
   296  	n, err := io.ReadFull(c, got)
   297  	if n < len(want) || err != nil {
   298  		t.Fatalf("io.ReadFull = %v, %v; want %v, nil", n, err, len(want))
   299  	}
   300  
   301  	if !bytes.Equal(got, want) {
   302  		t.Fatalf("io.ReadFull read %q, want %q", got, want)
   303  	}
   304  }
   305  
   306  func wantConnReadErr(t *testing.T, c *nettest.Conn, want error) {
   307  	t.Helper()
   308  	n, err := c.Read(make([]byte, 1))
   309  	if want == io.EOF {
   310  		if n != 0 || err != io.EOF {
   311  			t.Fatalf("c.Read() = %v, %v; want 0, io.EOF", n, err)
   312  		}
   313  	} else {
   314  		if n != 0 || !isOpError(err, want) {
   315  			t.Fatalf("c.Read() = %v, %v; want 0, OpError{Err: %q}", n, err, want)
   316  		}
   317  	}
   318  }
   319  
   320  func wantConnReadBlocked(t *testing.T, c *nettest.Conn) {
   321  	done := false
   322  	go func() {
   323  		n, err := c.Read(make([]byte, 1))
   324  		if n != 0 || !errors.Is(err, os.ErrDeadlineExceeded) {
   325  			t.Errorf("c.Read() = %v, %v; want 0, ErrDeadlineExceeded", n, err)
   326  		}
   327  		done = true
   328  	}()
   329  	synctest.Wait()
   330  	if done {
   331  		t.Fatalf("Read unexpectedly returned before setting deadline")
   332  	}
   333  	c.SetReadDeadline(time.Now().Add(-1 * time.Second))
   334  	synctest.Wait()
   335  	c.SetReadDeadline(time.Time{})
   336  	if !done {
   337  		t.Fatalf("Read unexpectedly did not return after setting deadline")
   338  	}
   339  }
   340  
   341  func TestConnSetReadError(t *testing.T) {
   342  	synctest.Test(t, func(t *testing.T) {
   343  		wantErr := errors.New("error")
   344  		rconn, wconn := nettest.NewConnPair()
   345  		rconn.SetReadError(wantErr)
   346  
   347  		// Consume buffer before returning error.
   348  		wconn.Write([]byte("one"))
   349  		wantConnReadBytes(t, rconn, []byte("one"))
   350  		wantConnReadErr(t, rconn, wantErr)
   351  
   352  		// Write more to the buffer, suppressing error until buffer drains again.
   353  		wconn.Write([]byte("two"))
   354  		wantConnReadBytes(t, rconn, []byte("two"))
   355  		wantConnReadErr(t, rconn, wantErr)
   356  
   357  		// Error may be cleared.
   358  		rconn.SetReadError(nil)
   359  		wantConnReadBlocked(t, rconn)
   360  
   361  		// Close overrides read error.
   362  		rconn.SetReadError(wantErr)
   363  		wconn.Write([]byte("three"))
   364  		wconn.Close()
   365  		wantConnReadBytes(t, rconn, []byte("three"))
   366  		wantConnReadErr(t, rconn, io.EOF)
   367  
   368  		// Setting another read error does not override Close.
   369  		rconn.SetReadError(nil)
   370  		wantConnReadErr(t, rconn, io.EOF)
   371  		rconn.SetReadError(wantErr)
   372  		wantConnReadErr(t, rconn, io.EOF)
   373  
   374  		// ErrClosed takes precedence over read error.
   375  		rconn.Close()
   376  		wantConnReadErr(t, rconn, net.ErrClosed)
   377  	})
   378  }
   379  
   380  func wantConnWriteBytes(t *testing.T, c *nettest.Conn, b []byte) {
   381  	t.Helper()
   382  	if n, err := c.Write(b); n != len(b) || err != nil {
   383  		t.Fatalf("c.Write() = %v, %v; want %v, nil", n, err, len(b))
   384  	}
   385  }
   386  
   387  func wantConnWriteErr(t *testing.T, c *nettest.Conn, want error) {
   388  	t.Helper()
   389  	n, err := c.Write(make([]byte, 1))
   390  	if n != 0 || !isOpError(err, want) {
   391  		t.Fatalf("c.Write() = %v, %v; want 0, OpError{Err: %q}", n, err, want)
   392  	}
   393  }
   394  
   395  func TestConnSetWriteError(t *testing.T) {
   396  	synctest.Test(t, func(t *testing.T) {
   397  		wantErr := errors.New("error")
   398  		rconn, wconn := nettest.NewConnPair()
   399  		wconn.SetWriteError(wantErr)
   400  
   401  		// Error blocks writes.
   402  		wantConnWriteErr(t, wconn, wantErr)
   403  		wantConnReadBlocked(t, rconn)
   404  
   405  		// Error may be cleared.
   406  		wconn.SetWriteError(nil)
   407  		wantConnWriteBytes(t, wconn, []byte("one"))
   408  
   409  		// Restoring error does not prevent reading buffered data.
   410  		wconn.SetWriteError(wantErr)
   411  		wantConnWriteErr(t, wconn, wantErr)
   412  		wantConnReadBytes(t, rconn, []byte("one"))
   413  
   414  		// Error does not interfere with closing the conn.
   415  		wconn.Close()
   416  		wantConnReadErr(t, rconn, io.EOF)
   417  	})
   418  }
   419  
   420  func TestConnSetCloseError(t *testing.T) {
   421  	synctest.Test(t, func(t *testing.T) {
   422  		wantErr := errors.New("error")
   423  		rconn, wconn := nettest.NewConnPair()
   424  
   425  		wconn.SetCloseError(wantErr)
   426  		if _, err := wconn.Write([]byte("one")); err != nil {
   427  			t.Fatalf("wconn.Write = %v, want success", err)
   428  		}
   429  		if err := wconn.Close(); !isOpError(err, wantErr) {
   430  			t.Fatalf("wconn.Close = %v, want OpError{Err: %v}", err, wantErr)
   431  		}
   432  		if err := wconn.Close(); !isOpError(err, net.ErrClosed) {
   433  			t.Fatalf("wconn.Close = %v, want OpError{Err: net.ErrClosed}", err)
   434  		}
   435  		wantConnReadBytes(t, rconn, []byte("one"))
   436  		wantConnReadErr(t, rconn, io.EOF)
   437  	})
   438  }
   439  
   440  func TestConnCloseReadWriteError(t *testing.T) {
   441  	synctest.Test(t, func(t *testing.T) {
   442  		conn, _ := nettest.NewConnPair()
   443  		conn.SetCloseError(errors.New("error"))
   444  		if err := conn.CloseRead(); err != nil {
   445  			t.Fatalf("conn.CloseRead = %v, want nil", err)
   446  		}
   447  		if err := conn.CloseWrite(); err != nil {
   448  			t.Fatalf("conn.CloseRead = %v, want nil", err)
   449  		}
   450  	})
   451  }
   452  

View as plain text