diff --git a/integration/helpers.go b/integration/helpers.go index 5e8cb144372..8ec0c9c73b6 100644 --- a/integration/helpers.go +++ b/integration/helpers.go @@ -66,7 +66,8 @@ import ( // SetTestTimeouts affects global timeouts inside Teleport, making connections // work faster but consuming more CPU (useful for integration testing) func SetTestTimeouts(t time.Duration) { - defaults.ReverseTunnelAgentHeartbeatPeriod = t + defaults.KeepAliveInterval = t + defaults.ResyncInterval = t defaults.ServerKeepAliveTTL = t defaults.SessionRefreshPeriod = t defaults.HeartbeatCheckPeriod = t diff --git a/lib/auth/clt.go b/lib/auth/clt.go index 14f704d5a46..a4b2391ec47 100644 --- a/lib/auth/clt.go +++ b/lib/auth/clt.go @@ -133,10 +133,10 @@ func DecodeClusterName(serverName string) (string, error) { } // NewAddrDialer returns new dialer from a list of addresses -func NewAddrDialer(addrs []utils.NetAddr) ContextDialer { +func NewAddrDialer(addrs []utils.NetAddr, keepAliveInterval time.Duration) ContextDialer { dialer := net.Dialer{ Timeout: defaults.DefaultDialTimeout, - KeepAlive: defaults.ReverseTunnelAgentHeartbeatPeriod, + KeepAlive: keepAliveInterval, } return ContextDialerFunc(func(in context.Context, network, _ string) (net.Conn, error) { var err error @@ -198,7 +198,7 @@ func (c *ClientConfig) CheckAndSetDefaults() error { c.KeepAliveCount = defaults.KeepAliveCountMax } if c.Dialer == nil { - c.Dialer = NewAddrDialer(c.Addrs) + c.Dialer = NewAddrDialer(c.Addrs, c.KeepAlivePeriod) } if c.TLS.ServerName == "" { c.TLS.ServerName = teleport.APIDomain diff --git a/lib/auth/trustedcluster.go b/lib/auth/trustedcluster.go index 59acff61cd7..b73899f4d55 100644 --- a/lib/auth/trustedcluster.go +++ b/lib/auth/trustedcluster.go @@ -22,6 +22,7 @@ import ( "net/http" "net/url" "strings" + "time" "github.com/gravitational/teleport" "github.com/gravitational/teleport/lib" @@ -339,6 +340,13 @@ func (a *AuthServer) GetRemoteCluster(clusterName string) (services.RemoteCluste } func (a *AuthServer) updateRemoteClusterStatus(remoteCluster services.RemoteCluster) error { + clusterConfig, err := a.GetClusterConfig() + if err != nil { + return trace.Wrap(err) + } + keepAliveCountMax := clusterConfig.GetKeepAliveCountMax() + keepAliveInterval := clusterConfig.GetKeepAliveInterval() + // fetch tunnel connections for the cluster to update runtime status connections, err := a.GetTunnelConnections(remoteCluster.GetName()) if err != nil { @@ -347,7 +355,9 @@ func (a *AuthServer) updateRemoteClusterStatus(remoteCluster services.RemoteClus remoteCluster.SetConnectionStatus(teleport.RemoteClusterStatusOffline) lastConn, err := services.LatestTunnelConnection(connections) if err == nil { - remoteCluster.SetConnectionStatus(services.TunnelConnectionStatus(a.clock, lastConn)) + offlineThreshold := time.Duration(keepAliveCountMax) * keepAliveInterval + tunnelStatus := services.TunnelConnectionStatus(a.clock, lastConn, offlineThreshold) + remoteCluster.SetConnectionStatus(tunnelStatus) remoteCluster.SetLastHeartbeat(lastConn.GetLastHeartbeat()) } return nil diff --git a/lib/defaults/defaults.go b/lib/defaults/defaults.go index 243ad6fc4cd..487635fb69e 100644 --- a/lib/defaults/defaults.go +++ b/lib/defaults/defaults.go @@ -233,12 +233,8 @@ const ( ) var ( - // ReverseTunnelAgentHeartbeatPeriod is the period between agent heartbeat messages - ReverseTunnelAgentHeartbeatPeriod = 5 * time.Second - - // ReverseTunnelOfflineThreshold is the threshold of missed heartbeats - // after which we are going to declare the reverse tunnel offline - ReverseTunnelOfflineThreshold = 5 * ReverseTunnelAgentHeartbeatPeriod + // ResyncInterval is how often tunnels are resynced. + ResyncInterval = 5 * time.Second // ServerAnnounceTTL is a period between heartbeats // Median sleep time between node pings is this value / 2 + random diff --git a/lib/reversetunnel/agentpool.go b/lib/reversetunnel/agentpool.go index 05009b4a95d..b6cbd3b63b4 100644 --- a/lib/reversetunnel/agentpool.go +++ b/lib/reversetunnel/agentpool.go @@ -161,7 +161,7 @@ func (m *AgentPool) Wait() error { } func (m *AgentPool) processDiscoveryRequests() { - ticker := time.NewTicker(defaults.ReverseTunnelAgentHeartbeatPeriod) + ticker := time.NewTicker(defaults.ResyncInterval) defer ticker.Stop() for { @@ -298,7 +298,7 @@ func (m *AgentPool) closeAgents(matchKey *agentKey) { } func (m *AgentPool) pollAndSyncAgents() { - ticker := time.NewTicker(defaults.ReverseTunnelAgentHeartbeatPeriod) + ticker := time.NewTicker(defaults.ResyncInterval) defer ticker.Stop() m.FetchAndSyncAgents() for { diff --git a/lib/reversetunnel/conn.go b/lib/reversetunnel/conn.go index dd7da025857..080c7a6558c 100644 --- a/lib/reversetunnel/conn.go +++ b/lib/reversetunnel/conn.go @@ -95,6 +95,10 @@ type connConfig struct { // nodeID is used when tunnelType is node and is set // to the node UUID dialing back nodeID string + + // offlineThreshold is how long to wait for a keep alive message before + // marking a reverse tunnel connection as invalid. + offlineThreshold time.Duration } func newRemoteConn(cfg *connConfig) *remoteConn { @@ -252,5 +256,6 @@ func (c *remoteConn) sendDiscoveryRequest(req discoveryRequest) error { } func (c *remoteConn) isOnline(conn services.TunnelConnection) bool { - return services.TunnelConnectionStatus(c.clock, conn) == teleport.RemoteClusterStatusOnline + tunnelStatus := services.TunnelConnectionStatus(c.clock, conn, c.offlineThreshold) + return tunnelStatus == teleport.RemoteClusterStatusOnline } diff --git a/lib/reversetunnel/localsite.go b/lib/reversetunnel/localsite.go index 9c0706c4c8d..0bf562117da 100644 --- a/lib/reversetunnel/localsite.go +++ b/lib/reversetunnel/localsite.go @@ -66,6 +66,7 @@ func newlocalSite(srv *server, domainName string, client auth.ClientI) (*localSi "cluster": domainName, }, }), + offlineThreshold: srv.offlineThreshold, }, nil } @@ -100,6 +101,10 @@ type localSite struct { // clock is used to control time in tests. clock clockwork.Clock + + // offlineThreshold is how long to wait for a keep alive message before + // marking a reverse tunnel connection as invalid. + offlineThreshold time.Duration } // GetTunnelsCount always the number of tunnel connections to this cluster. @@ -268,13 +273,14 @@ func (s *localSite) addConn(nodeID string, conn net.Conn, sconn ssh.Conn) (*remo defer s.Unlock() rconn := newRemoteConn(&connConfig{ - conn: conn, - sconn: sconn, - accessPoint: s.accessPoint, - tunnelType: string(services.NodeTunnel), - proxyName: s.srv.ID, - clusterName: s.domainName, - nodeID: nodeID, + conn: conn, + sconn: sconn, + accessPoint: s.accessPoint, + tunnelType: string(services.NodeTunnel), + proxyName: s.srv.ID, + clusterName: s.domainName, + nodeID: nodeID, + offlineThreshold: s.offlineThreshold, }) s.remoteConns[nodeID] = rconn @@ -349,10 +355,9 @@ func (s *localSite) handleHeartbeat(rconn *remoteConn, ch ssh.Channel, reqC <-ch } tm := time.Now().UTC() rconn.setLastHeartbeat(tm) - // Since we block on select, time.After is re-created everytime we process - // a request. - case <-time.After(defaults.ReverseTunnelOfflineThreshold): - rconn.markInvalid(trace.ConnectionProblem(nil, "no heartbeats for %v", defaults.ReverseTunnelOfflineThreshold)) + // Note that time.After is re-created everytime a request is processed. + case <-time.After(s.offlineThreshold): + rconn.markInvalid(trace.ConnectionProblem(nil, "no heartbeats for %v", s.offlineThreshold)) } } } diff --git a/lib/reversetunnel/peer.go b/lib/reversetunnel/peer.go index 5488e7941db..79cf312e3b0 100644 --- a/lib/reversetunnel/peer.go +++ b/lib/reversetunnel/peer.go @@ -23,10 +23,10 @@ import ( "github.com/gravitational/teleport" "github.com/gravitational/teleport/lib/auth" - "github.com/gravitational/teleport/lib/defaults" "github.com/gravitational/teleport/lib/services" "github.com/gravitational/trace" + "github.com/jonboulle/clockwork" log "github.com/sirupsen/logrus" ) @@ -133,7 +133,7 @@ func (p *clusterPeers) DialTCP(params DialParams) (conn net.Conn, err error) { } // newClusterPeer returns new cluster peer -func newClusterPeer(srv *server, connInfo services.TunnelConnection) (*clusterPeer, error) { +func newClusterPeer(srv *server, connInfo services.TunnelConnection, offlineThreshold time.Duration) (*clusterPeer, error) { clusterPeer := &clusterPeer{ srv: srv, connInfo: connInfo, @@ -143,6 +143,8 @@ func newClusterPeer(srv *server, connInfo services.TunnelConnection) (*clusterPe "cluster": connInfo.GetClusterName(), }, }), + clock: clockwork.NewRealClock(), + offlineThreshold: offlineThreshold, } return clusterPeer, nil @@ -154,6 +156,13 @@ type clusterPeer struct { log *log.Entry connInfo services.TunnelConnection srv *server + + // clock is used to control time in tests. + clock clockwork.Clock + + // offlineThreshold is how long to wait for a keep alive message before + // marking a reverse tunnel connection as invalid. + offlineThreshold time.Duration } func (s *clusterPeer) CachingAccessPoint() (auth.AccessPoint, error) { @@ -169,11 +178,7 @@ func (s *clusterPeer) String() string { } func (s *clusterPeer) GetStatus() string { - diff := time.Now().Sub(s.connInfo.GetLastHeartbeat()) - if diff > defaults.ReverseTunnelOfflineThreshold { - return teleport.RemoteClusterStatusOffline - } - return teleport.RemoteClusterStatusOnline + return services.TunnelConnectionStatus(s.clock, s.connInfo, s.offlineThreshold) } func (s *clusterPeer) GetName() string { diff --git a/lib/reversetunnel/remotesite.go b/lib/reversetunnel/remotesite.go index 370aa842793..a604ff6cd2c 100644 --- a/lib/reversetunnel/remotesite.go +++ b/lib/reversetunnel/remotesite.go @@ -27,7 +27,6 @@ import ( "github.com/gravitational/teleport" "github.com/gravitational/teleport/lib/auth" - "github.com/gravitational/teleport/lib/defaults" "github.com/gravitational/teleport/lib/services" "github.com/gravitational/teleport/lib/srv/forward" "github.com/gravitational/trace" @@ -78,6 +77,10 @@ type remoteSite struct { // state has been changed, the tunnel will reconnect to re-create the client // with new settings. remoteCA services.CertAuthority + + // offlineThreshold is how long to wait for a keep alive message before + // marking a reverse tunnel connection as invalid. + offlineThreshold time.Duration } func (s *remoteSite) getRemoteClient() (auth.ClientI, bool, error) { @@ -227,12 +230,13 @@ func (s *remoteSite) addConn(conn net.Conn, sconn ssh.Conn) (*remoteConn, error) defer s.Unlock() rconn := newRemoteConn(&connConfig{ - conn: conn, - sconn: sconn, - accessPoint: s.localAccessPoint, - tunnelType: string(services.ProxyTunnel), - proxyName: s.connInfo.GetProxyName(), - clusterName: s.domainName, + conn: conn, + sconn: sconn, + accessPoint: s.localAccessPoint, + tunnelType: string(services.ProxyTunnel), + proxyName: s.connInfo.GetProxyName(), + clusterName: s.domainName, + offlineThreshold: s.offlineThreshold, }) s.connections = append(s.connections, rconn) @@ -245,7 +249,7 @@ func (s *remoteSite) GetStatus() string { if err != nil { return teleport.RemoteClusterStatusOffline } - return services.TunnelConnectionStatus(s.clock, connInfo) + return services.TunnelConnectionStatus(s.clock, connInfo, s.offlineThreshold) } func (s *remoteSite) copyConnInfo() services.TunnelConnection { @@ -272,7 +276,7 @@ func (s *remoteSite) getLastConnInfo() (services.TunnelConnection, error) { func (s *remoteSite) registerHeartbeat(t time.Time) { connInfo := s.copyConnInfo() connInfo.SetLastHeartbeat(t) - connInfo.SetExpiry(s.clock.Now().Add(defaults.ReverseTunnelOfflineThreshold)) + connInfo.SetExpiry(s.clock.Now().Add(s.offlineThreshold)) s.setLastConnInfo(connInfo) err := s.localAccessPoint.UpsertTunnelConnection(connInfo) if err != nil { @@ -307,6 +311,7 @@ func (s *remoteSite) handleHeartbeat(conn *remoteConn, ch ssh.Channel, reqC <-ch s.Infof("Cluster connection closed.") conn.Close() }() + firstHeartbeat := true for { select { @@ -358,9 +363,9 @@ func (s *remoteSite) handleHeartbeat(conn *remoteConn, ch ssh.Channel, reqC <-ch tm := time.Now().UTC() conn.setLastHeartbeat(tm) go s.registerHeartbeat(tm) - // since we block on select, time.After is re-created everytime we process a request. - case <-time.After(defaults.ReverseTunnelOfflineThreshold): - conn.markInvalid(trace.ConnectionProblem(nil, "no heartbeats for %v", defaults.ReverseTunnelOfflineThreshold)) + // Note that time.After is re-created everytime a request is processed. + case <-time.After(s.offlineThreshold): + conn.markInvalid(trace.ConnectionProblem(nil, "no heartbeats for %v", s.offlineThreshold)) } } } diff --git a/lib/reversetunnel/srv.go b/lib/reversetunnel/srv.go index c406f8ef806..4238649c5b2 100644 --- a/lib/reversetunnel/srv.go +++ b/lib/reversetunnel/srv.go @@ -107,6 +107,10 @@ type server struct { // proxyWatcher monitors changes to the proxies // and broadcasts updates proxyWatcher *services.ProxyWatcher + + // offlineThreshold is how long to wait for a keep alive message before + // marking a reverse tunnel connection as invalid. + offlineThreshold time.Duration } // DirectCluster is used to access cluster directly @@ -222,6 +226,12 @@ func NewServer(cfg Config) (Server, error) { return nil, trace.Wrap(err) } + clusterConfig, err := cfg.LocalAccessPoint.GetClusterConfig() + if err != nil { + return nil, trace.Wrap(err) + } + offlineThreshold := time.Duration(clusterConfig.GetKeepAliveCountMax()) * clusterConfig.GetKeepAliveInterval() + ctx, cancel := context.WithCancel(cfg.Context) entry := log.WithFields(log.Fields{ @@ -252,6 +262,7 @@ func NewServer(cfg Config) (Server, error) { proxyWatcher: proxyWatcher, clusterPeers: make(map[string]*clusterPeers), Entry: entry, + offlineThreshold: offlineThreshold, } for _, clusterInfo := range cfg.DirectClusters { @@ -321,7 +332,7 @@ func (s *server) disconnectClusters() error { } func (s *server) periodicFunctions() { - ticker := time.NewTicker(defaults.ReverseTunnelAgentHeartbeatPeriod) + ticker := time.NewTicker(defaults.ResyncInterval) defer ticker.Stop() if err := s.fetchClusterPeers(); err != nil { @@ -402,7 +413,7 @@ func (s *server) reportClusterStats() error { func (s *server) addClusterPeers(conns map[string]services.TunnelConnection) error { for key := range conns { connInfo := conns[key] - peer, err := newClusterPeer(s, connInfo) + peer, err := newClusterPeer(s, connInfo, s.offlineThreshold) if err != nil { return trace.Wrap(err) } @@ -513,7 +524,7 @@ func (s *server) Shutdown(ctx context.Context) error { func (s *server) HandleNewChan(conn net.Conn, sconn *ssh.ServerConn, nch ssh.NewChannel) { // Apply read/write timeouts to the server connection. conn = utils.ObeyIdleTimeout(conn, - defaults.ReverseTunnelAgentHeartbeatPeriod*10, + s.offlineThreshold, "reverse tunnel server") channelType := nch.ChannelType() @@ -951,9 +962,10 @@ func newRemoteSite(srv *server, domainName string) (*remoteSite, error) { "cluster": domainName, }, }), - ctx: closeContext, - cancel: cancel, - clock: srv.Clock, + ctx: closeContext, + cancel: cancel, + clock: srv.Clock, + offlineThreshold: srv.offlineThreshold, } // configure access to the full Auth Server API and the cached subset for diff --git a/lib/services/tunnelconn.go b/lib/services/tunnelconn.go index 66cfb3177e8..b7ff64be615 100644 --- a/lib/services/tunnelconn.go +++ b/lib/services/tunnelconn.go @@ -75,9 +75,9 @@ func LatestTunnelConnection(conns []TunnelConnection) (TunnelConnection, error) // TunnelConnectionStatus returns tunnel connection status based on the last // heartbeat time recorded for a connection -func TunnelConnectionStatus(clock clockwork.Clock, conn TunnelConnection) string { +func TunnelConnectionStatus(clock clockwork.Clock, conn TunnelConnection, offlineThreshold time.Duration) string { diff := clock.Now().Sub(conn.GetLastHeartbeat()) - if diff < defaults.ReverseTunnelOfflineThreshold { + if diff < offlineThreshold { return teleport.RemoteClusterStatusOnline } return teleport.RemoteClusterStatusOffline