Add server-registry store methods with record-ID validation (step 1)

New IStore methods for a per-server registry (GetServers, GetServerByID,
CreateServer, DeleteServer, GetServerSettings/SaveServerSettings,
GetServerHashes/SaveServerHashes), all additive - existing single-server
methods untouched. Adds util.ValidateRecordID/ValidateInterfaceName and
applies them to every new method taking a server ID, closing a path-
traversal gap before serverID is ever driven by user input (flagged by
a security review pass).
This commit is contained in:
sysops
2026-07-11 22:39:53 +02:00
parent 40d7862702
commit 7eca2d01b8
3 changed files with 153 additions and 0 deletions
+119
View File
@@ -535,3 +535,122 @@ func (o *JsonDB) SaveHashes(hashes model.ClientServerHashes) error {
} }
return output 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
}
// 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 !util.ValidateRecordID(serverID) {
return server, fmt.Errorf("invalid server id: %s", serverID)
}
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 !util.ValidateRecordID(server.ID) {
return fmt.Errorf("invalid server id: %s", server.ID)
}
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 !util.ValidateRecordID(serverID) {
return fmt.Errorf("invalid server id: %s", serverID)
}
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 !util.ValidateRecordID(serverID) {
return settings, fmt.Errorf("invalid server id: %s", serverID)
}
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 !util.ValidateRecordID(serverID) {
return fmt.Errorf("invalid server id: %s", serverID)
}
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 !util.ValidateRecordID(serverID) {
return hashes, fmt.Errorf("invalid server id: %s", serverID)
}
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 !util.ValidateRecordID(serverID) {
return fmt.Errorf("invalid server id: %s", serverID)
}
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
}
+8
View File
@@ -27,4 +27,12 @@ type IStore interface {
GetPath() string GetPath() string
SaveHashes(hashes model.ClientServerHashes) error SaveHashes(hashes model.ClientServerHashes) error
GetHashes() (model.ClientServerHashes, error) GetHashes() (model.ClientServerHashes, error)
GetServers() ([]model.Server, error)
GetServerByID(serverID string) (model.Server, error)
CreateServer(server model.Server) error
DeleteServer(serverID string) error
GetServerSettings(serverID string) (model.ServerSetting, error)
SaveServerSettings(serverID string, settings model.ServerSetting) error
GetServerHashes(serverID string) (model.ClientServerHashes, error)
SaveServerHashes(serverID string, hashes model.ClientServerHashes) error
} }
+26
View File
@@ -15,6 +15,7 @@ import (
"os" "os"
"path" "path"
"path/filepath" "path/filepath"
"regexp"
"strconv" "strconv"
"strings" "strings"
"text/template" "text/template"
@@ -165,6 +166,31 @@ func ValidateServerAddresses(cidrs []string) bool {
return true return true
} }
// recordIDPattern restricts identifiers that end up as scribble/jsondb
// record keys (and therefore as filesystem path components) to a safe,
// predictable character set - no "/", "..", or other path metacharacters.
var recordIDPattern = regexp.MustCompile(`^[a-zA-Z0-9_-]{1,64}$`)
// ValidateRecordID validates an identifier before it is used as a jsondb
// record key (i.e. a filename component under the db directory), to
// prevent path traversal / arbitrary file writes via a crafted ID.
func ValidateRecordID(id string) bool {
return recordIDPattern.MatchString(id)
}
// interfaceNamePattern enforces Linux's IFNAMSIZ limit (16 bytes including
// the trailing NUL, so 15 usable characters) and a safe character set for
// any string that may later be passed as a network interface name to
// external tools (wg-quick, systemctl unit names, etc.).
var interfaceNamePattern = regexp.MustCompile(`^[a-zA-Z0-9_-]{1,15}$`)
// ValidateInterfaceName validates a WireGuard interface name before it is
// stored or ever used as an argument to exec'd tools such as wg-quick or
// systemctl, to prevent command/argument injection and invalid interfaces.
func ValidateInterfaceName(name string) bool {
return interfaceNamePattern.MatchString(name)
}
// ValidateIPAddress to validate the IPv4 and IPv6 address // ValidateIPAddress to validate the IPv4 and IPv6 address
func ValidateIPAddress(ip string) bool { func ValidateIPAddress(ip string) bool {
if net.ParseIP(ip) == nil { if net.ParseIP(ip) == nil {