From b864e3c6721b6b88d8b358b8ab940415caca19d6 Mon Sep 17 00:00:00 2001 From: Nitin Goyal Date: Tue, 29 Sep 2026 19:02:24 +0530 Subject: [PATCH] fix: implement SQS partial batch failure handling (#3) - Return events.SQSEventResponse with BatchItemFailures to avoid reprocessing succeeded records - Extract processMessage and processSQSEvent for granular per-message failure tracking - Configure FunctionResponseTypes: ['ReportBatchItemFailures'] in CloudFormation EventSourceMapping - Introduce S3Client and KMSClient interfaces for mockability - Fix conflicting dict,plain parquet tags on DatabaseActivityEvent - Add unit tests covering partial batch failure, all success, all fail, and error cases --- deploy/cloudformation/lambda.yaml | 2 + main.go | 244 +++++++++------- main_test.go | 467 ++++++++++++++++++++++++++++++ 3 files changed, 607 insertions(+), 106 deletions(-) create mode 100644 main_test.go diff --git a/deploy/cloudformation/lambda.yaml b/deploy/cloudformation/lambda.yaml index 01b1a9d..93ae874 100644 --- a/deploy/cloudformation/lambda.yaml +++ b/deploy/cloudformation/lambda.yaml @@ -140,6 +140,8 @@ Resources: FunctionName: !Ref ProcessorFunction BatchSize: 10 Enabled: true + FunctionResponseTypes: + - ReportBatchItemFailures Outputs: DestinationBucketName: diff --git a/main.go b/main.go index d3c52a4..af6d70b 100644 --- a/main.go +++ b/main.go @@ -49,9 +49,20 @@ type DASPayload struct { Key string `json:"key"` } +// S3Client defines the interface for S3 operations used by the processor. +type S3Client interface { + GetObject(ctx context.Context, in *s3.GetObjectInput, opt ...func(*s3.Options)) (*s3.GetObjectOutput, error) + PutObject(ctx context.Context, in *s3.PutObjectInput, opt ...func(*s3.Options)) (*s3.PutObjectOutput, error) +} + +// KMSClient defines the interface for KMS operations used by the processor. +type KMSClient interface { + Decrypt(ctx context.Context, in *kms.DecryptInput, opt ...func(*kms.Options)) (*kms.DecryptOutput, error) +} + var ( - s3Client *s3.Client - kmsClient *kms.Client + s3Client S3Client + kmsClient KMSClient ) func init() { @@ -65,7 +76,7 @@ func init() { } // handler is the main entry point for the Lambda function. -func handler(ctx context.Context, sqsEvent events.SQSEvent) error { +func handler(ctx context.Context, sqsEvent events.SQSEvent) (events.SQSEventResponse, error) { filterName := os.Getenv("DAS_FILTER_NAME") if filterName == "" { filterName = "default" @@ -73,115 +84,136 @@ func handler(ctx context.Context, sqsEvent events.SQSEvent) error { filterConfig, err := LoadFilterConfig(filterName) if err != nil { - return fmt.Errorf("failed to load filter config: %w", err) + slog.Error("Failed to load filter config", "error", err) + return events.SQSEventResponse{}, fmt.Errorf("failed to load filter config: %w", err) } rdsResourceID := os.Getenv("DAS_RDS_RESOURCE_ID") + return processSQSEvent(ctx, sqsEvent, func(ctx context.Context, msg events.SQSMessage) error { + return processMessage(ctx, msg, filterConfig, filterName, rdsResourceID) + }) +} + +// processSQSEvent processes each SQS message using the provided processor function, +// collecting failed message IDs into an SQSEventResponse without aborting the batch. +func processSQSEvent(ctx context.Context, sqsEvent events.SQSEvent, processFn func(context.Context, events.SQSMessage) error) (events.SQSEventResponse, error) { + var response events.SQSEventResponse + for _, message := range sqsEvent.Records { slog.Info("Processing SQS message", "MessageId", message.MessageId) - var s3Event S3EventNotification - - // Check if it's an SNS wrapped message - var snsMsg struct { - Type string `json:"Type"` - Message string `json:"Message"` - } - if err := json.Unmarshal([]byte(message.Body), &snsMsg); err == nil && snsMsg.Type == "Notification" && snsMsg.Message != "" { - // It's an SNS message - if err := json.Unmarshal([]byte(snsMsg.Message), &s3Event); err != nil { - slog.Error("Failed to parse S3 event from SNS message", "error", err) - continue - } - } else { - // Try parsing as direct S3 event - if err := json.Unmarshal([]byte(message.Body), &s3Event); err != nil { - slog.Error("Failed to parse S3 event", "error", err, "body", message.Body) - continue - } + if err := processFn(ctx, message); err != nil { + slog.Error("Failed to process SQS message", "MessageId", message.MessageId, "error", err) + response.BatchItemFailures = append(response.BatchItemFailures, events.SQSBatchItemFailure{ + ItemIdentifier: message.MessageId, + }) } + } - for _, record := range s3Event.Records { - bucket := record.S3.Bucket.Name - key := record.S3.Object.Key + return response, nil +} - slog.Info("Fetching S3 object", "Bucket", bucket, "Key", key) +// processMessage handles the processing of an individual SQS message. +func processMessage(ctx context.Context, message events.SQSMessage, filterConfig *FilterConfig, filterName, rdsResourceID string) error { + var s3Event S3EventNotification - // 1. Fetch object from S3 - getObjectOutput, err := s3Client.GetObject(ctx, &s3.GetObjectInput{ - Bucket: aws.String(bucket), - Key: aws.String(key), - }) - if err != nil { - return fmt.Errorf("failed to fetch object %s/%s: %w", bucket, key, err) - } + // Check if it's an SNS wrapped message + var snsMsg struct { + Type string `json:"Type"` + Message string `json:"Message"` + } + if err := json.Unmarshal([]byte(message.Body), &snsMsg); err == nil && snsMsg.Type == "Notification" && snsMsg.Message != "" { + // It's an SNS message + if err := json.Unmarshal([]byte(snsMsg.Message), &s3Event); err != nil { + return fmt.Errorf("failed to parse S3 event from SNS message: %w", err) + } + } else { + // Try parsing as direct S3 event + if err := json.Unmarshal([]byte(message.Body), &s3Event); err != nil { + return fmt.Errorf("failed to parse S3 event: %w", err) + } + } - bodyBytes, err := io.ReadAll(getObjectOutput.Body) - getObjectOutput.Body.Close() - if err != nil { - return fmt.Errorf("failed to read S3 object body: %w", err) - } + for _, record := range s3Event.Records { + bucket := record.S3.Bucket.Name + key := record.S3.Object.Key - // 2. Parse the DAS Payload - var payload DASPayload - if err := json.Unmarshal(bodyBytes, &payload); err != nil { - return fmt.Errorf("failed to parse DAS payload: %w", err) - } + slog.Info("Fetching S3 object", "Bucket", bucket, "Key", key) - // 3. Decrypt the KMS Data Key - decodedKmsKey, err := base64.StdEncoding.DecodeString(payload.Key) - if err != nil { - return fmt.Errorf("failed to base64 decode KMS key: %w", err) - } + // 1. Fetch object from S3 + getObjectOutput, err := s3Client.GetObject(ctx, &s3.GetObjectInput{ + Bucket: aws.String(bucket), + Key: aws.String(key), + }) + if err != nil { + return fmt.Errorf("failed to fetch object %s/%s: %w", bucket, key, err) + } - decryptOutput, err := kmsClient.Decrypt(ctx, &kms.DecryptInput{ - CiphertextBlob: decodedKmsKey, - EncryptionContext: map[string]string{ - "aws:rds:dbc-id": rdsResourceID, - }, - }) - if err != nil { - return fmt.Errorf("failed to decrypt KMS data key: %w", err) - } + bodyBytes, err := io.ReadAll(getObjectOutput.Body) + getObjectOutput.Body.Close() + if err != nil { + return fmt.Errorf("failed to read S3 object body: %w", err) + } - plaintextDataKey := decryptOutput.Plaintext + // 2. Parse the DAS Payload + var payload DASPayload + if err := json.Unmarshal(bodyBytes, &payload); err != nil { + return fmt.Errorf("failed to parse DAS payload: %w", err) + } - // 4. Decrypt the database activity events payload - decodedPayload, err := base64.StdEncoding.DecodeString(payload.DatabaseActivityEvents) - if err != nil { - return fmt.Errorf("failed to base64 decode events payload: %w", err) - } + // 3. Decrypt the KMS Data Key + decodedKmsKey, err := base64.StdEncoding.DecodeString(payload.Key) + if err != nil { + return fmt.Errorf("failed to base64 decode KMS key: %w", err) + } - decryptedPayload, err := decryptAWSEncryptionSDKPayload(ctx, decodedPayload, plaintextDataKey) - if err != nil { - return fmt.Errorf("failed to decrypt events payload: %w", err) - } + decryptOutput, err := kmsClient.Decrypt(ctx, &kms.DecryptInput{ + CiphertextBlob: decodedKmsKey, + EncryptionContext: map[string]string{ + "aws:rds:dbc-id": rdsResourceID, + }, + }) + if err != nil { + return fmt.Errorf("failed to decrypt KMS data key: %w", err) + } - // 5. Decompress the payload - decompressedPayload, err := decompressZlib(decryptedPayload) - if err != nil { - return fmt.Errorf("failed to decompress events payload: %w", err) - } + plaintextDataKey := decryptOutput.Plaintext - // 6. Convert to Parquet - parquetBytes, err := convertJSONToParquet(decompressedPayload, filterConfig) - if err != nil { - return fmt.Errorf("failed to convert JSON to Parquet: %w", err) - } + // 4. Decrypt the database activity events payload + decodedPayload, err := base64.StdEncoding.DecodeString(payload.DatabaseActivityEvents) + if err != nil { + return fmt.Errorf("failed to base64 decode events payload: %w", err) + } - // 7. Write to destination S3 (Fanout) - destKey := fmt.Sprintf("das/%s/%s-processed.parquet", filterName, key) - slog.Info("Writing processed data back to S3", "DestBucket", bucket, "DestKey", destKey) + decryptedPayload, err := decryptAWSEncryptionSDKPayload(ctx, decodedPayload, plaintextDataKey) + if err != nil { + return fmt.Errorf("failed to decrypt events payload: %w", err) + } - _, err = s3Client.PutObject(ctx, &s3.PutObjectInput{ - Bucket: aws.String(bucket), - Key: aws.String(destKey), - Body: bytes.NewReader(parquetBytes), - }) - if err != nil { - return fmt.Errorf("failed to write processed data: %w", err) - } + // 5. Decompress the payload + decompressedPayload, err := decompressZlib(decryptedPayload) + if err != nil { + return fmt.Errorf("failed to decompress events payload: %w", err) + } + + // 6. Convert to Parquet + parquetBytes, err := convertJSONToParquet(decompressedPayload, filterConfig) + if err != nil { + return fmt.Errorf("failed to convert JSON to Parquet: %w", err) + } + + // 7. Write to destination S3 (Fanout) + destKey := fmt.Sprintf("das/%s/%s-processed.parquet", filterName, key) + slog.Info("Writing processed data back to S3", "DestBucket", bucket, "DestKey", destKey) + + _, err = s3Client.PutObject(ctx, &s3.PutObjectInput{ + Bucket: aws.String(bucket), + Key: aws.String(destKey), + Body: bytes.NewReader(parquetBytes), + }) + if err != nil { + return fmt.Errorf("failed to write processed data: %w", err) } } @@ -240,27 +272,27 @@ func decompressZlib(data []byte) ([]byte, error) { // DatabaseActivityEvent represents a single DAS event type DatabaseActivityEvent struct { - LogTime string `parquet:"logTime,dict,plain" json:"logTime"` + LogTime string `parquet:"logTime,dict" json:"logTime"` StatementId int64 `parquet:"statementId" json:"statementId"` SubstatementId int64 `parquet:"substatementId" json:"substatementId"` - ObjectType string `parquet:"objectType,dict,plain" json:"objectType"` - Command string `parquet:"command,dict,plain" json:"command"` - ObjectName string `parquet:"objectName,dict,plain" json:"objectName"` - DatabaseName string `parquet:"databaseName,dict,plain" json:"databaseName"` - DbUserName string `parquet:"dbUserName,dict,plain" json:"dbUserName"` - RemoteHost string `parquet:"remoteHost,dict,plain" json:"remoteHost"` - SessionId string `parquet:"sessionId,dict,plain" json:"sessionId"` + ObjectType string `parquet:"objectType,dict" json:"objectType"` + Command string `parquet:"command,dict" json:"command"` + ObjectName string `parquet:"objectName,dict" json:"objectName"` + DatabaseName string `parquet:"databaseName,dict" json:"databaseName"` + DbUserName string `parquet:"dbUserName,dict" json:"dbUserName"` + RemoteHost string `parquet:"remoteHost,dict" json:"remoteHost"` + SessionId string `parquet:"sessionId,dict" json:"sessionId"` RowCount int64 `parquet:"rowCount" json:"rowCount"` - CommandText string `parquet:"commandText,dict,plain" json:"commandText"` + CommandText string `parquet:"commandText,dict" json:"commandText"` ParamList []string `parquet:"paramList,list" json:"paramList"` Pid int64 `parquet:"pid" json:"pid"` - ClientApplication string `parquet:"clientApplication,dict,plain" json:"clientApplication"` - ExitCode string `parquet:"exitCode,dict,plain" json:"exitCode"` - Class string `parquet:"class,dict,plain" json:"class"` - ServerHost string `parquet:"serverHost,dict,plain" json:"serverHost"` - Type string `parquet:"type,dict,plain" json:"type"` - StartTime string `parquet:"startTime,dict,plain" json:"startTime"` - ErrorMessage string `parquet:"errorMessage,dict,plain" json:"errorMessage"` + ClientApplication string `parquet:"clientApplication,dict" json:"clientApplication"` + ExitCode string `parquet:"exitCode,dict" json:"exitCode"` + Class string `parquet:"class,dict" json:"class"` + ServerHost string `parquet:"serverHost,dict" json:"serverHost"` + Type string `parquet:"type,dict" json:"type"` + StartTime string `parquet:"startTime,dict" json:"startTime"` + ErrorMessage string `parquet:"errorMessage,dict" json:"errorMessage"` } func convertJSONToParquet(decompressedJSON []byte, filterConfig *FilterConfig) ([]byte, error) { diff --git a/main_test.go b/main_test.go new file mode 100644 index 0000000..eefaf53 --- /dev/null +++ b/main_test.go @@ -0,0 +1,467 @@ +package main + +import ( + "bytes" + "compress/zlib" + "context" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "io" + "testing" + + mpl "github.com/aws/aws-cryptographic-material-providers-library/releases/go/mpl/awscryptographymaterialproviderssmithygenerated" + mpltypes "github.com/aws/aws-cryptographic-material-providers-library/releases/go/mpl/awscryptographymaterialproviderssmithygeneratedtypes" + client "github.com/aws/aws-encryption-sdk/releases/go/encryption-sdk/awscryptographyencryptionsdksmithygenerated" + esdktypes "github.com/aws/aws-encryption-sdk/releases/go/encryption-sdk/awscryptographyencryptionsdksmithygeneratedtypes" + "github.com/aws/aws-lambda-go/events" + "github.com/aws/aws-sdk-go-v2/service/kms" + "github.com/aws/aws-sdk-go-v2/service/s3" +) + +// mockS3Client implements S3Client for unit testing. +type mockS3Client struct { + getObjectFunc func(ctx context.Context, in *s3.GetObjectInput, opt ...func(*s3.Options)) (*s3.GetObjectOutput, error) + putObjectFunc func(ctx context.Context, in *s3.PutObjectInput, opt ...func(*s3.Options)) (*s3.PutObjectOutput, error) +} + +func (m *mockS3Client) GetObject(ctx context.Context, in *s3.GetObjectInput, opt ...func(*s3.Options)) (*s3.GetObjectOutput, error) { + if m.getObjectFunc != nil { + return m.getObjectFunc(ctx, in, opt...) + } + return &s3.GetObjectOutput{ + Body: io.NopCloser(bytes.NewReader([]byte("{}"))), + }, nil +} + +func (m *mockS3Client) PutObject(ctx context.Context, in *s3.PutObjectInput, opt ...func(*s3.Options)) (*s3.PutObjectOutput, error) { + if m.putObjectFunc != nil { + return m.putObjectFunc(ctx, in, opt...) + } + return &s3.PutObjectOutput{}, nil +} + +// mockKMSClient implements KMSClient for unit testing. +type mockKMSClient struct { + decryptFunc func(ctx context.Context, in *kms.DecryptInput, opt ...func(*kms.Options)) (*kms.DecryptOutput, error) +} + +func (m *mockKMSClient) Decrypt(ctx context.Context, in *kms.DecryptInput, opt ...func(*kms.Options)) (*kms.DecryptOutput, error) { + if m.decryptFunc != nil { + return m.decryptFunc(ctx, in, opt...) + } + return &kms.DecryptOutput{ + Plaintext: []byte("01234567890123456789012345678901"), + }, nil +} + +// TestProcessSQSEvent_PartialBatchFailure verifies that when a batch contains both +// succeeding and failing records, only the failing record's ID is included in BatchItemFailures. +func TestProcessSQSEvent_PartialBatchFailure(t *testing.T) { + ctx := context.Background() + + sqsEvent := events.SQSEvent{ + Records: []events.SQSMessage{ + { + MessageId: "msg-success-1", + Body: `{"Records":[]}`, + }, + { + MessageId: "msg-failure-2", + Body: `invalid-payload`, + }, + }, + } + + response, err := processSQSEvent(ctx, sqsEvent, func(ctx context.Context, msg events.SQSMessage) error { + if msg.MessageId == "msg-failure-2" { + return errors.New("simulated processing failure") + } + return nil + }) + + if err != nil { + t.Fatalf("expected nil error from processSQSEvent, got: %v", err) + } + + if len(response.BatchItemFailures) != 1 { + t.Fatalf("expected exactly 1 failure, got %d", len(response.BatchItemFailures)) + } + + if response.BatchItemFailures[0].ItemIdentifier != "msg-failure-2" { + t.Errorf("expected failed itemIdentifier to be 'msg-failure-2', got: %s", response.BatchItemFailures[0].ItemIdentifier) + } +} + +// TestProcessSQSEvent_AllSuccess verifies that when all records succeed, BatchItemFailures is empty. +func TestProcessSQSEvent_AllSuccess(t *testing.T) { + ctx := context.Background() + + sqsEvent := events.SQSEvent{ + Records: []events.SQSMessage{ + {MessageId: "msg-1"}, + {MessageId: "msg-2"}, + }, + } + + response, err := processSQSEvent(ctx, sqsEvent, func(ctx context.Context, msg events.SQSMessage) error { + return nil + }) + + if err != nil { + t.Fatalf("expected nil error, got: %v", err) + } + + if len(response.BatchItemFailures) != 0 { + t.Errorf("expected 0 failures, got %d", len(response.BatchItemFailures)) + } +} + +// TestProcessSQSEvent_AllFail verifies that when all records fail, all message IDs are returned. +func TestProcessSQSEvent_AllFail(t *testing.T) { + ctx := context.Background() + + sqsEvent := events.SQSEvent{ + Records: []events.SQSMessage{ + {MessageId: "msg-1"}, + {MessageId: "msg-2"}, + }, + } + + response, err := processSQSEvent(ctx, sqsEvent, func(ctx context.Context, msg events.SQSMessage) error { + return fmt.Errorf("failed: %s", msg.MessageId) + }) + + if err != nil { + t.Fatalf("expected nil error, got: %v", err) + } + + if len(response.BatchItemFailures) != 2 { + t.Fatalf("expected 2 failures, got %d", len(response.BatchItemFailures)) + } + + if response.BatchItemFailures[0].ItemIdentifier != "msg-1" || response.BatchItemFailures[1].ItemIdentifier != "msg-2" { + t.Errorf("unexpected failure identifiers: %+v", response.BatchItemFailures) + } +} + +// TestProcessSQSEvent_EmptyBatch verifies handling of empty event batch. +func TestProcessSQSEvent_EmptyBatch(t *testing.T) { + ctx := context.Background() + sqsEvent := events.SQSEvent{Records: []events.SQSMessage{}} + + response, err := processSQSEvent(ctx, sqsEvent, func(ctx context.Context, msg events.SQSMessage) error { + return nil + }) + + if err != nil { + t.Fatalf("expected nil error, got: %v", err) + } + + if len(response.BatchItemFailures) != 0 { + t.Errorf("expected 0 failures, got %d", len(response.BatchItemFailures)) + } +} + +// TestProcessMessage_InvalidJSON verifies that invalid message JSON returns a clear error. +func TestProcessMessage_InvalidJSON(t *testing.T) { + ctx := context.Background() + filterCfg := &FilterConfig{Drop: []string{}, Query: map[string]interface{}{}} + + msg := events.SQSMessage{ + MessageId: "test-msg-bad-json", + Body: "{bad json", + } + + err := processMessage(ctx, msg, filterCfg, "default", "db-123") + if err == nil { + t.Fatal("expected error for invalid JSON body, got nil") + } +} + +// TestProcessMessage_InvalidSNSEnvelope verifies that malformed inner SNS message returns error. +func TestProcessMessage_InvalidSNSEnvelope(t *testing.T) { + ctx := context.Background() + filterCfg := &FilterConfig{Drop: []string{}, Query: map[string]interface{}{}} + + msg := events.SQSMessage{ + MessageId: "test-msg-sns", + Body: `{"Type": "Notification", "Message": "malformed inner"}`, + } + + err := processMessage(ctx, msg, filterCfg, "default", "db-123") + if err == nil { + t.Fatal("expected error for invalid SNS inner JSON, got nil") + } +} + +// TestProcessMessage_S3GetObjectError verifies that S3 fetch errors are wrapped and returned. +func TestProcessMessage_S3GetObjectError(t *testing.T) { + ctx := context.Background() + oldS3 := s3Client + defer func() { s3Client = oldS3 }() + + s3Client = &mockS3Client{ + getObjectFunc: func(ctx context.Context, in *s3.GetObjectInput, opt ...func(*s3.Options)) (*s3.GetObjectOutput, error) { + return nil, errors.New("s3 connection reset") + }, + } + + s3Notification := S3EventNotification{ + Records: []struct { + S3 struct { + Bucket struct { + Name string `json:"name"` + } `json:"bucket"` + Object struct { + Key string `json:"key"` + } `json:"object"` + } `json:"s3"` + }{ + { + S3: struct { + Bucket struct { + Name string `json:"name"` + } `json:"bucket"` + Object struct { + Key string `json:"key"` + } `json:"object"` + }{ + Bucket: struct { + Name string `json:"name"` + }{Name: "test-bucket"}, + Object: struct { + Key string `json:"key"` + }{Key: "test-key"}, + }, + }, + }, + } + bodyBytes, _ := json.Marshal(s3Notification) + + msg := events.SQSMessage{ + MessageId: "msg-s3-fail", + Body: string(bodyBytes), + } + + filterCfg := &FilterConfig{Drop: []string{}, Query: map[string]interface{}{}} + err := processMessage(ctx, msg, filterCfg, "default", "db-123") + if err == nil { + t.Fatal("expected error when S3 GetObject fails, got nil") + } +} + +// TestProcessMessage_KMSDecryptError verifies that KMS decryption errors are wrapped and returned. +func TestProcessMessage_KMSDecryptError(t *testing.T) { + ctx := context.Background() + oldS3, oldKMS := s3Client, kmsClient + defer func() { + s3Client = oldS3 + kmsClient = oldKMS + }() + + dasPayload := DASPayload{ + Type: "DatabaseActivityMonitoringRecord", + Version: "1.0", + DatabaseActivityEvents: base64.StdEncoding.EncodeToString([]byte("dummy-events")), + Key: base64.StdEncoding.EncodeToString([]byte("dummy-key")), + } + dasPayloadBytes, _ := json.Marshal(dasPayload) + + s3Client = &mockS3Client{ + getObjectFunc: func(ctx context.Context, in *s3.GetObjectInput, opt ...func(*s3.Options)) (*s3.GetObjectOutput, error) { + return &s3.GetObjectOutput{ + Body: io.NopCloser(bytes.NewReader(dasPayloadBytes)), + }, nil + }, + } + + kmsClient = &mockKMSClient{ + decryptFunc: func(ctx context.Context, in *kms.DecryptInput, opt ...func(*kms.Options)) (*kms.DecryptOutput, error) { + return nil, errors.New("kms: AccessDeniedException") + }, + } + + s3Notification := S3EventNotification{ + Records: []struct { + S3 struct { + Bucket struct { + Name string `json:"name"` + } `json:"bucket"` + Object struct { + Key string `json:"key"` + } `json:"object"` + } `json:"s3"` + }{ + { + S3: struct { + Bucket struct { + Name string `json:"name"` + } `json:"bucket"` + Object struct { + Key string `json:"key"` + } `json:"object"` + }{ + Bucket: struct { + Name string `json:"name"` + }{Name: "test-bucket"}, + Object: struct { + Key string `json:"key"` + }{Key: "test-key"}, + }, + }, + }, + } + s3NotifBytes, _ := json.Marshal(s3Notification) + + msg := events.SQSMessage{ + MessageId: "msg-kms-fail", + Body: string(s3NotifBytes), + } + + filterCfg := &FilterConfig{Drop: []string{}, Query: map[string]interface{}{}} + err := processMessage(ctx, msg, filterCfg, "default", "db-123") + if err == nil { + t.Fatal("expected error when KMS Decrypt fails, got nil") + } +} + +// TestEndToEndProcessMessageSuccess creates an encrypted and compressed DAS payload and verifies +// full successful processing and Parquet output generation. +func TestEndToEndProcessMessageSuccess(t *testing.T) { + ctx := context.Background() + oldS3, oldKMS := s3Client, kmsClient + defer func() { + s3Client = oldS3 + kmsClient = oldKMS + }() + + // 1. Prepare raw JSON events + rawEventsJSON := `{"databaseActivityEventList":[{"type":"activity","dbUserName":"testuser","command":"SELECT","commandText":"SELECT 1","rowCount":1}]}` + + // 2. Compress with zlib + var zlibBuf bytes.Buffer + zw := zlib.NewWriter(&zlibBuf) + _, _ = zw.Write([]byte(rawEventsJSON)) + _ = zw.Close() + + // 3. Encrypt with AWS Encryption SDK using a 32-byte key + rawKey := []byte("01234567890123456789012345678901") + matProv, err := mpl.NewClient(mpltypes.MaterialProvidersConfig{}) + if err != nil { + t.Fatalf("failed to create material providers: %v", err) + } + aesKeyring, err := matProv.CreateRawAesKeyring(ctx, mpltypes.CreateRawAesKeyringInput{ + KeyName: "DataKey", + KeyNamespace: "RawMasterKeyProvider", + WrappingKey: rawKey, + WrappingAlg: mpltypes.AesWrappingAlgAlgAes256GcmIv12Tag16, + }) + if err != nil { + t.Fatalf("failed to create keyring: %v", err) + } + + esdkClient, err := client.NewClient(esdktypes.AwsEncryptionSdkConfig{}) + if err != nil { + t.Fatalf("failed to create esdk client: %v", err) + } + + encOutput, err := esdkClient.Encrypt(ctx, esdktypes.EncryptInput{ + Plaintext: zlibBuf.Bytes(), + Keyring: aesKeyring, + }) + if err != nil { + t.Fatalf("failed to encrypt test payload: %v", err) + } + + // 4. Create DASPayload + dasPayload := DASPayload{ + Type: "DatabaseActivityMonitoringRecord", + Version: "1.0", + DatabaseActivityEvents: base64.StdEncoding.EncodeToString(encOutput.Ciphertext), + Key: base64.StdEncoding.EncodeToString([]byte("encrypted-data-key")), + } + dasBytes, _ := json.Marshal(dasPayload) + + // 5. Mock S3 and KMS + var writtenKey string + var writtenBucket string + var writtenBody []byte + + s3Client = &mockS3Client{ + getObjectFunc: func(ctx context.Context, in *s3.GetObjectInput, opt ...func(*s3.Options)) (*s3.GetObjectOutput, error) { + return &s3.GetObjectOutput{ + Body: io.NopCloser(bytes.NewReader(dasBytes)), + }, nil + }, + putObjectFunc: func(ctx context.Context, in *s3.PutObjectInput, opt ...func(*s3.Options)) (*s3.PutObjectOutput, error) { + writtenBucket = *in.Bucket + writtenKey = *in.Key + writtenBody, _ = io.ReadAll(in.Body) + return &s3.PutObjectOutput{}, nil + }, + } + + kmsClient = &mockKMSClient{ + decryptFunc: func(ctx context.Context, in *kms.DecryptInput, opt ...func(*kms.Options)) (*kms.DecryptOutput, error) { + return &kms.DecryptOutput{Plaintext: rawKey}, nil + }, + } + + s3Notification := S3EventNotification{ + Records: []struct { + S3 struct { + Bucket struct { + Name string `json:"name"` + } `json:"bucket"` + Object struct { + Key string `json:"key"` + } `json:"object"` + } `json:"s3"` + }{ + { + S3: struct { + Bucket struct { + Name string `json:"name"` + } `json:"bucket"` + Object struct { + Key string `json:"key"` + } `json:"object"` + }{ + Bucket: struct { + Name string `json:"name"` + }{Name: "source-bucket"}, + Object: struct { + Key string `json:"key"` + }{Key: "das-log-key"}, + }, + }, + }, + } + s3NotifBytes, _ := json.Marshal(s3Notification) + + msg := events.SQSMessage{ + MessageId: "msg-success", + Body: string(s3NotifBytes), + } + + filterCfg := &FilterConfig{Drop: []string{}, Query: map[string]interface{}{}} + err = processMessage(ctx, msg, filterCfg, "default", "db-123") + if err != nil { + t.Fatalf("expected successful message processing, got error: %v", err) + } + + if writtenBucket != "source-bucket" { + t.Errorf("expected writtenBucket to be 'source-bucket', got %s", writtenBucket) + } + + expectedKey := "das/default/das-log-key-processed.parquet" + if writtenKey != expectedKey { + t.Errorf("expected writtenKey to be %s, got %s", expectedKey, writtenKey) + } + + if len(writtenBody) == 0 { + t.Errorf("expected non-empty written parquet bytes") + } +}