Avoid using deadlines in ObeyIdleTimeout (#35006)

* Avoid using deadlines in ObeyIdleTimeout

* wrap the error on Close

* add test and fix a bug

* switch to just including a clockwork.Timer
This commit is contained in:
Edoardo Spadolini
2023-11-29 11:36:42 +00:00
committed by GitHub
parent 1e0a7048a6
commit 731b9bf52b
4 changed files with 133 additions and 43 deletions
+1 -4
View File
@@ -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()
+2 -3
View File
@@ -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
}
+59 -36
View File
@@ -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
}
+71
View File
@@ -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)
}