parent
0c50253d1b
commit
dcf5c0fd45
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
@ -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")
|
||||
}
|
||||
|
|
@ -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) {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
8
main.go
8
main.go
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
@ -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>
|
||||
)}
|
||||
|
|
|
|||
|
|
@ -109,7 +109,6 @@ export interface ClientDefaults {
|
|||
AllowedIps: string[];
|
||||
ExtraAllowedIps: string[];
|
||||
UseServerDNS: boolean;
|
||||
EnableAfterCreation: boolean;
|
||||
}
|
||||
|
||||
export interface AppInfo {
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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 }))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 '[]',
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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) {
|
||||
|
|
|
|||
Loading…
Reference in New Issue