1
2
3
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
50
51
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
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")
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
348 wconn.Write([]byte("one"))
349 wantConnReadBytes(t, rconn, []byte("one"))
350 wantConnReadErr(t, rconn, wantErr)
351
352
353 wconn.Write([]byte("two"))
354 wantConnReadBytes(t, rconn, []byte("two"))
355 wantConnReadErr(t, rconn, wantErr)
356
357
358 rconn.SetReadError(nil)
359 wantConnReadBlocked(t, rconn)
360
361
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
369 rconn.SetReadError(nil)
370 wantConnReadErr(t, rconn, io.EOF)
371 rconn.SetReadError(wantErr)
372 wantConnReadErr(t, rconn, io.EOF)
373
374
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
402 wantConnWriteErr(t, wconn, wantErr)
403 wantConnReadBlocked(t, rconn)
404
405
406 wconn.SetWriteError(nil)
407 wantConnWriteBytes(t, wconn, []byte("one"))
408
409
410 wconn.SetWriteError(wantErr)
411 wantConnWriteErr(t, wconn, wantErr)
412 wantConnReadBytes(t, rconn, []byte("one"))
413
414
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