From 8e987a47690692fdcf4790876d5cc79bd33eaa8d Mon Sep 17 00:00:00 2001 From: Nitin Kumar Date: Sun, 27 Sep 2026 18:20:17 +0530 Subject: [PATCH 1/2] fix(aws-apigatewayv2): deployments, validation, tags, pagination and quick-create --- docs/coverage/README.md | 2 +- docs/coverage/aws/README.md | 2 +- docs/coverage/aws/apigatewayv2.md | 10 +- docs/coverage/coverage.json | 24 ++ providers/aws/apigatewayv2/apigatewayv2.go | 74 +++-- .../aws/apigatewayv2/apigatewayv2_test.go | 21 +- providers/aws/apigatewayv2/deployment_test.go | 208 ++++++++++++++ providers/aws/apigatewayv2/deployments.go | 245 +++++++++++++++++ providers/aws/apigatewayv2/integrations.go | 67 +++-- providers/aws/apigatewayv2/page.go | 66 +++++ providers/aws/apigatewayv2/quickcreate.go | 186 +++++++++++++ providers/aws/apigatewayv2/routes.go | 102 +++++-- providers/aws/apigatewayv2/snapshot.go | 14 +- providers/aws/apigatewayv2/snapshot_test.go | 47 ++++ providers/aws/apigatewayv2/stages.go | 121 +++++++-- providers/aws/apigatewayv2/tags.go | 178 ++++++++++++ providers/aws/apigatewayv2/validation.go | 253 ++++++++++++++++++ .../aws/apigatewayv2/deployment_e2e_test.go | 115 ++++++++ .../aws/apigatewayv2/deployments_handler.go | 67 +++++ server/aws/apigatewayv2/handler.go | 54 ++-- server/aws/apigatewayv2/helpers_e2e_test.go | 109 ++++++++ .../aws/apigatewayv2/integrations_handler.go | 6 +- .../aws/apigatewayv2/pagination_e2e_test.go | 93 +++++++ .../aws/apigatewayv2/quickcreate_e2e_test.go | 88 ++++++ server/aws/apigatewayv2/routes_handler.go | 5 +- server/aws/apigatewayv2/stages_handler.go | 5 +- server/aws/apigatewayv2/tags_e2e_test.go | 79 ++++++ server/aws/apigatewayv2/tags_handler.go | 50 ++++ server/aws/apigatewayv2/types.go | 84 +++++- .../aws/apigatewayv2/validation_e2e_test.go | 107 ++++++++ services/apigatewayv2/driver/driver.go | 96 ++++++- 31 files changed, 2432 insertions(+), 146 deletions(-) create mode 100644 providers/aws/apigatewayv2/deployment_test.go create mode 100644 providers/aws/apigatewayv2/deployments.go create mode 100644 providers/aws/apigatewayv2/page.go create mode 100644 providers/aws/apigatewayv2/quickcreate.go create mode 100644 providers/aws/apigatewayv2/tags.go create mode 100644 providers/aws/apigatewayv2/validation.go create mode 100644 server/aws/apigatewayv2/deployment_e2e_test.go create mode 100644 server/aws/apigatewayv2/deployments_handler.go create mode 100644 server/aws/apigatewayv2/helpers_e2e_test.go create mode 100644 server/aws/apigatewayv2/pagination_e2e_test.go create mode 100644 server/aws/apigatewayv2/quickcreate_e2e_test.go create mode 100644 server/aws/apigatewayv2/tags_e2e_test.go create mode 100644 server/aws/apigatewayv2/tags_handler.go create mode 100644 server/aws/apigatewayv2/validation_e2e_test.go diff --git a/docs/coverage/README.md b/docs/coverage/README.md index 96c4e64b3..37ea12b4f 100644 --- a/docs/coverage/README.md +++ b/docs/coverage/README.md @@ -17,7 +17,7 @@ code does not implement. Machine-readable: [`coverage.json`](./coverage.json). | `aoss` | [AOSS](./aws/aoss.md) | - | - | - | 18 | | `apigateway` | [APIGateway](./aws/apigateway.md) | - | - | - | 29 | | `apigatewaygcp` | - | - | [APIGateway](./gcp/apigateway.md) | - | 16 | -| `apigatewayv2` | [APIGatewayV2](./aws/apigatewayv2.md) | - | - | - | 20 | +| `apigatewayv2` | [APIGatewayV2](./aws/apigatewayv2.md) | - | - | - | 28 | | `appconfiguration` | - | [AppConfiguration](./azure/appconfiguration.md) | - | - | 9 | | `appflow` | [AppFlow](./aws/appflow.md) | - | - | - | 14 | | `appinsights` | - | [Appinsights](./azure/appinsights.md) | - | - | 6 | diff --git a/docs/coverage/aws/README.md b/docs/coverage/aws/README.md index 61790088a..b9dcbed94 100644 --- a/docs/coverage/aws/README.md +++ b/docs/coverage/aws/README.md @@ -8,7 +8,7 @@ Services cloudemu emulates for AWS, by native name. Back to the [cross-provider | [ACM](./acm.md) | `acm` | 17 | | [AOSS](./aoss.md) | `aoss` | 18 | | [APIGateway](./apigateway.md) | `apigateway` | 29 | -| [APIGatewayV2](./apigatewayv2.md) | `apigatewayv2` | 20 | +| [APIGatewayV2](./apigatewayv2.md) | `apigatewayv2` | 28 | | [APS](./aps.md) | `aps` | 21 | | [AppFlow](./appflow.md) | `appflow` | 14 | | [AppRunner](./apprunner.md) | `apprunner` | 29 | diff --git a/docs/coverage/aws/apigatewayv2.md b/docs/coverage/aws/apigatewayv2.md index 0a29e648f..e656200f7 100644 --- a/docs/coverage/aws/apigatewayv2.md +++ b/docs/coverage/aws/apigatewayv2.md @@ -3,27 +3,35 @@ AWS's `apigatewayv2` service · portable interface `driver.APIGatewayV2` · [AWS index](./README.md) -## Operations (20) +## Operations (28) | Operation | Description | | --- | --- | | `CreateAPI` | | +| `CreateDeployment` | | | `CreateIntegration` | | | `CreateRoute` | | | `CreateStage` | | | `DeleteAPI` | | +| `DeleteDeployment` | | | `DeleteIntegration` | | | `DeleteRoute` | | | `DeleteStage` | | | `GetAPI` | | | `GetAPIs` | | +| `GetDeployment` | | +| `GetDeployments` | | | `GetIntegration` | | | `GetIntegrations` | | | `GetRoute` | | | `GetRoutes` | | | `GetStage` | | | `GetStages` | | +| `GetTags` | | +| `TagResource` | | +| `UntagResource` | | | `UpdateAPI` | | +| `UpdateDeployment` | | | `UpdateIntegration` | | | `UpdateRoute` | | | `UpdateStage` | | diff --git a/docs/coverage/coverage.json b/docs/coverage/coverage.json index 6459156fb..697e03ae4 100644 --- a/docs/coverage/coverage.json +++ b/docs/coverage/coverage.json @@ -432,6 +432,9 @@ { "name": "CreateAPI" }, + { + "name": "CreateDeployment" + }, { "name": "CreateIntegration" }, @@ -444,6 +447,9 @@ { "name": "DeleteAPI" }, + { + "name": "DeleteDeployment" + }, { "name": "DeleteIntegration" }, @@ -459,6 +465,12 @@ { "name": "GetAPIs" }, + { + "name": "GetDeployment" + }, + { + "name": "GetDeployments" + }, { "name": "GetIntegration" }, @@ -477,9 +489,21 @@ { "name": "GetStages" }, + { + "name": "GetTags" + }, + { + "name": "TagResource" + }, + { + "name": "UntagResource" + }, { "name": "UpdateAPI" }, + { + "name": "UpdateDeployment" + }, { "name": "UpdateIntegration" }, diff --git a/providers/aws/apigatewayv2/apigatewayv2.go b/providers/aws/apigatewayv2/apigatewayv2.go index bd47e9270..acd8c6d12 100644 --- a/providers/aws/apigatewayv2/apigatewayv2.go +++ b/providers/aws/apigatewayv2/apigatewayv2.go @@ -1,7 +1,7 @@ // Package apigatewayv2 is an in-memory mock of Amazon API Gateway v2 (HTTP and // WebSocket APIs). It models the control plane only: an API and its Route, -// Integration and Stage sub-collections, reachable over the apigatewayv2 -// REST/JSON protocol. It is a separate service from API Gateway REST v1 +// Integration, Stage and Deployment sub-collections plus resource tags, +// reachable over the apigatewayv2 REST/JSON protocol. It is a separate service from API Gateway REST v1 // (providers/aws/apigateway), sharing no state or types. package apigatewayv2 @@ -52,6 +52,7 @@ type apiData struct { routes map[string]*driver.Route integrations map[string]*driver.Integration stages map[string]*driver.Stage + deployments map[string]*deploymentRecord } // Mock is an in-memory implementation of Amazon API Gateway v2. @@ -96,14 +97,15 @@ func (m *Mock) getAPI(id string) (*apiData, error) { } // CreateAPI creates a new API with defaulted selection expressions and a -// computed execute-api endpoint. +// computed execute-api endpoint. A Target quick-creates the default +// integration, route and auto-deployed $default stage. func (m *Mock) CreateAPI(_ context.Context, in *driver.CreateAPIInput) (*driver.API, error) { - if in.Name == "" { - return nil, cerrors.New(cerrors.InvalidArgument, "Name is required") + if in.ProtocolType != driver.ProtocolHTTP && in.ProtocolType != driver.ProtocolWebSocket { + return nil, badRequest("Invalid protocol type specified: %s", in.ProtocolType) } - if in.ProtocolType != driver.ProtocolHTTP && in.ProtocolType != driver.ProtocolWebSocket { - return nil, cerrors.Newf(cerrors.InvalidArgument, "Invalid protocol type specified: %s", in.ProtocolType) + if in.ProtocolType == driver.ProtocolWebSocket && in.RouteSelectionExpression == "" { + return nil, badRequest("RouteSelectionExpression is required for WEBSOCKET protocol") } apiID := genID() @@ -113,24 +115,54 @@ func (m *Mock) CreateAPI(_ context.Context, in *driver.CreateAPIInput) (*driver. RouteSelectionExpression: orDefault(in.RouteSelectionExpression, defaultRouteSelectionExpr), APIKeySelectionExpression: orDefault(in.APIKeySelectionExpression, defaultAPIKeySelectionExpr), DisableExecuteAPIEndpoint: in.DisableExecuteAPIEndpoint, - APIEndpoint: fmt.Sprintf("https://%s.execute-api.%s.amazonaws.com", apiID, m.region), + APIEndpoint: m.apiEndpoint(apiID, in.ProtocolType), CreatedDate: m.now(), Tags: copyStrMap(in.Tags), CorsConfiguration: copyCors(in.CorsConfiguration), } - m.apis.Set(apiID, &apiData{ + if err := validateAPIFields(&api); err != nil { + return nil, err + } + + if err := validateTags(in.Tags); err != nil { + return nil, err + } + + if err := checkQuickCreate(in.ProtocolType, in.Target, in.RouteKey, in.CredentialsArn); err != nil { + return nil, err + } + + ad := &apiData{ api: api, routes: map[string]*driver.Route{}, integrations: map[string]*driver.Integration{}, stages: map[string]*driver.Stage{}, - }) + deployments: map[string]*deploymentRecord{}, + } + + if in.Target != "" { + m.quickCreate(ad, in.Target, in.RouteKey, in.CredentialsArn) + } + + m.apis.Set(apiID, ad) out := copyAPI(&api) return &out, nil } +// apiEndpoint is the execute-api endpoint of an API: https for HTTP APIs and +// wss for WebSocket APIs. +func (m *Mock) apiEndpoint(apiID, protocol string) string { + scheme := "https" + if protocol == driver.ProtocolWebSocket { + scheme = "wss" + } + + return fmt.Sprintf("%s://%s.execute-api.%s.amazonaws.com", scheme, apiID, m.region) +} + // GetAPI returns a single API. func (m *Mock) GetAPI(_ context.Context, apiID string) (*driver.API, error) { ad, err := m.getAPI(apiID) @@ -146,8 +178,8 @@ func (m *Mock) GetAPI(_ context.Context, apiID string) (*driver.API, error) { return &out, nil } -// GetAPIs lists all APIs. -func (m *Mock) GetAPIs(_ context.Context) ([]driver.API, error) { +// GetAPIs lists one page of APIs, ordered by id. +func (m *Mock) GetAPIs(_ context.Context, page *driver.PageInput) ([]driver.API, string, error) { all := m.apis.All() out := make([]driver.API, 0, len(all)) @@ -157,10 +189,11 @@ func (m *Mock) GetAPIs(_ context.Context) ([]driver.API, error) { ad.mu.RUnlock() } - return out, nil + return pageOf(out, func(a, b driver.API) bool { return a.APIID < b.APIID }, page) } -// UpdateAPI applies the non-nil fields of in to the stored API (PATCH). +// UpdateAPI applies the non-nil fields of in to the stored API (PATCH). The +// quick-create fields update the managed integration and route. func (m *Mock) UpdateAPI(_ context.Context, apiID string, in *driver.UpdateAPIInput) (*driver.API, error) { ad, err := m.getAPI(apiID) if err != nil { @@ -170,7 +203,7 @@ func (m *Mock) UpdateAPI(_ context.Context, apiID string, in *driver.UpdateAPIIn ad.mu.Lock() defer ad.mu.Unlock() - a := &ad.api + a := ad.api setString(&a.Name, in.Name) setString(&a.Description, in.Description) setString(&a.Version, in.Version) @@ -182,7 +215,16 @@ func (m *Mock) UpdateAPI(_ context.Context, apiID string, in *driver.UpdateAPIIn a.CorsConfiguration = copyCors(in.CorsConfiguration) } - out := copyAPI(a) + if err := validateAPIFields(&a); err != nil { + return nil, err + } + + if err := m.updateQuickCreate(ad, in); err != nil { + return nil, err + } + + ad.api = a + out := copyAPI(&a) return &out, nil } diff --git a/providers/aws/apigatewayv2/apigatewayv2_test.go b/providers/aws/apigatewayv2/apigatewayv2_test.go index b5d2f5e43..35abeb591 100644 --- a/providers/aws/apigatewayv2/apigatewayv2_test.go +++ b/providers/aws/apigatewayv2/apigatewayv2_test.go @@ -97,7 +97,7 @@ func TestAPICRUD(t *testing.T) { t.Fatalf("PATCH clobbered untouched fields: %+v", upd) } - apis, err := m.GetAPIs(ctx()) + apis, _, err := m.GetAPIs(ctx(), nil) if err != nil || len(apis) != 1 { t.Fatalf("GetAPIs: %v, len=%d", err, len(apis)) } @@ -124,13 +124,20 @@ func TestRouteCRUD(t *testing.T) { t.Fatalf("CreateRoute: %v, %+v", err, rt) } - target := "integrations/abc" + ig, err := m.CreateIntegration(ctx(), api.APIID, &driver.CreateIntegrationInput{ + IntegrationType: driver.IntegrationHTTPProxy, IntegrationURI: "https://example.com", IntegrationMethod: "GET", + }) + if err != nil { + t.Fatalf("CreateIntegration: %v", err) + } + + target := "integrations/" + ig.IntegrationID upd, err := m.UpdateRoute(ctx(), api.APIID, rt.RouteID, &driver.UpdateRouteInput{Target: &target}) if err != nil || upd.Target != target || upd.RouteKey != "GET /items" { t.Fatalf("UpdateRoute: %v, %+v", err, upd) } - routes, err := m.GetRoutes(ctx(), api.APIID) + routes, _, err := m.GetRoutes(ctx(), api.APIID, nil) if err != nil || len(routes) != 1 { t.Fatalf("GetRoutes: %v, len=%d", err, len(routes)) } @@ -164,7 +171,9 @@ func TestIntegrationCRUDAndTimeoutDefault(t *testing.T) { } // WebSocket API defaults the integration timeout to 29000. - wsAPI, err := m.CreateAPI(ctx(), &driver.CreateAPIInput{Name: "ws", ProtocolType: driver.ProtocolWebSocket}) + wsAPI, err := m.CreateAPI(ctx(), &driver.CreateAPIInput{ + Name: "ws", ProtocolType: driver.ProtocolWebSocket, RouteSelectionExpression: "$request.body.action", + }) if err != nil { t.Fatalf("CreateAPI ws: %v", err) } @@ -202,7 +211,7 @@ func TestStageCRUDAndConflict(t *testing.T) { t.Fatalf("UpdateStage: %v, %+v", err, upd) } - stages, err := m.GetStages(ctx(), api.APIID) + stages, _, err := m.GetStages(ctx(), api.APIID, nil) if err != nil || len(stages) != 1 { t.Fatalf("GetStages: %v, len=%d", err, len(stages)) } @@ -219,7 +228,7 @@ func TestStageCRUDAndConflict(t *testing.T) { func TestSubResourcesOnMissingAPI(t *testing.T) { m := newMock(t) - if _, err := m.GetRoutes(ctx(), "nope"); !cerrors.IsNotFound(err) { + if _, _, err := m.GetRoutes(ctx(), "nope", nil); !cerrors.IsNotFound(err) { t.Fatalf("GetRoutes missing api err = %v, want NotFound", err) } diff --git a/providers/aws/apigatewayv2/deployment_test.go b/providers/aws/apigatewayv2/deployment_test.go new file mode 100644 index 000000000..b913690bd --- /dev/null +++ b/providers/aws/apigatewayv2/deployment_test.go @@ -0,0 +1,208 @@ +package apigatewayv2_test + +import ( + "testing" + + cerrors "github.com/stackshy/cloudemu/v2/errors" + "github.com/stackshy/cloudemu/v2/providers/aws/apigatewayv2" + "github.com/stackshy/cloudemu/v2/services/apigatewayv2/driver" +) + +// addRoute creates an HTTP_PROXY integration and a route targeting it. +func addRoute(t *testing.T, m *apigatewayv2.Mock, apiID, key string) *driver.Route { + t.Helper() + + ig, err := m.CreateIntegration(ctx(), apiID, &driver.CreateIntegrationInput{ + IntegrationType: driver.IntegrationHTTPProxy, IntegrationURI: "https://example.com", IntegrationMethod: "ANY", + }) + if err != nil { + t.Fatalf("CreateIntegration: %v", err) + } + + rt, err := m.CreateRoute(ctx(), apiID, &driver.CreateRouteInput{RouteKey: key, Target: "integrations/" + ig.IntegrationID}) + if err != nil { + t.Fatalf("CreateRoute: %v", err) + } + + return rt +} + +func TestDeploymentRequiresRoutesAndStage(t *testing.T) { + m := newMock(t) + api := createHTTPAPI(t, m) + + if _, err := m.CreateDeployment(ctx(), api.APIID, &driver.CreateDeploymentInput{}); !cerrors.IsInvalidArgument(err) { + t.Fatalf("deploy with no routes err = %v, want InvalidArgument", err) + } + + addRoute(t, m, api.APIID, "GET /a") + + if _, err := m.CreateDeployment(ctx(), api.APIID, &driver.CreateDeploymentInput{StageName: "x"}); !cerrors.IsNotFound(err) { + t.Fatalf("deploy to missing stage err = %v, want NotFound", err) + } + + d, err := m.CreateDeployment(ctx(), api.APIID, &driver.CreateDeploymentInput{Description: "d"}) + if err != nil || d.DeploymentStatus != driver.DeploymentStatusDeployed || d.AutoDeployed { + t.Fatalf("CreateDeployment: %v, %+v", err, d) + } +} + +func TestStageDeploymentPointerRules(t *testing.T) { + m := newMock(t) + api := createHTTPAPI(t, m) + addRoute(t, m, api.APIID, "GET /a") + + d, err := m.CreateDeployment(ctx(), api.APIID, &driver.CreateDeploymentInput{}) + if err != nil { + t.Fatalf("CreateDeployment: %v", err) + } + + if _, err := m.CreateStage(ctx(), api.APIID, &driver.CreateStageInput{StageName: "p", DeploymentID: "nope"}); !cerrors.IsInvalidArgument(err) { + t.Fatalf("bogus deploymentId err = %v, want InvalidArgument", err) + } + + st, err := m.CreateStage(ctx(), api.APIID, &driver.CreateStageInput{StageName: "p", DeploymentID: d.DeploymentID}) + if err != nil || st.DeploymentID != d.DeploymentID { + t.Fatalf("CreateStage: %v, %+v", err, st) + } + + if err := m.DeleteDeployment(ctx(), api.APIID, d.DeploymentID); !cerrors.IsInvalidArgument(err) { + t.Fatalf("delete active deployment err = %v, want InvalidArgument", err) + } + + if err := m.DeleteStage(ctx(), api.APIID, "p"); err != nil { + t.Fatalf("DeleteStage: %v", err) + } + + if err := m.DeleteDeployment(ctx(), api.APIID, d.DeploymentID); err != nil { + t.Fatalf("DeleteDeployment: %v", err) + } +} + +func TestAutoDeployOnStageCreateAndToggle(t *testing.T) { + m := newMock(t) + api := createHTTPAPI(t, m) + addRoute(t, m, api.APIID, "GET /a") + + st, err := m.CreateStage(ctx(), api.APIID, &driver.CreateStageInput{StageName: "auto", AutoDeploy: true}) + if err != nil || st.DeploymentID == "" { + t.Fatalf("autoDeploy CreateStage: %v, %+v", err, st) + } + + if _, err := m.CreateStage(ctx(), api.APIID, &driver.CreateStageInput{StageName: "manual"}); err != nil { + t.Fatalf("CreateStage manual: %v", err) + } + + on := true + + upd, err := m.UpdateStage(ctx(), api.APIID, "manual", &driver.UpdateStageInput{AutoDeploy: &on}) + if err != nil || upd.DeploymentID == "" { + t.Fatalf("toggle autoDeploy: %v, %+v", err, upd) + } + + deps, _, err := m.GetDeployments(ctx(), api.APIID, nil) + if err != nil || len(deps) != 2 || !deps[0].AutoDeployed { + t.Fatalf("GetDeployments: %v, %+v", err, deps) + } +} + +func TestQuickCreateManagedRules(t *testing.T) { + m := newMock(t) + + api, err := m.CreateAPI(ctx(), &driver.CreateAPIInput{ + Name: "q", ProtocolType: driver.ProtocolHTTP, + Target: "arn:aws:lambda:us-east-1:000000000000:function:f", RouteKey: "GET /pets", + }) + if err != nil { + t.Fatalf("CreateAPI quick: %v", err) + } + + routes, _, _ := m.GetRoutes(ctx(), api.APIID, nil) + if len(routes) != 1 || !routes[0].APIGatewayManaged { + t.Fatalf("routes = %+v", routes) + } + + key := "GET /cats" + if _, err := m.UpdateRoute(ctx(), api.APIID, routes[0].RouteID, &driver.UpdateRouteInput{RouteKey: &key}); !cerrors.IsInvalidArgument(err) { + t.Fatalf("managed route key change err = %v, want InvalidArgument", err) + } + + desc := "x" + if _, err := m.UpdateStage(ctx(), api.APIID, "$default", &driver.UpdateStageInput{Description: &desc}); !cerrors.IsInvalidArgument(err) { + t.Fatalf("managed stage update err = %v, want InvalidArgument", err) + } + + // UpdateApi's RouteKey is the supported way to rename the managed route. + if _, err := m.UpdateAPI(ctx(), api.APIID, &driver.UpdateAPIInput{RouteKey: &key}); err != nil { + t.Fatalf("UpdateAPI routeKey: %v", err) + } + + got, _ := m.GetRoute(ctx(), api.APIID, routes[0].RouteID) + if got.RouteKey != key { + t.Fatalf("managed route key = %q, want %q", got.RouteKey, key) + } + + bad := "nope" + if _, err := m.UpdateAPI(ctx(), api.APIID, &driver.UpdateAPIInput{Target: &bad}); !cerrors.IsInvalidArgument(err) { + t.Fatalf("bad target err = %v, want InvalidArgument", err) + } +} + +func TestWebSocketEndpointAndRoutes(t *testing.T) { + m := newMock(t) + + api, err := m.CreateAPI(ctx(), &driver.CreateAPIInput{ + Name: "ws", ProtocolType: driver.ProtocolWebSocket, RouteSelectionExpression: "$request.body.action", + }) + if err != nil { + t.Fatalf("CreateAPI ws: %v", err) + } + + if api.APIEndpoint != "wss://"+api.APIID+".execute-api.us-east-1.amazonaws.com" { + t.Fatalf("ws apiEndpoint = %q", api.APIEndpoint) + } + + if _, err := m.CreateRoute(ctx(), api.APIID, &driver.CreateRouteInput{RouteKey: "$connect", AuthorizationType: "JWT"}); !cerrors.IsInvalidArgument(err) { + t.Fatalf("ws JWT route err = %v, want InvalidArgument", err) + } +} + +func TestTagsOnStageAndLimits(t *testing.T) { + m := newMock(t) + api := createHTTPAPI(t, m) + arn := "arn:aws:apigateway:us-east-1::/apis/" + api.APIID + + if _, err := m.CreateStage(ctx(), api.APIID, &driver.CreateStageInput{StageName: "s", Tags: map[string]string{"a": "1"}}); err != nil { + t.Fatalf("CreateStage: %v", err) + } + + if err := m.TagResource(ctx(), arn+"/stages/s", map[string]string{"b": "2"}); err != nil { + t.Fatalf("TagResource stage: %v", err) + } + + tags, err := m.GetTags(ctx(), arn+"/stages/s") + if err != nil || tags["a"] != "1" || tags["b"] != "2" { + t.Fatalf("GetTags stage: %v, %v", err, tags) + } + + if _, err := m.GetTags(ctx(), arn+"/stages/missing"); !cerrors.IsNotFound(err) { + t.Fatalf("missing stage tags err = %v, want NotFound", err) + } + + if _, err := m.GetTags(ctx(), "arn:aws:apigateway:eu-west-1::/apis/"+api.APIID); !cerrors.IsNotFound(err) { + t.Fatalf("other region err = %v, want NotFound", err) + } + + many := map[string]string{} + for i := range 51 { + many[string(rune('a'+i%26))+string(rune('a'+i/26))] = "v" + } + + if err := m.TagResource(ctx(), arn, many); !cerrors.IsInvalidArgument(err) { + t.Fatalf("51 tags err = %v, want InvalidArgument", err) + } + + if err := m.UntagResource(ctx(), arn, nil); !cerrors.IsInvalidArgument(err) { + t.Fatalf("untag without keys err = %v, want InvalidArgument", err) + } +} diff --git a/providers/aws/apigatewayv2/deployments.go b/providers/aws/apigatewayv2/deployments.go new file mode 100644 index 000000000..b180df8d4 --- /dev/null +++ b/providers/aws/apigatewayv2/deployments.go @@ -0,0 +1,245 @@ +package apigatewayv2 + +import ( + "context" + "fmt" + + cerrors "github.com/stackshy/cloudemu/v2/errors" + "github.com/stackshy/cloudemu/v2/services/apigatewayv2/driver" +) + +// autoDeployDescription is the description API Gateway gives the deployments +// it creates for an autoDeploy stage. +const autoDeployDescription = "Automatic deployment triggered by changes to the Api configuration" + +// noRoutesMessage is the error for deploying an API that has no routes. +const noRoutesMessage = "Unable to deploy API because no routes exist in this API" + +// deploymentRecord is a deployment plus the routes and integrations it froze. +// Its fields are exported so a snapshot can serialize it directly. +type deploymentRecord struct { + Deployment driver.Deployment `json:"deployment"` + Routes map[string]*driver.Route `json:"routes,omitempty"` + Integrations map[string]*driver.Integration `json:"integrations,omitempty"` +} + +// newDeployment freezes the API's current routes and integrations into a new +// deployment and stores it. ad must be held for writing. +func (m *Mock) newDeployment(ad *apiData, description string, auto bool) *deploymentRecord { + rec := &deploymentRecord{ + Deployment: driver.Deployment{ + DeploymentID: genID(), Description: description, CreatedDate: m.now(), + DeploymentStatus: driver.DeploymentStatusDeployed, AutoDeployed: auto, + }, + Routes: make(map[string]*driver.Route, len(ad.routes)), + Integrations: make(map[string]*driver.Integration, len(ad.integrations)), + } + + for id, rt := range ad.routes { + cp := copyRoute(rt) + rec.Routes[id] = &cp + } + + for id, ig := range ad.integrations { + cp := copyIntegration(ig) + rec.Integrations[id] = &cp + } + + ad.deployments[rec.Deployment.DeploymentID] = rec + + return rec +} + +// pointStage moves a stage onto a deployment. ad must be held for writing. +func (m *Mock) pointStage(st *driver.Stage, deploymentID string) { + st.DeploymentID = deploymentID + st.LastDeploymentStatusMessage = fmt.Sprintf("Successfully deployed stage with deployment ID '%s'", deploymentID) + st.LastUpdatedDate = m.now() +} + +// autoDeploy gives every autoDeploy stage a fresh deployment after the API's +// routes or integrations change. An API with no routes cannot be deployed, so +// its stages keep what they have. ad must be held for writing. +func (m *Mock) autoDeploy(ad *apiData) { + for _, st := range ad.stages { + if st.AutoDeploy { + m.deployStage(ad, st) + } + } +} + +// deployStage creates an automatic deployment for one stage. ad must be held +// for writing. +func (m *Mock) deployStage(ad *apiData, st *driver.Stage) { + if len(ad.routes) == 0 { + return + } + + rec := m.newDeployment(ad, autoDeployDescription, true) + m.pointStage(st, rec.Deployment.DeploymentID) +} + +// CreateDeployment snapshots the API's routes and integrations. When StageName +// is set the stage moves onto the new deployment. +func (m *Mock) CreateDeployment( + _ context.Context, apiID string, in *driver.CreateDeploymentInput, +) (*driver.Deployment, error) { + ad, err := m.getAPI(apiID) + if err != nil { + return nil, err + } + + if len(in.Description) > maxDescriptionLen { + return nil, badRequest("Description must be at most %d characters", maxDescriptionLen) + } + + ad.mu.Lock() + defer ad.mu.Unlock() + + var st *driver.Stage + + if in.StageName != "" { + var ok bool + if st, ok = ad.stages[in.StageName]; !ok { + return nil, cerrors.New(cerrors.NotFound, "Invalid stage identifier specified") + } + } + + if len(ad.routes) == 0 { + return nil, badRequest(noRoutesMessage) + } + + rec := m.newDeployment(ad, in.Description, false) + + if st != nil { + m.pointStage(st, rec.Deployment.DeploymentID) + } + + out := rec.Deployment + + return &out, nil +} + +// GetDeployment returns a single Deployment. +func (m *Mock) GetDeployment(_ context.Context, apiID, deploymentID string) (*driver.Deployment, error) { + ad, err := m.getAPI(apiID) + if err != nil { + return nil, err + } + + ad.mu.RLock() + defer ad.mu.RUnlock() + + rec, err := findDeployment(ad, deploymentID) + if err != nil { + return nil, err + } + + out := rec.Deployment + + return &out, nil +} + +// GetDeployments lists one page of an API's Deployments, newest first. +func (m *Mock) GetDeployments(_ context.Context, apiID string, page *driver.PageInput) ([]driver.Deployment, string, error) { + return listPage(m, apiID, func(ad *apiData) map[string]*deploymentRecord { return ad.deployments }, + func(rec *deploymentRecord) driver.Deployment { return rec.Deployment }, + func(a, b driver.Deployment) bool { return newestFirst(&a, &b) }, page) +} + +// newestFirst orders deployments by creation time, newest first, then by id. +func newestFirst(a, b *driver.Deployment) bool { + if a.CreatedDate != b.CreatedDate { + return a.CreatedDate > b.CreatedDate + } + + return a.DeploymentID < b.DeploymentID +} + +// UpdateDeployment changes a Deployment's description. The frozen routes and +// integrations never change. +func (m *Mock) UpdateDeployment( + _ context.Context, apiID, deploymentID string, in *driver.UpdateDeploymentInput, +) (*driver.Deployment, error) { + ad, err := m.getAPI(apiID) + if err != nil { + return nil, err + } + + ad.mu.Lock() + defer ad.mu.Unlock() + + rec, err := findDeployment(ad, deploymentID) + if err != nil { + return nil, err + } + + if in.Description != nil { + if len(*in.Description) > maxDescriptionLen { + return nil, badRequest("Description must be at most %d characters", maxDescriptionLen) + } + + rec.Deployment.Description = *in.Description + } + + out := rec.Deployment + + return &out, nil +} + +// DeleteDeployment removes a Deployment no stage points at. +func (m *Mock) DeleteDeployment(_ context.Context, apiID, deploymentID string) error { + ad, err := m.getAPI(apiID) + if err != nil { + return err + } + + ad.mu.Lock() + defer ad.mu.Unlock() + + if _, err := findDeployment(ad, deploymentID); err != nil { + return err + } + + for _, st := range ad.stages { + if st.DeploymentID == deploymentID { + return badRequest("Active stages pointing to this deployment must be moved or deleted") + } + } + + delete(ad.deployments, deploymentID) + + return nil +} + +// findDeployment returns the stored deployment or a NotFound error. ad must +// be held. +func findDeployment(ad *apiData, deploymentID string) (*deploymentRecord, error) { + rec, ok := ad.deployments[deploymentID] + if !ok { + return nil, cerrors.Newf(cerrors.NotFound, "Invalid deployment identifier specified %s", deploymentID) + } + + return rec, nil +} + +// copyDeploymentRecord returns a deep copy of a deployment record. +func copyDeploymentRecord(rec *deploymentRecord) *deploymentRecord { + out := &deploymentRecord{ + Deployment: rec.Deployment, + Routes: make(map[string]*driver.Route, len(rec.Routes)), + Integrations: make(map[string]*driver.Integration, len(rec.Integrations)), + } + + for id, rt := range rec.Routes { + cp := copyRoute(rt) + out.Routes[id] = &cp + } + + for id, ig := range rec.Integrations { + cp := copyIntegration(ig) + out.Integrations[id] = &cp + } + + return out +} diff --git a/providers/aws/apigatewayv2/integrations.go b/providers/aws/apigatewayv2/integrations.go index 674695ad2..ee1f9f96e 100644 --- a/providers/aws/apigatewayv2/integrations.go +++ b/providers/aws/apigatewayv2/integrations.go @@ -16,10 +16,6 @@ func (m *Mock) CreateIntegration( return nil, err } - if in.IntegrationType == "" { - return nil, cerrors.New(cerrors.InvalidArgument, "IntegrationType is required") - } - ad.mu.Lock() defer ad.mu.Unlock() @@ -32,8 +28,15 @@ func (m *Mock) CreateIntegration( TimeoutInMillis: integrationTimeout(ad, in.TimeoutInMillis), Description: in.Description, RequestParameters: copyStrMap(in.RequestParameters), + CredentialsArn: in.CredentialsArn, } + + if err := validateIntegration(ad.api.ProtocolType, ig); err != nil { + return nil, err + } + ad.integrations[ig.IntegrationID] = ig + m.autoDeploy(ad) out := copyIntegration(ig) @@ -74,22 +77,12 @@ func (m *Mock) GetIntegration(_ context.Context, apiID, integrationID string) (* return &out, nil } -// GetIntegrations lists an API's Integrations. -func (m *Mock) GetIntegrations(_ context.Context, apiID string) ([]driver.Integration, error) { - ad, err := m.getAPI(apiID) - if err != nil { - return nil, err - } - - ad.mu.RLock() - defer ad.mu.RUnlock() - - out := make([]driver.Integration, 0, len(ad.integrations)) - for _, ig := range ad.integrations { - out = append(out, copyIntegration(ig)) - } - - return out, nil +// GetIntegrations lists one page of an API's Integrations, ordered by id. +func (m *Mock) GetIntegrations( + _ context.Context, apiID string, page *driver.PageInput, +) ([]driver.Integration, string, error) { + return listPage(m, apiID, func(ad *apiData) map[string]*driver.Integration { return ad.integrations }, copyIntegration, + func(a, b driver.Integration) bool { return a.IntegrationID < b.IntegrationID }, page) } // UpdateIntegration applies the non-nil fields of in to a stored Integration. @@ -109,21 +102,31 @@ func (m *Mock) UpdateIntegration( return nil, cerrors.Newf(cerrors.NotFound, "Invalid integration identifier specified %s", integrationID) } - setString(&ig.IntegrationType, in.IntegrationType) - setString(&ig.IntegrationURI, in.IntegrationURI) - setString(&ig.IntegrationMethod, in.IntegrationMethod) - setString(&ig.ConnectionType, in.ConnectionType) - setString(&ig.PayloadFormatVersion, in.PayloadFormatVersion) - setString(&ig.Description, in.Description) + next := copyIntegration(ig) + setString(&next.IntegrationType, in.IntegrationType) + setString(&next.IntegrationURI, in.IntegrationURI) + setString(&next.IntegrationMethod, in.IntegrationMethod) + setString(&next.ConnectionType, in.ConnectionType) + setString(&next.PayloadFormatVersion, in.PayloadFormatVersion) + setString(&next.Description, in.Description) + setString(&next.CredentialsArn, in.CredentialsArn) if in.TimeoutInMillis != nil { - ig.TimeoutInMillis = *in.TimeoutInMillis + next.TimeoutInMillis = *in.TimeoutInMillis } if in.RequestParameters != nil { - ig.RequestParameters = copyStrMap(in.RequestParameters) + next.RequestParameters = copyStrMap(in.RequestParameters) + } + + if err := validateIntegration(ad.api.ProtocolType, &next); err != nil { + return nil, err } + *ig = next + + m.autoDeploy(ad) + out := copyIntegration(ig) return &out, nil @@ -139,11 +142,17 @@ func (m *Mock) DeleteIntegration(_ context.Context, apiID, integrationID string) ad.mu.Lock() defer ad.mu.Unlock() - if _, ok := ad.integrations[integrationID]; !ok { + ig, ok := ad.integrations[integrationID] + if !ok { return cerrors.Newf(cerrors.NotFound, "Invalid integration identifier specified %s", integrationID) } + if ig.APIGatewayManaged { + return badRequest("Cannot delete an integration managed by API Gateway") + } + delete(ad.integrations, integrationID) + m.autoDeploy(ad) return nil } diff --git a/providers/aws/apigatewayv2/page.go b/providers/aws/apigatewayv2/page.go new file mode 100644 index 000000000..e6d1175c2 --- /dev/null +++ b/providers/aws/apigatewayv2/page.go @@ -0,0 +1,66 @@ +package apigatewayv2 + +import ( + "strconv" + + "github.com/stackshy/cloudemu/v2/internal/pagination" + "github.com/stackshy/cloudemu/v2/services/apigatewayv2/driver" +) + +// maxPageSize is the largest MaxResults a list call accepts. +const maxPageSize = 500 + +// listPage reads one of an API's sub-collections under its read lock and +// returns the page that in selects. +func listPage[V, T any]( + m *Mock, apiID string, pick func(*apiData) map[string]*V, render func(*V) T, + less func(a, b T) bool, in *driver.PageInput, +) (items []T, next string, err error) { + ad, err := m.getAPI(apiID) + if err != nil { + return nil, "", err + } + + ad.mu.RLock() + src := pick(ad) + all := make([]T, 0, len(src)) + + for _, v := range src { + all = append(all, render(v)) + } + ad.mu.RUnlock() + + return pageOf(all, less, in) +} + +// pageOf sorts items by less and returns the page that in selects plus the +// next token. An omitted MaxResults returns every remaining item. +func pageOf[T any](items []T, less func(a, b T) bool, in *driver.PageInput) (page []T, next string, err error) { + var maxResults, token string + if in != nil { + maxResults, token = in.MaxResults, in.NextToken + } + + size := len(items) + + if maxResults != "" { + n, convErr := strconv.Atoi(maxResults) + if convErr != nil || n < 1 || n > maxPageSize { + return nil, "", badRequest("MaxResults must be an integer between 1 and %d", maxPageSize) + } + + size = n + } + + pg, err := pagination.PaginateSorted(items, less, token, max(size, 1)) + if err != nil { + return nil, "", badRequest("Invalid NextToken specified") + } + + page = pg.Items + if page == nil { + page = []T{} + } + + return page, pg.NextPageToken, nil +} diff --git a/providers/aws/apigatewayv2/quickcreate.go b/providers/aws/apigatewayv2/quickcreate.go new file mode 100644 index 000000000..c09b13776 --- /dev/null +++ b/providers/aws/apigatewayv2/quickcreate.go @@ -0,0 +1,186 @@ +package apigatewayv2 + +import ( + "strings" + + "github.com/stackshy/cloudemu/v2/services/apigatewayv2/driver" +) + +// Integration methods quick create gives each target kind. +const ( + quickLambdaMethod = "POST" + quickHTTPMethod = "ANY" +) + +// quickTarget is the integration shape a quick-create Target maps to. +type quickTarget struct { + integrationType string + method string + payloadFormat string +} + +// resolveQuickTarget maps a quick-create Target to its integration: a Lambda +// function ARN becomes AWS_PROXY with payload 2.0, and an HTTP(S) URL becomes +// HTTP_PROXY with payload 1.0. +func resolveQuickTarget(target string) (quickTarget, error) { + switch { + case strings.HasPrefix(target, "arn:") && strings.Contains(target, ":lambda:"): + return quickTarget{driver.IntegrationAWSProxy, quickLambdaMethod, payloadFormatV2}, nil + case strings.HasPrefix(target, "https://"), strings.HasPrefix(target, "http://"): + return quickTarget{driver.IntegrationHTTPProxy, quickHTTPMethod, defaultPayloadFormat}, nil + default: + return quickTarget{}, badRequest("Target must be a Lambda function ARN or an HTTP(S) URL") + } +} + +// checkQuickCreate validates the quick-create fields of a CreateApi request. +func checkQuickCreate(protocol, target, routeKey, credentialsArn string) error { + if target == "" { + if routeKey != "" || credentialsArn != "" { + return badRequest("RouteKey and CredentialsArn can only be specified with Target") + } + + return nil + } + + if protocol != driver.ProtocolHTTP { + return badRequest("Quick create is only supported for HTTP APIs") + } + + if _, err := resolveQuickTarget(target); err != nil { + return err + } + + return validateRouteKey(protocol, orDefault(routeKey, defaultKey)) +} + +// quickCreate builds the managed integration, route and auto-deployed $default +// stage for a validated quick-create request. ad must be held for writing. +func (m *Mock) quickCreate(ad *apiData, target, routeKey, credentialsArn string) { + qt, _ := resolveQuickTarget(target) + + ig := &driver.Integration{ + IntegrationID: genID(), IntegrationType: qt.integrationType, IntegrationURI: target, + IntegrationMethod: qt.method, ConnectionType: defaultConnectionType, + PayloadFormatVersion: qt.payloadFormat, TimeoutInMillis: defaultHTTPTimeoutMillis, + CredentialsArn: credentialsArn, APIGatewayManaged: true, + } + ad.integrations[ig.IntegrationID] = ig + + rt := &driver.Route{ + RouteID: genID(), RouteKey: orDefault(routeKey, defaultKey), + Target: integrationsPrefix + ig.IntegrationID, + AuthorizationType: defaultAuthorizationType, APIGatewayManaged: true, + } + ad.routes[rt.RouteID] = rt + + if _, ok := ad.stages[defaultKey]; !ok { + now := m.now() + ad.stages[defaultKey] = &driver.Stage{ + StageName: defaultKey, AutoDeploy: true, APIGatewayManaged: true, + CreatedDate: now, LastUpdatedDate: now, + } + } + + m.autoDeploy(ad) +} + +// managedResources returns the quick-create integration and route of an API, +// or nils when it was not quick created. ad must be held. +func managedResources(ad *apiData) (*driver.Integration, *driver.Route) { + var ig *driver.Integration + + for _, i := range ad.integrations { + if i.APIGatewayManaged { + ig = i + } + } + + var rt *driver.Route + + for _, r := range ad.routes { + if r.APIGatewayManaged { + rt = r + } + } + + return ig, rt +} + +// updateQuickCreate applies UpdateApi's Target, RouteKey and CredentialsArn to +// the managed integration and route, creating them when the API was not quick +// created yet. ad must be held for writing. +func (m *Mock) updateQuickCreate(ad *apiData, in *driver.UpdateAPIInput) error { + if in.Target == nil && in.RouteKey == nil && in.CredentialsArn == nil { + return nil + } + + ig, rt := managedResources(ad) + if ig == nil || rt == nil { + err := checkQuickCreate(ad.api.ProtocolType, deref(in.Target), deref(in.RouteKey), deref(in.CredentialsArn)) + if err != nil { + return err + } + + if err := checkRouteKeyFree(ad, orDefault(deref(in.RouteKey), defaultKey), ""); err != nil { + return err + } + + if in.Target != nil { + m.quickCreate(ad, *in.Target, deref(in.RouteKey), deref(in.CredentialsArn)) + } + + return nil + } + + if err := applyQuickCreate(ad, ig, rt, in); err != nil { + return err + } + + m.autoDeploy(ad) + + return nil +} + +// applyQuickCreate validates and applies the quick-create updates to the +// managed integration and route. ad must be held for writing. +func applyQuickCreate(ad *apiData, ig *driver.Integration, rt *driver.Route, in *driver.UpdateAPIInput) error { + var qt quickTarget + + if in.Target != nil { + var err error + if qt, err = resolveQuickTarget(*in.Target); err != nil { + return err + } + } + + if in.RouteKey != nil { + if err := validateRouteKey(ad.api.ProtocolType, *in.RouteKey); err != nil { + return err + } + + if err := checkRouteKeyFree(ad, *in.RouteKey, rt.RouteID); err != nil { + return err + } + + rt.RouteKey = *in.RouteKey + } + + if in.Target != nil { + ig.IntegrationType, ig.IntegrationURI = qt.integrationType, *in.Target + ig.IntegrationMethod, ig.PayloadFormatVersion = qt.method, qt.payloadFormat + } + + setString(&ig.CredentialsArn, in.CredentialsArn) + + return nil +} + +// deref returns *s, or "" for a nil pointer. +func deref(s *string) string { + if s == nil { + return "" + } + + return *s +} diff --git a/providers/aws/apigatewayv2/routes.go b/providers/aws/apigatewayv2/routes.go index 63fbe8b47..41cc58f24 100644 --- a/providers/aws/apigatewayv2/routes.go +++ b/providers/aws/apigatewayv2/routes.go @@ -14,22 +14,38 @@ func (m *Mock) CreateRoute(_ context.Context, apiID string, in *driver.CreateRou return nil, err } - if in.RouteKey == "" { - return nil, cerrors.New(cerrors.InvalidArgument, "RouteKey is required") - } - ad.mu.Lock() defer ad.mu.Unlock() + protocol := ad.api.ProtocolType + authType := orDefault(in.AuthorizationType, defaultAuthorizationType) + + if err := validateRouteKey(protocol, in.RouteKey); err != nil { + return nil, err + } + + if err := validateAuthorizationType(protocol, authType); err != nil { + return nil, err + } + + if err := validateRouteTarget(ad, in.Target); err != nil { + return nil, err + } + + if err := checkRouteKeyFree(ad, in.RouteKey, ""); err != nil { + return nil, err + } + rt := &driver.Route{ RouteID: genID(), RouteKey: in.RouteKey, Target: in.Target, - AuthorizationType: orDefault(in.AuthorizationType, defaultAuthorizationType), + AuthorizationType: authType, APIKeyRequired: in.APIKeyRequired, AuthorizerID: in.AuthorizerID, AuthorizationScopes: append([]string(nil), in.AuthorizationScopes...), OperationName: in.OperationName, } ad.routes[rt.RouteID] = rt + m.autoDeploy(ad) out := copyRoute(rt) @@ -56,22 +72,10 @@ func (m *Mock) GetRoute(_ context.Context, apiID, routeID string) (*driver.Route return &out, nil } -// GetRoutes lists an API's Routes. -func (m *Mock) GetRoutes(_ context.Context, apiID string) ([]driver.Route, error) { - ad, err := m.getAPI(apiID) - if err != nil { - return nil, err - } - - ad.mu.RLock() - defer ad.mu.RUnlock() - - out := make([]driver.Route, 0, len(ad.routes)) - for _, rt := range ad.routes { - out = append(out, copyRoute(rt)) - } - - return out, nil +// GetRoutes lists one page of an API's Routes, ordered by id. +func (m *Mock) GetRoutes(_ context.Context, apiID string, page *driver.PageInput) ([]driver.Route, string, error) { + return listPage(m, apiID, func(ad *apiData) map[string]*driver.Route { return ad.routes }, copyRoute, + func(a, b driver.Route) bool { return a.RouteID < b.RouteID }, page) } // UpdateRoute applies the non-nil fields of in to a stored Route (PATCH). @@ -89,12 +93,25 @@ func (m *Mock) UpdateRoute(_ context.Context, apiID, routeID string, in *driver. return nil, cerrors.Newf(cerrors.NotFound, "Invalid route identifier specified %s", routeID) } - setString(&rt.RouteKey, in.RouteKey) - setString(&rt.Target, in.Target) - setString(&rt.AuthorizationType, in.AuthorizationType) - setString(&rt.AuthorizerID, in.AuthorizerID) - setString(&rt.OperationName, in.OperationName) - setBool(&rt.APIKeyRequired, in.APIKeyRequired) + next := copyRoute(rt) + setString(&next.RouteKey, in.RouteKey) + setString(&next.Target, in.Target) + setString(&next.AuthorizationType, in.AuthorizationType) + setString(&next.AuthorizerID, in.AuthorizerID) + setString(&next.OperationName, in.OperationName) + setBool(&next.APIKeyRequired, in.APIKeyRequired) + + if in.AuthorizationScopes != nil { + next.AuthorizationScopes = append([]string(nil), in.AuthorizationScopes...) + } + + if err := validateRouteUpdate(ad, rt, &next); err != nil { + return nil, err + } + + *rt = next + + m.autoDeploy(ad) out := copyRoute(rt) @@ -116,10 +133,41 @@ func (m *Mock) DeleteRoute(_ context.Context, apiID, routeID string) error { } delete(ad.routes, routeID) + m.autoDeploy(ad) return nil } +// validateRouteUpdate checks the route an UpdateRoute would produce. A managed +// (quick-create) route keeps its route key. ad must be held. +func validateRouteUpdate(ad *apiData, cur, next *driver.Route) error { + protocol := ad.api.ProtocolType + + if next.RouteKey != cur.RouteKey { + if cur.APIGatewayManaged { + return badRequest("Cannot modify the route key of a route managed by API Gateway") + } + + if err := validateRouteKey(protocol, next.RouteKey); err != nil { + return err + } + + if err := checkRouteKeyFree(ad, next.RouteKey, cur.RouteID); err != nil { + return err + } + } + + if err := validateAuthorizationType(protocol, next.AuthorizationType); err != nil { + return err + } + + if next.Target == cur.Target { + return nil + } + + return validateRouteTarget(ad, next.Target) +} + // copyRoute returns a deep copy of a Route. func copyRoute(r *driver.Route) driver.Route { out := *r diff --git a/providers/aws/apigatewayv2/snapshot.go b/providers/aws/apigatewayv2/snapshot.go index 0363ec14b..971b82b3f 100644 --- a/providers/aws/apigatewayv2/snapshot.go +++ b/providers/aws/apigatewayv2/snapshot.go @@ -20,12 +20,14 @@ type apigatewayV2Snapshot struct { } // apiSnapshot is the exported form of apiData: the API plus its routes, -// integrations and stages, all under their original identities. +// integrations, stages and deployments (with their frozen route and +// integration snapshots), all under their original identities. type apiSnapshot struct { API driver.API `json:"api"` Routes map[string]*driver.Route `json:"routes,omitempty"` Integrations map[string]*driver.Integration `json:"integrations,omitempty"` Stages map[string]*driver.Stage `json:"stages,omitempty"` + Deployments map[string]*deploymentRecord `json:"deployments,omitempty"` } // Snapshot captures the mock's entire state as JSON. includeAssets is unused. API Gateway v2 holds @@ -55,6 +57,11 @@ func snapshotAPI(ad *apiData) *apiSnapshot { Routes: make(map[string]*driver.Route, len(ad.routes)), Integrations: make(map[string]*driver.Integration, len(ad.integrations)), Stages: make(map[string]*driver.Stage, len(ad.stages)), + Deployments: make(map[string]*deploymentRecord, len(ad.deployments)), + } + + for id, rec := range ad.deployments { + as.Deployments[id] = copyDeploymentRecord(rec) } for id, r := range ad.routes { @@ -97,6 +104,11 @@ func restoreAPI(as *apiSnapshot) *apiData { routes: make(map[string]*driver.Route, len(as.Routes)), integrations: make(map[string]*driver.Integration, len(as.Integrations)), stages: make(map[string]*driver.Stage, len(as.Stages)), + deployments: make(map[string]*deploymentRecord, len(as.Deployments)), + } + + for id, rec := range as.Deployments { + ad.deployments[id] = rec } for id, r := range as.Routes { diff --git a/providers/aws/apigatewayv2/snapshot_test.go b/providers/aws/apigatewayv2/snapshot_test.go index f4044fa05..124df0d85 100644 --- a/providers/aws/apigatewayv2/snapshot_test.go +++ b/providers/aws/apigatewayv2/snapshot_test.go @@ -58,3 +58,50 @@ func TestSnapshotRestoreRoundTrip(t *testing.T) { t.Fatalf("restored GetStage: %v, %+v", err, gotSt) } } + +func TestSnapshotRestoreKeepsDeploymentsAndTags(t *testing.T) { + src := newMock(t) + + api, err := src.CreateAPI(ctx(), &driver.CreateAPIInput{ + Name: "q", ProtocolType: driver.ProtocolHTTP, Target: "https://example.com", + }) + if err != nil { + t.Fatalf("CreateAPI: %v", err) + } + + stageARN := "arn:aws:apigateway:us-east-1::/apis/" + api.APIID + "/stages/$default" + if err := src.TagResource(ctx(), stageARN, map[string]string{"k": "v"}); err != nil { + t.Fatalf("TagResource: %v", err) + } + + st, _ := src.GetStage(ctx(), api.APIID, "$default") + + data, err := src.Snapshot(ctx(), false) + if err != nil { + t.Fatalf("Snapshot: %v", err) + } + + dst := newMock(t) + if err := dst.Restore(ctx(), data); err != nil { + t.Fatalf("Restore: %v", err) + } + + d, err := dst.GetDeployment(ctx(), api.APIID, st.DeploymentID) + if err != nil || !d.AutoDeployed { + t.Fatalf("restored deployment: %v, %+v", err, d) + } + + if err := dst.DeleteDeployment(ctx(), api.APIID, st.DeploymentID); err == nil { + t.Fatalf("restored active deployment was deletable") + } + + tags, err := dst.GetTags(ctx(), stageARN) + if err != nil || tags["k"] != "v" { + t.Fatalf("restored stage tags: %v, %v", err, tags) + } + + gotSt, _ := dst.GetStage(ctx(), api.APIID, "$default") + if !gotSt.APIGatewayManaged { + t.Fatalf("restored stage lost apiGatewayManaged: %+v", gotSt) + } +} diff --git a/providers/aws/apigatewayv2/stages.go b/providers/aws/apigatewayv2/stages.go index 4fe88cb94..126a99bad 100644 --- a/providers/aws/apigatewayv2/stages.go +++ b/providers/aws/apigatewayv2/stages.go @@ -15,8 +15,16 @@ func (m *Mock) CreateStage(_ context.Context, apiID string, in *driver.CreateSta return nil, err } - if in.StageName == "" { - return nil, cerrors.New(cerrors.InvalidArgument, "StageName is required") + if err := validateStageName(in.StageName); err != nil { + return nil, err + } + + if len(in.Description) > maxDescriptionLen { + return nil, badRequest("Description must be at most %d characters", maxDescriptionLen) + } + + if err := validateTags(in.Tags); err != nil { + return nil, err } ad.mu.Lock() @@ -26,6 +34,10 @@ func (m *Mock) CreateStage(_ context.Context, apiID string, in *driver.CreateSta return nil, cerrors.Newf(cerrors.AlreadyExists, "Stage already exists: %s", in.StageName) } + if err := checkDeploymentID(ad, in.DeploymentID); err != nil { + return nil, err + } + now := m.now() st := &driver.Stage{ StageName: in.StageName, Description: in.Description, AutoDeploy: in.AutoDeploy, @@ -33,9 +45,18 @@ func (m *Mock) CreateStage(_ context.Context, apiID string, in *driver.CreateSta StageVariables: copyStrMap(in.StageVariables), DefaultRouteSettings: copyRouteSettings(in.DefaultRouteSettings), CreatedDate: now, LastUpdatedDate: now, + Tags: copyStrMap(in.Tags), } ad.stages[in.StageName] = st + if in.DeploymentID != "" { + m.pointStage(st, in.DeploymentID) + } + + if st.AutoDeploy { + m.deployStage(ad, st) + } + out := copyStage(st) return &out, nil @@ -61,22 +82,10 @@ func (m *Mock) GetStage(_ context.Context, apiID, stageName string) (*driver.Sta return &out, nil } -// GetStages lists an API's Stages. -func (m *Mock) GetStages(_ context.Context, apiID string) ([]driver.Stage, error) { - ad, err := m.getAPI(apiID) - if err != nil { - return nil, err - } - - ad.mu.RLock() - defer ad.mu.RUnlock() - - out := make([]driver.Stage, 0, len(ad.stages)) - for _, st := range ad.stages { - out = append(out, copyStage(st)) - } - - return out, nil +// GetStages lists one page of an API's Stages, ordered by name. +func (m *Mock) GetStages(_ context.Context, apiID string, page *driver.PageInput) ([]driver.Stage, string, error) { + return listPage(m, apiID, func(ad *apiData) map[string]*driver.Stage { return ad.stages }, copyStage, + func(a, b driver.Stage) bool { return a.StageName < b.StageName }, page) } // UpdateStage applies the non-nil fields of in to a stored Stage (PATCH). @@ -89,15 +98,24 @@ func (m *Mock) UpdateStage(_ context.Context, apiID, stageName string, in *drive ad.mu.Lock() defer ad.mu.Unlock() - st, ok := ad.stages[stageName] - if !ok { - return nil, cerrors.Newf(cerrors.NotFound, "Invalid stage name specified %s", stageName) + st, err := findStage(ad, stageName) + if err != nil { + return nil, err + } + + if err := checkStageUpdate(ad, st, in); err != nil { + return nil, err } + wasAuto := st.AutoDeploy + setString(&st.Description, in.Description) - setString(&st.DeploymentID, in.DeploymentID) setBool(&st.AutoDeploy, in.AutoDeploy) + if in.DeploymentID != nil && *in.DeploymentID != "" { + m.pointStage(st, *in.DeploymentID) + } + if in.StageVariables != nil { st.StageVariables = copyStrMap(in.StageVariables) } @@ -108,6 +126,10 @@ func (m *Mock) UpdateStage(_ context.Context, apiID, stageName string, in *drive st.LastUpdatedDate = m.now() + if st.AutoDeploy && !wasAuto { + m.deployStage(ad, st) + } + out := copyStage(st) return &out, nil @@ -123,8 +145,13 @@ func (m *Mock) DeleteStage(_ context.Context, apiID, stageName string) error { ad.mu.Lock() defer ad.mu.Unlock() - if _, ok := ad.stages[stageName]; !ok { - return cerrors.Newf(cerrors.NotFound, "Invalid stage name specified %s", stageName) + st, err := findStage(ad, stageName) + if err != nil { + return err + } + + if st.APIGatewayManaged { + return badRequest(managedStageMessage) } delete(ad.stages, stageName) @@ -132,6 +159,51 @@ func (m *Mock) DeleteStage(_ context.Context, apiID, stageName string) error { return nil } +// managedStageMessage is the error for changing a quick-create $default stage. +const managedStageMessage = "Cannot modify or delete a stage managed by API Gateway" + +// findStage returns the stored stage or a NotFound error. ad must be held. +func findStage(ad *apiData, stageName string) (*driver.Stage, error) { + st, ok := ad.stages[stageName] + if !ok { + return nil, cerrors.Newf(cerrors.NotFound, "Invalid stage name specified %s", stageName) + } + + return st, nil +} + +// checkStageUpdate validates an UpdateStage request against the stored stage. +// ad must be held. +func checkStageUpdate(ad *apiData, st *driver.Stage, in *driver.UpdateStageInput) error { + if st.APIGatewayManaged { + return badRequest(managedStageMessage) + } + + if in.Description != nil && len(*in.Description) > maxDescriptionLen { + return badRequest("Description must be at most %d characters", maxDescriptionLen) + } + + if in.DeploymentID != nil { + return checkDeploymentID(ad, *in.DeploymentID) + } + + return nil +} + +// checkDeploymentID rejects a stage deploymentId that names no deployment of +// the API. An empty id is allowed. ad must be held. +func checkDeploymentID(ad *apiData, deploymentID string) error { + if deploymentID == "" { + return nil + } + + if _, ok := ad.deployments[deploymentID]; !ok { + return badRequest("Invalid deployment identifier specified %s", deploymentID) + } + + return nil +} + // copyRouteSettings returns a deep copy of a RouteSettings, or nil. func copyRouteSettings(rs *driver.RouteSettings) *driver.RouteSettings { if rs == nil { @@ -146,6 +218,7 @@ func copyRouteSettings(rs *driver.RouteSettings) *driver.RouteSettings { // copyStage returns a deep copy of a Stage. func copyStage(s *driver.Stage) driver.Stage { out := *s + out.Tags = copyStrMap(s.Tags) out.StageVariables = copyStrMap(s.StageVariables) out.DefaultRouteSettings = copyRouteSettings(s.DefaultRouteSettings) diff --git a/providers/aws/apigatewayv2/tags.go b/providers/aws/apigatewayv2/tags.go new file mode 100644 index 000000000..aaa290fe5 --- /dev/null +++ b/providers/aws/apigatewayv2/tags.go @@ -0,0 +1,178 @@ +package apigatewayv2 + +import ( + "context" + "strings" + + cerrors "github.com/stackshy/cloudemu/v2/errors" + "github.com/stackshy/cloudemu/v2/services/apigatewayv2/driver" +) + +// Resource path shapes inside an apigatewayv2 ARN +// (arn:aws:apigateway:{region}::{path}). +const ( + arnPathAPIs = "apis" + arnPathStages = "stages" + arnPathDomainNames = "domainnames" + arnPathVpcLinks = "vpclinks" +) + +// Segment counts of the taggable resource paths. +const ( + segsTagAPI = 2 // apis/{apiId} + segsTagStage = 4 // apis/{apiId}/stages/{stageName} + segsTagOther = 2 // domainnames/{name} or vpclinks/{id} +) + +// arnParts is the number of colon-separated fields in an ARN. +const arnParts = 6 + +// tagHolder is a locked API or stage whose tags a tagging call reads or +// replaces. The caller must call unlock. +type tagHolder struct { + api *driver.API + stage *driver.Stage + unlock func() +} + +func (h tagHolder) get() map[string]string { + if h.stage != nil { + return h.stage.Tags + } + + return h.api.Tags +} + +func (h tagHolder) set(tags map[string]string) { + if h.stage != nil { + h.stage.Tags = tags + return + } + + h.api.Tags = tags +} + +// tagTarget locks the resource an ARN names and returns it. +func (m *Mock) tagTarget(resourceARN string) (tagHolder, error) { + path, err := m.parseARN(resourceARN) + if err != nil { + return tagHolder{}, err + } + + segs := strings.Split(strings.Trim(path, "/"), "/") + + switch { + case segs[0] == arnPathAPIs && (len(segs) == segsTagAPI || (len(segs) == segsTagStage && segs[2] == arnPathStages)): + return m.apiOrStageTags(segs) + case segs[0] == arnPathDomainNames && len(segs) == segsTagOther: + return tagHolder{}, cerrors.Newf(cerrors.NotFound, "Invalid domain name identifier specified %s", segs[1]) + case segs[0] == arnPathVpcLinks && len(segs) == segsTagOther: + return tagHolder{}, cerrors.Newf(cerrors.NotFound, "Invalid VpcLink identifier specified %s", segs[1]) + default: + return tagHolder{}, badRequest("Invalid resource ARN specified %s", resourceARN) + } +} + +// apiOrStageTags locks the API named by segs and returns the API or stage. +func (m *Mock) apiOrStageTags(segs []string) (tagHolder, error) { + ad, err := m.getAPI(segs[1]) + if err != nil { + return tagHolder{}, err + } + + ad.mu.Lock() + + if len(segs) == segsTagAPI { + return tagHolder{api: &ad.api, unlock: ad.mu.Unlock}, nil + } + + st, ok := ad.stages[segs[3]] + if !ok { + ad.mu.Unlock() + + return tagHolder{}, cerrors.Newf(cerrors.NotFound, "Invalid stage name specified %s", segs[3]) + } + + return tagHolder{stage: st, unlock: ad.mu.Unlock}, nil +} + +// parseARN checks an apigateway ARN in this region and returns its resource +// path. +func (m *Mock) parseARN(resourceARN string) (string, error) { + parts := strings.SplitN(resourceARN, ":", arnParts) + if len(parts) != arnParts || parts[0] != "arn" || parts[2] != "apigateway" || !strings.HasPrefix(parts[5], "/") { + return "", badRequest("Invalid resource ARN specified %s", resourceARN) + } + + if parts[3] != m.region { + return "", cerrors.Newf(cerrors.NotFound, "Invalid resource ARN specified %s", resourceARN) + } + + return parts[5], nil +} + +// TagResource adds or overwrites tags on an API, stage, domain name or VPC link. +func (m *Mock) TagResource(_ context.Context, resourceARN string, tags map[string]string) error { + if err := validateTags(tags); err != nil { + return err + } + + h, err := m.tagTarget(resourceARN) + if err != nil { + return err + } + defer h.unlock() + + merged := copyStrMap(h.get()) + if merged == nil { + merged = make(map[string]string, len(tags)) + } + + for k, v := range tags { + merged[k] = v + } + + if err := validateTags(merged); err != nil { + return err + } + + h.set(merged) + + return nil +} + +// UntagResource removes tag keys from a resource. Unknown keys are ignored. +func (m *Mock) UntagResource(_ context.Context, resourceARN string, tagKeys []string) error { + if len(tagKeys) == 0 { + return badRequest("TagKeys is required") + } + + h, err := m.tagTarget(resourceARN) + if err != nil { + return err + } + defer h.unlock() + + tags := h.get() + for _, k := range tagKeys { + delete(tags, k) + } + + return nil +} + +// GetTags returns a resource's tags (an empty map when it has none). +func (m *Mock) GetTags(_ context.Context, resourceARN string) (map[string]string, error) { + h, err := m.tagTarget(resourceARN) + if err != nil { + return nil, err + } + defer h.unlock() + + out := copyStrMap(h.get()) + if out == nil { + out = map[string]string{} + } + + return out, nil +} diff --git a/providers/aws/apigatewayv2/validation.go b/providers/aws/apigatewayv2/validation.go new file mode 100644 index 000000000..1ca104d3b --- /dev/null +++ b/providers/aws/apigatewayv2/validation.go @@ -0,0 +1,253 @@ +package apigatewayv2 + +import ( + "regexp" + "strings" + + cerrors "github.com/stackshy/cloudemu/v2/errors" + "github.com/stackshy/cloudemu/v2/services/apigatewayv2/driver" +) + +// Well-known keys and expressions. +const ( + defaultKey = "$default" + integrationsPrefix = "integrations/" + usageIdentifierKeyExp = "$context.authorizer.usageIdentifierKey" + payloadFormatV2 = "2.0" +) + +// Length limits from the apigatewayv2 model. +const ( + maxNameLen = 128 + maxDescriptionLen = 1024 + minTimeoutMillis = 50 +) + +// Tag limits. +const ( + maxTags = 50 + maxTagKeyLen = 128 + maxTagValueLen = 1600 + reservedTagPfx = "aws:" +) + +// stageNameRe is the character set a stage name may use; "$default" is the +// only other accepted value. +var stageNameRe = regexp.MustCompile(`^[a-zA-Z0-9._-]+$`) + +// httpRouteMethods are the methods an HTTP API route key may start with. +// +//nolint:gochecknoglobals // read-only lookup table +var httpRouteMethods = map[string]bool{ + "ANY": true, "DELETE": true, "GET": true, "HEAD": true, + "OPTIONS": true, "PATCH": true, "POST": true, "PUT": true, +} + +// integrationTypes is every integration type the service knows. +// +//nolint:gochecknoglobals // read-only lookup table +var integrationTypes = map[string]bool{ + driver.IntegrationAWS: true, driver.IntegrationAWSProxy: true, + driver.IntegrationHTTP: true, driver.IntegrationHTTPProxy: true, driver.IntegrationMock: true, +} + +// authorizationTypes lists the authorization types each protocol accepts. +// +//nolint:gochecknoglobals // read-only lookup table +var authorizationTypes = map[string]map[string]bool{ + driver.ProtocolHTTP: {"NONE": true, "AWS_IAM": true, "CUSTOM": true, "JWT": true}, + driver.ProtocolWebSocket: {"NONE": true, "AWS_IAM": true, "CUSTOM": true}, +} + +func badRequest(format string, args ...any) error { + return cerrors.Newf(cerrors.InvalidArgument, format, args...) +} + +// validateAPIFields checks the name, description and selection expressions an +// API ends up with after a create or update. +func validateAPIFields(a *driver.API) error { + if a.Name == "" || len(a.Name) > maxNameLen { + return badRequest("Name must be between 1 and %d characters", maxNameLen) + } + + if len(a.Description) > maxDescriptionLen { + return badRequest("Description must be at most %d characters", maxDescriptionLen) + } + + if a.APIKeySelectionExpression != defaultAPIKeySelectionExpr && a.APIKeySelectionExpression != usageIdentifierKeyExp { + return badRequest("Invalid API key selection expression specified: %s", a.APIKeySelectionExpression) + } + + return validateRouteSelection(a.ProtocolType, a.RouteSelectionExpression) +} + +// validateRouteSelection enforces the route selection expression rules: HTTP +// APIs accept only the fixed method-and-path form, and WebSocket APIs need a +// $request expression. +func validateRouteSelection(protocol, expr string) error { + if protocol == driver.ProtocolHTTP { + if expr != defaultRouteSelectionExpr { + return badRequest("Only %s is supported for HTTP APIs", defaultRouteSelectionExpr) + } + + return nil + } + + if !strings.HasPrefix(expr, "$request.") && !strings.HasPrefix(expr, "${request.") { + return badRequest("Invalid route selection expression specified: %s", expr) + } + + return nil +} + +// validateRouteKey checks a route key against its API's protocol. HTTP APIs +// take "$default" or "METHOD /path"; WebSocket APIs take any non-empty key. +func validateRouteKey(protocol, key string) error { + if key == "" { + return badRequest("RouteKey is required") + } + + if protocol != driver.ProtocolHTTP || key == defaultKey { + return nil + } + + method, path, ok := strings.Cut(key, " ") + if !ok || !httpRouteMethods[method] || !strings.HasPrefix(path, "/") || strings.ContainsAny(path, " \t") { + return badRequest(`The provided route key is not formatted properly for HTTP protocol. ` + + `Format should be " /" or "$default"`) + } + + return nil +} + +// validateRouteTarget checks that a route target names an integration on the +// API. ad must be held. +func validateRouteTarget(ad *apiData, target string) error { + if target == "" { + return nil + } + + id, ok := strings.CutPrefix(target, integrationsPrefix) + if !ok { + return badRequest("Invalid Integration identifier specified") + } + + if _, ok := ad.integrations[id]; !ok { + return badRequest("Invalid Integration identifier specified") + } + + return nil +} + +// validateAuthorizationType checks a route's authorization type against its +// API's protocol. +func validateAuthorizationType(protocol, authType string) error { + if !authorizationTypes[protocol][authType] { + return badRequest("Invalid authorization type specified: %s", authType) + } + + return nil +} + +// checkRouteKeyFree returns a Conflict error when another route on the API +// already uses key. ad must be held. +func checkRouteKeyFree(ad *apiData, key, selfID string) error { + for id, rt := range ad.routes { + if id != selfID && rt.RouteKey == key { + return cerrors.Newf(cerrors.AlreadyExists, "Route with key %s already exists for this API", key) + } + } + + return nil +} + +// validateIntegration checks an integration's type, payload format version +// and timeout against its API's protocol. +func validateIntegration(protocol string, ig *driver.Integration) error { + if ig.IntegrationType == "" { + return badRequest("IntegrationType is required") + } + + if !integrationTypes[ig.IntegrationType] { + return badRequest("Invalid integration type specified: %s", ig.IntegrationType) + } + + if len(ig.Description) > maxDescriptionLen { + return badRequest("Description must be at most %d characters", maxDescriptionLen) + } + + if err := validatePayloadFormat(protocol, ig); err != nil { + return err + } + + maxTimeout := defaultHTTPTimeoutMillis + if protocol == driver.ProtocolWebSocket { + maxTimeout = defaultWebSocketTimeoutMillis + } + + if ig.TimeoutInMillis < minTimeoutMillis || ig.TimeoutInMillis > maxTimeout { + return badRequest("TimeoutInMillis must be between %d and %d", minTimeoutMillis, maxTimeout) + } + + return nil +} + +// validatePayloadFormat applies the per-protocol payload format rules: HTTP +// APIs support only the two proxy types and 2.0 only on AWS_PROXY; WebSocket +// APIs support only 1.0. +func validatePayloadFormat(protocol string, ig *driver.Integration) error { + v := ig.PayloadFormatVersion + if v != defaultPayloadFormat && v != payloadFormatV2 { + return badRequest("Invalid payload format version specified: %s", v) + } + + if protocol == driver.ProtocolWebSocket { + if v != defaultPayloadFormat { + return badRequest("Payload format version %s is not supported for WEBSOCKET APIs", v) + } + + return nil + } + + if ig.IntegrationType != driver.IntegrationAWSProxy && ig.IntegrationType != driver.IntegrationHTTPProxy { + return badRequest("Integration type %s is not supported for HTTP APIs", ig.IntegrationType) + } + + if v == payloadFormatV2 && ig.IntegrationType != driver.IntegrationAWSProxy { + return badRequest("Payload format version 2.0 is only supported for AWS_PROXY integrations") + } + + return nil +} + +// validateStageName accepts "$default" or a name drawn from a-zA-Z0-9._-. +func validateStageName(name string) error { + if name == "" { + return badRequest("StageName is required") + } + + if name != defaultKey && (len(name) > maxNameLen || !stageNameRe.MatchString(name)) { + return badRequest("Stage name only allows a-zA-Z0-9._- or $default") + } + + return nil +} + +// validateTags applies the tag count, length and reserved-prefix rules. +func validateTags(tags map[string]string) error { + if len(tags) > maxTags { + return badRequest("A resource can have at most %d tags", maxTags) + } + + for k, v := range tags { + if strings.HasPrefix(strings.ToLower(k), reservedTagPfx) { + return badRequest("Tag keys cannot start with the reserved prefix %s", reservedTagPfx) + } + + if k == "" || len(k) > maxTagKeyLen || len(v) > maxTagValueLen { + return badRequest("Tag keys must be 1 to %d characters and values at most %d", maxTagKeyLen, maxTagValueLen) + } + } + + return nil +} diff --git a/server/aws/apigatewayv2/deployment_e2e_test.go b/server/aws/apigatewayv2/deployment_e2e_test.go new file mode 100644 index 000000000..1ca8c141c --- /dev/null +++ b/server/aws/apigatewayv2/deployment_e2e_test.go @@ -0,0 +1,115 @@ +package apigatewayv2_test + +import ( + "net/http" + "testing" +) + +// TestE2E_ManualDeploymentLifecycle covers CreateDeployment, GetDeployment(s), +// UpdateDeployment and DeleteDeployment, and the stage that points at it. +func TestE2E_ManualDeploymentLifecycle(t *testing.T) { + ts := newE2E(t) + apiID := newHTTPAPI(t, ts.URL) + apiBase := ts.URL + "/v2/apis/" + apiID + + wantErr(t, http.MethodPost, apiBase+"/deployments", `{}`, + http.StatusBadRequest, "BadRequestException", "Unable to deploy API because no routes exist in this API") + + igID := newLambdaIntegration(t, apiBase) + mustDo(t, http.MethodPost, apiBase+"/routes", `{"routeKey":"GET /items","target":"integrations/`+igID+`"}`, http.StatusCreated) + mustDo(t, http.MethodPost, apiBase+"/stages", `{"stageName":"prod"}`, http.StatusCreated) + + wantErr(t, http.MethodPost, apiBase+"/deployments", `{"stageName":"nope"}`, + http.StatusNotFound, "NotFoundException", "Invalid stage identifier specified") + + dep := mustDo(t, http.MethodPost, apiBase+"/deployments", `{"stageName":"prod","description":"v1"}`, http.StatusCreated) + + depID, _ := dep["deploymentId"].(string) + if depID == "" || dep["deploymentStatus"] != "DEPLOYED" || dep["autoDeployed"] != false || dep["description"] != "v1" { + t.Fatalf("CreateDeployment body = %v", dep) + } + + stage := mustDo(t, http.MethodGet, apiBase+"/stages/prod", "", http.StatusOK) + if stage["deploymentId"] != depID { + t.Fatalf("stage deploymentId = %v, want %s", stage["deploymentId"], depID) + } + + if stage["lastDeploymentStatusMessage"] != "Successfully deployed stage with deployment ID '"+depID+"'" { + t.Fatalf("lastDeploymentStatusMessage = %v", stage["lastDeploymentStatusMessage"]) + } + + got := mustDo(t, http.MethodGet, apiBase+"/deployments/"+depID, "", http.StatusOK) + if got["deploymentId"] != depID || got["createdDate"] == nil { + t.Fatalf("GetDeployment = %v", got) + } + + upd := mustDo(t, http.MethodPatch, apiBase+"/deployments/"+depID, `{"description":"v1b"}`, http.StatusOK) + if upd["description"] != "v1b" { + t.Fatalf("UpdateDeployment = %v", upd) + } + + if n := len(items(mustDo(t, http.MethodGet, apiBase+"/deployments", "", http.StatusOK))); n != 1 { + t.Fatalf("GetDeployments items = %d, want 1", n) + } + + wantErr(t, http.MethodDelete, apiBase+"/deployments/"+depID, "", + http.StatusBadRequest, "BadRequestException", "Active stages pointing to this deployment must be moved or deleted") + + mustDo(t, http.MethodDelete, apiBase+"/stages/prod", "", http.StatusNoContent) + mustDo(t, http.MethodDelete, apiBase+"/deployments/"+depID, "", http.StatusNoContent) + + wantErr(t, http.MethodGet, apiBase+"/deployments/"+depID, "", + http.StatusNotFound, "NotFoundException", "Invalid deployment identifier specified "+depID) +} + +// TestE2E_StageRejectsUnknownDeploymentID proves a stage cannot point at a +// deployment that does not exist, on create or update. +func TestE2E_StageRejectsUnknownDeploymentID(t *testing.T) { + ts := newE2E(t) + apiID := newHTTPAPI(t, ts.URL) + apiBase := ts.URL + "/v2/apis/" + apiID + + wantErr(t, http.MethodPost, apiBase+"/stages", `{"stageName":"prod","deploymentId":"bogus12345"}`, + http.StatusBadRequest, "BadRequestException", "Invalid deployment identifier specified bogus12345") + + mustDo(t, http.MethodPost, apiBase+"/stages", `{"stageName":"prod"}`, http.StatusCreated) + + wantErr(t, http.MethodPatch, apiBase+"/stages/prod", `{"deploymentId":"bogus12345"}`, + http.StatusBadRequest, "BadRequestException", "Invalid deployment identifier specified bogus12345") +} + +// TestE2E_AutoDeployStageRedeploysOnChange proves an autoDeploy stage gets a +// new automatic deployment whenever a route or integration changes. +func TestE2E_AutoDeployStageRedeploysOnChange(t *testing.T) { + ts := newE2E(t) + apiID := newHTTPAPI(t, ts.URL) + apiBase := ts.URL + "/v2/apis/" + apiID + + stage := mustDo(t, http.MethodPost, apiBase+"/stages", `{"stageName":"$default","autoDeploy":true}`, http.StatusCreated) + if _, ok := stage["deploymentId"]; ok { + t.Fatalf("autoDeploy stage on an API with no routes has a deployment: %v", stage) + } + + igID := newLambdaIntegration(t, apiBase) + mustDo(t, http.MethodPost, apiBase+"/routes", `{"routeKey":"GET /a","target":"integrations/`+igID+`"}`, http.StatusCreated) + + stage = mustDo(t, http.MethodGet, apiBase+"/stages/$default", "", http.StatusOK) + + first, _ := stage["deploymentId"].(string) + if first == "" { + t.Fatalf("autoDeploy stage not deployed after CreateRoute: %v", stage) + } + + dep := mustDo(t, http.MethodGet, apiBase+"/deployments/"+first, "", http.StatusOK) + if dep["autoDeployed"] != true || dep["deploymentStatus"] != "DEPLOYED" || + dep["description"] != "Automatic deployment triggered by changes to the Api configuration" { + t.Fatalf("auto deployment = %v", dep) + } + + mustDo(t, http.MethodPost, apiBase+"/routes", `{"routeKey":"GET /b","target":"integrations/`+igID+`"}`, http.StatusCreated) + + stage = mustDo(t, http.MethodGet, apiBase+"/stages/$default", "", http.StatusOK) + if second, _ := stage["deploymentId"].(string); second == "" || second == first { + t.Fatalf("autoDeploy stage not redeployed: first=%s now=%v", first, stage["deploymentId"]) + } +} diff --git a/server/aws/apigatewayv2/deployments_handler.go b/server/aws/apigatewayv2/deployments_handler.go new file mode 100644 index 000000000..6cf1bbb3f --- /dev/null +++ b/server/aws/apigatewayv2/deployments_handler.go @@ -0,0 +1,67 @@ +package apigatewayv2 + +import ( + "net/http" + + "github.com/stackshy/cloudemu/v2/services/apigatewayv2/driver" +) + +// serveDeployments handles /v2/apis/{apiId}/deployments: GET=GetDeployments, +// POST=CreateDeployment. +func (h *Handler) serveDeployments(w http.ResponseWriter, r *http.Request, apiID string) { + switch r.Method { + case http.MethodGet: + serveList(w, func() ([]driver.Deployment, string, error) { + return h.ag.GetDeployments(r.Context(), apiID, pageInput(r)) + }, toDeploymentResponse) + case http.MethodPost: + h.createDeployment(w, r, apiID) + default: + writeMethodNotAllowed(w) + } +} + +func (h *Handler) createDeployment(w http.ResponseWriter, r *http.Request, apiID string) { + var req deploymentRequest + if !decodeJSON(w, r, &req) { + return + } + + d, err := h.ag.CreateDeployment(r.Context(), apiID, &driver.CreateDeploymentInput{ + Description: req.Description, StageName: req.StageName, + }) + if err != nil { + writeErr(w, err) + return + } + + writeJSON(w, http.StatusCreated, toDeploymentResponse(d)) +} + +// serveDeploymentItem handles /v2/apis/{apiId}/deployments/{deploymentId}: +// GET, PATCH, DELETE. +func (h *Handler) serveDeploymentItem(w http.ResponseWriter, r *http.Request, apiID, deploymentID string) { + serveItem(w, r, + func() (*driver.Deployment, error) { return h.ag.GetDeployment(r.Context(), apiID, deploymentID) }, + toDeploymentResponse, + func() { h.updateDeployment(w, r, apiID, deploymentID) }, + func() error { return h.ag.DeleteDeployment(r.Context(), apiID, deploymentID) }, + ) +} + +func (h *Handler) updateDeployment(w http.ResponseWriter, r *http.Request, apiID, deploymentID string) { + var req updateDeploymentRequest + if !decodeJSON(w, r, &req) { + return + } + + d, err := h.ag.UpdateDeployment(r.Context(), apiID, deploymentID, &driver.UpdateDeploymentInput{ + Description: req.Description, + }) + if err != nil { + writeErr(w, err) + return + } + + writeJSON(w, http.StatusOK, toDeploymentResponse(d)) +} diff --git a/server/aws/apigatewayv2/handler.go b/server/aws/apigatewayv2/handler.go index 05bb91863..fd6163366 100644 --- a/server/aws/apigatewayv2/handler.go +++ b/server/aws/apigatewayv2/handler.go @@ -1,7 +1,7 @@ // Package apigatewayv2 implements the Amazon API Gateway v2 (HTTP/WebSocket // APIs) control-plane protocol as a server.Handler. It serves the restJson1 -// management API rooted at /v2/apis: Api CRUD plus its Route, Integration and -// Stage sub-collections. +// management API rooted at /v2/apis: Api CRUD plus its Route, Integration, +// Stage and Deployment sub-collections, and resource tagging at /v2/tags. // // This is a distinct service from API Gateway REST v1 (server/aws/apigateway, // rooted at /restapis): the two share no path prefix, so registering this @@ -20,6 +20,7 @@ import ( const ( controlPrefix = "/v2/apis" + tagsPrefix = "/v2/tags/" contentTypeJSON = "application/json" maxBodyBytes = 6 << 20 ) @@ -29,6 +30,7 @@ const ( subRoutes = "routes" subIntegrations = "integrations" subStages = "stages" + subDeployments = "deployments" ) // Path segment counts after the /v2/apis prefix is stripped. @@ -48,15 +50,23 @@ func New(d driver.APIGatewayV2) *Handler { return &Handler{ag: d} } -// Matches claims control-plane requests under /v2/apis. This prefix is disjoint -// from API Gateway REST v1 (/restapis) and every other AWS handler; it must -// register before S3's permissive REST catch-all. +// Matches claims control-plane requests under /v2/apis and the tagging API +// under /v2/tags/{arn}. No other AWS service uses these prefixes, and they are +// disjoint from API Gateway REST v1 (/restapis, /tags); they must register +// before S3's permissive REST catch-all. func (*Handler) Matches(r *http.Request) bool { - return r.URL.Path == "/v2/apis" || strings.HasPrefix(r.URL.Path, controlPrefix+"/") + return r.URL.Path == "/v2/apis" || strings.HasPrefix(r.URL.Path, controlPrefix+"/") || + strings.HasPrefix(r.URL.Path, tagsPrefix) } -// ServeHTTP routes the restJson1 management API under /v2/apis by segment count. +// ServeHTTP routes the restJson1 management API under /v2/apis by segment +// count, and /v2/tags/{arn} to the tagging API. func (h *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) { + if strings.HasPrefix(r.URL.Path, tagsPrefix) { + h.serveTags(w, r) + return + } + rest := strings.Trim(strings.TrimPrefix(r.URL.Path, controlPrefix), "/") if rest == "" { @@ -81,7 +91,7 @@ func (h *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) { func (h *Handler) serveCollection(w http.ResponseWriter, r *http.Request) { switch r.Method { case http.MethodGet: - serveList(w, func() ([]driver.API, error) { return h.ag.GetAPIs(r.Context()) }, toAPIResponse) + serveList(w, func() ([]driver.API, string, error) { return h.ag.GetAPIs(r.Context(), pageInput(r)) }, toAPIResponse) case http.MethodPost: h.createAPI(w, r) default: @@ -103,6 +113,7 @@ func (h *Handler) createAPI(w http.ResponseWriter, r *http.Request) { DisableExecuteAPIEndpoint: req.DisableExecuteAPIEndpoint, Tags: req.Tags, CorsConfiguration: corsToDriver(req.CorsConfiguration), + Target: req.Target, RouteKey: req.RouteKey, CredentialsArn: req.CredentialsArn, }) if err != nil { writeErr(w, err) @@ -149,6 +160,7 @@ func (h *Handler) updateAPI(w http.ResponseWriter, r *http.Request, apiID string APIKeySelectionExpression: req.APIKeySelectionExpression, DisableExecuteAPIEndpoint: req.DisableExecuteAPIEndpoint, CorsConfiguration: corsToDriver(req.CorsConfiguration), + Target: req.Target, RouteKey: req.RouteKey, CredentialsArn: req.CredentialsArn, }) if err != nil { writeErr(w, err) @@ -167,6 +179,8 @@ func (h *Handler) serveSubCollection(w http.ResponseWriter, r *http.Request, api h.serveIntegrations(w, r, apiID) case subStages: h.serveStages(w, r, apiID) + case subDeployments: + h.serveDeployments(w, r, apiID) default: writeError(w, http.StatusNotFound, "NotFoundException", "unsupported apigatewayv2 path") } @@ -181,21 +195,31 @@ func (h *Handler) serveSubItem(w http.ResponseWriter, r *http.Request, apiID, su h.serveIntegrationItem(w, r, apiID, item) case subStages: h.serveStageItem(w, r, apiID, item) + case subDeployments: + h.serveDeploymentItem(w, r, apiID, item) default: writeError(w, http.StatusNotFound, "NotFoundException", "unsupported apigatewayv2 path") } } -// collectionResponse is the {"items":[...]} envelope every apigatewayv2 list -// operation (GetApis/GetRoutes/GetIntegrations/GetStages) returns. +// collectionResponse is the {"items":[...],"nextToken":...} envelope every +// apigatewayv2 list operation returns. type collectionResponse[R any] struct { - Items []R `json:"items"` + Items []R `json:"items"` + NextToken string `json:"nextToken,omitempty"` +} + +// pageInput reads the maxResults and nextToken query parameters. +func pageInput(r *http.Request) *driver.PageInput { + q := r.URL.Query() + + return &driver.PageInput{MaxResults: q.Get("maxResults"), NextToken: q.Get("nextToken")} } // serveList renders a GET list: it runs list, maps each element through render -// and writes the {"items":[...]} envelope. -func serveList[T, R any](w http.ResponseWriter, list func() ([]T, error), render func(*T) R) { - items, err := list() +// and writes the {"items":[...]} envelope with the next page token. +func serveList[T, R any](w http.ResponseWriter, list func() ([]T, string, error), render func(*T) R) { + items, next, err := list() if err != nil { writeErr(w, err) return @@ -206,7 +230,7 @@ func serveList[T, R any](w http.ResponseWriter, list func() ([]T, error), render out = append(out, render(&items[i])) } - writeJSON(w, http.StatusOK, collectionResponse[R]{Items: out}) + writeJSON(w, http.StatusOK, collectionResponse[R]{Items: out, NextToken: next}) } // serveItem dispatches the GET/PATCH/DELETE shape shared by every apigatewayv2 diff --git a/server/aws/apigatewayv2/helpers_e2e_test.go b/server/aws/apigatewayv2/helpers_e2e_test.go new file mode 100644 index 000000000..2c0dd466a --- /dev/null +++ b/server/aws/apigatewayv2/helpers_e2e_test.go @@ -0,0 +1,109 @@ +package apigatewayv2_test + +import ( + "context" + "encoding/json" + "io" + "net/http" + "strings" + "testing" +) + +// wireErr is the decoded shape of an apigatewayv2 error response. +type wireErr struct { + status int + errType string + message string +} + +// doErr issues a request that is expected to fail and returns its status, +// X-Amzn-Errortype header and message. +func doErr(t *testing.T, method, url, body string) wireErr { + t.Helper() + + req, err := http.NewRequestWithContext(context.Background(), method, url, strings.NewReader(body)) + if err != nil { + t.Fatalf("new request: %v", err) + } + + req.Header.Set("Content-Type", "application/json") + + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("%s %s: %v", method, url, err) + } + defer resp.Body.Close() + + raw, _ := io.ReadAll(resp.Body) + + var out struct { + Message string `json:"message"` + } + _ = json.Unmarshal(raw, &out) + + return wireErr{status: resp.StatusCode, errType: resp.Header.Get("X-Amzn-Errortype"), message: out.Message} +} + +// wantErr asserts a request fails with the given status, error type and message. +func wantErr(t *testing.T, method, url, body string, status int, errType, message string) { + t.Helper() + + got := doErr(t, method, url, body) + if got.status != status || got.errType != errType || got.message != message { + t.Fatalf("%s %s = %d %s %q, want %d %s %q", method, url, + got.status, got.errType, got.message, status, errType, message) + } +} + +// mustDo issues a request and fails the test unless it returns wantStatus. +func mustDo(t *testing.T, method, url, body string, wantStatus int) map[string]any { + t.Helper() + + status, out := do(t, method, url, body) + if status != wantStatus { + t.Fatalf("%s %s = %d, want %d: %v", method, url, status, wantStatus, out) + } + + return out +} + +// newAPI creates an API of the given protocol and returns its id. +func newAPI(t *testing.T, base, body string) string { + t.Helper() + + out := mustDo(t, http.MethodPost, base+"/v2/apis", body, http.StatusCreated) + + id, _ := out["apiId"].(string) + if id == "" { + t.Fatalf("CreateApi gave no apiId: %v", out) + } + + return id +} + +// newHTTPAPI creates a plain HTTP API and returns its id. +func newHTTPAPI(t *testing.T, base string) string { + t.Helper() + + return newAPI(t, base, `{"name":"http","protocolType":"HTTP"}`) +} + +// newLambdaIntegration adds an AWS_PROXY integration and returns its id. +func newLambdaIntegration(t *testing.T, apiBase string) string { + t.Helper() + + out := mustDo(t, http.MethodPost, apiBase+"/integrations", + `{"integrationType":"AWS_PROXY","integrationUri":"arn:aws:lambda:us-east-1:000000000000:function:f",`+ + `"payloadFormatVersion":"2.0"}`, http.StatusCreated) + + id, _ := out["integrationId"].(string) + + return id +} + +// items returns the items array of a list response. +func items(out map[string]any) []any { + its, _ := out["items"].([]any) + + return its +} diff --git a/server/aws/apigatewayv2/integrations_handler.go b/server/aws/apigatewayv2/integrations_handler.go index 79b9522b2..2a7d09d4a 100644 --- a/server/aws/apigatewayv2/integrations_handler.go +++ b/server/aws/apigatewayv2/integrations_handler.go @@ -12,7 +12,9 @@ func (h *Handler) serveIntegrations(w http.ResponseWriter, r *http.Request, apiI switch r.Method { case http.MethodGet: serveList(w, - func() ([]driver.Integration, error) { return h.ag.GetIntegrations(r.Context(), apiID) }, + func() ([]driver.Integration, string, error) { + return h.ag.GetIntegrations(r.Context(), apiID, pageInput(r)) + }, toIntegrationResponse) case http.MethodPost: h.createIntegration(w, r, apiID) @@ -32,6 +34,7 @@ func (h *Handler) createIntegration(w http.ResponseWriter, r *http.Request, apiI IntegrationMethod: req.IntegrationMethod, ConnectionType: req.ConnectionType, PayloadFormatVersion: req.PayloadFormatVersion, TimeoutInMillis: req.TimeoutInMillis, Description: req.Description, RequestParameters: req.RequestParameters, + CredentialsArn: req.CredentialsArn, }) if err != nil { writeErr(w, err) @@ -63,6 +66,7 @@ func (h *Handler) updateIntegration(w http.ResponseWriter, r *http.Request, apiI IntegrationMethod: req.IntegrationMethod, ConnectionType: req.ConnectionType, PayloadFormatVersion: req.PayloadFormatVersion, TimeoutInMillis: req.TimeoutInMillis, Description: req.Description, RequestParameters: req.RequestParameters, + CredentialsArn: req.CredentialsArn, }) if err != nil { writeErr(w, err) diff --git a/server/aws/apigatewayv2/pagination_e2e_test.go b/server/aws/apigatewayv2/pagination_e2e_test.go new file mode 100644 index 000000000..dbe8a97b8 --- /dev/null +++ b/server/aws/apigatewayv2/pagination_e2e_test.go @@ -0,0 +1,93 @@ +package apigatewayv2_test + +import ( + "net/http" + "net/url" + "testing" +) + +// TestE2E_PaginationAcrossCollections walks GetApis and every sub-collection +// with maxResults=2, the way the CLI paginator does. +func TestE2E_PaginationAcrossCollections(t *testing.T) { + ts := newE2E(t) + + for range 4 { + newHTTPAPI(t, ts.URL) + } + + apiID := newHTTPAPI(t, ts.URL) + apiBase := ts.URL + "/v2/apis/" + apiID + igID := newLambdaIntegration(t, apiBase) + + for _, k := range []string{"GET /a", "GET /b", "GET /c"} { + mustDo(t, http.MethodPost, apiBase+"/routes", `{"routeKey":"`+k+`","target":"integrations/`+igID+`"}`, http.StatusCreated) + } + + newLambdaIntegration(t, apiBase) + newLambdaIntegration(t, apiBase) + + for _, s := range []string{"a", "b", "c"} { + mustDo(t, http.MethodPost, apiBase+"/stages", `{"stageName":"`+s+`"}`, http.StatusCreated) + mustDo(t, http.MethodPost, apiBase+"/deployments", `{"stageName":"`+s+`"}`, http.StatusCreated) + } + + cases := []struct { + coll, key string + want int + }{ + {ts.URL + "/v2/apis", "apiId", 5}, + {apiBase + "/routes", "routeId", 3}, + {apiBase + "/integrations", "integrationId", 3}, + {apiBase + "/stages", "stageName", 3}, + {apiBase + "/deployments", "deploymentId", 3}, + } + + for _, tc := range cases { + coll, want := tc.coll, tc.want + seen := map[string]bool{} + token := "" + + for page := 0; ; page++ { + u := coll + "?maxResults=2" + if token != "" { + u += "&nextToken=" + url.QueryEscape(token) + } + + out := mustDo(t, http.MethodGet, u, "", http.StatusOK) + if n := len(items(out)); n > 2 { + t.Fatalf("%s page %d has %d items, want <= 2", coll, page, n) + } + + for _, it := range items(out) { + m, _ := it.(map[string]any) + if v, _ := m[tc.key].(string); v != "" { + seen[v] = true + } + } + + token, _ = out["nextToken"].(string) + if token == "" { + break + } + } + + if len(seen) != want { + t.Fatalf("%s paged %d distinct items, want %d", coll, len(seen), want) + } + } +} + +// TestE2E_PaginationBounds covers malformed maxResults and nextToken values. +func TestE2E_PaginationBounds(t *testing.T) { + ts := newE2E(t) + + for _, q := range []string{"maxResults=abc", "maxResults=0", "maxResults=501"} { + wantErr(t, http.MethodGet, ts.URL+"/v2/apis?"+q, "", + http.StatusBadRequest, "BadRequestException", "MaxResults must be an integer between 1 and 500") + } + + wantErr(t, http.MethodGet, ts.URL+"/v2/apis?nextToken=%21%21bad", "", + http.StatusBadRequest, "BadRequestException", "Invalid NextToken specified") + + mustDo(t, http.MethodGet, ts.URL+"/v2/apis?maxResults=500", "", http.StatusOK) +} diff --git a/server/aws/apigatewayv2/quickcreate_e2e_test.go b/server/aws/apigatewayv2/quickcreate_e2e_test.go new file mode 100644 index 000000000..6edfe6bc0 --- /dev/null +++ b/server/aws/apigatewayv2/quickcreate_e2e_test.go @@ -0,0 +1,88 @@ +package apigatewayv2_test + +import ( + "net/http" + "testing" +) + +const lambdaARN = "arn:aws:lambda:us-east-1:000000000000:function:quick" + +// TestE2E_QuickCreateLambda proves CreateApi with Target builds the managed +// integration, route and auto-deployed $default stage. +func TestE2E_QuickCreateLambda(t *testing.T) { + ts := newE2E(t) + apiID := newAPI(t, ts.URL, `{"name":"q","protocolType":"HTTP","target":"`+lambdaARN+`","routeKey":"GET /pets",`+ + `"credentialsArn":"arn:aws:iam::000000000000:role/apigw"}`) + apiBase := ts.URL + "/v2/apis/" + apiID + + igs := items(mustDo(t, http.MethodGet, apiBase+"/integrations", "", http.StatusOK)) + if len(igs) != 1 { + t.Fatalf("integrations = %v, want 1", igs) + } + + ig, _ := igs[0].(map[string]any) + if ig["integrationType"] != "AWS_PROXY" || ig["integrationUri"] != lambdaARN || ig["payloadFormatVersion"] != "2.0" || + ig["apiGatewayManaged"] != true || ig["credentialsArn"] != "arn:aws:iam::000000000000:role/apigw" { + t.Fatalf("quick-create integration = %v", ig) + } + + routes := items(mustDo(t, http.MethodGet, apiBase+"/routes", "", http.StatusOK)) + if len(routes) != 1 { + t.Fatalf("routes = %v, want 1", routes) + } + + rt, _ := routes[0].(map[string]any) + if rt["routeKey"] != "GET /pets" || rt["target"] != "integrations/"+ig["integrationId"].(string) || rt["apiGatewayManaged"] != true { + t.Fatalf("quick-create route = %v", rt) + } + + stage := mustDo(t, http.MethodGet, apiBase+"/stages/$default", "", http.StatusOK) + if stage["autoDeploy"] != true || stage["apiGatewayManaged"] != true || stage["deploymentId"] == nil { + t.Fatalf("quick-create stage = %v", stage) + } + + wantErr(t, http.MethodDelete, apiBase+"/integrations/"+ig["integrationId"].(string), "", + http.StatusBadRequest, "BadRequestException", "Cannot delete an integration managed by API Gateway") + + wantErr(t, http.MethodDelete, apiBase+"/stages/$default", "", + http.StatusBadRequest, "BadRequestException", "Cannot modify or delete a stage managed by API Gateway") + + // UpdateApi with a URL target switches the managed integration to HTTP_PROXY. + mustDo(t, http.MethodPatch, apiBase, `{"target":"https://example.com/pets"}`, http.StatusOK) + + upd := mustDo(t, http.MethodGet, apiBase+"/integrations/"+ig["integrationId"].(string), "", http.StatusOK) + if upd["integrationType"] != "HTTP_PROXY" || upd["integrationUri"] != "https://example.com/pets" || upd["payloadFormatVersion"] != "1.0" { + t.Fatalf("UpdateApi target integration = %v", upd) + } +} + +// TestE2E_QuickCreateDefaultsAndErrors covers the $default route key default +// and the quick-create input rules. +func TestE2E_QuickCreateDefaultsAndErrors(t *testing.T) { + ts := newE2E(t) + apiID := newAPI(t, ts.URL, `{"name":"q","protocolType":"HTTP","target":"https://example.com"}`) + + routes := items(mustDo(t, http.MethodGet, ts.URL+"/v2/apis/"+apiID+"/routes", "", http.StatusOK)) + rt, _ := routes[0].(map[string]any) + + if rt["routeKey"] != "$default" { + t.Fatalf("quick-create default route key = %v", rt["routeKey"]) + } + + igs := items(mustDo(t, http.MethodGet, ts.URL+"/v2/apis/"+apiID+"/integrations", "", http.StatusOK)) + ig, _ := igs[0].(map[string]any) + + if ig["integrationType"] != "HTTP_PROXY" || ig["integrationMethod"] != "ANY" || ig["payloadFormatVersion"] != "1.0" { + t.Fatalf("quick-create HTTP integration = %v", ig) + } + + wantErr(t, http.MethodPost, ts.URL+"/v2/apis", `{"name":"q","protocolType":"HTTP","routeKey":"GET /x"}`, + http.StatusBadRequest, "BadRequestException", "RouteKey and CredentialsArn can only be specified with Target") + + wantErr(t, http.MethodPost, ts.URL+"/v2/apis", + `{"name":"q","protocolType":"WEBSOCKET","routeSelectionExpression":"$request.body.action","target":"`+lambdaARN+`"}`, + http.StatusBadRequest, "BadRequestException", "Quick create is only supported for HTTP APIs") + + wantErr(t, http.MethodPost, ts.URL+"/v2/apis", `{"name":"q","protocolType":"HTTP","target":"not-a-target"}`, + http.StatusBadRequest, "BadRequestException", "Target must be a Lambda function ARN or an HTTP(S) URL") +} diff --git a/server/aws/apigatewayv2/routes_handler.go b/server/aws/apigatewayv2/routes_handler.go index 3ba48b9a0..d511f7b86 100644 --- a/server/aws/apigatewayv2/routes_handler.go +++ b/server/aws/apigatewayv2/routes_handler.go @@ -10,7 +10,9 @@ import ( func (h *Handler) serveRoutes(w http.ResponseWriter, r *http.Request, apiID string) { switch r.Method { case http.MethodGet: - serveList(w, func() ([]driver.Route, error) { return h.ag.GetRoutes(r.Context(), apiID) }, toRouteResponse) + serveList(w, func() ([]driver.Route, string, error) { + return h.ag.GetRoutes(r.Context(), apiID, pageInput(r)) + }, toRouteResponse) case http.MethodPost: h.createRoute(w, r, apiID) default: @@ -58,6 +60,7 @@ func (h *Handler) updateRoute(w http.ResponseWriter, r *http.Request, apiID, rou RouteKey: req.RouteKey, Target: req.Target, AuthorizationType: req.AuthorizationType, APIKeyRequired: req.APIKeyRequired, AuthorizerID: req.AuthorizerID, OperationName: req.OperationName, + AuthorizationScopes: req.AuthorizationScopes, }) if err != nil { writeErr(w, err) diff --git a/server/aws/apigatewayv2/stages_handler.go b/server/aws/apigatewayv2/stages_handler.go index c3b33fc9c..5ddda53fd 100644 --- a/server/aws/apigatewayv2/stages_handler.go +++ b/server/aws/apigatewayv2/stages_handler.go @@ -10,7 +10,9 @@ import ( func (h *Handler) serveStages(w http.ResponseWriter, r *http.Request, apiID string) { switch r.Method { case http.MethodGet: - serveList(w, func() ([]driver.Stage, error) { return h.ag.GetStages(r.Context(), apiID) }, toStageResponse) + serveList(w, func() ([]driver.Stage, string, error) { + return h.ag.GetStages(r.Context(), apiID, pageInput(r)) + }, toStageResponse) case http.MethodPost: h.createStage(w, r, apiID) default: @@ -28,6 +30,7 @@ func (h *Handler) createStage(w http.ResponseWriter, r *http.Request, apiID stri StageName: req.StageName, Description: req.Description, AutoDeploy: req.AutoDeploy, DeploymentID: req.DeploymentID, StageVariables: req.StageVariables, DefaultRouteSettings: routeSettingsToDriver(req.DefaultRouteSettings), + Tags: req.Tags, }) if err != nil { writeErr(w, err) diff --git a/server/aws/apigatewayv2/tags_e2e_test.go b/server/aws/apigatewayv2/tags_e2e_test.go new file mode 100644 index 000000000..ed74b487d --- /dev/null +++ b/server/aws/apigatewayv2/tags_e2e_test.go @@ -0,0 +1,79 @@ +package apigatewayv2_test + +import ( + "net/http" + "net/url" + "testing" +) + +// tagsURL builds the /v2/tags/{resource-arn} URL with the ARN path-escaped the +// way the SDKs send it. +func tagsURL(base, arn string) string { + return base + "/v2/tags/" + url.PathEscape(arn) +} + +// TestE2E_TagsOnAPIAndStage covers Create* tags plus TagResource, UntagResource +// and GetTags on api and stage ARNs. +func TestE2E_TagsOnAPIAndStage(t *testing.T) { + ts := newE2E(t) + apiID := newAPI(t, ts.URL, `{"name":"t","protocolType":"HTTP","tags":{"env":"dev"}}`) + apiBase := ts.URL + "/v2/apis/" + apiID + apiARN := "arn:aws:apigateway:us-east-1::/apis/" + apiID + + got := mustDo(t, http.MethodGet, tagsURL(ts.URL, apiARN), "", http.StatusOK) + if tags, _ := got["tags"].(map[string]any); tags["env"] != "dev" { + t.Fatalf("GetTags(api) = %v", got) + } + + mustDo(t, http.MethodPost, tagsURL(ts.URL, apiARN), `{"tags":{"team":"core","env":"prod"}}`, http.StatusCreated) + + api := mustDo(t, http.MethodGet, apiBase, "", http.StatusOK) + if tags, _ := api["tags"].(map[string]any); tags["team"] != "core" || tags["env"] != "prod" { + t.Fatalf("GetApi tags = %v", api["tags"]) + } + + list := mustDo(t, http.MethodGet, ts.URL+"/v2/apis", "", http.StatusOK) + first, _ := items(list)[0].(map[string]any) + + if tags, _ := first["tags"].(map[string]any); tags["team"] != "core" { + t.Fatalf("GetApis tags = %v", first["tags"]) + } + + mustDo(t, http.MethodDelete, tagsURL(ts.URL, apiARN)+"?tagKeys=env&tagKeys=team", "", http.StatusNoContent) + + got = mustDo(t, http.MethodGet, tagsURL(ts.URL, apiARN), "", http.StatusOK) + if tags, ok := got["tags"].(map[string]any); !ok || len(tags) != 0 { + t.Fatalf("GetTags after untag = %v", got) + } + + mustDo(t, http.MethodPost, apiBase+"/stages", `{"stageName":"prod","tags":{"tier":"gold"}}`, http.StatusCreated) + stageARN := apiARN + "/stages/prod" + + mustDo(t, http.MethodPost, tagsURL(ts.URL, stageARN), `{"tags":{"owner":"me"}}`, http.StatusCreated) + + stage := mustDo(t, http.MethodGet, apiBase+"/stages/prod", "", http.StatusOK) + if tags, _ := stage["tags"].(map[string]any); tags["tier"] != "gold" || tags["owner"] != "me" { + t.Fatalf("GetStage tags = %v", stage["tags"]) + } +} + +// TestE2E_TagErrors covers unknown resources and malformed ARNs. +func TestE2E_TagErrors(t *testing.T) { + ts := newE2E(t) + + wantErr(t, http.MethodGet, tagsURL(ts.URL, "arn:aws:apigateway:us-east-1::/apis/missing123"), "", + http.StatusNotFound, "NotFoundException", "Invalid API identifier specified missing123") + + wantErr(t, http.MethodGet, tagsURL(ts.URL, "arn:aws:apigateway:us-east-1::/domainnames/api.example.com"), "", + http.StatusNotFound, "NotFoundException", "Invalid domain name identifier specified api.example.com") + + wantErr(t, http.MethodGet, tagsURL(ts.URL, "arn:aws:apigateway:us-east-1::/vpclinks/abc123"), "", + http.StatusNotFound, "NotFoundException", "Invalid VpcLink identifier specified abc123") + + wantErr(t, http.MethodGet, tagsURL(ts.URL, "not-an-arn"), "", + http.StatusBadRequest, "BadRequestException", "Invalid resource ARN specified not-an-arn") + + apiID := newHTTPAPI(t, ts.URL) + wantErr(t, http.MethodPost, tagsURL(ts.URL, "arn:aws:apigateway:us-east-1::/apis/"+apiID), `{"tags":{"aws:x":"y"}}`, + http.StatusBadRequest, "BadRequestException", "Tag keys cannot start with the reserved prefix aws:") +} diff --git a/server/aws/apigatewayv2/tags_handler.go b/server/aws/apigatewayv2/tags_handler.go new file mode 100644 index 000000000..c5c8d4622 --- /dev/null +++ b/server/aws/apigatewayv2/tags_handler.go @@ -0,0 +1,50 @@ +package apigatewayv2 + +import ( + "net/http" + "net/url" + "strings" +) + +// serveTags handles /v2/tags/{resource-arn}: GET=GetTags, POST=TagResource, +// DELETE=UntagResource (tagKeys in the query string). The ARN is one +// path-escaped label, so it is read from the escaped path. +func (h *Handler) serveTags(w http.ResponseWriter, r *http.Request) { + arn, err := url.PathUnescape(strings.TrimPrefix(r.URL.EscapedPath(), tagsPrefix)) + if err != nil { + writeError(w, http.StatusBadRequest, "BadRequestException", "Invalid resource ARN specified") + return + } + + switch r.Method { + case http.MethodGet: + tags, err := h.ag.GetTags(r.Context(), arn) + if err != nil { + writeErr(w, err) + return + } + + writeJSON(w, http.StatusOK, tagsBody{Tags: tags}) + case http.MethodPost: + var req tagsBody + if !decodeJSON(w, r, &req) { + return + } + + if err := h.ag.TagResource(r.Context(), arn, req.Tags); err != nil { + writeErr(w, err) + return + } + + writeJSON(w, http.StatusCreated, struct{}{}) + case http.MethodDelete: + if err := h.ag.UntagResource(r.Context(), arn, r.URL.Query()["tagKeys"]); err != nil { + writeErr(w, err) + return + } + + w.WriteHeader(http.StatusNoContent) + default: + writeMethodNotAllowed(w) + } +} diff --git a/server/aws/apigatewayv2/types.go b/server/aws/apigatewayv2/types.go index 6888962fa..ef68e5b40 100644 --- a/server/aws/apigatewayv2/types.go +++ b/server/aws/apigatewayv2/types.go @@ -51,6 +51,9 @@ type createAPIRequest struct { DisableExecuteAPIEndpoint bool `json:"disableExecuteApiEndpoint"` CorsConfiguration *corsWire `json:"corsConfiguration"` Tags map[string]string `json:"tags"` + Target string `json:"target"` + RouteKey string `json:"routeKey"` + CredentialsArn string `json:"credentialsArn"` } // updateAPIRequest is the UpdateApi (PATCH) request body. Pointer fields @@ -64,6 +67,9 @@ type updateAPIRequest struct { APIKeySelectionExpression *string `json:"apiKeySelectionExpression"` DisableExecuteAPIEndpoint *bool `json:"disableExecuteApiEndpoint"` CorsConfiguration *corsWire `json:"corsConfiguration"` + Target *string `json:"target"` + RouteKey *string `json:"routeKey"` + CredentialsArn *string `json:"credentialsArn"` } // apiResponse is the API wire object. @@ -79,7 +85,7 @@ type apiResponse struct { APIEndpoint string `json:"apiEndpoint"` CreatedDate string `json:"createdDate"` CorsConfiguration *corsWire `json:"corsConfiguration,omitempty"` - Tags map[string]string `json:"tags,omitempty"` + Tags map[string]string `json:"tags"` } func toAPIResponse(a *driver.API) apiResponse { @@ -92,10 +98,19 @@ func toAPIResponse(a *driver.API) apiResponse { APIEndpoint: a.APIEndpoint, CreatedDate: isoTime(a.CreatedDate), CorsConfiguration: corsFromDriver(a.CorsConfiguration), - Tags: a.Tags, + Tags: tagsOrEmpty(a.Tags), } } +// tagsOrEmpty renders a nil tag map as {} so every resource carries "tags". +func tagsOrEmpty(tags map[string]string) map[string]string { + if tags == nil { + return map[string]string{} + } + + return tags +} + // routeRequest is the CreateRoute/UpdateRoute request body. On create the // pointer/value split does not matter; UpdateRoute reads only the pointers. type routeRequest struct { @@ -110,12 +125,13 @@ type routeRequest struct { // updateRouteRequest is the UpdateRoute (PATCH) request body. type updateRouteRequest struct { - RouteKey *string `json:"routeKey"` - Target *string `json:"target"` - AuthorizationType *string `json:"authorizationType"` - APIKeyRequired *bool `json:"apiKeyRequired"` - AuthorizerID *string `json:"authorizerId"` - OperationName *string `json:"operationName"` + RouteKey *string `json:"routeKey"` + Target *string `json:"target"` + AuthorizationType *string `json:"authorizationType"` + APIKeyRequired *bool `json:"apiKeyRequired"` + AuthorizerID *string `json:"authorizerId"` + OperationName *string `json:"operationName"` + AuthorizationScopes []string `json:"authorizationScopes"` } // routeResponse is the Route wire object. @@ -128,6 +144,7 @@ type routeResponse struct { AuthorizerID string `json:"authorizerId,omitempty"` AuthorizationScopes []string `json:"authorizationScopes,omitempty"` OperationName string `json:"operationName,omitempty"` + APIGatewayManaged bool `json:"apiGatewayManaged,omitempty"` } func toRouteResponse(r *driver.Route) routeResponse { @@ -135,7 +152,7 @@ func toRouteResponse(r *driver.Route) routeResponse { RouteID: r.RouteID, RouteKey: r.RouteKey, Target: r.Target, AuthorizationType: r.AuthorizationType, APIKeyRequired: r.APIKeyRequired, AuthorizerID: r.AuthorizerID, AuthorizationScopes: r.AuthorizationScopes, - OperationName: r.OperationName, + OperationName: r.OperationName, APIGatewayManaged: r.APIGatewayManaged, } } @@ -149,6 +166,7 @@ type integrationRequest struct { TimeoutInMillis int `json:"timeoutInMillis"` Description string `json:"description"` RequestParameters map[string]string `json:"requestParameters"` + CredentialsArn string `json:"credentialsArn"` } // updateIntegrationRequest is the UpdateIntegration (PATCH) request body. @@ -161,6 +179,7 @@ type updateIntegrationRequest struct { TimeoutInMillis *int `json:"timeoutInMillis"` Description *string `json:"description"` RequestParameters map[string]string `json:"requestParameters"` + CredentialsArn *string `json:"credentialsArn"` } // integrationResponse is the Integration wire object. @@ -174,6 +193,8 @@ type integrationResponse struct { TimeoutInMillis int `json:"timeoutInMillis,omitempty"` Description string `json:"description,omitempty"` RequestParameters map[string]string `json:"requestParameters,omitempty"` + CredentialsArn string `json:"credentialsArn,omitempty"` + APIGatewayManaged bool `json:"apiGatewayManaged,omitempty"` } func toIntegrationResponse(i *driver.Integration) integrationResponse { @@ -182,7 +203,8 @@ func toIntegrationResponse(i *driver.Integration) integrationResponse { IntegrationURI: i.IntegrationURI, IntegrationMethod: i.IntegrationMethod, ConnectionType: i.ConnectionType, PayloadFormatVersion: i.PayloadFormatVersion, TimeoutInMillis: i.TimeoutInMillis, Description: i.Description, - RequestParameters: i.RequestParameters, + RequestParameters: i.RequestParameters, CredentialsArn: i.CredentialsArn, + APIGatewayManaged: i.APIGatewayManaged, } } @@ -227,6 +249,7 @@ type stageRequest struct { DeploymentID string `json:"deploymentId"` StageVariables map[string]string `json:"stageVariables"` DefaultRouteSettings *routeSettingsWire `json:"defaultRouteSettings"` + Tags map[string]string `json:"tags"` } // updateStageRequest is the UpdateStage (PATCH) request body. @@ -248,6 +271,10 @@ type stageResponse struct { DefaultRouteSettings *routeSettingsWire `json:"defaultRouteSettings,omitempty"` CreatedDate string `json:"createdDate"` LastUpdatedDate string `json:"lastUpdatedDate"` + Tags map[string]string `json:"tags"` + APIGatewayManaged bool `json:"apiGatewayManaged,omitempty"` + + LastDeploymentStatusMessage string `json:"lastDeploymentStatusMessage,omitempty"` } func toStageResponse(s *driver.Stage) stageResponse { @@ -256,9 +283,46 @@ func toStageResponse(s *driver.Stage) stageResponse { DeploymentID: s.DeploymentID, StageVariables: s.StageVariables, DefaultRouteSettings: routeSettingsFromDriver(s.DefaultRouteSettings), CreatedDate: isoTime(s.CreatedDate), LastUpdatedDate: isoTime(s.LastUpdatedDate), + Tags: tagsOrEmpty(s.Tags), APIGatewayManaged: s.APIGatewayManaged, + + LastDeploymentStatusMessage: s.LastDeploymentStatusMessage, } } +// deploymentRequest is the CreateDeployment request body. +type deploymentRequest struct { + Description string `json:"description"` + StageName string `json:"stageName"` +} + +// updateDeploymentRequest is the UpdateDeployment (PATCH) request body. +type updateDeploymentRequest struct { + Description *string `json:"description"` +} + +// deploymentResponse is the Deployment wire object. +type deploymentResponse struct { + DeploymentID string `json:"deploymentId"` + Description string `json:"description,omitempty"` + CreatedDate string `json:"createdDate"` + DeploymentStatus string `json:"deploymentStatus"` + DeploymentStatusMessage string `json:"deploymentStatusMessage,omitempty"` + AutoDeployed bool `json:"autoDeployed"` +} + +func toDeploymentResponse(d *driver.Deployment) deploymentResponse { + return deploymentResponse{ + DeploymentID: d.DeploymentID, Description: d.Description, CreatedDate: isoTime(d.CreatedDate), + DeploymentStatus: d.DeploymentStatus, DeploymentStatusMessage: d.DeploymentStatusMessage, + AutoDeployed: d.AutoDeployed, + } +} + +// tagsBody is the TagResource request body and the GetTags response body. +type tagsBody struct { + Tags map[string]string `json:"tags"` +} + // isoTime renders a unix-seconds timestamp as the ISO8601 string the // apigatewayv2 restJson1 protocol uses for timestamp fields (createdDate, // lastUpdatedDate). diff --git a/server/aws/apigatewayv2/validation_e2e_test.go b/server/aws/apigatewayv2/validation_e2e_test.go new file mode 100644 index 000000000..9face1484 --- /dev/null +++ b/server/aws/apigatewayv2/validation_e2e_test.go @@ -0,0 +1,107 @@ +package apigatewayv2_test + +import ( + "net/http" + "testing" +) + +const badHTTPRouteKey = `The provided route key is not formatted properly for HTTP protocol. ` + + `Format should be " /" or "$default"` + +// TestE2E_RouteValidation covers the RouteKey format, duplicate keys and the +// integration a route targets. +func TestE2E_RouteValidation(t *testing.T) { + ts := newE2E(t) + apiID := newHTTPAPI(t, ts.URL) + apiBase := ts.URL + "/v2/apis/" + apiID + igID := newLambdaIntegration(t, apiBase) + + for _, key := range []string{"GET", "/items", "FETCH /items", "GET items", "$connect"} { + wantErr(t, http.MethodPost, apiBase+"/routes", `{"routeKey":"`+key+`"}`, + http.StatusBadRequest, "BadRequestException", badHTTPRouteKey) + } + + for _, key := range []string{"$default", "ANY /{proxy+}", "GET /items/{id}"} { + mustDo(t, http.MethodPost, apiBase+"/routes", `{"routeKey":"`+key+`"}`, http.StatusCreated) + } + + wantErr(t, http.MethodPost, apiBase+"/routes", `{"routeKey":"GET /items/{id}"}`, + http.StatusConflict, "ConflictException", "Route with key GET /items/{id} already exists for this API") + + wantErr(t, http.MethodPost, apiBase+"/routes", `{"routeKey":"POST /x","target":"integrations/missing123"}`, + http.StatusBadRequest, "BadRequestException", "Invalid Integration identifier specified") + + rt := mustDo(t, http.MethodPost, apiBase+"/routes", + `{"routeKey":"POST /x","target":"integrations/`+igID+`"}`, http.StatusCreated) + rtID, _ := rt["routeId"].(string) + + wantErr(t, http.MethodPatch, apiBase+"/routes/"+rtID, `{"routeKey":"$default"}`, + http.StatusConflict, "ConflictException", "Route with key $default already exists for this API") + + wantErr(t, http.MethodPatch, apiBase+"/routes/"+rtID, `{"target":"integrations/missing123"}`, + http.StatusBadRequest, "BadRequestException", "Invalid Integration identifier specified") + + upd := mustDo(t, http.MethodPatch, apiBase+"/routes/"+rtID, + `{"authorizationType":"JWT","authorizationScopes":["read","write"]}`, http.StatusOK) + + scopes, _ := upd["authorizationScopes"].([]any) + if upd["authorizationType"] != "JWT" || len(scopes) != 2 { + t.Fatalf("UpdateRoute scopes = %v", upd) + } +} + +// TestE2E_IntegrationValidation covers the integration types and payload +// format versions each protocol accepts. +func TestE2E_IntegrationValidation(t *testing.T) { + ts := newE2E(t) + httpBase := ts.URL + "/v2/apis/" + newHTTPAPI(t, ts.URL) + wsBase := ts.URL + "/v2/apis/" + newAPI(t, ts.URL, + `{"name":"ws","protocolType":"WEBSOCKET","routeSelectionExpression":"$request.body.action"}`) + + wantErr(t, http.MethodPost, httpBase+"/integrations", `{"integrationType":"MOCK"}`, + http.StatusBadRequest, "BadRequestException", "Integration type MOCK is not supported for HTTP APIs") + + wantErr(t, http.MethodPost, httpBase+"/integrations", `{"integrationType":"LAMBDA"}`, + http.StatusBadRequest, "BadRequestException", "Invalid integration type specified: LAMBDA") + + wantErr(t, http.MethodPost, httpBase+"/integrations", + `{"integrationType":"HTTP_PROXY","integrationUri":"https://example.com","integrationMethod":"GET","payloadFormatVersion":"2.0"}`, + http.StatusBadRequest, "BadRequestException", "Payload format version 2.0 is only supported for AWS_PROXY integrations") + + wantErr(t, http.MethodPost, httpBase+"/integrations", + `{"integrationType":"AWS_PROXY","integrationUri":"arn:aws:lambda:us-east-1:000000000000:function:f","payloadFormatVersion":"3.0"}`, + http.StatusBadRequest, "BadRequestException", "Invalid payload format version specified: 3.0") + + wantErr(t, http.MethodPost, httpBase+"/integrations", + `{"integrationType":"AWS_PROXY","integrationUri":"arn:aws:lambda:us-east-1:000000000000:function:f","timeoutInMillis":40000}`, + http.StatusBadRequest, "BadRequestException", "TimeoutInMillis must be between 50 and 30000") + + mustDo(t, http.MethodPost, wsBase+"/integrations", `{"integrationType":"MOCK"}`, http.StatusCreated) + + wantErr(t, http.MethodPost, wsBase+"/integrations", `{"integrationType":"AWS_PROXY","payloadFormatVersion":"2.0"}`, + http.StatusBadRequest, "BadRequestException", "Payload format version 2.0 is not supported for WEBSOCKET APIs") +} + +// TestE2E_APIAndStageValidation covers the WebSocket route selection +// expression and the stage name rules. +func TestE2E_APIAndStageValidation(t *testing.T) { + ts := newE2E(t) + + wantErr(t, http.MethodPost, ts.URL+"/v2/apis", `{"name":"ws","protocolType":"WEBSOCKET"}`, + http.StatusBadRequest, "BadRequestException", "RouteSelectionExpression is required for WEBSOCKET protocol") + + wantErr(t, http.MethodPost, ts.URL+"/v2/apis", + `{"name":"h","protocolType":"HTTP","routeSelectionExpression":"$request.body.action"}`, + http.StatusBadRequest, "BadRequestException", "Only $request.method $request.path is supported for HTTP APIs") + + wsID := newAPI(t, ts.URL, `{"name":"ws","protocolType":"WEBSOCKET","routeSelectionExpression":"$request.body.action"}`) + mustDo(t, http.MethodPost, ts.URL+"/v2/apis/"+wsID+"/routes", `{"routeKey":"$connect"}`, http.StatusCreated) + mustDo(t, http.MethodPost, ts.URL+"/v2/apis/"+wsID+"/routes", `{"routeKey":"sendMessage"}`, http.StatusCreated) + + apiBase := ts.URL + "/v2/apis/" + newHTTPAPI(t, ts.URL) + + wantErr(t, http.MethodPost, apiBase+"/stages", `{"stageName":"bad name!"}`, + http.StatusBadRequest, "BadRequestException", "Stage name only allows a-zA-Z0-9._- or $default") + + mustDo(t, http.MethodPost, apiBase+"/stages", `{"stageName":"prod.v1_a-b"}`, http.StatusCreated) +} diff --git a/services/apigatewayv2/driver/driver.go b/services/apigatewayv2/driver/driver.go index 686cd4c36..c256a83e7 100644 --- a/services/apigatewayv2/driver/driver.go +++ b/services/apigatewayv2/driver/driver.go @@ -25,6 +25,21 @@ const ( IntegrationMock = "MOCK" ) +// Deployment statuses. +const ( + DeploymentStatusPending = "PENDING" + DeploymentStatusFailed = "FAILED" + DeploymentStatusDeployed = "DEPLOYED" +) + +// PageInput carries the MaxResults and NextToken query parameters every +// apigatewayv2 list operation accepts. MaxResults is a string on the wire, so +// it stays one here and the provider validates it. +type PageInput struct { + MaxResults string + NextToken string +} + // API is an apigatewayv2 API (HTTP or WebSocket) control-plane object. type API struct { APIID string @@ -62,6 +77,7 @@ type Route struct { AuthorizerID string AuthorizationScopes []string OperationName string + APIGatewayManaged bool } // Integration is a backend an API's routes forward to. @@ -75,6 +91,8 @@ type Integration struct { TimeoutInMillis int Description string RequestParameters map[string]string + CredentialsArn string + APIGatewayManaged bool } // Stage is a named deployment stage of an API (e.g. "$default", "prod"). @@ -87,6 +105,33 @@ type Stage struct { DefaultRouteSettings *RouteSettings CreatedDate int64 // unix seconds LastUpdatedDate int64 // unix seconds + Tags map[string]string + APIGatewayManaged bool + + LastDeploymentStatusMessage string +} + +// Deployment is an immutable snapshot of an API's routes and integrations +// that a stage serves. +type Deployment struct { + DeploymentID string + Description string + CreatedDate int64 // unix seconds + DeploymentStatus string + DeploymentStatusMessage string + AutoDeployed bool +} + +// CreateDeploymentInput carries the fields CreateDeployment accepts. A +// non-empty StageName also points that stage at the new deployment. +type CreateDeploymentInput struct { + Description string + StageName string +} + +// UpdateDeploymentInput carries the mutable fields UpdateDeployment accepts. +type UpdateDeploymentInput struct { + Description *string } // RouteSettings are the per-route (or default) execution settings on a Stage. @@ -109,6 +154,12 @@ type CreateAPIInput struct { DisableExecuteAPIEndpoint bool Tags map[string]string CorsConfiguration *Cors + + // Quick create: Target creates a managed integration, a route for + // RouteKey (default "$default") and an auto-deployed $default stage. + Target string + RouteKey string + CredentialsArn string } // UpdateAPIInput carries the mutable fields UpdateApi accepts. A nil pointer @@ -121,6 +172,11 @@ type UpdateAPIInput struct { APIKeySelectionExpression *string DisableExecuteAPIEndpoint *bool CorsConfiguration *Cors + + // Quick create fields; they update the managed integration and route. + Target *string + RouteKey *string + CredentialsArn *string } // CreateRouteInput carries the fields CreateRoute accepts. @@ -134,14 +190,16 @@ type CreateRouteInput struct { OperationName string } -// UpdateRouteInput carries the mutable fields UpdateRoute accepts. +// UpdateRouteInput carries the mutable fields UpdateRoute accepts. A nil +// AuthorizationScopes leaves the stored scopes unchanged. type UpdateRouteInput struct { - RouteKey *string - Target *string - AuthorizationType *string - APIKeyRequired *bool - AuthorizerID *string - OperationName *string + RouteKey *string + Target *string + AuthorizationType *string + APIKeyRequired *bool + AuthorizerID *string + OperationName *string + AuthorizationScopes []string } // CreateIntegrationInput carries the fields CreateIntegration accepts. @@ -154,6 +212,7 @@ type CreateIntegrationInput struct { TimeoutInMillis int Description string RequestParameters map[string]string + CredentialsArn string } // UpdateIntegrationInput carries the mutable fields UpdateIntegration accepts. @@ -166,6 +225,7 @@ type UpdateIntegrationInput struct { TimeoutInMillis *int Description *string RequestParameters map[string]string + CredentialsArn *string } // CreateStageInput carries the fields CreateStage accepts. @@ -176,6 +236,7 @@ type CreateStageInput struct { DeploymentID string StageVariables map[string]string DefaultRouteSettings *RouteSettings + Tags map[string]string } // UpdateStageInput carries the mutable fields UpdateStage accepts. @@ -188,29 +249,40 @@ type UpdateStageInput struct { } // APIGatewayV2 is the apigatewayv2 control-plane contract: API CRUD plus its -// Route, Integration and Stage sub-collections. +// Route, Integration, Stage and Deployment sub-collections, and resource +// tagging. Every list returns one page plus the next token ("" on the last). type APIGatewayV2 interface { CreateAPI(ctx context.Context, in *CreateAPIInput) (*API, error) GetAPI(ctx context.Context, apiID string) (*API, error) - GetAPIs(ctx context.Context) ([]API, error) + GetAPIs(ctx context.Context, page *PageInput) ([]API, string, error) UpdateAPI(ctx context.Context, apiID string, in *UpdateAPIInput) (*API, error) DeleteAPI(ctx context.Context, apiID string) error CreateRoute(ctx context.Context, apiID string, in *CreateRouteInput) (*Route, error) GetRoute(ctx context.Context, apiID, routeID string) (*Route, error) - GetRoutes(ctx context.Context, apiID string) ([]Route, error) + GetRoutes(ctx context.Context, apiID string, page *PageInput) ([]Route, string, error) UpdateRoute(ctx context.Context, apiID, routeID string, in *UpdateRouteInput) (*Route, error) DeleteRoute(ctx context.Context, apiID, routeID string) error CreateIntegration(ctx context.Context, apiID string, in *CreateIntegrationInput) (*Integration, error) GetIntegration(ctx context.Context, apiID, integrationID string) (*Integration, error) - GetIntegrations(ctx context.Context, apiID string) ([]Integration, error) + GetIntegrations(ctx context.Context, apiID string, page *PageInput) ([]Integration, string, error) UpdateIntegration(ctx context.Context, apiID, integrationID string, in *UpdateIntegrationInput) (*Integration, error) DeleteIntegration(ctx context.Context, apiID, integrationID string) error CreateStage(ctx context.Context, apiID string, in *CreateStageInput) (*Stage, error) GetStage(ctx context.Context, apiID, stageName string) (*Stage, error) - GetStages(ctx context.Context, apiID string) ([]Stage, error) + GetStages(ctx context.Context, apiID string, page *PageInput) ([]Stage, string, error) UpdateStage(ctx context.Context, apiID, stageName string, in *UpdateStageInput) (*Stage, error) DeleteStage(ctx context.Context, apiID, stageName string) error + + CreateDeployment(ctx context.Context, apiID string, in *CreateDeploymentInput) (*Deployment, error) + GetDeployment(ctx context.Context, apiID, deploymentID string) (*Deployment, error) + GetDeployments(ctx context.Context, apiID string, page *PageInput) ([]Deployment, string, error) + UpdateDeployment(ctx context.Context, apiID, deploymentID string, in *UpdateDeploymentInput) (*Deployment, error) + DeleteDeployment(ctx context.Context, apiID, deploymentID string) error + + TagResource(ctx context.Context, resourceARN string, tags map[string]string) error + UntagResource(ctx context.Context, resourceARN string, tagKeys []string) error + GetTags(ctx context.Context, resourceARN string) (map[string]string, error) } From a3d4d85f03990e22f5c0e45ef9f14c3de3a51db9 Mon Sep 17 00:00:00 2001 From: Nitin Kumar Date: Sun, 27 Sep 2026 18:43:49 +0530 Subject: [PATCH 2/2] fix(aws-apigatewayv2): braced selection expressions, doc-backed limits, stage and CORS rules, integration templates --- providers/aws/apigatewayv2/integrations.go | 11 +++ providers/aws/apigatewayv2/page.go | 7 +- providers/aws/apigatewayv2/stages.go | 15 +++- providers/aws/apigatewayv2/validation.go | 23 ++++-- .../aws/apigatewayv2/deployment_e2e_test.go | 9 +++ .../aws/apigatewayv2/integrations_handler.go | 8 +++ .../aws/apigatewayv2/pagination_e2e_test.go | 5 +- server/aws/apigatewayv2/types.go | 16 +++++ .../aws/apigatewayv2/validation_e2e_test.go | 72 ++++++++++++++++++- services/apigatewayv2/driver/driver.go | 12 ++++ 10 files changed, 158 insertions(+), 20 deletions(-) diff --git a/providers/aws/apigatewayv2/integrations.go b/providers/aws/apigatewayv2/integrations.go index ee1f9f96e..2095a7c55 100644 --- a/providers/aws/apigatewayv2/integrations.go +++ b/providers/aws/apigatewayv2/integrations.go @@ -29,6 +29,10 @@ func (m *Mock) CreateIntegration( Description: in.Description, RequestParameters: copyStrMap(in.RequestParameters), CredentialsArn: in.CredentialsArn, + + RequestTemplates: copyStrMap(in.RequestTemplates), + TemplateSelectionExpression: in.TemplateSelectionExpression, + PassthroughBehavior: in.PassthroughBehavior, } if err := validateIntegration(ad.api.ProtocolType, ig); err != nil { @@ -110,6 +114,12 @@ func (m *Mock) UpdateIntegration( setString(&next.PayloadFormatVersion, in.PayloadFormatVersion) setString(&next.Description, in.Description) setString(&next.CredentialsArn, in.CredentialsArn) + setString(&next.TemplateSelectionExpression, in.TemplateSelectionExpression) + setString(&next.PassthroughBehavior, in.PassthroughBehavior) + + if in.RequestTemplates != nil { + next.RequestTemplates = copyStrMap(in.RequestTemplates) + } if in.TimeoutInMillis != nil { next.TimeoutInMillis = *in.TimeoutInMillis @@ -161,6 +171,7 @@ func (m *Mock) DeleteIntegration(_ context.Context, apiID, integrationID string) func copyIntegration(i *driver.Integration) driver.Integration { out := *i out.RequestParameters = copyStrMap(i.RequestParameters) + out.RequestTemplates = copyStrMap(i.RequestTemplates) return out } diff --git a/providers/aws/apigatewayv2/page.go b/providers/aws/apigatewayv2/page.go index e6d1175c2..7fcaffdfe 100644 --- a/providers/aws/apigatewayv2/page.go +++ b/providers/aws/apigatewayv2/page.go @@ -7,9 +7,6 @@ import ( "github.com/stackshy/cloudemu/v2/services/apigatewayv2/driver" ) -// maxPageSize is the largest MaxResults a list call accepts. -const maxPageSize = 500 - // listPage reads one of an API's sub-collections under its read lock and // returns the page that in selects. func listPage[V, T any]( @@ -45,8 +42,8 @@ func pageOf[T any](items []T, less func(a, b T) bool, in *driver.PageInput) (pag if maxResults != "" { n, convErr := strconv.Atoi(maxResults) - if convErr != nil || n < 1 || n > maxPageSize { - return nil, "", badRequest("MaxResults must be an integer between 1 and %d", maxPageSize) + if convErr != nil || n < 1 { + return nil, "", badRequest("MaxResults must be a positive integer") } size = n diff --git a/providers/aws/apigatewayv2/stages.go b/providers/aws/apigatewayv2/stages.go index 126a99bad..c7312cd4a 100644 --- a/providers/aws/apigatewayv2/stages.go +++ b/providers/aws/apigatewayv2/stages.go @@ -183,11 +183,20 @@ func checkStageUpdate(ad *apiData, st *driver.Stage, in *driver.UpdateStageInput return badRequest("Description must be at most %d characters", maxDescriptionLen) } - if in.DeploymentID != nil { - return checkDeploymentID(ad, *in.DeploymentID) + if in.DeploymentID == nil || *in.DeploymentID == "" { + return nil } - return nil + autoDeploy := st.AutoDeploy + if in.AutoDeploy != nil { + autoDeploy = *in.AutoDeploy + } + + if autoDeploy { + return badRequest("DeploymentId can't be updated if autoDeploy is enabled") + } + + return checkDeploymentID(ad, *in.DeploymentID) } // checkDeploymentID rejects a stage deploymentId that names no deployment of diff --git a/providers/aws/apigatewayv2/validation.go b/providers/aws/apigatewayv2/validation.go index 1ca104d3b..5c859e6d8 100644 --- a/providers/aws/apigatewayv2/validation.go +++ b/providers/aws/apigatewayv2/validation.go @@ -13,6 +13,9 @@ const ( defaultKey = "$default" integrationsPrefix = "integrations/" usageIdentifierKeyExp = "$context.authorizer.usageIdentifierKey" + httpRouteSelectionAlt = "${request.method} ${request.path}" + apiKeyHeaderAlt = "${request.header.x-api-key}" + usageIdentifierAlt = "${context.authorizer.usageIdentifierKey}" payloadFormatV2 = "2.0" ) @@ -27,13 +30,13 @@ const ( const ( maxTags = 50 maxTagKeyLen = 128 - maxTagValueLen = 1600 + maxTagValueLen = 256 reservedTagPfx = "aws:" ) // stageNameRe is the character set a stage name may use; "$default" is the // only other accepted value. -var stageNameRe = regexp.MustCompile(`^[a-zA-Z0-9._-]+$`) +var stageNameRe = regexp.MustCompile(`^[a-zA-Z0-9_-]+$`) // httpRouteMethods are the methods an HTTP API route key may start with. // @@ -74,19 +77,25 @@ func validateAPIFields(a *driver.API) error { return badRequest("Description must be at most %d characters", maxDescriptionLen) } - if a.APIKeySelectionExpression != defaultAPIKeySelectionExpr && a.APIKeySelectionExpression != usageIdentifierKeyExp { + switch a.APIKeySelectionExpression { + case defaultAPIKeySelectionExpr, apiKeyHeaderAlt, usageIdentifierKeyExp, usageIdentifierAlt: + default: return badRequest("Invalid API key selection expression specified: %s", a.APIKeySelectionExpression) } + if a.ProtocolType == driver.ProtocolWebSocket && a.CorsConfiguration != nil { + return badRequest("CORS configuration is not supported for WEBSOCKET protocol") + } + return validateRouteSelection(a.ProtocolType, a.RouteSelectionExpression) } // validateRouteSelection enforces the route selection expression rules: HTTP -// APIs accept only the fixed method-and-path form, and WebSocket APIs need a -// $request expression. +// APIs accept only the method-and-path expression (with or without braces), +// and WebSocket APIs need a $request expression. func validateRouteSelection(protocol, expr string) error { if protocol == driver.ProtocolHTTP { - if expr != defaultRouteSelectionExpr { + if expr != defaultRouteSelectionExpr && expr != httpRouteSelectionAlt { return badRequest("Only %s is supported for HTTP APIs", defaultRouteSelectionExpr) } @@ -227,7 +236,7 @@ func validateStageName(name string) error { } if name != defaultKey && (len(name) > maxNameLen || !stageNameRe.MatchString(name)) { - return badRequest("Stage name only allows a-zA-Z0-9._- or $default") + return badRequest("Stage name only allows a-zA-Z0-9_- or $default") } return nil diff --git a/server/aws/apigatewayv2/deployment_e2e_test.go b/server/aws/apigatewayv2/deployment_e2e_test.go index 1ca8c141c..2f416e960 100644 --- a/server/aws/apigatewayv2/deployment_e2e_test.go +++ b/server/aws/apigatewayv2/deployment_e2e_test.go @@ -112,4 +112,13 @@ func TestE2E_AutoDeployStageRedeploysOnChange(t *testing.T) { if second, _ := stage["deploymentId"].(string); second == "" || second == first { t.Fatalf("autoDeploy stage not redeployed: first=%s now=%v", first, stage["deploymentId"]) } + + wantErr(t, http.MethodPatch, apiBase+"/stages/$default", `{"deploymentId":"`+first+`"}`, + http.StatusBadRequest, "BadRequestException", "DeploymentId can't be updated if autoDeploy is enabled") + + // Switching autoDeploy off in the same call lets the stage pin a deployment. + pinned := mustDo(t, http.MethodPatch, apiBase+"/stages/$default", `{"autoDeploy":false,"deploymentId":"`+first+`"}`, http.StatusOK) + if pinned["deploymentId"] != first || pinned["autoDeploy"] != false { + t.Fatalf("pin deployment = %v", pinned) + } } diff --git a/server/aws/apigatewayv2/integrations_handler.go b/server/aws/apigatewayv2/integrations_handler.go index 2a7d09d4a..2e244c7e8 100644 --- a/server/aws/apigatewayv2/integrations_handler.go +++ b/server/aws/apigatewayv2/integrations_handler.go @@ -35,6 +35,10 @@ func (h *Handler) createIntegration(w http.ResponseWriter, r *http.Request, apiI PayloadFormatVersion: req.PayloadFormatVersion, TimeoutInMillis: req.TimeoutInMillis, Description: req.Description, RequestParameters: req.RequestParameters, CredentialsArn: req.CredentialsArn, + + RequestTemplates: req.RequestTemplates, + TemplateSelectionExpression: req.TemplateSelectionExpression, + PassthroughBehavior: req.PassthroughBehavior, }) if err != nil { writeErr(w, err) @@ -67,6 +71,10 @@ func (h *Handler) updateIntegration(w http.ResponseWriter, r *http.Request, apiI PayloadFormatVersion: req.PayloadFormatVersion, TimeoutInMillis: req.TimeoutInMillis, Description: req.Description, RequestParameters: req.RequestParameters, CredentialsArn: req.CredentialsArn, + + RequestTemplates: req.RequestTemplates, + TemplateSelectionExpression: req.TemplateSelectionExpression, + PassthroughBehavior: req.PassthroughBehavior, }) if err != nil { writeErr(w, err) diff --git a/server/aws/apigatewayv2/pagination_e2e_test.go b/server/aws/apigatewayv2/pagination_e2e_test.go index dbe8a97b8..dc0eaba36 100644 --- a/server/aws/apigatewayv2/pagination_e2e_test.go +++ b/server/aws/apigatewayv2/pagination_e2e_test.go @@ -81,13 +81,14 @@ func TestE2E_PaginationAcrossCollections(t *testing.T) { func TestE2E_PaginationBounds(t *testing.T) { ts := newE2E(t) - for _, q := range []string{"maxResults=abc", "maxResults=0", "maxResults=501"} { + for _, q := range []string{"maxResults=abc", "maxResults=0", "maxResults=-1"} { wantErr(t, http.MethodGet, ts.URL+"/v2/apis?"+q, "", - http.StatusBadRequest, "BadRequestException", "MaxResults must be an integer between 1 and 500") + http.StatusBadRequest, "BadRequestException", "MaxResults must be a positive integer") } wantErr(t, http.MethodGet, ts.URL+"/v2/apis?nextToken=%21%21bad", "", http.StatusBadRequest, "BadRequestException", "Invalid NextToken specified") mustDo(t, http.MethodGet, ts.URL+"/v2/apis?maxResults=500", "", http.StatusOK) + mustDo(t, http.MethodGet, ts.URL+"/v2/apis?maxResults=5000", "", http.StatusOK) } diff --git a/server/aws/apigatewayv2/types.go b/server/aws/apigatewayv2/types.go index ef68e5b40..1add99342 100644 --- a/server/aws/apigatewayv2/types.go +++ b/server/aws/apigatewayv2/types.go @@ -167,6 +167,10 @@ type integrationRequest struct { Description string `json:"description"` RequestParameters map[string]string `json:"requestParameters"` CredentialsArn string `json:"credentialsArn"` + + RequestTemplates map[string]string `json:"requestTemplates"` + TemplateSelectionExpression string `json:"templateSelectionExpression"` + PassthroughBehavior string `json:"passthroughBehavior"` } // updateIntegrationRequest is the UpdateIntegration (PATCH) request body. @@ -180,6 +184,10 @@ type updateIntegrationRequest struct { Description *string `json:"description"` RequestParameters map[string]string `json:"requestParameters"` CredentialsArn *string `json:"credentialsArn"` + + RequestTemplates map[string]string `json:"requestTemplates"` + TemplateSelectionExpression *string `json:"templateSelectionExpression"` + PassthroughBehavior *string `json:"passthroughBehavior"` } // integrationResponse is the Integration wire object. @@ -195,6 +203,10 @@ type integrationResponse struct { RequestParameters map[string]string `json:"requestParameters,omitempty"` CredentialsArn string `json:"credentialsArn,omitempty"` APIGatewayManaged bool `json:"apiGatewayManaged,omitempty"` + + RequestTemplates map[string]string `json:"requestTemplates,omitempty"` + TemplateSelectionExpression string `json:"templateSelectionExpression,omitempty"` + PassthroughBehavior string `json:"passthroughBehavior,omitempty"` } func toIntegrationResponse(i *driver.Integration) integrationResponse { @@ -205,6 +217,10 @@ func toIntegrationResponse(i *driver.Integration) integrationResponse { TimeoutInMillis: i.TimeoutInMillis, Description: i.Description, RequestParameters: i.RequestParameters, CredentialsArn: i.CredentialsArn, APIGatewayManaged: i.APIGatewayManaged, + + RequestTemplates: i.RequestTemplates, + TemplateSelectionExpression: i.TemplateSelectionExpression, + PassthroughBehavior: i.PassthroughBehavior, } } diff --git a/server/aws/apigatewayv2/validation_e2e_test.go b/server/aws/apigatewayv2/validation_e2e_test.go index 9face1484..a53ab50a2 100644 --- a/server/aws/apigatewayv2/validation_e2e_test.go +++ b/server/aws/apigatewayv2/validation_e2e_test.go @@ -100,8 +100,74 @@ func TestE2E_APIAndStageValidation(t *testing.T) { apiBase := ts.URL + "/v2/apis/" + newHTTPAPI(t, ts.URL) - wantErr(t, http.MethodPost, apiBase+"/stages", `{"stageName":"bad name!"}`, - http.StatusBadRequest, "BadRequestException", "Stage name only allows a-zA-Z0-9._- or $default") + for _, name := range []string{"bad name!", "prod.v1"} { + wantErr(t, http.MethodPost, apiBase+"/stages", `{"stageName":"`+name+`"}`, + http.StatusBadRequest, "BadRequestException", "Stage name only allows a-zA-Z0-9_- or $default") + } + + mustDo(t, http.MethodPost, apiBase+"/stages", `{"stageName":"prod_v1-b"}`, http.StatusCreated) +} + +// TestE2E_SelectionExpressionForms proves both the bare and the braced forms +// of the HTTP route and API key selection expressions are accepted, the way +// CDK sends them. +func TestE2E_SelectionExpressionForms(t *testing.T) { + ts := newE2E(t) + + api := mustDo(t, http.MethodPost, ts.URL+"/v2/apis", `{"name":"cdk","protocolType":"HTTP",`+ + `"routeSelectionExpression":"${request.method} ${request.path}",`+ + `"apiKeySelectionExpression":"${request.header.x-api-key}"}`, http.StatusCreated) + + if api["routeSelectionExpression"] != "${request.method} ${request.path}" || + api["apiKeySelectionExpression"] != "${request.header.x-api-key}" { + t.Fatalf("braced expressions not stored as sent: %v", api) + } + + mustDo(t, http.MethodPost, ts.URL+"/v2/apis", `{"name":"ws","protocolType":"WEBSOCKET",`+ + `"routeSelectionExpression":"$request.body.action",`+ + `"apiKeySelectionExpression":"${context.authorizer.usageIdentifierKey}"}`, http.StatusCreated) + + wantErr(t, http.MethodPost, ts.URL+"/v2/apis", `{"name":"h","protocolType":"HTTP","apiKeySelectionExpression":"$request.header.foo"}`, + http.StatusBadRequest, "BadRequestException", "Invalid API key selection expression specified: $request.header.foo") +} + +// TestE2E_WebSocketRejectsCORS covers CORS on create and update of a WebSocket API. +func TestE2E_WebSocketRejectsCORS(t *testing.T) { + ts := newE2E(t) + const msg = "CORS configuration is not supported for WEBSOCKET protocol" + + wantErr(t, http.MethodPost, ts.URL+"/v2/apis", `{"name":"ws","protocolType":"WEBSOCKET",`+ + `"routeSelectionExpression":"$request.body.action","corsConfiguration":{"allowOrigins":["*"]}}`, + http.StatusBadRequest, "BadRequestException", msg) - mustDo(t, http.MethodPost, apiBase+"/stages", `{"stageName":"prod.v1_a-b"}`, http.StatusCreated) + wsID := newAPI(t, ts.URL, `{"name":"ws","protocolType":"WEBSOCKET","routeSelectionExpression":"$request.body.action"}`) + + wantErr(t, http.MethodPatch, ts.URL+"/v2/apis/"+wsID, `{"corsConfiguration":{"allowOrigins":["*"]}}`, + http.StatusBadRequest, "BadRequestException", msg) +} + +// TestE2E_IntegrationTemplatesRoundTrip proves request templates, the template +// selection expression and passthrough behavior are stored and echoed. +func TestE2E_IntegrationTemplatesRoundTrip(t *testing.T) { + ts := newE2E(t) + wsBase := ts.URL + "/v2/apis/" + newAPI(t, ts.URL, + `{"name":"ws","protocolType":"WEBSOCKET","routeSelectionExpression":"$request.body.action"}`) + + ig := mustDo(t, http.MethodPost, wsBase+"/integrations", `{"integrationType":"MOCK",`+ + `"requestTemplates":{"200":"{\"statusCode\":200}"},"templateSelectionExpression":"200",`+ + `"passthroughBehavior":"WHEN_NO_MATCH"}`, http.StatusCreated) + + igID, _ := ig["integrationId"].(string) + + got := mustDo(t, http.MethodGet, wsBase+"/integrations/"+igID, "", http.StatusOK) + tpl, _ := got["requestTemplates"].(map[string]any) + + if tpl["200"] != `{"statusCode":200}` || got["templateSelectionExpression"] != "200" || got["passthroughBehavior"] != "WHEN_NO_MATCH" { + t.Fatalf("GetIntegration templates = %v", got) + } + + upd := mustDo(t, http.MethodPatch, wsBase+"/integrations/"+igID, `{"requestTemplates":{"201":"x"}}`, http.StatusOK) + if tpl, _ := upd["requestTemplates"].(map[string]any); tpl["201"] != "x" || len(tpl) != 1 { + t.Fatalf("UpdateIntegration templates = %v", upd) + } } diff --git a/services/apigatewayv2/driver/driver.go b/services/apigatewayv2/driver/driver.go index c256a83e7..24a5a2e5f 100644 --- a/services/apigatewayv2/driver/driver.go +++ b/services/apigatewayv2/driver/driver.go @@ -93,6 +93,10 @@ type Integration struct { RequestParameters map[string]string CredentialsArn string APIGatewayManaged bool + + RequestTemplates map[string]string + TemplateSelectionExpression string + PassthroughBehavior string } // Stage is a named deployment stage of an API (e.g. "$default", "prod"). @@ -213,6 +217,10 @@ type CreateIntegrationInput struct { Description string RequestParameters map[string]string CredentialsArn string + + RequestTemplates map[string]string + TemplateSelectionExpression string + PassthroughBehavior string } // UpdateIntegrationInput carries the mutable fields UpdateIntegration accepts. @@ -226,6 +234,10 @@ type UpdateIntegrationInput struct { Description *string RequestParameters map[string]string CredentialsArn *string + + RequestTemplates map[string]string + TemplateSelectionExpression *string + PassthroughBehavior *string } // CreateStageInput carries the fields CreateStage accepts.