package app import ( "context" "database/sql" "encoding/json" "errors" "fmt" "net/http" "net/url" "regexp" "sort" "strconv" "strings" "sync" "time" "unicode/utf8" "github.com/gorilla/websocket" ) type messageView struct { ID int64 `json:"id"` ConversationID int64 `json:"conversationId"` Seq int64 `json:"seq"` SenderID int64 `json:"senderId"` ClientMsgID string `json:"clientMsgId"` Type int `json:"type"` Content any `json:"content"` Recalled bool `json:"recalled"` MediaDeleted bool `json:"mediaDeleted,omitempty"` CreatedAt time.Time `json:"createdAt"` } const ( maxMessageBodyBytes = 16 << 10 maxMessageTextRunes = 2000 ) var clientMessageIDPattern = regexp.MustCompile(`^[A-Za-z0-9_.:-]{1,64}$`) func validMessageMediaURL(raw string, production bool) bool { raw = strings.TrimSpace(raw) if raw == "" || len(raw) > 2048 { return false } if strings.HasPrefix(raw, "/uploads/") { return !strings.Contains(raw, "..") } parsed, err := url.Parse(raw) if err != nil || parsed.Host == "" || (parsed.Scheme != "http" && parsed.Scheme != "https") { return false } return !production || parsed.Scheme == "https" } func numericDuration(value any) (int, bool) { switch duration := value.(type) { case float64: return int(duration), duration == float64(int(duration)) case int: return duration, true case int64: return int(duration), true case json.Number: parsed, err := strconv.Atoi(duration.String()) return parsed, err == nil default: return 0, false } } func (a *App) validateMessagePayload(clientMsgID string, messageType int, content any) (string, any, []byte, error) { clientMsgID = strings.TrimSpace(clientMsgID) if clientMsgID == "" { clientMsgID = randomToken()[:26] } if !clientMessageIDPattern.MatchString(clientMsgID) { return "", nil, nil, fmt.Errorf("客户端消息 ID 格式无效") } if messageType == 0 { messageType = 1 } data, ok := content.(map[string]any) if !ok { return "", nil, nil, fmt.Errorf("消息内容格式错误") } switch messageType { case 1: text, ok := data["text"].(string) text = strings.TrimSpace(text) if !ok || text == "" || !utf8.ValidString(text) || utf8.RuneCountInString(text) > maxMessageTextRunes { return "", nil, nil, fmt.Errorf("文本消息应为 1 至 %d 个字符", maxMessageTextRunes) } data = map[string]any{"text": text} case 2: mediaURL, ok := data["url"].(string) if !ok || !validMessageMediaURL(mediaURL, a.config.Environment == "production") { return "", nil, nil, fmt.Errorf("图片地址无效") } data = map[string]any{"url": strings.TrimSpace(mediaURL)} case 3: mediaURL, ok := data["url"].(string) duration, durationOK := numericDuration(data["duration"]) if !ok || !validMessageMediaURL(mediaURL, a.config.Environment == "production") || !durationOK || duration < 1 || duration > 60 { return "", nil, nil, fmt.Errorf("语音消息地址或时长无效") } data = map[string]any{"duration": duration, "url": strings.TrimSpace(mediaURL)} default: return "", nil, nil, fmt.Errorf("不支持的消息类型") } body, err := json.Marshal(data) if err != nil || len(body) > maxMessageBodyBytes { return "", nil, nil, fmt.Errorf("消息内容过大") } return clientMsgID, data, body, nil } type wsClient struct { userID int64 conn *websocket.Conn mu sync.Mutex } func (c *wsClient) writeJSON(payload any) error { c.mu.Lock() defer c.mu.Unlock() _ = c.conn.SetWriteDeadline(time.Now().Add(5 * time.Second)) return c.conn.WriteJSON(payload) } type Hub struct { mu sync.RWMutex clients map[int64]map[*wsClient]struct{} } func NewHub() *Hub { return &Hub{clients: make(map[int64]map[*wsClient]struct{})} } func (h *Hub) add(client *wsClient) bool { h.mu.Lock() defer h.mu.Unlock() becameOnline := len(h.clients[client.userID]) == 0 if h.clients[client.userID] == nil { h.clients[client.userID] = make(map[*wsClient]struct{}) } h.clients[client.userID][client] = struct{}{} return becameOnline } func (h *Hub) remove(client *wsClient) bool { h.mu.Lock() defer h.mu.Unlock() if _, exists := h.clients[client.userID][client]; !exists { return false } delete(h.clients[client.userID], client) if len(h.clients[client.userID]) == 0 { delete(h.clients, client.userID) return true } return false } func (h *Hub) online(userID int64) bool { h.mu.RLock() defer h.mu.RUnlock() return len(h.clients[userID]) > 0 } func (h *Hub) onlineUsers() []int64 { h.mu.RLock() defer h.mu.RUnlock() users := make([]int64, 0, len(h.clients)) for userID, clients := range h.clients { if len(clients) > 0 { users = append(users, userID) } } return users } func (h *Hub) broadcast(userIDs []int64, payload any) { h.mu.RLock() targets := []*wsClient{} for _, userID := range userIDs { for client := range h.clients[userID] { targets = append(targets, client) } } h.mu.RUnlock() for _, client := range targets { _ = client.writeJSON(payload) } } func (h *Hub) disconnect(userID int64) { h.mu.RLock() targets := []*wsClient{} for client := range h.clients[userID] { targets = append(targets, client) } h.mu.RUnlock() for _, client := range targets { client.mu.Lock() _ = client.conn.WriteControl(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.ClosePolicyViolation, "登录状态已失效"), time.Now().Add(time.Second)) _ = client.conn.Close() client.mu.Unlock() } } func (a *App) onlineNow(userID int64) bool { return a.hub.online(userID) } func (a *App) publishPresence(userID int64, online bool) { ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) defer cancel() var visible int if err := a.db.QueryRowContext(ctx, `SELECT online_visible FROM user_privacy_settings WHERE user_id=?`, userID).Scan(&visible); err != nil { return } rows, err := a.db.QueryContext(ctx, `SELECT DISTINCT peer.user_id FROM im_conversation_members mine JOIN im_conversation_members peer ON peer.conversation_id=mine.conversation_id AND peer.user_id<>mine.user_id JOIN im_conversations conversation ON conversation.id=mine.conversation_id WHERE mine.user_id=? AND mine.status=1 AND peer.status=1 AND conversation.status=1`, userID) if err != nil { return } defer rows.Close() peers := []int64{} for rows.Next() { var peerID int64 if rows.Scan(&peerID) == nil { peers = append(peers, peerID) } } if len(peers) == 0 { return } a.hub.broadcast(peers, map[string]any{"command": "PRESENCE", "data": map[string]any{ "userId": userID, "online": online && visible == 1, "lastActiveAt": time.Now(), }}) } func (a *App) directConversation(w http.ResponseWriter, r *http.Request) { var req struct { UserID int64 `json:"userId"` } if decode(r, &req) != nil || req.UserID == 0 || req.UserID == current(r).ID { fail(w, 400, 20001, "无效的聊天对象") return } first, second := current(r).ID, req.UserID if first > second { first, second = second, first } var targetStatus, allowStranger int if err := a.db.QueryRowContext(r.Context(), `SELECT u.status,p.allow_stranger_message FROM users u JOIN user_privacy_settings p ON p.user_id=u.id WHERE u.id=? AND u.deleted_at IS NULL`, req.UserID).Scan(&targetStatus, &allowStranger); err != nil || targetStatus != 1 { fail(w, http.StatusNotFound, 30001, "聊天对象不存在或不可用") return } var blocked int if err := a.db.QueryRowContext(r.Context(), `SELECT EXISTS(SELECT 1 FROM user_blocks WHERE (user_id=? AND blocked_user_id=?) OR (user_id=? AND blocked_user_id=?))`, current(r).ID, req.UserID, req.UserID, current(r).ID).Scan(&blocked); err != nil || blocked == 1 { fail(w, http.StatusForbidden, 30002, "当前无法与该用户聊天") return } var id int64 err := a.db.QueryRowContext(r.Context(), `SELECT conversation_id FROM im_direct_conversations WHERE user1_id=? AND user2_id=?`, first, second).Scan(&id) if err == nil { reply(w, map[string]any{"id": id}) return } if allowStranger == 0 { var mutualFollow int _ = a.db.QueryRowContext(r.Context(), `SELECT EXISTS( SELECT 1 FROM user_follows mine JOIN user_follows target ON target.user_id=mine.target_user_id AND target.target_user_id=mine.user_id WHERE mine.user_id=? AND mine.target_user_id=?)`, current(r).ID, req.UserID).Scan(&mutualFollow) if mutualFollow != 1 { fail(w, http.StatusForbidden, 30002, "对方仅允许互相关注的人发起私信") return } } tx, beginErr := a.db.BeginTx(r.Context(), nil) if beginErr != nil { fail(w, http.StatusInternalServerError, 50001, "创建会话失败") return } defer func() { _ = tx.Rollback() }() result, err := tx.ExecContext(r.Context(), `INSERT INTO im_conversations(conversation_type)VALUES(1)`) if err != nil { _ = tx.Rollback() // Concurrent requests may have created the same unique direct pair. if a.db.QueryRowContext(r.Context(), `SELECT conversation_id FROM im_direct_conversations WHERE user1_id=? AND user2_id=?`, first, second).Scan(&id) == nil { reply(w, map[string]any{"id": id}) return } fail(w, 500, 50001, "创建会话失败") return } id, _ = result.LastInsertId() _, err = tx.ExecContext(r.Context(), `INSERT INTO im_direct_conversations(conversation_id,user1_id,user2_id)VALUES(?,?,?)`, id, first, second) if err == nil { _, err = tx.ExecContext(r.Context(), `INSERT INTO im_conversation_members(conversation_id,user_id)VALUES(?,?),(?,?)`, id, first, id, second) } if err != nil { _ = tx.Rollback() fail(w, 500, 50001, "创建会话失败") return } _ = tx.Commit() reply(w, map[string]any{"id": id}) } func (a *App) conversations(w http.ResponseWriter, r *http.Request) { who := current(r) // A sequence is shared by both senders, so last_seq-read_seq is not an unread count. rows, err := a.db.QueryContext(r.Context(), `SELECT c.id,c.last_seq,c.last_message_at, (SELECT COUNT(*) FROM im_messages unread_msg WHERE unread_msg.conversation_id=c.id AND unread_msg.seq>GREATEST(m.read_seq,m.clear_seq,m.join_seq) AND unread_msg.sender_id<>m.user_id AND unread_msg.recalled_at IS NULL AND unread_msg.admin_removed_at IS NULL),m.pinned,m.muted, other.user_id,p.nickname,p.avatar_url,p.is_vip,p.last_active_at,privacy.online_visible,COALESCE(CAST(msg.body AS CHAR CHARACTER SET utf8mb4),''),COALESCE(msg.message_type,0),msg.recalled_at,msg.admin_removed_at,COALESCE(last_media.status,1),COALESCE((SELECT CAST(MAX(changed.updated_at) AS CHAR) FROM im_message_media changed JOIN im_messages changed_message ON changed_message.id=changed.message_id WHERE changed_message.conversation_id=c.id),'') FROM im_conversation_members m JOIN im_conversations c ON c.id=m.conversation_id JOIN im_conversation_members other ON other.conversation_id=c.id AND other.user_id<>m.user_id JOIN user_profiles p ON p.user_id=other.user_id JOIN user_privacy_settings privacy ON privacy.user_id=other.user_id LEFT JOIN im_messages msg ON msg.id=c.last_message_id LEFT JOIN im_message_media last_media ON last_media.message_id=msg.id WHERE m.user_id=? AND m.status=1 AND c.status=1 ORDER BY m.pinned DESC,c.last_message_at DESC`, who.ID) if err != nil { fail(w, 500, 50001, err.Error()) return } defer rows.Close() items := []map[string]any{} for rows.Next() { var id, lastSeq, unread, otherID int64 var lastAt sql.NullTime var pinned, muted, vip int var messageType, mediaStatus int var nick, avatar, body string var mediaRevision string var active sql.NullTime var recalledAt, adminRemovedAt sql.NullTime var onlineVisible int if err := rows.Scan(&id, &lastSeq, &lastAt, &unread, &pinned, &muted, &otherID, &nick, &avatar, &vip, &active, &onlineVisible, &body, &messageType, &recalledAt, &adminRemovedAt, &mediaStatus, &mediaRevision); err != nil { fail(w, 500, 50001, "读取会话失败") return } preview := "开始聊天吧" var content map[string]any if (messageType == 2 || messageType == 3) && mediaStatus == 0 { if messageType == 2 { preview = "图片已清理" } else { preview = "语音已清理" } } else if recalledAt.Valid || adminRemovedAt.Valid { preview = "消息已撤回" } else if json.Unmarshal([]byte(body), &content) == nil { if text, ok := content["text"].(string); ok { preview = text } else if messageType == 2 { preview = "[图片]" } else if messageType == 3 { preview = "[语音]" } } items = append(items, map[string]any{"id": id, "lastSeq": lastSeq, "unread": unread, "lastMessageAt": nullableTime(lastAt), "mediaRevision": mediaRevision, "pinned": pinned == 1, "muted": muted == 1, "lastMessage": preview, "user": map[string]any{"id": otherID, "nickname": nick, "avatar": avatar, "avatarThumbnail": avatarThumbnailURL(avatar), "vip": vip == 1, "online": onlineVisible == 1 && a.onlineNow(otherID)}}) } if err := rows.Err(); err != nil { fail(w, 500, 50001, "读取会话失败") return } reply(w, map[string]any{"items": items, "total": len(items)}) } func (a *App) messages(w http.ResponseWriter, r *http.Request) { id, err := pathID(r) if err != nil { fail(w, 400, 20001, "invalid conversation") return } if !a.isMember(r, id, current(r).ID) { fail(w, 403, 30002, "不是会话成员") return } beforeSeq, _ := strconv.ParseInt(r.URL.Query().Get("beforeSeq"), 10, 64) afterSeq, _ := strconv.ParseInt(r.URL.Query().Get("afterSeq"), 10, 64) if beforeSeq > 0 && afterSeq > 0 { fail(w, http.StatusBadRequest, 20001, "beforeSeq 和 afterSeq 不能同时使用") return } limit, _ := strconv.Atoi(r.URL.Query().Get("limit")) if limit <= 0 || limit > 100 { limit = 50 } query := `SELECT m.id,m.conversation_id,m.seq,m.sender_id,m.client_msg_id,m.message_type,m.body,m.recalled_at,m.admin_removed_at,m.created_at,COALESCE(mm.status,1) FROM im_messages m LEFT JOIN im_message_media mm ON mm.message_id=m.id WHERE m.conversation_id=?` args := []any{id} if beforeSeq > 0 { query += ` AND m.seq 0 { query += ` AND m.seq>?` args = append(args, afterSeq) query += ` ORDER BY m.seq ASC LIMIT ?` } else { query += ` ORDER BY m.seq DESC LIMIT ?` } args = append(args, limit+1) rows, err := a.db.QueryContext(r.Context(), query, args...) if err != nil { fail(w, 500, 50001, "查询消息失败") return } defer rows.Close() items := []messageView{} for rows.Next() { var item messageView var body []byte var recalledAt, adminRemovedAt sql.NullTime var mediaStatus int if err := rows.Scan(&item.ID, &item.ConversationID, &item.Seq, &item.SenderID, &item.ClientMsgID, &item.Type, &body, &recalledAt, &adminRemovedAt, &item.CreatedAt, &mediaStatus); err != nil { fail(w, 500, 50001, "读取消息失败") return } var content any item.MediaDeleted = (item.Type == 2 || item.Type == 3) && mediaStatus == 0 item.Recalled = recalledAt.Valid || adminRemovedAt.Valid || item.MediaDeleted if !item.Recalled { _ = json.Unmarshal(body, &content) } item.Content = content items = append(items, item) } if err := rows.Err(); err != nil { fail(w, 500, 50001, "读取消息失败") return } _ = rows.Close() hasMore := len(items) > limit if hasMore { items = items[:limit] } if afterSeq <= 0 { sort.Slice(items, func(i, j int) bool { return items[i].Seq < items[j].Seq }) } if len(items) > 0 { last := items[len(items)-1].Seq if _, _, err := a.markConversationRead(r.Context(), id, current(r).ID, last); err != nil { fail(w, 500, 50001, "更新已读状态失败") return } } nextBeforeSeq := int64(0) if len(items) > 0 { nextBeforeSeq = items[0].Seq } nextAfterSeq := afterSeq if len(items) > 0 { nextAfterSeq = items[len(items)-1].Seq } // The chat page shows how many opening messages are left before the other // side answers, so the reader learns the rule before a send is refused. reply(w, map[string]any{"items": items, "hasMore": hasMore, "nextBeforeSeq": nextBeforeSeq, "nextAfterSeq": nextAfterSeq, "unansweredQuota": a.unansweredQuotaView(r.Context(), id, current(r).ID)}) } // unansweredQuotaView reports the opening-message allowance for this // conversation: remaining -1 means "no limit applies", either because the other // side already replied, the reader is a member, or the switch is off. func (a *App) unansweredQuotaView(ctx context.Context, conversationID, userID int64) map[string]any { limit := a.resolveUnansweredMessageLimit(ctx, a.db) unlimited := map[string]any{"limit": limit, "remaining": -1, "unlimited": true} if limit <= 0 { return unlimited } var mine, theirs int if err := a.db.QueryRowContext(ctx, `SELECT COALESCE(SUM(sender_id=?),0),COALESCE(SUM(sender_id<>?),0) FROM im_messages WHERE conversation_id=?`, userID, userID, conversationID).Scan(&mine, &theirs); err != nil { return unlimited } if theirs > 0 || a.hasActiveMembership(ctx, userID) { return unlimited } remaining := limit - mine if remaining < 0 { remaining = 0 } return map[string]any{"limit": limit, "remaining": remaining, "unlimited": false} } func (a *App) sendMessageHTTP(w http.ResponseWriter, r *http.Request) { conversationID, err := pathID(r) if err != nil { fail(w, 400, 20001, "invalid conversation") return } var req struct { ClientMsgID string `json:"clientMsgId"` Type int `json:"type"` Content any `json:"content"` } if decode(r, &req) != nil { fail(w, 400, 20001, "消息格式错误") return } item, members, err := a.persistMessage(r, conversationID, current(r).ID, req.ClientMsgID, req.Type, req.Content) if err != nil { var limitErr *dailyActiveChatLimitError if errors.As(err, &limitErr) { fail(w, http.StatusTooManyRequests, 30005, limitErr.Error()) return } // Its own code and status: the client turns this one into an upgrade // prompt rather than a plain "send failed" toast. var gateErr *messageMembershipGateError if errors.As(err, &gateErr) { fail(w, http.StatusPaymentRequired, 30006, gateErr.Error()) return } var knockErr *unansweredMessageLimitError if errors.As(err, &knockErr) { fail(w, http.StatusPaymentRequired, 30007, knockErr.Error()) return } fail(w, 400, 30004, err.Error()) return } a.hub.broadcast(members, map[string]any{"command": "MESSAGE_PUSH", "data": item}) reply(w, item) } func (a *App) recallMessage(w http.ResponseWriter, r *http.Request) { messageID, err := pathID(r) if err != nil { fail(w, http.StatusBadRequest, 20001, "消息编号无效") return } window, _ := strconv.Atoi(a.configPlain(r.Context(), "im.recall_seconds", "120")) if window < 1 || window > 86400 { window = 120 } tx, err := a.db.BeginTx(r.Context(), nil) if err != nil { fail(w, 500, 50001, "撤回失败") return } defer func() { _ = tx.Rollback() }() var conversationID, seq int64 var senderID int64 var created time.Time var recalled sql.NullTime if err = tx.QueryRowContext(r.Context(), `SELECT conversation_id,seq,sender_id,created_at,recalled_at FROM im_messages WHERE id=? FOR UPDATE`, messageID).Scan(&conversationID, &seq, &senderID, &created, &recalled); err != nil { fail(w, 404, 30001, "消息不存在") return } if senderID != current(r).ID { fail(w, 403, 30002, "只能撤回自己发送的消息") return } if recalled.Valid { reply(w, map[string]bool{"success": true}) return } if time.Since(created) > time.Duration(window)*time.Second { fail(w, 409, 30004, "已超过消息撤回时限") return } _, err = tx.ExecContext(r.Context(), `UPDATE im_messages SET recalled_at=NOW(3) WHERE id=?`, messageID) if err != nil || tx.Commit() != nil { fail(w, 500, 50001, "撤回失败") return } rows, _ := a.db.QueryContext(r.Context(), `SELECT user_id FROM im_conversation_members WHERE conversation_id=? AND status=1`, conversationID) members := []int64{} if rows != nil { defer rows.Close() for rows.Next() { var id int64 _ = rows.Scan(&id) members = append(members, id) } } a.hub.broadcast(members, map[string]any{"command": "MESSAGE_RECALLED", "data": map[string]any{"id": messageID, "conversationId": conversationID, "seq": seq}}) reply(w, map[string]bool{"success": true}) } // The reply worker has no *http.Request, so message writing is expressed in // terms of a context and the HTTP handlers wrap it. func (a *App) persistMessage(r *http.Request, conversationID, senderID int64, clientMsgID string, messageType int, content any) (messageView, []int64, error) { return a.persistMessageContext(r.Context(), conversationID, senderID, clientMsgID, messageType, content) } func (a *App) persistMessageContext(ctx context.Context, conversationID, senderID int64, clientMsgID string, messageType int, content any) (messageView, []int64, error) { item := messageView{} clientMsgID, content, body, err := a.validateMessagePayload(clientMsgID, messageType, content) if err != nil { return item, nil, err } if messageType == 0 { messageType = 1 } if err = a.ensureMessageMembership(ctx, senderID, messageType); err != nil { return item, nil, err } // 词库对文字消息生效:拦截词直接发不出去,转人工的词放行但记一条风险事件, // 免得把用户挡在门外却没人知道发生过什么。 if messageType == 1 { if word, action := a.matchKeyword(ctx, "message", string(body)); word != "" { if action == "block" { return item, nil, fmt.Errorf("消息包含平台不允许的内容") } a.recordKeywordRisk(ctx, senderID, "message", word) } } var mediaAssetID int64 if messageType == 2 || messageType == 3 { contentMap, _ := content.(map[string]any) mediaURL, _ := contentMap["url"].(string) expectedType := "image" if messageType == 3 { expectedType = "audio" } if queryErr := a.db.QueryRowContext(ctx, `SELECT id FROM media_assets WHERE owner_user_id=? AND public_url=? AND media_type=? AND status=1 AND moderation_status=1 ORDER BY id DESC LIMIT 1`, senderID, mediaURL, expectedType).Scan(&mediaAssetID); queryErr != nil { return item, nil, fmt.Errorf("消息媒体必须由当前账号上传") } } if !a.allowRequest(ctx, "message_send", fmt.Sprintf("%d", senderID), 120, time.Minute) { return item, nil, fmt.Errorf("消息发送过于频繁,请稍后再试") } if a.isSanctionActive(ctx, senderID, "MUTE") { return item, nil, fmt.Errorf("账号处于禁言期,暂时无法发送消息") } tx, err := a.db.BeginTx(ctx, nil) if err != nil { return item, nil, err } defer func() { _ = tx.Rollback() }() var lastSeq int64 if err = tx.QueryRowContext(ctx, `SELECT last_seq FROM im_conversations WHERE id=? AND status=1 FOR UPDATE`, conversationID).Scan(&lastSeq); err != nil { return item, nil, fmt.Errorf("会话不存在") } var memberCount int if err = tx.QueryRowContext(ctx, `SELECT COUNT(*) FROM im_conversation_members WHERE conversation_id=? AND user_id=? AND status=1`, conversationID, senderID).Scan(&memberCount); err != nil || memberCount == 0 { return item, nil, fmt.Errorf("不是会话成员") } var blocked int if err = tx.QueryRowContext(ctx, `SELECT EXISTS( SELECT 1 FROM im_direct_conversations d JOIN user_blocks b ON (b.user_id=d.user1_id AND b.blocked_user_id=d.user2_id) OR (b.user_id=d.user2_id AND b.blocked_user_id=d.user1_id) WHERE d.conversation_id=?)`, conversationID).Scan(&blocked); err != nil { return item, nil, err } if blocked == 1 { return item, nil, fmt.Errorf("当前无法向该用户发送消息") } if err = a.reserveDailyActiveChat(ctx, tx, conversationID, senderID); err != nil { return item, nil, err } // The AI worker answers on behalf of a managed account; its reply is the // answer, never an unanswered knock. if !isAIGenerated(ctx) { if err = a.reserveUnansweredMessage(ctx, tx, conversationID, senderID); err != nil { return item, nil, err } } seq := lastSeq + 1 result, err := tx.ExecContext(ctx, `INSERT INTO im_messages(conversation_id,seq,sender_id,client_msg_id,message_type,body)VALUES(?,?,?,?,?,?)`, conversationID, seq, senderID, clientMsgID, messageType, body) if err != nil { var existingID, existingSeq int64 existingErr := tx.QueryRowContext(ctx, `SELECT id,seq FROM im_messages WHERE sender_id=? AND client_msg_id=?`, senderID, clientMsgID).Scan(&existingID, &existingSeq) if existingErr == nil { _ = tx.Rollback() return a.loadMessageContext(ctx, existingID), nil, nil } return item, nil, err } messageID, _ := result.LastInsertId() if mediaAssetID > 0 { contentMap, _ := content.(map[string]any) mediaURL, _ := contentMap["url"].(string) mediaType := "image" var durationMS any if messageType == 3 { mediaType = "voice" duration, _ := numericDuration(contentMap["duration"]) durationMS = duration * 1000 } if _, err = tx.ExecContext(ctx, `INSERT INTO im_message_media(message_id,media_asset_id,media_type,public_url,duration_ms,status) VALUES(?,?,?,?,?,1)`, messageID, mediaAssetID, mediaType, mediaURL, durationMS); err != nil { return item, nil, err } } _, err = tx.ExecContext(ctx, `UPDATE im_conversations SET last_seq=?,last_message_id=?,last_message_at=NOW(3) WHERE id=?`, seq, messageID, conversationID) if err != nil { return item, nil, err } if _, err = tx.ExecContext(ctx, `UPDATE im_conversation_members SET delivered_seq=GREATEST(delivered_seq,?),updated_at=NOW(3) WHERE conversation_id=?`, seq, conversationID); err != nil { return item, nil, err } memberRows, err := tx.QueryContext(ctx, `SELECT user_id FROM im_conversation_members WHERE conversation_id=? AND status=1`, conversationID) if err != nil { return item, nil, err } members := []int64{} for memberRows.Next() { var userID int64 if err = memberRows.Scan(&userID); err != nil { _ = memberRows.Close() return item, nil, err } members = append(members, userID) } if err = memberRows.Err(); err != nil { _ = memberRows.Close() return item, nil, err } if err = memberRows.Close(); err != nil { return item, nil, err } sort.Slice(members, func(i, j int) bool { return members[i] < members[j] }) for _, userID := range members { var lockedUserID int64 if err = tx.QueryRowContext(ctx, `SELECT id FROM users WHERE id=? FOR UPDATE`, userID).Scan(&lockedUserID); err != nil { return item, nil, err } } for _, userID := range members { var next int64 if err = tx.QueryRowContext(ctx, `SELECT COALESCE(MAX(event_seq),0)+1 FROM im_user_sync_events WHERE user_id=?`, userID).Scan(&next); err != nil { return item, nil, err } if _, err = tx.ExecContext(ctx, `INSERT INTO im_user_sync_events(user_id,event_seq,event_type,conversation_id,message_seq,event_data)VALUES(?, ?,12,?,?,?)`, userID, next, conversationID, seq, body); err != nil { return item, nil, err } } if err = tx.Commit(); err != nil { return item, nil, err } a.enqueueAIReply(ctx, conversationID, senderID, members, messageID) item = messageView{ID: messageID, ConversationID: conversationID, Seq: seq, SenderID: senderID, ClientMsgID: clientMsgID, Type: messageType, Content: content, CreatedAt: time.Now()} // Everyone connected gets this over the socket; this is for the ones who are // not, and it runs on its own so a slow push service cannot delay the send. a.notifyOfflineMessage(conversationID, senderID, members, item) return item, members, nil } func (a *App) loadMessage(r *http.Request, id int64) messageView { return a.loadMessageContext(r.Context(), id) } func (a *App) loadMessageContext(ctx context.Context, id int64) messageView { var item messageView var body []byte var recalledAt, adminRemovedAt sql.NullTime var mediaStatus int _ = a.db.QueryRowContext(ctx, `SELECT m.id,m.conversation_id,m.seq,m.sender_id,m.client_msg_id,m.message_type,m.body,m.recalled_at,m.admin_removed_at,m.created_at,COALESCE(mm.status,1) FROM im_messages m LEFT JOIN im_message_media mm ON mm.message_id=m.id WHERE m.id=?`, id).Scan(&item.ID, &item.ConversationID, &item.Seq, &item.SenderID, &item.ClientMsgID, &item.Type, &body, &recalledAt, &adminRemovedAt, &item.CreatedAt, &mediaStatus) item.MediaDeleted = (item.Type == 2 || item.Type == 3) && mediaStatus == 0 item.Recalled = recalledAt.Valid || adminRemovedAt.Valid || item.MediaDeleted if !item.Recalled { _ = json.Unmarshal(body, &item.Content) } return item } func (a *App) isMember(r *http.Request, conversationID, userID int64) bool { var count int _ = a.db.QueryRowContext(r.Context(), `SELECT COUNT(*) FROM im_conversation_members WHERE conversation_id=? AND user_id=? AND status=1`, conversationID, userID).Scan(&count) return count > 0 } func (a *App) conversationSettings(w http.ResponseWriter, r *http.Request) { id, pathErr := pathID(r) if pathErr != nil { fail(w, http.StatusBadRequest, 20001, "会话 ID 无效") return } var req struct { Pinned *bool `json:"pinned"` Muted *bool `json:"muted"` ReadSeq int64 `json:"readSeq"` } if decode(r, &req) != nil || req.ReadSeq < 0 { fail(w, 400, 20001, "invalid settings") return } pinned, muted := -1, -1 if req.Pinned != nil { pinned = btoi(*req.Pinned) } if req.Muted != nil { muted = btoi(*req.Muted) } var lastSeq int64 if err := a.db.QueryRowContext(r.Context(), `SELECT c.last_seq FROM im_conversations c JOIN im_conversation_members m ON m.conversation_id=c.id WHERE c.id=? AND m.user_id=? AND m.status=1`, id, current(r).ID).Scan(&lastSeq); err != nil { fail(w, http.StatusForbidden, 30002, "不是会话成员") return } readSeq := req.ReadSeq if readSeq > lastSeq { readSeq = lastSeq } _, err := a.db.ExecContext(r.Context(), `UPDATE im_conversation_members SET pinned=IF(?>=0,?,pinned),muted=IF(?>=0,?,muted),read_seq=GREATEST(read_seq,?) WHERE conversation_id=? AND user_id=?`, pinned, pinned, muted, muted, readSeq, id, current(r).ID) if err != nil { fail(w, 500, 50001, "保存失败") return } savedReadSeq, unread, err := a.conversationReadState(r.Context(), id, current(r).ID) if err != nil { fail(w, 500, 50001, "读取已读状态失败") return } if req.ReadSeq > 0 { a.publishConversationRead(id, current(r).ID, savedReadSeq, unread) } reply(w, map[string]any{"success": true, "readSeq": savedReadSeq, "unread": unread}) } // Notify this user's other devices after a successful read; never notify the peer // or advance the sender's read marker merely because they sent a message. func (a *App) publishConversationRead(conversationID, userID, readSeq, unread int64) { a.hub.broadcast([]int64{userID}, map[string]any{"command": "READ_ACK", "data": map[string]any{ "conversationId": conversationID, "readSeq": readSeq, "unread": unread, }}) } func (a *App) conversationReadState(ctx context.Context, conversationID, userID int64) (int64, int64, error) { var savedReadSeq, lastSeq, unread int64 err := a.db.QueryRowContext(ctx, `SELECT m.read_seq,c.last_seq, (SELECT COUNT(*) FROM im_messages unread_msg WHERE unread_msg.conversation_id=m.conversation_id AND unread_msg.seq>GREATEST(m.read_seq,m.clear_seq,m.join_seq) AND unread_msg.sender_id<>m.user_id AND unread_msg.recalled_at IS NULL AND unread_msg.admin_removed_at IS NULL) FROM im_conversation_members m JOIN im_conversations c ON c.id=m.conversation_id WHERE m.conversation_id=? AND m.user_id=? AND m.status=1`, conversationID, userID).Scan(&savedReadSeq, &lastSeq, &unread) // Older clients compare READ_ACK with the conversation sequence, which also // includes the user's own messages. Cover those sequences only when the // authoritative unread count is zero; a newer incoming message stays safe. if err == nil && unread == 0 && lastSeq > savedReadSeq { savedReadSeq = lastSeq } return savedReadSeq, unread, err } func (a *App) markConversationRead(ctx context.Context, conversationID, userID, readSeq int64) (int64, int64, error) { if readSeq < 0 { return 0, 0, fmt.Errorf("已读序号无效") } _, err := a.db.ExecContext(ctx, `UPDATE im_conversation_members m JOIN im_conversations c ON c.id=m.conversation_id SET m.read_seq=GREATEST(m.read_seq,LEAST(?,c.last_seq)) WHERE m.conversation_id=? AND m.user_id=? AND m.status=1`, readSeq, conversationID, userID) if err != nil { return 0, 0, err } savedReadSeq, unread, err := a.conversationReadState(ctx, conversationID, userID) if err != nil { return 0, 0, err } a.publishConversationRead(conversationID, userID, savedReadSeq, unread) return savedReadSeq, unread, nil } func (a *App) websocket(w http.ResponseWriter, r *http.Request) { if !a.originAllowed(r.Header.Get("Origin")) { fail(w, http.StatusForbidden, 10006, "WebSocket 来源不允许") return } selectedProtocol := "" raw := "" for _, protocol := range websocket.Subprotocols(r) { if strings.HasPrefix(protocol, "xingyu.jwt.") { selectedProtocol = protocol raw = strings.TrimPrefix(protocol, "xingyu.jwt.") break } } // Query-token compatibility is development-only because URLs may be written to proxy logs. if raw == "" && a.config.Environment != "production" { raw = r.URL.Query().Get("token") } who, err := a.parseToken(raw) if err != nil || who.Role != "user" { fail(w, 401, 10001, "invalid token") return } var userStatus int statusErr := a.db.QueryRowContext(r.Context(), `SELECT status FROM users WHERE id=? AND deleted_at IS NULL`, who.ID).Scan(&userStatus) if statusErr == nil { userStatus = a.normalizeUserStatus(r.Context(), who.ID, userStatus) } if statusErr != nil || userStatus != 1 { fail(w, 403, 10006, "账号已被冻结或封禁") return } var tokenVersion int _ = a.db.QueryRowContext(r.Context(), `SELECT token_version FROM user_security_controls WHERE user_id=?`, who.ID).Scan(&tokenVersion) if who.Version != tokenVersion { fail(w, 401, 10001, "登录状态已失效") return } upgrader := websocket.Upgrader{CheckOrigin: func(request *http.Request) bool { return a.originAllowed(request.Header.Get("Origin")) }} if selectedProtocol != "" { upgrader.Subprotocols = []string{selectedProtocol} } conn, err := upgrader.Upgrade(w, r, nil) if err != nil { return } client := &wsClient{userID: who.ID, conn: conn} conn.SetReadLimit(64 << 10) _ = conn.SetReadDeadline(time.Now().Add(75 * time.Second)) a.touchUserActivity(r.Context(), who.ID) if a.hub.add(client) { a.publishPresence(who.ID, true) } defer func() { if a.hub.remove(client) { a.publishPresence(who.ID, false) } _ = conn.Close() }() _ = client.writeJSON(map[string]any{"command": "AUTH_ACK", "data": map[string]any{"heartbeatSeconds": 25, "serverTime": time.Now().UnixMilli()}}) for { var frame struct { Command string `json:"command"` ConversationID int64 `json:"conversationId"` ClientMsgID string `json:"clientMsgId"` Type int `json:"type"` Content any `json:"content"` ReadSeq int64 `json:"readSeq"` } if conn.ReadJSON(&frame) != nil { return } switch frame.Command { case "PING": a.touchUserActivity(r.Context(), who.ID) _ = conn.SetReadDeadline(time.Now().Add(75 * time.Second)) _ = client.writeJSON(map[string]any{"command": "PONG", "timestamp": time.Now().UnixMilli()}) case "SEND_MESSAGE": item, members, persistErr := a.persistMessage(r, frame.ConversationID, who.ID, frame.ClientMsgID, frame.Type, frame.Content) if persistErr != nil { _ = client.writeJSON(map[string]any{"command": "ERROR", "message": persistErr.Error()}) continue } a.hub.broadcast(members, map[string]any{"command": "MESSAGE_PUSH", "data": item}) case "READ": if frame.ReadSeq < 0 { _ = client.writeJSON(map[string]any{"command": "ERROR", "message": "已读序号无效"}) continue } var lastSeq int64 if a.db.QueryRowContext(r.Context(), `SELECT c.last_seq FROM im_conversations c JOIN im_conversation_members m ON m.conversation_id=c.id WHERE c.id=? AND m.user_id=? AND m.status=1`, frame.ConversationID, who.ID).Scan(&lastSeq) != nil { _ = client.writeJSON(map[string]any{"command": "ERROR", "message": "不是会话成员"}) continue } if frame.ReadSeq > lastSeq { frame.ReadSeq = lastSeq } if _, _, err := a.markConversationRead(r.Context(), frame.ConversationID, who.ID, frame.ReadSeq); err != nil { _ = client.writeJSON(map[string]any{"command": "ERROR", "message": "更新已读状态失败"}) continue } } } }