added support for routes

Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
This commit is contained in:
Dmitri Dolguikh
2026-07-29 17:37:00 +02:00
parent 25e882004f
commit 23a5c0de4b
7 changed files with 535 additions and 174 deletions
@@ -2,10 +2,21 @@ package networkmapdb
import ( import (
"context" "context"
"database/sql"
"encoding/json"
"errors"
"reflect"
"strings"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata" "github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
) )
const (
NMAP_STRUCT_TAG = "nmap"
NMAP_SKIP = "skip"
NMAP_MAP_TO = "map_to"
)
type NetworkMapDBStore interface { type NetworkMapDBStore interface {
GetGroups(ctx context.Context, accountId string) ([]nmdata.Group, error) // TODO: join/populate peers GetGroups(ctx context.Context, accountId string) ([]nmdata.Group, error) // TODO: join/populate peers
GetPeers(ctx context.Context, accountId string) ([]nmdata.Peer, error) GetPeers(ctx context.Context, accountId string) ([]nmdata.Peer, error)
@@ -15,3 +26,86 @@ type NetworkMapDBStore interface {
type NetworkMapDBStoreImpl struct { type NetworkMapDBStoreImpl struct {
store NetworkMapDBStore store NetworkMapDBStore
} }
func FromSqlTypesToSharedTypes(src reflect.Value, dst reflect.Value) error {
typ := src.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.FieldByName(dstFieldName)
if !dstField.IsValid() {
return errors.New("invalid field in destination type: " + dstFieldName)
}
srcField := src.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)
}
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)
json.Unmarshal(s, dstField.Addr().Interface())
}
}
return nil
}
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])}
}
@@ -0,0 +1,141 @@
package networkmap_pgsql
import (
"database/sql"
"encoding/json"
"reflect"
"testing"
"time"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
"github.com/stretchr/testify/assert"
)
type g1 struct {
Name sql.NullString
}
type g2 struct {
Name string
}
type g11 struct {
Name sql.NullString
PublicId sql.NullString
}
type g22 struct {
Name string
PublicId string
}
type g111 struct {
Name sql.NullString
TrueOrFalse sql.NullBool
}
type g222 struct {
Name string
TrueOrFalse bool
}
type g1111 struct {
Time sql.NullTime
}
type g2222 struct {
Time time.Time
}
type g11111 struct {
Blob json.RawMessage
}
type embeddedS struct {
Name string
SomeField int
}
type g22222 struct {
Blob embeddedS
}
type t1 struct {
Field string `nmap:"skip"`
}
type dt1 struct {
Field string
}
type o1 struct {
Field string `nmap:"map_to:AnotherField"`
}
type do1 struct {
AnotherField string
}
type ptr1 struct {
Field sql.NullString
}
type ptro1 struct {
Field *string
}
type i1 struct {
Field sql.NullInt64
}
type io1 struct {
Field int
}
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.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.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.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.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.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.Equal(t, g22222{}, dst4)
src5 := t1{Field: "shouldskip"}
dst5 := dt1{}
assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src5).Elem(), reflect.ValueOf(&dst5).Elem()))
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.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.Equal(t, io1{Field: 1}, dst7)
}
@@ -4,14 +4,22 @@ import (
"context" "context"
"database/sql" "database/sql"
"encoding/json" "encoding/json"
"reflect"
"github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata" "github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
) )
const ( const (
GetGroupsQuery = ` GetGroupsQuery = `
select name, public_id, resources from groups where account_id=$1 select name, public_id, resources,
(
select array_agg(group_peers.peer_id)
from group_peers
where group_peers.group_id = groups.id
) as peers
from groups where account_id=$1
` `
) )
@@ -21,33 +29,24 @@ func (pg *PgStore) GetGroups(ctx context.Context, accountId string) ([]nmdata.Gr
return nil, err return nil, err
} }
var ( groups, err := pgx.CollectRows(rows, pgx.RowToStructByName[group])
group nmdata.Group toret := make([]nmdata.Group, 0, len(groups))
name, publicId, resources sql.NullString for _, g := range groups {
) dg := nmdata.Group{}
err := networkmapdb.FromSqlTypesToSharedTypes(
groups, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (nmdata.Group, error) { reflect.ValueOf(&g).Elem(), reflect.ValueOf(&dg).Elem())
err := row.Scan(&name, &publicId, &resources)
if err != nil { if err != nil {
return group, err return nil, err
} }
toret = append(toret, dg)
}
if name.Valid { return toret, err
group.Name = name.String }
}
if publicId.Valid {
group.PublicID = publicId.String
}
if resources.Valid {
groupResources := make([]nmdata.Resource, 0)
err := json.Unmarshal([]byte(resources.String), &groupResources)
if err != nil {
return group, err
}
group.Resources = groupResources
}
return group, nil
})
return groups, err type group struct {
Name sql.NullString
PublicID sql.NullString
Resources json.RawMessage
Peers []string
} }
@@ -2,6 +2,7 @@ package networkmap_pgsql
import ( import (
"context" "context"
"fmt"
"strings" "strings"
"testing" "testing"
@@ -15,7 +16,7 @@ import (
//go:embed test_db.sql //go:embed test_db.sql
var initDb string var initDb string
func TestXxx(t *testing.T) { func TestGetGroups(t *testing.T) {
ctx := context.TODO() ctx := context.TODO()
s, err := NewPostgresqlStore(ctx, "postgresql://root:netbird@localhost:5432/netbird") s, err := NewPostgresqlStore(ctx, "postgresql://root:netbird@localhost:5432/netbird")
@@ -26,7 +27,7 @@ func TestXxx(t *testing.T) {
_, err = s.pool.Query(ctx, "insert into groups (id, account_id, name, resources, public_id) VALUES('test-group-id-1','ck7bnf2t2r9s739pkug0','test-group-1', '[{\"ID\":\"cui7q2jl0ubs73d8qpi0\",\"Type\":\"host\"}]','public-id-1')") _, err = s.pool.Query(ctx, "insert into groups (id, account_id, name, resources, public_id) VALUES('test-group-id-1','ck7bnf2t2r9s739pkug0','test-group-1', '[{\"ID\":\"cui7q2jl0ubs73d8qpi0\",\"Type\":\"host\"}]','public-id-1')")
assert.NoError(t, err) assert.NoError(t, err)
groups, err := s.GetGroups(ctx, "ck7bnf2t2r9s739pkug0") groups, err := s.GetGroups(ctx, "ck7bnf2t2r9s739pkug0") //"ckd7ee2fic3c73dtendg")
assert.NoError(t, err) assert.NoError(t, err)
assert.Contains(t, assert.Contains(t,
groups, groups,
@@ -38,6 +39,72 @@ func TestXxx(t *testing.T) {
) )
} }
func TestGetPeers(t *testing.T) {
ctx := context.TODO()
s, err := NewPostgresqlStore(ctx, "postgresql://root:netbird@localhost:5432/netbird")
assert.NoError(t, err)
// err = loadSQL(ctx, s.pool, initDb)
//assert.NoError(t, err)
peers, err := s.GetPeers(ctx, "ck7bnf2t2r9s739pkug0") //"ckd7ee2fic3c73dtendg")
assert.NoError(t, err)
fmt.Print(peers)
// assert.Contains(t,
// groups,
// nmdata.Group{Name: "test-group-1", PublicID: "public-id-1", Resources: []nmdata.Resource{{ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
// )
// assert.Contains(t,
// groups,
// nmdata.Group{Name: "All", PublicID: "d9aejspvcsu517nkh4a0", Resources: []nmdata.Resource{{ID: "cui7olrl0ubs73d8qpe0", Type: "subnet"}, {ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
// )
}
func TestGetPolocies(t *testing.T) {
ctx := context.TODO()
s, err := NewPostgresqlStore(ctx, "postgresql://root:netbird@localhost:5432/netbird")
assert.NoError(t, err)
// err = loadSQL(ctx, s.pool, initDb)
//assert.NoError(t, err)
peers, err := s.GetPolicies(ctx, "ck7bnf2t2r9s739pkug0") //"ckd7ee2fic3c73dtendg")
assert.NoError(t, err)
fmt.Print(peers)
// assert.Contains(t,
// groups,
// nmdata.Group{Name: "test-group-1", PublicID: "public-id-1", Resources: []nmdata.Resource{{ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
// )
// assert.Contains(t,
// groups,
// nmdata.Group{Name: "All", PublicID: "d9aejspvcsu517nkh4a0", Resources: []nmdata.Resource{{ID: "cui7olrl0ubs73d8qpe0", Type: "subnet"}, {ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
// )
}
func TestGetRoutes(t *testing.T) {
ctx := context.TODO()
s, err := NewPostgresqlStore(ctx, "postgresql://root:netbird@localhost:5432/netbird")
assert.NoError(t, err)
// err = loadSQL(ctx, s.pool, initDb)
//assert.NoError(t, err)
peers, err := s.GetRoutes(ctx, "csg5iabl0ubs7398nf1g") //"ckd7ee2fic3c73dtendg")
assert.NoError(t, err)
fmt.Print(peers)
// assert.Contains(t,
// groups,
// nmdata.Group{Name: "test-group-1", PublicID: "public-id-1", Resources: []nmdata.Resource{{ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
// )
// assert.Contains(t,
// groups,
// nmdata.Group{Name: "All", PublicID: "d9aejspvcsu517nkh4a0", Resources: []nmdata.Resource{{ID: "cui7olrl0ubs73d8qpe0", Type: "subnet"}, {ID: "cui7q2jl0ubs73d8qpi0", Type: "host"}}},
// )
}
func loadSQL(ctx context.Context, pool *pgxpool.Pool, initdb string) error { func loadSQL(ctx context.Context, pool *pgxpool.Pool, initdb string) error {
queries := strings.Split(string(initdb), ";") queries := strings.Split(string(initdb), ";")
@@ -4,8 +4,10 @@ import (
"context" "context"
"database/sql" "database/sql"
"encoding/json" "encoding/json"
"reflect"
"github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata" "github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
) )
@@ -25,108 +27,94 @@ func (pg *PgStore) GetPeers(ctx context.Context, accountId string) ([]nmdata.Pee
return nil, err return nil, err
} }
var ( peers, err := pgx.CollectRows(rows, pgx.RowToStructByName[peer])
key, sshKey, dnsLabel, userId sql.NullString if err != nil {
lastLogin sql.NullTime return nil, err
sshEnabled, loginExpirationEnabled sql.NullBool }
ip, ipv6, locationConnectionIp []byte
metaFiles, metaCapabilities, metaFlags, metaNetworkAddresses []byte
metaWtVersion, metaGoOS, metaOSVersion, metaKernelVersion sql.NullString
locationCountryCode, locationCityName sql.NullString
)
peers, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (nmdata.Peer, error) { toret := make([]nmdata.Peer, 0, len(peers))
var peer nmdata.Peer for _, p := range peers {
err := row.Scan(&peer.ID, &key, &sshKey, &dnsLabel, &userId, &sshEnabled, &loginExpirationEnabled, &lastLogin, &ip, &ipv6, dp := nmdata.Peer{}
&metaWtVersion, &metaGoOS, &metaOSVersion, &metaKernelVersion, &metaNetworkAddresses, &metaFiles, &metaCapabilities, &metaFlags, err := networkmapdb.FromSqlTypesToSharedTypes(
&locationCountryCode, &locationCityName, &locationConnectionIp) reflect.ValueOf(&p).Elem(), reflect.ValueOf(&dp).Elem())
if err != nil { if err != nil {
return peer, err return nil, err
} }
if key.Valid { if p.MetaWtVersion.Valid {
peer.Key = key.String dp.Meta.WtVersion = p.MetaWtVersion.String
} }
if sshKey.Valid { if p.MetaGoOS.Valid {
peer.SSHKey = sshKey.String dp.Meta.GoOS = p.MetaGoOS.String
} }
if dnsLabel.Valid { if p.MetaOSVersion.Valid {
peer.DNSLabel = dnsLabel.String dp.Meta.OSVersion = p.MetaOSVersion.String
} }
if userId.Valid { if p.MetaKernelVersion.Valid {
peer.UserID = userId.String dp.Meta.KernelVersion = p.MetaKernelVersion.String
} }
if lastLogin.Valid { if p.LocationCountryCode.Valid {
peer.LastLogin = &lastLogin.Time dp.Location.CountryCode = p.LocationCountryCode.String
} }
if sshEnabled.Valid { if p.LocationCityName.Valid {
peer.SSHEnabled = sshEnabled.Bool dp.Location.CityName = p.LocationCityName.String
} }
if loginExpirationEnabled.Valid { if p.LocationConnectionIp != nil {
peer.LoginExpirationEnabled = loginExpirationEnabled.Bool err := json.Unmarshal(p.LocationConnectionIp, &dp.Location.ConnectionIP)
}
if metaWtVersion.Valid {
peer.Meta.WtVersion = metaWtVersion.String
}
if metaGoOS.Valid {
peer.Meta.GoOS = metaGoOS.String
}
if metaOSVersion.Valid {
peer.Meta.OSVersion = metaOSVersion.String
}
if metaKernelVersion.Valid {
peer.Meta.KernelVersion = metaKernelVersion.String
}
if locationCountryCode.Valid {
peer.Location.CountryCode = locationCountryCode.String
}
if locationCityName.Valid {
peer.Location.CityName = locationCityName.String
}
if ip != nil {
err := json.Unmarshal(ip, &peer.IP)
if err != nil { if err != nil {
return peer, err return toret, err
} }
} }
if ipv6 != nil { if p.MetaFiles != nil {
err := json.Unmarshal(ipv6, &peer.IPv6) err := json.Unmarshal(p.MetaFiles, &dp.Meta.Files)
if err != nil { if err != nil {
return peer, err return toret, err
} }
} }
if locationConnectionIp != nil { if p.MetaCapabilities != nil {
err := json.Unmarshal(locationConnectionIp, &peer.Location.ConnectionIP) err := json.Unmarshal(p.MetaCapabilities, &dp.Meta.Capabilities)
if err != nil { if err != nil {
return peer, err return toret, err
} }
} }
if metaFiles != nil { if p.MetaFlags != nil {
err := json.Unmarshal(metaFiles, &peer.Meta.Files) err := json.Unmarshal(p.MetaFlags, &dp.Meta.Flags)
if err != nil { if err != nil {
return peer, err return toret, err
} }
} }
if metaCapabilities != nil { if p.MetaNetworkAddresses != nil {
err := json.Unmarshal(metaCapabilities, &peer.Meta.Capabilities) err := json.Unmarshal(p.MetaNetworkAddresses, &dp.Meta.NetworkAddresses)
if err != nil { if err != nil {
return peer, err return toret, err
} }
} }
if metaFlags != nil { }
err := json.Unmarshal(metaFlags, &peer.Meta.Flags)
if err != nil {
return peer, err
}
}
if metaNetworkAddresses != nil {
err := json.Unmarshal(metaNetworkAddresses, &peer.Meta.NetworkAddresses)
if err != nil {
return peer, err
}
}
return peer, nil
})
return peers, err return toret, nil
}
// TODO add support for creating struct fields from denormalized fields
type peer struct {
ID string
Key sql.NullString
SSHKey sql.NullString
DNSLabel sql.NullString
UserID sql.NullString
LastLogin sql.NullTime
SSHEnabled sql.NullBool
LoginExpirationEnabled sql.NullBool
IP json.RawMessage
IPv6 json.RawMessage
LocationConnectionIp json.RawMessage `nmap:"skip"`
MetaFiles json.RawMessage `nmap:"skip"`
MetaCapabilities json.RawMessage `nmap:"skip"`
MetaFlags json.RawMessage `nmap:"skip"`
MetaNetworkAddresses json.RawMessage `nmap:"skip"`
MetaWtVersion sql.NullString `nmap:"skip"`
MetaGoOS sql.NullString `nmap:"skip"`
MetaOSVersion sql.NullString `nmap:"skip"`
MetaKernelVersion sql.NullString `nmap:"skip"`
LocationCountryCode sql.NullString `nmap:"skip"`
LocationCityName sql.NullString `nmap:"skip"`
} }
@@ -4,14 +4,16 @@ import (
"context" "context"
"database/sql" "database/sql"
"encoding/json" "encoding/json"
"reflect"
"github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata" "github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
) )
const ( const (
GetPoliciesQuery = ` GetPoliciesQuery = `
select p.id, p.public_id, p.enabled, p.source_posture_checks, pr.enabled, pr.action, pr.protocol, pr.bidirectional, select p.id, p.public_id, p.enabled, p.source_posture_checks, pr.enabled as rule_enabled, pr.action, pr.protocol, pr.bidirectional,
pr.sources, pr.destinations, pr.source_resource, pr.destination_resource, pr.ports, pr.port_ranges, pr.sources, pr.destinations, pr.source_resource, pr.destination_resource, pr.ports, pr.port_ranges,
pr.authorized_groups, pr.authorized_user pr.authorized_groups, pr.authorized_user
from policies as p from policies as p
@@ -26,18 +28,21 @@ func (pg *PgStore) GetPolicies(ctx context.Context, accountId string) ([]nmdata.
return nil, err return nil, err
} }
var ( policies, err := pgx.CollectRows(rows, pgx.RowToStructByName[policy])
publicId, sourcePostureChecks sql.NullString if err != nil {
enabled, ruleEnabled, bidirectional sql.NullBool return nil, err
action, protocol, sources, destinations sql.NullString }
sourceResource, destinationResource, ports, portRanges sql.NullString
authorizedGroups, authorizedUser sql.NullString toret := make([]nmdata.Policy, 0, len(policies))
) for _, p := range policies {
policy := nmdata.Policy{}
err := networkmapdb.FromSqlTypesToSharedTypes(
reflect.ValueOf(&p).Elem(), reflect.ValueOf(&policy).Elem())
if err != nil {
return nil, err
}
policies, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (nmdata.Policy, error) {
var policy nmdata.Policy
var policyRule *nmdata.PolicyRule var policyRule *nmdata.PolicyRule
pr := func() *nmdata.PolicyRule { pr := func() *nmdata.PolicyRule {
if policyRule != nil { if policyRule != nil {
return policyRule return policyRule
@@ -47,91 +52,91 @@ func (pg *PgStore) GetPolicies(ctx context.Context, accountId string) ([]nmdata.
return policyRule return policyRule
} }
err := row.Scan(&policy.ID, &publicId, &enabled, &sourcePostureChecks, &ruleEnabled, &action, &protocol, &bidirectional, if p.RuleEnabled.Valid {
&sources, &destinations, &sourceResource, &destinationResource, &ports, &portRanges, pr().Enabled = p.RuleEnabled.Bool
&authorizedGroups, &authorizedUser)
if err != nil {
return policy, err
} }
if p.Action.Valid {
if publicId.Valid { pr().Action = p.Action.String
policy.PublicID = publicId.String
} }
if sourcePostureChecks.Valid { if p.Protocol.Valid {
err := json.Unmarshal([]byte(sourcePostureChecks.String), &policy.SourcePostureChecks) pr().Protocol = p.Protocol.String
}
if p.Bidirectional.Valid {
pr().Bidirectional = p.Bidirectional.Bool
}
if len(p.Sources) > 0 {
err := json.Unmarshal([]byte(p.Sources), &pr().Sources)
if err != nil { if err != nil {
return policy, err return toret, err
} }
} }
if enabled.Valid { if len(p.Destinations) > 0 {
policy.Enabled = enabled.Bool err := json.Unmarshal([]byte(p.Destinations), &pr().Destinations)
}
if ruleEnabled.Valid {
pr().Enabled = ruleEnabled.Bool
}
if action.Valid {
pr().Action = action.String
}
if protocol.Valid {
pr().Protocol = protocol.String
}
if bidirectional.Valid {
pr().Bidirectional = bidirectional.Bool
}
if sources.Valid {
err := json.Unmarshal([]byte(sources.String), &pr().Sources)
if err != nil { if err != nil {
return policy, err return toret, err
} }
} }
if destinations.Valid { if len(p.SourceResource) > 0 {
err := json.Unmarshal([]byte(destinations.String), &pr().Destinations) err := json.Unmarshal([]byte(p.SourceResource), &pr().SourceResource)
if err != nil { if err != nil {
return policy, err return toret, err
} }
} }
if sourceResource.Valid { if len(p.DestinationResource) > 0 {
err := json.Unmarshal([]byte(sourceResource.String), &pr().SourceResource) err := json.Unmarshal([]byte(p.DestinationResource), &pr().DestinationResource)
if err != nil { if err != nil {
return policy, err return toret, err
} }
} }
if destinationResource.Valid { if len(p.Ports) > 0 {
err := json.Unmarshal([]byte(destinationResource.String), &pr().DestinationResource) err := json.Unmarshal([]byte(p.Ports), &pr().Ports)
if err != nil { if err != nil {
return policy, err return toret, err
} }
} }
if ports.Valid { if len(p.PortRanges) > 0 {
err := json.Unmarshal([]byte(ports.String), &pr().Ports) err := json.Unmarshal([]byte(p.PortRanges), &pr().PortRanges)
if err != nil { if err != nil {
return policy, err return toret, err
} }
} }
if portRanges.Valid { if len(p.AuthorizedGroups) > 0 {
err := json.Unmarshal([]byte(portRanges.String), &pr().PortRanges) err := json.Unmarshal([]byte(p.AuthorizedGroups), &pr().AuthorizedGroups)
if err != nil { if err != nil {
return policy, err return toret, err
} }
} }
if authorizedGroups.Valid { if p.AuthorizedUser.Valid {
err := json.Unmarshal([]byte(authorizedGroups.String), &pr().AuthorizedGroups) pr().AuthorizedUser = p.AuthorizedUser.String
if err != nil {
return policy, err
}
}
if authorizedUser.Valid {
pr().AuthorizedUser = authorizedUser.String
} }
if policyRule != nil { if policyRule != nil {
policyRule.ID = policy.ID policyRule.ID = p.ID
policyRule.PolicyID = policy.ID policyRule.PolicyID = p.ID
policy.Rules = []*nmdata.PolicyRule{policyRule} policy.Rules = []*nmdata.PolicyRule{policyRule}
} }
return policy, nil toret = append(toret, policy)
}) }
return policies, err return toret, err
}
type policy struct {
ID string
PublicID sql.NullString
SourcePostureChecks json.RawMessage
Enabled sql.NullBool
RuleEnabled sql.NullBool `nmap:"skip"`
Bidirectional sql.NullBool `nmap:"skip"`
Action sql.NullString `nmap:"skip"`
Protocol sql.NullString `nmap:"skip"`
Sources json.RawMessage `nmap:"skip"`
Destinations json.RawMessage `nmap:"skip"`
SourceResource json.RawMessage `nmap:"skip"`
DestinationResource json.RawMessage `nmap:"skip"`
Ports json.RawMessage `nmap:"skip"`
PortRanges json.RawMessage `nmap:"skip"`
AuthorizedGroups json.RawMessage `nmap:"skip"`
AuthorizedUser sql.NullString `nmap:"skip"`
} }
@@ -0,0 +1,67 @@
package networkmap_pgsql
import (
"context"
"database/sql"
"encoding/json"
"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 (
GetRoutesQuery = `
select id, account_id, public_id, network, domains, keep_route, net_id, description,
peer, peer as peer_id, peer_groups, network_type, masquerade, metric, enabled,
groups, access_control_groups, skip_auto_apply
from routes
where account_id=$1
`
)
func (pg *PgStore) GetRoutes(ctx context.Context, accountId string) ([]nmdata.Route, error) {
rows, err := pg.pool.Query(ctx, GetRoutesQuery, accountId)
if err != nil {
return nil, err
}
routes, err := pgx.CollectRows(rows, pgx.RowToStructByName[route])
if err != nil {
return nil, err
}
toret := make([]nmdata.Route, 0, len(routes))
for _, r := range routes {
route := nmdata.Route{}
err := networkmapdb.FromSqlTypesToSharedTypes(
reflect.ValueOf(&r).Elem(), reflect.ValueOf(&route).Elem())
if err != nil {
return nil, err
}
toret = append(toret, route)
}
return toret, nil
}
type route struct {
ID string
AccountID sql.NullString
PublicID sql.NullString
Network json.RawMessage
Domains json.RawMessage
KeepRoute sql.NullBool
NetID sql.NullString
Description sql.NullString
Peer sql.NullString
PeerID sql.NullString
PeerGroups json.RawMessage
NetworkType sql.NullInt64
Masquerade sql.NullBool
Metric sql.NullInt64
Enabled sql.NullBool
Groups json.RawMessage
AccessControlGroups json.RawMessage
SkipAutoApply sql.NullBool
}