diff --git a/cmd/server/main.go b/cmd/server/main.go index 958c189..12006d1 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -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) diff --git a/internal/api/attachment_handler.go b/internal/api/attachment_handler.go index 568951e..678909a 100644 --- a/internal/api/attachment_handler.go +++ b/internal/api/attachment_handler.go @@ -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() diff --git a/internal/api/conversation_handler.go b/internal/api/conversation_handler.go index 9d4175b..cce4e27 100644 --- a/internal/api/conversation_handler.go +++ b/internal/api/conversation_handler.go @@ -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, "获取成功") +} diff --git a/internal/model/types.go b/internal/model/types.go index c3cbfd1..50da703 100644 --- a/internal/model/types.go +++ b/internal/model/types.go @@ -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 // 文本 diff --git a/internal/service/attachment_service.go b/internal/service/attachment_service.go index a209e24..15e2ea0 100644 --- a/internal/service/attachment_service.go +++ b/internal/service/attachment_service.go @@ -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"} } // 辅助函数:检查字符串是否在切片中 diff --git a/internal/service/conversation_service.go b/internal/service/conversation_service.go index d3b70cc..c81af01 100644 --- a/internal/service/conversation_service.go +++ b/internal/service/conversation_service.go @@ -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 +} +