fix handling of string slices

Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
This commit is contained in:
Dmitri Dolguikh
2026-07-30 11:31:53 +02:00
parent 42ce83a8f3
commit 5a10561ca1
11 changed files with 67 additions and 22 deletions

View File

@@ -33,7 +33,7 @@ type NetworkMapDBStoreImpl struct {
}
func FromSqlTypesToSharedTypes(src reflect.Value, dst reflect.Value) error {
typ := src.Type()
typ := src.Elem().Type()
for i := 0; i < typ.NumField(); i++ {
f := typ.Field(i)
@@ -56,13 +56,12 @@ func FromSqlTypesToSharedTypes(src reflect.Value, dst reflect.Value) error {
dstFieldName = override
}
dstField := dst.FieldByName(dstFieldName)
dstField := dst.Elem().FieldByName(dstFieldName)
if !dstField.IsValid() {
return errors.New("invalid field in destination type: " + dstFieldName)
return errors.New("unsupported type in destination field: " + dstFieldName)
}
srcField := src.Field(i)
srcField := src.Elem().Field(i)
srcFieldType := srcField.Type().String()
switch srcFieldType {
case "string":
@@ -96,6 +95,10 @@ func FromSqlTypesToSharedTypes(src reflect.Value, dst reflect.Value) error {
case "json.RawMessage":
s := srcField.Interface().(json.RawMessage)
json.Unmarshal(s, dstField.Addr().Interface())
case "[]string":
dstv := reflect.MakeSlice(dstField.Type(), srcField.Len(), srcField.Cap())
reflect.Copy(dstv, srcField)
dstField.Set(dstv)
}
}

View File

@@ -95,47 +95,47 @@ func TestOne(t *testing.T) {
src := g1{Name: sql.NullString{String: "string", Valid: true}}
dst := g2{}
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src).Elem(), reflect.ValueOf(&dst).Elem()))
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst)))
assert.Equal(t, g2{Name: "string"}, dst)
src = g1{Name: sql.NullString{Valid: false}}
dst = g2{}
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src).Elem(), reflect.ValueOf(&dst).Elem()))
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst)))
assert.Equal(t, g2{Name: ""}, dst)
src1 := g11{Name: sql.NullString{String: "aaa", Valid: true}, PublicId: sql.NullString{String: "id", Valid: true}}
dst1 := g22{}
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src1).Elem(), reflect.ValueOf(&dst1).Elem()))
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src1), reflect.ValueOf(&dst1)))
assert.Equal(t, g22{Name: "aaa", PublicId: "id"}, dst1)
src2 := g111{Name: sql.NullString{String: "aaa", Valid: true}, TrueOrFalse: sql.NullBool{Bool: true, Valid: true}}
dst2 := g222{}
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src2).Elem(), reflect.ValueOf(&dst2).Elem()))
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src2), reflect.ValueOf(&dst2)))
assert.Equal(t, g222{Name: "aaa", TrueOrFalse: true}, dst2)
jb, _ := json.Marshal(embeddedS{Name: "blob-name", SomeField: 1})
src3 := g11111{Blob: json.RawMessage(jb)}
dst3 := g22222{}
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src3).Elem(), reflect.ValueOf(&dst3).Elem()))
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src3), reflect.ValueOf(&dst3)))
assert.Equal(t, g22222{Blob: embeddedS{Name: "blob-name", SomeField: 1}}, dst3)
src4 := g11111{}
dst4 := g22222{}
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src4).Elem(), reflect.ValueOf(&dst4).Elem()))
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src4), reflect.ValueOf(&dst4)))
assert.Equal(t, g22222{}, dst4)
src5 := t1{Field: "shouldskip"}
dst5 := dt1{}
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src5).Elem(), reflect.ValueOf(&dst5).Elem()))
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src5), reflect.ValueOf(&dst5)))
assert.Equal(t, dt1{}, dst5)
src6 := o1{Field: "fieldvalue"}
dst6 := do1{}
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src6).Elem(), reflect.ValueOf(&dst6).Elem()))
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src6), reflect.ValueOf(&dst6)))
assert.Equal(t, do1{AnotherField: "fieldvalue"}, dst6)
src7 := i1{Field: sql.NullInt64{Int64: int64(1), Valid: true}}
dst7 := io1{Field: 1}
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src7).Elem(), reflect.ValueOf(&dst7).Elem()))
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src7), reflect.ValueOf(&dst7)))
assert.Equal(t, io1{Field: 1}, dst7)
}

View File

@@ -0,0 +1,42 @@
package networkmap_pgsql
import (
"context"
"reflect"
"github.com/jackc/pgx/v5"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
)
const (
GetCustomZonesQuery = `
select id, public_id, name, description, name_servers, groups, "primary", domains, enabled, search_domains_enabled
from name_server_groups
where account_id=$1
`
)
func (pg *PgStore) GetCustomZones(ctx context.Context, accountId string) ([]nmdata.NameServerGroup, error) {
rows, err := pg.pool.Query(ctx, GetNameserversQuery, accountId)
if err != nil {
return nil, err
}
nsgroups, err := pgx.CollectRows(rows, pgx.RowToStructByName[nameserverGroup])
if err != nil {
return nil, err
}
toret := make([]nmdata.NameServerGroup, 0, len(nsgroups))
for _, nsg := range nsgroups {
group := nmdata.NameServerGroup{}
err := networkmapdb.FromSqlTypesToSharedTypes(
reflect.ValueOf(&nsg), reflect.ValueOf(&group))
if err != nil {
return nil, err
}
toret = append(toret, group)
}
return toret, nil
}

View File

@@ -34,7 +34,7 @@ func (pg *PgStore) GetGroups(ctx context.Context, accountId string) ([]nmdata.Gr
for _, g := range groups {
dg := nmdata.Group{}
err := networkmapdb.FromSqlTypesToSharedTypes(
reflect.ValueOf(&g).Elem(), reflect.ValueOf(&dg).Elem())
reflect.ValueOf(&g), reflect.ValueOf(&dg))
if err != nil {
return nil, err
}

View File

@@ -34,7 +34,7 @@ func (pg *PgStore) GetNameServerGroups(ctx context.Context, accountId string) ([
for _, nsg := range nsgroups {
group := nmdata.NameServerGroup{}
err := networkmapdb.FromSqlTypesToSharedTypes(
reflect.ValueOf(&nsg).Elem(), reflect.ValueOf(&group).Elem())
reflect.ValueOf(&nsg), reflect.ValueOf(&group))
if err != nil {
return nil, err
}

View File

@@ -32,7 +32,7 @@ func (pg *PgStore) GetNetwork(ctx context.Context, accountId string) (nmdata.Net
toret := nmdata.Network{}
err = networkmapdb.FromSqlTypesToSharedTypes(
reflect.ValueOf(&n).Elem(), reflect.ValueOf(&toret).Elem())
reflect.ValueOf(&n), reflect.ValueOf(&toret))
if err != nil {
return nmdata.Network{}, err
}

View File

@@ -34,7 +34,7 @@ func (pg *PgStore) GetNetworkResources(ctx context.Context, accountId string) ([
for _, nres := range netresorces {
resource := nmdata.NetworkResource{}
err := networkmapdb.FromSqlTypesToSharedTypes(
reflect.ValueOf(&nres).Elem(), reflect.ValueOf(&resource).Elem())
reflect.ValueOf(&nres), reflect.ValueOf(&resource))
if err != nil {
return nil, err
}

View File

@@ -34,7 +34,7 @@ func (pg *PgStore) GetNetworkRouters(ctx context.Context, accountId string) ([]n
for _, nrt := range netrouters {
router := nmdata.NetworkRouter{}
err := networkmapdb.FromSqlTypesToSharedTypes(
reflect.ValueOf(&nrt).Elem(), reflect.ValueOf(&router).Elem())
reflect.ValueOf(&nrt), reflect.ValueOf(&router))
if err != nil {
return nil, err
}

View File

@@ -36,7 +36,7 @@ func (pg *PgStore) GetPeers(ctx context.Context, accountId string) ([]nmdata.Pee
for _, p := range peers {
dp := nmdata.Peer{}
err := networkmapdb.FromSqlTypesToSharedTypes(
reflect.ValueOf(&p).Elem(), reflect.ValueOf(&dp).Elem())
reflect.ValueOf(&p), reflect.ValueOf(&dp))
if err != nil {
return nil, err
}

View File

@@ -37,7 +37,7 @@ func (pg *PgStore) GetPolicies(ctx context.Context, accountId string) ([]nmdata.
for _, p := range policies {
policy := nmdata.Policy{}
err := networkmapdb.FromSqlTypesToSharedTypes(
reflect.ValueOf(&p).Elem(), reflect.ValueOf(&policy).Elem())
reflect.ValueOf(&p), reflect.ValueOf(&policy))
if err != nil {
return nil, err
}

View File

@@ -36,7 +36,7 @@ func (pg *PgStore) GetRoutes(ctx context.Context, accountId string) ([]nmdata.Ro
for _, r := range routes {
route := nmdata.Route{}
err := networkmapdb.FromSqlTypesToSharedTypes(
reflect.ValueOf(&r).Elem(), reflect.ValueOf(&route).Elem())
reflect.ValueOf(&r), reflect.ValueOf(&route))
if err != nil {
return nil, err
}