mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-04 01:06:36 +08:00
Merge branch 'origin/main' into opencode/gentle-comet
This commit is contained in:
@@ -889,11 +889,11 @@ func firstPortFromRange(portRange string) int {
|
||||
|
||||
func (h *Handler) listChainNodesForTunnel(tunnelID int64) ([]chainNodeRecord, error) {
|
||||
rows, err := h.repo.DB().Query(`
|
||||
SELECT ct.chain_type, COALESCE(ct.inx, 0), ct.node_id, COALESCE(ct.port, 0), n.name, ct.protocol, ct.strategy
|
||||
SELECT CAST(ct.chain_type AS INTEGER), COALESCE(ct.inx, 0), ct.node_id, COALESCE(ct.port, 0), n.name, ct.protocol, ct.strategy
|
||||
FROM chain_tunnel ct
|
||||
LEFT JOIN node n ON n.id = ct.node_id
|
||||
WHERE ct.tunnel_id = ?
|
||||
ORDER BY ct.chain_type ASC, COALESCE(ct.inx, 0) ASC, ct.id ASC
|
||||
ORDER BY CAST(ct.chain_type AS INTEGER) ASC, COALESCE(ct.inx, 0) ASC, ct.id ASC
|
||||
`, tunnelID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"net/http"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/http/client"
|
||||
@@ -40,6 +41,17 @@ type resetPeerShareFlowRequest struct {
|
||||
ID int64 `json:"id"`
|
||||
}
|
||||
|
||||
type updatePeerShareRequest struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
MaxBandwidth int64 `json:"maxBandwidth"`
|
||||
ExpiryTime int64 `json:"expiryTime"`
|
||||
PortRangeStart int `json:"portRangeStart"`
|
||||
PortRangeEnd int `json:"portRangeEnd"`
|
||||
AllowedDomains string `json:"allowedDomains"`
|
||||
AllowedIPs string `json:"allowedIps"`
|
||||
}
|
||||
|
||||
type nodeImportRequest struct {
|
||||
RemoteURL string `json:"remoteUrl"`
|
||||
Token string `json:"token"`
|
||||
@@ -325,6 +337,80 @@ func (h *Handler) federationShareResetFlow(w http.ResponseWriter, r *http.Reques
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
func (h *Handler) federationShareUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("Invalid method"))
|
||||
return
|
||||
}
|
||||
|
||||
var req updatePeerShareRequest
|
||||
if err := decodeJSON(r.Body, &req); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("Invalid JSON"))
|
||||
return
|
||||
}
|
||||
if req.ID <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("Share ID is required"))
|
||||
return
|
||||
}
|
||||
|
||||
share, err := h.repo.GetPeerShare(req.ID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if share == nil {
|
||||
response.WriteJSON(w, response.ErrDefault("Share not found"))
|
||||
return
|
||||
}
|
||||
|
||||
if req.Name == "" {
|
||||
response.WriteJSON(w, response.ErrDefault("Name is required"))
|
||||
return
|
||||
}
|
||||
|
||||
if req.MaxBandwidth < 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("Max bandwidth cannot be negative"))
|
||||
return
|
||||
}
|
||||
|
||||
if req.ExpiryTime < 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("Expiry time cannot be negative"))
|
||||
return
|
||||
}
|
||||
|
||||
if req.PortRangeStart < 0 || req.PortRangeStart > 65535 || req.PortRangeEnd < 0 || req.PortRangeEnd > 65535 {
|
||||
response.WriteJSON(w, response.ErrDefault("Invalid port range"))
|
||||
return
|
||||
}
|
||||
|
||||
if req.PortRangeStart > req.PortRangeEnd {
|
||||
response.WriteJSON(w, response.ErrDefault("Port range start cannot be greater than end"))
|
||||
return
|
||||
}
|
||||
|
||||
allowedIPs, err := normalizePeerShareAllowedIPs(req.AllowedIPs)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
share.Name = req.Name
|
||||
share.MaxBandwidth = req.MaxBandwidth
|
||||
share.ExpiryTime = req.ExpiryTime
|
||||
share.PortRangeStart = req.PortRangeStart
|
||||
share.PortRangeEnd = req.PortRangeEnd
|
||||
share.AllowedDomains = req.AllowedDomains
|
||||
share.AllowedIPs = allowedIPs
|
||||
share.UpdatedTime = time.Now().UnixMilli()
|
||||
|
||||
if err := h.repo.UpdatePeerShare(share); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
func (h *Handler) federationRemoteUsageList(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("Invalid method"))
|
||||
@@ -700,21 +786,20 @@ func (h *Handler) federationTunnelCreate(w http.ResponseWriter, r *http.Request)
|
||||
defer tx.Rollback()
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
var tunnelID int64
|
||||
err = tx.QueryRow(`INSERT INTO tunnel (name, type, protocol, flow, created_time, updated_time, status, in_ip) VALUES (?, ?, ?, 0, ?, ?, 1, ?) RETURNING id`,
|
||||
tunnelID, err := tx.ExecReturningID(`INSERT INTO tunnel (name, type, protocol, flow, created_time, updated_time, status, in_ip) VALUES (?, ?, ?, 0, ?, ?, 1, ?)`,
|
||||
fmt.Sprintf("Share-%d-Port-%d", share.ID, req.RemotePort),
|
||||
tunnelType,
|
||||
req.Protocol,
|
||||
now,
|
||||
now,
|
||||
"",
|
||||
).Scan(&tunnelID)
|
||||
)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
_, err = tx.Exec(`INSERT INTO chain_tunnel (tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES (?, 1, ?, ?, 'fifo', 0, ?)`,
|
||||
_, err = tx.Exec(`INSERT INTO chain_tunnel (tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES (?, '1', ?, ?, 'fifo', 0, ?)`,
|
||||
tunnelID,
|
||||
share.NodeID,
|
||||
req.RemotePort,
|
||||
@@ -1316,6 +1401,74 @@ func isPeerIPAllowed(clientIP net.IP, whitelist string) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func (h *Handler) syncRemoteNodeStatuses(items []map[string]interface{}) {
|
||||
type remoteEntry struct {
|
||||
index int
|
||||
remoteURL string
|
||||
remoteToken string
|
||||
}
|
||||
|
||||
var remotes []remoteEntry
|
||||
for i, item := range items {
|
||||
isRemote, _ := item["isRemote"].(int)
|
||||
if isRemote != 1 {
|
||||
continue
|
||||
}
|
||||
url, _ := item["remoteUrl"].(string)
|
||||
token, _ := item["remoteToken"].(string)
|
||||
url = strings.TrimSpace(url)
|
||||
token = strings.TrimSpace(token)
|
||||
if url == "" || token == "" {
|
||||
continue
|
||||
}
|
||||
remotes = append(remotes, remoteEntry{index: i, remoteURL: url, remoteToken: token})
|
||||
}
|
||||
if len(remotes) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
localDomain := h.federationLocalDomain()
|
||||
fc := client.NewFederationClientWithTimeout(5 * time.Second)
|
||||
|
||||
type syncResult struct {
|
||||
index int
|
||||
status int
|
||||
syncError string
|
||||
}
|
||||
|
||||
results := make([]syncResult, len(remotes))
|
||||
var wg sync.WaitGroup
|
||||
for i, entry := range remotes {
|
||||
wg.Add(1)
|
||||
go func(idx int, e remoteEntry) {
|
||||
defer wg.Done()
|
||||
info, err := fc.Connect(e.remoteURL, e.remoteToken, localDomain)
|
||||
if err != nil {
|
||||
errMsg := err.Error()
|
||||
if strings.Contains(errMsg, "401") || strings.Contains(errMsg, "Invalid token") || strings.Contains(errMsg, "Unauthorized") {
|
||||
results[idx] = syncResult{index: e.index, status: 0, syncError: "provider_share_deleted"}
|
||||
} else if strings.Contains(errMsg, "403") || strings.Contains(errMsg, "Share is disabled") {
|
||||
results[idx] = syncResult{index: e.index, status: 0, syncError: "provider_share_disabled"}
|
||||
} else if strings.Contains(errMsg, "Share expired") {
|
||||
results[idx] = syncResult{index: e.index, status: 0, syncError: "provider_share_expired"}
|
||||
} else {
|
||||
results[idx] = syncResult{index: e.index, status: 0, syncError: errMsg}
|
||||
}
|
||||
} else {
|
||||
results[idx] = syncResult{index: e.index, status: info.Status, syncError: ""}
|
||||
}
|
||||
}(i, entry)
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
for _, r := range results {
|
||||
items[r.index]["status"] = r.status
|
||||
if r.syncError != "" {
|
||||
items[r.index]["syncError"] = r.syncError
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (h *Handler) cleanupPeerShareRuntimes(shareID int64) {
|
||||
if h == nil || h.repo == nil || shareID <= 0 {
|
||||
return
|
||||
|
||||
@@ -105,6 +105,10 @@ func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/v1/node/update-order", h.nodeUpdateOrder)
|
||||
mux.HandleFunc("/api/v1/node/batch-delete", h.nodeBatchDelete)
|
||||
mux.HandleFunc("/api/v1/node/check-status", h.nodeCheckStatus)
|
||||
mux.HandleFunc("/api/v1/node/upgrade", h.nodeUpgrade)
|
||||
mux.HandleFunc("/api/v1/node/batch-upgrade", h.nodeBatchUpgrade)
|
||||
mux.HandleFunc("/api/v1/node/rollback", h.nodeRollback)
|
||||
mux.HandleFunc("/api/v1/node/releases", h.listReleases)
|
||||
mux.HandleFunc("/api/v1/tunnel/list", h.tunnelList)
|
||||
mux.HandleFunc("/api/v1/tunnel/create", h.tunnelCreate)
|
||||
mux.HandleFunc("/api/v1/tunnel/get", h.tunnelGet)
|
||||
@@ -155,6 +159,7 @@ func (h *Handler) Register(mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/v1/open_api/sub_store", h.openAPISubStore)
|
||||
mux.HandleFunc("/api/v1/federation/share/list", h.federationShareList)
|
||||
mux.HandleFunc("/api/v1/federation/share/create", h.federationShareCreate)
|
||||
mux.HandleFunc("/api/v1/federation/share/update", h.federationShareUpdate)
|
||||
mux.HandleFunc("/api/v1/federation/share/delete", h.federationShareDelete)
|
||||
mux.HandleFunc("/api/v1/federation/share/reset-flow", h.federationShareResetFlow)
|
||||
mux.HandleFunc("/api/v1/federation/share/remote-usage/list", h.federationRemoteUsageList)
|
||||
@@ -320,6 +325,9 @@ func (h *Handler) nodeList(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
|
||||
h.syncRemoteNodeStatuses(items)
|
||||
|
||||
response.WriteJSON(w, response.OK(items))
|
||||
}
|
||||
|
||||
|
||||
@@ -18,6 +18,7 @@ import (
|
||||
"go-backend/internal/http/client"
|
||||
"go-backend/internal/http/response"
|
||||
"go-backend/internal/security"
|
||||
"go-backend/internal/store"
|
||||
"go-backend/internal/store/sqlite"
|
||||
)
|
||||
|
||||
@@ -413,7 +414,7 @@ func (h *Handler) nodeInstall(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
cmd := fmt.Sprintf("curl -L https://github.com/Sagit-chu/flux-panel/releases/latest/download/install.sh -o ./install.sh && chmod +x ./install.sh && ./install.sh -a %s -s %s", processServerAddress(panelAddr), secret)
|
||||
cmd := fmt.Sprintf("curl -L https://gcode.hostcentral.cc/https://github.com/Sagit-chu/flvx/releases/latest/download/install.sh -o ./install.sh && chmod +x ./install.sh && ./install.sh -a %s -s %s", processServerAddress(panelAddr), secret)
|
||||
response.WriteJSON(w, response.OK(cmd))
|
||||
}
|
||||
|
||||
@@ -559,9 +560,8 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
}
|
||||
|
||||
var tunnelID int64
|
||||
err = tx.QueryRow(`INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?) RETURNING id`,
|
||||
name, trafficRatio, typeVal, "tls", flow, now, now, status, nullableText(inIP), inx).Scan(&tunnelID)
|
||||
tunnelID, err := tx.ExecReturningID(`INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`,
|
||||
name, trafficRatio, typeVal, "tls", flow, now, now, status, nullableText(inIP), inx)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -1122,11 +1122,10 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
defer func() { _ = tx.Rollback() }()
|
||||
var forwardID int64
|
||||
err = tx.QueryRow(`
|
||||
forwardID, err := tx.ExecReturningID(`
|
||||
INSERT INTO forward(user_id, user_name, name, tunnel_id, remote_addr, strategy, in_flow, out_flow, created_time, updated_time, status, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, 0, 0, ?, ?, 1, ?) RETURNING id
|
||||
`, userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, now, inx).Scan(&forwardID)
|
||||
VALUES(?, ?, ?, ?, ?, ?, 0, 0, ?, ?, 1, ?)
|
||||
`, userID, userName, name, tunnelID, remoteAddr, defaultString(asString(req["strategy"]), "fifo"), now, now, inx)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -1590,9 +1589,8 @@ func (h *Handler) speedLimitCreate(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
now := time.Now().UnixMilli()
|
||||
speed := asInt(req["speed"], 100)
|
||||
var id int64
|
||||
err := h.repo.DB().QueryRow(`INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status) VALUES(?, ?, ?, ?, ?, ?, ?) RETURNING id`,
|
||||
name, speed, tunnelID, tunnelName, now, now, asInt(req["status"], 1)).Scan(&id)
|
||||
id, err := h.repo.DB().ExecReturningID(`INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status) VALUES(?, ?, ?, ?, ?, ?, ?)`,
|
||||
name, speed, tunnelID, tunnelName, now, now, asInt(req["status"], 1))
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -1690,7 +1688,7 @@ func (h *Handler) groupTunnelAssign(w http.ResponseWriter, r *http.Request) {
|
||||
defer func() { _ = tx.Rollback() }()
|
||||
_, _ = tx.Exec(`DELETE FROM tunnel_group_tunnel WHERE tunnel_group_id = ?`, req.GroupID)
|
||||
for _, tid := range req.TunnelIDs {
|
||||
_, _ = tx.Exec(`INSERT INTO tunnel_group_tunnel(tunnel_group_id, tunnel_id, created_time) VALUES(?, ?, ?) ON CONFLICT(tunnel_group_id, tunnel_id) DO NOTHING`, req.GroupID, tid, time.Now().UnixMilli())
|
||||
_, _ = tx.Exec(`INSERT INTO tunnel_group_tunnel(tunnel_group_id, tunnel_id, created_time) VALUES(?, ?, ?) ON CONFLICT DO NOTHING`, req.GroupID, tid, time.Now().UnixMilli())
|
||||
}
|
||||
if err := tx.Commit(); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
@@ -1717,7 +1715,7 @@ func (h *Handler) groupUserAssign(w http.ResponseWriter, r *http.Request) {
|
||||
defer func() { _ = tx.Rollback() }()
|
||||
_, _ = tx.Exec(`DELETE FROM user_group_user WHERE user_group_id = ?`, req.GroupID)
|
||||
for _, uid := range req.UserIDs {
|
||||
_, _ = tx.Exec(`INSERT INTO user_group_user(user_group_id, user_id, created_time) VALUES(?, ?, ?) ON CONFLICT(user_group_id, user_id) DO NOTHING`, req.GroupID, uid, time.Now().UnixMilli())
|
||||
_, _ = tx.Exec(`INSERT INTO user_group_user(user_group_id, user_id, created_time) VALUES(?, ?, ?) ON CONFLICT DO NOTHING`, req.GroupID, uid, time.Now().UnixMilli())
|
||||
}
|
||||
if err := tx.Commit(); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
@@ -1736,7 +1734,7 @@ func (h *Handler) groupPermissionAssign(w http.ResponseWriter, r *http.Request)
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
_, err := h.repo.DB().Exec(`INSERT INTO group_permission(user_group_id, tunnel_group_id, created_time) VALUES(?, ?, ?) ON CONFLICT(user_group_id, tunnel_group_id) DO NOTHING`, req.UserGroupID, req.TunnelGroupID, time.Now().UnixMilli())
|
||||
_, err := h.repo.DB().Exec(`INSERT INTO group_permission(user_group_id, tunnel_group_id, created_time) VALUES(?, ?, ?) ON CONFLICT DO NOTHING`, req.UserGroupID, req.TunnelGroupID, time.Now().UnixMilli())
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
@@ -1838,7 +1836,7 @@ func (h *Handler) applyGroupPermission(userGroupID, tunnelGroupID int64) error {
|
||||
if created {
|
||||
createdByGroup = 1
|
||||
}
|
||||
_, _ = db.Exec(`INSERT INTO group_permission_grant(user_group_id, tunnel_group_id, user_tunnel_id, created_by_group, created_time) VALUES(?, ?, ?, ?, ?) ON CONFLICT(user_group_id, tunnel_group_id, user_tunnel_id) DO NOTHING`,
|
||||
_, _ = db.Exec(`INSERT INTO group_permission_grant(user_group_id, tunnel_group_id, user_tunnel_id, created_by_group, created_time) VALUES(?, ?, ?, ?, ?) ON CONFLICT DO NOTHING`,
|
||||
userGroupID, tunnelGroupID, utID, createdByGroup, time.Now().UnixMilli())
|
||||
}
|
||||
}
|
||||
@@ -1869,7 +1867,7 @@ func (h *Handler) syncPermissionsByTunnelGroup(tunnelGroupID int64) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func ensureUserTunnelGrant(db *sql.DB, userID, tunnelID int64) (int64, bool, error) {
|
||||
func ensureUserTunnelGrant(db *store.DB, userID, tunnelID int64) (int64, bool, error) {
|
||||
var id int64
|
||||
err := db.QueryRow(`SELECT id FROM user_tunnel WHERE user_id = ? AND tunnel_id = ? LIMIT 1`, userID, tunnelID).Scan(&id)
|
||||
if err == nil {
|
||||
@@ -1885,15 +1883,15 @@ func ensureUserTunnelGrant(db *sql.DB, userID, tunnelID int64) (int64, bool, err
|
||||
if err := db.QueryRow(`SELECT flow, num, exp_time, flow_reset_time FROM user WHERE id = ?`, userID).Scan(&flow, &num, &expTime, &flowReset); err != nil {
|
||||
return 0, false, err
|
||||
}
|
||||
err = db.QueryRow(`INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(?, ?, NULL, ?, ?, 0, 0, ?, ?, 1) RETURNING id`,
|
||||
userID, tunnelID, num, flow, flowReset, expTime).Scan(&id)
|
||||
id, err = db.ExecReturningID(`INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(?, ?, NULL, ?, ?, 0, 0, ?, ?, 1)`,
|
||||
userID, tunnelID, num, flow, flowReset, expTime)
|
||||
if err != nil {
|
||||
return 0, false, err
|
||||
}
|
||||
return id, true, nil
|
||||
}
|
||||
|
||||
func queryInt64List(db *sql.DB, q string, args ...interface{}) ([]int64, error) {
|
||||
func queryInt64List(db *store.DB, q string, args ...interface{}) ([]int64, error) {
|
||||
rows, err := db.Query(q, args...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -1910,7 +1908,7 @@ func queryInt64List(db *sql.DB, q string, args ...interface{}) ([]int64, error)
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
func queryPairs(db *sql.DB, q string, args ...interface{}) ([][2]int64, error) {
|
||||
func queryPairs(db *store.DB, q string, args ...interface{}) ([][2]int64, error) {
|
||||
rows, err := db.Query(q, args...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -1946,7 +1944,7 @@ type tunnelCreateState struct {
|
||||
NodeIDList []int64
|
||||
}
|
||||
|
||||
func (h *Handler) prepareTunnelCreateState(tx *sql.Tx, req map[string]interface{}, tunnelType int, excludeTunnelID int64) (*tunnelCreateState, error) {
|
||||
func (h *Handler) prepareTunnelCreateState(tx *store.Tx, req map[string]interface{}, tunnelType int, excludeTunnelID int64) (*tunnelCreateState, error) {
|
||||
state := &tunnelCreateState{
|
||||
Type: tunnelType,
|
||||
InNodes: make([]tunnelRuntimeNode, 0),
|
||||
@@ -2405,7 +2403,7 @@ func (h *Handler) cleanupFederationRuntime(tunnelID int64) {
|
||||
_ = h.repo.DeleteFederationTunnelBindingsByTunnel(tunnelID)
|
||||
}
|
||||
|
||||
func replaceFederationTunnelBindingsTx(tx *sql.Tx, tunnelID int64, bindings []sqlite.FederationTunnelBinding) error {
|
||||
func replaceFederationTunnelBindingsTx(tx *store.Tx, tunnelID int64, bindings []sqlite.FederationTunnelBinding) error {
|
||||
if tx == nil {
|
||||
return errors.New("database unavailable")
|
||||
}
|
||||
@@ -2715,7 +2713,7 @@ func pickNodeAddressV6(node *nodeRecord) string {
|
||||
return strings.TrimSpace(node.ServerIP)
|
||||
}
|
||||
|
||||
func isRemoteNodeTx(tx *sql.Tx, nodeID int64) (bool, error) {
|
||||
func isRemoteNodeTx(tx *store.Tx, nodeID int64) (bool, error) {
|
||||
if tx == nil {
|
||||
return false, errors.New("database unavailable")
|
||||
}
|
||||
@@ -2732,7 +2730,7 @@ func isRemoteNodeTx(tx *sql.Tx, nodeID int64) (bool, error) {
|
||||
return isRemote == 1, nil
|
||||
}
|
||||
|
||||
func pickNodePortTx(tx *sql.Tx, nodeID int64, allocated map[int64]int, excludeTunnelID int64) (int, error) {
|
||||
func pickNodePortTx(tx *store.Tx, nodeID int64, allocated map[int64]int, excludeTunnelID int64) (int, error) {
|
||||
if tx == nil {
|
||||
return 0, errors.New("database unavailable")
|
||||
}
|
||||
@@ -2843,7 +2841,7 @@ func parsePortRangeSpec(input string) []int {
|
||||
return out
|
||||
}
|
||||
|
||||
func replaceTunnelChainsTx(tx *sql.Tx, tunnelID int64, req map[string]interface{}) error {
|
||||
func replaceTunnelChainsTx(tx *store.Tx, tunnelID int64, req map[string]interface{}) error {
|
||||
allocated := map[int64]int{}
|
||||
inNodes := asMapSlice(req["inNodeId"])
|
||||
for _, n := range inNodes {
|
||||
@@ -2851,7 +2849,7 @@ func replaceTunnelChainsTx(tx *sql.Tx, tunnelID int64, req map[string]interface{
|
||||
if nodeID <= 0 {
|
||||
continue
|
||||
}
|
||||
_, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 1, ?, NULL, NULL, 0, ?)`,
|
||||
_, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, '1', ?, NULL, NULL, 0, ?)`,
|
||||
tunnelID, nodeID, defaultString(asString(n["protocol"]), "tls"))
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -2870,7 +2868,7 @@ func replaceTunnelChainsTx(tx *sql.Tx, tunnelID int64, req map[string]interface{
|
||||
return pickErr
|
||||
}
|
||||
}
|
||||
_, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 3, ?, ?, ?, 0, ?)`,
|
||||
_, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, '3', ?, ?, ?, 0, ?)`,
|
||||
tunnelID, nodeID, port, defaultString(asString(n["strategy"]), "round"), defaultString(asString(n["protocol"]), "tls"))
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -2891,7 +2889,7 @@ func replaceTunnelChainsTx(tx *sql.Tx, tunnelID int64, req map[string]interface{
|
||||
return pickErr
|
||||
}
|
||||
}
|
||||
_, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 2, ?, ?, ?, ?, ?)`,
|
||||
_, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, '2', ?, ?, ?, ?, ?)`,
|
||||
tunnelID, nodeID, port, defaultString(asString(n["strategy"]), "round"), i+1, defaultString(asString(n["protocol"]), "tls"))
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -2977,7 +2975,7 @@ func (h *Handler) batchForwardStatus(ids []int64, status int) (int, int) {
|
||||
}
|
||||
|
||||
func (h *Handler) tunnelEntryNodeIDs(tunnelID int64) ([]int64, error) {
|
||||
rows, err := h.repo.DB().Query(`SELECT node_id FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = 1 ORDER BY inx ASC, id ASC`, tunnelID)
|
||||
rows, err := h.repo.DB().Query(`SELECT node_id FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = '1' ORDER BY inx ASC, id ASC`, tunnelID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -3464,7 +3462,7 @@ func randomToken(n int) string {
|
||||
return hex.EncodeToString(buf)
|
||||
}
|
||||
|
||||
func nextIndex(db *sql.DB, table string) int {
|
||||
func nextIndex(db *store.DB, table string) int {
|
||||
if db == nil {
|
||||
return 0
|
||||
}
|
||||
|
||||
@@ -0,0 +1,291 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"go-backend/internal/http/response"
|
||||
)
|
||||
|
||||
const (
|
||||
githubRepo = "Sagit-chu/flvx"
|
||||
githubProxy = "https://gcode.hostcentral.cc"
|
||||
githubAPIBase = "https://api.github.com"
|
||||
githubHTMLBase = "https://github.com"
|
||||
upgradeTimeout = 5 * time.Minute
|
||||
batchWorkers = 5
|
||||
)
|
||||
|
||||
func (h *Handler) nodeUpgrade(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
ID int64 `json:"id"`
|
||||
Version string `json:"version"`
|
||||
}
|
||||
if err := decodeJSON(r.Body, &req); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
if req.ID <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("节点ID无效"))
|
||||
return
|
||||
}
|
||||
|
||||
version := strings.TrimSpace(req.Version)
|
||||
if version == "" {
|
||||
var err error
|
||||
version, err = resolveLatestRelease()
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("获取最新版本失败: %v", err)))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
downloadURL := fmt.Sprintf(
|
||||
githubProxy+"/%s/%s/releases/download/%s/gost-{ARCH}",
|
||||
githubHTMLBase, githubRepo, version,
|
||||
)
|
||||
checksumURL := fmt.Sprintf(
|
||||
githubProxy+"/%s/%s/releases/download/%s/gost-{ARCH}.sha256",
|
||||
githubHTMLBase, githubRepo, version,
|
||||
)
|
||||
|
||||
result, err := h.wsServer.SendCommand(req.ID, "UpgradeAgent", map[string]interface{}{
|
||||
"downloadUrl": downloadURL,
|
||||
"checksumUrl": checksumURL,
|
||||
}, upgradeTimeout)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("升级失败: %v", err)))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(map[string]interface{}{
|
||||
"version": version,
|
||||
"message": result.Message,
|
||||
}))
|
||||
}
|
||||
|
||||
func resolveLatestRelease() (string, error) {
|
||||
client := &http.Client{
|
||||
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
||||
return http.ErrUseLastResponse
|
||||
},
|
||||
Timeout: 10 * time.Second,
|
||||
}
|
||||
|
||||
resp, err := client.Get(githubProxy + "/" + githubHTMLBase + "/" + githubRepo + "/releases/latest")
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("请求GitHub失败: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusFound && resp.StatusCode != http.StatusMovedPermanently {
|
||||
return resolveLatestReleaseAPI()
|
||||
}
|
||||
|
||||
location := resp.Header.Get("Location")
|
||||
if location == "" {
|
||||
return resolveLatestReleaseAPI()
|
||||
}
|
||||
|
||||
parts := strings.Split(location, "/")
|
||||
tag := parts[len(parts)-1]
|
||||
if tag == "" || tag == "latest" {
|
||||
return resolveLatestReleaseAPI()
|
||||
}
|
||||
|
||||
return tag, nil
|
||||
}
|
||||
|
||||
func resolveLatestReleaseAPI() (string, error) {
|
||||
client := &http.Client{Timeout: 10 * time.Second}
|
||||
resp, err := client.Get(githubAPIBase + "/repos/" + githubRepo + "/releases/latest")
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("请求GitHub API失败: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
body, _ := io.ReadAll(io.LimitReader(resp.Body, 512))
|
||||
return "", fmt.Errorf("GitHub API返回 %d: %s", resp.StatusCode, string(body))
|
||||
}
|
||||
|
||||
var release struct {
|
||||
TagName string `json:"tag_name"`
|
||||
}
|
||||
if err := json.NewDecoder(resp.Body).Decode(&release); err != nil {
|
||||
return "", fmt.Errorf("解析GitHub API响应失败: %v", err)
|
||||
}
|
||||
if strings.TrimSpace(release.TagName) == "" {
|
||||
return "", fmt.Errorf("无法从GitHub获取最新版本号")
|
||||
}
|
||||
|
||||
return release.TagName, nil
|
||||
}
|
||||
|
||||
func (h *Handler) nodeBatchUpgrade(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
IDs []int64 `json:"ids"`
|
||||
Version string `json:"version"`
|
||||
}
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
if len(req.IDs) == 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("ids不能为空"))
|
||||
return
|
||||
}
|
||||
|
||||
version := strings.TrimSpace(req.Version)
|
||||
if version == "" {
|
||||
var err error
|
||||
version, err = resolveLatestRelease()
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("获取最新版本失败: %v", err)))
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
downloadURL := fmt.Sprintf(
|
||||
githubProxy+"/%s/%s/releases/download/%s/gost-{ARCH}",
|
||||
githubHTMLBase, githubRepo, version,
|
||||
)
|
||||
checksumURL := fmt.Sprintf(
|
||||
githubProxy+"/%s/%s/releases/download/%s/gost-{ARCH}.sha256",
|
||||
githubHTMLBase, githubRepo, version,
|
||||
)
|
||||
|
||||
type upgradeResult struct {
|
||||
ID int64 `json:"id"`
|
||||
Success bool `json:"success"`
|
||||
Message string `json:"message"`
|
||||
}
|
||||
|
||||
results := make([]upgradeResult, len(req.IDs))
|
||||
sem := make(chan struct{}, batchWorkers)
|
||||
var wg sync.WaitGroup
|
||||
|
||||
for i, id := range req.IDs {
|
||||
wg.Add(1)
|
||||
go func(index int, nodeID int64) {
|
||||
defer wg.Done()
|
||||
sem <- struct{}{}
|
||||
defer func() { <-sem }()
|
||||
|
||||
result, err := h.wsServer.SendCommand(nodeID, "UpgradeAgent", map[string]interface{}{
|
||||
"downloadUrl": downloadURL,
|
||||
"checksumUrl": checksumURL,
|
||||
}, upgradeTimeout)
|
||||
if err != nil {
|
||||
results[index] = upgradeResult{ID: nodeID, Success: false, Message: err.Error()}
|
||||
return
|
||||
}
|
||||
results[index] = upgradeResult{ID: nodeID, Success: true, Message: result.Message}
|
||||
}(i, id)
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
response.WriteJSON(w, response.OK(map[string]interface{}{
|
||||
"version": version,
|
||||
"results": results,
|
||||
}))
|
||||
}
|
||||
|
||||
func (h *Handler) listReleases(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
|
||||
client := &http.Client{Timeout: 15 * time.Second}
|
||||
resp, err := client.Get(githubAPIBase + "/repos/" + githubRepo + "/releases?per_page=20")
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("获取版本列表失败: %v", err)))
|
||||
return
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
body, _ := io.ReadAll(io.LimitReader(resp.Body, 512))
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("获取版本列表失败: GitHub API返回 %d: %s", resp.StatusCode, string(body))))
|
||||
return
|
||||
}
|
||||
|
||||
var releases []struct {
|
||||
TagName string `json:"tag_name"`
|
||||
Name string `json:"name"`
|
||||
PublishedAt string `json:"published_at"`
|
||||
Prerelease bool `json:"prerelease"`
|
||||
Draft bool `json:"draft"`
|
||||
}
|
||||
if err := json.NewDecoder(resp.Body).Decode(&releases); err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("解析版本列表失败: %v", err)))
|
||||
return
|
||||
}
|
||||
|
||||
type releaseItem struct {
|
||||
Version string `json:"version"`
|
||||
Name string `json:"name"`
|
||||
PublishedAt string `json:"publishedAt"`
|
||||
Prerelease bool `json:"prerelease"`
|
||||
}
|
||||
|
||||
items := make([]releaseItem, 0, len(releases))
|
||||
for _, r := range releases {
|
||||
if r.Draft {
|
||||
continue
|
||||
}
|
||||
items = append(items, releaseItem{
|
||||
Version: r.TagName,
|
||||
Name: r.Name,
|
||||
PublishedAt: r.PublishedAt,
|
||||
Prerelease: r.Prerelease,
|
||||
})
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(items))
|
||||
}
|
||||
|
||||
func (h *Handler) nodeRollback(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
response.WriteJSON(w, response.ErrDefault("请求失败"))
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
ID int64 `json:"id"`
|
||||
}
|
||||
if err := decodeJSON(r.Body, &req); err != nil {
|
||||
response.WriteJSON(w, response.ErrDefault("请求参数错误"))
|
||||
return
|
||||
}
|
||||
if req.ID <= 0 {
|
||||
response.WriteJSON(w, response.ErrDefault("节点ID无效"))
|
||||
return
|
||||
}
|
||||
|
||||
result, err := h.wsServer.SendCommand(req.ID, "RollbackAgent", map[string]interface{}{}, 30*time.Second)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, fmt.Sprintf("回退失败: %v", err)))
|
||||
return
|
||||
}
|
||||
|
||||
response.WriteJSON(w, response.OK(map[string]interface{}{
|
||||
"message": result.Message,
|
||||
}))
|
||||
}
|
||||
Reference in New Issue
Block a user