Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 6 additions & 3 deletions conncheck.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,9 +2,9 @@ package conncheck

import (
"crypto/tls"
"errors"
"net"
"syscall"
"time"
)

type Status int
Expand All @@ -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
}

Expand Down
86 changes: 65 additions & 21 deletions conncheck_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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))

Expand Down Expand Up @@ -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)
Expand Down
Loading