修复了一些问题
This commit is contained in:
@@ -13,6 +13,8 @@ import (
|
||||
"xk-websocket-v2/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
)
|
||||
|
||||
/**
|
||||
@@ -58,16 +60,49 @@ func UpdateUserSettingsHandler(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
// 上限校验:设置项数量与 key/value 长度设置合理上限,
|
||||
// 防止单次提交超大 map 触发大量写库(滥用/误用防护)
|
||||
const (
|
||||
maxSettingsPerRequest = 100
|
||||
maxKeyLen = 100
|
||||
maxValueLen = 4000
|
||||
)
|
||||
if len(req.Settings) > maxSettingsPerRequest {
|
||||
utils.BadRequest(c, "设置项过多")
|
||||
return
|
||||
}
|
||||
for key, val := range req.Settings {
|
||||
if key == "" || len(key) > maxKeyLen {
|
||||
utils.BadRequest(c, "设置键不合法")
|
||||
return
|
||||
}
|
||||
if len(val) > maxValueLen {
|
||||
utils.BadRequest(c, "设置值过长")
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
uid := userID.(string)
|
||||
now := time.Now()
|
||||
for key, val := range req.Settings {
|
||||
setting := model.UserSetting{
|
||||
UserID: uid,
|
||||
SettingKey: key,
|
||||
SettingValue: val,
|
||||
UpdatedAt: now,
|
||||
// 事务包裹批量写入:原实现循环逐条 Save,中途某条失败时前面的 key 已落库、
|
||||
// 接口却返回错误,用户设置停留在"新旧混合"的中间状态;
|
||||
// 改为同一事务内全部成功才提交,任一失败整体回滚,保证批量更新的原子性
|
||||
if err := service.ChatSvc.DB.Transaction(func(tx *gorm.DB) error {
|
||||
for key, val := range req.Settings {
|
||||
setting := model.UserSetting{
|
||||
UserID: uid,
|
||||
SettingKey: key,
|
||||
SettingValue: val,
|
||||
UpdatedAt: now,
|
||||
}
|
||||
if err := tx.Save(&setting).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
service.ChatSvc.DB.Save(&setting)
|
||||
return nil
|
||||
}); err != nil {
|
||||
utils.InternalError(c, "保存设置失败")
|
||||
return
|
||||
}
|
||||
|
||||
GetUserSettingsHandler(c)
|
||||
@@ -94,31 +129,78 @@ func MarkMessagesReadHandler(c *gin.Context) {
|
||||
|
||||
uid := userID.(string)
|
||||
now := time.Now()
|
||||
readIDs := make([]uint, 0, len(req.MessageIDs))
|
||||
for _, mid := range req.MessageIDs {
|
||||
receipt := model.MessageReadReceipt{
|
||||
MessageID: mid,
|
||||
|
||||
// 批量查询消息:一条 IN 查询取回,
|
||||
// 原实现每条消息 FirstOrCreate + First 各一次(N+1),一次标记50条已读要打100+次库
|
||||
var msgs []model.ChatMessage
|
||||
if err := service.ChatSvc.DB.Where("id IN ?", req.MessageIDs).Find(&msgs).Error; err != nil {
|
||||
utils.InternalError(c, "db error")
|
||||
return
|
||||
}
|
||||
|
||||
// 越权校验:只允许标记"自己所在房间"的消息为已读(口径与 HistoryHandler 的
|
||||
// IsRoomMember 鉴权对称)。消息ID由客户端任意提交,不校验成员身份的话,
|
||||
// 任何登录用户都能对别人房间的消息伪造已读回执,发送者会收到虚假的 messages_read 推送。
|
||||
// 同批消息通常集中在少数房间,按房间去重后每个房间只查一次成员关系,避免逐条 N 次查库;
|
||||
// 非成员房间的消息静默剔除(批量接口部分生效语义),readIDs 只返回真正标记成功的部分。
|
||||
roomAllowed := make(map[string]bool)
|
||||
visibleMsgs := msgs[:0]
|
||||
for _, msg := range msgs {
|
||||
allowed, checked := roomAllowed[msg.RoomID]
|
||||
if !checked {
|
||||
allowed = service.RoomSvc.IsRoomMember(msg.RoomID, uid)
|
||||
roomAllowed[msg.RoomID] = allowed
|
||||
}
|
||||
if allowed {
|
||||
visibleMsgs = append(visibleMsgs, msg)
|
||||
}
|
||||
}
|
||||
msgs = visibleMsgs
|
||||
|
||||
// 批量插入已读回执:靠 (message_id, user_id) 唯一索引 + OnConflict DoNothing
|
||||
// 实现幂等,重复标记已读不报错也不更新(保留首次已读时间)
|
||||
receipts := make([]model.MessageReadReceipt, 0, len(msgs))
|
||||
readIDs := make([]uint, 0, len(msgs))
|
||||
for _, msg := range msgs {
|
||||
receipts = append(receipts, model.MessageReadReceipt{
|
||||
MessageID: msg.ID,
|
||||
UserID: uid,
|
||||
ReadAt: now,
|
||||
})
|
||||
readIDs = append(readIDs, msg.ID)
|
||||
}
|
||||
if len(receipts) > 0 {
|
||||
if err := service.ChatSvc.DB.Clauses(clause.OnConflict{DoNothing: true}).Create(&receipts).Error; err != nil {
|
||||
utils.InternalError(c, "标记已读失败")
|
||||
return
|
||||
}
|
||||
service.ChatSvc.DB.Where("message_id = ? AND user_id = ?", mid, uid).
|
||||
Assign(receipt).FirstOrCreate(&receipt)
|
||||
readIDs = append(readIDs, mid)
|
||||
}
|
||||
|
||||
// 通知消息发送者(已读回执)
|
||||
var msg model.ChatMessage
|
||||
if err := service.ChatSvc.DB.First(&msg, mid).Error; err == nil && msg.SenderUserID != "" && msg.SenderUserID != uid {
|
||||
payload := model.WsPayload{
|
||||
RequestType: "messages_read",
|
||||
Data: map[string]interface{}{
|
||||
"message_ids": []uint{mid},
|
||||
"reader_id": uid,
|
||||
"room_id": msg.RoomID,
|
||||
},
|
||||
}
|
||||
if bytes, err := json.Marshal(payload); err == nil {
|
||||
service.ChatSvc.DispatchMessage(msg.SenderUserID, bytes)
|
||||
}
|
||||
// 按"发送者+房间"分组推送已读回执:同一发送者在同一房间的多条消息合并成一条推送,
|
||||
// 原实现每条消息推一次,批量已读时会向发送者刷屏式推送
|
||||
type senderRoomKey struct {
|
||||
sender string
|
||||
room string
|
||||
}
|
||||
grouped := make(map[senderRoomKey][]uint)
|
||||
for _, msg := range msgs {
|
||||
if msg.SenderUserID == "" || msg.SenderUserID == uid {
|
||||
continue
|
||||
}
|
||||
key := senderRoomKey{sender: msg.SenderUserID, room: msg.RoomID}
|
||||
grouped[key] = append(grouped[key], msg.ID)
|
||||
}
|
||||
for key, ids := range grouped {
|
||||
payload := model.WsPayload{
|
||||
RequestType: "messages_read",
|
||||
Data: map[string]interface{}{
|
||||
"message_ids": ids,
|
||||
"reader_id": uid,
|
||||
"room_id": key.room,
|
||||
},
|
||||
}
|
||||
if bytes, err := json.Marshal(payload); err == nil {
|
||||
service.ChatSvc.DispatchMessage(key.sender, bytes)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user