1969 lines
60 KiB
Go
1969 lines
60 KiB
Go
package sqlitedb
|
|
|
|
import (
|
|
"fmt"
|
|
"net"
|
|
"os"
|
|
"path/filepath"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/stretchr/testify/assert"
|
|
"github.com/stretchr/testify/require"
|
|
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
|
|
|
|
"github.com/DigitalTolk/wireguard-ui/model"
|
|
"github.com/DigitalTolk/wireguard-ui/util"
|
|
)
|
|
|
|
func newTestDB(t *testing.T) *SqliteDB {
|
|
t.Helper()
|
|
dir := t.TempDir()
|
|
dbPath := filepath.Join(dir, "test.db")
|
|
db, err := New(dbPath)
|
|
require.NoError(t, err)
|
|
return db
|
|
}
|
|
|
|
func initTestDB(t *testing.T) *SqliteDB {
|
|
t.Helper()
|
|
// set env vars so Init doesn't try to detect public IP
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
db := newTestDB(t)
|
|
err := db.Init()
|
|
require.NoError(t, err)
|
|
return db
|
|
}
|
|
|
|
// --- User Tests ---
|
|
|
|
func TestSaveAndGetUser(t *testing.T) {
|
|
db := newTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
user := model.User{
|
|
Username: "testuser",
|
|
Email: "test@example.com",
|
|
Admin: true,
|
|
CreatedAt: now,
|
|
UpdatedAt: now,
|
|
}
|
|
|
|
err := db.SaveUser(user)
|
|
require.NoError(t, err)
|
|
|
|
got, err := db.GetUserByName("testuser")
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "testuser", got.Username)
|
|
assert.Equal(t, "test@example.com", got.Email)
|
|
assert.True(t, got.Admin)
|
|
}
|
|
|
|
func TestGetUsers(t *testing.T) {
|
|
db := newTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
db.SaveUser(model.User{Username: "user1", CreatedAt: now, UpdatedAt: now})
|
|
db.SaveUser(model.User{Username: "user2", CreatedAt: now, UpdatedAt: now})
|
|
|
|
users, err := db.GetUsers()
|
|
require.NoError(t, err)
|
|
assert.Len(t, users, 2)
|
|
}
|
|
|
|
func TestDeleteUser(t *testing.T) {
|
|
db := newTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
db.SaveUser(model.User{Username: "delme", CreatedAt: now, UpdatedAt: now})
|
|
|
|
err := db.DeleteUser("delme")
|
|
require.NoError(t, err)
|
|
|
|
_, err = db.GetUserByName("delme")
|
|
assert.Error(t, err)
|
|
}
|
|
|
|
func TestSaveUser_Upsert(t *testing.T) {
|
|
db := newTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
db.SaveUser(model.User{Username: "user1", Email: "old@test.com", CreatedAt: now, UpdatedAt: now})
|
|
db.SaveUser(model.User{Username: "user1", Email: "new@test.com", CreatedAt: now, UpdatedAt: now})
|
|
|
|
got, err := db.GetUserByName("user1")
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "new@test.com", got.Email)
|
|
}
|
|
|
|
func TestGetUserByName_NotFound(t *testing.T) {
|
|
db := newTestDB(t)
|
|
_, err := db.GetUserByName("nonexistent")
|
|
assert.Error(t, err)
|
|
}
|
|
|
|
func TestSaveUser_ZeroTimestamps(t *testing.T) {
|
|
db := newTestDB(t)
|
|
|
|
// Save a user without setting timestamps - they should be auto-filled
|
|
user := model.User{
|
|
Username: "autotime",
|
|
Email: "auto@test.com",
|
|
Admin: false,
|
|
}
|
|
|
|
err := db.SaveUser(user)
|
|
require.NoError(t, err)
|
|
|
|
got, err := db.GetUserByName("autotime")
|
|
require.NoError(t, err)
|
|
assert.False(t, got.CreatedAt.IsZero())
|
|
assert.False(t, got.UpdatedAt.IsZero())
|
|
}
|
|
|
|
func TestSaveUser_WithOIDCSub(t *testing.T) {
|
|
db := newTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
user := model.User{
|
|
Username: "oidcuser",
|
|
Email: "oidc@test.com",
|
|
OIDCSub: "sub-12345",
|
|
Admin: false,
|
|
CreatedAt: now,
|
|
UpdatedAt: now,
|
|
}
|
|
|
|
err := db.SaveUser(user)
|
|
require.NoError(t, err)
|
|
|
|
got, err := db.GetUserByOIDCSub("sub-12345")
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "oidcuser", got.Username)
|
|
}
|
|
|
|
func TestSaveUser_WithDisplayName(t *testing.T) {
|
|
db := newTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
user := model.User{
|
|
Username: "displayuser",
|
|
Email: "display@test.com",
|
|
DisplayName: "Display Name",
|
|
CreatedAt: now,
|
|
UpdatedAt: now,
|
|
}
|
|
|
|
err := db.SaveUser(user)
|
|
require.NoError(t, err)
|
|
|
|
got, err := db.GetUserByName("displayuser")
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "Display Name", got.DisplayName)
|
|
}
|
|
|
|
// --- Client Tests ---
|
|
|
|
func TestSaveAndGetClient(t *testing.T) {
|
|
db := newTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
client := model.Client{
|
|
ID: "test123",
|
|
Name: "Test Client",
|
|
Email: "client@test.com",
|
|
PublicKey: "pubkey123",
|
|
PrivateKey: "privkey123",
|
|
PresharedKey: "psk123",
|
|
AllocatedIPs: []string{"10.0.0.2/32"},
|
|
AllowedIPs: []string{"0.0.0.0/0"},
|
|
ExtraAllowedIPs: []string{},
|
|
SubnetRanges: []string{},
|
|
UseServerDNS: true,
|
|
Enabled: true,
|
|
CreatedAt: now,
|
|
UpdatedAt: now,
|
|
}
|
|
|
|
err := db.SaveClient(client)
|
|
require.NoError(t, err)
|
|
|
|
got, err := db.GetClientByID("test123", model.QRCodeSettings{Enabled: false})
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "Test Client", got.Client.Name)
|
|
assert.Equal(t, "client@test.com", got.Client.Email)
|
|
assert.Equal(t, []string{"10.0.0.2/32"}, got.Client.AllocatedIPs)
|
|
assert.True(t, got.Client.Enabled)
|
|
}
|
|
|
|
func TestGetClients(t *testing.T) {
|
|
db := initTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
db.SaveClient(model.Client{ID: "c1", Name: "Client 1", AllocatedIPs: []string{}, AllowedIPs: []string{}, ExtraAllowedIPs: []string{}, SubnetRanges: []string{}, CreatedAt: now, UpdatedAt: now})
|
|
db.SaveClient(model.Client{ID: "c2", Name: "Client 2", AllocatedIPs: []string{}, AllowedIPs: []string{}, ExtraAllowedIPs: []string{}, SubnetRanges: []string{}, CreatedAt: now, UpdatedAt: now})
|
|
|
|
clients, err := db.GetClients(false)
|
|
require.NoError(t, err)
|
|
assert.Len(t, clients, 2)
|
|
}
|
|
|
|
func TestDeleteClient(t *testing.T) {
|
|
db := newTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
db.SaveClient(model.Client{ID: "delclient", Name: "Del", AllocatedIPs: []string{}, AllowedIPs: []string{}, ExtraAllowedIPs: []string{}, SubnetRanges: []string{}, CreatedAt: now, UpdatedAt: now})
|
|
|
|
err := db.DeleteClient("delclient")
|
|
require.NoError(t, err)
|
|
|
|
_, err = db.GetClientByID("delclient", model.QRCodeSettings{Enabled: false})
|
|
assert.Error(t, err)
|
|
}
|
|
|
|
func TestSaveClient_Upsert(t *testing.T) {
|
|
db := newTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
db.SaveClient(model.Client{ID: "c1", Name: "Original", AllocatedIPs: []string{}, AllowedIPs: []string{}, ExtraAllowedIPs: []string{}, SubnetRanges: []string{}, CreatedAt: now, UpdatedAt: now})
|
|
db.SaveClient(model.Client{ID: "c1", Name: "Updated", AllocatedIPs: []string{}, AllowedIPs: []string{}, ExtraAllowedIPs: []string{}, SubnetRanges: []string{}, CreatedAt: now, UpdatedAt: now})
|
|
|
|
got, err := db.GetClientByID("c1", model.QRCodeSettings{Enabled: false})
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "Updated", got.Client.Name)
|
|
}
|
|
|
|
// --- Server Tests ---
|
|
|
|
func TestGetServer(t *testing.T) {
|
|
db := initTestDB(t)
|
|
|
|
server, err := db.GetServer()
|
|
require.NoError(t, err)
|
|
assert.NotNil(t, server.Interface)
|
|
assert.NotNil(t, server.KeyPair)
|
|
assert.NotEmpty(t, server.KeyPair.PublicKey)
|
|
assert.NotEmpty(t, server.KeyPair.PrivateKey)
|
|
assert.Greater(t, server.Interface.ListenPort, 0)
|
|
}
|
|
|
|
func TestSaveServerInterface(t *testing.T) {
|
|
db := initTestDB(t)
|
|
|
|
iface := model.ServerInterface{
|
|
Addresses: []string{"10.0.0.0/24", "fd00::1/64"},
|
|
ListenPort: 12345,
|
|
PostUp: "iptables -A",
|
|
PostDown: "iptables -D",
|
|
UpdatedAt: time.Now().UTC(),
|
|
}
|
|
|
|
err := db.SaveServerInterface(iface)
|
|
require.NoError(t, err)
|
|
|
|
server, err := db.GetServer()
|
|
require.NoError(t, err)
|
|
assert.Equal(t, []string{"10.0.0.0/24", "fd00::1/64"}, server.Interface.Addresses)
|
|
assert.Equal(t, 12345, server.Interface.ListenPort)
|
|
assert.Equal(t, "iptables -A", server.Interface.PostUp)
|
|
}
|
|
|
|
func TestSaveServerKeyPair(t *testing.T) {
|
|
db := initTestDB(t)
|
|
|
|
kp := model.ServerKeypair{
|
|
PrivateKey: "newpriv",
|
|
PublicKey: "newpub",
|
|
UpdatedAt: time.Now().UTC(),
|
|
}
|
|
|
|
err := db.SaveServerKeyPair(kp)
|
|
require.NoError(t, err)
|
|
|
|
server, err := db.GetServer()
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "newpub", server.KeyPair.PublicKey)
|
|
assert.Equal(t, "newpriv", server.KeyPair.PrivateKey)
|
|
}
|
|
|
|
func TestSaveServerInterface_WithPreDown(t *testing.T) {
|
|
db := initTestDB(t)
|
|
|
|
iface := model.ServerInterface{
|
|
Addresses: []string{"10.0.0.0/24"},
|
|
ListenPort: 51820,
|
|
PostUp: "iptables -A FORWARD -i wg0 -j ACCEPT",
|
|
PreDown: "iptables -D FORWARD -i wg0 -j ACCEPT",
|
|
PostDown: "iptables -D FORWARD -i wg0 -j ACCEPT",
|
|
UpdatedAt: time.Now().UTC(),
|
|
}
|
|
|
|
err := db.SaveServerInterface(iface)
|
|
require.NoError(t, err)
|
|
|
|
server, err := db.GetServer()
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "iptables -D FORWARD -i wg0 -j ACCEPT", server.Interface.PreDown)
|
|
}
|
|
|
|
// --- Global Settings Tests ---
|
|
|
|
func TestGetGlobalSettings(t *testing.T) {
|
|
db := initTestDB(t)
|
|
|
|
gs, err := db.GetGlobalSettings()
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "10.0.0.1", gs.EndpointAddress)
|
|
assert.NotEmpty(t, gs.DNSServers)
|
|
assert.Greater(t, gs.MTU, 0)
|
|
}
|
|
|
|
func TestSaveGlobalSettings(t *testing.T) {
|
|
db := initTestDB(t)
|
|
|
|
gs := model.GlobalSetting{
|
|
EndpointAddress: "vpn.example.com",
|
|
DNSServers: []string{"8.8.8.8", "8.8.4.4"},
|
|
MTU: 1400,
|
|
PersistentKeepalive: 25,
|
|
FirewallMark: "0x1234",
|
|
Table: "auto",
|
|
ConfigFilePath: "/etc/wireguard/wg0.conf",
|
|
UpdatedAt: time.Now().UTC(),
|
|
}
|
|
|
|
err := db.SaveGlobalSettings(gs)
|
|
require.NoError(t, err)
|
|
|
|
got, err := db.GetGlobalSettings()
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "vpn.example.com", got.EndpointAddress)
|
|
assert.Equal(t, []string{"8.8.8.8", "8.8.4.4"}, got.DNSServers)
|
|
assert.Equal(t, 1400, got.MTU)
|
|
assert.Equal(t, 25, got.PersistentKeepalive)
|
|
}
|
|
|
|
// --- Hashes Tests ---
|
|
|
|
func TestSaveAndGetHashes(t *testing.T) {
|
|
db := initTestDB(t)
|
|
|
|
hashes := model.ClientServerHashes{
|
|
Client: "abc123",
|
|
Server: "def456",
|
|
}
|
|
|
|
err := db.SaveHashes(hashes)
|
|
require.NoError(t, err)
|
|
|
|
got, err := db.GetHashes()
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "abc123", got.Client)
|
|
assert.Equal(t, "def456", got.Server)
|
|
}
|
|
|
|
// --- AllocatedIPs Tests ---
|
|
|
|
func TestGetAllocatedIPs(t *testing.T) {
|
|
db := initTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
// save a client with allocated IPs
|
|
db.SaveClient(model.Client{
|
|
ID: "c1",
|
|
Name: "Client1",
|
|
AllocatedIPs: []string{"10.252.1.2/32"},
|
|
AllowedIPs: []string{"0.0.0.0/0"},
|
|
ExtraAllowedIPs: []string{},
|
|
SubnetRanges: []string{},
|
|
CreatedAt: now,
|
|
UpdatedAt: now,
|
|
})
|
|
|
|
ips, err := db.GetAllocatedIPs("")
|
|
require.NoError(t, err)
|
|
// should include server address + client address
|
|
assert.Contains(t, ips, "10.252.1.2")
|
|
}
|
|
|
|
func TestGetAllocatedIPs_ExcludeClient(t *testing.T) {
|
|
db := initTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
db.SaveClient(model.Client{
|
|
ID: "c1", Name: "Client A", AllocatedIPs: []string{"10.252.1.2/32"},
|
|
AllowedIPs: []string{}, ExtraAllowedIPs: []string{}, SubnetRanges: []string{},
|
|
CreatedAt: now, UpdatedAt: now,
|
|
})
|
|
db.SaveClient(model.Client{
|
|
ID: "c2", Name: "Client B", AllocatedIPs: []string{"10.252.1.3/32"},
|
|
AllowedIPs: []string{}, ExtraAllowedIPs: []string{}, SubnetRanges: []string{},
|
|
CreatedAt: now, UpdatedAt: now,
|
|
})
|
|
|
|
ips, err := db.GetAllocatedIPs("c1")
|
|
require.NoError(t, err)
|
|
assert.NotContains(t, ips, "10.252.1.2")
|
|
assert.Contains(t, ips, "10.252.1.3")
|
|
}
|
|
|
|
// --- GetPath ---
|
|
|
|
func TestGetPath(t *testing.T) {
|
|
db := initTestDB(t)
|
|
path := db.GetPath()
|
|
assert.NotEmpty(t, path)
|
|
}
|
|
|
|
// --- Init Tests ---
|
|
|
|
func TestInit_CreatesDefaults(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
db := newTestDB(t)
|
|
err := db.Init()
|
|
require.NoError(t, err)
|
|
|
|
// no default user — first OIDC login creates admin
|
|
users, err := db.GetUsers()
|
|
require.NoError(t, err)
|
|
assert.Len(t, users, 0)
|
|
|
|
// should have created server config
|
|
server, err := db.GetServer()
|
|
require.NoError(t, err)
|
|
assert.NotEmpty(t, server.KeyPair.PublicKey)
|
|
|
|
// should have created global settings
|
|
gs, err := db.GetGlobalSettings()
|
|
require.NoError(t, err)
|
|
assert.NotEmpty(t, gs.EndpointAddress)
|
|
|
|
// should have created hashes
|
|
h, err := db.GetHashes()
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "none", h.Client)
|
|
}
|
|
|
|
func TestInit_Idempotent(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
db := newTestDB(t)
|
|
require.NoError(t, db.Init())
|
|
require.NoError(t, db.Init()) // second call should not error
|
|
}
|
|
|
|
// --- QR Code Tests ---
|
|
|
|
func TestGetClients_WithQRCode(t *testing.T) {
|
|
db := initTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
db.SaveClient(model.Client{
|
|
ID: "qr1",
|
|
Name: "QR Client",
|
|
PrivateKey: "privkey123",
|
|
PublicKey: "pubkey123",
|
|
AllocatedIPs: []string{"10.252.1.10/32"},
|
|
AllowedIPs: []string{"0.0.0.0/0"},
|
|
ExtraAllowedIPs: []string{},
|
|
SubnetRanges: []string{},
|
|
Enabled: true,
|
|
CreatedAt: now,
|
|
UpdatedAt: now,
|
|
})
|
|
|
|
clients, err := db.GetClients(true)
|
|
require.NoError(t, err)
|
|
require.Len(t, clients, 1)
|
|
// Client with a private key should get a QR code when hasQRCode=true
|
|
assert.NotEmpty(t, clients[0].QRCode)
|
|
assert.Contains(t, clients[0].QRCode, "data:image/png;base64,")
|
|
}
|
|
|
|
func TestGetClients_WithQRCode_NoPrivateKey(t *testing.T) {
|
|
db := initTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
db.SaveClient(model.Client{
|
|
ID: "noqr1",
|
|
Name: "No QR Client",
|
|
PrivateKey: "", // empty private key
|
|
PublicKey: "pubkey123",
|
|
AllocatedIPs: []string{"10.252.1.10/32"},
|
|
AllowedIPs: []string{"0.0.0.0/0"},
|
|
ExtraAllowedIPs: []string{},
|
|
SubnetRanges: []string{},
|
|
Enabled: true,
|
|
CreatedAt: now,
|
|
UpdatedAt: now,
|
|
})
|
|
|
|
clients, err := db.GetClients(true)
|
|
require.NoError(t, err)
|
|
require.Len(t, clients, 1)
|
|
// Client without private key should have no QR code
|
|
assert.Empty(t, clients[0].QRCode)
|
|
}
|
|
|
|
func TestGetClientByID_WithQRCode(t *testing.T) {
|
|
db := initTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
db.SaveClient(model.Client{
|
|
ID: "qr2",
|
|
Name: "QR Client 2",
|
|
PrivateKey: "privkey456",
|
|
PublicKey: "pubkey456",
|
|
AllocatedIPs: []string{"10.252.1.11/32"},
|
|
AllowedIPs: []string{"0.0.0.0/0"},
|
|
ExtraAllowedIPs: []string{},
|
|
SubnetRanges: []string{},
|
|
Enabled: true,
|
|
CreatedAt: now,
|
|
UpdatedAt: now,
|
|
})
|
|
|
|
// With QR enabled
|
|
clientData, err := db.GetClientByID("qr2", model.QRCodeSettings{
|
|
Enabled: true,
|
|
IncludeDNS: true,
|
|
IncludeMTU: true,
|
|
})
|
|
require.NoError(t, err)
|
|
assert.NotEmpty(t, clientData.QRCode)
|
|
|
|
// With QR disabled
|
|
clientData2, err := db.GetClientByID("qr2", model.QRCodeSettings{
|
|
Enabled: false,
|
|
})
|
|
require.NoError(t, err)
|
|
assert.Empty(t, clientData2.QRCode)
|
|
}
|
|
|
|
func TestGetClientByID_QRCode_NoDNS_NoMTU(t *testing.T) {
|
|
db := initTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
db.SaveClient(model.Client{
|
|
ID: "qr3",
|
|
Name: "QR Client 3",
|
|
PrivateKey: "privkey789",
|
|
PublicKey: "pubkey789",
|
|
AllocatedIPs: []string{"10.252.1.12/32"},
|
|
AllowedIPs: []string{"0.0.0.0/0"},
|
|
ExtraAllowedIPs: []string{},
|
|
SubnetRanges: []string{},
|
|
Enabled: true,
|
|
CreatedAt: now,
|
|
UpdatedAt: now,
|
|
})
|
|
|
|
// With QR enabled but DNS and MTU excluded
|
|
clientData, err := db.GetClientByID("qr3", model.QRCodeSettings{
|
|
Enabled: true,
|
|
IncludeDNS: false,
|
|
IncludeMTU: false,
|
|
})
|
|
require.NoError(t, err)
|
|
assert.NotEmpty(t, clientData.QRCode)
|
|
}
|
|
|
|
func TestGetClientByID_NotFound(t *testing.T) {
|
|
db := newTestDB(t)
|
|
|
|
_, err := db.GetClientByID("nonexistent", model.QRCodeSettings{Enabled: false})
|
|
assert.Error(t, err)
|
|
}
|
|
|
|
// --- DB method ---
|
|
|
|
func TestDB(t *testing.T) {
|
|
db := newTestDB(t)
|
|
sqlDB := db.DB()
|
|
assert.NotNil(t, sqlDB)
|
|
|
|
// verify it can execute a query
|
|
var n int
|
|
err := sqlDB.QueryRow("SELECT 1").Scan(&n)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, 1, n)
|
|
}
|
|
|
|
// --- Init with custom env vars ---
|
|
|
|
func TestInit_CustomEnvVars(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "vpn.test.com")
|
|
os.Setenv("WGUI_DNS", "8.8.8.8,8.8.4.4")
|
|
os.Setenv("WGUI_MTU", "1400")
|
|
os.Setenv("WGUI_PERSISTENT_KEEPALIVE", "30")
|
|
os.Setenv("WGUI_SERVER_INTERFACE_ADDRESSES", "10.10.0.0/24")
|
|
os.Setenv("WGUI_SERVER_LISTEN_PORT", "12345")
|
|
defer func() {
|
|
os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
os.Unsetenv("WGUI_DNS")
|
|
os.Unsetenv("WGUI_MTU")
|
|
os.Unsetenv("WGUI_PERSISTENT_KEEPALIVE")
|
|
os.Unsetenv("WGUI_SERVER_INTERFACE_ADDRESSES")
|
|
os.Unsetenv("WGUI_SERVER_LISTEN_PORT")
|
|
}()
|
|
|
|
db := newTestDB(t)
|
|
err := db.Init()
|
|
require.NoError(t, err)
|
|
|
|
gs, err := db.GetGlobalSettings()
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "vpn.test.com", gs.EndpointAddress)
|
|
assert.Equal(t, []string{"8.8.8.8", "8.8.4.4"}, gs.DNSServers)
|
|
assert.Equal(t, 1400, gs.MTU)
|
|
assert.Equal(t, 30, gs.PersistentKeepalive)
|
|
|
|
server, err := db.GetServer()
|
|
require.NoError(t, err)
|
|
assert.Equal(t, []string{"10.10.0.0/24"}, server.Interface.Addresses)
|
|
assert.Equal(t, 12345, server.Interface.ListenPort)
|
|
}
|
|
|
|
// --- Hash operations via store ---
|
|
|
|
func TestGetCurrentHash_ViaStore(t *testing.T) {
|
|
db := initTestDB(t)
|
|
|
|
clientHash, serverHash := util.GetCurrentHash(db)
|
|
assert.NotEmpty(t, clientHash)
|
|
assert.NotEmpty(t, serverHash)
|
|
assert.NotEqual(t, "error", clientHash)
|
|
assert.NotEqual(t, "error", serverHash)
|
|
}
|
|
|
|
func TestHashesChanged_ViaStore(t *testing.T) {
|
|
db := initTestDB(t)
|
|
|
|
// Initially hashes should differ (db has 'none', computed is real)
|
|
changed := util.HashesChanged(db)
|
|
assert.True(t, changed)
|
|
|
|
// After updating, they should match
|
|
err := util.UpdateHashes(db)
|
|
require.NoError(t, err)
|
|
|
|
changed = util.HashesChanged(db)
|
|
assert.False(t, changed)
|
|
}
|
|
|
|
func TestUpdateHashes_ThenChange(t *testing.T) {
|
|
db := initTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
// Set hashes
|
|
err := util.UpdateHashes(db)
|
|
require.NoError(t, err)
|
|
assert.False(t, util.HashesChanged(db))
|
|
|
|
// Add a client to change the hash
|
|
db.SaveClient(model.Client{
|
|
ID: "hashclient", Name: "Hash Client",
|
|
AllocatedIPs: []string{"10.252.1.50/32"},
|
|
AllowedIPs: []string{"0.0.0.0/0"},
|
|
ExtraAllowedIPs: []string{}, SubnetRanges: []string{},
|
|
CreatedAt: now, UpdatedAt: now,
|
|
})
|
|
|
|
// Now hashes should differ
|
|
assert.True(t, util.HashesChanged(db))
|
|
}
|
|
|
|
// --- ValidateAndFixSubnetRanges ---
|
|
|
|
func TestValidateAndFixSubnetRanges(t *testing.T) {
|
|
db := initTestDB(t)
|
|
|
|
// Set up subnet ranges - the server has 10.252.1.0/24
|
|
util.SubnetRanges = map[string][]*net.IPNet{}
|
|
util.SubnetRangesOrder = nil
|
|
util.IPToSubnetRange = map[string]uint16{}
|
|
|
|
util.SubnetRanges = util.ParseSubnetRanges("valid:10.252.1.0/26;invalid:192.168.99.0/24")
|
|
|
|
err := util.ValidateAndFixSubnetRanges(db)
|
|
require.NoError(t, err)
|
|
|
|
// valid range should remain
|
|
assert.NotNil(t, util.SubnetRanges["valid"])
|
|
// invalid range should be removed (192.168.99.0/24 is outside 10.252.1.0/24)
|
|
_, hasInvalid := util.SubnetRanges["invalid"]
|
|
assert.False(t, hasInvalid)
|
|
}
|
|
|
|
func TestValidateAndFixSubnetRanges_Empty(t *testing.T) {
|
|
db := initTestDB(t)
|
|
|
|
util.SubnetRangesOrder = nil
|
|
util.SubnetRanges = map[string][]*net.IPNet{}
|
|
|
|
err := util.ValidateAndFixSubnetRanges(db)
|
|
require.NoError(t, err)
|
|
}
|
|
|
|
// --- Migration from JSON ---
|
|
|
|
func TestMigrateFromJSON_NoOldDB(t *testing.T) {
|
|
db := newTestDB(t)
|
|
dir := t.TempDir()
|
|
|
|
// No old JSON DB directory - should return nil
|
|
err := MigrateFromJSON(db, dir)
|
|
assert.NoError(t, err)
|
|
}
|
|
|
|
func TestMigrateFromJSON_FullMigration(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
db := newTestDB(t)
|
|
require.NoError(t, db.Init())
|
|
|
|
// Create a fake JSON database structure
|
|
jsonDBPath := t.TempDir()
|
|
|
|
// Create server directory
|
|
serverDir := filepath.Join(jsonDBPath, "server")
|
|
require.NoError(t, os.MkdirAll(serverDir, 0755))
|
|
|
|
// Create interfaces.json (ListenPort uses json:",string" tag)
|
|
ifaceJSON := `{"addresses":["10.0.0.1/24"],"listen_port":"51820","post_up":"iptables up","post_down":"iptables down"}`
|
|
require.NoError(t, os.WriteFile(filepath.Join(serverDir, "interfaces.json"), []byte(ifaceJSON), 0644))
|
|
|
|
// Create keypair.json
|
|
keypairJSON := `{"private_key":"migprivkey","public_key":"migpubkey"}`
|
|
require.NoError(t, os.WriteFile(filepath.Join(serverDir, "keypair.json"), []byte(keypairJSON), 0644))
|
|
|
|
// Create global_settings.json (MTU and PersistentKeepalive use json:",string" tags)
|
|
gsJSON := `{"endpoint_address":"migvpn.test.com","dns_servers":["8.8.8.8"],"mtu":"1400","persistent_keepalive":"25","firewall_mark":"0xca6c","table":"auto","config_file_path":"/etc/wireguard/wg0.conf"}`
|
|
require.NoError(t, os.WriteFile(filepath.Join(serverDir, "global_settings.json"), []byte(gsJSON), 0644))
|
|
|
|
// Create hashes.json
|
|
hashesJSON := `{"client":"abc","server":"def"}`
|
|
require.NoError(t, os.WriteFile(filepath.Join(serverDir, "hashes.json"), []byte(hashesJSON), 0644))
|
|
|
|
// Create users directory
|
|
usersDir := filepath.Join(jsonDBPath, "users")
|
|
require.NoError(t, os.MkdirAll(usersDir, 0755))
|
|
userJSON := `{"username":"miguser","email":"mig@test.com","admin":true}`
|
|
require.NoError(t, os.WriteFile(filepath.Join(usersDir, "miguser.json"), []byte(userJSON), 0644))
|
|
|
|
// Create clients directory
|
|
clientsDir := filepath.Join(jsonDBPath, "clients")
|
|
require.NoError(t, os.MkdirAll(clientsDir, 0755))
|
|
clientJSON := `{"id":"migclient","name":"Mig Client","public_key":"migclientpub","allocated_ips":["10.0.0.2/32"],"allowed_ips":["0.0.0.0/0"],"extra_allowed_ips":[],"subnet_ranges":[],"enabled":true}`
|
|
require.NoError(t, os.WriteFile(filepath.Join(clientsDir, "migclient.json"), []byte(clientJSON), 0644))
|
|
|
|
// Create wake_on_lan_hosts directory
|
|
wolDir := filepath.Join(jsonDBPath, "wake_on_lan_hosts")
|
|
require.NoError(t, os.MkdirAll(wolDir, 0755))
|
|
wolJSON := `{"MacAddress":"AA:BB:CC:DD:EE:FF","Name":"Test WOL"}`
|
|
require.NoError(t, os.WriteFile(filepath.Join(wolDir, "host1.json"), []byte(wolJSON), 0644))
|
|
|
|
// Run migration
|
|
err := MigrateFromJSON(db, jsonDBPath)
|
|
require.NoError(t, err)
|
|
|
|
// Verify data was migrated
|
|
server, err := db.GetServer()
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "migpubkey", server.KeyPair.PublicKey)
|
|
assert.Equal(t, "migprivkey", server.KeyPair.PrivateKey)
|
|
|
|
gs, err := db.GetGlobalSettings()
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "migvpn.test.com", gs.EndpointAddress)
|
|
|
|
hashes, err := db.GetHashes()
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "abc", hashes.Client)
|
|
assert.Equal(t, "def", hashes.Server)
|
|
|
|
// Users are no longer migrated (SSO-only auth makes legacy users useless)
|
|
_, err = db.GetUserByName("miguser")
|
|
assert.Error(t, err, "Legacy users should not be migrated")
|
|
|
|
clientData, err := db.GetClientByID("migclient", model.QRCodeSettings{Enabled: false})
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "Mig Client", clientData.Client.Name)
|
|
|
|
// Verify old directory was renamed
|
|
_, err = os.Stat(jsonDBPath + ".json.bak")
|
|
assert.NoError(t, err)
|
|
}
|
|
|
|
func TestMigrateFromJSON_MissingOptionalDirs(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
db := newTestDB(t)
|
|
require.NoError(t, db.Init())
|
|
|
|
// Create only the server directory (no users, clients, or wol)
|
|
jsonDBPath := t.TempDir()
|
|
serverDir := filepath.Join(jsonDBPath, "server")
|
|
require.NoError(t, os.MkdirAll(serverDir, 0755))
|
|
|
|
// Run migration - should succeed even without optional directories
|
|
err := MigrateFromJSON(db, jsonDBPath)
|
|
require.NoError(t, err)
|
|
}
|
|
|
|
func TestMigrateFromJSON_InvalidJSON(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
db := newTestDB(t)
|
|
require.NoError(t, db.Init())
|
|
|
|
jsonDBPath := t.TempDir()
|
|
serverDir := filepath.Join(jsonDBPath, "server")
|
|
require.NoError(t, os.MkdirAll(serverDir, 0755))
|
|
|
|
// Write invalid JSON for interfaces
|
|
require.NoError(t, os.WriteFile(filepath.Join(serverDir, "interfaces.json"), []byte("{invalid"), 0644))
|
|
|
|
err := MigrateFromJSON(db, jsonDBPath)
|
|
assert.Error(t, err)
|
|
}
|
|
|
|
// --- Client with AdditionalNotes ---
|
|
|
|
func TestSaveClient_WithAdditionalNotes(t *testing.T) {
|
|
db := newTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
client := model.Client{
|
|
ID: "notes-c1", Name: "Notes Client",
|
|
AdditionalNotes: "This client has notes\nMultiple lines",
|
|
AllocatedIPs: []string{"10.0.0.5/32"},
|
|
AllowedIPs: []string{"0.0.0.0/0"},
|
|
ExtraAllowedIPs: []string{}, SubnetRanges: []string{},
|
|
Enabled: true, CreatedAt: now, UpdatedAt: now,
|
|
}
|
|
err := db.SaveClient(client)
|
|
require.NoError(t, err)
|
|
|
|
got, err := db.GetClientByID("notes-c1", model.QRCodeSettings{Enabled: false})
|
|
require.NoError(t, err)
|
|
assert.Contains(t, got.Client.AdditionalNotes, "Multiple lines")
|
|
}
|
|
|
|
func TestSaveClient_WithEndpoint(t *testing.T) {
|
|
db := newTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
client := model.Client{
|
|
ID: "ep-c1", Name: "Endpoint Client",
|
|
Endpoint: "vpn.example.com:51820",
|
|
AllocatedIPs: []string{"10.0.0.6/32"},
|
|
AllowedIPs: []string{"0.0.0.0/0"},
|
|
ExtraAllowedIPs: []string{}, SubnetRanges: []string{},
|
|
Enabled: true, CreatedAt: now, UpdatedAt: now,
|
|
}
|
|
err := db.SaveClient(client)
|
|
require.NoError(t, err)
|
|
|
|
got, err := db.GetClientByID("ep-c1", model.QRCodeSettings{Enabled: false})
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "vpn.example.com:51820", got.Client.Endpoint)
|
|
}
|
|
|
|
// --- GetUserByOIDCSub Tests ---
|
|
|
|
func TestGetUserByOIDCSub_Found(t *testing.T) {
|
|
db := newTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
db.SaveUser(model.User{
|
|
Username: "oidcuser",
|
|
Email: "oidc@test.com",
|
|
OIDCSub: "sub-12345",
|
|
Admin: true,
|
|
CreatedAt: now,
|
|
UpdatedAt: now,
|
|
})
|
|
|
|
user, err := db.GetUserByOIDCSub("sub-12345")
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "oidcuser", user.Username)
|
|
assert.Equal(t, "oidc@test.com", user.Email)
|
|
assert.Equal(t, "sub-12345", user.OIDCSub)
|
|
assert.True(t, user.Admin)
|
|
}
|
|
|
|
func TestGetUserByOIDCSub_NotFound(t *testing.T) {
|
|
db := newTestDB(t)
|
|
|
|
_, err := db.GetUserByOIDCSub("nonexistent-sub")
|
|
assert.Error(t, err)
|
|
}
|
|
|
|
// --- GetClients with search by notes ---
|
|
|
|
func TestGetClients_SearchByNotes(t *testing.T) {
|
|
db := initTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
db.SaveClient(model.Client{
|
|
ID: "notes1", Name: "Client With Notes",
|
|
AdditionalNotes: "special deployment note",
|
|
AllocatedIPs: []string{}, AllowedIPs: []string{},
|
|
ExtraAllowedIPs: []string{}, SubnetRanges: []string{},
|
|
CreatedAt: now, UpdatedAt: now,
|
|
})
|
|
db.SaveClient(model.Client{
|
|
ID: "notes2", Name: "Other Client",
|
|
AdditionalNotes: "regular note",
|
|
AllocatedIPs: []string{}, AllowedIPs: []string{},
|
|
ExtraAllowedIPs: []string{}, SubnetRanges: []string{},
|
|
CreatedAt: now, UpdatedAt: now,
|
|
})
|
|
|
|
clients, err := db.GetClients(false)
|
|
require.NoError(t, err)
|
|
assert.Len(t, clients, 2)
|
|
}
|
|
|
|
// --- migrate: public key derivation from private key ---
|
|
|
|
func TestMigrate_DerivesPublicKeyFromPrivateKey(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
db := newTestDB(t)
|
|
require.NoError(t, db.Init())
|
|
|
|
// Insert a client with a valid private key but empty public key
|
|
key, err := wgtypes.GeneratePrivateKey()
|
|
require.NoError(t, err)
|
|
now := time.Now().UTC()
|
|
|
|
_, err = db.db.Exec(
|
|
`INSERT INTO clients (id, private_key, public_key, preshared_key, name, email,
|
|
subnet_ranges, allocated_ips, allowed_ips, extra_allowed_ips,
|
|
endpoint, additional_notes, use_server_dns, enabled, created_at, updated_at)
|
|
VALUES (?, ?, '', '', 'Derive Test', 'derive@test.com', '[]', '[]', '[]', '[]', '', '', 0, 1, ?, ?)`,
|
|
"derive-test", key.String(), now, now,
|
|
)
|
|
require.NoError(t, err)
|
|
|
|
// Run migrate again to trigger public key derivation
|
|
db.migrate()
|
|
|
|
// Verify public key was derived
|
|
clientData, err := db.GetClientByID("derive-test", model.QRCodeSettings{Enabled: false})
|
|
require.NoError(t, err)
|
|
assert.Equal(t, key.PublicKey().String(), clientData.Client.PublicKey)
|
|
}
|
|
|
|
func TestMigrate_SkipsInvalidPrivateKey(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
db := newTestDB(t)
|
|
require.NoError(t, db.Init())
|
|
|
|
now := time.Now().UTC()
|
|
|
|
// Insert a client with an invalid private key and empty public key
|
|
_, err := db.db.Exec(
|
|
`INSERT INTO clients (id, private_key, public_key, preshared_key, name, email,
|
|
subnet_ranges, allocated_ips, allowed_ips, extra_allowed_ips,
|
|
endpoint, additional_notes, use_server_dns, enabled, created_at, updated_at)
|
|
VALUES (?, 'invalid-key', '', '', 'Bad Key', 'bad@test.com', '[]', '[]', '[]', '[]', '', '', 0, 1, ?, ?)`,
|
|
"bad-key-test", now, now,
|
|
)
|
|
require.NoError(t, err)
|
|
|
|
// Run migrate - should not error, just skip
|
|
db.migrate()
|
|
|
|
// Public key should remain empty
|
|
clientData, err := db.GetClientByID("bad-key-test", model.QRCodeSettings{Enabled: false})
|
|
require.NoError(t, err)
|
|
assert.Empty(t, clientData.Client.PublicKey)
|
|
}
|
|
|
|
// --- New error path test ---
|
|
|
|
func TestNew_InvalidPath(t *testing.T) {
|
|
// Try to create a DB in a path that can't exist (null byte in path)
|
|
_, err := New("/dev/null/impossible/path/db.sqlite")
|
|
assert.Error(t, err)
|
|
}
|
|
|
|
// --- GetAllocatedIPs with multiple client IPs ---
|
|
|
|
func TestGetAllocatedIPs_MultipleClientIPs(t *testing.T) {
|
|
db := initTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
db.SaveClient(model.Client{
|
|
ID: "multi-ip", Name: "Multi IP",
|
|
AllocatedIPs: []string{"10.252.1.10/32", "10.252.1.11/32"},
|
|
AllowedIPs: []string{"0.0.0.0/0"},
|
|
ExtraAllowedIPs: []string{}, SubnetRanges: []string{},
|
|
CreatedAt: now, UpdatedAt: now,
|
|
})
|
|
|
|
ips, err := db.GetAllocatedIPs("")
|
|
require.NoError(t, err)
|
|
assert.Contains(t, ips, "10.252.1.10")
|
|
assert.Contains(t, ips, "10.252.1.11")
|
|
}
|
|
|
|
func TestGetAllocatedIPs_NoClients(t *testing.T) {
|
|
db := initTestDB(t)
|
|
|
|
ips, err := db.GetAllocatedIPs("")
|
|
require.NoError(t, err)
|
|
// Should have at least server address
|
|
assert.NotEmpty(t, ips)
|
|
}
|
|
|
|
// --- Init with existing users (for CRC32 cache) ---
|
|
|
|
func TestInit_CachesExistingUsers(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
db := newTestDB(t)
|
|
require.NoError(t, db.Init())
|
|
|
|
// Save a user
|
|
now := time.Now().UTC()
|
|
db.SaveUser(model.User{Username: "cached", Email: "cached@test.com", Admin: true, OIDCSub: "sub-cache", CreatedAt: now, UpdatedAt: now})
|
|
|
|
// Re-init to trigger cache rebuild
|
|
require.NoError(t, db.Init())
|
|
|
|
util.DBUsersToCRC32Mutex.RLock()
|
|
_, ok := util.DBUsersToCRC32["cached"]
|
|
util.DBUsersToCRC32Mutex.RUnlock()
|
|
assert.True(t, ok, "User should be in CRC32 cache after Init")
|
|
}
|
|
|
|
// --- GetClients without QR code (no private key) ---
|
|
|
|
func TestGetClients_WithoutQRCode(t *testing.T) {
|
|
db := initTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
db.SaveClient(model.Client{
|
|
ID: "noqr", Name: "No QR", PublicKey: "pub1",
|
|
AllocatedIPs: []string{}, AllowedIPs: []string{},
|
|
ExtraAllowedIPs: []string{}, SubnetRanges: []string{},
|
|
CreatedAt: now, UpdatedAt: now,
|
|
})
|
|
|
|
clients, err := db.GetClients(false)
|
|
require.NoError(t, err)
|
|
assert.Len(t, clients, 1)
|
|
assert.Empty(t, clients[0].QRCode)
|
|
}
|
|
|
|
// --- Server interface with multiple addresses ---
|
|
|
|
func TestGetServer_MultipleAddresses(t *testing.T) {
|
|
db := initTestDB(t)
|
|
|
|
iface := model.ServerInterface{
|
|
Addresses: []string{"10.0.0.0/24", "fd00::1/64"},
|
|
ListenPort: 51820,
|
|
UpdatedAt: time.Now().UTC(),
|
|
}
|
|
require.NoError(t, db.SaveServerInterface(iface))
|
|
|
|
server, err := db.GetServer()
|
|
require.NoError(t, err)
|
|
assert.Len(t, server.Interface.Addresses, 2)
|
|
assert.Contains(t, server.Interface.Addresses, "10.0.0.0/24")
|
|
assert.Contains(t, server.Interface.Addresses, "fd00::1/64")
|
|
}
|
|
|
|
// --- GetAllocatedIPs edge cases ---
|
|
|
|
func TestGetAllocatedIPs_ExcludeNonExistentClient(t *testing.T) {
|
|
db := initTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
db.SaveClient(model.Client{
|
|
ID: "c1", Name: "Client",
|
|
AllocatedIPs: []string{"10.252.1.5/32"},
|
|
AllowedIPs: []string{}, ExtraAllowedIPs: []string{}, SubnetRanges: []string{},
|
|
CreatedAt: now, UpdatedAt: now,
|
|
})
|
|
|
|
// Exclude a non-existent client - should still return all IPs
|
|
ips, err := db.GetAllocatedIPs("nonexistent")
|
|
require.NoError(t, err)
|
|
assert.Contains(t, ips, "10.252.1.5")
|
|
}
|
|
|
|
// --- SaveClient with all fields populated ---
|
|
|
|
func TestSaveClient_AllFields(t *testing.T) {
|
|
db := newTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
client := model.Client{
|
|
ID: "full-client",
|
|
Name: "Full Client",
|
|
Email: "full@test.com",
|
|
PublicKey: "fullpub",
|
|
PrivateKey: "fullpriv",
|
|
PresharedKey: "fullpsk",
|
|
AllocatedIPs: []string{"10.0.0.2/32", "10.0.0.3/32"},
|
|
AllowedIPs: []string{"0.0.0.0/0", "::/0"},
|
|
ExtraAllowedIPs: []string{"192.168.1.0/24"},
|
|
SubnetRanges: []string{"range1"},
|
|
Endpoint: "vpn.example.com:51820",
|
|
AdditionalNotes: "Test notes\nLine 2",
|
|
UseServerDNS: true,
|
|
Enabled: true,
|
|
CreatedAt: now,
|
|
UpdatedAt: now,
|
|
}
|
|
|
|
err := db.SaveClient(client)
|
|
require.NoError(t, err)
|
|
|
|
got, err := db.GetClientByID("full-client", model.QRCodeSettings{Enabled: false})
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "Full Client", got.Client.Name)
|
|
assert.Equal(t, "full@test.com", got.Client.Email)
|
|
assert.Equal(t, []string{"10.0.0.2/32", "10.0.0.3/32"}, got.Client.AllocatedIPs)
|
|
assert.Equal(t, []string{"0.0.0.0/0", "::/0"}, got.Client.AllowedIPs)
|
|
assert.Equal(t, []string{"192.168.1.0/24"}, got.Client.ExtraAllowedIPs)
|
|
assert.Equal(t, []string{"range1"}, got.Client.SubnetRanges)
|
|
assert.Equal(t, "vpn.example.com:51820", got.Client.Endpoint)
|
|
assert.True(t, got.Client.UseServerDNS)
|
|
}
|
|
|
|
// --- Migrate with no missing keys (no-op path) ---
|
|
|
|
func TestMigrate_NoMissingKeys(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
db := newTestDB(t)
|
|
require.NoError(t, db.Init())
|
|
|
|
// All clients already have public keys - migrate should be a no-op
|
|
now := time.Now().UTC()
|
|
db.SaveClient(model.Client{
|
|
ID: "has-key", Name: "Has Key", PublicKey: "pubexists", PrivateKey: "privexists",
|
|
AllocatedIPs: []string{}, AllowedIPs: []string{},
|
|
ExtraAllowedIPs: []string{}, SubnetRanges: []string{},
|
|
CreatedAt: now, UpdatedAt: now,
|
|
})
|
|
|
|
// Running migrate again should not change anything
|
|
db.migrate()
|
|
|
|
got, err := db.GetClientByID("has-key", model.QRCodeSettings{Enabled: false})
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "pubexists", got.Client.PublicKey)
|
|
}
|
|
|
|
// --- GetWakeOnLanHosts with many hosts ---
|
|
|
|
func TestGetWakeOnLanHosts_ManyHosts(t *testing.T) {
|
|
db := initTestDB(t)
|
|
|
|
for i := 1; i <= 5; i++ {
|
|
mac := fmt.Sprintf("AA:BB:CC:DD:%02X:%02X", i/256, i%256)
|
|
db.SaveWakeOnLanHost(model.WakeOnLanHost{MacAddress: mac, Name: fmt.Sprintf("Host%d", i)})
|
|
}
|
|
|
|
hosts, err := db.GetWakeOnLanHosts()
|
|
require.NoError(t, err)
|
|
assert.Len(t, hosts, 5)
|
|
}
|
|
|
|
// --- GetUsers with multiple users ---
|
|
|
|
func TestGetUsers_Multiple(t *testing.T) {
|
|
db := newTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
db.SaveUser(model.User{Username: "u1", Email: "u1@test.com", OIDCSub: "sub-u1", Admin: true, CreatedAt: now, UpdatedAt: now})
|
|
db.SaveUser(model.User{Username: "u2", Email: "u2@test.com", OIDCSub: "sub-u2", Admin: false, CreatedAt: now, UpdatedAt: now})
|
|
db.SaveUser(model.User{Username: "u3", Email: "u3@test.com", OIDCSub: "sub-u3", Admin: false, CreatedAt: now, UpdatedAt: now})
|
|
db.SaveUser(model.User{Username: "u4", Email: "u4@test.com", OIDCSub: "sub-u4", Admin: true, CreatedAt: now, UpdatedAt: now})
|
|
|
|
users, err := db.GetUsers()
|
|
require.NoError(t, err)
|
|
assert.Len(t, users, 4)
|
|
|
|
// Verify all fields are populated
|
|
for _, u := range users {
|
|
assert.NotEmpty(t, u.Username)
|
|
assert.NotEmpty(t, u.Email)
|
|
assert.False(t, u.CreatedAt.IsZero())
|
|
assert.False(t, u.UpdatedAt.IsZero())
|
|
}
|
|
}
|
|
|
|
func TestGetUsers_Empty(t *testing.T) {
|
|
db := newTestDB(t)
|
|
|
|
users, err := db.GetUsers()
|
|
require.NoError(t, err)
|
|
assert.Nil(t, users) // no users -> nil slice
|
|
}
|
|
|
|
// --- Double Init idempotency with data ---
|
|
|
|
func TestInit_Idempotent_WithData(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
db := newTestDB(t)
|
|
require.NoError(t, db.Init())
|
|
|
|
// Save some data
|
|
now := time.Now().UTC()
|
|
db.SaveClient(model.Client{
|
|
ID: "idempotent-test", Name: "Idempotent",
|
|
AllocatedIPs: []string{"10.252.1.100/32"}, AllowedIPs: []string{"0.0.0.0/0"},
|
|
ExtraAllowedIPs: []string{}, SubnetRanges: []string{},
|
|
CreatedAt: now, UpdatedAt: now,
|
|
})
|
|
|
|
// Init again - should not lose data
|
|
require.NoError(t, db.Init())
|
|
|
|
// Verify client is still there
|
|
cd, err := db.GetClientByID("idempotent-test", model.QRCodeSettings{Enabled: false})
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "Idempotent", cd.Client.Name)
|
|
|
|
// Verify server config is still intact
|
|
server, err := db.GetServer()
|
|
require.NoError(t, err)
|
|
assert.NotEmpty(t, server.KeyPair.PublicKey)
|
|
}
|
|
|
|
// --- GetAllocatedIPs with IPv6 ---
|
|
|
|
func TestGetAllocatedIPs_WithIPv6(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
os.Setenv("WGUI_SERVER_INTERFACE_ADDRESSES", "10.252.1.0/24,fd00::1/64")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
defer os.Unsetenv("WGUI_SERVER_INTERFACE_ADDRESSES")
|
|
|
|
db := newTestDB(t)
|
|
require.NoError(t, db.Init())
|
|
|
|
now := time.Now().UTC()
|
|
db.SaveClient(model.Client{
|
|
ID: "ipv6-client", Name: "IPv6 Client",
|
|
AllocatedIPs: []string{"10.252.1.5/32", "fd00::5/128"},
|
|
AllowedIPs: []string{}, ExtraAllowedIPs: []string{}, SubnetRanges: []string{},
|
|
CreatedAt: now, UpdatedAt: now,
|
|
})
|
|
|
|
ips, err := db.GetAllocatedIPs("")
|
|
require.NoError(t, err)
|
|
assert.Contains(t, ips, "10.252.1.5")
|
|
assert.Contains(t, ips, "fd00::5")
|
|
// Should also contain server IPs
|
|
assert.Contains(t, ips, "10.252.1.0")
|
|
assert.Contains(t, ips, "fd00::1")
|
|
}
|
|
|
|
// --- GetGlobalSettings after update ---
|
|
|
|
func TestGetGlobalSettings_AfterUpdate(t *testing.T) {
|
|
db := initTestDB(t)
|
|
|
|
gs := model.GlobalSetting{
|
|
EndpointAddress: "updated.vpn.com",
|
|
DNSServers: []string{"1.1.1.1", "1.0.0.1"},
|
|
MTU: 1300,
|
|
PersistentKeepalive: 30,
|
|
FirewallMark: "0xbeef",
|
|
Table: "100",
|
|
ConfigFilePath: "/custom/wg.conf",
|
|
UpdatedAt: time.Now().UTC(),
|
|
}
|
|
require.NoError(t, db.SaveGlobalSettings(gs))
|
|
|
|
got, err := db.GetGlobalSettings()
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "updated.vpn.com", got.EndpointAddress)
|
|
assert.Equal(t, []string{"1.1.1.1", "1.0.0.1"}, got.DNSServers)
|
|
assert.Equal(t, 1300, got.MTU)
|
|
assert.Equal(t, 30, got.PersistentKeepalive)
|
|
assert.Equal(t, "0xbeef", got.FirewallMark)
|
|
assert.Equal(t, "100", got.Table)
|
|
assert.Equal(t, "/custom/wg.conf", got.ConfigFilePath)
|
|
}
|
|
|
|
// --- GetServer after updates ---
|
|
|
|
func TestGetServer_AfterInterfaceAndKeypairUpdates(t *testing.T) {
|
|
db := initTestDB(t)
|
|
|
|
now := time.Now().UTC()
|
|
iface := model.ServerInterface{
|
|
Addresses: []string{"172.16.0.0/16"},
|
|
ListenPort: 12345,
|
|
PostUp: "iptup",
|
|
PreDown: "iptpre",
|
|
PostDown: "iptdown",
|
|
UpdatedAt: now,
|
|
}
|
|
require.NoError(t, db.SaveServerInterface(iface))
|
|
|
|
kp := model.ServerKeypair{
|
|
PrivateKey: "serverprivkey",
|
|
PublicKey: "serverpubkey",
|
|
UpdatedAt: now,
|
|
}
|
|
require.NoError(t, db.SaveServerKeyPair(kp))
|
|
|
|
server, err := db.GetServer()
|
|
require.NoError(t, err)
|
|
assert.Equal(t, []string{"172.16.0.0/16"}, server.Interface.Addresses)
|
|
assert.Equal(t, 12345, server.Interface.ListenPort)
|
|
assert.Equal(t, "iptup", server.Interface.PostUp)
|
|
assert.Equal(t, "iptpre", server.Interface.PreDown)
|
|
assert.Equal(t, "iptdown", server.Interface.PostDown)
|
|
assert.Equal(t, "serverprivkey", server.KeyPair.PrivateKey)
|
|
assert.Equal(t, "serverpubkey", server.KeyPair.PublicKey)
|
|
}
|
|
|
|
// --- MigrateFromJSON with valid clients and WoL hosts ---
|
|
|
|
func TestMigrateFromJSON_MultipleClients(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
db := newTestDB(t)
|
|
require.NoError(t, db.Init())
|
|
|
|
jsonDBPath := t.TempDir()
|
|
serverDir := filepath.Join(jsonDBPath, "server")
|
|
require.NoError(t, os.MkdirAll(serverDir, 0755))
|
|
|
|
// Create clients directory with multiple clients
|
|
clientsDir := filepath.Join(jsonDBPath, "clients")
|
|
require.NoError(t, os.MkdirAll(clientsDir, 0755))
|
|
|
|
for i := 1; i <= 3; i++ {
|
|
clientJSON := fmt.Sprintf(`{"id":"migclient%d","name":"Mig Client %d","public_key":"migclientpub%d","allocated_ips":["10.0.0.%d/32"],"allowed_ips":["0.0.0.0/0"],"extra_allowed_ips":[],"subnet_ranges":[],"enabled":true}`, i, i, i, i+1)
|
|
require.NoError(t, os.WriteFile(filepath.Join(clientsDir, fmt.Sprintf("client%d.json", i)), []byte(clientJSON), 0644))
|
|
}
|
|
|
|
err := MigrateFromJSON(db, jsonDBPath)
|
|
require.NoError(t, err)
|
|
|
|
clients, err := db.GetClients(false)
|
|
require.NoError(t, err)
|
|
assert.Len(t, clients, 3)
|
|
}
|
|
|
|
func TestMigrateFromJSON_InvalidKeypairJSON(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
db := newTestDB(t)
|
|
require.NoError(t, db.Init())
|
|
|
|
jsonDBPath := t.TempDir()
|
|
serverDir := filepath.Join(jsonDBPath, "server")
|
|
require.NoError(t, os.MkdirAll(serverDir, 0755))
|
|
|
|
// Valid interfaces.json
|
|
require.NoError(t, os.WriteFile(filepath.Join(serverDir, "interfaces.json"),
|
|
[]byte(`{"addresses":["10.0.0.1/24"],"listen_port":"51820"}`), 0644))
|
|
|
|
// Invalid keypair.json
|
|
require.NoError(t, os.WriteFile(filepath.Join(serverDir, "keypair.json"),
|
|
[]byte("{invalid}"), 0644))
|
|
|
|
err := MigrateFromJSON(db, jsonDBPath)
|
|
assert.Error(t, err)
|
|
assert.Contains(t, err.Error(), "keypair")
|
|
}
|
|
|
|
func TestMigrateFromJSON_InvalidGlobalSettingsJSON(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
db := newTestDB(t)
|
|
require.NoError(t, db.Init())
|
|
|
|
jsonDBPath := t.TempDir()
|
|
serverDir := filepath.Join(jsonDBPath, "server")
|
|
require.NoError(t, os.MkdirAll(serverDir, 0755))
|
|
|
|
// Valid interfaces and keypair
|
|
require.NoError(t, os.WriteFile(filepath.Join(serverDir, "interfaces.json"),
|
|
[]byte(`{"addresses":["10.0.0.1/24"],"listen_port":"51820"}`), 0644))
|
|
require.NoError(t, os.WriteFile(filepath.Join(serverDir, "keypair.json"),
|
|
[]byte(`{"private_key":"pk","public_key":"pub"}`), 0644))
|
|
|
|
// Invalid global_settings.json
|
|
require.NoError(t, os.WriteFile(filepath.Join(serverDir, "global_settings.json"),
|
|
[]byte("{invalid}"), 0644))
|
|
|
|
err := MigrateFromJSON(db, jsonDBPath)
|
|
assert.Error(t, err)
|
|
assert.Contains(t, err.Error(), "global settings")
|
|
}
|
|
|
|
func TestMigrateFromJSON_InvalidHashesJSON(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
db := newTestDB(t)
|
|
require.NoError(t, db.Init())
|
|
|
|
jsonDBPath := t.TempDir()
|
|
serverDir := filepath.Join(jsonDBPath, "server")
|
|
require.NoError(t, os.MkdirAll(serverDir, 0755))
|
|
|
|
// Valid interfaces, keypair, global_settings
|
|
require.NoError(t, os.WriteFile(filepath.Join(serverDir, "interfaces.json"),
|
|
[]byte(`{"addresses":["10.0.0.1/24"],"listen_port":"51820"}`), 0644))
|
|
require.NoError(t, os.WriteFile(filepath.Join(serverDir, "keypair.json"),
|
|
[]byte(`{"private_key":"pk","public_key":"pub"}`), 0644))
|
|
require.NoError(t, os.WriteFile(filepath.Join(serverDir, "global_settings.json"),
|
|
[]byte(`{"endpoint_address":"10.0.0.1","dns_servers":["8.8.8.8"],"mtu":"1420","persistent_keepalive":"25"}`), 0644))
|
|
|
|
// Invalid hashes.json
|
|
require.NoError(t, os.WriteFile(filepath.Join(serverDir, "hashes.json"),
|
|
[]byte("{invalid}"), 0644))
|
|
|
|
err := MigrateFromJSON(db, jsonDBPath)
|
|
assert.Error(t, err)
|
|
assert.Contains(t, err.Error(), "hashes")
|
|
}
|
|
|
|
// --- SaveHashes and read back ---
|
|
|
|
func TestSaveHashes_MultipleTimes(t *testing.T) {
|
|
db := initTestDB(t)
|
|
|
|
h1 := model.ClientServerHashes{Client: "hash1c", Server: "hash1s"}
|
|
require.NoError(t, db.SaveHashes(h1))
|
|
|
|
got1, err := db.GetHashes()
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "hash1c", got1.Client)
|
|
assert.Equal(t, "hash1s", got1.Server)
|
|
|
|
h2 := model.ClientServerHashes{Client: "hash2c", Server: "hash2s"}
|
|
require.NoError(t, db.SaveHashes(h2))
|
|
|
|
got2, err := db.GetHashes()
|
|
require.NoError(t, err)
|
|
assert.Equal(t, "hash2c", got2.Client)
|
|
assert.Equal(t, "hash2s", got2.Server)
|
|
}
|
|
|
|
func TestMigrateFromJSON_RenameFailure(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
db := newTestDB(t)
|
|
require.NoError(t, db.Init())
|
|
|
|
jsonDBPath := t.TempDir()
|
|
serverDir := filepath.Join(jsonDBPath, "server")
|
|
require.NoError(t, os.MkdirAll(serverDir, 0755))
|
|
|
|
// Create the backup path as a non-empty directory so rename fails
|
|
backupPath := jsonDBPath + ".json.bak"
|
|
require.NoError(t, os.MkdirAll(filepath.Join(backupPath, "blocker"), 0755))
|
|
|
|
// Run migration - should succeed even if rename fails (rename failure is a warning)
|
|
err := MigrateFromJSON(db, jsonDBPath)
|
|
require.NoError(t, err)
|
|
}
|
|
|
|
// --- MigrateFromJSON: missing interfaces.json (skip) ---
|
|
|
|
func TestMigrateFromJSON_MissingInterfacesJSON(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
db := newTestDB(t)
|
|
require.NoError(t, db.Init())
|
|
|
|
jsonDBPath := t.TempDir()
|
|
serverDir := filepath.Join(jsonDBPath, "server")
|
|
require.NoError(t, os.MkdirAll(serverDir, 0755))
|
|
|
|
// No interfaces.json - should skip gracefully
|
|
err := MigrateFromJSON(db, jsonDBPath)
|
|
require.NoError(t, err)
|
|
}
|
|
|
|
// --- MigrateFromJSON: missing keypair.json (skip) ---
|
|
|
|
func TestMigrateFromJSON_MissingKeypairJSON(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
db := newTestDB(t)
|
|
require.NoError(t, db.Init())
|
|
|
|
jsonDBPath := t.TempDir()
|
|
serverDir := filepath.Join(jsonDBPath, "server")
|
|
require.NoError(t, os.MkdirAll(serverDir, 0755))
|
|
|
|
// Create valid interfaces but skip keypair
|
|
require.NoError(t, os.WriteFile(filepath.Join(serverDir, "interfaces.json"),
|
|
[]byte(`{"addresses":["10.0.0.1/24"],"listen_port":"51820"}`), 0644))
|
|
|
|
// No keypair.json - should skip gracefully
|
|
err := MigrateFromJSON(db, jsonDBPath)
|
|
require.NoError(t, err)
|
|
}
|
|
|
|
// --- MigrateFromJSON: missing global_settings.json (skip) ---
|
|
|
|
func TestMigrateFromJSON_MissingGlobalSettings(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
db := newTestDB(t)
|
|
require.NoError(t, db.Init())
|
|
|
|
jsonDBPath := t.TempDir()
|
|
serverDir := filepath.Join(jsonDBPath, "server")
|
|
require.NoError(t, os.MkdirAll(serverDir, 0755))
|
|
|
|
// interfaces + keypair but no global_settings
|
|
require.NoError(t, os.WriteFile(filepath.Join(serverDir, "interfaces.json"),
|
|
[]byte(`{"addresses":["10.0.0.1/24"],"listen_port":"51820"}`), 0644))
|
|
require.NoError(t, os.WriteFile(filepath.Join(serverDir, "keypair.json"),
|
|
[]byte(`{"private_key":"pk","public_key":"pub"}`), 0644))
|
|
|
|
err := MigrateFromJSON(db, jsonDBPath)
|
|
require.NoError(t, err)
|
|
}
|
|
|
|
// --- MigrateFromJSON: missing hashes.json (skip) ---
|
|
|
|
func TestMigrateFromJSON_MissingHashes(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
db := newTestDB(t)
|
|
require.NoError(t, db.Init())
|
|
|
|
jsonDBPath := t.TempDir()
|
|
serverDir := filepath.Join(jsonDBPath, "server")
|
|
require.NoError(t, os.MkdirAll(serverDir, 0755))
|
|
|
|
// All server files except hashes
|
|
require.NoError(t, os.WriteFile(filepath.Join(serverDir, "interfaces.json"),
|
|
[]byte(`{"addresses":["10.0.0.1/24"],"listen_port":"51820"}`), 0644))
|
|
require.NoError(t, os.WriteFile(filepath.Join(serverDir, "keypair.json"),
|
|
[]byte(`{"private_key":"pk","public_key":"pub"}`), 0644))
|
|
require.NoError(t, os.WriteFile(filepath.Join(serverDir, "global_settings.json"),
|
|
[]byte(`{"endpoint_address":"10.0.0.1","dns_servers":["8.8.8.8"],"mtu":"1420","persistent_keepalive":"25"}`), 0644))
|
|
|
|
err := MigrateFromJSON(db, jsonDBPath)
|
|
require.NoError(t, err)
|
|
}
|
|
|
|
func TestMigrateFromJSON_WolHostSaveFailure(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
db := newTestDB(t)
|
|
require.NoError(t, db.Init())
|
|
|
|
jsonDBPath := t.TempDir()
|
|
serverDir := filepath.Join(jsonDBPath, "server")
|
|
require.NoError(t, os.MkdirAll(serverDir, 0755))
|
|
|
|
// Create WoL host with empty MAC address (valid JSON but SaveWakeOnLanHost fails)
|
|
wolDir := filepath.Join(jsonDBPath, "wake_on_lan_hosts")
|
|
require.NoError(t, os.MkdirAll(wolDir, 0755))
|
|
wolJSON := `{"MacAddress":"","Name":"Bad Host"}`
|
|
require.NoError(t, os.WriteFile(filepath.Join(wolDir, "bad.json"), []byte(wolJSON), 0644))
|
|
|
|
err := MigrateFromJSON(db, jsonDBPath)
|
|
assert.Error(t, err)
|
|
assert.Contains(t, err.Error(), "WoL host")
|
|
}
|
|
|
|
func TestMigrateFromJSON_ClientSaveFailure(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
db := newTestDB(t)
|
|
require.NoError(t, db.Init())
|
|
|
|
jsonDBPath := t.TempDir()
|
|
serverDir := filepath.Join(jsonDBPath, "server")
|
|
require.NoError(t, os.MkdirAll(serverDir, 0755))
|
|
|
|
// Create a first valid client
|
|
clientsDir := filepath.Join(jsonDBPath, "clients")
|
|
require.NoError(t, os.MkdirAll(clientsDir, 0755))
|
|
clientJSON1 := `{"id":"dup1","name":"DupClient","public_key":"pub1","allocated_ips":[],"allowed_ips":[],"extra_allowed_ips":[],"subnet_ranges":[],"enabled":true}`
|
|
require.NoError(t, os.WriteFile(filepath.Join(clientsDir, "client1.json"), []byte(clientJSON1), 0644))
|
|
|
|
// Create a second client with duplicate name (unique index on name will cause failure)
|
|
clientJSON2 := `{"id":"dup2","name":"DupClient","public_key":"pub2","allocated_ips":[],"allowed_ips":[],"extra_allowed_ips":[],"subnet_ranges":[],"enabled":true}`
|
|
require.NoError(t, os.WriteFile(filepath.Join(clientsDir, "client2.json"), []byte(clientJSON2), 0644))
|
|
|
|
err := MigrateFromJSON(db, jsonDBPath)
|
|
assert.Error(t, err)
|
|
assert.Contains(t, err.Error(), "migrate client")
|
|
}
|
|
|
|
// --- Migration integration tests (simulate production upgrade scenario) ---
|
|
|
|
// TestMigrate_RunsOnInit verifies that reopening an existing DB runs migrate()
|
|
// which deletes legacy users without oidc_sub.
|
|
func TestMigrate_RunsOnInit(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
dir := t.TempDir()
|
|
dbPath := filepath.Join(dir, "test.db")
|
|
|
|
// Phase 1: create DB and insert a legacy user (no oidc_sub)
|
|
db1, err := New(dbPath)
|
|
require.NoError(t, err)
|
|
require.NoError(t, db1.Init())
|
|
|
|
_, err = db1.db.Exec(
|
|
`INSERT INTO users (username, email, display_name, oidc_sub, admin, created_at, updated_at)
|
|
VALUES ('legacy', 'legacy@test.com', 'Legacy User', '', 1, ?, ?)`,
|
|
time.Now().UTC(), time.Now().UTC(),
|
|
)
|
|
require.NoError(t, err)
|
|
|
|
users, err := db1.GetUsers()
|
|
require.NoError(t, err)
|
|
assert.Len(t, users, 1, "legacy user must exist before migration")
|
|
|
|
db1.db.Close()
|
|
|
|
// Phase 2: reopen the same DB — Init() must run migrate() and delete legacy user
|
|
db2, err := New(dbPath)
|
|
require.NoError(t, err)
|
|
require.NoError(t, db2.Init())
|
|
|
|
users, err = db2.GetUsers()
|
|
require.NoError(t, err)
|
|
assert.Len(t, users, 0, "legacy user without oidc_sub must be deleted by migration")
|
|
}
|
|
|
|
func TestMigrate_PreservesOIDCUsers(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
dir := t.TempDir()
|
|
dbPath := filepath.Join(dir, "test.db")
|
|
|
|
db1, err := New(dbPath)
|
|
require.NoError(t, err)
|
|
require.NoError(t, db1.Init())
|
|
|
|
now := time.Now().UTC()
|
|
db1.db.Exec(
|
|
`INSERT INTO users (username, email, display_name, oidc_sub, admin, created_at, updated_at)
|
|
VALUES ('oidcuser', 'oidc@test.com', 'OIDC User', 'sub-12345', 1, ?, ?)`, now, now)
|
|
db1.db.Exec(
|
|
`INSERT INTO users (username, email, display_name, oidc_sub, admin, created_at, updated_at)
|
|
VALUES ('legacy', 'legacy@test.com', 'Legacy', '', 0, ?, ?)`, now, now)
|
|
db1.db.Close()
|
|
|
|
db2, err := New(dbPath)
|
|
require.NoError(t, err)
|
|
require.NoError(t, db2.Init())
|
|
|
|
users, err := db2.GetUsers()
|
|
require.NoError(t, err)
|
|
assert.Len(t, users, 1, "only the OIDC user should survive migration")
|
|
assert.Equal(t, "oidcuser", users[0].Username)
|
|
assert.Equal(t, "sub-12345", users[0].OIDCSub)
|
|
}
|
|
|
|
func TestMigrate_DeletesNullOIDCSub(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
dir := t.TempDir()
|
|
dbPath := filepath.Join(dir, "test.db")
|
|
|
|
db1, err := New(dbPath)
|
|
require.NoError(t, err)
|
|
require.NoError(t, db1.Init())
|
|
|
|
db1.db.Exec(
|
|
`INSERT INTO users (username, email, display_name, oidc_sub, admin, created_at, updated_at)
|
|
VALUES ('nulluser', 'null@test.com', 'Null', NULL, 1, ?, ?)`,
|
|
time.Now().UTC(), time.Now().UTC())
|
|
db1.db.Close()
|
|
|
|
db2, err := New(dbPath)
|
|
require.NoError(t, err)
|
|
require.NoError(t, db2.Init())
|
|
|
|
users, err := db2.GetUsers()
|
|
require.NoError(t, err)
|
|
assert.Len(t, users, 0, "user with NULL oidc_sub must be deleted")
|
|
}
|
|
|
|
// TestMigrate_EmptyUsersEnablesFirstOIDCAdmin simulates the full production flow:
|
|
// 1. Existing DB has legacy password-only users
|
|
// 2. Upgrade runs Init() which deletes them
|
|
// 3. First OIDC login sees empty users table and becomes admin
|
|
func TestMigrate_EmptyUsersEnablesFirstOIDCAdmin(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
dir := t.TempDir()
|
|
dbPath := filepath.Join(dir, "test.db")
|
|
|
|
db1, err := New(dbPath)
|
|
require.NoError(t, err)
|
|
require.NoError(t, db1.Init())
|
|
|
|
db1.db.Exec(
|
|
`INSERT INTO users (username, email, display_name, oidc_sub, admin, created_at, updated_at)
|
|
VALUES ('old-admin', 'old@test.com', 'Old Admin', '', 1, ?, ?)`,
|
|
time.Now().UTC(), time.Now().UTC())
|
|
db1.db.Close()
|
|
|
|
db2, err := New(dbPath)
|
|
require.NoError(t, err)
|
|
require.NoError(t, db2.Init())
|
|
|
|
users, err := db2.GetUsers()
|
|
require.NoError(t, err)
|
|
require.Len(t, users, 0, "users table must be empty after migration")
|
|
|
|
// Simulate first OIDC login (same logic as findOrCreateOIDCUser)
|
|
now := time.Now().UTC()
|
|
newUser := model.User{
|
|
Username: "first-oidc", Email: "first@co.com", DisplayName: "First",
|
|
OIDCSub: "oidc-sub-first", Admin: len(users) == 0,
|
|
CreatedAt: now, UpdatedAt: now,
|
|
}
|
|
require.NoError(t, db2.SaveUser(newUser))
|
|
|
|
saved, err := db2.GetUserByName("first-oidc")
|
|
require.NoError(t, err)
|
|
assert.True(t, saved.Admin, "first OIDC user after migration must be admin")
|
|
}
|
|
|
|
// TestMigrate_PromotesOnlyUserToAdmin simulates the exact production bug:
|
|
// OIDC user was created with admin=false while legacy users still existed,
|
|
// then migration deletes legacy users, leaving the OIDC user as the only
|
|
// user but not admin. Migration must promote them.
|
|
func TestMigrate_PromotesOnlyUserToAdmin(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
dir := t.TempDir()
|
|
dbPath := filepath.Join(dir, "test.db")
|
|
|
|
db1, err := New(dbPath)
|
|
require.NoError(t, err)
|
|
require.NoError(t, db1.Init())
|
|
|
|
now := time.Now().UTC()
|
|
// OIDC user created with admin=false (because legacy users existed at login time)
|
|
db1.db.Exec(
|
|
`INSERT INTO users (username, email, display_name, oidc_sub, admin, created_at, updated_at)
|
|
VALUES ('gunter', 'gunter@co.com', 'Gunter', 'oidc-sub-123', 0, ?, ?)`, now, now)
|
|
db1.db.Close()
|
|
|
|
// reopen — migrate should see no admin exists and promote gunter
|
|
db2, err := New(dbPath)
|
|
require.NoError(t, err)
|
|
require.NoError(t, db2.Init())
|
|
|
|
user, err := db2.GetUserByName("gunter")
|
|
require.NoError(t, err)
|
|
assert.True(t, user.Admin, "only OIDC user must be promoted to admin when no admin exists")
|
|
}
|
|
|
|
func TestMigrate_DoesNotPromoteWhenAdminExists(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
dir := t.TempDir()
|
|
dbPath := filepath.Join(dir, "test.db")
|
|
|
|
db1, err := New(dbPath)
|
|
require.NoError(t, err)
|
|
require.NoError(t, db1.Init())
|
|
|
|
now := time.Now().UTC()
|
|
db1.db.Exec(
|
|
`INSERT INTO users (username, email, display_name, oidc_sub, admin, created_at, updated_at)
|
|
VALUES ('admin-user', 'admin@co.com', 'Admin', 'oidc-admin', 1, ?, ?)`, now, now)
|
|
db1.db.Exec(
|
|
`INSERT INTO users (username, email, display_name, oidc_sub, admin, created_at, updated_at)
|
|
VALUES ('regular', 'reg@co.com', 'Regular', 'oidc-reg', 0, ?, ?)`, now, now)
|
|
db1.db.Close()
|
|
|
|
db2, err := New(dbPath)
|
|
require.NoError(t, err)
|
|
require.NoError(t, db2.Init())
|
|
|
|
user, err := db2.GetUserByName("regular")
|
|
require.NoError(t, err)
|
|
assert.False(t, user.Admin, "should not promote when an admin already exists")
|
|
}
|
|
|
|
func TestMigrate_UniqueIndexOnName(t *testing.T) {
|
|
db := initTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
c1 := model.Client{
|
|
ID: "c1", Name: "same-name", Email: "a@test.com",
|
|
AllocatedIPs: []string{}, AllowedIPs: []string{},
|
|
ExtraAllowedIPs: []string{}, SubnetRanges: []string{},
|
|
Enabled: true, CreatedAt: now, UpdatedAt: now,
|
|
}
|
|
require.NoError(t, db.SaveClient(c1))
|
|
|
|
c2 := model.Client{
|
|
ID: "c2", Name: "same-name", Email: "b@test.com",
|
|
AllocatedIPs: []string{}, AllowedIPs: []string{},
|
|
ExtraAllowedIPs: []string{}, SubnetRanges: []string{},
|
|
Enabled: true, CreatedAt: now, UpdatedAt: now,
|
|
}
|
|
assert.Error(t, db.SaveClient(c2), "duplicate name must be rejected by unique index")
|
|
}
|
|
|
|
func TestMigrate_UniqueIndexOnPublicKey(t *testing.T) {
|
|
db := initTestDB(t)
|
|
key, _ := wgtypes.GeneratePrivateKey()
|
|
pubKey := key.PublicKey().String()
|
|
now := time.Now().UTC()
|
|
|
|
c1 := model.Client{
|
|
ID: "c1", Name: "a", PublicKey: pubKey, Email: "a@test.com",
|
|
AllocatedIPs: []string{}, AllowedIPs: []string{},
|
|
ExtraAllowedIPs: []string{}, SubnetRanges: []string{},
|
|
Enabled: true, CreatedAt: now, UpdatedAt: now,
|
|
}
|
|
require.NoError(t, db.SaveClient(c1))
|
|
|
|
c2 := model.Client{
|
|
ID: "c2", Name: "b", PublicKey: pubKey, Email: "b@test.com",
|
|
AllocatedIPs: []string{}, AllowedIPs: []string{},
|
|
ExtraAllowedIPs: []string{}, SubnetRanges: []string{},
|
|
Enabled: true, CreatedAt: now, UpdatedAt: now,
|
|
}
|
|
assert.Error(t, db.SaveClient(c2), "duplicate public_key must be rejected by unique index")
|
|
}
|
|
|
|
func TestMigrate_EmptyPublicKeysAllowed(t *testing.T) {
|
|
db := initTestDB(t)
|
|
now := time.Now().UTC()
|
|
|
|
c1 := model.Client{
|
|
ID: "c1", Name: "a", PublicKey: "", Email: "a@test.com",
|
|
AllocatedIPs: []string{}, AllowedIPs: []string{},
|
|
ExtraAllowedIPs: []string{}, SubnetRanges: []string{},
|
|
Enabled: true, CreatedAt: now, UpdatedAt: now,
|
|
}
|
|
require.NoError(t, db.SaveClient(c1))
|
|
|
|
c2 := model.Client{
|
|
ID: "c2", Name: "b", PublicKey: "", Email: "b@test.com",
|
|
AllocatedIPs: []string{}, AllowedIPs: []string{},
|
|
ExtraAllowedIPs: []string{}, SubnetRanges: []string{},
|
|
Enabled: true, CreatedAt: now, UpdatedAt: now,
|
|
}
|
|
require.NoError(t, db.SaveClient(c2), "partial unique index must allow multiple empty public keys")
|
|
}
|
|
|
|
func TestMigrate_DerivesPublicKeys_OnReopen(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
dir := t.TempDir()
|
|
dbPath := filepath.Join(dir, "test.db")
|
|
|
|
db1, err := New(dbPath)
|
|
require.NoError(t, err)
|
|
require.NoError(t, db1.Init())
|
|
|
|
key1, _ := wgtypes.GeneratePrivateKey()
|
|
key2, _ := wgtypes.GeneratePrivateKey()
|
|
now := time.Now().UTC()
|
|
|
|
db1.db.Exec(
|
|
`INSERT INTO clients (id, private_key, public_key, preshared_key, name, email,
|
|
subnet_ranges, allocated_ips, allowed_ips, extra_allowed_ips,
|
|
endpoint, additional_notes, use_server_dns, enabled, created_at, updated_at)
|
|
VALUES (?, ?, '', '', 'C1', 'c1@t.com', '[]', '[]', '[]', '[]', '', '', 1, 1, ?, ?)`,
|
|
"c1", key1.String(), now, now)
|
|
db1.db.Exec(
|
|
`INSERT INTO clients (id, private_key, public_key, preshared_key, name, email,
|
|
subnet_ranges, allocated_ips, allowed_ips, extra_allowed_ips,
|
|
endpoint, additional_notes, use_server_dns, enabled, created_at, updated_at)
|
|
VALUES (?, ?, '', '', 'C2', 'c2@t.com', '[]', '[]', '[]', '[]', '', '', 1, 1, ?, ?)`,
|
|
"c2", key2.String(), now, now)
|
|
db1.db.Close()
|
|
|
|
db2, err := New(dbPath)
|
|
require.NoError(t, err)
|
|
require.NoError(t, db2.Init())
|
|
|
|
c1Data, err := db2.GetClientByID("c1", model.QRCodeSettings{Enabled: false})
|
|
require.NoError(t, err)
|
|
assert.Equal(t, key1.PublicKey().String(), c1Data.Client.PublicKey, "public key must be derived on reopen")
|
|
|
|
c2Data, err := db2.GetClientByID("c2", model.QRCodeSettings{Enabled: false})
|
|
require.NoError(t, err)
|
|
assert.Equal(t, key2.PublicKey().String(), c2Data.Client.PublicKey, "public key must be derived on reopen")
|
|
}
|
|
|
|
func TestMigrateFromJSON_SkipsNonJSON(t *testing.T) {
|
|
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
|
|
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")
|
|
|
|
db := newTestDB(t)
|
|
require.NoError(t, db.Init())
|
|
|
|
jsonDBPath := t.TempDir()
|
|
serverDir := filepath.Join(jsonDBPath, "server")
|
|
require.NoError(t, os.MkdirAll(serverDir, 0755))
|
|
|
|
// Create users dir with non-JSON files
|
|
usersDir := filepath.Join(jsonDBPath, "users")
|
|
require.NoError(t, os.MkdirAll(usersDir, 0755))
|
|
require.NoError(t, os.WriteFile(filepath.Join(usersDir, "readme.txt"), []byte("not json"), 0644))
|
|
require.NoError(t, os.MkdirAll(filepath.Join(usersDir, "subdir"), 0755))
|
|
|
|
// Create clients dir with invalid JSON (should warn and skip)
|
|
clientsDir := filepath.Join(jsonDBPath, "clients")
|
|
require.NoError(t, os.MkdirAll(clientsDir, 0755))
|
|
require.NoError(t, os.WriteFile(filepath.Join(clientsDir, "bad.json"), []byte("{invalid"), 0644))
|
|
|
|
// Create wol dir with invalid JSON (should warn and skip)
|
|
wolDir := filepath.Join(jsonDBPath, "wake_on_lan_hosts")
|
|
require.NoError(t, os.MkdirAll(wolDir, 0755))
|
|
require.NoError(t, os.WriteFile(filepath.Join(wolDir, "bad.json"), []byte("{invalid"), 0644))
|
|
|
|
err := MigrateFromJSON(db, jsonDBPath)
|
|
require.NoError(t, err)
|
|
}
|