初始提交:姜十三论坛 Jiang13 Forum

轻量自用论坛,Go 单二进制 + React SPA 内嵌 + SQLite。

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
freefire
2026-06-15 21:08:52 +08:00
commit e1c1708715
140 changed files with 16115 additions and 0 deletions

114
service/auth.go Normal file
View File

@@ -0,0 +1,114 @@
package service
import (
"errors"
"time"
"github.com/golang-jwt/jwt/v5"
"github.com/jiang13/forum/model"
)
const TokenExpire = 7 * 24 * time.Hour
type Claims struct {
UserID uint `json:"user_id"`
Username string `json:"username"`
Role model.Role `json:"role"`
jwt.RegisteredClaims
}
type AuthService struct {
jwtSecret string
filter *SensitiveFilter
}
func NewAuthService(jwtSecret string, filter *SensitiveFilter) *AuthService {
return &AuthService{jwtSecret: jwtSecret, filter: filter}
}
// Register 用户注册
func (s *AuthService) Register(username, password, nickname string) (*model.User, error) {
if err := ValidateUsername(username); err != nil {
return nil, err
}
if err := ValidatePassword(password); err != nil {
return nil, err
}
var exist model.User
if err := model.DB.Where("username = ?", username).First(&exist).Error; err == nil {
return nil, ErrUserExists
}
hash, err := HashPassword(password)
if err != nil {
return nil, err
}
if nickname == "" {
nickname = username
}
nickname = s.filter.Filter(nickname)
// 首个注册用户自动成为管理员
role := model.RoleUser
var userCount int64
model.DB.Model(&model.User{}).Count(&userCount)
if userCount == 0 {
role = model.RoleAdmin
}
user := &model.User{
Username: username,
Password: hash,
Nickname: nickname,
Role: role,
}
if err := model.DB.Create(user).Error; err != nil {
return nil, err
}
return user, nil
}
// Login 用户登录,返回 JWT token
func (s *AuthService) Login(username, password string) (string, *model.User, error) {
var user model.User
if err := model.DB.Where("username = ?", username).First(&user).Error; err != nil {
return "", nil, ErrInvalidCred
}
if user.Banned {
return "", nil, ErrUserBanned
}
if !CheckPassword(user.Password, password) {
return "", nil, ErrInvalidCred
}
token, err := s.GenerateToken(&user)
return token, &user, err
}
// GenerateToken 生成 JWT
func (s *AuthService) GenerateToken(user *model.User) (string, error) {
claims := Claims{
UserID: user.ID,
Username: user.Username,
Role: user.Role,
RegisteredClaims: jwt.RegisteredClaims{
ExpiresAt: jwt.NewNumericDate(time.Now().Add(TokenExpire)),
IssuedAt: jwt.NewNumericDate(time.Now()),
},
}
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
return token.SignedString([]byte(s.jwtSecret))
}
// ParseToken 解析 JWT
func (s *AuthService) ParseToken(tokenStr string) (*Claims, error) {
token, err := jwt.ParseWithClaims(tokenStr, &Claims{}, func(t *jwt.Token) (interface{}, error) {
return []byte(s.jwtSecret), nil
})
if err != nil {
return nil, err
}
claims, ok := token.Claims.(*Claims)
if !ok || !token.Valid {
return nil, errors.New("invalid token")
}
return claims, nil
}

52
service/backup.go Normal file
View File

@@ -0,0 +1,52 @@
package service
import (
"fmt"
"io"
"os"
"path/filepath"
"time"
)
type BackupService struct {
dbPath string
dataDir string
}
func NewBackupService(dbPath, dataDir string) *BackupService {
return &BackupService{dbPath: dbPath, dataDir: dataDir}
}
// ExportSQLite 导出 SQLite 备份文件到 data 目录
func (s *BackupService) ExportSQLite() (string, error) {
src, err := os.Open(s.dbPath)
if err != nil {
return "", fmt.Errorf("打开数据库失败: %w", err)
}
defer src.Close()
filename := fmt.Sprintf("jiang13_backup_%s.db", time.Now().Format("20060102_150405"))
destPath := filepath.Join(s.dataDir, filename)
dst, err := os.Create(destPath)
if err != nil {
return "", err
}
defer dst.Close()
if _, err := io.Copy(dst, src); err != nil {
return "", err
}
return destPath, nil
}
// WriteDefaultFilterWords 写入默认敏感词配置
func WriteDefaultFilterWords(path string) error {
if _, err := os.Stat(path); err == nil {
return nil
}
content := `# 姜十三论坛敏感词配置,每行一个词,# 开头为注释
违禁词示例
广告刷单
`
return os.WriteFile(path, []byte(content), 0644)
}

66
service/board.go Normal file
View File

@@ -0,0 +1,66 @@
package service
import (
"errors"
"github.com/jiang13/forum/model"
)
type BoardService struct{}
func NewBoardService() *BoardService {
return &BoardService{}
}
// BoardWithStats 板块及帖子数量
type BoardWithStats struct {
model.Board
PostCount int `json:"post_count"`
}
func (s *BoardService) List() ([]model.Board, error) {
var boards []model.Board
err := model.DB.Order("sort_order asc, id asc").Find(&boards).Error
return boards, err
}
func (s *BoardService) ListWithStats() ([]BoardWithStats, error) {
boards, err := s.List()
if err != nil {
return nil, err
}
result := make([]BoardWithStats, len(boards))
for i, b := range boards {
var count int64
model.DB.Model(&model.Post{}).Where("board_id = ?", b.ID).Count(&count)
result[i] = BoardWithStats{Board: b, PostCount: int(count)}
}
return result, nil
}
func (s *BoardService) GetByID(id uint) (*model.Board, error) {
var board model.Board
if err := model.DB.First(&board, id).Error; err != nil {
return nil, ErrBoardNotFound
}
return &board, nil
}
func (s *BoardService) Create(name, desc string, sortOrder int) (*model.Board, error) {
board := &model.Board{Name: name, Description: desc, SortOrder: sortOrder}
return board, model.DB.Create(board).Error
}
func (s *BoardService) Update(id uint, name, desc string, sortOrder int) error {
return model.DB.Model(&model.Board{}).Where("id = ?", id).Updates(map[string]interface{}{
"name": name, "description": desc, "sort_order": sortOrder,
}).Error
}
func (s *BoardService) Delete(id uint) error {
var count int64
model.DB.Model(&model.Post{}).Where("board_id = ?", id).Count(&count)
if count > 0 {
return errors.New("该板块下还有帖子,无法删除")
}
return model.DB.Delete(&model.Board{}, id).Error
}

193
service/comment.go Normal file
View File

@@ -0,0 +1,193 @@
package service
import (
"errors"
"net/mail"
"net/url"
"strings"
"github.com/jiang13/forum/model"
)
type CommentService struct {
filter *SensitiveFilter
}
func NewCommentService(filter *SensitiveFilter) *CommentService {
return &CommentService{filter: filter}
}
type CommentCreateInput struct {
UserID uint
PostID uint
Content string
ReplyTo *uint
GuestNick string
GuestEmail string
GuestURL string
IsPrivate bool
}
func (s *CommentService) canViewPrivate(c model.Comment, viewerID uint, isAdmin bool, postAuthorID uint, guestSet map[uint]struct{}) bool {
if !c.IsPrivate {
return true
}
if isAdmin {
return true
}
if viewerID > 0 && viewerID == postAuthorID {
return true
}
if c.UserID > 0 && viewerID == c.UserID {
return true
}
if _, ok := guestSet[c.ID]; ok {
return true
}
return false
}
func (s *CommentService) fillReplyTargets(comments []model.Comment, loadMissing bool) {
idMap := make(map[uint]model.Comment, len(comments))
for _, c := range comments {
idMap[c.ID] = c
}
for i := range comments {
if comments[i].ReplyTo == nil {
continue
}
if target, ok := idMap[*comments[i].ReplyTo]; ok {
t := target
comments[i].ReplyTarget = &t
continue
}
if loadMissing {
var target model.Comment
if model.DB.Preload("User").First(&target, *comments[i].ReplyTo).Error == nil {
comments[i].ReplyTarget = &target
}
}
}
}
func (s *CommentService) ListByPost(postID, viewerID uint, isAdmin bool, postAuthorID uint, visibleGuestIDs []uint) ([]model.Comment, error) {
var comments []model.Comment
err := model.DB.Preload("User").Where("post_id = ?", postID).Order("floor asc").Find(&comments).Error
if err != nil {
return nil, err
}
guestSet := make(map[uint]struct{}, len(visibleGuestIDs))
for _, id := range visibleGuestIDs {
guestSet[id] = struct{}{}
}
for i := range comments {
if comments[i].IsPrivate && !s.canViewPrivate(comments[i], viewerID, isAdmin, postAuthorID, guestSet) {
comments[i].ContentHidden = true
comments[i].Content = ""
}
}
s.fillReplyTargets(comments, false)
return comments, nil
}
func (s *CommentService) Create(in CommentCreateInput) (*model.Comment, error) {
content := s.filter.Filter(strings.TrimSpace(in.Content))
if content == "" {
return nil, errors.New("评论内容不能为空")
}
var post model.Post
if err := model.DB.First(&post, in.PostID).Error; err != nil {
return nil, ErrPostNotFound
}
if in.UserID > 0 {
var user model.User
if err := model.DB.First(&user, in.UserID).Error; err != nil {
return nil, errors.New("用户不存在")
}
if user.Banned {
return nil, errors.New("账号已被禁言")
}
} else {
nick := strings.TrimSpace(in.GuestNick)
if nick == "" {
return nil, errors.New("请填写昵称")
}
if len([]rune(nick)) > 32 {
return nil, errors.New("昵称过长")
}
if email := strings.TrimSpace(in.GuestEmail); email != "" {
if _, err := mail.ParseAddress(email); err != nil {
return nil, errors.New("邮箱格式不正确")
}
}
if rawURL := strings.TrimSpace(in.GuestURL); rawURL != "" {
u, err := url.ParseRequestURI(rawURL)
if err != nil || u.Scheme == "" || u.Host == "" {
return nil, errors.New("网址格式不正确")
}
}
}
var maxFloor int
model.DB.Model(&model.Comment{}).Where("post_id = ?", in.PostID).Select("COALESCE(MAX(floor), 0)").Scan(&maxFloor)
if in.ReplyTo != nil {
var target model.Comment
if err := model.DB.Where("id = ? AND post_id = ?", *in.ReplyTo, in.PostID).First(&target).Error; err != nil {
return nil, ErrCommentNotFound
}
}
comment := &model.Comment{
PostID: in.PostID,
UserID: in.UserID,
Floor: maxFloor + 1,
Content: content,
ReplyTo: in.ReplyTo,
GuestNick: strings.TrimSpace(in.GuestNick),
GuestEmail: strings.TrimSpace(in.GuestEmail),
GuestURL: strings.TrimSpace(in.GuestURL),
IsPrivate: in.IsPrivate,
}
return comment, model.DB.Create(comment).Error
}
func (s *CommentService) Delete(userID, commentID uint, isAdmin bool) error {
var comment model.Comment
if err := model.DB.First(&comment, commentID).Error; err != nil {
return ErrCommentNotFound
}
if !isAdmin && comment.UserID != userID {
return ErrPermissionDenied
}
return model.DB.Delete(&comment).Error
}
func (s *CommentService) AdminDelete(commentID uint) error {
return model.DB.Delete(&model.Comment{}, commentID).Error
}
// ListRecent 管理员查看最近评论
func (s *CommentService) ListRecent(page, size int) ([]model.Comment, int64, error) {
if page < 1 {
page = 1
}
if size < 1 {
size = 20
}
var total int64
model.DB.Model(&model.Comment{}).Count(&total)
var comments []model.Comment
err := model.DB.Preload("User").Preload("Post").
Order("id desc").Offset((page-1)*size).Limit(size).Find(&comments).Error
if err != nil {
return nil, 0, err
}
s.fillReplyTargets(comments, true)
return comments, total, err
}

109
service/common.go Normal file
View File

@@ -0,0 +1,109 @@
package service
import (
"errors"
"regexp"
"strings"
"sync"
"unicode/utf8"
"golang.org/x/crypto/bcrypt"
)
var (
ErrUserExists = errors.New("用户名已存在")
ErrInvalidCred = errors.New("用户名或密码错误")
ErrUserBanned = errors.New("账号已被禁言")
ErrWeakPassword = errors.New("密码至少 6 位")
ErrInvalidUsername = errors.New("用户名 3-32 位字母数字下划线")
ErrPostNotFound = errors.New("帖子不存在")
ErrCommentNotFound = errors.New("评论不存在")
ErrPermissionDenied = errors.New("无权操作")
ErrBoardNotFound = errors.New("板块不存在")
)
var usernameRe = regexp.MustCompile(`^[a-zA-Z0-9_]{3,32}$`)
// HashPassword 使用 bcrypt 加密密码
func HashPassword(password string) (string, error) {
bytes, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
return string(bytes), err
}
// CheckPassword 校验密码
func CheckPassword(hash, password string) bool {
return bcrypt.CompareHashAndPassword([]byte(hash), []byte(password)) == nil
}
// ValidateUsername 校验用户名格式
func ValidateUsername(username string) error {
if !usernameRe.MatchString(username) {
return ErrInvalidUsername
}
return nil
}
// ValidatePassword 校验密码强度
func ValidatePassword(password string) error {
if utf8.RuneCountInString(password) < 6 {
return ErrWeakPassword
}
return nil
}
// SensitiveFilter 敏感词过滤器
type SensitiveFilter struct {
mu sync.RWMutex
words []string
}
func NewSensitiveFilter() *SensitiveFilter {
return &SensitiveFilter{
words: []string{"违禁词示例", "广告刷单"},
}
}
// LoadFromFile 从配置文件加载敏感词,每行一个词
func (f *SensitiveFilter) LoadFromFile(path string) {
data, err := osReadFile(path)
if err != nil {
return
}
lines := strings.Split(string(data), "\n")
var words []string
for _, line := range lines {
line = strings.TrimSpace(line)
if line != "" && !strings.HasPrefix(line, "#") {
words = append(words, line)
}
}
if len(words) > 0 {
f.mu.Lock()
f.words = words
f.mu.Unlock()
}
}
func (f *SensitiveFilter) Filter(text string) string {
f.mu.RLock()
defer f.mu.RUnlock()
result := text
for _, w := range f.words {
if w == "" {
continue
}
replacement := strings.Repeat("*", utf8.RuneCountInString(w))
result = strings.ReplaceAll(result, w, replacement)
}
return result
}
// osReadFile 避免循环依赖,简单封装
func osReadFile(path string) ([]byte, error) {
return readFile(path)
}
// readFile 由 filter_io.go 实现
var readFile = func(path string) ([]byte, error) {
return nil, errors.New("not implemented")
}

37
service/content.go Normal file
View File

@@ -0,0 +1,37 @@
package service
import (
"regexp"
"strconv"
"strings"
"unicode/utf8"
)
var (
membersOnlyBlockRe = regexp.MustCompile(`(?is)<members-only\b[^>]*>([\s\S]*?)</members-only>`)
htmlTagRe = regexp.MustCompile(`<[^>]+>`)
)
// RedactMembersOnlyHTML 未登录时移除会员专属区块内的正文,保留长度提示供前端展示
func RedactMembersOnlyHTML(html string) string {
if html == "" {
return html
}
return membersOnlyBlockRe.ReplaceAllStringFunc(html, func(full string) string {
m := membersOnlyBlockRe.FindStringSubmatch(full)
inner := ""
if len(m) > 1 {
inner = m[1]
}
length := membersContentLength(inner)
return `<members-only data-locked="true" data-length="` + strconv.Itoa(length) + `"></members-only>`
})
}
func membersContentLength(html string) int {
text := strings.TrimSpace(htmlTagRe.ReplaceAllString(html, ""))
if text == "" {
return 0
}
return utf8.RuneCountInString(text)
}

7
service/filter_io.go Normal file
View File

@@ -0,0 +1,7 @@
package service
import "os"
func init() {
readFile = os.ReadFile
}

133
service/online.go Normal file
View File

@@ -0,0 +1,133 @@
package service
import (
"sync"
"time"
"github.com/jiang13/forum/model"
)
const onlineTTL = 5 * time.Minute
// OnlineService 在线浏览追踪(内存):登录会员 + 游客
type OnlineService struct {
mu sync.RWMutex
seen map[uint]time.Time // 登录用户
guests map[string]time.Time // 游客访客标识
}
func NewOnlineService() *OnlineService {
s := &OnlineService{
seen: make(map[uint]time.Time),
guests: make(map[string]time.Time),
}
go s.cleanup()
return s
}
func (s *OnlineService) Ping(userID uint) {
if userID == 0 {
return
}
s.mu.Lock()
s.seen[userID] = time.Now()
s.mu.Unlock()
}
func (s *OnlineService) PingGuest(visitorID string) {
if visitorID == "" {
return
}
s.mu.Lock()
s.guests[visitorID] = time.Now()
s.mu.Unlock()
}
type OnlineUser struct {
ID uint `json:"id"`
Nickname string `json:"nickname"`
Avatar string `json:"avatar"`
}
func (s *OnlineService) List(limit int) []OnlineUser {
if limit <= 0 {
limit = 20
}
cutoff := time.Now().Add(-onlineTTL)
s.mu.RLock()
var ids []uint
for id, t := range s.seen {
if t.After(cutoff) {
ids = append(ids, id)
}
}
s.mu.RUnlock()
if len(ids) == 0 {
return nil
}
var users []model.User
model.DB.Where("id IN ?", ids).Limit(limit).Find(&users)
out := make([]OnlineUser, 0, len(users))
for _, u := range users {
out = append(out, OnlineUser{ID: u.ID, Nickname: u.Nickname, Avatar: u.Avatar})
}
return out
}
func (s *OnlineService) CountMembers() int {
return s.countSeen(s.seen)
}
func (s *OnlineService) CountGuests() int {
return s.countSeenString(s.guests)
}
// Count 当前浏览总人数(会员 + 游客)
func (s *OnlineService) Count() int {
return s.CountMembers() + s.CountGuests()
}
func (s *OnlineService) countSeen(m map[uint]time.Time) int {
cutoff := time.Now().Add(-onlineTTL)
s.mu.RLock()
defer s.mu.RUnlock()
n := 0
for _, t := range m {
if t.After(cutoff) {
n++
}
}
return n
}
func (s *OnlineService) countSeenString(m map[string]time.Time) int {
cutoff := time.Now().Add(-onlineTTL)
s.mu.RLock()
defer s.mu.RUnlock()
n := 0
for _, t := range m {
if t.After(cutoff) {
n++
}
}
return n
}
func (s *OnlineService) cleanup() {
ticker := time.NewTicker(time.Minute)
for range ticker.C {
cutoff := time.Now().Add(-onlineTTL * 2)
s.mu.Lock()
for id, t := range s.seen {
if t.Before(cutoff) {
delete(s.seen, id)
}
}
for id, t := range s.guests {
if t.Before(cutoff) {
delete(s.guests, id)
}
}
s.mu.Unlock()
}
}

253
service/post.go Normal file
View File

@@ -0,0 +1,253 @@
package service
import (
"errors"
"strings"
"github.com/jiang13/forum/model"
"gorm.io/gorm"
)
type PostService struct {
filter *SensitiveFilter
}
func NewPostService(filter *SensitiveFilter) *PostService {
return &PostService{filter: filter}
}
type PostListQuery struct {
BoardID uint
Page int
Size int
Keyword string
}
// PostListItem 帖子列表项(含评论数等扩展字段)
type PostListItem struct {
model.Post
CommentCount int `json:"comment_count"`
}
func (s *PostService) ListItems(q PostListQuery) ([]PostListItem, int64, error) {
posts, total, err := s.List(q)
if err != nil {
return nil, 0, err
}
if len(posts) == 0 {
return []PostListItem{}, total, nil
}
ids := make([]uint, len(posts))
for i, p := range posts {
ids[i] = p.ID
}
countMap := s.commentCountMap(ids)
items := make([]PostListItem, len(posts))
for i, p := range posts {
items[i] = PostListItem{
Post: p,
CommentCount: countMap[p.ID],
}
}
return items, total, nil
}
func (s *PostService) commentCountMap(postIDs []uint) map[uint]int {
type row struct {
PostID uint
Count int
}
var rows []row
model.DB.Model(&model.Comment{}).Select("post_id, count(*) as count").
Where("post_id IN ?", postIDs).Group("post_id").Scan(&rows)
m := make(map[uint]int)
for _, r := range rows {
m[r.PostID] = r.Count
}
return m
}
func (s *PostService) HotPosts(limit int) ([]PostListItem, error) {
if limit <= 0 {
limit = 10
}
var posts []model.Post
err := model.DB.Preload("User").Preload("Board").
Order("like_count desc, view_count desc").Limit(limit).Find(&posts).Error
if err != nil {
return nil, err
}
ids := make([]uint, len(posts))
for i, p := range posts {
ids[i] = p.ID
}
countMap := s.commentCountMap(ids)
items := make([]PostListItem, len(posts))
for i, p := range posts {
items[i] = PostListItem{Post: p, CommentCount: countMap[p.ID]}
}
return items, nil
}
func (s *PostService) CommentCount(postID uint) int {
var count int64
model.DB.Model(&model.Comment{}).Where("post_id = ?", postID).Count(&count)
return int(count)
}
func (s *PostService) List(q PostListQuery) ([]model.Post, int64, error) {
if q.Page < 1 {
q.Page = 1
}
if q.Size < 1 {
q.Size = 20
}
db := model.DB.Model(&model.Post{}).Preload("User").Preload("Board")
if q.BoardID > 0 {
db = db.Where("board_id = ?", q.BoardID)
}
if q.Keyword != "" {
kw := "%" + q.Keyword + "%"
db = db.Where("title LIKE ? OR content LIKE ? OR tags LIKE ?", kw, kw, kw)
}
var total int64
db.Count(&total)
var posts []model.Post
err := db.Order("pinned desc, id desc").Offset((q.Page - 1) * q.Size).Limit(q.Size).Find(&posts).Error
return posts, total, err
}
func (s *PostService) GetByID(id uint) (*model.Post, error) {
var post model.Post
err := model.DB.Preload("User").Preload("Board").First(&post, id).Error
if err != nil {
return nil, ErrPostNotFound
}
model.DB.Model(&post).UpdateColumn("view_count", gorm.Expr("view_count + 1"))
return &post, nil
}
func (s *PostService) Create(userID, boardID uint, title, content, tags string) (*model.Post, error) {
title = s.filter.Filter(strings.TrimSpace(title))
content = s.filter.Filter(content)
tags = s.filter.Filter(strings.TrimSpace(tags))
if title == "" || content == "" {
return nil, errors.New("标题和内容不能为空")
}
if _, err := NewBoardService().GetByID(boardID); err != nil {
return nil, err
}
post := &model.Post{
BoardID: boardID, UserID: userID,
Title: title, Content: content, Tags: tags,
}
return post, model.DB.Create(post).Error
}
func (s *PostService) Update(userID, postID uint, isAdmin bool, title, content, tags string) error {
var post model.Post
if err := model.DB.First(&post, postID).Error; err != nil {
return ErrPostNotFound
}
if !isAdmin && post.UserID != userID {
return ErrPermissionDenied
}
title = s.filter.Filter(strings.TrimSpace(title))
content = s.filter.Filter(content)
tags = s.filter.Filter(strings.TrimSpace(tags))
return model.DB.Model(&post).Updates(map[string]interface{}{
"title": title, "content": content, "tags": tags,
}).Error
}
func (s *PostService) Delete(userID, postID uint, isAdmin bool) error {
var post model.Post
if err := model.DB.First(&post, postID).Error; err != nil {
return ErrPostNotFound
}
if !isAdmin && post.UserID != userID {
return ErrPermissionDenied
}
return model.DB.Transaction(func(tx *gorm.DB) error {
if err := tx.Where("post_id = ?", postID).Delete(&model.Comment{}).Error; err != nil {
return err
}
if err := tx.Where("post_id = ?", postID).Delete(&model.PostLike{}).Error; err != nil {
return err
}
if err := tx.Where("post_id = ?", postID).Delete(&model.PostFavorite{}).Error; err != nil {
return err
}
return tx.Delete(&post).Error
})
}
func (s *PostService) SetPinned(postID uint, pinned bool) error {
return model.DB.Model(&model.Post{}).Where("id = ?", postID).Update("pinned", pinned).Error
}
func (s *PostService) ToggleLike(userID, postID uint) (liked bool, err error) {
var like model.PostLike
result := model.DB.Where("post_id = ? AND user_id = ?", postID, userID).Limit(1).Find(&like)
if result.Error != nil {
return false, result.Error
}
if result.RowsAffected > 0 {
model.DB.Delete(&like)
model.DB.Model(&model.Post{}).Where("id = ?", postID).UpdateColumn("like_count", gorm.Expr("like_count - 1"))
return false, nil
}
like = model.PostLike{PostID: postID, UserID: userID}
if err := model.DB.Create(&like).Error; err != nil {
return false, err
}
model.DB.Model(&model.Post{}).Where("id = ?", postID).UpdateColumn("like_count", gorm.Expr("like_count + 1"))
return true, nil
}
func (s *PostService) IsLiked(userID, postID uint) bool {
var count int64
model.DB.Model(&model.PostLike{}).Where("post_id = ? AND user_id = ?", postID, userID).Count(&count)
return count > 0
}
func (s *PostService) ToggleFavorite(userID, postID uint) (faved bool, err error) {
var fav model.PostFavorite
result := model.DB.Where("post_id = ? AND user_id = ?", postID, userID).Limit(1).Find(&fav)
if result.Error != nil {
return false, result.Error
}
if result.RowsAffected > 0 {
if err := model.DB.Delete(&fav).Error; err != nil {
return false, err
}
return false, nil
}
fav = model.PostFavorite{PostID: postID, UserID: userID}
if err := model.DB.Create(&fav).Error; err != nil {
return false, err
}
return true, nil
}
func (s *PostService) IsFavorited(userID, postID uint) bool {
var count int64
model.DB.Model(&model.PostFavorite{}).Where("post_id = ? AND user_id = ?", postID, userID).Count(&count)
return count > 0
}
func (s *PostService) ListFavorites(userID uint, page, size int) ([]model.PostFavorite, int64, error) {
if page < 1 {
page = 1
}
if size < 1 {
size = 20
}
var total int64
model.DB.Model(&model.PostFavorite{}).Where("user_id = ?", userID).Count(&total)
var favs []model.PostFavorite
err := model.DB.Preload("Post.User").Preload("Post.Board").
Where("user_id = ?", userID).Order("id desc").
Offset((page - 1) * size).Limit(size).Find(&favs).Error
return favs, total, err
}

69
service/ratelimit.go Normal file
View File

@@ -0,0 +1,69 @@
package service
import (
"sync"
"time"
)
// RateLimiter 简单内存限流器,防止重复刷屏
type RateLimiter struct {
mu sync.Mutex
records map[string][]time.Time
limit int // 窗口内最大次数
window time.Duration // 时间窗口
}
func NewRateLimiter(limit int, window time.Duration) *RateLimiter {
r := &RateLimiter{
records: make(map[string][]time.Time),
limit: limit,
window: window,
}
go r.cleanup()
return r
}
// Allow 检查 key如 userID+action是否允许操作
func (r *RateLimiter) Allow(key string) bool {
r.mu.Lock()
defer r.mu.Unlock()
now := time.Now()
cutoff := now.Add(-r.window)
times := r.records[key]
var valid []time.Time
for _, t := range times {
if t.After(cutoff) {
valid = append(valid, t)
}
}
if len(valid) >= r.limit {
r.records[key] = valid
return false
}
valid = append(valid, now)
r.records[key] = valid
return true
}
func (r *RateLimiter) cleanup() {
ticker := time.NewTicker(5 * time.Minute)
for range ticker.C {
r.mu.Lock()
now := time.Now()
cutoff := now.Add(-r.window * 2)
for k, times := range r.records {
var valid []time.Time
for _, t := range times {
if t.After(cutoff) {
valid = append(valid, t)
}
}
if len(valid) == 0 {
delete(r.records, k)
} else {
r.records[k] = valid
}
}
r.mu.Unlock()
}
}

126
service/user.go Normal file
View File

@@ -0,0 +1,126 @@
package service
import (
"errors"
"fmt"
"io"
"mime/multipart"
"os"
"path/filepath"
"strings"
"time"
"github.com/jiang13/forum/model"
)
type UserService struct {
filter *SensitiveFilter
}
func NewUserService(filter *SensitiveFilter) *UserService {
return &UserService{filter: filter}
}
// GetByID 获取用户信息
func (s *UserService) GetByID(id uint) (*model.User, error) {
var user model.User
if err := model.DB.First(&user, id).Error; err != nil {
return nil, err
}
return &user, nil
}
// GetByUsername 按用户名查询
func (s *UserService) GetByUsername(username string) (*model.User, error) {
var user model.User
if err := model.DB.Where("username = ?", username).First(&user).Error; err != nil {
return nil, err
}
return &user, nil
}
// UpdateNickname 修改昵称
func (s *UserService) UpdateNickname(userID uint, nickname string) error {
nickname = strings.TrimSpace(nickname)
if nickname == "" {
return errors.New("昵称不能为空")
}
nickname = s.filter.Filter(nickname)
return model.DB.Model(&model.User{}).Where("id = ?", userID).Update("nickname", nickname).Error
}
// UpdatePassword 修改密码
func (s *UserService) UpdatePassword(userID uint, oldPass, newPass string) error {
if err := ValidatePassword(newPass); err != nil {
return err
}
var user model.User
if err := model.DB.First(&user, userID).Error; err != nil {
return err
}
if !CheckPassword(user.Password, oldPass) {
return errors.New("原密码错误")
}
hash, err := HashPassword(newPass)
if err != nil {
return err
}
return model.DB.Model(&user).Update("password", hash).Error
}
// UploadAvatar 上传头像到本地目录
func (s *UserService) UploadAvatar(userID uint, file *multipart.FileHeader, uploadDir string) (string, error) {
ext := strings.ToLower(filepath.Ext(file.Filename))
allowed := map[string]bool{".jpg": true, ".jpeg": true, ".png": true, ".gif": true, ".webp": true}
if !allowed[ext] {
return "", errors.New("仅支持 jpg/png/gif/webp 格式")
}
filename := fmt.Sprintf("%d%s", userID, ext)
destPath := filepath.Join(uploadDir, filename)
src, err := file.Open()
if err != nil {
return "", err
}
defer src.Close()
dst, err := os.Create(destPath)
if err != nil {
return "", err
}
defer dst.Close()
if _, err := io.Copy(dst, src); err != nil {
return "", err
}
avatarURL := "/uploads/avatars/" + filename
return avatarURL, model.DB.Model(&model.User{}).Where("id = ?", userID).Update("avatar", avatarURL).Error
}
// ListUsers 管理员列出用户
func (s *UserService) ListUsers(page, size int) ([]model.User, int64, error) {
var users []model.User
var total int64
model.DB.Model(&model.User{}).Count(&total)
offset := (page - 1) * size
err := model.DB.Order("id desc").Offset(offset).Limit(size).Find(&users).Error
return users, total, err
}
// BanUser 禁言用户
func (s *UserService) BanUser(userID uint, banned bool) error {
var user model.User
if err := model.DB.First(&user, userID).Error; err != nil {
return errors.New("用户不存在")
}
if user.Role == model.RoleAdmin {
return errors.New("不能禁言管理员账号")
}
now := time.Now()
updates := map[string]interface{}{"banned": banned}
if banned {
updates["banned_at"] = &now
}
return model.DB.Model(&model.User{}).Where("id = ?", userID).Updates(updates).Error
}