后端接入 AI 模型并支持测试账号托管回复
按协议而不是按厂商做适配,与现有短信、对象存储的做法一致:openai 兼容格式 一个适配器即可覆盖 DeepSeek、通义、智谱、火山方舟、Ollama 等,新增厂商通常 只需在后台加一条模型记录;anthropic、gemini、webhook、debug 各一个适配器。 密钥沿用 integration.go 的 AES-GCM 加密存储。 测试账号托管的回复走 persistMessage 同一条落库和推送路径,因此未读数、 WebSocket 推送和会话排序全部复用现有逻辑,客户端无需改动。为此把 persistMessage 抽出 persistMessageContext,因为 worker 没有 *http.Request。 任务队列用数据库表而非内存:重启不丢回复,多实例不重复消费。领取用 UPDATE 打 claim_token 再回读,没有用 SELECT ... FOR UPDATE SKIP LOCKED——生产存在 MySQL 5.7 环境,那里该语法无法解析。同一会话同时只允许一个待处理任务 (pending_key 生成列 + 唯一键),所以用户连发多条消息只会得到一条回复。 闸门:仅 is_test=1 且在允许批次内的账号生效,托管账号之间不互相触发,三层 配额,调用失败默认静默。会话内是否标注 AI 身份由 ai.disclose_in_chat 控制 并默认开启,关闭前需确认所在地区的监管要求。 app.go、integration.go、im.go 三个文件同时包含本次改动之前工作区里就已存在 的未提交修改,一并带入。 Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 5
parent
334890171e
commit
a024d59827
+533
-73
@@ -1,14 +1,20 @@
|
||||
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"
|
||||
)
|
||||
@@ -21,15 +27,107 @@ type messageView struct {
|
||||
ClientMsgID string `json:"clientMsgId"`
|
||||
Type int `json:"type"`
|
||||
Content any `json:"content"`
|
||||
Recalled bool `json:"recalled"`
|
||||
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{}
|
||||
@@ -37,22 +135,47 @@ type Hub struct {
|
||||
|
||||
func NewHub() *Hub { return &Hub{clients: make(map[int64]map[*wsClient]struct{})} }
|
||||
|
||||
func (h *Hub) add(client *wsClient) {
|
||||
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) {
|
||||
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) {
|
||||
@@ -65,9 +188,7 @@ func (h *Hub) broadcast(userIDs []int64, payload any) {
|
||||
}
|
||||
h.mu.RUnlock()
|
||||
for _, client := range targets {
|
||||
client.mu.Lock()
|
||||
_ = client.conn.WriteJSON(payload)
|
||||
client.mu.Unlock()
|
||||
_ = client.writeJSON(payload)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -86,6 +207,41 @@ func (h *Hub) disconnect(userID int64) {
|
||||
}
|
||||
}
|
||||
|
||||
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"`
|
||||
@@ -98,16 +254,47 @@ func (a *App) directConversation(w http.ResponseWriter, r *http.Request) {
|
||||
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
|
||||
}
|
||||
tx, _ := a.db.BeginTx(r.Context(), nil)
|
||||
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
|
||||
}
|
||||
@@ -127,12 +314,16 @@ func (a *App) directConversation(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
func (a *App) conversations(w http.ResponseWriter, r *http.Request) {
|
||||
who := current(r)
|
||||
rows, err := a.db.QueryContext(r.Context(), `SELECT c.id,c.last_seq,c.last_message_at,m.read_seq,m.pinned,m.muted,
|
||||
other.user_id,p.nickname,p.avatar_url,p.is_vip,p.last_active_at,COALESCE(CAST(msg.body AS CHAR CHARACTER SET utf8mb4),'')
|
||||
// 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),''),msg.recalled_at,msg.admin_removed_at
|
||||
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 LEFT JOIN im_messages msg ON msg.id=c.last_message_id
|
||||
WHERE m.user_id=? AND m.status=1 ORDER BY m.pinned DESC,c.last_message_at DESC`, who.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
|
||||
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
|
||||
@@ -140,20 +331,31 @@ func (a *App) conversations(w http.ResponseWriter, r *http.Request) {
|
||||
defer rows.Close()
|
||||
items := []map[string]any{}
|
||||
for rows.Next() {
|
||||
var id, lastSeq, readSeq, otherID int64
|
||||
var id, lastSeq, unread, otherID int64
|
||||
var lastAt sql.NullTime
|
||||
var pinned, muted, vip int
|
||||
var nick, avatar, body string
|
||||
var active sql.NullTime
|
||||
_ = rows.Scan(&id, &lastSeq, &lastAt, &readSeq, &pinned, &muted, &otherID, &nick, &avatar, &vip, &active, &body)
|
||||
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, &recalledAt, &adminRemovedAt); err != nil {
|
||||
fail(w, 500, 50001, "读取会话失败")
|
||||
return
|
||||
}
|
||||
preview := "开始聊天吧"
|
||||
var content map[string]any
|
||||
if json.Unmarshal([]byte(body), &content) == nil {
|
||||
if recalledAt.Valid || adminRemovedAt.Valid {
|
||||
preview = "消息已撤回"
|
||||
} else if json.Unmarshal([]byte(body), &content) == nil {
|
||||
if text, ok := content["text"].(string); ok {
|
||||
preview = text
|
||||
}
|
||||
}
|
||||
items = append(items, map[string]any{"id": id, "lastSeq": lastSeq, "unread": max64(lastSeq-readSeq, 0), "lastMessageAt": lastAt, "pinned": pinned == 1, "muted": muted == 1, "lastMessage": preview, "user": map[string]any{"id": otherID, "nickname": nick, "avatar": avatar, "vip": vip == 1, "online": active.Valid && time.Since(active.Time) < 15*time.Minute}})
|
||||
items = append(items, map[string]any{"id": id, "lastSeq": lastSeq, "unread": unread, "lastMessageAt": nullableTime(lastAt), "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)})
|
||||
}
|
||||
@@ -168,7 +370,31 @@ func (a *App) messages(w http.ResponseWriter, r *http.Request) {
|
||||
fail(w, 403, 30002, "不是会话成员")
|
||||
return
|
||||
}
|
||||
rows, err := a.db.QueryContext(r.Context(), `SELECT id,conversation_id,seq,sender_id,client_msg_id,message_type,body,created_at FROM im_messages WHERE conversation_id=? AND recalled_at IS NULL ORDER BY seq DESC LIMIT 100`, id)
|
||||
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 id,conversation_id,seq,sender_id,client_msg_id,message_type,body,recalled_at,admin_removed_at,created_at FROM im_messages WHERE conversation_id=?`
|
||||
args := []any{id}
|
||||
if beforeSeq > 0 {
|
||||
query += ` AND seq<?`
|
||||
args = append(args, beforeSeq)
|
||||
}
|
||||
if afterSeq > 0 {
|
||||
query += ` AND seq>?`
|
||||
args = append(args, afterSeq)
|
||||
query += ` ORDER BY seq ASC LIMIT ?`
|
||||
} else {
|
||||
query += ` ORDER BY 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
|
||||
@@ -178,18 +404,47 @@ func (a *App) messages(w http.ResponseWriter, r *http.Request) {
|
||||
for rows.Next() {
|
||||
var item messageView
|
||||
var body []byte
|
||||
_ = rows.Scan(&item.ID, &item.ConversationID, &item.Seq, &item.SenderID, &item.ClientMsgID, &item.Type, &body, &item.CreatedAt)
|
||||
var recalledAt, adminRemovedAt sql.NullTime
|
||||
if err := rows.Scan(&item.ID, &item.ConversationID, &item.Seq, &item.SenderID, &item.ClientMsgID, &item.Type, &body, &recalledAt, &adminRemovedAt, &item.CreatedAt); err != nil {
|
||||
fail(w, 500, 50001, "读取消息失败")
|
||||
return
|
||||
}
|
||||
var content any
|
||||
_ = json.Unmarshal(body, &content)
|
||||
item.Recalled = recalledAt.Valid || adminRemovedAt.Valid
|
||||
if !item.Recalled {
|
||||
_ = json.Unmarshal(body, &content)
|
||||
}
|
||||
item.Content = content
|
||||
items = append(items, item)
|
||||
}
|
||||
sort.Slice(items, func(i, j int) bool { return items[i].Seq < items[j].Seq })
|
||||
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
|
||||
_, _ = a.db.ExecContext(r.Context(), `UPDATE im_conversation_members SET read_seq=GREATEST(read_seq,?) WHERE conversation_id=? AND user_id=?`, last, id, current(r).ID)
|
||||
if _, _, err := a.markConversationRead(r.Context(), id, current(r).ID, last); err != nil {
|
||||
fail(w, 500, 50001, "更新已读状态失败")
|
||||
return
|
||||
}
|
||||
}
|
||||
reply(w, map[string]any{"items": items, "hasMore": false})
|
||||
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) {
|
||||
@@ -207,16 +462,6 @@ func (a *App) sendMessageHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
fail(w, 400, 20001, "消息格式错误")
|
||||
return
|
||||
}
|
||||
if req.ClientMsgID == "" {
|
||||
req.ClientMsgID = fmt.Sprintf("%026d", time.Now().UnixNano())
|
||||
}
|
||||
if len(req.ClientMsgID) > 64 {
|
||||
fail(w, 400, 20001, "客户端消息 ID 长度不能超过 64 个字符")
|
||||
return
|
||||
}
|
||||
if req.Type == 0 {
|
||||
req.Type = 1
|
||||
}
|
||||
item, members, err := a.persistMessage(r, conversationID, current(r).ID, req.ClientMsgID, req.Type, req.Content)
|
||||
if err != nil {
|
||||
var limitErr *dailyActiveChatLimitError
|
||||
@@ -231,78 +476,195 @@ func (a *App) sendMessageHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
reply(w, item)
|
||||
}
|
||||
|
||||
func (a *App) persistMessage(r *http.Request, conversationID, senderID int64, clientMsgID string, messageType int, content any) (messageView, []int64, error) {
|
||||
item := messageView{}
|
||||
if !a.allowRequest(r.Context(), "message_send", fmt.Sprintf("%d", senderID), 120, time.Minute) {
|
||||
return item, nil, fmt.Errorf("消息发送过于频繁,请稍后再试")
|
||||
func (a *App) recallMessage(w http.ResponseWriter, r *http.Request) {
|
||||
messageID, err := pathID(r)
|
||||
if err != nil {
|
||||
fail(w, http.StatusBadRequest, 20001, "消息编号无效")
|
||||
return
|
||||
}
|
||||
if a.isSanctionActive(r.Context(), senderID, "MUTE") {
|
||||
return item, nil, fmt.Errorf("账号处于禁言期,暂时无法发送消息")
|
||||
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 messageType == 2 || messageType == 3 {
|
||||
contentMap, _ := content.(map[string]any)
|
||||
mediaURL, _ := contentMap["url"].(string)
|
||||
expectedType := "image"
|
||||
if messageType == 3 {
|
||||
expectedType = "audio"
|
||||
}
|
||||
var owned int
|
||||
if queryErr := a.db.QueryRowContext(ctx, `SELECT EXISTS(SELECT 1 FROM media_assets WHERE owner_user_id=? AND public_url=? AND media_type=? AND status=1 AND moderation_status=1)`, senderID, mediaURL, expectedType).Scan(&owned); queryErr != nil || owned != 1 {
|
||||
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(r.Context(), `SELECT last_seq FROM im_conversations WHERE id=? AND status=1 FOR UPDATE`, conversationID).Scan(&lastSeq); err != nil {
|
||||
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(r.Context(), `SELECT COUNT(*) FROM im_conversation_members WHERE conversation_id=? AND user_id=? AND status=1`, conversationID, senderID).Scan(&memberCount); err != nil || memberCount == 0 {
|
||||
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("不是会话成员")
|
||||
}
|
||||
if err = a.reserveDailyActiveChat(r.Context(), tx, conversationID, senderID); err != nil {
|
||||
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
|
||||
}
|
||||
body, err := json.Marshal(content)
|
||||
if err != nil {
|
||||
return item, nil, fmt.Errorf("消息内容错误")
|
||||
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(r.Context(), `INSERT INTO im_messages(conversation_id,seq,sender_id,client_msg_id,message_type,body)VALUES(?,?,?,?,?,?)`, conversationID, seq, senderID, clientMsgID, messageType, body)
|
||||
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(r.Context(), `SELECT id,seq FROM im_messages WHERE sender_id=? AND client_msg_id=?`, senderID, clientMsgID).Scan(&existingID, &existingSeq)
|
||||
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.loadMessage(r, existingID), nil, nil
|
||||
return a.loadMessageContext(ctx, existingID), nil, nil
|
||||
}
|
||||
return item, nil, err
|
||||
}
|
||||
messageID, _ := result.LastInsertId()
|
||||
_, err = tx.ExecContext(r.Context(), `UPDATE im_conversations SET last_seq=?,last_message_id=?,last_message_at=NOW(3) WHERE id=?`, seq, messageID, conversationID)
|
||||
_, 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
|
||||
}
|
||||
_, _ = tx.ExecContext(r.Context(), `UPDATE im_conversation_members SET delivered_seq=GREATEST(delivered_seq,?),updated_at=NOW(3) WHERE conversation_id=?`, seq, conversationID)
|
||||
memberRows, err := tx.QueryContext(r.Context(), `SELECT user_id FROM im_conversation_members WHERE conversation_id=? AND status=1`, conversationID)
|
||||
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
|
||||
_ = memberRows.Scan(&userID)
|
||||
if err = memberRows.Scan(&userID); err != nil {
|
||||
_ = memberRows.Close()
|
||||
return item, nil, err
|
||||
}
|
||||
members = append(members, userID)
|
||||
}
|
||||
_ = memberRows.Close()
|
||||
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
|
||||
_ = tx.QueryRowContext(r.Context(), `SELECT COALESCE(MAX(event_seq),0)+1 FROM im_user_sync_events WHERE user_id=?`, userID).Scan(&next)
|
||||
_, _ = tx.ExecContext(r.Context(), `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)
|
||||
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
|
||||
_ = a.db.QueryRowContext(r.Context(), `SELECT id,conversation_id,seq,sender_id,client_msg_id,message_type,body,created_at FROM im_messages WHERE id=?`, id).Scan(&item.ID, &item.ConversationID, &item.Seq, &item.SenderID, &item.ClientMsgID, &item.Type, &body, &item.CreatedAt)
|
||||
_ = json.Unmarshal(body, &item.Content)
|
||||
var recalledAt, adminRemovedAt sql.NullTime
|
||||
_ = a.db.QueryRowContext(ctx, `SELECT id,conversation_id,seq,sender_id,client_msg_id,message_type,body,recalled_at,admin_removed_at,created_at FROM im_messages WHERE id=?`, id).Scan(&item.ID, &item.ConversationID, &item.Seq, &item.SenderID, &item.ClientMsgID, &item.Type, &body, &recalledAt, &adminRemovedAt, &item.CreatedAt)
|
||||
item.Recalled = recalledAt.Valid || adminRemovedAt.Valid
|
||||
if !item.Recalled {
|
||||
_ = json.Unmarshal(body, &item.Content)
|
||||
}
|
||||
return item
|
||||
}
|
||||
func (a *App) isMember(r *http.Request, conversationID, userID int64) bool {
|
||||
@@ -312,13 +674,17 @@ func (a *App) isMember(r *http.Request, conversationID, userID int64) bool {
|
||||
}
|
||||
|
||||
func (a *App) conversationSettings(w http.ResponseWriter, r *http.Request) {
|
||||
id, _ := pathID(r)
|
||||
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 {
|
||||
if decode(r, &req) != nil || req.ReadSeq < 0 {
|
||||
fail(w, 400, 20001, "invalid settings")
|
||||
return
|
||||
}
|
||||
@@ -329,12 +695,72 @@ func (a *App) conversationSettings(w http.ResponseWriter, r *http.Request) {
|
||||
if req.Muted != nil {
|
||||
muted = btoi(*req.Muted)
|
||||
}
|
||||
_, 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, req.ReadSeq, id, current(r).ID)
|
||||
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
|
||||
}
|
||||
reply(w, map[string]bool{"success": true})
|
||||
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) {
|
||||
@@ -342,7 +768,19 @@ func (a *App) websocket(w http.ResponseWriter, r *http.Request) {
|
||||
fail(w, http.StatusForbidden, 10006, "WebSocket 来源不允许")
|
||||
return
|
||||
}
|
||||
raw := r.URL.Query().Get("token")
|
||||
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")
|
||||
@@ -366,14 +804,27 @@ func (a *App) websocket(w http.ResponseWriter, r *http.Request) {
|
||||
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}
|
||||
a.hub.add(client)
|
||||
defer func() { a.hub.remove(client); _ = conn.Close() }()
|
||||
_ = conn.WriteJSON(map[string]any{"command": "AUTH_ACK", "data": map[string]any{"heartbeatSeconds": 25, "serverTime": time.Now().UnixMilli()}})
|
||||
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"`
|
||||
@@ -388,24 +839,33 @@ func (a *App) websocket(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
switch frame.Command {
|
||||
case "PING":
|
||||
_ = conn.WriteJSON(map[string]any{"command": "PONG", "timestamp": time.Now().UnixMilli()})
|
||||
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 {
|
||||
_ = conn.WriteJSON(map[string]any{"command": "ERROR", "message": persistErr.Error()})
|
||||
_ = 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":
|
||||
_, _ = a.db.ExecContext(r.Context(), `UPDATE im_conversation_members SET read_seq=GREATEST(read_seq,?) WHERE conversation_id=? AND user_id=?`, frame.ReadSeq, frame.ConversationID, who.ID)
|
||||
a.hub.broadcast([]int64{who.ID}, map[string]any{"command": "READ_ACK", "data": frame})
|
||||
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
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func max64(a, b int64) int64 {
|
||||
if a > b {
|
||||
return a
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user