Files
wireguard-ui-multi/store/jsondb/jsondb.go
T
sysops 5b72a7c118 Add per-server access control (User.ServerIDs)
Non-admin users are now restricted to servers explicitly listed in
their new ServerIDs field; empty means no access (secure by default).
Admins always have full access. Migration backfills existing users'
ServerIDs with the migrated legacy server so nobody is locked out on
upgrade. New RequireServerAccess middleware enforces this on
/servers/:id/... routes (applied to GET /servers/:id/clients so far);
GET /servers also filters its list for non-admins.
2026-07-11 23:20:52 +02:00

687 lines
22 KiB
Go

package jsondb
import (
"encoding/base64"
"encoding/json"
"fmt"
"log"
"os"
"path"
"strconv"
"strings"
"time"
"github.com/sdomino/scribble"
"github.com/skip2/go-qrcode"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
"github.com/ngoduykhanh/wireguard-ui/model"
"github.com/ngoduykhanh/wireguard-ui/util"
)
// legacyDefaultServerID is the synthetic ID assigned to a pre-existing
// single-server installation when it is migrated to the multi-server layout.
const legacyDefaultServerID = "wg0"
type JsonDB struct {
conn *scribble.Driver
dbPath string
}
// New returns a new pointer JsonDB
func New(dbPath string) (*JsonDB, error) {
conn, err := scribble.New(dbPath, nil)
if err != nil {
return nil, err
}
ans := JsonDB{
conn: conn,
dbPath: dbPath,
}
return &ans, nil
}
func (o *JsonDB) Init() error {
var clientPath = path.Join(o.dbPath, "clients")
var serverPath = path.Join(o.dbPath, "server")
var userPath = path.Join(o.dbPath, "users")
var wakeOnLanHostsPath = path.Join(o.dbPath, "wake_on_lan_hosts")
var serverInterfacePath = path.Join(serverPath, "interfaces.json")
var serverKeyPairPath = path.Join(serverPath, "keypair.json")
var globalSettingPath = path.Join(serverPath, "global_settings.json")
var hashesPath = path.Join(serverPath, "hashes.json")
// create directories if they do not exist
if _, err := os.Stat(clientPath); os.IsNotExist(err) {
os.MkdirAll(clientPath, os.ModePerm)
}
if _, err := os.Stat(serverPath); os.IsNotExist(err) {
os.MkdirAll(serverPath, os.ModePerm)
}
if _, err := os.Stat(userPath); os.IsNotExist(err) {
os.MkdirAll(userPath, os.ModePerm)
}
if _, err := os.Stat(wakeOnLanHostsPath); os.IsNotExist(err) {
os.MkdirAll(wakeOnLanHostsPath, os.ModePerm)
}
// server's interface
if _, err := os.Stat(serverInterfacePath); os.IsNotExist(err) {
serverInterface := new(model.ServerInterface)
serverInterface.Addresses = util.LookupEnvOrStrings(util.ServerAddressesEnvVar, []string{util.DefaultServerAddress})
serverInterface.ListenPort = util.LookupEnvOrInt(util.ServerListenPortEnvVar, util.DefaultServerPort)
serverInterface.PostUp = util.LookupEnvOrString(util.ServerPostUpScriptEnvVar, "")
serverInterface.PostDown = util.LookupEnvOrString(util.ServerPostDownScriptEnvVar, "")
serverInterface.UpdatedAt = time.Now().UTC()
o.conn.Write("server", "interfaces", serverInterface)
err := util.ManagePerms(serverInterfacePath)
if err != nil {
return err
}
}
// server's key pair
if _, err := os.Stat(serverKeyPairPath); os.IsNotExist(err) {
key, err := wgtypes.GeneratePrivateKey()
if err != nil {
return scribble.ErrMissingCollection
}
serverKeyPair := new(model.ServerKeypair)
serverKeyPair.PrivateKey = key.String()
serverKeyPair.PublicKey = key.PublicKey().String()
serverKeyPair.UpdatedAt = time.Now().UTC()
o.conn.Write("server", "keypair", serverKeyPair)
err = util.ManagePerms(serverKeyPairPath)
if err != nil {
return err
}
}
// global settings
if _, err := os.Stat(globalSettingPath); os.IsNotExist(err) {
endpointAddress := util.LookupEnvOrString(util.EndpointAddressEnvVar, "")
if endpointAddress == "" {
// automatically find an external IP address
publicInterface, err := util.GetPublicIP()
if err != nil {
return err
}
endpointAddress = publicInterface.IPAddress
}
globalSetting := new(model.GlobalSetting)
globalSetting.EndpointAddress = endpointAddress
globalSetting.DNSServers = util.LookupEnvOrStrings(util.DNSEnvVar, []string{util.DefaultDNS})
globalSetting.MTU = util.LookupEnvOrInt(util.MTUEnvVar, util.DefaultMTU)
globalSetting.PersistentKeepalive = util.LookupEnvOrInt(util.PersistentKeepaliveEnvVar, util.DefaultPersistentKeepalive)
globalSetting.FirewallMark = util.LookupEnvOrString(util.FirewallMarkEnvVar, util.DefaultFirewallMark)
globalSetting.Table = util.LookupEnvOrString(util.TableEnvVar, util.DefaultTable)
globalSetting.ConfigFilePath = util.LookupEnvOrString(util.ConfigFilePathEnvVar, util.DefaultConfigFilePath)
globalSetting.UpdatedAt = time.Now().UTC()
o.conn.Write("server", "global_settings", globalSetting)
err := util.ManagePerms(globalSettingPath)
if err != nil {
return err
}
}
// hashes
if _, err := os.Stat(hashesPath); os.IsNotExist(err) {
clientServerHashes := new(model.ClientServerHashes)
clientServerHashes.Client = "none"
clientServerHashes.Server = "none"
o.conn.Write("server", "hashes", clientServerHashes)
err := util.ManagePerms(hashesPath)
if err != nil {
return err
}
}
// user info
results, err := o.conn.ReadAll("users")
if err != nil || len(results) < 1 {
user := new(model.User)
user.Username = util.LookupEnvOrString(util.UsernameEnvVar, util.DefaultUsername)
user.Admin = util.DefaultIsAdmin
user.PasswordHash = util.LookupEnvOrString(util.PasswordHashEnvVar, "")
if user.PasswordHash == "" {
user.PasswordHash = util.LookupEnvOrFile(util.PasswordHashFileEnvVar, "")
if user.PasswordHash == "" {
plaintext := util.LookupEnvOrString(util.PasswordEnvVar, util.DefaultPassword)
if plaintext == util.DefaultPassword {
plaintext = util.LookupEnvOrFile(util.PasswordFileEnvVar, util.DefaultPassword)
}
hash, err := util.HashPassword(plaintext)
if err != nil {
return err
}
user.PasswordHash = hash
}
}
o.conn.Write("users", user.Username, user)
results, _ = o.conn.ReadAll("users")
err = util.ManagePerms(path.Join(path.Join(o.dbPath, "users"), user.Username+".json"))
if err != nil {
return err
}
}
// migrate legacy single-server layout to the new multi-server layout
if err := o.migrateLegacyServer(); err != nil {
return err
}
// init cache
for _, i := range results {
user := model.User{}
if err := json.Unmarshal([]byte(i), &user); err == nil {
util.DBUsersToCRC32[user.Username] = util.GetDBUserCRC32(user)
}
}
clients, err := o.GetClients(false)
if err != nil {
return nil
}
for _, cl := range clients {
client := cl.Client
if client.Enabled && len(client.TgUserid) > 0 {
if userid, err := strconv.ParseInt(client.TgUserid, 10, 64); err == nil {
util.UpdateTgToClientID(userid, client.ID)
}
}
}
return nil
}
// migrateLegacyServer detects a pre-multi-server DB layout (single
// "server" collection, no "servers" registry yet) and migrates it into
// the new per-server layout ("servers/<id>.json", "server_settings/<id>.json",
// "server_hashes/<id>.json"), backfilling ServerID on every existing client.
// It is idempotent: it only runs once, guarded by the absence of the
// "servers" directory, and never deletes or modifies the legacy "server"
// collection - existing store methods keep reading/writing it directly
// until they are switched over to the new layout in a later step.
func (o *JsonDB) migrateLegacyServer() error {
legacyServerPath := path.Join(o.dbPath, "server")
legacyInterfacePath := path.Join(legacyServerPath, "interfaces.json")
newServersPath := path.Join(o.dbPath, "servers")
if _, err := os.Stat(legacyInterfacePath); os.IsNotExist(err) {
// nothing to migrate (fresh install)
return nil
}
if _, err := os.Stat(newServersPath); err == nil {
// already migrated
return nil
}
var legacyInterface model.ServerInterface
if err := o.conn.Read("server", "interfaces", &legacyInterface); err != nil {
return fmt.Errorf("migration: cannot read legacy server interface: %v", err)
}
var legacyKeyPair model.ServerKeypair
if err := o.conn.Read("server", "keypair", &legacyKeyPair); err != nil {
return fmt.Errorf("migration: cannot read legacy server keypair: %v", err)
}
var legacyGlobalSettings model.GlobalSetting
if err := o.conn.Read("server", "global_settings", &legacyGlobalSettings); err != nil {
return fmt.Errorf("migration: cannot read legacy global settings: %v", err)
}
var legacyHashes model.ClientServerHashes
hasHashes := true
if err := o.conn.Read("server", "hashes", &legacyHashes); err != nil {
hasHashes = false
}
serverID := legacyDefaultServerID
ifaceName := legacyDefaultServerID
if legacyGlobalSettings.ConfigFilePath != "" {
base := path.Base(legacyGlobalSettings.ConfigFilePath)
if strings.HasSuffix(base, ".conf") {
ifaceName = strings.TrimSuffix(base, ".conf")
}
}
legacyInterface.Name = ifaceName
server := model.Server{
ID: serverID,
Name: "Default Server",
KeyPair: &legacyKeyPair,
Interface: &legacyInterface,
}
if err := o.conn.Write("servers", serverID, server); err != nil {
return fmt.Errorf("migration: cannot write servers/%s.json: %v", serverID, err)
}
if err := util.ManagePerms(path.Join(newServersPath, serverID+".json")); err != nil {
return err
}
serverSetting := model.ServerSetting{
EndpointAddress: legacyGlobalSettings.EndpointAddress,
FirewallMark: legacyGlobalSettings.FirewallMark,
Table: legacyGlobalSettings.Table,
ConfigFilePath: legacyGlobalSettings.ConfigFilePath,
UpdatedAt: time.Now().UTC(),
}
if err := o.conn.Write("server_settings", serverID, serverSetting); err != nil {
return fmt.Errorf("migration: cannot write server_settings/%s.json: %v", serverID, err)
}
if err := util.ManagePerms(path.Join(o.dbPath, "server_settings", serverID+".json")); err != nil {
return err
}
if hasHashes {
if err := o.conn.Write("server_hashes", serverID, legacyHashes); err != nil {
return fmt.Errorf("migration: cannot write server_hashes/%s.json: %v", serverID, err)
}
if err := util.ManagePerms(path.Join(o.dbPath, "server_hashes", serverID+".json")); err != nil {
return err
}
}
// backfill ServerID on every existing client
clientRecords, err := o.conn.ReadAll("clients")
if err != nil && err != scribble.ErrMissingCollection {
return fmt.Errorf("migration: cannot read clients: %v", err)
}
for _, rec := range clientRecords {
var client model.Client
if err := json.Unmarshal(rec, &client); err != nil {
return fmt.Errorf("migration: cannot decode client json: %v", err)
}
if client.ServerID != "" {
continue
}
client.ServerID = serverID
if err := o.conn.Write("clients", client.ID, client); err != nil {
return fmt.Errorf("migration: cannot backfill server_id on client %s: %v", client.ID, err)
}
}
// backfill ServerIDs on every existing user so nobody is locked out of
// the migrated server by the new per-server access control (users
// created after migration default to no access, per model.User).
userRecords, err := o.conn.ReadAll("users")
if err != nil && err != scribble.ErrMissingCollection {
return fmt.Errorf("migration: cannot read users: %v", err)
}
for _, rec := range userRecords {
var user model.User
if err := json.Unmarshal(rec, &user); err != nil {
return fmt.Errorf("migration: cannot decode user json: %v", err)
}
if len(user.ServerIDs) > 0 {
continue
}
user.ServerIDs = []string{serverID}
if err := o.conn.Write("users", user.Username, user); err != nil {
return fmt.Errorf("migration: cannot backfill server_ids on user %s: %v", user.Username, err)
}
}
// NOTE: the legacy "server" directory is intentionally left in place
// (not renamed/removed) at this stage. Existing store methods (GetServer,
// GetGlobalSettings, SaveServerInterface, ...) still read/write it
// directly until they are switched over to the new per-server layout;
// removing it now would break the running app. The rename-to-backup
// step happens once those methods are migrated.
log.Printf("Migrated legacy single-server DB to multi-server layout (server id: %s); legacy db/server/ kept in place for now", serverID)
return nil
}
// GetUsers func to get all users from the database
func (o *JsonDB) GetUsers() ([]model.User, error) {
var users []model.User
results, err := o.conn.ReadAll("users")
if err != nil {
return users, err
}
for _, i := range results {
user := model.User{}
if err := json.Unmarshal(i, &user); err != nil {
return users, fmt.Errorf("cannot decode user json structure: %v", err)
}
users = append(users, user)
}
return users, err
}
// GetUserByName func to get single user from the database
func (o *JsonDB) GetUserByName(username string) (model.User, error) {
user := model.User{}
if err := o.conn.Read("users", username, &user); err != nil {
return user, err
}
return user, nil
}
// SaveUser func to save user in the database
func (o *JsonDB) SaveUser(user model.User) error {
userPath := path.Join(path.Join(o.dbPath, "users"), user.Username+".json")
output := o.conn.Write("users", user.Username, user)
err := util.ManagePerms(userPath)
if err != nil {
return err
}
util.DBUsersToCRC32[user.Username] = util.GetDBUserCRC32(user)
return output
}
// DeleteUser func to remove user from the database
func (o *JsonDB) DeleteUser(username string) error {
delete(util.DBUsersToCRC32, username)
return o.conn.Delete("users", username)
}
// GetGlobalSettings func to query global settings from the database
func (o *JsonDB) GetGlobalSettings() (model.GlobalSetting, error) {
settings := model.GlobalSetting{}
return settings, o.conn.Read("server", "global_settings", &settings)
}
// GetServer func to query Server settings from the database
func (o *JsonDB) GetServer() (model.Server, error) {
server := model.Server{}
// read server interface information
serverInterface := model.ServerInterface{}
if err := o.conn.Read("server", "interfaces", &serverInterface); err != nil {
return server, err
}
// read server key pair information
serverKeyPair := model.ServerKeypair{}
if err := o.conn.Read("server", "keypair", &serverKeyPair); err != nil {
return server, err
}
// create Server object and return
server.Interface = &serverInterface
server.KeyPair = &serverKeyPair
return server, nil
}
func (o *JsonDB) GetClients(hasQRCode bool) ([]model.ClientData, error) {
var clients []model.ClientData
// read all client json files in "clients" directory
records, err := o.conn.ReadAll("clients")
if err != nil {
return clients, err
}
// build the ClientData list
for _, f := range records {
client := model.Client{}
clientData := model.ClientData{}
// get client info
if err := json.Unmarshal(f, &client); err != nil {
return clients, fmt.Errorf("cannot decode client json structure: %v", err)
}
// generate client qrcode image in base64
if hasQRCode && client.PrivateKey != "" {
server, _ := o.GetServer()
globalSettings, _ := o.GetGlobalSettings()
png, err := qrcode.Encode(util.BuildClientConfig(client, server, globalSettings), qrcode.Medium, 256)
if err == nil {
clientData.QRCode = "data:image/png;base64," + base64.StdEncoding.EncodeToString(png)
} else {
fmt.Print("Cannot generate QR code: ", err)
}
}
// create the list of clients and their qrcode data
clientData.Client = &client
clients = append(clients, clientData)
}
return clients, nil
}
func (o *JsonDB) GetClientByID(clientID string, qrCodeSettings model.QRCodeSettings) (model.ClientData, error) {
client := model.Client{}
clientData := model.ClientData{}
// read client information
if err := o.conn.Read("clients", clientID, &client); err != nil {
return clientData, err
}
// generate client qrcode image in base64
if qrCodeSettings.Enabled && client.PrivateKey != "" {
server, _ := o.GetServer()
globalSettings, _ := o.GetGlobalSettings()
client := client
if !qrCodeSettings.IncludeDNS {
globalSettings.DNSServers = []string{}
}
if !qrCodeSettings.IncludeMTU {
globalSettings.MTU = 0
}
png, err := qrcode.Encode(util.BuildClientConfig(client, server, globalSettings), qrcode.Medium, 256)
if err == nil {
clientData.QRCode = "data:image/png;base64," + base64.StdEncoding.EncodeToString(png)
} else {
fmt.Print("Cannot generate QR code: ", err)
}
}
clientData.Client = &client
return clientData, nil
}
func (o *JsonDB) SaveClient(client model.Client) error {
clientPath := path.Join(path.Join(o.dbPath, "clients"), client.ID+".json")
output := o.conn.Write("clients", client.ID, client)
if output == nil {
if client.Enabled && len(client.TgUserid) > 0 {
if userid, err := strconv.ParseInt(client.TgUserid, 10, 64); err == nil {
util.UpdateTgToClientID(userid, client.ID)
}
} else {
util.RemoveTgToClientID(client.ID)
}
} else {
util.RemoveTgToClientID(client.ID)
}
err := util.ManagePerms(clientPath)
if err != nil {
return err
}
return output
}
func (o *JsonDB) DeleteClient(clientID string) error {
util.RemoveTgToClientID(clientID)
return o.conn.Delete("clients", clientID)
}
func (o *JsonDB) SaveServerInterface(serverInterface model.ServerInterface) error {
serverInterfacePath := path.Join(path.Join(o.dbPath, "server"), "interfaces.json")
output := o.conn.Write("server", "interfaces", serverInterface)
err := util.ManagePerms(serverInterfacePath)
if err != nil {
return err
}
return output
}
func (o *JsonDB) SaveServerKeyPair(serverKeyPair model.ServerKeypair) error {
serverKeyPairPath := path.Join(path.Join(o.dbPath, "server"), "keypair.json")
output := o.conn.Write("server", "keypair", serverKeyPair)
err := util.ManagePerms(serverKeyPairPath)
if err != nil {
return err
}
return output
}
func (o *JsonDB) SaveGlobalSettings(globalSettings model.GlobalSetting) error {
globalSettingsPath := path.Join(path.Join(o.dbPath, "server"), "global_settings.json")
output := o.conn.Write("server", "global_settings", globalSettings)
err := util.ManagePerms(globalSettingsPath)
if err != nil {
return err
}
return output
}
func (o *JsonDB) GetPath() string {
return o.dbPath
}
func (o *JsonDB) GetHashes() (model.ClientServerHashes, error) {
hashes := model.ClientServerHashes{}
return hashes, o.conn.Read("server", "hashes", &hashes)
}
func (o *JsonDB) SaveHashes(hashes model.ClientServerHashes) error {
hashesPath := path.Join(path.Join(o.dbPath, "server"), "hashes.json")
output := o.conn.Write("server", "hashes", hashes)
err := util.ManagePerms(hashesPath)
if err != nil {
return err
}
return output
}
// GetServers func to get all servers from the database
func (o *JsonDB) GetServers() ([]model.Server, error) {
var servers []model.Server
results, err := o.conn.ReadAll("servers")
if err != nil {
if err == scribble.ErrMissingCollection {
return servers, nil
}
return servers, err
}
for _, i := range results {
server := model.Server{}
if err := json.Unmarshal([]byte(i), &server); err != nil {
return servers, fmt.Errorf("cannot decode server json structure: %v", err)
}
servers = append(servers, server)
}
return servers, nil
}
// validateServerID checks a server ID before it is ever used as a jsondb
// record key, shared by every server-scoped store method below.
func validateServerID(serverID string) error {
if !util.ValidateRecordID(serverID) {
return fmt.Errorf("invalid server id: %s", serverID)
}
return nil
}
// GetServerByID func to query a single server by id from the database
func (o *JsonDB) GetServerByID(serverID string) (model.Server, error) {
server := model.Server{}
if err := validateServerID(serverID); err != nil {
return server, err
}
err := o.conn.Read("servers", serverID, &server)
return server, err
}
// CreateServer func to create a new server in the database
func (o *JsonDB) CreateServer(server model.Server) error {
if server.ID == "" {
return fmt.Errorf("cannot create server: missing id")
}
if err := validateServerID(server.ID); err != nil {
return err
}
if server.Interface != nil && server.Interface.Name != "" && !util.ValidateInterfaceName(server.Interface.Name) {
return fmt.Errorf("invalid interface name: %s", server.Interface.Name)
}
if err := o.conn.Write("servers", server.ID, server); err != nil {
return err
}
return util.ManagePerms(path.Join(o.dbPath, "servers", server.ID+".json"))
}
// DeleteServer func to remove a server from the database, refusing if
// any client still references it
func (o *JsonDB) DeleteServer(serverID string) error {
if err := validateServerID(serverID); err != nil {
return err
}
records, err := o.conn.ReadAll("clients")
if err != nil && err != scribble.ErrMissingCollection {
return err
}
count := 0
for _, rec := range records {
var client model.Client
if err := json.Unmarshal(rec, &client); err != nil {
return fmt.Errorf("cannot decode client json structure: %v", err)
}
if client.ServerID == serverID {
count++
}
}
if count > 0 {
return fmt.Errorf("cannot delete server %s: %d client(s) still reference it", serverID, count)
}
return o.conn.Delete("servers", serverID)
}
// GetServerSettings func to query per-server settings from the database
func (o *JsonDB) GetServerSettings(serverID string) (model.ServerSetting, error) {
settings := model.ServerSetting{}
if err := validateServerID(serverID); err != nil {
return settings, err
}
return settings, o.conn.Read("server_settings", serverID, &settings)
}
// SaveServerSettings func to save per-server settings in the database
func (o *JsonDB) SaveServerSettings(serverID string, settings model.ServerSetting) error {
if err := validateServerID(serverID); err != nil {
return err
}
settingsPath := path.Join(o.dbPath, "server_settings", serverID+".json")
output := o.conn.Write("server_settings", serverID, settings)
err := util.ManagePerms(settingsPath)
if err != nil {
return err
}
return output
}
// GetServerHashes func to query per-server client/server hashes from the database
func (o *JsonDB) GetServerHashes(serverID string) (model.ClientServerHashes, error) {
hashes := model.ClientServerHashes{}
if err := validateServerID(serverID); err != nil {
return hashes, err
}
return hashes, o.conn.Read("server_hashes", serverID, &hashes)
}
// SaveServerHashes func to save per-server client/server hashes in the database
func (o *JsonDB) SaveServerHashes(serverID string, hashes model.ClientServerHashes) error {
if err := validateServerID(serverID); err != nil {
return err
}
hashesPath := path.Join(o.dbPath, "server_hashes", serverID+".json")
output := o.conn.Write("server_hashes", serverID, hashes)
err := util.ManagePerms(hashesPath)
if err != nil {
return err
}
return output
}