Update store.go

This commit is contained in:
İsmail
2024-11-04 22:05:51 +03:00
parent 3fd180c173
commit c6116441be
+17 -7
View File
@@ -20,6 +20,7 @@ import (
"github.com/netbirdio/netbird/dns" "github.com/netbirdio/netbird/dns"
nbgroup "github.com/netbirdio/netbird/management/server/group" nbgroup "github.com/netbirdio/netbird/management/server/group"
"github.com/netbirdio/netbird/management/server/testutil"
"github.com/netbirdio/netbird/management/server/telemetry" "github.com/netbirdio/netbird/management/server/telemetry"
"github.com/netbirdio/netbird/util" "github.com/netbirdio/netbird/util"
@@ -27,7 +28,6 @@ import (
"github.com/netbirdio/netbird/management/server/migration" "github.com/netbirdio/netbird/management/server/migration"
nbpeer "github.com/netbirdio/netbird/management/server/peer" nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/server/posture" "github.com/netbirdio/netbird/management/server/posture"
"github.com/netbirdio/netbird/management/server/testutil"
"github.com/netbirdio/netbird/route" "github.com/netbirdio/netbird/route"
) )
@@ -284,12 +284,14 @@ func NewTestStoreFromSQL(ctx context.Context, filename string, dataDir string) (
if err != nil { if err != nil {
return nil, nil, fmt.Errorf("failed to create test store: %v", err) return nil, nil, fmt.Errorf("failed to create test store: %v", err)
} }
cleanUp := func() {
store.Close(ctx) return getSqlStoreEngine(ctx, store, kind)
} }
func getSqlStoreEngine(ctx context.Context, store *SqlStore, kind StoreEngine) (Store, func(), error) {
if kind == PostgresStoreEngine { if kind == PostgresStoreEngine {
cleanUp, err = testutil.CreatePGDB() cleanUp, err := testutil.CreatePGDB()
if err != nil { if err != nil {
return nil, nil, err return nil, nil, err
} }
@@ -303,10 +305,12 @@ func NewTestStoreFromSQL(ctx context.Context, filename string, dataDir string) (
if err != nil { if err != nil {
return nil, nil, err return nil, nil, err
} }
return store, cleanUp, nil
} }
if kind == MysqlStoreEngine { if kind == MysqlStoreEngine {
cleanUp, err = testutil.CreateMyDB() cleanUp, err := testutil.CreateMyDB()
if err != nil { if err != nil {
return nil, nil, err return nil, nil, err
} }
@@ -320,9 +324,15 @@ func NewTestStoreFromSQL(ctx context.Context, filename string, dataDir string) (
if err != nil { if err != nil {
return nil, nil, err return nil, nil, err
} }
return store, cleanUp, nil
} }
return store, cleanUp, nil closeConnection := func() {
store.Close(ctx)
}
return store, closeConnection, nil
} }
func loadSQL(db *gorm.DB, filepath string) error { func loadSQL(db *gorm.DB, filepath string) error {