feat(backend): orchestrate federation runtime for shared middle and exit nodes

This commit is contained in:
sagit
2026-02-10 02:34:20 +00:00
parent 5f42c50689
commit 79f8aab600
8 changed files with 1406 additions and 30 deletions
@@ -51,6 +51,10 @@ type nodeRecord struct {
TCPListenAddr string
UDPListenAddr string
InterfaceName string
IsRemote int
RemoteURL string
RemoteToken string
RemoteConfig string
}
type chainNodeRecord struct {
@@ -197,7 +201,7 @@ func (h *Handler) listForwardPorts(forwardID int64) ([]forwardPortRecord, error)
func (h *Handler) getNodeRecord(nodeID int64) (*nodeRecord, error) {
row := h.repo.DB().QueryRow(`
SELECT id, name, server_ip, server_ip_v4, server_ip_v6, status, port, tcp_listen_addr, udp_listen_addr, interface_name
SELECT id, name, server_ip, server_ip_v4, server_ip_v6, status, port, tcp_listen_addr, udp_listen_addr, interface_name, is_remote, remote_url, remote_token, remote_config
FROM node
WHERE id = ?
LIMIT 1
@@ -209,7 +213,10 @@ func (h *Handler) getNodeRecord(nodeID int64) (*nodeRecord, error) {
var tcpListen sql.NullString
var udpListen sql.NullString
var iface sql.NullString
err := row.Scan(&n.ID, &n.Name, &n.ServerIP, &serverIPv4, &serverIPv6, &n.Status, &portRange, &tcpListen, &udpListen, &iface)
var remoteURL sql.NullString
var remoteToken sql.NullString
var remoteConfig sql.NullString
err := row.Scan(&n.ID, &n.Name, &n.ServerIP, &serverIPv4, &serverIPv6, &n.Status, &portRange, &tcpListen, &udpListen, &iface, &n.IsRemote, &remoteURL, &remoteToken, &remoteConfig)
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return nil, errors.New("节点不存在")
@@ -222,6 +229,9 @@ func (h *Handler) getNodeRecord(nodeID int64) (*nodeRecord, error) {
n.TCPListenAddr = strings.TrimSpace(tcpListen.String)
n.UDPListenAddr = strings.TrimSpace(udpListen.String)
n.InterfaceName = strings.TrimSpace(iface.String)
n.RemoteURL = strings.TrimSpace(remoteURL.String)
n.RemoteToken = strings.TrimSpace(remoteToken.String)
n.RemoteConfig = strings.TrimSpace(remoteConfig.String)
if n.TCPListenAddr == "" {
n.TCPListenAddr = "[::]"
}
+410 -1
View File
@@ -1,6 +1,7 @@
package handler
import (
"database/sql"
"encoding/json"
"fmt"
"net/http"
@@ -37,6 +38,33 @@ type nodeImportRequest struct {
Token string `json:"token"`
}
type federationRuntimeReservePortRequest struct {
ResourceKey string `json:"resourceKey"`
Protocol string `json:"protocol"`
RequestedPort int `json:"requestedPort"`
}
type federationRuntimeTarget struct {
Host string `json:"host"`
Port int `json:"port"`
Protocol string `json:"protocol"`
}
type federationRuntimeApplyRoleRequest struct {
ReservationID string `json:"reservationId"`
ResourceKey string `json:"resourceKey"`
Role string `json:"role"`
Protocol string `json:"protocol"`
Strategy string `json:"strategy"`
Targets []federationRuntimeTarget `json:"targets"`
}
type federationRuntimeReleaseRoleRequest struct {
BindingID string `json:"bindingId"`
ReservationID string `json:"reservationId"`
ResourceKey string `json:"resourceKey"`
}
func (h *Handler) federationShareList(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("Invalid method"))
@@ -183,6 +211,11 @@ func (h *Handler) nodeImport(w http.ResponseWriter, r *http.Request) {
}
configBytes, _ := json.Marshal(configData)
portRange := "0"
if info.PortRangeStart > 0 && info.PortRangeEnd >= info.PortRangeStart {
portRange = fmt.Sprintf("%d-%d", info.PortRangeStart, info.PortRangeEnd)
}
db := h.repo.DB()
inx := nextIndex(db, "node")
now := time.Now().UnixMilli()
@@ -195,7 +228,7 @@ func (h *Handler) nodeImport(w http.ResponseWriter, r *http.Request) {
randomToken(16), // Dummy secret
info.ServerIP,
"", "", // v4/v6 unknown, use server_ip
"0", // port range not applicable for remote
portRange,
"",
"",
now, now,
@@ -386,6 +419,382 @@ func (h *Handler) federationTunnelCreate(w http.ResponseWriter, r *http.Request)
}))
}
func (h *Handler) federationRuntimeReservePort(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("Invalid method"))
return
}
token := extractBearerToken(r)
share, err := h.repo.GetPeerShareByToken(token)
if err != nil || share == nil {
response.WriteJSON(w, response.Err(401, "Unauthorized"))
return
}
var req federationRuntimeReservePortRequest
if err := decodeJSON(r.Body, &req); err != nil {
response.WriteJSON(w, response.ErrDefault("Invalid JSON"))
return
}
req.ResourceKey = strings.TrimSpace(req.ResourceKey)
if req.ResourceKey == "" {
response.WriteJSON(w, response.ErrDefault("resourceKey is required"))
return
}
existing, err := h.repo.GetPeerShareRuntimeByResourceKey(share.ID, req.ResourceKey)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if existing != nil && existing.Status == 1 {
response.WriteJSON(w, response.OK(map[string]interface{}{
"reservationId": existing.ReservationID,
"allocatedPort": existing.Port,
"bindingId": existing.BindingID,
}))
return
}
allocatedPort, err := h.pickPeerSharePort(share, req.RequestedPort)
if err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
now := time.Now().UnixMilli()
if existing != nil {
existing.Protocol = defaultString(req.Protocol, "tls")
existing.Port = allocatedPort
existing.BindingID = ""
existing.Role = ""
existing.ChainName = ""
existing.ServiceName = ""
existing.Strategy = "round"
existing.Target = ""
existing.Applied = 0
existing.Status = 1
existing.UpdatedTime = now
if err := h.repo.UpdatePeerShareRuntime(existing); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OK(map[string]interface{}{
"reservationId": existing.ReservationID,
"allocatedPort": existing.Port,
"bindingId": existing.BindingID,
}))
return
}
runtime := &sqlite.PeerShareRuntime{
ShareID: share.ID,
NodeID: share.NodeID,
ReservationID: randomToken(24),
ResourceKey: req.ResourceKey,
BindingID: "",
Role: "",
ChainName: "",
ServiceName: "",
Protocol: defaultString(req.Protocol, "tls"),
Strategy: "round",
Port: allocatedPort,
Target: "",
Applied: 0,
Status: 1,
CreatedTime: now,
UpdatedTime: now,
}
if err := h.repo.CreatePeerShareRuntime(runtime); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OK(map[string]interface{}{
"reservationId": runtime.ReservationID,
"allocatedPort": runtime.Port,
"bindingId": runtime.BindingID,
}))
}
func (h *Handler) federationRuntimeApplyRole(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("Invalid method"))
return
}
token := extractBearerToken(r)
share, err := h.repo.GetPeerShareByToken(token)
if err != nil || share == nil {
response.WriteJSON(w, response.Err(401, "Unauthorized"))
return
}
var req federationRuntimeApplyRoleRequest
if err := decodeJSON(r.Body, &req); err != nil {
response.WriteJSON(w, response.ErrDefault("Invalid JSON"))
return
}
req.Role = strings.ToLower(strings.TrimSpace(req.Role))
if req.Role != "middle" && req.Role != "exit" {
response.WriteJSON(w, response.ErrDefault("Invalid role"))
return
}
var runtime *sqlite.PeerShareRuntime
if strings.TrimSpace(req.ReservationID) != "" {
runtime, err = h.repo.GetPeerShareRuntimeByReservationID(share.ID, strings.TrimSpace(req.ReservationID))
} else {
runtime, err = h.repo.GetPeerShareRuntimeByResourceKey(share.ID, strings.TrimSpace(req.ResourceKey))
}
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if runtime == nil || runtime.Status == 0 {
response.WriteJSON(w, response.ErrDefault("Reservation not found"))
return
}
if runtime.Applied == 1 && strings.TrimSpace(runtime.BindingID) != "" {
response.WriteJSON(w, response.OK(map[string]interface{}{
"bindingId": runtime.BindingID,
"allocatedPort": runtime.Port,
"reservationId": runtime.ReservationID,
}))
return
}
node, err := h.getNodeRecord(share.NodeID)
if err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
protocol := defaultString(req.Protocol, runtime.Protocol)
strategy := defaultString(req.Strategy, "round")
chainName := fmt.Sprintf("fed_chain_%d", runtime.ID)
serviceName := fmt.Sprintf("fed_svc_%d", runtime.ID)
if req.Role == "middle" {
if len(req.Targets) == 0 {
response.WriteJSON(w, response.ErrDefault("targets are required for middle role"))
return
}
nodeItems := make([]map[string]interface{}, 0, len(req.Targets))
for i, target := range req.Targets {
host := strings.TrimSpace(target.Host)
if host == "" || target.Port <= 0 {
response.WriteJSON(w, response.ErrDefault("Invalid target"))
return
}
nodeItems = append(nodeItems, map[string]interface{}{
"name": fmt.Sprintf("node_%d", i+1),
"addr": processServerAddress(fmt.Sprintf("%s:%d", host, target.Port)),
"connector": map[string]interface{}{
"type": "relay",
},
"dialer": map[string]interface{}{
"type": defaultString(target.Protocol, protocol),
},
})
}
chainData := map[string]interface{}{
"name": chainName,
"hops": []map[string]interface{}{
{
"name": fmt.Sprintf("hop_%d", runtime.ID),
"selector": map[string]interface{}{
"strategy": strategy,
"maxFails": 1,
"failTimeout": int64(600000000000),
},
"nodes": nodeItems,
},
},
}
if strings.TrimSpace(node.InterfaceName) != "" {
hops := chainData["hops"].([]map[string]interface{})
hops[0]["interface"] = node.InterfaceName
}
if _, err := h.sendNodeCommand(share.NodeID, "AddChains", chainData, true, false); err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
}
service := map[string]interface{}{
"name": serviceName,
"addr": fmt.Sprintf("%s:%d", node.TCPListenAddr, runtime.Port),
"handler": map[string]interface{}{
"type": "relay",
},
"listener": map[string]interface{}{
"type": protocol,
},
}
if req.Role == "middle" {
service["handler"].(map[string]interface{})["chain"] = chainName
}
if req.Role == "exit" && strings.TrimSpace(node.InterfaceName) != "" {
service["metadata"] = map[string]interface{}{"interface": node.InterfaceName}
}
if _, err := h.sendNodeCommand(share.NodeID, "AddService", []map[string]interface{}{service}, true, false); err != nil {
if req.Role == "middle" {
_, _ = h.sendNodeCommand(share.NodeID, "DeleteChains", map[string]interface{}{"chain": chainName}, false, true)
}
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
targetBytes, _ := json.Marshal(req.Targets)
runtime.BindingID = fmt.Sprintf("%d", runtime.ID)
runtime.Role = req.Role
runtime.ChainName = ""
if req.Role == "middle" {
runtime.ChainName = chainName
}
runtime.ServiceName = serviceName
runtime.Protocol = protocol
runtime.Strategy = strategy
runtime.Target = string(targetBytes)
runtime.Applied = 1
runtime.Status = 1
runtime.UpdatedTime = time.Now().UnixMilli()
if err := h.repo.UpdatePeerShareRuntime(runtime); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OK(map[string]interface{}{
"bindingId": runtime.BindingID,
"reservationId": runtime.ReservationID,
"allocatedPort": runtime.Port,
}))
}
func (h *Handler) federationRuntimeReleaseRole(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("Invalid method"))
return
}
token := extractBearerToken(r)
share, err := h.repo.GetPeerShareByToken(token)
if err != nil || share == nil {
response.WriteJSON(w, response.Err(401, "Unauthorized"))
return
}
var req federationRuntimeReleaseRoleRequest
if err := decodeJSON(r.Body, &req); err != nil {
response.WriteJSON(w, response.ErrDefault("Invalid JSON"))
return
}
var runtime *sqlite.PeerShareRuntime
if strings.TrimSpace(req.BindingID) != "" {
runtime, err = h.repo.GetPeerShareRuntimeByBindingID(share.ID, strings.TrimSpace(req.BindingID))
} else if strings.TrimSpace(req.ReservationID) != "" {
runtime, err = h.repo.GetPeerShareRuntimeByReservationID(share.ID, strings.TrimSpace(req.ReservationID))
} else if strings.TrimSpace(req.ResourceKey) != "" {
runtime, err = h.repo.GetPeerShareRuntimeByResourceKey(share.ID, strings.TrimSpace(req.ResourceKey))
} else {
response.WriteJSON(w, response.ErrDefault("bindingId or reservationId or resourceKey is required"))
return
}
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if runtime == nil {
response.WriteJSON(w, response.OKEmpty())
return
}
if runtime.Applied == 1 {
if strings.TrimSpace(runtime.ServiceName) != "" {
_, _ = h.sendNodeCommand(share.NodeID, "DeleteService", map[string]interface{}{"services": []string{runtime.ServiceName}}, false, true)
}
if strings.TrimSpace(runtime.Role) == "middle" && strings.TrimSpace(runtime.ChainName) != "" {
_, _ = h.sendNodeCommand(share.NodeID, "DeleteChains", map[string]interface{}{"chain": runtime.ChainName}, false, true)
}
}
if err := h.repo.MarkPeerShareRuntimeReleased(runtime.ID, time.Now().UnixMilli()); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) pickPeerSharePort(share *sqlite.PeerShare, requestedPort int) (int, error) {
if share == nil {
return 0, fmt.Errorf("share not found")
}
if share.PortRangeStart <= 0 || share.PortRangeEnd <= 0 || share.PortRangeEnd < share.PortRangeStart {
return 0, fmt.Errorf("No available port")
}
used := make(map[int]struct{})
rows, err := h.repo.DB().Query(`SELECT port FROM chain_tunnel WHERE node_id = ? AND port IS NOT NULL AND port > 0`, share.NodeID)
if err != nil {
return 0, err
}
for rows.Next() {
var p sql.NullInt64
if scanErr := rows.Scan(&p); scanErr == nil && p.Valid && p.Int64 > 0 {
used[int(p.Int64)] = struct{}{}
}
}
_ = rows.Close()
rows, err = h.repo.DB().Query(`SELECT port FROM forward_port WHERE node_id = ? AND port > 0`, share.NodeID)
if err != nil {
return 0, err
}
for rows.Next() {
var p sql.NullInt64
if scanErr := rows.Scan(&p); scanErr == nil && p.Valid && p.Int64 > 0 {
used[int(p.Int64)] = struct{}{}
}
}
_ = rows.Close()
ports, err := h.repo.ListActivePeerShareRuntimePorts(share.ID, share.NodeID)
if err != nil {
return 0, err
}
for _, p := range ports {
if p > 0 {
used[p] = struct{}{}
}
}
if requestedPort > 0 {
if requestedPort < share.PortRangeStart || requestedPort > share.PortRangeEnd {
return 0, fmt.Errorf("Port out of range")
}
if _, ok := used[requestedPort]; ok {
return 0, fmt.Errorf("No available port")
}
return requestedPort, nil
}
for p := share.PortRangeStart; p <= share.PortRangeEnd; p++ {
if _, ok := used[p]; ok {
continue
}
return p, nil
}
return 0, fmt.Errorf("No available port")
}
func extractBearerToken(r *http.Request) string {
authHeader := r.Header.Get("Authorization")
parts := strings.Split(authHeader, " ")
+125 -12
View File
@@ -27,6 +27,9 @@ type Handler struct {
jwtSecret string
wsServer *ws.Server
captchaMu sync.Mutex
captchaTokens map[string]int64
jobsMu sync.Mutex
jobsCancel context.CancelFunc
jobsStarted bool
@@ -39,6 +42,11 @@ type loginRequest struct {
CaptchaID string `json:"captchaId"`
}
type captchaVerifyRequest struct {
ID string `json:"id"`
Data string `json:"data"`
}
type nameRequest struct {
Name string `json:"name"`
}
@@ -63,9 +71,10 @@ type flowItem struct {
func New(repo *sqlite.Repository, jwtSecret string) *Handler {
return &Handler{
repo: repo,
jwtSecret: jwtSecret,
wsServer: ws.NewServer(repo, jwtSecret),
repo: repo,
jwtSecret: jwtSecret,
wsServer: ws.NewServer(repo, jwtSecret),
captchaTokens: make(map[string]int64),
}
}
@@ -85,6 +94,7 @@ func (h *Handler) Register(mux *http.ServeMux) {
mux.HandleFunc("/api/v1/config/update", h.updateConfigs)
mux.HandleFunc("/api/v1/config/update-single", h.updateSingleConfig)
mux.HandleFunc("/api/v1/captcha/check", h.checkCaptcha)
mux.HandleFunc("/api/v1/captcha/verify", h.captchaVerify)
mux.HandleFunc("/api/v1/user/package", h.userPackage)
mux.HandleFunc("/api/v1/user/updatePassword", h.updatePassword)
mux.HandleFunc("/api/v1/node/list", h.nodeList)
@@ -148,6 +158,9 @@ func (h *Handler) Register(mux *http.ServeMux) {
mux.HandleFunc("/api/v1/federation/share/delete", h.federationShareDelete)
mux.HandleFunc("/api/v1/federation/connect", h.authPeer(h.federationConnect))
mux.HandleFunc("/api/v1/federation/tunnel/create", h.authPeer(h.federationTunnelCreate))
mux.HandleFunc("/api/v1/federation/runtime/reserve-port", h.authPeer(h.federationRuntimeReservePort))
mux.HandleFunc("/api/v1/federation/runtime/apply-role", h.authPeer(h.federationRuntimeApplyRole))
mux.HandleFunc("/api/v1/federation/runtime/release-role", h.authPeer(h.federationRuntimeReleaseRole))
mux.HandleFunc("/api/v1/federation/node/import", h.nodeImport)
mux.HandleFunc("/flow/test", h.flowTest)
@@ -183,20 +196,23 @@ func (h *Handler) login(w http.ResponseWriter, r *http.Request) {
return
}
if captchaEnabled {
if strings.TrimSpace(req.CaptchaID) == "" {
captchaID := strings.TrimSpace(req.CaptchaID)
if captchaID == "" {
response.WriteJSON(w, response.ErrDefault("验证码校验失败"))
return
}
secretCfg, err := h.repo.GetConfigByName("cloudflare_secret_key")
if err != nil || secretCfg == nil || secretCfg.Value == "" {
response.WriteJSON(w, response.ErrDefault("验证码配置错误:未配置Secret Key"))
return
}
if !h.consumeCaptchaToken(captchaID) {
secretCfg, err := h.repo.GetConfigByName("cloudflare_secret_key")
if err != nil || secretCfg == nil || strings.TrimSpace(secretCfg.Value) == "" {
response.WriteJSON(w, response.ErrDefault("验证码校验失败"))
return
}
if !h.verifyCloudflareTurnstile(req.CaptchaID, secretCfg.Value) {
response.WriteJSON(w, response.ErrDefault("验证码校验失败"))
return
if !h.verifyCloudflareTurnstile(captchaID, strings.TrimSpace(secretCfg.Value)) {
response.WriteJSON(w, response.ErrDefault("验证码校验失败"))
return
}
}
}
@@ -585,6 +601,40 @@ func (h *Handler) checkCaptcha(w http.ResponseWriter, r *http.Request) {
response.WriteJSON(w, response.OK(0))
}
func (h *Handler) captchaVerify(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
return
}
var req captchaVerifyRequest
if err := decodeJSON(r.Body, &req); err != nil {
h.writeCaptchaVerifyResult(w, false, "")
return
}
id := strings.TrimSpace(req.ID)
data := strings.TrimSpace(req.Data)
if id == "" || data == "" {
h.writeCaptchaVerifyResult(w, false, "")
return
}
verified := false
secretCfg, err := h.repo.GetConfigByName("cloudflare_secret_key")
if err == nil && secretCfg != nil && strings.TrimSpace(secretCfg.Value) != "" {
verified = h.verifyCloudflareTurnstile(data, strings.TrimSpace(secretCfg.Value))
} else {
verified = data == "ok"
}
if !verified {
h.writeCaptchaVerifyResult(w, false, "")
return
}
h.markCaptchaToken(id)
h.writeCaptchaVerifyResult(w, true, id)
}
func (h *Handler) flowTest(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("Content-Type", "text/plain; charset=utf-8")
_, _ = w.Write([]byte("test"))
@@ -899,6 +949,69 @@ func (h *Handler) captchaEnabled() (bool, error) {
return strings.EqualFold(cfg.Value, "true"), nil
}
func (h *Handler) markCaptchaToken(token string) {
if h == nil {
return
}
token = strings.TrimSpace(token)
if token == "" {
return
}
now := time.Now().UnixMilli()
exp := now + int64(5*time.Minute/time.Millisecond)
h.captchaMu.Lock()
defer h.captchaMu.Unlock()
if h.captchaTokens == nil {
h.captchaTokens = make(map[string]int64)
}
for k, v := range h.captchaTokens {
if v <= now {
delete(h.captchaTokens, k)
}
}
h.captchaTokens[token] = exp
}
func (h *Handler) consumeCaptchaToken(token string) bool {
if h == nil {
return false
}
token = strings.TrimSpace(token)
if token == "" {
return false
}
now := time.Now().UnixMilli()
h.captchaMu.Lock()
defer h.captchaMu.Unlock()
if h.captchaTokens == nil {
return false
}
for k, v := range h.captchaTokens {
if v <= now {
delete(h.captchaTokens, k)
}
}
exp, ok := h.captchaTokens[token]
if !ok {
return false
}
delete(h.captchaTokens, token)
return exp > now
}
func (h *Handler) writeCaptchaVerifyResult(w http.ResponseWriter, success bool, token string) {
w.Header().Set("Content-Type", "application/json; charset=utf-8")
payload := map[string]interface{}{
"success": success,
"data": map[string]interface{}{
"validToken": token,
},
}
_ = json.NewEncoder(w).Encode(payload)
}
func decodeJSON(body io.ReadCloser, out interface{}) error {
defer body.Close()
decoder := json.NewDecoder(body)
+443 -15
View File
@@ -18,6 +18,7 @@ import (
"go-backend/internal/http/client"
"go-backend/internal/http/response"
"go-backend/internal/security"
"go-backend/internal/store/sqlite"
)
func (h *Handler) userCreate(w http.ResponseWriter, r *http.Request) {
@@ -505,7 +506,7 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
}
defer func() { _ = tx.Rollback() }()
runtimeState, err := h.prepareTunnelCreateState(tx, req, typeVal)
runtimeState, err := h.prepareTunnelCreateState(tx, req, typeVal, 0)
if err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
@@ -566,12 +567,28 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
}
tunnelID, _ := res.LastInsertId()
runtimeState.TunnelID = tunnelID
var federationBindings []sqlite.FederationTunnelBinding
var federationReleaseRefs []federationRuntimeReleaseRef
if typeVal == 2 {
federationBindings, federationReleaseRefs, err = h.applyFederationRuntime(runtimeState)
if err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
}
applyTunnelPortsToRequest(req, runtimeState)
if err := replaceTunnelChainsTx(tx, tunnelID, req); err != nil {
h.releaseFederationRuntimeRefs(federationReleaseRefs)
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if err := replaceFederationTunnelBindingsTx(tx, tunnelID, federationBindings); err != nil {
h.releaseFederationRuntimeRefs(federationReleaseRefs)
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if err := tx.Commit(); err != nil {
h.releaseFederationRuntimeRefs(federationReleaseRefs)
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
@@ -579,6 +596,7 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
createdChains, createdServices, applyErr := h.applyTunnelRuntime(runtimeState)
if applyErr != nil {
h.rollbackTunnelRuntime(createdChains, createdServices, tunnelID)
h.releaseFederationRuntimeRefs(federationReleaseRefs)
_ = h.deleteTunnelByID(tunnelID)
response.WriteJSON(w, response.ErrDefault(applyErr.Error()))
return
@@ -652,6 +670,7 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
}
h.cleanupTunnelRuntime(id)
h.cleanupFederationRuntime(id)
now := time.Now().UnixMilli()
typeVal := asInt(req["type"], 1)
@@ -663,12 +682,21 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
}
defer func() { _ = tx.Rollback() }()
runtimeState, err := h.prepareTunnelCreateState(tx, req, typeVal)
runtimeState, err := h.prepareTunnelCreateState(tx, req, typeVal, id)
if err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
runtimeState.TunnelID = id
var federationBindings []sqlite.FederationTunnelBinding
var federationReleaseRefs []federationRuntimeReleaseRef
if typeVal == 2 {
federationBindings, federationReleaseRefs, err = h.applyFederationRuntime(runtimeState)
if err != nil {
response.WriteJSON(w, response.ErrDefault(err.Error()))
return
}
}
applyTunnelPortsToRequest(req, runtimeState)
_, err = tx.Exec(`UPDATE tunnel SET name=?, type=?, flow=?, traffic_ratio=?, status=?, in_ip=?, updated_time=? WHERE id=?`,
@@ -683,10 +711,17 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
return
}
if err := replaceTunnelChainsTx(tx, id, req); err != nil {
h.releaseFederationRuntimeRefs(federationReleaseRefs)
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if err := replaceFederationTunnelBindingsTx(tx, id, federationBindings); err != nil {
h.releaseFederationRuntimeRefs(federationReleaseRefs)
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
if err := tx.Commit(); err != nil {
h.releaseFederationRuntimeRefs(federationReleaseRefs)
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
@@ -695,6 +730,12 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
createdChains, createdServices, applyErr := h.applyTunnelRuntime(runtimeState)
if applyErr != nil {
h.rollbackTunnelRuntime(createdChains, createdServices, id)
h.releaseFederationRuntimeRefs(federationReleaseRefs)
_ = h.repo.DeleteFederationTunnelBindingsByTunnel(id)
if len(federationReleaseRefs) == 0 && shouldDeferTunnelRuntimeApplyError(applyErr) {
response.WriteJSON(w, response.OKEmpty())
return
}
response.WriteJSON(w, response.ErrDefault(applyErr.Error()))
return
}
@@ -713,6 +754,7 @@ func (h *Handler) tunnelDelete(w http.ResponseWriter, r *http.Request) {
return
}
h.cleanupTunnelRuntime(id)
h.cleanupFederationRuntime(id)
if err := h.deleteTunnelByID(id); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
@@ -771,6 +813,7 @@ func (h *Handler) tunnelBatchDelete(w http.ResponseWriter, r *http.Request) {
fail := 0
for _, id := range ids {
h.cleanupTunnelRuntime(id)
h.cleanupFederationRuntime(id)
if err := h.deleteTunnelByID(id); err != nil {
fail++
} else {
@@ -872,13 +915,38 @@ func (h *Handler) tunnelBatchRedeploy(w http.ResponseWriter, r *http.Request) {
if tunnel.Type == 2 {
h.cleanupTunnelRuntime(tunnelID)
h.cleanupFederationRuntime(tunnelID)
state, err := h.reconstructTunnelState(tunnelID)
if err != nil {
fail++
continue
}
federationBindings, federationReleaseRefs, fedErr := h.applyFederationRuntime(state)
if fedErr != nil {
fail++
continue
}
tx, txErr := h.repo.DB().Begin()
if txErr != nil {
h.releaseFederationRuntimeRefs(federationReleaseRefs)
fail++
continue
}
if replaceErr := replaceFederationTunnelBindingsTx(tx, tunnelID, federationBindings); replaceErr != nil {
_ = tx.Rollback()
h.releaseFederationRuntimeRefs(federationReleaseRefs)
fail++
continue
}
if commitErr := tx.Commit(); commitErr != nil {
h.releaseFederationRuntimeRefs(federationReleaseRefs)
fail++
continue
}
_, _, applyErr := h.applyTunnelRuntime(state)
if applyErr != nil {
h.releaseFederationRuntimeRefs(federationReleaseRefs)
_ = h.repo.DeleteFederationTunnelBindingsByTunnel(tunnelID)
fail++
continue
}
@@ -1876,7 +1944,7 @@ type tunnelCreateState struct {
NodeIDList []int64
}
func (h *Handler) prepareTunnelCreateState(tx *sql.Tx, req map[string]interface{}, tunnelType int) (*tunnelCreateState, error) {
func (h *Handler) prepareTunnelCreateState(tx *sql.Tx, req map[string]interface{}, tunnelType int, excludeTunnelID int64) (*tunnelCreateState, error) {
state := &tunnelCreateState{
Type: tunnelType,
InNodes: make([]tunnelRuntimeNode, 0),
@@ -1918,10 +1986,16 @@ func (h *Handler) prepareTunnelCreateState(tx *sql.Tx, req map[string]interface{
nodeIDs = append(nodeIDs, nodeID)
port := asInt(item["port"], 0)
if port <= 0 {
var err error
port, err = pickNodePortTx(tx, nodeID, allocated)
if err != nil {
return nil, err
isRemote, remoteErr := isRemoteNodeTx(tx, nodeID)
if remoteErr != nil {
return nil, remoteErr
}
if !isRemote {
var err error
port, err = pickNodePortTx(tx, nodeID, allocated, excludeTunnelID)
if err != nil {
return nil, err
}
}
}
state.OutNodes = append(state.OutNodes, tunnelRuntimeNode{
@@ -1946,10 +2020,16 @@ func (h *Handler) prepareTunnelCreateState(tx *sql.Tx, req map[string]interface{
nodeIDs = append(nodeIDs, nodeID)
port := asInt(item["port"], 0)
if port <= 0 {
var err error
port, err = pickNodePortTx(tx, nodeID, allocated)
if err != nil {
return nil, err
isRemote, remoteErr := isRemoteNodeTx(tx, nodeID)
if remoteErr != nil {
return nil, remoteErr
}
if !isRemote {
var err error
port, err = pickNodePortTx(tx, nodeID, allocated, excludeTunnelID)
if err != nil {
return nil, err
}
}
}
hop = append(hop, tunnelRuntimeNode{
@@ -2053,6 +2133,303 @@ func applyTunnelPortsToRequest(req map[string]interface{}, state *tunnelCreateSt
}
}
type federationRuntimeReleaseRef struct {
RemoteURL string
RemoteToken string
BindingID string
ReservationID string
ResourceKey string
}
func federationRuntimeResourceKey(tunnelID int64, nodeID int64, chainType int, hopInx int) string {
return fmt.Sprintf("tunnel:%d:node:%d:type:%d:hop:%d", tunnelID, nodeID, chainType, hopInx)
}
func remoteShareIDFromConfig(raw string) int64 {
raw = strings.TrimSpace(raw)
if raw == "" {
return 0
}
var cfg map[string]interface{}
if err := json.Unmarshal([]byte(raw), &cfg); err != nil {
return 0
}
return asInt64(cfg["shareId"], 0)
}
func (h *Handler) federationLocalDomain() string {
cfg, _ := h.repo.GetConfigByName("panel_domain")
if cfg == nil {
return ""
}
return strings.TrimSpace(cfg.Value)
}
func (h *Handler) applyFederationRuntime(state *tunnelCreateState) ([]sqlite.FederationTunnelBinding, []federationRuntimeReleaseRef, error) {
bindings := make([]sqlite.FederationTunnelBinding, 0)
releaseRefs := make([]federationRuntimeReleaseRef, 0)
if h == nil || state == nil || state.Type != 2 {
return bindings, releaseRefs, nil
}
fc := client.NewFederationClient()
localDomain := h.federationLocalDomain()
now := time.Now().UnixMilli()
for outIdx := range state.OutNodes {
outNode := state.OutNodes[outIdx]
node := state.Nodes[outNode.NodeID]
if node == nil || node.IsRemote != 1 {
continue
}
remoteURL := strings.TrimSpace(node.RemoteURL)
remoteToken := strings.TrimSpace(node.RemoteToken)
if remoteURL == "" || remoteToken == "" {
h.releaseFederationRuntimeRefs(releaseRefs)
return nil, nil, fmt.Errorf("远程节点 %s 缺少共享配置", nodeDisplayName(node))
}
resourceKey := federationRuntimeResourceKey(state.TunnelID, outNode.NodeID, 3, 0)
reserveReq := client.RuntimeReservePortRequest{
ResourceKey: resourceKey,
Protocol: defaultString(outNode.Protocol, "tls"),
RequestedPort: outNode.Port,
}
reserveRes, err := fc.ReservePort(remoteURL, remoteToken, localDomain, reserveReq)
if err != nil && reserveReq.RequestedPort > 0 {
reserveReq.RequestedPort = 0
reserveRes, err = fc.ReservePort(remoteURL, remoteToken, localDomain, reserveReq)
}
if err != nil {
h.releaseFederationRuntimeRefs(releaseRefs)
return nil, nil, fmt.Errorf("远程节点 %s 端口分配失败: %w", nodeDisplayName(node), err)
}
state.OutNodes[outIdx].Port = reserveRes.AllocatedPort
outNode = state.OutNodes[outIdx]
applyReq := client.RuntimeApplyRoleRequest{
ReservationID: reserveRes.ReservationID,
ResourceKey: resourceKey,
Role: "exit",
Protocol: defaultString(outNode.Protocol, "tls"),
Strategy: defaultString(outNode.Strategy, "round"),
}
applyRes, err := fc.ApplyRole(remoteURL, remoteToken, localDomain, applyReq)
if err != nil {
h.releaseFederationRuntimeRefs(releaseRefs)
return nil, nil, fmt.Errorf("远程节点 %s 运行时下发失败: %w", nodeDisplayName(node), err)
}
if applyRes.AllocatedPort > 0 {
state.OutNodes[outIdx].Port = applyRes.AllocatedPort
outNode = state.OutNodes[outIdx]
}
bindings = append(bindings, sqlite.FederationTunnelBinding{
TunnelID: state.TunnelID,
NodeID: outNode.NodeID,
ChainType: 3,
HopInx: 0,
RemoteURL: remoteURL,
ResourceKey: resourceKey,
RemoteBindingID: defaultString(applyRes.BindingID, reserveRes.BindingID),
AllocatedPort: outNode.Port,
Status: 1,
CreatedTime: now,
UpdatedTime: now,
})
releaseRefs = append(releaseRefs, federationRuntimeReleaseRef{
RemoteURL: remoteURL,
RemoteToken: remoteToken,
BindingID: applyRes.BindingID,
ReservationID: reserveRes.ReservationID,
ResourceKey: resourceKey,
})
}
for hopIdx := len(state.ChainHops) - 1; hopIdx >= 0; hopIdx-- {
for nodeIdx := range state.ChainHops[hopIdx] {
chainNode := state.ChainHops[hopIdx][nodeIdx]
node := state.Nodes[chainNode.NodeID]
if node == nil || node.IsRemote != 1 {
continue
}
remoteURL := strings.TrimSpace(node.RemoteURL)
remoteToken := strings.TrimSpace(node.RemoteToken)
if remoteURL == "" || remoteToken == "" {
h.releaseFederationRuntimeRefs(releaseRefs)
return nil, nil, fmt.Errorf("远程节点 %s 缺少共享配置", nodeDisplayName(node))
}
resourceKey := federationRuntimeResourceKey(state.TunnelID, chainNode.NodeID, 2, hopIdx+1)
reserveReq := client.RuntimeReservePortRequest{
ResourceKey: resourceKey,
Protocol: defaultString(chainNode.Protocol, "tls"),
RequestedPort: chainNode.Port,
}
reserveRes, err := fc.ReservePort(remoteURL, remoteToken, localDomain, reserveReq)
if err != nil && reserveReq.RequestedPort > 0 {
reserveReq.RequestedPort = 0
reserveRes, err = fc.ReservePort(remoteURL, remoteToken, localDomain, reserveReq)
}
if err != nil {
h.releaseFederationRuntimeRefs(releaseRefs)
return nil, nil, fmt.Errorf("远程节点 %s 端口分配失败: %w", nodeDisplayName(node), err)
}
state.ChainHops[hopIdx][nodeIdx].Port = reserveRes.AllocatedPort
chainNode = state.ChainHops[hopIdx][nodeIdx]
nextTargets := state.OutNodes
if hopIdx+1 < len(state.ChainHops) {
nextTargets = state.ChainHops[hopIdx+1]
}
applyTargets := make([]client.RuntimeTarget, 0, len(nextTargets))
for _, target := range nextTargets {
targetNode := state.Nodes[target.NodeID]
if targetNode == nil {
h.releaseFederationRuntimeRefs(releaseRefs)
return nil, nil, errors.New("节点不存在")
}
host, hostErr := selectTunnelDialHost(node, targetNode)
if hostErr != nil {
h.releaseFederationRuntimeRefs(releaseRefs)
return nil, nil, hostErr
}
if target.Port <= 0 {
h.releaseFederationRuntimeRefs(releaseRefs)
return nil, nil, errors.New("节点端口不能为空")
}
applyTargets = append(applyTargets, client.RuntimeTarget{
Host: host,
Port: target.Port,
Protocol: defaultString(target.Protocol, "tls"),
})
}
applyReq := client.RuntimeApplyRoleRequest{
ReservationID: reserveRes.ReservationID,
ResourceKey: resourceKey,
Role: "middle",
Protocol: defaultString(chainNode.Protocol, "tls"),
Strategy: defaultString(chainNode.Strategy, "round"),
Targets: applyTargets,
}
applyRes, err := fc.ApplyRole(remoteURL, remoteToken, localDomain, applyReq)
if err != nil {
h.releaseFederationRuntimeRefs(releaseRefs)
return nil, nil, fmt.Errorf("远程节点 %s 运行时下发失败: %w", nodeDisplayName(node), err)
}
if applyRes.AllocatedPort > 0 {
state.ChainHops[hopIdx][nodeIdx].Port = applyRes.AllocatedPort
chainNode = state.ChainHops[hopIdx][nodeIdx]
}
bindings = append(bindings, sqlite.FederationTunnelBinding{
TunnelID: state.TunnelID,
NodeID: chainNode.NodeID,
ChainType: 2,
HopInx: hopIdx + 1,
RemoteURL: remoteURL,
ResourceKey: resourceKey,
RemoteBindingID: defaultString(applyRes.BindingID, reserveRes.BindingID),
AllocatedPort: chainNode.Port,
Status: 1,
CreatedTime: now,
UpdatedTime: now,
})
releaseRefs = append(releaseRefs, federationRuntimeReleaseRef{
RemoteURL: remoteURL,
RemoteToken: remoteToken,
BindingID: applyRes.BindingID,
ReservationID: reserveRes.ReservationID,
ResourceKey: resourceKey,
})
}
}
return bindings, releaseRefs, nil
}
func (h *Handler) releaseFederationRuntimeRefs(refs []federationRuntimeReleaseRef) {
if h == nil || len(refs) == 0 {
return
}
fc := client.NewFederationClient()
localDomain := h.federationLocalDomain()
for i := len(refs) - 1; i >= 0; i-- {
ref := refs[i]
if strings.TrimSpace(ref.RemoteURL) == "" || strings.TrimSpace(ref.RemoteToken) == "" {
continue
}
req := client.RuntimeReleaseRoleRequest{
BindingID: ref.BindingID,
ReservationID: ref.ReservationID,
ResourceKey: ref.ResourceKey,
}
_ = fc.ReleaseRole(ref.RemoteURL, ref.RemoteToken, localDomain, req)
}
}
func (h *Handler) cleanupFederationRuntime(tunnelID int64) {
if h == nil || tunnelID <= 0 {
return
}
bindings, err := h.repo.ListActiveFederationTunnelBindingsByTunnel(tunnelID)
if err != nil || len(bindings) == 0 {
return
}
fc := client.NewFederationClient()
localDomain := h.federationLocalDomain()
for _, b := range bindings {
node, nodeErr := h.repo.GetNodeByID(b.NodeID)
if nodeErr != nil || node == nil {
continue
}
remoteURL := strings.TrimSpace(node.RemoteURL.String)
if remoteURL == "" {
remoteURL = strings.TrimSpace(b.RemoteURL)
}
remoteToken := strings.TrimSpace(node.RemoteToken.String)
if remoteURL == "" || remoteToken == "" {
continue
}
req := client.RuntimeReleaseRoleRequest{
BindingID: strings.TrimSpace(b.RemoteBindingID),
ResourceKey: strings.TrimSpace(b.ResourceKey),
}
_ = fc.ReleaseRole(remoteURL, remoteToken, localDomain, req)
}
_ = h.repo.DeleteFederationTunnelBindingsByTunnel(tunnelID)
}
func replaceFederationTunnelBindingsTx(tx *sql.Tx, tunnelID int64, bindings []sqlite.FederationTunnelBinding) error {
if tx == nil {
return errors.New("database unavailable")
}
if _, err := tx.Exec(`DELETE FROM federation_tunnel_binding WHERE tunnel_id = ?`, tunnelID); err != nil {
return err
}
for _, b := range bindings {
created := b.CreatedTime
if created <= 0 {
created = time.Now().UnixMilli()
}
updated := b.UpdatedTime
if updated <= 0 {
updated = created
}
_, err := tx.Exec(`
INSERT INTO federation_tunnel_binding(tunnel_id, node_id, chain_type, hop_inx, remote_url, resource_key, remote_binding_id, allocated_port, status, created_time, updated_time)
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, tunnelID, b.NodeID, b.ChainType, b.HopInx, b.RemoteURL, b.ResourceKey, b.RemoteBindingID, b.AllocatedPort, b.Status, created, updated)
if err != nil {
return err
}
}
return nil
}
func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64, error) {
if h == nil || state == nil {
return nil, nil, errors.New("invalid tunnel runtime state")
@@ -2064,6 +2441,9 @@ func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64
}
for _, inNode := range state.InNodes {
if node := state.Nodes[inNode.NodeID]; node != nil && node.IsRemote == 1 {
continue
}
targets := state.OutNodes
if len(state.ChainHops) > 0 {
targets = state.ChainHops[0]
@@ -2084,6 +2464,9 @@ func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64
nextTargets = state.ChainHops[i+1]
}
for _, chainNode := range hop {
if node := state.Nodes[chainNode.NodeID]; node != nil && node.IsRemote == 1 {
continue
}
chainData, err := buildTunnelChainConfig(state.TunnelID, chainNode.NodeID, nextTargets, state.Nodes)
if err != nil {
return createdChains, createdServices, err
@@ -2102,6 +2485,9 @@ func (h *Handler) applyTunnelRuntime(state *tunnelCreateState) ([]int64, []int64
}
for _, outNode := range state.OutNodes {
if node := state.Nodes[outNode.NodeID]; node != nil && node.IsRemote == 1 {
continue
}
serviceData := buildTunnelChainServiceConfig(state.TunnelID, outNode, state.Nodes[outNode.NodeID])
if _, err := h.sendNodeCommand(outNode.NodeID, "AddService", serviceData, true, false); err != nil {
return createdChains, createdServices, fmt.Errorf("出口节点 %s 下发服务失败: %w", nodeDisplayName(state.Nodes[outNode.NodeID]), err)
@@ -2139,6 +2525,23 @@ func (h *Handler) rollbackTunnelRuntime(chainNodeIDs, serviceNodeIDs []int64, tu
}
}
func shouldDeferTunnelRuntimeApplyError(err error) bool {
if err == nil {
return false
}
msg := strings.ToLower(strings.TrimSpace(err.Error()))
if msg == "" {
return false
}
if strings.Contains(msg, "节点不在线") {
return true
}
if strings.Contains(msg, "等待节点响应超时") || strings.Contains(msg, "timeout") || strings.Contains(msg, "超时") {
return true
}
return false
}
func buildTunnelChainConfig(tunnelID int64, fromNodeID int64, targets []tunnelRuntimeNode, nodes map[int64]*nodeRecord) (map[string]interface{}, error) {
fromNode := nodes[fromNodeID]
if fromNode == nil {
@@ -2310,7 +2713,24 @@ func pickNodeAddressV6(node *nodeRecord) string {
return strings.TrimSpace(node.ServerIP)
}
func pickNodePortTx(tx *sql.Tx, nodeID int64, allocated map[int64]int) (int, error) {
func isRemoteNodeTx(tx *sql.Tx, nodeID int64) (bool, error) {
if tx == nil {
return false, errors.New("database unavailable")
}
if nodeID <= 0 {
return false, errors.New("节点不存在")
}
var isRemote int
if err := tx.QueryRow(`SELECT is_remote FROM node WHERE id = ? LIMIT 1`, nodeID).Scan(&isRemote); err != nil {
if errors.Is(err, sql.ErrNoRows) {
return false, errors.New("节点不存在")
}
return false, err
}
return isRemote == 1, nil
}
func pickNodePortTx(tx *sql.Tx, nodeID int64, allocated map[int64]int, excludeTunnelID int64) (int, error) {
if tx == nil {
return 0, errors.New("database unavailable")
}
@@ -2334,7 +2754,13 @@ func pickNodePortTx(tx *sql.Tx, nodeID int64, allocated map[int64]int) (int, err
}
used := map[int]struct{}{}
chainRows, err := tx.Query(`SELECT port FROM chain_tunnel WHERE node_id = ? AND port IS NOT NULL`, nodeID)
var chainRows *sql.Rows
var err error
if excludeTunnelID > 0 {
chainRows, err = tx.Query(`SELECT port FROM chain_tunnel WHERE node_id = ? AND port IS NOT NULL AND tunnel_id != ?`, nodeID, excludeTunnelID)
} else {
chainRows, err = tx.Query(`SELECT port FROM chain_tunnel WHERE node_id = ? AND port IS NOT NULL`, nodeID)
}
if err != nil {
return 0, err
}
@@ -2437,7 +2863,7 @@ func replaceTunnelChainsTx(tx *sql.Tx, tunnelID int64, req map[string]interface{
port := asInt(n["port"], 0)
if port <= 0 {
var pickErr error
port, pickErr = pickNodePortTx(tx, nodeID, allocated)
port, pickErr = pickNodePortTx(tx, nodeID, allocated, 0)
if pickErr != nil {
return pickErr
}
@@ -2458,7 +2884,7 @@ func replaceTunnelChainsTx(tx *sql.Tx, tunnelID int64, req map[string]interface{
port := asInt(n["port"], 0)
if port <= 0 {
var pickErr error
port, pickErr = pickNodePortTx(tx, nodeID, allocated)
port, pickErr = pickNodePortTx(tx, nodeID, allocated, 0)
if pickErr != nil {
return pickErr
}
@@ -2481,6 +2907,7 @@ func (h *Handler) deleteNodeByID(id int64) error {
defer func() { _ = tx.Rollback() }()
_, _ = tx.Exec(`DELETE FROM forward_port WHERE node_id = ?`, id)
_, _ = tx.Exec(`DELETE FROM chain_tunnel WHERE node_id = ?`, id)
_, _ = tx.Exec(`DELETE FROM federation_tunnel_binding WHERE node_id = ?`, id)
_, err = tx.Exec(`DELETE FROM node WHERE id = ?`, id)
if err != nil {
return err
@@ -2499,6 +2926,7 @@ func (h *Handler) deleteTunnelByID(id int64) error {
_, _ = tx.Exec(`DELETE FROM user_tunnel WHERE tunnel_id = ?`, id)
_, _ = tx.Exec(`DELETE FROM speed_limit WHERE tunnel_id = ?`, id)
_, _ = tx.Exec(`DELETE FROM chain_tunnel WHERE tunnel_id = ?`, id)
_, _ = tx.Exec(`DELETE FROM federation_tunnel_binding WHERE tunnel_id = ?`, id)
_, err = tx.Exec(`DELETE FROM tunnel WHERE id = ?`, id)
if err != nil {
return err