diff --git a/go-backend/internal/http/handler/handler.go b/go-backend/internal/http/handler/handler.go index bec3b3b..e010fda 100644 --- a/go-backend/internal/http/handler/handler.go +++ b/go-backend/internal/http/handler/handler.go @@ -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) diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index 815ab0d..de3f506 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -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 +} diff --git a/go-backend/internal/store/repo/repository_groups.go b/go-backend/internal/store/repo/repository_groups.go index 24683b3..88ee6b8 100644 --- a/go-backend/internal/store/repo/repository_groups.go +++ b/go-backend/internal/store/repo/repository_groups.go @@ -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") diff --git a/go-backend/internal/store/repo/repository_mutations.go b/go-backend/internal/store/repo/repository_mutations.go index 21e2817..6332e1e 100644 --- a/go-backend/internal/store/repo/repository_mutations.go +++ b/go-backend/internal/store/repo/repository_mutations.go @@ -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 +} diff --git a/vite-frontend/src/api/index.ts b/vite-frontend/src/api/index.ts index dc33537..5cdba72 100644 --- a/vite-frontend/src/api/index.ts +++ b/vite-frontend/src/api/index.ts @@ -131,6 +131,9 @@ export const updatePassword = (data: any) => export const resetUserFlow = (data: { id: number; type: number }) => Network.post("/user/reset", data); +export const getUserGroups = (id: number) => + Network.post("/user/groups", { id }); + // 网站配置相关接口 export const getConfigs = () => Network.post("/config/list"); export const getConfigByName = (name: string) => @@ -141,8 +144,10 @@ export const updateConfig = (name: string, value: string) => Network.post("/config/update-single", { name, value }); export const exportBackupData = () => Network.post("/backup/export"); -export const importBackupData = (data: any) => Network.post("/backup/import", data); -export const restoreBackupData = (data: any) => Network.post("/backup/restore", data); +export const importBackupData = (data: any) => + Network.post("/backup/import", data); +export const restoreBackupData = (data: any) => + Network.post("/backup/restore", data); // 验证码相关接口 export const checkCaptcha = () => Network.post("/captcha/check"); @@ -291,5 +296,7 @@ export interface AnnouncementData { enabled: number; } -export const getAnnouncement = () => Network.get("/announcement/get"); -export const updateAnnouncement = (data: AnnouncementData) => Network.post("/announcement/update", data); +export const getAnnouncement = () => + Network.get("/announcement/get"); +export const updateAnnouncement = (data: AnnouncementData) => + Network.post("/announcement/update", data); diff --git a/vite-frontend/src/pages/config.tsx b/vite-frontend/src/pages/config.tsx index 1776280..6979199 100644 --- a/vite-frontend/src/pages/config.tsx +++ b/vite-frontend/src/pages/config.tsx @@ -18,7 +18,14 @@ import { } from "@heroui/modal"; import toast from "react-hot-toast"; -import { updateConfigs, exportBackup, importBackup, getAnnouncement, updateAnnouncement, type AnnouncementData } from "@/api"; +import { + updateConfigs, + exportBackup, + importBackup, + getAnnouncement, + updateAnnouncement, + type AnnouncementData, +} from "@/api"; import { SettingsIcon } from "@/components/icons"; import { isAdmin } from "@/utils/auth"; import { @@ -501,7 +508,9 @@ export default function ConfigPage() { @@ -675,7 +684,10 @@ export default function ConfigPage() { - setAnnouncement({ ...announcement, enabled: checked ? 1 : 0 }) + setAnnouncement({ + ...announcement, + enabled: checked ? 1 : 0, + }) } > @@ -689,10 +701,10 @@ export default function ConfigPage() {