Files
wireguard-ui-multi/store/jsondb/jsondb.go
T
sysopsandClaude Sonnet 5 0dbb916866 Fix nil-pointer crash generating QR codes for non-default servers
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>
2026-07-12 23:42:36 +02:00

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)
}