diff --git a/cmd/management.go b/cmd/management.go index 077c70b..f5115a7 100644 --- a/cmd/management.go +++ b/cmd/management.go @@ -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") diff --git a/go.mod b/go.mod index ee08ad0..0559b83 100644 --- a/go.mod +++ b/go.mod @@ -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 diff --git a/go.sum b/go.sum index 780bc7b..220a621 100644 --- a/go.sum +++ b/go.sum @@ -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= diff --git a/internals/controllers/ephemeral_peers/controller.go b/internals/controllers/ephemeral_peers/controller.go index 46201a4..bf85c43 100644 --- a/internals/controllers/ephemeral_peers/controller.go +++ b/internals/controllers/ephemeral_peers/controller.go @@ -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) } diff --git a/internals/controllers/ephemeral_peers/controller_test.go b/internals/controllers/ephemeral_peers/controller_test.go deleted file mode 100644 index 818719c..0000000 --- a/internals/controllers/ephemeral_peers/controller_test.go +++ /dev/null @@ -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 - } -} diff --git a/internals/modules/networks/interface.go b/internals/modules/networks/interface.go index e9f5485..5e7d3ae 100644 --- a/internals/modules/networks/interface.go +++ b/internals/modules/networks/interface.go @@ -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] } diff --git a/internals/modules/networks/manager/manager.go b/internals/modules/networks/manager/manager.go index 5d14a3c..f26a226 100644 --- a/internals/modules/networks/manager/manager.go +++ b/internals/modules/networks/manager/manager.go @@ -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, } } diff --git a/internals/modules/networks/manager/repository.go b/internals/modules/networks/manager/repository.go index d6c0e3e..cc282a4 100644 --- a/internals/modules/networks/manager/repository.go +++ b/internals/modules/networks/manager/repository.go @@ -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} } diff --git a/internals/modules/networks/networkevent.go b/internals/modules/networks/networkevent.go deleted file mode 100644 index aa465e6..0000000 --- a/internals/modules/networks/networkevent.go +++ /dev/null @@ -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 -} diff --git a/internals/modules/networks/resources/manager/manager.go b/internals/modules/networks/resources/manager/manager.go index 777487b..d1eb250 100644 --- a/internals/modules/networks/resources/manager/manager.go +++ b/internals/modules/networks/resources/manager/manager.go @@ -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) { diff --git a/internals/modules/networks/resources/manager/repository.go b/internals/modules/networks/resources/manager/repository.go index 62b180c..8948db7 100644 --- a/internals/modules/networks/resources/manager/repository.go +++ b/internals/modules/networks/resources/manager/repository.go @@ -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} } diff --git a/internals/modules/peers/interface.go b/internals/modules/peers/interface.go index e1719a0..b0f560d 100644 --- a/internals/modules/peers/interface.go +++ b/internals/modules/peers/interface.go @@ -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 } diff --git a/internals/modules/peers/manager/manager.go b/internals/modules/peers/manager/manager.go index 301d98f..bc3a3eb 100644 --- a/internals/modules/peers/manager/manager.go +++ b/internals/modules/peers/manager/manager.go @@ -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") +} diff --git a/internals/modules/peers/manager/repository.go b/internals/modules/peers/manager/repository.go index 73ed575..3aa63a6 100644 --- a/internals/modules/peers/manager/repository.go +++ b/internals/modules/peers/manager/repository.go @@ -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} } diff --git a/internals/modules/users/manager/manager.go b/internals/modules/users/manager/manager.go index 5d89c10..08c2c80 100644 --- a/internals/modules/users/manager/manager.go +++ b/internals/modules/users/manager/manager.go @@ -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) { diff --git a/internals/modules/users/manager/repository.go b/internals/modules/users/manager/repository.go index c91491c..14d2c6a 100644 --- a/internals/modules/users/manager/repository.go +++ b/internals/modules/users/manager/repository.go @@ -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) diff --git a/internals/server/modules.go b/internals/server/modules.go index 313cd27..0c539b0 100644 --- a/internals/server/modules.go +++ b/internals/server/modules.go @@ -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) }) diff --git a/internals/server/server.go b/internals/server/server.go index 0356436..c552dc5 100644 --- a/internals/server/server.go +++ b/internals/server/server.go @@ -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), } } diff --git a/internals/shared/activity/sqlite/crypt_test.go b/internals/shared/activity/sqlite/crypt_test.go deleted file mode 100644 index aff3a08..0000000 --- a/internals/shared/activity/sqlite/crypt_test.go +++ /dev/null @@ -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) - } - } - }) - } -} diff --git a/internals/shared/activity/sqlite/migration_test.go b/internals/shared/activity/sqlite/migration_test.go deleted file mode 100644 index c157e04..0000000 --- a/internals/shared/activity/sqlite/migration_test.go +++ /dev/null @@ -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") -} diff --git a/internals/shared/activity/sqlite/sqlite_test.go b/internals/shared/activity/sqlite/sqlite_test.go deleted file mode 100644 index 76e181e..0000000 --- a/internals/shared/activity/sqlite/sqlite_test.go +++ /dev/null @@ -1,57 +0,0 @@ -package sqlite - -import ( - "context" - "fmt" - "testing" - "time" - - "github.com/stretchr/testify/assert" - - "github.com/netbirdio/management-refactor/management/server/activity" -) - -func TestNewSQLiteStore(t *testing.T) { - dataDir := t.TempDir() - key, _ := GenerateKey() - store, err := NewSQLiteStore(context.Background(), dataDir, key) - if err != nil { - t.Fatal(err) - return - } - defer store.Close(context.Background()) //nolint - - accountID := "account_1" - - for i := 0; i < 10; i++ { - _, err = store.Save(context.Background(), &activity.Event{ - Timestamp: time.Now().UTC(), - Activity: activity.PeerAddedByUser, - InitiatorID: "user_" + fmt.Sprint(i), - TargetID: "peer_" + fmt.Sprint(i), - AccountID: accountID, - }) - if err != nil { - t.Fatal(err) - return - } - } - - result, err := store.Get(context.Background(), accountID, 0, 10, false) - if err != nil { - t.Fatal(err) - return - } - - assert.Len(t, result, 10) - assert.True(t, result[0].Timestamp.Before(result[len(result)-1].Timestamp)) - - result, err = store.Get(context.Background(), accountID, 0, 5, true) - if err != nil { - t.Fatal(err) - return - } - - assert.Len(t, result, 5) - assert.True(t, result[0].Timestamp.After(result[len(result)-1].Timestamp)) -} diff --git a/internals/shared/api/grpc/grpcserver.go b/internals/shared/api/grpc/grpcserver.go deleted file mode 100644 index 432e7b5..0000000 --- a/internals/shared/api/grpc/grpcserver.go +++ /dev/null @@ -1,902 +0,0 @@ -package server - -import ( - "context" - "fmt" - "net" - "net/netip" - "strings" - "sync" - "time" - - pb "github.com/golang/protobuf/proto" // nolint - "github.com/golang/protobuf/ptypes/timestamp" - "github.com/grpc-ecosystem/go-grpc-middleware/v2/interceptors/realip" - log "github.com/sirupsen/logrus" - "golang.zx2c4.com/wireguard/wgctrl/wgtypes" - "google.golang.org/grpc/codes" - "google.golang.org/grpc/peer" - "google.golang.org/grpc/status" - - integrationsConfig "github.com/netbirdio/management-integrations/integrations/config" - - "github.com/netbirdio/management-refactor/encryption" - "github.com/netbirdio/management-refactor/management/proto" - "github.com/netbirdio/management-refactor/management/server/account" - "github.com/netbirdio/management-refactor/management/server/activity" - "github.com/netbirdio/management-refactor/management/server/auth" - nbContext "github.com/netbirdio/management-refactor/management/server/context" - nbpeer "github.com/netbirdio/management-refactor/management/server/peer" - "github.com/netbirdio/management-refactor/management/server/posture" - "github.com/netbirdio/management-refactor/management/server/settings" - internalStatus "github.com/netbirdio/management-refactor/management/server/status" - "github.com/netbirdio/management-refactor/management/server/telemetry" - "github.com/netbirdio/management-refactor/management/server/types" -) - -// GRPCServer an instance of a Management gRPC API server -type GRPCServer struct { - // accountManager account.Manager - settingsManager settings.Manager - wgKey wgtypes.Key - proto.UnimplementedManagementServiceServer - updateChannel *UpdateChannel - config *types.Config - secretsManager SecretsManager - appMetrics telemetry.AppMetrics - ephemeralManager *EphemeralManager - peerLocks sync.Map - authManager auth.Manager -} - -// NewServer creates a new Management server -func NewServer( - ctx context.Context, - config *types.Config, - accountManager account.Manager, - settingsManager settings.Manager, - updateChannel *UpdateChannel, - secretsManager SecretsManager, - appMetrics telemetry.AppMetrics, - ephemeralManager *EphemeralManager, - authManager auth.Manager, -) (*GRPCServer, error) { - key, err := wgtypes.GeneratePrivateKey() - if err != nil { - return nil, err - } - - if appMetrics != nil { - // update gauge based on number of connected peers which is equal to open gRPC streams - err = appMetrics.GRPCMetrics().RegisterConnectedStreams(func() int64 { - return int64(len(updateChannel.peerChannels)) - }) - if err != nil { - return nil, err - } - } - - return &GRPCServer{ - wgKey: key, - // peerKey -> event channel - updateChannel: updateChannel, - accountManager: accountManager, - settingsManager: settingsManager, - config: config, - secretsManager: secretsManager, - authManager: authManager, - appMetrics: appMetrics, - ephemeralManager: ephemeralManager, - }, nil -} - -func (s *GRPCServer) GetServerKey(ctx context.Context, req *proto.Empty) (*proto.ServerKeyResponse, error) { - ip := "" - p, ok := peer.FromContext(ctx) - if ok { - ip = p.Addr.String() - } - - log.WithContext(ctx).Tracef("GetServerKey request from %s", ip) - start := time.Now() - defer func() { - log.WithContext(ctx).Tracef("GetServerKey from %s took %v", ip, time.Since(start)) - }() - - // todo introduce something more meaningful with the key expiration/rotation - if s.appMetrics != nil { - s.appMetrics.GRPCMetrics().CountGetKeyRequest() - } - now := time.Now().Add(24 * time.Hour) - secs := int64(now.Second()) - nanos := int32(now.Nanosecond()) - expiresAt := ×tamp.Timestamp{Seconds: secs, Nanos: nanos} - - return &proto.ServerKeyResponse{ - Key: s.wgKey.PublicKey().String(), - ExpiresAt: expiresAt, - }, nil -} - -func getRealIP(ctx context.Context) net.IP { - if addr, ok := realip.FromContext(ctx); ok { - return net.IP(addr.AsSlice()) - } - return nil -} - -// Sync validates the existence of a connecting peer, sends an initial state (all available for the connecting peers) and -// notifies the connected peer of any updates (e.g. new peers under the same account) -func (s *GRPCServer) Sync(req *proto.EncryptedMessage, srv proto.ManagementService_SyncServer) error { - reqStart := time.Now() - if s.appMetrics != nil { - s.appMetrics.GRPCMetrics().CountSyncRequest() - } - - ctx := srv.Context() - - syncReq := &proto.SyncRequest{} - peerKey, err := s.parseRequest(ctx, req, syncReq) - if err != nil { - return err - } - - // nolint:staticcheck - ctx = context.WithValue(ctx, nbContext.PeerIDKey, peerKey.String()) - - unlock := s.acquirePeerLockByUID(ctx, peerKey.String()) - defer func() { - if unlock != nil { - unlock() - } - }() - - accountID, err := s.accountManager.GetAccountIDForPeerKey(ctx, peerKey.String()) - if err != nil { - // nolint:staticcheck - ctx = context.WithValue(ctx, nbContext.AccountIDKey, "UNKNOWN") - log.WithContext(ctx).Tracef("peer %s is not registered", peerKey.String()) - if errStatus, ok := internalStatus.FromError(err); ok && errStatus.Type() == internalStatus.NotFound { - return status.Errorf(codes.PermissionDenied, "peer is not registered") - } - return err - } - - // nolint:staticcheck - ctx = context.WithValue(ctx, nbContext.AccountIDKey, accountID) - - realIP := getRealIP(ctx) - log.WithContext(ctx).Debugf("Sync request from peer [%s] [%s]", req.WgPubKey, realIP.String()) - - if syncReq.GetMeta() == nil { - log.WithContext(ctx).Tracef("peer system meta has to be provided on sync. Peer %s, remote addr %s", peerKey.String(), realIP) - } - - peer, netMap, postureChecks, err := s.accountManager.SyncAndMarkPeer(ctx, accountID, peerKey.String(), extractPeerMeta(ctx, syncReq.GetMeta()), realIP) - if err != nil { - log.WithContext(ctx).Debugf("error while syncing peer %s: %v", peerKey.String(), err) - return mapError(ctx, err) - } - - err = s.sendInitialSync(ctx, peerKey, peer, netMap, postureChecks, srv) - if err != nil { - log.WithContext(ctx).Debugf("error while sending initial sync for %s: %v", peerKey.String(), err) - return err - } - - updates := s.updateChannel.CreateChannel(ctx, peer.ID) - - s.ephemeralManager.OnPeerConnected(ctx, peer) - - s.secretsManager.SetupRefresh(ctx, accountID, peer.ID) - - if s.appMetrics != nil { - s.appMetrics.GRPCMetrics().CountSyncRequestDuration(time.Since(reqStart)) - } - - unlock() - unlock = nil - - log.WithContext(ctx).Debugf("Sync: took %v", time.Since(reqStart)) - - return s.handleUpdates(ctx, accountID, peerKey, peer, updates, srv) -} - -// handleUpdates sends updates to the connected peer until the updates channel is closed. -func (s *GRPCServer) handleUpdates(ctx context.Context, accountID string, peerKey wgtypes.Key, peer *nbpeer.Peer, updates chan *UpdateMessage, srv proto.ManagementService_SyncServer) error { - log.WithContext(ctx).Tracef("starting to handle updates for peer %s", peerKey.String()) - for { - select { - // condition when there are some updates - case update, open := <-updates: - if s.appMetrics != nil { - s.appMetrics.GRPCMetrics().UpdateChannelQueueLength(len(updates) + 1) - } - - if !open { - log.WithContext(ctx).Debugf("updates channel for peer %s was closed", peerKey.String()) - s.cancelPeerRoutines(ctx, accountID, peer) - return nil - } - log.WithContext(ctx).Debugf("received an update for peer %s", peerKey.String()) - - if err := s.sendUpdate(ctx, accountID, peerKey, peer, update, srv); err != nil { - return err - } - - // condition when client <-> server connection has been terminated - case <-srv.Context().Done(): - // happens when connection drops, e.g. client disconnects - log.WithContext(ctx).Debugf("stream of peer %s has been closed", peerKey.String()) - s.cancelPeerRoutines(ctx, accountID, peer) - return srv.Context().Err() - } - } -} - -// sendUpdate encrypts the update message using the peer key and the server's wireguard key, -// then sends the encrypted message to the connected peer via the sync server. -func (s *GRPCServer) sendUpdate(ctx context.Context, accountID string, peerKey wgtypes.Key, peer *nbpeer.Peer, update *UpdateMessage, srv proto.ManagementService_SyncServer) error { - encryptedResp, err := encryption.EncryptMessage(peerKey, s.wgKey, update.Update) - if err != nil { - s.cancelPeerRoutines(ctx, accountID, peer) - return status.Errorf(codes.Internal, "failed processing update message") - } - err = srv.SendMsg(&proto.EncryptedMessage{ - WgPubKey: s.wgKey.PublicKey().String(), - Body: encryptedResp, - }) - if err != nil { - s.cancelPeerRoutines(ctx, accountID, peer) - return status.Errorf(codes.Internal, "failed sending update message") - } - log.WithContext(ctx).Debugf("sent an update to peer %s", peerKey.String()) - return nil -} - -func (s *GRPCServer) cancelPeerRoutines(ctx context.Context, accountID string, peer *nbpeer.Peer) { - unlock := s.acquirePeerLockByUID(ctx, peer.Key) - defer unlock() - - err := s.accountManager.OnPeerDisconnected(ctx, accountID, peer.Key) - if err != nil { - log.WithContext(ctx).Errorf("failed to disconnect peer %s properly: %v", peer.Key, err) - } - s.updateChannel.CloseChannel(ctx, peer.ID) - s.secretsManager.CancelRefresh(peer.ID) - s.ephemeralManager.OnPeerDisconnected(ctx, peer) - - log.WithContext(ctx).Tracef("peer %s has been disconnected", peer.Key) -} - -func (s *GRPCServer) validateToken(ctx context.Context, jwtToken string) (string, error) { - if s.authManager == nil { - return "", status.Errorf(codes.Internal, "missing auth manager") - } - - userAuth, token, err := s.authManager.ValidateAndParseToken(ctx, jwtToken) - if err != nil { - return "", status.Errorf(codes.InvalidArgument, "invalid jwt token, err: %v", err) - } - - // we need to call this method because if user is new, we will automatically add it to existing or create a new account - accountId, _, err := s.accountManager.GetAccountIDFromUserAuth(ctx, userAuth) - if err != nil { - return "", status.Errorf(codes.Internal, "unable to fetch account with claims, err: %v", err) - } - - if userAuth.AccountId != accountId { - log.WithContext(ctx).Debugf("gRPC server sets accountId from ensure, before %s, now %s", userAuth.AccountId, accountId) - userAuth.AccountId = accountId - } - - userAuth, err = s.authManager.EnsureUserAccessByJWTGroups(ctx, userAuth, token) - if err != nil { - return "", status.Error(codes.PermissionDenied, err.Error()) - } - - err = s.accountManager.SyncUserJWTGroups(ctx, userAuth) - if err != nil { - log.WithContext(ctx).Errorf("gRPC server failed to sync user JWT groups: %s", err) - } - - return userAuth.UserId, nil -} - -func (s *GRPCServer) acquirePeerLockByUID(ctx context.Context, uniqueID string) (unlock func()) { - log.WithContext(ctx).Tracef("acquiring peer lock for ID %s", uniqueID) - - start := time.Now() - value, _ := s.peerLocks.LoadOrStore(uniqueID, &sync.RWMutex{}) - mtx := value.(*sync.RWMutex) - mtx.Lock() - log.WithContext(ctx).Tracef("acquired peer lock for ID %s in %v", uniqueID, time.Since(start)) - start = time.Now() - - unlock = func() { - mtx.Unlock() - log.WithContext(ctx).Tracef("released peer lock for ID %s in %v", uniqueID, time.Since(start)) - } - - return unlock -} - -// maps internal internalStatus.Error to gRPC status.Error -func mapError(ctx context.Context, err error) error { - if e, ok := internalStatus.FromError(err); ok { - switch e.Type() { - case internalStatus.PermissionDenied: - return status.Error(codes.PermissionDenied, e.Message) - case internalStatus.Unauthorized: - return status.Error(codes.PermissionDenied, e.Message) - case internalStatus.Unauthenticated: - return status.Error(codes.PermissionDenied, e.Message) - case internalStatus.PreconditionFailed: - return status.Error(codes.FailedPrecondition, e.Message) - case internalStatus.NotFound: - return status.Error(codes.NotFound, e.Message) - default: - } - } - log.WithContext(ctx).Errorf("got an unhandled error: %s", err) - return status.Errorf(codes.Internal, "failed handling request") -} - -func extractPeerMeta(ctx context.Context, meta *proto.PeerSystemMeta) nbpeer.PeerSystemMeta { - if meta == nil { - return nbpeer.PeerSystemMeta{} - } - - osVersion := meta.GetOSVersion() - if osVersion == "" { - osVersion = meta.GetCore() - } - - networkAddresses := make([]nbpeer.NetworkAddress, 0, len(meta.GetNetworkAddresses())) - for _, addr := range meta.GetNetworkAddresses() { - netAddr, err := netip.ParsePrefix(addr.GetNetIP()) - if err != nil { - log.WithContext(ctx).Warnf("failed to parse netip address, %s: %v", addr.GetNetIP(), err) - continue - } - networkAddresses = append(networkAddresses, nbpeer.NetworkAddress{ - NetIP: netAddr, - Mac: addr.GetMac(), - }) - } - - files := make([]nbpeer.File, 0, len(meta.GetFiles())) - for _, file := range meta.GetFiles() { - files = append(files, nbpeer.File{ - Path: file.GetPath(), - Exist: file.GetExist(), - ProcessIsRunning: file.GetProcessIsRunning(), - }) - } - - return nbpeer.PeerSystemMeta{ - Hostname: meta.GetHostname(), - GoOS: meta.GetGoOS(), - Kernel: meta.GetKernel(), - Platform: meta.GetPlatform(), - OS: meta.GetOS(), - OSVersion: osVersion, - WtVersion: meta.GetNetbirdVersion(), - UIVersion: meta.GetUiVersion(), - KernelVersion: meta.GetKernelVersion(), - NetworkAddresses: networkAddresses, - SystemSerialNumber: meta.GetSysSerialNumber(), - SystemProductName: meta.GetSysProductName(), - SystemManufacturer: meta.GetSysManufacturer(), - Environment: nbpeer.Environment{ - Cloud: meta.GetEnvironment().GetCloud(), - Platform: meta.GetEnvironment().GetPlatform(), - }, - Files: files, - } -} - -func (s *GRPCServer) parseRequest(ctx context.Context, req *proto.EncryptedMessage, parsed pb.Message) (wgtypes.Key, error) { - peerKey, err := wgtypes.ParseKey(req.GetWgPubKey()) - if err != nil { - log.WithContext(ctx).Warnf("error while parsing peer's WireGuard public key %s.", req.WgPubKey) - return wgtypes.Key{}, status.Errorf(codes.InvalidArgument, "provided wgPubKey %s is invalid", req.WgPubKey) - } - - err = encryption.DecryptMessage(peerKey, s.wgKey, req.Body, parsed) - if err != nil { - return wgtypes.Key{}, status.Errorf(codes.InvalidArgument, "invalid request message") - } - - return peerKey, nil -} - -// Login endpoint first checks whether peer is registered under any account -// In case it is, the login is successful -// In case it isn't, the endpoint checks whether setup key is provided within the request and tries to register a peer. -// In case of the successful registration login is also successful -func (s *GRPCServer) Login(ctx context.Context, req *proto.EncryptedMessage) (*proto.EncryptedMessage, error) { - reqStart := time.Now() - defer func() { - if s.appMetrics != nil { - s.appMetrics.GRPCMetrics().CountLoginRequestDuration(time.Since(reqStart)) - } - }() - if s.appMetrics != nil { - s.appMetrics.GRPCMetrics().CountLoginRequest() - } - realIP := getRealIP(ctx) - log.WithContext(ctx).Debugf("Login request from peer [%s] [%s]", req.WgPubKey, realIP.String()) - - loginReq := &proto.LoginRequest{} - peerKey, err := s.parseRequest(ctx, req, loginReq) - if err != nil { - return nil, err - } - - //nolint - ctx = context.WithValue(ctx, nbContext.PeerIDKey, peerKey.String()) - accountID, err := s.accountManager.GetAccountIDForPeerKey(ctx, peerKey.String()) - if err != nil { - // this case should not happen and already indicates an issue but we don't want the system to fail due to being unable to log in detail - accountID = "UNKNOWN" - } - //nolint - ctx = context.WithValue(ctx, nbContext.AccountIDKey, accountID) - - if loginReq.GetMeta() == nil { - msg := status.Errorf(codes.FailedPrecondition, - "peer system meta has to be provided to log in. Peer %s, remote addr %s", peerKey.String(), realIP) - log.WithContext(ctx).Warn(msg) - return nil, msg - } - - userID, err := s.processJwtToken(ctx, loginReq, peerKey) - if err != nil { - return nil, err - } - - var sshKey []byte - if loginReq.GetPeerKeys() != nil { - sshKey = loginReq.GetPeerKeys().GetSshPubKey() - } - - peer, netMap, postureChecks, err := s.accountManager.LoginPeer(ctx, types.PeerLogin{ - WireGuardPubKey: peerKey.String(), - SSHKey: string(sshKey), - Meta: extractPeerMeta(ctx, loginReq.GetMeta()), - UserID: userID, - SetupKey: loginReq.GetSetupKey(), - ConnectionIP: realIP, - ExtraDNSLabels: loginReq.GetDnsLabels(), - }) - if err != nil { - log.WithContext(ctx).Warnf("failed logging in peer %s: %s", peerKey, err) - return nil, mapError(ctx, err) - } - - // if the login request contains setup key then it is a registration request - if loginReq.GetSetupKey() != "" { - s.ephemeralManager.OnPeerDisconnected(ctx, peer) - } - - var relayToken *Token - if s.config.Relay != nil && len(s.config.Relay.Addresses) > 0 { - relayToken, err = s.secretsManager.GenerateRelayToken() - if err != nil { - log.Errorf("failed generating Relay token: %v", err) - } - } - - // if peer has reached this point then it has logged in - loginResp := &proto.LoginResponse{ - NetbirdConfig: toNetbirdConfig(s.config, nil, relayToken, nil), - PeerConfig: toPeerConfig(peer, netMap.Network, s.accountManager.GetDNSDomain(), false), - Checks: toProtocolChecks(ctx, postureChecks), - } - encryptedResp, err := encryption.EncryptMessage(peerKey, s.wgKey, loginResp) - if err != nil { - log.WithContext(ctx).Warnf("failed encrypting peer %s message", peer.ID) - return nil, status.Errorf(codes.Internal, "failed logging in peer") - } - - return &proto.EncryptedMessage{ - WgPubKey: s.wgKey.PublicKey().String(), - Body: encryptedResp, - }, nil -} - -// processJwtToken validates the existence of a JWT token in the login request, and returns the corresponding user ID if -// the token is valid. -// -// The user ID can be empty if the token is not provided, which is acceptable if the peer is already -// registered or if it uses a setup key to register. -func (s *GRPCServer) processJwtToken(ctx context.Context, loginReq *proto.LoginRequest, peerKey wgtypes.Key) (string, error) { - userID := "" - if loginReq.GetJwtToken() != "" { - var err error - for i := 0; i < 3; i++ { - userID, err = s.validateToken(ctx, loginReq.GetJwtToken()) - if err == nil { - break - } - log.WithContext(ctx).Warnf("failed validating JWT token sent from peer %s with error %v. "+ - "Trying again as it may be due to the IdP cache issue", peerKey.String(), err) - time.Sleep(200 * time.Millisecond) - } - if err != nil { - return "", err - } - } - return userID, nil -} - -func ToResponseProto(configProto types.Protocol) proto.HostConfig_Protocol { - switch configProto { - case types.UDP: - return proto.HostConfig_UDP - case types.DTLS: - return proto.HostConfig_DTLS - case types.HTTP: - return proto.HostConfig_HTTP - case types.HTTPS: - return proto.HostConfig_HTTPS - case types.TCP: - return proto.HostConfig_TCP - default: - panic(fmt.Errorf("unexpected config protocol type %v", configProto)) - } -} - -func toNetbirdConfig(config *types.Config, turnCredentials *Token, relayToken *Token, extraSettings *types.ExtraSettings) *proto.NetbirdConfig { - if config == nil { - return nil - } - - var stuns []*proto.HostConfig - for _, stun := range config.Stuns { - stuns = append(stuns, &proto.HostConfig{ - Uri: stun.URI, - Protocol: ToResponseProto(stun.Proto), - }) - } - - var turns []*proto.ProtectedHostConfig - if config.TURNConfig != nil { - for _, turn := range config.TURNConfig.Turns { - var username string - var password string - if turnCredentials != nil { - username = turnCredentials.Payload - password = turnCredentials.Signature - } else { - username = turn.Username - password = turn.Password - } - turns = append(turns, &proto.ProtectedHostConfig{ - HostConfig: &proto.HostConfig{ - Uri: turn.URI, - Protocol: ToResponseProto(turn.Proto), - }, - User: username, - Password: password, - }) - } - } - - var relayCfg *proto.RelayConfig - if config.Relay != nil && len(config.Relay.Addresses) > 0 { - relayCfg = &proto.RelayConfig{ - Urls: config.Relay.Addresses, - } - - if relayToken != nil { - relayCfg.TokenPayload = relayToken.Payload - relayCfg.TokenSignature = relayToken.Signature - } - } - - var signalCfg *proto.HostConfig - if config.Signal != nil { - signalCfg = &proto.HostConfig{ - Uri: config.Signal.URI, - Protocol: ToResponseProto(config.Signal.Proto), - } - } - - nbConfig := &proto.NetbirdConfig{ - Stuns: stuns, - Turns: turns, - Signal: signalCfg, - Relay: relayCfg, - } - - return nbConfig -} - -func toPeerConfig(peer *nbpeer.Peer, network *types.Network, dnsName string, dnsResolutionOnRoutingPeerEnabled bool) *proto.PeerConfig { - netmask, _ := network.Net.Mask.Size() - fqdn := peer.FQDN(dnsName) - return &proto.PeerConfig{ - Address: fmt.Sprintf("%s/%d", peer.IP.String(), netmask), // take it from the network - SshConfig: &proto.SSHConfig{SshEnabled: peer.SSHEnabled}, - Fqdn: fqdn, - RoutingPeerDnsResolutionEnabled: dnsResolutionOnRoutingPeerEnabled, - } -} - -func toSyncResponse(ctx context.Context, config *types.Config, peer *nbpeer.Peer, turnCredentials *Token, relayCredentials *Token, networkMap *types.NetworkMap, dnsName string, checks []*posture.Checks, dnsCache *DNSConfigCache, dnsResolutionOnRoutingPeerEnabled bool, extraSettings *types.ExtraSettings) *proto.SyncResponse { - response := &proto.SyncResponse{ - PeerConfig: toPeerConfig(peer, networkMap.Network, dnsName, dnsResolutionOnRoutingPeerEnabled), - NetworkMap: &proto.NetworkMap{ - Serial: networkMap.Network.CurrentSerial(), - Routes: toProtocolRoutes(networkMap.Routes), - DNSConfig: toProtocolDNSConfig(networkMap.DNSConfig, dnsCache), - }, - Checks: toProtocolChecks(ctx, checks), - } - - nbConfig := toNetbirdConfig(config, turnCredentials, relayCredentials, extraSettings) - extendedConfig := integrationsConfig.ExtendNetBirdConfig(peer.ID, nbConfig, extraSettings) - response.NetbirdConfig = extendedConfig - - response.NetworkMap.PeerConfig = response.PeerConfig - - allPeers := make([]*proto.RemotePeerConfig, 0, len(networkMap.Peers)+len(networkMap.OfflinePeers)) - allPeers = appendRemotePeerConfig(allPeers, networkMap.Peers, dnsName) - response.RemotePeers = allPeers - response.NetworkMap.RemotePeers = allPeers - response.RemotePeersIsEmpty = len(allPeers) == 0 - response.NetworkMap.RemotePeersIsEmpty = response.RemotePeersIsEmpty - - response.NetworkMap.OfflinePeers = appendRemotePeerConfig(nil, networkMap.OfflinePeers, dnsName) - - firewallRules := toProtocolFirewallRules(networkMap.FirewallRules) - response.NetworkMap.FirewallRules = firewallRules - response.NetworkMap.FirewallRulesIsEmpty = len(firewallRules) == 0 - - routesFirewallRules := toProtocolRoutesFirewallRules(networkMap.RoutesFirewallRules) - response.NetworkMap.RoutesFirewallRules = routesFirewallRules - response.NetworkMap.RoutesFirewallRulesIsEmpty = len(routesFirewallRules) == 0 - - if networkMap.ForwardingRules != nil { - forwardingRules := make([]*proto.ForwardingRule, 0, len(networkMap.ForwardingRules)) - for _, rule := range networkMap.ForwardingRules { - forwardingRules = append(forwardingRules, rule.ToProto()) - } - response.NetworkMap.ForwardingRules = forwardingRules - } - - return response -} - -func appendRemotePeerConfig(dst []*proto.RemotePeerConfig, peers []*nbpeer.Peer, dnsName string) []*proto.RemotePeerConfig { - for _, rPeer := range peers { - dst = append(dst, &proto.RemotePeerConfig{ - WgPubKey: rPeer.Key, - AllowedIps: []string{rPeer.IP.String() + "/32"}, - SshConfig: &proto.SSHConfig{SshPubKey: []byte(rPeer.SSHKey)}, - Fqdn: rPeer.FQDN(dnsName), - }) - } - return dst -} - -// IsHealthy indicates whether the service is healthy -func (s *GRPCServer) IsHealthy(ctx context.Context, req *proto.Empty) (*proto.Empty, error) { - return &proto.Empty{}, nil -} - -// sendInitialSync sends initial proto.SyncResponse to the peer requesting synchronization -func (s *GRPCServer) sendInitialSync(ctx context.Context, peerKey wgtypes.Key, peer *nbpeer.Peer, networkMap *types.NetworkMap, postureChecks []*posture.Checks, srv proto.ManagementService_SyncServer) error { - var err error - - var turnToken *Token - if s.config.TURNConfig != nil && s.config.TURNConfig.TimeBasedCredentials { - turnToken, err = s.secretsManager.GenerateTurnToken() - if err != nil { - log.Errorf("failed generating TURN token: %v", err) - } - } - - var relayToken *Token - if s.config.Relay != nil && len(s.config.Relay.Addresses) > 0 { - relayToken, err = s.secretsManager.GenerateRelayToken() - if err != nil { - log.Errorf("failed generating Relay token: %v", err) - } - } - - settings, err := s.settingsManager.GetSettings(ctx, peer.AccountID, activity.SystemInitiator) - if err != nil { - return status.Errorf(codes.Internal, "error handling request") - } - - plainResp := toSyncResponse(ctx, s.config, peer, turnToken, relayToken, networkMap, s.accountManager.GetDNSDomain(), postureChecks, nil, settings.RoutingPeerDNSResolutionEnabled, settings.Extra) - - encryptedResp, err := encryption.EncryptMessage(peerKey, s.wgKey, plainResp) - if err != nil { - return status.Errorf(codes.Internal, "error handling request") - } - - err = srv.Send(&proto.EncryptedMessage{ - WgPubKey: s.wgKey.PublicKey().String(), - Body: encryptedResp, - }) - - if err != nil { - log.WithContext(ctx).Errorf("failed sending SyncResponse %v", err) - return status.Errorf(codes.Internal, "error handling request") - } - - return nil -} - -// GetDeviceAuthorizationFlow returns a device authorization flow information -// This is used for initiating an Oauth 2 device authorization grant flow -// which will be used by our clients to Login -func (s *GRPCServer) GetDeviceAuthorizationFlow(ctx context.Context, req *proto.EncryptedMessage) (*proto.EncryptedMessage, error) { - log.WithContext(ctx).Tracef("GetDeviceAuthorizationFlow request for pubKey: %s", req.WgPubKey) - start := time.Now() - defer func() { - log.WithContext(ctx).Tracef("GetDeviceAuthorizationFlow for pubKey: %s took %v", req.WgPubKey, time.Since(start)) - }() - - peerKey, err := wgtypes.ParseKey(req.GetWgPubKey()) - if err != nil { - errMSG := fmt.Sprintf("error while parsing peer's Wireguard public key %s on GetDeviceAuthorizationFlow request.", req.WgPubKey) - log.WithContext(ctx).Warn(errMSG) - return nil, status.Error(codes.InvalidArgument, errMSG) - } - - err = encryption.DecryptMessage(peerKey, s.wgKey, req.Body, &proto.DeviceAuthorizationFlowRequest{}) - if err != nil { - errMSG := fmt.Sprintf("error while decrypting peer's message with Wireguard public key %s.", req.WgPubKey) - log.WithContext(ctx).Warn(errMSG) - return nil, status.Error(codes.InvalidArgument, errMSG) - } - - if s.config.DeviceAuthorizationFlow == nil || s.config.DeviceAuthorizationFlow.Provider == string(types.NONE) { - return nil, status.Error(codes.NotFound, "no device authorization flow information available") - } - - provider, ok := proto.DeviceAuthorizationFlowProvider_value[strings.ToUpper(s.config.DeviceAuthorizationFlow.Provider)] - if !ok { - return nil, status.Errorf(codes.InvalidArgument, "no provider found in the protocol for %s", s.config.DeviceAuthorizationFlow.Provider) - } - - flowInfoResp := &proto.DeviceAuthorizationFlow{ - Provider: proto.DeviceAuthorizationFlowProvider(provider), - ProviderConfig: &proto.ProviderConfig{ - ClientID: s.config.DeviceAuthorizationFlow.ProviderConfig.ClientID, - ClientSecret: s.config.DeviceAuthorizationFlow.ProviderConfig.ClientSecret, - Domain: s.config.DeviceAuthorizationFlow.ProviderConfig.Domain, - Audience: s.config.DeviceAuthorizationFlow.ProviderConfig.Audience, - DeviceAuthEndpoint: s.config.DeviceAuthorizationFlow.ProviderConfig.DeviceAuthEndpoint, - TokenEndpoint: s.config.DeviceAuthorizationFlow.ProviderConfig.TokenEndpoint, - Scope: s.config.DeviceAuthorizationFlow.ProviderConfig.Scope, - UseIDToken: s.config.DeviceAuthorizationFlow.ProviderConfig.UseIDToken, - }, - } - - encryptedResp, err := encryption.EncryptMessage(peerKey, s.wgKey, flowInfoResp) - if err != nil { - return nil, status.Error(codes.Internal, "failed to encrypt no device authorization flow information") - } - - return &proto.EncryptedMessage{ - WgPubKey: s.wgKey.PublicKey().String(), - Body: encryptedResp, - }, nil -} - -// GetPKCEAuthorizationFlow returns a pkce authorization flow information -// This is used for initiating an Oauth 2 pkce authorization grant flow -// which will be used by our clients to Login -func (s *GRPCServer) GetPKCEAuthorizationFlow(ctx context.Context, req *proto.EncryptedMessage) (*proto.EncryptedMessage, error) { - log.WithContext(ctx).Tracef("GetPKCEAuthorizationFlow request for pubKey: %s", req.WgPubKey) - start := time.Now() - defer func() { - log.WithContext(ctx).Tracef("GetPKCEAuthorizationFlow for pubKey %s took %v", req.WgPubKey, time.Since(start)) - }() - - peerKey, err := wgtypes.ParseKey(req.GetWgPubKey()) - if err != nil { - errMSG := fmt.Sprintf("error while parsing peer's Wireguard public key %s on GetPKCEAuthorizationFlow request.", req.WgPubKey) - log.WithContext(ctx).Warn(errMSG) - return nil, status.Error(codes.InvalidArgument, errMSG) - } - - err = encryption.DecryptMessage(peerKey, s.wgKey, req.Body, &proto.PKCEAuthorizationFlowRequest{}) - if err != nil { - errMSG := fmt.Sprintf("error while decrypting peer's message with Wireguard public key %s.", req.WgPubKey) - log.WithContext(ctx).Warn(errMSG) - return nil, status.Error(codes.InvalidArgument, errMSG) - } - - if s.config.PKCEAuthorizationFlow == nil { - return nil, status.Error(codes.NotFound, "no pkce authorization flow information available") - } - - flowInfoResp := &proto.PKCEAuthorizationFlow{ - ProviderConfig: &proto.ProviderConfig{ - Audience: s.config.PKCEAuthorizationFlow.ProviderConfig.Audience, - ClientID: s.config.PKCEAuthorizationFlow.ProviderConfig.ClientID, - ClientSecret: s.config.PKCEAuthorizationFlow.ProviderConfig.ClientSecret, - TokenEndpoint: s.config.PKCEAuthorizationFlow.ProviderConfig.TokenEndpoint, - AuthorizationEndpoint: s.config.PKCEAuthorizationFlow.ProviderConfig.AuthorizationEndpoint, - Scope: s.config.PKCEAuthorizationFlow.ProviderConfig.Scope, - RedirectURLs: s.config.PKCEAuthorizationFlow.ProviderConfig.RedirectURLs, - UseIDToken: s.config.PKCEAuthorizationFlow.ProviderConfig.UseIDToken, - DisablePromptLogin: s.config.PKCEAuthorizationFlow.ProviderConfig.DisablePromptLogin, - }, - } - - encryptedResp, err := encryption.EncryptMessage(peerKey, s.wgKey, flowInfoResp) - if err != nil { - return nil, status.Error(codes.Internal, "failed to encrypt no pkce authorization flow information") - } - - return &proto.EncryptedMessage{ - WgPubKey: s.wgKey.PublicKey().String(), - Body: encryptedResp, - }, nil -} - -// SyncMeta endpoint is used to synchronize peer's system metadata and notifies the connected, -// peer's under the same account of any updates. -func (s *GRPCServer) SyncMeta(ctx context.Context, req *proto.EncryptedMessage) (*proto.Empty, error) { - realIP := getRealIP(ctx) - log.WithContext(ctx).Debugf("Sync meta request from peer [%s] [%s]", req.WgPubKey, realIP.String()) - - syncMetaReq := &proto.SyncMetaRequest{} - peerKey, err := s.parseRequest(ctx, req, syncMetaReq) - if err != nil { - return nil, err - } - - if syncMetaReq.GetMeta() == nil { - msg := status.Errorf(codes.FailedPrecondition, - "peer system meta has to be provided on sync. Peer %s, remote addr %s", peerKey.String(), realIP) - log.WithContext(ctx).Warn(msg) - return nil, msg - } - - err = s.accountManager.SyncPeerMeta(ctx, peerKey.String(), extractPeerMeta(ctx, syncMetaReq.GetMeta())) - if err != nil { - return nil, mapError(ctx, err) - } - - return &proto.Empty{}, nil -} - -// toProtocolChecks converts posture checks to protocol checks. -func toProtocolChecks(ctx context.Context, postureChecks []*posture.Checks) []*proto.Checks { - protoChecks := make([]*proto.Checks, 0, len(postureChecks)) - for _, postureCheck := range postureChecks { - protoChecks = append(protoChecks, toProtocolCheck(postureCheck)) - } - - return protoChecks -} - -// toProtocolCheck converts a posture.Checks to a proto.Checks. -func toProtocolCheck(postureCheck *posture.Checks) *proto.Checks { - protoCheck := &proto.Checks{} - - if check := postureCheck.Checks.ProcessCheck; check != nil { - for _, process := range check.Processes { - if process.LinuxPath != "" { - protoCheck.Files = append(protoCheck.Files, process.LinuxPath) - } - if process.MacPath != "" { - protoCheck.Files = append(protoCheck.Files, process.MacPath) - } - if process.WindowsPath != "" { - protoCheck.Files = append(protoCheck.Files, process.WindowsPath) - } - } - } - - return protoCheck -} diff --git a/internals/shared/api/grpc/updatechannel.go b/internals/shared/api/grpc/updatechannel.go deleted file mode 100644 index 63cf79f..0000000 --- a/internals/shared/api/grpc/updatechannel.go +++ /dev/null @@ -1,178 +0,0 @@ -package server - -import ( - "context" - "sync" - "time" - - log "github.com/sirupsen/logrus" - - "github.com/netbirdio/management-refactor/management/proto" - "github.com/netbirdio/management-refactor/management/server/telemetry" - "github.com/netbirdio/management-refactor/management/server/types" -) - -const channelBufferSize = 100 - -type UpdateMessage struct { - Update *proto.SyncResponse - NetworkMap *types.NetworkMap -} - -type UpdateChannel struct { - // peerChannels is an update channel indexed by Peer.ID - peerChannels map[string]chan *UpdateMessage - // channelsMux keeps the mutex to access peerChannels - channelsMux *sync.RWMutex - // metrics provides method to collect application metrics - metrics telemetry.AppMetrics -} - -// NewUpdateChannel returns a new instance of UpdateChannel -func NewUpdateChannel(metrics telemetry.AppMetrics) *UpdateChannel { - return &UpdateChannel{ - peerChannels: make(map[string]chan *UpdateMessage), - channelsMux: &sync.RWMutex{}, - metrics: metrics, - } -} - -// SendUpdate sends update message to the peer's channel -func (p *UpdateChannel) SendUpdate(ctx context.Context, peerID string, update *UpdateMessage) { - start := time.Now() - var found, dropped bool - - p.channelsMux.RLock() - - defer func() { - p.channelsMux.RUnlock() - if p.metrics != nil { - p.metrics.UpdateChannelMetrics().CountSendUpdateDuration(time.Since(start), found, dropped) - } - }() - - if channel, ok := p.peerChannels[peerID]; ok { - found = true - select { - case channel <- update: - log.WithContext(ctx).Debugf("update was sent to channel for peer %s", peerID) - default: - dropped = true - log.WithContext(ctx).Warnf("channel for peer %s is %d full or closed", peerID, len(channel)) - } - } else { - log.WithContext(ctx).Debugf("peer %s has no channel", peerID) - } -} - -// CreateChannel creates a go channel for a given peer used to deliver updates relevant to the peer. -func (p *UpdateChannel) CreateChannel(ctx context.Context, peerID string) chan *UpdateMessage { - start := time.Now() - - closed := false - - p.channelsMux.Lock() - defer func() { - p.channelsMux.Unlock() - if p.metrics != nil { - p.metrics.UpdateChannelMetrics().CountCreateChannelDuration(time.Since(start), closed) - } - }() - - if channel, ok := p.peerChannels[peerID]; ok { - closed = true - delete(p.peerChannels, peerID) - close(channel) - } - // mbragin: todo shouldn't it be more? or configurable? - channel := make(chan *UpdateMessage, channelBufferSize) - p.peerChannels[peerID] = channel - - log.WithContext(ctx).Debugf("opened updates channel for a peer %s", peerID) - - return channel -} - -func (p *UpdateChannel) closeChannel(ctx context.Context, peerID string) { - if channel, ok := p.peerChannels[peerID]; ok { - delete(p.peerChannels, peerID) - close(channel) - - log.WithContext(ctx).Debugf("closed updates channel of a peer %s", peerID) - return - } - - log.WithContext(ctx).Debugf("closing updates channel: peer %s has no channel", peerID) -} - -// CloseChannels closes updates channel for each given peer -func (p *UpdateChannel) CloseChannels(ctx context.Context, peerIDs []string) { - start := time.Now() - - p.channelsMux.Lock() - defer func() { - p.channelsMux.Unlock() - if p.metrics != nil { - p.metrics.UpdateChannelMetrics().CountCloseChannelsDuration(time.Since(start), len(peerIDs)) - } - }() - - for _, id := range peerIDs { - p.closeChannel(ctx, id) - } -} - -// CloseChannel closes updates channel of a given peer -func (p *UpdateChannel) CloseChannel(ctx context.Context, peerID string) { - start := time.Now() - - p.channelsMux.Lock() - defer func() { - p.channelsMux.Unlock() - if p.metrics != nil { - p.metrics.UpdateChannelMetrics().CountCloseChannelDuration(time.Since(start)) - } - }() - - p.closeChannel(ctx, peerID) -} - -// GetAllConnectedPeers returns a copy of the connected peers map -func (p *UpdateChannel) GetAllConnectedPeers() map[string]struct{} { - start := time.Now() - - p.channelsMux.RLock() - - m := make(map[string]struct{}) - - defer func() { - p.channelsMux.RUnlock() - if p.metrics != nil { - p.metrics.UpdateChannelMetrics().CountGetAllConnectedPeersDuration(time.Since(start), len(m)) - } - }() - - for ID := range p.peerChannels { - m[ID] = struct{}{} - } - - return m -} - -// HasChannel returns true if peers has channel in update manager, otherwise false -func (p *UpdateChannel) HasChannel(peerID string) bool { - start := time.Now() - - p.channelsMux.RLock() - - defer func() { - p.channelsMux.RUnlock() - if p.metrics != nil { - p.metrics.UpdateChannelMetrics().CountHasChannelDuration(time.Since(start)) - } - }() - - _, ok := p.peerChannels[peerID] - - return ok -} diff --git a/internals/shared/api/grpc/updatechannel_test.go b/internals/shared/api/grpc/updatechannel_test.go deleted file mode 100644 index 5e966d4..0000000 --- a/internals/shared/api/grpc/updatechannel_test.go +++ /dev/null @@ -1,79 +0,0 @@ -package server - -import ( - "context" - "testing" - "time" - - "github.com/netbirdio/management-refactor/management/proto" -) - -// var peersUpdater *UpdateChannel - -func TestCreateChannel(t *testing.T) { - peer := "test-create" - peersUpdater := NewUpdateChannel(nil) - defer peersUpdater.CloseChannel(context.Background(), peer) - - _ = peersUpdater.CreateChannel(context.Background(), peer) - if _, ok := peersUpdater.peerChannels[peer]; !ok { - t.Error("Error creating the channel") - } -} - -func TestSendUpdate(t *testing.T) { - peer := "test-sendupdate" - peersUpdater := NewUpdateChannel(nil) - update1 := &UpdateMessage{Update: &proto.SyncResponse{ - NetworkMap: &proto.NetworkMap{ - Serial: 0, - }, - }} - _ = peersUpdater.CreateChannel(context.Background(), peer) - if _, ok := peersUpdater.peerChannels[peer]; !ok { - t.Error("Error creating the channel") - } - peersUpdater.SendUpdate(context.Background(), peer, update1) - select { - case <-peersUpdater.peerChannels[peer]: - default: - t.Error("Update wasn't send") - } - - for range [channelBufferSize]int{} { - peersUpdater.SendUpdate(context.Background(), peer, update1) - } - - update2 := &UpdateMessage{Update: &proto.SyncResponse{ - NetworkMap: &proto.NetworkMap{ - Serial: 10, - }, - }} - - peersUpdater.SendUpdate(context.Background(), peer, update2) - timeout := time.After(5 * time.Second) - for range [channelBufferSize]int{} { - select { - case <-timeout: - t.Error("timed out reading previously sent updates") - case updateReader := <-peersUpdater.peerChannels[peer]: - if updateReader.Update.NetworkMap.Serial == update2.Update.NetworkMap.Serial { - t.Error("got the update that shouldn't have been sent") - } - } - } - -} - -func TestCloseChannel(t *testing.T) { - peer := "test-close" - peersUpdater := NewUpdateChannel(nil) - _ = peersUpdater.CreateChannel(context.Background(), peer) - if _, ok := peersUpdater.peerChannels[peer]; !ok { - t.Error("Error creating the channel") - } - peersUpdater.CloseChannel(context.Background(), peer) - if _, ok := peersUpdater.peerChannels[peer]; ok { - t.Error("Error closing the channel") - } -} diff --git a/internals/shared/api/rest/middleware/auth_middleware_test.go b/internals/shared/api/rest/middleware/auth_middleware_test.go deleted file mode 100644 index 94ed4c4..0000000 --- a/internals/shared/api/rest/middleware/auth_middleware_test.go +++ /dev/null @@ -1,325 +0,0 @@ -package middleware - -import ( - "context" - "fmt" - "net/http" - "net/http/httptest" - "testing" - "time" - - "github.com/golang-jwt/jwt" - "github.com/stretchr/testify/assert" - - "github.com/netbirdio/management-refactor/management/server/auth" - nbjwt "github.com/netbirdio/management-refactor/management/server/auth/jwt" - nbcontext "github.com/netbirdio/management-refactor/management/server/context" - "github.com/netbirdio/management-refactor/management/server/util" - - "github.com/netbirdio/management-refactor/management/server/http/middleware/bypass" - "github.com/netbirdio/management-refactor/management/server/types" -) - -const ( - audience = "audience" - userIDClaim = "userIDClaim" - accountID = "accountID" - domain = "domain" - domainCategory = "domainCategory" - userID = "userID" - tokenID = "tokenID" - PAT = "nbp_PAT" - JWT = "JWT" - wrongToken = "wrongToken" -) - -var testAccount = &types.Account{ - Id: accountID, - Domain: domain, - Users: map[string]*types.User{ - userID: { - Id: userID, - AccountID: accountID, - PATs: map[string]*types.PersonalAccessToken{ - tokenID: { - ID: tokenID, - Name: "My first token", - HashedToken: "someHash", - ExpirationDate: util.ToPtr(time.Now().UTC().AddDate(0, 0, 7)), - CreatedBy: userID, - CreatedAt: time.Now().UTC(), - LastUsed: util.ToPtr(time.Now().UTC()), - }, - }, - }, - }, -} - -func mockGetAccountInfoFromPAT(_ context.Context, token string) (user *types.User, pat *types.PersonalAccessToken, domain string, category string, err error) { - if token == PAT { - return testAccount.Users[userID], testAccount.Users[userID].PATs[tokenID], testAccount.Domain, testAccount.DomainCategory, nil - } - return nil, nil, "", "", fmt.Errorf("PAT invalid") -} - -func mockValidateAndParseToken(_ context.Context, token string) (nbcontext.UserAuth, *jwt.Token, error) { - if token == JWT { - return nbcontext.UserAuth{ - UserId: userID, - AccountId: accountID, - Domain: testAccount.Domain, - DomainCategory: testAccount.DomainCategory, - }, - &jwt.Token{ - Claims: jwt.MapClaims{ - userIDClaim: userID, - audience + nbjwt.AccountIDSuffix: accountID, - }, - Valid: true, - }, nil - } - return nbcontext.UserAuth{}, nil, fmt.Errorf("JWT invalid") -} - -func mockMarkPATUsed(_ context.Context, token string) error { - if token == tokenID { - return nil - } - return fmt.Errorf("Should never get reached") -} - -func mockEnsureUserAccessByJWTGroups(_ context.Context, userAuth nbcontext.UserAuth, token *jwt.Token) (nbcontext.UserAuth, error) { - if userAuth.IsChild || userAuth.IsPAT { - return userAuth, nil - } - - if testAccount.Id != userAuth.AccountId { - return userAuth, fmt.Errorf("account with id %s does not exist", userAuth.AccountId) - } - - if _, ok := testAccount.Users[userAuth.UserId]; !ok { - return userAuth, fmt.Errorf("user with id %s does not exist", userAuth.UserId) - } - - return userAuth, nil -} - -func TestAuthMiddleware_Handler(t *testing.T) { - tt := []struct { - name string - path string - authHeader string - expectedStatusCode int - shouldBypassAuth bool - }{ - { - name: "Valid PAT Token", - path: "/test", - authHeader: "Token " + PAT, - expectedStatusCode: 200, - }, - { - name: "Invalid PAT Token", - path: "/test", - authHeader: "Token " + wrongToken, - expectedStatusCode: 401, - }, - { - name: "Fallback to PAT Token", - path: "/test", - authHeader: "Bearer " + PAT, - expectedStatusCode: 200, - }, - { - name: "Valid JWT Token", - path: "/test", - authHeader: "Bearer " + JWT, - expectedStatusCode: 200, - }, - { - name: "Invalid JWT Token", - path: "/test", - authHeader: "Bearer " + wrongToken, - expectedStatusCode: 401, - }, - { - name: "Basic Auth", - path: "/test", - authHeader: "Basic " + PAT, - expectedStatusCode: 401, - }, - { - name: "Webhook Path Bypass", - path: "/webhook", - authHeader: "", - expectedStatusCode: 200, - shouldBypassAuth: true, - }, - { - name: "Webhook Path Bypass with Subpath", - path: "/webhook/test", - authHeader: "", - expectedStatusCode: 200, - shouldBypassAuth: true, - }, - { - name: "Different Webhook Path", - path: "/webhooktest", - authHeader: "", - expectedStatusCode: 401, - shouldBypassAuth: false, - }, - } - - nextHandler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - - }) - - mockAuth := &auth.MockManager{ - ValidateAndParseTokenFunc: mockValidateAndParseToken, - EnsureUserAccessByJWTGroupsFunc: mockEnsureUserAccessByJWTGroups, - MarkPATUsedFunc: mockMarkPATUsed, - GetPATInfoFunc: mockGetAccountInfoFromPAT, - } - - authMiddleware := NewAuthMiddleware( - mockAuth, - func(ctx context.Context, userAuth nbcontext.UserAuth) (string, string, error) { - return userAuth.AccountId, userAuth.UserId, nil - }, - func(ctx context.Context, userAuth nbcontext.UserAuth) error { - return nil - }, - ) - - handlerToTest := authMiddleware.Handler(nextHandler) - - for _, tc := range tt { - t.Run(tc.name, func(t *testing.T) { - if tc.shouldBypassAuth { - err := bypass.AddBypassPath(tc.path) - if err != nil { - t.Fatalf("failed to add bypass path: %v", err) - } - } - - req := httptest.NewRequest("GET", "http://testing"+tc.path, nil) - req.Header.Set("Authorization", tc.authHeader) - rec := httptest.NewRecorder() - - handlerToTest.ServeHTTP(rec, req) - - result := rec.Result() - defer result.Body.Close() - - if result.StatusCode != tc.expectedStatusCode { - t.Errorf("expected status code %d, got %d", tc.expectedStatusCode, result.StatusCode) - } - }) - } -} - -func TestAuthMiddleware_Handler_Child(t *testing.T) { - tt := []struct { - name string - path string - authHeader string - expectedUserAuth *nbcontext.UserAuth // nil expects 401 response status - }{ - { - name: "Valid PAT Token", - path: "/test", - authHeader: "Token " + PAT, - expectedUserAuth: &nbcontext.UserAuth{ - AccountId: accountID, - UserId: userID, - Domain: testAccount.Domain, - DomainCategory: testAccount.DomainCategory, - IsPAT: true, - }, - }, - { - name: "Valid PAT Token ignores child", - path: "/test?account=xyz", - authHeader: "Token " + PAT, - expectedUserAuth: &nbcontext.UserAuth{ - AccountId: accountID, - UserId: userID, - Domain: testAccount.Domain, - DomainCategory: testAccount.DomainCategory, - IsPAT: true, - }, - }, - { - name: "Valid JWT Token", - path: "/test", - authHeader: "Bearer " + JWT, - expectedUserAuth: &nbcontext.UserAuth{ - AccountId: accountID, - UserId: userID, - Domain: testAccount.Domain, - DomainCategory: testAccount.DomainCategory, - }, - }, - - { - name: "Valid JWT Token with child", - path: "/test?account=xyz", - authHeader: "Bearer " + JWT, - expectedUserAuth: &nbcontext.UserAuth{ - AccountId: "xyz", - UserId: userID, - Domain: testAccount.Domain, - DomainCategory: testAccount.DomainCategory, - IsChild: true, - }, - }, - } - - mockAuth := &auth.MockManager{ - ValidateAndParseTokenFunc: mockValidateAndParseToken, - EnsureUserAccessByJWTGroupsFunc: mockEnsureUserAccessByJWTGroups, - MarkPATUsedFunc: mockMarkPATUsed, - GetPATInfoFunc: mockGetAccountInfoFromPAT, - } - - authMiddleware := NewAuthMiddleware( - mockAuth, - func(ctx context.Context, userAuth nbcontext.UserAuth) (string, string, error) { - return userAuth.AccountId, userAuth.UserId, nil - }, - func(ctx context.Context, userAuth nbcontext.UserAuth) error { - return nil - }, - ) - - for _, tc := range tt { - t.Run(tc.name, func(t *testing.T) { - handlerToTest := authMiddleware.Handler(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - userAuth, err := nbcontext.GetUserAuthFromRequest(r) - if tc.expectedUserAuth != nil { - assert.NoError(t, err) - assert.Equal(t, *tc.expectedUserAuth, userAuth) - } else { - assert.Error(t, err) - assert.Empty(t, userAuth) - } - })) - - req := httptest.NewRequest("GET", "http://testing"+tc.path, nil) - req.Header.Set("Authorization", tc.authHeader) - rec := httptest.NewRecorder() - - handlerToTest.ServeHTTP(rec, req) - - result := rec.Result() - defer result.Body.Close() - - if tc.expectedUserAuth != nil { - assert.Equal(t, 200, result.StatusCode) - } else { - assert.Equal(t, 401, result.StatusCode) - } - }) - } -} diff --git a/internals/shared/api/rest/middleware/bypass/bypass_test.go b/internals/shared/api/rest/middleware/bypass/bypass_test.go deleted file mode 100644 index 01c3689..0000000 --- a/internals/shared/api/rest/middleware/bypass/bypass_test.go +++ /dev/null @@ -1,131 +0,0 @@ -package bypass_test - -import ( - "net/http" - "net/http/httptest" - "testing" - - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" - - "github.com/netbirdio/management-refactor/management/server/http/middleware/bypass" -) - -func TestGetList(t *testing.T) { - bypassPaths := []string{"/path1", "/path2", "/path3"} - - for _, path := range bypassPaths { - err := bypass.AddBypassPath(path) - require.NoError(t, err, "Adding bypass path should not fail") - } - - list := bypass.GetList() - - assert.ElementsMatch(t, bypassPaths, list, "Bypass path list did not match expected paths") -} - -func TestAuthBypass(t *testing.T) { - dummyHandler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusOK) - }) - - tests := []struct { - name string - pathToAdd string - pathToRemove string - testPath string - expectBypass bool - expectHTTPCode int - }{ - { - name: "Path added to bypass", - pathToAdd: "/bypass", - testPath: "/bypass", - expectBypass: true, - expectHTTPCode: http.StatusOK, - }, - { - name: "Wildcard path added to bypass", - pathToAdd: "/bypass/*", - testPath: "/bypass/extra", - expectBypass: true, - expectHTTPCode: http.StatusOK, - }, - { - name: "Path not added to bypass", - testPath: "/no-bypass", - expectBypass: false, - expectHTTPCode: http.StatusOK, - }, - { - name: "Path removed from bypass", - pathToAdd: "/remove-bypass", - pathToRemove: "/remove-bypass", - testPath: "/remove-bypass", - expectBypass: false, - expectHTTPCode: http.StatusOK, - }, - { - name: "Exact path matches bypass", - pathToAdd: "/webhook", - testPath: "/webhook", - expectBypass: true, - expectHTTPCode: http.StatusOK, - }, - { - name: "Subpath does not match bypass", - pathToAdd: "/webhook", - testPath: "/webhook/extra", - expectBypass: false, - expectHTTPCode: http.StatusOK, - }, - { - name: "Wildcard subpath does not match bypass", - pathToAdd: "/webhook/*", - testPath: "/webhook/extra/path", - expectBypass: false, - expectHTTPCode: http.StatusOK, - }, - { - name: "Similar path does not match bypass", - pathToAdd: "/webhook", - testPath: "/webhooking", - expectBypass: false, - expectHTTPCode: http.StatusOK, - }, - { - name: "Prefix path does not match bypass", - pathToAdd: "/webhook", - testPath: "/web", - expectBypass: false, - expectHTTPCode: http.StatusOK, - }, - } - - for _, tc := range tests { - t.Run(tc.name, func(t *testing.T) { - if tc.pathToAdd != "" { - err := bypass.AddBypassPath(tc.pathToAdd) - require.NoError(t, err, "Adding bypass path should not fail") - defer bypass.RemovePath(tc.pathToAdd) - } - - if tc.pathToRemove != "" { - bypass.RemovePath(tc.pathToRemove) - } - - request, err := http.NewRequest("GET", tc.testPath, nil) - require.NoError(t, err, "Creating request should not fail") - - recorder := httptest.NewRecorder() - - bypassed := bypass.ShouldBypass(tc.testPath, dummyHandler, recorder, request) - - assert.Equal(t, tc.expectBypass, bypassed, "Bypass check did not match expectation") - - if tc.expectBypass { - assert.Equal(t, tc.expectHTTPCode, recorder.Code, "HTTP status code did not match expectation for bypassed path") - } - }) - } -} diff --git a/internals/shared/auth/jwt/extractor.go b/internals/shared/auth/jwt/extractor.go deleted file mode 100644 index 58e74a5..0000000 --- a/internals/shared/auth/jwt/extractor.go +++ /dev/null @@ -1,144 +0,0 @@ -package jwt - -import ( - "errors" - "net/url" - "time" - - "github.com/golang-jwt/jwt" - log "github.com/sirupsen/logrus" - - nbcontext "github.com/netbirdio/management-refactor/management/server/context" -) - -const ( - // AccountIDSuffix suffix for the account id claim - AccountIDSuffix = "wt_account_id" - // DomainIDSuffix suffix for the domain id claim - DomainIDSuffix = "wt_account_domain" - // DomainCategorySuffix suffix for the domain category claim - DomainCategorySuffix = "wt_account_domain_category" - // UserIDClaim claim for the user id - UserIDClaim = "sub" - // LastLoginSuffix claim for the last login - LastLoginSuffix = "nb_last_login" - // Invited claim indicates that an incoming JWT is from a user that just accepted an invitation - Invited = "nb_invited" -) - -var ( - errUserIDClaimEmpty = errors.New("user ID claim token value is empty") -) - -// ClaimsExtractor struct that holds the extract function -type ClaimsExtractor struct { - authAudience string - userIDClaim string -} - -// ClaimsExtractorOption is a function that configures the ClaimsExtractor -type ClaimsExtractorOption func(*ClaimsExtractor) - -// WithAudience sets the audience for the extractor -func WithAudience(audience string) ClaimsExtractorOption { - return func(c *ClaimsExtractor) { - c.authAudience = audience - } -} - -// WithUserIDClaim sets the user id claim for the extractor -func WithUserIDClaim(userIDClaim string) ClaimsExtractorOption { - return func(c *ClaimsExtractor) { - c.userIDClaim = userIDClaim - } -} - -// NewClaimsExtractor returns an extractor, and if provided with a function with ExtractClaims signature, -// then it will use that logic. Uses ExtractClaimsFromRequestContext by default -func NewClaimsExtractor(options ...ClaimsExtractorOption) *ClaimsExtractor { - ce := &ClaimsExtractor{} - for _, option := range options { - option(ce) - } - - if ce.userIDClaim == "" { - ce.userIDClaim = UserIDClaim - } - return ce -} - -func parseTime(timeString string) time.Time { - if timeString == "" { - return time.Time{} - } - parsedTime, err := time.Parse(time.RFC3339, timeString) - if err != nil { - return time.Time{} - } - return parsedTime -} - -func (c ClaimsExtractor) audienceClaim(claimName string) string { - url, err := url.JoinPath(c.authAudience, claimName) - if err != nil { - return c.authAudience + claimName // as it was previously - } - - return url -} - -func (c *ClaimsExtractor) ToUserAuth(token *jwt.Token) (nbcontext.UserAuth, error) { - claims := token.Claims.(jwt.MapClaims) - userAuth := nbcontext.UserAuth{} - - userID, ok := claims[c.userIDClaim].(string) - if !ok { - return userAuth, errUserIDClaimEmpty - } - userAuth.UserId = userID - - if accountIDClaim, ok := claims[c.audienceClaim(AccountIDSuffix)]; ok { - userAuth.AccountId = accountIDClaim.(string) - } - - if domainClaim, ok := claims[c.audienceClaim(DomainIDSuffix)]; ok { - userAuth.Domain = domainClaim.(string) - } - - if domainCategoryClaim, ok := claims[c.audienceClaim(DomainCategorySuffix)]; ok { - userAuth.DomainCategory = domainCategoryClaim.(string) - } - - if lastLoginClaimString, ok := claims[c.audienceClaim(LastLoginSuffix)]; ok { - userAuth.LastLogin = parseTime(lastLoginClaimString.(string)) - } - - if invitedBool, ok := claims[c.audienceClaim(Invited)]; ok { - if value, ok := invitedBool.(bool); ok { - userAuth.Invited = value - } - } - - return userAuth, nil -} - -func (c *ClaimsExtractor) ToGroups(token *jwt.Token, claimName string) []string { - claims := token.Claims.(jwt.MapClaims) - userJWTGroups := make([]string, 0) - - if claim, ok := claims[claimName]; ok { - if claimGroups, ok := claim.([]interface{}); ok { - for _, g := range claimGroups { - if group, ok := g.(string); ok { - userJWTGroups = append(userJWTGroups, group) - } else { - log.Debugf("JWT claim %q contains a non-string group (type: %T): %v", claimName, g, g) - } - } - } - } else { - log.Debugf("JWT claim %q is not a string array", claimName) - } - - return userJWTGroups -} diff --git a/internals/shared/auth/jwt/validator.go b/internals/shared/auth/jwt/validator.go deleted file mode 100644 index 5b38ca7..0000000 --- a/internals/shared/auth/jwt/validator.go +++ /dev/null @@ -1,302 +0,0 @@ -package jwt - -import ( - "context" - "crypto/ecdsa" - "crypto/elliptic" - "crypto/rsa" - "encoding/base64" - "encoding/json" - "errors" - "fmt" - "math/big" - "net/http" - "net/url" - "strconv" - "strings" - "sync" - "time" - - "github.com/golang-jwt/jwt" - - log "github.com/sirupsen/logrus" -) - -// Jwks is a collection of JSONWebKey obtained from Config.HttpServerConfig.AuthKeysLocation -type Jwks struct { - Keys []JSONWebKey `json:"keys"` - expiresInTime time.Time -} - -// The supported elliptic curves types -const ( - // p256 represents a cryptographic elliptical curve type. - p256 = "P-256" - - // p384 represents a cryptographic elliptical curve type. - p384 = "P-384" - - // p521 represents a cryptographic elliptical curve type. - p521 = "P-521" -) - -// JSONWebKey is a representation of a Jason Web Key -type JSONWebKey struct { - Kty string `json:"kty"` - Kid string `json:"kid"` - Use string `json:"use"` - N string `json:"n"` - E string `json:"e"` - Crv string `json:"crv"` - X string `json:"x"` - Y string `json:"y"` - X5c []string `json:"x5c"` -} - -type Validator struct { - lock sync.Mutex - issuer string - audienceList []string - keysLocation string - idpSignkeyRefreshEnabled bool - keys *Jwks -} - -var ( - errKeyNotFound = errors.New("unable to find appropriate key") - errInvalidAudience = errors.New("invalid audience") - errInvalidIssuer = errors.New("invalid issuer") - errTokenEmpty = errors.New("required authorization token not found") - errTokenInvalid = errors.New("token is invalid") - errTokenParsing = errors.New("token could not be parsed") -) - -func NewValidator(issuer string, audienceList []string, keysLocation string, idpSignkeyRefreshEnabled bool) *Validator { - keys, err := getPemKeys(keysLocation) - if err != nil { - log.WithField("keysLocation", keysLocation).Errorf("could not get keys from location: %s", err) - } - - return &Validator{ - keys: keys, - issuer: issuer, - audienceList: audienceList, - keysLocation: keysLocation, - idpSignkeyRefreshEnabled: idpSignkeyRefreshEnabled, - } -} - -func (v *Validator) getKeyFunc(ctx context.Context) jwt.Keyfunc { - return func(token *jwt.Token) (interface{}, error) { - // Verify 'aud' claim - var checkAud bool - for _, audience := range v.audienceList { - checkAud = token.Claims.(jwt.MapClaims).VerifyAudience(audience, false) - if checkAud { - break - } - } - if !checkAud { - return token, errInvalidAudience - } - - // Verify 'issuer' claim - checkIss := token.Claims.(jwt.MapClaims).VerifyIssuer(v.issuer, false) - if !checkIss { - return token, errInvalidIssuer - } - - // If keys are rotated, verify the keys prior to token validation - if v.idpSignkeyRefreshEnabled { - // If the keys are invalid, retrieve new ones - // @todo propose a separate go routine to regularly check these to prevent blocking when actually - // validating the token - if !v.keys.stillValid() { - v.lock.Lock() - defer v.lock.Unlock() - - refreshedKeys, err := getPemKeys(v.keysLocation) - if err != nil { - log.WithContext(ctx).Debugf("cannot get JSONWebKey: %v, falling back to old keys", err) - refreshedKeys = v.keys - } - - log.WithContext(ctx).Debugf("keys refreshed, new UTC expiration time: %s", refreshedKeys.expiresInTime.UTC()) - - v.keys = refreshedKeys - } - } - - publicKey, err := getPublicKey(token, v.keys) - if err == nil { - return publicKey, nil - } - - msg := fmt.Sprintf("getPublicKey error: %s", err) - if errors.Is(err, errKeyNotFound) && !v.idpSignkeyRefreshEnabled { - msg = fmt.Sprintf("getPublicKey error: %s. You can enable key refresh by setting HttpServerConfig.IdpSignKeyRefreshEnabled to true in your management.json file and restart the service", err) - } - - log.WithContext(ctx).Error(msg) - - return nil, err - } -} - -// ValidateAndParse validates the token and returns the parsed token -func (m *Validator) ValidateAndParse(ctx context.Context, token string) (*jwt.Token, error) { - // If the token is empty... - if token == "" { - // If we get here, the required token is missing - log.WithContext(ctx).Debugf(" Error: No credentials found (CredentialsOptional=false)") - return nil, errTokenEmpty - } - - // Now parse the token - parsedToken, err := jwt.Parse(token, m.getKeyFunc(ctx)) - - // Check if there was an error in parsing... - if err != nil { - err = fmt.Errorf("%w: %s", errTokenParsing, err) - log.WithContext(ctx).Error(err.Error()) - return nil, err - } - - // Check if the parsed token is valid... - if !parsedToken.Valid { - log.WithContext(ctx).Debug(errTokenInvalid.Error()) - return nil, errTokenInvalid - } - - return parsedToken, nil -} - -// stillValid returns true if the JSONWebKey still valid and have enough time to be used -func (jwks *Jwks) stillValid() bool { - return !jwks.expiresInTime.IsZero() && time.Now().Add(5*time.Second).Before(jwks.expiresInTime) -} - -func getPemKeys(keysLocation string) (*Jwks, error) { - jwks := &Jwks{} - - url, err := url.ParseRequestURI(keysLocation) - if err != nil { - return jwks, err - } - - resp, err := http.Get(url.String()) - if err != nil { - return jwks, err - } - defer resp.Body.Close() - - err = json.NewDecoder(resp.Body).Decode(jwks) - if err != nil { - return jwks, err - } - - cacheControlHeader := resp.Header.Get("Cache-Control") - expiresIn := getMaxAgeFromCacheHeader(cacheControlHeader) - jwks.expiresInTime = time.Now().Add(time.Duration(expiresIn) * time.Second) - - return jwks, nil -} - -func getPublicKey(token *jwt.Token, jwks *Jwks) (interface{}, error) { - // todo as we load the jkws when the server is starting, we should build a JKS map with the pem cert at the boot time - for k := range jwks.Keys { - if token.Header["kid"] != jwks.Keys[k].Kid { - continue - } - - if len(jwks.Keys[k].X5c) != 0 { - cert := "-----BEGIN CERTIFICATE-----\n" + jwks.Keys[k].X5c[0] + "\n-----END CERTIFICATE-----" - return jwt.ParseRSAPublicKeyFromPEM([]byte(cert)) - } - - if jwks.Keys[k].Kty == "RSA" { - return getPublicKeyFromRSA(jwks.Keys[k]) - } - if jwks.Keys[k].Kty == "EC" { - return getPublicKeyFromECDSA(jwks.Keys[k]) - } - } - - return nil, errKeyNotFound -} - -func getPublicKeyFromECDSA(jwk JSONWebKey) (publicKey *ecdsa.PublicKey, err error) { - if jwk.X == "" || jwk.Y == "" || jwk.Crv == "" { - return nil, fmt.Errorf("ecdsa key incomplete") - } - - var xCoordinate []byte - if xCoordinate, err = base64.RawURLEncoding.DecodeString(jwk.X); err != nil { - return nil, err - } - - var yCoordinate []byte - if yCoordinate, err = base64.RawURLEncoding.DecodeString(jwk.Y); err != nil { - return nil, err - } - - publicKey = &ecdsa.PublicKey{} - - var curve elliptic.Curve - switch jwk.Crv { - case p256: - curve = elliptic.P256() - case p384: - curve = elliptic.P384() - case p521: - curve = elliptic.P521() - } - - publicKey.Curve = curve - publicKey.X = big.NewInt(0).SetBytes(xCoordinate) - publicKey.Y = big.NewInt(0).SetBytes(yCoordinate) - - return publicKey, nil -} - -func getPublicKeyFromRSA(jwk JSONWebKey) (*rsa.PublicKey, error) { - decodedE, err := base64.RawURLEncoding.DecodeString(jwk.E) - if err != nil { - return nil, err - } - decodedN, err := base64.RawURLEncoding.DecodeString(jwk.N) - if err != nil { - return nil, err - } - - var n, e big.Int - e.SetBytes(decodedE) - n.SetBytes(decodedN) - - return &rsa.PublicKey{ - E: int(e.Int64()), - N: &n, - }, nil -} - -// getMaxAgeFromCacheHeader extracts max-age directive from the Cache-Control header -func getMaxAgeFromCacheHeader(cacheControl string) int { - // Split into individual directives - directives := strings.Split(cacheControl, ",") - - for _, directive := range directives { - directive = strings.TrimSpace(directive) - if strings.HasPrefix(directive, "max-age=") { - // Extract the max-age value - maxAgeStr := strings.TrimPrefix(directive, "max-age=") - maxAge, err := strconv.Atoi(maxAgeStr) - if err != nil { - return 0 - } - - return maxAge - } - } - - return 0 -} diff --git a/internals/shared/auth/manager.go b/internals/shared/auth/manager.go deleted file mode 100644 index b6c0649..0000000 --- a/internals/shared/auth/manager.go +++ /dev/null @@ -1,177 +0,0 @@ -package auth - -import ( - "context" - "crypto/sha256" - "encoding/base64" - "fmt" - "hash/crc32" - - "github.com/golang-jwt/jwt" - "github.com/netbirdio/management-refactor/base62" - nbjwt "github.com/netbirdio/management-refactor/management/server/auth/jwt" - nbcontext "github.com/netbirdio/management-refactor/management/server/context" - "github.com/netbirdio/management-refactor/management/server/store" - "github.com/netbirdio/management-refactor/management/server/types" - - "github.com/netbirdio/management-refactor/internals/modules/accounts/settings" - "github.com/netbirdio/management-refactor/internals/modules/users" - "github.com/netbirdio/management-refactor/internals/modules/users/pats" - - "github.com/netbirdio/management-refactor/internals/shared/db" -) - -var _ Manager = (*manager)(nil) - -type Manager interface { - ValidateAndParseToken(ctx context.Context, value string) (nbcontext.UserAuth, *jwt.Token, error) - EnsureUserAccessByJWTGroups(ctx context.Context, userAuth nbcontext.UserAuth, token *jwt.Token) (nbcontext.UserAuth, error) - MarkPATUsed(ctx context.Context, tokenID string) error - GetPATInfo(ctx context.Context, token string) (user *users.User, pat *pats.PersonalAccessToken, domain string, category string, err error) -} - -type manager struct { - userManager *users.Manager - settingsManager *settings.Manager - - validator *nbjwt.Validator - extractor *nbjwt.ClaimsExtractor -} - -func NewManager(userManager *users.Manager, settingsManager *settings.Manager, issuer, audience, keysLocation, userIdClaim string, allAudiences []string, idpRefreshKeys bool) Manager { - // @note if invalid/missing parameters are sent the validator will instantiate - // but it will fail when validating and parsing the token - jwtValidator := nbjwt.NewValidator( - issuer, - allAudiences, - keysLocation, - idpRefreshKeys, - ) - - claimsExtractor := nbjwt.NewClaimsExtractor( - nbjwt.WithAudience(audience), - nbjwt.WithUserIDClaim(userIdClaim), - ) - - return &manager{ - userManager: userManager, - settingsManager: settingsManager, - - validator: jwtValidator, - extractor: claimsExtractor, - } -} - -func (m *manager) ValidateAndParseToken(ctx context.Context, value string) (nbcontext.UserAuth, *jwt.Token, error) { - token, err := m.validator.ValidateAndParse(ctx, value) - if err != nil { - return nbcontext.UserAuth{}, nil, err - } - - userAuth, err := m.extractor.ToUserAuth(token) - if err != nil { - return nbcontext.UserAuth{}, nil, err - } - return userAuth, token, err -} - -func (m *manager) EnsureUserAccessByJWTGroups(ctx context.Context, userAuth nbcontext.UserAuth, token *jwt.Token) (nbcontext.UserAuth, error) { - if userAuth.IsChild || userAuth.IsPAT { - return userAuth, nil - } - - settings, err := m.settingsManager.GetSettings(ctx, nil, db.LockingStrengthShare, userAuth.AccountId, userAuth.UserId) - if err != nil { - return userAuth, err - } - - // Ensures JWT group synchronization to the management is enabled before, - // filtering access based on the allowed groups. - if settings != nil && settings.JWTGroupsEnabled { - userAuth.Groups = m.extractor.ToGroups(token, settings.JWTGroupsClaimName) - if allowedGroups := settings.JWTAllowGroups; len(allowedGroups) > 0 { - if !userHasAllowedGroup(allowedGroups, userAuth.Groups) { - return userAuth, fmt.Errorf("user does not belong to any of the allowed JWT groups") - } - } - } - - return userAuth, nil -} - -// MarkPATUsed marks a personal access token as used -func (am *manager) MarkPATUsed(ctx context.Context, tokenID string) error { - return am.store.MarkPATUsed(ctx, store.LockingStrengthUpdate, tokenID) -} - -// GetPATInfo retrieves user, personal access token, domain, and category details from a personal access token. -func (am *manager) GetPATInfo(ctx context.Context, token string) (user *types.User, pat *pattypes.PersonalAccessToken, domain string, category string, err error) { - user, pat, err = am.extractPATFromToken(ctx, token) - if err != nil { - return nil, nil, "", "", err - } - - domain, category, err = am.store.GetAccountDomainAndCategory(ctx, store.LockingStrengthShare, user.AccountID) - if err != nil { - return nil, nil, "", "", err - } - - return user, pat, domain, category, nil -} - -// extractPATFromToken validates the token structure and retrieves associated User and PAT. -func (am *manager) extractPATFromToken(ctx context.Context, token string) (*types.User, *pattypes.PersonalAccessToken, error) { - if len(token) != pattypes.PATLength { - return nil, nil, fmt.Errorf("PAT has incorrect length") - } - - prefix := token[:len(pattypes.PATPrefix)] - if prefix != pattypes.PATPrefix { - return nil, nil, fmt.Errorf("PAT has wrong prefix") - } - secret := token[len(pattypes.PATPrefix) : len(pattypes.PATPrefix)+pattypes.PATSecretLength] - encodedChecksum := token[len(pattypes.PATPrefix)+pattypes.PATSecretLength : len(pattypes.PATPrefix)+pattypes.PATSecretLength+pattypes.PATChecksumLength] - - verificationChecksum, err := base62.Decode(encodedChecksum) - if err != nil { - return nil, nil, fmt.Errorf("PAT checksum decoding failed: %w", err) - } - - secretChecksum := crc32.ChecksumIEEE([]byte(secret)) - if secretChecksum != verificationChecksum { - return nil, nil, fmt.Errorf("PAT checksum does not match") - } - - hashedToken := sha256.Sum256([]byte(token)) - encodedHashedToken := base64.StdEncoding.EncodeToString(hashedToken[:]) - - var user *types.User - var pat *pattypes.PersonalAccessToken - - err = am.store.ExecuteInTransaction(ctx, func(transaction store.Store) error { - pat, err = transaction.GetPATByHashedToken(ctx, store.LockingStrengthShare, encodedHashedToken) - if err != nil { - return err - } - - user, err = transaction.GetUserByPATID(ctx, store.LockingStrengthShare, pat.ID) - return err - }) - if err != nil { - return nil, nil, err - } - - return user, pat, nil -} - -// userHasAllowedGroup checks if a user belongs to any of the allowed groups. -func userHasAllowedGroup(allowedGroups []string, userGroups []string) bool { - for _, userGroup := range userGroups { - for _, allowedGroup := range allowedGroups { - if userGroup == allowedGroup { - return true - } - } - } - return false -} diff --git a/internals/shared/auth/manager_mock.go b/internals/shared/auth/manager_mock.go deleted file mode 100644 index 4d290c9..0000000 --- a/internals/shared/auth/manager_mock.go +++ /dev/null @@ -1,54 +0,0 @@ -package auth - -import ( - "context" - - "github.com/golang-jwt/jwt" - - nbcontext "github.com/netbirdio/management-refactor/management/server/context" - "github.com/netbirdio/management-refactor/management/server/types" -) - -var ( - _ Manager = (*MockManager)(nil) -) - -// @note really dislike this mocking approach but rather than have to do additional test refactoring. -type MockManager struct { - ValidateAndParseTokenFunc func(ctx context.Context, value string) (nbcontext.UserAuth, *jwt.Token, error) - EnsureUserAccessByJWTGroupsFunc func(ctx context.Context, userAuth nbcontext.UserAuth, token *jwt.Token) (nbcontext.UserAuth, error) - MarkPATUsedFunc func(ctx context.Context, tokenID string) error - GetPATInfoFunc func(ctx context.Context, token string) (user *types.User, pat *types.PersonalAccessToken, domain string, category string, err error) -} - -// EnsureUserAccessByJWTGroups implements Manager. -func (m *MockManager) EnsureUserAccessByJWTGroups(ctx context.Context, userAuth nbcontext.UserAuth, token *jwt.Token) (nbcontext.UserAuth, error) { - if m.EnsureUserAccessByJWTGroupsFunc != nil { - return m.EnsureUserAccessByJWTGroupsFunc(ctx, userAuth, token) - } - return nbcontext.UserAuth{}, nil -} - -// GetPATInfo implements Manager. -func (m *MockManager) GetPATInfo(ctx context.Context, token string) (user *types.User, pat *types.PersonalAccessToken, domain string, category string, err error) { - if m.GetPATInfoFunc != nil { - return m.GetPATInfoFunc(ctx, token) - } - return &types.User{}, &types.PersonalAccessToken{}, "", "", nil -} - -// MarkPATUsed implements Manager. -func (m *MockManager) MarkPATUsed(ctx context.Context, tokenID string) error { - if m.MarkPATUsedFunc != nil { - return m.MarkPATUsedFunc(ctx, tokenID) - } - return nil -} - -// ValidateAndParseToken implements Manager. -func (m *MockManager) ValidateAndParseToken(ctx context.Context, value string) (nbcontext.UserAuth, *jwt.Token, error) { - if m.ValidateAndParseTokenFunc != nil { - return m.ValidateAndParseTokenFunc(ctx, value) - } - return nbcontext.UserAuth{}, &jwt.Token{}, nil -} diff --git a/internals/shared/auth/manager_test.go b/internals/shared/auth/manager_test.go deleted file mode 100644 index 900953d..0000000 --- a/internals/shared/auth/manager_test.go +++ /dev/null @@ -1,407 +0,0 @@ -package auth_test - -import ( - "context" - "crypto/sha256" - "encoding/base64" - "fmt" - "net/http" - "net/http/httptest" - "os" - "strings" - "testing" - "time" - - "github.com/golang-jwt/jwt" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" - - "github.com/netbirdio/management-refactor/management/server/auth" - nbjwt "github.com/netbirdio/management-refactor/management/server/auth/jwt" - nbcontext "github.com/netbirdio/management-refactor/management/server/context" - "github.com/netbirdio/management-refactor/management/server/store" - "github.com/netbirdio/management-refactor/management/server/types" -) - -func TestAuthManager_GetAccountInfoFromPAT(t *testing.T) { - store, cleanup, err := store.NewTestStoreFromSQL(context.Background(), "", t.TempDir()) - if err != nil { - t.Fatalf("Error when creating store: %s", err) - } - t.Cleanup(cleanup) - - token := "nbp_9999EUDNdkeusjentDLSJEn1902u84390W6W" - hashedToken := sha256.Sum256([]byte(token)) - encodedHashedToken := base64.StdEncoding.EncodeToString(hashedToken[:]) - account := &types.Account{ - Id: "account_id", - Users: map[string]*types.User{"someUser": { - Id: "someUser", - PATs: map[string]*types.PersonalAccessToken{ - "tokenId": { - ID: "tokenId", - UserID: "someUser", - HashedToken: encodedHashedToken, - }, - }, - }}, - } - - err = store.SaveAccount(context.Background(), account) - if err != nil { - t.Fatalf("Error when saving account: %s", err) - } - - manager := auth.NewManager(store, "", "", "", "", []string{}, false) - - user, pat, _, _, err := manager.GetPATInfo(context.Background(), token) - if err != nil { - t.Fatalf("Error when getting Account from PAT: %s", err) - } - - assert.Equal(t, "account_id", user.AccountID) - assert.Equal(t, "someUser", user.Id) - assert.Equal(t, account.Users["someUser"].PATs["tokenId"].ID, pat.ID) -} - -func TestAuthManager_MarkPATUsed(t *testing.T) { - store, cleanup, err := store.NewTestStoreFromSQL(context.Background(), "", t.TempDir()) - if err != nil { - t.Fatalf("Error when creating store: %s", err) - } - t.Cleanup(cleanup) - - token := "nbp_9999EUDNdkeusjentDLSJEn1902u84390W6W" - hashedToken := sha256.Sum256([]byte(token)) - encodedHashedToken := base64.StdEncoding.EncodeToString(hashedToken[:]) - account := &types.Account{ - Id: "account_id", - Users: map[string]*types.User{"someUser": { - Id: "someUser", - PATs: map[string]*types.PersonalAccessToken{ - "tokenId": { - ID: "tokenId", - HashedToken: encodedHashedToken, - }, - }, - }}, - } - - err = store.SaveAccount(context.Background(), account) - if err != nil { - t.Fatalf("Error when saving account: %s", err) - } - - manager := auth.NewManager(store, "", "", "", "", []string{}, false) - - err = manager.MarkPATUsed(context.Background(), "tokenId") - if err != nil { - t.Fatalf("Error when marking PAT used: %s", err) - } - - account, err = store.GetAccount(context.Background(), "account_id") - if err != nil { - t.Fatalf("Error when getting account: %s", err) - } - assert.True(t, !account.Users["someUser"].PATs["tokenId"].GetLastUsed().IsZero()) -} - -func TestAuthManager_EnsureUserAccessByJWTGroups(t *testing.T) { - store, cleanup, err := store.NewTestStoreFromSQL(context.Background(), "", t.TempDir()) - if err != nil { - t.Fatalf("Error when creating store: %s", err) - } - t.Cleanup(cleanup) - - userId := "user-id" - domain := "test.domain" - - account := &types.Account{ - Id: "account_id", - Domain: domain, - Users: map[string]*types.User{"someUser": { - Id: "someUser", - }}, - Settings: &types.Settings{}, - } - - err = store.SaveAccount(context.Background(), account) - if err != nil { - t.Fatalf("Error when saving account: %s", err) - } - - // this has been validated and parsed by ValidateAndParseToken - userAuth := nbcontext.UserAuth{ - AccountId: account.Id, - Domain: domain, - UserId: userId, - DomainCategory: "test-category", - // Groups: []string{"group1", "group2"}, - } - - // these tests only assert groups are parsed from token as per account settings - token := jwt.NewWithClaims(jwt.SigningMethodHS256, jwt.MapClaims{"idp-groups": []interface{}{"group1", "group2"}}) - - manager := auth.NewManager(store, "", "", "", "", []string{}, false) - - t.Run("JWT groups disabled", func(t *testing.T) { - userAuth, err := manager.EnsureUserAccessByJWTGroups(context.Background(), userAuth, token) - require.NoError(t, err, "ensure user access by JWT groups failed") - require.Len(t, userAuth.Groups, 0, "account not enabled to ensure access by groups") - }) - - t.Run("User impersonated", func(t *testing.T) { - userAuth, err := manager.EnsureUserAccessByJWTGroups(context.Background(), userAuth, token) - require.NoError(t, err, "ensure user access by JWT groups failed") - require.Len(t, userAuth.Groups, 0, "account not enabled to ensure access by groups") - }) - - t.Run("User PAT", func(t *testing.T) { - userAuth, err := manager.EnsureUserAccessByJWTGroups(context.Background(), userAuth, token) - require.NoError(t, err, "ensure user access by JWT groups failed") - require.Len(t, userAuth.Groups, 0, "account not enabled to ensure access by groups") - }) - - t.Run("JWT groups enabled without claim name", func(t *testing.T) { - account.Settings.JWTGroupsEnabled = true - err := store.SaveAccount(context.Background(), account) - require.NoError(t, err, "save account failed") - - userAuth, err := manager.EnsureUserAccessByJWTGroups(context.Background(), userAuth, token) - require.NoError(t, err, "ensure user access by JWT groups failed") - require.Len(t, userAuth.Groups, 0, "account missing groups claim name") - }) - - t.Run("JWT groups enabled without allowed groups", func(t *testing.T) { - account.Settings.JWTGroupsEnabled = true - account.Settings.JWTGroupsClaimName = "idp-groups" - err := store.SaveAccount(context.Background(), account) - require.NoError(t, err, "save account failed") - - userAuth, err := manager.EnsureUserAccessByJWTGroups(context.Background(), userAuth, token) - require.NoError(t, err, "ensure user access by JWT groups failed") - require.Equal(t, []string{"group1", "group2"}, userAuth.Groups, "group parsed do not match") - }) - - t.Run("User in allowed JWT groups", func(t *testing.T) { - account.Settings.JWTGroupsEnabled = true - account.Settings.JWTGroupsClaimName = "idp-groups" - account.Settings.JWTAllowGroups = []string{"group1"} - err := store.SaveAccount(context.Background(), account) - require.NoError(t, err, "save account failed") - - userAuth, err := manager.EnsureUserAccessByJWTGroups(context.Background(), userAuth, token) - require.NoError(t, err, "ensure user access by JWT groups failed") - - require.Equal(t, []string{"group1", "group2"}, userAuth.Groups, "group parsed do not match") - }) - - t.Run("User not in allowed JWT groups", func(t *testing.T) { - account.Settings.JWTGroupsEnabled = true - account.Settings.JWTGroupsClaimName = "idp-groups" - account.Settings.JWTAllowGroups = []string{"not-a-group"} - err := store.SaveAccount(context.Background(), account) - require.NoError(t, err, "save account failed") - - _, err = manager.EnsureUserAccessByJWTGroups(context.Background(), userAuth, token) - require.Error(t, err, "ensure user access is not in allowed groups") - }) -} - -func TestAuthManager_ValidateAndParseToken(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Add("Cache-Control", "max-age=30") // set a 30s expiry to these keys - http.ServeFile(w, r, "test_data/jwks.json") - })) - defer server.Close() - - issuer := "http://issuer.local" - audience := "http://audience.local" - userIdClaim := "" // defaults to "sub" - - // we're only testing with RSA256 - keyData, _ := os.ReadFile("test_data/sample_key") - key, _ := jwt.ParseRSAPrivateKeyFromPEM(keyData) - keyId := "test-key" - - // note, we can use a nil store because ValidateAndParseToken does not use it in it's flow - manager := auth.NewManager(nil, issuer, audience, server.URL, userIdClaim, []string{audience}, false) - - customClaim := func(name string) string { - return fmt.Sprintf("%s/%s", audience, name) - } - - lastLogin := time.Date(2025, 2, 12, 14, 25, 26, 0, time.UTC) // "2025-02-12T14:25:26.186Z" - - tests := []struct { - name string - tokenFunc func() string - expected *nbcontext.UserAuth // nil indicates expected error - }{ - { - name: "Valid with custom claims", - tokenFunc: func() string { - token := jwt.New(jwt.SigningMethodRS256) - token.Header["kid"] = keyId - token.Claims = jwt.MapClaims{ - "iss": issuer, - "aud": []string{audience}, - "iat": time.Now().Unix(), - "exp": time.Now().Add(time.Hour * 1).Unix(), - "sub": "user-id|123", - customClaim(nbjwt.AccountIDSuffix): "account-id|567", - customClaim(nbjwt.DomainIDSuffix): "http://localhost", - customClaim(nbjwt.DomainCategorySuffix): "private", - customClaim(nbjwt.LastLoginSuffix): lastLogin.Format(time.RFC3339), - customClaim(nbjwt.Invited): false, - } - tokenString, _ := token.SignedString(key) - return tokenString - }, - expected: &nbcontext.UserAuth{ - UserId: "user-id|123", - AccountId: "account-id|567", - Domain: "http://localhost", - DomainCategory: "private", - LastLogin: lastLogin, - Invited: false, - }, - }, - { - name: "Valid without custom claims", - tokenFunc: func() string { - token := jwt.New(jwt.SigningMethodRS256) - token.Header["kid"] = keyId - token.Claims = jwt.MapClaims{ - "iss": issuer, - "aud": []string{audience}, - "iat": time.Now().Unix(), - "exp": time.Now().Add(time.Hour).Unix(), - "sub": "user-id|123", - } - tokenString, _ := token.SignedString(key) - return tokenString - }, - expected: &nbcontext.UserAuth{ - UserId: "user-id|123", - }, - }, - { - name: "Expired token", - tokenFunc: func() string { - token := jwt.New(jwt.SigningMethodRS256) - token.Header["kid"] = keyId - token.Claims = jwt.MapClaims{ - "iss": issuer, - "aud": []string{audience}, - "iat": time.Now().Add(time.Hour * -2).Unix(), - "exp": time.Now().Add(time.Hour * -1).Unix(), - "sub": "user-id|123", - } - tokenString, _ := token.SignedString(key) - return tokenString - }, - }, - { - name: "Not yet valid", - tokenFunc: func() string { - token := jwt.New(jwt.SigningMethodRS256) - token.Header["kid"] = keyId - token.Claims = jwt.MapClaims{ - "iss": issuer, - "aud": []string{audience}, - "iat": time.Now().Add(time.Hour).Unix(), - "exp": time.Now().Add(time.Hour * 2).Unix(), - "sub": "user-id|123", - } - tokenString, _ := token.SignedString(key) - return tokenString - }, - }, - { - name: "Invalid signature", - tokenFunc: func() string { - token := jwt.New(jwt.SigningMethodRS256) - token.Header["kid"] = keyId - token.Claims = jwt.MapClaims{ - "iss": issuer, - "aud": []string{audience}, - "iat": time.Now().Unix(), - "exp": time.Now().Add(time.Hour).Unix(), - "sub": "user-id|123", - } - tokenString, _ := token.SignedString(key) - parts := strings.Split(tokenString, ".") - parts[2] = "invalid-signature" - return strings.Join(parts, ".") - }, - }, - { - name: "Invalid issuer", - tokenFunc: func() string { - token := jwt.New(jwt.SigningMethodRS256) - token.Header["kid"] = keyId - token.Claims = jwt.MapClaims{ - "iss": "not-the-issuer", - "aud": []string{audience}, - "iat": time.Now().Unix(), - "exp": time.Now().Add(time.Hour).Unix(), - "sub": "user-id|123", - } - tokenString, _ := token.SignedString(key) - return tokenString - }, - }, - { - name: "Invalid audience", - tokenFunc: func() string { - token := jwt.New(jwt.SigningMethodRS256) - token.Header["kid"] = keyId - token.Claims = jwt.MapClaims{ - "iss": issuer, - "aud": []string{"not-the-audience"}, - "iat": time.Now().Unix(), - "exp": time.Now().Add(time.Hour).Unix(), - "sub": "user-id|123", - } - tokenString, _ := token.SignedString(key) - return tokenString - }, - }, - { - name: "Invalid user claim", - tokenFunc: func() string { - token := jwt.New(jwt.SigningMethodRS256) - token.Header["kid"] = keyId - token.Claims = jwt.MapClaims{ - "iss": issuer, - "aud": []string{audience}, - "iat": time.Now().Unix(), - "exp": time.Now().Add(time.Hour).Unix(), - "not-sub": "user-id|123", - } - tokenString, _ := token.SignedString(key) - return tokenString - }, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - tokenString := tt.tokenFunc() - - userAuth, token, err := manager.ValidateAndParseToken(context.Background(), tokenString) - - if tt.expected != nil { - assert.NoError(t, err) - assert.True(t, token.Valid) - assert.Equal(t, *tt.expected, userAuth) - } else { - assert.Error(t, err) - assert.Nil(t, token) - assert.Empty(t, userAuth) - } - }) - } - -} diff --git a/internals/shared/auth/test_data/jwks.json b/internals/shared/auth/test_data/jwks.json deleted file mode 100644 index 8080f55..0000000 --- a/internals/shared/auth/test_data/jwks.json +++ /dev/null @@ -1,11 +0,0 @@ -{ - "keys": [ - { - "kty": "RSA", - "kid": "test-key", - "use": "sig", - "n": "4f5wg5l2hKsTeNem_V41fGnJm6gOdrj8ym3rFkEU_wT8RDtnSgFEZOQpHEgQ7JL38xUfU0Y3g6aYw9QT0hJ7mCpz9Er5qLaMXJwZxzHzAahlfA0icqabvJOMvQtzD6uQv6wPEyZtDTWiQi9AXwBpHssPnpYGIn20ZZuNlX2BrClciHhCPUIIZOQn_MmqTD31jSyjoQoV7MhhMTATKJx2XrHhR-1DcKJzQBSTAGnpYVaqpsARap-nwRipr3nUTuxyGohBTSmjJ2usSeQXHI3bODIRe1AuTyHceAbewn8b462yEWKARdpd9AjQW5SIVPfdsz5B6GlYQ5LdYKtznTuy7w", - "e": "AQAB" - } - ] -} \ No newline at end of file diff --git a/internals/shared/auth/test_data/sample_key b/internals/shared/auth/test_data/sample_key deleted file mode 100644 index e69284a..0000000 --- a/internals/shared/auth/test_data/sample_key +++ /dev/null @@ -1,27 +0,0 @@ ------BEGIN RSA PRIVATE KEY----- -MIIEowIBAAKCAQEA4f5wg5l2hKsTeNem/V41fGnJm6gOdrj8ym3rFkEU/wT8RDtn -SgFEZOQpHEgQ7JL38xUfU0Y3g6aYw9QT0hJ7mCpz9Er5qLaMXJwZxzHzAahlfA0i -cqabvJOMvQtzD6uQv6wPEyZtDTWiQi9AXwBpHssPnpYGIn20ZZuNlX2BrClciHhC -PUIIZOQn/MmqTD31jSyjoQoV7MhhMTATKJx2XrHhR+1DcKJzQBSTAGnpYVaqpsAR -ap+nwRipr3nUTuxyGohBTSmjJ2usSeQXHI3bODIRe1AuTyHceAbewn8b462yEWKA -Rdpd9AjQW5SIVPfdsz5B6GlYQ5LdYKtznTuy7wIDAQABAoIBAQCwia1k7+2oZ2d3 -n6agCAbqIE1QXfCmh41ZqJHbOY3oRQG3X1wpcGH4Gk+O+zDVTV2JszdcOt7E5dAy -MaomETAhRxB7hlIOnEN7WKm+dGNrKRvV0wDU5ReFMRHg31/Lnu8c+5BvGjZX+ky9 -POIhFFYJqwCRlopGSUIxmVj5rSgtzk3iWOQXr+ah1bjEXvlxDOWkHN6YfpV5ThdE -KdBIPGEVqa63r9n2h+qazKrtiRqJqGnOrHzOECYbRFYhexsNFz7YT02xdfSHn7gM -IvabDDP/Qp0PjE1jdouiMaFHYnLBbgvlnZW9yuVf/rpXTUq/njxIXMmvmEyyvSDn -FcFikB8pAoGBAPF77hK4m3/rdGT7X8a/gwvZ2R121aBcdPwEaUhvj/36dx596zvY -mEOjrWfZhF083/nYWE2kVquj2wjs+otCLfifEEgXcVPTnEOPO9Zg3uNSL0nNQghj -FuD3iGLTUBCtM66oTe0jLSslHe8gLGEQqyMzHOzYxNqibxcOZIe8Qt0NAoGBAO+U -I5+XWjWEgDmvyC3TrOSf/KCGjtu0TSv30ipv27bDLMrpvPmD/5lpptTFwcxvVhCs -2b+chCjlghFSWFbBULBrfci2FtliClOVMYrlNBdUSJhf3aYSG2Doe6Bgt1n2CpNn -/iu37Y3NfemZBJA7hNl4dYe+f+uzM87cdQ214+jrAoGAXA0XxX8ll2+ToOLJsaNT -OvNB9h9Uc5qK5X5w+7G7O998BN2PC/MWp8H+2fVqpXgNENpNXttkRm1hk1dych86 -EunfdPuqsX+as44oCyJGFHVBnWpm33eWQw9YqANRI+pCJzP08I5WK3osnPiwshd+ -hR54yjgfYhBFNI7B95PmEQkCgYBzFSz7h1+s34Ycr8SvxsOBWxymG5zaCsUbPsL0 -4aCgLScCHb9J+E86aVbbVFdglYa5Id7DPTL61ixhl7WZjujspeXZGSbmq0Kcnckb -mDgqkLECiOJW2NHP/j0McAkDLL4tysF8TLDO8gvuvzNC+WQ6drO2ThrypLVZQ+ry -eBIPmwKBgEZxhqa0gVvHQG/7Od69KWj4eJP28kq13RhKay8JOoN0vPmspXJo1HY3 -CKuHRG+AP579dncdUnOMvfXOtkdM4vk0+hWASBQzM9xzVcztCa+koAugjVaLS9A+ -9uQoqEeVNTckxx0S2bYevRy7hGQmUJTyQm3j1zEUR5jpdbL83Fbq ------END RSA PRIVATE KEY----- \ No newline at end of file diff --git a/internals/shared/auth/test_data/sample_key.pub b/internals/shared/auth/test_data/sample_key.pub deleted file mode 100644 index d5b7f71..0000000 --- a/internals/shared/auth/test_data/sample_key.pub +++ /dev/null @@ -1,9 +0,0 @@ ------BEGIN PUBLIC KEY----- -MIIBIjANBgkqhkiG9w0BAQEFAAOCAQ8AMIIBCgKCAQEA4f5wg5l2hKsTeNem/V41 -fGnJm6gOdrj8ym3rFkEU/wT8RDtnSgFEZOQpHEgQ7JL38xUfU0Y3g6aYw9QT0hJ7 -mCpz9Er5qLaMXJwZxzHzAahlfA0icqabvJOMvQtzD6uQv6wPEyZtDTWiQi9AXwBp -HssPnpYGIn20ZZuNlX2BrClciHhCPUIIZOQn/MmqTD31jSyjoQoV7MhhMTATKJx2 -XrHhR+1DcKJzQBSTAGnpYVaqpsARap+nwRipr3nUTuxyGohBTSmjJ2usSeQXHI3b -ODIRe1AuTyHceAbewn8b462yEWKARdpd9AjQW5SIVPfdsz5B6GlYQ5LdYKtznTuy -7wIDAQAB ------END PUBLIC KEY----- \ No newline at end of file diff --git a/internals/shared/event-bus/bus.go b/internals/shared/event-bus/bus.go deleted file mode 100644 index 4e436a4..0000000 --- a/internals/shared/event-bus/bus.go +++ /dev/null @@ -1,16 +0,0 @@ -package event_bus - -type BaseEvent struct { - AccountID string -} - -type PeerCreated struct { - BaseEvent - PeerID string -} - -type PortAllocated struct { - BaseEvent - PeerID string - Port int -} diff --git a/internals/shared/hook/README.md b/internals/shared/hook/README.md deleted file mode 100644 index e1e2610..0000000 --- a/internals/shared/hook/README.md +++ /dev/null @@ -1,10 +0,0 @@ -# Pocketbase Hooks - -This package original authoring is by Gani Georgiev the creator and mantainer of Pocketbase, I've merely ported and adapted it by removing dependencies. You can read more about it and its use on Pocketbase [here](https://pocketbase.io/docs/go-event-hooks/). - -This is a very powerful event/hook system, especially this: - -> All hook handler functions share the same func(e T) error signature and expect the user to call e.Next() if they want to proceed with the execution chain. - -I've included it in our cloud repository and implemented the concept with a sub-set of billing events as a PoC of how we could re-architect Management for a cleaner developer experience and faster development by decoupling server bootstrapping (dependencies), serve (http and gRPC) and actual features (billing, idp, event streaming, etc). - diff --git a/internals/shared/hook/event.go b/internals/shared/hook/event.go deleted file mode 100644 index 12bceec..0000000 --- a/internals/shared/hook/event.go +++ /dev/null @@ -1,45 +0,0 @@ -package hook - -// Resolver defines a common interface for a Hook event (see [Event]). -type Resolver interface { - // Next triggers the next handler in the hook's chain (if any). - Next() error - - // note: kept only for the generic interface; may get removed in the future - nextFunc() func() error - setNextFunc(f func() error) -} - -var _ Resolver = (*Event)(nil) - -// Event implements [Resolver] and it is intended to be used as a base -// Hook event that you can embed in your custom typed event structs. -// -// Example: -// -// type CustomEvent struct { -// hook.Event -// -// SomeField int -// } -type Event struct { - next func() error -} - -// Next calls the next hook handler. -func (e *Event) Next() error { - if e.next != nil { - return e.next() - } - return nil -} - -// nextFunc returns the function that Next calls. -func (e *Event) nextFunc() func() error { - return e.next -} - -// setNextFunc sets the function that Next calls. -func (e *Event) setNextFunc(f func() error) { - e.next = f -} diff --git a/internals/shared/hook/event_test.go b/internals/shared/hook/event_test.go deleted file mode 100644 index 89b15ef..0000000 --- a/internals/shared/hook/event_test.go +++ /dev/null @@ -1,29 +0,0 @@ -package hook - -import "testing" - -func TestEventNext(t *testing.T) { - calls := 0 - - e := Event{} - - if e.nextFunc() != nil { - t.Fatalf("Expected nextFunc to be nil") - } - - e.setNextFunc(func() error { - calls++ - return nil - }) - - if e.nextFunc() == nil { - t.Fatalf("Expected nextFunc to be non-nil") - } - - e.Next() - e.Next() - - if calls != 2 { - t.Fatalf("Expected %d calls, got %d", 2, calls) - } -} diff --git a/internals/shared/hook/hook.go b/internals/shared/hook/hook.go deleted file mode 100644 index 07e0dde..0000000 --- a/internals/shared/hook/hook.go +++ /dev/null @@ -1,178 +0,0 @@ -package hook - -import ( - "sort" - "sync" - - "github.com/rs/xid" -) - -// Handler defines a single Hook handler. -// Multiple handlers can share the same id. -// If Id is not explicitly set it will be autogenerated by Hook.Add and Hook.AddHandler. -type Handler[T Resolver] struct { - // Func defines the handler function to execute. - // - // Note that users need to call e.Next() in order to proceed with - // the execution of the hook chain. - Func func(T) error - - // Id is the unique identifier of the handler. - // - // It could be used later to remove the handler from a hook via [Hook.Remove]. - // - // If missing, an autogenerated value will be assigned when adding - // the handler to a hook. - Id string - - // Priority allows changing the default exec priority of the handler within a hook. - // - // If 0, the handler will be executed in the same order it was registered. - Priority int -} - -// Hook defines a generic concurrent safe structure for managing event hooks. -// -// When using custom event it must embed the base [hook.Event]. -// -// Example: -// -// type CustomEvent struct { -// hook.Event -// SomeField int -// } -// -// h := Hook[*CustomEvent]{} -// -// h.BindFunc(func(e *CustomEvent) error { -// println(e.SomeField) -// -// return e.Next() -// }) -// -// h.Trigger(&CustomEvent{ SomeField: 123 }) -type Hook[T Resolver] struct { - handlers []*Handler[T] - mu sync.RWMutex -} - -// Bind registers the provided handler to the current hooks queue. -// -// If handler.Id is empty it is updated with autogenerated value. -// -// If a handler from the current hook list has Id matching handler.Id -// then the old handler is replaced with the new one. -func (h *Hook[T]) Bind(handler *Handler[T]) string { - h.mu.Lock() - defer h.mu.Unlock() - - var exists bool - - if handler.Id == "" { - handler.Id = generateHookId() - - // ensure that it doesn't exist - DUPLICATE_CHECK: - for _, existing := range h.handlers { - if existing.Id == handler.Id { - handler.Id = generateHookId() - goto DUPLICATE_CHECK - } - } - } else { - // replace existing - for i, existing := range h.handlers { - if existing.Id == handler.Id { - h.handlers[i] = handler - exists = true - break - } - } - } - - // append new - if !exists { - h.handlers = append(h.handlers, handler) - } - - // sort handlers by Priority, preserving the original order of equal items - sort.SliceStable(h.handlers, func(i, j int) bool { - return h.handlers[i].Priority < h.handlers[j].Priority - }) - - return handler.Id -} - -// BindFunc is similar to Bind but registers a new handler from just the provided function. -// -// The registered handler is added with a default 0 priority and the id will be autogenerated. -// -// If you want to register a handler with custom priority or id use the [Hook.Bind] method. -func (h *Hook[T]) BindFunc(fn func(e T) error) string { - return h.Bind(&Handler[T]{Func: fn}) -} - -// Unbind removes one or many hook handler by their id. -func (h *Hook[T]) Unbind(idsToRemove ...string) { - h.mu.Lock() - defer h.mu.Unlock() - - for _, id := range idsToRemove { - for i := len(h.handlers) - 1; i >= 0; i-- { - if h.handlers[i].Id == id { - h.handlers = append(h.handlers[:i], h.handlers[i+1:]...) - break // for now stop on the first occurrence since we don't allow handlers with duplicated ids - } - } - } -} - -// UnbindAll removes all registered handlers. -func (h *Hook[T]) UnbindAll() { - h.mu.Lock() - defer h.mu.Unlock() - - h.handlers = nil -} - -// Length returns to total number of registered hook handlers. -func (h *Hook[T]) Length() int { - h.mu.RLock() - defer h.mu.RUnlock() - - return len(h.handlers) -} - -// Trigger executes all registered hook handlers one by one -// with the specified event as an argument. -// -// Optionally, this method allows also to register additional one off -// handler funcs that will be temporary appended to the handlers queue. -// -// NB! Each hook handler must call event.Next() in order the hook chain to proceed. -func (h *Hook[T]) Trigger(event T, oneOffHandlerFuncs ...func(T) error) error { - h.mu.RLock() - handlers := make([]func(T) error, 0, len(h.handlers)+len(oneOffHandlerFuncs)) - for _, handler := range h.handlers { - handlers = append(handlers, handler.Func) - } - handlers = append(handlers, oneOffHandlerFuncs...) - h.mu.RUnlock() - - event.setNextFunc(nil) // reset in case the event is being reused - - for i := len(handlers) - 1; i >= 0; i-- { - i := i - old := event.nextFunc() - event.setNextFunc(func() error { - event.setNextFunc(old) - return handlers[i](event) - }) - } - - return event.Next() -} - -func generateHookId() string { - return xid.New().String() -} diff --git a/internals/shared/hook/hook_test.go b/internals/shared/hook/hook_test.go deleted file mode 100644 index 36bd945..0000000 --- a/internals/shared/hook/hook_test.go +++ /dev/null @@ -1,162 +0,0 @@ -package hook - -import ( - "errors" - "testing" -) - -func TestHookAddHandlerAndAdd(t *testing.T) { - calls := "" - - h := Hook[*Event]{} - - h.BindFunc(func(e *Event) error { calls += "1"; return e.Next() }) - h.BindFunc(func(e *Event) error { calls += "2"; return e.Next() }) - h3Id := h.BindFunc(func(e *Event) error { calls += "3"; return e.Next() }) - h.Bind(&Handler[*Event]{ - Id: h3Id, // should replace 3 - Func: func(e *Event) error { calls += "3'"; return e.Next() }, - }) - h.Bind(&Handler[*Event]{ - Func: func(e *Event) error { calls += "4"; return e.Next() }, - Priority: -2, - }) - h.Bind(&Handler[*Event]{ - Func: func(e *Event) error { calls += "5"; return e.Next() }, - Priority: -1, - }) - h.Bind(&Handler[*Event]{ - Func: func(e *Event) error { calls += "6"; return e.Next() }, - }) - h.Bind(&Handler[*Event]{ - Func: func(e *Event) error { calls += "7"; e.Next(); return errors.New("test") }, // error shouldn't stop the chain - }) - - h.Trigger( - &Event{}, - func(e *Event) error { calls += "8"; return e.Next() }, - func(e *Event) error { calls += "9"; return nil }, // skip next - func(e *Event) error { calls += "10"; return e.Next() }, - ) - - if total := len(h.handlers); total != 7 { - t.Fatalf("Expected %d handlers, found %d", 7, total) - } - - expectedCalls := "45123'6789" - - if calls != expectedCalls { - t.Fatalf("Expected calls sequence %q, got %q", expectedCalls, calls) - } -} - -func TestHookLength(t *testing.T) { - h := Hook[*Event]{} - - if l := h.Length(); l != 0 { - t.Fatalf("Expected 0 hook handlers, got %d", l) - } - - h.BindFunc(func(e *Event) error { return e.Next() }) - h.BindFunc(func(e *Event) error { return e.Next() }) - - if l := h.Length(); l != 2 { - t.Fatalf("Expected 2 hook handlers, got %d", l) - } -} - -func TestHookUnbind(t *testing.T) { - h := Hook[*Event]{} - - calls := "" - - id0 := h.BindFunc(func(e *Event) error { calls += "0"; return e.Next() }) - id1 := h.BindFunc(func(e *Event) error { calls += "1"; return e.Next() }) - h.BindFunc(func(e *Event) error { calls += "2"; return e.Next() }) - h.Bind(&Handler[*Event]{ - Func: func(e *Event) error { calls += "3"; return e.Next() }, - }) - - h.Unbind("missing") // should do nothing and not panic - - if total := len(h.handlers); total != 4 { - t.Fatalf("Expected %d handlers, got %d", 4, total) - } - - h.Unbind(id1, id0) - - if total := len(h.handlers); total != 2 { - t.Fatalf("Expected %d handlers, got %d", 2, total) - } - - err := h.Trigger(&Event{}, func(e *Event) error { calls += "4"; return e.Next() }) - if err != nil { - t.Fatal(err) - } - - expectedCalls := "234" - - if calls != expectedCalls { - t.Fatalf("Expected calls sequence %q, got %q", expectedCalls, calls) - } -} - -func TestHookUnbindAll(t *testing.T) { - h := Hook[*Event]{} - - h.UnbindAll() // should do nothing and not panic - - h.BindFunc(func(e *Event) error { return nil }) - h.BindFunc(func(e *Event) error { return nil }) - - if total := len(h.handlers); total != 2 { - t.Fatalf("Expected %d handlers before UnbindAll, found %d", 2, total) - } - - h.UnbindAll() - - if total := len(h.handlers); total != 0 { - t.Fatalf("Expected no handlers after UnbindAll, found %d", total) - } -} - -func TestHookTriggerErrorPropagation(t *testing.T) { - err := errors.New("test") - - scenarios := []struct { - name string - handlers []func(*Event) error - expectedError error - }{ - { - "without error", - []func(*Event) error{ - func(e *Event) error { return e.Next() }, - func(e *Event) error { return e.Next() }, - }, - nil, - }, - { - "with error", - []func(*Event) error{ - func(e *Event) error { return e.Next() }, - func(e *Event) error { e.Next(); return err }, - func(e *Event) error { return e.Next() }, - }, - err, - }, - } - - for _, s := range scenarios { - t.Run(s.name, func(t *testing.T) { - h := Hook[*Event]{} - for _, handler := range s.handlers { - h.BindFunc(handler) - } - result := h.Trigger(&Event{}) - if result != s.expectedError { - t.Fatalf("Expected %v, got %v", s.expectedError, result) - } - }) - } -} diff --git a/internals/shared/hook/tagged.go b/internals/shared/hook/tagged.go deleted file mode 100644 index e6a7764..0000000 --- a/internals/shared/hook/tagged.go +++ /dev/null @@ -1,84 +0,0 @@ -package hook - -import ( - "slices" -) - -// Tagger defines an interface for event data structs that support tags/groups/categories/etc. -// Usually used together with TaggedHook. -type Tagger interface { - Resolver - - Tags() []string -} - -// wrapped local Hook embedded struct to limit the public API surface. -type mainHook[T Tagger] struct { - *Hook[T] -} - -// NewTaggedHook creates a new TaggedHook with the provided main hook and optional tags. -func NewTaggedHook[T Tagger](hook *Hook[T], tags ...string) *TaggedHook[T] { - return &TaggedHook[T]{ - mainHook[T]{hook}, - tags, - } -} - -// TaggedHook defines a proxy hook which register handlers that are triggered only -// if the TaggedHook.tags are empty or includes at least one of the event data tag(s). -type TaggedHook[T Tagger] struct { - mainHook[T] - - tags []string -} - -// CanTriggerOn checks if the current TaggedHook can be triggered with -// the provided event data tags. -// -// It returns always true if the hook doens't have any tags. -func (h *TaggedHook[T]) CanTriggerOn(tagsToCheck []string) bool { - if len(h.tags) == 0 { - return true // match all - } - - for _, t := range tagsToCheck { - if slices.Contains(h.tags, t) { - return true - } - } - - return false -} - -// Bind registers the provided handler to the current hooks queue. -// -// It is similar to [Hook.Bind] with the difference that the handler -// function is invoked only if the event data tags satisfy h.CanTriggerOn. -func (h *TaggedHook[T]) Bind(handler *Handler[T]) string { - fn := handler.Func - - handler.Func = func(e T) error { - if h.CanTriggerOn(e.Tags()) { - return fn(e) - } - - return e.Next() - } - - return h.mainHook.Bind(handler) -} - -// BindFunc registers a new handler with the specified function. -// -// It is similar to [Hook.Bind] with the difference that the handler -// function is invoked only if the event data tags satisfy h.CanTriggerOn. -func (h *TaggedHook[T]) BindFunc(fn func(e T) error) string { - return h.mainHook.BindFunc(func(e T) error { - if h.CanTriggerOn(e.Tags()) { - return fn(e) - } - - return e.Next() - }) -} diff --git a/internals/shared/hook/tagged_test.go b/internals/shared/hook/tagged_test.go deleted file mode 100644 index 7fa4ce3..0000000 --- a/internals/shared/hook/tagged_test.go +++ /dev/null @@ -1,84 +0,0 @@ -package hook - -import ( - "strings" - "testing" -) - -type mockTagsEvent struct { - Event - tags []string -} - -func (m mockTagsEvent) Tags() []string { - return m.tags -} - -func TestTaggedHook(t *testing.T) { - calls := "" - - base := &Hook[*mockTagsEvent]{} - base.BindFunc(func(e *mockTagsEvent) error { calls += "f0"; return e.Next() }) - - hA := NewTaggedHook(base) - hA.BindFunc(func(e *mockTagsEvent) error { calls += "a1"; return e.Next() }) - hA.Bind(&Handler[*mockTagsEvent]{ - Func: func(e *mockTagsEvent) error { calls += "a2"; return e.Next() }, - Priority: -1, - }) - - hB := NewTaggedHook(base, "b1", "b2") - hB.BindFunc(func(e *mockTagsEvent) error { calls += "b1"; return e.Next() }) - hB.Bind(&Handler[*mockTagsEvent]{ - Func: func(e *mockTagsEvent) error { calls += "b2"; return e.Next() }, - Priority: -2, - }) - - hC := NewTaggedHook(base, "c1", "c2") - hC.BindFunc(func(e *mockTagsEvent) error { calls += "c1"; return e.Next() }) - hC.Bind(&Handler[*mockTagsEvent]{ - Func: func(e *mockTagsEvent) error { calls += "c2"; return e.Next() }, - Priority: -3, - }) - - scenarios := []struct { - event *mockTagsEvent - expectedCalls string - }{ - { - &mockTagsEvent{}, - "a2f0a1", - }, - { - &mockTagsEvent{tags: []string{"missing"}}, - "a2f0a1", - }, - { - &mockTagsEvent{tags: []string{"b2"}}, - "b2a2f0a1b1", - }, - { - &mockTagsEvent{tags: []string{"c1"}}, - "c2a2f0a1c1", - }, - { - &mockTagsEvent{tags: []string{"b1", "c2"}}, - "c2b2a2f0a1b1c1", - }, - } - - for _, s := range scenarios { - t.Run(strings.Join(s.event.tags, "_"), func(t *testing.T) { - calls = "" // reset - - err := base.Trigger(s.event) - if err != nil { - t.Fatalf("Unexpected trigger error: %v", err) - } - - if calls != s.expectedCalls { - t.Fatalf("Expected calls sequence %q, got %q", s.expectedCalls, calls) - } - }) - } -} diff --git a/pkg/logging/txt/format_gorutines.go b/pkg/logging/txt/format_gorutines.go index 2c383fb..cbef038 100644 --- a/pkg/logging/txt/format_gorutines.go +++ b/pkg/logging/txt/format_gorutines.go @@ -8,7 +8,7 @@ import ( "github.com/sirupsen/logrus" - "github.com/netbirdio/management-refactor/formatter/hook" + "github.com/netbirdio/management-refactor/pkg/logging/hook" ) func (f *TextFormatter) Format(entry *logrus.Entry) ([]byte, error) {