Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 24 additions & 0 deletions core/cmd/mildstack/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,9 @@ func main() {
if err := registerNativeS3Routes(router, root.Services); err != nil {
return recordingHTTPServer{server: failedHTTPServer{err: err}, storage: storage, port: port, instanceID: instanceID}
}
if err := registerNativeSQSRoutes(router, root.Services); err != nil {
return recordingHTTPServer{server: failedHTTPServer{err: err}, storage: storage, port: port, instanceID: instanceID}
}
registrar := instanceRegistrar{manager: manager, storage: storage, instanceID: instanceID}
return recordingHTTPServer{server: deliveryhttp.NewServer(registrar, router, port), storage: storage, port: port, instanceID: instanceID}
}
Expand Down Expand Up @@ -178,6 +181,27 @@ func registerNativeDynamoDBRoutes(router *deliveryhttp.Router, services []orches
return nil
}

func registerNativeSQSRoutes(router *deliveryhttp.Router, services []orchestrator.Service) error {
if router == nil {
return nil
}

for _, service := range services {
if service == nil || service.Metadata().Name != "sqs" {
continue
}

sqsService, ok := service.(deliveryhttp.SQSNativeService)
if !ok {
return fmt.Errorf("sqs service does not expose the native http surface")
}
deliveryhttp.RegisterSQSNativeRoutes(router.Engine(), sqsService)
return nil
}

return nil
}

func containsPort(ports []int, port int) bool {
for _, existing := range ports {
if existing == port {
Expand Down
90 changes: 90 additions & 0 deletions core/cmd/mildstack/main_test.go
Original file line number Diff line number Diff line change
@@ -1,9 +1,11 @@
package main

import (
"bytes"
"context"
"errors"
"fmt"
"io"
"net/http"
"net/http/httptest"
"os"
Expand All @@ -17,6 +19,7 @@ import (
"github.com/aws/aws-sdk-go-v2/credentials"
"github.com/aws/aws-sdk-go-v2/service/dynamodb"
"github.com/aws/aws-sdk-go-v2/service/dynamodb/types"
sqssdk "github.com/aws/aws-sdk-go-v2/service/sqs"
"github.com/michasdev/mildstack/core/internal/application/orchestrator"
"github.com/michasdev/mildstack/core/internal/application/runtime"
"github.com/michasdev/mildstack/core/internal/composition"
Expand Down Expand Up @@ -688,6 +691,93 @@ func TestInstanceRegistrarServeSkipsDuplicateLoadedPort(t *testing.T) {
}
}

func TestRegisterNativeSQSRoutesExposesAwsCompatibleSmokeSurface(t *testing.T) {
t.Helper()

root := composition.DefaultRoot("test-instance")
manager := runtime.New(root.Services)
router := deliveryhttp.NewRouter(deliveryhttp.DefaultConfig(), manager)

if err := registerNativeSQSRoutes(router, root.Services); err != nil {
t.Fatalf("register native sqs routes: %v", err)
}

healthRecorder := httptest.NewRecorder()
healthRequest := httptest.NewRequest(http.MethodGet, "/api/v1/runtime/health", nil)
router.Engine().ServeHTTP(healthRecorder, healthRequest)
if got, want := healthRecorder.Code, http.StatusOK; got != want {
t.Fatalf("unexpected health status: got %d want %d", got, want)
}

rootRecorder := httptest.NewRecorder()
rootRequest := httptest.NewRequest(http.MethodGet, "/?Action=ListQueues&Version=2012-11-05", nil)
router.Engine().ServeHTTP(rootRecorder, rootRequest)
if got, want := rootRecorder.Code, http.StatusBadRequest; got != want {
t.Fatalf("unexpected sqs root status: got %d want %d", got, want)
}
if !strings.Contains(rootRecorder.Body.String(), "<ErrorResponse>") {
t.Fatalf("expected sqs error response xml, got %q", rootRecorder.Body.String())
}
if !strings.Contains(rootRecorder.Body.String(), "UnsupportedOperation") {
t.Fatalf("expected unsupported operation xml, got %q", rootRecorder.Body.String())
}

server := httptest.NewServer(router.Engine())
t.Cleanup(server.Close)

ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
t.Cleanup(cancel)

transport := &captureTransport{base: http.DefaultTransport}
cfg, err := config.LoadDefaultConfig(ctx,
config.WithRegion("us-east-1"),
config.WithCredentialsProvider(credentials.NewStaticCredentialsProvider("test", "test", "test")),
config.WithHTTPClient(&http.Client{Transport: transport}),
)
if err != nil {
t.Fatalf("load aws config: %v", err)
}

client := sqssdk.NewFromConfig(cfg, func(o *sqssdk.Options) {
o.BaseEndpoint = aws.String(server.URL)
})

_, err = client.ListQueues(ctx, &sqssdk.ListQueuesInput{})
if err == nil {
t.Fatal("expected list queues to return an error")
}
if !strings.Contains(string(transport.body), "<ErrorResponse>") {
t.Fatalf("expected captured sqs xml body, got %q", string(transport.body))
}
if !strings.Contains(string(transport.body), "UnsupportedOperation") && !strings.Contains(string(transport.body), "InvalidQueryParameter") {
t.Fatalf("expected captured sqs xml body to contain an SQS error code, got %q", string(transport.body))
}
}

type captureTransport struct {
base http.RoundTripper
body []byte
}

func (t *captureTransport) RoundTrip(req *http.Request) (*http.Response, error) {
resp, err := t.base.RoundTrip(req)
if err != nil {
return nil, err
}
if resp.Body == nil {
return resp, nil
}

data, readErr := io.ReadAll(resp.Body)
_ = resp.Body.Close()
if readErr != nil {
return nil, readErr
}
t.body = append([]byte(nil), data...)
resp.Body = io.NopCloser(bytes.NewReader(data))
return resp, nil
}

func newDynamoDBSmokeClient(t *testing.T, endpoint string) *dynamodb.Client {
t.Helper()

Expand Down
14 changes: 13 additions & 1 deletion core/internal/composition/default_root.go
Original file line number Diff line number Diff line change
Expand Up @@ -9,12 +9,14 @@ import (
"github.com/michasdev/mildstack/core/internal/application/runtime"
"github.com/michasdev/mildstack/core/internal/resources/dynamodb"
"github.com/michasdev/mildstack/core/internal/resources/s3"
"github.com/michasdev/mildstack/core/internal/resources/sqs"
)

type DefaultRootConfig struct {
InstanceID string
S3StorageBaseDir string
DynamoDBStorageBaseDir string
SQSStorageBaseDir string
}

func DefaultRoot(instanceID string) Root {
Expand Down Expand Up @@ -48,7 +50,17 @@ func defaultRootWithHook(hook orchestrator.StateHook, config DefaultRootConfig)
panic(fmt.Sprintf("composition: init dynamodb service: %v", err))
}

services := []orchestrator.Service{s3Service, dynamoService}
sqsService, err := sqs.NewWithStorage(sqs.StorageConfig{
BaseDir: config.SQSStorageBaseDir,
InstanceID: instanceID,
})
if err != nil {
_ = s3Service.Stop(context.Background())
_ = dynamoService.Stop(context.Background())
panic(fmt.Sprintf("composition: init sqs service: %v", err))
}

services := []orchestrator.Service{s3Service, dynamoService, sqsService}
for _, service := range services {
if err := service.AttachState(hook); err != nil {
for _, candidate := range services {
Expand Down
31 changes: 28 additions & 3 deletions core/internal/composition/default_root_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@ import (
dynamodbapp "github.com/michasdev/mildstack/core/internal/resources/dynamodb/application"
dynamodbdomain "github.com/michasdev/mildstack/core/internal/resources/dynamodb/domain"
s3domain "github.com/michasdev/mildstack/core/internal/resources/s3/domain"
sqsdomain "github.com/michasdev/mildstack/core/internal/resources/sqs/domain"
)

type stateHookStub struct {
Expand Down Expand Up @@ -36,19 +37,24 @@ func TestDefaultRootIncludesS3AndDynamoDBWithDeterministicRoutes(t *testing.T) {
InstanceID: "test-instance",
S3StorageBaseDir: baseDir,
DynamoDBStorageBaseDir: baseDir,
SQSStorageBaseDir: baseDir,
})
if got, want := len(root.Services), 2; got != want {
if got, want := len(root.Services), 3; got != want {
t.Fatalf("unexpected service count: got %d want %d", got, want)
}

first := root.Services[0]
second := root.Services[1]
third := root.Services[2]
if got, want := first.Metadata().Name, "s3"; got != want {
t.Fatalf("unexpected first service name: got %q want %q", got, want)
}
if got, want := second.Metadata().Name, "dynamodb"; got != want {
t.Fatalf("unexpected second service name: got %q want %q", got, want)
}
if got, want := third.Metadata().Name, "sqs"; got != want {
t.Fatalf("unexpected third service name: got %q want %q", got, want)
}

registrar := deliveryhttp.NewRegistrar()
for _, service := range root.Services {
Expand All @@ -58,7 +64,7 @@ func TestDefaultRootIncludesS3AndDynamoDBWithDeterministicRoutes(t *testing.T) {
}

entries := registrar.Services()
if got, want := len(entries), 2; got != want {
if got, want := len(entries), 3; got != want {
t.Fatalf("unexpected catalog size: got %d want %d", got, want)
}
if got, want := entries[0].Name, "dynamodb"; got != want {
Expand All @@ -67,6 +73,9 @@ func TestDefaultRootIncludesS3AndDynamoDBWithDeterministicRoutes(t *testing.T) {
if got, want := entries[1].Name, "s3"; got != want {
t.Fatalf("unexpected second catalog service: got %q want %q", got, want)
}
if got, want := entries[2].Name, "sqs"; got != want {
t.Fatalf("unexpected third catalog service: got %q want %q", got, want)
}

s3Entry, ok := registrar.Service("s3")
if !ok {
Expand Down Expand Up @@ -125,6 +134,9 @@ func TestDefaultRootIncludesS3AndDynamoDBWithDeterministicRoutes(t *testing.T) {
if _, ok := root.Services[1].(deliveryhttp.DynamoDBNativeService); !ok {
t.Fatal("expected dynamodb service to expose the native http surface")
}
if _, ok := root.Services[2].(deliveryhttp.SQSNativeService); !ok {
t.Fatal("expected sqs service to expose the native http surface")
}

if value, ok := hook.Get(dynamodbdomain.StateKey); !ok {
t.Fatalf("expected state for %q to be present", dynamodbdomain.StateKey)
Expand All @@ -144,6 +156,18 @@ func TestDefaultRootIncludesS3AndDynamoDBWithDeterministicRoutes(t *testing.T) {
t.Fatalf("unexpected s3 state: got %v want %v", got, want)
}

if value, ok := hook.Get(sqsdomain.StateKey); !ok {
t.Fatalf("expected state for %q to be present", sqsdomain.StateKey)
} else {
state := value.(map[string]any)
if got, want := state["service"], "sqs"; got != want {
t.Fatalf("unexpected sqs state: got %v want %v", got, want)
}
if got, want := len(state["queues"].([]any)), 0; got != want {
t.Fatalf("unexpected sqs queue count: got %d want %d", got, want)
}
}

dynamoDBPath := filepath.Join(baseDir, "instances", "test-instance", "dynamodb", "state.db")
if _, err := os.Stat(dynamoDBPath); err != nil {
t.Fatalf("expected dynamodb database to exist at %s: %v", dynamoDBPath, err)
Expand Down Expand Up @@ -186,8 +210,9 @@ func TestDefaultRootUsesInstanceScopedDynamoDBStorage(t *testing.T) {
InstanceID: "instance-a",
S3StorageBaseDir: baseDir,
DynamoDBStorageBaseDir: baseDir,
SQSStorageBaseDir: baseDir,
})
if got, want := len(root.Services), 2; got != want {
if got, want := len(root.Services), 3; got != want {
t.Fatalf("unexpected service count: got %d want %d", got, want)
}

Expand Down
104 changes: 104 additions & 0 deletions core/internal/delivery/http/sqs_native.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,104 @@
package http

import (
"errors"
"net/http"
"strings"

"github.com/gin-gonic/gin"
"github.com/michasdev/mildstack/core/internal/application/orchestrator"
)

type SQSNativeService interface {
Policy() orchestrator.EmulationPolicy
Metadata() orchestrator.Metadata
}

func RegisterSQSNativeRoutes(engine *gin.Engine, service SQSNativeService) {
if engine == nil || service == nil {
return
}

handler := newSQSNativeHandler(service)
engine.Use(func(c *gin.Context) {
if handled := handler.dispatch(c); handled {
c.Abort()
return
}
c.Next()
})
}

type sqsNativeHandler struct {
service SQSNativeService
registry SQSRegistry
supported map[string]struct{}
}

func newSQSNativeHandler(service SQSNativeService) sqsNativeHandler {
supported := make(map[string]struct{})
if service != nil {
for _, action := range service.Policy().Supported {
supported[action] = struct{}{}
}
}

return sqsNativeHandler{
service: service,
registry: NewSQSRegistry(),
supported: supported,
}
}

func (h sqsNativeHandler) dispatch(c *gin.Context) bool {
if c == nil || c.Request == nil || c.Request.URL == nil {
return false
}

path := strings.TrimSpace(c.Request.URL.Path)
if path == "" || strings.HasPrefix(path, "/api/") {
return false
}
switch c.Request.Method {
case http.MethodGet, http.MethodPost:
default:
return false
}

ctx, err := ParseSQSRequest(c.Request)
if err != nil {
if errors.Is(err, ErrSQSNotOwned) {
return false
}
writeSQSError(c, err, requestIDFromContext(c))
return true
}

spec, err := h.registry.Resolve(ctx)
if err != nil {
writeSQSError(c, err, requestIDFromContext(c))
return true
}

if _, ok := h.supported[spec.Action]; !ok || spec.DomainDeferred {
writeSQSError(c, ErrSQSUnsupported, requestIDFromContext(c))
return true
}

writeSQSError(c, ErrSQSUnsupported, requestIDFromContext(c))
return true
}

func requestIDFromContext(c *gin.Context) string {
if c == nil {
return "mildstack-sqs-request"
}

for _, key := range []string{"x-amzn-requestid", "X-Amzn-RequestId", "x-amz-request-id"} {
if requestID := strings.TrimSpace(c.GetHeader(key)); requestID != "" {
return requestID
}
}

return "mildstack-sqs-request"
}
Loading