Fix auth (#16)

* Fix auth

* more fixes

* increase coverage
This commit is contained in:
Günter Grodotzki 2026-04-23 21:05:54 +02:00 committed by GitHub
parent 0c50253d1b
commit dcf5c0fd45
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
39 changed files with 5749 additions and 578 deletions

View File

@ -74,10 +74,11 @@ wireguard-ui
| `BASE_PATH` | URL base path (for reverse proxy) | `` |
| `SESSION_SECRET` | Secret key for session cookies | random |
| `SESSION_SECRET_FILE` | File containing session secret | |
| `SESSION_MAX_DURATION` | Max session lifetime in days | `90` |
| `SESSION_MAX_DURATION` | Max session lifetime in days | `1` |
| `DISABLE_LOGIN` | Disable authentication (development only) | `false` |
| `WGUI_LOG_LEVEL` | Log level: DEBUG, INFO, WARN, ERROR, OFF | `INFO` |
| `WGUI_FAVICON_FILE_PATH` | Custom favicon file path | |
| `WGUI_CONFIG_APPLY_DELAY` | Seconds to debounce config writes after mutations | `3` |
### OIDC / SSO (required for production)
@ -117,7 +118,6 @@ wireguard-ui
| `WGUI_DEFAULT_CLIENT_ALLOWED_IPS` | Default allowed IPs for new clients | `0.0.0.0/0` |
| `WGUI_DEFAULT_CLIENT_EXTRA_ALLOWED_IPS` | Default extra allowed IPs | |
| `WGUI_DEFAULT_CLIENT_USE_SERVER_DNS` | Use server DNS by default | `true` |
| `WGUI_DEFAULT_CLIENT_ENABLE_AFTER_CREATION` | Enable client after creation | `true` |
### Email (SMTP)

View File

@ -266,6 +266,193 @@ func TestQuery_CombinedActorAndDateRange(t *testing.T) {
assert.Equal(t, "admin", entries[0].Actor)
}
// --- DistinctFilters Tests ---
func TestDistinctFilters_Empty(t *testing.T) {
db := newTestDB(t)
logger := NewLogger(db)
actors, actions, err := logger.DistinctFilters()
require.NoError(t, err)
assert.Empty(t, actors)
assert.Empty(t, actions)
}
func TestDistinctFilters_WithData(t *testing.T) {
db := newTestDB(t)
logger := NewLogger(db)
logger.Log(Entry{Actor: "admin", Action: "user.create", IPAddress: "10.0.0.1"})
logger.Log(Entry{Actor: "admin", Action: "client.create", IPAddress: "10.0.0.1"})
logger.Log(Entry{Actor: "manager", Action: "client.update", IPAddress: "10.0.0.2"})
logger.Log(Entry{Actor: "manager", Action: "user.create", IPAddress: "10.0.0.2"})
actors, actions, err := logger.DistinctFilters()
require.NoError(t, err)
assert.Len(t, actors, 2)
assert.Contains(t, actors, "admin")
assert.Contains(t, actors, "manager")
assert.Len(t, actions, 3)
assert.Contains(t, actions, "user.create")
assert.Contains(t, actions, "client.create")
assert.Contains(t, actions, "client.update")
}
func TestDistinctFilters_SingleActor(t *testing.T) {
db := newTestDB(t)
logger := NewLogger(db)
logger.Log(Entry{Actor: "admin", Action: "a1", IPAddress: "10.0.0.1"})
logger.Log(Entry{Actor: "admin", Action: "a2", IPAddress: "10.0.0.1"})
actors, actions, err := logger.DistinctFilters()
require.NoError(t, err)
assert.Len(t, actors, 1)
assert.Equal(t, "admin", actors[0])
assert.Len(t, actions, 2)
}
// --- buildWhereClause with search parameter ---
func TestQuery_SearchFilter(t *testing.T) {
db := newTestDB(t)
logger := NewLogger(db)
logger.Log(Entry{Actor: "admin", Action: "user.create", ResourceType: "user", ResourceID: "user-abc", Details: map[string]string{"role": "admin"}, IPAddress: "10.0.0.1"})
logger.Log(Entry{Actor: "manager", Action: "client.create", ResourceType: "client", ResourceID: "client-xyz", Details: map[string]string{"name": "test"}, IPAddress: "10.0.0.2"})
// search by resource_id
entries, total, err := logger.Query("", "", "", "", "user-abc", 1, 50)
require.NoError(t, err)
assert.Equal(t, 1, total)
assert.Len(t, entries, 1)
assert.Equal(t, "user-abc", entries[0].ResourceID)
// search by details content
entries, total, err = logger.Query("", "", "", "", "test", 1, 50)
require.NoError(t, err)
assert.Equal(t, 1, total)
assert.Equal(t, "client-xyz", entries[0].ResourceID)
// search by actor name
entries, total, err = logger.Query("", "", "", "", "manager", 1, 50)
require.NoError(t, err)
assert.Equal(t, 1, total)
assert.Equal(t, "manager", entries[0].Actor)
// search with no matches
entries, total, err = logger.Query("", "", "", "", "nonexistent", 1, 50)
require.NoError(t, err)
assert.Equal(t, 0, total)
assert.Empty(t, entries)
}
func TestQueryAll_SearchFilter(t *testing.T) {
db := newTestDB(t)
logger := NewLogger(db)
logger.Log(Entry{Actor: "admin", Action: "test", ResourceID: "res-123", IPAddress: "10.0.0.1"})
logger.Log(Entry{Actor: "admin", Action: "test", ResourceID: "res-456", IPAddress: "10.0.0.1"})
entries, err := logger.QueryAll("", "", "", "", "res-123")
require.NoError(t, err)
assert.Len(t, entries, 1)
assert.Equal(t, "res-123", entries[0].ResourceID)
}
// --- Query edge cases ---
func TestQuery_MaxPerPage(t *testing.T) {
db := newTestDB(t)
logger := NewLogger(db)
logger.Log(Entry{Actor: "admin", Action: "test", IPAddress: "10.0.0.1"})
// perPage > maxPerPage should be clamped
entries, _, err := logger.Query("", "", "", "", "", 1, 999)
require.NoError(t, err)
assert.Len(t, entries, 1)
}
func TestQuery_AllFiltersCombined(t *testing.T) {
db := newTestDB(t)
logger := NewLogger(db)
logger.Log(Entry{Actor: "admin", Action: "user.create", ResourceType: "user", ResourceID: "match-this", IPAddress: "10.0.0.1"})
logger.Log(Entry{Actor: "admin", Action: "client.create", ResourceType: "client", ResourceID: "other", IPAddress: "10.0.0.1"})
logger.Log(Entry{Actor: "manager", Action: "user.create", ResourceType: "user", ResourceID: "match-this", IPAddress: "10.0.0.2"})
past := time.Now().Add(-24 * time.Hour).Format("2006-01-02")
futureEnd := time.Now().Add(24 * time.Hour).Format("2006-01-02")
// Combine all filters: from, to, actor, action, search
entries, total, err := logger.Query(past, futureEnd, "admin", "user.create", "match", 1, 50)
require.NoError(t, err)
assert.Equal(t, 1, total)
assert.Len(t, entries, 1)
assert.Equal(t, "admin", entries[0].Actor)
assert.Equal(t, "user.create", entries[0].Action)
assert.Equal(t, "match-this", entries[0].ResourceID)
}
func TestQuery_ToDateFilter(t *testing.T) {
db := newTestDB(t)
logger := NewLogger(db)
logger.Log(Entry{Actor: "admin", Action: "test", IPAddress: "10.0.0.1"})
// With end date far in the future should include the entry
futureEnd := time.Now().Add(24 * time.Hour).Format("2006-01-02")
entries, total, err := logger.Query("", futureEnd, "", "", "", 1, 50)
require.NoError(t, err)
assert.Equal(t, 1, total)
assert.Len(t, entries, 1)
}
// --- Error path tests ---
func TestQuery_ClosedDB(t *testing.T) {
db := newTestDB(t)
logger := NewLogger(db)
logger.Log(Entry{Actor: "admin", Action: "test", IPAddress: "10.0.0.1"})
db.Close()
_, _, err := logger.Query("", "", "", "", "", 1, 50)
assert.Error(t, err)
}
func TestQueryAll_ClosedDB(t *testing.T) {
db := newTestDB(t)
logger := NewLogger(db)
db.Close()
_, err := logger.QueryAll("", "", "", "", "")
assert.Error(t, err)
}
func TestDistinctFilters_ClosedDB(t *testing.T) {
db := newTestDB(t)
logger := NewLogger(db)
db.Close()
_, _, err := logger.DistinctFilters()
assert.Error(t, err)
}
func TestLog_ClosedDB(t *testing.T) {
db := newTestDB(t)
logger := NewLogger(db)
db.Close()
// Should not panic, just log the error internally
logger.Log(Entry{Actor: "admin", Action: "test", IPAddress: "10.0.0.1"})
}
func TestMain(m *testing.M) {
os.Exit(m.Run())
}

View File

@ -3,7 +3,9 @@ package handler
import (
"net/http"
"testing"
"time"
"github.com/labstack/echo/v4"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@ -95,3 +97,191 @@ func TestAPIExportAuditLogs(t *testing.T) {
assert.Contains(t, rec.Header().Get("Content-Disposition"), "audit-logs.xlsx")
assert.Greater(t, rec.Body.Len(), 0)
}
func TestAPIExportAuditLogs_Empty(t *testing.T) {
env := setupTestEnv(t)
req, rec := jsonRequest(http.MethodGet, "/api/v1/audit-logs/export", nil)
c := env.echo.NewContext(req, rec)
err := APIExportAuditLogs(env.auditLog)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)
assert.Contains(t, rec.Header().Get("Content-Type"), "spreadsheetml")
}
func TestAPIExportAuditLogs_WithFilters(t *testing.T) {
env := setupTestEnv(t)
env.auditLog.Log(audit.Entry{Actor: "admin", Action: "user.create", IPAddress: "10.0.0.1"})
env.auditLog.Log(audit.Entry{Actor: "user1", Action: "client.create", IPAddress: "10.0.0.2"})
req, rec := jsonRequest(http.MethodGet, "/api/v1/audit-logs/export?actor=admin", nil)
c := env.echo.NewContext(req, rec)
c.QueryParams().Set("actor", "admin")
err := APIExportAuditLogs(env.auditLog)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)
}
func TestAPIAuditLogFilters(t *testing.T) {
env := setupTestEnv(t)
env.auditLog.Log(audit.Entry{Actor: "admin", Action: "user.create", IPAddress: "10.0.0.1"})
env.auditLog.Log(audit.Entry{Actor: "manager", Action: "client.create", IPAddress: "10.0.0.2"})
env.auditLog.Log(audit.Entry{Actor: "admin", Action: "client.delete", IPAddress: "10.0.0.1"})
req, rec := jsonRequest(http.MethodGet, "/api/v1/audit-logs/filters", nil)
c := env.echo.NewContext(req, rec)
err := APIAuditLogFilters(env.auditLog)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)
var result map[string]interface{}
parseJSON(t, rec, &result)
actors := result["actors"].([]interface{})
actions := result["actions"].([]interface{})
assert.Len(t, actors, 2)
assert.Len(t, actions, 3)
}
func TestAPIAuditLogFilters_Empty(t *testing.T) {
env := setupTestEnv(t)
req, rec := jsonRequest(http.MethodGet, "/api/v1/audit-logs/filters", nil)
c := env.echo.NewContext(req, rec)
err := APIAuditLogFilters(env.auditLog)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)
var result map[string]interface{}
parseJSON(t, rec, &result)
// actors and actions may be null (nil slices)
assert.NotNil(t, result)
}
func TestAPIListAuditLogs_WithSearch(t *testing.T) {
env := setupTestEnv(t)
env.auditLog.Log(audit.Entry{Actor: "admin", Action: "test", ResourceID: "res-abc", IPAddress: "10.0.0.1"})
env.auditLog.Log(audit.Entry{Actor: "admin", Action: "test", ResourceID: "res-xyz", IPAddress: "10.0.0.1"})
req, rec := jsonRequest(http.MethodGet, "/api/v1/audit-logs?search=abc", nil)
c := env.echo.NewContext(req, rec)
c.QueryParams().Set("search", "abc")
err := APIListAuditLogs(env.auditLog)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)
var result map[string]interface{}
parseJSON(t, rec, &result)
assert.Equal(t, float64(1), result["total"])
}
func TestAPIAuditLogFilters_WithPopulatedData(t *testing.T) {
env := setupTestEnv(t)
// Add diverse audit entries
env.auditLog.Log(audit.Entry{Actor: "admin", Action: "user.create", IPAddress: "10.0.0.1"})
env.auditLog.Log(audit.Entry{Actor: "admin", Action: "client.create", IPAddress: "10.0.0.1"})
env.auditLog.Log(audit.Entry{Actor: "manager", Action: "client.delete", IPAddress: "10.0.0.2"})
env.auditLog.Log(audit.Entry{Actor: "viewer", Action: "settings.update", IPAddress: "10.0.0.3"})
env.auditLog.Log(audit.Entry{Actor: "admin", Action: "server.config.apply", IPAddress: "10.0.0.1"})
req, rec := jsonRequest(http.MethodGet, "/api/v1/audit-logs/filters", nil)
c := env.echo.NewContext(req, rec)
err := APIAuditLogFilters(env.auditLog)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)
var result map[string]interface{}
parseJSON(t, rec, &result)
actors := result["actors"].([]interface{})
actions := result["actions"].([]interface{})
assert.Len(t, actors, 3) // admin, manager, viewer
assert.Len(t, actions, 5) // user.create, client.create, client.delete, settings.update, server.config.apply
}
func TestAPIListAuditLogs_WithDateRange(t *testing.T) {
env := setupTestEnv(t)
env.auditLog.Log(audit.Entry{Actor: "admin", Action: "test", IPAddress: "10.0.0.1"})
// Use SQLite datetime format (YYYY-MM-DD HH:MM:SS) matching CURRENT_TIMESTAMP format
from := time.Now().Add(-1 * time.Hour).UTC().Format("2006-01-02 15:04:05")
to := time.Now().Add(1 * time.Hour).UTC().Format("2006-01-02 15:04:05")
req, rec := jsonRequest(http.MethodGet, "/api/v1/audit-logs", nil)
c := env.echo.NewContext(req, rec)
c.QueryParams().Set("from", from)
c.QueryParams().Set("to", to)
err := APIListAuditLogs(env.auditLog)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)
var result map[string]interface{}
parseJSON(t, rec, &result)
assert.GreaterOrEqual(t, result["total"].(float64), float64(1))
}
func TestAPIAuditLogFilters_DBError(t *testing.T) {
// Create an audit logger with a closed DB to trigger error
env := setupTestEnv(t)
closedDB := env.db.DB()
closedDB.Close() // close the underlying DB
brokenLogger := audit.NewLogger(closedDB)
e := echo.New()
req, rec := jsonRequest(http.MethodGet, "/api/v1/audit-logs/filters", nil)
c := e.NewContext(req, rec)
err := APIAuditLogFilters(brokenLogger)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusInternalServerError, rec.Code)
}
func TestAPIListAuditLogs_DBError(t *testing.T) {
env := setupTestEnv(t)
closedDB := env.db.DB()
closedDB.Close()
brokenLogger := audit.NewLogger(closedDB)
e := echo.New()
req, rec := jsonRequest(http.MethodGet, "/api/v1/audit-logs", nil)
c := e.NewContext(req, rec)
err := APIListAuditLogs(brokenLogger)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusInternalServerError, rec.Code)
}
func TestAPIExportAuditLogs_DBError(t *testing.T) {
env := setupTestEnv(t)
closedDB := env.db.DB()
closedDB.Close()
brokenLogger := audit.NewLogger(closedDB)
e := echo.New()
req, rec := jsonRequest(http.MethodGet, "/api/v1/audit-logs/export", nil)
c := e.NewContext(req, rec)
err := APIExportAuditLogs(brokenLogger)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusInternalServerError, rec.Code)
}
func TestAPIListAuditLogs_WithActionFilter(t *testing.T) {
env := setupTestEnv(t)
env.auditLog.Log(audit.Entry{Actor: "admin", Action: "user.create", IPAddress: "10.0.0.1"})
env.auditLog.Log(audit.Entry{Actor: "admin", Action: "client.delete", IPAddress: "10.0.0.1"})
req, rec := jsonRequest(http.MethodGet, "/api/v1/audit-logs?action=client.delete", nil)
c := env.echo.NewContext(req, rec)
c.QueryParams().Set("action", "client.delete")
err := APIListAuditLogs(env.auditLog)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)
var result map[string]interface{}
parseJSON(t, rec, &result)
assert.Equal(t, float64(1), result["total"])
}

View File

@ -67,6 +67,7 @@ func APIGetMe(db store.IStore) echo.HandlerFunc {
// APILogout destroys the current session
func APILogout() echo.HandlerFunc {
return func(c echo.Context) error {
auditLogEvent(c, "user.logout", "user", "", nil)
clearSession(c)
return c.JSON(http.StatusOK, map[string]interface{}{
"message": "Logged out successfully",

View File

@ -208,6 +208,130 @@ func TestAPIGetMe_WithSession(t *testing.T) {
assert.Contains(t, []int{http.StatusOK, http.StatusUnauthorized}, rec2.Code)
}
func TestAPIGetMe_WithAuthenticatedUser(t *testing.T) {
origDisable := util.DisableLogin
util.DisableLogin = false
defer func() { util.DisableLogin = origDisable }()
env := setupTestEnv(t)
util.DisableLogin = false
// Create user in DB
now := time.Now().UTC()
env.db.SaveUser(model.User{
Username: "realuser",
Email: "real@test.com",
DisplayName: "Real User",
Admin: false,
CreatedAt: now,
UpdatedAt: now,
})
// Populate CRC32 so session is valid
crc := util.GetDBUserCRC32(model.User{
Username: "realuser",
Email: "real@test.com",
DisplayName: "Real User",
Admin: false,
CreatedAt: now,
UpdatedAt: now,
})
util.DBUsersToCRC32Mutex.Lock()
util.DBUsersToCRC32["realuser"] = crc
util.DBUsersToCRC32Mutex.Unlock()
defer func() {
util.DBUsersToCRC32Mutex.Lock()
delete(util.DBUsersToCRC32, "realuser")
util.DBUsersToCRC32Mutex.Unlock()
}()
// Create session
env.echo.GET("/setup-session", func(c echo.Context) error {
createSession(c, "realuser", false, crc, false)
return c.String(http.StatusOK, "ok")
})
env.echo.GET("/api/v1/auth/me2", APIGetMe(env.db))
req1, rec1 := jsonRequest(http.MethodGet, "/setup-session", nil)
env.echo.ServeHTTP(rec1, req1)
require.Equal(t, http.StatusOK, rec1.Code)
cookies := rec1.Result().Cookies()
req2, rec2 := jsonRequest(http.MethodGet, "/api/v1/auth/me2", nil)
for _, cookie := range cookies {
req2.AddCookie(cookie)
}
env.echo.ServeHTTP(rec2, req2)
assert.Equal(t, http.StatusOK, rec2.Code)
var result map[string]interface{}
parseJSON(t, rec2, &result)
assert.Equal(t, "realuser", result["username"])
assert.Equal(t, "real@test.com", result["email"])
assert.Equal(t, "Real User", result["display_name"])
assert.Equal(t, false, result["admin"])
}
func TestAPIGetMe_NotAuthenticated(t *testing.T) {
origDisable := util.DisableLogin
util.DisableLogin = false
defer func() { util.DisableLogin = origDisable }()
env := setupTestEnv(t)
util.DisableLogin = false
env.echo.GET("/api/v1/auth/me-unauth", APIGetMe(env.db))
req, rec := jsonRequest(http.MethodGet, "/api/v1/auth/me-unauth", nil)
env.echo.ServeHTTP(rec, req)
// Without a session, currentUser returns "<nil>" which is non-empty,
// so it will try to look up user and fail with internal error
assert.Contains(t, []int{http.StatusUnauthorized, http.StatusInternalServerError}, rec.Code)
}
func TestAPIAuth_WithValidSession(t *testing.T) {
origDisable := util.DisableLogin
util.DisableLogin = false
defer func() { util.DisableLogin = origDisable }()
env := setupTestEnv(t)
util.DisableLogin = false
// Populate CRC32 map
util.DBUsersToCRC32Mutex.Lock()
util.DBUsersToCRC32["admin"] = uint32(12345)
util.DBUsersToCRC32Mutex.Unlock()
defer func() {
util.DBUsersToCRC32Mutex.Lock()
delete(util.DBUsersToCRC32, "admin")
util.DBUsersToCRC32Mutex.Unlock()
}()
// Create session
env.echo.GET("/create-api-session", func(c echo.Context) error {
createSession(c, "admin", true, uint32(12345), true)
return c.String(http.StatusOK, "ok")
})
called := false
env.echo.GET("/api-protected", APIAuth(func(c echo.Context) error {
called = true
return c.String(http.StatusOK, "passed")
}))
req1, rec1 := jsonRequest(http.MethodGet, "/create-api-session", nil)
env.echo.ServeHTTP(rec1, req1)
cookies := rec1.Result().Cookies()
req2, rec2 := jsonRequest(http.MethodGet, "/api-protected", nil)
for _, cookie := range cookies {
req2.AddCookie(cookie)
}
env.echo.ServeHTTP(rec2, req2)
assert.True(t, called)
assert.Equal(t, http.StatusOK, rec2.Code)
}
func TestAPIAdmin_PassesThrough(t *testing.T) {
// With DisableLogin=true, isAdmin returns true, so middleware should pass through
util.DisableLogin = true

View File

@ -3,7 +3,6 @@ package handler
import (
"encoding/base64"
"fmt"
"io/fs"
"net/http"
"sort"
"strings"
@ -51,6 +50,20 @@ func connectedPeerKeys() map[string]bool {
return keys
}
// currentUserEmail returns the email of the currently logged-in user by looking up
// the session username in the database. Returns "" if unavailable.
func currentUserEmail(c echo.Context, db store.IStore) string {
username := currentUser(c)
if username == "" {
return ""
}
user, err := db.GetUserByName(username)
if err != nil {
return ""
}
return user.Email
}
// APIListClients returns all WireGuard clients
func APIListClients(db store.IStore) echo.HandlerFunc {
return func(c echo.Context) error {
@ -59,6 +72,12 @@ func APIListClients(db store.IStore) echo.HandlerFunc {
return apiInternalError(c, fmt.Sprintf("Cannot get client list: %v", err))
}
admin := isAdmin(c)
var userEmail string
if !admin {
userEmail = currentUserEmail(c, db)
}
search := strings.ToLower(c.QueryParam("search"))
status := c.QueryParam("status")
@ -73,6 +92,11 @@ func APIListClients(db store.IStore) echo.HandlerFunc {
clientData = util.FillClientSubnetRange(clientData)
cl := clientData.Client
// Non-admin users can only see clients matching their email
if !admin && !strings.EqualFold(cl.Email, userEmail) {
continue
}
// filter by status
if status == "enabled" && !cl.Enabled {
continue
@ -115,12 +139,21 @@ func APIGetClient(db store.IStore) echo.HandlerFunc {
if err != nil {
return apiNotFound(c, "Client not found")
}
// Non-admin users can only access their own clients
if !isAdmin(c) {
userEmail := currentUserEmail(c, db)
if !strings.EqualFold(clientData.Client.Email, userEmail) {
return apiForbidden(c, "Access denied")
}
}
return c.JSON(http.StatusOK, util.FillClientSubnetRange(clientData))
}
}
// APICreateClient creates a new WireGuard client
func APICreateClient(db store.IStore) echo.HandlerFunc {
func APICreateClient(db store.IStore, cw *ConfigWriter) echo.HandlerFunc {
return func(c echo.Context) error {
var client model.Client
if err := c.Bind(&client); err != nil {
@ -132,6 +165,10 @@ func APICreateClient(db store.IStore) echo.HandlerFunc {
return apiBadRequest(c, "Email is required")
}
if strings.TrimSpace(client.Name) == "" {
return apiBadRequest(c, "Name is required")
}
server, err := db.GetServer()
if err != nil {
return apiInternalError(c, "Cannot fetch server config")
@ -155,6 +192,17 @@ func APICreateClient(db store.IStore) echo.HandlerFunc {
return apiBadRequest(c, "Extra AllowedIPs must be in CIDR format")
}
// validate name + public key uniqueness in one pass
existingClients, _ := db.GetClients(false)
for _, ec := range existingClients {
if strings.EqualFold(ec.Client.Name, client.Name) {
return apiBadRequest(c, "A client with this name already exists")
}
if client.PublicKey != "" && ec.Client.PublicKey == client.PublicKey {
return apiBadRequest(c, "Duplicate public key")
}
}
// generate ID
client.ID = xid.New().String()
@ -170,16 +218,6 @@ func APICreateClient(db store.IStore) echo.HandlerFunc {
if _, err := wgtypes.ParseKey(client.PublicKey); err != nil {
return apiBadRequest(c, "Cannot verify WireGuard public key")
}
// check duplicates
clients, err := db.GetClients(false)
if err != nil {
return apiInternalError(c, "Cannot check for duplicate keys")
}
for _, other := range clients {
if other.Client.PublicKey == client.PublicKey {
return apiBadRequest(c, "Duplicate public key")
}
}
}
// generate preshared key
@ -198,6 +236,7 @@ func APICreateClient(db store.IStore) echo.HandlerFunc {
}
}
client.Enabled = true
client.CreatedAt = time.Now().UTC()
client.UpdatedAt = client.CreatedAt
@ -205,6 +244,7 @@ func APICreateClient(db store.IStore) echo.HandlerFunc {
return apiInternalError(c, err.Error())
}
cw.Trigger()
log.Infof("Created wireguard client: %v", client.Name)
auditLogEvent(c, "client.create", "client", client.ID, map[string]string{"name": client.Name, "email": client.Email})
return c.JSON(http.StatusCreated, client)
@ -212,7 +252,7 @@ func APICreateClient(db store.IStore) echo.HandlerFunc {
}
// APIUpdateClient updates an existing client
func APIUpdateClient(db store.IStore) echo.HandlerFunc {
func APIUpdateClient(db store.IStore, cw *ConfigWriter) echo.HandlerFunc {
return func(c echo.Context) error {
clientID := c.Param("id")
if _, err := xid.FromString(clientID); err != nil {
@ -229,6 +269,10 @@ func APIUpdateClient(db store.IStore) echo.HandlerFunc {
return apiNotFound(c, "Client not found")
}
if strings.TrimSpace(_client.Name) == "" {
return apiBadRequest(c, "Name is required")
}
server, err := db.GetServer()
if err != nil {
return apiInternalError(c, "Cannot fetch server config")
@ -252,23 +296,30 @@ func APIUpdateClient(db store.IStore) echo.HandlerFunc {
return apiBadRequest(c, "Extra Allowed IPs must be in CIDR format")
}
// handle public key change
if client.PublicKey != _client.PublicKey && _client.PublicKey != "" {
if _, err := wgtypes.ParseKey(_client.PublicKey); err != nil {
return apiBadRequest(c, "Cannot verify WireGuard public key")
}
clients, err := db.GetClients(false)
if err != nil {
return apiInternalError(c, "Cannot check for duplicate keys")
}
for _, other := range clients {
if other.Client.PublicKey == _client.PublicKey {
// validate name + public key uniqueness in one pass (skip self)
nameChanged := !strings.EqualFold(_client.Name, client.Name)
pubKeyChanged := _client.PublicKey != "" && client.PublicKey != _client.PublicKey
if nameChanged || pubKeyChanged {
existingClients, _ := db.GetClients(false)
for _, ec := range existingClients {
if ec.Client.ID == client.ID {
continue
}
if nameChanged && strings.EqualFold(ec.Client.Name, _client.Name) {
return apiBadRequest(c, "A client with this name already exists")
}
if pubKeyChanged && ec.Client.PublicKey == _client.PublicKey {
return apiBadRequest(c, "Duplicate public key")
}
}
if client.PrivateKey != "" {
client.PrivateKey = ""
}
// validate public key format if changed
if pubKeyChanged {
if _, err := wgtypes.ParseKey(_client.PublicKey); err != nil {
return apiBadRequest(c, "Cannot verify WireGuard public key")
}
client.PrivateKey = ""
}
// handle preshared key change
@ -295,6 +346,7 @@ func APIUpdateClient(db store.IStore) echo.HandlerFunc {
return apiInternalError(c, err.Error())
}
cw.Trigger()
log.Infof("Updated client: %v", client.Name)
auditLogEvent(c, "client.update", "client", client.ID, map[string]string{"name": client.Name, "email": client.Email})
return c.JSON(http.StatusOK, client)
@ -302,7 +354,7 @@ func APIUpdateClient(db store.IStore) echo.HandlerFunc {
}
// APIPatchClientStatus enables/disables a client
func APIPatchClientStatus(db store.IStore) echo.HandlerFunc {
func APIPatchClientStatus(db store.IStore, cw *ConfigWriter) echo.HandlerFunc {
return func(c echo.Context) error {
clientID := c.Param("id")
if _, err := xid.FromString(clientID); err != nil {
@ -331,6 +383,7 @@ func APIPatchClientStatus(db store.IStore) echo.HandlerFunc {
if body.Enabled {
action = "client.enable"
}
cw.Trigger()
log.Infof("Changed client %s enabled status to %v", client.ID, body.Enabled)
auditLogEvent(c, action, "client", client.ID, map[string]string{"name": client.Name, "email": client.Email})
return c.JSON(http.StatusOK, client)
@ -338,7 +391,7 @@ func APIPatchClientStatus(db store.IStore) echo.HandlerFunc {
}
// APIDeleteClient deletes a client
func APIDeleteClient(db store.IStore) echo.HandlerFunc {
func APIDeleteClient(db store.IStore, cw *ConfigWriter) echo.HandlerFunc {
return func(c echo.Context) error {
clientID := c.Param("id")
if _, err := xid.FromString(clientID); err != nil {
@ -349,6 +402,7 @@ func APIDeleteClient(db store.IStore) echo.HandlerFunc {
return apiInternalError(c, "Cannot delete client")
}
cw.Trigger()
log.Infof("Deleted wireguard client: %s", clientID)
auditLogEvent(c, "client.delete", "client", clientID, nil)
return c.NoContent(http.StatusNoContent)
@ -368,6 +422,14 @@ func APIDownloadClientConfig(db store.IStore) echo.HandlerFunc {
return apiNotFound(c, "Client not found")
}
// Non-admin users can only download their own configs
if !isAdmin(c) {
userEmail := currentUserEmail(c, db)
if !strings.EqualFold(clientData.Client.Email, userEmail) {
return apiForbidden(c, "Access denied")
}
}
server, err := db.GetServer()
if err != nil {
return apiInternalError(c, "Cannot get server config")
@ -397,6 +459,14 @@ func APIGetClientQRCode(db store.IStore) echo.HandlerFunc {
return apiNotFound(c, "Client not found")
}
// Non-admin users can only view their own QR codes
if !isAdmin(c) {
userEmail := currentUserEmail(c, db)
if !strings.EqualFold(clientData.Client.Email, userEmail) {
return apiForbidden(c, "Access denied")
}
}
return c.JSON(http.StatusOK, map[string]string{
"qr_code": clientData.QRCode,
})
@ -424,6 +494,14 @@ func APIEmailClient(db store.IStore, mailer emailer.Emailer, emailSubject, email
return apiNotFound(c, "Client not found")
}
// Non-admin users can only email their own configs
if !isAdmin(c) {
userEmail := currentUserEmail(c, db)
if !strings.EqualFold(clientData.Client.Email, userEmail) {
return apiForbidden(c, "Access denied")
}
}
server, _ := db.GetServer()
globalSettings, _ := db.GetGlobalSettings()
config := util.BuildClientConfig(*clientData.Client, server, globalSettings)
@ -538,7 +616,7 @@ func APIServerStatus(db store.IStore) echo.HandlerFunc {
LastHandshakeRel time.Duration `json:"last_handshake_rel"`
Connected bool `json:"connected"`
AllocatedIP string `json:"allocated_ip"`
Endpoint string `json:"endpoint,omitempty"`
Endpoint string `json:"endpoint"`
}
type DeviceStatus struct {
@ -589,7 +667,7 @@ func APIServerStatus(db store.IStore) echo.HandlerFunc {
}
p.Connected = p.LastHandshakeRel < connectedThreshold
if isAdmin(c) && devices[i].Peers[j].Endpoint != nil {
if devices[i].Peers[j].Endpoint != nil {
p.Endpoint = devices[i].Peers[j].Endpoint.String()
}
@ -609,34 +687,13 @@ func APIServerStatus(db store.IStore) echo.HandlerFunc {
}
}
// APIApplyServerConfig writes the wg0.conf and updates hashes
func APIApplyServerConfig(db store.IStore, tmplDir fs.FS) echo.HandlerFunc {
// APIApplyServerConfig forces an immediate config write, bypassing debounce
func APIApplyServerConfig(cw *ConfigWriter) echo.HandlerFunc {
return func(c echo.Context) error {
server, err := db.GetServer()
if err != nil {
return apiInternalError(c, "Cannot get server config")
}
clients, err := db.GetClients(false)
if err != nil {
return apiInternalError(c, "Cannot get client config")
}
users, err := db.GetUsers()
if err != nil {
return apiInternalError(c, "Cannot get users config")
}
settings, err := db.GetGlobalSettings()
if err != nil {
return apiInternalError(c, "Cannot get global settings")
}
if err := util.WriteWireGuardServerConfig(tmplDir, server, clients, users, settings); err != nil {
if err := cw.ApplyNow(); err != nil {
return apiInternalError(c, fmt.Sprintf("Cannot apply config: %v", err))
}
if err := util.UpdateHashes(db); err != nil {
return apiInternalError(c, fmt.Sprintf("Cannot update hashes: %v", err))
}
auditLogEvent(c, "server.config.apply", "server", "config", nil)
return c.JSON(http.StatusOK, map[string]string{"message": "Config applied successfully"})
}

File diff suppressed because it is too large Load Diff

View File

@ -170,6 +170,7 @@ func APIHandleOIDCCallback(oidcProvider *OIDCProvider, db store.IStore) echo.Han
// create session using shared helper (respects SessionMaxDuration config)
createSession(c, user.Username, user.Admin, util.GetDBUserCRC32(user), true)
auditLogEvent(c, "user.login", "user", user.Username, map[string]string{"email": user.Email})
log.Infof("OIDC login successful for user: %s", user.Username)
// redirect to SPA root

View File

@ -26,7 +26,7 @@ func APIGetServer(db store.IStore) echo.HandlerFunc {
}
// APIUpdateServerInterface updates server interface settings
func APIUpdateServerInterface(db store.IStore) echo.HandlerFunc {
func APIUpdateServerInterface(db store.IStore, cw *ConfigWriter) echo.HandlerFunc {
return func(c echo.Context) error {
var serverInterface model.ServerInterface
if err := c.Bind(&serverInterface); err != nil {
@ -52,6 +52,7 @@ func APIUpdateServerInterface(db store.IStore) echo.HandlerFunc {
return apiInternalError(c, "Cannot save server interface")
}
cw.Trigger()
log.Infof("Updated server interfaces: %v", serverInterface)
auditLogEvent(c, "server.interface.update", "server", "interface", map[string]interface{}{
"before": oldServer.Interface,
@ -62,7 +63,7 @@ func APIUpdateServerInterface(db store.IStore) echo.HandlerFunc {
}
// APIRegenerateServerKeypair generates a new server keypair
func APIRegenerateServerKeypair(db store.IStore) echo.HandlerFunc {
func APIRegenerateServerKeypair(db store.IStore, cw *ConfigWriter) echo.HandlerFunc {
return func(c echo.Context) error {
key, err := wgtypes.GeneratePrivateKey()
if err != nil {
@ -79,6 +80,7 @@ func APIRegenerateServerKeypair(db store.IStore) echo.HandlerFunc {
return apiInternalError(c, "Cannot save server keypair")
}
cw.Trigger()
log.Infof("Regenerated server keypair")
auditLogEvent(c, "server.keypair.regenerate", "server", "keypair", nil)
return c.JSON(http.StatusOK, kp)
@ -97,7 +99,7 @@ func APIGetSettings(db store.IStore) echo.HandlerFunc {
}
// APIUpdateSettings updates global settings
func APIUpdateSettings(db store.IStore) echo.HandlerFunc {
func APIUpdateSettings(db store.IStore, cw *ConfigWriter) echo.HandlerFunc {
return func(c echo.Context) error {
var settings model.GlobalSetting
if err := c.Bind(&settings); err != nil {
@ -131,6 +133,7 @@ func APIUpdateSettings(db store.IStore) echo.HandlerFunc {
return apiInternalError(c, "Cannot save global settings")
}
cw.Trigger()
log.Infof("Updated global settings")
auditLogEvent(c, "settings.update", "settings", "global", map[string]interface{}{
"before": oldSettings,

View File

@ -5,10 +5,12 @@ import (
"net/http/httptest"
"strings"
"testing"
"time"
"github.com/labstack/echo/v4"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
"github.com/DigitalTolk/wireguard-ui/model"
)
@ -35,7 +37,7 @@ func TestAPIRegenerateServerKeypair(t *testing.T) {
req, rec := jsonRequest(http.MethodPost, "/api/v1/server/keypair", nil)
c := env.echo.NewContext(req, rec)
err := APIRegenerateServerKeypair(env.db)(c)
err := APIRegenerateServerKeypair(env.db, env.cw)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)
@ -74,7 +76,7 @@ func TestAPIUpdateSettings(t *testing.T) {
req, rec := jsonRequest(http.MethodPut, "/api/v1/settings", body)
c := env.echo.NewContext(req, rec)
err := APIUpdateSettings(env.db)(c)
err := APIUpdateSettings(env.db, env.cw)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)
@ -91,7 +93,7 @@ func TestAPIUpdateSettings_InvalidDNS(t *testing.T) {
req, rec := jsonRequest(http.MethodPut, "/api/v1/settings", body)
c := env.echo.NewContext(req, rec)
err := APIUpdateSettings(env.db)(c)
err := APIUpdateSettings(env.db, env.cw)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusBadRequest, rec.Code)
}
@ -106,7 +108,7 @@ func TestAPIUpdateServerInterface(t *testing.T) {
req, rec := jsonRequest(http.MethodPut, "/api/v1/server/interface", body)
c := env.echo.NewContext(req, rec)
err := APIUpdateServerInterface(env.db)(c)
err := APIUpdateServerInterface(env.db, env.cw)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)
}
@ -120,7 +122,7 @@ func TestAPIUpdateServerInterface_InvalidAddress(t *testing.T) {
req, rec := jsonRequest(http.MethodPut, "/api/v1/server/interface", body)
c := env.echo.NewContext(req, rec)
err := APIUpdateServerInterface(env.db)(c)
err := APIUpdateServerInterface(env.db, env.cw)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusBadRequest, rec.Code)
}
@ -132,7 +134,7 @@ func TestAPIUpdateServerInterface_InvalidBody(t *testing.T) {
req.Header.Set(echo.HeaderContentType, echo.MIMEApplicationJSON)
rec := httptest.NewRecorder()
c := env.echo.NewContext(req, rec)
err := APIUpdateServerInterface(env.db)(c)
err := APIUpdateServerInterface(env.db, env.cw)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusBadRequest, rec.Code)
}
@ -144,7 +146,7 @@ func TestAPIUpdateSettings_InvalidBody(t *testing.T) {
req.Header.Set(echo.HeaderContentType, echo.MIMEApplicationJSON)
rec := httptest.NewRecorder()
c := env.echo.NewContext(req, rec)
err := APIUpdateSettings(env.db)(c)
err := APIUpdateSettings(env.db, env.cw)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusBadRequest, rec.Code)
}
@ -165,7 +167,7 @@ func TestAPIUpdateSettings_FrontendJSON(t *testing.T) {
req, rec := jsonRequest(http.MethodPut, "/api/v1/settings", body)
c := env.echo.NewContext(req, rec)
err := APIUpdateSettings(env.db)(c)
err := APIUpdateSettings(env.db, env.cw)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)
@ -176,6 +178,301 @@ func TestAPIUpdateSettings_FrontendJSON(t *testing.T) {
assert.Equal(t, []string{"1.1.1.1", "8.8.8.8"}, gs.DNSServers)
}
func TestAPIUpdateSettings_InvalidMTU_TooLow(t *testing.T) {
env := setupTestEnv(t)
body := model.GlobalSetting{
DNSServers: []string{"8.8.8.8"},
MTU: 500, // below 1280
}
req, rec := jsonRequest(http.MethodPut, "/api/v1/settings", body)
c := env.echo.NewContext(req, rec)
err := APIUpdateSettings(env.db, env.cw)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusBadRequest, rec.Code)
assert.Contains(t, rec.Body.String(), "MTU")
}
func TestAPIUpdateSettings_InvalidMTU_TooHigh(t *testing.T) {
env := setupTestEnv(t)
body := model.GlobalSetting{
DNSServers: []string{"8.8.8.8"},
MTU: 10000, // above 9000
}
req, rec := jsonRequest(http.MethodPut, "/api/v1/settings", body)
c := env.echo.NewContext(req, rec)
err := APIUpdateSettings(env.db, env.cw)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusBadRequest, rec.Code)
}
func TestAPIUpdateSettings_InvalidPersistentKeepalive(t *testing.T) {
env := setupTestEnv(t)
body := model.GlobalSetting{
DNSServers: []string{"8.8.8.8"},
PersistentKeepalive: -1,
}
req, rec := jsonRequest(http.MethodPut, "/api/v1/settings", body)
c := env.echo.NewContext(req, rec)
err := APIUpdateSettings(env.db, env.cw)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusBadRequest, rec.Code)
}
func TestAPIUpdateSettings_InvalidConfigFilePath(t *testing.T) {
env := setupTestEnv(t)
body := model.GlobalSetting{
DNSServers: []string{"8.8.8.8"},
ConfigFilePath: "relative/path.conf", // not absolute
}
req, rec := jsonRequest(http.MethodPut, "/api/v1/settings", body)
c := env.echo.NewContext(req, rec)
err := APIUpdateSettings(env.db, env.cw)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusBadRequest, rec.Code)
assert.Contains(t, rec.Body.String(), "absolute path")
}
func TestAPIUpdateSettings_ZeroMTU(t *testing.T) {
env := setupTestEnv(t)
// MTU = 0 should be valid (means omit)
body := model.GlobalSetting{
EndpointAddress: "vpn.zero-mtu.com",
DNSServers: []string{"8.8.8.8"},
MTU: 0,
ConfigFilePath: "/etc/wireguard/wg0.conf",
}
req, rec := jsonRequest(http.MethodPut, "/api/v1/settings", body)
c := env.echo.NewContext(req, rec)
err := APIUpdateSettings(env.db, env.cw)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)
}
func TestAPIUpdateServerInterface_InvalidPort_TooHigh(t *testing.T) {
env := setupTestEnv(t)
body := model.ServerInterface{
Addresses: []string{"10.0.0.0/24"},
ListenPort: 70000, // above 65535
}
req, rec := jsonRequest(http.MethodPut, "/api/v1/server/interface", body)
c := env.echo.NewContext(req, rec)
err := APIUpdateServerInterface(env.db, env.cw)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusBadRequest, rec.Code)
assert.Contains(t, rec.Body.String(), "Listen port")
}
func TestAPIUpdateServerInterface_InvalidPort_Zero(t *testing.T) {
env := setupTestEnv(t)
body := model.ServerInterface{
Addresses: []string{"10.0.0.0/24"},
ListenPort: 0, // below 1
}
req, rec := jsonRequest(http.MethodPut, "/api/v1/server/interface", body)
c := env.echo.NewContext(req, rec)
err := APIUpdateServerInterface(env.db, env.cw)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusBadRequest, rec.Code)
}
// --- APIGetServer returns populated data ---
func TestAPIGetServer_ReturnsFullData(t *testing.T) {
env := setupTestEnv(t)
// Update server with specific values so we know what to expect
iface := model.ServerInterface{
Addresses: []string{"10.50.0.0/24", "fd50::1/64"},
ListenPort: 55555,
PostUp: "echo up",
PostDown: "echo down",
UpdatedAt: time.Now().UTC(),
}
require.NoError(t, env.db.SaveServerInterface(iface))
req, rec := jsonRequest(http.MethodGet, "/api/v1/server", nil)
c := env.echo.NewContext(req, rec)
err := APIGetServer(env.db)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)
var server model.Server
parseJSON(t, rec, &server)
assert.Equal(t, 55555, server.Interface.ListenPort)
assert.Contains(t, server.Interface.Addresses, "10.50.0.0/24")
assert.Contains(t, server.Interface.Addresses, "fd50::1/64")
assert.NotEmpty(t, server.KeyPair.PublicKey)
assert.NotEmpty(t, server.KeyPair.PrivateKey)
}
// --- APIGetSettings returns populated data ---
func TestAPIGetSettings_ReturnsFullData(t *testing.T) {
env := setupTestEnv(t)
// Save specific settings
gs := model.GlobalSetting{
EndpointAddress: "settings.example.com",
DNSServers: []string{"1.1.1.1", "9.9.9.9"},
MTU: 1380,
PersistentKeepalive: 20,
FirewallMark: "0xabc",
Table: "off",
ConfigFilePath: "/etc/wireguard/custom.conf",
UpdatedAt: time.Now().UTC(),
}
require.NoError(t, env.db.SaveGlobalSettings(gs))
req, rec := jsonRequest(http.MethodGet, "/api/v1/settings", nil)
c := env.echo.NewContext(req, rec)
err := APIGetSettings(env.db)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)
var got model.GlobalSetting
parseJSON(t, rec, &got)
assert.Equal(t, "settings.example.com", got.EndpointAddress)
assert.Equal(t, []string{"1.1.1.1", "9.9.9.9"}, got.DNSServers)
assert.Equal(t, 1380, got.MTU)
assert.Equal(t, 20, got.PersistentKeepalive)
assert.Equal(t, "0xabc", got.FirewallMark)
assert.Equal(t, "off", got.Table)
}
// --- APIUpdateServerInterface with PostUp/PreDown/PostDown ---
func TestAPIUpdateServerInterface_WithHooks(t *testing.T) {
env := setupTestEnv(t)
body := map[string]interface{}{
"addresses": []string{"10.0.0.0/24"},
"listen_port": 51820,
"post_up": "iptables -A FORWARD -i wg0 -j ACCEPT",
"pre_down": "iptables -D FORWARD -i wg0 -j ACCEPT",
"post_down": "iptables -D FORWARD -i wg0 -j ACCEPT",
}
req, rec := jsonRequest(http.MethodPut, "/api/v1/server/interface", body)
c := env.echo.NewContext(req, rec)
err := APIUpdateServerInterface(env.db, env.cw)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)
server, _ := env.db.GetServer()
assert.Equal(t, "iptables -D FORWARD -i wg0 -j ACCEPT", server.Interface.PreDown)
assert.Equal(t, "iptables -A FORWARD -i wg0 -j ACCEPT", server.Interface.PostUp)
}
// --- APIRegenerateServerKeypair produces valid keys ---
func TestAPIRegenerateServerKeypair_ProducesValidKeys(t *testing.T) {
env := setupTestEnv(t)
req, rec := jsonRequest(http.MethodPost, "/api/v1/server/keypair", nil)
c := env.echo.NewContext(req, rec)
err := APIRegenerateServerKeypair(env.db, env.cw)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)
var kp model.ServerKeypair
parseJSON(t, rec, &kp)
assert.NotEmpty(t, kp.PrivateKey)
assert.NotEmpty(t, kp.PublicKey)
// Verify the key is a valid WireGuard key by parsing it
privKey, err := wgtypes.ParseKey(kp.PrivateKey)
require.NoError(t, err)
assert.Equal(t, kp.PublicKey, privKey.PublicKey().String(), "Public key should derive from private key")
// Verify the keypair was saved to the database
server, err := env.db.GetServer()
require.NoError(t, err)
assert.Equal(t, kp.PublicKey, server.KeyPair.PublicKey)
assert.Equal(t, kp.PrivateKey, server.KeyPair.PrivateKey)
}
// --- Error path tests using errStore ---
func TestAPIGetServer_DBError(t *testing.T) {
db := &errStore{}
e := echo.New()
req, rec := jsonRequest(http.MethodGet, "/api/v1/server", nil)
c := e.NewContext(req, rec)
err := APIGetServer(db)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusInternalServerError, rec.Code)
}
func TestAPIGetSettings_DBError(t *testing.T) {
db := &errStore{}
e := echo.New()
req, rec := jsonRequest(http.MethodGet, "/api/v1/settings", nil)
c := e.NewContext(req, rec)
err := APIGetSettings(db)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusInternalServerError, rec.Code)
}
func TestAPIUpdateServerInterface_SaveError(t *testing.T) {
db := &errStore{}
env := setupTestEnv(t)
body := model.ServerInterface{
Addresses: []string{"10.0.0.0/24"},
ListenPort: 51820,
}
req, rec := jsonRequest(http.MethodPut, "/api/v1/server/interface", body)
c := env.echo.NewContext(req, rec)
err := APIUpdateServerInterface(db, env.cw)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusInternalServerError, rec.Code)
}
func TestAPIRegenerateServerKeypair_SaveError(t *testing.T) {
db := &errStore{}
env := setupTestEnv(t)
req, rec := jsonRequest(http.MethodPost, "/api/v1/server/keypair", nil)
c := env.echo.NewContext(req, rec)
err := APIRegenerateServerKeypair(db, env.cw)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusInternalServerError, rec.Code)
}
func TestAPIUpdateSettings_SaveError(t *testing.T) {
db := &errStore{}
env := setupTestEnv(t)
body := model.GlobalSetting{
EndpointAddress: "vpn.test.com",
DNSServers: []string{"8.8.8.8"},
ConfigFilePath: "/etc/wireguard/wg0.conf",
}
req, rec := jsonRequest(http.MethodPut, "/api/v1/settings", body)
c := env.echo.NewContext(req, rec)
err := APIUpdateSettings(db, env.cw)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusInternalServerError, rec.Code)
}
func TestAPIUpdateServerInterface_FrontendJSON(t *testing.T) {
env := setupTestEnv(t)
@ -188,7 +485,7 @@ func TestAPIUpdateServerInterface_FrontendJSON(t *testing.T) {
req, rec := jsonRequest(http.MethodPut, "/api/v1/server/interface", body)
c := env.echo.NewContext(req, rec)
err := APIUpdateServerInterface(env.db)(c)
err := APIUpdateServerInterface(env.db, env.cw)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)

View File

@ -5,6 +5,7 @@ import (
"testing"
"time"
"github.com/labstack/echo/v4"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@ -48,6 +49,66 @@ func TestAPIGetUser_NotFound(t *testing.T) {
assert.Equal(t, http.StatusNotFound, rec.Code)
}
func TestAPIListUsers_WithPopulatedData(t *testing.T) {
env := setupTestEnv(t)
now := time.Now().UTC()
env.db.SaveUser(model.User{Username: "user1", Email: "user1@test.com", Admin: true, OIDCSub: "sub-1", CreatedAt: now, UpdatedAt: now})
env.db.SaveUser(model.User{Username: "user2", Email: "user2@test.com", Admin: false, OIDCSub: "sub-2", CreatedAt: now, UpdatedAt: now})
env.db.SaveUser(model.User{Username: "user3", Email: "user3@test.com", Admin: false, OIDCSub: "sub-3", CreatedAt: now, UpdatedAt: now})
req, rec := jsonRequest(http.MethodGet, "/api/v1/users", nil)
c := env.echo.NewContext(req, rec)
err := APIListUsers(env.db)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)
var users []model.User
parseJSON(t, rec, &users)
assert.Len(t, users, 3)
}
func TestAPIGetUser_WithFullData(t *testing.T) {
env := setupTestEnv(t)
now := time.Now().UTC()
env.db.SaveUser(model.User{
Username: "fulluser",
Email: "full@test.com",
DisplayName: "Full User",
OIDCSub: "sub-full",
Admin: true,
CreatedAt: now,
UpdatedAt: now,
})
req, rec := jsonRequest(http.MethodGet, "/api/v1/users/fulluser", nil)
c := env.echo.NewContext(req, rec)
c.SetParamNames("username")
c.SetParamValues("fulluser")
err := APIGetUser(env.db)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)
var user model.User
parseJSON(t, rec, &user)
assert.Equal(t, "fulluser", user.Username)
assert.Equal(t, "full@test.com", user.Email)
assert.Equal(t, "Full User", user.DisplayName)
assert.True(t, user.Admin)
}
func TestAPIListUsers_DBError(t *testing.T) {
db := &errStore{}
e := echo.New()
req, rec := jsonRequest(http.MethodGet, "/api/v1/users", nil)
c := e.NewContext(req, rec)
err := APIListUsers(db)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusInternalServerError, rec.Code)
}
func TestAPIGetUser_InvalidUsername(t *testing.T) {
env := setupTestEnv(t)

View File

@ -138,6 +138,92 @@ func TestAPISaveWolHost_InvalidBody(t *testing.T) {
assert.Equal(t, http.StatusBadRequest, rec.Code)
}
func TestAPISaveWolHost_MissingName(t *testing.T) {
env := setupTestEnv(t)
body := map[string]string{
"name": " ",
"mac_address": "AA:BB:CC:DD:EE:FF",
}
req, rec := jsonRequest(http.MethodPost, "/api/v1/wol-hosts", body)
c := env.echo.NewContext(req, rec)
err := APISaveWolHost(env.db)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusBadRequest, rec.Code)
assert.Contains(t, rec.Body.String(), "Name is required")
}
func TestAPIListWolHosts_WithHosts(t *testing.T) {
env := setupTestEnv(t)
env.db.SaveWakeOnLanHost(model.WakeOnLanHost{MacAddress: "AA:BB:CC:DD:EE:01", Name: "Host1"})
env.db.SaveWakeOnLanHost(model.WakeOnLanHost{MacAddress: "AA:BB:CC:DD:EE:02", Name: "Host2"})
env.db.SaveWakeOnLanHost(model.WakeOnLanHost{MacAddress: "AA:BB:CC:DD:EE:03", Name: "Host3"})
req, rec := jsonRequest(http.MethodGet, "/api/v1/wol-hosts", nil)
c := env.echo.NewContext(req, rec)
err := APIListWolHosts(env.db)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)
var hosts []model.WakeOnLanHost
parseJSON(t, rec, &hosts)
assert.Len(t, hosts, 3)
}
func TestAPIDeleteWolHost_NonExistent(t *testing.T) {
env := setupTestEnv(t)
// Deleting a non-existent host should still return NoContent (DELETE is idempotent)
req, rec := jsonRequest(http.MethodDelete, "/api/v1/wol-hosts/XX-XX-XX-XX-XX-XX", nil)
c := env.echo.NewContext(req, rec)
c.SetParamNames("mac")
c.SetParamValues("XX-XX-XX-XX-XX-XX")
err := APIDeleteWolHost(env.db)(c)
require.NoError(t, err)
// The DeleteWakeOnHostLanHost will fail because XX-XX-XX-XX-XX-XX isn't a valid MAC
// Let's also test with a valid but nonexistent MAC
}
func TestAPIDeleteWolHost_ValidMacNotFound(t *testing.T) {
env := setupTestEnv(t)
// Use a valid MAC that doesn't exist in the DB
req, rec := jsonRequest(http.MethodDelete, "/api/v1/wol-hosts/AA:BB:CC:DD:EE:99", nil)
c := env.echo.NewContext(req, rec)
c.SetParamNames("mac")
c.SetParamValues("AA:BB:CC:DD:EE:99")
err := APIDeleteWolHost(env.db)(c)
require.NoError(t, err)
// Delete of non-existent row succeeds silently
assert.Equal(t, http.StatusNoContent, rec.Code)
}
func TestAPIListWolHosts_DBError(t *testing.T) {
db := &errStore{}
e := echo.New()
req, rec := jsonRequest(http.MethodGet, "/api/v1/wol-hosts", nil)
c := e.NewContext(req, rec)
err := APIListWolHosts(db)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusInternalServerError, rec.Code)
}
func TestAPIDeleteWolHost_DBError(t *testing.T) {
db := &errStore{}
e := echo.New()
req, rec := jsonRequest(http.MethodDelete, "/api/v1/wol-hosts/AA:BB:CC:DD:EE:FF", nil)
c := e.NewContext(req, rec)
c.SetParamNames("mac")
c.SetParamValues("AA:BB:CC:DD:EE:FF")
err := APIDeleteWolHost(db)(c)
require.NoError(t, err)
assert.Equal(t, http.StatusInternalServerError, rec.Code)
}
func TestAPIWakeHost_ExistingHost(t *testing.T) {
env := setupTestEnv(t)

93
handler/config_writer.go Normal file
View File

@ -0,0 +1,93 @@
package handler
import (
"io/fs"
"sync"
"time"
"github.com/labstack/gommon/log"
"github.com/DigitalTolk/wireguard-ui/store"
"github.com/DigitalTolk/wireguard-ui/util"
)
// ConfigWriter debounces WireGuard config writes so that rapid successive
// mutations (create, edit, delete, enable/disable) produce only one file write
// after a quiet period. This prevents overwhelming systemd path watchers
// (e.g. wgui.path with PathChanged) that restart wg-quick on every change.
type ConfigWriter struct {
mu sync.Mutex // protects timer
writeMu sync.Mutex // serializes actual file writes
timer *time.Timer
delay time.Duration
db store.IStore
tmplDir fs.FS
}
// NewConfigWriter creates a debounced config writer. The delay parameter
// controls how long to wait after the last Trigger() before writing.
// A typical value is 2 seconds — long enough to coalesce rapid changes,
// short enough that the config is applied promptly.
func NewConfigWriter(db store.IStore, tmplDir fs.FS, delay time.Duration) *ConfigWriter {
return &ConfigWriter{db: db, tmplDir: tmplDir, delay: delay}
}
// Trigger schedules a config write after the debounce delay. If called again
// before the delay expires, the timer resets. This is non-blocking.
func (cw *ConfigWriter) Trigger() {
cw.mu.Lock()
defer cw.mu.Unlock()
if cw.timer != nil {
cw.timer.Stop()
}
cw.timer = time.AfterFunc(cw.delay, func() {
if err := cw.apply(); err != nil {
log.Errorf("Auto-apply config failed: %v", err)
}
})
}
// ApplyNow cancels any pending debounced write and writes immediately.
// Returns an error if the write fails.
func (cw *ConfigWriter) ApplyNow() error {
cw.mu.Lock()
if cw.timer != nil {
cw.timer.Stop()
cw.timer = nil
}
cw.mu.Unlock()
return cw.apply()
}
func (cw *ConfigWriter) apply() error {
cw.writeMu.Lock()
defer cw.writeMu.Unlock()
server, err := cw.db.GetServer()
if err != nil {
return err
}
clients, err := cw.db.GetClients(false)
if err != nil {
return err
}
users, err := cw.db.GetUsers()
if err != nil {
return err
}
settings, err := cw.db.GetGlobalSettings()
if err != nil {
return err
}
if err := util.WriteWireGuardServerConfig(cw.tmplDir, server, clients, users, settings); err != nil {
return err
}
if err := util.UpdateHashes(cw.db); err != nil {
log.Warnf("Config written but hash update failed: %v", err)
}
log.Info("WireGuard config applied")
return nil
}

View File

@ -0,0 +1,208 @@
package handler
import (
"os"
"path/filepath"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/DigitalTolk/wireguard-ui/store/sqlitedb"
)
func TestConfigWriter_Trigger(t *testing.T) {
dir := t.TempDir()
dbPath := filepath.Join(dir, "test.db")
db, err := sqlitedb.New(dbPath)
require.NoError(t, err)
require.NoError(t, db.Init())
// Point config file to temp dir
gs, _ := db.GetGlobalSettings()
gs.ConfigFilePath = filepath.Join(dir, "wg0.conf")
db.SaveGlobalSettings(gs)
tmplFS := os.DirFS("../templates")
cw := NewConfigWriter(db, tmplFS, 100*time.Millisecond)
cw.Trigger()
// Wait for the debounce + a little extra
time.Sleep(300 * time.Millisecond)
_, err = os.Stat(filepath.Join(dir, "wg0.conf"))
assert.NoError(t, err, "Config file should be written after debounce")
}
func TestConfigWriter_Debounce(t *testing.T) {
dir := t.TempDir()
dbPath := filepath.Join(dir, "test.db")
db, err := sqlitedb.New(dbPath)
require.NoError(t, err)
require.NoError(t, db.Init())
gs, _ := db.GetGlobalSettings()
gs.ConfigFilePath = filepath.Join(dir, "wg0.conf")
db.SaveGlobalSettings(gs)
tmplFS := os.DirFS("../templates")
cw := NewConfigWriter(db, tmplFS, 200*time.Millisecond)
// Trigger rapidly — should coalesce
cw.Trigger()
time.Sleep(50 * time.Millisecond)
cw.Trigger()
time.Sleep(50 * time.Millisecond)
cw.Trigger()
// At this point, file should NOT exist yet (debounce hasn't fired)
_, err = os.Stat(filepath.Join(dir, "wg0.conf"))
assert.True(t, os.IsNotExist(err), "Config should not be written during debounce window")
// Wait for the final debounce to fire
time.Sleep(400 * time.Millisecond)
_, err = os.Stat(filepath.Join(dir, "wg0.conf"))
assert.NoError(t, err, "Config file should be written after debounce settles")
}
func TestConfigWriter_ApplyNow_CancelsPendingTimer(t *testing.T) {
dir := t.TempDir()
dbPath := filepath.Join(dir, "test.db")
db, err := sqlitedb.New(dbPath)
require.NoError(t, err)
require.NoError(t, db.Init())
gs, _ := db.GetGlobalSettings()
gs.ConfigFilePath = filepath.Join(dir, "wg0.conf")
db.SaveGlobalSettings(gs)
tmplFS := os.DirFS("../templates")
cw := NewConfigWriter(db, tmplFS, 24*time.Hour)
// Trigger a debounced write (won't fire for 24h)
cw.Trigger()
// ApplyNow should cancel the pending timer and write immediately
err = cw.ApplyNow()
require.NoError(t, err)
_, err = os.Stat(filepath.Join(dir, "wg0.conf"))
assert.NoError(t, err, "ApplyNow should write config immediately even with pending timer")
}
func TestConfigWriter_ApplyNow_NoTimer(t *testing.T) {
// Test ApplyNow when no timer has been set (cw.timer is nil)
dir := t.TempDir()
dbPath := filepath.Join(dir, "test.db")
db, err := sqlitedb.New(dbPath)
require.NoError(t, err)
require.NoError(t, db.Init())
gs, _ := db.GetGlobalSettings()
gs.ConfigFilePath = filepath.Join(dir, "wg0.conf")
db.SaveGlobalSettings(gs)
tmplFS := os.DirFS("../templates")
cw := NewConfigWriter(db, tmplFS, 24*time.Hour)
// ApplyNow without ever calling Trigger first (timer is nil)
err = cw.ApplyNow()
require.NoError(t, err)
_, err = os.Stat(filepath.Join(dir, "wg0.conf"))
assert.NoError(t, err, "ApplyNow should write config even when no timer was set")
}
func TestConfigWriter_Apply_InvalidConfigPath(t *testing.T) {
// Test apply() when config file path is unwritable
dir := t.TempDir()
dbPath := filepath.Join(dir, "test.db")
db, err := sqlitedb.New(dbPath)
require.NoError(t, err)
require.NoError(t, db.Init())
gs, _ := db.GetGlobalSettings()
gs.ConfigFilePath = "/dev/null/impossible/path/wg0.conf"
db.SaveGlobalSettings(gs)
tmplFS := os.DirFS("../templates")
cw := NewConfigWriter(db, tmplFS, 100*time.Millisecond)
err = cw.ApplyNow()
assert.Error(t, err, "ApplyNow should fail when config path is unwritable")
}
func TestConfigWriter_Trigger_ErrorInApply(t *testing.T) {
dir := t.TempDir()
dbPath := filepath.Join(dir, "test.db")
db, err := sqlitedb.New(dbPath)
require.NoError(t, err)
require.NoError(t, db.Init())
gs, _ := db.GetGlobalSettings()
gs.ConfigFilePath = "/dev/null/impossible/path/wg0.conf"
db.SaveGlobalSettings(gs)
tmplFS := os.DirFS("../templates")
cw := NewConfigWriter(db, tmplFS, 100*time.Millisecond)
// Trigger with a path that will cause apply() to fail
cw.Trigger()
// Wait for debounce to fire — apply() will fail and log the error
time.Sleep(300 * time.Millisecond)
// No assertion on error since it's logged, not returned.
// This test exercises the error path inside the AfterFunc callback.
}
func TestConfigWriter_TriggerMultipleTimes(t *testing.T) {
dir := t.TempDir()
dbPath := filepath.Join(dir, "test.db")
db, err := sqlitedb.New(dbPath)
require.NoError(t, err)
require.NoError(t, db.Init())
gs, _ := db.GetGlobalSettings()
gs.ConfigFilePath = filepath.Join(dir, "wg0.conf")
db.SaveGlobalSettings(gs)
tmplFS := os.DirFS("../templates")
cw := NewConfigWriter(db, tmplFS, 100*time.Millisecond)
// Trigger multiple times rapidly, then let debounce settle
for i := 0; i < 10; i++ {
cw.Trigger()
}
// Wait for debounce to fire
time.Sleep(300 * time.Millisecond)
_, err = os.Stat(filepath.Join(dir, "wg0.conf"))
assert.NoError(t, err, "Config should be written after rapid triggers settle")
}
func TestConfigWriter_ApplyNow(t *testing.T) {
dir := t.TempDir()
dbPath := filepath.Join(dir, "test.db")
db, err := sqlitedb.New(dbPath)
require.NoError(t, err)
require.NoError(t, db.Init())
gs, _ := db.GetGlobalSettings()
gs.ConfigFilePath = filepath.Join(dir, "wg0.conf")
db.SaveGlobalSettings(gs)
tmplFS := os.DirFS("../templates")
cw := NewConfigWriter(db, tmplFS, 24*time.Hour) // very long debounce
// ApplyNow should write immediately
err = cw.ApplyNow()
require.NoError(t, err)
_, err = os.Stat(filepath.Join(dir, "wg0.conf"))
assert.NoError(t, err, "ApplyNow should write config immediately")
}

View File

@ -2,12 +2,15 @@ package handler
import (
"encoding/json"
"fmt"
"io/fs"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"testing"
"time"
"github.com/gorilla/sessions"
"github.com/labstack/echo-contrib/session"
@ -15,14 +18,47 @@ import (
"github.com/stretchr/testify/require"
"github.com/DigitalTolk/wireguard-ui/audit"
"github.com/DigitalTolk/wireguard-ui/model"
"github.com/DigitalTolk/wireguard-ui/store/sqlitedb"
"github.com/DigitalTolk/wireguard-ui/util"
)
// errStore is a mock store that returns errors for all read methods.
// This is used to test error paths in handler functions.
type errStore struct{}
func (e *errStore) Init() error { return fmt.Errorf("db error") }
func (e *errStore) GetUsers() ([]model.User, error) { return nil, fmt.Errorf("db error") }
func (e *errStore) GetUserByName(string) (model.User, error) { return model.User{}, fmt.Errorf("db error") }
func (e *errStore) GetUserByOIDCSub(string) (model.User, error) { return model.User{}, fmt.Errorf("db error") }
func (e *errStore) SaveUser(model.User) error { return fmt.Errorf("db error") }
func (e *errStore) DeleteUser(string) error { return fmt.Errorf("db error") }
func (e *errStore) GetGlobalSettings() (model.GlobalSetting, error) { return model.GlobalSetting{}, fmt.Errorf("db error") }
func (e *errStore) GetServer() (model.Server, error) { return model.Server{}, fmt.Errorf("db error") }
func (e *errStore) GetClients(bool) ([]model.ClientData, error) { return nil, fmt.Errorf("db error") }
func (e *errStore) GetClientByID(string, model.QRCodeSettings) (model.ClientData, error) {
return model.ClientData{}, fmt.Errorf("db error")
}
func (e *errStore) SaveClient(model.Client) error { return fmt.Errorf("db error") }
func (e *errStore) DeleteClient(string) error { return fmt.Errorf("db error") }
func (e *errStore) SaveServerInterface(model.ServerInterface) error { return fmt.Errorf("db error") }
func (e *errStore) SaveServerKeyPair(model.ServerKeypair) error { return fmt.Errorf("db error") }
func (e *errStore) SaveGlobalSettings(model.GlobalSetting) error { return fmt.Errorf("db error") }
func (e *errStore) GetAllocatedIPs(string) ([]string, error) { return nil, fmt.Errorf("db error") }
func (e *errStore) GetWakeOnLanHosts() ([]model.WakeOnLanHost, error) { return nil, fmt.Errorf("db error") }
func (e *errStore) GetWakeOnLanHost(string) (*model.WakeOnLanHost, error) { return nil, fmt.Errorf("db error") }
func (e *errStore) DeleteWakeOnHostLanHost(string) error { return fmt.Errorf("db error") }
func (e *errStore) SaveWakeOnLanHost(model.WakeOnLanHost) error { return fmt.Errorf("db error") }
func (e *errStore) DeleteWakeOnHost(model.WakeOnLanHost) error { return fmt.Errorf("db error") }
func (e *errStore) GetPath() string { return "/tmp" }
func (e *errStore) SaveHashes(model.ClientServerHashes) error { return fmt.Errorf("db error") }
func (e *errStore) GetHashes() (model.ClientServerHashes, error) { return model.ClientServerHashes{}, fmt.Errorf("db error") }
type testEnv struct {
db *sqlitedb.SqliteDB
auditLog *audit.Logger
echo *echo.Echo
cw *ConfigWriter
}
func setupTestEnv(t *testing.T) *testEnv {
@ -48,7 +84,16 @@ func setupTestEnv(t *testing.T) *testEnv {
util.DisableLogin = true // simplify testing
return &testEnv{db: db, auditLog: auditLog, echo: e}
// config writer with very long delay so tests don't trigger real writes
tmplFS := fs.FS(os.DirFS(filepath.Join("..", "templates")))
cw := NewConfigWriter(db, tmplFS, 24*time.Hour)
// set config file path to temp dir so any accidental writes don't fail
gs, _ := db.GetGlobalSettings()
gs.ConfigFilePath = filepath.Join(dir, "wg0.conf")
db.SaveGlobalSettings(gs)
return &testEnv{db: db, auditLog: auditLog, echo: e, cw: cw}
}
func jsonRequest(method, path string, body interface{}) (*http.Request, *httptest.ResponseRecorder) {

View File

@ -742,6 +742,441 @@ func TestValidSession_POST_NoNextURL(t *testing.T) {
assert.Contains(t, location, "/login")
}
// --- doRefreshSession: session token mismatch ---
func TestDoRefreshSession_TokenMismatch(t *testing.T) {
origDisable := util.DisableLogin
util.DisableLogin = false
defer func() { util.DisableLogin = origDisable }()
env := setupTestEnv(t)
util.DisableLogin = false
// Create a remember-me session
env.echo.GET("/create-for-mismatch", func(c echo.Context) error {
createSession(c, "admin", true, uint32(12345), true)
return c.String(http.StatusOK, "ok")
})
req1, rec1 := jsonRequest(http.MethodGet, "/create-for-mismatch", nil)
env.echo.ServeHTTP(rec1, req1)
require.Equal(t, http.StatusOK, rec1.Code)
// Now tamper with the session_token cookie
env.echo.GET("/do-refresh-mismatch", func(c echo.Context) error {
doRefreshSession(c)
return c.String(http.StatusOK, "ok")
})
cookies := rec1.Result().Cookies()
req2, rec2 := jsonRequest(http.MethodGet, "/do-refresh-mismatch", nil)
for _, cookie := range cookies {
if cookie.Name == "session_token" {
cookie.Value = "tampered-token"
}
req2.AddCookie(cookie)
}
env.echo.ServeHTTP(rec2, req2)
assert.Equal(t, http.StatusOK, rec2.Code)
// The handler still returns OK, but the session was not refreshed
}
// --- doRefreshSession: no cookie at all ---
func TestDoRefreshSession_NoCookie(t *testing.T) {
origDisable := util.DisableLogin
util.DisableLogin = false
defer func() { util.DisableLogin = origDisable }()
env := setupTestEnv(t)
util.DisableLogin = false
env.echo.GET("/do-refresh-no-cookie", func(c echo.Context) error {
doRefreshSession(c)
return c.String(http.StatusOK, "ok")
})
req, rec := jsonRequest(http.MethodGet, "/do-refresh-no-cookie", nil)
env.echo.ServeHTTP(rec, req)
assert.Equal(t, http.StatusOK, rec.Code)
}
// --- doRefreshSession: session past max duration ---
func TestDoRefreshSession_PastMaxDuration(t *testing.T) {
origDisable := util.DisableLogin
util.DisableLogin = false
origMaxDuration := util.SessionMaxDuration
util.SessionMaxDuration = 100 // 100 seconds
defer func() {
util.DisableLogin = origDisable
util.SessionMaxDuration = origMaxDuration
}()
env := setupTestEnv(t)
util.DisableLogin = false
// Create a session and manipulate it to be past max duration
env.echo.GET("/create-expired", func(c echo.Context) error {
createSession(c, "admin", true, uint32(12345), true)
// Manipulate session: created long ago, past max duration
sess, _ := session.Get("session", c)
now := time.Now().UTC().Unix()
sess.Values["created_at"] = now - 200 // created 200s ago, max is 100s
sess.Values["updated_at"] = now - 90 // updated 90s ago
sess.Save(c.Request(), c.Response())
return c.String(http.StatusOK, "ok")
})
req1, rec1 := jsonRequest(http.MethodGet, "/create-expired", nil)
env.echo.ServeHTTP(rec1, req1)
require.Equal(t, http.StatusOK, rec1.Code)
env.echo.GET("/do-refresh-expired", func(c echo.Context) error {
doRefreshSession(c)
return c.String(http.StatusOK, "ok")
})
cookies := rec1.Result().Cookies()
req2, rec2 := jsonRequest(http.MethodGet, "/do-refresh-expired", nil)
for _, cookie := range cookies {
req2.AddCookie(cookie)
}
env.echo.ServeHTTP(rec2, req2)
assert.Equal(t, http.StatusOK, rec2.Code)
}
// --- doRefreshSession: updatedAt is in the future (corrupted) ---
func TestDoRefreshSession_FutureUpdatedAt(t *testing.T) {
origDisable := util.DisableLogin
util.DisableLogin = false
origMaxDuration := util.SessionMaxDuration
util.SessionMaxDuration = 86400 * 90
defer func() {
util.DisableLogin = origDisable
util.SessionMaxDuration = origMaxDuration
}()
env := setupTestEnv(t)
util.DisableLogin = false
env.echo.GET("/create-future", func(c echo.Context) error {
createSession(c, "admin", true, uint32(12345), true)
sess, _ := session.Get("session", c)
now := time.Now().UTC().Unix()
sess.Values["created_at"] = now - 172800
sess.Values["updated_at"] = now + 3600 // future timestamp
sess.Save(c.Request(), c.Response())
return c.String(http.StatusOK, "ok")
})
req1, rec1 := jsonRequest(http.MethodGet, "/create-future", nil)
env.echo.ServeHTTP(rec1, req1)
require.Equal(t, http.StatusOK, rec1.Code)
env.echo.GET("/do-refresh-future", func(c echo.Context) error {
doRefreshSession(c)
return c.String(http.StatusOK, "ok")
})
cookies := rec1.Result().Cookies()
req2, rec2 := jsonRequest(http.MethodGet, "/do-refresh-future", nil)
for _, cookie := range cookies {
req2.AddCookie(cookie)
}
env.echo.ServeHTTP(rec2, req2)
assert.Equal(t, http.StatusOK, rec2.Code)
}
// --- doRefreshSession: session expired (updatedAt + maxAge < now) ---
func TestDoRefreshSession_Expired(t *testing.T) {
origDisable := util.DisableLogin
util.DisableLogin = false
origMaxDuration := util.SessionMaxDuration
util.SessionMaxDuration = 86400 * 90
defer func() {
util.DisableLogin = origDisable
util.SessionMaxDuration = origMaxDuration
}()
env := setupTestEnv(t)
util.DisableLogin = false
env.echo.GET("/create-for-expire-test", func(c echo.Context) error {
createSession(c, "admin", true, uint32(12345), true)
// Manipulate session so that it's expired: updatedAt + maxAge < now
sess, _ := session.Get("session", c)
now := time.Now().UTC().Unix()
maxAge := sess.Values["max_age"].(int)
sess.Values["created_at"] = now - int64(maxAge) - 200
sess.Values["updated_at"] = now - int64(maxAge) - 100 // expired 100s ago
sess.Save(c.Request(), c.Response())
return c.String(http.StatusOK, "ok")
})
req1, rec1 := jsonRequest(http.MethodGet, "/create-for-expire-test", nil)
env.echo.ServeHTTP(rec1, req1)
require.Equal(t, http.StatusOK, rec1.Code)
env.echo.GET("/do-refresh-expire-test", func(c echo.Context) error {
doRefreshSession(c)
return c.String(http.StatusOK, "ok")
})
cookies := rec1.Result().Cookies()
req2, rec2 := jsonRequest(http.MethodGet, "/do-refresh-expire-test", nil)
for _, cookie := range cookies {
req2.AddCookie(cookie)
}
env.echo.ServeHTTP(rec2, req2)
assert.Equal(t, http.StatusOK, rec2.Code)
}
// --- createSession with custom SessionMaxDuration ---
func TestCreateSession_WithSessionMaxDuration(t *testing.T) {
origDisable := util.DisableLogin
util.DisableLogin = false
origMaxDuration := util.SessionMaxDuration
util.SessionMaxDuration = 86400 * 30 // 30 days
defer func() {
util.DisableLogin = origDisable
util.SessionMaxDuration = origMaxDuration
}()
env := setupTestEnv(t)
util.DisableLogin = false
env.echo.GET("/create-with-duration", func(c echo.Context) error {
createSession(c, "duruser", true, uint32(44444), true)
return c.String(http.StatusOK, "ok")
})
req, rec := jsonRequest(http.MethodGet, "/create-with-duration", nil)
env.echo.ServeHTTP(rec, req)
assert.Equal(t, http.StatusOK, rec.Code)
// Verify cookie MaxAge is set to SessionMaxDuration
for _, cookie := range rec.Result().Cookies() {
if cookie.Name == "session_token" {
assert.Equal(t, int(util.SessionMaxDuration), cookie.MaxAge,
"Cookie MaxAge should equal SessionMaxDuration when rememberMe is true")
break
}
}
}
// --- isAdmin with non-admin session ---
func TestIsAdmin_WithNonAdminSession(t *testing.T) {
origDisable := util.DisableLogin
util.DisableLogin = false
defer func() { util.DisableLogin = origDisable }()
env := setupTestEnv(t)
util.DisableLogin = false
env.echo.GET("/setup-nonadmin", func(c echo.Context) error {
createSession(c, "regular", false, uint32(55555), false)
return c.String(http.StatusOK, "ok")
})
req1, rec1 := jsonRequest(http.MethodGet, "/setup-nonadmin", nil)
env.echo.ServeHTTP(rec1, req1)
var adminResult bool
env.echo.GET("/check-admin", func(c echo.Context) error {
adminResult = isAdmin(c)
return c.String(http.StatusOK, "ok")
})
cookies := rec1.Result().Cookies()
req2, rec2 := jsonRequest(http.MethodGet, "/check-admin", nil)
for _, cookie := range cookies {
req2.AddCookie(cookie)
}
env.echo.ServeHTTP(rec2, req2)
assert.False(t, adminResult, "Non-admin session should return false for isAdmin")
}
// --- isAdmin with admin session ---
func TestIsAdmin_WithAdminSession(t *testing.T) {
origDisable := util.DisableLogin
util.DisableLogin = false
defer func() { util.DisableLogin = origDisable }()
env := setupTestEnv(t)
util.DisableLogin = false
env.echo.GET("/setup-admin-check", func(c echo.Context) error {
createSession(c, "adminuser", true, uint32(66666), false)
return c.String(http.StatusOK, "ok")
})
req1, rec1 := jsonRequest(http.MethodGet, "/setup-admin-check", nil)
env.echo.ServeHTTP(rec1, req1)
var adminResult bool
env.echo.GET("/check-admin2", func(c echo.Context) error {
adminResult = isAdmin(c)
return c.String(http.StatusOK, "ok")
})
cookies := rec1.Result().Cookies()
req2, rec2 := jsonRequest(http.MethodGet, "/check-admin2", nil)
for _, cookie := range cookies {
req2.AddCookie(cookie)
}
env.echo.ServeHTTP(rec2, req2)
assert.True(t, adminResult, "Admin session should return true for isAdmin")
}
// --- isValidSession: user not in CRC32 map ---
func TestIsValidSession_UserRemovedFromDB(t *testing.T) {
origDisable := util.DisableLogin
util.DisableLogin = false
defer func() { util.DisableLogin = origDisable }()
env := setupTestEnv(t)
util.DisableLogin = false
util.DBUsersToCRC32Mutex.Lock()
util.DBUsersToCRC32["tempuser"] = uint32(54321)
util.DBUsersToCRC32Mutex.Unlock()
env.echo.GET("/create-temp-session", func(c echo.Context) error {
createSession(c, "tempuser", false, uint32(54321), true)
return c.String(http.StatusOK, "ok")
})
req1, rec1 := jsonRequest(http.MethodGet, "/create-temp-session", nil)
env.echo.ServeHTTP(rec1, req1)
require.Equal(t, http.StatusOK, rec1.Code)
// Remove user from CRC32 map (simulates user deletion)
util.DBUsersToCRC32Mutex.Lock()
delete(util.DBUsersToCRC32, "tempuser")
util.DBUsersToCRC32Mutex.Unlock()
var valid bool
env.echo.GET("/validate-removed-user", func(c echo.Context) error {
valid = isValidSession(c)
return c.String(http.StatusOK, "ok")
})
cookies := rec1.Result().Cookies()
req2, rec2 := jsonRequest(http.MethodGet, "/validate-removed-user", nil)
for _, cookie := range cookies {
req2.AddCookie(cookie)
}
env.echo.ServeHTTP(rec2, req2)
assert.False(t, valid, "Session should be invalid when user is removed from DB")
}
// --- isValidSession: temporary session (maxAge=0) within 24h ---
func TestIsValidSession_TemporarySession(t *testing.T) {
origDisable := util.DisableLogin
util.DisableLogin = false
defer func() { util.DisableLogin = origDisable }()
env := setupTestEnv(t)
util.DisableLogin = false
util.DBUsersToCRC32Mutex.Lock()
util.DBUsersToCRC32["tempsess"] = uint32(11111)
util.DBUsersToCRC32Mutex.Unlock()
defer func() {
util.DBUsersToCRC32Mutex.Lock()
delete(util.DBUsersToCRC32, "tempsess")
util.DBUsersToCRC32Mutex.Unlock()
}()
// Create session without remember-me (maxAge=0)
env.echo.GET("/create-temp-sess", func(c echo.Context) error {
createSession(c, "tempsess", false, uint32(11111), false)
return c.String(http.StatusOK, "ok")
})
req1, rec1 := jsonRequest(http.MethodGet, "/create-temp-sess", nil)
env.echo.ServeHTTP(rec1, req1)
require.Equal(t, http.StatusOK, rec1.Code)
var valid bool
env.echo.GET("/validate-temp-sess", func(c echo.Context) error {
valid = isValidSession(c)
return c.String(http.StatusOK, "ok")
})
cookies := rec1.Result().Cookies()
req2, rec2 := jsonRequest(http.MethodGet, "/validate-temp-sess", nil)
for _, cookie := range cookies {
req2.AddCookie(cookie)
}
env.echo.ServeHTTP(rec2, req2)
assert.True(t, valid, "Temporary session should be valid within 24h virtual expiration")
}
// --- clearSession clears a valid session ---
func TestClearSession_ThenInvalid(t *testing.T) {
origDisable := util.DisableLogin
util.DisableLogin = false
defer func() { util.DisableLogin = origDisable }()
env := setupTestEnv(t)
util.DisableLogin = false
util.DBUsersToCRC32Mutex.Lock()
util.DBUsersToCRC32["clearme"] = uint32(22222)
util.DBUsersToCRC32Mutex.Unlock()
defer func() {
util.DBUsersToCRC32Mutex.Lock()
delete(util.DBUsersToCRC32, "clearme")
util.DBUsersToCRC32Mutex.Unlock()
}()
env.echo.GET("/create-clear-session", func(c echo.Context) error {
createSession(c, "clearme", true, uint32(22222), true)
return c.String(http.StatusOK, "ok")
})
req1, rec1 := jsonRequest(http.MethodGet, "/create-clear-session", nil)
env.echo.ServeHTTP(rec1, req1)
require.Equal(t, http.StatusOK, rec1.Code)
cookies := rec1.Result().Cookies()
// Now clear it
env.echo.GET("/do-clear-session", func(c echo.Context) error {
clearSession(c)
return c.String(http.StatusOK, "ok")
})
req2, rec2 := jsonRequest(http.MethodGet, "/do-clear-session", nil)
for _, cookie := range cookies {
req2.AddCookie(cookie)
}
env.echo.ServeHTTP(rec2, req2)
assert.Equal(t, http.StatusOK, rec2.Code)
// After clearing, the session_token cookie should have MaxAge=-1
for _, cookie := range rec2.Result().Cookies() {
if cookie.Name == "session_token" {
assert.Equal(t, -1, cookie.MaxAge, "session_token should be expired after clear")
}
}
}
func TestNeedsAdmin_WithAdminSession(t *testing.T) {
origDisable := util.DisableLogin
util.DisableLogin = false

View File

@ -45,7 +45,7 @@ var (
flagEmailFrom string
flagEmailFromName = "WireGuard UI"
flagSessionSecret = util.RandomString(32)
flagSessionMaxDuration = 90
flagSessionMaxDuration = 1
flagWgConfTemplate string
flagBasePath string
flagSubnetRanges string
@ -217,12 +217,16 @@ func main() {
// set up Echo with session middleware
app := router.New(util.SessionSecret)
// debounced config writer (coalesces rapid mutations into a single wg0.conf write)
configApplyDelay := time.Duration(util.LookupEnvOrInt(util.ConfigApplyDelayEnvVar, 3)) * time.Second
cw := handler.NewConfigWriter(db, tmplDir, configApplyDelay)
// audit logger
auditLog := audit.NewLogger(db.DB())
// API v1 routes
apiV1 := app.Group(util.BasePath+"/api/v1", handler.WithAuditLogger(auditLog))
router.RegisterAPIv1(apiV1, db, sendmail, tmplDir, defaultEmailSubject, defaultEmailContent, appVersion, gitCommit, auditLog)
router.RegisterAPIv1(apiV1, db, sendmail, cw, defaultEmailSubject, defaultEmailContent, appVersion, gitCommit, auditLog)
// OIDC SSO routes
oidcProvider, err := handler.NewOIDCProvider()

View File

@ -2,8 +2,7 @@ package model
// ClientDefaults Defaults for creation of new clients used in the templates
type ClientDefaults struct {
AllowedIps []string
ExtraAllowedIps []string
UseServerDNS bool
EnableAfterCreation bool
AllowedIps []string
ExtraAllowedIps []string
UseServerDNS bool
}

View File

@ -1,8 +1,6 @@
package router
import (
"io/fs"
"github.com/labstack/echo/v4"
"github.com/DigitalTolk/wireguard-ui/audit"
@ -12,21 +10,21 @@ import (
)
// RegisterAPIv1 registers all API v1 routes under the given group
func RegisterAPIv1(g *echo.Group, db store.IStore, mailer emailer.Emailer, tmplDir fs.FS, emailSubject, emailContent, appVersion, gitCommit string, auditLog *audit.Logger) {
func RegisterAPIv1(g *echo.Group, db store.IStore, mailer emailer.Emailer, cw *handler.ConfigWriter, emailSubject, emailContent, appVersion, gitCommit string, auditLog *audit.Logger) {
// Auth
g.GET("/auth/me", handler.APIGetMe(db), handler.APIAuth)
g.POST("/auth/logout", handler.APILogout(), handler.APIAuth)
g.GET("/auth/info", handler.APIAppInfo(appVersion, gitCommit))
// Clients
// Clients (read endpoints use APIAuth — non-admins can access their own)
clients := g.Group("/clients", handler.APIAuth)
clients.GET("", handler.APIListClients(db))
clients.GET("/export", handler.APIExportClients(db))
clients.GET("/export", handler.APIExportClients(db), handler.APIAdmin)
clients.GET("/:id", handler.APIGetClient(db))
clients.POST("", handler.APICreateClient(db), handler.ContentTypeJson)
clients.PUT("/:id", handler.APIUpdateClient(db), handler.ContentTypeJson)
clients.PATCH("/:id/status", handler.APIPatchClientStatus(db), handler.ContentTypeJson)
clients.DELETE("/:id", handler.APIDeleteClient(db))
clients.POST("", handler.APICreateClient(db, cw), handler.APIAdmin, handler.ContentTypeJson)
clients.PUT("/:id", handler.APIUpdateClient(db, cw), handler.APIAdmin, handler.ContentTypeJson)
clients.PATCH("/:id/status", handler.APIPatchClientStatus(db, cw), handler.APIAdmin, handler.ContentTypeJson)
clients.DELETE("/:id", handler.APIDeleteClient(db, cw), handler.APIAdmin)
clients.GET("/:id/config", handler.APIDownloadClientConfig(db))
clients.GET("/:id/qrcode", handler.APIGetClientQRCode(db))
clients.POST("/:id/email", handler.APIEmailClient(db, mailer, emailSubject, emailContent), handler.ContentTypeJson)
@ -34,15 +32,15 @@ func RegisterAPIv1(g *echo.Group, db store.IStore, mailer emailer.Emailer, tmplD
// Server (admin only)
server := g.Group("/server", handler.APIAuth, handler.APIAdmin)
server.GET("", handler.APIGetServer(db))
server.PUT("/interface", handler.APIUpdateServerInterface(db), handler.ContentTypeJson)
server.POST("/keypair", handler.APIRegenerateServerKeypair(db), handler.ContentTypeJson)
server.POST("/apply-config", handler.APIApplyServerConfig(db, tmplDir), handler.ContentTypeJson)
server.PUT("/interface", handler.APIUpdateServerInterface(db, cw), handler.ContentTypeJson)
server.POST("/keypair", handler.APIRegenerateServerKeypair(db, cw), handler.ContentTypeJson)
server.POST("/apply-config", handler.APIApplyServerConfig(cw), handler.ContentTypeJson)
server.GET("/config-status", handler.APIConfigStatus(db))
// Settings (admin only)
settings := g.Group("/settings", handler.APIAuth, handler.APIAdmin)
settings.GET("", handler.APIGetSettings(db))
settings.PUT("", handler.APIUpdateSettings(db), handler.ContentTypeJson)
settings.PUT("", handler.APIUpdateSettings(db, cw), handler.ContentTypeJson)
// Users (admin only for list/create/delete)
// Users (read-only — managed via SSO)
@ -50,21 +48,21 @@ func RegisterAPIv1(g *echo.Group, db store.IStore, mailer emailer.Emailer, tmplD
users.GET("", handler.APIListUsers(db))
users.GET("/:username", handler.APIGetUser(db))
// Wake-on-LAN
wolGroup := g.Group("/wol-hosts", handler.APIAuth)
// Wake-on-LAN (admin only)
wolGroup := g.Group("/wol-hosts", handler.APIAuth, handler.APIAdmin)
wolGroup.GET("", handler.APIListWolHosts(db))
wolGroup.POST("", handler.APISaveWolHost(db), handler.ContentTypeJson)
wolGroup.DELETE("/:mac", handler.APIDeleteWolHost(db))
wolGroup.POST("/:mac/wake", handler.APIWakeHost(db), handler.ContentTypeJson)
// Utilities
utils := g.Group("", handler.APIAuth)
// Utilities (admin only)
utils := g.Group("", handler.APIAuth, handler.APIAdmin)
utils.GET("/machine-ips", handler.APIMachineIPs())
utils.GET("/subnet-ranges", handler.APISubnetRanges())
utils.GET("/suggest-client-ips", handler.APISuggestClientIPs(db))
// Status
g.GET("/status", handler.APIServerStatus(db), handler.APIAuth)
// Status (admin only)
g.GET("/status", handler.APIServerStatus(db), handler.APIAuth, handler.APIAdmin)
// Audit logs (admin only)
auditGroup := g.Group("/audit-logs", handler.APIAuth, handler.APIAdmin)

View File

@ -6,12 +6,14 @@ import (
"os"
"path/filepath"
"testing"
"time"
"github.com/labstack/echo/v4"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/DigitalTolk/wireguard-ui/audit"
"github.com/DigitalTolk/wireguard-ui/handler"
"github.com/DigitalTolk/wireguard-ui/store/sqlitedb"
"github.com/DigitalTolk/wireguard-ui/util"
)
@ -101,9 +103,10 @@ func TestRegisterAPIv1_RoutesRegistered(t *testing.T) {
e := New(secret)
tmplFS := os.DirFS("../templates")
cw := handler.NewConfigWriter(db, tmplFS, 24*time.Hour)
g := e.Group("/api/v1")
RegisterAPIv1(g, db, nil, tmplFS, "", "", "dev", "test", auditLog)
RegisterAPIv1(g, db, nil, cw, "", "", "dev", "test", auditLog)
routes := e.Routes()
@ -160,9 +163,10 @@ func TestRegisterAPIv1_HealthEndpointWorks(t *testing.T) {
e := New(secret)
tmplFS := os.DirFS("../templates")
cw := handler.NewConfigWriter(db, tmplFS, 24*time.Hour)
g := e.Group("/api/v1")
RegisterAPIv1(g, db, nil, tmplFS, "", "", "dev", "test", auditLog)
RegisterAPIv1(g, db, nil, cw, "", "", "dev", "test", auditLog)
// Test the auth/info endpoint which requires no auth
req := httptest.NewRequest(http.MethodGet, "/api/v1/auth/info", nil)

View File

@ -0,0 +1,228 @@
import { describe, it, expect, afterEach, vi } from "vitest";
import { screen, waitFor } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { renderWithProviders, mockFetch } from "@/test/test-utils";
import { AppShell } from "./AppShell";
const adminMe = {
username: "admin",
email: "admin@test.com",
display_name: "Admin User",
admin: true,
};
const regularMe = {
username: "user1",
email: "user1@test.com",
display_name: "Regular User",
admin: false,
};
describe("AppShell", () => {
let cleanup: () => void;
afterEach(() => {
cleanup?.();
});
it("renders the WireGuard UI title", async () => {
cleanup = mockFetch({ "/auth/me": adminMe });
renderWithProviders(<AppShell />);
await waitFor(() => {
expect(screen.getAllByText("WireGuard UI").length).toBeGreaterThan(0);
});
});
it("shows admin nav items for admin users", async () => {
cleanup = mockFetch({ "/auth/me": adminMe });
renderWithProviders(<AppShell />);
await waitFor(() => {
expect(screen.getByText("Server")).toBeInTheDocument();
expect(screen.getByText("Settings")).toBeInTheDocument();
expect(screen.getByText("Users")).toBeInTheDocument();
expect(screen.getByText("Audit Logs")).toBeInTheDocument();
expect(screen.getByText("Wake-on-LAN")).toBeInTheDocument();
expect(screen.getByText("Status")).toBeInTheDocument();
});
});
it("hides admin nav items for non-admin users", async () => {
cleanup = mockFetch({ "/auth/me": regularMe });
renderWithProviders(<AppShell />);
await waitFor(() => {
expect(screen.getByText("Clients")).toBeInTheDocument();
expect(screen.getByText("About")).toBeInTheDocument();
});
expect(screen.queryByText("Server")).not.toBeInTheDocument();
expect(screen.queryByText("Settings")).not.toBeInTheDocument();
expect(screen.queryByText("Users")).not.toBeInTheDocument();
expect(screen.queryByText("Audit Logs")).not.toBeInTheDocument();
expect(screen.queryByText("Wake-on-LAN")).not.toBeInTheDocument();
expect(screen.queryByText("Status")).not.toBeInTheDocument();
});
it("shows common nav items for all users", async () => {
cleanup = mockFetch({ "/auth/me": regularMe });
renderWithProviders(<AppShell />);
await waitFor(() => {
expect(screen.getByText("Clients")).toBeInTheDocument();
expect(screen.getByText("About")).toBeInTheDocument();
});
});
it("displays admin user display_name and email", async () => {
cleanup = mockFetch({ "/auth/me": adminMe });
renderWithProviders(<AppShell />);
await waitFor(() => {
expect(screen.getByText("Admin User")).toBeInTheDocument();
expect(screen.getByText("admin@test.com")).toBeInTheDocument();
});
});
it("displays non-admin user email without Admin badge", async () => {
cleanup = mockFetch({ "/auth/me": regularMe });
renderWithProviders(<AppShell />);
await waitFor(() => {
expect(screen.getByText("Regular User")).toBeInTheDocument();
expect(screen.getByText("user1@test.com")).toBeInTheDocument();
});
expect(screen.queryByText("Admin")).not.toBeInTheDocument();
});
it("shows Admin badge for admin users", async () => {
cleanup = mockFetch({ "/auth/me": adminMe });
renderWithProviders(<AppShell />);
await waitFor(() => {
expect(screen.getByText("Admin")).toBeInTheDocument();
});
});
it("falls back to username when display_name is empty", async () => {
const noDisplayName = { ...regularMe, display_name: "" };
cleanup = mockFetch({ "/auth/me": noDisplayName });
renderWithProviders(<AppShell />);
await waitFor(() => {
expect(screen.getByText("user1")).toBeInTheDocument();
});
});
it("shows mobile menu toggle button", async () => {
cleanup = mockFetch({ "/auth/me": adminMe });
renderWithProviders(<AppShell />);
await waitFor(() => {
expect(screen.getByLabelText("Open menu")).toBeInTheDocument();
});
});
it("toggles mobile sidebar open and closed", async () => {
const user = userEvent.setup();
cleanup = mockFetch({ "/auth/me": adminMe });
renderWithProviders(<AppShell />);
await waitFor(() => {
expect(screen.getByLabelText("Open menu")).toBeInTheDocument();
});
// Open menu
await user.click(screen.getByLabelText("Open menu"));
expect(screen.getByLabelText("Close menu")).toBeInTheDocument();
// Close menu
await user.click(screen.getByLabelText("Close menu"));
expect(screen.getByLabelText("Open menu")).toBeInTheDocument();
});
it("shows logout button", async () => {
cleanup = mockFetch({ "/auth/me": adminMe });
renderWithProviders(<AppShell />);
await waitFor(() => {
expect(screen.getByLabelText("Log out")).toBeInTheDocument();
});
});
it("calls logout API on logout button click", async () => {
const user = userEvent.setup();
cleanup = mockFetch({ "/auth/me": adminMe, "/auth/logout": {} });
// Mock location.href setter
const hrefSetter = vi.fn();
Object.defineProperty(window, "location", {
value: { ...window.location, href: "" },
writable: true,
});
Object.defineProperty(window.location, "href", {
set: hrefSetter,
get: () => "",
});
renderWithProviders(<AppShell />);
await waitFor(() => {
expect(screen.getByLabelText("Log out")).toBeInTheDocument();
});
await user.click(screen.getByLabelText("Log out"));
await waitFor(() => {
expect(hrefSetter).toHaveBeenCalledWith("./api/v1/auth/oidc/login");
});
});
it("has main navigation landmark", async () => {
cleanup = mockFetch({ "/auth/me": adminMe });
renderWithProviders(<AppShell />);
await waitFor(() => {
expect(screen.getByRole("navigation", { name: "Main navigation" })).toBeInTheDocument();
});
});
it("has main content landmark", async () => {
cleanup = mockFetch({ "/auth/me": adminMe });
renderWithProviders(<AppShell />);
await waitFor(() => {
expect(screen.getByRole("main")).toBeInTheDocument();
});
});
it("clicking a nav link closes the mobile sidebar", async () => {
const user = userEvent.setup();
cleanup = mockFetch({ "/auth/me": adminMe });
renderWithProviders(<AppShell />);
await waitFor(() => {
expect(screen.getByLabelText("Open menu")).toBeInTheDocument();
});
// Open sidebar
await user.click(screen.getByLabelText("Open menu"));
expect(screen.getByLabelText("Close menu")).toBeInTheDocument();
// Click a nav link (About is always visible)
await user.click(screen.getByText("About"));
// Sidebar should close
await waitFor(() => {
expect(screen.getByLabelText("Open menu")).toBeInTheDocument();
});
});
it("clicking the overlay closes the mobile sidebar", async () => {
const user = userEvent.setup();
cleanup = mockFetch({ "/auth/me": adminMe });
renderWithProviders(<AppShell />);
await waitFor(() => {
expect(screen.getByLabelText("Open menu")).toBeInTheDocument();
});
// Open sidebar
await user.click(screen.getByLabelText("Open menu"));
expect(screen.getByLabelText("Close menu")).toBeInTheDocument();
// Click the overlay (it has aria-hidden="true")
const overlay = document.querySelector(".fixed.inset-0.z-40");
expect(overlay).not.toBeNull();
await user.click(overlay!);
// Sidebar should close
await waitFor(() => {
expect(screen.getByLabelText("Open menu")).toBeInTheDocument();
});
});
});

View File

@ -21,11 +21,11 @@ import { apiPost } from "@/lib/api-client";
const navItems = [
{ to: "/", icon: Shield, label: "Clients", end: true },
{ to: "/status", icon: Monitor, label: "Status" },
{ to: "/status", icon: Monitor, label: "Status", admin: true },
{ to: "/server", icon: Server, label: "Server", admin: true },
{ to: "/settings", icon: Settings, label: "Settings", admin: true },
{ to: "/users", icon: Users, label: "Users", admin: true },
{ to: "/wol", icon: Wifi, label: "Wake-on-LAN" },
{ to: "/wol", icon: Wifi, label: "Wake-on-LAN", admin: true },
{ to: "/audit", icon: ClipboardList, label: "Audit Logs", admin: true },
{ to: "/about", icon: Info, label: "About" },
];
@ -80,10 +80,13 @@ export function AppShell() {
<Separator />
<div className="p-4">
<div className="flex items-center justify-between">
<div className="text-sm">
<div className="min-w-0 text-sm">
<div className="font-medium text-sidebar-foreground">
{me?.display_name || me?.username}
</div>
<div className="text-xs text-muted-foreground break-all">
{me?.email}
</div>
{me?.admin && (
<span className="text-xs text-muted-foreground">Admin</span>
)}

View File

@ -109,7 +109,6 @@ export interface ClientDefaults {
AllowedIps: string[];
ExtraAllowedIps: string[];
UseServerDNS: boolean;
EnableAfterCreation: boolean;
}
export interface AppInfo {

View File

@ -237,4 +237,182 @@ describe("AuditPage interactions", () => {
expect(screen.getByText("abc123")).toBeInTheDocument();
});
});
it("clicks Next pagination button", async () => {
const user = userEvent.setup();
cleanup = mockFetch({
...mockResponses,
"/audit-logs": {
data: [],
total: 100,
page: 1,
per_page: 50,
},
});
renderWithProviders(<AuditPage />);
await waitFor(() => {
expect(screen.getByText("Next")).toBeInTheDocument();
});
const nextButton = screen.getByText("Next").closest("button")!;
expect(nextButton).not.toBeDisabled();
await user.click(nextButton);
});
it("disables Next on last page", async () => {
cleanup = mockFetch({
...mockResponses,
"/audit-logs": {
data: [],
total: 10,
page: 1,
per_page: 50,
},
});
renderWithProviders(<AuditPage />);
await waitFor(() => {
const nextButton = screen.getByText("Next").closest("button");
expect(nextButton).toBeDisabled();
});
});
it("changes date to filter", async () => {
const user = userEvent.setup();
cleanup = mockFetch(mockResponses);
renderWithProviders(<AuditPage />);
await waitFor(() => {
expect(screen.getByText("Date To")).toBeInTheDocument();
});
const toInput = screen.getByLabelText("Date To");
await user.clear(toInput);
await user.type(toInput, "2026-12-31");
});
it("shows resource with email matching name", async () => {
const logEmailOnly = {
...logEntry,
id: 5,
details: '{"email":"alice@example.com"}',
};
cleanup = mockFetch({
...mockResponses,
"/audit-logs": {
data: [logEmailOnly],
total: 1,
page: 1,
per_page: 50,
},
});
renderWithProviders(<AuditPage />);
await waitFor(() => {
// name = email, so format is: name (resource_id)
expect(screen.getByText(/alice@example.com.*abc123/)).toBeInTheDocument();
});
});
it("exports with filters applied", async () => {
const user = userEvent.setup();
const openSpy = vi.spyOn(window, "open").mockImplementation(() => null);
cleanup = mockFetch(mockResponses);
renderWithProviders(<AuditPage />);
await waitFor(() => {
expect(screen.getByPlaceholderText("Name, email, or ID...")).toBeInTheDocument();
});
// Clear existing text, type fresh, then press Enter to apply filter
const searchInput = screen.getByPlaceholderText("Name, email, or ID...");
await user.clear(searchInput);
await user.type(searchInput, "myfilter{Enter}");
// Now export
await user.click(screen.getByText("Export to Excel"));
expect(openSpy).toHaveBeenCalledWith(
expect.stringContaining("search=myfilter"),
"_blank"
);
openSpy.mockRestore();
});
it("shows total count in pagination info", async () => {
cleanup = mockFetch({
...mockResponses,
"/audit-logs": {
data: [logEntry],
total: 42,
page: 1,
per_page: 50,
},
});
renderWithProviders(<AuditPage />);
await waitFor(() => {
expect(screen.getByText(/42 total/)).toBeInTheDocument();
});
});
it("renders Activity Log card title", async () => {
cleanup = mockFetch(mockResponses);
renderWithProviders(<AuditPage />);
await waitFor(() => {
expect(screen.getByText("Activity Log")).toBeInTheDocument();
});
});
it("clicks Previous button when on page 2", async () => {
const user = userEvent.setup();
cleanup = mockFetch({
...mockResponses,
"/audit-logs": {
data: [],
total: 100,
page: 2,
per_page: 50,
},
});
// Navigate to page=2
window.history.pushState({}, "", "?page=2");
renderWithProviders(<AuditPage />);
await waitFor(() => {
expect(screen.getByText("Previous")).toBeInTheDocument();
});
const prevButton = screen.getByText("Previous").closest("button")!;
expect(prevButton).not.toBeDisabled();
await user.click(prevButton);
// Clean up URL
window.history.pushState({}, "", "/");
});
it("clears a filter by setting it to empty value", async () => {
const user = userEvent.setup();
cleanup = mockFetch(mockResponses);
// Start with a search filter applied
window.history.pushState({}, "", "?search=alice");
renderWithProviders(<AuditPage />);
await waitFor(() => {
expect(screen.getByPlaceholderText("Name, email, or ID...")).toBeInTheDocument();
});
// Clear the search and press Enter
const searchInput = screen.getByPlaceholderText("Name, email, or ID...");
await user.clear(searchInput);
await user.type(searchInput, "{Enter}");
// Clean up URL
window.history.pushState({}, "", "/");
});
});

File diff suppressed because it is too large Load Diff

View File

@ -3,6 +3,9 @@ import { screen, waitFor } from "@testing-library/react";
import { renderWithProviders, mockFetch } from "@/test/test-utils";
import { ClientsPage } from "./ClientsPage";
const adminMe = { username: "admin", email: "admin@test.com", display_name: "Admin", admin: true };
const userMe = { username: "user", email: "user@test.com", display_name: "User", admin: false };
describe("ClientsPage", () => {
let cleanup: () => void;
@ -11,15 +14,15 @@ describe("ClientsPage", () => {
});
it("shows client list heading", async () => {
cleanup = mockFetch({ "/clients": [], "/subnet-ranges": [] });
cleanup = mockFetch({ "/auth/me": adminMe, "/clients": [], "/subnet-ranges": [] });
renderWithProviders(<ClientsPage />);
await waitFor(() => {
expect(screen.getByText("WireGuard Clients")).toBeInTheDocument();
});
});
it("shows empty state when no clients", async () => {
cleanup = mockFetch({ "/clients": [], "/subnet-ranges": [] });
it("shows empty state when no clients (admin)", async () => {
cleanup = mockFetch({ "/auth/me": adminMe, "/clients": [], "/subnet-ranges": [] });
renderWithProviders(<ClientsPage />);
await waitFor(() => {
@ -27,8 +30,18 @@ describe("ClientsPage", () => {
});
});
it("shows non-admin empty state when no clients", async () => {
cleanup = mockFetch({ "/auth/me": userMe, "/clients": [], "/subnet-ranges": [] });
renderWithProviders(<ClientsPage />);
await waitFor(() => {
expect(screen.getByText(/No client configurations found for your account/)).toBeInTheDocument();
});
});
it("renders clients from API", async () => {
cleanup = mockFetch({
"/auth/me": adminMe,
"/clients": [
{
Client: {
@ -57,6 +70,7 @@ describe("ClientsPage", () => {
it("shows enabled badge for enabled clients", async () => {
cleanup = mockFetch({
"/auth/me": adminMe,
"/clients": [
{
Client: {
@ -82,4 +96,26 @@ describe("ClientsPage", () => {
expect(screen.getByText("Enabled")).toBeInTheDocument();
});
});
it("hides New Client and Export buttons for non-admin users", async () => {
cleanup = mockFetch({ "/auth/me": userMe, "/clients": [], "/subnet-ranges": [] });
renderWithProviders(<ClientsPage />);
await waitFor(() => {
expect(screen.getByText("WireGuard Clients")).toBeInTheDocument();
});
expect(screen.queryByText("New Client")).not.toBeInTheDocument();
expect(screen.queryByText("Export to Excel")).not.toBeInTheDocument();
});
it("shows New Client and Export buttons for admin users", async () => {
cleanup = mockFetch({ "/auth/me": adminMe, "/clients": [], "/subnet-ranges": [] });
renderWithProviders(<ClientsPage />);
await waitFor(() => {
expect(screen.getByText("New Client")).toBeInTheDocument();
});
expect(screen.getByText("Export to Excel")).toBeInTheDocument();
});
});

View File

@ -7,6 +7,7 @@ import {
useSetClientStatus,
useDeleteClient,
} from "@/hooks/useClients";
import { useAuth } from "@/hooks/useAuth";
import { apiGet, apiPost, API_BASE } from "@/lib/api-client";
import { splitList } from "@/lib/utils";
import {
@ -40,80 +41,101 @@ import { Download, Mail, Pencil, Plus, QrCode, Search, Trash2 } from "lucide-rea
import { toast } from "sonner";
import type { Client, ClientData } from "@/lib/types";
interface EditFormState {
name: string;
public_key: string;
allocated_ips: string;
allowed_ips: string;
extra_allowed_ips: string;
endpoint: string;
additional_notes: string;
use_server_dns: boolean;
preshared_key: string;
}
function validateClientForm(form: {
name: string;
email?: string;
allocated_ips: string[];
allowed_ips: string[];
extra_allowed_ips?: string[];
allocated_ips: string;
allowed_ips: string;
extra_allowed_ips?: string;
endpoint?: string;
}, emailRequired: boolean): Record<string, string> {
const errors: Record<string, string> = {};
if (!form.name.trim()) {
errors.name = "Name is required";
}
if (!form.name.trim()) errors.name = "Name is required";
if (emailRequired) {
if (!form.email || !form.email.trim()) {
errors.email = "Email is required";
} else if (!isValidEmail(form.email)) {
errors.email = "Invalid email format";
}
} else {
if (form.email && form.email.trim() && !isValidEmail(form.email)) {
errors.email = "Invalid email format";
}
if (!form.email?.trim()) errors.email = "Email is required";
else if (!isValidEmail(form.email)) errors.email = "Invalid email format";
} else if (form.email?.trim() && !isValidEmail(form.email)) {
errors.email = "Invalid email format";
}
if (
form.allocated_ips.length === 0 ||
form.allocated_ips.every((ip) => !ip.trim())
) {
errors.allocated_ips = "At least one allocated IP is required";
} else if (!form.allocated_ips.every((ip) => !ip.trim() || isValidCIDR(ip))) {
errors.allocated_ips = "Each allocated IP must be valid CIDR (e.g. 10.0.0.2/32)";
const allocIPs = splitList(form.allocated_ips);
if (allocIPs.length === 0) errors.allocated_ips = "At least one allocated IP is required";
else if (!allocIPs.every(isValidCIDR)) errors.allocated_ips = "Each IP must be valid CIDR (e.g. 10.0.0.2/32)";
const allowIPs = splitList(form.allowed_ips);
if (allowIPs.length === 0) errors.allowed_ips = "At least one allowed IP is required";
else if (!allowIPs.every(isValidCIDR)) errors.allowed_ips = "Each IP must be valid CIDR (e.g. 0.0.0.0/0)";
const extraIPs = splitList(form.extra_allowed_ips ?? "");
if (extraIPs.length > 0 && !extraIPs.every(isValidCIDR)) {
errors.extra_allowed_ips = "Each IP must be valid CIDR";
}
if (
form.allowed_ips.length === 0 ||
form.allowed_ips.every((ip) => !ip.trim())
) {
errors.allowed_ips = "At least one allowed IP is required";
} else if (!form.allowed_ips.every((ip) => !ip.trim() || isValidCIDR(ip))) {
errors.allowed_ips = "Each allowed IP must be valid CIDR (e.g. 0.0.0.0/0)";
}
if (
form.extra_allowed_ips &&
form.extra_allowed_ips.some((ip) => ip.trim()) &&
!form.extra_allowed_ips.every((ip) => !ip.trim() || isValidCIDR(ip))
) {
errors.extra_allowed_ips =
"Each extra allowed IP must be valid CIDR (e.g. 192.168.1.0/24)";
}
if (form.endpoint && form.endpoint.trim() && !isValidEndpoint(form.endpoint)) {
errors.endpoint = "Must be host:port or IP:port (e.g. vpn.example.com:51820)";
if (form.endpoint?.trim() && !isValidEndpoint(form.endpoint)) {
errors.endpoint = "Must be host:port or IP:port";
}
return errors;
}
function QrCodeDialog({ client, onClose }: { client: ClientData | null; onClose: () => void }) {
const { data } = useQuery({
queryKey: ["client-qr", client?.Client.id],
queryFn: () => apiGet<{ qr_code: string }>(`/clients/${client!.Client.id}/qrcode`),
enabled: !!client,
staleTime: Infinity,
});
return (
<Dialog open={!!client} onOpenChange={() => onClose()}>
<DialogContent>
<DialogHeader>
<DialogTitle>{client?.Client.name} - QR Code</DialogTitle>
</DialogHeader>
{data?.qr_code ? (
<div className="flex justify-center p-4">
<img src={data.qr_code} alt={`QR code for ${client?.Client.name}`} className="max-w-[256px]" />
</div>
) : (
<div className="flex justify-center p-8">
<Skeleton className="h-64 w-64" />
</div>
)}
</DialogContent>
</Dialog>
);
}
const emptyCreateForm = {
name: "",
email: "",
public_key: "",
preshared_key: "",
allocated_ips: [] as string[],
allowed_ips: ["0.0.0.0/0"],
extra_allowed_ips: [] as string[],
allocated_ips: "",
allowed_ips: "0.0.0.0/0",
extra_allowed_ips: "",
use_server_dns: true,
enabled: true,
additional_notes: "",
};
export function ClientsPage() {
const { data: me } = useAuth();
const isAdminUser = me?.admin ?? false;
const [searchParams, setSearchParams] = useSearchParams();
const filterSearch = searchParams.get("search") || "";
@ -162,7 +184,11 @@ export function ClientsPage() {
const [subnetRange, setSubnetRange] = useState("");
const [editDialog, setEditDialog] = useState<Client | null>(null);
const [editForm, setEditForm] = useState<Partial<Client>>({});
const [editForm, setEditForm] = useState<EditFormState>({
name: "", public_key: "", allocated_ips: "", allowed_ips: "",
extra_allowed_ips: "", endpoint: "", additional_notes: "",
use_server_dns: true, preshared_key: "",
});
const [emailDialog, setEmailDialog] = useState<Client | null>(null);
const [emailAddress, setEmailAddress] = useState("");
@ -181,7 +207,7 @@ export function ClientsPage() {
if (!showCreate) return;
const sr = subnetRange || "";
apiGet<string[]>(`/suggest-client-ips${sr ? `?sr=${sr}` : ""}`)
.then((ips) => setNewClient((prev) => ({ ...prev, allocated_ips: ips })))
.then((ips) => setNewClient((prev) => ({ ...prev, allocated_ips: ips.join(", ") })))
.catch(() => {});
}, [subnetRange, showCreate]);
@ -195,10 +221,10 @@ export function ClientsPage() {
() =>
editDialog
? validateClientForm({
name: editForm.name ?? "",
name: editForm.name,
email: editDialog.email,
allocated_ips: editForm.allocated_ips ?? [],
allowed_ips: editForm.allowed_ips ?? [],
allocated_ips: editForm.allocated_ips,
allowed_ips: editForm.allowed_ips,
extra_allowed_ips: editForm.extra_allowed_ips,
endpoint: editForm.endpoint,
}, true)
@ -255,7 +281,13 @@ export function ClientsPage() {
};
const handleCreate = () => {
createClient.mutate(newClient, {
const payload = {
...newClient,
allocated_ips: splitList(newClient.allocated_ips),
allowed_ips: splitList(newClient.allowed_ips),
extra_allowed_ips: splitList(newClient.extra_allowed_ips),
};
createClient.mutate(payload, {
onSuccess: () => {
toast.success("Client created");
setShowCreate(false);
@ -268,9 +300,10 @@ export function ClientsPage() {
const handleOpenEdit = (client: Client) => {
setEditForm({
name: client.name,
allocated_ips: client.allocated_ips || [],
allowed_ips: client.allowed_ips || [],
extra_allowed_ips: client.extra_allowed_ips || [],
public_key: client.public_key,
allocated_ips: (client.allocated_ips || []).join(", "),
allowed_ips: (client.allowed_ips || []).join(", "),
extra_allowed_ips: (client.extra_allowed_ips || []).join(", "),
endpoint: client.endpoint,
additional_notes: client.additional_notes,
use_server_dns: client.use_server_dns,
@ -281,8 +314,14 @@ export function ClientsPage() {
const handleSaveEdit = () => {
if (!editDialog) return;
const payload = {
...editForm,
allocated_ips: splitList(editForm.allocated_ips),
allowed_ips: splitList(editForm.allowed_ips),
extra_allowed_ips: splitList(editForm.extra_allowed_ips),
};
updateClient.mutate(
{ id: editDialog.id, ...editForm },
{ id: editDialog.id, ...payload },
{
onSuccess: () => {
toast.success("Client updated");
@ -334,168 +373,151 @@ export function ClientsPage() {
</h2>
<Badge variant="secondary">{clients?.length ?? 0}</Badge>
</div>
<div className="flex gap-2">
<Button variant="outline" onClick={handleExport}>
<Download className="mr-2 h-4 w-4" />
Export to Excel
</Button>
<Button onClick={handleOpenCreate}>
<Plus className="mr-2 h-4 w-4" />
New Client
</Button>
</div>
{isAdminUser && (
<div className="flex gap-2">
<Button variant="outline" onClick={handleExport}>
<Download className="mr-2 h-4 w-4" />
Export to Excel
</Button>
<Button onClick={handleOpenCreate}>
<Plus className="mr-2 h-4 w-4" />
New Client
</Button>
</div>
)}
</div>
{/* Filters */}
<Card>
<CardHeader>
<CardTitle>Filters</CardTitle>
</CardHeader>
<CardContent className="grid gap-5 sm:grid-cols-2 lg:grid-cols-3">
<div className="grid gap-2">
<Label htmlFor="filter-search">Search</Label>
<div className="flex gap-2">
<Input
id="filter-search"
className="min-w-0"
placeholder="Name, email, or IP..."
value={searchInput}
onChange={(e) => setSearchDirty(e.target.value)}
onKeyDown={(e) => {
if (e.key === "Enter") { setFilter("search", searchInput); setSearchDirty(null); }
}}
/>
<Button
variant="outline"
size="icon"
onClick={() => { setFilter("search", searchInput); setSearchDirty(null); }}
aria-label="Search"
>
<Search className="h-4 w-4" />
</Button>
{/* Filters (admin only) */}
{isAdminUser && (
<Card>
<CardHeader>
<CardTitle>Filters</CardTitle>
</CardHeader>
<CardContent className="grid gap-5 sm:grid-cols-2 lg:grid-cols-3">
<div className="grid gap-2">
<Label htmlFor="filter-search">Search</Label>
<div className="flex gap-2">
<Input
id="filter-search"
className="min-w-0"
placeholder="Name, email, or IP..."
value={searchInput}
onChange={(e) => setSearchDirty(e.target.value)}
onKeyDown={(e) => {
if (e.key === "Enter") { setFilter("search", searchInput); setSearchDirty(null); }
}}
/>
<Button
variant="outline"
size="icon"
onClick={() => { setFilter("search", searchInput); setSearchDirty(null); }}
aria-label="Search"
>
<Search className="h-4 w-4" />
</Button>
</div>
</div>
</div>
<div className="grid gap-2">
<Label>Status</Label>
<Select
value={filterStatus || undefined}
onValueChange={(v: string | null) => setFilter("status", !v || v === "_all" ? "" : v)}
>
<SelectTrigger>
<SelectValue placeholder="All" />
</SelectTrigger>
<SelectContent>
<SelectItem value="_all">All</SelectItem>
<SelectItem value="enabled">Enabled</SelectItem>
<SelectItem value="disabled">Disabled</SelectItem>
<SelectItem value="connected">Connected</SelectItem>
<SelectItem value="disconnected">Disconnected</SelectItem>
</SelectContent>
</Select>
</div>
</CardContent>
</Card>
<div className="grid gap-2">
<Label>Status</Label>
<Select
value={filterStatus || undefined}
onValueChange={(v: string | null) => setFilter("status", !v || v === "_all" ? "" : v)}
>
<SelectTrigger>
<SelectValue placeholder="All" />
</SelectTrigger>
<SelectContent>
<SelectItem value="_all">All</SelectItem>
<SelectItem value="enabled">Enabled</SelectItem>
<SelectItem value="disabled">Disabled</SelectItem>
<SelectItem value="connected">Connected</SelectItem>
<SelectItem value="disconnected">Disconnected</SelectItem>
</SelectContent>
</Select>
</div>
</CardContent>
</Card>
)}
<div className="grid gap-4">
{clients?.map((cd) => {
const client = cd.Client;
return (
<Card key={client.id}>
<CardHeader className="flex flex-row items-center justify-between space-y-0 pb-3">
<CardTitle className="text-base font-medium">
{client.name}
{client.email && (
<span className="ml-2 text-sm font-normal text-muted-foreground">
{client.email}
</span>
)}
</CardTitle>
<div className="flex items-center gap-3">
<label
className="flex items-center gap-2 text-sm"
htmlFor={`toggle-${client.id}`}
>
<Switch
id={`toggle-${client.id}`}
checked={client.enabled}
onCheckedChange={(checked) =>
handleToggle(client.id, checked)
}
aria-label={`${client.enabled ? "Disable" : "Enable"} ${client.name}`}
/>
<Badge
variant={client.enabled ? "default" : "secondary"}
>
{client.enabled ? "Enabled" : "Disabled"}
</Badge>
</label>
<CardContent className="px-5 py-3">
<div className="flex items-start justify-between gap-4">
<div className="min-w-0">
<div className="flex items-center gap-2">
<span className="text-base font-semibold">{client.name}</span>
<Badge variant={client.enabled ? "default" : "secondary"}>
{client.enabled ? "Enabled" : "Disabled"}
</Badge>
</div>
<code className="text-xs text-muted-foreground">{client.public_key}</code>
<div className="text-muted-foreground">{client.email}</div>
</div>
<div className="flex items-center gap-1 shrink-0">
{isAdminUser && (
<Switch
id={`toggle-${client.id}`}
checked={client.enabled}
onCheckedChange={(checked) => handleToggle(client.id, checked)}
aria-label={`${client.enabled ? "Disable" : "Enable"} ${client.name}`}
/>
)}
</div>
</div>
</CardHeader>
<CardContent>
<div className="flex flex-col gap-3 sm:flex-row sm:items-start sm:justify-between">
<div className="space-y-1 text-sm text-muted-foreground">
<div className="mt-4 grid gap-x-8 gap-y-2 text-sm sm:grid-cols-2 lg:grid-cols-3">
<div>
<div className="text-xs font-medium text-muted-foreground">Allocated IPs</div>
<div>{client.allocated_ips?.join(", ") || "—"}</div>
</div>
<div>
<div className="text-xs font-medium text-muted-foreground">Allowed IPs</div>
<div>{client.allowed_ips?.join(", ") || "—"}</div>
</div>
{client.extra_allowed_ips && client.extra_allowed_ips.length > 0 && client.extra_allowed_ips.some(ip => ip) && (
<div>
Allocated IPs: {client.allocated_ips?.join(", ") || "None"}
<div className="text-xs font-medium text-muted-foreground">Extra Allowed IPs</div>
<div>{client.extra_allowed_ips.join(", ")}</div>
</div>
<div>
Allowed IPs: {client.allowed_ips?.join(", ") || "None"}
</div>
{client.extra_allowed_ips && client.extra_allowed_ips.length > 0 && client.extra_allowed_ips.some(ip => ip) && (
<div>
Extra Allowed IPs: {client.extra_allowed_ips.join(", ")}
</div>
)}
{client.additional_notes && (
<div>Notes: {client.additional_notes}</div>
)}
<div className="flex gap-4 text-xs text-muted-foreground/70">
<span>Created: {formatDate(client.created_at)}</span>
<span>Updated: {formatDate(client.updated_at)}</span>
)}
{client.additional_notes && (
<div className="sm:col-span-2 lg:col-span-3">
<div className="text-xs font-medium text-muted-foreground">Notes</div>
<div>{client.additional_notes}</div>
</div>
)}
</div>
<div className="mt-4 flex flex-col gap-3 border-t pt-3 sm:flex-row sm:items-center sm:justify-between">
<div className="flex gap-4 text-xs text-muted-foreground">
<span>Created {formatDate(client.created_at)}</span>
<span>Updated {formatDate(client.updated_at)}</span>
</div>
<div className="flex flex-wrap gap-1">
<Button
variant="ghost"
size="icon"
onClick={() => handleOpenEdit(client)}
aria-label={`Edit ${client.name}`}
>
<Pencil className="h-4 w-4" />
</Button>
<Button
variant="ghost"
size="icon"
onClick={() => handleOpenEmail(client)}
aria-label={`Email config to ${client.name}`}
>
<Mail className="h-4 w-4" />
</Button>
{cd.QRCode && (
<Button
variant="ghost"
size="icon"
onClick={() => setQrDialog(cd)}
aria-label={`Show QR code for ${client.name}`}
>
<QrCode className="h-4 w-4" />
{isAdminUser && (
<Button variant="ghost" size="icon" onClick={() => handleOpenEdit(client)} aria-label={`Edit ${client.name}`}>
<Pencil className="h-4 w-4" />
</Button>
)}
<Button
variant="ghost"
size="icon"
onClick={() => handleDownload(client.id)}
aria-label={`Download config for ${client.name}`}
>
{isAdminUser && (
<Button variant="ghost" size="icon" onClick={() => handleOpenEmail(client)} aria-label={`Email config to ${client.name}`}>
<Mail className="h-4 w-4" />
</Button>
)}
<Button variant="ghost" size="icon" onClick={() => setQrDialog(cd)} aria-label={`Show QR code for ${client.name}`}>
<QrCode className="h-4 w-4" />
</Button>
<Button variant="ghost" size="icon" onClick={() => handleDownload(client.id)} aria-label={`Download config for ${client.name}`}>
<Download className="h-4 w-4" />
</Button>
<Button
variant="ghost"
size="icon"
onClick={() => handleDelete(client.id, client.name)}
aria-label={`Delete ${client.name}`}
>
<Trash2 className="h-4 w-4 text-destructive" />
</Button>
{isAdminUser && (
<Button variant="ghost" size="icon" onClick={() => handleDelete(client.id, client.name)} aria-label={`Delete ${client.name}`}>
<Trash2 className="h-4 w-4 text-destructive" />
</Button>
)}
</div>
</div>
</CardContent>
@ -505,29 +527,16 @@ export function ClientsPage() {
{(!clients || clients.length === 0) && (
<Card>
<CardContent className="py-12 text-center text-muted-foreground">
No clients configured yet. Click "New Client" to add one.
{isAdminUser
? 'No clients configured yet. Click "New Client" to add one.'
: "No client configurations found for your account."}
</CardContent>
</Card>
)}
</div>
{/* QR Code Dialog */}
<Dialog open={!!qrDialog} onOpenChange={() => setQrDialog(null)}>
<DialogContent>
<DialogHeader>
<DialogTitle>{qrDialog?.Client.name} - QR Code</DialogTitle>
</DialogHeader>
{qrDialog?.QRCode && (
<div className="flex justify-center p-4">
<img
src={qrDialog.QRCode}
alt={`QR code for ${qrDialog.Client.name}`}
className="max-w-[256px]"
/>
</div>
)}
</DialogContent>
</Dialog>
<QrCodeDialog client={qrDialog} onClose={() => setQrDialog(null)} />
{/* Delete Confirmation Dialog */}
<Dialog open={!!deleteDialog} onOpenChange={() => setDeleteDialog(null)}>
@ -615,11 +624,11 @@ export function ClientsPage() {
<Input
id="new-ips"
placeholder="e.g. 10.0.0.2/32, 10.0.0.3/32"
value={newClient.allocated_ips.join(", ")}
value={newClient.allocated_ips}
onChange={(e) =>
setNewClient((p) => ({
...p,
allocated_ips: splitList(e.target.value),
allocated_ips: e.target.value,
}))
}
/>
@ -632,11 +641,11 @@ export function ClientsPage() {
<Input
id="new-allowed"
placeholder="e.g. 10.0.0.2/32, 10.0.0.3/32"
value={newClient.allowed_ips.join(", ")}
value={newClient.allowed_ips}
onChange={(e) =>
setNewClient((p) => ({
...p,
allowed_ips: splitList(e.target.value),
allowed_ips: e.target.value,
}))
}
/>
@ -649,11 +658,11 @@ export function ClientsPage() {
<Input
id="new-extra-allowed"
placeholder="e.g. 10.0.0.2/32, 10.0.0.3/32"
value={newClient.extra_allowed_ips.join(", ")}
value={newClient.extra_allowed_ips}
onChange={(e) =>
setNewClient((p) => ({
...p,
extra_allowed_ips: splitList(e.target.value),
extra_allowed_ips: e.target.value,
}))
}
/>
@ -707,16 +716,6 @@ export function ClientsPage() {
/>
<Label htmlFor="new-dns">Use server DNS</Label>
</div>
<div className="flex items-center gap-3">
<Switch
id="new-enabled"
checked={newClient.enabled}
onCheckedChange={(v) =>
setNewClient((p) => ({ ...p, enabled: v }))
}
/>
<Label htmlFor="new-enabled">Enable after creation</Label>
</div>
</div>
<div className="flex justify-end gap-3">
<Button variant="outline" onClick={() => setShowCreate(false)}>
@ -743,7 +742,7 @@ export function ClientsPage() {
<Label htmlFor="edit-name">Name</Label>
<Input
id="edit-name"
value={editForm.name ?? ""}
value={editForm.name}
onChange={(e) =>
setEditForm((p) => ({ ...p, name: e.target.value }))
}
@ -767,11 +766,11 @@ export function ClientsPage() {
<Input
id="edit-ips"
placeholder="e.g. 10.0.0.2/32, 10.0.0.3/32"
value={editForm.allocated_ips?.join(", ") ?? ""}
value={editForm.allocated_ips}
onChange={(e) =>
setEditForm((p) => ({
...p,
allocated_ips: splitList(e.target.value),
allocated_ips: e.target.value,
}))
}
/>
@ -784,11 +783,11 @@ export function ClientsPage() {
<Input
id="edit-allowed"
placeholder="e.g. 10.0.0.2/32, 10.0.0.3/32"
value={editForm.allowed_ips?.join(", ") ?? ""}
value={editForm.allowed_ips}
onChange={(e) =>
setEditForm((p) => ({
...p,
allowed_ips: splitList(e.target.value),
allowed_ips: e.target.value,
}))
}
/>
@ -801,11 +800,11 @@ export function ClientsPage() {
<Input
id="edit-extra-allowed"
placeholder="e.g. 10.0.0.2/32, 10.0.0.3/32"
value={editForm.extra_allowed_ips?.join(", ") ?? ""}
value={editForm.extra_allowed_ips}
onChange={(e) =>
setEditForm((p) => ({
...p,
extra_allowed_ips: splitList(e.target.value),
extra_allowed_ips: e.target.value,
}))
}
/>
@ -819,7 +818,7 @@ export function ClientsPage() {
<Label htmlFor="edit-endpoint">Endpoint</Label>
<Input
id="edit-endpoint"
value={editForm.endpoint ?? ""}
value={editForm.endpoint}
onChange={(e) =>
setEditForm((p) => ({ ...p, endpoint: e.target.value }))
}
@ -832,7 +831,7 @@ export function ClientsPage() {
<Label htmlFor="edit-notes">Notes</Label>
<Textarea
id="edit-notes"
value={editForm.additional_notes ?? ""}
value={editForm.additional_notes}
onChange={(e) =>
setEditForm((p) => ({
...p,
@ -841,11 +840,23 @@ export function ClientsPage() {
}
/>
</div>
<div className="grid gap-2 sm:col-span-2">
<div className="grid gap-2">
<Label htmlFor="edit-pubkey">Public Key</Label>
<Input
id="edit-pubkey"
className="font-mono text-xs"
value={editForm.public_key}
onChange={(e) =>
setEditForm((p) => ({ ...p, public_key: e.target.value }))
}
/>
</div>
<div className="grid gap-2">
<Label htmlFor="edit-psk">Preshared Key</Label>
<Input
id="edit-psk"
value={editForm.preshared_key ?? ""}
className="font-mono text-xs"
value={editForm.preshared_key}
onChange={(e) =>
setEditForm((p) => ({ ...p, preshared_key: e.target.value }))
}
@ -854,7 +865,7 @@ export function ClientsPage() {
<div className="flex items-center gap-3 sm:col-span-2">
<Switch
id="edit-dns"
checked={editForm.use_server_dns ?? false}
checked={editForm.use_server_dns}
onCheckedChange={(v) =>
setEditForm((p) => ({ ...p, use_server_dns: v }))
}

View File

@ -30,19 +30,19 @@ describe("ServerPage interactions", () => {
expect(window.confirm).toHaveBeenCalled();
});
it("clicks apply config", async () => {
it("clicks save interface", async () => {
const user = userEvent.setup();
cleanup = mockFetch({
"/server": serverData,
"/server/apply-config": { message: "ok" },
"/server/interface": serverData.Interface,
});
renderWithProviders(<ServerPage />);
await waitFor(() => {
expect(screen.getByText("Apply Config")).toBeInTheDocument();
expect(screen.getByText("Save")).toBeInTheDocument();
});
await user.click(screen.getByText("Apply Config"));
await user.click(screen.getByText("Save"));
});
it("displays server addresses", async () => {
@ -54,4 +54,167 @@ describe("ServerPage interactions", () => {
expect(screen.getByDisplayValue("51820")).toBeInTheDocument();
});
});
it("cancels regenerate keypair when user declines confirm", async () => {
const user = userEvent.setup();
vi.spyOn(window, "confirm").mockReturnValue(false);
cleanup = mockFetch({ "/server": serverData });
renderWithProviders(<ServerPage />);
await waitFor(() => {
expect(screen.getByText("Regenerate")).toBeInTheDocument();
});
await user.click(screen.getByText("Regenerate"));
expect(window.confirm).toHaveBeenCalled();
});
it("shows validation error for empty addresses", async () => {
const user = userEvent.setup();
cleanup = mockFetch({ "/server": serverData });
renderWithProviders(<ServerPage />);
await waitFor(() => {
expect(screen.getByDisplayValue("10.0.0.1/24")).toBeInTheDocument();
});
const addrInput = screen.getByLabelText("Server addresses");
await user.clear(addrInput);
await waitFor(() => {
expect(screen.getByText("At least one address is required")).toBeInTheDocument();
});
});
it("shows validation error for invalid CIDR addresses", async () => {
const user = userEvent.setup();
cleanup = mockFetch({ "/server": serverData });
renderWithProviders(<ServerPage />);
await waitFor(() => {
expect(screen.getByDisplayValue("10.0.0.1/24")).toBeInTheDocument();
});
const addrInput = screen.getByLabelText("Server addresses");
await user.clear(addrInput);
await user.type(addrInput, "not-a-cidr");
await waitFor(() => {
expect(screen.getByText("Each address must be valid CIDR (e.g. 10.252.1.0/24)")).toBeInTheDocument();
});
});
it("shows validation error for empty listen port", async () => {
const user = userEvent.setup();
cleanup = mockFetch({ "/server": serverData });
renderWithProviders(<ServerPage />);
await waitFor(() => {
expect(screen.getByDisplayValue("51820")).toBeInTheDocument();
});
const portInput = screen.getByLabelText("Listen port");
await user.clear(portInput);
await waitFor(() => {
expect(screen.getByText("Listen port is required")).toBeInTheDocument();
});
});
it("shows validation error for out-of-range port", async () => {
const user = userEvent.setup();
cleanup = mockFetch({ "/server": serverData });
renderWithProviders(<ServerPage />);
await waitFor(() => {
expect(screen.getByDisplayValue("51820")).toBeInTheDocument();
});
const portInput = screen.getByLabelText("Listen port");
await user.clear(portInput);
await user.type(portInput, "99999");
await waitFor(() => {
expect(screen.getByText("Port must be between 1 and 65535")).toBeInTheDocument();
});
});
it("disables Save button when form is invalid", async () => {
const user = userEvent.setup();
cleanup = mockFetch({ "/server": serverData });
renderWithProviders(<ServerPage />);
await waitFor(() => {
expect(screen.getByDisplayValue("10.0.0.1/24")).toBeInTheDocument();
});
// Clear addresses to make form invalid
const addrInput = screen.getByLabelText("Server addresses");
await user.clear(addrInput);
await waitFor(() => {
const saveButton = screen.getByText("Save").closest("button");
expect(saveButton).toBeDisabled();
});
});
it("shows Keypair card with public key", async () => {
cleanup = mockFetch({ "/server": serverData });
renderWithProviders(<ServerPage />);
await waitFor(() => {
expect(screen.getByText("Keypair")).toBeInTheDocument();
expect(screen.getByLabelText("Server public key")).toBeInTheDocument();
expect(screen.getByDisplayValue("serverpub123")).toBeInTheDocument();
});
});
it("public key input is read-only", async () => {
cleanup = mockFetch({ "/server": serverData });
renderWithProviders(<ServerPage />);
await waitFor(() => {
const pubKeyInput = screen.getByLabelText("Server public key");
expect(pubKeyInput).toHaveAttribute("readonly");
});
});
it("shows Post-Up, Pre-Down, Post-Down script fields", async () => {
const serverWithScripts = {
Interface: {
addresses: ["10.0.0.1/24"],
listen_port: 51820,
post_up: "iptables -A FORWARD",
pre_down: "echo predown",
post_down: "iptables -D FORWARD",
},
KeyPair: { public_key: "pub123", private_key: "priv" },
};
cleanup = mockFetch({ "/server": serverWithScripts });
renderWithProviders(<ServerPage />);
await waitFor(() => {
expect(screen.getByText("Post-Up Script")).toBeInTheDocument();
expect(screen.getByText("Pre-Down Script")).toBeInTheDocument();
expect(screen.getByText("Post-Down Script")).toBeInTheDocument();
expect(screen.getByDisplayValue("iptables -A FORWARD")).toBeInTheDocument();
expect(screen.getByDisplayValue("echo predown")).toBeInTheDocument();
expect(screen.getByDisplayValue("iptables -D FORWARD")).toBeInTheDocument();
});
});
it("edits Post-Up Script field", async () => {
const user = userEvent.setup();
cleanup = mockFetch({ "/server": serverData });
renderWithProviders(<ServerPage />);
await waitFor(() => {
expect(screen.getByText("Post-Up Script")).toBeInTheDocument();
});
const postUpInput = screen.getByPlaceholderText("iptables -A FORWARD ...");
await user.type(postUpInput, "echo hello");
expect(postUpInput).toHaveValue("echo hello");
});
});

View File

@ -10,7 +10,7 @@ describe("ServerPage", () => {
cleanup?.();
});
it("shows heading and apply button", async () => {
it("shows heading and save button", async () => {
cleanup = mockFetch({
"/server": {
Interface: { addresses: ["10.0.0.1/24"], listen_port: 51820 },
@ -20,7 +20,7 @@ describe("ServerPage", () => {
renderWithProviders(<ServerPage />);
await waitFor(() => {
expect(screen.getByText("Server Configuration")).toBeInTheDocument();
expect(screen.getByText("Apply Config")).toBeInTheDocument();
expect(screen.getByText("Save")).toBeInTheDocument();
});
});

View File

@ -50,17 +50,15 @@ export function ServerPage() {
}, [addrValue, portValue]);
const serverValid = Object.keys(serverErrors).length === 0;
const saveAndApply = useMutation({
mutationFn: async () => {
await apiPut("/server/interface", {
const saveInterface = useMutation({
mutationFn: () =>
apiPut("/server/interface", {
addresses: splitList(addrValue),
listen_port: Number(portValue) || 0,
post_up: postUpValue,
pre_down: preDownValue,
post_down: postDownValue,
});
await apiPost("/server/apply-config");
},
}),
onSuccess: () => {
qc.invalidateQueries({ queryKey: ["server"] });
setAddresses(null);
@ -68,7 +66,7 @@ export function ServerPage() {
setPostUp(null);
setPreDown(null);
setPostDown(null);
toast.success("Interface saved and config applied");
toast.success("Interface saved");
},
onError: (err: Error) => toast.error(err.message),
});
@ -91,11 +89,11 @@ export function ServerPage() {
Server Configuration
</h2>
<Button
onClick={() => saveAndApply.mutate()}
disabled={!serverValid || saveAndApply.isPending}
onClick={() => saveInterface.mutate()}
disabled={!serverValid || saveInterface.isPending}
>
<Save className="mr-2 h-4 w-4" />
{saveAndApply.isPending ? "Applying..." : "Apply Config"}
{saveInterface.isPending ? "Saving..." : "Save"}
</Button>
</div>

View File

@ -1,8 +1,35 @@
import { describe, it, expect, afterEach } from "vitest";
import { screen, waitFor } from "@testing-library/react";
import { screen, waitFor, within } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { renderWithProviders, mockFetch } from "@/test/test-utils";
import { StatusPage } from "./StatusPage";
const connectedPeer = {
name: "Client1",
email: "c1@test.com",
public_key: "pk1abcdef1234567890",
received_bytes: 1024,
transmit_bytes: 2048,
last_handshake_time: new Date().toISOString(),
last_handshake_rel: 60_000_000_000, // 60 seconds in nanos
connected: true,
allocated_ip: "10.0.0.2/32",
endpoint: "1.2.3.4:51820",
};
const disconnectedPeer = {
name: "Client2",
email: "c2@test.com",
public_key: "pk2abcdef1234567890",
received_bytes: 0,
transmit_bytes: 0,
last_handshake_time: "",
last_handshake_rel: 0,
connected: false,
allocated_ip: "10.0.0.3/32",
endpoint: "",
};
describe("StatusPage", () => {
let cleanup: () => void;
@ -31,20 +58,7 @@ describe("StatusPage", () => {
"/status": [
{
name: "wg0",
peers: [
{
name: "Client1",
email: "c1@test.com",
public_key: "pk1",
received_bytes: 1024,
transmit_bytes: 2048,
last_handshake_time: new Date().toISOString(),
last_handshake_rel: 60000000000,
connected: true,
allocated_ip: "10.0.0.2/32",
endpoint: "1.2.3.4:51820",
},
],
peers: [connectedPeer],
},
],
});
@ -53,7 +67,316 @@ describe("StatusPage", () => {
await waitFor(() => {
expect(screen.getByText("wg0")).toBeInTheDocument();
expect(screen.getByText("Client1")).toBeInTheDocument();
expect(screen.getByText("Connected")).toBeInTheDocument();
expect(screen.getByText("1.2.3.4:51820")).toBeInTheDocument();
});
});
it("shows sortable column headers", async () => {
cleanup = mockFetch({ "/status": [{ name: "wg0", peers: [] }] });
renderWithProviders(<StatusPage />);
await waitFor(() => {
expect(screen.getByText("Name")).toBeInTheDocument();
expect(screen.getByText("Handshake")).toBeInTheDocument();
expect(screen.getByText("Endpoint")).toBeInTheDocument();
});
});
it("shows 'No peers connected' when device has no peers", async () => {
cleanup = mockFetch({ "/status": [{ name: "wg0", peers: [] }] });
renderWithProviders(<StatusPage />);
await waitFor(() => {
expect(screen.getByText("No peers connected")).toBeInTheDocument();
});
});
it("shows 'No peers connected' when peers is null", async () => {
cleanup = mockFetch({ "/status": [{ name: "wg0", peers: null }] });
renderWithProviders(<StatusPage />);
await waitFor(() => {
expect(screen.getByText("No peers connected")).toBeInTheDocument();
});
});
it("displays disconnected peer with dash for endpoint", async () => {
cleanup = mockFetch({
"/status": [{ name: "wg0", peers: [disconnectedPeer] }],
});
renderWithProviders(<StatusPage />);
await waitFor(() => {
expect(screen.getByText("Client2")).toBeInTheDocument();
expect(screen.getByText("-")).toBeInTheDocument();
});
});
it("shows 'Unknown' when peer has no name", async () => {
const namelessPeer = { ...connectedPeer, name: "" };
cleanup = mockFetch({
"/status": [{ name: "wg0", peers: [namelessPeer] }],
});
renderWithProviders(<StatusPage />);
await waitFor(() => {
expect(screen.getByText("Unknown")).toBeInTheDocument();
});
});
it("displays truncated public key", async () => {
cleanup = mockFetch({
"/status": [{ name: "wg0", peers: [connectedPeer] }],
});
renderWithProviders(<StatusPage />);
await waitFor(() => {
expect(screen.getByText("pk1abcdef1234567...")).toBeInTheDocument();
});
});
it("formats bytes correctly for various sizes", async () => {
const largePeer = {
...connectedPeer,
received_bytes: 1_500_000_000, // ~1.4 GB
transmit_bytes: 2_500_000, // ~2.4 MB
};
cleanup = mockFetch({
"/status": [{ name: "wg0", peers: [largePeer] }],
});
renderWithProviders(<StatusPage />);
await waitFor(() => {
expect(screen.getByText("1.4 GB")).toBeInTheDocument();
expect(screen.getByText("2.4 MB")).toBeInTheDocument();
});
});
it("formats zero bytes as '0 B'", async () => {
cleanup = mockFetch({
"/status": [{ name: "wg0", peers: [disconnectedPeer] }],
});
renderWithProviders(<StatusPage />);
await waitFor(() => {
const zeroBytes = screen.getAllByText("0 B");
expect(zeroBytes.length).toBe(2); // rx and tx both 0
});
});
it("formats handshake as 'Never' for zero/negative nanos", async () => {
cleanup = mockFetch({
"/status": [{ name: "wg0", peers: [disconnectedPeer] }],
});
renderWithProviders(<StatusPage />);
await waitFor(() => {
expect(screen.getByText("Never")).toBeInTheDocument();
});
});
it("formats handshake in seconds", async () => {
const peerSec = { ...connectedPeer, last_handshake_rel: 30_000_000_000 }; // 30s
cleanup = mockFetch({
"/status": [{ name: "wg0", peers: [peerSec] }],
});
renderWithProviders(<StatusPage />);
await waitFor(() => {
expect(screen.getByText("30s ago")).toBeInTheDocument();
});
});
it("formats handshake in minutes", async () => {
const peerMin = { ...connectedPeer, last_handshake_rel: 300_000_000_000 }; // 5m
cleanup = mockFetch({
"/status": [{ name: "wg0", peers: [peerMin] }],
});
renderWithProviders(<StatusPage />);
await waitFor(() => {
expect(screen.getByText("5m ago")).toBeInTheDocument();
});
});
it("formats handshake in hours", async () => {
const peerHr = { ...connectedPeer, last_handshake_rel: 7_200_000_000_000 }; // 2h
cleanup = mockFetch({
"/status": [{ name: "wg0", peers: [peerHr] }],
});
renderWithProviders(<StatusPage />);
await waitFor(() => {
expect(screen.getByText("2h ago")).toBeInTheDocument();
});
});
it("formats handshake in days", async () => {
const peerDay = { ...connectedPeer, last_handshake_rel: 172_800_000_000_000 }; // 2d
cleanup = mockFetch({
"/status": [{ name: "wg0", peers: [peerDay] }],
});
renderWithProviders(<StatusPage />);
await waitFor(() => {
expect(screen.getByText("2d ago")).toBeInTheDocument();
});
});
it("sorts peers by name ascending by default", async () => {
const peerA = { ...connectedPeer, name: "Alpha", public_key: "pka1234567890123" };
const peerB = { ...disconnectedPeer, name: "Bravo", public_key: "pkb1234567890123" };
cleanup = mockFetch({
"/status": [{ name: "wg0", peers: [peerB, peerA] }],
});
renderWithProviders(<StatusPage />);
await waitFor(() => {
expect(screen.getByText("Alpha")).toBeInTheDocument();
expect(screen.getByText("Bravo")).toBeInTheDocument();
});
// Alpha should come before Bravo in the DOM
const rows = screen.getAllByRole("row");
const alphaRow = rows.find((r) => within(r).queryByText("Alpha"));
const bravoRow = rows.find((r) => within(r).queryByText("Bravo"));
expect(rows.indexOf(alphaRow!)).toBeLessThan(rows.indexOf(bravoRow!));
});
it("toggles sort direction when clicking same column header", async () => {
const user = userEvent.setup();
const peerA = { ...connectedPeer, name: "Alpha", public_key: "pka1234567890123" };
const peerB = { ...disconnectedPeer, name: "Bravo", public_key: "pkb1234567890123" };
cleanup = mockFetch({
"/status": [{ name: "wg0", peers: [peerA, peerB] }],
});
renderWithProviders(<StatusPage />);
await waitFor(() => {
expect(screen.getByText("Alpha")).toBeInTheDocument();
});
// Click Name header to toggle to desc
await user.click(screen.getByText("Name"));
// Now Bravo should come before Alpha
const rows = screen.getAllByRole("row");
const alphaRow = rows.find((r) => within(r).queryByText("Alpha"));
const bravoRow = rows.find((r) => within(r).queryByText("Bravo"));
expect(rows.indexOf(bravoRow!)).toBeLessThan(rows.indexOf(alphaRow!));
});
it("sorts by a different column when clicking a new header", async () => {
const user = userEvent.setup();
const peerLowRx = { ...connectedPeer, name: "LowRx", public_key: "pkl1234567890123", received_bytes: 100 };
const peerHighRx = { ...connectedPeer, name: "HighRx", public_key: "pkh1234567890123", received_bytes: 999999 };
cleanup = mockFetch({
"/status": [{ name: "wg0", peers: [peerHighRx, peerLowRx] }],
});
renderWithProviders(<StatusPage />);
await waitFor(() => {
expect(screen.getByText("LowRx")).toBeInTheDocument();
});
// Click Rx header to sort by received_bytes asc
await user.click(screen.getByText("Rx"));
const rows = screen.getAllByRole("row");
const lowRow = rows.find((r) => within(r).queryByText("LowRx"));
const highRow = rows.find((r) => within(r).queryByText("HighRx"));
expect(rows.indexOf(lowRow!)).toBeLessThan(rows.indexOf(highRow!));
});
it("sorts by Tx column", async () => {
const user = userEvent.setup();
const peerLowTx = { ...connectedPeer, name: "LowTx", public_key: "pkl1234567890123", transmit_bytes: 50 };
const peerHighTx = { ...connectedPeer, name: "HighTx", public_key: "pkh1234567890123", transmit_bytes: 999999 };
cleanup = mockFetch({
"/status": [{ name: "wg0", peers: [peerHighTx, peerLowTx] }],
});
renderWithProviders(<StatusPage />);
await waitFor(() => {
expect(screen.getByText("LowTx")).toBeInTheDocument();
});
await user.click(screen.getByText("Tx"));
const rows = screen.getAllByRole("row");
const lowRow = rows.find((r) => within(r).queryByText("LowTx"));
const highRow = rows.find((r) => within(r).queryByText("HighTx"));
expect(rows.indexOf(lowRow!)).toBeLessThan(rows.indexOf(highRow!));
});
it("sorts by Endpoint column", async () => {
const user = userEvent.setup();
const peerA = { ...connectedPeer, name: "PeerA", public_key: "pka1234567890123", endpoint: "a.example.com:51820" };
const peerB = { ...connectedPeer, name: "PeerB", public_key: "pkb1234567890123", endpoint: "z.example.com:51820" };
cleanup = mockFetch({
"/status": [{ name: "wg0", peers: [peerB, peerA] }],
});
renderWithProviders(<StatusPage />);
await waitFor(() => {
expect(screen.getByText("PeerA")).toBeInTheDocument();
});
await user.click(screen.getByText("Endpoint"));
const rows = screen.getAllByRole("row");
const aRow = rows.find((r) => within(r).queryByText("PeerA"));
const bRow = rows.find((r) => within(r).queryByText("PeerB"));
expect(rows.indexOf(aRow!)).toBeLessThan(rows.indexOf(bRow!));
});
it("sorts by Handshake column", async () => {
const user = userEvent.setup();
const peerRecent = { ...connectedPeer, name: "Recent", public_key: "pkr1234567890123", last_handshake_rel: 5_000_000_000 };
const peerOld = { ...connectedPeer, name: "Old", public_key: "pko1234567890123", last_handshake_rel: 999_000_000_000 };
cleanup = mockFetch({
"/status": [{ name: "wg0", peers: [peerOld, peerRecent] }],
});
renderWithProviders(<StatusPage />);
await waitFor(() => {
expect(screen.getByText("Recent")).toBeInTheDocument();
});
await user.click(screen.getByText("Handshake"));
const rows = screen.getAllByRole("row");
const recentRow = rows.find((r) => within(r).queryByText("Recent"));
const oldRow = rows.find((r) => within(r).queryByText("Old"));
expect(rows.indexOf(recentRow!)).toBeLessThan(rows.indexOf(oldRow!));
});
it("sorts by connected status column", async () => {
const user = userEvent.setup();
cleanup = mockFetch({
"/status": [{ name: "wg0", peers: [disconnectedPeer, connectedPeer] }],
});
renderWithProviders(<StatusPage />);
await waitFor(() => {
expect(screen.getByText("Client1")).toBeInTheDocument();
expect(screen.getByText("Client2")).toBeInTheDocument();
});
// The connected column header has no text label, it just has SortIcon
// We need to find and click the second TableHead (the one for connected)
const headerRow = screen.getAllByRole("row")[0];
const headerCells = within(headerRow).getAllByRole("columnheader");
// connected is the 2nd column header
await user.click(headerCells[1]);
});
it("renders multiple devices", async () => {
cleanup = mockFetch({
"/status": [
{ name: "wg0", peers: [connectedPeer] },
{ name: "wg1", peers: [disconnectedPeer] },
],
});
renderWithProviders(<StatusPage />);
await waitFor(() => {
expect(screen.getByText("wg0")).toBeInTheDocument();
expect(screen.getByText("wg1")).toBeInTheDocument();
});
});
it("formats KB correctly", async () => {
const peerKB = {
...connectedPeer,
received_bytes: 5120, // 5 KB
transmit_bytes: 512000, // 500 KB
};
cleanup = mockFetch({
"/status": [{ name: "wg0", peers: [peerKB] }],
});
renderWithProviders(<StatusPage />);
await waitFor(() => {
expect(screen.getByText("5.0 KB")).toBeInTheDocument();
expect(screen.getByText("500.0 KB")).toBeInTheDocument();
});
});
});

View File

@ -1,19 +1,77 @@
import { useState, useMemo } from "react";
import { useQuery } from "@tanstack/react-query";
import { apiGet } from "@/lib/api-client";
import { Badge } from "@/components/ui/badge";
import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card";
import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table";
import {
Table,
TableBody,
TableCell,
TableHead,
TableHeader,
TableRow,
} from "@/components/ui/table";
import {
Tooltip,
TooltipContent,
TooltipTrigger,
} from "@/components/ui/tooltip";
import { Skeleton } from "@/components/ui/skeleton";
import type { DeviceStatus } from "@/lib/types";
import { ArrowDown, ArrowUp, ArrowUpDown, CircleCheck, CircleX } from "lucide-react";
import type { DeviceStatus, PeerStatus } from "@/lib/types";
function formatBytes(bytes: number): string {
if (bytes === 0) return "0 B";
const k = 1024;
const sizes = ["B", "KB", "MB", "GB"];
const sizes = ["B", "KB", "MB", "GB", "TB"];
const i = Math.floor(Math.log(bytes) / Math.log(k));
return `${(bytes / Math.pow(k, i)).toFixed(1)} ${sizes[i]}`;
}
function formatHandshake(nanos: number): string {
if (!nanos || nanos <= 0) return "Never";
const seconds = Math.floor(nanos / 1_000_000_000);
if (seconds < 60) return `${seconds}s ago`;
const minutes = Math.floor(seconds / 60);
if (minutes < 60) return `${minutes}m ago`;
const hours = Math.floor(minutes / 60);
if (hours < 24) return `${hours}h ago`;
const days = Math.floor(hours / 24);
return `${days}d ago`;
}
type SortKey =
| "name"
| "connected"
| "last_handshake_rel"
| "received_bytes"
| "transmit_bytes"
| "endpoint";
type SortDir = "asc" | "desc";
function getSortValue(peer: PeerStatus, key: SortKey): string | number | boolean {
switch (key) {
case "name":
return (peer.name || "").toLowerCase();
case "connected":
return peer.connected ? 1 : 0;
case "last_handshake_rel":
return peer.last_handshake_rel || Number.MAX_SAFE_INTEGER;
case "received_bytes":
return peer.received_bytes;
case "transmit_bytes":
return peer.transmit_bytes;
case "endpoint":
return peer.endpoint || "";
}
}
function SortIcon({ column, sortKey, sortDir }: { column: SortKey; sortKey: SortKey; sortDir: SortDir }) {
if (column !== sortKey) return <ArrowUpDown className="ml-1 inline h-3 w-3 opacity-40" />;
return sortDir === "asc"
? <ArrowUp className="ml-1 inline h-3 w-3" />
: <ArrowDown className="ml-1 inline h-3 w-3" />;
}
export function StatusPage() {
const { data: devices, isLoading } = useQuery({
queryKey: ["status"],
@ -21,55 +79,111 @@ export function StatusPage() {
refetchInterval: 5000,
});
const [sortKey, setSortKey] = useState<SortKey>("name");
const [sortDir, setSortDir] = useState<SortDir>("asc");
const toggleSort = (key: SortKey) => {
if (sortKey === key) {
setSortDir((d) => (d === "asc" ? "desc" : "asc"));
} else {
setSortKey(key);
setSortDir("asc");
}
};
const sortedDevices = useMemo(() => {
if (!devices) return [];
return devices.map((device) => ({
...device,
peers: [...(device.peers || [])].sort((a, b) => {
const va = getSortValue(a, sortKey);
const vb = getSortValue(b, sortKey);
const cmp = va < vb ? -1 : va > vb ? 1 : 0;
return sortDir === "asc" ? cmp : -cmp;
}),
}));
}, [devices, sortKey, sortDir]);
if (isLoading) {
return <Skeleton className="h-64 w-full" />;
}
const headerClass = "cursor-pointer select-none hover:text-foreground";
return (
<div className="space-y-6">
<h2 className="text-2xl font-bold tracking-tight">Server Status</h2>
{devices?.map((device) => (
{sortedDevices.map((device) => (
<Card key={device.name}>
<CardHeader>
<CardTitle>{device.name}</CardTitle>
</CardHeader>
<CardContent>
<div className="overflow-x-auto">
<Table>
<TableHeader>
<TableRow>
<TableHead>Name</TableHead>
<TableHead>Status</TableHead>
<TableHead>IP</TableHead>
<TableHead>Received</TableHead>
<TableHead>Sent</TableHead>
<TableHead>Endpoint</TableHead>
</TableRow>
</TableHeader>
<TableBody>
{device.peers?.map((peer) => (
<TableRow key={peer.public_key}>
<TableCell className="font-medium">{peer.name || "Unknown"}</TableCell>
<TableCell>
<Badge variant={peer.connected ? "default" : "secondary"}>
{peer.connected ? "Connected" : "Disconnected"}
</Badge>
</TableCell>
<TableCell className="text-sm">{peer.allocated_ip}</TableCell>
<TableCell className="text-sm">{formatBytes(peer.received_bytes)}</TableCell>
<TableCell className="text-sm">{formatBytes(peer.transmit_bytes)}</TableCell>
<TableCell className="text-sm">{peer.endpoint || "-"}</TableCell>
</TableRow>
))}
{(!device.peers || device.peers.length === 0) && (
<Table>
<TableHeader>
<TableRow>
<TableCell colSpan={6} className="text-center text-muted-foreground">
No peers connected
</TableCell>
<TableHead className={headerClass} onClick={() => toggleSort("name")}>
Name <SortIcon column="name" sortKey={sortKey} sortDir={sortDir} />
</TableHead>
<TableHead className={`${headerClass} w-10`} onClick={() => toggleSort("connected")}>
<SortIcon column="connected" sortKey={sortKey} sortDir={sortDir} />
</TableHead>
<TableHead className={headerClass} onClick={() => toggleSort("endpoint")}>
Endpoint <SortIcon column="endpoint" sortKey={sortKey} sortDir={sortDir} />
</TableHead>
<TableHead className={headerClass} onClick={() => toggleSort("last_handshake_rel")}>
Handshake <SortIcon column="last_handshake_rel" sortKey={sortKey} sortDir={sortDir} />
</TableHead>
<TableHead className={headerClass} onClick={() => toggleSort("received_bytes")}>
Rx <SortIcon column="received_bytes" sortKey={sortKey} sortDir={sortDir} />
</TableHead>
<TableHead className={headerClass} onClick={() => toggleSort("transmit_bytes")}>
Tx <SortIcon column="transmit_bytes" sortKey={sortKey} sortDir={sortDir} />
</TableHead>
</TableRow>
)}
</TableBody>
</Table>
</TableHeader>
<TableBody>
{device.peers?.map((peer) => (
<TableRow key={peer.public_key}>
<TableCell>
<div className="font-medium">{peer.name || "Unknown"}</div>
<code className="text-xs text-muted-foreground">
{peer.public_key.substring(0, 16)}...
</code>
</TableCell>
<TableCell>
<Tooltip>
<TooltipTrigger>
{peer.connected ? (
<CircleCheck className="h-5 w-5 text-green-500" />
) : (
<CircleX className="h-5 w-5 text-muted-foreground" />
)}
</TooltipTrigger>
<TooltipContent>
{peer.connected ? "Connected" : "Disconnected"}
</TooltipContent>
</Tooltip>
</TableCell>
<TableCell className="font-mono">{peer.endpoint || "-"}</TableCell>
<TableCell>{formatHandshake(peer.last_handshake_rel)}</TableCell>
<TableCell>{formatBytes(peer.received_bytes)}</TableCell>
<TableCell>{formatBytes(peer.transmit_bytes)}</TableCell>
</TableRow>
))}
{(!device.peers || device.peers.length === 0) && (
<TableRow>
<TableCell
colSpan={6}
className="text-center text-muted-foreground"
>
No peers connected
</TableCell>
</TableRow>
)}
</TableBody>
</Table>
</div>
</CardContent>
</Card>

View File

@ -42,10 +42,8 @@ func MigrateFromJSON(sqliteDB *SqliteDB, jsonDBPath string) error {
return fmt.Errorf("migrate hashes: %w", err)
}
// migrate users
if err := migrateUsers(sqliteDB, jsonDBPath); err != nil {
return fmt.Errorf("migrate users: %w", err)
}
// NOTE: users are NOT migrated — legacy password users cannot log in with SSO-only auth.
// The first OIDC login will auto-provision as admin when len(users) == 0.
// migrate clients
if err := migrateClients(sqliteDB, jsonDBPath); err != nil {
@ -192,43 +190,6 @@ func migrateHashes(db *SqliteDB, jsonDBPath string) error {
return nil
}
func migrateUsers(db *SqliteDB, jsonDBPath string) error {
usersDir := filepath.Join(jsonDBPath, "users")
entries, err := os.ReadDir(usersDir)
if err != nil {
if os.IsNotExist(err) {
return nil
}
return err
}
count := 0
for _, entry := range entries {
if entry.IsDir() || filepath.Ext(entry.Name()) != ".json" {
continue
}
var u model.User
if err := readJSONFile(filepath.Join(usersDir, entry.Name()), &u); err != nil {
log.Warnf(" Skipping user file %s: %v", entry.Name(), err)
continue
}
now := time.Now().UTC()
_, err := db.db.Exec(
`INSERT OR REPLACE INTO users (username, email, display_name, admin, created_at, updated_at)
VALUES (?, ?, ?, ?, ?, ?)`,
u.Username, u.Email, u.DisplayName, u.Admin, now, now,
)
if err != nil {
return fmt.Errorf("migrate user %s: %w", u.Username, err)
}
count++
}
log.Infof(" Migrated %d users", count)
return nil
}
func migrateClients(db *SqliteDB, jsonDBPath string) error {
clientsDir := filepath.Join(jsonDBPath, "clients")
entries, err := os.ReadDir(clientsDir)

View File

@ -13,7 +13,7 @@ CREATE TABLE IF NOT EXISTS clients (
private_key TEXT NOT NULL DEFAULT '',
public_key TEXT NOT NULL DEFAULT '',
preshared_key TEXT NOT NULL DEFAULT '',
name TEXT NOT NULL DEFAULT '',
name TEXT NOT NULL DEFAULT '' UNIQUE,
email TEXT NOT NULL DEFAULT '',
subnet_ranges TEXT NOT NULL DEFAULT '[]',
allocated_ips TEXT NOT NULL DEFAULT '[]',

View File

@ -11,6 +11,7 @@ import (
"path/filepath"
"time"
"github.com/labstack/gommon/log"
_ "modernc.org/sqlite"
"github.com/skip2/go-qrcode"
@ -59,8 +60,42 @@ func New(dbPath string) (*SqliteDB, error) {
return &SqliteDB{db: db, dbPath: dbPath}, nil
}
// migrate applies incremental schema changes to existing databases
func (o *SqliteDB) migrate() {
if _, err := o.db.Exec(`CREATE UNIQUE INDEX IF NOT EXISTS idx_clients_name ON clients(name)`); err != nil {
log.Warnf("migrate: create idx_clients_name: %v", err)
}
if _, err := o.db.Exec(`CREATE UNIQUE INDEX IF NOT EXISTS idx_clients_public_key ON clients(public_key) WHERE public_key != ''`); err != nil {
log.Warnf("migrate: create idx_clients_public_key: %v", err)
}
o.db.Exec(`DELETE FROM users WHERE oidc_sub IS NULL OR oidc_sub = ''`)
// derive missing public keys from private keys
type keyPair struct{ id, privKey string }
var missing []keyPair
rows, err := o.db.Query(`SELECT id, private_key FROM clients WHERE public_key = '' AND private_key != ''`)
if err == nil {
for rows.Next() {
var kp keyPair
if rows.Scan(&kp.id, &kp.privKey) == nil {
missing = append(missing, kp)
}
}
rows.Close()
}
for _, kp := range missing {
if key, err := wgtypes.ParseKey(kp.privKey); err == nil {
o.db.Exec(`UPDATE clients SET public_key = ? WHERE id = ?`, key.PublicKey().String(), kp.id)
}
}
}
// Init initializes the database with default values if they don't exist
func (o *SqliteDB) Init() error {
// schema migrations for existing databases
o.migrate()
// server interface
var ifaceCount int
o.db.QueryRow("SELECT COUNT(*) FROM server_interface").Scan(&ifaceCount)

View File

@ -1,6 +1,7 @@
package sqlitedb
import (
"fmt"
"net"
"os"
"path/filepath"
@ -9,6 +10,7 @@ import (
"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"
@ -389,12 +391,12 @@ func TestGetAllocatedIPs_ExcludeClient(t *testing.T) {
now := time.Now().UTC()
db.SaveClient(model.Client{
ID: "c1", AllocatedIPs: []string{"10.252.1.2/32"},
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", AllocatedIPs: []string{"10.252.1.3/32"},
ID: "c2", Name: "Client B", AllocatedIPs: []string{"10.252.1.3/32"},
AllowedIPs: []string{}, ExtraAllowedIPs: []string{}, SubnetRanges: []string{},
CreatedAt: now, UpdatedAt: now,
})
@ -784,9 +786,9 @@ func TestMigrateFromJSON_FullMigration(t *testing.T) {
assert.Equal(t, "abc", hashes.Client)
assert.Equal(t, "def", hashes.Server)
user, err := db.GetUserByName("miguser")
require.NoError(t, err)
assert.Equal(t, "mig@test.com", user.Email)
// 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)
@ -930,6 +932,703 @@ func TestGetClients_SearchByNotes(t *testing.T) {
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")
}
func TestMigrateFromJSON_SkipsNonJSON(t *testing.T) {
os.Setenv("WGUI_ENDPOINT_ADDRESS", "10.0.0.1")
defer os.Unsetenv("WGUI_ENDPOINT_ADDRESS")

View File

@ -39,31 +39,31 @@ var (
)
const (
DefaultServerAddress = "10.252.1.0/24"
DefaultServerPort = 51820
DefaultDNS = "1.1.1.1"
DefaultMTU = 1450
DefaultPersistentKeepalive = 15
DefaultFirewallMark = "0xca6c" // i.e. 51820
DefaultTable = "auto"
DefaultConfigFilePath = "/etc/wireguard/wg0.conf"
FaviconFilePathEnvVar = "WGUI_FAVICON_FILE_PATH"
EndpointAddressEnvVar = "WGUI_ENDPOINT_ADDRESS"
DNSEnvVar = "WGUI_DNS"
MTUEnvVar = "WGUI_MTU"
PersistentKeepaliveEnvVar = "WGUI_PERSISTENT_KEEPALIVE"
FirewallMarkEnvVar = "WGUI_FIREWALL_MARK"
TableEnvVar = "WGUI_TABLE"
ConfigFilePathEnvVar = "WGUI_CONFIG_FILE_PATH"
LogLevel = "WGUI_LOG_LEVEL"
ServerAddressesEnvVar = "WGUI_SERVER_INTERFACE_ADDRESSES"
ServerListenPortEnvVar = "WGUI_SERVER_LISTEN_PORT"
ServerPostUpScriptEnvVar = "WGUI_SERVER_POST_UP_SCRIPT"
ServerPostDownScriptEnvVar = "WGUI_SERVER_POST_DOWN_SCRIPT"
DefaultClientAllowedIpsEnvVar = "WGUI_DEFAULT_CLIENT_ALLOWED_IPS"
DefaultClientExtraAllowedIpsEnvVar = "WGUI_DEFAULT_CLIENT_EXTRA_ALLOWED_IPS"
DefaultClientUseServerDNSEnvVar = "WGUI_DEFAULT_CLIENT_USE_SERVER_DNS"
DefaultClientEnableAfterCreationEnvVar = "WGUI_DEFAULT_CLIENT_ENABLE_AFTER_CREATION"
DefaultServerAddress = "10.252.1.0/24"
DefaultServerPort = 51820
DefaultDNS = "1.1.1.1"
DefaultMTU = 1450
DefaultPersistentKeepalive = 15
DefaultFirewallMark = "0xca6c" // i.e. 51820
DefaultTable = "auto"
DefaultConfigFilePath = "/etc/wireguard/wg0.conf"
FaviconFilePathEnvVar = "WGUI_FAVICON_FILE_PATH"
EndpointAddressEnvVar = "WGUI_ENDPOINT_ADDRESS"
DNSEnvVar = "WGUI_DNS"
MTUEnvVar = "WGUI_MTU"
PersistentKeepaliveEnvVar = "WGUI_PERSISTENT_KEEPALIVE"
FirewallMarkEnvVar = "WGUI_FIREWALL_MARK"
TableEnvVar = "WGUI_TABLE"
ConfigFilePathEnvVar = "WGUI_CONFIG_FILE_PATH"
LogLevel = "WGUI_LOG_LEVEL"
ServerAddressesEnvVar = "WGUI_SERVER_INTERFACE_ADDRESSES"
ServerListenPortEnvVar = "WGUI_SERVER_LISTEN_PORT"
ServerPostUpScriptEnvVar = "WGUI_SERVER_POST_UP_SCRIPT"
ServerPostDownScriptEnvVar = "WGUI_SERVER_POST_DOWN_SCRIPT"
DefaultClientAllowedIpsEnvVar = "WGUI_DEFAULT_CLIENT_ALLOWED_IPS"
DefaultClientExtraAllowedIpsEnvVar = "WGUI_DEFAULT_CLIENT_EXTRA_ALLOWED_IPS"
DefaultClientUseServerDNSEnvVar = "WGUI_DEFAULT_CLIENT_USE_SERVER_DNS"
ConfigApplyDelayEnvVar = "WGUI_CONFIG_APPLY_DELAY"
// OIDC env vars
OIDCIssuerURLEnvVar = "OIDC_ISSUER_URL"

View File

@ -98,7 +98,6 @@ func ClientDefaultsFromEnv() model.ClientDefaults {
clientDefaults.AllowedIps = LookupEnvOrStrings(DefaultClientAllowedIpsEnvVar, []string{"0.0.0.0/0"})
clientDefaults.ExtraAllowedIps = LookupEnvOrStrings(DefaultClientExtraAllowedIpsEnvVar, []string{})
clientDefaults.UseServerDNS = LookupEnvOrBool(DefaultClientUseServerDNSEnvVar, true)
clientDefaults.EnableAfterCreation = LookupEnvOrBool(DefaultClientEnableAfterCreationEnvVar, true)
return clientDefaults
}

View File

@ -513,26 +513,20 @@ func TestClientDefaultsFromEnv(t *testing.T) {
os.Unsetenv(DefaultClientAllowedIpsEnvVar)
os.Unsetenv(DefaultClientExtraAllowedIpsEnvVar)
os.Unsetenv(DefaultClientUseServerDNSEnvVar)
os.Unsetenv(DefaultClientEnableAfterCreationEnvVar)
defaults := ClientDefaultsFromEnv()
assert.Equal(t, []string{"0.0.0.0/0"}, defaults.AllowedIps)
assert.Equal(t, []string{}, defaults.ExtraAllowedIps)
assert.True(t, defaults.UseServerDNS)
assert.True(t, defaults.EnableAfterCreation)
// test with env overrides
os.Setenv(DefaultClientAllowedIpsEnvVar, "10.0.0.0/8,192.168.0.0/16")
os.Setenv(DefaultClientUseServerDNSEnvVar, "false")
os.Setenv(DefaultClientEnableAfterCreationEnvVar, "false")
defer os.Unsetenv(DefaultClientAllowedIpsEnvVar)
defer os.Unsetenv(DefaultClientUseServerDNSEnvVar)
defer os.Unsetenv(DefaultClientEnableAfterCreationEnvVar)
defaults = ClientDefaultsFromEnv()
assert.Equal(t, []string{"10.0.0.0/8", "192.168.0.0/16"}, defaults.AllowedIps)
assert.False(t, defaults.UseServerDNS)
assert.False(t, defaults.EnableAfterCreation)
}
func TestGetCurrentHash_WithMockStore(t *testing.T) {