package ws import ( "allapp-go/internal/httpx" "allapp-go/internal/middleware" "allapp-go/pkg/db" "allapp-go/pkg/jwtx" "allapp-go/pkg/uniqueid" "context" "encoding/json" "log" "strconv" "strings" "time" "github.com/gofiber/contrib/v3/websocket" "github.com/gofiber/fiber/v3" ) // Client 单个 WebSocket 连接 type Client struct { conn *websocket.Conn hub *Hub userID int64 familyID int64 nickname string avatar string role int64 send chan []byte } // clientSendMsg 客户端发送的原始消息 type clientSendMsg struct { Type string `json:"type"` // send | recall | delete Content string `json:"content"` MsgType int `json:"msg_type"` // 0=text 1=image 2=video 3=file 4=forward SourceTable string `json:"source_table"` SourceID int64 `json:"source_id"` MessageID int64 `json:"message_id"` Attachments []attachmentInput `json:"attachments"` } type attachmentInput struct { FileKey string `json:"file_key"` FileName string `json:"file_name"` FileType int `json:"file_type"` MimeType string `json:"mime_type"` FileSize int64 `json:"file_size"` Width int `json:"width"` Height int `json:"height"` Duration int `json:"duration"` } // AuthWS WebSocket 鉴权中间件 func AuthWS() fiber.Handler { return func(c fiber.Ctx) error { if !websocket.IsWebSocketUpgrade(c) { return fiber.ErrUpgradeRequired } token := strings.TrimSpace(c.Query("token")) familyIDStr := c.Query("family_id") if token == "" || familyIDStr == "" { return httpx.Unauthorized(c, "缺少 token 或 family_id") } claims, ok := jwtx.VerifyToken(c.Context(), token) if !ok { return httpx.Unauthorized(c, "登录已过期") } userID := claims.Data.UserID familyID, err := strconv.ParseInt(familyIDStr, 10, 64) if err != nil { return httpx.Unauthorized(c, "family_id 格式错误") } // 验证用户属于该家庭,同时获取 nickname / avatar / role // JOIN b_user 确保能拿到 avatar(v_family_member 视图可能不含此字段) dbClient := db.New() rows, err := dbClient.LoadDataBySQL(c.Context(), `SELECT m.role, m.nickname, u.avatar FROM b_family_member m JOIN b_user u ON m.user_id = u.id WHERE m.user_id = ? AND m.family_id = ?`, []any{userID, familyID}) if err != nil || len(rows) == 0 { return httpx.Unauthorized(c, "不属于该家庭") } member := rows[0] // 存入 Locals,WebSocket handler 中通过 conn.Locals 读取 c.Locals("user_id", userID) c.Locals("family_id", familyID) c.Locals("nickname", toString(member["nickname"])) c.Locals("avatar", toString(member["avatar"])) c.Locals("role", toInt64(member["role"])) return c.Next() } } // HandleWebSocket WebSocket 连接处理 func HandleWebSocket(conn *websocket.Conn) { userID := conn.Locals("user_id").(int64) familyID := conn.Locals("family_id").(int64) nickname := conn.Locals("nickname").(string) avatar := conn.Locals("avatar").(string) role := conn.Locals("role").(int64) client := &Client{ conn: conn, hub: GetHub(), userID: userID, familyID: familyID, nickname: nickname, avatar: avatar, role: role, send: make(chan []byte, 64), } client.hub.Join(familyID, client) defer client.hub.Leave(familyID, userID) go client.writePump() client.readPump() } const ( // 心跳间隔 pingPeriod = 10 * time.Second // 读超时(超时未收到任何消息则断开) readDeadline = 30 * time.Second ) // writePump 从 send 通道写入 WebSocket,同时定期发 ping func (c *Client) writePump() { ticker := time.NewTicker(pingPeriod) defer func() { ticker.Stop() c.conn.Close() }() for { select { case msg, ok := <-c.send: if !ok { return } if err := c.conn.WriteMessage(websocket.TextMessage, msg); err != nil { return } case <-ticker.C: if err := c.conn.WriteMessage(websocket.PingMessage, nil); err != nil { return } } } } // readPump 读取 WebSocket 消息并分发 func (c *Client) readPump() { defer close(c.send) // 设置 pong 处理器:收到 pong 时重置读超时 c.conn.SetPongHandler(func(string) error { c.conn.SetReadDeadline(time.Now().Add(readDeadline)) return nil }) c.conn.SetReadDeadline(time.Now().Add(readDeadline)) for { _, raw, err := c.conn.ReadMessage() if err != nil { break } var msg clientSendMsg if err := json.Unmarshal(raw, &msg); err != nil { continue } switch msg.Type { case "send": c.handleSend(msg) case "recall": c.handleRecall(msg) case "delete": c.handleDelete(msg) } } } // handleSend 处理发送消息 func (c *Client) handleSend(msg clientSendMsg) { ctx := context.WithValue(context.Background(), middleware.CtxUserIDKey, c.userID) dbClient := db.New() msgID := uniqueid.NextId() now := time.Now() // 转发:查源表生成 source_data var sourceData string if msg.MsgType == 4 && msg.SourceTable != "" && msg.SourceID > 0 { rows, err := dbClient.LoadData(ctx, msg.SourceTable, "id = ?", "", nil, []any{msg.SourceID}) if err == nil && len(rows) > 0 { data, _ := json.Marshal(rows[0]) sourceData = string(data) } } // 插入 b_chat chatData := map[string]any{ "id": msgID, "family_id": c.familyID, "user_id": c.userID, "content": msg.Content, "type": msg.MsgType, "source_table": msg.SourceTable, "source_id": msg.SourceID, "source_data": sourceData, } if err := dbClient.Insert(ctx, "b_chat", "id", chatData); err != nil { log.Printf("[chat] insert b_chat error: %v", err) return } // 附件 → 批量插入 b_file if len(msg.Attachments) > 0 { fileInserts := make([]map[string]any, 0, len(msg.Attachments)) for i, att := range msg.Attachments { fileInserts = append(fileInserts, map[string]any{ "id": uniqueid.NextId(), "target_id": msgID, "file_key": att.FileKey, "file_name": att.FileName, "file_type": att.FileType, "mime_type": att.MimeType, "file_size": att.FileSize, "width": att.Width, "height": att.Height, "duration": att.Duration, "sort_order": i, }) } if err := dbClient.BatchInsert(ctx, "b_chat_file", "id", fileInserts); err != nil { log.Printf("[chat] insert b_chat_file error: %v", err) } } // 组装附件输出 attachmentsOut := make([]map[string]any, 0) for _, att := range msg.Attachments { attachmentsOut = append(attachmentsOut, map[string]any{ "file_key": att.FileKey, "file_name": att.FileName, "file_type": att.FileType, "mime_type": att.MimeType, "file_size": att.FileSize, "width": att.Width, "height": att.Height, "duration": att.Duration, }) } if attachmentsOut == nil { attachmentsOut = []map[string]any{} } // 广播 c.hub.Broadcast(c.familyID, map[string]any{ "type": "message", "data": map[string]any{ "id": msgID, "family_id": c.familyID, "user_id": c.userID, "nickname": c.nickname, "avatar": c.avatar, "content": msg.Content, "msg_type": msg.MsgType, "source_table": msg.SourceTable, "source_id": msg.SourceID, "source_data": sourceData, "attachments": attachmentsOut, "recalled": 0, "deleted": 0, "create_time": now.Format("2006-01-02 15:04:05"), }, }) } // handleRecall 处理撤回(发送者 2 分钟内) func (c *Client) handleRecall(msg clientSendMsg) { ctx := context.WithValue(context.Background(), middleware.CtxUserIDKey, c.userID) dbClient := db.New() rows, err := dbClient.LoadData(ctx, "b_chat", "id = ?", "", nil, []any{msg.MessageID}) if err != nil || len(rows) == 0 { return } chatMsg := rows[0] // 校验:只能撤回自己的消息 msgUserID, _ := chatMsg["user_id"].(int64) if msgUserID != c.userID { return } // 校验:2 分钟内 createTime, _ := chatMsg["create_time"].(time.Time) if time.Since(createTime) > 2*time.Minute { return } // 更新 if err := dbClient.Update(ctx, "b_chat", "id", map[string]any{ "id": msg.MessageID, "recalled": 1, }); err != nil { log.Printf("[chat] recall error: %v", err) return } c.hub.Broadcast(c.familyID, map[string]any{ "type": "message_recalled", "message_id": msg.MessageID, }) } // handleDelete 处理删除(仅管理员 role=0) func (c *Client) handleDelete(msg clientSendMsg) { if c.role != 0 { return } ctx := context.WithValue(context.Background(), middleware.CtxUserIDKey, c.userID) dbClient := db.New() if err := dbClient.Update(ctx, "b_chat", "id", map[string]any{ "id": msg.MessageID, "deleted": 1, }); err != nil { log.Printf("[chat] delete error: %v", err) return } c.hub.Broadcast(c.familyID, map[string]any{ "type": "message_deleted", "message_id": msg.MessageID, }) } func toString(v any) string { if s, ok := v.(string); ok { return s } return "" } func toInt64(v any) int64 { switch val := v.(type) { case int64: return val case int: return int64(val) case int32: return int64(val) case float64: return int64(val) } return 0 }