mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-28 23:56:36 +08:00
Compare commits
6 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| b93c259fac | |||
| 2e3d5c9249 | |||
| c8c1841058 | |||
| 1c596fae4b | |||
| 2ff52e3275 | |||
| 7efb49bdab |
@@ -223,21 +223,27 @@ func (h *Handler) listUserTunnelIDsByUser(userID int64) ([]int64, error) {
|
||||
}
|
||||
|
||||
func (h *Handler) syncForwardServices(forward *forwardRecord, method string, allowFallbackAdd bool) error {
|
||||
_, err := h.syncForwardServicesWithWarnings(forward, method, allowFallbackAdd)
|
||||
return err
|
||||
}
|
||||
|
||||
func (h *Handler) syncForwardServicesWithWarnings(forward *forwardRecord, method string, allowFallbackAdd bool) ([]string, error) {
|
||||
if h == nil || forward == nil {
|
||||
return errors.New("invalid forward sync context")
|
||||
return nil, errors.New("invalid forward sync context")
|
||||
}
|
||||
|
||||
tunnel, err := h.getTunnelRecord(forward.TunnelID)
|
||||
if err != nil {
|
||||
return err
|
||||
return nil, err
|
||||
}
|
||||
ports, err := h.listForwardPorts(forward.ID)
|
||||
if err != nil {
|
||||
return err
|
||||
return nil, err
|
||||
}
|
||||
if len(ports) == 0 {
|
||||
return errors.New("转发入口端口不存在")
|
||||
return nil, errors.New("转发入口端口不存在")
|
||||
}
|
||||
warnings := make([]string, 0)
|
||||
|
||||
// Determine limiter from forward's SpeedID first, fallback to UserTunnel's limiter
|
||||
var limiterID *int64
|
||||
@@ -258,7 +264,7 @@ func (h *Handler) syncForwardServices(forward *forwardRecord, method string, all
|
||||
var utSpeed *int
|
||||
_, utLimiterID, utSpeed, err = h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID)
|
||||
if err != nil {
|
||||
return err
|
||||
return nil, err
|
||||
}
|
||||
limiterID = utLimiterID
|
||||
speed = utSpeed
|
||||
@@ -267,28 +273,138 @@ func (h *Handler) syncForwardServices(forward *forwardRecord, method string, all
|
||||
serviceBase := buildForwardServiceBase(forward.ID, forward.UserID, 0)
|
||||
tunnelTLSProtocol, err := h.isTunnelSelectedTLSProtocol(forward.TunnelID)
|
||||
if err != nil {
|
||||
return err
|
||||
return nil, err
|
||||
}
|
||||
|
||||
for _, fp := range ports {
|
||||
if limiterID != nil && speed != nil {
|
||||
if err := h.ensureLimiterOnNode(fp.NodeID, *limiterID, *speed); err != nil {
|
||||
return err
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
node, err := h.getNodeRecord(fp.NodeID)
|
||||
if err != nil {
|
||||
return err
|
||||
return nil, err
|
||||
}
|
||||
services := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, strings.TrimSpace(fp.InIP), limiterID, tunnelTLSProtocol)
|
||||
_, err = h.sendNodeCommand(node.ID, method, services, true, false)
|
||||
if err != nil && allowFallbackAdd && method == "UpdateService" {
|
||||
_, err = h.sendNodeCommand(node.ID, "AddService", services, true, false)
|
||||
}
|
||||
if err != nil {
|
||||
return fmt.Errorf("节点 %s 下发失败: %w", node.Name, err)
|
||||
if err != nil && strings.EqualFold(strings.TrimSpace(method), "UpdateService") && isAddressAlreadyInUseError(err) {
|
||||
err = h.rebindForwardServiceOnSelfOccupiedPort(forward, node, fp.Port, services)
|
||||
}
|
||||
if err != nil && strings.EqualFold(strings.TrimSpace(method), "UpdateService") && isCannotAssignRequestedAddressError(err) {
|
||||
var warning string
|
||||
warning, err = h.fallbackForwardPortToDefaultBind(forward, tunnel, node, fp, serviceBase, limiterID, tunnelTLSProtocol)
|
||||
if err == nil && warning != "" {
|
||||
warnings = append(warnings, warning)
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
return warnings, fmt.Errorf("节点 %s 下发失败: %w", node.Name, err)
|
||||
}
|
||||
}
|
||||
return warnings, nil
|
||||
}
|
||||
|
||||
func (h *Handler) fallbackForwardPortToDefaultBind(forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, fp forwardPortRecord, serviceBase string, limiterID *int64, tunnelTLSProtocol bool) (string, error) {
|
||||
if h == nil || forward == nil || tunnel == nil || node == nil {
|
||||
return "", errors.New("invalid bind fallback context")
|
||||
}
|
||||
if fp.Port <= 0 {
|
||||
return "", errors.New("invalid forward port")
|
||||
}
|
||||
explicitBindIP := strings.TrimSpace(fp.InIP)
|
||||
if explicitBindIP == "" {
|
||||
return "", errors.New("default bind address cannot be assigned")
|
||||
}
|
||||
|
||||
if err := h.deleteForwardServicesOnNode(forward, node.ID); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
time.Sleep(150 * time.Millisecond)
|
||||
defaultServices := buildForwardServiceConfigs(serviceBase, forward, tunnel, node, fp.Port, "", limiterID, tunnelTLSProtocol)
|
||||
if _, err := h.sendNodeCommand(node.ID, "AddService", defaultServices, true, false); err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err := h.repo.UpdateForwardPortBindIP(forward.ID, node.ID, fp.Port, ""); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
warning := fmt.Sprintf("节点 %s 监听IP %s 不在主机网卡地址中,已自动回退为默认监听IP", strings.TrimSpace(node.Name), explicitBindIP)
|
||||
return warning, nil
|
||||
}
|
||||
|
||||
func (h *Handler) rebindForwardServiceOnSelfOccupiedPort(forward *forwardRecord, node *nodeRecord, port int, services []map[string]interface{}) error {
|
||||
if h == nil || forward == nil || node == nil {
|
||||
return errors.New("invalid self-occupy rebind context")
|
||||
}
|
||||
if port <= 0 {
|
||||
return errors.New("invalid forward port")
|
||||
}
|
||||
|
||||
hasOtherForward, err := h.repo.HasOtherForwardOnNodePort(node.ID, port, forward.ID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if hasOtherForward {
|
||||
return fmt.Errorf("端口 %d 已被其他转发占用", port)
|
||||
}
|
||||
|
||||
if err := h.deleteForwardServicesOnNode(forward, node.ID); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
time.Sleep(150 * time.Millisecond)
|
||||
|
||||
_, err = h.sendNodeCommand(node.ID, "AddService", services, true, false)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Handler) deleteForwardServicesOnNode(forward *forwardRecord, nodeID int64) error {
|
||||
if h == nil || forward == nil {
|
||||
return errors.New("invalid forward delete context")
|
||||
}
|
||||
|
||||
userTunnelID, _, _, err := h.resolveUserTunnelAndLimiter(forward.UserID, forward.TunnelID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
userTunnelIDs, err := h.listUserTunnelIDs(forward.UserID, forward.TunnelID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
allUserTunnelIDs, err := h.listUserTunnelIDsByUser(forward.UserID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
candidateTunnelIDs := make([]int64, 0, len(userTunnelIDs)+len(allUserTunnelIDs))
|
||||
candidateTunnelIDs = append(candidateTunnelIDs, userTunnelIDs...)
|
||||
candidateTunnelIDs = append(candidateTunnelIDs, allUserTunnelIDs...)
|
||||
bases := buildForwardServiceBaseCandidates(forward.ID, forward.UserID, userTunnelID, candidateTunnelIDs)
|
||||
|
||||
var lastErr error
|
||||
for _, base := range bases {
|
||||
names := buildForwardControlServiceNames(base, "DeleteService")
|
||||
payload := map[string]interface{}{
|
||||
"services": names,
|
||||
}
|
||||
_, cmdErr := h.sendNodeCommand(nodeID, "DeleteService", payload, false, true)
|
||||
if cmdErr == nil {
|
||||
return nil
|
||||
}
|
||||
lastErr = cmdErr
|
||||
}
|
||||
|
||||
if lastErr != nil {
|
||||
return lastErr
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -1303,6 +1419,42 @@ func isAlreadyExistsMessage(message string) bool {
|
||||
return strings.Contains(msg, "already exists") || strings.Contains(msg, "已存在")
|
||||
}
|
||||
|
||||
func isBindAddressInUseError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
msg := strings.ToLower(strings.TrimSpace(err.Error()))
|
||||
if msg == "" {
|
||||
return false
|
||||
}
|
||||
return isAddressAlreadyInUseMessage(msg) || strings.Contains(msg, "cannot assign requested address")
|
||||
}
|
||||
|
||||
func isAddressAlreadyInUseError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
return isAddressAlreadyInUseMessage(strings.ToLower(strings.TrimSpace(err.Error())))
|
||||
}
|
||||
|
||||
func isAddressAlreadyInUseMessage(msg string) bool {
|
||||
if msg == "" {
|
||||
return false
|
||||
}
|
||||
return strings.Contains(msg, "address already in use")
|
||||
}
|
||||
|
||||
func isCannotAssignRequestedAddressError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
msg := strings.ToLower(strings.TrimSpace(err.Error()))
|
||||
if msg == "" {
|
||||
return false
|
||||
}
|
||||
return strings.Contains(msg, "cannot assign requested address")
|
||||
}
|
||||
|
||||
func buildForwardServiceConfigs(baseName string, forward *forwardRecord, tunnel *tunnelRecord, node *nodeRecord, port int, bindIP string, limiterID *int64, tunnelTLSProtocol bool) []map[string]interface{} {
|
||||
protocols := []string{"tcp", "udp"}
|
||||
services := make([]map[string]interface{}, 0, 2)
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"reflect"
|
||||
"testing"
|
||||
)
|
||||
@@ -66,6 +67,39 @@ func TestIsAlreadyExistsMessage(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsBindAddressInUseError(t *testing.T) {
|
||||
if !isBindAddressInUseError(errors.New("listen tcp [::]:10001: bind: address already in use")) {
|
||||
t.Fatalf("address already in use should be detected")
|
||||
}
|
||||
if !isBindAddressInUseError(errors.New("listen tcp4 13.228.170.187:16765: bind: cannot assign requested address")) {
|
||||
t.Fatalf("cannot assign requested address should be detected")
|
||||
}
|
||||
if isBindAddressInUseError(errors.New("service demo already exists")) {
|
||||
t.Fatalf("already exists should not be treated as bind conflict")
|
||||
}
|
||||
if isBindAddressInUseError(nil) {
|
||||
t.Fatalf("nil error should not be treated as bind conflict")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsAddressAlreadyInUseError(t *testing.T) {
|
||||
if !isAddressAlreadyInUseError(errors.New("listen tcp [::]:10001: bind: address already in use")) {
|
||||
t.Fatalf("address already in use should be detected")
|
||||
}
|
||||
if isAddressAlreadyInUseError(errors.New("listen tcp4 13.228.170.187:16765: bind: cannot assign requested address")) {
|
||||
t.Fatalf("cannot assign requested address should not be treated as address-in-use")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsCannotAssignRequestedAddressError(t *testing.T) {
|
||||
if !isCannotAssignRequestedAddressError(errors.New("listen tcp4 13.228.170.187:16765: bind: cannot assign requested address")) {
|
||||
t.Fatalf("cannot assign requested address should be detected")
|
||||
}
|
||||
if isCannotAssignRequestedAddressError(errors.New("listen tcp [::]:10001: bind: address already in use")) {
|
||||
t.Fatalf("address already in use should not be treated as cannot-assign")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildForwardServiceConfigs_UsesBindIPForListen(t *testing.T) {
|
||||
forward := &forwardRecord{RemoteAddr: "1.2.3.4:80", Strategy: "fifo", TunnelID: 7}
|
||||
node := &nodeRecord{TCPListenAddr: "[::]", UDPListenAddr: "[::]"}
|
||||
|
||||
@@ -1057,7 +1057,8 @@ func (h *Handler) userTunnelUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
speedID := asAnyToInt64Ptr(req["speedId"])
|
||||
if err := h.validateSpeedLimitReference(speedID); err != nil {
|
||||
speedID, err := h.normalizeSpeedLimitReference(speedID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
@@ -1145,16 +1146,10 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
speedID := asAnyToInt64Ptr(req["speedId"])
|
||||
if speedID != nil {
|
||||
exists, speedErr := h.repo.SpeedLimitExists(*speedID)
|
||||
if speedErr != nil {
|
||||
response.WriteJSON(w, response.Err(-2, speedErr.Error()))
|
||||
return
|
||||
}
|
||||
if !exists {
|
||||
response.WriteJSON(w, response.ErrDefault("限速规则不存在"))
|
||||
return
|
||||
}
|
||||
speedID, err = h.normalizeSpeedLimitReference(speedID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
port := asInt(req["inPort"], 0)
|
||||
if port <= 0 {
|
||||
@@ -1257,16 +1252,10 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
strategy = forward.Strategy
|
||||
}
|
||||
speedID := asAnyToInt64Ptr(req["speedId"])
|
||||
if speedID != nil {
|
||||
exists, speedErr := h.repo.SpeedLimitExists(*speedID)
|
||||
if speedErr != nil {
|
||||
response.WriteJSON(w, response.Err(-2, speedErr.Error()))
|
||||
return
|
||||
}
|
||||
if !exists {
|
||||
response.WriteJSON(w, response.ErrDefault("限速规则不存在"))
|
||||
return
|
||||
}
|
||||
speedID, err = h.normalizeSpeedLimitReference(speedID)
|
||||
if err != nil {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
newSpeedID := forward.SpeedID
|
||||
if speedID != nil {
|
||||
@@ -1325,11 +1314,16 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
|
||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||
return
|
||||
}
|
||||
if err := h.syncForwardServices(updatedForward, "UpdateService", true); err != nil {
|
||||
warnings, err := h.syncForwardServicesWithWarnings(updatedForward, "UpdateService", true)
|
||||
if err != nil {
|
||||
h.rollbackForwardMutation(forward, oldPorts)
|
||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||
return
|
||||
}
|
||||
if len(warnings) > 0 {
|
||||
response.WriteJSON(w, response.OK(map[string]interface{}{"warnings": warnings}))
|
||||
return
|
||||
}
|
||||
response.WriteJSON(w, response.OKEmpty())
|
||||
}
|
||||
|
||||
@@ -3118,7 +3112,8 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
|
||||
h.repo.GetExistingUserTunnel(userID, tunnelID)
|
||||
|
||||
speedID := asAnyToInt64Ptr(req["speedId"])
|
||||
if err := h.validateSpeedLimitReference(speedID); err != nil {
|
||||
speedID, err = h.normalizeSpeedLimitReference(speedID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -3258,20 +3253,20 @@ func (h *Handler) syncUserTunnelForwards(userID, tunnelID int64) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *Handler) validateSpeedLimitReference(speedID *int64) error {
|
||||
func (h *Handler) normalizeSpeedLimitReference(speedID *int64) (*int64, error) {
|
||||
if speedID == nil {
|
||||
return nil
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
exists, err := h.repo.SpeedLimitExists(*speedID)
|
||||
if err != nil {
|
||||
return err
|
||||
return nil, err
|
||||
}
|
||||
if !exists {
|
||||
return errors.New("限速规则不存在")
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
return nil
|
||||
return speedID, nil
|
||||
}
|
||||
|
||||
func asAnySlice(v interface{}) []interface{} {
|
||||
|
||||
@@ -111,6 +111,25 @@ func (r *Repository) ListForwardPorts(forwardID int64) ([]model.ForwardPortRecor
|
||||
return rows, nil
|
||||
}
|
||||
|
||||
func (r *Repository) HasOtherForwardOnNodePort(nodeID int64, port int, currentForwardID int64) (bool, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return false, errors.New("repository not initialized")
|
||||
}
|
||||
if nodeID <= 0 || port <= 0 {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
var count int64
|
||||
err := r.db.Model(&model.ForwardPort{}).
|
||||
Where("node_id = ? AND port = ? AND forward_id <> ?", nodeID, port, currentForwardID).
|
||||
Count(&count).Error
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
func (r *Repository) GetTunnelOutProtocol(tunnelID int64) (string, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return "", errors.New("repository not initialized")
|
||||
|
||||
@@ -720,6 +720,18 @@ func (r *Repository) ReplaceForwardPorts(forwardID int64, entries []struct {
|
||||
})
|
||||
}
|
||||
|
||||
func (r *Repository) UpdateForwardPortBindIP(forwardID, nodeID int64, port int, inIP string) error {
|
||||
if r == nil || r.db == nil {
|
||||
return errors.New("repository not initialized")
|
||||
}
|
||||
if forwardID <= 0 || nodeID <= 0 || port <= 0 {
|
||||
return nil
|
||||
}
|
||||
return r.db.Model(&model.ForwardPort{}).
|
||||
Where("forward_id = ? AND node_id = ? AND port = ?", forwardID, nodeID, port).
|
||||
Update("in_ip", sql.NullString{String: inIP, Valid: strings.TrimSpace(inIP) != ""}).Error
|
||||
}
|
||||
|
||||
func (r *Repository) RollbackForwardFields(id, userID int64, userName, name string, tunnelID int64, remoteAddr, strategy string, status int, speedID interface{}, now int64) {
|
||||
if r == nil || r.db == nil {
|
||||
return
|
||||
|
||||
@@ -480,6 +480,113 @@ func TestUserTunnelReassignmentKeepsStableID(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestUserTunnelSaveIgnoresDeletedSpeedLimitContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
now := time.Now().UnixMilli()
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user(id, user, pwd, role_id, exp_time, flow, in_flow, out_flow, flow_reset_time, num, created_time, updated_time, status)
|
||||
VALUES(101, 'user_tunnel_speed_user_a', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||
`, now, now).Error; err != nil {
|
||||
t.Fatalf("insert user a: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "user-tunnel-missing-speed-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, repo, "user-tunnel-missing-speed-tunnel")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status)
|
||||
VALUES(?, ?, NULL, NULL, ?, NULL, ?)
|
||||
`, "user-tunnel-missing-speed-limit", 2048, now, 1).Error; err != nil {
|
||||
t.Fatalf("insert speed limit: %v", err)
|
||||
}
|
||||
speedID := mustLastInsertID(t, repo, "user-tunnel-missing-speed-limit")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO user_tunnel(id, user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status)
|
||||
VALUES(31, 101, ?, ?, 999, 99999, 0, 0, 1, 2727251700000, 1)
|
||||
`, tunnelID, speedID).Error; err != nil {
|
||||
t.Fatalf("insert user_tunnel: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`DELETE FROM speed_limit WHERE id = ?`, speedID).Error; err != nil {
|
||||
t.Fatalf("delete speed limit: %v", err)
|
||||
}
|
||||
|
||||
t.Run("user tunnel update auto clears missing speed", func(t *testing.T) {
|
||||
updatePayload := map[string]interface{}{
|
||||
"id": 31,
|
||||
"flow": 99999,
|
||||
"num": 999,
|
||||
"expTime": int64(2727251700000),
|
||||
"flowResetTime": 1,
|
||||
"status": 1,
|
||||
"speedId": speedID,
|
||||
}
|
||||
updateBody, err := json.Marshal(updatePayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal update payload: %v", err)
|
||||
}
|
||||
updateReq := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/user/update", bytes.NewReader(updateBody))
|
||||
updateReq.Header.Set("Authorization", adminToken)
|
||||
updateReq.Header.Set("Content-Type", "application/json")
|
||||
updateRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(updateRes, updateReq)
|
||||
assertCode(t, updateRes, 0)
|
||||
|
||||
var updatedSpeed sql.NullInt64
|
||||
if err := repo.DB().Raw(`SELECT speed_id FROM user_tunnel WHERE id = 31`).Row().Scan(&updatedSpeed); err != nil {
|
||||
t.Fatalf("query updated user_tunnel speed_id: %v", err)
|
||||
}
|
||||
if updatedSpeed.Valid {
|
||||
t.Fatalf("expected updated user_tunnel speed_id to be NULL, got %d", updatedSpeed.Int64)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("user tunnel batch assign auto clears missing speed", func(t *testing.T) {
|
||||
if err := repo.DB().Exec(`UPDATE user_tunnel SET speed_id = ? WHERE id = 31`, speedID).Error; err != nil {
|
||||
t.Fatalf("prepare user_tunnel speed_id for batch assign: %v", err)
|
||||
}
|
||||
|
||||
assignPayload := map[string]interface{}{
|
||||
"userId": 101,
|
||||
"tunnels": []map[string]interface{}{{
|
||||
"tunnelId": tunnelID,
|
||||
"speedId": speedID,
|
||||
}},
|
||||
}
|
||||
assignBody, err := json.Marshal(assignPayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal assign payload: %v", err)
|
||||
}
|
||||
assignReq := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/user/batch-assign", bytes.NewReader(assignBody))
|
||||
assignReq.Header.Set("Authorization", adminToken)
|
||||
assignReq.Header.Set("Content-Type", "application/json")
|
||||
assignRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(assignRes, assignReq)
|
||||
assertCode(t, assignRes, 0)
|
||||
|
||||
var assignedSpeed sql.NullInt64
|
||||
if err := repo.DB().Raw(`SELECT speed_id FROM user_tunnel WHERE id = 31`).Row().Scan(&assignedSpeed); err != nil {
|
||||
t.Fatalf("query assigned user_tunnel speed_id: %v", err)
|
||||
}
|
||||
if assignedSpeed.Valid {
|
||||
t.Fatalf("expected assigned user_tunnel speed_id to be NULL, got %d", assignedSpeed.Int64)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestForwardSpeedIDWriteAndClearContracts(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
@@ -618,6 +725,105 @@ func TestForwardSpeedIDWriteAndClearContracts(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardUpdateIgnoresDeletedSpeedLimitContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
|
||||
adminToken, err := auth.GenerateToken(1, "admin_user", 0, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("generate admin token: %v", err)
|
||||
}
|
||||
|
||||
now := time.Now().UnixMilli()
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "forward-update-missing-speed-tunnel", 1.0, 1, "tls", 99999, now, now, 1, nil, 0).Error; err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
tunnelID := mustLastInsertID(t, repo, "forward-update-missing-speed-tunnel")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO node(name, secret, server_ip, server_ip_v4, server_ip_v6, port, interface_name, version, http, tls, socks, created_time, updated_time, status, tcp_listen_addr, udp_listen_addr, inx)
|
||||
VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`, "forward-update-missing-speed-node", "forward-update-missing-speed-secret", "10.32.0.1", "10.32.0.1", "", "42000-42010", "", "v1", 1, 1, 1, now, now, 1, "[::]", "[::]", 0).Error; err != nil {
|
||||
t.Fatalf("insert node: %v", err)
|
||||
}
|
||||
nodeID := mustLastInsertID(t, repo, "forward-update-missing-speed-node")
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol)
|
||||
VALUES(?, 1, ?, 42001, 'round', 1, 'tls')
|
||||
`, tunnelID, nodeID).Error; err != nil {
|
||||
t.Fatalf("insert chain_tunnel: %v", err)
|
||||
}
|
||||
|
||||
if err := repo.DB().Exec(`
|
||||
INSERT INTO speed_limit(name, speed, tunnel_id, tunnel_name, created_time, updated_time, status)
|
||||
VALUES(?, ?, NULL, NULL, ?, NULL, ?)
|
||||
`, "forward-update-missing-speed-limit", 2048, now, 1).Error; err != nil {
|
||||
t.Fatalf("insert speed limit: %v", err)
|
||||
}
|
||||
speedID := mustLastInsertID(t, repo, "forward-update-missing-speed-limit")
|
||||
|
||||
server := httptest.NewServer(router)
|
||||
defer server.Close()
|
||||
stopNode := startMockNodeSession(t, server.URL, "forward-update-missing-speed-secret")
|
||||
defer stopNode()
|
||||
|
||||
createPayload := map[string]interface{}{
|
||||
"name": "forward-update-missing-speed-target",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "1.1.1.1:443",
|
||||
"strategy": "fifo",
|
||||
"speedId": speedID,
|
||||
}
|
||||
createBody, err := json.Marshal(createPayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal create payload: %v", err)
|
||||
}
|
||||
createReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/create", bytes.NewReader(createBody))
|
||||
createReq.Header.Set("Authorization", adminToken)
|
||||
createReq.Header.Set("Content-Type", "application/json")
|
||||
createRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(createRes, createReq)
|
||||
assertCode(t, createRes, 0)
|
||||
|
||||
forwardID := mustLastInsertID(t, repo, "forward-update-missing-speed-target")
|
||||
|
||||
if err := repo.DB().Exec(`DELETE FROM speed_limit WHERE id = ?`, speedID).Error; err != nil {
|
||||
t.Fatalf("delete speed limit: %v", err)
|
||||
}
|
||||
|
||||
updatePayload := map[string]interface{}{
|
||||
"id": forwardID,
|
||||
"name": "forward-update-missing-speed-target-updated",
|
||||
"tunnelId": tunnelID,
|
||||
"remoteAddr": "1.1.1.1:443",
|
||||
"strategy": "fifo",
|
||||
"speedId": speedID,
|
||||
}
|
||||
updateBody, err := json.Marshal(updatePayload)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal update payload: %v", err)
|
||||
}
|
||||
updateReq := httptest.NewRequest(http.MethodPost, "/api/v1/forward/update", bytes.NewReader(updateBody))
|
||||
updateReq.Header.Set("Authorization", adminToken)
|
||||
updateReq.Header.Set("Content-Type", "application/json")
|
||||
updateRes := httptest.NewRecorder()
|
||||
router.ServeHTTP(updateRes, updateReq)
|
||||
assertCode(t, updateRes, 0)
|
||||
|
||||
storedSpeed := repo.DB().Raw(`SELECT speed_id FROM forward WHERE id = ?`, forwardID).Row()
|
||||
var updatedSpeed sql.NullInt64
|
||||
if err := storedSpeed.Scan(&updatedSpeed); err != nil {
|
||||
t.Fatalf("query updated forward speed_id: %v", err)
|
||||
}
|
||||
if updatedSpeed.Valid {
|
||||
t.Fatalf("expected updated speed_id to be NULL after missing speed limit, got %d", updatedSpeed.Int64)
|
||||
}
|
||||
}
|
||||
|
||||
func TestForwardCreateThenPauseResumeContract(t *testing.T) {
|
||||
secret := "contract-jwt-secret"
|
||||
router, repo := setupContractRouter(t, secret)
|
||||
|
||||
@@ -0,0 +1,11 @@
|
||||
# 004 Forward Explicit Bind Self-Occupy Release
|
||||
|
||||
## Checklist
|
||||
|
||||
- [x] Confirm current forward edit/save failure path and lock strategy: explicit bind always stays explicit.
|
||||
- [x] Add repository query to detect whether a node+port is occupied by other forwards (excluding current forward).
|
||||
- [x] Enhance forward service sync to treat address-in-use as a recoverable case when only self occupies the port.
|
||||
- [x] On self-occupy conflict, proactively delete current forward services on target node and retry AddService.
|
||||
- [x] Keep hard failure when the same node+port is occupied by other forwards.
|
||||
- [x] Add focused unit tests for new error classification helpers.
|
||||
- [x] Run focused backend tests for touched handler/repo packages.
|
||||
@@ -0,0 +1,11 @@
|
||||
# 005 Forward Invalid BindIP Fallback Default
|
||||
|
||||
## Checklist
|
||||
|
||||
- [x] Split forward service bind failures into address-in-use and cannot-assign classes.
|
||||
- [x] Keep self-occupy release/rebind only for address-in-use conflicts.
|
||||
- [x] Add fallback path for cannot-assign: switch to default listener bind and retry service creation.
|
||||
- [x] Persist fallback result to DB by clearing `forward_port.in_ip` for affected node+port.
|
||||
- [x] Return non-blocking warning in forward update response when fallback occurs.
|
||||
- [x] Show warning toast in forward edit UI while still treating operation as success.
|
||||
- [x] Run focused backend tests for touched handler/repo packages.
|
||||
@@ -0,0 +1,8 @@
|
||||
# 006 Forward Save Missing Speed Limit Auto Clear
|
||||
|
||||
## Checklist
|
||||
|
||||
- [x] Locate forward create/update speed limit validation path that blocks save when speed rule is deleted.
|
||||
- [x] Change forward save behavior to auto-clear missing `speedId` instead of returning "限速规则不存在".
|
||||
- [x] Add contract test coverage for editing a forward after its referenced speed limit is deleted.
|
||||
- [x] Run focused contract tests for forward save behavior.
|
||||
@@ -0,0 +1,8 @@
|
||||
# 007 User Tunnel Save Missing Speed Limit Auto Clear
|
||||
|
||||
## Checklist
|
||||
|
||||
- [x] Locate user tunnel speed limit validation paths for assign/update flows.
|
||||
- [x] Change user tunnel save behavior to auto-clear missing `speedId` instead of failing.
|
||||
- [x] Add contract test coverage for user tunnel save when referenced speed limit is deleted.
|
||||
- [x] Run focused contract tests for user tunnel save behavior.
|
||||
@@ -0,0 +1,8 @@
|
||||
# 008 Frontend Missing Speed Limit Consistency
|
||||
|
||||
## Checklist
|
||||
|
||||
- [x] Review forward and user tunnel submit flows for missing speed limit behavior.
|
||||
- [x] Make frontend normalize deleted `speedId` to `null` before submit in both pages.
|
||||
- [x] Add consistent non-blocking warning toast when deleted speed rule is auto-cleared.
|
||||
- [x] Verify touched frontend files pass lint checks.
|
||||
@@ -1202,6 +1202,10 @@ export default function ForwardPage() {
|
||||
);
|
||||
}, [speedLimits]);
|
||||
|
||||
const speedLimitIds = useMemo(() => {
|
||||
return new Set(speedLimits.map((speedLimit) => speedLimit.id));
|
||||
}, [speedLimits]);
|
||||
|
||||
const availableSpeedLimits = useMemo(() => {
|
||||
return speedLimits.filter(
|
||||
(speedLimit) => !noLimitSpeedLimitIds.has(speedLimit.id),
|
||||
@@ -1213,7 +1217,27 @@ export default function ForwardPage() {
|
||||
return null;
|
||||
}
|
||||
|
||||
return noLimitSpeedLimitIds.has(speedId) ? null : speedId;
|
||||
if (noLimitSpeedLimitIds.has(speedId)) {
|
||||
return null;
|
||||
}
|
||||
|
||||
if (speedLimits.length > 0 && !speedLimitIds.has(speedId)) {
|
||||
return null;
|
||||
}
|
||||
|
||||
return speedId;
|
||||
};
|
||||
|
||||
const isMissingSpeedLimit = (speedId?: number | null): boolean => {
|
||||
if (speedId === null || speedId === undefined) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (speedLimits.length === 0 || noLimitSpeedLimitIds.has(speedId)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
return !speedLimitIds.has(speedId);
|
||||
};
|
||||
|
||||
const selectedSpeedId = normalizeSpeedId(form.speedId);
|
||||
@@ -1388,6 +1412,8 @@ export default function ForwardPage() {
|
||||
const addressCount = processedRemoteAddr.split(",").length;
|
||||
|
||||
let res: { code: number; msg: string };
|
||||
const normalizedSpeedId = normalizeSpeedId(form.speedId);
|
||||
const speedLimitAutoCleared = isMissingSpeedLimit(form.speedId);
|
||||
|
||||
if (isEdit) {
|
||||
// 更新时确保包含必要字段
|
||||
@@ -1400,7 +1426,7 @@ export default function ForwardPage() {
|
||||
...(inIpTouched ? { inIp: form.inIp || "" } : {}),
|
||||
remoteAddr: processedRemoteAddr,
|
||||
strategy: addressCount > 1 ? form.strategy : "fifo",
|
||||
speedId: normalizeSpeedId(form.speedId),
|
||||
speedId: normalizedSpeedId,
|
||||
};
|
||||
|
||||
res = await updateForward(updateData);
|
||||
@@ -1412,13 +1438,33 @@ export default function ForwardPage() {
|
||||
inIp: form.inIp || undefined,
|
||||
remoteAddr: processedRemoteAddr,
|
||||
strategy: addressCount > 1 ? form.strategy : "fifo",
|
||||
speedId: normalizeSpeedId(form.speedId),
|
||||
speedId: normalizedSpeedId,
|
||||
};
|
||||
|
||||
res = await createForward(createData);
|
||||
}
|
||||
|
||||
if (res.code === 0) {
|
||||
const warningItems = Array.isArray((res as any).data?.warnings)
|
||||
? (res as any).data.warnings
|
||||
.map((item: unknown) =>
|
||||
typeof item === "string" ? item.trim() : "",
|
||||
)
|
||||
.filter((item: string) => item)
|
||||
: [];
|
||||
|
||||
warningItems.forEach((warning: string) => {
|
||||
toast(warning, {
|
||||
icon: "⚠️",
|
||||
duration: 5000,
|
||||
});
|
||||
});
|
||||
if (speedLimitAutoCleared) {
|
||||
toast("所选限速规则不存在,已自动清除为不限速", {
|
||||
icon: "⚠️",
|
||||
duration: 5000,
|
||||
});
|
||||
}
|
||||
toast.success(isEdit ? "修改成功" : "创建成功");
|
||||
setModalOpen(false);
|
||||
loadData();
|
||||
|
||||
@@ -227,12 +227,36 @@ export default function UserPage() {
|
||||
);
|
||||
}, [speedLimits]);
|
||||
|
||||
const speedLimitIds = useMemo(() => {
|
||||
return new Set(speedLimits.map((speedLimit) => speedLimit.id));
|
||||
}, [speedLimits]);
|
||||
|
||||
const normalizeSpeedId = (speedId?: number | null): number | null => {
|
||||
if (speedId === null || speedId === undefined) {
|
||||
return null;
|
||||
}
|
||||
|
||||
return noLimitSpeedLimitIds.has(speedId) ? null : speedId;
|
||||
if (noLimitSpeedLimitIds.has(speedId)) {
|
||||
return null;
|
||||
}
|
||||
|
||||
if (speedLimits.length > 0 && !speedLimitIds.has(speedId)) {
|
||||
return null;
|
||||
}
|
||||
|
||||
return speedId;
|
||||
};
|
||||
|
||||
const isMissingSpeedLimit = (speedId?: number | null): boolean => {
|
||||
if (speedId === null || speedId === undefined) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (speedLimits.length === 0 || noLimitSpeedLimitIds.has(speedId)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
return !speedLimitIds.has(speedId);
|
||||
};
|
||||
|
||||
// 生命周期
|
||||
@@ -446,11 +470,20 @@ export default function UserPage() {
|
||||
|
||||
setAssignLoading(true);
|
||||
try {
|
||||
let speedLimitAutoCleared = false;
|
||||
const tunnelsToAssign: TunnelAssignItem[] = Array.from(
|
||||
batchTunnelSelections.entries(),
|
||||
).map(([tunnelId, speedId]) => ({
|
||||
tunnelId,
|
||||
speedId: normalizeSpeedId(speedId),
|
||||
speedId: (() => {
|
||||
const cleared = normalizeSpeedId(speedId);
|
||||
|
||||
if (isMissingSpeedLimit(speedId)) {
|
||||
speedLimitAutoCleared = true;
|
||||
}
|
||||
|
||||
return cleared;
|
||||
})(),
|
||||
}));
|
||||
|
||||
const response = await batchAssignUserTunnel({
|
||||
@@ -459,6 +492,12 @@ export default function UserPage() {
|
||||
});
|
||||
|
||||
if (response.code === 0) {
|
||||
if (speedLimitAutoCleared) {
|
||||
toast("所选限速规则不存在,已自动清除为不限速", {
|
||||
icon: "⚠️",
|
||||
duration: 5000,
|
||||
});
|
||||
}
|
||||
toast.success(response.msg || "分配成功");
|
||||
setBatchTunnelSelections(new Map());
|
||||
loadUserTunnels(currentUser.id);
|
||||
@@ -486,6 +525,7 @@ export default function UserPage() {
|
||||
|
||||
setEditTunnelLoading(true);
|
||||
try {
|
||||
const speedLimitAutoCleared = isMissingSpeedLimit(editTunnelForm.speedId);
|
||||
const response = await updateUserTunnel({
|
||||
id: editTunnelForm.id,
|
||||
flow: editTunnelForm.flow,
|
||||
@@ -497,6 +537,12 @@ export default function UserPage() {
|
||||
});
|
||||
|
||||
if (response.code === 0) {
|
||||
if (speedLimitAutoCleared) {
|
||||
toast("所选限速规则不存在,已自动清除为不限速", {
|
||||
icon: "⚠️",
|
||||
duration: 5000,
|
||||
});
|
||||
}
|
||||
toast.success("更新成功");
|
||||
onEditTunnelModalClose();
|
||||
if (currentUser) {
|
||||
|
||||
Reference in New Issue
Block a user