This commit is contained in:
oneao committed 2026-06-25 16:47:54 +08:00
1 parent 513a991b57
commit 8f8e8d03c8
22 files changed
+2212 -213

No files matched your search

+372
View File
@@ -0,0 +1,372 @@
package ws
import (
"allapp-go/internal/httpx"
"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
dbClient := db.New()
rows, err := dbClient.LoadData(c.Context(), "b_family_member",
"user_id = ? AND family_id = ?", "", nil,
[]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.Background()
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_file", "id", fileInserts); err != nil {
log.Printf("[chat] insert b_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.Background()
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.Background()
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
}
+132
View File
@@ -0,0 +1,132 @@
package ws
import (
"encoding/json"
"sync"
)
// OnlineUser 在线成员信息
type OnlineUser struct {
UserID int64 `json:"user_id"`
Nickname string `json:"nickname"`
Avatar string `json:"avatar"`
}
// Room 按 family_id 分组的聊天室
type Room struct {
clients map[int64]*Client // user_id → Client
mu sync.RWMutex
}
// Hub 管理所有 Room
type Hub struct {
rooms map[int64]*Room // family_id → Room
mu sync.RWMutex
}
var defaultHub = &Hub{
rooms: make(map[int64]*Room),
}
func GetHub() *Hub {
return defaultHub
}
// getOrCreateRoom 获取或创建 Room
func (h *Hub) getOrCreateRoom(familyID int64) *Room {
h.mu.Lock()
defer h.mu.Unlock()
room, ok := h.rooms[familyID]
if !ok {
room = &Room{
clients: make(map[int64]*Client),
}
h.rooms[familyID] = room
}
return room
}
// Join 用户加入房间,广播 online 事件
func (h *Hub) Join(familyID int64, client *Client) {
room := h.getOrCreateRoom(familyID)
room.mu.Lock()
room.clients[client.userID] = client
onlineUsers := h.collectOnlineUsers(room)
room.mu.Unlock()
h.broadcastOnlineUsers(familyID, onlineUsers)
}
// Leave 用户离开房间,广播 offline 事件
func (h *Hub) Leave(familyID int64, userID int64) {
room := h.getOrCreateRoom(familyID)
room.mu.Lock()
delete(room.clients, userID)
onlineUsers := h.collectOnlineUsers(room)
room.mu.Unlock()
// 广播 offline 给房间内其他人
msg, _ := json.Marshal(map[string]any{
"type": "offline",
"user_id": userID,
})
h.broadcastRaw(familyID, msg)
// 广播更新后的在线列表
h.broadcastOnlineUsers(familyID, onlineUsers)
}
// Broadcast 向房间内所有连接广播消息
func (h *Hub) Broadcast(familyID int64, msg any) {
data, err := json.Marshal(msg)
if err != nil {
return
}
h.broadcastRaw(familyID, data)
}
// broadcastRaw 向房间内所有连接广播原始字节
func (h *Hub) broadcastRaw(familyID int64, data []byte) {
h.mu.RLock()
room, ok := h.rooms[familyID]
h.mu.RUnlock()
if !ok {
return
}
room.mu.RLock()
defer room.mu.RUnlock()
for _, client := range room.clients {
select {
case client.send <- data:
default:
// 发送通道满则跳过
}
}
}
// collectOnlineUsers 收集在线用户列表(调用方需持有 room.mu 锁)
func (h *Hub) collectOnlineUsers(room *Room) []OnlineUser {
users := make([]OnlineUser, 0, len(room.clients))
for _, client := range room.clients {
users = append(users, OnlineUser{
UserID: client.userID,
Nickname: client.nickname,
Avatar: client.avatar,
})
}
return users
}
// broadcastOnlineUsers 广播在线用户列表
func (h *Hub) broadcastOnlineUsers(familyID int64, users []OnlineUser) {
msg, _ := json.Marshal(map[string]any{
"type": "online",
"users": users,
})
h.broadcastRaw(familyID, msg)
}