package test import ( "context" "fmt" "queryorchestration/internal/server/queue" "queryorchestration/internal/serviceconfig" "testing" "time" "github.com/aws/aws-sdk-go-v2/aws" "github.com/aws/aws-sdk-go-v2/config" "github.com/aws/aws-sdk-go-v2/credentials" "github.com/aws/aws-sdk-go-v2/service/sqs" "github.com/aws/aws-sdk-go-v2/service/sqs/types" "github.com/docker/go-connections/nat" "github.com/stretchr/testify/assert" "github.com/testcontainers/testcontainers-go" "github.com/testcontainers/testcontainers-go/wait" ) type queueContainerConfig struct { Container testcontainers.Container ExternalEndpoint string NetworkEndpoint string } type CreateQueueConfig struct { Network *testcontainers.DockerNetwork Cfg *serviceconfig.BaseConfig } func createQueueContainer(t *testing.T, ctx context.Context, cfg *CreateQueueConfig) (*queueContainerConfig, func()) { alias := "localstack" provider := credentials.NewStaticCredentialsProvider(cfg.Cfg.AWSKeyID, cfg.Cfg.AWSSecretKey, cfg.Cfg.AWSSessionToken) port, err := nat.NewPort("tcp", "4566") if err != nil { t.Fatalf("Failed to create port: %v", err) } req := testcontainers.ContainerRequest{ Image: "localstack/localstack:4.0.3", Env: map[string]string{ "AWS_ACCESS_KEY_ID": provider.Value.AccessKeyID, "AWS_SECRET_ACCESS_KEY": provider.Value.SecretAccessKey, "AWS_SESSION_TOKEN": provider.Value.SessionToken, "AWS_REGION": cfg.Cfg.AWSRegion, "SERVICES": "sqs", "SKIP_SSL_CERT_DOWNLOAD": "1", "LOCALSTACK_HOST": alias, "SQS_ENDPOINT_STRATEGY": "path", }, ExposedPorts: []string{port.Port()}, WaitingFor: wait.ForAll( // wait.ForExposedPort(), wait.ForListeningPort(port), wait.ForLog("Ready."), ), } if cfg.Network != nil { req.Networks = []string{cfg.Network.Name} req.NetworkAliases = map[string][]string{ cfg.Network.Name: {alias}, } } container, err := testcontainers.GenericContainer(ctx, testcontainers.GenericContainerRequest{ ContainerRequest: req, Started: true, }) if err != nil { t.Fatalf("Failed to start container: %v", err) } host, err := container.Host(ctx) if err != nil { t.Fatalf("Failed to extract host: %v", err) } mappedPort, err := container.MappedPort(ctx, port) if err != nil { t.Fatalf("Failed to extract port: %v", err) } extEndpoint := fmt.Sprintf("http://%s:%s", host, mappedPort.Port()) t.Setenv("AWS_ENDPOINT_URL_SQS", extEndpoint) cfg.Cfg.AWSEndpointUrlSQS = extEndpoint var endpoint string if cfg.Network != nil { endpoint = fmt.Sprintf("http://%s:%s", alias, port.Port()) cfg.Cfg.AWSEndpointUrlSQS = endpoint } return &queueContainerConfig{ Container: container, ExternalEndpoint: extEndpoint, NetworkEndpoint: endpoint, }, func() { err := container.Terminate(ctx) if err != nil { t.Error(err) } } } func CreateQueueClient(t *testing.T, ctx context.Context, cfg *CreateQueueConfig) func() { qcfg, clean := createQueueContainer(t, ctx, cfg) provider := credentials.NewStaticCredentialsProvider(cfg.Cfg.AWSKeyID, cfg.Cfg.AWSSecretKey, cfg.Cfg.AWSSessionToken) sqsCfg, err := config.LoadDefaultConfig(ctx, config.WithRegion(cfg.Cfg.AWSRegion), config.WithCredentialsProvider(provider), config.WithBaseEndpoint(qcfg.ExternalEndpoint)) if err != nil { t.Fatal(err) } cfg.Cfg.QueueClient = sqs.NewFromConfig(sqsCfg) return clean } func CreateQueue(t *testing.T, ctx context.Context, cfg *serviceconfig.BaseConfig, name string) string { queueM, err := cfg.QueueClient.CreateQueue(ctx, &sqs.CreateQueueInput{ QueueName: aws.String(name), }) if err != nil { t.Fatal(err) } return *queueM.QueueUrl } func AssertMessageWait(t *testing.T, ctx context.Context, cfg *queue.Config, attrs []string) types.Message { time.Sleep(1 * time.Second) result, err := queue.Receive(ctx, cfg, attrs) assert.NoError(t, err) assert.NotNil(t, result.Messages) assert.Len(t, result.Messages, 1) assert.NotNil(t, result.Messages[0]) return result.Messages[0] } func AssertMessageBodyWait(t *testing.T, ctx context.Context, cfg *queue.Config, body string) { message := AssertMessageWait(t, ctx, cfg, []string{}) assert.Equal(t, body, *message.Body) } func AssertMessageAttrWait(t *testing.T, ctx context.Context, cfg *queue.Config, name string, value string) { message := AssertMessageWait(t, ctx, cfg, []string{name}) assert.NotNil(t, message.MessageAttributes[name]) assert.Equal(t, value, *(message.MessageAttributes[name]).StringValue) }