From 78adbd0b66746db719b859f5f60c11a2812f3810 Mon Sep 17 00:00:00 2001 From: Nitin Kumar Date: Sun, 27 Sep 2026 21:27:57 +0530 Subject: [PATCH 1/2] fix(ec2): add DescribePrefixLists and gateway endpoint prefix-list routes --- docs/coverage/aws/vpc.md | 9 + docs/coverage/coverage.json | 14 + providers/aws/vpc/endpoint.go | 20 +- providers/aws/vpc/endpoint_test.go | 29 ++ providers/aws/vpc/route_table.go | 6 +- providers/aws/vpc/service_prefix_list.go | 211 +++++++++++ providers/aws/vpc/service_prefix_list_test.go | 189 ++++++++++ providers/aws/vpc/tags.go | 10 +- server/aws/ec2/endpoint.go | 60 ++- server/aws/ec2/handler.go | 1 + server/aws/ec2/prefix_list.go | 74 +++- server/aws/ec2/route_table.go | 25 +- server/aws/ec2/service_prefix_list.go | 122 ++++++ server/aws/ec2/service_prefix_list_test.go | 349 ++++++++++++++++++ server/aws/ec2/tags.go | 2 +- .../networking/driver/aws_capabilities.go | 25 ++ services/networking/driver/driver.go | 9 +- 17 files changed, 1111 insertions(+), 44 deletions(-) create mode 100644 providers/aws/vpc/service_prefix_list.go create mode 100644 providers/aws/vpc/service_prefix_list_test.go create mode 100644 server/aws/ec2/service_prefix_list.go create mode 100644 server/aws/ec2/service_prefix_list_test.go diff --git a/docs/coverage/aws/vpc.md b/docs/coverage/aws/vpc.md index 4f52606a6..8c272ac6e 100644 --- a/docs/coverage/aws/vpc.md +++ b/docs/coverage/aws/vpc.md @@ -325,6 +325,15 @@ PrefixLists is an OPTIONAL AWS capability (type-asserted). | `GetManagedPrefixListEntries` | | | `ModifyManagedPrefixList` | | +### ServicePrefixLists + +ServicePrefixLists is an OPTIONAL AWS capability (type-asserted). The lists + +| Operation | Description | +| --- | --- | +| `DescribeAWSManagedPrefixLists` | DescribeAWSManagedPrefixLists returns the same lists in the managed prefix | +| `DescribePrefixLists` | DescribePrefixLists returns the service lists for region, narrowed to ids | + ### SubnetAttributes SubnetAttributes is an OPTIONAL capability, discovered by type assertion. diff --git a/docs/coverage/coverage.json b/docs/coverage/coverage.json index a283afa95..79246683d 100644 --- a/docs/coverage/coverage.json +++ b/docs/coverage/coverage.json @@ -11757,6 +11757,20 @@ } ] }, + { + "name": "ServicePrefixLists", + "doc": "ServicePrefixLists is an OPTIONAL AWS capability (type-asserted). The lists", + "operations": [ + { + "name": "DescribeAWSManagedPrefixLists", + "doc": "DescribeAWSManagedPrefixLists returns the same lists in the managed prefix" + }, + { + "name": "DescribePrefixLists", + "doc": "DescribePrefixLists returns the service lists for region, narrowed to ids" + } + ] + }, { "name": "SubnetAttributes", "doc": "SubnetAttributes is an OPTIONAL capability, discovered by type assertion.", diff --git a/providers/aws/vpc/endpoint.go b/providers/aws/vpc/endpoint.go index 8ee0c5bbc..282da2dcf 100644 --- a/providers/aws/vpc/endpoint.go +++ b/providers/aws/vpc/endpoint.go @@ -12,6 +12,10 @@ import ( // each specified subnet. Gateway-type endpoints hold no interfaces. const vpcEndpointTypeInterface = "Interface" +// vpcEndpointTypeGateway is the endpoint type that routes to the service +// through a prefix-list route in each of its route tables. +const vpcEndpointTypeGateway = "Gateway" + // endpointENIDescription is the description stamped on the ENIs an Interface // endpoint occupies, so DeleteVpcEndpoint can release exactly this endpoint's set. func endpointENIDescription(endpointID string) string { @@ -80,16 +84,19 @@ func (m *Mock) CreateVPCEndpoint( } m.endpoints.Set(id, ep) + m.syncEndpointRoutes(ep) return copyEndpoint(ep), nil } // DeleteVPCEndpoint deletes the VPC endpoint with the given ID, releasing any -// backing ENIs an Interface endpoint provisioned. +// backing ENIs an Interface endpoint provisioned and the prefix-list routes a +// Gateway endpoint added. func (m *Mock) DeleteVPCEndpoint( _ context.Context, id string, ) error { - if !m.endpoints.Has(id) { + ep, ok := m.endpoints.Get(id) + if !ok { return errors.Newf( errors.NotFound, "vpc endpoint %q not found", id, @@ -99,6 +106,10 @@ func (m *Mock) DeleteVPCEndpoint( m.endpoints.Delete(id) m.releaseManagedENIs(endpointENIDescription(id)) + gone := *ep + gone.RouteTableIDs = nil + m.syncEndpointRoutes(&gone) + return nil } @@ -145,8 +156,11 @@ func (m *Mock) ModifyVPCEndpoint( ep.SecurityGroupIDs = copyStringSlice(cfg.SecurityGroupIDs) } - if len(cfg.RouteTableIDs) > 0 { + // A non-nil empty set removes every route table, which is how + // ModifyVpcEndpoint with only RemoveRouteTableId arrives here. + if cfg.RouteTableIDs != nil { ep.RouteTableIDs = copyStringSlice(cfg.RouteTableIDs) + m.syncEndpointRoutes(ep) } if len(cfg.Tags) > 0 { diff --git a/providers/aws/vpc/endpoint_test.go b/providers/aws/vpc/endpoint_test.go index cdfd5b596..ed26fec9f 100644 --- a/providers/aws/vpc/endpoint_test.go +++ b/providers/aws/vpc/endpoint_test.go @@ -77,3 +77,32 @@ func TestGatewayEndpointHasNoENIs(t *testing.T) { assertEqual(t, 0, len(ep.NetworkInterfaceIDs)) assertEqual(t, 0, countENIsInSubnet(m, sub.ID)) } + +// TestVPCEndpointResourceTags pins that the generic CreateTags/DeleteTags path +// reaches VPC endpoints and endpoint services. +func TestVPCEndpointResourceTags(t *testing.T) { + ctx := context.Background() + m := newTestMock() + v := createTestVPC(m) + + ep, err := m.CreateVPCEndpoint(ctx, driver.VPCEndpointConfig{VPCID: v.ID, ServiceName: "com.amazonaws.us-east-1.s3"}) + requireNoError(t, err) + + requireNoError(t, m.UpdateResourceTags(ctx, ep.ID, map[string]string{"env": "prod", "team": "net"})) + requireNoError(t, m.RemoveResourceTags(ctx, ep.ID, []string{"team"})) + + got, err := m.DescribeVPCEndpoints(ctx, []string{ep.ID}) + requireNoError(t, err) + assertEqual(t, 1, len(got[0].Tags)) + assertEqual(t, "prod", got[0].Tags["env"]) + + svc, err := m.CreateVPCEndpointServiceConfiguration(ctx, driver.EndpointServiceConfig{ + NetworkLoadBalancerARNs: []string{"arn:aws:elasticloadbalancing:us-east-1:123456789012:loadbalancer/net/n/1"}, + }) + requireNoError(t, err) + requireNoError(t, m.UpdateResourceTags(ctx, svc.ID, map[string]string{"env": "prod"})) + + if err := m.UpdateResourceTags(ctx, "vpce-missing", map[string]string{"a": "b"}); err == nil { + t.Fatal("tagging an unknown vpce- id should fail") + } +} diff --git a/providers/aws/vpc/route_table.go b/providers/aws/vpc/route_table.go index 626d07037..b7bc56ed5 100644 --- a/providers/aws/vpc/route_table.go +++ b/providers/aws/vpc/route_table.go @@ -151,7 +151,7 @@ func (m *Mock) CreateRoute( } for _, r := range rt.Routes { - if r.DestinationCIDR == destinationCIDR { + if r.DestinationCIDR != "" && r.DestinationCIDR == destinationCIDR { return errors.Newf(errors.AlreadyExists, "route for %q already exists in route table %q", destinationCIDR, routeTableID) } @@ -209,7 +209,9 @@ func (m *Mock) DeleteRoute(_ context.Context, routeTableID, destinationCIDR stri } for i, r := range rt.Routes { - if r.DestinationCIDR == destinationCIDR { + // Prefix-list routes carry no CIDR, so an empty destination never + // matches one of them. + if r.DestinationCIDR != "" && r.DestinationCIDR == destinationCIDR { rt.Routes = append(rt.Routes[:i], rt.Routes[i+1:]...) return nil } diff --git a/providers/aws/vpc/service_prefix_list.go b/providers/aws/vpc/service_prefix_list.go new file mode 100644 index 000000000..c7699fe2c --- /dev/null +++ b/providers/aws/vpc/service_prefix_list.go @@ -0,0 +1,211 @@ +package vpc + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "strings" + + "github.com/stackshy/cloudemu/v2/errors" + "github.com/stackshy/cloudemu/v2/services/networking/driver" +) + +const ( + // servicePrefixListOwner is the ownerId EC2 reports on the AWS-managed + // service prefix lists. + servicePrefixListOwner = "AWS" + // servicePrefixListIDHexLen is the length of the hex part of a service + // list id (pl-63a5400a style). + servicePrefixListIDHexLen = 8 + // serviceNamePrefix starts every AWS service endpoint name. + serviceNamePrefix = "com.amazonaws." + + routeStateActive = "active" +) + +// RouteTargetVPCEndpoint marks the route a Gateway endpoint adds. It encodes +// as gatewayId on the wire, like real EC2. +const RouteTargetVPCEndpoint = "vpc-endpoint" + +// gatewayServiceCIDRs returns the services that have a gateway endpoint (and so +// a prefix list), each with the address ranges its list holds. The ranges are +// the published us-east-1 ones; every region reuses them. +func gatewayServiceCIDRs() map[string][]string { + return map[string][]string{ + "s3": { + "3.5.0.0/19", "16.182.0.0/16", "18.34.0.0/19", + "18.34.232.0/21", "52.216.0.0/15", "54.231.0.0/16", + }, + "dynamodb": { + "3.218.180.0/22", "3.218.184.0/22", "52.94.0.0/22", "52.119.224.0/20", + }, + } +} + +// gatewayServices is the fixed order the service lists are returned in. +func gatewayServices() []string { + return []string{"s3", "dynamodb"} +} + +// servicePrefixListID derives the pl- id for a service list from its name, so +// the id is the same on every call and differs between regions. +func servicePrefixListID(name string) string { + sum := sha256.Sum256([]byte(name)) + + return "pl-" + hex.EncodeToString(sum[:])[:servicePrefixListIDHexLen] +} + +// servicePrefixListForEndpoint returns the pl- id a Gateway endpoint for +// serviceName routes to, or "" when the service has no prefix list. +func servicePrefixListForEndpoint(serviceName string) string { + rest, ok := strings.CutPrefix(serviceName, serviceNamePrefix) + if !ok { + return "" + } + + dot := strings.LastIndex(rest, ".") + if dot <= 0 { + return "" + } + + if _, known := gatewayServiceCIDRs()[rest[dot+1:]]; !known { + return "" + } + + return servicePrefixListID(serviceName) +} + +func (m *Mock) servicePrefixLists(region string) []driver.ServicePrefixList { + region = orDefaultStr(region, m.opts.Region) + cidrs := gatewayServiceCIDRs() + + out := make([]driver.ServicePrefixList, 0, len(cidrs)) + + for _, svc := range gatewayServices() { + name := serviceNamePrefix + region + "." + svc + out = append(out, driver.ServicePrefixList{ + ID: servicePrefixListID(name), + Name: name, + CIDRs: append([]string(nil), cidrs[svc]...), + }) + } + + return out +} + +// DescribePrefixLists returns the AWS service prefix lists for region. +func (m *Mock) DescribePrefixLists( + _ context.Context, region string, ids []string, +) ([]driver.ServicePrefixList, error) { + all := m.servicePrefixLists(region) + if len(ids) == 0 { + return all, nil + } + + byID := make(map[string]driver.ServicePrefixList, len(all)) + for _, pl := range all { + byID[pl.ID] = pl + } + + out := make([]driver.ServicePrefixList, 0, len(ids)) + + for _, id := range ids { + pl, ok := byID[id] + if !ok { + return nil, errors.Newf(errors.NotFound, "The prefix list ID '%s' does not exist", id) + } + + out = append(out, pl) + } + + return out, nil +} + +// DescribeAWSManagedPrefixLists returns the service lists in the managed +// prefix list shape. Unknown ids are skipped. +func (m *Mock) DescribeAWSManagedPrefixLists( + _ context.Context, region string, ids []string, +) ([]driver.PrefixList, error) { + want := make(map[string]bool, len(ids)) + for _, id := range ids { + want[id] = true + } + + var out []driver.PrefixList + + for _, pl := range m.servicePrefixLists(region) { + if len(ids) > 0 && !want[pl.ID] { + continue + } + + entries := make([]driver.PrefixListEntry, 0, len(pl.CIDRs)) + for _, c := range pl.CIDRs { + entries = append(entries, driver.PrefixListEntry{CIDR: c}) + } + + out = append(out, driver.PrefixList{ + ID: pl.ID, Name: pl.Name, AddressFamily: "IPv4", + MaxEntries: len(entries), State: "create-complete", Version: 1, + Entries: entries, OwnerID: servicePrefixListOwner, + }) + } + + return out, nil +} + +// syncEndpointRoutes makes the prefix-list routes of a Gateway endpoint match +// its route table set: one route per table to the service's pl- id, removed +// from tables the endpoint no longer uses. Tables that do not exist are +// skipped. Other endpoint types hold no routes. +func (m *Mock) syncEndpointRoutes(ep *driver.VPCEndpoint) { + var plID string + if ep.EndpointType == "" || ep.EndpointType == vpcEndpointTypeGateway { + plID = servicePrefixListForEndpoint(ep.ServiceName) + } + + want := map[string]bool{} + + if plID != "" { + for _, id := range ep.RouteTableIDs { + want[id] = true + } + } + + m.mu.Lock() + defer m.mu.Unlock() + + for _, rt := range m.routeTables.All() { + rt.Routes = endpointRoutes(rt.Routes, ep.ID, plID, want[rt.ID]) + } +} + +// endpointRoutes returns routes with at most one route to endpointID, kept +// (or added, pointing at plID) only when keep is set. +func endpointRoutes(routes []driver.Route, endpointID, plID string, keep bool) []driver.Route { + out := routes[:0:0] + has := false + + for _, r := range routes { + if r.TargetID != endpointID { + out = append(out, r) + continue + } + + if keep && !has { + has = true + + out = append(out, r) + } + } + + if keep && !has { + out = append(out, driver.Route{ + DestinationPrefixListID: plID, + TargetID: endpointID, + TargetType: RouteTargetVPCEndpoint, + State: routeStateActive, + }) + } + + return out +} diff --git a/providers/aws/vpc/service_prefix_list_test.go b/providers/aws/vpc/service_prefix_list_test.go new file mode 100644 index 000000000..cc3836b64 --- /dev/null +++ b/providers/aws/vpc/service_prefix_list_test.go @@ -0,0 +1,189 @@ +package vpc + +import ( + "context" + "strings" + "testing" + + "github.com/stackshy/cloudemu/v2/errors" + "github.com/stackshy/cloudemu/v2/services/networking/driver" +) + +// TestServicePrefixListsPerRegion pins that DescribePrefixLists returns the +// AWS-managed s3 and dynamodb lists for the region asked about, with pl- ids +// that stay the same across calls and differ between regions. +func TestServicePrefixListsPerRegion(t *testing.T) { + ctx := context.Background() + m := newTestMock() + + east, err := m.DescribePrefixLists(ctx, "us-east-1", nil) + requireNoError(t, err) + assertEqual(t, 2, len(east)) + + names := map[string]string{} + + for _, pl := range east { + if !strings.HasPrefix(pl.ID, "pl-") { + t.Errorf("prefix list id %q lacks the pl- prefix", pl.ID) + } + + if len(pl.CIDRs) == 0 { + t.Errorf("prefix list %s has no cidrs", pl.Name) + } + + names[pl.Name] = pl.ID + } + + if names["com.amazonaws.us-east-1.s3"] == "" || names["com.amazonaws.us-east-1.dynamodb"] == "" { + t.Fatalf("missing s3/dynamodb lists, got %v", names) + } + + again, err := m.DescribePrefixLists(ctx, "us-east-1", nil) + requireNoError(t, err) + + for _, pl := range again { + assertEqual(t, names[pl.Name], pl.ID) + } + + west, err := m.DescribePrefixLists(ctx, "us-west-2", nil) + requireNoError(t, err) + + for _, pl := range west { + if !strings.HasPrefix(pl.Name, "com.amazonaws.us-west-2.") { + t.Errorf("us-west-2 list named %q", pl.Name) + } + + if pl.ID == names["com.amazonaws.us-east-1.s3"] || pl.ID == names["com.amazonaws.us-east-1.dynamodb"] { + t.Errorf("us-west-2 list %s reuses a us-east-1 id", pl.ID) + } + } +} + +// TestServicePrefixListsByID pins id lookup and the NotFound for an unknown id. +func TestServicePrefixListsByID(t *testing.T) { + ctx := context.Background() + m := newTestMock() + + all, err := m.DescribePrefixLists(ctx, "", nil) + requireNoError(t, err) + + one, err := m.DescribePrefixLists(ctx, "", []string{all[0].ID}) + requireNoError(t, err) + assertEqual(t, 1, len(one)) + assertEqual(t, all[0].Name, one[0].Name) + + _, err = m.DescribePrefixLists(ctx, "", []string{"pl-00000000"}) + if !errors.IsNotFound(err) { + t.Fatalf("unknown id: want NotFound, got %v", err) + } +} + +// TestAWSManagedPrefixListsShape pins the managed-prefix-list view of the +// service lists: owner AWS, create-complete, entries equal to the cidrs. +func TestAWSManagedPrefixListsShape(t *testing.T) { + ctx := context.Background() + m := newTestMock() + + svc, err := m.DescribePrefixLists(ctx, "us-east-1", nil) + requireNoError(t, err) + + managed, err := m.DescribeAWSManagedPrefixLists(ctx, "us-east-1", nil) + requireNoError(t, err) + assertEqual(t, len(svc), len(managed)) + + for i := range managed { + pl := managed[i] + assertEqual(t, "AWS", pl.OwnerID) + assertEqual(t, "create-complete", pl.State) + assertEqual(t, "IPv4", pl.AddressFamily) + assertEqual(t, svc[i].ID, pl.ID) + assertEqual(t, len(svc[i].CIDRs), len(pl.Entries)) + } + + none, err := m.DescribeAWSManagedPrefixLists(ctx, "us-east-1", []string{"pl-00000000"}) + requireNoError(t, err) + assertEqual(t, 0, len(none)) +} + +// routesTo returns the routes in rtID whose target is targetID. +func routesTo(t *testing.T, m *Mock, rtID, targetID string) []driver.Route { + t.Helper() + + rts, err := m.DescribeRouteTables(context.Background(), []string{rtID}) + requireNoError(t, err) + + var out []driver.Route + + for _, r := range rts[0].Routes { + if r.TargetID == targetID { + out = append(out, r) + } + } + + return out +} + +// TestGatewayEndpointPrefixListRoutes pins that a Gateway endpoint adds a +// DestinationPrefixListId route to each of its route tables, pointing at the +// service's pl- id, and that modify and delete take the routes away again. +func TestGatewayEndpointPrefixListRoutes(t *testing.T) { + ctx := context.Background() + m := newTestMock() + v := createTestVPC(m) + + rtA, err := m.CreateRouteTable(ctx, driver.RouteTableConfig{VPCID: v.ID}) + requireNoError(t, err) + rtB, err := m.CreateRouteTable(ctx, driver.RouteTableConfig{VPCID: v.ID}) + requireNoError(t, err) + + pls, err := m.DescribePrefixLists(ctx, "us-east-1", nil) + requireNoError(t, err) + + var s3ID string + + for _, pl := range pls { + if pl.Name == "com.amazonaws.us-east-1.s3" { + s3ID = pl.ID + } + } + + ep, err := m.CreateVPCEndpoint(ctx, driver.VPCEndpointConfig{ + VPCID: v.ID, ServiceName: "com.amazonaws.us-east-1.s3", + EndpointType: "Gateway", RouteTableIDs: []string{rtA.ID, rtB.ID}, + }) + requireNoError(t, err) + + for _, rt := range []string{rtA.ID, rtB.ID} { + got := routesTo(t, m, rt, ep.ID) + assertEqual(t, 1, len(got)) + assertEqual(t, s3ID, got[0].DestinationPrefixListID) + assertEqual(t, "", got[0].DestinationCIDR) + assertEqual(t, "active", got[0].State) + } + + _, err = m.ModifyVPCEndpoint(ctx, ep.ID, driver.VPCEndpointConfig{RouteTableIDs: []string{rtA.ID}}) + requireNoError(t, err) + assertEqual(t, 1, len(routesTo(t, m, rtA.ID, ep.ID))) + assertEqual(t, 0, len(routesTo(t, m, rtB.ID, ep.ID))) + + requireNoError(t, m.DeleteVPCEndpoint(ctx, ep.ID)) + assertEqual(t, 0, len(routesTo(t, m, rtA.ID, ep.ID))) +} + +// TestInterfaceEndpointAddsNoRoutes pins that only Gateway endpoints touch +// route tables. +func TestInterfaceEndpointAddsNoRoutes(t *testing.T) { + ctx := context.Background() + m := newTestMock() + v := createTestVPC(m) + + rt, err := m.CreateRouteTable(ctx, driver.RouteTableConfig{VPCID: v.ID}) + requireNoError(t, err) + + ep, err := m.CreateVPCEndpoint(ctx, driver.VPCEndpointConfig{ + VPCID: v.ID, ServiceName: "com.amazonaws.us-east-1.ssm", + EndpointType: "Interface", RouteTableIDs: []string{rt.ID}, + }) + requireNoError(t, err) + assertEqual(t, 0, len(routesTo(t, m, rt.ID, ep.ID))) +} diff --git a/providers/aws/vpc/tags.go b/providers/aws/vpc/tags.go index be9d1d12b..1d0d03acc 100644 --- a/providers/aws/vpc/tags.go +++ b/providers/aws/vpc/tags.go @@ -11,7 +11,8 @@ import ( // UpdateResourceTags merges tags onto a VPC-family resource that has no // dedicated Update*Tags method: route tables, internet gateways, NAT // gateways, network ACLs, DHCP option sets, peering connections, managed -// prefix lists, and egress-only internet gateways. An unknown or missing id +// prefix lists, VPC endpoints and endpoint services, and egress-only internet +// gateways. An unknown or missing id // is NotFound, so the wire layer can map it to the InvalidID.NotFound code // real EC2 returns for CreateTags on a non-existent resource. func (m *Mock) UpdateResourceTags(_ context.Context, id string, tags map[string]string) error { @@ -56,6 +57,13 @@ func (m *Mock) mutateResourceTags(id string, transform func(map[string]string) m return m.peerings.Update(id, func(v *peeringData) *peeringData { v.Tags = transform(v.Tags); return v }) case strings.HasPrefix(id, "pl-"): return m.prefixLists.Update(id, func(v *driver.PrefixList) *driver.PrefixList { v.Tags = transform(v.Tags); return v }) + case strings.HasPrefix(id, "vpce-svc-"): + return m.endpointServices.Update(id, func(v *driver.EndpointService) *driver.EndpointService { + v.Tags = transform(v.Tags) + return v + }) + case strings.HasPrefix(id, "vpce-"): + return m.endpoints.Update(id, func(v *driver.VPCEndpoint) *driver.VPCEndpoint { v.Tags = transform(v.Tags); return v }) case strings.HasPrefix(id, "eigw-"): return m.egressOnlyIGWs.Update(id, func(v *driver.EgressOnlyInternetGateway) *driver.EgressOnlyInternetGateway { v.Tags = transform(v.Tags) diff --git a/server/aws/ec2/endpoint.go b/server/aws/ec2/endpoint.go index e9f9416ca..6e1ad28c6 100644 --- a/server/aws/ec2/endpoint.go +++ b/server/aws/ec2/endpoint.go @@ -14,17 +14,24 @@ import ( const defaultVPCEndpointType = "Gateway" type vpcEndpointXML struct { - VpcEndpointID string `xml:"vpcEndpointId"` - VpcEndpointType string `xml:"vpcEndpointType"` - VpcID string `xml:"vpcId"` - ServiceName string `xml:"serviceName"` - State string `xml:"state"` - RouteTableIDs []string `xml:"routeTableIdSet>item,omitempty"` - SubnetIDs []string `xml:"subnetIdSet>item,omitempty"` - Groups []string `xml:"groupSet>item,omitempty"` - NetworkInterfaceIDs []string `xml:"networkInterfaceIdSet>item,omitempty"` - CreationTime string `xml:"creationTimestamp,omitempty"` - Tags []tagItem `xml:"tagSet>item,omitempty"` + VpcEndpointID string `xml:"vpcEndpointId"` + VpcEndpointType string `xml:"vpcEndpointType"` + VpcID string `xml:"vpcId"` + ServiceName string `xml:"serviceName"` + State string `xml:"state"` + RouteTableIDs []string `xml:"routeTableIdSet>item,omitempty"` + SubnetIDs []string `xml:"subnetIdSet>item,omitempty"` + Groups []endpointGroupXML `xml:"groupSet>item,omitempty"` + NetworkInterfaceIDs []string `xml:"networkInterfaceIdSet>item,omitempty"` + CreationTime string `xml:"creationTimestamp,omitempty"` + Tags []tagItem `xml:"tagSet>item,omitempty"` +} + +// endpointGroupXML is one groupSet item (SecurityGroupIdentifier). Terraform +// reads security_group_ids from its groupId. +type endpointGroupXML struct { + GroupID string `xml:"groupId"` + GroupName string `xml:"groupName,omitempty"` } func (h *Handler) routeVPCEndpoints(w http.ResponseWriter, r *http.Request, action string) bool { @@ -69,7 +76,7 @@ func (h *Handler) createVPCEndpoint(w http.ResponseWriter, r *http.Request) { Xmlns string `xml:"xmlns,attr"` Req string `xml:"requestId"` Endpoint vpcEndpointXML `xml:"vpcEndpoint"` - }{Xmlns: awsquery.Namespace, Req: awsquery.RequestID, Endpoint: toVPCEndpointXML(ep)}) + }{Xmlns: awsquery.Namespace, Req: awsquery.RequestID, Endpoint: h.toVPCEndpointXML(r, ep)}) } // deleteVPCEndpoints is idempotent: like real EC2 it always returns HTTP 200 @@ -108,7 +115,7 @@ func (h *Handler) describeVPCEndpoints(w http.ResponseWriter, r *http.Request) { for i := range items { if vpcEndpointMatchesFilters(&items[i], filters) { - out = append(out, toVPCEndpointXML(&items[i])) + out = append(out, h.toVPCEndpointXML(r, &items[i])) } } @@ -214,7 +221,7 @@ func vpcEndpointMatchesFilter(ep *netdriver.VPCEndpoint, f awsquery.Filter) bool } } -func toVPCEndpointXML(ep *netdriver.VPCEndpoint) vpcEndpointXML { +func (h *Handler) toVPCEndpointXML(r *http.Request, ep *netdriver.VPCEndpoint) vpcEndpointXML { return vpcEndpointXML{ VpcEndpointID: ep.ID, VpcEndpointType: nonEmpty(ep.EndpointType, defaultVPCEndpointType), @@ -223,7 +230,7 @@ func toVPCEndpointXML(ep *netdriver.VPCEndpoint) vpcEndpointXML { State: nonEmpty(ep.State, stateAvailable), RouteTableIDs: ep.RouteTableIDs, SubnetIDs: ep.SubnetIDs, - Groups: ep.SecurityGroupIDs, + Groups: h.endpointGroups(r, ep.SecurityGroupIDs), NetworkInterfaceIDs: ep.NetworkInterfaceIDs, CreationTime: ep.CreatedAt, Tags: toTagItems(ep.Tags), @@ -233,3 +240,26 @@ func toVPCEndpointXML(ep *netdriver.VPCEndpoint) vpcEndpointXML { func writeVPCEndpointErr(w http.ResponseWriter, err error) { writeErrWithNotFound(w, err, "InvalidVpcEndpointId.NotFound", "DependencyViolation") } + +// endpointGroups pairs each security group id with its name, as the groupSet of +// DescribeVpcEndpoints does. A group that no longer exists keeps its id. +func (h *Handler) endpointGroups(r *http.Request, ids []string) []endpointGroupXML { + if len(ids) == 0 { + return nil + } + + names := map[string]string{} + + if groups, err := h.vpc.DescribeSecurityGroups(r.Context(), nil); err == nil { + for i := range groups { + names[groups[i].ID] = groups[i].Name + } + } + + out := make([]endpointGroupXML, 0, len(ids)) + for _, id := range ids { + out = append(out, endpointGroupXML{GroupID: id, GroupName: names[id]}) + } + + return out +} diff --git a/server/aws/ec2/handler.go b/server/aws/ec2/handler.go index 4dc539e8c..abb943ff7 100644 --- a/server/aws/ec2/handler.go +++ b/server/aws/ec2/handler.go @@ -167,6 +167,7 @@ func (h *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) { h.routeTransitGateways, h.routeVPN, h.routeDHCPOptions, + h.routeServicePrefixLists, h.routePrefixLists, h.routeEgressOnlyIGW, h.routeEndpointServices, diff --git a/server/aws/ec2/prefix_list.go b/server/aws/ec2/prefix_list.go index b4fd49c57..10ee97a5a 100644 --- a/server/aws/ec2/prefix_list.go +++ b/server/aws/ec2/prefix_list.go @@ -95,30 +95,79 @@ func (h *Handler) deletePrefixList(w http.ResponseWriter, r *http.Request, p net }{Xmlns: awsquery.Namespace, Req: awsquery.RequestID, PL: h.toPrefixListXML(regionFromRequest(r), out)}) } +// describePrefixLists answers DescribeManagedPrefixLists. Like real EC2 it +// returns the AWS-owned service lists (owner AWS) next to the account's own, +// and an explicitly named id that is neither is InvalidPrefixListID.NotFound. func (h *Handler) describePrefixLists(w http.ResponseWriter, r *http.Request, p netdriver.PrefixLists) { - items, err := p.DescribeManagedPrefixLists(r.Context(), awsquery.ListStrings(r.Form, "PrefixListId")) + filters := awsquery.Filters(r.Form) + if err := validateNetworkingFilters(filters, h.matchManagedPrefixListFilter); err != nil { + writePrefixListErr(w, err) + return + } + + ids := awsquery.ListStrings(r.Form, "PrefixListId") + + items, err := p.DescribeManagedPrefixLists(r.Context(), ids) if err != nil { writePrefixListErr(w, err) return } + items = append(items, h.awsManagedPrefixLists(r, ids)...) + + if missing := missingPrefixListID(ids, items); missing != "" { + writePrefixListErr(w, cerrors.Newf(cerrors.NotFound, "The prefix list ID '%s' does not exist", missing)) + return + } + region := regionFromRequest(r) out := make([]prefixListXML, 0, len(items)) + for i := range items { - out = append(out, h.toPrefixListXML(region, &items[i])) + if matchNetworkingFilters(&items[i], filters, h.matchManagedPrefixListFilter) { + out = append(out, h.toPrefixListXML(region, &items[i])) + } } + page, next := pageNetworkingXML(out, r, func(x prefixListXML) string { return x.PrefixListID }) + awsquery.WriteXMLResponse(w, struct { XMLName xml.Name `xml:"DescribeManagedPrefixListsResponse"` Xmlns string `xml:"xmlns,attr"` Req string `xml:"requestId"` Set []prefixListXML `xml:"prefixListSet>item"` - }{Xmlns: awsquery.Namespace, Req: awsquery.RequestID, Set: out}) + Next string `xml:"nextToken,omitempty"` + }{Xmlns: awsquery.Namespace, Req: awsquery.RequestID, Set: page, Next: next}) } -func (*Handler) getPrefixListEntries(w http.ResponseWriter, r *http.Request, p netdriver.PrefixLists) { - entries, err := p.GetManagedPrefixListEntries(r.Context(), r.Form.Get("PrefixListId")) +// missingPrefixListID returns the first of ids with no list in items. +func missingPrefixListID(ids []string, items []netdriver.PrefixList) string { + found := make(map[string]bool, len(items)) + for i := range items { + found[items[i].ID] = true + } + + for _, id := range ids { + if !found[id] { + return id + } + } + + return "" +} + +func (h *Handler) getPrefixListEntries(w http.ResponseWriter, r *http.Request, p netdriver.PrefixLists) { + id := r.Form.Get("PrefixListId") + + entries, err := p.GetManagedPrefixListEntries(r.Context(), id) + if cerrors.IsNotFound(err) { + // The AWS-owned service lists are readable too. + if owned := h.awsManagedPrefixLists(r, []string{id}); len(owned) == 1 { + entries, err = owned[0].Entries, nil + } + } + if err != nil { writePrefixListErr(w, err) return @@ -265,22 +314,29 @@ func parsePrefixListEntries(r *http.Request) []netdriver.PrefixListEntry { } func (h *Handler) toPrefixListXML(region string, p *netdriver.PrefixList) prefixListXML { + owner := nonEmpty(p.OwnerID, h.accountID) + return prefixListXML{ - PrefixListID: p.ID, PrefixListArn: h.prefixListARN(region, p.ID), + PrefixListID: p.ID, PrefixListArn: prefixListARN(region, owner, p.ID), PrefixListName: p.Name, AddressFamily: p.AddressFamily, MaxEntries: p.MaxEntries, State: p.State, Version: p.Version, - OwnerID: h.accountID, Tags: toTagItems(p.Tags), + OwnerID: owner, Tags: toTagItems(p.Tags), } } // prefixListARN builds the managed-prefix-list ARN AWS returns; the SDK and // Terraform read prefixListArn to reference the list in policies and rules. -func (h *Handler) prefixListARN(region, id string) string { +// The AWS-owned lists carry "aws" in the account field. +func prefixListARN(region, owner, id string) string { if id == "" { return "" } - return "arn:aws:ec2:" + region + ":" + h.accountID + ":prefix-list/" + id + if owner == "AWS" { + owner = "aws" + } + + return "arn:aws:ec2:" + region + ":" + owner + ":prefix-list/" + id } func writePrefixListErr(w http.ResponseWriter, err error) { diff --git a/server/aws/ec2/route_table.go b/server/aws/ec2/route_table.go index a98b1a5ee..a140ebc24 100644 --- a/server/aws/ec2/route_table.go +++ b/server/aws/ec2/route_table.go @@ -17,6 +17,9 @@ const ( targetTypeNatGateway = "nat-gateway" targetTypePeering = "peering" targetTypeLocal = "local" + // targetTypeVPCEndpoint is the route a Gateway VPC endpoint adds; its + // target is reported as gatewayId. + targetTypeVPCEndpoint = "vpc-endpoint" // routeOriginCreateRouteTable is the origin AWS reports for the implicit // local route created with the table; routeOriginCreateRoute is what it @@ -30,12 +33,13 @@ const ( ) type routeXML struct { - DestinationCIDR string `xml:"destinationCidrBlock"` - GatewayID string `xml:"gatewayId,omitempty"` - NatGatewayID string `xml:"natGatewayId,omitempty"` - VpcPeeringConnection string `xml:"vpcPeeringConnectionId,omitempty"` - State string `xml:"state"` - Origin string `xml:"origin,omitempty"` + DestinationCIDR string `xml:"destinationCidrBlock,omitempty"` + DestinationPrefixList string `xml:"destinationPrefixListId,omitempty"` + GatewayID string `xml:"gatewayId,omitempty"` + NatGatewayID string `xml:"natGatewayId,omitempty"` + VpcPeeringConnection string `xml:"vpcPeeringConnectionId,omitempty"` + State string `xml:"state"` + Origin string `xml:"origin,omitempty"` } type rtAssociationStateXML struct { @@ -432,13 +436,14 @@ func (h *Handler) toRouteTableXML(rt *netdriver.RouteTable) routeTableXML { for _, route := range rt.Routes { rx := routeXML{ - DestinationCIDR: route.DestinationCIDR, - State: nonEmpty(route.State, "active"), - Origin: routeOrigin(route.TargetType), + DestinationCIDR: route.DestinationCIDR, + DestinationPrefixList: route.DestinationPrefixListID, + State: nonEmpty(route.State, "active"), + Origin: routeOrigin(route.TargetType), } switch route.TargetType { - case targetTypeGateway: + case targetTypeGateway, targetTypeVPCEndpoint: rx.GatewayID = route.TargetID case targetTypeNatGateway: rx.NatGatewayID = route.TargetID diff --git a/server/aws/ec2/service_prefix_list.go b/server/aws/ec2/service_prefix_list.go new file mode 100644 index 000000000..98474b5ba --- /dev/null +++ b/server/aws/ec2/service_prefix_list.go @@ -0,0 +1,122 @@ +package ec2 + +import ( + "encoding/xml" + "net/http" + + "github.com/stackshy/cloudemu/v2/server/wire/awsquery" + netdriver "github.com/stackshy/cloudemu/v2/services/networking/driver" +) + +const ( + filterPrefixListID = "prefix-list-id" + filterPrefixListName = "prefix-list-name" +) + +type servicePrefixListXML struct { + PrefixListID string `xml:"prefixListId"` + PrefixListName string `xml:"prefixListName"` + Cidrs []string `xml:"cidrSet>item"` +} + +func (h *Handler) servicePrefixLists() (netdriver.ServicePrefixLists, bool) { + p, ok := h.vpc.(netdriver.ServicePrefixLists) + + return p, ok +} + +func (h *Handler) routeServicePrefixLists(w http.ResponseWriter, r *http.Request, action string) bool { + if action != "DescribePrefixLists" { + return false + } + + p, ok := h.servicePrefixLists() + if !ok { + return false + } + + h.describeServicePrefixLists(w, r, p) + + return true +} + +// describeServicePrefixLists answers DescribePrefixLists: the AWS service +// prefix lists for the caller's region. Terraform's aws_vpc_endpoint read +// looks one up by prefix-list-name to fill prefix_list_id and cidr_blocks. +func (*Handler) describeServicePrefixLists(w http.ResponseWriter, r *http.Request, p netdriver.ServicePrefixLists) { + filters := awsquery.Filters(r.Form) + if err := validateNetworkingFilters(filters, matchServicePrefixListFilter); err != nil { + writePrefixListErr(w, err) + return + } + + lists, err := p.DescribePrefixLists(r.Context(), regionFromRequest(r), awsquery.ListStrings(r.Form, "PrefixListId")) + if err != nil { + writePrefixListErr(w, err) + return + } + + out := make([]servicePrefixListXML, 0, len(lists)) + + for i := range lists { + if matchNetworkingFilters(&lists[i], filters, matchServicePrefixListFilter) { + out = append(out, servicePrefixListXML{ + PrefixListID: lists[i].ID, PrefixListName: lists[i].Name, Cidrs: lists[i].CIDRs, + }) + } + } + + page, next := pageNetworkingXML(out, r, func(x servicePrefixListXML) string { return x.PrefixListID }) + + awsquery.WriteXMLResponse(w, struct { + XMLName xml.Name `xml:"DescribePrefixListsResponse"` + Xmlns string `xml:"xmlns,attr"` + Req string `xml:"requestId"` + Set []servicePrefixListXML `xml:"prefixListSet>item"` + Next string `xml:"nextToken,omitempty"` + }{Xmlns: awsquery.Namespace, Req: awsquery.RequestID, Set: page, Next: next}) +} + +func matchServicePrefixListFilter(pl *netdriver.ServicePrefixList, f awsquery.Filter) (matched, known bool) { + switch f.Name { + case filterPrefixListID: + return containsString(f.Values, pl.ID), true + case filterPrefixListName: + return containsString(f.Values, pl.Name), true + default: + return false, false + } +} + +func (h *Handler) matchManagedPrefixListFilter(pl *netdriver.PrefixList, f awsquery.Filter) (matched, known bool) { + switch f.Name { + case filterPrefixListID: + return containsString(f.Values, pl.ID), true + case filterPrefixListName: + return containsString(f.Values, pl.Name), true + case filterOwnerID: + return containsString(f.Values, nonEmpty(pl.OwnerID, h.accountID)), true + default: + if matched, isTag := matchStorageTagFilter(pl.Tags, f); isTag { + return matched, true + } + + return false, false + } +} + +// awsManagedPrefixLists returns the AWS-owned lists matching ids for the +// caller's region, or nil when the backend has none. +func (h *Handler) awsManagedPrefixLists(r *http.Request, ids []string) []netdriver.PrefixList { + p, ok := h.servicePrefixLists() + if !ok { + return nil + } + + lists, err := p.DescribeAWSManagedPrefixLists(r.Context(), regionFromRequest(r), ids) + if err != nil { + return nil + } + + return lists +} diff --git a/server/aws/ec2/service_prefix_list_test.go b/server/aws/ec2/service_prefix_list_test.go new file mode 100644 index 000000000..35f262707 --- /dev/null +++ b/server/aws/ec2/service_prefix_list_test.go @@ -0,0 +1,349 @@ +package ec2_test + +import ( + "context" + "errors" + "strings" + "testing" + + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/service/ec2" + ec2types "github.com/aws/aws-sdk-go-v2/service/ec2/types" + "github.com/aws/smithy-go" +) + +const ( + s3PrefixListName = "com.amazonaws.us-east-1.s3" + dynamoPrefixListName = "com.amazonaws.us-east-1.dynamodb" +) + +// servicePrefixListIDs maps each AWS service prefix list name to its pl- id. +func servicePrefixListIDs(t *testing.T, c *ec2.Client) map[string]string { + t.Helper() + + out, err := c.DescribePrefixLists(context.Background(), &ec2.DescribePrefixListsInput{}) + if err != nil { + t.Fatalf("DescribePrefixLists: %v", err) + } + + ids := map[string]string{} + + for _, pl := range out.PrefixLists { + if len(pl.Cidrs) == 0 { + t.Errorf("prefix list %s has no cidrs", aws.ToString(pl.PrefixListName)) + } + + ids[aws.ToString(pl.PrefixListName)] = aws.ToString(pl.PrefixListId) + } + + return ids +} + +func requireAPIErrorCode(t *testing.T, err error, want string) { + t.Helper() + + var apiErr smithy.APIError + if !errors.As(err, &apiErr) { + t.Fatalf("want API error %s, got %v", want, err) + } + + if apiErr.ErrorCode() != want { + t.Fatalf("error code = %s, want %s", apiErr.ErrorCode(), want) + } +} + +// TestDescribePrefixListsServiceLists pins the unfiltered answer: the s3 and +// dynamodb lists for the caller's region, with stable pl- ids. +func TestDescribePrefixListsServiceLists(t *testing.T) { + c := newRoutingEdgeEC2(t) + + ids := servicePrefixListIDs(t, c) + if len(ids) != 2 { + t.Fatalf("got %d prefix lists, want 2: %v", len(ids), ids) + } + + for _, name := range []string{s3PrefixListName, dynamoPrefixListName} { + if !strings.HasPrefix(ids[name], "pl-") { + t.Errorf("%s id = %q, want a pl- id", name, ids[name]) + } + } + + again := servicePrefixListIDs(t, c) + if again[s3PrefixListName] != ids[s3PrefixListName] { + t.Errorf("s3 prefix list id changed between calls: %s then %s", ids[s3PrefixListName], again[s3PrefixListName]) + } +} + +// TestDescribePrefixListsFilters covers the prefix-list-name and +// prefix-list-id filters, PrefixListIds, and MaxResults/NextToken paging. +// Terraform's aws_vpc_endpoint read looks the list up by prefix-list-name. +func TestDescribePrefixListsFilters(t *testing.T) { + ctx := context.Background() + c := newRoutingEdgeEC2(t) + ids := servicePrefixListIDs(t, c) + + byName, err := c.DescribePrefixLists(ctx, &ec2.DescribePrefixListsInput{ + Filters: []ec2types.Filter{{Name: aws.String("prefix-list-name"), Values: []string{s3PrefixListName}}}, + }) + if err != nil { + t.Fatalf("filter by name: %v", err) + } + + if len(byName.PrefixLists) != 1 || aws.ToString(byName.PrefixLists[0].PrefixListId) != ids[s3PrefixListName] { + t.Fatalf("filter by name = %+v", byName.PrefixLists) + } + + byIDFilter, err := c.DescribePrefixLists(ctx, &ec2.DescribePrefixListsInput{ + Filters: []ec2types.Filter{{Name: aws.String("prefix-list-id"), Values: []string{ids[dynamoPrefixListName]}}}, + }) + if err != nil { + t.Fatalf("filter by id: %v", err) + } + + if len(byIDFilter.PrefixLists) != 1 || aws.ToString(byIDFilter.PrefixLists[0].PrefixListName) != dynamoPrefixListName { + t.Fatalf("filter by id = %+v", byIDFilter.PrefixLists) + } + + byIDs, err := c.DescribePrefixLists(ctx, &ec2.DescribePrefixListsInput{PrefixListIds: []string{ids[s3PrefixListName]}}) + if err != nil { + t.Fatalf("PrefixListIds: %v", err) + } + + if len(byIDs.PrefixLists) != 1 { + t.Fatalf("PrefixListIds returned %d lists, want 1", len(byIDs.PrefixLists)) + } + + none, err := c.DescribePrefixLists(ctx, &ec2.DescribePrefixListsInput{ + Filters: []ec2types.Filter{{Name: aws.String("prefix-list-name"), Values: []string{"com.amazonaws.us-east-1.ssm"}}}, + }) + if err != nil || len(none.PrefixLists) != 0 { + t.Fatalf("filter by non-gateway service = %+v, %v; want empty", none, err) + } + + first, err := c.DescribePrefixLists(ctx, &ec2.DescribePrefixListsInput{MaxResults: aws.Int32(1)}) + if err != nil { + t.Fatalf("page 1: %v", err) + } + + if len(first.PrefixLists) != 1 || first.NextToken == nil { + t.Fatalf("page 1 = %d lists, next %v", len(first.PrefixLists), first.NextToken) + } + + second, err := c.DescribePrefixLists(ctx, &ec2.DescribePrefixListsInput{MaxResults: aws.Int32(1), NextToken: first.NextToken}) + if err != nil { + t.Fatalf("page 2: %v", err) + } + + if len(second.PrefixLists) != 1 || second.NextToken != nil || + aws.ToString(second.PrefixLists[0].PrefixListId) == aws.ToString(first.PrefixLists[0].PrefixListId) { + t.Fatalf("page 2 = %+v next %v", second.PrefixLists, second.NextToken) + } +} + +// TestDescribePrefixListsUnknownID pins the EC2 error code for an unknown id. +func TestDescribePrefixListsUnknownID(t *testing.T) { + c := newRoutingEdgeEC2(t) + + _, err := c.DescribePrefixLists(context.Background(), &ec2.DescribePrefixListsInput{ + PrefixListIds: []string{"pl-00000000"}, + }) + requireAPIErrorCode(t, err, "InvalidPrefixListID.NotFound") +} + +// TestGatewayEndpointPrefixListRoute pins that a Gateway s3 endpoint writes a +// route with destinationPrefixListId (the s3 pl- id) and gatewayId (the +// vpce- id) into its route table, and that deleting the endpoint removes it. +func TestGatewayEndpointPrefixListRoute(t *testing.T) { + ctx := context.Background() + c := newRoutingEdgeEC2(t) + vpcID, _ := mkVPCSubnet(t, c) + ids := servicePrefixListIDs(t, c) + + rt, err := c.CreateRouteTable(ctx, &ec2.CreateRouteTableInput{VpcId: aws.String(vpcID)}) + if err != nil { + t.Fatalf("CreateRouteTable: %v", err) + } + + rtID := aws.ToString(rt.RouteTable.RouteTableId) + + ep, err := c.CreateVpcEndpoint(ctx, &ec2.CreateVpcEndpointInput{ + VpcId: aws.String(vpcID), ServiceName: aws.String(s3PrefixListName), + VpcEndpointType: ec2types.VpcEndpointTypeGateway, RouteTableIds: []string{rtID}, + }) + if err != nil { + t.Fatalf("CreateVpcEndpoint: %v", err) + } + + vpceID := aws.ToString(ep.VpcEndpoint.VpcEndpointId) + + route := findVPCERoute(t, c, rtID, vpceID) + if route == nil { + t.Fatal("no route to the gateway endpoint in the route table") + } + + if got := aws.ToString(route.DestinationPrefixListId); got != ids[s3PrefixListName] { + t.Errorf("destinationPrefixListId = %q, want %q", got, ids[s3PrefixListName]) + } + + if route.DestinationCidrBlock != nil { + t.Errorf("destinationCidrBlock = %q, want unset", aws.ToString(route.DestinationCidrBlock)) + } + + if route.State != ec2types.RouteStateActive { + t.Errorf("route state = %s, want active", route.State) + } + + if _, err := c.DeleteVpcEndpoints(ctx, &ec2.DeleteVpcEndpointsInput{VpcEndpointIds: []string{vpceID}}); err != nil { + t.Fatalf("DeleteVpcEndpoints: %v", err) + } + + if findVPCERoute(t, c, rtID, vpceID) != nil { + t.Error("route to the deleted endpoint is still in the route table") + } +} + +func findVPCERoute(t *testing.T, c *ec2.Client, rtID, vpceID string) *ec2types.Route { + t.Helper() + + out, err := c.DescribeRouteTables(context.Background(), &ec2.DescribeRouteTablesInput{RouteTableIds: []string{rtID}}) + if err != nil { + t.Fatalf("DescribeRouteTables: %v", err) + } + + for i := range out.RouteTables[0].Routes { + r := &out.RouteTables[0].Routes[i] + if aws.ToString(r.GatewayId) == vpceID { + return r + } + } + + return nil +} + +// TestInterfaceEndpointGroupSet pins that DescribeVpcEndpoints returns each +// security group as a groupSet item with a groupId. Terraform reads +// security_group_ids from Groups[].GroupId, so a bare-string item drifts. +func TestInterfaceEndpointGroupSet(t *testing.T) { + ctx := context.Background() + c := newRoutingEdgeEC2(t) + vpcID, subnetID := mkVPCSubnet(t, c) + + sg, err := c.CreateSecurityGroup(ctx, &ec2.CreateSecurityGroupInput{ + GroupName: aws.String("vpce-sg"), Description: aws.String("vpce"), VpcId: aws.String(vpcID), + }) + if err != nil { + t.Fatalf("CreateSecurityGroup: %v", err) + } + + sgID := aws.ToString(sg.GroupId) + + ep, err := c.CreateVpcEndpoint(ctx, &ec2.CreateVpcEndpointInput{ + VpcId: aws.String(vpcID), ServiceName: aws.String("com.amazonaws.us-east-1.ssm"), + VpcEndpointType: ec2types.VpcEndpointTypeInterface, + SubnetIds: []string{subnetID}, SecurityGroupIds: []string{sgID}, + }) + if err != nil { + t.Fatalf("CreateVpcEndpoint: %v", err) + } + + out, err := c.DescribeVpcEndpoints(ctx, &ec2.DescribeVpcEndpointsInput{ + VpcEndpointIds: []string{aws.ToString(ep.VpcEndpoint.VpcEndpointId)}, + }) + if err != nil { + t.Fatalf("DescribeVpcEndpoints: %v", err) + } + + groups := out.VpcEndpoints[0].Groups + if len(groups) != 1 || aws.ToString(groups[0].GroupId) != sgID { + t.Fatalf("groups = %+v, want one entry with groupId %s", groups, sgID) + } + + if aws.ToString(groups[0].GroupName) != "vpce-sg" { + t.Errorf("groupName = %q, want vpce-sg", aws.ToString(groups[0].GroupName)) + } +} + +// TestManagedPrefixListsIncludeAWSOwned pins that DescribeManagedPrefixLists +// also lists the AWS-owned service lists (ownerId AWS) and that +// GetManagedPrefixListEntries reads their cidrs. +func TestManagedPrefixListsIncludeAWSOwned(t *testing.T) { + ctx := context.Background() + c := newRoutingEdgeEC2(t) + ids := servicePrefixListIDs(t, c) + + out, err := c.DescribeManagedPrefixLists(ctx, &ec2.DescribeManagedPrefixListsInput{ + Filters: []ec2types.Filter{{Name: aws.String("prefix-list-name"), Values: []string{s3PrefixListName}}}, + }) + if err != nil { + t.Fatalf("DescribeManagedPrefixLists: %v", err) + } + + if len(out.PrefixLists) != 1 { + t.Fatalf("got %d lists, want 1", len(out.PrefixLists)) + } + + pl := out.PrefixLists[0] + if aws.ToString(pl.PrefixListId) != ids[s3PrefixListName] || aws.ToString(pl.OwnerId) != "AWS" { + t.Errorf("list = id %s owner %s", aws.ToString(pl.PrefixListId), aws.ToString(pl.OwnerId)) + } + + if want := "arn:aws:ec2:us-east-1:aws:prefix-list/" + ids[s3PrefixListName]; aws.ToString(pl.PrefixListArn) != want { + t.Errorf("arn = %s, want %s", aws.ToString(pl.PrefixListArn), want) + } + + entries, err := c.GetManagedPrefixListEntries(ctx, &ec2.GetManagedPrefixListEntriesInput{PrefixListId: pl.PrefixListId}) + if err != nil { + t.Fatalf("GetManagedPrefixListEntries: %v", err) + } + + if len(entries.Entries) == 0 { + t.Error("AWS-owned list has no entries") + } + + _, err = c.DescribeManagedPrefixLists(ctx, &ec2.DescribeManagedPrefixListsInput{PrefixListIds: []string{"pl-00000000"}}) + requireAPIErrorCode(t, err, "InvalidPrefixListID.NotFound") +} + +// TestVPCEndpointCreateDeleteTags pins that CreateTags and DeleteTags work on +// a vpce- id. Terraform's aws_vpc_endpoint updates tags this way. +func TestVPCEndpointCreateDeleteTags(t *testing.T) { + ctx := context.Background() + c := newRoutingEdgeEC2(t) + vpcID, _ := mkVPCSubnet(t, c) + + ep, err := c.CreateVpcEndpoint(ctx, &ec2.CreateVpcEndpointInput{ + VpcId: aws.String(vpcID), ServiceName: aws.String(s3PrefixListName), + TagSpecifications: []ec2types.TagSpecification{{ + ResourceType: ec2types.ResourceTypeVpcEndpoint, + Tags: []ec2types.Tag{{Key: aws.String("env"), Value: aws.String("dev")}}, + }}, + }) + if err != nil { + t.Fatalf("CreateVpcEndpoint: %v", err) + } + + id := aws.ToString(ep.VpcEndpoint.VpcEndpointId) + + if _, err := c.CreateTags(ctx, &ec2.CreateTagsInput{ + Resources: []string{id}, + Tags: []ec2types.Tag{{Key: aws.String("env"), Value: aws.String("prod")}, {Key: aws.String("team"), Value: aws.String("net")}}, + }); err != nil { + t.Fatalf("CreateTags: %v", err) + } + + if _, err := c.DeleteTags(ctx, &ec2.DeleteTagsInput{ + Resources: []string{id}, Tags: []ec2types.Tag{{Key: aws.String("team")}}, + }); err != nil { + t.Fatalf("DeleteTags: %v", err) + } + + out, err := c.DescribeVpcEndpoints(ctx, &ec2.DescribeVpcEndpointsInput{VpcEndpointIds: []string{id}}) + if err != nil { + t.Fatalf("DescribeVpcEndpoints: %v", err) + } + + tags := out.VpcEndpoints[0].Tags + if len(tags) != 1 || aws.ToString(tags[0].Key) != "env" || aws.ToString(tags[0].Value) != "prod" { + t.Fatalf("tags = %+v, want only env=prod", tags) + } +} diff --git a/server/aws/ec2/tags.go b/server/aws/ec2/tags.go index 84127147d..0a4958335 100644 --- a/server/aws/ec2/tags.go +++ b/server/aws/ec2/tags.go @@ -332,7 +332,7 @@ func tagNotFoundCode(id string) string { // their own methods and are handled separately. // //nolint:gochecknoglobals // static id-prefix routing table -var networkResourceTagPrefixes = []string{"rtb-", "igw-", "nat-", "acl-", "dopt-", "pcx-", "pl-", "eigw-", "sgr-"} +var networkResourceTagPrefixes = []string{"rtb-", "igw-", "nat-", "acl-", "dopt-", "pcx-", "pl-", "eigw-", "sgr-", "vpce-"} // networkTaggableID reports whether id belongs to a resource tagged via the // NetworkResourceTagger optional interface. diff --git a/services/networking/driver/aws_capabilities.go b/services/networking/driver/aws_capabilities.go index 323c338d6..ddf5a1726 100644 --- a/services/networking/driver/aws_capabilities.go +++ b/services/networking/driver/aws_capabilities.go @@ -225,6 +225,9 @@ type PrefixList struct { Version int Entries []PrefixListEntry Tags map[string]string + // OwnerID is "AWS" for the AWS-managed service lists and empty for a + // customer-managed list, which the caller's account owns. + OwnerID string } // PrefixListConfig is the input to CreateManagedPrefixList. @@ -245,6 +248,28 @@ type PrefixLists interface { ModifyManagedPrefixList(ctx context.Context, id string, addEntries []PrefixListEntry, removeCIDRs []string) (*PrefixList, error) } +// ServicePrefixList is an AWS-managed prefix list for a gateway endpoint +// service (com.amazonaws..s3 or .dynamodb), as DescribePrefixLists +// returns it. +type ServicePrefixList struct { + ID string + Name string + CIDRs []string +} + +// ServicePrefixLists is an OPTIONAL AWS capability (type-asserted). The lists +// are per region, so every call names the region; an empty region means the +// provider's own. +type ServicePrefixLists interface { + // DescribePrefixLists returns the service lists for region, narrowed to ids + // when given. An unknown id is NotFound. + DescribePrefixLists(ctx context.Context, region string, ids []string) ([]ServicePrefixList, error) + // DescribeAWSManagedPrefixLists returns the same lists in the managed prefix + // list shape (owner AWS, entries = cidrs). Unknown ids are skipped so the + // caller can merge the result with customer-managed lists. + DescribeAWSManagedPrefixLists(ctx context.Context, region string, ids []string) ([]PrefixList, error) +} + // ---- Egress-only Internet Gateway (IPv6) ---- // EgressOnlyInternetGateway provides outbound-only IPv6 for private subnets. diff --git a/services/networking/driver/driver.go b/services/networking/driver/driver.go index 1d4996df4..8c39e1455 100644 --- a/services/networking/driver/driver.go +++ b/services/networking/driver/driver.go @@ -239,9 +239,12 @@ type RouteTable struct { // Route represents a route in a route table. type Route struct { DestinationCIDR string - TargetID string // gateway ID, NAT gateway ID, peering connection ID, etc. - TargetType string // "gateway", "nat-gateway", "peering", "local" - State string // "active", "blackhole" + // DestinationPrefixListID is set instead of DestinationCIDR on the routes a + // Gateway VPC endpoint adds (AWS only). + DestinationPrefixListID string + TargetID string // gateway ID, NAT gateway ID, peering connection ID, etc. + TargetType string // "gateway", "nat-gateway", "peering", "local" + State string // "active", "blackhole" } // RouteTableConfig configures a route table. From 093b20ad1cf83d540df3d400df97cd3912012ca2 Mon Sep 17 00:00:00 2001 From: Nitin Kumar Date: Sun, 27 Sep 2026 22:19:20 +0530 Subject: [PATCH 2/2] fix(ec2): apply ModifyVpcEndpoint set changes under the provider lock --- docs/coverage/aws/vpc.md | 8 + docs/coverage/coverage.json | 9 + providers/aws/vpc/endpoint.go | 37 +-- providers/aws/vpc/endpoint_sets.go | 153 ++++++++++++ providers/aws/vpc/endpoint_sets_test.go | 227 ++++++++++++++++++ providers/aws/vpc/route_table.go | 5 + providers/aws/vpc/service_prefix_list.go | 21 +- providers/aws/vpc/service_prefix_list_test.go | 3 + server/aws/ec2/endpoint.go | 27 +++ server/aws/ec2/endpoint_modify_test.go | 106 ++++++++ server/aws/ec2/prefix_list.go | 29 ++- server/aws/ec2/service_prefix_list_test.go | 26 ++ .../networking/driver/aws_capabilities.go | 18 ++ 13 files changed, 639 insertions(+), 30 deletions(-) create mode 100644 providers/aws/vpc/endpoint_sets.go create mode 100644 providers/aws/vpc/endpoint_sets_test.go create mode 100644 server/aws/ec2/endpoint_modify_test.go diff --git a/docs/coverage/aws/vpc.md b/docs/coverage/aws/vpc.md index 8c272ac6e..c94302688 100644 --- a/docs/coverage/aws/vpc.md +++ b/docs/coverage/aws/vpc.md @@ -419,6 +419,14 @@ VPCEndpointServices is an OPTIONAL AWS capability (type-asserted). | `DescribeVPCEndpointServicePermissions` | | | `ModifyVPCEndpointServicePermissions` | | +### VPCEndpointSetModifier + +VPCEndpointSetModifier is an OPTIONAL AWS capability (type-asserted). It + +| Operation | Description | +| --- | --- | +| `ModifyVPCEndpointSets` | | + ### VPNConnections VPNConnections is an OPTIONAL AWS capability (type-asserted). diff --git a/docs/coverage/coverage.json b/docs/coverage/coverage.json index 0b0c52ffe..272fda1b9 100644 --- a/docs/coverage/coverage.json +++ b/docs/coverage/coverage.json @@ -11984,6 +11984,15 @@ } ] }, + { + "name": "VPCEndpointSetModifier", + "doc": "VPCEndpointSetModifier is an OPTIONAL AWS capability (type-asserted). It", + "operations": [ + { + "name": "ModifyVPCEndpointSets" + } + ] + }, { "name": "VPNConnections", "doc": "VPNConnections is an OPTIONAL AWS capability (type-asserted).", diff --git a/providers/aws/vpc/endpoint.go b/providers/aws/vpc/endpoint.go index 897397a51..b259f23e5 100644 --- a/providers/aws/vpc/endpoint.go +++ b/providers/aws/vpc/endpoint.go @@ -68,11 +68,20 @@ func (m *Mock) CreateVPCEndpoint( State: "available", SubnetIDs: copyStringSlice(cfg.SubnetIDs), SecurityGroupIDs: copyStringSlice(cfg.SecurityGroupIDs), - RouteTableIDs: copyStringSlice(cfg.RouteTableIDs), Tags: copyTags(cfg.Tags), CreatedAt: m.opts.Clock.Now().Format(timeFormat), } + // The route-table check, the store write and the route sync run under one + // lock hold, so a concurrent Delete or a second endpoint for the same + // service cannot slip in between. + m.mu.Lock() + defer m.mu.Unlock() + + if err := m.setEndpointRouteTables(ep, cfg.RouteTableIDs); err != nil { + return nil, err + } + // An Interface endpoint provisions one requester-managed ENI per subnet, which // consumes a subnet IP and (like a NAT gateway's ENI) blocks a premature subnet // or VPC delete. Gateway endpoints hold none. @@ -84,10 +93,7 @@ func (m *Mock) CreateVPCEndpoint( } m.endpoints.Set(id, ep) - - m.mu.Lock() m.syncEndpointRoutes(ep) - m.mu.Unlock() return copyEndpoint(ep), nil } @@ -98,6 +104,9 @@ func (m *Mock) CreateVPCEndpoint( func (m *Mock) DeleteVPCEndpoint( _ context.Context, id string, ) error { + m.mu.Lock() + defer m.mu.Unlock() + ep, ok := m.endpoints.Get(id) if !ok { return errors.Newf( @@ -109,11 +118,9 @@ func (m *Mock) DeleteVPCEndpoint( m.endpoints.Delete(id) m.releaseManagedENIs(endpointENIDescription(id)) - m.mu.Lock() gone := *ep gone.RouteTableIDs = nil m.syncEndpointRoutes(&gone) - m.mu.Unlock() return nil } @@ -142,7 +149,8 @@ func (m *Mock) DescribeVPCEndpoints( ), nil } -// ModifyVPCEndpoint updates a VPC endpoint configuration. +// ModifyVPCEndpoint replaces an endpoint's id sets and tags. A nil set leaves +// that set unchanged. The AWS wire layer uses ModifyVPCEndpointSets instead. // //nolint:gocritic // hugeParam: interface method signature cannot be changed. func (m *Mock) ModifyVPCEndpoint( @@ -161,6 +169,14 @@ func (m *Mock) ModifyVPCEndpoint( ) } + if cfg.RouteTableIDs != nil { + if err := m.setEndpointRouteTables(ep, cfg.RouteTableIDs); err != nil { + return nil, err + } + + m.syncEndpointRoutes(ep) + } + if len(cfg.SubnetIDs) > 0 { ep.SubnetIDs = copyStringSlice(cfg.SubnetIDs) } @@ -169,13 +185,6 @@ func (m *Mock) ModifyVPCEndpoint( ep.SecurityGroupIDs = copyStringSlice(cfg.SecurityGroupIDs) } - // A non-nil empty set removes every route table, which is how - // ModifyVpcEndpoint with only RemoveRouteTableId arrives here. - if cfg.RouteTableIDs != nil { - ep.RouteTableIDs = copyStringSlice(cfg.RouteTableIDs) - m.syncEndpointRoutes(ep) - } - if len(cfg.Tags) > 0 { ep.Tags = copyTags(cfg.Tags) } diff --git a/providers/aws/vpc/endpoint_sets.go b/providers/aws/vpc/endpoint_sets.go new file mode 100644 index 000000000..b1215dc31 --- /dev/null +++ b/providers/aws/vpc/endpoint_sets.go @@ -0,0 +1,153 @@ +package vpc + +import ( + "context" + "sort" + + "github.com/stackshy/cloudemu/v2/errors" + "github.com/stackshy/cloudemu/v2/services/networking/driver" +) + +// ModifyVPCEndpointSets applies the Add*/Remove* part of ModifyVpcEndpoint to +// the endpoint's current route tables, subnets and security groups. The read +// and the write happen under one lock hold, so parallel modifies of the same +// endpoint (Terraform creates route table associations concurrently) all land. +func (m *Mock) ModifyVPCEndpointSets( + _ context.Context, id string, change *driver.VPCEndpointSetChange, +) (*driver.VPCEndpoint, error) { + m.mu.Lock() + defer m.mu.Unlock() + + ep, ok := m.endpoints.Get(id) + if !ok { + return nil, errors.Newf(errors.NotFound, "vpc endpoint %q not found", id) + } + + rts := applyIDDelta(ep.RouteTableIDs, change.AddRouteTableIDs, change.RemoveRouteTableIDs) + if err := m.setEndpointRouteTables(ep, rts); err != nil { + return nil, err + } + + m.syncEndpointRoutes(ep) + + subnets := applyIDDelta(ep.SubnetIDs, change.AddSubnetIDs, change.RemoveSubnetIDs) + if ep.EndpointType == vpcEndpointTypeInterface { + m.syncEndpointENIs(ep, subnets) + } + + ep.SubnetIDs = subnets + ep.SecurityGroupIDs = applyIDDelta(ep.SecurityGroupIDs, change.AddSecurityGroupIDs, change.RemoveSecurityGroupIDs) + + return copyEndpoint(ep), nil +} + +// applyIDDelta returns cur without the removed ids and with the added ones +// appended, deduplicated. An id named in both lists is dropped. +func applyIDDelta(cur, add, remove []string) []string { + drop := make(map[string]bool, len(remove)) + for _, id := range remove { + drop[id] = true + } + + seen := map[string]bool{} + out := make([]string, 0, len(cur)+len(add)) + + for _, list := range [][]string{cur, add} { + for _, id := range list { + if drop[id] || seen[id] { + continue + } + + seen[id] = true + + out = append(out, id) + } + } + + return out +} + +// setEndpointRouteTables sets ep's route tables to ids (deduplicated). A +// Gateway endpoint cannot take a table that already routes the same service +// through another endpoint: EC2 allows one endpoint route per service per +// route table and answers RouteAlreadyExists. The caller holds m.mu and syncs +// the routes afterwards. +func (m *Mock) setEndpointRouteTables(ep *driver.VPCEndpoint, ids []string) error { + ids = applyIDDelta(nil, ids, nil) + + if plID := endpointPrefixList(ep); plID != "" { + for _, rtID := range ids { + if m.routeTableHasOtherEndpointRoute(rtID, plID, ep.ID) { + return errors.Newf(errors.AlreadyExists, + "route table %s already has a route with destination-prefix-list-id %s", rtID, plID) + } + } + } + + ep.RouteTableIDs = ids + + return nil +} + +// routeTableHasOtherEndpointRoute reports whether rtID routes plID to an +// endpoint other than endpointID. The caller holds m.mu. +func (m *Mock) routeTableHasOtherEndpointRoute(rtID, plID, endpointID string) bool { + rt, ok := m.routeTables.Get(rtID) + if !ok { + return false + } + + for _, r := range rt.Routes { + if r.DestinationPrefixListID == plID && r.TargetID != endpointID { + return true + } + } + + return false +} + +// syncEndpointENIs gives an Interface endpoint exactly one ENI in each of +// subnets, releasing the ENIs of subnets it left. The caller holds m.mu. +func (m *Mock) syncEndpointENIs(ep *driver.VPCEndpoint, subnets []string) { + want := make(map[string]bool, len(subnets)) + for _, s := range subnets { + want[s] = true + } + + desc := endpointENIDescription(ep.ID) + have := map[string]bool{} + + var ids []string + + for id, eni := range m.enis.All() { + if eni.Description != desc { + continue + } + + if !want[eni.SubnetID] || have[eni.SubnetID] { + m.enis.Delete(id) + continue + } + + have[eni.SubnetID] = true + + ids = append(ids, id) + } + + for _, s := range subnets { + if !have[s] { + ids = append(ids, m.attachManagedENI(ep.VPCID, s, desc).ID) + } + } + + sort.Strings(ids) + ep.NetworkInterfaceIDs = ids +} + +// dropRouteTableFromEndpoints removes a deleted route table from every +// endpoint that listed it. The caller holds m.mu. +func (m *Mock) dropRouteTableFromEndpoints(rtID string) { + for _, ep := range m.endpoints.All() { + ep.RouteTableIDs = applyIDDelta(ep.RouteTableIDs, nil, []string{rtID}) + } +} diff --git a/providers/aws/vpc/endpoint_sets_test.go b/providers/aws/vpc/endpoint_sets_test.go new file mode 100644 index 000000000..05dc25cf4 --- /dev/null +++ b/providers/aws/vpc/endpoint_sets_test.go @@ -0,0 +1,227 @@ +package vpc + +import ( + "context" + "sync" + "testing" + + "github.com/stackshy/cloudemu/v2/errors" + "github.com/stackshy/cloudemu/v2/services/networking/driver" +) + +func newGatewayEndpoint(t *testing.T, m *Mock, vpcID, service string, rts ...string) *driver.VPCEndpoint { + t.Helper() + + ep, err := m.CreateVPCEndpoint(context.Background(), driver.VPCEndpointConfig{ + VPCID: vpcID, ServiceName: service, EndpointType: "Gateway", RouteTableIDs: rts, + }) + requireNoError(t, err) + + return ep +} + +func newRouteTables(t *testing.T, m *Mock, vpcID string, n int) []string { + t.Helper() + + ids := make([]string, 0, n) + + for range n { + rt, err := m.CreateRouteTable(context.Background(), driver.RouteTableConfig{VPCID: vpcID}) + requireNoError(t, err) + + ids = append(ids, rt.ID) + } + + return ids +} + +// TestModifyVPCEndpointSetsConcurrent pins that concurrent AddRouteTableId +// changes on one endpoint all land. Each change is applied as a delta under the +// provider lock, so no caller's read-modify-write drops another's table. +func TestModifyVPCEndpointSetsConcurrent(t *testing.T) { + ctx := context.Background() + m := newTestMock() + v := createTestVPC(m) + rts := newRouteTables(t, m, v.ID, 12) + ep := newGatewayEndpoint(t, m, v.ID, "com.amazonaws.us-east-1.s3") + + var wg sync.WaitGroup + + for _, rt := range rts { + wg.Add(1) + + go func(rt string) { + defer wg.Done() + + _, err := m.ModifyVPCEndpointSets(ctx, ep.ID, &driver.VPCEndpointSetChange{AddRouteTableIDs: []string{rt}}) + if err != nil { + t.Errorf("add %s: %v", rt, err) + } + }(rt) + } + + wg.Wait() + + got, err := m.DescribeVPCEndpoints(ctx, []string{ep.ID}) + requireNoError(t, err) + assertEqual(t, len(rts), len(got[0].RouteTableIDs)) + + for _, rt := range rts { + assertEqual(t, 1, len(routesTo(t, m, rt, ep.ID))) + } + + for _, rt := range rts { + wg.Add(1) + + go func(rt string) { + defer wg.Done() + + if _, err := m.ModifyVPCEndpointSets(ctx, ep.ID, &driver.VPCEndpointSetChange{RemoveRouteTableIDs: []string{rt}}); err != nil { + t.Errorf("remove %s: %v", rt, err) + } + }(rt) + } + + wg.Wait() + + got, err = m.DescribeVPCEndpoints(ctx, []string{ep.ID}) + requireNoError(t, err) + assertEqual(t, 0, len(got[0].RouteTableIDs)) + + for _, rt := range rts { + assertEqual(t, 0, len(routesTo(t, m, rt, ep.ID))) + } +} + +// TestModifyVPCEndpointSetsSubnetsAndGroups covers the subnet and security +// group deltas: an Interface endpoint gains or loses one ENI per subnet. +func TestModifyVPCEndpointSetsSubnetsAndGroups(t *testing.T) { + ctx := context.Background() + m := newTestMock() + v := createTestVPC(m) + + subA, err := m.CreateSubnet(ctx, driver.SubnetConfig{VPCID: v.ID, CIDRBlock: "10.0.1.0/24"}) + requireNoError(t, err) + subB, err := m.CreateSubnet(ctx, driver.SubnetConfig{VPCID: v.ID, CIDRBlock: "10.0.2.0/24"}) + requireNoError(t, err) + + ep, err := m.CreateVPCEndpoint(ctx, driver.VPCEndpointConfig{ + VPCID: v.ID, ServiceName: "com.amazonaws.us-east-1.ssm", EndpointType: "Interface", + SubnetIDs: []string{subA.ID}, SecurityGroupIDs: []string{"sg-a"}, + }) + requireNoError(t, err) + + out, err := m.ModifyVPCEndpointSets(ctx, ep.ID, &driver.VPCEndpointSetChange{ + AddSubnetIDs: []string{subB.ID}, RemoveSubnetIDs: []string{subA.ID}, + AddSecurityGroupIDs: []string{"sg-b"}, RemoveSecurityGroupIDs: []string{"sg-a"}, + }) + requireNoError(t, err) + + assertEqual(t, 1, len(out.SubnetIDs)) + assertEqual(t, subB.ID, out.SubnetIDs[0]) + assertEqual(t, 1, len(out.SecurityGroupIDs)) + assertEqual(t, "sg-b", out.SecurityGroupIDs[0]) + assertEqual(t, 1, len(out.NetworkInterfaceIDs)) + assertEqual(t, 0, countENIsInSubnet(m, subA.ID)) + assertEqual(t, 1, countENIsInSubnet(m, subB.ID)) + + if _, err := m.ModifyVPCEndpointSets(ctx, "vpce-missing", &driver.VPCEndpointSetChange{}); !errors.IsNotFound(err) { + t.Fatalf("unknown endpoint: want NotFound, got %v", err) + } +} + +// TestGatewayEndpointOneRoutePerServicePerTable pins the EC2 rule that a route +// table holds at most one endpoint route per service: a second s3 endpoint on +// the same table is refused with AlreadyExists, on create and on modify. +func TestGatewayEndpointOneRoutePerServicePerTable(t *testing.T) { + ctx := context.Background() + m := newTestMock() + v := createTestVPC(m) + rts := newRouteTables(t, m, v.ID, 2) + + first := newGatewayEndpoint(t, m, v.ID, "com.amazonaws.us-east-1.s3", rts[0]) + + _, err := m.CreateVPCEndpoint(ctx, driver.VPCEndpointConfig{ + VPCID: v.ID, ServiceName: "com.amazonaws.us-east-1.s3", EndpointType: "Gateway", RouteTableIDs: []string{rts[0]}, + }) + if !errors.IsAlreadyExists(err) { + t.Fatalf("second s3 endpoint on the same table: want AlreadyExists, got %v", err) + } + + eps, err := m.DescribeVPCEndpoints(ctx, nil) + requireNoError(t, err) + assertEqual(t, 1, len(eps)) + + // dynamodb on the same table is fine, and so is s3 on another table. + newGatewayEndpoint(t, m, v.ID, "com.amazonaws.us-east-1.dynamodb", rts[0]) + second := newGatewayEndpoint(t, m, v.ID, "com.amazonaws.us-east-1.s3", rts[1]) + + _, err = m.ModifyVPCEndpointSets(ctx, second.ID, &driver.VPCEndpointSetChange{AddRouteTableIDs: []string{rts[0]}}) + if !errors.IsAlreadyExists(err) { + t.Fatalf("modify adding a table with an s3 route: want AlreadyExists, got %v", err) + } + + assertEqual(t, 1, len(routesTo(t, m, rts[0], first.ID))) + assertEqual(t, 0, len(routesTo(t, m, rts[0], second.ID))) +} + +// TestDeleteRouteTableLeavesGatewayEndpoint pins that deleting a route table +// drops it from the endpoints that used it. +func TestDeleteRouteTableLeavesGatewayEndpoint(t *testing.T) { + ctx := context.Background() + m := newTestMock() + v := createTestVPC(m) + rts := newRouteTables(t, m, v.ID, 2) + ep := newGatewayEndpoint(t, m, v.ID, "com.amazonaws.us-east-1.s3", rts...) + + requireNoError(t, m.DeleteRouteTable(ctx, rts[0])) + + got, err := m.DescribeVPCEndpoints(ctx, []string{ep.ID}) + requireNoError(t, err) + assertEqual(t, 1, len(got[0].RouteTableIDs)) + assertEqual(t, rts[1], got[0].RouteTableIDs[0]) +} + +// TestCreateDeleteEndpointConcurrentRoutes runs creates and deletes in +// parallel and checks that no endpoint route outlives its endpoint. +func TestCreateDeleteEndpointConcurrentRoutes(t *testing.T) { + ctx := context.Background() + m := newTestMock() + v := createTestVPC(m) + rts := newRouteTables(t, m, v.ID, 8) + + var wg sync.WaitGroup + + for _, rt := range rts { + wg.Add(1) + + go func(rt string) { + defer wg.Done() + + ep, err := m.CreateVPCEndpoint(ctx, driver.VPCEndpointConfig{ + VPCID: v.ID, ServiceName: "com.amazonaws.us-east-1.s3", EndpointType: "Gateway", RouteTableIDs: []string{rt}, + }) + if err != nil { + t.Errorf("create: %v", err) + return + } + + if err := m.DeleteVPCEndpoint(ctx, ep.ID); err != nil { + t.Errorf("delete: %v", err) + } + }(rt) + } + + wg.Wait() + + tables, err := m.DescribeRouteTables(ctx, rts) + requireNoError(t, err) + + for _, rt := range tables { + for _, r := range rt.Routes { + if r.DestinationPrefixListID != "" { + t.Errorf("orphan endpoint route %+v in %s", r, rt.ID) + } + } + } +} diff --git a/providers/aws/vpc/route_table.go b/providers/aws/vpc/route_table.go index b7bc56ed5..6b31858ab 100644 --- a/providers/aws/vpc/route_table.go +++ b/providers/aws/vpc/route_table.go @@ -95,7 +95,12 @@ func (m *Mock) DeleteRouteTable(_ context.Context, id string) error { } } + m.mu.Lock() m.routeTables.Delete(id) + // A Gateway endpoint that used the table loses it, the same as a + // ModifyVpcEndpoint RemoveRouteTableId. + m.dropRouteTableFromEndpoints(id) + m.mu.Unlock() return nil } diff --git a/providers/aws/vpc/service_prefix_list.go b/providers/aws/vpc/service_prefix_list.go index 1a4abbcb7..fd0f04c22 100644 --- a/providers/aws/vpc/service_prefix_list.go +++ b/providers/aws/vpc/service_prefix_list.go @@ -122,7 +122,8 @@ func (m *Mock) DescribePrefixLists( } // DescribeAWSManagedPrefixLists returns the service lists in the managed -// prefix list shape. Unknown ids are skipped. +// prefix list shape. Unknown ids are skipped. MaxEntries and Version stay +// zero: EC2 reports neither for an AWS-owned list. func (m *Mock) DescribeAWSManagedPrefixLists( _ context.Context, region string, ids []string, ) ([]driver.PrefixList, error) { @@ -145,8 +146,7 @@ func (m *Mock) DescribeAWSManagedPrefixLists( out = append(out, driver.PrefixList{ ID: pl.ID, Name: pl.Name, AddressFamily: "IPv4", - MaxEntries: len(entries), State: "create-complete", Version: 1, - Entries: entries, OwnerID: servicePrefixListOwner, + State: "create-complete", Entries: entries, OwnerID: servicePrefixListOwner, }) } @@ -158,10 +158,7 @@ func (m *Mock) DescribeAWSManagedPrefixLists( // from tables the endpoint no longer uses. Tables that do not exist are // skipped. Other endpoint types hold no routes. The caller holds m.mu. func (m *Mock) syncEndpointRoutes(ep *driver.VPCEndpoint) { - var plID string - if ep.EndpointType == "" || ep.EndpointType == vpcEndpointTypeGateway { - plID = servicePrefixListForEndpoint(ep.ServiceName) - } + plID := endpointPrefixList(ep) want := map[string]bool{} @@ -176,6 +173,16 @@ func (m *Mock) syncEndpointRoutes(ep *driver.VPCEndpoint) { } } +// endpointPrefixList returns the pl- id a Gateway endpoint routes to, or "" +// for other endpoint types and services without a prefix list. +func endpointPrefixList(ep *driver.VPCEndpoint) string { + if ep.EndpointType != "" && ep.EndpointType != vpcEndpointTypeGateway { + return "" + } + + return servicePrefixListForEndpoint(ep.ServiceName) +} + // endpointRoutes returns routes with at most one route to endpointID, kept // (or added, pointing at plID) only when keep is set. func endpointRoutes(routes []driver.Route, endpointID, plID string, keep bool) []driver.Route { diff --git a/providers/aws/vpc/service_prefix_list_test.go b/providers/aws/vpc/service_prefix_list_test.go index cc3836b64..be884ffad 100644 --- a/providers/aws/vpc/service_prefix_list_test.go +++ b/providers/aws/vpc/service_prefix_list_test.go @@ -98,6 +98,9 @@ func TestAWSManagedPrefixListsShape(t *testing.T) { assertEqual(t, "IPv4", pl.AddressFamily) assertEqual(t, svc[i].ID, pl.ID) assertEqual(t, len(svc[i].CIDRs), len(pl.Entries)) + // AWS-owned lists carry no maxEntries or version on the wire. + assertEqual(t, 0, pl.MaxEntries) + assertEqual(t, 0, pl.Version) } none, err := m.DescribeAWSManagedPrefixLists(ctx, "us-east-1", []string{"pl-00000000"}) diff --git a/server/aws/ec2/endpoint.go b/server/aws/ec2/endpoint.go index cf983b56f..549631929 100644 --- a/server/aws/ec2/endpoint.go +++ b/server/aws/ec2/endpoint.go @@ -128,6 +128,26 @@ func (h *Handler) describeVPCEndpoints(w http.ResponseWriter, r *http.Request) { func (h *Handler) modifyVPCEndpoint(w http.ResponseWriter, r *http.Request) { id := r.Form.Get("VpcEndpointId") + // A backend that applies the change as a delta does the read-modify-write + // under its own lock, so parallel modifies of one endpoint do not race. + if sets, ok := h.vpc.(netdriver.VPCEndpointSetModifier); ok { + if _, err := sets.ModifyVPCEndpointSets(r.Context(), id, &netdriver.VPCEndpointSetChange{ + AddRouteTableIDs: awsquery.ListStrings(r.Form, "AddRouteTableId"), + RemoveRouteTableIDs: awsquery.ListStrings(r.Form, "RemoveRouteTableId"), + AddSubnetIDs: awsquery.ListStrings(r.Form, "AddSubnetId"), + RemoveSubnetIDs: awsquery.ListStrings(r.Form, "RemoveSubnetId"), + AddSecurityGroupIDs: awsquery.ListStrings(r.Form, "AddSecurityGroupId"), + RemoveSecurityGroupIDs: awsquery.ListStrings(r.Form, "RemoveSecurityGroupId"), + }); err != nil { + writeVPCEndpointErr(w, err) + return + } + + writeReturnTrue(w, "ModifyVpcEndpointResponse") + + return + } + current, err := h.vpc.DescribeVPCEndpoints(r.Context(), []string{id}) if err != nil { writeVPCEndpointErr(w, err) @@ -231,5 +251,12 @@ func toVPCEndpointXML(ep *netdriver.VPCEndpoint) vpcEndpointXML { } func writeVPCEndpointErr(w http.ResponseWriter, err error) { + // The only AlreadyExists an endpoint call raises is a second endpoint route + // for the same service in one route table. + if cerrors.IsAlreadyExists(err) { + awsquery.WriteXMLError(w, http.StatusBadRequest, "RouteAlreadyExists", cerrors.Message(err)) + return + } + writeErrWithNotFound(w, err, codeInvalidVpcEndpointID, "DependencyViolation") } diff --git a/server/aws/ec2/endpoint_modify_test.go b/server/aws/ec2/endpoint_modify_test.go new file mode 100644 index 000000000..7b60815fb --- /dev/null +++ b/server/aws/ec2/endpoint_modify_test.go @@ -0,0 +1,106 @@ +package ec2_test + +import ( + "context" + "sync" + "testing" + + "github.com/aws/aws-sdk-go-v2/aws" + "github.com/aws/aws-sdk-go-v2/service/ec2" + ec2types "github.com/aws/aws-sdk-go-v2/service/ec2/types" +) + +func mkGatewayEndpoint(t *testing.T, c *ec2.Client, vpcID string, rts ...string) (string, error) { + t.Helper() + + ep, err := c.CreateVpcEndpoint(context.Background(), &ec2.CreateVpcEndpointInput{ + VpcId: aws.String(vpcID), ServiceName: aws.String(s3PrefixListName), + VpcEndpointType: ec2types.VpcEndpointTypeGateway, RouteTableIds: rts, + }) + if err != nil { + return "", err + } + + return aws.ToString(ep.VpcEndpoint.VpcEndpointId), nil +} + +// TestModifyVpcEndpointConcurrentAddRouteTable pins that parallel +// AddRouteTableId calls on one endpoint all land. Terraform creates several +// aws_vpc_endpoint_route_table_association resources at once, and a lost update +// here left tables (and their pl- routes) missing. +func TestModifyVpcEndpointConcurrentAddRouteTable(t *testing.T) { + ctx := context.Background() + c := newRoutingEdgeEC2(t) + vpcID, _ := mkVPCSubnet(t, c) + + rts := make([]string, 0, 6) + for range 6 { + rts = append(rts, mkRouteTable(t, c, vpcID)) + } + + vpceID, err := mkGatewayEndpoint(t, c, vpcID) + if err != nil { + t.Fatalf("CreateVpcEndpoint: %v", err) + } + + var wg sync.WaitGroup + + for _, rt := range rts { + wg.Add(1) + + go func(rt string) { + defer wg.Done() + + if _, err := c.ModifyVpcEndpoint(ctx, &ec2.ModifyVpcEndpointInput{ + VpcEndpointId: aws.String(vpceID), AddRouteTableIds: []string{rt}, + }); err != nil { + t.Errorf("ModifyVpcEndpoint add %s: %v", rt, err) + } + }(rt) + } + + wg.Wait() + + out, err := c.DescribeVpcEndpoints(ctx, &ec2.DescribeVpcEndpointsInput{VpcEndpointIds: []string{vpceID}}) + if err != nil { + t.Fatalf("DescribeVpcEndpoints: %v", err) + } + + if got := len(out.VpcEndpoints[0].RouteTableIds); got != len(rts) { + t.Fatalf("endpoint has %d route tables, want %d", got, len(rts)) + } + + for _, rt := range rts { + if findVPCERoute(t, c, rt, vpceID) == nil { + t.Errorf("route table %s has no route to %s", rt, vpceID) + } + } +} + +// TestGatewayEndpointDuplicateServiceRouteRejected pins the EC2 rule that a +// route table holds one endpoint route per service: a second s3 Gateway +// endpoint on the same table fails with RouteAlreadyExists. +func TestGatewayEndpointDuplicateServiceRouteRejected(t *testing.T) { + ctx := context.Background() + c := newRoutingEdgeEC2(t) + vpcID, _ := mkVPCSubnet(t, c) + rtA := mkRouteTable(t, c, vpcID) + rtB := mkRouteTable(t, c, vpcID) + + if _, err := mkGatewayEndpoint(t, c, vpcID, rtA); err != nil { + t.Fatalf("first endpoint: %v", err) + } + + _, err := mkGatewayEndpoint(t, c, vpcID, rtA) + requireAPIErrorCode(t, err, "RouteAlreadyExists") + + second, err := mkGatewayEndpoint(t, c, vpcID, rtB) + if err != nil { + t.Fatalf("s3 endpoint on another table: %v", err) + } + + _, err = c.ModifyVpcEndpoint(ctx, &ec2.ModifyVpcEndpointInput{ + VpcEndpointId: aws.String(second), AddRouteTableIds: []string{rtA}, + }) + requireAPIErrorCode(t, err, "RouteAlreadyExists") +} diff --git a/server/aws/ec2/prefix_list.go b/server/aws/ec2/prefix_list.go index 10ee97a5a..40f06aaf1 100644 --- a/server/aws/ec2/prefix_list.go +++ b/server/aws/ec2/prefix_list.go @@ -21,9 +21,9 @@ type prefixListXML struct { PrefixListArn string `xml:"prefixListArn,omitempty"` PrefixListName string `xml:"prefixListName"` AddressFamily string `xml:"addressFamily"` - MaxEntries int `xml:"maxEntries"` + MaxEntries *int `xml:"maxEntries,omitempty"` State string `xml:"state"` - Version int `xml:"version"` + Version *int `xml:"version,omitempty"` OwnerID string `xml:"ownerId,omitempty"` Tags []tagItem `xml:"tagSet>item,omitempty"` } @@ -178,12 +178,18 @@ func (h *Handler) getPrefixListEntries(w http.ResponseWriter, r *http.Request, p out = append(out, prefixListEntryXML{Cidr: entries[i].CIDR, Description: entries[i].Description}) } + // Entries keep their list order; the cidr is unique within a list, so it + // doubles as the page token key. + page, next := paginateXML(out, r.Form.Get("MaxResults"), r.Form.Get("NextToken"), + func(e prefixListEntryXML) string { return e.Cidr }) + awsquery.WriteXMLResponse(w, struct { XMLName xml.Name `xml:"GetManagedPrefixListEntriesResponse"` Xmlns string `xml:"xmlns,attr"` Req string `xml:"requestId"` Set []prefixListEntryXML `xml:"entrySet>item"` - }{Xmlns: awsquery.Namespace, Req: awsquery.RequestID, Set: out}) + Next string `xml:"nextToken,omitempty"` + }{Xmlns: awsquery.Namespace, Req: awsquery.RequestID, Set: page, Next: next}) } func (h *Handler) modifyPrefixList(w http.ResponseWriter, r *http.Request, p netdriver.PrefixLists) { @@ -314,14 +320,19 @@ func parsePrefixListEntries(r *http.Request) []netdriver.PrefixListEntry { } func (h *Handler) toPrefixListXML(region string, p *netdriver.PrefixList) prefixListXML { - owner := nonEmpty(p.OwnerID, h.accountID) + x := prefixListXML{ + PrefixListID: p.ID, PrefixListName: p.Name, AddressFamily: p.AddressFamily, + State: p.State, OwnerID: nonEmpty(p.OwnerID, h.accountID), Tags: toTagItems(p.Tags), + } + x.PrefixListArn = prefixListARN(region, x.OwnerID, p.ID) - return prefixListXML{ - PrefixListID: p.ID, PrefixListArn: prefixListARN(region, owner, p.ID), - PrefixListName: p.Name, AddressFamily: p.AddressFamily, - MaxEntries: p.MaxEntries, State: p.State, Version: p.Version, - OwnerID: owner, Tags: toTagItems(p.Tags), + // AWS-owned lists carry no maxEntries or version; customer lists always do. + if p.OwnerID == "" { + maxEntries, version := p.MaxEntries, p.Version + x.MaxEntries, x.Version = &maxEntries, &version } + + return x } // prefixListARN builds the managed-prefix-list ARN AWS returns; the SDK and diff --git a/server/aws/ec2/service_prefix_list_test.go b/server/aws/ec2/service_prefix_list_test.go index 0ec3f1e71..dadc4853c 100644 --- a/server/aws/ec2/service_prefix_list_test.go +++ b/server/aws/ec2/service_prefix_list_test.go @@ -257,6 +257,32 @@ func TestManagedPrefixListsIncludeAWSOwned(t *testing.T) { t.Error("AWS-owned list has no entries") } + if pl.MaxEntries != nil || pl.Version != nil { + t.Errorf("AWS-owned list carries maxEntries %v / version %v, want neither", pl.MaxEntries, pl.Version) + } + + page1, err := c.GetManagedPrefixListEntries(ctx, &ec2.GetManagedPrefixListEntriesInput{ + PrefixListId: pl.PrefixListId, MaxResults: aws.Int32(5), + }) + if err != nil { + t.Fatalf("GetManagedPrefixListEntries page 1: %v", err) + } + + if len(page1.Entries) != 5 || page1.NextToken == nil { + t.Fatalf("page 1 = %d entries, next %v; want 5 and a token", len(page1.Entries), page1.NextToken) + } + + page2, err := c.GetManagedPrefixListEntries(ctx, &ec2.GetManagedPrefixListEntriesInput{ + PrefixListId: pl.PrefixListId, MaxResults: aws.Int32(5), NextToken: page1.NextToken, + }) + if err != nil { + t.Fatalf("GetManagedPrefixListEntries page 2: %v", err) + } + + if len(page1.Entries)+len(page2.Entries) != len(entries.Entries) || page2.NextToken != nil { + t.Fatalf("page 2 = %d entries, next %v", len(page2.Entries), page2.NextToken) + } + _, err = c.DescribeManagedPrefixLists(ctx, &ec2.DescribeManagedPrefixListsInput{PrefixListIds: []string{"pl-00000000"}}) requireAPIErrorCode(t, err, "InvalidPrefixListID.NotFound") } diff --git a/services/networking/driver/aws_capabilities.go b/services/networking/driver/aws_capabilities.go index a77a8ecb4..fdb36bafd 100644 --- a/services/networking/driver/aws_capabilities.go +++ b/services/networking/driver/aws_capabilities.go @@ -270,6 +270,24 @@ type ServicePrefixLists interface { DescribeAWSManagedPrefixLists(ctx context.Context, region string, ids []string) ([]PrefixList, error) } +// VPCEndpointSetChange is the Add*/Remove* id-set part of ModifyVpcEndpoint. +// An id named in both an Add and its Remove list is dropped. +type VPCEndpointSetChange struct { + AddRouteTableIDs []string + RemoveRouteTableIDs []string + AddSubnetIDs []string + RemoveSubnetIDs []string + AddSecurityGroupIDs []string + RemoveSecurityGroupIDs []string +} + +// VPCEndpointSetModifier is an OPTIONAL AWS capability (type-asserted). It +// applies a ModifyVpcEndpoint set change to the endpoint's current sets in one +// step, so concurrent modifies of one endpoint do not overwrite each other. +type VPCEndpointSetModifier interface { + ModifyVPCEndpointSets(ctx context.Context, id string, change *VPCEndpointSetChange) (*VPCEndpoint, error) +} + // ---- Egress-only Internet Gateway (IPv6) ---- // EgressOnlyInternetGateway provides outbound-only IPv6 for private subnets.