From 7eca2d01b81a8ccea854787f50e999753b7ce9ce Mon Sep 17 00:00:00 2001 From: sysops Date: Sat, 11 Jul 2026 22:39:53 +0200 Subject: [PATCH] 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). --- store/jsondb/jsondb.go | 119 +++++++++++++++++++++++++++++++++++++++++ store/store.go | 8 +++ util/util.go | 26 +++++++++ 3 files changed, 153 insertions(+) diff --git a/store/jsondb/jsondb.go b/store/jsondb/jsondb.go index 6468d8a..c2392ce 100644 --- a/store/jsondb/jsondb.go +++ b/store/jsondb/jsondb.go @@ -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 +} diff --git a/store/store.go b/store/store.go index ef6d723..69e346c 100644 --- a/store/store.go +++ b/store/store.go @@ -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 } diff --git a/util/util.go b/util/util.go index ec700ff..a0a0648 100644 --- a/util/util.go +++ b/util/util.go @@ -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 {