diff --git a/go-backend/internal/http/client/federation.go b/go-backend/internal/http/client/federation.go index e1dd5fd..4fce731 100644 --- a/go-backend/internal/http/client/federation.go +++ b/go-backend/internal/http/client/federation.go @@ -30,6 +30,45 @@ type RemoteTunnelResponse struct { TunnelID int64 `json:"tunnelId"` } +type RuntimeReservePortRequest struct { + ResourceKey string `json:"resourceKey"` + Protocol string `json:"protocol"` + RequestedPort int `json:"requestedPort"` +} + +type RuntimeReservePortResponse struct { + ReservationID string `json:"reservationId"` + BindingID string `json:"bindingId"` + AllocatedPort int `json:"allocatedPort"` +} + +type RuntimeTarget struct { + Host string `json:"host"` + Port int `json:"port"` + Protocol string `json:"protocol"` +} + +type RuntimeApplyRoleRequest struct { + ReservationID string `json:"reservationId"` + ResourceKey string `json:"resourceKey"` + Role string `json:"role"` + Protocol string `json:"protocol"` + Strategy string `json:"strategy"` + Targets []RuntimeTarget `json:"targets"` +} + +type RuntimeApplyRoleResponse struct { + BindingID string `json:"bindingId"` + ReservationID string `json:"reservationId"` + AllocatedPort int `json:"allocatedPort"` +} + +type RuntimeReleaseRoleRequest struct { + BindingID string `json:"bindingId"` + ReservationID string `json:"reservationId"` + ResourceKey string `json:"resourceKey"` +} + func NewFederationClient() *FederationClient { return &FederationClient{ client: &http.Client{ @@ -119,3 +158,119 @@ func (c *FederationClient) CreateTunnel(url, token, localDomain, protocol string return &res.Data, nil } + +func (c *FederationClient) ReservePort(url, token, localDomain string, reqData RuntimeReservePortRequest) (*RuntimeReservePortResponse, error) { + url = strings.TrimSuffix(url, "/") + bodyBytes, _ := json.Marshal(reqData) + req, err := http.NewRequest("POST", url+"/api/v1/federation/runtime/reserve-port", strings.NewReader(string(bodyBytes))) + if err != nil { + return nil, err + } + req.Header.Set("Authorization", "Bearer "+token) + if localDomain != "" { + req.Header.Set("X-Panel-Domain", localDomain) + } + req.Header.Set("Content-Type", "application/json") + + resp, err := c.client.Do(req) + if err != nil { + return nil, err + } + defer resp.Body.Close() + + if resp.StatusCode != 200 { + body, _ := io.ReadAll(resp.Body) + return nil, fmt.Errorf("remote error %d: %s", resp.StatusCode, string(body)) + } + + var res struct { + Code int `json:"code"` + Msg string `json:"msg"` + Data RuntimeReservePortResponse `json:"data"` + } + if err := json.NewDecoder(resp.Body).Decode(&res); err != nil { + return nil, err + } + if res.Code != 0 { + return nil, fmt.Errorf("remote api error: %s", res.Msg) + } + + return &res.Data, nil +} + +func (c *FederationClient) ApplyRole(url, token, localDomain string, reqData RuntimeApplyRoleRequest) (*RuntimeApplyRoleResponse, error) { + url = strings.TrimSuffix(url, "/") + bodyBytes, _ := json.Marshal(reqData) + req, err := http.NewRequest("POST", url+"/api/v1/federation/runtime/apply-role", strings.NewReader(string(bodyBytes))) + if err != nil { + return nil, err + } + req.Header.Set("Authorization", "Bearer "+token) + if localDomain != "" { + req.Header.Set("X-Panel-Domain", localDomain) + } + req.Header.Set("Content-Type", "application/json") + + resp, err := c.client.Do(req) + if err != nil { + return nil, err + } + defer resp.Body.Close() + + if resp.StatusCode != 200 { + body, _ := io.ReadAll(resp.Body) + return nil, fmt.Errorf("remote error %d: %s", resp.StatusCode, string(body)) + } + + var res struct { + Code int `json:"code"` + Msg string `json:"msg"` + Data RuntimeApplyRoleResponse `json:"data"` + } + if err := json.NewDecoder(resp.Body).Decode(&res); err != nil { + return nil, err + } + if res.Code != 0 { + return nil, fmt.Errorf("remote api error: %s", res.Msg) + } + + return &res.Data, nil +} + +func (c *FederationClient) ReleaseRole(url, token, localDomain string, reqData RuntimeReleaseRoleRequest) error { + url = strings.TrimSuffix(url, "/") + bodyBytes, _ := json.Marshal(reqData) + req, err := http.NewRequest("POST", url+"/api/v1/federation/runtime/release-role", strings.NewReader(string(bodyBytes))) + if err != nil { + return err + } + req.Header.Set("Authorization", "Bearer "+token) + if localDomain != "" { + req.Header.Set("X-Panel-Domain", localDomain) + } + req.Header.Set("Content-Type", "application/json") + + resp, err := c.client.Do(req) + if err != nil { + return err + } + defer resp.Body.Close() + + if resp.StatusCode != 200 { + body, _ := io.ReadAll(resp.Body) + return fmt.Errorf("remote error %d: %s", resp.StatusCode, string(body)) + } + + var res struct { + Code int `json:"code"` + Msg string `json:"msg"` + } + if err := json.NewDecoder(resp.Body).Decode(&res); err != nil { + return err + } + if res.Code != 0 { + return fmt.Errorf("remote api error: %s", res.Msg) + } + + return nil +} diff --git a/go-backend/internal/http/handler/control_plane.go b/go-backend/internal/http/handler/control_plane.go index 75452eb..c17cab3 100644 --- a/go-backend/internal/http/handler/control_plane.go +++ b/go-backend/internal/http/handler/control_plane.go @@ -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 = "[::]" } diff --git a/go-backend/internal/http/handler/federation.go b/go-backend/internal/http/handler/federation.go index eeb6f3e..2a741a9 100644 --- a/go-backend/internal/http/handler/federation.go +++ b/go-backend/internal/http/handler/federation.go @@ -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, " ") diff --git a/go-backend/internal/http/handler/handler.go b/go-backend/internal/http/handler/handler.go index f0bedad..d63ecb8 100644 --- a/go-backend/internal/http/handler/handler.go +++ b/go-backend/internal/http/handler/handler.go @@ -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) diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index 6324371..9784d04 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -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 diff --git a/go-backend/internal/http/middleware/auth.go b/go-backend/internal/http/middleware/auth.go index a102aab..93e8524 100644 --- a/go-backend/internal/http/middleware/auth.go +++ b/go-backend/internal/http/middleware/auth.go @@ -85,6 +85,12 @@ func shouldSkip(path string) bool { return true case path == "/api/v1/federation/tunnel/create": return true + case path == "/api/v1/federation/runtime/reserve-port": + return true + case path == "/api/v1/federation/runtime/apply-role": + return true + case path == "/api/v1/federation/runtime/release-role": + return true default: return false } diff --git a/go-backend/internal/store/sqlite/repository.go b/go-backend/internal/store/sqlite/repository.go index 39cab75..66b8db2 100644 --- a/go-backend/internal/store/sqlite/repository.go +++ b/go-backend/internal/store/sqlite/repository.go @@ -124,6 +124,41 @@ type PeerShare struct { AllowedDomains string `json:"allowedDomains"` } +type PeerShareRuntime struct { + ID int64 + ShareID int64 + NodeID int64 + ReservationID string + ResourceKey string + BindingID string + Role string + ChainName string + ServiceName string + Protocol string + Strategy string + Port int + Target string + Applied int + Status int + CreatedTime int64 + UpdatedTime int64 +} + +type FederationTunnelBinding struct { + ID int64 + TunnelID int64 + NodeID int64 + ChainType int + HopInx int + RemoteURL string + ResourceKey string + RemoteBindingID string + AllocatedPort int + Status int + CreatedTime int64 + UpdatedTime int64 +} + func Open(path string) (*Repository, error) { if err := ensureParentDir(path); err != nil { return nil, err @@ -1342,6 +1377,186 @@ func (r *Repository) ListPeerShares() ([]PeerShare, error) { return shares, nil } +func (r *Repository) GetPeerShareRuntimeByResourceKey(shareID int64, resourceKey string) (*PeerShareRuntime, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + row := r.db.QueryRow(` + SELECT id, share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time + FROM peer_share_runtime + WHERE share_id = ? AND resource_key = ? + LIMIT 1 + `, shareID, resourceKey) + var item PeerShareRuntime + if err := row.Scan(&item.ID, &item.ShareID, &item.NodeID, &item.ReservationID, &item.ResourceKey, &item.BindingID, &item.Role, &item.ChainName, &item.ServiceName, &item.Protocol, &item.Strategy, &item.Port, &item.Target, &item.Applied, &item.Status, &item.CreatedTime, &item.UpdatedTime); err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return nil, err + } + return &item, nil +} + +func (r *Repository) GetPeerShareRuntimeByReservationID(shareID int64, reservationID string) (*PeerShareRuntime, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + row := r.db.QueryRow(` + SELECT id, share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time + FROM peer_share_runtime + WHERE share_id = ? AND reservation_id = ? + LIMIT 1 + `, shareID, reservationID) + var item PeerShareRuntime + if err := row.Scan(&item.ID, &item.ShareID, &item.NodeID, &item.ReservationID, &item.ResourceKey, &item.BindingID, &item.Role, &item.ChainName, &item.ServiceName, &item.Protocol, &item.Strategy, &item.Port, &item.Target, &item.Applied, &item.Status, &item.CreatedTime, &item.UpdatedTime); err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return nil, err + } + return &item, nil +} + +func (r *Repository) GetPeerShareRuntimeByBindingID(shareID int64, bindingID string) (*PeerShareRuntime, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + row := r.db.QueryRow(` + SELECT id, share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time + FROM peer_share_runtime + WHERE share_id = ? AND binding_id = ? + LIMIT 1 + `, shareID, bindingID) + var item PeerShareRuntime + if err := row.Scan(&item.ID, &item.ShareID, &item.NodeID, &item.ReservationID, &item.ResourceKey, &item.BindingID, &item.Role, &item.ChainName, &item.ServiceName, &item.Protocol, &item.Strategy, &item.Port, &item.Target, &item.Applied, &item.Status, &item.CreatedTime, &item.UpdatedTime); err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return nil, err + } + return &item, nil +} + +func (r *Repository) CreatePeerShareRuntime(item *PeerShareRuntime) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + if item == nil { + return errors.New("runtime item is nil") + } + _, err := r.db.Exec(` + INSERT INTO peer_share_runtime(share_id, node_id, reservation_id, resource_key, binding_id, role, chain_name, service_name, protocol, strategy, port, target, applied, status, created_time, updated_time) + VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + `, item.ShareID, item.NodeID, item.ReservationID, item.ResourceKey, item.BindingID, item.Role, item.ChainName, item.ServiceName, item.Protocol, item.Strategy, item.Port, item.Target, item.Applied, item.Status, item.CreatedTime, item.UpdatedTime) + return err +} + +func (r *Repository) UpdatePeerShareRuntime(item *PeerShareRuntime) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + if item == nil { + return errors.New("runtime item is nil") + } + _, err := r.db.Exec(` + UPDATE peer_share_runtime + SET binding_id = ?, role = ?, chain_name = ?, service_name = ?, protocol = ?, strategy = ?, port = ?, target = ?, applied = ?, status = ?, updated_time = ? + WHERE id = ? + `, item.BindingID, item.Role, item.ChainName, item.ServiceName, item.Protocol, item.Strategy, item.Port, item.Target, item.Applied, item.Status, item.UpdatedTime, item.ID) + return err +} + +func (r *Repository) MarkPeerShareRuntimeReleased(id int64, updatedTime int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + _, err := r.db.Exec(`UPDATE peer_share_runtime SET status = 0, updated_time = ? WHERE id = ?`, updatedTime, id) + return err +} + +func (r *Repository) ListActivePeerShareRuntimePorts(shareID int64, nodeID int64) ([]int, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + rows, err := r.db.Query(`SELECT port FROM peer_share_runtime WHERE share_id = ? AND node_id = ? AND status = 1 AND port > 0`, shareID, nodeID) + if err != nil { + return nil, err + } + defer rows.Close() + out := make([]int, 0) + for rows.Next() { + var port int + if err := rows.Scan(&port); err != nil { + return nil, err + } + if port > 0 { + out = append(out, port) + } + } + if err := rows.Err(); err != nil { + return nil, err + } + return out, nil +} + +func (r *Repository) UpsertFederationTunnelBinding(item *FederationTunnelBinding) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + if item == nil { + return errors.New("binding item is nil") + } + _, err := r.db.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(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(tunnel_id, node_id, chain_type, hop_inx) + DO UPDATE SET + remote_url = excluded.remote_url, + resource_key = excluded.resource_key, + remote_binding_id = excluded.remote_binding_id, + allocated_port = excluded.allocated_port, + status = excluded.status, + updated_time = excluded.updated_time + `, item.TunnelID, item.NodeID, item.ChainType, item.HopInx, item.RemoteURL, item.ResourceKey, item.RemoteBindingID, item.AllocatedPort, item.Status, item.CreatedTime, item.UpdatedTime) + return err +} + +func (r *Repository) ListActiveFederationTunnelBindingsByTunnel(tunnelID int64) ([]FederationTunnelBinding, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + rows, err := r.db.Query(` + SELECT id, tunnel_id, node_id, chain_type, hop_inx, remote_url, resource_key, remote_binding_id, allocated_port, status, created_time, updated_time + FROM federation_tunnel_binding + WHERE tunnel_id = ? AND status = 1 + ORDER BY chain_type ASC, hop_inx ASC, id ASC + `, tunnelID) + if err != nil { + return nil, err + } + defer rows.Close() + out := make([]FederationTunnelBinding, 0) + for rows.Next() { + var item FederationTunnelBinding + if err := rows.Scan(&item.ID, &item.TunnelID, &item.NodeID, &item.ChainType, &item.HopInx, &item.RemoteURL, &item.ResourceKey, &item.RemoteBindingID, &item.AllocatedPort, &item.Status, &item.CreatedTime, &item.UpdatedTime); err != nil { + return nil, err + } + out = append(out, item) + } + if err := rows.Err(); err != nil { + return nil, err + } + return out, nil +} + +func (r *Repository) DeleteFederationTunnelBindingsByTunnel(tunnelID int64) error { + if r == nil || r.db == nil { + return errors.New("repository not initialized") + } + _, err := r.db.Exec(`DELETE FROM federation_tunnel_binding WHERE tunnel_id = ?`, tunnelID) + return err +} + var osMkdirAll = func(path string) error { return os.MkdirAll(path, 0o755) } diff --git a/go-backend/internal/store/sqlite/sql/schema.sql b/go-backend/internal/store/sqlite/sql/schema.sql index bb8a100..74fcf0e 100644 --- a/go-backend/internal/store/sqlite/sql/schema.sql +++ b/go-backend/internal/store/sqlite/sql/schema.sql @@ -202,3 +202,43 @@ CREATE TABLE IF NOT EXISTS peer_share ( allowed_domains TEXT DEFAULT '' ); +CREATE TABLE IF NOT EXISTS peer_share_runtime ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + share_id INTEGER NOT NULL, + node_id INTEGER NOT NULL, + reservation_id TEXT NOT NULL UNIQUE, + resource_key TEXT NOT NULL UNIQUE, + binding_id TEXT NOT NULL DEFAULT '', + role TEXT NOT NULL DEFAULT '', + chain_name TEXT NOT NULL DEFAULT '', + service_name TEXT NOT NULL DEFAULT '', + protocol TEXT NOT NULL DEFAULT 'tls', + strategy TEXT NOT NULL DEFAULT 'round', + port INTEGER NOT NULL DEFAULT 0, + target TEXT NOT NULL DEFAULT '', + applied INTEGER NOT NULL DEFAULT 0, + status INTEGER NOT NULL DEFAULT 1, + created_time INTEGER NOT NULL, + updated_time INTEGER NOT NULL +); + +CREATE INDEX IF NOT EXISTS idx_peer_share_runtime_share_node_status ON peer_share_runtime(share_id, node_id, status); +CREATE INDEX IF NOT EXISTS idx_peer_share_runtime_binding_id ON peer_share_runtime(binding_id); + +CREATE TABLE IF NOT EXISTS federation_tunnel_binding ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + tunnel_id INTEGER NOT NULL, + node_id INTEGER NOT NULL, + chain_type INTEGER NOT NULL, + hop_inx INTEGER NOT NULL DEFAULT 0, + remote_url TEXT NOT NULL, + resource_key TEXT NOT NULL UNIQUE, + remote_binding_id TEXT NOT NULL, + allocated_port INTEGER NOT NULL, + status INTEGER NOT NULL DEFAULT 1, + created_time INTEGER NOT NULL, + updated_time INTEGER NOT NULL +); + +CREATE UNIQUE INDEX IF NOT EXISTS idx_federation_tunnel_binding_unique ON federation_tunnel_binding(tunnel_id, node_id, chain_type, hop_inx); +CREATE INDEX IF NOT EXISTS idx_federation_tunnel_binding_tunnel ON federation_tunnel_binding(tunnel_id, status);