diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index 7337c32..4b22426 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -1557,6 +1557,12 @@ func (h *Handler) userTunnelRemove(w http.ResponseWriter, r *http.Request) { if id <= 0 { return } + userID, tunnelID, lookupErr := h.repo.GetUserTunnelUserAndTunnel(id) + if lookupErr != nil { + response.WriteJSON(w, response.Err(-2, lookupErr.Error())) + return + } + h.cleanupForwardsForUserTunnel(userID, tunnelID) if err := h.repo.DeleteUserTunnel(id); err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return @@ -2483,14 +2489,18 @@ func (h *Handler) groupUserAssign(w http.ResponseWriter, r *http.Request) { response.WriteJSON(w, response.Err(-2, err.Error())) return } - if err := h.repo.RevokeGroupGrantsForRemovedUsersTx(tx, req.GroupID, previousUserIDs, req.UserIDs); err != nil { - response.WriteJSON(w, response.Err(-2, err.Error())) + revokedPairs, revokeErr := h.repo.RevokeGroupGrantsForRemovedUsersTx(tx, req.GroupID, previousUserIDs, req.UserIDs) + if revokeErr != nil { + response.WriteJSON(w, response.Err(-2, revokeErr.Error())) return } if err := tx.Commit().Error; err != nil { response.WriteJSON(w, response.Err(-2, err.Error())) return } + for _, pair := range revokedPairs { + h.cleanupForwardsForUserTunnel(pair.UserID, pair.TunnelID) + } _ = h.syncPermissionsByUserGroup(req.GroupID) response.WriteJSON(w, response.OKEmpty()) } @@ -2534,9 +2544,12 @@ func (h *Handler) groupPermissionRemove(w http.ResponseWriter, r *http.Request) response.WriteJSON(w, response.Err(-2, err.Error())) return } + var revokedPairs []repo.RevokedUserTunnelPair if exists { - if err := h.repo.RevokeGroupPermissionPairTx(tx, ug, tg); err != nil { - response.WriteJSON(w, response.Err(-2, err.Error())) + var revokeErr error + revokedPairs, revokeErr = h.repo.RevokeGroupPermissionPairTx(tx, ug, tg) + if revokeErr != nil { + response.WriteJSON(w, response.Err(-2, revokeErr.Error())) return } } @@ -2545,6 +2558,9 @@ func (h *Handler) groupPermissionRemove(w http.ResponseWriter, r *http.Request) response.WriteJSON(w, response.Err(-2, err.Error())) return } + for _, pair := range revokedPairs { + h.cleanupForwardsForUserTunnel(pair.UserID, pair.TunnelID) + } response.WriteJSON(w, response.OKEmpty()) } @@ -4000,6 +4016,27 @@ func (h *Handler) syncUserTunnelForwards(userID, tunnelID int64) error { return nil } +// cleanupForwardsForUserTunnel deletes all forwarding rules belonging to a +// specific user+tunnel pair. It notifies nodes to remove the runtime services +// first, then deletes the DB records. This is best-effort: individual failures +// do not abort the overall cleanup so that remaining forwards are still cleaned. +func (h *Handler) cleanupForwardsForUserTunnel(userID, tunnelID int64) { + if userID <= 0 || tunnelID <= 0 { + return + } + forwards, err := h.repo.ListForwardsByUserAndTunnel(userID, tunnelID) + if err != nil || len(forwards) == 0 { + return + } + for i := range forwards { + f := &forwards[i] + if f.Status == 1 { + _ = h.controlForwardServices(f, "DeleteService", true) + } + _ = h.deleteForwardByID(f.ID) + } +} + func (h *Handler) normalizeSpeedLimitReference(speedID *int64) (*int64, error) { if speedID == nil { return nil, nil diff --git a/go-backend/internal/store/repo/repository_flow.go b/go-backend/internal/store/repo/repository_flow.go index dcde7a1..603ab46 100644 --- a/go-backend/internal/store/repo/repository_flow.go +++ b/go-backend/internal/store/repo/repository_flow.go @@ -80,6 +80,37 @@ func (r *Repository) ListActiveForwardsByUserTunnel(userID, tunnelID int64) ([]m return rows, nil } +func (r *Repository) ListForwardsByUserAndTunnel(userID, tunnelID int64) ([]model.ForwardRecord, error) { + if r == nil || r.db == nil { + return nil, errors.New("repository not initialized") + } + var forwards []model.Forward + err := r.db.Where("user_id = ? AND tunnel_id = ?", userID, tunnelID).Order("id ASC").Find(&forwards).Error + if err != nil { + return nil, err + } + rows := make([]model.ForwardRecord, 0, len(forwards)) + for _, f := range forwards { + rows = append(rows, model.ForwardRecord{ + ID: f.ID, + UserID: f.UserID, + UserName: f.UserName, + Name: f.Name, + TunnelID: f.TunnelID, + RemoteAddr: f.RemoteAddr, + Strategy: f.Strategy, + Status: f.Status, + SpeedID: f.SpeedID, + }) + } + for i := range rows { + if strings.TrimSpace(rows[i].Strategy) == "" { + rows[i].Strategy = "fifo" + } + } + return rows, nil +} + func (r *Repository) GetForwardRecord(forwardID int64) (*model.ForwardRecord, error) { if r == nil || r.db == nil { return nil, errors.New("repository not initialized") diff --git a/go-backend/internal/store/repo/repository_mutations.go b/go-backend/internal/store/repo/repository_mutations.go index 8926064..b1b4876 100644 --- a/go-backend/internal/store/repo/repository_mutations.go +++ b/go-backend/internal/store/repo/repository_mutations.go @@ -997,9 +997,16 @@ func (r *Repository) DeleteGroupPermissionByIDTx(tx *gorm.DB, id int64) error { return tx.Where("id = ?", id).Delete(&model.GroupPermission{}).Error } -func (r *Repository) RevokeGroupGrantsForRemovedUsersTx(tx *gorm.DB, userGroupID int64, previousUserIDs, currentUserIDs []int64) error { +// RevokedUserTunnelPair holds the (userID, tunnelID) of a deleted user_tunnel row, +// so the handler layer can clean up associated forwarding rules. +type RevokedUserTunnelPair struct { + UserID int64 + TunnelID int64 +} + +func (r *Repository) RevokeGroupGrantsForRemovedUsersTx(tx *gorm.DB, userGroupID int64, previousUserIDs, currentUserIDs []int64) ([]RevokedUserTunnelPair, error) { if tx == nil { - return errors.New("database unavailable") + return nil, errors.New("database unavailable") } currentSet := make(map[int64]struct{}, len(currentUserIDs)) for _, uid := range currentUserIDs { @@ -1018,7 +1025,7 @@ func (r *Repository) RevokeGroupGrantsForRemovedUsersTx(tx *gorm.DB, userGroupID } } if len(removedUserIDs) == 0 { - return nil + return nil, nil } type grantRow struct { @@ -1026,6 +1033,8 @@ func (r *Repository) RevokeGroupGrantsForRemovedUsersTx(tx *gorm.DB, userGroupID CreatedByGroup int } + var revoked []RevokedUserTunnelPair + for _, userID := range removedUserIDs { var rows []grantRow if err := tx.Model(&model.GroupPermissionGrant{}). @@ -1033,7 +1042,7 @@ func (r *Repository) RevokeGroupGrantsForRemovedUsersTx(tx *gorm.DB, userGroupID Joins("JOIN user_tunnel ON user_tunnel.id = group_permission_grant.user_tunnel_id"). Where("group_permission_grant.user_group_id = ? AND user_tunnel.user_id = ?", userGroupID, userID). Find(&rows).Error; err != nil { - return err + return revoked, err } groupCreatedTunnelIDs := make(map[int64]struct{}) @@ -1046,28 +1055,32 @@ func (r *Repository) RevokeGroupGrantsForRemovedUsersTx(tx *gorm.DB, userGroupID userTunnelIDs := tx.Model(&model.UserTunnel{}).Select("id").Where("user_id = ?", userID) if err := tx.Where("user_group_id = ? AND user_tunnel_id IN (?)", userGroupID, userTunnelIDs). Delete(&model.GroupPermissionGrant{}).Error; err != nil { - return err + return revoked, err } for userTunnelID := range groupCreatedTunnelIDs { var remaining int64 if err := tx.Model(&model.GroupPermissionGrant{}).Where("user_tunnel_id = ?", userTunnelID).Count(&remaining).Error; err != nil { - return err + return revoked, err } if remaining == 0 { + var ut model.UserTunnel + if lookupErr := tx.Select("user_id", "tunnel_id").Where("id = ?", userTunnelID).First(&ut).Error; lookupErr == nil { + revoked = append(revoked, RevokedUserTunnelPair{UserID: ut.UserID, TunnelID: ut.TunnelID}) + } if err := tx.Where("id = ?", userTunnelID).Delete(&model.UserTunnel{}).Error; err != nil { - return err + return revoked, err } } } } - return nil + return revoked, nil } -func (r *Repository) RevokeGroupPermissionPairTx(tx *gorm.DB, userGroupID, tunnelGroupID int64) error { +func (r *Repository) RevokeGroupPermissionPairTx(tx *gorm.DB, userGroupID, tunnelGroupID int64) ([]RevokedUserTunnelPair, error) { if tx == nil { - return errors.New("database unavailable") + return nil, errors.New("database unavailable") } type grantRow struct { @@ -1080,7 +1093,7 @@ func (r *Repository) RevokeGroupPermissionPairTx(tx *gorm.DB, userGroupID, tunne Select("user_tunnel_id, created_by_group"). Where("user_group_id = ? AND tunnel_group_id = ?", userGroupID, tunnelGroupID). Find(&rows).Error; err != nil { - return err + return nil, err } groupCreatedTunnelIDs := make(map[int64]struct{}) @@ -1092,22 +1105,27 @@ func (r *Repository) RevokeGroupPermissionPairTx(tx *gorm.DB, userGroupID, tunne if err := tx.Where("user_group_id = ? AND tunnel_group_id = ?", userGroupID, tunnelGroupID). Delete(&model.GroupPermissionGrant{}).Error; err != nil { - return err + return nil, err } + var revoked []RevokedUserTunnelPair for userTunnelID := range groupCreatedTunnelIDs { var remaining int64 if err := tx.Model(&model.GroupPermissionGrant{}).Where("user_tunnel_id = ?", userTunnelID).Count(&remaining).Error; err != nil { - return err + return revoked, err } if remaining == 0 { + var ut model.UserTunnel + if lookupErr := tx.Select("user_id", "tunnel_id").Where("id = ?", userTunnelID).First(&ut).Error; lookupErr == nil { + revoked = append(revoked, RevokedUserTunnelPair{UserID: ut.UserID, TunnelID: ut.TunnelID}) + } if err := tx.Where("id = ?", userTunnelID).Delete(&model.UserTunnel{}).Error; err != nil { - return err + return revoked, err } } } - return nil + return revoked, nil } func (r *Repository) ReplaceFederationTunnelBindingsTx(tx *gorm.DB, tunnelID int64, bindings []FederationTunnelBinding) error {