support for GetAppliedZoneCandidates in sqlite

Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
This commit is contained in:
Dmitri Dolguikh
2026-08-10 14:59:03 +02:00
parent 166b3d7739
commit e530c25812
7 changed files with 112 additions and 79 deletions

View File

@@ -7,17 +7,12 @@ import (
"testing"
"github.com/miekg/dns"
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/shared/management/networkmap"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/stretchr/testify/assert"
)
func TestGetAppliedZoneCandidatesViaPgxConnection(t *testing.T) {
if engine == string(types.SqliteStoreEngine) {
t.Skip()
}
ctx := context.TODO()
execQuery(t, ctx,

View File

@@ -2,15 +2,10 @@ package networkmap_pgsql
import (
"context"
"database/sql"
"encoding/json"
"errors"
"reflect"
"github.com/jackc/pgx/v5"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
"github.com/netbirdio/netbird/shared/management/networkmap"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
)
const (
@@ -29,71 +24,10 @@ func (pgc *PgStoreConn) GetAppliedZoneCandidates(ctx context.Context, accountId
return nil, err
}
zones, err := pgx.CollectRows(rows, pgx.RowToStructByName[zone])
zones, err := pgx.CollectRows(rows, pgx.RowToStructByName[networkmapdb.Zone])
if err != nil {
return nil, err
}
toret := make([]networkmap.AppliedZoneCandidate, 0, len(zones))
currentZoneId := ""
for _, z := range zones {
if !z.RecordType.Valid {
continue
}
zone := nmdata.CustomZone{}
err := networkmapdb.FromSqlTypesToSharedTypes(
reflect.ValueOf(&z), reflect.ValueOf(&zone))
if err != nil {
return nil, err
}
var distributionGroups []string
if err := json.Unmarshal(z.DistributionGroups, &distributionGroups); err != nil {
return nil, err
}
if z.Id != currentZoneId {
zone.Records = []nmdata.SimpleRecord{}
toret = append(toret, appliedZoneCandidateFromZone(zone, distributionGroups))
currentZoneId = z.Id
}
rtype, rdata, err := networkmapdb.RecordTypeAndRdata(z.RecordType.String, z.RecordRData.String)
if err != nil {
if errors.Is(err, networkmapdb.ErrDnsUnsupportedRecordType) {
continue
}
return nil, err
}
lastZone := &toret[len(toret)-1]
lastZone.Zone.Records = append(lastZone.Zone.Records, nmdata.SimpleRecord{
Name: z.RecordName.String,
Class: z.RecordClass.String,
TTL: int(z.RecordTTL.Int64),
RData: rdata,
Type: rtype,
})
}
return toret, nil
}
type zone struct {
Id string `nmap:"skip"`
DistributionGroups json.RawMessage `nmap:"skip"`
Domain sql.NullString
SearchDomainDisabled sql.NullBool
RecordName sql.NullString `nmap:"skip"`
RecordType sql.NullString `nmap:"skip"`
RecordClass sql.NullString `nmap:"skip"`
RecordTTL sql.NullInt64 `nmap:"skip"`
RecordRData sql.NullString `nmap:"skip"`
}
func appliedZoneCandidateFromZone(z nmdata.CustomZone, distributionGroups []string) networkmap.AppliedZoneCandidate {
return networkmap.AppliedZoneCandidate{
DistributionGroups: distributionGroups,
Zone: z,
}
return networkmapdb.ZonesToAppliedZoneCandidates(zones)
}

View File

@@ -2,10 +2,14 @@ package networkmapdb
import (
"database/sql"
"encoding/json"
"errors"
"fmt"
"reflect"
"github.com/miekg/dns"
"github.com/netbirdio/netbird/shared/management/networkmap"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
)
var ErrDnsUnsupportedRecordType = errors.New("unsupported record type")
@@ -23,6 +27,18 @@ type Service struct {
Domain sql.NullString
}
type Zone struct {
Id string `nmap:"skip"`
Domain sql.NullString
SearchDomainDisabled sql.NullBool
DistributionGroups []byte `nmap:"skip,json"`
RecordName sql.NullString `nmap:"skip"`
RecordType sql.NullString `nmap:"skip"`
RecordClass sql.NullString `nmap:"skip"`
RecordTTL sql.NullInt64 `nmap:"skip"`
RecordRData sql.NullString `nmap:"skip"`
}
func RecordTypeAndRdata(t, rdata string) (int, string, error) {
switch t {
case "A":
@@ -35,3 +51,56 @@ func RecordTypeAndRdata(t, rdata string) (int, string, error) {
return 0, "", fmt.Errorf("record type: %s %w", t, ErrDnsUnsupportedRecordType)
}
}
func ZonesToAppliedZoneCandidates(zones []Zone) ([]networkmap.AppliedZoneCandidate, error) {
toret := make([]networkmap.AppliedZoneCandidate, 0, len(zones))
currentZoneId := ""
for _, z := range zones {
if !z.RecordType.Valid {
continue
}
zone := nmdata.CustomZone{}
err := FromSqlTypesToSharedTypes(
reflect.ValueOf(&z), reflect.ValueOf(&zone))
if err != nil {
return nil, err
}
var distributionGroups []string
if err := json.Unmarshal(z.DistributionGroups, &distributionGroups); err != nil {
return nil, err
}
if z.Id != currentZoneId {
zone.Records = []nmdata.SimpleRecord{}
toret = append(toret, AppliedZoneCandidateFromZone(zone, distributionGroups))
currentZoneId = z.Id
}
rtype, rdata, err := RecordTypeAndRdata(z.RecordType.String, z.RecordRData.String)
if err != nil {
if errors.Is(err, ErrDnsUnsupportedRecordType) {
continue
}
return nil, err
}
lastZone := &toret[len(toret)-1]
lastZone.Zone.Records = append(lastZone.Zone.Records, nmdata.SimpleRecord{
Name: z.RecordName.String,
Class: z.RecordClass.String,
TTL: int(z.RecordTTL.Int64),
RData: rdata,
Type: rtype,
})
}
return toret, nil
}
func AppliedZoneCandidateFromZone(z nmdata.CustomZone, distributionGroups []string) networkmap.AppliedZoneCandidate {
return networkmap.AppliedZoneCandidate{
DistributionGroups: distributionGroups,
Zone: z,
}
}

View File

@@ -36,7 +36,7 @@ func (sc *SqliteStoreConn) GetAccountSettings(ctx context.Context, accountId str
rows.Next()
a := account{}
err = rows.Scan(networkmapdb.StructFields(reflect.ValueOf(&a))...)
err = rows.Scan(networkmapdb.StructFields(&a)...)
if err != nil {
return nmdata.AccountSettingsInfo{}, err
}

View File

@@ -0,0 +1,38 @@
package networkmap_sqlite
import (
"context"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
"github.com/netbirdio/netbird/shared/management/networkmap"
)
const (
GetAccountZonesQuery = `
select zones.id as id, domain, not enable_search_domain as search_domain_disabled, distribution_groups,
r.name as record_name, r.type as record_type, 'IN' record_class, r.ttl as record_ttl, r.content as record_rdata
from zones
left join records as r on r.zone_id = zones.id
where zones.account_id=$1
`
)
func (sc *SqliteStoreConn) GetAppliedZoneCandidates(ctx context.Context, accountId string) ([]networkmap.AppliedZoneCandidate, error) {
rows, err := sc.Conn.QueryContext(ctx, GetAccountZonesQuery, accountId)
if err != nil {
return nil, err
}
defer rows.Close()
zones := make([]networkmapdb.Zone, 0)
for rows.Next() {
z := networkmapdb.Zone{}
err := rows.Scan(networkmapdb.StructFields(&z)...)
if err != nil {
return nil, err
}
zones = append(zones, z)
}
return networkmapdb.ZonesToAppliedZoneCandidates(zones)
}

View File

@@ -11,7 +11,6 @@ import (
"database/sql"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
"github.com/netbirdio/netbird/shared/management/networkmap"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
)
@@ -110,9 +109,6 @@ func (s *SqliteStoreConn) GetNetworkRouters(ctx context.Context, accountId strin
func (s *SqliteStoreConn) GetNetwork(ctx context.Context, accountId string) (nmdata.Network, error) {
return nmdata.Network{}, nil
}
func (s *SqliteStoreConn) GetAppliedZoneCandidates(ctx context.Context, accountId string) ([]networkmap.AppliedZoneCandidate, error) {
return nil, nil
}
func (s *SqliteStoreConn) GetPostureChecks(ctx context.Context, accountId string) ([]nmdata.PostureChecks, map[string]string, error) {
return nil, nil, nil
}

View File

@@ -122,7 +122,8 @@ func FromSqlTypesToSharedTypes(src reflect.Value, dst reflect.Value) error {
return nil
}
func StructFields(src reflect.Value) []any {
func StructFields(s any) []any {
src := reflect.ValueOf(s)
toret := make([]any, 0)
typ := src.Elem().Type()