初始提交:姜十三论坛 Jiang13 Forum
轻量自用论坛,Go 单二进制 + React SPA 内嵌 + SQLite。 Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
114
service/auth.go
Normal file
114
service/auth.go
Normal 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
52
service/backup.go
Normal 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
66
service/board.go
Normal 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
193
service/comment.go
Normal 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
109
service/common.go
Normal 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
37
service/content.go
Normal 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
7
service/filter_io.go
Normal file
@@ -0,0 +1,7 @@
|
||||
package service
|
||||
|
||||
import "os"
|
||||
|
||||
func init() {
|
||||
readFile = os.ReadFile
|
||||
}
|
||||
133
service/online.go
Normal file
133
service/online.go
Normal 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
253
service/post.go
Normal 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
69
service/ratelimit.go
Normal 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
126
service/user.go
Normal 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
|
||||
}
|
||||
Reference in New Issue
Block a user