mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-07 02:06:38 +08:00
feat: add nftables forwarding mode
This commit is contained in:
@@ -240,6 +240,19 @@ func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
nftMode, entryNodeIDs, err := h.tunnelUsesNftables(forward.TunnelID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if nftMode {
|
||||
if err := h.validateNftablesForwardRequest(tunnel, forward.RemoteAddr, entryNodeIDs); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(entryNodeIDs) == 0 {
|
||||
return nil, errors.New("nftables 转发缺少入口节点")
|
||||
}
|
||||
return nil, h.syncNftablesNode(entryNodeIDs[0])
|
||||
}
|
||||
ports, err := h.listForwardPorts(forward.ID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
|
||||
@@ -23,6 +23,7 @@ import (
|
||||
"go-backend/internal/license"
|
||||
"go-backend/internal/metrics"
|
||||
"go-backend/internal/monitoring"
|
||||
runtimenft "go-backend/internal/runtime/nftables"
|
||||
"go-backend/internal/security"
|
||||
"go-backend/internal/store/repo"
|
||||
"go-backend/internal/ws"
|
||||
@@ -31,11 +32,12 @@ import (
|
||||
)
|
||||
|
||||
type Handler struct {
|
||||
repo *repo.Repository
|
||||
jwtSecret string
|
||||
wsServer *ws.Server
|
||||
metrics *metrics.IngestionService
|
||||
healthCheck *health.Checker
|
||||
repo *repo.Repository
|
||||
jwtSecret string
|
||||
wsServer *ws.Server
|
||||
metrics *metrics.IngestionService
|
||||
healthCheck *health.Checker
|
||||
nftablesManager nftablesRuntimeManager
|
||||
|
||||
captchaMu sync.Mutex
|
||||
captchaTokens map[string]int64
|
||||
@@ -108,6 +110,7 @@ func New(repo *repo.Repository, jwtSecret string) *Handler {
|
||||
wsServer: ws.NewServer(repo, jwtSecret),
|
||||
metrics: metrics.NewIngestionService(repo),
|
||||
healthCheck: nil,
|
||||
nftablesManager: runtimenft.NewManager(nil),
|
||||
captchaTokens: make(map[string]int64),
|
||||
pendingUpgradeRedeploy: make(map[int64]struct{}),
|
||||
nodeOnlineRedeployAt: make(map[int64]time.Time),
|
||||
@@ -190,6 +193,9 @@ func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/v1/node/batch-upgrade", h.nodeBatchUpgrade)
|
||||
mux.HandleFunc("/api/v1/node/rollback", h.nodeRollback)
|
||||
mux.HandleFunc("/api/v1/node/releases", h.listReleases)
|
||||
mux.HandleFunc("/api/v1/node/nftables/test", h.nodeNftablesTest)
|
||||
mux.HandleFunc("/api/v1/node/nftables/reconcile", h.nodeNftablesReconcile)
|
||||
mux.HandleFunc("/api/v1/node/nftables/clear", h.nodeNftablesClear)
|
||||
mux.HandleFunc("/api/v1/tunnel/list", h.tunnelList)
|
||||
mux.HandleFunc("/api/v1/tunnel/create", h.tunnelCreate)
|
||||
mux.HandleFunc("/api/v1/tunnel/get", h.tunnelGet)
|
||||
|
||||
@@ -341,6 +341,11 @@ func (h *Handler) nodeCreate(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
inx := h.repo.NextIndex("node")
|
||||
forwardMode := defaultNodeForwardMode(asString(req["forwardMode"]))
|
||||
status := 0
|
||||
if forwardMode == "nftables" {
|
||||
status = 1
|
||||
}
|
||||
if err := h.repo.CreateNode(
|
||||
name,
|
||||
randomToken(16),
|
||||
@@ -357,7 +362,7 @@ func (h *Handler) nodeCreate(w http.ResponseWriter, r *http.Request) {
|
||||
asInt(req["tls"], 0),
|
||||
asInt(req["socks"], 0),
|
||||
now,
|
||||
0,
|
||||
status,
|
||||
defaultString(asString(req["tcpListenAddr"]), "[::]"),
|
||||
defaultString(asString(req["udpListenAddr"]), "[::]"),
|
||||
inx,
|
||||
@@ -366,10 +371,20 @@ func (h *Handler) nodeCreate(w http.ResponseWriter, r *http.Request) {
|
||||
nullableText(asString(req["remoteToken"])),
|
||||
nullableText(asString(req["remoteConfig"])),
|
||||
nullableText(asString(req["extraIPs"])),
|
||||
forwardMode,
|
||||
); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
nodeID, err := h.findCreatedNodeID(name, serverIP)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.persistNodeSSHConfig(nodeID, req, forwardMode, now, false); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
@@ -417,6 +432,7 @@ func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
forwardMode := defaultNodeForwardMode(strings.TrimSpace(asString(req["forwardMode"])))
|
||||
if err := h.repo.UpdateNode(id,
|
||||
asString(req["name"]),
|
||||
serverIP,
|
||||
@@ -428,6 +444,7 @@ func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
nullableText(strings.TrimSpace(asString(req["remark"]))),
|
||||
nullableUnixMilli(asInt64(req["expiryTime"], 0)),
|
||||
nullableText(normalizeNodeRenewalCycle(asString(req["renewalCycle"]))),
|
||||
forwardMode,
|
||||
newHTTP,
|
||||
newTLS,
|
||||
newSocks,
|
||||
@@ -438,9 +455,142 @@ func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.persistNodeSSHConfig(id, req, forwardMode, now, true); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if forwardMode == "nftables" && currentStatus != 1 {
|
||||
if err := h.repo.UpdateNodeStatus(id, 1); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
func (h *Handler) findCreatedNodeID(name, serverIP string) (int64, error) {
|
||||
if h == nil || h.repo == nil {
|
||||
return 0, errors.New("handler not initialized")
|
||||
}
|
||||
nodes, err := h.repo.ListNodes()
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
for i := len(nodes) - 1; i >= 0; i-- {
|
||||
item := nodes[i]
|
||||
if asString(item["name"]) != name {
|
||||
continue
|
||||
}
|
||||
if asString(item["serverIp"]) != serverIP {
|
||||
continue
|
||||
}
|
||||
if nodeID := asInt64(item["id"], 0); nodeID > 0 {
|
||||
return nodeID, nil
|
||||
}
|
||||
}
|
||||
return 0, errors.New("节点创建成功,但未能查询到节点记录")
|
||||
}
|
||||
|
||||
func (h *Handler) persistNodeSSHConfig(nodeID int64, req map[string]interface{}, forwardMode string, now int64, preserveSecrets bool) error {
|
||||
if h == nil || h.repo == nil {
|
||||
return errors.New("handler not initialized")
|
||||
}
|
||||
if nodeID <= 0 {
|
||||
return errors.New("节点ID不能为空")
|
||||
}
|
||||
if forwardMode != "nftables" {
|
||||
return h.repo.DeleteNodeSSHConfig(nodeID)
|
||||
}
|
||||
cfgMap := asMap(req["sshConfig"])
|
||||
host := strings.TrimSpace(asString(cfgMap["host"]))
|
||||
if host == "" {
|
||||
host = strings.TrimSpace(asString(req["serverIp"]))
|
||||
}
|
||||
port := asInt(cfgMap["port"], 22)
|
||||
username := strings.TrimSpace(asString(cfgMap["username"]))
|
||||
authType := strings.TrimSpace(asString(cfgMap["authType"]))
|
||||
password := asString(cfgMap["password"])
|
||||
privateKey := asString(cfgMap["privateKey"])
|
||||
passphrase := asString(cfgMap["passphrase"])
|
||||
sudoMode := strings.TrimSpace(asString(cfgMap["sudoMode"]))
|
||||
|
||||
if preserveSecrets {
|
||||
existing, err := h.repo.GetNodeSSHConfig(nodeID)
|
||||
if err != nil && !errors.Is(err, sql.ErrNoRows) {
|
||||
return err
|
||||
}
|
||||
if existing != nil {
|
||||
if host == "" {
|
||||
host = strings.TrimSpace(existing.Host)
|
||||
}
|
||||
if port <= 0 {
|
||||
port = existing.Port
|
||||
}
|
||||
if username == "" {
|
||||
username = strings.TrimSpace(existing.Username)
|
||||
}
|
||||
if authType == "" {
|
||||
authType = strings.TrimSpace(existing.AuthType)
|
||||
}
|
||||
if strings.TrimSpace(password) == "" && existing.Password.Valid {
|
||||
password = existing.Password.String
|
||||
}
|
||||
if strings.TrimSpace(privateKey) == "" && existing.PrivateKey.Valid {
|
||||
privateKey = existing.PrivateKey.String
|
||||
}
|
||||
if strings.TrimSpace(passphrase) == "" && existing.Passphrase.Valid {
|
||||
passphrase = existing.Passphrase.String
|
||||
}
|
||||
if sudoMode == "" {
|
||||
sudoMode = strings.TrimSpace(existing.SudoMode)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if host == "" || username == "" {
|
||||
return errors.New("nftables 节点 SSH 配置不完整")
|
||||
}
|
||||
if port <= 0 || port > 65535 {
|
||||
return errors.New("nftables 节点 SSH 端口无效")
|
||||
}
|
||||
|
||||
authType = strings.ToLower(authType)
|
||||
switch authType {
|
||||
case "password":
|
||||
if strings.TrimSpace(password) == "" {
|
||||
return errors.New("nftables 节点 SSH 密码不能为空")
|
||||
}
|
||||
privateKey = ""
|
||||
case "private_key", "":
|
||||
authType = "private_key"
|
||||
if strings.TrimSpace(privateKey) == "" {
|
||||
return errors.New("nftables 节点 SSH 私钥不能为空")
|
||||
}
|
||||
password = ""
|
||||
default:
|
||||
return errors.New("nftables 节点 SSH 认证方式无效")
|
||||
}
|
||||
|
||||
switch strings.ToLower(sudoMode) {
|
||||
case "", "none":
|
||||
sudoMode = "none"
|
||||
case "sudo", "sudo_su":
|
||||
default:
|
||||
return errors.New("nftables 节点 sudo 模式无效")
|
||||
}
|
||||
|
||||
return h.repo.UpsertNodeSSHConfig(nodeID, repo.NftSSHConfigInput{
|
||||
Host: host,
|
||||
Port: port,
|
||||
Username: username,
|
||||
AuthType: authType,
|
||||
Password: password,
|
||||
PrivateKey: privateKey,
|
||||
Passphrase: passphrase,
|
||||
SudoMode: sudoMode,
|
||||
}, now)
|
||||
}
|
||||
|
||||
func (h *Handler) nodeDelete(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
@@ -648,6 +798,16 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
|
||||
if strings.TrimSpace(inIP) == "" {
|
||||
inIP = buildTunnelInIP(runtimeState.InNodes, runtimeState.Nodes, ipPreference)
|
||||
}
|
||||
entryNodeIDs := make([]int64, 0, len(runtimeState.InNodes))
|
||||
for _, inNode := range runtimeState.InNodes {
|
||||
if inNode.NodeID > 0 {
|
||||
entryNodeIDs = append(entryNodeIDs, inNode.NodeID)
|
||||
}
|
||||
}
|
||||
if err := h.validateNftablesTunnelStateTx(tx, entryNodeIDs); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
if len(runtimeState.InNodes) > 0 {
|
||||
firstNodeID := runtimeState.InNodes[0].NodeID
|
||||
@@ -934,6 +1094,16 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
runtimeState.TunnelID = id
|
||||
runtimeState.IPPreference = ipPreference
|
||||
entryNodeIDs := make([]int64, 0, len(runtimeState.InNodes))
|
||||
for _, inNode := range runtimeState.InNodes {
|
||||
if inNode.NodeID > 0 {
|
||||
entryNodeIDs = append(entryNodeIDs, inNode.NodeID)
|
||||
}
|
||||
}
|
||||
if err := h.validateNftablesTunnelState(entryNodeIDs); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
inIp := buildTunnelInIP(runtimeState.InNodes, runtimeState.Nodes, ipPreference)
|
||||
|
||||
@@ -1700,6 +1870,24 @@ func (h *Handler) tunnelBatchRedeploy(w http.ResponseWriter, r *http.Request) {
|
||||
failures = appendBatchFailure(failures, tunnelID, tunnelName, tunnelErr)
|
||||
continue
|
||||
}
|
||||
if nftMode, entryNodeIDs, modeErr := h.tunnelUsesNftables(tunnelID); modeErr != nil {
|
||||
fail++
|
||||
failures = appendBatchFailure(failures, tunnelID, tunnelName, modeErr)
|
||||
continue
|
||||
} else if nftMode {
|
||||
if len(entryNodeIDs) == 0 {
|
||||
fail++
|
||||
failures = appendBatchFailureReason(failures, tunnelID, tunnelName, "nftables 转发缺少入口节点")
|
||||
continue
|
||||
}
|
||||
if reconcileErr := h.reconcileNftablesNodeByRequest(entryNodeIDs[0]); reconcileErr != nil {
|
||||
fail++
|
||||
failures = appendBatchFailure(failures, tunnelID, tunnelName, reconcileErr)
|
||||
continue
|
||||
}
|
||||
success++
|
||||
continue
|
||||
}
|
||||
if err := h.redeployTunnelAndForwards(tunnelID); err != nil {
|
||||
fail++
|
||||
failures = appendBatchFailure(failures, tunnelID, tunnelName, err)
|
||||
@@ -1909,6 +2097,17 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
|
||||
port = 10000
|
||||
}
|
||||
entryNodes, _ := h.tunnelEntryNodeIDs(tunnelID)
|
||||
isNftTunnel, _, err := h.tunnelUsesNftables(tunnelID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if isNftTunnel {
|
||||
if err := h.validateNftablesForwardRequest(tunnel, remoteAddr, entryNodes); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
inIp := strings.TrimSpace(asString(req["inIp"]))
|
||||
if inIp != "" && len(entryNodes) > 1 {
|
||||
response.WriteJSON(w, response.ErrDefault("多入口隧道的转发不支持自定义监听IP"))
|
||||
@@ -2016,6 +2215,17 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
if remoteAddr == "" {
|
||||
remoteAddr = forward.RemoteAddr
|
||||
}
|
||||
isNftTunnel, entryNodes, err := h.tunnelUsesNftables(tunnelID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if isNftTunnel {
|
||||
if err := h.validateNftablesForwardRequest(tunnel, remoteAddr, entryNodes); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
if actorRole != 0 && !h.allowLocalRemoteAddr() {
|
||||
if err := IsSafeRemoteAddr(remoteAddr); err != nil {
|
||||
response.WriteJSON(w, response.Err(403, err.Error()))
|
||||
@@ -2206,6 +2416,15 @@ func (h *Handler) forwardDelete(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
if nftMode, entryNodeIDs, modeErr := h.tunnelUsesNftables(forward.TunnelID); modeErr != nil {
|
||||
response.WriteJSON(w, response.Err(-2, modeErr.Error()))
|
||||
return
|
||||
} else if nftMode && len(entryNodeIDs) > 0 {
|
||||
if err := h.reconcileNftablesNodeByRequest(entryNodeIDs[0]); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
if err := h.deleteForwardByID(id); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -2218,7 +2437,7 @@ func (h *Handler) forwardForceDelete(w http.ResponseWriter, r *http.Request) {
|
||||
if id <= 0 {
|
||||
return
|
||||
}
|
||||
_, _, _, err := h.resolveForwardAccess(r, id)
|
||||
forward, _, _, err := h.resolveForwardAccess(r, id)
|
||||
if err != nil {
|
||||
if errors.Is(err, errForwardNotFound) {
|
||||
response.WriteJSON(w, response.ErrDefault("转发不存在"))
|
||||
@@ -2234,6 +2453,13 @@ func (h *Handler) forwardForceDelete(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
_ = h.repo.DeleteNftRuleBindingsByForward(id)
|
||||
if nftMode, entryNodeIDs, modeErr := h.tunnelUsesNftables(forward.TunnelID); modeErr == nil && nftMode && len(entryNodeIDs) > 0 {
|
||||
if err := h.reconcileNftablesNodeByRequest(entryNodeIDs[0]); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
@@ -2462,6 +2688,24 @@ func (h *Handler) forwardBatchRedeploy(w http.ResponseWriter, r *http.Request) {
|
||||
failures = appendBatchFailure(failures, id, "", accessErr)
|
||||
continue
|
||||
}
|
||||
if nftMode, entryNodeIDs, modeErr := h.tunnelUsesNftables(forward.TunnelID); modeErr != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, modeErr)
|
||||
continue
|
||||
} else if nftMode {
|
||||
if len(entryNodeIDs) == 0 {
|
||||
f++
|
||||
failures = appendBatchFailureReason(failures, id, forward.Name, "nftables 转发缺少入口节点")
|
||||
continue
|
||||
}
|
||||
if err := h.reconcileNftablesNodeByRequest(entryNodeIDs[0]); err != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
} else {
|
||||
s++
|
||||
}
|
||||
continue
|
||||
}
|
||||
if err := h.syncForwardServices(forward, "UpdateService", true); err != nil {
|
||||
f++
|
||||
failures = appendBatchFailure(failures, id, forward.Name, err)
|
||||
@@ -4705,6 +4949,13 @@ func asMapSlice(v interface{}) []map[string]interface{} {
|
||||
return out
|
||||
}
|
||||
|
||||
func asMap(v interface{}) map[string]interface{} {
|
||||
if m, ok := v.(map[string]interface{}); ok && m != nil {
|
||||
return m
|
||||
}
|
||||
return map[string]interface{}{}
|
||||
}
|
||||
|
||||
func asString(v interface{}) string {
|
||||
switch t := v.(type) {
|
||||
case nil:
|
||||
@@ -4843,6 +5094,15 @@ func normalizeNodeRenewalCycle(v string) string {
|
||||
}
|
||||
}
|
||||
|
||||
func defaultNodeForwardMode(mode string) string {
|
||||
switch strings.TrimSpace(strings.ToLower(mode)) {
|
||||
case "nftables":
|
||||
return "nftables"
|
||||
default:
|
||||
return "agent"
|
||||
}
|
||||
}
|
||||
|
||||
func nullableInt(v *int64) interface{} {
|
||||
if v == nil {
|
||||
return nil
|
||||
|
||||
@@ -0,0 +1,365 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/http/response"
|
||||
runtimenft "go-backend/internal/runtime/nftables"
|
||||
"go-backend/internal/store/model"
|
||||
"go-backend/internal/store/repo"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type nftablesRuntimeManager interface {
|
||||
Test(ctx context.Context, cfg runtimenft.SSHConfig) error
|
||||
Reconcile(ctx context.Context, cfg runtimenft.SSHConfig, plan runtimenft.NodePlan) (runtimenft.ApplyResult, error)
|
||||
Clear(ctx context.Context, cfg runtimenft.SSHConfig) error
|
||||
}
|
||||
|
||||
func isNftablesForwardMode(mode string) bool {
|
||||
return strings.EqualFold(strings.TrimSpace(mode), runtimenft.ModeNftables)
|
||||
}
|
||||
|
||||
func (h *Handler) nodeUsesNftables(nodeID int64) (bool, error) {
|
||||
return h.nodeUsesNftablesTx(nil, nodeID)
|
||||
}
|
||||
|
||||
func (h *Handler) nodeUsesNftablesTx(tx *gorm.DB, nodeID int64) (bool, error) {
|
||||
if h == nil || h.repo == nil {
|
||||
return false, errors.New("handler not initialized")
|
||||
}
|
||||
var (
|
||||
mode string
|
||||
err error
|
||||
)
|
||||
if tx != nil {
|
||||
mode, err = h.repo.GetNodeForwardModeTx(tx, nodeID)
|
||||
} else {
|
||||
mode, err = h.repo.GetNodeForwardMode(nodeID)
|
||||
}
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return isNftablesForwardMode(mode), nil
|
||||
}
|
||||
|
||||
func (h *Handler) tunnelUsesNftables(tunnelID int64) (bool, []int64, error) {
|
||||
entryNodeIDs, err := h.tunnelEntryNodeIDs(tunnelID)
|
||||
if err != nil {
|
||||
return false, nil, err
|
||||
}
|
||||
for _, nodeID := range entryNodeIDs {
|
||||
ok, modeErr := h.nodeUsesNftables(nodeID)
|
||||
if modeErr != nil {
|
||||
return false, nil, modeErr
|
||||
}
|
||||
if ok {
|
||||
return true, entryNodeIDs, nil
|
||||
}
|
||||
}
|
||||
return false, entryNodeIDs, nil
|
||||
}
|
||||
|
||||
func (h *Handler) validateNftablesForwardRequest(tunnel *tunnelRecord, remoteAddr string, entryNodeIDs []int64) error {
|
||||
if tunnel == nil {
|
||||
return errors.New("隧道不存在")
|
||||
}
|
||||
if tunnel.Type != 1 {
|
||||
return errors.New("nftables 节点仅支持直连隧道")
|
||||
}
|
||||
if len(entryNodeIDs) != 1 {
|
||||
return errors.New("nftables 节点仅支持单入口隧道")
|
||||
}
|
||||
if _, err := runtimenft.ParseSingleTarget(remoteAddr); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func sshConfigFromModel(cfg *model.NodeSSHConfig) (runtimenft.SSHConfig, error) {
|
||||
if cfg == nil {
|
||||
return runtimenft.SSHConfig{}, errors.New("节点缺少 SSH 配置")
|
||||
}
|
||||
if strings.TrimSpace(cfg.Host) == "" || strings.TrimSpace(cfg.Username) == "" {
|
||||
return runtimenft.SSHConfig{}, errors.New("节点 SSH 配置不完整")
|
||||
}
|
||||
return runtimenft.SSHConfig{
|
||||
Host: strings.TrimSpace(cfg.Host),
|
||||
Port: cfg.Port,
|
||||
Username: strings.TrimSpace(cfg.Username),
|
||||
AuthType: strings.TrimSpace(cfg.AuthType),
|
||||
Password: cfg.Password.String,
|
||||
PrivateKey: cfg.PrivateKey.String,
|
||||
Passphrase: cfg.Passphrase.String,
|
||||
SudoMode: strings.TrimSpace(cfg.SudoMode),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (h *Handler) validateNftablesTunnelState(entryNodeIDs []int64) error {
|
||||
return h.validateNftablesTunnelStateTx(nil, entryNodeIDs)
|
||||
}
|
||||
|
||||
func (h *Handler) validateNftablesTunnelStateTx(tx *gorm.DB, entryNodeIDs []int64) error {
|
||||
if h == nil || h.repo == nil {
|
||||
return errors.New("handler not initialized")
|
||||
}
|
||||
for _, nodeID := range entryNodeIDs {
|
||||
isNft, err := h.nodeUsesNftablesTx(tx, nodeID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !isNft {
|
||||
continue
|
||||
}
|
||||
var cfg *model.NodeSSHConfig
|
||||
if tx != nil {
|
||||
cfg, err = h.repo.GetNodeSSHConfigTx(tx, nodeID)
|
||||
} else {
|
||||
cfg, err = h.repo.GetNodeSSHConfig(nodeID)
|
||||
}
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return errors.New("nftables 节点缺少 SSH 配置")
|
||||
}
|
||||
return err
|
||||
}
|
||||
sshCfg, err := sshConfigFromModel(cfg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if h.nftablesManager == nil {
|
||||
return errors.New("nftables manager not initialized")
|
||||
}
|
||||
if err := h.nftablesManager.Test(context.Background(), sshCfg); err != nil {
|
||||
return fmt.Errorf("nftables 节点能力校验失败: %w", err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Handler) buildNftablesNodePlan(nodeID int64) (runtimenft.NodePlan, *model.NodeSSHConfig, error) {
|
||||
cfg, err := h.repo.GetNodeSSHConfig(nodeID)
|
||||
if err != nil {
|
||||
return runtimenft.NodePlan{}, nil, err
|
||||
}
|
||||
forwards, err := h.repo.ListActiveForwardsByNode(nodeID)
|
||||
if err != nil {
|
||||
return runtimenft.NodePlan{}, nil, err
|
||||
}
|
||||
plan := runtimenft.NodePlan{NodeID: nodeID, Rules: make([]runtimenft.Rule, 0, len(forwards))}
|
||||
for i := range forwards {
|
||||
forward := &forwards[i]
|
||||
tunnel, err := h.getTunnelRecord(forward.TunnelID)
|
||||
if err != nil || tunnel == nil || tunnel.Status != 1 {
|
||||
continue
|
||||
}
|
||||
entryNodeIDs, err := h.tunnelEntryNodeIDs(forward.TunnelID)
|
||||
if err != nil {
|
||||
return runtimenft.NodePlan{}, nil, err
|
||||
}
|
||||
if len(entryNodeIDs) != 1 || entryNodeIDs[0] != nodeID {
|
||||
continue
|
||||
}
|
||||
if err := h.validateNftablesForwardRequest(tunnel, forward.RemoteAddr, entryNodeIDs); err != nil {
|
||||
return runtimenft.NodePlan{}, nil, err
|
||||
}
|
||||
ports, err := h.listForwardPorts(forward.ID)
|
||||
if err != nil {
|
||||
return runtimenft.NodePlan{}, nil, err
|
||||
}
|
||||
for _, fp := range ports {
|
||||
if fp.NodeID != nodeID {
|
||||
continue
|
||||
}
|
||||
target, err := runtimenft.ParseSingleTarget(forward.RemoteAddr)
|
||||
if err != nil {
|
||||
return runtimenft.NodePlan{}, nil, err
|
||||
}
|
||||
plan.Rules = append(plan.Rules, runtimenft.Rule{
|
||||
ForwardID: forward.ID,
|
||||
InPort: fp.Port,
|
||||
BindIP: strings.TrimSpace(fp.InIP),
|
||||
TargetHost: target.Host,
|
||||
TargetPort: target.Port,
|
||||
Protocols: []string{"tcp", "udp"},
|
||||
})
|
||||
}
|
||||
}
|
||||
return plan, cfg, nil
|
||||
}
|
||||
|
||||
func (h *Handler) syncNftablesNode(nodeID int64) error {
|
||||
if h == nil || h.repo == nil {
|
||||
return errors.New("handler not initialized")
|
||||
}
|
||||
plan, cfgModel, err := h.buildNftablesNodePlan(nodeID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
sshCfg, err := sshConfigFromModel(cfgModel)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
result, err := h.nftablesManager.Reconcile(context.Background(), sshCfg, plan)
|
||||
now := time.Now().UnixMilli()
|
||||
if err != nil {
|
||||
bindings, _ := h.repo.ListNftRuleBindingsByNode(nodeID)
|
||||
for _, binding := range bindings {
|
||||
_ = h.repo.MarkNftRuleBindingError(binding.ForwardID, nodeID, err.Error(), now)
|
||||
}
|
||||
return err
|
||||
}
|
||||
activeForwardIDs := make(map[int64]struct{}, len(plan.Rules))
|
||||
for _, rule := range plan.Rules {
|
||||
activeForwardIDs[rule.ForwardID] = struct{}{}
|
||||
hash := result.Hashes[rule.ForwardID]
|
||||
_ = h.repo.UpsertNftRuleBinding(modelToRuleBindingInput(nodeID, rule, hash), now)
|
||||
}
|
||||
bindings, _ := h.repo.ListNftRuleBindingsByNode(nodeID)
|
||||
for _, binding := range bindings {
|
||||
if _, ok := activeForwardIDs[binding.ForwardID]; ok {
|
||||
continue
|
||||
}
|
||||
_ = h.repo.DeleteNftRuleBindingsByForward(binding.ForwardID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func modelToRuleBindingInput(nodeID int64, rule runtimenft.Rule, hash string) repo.NftRuleBindingInput {
|
||||
return repo.NftRuleBindingInput{
|
||||
ForwardID: rule.ForwardID,
|
||||
NodeID: nodeID,
|
||||
InPort: rule.InPort,
|
||||
Protocols: strings.Join(rule.Protocols, ","),
|
||||
TargetAddr: fmt.Sprintf("%s:%d", rule.TargetHost, rule.TargetPort),
|
||||
BindIP: rule.BindIP,
|
||||
RuleHash: hash,
|
||||
Status: runtimenft.StatusApplied,
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) nftablesNodeIDFromRequest(r *http.Request, w http.ResponseWriter) (int64, bool) {
|
||||
nodeID := asInt64FromBodyKey(r, w, "nodeId")
|
||||
if nodeID <= 0 {
|
||||
return 0, false
|
||||
}
|
||||
return nodeID, true
|
||||
}
|
||||
|
||||
func (h *Handler) loadNftablesSSHConfig(nodeID int64) (runtimenft.SSHConfig, error) {
|
||||
cfg, err := h.repo.GetNodeSSHConfig(nodeID)
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return runtimenft.SSHConfig{}, errors.New("nftables 节点缺少 SSH 配置")
|
||||
}
|
||||
return runtimenft.SSHConfig{}, err
|
||||
}
|
||||
return sshConfigFromModel(cfg)
|
||||
}
|
||||
|
||||
func (h *Handler) clearNftablesNode(nodeID int64) error {
|
||||
if h == nil || h.repo == nil {
|
||||
return errors.New("handler not initialized")
|
||||
}
|
||||
sshCfg, err := h.loadNftablesSSHConfig(nodeID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if h.nftablesManager == nil {
|
||||
return errors.New("nftables manager not initialized")
|
||||
}
|
||||
if err := h.nftablesManager.Clear(context.Background(), sshCfg); err != nil {
|
||||
return err
|
||||
}
|
||||
bindings, listErr := h.repo.ListNftRuleBindingsByNode(nodeID)
|
||||
if listErr != nil {
|
||||
return listErr
|
||||
}
|
||||
for _, binding := range bindings {
|
||||
if err := h.repo.DeleteNftRuleBindingsByForward(binding.ForwardID); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Handler) reconcileNftablesNodeByRequest(nodeID int64) error {
|
||||
usesNft, err := h.nodeUsesNftables(nodeID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !usesNft {
|
||||
return errors.New("节点未启用 nftables 转发模式")
|
||||
}
|
||||
return h.syncNftablesNode(nodeID)
|
||||
}
|
||||
|
||||
func (h *Handler) nodeNftablesTest(w http.ResponseWriter, r *http.Request) {
|
||||
nodeID, ok := h.nftablesNodeIDFromRequest(r, w)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
usesNft, err := h.nodeUsesNftables(nodeID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if !usesNft {
|
||||
response.WriteJSON(w, response.ErrDefault("节点未启用 nftables 转发模式"))
|
||||
return
|
||||
}
|
||||
sshCfg, err := h.loadNftablesSSHConfig(nodeID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
if h.nftablesManager == nil {
|
||||
response.WriteJSON(w, response.Err(-2, "nftables manager not initialized"))
|
||||
return
|
||||
}
|
||||
if err := h.nftablesManager.Test(context.Background(), sshCfg); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
func (h *Handler) nodeNftablesReconcile(w http.ResponseWriter, r *http.Request) {
|
||||
nodeID, ok := h.nftablesNodeIDFromRequest(r, w)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if err := h.reconcileNftablesNodeByRequest(nodeID); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
func (h *Handler) nodeNftablesClear(w http.ResponseWriter, r *http.Request) {
|
||||
nodeID, ok := h.nftablesNodeIDFromRequest(r, w)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
usesNft, err := h.nodeUsesNftables(nodeID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if !usesNft {
|
||||
response.WriteJSON(w, response.ErrDefault("节点未启用 nftables 转发模式"))
|
||||
return
|
||||
}
|
||||
if err := h.clearNftablesNode(nodeID); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
@@ -0,0 +1,478 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/auth"
|
||||
"go-backend/internal/http/middleware"
|
||||
runtimenft "go-backend/internal/runtime/nftables"
|
||||
"go-backend/internal/store/repo"
|
||||
)
|
||||
|
||||
type fakeNftablesManager struct {
|
||||
testErr error
|
||||
reconcileErr error
|
||||
reconcileHit int
|
||||
clearErr error
|
||||
clearHit int
|
||||
lastConfig runtimenft.SSHConfig
|
||||
lastPlan runtimenft.NodePlan
|
||||
}
|
||||
|
||||
func (f *fakeNftablesManager) Test(_ context.Context, cfg runtimenft.SSHConfig) error {
|
||||
f.lastConfig = cfg
|
||||
return f.testErr
|
||||
}
|
||||
|
||||
func (f *fakeNftablesManager) Reconcile(_ context.Context, cfg runtimenft.SSHConfig, plan runtimenft.NodePlan) (runtimenft.ApplyResult, error) {
|
||||
f.reconcileHit++
|
||||
f.lastConfig = cfg
|
||||
f.lastPlan = plan
|
||||
if f.reconcileErr != nil {
|
||||
return runtimenft.ApplyResult{}, f.reconcileErr
|
||||
}
|
||||
return runtimenft.ApplyResult{
|
||||
NodeID: plan.NodeID,
|
||||
Script: "table inet flvx {}",
|
||||
Hashes: map[int64]string{plan.NodeID: "hash"},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (f *fakeNftablesManager) Clear(context.Context, runtimenft.SSHConfig) error {
|
||||
f.clearHit++
|
||||
return f.clearErr
|
||||
}
|
||||
|
||||
type nftablesTestFixture struct {
|
||||
handler *Handler
|
||||
nodeID int64
|
||||
}
|
||||
|
||||
func TestTunnelCreateRejectsNftablesEntryNodeWithoutSSHConfig(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
err := fixture.handler.validateNftablesTunnelState([]int64{fixture.nodeID})
|
||||
if err == nil {
|
||||
t.Fatalf("expected validation failure")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "SSH") {
|
||||
t.Fatalf("expected SSH config validation error, got %q", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelUpdateRejectsNftablesEntryNodeWhenCapabilityTestFails(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
manager := &fakeNftablesManager{testErr: errors.New("ssh failed")}
|
||||
h.nftablesManager = manager
|
||||
seedNftablesSSHConfig(t, h, fixture.nodeID)
|
||||
err := h.validateNftablesTunnelState([]int64{fixture.nodeID})
|
||||
if err == nil {
|
||||
t.Fatalf("expected validation failure")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "ssh failed") {
|
||||
t.Fatalf("expected capability error in response, got %q", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSyncForwardServicesWithWarningsUsesNftablesRuntime(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
manager := &fakeNftablesManager{}
|
||||
h.nftablesManager = manager
|
||||
seedNftablesSSHConfig(t, h, fixture.nodeID)
|
||||
tunnelID := seedTunnelForNftables(t, h, "nft-tunnel", fixture.nodeID)
|
||||
forward := seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080")
|
||||
|
||||
warnings, err := h.syncForwardServicesWithWarnings(forward, "UpdateService", true)
|
||||
if err != nil {
|
||||
t.Fatalf("sync forward services: %v", err)
|
||||
}
|
||||
if len(warnings) != 0 {
|
||||
t.Fatalf("expected no warnings, got %v", warnings)
|
||||
}
|
||||
if manager.reconcileHit != 1 {
|
||||
t.Fatalf("expected nftables reconcile to run once, got %d", manager.reconcileHit)
|
||||
}
|
||||
if manager.lastPlan.NodeID != fixture.nodeID {
|
||||
t.Fatalf("expected plan for node %d, got %+v", fixture.nodeID, manager.lastPlan)
|
||||
}
|
||||
if len(manager.lastPlan.Rules) != 1 || manager.lastPlan.Rules[0].ForwardID != forward.ID {
|
||||
t.Fatalf("unexpected plan: %+v", manager.lastPlan)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeNftablesTestEndpointRunsCapabilityCheck(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
seedNftablesSSHConfig(t, fixture.handler, fixture.nodeID)
|
||||
manager := &fakeNftablesManager{}
|
||||
fixture.handler.nftablesManager = manager
|
||||
|
||||
res := postJSONToHandler(t, fixture.handler.nodeNftablesTest, map[string]int64{"nodeId": fixture.nodeID})
|
||||
assertNftablesSuccess(t, res)
|
||||
if manager.lastConfig.Host != "203.0.113.10" {
|
||||
t.Fatalf("expected SSH config to be passed to manager, got %+v", manager.lastConfig)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeNftablesReconcileEndpointPersistsBindings(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
seedNftablesSSHConfig(t, fixture.handler, fixture.nodeID)
|
||||
tunnelID := seedTunnelForNftables(t, fixture.handler, "nft-tunnel", fixture.nodeID)
|
||||
forward := seedForwardForNftables(t, fixture.handler, tunnelID, fixture.nodeID, "203.0.113.9:8080")
|
||||
manager := &fakeNftablesManager{}
|
||||
fixture.handler.nftablesManager = manager
|
||||
|
||||
res := postJSONToHandler(t, fixture.handler.nodeNftablesReconcile, map[string]int64{"nodeId": fixture.nodeID})
|
||||
assertNftablesSuccess(t, res)
|
||||
if manager.reconcileHit != 1 {
|
||||
t.Fatalf("expected reconcile once, got %d", manager.reconcileHit)
|
||||
}
|
||||
bindings, err := fixture.handler.repo.ListNftRuleBindingsByNode(fixture.nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("list bindings: %v", err)
|
||||
}
|
||||
if len(bindings) != 1 || bindings[0].ForwardID != forward.ID {
|
||||
t.Fatalf("unexpected bindings: %+v", bindings)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeNftablesClearEndpointClearsBindings(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
seedNftablesSSHConfig(t, fixture.handler, fixture.nodeID)
|
||||
now := time.Now().UnixMilli()
|
||||
if err := fixture.handler.repo.UpsertNftRuleBinding(repo.NftRuleBindingInput{
|
||||
ForwardID: 99,
|
||||
NodeID: fixture.nodeID,
|
||||
InPort: 24000,
|
||||
Protocols: "tcp",
|
||||
TargetAddr: "203.0.113.9:8080",
|
||||
Status: runtimenft.StatusApplied,
|
||||
}, now); err != nil {
|
||||
t.Fatalf("seed binding: %v", err)
|
||||
}
|
||||
manager := &fakeNftablesManager{}
|
||||
fixture.handler.nftablesManager = manager
|
||||
|
||||
res := postJSONToHandler(t, fixture.handler.nodeNftablesClear, map[string]int64{"nodeId": fixture.nodeID})
|
||||
assertNftablesSuccess(t, res)
|
||||
if manager.clearHit != 1 {
|
||||
t.Fatalf("expected clear once, got %d", manager.clearHit)
|
||||
}
|
||||
if bindings, err := fixture.handler.repo.ListNftRuleBindingsByNode(fixture.nodeID); err != nil {
|
||||
t.Fatalf("list bindings after clear: %v", err)
|
||||
} else if len(bindings) != 0 {
|
||||
t.Fatalf("expected bindings to be cleared, got %+v", bindings)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeCreatePersistsNftablesSSHConfig(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
req := newAuthenticatedJSONRequest(t, map[string]interface{}{
|
||||
"name": "nft-node-created",
|
||||
"serverIp": "203.0.113.20",
|
||||
"serverIpV4": "203.0.113.20",
|
||||
"port": "20000-20100",
|
||||
"forwardMode": "nftables",
|
||||
"sshConfig": map[string]interface{}{
|
||||
"host": "203.0.113.21",
|
||||
"port": 2222,
|
||||
"username": "root",
|
||||
"authType": "private_key",
|
||||
"privateKey": "TEST-PRIVATE-KEY",
|
||||
"passphrase": "secret",
|
||||
"sudoMode": "sudo",
|
||||
},
|
||||
})
|
||||
res := httptest.NewRecorder()
|
||||
fixture.handler.nodeCreate(res, req)
|
||||
assertNftablesSuccessWithBody(t, res)
|
||||
|
||||
nodes, err := fixture.handler.repo.ListNodes()
|
||||
if err != nil {
|
||||
t.Fatalf("list nodes: %v", err)
|
||||
}
|
||||
var createdNodeID int64
|
||||
for _, item := range nodes {
|
||||
if item["name"] == "nft-node-created" {
|
||||
createdNodeID = item["id"].(int64)
|
||||
break
|
||||
}
|
||||
}
|
||||
if createdNodeID <= 0 {
|
||||
t.Fatalf("expected created node to exist")
|
||||
}
|
||||
createdNode, err := fixture.handler.repo.GetNodeRecord(createdNodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("load created node: %v", err)
|
||||
}
|
||||
if createdNode == nil {
|
||||
t.Fatal("expected created node record, got nil")
|
||||
}
|
||||
if createdNode.Status != 1 {
|
||||
t.Fatalf("expected nftables node to be online, got status %d", createdNode.Status)
|
||||
}
|
||||
cfg, err := fixture.handler.repo.GetNodeSSHConfig(createdNodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("load ssh config: %v", err)
|
||||
}
|
||||
if cfg.Host != "203.0.113.21" || cfg.Port != 2222 || cfg.Username != "root" || cfg.AuthType != "private_key" {
|
||||
t.Fatalf("unexpected ssh config: %+v", cfg)
|
||||
}
|
||||
if !cfg.PrivateKey.Valid || cfg.PrivateKey.String != "TEST-PRIVATE-KEY" {
|
||||
t.Fatalf("expected private key to persist, got %+v", cfg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNodeUpdatePreservesExistingNftablesSecretsWhenFieldsOmitted(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
seedNftablesSSHConfig(t, fixture.handler, fixture.nodeID)
|
||||
|
||||
req := newAuthenticatedJSONRequest(t, map[string]interface{}{
|
||||
"id": fixture.nodeID,
|
||||
"name": "nft-node-updated",
|
||||
"serverIp": "198.51.100.10",
|
||||
"serverIpV4": "198.51.100.10",
|
||||
"port": "1000-65535",
|
||||
"forwardMode": "nftables",
|
||||
"sshConfig": map[string]interface{}{
|
||||
"host": "203.0.113.30",
|
||||
"port": 22,
|
||||
"username": "admin",
|
||||
"authType": "password",
|
||||
"sudoMode": "none",
|
||||
},
|
||||
})
|
||||
res := httptest.NewRecorder()
|
||||
fixture.handler.nodeUpdate(res, req)
|
||||
assertNftablesSuccessWithBody(t, res)
|
||||
|
||||
cfg, err := fixture.handler.repo.GetNodeSSHConfig(fixture.nodeID)
|
||||
if err != nil {
|
||||
t.Fatalf("load ssh config: %v", err)
|
||||
}
|
||||
if cfg.Host != "203.0.113.30" || cfg.Username != "admin" || cfg.AuthType != "password" {
|
||||
t.Fatalf("unexpected ssh config after update: %+v", cfg)
|
||||
}
|
||||
if !cfg.Password.Valid || cfg.Password.String != "secret" {
|
||||
t.Fatalf("expected password secret to be preserved, got %+v", cfg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardForceDeleteRemovesNftablesBindingAndReconciles(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
seedNftablesSSHConfig(t, h, fixture.nodeID)
|
||||
tunnelID := seedTunnelForNftables(t, h, "nft-tunnel", fixture.nodeID)
|
||||
forward := seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080")
|
||||
if err := h.repo.UpsertNftRuleBinding(repo.NftRuleBindingInput{
|
||||
ForwardID: forward.ID,
|
||||
NodeID: fixture.nodeID,
|
||||
InPort: 20000,
|
||||
Protocols: "tcp,udp",
|
||||
TargetAddr: "203.0.113.9:8080",
|
||||
Status: runtimenft.StatusApplied,
|
||||
}, time.Now().UnixMilli()); err != nil {
|
||||
t.Fatalf("seed binding: %v", err)
|
||||
}
|
||||
manager := &fakeNftablesManager{}
|
||||
h.nftablesManager = manager
|
||||
|
||||
req := newAuthenticatedJSONRequest(t, map[string]int64{"id": forward.ID})
|
||||
req.URL.Path = "/api/v1/forward/force-delete"
|
||||
res := httptest.NewRecorder()
|
||||
mux := http.NewServeMux()
|
||||
h.Register(mux)
|
||||
mux.ServeHTTP(res, req)
|
||||
assertNftablesSuccessWithBody(t, res)
|
||||
if manager.reconcileHit != 1 {
|
||||
t.Fatalf("expected reconcile once, got %d", manager.reconcileHit)
|
||||
}
|
||||
if _, err := h.getForwardRecord(forward.ID); !errors.Is(err, errForwardNotFound) {
|
||||
t.Fatalf("expected forward to be deleted, got %v", err)
|
||||
}
|
||||
if bindings, err := h.repo.ListNftRuleBindingsByNode(fixture.nodeID); err != nil {
|
||||
t.Fatalf("list bindings after delete: %v", err)
|
||||
} else if len(bindings) != 0 {
|
||||
t.Fatalf("expected no bindings after delete, got %+v", bindings)
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardBatchRedeployUsesNftablesReconcile(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
seedNftablesSSHConfig(t, h, fixture.nodeID)
|
||||
tunnelID := seedTunnelForNftables(t, h, "nft-tunnel", fixture.nodeID)
|
||||
forward := seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080")
|
||||
manager := &fakeNftablesManager{}
|
||||
h.nftablesManager = manager
|
||||
|
||||
req := newAuthenticatedJSONRequest(t, map[string][]int64{"ids": {forward.ID}})
|
||||
res := httptest.NewRecorder()
|
||||
h.forwardBatchRedeploy(res, req)
|
||||
assertNftablesSuccessWithBody(t, res)
|
||||
if manager.reconcileHit != 1 {
|
||||
t.Fatalf("expected reconcile once, got %d", manager.reconcileHit)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTunnelBatchRedeployUsesNftablesReconcile(t *testing.T) {
|
||||
fixture := setupNftablesHandler(t)
|
||||
h := fixture.handler
|
||||
seedNftablesSSHConfig(t, h, fixture.nodeID)
|
||||
tunnelID := seedTunnelForNftables(t, h, "nft-tunnel", fixture.nodeID)
|
||||
seedForwardForNftables(t, h, tunnelID, fixture.nodeID, "203.0.113.9:8080")
|
||||
manager := &fakeNftablesManager{}
|
||||
h.nftablesManager = manager
|
||||
|
||||
req := newAuthenticatedJSONRequest(t, map[string][]int64{"ids": {tunnelID}})
|
||||
res := httptest.NewRecorder()
|
||||
h.tunnelBatchRedeploy(res, req)
|
||||
assertNftablesSuccessWithBody(t, res)
|
||||
if manager.reconcileHit != 1 {
|
||||
t.Fatalf("expected reconcile once, got %d", manager.reconcileHit)
|
||||
}
|
||||
}
|
||||
|
||||
func setupNftablesHandler(t *testing.T) nftablesTestFixture {
|
||||
t.Helper()
|
||||
|
||||
dbPath := filepath.Join(t.TempDir(), "handler-nftables.sqlite")
|
||||
r, err := repo.Open(dbPath)
|
||||
if err != nil {
|
||||
t.Fatalf("open repo: %v", err)
|
||||
}
|
||||
|
||||
h := New(r, "test-secret")
|
||||
now := time.Now().UnixMilli()
|
||||
if _, err := r.CreateUser("admin", "hash", 0, now+86400000, 1, 1, 100, 1, 0, now); err != nil {
|
||||
t.Fatalf("create user: %v", err)
|
||||
}
|
||||
if err := r.CreateNode("nft-node", "secret", "198.51.100.10", nil, nil, "1000-65535", nil, nil, nil, nil, nil, 0, 0, 0, now, 1, "", "", 1, 0, nil, nil, nil, nil, "nftables"); err != nil {
|
||||
t.Fatalf("create node: %v", err)
|
||||
}
|
||||
node, err := r.GetNodeRecord(1)
|
||||
if err != nil || node == nil {
|
||||
t.Fatalf("get node: %v", err)
|
||||
}
|
||||
return nftablesTestFixture{handler: h, nodeID: node.ID}
|
||||
}
|
||||
|
||||
func seedNftablesSSHConfig(t *testing.T, h *Handler, nodeID int64) {
|
||||
t.Helper()
|
||||
if err := h.repo.UpsertNodeSSHConfig(nodeID, repo.NftSSHConfigInput{
|
||||
Host: "203.0.113.10",
|
||||
Port: 22,
|
||||
Username: "root",
|
||||
AuthType: "password",
|
||||
Password: "secret",
|
||||
SudoMode: "none",
|
||||
}, time.Now().UnixMilli()); err != nil {
|
||||
t.Fatalf("upsert ssh config: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func seedTunnelForNftables(t *testing.T, h *Handler, name string, nodeID int64) int64 {
|
||||
t.Helper()
|
||||
now := time.Now().UnixMilli()
|
||||
tx := h.repo.BeginTx()
|
||||
if tx == nil {
|
||||
t.Fatal("begin tx: nil transaction")
|
||||
}
|
||||
if tx.Error != nil {
|
||||
t.Fatalf("begin tx: %v", tx.Error)
|
||||
}
|
||||
tunnelID, err := h.repo.CreateTunnelTx(tx, name, 1, 1, 1, now, 1, nil, 1, "", "", 0)
|
||||
if err != nil {
|
||||
_ = tx.Rollback().Error
|
||||
t.Fatalf("create tunnel: %v", err)
|
||||
}
|
||||
if err := h.repo.CreateChainTunnelTx(tx, tunnelID, "1", nodeID, sql.NullInt64{}, "", 1, "tls", ""); err != nil {
|
||||
_ = tx.Rollback().Error
|
||||
t.Fatalf("create chain tunnel: %v", err)
|
||||
}
|
||||
if err := tx.Commit().Error; err != nil {
|
||||
_ = tx.Rollback().Error
|
||||
t.Fatalf("commit tx: %v", err)
|
||||
}
|
||||
return tunnelID
|
||||
}
|
||||
|
||||
func seedForwardForNftables(t *testing.T, h *Handler, tunnelID, nodeID int64, remoteAddr string) *forwardRecord {
|
||||
t.Helper()
|
||||
now := time.Now().UnixMilli()
|
||||
forwardID, err := h.repo.CreateForwardTx(
|
||||
1, "admin", "nft-forward", tunnelID, remoteAddr, "fifo", now, 1,
|
||||
[]int64{nodeID}, 20000, "", nil, 0, 0, nil, 0,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("create forward: %v", err)
|
||||
}
|
||||
forward, err := h.getForwardRecord(forwardID)
|
||||
if err != nil {
|
||||
t.Fatalf("get forward: %v", err)
|
||||
}
|
||||
return forward
|
||||
}
|
||||
|
||||
func postJSONToHandler(t *testing.T, fn func(http.ResponseWriter, *http.Request), payload any) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
body, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal payload: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/", bytes.NewReader(body))
|
||||
res := httptest.NewRecorder()
|
||||
fn(res, req)
|
||||
return res
|
||||
}
|
||||
|
||||
func newAuthenticatedJSONRequest(t *testing.T, payload any) *http.Request {
|
||||
t.Helper()
|
||||
body, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal payload: %v", err)
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/", bytes.NewReader(body))
|
||||
token, err := auth.GenerateToken(1, "admin", 0, "test-secret")
|
||||
if err != nil {
|
||||
t.Fatalf("create token: %v", err)
|
||||
}
|
||||
req.Header.Set("Authorization", token)
|
||||
claims, ok := auth.ValidateToken(token, "test-secret")
|
||||
if !ok {
|
||||
t.Fatalf("validate token failed")
|
||||
}
|
||||
return req.WithContext(context.WithValue(req.Context(), middleware.ClaimsContextKey, claims))
|
||||
}
|
||||
|
||||
func assertNftablesSuccess(t *testing.T, res *httptest.ResponseRecorder) {
|
||||
t.Helper()
|
||||
assertNftablesSuccessWithBody(t, res)
|
||||
}
|
||||
|
||||
func assertNftablesSuccessWithBody(t *testing.T, res *httptest.ResponseRecorder) {
|
||||
t.Helper()
|
||||
var payload struct {
|
||||
Code int `json:"code"`
|
||||
Msg string `json:"msg"`
|
||||
}
|
||||
if res.Code != http.StatusOK {
|
||||
t.Fatalf("expected HTTP %d, got %d", http.StatusOK, res.Code)
|
||||
}
|
||||
if err := json.NewDecoder(res.Body).Decode(&payload); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if payload.Code != 0 {
|
||||
t.Fatalf("expected success, got %+v", payload)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user