377 lines
8.8 KiB
Go
377 lines
8.8 KiB
Go
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
|
||
// 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.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
|
||
}
|