mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-06 01:56:37 +08:00
fix: implement random port assignment with conflict check for forward creation
This commit is contained in:
@@ -7,6 +7,7 @@ import (
|
|||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"math/big"
|
||||||
"net"
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"sort"
|
"sort"
|
||||||
@@ -2399,30 +2400,129 @@ func (h *Handler) tunnelEntryNodeIDs(tunnelID int64) ([]int64, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (h *Handler) pickTunnelPort(tunnelID int64) int {
|
func (h *Handler) pickTunnelPort(tunnelID int64) int {
|
||||||
entry, _ := h.tunnelEntryNodeIDs(tunnelID)
|
entryNodes, err := h.tunnelEntryNodeIDs(tunnelID)
|
||||||
if len(entry) == 0 {
|
if err != nil || len(entryNodes) == 0 {
|
||||||
return 10000
|
return 10000
|
||||||
}
|
}
|
||||||
var portRange string
|
|
||||||
_ = h.repo.DB().QueryRow(`SELECT port FROM node WHERE id = ?`, entry[0]).Scan(&portRange)
|
var commonAvailable []int
|
||||||
if portRange == "" {
|
firstNode := true
|
||||||
return 10000
|
|
||||||
}
|
for _, nodeID := range entryNodes {
|
||||||
first := strings.Split(portRange, ",")[0]
|
var portRange string
|
||||||
first = strings.TrimSpace(first)
|
if err := h.repo.DB().QueryRow("SELECT port FROM node WHERE id = ?", nodeID).Scan(&portRange); err != nil {
|
||||||
if strings.Contains(first, "-") {
|
continue
|
||||||
parts := strings.SplitN(first, "-", 2)
|
}
|
||||||
p, _ := strconv.Atoi(strings.TrimSpace(parts[0]))
|
if portRange == "" {
|
||||||
if p > 0 {
|
portRange = "1000-65535"
|
||||||
return p
|
}
|
||||||
|
|
||||||
|
nodePorts, err := parsePorts(portRange)
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
used, err := h.getUsedPorts(nodeID)
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
var available []int
|
||||||
|
for _, p := range nodePorts {
|
||||||
|
if !used[p] {
|
||||||
|
available = append(available, p)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if firstNode {
|
||||||
|
commonAvailable = available
|
||||||
|
firstNode = false
|
||||||
|
} else {
|
||||||
|
set := make(map[int]bool)
|
||||||
|
for _, p := range available {
|
||||||
|
set[p] = true
|
||||||
|
}
|
||||||
|
var newCommon []int
|
||||||
|
for _, p := range commonAvailable {
|
||||||
|
if set[p] {
|
||||||
|
newCommon = append(newCommon, p)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
commonAvailable = newCommon
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(commonAvailable) == 0 {
|
||||||
|
break
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if p, err := strconv.Atoi(first); err == nil && p > 0 {
|
|
||||||
return p
|
if len(commonAvailable) > 0 {
|
||||||
|
idx, _ := rand.Int(rand.Reader, big.NewInt(int64(len(commonAvailable))))
|
||||||
|
return commonAvailable[idx.Int64()]
|
||||||
}
|
}
|
||||||
|
|
||||||
return 10000
|
return 10000
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (h *Handler) getUsedPorts(nodeID int64) (map[int]bool, error) {
|
||||||
|
used := make(map[int]bool)
|
||||||
|
rows, err := h.repo.DB().Query("SELECT port FROM forward_port WHERE node_id = ?", nodeID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
for rows.Next() {
|
||||||
|
var p int
|
||||||
|
if err := rows.Scan(&p); err == nil {
|
||||||
|
used[p] = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
rows2, err := h.repo.DB().Query("SELECT port FROM chain_tunnel WHERE node_id = ? AND port > 0", nodeID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer rows2.Close()
|
||||||
|
for rows2.Next() {
|
||||||
|
var p int
|
||||||
|
if err := rows2.Scan(&p); err == nil {
|
||||||
|
used[p] = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return used, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func parsePorts(portRange string) ([]int, error) {
|
||||||
|
var ports []int
|
||||||
|
parts := strings.Split(portRange, ",")
|
||||||
|
for _, part := range parts {
|
||||||
|
part = strings.TrimSpace(part)
|
||||||
|
if part == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if strings.Contains(part, "-") {
|
||||||
|
rangeParts := strings.Split(part, "-")
|
||||||
|
if len(rangeParts) != 2 {
|
||||||
|
return nil, errors.New("invalid port range format")
|
||||||
|
}
|
||||||
|
start, err1 := strconv.Atoi(strings.TrimSpace(rangeParts[0]))
|
||||||
|
end, err2 := strconv.Atoi(strings.TrimSpace(rangeParts[1]))
|
||||||
|
if err1 != nil || err2 != nil || start > end {
|
||||||
|
return nil, errors.New("invalid port range values")
|
||||||
|
}
|
||||||
|
for i := start; i <= end; i++ {
|
||||||
|
ports = append(ports, i)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
p, err := strconv.Atoi(part)
|
||||||
|
if err != nil {
|
||||||
|
return nil, errors.New("invalid port value")
|
||||||
|
}
|
||||||
|
ports = append(ports, p)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return ports, nil
|
||||||
|
}
|
||||||
|
|
||||||
func (h *Handler) replaceForwardPorts(forwardID, tunnelID int64, port int) error {
|
func (h *Handler) replaceForwardPorts(forwardID, tunnelID int64, port int) error {
|
||||||
tx, err := h.repo.DB().Begin()
|
tx, err := h.repo.DB().Begin()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
Reference in New Issue
Block a user