修复了一些逻辑漏洞
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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, "文件类型必须是image或video")
|
||||
if fileType != "image" && fileType != "video" && fileType != "audio" && fileType != "file" {
|
||||
utils.BadRequest(c, "文件类型必须是image、video、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()
|
||||
|
||||
@@ -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, "获取成功")
|
||||
}
|
||||
|
||||
|
||||
@@ -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 // 文本
|
||||
|
||||
@@ -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("文件类型必须是image或video")
|
||||
if fileType != "image" && fileType != "video" && fileType != "audio" && fileType != "file" {
|
||||
return nil, errors.New("文件类型必须是image、video、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"}
|
||||
}
|
||||
|
||||
// 辅助函数:检查字符串是否在切片中
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user