修复了一些逻辑漏洞

This commit is contained in:
2025-12-05 16:21:43 +08:00
parent 3dfc468f50
commit d3fc16b546
6 changed files with 220 additions and 18 deletions

View File

@@ -343,6 +343,7 @@ func main() {
authGroup.POST("/contacts/delete/:id", api.DeleteContactHandler)
// 会话管理
authGroup.GET("/conversations/by-room/:room_id", api.GetConversationByRoomHandler)
authGroup.GET("/conversations", api.GetConversationListHandler)
authGroup.POST("/conversations/reset-unread", api.ResetConversationUnreadHandler)
authGroup.POST("/conversations/update", api.UpdateConversationHandler)

View File

@@ -5,7 +5,9 @@
package api
import (
"path/filepath"
"strconv"
"strings"
"xk-websocket-v2/internal/model"
"xk-websocket-v2/internal/service"
"xk-websocket-v2/internal/utils"
@@ -32,17 +34,21 @@ func UploadAttachmentHandler(c *gin.Context) {
fileType := c.PostForm("type")
if fileType == "" {
// 根据文件扩展名推断类型
ext := file.Filename[len(file.Filename)-4:]
if ext == ".jpg" || ext == ".png" || ext == ".gif" || ext == "webp" || ext == "jpeg" {
ext := strings.ToLower(filepath.Ext(file.Filename))
if ext == ".jpg" || ext == ".jpeg" || ext == ".png" || ext == ".gif" || ext == ".webp" {
fileType = "image"
} else {
} else if ext == ".mp4" || ext == ".avi" || ext == ".mov" || ext == ".wmv" || ext == ".flv" || ext == ".mkv" {
fileType = "video"
} else if ext == ".mp3" || ext == ".wav" || ext == ".m4a" || ext == ".webm" || ext == ".ogg" || ext == ".aac" {
fileType = "audio"
} else {
fileType = "file"
}
}
// 验证文件类型
if fileType != "image" && fileType != "video" {
utils.BadRequest(c, "文件类型必须是imagevideo")
if fileType != "image" && fileType != "video" && fileType != "audio" && fileType != "file" {
utils.BadRequest(c, "文件类型必须是imagevideo、audio或file")
return
}
@@ -55,6 +61,14 @@ func UploadAttachmentHandler(c *gin.Context) {
utils.BadRequest(c, "视频大小不能超过500MB")
return
}
if fileType == "audio" && file.Size > model.MaxAudioSize {
utils.BadRequest(c, "音频大小不能超过50MB")
return
}
if fileType == "file" && file.Size > model.MaxFileSize {
utils.BadRequest(c, "文件大小不能超过100MB")
return
}
// 打开文件
src, err := file.Open()

View File

@@ -9,6 +9,7 @@ import (
"xk-websocket-v2/internal/utils"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
/**
@@ -119,4 +120,30 @@ func DeleteConversationHandler(c *gin.Context) {
utils.Success(c, "删除成功")
}
/**
* GetConversationByRoomHandler
* 功能:根据 room_id 获取或创建会话
* 路径GET /api/conversations/by-room/:room_id
*/
func GetConversationByRoomHandler(c *gin.Context) {
userID, _ := c.Get("user_id")
roomID := c.Param("room_id")
if roomID == "" {
utils.BadRequest(c, "room_id 不能为空")
return
}
conv, err := service.ConversationSvc.GetOrCreateConversationByRoom(userID.(string), roomID)
if err != nil {
if err == gorm.ErrRecordNotFound {
utils.NotFound(c, "房间不存在或无权访问")
} else {
utils.InternalError(c, "查询失败")
}
return
}
utils.SuccessWithData(c, conv, "获取成功")
}

View File

@@ -581,8 +581,10 @@ const (
ChanClusterBroadcast = "ws:cluster:broadcast"
// 文件大小限制(字节)
MaxImageSize = 10 * 1024 * 1024 // 10MB
MaxVideoSize = 500 * 1024 * 1024 // 500MB
MaxImageSize = 10 * 1024 * 1024 // 10MB
MaxVideoSize = 500 * 1024 * 1024 // 500MB
MaxAudioSize = 50 * 1024 * 1024 // 50MB
MaxFileSize = 100 * 1024 * 1024 // 100MB
// 消息类型常量
MessageTypeText = 0 // 文本

View File

@@ -35,6 +35,8 @@ func InitAttachmentService(db *gorm.DB) {
// 创建上传目录
os.MkdirAll("./uploads/images", os.ModePerm)
os.MkdirAll("./uploads/videos", os.ModePerm)
os.MkdirAll("./uploads/audio", os.ModePerm)
os.MkdirAll("./uploads/files", os.ModePerm)
}
/**
@@ -49,8 +51,8 @@ func InitAttachmentService(db *gorm.DB) {
*/
func (s *AttachmentService) UploadFile(userID, fileName, fileType string, fileSize int64, fileData io.Reader) (*model.Attachment, error) {
// 验证文件类型
if fileType != "image" && fileType != "video" {
return nil, errors.New("文件类型必须是imagevideo")
if fileType != "image" && fileType != "video" && fileType != "audio" && fileType != "file" {
return nil, errors.New("文件类型必须是imagevideo、audio或file")
}
// 验证文件大小
@@ -60,12 +62,22 @@ func (s *AttachmentService) UploadFile(userID, fileName, fileType string, fileSi
if fileType == "video" && fileSize > model.MaxVideoSize {
return nil, fmt.Errorf("视频大小不能超过%dMB", model.MaxVideoSize/(1024*1024))
}
if fileType == "audio" && fileSize > model.MaxAudioSize {
return nil, fmt.Errorf("音频大小不能超过%dMB", model.MaxAudioSize/(1024*1024))
}
if fileType == "file" && fileSize > model.MaxFileSize {
return nil, fmt.Errorf("文件大小不能超过%dMB", model.MaxFileSize/(1024*1024))
}
// 验证文件扩展名
// 获取文件扩展名
ext := strings.ToLower(filepath.Ext(fileName))
allowedExts := s.getAllowedExtensions(fileType)
if !contains(allowedExts, ext) {
return nil, fmt.Errorf("不支持的文件类型,允许的类型: %v", allowedExts)
// 验证文件扩展名file类型不限制扩展名
if fileType != "file" {
allowedExts := s.getAllowedExtensions(fileType)
if len(allowedExts) > 0 && !contains(allowedExts, ext) {
return nil, fmt.Errorf("不支持的文件类型,允许的类型: %v", allowedExts)
}
}
// 生成唯一文件名
@@ -75,10 +87,17 @@ func (s *AttachmentService) UploadFile(userID, fileName, fileType string, fileSi
// 确定保存路径
var saveDir string
if fileType == "image" {
switch fileType {
case "image":
saveDir = "./uploads/images"
} else {
case "video":
saveDir = "./uploads/videos"
case "audio":
saveDir = "./uploads/audio"
case "file":
saveDir = "./uploads/files"
default:
saveDir = "./uploads/files"
}
filePath := filepath.Join(saveDir, newFileName)
@@ -97,7 +116,20 @@ func (s *AttachmentService) UploadFile(userID, fileName, fileType string, fileSi
}
// 生成访问URL
fileURL := fmt.Sprintf("/uploads/%s/%s", fileType+"s", newFileName)
var urlPrefix string
switch fileType {
case "image":
urlPrefix = "images"
case "video":
urlPrefix = "videos"
case "audio":
urlPrefix = "audio"
case "file":
urlPrefix = "files"
default:
urlPrefix = "files"
}
fileURL := fmt.Sprintf("/uploads/%s/%s", urlPrefix, newFileName)
// 获取MIME类型
mimeType := mime.TypeByExtension(ext)
@@ -189,10 +221,19 @@ func (s *AttachmentService) GetUserAttachments(userID, fileType string, page, pa
// 辅助函数:获取允许的文件扩展名
func (s *AttachmentService) getAllowedExtensions(fileType string) []string {
if fileType == "image" {
switch fileType {
case "image":
return []string{".jpg", ".jpeg", ".png", ".gif", ".webp"}
case "video":
return []string{".mp4", ".avi", ".mov", ".wmv", ".flv", ".mkv"}
case "audio":
return []string{".mp3", ".wav", ".m4a", ".webm", ".ogg", ".aac", ".flac"}
case "file":
// 文件类型允许所有扩展名,但排除图片、视频、音频的扩展名
return []string{} // 空数组表示不限制扩展名
default:
return []string{}
}
return []string{".mp4", ".avi", ".mov", ".wmv", ".flv", ".mkv"}
}
// 辅助函数:检查字符串是否在切片中

View File

@@ -50,6 +50,9 @@ func (s *ConversationService) GetConversations(userID string) ([]model.ChatConve
var room model.ChatRoom
if err := s.DB.Where("room_id = ?", list[i].RoomID).First(&room).Error; err == nil {
list[i].Room = &room
} else {
// 记录警告:群聊信息查询失败
log.Printf("⚠️ 群聊信息查询失败 (room_id: %s): %v", list[i].RoomID, err)
}
}
}
@@ -257,4 +260,118 @@ func buildMessageSummary(msg model.ChatMessage, isGroupChat bool, senderUserID s
return content
}
/**
* GetOrCreateConversationByRoom
* 功能:根据 room_id 获取或创建会话
* @param userID 用户ID
* @param roomID 房间ID
* @returns 会话对象和错误
*/
func (s *ConversationService) GetOrCreateConversationByRoom(userID, roomID string) (*model.ChatConversation, error) {
if s == nil {
return nil, gorm.ErrRecordNotFound
}
// 查询房间信息以确定类型
var room model.ChatRoom
roomErr := s.DB.Where("room_id = ?", roomID).First(&room).Error
var conversationType int
var targetID string
if roomErr == nil {
// 房间存在,根据房间类型判断
if room.RoomType == "group" {
conversationType = 2 // 群聊
targetID = roomID // 群聊时 targetID = roomID
} else {
conversationType = 1 // 私聊
// 私聊需要确定 targetID对方用户ID
// 从 room_members 中查找另一个成员
var members []model.RoomMember
if err := s.DB.Where("room_id = ?", roomID).Find(&members).Error; err == nil {
for _, member := range members {
if member.UserID != userID {
targetID = member.UserID
break
}
}
}
if targetID == "" {
return nil, gorm.ErrRecordNotFound
}
}
} else {
// 房间不存在,尝试通过 roomID 格式判断
if len(roomID) > 6 && roomID[:6] == "group_" {
conversationType = 2 // 群聊
targetID = roomID
} else {
// 可能是私聊,尝试从 room_members 查找
var members []model.RoomMember
if err := s.DB.Where("room_id = ?", roomID).Find(&members).Error; err == nil && len(members) == 2 {
conversationType = 1 // 私聊
for _, member := range members {
if member.UserID != userID {
targetID = member.UserID
break
}
}
} else {
return nil, gorm.ErrRecordNotFound
}
}
}
// 查询会话
var conv model.ChatConversation
tx := s.DB.Where("user_id = ? AND target_id = ? AND type = ?", userID, targetID, conversationType).First(&conv)
if tx.Error != nil {
if tx.Error == gorm.ErrRecordNotFound {
// 创建新会话
conv = model.ChatConversation{
UserID: userID,
TargetID: targetID,
RoomID: roomID,
Type: conversationType,
LastMessage: "",
LastTime: time.Now(),
UnreadCount: 0,
}
if err := s.DB.Create(&conv).Error; err != nil {
return nil, err
}
} else {
return nil, tx.Error
}
}
// 加载关联信息
if conv.Type == 1 { // 私聊:加载目标用户信息
var user model.User
if err := s.DB.Where("id = ?", conv.TargetID).First(&user).Error; err == nil {
conv.TargetUser = &user
} else {
log.Printf("⚠️ 用户信息查询失败 (user_id: %s): %v", conv.TargetID, err)
}
} else if conv.Type == 2 { // 群聊:加载群信息
if roomErr == nil {
// 之前查询成功,直接使用
conv.Room = &room
} else {
// 如果之前查询失败,再次尝试查询
var retryRoom model.ChatRoom
if err := s.DB.Where("room_id = ?", roomID).First(&retryRoom).Error; err == nil {
conv.Room = &retryRoom
} else {
log.Printf("⚠️ 群聊信息查询失败 (room_id: %s): %v", roomID, err)
// 即使查询失败,也返回会话,但 Room 字段为 nil
}
}
}
return &conv, nil
}