diff --git a/transport/client.go b/transport/client.go index cb1d712f..47ad3671 100644 --- a/transport/client.go +++ b/transport/client.go @@ -20,9 +20,7 @@ package getty import ( "crypto/tls" "crypto/x509" - "encoding/pem" "fmt" - "math" "net" "os" "strings" @@ -157,29 +155,41 @@ func (c *client) dialTCP() Session { conn net.Conn ) - for { - if c.IsClosed() { - return nil - } - if c.sslEnabled { - if sslConfig, buildTlsConfErr := c.tlsConfigBuilder.BuildTlsConfig(); buildTlsConfErr == nil && sslConfig != nil { - d := &net.Dialer{Timeout: connectTimeout} - conn, err = tls.DialWithDialer(d, "tcp", c.addr, sslConfig) - } + // #106: this function performs a SINGLE dial attempt and returns nil on + // failure. The bounded retry/back-off is owned by reConnect() (which + // honors maxReconnectAttempts). Previously dialTCP had an unbounded + // for{} loop here that never returned on failure, making + // WithReconnectAttempts a no-op and hanging clients/tests forever when + // the target was unreachable. + if c.IsClosed() { + return nil + } + if c.sslEnabled { + // #101: guard against a TLS config builder that returns (nil, nil); + // previously a nil config with nil err fell through and the nil conn + // was dereferenced below (conn.RemoteAddr) -> panic. + sslConfig, buildTlsConfErr := c.tlsConfigBuilder.BuildTlsConfig() + if buildTlsConfErr != nil { + err = buildTlsConfErr + } else if sslConfig == nil { + err = fmt.Errorf("tlsConfigBuilder returned nil config without error") } else { - conn, err = net.DialTimeout("tcp", c.addr, connectTimeout) - } - if err == nil && gxnet.IsSameAddr(conn.RemoteAddr(), conn.LocalAddr()) { - _ = conn.Close() - err = errSelfConnect + d := &net.Dialer{Timeout: connectTimeout} + conn, err = tls.DialWithDialer(d, "tcp", c.addr, sslConfig) } - if err == nil { - return newTCPSession(conn, c) - } - - log.Infof("net.DialTimeout(addr:%s, timeout:%v) = error:%+v", c.addr, connectTimeout, perrors.WithStack(err)) - <-gxtime.After(connectInterval) + } else { + conn, err = net.DialTimeout("tcp", c.addr, connectTimeout) + } + if err == nil && gxnet.IsSameAddr(conn.RemoteAddr(), conn.LocalAddr()) { + _ = conn.Close() + err = errSelfConnect } + if err == nil { + return newTCPSession(conn, c) + } + + log.Infof("net.DialTimeout(addr:%s, timeout:%v) = error:%+v", c.addr, connectTimeout, perrors.WithStack(err)) + return nil } func (c *client) dialUDP() Session { @@ -193,51 +203,49 @@ func (c *client) dialUDP() Session { buf []byte ) + // #106: single attempt; reConnect() owns bounded retry/back-off. + if c.IsClosed() { + return nil + } bufp = gxbytes.GetBytes(128) defer gxbytes.PutBytes(bufp) buf = *bufp localAddr = &net.UDPAddr{IP: net.IPv4zero, Port: 0} peerAddr, _ = net.ResolveUDPAddr("udp", c.addr) - for { - if c.IsClosed() { - return nil - } - conn, err = net.DialUDP("udp", localAddr, peerAddr) - if err == nil && gxnet.IsSameAddr(conn.RemoteAddr(), conn.LocalAddr()) { - _ = conn.Close() - err = errSelfConnect - } - if err != nil { - log.Warnf("net.DialTimeout(addr:%s, timeout:%v) = error:%+v", c.addr, perrors.WithStack(err)) - <-gxtime.After(connectInterval) - continue - } + conn, err = net.DialUDP("udp", localAddr, peerAddr) + if err == nil && gxnet.IsSameAddr(conn.RemoteAddr(), conn.LocalAddr()) { + _ = conn.Close() + err = errSelfConnect + } + if err != nil { + // #104: the format string referenced a timeout verb but no argument + // was supplied; UDP dial has no timeout, so drop the verb. + log.Warnf("net.DialUDP(addr:%s) = error:%+v", c.addr, perrors.WithStack(err)) + return nil + } - // check connection alive by write/read action - if err := conn.SetWriteDeadline(time.Now().Add(1e9)); err != nil { - log.Warnf("failed to set write deadline: %+v", err) - } - if length, err = conn.Write(connectPingPackage[:]); err != nil { - _ = conn.Close() - log.Warnf("conn.Write(%s) = {length:%d, err:%+v}", string(connectPingPackage), length, perrors.WithStack(err)) - <-gxtime.After(connectInterval) - continue - } - if err := conn.SetReadDeadline(time.Now().Add(1e9)); err != nil { - log.Warnf("failed to set read deadline: %+v", err) - } - length, err = conn.Read(buf) - if netErr, ok := perrors.Cause(err).(net.Error); ok && netErr.Timeout() { - err = nil - } - if err != nil { - log.Infof("conn{%#v}.Read() = {length:%d, err:%+v}", conn, length, perrors.WithStack(err)) - _ = conn.Close() - <-gxtime.After(connectInterval) - continue - } - return newUDPSession(conn, c) + // check connection alive by write/read action + if err := conn.SetWriteDeadline(time.Now().Add(1e9)); err != nil { + log.Warnf("failed to set write deadline: %+v", err) } + if length, err = conn.Write(connectPingPackage[:]); err != nil { + _ = conn.Close() + log.Warnf("conn.Write(%s) = {length:%d, err:%+v}", string(connectPingPackage), length, perrors.WithStack(err)) + return nil + } + if err := conn.SetReadDeadline(time.Now().Add(1e9)); err != nil { + log.Warnf("failed to set read deadline: %+v", err) + } + length, err = conn.Read(buf) + if netErr, ok := perrors.Cause(err).(net.Error); ok && netErr.Timeout() { + err = nil + } + if err != nil { + log.Infof("conn{%#v}.Read() = {length:%d, err:%+v}", conn, length, perrors.WithStack(err)) + _ = conn.Close() + return nil + } + return newUDPSession(conn, c) } func (c *client) dialWS() Session { @@ -248,107 +256,86 @@ func (c *client) dialWS() Session { ss Session ) + // #106: single attempt; reConnect() owns bounded retry/back-off. + if c.IsClosed() { + return nil + } dialer.EnableCompression = true - for { - if c.IsClosed() { - return nil - } - conn, _, err = dialer.Dial(c.addr, nil) - log.Infof("websocket.dialer.Dial(addr:%s) = error:%+v", c.addr, perrors.WithStack(err)) - if err == nil && gxnet.IsSameAddr(conn.RemoteAddr(), conn.LocalAddr()) { - _ = conn.Close() - err = errSelfConnect + conn, _, err = dialer.Dial(c.addr, nil) + if err == nil && gxnet.IsSameAddr(conn.RemoteAddr(), conn.LocalAddr()) { + _ = conn.Close() + err = errSelfConnect + } + if err == nil { + ss = newWSSession(conn, c) + if ss.(*session).maxMsgLen > 0 { + conn.SetReadLimit(int64(ss.(*session).maxMsgLen)) } - if err == nil { - ss = newWSSession(conn, c) - if ss.(*session).maxMsgLen > 0 { - conn.SetReadLimit(int64(ss.(*session).maxMsgLen)) - } - return ss - } + return ss + } - log.Infof("websocket.dialer.Dial(addr:%s) = error:%+v", c.addr, perrors.WithStack(err)) - <-gxtime.After(connectInterval) + log.Infof("websocket.dialer.Dial(addr:%s) = error:%+v", c.addr, perrors.WithStack(err)) + return nil +} + +func (c *client) buildWSSClientTLSConfig() (*tls.Config, error) { + config := &tls.Config{MinVersion: tls.VersionTLS12} + if c.cert == "" { + return config, nil + } + + certPEM, err := os.ReadFile(c.cert) + if err != nil { + return nil, perrors.Wrapf(err, "os.ReadFile(cert:%s)", c.cert) } + certPool := x509.NewCertPool() + if !certPool.AppendCertsFromPEM(certPEM) { + return nil, fmt.Errorf("failed to parse root certificate: %s", c.cert) + } + config.RootCAs = certPool + return config, nil } func (c *client) dialWSS() Session { var ( - err error - root *x509.Certificate - roots []*x509.Certificate - certPool *x509.CertPool - config *tls.Config - dialer websocket.Dialer - conn *websocket.Conn - ss Session + err error + config *tls.Config + dialer websocket.Dialer + conn *websocket.Conn + ss Session ) - dialer.EnableCompression = true - - config = &tls.Config{ - InsecureSkipVerify: true, - } - - if c.cert != "" { - certPEMBlock, err := os.ReadFile(c.cert) - if err != nil { - panic(fmt.Sprintf("os.ReadFile(cert:%s) = error:%+v", c.cert, perrors.WithStack(err))) - } - - var cert tls.Certificate - for { - var certDERBlock *pem.Block - certDERBlock, certPEMBlock = pem.Decode(certPEMBlock) - if certDERBlock == nil { - break - } - if certDERBlock.Type == "CERTIFICATE" { - cert.Certificate = append(cert.Certificate, certDERBlock.Bytes) - } - } - config.Certificates = make([]tls.Certificate, 1) - config.Certificates[0] = cert + // #106: single attempt; reConnect() owns bounded retry/back-off. + if c.IsClosed() { + return nil } + dialer.EnableCompression = true - certPool = x509.NewCertPool() - for _, c := range config.Certificates { - roots, err = x509.ParseCertificates(c.Certificate[len(c.Certificate)-1]) - if err != nil { - panic(fmt.Sprintf("error parsing server's root cert: %+v\n", perrors.WithStack(err))) - } - for _, root = range roots { - certPool.AddCert(root) - } + config, err = c.buildWSSClientTLSConfig() + if err != nil { + log.Errorf("build WSS client TLS config = error:%+v", perrors.WithStack(err)) + return nil } - config.InsecureSkipVerify = true - config.RootCAs = certPool - // dialer.EnableCompression = true dialer.TLSClientConfig = config - for { - if c.IsClosed() { - return nil - } - conn, _, err = dialer.Dial(c.addr, nil) - if err == nil && gxnet.IsSameAddr(conn.RemoteAddr(), conn.LocalAddr()) { - _ = conn.Close() - err = errSelfConnect - } - if err == nil { - ss = newWSSession(conn, c) - if ss.(*session).maxMsgLen > 0 { - conn.SetReadLimit(int64(ss.(*session).maxMsgLen)) - } - ss.SetName(defaultWSSSessionName) - - return ss + conn, _, err = dialer.Dial(c.addr, nil) + if err == nil && gxnet.IsSameAddr(conn.RemoteAddr(), conn.LocalAddr()) { + _ = conn.Close() + err = errSelfConnect + } + if err == nil { + ss = newWSSession(conn, c) + if ss.(*session).maxMsgLen > 0 { + conn.SetReadLimit(int64(ss.(*session).maxMsgLen)) } + ss.SetName(defaultWSSSessionName) - log.Infof("websocket.dialer.Dial(addr:%s) = error:%+v", c.addr, perrors.WithStack(err)) - <-gxtime.After(connectInterval) + return ss } + + log.Infof("websocket.dialer.Dial(addr:%s) = error:%+v", c.addr, perrors.WithStack(err)) + return nil } func (c *client) dial() Session { @@ -430,15 +417,36 @@ func (c *client) RunEventLoop(newSession NewSessionCallback) { c.Lock() c.newSession = newSession c.Unlock() - c.reConnect() + <-c.runReconnect() +} + +// runReconnect starts a reconnect loop only while the client is open. The +// client lock serializes WaitGroup.Add with stop, which prevents Add racing +// with Close's Wait. +func (c *client) runReconnect() <-chan struct{} { + done := make(chan struct{}) + c.Lock() + select { + case <-c.done: + c.Unlock() + close(done) + return done + default: + c.wg.Add(1) + } + c.Unlock() + + go func() { + defer c.wg.Done() + defer close(done) + c.reConnect() + }() + return done } // a for-loop connect to make sure the connection pool is valid func (c *client) reConnect() { - var ( - sessionNum, reconnectAttempts int - maxReconnectInterval int64 - ) + var reconnectAttempts int reconnectInterval := c.reconnectInterval if reconnectInterval == 0 { reconnectInterval = defaultReconnectInterval @@ -448,42 +456,48 @@ func (c *client) reConnect() { maxReconnectAttempts = defaultMaxReconnectAttempts } connPoolSize := c.number - for { + for reconnectAttempts < maxReconnectAttempts { if c.IsClosed() { log.Warnf("client{peer:%s} goroutine exit now.", c.addr) - break + return } - sessionNum = c.sessionNum() - if connPoolSize <= sessionNum || maxReconnectAttempts < reconnectAttempts { - //exit reconnect when the number of connection pools is sufficient or the current reconnection attempts exceeds the max reconnection attempts. - break + if connPoolSize <= c.sessionNum() { + return } c.connect() reconnectAttempts++ - maxReconnectInterval = int64(math.Min(float64(reconnectAttempts), float64(maxBackOffTimes))) * int64(reconnectInterval) - <-gxtime.After(time.Duration(maxReconnectInterval)) + if c.IsClosed() || connPoolSize <= c.sessionNum() || reconnectAttempts >= maxReconnectAttempts { + return + } + + backOffTimes := reconnectAttempts + if maxBackOffTimes < backOffTimes { + backOffTimes = maxBackOffTimes + } + backoff := time.Duration(int64(backOffTimes) * int64(reconnectInterval)) + select { + case <-c.done: + return + case <-gxtime.After(backoff): + } } } func (c *client) stop() { - select { - case <-c.done: - return - default: - c.Do(func() { - close(c.done) - c.Lock() - for s := range c.ssMap { - s.RemoveAttribute(sessionClientKey) - s.RemoveAttribute(ignoreReconnectKey) - s.Close() - } - c.ssMap = nil - - c.Unlock() - }) - } + c.Do(func() { + c.Lock() + close(c.done) + sessions := c.ssMap + c.ssMap = nil + c.Unlock() + + for s := range sessions { + s.RemoveAttribute(sessionClientKey) + s.RemoveAttribute(ignoreReconnectKey) + s.Close() + } + }) } func (c *client) IsClosed() bool { diff --git a/transport/client_test.go b/transport/client_test.go index 8f982307..5b8afd95 100644 --- a/transport/client_test.go +++ b/transport/client_test.go @@ -19,11 +19,15 @@ package getty import ( "bytes" + "crypto/tls" + "errors" "net" "net/http" "os" + "path/filepath" "strconv" "sync" + "sync/atomic" "testing" "time" ) @@ -36,6 +40,150 @@ import ( type PackageHandler struct{} +var errTestTLSConfig = errors.New("test TLS config failure") + +type countingTLSConfigBuilder struct { + calls atomic.Int32 + entered chan struct{} + release <-chan struct{} +} + +func (b *countingTLSConfigBuilder) BuildTlsConfig() (*tls.Config, error) { + b.calls.Add(1) + select { + case b.entered <- struct{}{}: + default: + } + if b.release != nil { + <-b.release + } + return nil, errTestTLSConfig +} + +func newFailingReconnectClient(builder TlsConfigBuilder, interval time.Duration, attempts int) *client { + return newClient(TCP_CLIENT, + WithServerAddress("127.0.0.1:1"), + WithConnectionNumber(1), + WithReconnectInterval(int(interval)), + WithReconnectAttempts(attempts), + WithClientSslEnabled(true), + WithClientTlsConfigBuilder(builder), + ) +} + +func TestReconnectAttemptsAreExactAndSkipFinalBackoff(t *testing.T) { + builder := &countingTLSConfigBuilder{entered: make(chan struct{}, 4)} + clt := newFailingReconnectClient(builder, 100*time.Millisecond, 3) + + started := time.Now() + clt.RunEventLoop(func(Session) error { return nil }) + elapsed := time.Since(started) + + if got := builder.calls.Load(); got != 3 { + t.Fatalf("TLS config build calls = %d, want exactly 3 reconnect attempts", got) + } + // Correct behavior waits after attempts one and two only: 100ms + 200ms. + // A final backoff adds another 300ms. + if elapsed >= 500*time.Millisecond { + t.Fatalf("reconnect loop took %v; it appears to wait after the final attempt", elapsed) + } +} + +func TestReconnectBackoffIsCancelledByClose(t *testing.T) { + builder := &countingTLSConfigBuilder{entered: make(chan struct{}, 4)} + clt := newFailingReconnectClient(builder, 2*time.Second, 3) + + eventLoopDone := make(chan struct{}) + go func() { + clt.RunEventLoop(func(Session) error { return nil }) + close(eventLoopDone) + }() + select { + case <-builder.entered: + case <-time.After(time.Second): + t.Fatal("first reconnect attempt did not start") + } + time.Sleep(20 * time.Millisecond) + + closeDone := make(chan struct{}) + go func() { + clt.Close() + close(closeDone) + }() + + timedOut := false + select { + case <-closeDone: + case <-time.After(200 * time.Millisecond): + timedOut = true + } + if timedOut { + select { + case <-closeDone: + case <-time.After(3 * time.Second): + t.Fatal("Close did not return after the current backoff elapsed") + } + t.Fatal("Close did not cancel the reconnect backoff") + } + select { + case <-eventLoopDone: + case <-time.After(time.Second): + t.Fatal("RunEventLoop did not return after Close") + } +} + +func TestSessionReconnectIsTrackedByClose(t *testing.T) { + releaseReconnect := make(chan struct{}) + builder := &countingTLSConfigBuilder{ + entered: make(chan struct{}, 4), + release: releaseReconnect, + } + clt := newFailingReconnectClient(builder, 2*time.Second, 3) + localConn, peerConn := net.Pipe() + defer func() { + _ = localConn.Close() + _ = peerConn.Close() + }() + ss := newTCPSession(localConn, clt).(*session) + ss.SetAttribute(sessionClientKey, clt) + ss.SetAttribute(ignoreReconnectKey, false) + + sessionStopDone := make(chan struct{}) + go func() { + ss.stop() + close(sessionStopDone) + }() + select { + case <-builder.entered: + case <-time.After(time.Second): + t.Fatal("session-triggered reconnect did not start") + } + + closeDone := make(chan struct{}) + go func() { + clt.Close() + close(closeDone) + }() + select { + case <-closeDone: + close(releaseReconnect) + t.Fatal("Close returned while the session-triggered reconnect was still running") + case <-time.After(100 * time.Millisecond): + } + + close(releaseReconnect) + select { + case <-closeDone: + case <-time.After(time.Second): + t.Fatal("Close did not return after the session-triggered reconnect completed") + } + select { + case <-sessionStopDone: + case <-time.After(time.Second): + t.Fatal("session stop did not return") + } +} + func (h *PackageHandler) Read(ss Session, data []byte) (any, int, error) { return nil, 0, nil } @@ -93,7 +241,11 @@ func newSessionCallback(session Session, handler *MessageHandler) error { func TestTCPClient(t *testing.T) { listenLocalServer := func() (net.Listener, error) { - listener, err := net.Listen("tcp", ":0") + // #106: bind a concrete loopback address instead of ":0" (which + // resolves to the unspecified "[::]" address). Dialing "[::]:port" + // is not a valid connect destination, so every dial failed and the + // client's reconnect loop hung the test forever. + listener, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { return nil, err } @@ -459,6 +611,58 @@ func DownloadFile(filepath string, content []byte) error { return err } +func TestBuildWSSClientTLSConfig(t *testing.T) { + t.Run("system roots", func(t *testing.T) { + config, err := (&client{}).buildWSSClientTLSConfig() + if err != nil { + t.Fatal(err) + } + if config.RootCAs != nil { + t.Fatal("RootCAs must be nil when no custom root certificate is configured") + } + if config.MinVersion != tls.VersionTLS12 { + t.Fatalf("MinVersion = %d, want TLS 1.2 (%d)", config.MinVersion, tls.VersionTLS12) + } + }) + + t.Run("custom root certificate", func(t *testing.T) { + certPath := filepath.Join(t.TempDir(), "root.crt") + if err := os.WriteFile(certPath, WssClientCRT, 0o600); err != nil { + t.Fatal(err) + } + config, err := (&client{ClientOptions: ClientOptions{cert: certPath}}).buildWSSClientTLSConfig() + if err != nil { + t.Fatal(err) + } + if config.RootCAs == nil { + t.Fatal("RootCAs is nil with a configured root certificate") + } + if len(config.Certificates) != 0 { + t.Fatalf("client Certificates contains %d entries, want 0 for a root-only option", len(config.Certificates)) + } + if config.MinVersion != tls.VersionTLS12 { + t.Fatalf("MinVersion = %d, want TLS 1.2 (%d)", config.MinVersion, tls.VersionTLS12) + } + }) + + t.Run("invalid PEM", func(t *testing.T) { + certPath := filepath.Join(t.TempDir(), "invalid.crt") + if err := os.WriteFile(certPath, []byte("not a certificate"), 0o600); err != nil { + t.Fatal(err) + } + if _, err := (&client{ClientOptions: ClientOptions{cert: certPath}}).buildWSSClientTLSConfig(); err == nil { + t.Fatal("invalid PEM returned nil error") + } + }) + + t.Run("missing file", func(t *testing.T) { + certPath := filepath.Join(t.TempDir(), "missing.crt") + if _, err := (&client{ClientOptions: ClientOptions{cert: certPath}}).buildWSSClientTLSConfig(); err == nil { + t.Fatal("missing root certificate returned nil error") + } + }) +} + func TestNewWSSClient(t *testing.T) { var ( err error diff --git a/transport/connection.go b/transport/connection.go index 9f3ff501..4fd89023 100644 --- a/transport/connection.go +++ b/transport/connection.go @@ -127,7 +127,7 @@ func (c *gettyConn) GetActive() time.Time { // removed unused methods send/close -func (c gettyConn) ReadTimeout() time.Duration { +func (c *gettyConn) ReadTimeout() time.Duration { return c.rTimeout.Load() } @@ -150,7 +150,7 @@ func (c *gettyConn) SetReadTimeout(rTimeout time.Duration) { } } -func (c gettyConn) WriteTimeout() time.Duration { +func (c *gettyConn) WriteTimeout() time.Duration { return c.wTimeout.Load() } @@ -227,12 +227,44 @@ func (t *writeFlusher) Write(p []byte) (int, error) { return n, perrors.WithStack(err) } if err := t.flusher.Flush(); err != nil { - return 0, perrors.WithStack(err) + return n, perrors.WithStack(err) } return n, nil } +// for snappy compress. #102: snappy.NewBufferedWriter buffers writes and only +// emits data on Flush, so small packets would sit in the buffer forever if not +// flushed after every Write. This wrapper flushes on every Write, mirroring the +// flate writeFlusher behavior above. +type snappyWriteFlusher struct { + writer *snappy.Writer + lock sync.Mutex +} + +func newSnappyWriteFlusher(w *snappy.Writer) *snappyWriteFlusher { + return &snappyWriteFlusher{writer: w} +} + +func (s *snappyWriteFlusher) Write(p []byte) (int, error) { + s.lock.Lock() + defer s.lock.Unlock() + n, err := s.writer.Write(p) + if err != nil { + return n, perrors.WithStack(err) + } + if err := s.writer.Flush(); err != nil { + return n, perrors.WithStack(err) + } + return n, nil +} + +func (s *snappyWriteFlusher) Close() error { + s.lock.Lock() + defer s.lock.Unlock() + return perrors.WithStack(s.writer.Close()) +} + // SetCompressType set compress type(tcp: zip/snappy, websocket:zip) func (t *gettyTCPConn) SetCompressType(c CompressType) { switch c { @@ -251,7 +283,9 @@ func (t *gettyTCPConn) SetCompressType(c CompressType) { ioReader := io.Reader(t.conn) t.reader = snappy.NewReader(ioReader) ioWriter := io.Writer(t.conn) - t.writer = snappy.NewBufferedWriter(ioWriter) + // #102: wrap the buffered snappy writer so every Write is flushed, + // otherwise small packets never leave the internal buffer. + t.writer = newSnappyWriteFlusher(snappy.NewBufferedWriter(ioWriter)) default: panic(fmt.Sprintf("illegal comparess type %d", c)) @@ -306,8 +340,23 @@ func (t *gettyTCPConn) Send(pkg any) (int, error) { } if buffers, ok := pkg.([][]byte); ok { - netBuf := net.Buffers(buffers) - lg, err = netBuf.WriteTo(t.conn) + // #102: when compression is enabled the [][]byte path must go through + // t.writer (the compress writer), otherwise it writes raw frames + // directly to t.conn and the peer receives a corrupt mix of + // compressed and uncompressed data. + if t.compress == CompressNone { + netBuf := net.Buffers(buffers) + lg, err = netBuf.WriteTo(t.conn) + } else { + for _, b := range buffers { + var n int + n, err = t.writer.Write(b) + if err != nil { + break + } + lg += int64(n) + } + } if err == nil { t.writeBytes.Add((uint32)(lg)) t.writePkgNum.Add((uint32)(len(buffers))) @@ -338,16 +387,21 @@ func (t *gettyTCPConn) CloseConn(waitSec int) { // } if t.conn != nil { - if writer, ok := t.writer.(*snappy.Writer); ok { + // #102: snappy writer is now wrapped in *snappyWriteFlusher. + if writer, ok := t.writer.(*snappyWriteFlusher); ok { if err := writer.Close(); err != nil { log.Errorf("snappy.Writer.Close() = error:%+v", err) } } + // #103: do not hard-assert *tls.Conn; use safe type assertions so a + // non-TLS, non-TCP conn does not panic here. if conn, ok := t.conn.(*net.TCPConn); ok { _ = conn.SetLinger(waitSec) _ = conn.Close() + } else if tlsConn, ok := t.conn.(*tls.Conn); ok { + _ = tlsConn.Close() } else { - _ = t.conn.(*tls.Conn).Close() + _ = t.conn.Close() } t.conn = nil } diff --git a/transport/connection_test.go b/transport/connection_test.go new file mode 100644 index 00000000..ca601f3e --- /dev/null +++ b/transport/connection_test.go @@ -0,0 +1,169 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package getty + +import ( + "compress/flate" + "errors" + "io" + "net" + "sync" + "testing" + "time" +) + +import ( + "github.com/golang/snappy" +) + +type blockingSnappyWriter struct { + entered chan struct{} + release chan struct{} +} + +var errFlushWriter = errors.New("flush writer failure") + +type flushErrorWriter struct{} + +func (flushErrorWriter) Write([]byte) (int, error) { + return 0, errFlushWriter +} + +func (w *blockingSnappyWriter) Write(p []byte) (int, error) { + select { + case w.entered <- struct{}{}: + default: + } + <-w.release + return len(p), nil +} + +type timeoutAccessorNetConn struct{} + +func (*timeoutAccessorNetConn) Read([]byte) (int, error) { return 0, io.EOF } +func (*timeoutAccessorNetConn) Write(p []byte) (int, error) { return len(p), nil } +func (*timeoutAccessorNetConn) Close() error { return nil } +func (*timeoutAccessorNetConn) LocalAddr() net.Addr { return &net.TCPAddr{} } +func (*timeoutAccessorNetConn) RemoteAddr() net.Addr { return &net.TCPAddr{} } +func (*timeoutAccessorNetConn) SetDeadline(time.Time) error { return nil } +func (*timeoutAccessorNetConn) SetReadDeadline(time.Time) error { return nil } +func (*timeoutAccessorNetConn) SetWriteDeadline(time.Time) error { return nil } + +func TestConnectionTimeoutAccessorsDoNotCopyAtomicState(t *testing.T) { + conn := newGettyTCPConn(&timeoutAccessorNetConn{}) + start := make(chan struct{}) + var wg sync.WaitGroup + wg.Add(2) + + go func() { + defer wg.Done() + <-start + for i := 0; i < 10000; i++ { + conn.rLastDeadline.Store(time.Unix(0, int64(i))) + conn.wLastDeadline.Store(time.Unix(0, int64(i))) + } + }() + go func() { + defer wg.Done() + <-start + for i := 0; i < 10000; i++ { + _ = conn.ReadTimeout() + _ = conn.WriteTimeout() + } + }() + + close(start) + wg.Wait() +} + +func TestWriteFlushersReturnConsumedBytesOnFlushError(t *testing.T) { + payload := []byte("payload") + tests := []struct { + name string + writer io.Writer + }{ + { + name: "flate", + writer: func() io.Writer { + writer, err := flate.NewWriter(flushErrorWriter{}, flate.DefaultCompression) + if err != nil { + t.Fatal(err) + } + return &writeFlusher{flusher: writer} + }(), + }, + { + name: "snappy", + writer: newSnappyWriteFlusher(snappy.NewBufferedWriter(flushErrorWriter{})), + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + n, err := test.writer.Write(payload) + if !errors.Is(err, errFlushWriter) { + t.Fatalf("Write error = %v, want %v", err, errFlushWriter) + } + if n != len(payload) { + t.Fatalf("Write returned %d bytes after consuming %d", n, len(payload)) + } + }) + } +} + +func TestSnappyWriteFlusherCloseWaitsForWrite(t *testing.T) { + underlying := &blockingSnappyWriter{ + entered: make(chan struct{}, 1), + release: make(chan struct{}), + } + var releaseOnce sync.Once + release := func() { releaseOnce.Do(func() { close(underlying.release) }) } + defer release() + + writer := newSnappyWriteFlusher(snappy.NewBufferedWriter(underlying)) + writeDone := make(chan error, 1) + go func() { + _, err := writer.Write([]byte("payload")) + writeDone <- err + }() + + select { + case <-underlying.entered: + case <-time.After(time.Second): + t.Fatal("snappy write did not reach the underlying writer") + } + + closeDone := make(chan error, 1) + go func() { + closeDone <- writer.Close() + }() + + select { + case err := <-closeDone: + t.Fatalf("Close returned before the active Write completed: %v", err) + case <-time.After(50 * time.Millisecond): + } + + release() + if err := <-writeDone; err != nil { + t.Fatalf("Write failed: %v", err) + } + if err := <-closeDone; err != nil { + t.Fatalf("Close failed: %v", err) + } +} diff --git a/transport/server.go b/transport/server.go index 625e80e6..f5f60808 100644 --- a/transport/server.go +++ b/transport/server.go @@ -48,6 +48,7 @@ import ( var ( errSelfConnect = perrors.New("connect self!") + errServerClosed = perrors.New("server closed") serverFastFailTimeout = time.Second * 1 serverID uatomic.Int32 @@ -81,7 +82,7 @@ type server struct { // net pktListener net.PacketConn streamListener net.Listener - lock sync.Mutex // for server + lock sync.RWMutex // for server endPointType EndPointType server *http.Server // for ws or wss server sync.Once @@ -160,15 +161,18 @@ func (s *server) stop() { cancel() } s.server = nil + // Snapshot the listeners under s.lock. Keep the published listener + // objects in place so event loops cannot race with a nil assignment. + streamListener := s.streamListener + pktListener := s.pktListener s.lock.Unlock() - if s.streamListener != nil { + // close outside the lock to avoid blocking other lock holders. + if streamListener != nil { // let the server exit asap when got error from RunEventLoop. - _ = s.streamListener.Close() - s.streamListener = nil + _ = streamListener.Close() } - if s.pktListener != nil { - _ = s.pktListener.Close() - s.pktListener = nil + if pktListener != nil { + _ = pktListener.Close() } }) } @@ -195,6 +199,9 @@ func (s *server) listenTCP() error { err error streamListener net.Listener ) + if s.IsClosed() { + return errServerClosed + } if len(s.addr) == 0 || !strings.Contains(s.addr, ":") { streamListener, err = gxnet.ListenOnTCPRandomPort(s.addr) @@ -203,9 +210,17 @@ func (s *server) listenTCP() error { } } else { if s.sslEnabled { - if sslConfig, buildTlsConfErr := s.tlsConfigBuilder.BuildTlsConfig(); buildTlsConfErr == nil && sslConfig != nil { - streamListener, err = tls.Listen("tcp", s.addr, sslConfig) + // #101: guard against a TLS config builder that returns (nil, nil); + // previously a nil config with nil err fell through, leaving + // streamListener nil, and s.streamListener.Addr() panicked below. + sslConfig, buildTlsConfErr := s.tlsConfigBuilder.BuildTlsConfig() + if buildTlsConfErr != nil { + return perrors.Wrapf(buildTlsConfErr, "BuildTlsConfig(addr:%s)", s.addr) + } + if sslConfig == nil { + return fmt.Errorf("BuildTlsConfig returned nil config without error for addr:%s", s.addr) } + streamListener, err = tls.Listen("tcp", s.addr, sslConfig) } else { streamListener, err = net.Listen("tcp", s.addr) } @@ -214,8 +229,22 @@ func (s *server) listenTCP() error { } } + return s.publishStreamListener(streamListener) +} + +func (s *server) publishStreamListener(streamListener net.Listener) error { + addr := streamListener.Addr().String() + s.lock.Lock() + select { + case <-s.done: + s.lock.Unlock() + _ = streamListener.Close() + return errServerClosed + default: + } s.streamListener = streamListener - s.addr = s.streamListener.Addr().String() + s.addr = addr + s.lock.Unlock() return nil } @@ -226,6 +255,9 @@ func (s *server) listenUDP() error { localAddr *net.UDPAddr pktListener *net.UDPConn ) + if s.IsClosed() { + return errServerClosed + } if len(s.addr) == 0 || !strings.Contains(s.addr, ":") { pktListener, err = gxnet.ListenOnUDPRandomPort(s.addr) @@ -243,8 +275,22 @@ func (s *server) listenUDP() error { } } + return s.publishPacketListener(pktListener) +} + +func (s *server) publishPacketListener(pktListener net.PacketConn) error { + addr := pktListener.LocalAddr().String() + s.lock.Lock() + select { + case <-s.done: + s.lock.Unlock() + _ = pktListener.Close() + return errServerClosed + default: + } s.pktListener = pktListener - s.addr = s.pktListener.LocalAddr().String() + s.addr = addr + s.lock.Unlock() return nil } @@ -267,7 +313,7 @@ func (s *server) accept(newSession NewSessionCallback) (Session, error) { return nil, perrors.WithStack(err) } if gxnet.IsSameAddr(conn.RemoteAddr(), conn.LocalAddr()) { - log.Warnf("conn.localAddr{%s} == conn.RemoteAddr", conn.LocalAddr().String(), conn.RemoteAddr().String()) + log.Warnf("conn.localAddr{%s} == conn.RemoteAddr{%s}", conn.LocalAddr().String(), conn.RemoteAddr().String()) return nil, perrors.WithStack(errSelfConnect) } @@ -384,7 +430,7 @@ func (s *wsHandler) serveWSRequest(w http.ResponseWriter, r *http.Request) { } if conn.RemoteAddr().String() == conn.LocalAddr().String() { _ = conn.Close() - log.Warnf("conn.localAddr{%s} == conn.RemoteAddr", conn.LocalAddr().String(), conn.RemoteAddr().String()) + log.Warnf("conn.localAddr{%s} == conn.RemoteAddr{%s}", conn.LocalAddr().String(), conn.RemoteAddr().String()) return } // conn.SetReadLimit(int64(handler.maxMsgLen)) @@ -495,6 +541,9 @@ func (s *server) runWSSEventLoop(newSession NewSessionCallback) { // @newSession: new connection callback func (s *server) RunEventLoop(newSession NewSessionCallback) { if err := s.listen(); err != nil { + if perrors.Cause(err) == errServerClosed { + return + } panic(fmt.Errorf("server.listen() = error:%+v", perrors.WithStack(err))) } @@ -513,10 +562,14 @@ func (s *server) RunEventLoop(newSession NewSessionCallback) { } func (s *server) Listener() net.Listener { + s.lock.RLock() + defer s.lock.RUnlock() return s.streamListener } func (s *server) PacketConn() net.PacketConn { + s.lock.RLock() + defer s.lock.RUnlock() return s.pktListener } diff --git a/transport/server_test.go b/transport/server_test.go index b3fb689c..605981a5 100644 --- a/transport/server_test.go +++ b/transport/server_test.go @@ -177,6 +177,127 @@ func testUDPServer(t *testing.T, address string) { t.Logf("@address:%s, udp server addr: %v", address, addr) } +func TestServerCloseKeepsPublishedListener(t *testing.T) { + t.Run("TCP", func(t *testing.T) { + server := newServer(TCP_SERVER, WithLocalAddress("127.0.0.1:0")) + if err := server.listen(); err != nil { + t.Fatal(err) + } + listener := server.Listener() + if listener == nil { + t.Fatal("listen did not publish the TCP listener") + } + + server.Close() + + if got := server.Listener(); got != listener { + t.Fatalf("Listener() after Close = %v, want the published listener %v", got, listener) + } + }) + + t.Run("UDP", func(t *testing.T) { + server := newServer(UDP_ENDPOINT, WithLocalAddress("127.0.0.1:0")) + if err := server.listen(); err != nil { + t.Fatal(err) + } + listener := server.PacketConn() + if listener == nil { + t.Fatal("listen did not publish the UDP listener") + } + + server.Close() + + if got := server.PacketConn(); got != listener { + t.Fatalf("PacketConn() after Close = %v, want the published listener %v", got, listener) + } + }) +} + +func TestServerClosePreventsLateListenerPublication(t *testing.T) { + t.Run("already closed", func(t *testing.T) { + server := newServer(TCP_SERVER, WithLocalAddress("127.0.0.1:0")) + server.Close() + server.RunEventLoop(func(Session) error { return nil }) + if server.Listener() != nil { + t.Fatal("closed server opened a listener") + } + }) + + t.Run("TCP", func(t *testing.T) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + server := newServer(TCP_SERVER) + server.lock.Lock() + closeDone := make(chan struct{}) + go func() { + server.Close() + close(closeDone) + }() + select { + case <-server.done: + case <-time.After(time.Second): + server.lock.Unlock() + t.Fatal("Close did not start shutdown") + } + + publishDone := make(chan error, 1) + go func() { + publishDone <- server.publishStreamListener(listener) + }() + server.lock.Unlock() + + if err := <-publishDone; !errors.Is(err, errServerClosed) { + t.Fatalf("publishStreamListener returned %v, want %v", err, errServerClosed) + } + <-closeDone + if server.Listener() != nil { + t.Fatal("closed server published a late TCP listener") + } + if _, err := listener.Accept(); err == nil { + t.Fatal("late TCP listener remained open after rejected publication") + } + }) + + t.Run("UDP", func(t *testing.T) { + listener, err := net.ListenPacket("udp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + server := newServer(UDP_ENDPOINT) + server.lock.Lock() + closeDone := make(chan struct{}) + go func() { + server.Close() + close(closeDone) + }() + select { + case <-server.done: + case <-time.After(time.Second): + server.lock.Unlock() + t.Fatal("Close did not start shutdown") + } + + publishDone := make(chan error, 1) + go func() { + publishDone <- server.publishPacketListener(listener) + }() + server.lock.Unlock() + + if err := <-publishDone; !errors.Is(err, errServerClosed) { + t.Fatalf("publishPacketListener returned %v, want %v", err, errServerClosed) + } + <-closeDone + if server.PacketConn() != nil { + t.Fatal("closed server published a late UDP listener") + } + if _, _, err := listener.ReadFrom(make([]byte, 1)); err == nil { + t.Fatal("late UDP listener remained open after rejected publication") + } + }) +} + func TestServer(t *testing.T) { var addr string diff --git a/transport/session.go b/transport/session.go index 9d0d9903..ff1634f0 100644 --- a/transport/session.go +++ b/transport/session.go @@ -124,7 +124,9 @@ type session struct { maxMsgLen int32 // heartbeat - period time.Duration + period time.Duration + heartbeatTimer *gxtime.Timer + lifecycle *sessionLifecycle // done wait time.Duration @@ -136,6 +138,7 @@ type session struct { // goroutines sync grNum uatomic.Int32 + grWG sync.WaitGroup lock sync.RWMutex packetLock sync.RWMutex @@ -144,6 +147,51 @@ type session struct { closeCallbackMutex sync.RWMutex } +type sessionLifecycle struct { + lock sync.Mutex + closed bool + wg sync.WaitGroup +} + +func newSessionLifecycle() *sessionLifecycle { + return &sessionLifecycle{} +} + +func (l *sessionLifecycle) acquire() bool { + l.lock.Lock() + defer l.lock.Unlock() + if l.closed { + return false + } + l.wg.Add(1) + return true +} + +func (l *sessionLifecycle) release() { + l.wg.Done() +} + +func (l *sessionLifecycle) addClosingTask() { + l.lock.Lock() + l.wg.Add(1) + l.lock.Unlock() +} + +func (l *sessionLifecycle) close() { + l.lock.Lock() + l.closed = true + l.lock.Unlock() +} + +func (l *sessionLifecycle) wait() { + l.wg.Wait() +} + +type heartbeatContext struct { + session *session + lifecycle *sessionLifecycle +} + func newSession(endPoint EndPoint, conn Connection) *session { ss := &session{ name: defaultSessionName, @@ -153,7 +201,8 @@ func newSession(endPoint EndPoint, conn Connection) *session { maxMsgLen: maxReadBufLen, - period: period, + period: period, + lifecycle: newSessionLifecycle(), once: &sync.Once{}, done: make(chan struct{}), @@ -193,14 +242,35 @@ func newWSSession(conn *websocket.Conn, endPoint EndPoint) Session { } func (s *session) Reset() { - *s = session{ - name: defaultSessionName, - once: &sync.Once{}, - done: make(chan struct{}), - period: period, - wait: pendingDuration, - attrs: gxcontext.NewValuesContext(context.Background()), - } + lifecycle := s.stopHeartbeat() + s.grWG.Wait() + if lifecycle != nil { + lifecycle.wait() + } + // #105: Reset() previously did `*s = session{...}` without holding s.lock, + // racing with concurrent readers. Reset the fields individually under the + // lock instead; replacing the whole struct would also copy the mutexes + // (flagged by `go vet`). + s.closeCallbackMutex.Lock() + defer s.closeCallbackMutex.Unlock() + s.lock.Lock() + defer s.lock.Unlock() + s.name = defaultSessionName + s.endPoint = nil + s.Connection = nil + s.listener = nil + s.reader = nil + s.writer = nil + s.maxMsgLen = 0 + s.once = &sync.Once{} + s.done = make(chan struct{}) + s.period = period + s.heartbeatTimer = nil + s.lifecycle = newSessionLifecycle() + s.wait = pendingDuration + s.attrs = gxcontext.NewValuesContext(context.Background()) + s.grNum.Store(0) + s.closeCallback = callbacks{} } func (s *session) Conn() net.Conn { @@ -220,6 +290,8 @@ func (s *session) Conn() net.Conn { } func (s *session) EndPoint() EndPoint { + s.lock.RLock() + defer s.lock.RUnlock() return s.endPoint } @@ -257,8 +329,11 @@ func (s *session) Stat() string { // IsClosed check whether the session has been closed. func (s *session) IsClosed() bool { + s.lock.RLock() + done := s.done + s.lock.RUnlock() select { - case <-s.done: + case <-done: return true default: @@ -371,12 +446,24 @@ func (s *session) RemoveAttribute(key any) { } func (s *session) sessionToken() string { - if s.IsClosed() || s.Connection == nil { + s.lock.RLock() + done := s.done + conn := s.Connection + name := s.name + endPoint := s.endPoint + s.lock.RUnlock() + + select { + case <-done: + return "session-closed" + default: + } + if conn == nil || endPoint == nil { return "session-closed" } return fmt.Sprintf("{%s:%s:%d:%s<->%s}", - s.name, s.EndPoint().EndPointType(), s.ID(), s.LocalAddr(), s.RemoteAddr()) + name, endPoint.EndPointType(), conn.ID(), conn.LocalAddr(), conn.RemoteAddr()) } func (s *session) WritePkg(pkg any, timeout time.Duration) (pkgBytesLenth int, successCount int, err error) { @@ -414,12 +501,32 @@ func (s *session) WritePkg(pkg any, timeout time.Duration) (pkgBytesLenth int, s } else { pkg = pkgBytes } - s.packetLock.RLock() - defer s.packetLock.RUnlock() if 0 < timeout { - s.gettyConn().SetWriteTimeout(timeout) + s.packetLock.Lock() + defer s.packetLock.Unlock() + } else { + s.packetLock.RLock() + defer s.packetLock.RUnlock() + } + // #103: read s.Connection under s.lock to guard against concurrent gc() + // which sets s.Connection = nil; a bare s.Connection.Send below could + // otherwise nil-deref. The obtained conn/gc pointers stay valid even if + // gc() later nils the field, because they reference the underlying obj. + s.lock.RLock() + conn := s.Connection + gc := s.gettyConn() + s.lock.RUnlock() + if conn == nil || gc == nil { + return 0, 0, ErrSessionClosed + } + if 0 < timeout { + // #103: save & restore so a per-call timeout does not permanently + // rewrite the connection's write deadline for subsequent writes. + origWriteTimeout := gc.WriteTimeout() + gc.SetWriteTimeout(timeout) + defer gc.SetWriteTimeout(origWriteTimeout) } - successCount, err = s.Connection.Send(pkg) + successCount, err = conn.Send(pkg) if err != nil { log.Warnf("%s, [session.WritePkg] @s.Connection.Write(pkg:%#v) = err:%+v", s.Stat(), pkg, err) return len(pkgBytes), successCount, perrors.WithStack(err) @@ -442,8 +549,16 @@ func (s *session) WriteBytes(pkg []byte) (int, error) { defer s.packetLock.RUnlock() } + // #103: guard s.Connection against concurrent gc() nil-ing it. + s.lock.RLock() + conn := s.Connection + s.lock.RUnlock() + if conn == nil { + return 0, ErrSessionClosed + } + for leftPackageSize > maxPacketLen { - _, err := s.Connection.Send(pkg[writeSize:(writeSize + maxPacketLen)]) + _, err := conn.Send(pkg[writeSize:(writeSize + maxPacketLen)]) if err != nil { return writeSize, perrors.Wrapf(err, "s.Connection.Write(pkg len:%d)", len(pkg)) } @@ -455,7 +570,7 @@ func (s *session) WriteBytes(pkg []byte) (int, error) { return writeSize, nil } - _, err := s.Connection.Send(pkg[writeSize:]) + _, err := conn.Send(pkg[writeSize:]) if err != nil { return writeSize, perrors.Wrapf(err, "s.Connection.Write(pkg len:%d)", len(pkg)) } @@ -473,10 +588,17 @@ func (s *session) WriteBytesArray(pkgs ...[]byte) (int, error) { } // reduce syscall and memcopy for multiple packages - if _, ok := s.Connection.(*gettyTCPConn); ok { + // #103: guard s.Connection against concurrent gc() nil-ing it. + s.lock.RLock() + conn := s.Connection + s.lock.RUnlock() + if conn == nil { + return 0, ErrSessionClosed + } + if _, ok := conn.(*gettyTCPConn); ok { s.packetLock.RLock() defer s.packetLock.RUnlock() - lg, err := s.Connection.Send(pkgs) + lg, err := conn.Send(pkgs) if err != nil { return 0, perrors.Wrapf(err, "s.Connection.Write(pkgs num:%d)", len(pkgs)) } @@ -522,13 +644,27 @@ func (s *session) WriteBytesArray(pkgs ...[]byte) (int, error) { } func heartbeat(_ gxtime.TimerID, _ time.Time, arg any) error { - ss, _ := arg.(*session) - if ss == nil || ss.IsClosed() { + ctx, _ := arg.(*heartbeatContext) + if ctx == nil || ctx.session == nil || ctx.lifecycle == nil { return ErrSessionClosed } + ss := ctx.session f := func() { - wsConn, wsFlag := ss.Connection.(*gettyWSConn) + if !ctx.lifecycle.acquire() { + return + } + defer ctx.lifecycle.release() + + ss.lock.RLock() + conn := ss.Connection + listener := ss.listener + ss.lock.RUnlock() + if conn == nil || listener == nil { + return + } + + wsConn, wsFlag := conn.(*gettyWSConn) if wsFlag { err := wsConn.writePing() if err != nil { @@ -536,11 +672,17 @@ func heartbeat(_ gxtime.TimerID, _ time.Time, arg any) error { } } - ss.listener.OnCron(ss) + listener.OnCron(ss) } // if enable task pool, run @f asynchronously. - if taskPool := ss.EndPoint().GetTaskPool(); taskPool != nil { + ss.lock.RLock() + endPoint := ss.endPoint + ss.lock.RUnlock() + if endPoint == nil { + return ErrSessionClosed + } + if taskPool := endPoint.GetTaskPool(); taskPool != nil { taskPool.AddTaskAlways(f) return nil } @@ -565,28 +707,62 @@ func (s *session) run() { return } - if _, err := defaultTimerWheel.AddTimer(heartbeat, gxtime.TimerLoop, s.period, s); err != nil { + s.lock.Lock() + lifecycle := s.lifecycle + timer, err := defaultTimerWheel.AddTimer( + heartbeat, + gxtime.TimerLoop, + s.period, + &heartbeatContext{session: s, lifecycle: lifecycle}, + ) + if err == nil { + s.heartbeatTimer = timer + } + s.lock.Unlock() + if err != nil { panic(fmt.Sprintf("failed to add session %s to defaultTimerWheel err:%v", s.Stat(), err)) } s.grNum.Add(1) + s.grWG.Add(1) // start read gr - go s.handlePackage() + go func() { + defer s.grWG.Done() + s.handlePackage() + }() } func (s *session) addTask(pkg any) { + s.lock.RLock() + lifecycle := s.lifecycle + endPoint := s.endPoint + s.lock.RUnlock() + f := func() { + if lifecycle == nil || !lifecycle.acquire() { + return + } + defer lifecycle.release() + // If the session is closed, there is no need to perform CPU-intensive operations. if s.IsClosed() { - log.Errorf("[Id:%d, name=%s, endpoint=%s] Session is closed", s.ID(), s.name, s.EndPoint()) + log.Errorf("%s Session is closed", s.sessionToken()) + return + } + s.lock.RLock() + listener := s.listener + s.lock.RUnlock() + if listener == nil { return } - s.listener.OnMessage(s, pkg) + listener.OnMessage(s, pkg) s.IncReadPkgNum() } - if taskPool := s.EndPoint().GetTaskPool(); taskPool != nil { - taskPool.AddTaskAlways(f) - return + if endPoint != nil { + if taskPool := endPoint.GetTaskPool(); taskPool != nil { + taskPool.AddTaskAlways(f) + return + } } f() } @@ -688,11 +864,9 @@ func (s *session) handleTCPPackage() error { err = nil exit = true if bufLen != 0 { - // as https://github.com/apache/dubbo-getty/issues/77#issuecomment-939652203 - // this branch is impossible. Even if it happens, the bufLen will be zero and the error - // is io.EOF when getty continues to read the socket. - exit = false - log.Infof("%s, session.conn read EOF, while the bufLen(%d) is non-zero.", s.sessionToken()) + // Process the bytes returned with EOF below, then exit + // without issuing another read on the closed stream. + log.Infof("%s, session.conn read EOF, while the bufLen(%d) is non-zero.", s.sessionToken(), bufLen) } break } @@ -869,8 +1043,19 @@ func (s *session) stop() { } } close(s.done) + lifecycle := s.stopHeartbeat() - go func(sessionToken string) { + s.closeCallbackMutex.RLock() + closeCallbacks := s.closeCallback + s.closeCallbackMutex.RUnlock() + if lifecycle != nil { + lifecycle.addClosingTask() + } + + go func(sessionToken string, closeCallbacks callbacks, lifecycle *sessionLifecycle) { + if lifecycle != nil { + defer lifecycle.release() + } defer func() { if r := recover(); r != nil { const size = 64 << 10 @@ -882,34 +1067,60 @@ func (s *session) stop() { } }() - s.invokeCloseCallbacks() - }(s.sessionToken()) + closeCallbacks.Invoke() + }(s.sessionToken(), closeCallbacks, lifecycle) clt, cltFound := s.GetAttribute(sessionClientKey).(*client) ignoreReconnect, flagFound := s.GetAttribute(ignoreReconnectKey).(bool) if cltFound && flagFound && !ignoreReconnect { - clt.reConnect() + clt.runReconnect() } }) } } +func (s *session) stopHeartbeat() *sessionLifecycle { + s.lock.Lock() + timer := s.heartbeatTimer + s.heartbeatTimer = nil + lifecycle := s.lifecycle + s.lock.Unlock() + + if timer != nil { + timer.Stop() + } + if lifecycle != nil { + lifecycle.close() + } + return lifecycle +} + func (s *session) gc() { var conn Connection + var wait time.Duration + var lifecycle *sessionLifecycle s.lock.Lock() if s.attrs != nil { s.attrs = nil conn = s.Connection s.Connection = nil + wait = s.wait + lifecycle = s.lifecycle } s.lock.Unlock() + if conn != nil && lifecycle != nil { + lifecycle.addClosingTask() + } - go func() { + go func(conn Connection, wait time.Duration, lifecycle *sessionLifecycle) { + if conn != nil && lifecycle != nil { + defer lifecycle.release() + } if conn != nil { - conn.CloseConn(int(s.wait)) + conn.CloseConn(int(wait)) } - }() + }(conn, wait, lifecycle) } // Close will be invoked by NewSessionCallback(if return error is not nil) diff --git a/transport/session_callback_test.go b/transport/session_callback_test.go index 68228fb2..1bb538fa 100644 --- a/transport/session_callback_test.go +++ b/transport/session_callback_test.go @@ -237,3 +237,54 @@ func TestSessionCallback(t *testing.T) { } }) } + +func TestResetWaitsForCloseCallbacks(t *testing.T) { + s := &session{ + once: &sync.Once{}, + done: make(chan struct{}), + closeCallback: callbacks{}, + } + + callbackStarted := make(chan struct{}) + releaseCallback := make(chan struct{}) + s.AddCloseCallback("test", "blocking", func() { + close(callbackStarted) + <-releaseCallback + }) + + callbackDone := make(chan struct{}) + go func() { + s.invokeCloseCallbacks() + close(callbackDone) + }() + <-callbackStarted + + resetDone := make(chan struct{}) + go func() { + s.Reset() + close(resetDone) + }() + + select { + case <-resetDone: + close(releaseCallback) + <-callbackDone + t.Fatal("Reset returned while a close callback was still running") + case <-time.After(50 * time.Millisecond): + } + + close(releaseCallback) + select { + case <-callbackDone: + case <-time.After(time.Second): + t.Fatal("close callback did not finish") + } + select { + case <-resetDone: + case <-time.After(time.Second): + t.Fatal("Reset did not finish after the close callback completed") + } + if got := s.closeCallback.Len(); got != 0 { + t.Fatalf("callback count after Reset = %d, want 0", got) + } +} diff --git a/transport/session_test.go b/transport/session_test.go index bbbc9a26..c7e5e3f8 100644 --- a/transport/session_test.go +++ b/transport/session_test.go @@ -19,12 +19,17 @@ package getty import ( "errors" + "io" "net" + "sync" "testing" "time" ) -var errTestReadFailure = errors.New("test read failure") +var ( + errTestReadFailure = errors.New("test read failure") + errUnexpectedSecondRead = errors.New("unexpected second read") +) type errorReader struct{} @@ -32,6 +37,291 @@ func (errorReader) Read(Session, []byte) (any, int, error) { return nil, 0, errTestReadFailure } +type timeoutTestWriter struct{} + +func (timeoutTestWriter) Write(Session, any) ([]byte, error) { + return []byte("x"), nil +} + +type timeoutTestCall struct { + observed time.Duration + release chan struct{} +} + +type timeoutTestNetConn struct { + owner *gettyTCPConn + entered chan *timeoutTestCall +} + +func (c *timeoutTestNetConn) Write(p []byte) (int, error) { + call := &timeoutTestCall{ + observed: c.owner.WriteTimeout(), + release: make(chan struct{}), + } + c.entered <- call + <-call.release + return len(p), nil +} + +func (*timeoutTestNetConn) Read([]byte) (int, error) { return 0, io.EOF } +func (*timeoutTestNetConn) Close() error { return nil } +func (*timeoutTestNetConn) LocalAddr() net.Addr { return &net.TCPAddr{} } +func (*timeoutTestNetConn) RemoteAddr() net.Addr { return &net.TCPAddr{} } +func (*timeoutTestNetConn) SetDeadline(time.Time) error { return nil } +func (*timeoutTestNetConn) SetReadDeadline(time.Time) error { return nil } +func (*timeoutTestNetConn) SetWriteDeadline(time.Time) error { return nil } + +type eofDataNetConn struct { + data []byte + reads int +} + +func (c *eofDataNetConn) Read(p []byte) (int, error) { + c.reads++ + if c.reads == 1 { + return copy(p, c.data), io.EOF + } + return 0, errUnexpectedSecondRead +} + +func (*eofDataNetConn) Write(p []byte) (int, error) { return len(p), nil } +func (*eofDataNetConn) Close() error { return nil } +func (*eofDataNetConn) LocalAddr() net.Addr { return &net.TCPAddr{} } +func (*eofDataNetConn) RemoteAddr() net.Addr { return &net.TCPAddr{} } +func (*eofDataNetConn) SetDeadline(time.Time) error { return nil } +func (*eofDataNetConn) SetReadDeadline(time.Time) error { return nil } +func (*eofDataNetConn) SetWriteDeadline(time.Time) error { return nil } + +type wholeFrameReader struct{} + +func (wholeFrameReader) Read(_ Session, data []byte) (any, int, error) { + return string(data), len(data), nil +} + +type recordingEventListener struct { + messages []any +} + +func (*recordingEventListener) OnOpen(Session) error { return nil } +func (*recordingEventListener) OnClose(Session) {} +func (*recordingEventListener) OnError(Session, error) {} +func (*recordingEventListener) OnCron(Session) {} +func (l *recordingEventListener) OnMessage(_ Session, v any) { + l.messages = append(l.messages, v) +} + +type blockingCronEventListener struct { + entered chan struct{} + release chan struct{} + enteredOnce sync.Once +} + +func (*blockingCronEventListener) OnOpen(Session) error { return nil } +func (*blockingCronEventListener) OnClose(Session) {} +func (*blockingCronEventListener) OnError(Session, error) {} +func (l *blockingCronEventListener) OnCron(Session) { + l.enteredOnce.Do(func() { close(l.entered) }) + <-l.release +} +func (*blockingCronEventListener) OnMessage(Session, any) {} + +type resetBarrierNetConn struct { + entered chan struct{} + release chan struct{} + enteredOnce sync.Once + releaseOnce sync.Once +} + +func (c *resetBarrierNetConn) Read([]byte) (int, error) { + c.enteredOnce.Do(func() { close(c.entered) }) + <-c.release + return 0, io.EOF +} + +func (*resetBarrierNetConn) Write(p []byte) (int, error) { return len(p), nil } +func (*resetBarrierNetConn) Close() error { return nil } +func (*resetBarrierNetConn) LocalAddr() net.Addr { return &net.TCPAddr{} } +func (*resetBarrierNetConn) RemoteAddr() net.Addr { return &net.TCPAddr{} } +func (*resetBarrierNetConn) SetDeadline(time.Time) error { return nil } +func (*resetBarrierNetConn) SetReadDeadline(time.Time) error { return nil } +func (*resetBarrierNetConn) SetWriteDeadline(time.Time) error { return nil } + +func (c *resetBarrierNetConn) releaseRead() { + c.releaseOnce.Do(func() { close(c.release) }) +} + +func TestConcurrentWritePkgTimeoutRestoration(t *testing.T) { + netConn := &timeoutTestNetConn{entered: make(chan *timeoutTestCall, 2)} + ss := newTCPSession(netConn, nil).(*session) + ss.writer = timeoutTestWriter{} + conn := ss.Connection.(*gettyTCPConn) + netConn.owner = conn + initialTimeout := conn.WriteTimeout() + + firstDone := make(chan error, 1) + go func() { + _, _, err := ss.WritePkg("first", 3*time.Second) + firstDone <- err + }() + firstCall := <-netConn.entered + + secondDone := make(chan error, 1) + go func() { + _, _, err := ss.WritePkg("second", 5*time.Second) + secondDone <- err + }() + + select { + case secondCall := <-netConn.entered: + close(firstCall.release) + <-firstDone + close(secondCall.release) + <-secondDone + t.Fatal("second timed write entered while the first still owned the shared write timeout") + case <-time.After(50 * time.Millisecond): + } + + if firstCall.observed != 3*time.Second { + t.Fatalf("first write observed timeout %v, want %v", firstCall.observed, 3*time.Second) + } + close(firstCall.release) + if err := <-firstDone; err != nil { + t.Fatalf("first write failed: %v", err) + } + + var secondCall *timeoutTestCall + select { + case secondCall = <-netConn.entered: + case <-time.After(time.Second): + t.Fatal("second timed write did not enter after the first completed") + } + if secondCall.observed != 5*time.Second { + t.Fatalf("second write observed timeout %v, want %v", secondCall.observed, 5*time.Second) + } + close(secondCall.release) + if err := <-secondDone; err != nil { + t.Fatalf("second write failed: %v", err) + } + if got := conn.WriteTimeout(); got != initialTimeout { + t.Fatalf("write timeout after concurrent calls = %v, want %v", got, initialTimeout) + } +} + +func TestResetWaitsForPackageLoop(t *testing.T) { + netConn := &resetBarrierNetConn{ + entered: make(chan struct{}), + release: make(chan struct{}), + } + defer netConn.releaseOnce.Do(func() { close(netConn.release) }) + + ss := newTCPSession(netConn, newServer(TCP_SERVER)).(*session) + ss.SetReader(wholeFrameReader{}) + ss.SetWriter(timeoutTestWriter{}) + ss.SetEventListener(&recordingEventListener{}) + ss.run() + + select { + case <-netConn.entered: + case <-time.After(time.Second): + t.Fatal("package loop did not enter Read") + } + + resetDone := make(chan struct{}) + go func() { + ss.Reset() + close(resetDone) + }() + + select { + case <-resetDone: + t.Fatal("Reset returned while the package loop was still running") + case <-time.After(50 * time.Millisecond): + } + + ss.Close() + netConn.releaseRead() + select { + case <-resetDone: + case <-time.After(time.Second): + t.Fatal("Reset did not return after Close released the package loop") + } + if ss.Connection != nil { + t.Fatal("Reset did not clear the session connection") + } + if got := ss.grNum.Load(); got != 0 { + t.Fatalf("goroutine count after Reset = %d, want 0", got) + } +} + +func TestResetWaitsForActiveHeartbeat(t *testing.T) { + listener := &blockingCronEventListener{ + entered: make(chan struct{}), + release: make(chan struct{}), + } + ss := newTCPSession(&eofDataNetConn{}, newServer(TCP_SERVER)).(*session) + ss.SetEventListener(listener) + lifecycle := ss.lifecycle + ctx := &heartbeatContext{session: ss, lifecycle: lifecycle} + + heartbeatDone := make(chan error, 1) + go func() { + heartbeatDone <- heartbeat(0, time.Time{}, ctx) + }() + select { + case <-listener.entered: + case <-time.After(time.Second): + t.Fatal("heartbeat did not enter OnCron") + } + + resetDone := make(chan struct{}) + go func() { + ss.Reset() + close(resetDone) + }() + select { + case <-resetDone: + t.Fatal("Reset returned while the heartbeat callback was still running") + case <-time.After(50 * time.Millisecond): + } + + close(listener.release) + select { + case err := <-heartbeatDone: + if err != nil { + t.Fatalf("heartbeat returned error: %v", err) + } + case <-time.After(time.Second): + t.Fatal("heartbeat did not return") + } + select { + case <-resetDone: + case <-time.After(time.Second): + t.Fatal("Reset did not return after the heartbeat callback completed") + } + + if err := heartbeat(0, time.Time{}, ctx); !errors.Is(err, ErrSessionClosed) { + t.Fatalf("stale heartbeat returned %v, want %v", err, ErrSessionClosed) + } +} + +func TestHandleTCPPackageProcessesDataAndStopsOnEOF(t *testing.T) { + netConn := &eofDataNetConn{data: []byte("final frame")} + ss := newTCPSession(netConn, newServer(TCP_SERVER)).(*session) + listener := &recordingEventListener{} + ss.SetReader(wholeFrameReader{}) + ss.SetEventListener(listener) + + if err := ss.handleTCPPackage(); err != nil { + t.Fatalf("handleTCPPackage returned error: %v", err) + } + if netConn.reads != 1 { + t.Fatalf("underlying Read calls = %d, want 1", netConn.reads) + } + if len(listener.messages) != 1 || listener.messages[0] != "final frame" { + t.Fatalf("delivered messages = %#v, want [\"final frame\"]", listener.messages) + } +} + func TestHandlePackageWithNilListenerDoesNotPanicOnError(t *testing.T) { listener, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { diff --git a/transport/tls.go b/transport/tls.go index 86323eb3..6c37473d 100644 --- a/transport/tls.go +++ b/transport/tls.go @@ -74,7 +74,10 @@ func (s *ServerTlsConfigBuilder) BuildTlsConfig() (*tls.Config, error) { certPool = x509.NewCertPool() if ok := certPool.AppendCertsFromPEM(certPem); !ok { log.Error("failed to parse root certificate file") - return nil, err + // #101: err is nil here (from a successful os.ReadFile above), + // returning (nil, nil) causes a nil-pointer panic at the caller + // (listenTCP). Return an explicit error instead. + return nil, fmt.Errorf("failed to parse root certificate file: %s", s.ServerTrustCertCollectionPath) } config.ClientCAs = certPool config.ClientAuth = tls.RequireAnyClientCert @@ -107,11 +110,17 @@ func (c *ClientTlsConfigBuilder) BuildTlsConfig() (*tls.Config, error) { ok := clientCertPool.AppendCertsFromPEM(certBytes) if !ok { log.Error("failed to parse root certificate") - return nil, err + // #101: err is nil here (from a successful os.ReadFile above), + // returning (nil, nil) causes a nil-pointer panic at the caller + // (dialTCP/dialWSS). Return an explicit error instead. + return nil, fmt.Errorf("failed to parse root certificate: %s", c.ClientTrustCertCollectionPath) } return &tls.Config{ - RootCAs: clientCertPool, - Certificates: []tls.Certificate{cert}, - InsecureSkipVerify: true, + MinVersion: tls.VersionTLS12, + RootCAs: clientCertPool, + Certificates: []tls.Certificate{cert}, + // #100: do NOT set InsecureSkipVerify=true here; it disables + // certificate verification entirely and makes the RootCAs configured + // above useless, exposing the client to MITM attacks. }, nil } diff --git a/transport/tls_test.go b/transport/tls_test.go new file mode 100644 index 00000000..29fb4d68 --- /dev/null +++ b/transport/tls_test.go @@ -0,0 +1,129 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package getty + +import ( + "bytes" + "crypto/tls" + "crypto/x509" + "encoding/pem" + "os" + "path/filepath" + "testing" +) + +var tlsTestRootCertificate = []byte(`-----BEGIN CERTIFICATE----- +MIIBiDCCAS+gAwIBAgIUMaJuA5AGTTBvqSWb4fhJCC7UY4wwCgYIKoZIzj0EAwIw +GjEYMBYGA1UEAwwPZ2V0dHktdGVzdC1yb290MB4XDTI2MDczMTIzMTYyN1oXDTM2 +MDcyODIzMTYyN1owGjEYMBYGA1UEAwwPZ2V0dHktdGVzdC1yb290MFkwEwYHKoZI +zj0CAQYIKoZIzj0DAQcDQgAEWZNS+42M+wb2AmNunl7ccsdoaRYanWn1kgt5Rj7X +50hqE1aA8Wdl7dbbDmCwSrwLRNus1ebi2571N0XJNXn536NTMFEwHQYDVR0OBBYE +FJYJbIsdqMVkz65eVtuLmz41l4IjMB8GA1UdIwQYMBaAFJYJbIsdqMVkz65eVtuL +mz41l4IjMA8GA1UdEwEB/wQFMAMBAf8wCgYIKoZIzj0EAwIDRwAwRAIgX6EFP2GN +UF0MEbozG6tzqvrF1R8NUNEUEF4ThXnQMpMCIB/191gSjtjhiuDKu/pT5cCXe9ka +Wf17jc2sFoJ9DUsb +-----END CERTIFICATE-----`) + +func TestClientTLSConfigBuilderMinimumVersion(t *testing.T) { + tempDir := t.TempDir() + certPath := filepath.Join(tempDir, "client.crt") + keyPath := filepath.Join(tempDir, "client.key") + caPath := filepath.Join(tempDir, "ca.crt") + for path, data := range map[string][]byte{ + certPath: WssServerCRT, + keyPath: WssServerKEY, + caPath: tlsTestRootCertificate, + } { + if err := os.WriteFile(path, data, 0o600); err != nil { + t.Fatal(err) + } + } + + config, err := (&ClientTlsConfigBuilder{ + ClientKeyCertChainPath: certPath, + ClientPrivateKeyPath: keyPath, + ClientTrustCertCollectionPath: caPath, + }).BuildTlsConfig() + if err != nil { + t.Fatal(err) + } + if config.MinVersion != tls.VersionTLS12 { + t.Fatalf("MinVersion = %d, want TLS 1.2 (%d)", config.MinVersion, tls.VersionTLS12) + } + if config.InsecureSkipVerify { + t.Fatal("InsecureSkipVerify is true; certificate verification must stay enabled") + } + if config.RootCAs == nil { + t.Fatal("RootCAs is nil; the configured trust collection was not loaded") + } + if len(config.Certificates) != 1 { + t.Fatalf("Certificates contains %d entries, want 1", len(config.Certificates)) + } + expectedRootCAs := x509.NewCertPool() + if !expectedRootCAs.AppendCertsFromPEM(tlsTestRootCertificate) { + t.Fatal("failed to parse the expected root certificate") + } + if !config.RootCAs.Equal(expectedRootCAs) { + t.Fatal("RootCAs does not contain the configured trust certificate") + } + expectedClientCertificate, _ := pem.Decode(WssServerCRT) + if expectedClientCertificate == nil { + t.Fatal("failed to decode the expected client certificate") + } + expectedRootCertificate, _ := pem.Decode(tlsTestRootCertificate) + if expectedRootCertificate == nil { + t.Fatal("failed to decode the expected root certificate") + } + if bytes.Equal(expectedClientCertificate.Bytes, expectedRootCertificate.Bytes) { + t.Fatal("client and root certificate fixtures must be distinct") + } + if len(config.Certificates[0].Certificate) == 0 { + t.Fatal("configured client certificate has an empty certificate chain") + } + if !bytes.Equal(config.Certificates[0].Certificate[0], expectedClientCertificate.Bytes) { + t.Fatal("configured client certificate does not match the requested certificate") + } +} + +func TestClientTLSConfigBuilderRejectsInvalidTrustCollection(t *testing.T) { + tempDir := t.TempDir() + certPath := filepath.Join(tempDir, "client.crt") + keyPath := filepath.Join(tempDir, "client.key") + caPath := filepath.Join(tempDir, "ca.crt") + for path, data := range map[string][]byte{ + certPath: WssServerCRT, + keyPath: WssServerKEY, + caPath: []byte("not a certificate"), + } { + if err := os.WriteFile(path, data, 0o600); err != nil { + t.Fatal(err) + } + } + + config, err := (&ClientTlsConfigBuilder{ + ClientKeyCertChainPath: certPath, + ClientPrivateKeyPath: keyPath, + ClientTrustCertCollectionPath: caPath, + }).BuildTlsConfig() + if err == nil { + t.Fatal("invalid trust collection returned nil error") + } + if config != nil { + t.Fatal("config must be nil when BuildTlsConfig fails") + } +}