feat(user): support user-group assignment in user management

This commit is contained in:
Sagit
2026-02-18 14:47:12 +00:00
parent 42701e6c01
commit d12c5bf2e1
10 changed files with 260 additions and 40 deletions
@@ -89,6 +89,7 @@ func (h *Handler) Register(mux *http.ServeMux) {
mux.HandleFunc("/api/v1/user/update", h.userUpdate)
mux.HandleFunc("/api/v1/user/delete", h.userDelete)
mux.HandleFunc("/api/v1/user/reset", h.userResetFlow)
mux.HandleFunc("/api/v1/user/groups", h.userGroups)
mux.HandleFunc("/api/v1/config/get", h.getConfigByName)
mux.HandleFunc("/api/v1/config/list", h.getConfigs)
mux.HandleFunc("/api/v1/config/update", h.updateConfigs)
+53 -1
View File
@@ -60,10 +60,21 @@ func (h *Handler) userCreate(w http.ResponseWriter, r *http.Request) {
roleID := 1
now := time.Now().UnixMilli()
if err := h.repo.CreateUser(username, security.MD5(pwd), roleID, expTime, flow, flowResetTime, num, status, now); err != nil {
userID, err := h.repo.CreateUser(username, security.MD5(pwd), roleID, expTime, flow, flowResetTime, num, status, now)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
groupIDs := asInt64Slice(req["groupIds"])
if len(groupIDs) > 0 {
if addErr := h.repo.AddUserToGroups(userID, groupIDs, now); addErr == nil {
for _, gid := range groupIDs {
_ = h.syncPermissionsByUserGroup(gid)
}
}
}
response.WriteJSON(w, response.OKEmpty())
}
@@ -133,9 +144,36 @@ func (h *Handler) userUpdate(w http.ResponseWriter, r *http.Request) {
}
h.repo.PropagateUserFlowToTunnels(id, flow, num, expTime, flowResetTime)
if groupIDsRaw, ok := req["groupIds"]; ok {
newGroupIDs := asInt64Slice(groupIDsRaw)
if affected, replaceErr := h.repo.ReplaceUserGroupsByUserID(id, newGroupIDs, now); replaceErr == nil {
for _, gid := range affected {
_ = h.syncPermissionsByUserGroup(gid)
}
}
}
response.WriteJSON(w, response.OKEmpty())
}
func (h *Handler) userGroups(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
return
}
id := idFromBody(r, w)
if id <= 0 {
return
}
ids, err := h.repo.GetUserGroupIDsByUserID(id)
if err != nil {
response.WriteJSON(w, response.Err(-2, err.Error()))
return
}
response.WriteJSON(w, response.OK(ids))
}
func (h *Handler) userDelete(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
response.WriteJSON(w, response.ErrDefault("请求失败"))
@@ -3204,3 +3242,17 @@ func randomToken(n int) string {
}
return hex.EncodeToString(buf)
}
func asInt64Slice(v interface{}) []int64 {
arr := asAnySlice(v)
if len(arr) == 0 {
return nil
}
ids := make([]int64, 0, len(arr))
for _, x := range arr {
if id := asInt64(x, 0); id > 0 {
ids = append(ids, id)
}
}
return ids
}
@@ -50,8 +50,17 @@ func (r *Repository) ListGroupPermissionPairsByUserGroup(userGroupID int64) ([][
return result, err
}
// ListGroupPermissionPairsByTunnelGroup returns [userGroupID, tunnelGroupID] pairs
// for all group permissions associated with a tunnel group.
func (r *Repository) GetUserGroupIDsByUserID(userID int64) ([]int64, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var ids []int64
err := r.db.Model(&model.UserGroupUser{}).
Where("user_id = ?", userID).
Pluck("user_group_id", &ids).Error
return ids, err
}
func (r *Repository) ListGroupPermissionPairsByTunnelGroup(tunnelGroupID int64) ([][2]int64, error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
@@ -34,9 +34,9 @@ func (r *Repository) UserExistsExcluding(username string, excludeID int64) (bool
return cnt > 0, err
}
func (r *Repository) CreateUser(username, pwdHash string, roleID int, expTime, flow, flowResetTime int64, num, status int, now int64) error {
func (r *Repository) CreateUser(username, pwdHash string, roleID int, expTime, flow, flowResetTime int64, num, status int, now int64) (int64, error) {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
return 0, errors.New("repository not initialized")
}
user := model.User{
User: username,
@@ -52,7 +52,10 @@ func (r *Repository) CreateUser(username, pwdHash string, roleID int, expTime, f
UpdatedTime: sql.NullInt64{Int64: now, Valid: true},
Status: status,
}
return r.db.Create(&user).Error
if err := r.db.Create(&user).Error; err != nil {
return 0, err
}
return user.ID, nil
}
func (r *Repository) GetUserRoleID(userID int64) (int, error) {
@@ -1406,3 +1409,70 @@ func parsePortRangeSpec(input string) []int {
sort.Ints(out)
return out
}
func (r *Repository) AddUserToGroups(userID int64, groupIDs []int64, now int64) error {
if r == nil || r.db == nil {
return errors.New("repository not initialized")
}
if len(groupIDs) == 0 {
return nil
}
rows := make([]model.UserGroupUser, 0, len(groupIDs))
for _, gid := range groupIDs {
if gid <= 0 {
continue
}
rows = append(rows, model.UserGroupUser{UserGroupID: gid, UserID: userID, CreatedTime: now})
}
if len(rows) == 0 {
return nil
}
return r.db.Clauses(clause.OnConflict{DoNothing: true}).Create(&rows).Error
}
func (r *Repository) ReplaceUserGroupsByUserID(userID int64, newGroupIDs []int64, now int64) (affectedGroupIDs []int64, err error) {
if r == nil || r.db == nil {
return nil, errors.New("repository not initialized")
}
var oldGroupIDs []int64
if err = r.db.Model(&model.UserGroupUser{}).
Where("user_id = ?", userID).
Pluck("user_group_id", &oldGroupIDs).Error; err != nil {
return nil, err
}
seen := make(map[int64]struct{})
for _, id := range oldGroupIDs {
seen[id] = struct{}{}
}
for _, id := range newGroupIDs {
if id > 0 {
seen[id] = struct{}{}
}
}
for id := range seen {
affectedGroupIDs = append(affectedGroupIDs, id)
}
if err = r.db.Where("user_id = ?", userID).Delete(&model.UserGroupUser{}).Error; err != nil {
return nil, err
}
if len(newGroupIDs) == 0 {
return affectedGroupIDs, nil
}
rows := make([]model.UserGroupUser, 0, len(newGroupIDs))
for _, gid := range newGroupIDs {
if gid <= 0 {
continue
}
rows = append(rows, model.UserGroupUser{UserGroupID: gid, UserID: userID, CreatedTime: now})
}
if len(rows) > 0 {
if err = r.db.Clauses(clause.OnConflict{DoNothing: true}).Create(&rows).Error; err != nil {
return nil, err
}
}
return affectedGroupIDs, nil
}