From fcb4135b143ff20b243de59a2dc162b044a0b074 Mon Sep 17 00:00:00 2001 From: Marin Nozhchev Date: Tue, 19 May 2026 09:12:09 +0300 Subject: [PATCH 1/4] fix: more tests and edge case fixes --- conncheck.go | 15 +++++++-- conncheck_test.go | 83 +++++++++++++++++++++++++++++++++++------------ 2 files changed, 75 insertions(+), 23 deletions(-) diff --git a/conncheck.go b/conncheck.go index b56525e..86b89b5 100644 --- a/conncheck.go +++ b/conncheck.go @@ -2,6 +2,7 @@ package conncheck import ( "crypto/tls" + "errors" "net" "syscall" "time" @@ -17,18 +18,28 @@ 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() + // we need to reset the deadline for the checks to work + err := conn.SetReadDeadline(time.Time{}) + if err != nil { + // only happens if closing + return StatusNotOpen + } + rawConn, err := sc.SyscallConn() if err != nil { + if errors.Is(err, syscall.EINVAL) { + return StatusNotOpen + } return StatusUnknown } diff --git a/conncheck_test.go b/conncheck_test.go index e43f133..befcc4d 100644 --- a/conncheck_test.go +++ b/conncheck_test.go @@ -24,43 +24,60 @@ 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() + }() + <-accepted + + 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 +113,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) From a0b4ca998101075a5829f898ff507ebaa567092f Mon Sep 17 00:00:00 2001 From: Marin Nozhchev Date: Fri, 22 May 2026 19:31:52 +0300 Subject: [PATCH 2/4] Apply suggestions from code review Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> --- conncheck.go | 2 +- conncheck_test.go | 3 ++- 2 files changed, 3 insertions(+), 2 deletions(-) diff --git a/conncheck.go b/conncheck.go index 86b89b5..dbfd59e 100644 --- a/conncheck.go +++ b/conncheck.go @@ -37,7 +37,7 @@ func Do(conn net.Conn) Status { rawConn, err := sc.SyscallConn() if err != nil { - if errors.Is(err, syscall.EINVAL) { + 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 befcc4d..4903274 100644 --- a/conncheck_test.go +++ b/conncheck_test.go @@ -40,7 +40,8 @@ func testClientClose(t *testing.T) { defer func() { _ = conn.Close() }() - <-accepted + sconn := <-accepted + defer sconn.Close() require.Equal(t, conncheck.StatusOpen, conncheck.Do(conn)) From d40d1cf76fb63363de2321d5b319a166763770e8 Mon Sep 17 00:00:00 2001 From: Marin Nozhchev Date: Fri, 22 May 2026 19:34:57 +0300 Subject: [PATCH 3/4] fixup --- conncheck.go | 8 -------- 1 file changed, 8 deletions(-) diff --git a/conncheck.go b/conncheck.go index dbfd59e..5a79e76 100644 --- a/conncheck.go +++ b/conncheck.go @@ -5,7 +5,6 @@ import ( "errors" "net" "syscall" - "time" ) type Status int @@ -28,13 +27,6 @@ func Do(conn net.Conn) Status { return StatusUnknown } - // we need to reset the deadline for the checks to work - err := conn.SetReadDeadline(time.Time{}) - if err != nil { - // only happens if closing - return StatusNotOpen - } - rawConn, err := sc.SyscallConn() if err != nil { if errors.Is(err, syscall.EINVAL) || errors.Is(err, net.ErrClosed) { From 7e1262581b8706c352e4a9ed09eb60d2f6961b8a Mon Sep 17 00:00:00 2001 From: Marin Nozhchev Date: Fri, 22 May 2026 19:38:40 +0300 Subject: [PATCH 4/4] fixup --- conncheck_test.go | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/conncheck_test.go b/conncheck_test.go index 4903274..750c129 100644 --- a/conncheck_test.go +++ b/conncheck_test.go @@ -41,7 +41,9 @@ func testClientClose(t *testing.T) { _ = conn.Close() }() sconn := <-accepted - defer sconn.Close() + defer func() { + _ = sconn.Close() + }() require.Equal(t, conncheck.StatusOpen, conncheck.Do(conn))