diff --git a/convert_test.go b/convert_test.go index b776d08..136d216 100644 --- a/convert_test.go +++ b/convert_test.go @@ -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) + } +} diff --git a/main.go b/main.go index 5bcef3f..465f1c4 100644 --- a/main.go +++ b/main.go @@ -10,6 +10,7 @@ import ( "io" "log/slog" "os" + "strconv" "github.com/aws/aws-lambda-go/events" "github.com/aws/aws-lambda-go/lambda" @@ -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 { @@ -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 { @@ -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) } } }