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.PostUp = util.LookupEnvOrString(util.ServerPostUpScriptEnvVar, "") 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/.json", "server_settings/.json", // "server_hashes/.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/.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 != "" { server, _ := o.GetServerByID(util.DefaultServerID) 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 != "" { server, _ := o.GetServerByID(util.DefaultServerID) 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 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 } // 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 }