diff --git a/go.mod b/go.mod index 05ef098..17ccbab 100644 --- a/go.mod +++ b/go.mod @@ -3,6 +3,7 @@ go 1.16 require ( + github.com/DATA-DOG/go-sqlmock v1.5.0 github.com/kr/text v0.2.0 // indirect github.com/niemeyer/pretty v0.0.0-20200227124842-a10e7caefd8e // indirect github.com/pkg/errors v0.9.1 @@ -13,5 +14,6 @@ golang.zx2c4.com/wireguard/wgctrl v0.0.0-20210506160403-92e472f520a5 gopkg.in/check.v1 v1.0.0-20200227125254-8fa46927fb4f // indirect gopkg.in/yaml.v3 v3.0.0-20210107192922-496545a6307b // indirect + gorm.io/driver/mysql v1.1.2 gorm.io/gorm v1.21.14 ) diff --git a/internal/lowlevel/wrappers.go b/internal/lowlevel/wrappers.go new file mode 100644 index 0000000..0aa57ab --- /dev/null +++ b/internal/lowlevel/wrappers.go @@ -0,0 +1,55 @@ +package lowlevel + +import ( + "io" + + "github.com/vishvananda/netlink" + "golang.zx2c4.com/wireguard/wgctrl/wgtypes" +) + +// A WireGuardClient is a type which can control a WireGuard device. +type WireGuardClient interface { + io.Closer + Devices() ([]*wgtypes.Device, error) + Device(name string) (*wgtypes.Device, error) + ConfigureDevice(name string, cfg wgtypes.Config) error +} + +// A NetlinkClient is a type which can control a netlink device. +type NetlinkClient interface { + LinkAdd(link netlink.Link) error + LinkDel(link netlink.Link) error + LinkByName(name string) (netlink.Link, error) + LinkSetUp(link netlink.Link) error + LinkSetDown(link netlink.Link) error + LinkSetMTU(link netlink.Link, mtu int) error + AddrReplace(link netlink.Link, addr *netlink.Addr) error + AddrAdd(link netlink.Link, addr *netlink.Addr) error +} + +type NetlinkManager struct { +} + +func (n NetlinkManager) LinkAdd(link netlink.Link) error { return netlink.LinkAdd(link) } + +func (n NetlinkManager) LinkDel(link netlink.Link) error { return netlink.LinkDel(link) } + +func (n NetlinkManager) LinkByName(name string) (netlink.Link, error) { + return netlink.LinkByName(name) +} + +func (n NetlinkManager) LinkSetUp(link netlink.Link) error { return netlink.LinkSetUp(link) } + +func (n NetlinkManager) LinkSetDown(link netlink.Link) error { return netlink.LinkSetDown(link) } + +func (n NetlinkManager) LinkSetMTU(link netlink.Link, mtu int) error { + return netlink.LinkSetMTU(link, mtu) +} + +func (n NetlinkManager) AddrReplace(link netlink.Link, addr *netlink.Addr) error { + return netlink.AddrReplace(link, addr) +} + +func (n NetlinkManager) AddrAdd(link netlink.Link, addr *netlink.Addr) error { + return netlink.AddrAdd(link, addr) +} diff --git a/internal/wireguard/backend_db.go b/internal/wireguard/backend_db.go new file mode 100644 index 0000000..414b9e9 --- /dev/null +++ b/internal/wireguard/backend_db.go @@ -0,0 +1,396 @@ +package wireguard + +import ( + "database/sql" + "time" + + "gorm.io/gorm/clause" + + "github.com/pkg/errors" + "gorm.io/gorm" +) + +type DatabaseBackend struct { + db *gorm.DB +} + +func NewDatabaseBackend(db *gorm.DB) (*DatabaseBackend, error) { + backend := &DatabaseBackend{db: db} + + // Auto-Migrate Gorm models + err := db.AutoMigrate(&dbInterfaceConfig{}, &dbDefaultPeerConfig{}, &dbPeerConfig{}) + if err != nil { + return nil, errors.Wrap(err, "failed to migrate WireGuard database") + } + + return backend, nil +} + +func (d DatabaseBackend) SaveInterface(cfg InterfaceConfig, _ []PeerConfig) error { + iface, peerDefaults := convertInterface(cfg) + + if err := d.db.Clauses(clause.OnConflict{UpdateAll: true}).Create(&iface).Error; err != nil { + return errors.Wrapf(err, "failed to save interface %s to db", cfg.DeviceName) + } + if err := d.db.Clauses(clause.OnConflict{UpdateAll: true}).Create(&peerDefaults).Error; err != nil { + return errors.Wrapf(err, "failed to save peer defaults of %s to db", cfg.DeviceName) + } + + return nil +} + +func (d DatabaseBackend) SavePeer(cfg PeerConfig, iface InterfaceConfig) error { + peer := convertPeer(cfg, iface.DeviceName) + + if err := d.db.Clauses(clause.OnConflict{UpdateAll: true}).Create(&peer).Error; err != nil { + return errors.Wrapf(err, "failed to save peer %s to db", cfg.Uid) + } + + return nil +} + +func (d DatabaseBackend) DeleteInterface(cfg InterfaceConfig, _ []PeerConfig) error { + // Delete peers + if err := d.db.Where("device_name = ?", cfg.DeviceName).Delete(&dbPeerConfig{}).Error; err != nil { + return errors.Wrapf(err, "failed to delete peer for %s from db", cfg.DeviceName) + } + // Delete peer default settings + if err := d.db.Where("device_name = ?", cfg.DeviceName).Delete(&dbDefaultPeerConfig{}).Error; err != nil { + return errors.Wrapf(err, "failed to delete peer defaults for %s from db", cfg.DeviceName) + } + // Delete interface config + if err := d.db.Where("device_name = ?", cfg.DeviceName).Delete(&dbInterfaceConfig{}).Error; err != nil { + return errors.Wrapf(err, "failed to delete interface %s from db", cfg.DeviceName) + } + return nil +} + +func (d DatabaseBackend) DeletePeer(cfg PeerConfig, iface InterfaceConfig) error { + err := d.db.Where("device_name = ? AND uid = ?", iface.DeviceName, cfg.Uid).Delete(&dbPeerConfig{}).Error + if err != nil { + return errors.Wrapf(err, "failed to delete peer %s from db", cfg.Uid) + } + return nil +} + +func (d DatabaseBackend) Load(identifier DeviceIdentifier) (InterfaceConfig, []PeerConfig, error) { + var iface dbInterfaceConfig + var peerDefaults dbDefaultPeerConfig + var peers []dbPeerConfig + + if err := d.db.Where("device_name = ?", identifier).First(&iface).Error; err != nil { + return InterfaceConfig{}, nil, errors.Wrapf(err, "failed to load interface %s from db", identifier) + } + if err := d.db.Where("device_name = ?", identifier).First(&peerDefaults).Error; err != nil { + return InterfaceConfig{}, nil, errors.Wrapf(err, "failed to load peer defaults for %s from db", identifier) + } + if err := d.db.Where("device_name = ?", identifier).Find(&peers).Error; err != nil { + return InterfaceConfig{}, nil, errors.Wrapf(err, "failed to load peers for %s from db", identifier) + } + + interfaceConfig := InterfaceConfig{ + DeviceName: DeviceIdentifier(iface.DeviceName), + KeyPair: KeyPair{PrivateKey: iface.PrivateKey, PublicKey: iface.PublicKey}, + ListenPort: iface.ListenPort, + AddressStr: iface.AddressStr, + DnsStr: iface.DnsStr, + Mtu: iface.Mtu, + FirewallMark: int32(iface.FirewallMark), + RoutingTable: iface.RoutingTable, + PreUp: iface.PreUp, + PostUp: iface.PostUp, + PreDown: iface.PreDown, + PostDown: iface.PostDown, + SaveConfig: iface.SaveConfig, + Enabled: iface.Enabled, + DisplayName: iface.DisplayName, + Type: InterfaceType(iface.Type), + DriverType: iface.DriverType, + + PeerDefNetworkStr: peerDefaults.NetworkStr, + PeerDefDnsStr: peerDefaults.DnsStr, + PeerDefEndpoint: peerDefaults.Endpoint, + PeerDefAllowedIPsStr: peerDefaults.AllowedIPsStr, + PeerDefMtu: peerDefaults.Mtu, + PeerDefPersistentKeepalive: peerDefaults.PersistentKeepalive, + PeerDefFirewallMark: int32(peerDefaults.FirewallMark), + PeerDefRoutingTable: peerDefaults.RoutingTable, + PeerDefPreUp: peerDefaults.PreUp, + PeerDefPostUp: peerDefaults.PostUp, + PeerDefPreDown: peerDefaults.PreDown, + PeerDefPostDown: peerDefaults.PostDown, + + DisabledAt: nil, + BaseConfig: BaseConfig{ + CreatedAt: iface.CreatedAt, + UpdatedAt: iface.UpdatedAt, + CreatedBy: iface.CreatedBy, + UpdatedBy: iface.UpdatedBy, + }, + } + if iface.DisabledAt.Valid { + interfaceConfig.DisabledAt = &iface.DisabledAt.Time + } + + peerConfigs := make([]PeerConfig, len(peers)) + for i, peer := range peers { + peerConfigs[i] = PeerConfig{ + Endpoint: NewStringConfigOption(peer.Endpoint, peer.OvrEndpoint), + AllowedIPsStr: NewStringConfigOption(peer.AllowedIPsStr, peer.OvrAllowedIPsStr), + ExtraAllowedIPsStr: peer.ExtraAllowedIPsStr, + KeyPair: KeyPair{PrivateKey: peer.PrivateKey, PublicKey: peer.PublicKey}, + PresharedKey: peer.PresharedKey, + PersistentKeepalive: NewIntConfigOption(peer.PersistentKeepalive, peer.OvrPersistentKeepalive), + Identifier: peer.Identifier, + Uid: PeerIdentifier(peer.Uid), + AddressStr: NewStringConfigOption(peer.AddressStr, peer.OvrAddressStr), + DnsStr: NewStringConfigOption(peer.DnsStr, peer.OvrDnsStr), + Mtu: NewIntConfigOption(peer.Mtu, peer.OvrMtu), + FirewallMark: NewInt32ConfigOption(int32(peer.FirewallMark), peer.OvrFirewallMark), + RoutingTable: NewStringConfigOption(peer.RoutingTable, peer.OvrRoutingTable), + PreUp: NewStringConfigOption(peer.PreUp, peer.OvrPreUp), + PostUp: NewStringConfigOption(peer.PostUp, peer.OvrPostUp), + PreDown: NewStringConfigOption(peer.PreDown, peer.OvrPreDown), + PostDown: NewStringConfigOption(peer.PostDown, peer.OvrPostDown), + + DisabledAt: nil, + BaseConfig: BaseConfig{ + CreatedAt: iface.CreatedAt, + UpdatedAt: iface.UpdatedAt, + CreatedBy: iface.CreatedBy, + UpdatedBy: iface.UpdatedBy, + }, + } + + if peer.DisabledAt.Valid { + peerConfigs[i].DisabledAt = &peer.DisabledAt.Time + } + } + + return interfaceConfig, peerConfigs, nil +} + +func (d DatabaseBackend) LoadAll(ignored ...DeviceIdentifier) (map[InterfaceConfig][]PeerConfig, error) { + interfaceIdentifiers := []DeviceIdentifier{} // TODO: fill this ?! + + result := make(map[InterfaceConfig][]PeerConfig) + for _, identifier := range interfaceIdentifiers { + iface, peers, err := d.Load(identifier) + if err != nil { + return nil, errors.Wrapf(err, "failed to load data for %s", identifier) + } + result[iface] = peers + } + + return result, nil +} + +// +// --- Models +// + +type dbBaseModel struct { + CreatedBy string + UpdatedBy string + CreatedAt time.Time + UpdatedAt time.Time +} + +type dbInterfaceConfig struct { + dbBaseModel + DisabledAt sql.NullTime + + // WireGuard specific (for the [interface] section of the config file) + + DeviceName string `gorm:"primaryKey"` + PrivateKey string + PublicKey string + ListenPort int + + AddressStr string + DnsStr string + + Mtu int + FirewallMark int + RoutingTable string + + PreUp string + PostUp string + PreDown string + PostDown string + + SaveConfig bool + + // WG Portal specific + Enabled bool + DisplayName string + Type string + DriverType string + + // Default settings for the peer, used for new peers, those settings will be published to ConfigOption options of + // the peer config + + dbDefaultPeerConfig dbDefaultPeerConfig +} + +func (d dbInterfaceConfig) TableName() string { + return "interface" +} + +type dbDefaultPeerConfig struct { + dbBaseModel + + DeviceName string `gorm:"primaryKey"` // Foreign key + + NetworkStr string // the default subnets from which peers will get their IP addresses, comma seperated + DnsStr string // the default dns server for the peer + Endpoint string // the default endpoint for the peer + AllowedIPsStr string // the default allowed IP string for the peer + Mtu int // the default device MTU + PersistentKeepalive int // the default persistent keep-alive Value + FirewallMark int // default firewall mark + RoutingTable string // the default routing table + + PreUp string // default action that is executed before the device is up + PostUp string // default action that is executed after the device is up + PreDown string // default action that is executed before the device is down + PostDown string // default action that is executed after the device is down +} + +func (d dbDefaultPeerConfig) TableName() string { + return "peer_defaults" +} + +type dbPeerConfig struct { + dbBaseModel + DisabledAt sql.NullTime + + DeviceName string `gorm:"primaryKey"` + Endpoint string + OvrEndpoint bool + AllowedIPsStr string + OvrAllowedIPsStr bool + ExtraAllowedIPsStr string + PrivateKey string + PublicKey string + PresharedKey string + PersistentKeepalive int + OvrPersistentKeepalive bool + + // WG Portal specific + + Identifier string + Uid string `gorm:"primaryKey"` + + // Interface settings for the peer, used to generate the [interface] section in the peer config file + + AddressStr string + OvrAddressStr bool + DnsStr string + OvrDnsStr bool + Mtu int + OvrMtu bool + FirewallMark int + OvrFirewallMark bool + RoutingTable string + OvrRoutingTable bool + + PreUp string + OvrPreUp bool + PostUp string + OvrPostUp bool + PreDown string + OvrPreDown bool + PostDown string + OvrPostDown bool +} + +func (d dbPeerConfig) TableName() string { + return "peer" +} + +func convertPeer(peer PeerConfig, devName DeviceIdentifier) dbPeerConfig { + cfg := dbPeerConfig{ + DeviceName: string(devName), + Endpoint: peer.Endpoint.GetValue(), + OvrEndpoint: peer.Endpoint.Overridable, + AllowedIPsStr: peer.AllowedIPsStr.GetValue(), + OvrAllowedIPsStr: peer.AllowedIPsStr.Overridable, + ExtraAllowedIPsStr: peer.ExtraAllowedIPsStr, + PrivateKey: peer.KeyPair.PrivateKey, + PublicKey: peer.KeyPair.PublicKey, + PresharedKey: peer.PresharedKey, + PersistentKeepalive: peer.PersistentKeepalive.GetValue(), + OvrPersistentKeepalive: peer.PersistentKeepalive.Overridable, + Identifier: peer.Identifier, + Uid: string(peer.Uid), + AddressStr: peer.AddressStr.GetValue(), + OvrAddressStr: peer.AddressStr.Overridable, + DnsStr: peer.DnsStr.GetValue(), + OvrDnsStr: peer.DnsStr.Overridable, + Mtu: peer.Mtu.GetValue(), + OvrMtu: peer.Mtu.Overridable, + FirewallMark: int(peer.FirewallMark.GetValue()), + OvrFirewallMark: peer.FirewallMark.Overridable, + RoutingTable: peer.RoutingTable.GetValue(), + OvrRoutingTable: peer.RoutingTable.Overridable, + PreUp: peer.PreUp.GetValue(), + OvrPreUp: peer.PreUp.Overridable, + PostUp: peer.PostUp.GetValue(), + OvrPostUp: peer.PostUp.Overridable, + PreDown: peer.PreDown.GetValue(), + OvrPreDown: peer.PreDown.Overridable, + PostDown: peer.PostDown.GetValue(), + OvrPostDown: peer.PostDown.Overridable, + DisabledAt: sql.NullTime{Time: time.Time{}, Valid: peer.DisabledAt != nil}, + } + if peer.DisabledAt != nil { + cfg.DisabledAt.Time = *peer.DisabledAt + } + + return cfg +} + +func convertInterface(iface InterfaceConfig) (dbInterfaceConfig, dbDefaultPeerConfig) { + cfg := dbInterfaceConfig{ + DeviceName: string(iface.DeviceName), + PrivateKey: iface.KeyPair.PrivateKey, + PublicKey: iface.KeyPair.PublicKey, + ListenPort: iface.ListenPort, + AddressStr: iface.AddressStr, + DnsStr: iface.DnsStr, + Mtu: iface.Mtu, + FirewallMark: int(iface.FirewallMark), + RoutingTable: iface.RoutingTable, + PreUp: iface.PreUp, + PostUp: iface.PostUp, + PreDown: iface.PreDown, + PostDown: iface.PostDown, + SaveConfig: iface.SaveConfig, + Enabled: iface.Enabled, + DisplayName: iface.DisplayName, + Type: string(iface.Type), + DriverType: iface.DriverType, + DisabledAt: sql.NullTime{Time: time.Time{}, Valid: iface.DisabledAt != nil}, + } + if iface.DisabledAt != nil { + cfg.DisabledAt.Time = *iface.DisabledAt + } + peerDefaults := dbDefaultPeerConfig{ + DeviceName: string(iface.DeviceName), + NetworkStr: iface.PeerDefNetworkStr, + DnsStr: iface.PeerDefDnsStr, + Endpoint: iface.PeerDefEndpoint, + AllowedIPsStr: iface.PeerDefAllowedIPsStr, + Mtu: iface.PeerDefMtu, + PersistentKeepalive: iface.PeerDefPersistentKeepalive, + FirewallMark: int(iface.PeerDefFirewallMark), + RoutingTable: iface.PeerDefRoutingTable, + PreUp: iface.PeerDefPreUp, + PostUp: iface.PeerDefPostUp, + PreDown: iface.PeerDefPreDown, + PostDown: iface.PeerDefPostDown, + } + + return cfg, peerDefaults +} diff --git a/internal/wireguard/backend_db_test.go b/internal/wireguard/backend_db_test.go new file mode 100644 index 0000000..b8b6bbc --- /dev/null +++ b/internal/wireguard/backend_db_test.go @@ -0,0 +1,422 @@ +package wireguard + +import ( + "database/sql" + "database/sql/driver" + "reflect" + "testing" + "time" + + "github.com/pkg/errors" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/DATA-DOG/go-sqlmock" + "gorm.io/driver/mysql" + "gorm.io/gorm" +) + +type AnyTime struct{} + +// Match satisfies sqlmock.Argument interface +func (a AnyTime) Match(v driver.Value) bool { + _, ok := v.(time.Time) + return ok +} + +func getMockedGorm() (*gorm.DB, sqlmock.Sqlmock, error) { + // Default mock with regex matching (https://tienbm90.medium.com/unit-test-for-gorm-application-with-go-sqlmock-ecb5c369e570) + db, mock, err := sqlmock.New() + if err != nil { + return nil, nil, err + } + + gdb, err := gorm.Open(mysql.New(mysql.Config{ + Conn: db, + SkipInitializeWithVersion: true, + }), &gorm.Config{ + SkipDefaultTransaction: true, + }) // open gorm db + if err != nil { + return nil, nil, err + } + return gdb, mock, nil +} + +func TestDatabaseBackend_DeleteInterface(t *testing.T) { + db, mock, err := getMockedGorm() + require.NoError(t, err) + backend := &DatabaseBackend{db: db} + + type args struct { + iface InterfaceConfig + peers []PeerConfig + } + tests := []struct { + name string + mock func() + args args + wantErr bool + }{ + { + name: "Success", + mock: func() { + mock.ExpectExec("DELETE FROM `peer` WHERE device_name = \\?"). + WithArgs("wg0").WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectExec("DELETE FROM `peer_defaults` WHERE device_name = \\?"). + WithArgs("wg0").WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectExec("DELETE FROM `interface` WHERE device_name = \\?"). + WithArgs("wg0").WillReturnResult(sqlmock.NewResult(1, 1)) + }, + args: args{ + iface: InterfaceConfig{DeviceName: "wg0"}, + peers: nil, + }, + wantErr: false, + }, + { + name: "Peer Delete Failure", + mock: func() { + mock.ExpectExec("DELETE FROM `peer` WHERE device_name = \\?"). + WithArgs("wg0").WillReturnError(errors.New("peererr")) + }, + args: args{ + iface: InterfaceConfig{DeviceName: "wg0"}, + peers: nil, + }, + wantErr: true, + }, + { + name: "Peer Defaults Delete Failure", + mock: func() { + mock.ExpectExec("DELETE FROM `peer` WHERE device_name = \\?"). + WithArgs("wg0").WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectExec("DELETE FROM `peer_defaults` WHERE device_name = \\?"). + WithArgs("wg0").WillReturnError(errors.New("defaultserr")) + }, + args: args{ + iface: InterfaceConfig{DeviceName: "wg0"}, + peers: nil, + }, + wantErr: true, + }, + { + name: "Interface Delete Failure", + mock: func() { + mock.ExpectExec("DELETE FROM `peer` WHERE device_name = \\?"). + WithArgs("wg0").WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectExec("DELETE FROM `peer_defaults` WHERE device_name = \\?"). + WithArgs("wg0").WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectExec("DELETE FROM `interface` WHERE device_name = \\?"). + WithArgs("wg0").WillReturnError(errors.New("ifaceerr")) + }, + args: args{ + iface: InterfaceConfig{DeviceName: "wg0"}, + peers: nil, + }, + wantErr: true, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + tt.mock() + if err := backend.DeleteInterface(tt.args.iface, tt.args.peers); (err != nil) != tt.wantErr { + t.Errorf("DeleteInterface() error = %v, wantErr %v", err, tt.wantErr) + } + assert.NoError(t, mock.ExpectationsWereMet()) + }) + } +} + +func TestDatabaseBackend_DeletePeer(t *testing.T) { + db, mock, err := getMockedGorm() + require.NoError(t, err) + backend := &DatabaseBackend{db: db} + + type args struct { + peer PeerConfig + iface InterfaceConfig + } + tests := []struct { + name string + mock func() + args args + wantErr bool + }{ + { + name: "Success", + mock: func() { + mock.ExpectExec("DELETE FROM `peer` WHERE device_name = \\? AND uid = \\?"). + WithArgs("wg0", "peer0").WillReturnResult(sqlmock.NewResult(1, 1)) + }, + args: args{ + peer: PeerConfig{Uid: "peer0"}, + iface: InterfaceConfig{DeviceName: "wg0"}, + }, + wantErr: false, + }, + { + name: "Peer Delete Failure", + mock: func() { + mock.ExpectExec("DELETE FROM `peer` WHERE device_name = \\? AND uid = \\?"). + WithArgs("wg0", "peer0").WillReturnError(errors.New("peererr")) + }, + args: args{ + peer: PeerConfig{Uid: "peer0"}, + iface: InterfaceConfig{DeviceName: "wg0"}, + }, + wantErr: true, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + tt.mock() + if err := backend.DeletePeer(tt.args.peer, tt.args.iface); (err != nil) != tt.wantErr { + t.Errorf("DeletePeer() error = %v, wantErr %v", err, tt.wantErr) + } + assert.NoError(t, mock.ExpectationsWereMet()) + }) + } +} + +func TestDatabaseBackend_Load(t *testing.T) { + type fields struct { + db *gorm.DB + } + type args struct { + identifier DeviceIdentifier + } + tests := []struct { + name string + fields fields + args args + want InterfaceConfig + want1 []PeerConfig + wantErr bool + }{ + // TODO: Add test cases. + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + d := DatabaseBackend{ + db: tt.fields.db, + } + got, got1, err := d.Load(tt.args.identifier) + if (err != nil) != tt.wantErr { + t.Errorf("Load() error = %v, wantErr %v", err, tt.wantErr) + return + } + if !reflect.DeepEqual(got, tt.want) { + t.Errorf("Load() got = %v, want %v", got, tt.want) + } + if !reflect.DeepEqual(got1, tt.want1) { + t.Errorf("Load() got1 = %v, want %v", got1, tt.want1) + } + }) + } +} + +func TestDatabaseBackend_LoadAll(t *testing.T) { + type fields struct { + db *gorm.DB + } + type args struct { + ignored []DeviceIdentifier + } + tests := []struct { + name string + fields fields + args args + want map[InterfaceConfig][]PeerConfig + wantErr bool + }{ + // TODO: Add test cases. + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + d := DatabaseBackend{ + db: tt.fields.db, + } + got, err := d.LoadAll(tt.args.ignored...) + if (err != nil) != tt.wantErr { + t.Errorf("LoadAll() error = %v, wantErr %v", err, tt.wantErr) + return + } + if !reflect.DeepEqual(got, tt.want) { + t.Errorf("LoadAll() got = %v, want %v", got, tt.want) + } + }) + } +} + +func TestDatabaseBackend_SaveInterface(t *testing.T) { + db, mock, err := getMockedGorm() + require.NoError(t, err) + backend := &DatabaseBackend{db: db} + + type args struct { + cfg InterfaceConfig + peers []PeerConfig + } + tests := []struct { + name string + mock func() + args args + wantErr bool + }{ + { + name: "Success Create", + mock: func() { + mock.ExpectExec("INSERT INTO `interface` .*"). + WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectExec("INSERT INTO `peer_defaults` .*"). + WillReturnResult(sqlmock.NewResult(1, 1)) + }, + args: args{ + cfg: InterfaceConfig{DeviceName: "wg0"}, + peers: nil, + }, + wantErr: false, + }, + { + name: "Error Interface", + mock: func() { + mock.ExpectExec("INSERT INTO `interface` .*"). + WillReturnError(errors.New("ifaceerr")) + }, + args: args{ + cfg: InterfaceConfig{DeviceName: "wg0"}, + peers: nil, + }, + wantErr: true, + }, + { + name: "Error Peer Defaults", + mock: func() { + mock.ExpectExec("INSERT INTO `interface` .*"). + WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectExec("INSERT INTO `peer_defaults` .*"). + WillReturnError(errors.New("ifaceerr")) + }, + args: args{ + cfg: InterfaceConfig{DeviceName: "wg0"}, + peers: nil, + }, + wantErr: true, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + tt.mock() + if err := backend.SaveInterface(tt.args.cfg, tt.args.peers); (err != nil) != tt.wantErr { + t.Errorf("SaveInterface() error = %v, wantErr %v", err, tt.wantErr) + } + assert.NoError(t, mock.ExpectationsWereMet()) + }) + } +} + +func TestDatabaseBackend_SavePeer(t *testing.T) { + db, mock, err := getMockedGorm() + require.NoError(t, err) + backend := &DatabaseBackend{db: db} + + type args struct { + peer PeerConfig + iface InterfaceConfig + } + tests := []struct { + name string + mock func() + args args + wantErr bool + }{ + { + name: "Success Create", + mock: func() { + mock.ExpectExec("INSERT INTO `peer` .*"). + WillReturnResult(sqlmock.NewResult(1, 1)) + }, + args: args{ + peer: PeerConfig{Uid: "peer0"}, + iface: InterfaceConfig{DeviceName: "wg0"}, + }, + wantErr: false, + }, + { + name: "Error", + mock: func() { + mock.ExpectExec("INSERT INTO `peer` .*"). + WillReturnError(errors.New("peererr")) + }, + args: args{ + peer: PeerConfig{Uid: "peer0"}, + iface: InterfaceConfig{DeviceName: "wg0"}, + }, + wantErr: true, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + tt.mock() + if err := backend.SavePeer(tt.args.peer, tt.args.iface); (err != nil) != tt.wantErr { + t.Errorf("SavePeer() error = %v, wantErr %v", err, tt.wantErr) + } + assert.NoError(t, mock.ExpectationsWereMet()) + }) + } +} + +func TestNewDatabaseBackend(t *testing.T) { + db, mock, err := getMockedGorm() + require.NoError(t, err) + + // Success + mock.ExpectExec("CREATE TABLE `interface` .*").WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectExec("CREATE TABLE `peer_defaults` .*").WillReturnResult(sqlmock.NewResult(1, 1)) + mock.ExpectExec("CREATE TABLE `peer` .*").WillReturnResult(sqlmock.NewResult(1, 1)) + backend, err := NewDatabaseBackend(db) + assert.NoError(t, err) + assert.NotNil(t, backend) + assert.NoError(t, mock.ExpectationsWereMet()) + + // Migration failure + mock.ExpectExec("CREATE TABLE `interface` .*").WillReturnError(errors.New("migerr")) + backend, err = NewDatabaseBackend(db) + assert.Error(t, err) + assert.Nil(t, backend) + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func Test_convertInterface(t *testing.T) { + config, peerDefaultConfig := convertInterface(InterfaceConfig{}) + assert.Equal(t, dbInterfaceConfig{}, config) + assert.Equal(t, dbDefaultPeerConfig{}, peerDefaultConfig) + + now := time.Now() + config, peerDefaultConfig = convertInterface(InterfaceConfig{DisabledAt: &now}) + assert.Equal(t, dbInterfaceConfig{DisabledAt: sql.NullTime{Time: now, Valid: true}}, config) + assert.Equal(t, dbDefaultPeerConfig{}, peerDefaultConfig) +} + +func Test_convertPeer(t *testing.T) { + peer := convertPeer(PeerConfig{}, "wg0") + assert.Equal(t, dbPeerConfig{DeviceName: "wg0"}, peer) + + now := time.Now() + peer = convertPeer(PeerConfig{DisabledAt: &now}, "wg0") + assert.Equal(t, dbPeerConfig{DeviceName: "wg0", DisabledAt: sql.NullTime{Time: now, Valid: true}}, peer) +} + +func Test_dbDefaultPeerConfig_TableName(t *testing.T) { + assert.Equal(t, "peer_defaults", dbDefaultPeerConfig{}.TableName()) +} + +func Test_dbInterfaceConfig_TableName(t *testing.T) { + assert.Equal(t, "interface", dbInterfaceConfig{}.TableName()) +} + +func Test_dbPeerConfig_TableName(t *testing.T) { + assert.Equal(t, "peer", dbPeerConfig{}.TableName()) +} diff --git a/internal/wireguard/backend_file.go b/internal/wireguard/backend_file.go new file mode 100644 index 0000000..6b330a4 --- /dev/null +++ b/internal/wireguard/backend_file.go @@ -0,0 +1,66 @@ +package wireguard + +import ( + "io" + "os" + "path/filepath" + + "github.com/pkg/errors" +) + +type FileBackend struct { + configurationPath string + fileGenerator ConfigFileGenerator +} + +func NewFileBackend(configStoragePath string, fileGenerator ConfigFileGenerator) (*FileBackend, error) { + backend := &FileBackend{configurationPath: configStoragePath, fileGenerator: fileGenerator} + return backend, nil +} + +func (f FileBackend) SaveInterface(cfg InterfaceConfig, peers []PeerConfig) error { + configContents, err := f.fileGenerator.GetInterfaceConfig(cfg, peers) + if err != nil { + return errors.Wrapf(err, "failed to generate config file contents for %s", cfg.DeviceName) + } + + configFilePath := filepath.Join(f.configurationPath, string(cfg.DeviceName)+".conf") + configFile, err := os.OpenFile(configFilePath, os.O_RDWR|os.O_CREATE|os.O_TRUNC, 0640) + if err != nil { + return errors.Wrapf(err, "failed to create config file for %s", cfg.DeviceName) + } + defer configFile.Close() + + _, err = io.Copy(configFile, configContents) + if err != nil { + return errors.Wrapf(err, "failed to write config file for %s", cfg.DeviceName) + } + + return nil +} + +func (f FileBackend) SavePeer(_ PeerConfig, _ InterfaceConfig) error { + return nil // the file backend will only store changed interfaces +} + +func (f FileBackend) DeleteInterface(cfg InterfaceConfig, _ []PeerConfig) error { + configFilePath := filepath.Join(f.configurationPath, string(cfg.DeviceName)+".conf") + + err := os.Remove(configFilePath) + if err != nil { + return errors.Wrapf(err, "failed to delete config file for %s", cfg.DeviceName) + } + return nil +} + +func (f FileBackend) DeletePeer(_ PeerConfig, _ InterfaceConfig) error { + return nil // the file backend will only store changed interfaces +} + +func (f FileBackend) Load(identifier DeviceIdentifier) (InterfaceConfig, []PeerConfig, error) { + panic("implement me") +} + +func (f FileBackend) LoadAll(ignored ...DeviceIdentifier) (map[InterfaceConfig][]PeerConfig, error) { + panic("implement me") +} diff --git a/internal/wireguard/backend_file_test.go b/internal/wireguard/backend_file_test.go new file mode 100644 index 0000000..4fab92c --- /dev/null +++ b/internal/wireguard/backend_file_test.go @@ -0,0 +1,208 @@ +package wireguard + +import ( + "bytes" + "io" + "io/ioutil" + "os" + "path/filepath" + "reflect" + "strings" + "testing" + + "github.com/pkg/errors" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" + "github.com/stretchr/testify/require" +) + +type MockFileGenerator struct { + mock.Mock +} + +func (m *MockFileGenerator) GetInterfaceConfig(cfg InterfaceConfig, peers []PeerConfig) (io.Reader, error) { + args := m.Called(cfg, peers) + return args.Get(0).(io.Reader), args.Error(1) +} + +func (m *MockFileGenerator) GetPeerConfig(peer PeerConfig, iface InterfaceConfig) (io.Reader, error) { + args := m.Called(peer, iface) + return args.Get(0).(io.Reader), args.Error(1) +} + +func TestFileBackend_DeleteInterface(t *testing.T) { + // setup + tmpDir := os.TempDir() + tmpFile, err := ioutil.TempFile(tmpDir, "wg*.conf") + require.NoError(t, err) + defer os.Remove(tmpFile.Name()) + + f := FileBackend{ + configurationPath: tmpDir, + } + + // Successful delete + err = f.DeleteInterface(InterfaceConfig{ + DeviceName: DeviceIdentifier(strings.ReplaceAll(filepath.Base(tmpFile.Name()), ".conf", "")), + }, nil) + assert.NoError(t, err) + + // Unsuccessful delete + err = f.DeleteInterface(InterfaceConfig{ + DeviceName: DeviceIdentifier(strings.ReplaceAll(filepath.Base(tmpFile.Name()), ".conf", "")), + }, nil) + assert.Error(t, err) +} + +func TestFileBackend_DeletePeer(t *testing.T) { + assert.NoError(t, FileBackend{}.DeletePeer(PeerConfig{}, InterfaceConfig{})) +} + +func TestFileBackend_Load(t *testing.T) { + type fields struct { + configurationPath string + fileGenerator ConfigFileGenerator + } + type args struct { + identifier DeviceIdentifier + } + tests := []struct { + name string + fields fields + args args + want InterfaceConfig + want1 []PeerConfig + wantErr bool + }{ + // TODO: Add test cases. + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + f := FileBackend{ + configurationPath: tt.fields.configurationPath, + fileGenerator: tt.fields.fileGenerator, + } + got, got1, err := f.Load(tt.args.identifier) + if (err != nil) != tt.wantErr { + t.Errorf("Load() error = %v, wantErr %v", err, tt.wantErr) + return + } + if !reflect.DeepEqual(got, tt.want) { + t.Errorf("Load() got = %v, want %v", got, tt.want) + } + if !reflect.DeepEqual(got1, tt.want1) { + t.Errorf("Load() got1 = %v, want %v", got1, tt.want1) + } + }) + } +} + +func TestFileBackend_LoadAll(t *testing.T) { + type fields struct { + configurationPath string + fileGenerator ConfigFileGenerator + } + type args struct { + ignored []DeviceIdentifier + } + tests := []struct { + name string + fields fields + args args + want map[InterfaceConfig][]PeerConfig + wantErr bool + }{ + // TODO: Add test cases. + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + f := FileBackend{ + configurationPath: tt.fields.configurationPath, + fileGenerator: tt.fields.fileGenerator, + } + got, err := f.LoadAll(tt.args.ignored...) + if (err != nil) != tt.wantErr { + t.Errorf("LoadAll() error = %v, wantErr %v", err, tt.wantErr) + return + } + if !reflect.DeepEqual(got, tt.want) { + t.Errorf("LoadAll() got = %v, want %v", got, tt.want) + } + }) + } +} + +func TestFileBackend_SaveInterface(t *testing.T) { + // setup + tmpDir := os.TempDir() + tmpFile, err := ioutil.TempFile(tmpDir, "wg*.conf") + require.NoError(t, err) + defer os.Remove(tmpFile.Name()) + deviceName := strings.ReplaceAll(filepath.Base(tmpFile.Name()), ".conf", "") + + type fields struct { + prepare func(m *mock.Mock) + } + type args struct { + cfg InterfaceConfig + peers []PeerConfig + } + tests := []struct { + name string + fields fields + args args + wantErr bool + }{ + { + name: "FileGeneratorError", + fields: fields{ + prepare: func(m *mock.Mock) { + m.On("GetInterfaceConfig", mock.Anything, mock.Anything). + Return(&bytes.Buffer{}, errors.New("generr")) + }, + }, + args: args{}, + wantErr: true, + }, + { + name: "Success", + fields: fields{ + prepare: func(m *mock.Mock) { + m.On("GetInterfaceConfig", mock.Anything, mock.Anything). + Return(bytes.NewBuffer([]byte("hello world")), nil) + }, + }, + args: args{ + cfg: InterfaceConfig{DeviceName: DeviceIdentifier(deviceName)}, + peers: nil, + }, + wantErr: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + fg := new(MockFileGenerator) + f := FileBackend{ + configurationPath: tmpDir, + fileGenerator: fg, + } + tt.fields.prepare(&fg.Mock) + if err := f.SaveInterface(tt.args.cfg, tt.args.peers); (err != nil) != tt.wantErr { + t.Errorf("SaveInterface() error = %v, wantErr %v", err, tt.wantErr) + } + + fg.AssertExpectations(t) + }) + } +} + +func TestFileBackend_SavePeer(t *testing.T) { + assert.NoError(t, FileBackend{}.SavePeer(PeerConfig{}, InterfaceConfig{})) +} + +func TestNewFileBackend(t *testing.T) { + got, err := NewFileBackend("testing", nil) + assert.NoError(t, err) + assert.NotNil(t, got) +} diff --git a/internal/wireguard/configuration.go b/internal/wireguard/configuration.go index 83c3d68..a026201 100644 --- a/internal/wireguard/configuration.go +++ b/internal/wireguard/configuration.go @@ -1,7 +1,6 @@ package wireguard import ( - "database/sql" "time" ) @@ -16,6 +15,9 @@ } func (o StringConfigOption) GetValue() string { + if o.Value == nil { + return "" + } return o.Value.(string) } @@ -31,6 +33,9 @@ } func (o IntConfigOption) GetValue() int { + if o.Value == nil { + return 0 + } return o.Value.(int) } @@ -46,6 +51,10 @@ } func (o Int32ConfigOption) GetValue() int32 { + if o.Value == nil { + return 0 + } + return o.Value.(int32) } @@ -61,6 +70,10 @@ } func (o BoolConfigOption) GetValue() bool { + if o.Value == nil { + return false + } + return o.Value.(bool) } @@ -81,7 +94,16 @@ type DeviceIdentifier string type PeerIdentifier string +type BaseConfig struct { + CreatedBy string + UpdatedBy string + CreatedAt time.Time + UpdatedAt time.Time +} + type InterfaceConfig struct { + BaseConfig + // WireGuard specific (for the [interface] section of the config file) DeviceName DeviceIdentifier // device name, for example: wg0 @@ -89,7 +111,7 @@ ListenPort int // the listening port, for example: 51820 AddressStr string // the interface ip addresses, comma separated - Dns string // the dns server that should be set if the interface is up + DnsStr string // the dns server that should be set if the interface is up, comma separated Mtu int // the device MTU FirewallMark int32 // a firewall mark @@ -112,9 +134,9 @@ // the peer config PeerDefNetworkStr string // the default subnets from which peers will get their IP addresses, comma seperated - PeerDefDns string // the default dns server for the peer + PeerDefDnsStr string // the default dns server for the peer PeerDefEndpoint string // the default endpoint for the peer - PeerDefAllowedIPsString string // the default allowed IP string for the peer + PeerDefAllowedIPsStr string // the default allowed IP string for the peer PeerDefMtu int // the default device MTU PeerDefPersistentKeepalive int // the default persistent keep-alive Value PeerDefFirewallMark int32 // default firewall mark @@ -124,17 +146,23 @@ PeerDefPostUp string // default action that is executed after the device is up PeerDefPreDown string // default action that is executed before the device is down PeerDefPostDown string // default action that is executed after the device is down + + // Internal stats + + DisabledAt *time.Time } type PeerConfig struct { + BaseConfig + // WireGuard specific (for the [peer] section of the config file) - Endpoint StringConfigOption // the endpoint address - AllowedIPsString StringConfigOption // all allowed ip subnets, comma seperated - ExtraAllowedIPsString string // all allowed ip subnets on the server side, comma seperated - KeyPair KeyPair // private/public Key of the peer - PresharedKey string // the pre-shared Key of the peer - PersistentKeepalive IntConfigOption // the persistent keep-alive interval + Endpoint StringConfigOption // the endpoint address + AllowedIPsStr StringConfigOption // all allowed ip subnets, comma seperated + ExtraAllowedIPsStr string // all allowed ip subnets on the server side, comma seperated + KeyPair KeyPair // private/public Key of the peer + PresharedKey string // the pre-shared Key of the peer + PersistentKeepalive IntConfigOption // the persistent keep-alive interval // WG Portal specific @@ -144,7 +172,7 @@ // Interface settings for the peer, used to generate the [interface] section in the peer config file AddressStr StringConfigOption // the interface ip addresses, comma separated - Dns StringConfigOption // the dns server that should be set if the interface is up + DnsStr StringConfigOption // the dns server that should be set if the interface is up, comma separated Mtu IntConfigOption // the device MTU FirewallMark Int32ConfigOption // a firewall mark RoutingTable StringConfigOption // the routing table @@ -156,11 +184,7 @@ // Internal stats - DeactivatedAt sql.NullTime - CreatedBy string - UpdatedBy string - CreatedAt time.Time - UpdatedAt time.Time + DisabledAt *time.Time } // ConfigWriter provides methods for updating persistent backends (like a database or a WireGuard configuration file) @@ -174,5 +198,5 @@ // ConfigLoader provides methods to load interface and peer configurations from a persistent backend. type ConfigLoader interface { Load(identifier DeviceIdentifier) (InterfaceConfig, []PeerConfig, error) - LoadAll() (map[InterfaceConfig][]PeerConfig, error) + LoadAll(ignored ...DeviceIdentifier) (map[InterfaceConfig][]PeerConfig, error) } diff --git a/internal/wireguard/configuration_test.go b/internal/wireguard/configuration_test.go new file mode 100644 index 0000000..97ce95d --- /dev/null +++ b/internal/wireguard/configuration_test.go @@ -0,0 +1,259 @@ +package wireguard + +import ( + "reflect" + "testing" +) + +func TestBoolConfigOption_GetValue(t *testing.T) { + type fields struct { + ConfigOption ConfigOption + } + tests := []struct { + name string + fields fields + want bool + }{ + { + name: "Empty", + fields: fields{}, + want: false, + }, + { + name: "True", + fields: fields{ConfigOption: ConfigOption{Value: true}}, + want: true, + }, + { + name: "False", + fields: fields{ConfigOption: ConfigOption{Value: false}}, + want: false, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + o := BoolConfigOption{ + ConfigOption: tt.fields.ConfigOption, + } + if got := o.GetValue(); got != tt.want { + t.Errorf("GetValue() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestInt32ConfigOption_GetValue(t *testing.T) { + type fields struct { + ConfigOption ConfigOption + } + tests := []struct { + name string + fields fields + want int32 + }{ + { + name: "Empty", + fields: fields{}, + want: 0, + }, + { + name: "Leet", + fields: fields{ConfigOption: ConfigOption{Value: int32(1337)}}, + want: 1337, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + o := Int32ConfigOption{ + ConfigOption: tt.fields.ConfigOption, + } + if got := o.GetValue(); got != tt.want { + t.Errorf("GetValue() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestIntConfigOption_GetValue(t *testing.T) { + type fields struct { + ConfigOption ConfigOption + } + tests := []struct { + name string + fields fields + want int + }{ + { + name: "Empty", + fields: fields{}, + want: 0, + }, + { + name: "Leet", + fields: fields{ConfigOption: ConfigOption{Value: 1337}}, + want: 1337, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + o := IntConfigOption{ + ConfigOption: tt.fields.ConfigOption, + } + if got := o.GetValue(); got != tt.want { + t.Errorf("GetValue() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestStringConfigOption_GetValue(t *testing.T) { + type fields struct { + ConfigOption ConfigOption + } + tests := []struct { + name string + fields fields + want string + }{ + { + name: "Empty", + fields: fields{}, + want: "", + }, + { + name: "Leet", + fields: fields{ConfigOption: ConfigOption{Value: "leet"}}, + want: "leet", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + o := StringConfigOption{ + ConfigOption: tt.fields.ConfigOption, + } + if got := o.GetValue(); got != tt.want { + t.Errorf("GetValue() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestNewBoolConfigOption(t *testing.T) { + type args struct { + value bool + overridable bool + } + tests := []struct { + name string + args args + want BoolConfigOption + }{ + { + name: "Overridable", + args: args{value: false, overridable: true}, + want: BoolConfigOption{ConfigOption: ConfigOption{Value: false, Overridable: true}}, + }, + { + name: "Not Overridable", + args: args{value: true, overridable: false}, + want: BoolConfigOption{ConfigOption: ConfigOption{Value: true, Overridable: false}}, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := NewBoolConfigOption(tt.args.value, tt.args.overridable); !reflect.DeepEqual(got, tt.want) { + t.Errorf("NewBoolConfigOption() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestNewInt32ConfigOption(t *testing.T) { + type args struct { + value int32 + overridable bool + } + tests := []struct { + name string + args args + want Int32ConfigOption + }{ + { + name: "Overridable", + args: args{value: 1337, overridable: true}, + want: Int32ConfigOption{ConfigOption: ConfigOption{Value: int32(1337), Overridable: true}}, + }, + { + name: "Not Overridable", + args: args{value: 1337, overridable: false}, + want: Int32ConfigOption{ConfigOption: ConfigOption{Value: int32(1337), Overridable: false}}, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := NewInt32ConfigOption(tt.args.value, tt.args.overridable); !reflect.DeepEqual(got, tt.want) { + t.Errorf("NewInt32ConfigOption() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestNewIntConfigOption(t *testing.T) { + type args struct { + value int + overridable bool + } + tests := []struct { + name string + args args + want IntConfigOption + }{ + { + name: "Overridable", + args: args{value: 1337, overridable: true}, + want: IntConfigOption{ConfigOption: ConfigOption{Value: 1337, Overridable: true}}, + }, + { + name: "Not Overridable", + args: args{value: 1337, overridable: false}, + want: IntConfigOption{ConfigOption: ConfigOption{Value: 1337, Overridable: false}}, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := NewIntConfigOption(tt.args.value, tt.args.overridable); !reflect.DeepEqual(got, tt.want) { + t.Errorf("NewIntConfigOption() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestNewStringConfigOption(t *testing.T) { + type args struct { + value string + overridable bool + } + tests := []struct { + name string + args args + want StringConfigOption + }{ + { + name: "Overridable", + args: args{value: "leet", overridable: true}, + want: StringConfigOption{ConfigOption: ConfigOption{Value: "leet", Overridable: true}}, + }, + { + name: "Not Overridable", + args: args{value: "leet", overridable: false}, + want: StringConfigOption{ConfigOption: ConfigOption{Value: "leet", Overridable: false}}, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := NewStringConfigOption(tt.args.value, tt.args.overridable); !reflect.DeepEqual(got, tt.want) { + t.Errorf("NewStringConfigOption() = %v, want %v", got, tt.want) + } + }) + } +} diff --git a/internal/wireguard/manager.go b/internal/wireguard/manager.go index 0831e71..defb042 100644 --- a/internal/wireguard/manager.go +++ b/internal/wireguard/manager.go @@ -10,6 +10,8 @@ "sync" "time" + "github.com/h44z/wg-portal/internal/lowlevel" + "github.com/pkg/errors" "github.com/vishvananda/netlink" "golang.zx2c4.com/wireguard/wgctrl/wgtypes" @@ -42,8 +44,8 @@ type ManagementUtil struct { mux sync.RWMutex // mutex to synchronize access to maps - wg Client - nl NetlinkClient + wg lowlevel.WireGuardClient + nl lowlevel.NetlinkClient // config writers and loaders are used to populate the internal config maps cw []ConfigWriter @@ -315,16 +317,16 @@ var peerAllowedIPs []*netlink.Addr switch deviceType { case InterfaceTypeClient: - peerAllowedIPs, err = parseIpAddressString(peer.AllowedIPsString.GetValue()) + peerAllowedIPs, err = parseIpAddressString(peer.AllowedIPsStr.GetValue()) if err != nil { return wgtypes.PeerConfig{}, errors.Wrapf(err, "failed to parse allowed IP's for peer %s", peer.Uid) } case InterfaceTypeServer: - peerAllowedIPs, err = parseIpAddressString(peer.AllowedIPsString.GetValue()) + peerAllowedIPs, err = parseIpAddressString(peer.AllowedIPsStr.GetValue()) if err != nil { return wgtypes.PeerConfig{}, errors.Wrapf(err, "failed to parse allowed IP's for peer %s", peer.Uid) } - peerExtraAllowedIPs, err := parseIpAddressString(peer.ExtraAllowedIPsString) + peerExtraAllowedIPs, err := parseIpAddressString(peer.ExtraAllowedIPsStr) if err != nil { return wgtypes.PeerConfig{}, errors.Wrapf(err, "failed to parse extra allowed IP's for peer %s", peer.Uid) } @@ -438,7 +440,7 @@ if err != nil { continue } - interfaces[i].Dns = parsedInterface.Dns + interfaces[i].DnsStr = parsedInterface.DnsStr interfaces[i].DisplayName = parsedInterface.DisplayName interfaces[i].PostDown = parsedInterface.PostDown interfaces[i].PreDown = parsedInterface.PreDown @@ -528,7 +530,7 @@ } iface.Mtu = mtu case "dns": - iface.Dns = value + iface.DnsStr = value case "table": iface.RoutingTable = value case "fwmark": diff --git a/internal/wireguard/manager_int_test.go b/internal/wireguard/manager_int_test.go index 74fe4be..366b850 100644 --- a/internal/wireguard/manager_int_test.go +++ b/internal/wireguard/manager_int_test.go @@ -11,6 +11,8 @@ "os/exec" "testing" + "github.com/h44z/wg-portal/internal/lowlevel" + "golang.zx2c4.com/wireguard/wgctrl" "github.com/stretchr/testify/assert" @@ -29,7 +31,7 @@ func TestManagementUtil_CreateDevice(t *testing.T) { devName := DeviceIdentifier("wg666") prepareTest(devName) - m := ManagementUtil{interfaces: make(map[DeviceIdentifier]InterfaceConfig), nl: NetlinkManager{}} + m := ManagementUtil{interfaces: make(map[DeviceIdentifier]InterfaceConfig), nl: lowlevel.NetlinkManager{}} defer m.DeleteDevice(devName) err := m.CreateDevice(devName) @@ -44,7 +46,7 @@ func TestManagementUtil_DeleteDevice(t *testing.T) { devName := DeviceIdentifier("wg667") prepareTest(devName) - m := ManagementUtil{interfaces: make(map[DeviceIdentifier]InterfaceConfig), nl: NetlinkManager{}} + m := ManagementUtil{interfaces: make(map[DeviceIdentifier]InterfaceConfig), nl: lowlevel.NetlinkManager{}} err := m.CreateDevice(devName) assert.NoError(t, err) @@ -72,7 +74,7 @@ if !assert.NoError(t, err) { return } - m := ManagementUtil{interfaces: make(map[DeviceIdentifier]InterfaceConfig), nl: NetlinkManager{}, wg: wg} + m := ManagementUtil{interfaces: make(map[DeviceIdentifier]InterfaceConfig), nl: lowlevel.NetlinkManager{}, wg: wg} defer m.DeleteDevice(devName) err = m.CreateDevice(devName) diff --git a/internal/wireguard/persistant_db.go b/internal/wireguard/persistant_db.go deleted file mode 100644 index 736ef26..0000000 --- a/internal/wireguard/persistant_db.go +++ /dev/null @@ -1,354 +0,0 @@ -package wireguard - -import ( - "database/sql" - "time" - - "github.com/pkg/errors" - "gorm.io/gorm" -) - -type DatabaseBackend struct { - DB *gorm.DB -} - -func (d DatabaseBackend) SaveInterface(cfg InterfaceConfig, _ []PeerConfig) error { - iface, peerDefaults := convertInterface(cfg) - - if err := d.DB.Save(&iface).Error; err != nil { - return errors.Wrapf(err, "failed to save interface %s to db", cfg.DeviceName) - } - if err := d.DB.Save(&peerDefaults).Error; err != nil { - return errors.Wrapf(err, "failed to save peer defaults of %s to db", cfg.DeviceName) - } - - return nil -} - -func (d DatabaseBackend) SavePeer(cfg PeerConfig, iface InterfaceConfig) error { - peer := convertPeer(cfg, iface.DeviceName) - - if err := d.DB.Save(&peer).Error; err != nil { - return errors.Wrapf(err, "failed to save peer %s to db", cfg.Uid) - } - - return nil -} - -func (d DatabaseBackend) DeleteInterface(cfg InterfaceConfig, _ []PeerConfig) error { - // Delete peers - if err := d.DB.Where("device_name = ?", cfg.DeviceName).Delete(&dbPeerConfig{}).Error; err != nil { - return errors.Wrapf(err, "failed to delete peer for %s from db", cfg.DeviceName) - } - // Delete peer default settings - if err := d.DB.Where("device_name = ?", cfg.DeviceName).Delete(&dbDefaultPeerConfig{}).Error; err != nil { - return errors.Wrapf(err, "failed to delete peer defaults for %s from db", cfg.DeviceName) - } - // Delete interface config - if err := d.DB.Where("device_name = ?", cfg.DeviceName).Delete(&dbInterfaceConfig{}).Error; err != nil { - return errors.Wrapf(err, "failed to delete interface %s from db", cfg.DeviceName) - } - return nil -} - -func (d DatabaseBackend) DeletePeer(cfg PeerConfig, iface InterfaceConfig) error { - err := d.DB.Where("device_name = ? AND uid = ?", iface.DeviceName, cfg.Uid).Delete(&dbPeerConfig{}).Error - if err != nil { - return errors.Wrapf(err, "failed to delete peer %s from db", cfg.Uid) - } - return nil -} - -func (d DatabaseBackend) Load(identifier DeviceIdentifier) (InterfaceConfig, []PeerConfig, error) { - var iface dbInterfaceConfig - var peerDefaults dbDefaultPeerConfig - var peers []dbPeerConfig - - if err := d.DB.Where("device_name = ?", identifier).First(&iface).Error; err != nil { - return InterfaceConfig{}, nil, errors.Wrapf(err, "failed to load interface %s from db", identifier) - } - if err := d.DB.Where("device_name = ?", identifier).First(&peerDefaults).Error; err != nil { - return InterfaceConfig{}, nil, errors.Wrapf(err, "failed to load peer defaults for %s from db", identifier) - } - if err := d.DB.Where("device_name = ?", identifier).Find(&peers).Error; err != nil { - return InterfaceConfig{}, nil, errors.Wrapf(err, "failed to load peers for %s from db", identifier) - } - - interfaceConfig := InterfaceConfig{ - DeviceName: DeviceIdentifier(iface.DeviceName), - KeyPair: KeyPair{PrivateKey: iface.PrivateKey, PublicKey: iface.PublicKey}, - ListenPort: iface.ListenPort, - AddressStr: iface.AddressStr, - Dns: iface.Dns, - Mtu: iface.Mtu, - FirewallMark: int32(iface.FirewallMark), - RoutingTable: iface.RoutingTable, - PreUp: iface.PreUp, - PostUp: iface.PostUp, - PreDown: iface.PreDown, - PostDown: iface.PostDown, - SaveConfig: iface.SaveConfig, - Enabled: iface.Enabled, - DisplayName: iface.DisplayName, - Type: InterfaceType(iface.Type), - DriverType: iface.DriverType, - - PeerDefNetworkStr: peerDefaults.NetworkStr, - PeerDefDns: peerDefaults.Dns, - PeerDefEndpoint: peerDefaults.Endpoint, - PeerDefAllowedIPsString: peerDefaults.AllowedIPsString, - PeerDefMtu: peerDefaults.Mtu, - PeerDefPersistentKeepalive: peerDefaults.PersistentKeepalive, - PeerDefFirewallMark: int32(peerDefaults.FirewallMark), - PeerDefRoutingTable: peerDefaults.RoutingTable, - PeerDefPreUp: peerDefaults.PreUp, - PeerDefPostUp: peerDefaults.PostUp, - PeerDefPreDown: peerDefaults.PreDown, - PeerDefPostDown: peerDefaults.PostDown, - } - peerConfigs := make([]PeerConfig, len(peers)) - for i, peer := range peers { - peerConfigs[i] = PeerConfig{ - Endpoint: NewStringConfigOption(peer.Endpoint, peer.OvrEndpoint), - AllowedIPsString: NewStringConfigOption(peer.AllowedIPsString, peer.OvrAllowedIPsString), - ExtraAllowedIPsString: peer.ExtraAllowedIPsString, - KeyPair: KeyPair{PrivateKey: peer.PrivateKey, PublicKey: peer.PublicKey}, - PresharedKey: peer.PresharedKey, - PersistentKeepalive: NewIntConfigOption(peer.PersistentKeepalive, peer.OvrPersistentKeepalive), - Identifier: peer.Identifier, - Uid: PeerIdentifier(peer.Uid), - AddressStr: NewStringConfigOption(peer.AddressStr, peer.OvrAddressStr), - Dns: NewStringConfigOption(peer.Dns, peer.OvrDns), - Mtu: NewIntConfigOption(peer.Mtu, peer.OvrMtu), - FirewallMark: NewInt32ConfigOption(int32(peer.FirewallMark), peer.OvrFirewallMark), - RoutingTable: NewStringConfigOption(peer.RoutingTable, peer.OvrRoutingTable), - PreUp: NewStringConfigOption(peer.PreUp, peer.OvrPreUp), - PostUp: NewStringConfigOption(peer.PostUp, peer.OvrPostUp), - PreDown: NewStringConfigOption(peer.PreDown, peer.OvrPreDown), - PostDown: NewStringConfigOption(peer.PostDown, peer.OvrPostDown), - - DeactivatedAt: peer.DisabledAt, - CreatedBy: peer.CreatedBy, - UpdatedBy: peer.UpdatedBy, - CreatedAt: peer.CreatedAt, - UpdatedAt: peer.UpdatedAt, - } - } - - return interfaceConfig, peerConfigs, nil -} - -func (d DatabaseBackend) LoadAll() (map[InterfaceConfig][]PeerConfig, error) { - interfaceIdentifiers := []DeviceIdentifier{} // TODO: fill this ?! - - result := make(map[InterfaceConfig][]PeerConfig) - for _, identifier := range interfaceIdentifiers { - iface, peers, err := d.Load(identifier) - if err != nil { - return nil, errors.Wrapf(err, "failed to load data for %s", identifier) - } - result[iface] = peers - } - - return result, nil -} - -// -// --- Models -// - -type dbBaseModel struct { - CreatedBy string - UpdatedBy string - CreatedAt time.Time - UpdatedAt time.Time -} - -type dbInterfaceConfig struct { - dbBaseModel - DisabledAt sql.NullTime - - // WireGuard specific (for the [interface] section of the config file) - - DeviceName string `gorm:"primaryKey"` - PrivateKey string - PublicKey string - ListenPort int - - AddressStr string - Dns string - - Mtu int - FirewallMark int - RoutingTable string - - PreUp string - PostUp string - PreDown string - PostDown string - - SaveConfig bool - - // WG Portal specific - Enabled bool - DisplayName string - Type string - DriverType string - - // Default settings for the peer, used for new peers, those settings will be published to ConfigOption options of - // the peer config - - dbDefaultPeerConfig dbDefaultPeerConfig -} - -func (d dbInterfaceConfig) TableName() string { - return "interface" -} - -type dbDefaultPeerConfig struct { - dbBaseModel - - DeviceName string `gorm:"primaryKey"` // Foreign key - - NetworkStr string // the default subnets from which peers will get their IP addresses, comma seperated - Dns string // the default dns server for the peer - Endpoint string // the default endpoint for the peer - AllowedIPsString string // the default allowed IP string for the peer - Mtu int // the default device MTU - PersistentKeepalive int // the default persistent keep-alive Value - FirewallMark int // default firewall mark - RoutingTable string // the default routing table - - PreUp string // default action that is executed before the device is up - PostUp string // default action that is executed after the device is up - PreDown string // default action that is executed before the device is down - PostDown string // default action that is executed after the device is down -} - -func (d dbDefaultPeerConfig) TableName() string { - return "peer_defaults" -} - -type dbPeerConfig struct { - dbBaseModel - DisabledAt sql.NullTime - - DeviceName string `gorm:"primaryKey"` - Endpoint string - OvrEndpoint bool - AllowedIPsString string - OvrAllowedIPsString bool - ExtraAllowedIPsString string - PrivateKey string - PublicKey string - PresharedKey string - PersistentKeepalive int - OvrPersistentKeepalive bool - - // WG Portal specific - - Identifier string - Uid string `gorm:"primaryKey"` - - // Interface settings for the peer, used to generate the [interface] section in the peer config file - - AddressStr string - OvrAddressStr bool - Dns string - OvrDns bool - Mtu int - OvrMtu bool - FirewallMark int - OvrFirewallMark bool - RoutingTable string - OvrRoutingTable bool - - PreUp string - OvrPreUp bool - PostUp string - OvrPostUp bool - PreDown string - OvrPreDown bool - PostDown string - OvrPostDown bool -} - -func (d dbPeerConfig) TableName() string { - return "peer" -} - -func convertPeer(peer PeerConfig, devName DeviceIdentifier) dbPeerConfig { - return dbPeerConfig{ - DeviceName: string(devName), - Endpoint: peer.Endpoint.GetValue(), - OvrEndpoint: peer.Endpoint.Overridable, - AllowedIPsString: peer.AllowedIPsString.GetValue(), - OvrAllowedIPsString: peer.AllowedIPsString.Overridable, - ExtraAllowedIPsString: peer.ExtraAllowedIPsString, - PrivateKey: peer.KeyPair.PrivateKey, - PublicKey: peer.KeyPair.PublicKey, - PresharedKey: peer.PresharedKey, - PersistentKeepalive: peer.PersistentKeepalive.GetValue(), - OvrPersistentKeepalive: peer.PersistentKeepalive.Overridable, - Identifier: peer.Identifier, - Uid: string(peer.Uid), - AddressStr: peer.AddressStr.GetValue(), - OvrAddressStr: peer.AddressStr.Overridable, - Dns: peer.Dns.GetValue(), - OvrDns: peer.Dns.Overridable, - Mtu: peer.Mtu.GetValue(), - OvrMtu: peer.Mtu.Overridable, - FirewallMark: int(peer.FirewallMark.GetValue()), - OvrFirewallMark: peer.FirewallMark.Overridable, - RoutingTable: peer.RoutingTable.GetValue(), - OvrRoutingTable: peer.RoutingTable.Overridable, - PreUp: peer.PreUp.GetValue(), - OvrPreUp: peer.PreUp.Overridable, - PostUp: peer.PostUp.GetValue(), - OvrPostUp: peer.PostUp.Overridable, - PreDown: peer.PreDown.GetValue(), - OvrPreDown: peer.PreDown.Overridable, - PostDown: peer.PostDown.GetValue(), - OvrPostDown: peer.PostDown.Overridable, - } -} - -func convertInterface(iface InterfaceConfig) (dbInterfaceConfig, dbDefaultPeerConfig) { - cfg := dbInterfaceConfig{ - DeviceName: string(iface.DeviceName), - PrivateKey: iface.KeyPair.PrivateKey, - PublicKey: iface.KeyPair.PublicKey, - ListenPort: iface.ListenPort, - AddressStr: iface.AddressStr, - Dns: iface.Dns, - Mtu: iface.Mtu, - FirewallMark: int(iface.FirewallMark), - RoutingTable: iface.RoutingTable, - PreUp: iface.PreUp, - PostUp: iface.PostUp, - PreDown: iface.PreDown, - PostDown: iface.PostDown, - SaveConfig: iface.SaveConfig, - Enabled: iface.Enabled, - DisplayName: iface.DisplayName, - Type: string(iface.Type), - DriverType: iface.DriverType, - } - peerDefaults := dbDefaultPeerConfig{ - DeviceName: string(iface.DeviceName), - NetworkStr: iface.PeerDefNetworkStr, - Dns: iface.PeerDefDns, - Endpoint: iface.PeerDefEndpoint, - AllowedIPsString: iface.PeerDefAllowedIPsString, - Mtu: iface.PeerDefMtu, - PersistentKeepalive: iface.PeerDefPersistentKeepalive, - FirewallMark: int(iface.PeerDefFirewallMark), - RoutingTable: iface.PeerDefRoutingTable, - PreUp: iface.PeerDefPreUp, - PostUp: iface.PeerDefPostUp, - PreDown: iface.PeerDefPreDown, - PostDown: iface.PeerDefPostDown, - } - - return cfg, peerDefaults -} diff --git a/internal/wireguard/persistant_file.go b/internal/wireguard/persistant_file.go deleted file mode 100644 index f9fcf21..0000000 --- a/internal/wireguard/persistant_file.go +++ /dev/null @@ -1,29 +0,0 @@ -package wireguard - -type FileBackend struct { - ConfigurationPath string -} - -func (f FileBackend) SaveInterface(cfg InterfaceConfig, peers []PeerConfig) error { - panic("implement me") -} - -func (f FileBackend) SavePeer(peer PeerConfig, cfg InterfaceConfig) error { - panic("implement me") -} - -func (f FileBackend) DeleteInterface(cfg InterfaceConfig, peers []PeerConfig) error { - panic("implement me") -} - -func (f FileBackend) DeletePeer(peer PeerConfig, cfg InterfaceConfig) error { - panic("implement me") -} - -func (f FileBackend) Load(identifier DeviceIdentifier) (InterfaceConfig, []PeerConfig, error) { - panic("implement me") -} - -func (f FileBackend) LoadAll() (map[InterfaceConfig][]PeerConfig, error) { - panic("implement me") -} diff --git a/internal/wireguard/template.go b/internal/wireguard/template.go new file mode 100644 index 0000000..abea343 --- /dev/null +++ b/internal/wireguard/template.go @@ -0,0 +1,73 @@ +package wireguard + +import ( + "bytes" + "embed" + "io" + "text/template" + + "github.com/pkg/errors" +) + +//go:embed tpl_files/* +var TemplateFiles embed.FS + +type ConfigFileGenerator interface { + GetInterfaceConfig(cfg InterfaceConfig, peers []PeerConfig) (io.Reader, error) + GetPeerConfig(peer PeerConfig, iface InterfaceConfig) (io.Reader, error) +} + +type ConfigFileParser interface { + ParseConfig(fileContents io.Reader) (InterfaceConfig, []PeerConfig, error) +} + +type TemplateHandler struct { + templates *template.Template +} + +func NewTemplateHandler() (*TemplateHandler, error) { + templateCache, err := template.New("WireGuard").ParseFS(TemplateFiles, "tpl_files/*.tpl") + if err != nil { + return nil, errors.Wrapf(err, "failed to parse template files") + } + + handler := &TemplateHandler{ + templates: templateCache, + } + + return handler, nil +} + +func (c TemplateHandler) GetInterfaceConfig(cfg InterfaceConfig, peers []PeerConfig) (io.Reader, error) { + var tplBuff bytes.Buffer + + err := c.templates.ExecuteTemplate(&tplBuff, "interface.tpl", map[string]interface{}{ + "Interface": cfg, + "Peers": peers, + "Portal": map[string]interface{}{ + "Version": "unknown", + }, + }) + if err != nil { + return nil, errors.Wrapf(err, "failed to execute interface template for %s", cfg.DeviceName) + } + + return &tplBuff, nil +} + +func (c TemplateHandler) GetPeerConfig(peer PeerConfig, iface InterfaceConfig) (io.Reader, error) { + var tplBuff bytes.Buffer + + err := c.templates.ExecuteTemplate(&tplBuff, "peer.tpl", map[string]interface{}{ + "Peer": peer, + "Interface": iface, + "Portal": map[string]interface{}{ + "Version": "unknown", + }, + }) + if err != nil { + return nil, errors.Wrapf(err, "failed to execute peer template for %s", peer.Uid) + } + + return &tplBuff, nil +} diff --git a/internal/wireguard/template_test.go b/internal/wireguard/template_test.go new file mode 100644 index 0000000..9478ca4 --- /dev/null +++ b/internal/wireguard/template_test.go @@ -0,0 +1,129 @@ +package wireguard + +import ( + "bytes" + "io" + "reflect" + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestNewTemplateHandler(t *testing.T) { + got, err := NewTemplateHandler() + assert.NoError(t, err) + assert.NotNil(t, got) +} + +func TestTemplateHandler_GetInterfaceConfig(t *testing.T) { + type args struct { + cfg InterfaceConfig + peers []PeerConfig + } + tests := []struct { + name string + args args + want io.Reader + wantErr bool + }{ + { + name: "All Empty", + args: args{}, + want: bytes.NewBuffer([]byte(`# AUTOGENERATED FILE - DO NOT EDIT +# This file uses wg-quick format. See https://man7.org/linux/man-pages/man8/wg-quick.8.html#CONFIGURATION + +# -WGP- WIREGUARD PORTAL CONFIGURATION FILE, version unknown +# Lines starting with the -WGP- tag are used by the WireGuard Portal configuration parser. + +[Interface] +# -WGP- Interface: | Updated: 0001-01-01 00:00:00 +0000 UTC | Created: 0001-01-01 00:00:00 +0000 UTC +# -WGP- Display name: +# -WGP- Interface mode: +# -WGP- PublicKey = + +# Core settings +PrivateKey = +Address = + +# Misc. settings (optional) + +# Interface hooks (optional) + +# +# Peers +# + +`)), + wantErr: false, + }, + } + + c, _ := NewTemplateHandler() + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := c.GetInterfaceConfig(tt.args.cfg, tt.args.peers) + if (err != nil) != tt.wantErr { + t.Errorf("GetInterfaceConfig() error = %v, wantErr %v", err, tt.wantErr) + return + } + if !reflect.DeepEqual(got, tt.want) { + t.Errorf("GetInterfaceConfig() got = %v, want %v", got, tt.want) + } + }) + } +} + +func TestTemplateHandler_GetPeerConfig(t *testing.T) { + type args struct { + peer PeerConfig + iface InterfaceConfig + } + tests := []struct { + name string + args args + want io.Reader + wantErr bool + }{ + { + name: "All empty", + args: args{}, + want: bytes.NewBuffer([]byte(`# AUTOGENERATED FILE - DO NOT EDIT +# This file uses wg-quick format. See https://man7.org/linux/man-pages/man8/wg-quick.8.html#CONFIGURATION + +# -WGP- WIREGUARD PORTAL CONFIGURATION FILE, version unknown +# Lines starting with the -WGP- tag are used by the WireGuard Portal configuration parser. + +[Interface] +# -WGP- Peer: | Updated: 0001-01-01 00:00:00 +0000 UTC | Created: 0001-01-01 00:00:00 +0000 UTC +# -WGP- Display name: +# -WGP- PublicKey: +# -WGP- Peer type: server + +# Core settings +PrivateKey = +Address = + +# Misc. settings (optional) + +# Interface hooks (optional) + +[Peer] +PublicKey = +Endpoint = `)), + wantErr: false, + }, + } + c, _ := NewTemplateHandler() + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := c.GetPeerConfig(tt.args.peer, tt.args.iface) + if (err != nil) != tt.wantErr { + t.Errorf("GetPeerConfig() error = %v, wantErr %v", err, tt.wantErr) + return + } + if !reflect.DeepEqual(got, tt.want) { + t.Errorf("GetPeerConfig() got = %v, want %v", got, tt.want) + } + }) + } +} diff --git a/internal/wireguard/tpl_files/interface.tpl b/internal/wireguard/tpl_files/interface.tpl new file mode 100644 index 0000000..2abb63d --- /dev/null +++ b/internal/wireguard/tpl_files/interface.tpl @@ -0,0 +1,82 @@ +# AUTOGENERATED FILE - DO NOT EDIT +# This file uses wg-quick format. See https://man7.org/linux/man-pages/man8/wg-quick.8.html#CONFIGURATION + +# -WGP- WIREGUARD PORTAL CONFIGURATION FILE, version {{ .Portal.Version }} +# Lines starting with the -WGP- tag are used by the WireGuard Portal configuration parser. + +[Interface] +# -WGP- Interface: {{ .Interface.DeviceName }} | Updated: {{ .Interface.UpdatedAt }} | Created: {{ .Interface.CreatedAt }} +# -WGP- Display name: {{ .Interface.DisplayName }} +# -WGP- Interface mode: {{ .Interface.Type }} +# -WGP- PublicKey = {{ .Interface.KeyPair.PublicKey }} + +# Core settings +PrivateKey = {{ .Interface.KeyPair.PrivateKey }} +Address = {{ .Interface.AddressStr }} + +# Misc. settings (optional) +{{- if ne .Interface.ListenPort 0}} +ListenPort = {{ .Interface.ListenPort }} +{{- end}} +{{- if ne .Interface.Mtu 0}} +MTU = {{.Interface.Mtu}} +{{- end}} +{{- if and (ne .Interface.DnsStr "") (eq $.Interface.Type "client")}} +DNS = {{ .Interface.DnsStr }} +{{- end}} +{{- if ne .Interface.FirewallMark 0}} +FwMark = {{.Interface.FirewallMark}} +{{- end}} +{{- if ne .Interface.RoutingTable ""}} +Table = {{.Interface.RoutingTable}} +{{- end}} +{{- if .Interface.SaveConfig}} +SaveConfig = true +{{- end}} + +# Interface hooks (optional) +{{- if .Interface.PreUp}} +PreUp = {{ .Interface.PreUp }} +{{- end}} +{{- if .Interface.PostUp}} +PostUp = {{ .Interface.PostUp }} +{{- end}} +{{- if .Interface.PreDown}} +PreDown = {{ .Interface.PreDown }} +{{- end}} +{{- if .Interface.PostDown}} +PostDown = {{ .Interface.PostDown }} +{{- end}} + +# +# Peers +# + +{{range .Peers}} +{{- if not .DisabledAt}} +[Peer] +# -WGP- Peer: {{.Uid}} | Updated: {{.UpdatedAt}} | Created: {{.CreatedAt}} +# -WGP- Display name: {{ .Identifier }} +{{- if .KeyPair.PrivateKey}} +# -WGP- PrivateKey: {{.KeyPair.PrivateKey}} +{{- end}} +PublicKey = {{ .KeyPair.PublicKey }} +{{- if .PresharedKey}} +PresharedKey = {{ .PresharedKey }} +{{- end}} +{{- if eq $.Interface.Type "server"}} +AllowedIPs = {{ .AddressStr }}{{if ne .ExtraAllowedIPsStr ""}}, {{ .ExtraAllowedIPsStr }}{{end}} +{{- end}} +{{- if eq $.Interface.Type "client"}} +{{- if .AllowedIPsStr}} +AllowedIPs = {{ .AllowedIPsStr }} +{{- end}} +{{- end}} +{{- if and (ne .Endpoint "") (eq $.Interface.Type "client")}} +Endpoint = {{ .Endpoint }} +{{- end}} +{{- if ne .PersistentKeepalive 0}} +PersistentKeepalive = {{ .PersistentKeepalive }} +{{- end}} +{{- end}} +{{end}} \ No newline at end of file diff --git a/internal/wireguard/tpl_files/peer.tpl b/internal/wireguard/tpl_files/peer.tpl new file mode 100644 index 0000000..2278023 --- /dev/null +++ b/internal/wireguard/tpl_files/peer.tpl @@ -0,0 +1,60 @@ +# AUTOGENERATED FILE - DO NOT EDIT +# This file uses wg-quick format. See https://man7.org/linux/man-pages/man8/wg-quick.8.html#CONFIGURATION + +# -WGP- WIREGUARD PORTAL CONFIGURATION FILE, version {{ .Portal.Version }} +# Lines starting with the -WGP- tag are used by the WireGuard Portal configuration parser. + +[Interface] +# -WGP- Peer: {{.Peer.Uid}} | Updated: {{.Peer.UpdatedAt}} | Created: {{.Peer.CreatedAt}} +# -WGP- Display name: {{ .Peer.Identifier }} +# -WGP- PublicKey: {{ .Peer.KeyPair.PublicKey }} +{{- if eq $.Interface.Type "server"}} +# -WGP- Peer type: client +{{else}} +# -WGP- Peer type: server +{{- end}} + +# Core settings +PrivateKey = {{ .Peer.KeyPair.PrivateKey }} +Address = {{ .Peer.AddressStr.GetValue }} + +# Misc. settings (optional) +{{- if .Peer.DnsStr.GetValue}} +DNS = {{ .Peer.DnsStr.GetValue }} +{{- end}} +{{- if ne .Peer.Mtu.GetValue 0}} +MTU = {{ .Peer.Mtu.GetValue }} +{{- end}} +{{- if ne .Peer.FirewallMark.GetValue 0}} +FwMark = {{ .Peer.FirewallMark.GetValue }} +{{- end}} +{{- if ne .Peer.RoutingTable.GetValue ""}} +Table = {{ .Peer.RoutingTable.GetValue }} +{{- end}} + +# Interface hooks (optional) +{{- if .Peer.PreUp.GetValue}} +PreUp = {{ .Peer.PreUp.GetValue }} +{{- end}} +{{- if .Peer.PostUp.GetValue}} +PostUp = {{ .Peer.PostUp.GetValue }} +{{- end}} +{{- if .Peer.PreDown.GetValue}} +PreDown = {{ .Peer.PreDown.GetValue }} +{{- end}} +{{- if .Peer.PostDown.GetValue}} +PostDown = {{ .Peer.PostDown.GetValue }} +{{- end}} + +[Peer] +PublicKey = {{ .Interface.KeyPair.PublicKey }} +Endpoint = {{ .Peer.Endpoint.GetValue }} +{{- if .Peer.AllowedIPsStr.GetValue}} +AllowedIPs = {{ .Peer.AllowedIPsStr.GetValue }} +{{- end}} +{{- if .Peer.PresharedKey}} +PresharedKey = {{ .Peer.PresharedKey }} +{{- end}} +{{- if ne .Peer.PersistentKeepalive.GetValue 0}} +PersistentKeepalive = {{ .Peer.PersistentKeepalive.GetValue }} +{{- end}} \ No newline at end of file diff --git a/internal/wireguard/wrappers.go b/internal/wireguard/wrappers.go deleted file mode 100644 index 0ec5e47..0000000 --- a/internal/wireguard/wrappers.go +++ /dev/null @@ -1,55 +0,0 @@ -package wireguard - -import ( - "io" - - "github.com/vishvananda/netlink" - "golang.zx2c4.com/wireguard/wgctrl/wgtypes" -) - -// A Client is a type which can control a WireGuard device. -type Client interface { - io.Closer - Devices() ([]*wgtypes.Device, error) - Device(name string) (*wgtypes.Device, error) - ConfigureDevice(name string, cfg wgtypes.Config) error -} - -// A NetlinkClient is a type which can control a netlink device. -type NetlinkClient interface { - LinkAdd(link netlink.Link) error - LinkDel(link netlink.Link) error - LinkByName(name string) (netlink.Link, error) - LinkSetUp(link netlink.Link) error - LinkSetDown(link netlink.Link) error - LinkSetMTU(link netlink.Link, mtu int) error - AddrReplace(link netlink.Link, addr *netlink.Addr) error - AddrAdd(link netlink.Link, addr *netlink.Addr) error -} - -type NetlinkManager struct { -} - -func (n NetlinkManager) LinkAdd(link netlink.Link) error { return netlink.LinkAdd(link) } - -func (n NetlinkManager) LinkDel(link netlink.Link) error { return netlink.LinkDel(link) } - -func (n NetlinkManager) LinkByName(name string) (netlink.Link, error) { - return netlink.LinkByName(name) -} - -func (n NetlinkManager) LinkSetUp(link netlink.Link) error { return netlink.LinkSetUp(link) } - -func (n NetlinkManager) LinkSetDown(link netlink.Link) error { return netlink.LinkSetDown(link) } - -func (n NetlinkManager) LinkSetMTU(link netlink.Link, mtu int) error { - return netlink.LinkSetMTU(link, mtu) -} - -func (n NetlinkManager) AddrReplace(link netlink.Link, addr *netlink.Addr) error { - return netlink.AddrReplace(link, addr) -} - -func (n NetlinkManager) AddrAdd(link netlink.Link, addr *netlink.Addr) error { - return netlink.AddrAdd(link, addr) -}