Files
workspace/code/app/app-go/internal/ws/handler.go
T
2026-06-25 21:54:06 +08:00

377 lines
8.8 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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
}