Files
nl-blogs/server/repositories/post_repository.go
2026-06-26 17:00:02 +08:00

1027 lines
27 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package repositories
import (
"encoding/json"
"fmt"
"log"
"strconv"
"strings"
"time"
"github.com/niangaodev/art-code/config"
"github.com/niangaodev/art-code/models"
"github.com/sergi/go-diff/diffmatchpatch"
"gorm.io/gorm"
)
// TrendData 趋势数据
type TrendData struct {
Date string `json:"date"`
Count int `json:"value"`
YoY float64 `json:"yoy"`
MoM float64 `json:"mom"`
}
// GetPosts 获取所有博客文章(支持搜索、分类、标签筛选,支持分页)
func GetPosts(keyword string, categoryID uint, tagID uint, columnID uint, page int, pageSize int) ([]models.Post, int64, error) {
var posts []models.Post
var total int64
// 构建基础查询
query := config.DB.Model(&models.Post{}).
Where("is_published = ? AND deleted_at = ?", 1, 0)
if tagID > 0 {
query = query.Joins("JOIN post_tags pt ON posts.id = pt.post_id").
Where("pt.tag_id = ?", tagID)
}
if categoryID > 0 {
query = query.Where("category_id = ?", categoryID)
}
if columnID > 0 {
query = query.Where("column_id = ?", columnID)
}
if keyword != "" {
likeKeyword := "%" + keyword + "%"
query = query.Where("(MATCH(title, content) AGAINST(? IN BOOLEAN MODE) OR title LIKE ? OR content LIKE ?)",
keyword, likeKeyword, likeKeyword)
}
// 计算总数
countQuery := query
err := countQuery.Count(&total).Error
if err != nil {
log.Printf("Error counting posts: %v", err)
return nil, 0, err
}
// 应用分页
offset := (page - 1) * pageSize
err = query.
Preload("Category").
Preload("Column").
Preload("Tags").
Order("created_at DESC").
Limit(pageSize).
Offset(offset).
Find(&posts).Error
if err != nil {
log.Printf("Error querying posts: %v", err)
return nil, 0, err
}
return posts, total, nil
}
// GetPostsByUserID 获取用户已发布文章(用于公开资料页)
func GetPostsByUserID(userID uint, sort string, limit int) ([]models.Post, error) {
if limit <= 0 {
limit = 5
}
if limit > 20 {
limit = 20
}
orderClause := "created_at DESC"
if sort == "popular" {
orderClause = "read_count DESC, created_at DESC"
}
var posts []models.Post
err := config.DB.Model(&models.Post{}).
Where("user_id = ? AND is_published = ? AND deleted_at = ?", userID, 1, 0).
Preload("Category").
Preload("Column").
Preload("Tags").
Order(orderClause).
Limit(limit).
Find(&posts).Error
if err != nil {
log.Printf("Error querying posts by user ID: %v", err)
return nil, err
}
return posts, nil
}
// GetPostByID 根据ID获取博客文章
func GetPostByID(id uint) (*models.Post, error) {
var post models.Post
// 显式选择所有字段,确保 content 字段被加载
err := config.DB.Model(&models.Post{}).
Select("id", "title", "category_id", "column_id", "excerpt", "content", "read_count", "is_published", "user_id", "created_at", "updated_at", "deleted_at").
Preload("Category").
Preload("Column").
Preload("Tags").
Preload("Author", func(db *gorm.DB) *gorm.DB {
return db.Select("id", "username", "avatar", "email")
}).
Where("id = ? AND is_published = ? AND deleted_at = ?", id, 1, 0).
First(&post).Error
if err != nil {
if err == gorm.ErrRecordNotFound {
return nil, nil
}
log.Printf("Error getting post by ID: %v", err)
return nil, err
}
// 不再在此处更新阅读量,由 handler 按小时去重后递增
return &post, nil
}
// IncrementReadCount 递增文章阅读量
func IncrementReadCount(id uint) error {
return config.DB.Model(&models.Post{}).
Where("id = ?", id).
UpdateColumn("read_count", gorm.Expr("read_count + ?", 1)).Error
}
// GetAllPosts 获取所有博客文章(包括未发布的,后台用,支持搜索与筛选)
type AdminPostFilter struct {
Keyword string
CategoryID uint
ColumnID uint
IsPublished int // -1 表示全部
StartDate string
EndDate string
}
func GetAllPosts(page, pageSize int, filter AdminPostFilter) ([]models.Post, int64, error) {
offset := (page - 1) * pageSize
var posts []models.Post
var total int64
query := config.DB.Model(&models.Post{}).
Where("deleted_at = ?", 0)
if filter.Keyword != "" {
likeKeyword := "%" + filter.Keyword + "%"
query = query.Where("(MATCH(title, content) AGAINST(? IN BOOLEAN MODE) OR title LIKE ? OR content LIKE ?)",
filter.Keyword, likeKeyword, likeKeyword)
}
if filter.CategoryID > 0 {
query = query.Where("category_id = ?", filter.CategoryID)
}
if filter.ColumnID > 0 {
query = query.Where("column_id = ?", filter.ColumnID)
}
if filter.IsPublished >= 0 {
query = query.Where("is_published = ?", filter.IsPublished)
}
if filter.StartDate != "" {
if ts, err := parseDateStart(filter.StartDate); err == nil {
query = query.Where("created_at >= ?", ts)
}
}
if filter.EndDate != "" {
if ts, err := parseDateEnd(filter.EndDate); err == nil {
query = query.Where("created_at <= ?", ts)
}
}
err := query.Count(&total).Error
if err != nil {
log.Printf("Error counting posts: %v", err)
return nil, 0, err
}
// Get posts
err = query.
Preload("Category").
Preload("Column").
Preload("Tags").
Order("created_at DESC").
Limit(pageSize).
Offset(offset).
Find(&posts).Error
if err != nil {
log.Printf("Error querying all posts: %v", err)
return nil, 0, err
}
return posts, total, nil
}
// CreatePost 创建博客文章
func CreatePost(post *models.Post) error {
err := config.DB.Create(post).Error
if err != nil {
log.Printf("Error creating post: %v", err)
return err
}
// Insert Tags
if len(post.Tags) > 0 {
err = config.DB.Model(post).Association("Tags").Replace(post.Tags)
if err != nil {
log.Printf("Error associating tags: %v", err)
return err
}
}
// Handle column association if ColumnID is set
if post.ColumnID != nil && *post.ColumnID > 0 {
// Add to column_posts table
err = AddPostToColumn(*post.ColumnID, post.ID, 0)
if err != nil {
log.Printf("Error adding post to column: %v", err)
// Don't fail the whole operation, just log
}
}
return nil
}
// UpdatePost 更新博客文章
func UpdatePost(post *models.Post) error {
// 使用 map 更新,并明确指定要更新的字段,确保即使字段是空字符串也会被更新
updateData := map[string]interface{}{
"title": post.Title,
"category_id": post.CategoryID,
"excerpt": post.Excerpt,
"content": post.Content, // 明确包含 content即使为空字符串也会更新
"is_published": post.IsPublished,
"updated_at": time.Now().Unix(),
}
// Update column_id if provided
if post.ColumnID != nil {
updateData["column_id"] = post.ColumnID
} else {
// Set to NULL if explicitly set to nil
updateData["column_id"] = nil
}
// 使用 Select 明确指定要更新的字段,确保所有字段都被更新
err := config.DB.Model(&models.Post{}).
Where("id = ? AND deleted_at = ?", post.ID, 0).
Select("title", "category_id", "column_id", "excerpt", "content", "is_published", "updated_at").
Updates(updateData).Error
if err != nil {
log.Printf("Error updating post: %v", err)
return err
}
// Update Tags
if len(post.Tags) > 0 {
err = config.DB.Model(&models.Post{ID: post.ID}).Association("Tags").Replace(post.Tags)
if err != nil {
log.Printf("Error updating tags: %v", err)
return err
}
} else {
// Clear all tags
err = config.DB.Model(&models.Post{ID: post.ID}).Association("Tags").Clear()
if err != nil {
log.Printf("Error clearing tags: %v", err)
return err
}
}
// Handle column association
// Always sync column_posts table with column_id
// First, remove from all columns
err = config.DB.Where("post_id = ?", post.ID).Delete(&models.ColumnPost{}).Error
if err != nil {
log.Printf("Error removing post from columns: %v", err)
// Don't fail the whole operation
}
// Then add to new column if specified
if post.ColumnID != nil && *post.ColumnID > 0 {
err = AddPostToColumn(*post.ColumnID, post.ID, 0)
if err != nil {
log.Printf("Error adding post to column: %v", err)
// Don't fail the whole operation
}
}
return nil
}
// UpdatePostRelations 只更新文章的关联关系(分类、专栏、标签),不更新内容
func UpdatePostRelations(postID uint, categoryID uint, columnID *uint, tagIDs []uint) error {
// 更新分类和专栏
updateData := map[string]interface{}{
"category_id": categoryID,
"updated_at": time.Now().Unix(),
}
if columnID != nil {
updateData["column_id"] = columnID
} else {
updateData["column_id"] = nil
}
err := config.DB.Model(&models.Post{}).
Where("id = ? AND deleted_at = ?", postID, 0).
Updates(updateData).Error
if err != nil {
log.Printf("Error updating post relations: %v", err)
return err
}
// 更新标签关联
post := &models.Post{ID: postID}
if len(tagIDs) > 0 {
// 构建标签对象
tags := make([]models.Tag, 0, len(tagIDs))
for _, tagID := range tagIDs {
tags = append(tags, models.Tag{ID: tagID})
}
err = config.DB.Model(post).Association("Tags").Replace(tags)
if err != nil {
log.Printf("Error updating tags: %v", err)
return err
}
} else {
// 清除所有标签
err = config.DB.Model(post).Association("Tags").Clear()
if err != nil {
log.Printf("Error clearing tags: %v", err)
return err
}
}
// 处理专栏关联
// 先移除所有专栏关联
err = config.DB.Where("post_id = ?", postID).Delete(&models.ColumnPost{}).Error
if err != nil {
log.Printf("Error removing post from columns: %v", err)
// Don't fail the whole operation
}
// 然后添加新的专栏关联
if columnID != nil && *columnID > 0 {
err = AddPostToColumn(*columnID, postID, 0)
if err != nil {
log.Printf("Error adding post to column: %v", err)
// Don't fail the whole operation
}
}
return nil
}
// UpdatePostStatus 更新文章状态
func UpdatePostStatus(id uint, status int) error {
err := config.DB.Model(&models.Post{}).
Where("id = ? AND deleted_at = ?", id, 0).
Updates(map[string]interface{}{
"is_published": status,
"updated_at": time.Now().Unix(),
}).Error
return err
}
// DeletePost 删除博客文章 (Soft Delete)
func DeletePost(id uint) error {
err := config.DB.Model(&models.Post{}).
Where("id = ?", id).
Update("deleted_at", time.Now().Unix()).Error
if err != nil {
log.Printf("Error deleting post: %v", err)
return err
}
return nil
}
// GetPostCount 获取文章总数
func GetPostCount() (int, error) {
var count int64
err := config.DB.Model(&models.Post{}).
Where("deleted_at = ?", 0).
Count(&count).Error
if err != nil {
log.Printf("Error getting post count: %v", err)
return 0, err
}
return int(count), nil
}
type userBrief struct {
Username string
Avatar string
Email string
}
func resolveUserBrief(post *models.Post, userMap map[uint]userBrief) (username, avatar, email string) {
if post.Author != nil && post.Author.ID > 0 {
return post.Author.Username, post.Author.Avatar, post.Author.Email
}
if userMap != nil {
if u, ok := userMap[post.UserID]; ok {
return u.Username, u.Avatar, u.Email
}
}
if post.UserID > 0 {
return lookupUserBrief(post.UserID)
}
return "", "", ""
}
func collectPostUserIDs(posts []models.Post) []uint {
ids := make([]uint, 0, len(posts))
seen := make(map[uint]struct{})
for _, post := range posts {
if post.UserID == 0 {
continue
}
if _, ok := seen[post.UserID]; ok {
continue
}
seen[post.UserID] = struct{}{}
ids = append(ids, post.UserID)
}
return ids
}
func lookupUsersBrief(ids []uint) map[uint]userBrief {
result := make(map[uint]userBrief)
if len(ids) == 0 {
return result
}
var users []models.User
err := config.DB.Model(&models.User{}).
Select("id", "username", "avatar", "email").
Where("id IN ? AND deleted_at = ?", ids, 0).
Find(&users).Error
if err != nil {
log.Printf("Error batch loading users: %v", err)
return result
}
for _, user := range users {
result[user.ID] = userBrief{
Username: user.Username,
Avatar: user.Avatar,
Email: user.Email,
}
}
return result
}
// BuildPostResponse 构建博客文章响应
func BuildPostResponse(post *models.Post, includeContent bool) *models.PostResponse {
return buildPostResponse(post, includeContent, true, nil)
}
func buildPostResponse(post *models.Post, includeContent, includeSnippets bool, userMap map[uint]userBrief) *models.PostResponse {
// Format CreatedAt to Date string
dateStr := time.Unix(post.CreatedAt, 0).Format("2006-01-02")
catName := ""
catSlug := ""
if post.Category != nil {
catName = post.Category.Name
catSlug = post.Category.Slug
} else if post.CategoryID > 0 {
catName = lookupCategoryName(post.CategoryID)
}
colName := ""
colSlug := ""
if post.Column != nil {
colName = post.Column.Name
}
userName, userAvatar, userEmail := resolveUserBrief(post, userMap)
response := &models.PostResponse{
ID: post.ID,
Title: post.Title,
CategoryID: post.CategoryID,
CategoryName: catName,
CategorySlug: catSlug,
ColumnID: post.ColumnID,
ColumnName: colName,
ColumnSlug: colSlug,
Date: dateStr,
Excerpt: post.Excerpt,
Tags: post.Tags,
IsPublished: post.IsPublished,
ReadCount: post.ReadCount,
UserID: post.UserID,
UserName: userName,
UserAvatar: userAvatar,
UserEmail: userEmail,
}
if includeContent {
// 使用指针类型,确保即使内容为空字符串也会出现在 JSON 中
content := post.Content
response.Content = &content
}
if includeSnippets {
if snippets, err := GetSnippetsByPostID(post.ID); err == nil && len(snippets) > 0 {
response.Snippets = snippets
}
}
return response
}
// BuildPostsResponse 构建博客文章列表响应(不含正文与 snippets批量查作者
func BuildPostsResponse(posts []models.Post) []models.PostResponse {
if len(posts) == 0 {
return []models.PostResponse{}
}
userMap := lookupUsersBrief(collectPostUserIDs(posts))
responses := make([]models.PostResponse, 0, len(posts))
for i := range posts {
responses = append(responses, *buildPostResponse(&posts[i], false, false, userMap))
}
return responses
}
// SavePostHistory 保存文章历史记录
func SavePostHistory(post *models.Post, modifiedBy uint) error {
// 获取当前最大版本号
var maxVersion int
err := config.DB.Model(&models.PostHistory{}).
Where("post_id = ?", post.ID).
Select("COALESCE(MAX(version), 0)").
Scan(&maxVersion).Error
if err != nil {
log.Printf("Error getting max version: %v", err)
return err
}
tagIDs := make([]uint, 0, len(post.Tags))
for _, tag := range post.Tags {
tagIDs = append(tagIDs, tag.ID)
}
tagJSON, _ := json.Marshal(tagIDs)
history := &models.PostHistory{
PostID: post.ID,
Version: maxVersion + 1,
Title: post.Title,
CategoryID: post.CategoryID,
ColumnID: post.ColumnID,
TagIDs: string(tagJSON),
Excerpt: post.Excerpt,
Content: post.Content,
IsPublished: post.IsPublished,
ModifiedBy: modifiedBy,
}
err = config.DB.Create(history).Error
if err != nil {
log.Printf("Error saving post history: %v", err)
return err
}
return nil
}
// GetTopPosts 获取热门文章 (按阅读量)
func GetTopPosts(limit int) ([]models.Post, error) {
var posts []models.Post
err := config.DB.Model(&models.Post{}).
Preload("Category").
Preload("Tags").
Where("is_published = ? AND deleted_at = ?", 1, 0).
Order("read_count DESC").
Limit(limit).
Find(&posts).Error
if err != nil {
return nil, err
}
return posts, nil
}
// GetNewPostsTrend 获取新增文章趋势
func GetNewPostsTrend(startDate, endDate string) ([]TrendData, error) {
// 确定日期范围
var startTime, endTime time.Time
if startDate != "" {
startUnix := parseDateToUnix(startDate, false)
startTime = time.Unix(startUnix, 0)
} else {
startTime = time.Now().AddDate(0, 0, -6)
startTime = time.Date(startTime.Year(), startTime.Month(), startTime.Day(), 0, 0, 0, 0, time.Local)
}
if endDate != "" {
endUnix := parseDateToUnix(endDate, true)
endTime = time.Unix(endUnix, 0)
} else {
endTime = time.Now()
endTime = time.Date(endTime.Year(), endTime.Month(), endTime.Day(), 23, 59, 59, 0, time.Local)
}
// 查询数据库
query := config.DB.Model(&models.Post{}).
Select("FROM_UNIXTIME(created_at, '%Y-%m-%d') as date, COUNT(*) as count").
Where("deleted_at = ?", 0).
Where("created_at >= ?", startTime.Unix()).
Where("created_at <= ?", endTime.Unix())
var results []TrendData
err := query.Group("date").
Order("date ASC").
Scan(&results).Error
if err != nil {
return nil, err
}
// 创建日期到数据的映射
resultMap := make(map[string]TrendData)
for _, r := range results {
resultMap[r.Date] = r
}
// 生成完整日期列表
var fullResults []TrendData
current := time.Date(startTime.Year(), startTime.Month(), startTime.Day(), 0, 0, 0, 0, time.Local)
endDay := time.Date(endTime.Year(), endTime.Month(), endTime.Day(), 0, 0, 0, 0, time.Local)
for !current.After(endDay) {
dateStr := current.Format("2006-01-02")
if data, exists := resultMap[dateStr]; exists {
data.YoY = 0
data.MoM = 0
fullResults = append(fullResults, data)
} else {
// 填充0值
fullResults = append(fullResults, TrendData{
Date: dateStr,
Count: 0,
YoY: 0,
MoM: 0,
})
}
current = current.AddDate(0, 0, 1)
}
return fullResults, nil
}
// GetPostHistory 获取文章修改历史
func GetPostHistory(postID uint) ([]models.PostHistory, error) {
var history []models.PostHistory
err := config.DB.Model(&models.PostHistory{}).
Where("post_id = ?", postID).
Order("version DESC").
Find(&history).Error
if err != nil {
return nil, err
}
return history, nil
}
// GetPostHistoryByVersion 获取特定版本的历史记录
func GetPostHistoryByVersion(postID uint, version uint) (*models.PostHistory, error) {
var h models.PostHistory
err := config.DB.Model(&models.PostHistory{}).
Where("post_id = ? AND version = ?", postID, version).
First(&h).Error
if err != nil {
if err == gorm.ErrRecordNotFound {
return nil, nil
}
return nil, err
}
return &h, nil
}
// BuildPostHistoryResponse 构建历史记录响应
func BuildPostHistoryResponse(h *models.PostHistory, includeFull bool) *models.PostHistoryResponse {
opName, opAvatar, opEmail := lookupUserBrief(h.ModifiedBy)
resp := &models.PostHistoryResponse{
ID: h.ID,
PostID: h.PostID,
Version: h.Version,
Title: h.Title,
CategoryID: h.CategoryID,
CategoryName: lookupCategoryName(h.CategoryID),
ColumnID: h.ColumnID,
ColumnName: lookupColumnName(h.ColumnID),
TagIDs: h.GetTagIDList(),
TagNames: lookupTagNames(h.GetTagIDList()),
Date: time.Unix(h.CreatedAt, 0).Format("2006-01-02"),
IsPublished: h.IsPublished,
ModifiedBy: h.ModifiedBy,
ModifiedByName: opName,
UserID: h.ModifiedBy,
UserName: opName,
UserAvatar: opAvatar,
UserEmail: opEmail,
ModifiedAt: time.Unix(h.ModifiedAt, 0).Format("2006-01-02 15:04:05"),
CreatedAt: time.Unix(h.CreatedAt, 0).Format("2006-01-02 15:04:05"),
}
if includeFull {
resp.Excerpt = h.Excerpt
resp.Content = h.Content
}
return resp
}
func lookupCategoryName(id uint) string {
if id == 0 {
return ""
}
var name string
config.DB.Model(&models.Category{}).Where("id = ?", id).Pluck("name", &name)
return name
}
func lookupColumnName(id *uint) string {
if id == nil || *id == 0 {
return ""
}
var name string
config.DB.Model(&models.Column{}).Where("id = ?", *id).Pluck("name", &name)
return name
}
func lookupTagNames(ids []uint) []string {
if len(ids) == 0 {
return nil
}
var tags []models.Tag
config.DB.Model(&models.Tag{}).Where("id IN ?", ids).Find(&tags)
names := make([]string, 0, len(tags))
for _, t := range tags {
names = append(names, t.Name)
}
return names
}
func lookupUserBrief(id uint) (username, avatar, email string) {
if id == 0 {
return "", "", ""
}
var user models.User
err := config.DB.Model(&models.User{}).
Select("username", "avatar", "email").
Where("id = ? AND deleted_at = ?", id, 0).
First(&user).Error
if err != nil {
return "", "", ""
}
return user.Username, user.Avatar, user.Email
}
func lookupUsername(id uint) string {
name, _, _ := lookupUserBrief(id)
return name
}
func formatPublishedLabel(v int) string {
if v == 1 {
return "已发布"
}
return "草稿"
}
// BuildPostHistoryResponses 构建历史记录列表响应
func BuildPostHistoryResponses(history []models.PostHistory) []models.PostHistoryResponse {
var responses []models.PostHistoryResponse
for _, h := range history {
responses = append(responses, *BuildPostHistoryResponse(&h, false))
}
return responses
}
func buildFieldDiff(from, to string, withLineDiff bool) models.PostHistoryFieldDiff {
diff := models.PostHistoryFieldDiff{
From: from,
To: to,
Changed: from != to,
}
if withLineDiff && from != to {
dmp := diffmatchpatch.New()
patches := dmp.PatchMake(from, to)
diff.Diff = dmp.PatchToText(patches)
}
return diff
}
// GetPostHistoryDiff compares two history versions.
func GetPostHistoryDiff(postID uint, fromVersion, toVersion uint) (*models.PostHistoryDiffResponse, error) {
fromHistory, err := GetPostHistoryByVersion(postID, fromVersion)
if err != nil {
return nil, err
}
if fromHistory == nil {
return nil, fmt.Errorf("from version not found")
}
toHistory, err := GetPostHistoryByVersion(postID, toVersion)
if err != nil {
return nil, err
}
if toHistory == nil {
return nil, fmt.Errorf("to version not found")
}
fields := map[string]models.PostHistoryFieldDiff{
"title": buildFieldDiff(fromHistory.Title, toHistory.Title, false),
"excerpt": buildFieldDiff(fromHistory.Excerpt, toHistory.Excerpt, false),
"content": buildFieldDiff(fromHistory.Content, toHistory.Content, true),
"category": buildFieldDiff(lookupCategoryName(fromHistory.CategoryID), lookupCategoryName(toHistory.CategoryID), false),
"column": buildFieldDiff(lookupColumnName(fromHistory.ColumnID), lookupColumnName(toHistory.ColumnID), false),
"tags": buildFieldDiff(strings.Join(lookupTagNames(fromHistory.GetTagIDList()), ", "), strings.Join(lookupTagNames(toHistory.GetTagIDList()), ", "), false),
"isPublished": buildFieldDiff(formatPublishedLabel(fromHistory.IsPublished), formatPublishedLabel(toHistory.IsPublished), false),
}
return &models.PostHistoryDiffResponse{
FromVersion: int(fromVersion),
ToVersion: int(toVersion),
Fields: fields,
}, nil
}
func formatOptionalUint(v *uint) string {
if v == nil {
return ""
}
return strconv.FormatUint(uint64(*v), 10)
}
// RestorePostFromHistory restores a post from a history snapshot and saves a new history entry.
func RestorePostFromHistory(postID, version, modifiedBy uint) (*models.Post, error) {
history, err := GetPostHistoryByVersion(postID, version)
if err != nil {
return nil, err
}
if history == nil {
return nil, fmt.Errorf("history version not found")
}
var post models.Post
if err := config.DB.Preload("Tags").Where("id = ? AND deleted_at = ?", postID, 0).First(&post).Error; err != nil {
return nil, err
}
post.Title = history.Title
post.CategoryID = history.CategoryID
post.ColumnID = history.ColumnID
post.Excerpt = history.Excerpt
post.Content = history.Content
post.IsPublished = history.IsPublished
if err := UpdatePost(&post); err != nil {
return nil, err
}
tagIDs := history.GetTagIDList()
tags := make([]models.Tag, 0, len(tagIDs))
for _, id := range tagIDs {
tags = append(tags, models.Tag{ID: id})
}
if err := config.DB.Model(&post).Association("Tags").Replace(tags); err != nil {
return nil, err
}
post.Tags = tags
if err := SavePostHistory(&post, modifiedBy); err != nil {
return nil, err
}
reloaded, err := GetPostByIDAdmin(postID)
if err != nil {
return nil, err
}
return reloaded, nil
}
// GetPostByIDAdmin loads a post for admin use without incrementing read count.
func GetPostByIDAdmin(id uint) (*models.Post, error) {
var post models.Post
err := config.DB.Model(&models.Post{}).
Preload("Category").
Preload("Column").
Preload("Tags").
Where("id = ? AND deleted_at = ?", id, 0).
First(&post).Error
if err != nil {
if err == gorm.ErrRecordNotFound {
return nil, nil
}
return nil, err
}
return &post, nil
}
// GetRecommendedPostsByIP 基于IP的协同过滤推荐算法
// 查找相同IP访问过的其他文章按访问频率排序返回
func GetRecommendedPostsByIP(currentPostID uint, userIP string, limit int) ([]models.Post, error) {
if userIP == "" {
// 如果IP为空返回热门文章作为备选
return GetTopPosts(limit)
}
// 1. 查找相同IP访问过的其他文章ID及其访问次数
type PostAccessCount struct {
ArticleID uint
Count int64
}
var accessCounts []PostAccessCount
err := config.DB.Model(&models.UserAccessLog{}).
Select("article_id, COUNT(*) as count").
Where("user_ip = ? AND article_id != ? AND deleted_at = ?", userIP, currentPostID, 0).
Group("article_id").
Order("count DESC").
Limit(limit * 2). // 多查询一些,因为后面还要过滤未发布的文章
Scan(&accessCounts).Error
if err != nil {
log.Printf("Error querying recommended posts by IP: %v", err)
// 如果查询失败,返回热门文章作为备选
return GetTopPosts(limit)
}
if len(accessCounts) == 0 {
// 如果没有访问记录,返回热门文章作为备选
return GetTopPosts(limit)
}
// 2. 提取文章ID列表
articleIDs := make([]uint, 0, len(accessCounts))
for _, ac := range accessCounts {
articleIDs = append(articleIDs, ac.ArticleID)
}
// 3. 查询这些文章的详细信息(排除未发布的文章)
var posts []models.Post
err = config.DB.Model(&models.Post{}).
Preload("Category").
Preload("Tags").
Where("id IN ? AND is_published = ? AND deleted_at = ?", articleIDs, 1, 0).
Find(&posts).Error
if err != nil {
log.Printf("Error fetching recommended posts: %v", err)
return GetTopPosts(limit)
}
// 4. 按照访问次数排序(保持与 accessCounts 的顺序一致)
postMap := make(map[uint]models.Post)
for _, post := range posts {
postMap[post.ID] = post
}
var sortedPosts []models.Post
for _, ac := range accessCounts {
if post, exists := postMap[ac.ArticleID]; exists {
sortedPosts = append(sortedPosts, post)
if len(sortedPosts) >= limit {
break
}
}
}
// 5. 如果推荐的文章数量不足,用热门文章补充
if len(sortedPosts) < limit {
topPosts, err := GetTopPosts(limit - len(sortedPosts))
if err == nil {
// 排除已经推荐的文章
existingIDs := make(map[uint]bool)
for _, p := range sortedPosts {
existingIDs[p.ID] = true
}
for _, p := range topPosts {
if !existingIDs[p.ID] && p.ID != currentPostID {
sortedPosts = append(sortedPosts, p)
if len(sortedPosts) >= limit {
break
}
}
}
}
}
return sortedPosts, nil
}
func parseDateStart(dateStr string) (int64, error) {
t, err := time.ParseInLocation("2006-01-02", dateStr, time.Local)
if err != nil {
return 0, err
}
return t.Unix(), nil
}
func parseDateEnd(dateStr string) (int64, error) {
t, err := time.ParseInLocation("2006-01-02", dateStr, time.Local)
if err != nil {
return 0, err
}
return t.Add(24*time.Hour - time.Second).Unix(), nil
}