Upsert ServerInfos from discovery service (#27475)

This change adds the `labelReconciler` to the discovery service, which periodically
reconciles the labels it receives from discovered EC2 instances with the labels of the
corresponding SSH servers stored in the auth server.
This commit is contained in:
Andrew Burke
2023-07-12 23:56:32 +00:00
committed by GitHub
parent f6e9ca269e
commit fbba4c2bfa
12 changed files with 2044 additions and 1401 deletions
@@ -5850,6 +5850,8 @@ message ServerInfoSpecV1 {
}
// AWS matches an EC2 instance.
AWSInfo AWS = 1 [(gogoproto.jsontag) = "aws,omitempty"];
// NewLabels is the set of labels to add to nodes matching this ServerInfo.
map<string, string> NewLabels = 2 [(gogoproto.jsontag) = "new_labels,omitempty"];
}
// JamfSpecV1 is the base configuration for the Jamf MDM service.
+14
View File
@@ -28,6 +28,10 @@ import (
type ServerInfo interface {
// ResourceWithLabels provides common resource headers
ResourceWithLabels
// GetNewLabels gets the labels to apply to matched Nodes.
GetNewLabels() map[string]string
// SetNewLabels sets the labels to apply to matched Nodes.
SetNewLabels(map[string]string)
}
// NewServerInfo creates an instance of ServerInfo.
@@ -141,6 +145,16 @@ func (s *ServerInfoV1) MatchSearch(searchValues []string) bool {
return MatchSearch(fieldVals, searchValues, nil)
}
// GetNewLabels gets the labels to apply to matched Nodes.
func (s *ServerInfoV1) GetNewLabels() map[string]string {
return s.Spec.NewLabels
}
// SetNewLabels sets the labels to apply to matched Nodes.
func (s *ServerInfoV1) SetNewLabels(labels map[string]string) {
s.Spec.NewLabels = labels
}
func (s *ServerInfoV1) setStaticFields() {
s.Kind = KindServerInfo
s.Version = V1
+1554 -1395
View File
File diff suppressed because it is too large Load Diff
+7
View File
@@ -720,6 +720,8 @@ type DiscoveryAccessPoint interface {
UpdateDatabase(ctx context.Context, database types.Database) error
// DeleteDatabase deletes a database resource.
DeleteDatabase(ctx context.Context, name string) error
// UpsertServerInfo upserts a server info resource.
UpsertServerInfo(ctx context.Context, si types.ServerInfo) error
}
// ReadOktaAccessPoint is a read only API interface to be
@@ -1212,6 +1214,11 @@ func (w *DiscoveryWrapper) DeleteDatabase(ctx context.Context, name string) erro
return w.NoCache.DeleteDatabase(ctx, name)
}
// UpsertServerInfo upserts a server info resource.
func (w *DiscoveryWrapper) UpsertServerInfo(ctx context.Context, si types.ServerInfo) error {
return w.NoCache.UpsertServerInfo(ctx, si)
}
// Close closes all associated resources
func (w *DiscoveryWrapper) Close() error {
err := w.NoCache.Close()
+1
View File
@@ -863,6 +863,7 @@ func definitionForBuiltinRole(clusterName string, recConfig types.SessionRecordi
types.NewRule(types.KindNode, services.RO()),
types.NewRule(types.KindKubernetesCluster, services.RW()),
types.NewRule(types.KindDatabase, services.RW()),
types.NewRule(types.KindServerInfo, services.RW()),
},
// wildcard any cluster available.
KubernetesLabels: types.Labels{types.Wildcard: []string{types.Wildcard}},
+3 -3
View File
@@ -150,7 +150,7 @@ func (l *CloudImporter) Sync(ctx context.Context) error {
l.Log.Debugf("Skipping cloud tag %q, not a valid label key.", key)
continue
}
m[formatKey(l.namespace, key)] = value
m[FormatCloudLabelKey(l.namespace, key)] = value
}
l.muLabels.Lock()
@@ -180,7 +180,7 @@ func (l *CloudImporter) periodicUpdateLabels(ctx context.Context) {
}
}
// formatKey formats label keys coming from a cloud instance.
func formatKey(namespace, key string) string {
// FormatCloudLabelKey formats label keys coming from a cloud instance.
func FormatCloudLabelKey(namespace, key string) string {
return fmt.Sprintf("%s/%s", namespace, key)
}
+19
View File
@@ -127,6 +127,9 @@ type Server struct {
databaseFetchers []common.Fetcher
// caRotationCh receives nodes that need to have their CAs rotated.
caRotationCh chan []types.Server
// reconciler periodically reconciles the labels of discovered instances
// with the auth server.
reconciler *labelReconciler
}
// New initializes a discovery Server
@@ -181,6 +184,14 @@ func (s *Server) initAWSWatchers(matchers []types.AWSMatcher) error {
s.ec2Installer = server.NewSSMInstaller(server.SSMInstallerConfig{
Emitter: s.Emitter,
})
lr, err := newLabelReconciler(&labelReconcilerConfig{
log: s.Log,
accessPoint: s.AccessPoint,
})
if err != nil {
return trace.Wrap(err)
}
s.reconciler = lr
}
// Add database fetchers.
@@ -391,6 +402,13 @@ func (s *Server) handleEC2Instances(instances *server.EC2Instances) error {
if err != nil {
return trace.Wrap(err)
}
serverInfos, err := instances.ServerInfos()
if err != nil {
return trace.Wrap(err)
}
s.reconciler.queueServerInfos(serverInfos)
// instances.Rotation is true whenever the instances received need
// to be rotated, we don't want to filter out existing OpenSSH nodes as
// they all need to have the command run on them
@@ -607,6 +625,7 @@ func (s *Server) handleAzureDiscovery() {
func (s *Server) Start() error {
if s.ec2Watcher != nil {
go s.handleEC2Discovery()
go s.reconciler.run(s.ctx)
}
if s.azureWatcher != nil {
go s.handleAzureDiscovery()
+13
View File
@@ -1580,6 +1580,14 @@ type fakeAccessPoint struct {
auth.DiscoveryAccessPoint
updateKube bool
updateDatabase bool
upsertedServerInfos chan types.ServerInfo
}
func newFakeAccessPoint() *fakeAccessPoint {
return &fakeAccessPoint{
upsertedServerInfos: make(chan types.ServerInfo),
}
}
func (f *fakeAccessPoint) CreateDatabase(ctx context.Context, database types.Database) error {
@@ -1600,3 +1608,8 @@ func (f *fakeAccessPoint) UpdateKubernetesCluster(ctx context.Context, cluster t
f.updateKube = true
return nil
}
func (f *fakeAccessPoint) UpsertServerInfo(ctx context.Context, si types.ServerInfo) error {
f.upsertedServerInfos <- si
return nil
}
+154
View File
@@ -0,0 +1,154 @@
/*
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 discovery
import (
"context"
"sync"
"time"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
"github.com/sirupsen/logrus"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/utils/retryutils"
"github.com/gravitational/teleport/lib/utils"
)
// minBatchSize is the minimum batch size to send ServerInfos in for discovered
// instances.
const minBatchSize = 5
type serverInfoUpserter interface {
UpsertServerInfo(ctx context.Context, si types.ServerInfo) error
}
type labelReconcilerConfig struct {
clock clockwork.Clock
log logrus.FieldLogger
accessPoint serverInfoUpserter
}
func (c *labelReconcilerConfig) checkAndSetDefaults() error {
if c.accessPoint == nil {
return trace.BadParameter("missing parameter: accessPoint")
}
if c.clock == nil {
c.clock = clockwork.NewRealClock()
}
if c.log == nil {
c.log = logrus.New()
}
return nil
}
// labelReconciler periodically reconciles the labels of discovered instances
// with the auth server.
type labelReconciler struct {
cfg *labelReconcilerConfig
mu sync.Mutex
discoveredServers map[string]types.ServerInfo
serverInfoQueue []types.ServerInfo
lastBatchSize int
jitter retryutils.Jitter
}
func newLabelReconciler(cfg *labelReconcilerConfig) (*labelReconciler, error) {
if err := cfg.checkAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return &labelReconciler{
cfg: cfg,
discoveredServers: make(map[string]types.ServerInfo),
serverInfoQueue: make([]types.ServerInfo, 0, minBatchSize),
lastBatchSize: minBatchSize,
jitter: retryutils.NewSeventhJitter(),
}, nil
}
// getUpsertBatchSize calculates the size of batch to upsert ServerInfos in.
//
// Batches are sent once per second, and the goal is to upsert all ServerInfos
// within 15 minutes.
func getUpsertBatchSize(queueLen, lastBatchSize int) int {
batchSize := lastBatchSize
// Increase batch size so that all upserts can finish within 15 minutes.
if dynamicBatchSize := (queueLen / 900) + 1; dynamicBatchSize > batchSize {
batchSize = dynamicBatchSize
}
if batchSize < minBatchSize {
batchSize = minBatchSize
}
if batchSize > queueLen {
batchSize = queueLen
}
return batchSize
}
func (r *labelReconciler) run(ctx context.Context) {
for ctx.Err() == nil {
select {
case <-r.cfg.clock.After(time.Second):
r.mu.Lock()
if len(r.serverInfoQueue) == 0 {
r.mu.Unlock()
continue
}
batchSize := getUpsertBatchSize(len(r.serverInfoQueue), r.lastBatchSize)
r.lastBatchSize = batchSize
batch := r.serverInfoQueue[:batchSize]
r.serverInfoQueue = r.serverInfoQueue[batchSize:]
for _, si := range batch {
if err := r.cfg.accessPoint.UpsertServerInfo(ctx, si); err != nil {
r.cfg.log.WithError(err).Error("Failed to upsert server info.")
// Allow the server info to be queued again.
delete(r.discoveredServers, si.GetName())
}
}
r.mu.Unlock()
case <-ctx.Done():
return
}
}
}
// queueServerInfos queues a list of ServerInfos to be upserted.
func (r *labelReconciler) queueServerInfos(serverInfos []types.ServerInfo) {
r.mu.Lock()
defer r.mu.Unlock()
now := r.cfg.clock.Now()
for _, si := range serverInfos {
existingInfo, ok := r.discoveredServers[si.GetName()]
// ServerInfos should be upserted if
// - the instance is new
// - the instance's labels have changed
// - the existing ServerInfo will expire within 30 minutes
if !ok ||
!utils.StringMapsEqual(si.GetNewLabels(), existingInfo.GetNewLabels()) ||
existingInfo.Expiry().Before(now.Add(30*time.Minute)) {
si.SetExpiry(now.Add(r.jitter(90 * time.Minute)))
r.discoveredServers[si.GetName()] = si
r.serverInfoQueue = append(r.serverInfoQueue, si)
}
}
}
+207
View File
@@ -0,0 +1,207 @@
/*
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 discovery
import (
"context"
"fmt"
"testing"
"time"
"github.com/google/go-cmp/cmp"
"github.com/google/go-cmp/cmp/cmpopts"
"github.com/jonboulle/clockwork"
"github.com/stretchr/testify/require"
"github.com/gravitational/teleport/api/types"
)
func TestGetUpsertBatchSize(t *testing.T) {
t.Parallel()
tests := []struct {
name string
queueLen int
lastBatchSize int
expectedBatchSize int
}{
{
name: "small batches",
queueLen: 100,
lastBatchSize: 0,
expectedBatchSize: minBatchSize,
},
{
name: "continue previous batch size",
queueLen: 100,
lastBatchSize: 20,
expectedBatchSize: 20,
},
{
name: "large batches",
queueLen: 10000,
lastBatchSize: 0,
expectedBatchSize: 12,
},
{
name: "larger batch than previous",
queueLen: 10000,
lastBatchSize: 10,
expectedBatchSize: 12,
},
{
name: "last batch larger than queue size",
queueLen: 10,
lastBatchSize: 15,
expectedBatchSize: 10,
},
{
name: "short queue",
queueLen: 3,
lastBatchSize: 0,
expectedBatchSize: 3,
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
require.Equal(t, tc.expectedBatchSize, getUpsertBatchSize(tc.queueLen, tc.lastBatchSize))
})
}
}
func generateServerInfos(t *testing.T, n int) []types.ServerInfo {
serverInfos := make([]types.ServerInfo, 0, n)
for i := 0; i < n; i++ {
si, err := types.NewServerInfo(types.Metadata{
Name: fmt.Sprintf("instance-%d", i),
Labels: map[string]string{"foo": "bar"},
}, types.ServerInfoSpecV1{})
require.NoError(t, err)
serverInfos = append(serverInfos, si)
}
return serverInfos
}
func initLabelReconcilerForTests(t *testing.T, clock clockwork.Clock) (*labelReconciler, *fakeAccessPoint) {
ap := newFakeAccessPoint()
lr, err := newLabelReconciler(&labelReconcilerConfig{
clock: clock,
accessPoint: ap,
})
require.NoError(t, err)
return lr, ap
}
func TestLabelReconciler(t *testing.T) {
t.Parallel()
clock := clockwork.NewFakeClock()
lr, ap := initLabelReconcilerForTests(t, clock)
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
t.Cleanup(cancel)
go lr.run(ctx)
serverInfos := generateServerInfos(t, 25)
lr.queueServerInfos(serverInfos)
b := minBatchSize
for i := 0; i < 5; i++ {
clock.BlockUntil(1)
clock.Advance(time.Second)
var upsertedServerInfos []types.ServerInfo
outer:
for {
select {
case si := <-ap.upsertedServerInfos:
upsertedServerInfos = append(upsertedServerInfos, si)
case <-time.After(10 * time.Millisecond):
break outer
case <-ctx.Done():
require.Fail(t, "timed out waiting for server infos")
}
}
require.Len(t, upsertedServerInfos, b)
require.Equal(t, serverInfos[b*i:b*(i+1)], upsertedServerInfos)
}
}
func TestQueueServerInfos(t *testing.T) {
t.Parallel()
clock := clockwork.NewFakeClock()
nearFuture := clock.Now().Add(10 * time.Minute)
farFuture := clock.Now().Add(time.Hour)
newServerInfo := func(mod func(si types.ServerInfo)) types.ServerInfo {
defaultServerInfo, err := types.NewServerInfo(types.Metadata{
Name: "default",
Labels: map[string]string{"foo": "bar"},
Expires: &farFuture,
}, types.ServerInfoSpecV1{})
require.NoError(t, err)
if mod != nil {
mod(defaultServerInfo)
}
return defaultServerInfo
}
defaultServerInfos := []types.ServerInfo{newServerInfo(nil)}
tests := []struct {
name string
existingInfos []types.ServerInfo
newInfos []types.ServerInfo
expectedInfos []types.ServerInfo
}{
{
name: "new info",
newInfos: defaultServerInfos,
expectedInfos: defaultServerInfos,
},
{
name: "ignore existing info",
existingInfos: defaultServerInfos,
newInfos: defaultServerInfos,
expectedInfos: []types.ServerInfo{},
},
{
name: "re-queue updated labels",
existingInfos: []types.ServerInfo{newServerInfo(func(si types.ServerInfo) {
si.SetNewLabels(map[string]string{"foo": "baz"})
})},
newInfos: defaultServerInfos,
expectedInfos: defaultServerInfos,
},
{
name: "re-queue expiring soon",
existingInfos: []types.ServerInfo{newServerInfo(func(si types.ServerInfo) {
si.SetExpiry(nearFuture)
})},
newInfos: defaultServerInfos,
expectedInfos: defaultServerInfos,
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
lr, _ := initLabelReconcilerForTests(t, clock)
for _, si := range tc.existingInfos {
lr.discoveredServers[si.GetName()] = si
}
lr.queueServerInfos(tc.newInfos)
require.Empty(t, cmp.Diff(tc.expectedInfos, lr.serverInfoQueue,
cmpopts.IgnoreFields(types.Metadata{}, "Expires")))
})
}
}
+40 -3
View File
@@ -29,6 +29,7 @@ import (
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/cloud"
"github.com/gravitational/teleport/lib/labels"
"github.com/gravitational/teleport/lib/srv/db/common"
)
@@ -62,12 +63,20 @@ type EC2Instances struct {
// discovered.
type EC2Instance struct {
InstanceID string
Tags map[string]string
}
func toEC2Instance(inst *ec2.Instance) EC2Instance {
return EC2Instance{
InstanceID: aws.StringValue(inst.InstanceId),
func toEC2Instance(originalInst *ec2.Instance) EC2Instance {
inst := EC2Instance{
InstanceID: aws.StringValue(originalInst.InstanceId),
Tags: make(map[string]string, len(originalInst.Tags)),
}
for _, tag := range originalInst.Tags {
if key := aws.StringValue(tag.Key); key != "" {
inst.Tags[key] = aws.StringValue(tag.Value)
}
}
return inst
}
// ToEC2Instances converts aws []*ec2.Instance to []EC2Instance
@@ -81,6 +90,34 @@ func ToEC2Instances(insts []*ec2.Instance) []EC2Instance {
}
// ServerInfos creates a ServerInfo resource for each discovered instance.
func (i *EC2Instances) ServerInfos() ([]types.ServerInfo, error) {
serverInfos := make([]types.ServerInfo, 0, len(i.Instances))
for _, instance := range i.Instances {
name := i.AccountID + "-" + instance.InstanceID
tags := make(map[string]string, len(instance.Tags))
for k, v := range instance.Tags {
tags[labels.FormatCloudLabelKey(labels.AWSLabelNamespace, k)] = v
}
si, err := types.NewServerInfo(types.Metadata{
Name: name,
}, types.ServerInfoSpecV1{
AWS: &types.ServerInfoSpecV1_AWSInfo{
AccountID: i.AccountID,
InstanceID: instance.InstanceID,
},
NewLabels: tags,
})
if err != nil {
return nil, trace.Wrap(err)
}
serverInfos = append(serverInfos, si)
}
return serverInfos, nil
}
// NewEC2Watcher creates a new EC2 watcher instance.
func NewEC2Watcher(ctx context.Context, matchers []types.AWSMatcher, clients cloud.Clients, missedRotation <-chan []types.Server) (*Watcher, error) {
cancelCtx, cancelFn := context.WithCancel(ctx)
+30
View File
@@ -24,6 +24,7 @@ import (
"github.com/aws/aws-sdk-go/aws/request"
"github.com/aws/aws-sdk-go/service/ec2"
"github.com/aws/aws-sdk-go/service/ec2/ec2iface"
"github.com/google/go-cmp/cmp"
"github.com/stretchr/testify/require"
"github.com/gravitational/teleport/api/types"
@@ -233,3 +234,32 @@ func TestEC2Watcher(t *testing.T) {
Parameters: map[string]string{"token": "", "scriptName": ""},
}, *result.EC2Instances)
}
func TestConvertEC2InstancesToServerInfos(t *testing.T) {
t.Parallel()
expected, err := types.NewServerInfo(types.Metadata{
Name: "myaccount-myinstance",
}, types.ServerInfoSpecV1{
AWS: &types.ServerInfoSpecV1_AWSInfo{
AccountID: "myaccount",
InstanceID: "myinstance",
},
NewLabels: map[string]string{"aws/foo": "bar"},
})
require.NoError(t, err)
ec2Instances := &EC2Instances{
AccountID: "myaccount",
Instances: []EC2Instance{
{
InstanceID: "myinstance",
Tags: map[string]string{"foo": "bar"},
},
},
}
serverInfos, err := ec2Instances.ServerInfos()
require.NoError(t, err)
require.Len(t, serverInfos, 1)
require.Empty(t, cmp.Diff(expected, serverInfos[0]))
}