support for GetDomains in sqlite

Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
This commit is contained in:
Dmitri Dolguikh
2026-08-10 15:55:24 +02:00
parent e530c25812
commit 763b8f6933
6 changed files with 68 additions and 24 deletions

View File

@@ -8,15 +8,10 @@ import (
"testing"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
"github.com/netbirdio/netbird/management/server/types"
"github.com/stretchr/testify/assert"
)
func TestGetDomains(t *testing.T) {
if engine == string(types.SqliteStoreEngine) {
t.Skip()
}
ctx := context.TODO()
execQuery(t, ctx,

View File

@@ -32,14 +32,8 @@ func (sc *SqliteStoreConn) GetAccountSettings(ctx context.Context, accountId str
if err != nil {
return nmdata.AccountSettingsInfo{}, err
}
defer rows.Close()
rows.Next()
a := account{}
err = rows.Scan(networkmapdb.StructFields(&a)...)
if err != nil {
return nmdata.AccountSettingsInfo{}, err
}
a, err := networkmapdb.CollectOneRowForSqlite[account](rows)
settingsInfo := nmdata.AccountSettingsInfo{}
err = networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&a), reflect.ValueOf(&settingsInfo))

View File

@@ -22,16 +22,10 @@ func (sc *SqliteStoreConn) GetAppliedZoneCandidates(ctx context.Context, account
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)
zones, err := networkmapdb.CollectRowsForSqlite[networkmapdb.Zone](rows)
if err != nil {
return nil, err
}
return networkmapdb.ZonesToAppliedZoneCandidates(zones)

View File

@@ -0,0 +1,24 @@
package networkmap_sqlite
import (
"context"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
)
const (
GetDomainsQuery = `
select domain, target_cluster
from domains
where account_id=$1 and domain<>'' and target_cluster<>''
`
)
func (sc *SqliteStoreConn) GetDomains(ctx context.Context, accountId string) ([]networkmapdb.Domain, error) {
rows, err := sc.Conn.QueryContext(ctx, GetDomainsQuery, accountId)
if err != nil {
return nil, err
}
return networkmapdb.CollectRowsForSqlite[networkmapdb.Domain](rows)
}

View File

@@ -85,9 +85,6 @@ func (s *SqliteStore) UsingConn() *SqliteStoreConn {
func (s *SqliteStoreConn) GetGroups(ctx context.Context, accountId string) ([]nmdata.Group, map[string]map[string]any, error) {
return nil, nil, nil
}
func (s *SqliteStoreConn) GetDomains(ctx context.Context, accountId string) ([]networkmapdb.Domain, error) {
return nil, nil
}
func (s *SqliteStoreConn) GetPeers(ctx context.Context, accountId string) ([]nmdata.Peer, map[string][]*nmdata.Peer, error) {
return nil, nil, nil
}

View File

@@ -10,6 +10,8 @@ import (
"github.com/rs/xid"
)
var ErrNoRows = errors.New("no rows in result set")
const (
NMAP_STRUCT_TAG = "nmap"
NMAP_SKIP = "skip"
@@ -139,3 +141,41 @@ func StructFields(s any) []any {
return toret
}
func CollectOneRowForSqlite[T any](rows *sql.Rows) (T, error) {
defer rows.Close()
var r T
if !rows.Next() {
if err := rows.Err(); err != nil {
return r, err
}
return r, ErrNoRows
}
err := rows.Scan(StructFields(&r)...)
if err != nil {
return r, err
}
return r, nil
}
func CollectRowsForSqlite[T any](rows *sql.Rows) ([]T, error) {
defer rows.Close()
toret := make([]T, 0)
for rows.Next() {
var r T
err := rows.Scan(StructFields(&r)...)
if err != nil {
return nil, err
}
toret = append(toret, r)
}
if err := rows.Err(); err != nil {
return nil, err
}
return toret, nil
}