diff --git a/lib/reversetunnel/conn.go b/lib/reversetunnel/conn.go index 6034668651c..5c86bae1b14 100644 --- a/lib/reversetunnel/conn.go +++ b/lib/reversetunnel/conn.go @@ -64,8 +64,8 @@ type remoteConn struct { // discoveryCh is the SSH channel over which discovery requests are sent. discoveryCh ssh.Channel - // newProxiesC is a list used to nofity about new proxies - newProxiesC chan []types.Server + // proxyDiscoverySubscriber receives proxy updates. + proxyDiscoverySubscriber *proxyDiscoverySubscriber // invalid indicates the connection is invalid and connections can no longer // be made on it. @@ -112,17 +112,22 @@ type connConfig struct { // offlineThreshold is how long to wait for a keep alive message before // marking a reverse tunnel connection as invalid. offlineThreshold time.Duration + + // proxyDiscoverySubscriber receives proxy discovery events. + proxyDiscoverySubscriber *proxyDiscoverySubscriber } -func newRemoteConn(cfg *connConfig) *remoteConn { - c := &remoteConn{ - logger: slog.With(teleport.ComponentKey, "discovery"), - connConfig: cfg, - clock: clockwork.NewRealClock(), - newProxiesC: make(chan []types.Server, 100), +func newRemoteConn(cfg *connConfig) (*remoteConn, error) { + if cfg.proxyDiscoverySubscriber == nil { + return nil, trace.BadParameter("missing proxy discovery subscription") } - - return c + c := &remoteConn{ + logger: slog.With(teleport.ComponentKey, "discovery"), + connConfig: cfg, + clock: clockwork.NewRealClock(), + proxyDiscoverySubscriber: cfg.proxyDiscoverySubscriber, + } + return c, nil } func (c *remoteConn) String() string { @@ -256,19 +261,6 @@ func (c *remoteConn) openDiscoveryChannel() (ssh.Channel, error) { return c.discoveryCh, nil } -// updateProxies is a non-blocking call that puts the new proxies -// list so that remote connection can notify the remote agent -// about the list update -func (c *remoteConn) updateProxies(proxies []types.Server) { - select { - case c.newProxiesC <- proxies: - default: - // Missing proxies update is no longer critical with more permissive - // discovery protocol that tolerates conflicting, stale or missing updates - c.logger.WarnContext(context.Background(), "Discovery channel overflow", "new_proxy_count", len(c.newProxiesC)) - } -} - func (c *remoteConn) adviseReconnect() error { _, _, err := c.sconn.SendRequest(reconnectRequest, true, nil) return trace.Wrap(err) diff --git a/lib/reversetunnel/discovery.go b/lib/reversetunnel/discovery.go index 57bc4c78a14..444224e227d 100644 --- a/lib/reversetunnel/discovery.go +++ b/lib/reversetunnel/discovery.go @@ -19,12 +19,20 @@ package reversetunnel import ( + "context" "fmt" + "log/slog" + "os" "strconv" "strings" + "sync/atomic" + "time" + "github.com/gravitational/teleport/api/defaults" "github.com/gravitational/teleport/api/types" "github.com/gravitational/teleport/lib/reversetunnel/track" + "github.com/gravitational/teleport/lib/services" + "github.com/gravitational/teleport/lib/services/readonly" ) // discoveryRequest is the minimal structure that can be exchanged as JSON as a @@ -48,28 +56,9 @@ type discoveryProxy struct { ProxyGroupID string `json:"gid,omitempty"` ProxyGroupGeneration uint64 `json:"ggen,omitempty"` -} - -// SetProxies overwrites the proxy list in the discoveryRequest with data from -// the slice of [types.Server]s. -func (r *discoveryRequest) SetProxies(proxies []types.Server) { - r.Proxies = make([]discoveryProxy, 0, len(proxies)) - for _, proxy := range proxies { - d := discoveryProxy{ - Version: types.V2, - } - d.Metadata.Name = proxy.GetName() - d.ProxyGroupID, _ = proxy.GetLabel(types.ProxyGroupIDLabel) - proxyGroupGeneration, _ := proxy.GetLabel(types.ProxyGroupGenerationLabel) - var err error - d.ProxyGroupGeneration, err = strconv.ParseUint(proxyGroupGeneration, 10, 64) - if err != nil { - // ParseUint can return the maximum uint64 on ErrRange - d.ProxyGroupGeneration = 0 - } - - r.Proxies = append(r.Proxies, d) - } + // TTL is used by the agent [track.Tracker] for proxy expiry. This allows + // server side override of the trackers default expiry. + TTL time.Duration `json:"ttl,omitempty"` } // ProxyNames returns the names of all proxies carried in the request. @@ -91,6 +80,7 @@ func (r *discoveryRequest) TrackProxies() []track.Proxy { Name: p.Metadata.Name, Group: p.ProxyGroupID, Generation: p.ProxyGroupGeneration, + TTL: p.TTL, }) } return tp @@ -115,3 +105,208 @@ func (r *discoveryRequest) String() string { b.WriteRune(']') return b.String() } + +// proxyDiscoverySubscriber is a subscriber to proxy discovery events. +type proxyDiscoverySubscriber struct { + pb *proxyDiscoveryPublisher + notify chan struct{} + version uint64 +} + +// Wait returns a channel which is notified when there is an event to fetch. +func (s *proxyDiscoverySubscriber) Wait() <-chan struct{} { + return s.notify +} + +// Get returns each [discoveryProxy] fetches the latest set of proxies. If compaction +// is enabled ony the changes since the last fetch is returned. +func (s *proxyDiscoverySubscriber) Get() []discoveryProxy { + return s.pb.get(s, discoGetParams{sinceLastVersion: true}) +} + +// GetAll returns all [discoveryProxy]s. +func (s *proxyDiscoverySubscriber) GetAll() []discoveryProxy { + return s.pb.get(s, discoGetParams{sinceLastVersion: false}) +} + +// proxyDiscoveryPublisher broadcasts proxy watch events to many subscribers. +type proxyDiscoveryPublisher struct { + ctx context.Context + cancel func() + watcher *services.GenericWatcher[types.Server, readonly.Server] + compact bool + log *slog.Logger + + // state points to the current immutable discovery snapshot. + // writers must publish a fully rebuilt discoveryState and never mutate it. + state atomic.Pointer[discoveryState] + // alwaysClosed is a channel sent to subscribers when they are first created + // to notify them immediately to fetch the latest state. it is always closed + // and never replaced. + alwaysClosed chan struct{} +} + +// discoveryState is a single published discovery snapshot. +// its fields are read without locking and must be treated as immutable. +type discoveryState struct { + // proxies is the lastest set of servers received from the watcher. + proxies []types.Server + // proxy version associates a proxy by name with the version at which its + // expiry last changed. + versions map[string]proxyversion + // version is incremented each time a proxy watch event is receieved. + version uint64 + // notify is closed when this state is replaced to notify subscribers. + notify chan struct{} +} + +type proxyversion struct { + version uint64 + expiry time.Time + updated time.Time +} + +// newProxyDiscoveryPublisher constructs a [proxyDiscoveryPublisher] using the given [services.GenericWatcher]. +func newProxyDiscoveryPublisher(ctx context.Context, watcher *services.GenericWatcher[types.Server, readonly.Server], logger *slog.Logger) *proxyDiscoveryPublisher { + ctx, cancel := context.WithCancel(ctx) + v := os.Getenv("TELEPORT_UNSTABLE_PROXY_COMPACT_DISCOVERY") + compact, _ := strconv.ParseBool(v) + + dp := &proxyDiscoveryPublisher{ + ctx: ctx, + cancel: cancel, + watcher: watcher, + state: atomic.Pointer[discoveryState]{}, + alwaysClosed: make(chan struct{}), + compact: compact, + log: logger, + } + close(dp.alwaysClosed) + dp.state.Store(&discoveryState{ + versions: map[string]proxyversion{}, + notify: make(chan struct{}), + }) + go dp.run() + return dp +} + +func (dp *proxyDiscoveryPublisher) discoFromServer(s types.Server, ttl time.Duration) discoveryProxy { + p := discoveryProxy{ + Version: types.V2, + } + p.Metadata.Name = s.GetName() + p.TTL = ttl + p.ProxyGroupID, _ = s.GetLabel(types.ProxyGroupIDLabel) + proxyGroupGeneration, ok := s.GetLabel(types.ProxyGroupGenerationLabel) + if !ok { + return p + } + + var err error + p.ProxyGroupGeneration, err = strconv.ParseUint(proxyGroupGeneration, 10, 64) + if err != nil { + // ParseUint can return the maximum uint64 on ErrRange + p.ProxyGroupGeneration = 0 + dp.log.DebugContext(dp.ctx, "Failed to parse proxy group generation", "error", err, "value", proxyGroupGeneration) + } + return p +} + +func (dp *proxyDiscoveryPublisher) run() { + for { + select { + case <-dp.ctx.Done(): + return + case servers, ok := <-dp.watcher.ResourcesC: + if !ok { + dp.log.WarnContext(dp.ctx, "Proxy discovery watcher closed unexpectedly") + return + } + now := dp.watcher.Clock.Now() + + prev := dp.state.Load() + next := &discoveryState{ + proxies: servers, + versions: make(map[string]proxyversion, len(servers)), + version: prev.version + 1, + notify: make(chan struct{}), + } + for _, server := range servers { + nextpv := proxyversion{ + version: next.version, + expiry: server.Expiry(), + updated: now, + } + prevpv, ok := prev.versions[server.GetName()] + if ok && prevpv.expiry.Equal(nextpv.expiry) { + // Preserve previous version and updated timestamp if the + // server hasn't heartbeated. + nextpv.version = prevpv.version + nextpv.updated = prevpv.updated + } + next.versions[server.GetName()] = nextpv + } + dp.state.Store(next) + close(prev.notify) + } + } +} + +// Subscribe returns a new [proxyDiscoverySubscriber] for receiving proxy event updates. +func (dp *proxyDiscoveryPublisher) Subscribe() *proxyDiscoverySubscriber { + s := &proxyDiscoverySubscriber{ + pb: dp, + notify: dp.alwaysClosed, + } + state := dp.state.Load() + if state.version == 0 { + s.notify = state.notify + } + return s +} + +// Close cleans up resources allocated by a [proxyDiscoveryPublisher]. +func (dp *proxyDiscoveryPublisher) Close() { + dp.cancel() + dp.watcher.Close() +} + +// discoGetParams contains parameters for [proxyDiscoveryPublisher.get]. +type discoGetParams struct { + // sinceLastVersion indicates that only proxies that have been fetched since + // the last get by the subscriber will be returned. + sinceLastVersion bool +} + +// get fetches the latest set of [discoveryProxy]s. +func (dp *proxyDiscoveryPublisher) get(sub *proxyDiscoverySubscriber, params discoGetParams) []discoveryProxy { + compact := dp.compact + state := dp.state.Load() + + var ttl time.Duration + if compact { + ttl = defaults.ProxyAnnounceTTL() + } + + now := dp.watcher.Clock.Now() + disco := make([]discoveryProxy, 0, len(state.proxies)) + prevVersion := sub.version + sub.version = state.version + sub.notify = state.notify + for _, proxy := range state.proxies { + if compact { + pv, ok := state.versions[proxy.GetName()] + if !ok { + continue + } + if pv.updated.Add(defaults.ProxyAnnounceTTL()).Before(now) { + continue + } + if params.sinceLastVersion && prevVersion >= pv.version { + continue + } + } + disco = append(disco, dp.discoFromServer(proxy, ttl)) + } + return disco +} diff --git a/lib/reversetunnel/discovery_test.go b/lib/reversetunnel/discovery_test.go index 08832fe48bd..4880172c01e 100644 --- a/lib/reversetunnel/discovery_test.go +++ b/lib/reversetunnel/discovery_test.go @@ -19,17 +19,24 @@ package reversetunnel import ( + "context" "encoding/json" + "slices" "testing" + "testing/synctest" + "time" "github.com/google/go-cmp/cmp" "github.com/google/uuid" "github.com/gravitational/trace" + "github.com/jonboulle/clockwork" "github.com/stretchr/testify/require" + "github.com/gravitational/teleport/api/defaults" "github.com/gravitational/teleport/api/types" "github.com/gravitational/teleport/lib/services" "github.com/gravitational/teleport/lib/utils" + "github.com/gravitational/teleport/lib/utils/log/logtest" ) // discoveryRequestRaw is the legacy type that was used @@ -112,7 +119,10 @@ func TestDiscoveryRequestMarshalling(t *testing.T) { // create the request var req discoveryRequest - req.SetProxies(proxies) + dp := &proxyDiscoveryPublisher{} + for _, proxy := range proxies { + req.Proxies = append(req.Proxies, dp.discoFromServer(proxy, 0)) + } // test marshaling the request with the legacy mechanism and unmarshaling // with the new mechanism @@ -160,3 +170,220 @@ func TestDiscoveryRequestMarshalling(t *testing.T) { require.Empty(t, cmp.Diff(req.ProxyNames(), got)) }) } + +func TestTrackProxiesPreservesTTL(t *testing.T) { + req := discoveryRequest{ + Proxies: []discoveryProxy{ + { + Version: types.V2, + ProxyGroupID: "group-a", + ProxyGroupGeneration: 7, + TTL: 42 * time.Second, + }, + }, + } + req.Proxies[0].Metadata.Name = "proxy-a" + + got := req.TrackProxies() + require.Len(t, got, 1) + require.Equal(t, "proxy-a", got[0].Name) + require.Equal(t, "group-a", got[0].Group) + require.Equal(t, uint64(7), got[0].Generation) + require.Equal(t, 42*time.Second, got[0].TTL) +} + +func TestProxyPubSub(t *testing.T) { + type proxy struct { + name string + expiryOffset time.Duration + } + type update struct { + expiryAdvance time.Duration + update []proxy + wantGet []string + wantAll []string + } + tests := []struct { + name string + compact bool + initial []proxy + wantInitialGet []string + updates []update + }{ + { + name: "initial state triggers wait and get", + compact: false, + initial: []proxy{ + {name: "proxy-a"}, + {name: "proxy-b"}, + }, + wantInitialGet: []string{"proxy-a", "proxy-b"}, + }, + { + name: "default get returns full proxy set", + compact: false, + initial: []proxy{ + {name: "proxy-a"}, + {name: "proxy-b"}, + }, + wantInitialGet: []string{"proxy-a", "proxy-b"}, + updates: []update{ + { + update: []proxy{ + {name: "proxy-a"}, + {name: "proxy-b"}, + }, + wantGet: []string{"proxy-a", "proxy-b"}, + wantAll: []string{"proxy-a", "proxy-b"}, + }, + }, + }, + { + name: "compact get returns only updated proxies", + compact: true, + initial: []proxy{ + {name: "proxy-a"}, + {name: "proxy-b"}, + }, + wantInitialGet: []string{"proxy-a", "proxy-b"}, + updates: []update{ + { + update: []proxy{ + {name: "proxy-a", expiryOffset: time.Second}, + {name: "proxy-b"}, + }, + wantGet: []string{"proxy-a"}, + wantAll: []string{"proxy-a", "proxy-b"}, + }, + }, + }, + { + name: "compact get removes expired proxies", + compact: true, + initial: []proxy{ + {name: "proxy-a"}, + {name: "proxy-b"}, + }, + wantInitialGet: []string{"proxy-a", "proxy-b"}, + updates: []update{ + { + expiryAdvance: defaults.ProxyAnnounceTTL() + time.Second, + update: []proxy{ + {name: "proxy-a", expiryOffset: defaults.ProxyAnnounceTTL() + time.Second}, + {name: "proxy-b"}, + }, + wantGet: []string{"proxy-a"}, + wantAll: []string{"proxy-a"}, + }, + }, + }, + { + name: "default get ignores expiry and version", + compact: false, + initial: []proxy{ + {name: "proxy-a"}, + {name: "proxy-b"}, + }, + wantInitialGet: []string{"proxy-a", "proxy-b"}, + updates: []update{ + { + expiryAdvance: defaults.ProxyAnnounceTTL() + time.Second, + update: []proxy{ + {name: "proxy-a", expiryOffset: defaults.ProxyAnnounceTTL() + time.Second}, + {name: "proxy-b"}, + }, + wantGet: []string{"proxy-a", "proxy-b"}, + wantAll: []string{"proxy-a", "proxy-b"}, + }, + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + clock := clockwork.NewFakeClock() + baseTime := clock.Now() + mkServers := func(states []proxy) []types.Server { + servers := make([]types.Server, 0, len(states)) + for _, state := range states { + s, err := types.NewServer(state.name, types.KindProxy, types.ServerSpecV2{}) + require.NoError(t, err) + if state.expiryOffset != 0 { + s.SetExpiry(baseTime.Add(state.expiryOffset)) + } else { + s.SetExpiry(baseTime) + } + servers = append(servers, s) + } + return servers + } + proxyNames := func(proxies []discoveryProxy) []string { + names := make([]string, 0, len(proxies)) + for _, proxy := range proxies { + names = append(names, proxy.Metadata.Name) + } + slices.Sort(names) + return names + } + stateNames := func(states []proxy) []string { + names := make([]string, 0, len(states)) + for _, state := range states { + names = append(names, state.name) + } + slices.Sort(names) + return names + } + + client := &mockLocalClusterClient{ + proxies: mkServers(tt.initial), + } + watcher, err := services.NewProxyWatcher(ctx, services.ProxyWatcherConfig{ + ResourceWatcherConfig: services.ResourceWatcherConfig{ + Component: "test", + Clock: clock, + Logger: logtest.NewLogger(), + Client: client, + }, + ProxyGetter: client, + ProxiesC: make(chan []types.Server, len(tt.updates)+1), + }) + require.NoError(t, err) + require.NoError(t, watcher.WaitInitialization()) + + pb := newProxyDiscoveryPublisher(ctx, watcher, logtest.NewLogger()) + pb.compact = tt.compact + t.Cleanup(pb.Close) + + sub := pb.Subscribe() + + select { + case <-sub.Wait(): + case <-time.After(5 * time.Second): + t.Fatal("timed out waiting for initial proxy state") + } + + require.ElementsMatch(t, tt.wantInitialGet, proxyNames(sub.Get())) + require.ElementsMatch(t, stateNames(tt.initial), proxyNames(sub.GetAll())) + + for _, update := range tt.updates { + clock.Advance(update.expiryAdvance) + if len(update.update) > 0 { + watcher.ResourcesC <- mkServers(update.update) + select { + case <-sub.Wait(): + case <-time.After(5 * time.Second): + t.Fatal("timed out waiting for proxy update") + } + } + + require.ElementsMatch(t, update.wantGet, proxyNames(sub.Get())) + require.ElementsMatch(t, update.wantAll, proxyNames(sub.GetAll())) + } + }) + }) + } +} diff --git a/lib/reversetunnel/leaf_cluster.go b/lib/reversetunnel/leaf_cluster.go index f687d173ced..808bc71c0a2 100644 --- a/lib/reversetunnel/leaf_cluster.go +++ b/lib/reversetunnel/leaf_cluster.go @@ -312,14 +312,18 @@ func (s *leafCluster) addConn(conn net.Conn, sconn ssh.Conn) (*remoteConn, error s.Lock() defer s.Unlock() - rconn := newRemoteConn(&connConfig{ - conn: conn, - sconn: sconn, - tunnelType: string(types.ProxyTunnel), - proxyName: s.connInfo.GetProxyName(), - clusterName: s.domainName, - offlineThreshold: s.offlineThreshold, + rconn, err := newRemoteConn(&connConfig{ + conn: conn, + sconn: sconn, + tunnelType: string(types.ProxyTunnel), + proxyName: s.connInfo.GetProxyName(), + clusterName: s.domainName, + offlineThreshold: s.offlineThreshold, + proxyDiscoverySubscriber: s.srv.proxyDiscoveryPublisher.Subscribe(), }) + if err != nil { + return nil, trace.Wrap(err) + } s.connections = append(s.connections, rconn) s.lastUsed = 0 @@ -403,17 +407,6 @@ func (s *leafCluster) deleteConnectionRecord() { } } -// fanOutProxies is a non-blocking call that puts the new proxies -// list so that remote connection can notify the remote agent -// about the list update -func (s *leafCluster) fanOutProxies(proxies []types.Server) { - s.Lock() - defer s.Unlock() - for _, conn := range s.connections { - conn.updateProxies(proxies) - } -} - // handleHeartbeat receives heartbeat messages from the connected agent // if the agent has missed several heartbeats in a row, Proxy marks // the connection as invalid. @@ -432,7 +425,6 @@ func (s *leafCluster) handleHeartbeat(ctx context.Context, conn *remoteConn, ch }() } - firstHeartbeat := true proxyResyncTicker := s.clock.NewTicker(s.proxySyncInterval) defer func() { proxyResyncTicker.Stop() @@ -456,21 +448,19 @@ func (s *leafCluster) handleHeartbeat(ctx context.Context, conn *remoteConn, ch return case <-proxyResyncTicker.Chan(): var req discoveryRequest - proxies, err := s.srv.proxyWatcher.CurrentResources(s.srv.ctx) - if err != nil { - logger.WarnContext(ctx, "Failed to get proxy set", "error", err) - } - req.SetProxies(proxies) + req.Proxies = conn.proxyDiscoverySubscriber.GetAll() if err := conn.sendDiscoveryRequest(ctx, req); err != nil { logger.DebugContext(ctx, "Marking connection invalid on error", "error", err) conn.markInvalid(err) return } - case proxies := <-conn.newProxiesC: + case <-conn.proxyDiscoverySubscriber.Wait(): var req discoveryRequest - req.SetProxies(proxies) - + req.Proxies = conn.proxyDiscoverySubscriber.Get() + if len(req.Proxies) == 0 { + continue + } if err := conn.sendDiscoveryRequest(ctx, req); err != nil { logger.DebugContext(ctx, "Marking connection invalid on error", "error", err) conn.markInvalid(err) @@ -486,18 +476,6 @@ func (s *leafCluster) handleHeartbeat(ctx context.Context, conn *remoteConn, ch } return } - if firstHeartbeat { - // as soon as the agent connects and sends a first heartbeat - // send it the list of current proxies back - proxies, err := s.srv.proxyWatcher.CurrentResources(ctx) - if err != nil { - logger.WarnContext(ctx, "Failed to get proxy set", "error", err) - } - if len(proxies) > 0 { - conn.updateProxies(proxies) - } - firstHeartbeat = false - } var timeSent time.Time var roundtrip time.Duration if req.Payload != nil { diff --git a/lib/reversetunnel/local_cluster.go b/lib/reversetunnel/local_cluster.go index 6e289c82c85..b40ecfb8450 100644 --- a/lib/reversetunnel/local_cluster.go +++ b/lib/reversetunnel/local_cluster.go @@ -57,15 +57,26 @@ const ( // periodicFunctionInterval is the interval at which periodic stats are calculated. periodicFunctionInterval = 3 * time.Minute - // proxySyncInterval is the interval at which the current proxies are synchronized to - // connected agents via a discovery request. It is a function of track.DefaultProxyExpiry - // to ensure that the proxies are always synced before the tracker expiry. - proxySyncInterval = track.DefaultProxyExpiry * 2 / 3 - // missedHeartBeatThreshold is the number of missed heart beats needed to terminate a connection. missedHeartBeatThreshold = 3 ) +// proxySyncInterval is the interval at which the current proxies are synchronized to +// connected agents via a discovery request. It is a function of track.DefaultProxyExpiry +// to ensure that the proxies are always synced before the tracker expiry. +// +// With support of proxy discovery TTLs the tracker may expire proxies sooner +// than the default proxy expiry. In this case a lower sync interval is used +// to ensure that proxies are still always synced before expiry. +var proxySyncInterval = func() time.Duration { + defaultSyncInterval := track.DefaultProxyExpiry * 2 / 3 + calculatedSyncInteval := apidefaults.ProxyAnnounceTTL() * 2 / 3 + if calculatedSyncInteval < defaultSyncInterval { + return calculatedSyncInteval + } + return defaultSyncInterval +}() + // withPeriodicFunctionInterval adjusts the periodic function interval func withPeriodicFunctionInterval(interval time.Duration) func(cluster *localCluster) { return func(cluster *localCluster) { @@ -743,16 +754,20 @@ func (s *localCluster) addConn(nodeID, scope string, connType types.TunnelType, s.remoteConnsMtx.Lock() defer s.remoteConnsMtx.Unlock() - rconn := newRemoteConn(&connConfig{ - conn: conn, - sconn: sconn, - tunnelType: string(connType), - proxyName: s.srv.ID, - clusterName: s.domainName, - nodeID: nodeID, - scope: scope, - offlineThreshold: s.offlineThreshold, + rconn, err := newRemoteConn(&connConfig{ + conn: conn, + sconn: sconn, + tunnelType: string(connType), + proxyName: s.srv.ID, + clusterName: s.domainName, + nodeID: nodeID, + scope: scope, + offlineThreshold: s.offlineThreshold, + proxyDiscoverySubscriber: s.srv.proxyDiscoveryPublisher.Subscribe(), }) + if err != nil { + return nil, trace.Wrap(err) + } key := connKey{ uuid: nodeID, connType: connType, @@ -763,20 +778,6 @@ func (s *localCluster) addConn(nodeID, scope string, connType types.TunnelType, return rconn, nil } -// fanOutProxies is a non-blocking call that puts the new proxies -// list so that remote connection can notify the remote agent -// about the list update -func (s *localCluster) fanOutProxies(proxies []types.Server) { - s.remoteConnsMtx.Lock() - defer s.remoteConnsMtx.Unlock() - - for _, conns := range s.remoteConns { - for _, conn := range conns { - conn.updateProxies(proxies) - } - } -} - // handleHeartbeat receives heartbeat messages from the connected agent // if the agent has missed several heartbeats in a row, Proxy marks // the connection as invalid. @@ -795,8 +796,8 @@ func (s *localCluster) handleHeartbeat(ctx context.Context, rconn *remoteConn, c "addr", logutils.StringerAttr(rconn.conn.RemoteAddr()), ) - firstHeartbeat := true proxyResyncTicker := s.clock.NewTicker(s.proxySyncInterval) + reverseSSHTunnels.WithLabelValues(rconn.tunnelType).Inc() defer func() { proxyResyncTicker.Stop() logger.WarnContext(ctx, "Closing remote connection to agent") @@ -804,9 +805,7 @@ func (s *localCluster) handleHeartbeat(ctx context.Context, rconn *remoteConn, c if err := rconn.Close(); err != nil && !utils.IsOKNetworkError(err) { logger.WarnContext(ctx, "Failed to close remote connection", "error", err) } - if !firstHeartbeat { - reverseSSHTunnels.WithLabelValues(rconn.tunnelType).Dec() - } + reverseSSHTunnels.WithLabelValues(rconn.tunnelType).Dec() }() offlineThresholdTimer := s.clock.NewTimer(s.offlineThreshold) @@ -818,21 +817,18 @@ func (s *localCluster) handleHeartbeat(ctx context.Context, rconn *remoteConn, c return case <-proxyResyncTicker.Chan(): var req discoveryRequest - proxies, err := s.srv.proxyWatcher.CurrentResources(ctx) - if err != nil { - logger.WarnContext(ctx, "Failed to get proxy set", "error", err) - } - req.SetProxies(proxies) - + req.Proxies = rconn.proxyDiscoverySubscriber.GetAll() if err := rconn.sendDiscoveryRequest(ctx, req); err != nil { logger.DebugContext(ctx, "Marking connection invalid on error", "error", err) rconn.markInvalid(err) return } - case proxies := <-rconn.newProxiesC: + case <-rconn.proxyDiscoverySubscriber.Wait(): var req discoveryRequest - req.SetProxies(proxies) - + req.Proxies = rconn.proxyDiscoverySubscriber.Get() + if len(req.Proxies) == 0 { + continue + } if err := rconn.sendDiscoveryRequest(ctx, req); err != nil { logger.DebugContext(ctx, "Failed to send discovery request to agent", "error", err) rconn.markInvalid(err) @@ -844,19 +840,6 @@ func (s *localCluster) handleHeartbeat(ctx context.Context, rconn *remoteConn, c rconn.markInvalid(trace.ConnectionProblem(nil, "agent disconnected")) return } - if firstHeartbeat { - // as soon as the agent connects and sends a first heartbeat - // send it the list of current proxies back - proxies, err := s.srv.proxyWatcher.CurrentResources(s.srv.ctx) - if err != nil { - logger.WarnContext(ctx, "Failed to get proxy set", "error", err) - } - if len(proxies) > 0 { - rconn.updateProxies(proxies) - } - reverseSSHTunnels.WithLabelValues(rconn.tunnelType).Inc() - firstHeartbeat = false - } var timeSent time.Time var roundtrip time.Duration if req.Payload != nil { diff --git a/lib/reversetunnel/local_cluster_test.go b/lib/reversetunnel/local_cluster_test.go index f89bb44d9cb..0aea240ad3f 100644 --- a/lib/reversetunnel/local_cluster_test.go +++ b/lib/reversetunnel/local_cluster_test.go @@ -75,12 +75,12 @@ func TestRemoteConnCleanup(t *testing.T) { // set up the cluster srv := &server{ - ctx: ctx, - Config: Config{Clock: clock}, - localAuthClient: &mockLocalClusterClient{}, - logger: logtest.NewLogger(), - offlineThreshold: time.Second, - proxyWatcher: watcher, + ctx: ctx, + Config: Config{Clock: clock}, + localAuthClient: &mockLocalClusterClient{}, + logger: logtest.NewLogger(), + offlineThreshold: time.Second, + proxyDiscoveryPublisher: newProxyDiscoveryPublisher(ctx, watcher, logtest.NewLogger()), } cluster, err := newLocalCluster(srv, "clustername", nil, @@ -147,11 +147,26 @@ func TestRemoteConnCleanup(t *testing.T) { func TestLocalClusterOverlap(t *testing.T) { t.Parallel() + ctx := t.Context() + clock := clockwork.NewFakeClock() + clt := &mockLocalClusterClient{} + watcher, err := services.NewProxyWatcher(ctx, services.ProxyWatcherConfig{ + ResourceWatcherConfig: services.ResourceWatcherConfig{ + Component: "test", + Logger: logtest.NewLogger(), + Clock: clock, + Client: clt, + }, + ProxyGetter: clt, + ProxiesC: make(chan []types.Server, 2), + }) + require.NoError(t, err) srv := &server{ - Config: Config{Clock: clockwork.NewFakeClock()}, - ctx: context.Background(), - localAuthClient: &mockLocalClusterClient{}, + Config: Config{Clock: clockwork.NewFakeClock()}, + ctx: context.Background(), + localAuthClient: &mockLocalClusterClient{}, + proxyDiscoveryPublisher: newProxyDiscoveryPublisher(ctx, watcher, logtest.NewLogger()), } cluster, err := newLocalCluster(srv, "clustername", nil, @@ -270,12 +285,12 @@ func TestProxyResync(t *testing.T) { // set up the cluster srv := &server{ - ctx: ctx, - Config: Config{Clock: clock}, - localAuthClient: &mockLocalClusterClient{}, - logger: logtest.NewLogger(), - offlineThreshold: 24 * time.Hour, - proxyWatcher: watcher, + ctx: ctx, + Config: Config{Clock: clock}, + localAuthClient: &mockLocalClusterClient{}, + logger: logtest.NewLogger(), + offlineThreshold: 24 * time.Hour, + proxyDiscoveryPublisher: newProxyDiscoveryPublisher(ctx, watcher, logtest.NewLogger()), } cluster, err := newLocalCluster(srv, "clustername", nil, withProxySyncInterval(time.Second), diff --git a/lib/reversetunnel/srv.go b/lib/reversetunnel/srv.go index cb8b6faeccb..c7acec29430 100644 --- a/lib/reversetunnel/srv.go +++ b/lib/reversetunnel/srv.go @@ -117,10 +117,6 @@ type server struct { // logger specifies the logger logger *slog.Logger - // proxyWatcher monitors changes to the proxies - // and broadcasts updates - proxyWatcher *services.GenericWatcher[types.Server, readonly.Server] - // offlineThreshold is how long to wait for a keep alive message before // marking a reverse tunnel connection as invalid. offlineThreshold time.Duration @@ -130,6 +126,9 @@ type server struct { // gitKeyManager manages keys for git proxies. gitKeyManager *git.KeyManager + + // proxyDiscoveryPublisher publishes proxy discovery events to subscribers. + proxyDiscoveryPublisher *proxyDiscoveryPublisher } // EICESigner is a function that is used to obatin an [ssh.Signer] for an EICE instance. The @@ -366,18 +365,18 @@ func NewServer(cfg Config) (reversetunnelclient.Server, error) { } srv := &server{ - Config: cfg, - localAuthClient: cfg.LocalAuthClient, - localAccessPoint: cfg.LocalAccessPoint, - limiter: cfg.Limiter, - ctx: ctx, - cancel: cancel, - proxyWatcher: proxyWatcher, - expectedLeafClusters: make(map[string]*expectedLeafClusters), - logger: cfg.Logger, - offlineThreshold: offlineThreshold, - proxySigner: cfg.PROXYSigner, - gitKeyManager: gitKeyManager, + Config: cfg, + localAuthClient: cfg.LocalAuthClient, + localAccessPoint: cfg.LocalAccessPoint, + limiter: cfg.Limiter, + ctx: ctx, + cancel: cancel, + expectedLeafClusters: make(map[string]*expectedLeafClusters), + logger: cfg.Logger, + offlineThreshold: offlineThreshold, + proxySigner: cfg.PROXYSigner, + gitKeyManager: gitKeyManager, + proxyDiscoveryPublisher: newProxyDiscoveryPublisher(ctx, proxyWatcher, cfg.Logger), } srv.localCluster, err = newLocalCluster(srv, cfg.ClusterName, cfg.LocalAuthAddresses) @@ -464,9 +463,6 @@ func (s *server) periodicFunctions() { case <-s.ctx.Done(): s.logger.DebugContext(s.ctx, "Closing") return - // Proxies have been updated, notify connected agents about the update. - case proxies := <-s.proxyWatcher.ResourcesC: - s.fanOutProxies(proxies) case <-ticker.C: if err := s.fetchExpectedLeafClusters(); err != nil { s.logger.WarnContext(s.ctx, "Failed to fetch expected leaf clusters", "error", err) @@ -667,7 +663,7 @@ func (s *server) Start() error { func (s *server) Close() error { s.cancel() - s.proxyWatcher.Close() + s.proxyDiscoveryPublisher.Close() return s.srv.Close() } @@ -693,7 +689,7 @@ func (s *server) DrainConnections(ctx context.Context) error { func (s *server) Shutdown(ctx context.Context) error { err := s.srv.Shutdown(ctx) - s.proxyWatcher.Close() + s.proxyDiscoveryPublisher.Close() s.cancel() return trace.Wrap(err) @@ -1199,18 +1195,6 @@ func (s *server) onClusterTunnelClose(cluster clusterCloser) error { return trace.NotFound("cluster %q is not found", cluster.GetName()) } -// fanOutProxies is a non-blocking call that updated the watches proxies -// list and notifies all clusters about the proxy list change -func (s *server) fanOutProxies(proxies []types.Server) { - s.Lock() - defer s.Unlock() - s.localCluster.fanOutProxies(proxies) - - for _, cluster := range s.leafClusters { - cluster.fanOutProxies(proxies) - } -} - func (s *server) rejectRequest(ch ssh.NewChannel, reason ssh.RejectionReason, msg string) { if err := ch.Reject(reason, msg); err != nil { s.logger.WarnContext(s.ctx, "Failed rejecting new channel request", "error", err) diff --git a/lib/reversetunnel/track/tracker.go b/lib/reversetunnel/track/tracker.go index 8ed1c05fd49..2cce279b669 100644 --- a/lib/reversetunnel/track/tracker.go +++ b/lib/reversetunnel/track/tracker.go @@ -92,6 +92,7 @@ type Proxy struct { Name string Group string Generation uint64 + TTL time.Duration expiry time.Time }