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
2 changes: 2 additions & 0 deletions deploy/cloudformation/lambda.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -140,6 +140,8 @@ Resources:
FunctionName: !Ref ProcessorFunction
BatchSize: 10
Enabled: true
FunctionResponseTypes:
- ReportBatchItemFailures

Outputs:
DestinationBucketName:
Expand Down
244 changes: 138 additions & 106 deletions main.go
Original file line number Diff line number Diff line change
Expand Up @@ -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() {
Expand All @@ -65,123 +76,144 @@ 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"
}

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)
}
}

Expand Down Expand Up @@ -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) {
Expand Down
Loading
Loading