diff --git a/conncheck.go b/conncheck.go index b56525e..5a79e76 100644 --- a/conncheck.go +++ b/conncheck.go @@ -2,9 +2,9 @@ package conncheck import ( "crypto/tls" + "errors" "net" "syscall" - "time" ) type Status int @@ -17,18 +17,21 @@ const ( func Do(conn net.Conn) Status { if tlsConn, ok := conn.(*tls.Conn); ok { + // We will only peek from the underlying connection, so we don't corrupt the TLS session. conn = tlsConn.NetConn() } sc, ok := conn.(syscall.Conn) if !ok { + // This happens on WASM return StatusUnknown } - _ = conn.SetReadDeadline(time.Time{}) rawConn, err := sc.SyscallConn() - if err != nil { + if errors.Is(err, syscall.EINVAL) || errors.Is(err, net.ErrClosed) { + return StatusNotOpen + } return StatusUnknown } diff --git a/conncheck_test.go b/conncheck_test.go index e43f133..750c129 100644 --- a/conncheck_test.go +++ b/conncheck_test.go @@ -24,43 +24,63 @@ func TestDo(t *testing.T) { tlsCert := randomTLSCertificate(t) testSocketConn(t, tlsCert) }) + t.Run("clientClose", testClientClose) + t.Run("UDP", testUDP) } -func testSocketConn(t *testing.T, tlsCert *tls.Certificate) { +func testClientClose(t *testing.T) { + accepted, ln := createServer(t, nil) + defer func() { + assert.NoError(t, ln.Close()) + }() - accepted := make(chan net.Conn) + conn, err := net.Dial("tcp", ln.Addr().String()) + require.NoError(t, err) - ln, err := net.Listen("tcp", "localhost:0") + defer func() { + _ = conn.Close() + }() + sconn := <-accepted + defer func() { + _ = sconn.Close() + }() + + require.Equal(t, conncheck.StatusOpen, conncheck.Do(conn)) + + require.NoError(t, conn.Close()) + require.Equal(t, conncheck.StatusNotOpen, conncheck.Do(conn)) +} + +func testUDP(t *testing.T) { + conn, err := net.Dial("udp", "localhost:12345") require.NoError(t, err) defer func() { - assert.NoError(t, ln.Close()) + _ = conn.Close() }() + require.Equal(t, conncheck.StatusOpen, conncheck.Do(conn)) + require.NoError(t, conn.Close()) + require.Equal(t, conncheck.StatusNotOpen, conncheck.Do(conn)) +} - if tlsCert != nil { - ln = tls.NewListener(ln, &tls.Config{Certificates: []tls.Certificate{*tlsCert}}) - } - go func() { - for { - cn, err := ln.Accept() - if err != nil { - // This is usually caused by Listener being - // closed, not really an error. - t.Log("accept goroutine completed") - return - } - accepted <- cn - } +func testSocketConn(t *testing.T, tlsCert *tls.Certificate) { + + accepted, ln := createServer(t, tlsCert) + defer func() { + assert.NoError(t, ln.Close()) }() conn, err := net.Dial("tcp", ln.Addr().String()) - if err != nil { - t.Fatal(err) - } + require.NoError(t, err) + if tlsCert != nil { conn = tls.Client(conn, &tls.Config{ InsecureSkipVerify: true, }) } + defer func() { + // connection may be closed at this point + _ = conn.Close() + }() require.Equal(t, conncheck.StatusOpen, conncheck.Do(conn)) @@ -96,6 +116,30 @@ func testSocketConn(t *testing.T, tlsCert *tls.Certificate) { require.Equal(t, conncheck.StatusNotOpen, conncheck.Do(conn)) } +func createServer(t *testing.T, tlsCert *tls.Certificate) (chan net.Conn, net.Listener) { + accepted := make(chan net.Conn) + + ln, err := net.Listen("tcp", "localhost:0") + require.NoError(t, err) + + if tlsCert != nil { + ln = tls.NewListener(ln, &tls.Config{Certificates: []tls.Certificate{*tlsCert}}) + } + go func() { + for { + cn, err := ln.Accept() + if err != nil { + // This is usually caused by Listener being + // closed, not really an error. + t.Log("accept goroutine completed") + return + } + accepted <- cn + } + }() + return accepted, ln +} + func randomTLSCertificate(t *testing.T) *tls.Certificate { privateKey, err := rsa.GenerateKey(rand.Reader, 2048) require.NoError(t, err)