mirror of
https://github.com/gravitational/teleport.git
synced 2026-09-24 16:17:11 +08:00
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:
@@ -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()
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
Reference in New Issue
Block a user