From e530c25812cb85d0fb9ebec3a2622a4b3b8817c4 Mon Sep 17 00:00:00 2001 From: Dmitri Dolguikh Date: Mon, 10 Aug 2026 14:59:03 +0200 Subject: [PATCH] support for GetAppliedZoneCandidates in sqlite Signed-off-by: Dmitri Dolguikh --- .../network_map_db/pgsql/dns_test.go | 5 -- .../internals/network_map_db/pgsql/dns.go | 70 +------------------ .../internals/network_map_db/shared_types.go | 69 ++++++++++++++++++ .../network_map_db/sqlite/account_setting.go | 2 +- .../internals/network_map_db/sqlite/dns.go | 38 ++++++++++ .../network_map_db/sqlite/sqlite_store.go | 4 -- .../network_map_db/struct_helpers.go | 3 +- 7 files changed, 112 insertions(+), 79 deletions(-) create mode 100644 management/internals/network_map_db/sqlite/dns.go diff --git a/integration_tests/management/network_map_db/pgsql/dns_test.go b/integration_tests/management/network_map_db/pgsql/dns_test.go index 52539bf1b..63a609193 100644 --- a/integration_tests/management/network_map_db/pgsql/dns_test.go +++ b/integration_tests/management/network_map_db/pgsql/dns_test.go @@ -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, diff --git a/management/internals/network_map_db/pgsql/dns.go b/management/internals/network_map_db/pgsql/dns.go index dcae28556..c55aa4f5d 100644 --- a/management/internals/network_map_db/pgsql/dns.go +++ b/management/internals/network_map_db/pgsql/dns.go @@ -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) } diff --git a/management/internals/network_map_db/shared_types.go b/management/internals/network_map_db/shared_types.go index a48a8d06b..f97315331 100644 --- a/management/internals/network_map_db/shared_types.go +++ b/management/internals/network_map_db/shared_types.go @@ -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, + } +} diff --git a/management/internals/network_map_db/sqlite/account_setting.go b/management/internals/network_map_db/sqlite/account_setting.go index 3ea7d5cf3..8653da137 100644 --- a/management/internals/network_map_db/sqlite/account_setting.go +++ b/management/internals/network_map_db/sqlite/account_setting.go @@ -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 } diff --git a/management/internals/network_map_db/sqlite/dns.go b/management/internals/network_map_db/sqlite/dns.go new file mode 100644 index 000000000..0fc1d6ad2 --- /dev/null +++ b/management/internals/network_map_db/sqlite/dns.go @@ -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) +} diff --git a/management/internals/network_map_db/sqlite/sqlite_store.go b/management/internals/network_map_db/sqlite/sqlite_store.go index ecfb9ec08..7eaa5af4f 100644 --- a/management/internals/network_map_db/sqlite/sqlite_store.go +++ b/management/internals/network_map_db/sqlite/sqlite_store.go @@ -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 } diff --git a/management/internals/network_map_db/struct_helpers.go b/management/internals/network_map_db/struct_helpers.go index e2e20cfa7..5f67bc8e3 100644 --- a/management/internals/network_map_db/struct_helpers.go +++ b/management/internals/network_map_db/struct_helpers.go @@ -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()