From 9673e9d15852748e2465b9645a0c453390d90d44 Mon Sep 17 00:00:00 2001 From: Steve Kriss Date: Wed, 28 Feb 2018 15:24:46 -0800 Subject: [PATCH] AWS: copy tags from volume to snapshot, and snapshot to volume Signed-off-by: Steve Kriss --- docs/aws-config.md | 1 + pkg/cloudprovider/aws/block_store.go | 112 ++++++++++++++-------- pkg/cloudprovider/aws/block_store_test.go | 105 ++++++++++++++++++++ 3 files changed, 179 insertions(+), 39 deletions(-) diff --git a/docs/aws-config.md b/docs/aws-config.md index 9a12ba75c..83ff8e3d5 100644 --- a/docs/aws-config.md +++ b/docs/aws-config.md @@ -49,6 +49,7 @@ For more information, see [the AWS documentation on IAM users][14]. "Effect": "Allow", "Action": [ "ec2:DescribeVolumes", + "ec2:DescribeSnapshots", "ec2:CreateTags", "ec2:CreateVolume", "ec2:CreateSnapshot", diff --git a/pkg/cloudprovider/aws/block_store.go b/pkg/cloudprovider/aws/block_store.go index 0c8219adc..db0cf6720 100644 --- a/pkg/cloudprovider/aws/block_store.go +++ b/pkg/cloudprovider/aws/block_store.go @@ -78,10 +78,30 @@ func (b *blockStore) Init(config map[string]string) error { } func (b *blockStore) CreateVolumeFromSnapshot(snapshotID, volumeType, volumeAZ string, iops *int64) (volumeID string, err error) { + // describe the snapshot so we can apply its tags to the volume + snapReq := &ec2.DescribeSnapshotsInput{ + SnapshotIds: []*string{&snapshotID}, + } + + snapRes, err := b.ec2.DescribeSnapshots(snapReq) + if err != nil { + return "", errors.WithStack(err) + } + + if count := len(snapRes.Snapshots); count != 1 { + return "", errors.Errorf("expected 1 snapshot from DescribeSnapshots for %s, got %v", snapshotID, count) + } + req := &ec2.CreateVolumeInput{ SnapshotId: &snapshotID, AvailabilityZone: &volumeAZ, VolumeType: &volumeType, + TagSpecifications: []*ec2.TagSpecification{ + { + ResourceType: aws.String(ec2.ResourceTypeVolume), + Tags: snapRes.Snapshots[0].Tags, + }, + }, } if iopsVolumeTypes.Has(volumeType) && iops != nil { @@ -97,83 +117,97 @@ func (b *blockStore) CreateVolumeFromSnapshot(snapshotID, volumeType, volumeAZ s } func (b *blockStore) GetVolumeInfo(volumeID, volumeAZ string) (string, *int64, error) { - req := &ec2.DescribeVolumesInput{ - VolumeIds: []*string{&volumeID}, - } - - res, err := b.ec2.DescribeVolumes(req) + volumeInfo, err := b.describeVolume(volumeID) if err != nil { - return "", nil, errors.WithStack(err) + return "", nil, err } - if len(res.Volumes) != 1 { - return "", nil, errors.Errorf("Expected one volume from DescribeVolumes for volume ID %v, got %v", volumeID, len(res.Volumes)) - } - - vol := res.Volumes[0] - var ( volumeType string iops *int64 ) - if vol.VolumeType != nil { - volumeType = *vol.VolumeType + if volumeInfo.VolumeType != nil { + volumeType = *volumeInfo.VolumeType } - if iopsVolumeTypes.Has(volumeType) && vol.Iops != nil { - iops = vol.Iops + if iopsVolumeTypes.Has(volumeType) && volumeInfo.Iops != nil { + iops = volumeInfo.Iops } return volumeType, iops, nil } func (b *blockStore) IsVolumeReady(volumeID, volumeAZ string) (ready bool, err error) { + volumeInfo, err := b.describeVolume(volumeID) + if err != nil { + return false, err + } + + return *volumeInfo.State == ec2.VolumeStateAvailable, nil +} + +func (b *blockStore) describeVolume(volumeID string) (*ec2.Volume, error) { req := &ec2.DescribeVolumesInput{ VolumeIds: []*string{&volumeID}, } res, err := b.ec2.DescribeVolumes(req) if err != nil { - return false, errors.WithStack(err) + return nil, errors.WithStack(err) } - if len(res.Volumes) != 1 { - return false, errors.Errorf("Expected one volume from DescribeVolumes for volume ID %v, got %v", volumeID, len(res.Volumes)) + if count := len(res.Volumes); count != 1 { + return nil, errors.Errorf("Expected one volume from DescribeVolumes for volume ID %v, got %v", volumeID, count) } - return *res.Volumes[0].State == ec2.VolumeStateAvailable, nil + return res.Volumes[0], nil } func (b *blockStore) CreateSnapshot(volumeID, volumeAZ string, tags map[string]string) (string, error) { - req := &ec2.CreateSnapshotInput{ - VolumeId: &volumeID, - } - - res, err := b.ec2.CreateSnapshot(req) + res, err := b.ec2.CreateSnapshot(&ec2.CreateSnapshotInput{VolumeId: &volumeID}) if err != nil { return "", errors.WithStack(err) } - tagsReq := &ec2.CreateTagsInput{} - tagsReq.SetResources([]*string{res.SnapshotId}) - - ec2Tags := make([]*ec2.Tag, 0, len(tags)) - - for k, v := range tags { - key := k - val := v - - tag := &ec2.Tag{Key: &key, Value: &val} - ec2Tags = append(ec2Tags, tag) + // describe the volume so we can copy its tags to the snapshot + volumeInfo, err := b.describeVolume(volumeID) + if err != nil { + return "", err } - tagsReq.SetTags(ec2Tags) - - _, err = b.ec2.CreateTags(tagsReq) + _, err = b.ec2.CreateTags(getCreateTagsInput(*res.SnapshotId, tags, volumeInfo.Tags)) return *res.SnapshotId, errors.WithStack(err) } +func getCreateTagsInput(snapshotID string, arkTags map[string]string, volumeTags []*ec2.Tag) *ec2.CreateTagsInput { + tagsInput := &ec2.CreateTagsInput{ + Resources: []*string{&snapshotID}, + } + + // set Ark-assigned tags + for k, v := range arkTags { + tagsInput.Tags = append(tagsInput.Tags, ec2Tag(k, v)) + } + + // copy tags from volume to snapshot + for _, tag := range volumeTags { + // we want current Ark-assigned tags to overwrite any older versions + // of them that may exist due to prior snapshots/restores + if _, found := arkTags[*tag.Key]; found { + continue + } + + tagsInput.Tags = append(tagsInput.Tags, ec2Tag(*tag.Key, *tag.Value)) + } + + return tagsInput +} + +func ec2Tag(key, val string) *ec2.Tag { + return &ec2.Tag{Key: &key, Value: &val} +} + func (b *blockStore) DeleteSnapshot(snapshotID string) error { req := &ec2.DeleteSnapshotInput{ SnapshotId: &snapshotID, diff --git a/pkg/cloudprovider/aws/block_store_test.go b/pkg/cloudprovider/aws/block_store_test.go index e74e07cc3..20d3ed95d 100644 --- a/pkg/cloudprovider/aws/block_store_test.go +++ b/pkg/cloudprovider/aws/block_store_test.go @@ -17,8 +17,11 @@ limitations under the License. package aws import ( + "sort" "testing" + "github.com/aws/aws-sdk-go/aws" + "github.com/aws/aws-sdk-go/service/ec2" "github.com/heptio/ark/pkg/util/collections" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -84,3 +87,105 @@ func TestSetVolumeID(t *testing.T) { require.NoError(t, err) assert.Equal(t, "vol-updated", actual) } + +func TestGetCreateTagsInput(t *testing.T) { + tests := []struct { + name string + snapshotID string + arkTags map[string]string + volumeTags []*ec2.Tag + expected *ec2.CreateTagsInput + }{ + { + name: "degenerate case (no tags)", + snapshotID: "foo", + arkTags: nil, + volumeTags: nil, + expected: &ec2.CreateTagsInput{ + Resources: []*string{aws.String("foo")}, + Tags: nil, + }, + }, + { + name: "ark tags only get applied", + snapshotID: "foo", + arkTags: map[string]string{ + "ark-key1": "ark-val1", + "ark-key2": "ark-val2", + }, + volumeTags: nil, + expected: &ec2.CreateTagsInput{ + Resources: []*string{aws.String("foo")}, + Tags: []*ec2.Tag{ + ec2Tag("ark-key1", "ark-val1"), + ec2Tag("ark-key2", "ark-val2"), + }, + }, + }, + { + name: "volume tags only get applied", + snapshotID: "foo", + arkTags: nil, + volumeTags: []*ec2.Tag{ + ec2Tag("aws-key1", "aws-val1"), + ec2Tag("aws-key2", "aws-val2"), + }, + expected: &ec2.CreateTagsInput{ + Resources: []*string{aws.String("foo")}, + Tags: []*ec2.Tag{ + ec2Tag("aws-key1", "aws-val1"), + ec2Tag("aws-key2", "aws-val2"), + }, + }, + }, + { + name: "non-overlapping ark and volume tags both get applied", + snapshotID: "foo", + arkTags: map[string]string{"ark-key": "ark-val"}, + volumeTags: []*ec2.Tag{ec2Tag("aws-key", "aws-val")}, + expected: &ec2.CreateTagsInput{ + Resources: []*string{aws.String("foo")}, + Tags: []*ec2.Tag{ + ec2Tag("ark-key", "ark-val"), + ec2Tag("aws-key", "aws-val"), + }, + }, + }, + { + name: "when tags overlap, ark tags take precedence", + snapshotID: "foo", + arkTags: map[string]string{ + "ark-key": "ark-val", + "overlapping-key": "ark-val", + }, + volumeTags: []*ec2.Tag{ + ec2Tag("aws-key", "aws-val"), + ec2Tag("overlapping-key", "aws-val"), + }, + expected: &ec2.CreateTagsInput{ + Resources: []*string{aws.String("foo")}, + Tags: []*ec2.Tag{ + ec2Tag("ark-key", "ark-val"), + ec2Tag("overlapping-key", "ark-val"), + ec2Tag("aws-key", "aws-val"), + }, + }, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + res := getCreateTagsInput(test.snapshotID, test.arkTags, test.volumeTags) + + sort.Slice(res.Tags, func(i, j int) bool { + return *res.Tags[i].Key < *res.Tags[j].Key + }) + + sort.Slice(test.expected.Tags, func(i, j int) bool { + return *test.expected.Tags[i].Key < *test.expected.Tags[j].Key + }) + + assert.Equal(t, test.expected, res) + }) + } +}