GetClients, GetClientByID, and SendRequestedConfigsToTelegram all looked up the legacy "wg0" server unconditionally when building a client's QR code / config, ignoring which server the client actually belongs to. On any multi-server setup without a migrated wg0, the server lookup returned a zero-value model.Server (error discarded), and BuildClientConfig then dereferenced its nil KeyPair pointer, panicking whenever a client's QR code was rendered. Now looks up the client's own ServerID (falling back to wg0 only for legacy clients with no ServerID set) and surfaces the lookup error instead of silently continuing with an empty server. Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
798 lines
26 KiB
Go
798 lines
26 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 = util.DefaultServerID
|
|
|
|
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.PreUp = util.LookupEnvOrString(util.ServerPreUpScriptEnvVar, "")
|
|
serverInterface.PostUp = util.LookupEnvOrString(util.ServerPostUpScriptEnvVar, "")
|
|
serverInterface.PreDown = util.LookupEnvOrString(util.ServerPreDownScriptEnvVar, "")
|
|
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). GetGlobalSettings/SaveGlobalSettings still read
|
|
// and write the app-wide settings record there (DNS/MTU/
|
|
// PersistentKeepalive are not yet per-server concerns); every other
|
|
// server-scoped read/write now goes exclusively through the per-server
|
|
// registry (servers/<id>.json via GetServerByID/UpdateServerInterface/
|
|
// UpdateServerKeyPair/GetServerSettings), so this directory is kept only
|
|
// as a landing spot for one-time migration and for global settings.
|
|
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)
|
|
}
|
|
|
|
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 != "" {
|
|
serverID := client.ServerID
|
|
if serverID == "" {
|
|
serverID = util.DefaultServerID
|
|
}
|
|
server, err := o.GetServerByID(serverID)
|
|
if err != nil {
|
|
fmt.Printf("Cannot generate QR code: server %s not found: %v\n", serverID, err)
|
|
} else {
|
|
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 != "" {
|
|
serverID := client.ServerID
|
|
if serverID == "" {
|
|
serverID = util.DefaultServerID
|
|
}
|
|
server, err := o.GetServerByID(serverID)
|
|
if err != nil {
|
|
fmt.Printf("Cannot generate QR code: server %s not found: %v\n", serverID, err)
|
|
} else {
|
|
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) 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
|
|
}
|
|
|
|
// UpdateServerInterface func updates the Interface of an existing server
|
|
func (o *JsonDB) UpdateServerInterface(serverID string, serverInterface model.ServerInterface) error {
|
|
if err := validateServerID(serverID); err != nil {
|
|
return err
|
|
}
|
|
server, err := o.GetServerByID(serverID)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
server.Interface = &serverInterface
|
|
if err := o.conn.Write("servers", serverID, server); err != nil {
|
|
return err
|
|
}
|
|
return util.ManagePerms(path.Join(o.dbPath, "servers", serverID+".json"))
|
|
}
|
|
|
|
// UpdateServerKeyPair func updates the KeyPair of an existing server
|
|
func (o *JsonDB) UpdateServerKeyPair(serverID string, serverKeyPair model.ServerKeypair) error {
|
|
if err := validateServerID(serverID); err != nil {
|
|
return err
|
|
}
|
|
server, err := o.GetServerByID(serverID)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
server.KeyPair = &serverKeyPair
|
|
if err := o.conn.Write("servers", serverID, server); err != nil {
|
|
return err
|
|
}
|
|
return util.ManagePerms(path.Join(o.dbPath, "servers", serverID+".json"))
|
|
}
|
|
|
|
// 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 isMissingCollectionErr(err) {
|
|
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
|
|
}
|
|
|
|
// isMissingCollectionErr reports whether err from scribble.ReadAll just
|
|
// means "this collection doesn't exist yet" (nothing written there so
|
|
// far). scribble only returns its own ErrMissingCollection when the
|
|
// collection name itself is empty - if the collection's directory simply
|
|
// hasn't been created yet, ReadAll instead returns a raw os.IsNotExist
|
|
// error, which must be treated the same way (empty collection, not a
|
|
// real failure).
|
|
func isMissingCollectionErr(err error) bool {
|
|
return err == scribble.ErrMissingCollection || os.IsNotExist(err)
|
|
}
|
|
|
|
// 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
|
|
}
|
|
|
|
// GetFirewallRules func to query all firewall rules belonging to a server
|
|
func (o *JsonDB) GetFirewallRules(serverID string) ([]model.FirewallRule, error) {
|
|
if err := validateServerID(serverID); err != nil {
|
|
return nil, err
|
|
}
|
|
rules := make([]model.FirewallRule, 0)
|
|
records, err := o.conn.ReadAll("firewall_rules")
|
|
if err != nil {
|
|
if isMissingCollectionErr(err) {
|
|
return rules, nil
|
|
}
|
|
return nil, err
|
|
}
|
|
for _, rec := range records {
|
|
var rule model.FirewallRule
|
|
if err := json.Unmarshal(rec, &rule); err != nil {
|
|
return nil, fmt.Errorf("cannot decode firewall rule json structure: %v", err)
|
|
}
|
|
if rule.ServerID == serverID {
|
|
rules = append(rules, rule)
|
|
}
|
|
}
|
|
return rules, nil
|
|
}
|
|
|
|
// CreateFirewallRule func to add a new firewall rule for a server
|
|
func (o *JsonDB) CreateFirewallRule(rule model.FirewallRule) error {
|
|
if err := validateServerID(rule.ServerID); err != nil {
|
|
return err
|
|
}
|
|
return o.conn.Write("firewall_rules", rule.ID, rule)
|
|
}
|
|
|
|
// UpdateFirewallRule func to update an existing firewall rule, refusing to
|
|
// move it to a different server than it currently belongs to
|
|
func (o *JsonDB) UpdateFirewallRule(rule model.FirewallRule) error {
|
|
var existing model.FirewallRule
|
|
if err := o.conn.Read("firewall_rules", rule.ID, &existing); err != nil {
|
|
return fmt.Errorf("firewall rule not found: %v", err)
|
|
}
|
|
if existing.ServerID != rule.ServerID {
|
|
return fmt.Errorf("cannot change server_id of an existing firewall rule")
|
|
}
|
|
return o.conn.Write("firewall_rules", rule.ID, rule)
|
|
}
|
|
|
|
// DeleteFirewallRule func to remove a firewall rule, verifying it belongs to
|
|
// the given server first
|
|
func (o *JsonDB) DeleteFirewallRule(serverID, ruleID string) error {
|
|
var existing model.FirewallRule
|
|
if err := o.conn.Read("firewall_rules", ruleID, &existing); err != nil {
|
|
return fmt.Errorf("firewall rule not found: %v", err)
|
|
}
|
|
if existing.ServerID != serverID {
|
|
return fmt.Errorf("firewall rule does not belong to server %s", serverID)
|
|
}
|
|
return o.conn.Delete("firewall_rules", ruleID)
|
|
}
|
|
|
|
// GetIPListEntries func to query every host-wide allow/block list entry
|
|
func (o *JsonDB) GetIPListEntries() ([]model.IPListEntry, error) {
|
|
entries := make([]model.IPListEntry, 0)
|
|
records, err := o.conn.ReadAll("ip_list_entries")
|
|
if err != nil {
|
|
if isMissingCollectionErr(err) {
|
|
return entries, nil
|
|
}
|
|
return nil, err
|
|
}
|
|
for _, rec := range records {
|
|
var entry model.IPListEntry
|
|
if err := json.Unmarshal(rec, &entry); err != nil {
|
|
return nil, fmt.Errorf("cannot decode ip list entry json structure: %v", err)
|
|
}
|
|
entries = append(entries, entry)
|
|
}
|
|
return entries, nil
|
|
}
|
|
|
|
// CreateIPListEntry func to add a new host-wide allow/block list entry
|
|
func (o *JsonDB) CreateIPListEntry(entry model.IPListEntry) error {
|
|
return o.conn.Write("ip_list_entries", entry.ID, entry)
|
|
}
|
|
|
|
// DeleteIPListEntry func to remove a host-wide allow/block list entry
|
|
func (o *JsonDB) DeleteIPListEntry(id string) error {
|
|
return o.conn.Delete("ip_list_entries", id)
|
|
}
|