Files
kefu/im/backend/internal/app/im.go
T
Your NameandClaude Opus 5 acd8933dbb 后端:资料距离、语音图片会员限制、免短信注册、AI 托管回复修复
- 单用户资料接口补上距离:此前只有推荐/附近列表会算距离,资料页和
  聊天头部因此无内容可显示。沿用同一套 haversine 与隐私开关。
- 语音/图片消息可限定会员发送,两个开关在管理端「运营配置」中修改
  (迁移 032)。校验放在 persistMessageContext,HTTP 与 WebSocket
  两条发送路径都覆盖;文本消息永不受限。
- 短信服务关闭时注册不再要求验证码:关掉之后没人能拿到验证码,继续
  要求就等于关闭注册通道。重置密码不做同样放宽,那里缺验证码等于
  凭手机号夺号。app/config 增加 smsVerification 供客户端决定表单形态。
- 修复 AI 托管账号之间不回复:原规则按「发送方是否托管账号」拦截,
  把真人操作测试号的正常对话也挡了。改为标记 worker 自己写入的回复,
  只对 AI 生成的消息跳过入队。
- ai.default_model_id 同时接受模型 ID 与名称,填名称时不再被 MySQL
  静默转成 0 而使配置失效。
- 聊天媒体留存管理与清理任务(迁移 033,两台线上均已应用)。

新增集成测试均针对真实 MySQL:会员限制、免短信注册、AI 入队规则、
默认模型解析、资料距离。

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

914 lines
33 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
}
reply(w, map[string]any{"items": items, "hasMore": hasMore, "nextBeforeSeq": nextBeforeSeq, "nextAfterSeq": nextAfterSeq})
}
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
}
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
}
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)
return messageView{ID: messageID, ConversationID: conversationID, Seq: seq, SenderID: senderID, ClientMsgID: clientMsgID, Type: messageType, Content: content, CreatedAt: time.Now()}, 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
}
}
}
}