feat: 管理员跨租户管理 + 注册开关

- 管理员可管理任意用户笔记(读取/修改/删除/回收站/标签/图谱/FTS 全平台)
- 普通用户仍数据隔离, 越权返回404
- 新增注册开关: 管理员后台⚙设置可开/关, 支持REGISTRATION_ENABLED环境变量
- 注册关闭时前台/登录页隐藏注册入口, 注册接口返回400
- 冒烟测试扩展到91用例全过
This commit is contained in:
Your Name
2026-08-11 13:07:56 +08:00
parent d9793300f9
commit 13c53fea0a
12 changed files with 741 additions and 89 deletions
+205 -65
View File
@@ -20,6 +20,12 @@ import (
type NoteService struct {
repo *repository.NoteRepository
pageSize int
roles RoleChecker // 用户角色判定(管理员可跨租户管理)
}
// RoleChecker 提供用户角色判断,便于解耦(由 UserService 实现)
type RoleChecker interface {
IsAdmin(userID uint) bool
}
// ErrNotFound 记录不存在(映射为 HTTP 404)
@@ -30,6 +36,16 @@ func NewNoteService(repo *repository.NoteRepository, pageSize int) *NoteService
return &NoteService{repo: repo, pageSize: pageSize}
}
// RegisterRoleChecker 注入角色判定器(用于管理员跨租户管理)
func (s *NoteService) RegisterRoleChecker(rc RoleChecker) {
s.roles = rc
}
// isAdmin 便捷判断
func (s *NoteService) isAdmin(userID uint) bool {
return s.roles != nil && s.roles.IsAdmin(userID)
}
// CreateNote 创建笔记或目录(归属当前用户)
func (s *NoteService) CreateNote(userID uint, req model.NoteCreateRequest) (*model.Note, error) {
note := &model.Note{
@@ -66,12 +82,18 @@ func (s *NoteService) CreateNote(userID uint, req model.NoteCreateRequest) (*mod
return note, nil
}
// GetNote 获取单条笔记(后台/个人管理使用,校验归属
// GetNote 获取单条笔记(后台/个人管理使用,管理员可跨租户读取任意笔记
func (s *NoteService) GetNote(userID, id uint) (*model.Note, error) {
note, err := s.repo.GetByIDScoped(userID, id)
var note *model.Note
var err error
if s.isAdmin(userID) {
note, err = s.repo.GetByID(id) // 管理员直接按 id 读取
} else {
note, err = s.repo.GetByIDScoped(userID, id)
}
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, errors.New("笔记不存在")
return nil, ErrNotFound
}
return nil, err
}
@@ -108,9 +130,15 @@ func (s *NoteService) GetNoteContent(id uint, password string) (*model.Note, boo
return note, false, nil
}
// UpdateNote 更新笔记或目录(保存更新前快照到版本历史,校验归属
// UpdateNote 更新笔记或目录(保存更新前快照到版本历史,管理员可跨租户更新
func (s *NoteService) UpdateNote(userID, id uint, req model.NoteUpdateRequest) (*model.Note, error) {
note, err := s.repo.GetByIDScoped(userID, id)
var note *model.Note
var err error
if s.isAdmin(userID) {
note, err = s.repo.GetByID(id)
} else {
note, err = s.repo.GetByIDScoped(userID, id)
}
if err != nil {
return nil, ErrNotFound
}
@@ -166,9 +194,15 @@ func (s *NoteService) UpdateNote(userID, id uint, req model.NoteUpdateRequest) (
return note, nil
}
// DeleteNote 软删除笔记或目录(目录会软删除所有子项,校验归属
// DeleteNote 软删除笔记或目录(目录会软删除所有子项,管理员可跨租户删除
func (s *NoteService) DeleteNote(userID, id uint) error {
note, err := s.repo.GetByIDScoped(userID, id)
var note *model.Note
var err error
if s.isAdmin(userID) {
note, err = s.repo.GetByID(id)
} else {
note, err = s.repo.GetByIDScoped(userID, id)
}
if err != nil {
return ErrNotFound
}
@@ -178,17 +212,23 @@ func (s *NoteService) DeleteNote(userID, id uint) error {
return s.repo.Delete(id)
}
// GetAllTree 获取所有笔记和目录的树形结构(个人管理用)
// GetAllTree 获取所有笔记和目录的树形结构(个人管理用;管理员看全平台
func (s *NoteService) GetAllTree(userID uint) ([]model.NoteListItem, error) {
if s.isAdmin(userID) {
return s.repo.GetAllTreeAll()
}
return s.repo.GetAllTree(userID)
}
// GetPublicTree 获取树形结构(登录用户看自己全部;游客看所有公开笔记)
// GetPublicTree 获取树形结构(登录用户看自己全部;游客看所有公开笔记;管理员看全平台
func (s *NoteService) GetPublicTree(userID uint, loggedIn bool) ([]model.NoteListItem, error) {
if loggedIn && s.isAdmin(userID) {
return s.repo.GetPublicTreeAll()
}
return s.repo.GetPublicTree(userID, loggedIn)
}
// ListNotes 获取笔记列表(登录用户看自己的;游客看公开笔记)
// ListNotes 获取笔记列表(登录用户看自己的;游客看公开笔记;管理员跨租户全部
func (s *NoteService) ListNotes(userID uint, loggedIn bool, pageStr, pageSizeStr, category, tag string, pinned, favorite *bool) ([]model.NoteListItem, int64, int, error) {
page := parseInt(pageStr, 1)
pageSize := parseInt(pageSizeStr, s.pageSize)
@@ -201,14 +241,15 @@ func (s *NoteService) ListNotes(userID uint, loggedIn bool, pageStr, pageSizeStr
}
items, total, err := s.repo.List(repository.ListQuery{
UserID: userID,
LoggedIn: loggedIn,
Page: page,
PageSize: pageSize,
Category: category,
Tag: tag,
Pinned: pinned,
Favorite: favorite,
UserID: userID,
LoggedIn: loggedIn,
Admin: loggedIn && s.isAdmin(userID),
Page: page,
PageSize: pageSize,
Category: category,
Tag: tag,
Pinned: pinned,
Favorite: favorite,
})
if err != nil {
return nil, 0, 0, fmt.Errorf("获取笔记列表失败: %w", err)
@@ -222,12 +263,15 @@ func (s *NoteService) ListNotes(userID uint, loggedIn bool, pageStr, pageSizeStr
return items, total, totalPages, nil
}
// GetByParentID 获取指定目录下的所有项目(个人管理用)
// GetByParentID 获取指定目录下的所有项目(个人管理用;管理员看全平台
func (s *NoteService) GetByParentID(userID, parentID uint) ([]model.NoteListItem, error) {
if s.isAdmin(userID) {
return s.repo.GetByParentIDAll(parentID)
}
return s.repo.GetByParentID(userID, parentID)
}
// SearchNotes 搜索笔记(优先使用 FTS5 全文索引,含中文分词;按用户隔离
// SearchNotes 搜索笔记(优先使用 FTS5 全文索引,含中文分词;管理员跨租户搜索
func (s *NoteService) SearchNotes(userID uint, loggedIn bool, keyword, pageStr, pageSizeStr string) ([]model.NoteListItem, int64, int, error) {
if keyword == "" {
return nil, 0, 0, errors.New("搜索关键词不能为空")
@@ -246,6 +290,20 @@ func (s *NoteService) SearchNotes(userID uint, loggedIn bool, keyword, pageStr,
keyword = sanitizeFTS5(keyword)
if loggedIn {
if s.isAdmin(userID) {
items, total, err := s.repo.FTS5SearchAll(keyword, page, pageSize)
if err != nil {
items, total, err = s.repo.SearchAll(keyword, page, pageSize)
if err != nil {
return nil, 0, 0, fmt.Errorf("搜索笔记失败: %w", err)
}
}
totalPages := int(total) / pageSize
if int(total)%pageSize > 0 {
totalPages++
}
return items, total, totalPages, nil
}
items, total, err := s.repo.FTS5Search(userID, keyword, page, pageSize)
if err != nil {
// FTS5 失败时回退到传统 LIKE 搜索
@@ -288,8 +346,11 @@ func sanitizeFTS5(q string) string {
return b.String()
}
// GetCategories 获取所有分类(登录用户看自己的;游客看公开笔记的分类)
// GetCategories 获取所有分类(登录用户看自己的;游客看公开笔记的分类;管理员看全平台
func (s *NoteService) GetCategories(userID uint, loggedIn bool) ([]string, error) {
if loggedIn && s.isAdmin(userID) {
return s.repo.GetAllCategories()
}
if loggedIn {
return s.repo.GetCategories(userID)
}
@@ -297,8 +358,11 @@ func (s *NoteService) GetCategories(userID uint, loggedIn bool) ([]string, error
return s.repo.GetPublicCategories()
}
// GetTags 获取所有标签(登录用户看自己的;游客看公开笔记的标签)
// GetTags 获取所有标签(登录用户看自己的;游客看公开笔记的标签;管理员看全平台
func (s *NoteService) GetTags(userID uint, loggedIn bool) ([]string, error) {
if loggedIn && s.isAdmin(userID) {
return s.repo.GetAllTags()
}
if loggedIn {
return s.repo.GetTags(userID)
}
@@ -307,18 +371,21 @@ func (s *NoteService) GetTags(userID uint, loggedIn bool) ([]string, error) {
// ─────────────── 回收站 ───────────────
// ListTrash 获取回收站列表(当前用户)
// ListTrash 获取回收站列表(当前用户;管理员看全平台
func (s *NoteService) ListTrash(userID uint) ([]model.NoteListItem, error) {
if s.isAdmin(userID) {
return s.repo.ListTrashAll()
}
return s.repo.ListTrash(userID)
}
// RestoreNote 从回收站恢复笔记或目录(整棵子树,校验归属)
// RestoreNote 从回收站恢复笔记或目录(整棵子树,校验归属;管理员可恢复任意
func (s *NoteService) RestoreNote(userID, id uint) error {
note, err := s.repo.GetByIDIncludingDeleted(id)
if err != nil {
return errors.New("记录不存在")
}
if note.UserID != userID {
if !s.isAdmin(userID) && note.UserID != userID {
return errors.New("无权操作该记录")
}
if note.IsFolder {
@@ -327,13 +394,13 @@ func (s *NoteService) RestoreNote(userID, id uint) error {
return s.repo.Restore(id)
}
// PurgeNote 彻底删除笔记或目录(不可恢复,校验归属)
// PurgeNote 彻底删除笔记或目录(不可恢复,校验归属;管理员可彻底删除任意
func (s *NoteService) PurgeNote(userID, id uint) error {
note, err := s.repo.GetByIDIncludingDeleted(id)
if err != nil {
return errors.New("记录不存在")
}
if note.UserID != userID {
if !s.isAdmin(userID) && note.UserID != userID {
return errors.New("无权操作该记录")
}
if note.IsFolder {
@@ -350,9 +417,15 @@ func (s *NoteService) PurgeNote(userID, id uint) error {
return nil
}
// EmptyTrash 清空回收站(当前用户)
// EmptyTrash 清空回收站(当前用户;管理员清空全平台
func (s *NoteService) EmptyTrash(userID uint) error {
trash, err := s.repo.ListTrash(userID)
var trash []model.NoteListItem
var err error
if s.isAdmin(userID) {
trash, err = s.repo.ListTrashAll()
} else {
trash, err = s.repo.ListTrash(userID)
}
if err != nil {
return err
}
@@ -366,8 +439,15 @@ func (s *NoteService) EmptyTrash(userID uint) error {
// ─────────────── 版本历史 ───────────────
// ListVersions 获取笔记版本列表(校验归属)
// ListVersions 获取笔记版本列表(校验归属;管理员可看任意
func (s *NoteService) ListVersions(userID, noteID uint) ([]model.NoteVersion, error) {
if s.isAdmin(userID) {
_, err := s.repo.GetByID(noteID)
if err != nil {
return nil, errors.New("笔记不存在")
}
return s.repo.ListVersions(noteID)
}
note, err := s.repo.GetByIDScoped(userID, noteID)
if err != nil {
return nil, errors.New("笔记不存在")
@@ -376,9 +456,15 @@ func (s *NoteService) ListVersions(userID, noteID uint) ([]model.NoteVersion, er
return s.repo.ListVersions(noteID)
}
// RestoreVersion 将笔记恢复到指定版本(校验归属)
// RestoreVersion 将笔记恢复到指定版本(校验归属;管理员可操作任意
func (s *NoteService) RestoreVersion(userID, noteID, versionID uint) (*model.Note, error) {
note, err := s.repo.GetByIDScoped(userID, noteID)
var note *model.Note
var err error
if s.isAdmin(userID) {
note, err = s.repo.GetByID(noteID)
} else {
note, err = s.repo.GetByIDScoped(userID, noteID)
}
if err != nil {
return nil, errors.New("笔记不存在")
}
@@ -386,7 +472,6 @@ func (s *NoteService) RestoreVersion(userID, noteID, versionID uint) (*model.Not
if err != nil {
return nil, errors.New("版本不存在")
}
_ = note
// 保存当前状态为历史版本(防止覆盖)
_, _ = s.repo.SaveVersion(note)
// 回滚
@@ -402,9 +487,15 @@ func (s *NoteService) RestoreVersion(userID, noteID, versionID uint) (*model.Not
// ─────────────── 分享 ───────────────
// CreateShare 创建/更新分享令牌(校验归属)
// CreateShare 创建/更新分享令牌(校验归属;管理员可操作任意
func (s *NoteService) CreateShare(userID, noteID uint, expireHours int) (*model.Note, error) {
note, err := s.repo.GetByIDScoped(userID, noteID)
var note *model.Note
var err error
if s.isAdmin(userID) {
note, err = s.repo.GetByID(noteID)
} else {
note, err = s.repo.GetByIDScoped(userID, noteID)
}
if err != nil {
return nil, errors.New("笔记不存在")
}
@@ -424,9 +515,15 @@ func (s *NoteService) CreateShare(userID, noteID uint, expireHours int) (*model.
return note, nil
}
// RevokeShare 撤销分享(校验归属)
// RevokeShare 撤销分享(校验归属;管理员可操作任意
func (s *NoteService) RevokeShare(userID, noteID uint) error {
note, err := s.repo.GetByIDScoped(userID, noteID)
var note *model.Note
var err error
if s.isAdmin(userID) {
note, err = s.repo.GetByID(noteID)
} else {
note, err = s.repo.GetByIDScoped(userID, noteID)
}
if err != nil {
return errors.New("笔记不存在")
}
@@ -461,26 +558,35 @@ func (s *NoteService) UpgradePasswordHash(id uint, password string) error {
// ─────────────── 自动保存草稿 ───────────────
// SaveDraft 保存笔记草稿(仅更新草稿字段,不触发版本历史,校验归属)
// 返回是否有未保存草稿被记录
// SaveDraft 保存笔记草稿(仅更新草稿字段,不触发版本历史,校验归属;管理员可操作任意
func (s *NoteService) SaveDraft(userID, id uint, content string) error {
note, err := s.repo.GetByIDScoped(userID, id)
var note *model.Note
var err error
if s.isAdmin(userID) {
note, err = s.repo.GetByID(id)
} else {
note, err = s.repo.GetByIDScoped(userID, id)
}
if err != nil {
return errors.New("笔记不存在")
}
if note.IsFolder {
return errors.New("目录不支持草稿")
}
_ = note
// 直接更新草稿字段,保持 updated_at 不变(避免与正文保存混淆)
return s.repo.UpdateFields(id, map[string]interface{}{"draft_content": content})
}
// ClearDraft 清除笔记草稿(保存正文成功后调用,校验归属)
// ClearDraft 清除笔记草稿(保存正文成功后调用,校验归属;管理员可操作任意
func (s *NoteService) ClearDraft(userID, id uint) error {
_, err := s.repo.GetByIDScoped(userID, id)
if err != nil {
return errors.New("笔记不存在")
if s.isAdmin(userID) {
if _, err := s.repo.GetByID(id); err != nil {
return errors.New("笔记不存在")
}
} else {
if _, err := s.repo.GetByIDScoped(userID, id); err != nil {
return errors.New("笔记不存在")
}
}
return s.repo.UpdateFields(id, map[string]interface{}{"draft_content": ""})
}
@@ -490,27 +596,31 @@ func (s *NoteService) ClearDraft(userID, id uint) error {
// wikiLinkRe 匹配笔记正文中的 [[wiki链接]] 语法
var wikiLinkRe = regexp.MustCompile(`\[\[([^\[\]|]+)(?:\|[^\[\]]*)?\]\]`)
// GetBacklinks 获取指向指定笔记的所有笔记(反向链接,校验归属)
// GetBacklinks 获取指向指定笔记的所有笔记(反向链接,校验归属;管理员可查看任意
func (s *NoteService) GetBacklinks(userID, noteID uint, title string) ([]model.NoteListItem, error) {
var n *model.Note
if title == "" {
// 若未提供标题,先查一下
if s.isAdmin(userID) {
var err error
n, err = s.repo.GetByID(noteID)
if err != nil {
return nil, errors.New("笔记不存在")
}
} else {
var err error
n, err = s.repo.GetByIDScoped(userID, noteID)
if err != nil {
return nil, errors.New("笔记不存在")
}
} else {
// 校验归属(避免越权读取他人笔记反链)
scoped, err := s.repo.GetByIDScoped(userID, noteID)
if err != nil {
return nil, errors.New("笔记不存在")
}
n = scoped
}
title = n.Title
all, err := s.repo.GetAllNotesLight(userID)
var all []model.NoteListItem
var err error
if s.isAdmin(userID) {
all, err = s.repo.GetAllNotesLightAll()
} else {
all, err = s.repo.GetAllNotesLight(userID)
}
if err != nil {
return nil, err
}
@@ -521,8 +631,13 @@ func (s *NoteService) GetBacklinks(userID, noteID uint, title string) ([]model.N
if note.ID == noteID {
continue
}
full, err := s.repo.GetByIDScoped(userID, note.ID)
if err != nil {
var full *model.Note
if s.isAdmin(userID) {
full, _ = s.repo.GetByID(note.ID)
} else {
full, _ = s.repo.GetByIDScoped(userID, note.ID)
}
if full == nil {
continue
}
if strings.Contains(full.Content, "[["+title+"]]") {
@@ -544,9 +659,19 @@ type GraphEdge struct {
Target uint `json:"target"`
}
// GetKnowledgeGraph 构建当前用户的知识图谱(节点 + [[链接]] 边)
// GetKnowledgeGraph 构建当前用户的知识图谱(节点 + [[链接]] 边;管理员看全平台
func (s *NoteService) GetKnowledgeGraph(userID uint) (map[string]interface{}, error) {
all, err := s.repo.GetAllContentLight(userID)
var all []struct {
ID uint `gorm:"column:id"`
Title string `gorm:"column:title"`
Content string `gorm:"column:content"`
}
var err error
if s.isAdmin(userID) {
all, err = s.repo.GetAllContentLightAll()
} else {
all, err = s.repo.GetAllContentLight(userID)
}
if err != nil {
return nil, err
}
@@ -587,7 +712,7 @@ func (s *NoteService) GetKnowledgeGraph(userID uint) (map[string]interface{}, er
// ─────────────── 标签管理 ───────────────
// RenameTag 重命名标签(当前用户所有含该标签的笔记同步更新)
// RenameTag 重命名标签(当前用户所有含该标签的笔记同步更新;管理员跨全平台
func (s *NoteService) RenameTag(userID uint, oldTag, newTag string) (int64, error) {
if oldTag == "" || newTag == "" {
return 0, errors.New("标签名不能为空")
@@ -595,10 +720,13 @@ func (s *NoteService) RenameTag(userID uint, oldTag, newTag string) (int64, erro
if oldTag == newTag {
return 0, nil
}
if s.isAdmin(userID) {
return s.repo.UpdateTagAllAll(oldTag, newTag)
}
return s.repo.UpdateTagAll(userID, oldTag, newTag)
}
// MergeTag 将 from 标签合并到 to 标签(current 用户,from 消失)
// MergeTag 将 from 标签合并到 to 标签(current 用户,from 消失;管理员跨全平台
func (s *NoteService) MergeTag(userID uint, from, to string) (int64, error) {
if from == "" || to == "" {
return 0, errors.New("标签名不能为空")
@@ -606,22 +734,34 @@ func (s *NoteService) MergeTag(userID uint, from, to string) (int64, error) {
if from == to {
return 0, nil
}
if s.isAdmin(userID) {
return s.repo.UpdateTagAllAll(from, to)
}
return s.repo.UpdateTagAll(userID, from, to)
}
// DeleteTag 删除指定标签(当前用户所有笔记中移除)
// DeleteTag 删除指定标签(当前用户所有笔记中移除;管理员跨全平台
func (s *NoteService) DeleteTag(userID uint, tag string) (int64, error) {
if tag == "" {
return 0, errors.New("标签名不能为空")
}
if s.isAdmin(userID) {
return s.repo.UpdateTagAllAll(tag, "")
}
return s.repo.UpdateTagAll(userID, tag, "")
}
// GetTagUsage 获取当前用户每个标签及其使用次数
// GetTagUsage 获取当前用户每个标签及其使用次数(管理员跨全平台)
func (s *NoteService) GetTagUsage(userID uint) ([]model.TagUsage, error) {
var result []model.TagUsage
counts := make(map[string]int)
all, err := s.repo.GetAllNotesLight(userID)
var all []model.NoteListItem
var err error
if s.isAdmin(userID) {
all, err = s.repo.GetAllNotesLightAll()
} else {
all, err = s.repo.GetAllNotesLight(userID)
}
if err != nil {
return nil, err
}
+51 -2
View File
@@ -11,15 +11,37 @@ import (
// UserService 用户/账号业务逻辑(多租户认证)
type UserService struct {
userRepo *repository.UserRepository
// registrationEnabled 以 env 为准的注册默认开关;DB 设置项可覆盖(运行时切换)
registrationDefault string
}
// NewUserService 创建用户服务
func NewUserService(userRepo *repository.UserRepository) *UserService {
return &UserService{userRepo: userRepo}
return &UserService{userRepo: userRepo, registrationDefault: "true"}
}
// SetRegistrationDefault 设置注册开关的默认值(来自环境变量)
func (s *UserService) SetRegistrationDefault(v string) {
s.registrationDefault = v
}
// RegistrationEnabled 当前注册开关是否开启
func (s *UserService) RegistrationEnabled() bool {
val := s.userRepo.GetSetting("registration_enabled", s.registrationDefault)
return val == "1" || val == "true"
}
// SetRegistrationEnabled 切换注册开关(持久化到 DB)
func (s *UserService) SetRegistrationEnabled(on bool) error {
v := "false"
if on {
v = "true"
}
return s.userRepo.SetSetting("registration_enabled", v)
}
// Register 注册新用户
// 说明:第一个注册的用户自动成为 admin(拥有平台管理权限);其余为普通 user。
// 说明:个注册的用户自动成为 admin(拥有平台管理权限);其余为普通 user。
// 同时把历史遗留(user_id=0)的笔记迁移给首位注册用户。
func (s *UserService) Register(username, password, displayName string) (*model.User, error) {
username = strings.TrimSpace(strings.ToLower(username))
@@ -30,6 +52,9 @@ func (s *UserService) Register(username, password, displayName string) (*model.U
if len(password) < 6 {
return nil, errors.New("密码至少 6 位")
}
if !s.RegistrationEnabled() {
return nil, errors.New("注册功能已关闭")
}
if _, err := s.userRepo.GetByUsername(username); err == nil {
return nil, errors.New("用户名已存在")
}
@@ -66,6 +91,30 @@ func (s *UserService) Register(username, password, displayName string) (*model.U
return u, nil
}
// IsAdmin 判断用户是否为管理员
func (s *UserService) IsAdmin(userID uint) bool {
if userID == 0 {
return false
}
u, err := s.userRepo.GetByID(userID)
if err != nil || u == nil {
return false
}
return u.Role == "admin"
}
// GetRole 返回用户角色(admin/user/空)
func (s *UserService) GetRole(userID uint) string {
if userID == 0 {
return ""
}
u, err := s.userRepo.GetByID(userID)
if err != nil || u == nil {
return ""
}
return u.Role
}
// Login 校验用户名密码,返回用户
func (s *UserService) Login(username, password string) (*model.User, error) {
username = strings.TrimSpace(strings.ToLower(username))