mirror of
https://github.com/gravitational/teleport.git
synced 2026-09-24 16:17:11 +08:00
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:
@@ -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.
|
||||
|
||||
@@ -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
File diff suppressed because it is too large
Load Diff
@@ -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()
|
||||
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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")))
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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]))
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user