From 65f184141b02d0d19e4f0f35dc2a6b43741b895f Mon Sep 17 00:00:00 2001 From: Dmitri Dolguikh Date: Mon, 10 Aug 2026 13:31:04 +0200 Subject: [PATCH] added support for GetAccountSettings to sqlite store Signed-off-by: Dmitri Dolguikh --- .../pgsql/account_settings_test.go | 4 - .../network_map_db/pgsql/main_test.go | 42 ++++++++--- .../network_map_db/pgsql/pg_test_store.go | 9 +-- .../network_map_db/pgsql/service_test.go | 9 ++- .../network_map_db/pgsql/sqlite_test_store.go | 48 ++++++++++++ .../internals/network_map_db/db_store.go | 26 +++++++ .../pgsql/sql_type_conversion_test.go | 35 +++++++++ .../network_map_db/sqlite/account_setting.go | 65 ++++++++++++++++ .../network_map_db/sqlite/sqlite_store.go | 75 ++++++++++++++++++- 9 files changed, 284 insertions(+), 29 deletions(-) create mode 100644 integration_tests/management/network_map_db/pgsql/sqlite_test_store.go create mode 100644 management/internals/network_map_db/sqlite/account_setting.go diff --git a/integration_tests/management/network_map_db/pgsql/account_settings_test.go b/integration_tests/management/network_map_db/pgsql/account_settings_test.go index 5fe3cc1b7..d7927aaf1 100644 --- a/integration_tests/management/network_map_db/pgsql/account_settings_test.go +++ b/integration_tests/management/network_map_db/pgsql/account_settings_test.go @@ -7,15 +7,11 @@ import ( "testing" "time" - "github.com/netbirdio/netbird/management/server/types" "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" "github.com/stretchr/testify/assert" ) func TestGetAccountSettings(t *testing.T) { - if engine == string(types.SqliteStoreEngine) { - t.Skip() - } ctx := context.TODO() execQuery(t, ctx, diff --git a/integration_tests/management/network_map_db/pgsql/main_test.go b/integration_tests/management/network_map_db/pgsql/main_test.go index a8d482654..d47b52ffb 100644 --- a/integration_tests/management/network_map_db/pgsql/main_test.go +++ b/integration_tests/management/network_map_db/pgsql/main_test.go @@ -14,6 +14,7 @@ import ( networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db" networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql" + networkmap_sqlite "github.com/netbirdio/netbird/management/internals/network_map_db/sqlite" "github.com/netbirdio/netbird/management/server/types" ) @@ -21,24 +22,28 @@ import ( var baseData string var ( - pgstore *networkmap_pgsql.PgStore - engine string + pgstore *networkmap_pgsql.PgStore + sqlitestore *networkmap_sqlite.SqliteStore + engine string ) func TestMain(m *testing.M) { + var cleanup func() kind, _ := os.LookupEnv("NETBIRD_STORE_ENGINE") switch kind { case "": - engine = string(types.PostgresStoreEngine) - case string(types.PostgresStoreEngine), string(types.SqliteStoreEngine): + engine = string(types.SqliteStoreEngine) + sqlitestore, cleanup = createSqliteTestStore(baseData) + case string(types.PostgresStoreEngine): engine = kind + pgstore, cleanup = createPGTestStore(baseData) + case string(types.SqliteStoreEngine): + engine = kind + sqlitestore, cleanup = createSqliteTestStore(baseData) default: log.Fatalf("unsupported db '%s' in NETBIRD_STORE_ENGINE env var", kind) } - store, cleanup := createPGTestStore(baseData) - pgstore = store - code := m.Run() cleanup() @@ -47,15 +52,28 @@ func TestMain(m *testing.M) { func conn(t *testing.T, ctx context.Context) networkmapdb.NetworkMapDBStoreConn { t.Helper() - c, err := pgstore.Pool.Acquire(ctx) - assert.NoError(t, err) - return pgstore.UsingConnection(c.Conn()) + switch engine { + case string(types.PostgresStoreEngine): + c, err := pgstore.Pool.Acquire(ctx) + assert.NoError(t, err) + return pgstore.UsingConnection(c.Conn()) + case string(types.SqliteStoreEngine): + return sqlitestore.UsingConn() + } + log.Fatalf("unknown db engine kind %s", engine) + return nil } func execQuery(t *testing.T, ctx context.Context, q string) { t.Helper() - _, err := pgstore.Pool.Exec(ctx, q) - assert.NoError(t, err) + switch engine { + case string(types.PostgresStoreEngine): + _, err := pgstore.Pool.Exec(ctx, q) + assert.NoError(t, err) + case string(types.SqliteStoreEngine): + _, err := sqlitestore.Db.ExecContext(ctx, q) + assert.NoError(t, err) + } } // use to parse time in time.RFC3339Nano format diff --git a/integration_tests/management/network_map_db/pgsql/pg_test_store.go b/integration_tests/management/network_map_db/pgsql/pg_test_store.go index 2c7afacc5..1710747b5 100644 --- a/integration_tests/management/network_map_db/pgsql/pg_test_store.go +++ b/integration_tests/management/network_map_db/pgsql/pg_test_store.go @@ -58,22 +58,17 @@ func createPGTestStore(baseData string) (*networkmap_pgsql.PgStore, func()) { } ctx := context.TODO() - s, err := networkmap_pgsql.NewPostgresqlStore(ctx, dsn) + pgstore, err := networkmap_pgsql.NewPostgresqlStore(ctx, dsn) if err != nil { log.Fatal("error creating postgres store %w", err) } for _, query := range strings.Split(baseData, ";") { - if _, err := s.Pool.Exec(ctx, query); err != nil { + if _, err := pgstore.Pool.Exec(ctx, query); err != nil { log.Fatalf("error initializing db: %s", err.Error()) } } - pgstore, err := networkmap_pgsql.NewPostgresqlStore(ctx, dsn) - if err != nil { - log.Fatalf("error creating pg store %v", err.Error()) - } - return pgstore, cleanup } diff --git a/integration_tests/management/network_map_db/pgsql/service_test.go b/integration_tests/management/network_map_db/pgsql/service_test.go index 31cddcb29..ad818b628 100644 --- a/integration_tests/management/network_map_db/pgsql/service_test.go +++ b/integration_tests/management/network_map_db/pgsql/service_test.go @@ -19,18 +19,15 @@ func TestGetPrivateServicesViaPgxConnection(t *testing.T) { } ctx := context.TODO() - _, err := pgstore.Pool.Exec(ctx, + execQuery(t, ctx, `insert into services (id, account_id, enabled, private, access_groups, proxy_cluster, domain) values('service-1','account-1',true,true,'["group-one-resource-id"]','test-1.com','test-2.com')`) - assert.NoError(t, err) execQuery(t, ctx, `insert into services (id, account_id, enabled, private, access_groups, proxy_cluster, domain) values('service-2','account-1',true,true,'["group-one-resource-id","group-two-resources-id"]','test-3.com','test-4.com')`) - assert.NoError(t, err) execQuery(t, ctx, `insert into services (id, account_id, enabled, private, access_groups, proxy_cluster, domain) values('service-3','account-1',null,null,null,null,null)`) - assert.NoError(t, err) services, err := conn(t, ctx).GetPrivateServices(ctx, "account-1") assert.NoError(t, err) @@ -58,6 +55,10 @@ func TestGetPrivateServicesViaPgxConnection(t *testing.T) { } func TestGetProxyTargetedDomainResourceIDsViaPgxConnection(t *testing.T) { + if engine == string(types.SqliteStoreEngine) { + t.Skip() + } + ctx := context.TODO() execQuery(t, ctx, diff --git a/integration_tests/management/network_map_db/pgsql/sqlite_test_store.go b/integration_tests/management/network_map_db/pgsql/sqlite_test_store.go new file mode 100644 index 000000000..1c70c93d4 --- /dev/null +++ b/integration_tests/management/network_map_db/pgsql/sqlite_test_store.go @@ -0,0 +1,48 @@ +//go:build integration + +package networkmap_pgsql + +import ( + "context" + "fmt" + "runtime" + "strings" + + networkmap_sqlite "github.com/netbirdio/netbird/management/internals/network_map_db/sqlite" + gormstore "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/management/server/types" + log "github.com/sirupsen/logrus" + "gorm.io/driver/sqlite" + "gorm.io/gorm" +) + +func createSqliteTestStore(baseData string) (*networkmap_sqlite.SqliteStore, func()) { + storeSqliteFileName := ":memory:" + storeStr := fmt.Sprintf("%s?cache=shared", storeSqliteFileName) + if runtime.GOOS == "windows" { + // Vo avoid `The process cannot access the file because it is being used by another process` on Windows + storeStr = storeSqliteFileName + } + + db, err := gorm.Open(sqlite.Open(storeStr), &gorm.Config{}) + if err != nil { + log.Fatalf("error initializing db: %s", err.Error()) + } + _, err = gormstore.NewSqlStore(context.TODO(), db, types.SqliteStoreEngine, nil, false) + if err != nil { + log.Fatalf("error initializing db: %s", err.Error()) + } + + sqldb, err := db.DB() + if err != nil { + log.Fatalf("error initializing db: %s", err.Error()) + + } + for _, query := range strings.Split(baseData, ";") { + if _, err := sqldb.Exec(query); err != nil { + log.Fatalf("error initializing db: %s", err.Error()) + } + } + + return &networkmap_sqlite.SqliteStore{Db: sqldb}, func() {} +} diff --git a/management/internals/network_map_db/db_store.go b/management/internals/network_map_db/db_store.go index fe31c5a5f..c96ab9b9a 100644 --- a/management/internals/network_map_db/db_store.go +++ b/management/internals/network_map_db/db_store.go @@ -21,6 +21,7 @@ const ( NMAP_STRUCT_TAG = "nmap" NMAP_SKIP = "skip" NMAP_MAP_TO = "map_to" + NMAP_JSON = "json" ) type NetworkMapDBStore interface { //nolint:revive // established name across the codebase @@ -164,6 +165,14 @@ func FromSqlTypesToSharedTypes(src reflect.Value, dst reflect.Value) error { if err := json.Unmarshal(s, dstField.Addr().Interface()); err != nil { return err } + case "[]byte", "[]uint8": + s := srcField.Interface().([]byte) + if _, ok := fieldTags[NMAP_JSON]; !ok || len(s) == 0 { + continue + } + if err := json.Unmarshal(s, dstField.Addr().Interface()); err != nil { + return err + } case "[]string": if srcField.IsNil() { continue @@ -177,6 +186,23 @@ func FromSqlTypesToSharedTypes(src reflect.Value, dst reflect.Value) error { return nil } +func StructFields(src reflect.Value) []any { + toret := make([]any, 0) + typ := src.Elem().Type() + + for i := 0; i < typ.NumField(); i++ { + f := typ.Field(i) + if f.PkgPath != "" { // skip unexported fields + continue + } + + srcField := src.Elem().Field(i) + toret = append(toret, srcField.Addr().Interface()) + } + + return toret +} + type fieldTag struct { Key string Value string diff --git a/management/internals/network_map_db/pgsql/sql_type_conversion_test.go b/management/internals/network_map_db/pgsql/sql_type_conversion_test.go index 9cd7086a1..81e6590e2 100644 --- a/management/internals/network_map_db/pgsql/sql_type_conversion_test.go +++ b/management/internals/network_map_db/pgsql/sql_type_conversion_test.go @@ -124,6 +124,25 @@ func TestEmptyPublicIdsFilled(t *testing.T) { assert.NotEmpty(t, dst.PublicId) } +// only []byte and []uint8 slices with "json" tag are being parsed +func TestByteSliceSupport(t *testing.T) { + src := withByteSlice{ + Field: []byte("[\"one\",\"two\",\"three\"]"), + } + dst := byteSliceTarget{} + assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) + assert.Equal(t, []string{"one", "two", "three"}, dst.Field) +} + +func TestUint8SliceSupport(t *testing.T) { + src := withUint8Slice{ + Field: []uint8("[\"one\",\"two\",\"three\"]"), + } + dst := uint8SliceTarget{} + assert.NoError(t, networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&src), reflect.ValueOf(&dst))) + assert.Equal(t, []string{"one", "two", "three"}, dst.Field) +} + type withNullString struct { Name sql.NullString } @@ -217,3 +236,19 @@ type emptyPublicIdTarget struct { PublicID string PublicId string } + +type withByteSlice struct { + Field []byte `nmap:"json"` +} + +type byteSliceTarget struct { + Field []string +} + +type withUint8Slice struct { + Field []byte `nmap:"json"` +} + +type uint8SliceTarget struct { + Field []string +} diff --git a/management/internals/network_map_db/sqlite/account_setting.go b/management/internals/network_map_db/sqlite/account_setting.go new file mode 100644 index 000000000..3ea7d5cf3 --- /dev/null +++ b/management/internals/network_map_db/sqlite/account_setting.go @@ -0,0 +1,65 @@ +package networkmap_sqlite + +import ( + "context" + "database/sql" + "reflect" + + networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db" + "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" +) + +const ( + GetAccountSettingsQuery = ` + select settings_peer_login_expiration_enabled as peer_login_expiration_enabled, + settings_peer_login_expiration as peer_login_expiration, + settings_peer_inactivity_expiration_enabled as peer_inactivity_expiration_enabled, + settings_peer_inactivity_expiration as peer_inactivity_expiration, + settings_dns_domain as dns_domain, + settings_ipv6_enabled_groups as ipv6_enabled_groups, + settings_routing_peer_dns_resolution_enabled as routing_peer_dns_resolution_enabled, + settings_lazy_connection_enabled as lazy_connection_enabled, + settings_auto_update_version as auto_update_version, + settings_auto_update_always as auto_update_always, + settings_metrics_push_enabled as metrics_push_enabled + from accounts + where id=$1 + ` +) + +func (sc *SqliteStoreConn) GetAccountSettings(ctx context.Context, accountId string) (nmdata.AccountSettingsInfo, error) { + rows, err := sc.Conn.QueryContext(ctx, GetAccountSettingsQuery, accountId) + if err != nil { + return nmdata.AccountSettingsInfo{}, err + } + defer rows.Close() + + rows.Next() + a := account{} + err = rows.Scan(networkmapdb.StructFields(reflect.ValueOf(&a))...) + if err != nil { + return nmdata.AccountSettingsInfo{}, err + } + + settingsInfo := nmdata.AccountSettingsInfo{} + err = networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&a), reflect.ValueOf(&settingsInfo)) + if err != nil { + return nmdata.AccountSettingsInfo{}, err + } + + return settingsInfo, nil +} + +type account struct { + PeerLoginExpirationEnabled sql.NullBool + PeerLoginExpiration sql.NullInt64 + PeerInactivityExpirationEnabled sql.NullBool + PeerInactivityExpiration sql.NullInt64 + DNSDomain sql.NullString + IPv6EnabledGroups []byte `nmap:"json"` + RoutingPeerDNSResolutionEnabled sql.NullBool + LazyConnectionEnabled sql.NullBool + AutoUpdateVersion sql.NullString + AutoUpdateAlways sql.NullBool + MetricsPushEnabled sql.NullBool +} diff --git a/management/internals/network_map_db/sqlite/sqlite_store.go b/management/internals/network_map_db/sqlite/sqlite_store.go index 8d5226d6e..bd25ef5ab 100644 --- a/management/internals/network_map_db/sqlite/sqlite_store.go +++ b/management/internals/network_map_db/sqlite/sqlite_store.go @@ -8,13 +8,26 @@ import ( "strings" "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" ) type SqliteStore struct { Db *sql.DB } -func NewSqliteStore(ctx context.Context, storeFile, dataDir, dsn string) (*SqliteStore, error) { +type sqliteInterface interface { + QueryContext(ctx context.Context, query string, args ...any) (*sql.Rows, error) + ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error) +} + +type SqliteStoreConn struct { + Conn sqliteInterface +} + +func NewSqliteStore(ctx context.Context, storeFile, dataDir string) (*SqliteStore, error) { // storeFile := storeSqliteFileName // if envFile, ok := os.LookupEnv("NB_STORE_ENGINE_SQLITE_FILE"); ok && envFile != "" { // storeFile = envFile @@ -51,9 +64,67 @@ func NewSqliteStore(ctx context.Context, storeFile, dataDir, dsn string) (*Sqlit connStr += "?" + strings.Join(parts, "&") } - db, err := sql.Open("sqlite3", "") + db, err := sql.Open("sqlite3", connStr) if err != nil { return nil, err } + return &SqliteStore{Db: db}, nil } + +func (s *SqliteStore) WithTx(tx *sql.Tx) *SqliteStoreConn { + return &SqliteStoreConn{Conn: tx} +} + +func (s *SqliteStore) UsingConn() *SqliteStoreConn { + return &SqliteStoreConn{Conn: s.Db} +} + +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 +} +func (s *SqliteStoreConn) GetPolicies(ctx context.Context, accountId string) ([]nmdata.Policy, map[string]map[string]any, map[string]map[string]any, error) { + return nil, nil, nil, nil +} +func (s *SqliteStoreConn) GetRoutes(ctx context.Context, accountId string) ([]nmdata.Route, error) { + return nil, nil +} +func (s *SqliteStoreConn) GetNameServerGroups(ctx context.Context, accountId string) ([]nmdata.NameServerGroup, error) { + return nil, nil +} +func (s *SqliteStoreConn) GetNetworkResources(ctx context.Context, accountId string) ([]nmdata.NetworkResource, error) { + return nil, nil +} +func (s *SqliteStoreConn) GetNetworkRouters(ctx context.Context, accountId string) (map[string]map[string]*nmdata.NetworkRouter, error) { + return nil, nil +} +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 +} +func (s *SqliteStoreConn) GetAllowedUsers(ctx context.Context, accountId string) (map[string]struct{}, map[string][]string, error) { + return nil, nil, nil +} +func (s *SqliteStoreConn) GetDnsSettings(ctx context.Context, accountId string) (nmdata.DNSSettings, error) { + return nmdata.DNSSettings{}, nil +} +func (s *SqliteStoreConn) GetNetworkXIDToPublicIdMap(ctx context.Context, accountId string) (map[string]string, error) { + return nil, nil +} +func (s *SqliteStoreConn) GetPrivateServices(ctx context.Context, accountId string) ([]networkmapdb.Service, error) { + return nil, nil +} +func (s *SqliteStoreConn) GetProxyTargetedDomainResourceIDs(ctx context.Context, accountId string) (map[string]struct{}, error) { + return nil, nil +}