commit 67e1b9213bc1d749026e61effe0b45314c6bd5d4 Author: Pascal Fischer Date: Tue Apr 15 17:09:46 2025 +0200 add first structure diff --git a/cmd/management.go b/cmd/management.go new file mode 100644 index 0000000..52e4849 --- /dev/null +++ b/cmd/management.go @@ -0,0 +1,52 @@ +package cmd + +import ( + "os" + "os/signal" + "syscall" + + "github.com/spf13/cobra" + + "management/internal/server" + "management/pkg/logging" +) + +var log = logging.LoggerForThisPackage() + +// mgmtCmd starts the management server +var mgmtCmd = &cobra.Command{ + Use: "management", + Short: "start NetBird Management Server", + PreRunE: func(cmd *cobra.Command, args []string) error { + err := logging.Init("logging.yaml") + if err != nil { + log.Fatalf("Failed to init logging: %v", err) + } + + srv := server.NewServer() + + go func() { + log.Info("Starting server on :8080") + if err := srv.Start(); err != nil { + log.Fatalf("Server error: %v", err) + } + }() + + stopChan := make(chan os.Signal, 1) + signal.Notify(stopChan, os.Interrupt, syscall.SIGTERM) + <-stopChan + + log.Info("Shutting down server...") + if err := srv.Stop(); err != nil { + log.Errorf("Error stopping server: %v", err) + } + log.Info("Server stopped gracefully.") + + return nil + }, +} + +func init() { + // Attach serveCmd to the rootCmd + rootCmd.AddCommand(mgmtCmd) +} diff --git a/cmd/root.go b/cmd/root.go new file mode 100644 index 0000000..f86cc81 --- /dev/null +++ b/cmd/root.go @@ -0,0 +1,16 @@ +package cmd + +import "github.com/spf13/cobra" + +var rootCmd = &cobra.Command{ + Use: "netbird-mgmt", + Short: "", + Long: "", + Version: "", + SilenceUsage: true, +} + +// Execute is the entry point for all commands. +func Execute() error { + return rootCmd.Execute() +} diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..ed393da --- /dev/null +++ b/go.mod @@ -0,0 +1,47 @@ +module management + +go 1.23.0 + +toolchain go1.23.8 + +require ( + github.com/caarlos0/env/v11 v11.3.1 + github.com/gorilla/mux v1.8.1 + github.com/netbirdio/netbird v0.41.0 + github.com/petermattis/goid v0.0.0-20250319124200-ccd6737f222a + github.com/sirupsen/logrus v1.9.3 + github.com/spf13/cobra v1.9.1 + github.com/spf13/viper v1.20.1 + github.com/stretchr/testify v1.10.0 + gorm.io/driver/postgres v1.5.11 + gorm.io/driver/sqlite v1.5.7 + gorm.io/gorm v1.25.12 +) + +require ( + github.com/davecgh/go-spew v1.1.1 // indirect + github.com/fsnotify/fsnotify v1.8.0 // indirect + github.com/go-viper/mapstructure/v2 v2.2.1 // indirect + github.com/inconshreveable/mousetrap v1.1.0 // indirect + github.com/jackc/pgpassfile v1.0.0 // indirect + github.com/jackc/pgservicefile v0.0.0-20221227161230-091c0ba34f0a // indirect + github.com/jackc/pgx/v5 v5.5.5 // indirect + github.com/jackc/puddle/v2 v2.2.1 // indirect + github.com/jinzhu/inflection v1.0.0 // indirect + github.com/jinzhu/now v1.1.5 // indirect + github.com/mattn/go-sqlite3 v1.14.22 // indirect + github.com/pelletier/go-toml/v2 v2.2.3 // indirect + github.com/pmezard/go-difflib v1.0.0 // indirect + github.com/sagikazarmark/locafero v0.7.0 // indirect + github.com/sourcegraph/conc v0.3.0 // indirect + github.com/spf13/afero v1.12.0 // indirect + github.com/spf13/cast v1.7.1 // indirect + github.com/spf13/pflag v1.0.6 // indirect + github.com/subosito/gotenv v1.6.0 // indirect + go.uber.org/multierr v1.11.0 // indirect + golang.org/x/crypto v0.36.0 // indirect + golang.org/x/sync v0.12.0 // indirect + golang.org/x/sys v0.31.0 // indirect + golang.org/x/text v0.23.0 // indirect + gopkg.in/yaml.v3 v3.0.1 // indirect +) diff --git a/internal/controllers/ephemeral_peers/controller.go b/internal/controllers/ephemeral_peers/controller.go new file mode 100644 index 0000000..5b634be --- /dev/null +++ b/internal/controllers/ephemeral_peers/controller.go @@ -0,0 +1,226 @@ +package server + +import ( + "context" + "sync" + "time" + + log "github.com/sirupsen/logrus" + + nbAccount "github.com/netbirdio/netbird/management/server/account" + "github.com/netbirdio/netbird/management/server/activity" + nbpeer "github.com/netbirdio/netbird/management/server/peer" + "github.com/netbirdio/netbird/management/server/store" +) + +const ( + ephemeralLifeTime = 10 * time.Minute +) + +var ( + timeNow = time.Now +) + +type ephemeralPeer struct { + id string + accountID string + deadline time.Time + next *ephemeralPeer +} + +// todo: consider to remove peer from ephemeral list when the peer has been deleted via API. If we do not do it +// in worst case we will get invalid error message in this manager. + +// Controller keep a list of ephemeral peers. After ephemeralLifeTime inactivity the peer will be deleted +// automatically. Inactivity means the peer disconnected from the Management server. +type Controller struct { + store store.Store + + headPeer *ephemeralPeer + tailPeer *ephemeralPeer + peersLock sync.Mutex + timer *time.Timer +} + +// NewEphemeralManager instantiate new Controller +func NewEphemeralManager(peersManager) *Controller { + return &Controller{ + store: store, + } +} + +// LoadInitialPeers load from the database the ephemeral type of peers and schedule a cleanup procedure to the head +// of the linked list (to the most deprecated peer). At the end of cleanup it schedules the next cleanup to the new +// head. +func (e *Controller) LoadInitialPeers(ctx context.Context) { + e.peersLock.Lock() + defer e.peersLock.Unlock() + + e.loadEphemeralPeers(ctx) + if e.headPeer != nil { + e.timer = time.AfterFunc(ephemeralLifeTime, func() { + e.cleanup(ctx) + }) + } +} + +// Stop timer +func (e *Controller) Stop() { + e.peersLock.Lock() + defer e.peersLock.Unlock() + + if e.timer != nil { + e.timer.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) { + if !peer.Ephemeral { + return + } + + log.WithContext(ctx).Tracef("remove peer from ephemeral list: %s", peer.ID) + + e.peersLock.Lock() + defer e.peersLock.Unlock() + + e.removePeer(peer.ID) + + // stop the unnecessary timer + if e.headPeer == nil && e.timer != nil { + e.timer.Stop() + e.timer = nil + } +} + +// OnPeerDisconnected add the peer to the linked list of ephemeral peers. Because of the peer +// is inactive it will be deleted after the ephemeralLifeTime period. +func (e *Controller) OnPeerDisconnected(ctx context.Context, peer *nbpeer.Peer) { + if !peer.Ephemeral { + return + } + + log.WithContext(ctx).Tracef("add peer to ephemeral list: %s", peer.ID) + + e.peersLock.Lock() + defer e.peersLock.Unlock() + + if e.isPeerOnList(peer.ID) { + return + } + + e.addPeer(peer.AccountID, peer.ID, newDeadLine()) + if e.timer == nil { + e.timer = time.AfterFunc(e.headPeer.deadline.Sub(timeNow()), func() { + e.cleanup(ctx) + }) + } +} + +func (e *Controller) loadEphemeralPeers(ctx context.Context) { + peers, err := e.store.GetAllEphemeralPeers(ctx, store.LockingStrengthShare) + if err != nil { + log.WithContext(ctx).Debugf("failed to load ephemeral peers: %s", err) + return + } + + t := newDeadLine() + for _, p := range peers { + e.addPeer(p.AccountID, p.ID, t) + } + + log.WithContext(ctx).Debugf("loaded ephemeral peer(s): %d", len(peers)) +} + +func (e *Controller) cleanup(ctx context.Context) { + log.Tracef("on ephemeral cleanup") + deletePeers := make(map[string]*ephemeralPeer) + + e.peersLock.Lock() + now := timeNow() + for p := e.headPeer; p != nil; p = p.next { + if now.Before(p.deadline) { + break + } + + deletePeers[p.id] = p + e.headPeer = p.next + if p.next == nil { + e.tailPeer = nil + } + } + + if e.headPeer != nil { + e.timer = time.AfterFunc(e.headPeer.deadline.Sub(timeNow()), func() { + e.cleanup(ctx) + }) + } else { + e.timer = nil + } + + e.peersLock.Unlock() + + for id, p := range deletePeers { + log.WithContext(ctx).Debugf("delete ephemeral peer: %s", id) + err := e.accountManager.DeletePeer(ctx, p.accountID, id, activity.SystemInitiator) + if err != nil { + log.WithContext(ctx).Errorf("failed to delete ephemeral peer: %s", err) + } + } +} + +func (e *Controller) addPeer(accountID string, peerID string, deadline time.Time) { + ep := &ephemeralPeer{ + id: peerID, + accountID: accountID, + deadline: deadline, + } + + if e.headPeer == nil { + e.headPeer = ep + } + if e.tailPeer != nil { + e.tailPeer.next = ep + } + e.tailPeer = ep +} + +func (e *Controller) removePeer(id string) { + if e.headPeer == nil { + return + } + + if e.headPeer.id == id { + e.headPeer = e.headPeer.next + if e.tailPeer.id == id { + e.tailPeer = nil + } + return + } + + for p := e.headPeer; p.next != nil; p = p.next { + if p.next.id == id { + // if we remove the last element from the chain then set the last-1 as tail + if e.tailPeer.id == id { + e.tailPeer = p + } + p.next = p.next.next + return + } + } +} + +func (e *Controller) isPeerOnList(id string) bool { + for p := e.headPeer; p != nil; p = p.next { + if p.id == id { + return true + } + } + return false +} + +func newDeadLine() time.Time { + return timeNow().Add(ephemeralLifeTime) +} diff --git a/internal/controllers/ephemeral_peers/controller_test.go b/internal/controllers/ephemeral_peers/controller_test.go new file mode 100644 index 0000000..38477f7 --- /dev/null +++ b/internal/controllers/ephemeral_peers/controller_test.go @@ -0,0 +1,149 @@ +package server + +import ( + "context" + "fmt" + "testing" + "time" + + nbAccount "github.com/netbirdio/netbird/management/server/account" + nbpeer "github.com/netbirdio/netbird/management/server/peer" + "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/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/internal/modules/users/api.go b/internal/modules/users/api.go new file mode 100644 index 0000000..c834e1f --- /dev/null +++ b/internal/modules/users/api.go @@ -0,0 +1,84 @@ +package users + +import ( + "encoding/json" + "net/http" + + "github.com/gorilla/mux" + nbcontext "github.com/netbirdio/netbird/management/server/context" + "github.com/netbirdio/netbird/management/server/http/util" + + "management/internal/shared/db" + "management/internal/shared/errors" + "management/internal/shared/permissions" + "management/internal/shared/permissions/modules" + "management/internal/shared/permissions/operations" +) + +type handler struct { + manager *Manager + permissionsManager permissions.Manager +} + +func newHandler(manager *Manager, permissionsManager permissions.Manager) *handler { + return &handler{ + manager: manager, + permissionsManager: permissionsManager, + } +} + +func (h *handler) RegisterAPI(router *mux.Router) { + router.HandleFunc("/users", h.GetAllUsers).Methods("GET", "OPTIONS") + router.HandleFunc("/users/{userId}", h.GetUser).Methods("GET", "OPTIONS") +} + +func (h *handler) GetAllUsers(w http.ResponseWriter, r *http.Request) { + userAuth, err := nbcontext.GetUserAuthFromContext(r.Context()) + if err != nil { + util.WriteError(r.Context(), err, w) + return + } + + allowed, err := h.permissionsManager.ValidateUserPermissions(r.Context(), userAuth.AccountId, userAuth.UserId, modules.Users, operations.Read) + if err != nil { + util.WriteError(r.Context(), errors.NewPermissionValidationError(err), w) + return + } + if !allowed { + util.WriteError(r.Context(), errors.NewPermissionDeniedError(), w) + } + + users, err := h.manager.GetAllUsers(r.Context(), nil, db.LockingStrengthShare, userAuth.AccountId) + if err != nil { + http.Error(w, "Internal Server Error", http.StatusInternalServerError) + return + } + _ = json.NewEncoder(w).Encode(users) +} + +func (h *handler) GetUser(w http.ResponseWriter, r *http.Request) { + userAuth, err := nbcontext.GetUserAuthFromContext(r.Context()) + if err != nil { + util.WriteError(r.Context(), err, w) + return + } + + allowed, err := h.permissionsManager.ValidateUserPermissions(r.Context(), userAuth.AccountId, userAuth.UserId, modules.Users, operations.Read) + if err != nil { + util.WriteError(r.Context(), errors.NewPermissionValidationError(err), w) + return + } + if !allowed { + util.WriteError(r.Context(), errors.NewPermissionDeniedError(), w) + } + + vars := mux.Vars(r) + userId := vars["userId"] + + user, err := h.manager.GetUserByID(r.Context(), nil, db.LockingStrengthShare, userId) + if err != nil { + http.Error(w, "Not Found", http.StatusNotFound) + return + } + _ = json.NewEncoder(w).Encode(user) +} diff --git a/internal/modules/users/config.go b/internal/modules/users/config.go new file mode 100644 index 0000000..82abcb9 --- /dev/null +++ b/internal/modules/users/config.go @@ -0,0 +1 @@ +package users diff --git a/internal/modules/users/manager.go b/internal/modules/users/manager.go new file mode 100644 index 0000000..7b18599 --- /dev/null +++ b/internal/modules/users/manager.go @@ -0,0 +1,32 @@ +package users + +import ( + "context" + + "management/internal/modules/users/types" + "management/internal/shared/db" + "management/internal/shared/permissions" + "management/pkg/logging" +) + +var log = logging.LoggerForThisPackage() + +type Manager struct { + repo Repository + handler *handler +} + +func NewManager(store *db.Store, permissionsManager permissions.Manager) *Manager { + repo := newRepository(store) + m := &Manager{repo: repo} + m.handler = newHandler(m, permissionsManager) + return m +} + +func (m *Manager) GetAllUsers(ctx context.Context, tx db.Transaction, strength db.LockingStrength, accountID string) ([]types.User, error) { + return m.repo.GetAllUsers(tx, strength, accountID) +} + +func (m *Manager) GetUserByID(ctx context.Context, tx db.Transaction, strength db.LockingStrength, id string) (*types.User, error) { + return m.repo.GetUserByID(tx, strength, id) +} diff --git a/internal/modules/users/pats/types/personal_access_token.go b/internal/modules/users/pats/types/personal_access_token.go new file mode 100644 index 0000000..b39e5dd --- /dev/null +++ b/internal/modules/users/pats/types/personal_access_token.go @@ -0,0 +1,112 @@ +package types + +import ( + "crypto/sha256" + b64 "encoding/base64" + "fmt" + "hash/crc32" + "time" + + b "github.com/hashicorp/go-secure-stdlib/base62" + "github.com/rs/xid" + + "github.com/netbirdio/netbird/base62" +) + +const ( + // PATPrefix is the globally used, 4 char prefix for personal access tokens + PATPrefix = "nbp_" + // PATSecretLength number of characters used for the secret inside the token + PATSecretLength = 30 + // PATChecksumLength number of characters used for the encoded checksum of the secret inside the token + PATChecksumLength = 6 + // PATLength total number of characters used for the token + PATLength = 40 +) + +// PersonalAccessToken holds all information about a PAT including a hashed version of it for verification +type PersonalAccessToken struct { + ID string `gorm:"primaryKey"` + // User is a reference to Account that this object belongs + UserID string `gorm:"index"` + Name string + HashedToken string + ExpirationDate *time.Time + // scope could be added in future + CreatedBy string + CreatedAt time.Time + LastUsed *time.Time +} + +func (t *PersonalAccessToken) Copy() *PersonalAccessToken { + return &PersonalAccessToken{ + ID: t.ID, + Name: t.Name, + HashedToken: t.HashedToken, + ExpirationDate: t.ExpirationDate, + CreatedBy: t.CreatedBy, + CreatedAt: t.CreatedAt, + LastUsed: t.LastUsed, + } +} + +// GetExpirationDate returns the expiration time of the token. +func (t *PersonalAccessToken) GetExpirationDate() time.Time { + if t.ExpirationDate != nil { + return *t.ExpirationDate + } + return time.Time{} +} + +// GetLastUsed returns the last time the token was used. +func (t *PersonalAccessToken) GetLastUsed() time.Time { + if t.LastUsed != nil { + return *t.LastUsed + } + return time.Time{} +} + +// PersonalAccessTokenGenerated holds the new PersonalAccessToken and the plain text version of it +type PersonalAccessTokenGenerated struct { + PlainToken string + PersonalAccessToken +} + +// CreateNewPAT will generate a new PersonalAccessToken that can be assigned to a User. +// Additionally, it will return the token in plain text once, to give to the user and only save a hashed version +func CreateNewPAT(name string, expirationInDays int, targetID, createdBy string) (*PersonalAccessTokenGenerated, error) { + hashedToken, plainToken, err := generateNewToken() + if err != nil { + return nil, err + } + currentTime := time.Now() + expirationDate := currentTime.AddDate(0, 0, expirationInDays) + return &PersonalAccessTokenGenerated{ + PersonalAccessToken: PersonalAccessToken{ + ID: xid.New().String(), + UserID: targetID, + Name: name, + HashedToken: hashedToken, + ExpirationDate: &expirationDate, + CreatedBy: createdBy, + CreatedAt: currentTime, + }, + PlainToken: plainToken, + }, nil + +} + +func generateNewToken() (string, string, error) { + secret, err := b.Random(PATSecretLength) + if err != nil { + return "", "", err + } + + checksum := crc32.ChecksumIEEE([]byte(secret)) + encodedChecksum := base62.Encode(checksum) + paddedChecksum := fmt.Sprintf("%06s", encodedChecksum) + plainToken := PATPrefix + secret + paddedChecksum + hashedToken := sha256.Sum256([]byte(plainToken)) + encodedHashedToken := b64.StdEncoding.EncodeToString(hashedToken[:]) + return encodedHashedToken, plainToken, nil +} diff --git a/internal/modules/users/repository.go b/internal/modules/users/repository.go new file mode 100644 index 0000000..a33c4f5 --- /dev/null +++ b/internal/modules/users/repository.go @@ -0,0 +1,51 @@ +package users + +import ( + "management/internal/modules/users/types" + "management/internal/shared/db" +) + +type Repository interface { + RunInTx(fn func(tx db.Transaction) error) error + GetAllUsers(tx db.Transaction, strength db.LockingStrength, accountID string) ([]types.User, error) + GetUserByID(tx db.Transaction, strength db.LockingStrength, id string) (*types.User, error) + CreateUser(tx db.Transaction, u *types.User) error +} + +type repository struct { + store *db.Store +} + +func newRepository(s *db.Store) Repository { + err := s.AutoMigrate(types.User{}) + if err != nil { + log.Fatalf("Failed to auto migrate: %v", err) + } + return &repository{store: s} +} + +func (r *repository) RunInTx(fn func(tx db.Transaction) error) error { + return r.store.RunInTx(fn) +} + +func (r *repository) GetAllUsers(tx db.Transaction, strength db.LockingStrength, accountID string) ([]types.User, error) { + var users []types.User + err := r.store.GetMany(tx, strength, &users, "account_id = ?", accountID) + if err != nil { + return nil, err + } + return users, nil +} + +func (r *repository) GetUserByID(tx db.Transaction, strength db.LockingStrength, id string) (*types.User, error) { + var user types.User + err := r.store.GetOne(tx, strength, &user, "id = ?", id) + if err != nil { + return nil, err + } + return &user, nil +} + +func (r *repository) CreateUser(tx db.Transaction, u *types.User) error { + return r.store.Create(tx, u) +} diff --git a/internal/modules/users/types/user.go b/internal/modules/users/types/user.go new file mode 100644 index 0000000..9abbffe --- /dev/null +++ b/internal/modules/users/types/user.go @@ -0,0 +1,241 @@ +package types + +import ( + "fmt" + "strings" + "time" + + "github.com/netbirdio/netbird/management/server/idp" + "github.com/netbirdio/netbird/management/server/integration_reference" + + "management/internal/modules/users/pats/types" +) + +const ( + UserRoleOwner UserRole = "owner" + UserRoleAdmin UserRole = "admin" + UserRoleUser UserRole = "user" + UserRoleUnknown UserRole = "unknown" + UserRoleBillingAdmin UserRole = "billing_admin" + + UserStatusActive UserStatus = "active" + UserStatusDisabled UserStatus = "disabled" + UserStatusInvited UserStatus = "invited" + + UserIssuedAPI = "api" + UserIssuedIntegration = "integration" +) + +// StrRoleToUserRole returns UserRole for a given strRole or UserRoleUnknown if the specified role is unknown +func StrRoleToUserRole(strRole string) UserRole { + switch strings.ToLower(strRole) { + case "owner": + return UserRoleOwner + case "admin": + return UserRoleAdmin + case "user": + return UserRoleUser + case "billing_admin": + return UserRoleBillingAdmin + default: + return UserRoleUnknown + } +} + +// UserStatus is the status of a User +type UserStatus string + +// UserRole is the role of a User +type UserRole string + +type UserInfo struct { + ID string `json:"id"` + Email string `json:"email"` + Name string `json:"name"` + Role string `json:"role"` + AutoGroups []string `json:"auto_groups"` + Status string `json:"-"` + IsServiceUser bool `json:"is_service_user"` + IsBlocked bool `json:"is_blocked"` + NonDeletable bool `json:"non_deletable"` + LastLogin time.Time `json:"last_login"` + Issued string `json:"issued"` + IntegrationReference integration_reference.IntegrationReference `json:"-"` + Permissions UserPermissions `json:"permissions"` +} + +type UserPermissions struct { + DashboardView string `json:"dashboard_view"` +} + +// User represents a user of the system +type User struct { + Id string `gorm:"primaryKey"` + // AccountID is a reference to Account that this object belongs + AccountID string `json:"-" gorm:"index"` + Role UserRole + IsServiceUser bool + // NonDeletable indicates whether the service user can be deleted + NonDeletable bool + // ServiceUserName is only set if IsServiceUser is true + ServiceUserName string + // AutoGroups is a list of Group IDs to auto-assign to peers registered by this user + AutoGroups []string `gorm:"serializer:json"` + PATs map[string]*types.PersonalAccessToken `gorm:"-"` + PATsG []types.PersonalAccessToken `json:"-" gorm:"foreignKey:UserID;references:id;constraint:OnDelete:CASCADE;"` + // Blocked indicates whether the user is blocked. Blocked users can't use the system. + Blocked bool + // LastLogin is the last time the user logged in to IdP + LastLogin *time.Time + // CreatedAt records the time the user was created + CreatedAt time.Time + + // Issued of the user + Issued string `gorm:"default:api"` + + IntegrationReference integration_reference.IntegrationReference `gorm:"embedded;embeddedPrefix:integration_ref_"` +} + +// IsBlocked returns true if the user is blocked, false otherwise +func (u *User) IsBlocked() bool { + return u.Blocked +} + +func (u *User) LastDashboardLoginChanged(lastLogin time.Time) bool { + return lastLogin.After(u.GetLastLogin()) && !u.GetLastLogin().IsZero() +} + +// GetLastLogin returns the last login time of the user. +func (u *User) GetLastLogin() time.Time { + if u.LastLogin != nil { + return *u.LastLogin + } + return time.Time{} +} + +// HasAdminPower returns true if the user has admin or owner roles, false otherwise +func (u *User) HasAdminPower() bool { + return u.Role == UserRoleAdmin || u.Role == UserRoleOwner +} + +// IsAdminOrServiceUser checks if the user has admin power or is a service user. +func (u *User) IsAdminOrServiceUser() bool { + return u.HasAdminPower() || u.IsServiceUser +} + +// IsRegularUser checks if the user is a regular user. +func (u *User) IsRegularUser() bool { + return !u.HasAdminPower() && !u.IsServiceUser +} + +// ToUserInfo converts a User object to a UserInfo object. +func (u *User) ToUserInfo(userData *idp.UserData, settings *Settings) (*UserInfo, error) { + autoGroups := u.AutoGroups + if autoGroups == nil { + autoGroups = []string{} + } + + dashboardViewPermissions := "full" + if !u.HasAdminPower() { + dashboardViewPermissions = "limited" + if settings.RegularUsersViewBlocked { + dashboardViewPermissions = "blocked" + } + } + + if userData == nil { + return &UserInfo{ + ID: u.Id, + Email: "", + Name: u.ServiceUserName, + Role: string(u.Role), + AutoGroups: u.AutoGroups, + Status: string(UserStatusActive), + IsServiceUser: u.IsServiceUser, + IsBlocked: u.Blocked, + LastLogin: u.GetLastLogin(), + Issued: u.Issued, + Permissions: UserPermissions{ + DashboardView: dashboardViewPermissions, + }, + }, nil + } + if userData.ID != u.Id { + return nil, fmt.Errorf("wrong UserData provided for user %s", u.Id) + } + + userStatus := UserStatusActive + if userData.AppMetadata.WTPendingInvite != nil && *userData.AppMetadata.WTPendingInvite { + userStatus = UserStatusInvited + } + + return &UserInfo{ + ID: u.Id, + Email: userData.Email, + Name: userData.Name, + Role: string(u.Role), + AutoGroups: autoGroups, + Status: string(userStatus), + IsServiceUser: u.IsServiceUser, + IsBlocked: u.Blocked, + LastLogin: u.GetLastLogin(), + Issued: u.Issued, + Permissions: UserPermissions{ + DashboardView: dashboardViewPermissions, + }, + }, nil +} + +// Copy the user +func (u *User) Copy() *User { + autoGroups := make([]string, len(u.AutoGroups)) + copy(autoGroups, u.AutoGroups) + pats := make(map[string]*types.PersonalAccessToken, len(u.PATs)) + for k, v := range u.PATs { + pats[k] = v.Copy() + } + return &User{ + Id: u.Id, + AccountID: u.AccountID, + Role: u.Role, + AutoGroups: autoGroups, + IsServiceUser: u.IsServiceUser, + NonDeletable: u.NonDeletable, + ServiceUserName: u.ServiceUserName, + PATs: pats, + Blocked: u.Blocked, + LastLogin: u.LastLogin, + CreatedAt: u.CreatedAt, + Issued: u.Issued, + IntegrationReference: u.IntegrationReference, + } +} + +// NewUser creates a new user +func NewUser(id string, role UserRole, isServiceUser bool, nonDeletable bool, serviceUserName string, autoGroups []string, issued string) *User { + return &User{ + Id: id, + Role: role, + IsServiceUser: isServiceUser, + NonDeletable: nonDeletable, + ServiceUserName: serviceUserName, + AutoGroups: autoGroups, + Issued: issued, + CreatedAt: time.Now().UTC(), + } +} + +// NewRegularUser creates a new user with role UserRoleUser +func NewRegularUser(id string) *User { + return NewUser(id, UserRoleUser, false, false, "", []string{}, UserIssuedAPI) +} + +// NewAdminUser creates a new user with role UserRoleAdmin +func NewAdminUser(id string) *User { + return NewUser(id, UserRoleAdmin, false, false, "", []string{}, UserIssuedAPI) +} + +// NewOwnerUser creates a new user with role UserRoleOwner +func NewOwnerUser(id string) *User { + return NewUser(id, UserRoleOwner, false, false, "", []string{}, UserIssuedAPI) +} diff --git a/internal/server/server.go b/internal/server/server.go new file mode 100644 index 0000000..7f083aa --- /dev/null +++ b/internal/server/server.go @@ -0,0 +1,58 @@ +package server + +import ( + "context" + "net/http" + "time" + + "management/internal/modules/users" + "management/internal/shared/api" + "management/internal/shared/db" + "management/internal/shared/permissions" + "management/pkg/logging" +) + +// Server holds the HTTP server instance. +// Add any additional fields you need, such as database connections, config, etc. +type Server struct { + httpServer *http.Server +} + +var log = logging.LoggerForThisPackage() + +// NewServer initializes and configures a new Server instance +func NewServer() *Server { + ctx := context.Background() + + dbConn, err := db.NewDatabaseConn(ctx) + if err != nil { + log.Fatalf("error while creating database connection: %s", err) + } + + store := db.NewStore(ctx, dbConn) + + router := api.NewRouter() + + permissionsManager := permissions.NewManager(store) + userManager := users.NewManager(store, permissions.NewManager(store)) + + return &Server{ + httpServer: &http.Server{ + Addr: ":8080", // or from a config file + Handler: router, + }, + } +} + +// Start begins listening for HTTP requests on the configured address +func (s *Server) Start() error { + return s.httpServer.ListenAndServe() +} + +// Stop attempts a graceful shutdown, waiting up to 5 seconds for active connections to finish +func (s *Server) Stop() error { + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + return s.httpServer.Shutdown(ctx) +} diff --git a/internal/shared/activity/codes.go b/internal/shared/activity/codes.go new file mode 100644 index 0000000..46ae754 --- /dev/null +++ b/internal/shared/activity/codes.go @@ -0,0 +1,288 @@ +package activity + +import "maps" + +// Activity that triggered an Event +type Activity int + +// Code is an activity string representation +type Code struct { + Message string + Code string +} + +// Existing consts must not be changed, as this will break the compatibility with the existing data +const ( + // PeerAddedByUser indicates that a user added a new peer to the system + PeerAddedByUser Activity = 0 + // PeerAddedWithSetupKey indicates that a new peer joined the system using a setup key + PeerAddedWithSetupKey Activity = 1 + // UserJoined indicates that a new user joined the account + UserJoined Activity = 2 + // UserInvited indicates that a new user was invited to join the account + UserInvited Activity = 3 + // AccountCreated indicates that a new account has been created + AccountCreated Activity = 4 + // PeerRemovedByUser indicates that a user removed a peer from the system + PeerRemovedByUser Activity = 5 + // RuleAdded indicates that a user added a new rule + RuleAdded Activity = 6 + // RuleUpdated indicates that a user updated a rule + RuleUpdated Activity = 7 + // RuleRemoved indicates that a user removed a rule + RuleRemoved Activity = 8 + // PolicyAdded indicates that a user added a new policy + PolicyAdded Activity = 9 + // PolicyUpdated indicates that a user updated a policy + PolicyUpdated Activity = 10 + // PolicyRemoved indicates that a user removed a policy + PolicyRemoved Activity = 11 + // SetupKeyCreated indicates that a user created a new setup key + SetupKeyCreated Activity = 12 + // SetupKeyUpdated indicates that a user updated a setup key + SetupKeyUpdated Activity = 13 + // SetupKeyRevoked indicates that a user revoked a setup key + SetupKeyRevoked Activity = 14 + // SetupKeyOverused indicates that setup key usage exhausted + SetupKeyOverused Activity = 15 + // GroupCreated indicates that a user created a group + GroupCreated Activity = 16 + // GroupUpdated indicates that a user updated a group + GroupUpdated Activity = 17 + // GroupAddedToPeer indicates that a user added group to a peer + GroupAddedToPeer Activity = 18 + // GroupRemovedFromPeer indicates that a user removed peer group + GroupRemovedFromPeer Activity = 19 + // GroupAddedToUser indicates that a user added group to a user + GroupAddedToUser Activity = 20 + // GroupRemovedFromUser indicates that a user removed a group from a user + GroupRemovedFromUser Activity = 21 + // UserRoleUpdated indicates that a user changed the role of a user + UserRoleUpdated Activity = 22 + // GroupAddedToSetupKey indicates that a user added group to a setup key + GroupAddedToSetupKey Activity = 23 + // GroupRemovedFromSetupKey indicates that a user removed a group from a setup key + GroupRemovedFromSetupKey Activity = 24 + // GroupAddedToDisabledManagementGroups indicates that a user added a group to the DNS setting Disabled management groups + GroupAddedToDisabledManagementGroups Activity = 25 + // GroupRemovedFromDisabledManagementGroups indicates that a user removed a group from the DNS setting Disabled management groups + GroupRemovedFromDisabledManagementGroups Activity = 26 + // RouteCreated indicates that a user created a route + RouteCreated Activity = 27 + // RouteRemoved indicates that a user deleted a route + RouteRemoved Activity = 28 + // RouteUpdated indicates that a user updated a route + RouteUpdated Activity = 29 + // PeerSSHEnabled indicates that a user enabled SSH server on a peer + PeerSSHEnabled Activity = 30 + // PeerSSHDisabled indicates that a user disabled SSH server on a peer + PeerSSHDisabled Activity = 31 + // PeerRenamed indicates that a user renamed a peer + PeerRenamed Activity = 32 + // PeerLoginExpirationEnabled indicates that a user enabled login expiration of a peer + PeerLoginExpirationEnabled Activity = 33 + // PeerLoginExpirationDisabled indicates that a user disabled login expiration of a peer + PeerLoginExpirationDisabled Activity = 34 + // NameserverGroupCreated indicates that a user created a nameservers group + NameserverGroupCreated Activity = 35 + // NameserverGroupDeleted indicates that a user deleted a nameservers group + NameserverGroupDeleted Activity = 36 + // NameserverGroupUpdated indicates that a user updated a nameservers group + NameserverGroupUpdated Activity = 37 + // AccountPeerLoginExpirationEnabled indicates that a user enabled peer login expiration for the account + AccountPeerLoginExpirationEnabled Activity = 38 + // AccountPeerLoginExpirationDisabled indicates that a user disabled peer login expiration for the account + AccountPeerLoginExpirationDisabled Activity = 39 + // AccountPeerLoginExpirationDurationUpdated indicates that a user updated peer login expiration duration for the account + AccountPeerLoginExpirationDurationUpdated Activity = 40 + // PersonalAccessTokenCreated indicates that a user created a personal access token + PersonalAccessTokenCreated Activity = 41 + // PersonalAccessTokenDeleted indicates that a user deleted a personal access token + PersonalAccessTokenDeleted Activity = 42 + // ServiceUserCreated indicates that a user created a service user + ServiceUserCreated Activity = 43 + // ServiceUserDeleted indicates that a user deleted a service user + ServiceUserDeleted Activity = 44 + // UserBlocked indicates that a user blocked another user + UserBlocked Activity = 45 + // UserUnblocked indicates that a user unblocked another user + UserUnblocked Activity = 46 + // UserDeleted indicates that a user deleted another user + UserDeleted Activity = 47 + // GroupDeleted indicates that a user deleted group + GroupDeleted Activity = 48 + // UserLoggedInPeer indicates that user logged in their peer with an interactive SSO login + UserLoggedInPeer Activity = 49 + // PeerLoginExpired indicates that the user peer login has been expired and peer disconnected + PeerLoginExpired Activity = 50 + // DashboardLogin indicates that the user logged in to the dashboard + DashboardLogin Activity = 51 + // IntegrationCreated indicates that the user created an integration + IntegrationCreated Activity = 52 + // IntegrationUpdated indicates that the user updated an integration + IntegrationUpdated Activity = 53 + // IntegrationDeleted indicates that the user deleted an integration + IntegrationDeleted Activity = 54 + // AccountPeerApprovalEnabled indicates that the user enabled peer approval for the account + AccountPeerApprovalEnabled Activity = 55 + // AccountPeerApprovalDisabled indicates that the user disabled peer approval for the account + AccountPeerApprovalDisabled Activity = 56 + // PeerApproved indicates that the peer has been approved + PeerApproved Activity = 57 + // PeerApprovalRevoked indicates that the peer approval has been revoked + PeerApprovalRevoked Activity = 58 + // TransferredOwnerRole indicates that the user transferred the owner role of the account + TransferredOwnerRole Activity = 59 + // PostureCheckCreated indicates that the user created a posture check + PostureCheckCreated Activity = 60 + // PostureCheckUpdated indicates that the user updated a posture check + PostureCheckUpdated Activity = 61 + // PostureCheckDeleted indicates that the user deleted a posture check + PostureCheckDeleted Activity = 62 + + PeerInactivityExpirationEnabled Activity = 63 + PeerInactivityExpirationDisabled Activity = 64 + + AccountPeerInactivityExpirationEnabled Activity = 65 + AccountPeerInactivityExpirationDisabled Activity = 66 + AccountPeerInactivityExpirationDurationUpdated Activity = 67 + + SetupKeyDeleted Activity = 68 + + UserGroupPropagationEnabled Activity = 69 + UserGroupPropagationDisabled Activity = 70 + + AccountRoutingPeerDNSResolutionEnabled Activity = 71 + AccountRoutingPeerDNSResolutionDisabled Activity = 72 + + NetworkCreated Activity = 73 + NetworkUpdated Activity = 74 + NetworkDeleted Activity = 75 + + NetworkResourceCreated Activity = 76 + NetworkResourceUpdated Activity = 77 + NetworkResourceDeleted Activity = 78 + + NetworkRouterCreated Activity = 79 + NetworkRouterUpdated Activity = 80 + NetworkRouterDeleted Activity = 81 + + ResourceAddedToGroup Activity = 82 + ResourceRemovedFromGroup Activity = 83 +) + +var activityMap = map[Activity]Code{ + PeerAddedByUser: {"Peer added", "peer.user.add"}, + PeerAddedWithSetupKey: {"Peer added", "peer.setupkey.add"}, + UserJoined: {"User joined", "user.join"}, + UserInvited: {"User invited", "user.invite"}, + AccountCreated: {"Account created", "account.create"}, + PeerRemovedByUser: {"Peer deleted", "user.peer.delete"}, + RuleAdded: {"Rule added", "rule.add"}, + RuleUpdated: {"Rule updated", "rule.update"}, + RuleRemoved: {"Rule deleted", "rule.delete"}, + PolicyAdded: {"Policy added", "policy.add"}, + PolicyUpdated: {"Policy updated", "policy.update"}, + PolicyRemoved: {"Policy deleted", "policy.delete"}, + SetupKeyCreated: {"Setup key created", "setupkey.add"}, + SetupKeyUpdated: {"Setup key updated", "setupkey.update"}, + SetupKeyRevoked: {"Setup key revoked", "setupkey.revoke"}, + SetupKeyOverused: {"Setup key overused", "setupkey.overuse"}, + GroupCreated: {"Group created", "group.add"}, + GroupUpdated: {"Group updated", "group.update"}, + GroupAddedToPeer: {"Group added to peer", "peer.group.add"}, + GroupRemovedFromPeer: {"Group removed from peer", "peer.group.delete"}, + GroupAddedToUser: {"Group added to user", "user.group.add"}, + GroupRemovedFromUser: {"Group removed from user", "user.group.delete"}, + UserRoleUpdated: {"User role updated", "user.role.update"}, + GroupAddedToSetupKey: {"Group added to setup key", "setupkey.group.add"}, + GroupRemovedFromSetupKey: {"Group removed from user setup key", "setupkey.group.delete"}, + GroupAddedToDisabledManagementGroups: {"Group added to disabled management DNS setting", "dns.setting.disabled.management.group.add"}, + GroupRemovedFromDisabledManagementGroups: {"Group removed from disabled management DNS setting", "dns.setting.disabled.management.group.delete"}, + RouteCreated: {"Route created", "route.add"}, + RouteRemoved: {"Route deleted", "route.delete"}, + RouteUpdated: {"Route updated", "route.update"}, + PeerSSHEnabled: {"Peer SSH server enabled", "peer.ssh.enable"}, + PeerSSHDisabled: {"Peer SSH server disabled", "peer.ssh.disable"}, + PeerRenamed: {"Peer renamed", "peer.rename"}, + PeerLoginExpirationEnabled: {"Peer login expiration enabled", "peer.login.expiration.enable"}, + PeerLoginExpirationDisabled: {"Peer login expiration disabled", "peer.login.expiration.disable"}, + NameserverGroupCreated: {"Nameserver group created", "nameserver.group.add"}, + NameserverGroupDeleted: {"Nameserver group deleted", "nameserver.group.delete"}, + NameserverGroupUpdated: {"Nameserver group updated", "nameserver.group.update"}, + AccountPeerLoginExpirationDurationUpdated: {"Account peer login expiration duration updated", "account.setting.peer.login.expiration.update"}, + AccountPeerLoginExpirationEnabled: {"Account peer login expiration enabled", "account.setting.peer.login.expiration.enable"}, + AccountPeerLoginExpirationDisabled: {"Account peer login expiration disabled", "account.setting.peer.login.expiration.disable"}, + PersonalAccessTokenCreated: {"Personal access token created", "personal.access.token.create"}, + PersonalAccessTokenDeleted: {"Personal access token deleted", "personal.access.token.delete"}, + ServiceUserCreated: {"Service user created", "service.user.create"}, + ServiceUserDeleted: {"Service user deleted", "service.user.delete"}, + UserBlocked: {"User blocked", "user.block"}, + UserUnblocked: {"User unblocked", "user.unblock"}, + UserDeleted: {"User deleted", "user.delete"}, + GroupDeleted: {"Group deleted", "group.delete"}, + UserLoggedInPeer: {"User logged in peer", "user.peer.login"}, + PeerLoginExpired: {"Peer login expired", "peer.login.expire"}, + DashboardLogin: {"Dashboard login", "dashboard.login"}, + IntegrationCreated: {"Integration created", "integration.create"}, + IntegrationUpdated: {"Integration updated", "integration.update"}, + IntegrationDeleted: {"Integration deleted", "integration.delete"}, + AccountPeerApprovalEnabled: {"Account peer approval enabled", "account.setting.peer.approval.enable"}, + AccountPeerApprovalDisabled: {"Account peer approval disabled", "account.setting.peer.approval.disable"}, + PeerApproved: {"Peer approved", "peer.approve"}, + PeerApprovalRevoked: {"Peer approval revoked", "peer.approval.revoke"}, + TransferredOwnerRole: {"Transferred owner role", "transferred.owner.role"}, + PostureCheckCreated: {"Posture check created", "posture.check.create"}, + PostureCheckUpdated: {"Posture check updated", "posture.check.update"}, + PostureCheckDeleted: {"Posture check deleted", "posture.check.delete"}, + + PeerInactivityExpirationEnabled: {"Peer inactivity expiration enabled", "peer.inactivity.expiration.enable"}, + PeerInactivityExpirationDisabled: {"Peer inactivity expiration disabled", "peer.inactivity.expiration.disable"}, + + AccountPeerInactivityExpirationEnabled: {"Account peer inactivity expiration enabled", "account.peer.inactivity.expiration.enable"}, + AccountPeerInactivityExpirationDisabled: {"Account peer inactivity expiration disabled", "account.peer.inactivity.expiration.disable"}, + AccountPeerInactivityExpirationDurationUpdated: {"Account peer inactivity expiration duration updated", "account.peer.inactivity.expiration.update"}, + SetupKeyDeleted: {"Setup key deleted", "setupkey.delete"}, + + UserGroupPropagationEnabled: {"User group propagation enabled", "account.setting.group.propagation.enable"}, + UserGroupPropagationDisabled: {"User group propagation disabled", "account.setting.group.propagation.disable"}, + + AccountRoutingPeerDNSResolutionEnabled: {"Account routing peer DNS resolution enabled", "account.setting.routing.peer.dns.resolution.enable"}, + AccountRoutingPeerDNSResolutionDisabled: {"Account routing peer DNS resolution disabled", "account.setting.routing.peer.dns.resolution.disable"}, + + NetworkCreated: {"Network created", "network.create"}, + NetworkUpdated: {"Network updated", "network.update"}, + NetworkDeleted: {"Network deleted", "network.delete"}, + + NetworkResourceCreated: {"Network resource created", "network.resource.create"}, + NetworkResourceUpdated: {"Network resource updated", "network.resource.update"}, + NetworkResourceDeleted: {"Network resource deleted", "network.resource.delete"}, + + NetworkRouterCreated: {"Network router created", "network.router.create"}, + NetworkRouterUpdated: {"Network router updated", "network.router.update"}, + NetworkRouterDeleted: {"Network router deleted", "network.router.delete"}, + + ResourceAddedToGroup: {"Resource added to group", "resource.group.add"}, + ResourceRemovedFromGroup: {"Resource removed from group", "resource.group.delete"}, +} + +// StringCode returns a string code of the activity +func (a Activity) StringCode() string { + if code, ok := activityMap[a]; ok { + return code.Code + } + return "UNKNOWN_ACTIVITY" +} + +// Message returns a string representation of an activity +func (a Activity) Message() string { + if code, ok := activityMap[a]; ok { + return code.Message + } + return "UNKNOWN_ACTIVITY" +} + +// RegisterActivityMap adds new codes to the activity map +func RegisterActivityMap(codes map[Activity]Code) { + maps.Copy(activityMap, codes) +} diff --git a/internal/shared/activity/config.go b/internal/shared/activity/config.go new file mode 100644 index 0000000..8e9ab5a --- /dev/null +++ b/internal/shared/activity/config.go @@ -0,0 +1,5 @@ +package activity + +type config struct { + Enabled bool `env:"NB_EVENT_ACTIVITY_LOG_ENABLED" envDefault:"true"` +} diff --git a/internal/shared/activity/event.go b/internal/shared/activity/event.go new file mode 100644 index 0000000..0e819c3 --- /dev/null +++ b/internal/shared/activity/event.go @@ -0,0 +1,59 @@ +package activity + +import ( + "time" +) + +const ( + SystemInitiator = "sys" +) + +// ActivityDescriber is an interface that describes an activity +type ActivityDescriber interface { //nolint:revive + StringCode() string + Message() string +} + +// Event represents a network/system activity event. +type Event struct { + // Timestamp of the event + Timestamp time.Time + // Activity that was performed during the event + Activity ActivityDescriber + // ID of the event (can be empty, meaning that it wasn't yet generated) + ID uint64 + // InitiatorID is the ID of an object that initiated the event (e.g., a user) + InitiatorID string + // InitiatorName is the name of an object that initiated the event. + InitiatorName string + // InitiatorEmail is the email address of an object that initiated the event. + InitiatorEmail string + // TargetID is the ID of an object that was effected by the event (e.g., a peer) + TargetID string + // AccountID is the ID of an account where the event happened + AccountID string + + // Meta of the event, e.g. deleted peer information like name, IP, etc + Meta map[string]any +} + +// Copy the event +func (e *Event) Copy() *Event { + + meta := make(map[string]any, len(e.Meta)) + for key, value := range e.Meta { + meta[key] = value + } + + return &Event{ + Timestamp: e.Timestamp, + Activity: e.Activity, + ID: e.ID, + InitiatorID: e.InitiatorID, + InitiatorName: e.InitiatorName, + InitiatorEmail: e.InitiatorEmail, + TargetID: e.TargetID, + AccountID: e.AccountID, + Meta: meta, + } +} diff --git a/internal/shared/activity/manager.go b/internal/shared/activity/manager.go new file mode 100644 index 0000000..255067e --- /dev/null +++ b/internal/shared/activity/manager.go @@ -0,0 +1,48 @@ +package activity + +import ( + "context" + "time" + + "management/pkg/configuration" + "management/pkg/logging" +) + +var log = logging.LoggerForThisPackage() + +type Manager struct { + cfg *config + // eventStore is the event store + eventStore Store +} + +// NewManager creates a new activity manager +func NewManager(eventStore Store) *Manager { + cfg, err := configuration.Parse[config]() + if err != nil { + log.Fatalf("failed to parse activity config: %v", err) + } + return &Manager{ + cfg: cfg, + eventStore: eventStore, + } +} + +func (m *Manager) StoreEvent(ctx context.Context, initiatorID, targetID, accountID string, activityID ActivityDescriber, meta map[string]any) { + if m.cfg.Enabled { + go func() { + _, err := m.eventStore.Save(ctx, &Event{ + Timestamp: time.Now().UTC(), + Activity: activityID, + InitiatorID: initiatorID, + TargetID: targetID, + AccountID: accountID, + Meta: meta, + }) + if err != nil { + // todo add metric + log.WithContext(ctx).Errorf("received an error while storing an activity event, error: %s", err) + } + }() + } +} diff --git a/internal/shared/activity/sqlite/crypt.go b/internal/shared/activity/sqlite/crypt.go new file mode 100644 index 0000000..096f49e --- /dev/null +++ b/internal/shared/activity/sqlite/crypt.go @@ -0,0 +1,136 @@ +package sqlite + +import ( + "bytes" + "crypto/aes" + "crypto/cipher" + "crypto/rand" + "encoding/base64" + "errors" +) + +var iv = []byte{10, 22, 13, 79, 05, 8, 52, 91, 87, 98, 88, 98, 35, 25, 13, 05} + +type FieldEncrypt struct { + block cipher.Block + gcm cipher.AEAD +} + +func GenerateKey() (string, error) { + key := make([]byte, 32) + _, err := rand.Read(key) + if err != nil { + return "", err + } + readableKey := base64.StdEncoding.EncodeToString(key) + return readableKey, nil +} + +func NewFieldEncrypt(key string) (*FieldEncrypt, error) { + binKey, err := base64.StdEncoding.DecodeString(key) + if err != nil { + return nil, err + } + + block, err := aes.NewCipher(binKey) + if err != nil { + return nil, err + } + + gcm, err := cipher.NewGCM(block) + if err != nil { + return nil, err + } + + ec := &FieldEncrypt{ + block: block, + gcm: gcm, + } + + return ec, nil +} + +func (ec *FieldEncrypt) LegacyEncrypt(payload string) string { + plainText := pkcs5Padding([]byte(payload)) + cipherText := make([]byte, len(plainText)) + cbc := cipher.NewCBCEncrypter(ec.block, iv) + cbc.CryptBlocks(cipherText, plainText) + return base64.StdEncoding.EncodeToString(cipherText) +} + +// Encrypt encrypts plaintext using AES-GCM +func (ec *FieldEncrypt) Encrypt(payload string) (string, error) { + plaintext := []byte(payload) + nonceSize := ec.gcm.NonceSize() + + nonce := make([]byte, nonceSize, len(plaintext)+nonceSize+ec.gcm.Overhead()) + if _, err := rand.Read(nonce); err != nil { + return "", err + } + + ciphertext := ec.gcm.Seal(nonce, nonce, plaintext, nil) + + return base64.StdEncoding.EncodeToString(ciphertext), nil +} + +func (ec *FieldEncrypt) LegacyDecrypt(data string) (string, error) { + cipherText, err := base64.StdEncoding.DecodeString(data) + if err != nil { + return "", err + } + cbc := cipher.NewCBCDecrypter(ec.block, iv) + cbc.CryptBlocks(cipherText, cipherText) + payload, err := pkcs5UnPadding(cipherText) + if err != nil { + return "", err + } + + return string(payload), nil +} + +// Decrypt decrypts ciphertext using AES-GCM +func (ec *FieldEncrypt) Decrypt(data string) (string, error) { + cipherText, err := base64.StdEncoding.DecodeString(data) + if err != nil { + return "", err + } + + nonceSize := ec.gcm.NonceSize() + if len(cipherText) < nonceSize { + return "", errors.New("cipher text too short") + } + + nonce, cipherText := cipherText[:nonceSize], cipherText[nonceSize:] + plainText, err := ec.gcm.Open(nil, nonce, cipherText, nil) + if err != nil { + return "", err + } + + return string(plainText), nil +} + +func pkcs5Padding(ciphertext []byte) []byte { + padding := aes.BlockSize - len(ciphertext)%aes.BlockSize + padText := bytes.Repeat([]byte{byte(padding)}, padding) + return append(ciphertext, padText...) +} +func pkcs5UnPadding(src []byte) ([]byte, error) { + srcLen := len(src) + if srcLen == 0 { + return nil, errors.New("input data is empty") + } + + paddingLen := int(src[srcLen-1]) + if paddingLen == 0 || paddingLen > aes.BlockSize || paddingLen > srcLen { + return nil, errors.New("invalid padding size") + } + + // Verify that all padding bytes are the same + for i := 0; i < paddingLen; i++ { + if src[srcLen-1-i] != byte(paddingLen) { + return nil, errors.New("invalid padding") + } + } + + return src[:srcLen-paddingLen], nil +} diff --git a/internal/shared/activity/sqlite/crypt_test.go b/internal/shared/activity/sqlite/crypt_test.go new file mode 100644 index 0000000..aff3a08 --- /dev/null +++ b/internal/shared/activity/sqlite/crypt_test.go @@ -0,0 +1,310 @@ +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/internal/shared/activity/sqlite/migration.go b/internal/shared/activity/sqlite/migration.go new file mode 100644 index 0000000..28c5b30 --- /dev/null +++ b/internal/shared/activity/sqlite/migration.go @@ -0,0 +1,157 @@ +package sqlite + +import ( + "context" + "database/sql" + "fmt" + + log "github.com/sirupsen/logrus" +) + +func migrate(ctx context.Context, crypt *FieldEncrypt, db *sql.DB) error { + if _, err := db.Exec(createTableQuery); err != nil { + return err + } + + if _, err := db.Exec(creatTableDeletedUsersQuery); err != nil { + return err + } + + if err := updateDeletedUsersTable(ctx, db); err != nil { + return fmt.Errorf("failed to update deleted_users table: %v", err) + } + + return migrateLegacyEncryptedUsersToGCM(ctx, crypt, db) +} + +// updateDeletedUsersTable checks and updates the deleted_users table schema to ensure required columns exist. +func updateDeletedUsersTable(ctx context.Context, db *sql.DB) error { + exists, err := checkColumnExists(db, "deleted_users", "name") + if err != nil { + return err + } + + if !exists { + log.WithContext(ctx).Debug("Adding name column to the deleted_users table") + + _, err = db.Exec(`ALTER TABLE deleted_users ADD COLUMN name TEXT;`) + if err != nil { + return err + } + + log.WithContext(ctx).Debug("Successfully added name column to the deleted_users table") + } + + exists, err = checkColumnExists(db, "deleted_users", "enc_algo") + if err != nil { + return err + } + + if !exists { + log.WithContext(ctx).Debug("Adding enc_algo column to the deleted_users table") + + _, err = db.Exec(`ALTER TABLE deleted_users ADD COLUMN enc_algo TEXT;`) + if err != nil { + return err + } + + log.WithContext(ctx).Debug("Successfully added enc_algo column to the deleted_users table") + } + + return nil +} + +// migrateLegacyEncryptedUsersToGCM migrates previously encrypted data using, +// legacy CBC encryption with a static IV to the new GCM encryption method. +func migrateLegacyEncryptedUsersToGCM(ctx context.Context, crypt *FieldEncrypt, db *sql.DB) error { + log.WithContext(ctx).Debug("Migrating CBC encrypted deleted users to GCM") + + tx, err := db.Begin() + if err != nil { + return fmt.Errorf("failed to begin transaction: %v", err) + } + defer func() { + _ = tx.Rollback() + }() + + rows, err := tx.Query(fmt.Sprintf(`SELECT id, email, name FROM deleted_users where enc_algo IS NULL OR enc_algo != '%s'`, gcmEncAlgo)) + if err != nil { + return fmt.Errorf("failed to execute select query: %v", err) + } + defer rows.Close() + + updateStmt, err := tx.Prepare(`UPDATE deleted_users SET email = ?, name = ?, enc_algo = ? WHERE id = ?`) + if err != nil { + return fmt.Errorf("failed to prepare update statement: %v", err) + } + defer updateStmt.Close() + + if err = processUserRows(ctx, crypt, rows, updateStmt); err != nil { + return err + } + + if err = tx.Commit(); err != nil { + return fmt.Errorf("failed to commit transaction: %v", err) + } + + log.WithContext(ctx).Debug("Successfully migrated CBC encrypted deleted users to GCM") + return nil +} + +// processUserRows processes database rows of user data, decrypts legacy encryption fields, and re-encrypts them using GCM. +func processUserRows(ctx context.Context, crypt *FieldEncrypt, rows *sql.Rows, updateStmt *sql.Stmt) error { + for rows.Next() { + var ( + id, decryptedEmail, decryptedName string + email, name *string + ) + + err := rows.Scan(&id, &email, &name) + if err != nil { + return err + } + + if email != nil { + decryptedEmail, err = crypt.LegacyDecrypt(*email) + if err != nil { + log.WithContext(ctx).Warnf("skipping migrating deleted user %s: %v", + id, + fmt.Errorf("failed to decrypt email: %w", err), + ) + continue + } + } + + if name != nil { + decryptedName, err = crypt.LegacyDecrypt(*name) + if err != nil { + log.WithContext(ctx).Warnf("skipping migrating deleted user %s: %v", + id, + fmt.Errorf("failed to decrypt name: %w", err), + ) + continue + } + } + + encryptedEmail, err := crypt.Encrypt(decryptedEmail) + if err != nil { + return fmt.Errorf("failed to encrypt email: %w", err) + } + + encryptedName, err := crypt.Encrypt(decryptedName) + if err != nil { + return fmt.Errorf("failed to encrypt name: %w", err) + } + + _, err = updateStmt.Exec(encryptedEmail, encryptedName, gcmEncAlgo, id) + if err != nil { + return err + } + } + + if err := rows.Err(); err != nil { + return err + } + + return nil +} diff --git a/internal/shared/activity/sqlite/migration_test.go b/internal/shared/activity/sqlite/migration_test.go new file mode 100644 index 0000000..a03774f --- /dev/null +++ b/internal/shared/activity/sqlite/migration_test.go @@ -0,0 +1,84 @@ +package sqlite + +import ( + "context" + "database/sql" + "path/filepath" + "testing" + "time" + + _ "github.com/mattn/go-sqlite3" + "github.com/netbirdio/netbird/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/internal/shared/activity/sqlite/sqlite.go b/internal/shared/activity/sqlite/sqlite.go new file mode 100644 index 0000000..ffb863d --- /dev/null +++ b/internal/shared/activity/sqlite/sqlite.go @@ -0,0 +1,359 @@ +package sqlite + +import ( + "context" + "database/sql" + "encoding/json" + "fmt" + "path/filepath" + "runtime" + "time" + + _ "github.com/mattn/go-sqlite3" + log "github.com/sirupsen/logrus" + + "github.com/netbirdio/netbird/management/server/activity" +) + +const ( + // eventSinkDB is the default name of the events database + eventSinkDB = "events.db" + createTableQuery = "CREATE TABLE IF NOT EXISTS events " + + "(id INTEGER PRIMARY KEY AUTOINCREMENT, " + + "activity INTEGER, " + + "timestamp DATETIME, " + + "initiator_id TEXT," + + "account_id TEXT," + + "meta TEXT," + + " target_id TEXT);" + + creatTableDeletedUsersQuery = `CREATE TABLE IF NOT EXISTS deleted_users (id TEXT NOT NULL, email TEXT NOT NULL, name TEXT, enc_algo TEXT NOT NULL);` + + selectDescQuery = `SELECT events.id, activity, timestamp, initiator_id, i.name as "initiator_name", i.email as "initiator_email", target_id, t.name as "target_name", t.email as "target_email", account_id, meta + FROM events + LEFT JOIN ( + SELECT id, MAX(name) as name, MAX(email) as email + FROM deleted_users + GROUP BY id + ) i ON events.initiator_id = i.id + LEFT JOIN ( + SELECT id, MAX(name) as name, MAX(email) as email + FROM deleted_users + GROUP BY id + ) t ON events.target_id = t.id + WHERE account_id = ? + ORDER BY timestamp DESC LIMIT ? OFFSET ?;` + + selectAscQuery = `SELECT events.id, activity, timestamp, initiator_id, i.name as "initiator_name", i.email as "initiator_email", target_id, t.name as "target_name", t.email as "target_email", account_id, meta + FROM events + LEFT JOIN ( + SELECT id, MAX(name) as name, MAX(email) as email + FROM deleted_users + GROUP BY id + ) i ON events.initiator_id = i.id + LEFT JOIN ( + SELECT id, MAX(name) as name, MAX(email) as email + FROM deleted_users + GROUP BY id + ) t ON events.target_id = t.id + WHERE account_id = ? + ORDER BY timestamp ASC LIMIT ? OFFSET ?;` + + insertQuery = "INSERT INTO events(activity, timestamp, initiator_id, target_id, account_id, meta) " + + "VALUES(?, ?, ?, ?, ?, ?)" + + /* + TODO: + The insert should avoid duplicated IDs in the table. So the query should be changes to something like: + `INSERT INTO deleted_users(id, email, name) VALUES(?, ?, ?) ON CONFLICT (id) DO UPDATE SET email = EXCLUDED.email, name = EXCLUDED.name;` + For this to work we have to set the id column as primary key. But this is not possible because the id column is not unique + and some selfhosted deployments might have duplicates already so we need to clean the table first. + */ + + insertDeleteUserQuery = `INSERT INTO deleted_users(id, email, name, enc_algo) VALUES(?, ?, ?, ?)` + + fallbackName = "unknown" + fallbackEmail = "unknown@unknown.com" + + gcmEncAlgo = "GCM" +) + +// Store is the implementation of the activity.Store interface backed by SQLite +type Store struct { + db *sql.DB + fieldEncrypt *FieldEncrypt + + insertStatement *sql.Stmt + selectAscStatement *sql.Stmt + selectDescStatement *sql.Stmt + deleteUserStmt *sql.Stmt +} + +// NewSQLiteStore creates a new Store with an event table if not exists. +func NewSQLiteStore(ctx context.Context, dataDir string, encryptionKey string) (*Store, error) { + dbFile := filepath.Join(dataDir, eventSinkDB) + db, err := sql.Open("sqlite3", dbFile) + if err != nil { + return nil, err + } + db.SetMaxOpenConns(runtime.NumCPU()) + + crypt, err := NewFieldEncrypt(encryptionKey) + if err != nil { + _ = db.Close() + return nil, err + } + + if err = migrate(ctx, crypt, db); err != nil { + _ = db.Close() + return nil, fmt.Errorf("events database migration: %w", err) + } + + return createStore(crypt, db) +} + +func (store *Store) processResult(ctx context.Context, result *sql.Rows) ([]*activity.Event, error) { + events := make([]*activity.Event, 0) + var cryptErr error + for result.Next() { + var id int64 + var operation activity.Activity + var timestamp time.Time + var initiator string + var initiatorName *string + var initiatorEmail *string + var target string + var targetUserName *string + var targetEmail *string + var account string + var jsonMeta string + err := result.Scan(&id, &operation, ×tamp, &initiator, &initiatorName, &initiatorEmail, &target, &targetUserName, &targetEmail, &account, &jsonMeta) + if err != nil { + return nil, err + } + + meta := make(map[string]any) + if jsonMeta != "" { + err = json.Unmarshal([]byte(jsonMeta), &meta) + if err != nil { + return nil, err + } + } + + if targetUserName != nil { + name, err := store.fieldEncrypt.Decrypt(*targetUserName) + if err != nil { + cryptErr = fmt.Errorf("failed to decrypt username for target id: %s", target) + meta["username"] = fallbackName + } else { + meta["username"] = name + } + } + + if targetEmail != nil { + email, err := store.fieldEncrypt.Decrypt(*targetEmail) + if err != nil { + cryptErr = fmt.Errorf("failed to decrypt email address for target id: %s", target) + meta["email"] = fallbackEmail + } else { + meta["email"] = email + } + } + + event := &activity.Event{ + Timestamp: timestamp, + Activity: operation, + ID: uint64(id), + InitiatorID: initiator, + TargetID: target, + AccountID: account, + Meta: meta, + } + + if initiatorName != nil { + name, err := store.fieldEncrypt.Decrypt(*initiatorName) + if err != nil { + cryptErr = fmt.Errorf("failed to decrypt username of initiator: %s", initiator) + event.InitiatorName = fallbackName + } else { + event.InitiatorName = name + } + } + + if initiatorEmail != nil { + email, err := store.fieldEncrypt.Decrypt(*initiatorEmail) + if err != nil { + cryptErr = fmt.Errorf("failed to decrypt email address of initiator: %s", initiator) + event.InitiatorEmail = fallbackEmail + } else { + event.InitiatorEmail = email + } + } + + events = append(events, event) + } + + if cryptErr != nil { + log.WithContext(ctx).Warnf("%s", cryptErr) + } + + return events, nil +} + +// Get returns "limit" number of events from index ordered descending or ascending by a timestamp +func (store *Store) Get(ctx context.Context, accountID string, offset, limit int, descending bool) ([]*activity.Event, error) { + stmt := store.selectDescStatement + if !descending { + stmt = store.selectAscStatement + } + + result, err := stmt.Query(accountID, limit, offset) + if err != nil { + return nil, err + } + + defer result.Close() //nolint + return store.processResult(ctx, result) +} + +// Save an event in the SQLite events table end encrypt the "email" element in meta map +func (store *Store) Save(_ context.Context, event *activity.Event) (*activity.Event, error) { + var jsonMeta string + meta, err := store.saveDeletedUserEmailAndNameInEncrypted(event) + if err != nil { + return nil, err + } + + if meta != nil { + metaBytes, err := json.Marshal(event.Meta) + if err != nil { + return nil, err + } + jsonMeta = string(metaBytes) + } + + result, err := store.insertStatement.Exec(event.Activity, event.Timestamp, event.InitiatorID, event.TargetID, event.AccountID, jsonMeta) + if err != nil { + return nil, err + } + + id, err := result.LastInsertId() + if err != nil { + return nil, err + } + + eventCopy := event.Copy() + eventCopy.ID = uint64(id) + return eventCopy, nil +} + +// saveDeletedUserEmailAndNameInEncrypted if the meta contains email and name then store it in encrypted way and delete +// this item from meta map +func (store *Store) saveDeletedUserEmailAndNameInEncrypted(event *activity.Event) (map[string]any, error) { + email, ok := event.Meta["email"] + if !ok { + return event.Meta, nil + } + + name, ok := event.Meta["name"] + if !ok { + return event.Meta, nil + } + + encryptedEmail, err := store.fieldEncrypt.Encrypt(fmt.Sprintf("%s", email)) + if err != nil { + return nil, err + } + encryptedName, err := store.fieldEncrypt.Encrypt(fmt.Sprintf("%s", name)) + if err != nil { + return nil, err + } + + _, err = store.deleteUserStmt.Exec(event.TargetID, encryptedEmail, encryptedName, gcmEncAlgo) + if err != nil { + return nil, err + } + + if len(event.Meta) == 2 { + return nil, nil // nolint + } + delete(event.Meta, "email") + delete(event.Meta, "name") + return event.Meta, nil +} + +// Close the Store +func (store *Store) Close(_ context.Context) error { + if store.db != nil { + return store.db.Close() + } + return nil +} + +// createStore initializes and returns a new Store instance with prepared SQL statements. +func createStore(crypt *FieldEncrypt, db *sql.DB) (*Store, error) { + insertStmt, err := db.Prepare(insertQuery) + if err != nil { + _ = db.Close() + return nil, err + } + + selectDescStmt, err := db.Prepare(selectDescQuery) + if err != nil { + _ = db.Close() + return nil, err + } + + selectAscStmt, err := db.Prepare(selectAscQuery) + if err != nil { + _ = db.Close() + return nil, err + } + + deleteUserStmt, err := db.Prepare(insertDeleteUserQuery) + if err != nil { + _ = db.Close() + return nil, err + } + + return &Store{ + db: db, + fieldEncrypt: crypt, + insertStatement: insertStmt, + selectDescStatement: selectDescStmt, + selectAscStatement: selectAscStmt, + deleteUserStmt: deleteUserStmt, + }, nil +} + +// checkColumnExists checks if a column exists in a specified table +func checkColumnExists(db *sql.DB, tableName, columnName string) (bool, error) { + query := fmt.Sprintf("PRAGMA table_info(%s);", tableName) + rows, err := db.Query(query) + if err != nil { + return false, fmt.Errorf("failed to query table info: %w", err) + } + defer rows.Close() + + for rows.Next() { + var cid int + var name, ctype string + var notnull, pk int + var dfltValue sql.NullString + + err = rows.Scan(&cid, &name, &ctype, ¬null, &dfltValue, &pk) + if err != nil { + return false, fmt.Errorf("failed to scan row: %w", err) + } + + if name == columnName { + return true, nil + } + } + + if err = rows.Err(); err != nil { + return false, err + } + + return false, nil +} diff --git a/internal/shared/activity/sqlite/sqlite_test.go b/internal/shared/activity/sqlite/sqlite_test.go new file mode 100644 index 0000000..b10f9b5 --- /dev/null +++ b/internal/shared/activity/sqlite/sqlite_test.go @@ -0,0 +1,57 @@ +package sqlite + +import ( + "context" + "fmt" + "testing" + "time" + + "github.com/stretchr/testify/assert" + + "github.com/netbirdio/netbird/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/internal/shared/activity/store.go b/internal/shared/activity/store.go new file mode 100644 index 0000000..ef08e2b --- /dev/null +++ b/internal/shared/activity/store.go @@ -0,0 +1,57 @@ +package activity + +import ( + "context" + "sync" +) + +// Store provides an interface to store or stream events. +type Store interface { + // Save an event in the store + Save(ctx context.Context, event *Event) (*Event, error) + // Get returns "limit" number of events from the "offset" index ordered descending or ascending by a timestamp + Get(ctx context.Context, accountID string, offset, limit int, descending bool) ([]*Event, error) + // Close the sink flushing events if necessary + Close(ctx context.Context) error +} + +// InMemoryEventStore implements the Store interface storing data in-memory +type InMemoryEventStore struct { + mu sync.Mutex + nextID uint64 + events []*Event +} + +// Save sets the Event.ID to 1 +func (store *InMemoryEventStore) Save(_ context.Context, event *Event) (*Event, error) { + store.mu.Lock() + defer store.mu.Unlock() + if store.events == nil { + store.events = make([]*Event, 0) + } + event.ID = store.nextID + store.nextID++ + store.events = append(store.events, event) + return event, nil +} + +// Get returns a list of ALL events that belong to the given accountID without taking offset, limit and order into consideration +func (store *InMemoryEventStore) Get(_ context.Context, accountID string, offset, limit int, descending bool) ([]*Event, error) { + store.mu.Lock() + defer store.mu.Unlock() + events := make([]*Event, 0) + for _, event := range store.events { + if event.AccountID == accountID { + events = append(events, event) + } + } + return events, nil +} + +// Close cleans up the event list +func (store *InMemoryEventStore) Close(_ context.Context) error { + store.mu.Lock() + defer store.mu.Unlock() + store.events = make([]*Event, 0) + return nil +} diff --git a/internal/shared/api/middleware/auth_middleware.go b/internal/shared/api/middleware/auth_middleware.go new file mode 100644 index 0000000..e12a98f --- /dev/null +++ b/internal/shared/api/middleware/auth_middleware.go @@ -0,0 +1,189 @@ +package middleware + +import ( + "context" + "errors" + "fmt" + "net/http" + "strings" + "time" + + "github.com/google/uuid" + log "github.com/sirupsen/logrus" + + "github.com/netbirdio/netbird/management/server/auth" + nbcontext "github.com/netbirdio/netbird/management/server/context" + "github.com/netbirdio/netbird/management/server/http/middleware/bypass" + "github.com/netbirdio/netbird/management/server/http/util" + "github.com/netbirdio/netbird/management/server/status" + + "management/pkg/logging/hook" +) + +type EnsureAccountFunc func(ctx context.Context, userAuth nbcontext.UserAuth) (string, string, error) +type SyncUserJWTGroupsFunc func(ctx context.Context, userAuth nbcontext.UserAuth) error + +// AuthMiddleware middleware to verify personal access tokens (PAT) and JWT tokens +type AuthMiddleware struct { + authManager auth.Manager + ensureAccount EnsureAccountFunc + syncUserJWTGroups SyncUserJWTGroupsFunc +} + +// NewAuthMiddleware instance constructor +func NewAuthMiddleware( + authManager auth.Manager, + ensureAccount EnsureAccountFunc, + syncUserJWTGroups SyncUserJWTGroupsFunc, +) *AuthMiddleware { + return &AuthMiddleware{ + authManager: authManager, + ensureAccount: ensureAccount, + syncUserJWTGroups: syncUserJWTGroups, + } +} + +// Handler method of the middleware which authenticates a user either by JWT claims or by PAT +func (m *AuthMiddleware) Handler(h http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + //nolint + ctx := context.WithValue(r.Context(), hook.ExecutionContextKey, hook.HTTPSource) + + reqID := uuid.New().String() + //nolint + ctx = context.WithValue(ctx, nbcontext.RequestIDKey, reqID) + + if bypass.ShouldBypass(r.URL.Path, h, w, r) { + return + } + + auth := strings.Split(r.Header.Get("Authorization"), " ") + authType := strings.ToLower(auth[0]) + + // fallback to token when receive pat as bearer + if len(auth) >= 2 && authType == "bearer" && strings.HasPrefix(auth[1], "nbp_") { + authType = "token" + auth[0] = authType + } + + switch authType { + case "bearer": + request, err := m.checkJWTFromRequest(r, auth) + if err != nil { + log.WithContext(r.Context()).Errorf("Error when validating JWT: %s", err.Error()) + util.WriteError(r.Context(), status.Errorf(status.Unauthorized, "token invalid"), w) + return + } + + h.ServeHTTP(w, request) + case "token": + request, err := m.checkPATFromRequest(r, auth) + if err != nil { + log.WithContext(r.Context()).Debugf("Error when validating PAT: %s", err.Error()) + util.WriteError(r.Context(), status.Errorf(status.Unauthorized, "token invalid"), w) + return + } + h.ServeHTTP(w, request) + default: + util.WriteError(r.Context(), status.Errorf(status.Unauthorized, "no valid authentication provided"), w) + return + } + }) +} + +// CheckJWTFromRequest checks if the JWT is valid +func (m *AuthMiddleware) checkJWTFromRequest(r *http.Request, auth []string) (*http.Request, error) { + token, err := getTokenFromJWTRequest(auth) + + // If an error occurs, call the error handler and return an error + if err != nil { + return r, fmt.Errorf("error extracting token: %w", err) + } + + ctx := r.Context() + + userAuth, validatedToken, err := m.authManager.ValidateAndParseToken(ctx, token) + if err != nil { + return r, err + } + + if impersonate, ok := r.URL.Query()["account"]; ok && len(impersonate) == 1 { + userAuth.AccountId = impersonate[0] + userAuth.IsChild = ok + } + + // 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 := m.ensureAccount(ctx, userAuth) + if err != nil { + return r, err + } + + if userAuth.AccountId != accountId { + log.WithContext(ctx).Debugf("Auth middleware sets accountId from ensure, before %s, now %s", userAuth.AccountId, accountId) + userAuth.AccountId = accountId + } + + userAuth, err = m.authManager.EnsureUserAccessByJWTGroups(ctx, userAuth, validatedToken) + if err != nil { + return r, err + } + + err = m.syncUserJWTGroups(ctx, userAuth) + if err != nil { + log.WithContext(ctx).Errorf("HTTP server failed to sync user JWT groups: %s", err) + } + + return nbcontext.SetUserAuthInRequest(r, userAuth), nil +} + +// CheckPATFromRequest checks if the PAT is valid +func (m *AuthMiddleware) checkPATFromRequest(r *http.Request, auth []string) (*http.Request, error) { + token, err := getTokenFromPATRequest(auth) + if err != nil { + return r, fmt.Errorf("error extracting token: %w", err) + } + + ctx := r.Context() + user, pat, accDomain, accCategory, err := m.authManager.GetPATInfo(ctx, token) + if err != nil { + return r, fmt.Errorf("invalid Token: %w", err) + } + if time.Now().After(pat.GetExpirationDate()) { + return r, fmt.Errorf("token expired") + } + + err = m.authManager.MarkPATUsed(ctx, pat.ID) + if err != nil { + return r, err + } + + userAuth := nbcontext.UserAuth{ + UserId: user.Id, + AccountId: user.AccountID, + Domain: accDomain, + DomainCategory: accCategory, + IsPAT: true, + } + + return nbcontext.SetUserAuthInRequest(r, userAuth), nil +} + +// getTokenFromJWTRequest is a "TokenExtractor" that takes auth header parts and extracts +// the JWT token from the Authorization header. +func getTokenFromJWTRequest(authHeaderParts []string) (string, error) { + if len(authHeaderParts) != 2 || strings.ToLower(authHeaderParts[0]) != "bearer" { + return "", errors.New("authorization header format must be Bearer {token}") + } + + return authHeaderParts[1], nil +} + +// getTokenFromPATRequest is a "TokenExtractor" that takes auth header parts and extracts +// the PAT token from the Authorization header. +func getTokenFromPATRequest(authHeaderParts []string) (string, error) { + if len(authHeaderParts) != 2 || strings.ToLower(authHeaderParts[0]) != "token" { + return "", errors.New("authorization header format must be Token {token}") + } + + return authHeaderParts[1], nil +} diff --git a/internal/shared/api/middleware/auth_middleware_test.go b/internal/shared/api/middleware/auth_middleware_test.go new file mode 100644 index 0000000..3dc7d51 --- /dev/null +++ b/internal/shared/api/middleware/auth_middleware_test.go @@ -0,0 +1,325 @@ +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/netbird/management/server/auth" + nbjwt "github.com/netbirdio/netbird/management/server/auth/jwt" + nbcontext "github.com/netbirdio/netbird/management/server/context" + "github.com/netbirdio/netbird/management/server/util" + + "github.com/netbirdio/netbird/management/server/http/middleware/bypass" + "github.com/netbirdio/netbird/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/internal/shared/api/middleware/bypass/bypass.go b/internal/shared/api/middleware/bypass/bypass.go new file mode 100644 index 0000000..9447704 --- /dev/null +++ b/internal/shared/api/middleware/bypass/bypass.go @@ -0,0 +1,74 @@ +package bypass + +import ( + "fmt" + "net/http" + "path" + "sync" + + log "github.com/sirupsen/logrus" +) + +var byPassMutex sync.RWMutex + +// bypassPaths is a set of paths that should bypass middleware. +var bypassPaths = make(map[string]struct{}) + +// AddBypassPath adds an exact path to the list of paths that bypass middleware. +// Paths can include wildcards, such as /api/*. Paths are matched using path.Match. +// Returns an error if the path has invalid pattern. +func AddBypassPath(path string) error { + byPassMutex.Lock() + defer byPassMutex.Unlock() + if err := validatePath(path); err != nil { + return fmt.Errorf("validate: %w", err) + } + bypassPaths[path] = struct{}{} + return nil +} + +// RemovePath removes a path from the list of paths that bypass middleware. +func RemovePath(path string) { + byPassMutex.Lock() + defer byPassMutex.Unlock() + delete(bypassPaths, path) +} + +// GetList returns a list of all bypass paths. +func GetList() []string { + byPassMutex.RLock() + defer byPassMutex.RUnlock() + + list := make([]string, 0, len(bypassPaths)) + for k := range bypassPaths { + list = append(list, k) + } + + return list +} + +// ShouldBypass checks if the request path is one of the auth bypass paths and returns true if the middleware should be bypassed. +// This can be used to bypass authz/authn middlewares for certain paths, such as webhooks that implement their own authentication. +func ShouldBypass(requestPath string, h http.Handler, w http.ResponseWriter, r *http.Request) bool { + byPassMutex.RLock() + defer byPassMutex.RUnlock() + + for bypassPath := range bypassPaths { + matched, err := path.Match(bypassPath, requestPath) + if err != nil { + log.WithContext(r.Context()).Errorf("Error matching path %s with %s from %s: %v", bypassPath, requestPath, GetList(), err) + continue + } + if matched { + h.ServeHTTP(w, r) + return true + } + } + + return false +} + +func validatePath(p string) error { + _, err := path.Match(p, "") + return err +} diff --git a/internal/shared/api/middleware/bypass/bypass_test.go b/internal/shared/api/middleware/bypass/bypass_test.go new file mode 100644 index 0000000..c65e6fa --- /dev/null +++ b/internal/shared/api/middleware/bypass/bypass_test.go @@ -0,0 +1,131 @@ +package bypass_test + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/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/internal/shared/api/middleware/metrics_middleware.go b/internal/shared/api/middleware/metrics_middleware.go new file mode 100644 index 0000000..c870d7c --- /dev/null +++ b/internal/shared/api/middleware/metrics_middleware.go @@ -0,0 +1 @@ +package middleware diff --git a/internal/shared/api/router.go b/internal/shared/api/router.go new file mode 100644 index 0000000..9c56169 --- /dev/null +++ b/internal/shared/api/router.go @@ -0,0 +1,43 @@ +package api + +import ( + "net/http" + + "github.com/gorilla/mux" +) + +// NewRouter creates and returns a mux.Router configured with default middleware +// and placeholder endpoints. You can add your own handlers here or in other files. +func NewRouter() http.Handler { + r := mux.NewRouter() + + // Attach middlewares + r.Use(loggingMiddleware) + r.Use(recoveryMiddleware) + + // Example endpoint + // r.HandleFunc("/health", healthCheckHandler).Methods("GET") + + return r +} + +// loggingMiddleware is an example that logs each incoming request. +// Replace with your logger of choice. +func loggingMiddleware(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // log.Printf("[%s] %s", r.Method, r.URL.Path) + next.ServeHTTP(w, r) + }) +} + +// recoveryMiddleware recovers from panics and returns a 500 Internal Server Error. +func recoveryMiddleware(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + defer func() { + if rec := recover(); rec != nil { + http.Error(w, "Internal Server Error", http.StatusInternalServerError) + } + }() + next.ServeHTTP(w, r) + }) +} diff --git a/internal/shared/db/config.go b/internal/shared/db/config.go new file mode 100644 index 0000000..d4b2462 --- /dev/null +++ b/internal/shared/db/config.go @@ -0,0 +1,6 @@ +package db + +type config struct { + Engine string `env:"NB_STORE_ENGINE" envDefault:"sqlite"` + PostgresDsnEnv string `env:"NB_STORE_ENGINE_POSTGRES_DSN" envDefault:""` +} diff --git a/internal/shared/db/database_connection.go b/internal/shared/db/database_connection.go new file mode 100644 index 0000000..d3fff24 --- /dev/null +++ b/internal/shared/db/database_connection.go @@ -0,0 +1,120 @@ +package db + +import ( + "context" + "fmt" + "os" + "path/filepath" + "runtime" + + "gorm.io/driver/postgres" + "gorm.io/driver/sqlite" + "gorm.io/gorm" + "gorm.io/gorm/logger" + + "management/internal/shared/errors" + "management/pkg/configuration" +) + +const ( + storeSqliteFileName = "licenses.db" + storeDataDirEnv = "NB_STORE_DATA_DIR" + storeDefaultDataDir = "/var/lib/netbird" +) + +// DatabaseConn is a wrapper around the gorm database connection +type DatabaseConn struct { + DB *gorm.DB +} + +// NewDatabaseConn creates a new database connection based on the store engine +func NewDatabaseConn(ctx context.Context) (*DatabaseConn, error) { + cfg, err := configuration.Parse[config]() + if err != nil { + log.Fatalf("failed to parse config: %v", err) + } + + log.WithContext(ctx).Infof("using %s store engine", cfg.Engine) + + var db *gorm.DB + switch Engine(cfg.Engine) { + case SqliteStoreEngine: + db, err = openSQLiteDB() + case PostgresStoreEngine: + db, err = openPostgresDB(cfg) + case MemoryStoreEngine: + db, err = openMemoryDB() + default: + err = errors.NewUnsupportedStoreEngineConfigError(cfg.Engine) + } + + if err != nil || db == nil { + return nil, fmt.Errorf("error while opening database: %w", err) + } + + sql, err := db.DB() + if err != nil { + return nil, fmt.Errorf("error getting sql db connection: %w", err) + } + + conns := runtime.NumCPU() + sql.SetMaxOpenConns(conns) + return &DatabaseConn{ + DB: db, + }, nil +} + +// openMemoryDB opens a new connection to an in-memory SQLite database for testing +func openMemoryDB() (*gorm.DB, error) { + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) + if err != nil { + return nil, err + } + + return db, nil +} + +// openSQLiteDB opens a new connection to a SQLite database +func openSQLiteDB() (*gorm.DB, error) { + storeStr := fmt.Sprintf("%s?cache=shared", storeSqliteFileName) + if runtime.GOOS == "windows" { + // To avoid `The process cannot access the file because it is being used by another process` on Windows + storeStr = storeSqliteFileName + } + + dataDir, ok := os.LookupEnv(storeDataDirEnv) + if !ok { + dataDir = storeDefaultDataDir + } + + file := filepath.Join(dataDir, storeStr) + db, err := gorm.Open(sqlite.Open(file), getGormConfig()) + if err != nil { + return nil, err + } + + return db, nil +} + +// openPostgresDB opens a new connection to a Postgres database +func openPostgresDB(cfg *config) (*gorm.DB, error) { + dsn, ok := os.LookupEnv(cfg.PostgresDsnEnv) + if !ok { + return nil, fmt.Errorf("%s is not set", cfg.PostgresDsnEnv) + } + + db, err := gorm.Open(postgres.Open(dsn), getGormConfig()) + if err != nil { + return nil, err + } + return db, nil +} + +// getGormConfig returns the gorm configuration +func getGormConfig() *gorm.Config { + return &gorm.Config{ + Logger: logger.Default.LogMode(logger.Silent), + CreateBatchSize: 400, + PrepareStmt: false, + } +} diff --git a/internal/shared/db/store.go b/internal/shared/db/store.go new file mode 100644 index 0000000..cb2d4ad --- /dev/null +++ b/internal/shared/db/store.go @@ -0,0 +1,86 @@ +package db + +import ( + "context" + + "gorm.io/gorm" + "gorm.io/gorm/clause" + + "management/pkg/logging" +) + +var log = logging.LoggerForThisPackage() + +type Store struct { + db *gorm.DB +} + +func NewStore(ctx context.Context, dbConn *DatabaseConn) *Store { + return &Store{db: dbConn.DB} +} + +func (s *Store) Begin() (Transaction, error) { + tx := s.db.Begin() + if tx.Error != nil { + return nil, tx.Error + } + return &storeTx{db: tx}, nil +} + +// AutoMigrate automatically migrates your schema, to keep your database up to date. +func (s *Store) AutoMigrate(value interface{}) error { + return s.db.AutoMigrate(value) +} + +// Using picks the underlying db +func (s *Store) Using(tx Transaction) *gorm.DB { + if tx == nil { + return s.db + } + if st, ok := tx.(*storeTx); ok { + return st.db + } + return s.db +} + +// RunInTx is the new helper that starts a transaction, calls fn, and commits/rolls back automatically +func (s *Store) RunInTx(fn func(tx Transaction) error) error { + tx, err := s.Begin() + if err != nil { + return err + } + + if err := fn(tx); err != nil { + _ = tx.Rollback() + return err + } + return tx.Commit() +} + +func (s *Store) Create(tx Transaction, value interface{}) error { + return s.Using(tx).Create(value).Error +} + +func (s *Store) GetOne(tx Transaction, strength LockingStrength, dest interface{}, query string, args ...interface{}) error { + db := s.Using(tx).Clauses(clause.Locking{Strength: string(strength)}) + + if query != "" && len(args) > 0 { + db.Where(query, args...) + } + + return db.First(dest).Error +} + +func (s *Store) GetMany(tx Transaction, strength LockingStrength, dest interface{}, query string, args ...interface{}) error { + db := s.Using(tx).Clauses(clause.Locking{Strength: string(strength)}) + + if query != "" && len(args) > 0 { + db.Where(query, args...) + } + + return db.Find(dest).Error +} + +func (s *Store) Delete(value interface{}) error { + return s.db.Delete(value).Error +} diff --git a/internal/shared/db/store_engines.go b/internal/shared/db/store_engines.go new file mode 100644 index 0000000..29eb012 --- /dev/null +++ b/internal/shared/db/store_engines.go @@ -0,0 +1,19 @@ +package db + +import "slices" + +// Engine represents the db engine to use. +type Engine string + +const ( + SqliteStoreEngine Engine = "sqlite" + PostgresStoreEngine Engine = "postgres" + MemoryStoreEngine Engine = "memory" + MysqlStoreEngine Engine = "mysql" +) + +var supportedEngines = []Engine{SqliteStoreEngine, PostgresStoreEngine, MysqlStoreEngine} + +func IsSupportedEngine(engine Engine) bool { + return slices.Contains(supportedEngines, engine) +} diff --git a/internal/shared/db/transaction.go b/internal/shared/db/transaction.go new file mode 100644 index 0000000..49708b3 --- /dev/null +++ b/internal/shared/db/transaction.go @@ -0,0 +1,30 @@ +package db + +import "gorm.io/gorm" + +type LockingStrength string + +const ( + LockingStrengthUpdate LockingStrength = "UPDATE" // Strongest lock, preventing any changes by other transactions until your transaction completes. + LockingStrengthShare LockingStrength = "SHARE" // Allows reading but prevents changes by other transactions. + LockingStrengthNoKeyUpdate LockingStrength = "NO KEY UPDATE" // Similar to UPDATE but allows changes to related rows. + LockingStrengthKeyShare LockingStrength = "KEY SHARE" // Protects against changes to primary/unique keys but allows other updates. +) + +// Transaction interface (unchanged) +type Transaction interface { + Commit() error + Rollback() error +} + +type storeTx struct { + db *gorm.DB +} + +func (t *storeTx) Commit() error { + return t.db.Commit().Error +} + +func (t *storeTx) Rollback() error { + return t.db.Rollback().Error +} diff --git a/internal/shared/errors/errors.go b/internal/shared/errors/errors.go new file mode 100644 index 0000000..16b4a20 --- /dev/null +++ b/internal/shared/errors/errors.go @@ -0,0 +1,24 @@ +package errors + +import ( + "errors" + "fmt" +) + +var ( + UnsupportedStoreEngineConfigError = errors.New("unsupported store engine") + PermissionValidationError = errors.New("permission validation failed") + PermissionDeniedError = errors.New("permission denied") +) + +func NewUnsupportedStoreEngineConfigError(engine string) error { + return fmt.Errorf("%w: %s", UnsupportedStoreEngineConfigError, engine) +} + +func NewPermissionValidationError(err error) error { + return fmt.Errorf("%w: %s", PermissionValidationError, err) +} + +func NewPermissionDeniedError() error { + return fmt.Errorf("%w", PermissionDeniedError) +} diff --git a/internal/shared/permissions/manager.go b/internal/shared/permissions/manager.go new file mode 100644 index 0000000..c758a93 --- /dev/null +++ b/internal/shared/permissions/manager.go @@ -0,0 +1,104 @@ +package permissions + +//go:generate go run github.com/golang/mock/mockgen -package permissions -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod + +import ( + "context" + + "github.com/netbirdio/netbird/management/server/status" + "github.com/netbirdio/netbird/management/server/store" + + "management/internal/modules/users/types" + "management/internal/shared/activity" + "management/internal/shared/permissions/modules" + "management/internal/shared/permissions/operations" + "management/internal/shared/permissions/roles" + "management/pkg/logging" +) + +var log = logging.LoggerForThisPackage() + +type Manager interface { + ValidateUserPermissions(ctx context.Context, accountID, userID string, module modules.Module, operation operations.Operation) (bool, error) + ValidateRoleModuleAccess(ctx context.Context, accountID string, role roles.RolePermissions, module modules.Module, operation operations.Operation) bool + ValidateAccountAccess(ctx context.Context, accountID string, user *types.User, allowOwnerAndAdmin bool) error +} + +type userManager interface { + GetUserByUserID(ctx context.Context, lockingStrength store.LockingStrength, userID string) (*types.User, error) +} + +type managerImpl struct { + userManager userManager +} + +func NewManager(userManager userManager) Manager { + return &managerImpl{ + userManager: userManager, + } +} + +func (m *managerImpl) ValidateUserPermissions( + ctx context.Context, + accountID string, + userID string, + module modules.Module, + operation operations.Operation, +) (bool, error) { + if userID == activity.SystemInitiator { + return true, nil + } + + user, err := m.userManager.GetUserByUserID(ctx, store.LockingStrengthShare, userID) + if err != nil { + return false, err + } + + if user == nil { + return false, status.NewUserNotFoundError(userID) + } + + if user.IsBlocked() { + return false, status.NewUserBlockedError() + } + + if err := m.ValidateAccountAccess(ctx, accountID, user, false); err != nil { + return false, err + } + + if operation == operations.Read && user.IsServiceUser { + return true, nil // this should be replaced by proper granular access role + } + + role, ok := roles.RolesMap[user.Role] + if !ok { + return false, status.NewUserRoleNotFoundError(string(user.Role)) + } + + return m.ValidateRoleModuleAccess(ctx, accountID, role, module, operation), nil +} + +func (m *managerImpl) ValidateRoleModuleAccess( + ctx context.Context, + accountID string, + role roles.RolePermissions, + module modules.Module, + operation operations.Operation, +) bool { + if permissions, ok := role.Permissions[module]; ok { + if allowed, exists := permissions[operation]; exists { + return allowed + } + log.WithContext(ctx).Tracef("operation %s not found on module %s for role %s", operation, module, role.Role) + return false + } + + return role.AutoAllowNew[operation] +} + +func (m *managerImpl) ValidateAccountAccess(ctx context.Context, accountID string, user *types.User, allowOwnerAndAdmin bool) error { + if user.AccountID != accountID { + return status.NewUserNotPartOfAccountError() + } + return nil +} diff --git a/internal/shared/permissions/manager_mock.go b/internal/shared/permissions/manager_mock.go new file mode 100644 index 0000000..266a242 --- /dev/null +++ b/internal/shared/permissions/manager_mock.go @@ -0,0 +1,82 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: ./manager.go + +// Package permissions is a generated GoMock package. +package permissions + +import ( + context "context" + reflect "reflect" + + gomock "github.com/golang/mock/gomock" + modules "github.com/netbirdio/netbird/management/server/permissions/modules" + operations "github.com/netbirdio/netbird/management/server/permissions/operations" + roles "github.com/netbirdio/netbird/management/server/permissions/roles" + types "github.com/netbirdio/netbird/management/server/types" +) + +// MockManager is a mock of Manager interface. +type MockManager struct { + ctrl *gomock.Controller + recorder *MockManagerMockRecorder +} + +// MockManagerMockRecorder is the mock recorder for MockManager. +type MockManagerMockRecorder struct { + mock *MockManager +} + +// NewMockManager creates a new mock instance. +func NewMockManager(ctrl *gomock.Controller) *MockManager { + mock := &MockManager{ctrl: ctrl} + mock.recorder = &MockManagerMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockManager) EXPECT() *MockManagerMockRecorder { + return m.recorder +} + +// ValidateAccountAccess mocks base method. +func (m *MockManager) ValidateAccountAccess(ctx context.Context, accountID string, user *types.User, allowOwnerAndAdmin bool) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "ValidateAccountAccess", ctx, accountID, user, allowOwnerAndAdmin) + ret0, _ := ret[0].(error) + return ret0 +} + +// ValidateAccountAccess indicates an expected call of ValidateAccountAccess. +func (mr *MockManagerMockRecorder) ValidateAccountAccess(ctx, accountID, user, allowOwnerAndAdmin interface{}) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ValidateAccountAccess", reflect.TypeOf((*MockManager)(nil).ValidateAccountAccess), ctx, accountID, user, allowOwnerAndAdmin) +} + +// ValidateRoleModuleAccess mocks base method. +func (m *MockManager) ValidateRoleModuleAccess(ctx context.Context, accountID string, role roles.RolePermissions, module modules.Module, operation operations.Operation) bool { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "ValidateRoleModuleAccess", ctx, accountID, role, module, operation) + ret0, _ := ret[0].(bool) + return ret0 +} + +// ValidateRoleModuleAccess indicates an expected call of ValidateRoleModuleAccess. +func (mr *MockManagerMockRecorder) ValidateRoleModuleAccess(ctx, accountID, role, module, operation interface{}) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ValidateRoleModuleAccess", reflect.TypeOf((*MockManager)(nil).ValidateRoleModuleAccess), ctx, accountID, role, module, operation) +} + +// ValidateUserPermissions mocks base method. +func (m *MockManager) ValidateUserPermissions(ctx context.Context, accountID, userID string, module modules.Module, operation operations.Operation) (bool, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "ValidateUserPermissions", ctx, accountID, userID, module, operation) + ret0, _ := ret[0].(bool) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// ValidateUserPermissions indicates an expected call of ValidateUserPermissions. +func (mr *MockManagerMockRecorder) ValidateUserPermissions(ctx, accountID, userID, module, operation interface{}) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ValidateUserPermissions", reflect.TypeOf((*MockManager)(nil).ValidateUserPermissions), ctx, accountID, userID, module, operation) +} diff --git a/internal/shared/permissions/modules/module.go b/internal/shared/permissions/modules/module.go new file mode 100644 index 0000000..4c42b61 --- /dev/null +++ b/internal/shared/permissions/modules/module.go @@ -0,0 +1,19 @@ +package modules + +type Module string + +const ( + Networks Module = "networks" + Peers Module = "peers" + Groups Module = "groups" + Settings Module = "settings" + Accounts Module = "accounts" + Dns Module = "dns" + Nameservers Module = "nameservers" + Events Module = "events" + Policies Module = "policies" + Routes Module = "routes" + Users Module = "users" + SetupKeys Module = "setup_keys" + Pats Module = "pats" +) diff --git a/internal/shared/permissions/operations/operation.go b/internal/shared/permissions/operations/operation.go new file mode 100644 index 0000000..af709de --- /dev/null +++ b/internal/shared/permissions/operations/operation.go @@ -0,0 +1,8 @@ +package operations + +type Operation string + +const ( + Read Operation = "read" + Write Operation = "write" +) diff --git a/internal/shared/permissions/roles/admin.go b/internal/shared/permissions/roles/admin.go new file mode 100644 index 0000000..a826d18 --- /dev/null +++ b/internal/shared/permissions/roles/admin.go @@ -0,0 +1,21 @@ +package roles + +import ( + "github.com/netbirdio/netbird/management/server/permissions/modules" + "github.com/netbirdio/netbird/management/server/permissions/operations" + "github.com/netbirdio/netbird/management/server/types" +) + +var Admin = RolePermissions{ + Role: types.UserRoleAdmin, + AutoAllowNew: map[operations.Operation]bool{ + operations.Read: true, + operations.Write: true, + }, + Permissions: Permissions{ + modules.Accounts: { + operations.Read: true, + operations.Write: false, + }, + }, +} diff --git a/internal/shared/permissions/roles/owner.go b/internal/shared/permissions/roles/owner.go new file mode 100644 index 0000000..f739d18 --- /dev/null +++ b/internal/shared/permissions/roles/owner.go @@ -0,0 +1,14 @@ +package roles + +import ( + "github.com/netbirdio/netbird/management/server/permissions/operations" + "github.com/netbirdio/netbird/management/server/types" +) + +var Owner = RolePermissions{ + Role: types.UserRoleOwner, + AutoAllowNew: map[operations.Operation]bool{ + operations.Read: true, + operations.Write: true, + }, +} diff --git a/internal/shared/permissions/roles/role_permissions.go b/internal/shared/permissions/roles/role_permissions.go new file mode 100644 index 0000000..9995740 --- /dev/null +++ b/internal/shared/permissions/roles/role_permissions.go @@ -0,0 +1,21 @@ +package roles + +import ( + "management/internal/modules/users/types" + "management/internal/shared/permissions/modules" + "management/internal/shared/permissions/operations" +) + +type RolePermissions struct { + Role types.UserRole + Permissions Permissions + AutoAllowNew map[operations.Operation]bool +} + +type Permissions map[modules.Module]map[operations.Operation]bool + +var RolesMap = map[types.UserRole]RolePermissions{ + types.UserRoleOwner: Owner, + types.UserRoleAdmin: Admin, + types.UserRoleUser: User, +} diff --git a/internal/shared/permissions/roles/user.go b/internal/shared/permissions/roles/user.go new file mode 100644 index 0000000..8796b7d --- /dev/null +++ b/internal/shared/permissions/roles/user.go @@ -0,0 +1,15 @@ +package roles + +import ( + "github.com/netbirdio/netbird/management/server/types" + + "management/internal/shared/permissions/operations" +) + +var User = RolePermissions{ + Role: types.UserRoleUser, + AutoAllowNew: map[operations.Operation]bool{ + operations.Read: false, + operations.Write: false, + }, +} diff --git a/pkg/configuration/config.go b/pkg/configuration/config.go new file mode 100644 index 0000000..da17201 --- /dev/null +++ b/pkg/configuration/config.go @@ -0,0 +1,19 @@ +package configuration + +import ( + "errors" + + "github.com/caarlos0/env/v11" +) + +var ( + ErrFailedToParseConfig = errors.New("failed to parse config from env") +) + +func Parse[T any]() (*T, error) { + var cfg T + if err := env.Parse(&cfg); err != nil { + return &cfg, ErrFailedToParseConfig + } + return &cfg, nil +} diff --git a/pkg/logging/config.go b/pkg/logging/config.go new file mode 100644 index 0000000..256a7fc --- /dev/null +++ b/pkg/logging/config.go @@ -0,0 +1,5 @@ +package logging + +type LoggingConfig struct { + LogLevels map[string]string `mapstructure:"log_levels"` +} diff --git a/pkg/logging/hook/additional_empty.go b/pkg/logging/hook/additional_empty.go new file mode 100644 index 0000000..4f50694 --- /dev/null +++ b/pkg/logging/hook/additional_empty.go @@ -0,0 +1,9 @@ +//go:build !loggoroutine + +package hook + +import log "github.com/sirupsen/logrus" + +func additionalEntries(_ *log.Entry) { + // This function is empty and is used to demonstrate the use of additional hooks. +} diff --git a/pkg/logging/hook/additional_goroutine.go b/pkg/logging/hook/additional_goroutine.go new file mode 100644 index 0000000..fb4e09f --- /dev/null +++ b/pkg/logging/hook/additional_goroutine.go @@ -0,0 +1,12 @@ +//go:build loggoroutine + +package hook + +import ( + "github.com/petermattis/goid" + log "github.com/sirupsen/logrus" +) + +func additionalEntries(entry *log.Entry) { + entry.Data[EntryKeyGoroutineID] = goid.Get() +} diff --git a/pkg/logging/hook/hook.go b/pkg/logging/hook/hook.go new file mode 100644 index 0000000..290c337 --- /dev/null +++ b/pkg/logging/hook/hook.go @@ -0,0 +1,139 @@ +package hook + +import ( + "fmt" + "path" + "runtime" + "runtime/debug" + "strings" + + "github.com/sirupsen/logrus" + + "github.com/netbirdio/netbird/management/server/context" +) + +type ExecutionContext string + +const ( + ExecutionContextKey = "executionContext" + + HTTPSource ExecutionContext = "HTTP" + GRPCSource ExecutionContext = "GRPC" + SystemSource ExecutionContext = "SYSTEM" +) + +// ContextHook is a custom hook for add the source information for the entry +type ContextHook struct { + goModuleName string +} + +// NewContextHook instantiate a new context hook +func NewContextHook() *ContextHook { + hook := &ContextHook{} + hook.goModuleName = hook.moduleName() + "/" + return hook +} + +// Levels set the supported levels for this hook +func (hook ContextHook) Levels() []logrus.Level { + return logrus.AllLevels +} + +// Fire extend with the source information the entry.Data +func (hook ContextHook) Fire(entry *logrus.Entry) error { + caller := &runtime.Frame{Line: 0, File: "caller_not_available"} + if entry.Caller != nil { + caller = entry.Caller + } + src := hook.parseSrc(caller.File) + entry.Data[EntryKeySource] = fmt.Sprintf("%s:%v", src, caller.Line) + additionalEntries(entry) + + if entry.Context == nil { + return nil + } + + source, ok := entry.Context.Value(ExecutionContextKey).(ExecutionContext) + if !ok { + return nil + } + + entry.Data["context"] = source + + switch source { + case HTTPSource: + addHTTPFields(entry) + case GRPCSource: + addGRPCFields(entry) + case SystemSource: + addSystemFields(entry) + } + + return nil +} + +func (hook ContextHook) moduleName() string { + info, ok := debug.ReadBuildInfo() + if ok && info.Main.Path != "" { + return info.Main.Path + } + + return "netbird" +} + +func (hook ContextHook) parseSrc(filePath string) string { + netbirdPath := strings.SplitAfter(filePath, hook.goModuleName) + if len(netbirdPath) > 1 { + return netbirdPath[len(netbirdPath)-1] + } + + // in case of forked repo + netbirdPath = strings.SplitAfter(filePath, "netbird/") + if len(netbirdPath) > 1 { + return netbirdPath[len(netbirdPath)-1] + } + + // in case if log entry is come from external pkg + _, pkg := path.Split(path.Dir(filePath)) + file := path.Base(filePath) + return fmt.Sprintf("%s/%s", pkg, file) +} + +func addHTTPFields(entry *logrus.Entry) { + if ctxReqID, ok := entry.Context.Value(context.RequestIDKey).(string); ok { + entry.Data[context.RequestIDKey] = ctxReqID + } + if ctxAccountID, ok := entry.Context.Value(context.AccountIDKey).(string); ok { + entry.Data[context.AccountIDKey] = ctxAccountID + } + if ctxInitiatorID, ok := entry.Context.Value(context.UserIDKey).(string); ok { + entry.Data[context.UserIDKey] = ctxInitiatorID + } +} + +func addGRPCFields(entry *logrus.Entry) { + if ctxReqID, ok := entry.Context.Value(context.RequestIDKey).(string); ok { + entry.Data[context.RequestIDKey] = ctxReqID + } + if ctxAccountID, ok := entry.Context.Value(context.AccountIDKey).(string); ok { + entry.Data[context.AccountIDKey] = ctxAccountID + } + if ctxDeviceID, ok := entry.Context.Value(context.PeerIDKey).(string); ok { + entry.Data[context.PeerIDKey] = ctxDeviceID + } +} + +func addSystemFields(entry *logrus.Entry) { + if ctxReqID, ok := entry.Context.Value(context.RequestIDKey).(string); ok { + entry.Data[context.RequestIDKey] = ctxReqID + } + if ctxInitiatorID, ok := entry.Context.Value(context.UserIDKey).(string); ok { + entry.Data[context.UserIDKey] = ctxInitiatorID + } + if ctxAccountID, ok := entry.Context.Value(context.AccountIDKey).(string); ok { + entry.Data[context.AccountIDKey] = ctxAccountID + } + if ctxDeviceID, ok := entry.Context.Value(context.PeerIDKey).(string); ok { + entry.Data[context.PeerIDKey] = ctxDeviceID + } +} diff --git a/pkg/logging/hook/hook_test.go b/pkg/logging/hook/hook_test.go new file mode 100644 index 0000000..8021632 --- /dev/null +++ b/pkg/logging/hook/hook_test.go @@ -0,0 +1,39 @@ +package hook + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestFilePathParsing(t *testing.T) { + + testCases := []struct { + filePath string + expectedFileName string + }{ + // locally cloned repo + { + filePath: "/Users/user/Github/Netbird/netbird/formatter/formatter.go", + expectedFileName: "formatter/formatter.go", + }, + // locally cloned repo with duplicated name in path + { + filePath: "/Users/user/netbird/repos/netbird/formatter/formatter.go", + expectedFileName: "formatter/formatter.go", + }, + // locally cloned repo with renamed package root + { + filePath: "/Users/user/Github/MyOwnNetbirdClient/formatter/formatter.go", + expectedFileName: "formatter/formatter.go", + }, + } + + hook := NewContextHook() + + for _, testCase := range testCases { + parsedString := hook.parseSrc(testCase.filePath) + assert.Equal(t, testCase.expectedFileName, parsedString, "Parsed filepath does not match expected for %s", testCase.filePath) + } + +} diff --git a/pkg/logging/hook/keys.go b/pkg/logging/hook/keys.go new file mode 100644 index 0000000..09781a8 --- /dev/null +++ b/pkg/logging/hook/keys.go @@ -0,0 +1,6 @@ +package hook + +const ( + EntryKeySource = "source" + EntryKeyGoroutineID = "goroutine_id" +) diff --git a/pkg/logging/init.go b/pkg/logging/init.go new file mode 100644 index 0000000..a0f190f --- /dev/null +++ b/pkg/logging/init.go @@ -0,0 +1,210 @@ +package logging + +import ( + "fmt" + "io" + "log" + "os" + "path/filepath" + "runtime" + "slices" + "strconv" + "strings" + "sync" + + "github.com/sirupsen/logrus" + "github.com/spf13/viper" +) + +// global map of package paths to *logrus.Logger +var ( + mu sync.RWMutex + loggers = map[string]*logrus.Logger{} +) + +// Init reads the logging config from a YAML file and sets up per-package loggers. +func Init(configFilePath string) error { + v := viper.New() + v.SetConfigFile(configFilePath) + + // Read the YAML + if err := v.ReadInConfig(); err != nil { + return fmt.Errorf("failed to read logging config: %w", err) + } + + var cfg LoggingConfig + if err := v.Unmarshal(&cfg); err != nil { + return fmt.Errorf("failed to unmarshal logging config: %w", err) + } + + mu.Lock() + defer mu.Unlock() + + // For each package in our config, create a logger at the given level + for pkgPath, levelStr := range cfg.LogLevels { + l := logrus.New() // each package gets its own *logrus.Logger + l.SetLevel(parseLogrusLevel(levelStr)) + + // Optionally, set the formatter, output, etc.: + // l.SetFormatter(&logrus.JSONFormatter{}) + // l.SetOutput(os.Stdout) + + loggers[pkgPath] = l + } + + // Optionally, define a default logger for packages not explicitly listed + if _, ok := loggers["default"]; !ok { + defaultLogger := logrus.New() + defaultLogger.SetLevel(logrus.InfoLevel) + loggers["default"] = defaultLogger + } + + return nil +} + +// parseLogrusLevel is a helper that converts a string (e.g. "debug") to a logrus.Level. +func parseLogrusLevel(levelStr string) logrus.Level { + switch strings.ToLower(levelStr) { + case "trace": + return logrus.TraceLevel + case "debug": + return logrus.DebugLevel + case "info": + return logrus.InfoLevel + case "warn", "warning": + return logrus.WarnLevel + case "error": + return logrus.ErrorLevel + case "fatal": + return logrus.FatalLevel + case "panic": + return logrus.PanicLevel + } + // default + return logrus.InfoLevel +} + +// LoggerFor returns a *logrus.Logger for the specified package path. +// If there's no explicit logger, we return the "default" logger. +func LoggerFor(pkgPath string) *logrus.Logger { + mu.RLock() + defer mu.RUnlock() + + if l, ok := loggers[pkgPath]; ok { + return l + } + if l, ok := loggers["default"]; ok { + return l + } + + logrus.Printf("No logger configured for %q; using fallback (info-level) logger.\n", pkgPath) + fallback := logrus.New() + fallback.SetLevel(logrus.InfoLevel) + return fallback +} + +func LoggerForThisPackage() *logrus.Logger { + pc, _, _, ok := runtime.Caller(1) + if !ok { + return LoggerFor("default") + } + fn := runtime.FuncForPC(pc) + if fn == nil { + return LoggerFor("default") + } + + fullFuncName := fn.Name() + pkgPath := parsePackageFromFuncName(fullFuncName) + + return LoggerFor(pkgPath) +} + +func parsePackageFromFuncName(funcName string) string { + parts := strings.Split(funcName, "/") + if len(parts) == 0 { + return "default" + } + last := parts[len(parts)-1] + + base := strings.Join(parts[:len(parts)-1], "/") + + dotIdx := strings.IndexByte(last, '.') + var pkgName string + if dotIdx == -1 { + pkgName = last + } else { + pkgName = last[:dotIdx] + } + + return base + "/" + pkgName +} + +// InitLog parses and sets log-level input +func InitLog(logLevel string, logPath string) error { + level, err := logrus.ParseLevel(logLevel) + if err != nil { + logrus.Errorf("Failed parsing log-level %s: %s", logLevel, err) + return err + } + customOutputs := []string{"console", "syslog"} + + if logPath != "" && !slices.Contains(customOutputs, logPath) { + maxLogSize := getLogMaxSize() + lumberjackLogger := &lumberjack.Logger{ + // Log file absolute path, os agnostic + Filename: filepath.ToSlash(logPath), + MaxSize: maxLogSize, // MB + MaxBackups: 10, + MaxAge: 30, // days + Compress: true, + } + log.SetOutput(io.Writer(lumberjackLogger)) + } else if logPath == "syslog" { + AddSyslogHook() + } + + //nolint:gocritic + if os.Getenv("NB_LOG_FORMAT") == "json" { + SetJSONFormatter(logrus.StandardLogger()) + } else if logPath == "syslog" { + SetSyslogFormatter(logrus.StandardLogger()) + } else { + SetTextFormatter(logrus.StandardLogger()) + } + logrus.SetLevel(level) + + setGRPCLibLogger() + + return nil +} + +func setGRPCLibLogger() { + logOut := logrus.StandardLogger().Writer() + if os.Getenv("GRPC_GO_LOG_SEVERITY_LEVEL") != "info" { + grpclog.SetLoggerV2(grpclog.NewLoggerV2(io.Discard, logOut, logOut)) + return + } + + var v int + vLevel := os.Getenv("GRPC_GO_LOG_VERBOSITY_LEVEL") + if vl, err := strconv.Atoi(vLevel); err == nil { + v = vl + } + + grpclog.SetLoggerV2(grpclog.NewLoggerV2WithVerbosity(logOut, logOut, logOut, v)) +} + +func getLogMaxSize() int { + if sizeVar, ok := os.LookupEnv("NB_LOG_MAX_SIZE_MB"); ok { + size, err := strconv.ParseInt(sizeVar, 10, 64) + if err != nil { + log.Errorf("Failed parsing log-size %s: %s. Should be just an integer", sizeVar, err) + return defaultLogSize + } + + log.Infof("Setting log file max size to %d MB", size) + + return int(size) + } + return defaultLogSize +} diff --git a/pkg/logging/logcat/logcat.go b/pkg/logging/logcat/logcat.go new file mode 100644 index 0000000..c561d32 --- /dev/null +++ b/pkg/logging/logcat/logcat.go @@ -0,0 +1,50 @@ +package logcat + +import ( + "fmt" + "strings" + + "github.com/sirupsen/logrus" + + "github.com/netbirdio/netbird/formatter/levels" +) + +// Formatter formats logs into text what is fit for logcat +type Formatter struct { + levelDesc []string +} + +// NewLogcatFormatter create new LogcatFormatter instance +func NewLogcatFormatter() *Formatter { + return &Formatter{ + levelDesc: levels.ValidLevelDesc, + } +} + +// Format renders a single log entry +func (f *Formatter) Format(entry *logrus.Entry) ([]byte, error) { + var fields string + keys := make([]string, 0, len(entry.Data)) + for k, v := range entry.Data { + if k == "source" { + continue + } + keys = append(keys, fmt.Sprintf("%s: %v", k, v)) + } + + if len(keys) > 0 { + fields = fmt.Sprintf("[%s] ", strings.Join(keys, ", ")) + } + + level := f.parseLevel(entry.Level) + + return []byte(fmt.Sprintf("[%s] %s%s %s\n", level, fields, entry.Data["source"], entry.Message)), nil +} + +func (f *Formatter) parseLevel(level logrus.Level) string { + if len(f.levelDesc) < int(level) { + return "" + } + + return f.levelDesc[level] +} diff --git a/pkg/logging/logcat/logcat_test.go b/pkg/logging/logcat/logcat_test.go new file mode 100644 index 0000000..fd4d928 --- /dev/null +++ b/pkg/logging/logcat/logcat_test.go @@ -0,0 +1,29 @@ +package logcat + +import ( + "testing" + "time" + + "github.com/sirupsen/logrus" +) + +func TestLogcatMessageFormat(t *testing.T) { + + someEntry := &logrus.Entry{ + Data: logrus.Fields{"att1": 1, "att2": 2, "source": "some/fancy/path.go:46"}, + Time: time.Date(2021, time.Month(2), 21, 1, 10, 30, 0, time.UTC), + Level: 3, + Message: "Some Message", + } + + formatter := NewLogcatFormatter() + result, _ := formatter.Format(someEntry) + + expectedString := "[WARN] [att1: 1, att2: 2] some/fancy/path.go:46 Some Message\n" + expectedStringVariant := "[WARN] [att2: 2, att1: 1] some/fancy/path.go:46 Some Message\n" + parsedString := string(result) + if parsedString != expectedString && parsedString != expectedStringVariant { + t.Errorf("The log messages don't match. Expected: '%s', got: '%s'", expectedString, parsedString) + } + +} diff --git a/pkg/logging/logging.yaml b/pkg/logging/logging.yaml new file mode 100644 index 0000000..7be53f4 --- /dev/null +++ b/pkg/logging/logging.yaml @@ -0,0 +1,2 @@ +log_levels: + management-refactor/internal/shared/db: debug \ No newline at end of file diff --git a/pkg/logging/set.go b/pkg/logging/set.go new file mode 100644 index 0000000..5c38ee3 --- /dev/null +++ b/pkg/logging/set.go @@ -0,0 +1,38 @@ +package logging + +import ( + "github.com/sirupsen/logrus" + + "management/pkg/logging/hook" + "management/pkg/logging/logcat" + "management/pkg/logging/syslog" + "management/pkg/logging/txt" +) + +// SetTextFormatter set the text formatter for given logger. +func SetTextFormatter(logger *logrus.Logger) { + logger.Formatter = txt.NewTextFormatter() + logger.ReportCaller = true + logger.AddHook(hook.NewContextHook()) +} + +// SetSyslogFormatter set the text formatter for given logger. +func SetSyslogFormatter(logger *logrus.Logger) { + logger.Formatter = syslog.NewSyslogFormatter() + logger.ReportCaller = true + logger.AddHook(hook.NewContextHook()) +} + +// SetJSONFormatter set the JSON formatter for given logger. +func SetJSONFormatter(logger *logrus.Logger) { + logger.Formatter = &logrus.JSONFormatter{} + logger.ReportCaller = true + logger.AddHook(hook.NewContextHook()) +} + +// SetLogcatFormatter set the logcat formatter for given logger. +func SetLogcatFormatter(logger *logrus.Logger) { + logger.Formatter = logcat.NewLogcatFormatter() + logger.ReportCaller = true + logger.AddHook(hook.NewContextHook()) +} diff --git a/pkg/logging/syslog/formatter.go b/pkg/logging/syslog/formatter.go new file mode 100644 index 0000000..e72c303 --- /dev/null +++ b/pkg/logging/syslog/formatter.go @@ -0,0 +1,39 @@ +package syslog + +import ( + "fmt" + "strings" + + "github.com/sirupsen/logrus" + + "github.com/netbirdio/netbird/formatter/levels" +) + +// Formatter formats logs into text +type Formatter struct { + levelDesc []string +} + +// NewSyslogFormatter create new MySyslogFormatter instance +func NewSyslogFormatter() *Formatter { + return &Formatter{ + levelDesc: levels.ValidLevelDesc, + } +} + +// Format renders a single log entry +func (f *Formatter) Format(entry *logrus.Entry) ([]byte, error) { + var fields string + keys := make([]string, 0, len(entry.Data)) + for k, v := range entry.Data { + if k == "source" { + continue + } + keys = append(keys, fmt.Sprintf("%s: %v", k, v)) + } + + if len(keys) > 0 { + fields = fmt.Sprintf("[%s] ", strings.Join(keys, ", ")) + } + return []byte(fmt.Sprintf("%s%s\n", fields, entry.Message)), nil +} diff --git a/pkg/logging/syslog/formatter_test.go b/pkg/logging/syslog/formatter_test.go new file mode 100644 index 0000000..110a339 --- /dev/null +++ b/pkg/logging/syslog/formatter_test.go @@ -0,0 +1,26 @@ +package syslog + +import ( + "testing" + "time" + + "github.com/sirupsen/logrus" + "github.com/stretchr/testify/assert" +) + +func TestLogSyslogFormat(t *testing.T) { + + someEntry := &logrus.Entry{ + Data: logrus.Fields{"att1": 1, "att2": 2, "source": "some/fancy/path.go:46"}, + Time: time.Date(2021, time.Month(2), 21, 1, 10, 30, 0, time.UTC), + Level: 3, + Message: "Some Message", + } + + formatter := NewSyslogFormatter() + result, _ := formatter.Format(someEntry) + + parsedString := string(result) + expectedString := "^\\[(att1: 1, att2: 2|att2: 2, att1: 1)\\] Some Message\\s+$" + assert.Regexp(t, expectedString, parsedString) +} diff --git a/pkg/logging/txt/format.go b/pkg/logging/txt/format.go new file mode 100644 index 0000000..a88c410 --- /dev/null +++ b/pkg/logging/txt/format.go @@ -0,0 +1,31 @@ +//go:build !loggoroutine + +package txt + +import ( + "fmt" + "strings" + + "github.com/sirupsen/logrus" + + "github.com/netbirdio/netbird/formatter/hook" +) + +func (f *TextFormatter) Format(entry *logrus.Entry) ([]byte, error) { + var fields string + keys := make([]string, 0, len(entry.Data)) + for k, v := range entry.Data { + if k == hook.EntryKeySource { + continue + } + keys = append(keys, fmt.Sprintf("%s: %v", k, v)) + } + + if len(keys) > 0 { + fields = fmt.Sprintf("[%s] ", strings.Join(keys, ", ")) + } + + level := f.parseLevel(entry.Level) + + return []byte(fmt.Sprintf("%s %s %s%s: %s\n", entry.Time.Format(f.timestampFormat), level, fields, entry.Data[hook.EntryKeySource], entry.Message)), nil +} diff --git a/pkg/logging/txt/format_gorutines.go b/pkg/logging/txt/format_gorutines.go new file mode 100644 index 0000000..a39aee6 --- /dev/null +++ b/pkg/logging/txt/format_gorutines.go @@ -0,0 +1,35 @@ +//go:build loggoroutine + +package txt + +import ( + "fmt" + "strings" + + "github.com/sirupsen/logrus" + + "github.com/netbirdio/netbird/formatter/hook" +) + +func (f *TextFormatter) Format(entry *logrus.Entry) ([]byte, error) { + var fields string + keys := make([]string, 0, len(entry.Data)) + for k, v := range entry.Data { + if k == hook.EntryKeySource { + continue + } + + if k == hook.EntryKeyGoroutineID { + continue + } + keys = append(keys, fmt.Sprintf("%s: %v", k, v)) + } + + if len(keys) > 0 { + fields = fmt.Sprintf("[%s] ", strings.Join(keys, ", ")) + } + + level := f.parseLevel(entry.Level) + + return []byte(fmt.Sprintf("%s %s %d %s%s: %s\n", entry.Time.Format(f.timestampFormat), level, entry.Data[hook.EntryKeyGoroutineID], fields, entry.Data[hook.EntryKeySource], entry.Message)), nil +} diff --git a/pkg/logging/txt/formatter.go b/pkg/logging/txt/formatter.go new file mode 100644 index 0000000..3b2a3fb --- /dev/null +++ b/pkg/logging/txt/formatter.go @@ -0,0 +1,31 @@ +package txt + +import ( + "time" + + "github.com/sirupsen/logrus" + + "github.com/netbirdio/netbird/formatter/levels" +) + +// TextFormatter formats logs into text with included source code's path +type TextFormatter struct { + timestampFormat string + levelDesc []string +} + +// NewTextFormatter create new MyTextFormatter instance +func NewTextFormatter() *TextFormatter { + return &TextFormatter{ + levelDesc: levels.ValidLevelDesc, + timestampFormat: time.RFC3339, // or RFC3339 + } +} + +func (f *TextFormatter) parseLevel(level logrus.Level) string { + if len(f.levelDesc) < int(level) { + return "" + } + + return f.levelDesc[level] +} diff --git a/pkg/logging/txt/formatter_test.go b/pkg/logging/txt/formatter_test.go new file mode 100644 index 0000000..590af5d --- /dev/null +++ b/pkg/logging/txt/formatter_test.go @@ -0,0 +1,26 @@ +package txt + +import ( + "testing" + "time" + + "github.com/sirupsen/logrus" + "github.com/stretchr/testify/assert" +) + +func TestLogTextFormat(t *testing.T) { + + someEntry := &logrus.Entry{ + Data: logrus.Fields{"att1": 1, "att2": 2, "source": "some/fancy/path.go:46"}, + Time: time.Date(2021, time.Month(2), 21, 1, 10, 30, 0, time.UTC), + Level: 3, + Message: "Some Message", + } + + formatter := NewTextFormatter() + result, _ := formatter.Format(someEntry) + + parsedString := string(result) + expectedString := "^2021-02-21T01:10:30Z WARN \\[(att1: 1, att2: 2|att2: 2, att1: 1)\\] some/fancy/path.go:46: Some Message\\s+$" + assert.Regexp(t, expectedString, parsedString) +}