diff --git a/management/internals/network_map_db/db_store.go b/management/internals/network_map_db/db_store.go index c96ab9b9a..9fbdaf35f 100644 --- a/management/internals/network_map_db/db_store.go +++ b/management/internals/network_map_db/db_store.go @@ -3,12 +3,7 @@ package networkmapdb import ( "context" "database/sql" - "encoding/json" - "errors" - "reflect" - "strings" - "github.com/rs/xid" "golang.org/x/exp/maps" "github.com/netbirdio/netbird/management/server/integrations/integrated_validator" @@ -93,125 +88,3 @@ func (s *NetworkMapDBStoreImpl) GetNetworkMapData(ctx context.Context, accountId return nmdata, nil } - -func FromSqlTypesToSharedTypes(src reflect.Value, dst reflect.Value) error { - typ := src.Elem().Type() - - for i := 0; i < typ.NumField(); i++ { - f := typ.Field(i) - - fieldTags := make(map[string]string) - if v := f.Tag.Get(NMAP_STRUCT_TAG); v != "" { - for _, t := range strings.Split(v, ",") { - kv := tagFromString(t) - fieldTags[kv.Key] = kv.Value - } - } - if _, ok := fieldTags[NMAP_SKIP]; ok { - continue - } - if f.PkgPath != "" { // skip unexported fields - continue - } - dstFieldName := f.Name - if override, ok := fieldTags[NMAP_MAP_TO]; ok { - dstFieldName = override - } - - dstField := dst.Elem().FieldByName(dstFieldName) - if !dstField.IsValid() { - return errors.New("unsupported type in destination field: " + dstFieldName) - } - - srcField := src.Elem().Field(i) - srcFieldType := srcField.Type().String() - switch srcFieldType { - case "string": - s := srcField.Interface().(string) - dstField.SetString(s) - case "sql.NullString": - s := srcField.Interface().(sql.NullString) - if s.Valid { - dstField.SetString(s.String) - } - if (dstFieldName == "PublicId" || dstFieldName == "PublicID") && s.String == "" { - dstField.SetString(xid.New().String()) // TODO (dmitri) this needs to be removed to support delta updates - } - case "sql.NullTime": - s := srcField.Interface().(sql.NullTime) - if s.Valid { - if dstField.Kind() == reflect.Ptr { - t := reflect.ValueOf(&s.Time).Elem() - dstField.Set(t.Addr()) - } else { - dstField.Set(reflect.ValueOf(s.Time)) - } - } - case "sql.NullBool": - s := srcField.Interface().(sql.NullBool) - if s.Valid { - dstField.SetBool(s.Bool) - } - case "sql.NullInt64": - s := srcField.Interface().(sql.NullInt64) - if s.Valid { - dstField.SetInt(s.Int64) - } - case "json.RawMessage": - s := srcField.Interface().(json.RawMessage) - if len(s) == 0 { - continue - } - if err := json.Unmarshal(s, dstField.Addr().Interface()); err != nil { - return err - } - case "[]byte", "[]uint8": - s := srcField.Interface().([]byte) - if _, ok := fieldTags[NMAP_JSON]; !ok || len(s) == 0 { - continue - } - if err := json.Unmarshal(s, dstField.Addr().Interface()); err != nil { - return err - } - case "[]string": - if srcField.IsNil() { - continue - } - dstv := reflect.MakeSlice(dstField.Type(), srcField.Len(), srcField.Cap()) - reflect.Copy(dstv, srcField) - dstField.Set(dstv) - } - } - - return nil -} - -func StructFields(src reflect.Value) []any { - toret := make([]any, 0) - typ := src.Elem().Type() - - for i := 0; i < typ.NumField(); i++ { - f := typ.Field(i) - if f.PkgPath != "" { // skip unexported fields - continue - } - - srcField := src.Elem().Field(i) - toret = append(toret, srcField.Addr().Interface()) - } - - return toret -} - -type fieldTag struct { - Key string - Value string -} - -func tagFromString(t string) fieldTag { - kv := strings.Split(t, ":") - if len(kv) == 1 { - return fieldTag{Key: strings.TrimSpace(kv[0])} - } - return fieldTag{Key: strings.TrimSpace(kv[0]), Value: strings.TrimSpace(kv[1])} -} diff --git a/management/internals/network_map_db/pgsql/sql_type_conversion_test.go b/management/internals/network_map_db/sql_type_conversion_test.go similarity index 73% rename from management/internals/network_map_db/pgsql/sql_type_conversion_test.go rename to management/internals/network_map_db/sql_type_conversion_test.go index 81e6590e2..77dab93a8 100644 --- a/management/internals/network_map_db/pgsql/sql_type_conversion_test.go +++ b/management/internals/network_map_db/sql_type_conversion_test.go @@ -1,4 +1,4 @@ -package networkmap_pgsql +package networkmapdb import ( "database/sql" @@ -7,26 +7,25 @@ import ( "testing" "time" - networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db" "github.com/stretchr/testify/assert" ) func TestNullStringSupport(t *testing.T) { src := withNullString{Name: sql.NullString{String: "string", Valid: true}} dst := withString{} - assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) + assert.NoError(t, FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) assert.Equal(t, withString{Name: "string"}, dst) src = withNullString{Name: sql.NullString{Valid: false}} dst = withString{} - assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) + assert.NoError(t, FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) assert.Equal(t, withString{Name: ""}, dst) } func TestNullBoolSupport(t *testing.T) { src := withNullBool{TrueOrFalse: sql.NullBool{Bool: true, Valid: true}} dst := withBool{} - assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) + assert.NoError(t, FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) assert.Equal(t, withBool{TrueOrFalse: true}, dst) } @@ -35,19 +34,19 @@ func TestRawJsonSupport(t *testing.T) { jb, _ := json.Marshal(embeddedS{Name: "blob-name", SomeField: 1}) src := withRawJson{Blob: json.RawMessage(jb)} dst := fromJson{} - assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) + assert.NoError(t, FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) assert.Equal(t, fromJson{Blob: embeddedS{Name: "blob-name", SomeField: 1}}, dst) src1 := withRawJson{} dst1 := fromJson{} - assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src1), reflect.ValueOf(&dst1))) + assert.NoError(t, FromSqlTypesToSharedTypes(reflect.ValueOf(&src1), reflect.ValueOf(&dst1))) assert.Equal(t, fromJson{}, dst1) } func TestShouldSkipTag(t *testing.T) { src5 := withSkipTag{Field: "shouldskip"} dst5 := emptySkipTagTarget{} - assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src5), reflect.ValueOf(&dst5))) + assert.NoError(t, FromSqlTypesToSharedTypes(reflect.ValueOf(&src5), reflect.ValueOf(&dst5))) assert.Equal(t, emptySkipTagTarget{}, dst5) } @@ -55,14 +54,14 @@ func TestShouldSkipTag(t *testing.T) { func TestMapToTag(t *testing.T) { src6 := withMapToTag{Field: "fieldvalue"} dst6 := mapToTagTarget{} - assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src6), reflect.ValueOf(&dst6))) + assert.NoError(t, FromSqlTypesToSharedTypes(reflect.ValueOf(&src6), reflect.ValueOf(&dst6))) assert.Equal(t, mapToTagTarget{AnotherField: "fieldvalue"}, dst6) } func TestNullableInt64Support(t *testing.T) { src := withInt64{Field: sql.NullInt64{Int64: int64(1), Valid: true}} dst := int64Target{} - assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) + assert.NoError(t, FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) assert.Equal(t, int64Target{Field: 1}, dst) } @@ -70,7 +69,7 @@ func TestNullableTimeSupport(t *testing.T) { now := time.Now() src := withNullableTime{Field: sql.NullTime{Time: now, Valid: true}} dst := nullableTimeTarget{} - assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) + assert.NoError(t, FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) assert.Equal(t, nullableTimeTarget{Field: now}, dst) } @@ -78,21 +77,21 @@ func TestNullableTimePointerSupport(t *testing.T) { now := time.Now() src := withNullableTime{Field: sql.NullTime{Time: now, Valid: true}} dst := nullableTimePointerTarget{} - assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) + assert.NoError(t, FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) assert.Equal(t, nullableTimePointerTarget{Field: &now}, dst) } func TestStringSLiceSupport(t *testing.T) { src := withStringSlice{Field: []string{"one"}} dst := withStringSlice{} - assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) + assert.NoError(t, FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) assert.Equal(t, withStringSlice{Field: []string{"one"}}, dst) } func TestNullStringSLiceSupport(t *testing.T) { src := withStringSlice{} dst := withStringSlice{} - assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) + assert.NoError(t, FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) assert.Equal(t, withStringSlice{}, dst) } @@ -106,7 +105,7 @@ func TestWithMultipleFields(t *testing.T) { Field5: "another", } dst := multipleFieldsTarget{} - assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) + assert.NoError(t, FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) assert.Equal(t, multipleFieldsTarget{ Field1: "aaa", Field2: true, @@ -119,7 +118,7 @@ func TestWithMultipleFields(t *testing.T) { func TestEmptyPublicIdsFilled(t *testing.T) { src := withEmptyPublicIds{} dst := emptyPublicIdTarget{} - assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) + assert.NoError(t, FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) assert.NotEmpty(t, dst.PublicID) assert.NotEmpty(t, dst.PublicId) } @@ -130,7 +129,7 @@ func TestByteSliceSupport(t *testing.T) { Field: []byte("[\"one\",\"two\",\"three\"]"), } dst := byteSliceTarget{} - assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) + assert.NoError(t, FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) assert.Equal(t, []string{"one", "two", "three"}, dst.Field) } @@ -139,7 +138,7 @@ func TestUint8SliceSupport(t *testing.T) { Field: []uint8("[\"one\",\"two\",\"three\"]"), } dst := uint8SliceTarget{} - assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) + assert.NoError(t, FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) assert.Equal(t, []string{"one", "two", "three"}, dst.Field) } diff --git a/management/internals/network_map_db/struct_helpers.go b/management/internals/network_map_db/struct_helpers.go new file mode 100644 index 000000000..324c6df93 --- /dev/null +++ b/management/internals/network_map_db/struct_helpers.go @@ -0,0 +1,133 @@ +package networkmapdb + +import ( + "database/sql" + "encoding/json" + "errors" + "reflect" + "strings" + + "github.com/rs/xid" +) + +type fieldTag struct { + Key string + Value string +} + +func tagFromString(t string) fieldTag { + kv := strings.Split(t, ":") + if len(kv) == 1 { + return fieldTag{Key: strings.TrimSpace(kv[0])} + } + return fieldTag{Key: strings.TrimSpace(kv[0]), Value: strings.TrimSpace(kv[1])} +} + +func FromSqlTypesToSharedTypes(src reflect.Value, dst reflect.Value) error { + typ := src.Elem().Type() + + for i := 0; i < typ.NumField(); i++ { + f := typ.Field(i) + + fieldTags := make(map[string]string) + if v := f.Tag.Get(NMAP_STRUCT_TAG); v != "" { + for _, t := range strings.Split(v, ",") { + kv := tagFromString(t) + fieldTags[kv.Key] = kv.Value + } + } + if _, ok := fieldTags[NMAP_SKIP]; ok { + continue + } + if f.PkgPath != "" { // skip unexported fields + continue + } + dstFieldName := f.Name + if override, ok := fieldTags[NMAP_MAP_TO]; ok { + dstFieldName = override + } + + dstField := dst.Elem().FieldByName(dstFieldName) + if !dstField.IsValid() { + return errors.New("unsupported type in destination field: " + dstFieldName) + } + + srcField := src.Elem().Field(i) + srcFieldType := srcField.Type().String() + switch srcFieldType { + case "string": + s := srcField.Interface().(string) + dstField.SetString(s) + case "sql.NullString": + s := srcField.Interface().(sql.NullString) + if s.Valid { + dstField.SetString(s.String) + } + if (dstFieldName == "PublicId" || dstFieldName == "PublicID") && s.String == "" { + dstField.SetString(xid.New().String()) // TODO (dmitri) this needs to be removed to support delta updates + } + case "sql.NullTime": + s := srcField.Interface().(sql.NullTime) + if s.Valid { + if dstField.Kind() == reflect.Ptr { + t := reflect.ValueOf(&s.Time).Elem() + dstField.Set(t.Addr()) + } else { + dstField.Set(reflect.ValueOf(s.Time)) + } + } + case "sql.NullBool": + s := srcField.Interface().(sql.NullBool) + if s.Valid { + dstField.SetBool(s.Bool) + } + case "sql.NullInt64": + s := srcField.Interface().(sql.NullInt64) + if s.Valid { + dstField.SetInt(s.Int64) + } + case "json.RawMessage": + s := srcField.Interface().(json.RawMessage) + if len(s) == 0 { + continue + } + if err := json.Unmarshal(s, dstField.Addr().Interface()); err != nil { + return err + } + case "[]byte", "[]uint8": + s := srcField.Interface().([]byte) + if _, ok := fieldTags[NMAP_JSON]; !ok || len(s) == 0 { + continue + } + if err := json.Unmarshal(s, dstField.Addr().Interface()); err != nil { + return err + } + case "[]string": + if srcField.IsNil() { + continue + } + dstv := reflect.MakeSlice(dstField.Type(), srcField.Len(), srcField.Cap()) + reflect.Copy(dstv, srcField) + dstField.Set(dstv) + } + } + + return nil +} + +func StructFields(src reflect.Value) []any { + toret := make([]any, 0) + typ := src.Elem().Type() + + for i := 0; i < typ.NumField(); i++ { + f := typ.Field(i) + if f.PkgPath != "" { // skip unexported fields + continue + } + + srcField := src.Elem().Field(i) + toret = append(toret, srcField.Addr().Interface()) + } + + return toret +}