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
@@ -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
}
@@ -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
@@ -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
}
@@ -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)
}
@@ -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);