diff --git a/lib/reversetunnel/srv.go b/lib/reversetunnel/srv.go index d9a773b8787..85f156bdeb0 100644 --- a/lib/reversetunnel/srv.go +++ b/lib/reversetunnel/srv.go @@ -638,10 +638,7 @@ func (s *server) Shutdown(ctx context.Context) error { } func (s *server) HandleNewChan(ctx context.Context, ccx *sshutils.ConnectionContext, nch ssh.NewChannel) { - // Apply read/write timeouts to the server connection. - conn := utils.ObeyIdleTimeout(ccx.NetConn, - s.offlineThreshold, - "reverse tunnel server") + conn := ccx.NetConn sconn := ccx.ServerConn channelType := nch.ChannelType() diff --git a/lib/sshutils/server.go b/lib/sshutils/server.go index 32b182059a2..ae166a4e1bc 100644 --- a/lib/sshutils/server.go +++ b/lib/sshutils/server.go @@ -472,8 +472,7 @@ func (s *Server) HandleConnection(conn net.Conn) { // apply idle read/write timeout to this connection. conn = utils.ObeyIdleTimeout(conn, - defaults.DefaultIdleConnectionDuration, - s.component) + defaults.DefaultIdleConnectionDuration) // Wrap connection with a tracker used to monitor how much data was // transmitted and received over the connection. wconn := utils.NewTrackingConn(conn) @@ -487,7 +486,7 @@ func (s *Server) HandleConnection(conn net.Conn) { WithField("remote_addr", conn.RemoteAddr()). Warn("Error occurred in handshake for new SSH conn") } - conn.SetDeadline(time.Time{}) + conn.Close() return } diff --git a/lib/utils/timeout.go b/lib/utils/timeout.go index 9b47a035592..4b0f644cb13 100644 --- a/lib/utils/timeout.go +++ b/lib/utils/timeout.go @@ -19,52 +19,75 @@ package utils import ( "net" + "sync" "time" + + "github.com/gravitational/trace" + "github.com/jonboulle/clockwork" ) -// TimeoutConn wraps an existing net.Conn and adds read/write timeouts -// for it, allowing to implement "disconnect after XX of idle time" policy -// -// Usage example: -// tc := utils.ObeyIdleTimeout(conn, time.Second * 30, "ssh connection") -// io.Copy(tc, xxx) -type TimeoutConn struct { - net.Conn - TimeoutDuration time.Duration - - // Name is only useful for debugging/logging, it's a convenient - // way to tag every idle connection - OwnerName string +// ObeyIdleTimeout wraps an existing network connection, closing it if data +// isn't read often enough. The connection will be closed even if Read is never +// called, or if it's called on the underlying connection instead of the +// returned one. +func ObeyIdleTimeout(conn net.Conn, timeout time.Duration) net.Conn { + return obeyIdleTimeoutClock(conn, timeout, clockwork.NewRealClock()) } -// ObeyIdleTimeout wraps an existing network connection with timeout-obeying -// Write() and Read() - it will drop the connection after 'timeout' on idle -// -// Example: -// ObeyIdletimeout(conn, time.Second * 60, "api server"). -func ObeyIdleTimeout(conn net.Conn, timeout time.Duration, ownerName string) net.Conn { - return &TimeoutConn{ - Conn: conn, - TimeoutDuration: timeout, - OwnerName: ownerName, +// obeyIdleTimeoutClock is [ObeyIdleTimeout] but lets the caller specify an +// arbitrary [clockwork.Clock] to be used for the timer. +func obeyIdleTimeoutClock(conn net.Conn, timeout time.Duration, clock clockwork.Clock) net.Conn { + return &timeoutConn{ + Conn: conn, + timeout: timeout, + watchdog: clock.AfterFunc(timeout, func() { + conn.Close() + }), } } -// NetConn returns the underlying net.Conn. -func (tc *TimeoutConn) NetConn() net.Conn { - return tc.Conn +type timeoutConn struct { + net.Conn + + timeout time.Duration + + mu sync.Mutex + watchdog clockwork.Timer } -func (tc *TimeoutConn) Read(p []byte) (n int, err error) { - // note: checking for errors here does not buy anything: some net.Conn interface - // implementations (sshConn, pipe) simply return "not supported" error - tc.Conn.SetReadDeadline(time.Now().Add(tc.TimeoutDuration)) - return tc.Conn.Read(p) +func (c *timeoutConn) pet() { + c.mu.Lock() + defer c.mu.Unlock() + // if the timer has already fired the underlying net.Conn has been closed or + // will be closed shortly anyway + if c.watchdog.Stop() { + c.watchdog.Reset(c.timeout) + } } -func (tc *TimeoutConn) Write(p []byte) (n int, err error) { - // note: checking for errors here does not buy anything: some net.Conn interface - // implementations (sshConn, pipe) simply return "not supported" error - tc.Conn.SetWriteDeadline(time.Now().Add(tc.TimeoutDuration)) - return tc.Conn.Write(p) +// NetConn returns the underlying [net.Conn]. +func (c *timeoutConn) NetConn() net.Conn { + return c.Conn +} + +// Close implements [io.Closer] and [net.Conn] by closing the underlying +// connection and then stopping the watchdog, if it's still running. +func (c *timeoutConn) Close() error { + err := c.Conn.Close() + c.mu.Lock() + defer c.mu.Unlock() + c.watchdog.Stop() + return trace.Wrap(err) +} + +// Read implements [io.Reader] and [net.Conn], petting the watchdog timer if any +// data is successfully read. +func (c *timeoutConn) Read(p []byte) (n int, err error) { + n, err = c.Conn.Read(p) + if n > 0 { + c.pet() + } + // avoid trace.Wrap to maintain the exact errors from the underlying + // connection (like io.EOF) + return n, err } diff --git a/lib/utils/timeout_test.go b/lib/utils/timeout_test.go new file mode 100644 index 00000000000..42bd76ab843 --- /dev/null +++ b/lib/utils/timeout_test.go @@ -0,0 +1,71 @@ +// Copyright 2023 Gravitational, Inc +// +// Licensed 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 utils + +import ( + "io" + "net" + "testing" + "time" + + "github.com/jonboulle/clockwork" + "github.com/stretchr/testify/require" +) + +func TestObeyIdleTimeout(t *testing.T) { + clock := clockwork.NewFakeClock() + t.Cleanup(func() { clock.Advance(time.Hour) }) + + c1, c2 := net.Pipe() + c1 = obeyIdleTimeoutClock(c1, time.Minute, clock) + t.Cleanup(func() { c1.Close() }) + t.Cleanup(func() { c2.Close() }) + + go func() { + c2.Write([]byte{0}) + clock.Sleep(30 * time.Second) + c2.Write([]byte{0}) + }() + + errC := make(chan error, 3) + go func() { + var b [1]byte + for i := 0; i < 3; i++ { + _, err := io.ReadFull(c1, b[:]) + errC <- err + } + }() + + err1 := <-errC + // wait for the writing goroutine to be waiting as well (the watchdog counts + // as a waiter) + clock.BlockUntil(2) + clock.Advance(30 * time.Second) + err2 := <-errC + clock.Advance(30 * time.Second) + select { + case err := <-errC: + require.FailNow(t, "expected Read to block", "got err %v", err) + case <-time.After(50 * time.Millisecond): + } + clock.Advance(30 * time.Second) + err3 := <-errC + + require.NoError(t, err1) + require.NoError(t, err2) + // this should be net.ErrClosed, but net.Pipe uses io.ErrClosedPipe and it + // can't be changed + require.ErrorIs(t, err3, io.ErrClosedPipe) +}