gengx
This commit is contained in:
@@ -0,0 +1,411 @@
|
||||
package app
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"sort"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"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"`
|
||||
CreatedAt time.Time `json:"createdAt"`
|
||||
}
|
||||
|
||||
type wsClient struct {
|
||||
userID int64
|
||||
conn *websocket.Conn
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
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) {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
if h.clients[client.userID] == nil {
|
||||
h.clients[client.userID] = make(map[*wsClient]struct{})
|
||||
}
|
||||
h.clients[client.userID][client] = struct{}{}
|
||||
}
|
||||
|
||||
func (h *Hub) remove(client *wsClient) {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
delete(h.clients[client.userID], client)
|
||||
if len(h.clients[client.userID]) == 0 {
|
||||
delete(h.clients, client.userID)
|
||||
}
|
||||
}
|
||||
|
||||
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.mu.Lock()
|
||||
_ = client.conn.WriteJSON(payload)
|
||||
client.mu.Unlock()
|
||||
}
|
||||
}
|
||||
|
||||
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) 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 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)
|
||||
result, err := tx.ExecContext(r.Context(), `INSERT INTO im_conversations(conversation_type)VALUES(1)`)
|
||||
if err != nil {
|
||||
_ = tx.Rollback()
|
||||
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)
|
||||
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),'')
|
||||
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)
|
||||
if err != nil {
|
||||
fail(w, 500, 50001, err.Error())
|
||||
return
|
||||
}
|
||||
defer rows.Close()
|
||||
items := []map[string]any{}
|
||||
for rows.Next() {
|
||||
var id, lastSeq, readSeq, 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)
|
||||
preview := "开始聊天吧"
|
||||
var content map[string]any
|
||||
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}})
|
||||
}
|
||||
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
|
||||
}
|
||||
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)
|
||||
if err != nil {
|
||||
fail(w, 500, 50001, "查询消息失败")
|
||||
return
|
||||
}
|
||||
defer rows.Close()
|
||||
items := []messageView{}
|
||||
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 content any
|
||||
_ = 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 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)
|
||||
}
|
||||
reply(w, map[string]any{"items": items, "hasMore": 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
|
||||
}
|
||||
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
|
||||
if errors.As(err, &limitErr) {
|
||||
fail(w, http.StatusTooManyRequests, 30005, limitErr.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) 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("消息发送过于频繁,请稍后再试")
|
||||
}
|
||||
if a.isSanctionActive(r.Context(), senderID, "MUTE") {
|
||||
return item, nil, fmt.Errorf("账号处于禁言期,暂时无法发送消息")
|
||||
}
|
||||
tx, err := a.db.BeginTx(r.Context(), 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 {
|
||||
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 {
|
||||
return item, nil, fmt.Errorf("不是会话成员")
|
||||
}
|
||||
if err = a.reserveDailyActiveChat(r.Context(), tx, conversationID, senderID); err != nil {
|
||||
return item, nil, err
|
||||
}
|
||||
body, err := json.Marshal(content)
|
||||
if err != nil {
|
||||
return item, nil, fmt.Errorf("消息内容错误")
|
||||
}
|
||||
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)
|
||||
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)
|
||||
if existingErr == nil {
|
||||
_ = tx.Rollback()
|
||||
return a.loadMessage(r, 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)
|
||||
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 != nil {
|
||||
return item, nil, err
|
||||
}
|
||||
members := []int64{}
|
||||
for memberRows.Next() {
|
||||
var userID int64
|
||||
_ = memberRows.Scan(&userID)
|
||||
members = append(members, userID)
|
||||
}
|
||||
_ = memberRows.Close()
|
||||
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.Commit(); err != nil {
|
||||
return item, nil, err
|
||||
}
|
||||
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 {
|
||||
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)
|
||||
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, _ := pathID(r)
|
||||
var req struct {
|
||||
Pinned *bool `json:"pinned"`
|
||||
Muted *bool `json:"muted"`
|
||||
ReadSeq int64 `json:"readSeq"`
|
||||
}
|
||||
if decode(r, &req) != nil {
|
||||
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)
|
||||
}
|
||||
_, 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)
|
||||
if err != nil {
|
||||
fail(w, 500, 50001, "保存失败")
|
||||
return
|
||||
}
|
||||
reply(w, map[string]bool{"success": true})
|
||||
}
|
||||
|
||||
func (a *App) websocket(w http.ResponseWriter, r *http.Request) {
|
||||
if !a.originAllowed(r.Header.Get("Origin")) {
|
||||
fail(w, http.StatusForbidden, 10006, "WebSocket 来源不允许")
|
||||
return
|
||||
}
|
||||
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"))
|
||||
}}
|
||||
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()}})
|
||||
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":
|
||||
_ = conn.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()})
|
||||
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})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func max64(a, b int64) int64 {
|
||||
if a > b {
|
||||
return a
|
||||
}
|
||||
return b
|
||||
}
|
||||
Reference in New Issue
Block a user