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:
@@ -535,3 +535,122 @@ func (o *JsonDB) SaveHashes(hashes model.ClientServerHashes) error {
|
||||
}
|
||||
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
|
||||
}
|
||||
|
||||
@@ -27,4 +27,12 @@ type IStore interface {
|
||||
GetPath() string
|
||||
SaveHashes(hashes 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
|
||||
}
|
||||
|
||||
@@ -15,6 +15,7 @@ import (
|
||||
"os"
|
||||
"path"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
"text/template"
|
||||
@@ -165,6 +166,31 @@ func ValidateServerAddresses(cidrs []string) bool {
|
||||
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
|
||||
func ValidateIPAddress(ip string) bool {
|
||||
if net.ParseIP(ip) == nil {
|
||||
|
||||
Reference in New Issue
Block a user