diff --git a/options_test.go b/options_test.go index 584bfea..458e9dc 100644 --- a/options_test.go +++ b/options_test.go @@ -81,7 +81,7 @@ type wrappedConn struct { func (c *wrappedConn) Write(p []byte) (n int, err error) { n, err = c.Conn.Write(p) - atomic.AddInt32(&(c.written), int32(n)) //nolint:gosec // test code, overflow not possible + atomic.AddInt32(&c.written, int32(n)) //nolint:gosec // test code, overflow not possible return } @@ -108,7 +108,7 @@ func TestConnWrapping(t *testing.T) { if err := session.Shell(); err != nil { t.Fatal(err) } - if atomic.LoadInt32(&(wrapped.written)) == 0 { + if atomic.LoadInt32(&wrapped.written) == 0 { t.Fatal("wrapped conn not written to") } } diff --git a/server.go b/server.go index a95acc3..8a624b9 100644 --- a/server.go +++ b/server.go @@ -571,22 +571,25 @@ func (srv *Server) connectionKeepAlive( // next tick enforces the deadline, and if SendRequest // hangs forever it will be unblocked when the TimeIsUp // branch closes sshConn. - var err error + var ( + ok bool + err error + ) ch := openChans.any() if ch != nil { - _, err = ch.SendRequest(keepAliveRequestType, true, nil) + ok, err = ch.SendRequest(keepAliveRequestType, true, nil) if err != nil { openChans.remove(ch) ch = nil } } if ch == nil { - _, _, err = sshConn.SendRequest(keepAliveRequestType, true, nil) + ok, _, err = sshConn.SendRequest(keepAliveRequestType, true, nil) } - if err == nil { + if err == nil && ok { keepAlive.Reset() } else { - log.Printf("ssh: keepalive request failed: %v", err) + log.Printf("ssh: keepalive request failed: ok=%t err=%v", ok, err) } }() } diff --git a/server_test.go b/server_test.go index 298acc8..d3b20aa 100644 --- a/server_test.go +++ b/server_test.go @@ -373,6 +373,129 @@ func TestConnectionKeepAliveUsesChannelRequestWhenSessionOpen(t *testing.T) { } } +// TestConnectionKeepAliveNegativeGlobalReplyDoesNotReset verifies that a +// protocol-level negative response is not counted as a successful keepalive. +func TestConnectionKeepAliveNegativeGlobalReplyDoesNotReset(t *testing.T) { + t.Parallel() + + closingFired := make(chan struct{}) + srv := &Server{ + Handler: func(_ Session) {}, + ClientAliveInterval: 100 * time.Millisecond, + ClientAliveCountMax: 2, + ConnectionClosingCallback: func(_ Context, _ *gossh.ServerConn) { + close(closingFired) + }, + } + + l := newLocalTCPListener() + defer func() { _ = l.Close() }() + go func() { _ = srv.serveOnce(l) }() + + cfg := &gossh.ClientConfig{ + User: "testuser", + Auth: []gossh.AuthMethod{gossh.Password("testpass")}, + HostKeyCallback: gossh.InsecureIgnoreHostKey(), //nolint:gosec // test code + } + netConn, err := net.Dial("tcp", l.Addr().String()) + if err != nil { + t.Fatalf("dial: %v", err) + } + sshConn, chans, reqs, err := gossh.NewClientConn(netConn, l.Addr().String(), cfg) + if err != nil { + t.Fatalf("NewClientConn: %v", err) + } + defer func() { _ = sshConn.Close() }() + + go func() { + for range chans { //nolint:revive // intentional drain + } + }() + go func() { + for req := range reqs { + if req.Type == keepAliveRequestType { + _ = req.Reply(false, nil) + } else if req.WantReply { + _ = req.Reply(true, nil) + } + } + }() + + select { + case <-closingFired: + case <-time.After(5 * time.Second): + t.Fatal("negative global keepalive replies incorrectly reset the deadline") + } +} + +// TestConnectionKeepAliveNegativeChannelReplyDoesNotReset is the channel +// request counterpart to TestConnectionKeepAliveNegativeGlobalReplyDoesNotReset. +func TestConnectionKeepAliveNegativeChannelReplyDoesNotReset(t *testing.T) { + t.Parallel() + + closingFired := make(chan struct{}) + srv := &Server{ + Handler: func(s Session) { <-s.Context().Done() }, + ClientAliveInterval: 100 * time.Millisecond, + ClientAliveCountMax: 2, + ConnectionClosingCallback: func(_ Context, _ *gossh.ServerConn) { + close(closingFired) + }, + } + + l := newLocalTCPListener() + defer func() { _ = l.Close() }() + go func() { _ = srv.serveOnce(l) }() + + cfg := &gossh.ClientConfig{ + User: "testuser", + Auth: []gossh.AuthMethod{gossh.Password("testpass")}, + HostKeyCallback: gossh.InsecureIgnoreHostKey(), //nolint:gosec // test code + } + netConn, err := net.Dial("tcp", l.Addr().String()) + if err != nil { + t.Fatalf("dial: %v", err) + } + sshConn, chans, reqs, err := gossh.NewClientConn(netConn, l.Addr().String(), cfg) + if err != nil { + t.Fatalf("NewClientConn: %v", err) + } + defer func() { _ = sshConn.Close() }() + + go func() { + for range chans { //nolint:revive // intentional drain + } + }() + go func() { + for req := range reqs { + if req.WantReply { + _ = req.Reply(true, nil) + } + } + }() + + ch, chReqs, err := sshConn.OpenChannel("session", nil) + if err != nil { + t.Fatalf("OpenChannel: %v", err) + } + defer func() { _ = ch.Close() }() + go func() { + for req := range chReqs { + if req.Type == keepAliveRequestType { + _ = req.Reply(false, nil) + } else if req.WantReply { + _ = req.Reply(true, nil) + } + } + }() + + select { + case <-closingFired: + case <-time.After(5 * time.Second): + t.Fatal("negative channel keepalive replies incorrectly reset the deadline") + } +} + // TestConnectionKeepAlivePrunesClosedChannels verifies that the // per-channel close hook prunes channels from the openChannelSet as // soon as the client closes them, BEFORE the next keepalive probe