From 1d080904b03523a4f3512aa1dd1ddb94f64c6912 Mon Sep 17 00:00:00 2001 From: sysops Date: Sun, 12 Jul 2026 17:31:08 +0200 Subject: [PATCH] Add live firewall rule management per server (nftables) New model.FirewallRule + jsondb CRUD (GetFirewallRules/CreateFirewallRule/ UpdateFirewallRule/DeleteFirewallRule), scoped per server. firewall package now generates a full ruleset (baseline + enabled custom rules) and can apply it live via `nft -f` (firewall.Apply), scoped to a per-server nftables table (wireguard_ui_) so applying one server never touches another server's rules or any pre-existing firewall state. New endpoints: GET/POST /servers/:id/firewall/rules, POST .../rules/:ruleId, POST .../rules/:ruleId/delete, POST .../apply (live, admin-only). UI in the All Servers page: rule table with add/delete, ruleset preview, and an "Apply now (live)" button with an explicit confirm() warning before it touches the running firewall. Co-Authored-By: Claude Sonnet 5 --- firewall/apply.go | 44 ++++++++++ firewall/nftables.go | 78 +++++++++++++----- handler/routes.go | 157 +++++++++++++++++++++++++++++++++-- main.go | 5 ++ model/firewall.go | 21 +++++ store/jsondb/jsondb.go | 59 ++++++++++++++ store/store.go | 4 + templates/servers.html | 180 +++++++++++++++++++++++++++++++++++++---- 8 files changed, 507 insertions(+), 41 deletions(-) create mode 100644 firewall/apply.go create mode 100644 model/firewall.go diff --git a/firewall/apply.go b/firewall/apply.go new file mode 100644 index 0000000..ec3cb60 --- /dev/null +++ b/firewall/apply.go @@ -0,0 +1,44 @@ +package firewall + +import ( + "context" + "fmt" + "os" + "os/exec" + "time" +) + +// Apply writes ruleset to a temp file and loads it with `nft -f`, after +// first deleting the server's own table (ignoring the error - the table +// may not exist yet on first apply). Only ever touches the single table +// named by TableName(serverID), never any other nftables state. +// Returns combined nft output for display, and an error if the load failed. +func Apply(serverID, ruleset string) (string, error) { + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + + // best-effort: drop any previous version of this server's table so + // reapplying is idempotent. Error ignored - table may not exist yet. + _ = exec.CommandContext(ctx, "nft", "delete", "table", "inet", TableName(serverID)).Run() + + tmpFile, err := os.CreateTemp("", "wg-ui-multi-fw-*.nft") + if err != nil { + return "", fmt.Errorf("cannot create temp ruleset file: %w", err) + } + defer os.Remove(tmpFile.Name()) + + if _, err := tmpFile.WriteString(ruleset); err != nil { + tmpFile.Close() + return "", fmt.Errorf("cannot write temp ruleset file: %w", err) + } + if err := tmpFile.Close(); err != nil { + return "", fmt.Errorf("cannot close temp ruleset file: %w", err) + } + + cmd := exec.CommandContext(ctx, "nft", "-f", tmpFile.Name()) + out, err := cmd.CombinedOutput() + if err != nil { + return string(out), fmt.Errorf("nft -f failed: %w", err) + } + return string(out), nil +} diff --git a/firewall/nftables.go b/firewall/nftables.go index 9ed2bcb..7a9baaa 100644 --- a/firewall/nftables.go +++ b/firewall/nftables.go @@ -1,7 +1,7 @@ -// Package firewall generates an nftables ruleset preview for a WireGuard -// server. It never touches the live firewall - the output is text only, -// meant to be reviewed and applied manually (nft -f ) or copied into -// an existing ruleset. +// Package firewall builds and (optionally) applies an nftables ruleset for +// a WireGuard server. Every server gets its own table +// (inet wireguard_ui_) so applying/removing one server's rules +// never touches any other table on the system. package firewall import ( @@ -11,10 +11,38 @@ import ( "github.com/ngoduykhanh/wireguard-ui/model" ) -// GeneratePreview renders an nftables ruleset snippet for the given server. -// If settings.LanInterface is empty, only WireGuard-interface-local traffic -// rules are emitted; forwarding to a LAN interface is added when set. -func GeneratePreview(server model.Server, settings model.ServerSetting) string { +// TableName returns the dedicated nftables table name for a server. +func TableName(serverID string) string { + return "wireguard_ui_" + serverID +} + +// buildRuleLine renders one custom rule as an nftables statement. +func buildRuleLine(rule model.FirewallRule) string { + var parts []string + if rule.Source != "" { + parts = append(parts, fmt.Sprintf("ip saddr %s", rule.Source)) + } + switch { + case rule.Protocol != "" && rule.Port != "": + parts = append(parts, fmt.Sprintf("%s dport %s", rule.Protocol, rule.Port)) + case rule.Protocol != "": + parts = append(parts, fmt.Sprintf("meta l4proto %s", rule.Protocol)) + case rule.Port != "": + parts = append(parts, fmt.Sprintf("th dport %s", rule.Port)) + } + parts = append(parts, rule.Action) + comment := rule.Comment + if comment == "" { + comment = "wg-ui-multi custom rule" + } + parts = append(parts, fmt.Sprintf("comment %q", comment)) + return " " + strings.Join(parts, " ") +} + +// GenerateRuleset renders the full nftables ruleset for a server: the +// baseline (listen-port accept, WireGuard-interface forwarding, optional +// LAN forwarding) plus every enabled custom rule, grouped by chain. +func GenerateRuleset(server model.Server, settings model.ServerSetting, rules []model.FirewallRule) string { ifaceName := "wgX" listenPort := 0 if server.Interface != nil { @@ -24,16 +52,30 @@ func GeneratePreview(server model.Server, settings model.ServerSetting) string { listenPort = server.Interface.ListenPort } - var b strings.Builder - fmt.Fprintf(&b, "# nftables ruleset preview for server %q (%s)\n", server.Name, server.ID) - fmt.Fprintf(&b, "# Generated by wireguard-ui-multi - review before applying, e.g.:\n") - fmt.Fprintf(&b, "# nft -f this-file.nft\n") - fmt.Fprintf(&b, "# Not applied automatically.\n\n") + var inputExtra, forwardExtra []string + for _, rule := range rules { + if !rule.Enabled { + continue + } + line := buildRuleLine(rule) + if rule.Direction == "forward" { + forwardExtra = append(forwardExtra, line) + } else { + inputExtra = append(inputExtra, line) + } + } - fmt.Fprintf(&b, "table inet wireguard_ui_%s {\n", server.ID) + var b strings.Builder + fmt.Fprintf(&b, "# nftables ruleset for server %q (%s)\n", server.Name, server.ID) + fmt.Fprintf(&b, "# Generated by wireguard-ui-multi.\n\n") + + fmt.Fprintf(&b, "table inet %s {\n", TableName(server.ID)) fmt.Fprintf(&b, " chain input {\n") fmt.Fprintf(&b, " type filter hook input priority 0; policy accept;\n") fmt.Fprintf(&b, " udp dport %d accept comment \"wg-ui-multi: %s\"\n", listenPort, server.ID) + for _, line := range inputExtra { + fmt.Fprintf(&b, "%s\n", line) + } fmt.Fprintf(&b, " }\n\n") fmt.Fprintf(&b, " chain forward {\n") @@ -44,13 +86,11 @@ func GeneratePreview(server model.Server, settings model.ServerSetting) string { fmt.Fprintf(&b, " iifname \"%s\" oifname \"%s\" accept comment \"wg-ui-multi: %s -> lan\"\n", ifaceName, settings.LanInterface, server.ID) fmt.Fprintf(&b, " iifname \"%s\" oifname \"%s\" accept comment \"wg-ui-multi: lan -> %s\"\n", settings.LanInterface, ifaceName, server.ID) } + for _, line := range forwardExtra { + fmt.Fprintf(&b, "%s\n", line) + } fmt.Fprintf(&b, " }\n") fmt.Fprintf(&b, "}\n") - if settings.LanInterface == "" { - fmt.Fprintf(&b, "\n# No LAN interface configured for this server - peers can only reach\n") - fmt.Fprintf(&b, "# each other, not your LAN. Set one in Server Settings to add forwarding.\n") - } - return b.String() } diff --git a/handler/routes.go b/handler/routes.go index 12ed0f0..37b230c 100644 --- a/handler/routes.go +++ b/handler/routes.go @@ -6,6 +6,7 @@ import ( "encoding/json" "fmt" "io/fs" + "net" "net/http" "os" "regexp" @@ -723,20 +724,162 @@ func RemoveServer(db store.IStore) echo.HandlerFunc { } } -// GetServerFirewallPreview returns a generated nftables ruleset preview for -// a server as plain text. Never applied automatically - review-only. +// loadFirewallRuleset fetches everything needed to render a server's full +// nftables ruleset (baseline + custom rules). +func loadFirewallRuleset(db store.IStore, serverID string) (string, error) { + server, err := db.GetServerByID(serverID) + if err != nil { + return "", fmt.Errorf("server not found") + } + settings, err := db.GetServerSettings(serverID) + if err != nil { + return "", fmt.Errorf("server settings not found") + } + rules, err := db.GetFirewallRules(serverID) + if err != nil { + return "", fmt.Errorf("cannot load firewall rules: %v", err) + } + return firewall.GenerateRuleset(server, settings, rules), nil +} + +// GetServerFirewallPreview returns the generated nftables ruleset (baseline +// + custom rules) for a server as plain text. Preview only. func GetServerFirewallPreview(db store.IStore) echo.HandlerFunc { return func(c echo.Context) error { serverID := c.Param("id") - server, err := db.GetServerByID(serverID) + ruleset, err := loadFirewallRuleset(db, serverID) if err != nil { + return c.JSON(http.StatusNotFound, jsonHTTPResponse{false, err.Error()}) + } + return c.String(http.StatusOK, ruleset) + } +} + +var validFirewallDirections = map[string]bool{"input": true, "forward": true} +var validFirewallProtocols = map[string]bool{"": true, "tcp": true, "udp": true} +var validFirewallActions = map[string]bool{"accept": true, "drop": true, "reject": true} +var firewallPortRegexp = regexp.MustCompile(`^[0-9]{1,5}(-[0-9]{1,5})?$`) + +func validateFirewallRule(rule model.FirewallRule) error { + if !validFirewallDirections[rule.Direction] { + return fmt.Errorf("direction must be 'input' or 'forward'") + } + if !validFirewallProtocols[rule.Protocol] { + return fmt.Errorf("protocol must be 'tcp', 'udp', or empty") + } + if !validFirewallActions[rule.Action] { + return fmt.Errorf("action must be 'accept', 'drop', or 'reject'") + } + if rule.Port != "" && !firewallPortRegexp.MatchString(rule.Port) { + return fmt.Errorf("port must be a number or range like 8000-9000") + } + if rule.Source != "" { + if !util.ValidateServerAddresses([]string{rule.Source}) && net.ParseIP(rule.Source) == nil { + return fmt.Errorf("source must be a valid IP or CIDR") + } + } + if len(rule.Comment) > 200 { + return fmt.Errorf("comment too long") + } + return nil +} + +// GetFirewallRules lists a server's custom firewall rules. +func GetFirewallRules(db store.IStore) echo.HandlerFunc { + return func(c echo.Context) error { + serverID := c.Param("id") + rules, err := db.GetFirewallRules(serverID) + if err != nil { + return c.JSON(http.StatusBadRequest, jsonHTTPResponse{false, err.Error()}) + } + return c.JSON(http.StatusOK, rules) + } +} + +// CreateFirewallRuleHandler creates a new custom firewall rule for a server. +func CreateFirewallRuleHandler(db store.IStore) echo.HandlerFunc { + return func(c echo.Context) error { + serverID := c.Param("id") + if _, err := db.GetServerByID(serverID); err != nil { return c.JSON(http.StatusNotFound, jsonHTTPResponse{false, "Server not found"}) } - settings, err := db.GetServerSettings(serverID) - if err != nil { - return c.JSON(http.StatusNotFound, jsonHTTPResponse{false, "Server settings not found"}) + var rule model.FirewallRule + if err := c.Bind(&rule); err != nil { + return c.JSON(http.StatusBadRequest, jsonHTTPResponse{false, "Bad post data"}) } - return c.String(http.StatusOK, firewall.GeneratePreview(server, settings)) + rule.ServerID = serverID + if err := validateFirewallRule(rule); err != nil { + return c.JSON(http.StatusBadRequest, jsonHTTPResponse{false, err.Error()}) + } + rule.ID = xid.New().String() + rule.CreatedAt = time.Now().UTC() + rule.UpdatedAt = rule.CreatedAt + if err := db.CreateFirewallRule(rule); err != nil { + return c.JSON(http.StatusInternalServerError, jsonHTTPResponse{false, fmt.Sprintf("Cannot create rule: %v", err)}) + } + return c.JSON(http.StatusOK, rule) + } +} + +// UpdateFirewallRuleHandler updates an existing custom firewall rule. +func UpdateFirewallRuleHandler(db store.IStore) echo.HandlerFunc { + return func(c echo.Context) error { + serverID := c.Param("id") + ruleID := c.Param("ruleId") + var rule model.FirewallRule + if err := c.Bind(&rule); err != nil { + return c.JSON(http.StatusBadRequest, jsonHTTPResponse{false, "Bad post data"}) + } + rule.ID = ruleID + rule.ServerID = serverID + if err := validateFirewallRule(rule); err != nil { + return c.JSON(http.StatusBadRequest, jsonHTTPResponse{false, err.Error()}) + } + rule.UpdatedAt = time.Now().UTC() + if err := db.UpdateFirewallRule(rule); err != nil { + return c.JSON(http.StatusBadRequest, jsonHTTPResponse{false, err.Error()}) + } + return c.JSON(http.StatusOK, rule) + } +} + +// DeleteFirewallRuleHandler removes a custom firewall rule. +func DeleteFirewallRuleHandler(db store.IStore) echo.HandlerFunc { + return func(c echo.Context) error { + serverID := c.Param("id") + ruleID := c.Param("ruleId") + if err := db.DeleteFirewallRule(serverID, ruleID); err != nil { + return c.JSON(http.StatusBadRequest, jsonHTTPResponse{false, err.Error()}) + } + return c.JSON(http.StatusOK, jsonHTTPResponse{true, "Rule deleted successfully"}) + } +} + +// ApplyFirewallHandler generates the full ruleset for a server and loads it +// live via `nft -f`. This DOES modify the running firewall, scoped to this +// server's own nftables table only (see firewall.TableName). +func ApplyFirewallHandler(db store.IStore) echo.HandlerFunc { + return func(c echo.Context) error { + serverID := c.Param("id") + ruleset, err := loadFirewallRuleset(db, serverID) + if err != nil { + return c.JSON(http.StatusNotFound, jsonHTTPResponse{false, err.Error()}) + } + output, err := firewall.Apply(serverID, ruleset) + if err != nil { + log.Errorf("Failed to apply firewall rules for server %s: %v\n%s", serverID, err, output) + return c.JSON(http.StatusInternalServerError, map[string]interface{}{ + "success": false, + "message": err.Error(), + "output": output, + }) + } + log.Infof("Applied firewall rules for server %s", serverID) + return c.JSON(http.StatusOK, map[string]interface{}{ + "success": true, + "message": "Firewall rules applied successfully", + "output": output, + }) } } diff --git a/main.go b/main.go index 204f50e..059af48 100644 --- a/main.go +++ b/main.go @@ -266,6 +266,11 @@ func main() { app.POST(util.BasePath+"/servers/:id/keypair", handler.UpdateServerKeyPairHandler(db), handler.ValidSession, handler.ContentTypeJson, handler.NeedsAdmin) app.POST(util.BasePath+"/servers/:id/delete", handler.RemoveServer(db), handler.ValidSession, handler.ContentTypeJson, handler.NeedsAdmin) app.GET(util.BasePath+"/servers/:id/firewall-preview", handler.GetServerFirewallPreview(db), handler.ValidSession, handler.RequireServerAccess(db)) + app.GET(util.BasePath+"/servers/:id/firewall/rules", handler.GetFirewallRules(db), handler.ValidSession, handler.RequireServerAccess(db)) + app.POST(util.BasePath+"/servers/:id/firewall/rules", handler.CreateFirewallRuleHandler(db), handler.ValidSession, handler.ContentTypeJson, handler.NeedsAdmin) + app.POST(util.BasePath+"/servers/:id/firewall/rules/:ruleId", handler.UpdateFirewallRuleHandler(db), handler.ValidSession, handler.ContentTypeJson, handler.NeedsAdmin) + 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.POST(util.BasePath+"/backup/download", handler.DownloadBackup(db), handler.ValidSession, handler.ContentTypeJson, handler.NeedsAdmin) app.GET(util.BasePath+"/api/clients", handler.GetClients(db), handler.ValidSession) app.GET(util.BasePath+"/api/client/:id", handler.GetClient(db), handler.ValidSession) diff --git a/model/firewall.go b/model/firewall.go new file mode 100644 index 0000000..da3cade --- /dev/null +++ b/model/firewall.go @@ -0,0 +1,21 @@ +package model + +import "time" + +// FirewallRule is a single user-defined nftables rule scoped to one server. +// Rules are combined with the server's baseline (listen-port accept + +// WireGuard-interface forwarding) to build the full ruleset that gets +// applied via `nft -f`. +type FirewallRule struct { + ID string `json:"id"` + ServerID string `json:"server_id"` + Direction string `json:"direction"` // "input" or "forward" + Protocol string `json:"protocol"` // "tcp", "udp", or "" (any) + Port string `json:"port"` // e.g. "8080" or "8000-9000", "" = any + Source string `json:"source"` // optional CIDR, "" = any + Action string `json:"action"` // "accept", "drop", or "reject" + Comment string `json:"comment"` + Enabled bool `json:"enabled"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` +} diff --git a/store/jsondb/jsondb.go b/store/jsondb/jsondb.go index e2bb9cc..d27ee1b 100644 --- a/store/jsondb/jsondb.go +++ b/store/jsondb/jsondb.go @@ -677,3 +677,62 @@ func (o *JsonDB) SaveServerHashes(serverID string, hashes model.ClientServerHash } return output } + +// GetFirewallRules func to query all firewall rules belonging to a server +func (o *JsonDB) GetFirewallRules(serverID string) ([]model.FirewallRule, error) { + if err := validateServerID(serverID); err != nil { + return nil, err + } + rules := make([]model.FirewallRule, 0) + records, err := o.conn.ReadAll("firewall_rules") + if err != nil { + if err == scribble.ErrMissingCollection { + return rules, nil + } + return nil, err + } + for _, rec := range records { + var rule model.FirewallRule + if err := json.Unmarshal(rec, &rule); err != nil { + return nil, fmt.Errorf("cannot decode firewall rule json structure: %v", err) + } + if rule.ServerID == serverID { + rules = append(rules, rule) + } + } + return rules, nil +} + +// CreateFirewallRule func to add a new firewall rule for a server +func (o *JsonDB) CreateFirewallRule(rule model.FirewallRule) error { + if err := validateServerID(rule.ServerID); err != nil { + return err + } + return o.conn.Write("firewall_rules", rule.ID, rule) +} + +// UpdateFirewallRule func to update an existing firewall rule, refusing to +// move it to a different server than it currently belongs to +func (o *JsonDB) UpdateFirewallRule(rule model.FirewallRule) error { + var existing model.FirewallRule + if err := o.conn.Read("firewall_rules", rule.ID, &existing); err != nil { + return fmt.Errorf("firewall rule not found: %v", err) + } + if existing.ServerID != rule.ServerID { + return fmt.Errorf("cannot change server_id of an existing firewall rule") + } + return o.conn.Write("firewall_rules", rule.ID, rule) +} + +// DeleteFirewallRule func to remove a firewall rule, verifying it belongs to +// the given server first +func (o *JsonDB) DeleteFirewallRule(serverID, ruleID string) error { + var existing model.FirewallRule + if err := o.conn.Read("firewall_rules", ruleID, &existing); err != nil { + return fmt.Errorf("firewall rule not found: %v", err) + } + if existing.ServerID != serverID { + return fmt.Errorf("firewall rule does not belong to server %s", serverID) + } + return o.conn.Delete("firewall_rules", ruleID) +} diff --git a/store/store.go b/store/store.go index f39d1ce..60fa4dd 100644 --- a/store/store.go +++ b/store/store.go @@ -34,4 +34,8 @@ type IStore interface { SaveServerHashes(serverID string, hashes model.ClientServerHashes) error UpdateServerInterface(serverID string, serverInterface model.ServerInterface) error UpdateServerKeyPair(serverID string, serverKeyPair model.ServerKeypair) error + GetFirewallRules(serverID string) ([]model.FirewallRule, error) + CreateFirewallRule(rule model.FirewallRule) error + UpdateFirewallRule(rule model.FirewallRule) error + DeleteFirewallRule(serverID, ruleID string) error } diff --git a/templates/servers.html b/templates/servers.html index dbe0dae..13cf2eb 100644 --- a/templates/servers.html +++ b/templates/servers.html @@ -176,21 +176,55 @@ All Servers -