extracted struct-handling helpers into their own file; move sql_type_conversion_test to the parent module

Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
This commit is contained in:
Dmitri Dolguikh
2026-08-10 14:15:11 +02:00
parent 5594289924
commit 20f18d1479
3 changed files with 150 additions and 145 deletions

View File

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

View File

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

View File

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