Files
kefu/im/backend/internal/app/im.go
T
Your NameandClaude Opus 5 ce429afcf2 密码强度下限、离线推送下发
一、密码规则:至少 8 位且不能全是数字。短信验证可以在后台关掉,关掉之后
密码就是账号质量的唯一门槛,而此前 validUserPassword 只检查非空——"1"
也是合法密码。规则只在「设置密码」时校验(注册、找回、改密、后台建号与
后台重置),登录不再校验,已有账号照常使用。上下限都按字符数计算,否则
43 个汉字的密码会因为字节数超限被拒。各处错误提示改为直接说明规则,
而不是笼统的一句「请填写有效的密码」。

二、离线推送:此前客户端一直在上报 push token,服务端从未下发过任何东西,
App 退到后台或被杀掉时新消息完全没有提醒(MESSAGE_PUSH 只是 WebSocket
帧名)。补上服务端下发:
- 只发给「此刻不在线 + 未对该会话免打扰 + 未关闭消息通知」的接收者,
  在线的人已经从实时通道拿到了。
- 鉴权 token 按 provider 的过期时间缓存,个推的 auth 接口限流很紧。
- 失效的 cid(10001/10002)就地停用,不再每条消息重试一次。
- 整个过程在独立 goroutine 与独立 context 上进行,推送服务再慢也不会
  拖慢或拖垮一条已经发出的消息。
- 是否显示正文由 push.show_preview 控制,关闭后锁屏上不出现消息内容。
凭据在管理端「离线推送」页填写(迁移 035 先建出配置行——集成配置保存
走的是 UPDATE,行不存在会静默保存不上)。

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
2026-09-04 19:09:33 +08:00

957 lines
35 KiB
Go

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<?`
args = append(args, beforeSeq)
}
if afterSeq > 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
}
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
}
}
}
}