mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-29 07:56:37 +08:00
fix: resolve user_tunnel instability and service name drift
This commit is contained in:
@@ -2489,38 +2489,100 @@ func (h *Handler) upsertUserTunnel(req map[string]interface{}) error {
|
||||
}
|
||||
db := h.repo.DB()
|
||||
var existingID int64
|
||||
err := db.QueryRow(`SELECT id FROM user_tunnel WHERE user_id = ? AND tunnel_id = ? LIMIT 1`, userID, tunnelID).Scan(&existingID)
|
||||
flow := asInt64(req["flow"], -1)
|
||||
num := asInt(req["num"], -1)
|
||||
expTime := asInt64(req["expTime"], -1)
|
||||
flowReset := asInt64(req["flowResetTime"], -1)
|
||||
status := asInt(req["status"], 1)
|
||||
var currentFlow, currentNum, currentExpTime, currentFlowReset int64
|
||||
var currentSpeedID sql.NullInt64
|
||||
var currentStatus int
|
||||
|
||||
err := db.QueryRow(`
|
||||
SELECT id, flow, num, exp_time, flow_reset_time, speed_id, status
|
||||
FROM user_tunnel
|
||||
WHERE user_id = ? AND tunnel_id = ?
|
||||
LIMIT 1
|
||||
`, userID, tunnelID).Scan(&existingID, ¤tFlow, ¤tNum, ¤tExpTime, ¤tFlowReset, ¤tSpeedID, ¤tStatus)
|
||||
|
||||
speedID := asAnyToInt64Ptr(req["speedId"])
|
||||
reqFlow := asInt64(req["flow"], -1)
|
||||
reqNum := asInt(req["num"], -1)
|
||||
reqExpTime := asInt64(req["expTime"], -1)
|
||||
reqFlowReset := asInt64(req["flowResetTime"], -1)
|
||||
reqStatus := asInt(req["status"], -1)
|
||||
|
||||
if err == sql.ErrNoRows {
|
||||
if flow < 0 || num < 0 || expTime < 0 || flowReset < 0 {
|
||||
_ = db.QueryRow(`SELECT flow, num, exp_time, flow_reset_time FROM user WHERE id = ?`, userID).Scan(&flow, &num, &expTime, &flowReset)
|
||||
if reqFlow < 0 || reqNum < 0 || reqExpTime < 0 || reqFlowReset < 0 {
|
||||
var uFlow, uNum, uExp, uReset int64
|
||||
if uErr := db.QueryRow(`SELECT flow, num, exp_time, flow_reset_time FROM user WHERE id = ?`, userID).Scan(&uFlow, &uNum, &uExp, &uReset); uErr == nil {
|
||||
if reqFlow < 0 {
|
||||
reqFlow = uFlow
|
||||
}
|
||||
if reqNum < 0 {
|
||||
reqNum = int(uNum)
|
||||
}
|
||||
if reqExpTime < 0 {
|
||||
reqExpTime = uExp
|
||||
}
|
||||
if reqFlowReset < 0 {
|
||||
reqFlowReset = uReset
|
||||
}
|
||||
}
|
||||
}
|
||||
if reqFlow < 0 {
|
||||
reqFlow = 0
|
||||
}
|
||||
if reqNum < 0 {
|
||||
reqNum = 0
|
||||
}
|
||||
if reqExpTime < 0 {
|
||||
reqExpTime = time.Now().Add(365 * 24 * time.Hour).UnixMilli()
|
||||
}
|
||||
if reqFlowReset < 0 {
|
||||
reqFlowReset = 1
|
||||
}
|
||||
if reqStatus < 0 {
|
||||
reqStatus = 1
|
||||
}
|
||||
|
||||
_, err = db.Exec(`INSERT INTO user_tunnel(user_id, tunnel_id, speed_id, num, flow, in_flow, out_flow, flow_reset_time, exp_time, status) VALUES(?, ?, ?, ?, ?, 0, 0, ?, ?, ?)`,
|
||||
userID, tunnelID, nullableInt(speedID), num, flow, flowReset, expTime, status)
|
||||
userID, tunnelID, nullableInt(speedID), reqNum, reqFlow, reqFlowReset, reqExpTime, reqStatus)
|
||||
return err
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if flow < 0 {
|
||||
flow = 0
|
||||
|
||||
newFlow := currentFlow
|
||||
if reqFlow >= 0 {
|
||||
newFlow = reqFlow
|
||||
}
|
||||
if num < 0 {
|
||||
num = 0
|
||||
|
||||
newNum := int(currentNum)
|
||||
if reqNum >= 0 {
|
||||
newNum = reqNum
|
||||
}
|
||||
if expTime < 0 {
|
||||
expTime = time.Now().Add(365 * 24 * time.Hour).UnixMilli()
|
||||
|
||||
newExpTime := currentExpTime
|
||||
if reqExpTime >= 0 {
|
||||
newExpTime = reqExpTime
|
||||
}
|
||||
if flowReset < 0 {
|
||||
flowReset = 1
|
||||
|
||||
newFlowReset := currentFlowReset
|
||||
if reqFlowReset >= 0 {
|
||||
newFlowReset = reqFlowReset
|
||||
}
|
||||
|
||||
newStatus := currentStatus
|
||||
if reqStatus >= 0 {
|
||||
newStatus = reqStatus
|
||||
}
|
||||
|
||||
newSpeedID := currentSpeedID
|
||||
if speedID != nil {
|
||||
newSpeedID = sql.NullInt64{Int64: *speedID, Valid: true}
|
||||
} else if _, ok := req["speedId"]; ok {
|
||||
newSpeedID = sql.NullInt64{Valid: false}
|
||||
}
|
||||
|
||||
_, err = db.Exec(`UPDATE user_tunnel SET speed_id = ?, flow = ?, num = ?, exp_time = ?, flow_reset_time = ?, status = ? WHERE id = ?`,
|
||||
nullableInt(speedID), flow, num, expTime, flowReset, status, existingID)
|
||||
newSpeedID, newFlow, newNum, newExpTime, newFlowReset, newStatus, existingID)
|
||||
return err
|
||||
}
|
||||
|
||||
|
||||
@@ -173,6 +173,7 @@ CREATE UNIQUE INDEX IF NOT EXISTS idx_tunnel_group_tunnel_unique ON tunnel_group
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS idx_user_group_user_unique ON user_group_user(user_group_id, user_id);
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS idx_group_permission_unique ON group_permission(user_group_id, tunnel_group_id);
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS idx_group_permission_grant_unique ON group_permission_grant(user_group_id, tunnel_group_id, user_tunnel_id);
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS idx_user_tunnel_unique ON user_tunnel(user_id, tunnel_id);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS vite_config (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
|
||||
@@ -447,6 +447,89 @@ func TestForwardBatchChangeTunnelRollbackOnSyncFailure(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestUserTunnelReassignmentKeepsStableID(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(100, 'stable_user', 'pwd', 1, 2727251700000, 99999, 0, 0, 1, 99999, ?, ?, 1)
|
||||
`, now, now); err != nil {
|
||||
t.Fatalf("insert user: %v", err)
|
||||
}
|
||||
|
||||
tunnelRes, err := repo.DB().Exec(`
|
||||
INSERT INTO tunnel(name, traffic_ratio, type, protocol, flow, created_time, updated_time, status, in_ip, inx)
|
||||
VALUES('stable-tunnel', 1.0, 1, 'tls', 99999, ?, ?, 1, NULL, 0)
|
||||
`, now, now)
|
||||
if err != nil {
|
||||
t.Fatalf("insert tunnel: %v", err)
|
||||
}
|
||||
tunnelID, _ := tunnelRes.LastInsertId()
|
||||
|
||||
// 1. Assign permission (creates new user_tunnel)
|
||||
// userTunnelBatchAssign expects structure: {userId: 123, tunnels: [{tunnelId: 456, ...}]}
|
||||
assignPayload := `{"userId":100,"tunnels":[{"tunnelId":` + jsonNumber(tunnelID) + `}]}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/user/batch-assign", bytes.NewBufferString(assignPayload))
|
||||
req.Header.Set("Authorization", adminToken)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
res := httptest.NewRecorder()
|
||||
router.ServeHTTP(res, req)
|
||||
|
||||
var out response.R
|
||||
if err := json.NewDecoder(res.Body).Decode(&out); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if out.Code != 0 {
|
||||
t.Fatalf("expected code 0, got %d msg=%q", out.Code, out.Msg)
|
||||
}
|
||||
|
||||
var initialID int64
|
||||
if err := repo.DB().QueryRow(`SELECT id FROM user_tunnel WHERE user_id = 100 AND tunnel_id = ?`, tunnelID).Scan(&initialID); err != nil {
|
||||
t.Fatalf("query initial user_tunnel id: %v", err)
|
||||
}
|
||||
|
||||
// 2. Re-assign permission (should UPDATE, not INSERT)
|
||||
reassignPayload := `{"userId":100,"tunnels":[{"tunnelId":` + jsonNumber(tunnelID) + `}]}`
|
||||
req2 := httptest.NewRequest(http.MethodPost, "/api/v1/tunnel/user/batch-assign", bytes.NewBufferString(reassignPayload))
|
||||
req2.Header.Set("Authorization", adminToken)
|
||||
req2.Header.Set("Content-Type", "application/json")
|
||||
res2 := httptest.NewRecorder()
|
||||
router.ServeHTTP(res2, req2)
|
||||
|
||||
var out2 response.R
|
||||
if err := json.NewDecoder(res2.Body).Decode(&out2); err != nil {
|
||||
t.Fatalf("decode response 2: %v", err)
|
||||
}
|
||||
if out2.Code != 0 {
|
||||
t.Fatalf("expected code 0, got %d msg=%q", out2.Code, out2.Msg)
|
||||
}
|
||||
|
||||
// 3. Verify stable ID and no duplicates
|
||||
var count int
|
||||
if err := repo.DB().QueryRow(`SELECT COUNT(1) FROM user_tunnel WHERE user_id = 100 AND tunnel_id = ?`, tunnelID).Scan(&count); err != nil {
|
||||
t.Fatalf("query count: %v", err)
|
||||
}
|
||||
if count != 1 {
|
||||
t.Fatalf("expected exactly 1 user_tunnel record, got %d", count)
|
||||
}
|
||||
|
||||
var currentID int64
|
||||
if err := repo.DB().QueryRow(`SELECT id FROM user_tunnel WHERE user_id = 100 AND tunnel_id = ?`, tunnelID).Scan(¤tID); err != nil {
|
||||
t.Fatalf("query current user_tunnel: %v", err)
|
||||
}
|
||||
|
||||
if currentID != initialID {
|
||||
t.Fatalf("user_tunnel ID changed from %d to %d (unstable ID!)", initialID, currentID)
|
||||
}
|
||||
}
|
||||
|
||||
func jsonNumber(v int64) string {
|
||||
return strconv.FormatInt(v, 10)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user