u
This commit is contained in:
1 parent
513a991b57
commit
8f8e8d03c8
22 files changed
+2212
-213
No files matched your search
@@ -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
|
||||
}
|
||||
Reference in new issue
Block a user