diff --git a/handler/routes.go b/handler/routes.go index 2f5221b..63711cc 100644 --- a/handler/routes.go +++ b/handler/routes.go @@ -1151,7 +1151,12 @@ func ServerServiceStatus(db store.IStore) echo.HandlerFunc { // ServerServiceStart brings up a server's WireGuard interface via // `systemctl start wg-quick@.service`. Admin-only. -func ServerServiceStart(db store.IStore) echo.HandlerFunc { +// +// wg-quick refuses to start when /etc/wireguard/.conf doesn't exist +// yet - true for a freshly created or imported server that nobody has +// applied config for - so the config is (re)written from current DB state +// first. +func ServerServiceStart(db store.IStore, tmplDir fs.FS) echo.HandlerFunc { return func(c echo.Context) error { serverID := c.Param("id") server, err := db.GetServerByID(serverID) @@ -1159,6 +1164,13 @@ func ServerServiceStart(db store.IStore) echo.HandlerFunc { return c.JSON(http.StatusNotFound, jsonHTTPResponse{false, "Server not found"}) } + if err := writeServerConfigToDisk(db, tmplDir, serverID); err != nil { + log.Errorf("Failed to write config before starting server %s: %v", serverID, err) + return c.JSON(http.StatusInternalServerError, jsonHTTPResponse{ + false, fmt.Sprintf("Cannot write server config: %v", err), + }) + } + if err := wireguard.Start(c.Request().Context(), server.Interface.Name); err != nil { log.Errorf("Failed to start service for server %s: %v", serverID, err) return c.JSON(http.StatusInternalServerError, jsonHTTPResponse{false, err.Error()}) @@ -1981,62 +1993,65 @@ func SuggestIPAllocation(db store.IStore) echo.HandlerFunc { } } +// writeServerConfigToDisk regenerates a server's /etc/wireguard/.conf +// from current DB state (server, its own clients, users, effective +// settings). Shared by ApplyServerConfig and ServerServiceStart, since +// wg-quick refuses to start when the file is missing (e.g. right after a +// server is created or imported and no one has hit "Apply Config" yet). +func writeServerConfigToDisk(db store.IStore, tmplDir fs.FS, serverID string) error { + server, err := db.GetServerByID(serverID) + if err != nil { + return fmt.Errorf("cannot get server config: %w", err) + } + + allClients, err := db.GetClients(false) + if err != nil { + return fmt.Errorf("cannot get client config: %w", err) + } + // only include this server's own clients as peers - other servers' + // clients must never leak into this config file + clients := make([]model.ClientData, 0, len(allClients)) + for _, cd := range allClients { + clientServerID := cd.Client.ServerID + if clientServerID == "" { + clientServerID = util.DefaultServerID + } + if clientServerID == serverID { + clients = append(clients, cd) + } + } + + users, err := db.GetUsers() + if err != nil { + return fmt.Errorf("cannot get users config: %w", err) + } + + settings, err := buildEffectiveSettings(db, serverID) + if err != nil { + return fmt.Errorf("cannot get global settings: %w", err) + } + if serverID != util.DefaultServerID { + if serverSettings, sErr := db.GetServerSettings(serverID); sErr == nil && serverSettings.ConfigFilePath != "" { + settings.ConfigFilePath = serverSettings.ConfigFilePath + } + } + + return util.WriteWireGuardServerConfig(tmplDir, server, clients, users, settings) +} + // ApplyServerConfig handler to write config file and restart Wireguard server func ApplyServerConfig(db store.IStore, tmplDir fs.FS) echo.HandlerFunc { return func(c echo.Context) error { serverID := resolveServerID(c) - server, err := db.GetServerByID(serverID) - if err != nil { - log.Error("Cannot get server config: ", err) - return c.JSON(http.StatusInternalServerError, jsonHTTPResponse{false, "Cannot get server config"}) - } - - allClients, err := db.GetClients(false) - if err != nil { - log.Error("Cannot get client config: ", err) - return c.JSON(http.StatusInternalServerError, jsonHTTPResponse{false, "Cannot get client config"}) - } - // only include this server's own clients as peers - other servers' - // clients must never leak into this config file - clients := make([]model.ClientData, 0, len(allClients)) - for _, cd := range allClients { - clientServerID := cd.Client.ServerID - if clientServerID == "" { - clientServerID = util.DefaultServerID - } - if clientServerID == serverID { - clients = append(clients, cd) - } - } - - users, err := db.GetUsers() - if err != nil { - log.Error("Cannot get users config: ", err) - return c.JSON(http.StatusInternalServerError, jsonHTTPResponse{false, "Cannot get users config"}) - } - - settings, err := buildEffectiveSettings(db, serverID) - if err != nil { - log.Error("Cannot get global settings: ", err) - return c.JSON(http.StatusInternalServerError, jsonHTTPResponse{false, "Cannot get global settings"}) - } - if serverID != util.DefaultServerID { - if serverSettings, sErr := db.GetServerSettings(serverID); sErr == nil && serverSettings.ConfigFilePath != "" { - settings.ConfigFilePath = serverSettings.ConfigFilePath - } - } - - // Write config file - err = util.WriteWireGuardServerConfig(tmplDir, server, clients, users, settings) - if err != nil { + if err := writeServerConfigToDisk(db, tmplDir, serverID); err != nil { log.Error("Cannot apply server config: ", err) return c.JSON(http.StatusInternalServerError, jsonHTTPResponse{ false, fmt.Sprintf("Cannot apply server config: %v", err), }) } - err = util.UpdateHashes(db) + err := util.UpdateHashes(db) if err != nil { log.Error("Cannot update hashes: ", err) return c.JSON(http.StatusInternalServerError, jsonHTTPResponse{ diff --git a/main.go b/main.go index 631669c..0458ecd 100644 --- a/main.go +++ b/main.go @@ -297,7 +297,7 @@ func main() { app.POST(util.BasePath+"/servers/:id/firewall/rules/:ruleId/delete", handler.DeleteFirewallRuleHandler(db), handler.ValidSession, handler.ContentTypeJson, handler.NeedsAdmin) app.POST(util.BasePath+"/servers/:id/firewall/apply", handler.ApplyFirewallHandler(db), handler.ValidSession, handler.ContentTypeJson, handler.NeedsAdmin) app.GET(util.BasePath+"/servers/:id/service-status", handler.ServerServiceStatus(db), handler.ValidSession, handler.RequireServerAccess(db)) - app.POST(util.BasePath+"/servers/:id/service/start", handler.ServerServiceStart(db), handler.ValidSession, handler.ContentTypeJson, handler.NeedsAdmin) + app.POST(util.BasePath+"/servers/:id/service/start", handler.ServerServiceStart(db, tmplDir), handler.ValidSession, handler.ContentTypeJson, handler.NeedsAdmin) app.POST(util.BasePath+"/servers/:id/service/stop", handler.ServerServiceStop(db), handler.ValidSession, handler.ContentTypeJson, handler.NeedsAdmin) app.POST(util.BasePath+"/servers/:id/service/restart", handler.ServerServiceRestart(db), handler.ValidSession, handler.ContentTypeJson, handler.NeedsAdmin) app.POST(util.BasePath+"/backup/download", handler.DownloadBackup(db), handler.ValidSession, handler.ContentTypeJson, handler.NeedsAdmin)