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
175 changes: 175 additions & 0 deletions convert_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -369,3 +369,178 @@ func TestFixtures_Unmarshal(t *testing.T) {
t.Errorf("expected non-empty parquet schema header for empty fixture")
}
}

func TestMapToDatabaseActivityEvent_AllFields(t *testing.T) {
input := map[string]interface{}{
"logTime": "2026-09-29 12:00:00.000",
"statementId": float64(1001),
"substatementId": float64(2),
"objectType": "TABLE",
"command": "SELECT",
"objectName": "users",
"databaseName": "prod_db",
"dbUserName": "app_user",
"remoteHost": "10.0.1.45",
"sessionId": "sess-998877",
"rowCount": float64(42),
"commandText": "SELECT 1",
"paramList": []interface{}{"val1", "val2", 123},
"pid": float64(9876),
"clientApplication": "service-app",
"exitCode": 0,
"class": "READ",
"serverHost": "ip-10-0-2-10",
"type": "activity",
"startTime": "2026-09-29 12:00:00.000",
"errorMessage": "none",
}

event := mapToDatabaseActivityEvent(input)

if event.LogTime != "2026-09-29 12:00:00.000" {
t.Errorf("unexpected LogTime: %s", event.LogTime)
}
if event.StatementId != 1001 {
t.Errorf("unexpected StatementId: %d", event.StatementId)
}
if event.SubstatementId != 2 {
t.Errorf("unexpected SubstatementId: %d", event.SubstatementId)
}
if event.ObjectType != "TABLE" || event.Command != "SELECT" || event.ObjectName != "users" {
t.Errorf("unexpected object/command fields: %+v", event)
}
if event.DatabaseName != "prod_db" || event.DbUserName != "app_user" || event.RemoteHost != "10.0.1.45" {
t.Errorf("unexpected database connection fields: %+v", event)
}
if event.SessionId != "sess-998877" || event.RowCount != 42 || event.CommandText != "SELECT 1" {
t.Errorf("unexpected session/row fields: %+v", event)
}
if len(event.ParamList) != 3 || event.ParamList[0] != "val1" || event.ParamList[2] != "123" {
t.Errorf("unexpected ParamList: %+v", event.ParamList)
}
if event.Pid != 9876 || event.ClientApplication != "service-app" {
t.Errorf("unexpected pid/clientApp: %+v", event)
}
if event.ExitCode != "0" {
t.Errorf("unexpected ExitCode (expected string '0'): %q", event.ExitCode)
}
if event.Class != "READ" || event.ServerHost != "ip-10-0-2-10" || event.Type != "activity" {
t.Errorf("unexpected class/serverHost/type: %+v", event)
}
if event.StartTime != "2026-09-29 12:00:00.000" || event.ErrorMessage != "none" {
t.Errorf("unexpected startTime/errorMessage: %+v", event)
}
}

func TestGetHelpers_TypeConversions(t *testing.T) {
m := map[string]interface{}{
"str": "hello",
"numStr": 12345,
"int64Val": int64(999),
"intVal": int(888),
"strNum": "777",
"badNum": "not-a-number",
"nilVal": nil,
"sliceStr": []string{"a", "b"},
"sliceAny": []interface{}{"c", 4},
"notASlice": "scalar",
}

// getString
if getString(m, "str") != "hello" {
t.Errorf("expected 'hello', got %s", getString(m, "str"))
}
if getString(m, "numStr") != "12345" {
t.Errorf("expected '12345', got %s", getString(m, "numStr"))
}
if getString(m, "missing") != "" {
t.Errorf("expected empty string for missing key")
}
if getString(m, "nilVal") != "" {
t.Errorf("expected empty string for nil value")
}

// getInt64
if getInt64(m, "int64Val") != 999 {
t.Errorf("expected 999, got %d", getInt64(m, "int64Val"))
}
if getInt64(m, "intVal") != 888 {
t.Errorf("expected 888, got %d", getInt64(m, "intVal"))
}
if getInt64(m, "strNum") != 777 {
t.Errorf("expected 777, got %d", getInt64(m, "strNum"))
}
if getInt64(m, "badNum") != 0 {
t.Errorf("expected 0 for bad number string, got %d", getInt64(m, "badNum"))
}
if getInt64(m, "missing") != 0 {
t.Errorf("expected 0 for missing key")
}

// getStringSlice
s1 := getStringSlice(m, "sliceStr")
if len(s1) != 2 || s1[0] != "a" {
t.Errorf("unexpected string slice: %+v", s1)
}
s2 := getStringSlice(m, "sliceAny")
if len(s2) != 2 || s2[1] != "4" {
t.Errorf("unexpected any slice: %+v", s2)
}
if getStringSlice(m, "notASlice") != nil {
t.Errorf("expected nil for non-slice")
}
if getStringSlice(m, "missing") != nil {
t.Errorf("expected nil for missing key")
}
}

func BenchmarkConvertJSONToParquet(b *testing.B) {
dasData, err := os.ReadFile("testdata/das_events.json")
if err != nil {
b.Fatalf("failed to read fixture: %v", err)
}
filterCfg := &FilterConfig{Drop: []string{}, Query: map[string]interface{}{}}

b.ResetTimer()
b.ReportAllocs()

for i := 0; i < b.N; i++ {
_, err := convertJSONToParquet(dasData, filterCfg)
if err != nil {
b.Fatalf("convertJSONToParquet error: %v", err)
}
}
}

func BenchmarkMapToDatabaseActivityEvent(b *testing.B) {
m := map[string]interface{}{
"logTime": "2026-09-29 12:00:00.000",
"statementId": float64(1001),
"substatementId": float64(0),
"objectType": "TABLE",
"command": "SELECT",
"objectName": "users",
"databaseName": "prod_db",
"dbUserName": "app_user",
"remoteHost": "10.0.1.45",
"sessionId": "sess-998877",
"rowCount": float64(42),
"commandText": "SELECT id, name FROM users",
"paramList": []interface{}{"val1", "val2"},
"pid": float64(12345),
"clientApplication": "backend",
"exitCode": "0",
"class": "READ",
"serverHost": "ip-10-0-2-10",
"type": "activity",
"startTime": "2026-09-29 12:00:00.000",
"errorMessage": "",
}

b.ResetTimer()
b.ReportAllocs()

for i := 0; i < b.N; i++ {
_ = mapToDatabaseActivityEvent(m)
}
}
96 changes: 85 additions & 11 deletions main.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ import (
"io"
"log/slog"
"os"
"strconv"

"github.com/aws/aws-lambda-go/events"
"github.com/aws/aws-lambda-go/lambda"
Expand Down Expand Up @@ -298,6 +299,84 @@ type DatabaseActivityEvent struct {
ErrorMessage string `parquet:"errorMessage,dict" json:"errorMessage"`
}

// mapToDatabaseActivityEvent maps a generic map to DatabaseActivityEvent,
// avoiding expensive json.Marshal + json.Unmarshal serialization cycles.
func mapToDatabaseActivityEvent(m map[string]interface{}) DatabaseActivityEvent {
return DatabaseActivityEvent{
LogTime: getString(m, "logTime"),
StatementId: getInt64(m, "statementId"),
SubstatementId: getInt64(m, "substatementId"),
ObjectType: getString(m, "objectType"),
Command: getString(m, "command"),
ObjectName: getString(m, "objectName"),
DatabaseName: getString(m, "databaseName"),
DbUserName: getString(m, "dbUserName"),
RemoteHost: getString(m, "remoteHost"),
SessionId: getString(m, "sessionId"),
RowCount: getInt64(m, "rowCount"),
CommandText: getString(m, "commandText"),
ParamList: getStringSlice(m, "paramList"),
Pid: getInt64(m, "pid"),
ClientApplication: getString(m, "clientApplication"),
ExitCode: getString(m, "exitCode"),
Class: getString(m, "class"),
ServerHost: getString(m, "serverHost"),
Type: getString(m, "type"),
StartTime: getString(m, "startTime"),
ErrorMessage: getString(m, "errorMessage"),
}
}

func getString(m map[string]interface{}, key string) string {
if v, ok := m[key]; ok && v != nil {
if s, ok := v.(string); ok {
return s
}
return fmt.Sprintf("%v", v)
}
return ""
}

func getInt64(m map[string]interface{}, key string) int64 {
if v, ok := m[key]; ok && v != nil {
switch n := v.(type) {
case float64:
return int64(n)
case int64:
return n
case int:
return int64(n)
case json.Number:
i, _ := n.Int64()
return i
case string:
i, _ := strconv.ParseInt(n, 10, 64)
return i
}
}
return 0
}

func getStringSlice(m map[string]interface{}, key string) []string {
if v, ok := m[key]; ok && v != nil {
switch items := v.(type) {
case []interface{}:
result := make([]string, len(items))
for i, item := range items {
if s, ok := item.(string); ok {
result[i] = s
} else {
result[i] = fmt.Sprintf("%v", item)
}
}
return result
case []string:
return items
}
}
return nil
}

func convertJSONToParquet(decompressedJSON []byte, filterConfig *FilterConfig) ([]byte, error) {
// We unmarshal into a generic map to perform filtering and drops
var container struct {
Expand All @@ -308,6 +387,9 @@ func convertJSONToParquet(decompressedJSON []byte, filterConfig *FilterConfig) (
}

var buf bytes.Buffer
if len(decompressedJSON) > 0 {
buf.Grow(len(decompressedJSON) / 2)
}
writer := parquet.NewWriter(&buf, parquet.SchemaOf(new(DatabaseActivityEvent)), parquet.Compression(&snappy.Codec{}))

for _, eventMap := range container.DatabaseActivityEventList {
Expand All @@ -324,17 +406,9 @@ func convertJSONToParquet(decompressedJSON []byte, filterConfig *FilterConfig) (
delete(eventMap, d)
}

// Convert map back to bytes then struct to leverage strict typing for Parquet schema
eventBytes, err := json.Marshal(eventMap)
if err != nil {
continue
}

var event DatabaseActivityEvent
if err := json.Unmarshal(eventBytes, &event); err == nil {
if err := writer.Write(event); err != nil {
return nil, fmt.Errorf("failed to write parquet row: %w", err)
}
event := mapToDatabaseActivityEvent(eventMap)
if err := writer.Write(event); err != nil {
return nil, fmt.Errorf("failed to write parquet row: %w", err)
}
}
}
Expand Down
Loading