diff --git a/go-backend/internal/http/handler/federation.go b/go-backend/internal/http/handler/federation.go index 90c5cd9..b44ad00 100644 --- a/go-backend/internal/http/handler/federation.go +++ b/go-backend/internal/http/handler/federation.go @@ -544,6 +544,28 @@ func (h *Handler) federationRemoteUsageList(w http.ResponseWriter, r *http.Reque response.WriteJSON(w, response.OK(items)) } +func remoteNodePortRange(node *nodeRecord) (int, int) { + if node == nil || node.IsRemote != 1 || node.RemoteConfig == "" { + return 0, 0 + } + _, _, _, _, portRangeStart, portRangeEnd := parseRemoteShareUsageConfig(node.RemoteConfig) + return portRangeStart, portRangeEnd +} + +func validateRemoteNodePort(node *nodeRecord, port int) error { + if node == nil || node.IsRemote != 1 || port <= 0 { + return nil + } + start, end := remoteNodePortRange(node) + if start <= 0 || end <= 0 { + return nil + } + if port < start || port > end { + return fmt.Errorf("远程节点端口 %d 超出允许范围 %d-%d", port, start, end) + } + return nil +} + func parseRemoteShareUsageConfig(raw string) (int64, int64, int64, int64, int, int) { raw = strings.TrimSpace(raw) if raw == "" { @@ -980,6 +1002,13 @@ func (h *Handler) federationRuntimeApplyRole(w http.ResponseWriter, r *http.Requ return } + if share.PortRangeStart > 0 && share.PortRangeEnd > 0 && runtime.Port > 0 { + if runtime.Port < share.PortRangeStart || runtime.Port > share.PortRangeEnd { + response.WriteJSON(w, response.Err(403, fmt.Sprintf("port %d out of allowed range %d-%d", runtime.Port, share.PortRangeStart, share.PortRangeEnd))) + return + } + } + node, err := h.getNodeRecord(share.NodeID) if err != nil { response.WriteJSON(w, response.ErrDefault(err.Error())) @@ -1271,35 +1300,35 @@ func validateFederationCommandPorts(share *sqlite.PeerShare, data interface{}) e if !ok { return nil } - services, ok := dataMap["services"] - if !ok { - return nil - } - serviceList, ok := services.([]interface{}) - if !ok { - return nil - } - for _, svc := range serviceList { - svcMap, ok := svc.(map[string]interface{}) + + if services, ok := dataMap["services"]; ok { + serviceList, ok := services.([]interface{}) if !ok { - continue + return fmt.Errorf("invalid services format") } - addr, ok := svcMap["addr"].(string) - if !ok || addr == "" { - continue - } - _, portStr, err := net.SplitHostPort(addr) - if err != nil { - continue - } - port, err := strconv.Atoi(portStr) - if err != nil || port <= 0 { - continue - } - if port < share.PortRangeStart || port > share.PortRangeEnd { - return fmt.Errorf("port %d out of allowed range %d-%d", port, share.PortRangeStart, share.PortRangeEnd) + for _, svc := range serviceList { + svcMap, ok := svc.(map[string]interface{}) + if !ok { + return fmt.Errorf("invalid service entry format") + } + addr, ok := svcMap["addr"].(string) + if !ok || addr == "" { + continue + } + _, portStr, err := net.SplitHostPort(addr) + if err != nil { + return fmt.Errorf("invalid service address: %s", addr) + } + port, err := strconv.Atoi(portStr) + if err != nil || port <= 0 { + return fmt.Errorf("invalid port in service address: %s", addr) + } + if port < share.PortRangeStart || port > share.PortRangeEnd { + return fmt.Errorf("port %d out of allowed range %d-%d", port, share.PortRangeStart, share.PortRangeEnd) + } } } + return nil } diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index ba1d683..825a039 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -545,6 +545,11 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) { } if targetPort > 0 && targetAddr != "" { + inNodeRec := runtimeState.Nodes[firstNodeID] + if err := validateRemoteNodePort(inNodeRec, targetPort); err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } domainCfg, _ := h.repo.GetConfigByName("panel_domain") localDomain := "" if domainCfg != nil { @@ -1105,6 +1110,17 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) { if port <= 0 { port = 10000 } + entryNodes, _ := h.tunnelEntryNodeIDs(tunnelID) + for _, nodeID := range entryNodes { + node, nodeErr := h.getNodeRecord(nodeID) + if nodeErr != nil { + continue + } + if err := validateRemoteNodePort(node, port); err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + } now := time.Now().UnixMilli() inx := nextIndex(h.repo.DB(), "forward") var userName string @@ -1126,7 +1142,6 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.Err(-2, err.Error())) return } - entryNodes, _ := h.tunnelEntryNodeIDs(tunnelID) for _, nodeID := range entryNodes { _, _ = tx.Exec(`INSERT INTO forward_port(forward_id, node_id, port) VALUES(?, ?, ?)`, forwardID, nodeID, port) } @@ -1216,6 +1231,17 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) { port = h.pickTunnelPort(tunnelID) } } + fwdEntryNodes, _ := h.tunnelEntryNodeIDs(tunnelID) + for _, nodeID := range fwdEntryNodes { + node, nodeErr := h.getNodeRecord(nodeID) + if nodeErr != nil { + continue + } + if err := validateRemoteNodePort(node, port); err != nil { + response.WriteJSON(w, response.ErrDefault(err.Error())) + return + } + } now := time.Now().UnixMilli() _, err = h.repo.DB().Exec(` UPDATE forward SET name = ?, tunnel_id = ?, remote_addr = ?, strategy = ?, updated_time = ? WHERE id = ? @@ -1540,6 +1566,22 @@ func (h *Handler) forwardBatchChangeTunnel(w http.ResponseWriter, r *http.Reques if p <= 0 { p = h.pickTunnelPort(req.TargetTunnelID) } + bctEntryNodes, _ := h.tunnelEntryNodeIDs(req.TargetTunnelID) + portRangeOk := true + for _, nid := range bctEntryNodes { + nd, ndErr := h.getNodeRecord(nid) + if ndErr != nil { + continue + } + if validateRemoteNodePort(nd, p) != nil { + portRangeOk = false + break + } + } + if !portRangeOk { + fail++ + continue + } if err := h.replaceForwardPorts(id, req.TargetTunnelID, p); err != nil { h.rollbackForwardMutation(forward, oldPorts) fail++ @@ -2235,6 +2277,19 @@ func (h *Handler) prepareTunnelCreateState(tx *store.Tx, req map[string]interfac state.Nodes[nodeID] = node } + for _, outNode := range state.OutNodes { + if err := validateRemoteNodePort(state.Nodes[outNode.NodeID], outNode.Port); err != nil { + return nil, err + } + } + for _, hop := range state.ChainHops { + for _, chainNode := range hop { + if err := validateRemoteNodePort(state.Nodes[chainNode.NodeID], chainNode.Port); err != nil { + return nil, err + } + } + } + return state, nil }