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
}