extract repo

This commit is contained in:
Pascal Fischer
2025-06-11 23:24:35 +02:00
parent e662c5dd54
commit f612dba3d3
43 changed files with 49 additions and 4021 deletions
+9 -3
View File
@@ -7,14 +7,20 @@ import (
"github.com/spf13/cobra"
"github.com/netbirdio/management-integrations-refactor/integrations"
"github.com/netbirdio/management-refactor/internals/server"
"github.com/netbirdio/management-refactor/pkg/logging"
)
var log = logging.LoggerForThisPackage
var newServer = func() server.Server {
return server.NewServer()
}
func SetNewServer(fn func() server.Server) {
newServer = fn
}
// mgmtCmd starts the management server
var mgmtCmd = &cobra.Command{
Use: "management",
@@ -25,7 +31,7 @@ var mgmtCmd = &cobra.Command{
log().Debugf("Failed to init logging: %v", err)
}
srv := integrations.InitCloud(server.NewServer())
srv := newServer()
go func() {
log().Info("Starting server on :8080")
-1
View File
@@ -15,7 +15,6 @@ require (
github.com/grpc-ecosystem/go-grpc-middleware/v2 v2.0.2-0.20240212192251-757544f21357
github.com/hashicorp/go-secure-stdlib/base62 v0.1.2
github.com/mattn/go-sqlite3 v1.14.22
github.com/netbirdio/management-integrations-refactor/integrations v0.0.0-20250610135639-c2a4c2d5389e
github.com/netbirdio/management-integrations/integrations v0.0.0-20250330143713-7901e0a82203
github.com/netbirdio/netbird v0.41.0
github.com/petermattis/goid v0.0.0-20250319124200-ccd6737f222a
+2 -2
View File
@@ -222,8 +222,8 @@ github.com/moby/term v0.5.0 h1:xt8Q1nalod/v7BqbG21f8mQPqH+xAaC9C3N3wfWbVP0=
github.com/moby/term v0.5.0/go.mod h1:8FzsFHVUBGZdbDsJw/ot+X+d5HLUbvklYLJ9uGfcI3Y=
github.com/morikuni/aec v1.0.0 h1:nP9CBfwrvYnBRgY6qfDQkygYDmYwOilePFkwzv4dU8A=
github.com/morikuni/aec v1.0.0/go.mod h1:BbKIizmSmc5MMPqRYbxO4ZU0S0+P200+tUnFx7PXmsc=
github.com/netbirdio/management-integrations-refactor/integrations v0.0.0-20250610135639-c2a4c2d5389e h1:s5bU6r6sojx8lLCNY9AERx23DcaptZFZ3NHmUgPH/cg=
github.com/netbirdio/management-integrations-refactor/integrations v0.0.0-20250610135639-c2a4c2d5389e/go.mod h1:I3ZwjX7HGAUs8QdOMIAf155YS3G+A7y+/dtv8WB807Y=
github.com/netbirdio/management-refactor/integrations v0.0.0-20250610135639-c2a4c2d5389e h1:s5bU6r6sojx8lLCNY9AERx23DcaptZFZ3NHmUgPH/cg=
github.com/netbirdio/management-refactor/integrations v0.0.0-20250610135639-c2a4c2d5389e/go.mod h1:I3ZwjX7HGAUs8QdOMIAf155YS3G+A7y+/dtv8WB807Y=
github.com/netbirdio/management-integrations/integrations v0.0.0-20250330143713-7901e0a82203 h1:uxxbLPXQgC9VO15epNPtrD6zazyd5rZeqC5hQSmCdZU=
github.com/netbirdio/management-integrations/integrations v0.0.0-20250330143713-7901e0a82203/go.mod h1:2ZE6/tBBCKHQggPfO2UOQjyjXI7k+JDVl2ymorTOVQs=
github.com/netbirdio/netbird v0.41.0 h1:netVLMdYZyFGEOvzsCUcB8TrCgpPRJElPV9gNano2pg=
@@ -5,13 +5,11 @@ import (
"sync"
"time"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
log "github.com/sirupsen/logrus"
"github.com/netbirdio/management-refactor/management/server/activity"
nbpeer "github.com/netbirdio/management-refactor/management/server/peer"
"github.com/netbirdio/management-refactor/management/server/store"
"github.com/netbirdio/management-refactor/internals/modules/peers"
"github.com/netbirdio/management-refactor/internals/shared/db"
)
const (
@@ -77,7 +75,7 @@ func (e *Controller) Stop() {
// OnPeerConnected remove the peer from the linked list of ephemeral peers. Because it has been called when the peer
// is active the manager will not delete it while it is active.
func (e *Controller) OnPeerConnected(ctx context.Context, peer *nbpeer.Peer) {
func (e *Controller) OnPeerConnected(ctx context.Context, peer *peers.Peer) {
if !peer.Ephemeral {
return
}
@@ -121,7 +119,7 @@ func (e *Controller) OnPeerDisconnected(ctx context.Context, peer *nbpeer.Peer)
}
func (e *Controller) loadEphemeralPeers(ctx context.Context) {
peers, err := e.peersManager.GetAllEphemeralPeers(ctx, store.LockingStrengthShare)
peers, err := e.peersManager.GetAllEphemeralPeers(ctx, nil, db.LockingStrengthShare)
if err != nil {
log.WithContext(ctx).Debugf("failed to load ephemeral peers: %s", err)
return
@@ -165,7 +163,7 @@ func (e *Controller) cleanup(ctx context.Context) {
for id, p := range deletePeers {
log.WithContext(ctx).Debugf("delete ephemeral peer: %s", id)
err := e.peersManager.DeletePeer(ctx, p.accountID, id, activity.SystemInitiator)
err := e.peersManager.DeletePeer(ctx, nil, p.accountID, id)
if err != nil {
log.WithContext(ctx).Errorf("failed to delete ephemeral peer: %s", err)
}
@@ -1,149 +0,0 @@
package server
import (
"context"
"fmt"
"testing"
"time"
nbAccount "github.com/netbirdio/management-refactor/management/server/account"
nbpeer "github.com/netbirdio/management-refactor/management/server/peer"
"github.com/netbirdio/management-refactor/management/server/store"
"github.com/netbirdio/management-refactor/management/server/types"
)
type MockStore struct {
store.Store
account *types.Account
}
func (s *MockStore) GetAllEphemeralPeers(_ context.Context, _ store.LockingStrength) ([]*nbpeer.Peer, error) {
var peers []*nbpeer.Peer
for _, v := range s.account.Peers {
if v.Ephemeral {
peers = append(peers, v)
}
}
return peers, nil
}
type MocAccountManager struct {
nbAccount.Manager
store *MockStore
}
func (a MocAccountManager) DeletePeer(_ context.Context, accountID, peerID, userID string) error {
delete(a.store.account.Peers, peerID)
return nil //nolint:nil
}
func (a MocAccountManager) GetStore() store.Store {
return a.store
}
func TestNewManager(t *testing.T) {
startTime := time.Now()
timeNow = func() time.Time {
return startTime
}
store := &MockStore{}
am := MocAccountManager{
store: store,
}
numberOfPeers := 5
numberOfEphemeralPeers := 3
seedPeers(store, numberOfPeers, numberOfEphemeralPeers)
mgr := NewEphemeralManager(store, am)
mgr.loadEphemeralPeers(context.Background())
startTime = startTime.Add(ephemeralLifeTime + 1)
mgr.cleanup(context.Background())
if len(store.account.Peers) != numberOfPeers {
t.Errorf("failed to cleanup ephemeral peers, expected: %d, result: %d", numberOfPeers, len(store.account.Peers))
}
}
func TestNewManagerPeerConnected(t *testing.T) {
startTime := time.Now()
timeNow = func() time.Time {
return startTime
}
store := &MockStore{}
am := MocAccountManager{
store: store,
}
numberOfPeers := 5
numberOfEphemeralPeers := 3
seedPeers(store, numberOfPeers, numberOfEphemeralPeers)
mgr := NewEphemeralManager(store, am)
mgr.loadEphemeralPeers(context.Background())
mgr.OnPeerConnected(context.Background(), store.account.Peers["ephemeral_peer_0"])
startTime = startTime.Add(ephemeralLifeTime + 1)
mgr.cleanup(context.Background())
expected := numberOfPeers + 1
if len(store.account.Peers) != expected {
t.Errorf("failed to cleanup ephemeral peers, expected: %d, result: %d", expected, len(store.account.Peers))
}
}
func TestNewManagerPeerDisconnected(t *testing.T) {
startTime := time.Now()
timeNow = func() time.Time {
return startTime
}
store := &MockStore{}
am := MocAccountManager{
store: store,
}
numberOfPeers := 5
numberOfEphemeralPeers := 3
seedPeers(store, numberOfPeers, numberOfEphemeralPeers)
mgr := NewEphemeralManager(store, am)
mgr.loadEphemeralPeers(context.Background())
for _, v := range store.account.Peers {
mgr.OnPeerConnected(context.Background(), v)
}
mgr.OnPeerDisconnected(context.Background(), store.account.Peers["ephemeral_peer_0"])
startTime = startTime.Add(ephemeralLifeTime + 1)
mgr.cleanup(context.Background())
expected := numberOfPeers + numberOfEphemeralPeers - 1
if len(store.account.Peers) != expected {
t.Errorf("failed to cleanup ephemeral peers, expected: %d, result: %d", expected, len(store.account.Peers))
}
}
func seedPeers(store *MockStore, numberOfPeers int, numberOfEphemeralPeers int) {
store.account = newAccountWithId(context.Background(), "my account", "", "")
for i := 0; i < numberOfPeers; i++ {
peerId := fmt.Sprintf("peer_%d", i)
p := &nbpeer.Peer{
ID: peerId,
Ephemeral: false,
}
store.account.Peers[p.ID] = p
}
for i := 0; i < numberOfEphemeralPeers; i++ {
peerId := fmt.Sprintf("ephemeral_peer_%d", i)
p := &nbpeer.Peer{
ID: peerId,
Ephemeral: true,
}
store.account.Peers[p.ID] = p
}
}
-4
View File
@@ -4,7 +4,6 @@ import (
"context"
"github.com/netbirdio/management-refactor/internals/shared/db"
"github.com/netbirdio/management-refactor/internals/shared/hook"
)
type Manager interface {
@@ -13,7 +12,4 @@ type Manager interface {
CreateNetwork(ctx context.Context, tx db.Transaction, userID string, network *Network) (*Network, error)
UpdateNetwork(ctx context.Context, tx db.Transaction, userID string, network *Network) (*Network, error)
DeleteNetwork(ctx context.Context, tx db.Transaction, accountID, userID, networkID string) error
// events
OnNetworkDelete() *hook.Hook[*NetworkEvent]
}
@@ -13,13 +13,11 @@ import (
type managerImpl struct {
repo Repository
onNetworkDelete *hook.Hook[*networks.NetworkEvent]
}
func NewManager(store *db.Store) networks.Manager {
func NewManager(repo Repository) networks.Manager {
return &managerImpl{
repo: newRepository(store),
repo: repo,
}
}
@@ -18,7 +18,7 @@ type repository struct {
store *db.Store
}
func newRepository(s *db.Store) Repository {
func NewRepository(s *db.Store) Repository {
return &repository{store: s}
}
@@ -1,15 +0,0 @@
package networks
import (
"context"
"github.com/netbirdio/management-refactor/internals/shared/db"
"github.com/netbirdio/management-refactor/internals/shared/hook"
)
type NetworkEvent struct {
hook.Event
Context context.Context
Tx db.Transaction
Network *Network
}
@@ -16,24 +16,11 @@ type managerImpl struct {
networkManager networks.Manager
}
func NewManager(store *db.Store, router *mux.Router, networkManager networks.Manager) resources.Manager {
repo := newRepository(store)
m := &managerImpl{
func NewManager(repo Repository, router *mux.Router, networkManager networks.Manager) resources.Manager {
return &managerImpl{
repo: repo,
networkManager: networkManager,
}
networkManager.OnNetworkDelete().BindFunc(func(e *networks.NetworkEvent) error {
if err := m.DeleteResourcesInNetwork(e.Context, e.Tx, e.Network); err != nil {
return fmt.Errorf("failed to delete resources in network: %w", err)
}
return e.Next()
})
// api := newHandler(m, permissionsManager)
// api.RegisterEndpoints(router)
return m
}
func (m *managerImpl) GetNetworkResourcesByNetID(ctx context.Context, tx db.Transaction, lockingStrength db.LockingStrength, network *networks.Network) ([]*resources.NetworkResource, error) {
@@ -15,7 +15,7 @@ type repository struct {
store *db.Store
}
func newRepository(s *db.Store) Repository {
func NewRepository(s *db.Store) Repository {
return &repository{store: s}
}
+2
View File
@@ -13,4 +13,6 @@ type Manager interface {
GetPeers(ctx context.Context, tx db.Transaction, strength db.LockingStrength, accountID string) ([]*Peer, error)
GetFilteredPeers(ctx context.Context, tx db.Transaction, strength db.LockingStrength, accountID, nameFilter, ipFilter string) ([]*Peer, error)
UpdatePeer(ctx context.Context, tx db.Transaction, peer *Peer) error
GetAllEphemeralPeers(ctx context.Context, tx db.Transaction, strength db.LockingStrength) ([]*Peer, error)
DeletePeer(ctx context.Context, tx db.Transaction, accountID, peerID string) error
}
+12 -2
View File
@@ -20,8 +20,8 @@ type Manager struct {
networkMapController network_map.Controller
}
func NewManager(store *db.Store) *Manager {
return &Manager{repo: newRepository(store)}
func NewManager(repo Repository) *Manager {
return &Manager{repo: repo}
}
func (m *Manager) SetNetworkMapController(networkMapController network_map.Controller) {
@@ -51,3 +51,13 @@ func (m *Manager) UpdatePeer(ctx context.Context, tx db.Transaction, peer *peers
return nil
}
func (m *Manager) GetAllEphemeralPeers(ctx context.Context, tx db.Transaction, strength db.LockingStrength) ([]*peers.Peer, error) {
// TODO implement me
panic("implement me")
}
func (m *Manager) DeletePeer(ctx context.Context, tx db.Transaction, accountID, peerID string) error {
// TODO implement me
panic("implement me")
}
@@ -17,7 +17,7 @@ type repository struct {
store *db.Store
}
func newRepository(s *db.Store) Repository {
func NewRepository(s *db.Store) Repository {
return &repository{store: s}
}
+2 -4
View File
@@ -14,10 +14,8 @@ type Manager struct {
repo Repository
}
func NewManager(store *db.Store) *Manager {
repo := newRepository(store)
m := &Manager{repo: repo}
return m
func NewManager(repo Repository) *Manager {
return &Manager{repo: repo}
}
func (m *Manager) GetAllUsers(ctx context.Context, tx db.Transaction, strength db.LockingStrength, accountID string) ([]users.User, error) {
@@ -16,7 +16,7 @@ type repository struct {
store *db.Store
}
func newRepository(s *db.Store) Repository {
func NewRepository(s *db.Store) Repository {
err := s.AutoMigrate(users.User{})
if err != nil {
log.Fatalf("Failed to auto migrate: %v", err)
+8 -4
View File
@@ -14,13 +14,15 @@ import (
func (s *BaseServer) NetworksManager() networks.Manager {
return Create(s, func() networks.Manager {
return manager.NewManager(s.Store())
repo := manager.NewRepository(s.Store())
return manager.NewManager(repo)
})
}
func (s *BaseServer) ResourcesManager() resources.Manager {
return Create(s, func() resources.Manager {
manager := resourcesManager.NewManager(s.Store(), s.Router(), s.NetworksManager())
repo := resourcesManager.NewRepository(s.Store())
manager := resourcesManager.NewManager(repo, s.Router(), s.NetworksManager())
return manager
})
}
@@ -33,7 +35,8 @@ func (s *BaseServer) PermissionsManager() permissions.Manager {
func (s *BaseServer) PeersManager() peers.Manager {
return Create(s, func() peers.Manager {
manager := peersManager.NewManager(s.Store())
repo := peersManager.NewRepository(s.Store())
manager := peersManager.NewManager(repo)
s.AfterInit(func(s *BaseServer) {
peersManager.RegisterEndpoints(s.Router(), s.PermissionsManager(), manager)
manager.SetNetworkMapController(s.NetworkMapController())
@@ -44,7 +47,8 @@ func (s *BaseServer) PeersManager() peers.Manager {
func (s *BaseServer) UsersManager() users.Manager {
return Create(s, func() users.Manager {
manager := usersManager.NewManager(s.Store())
repo := usersManager.NewRepository(s.Store())
manager := usersManager.NewManager(repo)
s.AfterInit(func(s *BaseServer) {
usersManager.RegisterEndpoints(s.Router(), s.PermissionsManager(), manager)
})
-1
View File
@@ -28,7 +28,6 @@ var log = logging.LoggerForThisPackage()
// NewServer initializes and configures a new Server instance
func NewServer() *BaseServer {
return &BaseServer{
// @todo shared config
container: make(map[string]any),
}
}
@@ -1,310 +0,0 @@
package sqlite
import (
"bytes"
"testing"
)
func TestGenerateKey(t *testing.T) {
testData := "exampl@netbird.io"
key, err := GenerateKey()
if err != nil {
t.Fatalf("failed to generate key: %s", err)
}
ee, err := NewFieldEncrypt(key)
if err != nil {
t.Fatalf("failed to init email encryption: %s", err)
}
encrypted, err := ee.Encrypt(testData)
if err != nil {
t.Fatalf("failed to encrypt data: %s", err)
}
if encrypted == "" {
t.Fatalf("invalid encrypted text")
}
decrypted, err := ee.Decrypt(encrypted)
if err != nil {
t.Fatalf("failed to decrypt data: %s", err)
}
if decrypted != testData {
t.Fatalf("decrypted data is not match with test data: %s, %s", testData, decrypted)
}
}
func TestGenerateKeyLegacy(t *testing.T) {
testData := "exampl@netbird.io"
key, err := GenerateKey()
if err != nil {
t.Fatalf("failed to generate key: %s", err)
}
ee, err := NewFieldEncrypt(key)
if err != nil {
t.Fatalf("failed to init email encryption: %s", err)
}
encrypted := ee.LegacyEncrypt(testData)
if encrypted == "" {
t.Fatalf("invalid encrypted text")
}
decrypted, err := ee.LegacyDecrypt(encrypted)
if err != nil {
t.Fatalf("failed to decrypt data: %s", err)
}
if decrypted != testData {
t.Fatalf("decrypted data is not match with test data: %s, %s", testData, decrypted)
}
}
func TestCorruptKey(t *testing.T) {
testData := "exampl@netbird.io"
key, err := GenerateKey()
if err != nil {
t.Fatalf("failed to generate key: %s", err)
}
ee, err := NewFieldEncrypt(key)
if err != nil {
t.Fatalf("failed to init email encryption: %s", err)
}
encrypted, err := ee.Encrypt(testData)
if err != nil {
t.Fatalf("failed to encrypt data: %s", err)
}
if encrypted == "" {
t.Fatalf("invalid encrypted text")
}
newKey, err := GenerateKey()
if err != nil {
t.Fatalf("failed to generate key: %s", err)
}
ee, err = NewFieldEncrypt(newKey)
if err != nil {
t.Fatalf("failed to init email encryption: %s", err)
}
res, _ := ee.Decrypt(encrypted)
if res == testData {
t.Fatalf("incorrect decryption, the result is: %s", res)
}
}
func TestEncryptDecrypt(t *testing.T) {
// Generate a key for encryption/decryption
key, err := GenerateKey()
if err != nil {
t.Fatalf("Failed to generate key: %v", err)
}
// Initialize the FieldEncrypt with the generated key
ec, err := NewFieldEncrypt(key)
if err != nil {
t.Fatalf("Failed to create FieldEncrypt: %v", err)
}
// Test cases
testCases := []struct {
name string
input string
}{
{
name: "Empty String",
input: "",
},
{
name: "Short String",
input: "Hello",
},
{
name: "String with Spaces",
input: "Hello, World!",
},
{
name: "Long String",
input: "The quick brown fox jumps over the lazy dog.",
},
{
name: "Unicode Characters",
input: "こんにちは世界",
},
{
name: "Special Characters",
input: "!@#$%^&*()_+-=[]{}|;':\",./<>?",
},
{
name: "Numeric String",
input: "1234567890",
},
{
name: "Repeated Characters",
input: "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa",
},
{
name: "Multi-block String",
input: "This is a longer string that will span multiple blocks in the encryption algorithm.",
},
{
name: "Non-ASCII and ASCII Mix",
input: "Hello 世界 123",
},
}
for _, tc := range testCases {
t.Run(tc.name+" - Legacy", func(t *testing.T) {
// Legacy Encryption
encryptedLegacy := ec.LegacyEncrypt(tc.input)
if encryptedLegacy == "" {
t.Errorf("LegacyEncrypt returned empty string for input '%s'", tc.input)
}
// Legacy Decryption
decryptedLegacy, err := ec.LegacyDecrypt(encryptedLegacy)
if err != nil {
t.Errorf("LegacyDecrypt failed for input '%s': %v", tc.input, err)
}
// Verify that the decrypted value matches the original input
if decryptedLegacy != tc.input {
t.Errorf("LegacyDecrypt output '%s' does not match original input '%s'", decryptedLegacy, tc.input)
}
})
t.Run(tc.name+" - New", func(t *testing.T) {
// New Encryption
encryptedNew, err := ec.Encrypt(tc.input)
if err != nil {
t.Errorf("Encrypt failed for input '%s': %v", tc.input, err)
}
if encryptedNew == "" {
t.Errorf("Encrypt returned empty string for input '%s'", tc.input)
}
// New Decryption
decryptedNew, err := ec.Decrypt(encryptedNew)
if err != nil {
t.Errorf("Decrypt failed for input '%s': %v", tc.input, err)
}
// Verify that the decrypted value matches the original input
if decryptedNew != tc.input {
t.Errorf("Decrypt output '%s' does not match original input '%s'", decryptedNew, tc.input)
}
})
}
}
func TestPKCS5UnPadding(t *testing.T) {
tests := []struct {
name string
input []byte
expected []byte
expectError bool
}{
{
name: "Valid Padding",
input: append([]byte("Hello, World!"), bytes.Repeat([]byte{4}, 4)...),
expected: []byte("Hello, World!"),
},
{
name: "Empty Input",
input: []byte{},
expectError: true,
},
{
name: "Padding Length Zero",
input: append([]byte("Hello, World!"), bytes.Repeat([]byte{0}, 4)...),
expectError: true,
},
{
name: "Padding Length Exceeds Block Size",
input: append([]byte("Hello, World!"), bytes.Repeat([]byte{17}, 17)...),
expectError: true,
},
{
name: "Padding Length Exceeds Input Length",
input: []byte{5, 5, 5},
expectError: true,
},
{
name: "Invalid Padding Bytes",
input: append([]byte("Hello, World!"), []byte{2, 3, 4, 5}...),
expectError: true,
},
{
name: "Valid Single Byte Padding",
input: append([]byte("Hello, World!"), byte(1)),
expected: []byte("Hello, World!"),
},
{
name: "Invalid Mixed Padding Bytes",
input: append([]byte("Hello, World!"), []byte{3, 3, 2}...),
expectError: true,
},
{
name: "Valid Full Block Padding",
input: append([]byte("Hello, World!"), bytes.Repeat([]byte{16}, 16)...),
expected: []byte("Hello, World!"),
},
{
name: "Non-Padding Byte at End",
input: append([]byte("Hello, World!"), []byte{4, 4, 4, 5}...),
expectError: true,
},
{
name: "Valid Padding with Different Text Length",
input: append([]byte("Test"), bytes.Repeat([]byte{12}, 12)...),
expected: []byte("Test"),
},
{
name: "Padding Length Equal to Input Length",
input: bytes.Repeat([]byte{8}, 8),
expected: []byte{},
},
{
name: "Invalid Padding Length Zero (Again)",
input: append([]byte("Test"), byte(0)),
expectError: true,
},
{
name: "Padding Length Greater Than Input",
input: []byte{10},
expectError: true,
},
{
name: "Input Length Not Multiple of Block Size",
input: append([]byte("Invalid Length"), byte(1)),
expected: []byte("Invalid Length"),
},
{
name: "Valid Padding with Non-ASCII Characters",
input: append([]byte("こんにちは"), bytes.Repeat([]byte{2}, 2)...),
expected: []byte("こんにちは"),
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
result, err := pkcs5UnPadding(tt.input)
if tt.expectError {
if err == nil {
t.Errorf("Expected error but got nil")
}
} else {
if err != nil {
t.Errorf("Did not expect error but got: %v", err)
}
if !bytes.Equal(result, tt.expected) {
t.Errorf("Expected output %v, got %v", tt.expected, result)
}
}
})
}
}
@@ -1,84 +0,0 @@
package sqlite
import (
"context"
"database/sql"
"path/filepath"
"testing"
"time"
_ "github.com/mattn/go-sqlite3"
"github.com/netbirdio/management-refactor/management/server/activity"
"github.com/stretchr/testify/require"
)
func setupDatabase(t *testing.T) *sql.DB {
t.Helper()
dbFile := filepath.Join(t.TempDir(), eventSinkDB)
db, err := sql.Open("sqlite3", dbFile)
require.NoError(t, err, "Failed to open database")
t.Cleanup(func() {
_ = db.Close()
})
_, err = db.Exec(createTableQuery)
require.NoError(t, err, "Failed to create events table")
_, err = db.Exec(`CREATE TABLE deleted_users (id TEXT NOT NULL, email TEXT NOT NULL, name TEXT);`)
require.NoError(t, err, "Failed to create deleted_users table")
return db
}
func TestMigrate(t *testing.T) {
db := setupDatabase(t)
key, err := GenerateKey()
require.NoError(t, err, "Failed to generate key")
crypt, err := NewFieldEncrypt(key)
require.NoError(t, err, "Failed to initialize FieldEncrypt")
legacyEmail := crypt.LegacyEncrypt("testaccount@test.com")
legacyName := crypt.LegacyEncrypt("Test Account")
_, err = db.Exec(`INSERT INTO events(activity, timestamp, initiator_id, target_id, account_id, meta) VALUES(?, ?, ?, ?, ?, ?)`,
activity.UserDeleted, time.Now(), "initiatorID", "targetID", "accountID", "")
require.NoError(t, err, "Failed to insert event")
_, err = db.Exec(`INSERT INTO deleted_users(id, email, name) VALUES(?, ?, ?)`, "targetID", legacyEmail, legacyName)
require.NoError(t, err, "Failed to insert legacy encrypted data")
colExists, err := checkColumnExists(db, "deleted_users", "enc_algo")
require.NoError(t, err, "Failed to check if enc_algo column exists")
require.False(t, colExists, "enc_algo column should not exist before migration")
err = migrate(context.Background(), crypt, db)
require.NoError(t, err, "Migration failed")
colExists, err = checkColumnExists(db, "deleted_users", "enc_algo")
require.NoError(t, err, "Failed to check if enc_algo column exists after migration")
require.True(t, colExists, "enc_algo column should exist after migration")
var encAlgo string
err = db.QueryRow(`SELECT enc_algo FROM deleted_users LIMIT 1`, "").Scan(&encAlgo)
require.NoError(t, err, "Failed to select updated data")
require.Equal(t, gcmEncAlgo, encAlgo, "enc_algo should be set to 'GCM' after migration")
store, err := createStore(crypt, db)
require.NoError(t, err, "Failed to create store")
events, err := store.Get(context.Background(), "accountID", 0, 1, false)
require.NoError(t, err, "Failed to get events")
require.Len(t, events, 1, "Should have one event")
require.Equal(t, activity.UserDeleted, events[0].Activity, "activity should match")
require.Equal(t, "initiatorID", events[0].InitiatorID, "initiator id should match")
require.Equal(t, "targetID", events[0].TargetID, "target id should match")
require.Equal(t, "accountID", events[0].AccountID, "account id should match")
require.Equal(t, "testaccount@test.com", events[0].Meta["email"], "email should match")
require.Equal(t, "Test Account", events[0].Meta["username"], "username should match")
}

Some files were not shown because too many files have changed in this diff Show More