diff --git a/.idea/go.imports.xml b/.idea/go.imports.xml
new file mode 100644
index 0000000..d7202f0
--- /dev/null
+++ b/.idea/go.imports.xml
@@ -0,0 +1,11 @@
+
+
+
+
+
+
\ No newline at end of file
diff --git a/cmd/server/main.go b/cmd/server/main.go
index b3d48a0..3e6770a 100644
--- a/cmd/server/main.go
+++ b/cmd/server/main.go
@@ -20,9 +20,14 @@
package main
import (
+ "context"
"fmt"
"log"
"net/http"
+ "os"
+ "os/signal"
+ "strings"
+ "syscall"
"time"
"xk-websocket-v2/internal/api"
@@ -43,6 +48,15 @@ import (
"gorm.io/gorm"
)
+// WebSocket 读端防护参数
+// wsMaxMessageBytes: 单帧最大字节数,超过直接断开,防止超大帧耗尽内存
+// wsReadWait: 读超时时间,须大于 writePump 的 Ping 周期(50s),
+// 收到 Pong 或业务消息时刷新;超时未收到任何数据则判定连接假死并回收
+const (
+ wsMaxMessageBytes = 512 * 1024
+ wsReadWait = 90 * time.Second
+)
+
/**
* initConfig
*
@@ -54,16 +68,56 @@ import (
* 3. 添加配置文件搜索路径(configs目录和当前目录)
* 4. 读取配置文件,如果失败则终止程序
*/
+// parseCSV 解析逗号分隔的字符串为去空白的切片(用于读取来源白名单等配置)
+func parseCSV(s string) []string {
+ if strings.TrimSpace(s) == "" {
+ return nil
+ }
+ parts := strings.Split(s, ",")
+ out := make([]string, 0, len(parts))
+ for _, p := range parts {
+ if v := strings.TrimSpace(p); v != "" {
+ out = append(out, v)
+ }
+ }
+ return out
+}
+
+// isOriginAllowed 校验 WebSocket 来源(Origin)是否在白名单内;白名单为空时放行(仅建议本地开发)
+func isOriginAllowed(origin string, allowed []string) bool {
+ if len(allowed) == 0 {
+ return true
+ }
+ for _, a := range allowed {
+ if a == origin {
+ return true
+ }
+ }
+ return false
+}
+
func initConfig() {
viper.SetConfigName("config")
viper.SetConfigType("yaml")
viper.AddConfigPath("configs")
viper.AddConfigPath(".")
+ // 允许用环境变量覆盖配置中的敏感项(如 JWT_SECRET、DATABASE_DSN、REDIS_PASSWORD、TURN_SHARED_SECRET)。
+ // 生产环境应通过环境变量注入密钥/密码,config.yaml 仅保留本地开发默认值。
+ viper.AutomaticEnv()
+ viper.SetEnvKeyReplacer(strings.NewReplacer(".", "_"))
if err := viper.ReadInConfig(); err != nil {
log.Fatalf("❌ 无法读取配置文件: %v", err)
}
}
+// isProdEnv 判断是否为生产环境(app.env 为 production/prod),
+// 口径与 internal/api/auth_handler.go 的同名函数一致(包不同无法复用,各自维护4行)。
+// 用途:数据库自动迁移只允许在开发环境执行,生产环境启动必须零 DDL。
+func isProdEnv() bool {
+ env := viper.GetString("app.env")
+ return env == "production" || env == "prod"
+}
+
/**
* initDB
*
@@ -73,9 +127,8 @@ func initConfig() {
* 1. 从配置文件中读取数据库连接字符串(DSN)
* 2. 使用GORM连接MySQL数据库
* 3. 如果连接失败,终止程序
- * 4. 执行自动数据库迁移,创建所有表结构
- * 5. 为所有表添加中文注释
- * 6. 返回数据库连接实例
+ * 4. 仅开发环境:执行 schema 自愈 + 自动迁移 + 表注释(生产环境启动零 DDL)
+ * 5. 返回数据库连接实例
*
* @returns *gorm.DB 数据库连接实例
*/
@@ -84,11 +137,53 @@ func initDB() *gorm.DB {
dsn := viper.GetString("database.dsn")
// 步骤2: 使用GORM连接MySQL数据库
- db, err := gorm.Open(mysql.Open(dsn), &gorm.Config{})
+ // DisableForeignKeyConstraintWhenMigrating: 禁止 AutoMigrate 自动创建外键约束。
+ // 原因:chat_conversations.target_id 既可能是用户ID也可能是群ID,历史上 GORM
+ // 自动建的外键导致群会话插入失败,代码只能靠 SET FOREIGN_KEY_CHECKS=0 绕过(已移除)。
+ // 表结构与约束统一以 nl_im_plus.sql / migrations 下的 SQL 脚本为准。
+ db, err := gorm.Open(mysql.Open(dsn), &gorm.Config{
+ DisableForeignKeyConstraintWhenMigrating: true,
+ })
if err != nil {
log.Fatalf("❌ 数据库连接失败: %v", err)
}
+ // 步骤2.5: 应用连接池配置。
+ // config.yaml 中声明了 max_idle_conns/max_open_conns,但此前从未真正设置到
+ // 底层 sql.DB 上(配置形同虚设,走的是驱动默认值),高并发下连接数不可控
+ sqlDB, err := db.DB()
+ if err != nil {
+ log.Fatalf("❌ 获取底层数据库连接失败: %v", err)
+ }
+ maxIdleConns := viper.GetInt("database.max_idle_conns")
+ if maxIdleConns <= 0 {
+ maxIdleConns = 10
+ }
+ maxOpenConns := viper.GetInt("database.max_open_conns")
+ if maxOpenConns <= 0 {
+ maxOpenConns = 100
+ }
+ sqlDB.SetMaxIdleConns(maxIdleConns)
+ sqlDB.SetMaxOpenConns(maxOpenConns)
+ // 连接最长存活1小时:避免被 MySQL wait_timeout 掐断后拿到失效连接
+ sqlDB.SetConnMaxLifetime(time.Hour)
+
+ // 步骤2.6: 环境闸门——数据库迁移(ensureSchema/AutoMigrate/表注释,均含 DDL)只在开发环境执行。
+ // 为什么:打包部署的生产环境启动绝不允许改表结构——多节点同时启动会并发迁移互相冲突、
+ // 大表 DDL 锁表可能拖垮线上服务、模型与线上库的历史差异还可能触发误改。
+ // 生产库结构变更统一由 DBA 手动执行 migrations/ 下的 SQL 脚本(01~07 均幂等可重复执行)。
+ if isProdEnv() {
+ log.Println("🔒 生产环境跳过数据库自动迁移,表结构以 migrations/ SQL 脚本为准")
+ return db
+ }
+
+ // 步骤2.7: schema 自愈(仅开发环境)——把 migrations/ 下需人工执行的结构变更在启动时幂等补齐。
+ // 为什么必须在 AutoMigrate 之前:moment_likes 的历史重复数据不先清理,
+ // AutoMigrate 建 uk_moment_user 唯一索引会直接失败导致服务启不来
+ if err := ensureSchema(db); err != nil {
+ log.Fatalf("❌ 数据库结构自愈失败: %v", err)
+ }
+
// 步骤3: 执行自动数据库迁移,创建所有表结构
err = db.AutoMigrate(
&model.ChatMessage{},
@@ -108,6 +203,10 @@ func initDB() *gorm.DB {
&model.MomentNotification{},
&model.MessageReadReceipt{},
&model.UserSetting{},
+ // AI 机器人两张表(等价 migrations/07):此前遗漏,漏执行脚本的开发库 /api/ai/* 会运行时报错;
+ // users.is_bot 列已随 User 模型由上方 AutoMigrate 自动补齐
+ &model.AIConfig{},
+ &model.AIBot{},
)
if err != nil {
log.Fatalf("❌ 数据库迁移失败: %v", err)
@@ -135,6 +234,10 @@ func initDB() *gorm.DB {
{"moment_notifications", "朋友圈通知表"},
{"message_read_receipts", "消息已读回执表"},
{"user_settings", "用户个性化设置表"},
+ {"room_members", "房间成员表,存储房间与用户的关联关系"},
+ {"message_deletions", "消息删除记录表,记录哪个用户删除了哪条消息(仅影响本人可见性)"},
+ {"ai_configs", "AI提供商配置表(全局一条,仅管理员维护)"},
+ {"ai_bots", "AI机器人定义表"},
}
for _, tc := range tableComments {
@@ -148,6 +251,109 @@ func initDB() *gorm.DB {
return db
}
+/**
+ * ensureSchema
+ *
+ * 功能:启动时的数据库结构自愈——把 migrations/ 下需人工执行的变更以幂等 SQL 补齐。
+ * 注意:仅开发环境执行(initDB 里被 isProdEnv 闸门拦截),生产环境启动零 DDL,
+ * 线上结构变更必须由 DBA 手动执行 migrations/ 脚本。
+ *
+ * 为什么需要:项目约定表结构以 SQL 脚本为准(新表/新列不走 GORM 模型迁移),
+ * 但脚本依赖人工执行,漏执行的开发库会出现"代码引用了不存在的表/列"的运行时错误,
+ * 甚至 AutoMigrate 因存量脏数据建唯一索引失败而直接启动崩溃。
+ * 这里在启动路径上自动补齐,等价脚本(03/04/05)保留供 DBA 手动可控升级,
+ * 已手动执行过的环境下各步骤经 information_schema / IF NOT EXISTS 判断后跳过。
+ *
+ * 自愈内容:
+ * 1. room_members 表(等价 migrations/02+03):建表含 nickname 列;老表缺列则补列
+ * 2. message_deletions 表(等价 migrations/04):删除仅对我生效的记录表
+ * 3. moment_likes 唯一索引(等价 migrations/05):先清理重复点赞、修正 like_count,
+ * 再建 uk_moment_user——否则 AutoMigrate 在有重复数据的库上建索引必失败
+ */
+func ensureSchema(db *gorm.DB) error {
+ // ---- 1. room_members:建表(含 nickname)或给老表补 nickname 列 ----
+ // 不显式指定 COLLATE:沿用库默认排序规则,与 AutoMigrate 建的其他表保持一致,
+ // 避免 JOIN users 等表时出现 "Illegal mix of collations"
+ if err := db.Exec("CREATE TABLE IF NOT EXISTS `room_members` (" +
+ "`room_id` varchar(100) NOT NULL COMMENT '房间ID'," +
+ "`user_id` varchar(100) NOT NULL COMMENT '用户ID'," +
+ "`role` tinyint(1) NOT NULL DEFAULT 0 COMMENT '成员角色(0=成员,1=管理员,2=群主)'," +
+ "`joined_at` datetime(3) NULL DEFAULT NULL COMMENT '加入时间'," +
+ "`muted_until` datetime NULL DEFAULT NULL COMMENT '禁言到期时间'," +
+ "`nickname` varchar(100) NOT NULL DEFAULT '' COMMENT '群名片(群内显示昵称)'," +
+ "PRIMARY KEY (`room_id`, `user_id`)," +
+ "INDEX `idx_room_members_user_id`(`user_id`)" +
+ ") ENGINE = InnoDB CHARACTER SET = utf8mb4 COMMENT = '房间成员表,存储房间与用户的关联关系'").Error; err != nil {
+ return fmt.Errorf("创建 room_members 表失败: %w", err)
+ }
+ // 表已存在的老库:检查 nickname 列,缺则补(MySQL 无 ADD COLUMN IF NOT EXISTS)
+ var nickCount int64
+ if err := db.Raw("SELECT COUNT(*) FROM information_schema.COLUMNS " +
+ "WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = 'room_members' AND COLUMN_NAME = 'nickname'").
+ Scan(&nickCount).Error; err != nil {
+ return fmt.Errorf("检查 room_members.nickname 列失败: %w", err)
+ }
+ if nickCount == 0 {
+ if err := db.Exec("ALTER TABLE `room_members` ADD COLUMN `nickname` varchar(100) NOT NULL DEFAULT '' " +
+ "COMMENT '群名片(群内显示昵称)' AFTER `muted_until`").Error; err != nil {
+ return fmt.Errorf("补充 room_members.nickname 列失败: %w", err)
+ }
+ log.Println("🔧 [Schema] room_members 已补充 nickname 列")
+ }
+
+ // ---- 2. message_deletions:消息删除记录表(删除仅对我生效) ----
+ if err := db.Exec("CREATE TABLE IF NOT EXISTS `message_deletions` (" +
+ "`id` bigint UNSIGNED NOT NULL AUTO_INCREMENT COMMENT '自增主键'," +
+ "`message_id` bigint UNSIGNED NOT NULL COMMENT '被删除的消息ID'," +
+ "`user_id` varchar(100) NOT NULL COMMENT '执行删除的用户ID(删除仅对该用户生效)'," +
+ "`created_at` datetime(3) NULL DEFAULT NULL COMMENT '删除时间'," +
+ "PRIMARY KEY (`id`)," +
+ // 唯一索引防止同一用户对同一消息重复插入删除记录
+ "UNIQUE INDEX `uk_message_user`(`message_id`, `user_id`)," +
+ "INDEX `idx_message_deletions_user_id`(`user_id`)" +
+ ") ENGINE = InnoDB CHARACTER SET = utf8mb4 COMMENT = '消息删除记录表,记录哪个用户删除了哪条消息(仅影响本人可见性)'").Error; err != nil {
+ return fmt.Errorf("创建 message_deletions 表失败: %w", err)
+ }
+
+ // ---- 3. moment_likes:清理重复点赞后补 uk_moment_user 唯一索引 ----
+ // 仅处理"表已存在且索引缺失"的存量库;全新环境表还不存在,
+ // 跳过后由 AutoMigrate 直接建表+索引(空表无重复数据,安全)
+ var likeTableCount int64
+ if err := db.Raw("SELECT COUNT(*) FROM information_schema.TABLES " +
+ "WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = 'moment_likes'").
+ Scan(&likeTableCount).Error; err != nil {
+ return fmt.Errorf("检查 moment_likes 表失败: %w", err)
+ }
+ if likeTableCount > 0 {
+ var idxCount int64
+ if err := db.Raw("SELECT COUNT(*) FROM information_schema.STATISTICS " +
+ "WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = 'moment_likes' AND INDEX_NAME = 'uk_moment_user'").
+ Scan(&idxCount).Error; err != nil {
+ return fmt.Errorf("检查 moment_likes 唯一索引失败: %w", err)
+ }
+ if idxCount == 0 {
+ // 3.1 清理历史重复数据(保留每组最早一条),否则建唯一索引会失败
+ if err := db.Exec("DELETE ml FROM `moment_likes` ml " +
+ "INNER JOIN `moment_likes` ml2 ON ml.moment_id = ml2.moment_id " +
+ "AND ml.user_id = ml2.user_id AND ml.id > ml2.id").Error; err != nil {
+ return fmt.Errorf("清理 moment_likes 重复数据失败: %w", err)
+ }
+ // 3.2 修正因重复点赞虚高的 like_count(以实际点赞记录数为准)
+ if err := db.Exec("UPDATE `moments` m SET m.like_count = " +
+ "(SELECT COUNT(*) FROM `moment_likes` ml WHERE ml.moment_id = m.id)").Error; err != nil {
+ return fmt.Errorf("修正 moments.like_count 失败: %w", err)
+ }
+ // 3.3 建唯一索引:并发点赞由数据库兜底,代码层把唯一冲突当"已点赞"幂等处理
+ if err := db.Exec("ALTER TABLE `moment_likes` ADD UNIQUE INDEX `uk_moment_user`(`moment_id`, `user_id`)").Error; err != nil {
+ return fmt.Errorf("创建 moment_likes 唯一索引失败: %w", err)
+ }
+ log.Println("🔧 [Schema] moment_likes 已去重并补充 uk_moment_user 唯一索引")
+ }
+ }
+
+ return nil
+}
+
/**
* initRedis
*
@@ -197,7 +403,10 @@ func initRedis() *redis.Client {
*/
func main() {
// 步骤1: 初始化基础服务
- initConfig() // 读取配置文件
+ initConfig() // 读取配置文件
+ // 配置加载完成后再初始化 JWT 密钥,确保 config/env 中的 jwt.secret 真正生效
+ // (历史问题:utils 包级 init 早于配置读取,导致永远使用默认密钥)
+ utils.InitJWT(viper.GetString("jwt.secret"))
db := initDB() // 连接MySQL数据库
rdb := initRedis() // 连接Redis
@@ -239,6 +448,8 @@ func main() {
service.InitLoginLogService(db) // 登录日志服务(记录登录历史)
service.InitMomentService(db) // 朋友圈服务(动态、点赞、评论)
service.InitSearchService(db) // 搜索服务(聚合搜索)
+ service.InitQRCodeLoginService(rdb) // 扫码登录服务(App 扫码登录 PC,Redis 状态机)
+ service.InitAIBotService(db) // AI机器人服务(工厂模式接入大模型,机器人聊天应答)
// 步骤5: 启动TURN服务器(用于WebRTC音视频通话)
go turnserver.Start()
@@ -261,10 +472,11 @@ func main() {
r.Use(func(c *gin.Context) {
// 设置允许的源(*表示允许所有源)
c.Writer.Header().Set("Access-Control-Allow-Origin", "*")
- // 设置允许的HTTP方法
- c.Writer.Header().Set("Access-Control-Allow-Methods", "POST, GET, OPTIONS")
- // 设置允许的请求头
- c.Writer.Header().Set("Access-Control-Allow-Headers", "Content-Type, Content-Length, Accept-Encoding, X-CSRF-Token, Authorization, X-User-ID")
+ // 设置允许的HTTP方法(路由中大量使用 PUT/DELETE/PATCH,如删好友、改群设置,
+ // 缺失会导致浏览器预检失败、相关功能在跨域场景下全部不可用)
+ c.Writer.Header().Set("Access-Control-Allow-Methods", "POST, GET, PUT, DELETE, PATCH, OPTIONS")
+ // 设置允许的请求头(X-User-ID 已废弃:身份一律取自 JWT,不再信任该头)
+ c.Writer.Header().Set("Access-Control-Allow-Headers", "Content-Type, Content-Length, Accept-Encoding, X-CSRF-Token, Authorization")
// 处理OPTIONS预检请求
if c.Request.Method == "OPTIONS" {
c.AbortWithStatus(204) // 返回204 No Content
@@ -277,20 +489,55 @@ func main() {
// 访问路径:/uploads/xxx -> ./uploads/xxx
r.Static("/uploads", "./uploads")
- // 步骤8: 注册WebSocket路由
- // 路径:GET /ws?user_id=xxx
- r.GET("/ws", func(c *gin.Context) {
- // 步骤1: 创建WebSocket升级器,允许所有来源连接
- upgrader := websocket.Upgrader{CheckOrigin: func(r *http.Request) bool { return true }}
+ // PC 桌面端(Electron)更新包托管目录
+ // electron-updater 会到 /updates/latest.yml 检查版本,并下载同目录下的安装包与 .blockmap(增量更新)
+ // 发版流程:electron-builder 打包后,把 release/ 下的 latest.yml、NL-IM-Setup-x.y.z.exe、.blockmap 上传到 ./updates
+ r.Static("/updates", "./updates")
- // 步骤2: 将HTTP连接升级为WebSocket连接
+ // 步骤8: 注册WebSocket路由
+ // 路径:GET /ws?token=xxx(用户身份以 JWT 为准)
+ // WebSocket 允许来源白名单(为空表示不限制,仅建议本地开发使用)
+ wsAllowedOrigins := parseCSV(viper.GetString("app.ws_allowed_origins"))
+ r.GET("/ws", func(c *gin.Context) {
+ // 步骤1: 握手鉴权——必须携带有效 JWT(?token= 或 Authorization),
+ // 用户身份以 Token 为准,禁止用 user_id 查询参数伪造身份(原实现无鉴权可冒充任意人)
+ token := c.Query("token")
+ if token == "" {
+ token = c.GetHeader("Authorization")
+ }
+ if len(token) > 7 && token[:7] == "Bearer " {
+ token = token[7:]
+ }
+ userID, err := utils.ValidateToken(token)
+ if err != nil || userID == "" {
+ // 统一响应结构:保持 HTTP 401(浏览器 WebSocket 握手失败只看状态码),
+ // 响应体复用与业务接口一致的 ApiResponse 格式,不再返回裸 gin.H
+ middleware.AbortUnauthorized(c, "WebSocket 认证失败")
+ return
+ }
+
+ // 步骤2: 创建WebSocket升级器,按白名单校验来源(Origin)
+ upgrader := websocket.Upgrader{CheckOrigin: func(r *http.Request) bool {
+ return isOriginAllowed(r.Header.Get("Origin"), wsAllowedOrigins)
+ }}
+
+ // 步骤3: 将HTTP连接升级为WebSocket连接
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
if err != nil {
return // 升级失败,直接返回
}
- // 步骤3: 获取用户ID(从查询参数中获取)
- userID := c.Query("user_id")
+ // 步骤3.1: 读端加固
+ // 1) SetReadLimit 限制单帧大小,防止恶意客户端发送超大帧一次性撑爆内存;
+ // 2) SetReadDeadline + SetPongHandler 组成读超时机制:writePump 每 50s 发一次 Ping,
+ // 正常客户端自动回 Pong 会刷新读截止时间;而"半开"连接(TCP 未断、应用层无响应)
+ // 收不到任何数据/Pong,90s 后 ReadMessage 会超时返回,连接及其 goroutine/缓冲得以及时回收,
+ // 避免假死连接长期堆积。
+ conn.SetReadLimit(wsMaxMessageBytes)
+ _ = conn.SetReadDeadline(time.Now().Add(wsReadWait))
+ conn.SetPongHandler(func(string) error {
+ return conn.SetReadDeadline(time.Now().Add(wsReadWait))
+ })
// 步骤4: 生成唯一的客户端ID(节点ID + 时间戳)
clientID := fmt.Sprintf("%s-%d", viper.GetString("app.node_id"), time.Now().UnixNano())
@@ -301,10 +548,8 @@ func main() {
// 步骤6: 注册客户端到管理器
manager.Manager.Register(client)
- // 步骤7: 如果提供了用户ID,绑定用户到客户端
- if userID != "" {
- service.ChatSvc.BindUser(client, userID)
- }
+ // 步骤7: 用 Token 中解析出的用户ID 绑定客户端(忽略查询参数中的 user_id)
+ service.ChatSvc.BindUser(client, userID)
// 步骤8: 发送客户端ID给前端
client.SendQueue <- []byte(fmt.Sprintf(`{"clientId": "%s"}`, clientID))
@@ -313,13 +558,18 @@ func main() {
for {
_, message, err := conn.ReadMessage()
if err != nil {
- // 连接断开,注销客户端
- manager.Manager.Unregister(client)
+ // 连接断开:注销客户端并清理 Redis 路由(该用户在本节点最后一个连接下线时)
+ service.ChatSvc.HandleClientOffline(client)
break
}
+ // 收到业务消息同样视为连接存活,刷新读截止时间
+ _ = conn.SetReadDeadline(time.Now().Add(wsReadWait))
// WebSocket 仅处理信令和心跳,不再处理 send_message (改走API)
// 但为了兼容,仍保留 PushTask
- ws.PushTask(client, message)
+ // 协程池满时 Invoke 会返回错误,记录日志避免消息静默丢失(无观测)
+ if err := ws.PushTask(client, message); err != nil {
+ log.Printf("⚠️ [WS] 任务入池失败(协程池可能已满): client=%s err=%v", client.ID, err)
+ }
}
})
@@ -337,15 +587,10 @@ func main() {
apiGroup.POST("/send-sms-code", api.SendSmsCodeHandler)
apiGroup.GET("/check-token", api.CheckTokenHandler)
apiGroup.GET("/health", api.HealthHandler)
- apiGroup.GET("/ice-servers", api.ICEHandler)
- // 消息相关(可选认证,兼容旧代码)
- apiGroup.POST("/send", api.SendHandler)
- apiGroup.POST("/send-to-user", api.SendToUserHandler)
- apiGroup.POST("/bind", api.BindHandler)
- apiGroup.GET("/check-user-online", api.CheckUserOnlineHandler)
- apiGroup.GET("/messages", api.HistoryHandler)
- apiGroup.GET("/messages/sync", api.SyncMessagesHandler)
+ // 扫码登录(PC 侧,未登录即可访问:生成二维码 + 轮询状态)
+ apiGroup.POST("/qrcode/generate", api.GenerateQRCodeHandler)
+ apiGroup.GET("/qrcode/status", api.QRCodeStatusHandler)
// 需要认证的接口组(必须携带有效的JWT Token)
authGroup := apiGroup.Group("")
@@ -354,6 +599,17 @@ func main() {
// 聚合搜索
authGroup.GET("/search", api.GlobalSearchHandler)
+ // 消息发送 / 绑定 / 历史(强制 JWT:发送者与绑定用户以 Token 为准,禁止伪造身份;
+ // 历史消息仅房间成员可读)。原先放在公开组且靠 X-User-ID 兜底,存在严重越权风险,已收敛到此。
+ authGroup.POST("/send", api.SendHandler)
+ authGroup.POST("/send-to-user", api.SendToUserHandler)
+ authGroup.POST("/bind", api.BindHandler)
+ authGroup.GET("/check-user-online", api.CheckUserOnlineHandler)
+ authGroup.GET("/messages", api.HistoryHandler)
+ authGroup.GET("/messages/sync", api.SyncMessagesHandler)
+ // TURN/STUN 临时凭证:需登录,凭证绑定当前用户
+ authGroup.GET("/ice-servers", api.ICEHandler)
+
// 用户管理
authGroup.GET("/user/my-info", api.GetMyInfoHandler)
authGroup.GET("/user/list", api.GetUserListHandler)
@@ -377,8 +633,9 @@ func main() {
authGroup.POST("/contacts/update/:id", api.UpdateContactHandler)
authGroup.POST("/contacts/delete/:id", api.DeleteContactHandler)
- // 消息撤回、登出、已读回执
+ // 消息撤回、删除、登出、已读回执
authGroup.POST("/messages/recall", api.RecallMessageHandler)
+ authGroup.POST("/messages/delete", api.DeleteMessageHandler)
authGroup.POST("/messages/read-receipts", api.MarkMessagesReadHandler)
authGroup.POST("/logout", api.LogoutHandler)
@@ -414,9 +671,24 @@ func main() {
authGroup.GET("/groups/:room_id/members/:user_id/mute", api.GetMemberMuteStatusHandler)
authGroup.GET("/groups/:room_id/settings", api.GetGroupSettingsHandler)
authGroup.POST("/groups/:room_id/settings", api.UpdateGroupSettingsHandler)
+ authGroup.POST("/groups/:room_id/nickname", api.UpdateMyNicknameHandler) // 修改我在群里的昵称(群名片)
authGroup.GET("/group-notifications", api.GetGroupNotificationsHandler)
authGroup.GET("/calls/history", api.GetCallHistoryHandler)
+ // 扫码登录(App 侧,需登录:扫码/确认/取消)
+ authGroup.POST("/qrcode/scan", api.ScanQRCodeHandler)
+ authGroup.POST("/qrcode/confirm", api.ConfirmQRCodeHandler)
+ authGroup.POST("/qrcode/cancel", api.CancelQRCodeHandler)
+
+ // AI 机器人:配置读写/测试/机器人增删改仅管理员(id=1),列表所有登录用户可见
+ authGroup.GET("/ai/config", api.GetAIConfigHandler)
+ authGroup.POST("/ai/config", api.SaveAIConfigHandler)
+ authGroup.POST("/ai/config/test", api.TestAIConfigHandler)
+ authGroup.GET("/ai/bots", api.ListAIBotsHandler)
+ authGroup.POST("/ai/bots", api.CreateAIBotHandler)
+ authGroup.POST("/ai/bots/update/:id", api.UpdateAIBotHandler)
+ authGroup.POST("/ai/bots/delete/:id", api.DeleteAIBotHandler)
+
// 附件管理
authGroup.POST("/attachments/upload", api.UploadAttachmentHandler)
authGroup.GET("/attachments", api.GetAttachmentsHandler)
@@ -434,8 +706,53 @@ func main() {
}
}
- // 步骤10: 启动HTTP服务器
+ // 步骤10: 启动HTTP服务器(支持优雅关闭)
+ // 原实现 r.Run 阻塞至进程被信号强杀,defer 的清理逻辑(StopWorkerPool 等)不会执行,
+ // MediaServer(SFU/RTMP/FFmpeg 子进程)、协程池、Redis/DB 连接都得不到有序释放,
+ // 可能残留端口占用与孤儿 FFmpeg 进程。
+ // 改为 http.Server + 信号监听:收到 SIGINT/SIGTERM 后先停止接收新请求并限时等待存量请求,
+ // 再依次关闭各组件,最后退出。
port := viper.GetString("app.port") // 从配置文件读取端口号
- log.Printf("🚀 服务启动在端口: %s", port)
- r.Run(":" + port) // 启动服务器并监听指定端口
+ srv := &http.Server{
+ Addr: ":" + port,
+ Handler: r,
+ }
+
+ // HTTP 服务放到独立协程启动,主协程留下来等退出信号
+ go func() {
+ log.Printf("🚀 服务启动在端口: %s", port)
+ if err := srv.ListenAndServe(); err != nil && err != http.ErrServerClosed {
+ log.Fatalf("❌ HTTP 服务启动失败: %v", err)
+ }
+ }()
+
+ // 阻塞等待 SIGINT(Ctrl+C) / SIGTERM(kill) 退出信号
+ quit := make(chan os.Signal, 1)
+ signal.Notify(quit, syscall.SIGINT, syscall.SIGTERM)
+ <-quit
+ log.Println("🛑 收到退出信号,开始优雅关闭...")
+
+ // 限时 10 秒等待存量 HTTP 请求处理完毕(超时则强制返回)
+ ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
+ defer cancel()
+ if err := srv.Shutdown(ctx); err != nil {
+ log.Printf("⚠️ HTTP 服务关闭异常: %v", err)
+ }
+
+ // 有序释放各组件资源:
+ // 1. 媒体服务器(关闭所有房间、SFU、RTMP 监听、FFmpeg 子进程)
+ mediaserver.GetServer().Stop()
+ // 2. WebSocket 协程池(defer 中也有一次,ants 的 Release 可重复调用,安全)
+ ws.StopWorkerPool()
+ // 3. Redis 连接
+ if err := rdb.Close(); err != nil {
+ log.Printf("⚠️ Redis 关闭异常: %v", err)
+ }
+ // 4. MySQL 连接池
+ if sqlDB, err := db.DB(); err == nil {
+ if err := sqlDB.Close(); err != nil {
+ log.Printf("⚠️ MySQL 关闭异常: %v", err)
+ }
+ }
+ log.Println("✅ 服务已优雅退出")
}
diff --git a/configs/config.yaml b/configs/config.yaml
index d6852f1..2b225fe 100644
--- a/configs/config.yaml
+++ b/configs/config.yaml
@@ -3,12 +3,16 @@
# ==========================================
app:
name: "xk-websocket-v2"
+ # 运行环境: dev / production。生产环境(production)下不会把验证码等敏感信息返回给客户端
+ env: "dev"
# 服务监听端口 (HTTP & WebSocket)
port: "12080"
# 节点唯一标识符 (集群模式下必须唯一,用于路由消息)
node_id: "node-01"
# 日志级别: debug, info, warn, error
log_level: "debug"
+ # WebSocket 允许的来源(Origin)白名单,逗号分隔;为空表示不限制(仅建议本地开发)
+ ws_allowed_origins: ""
# ==========================================
# MySQL 数据库配置
@@ -16,8 +20,8 @@ app:
database:
# 数据库连接字符串 (DSN)
# 格式: user:password@tcp(host:port)/dbname?charset=utf8mb4&parseTime=True&loc=Local
-# dsn: "root:root@tcp(127.0.0.1:3306)/nl_im?charset=utf8mb4&parseTime=True&loc=Local"
- dsn: "root:mysql_PKC65h@tcp(101.43.12.11:3306)/nl_im_plus?charset=utf8mb4&parseTime=True&loc=Local"
+ dsn: "root:root@tcp(127.0.0.1:3306)/nl_im_plus?charset=utf8mb4&parseTime=True&loc=Local"
+# dsn: "root:mysql_PKC65h@tcp(101.43.12.11:3306)/nl_im_plus?charset=utf8mb4&parseTime=True&loc=Local"
# 连接池最大空闲连接数
max_idle_conns: 10
# 连接池最大打开连接数
@@ -27,16 +31,29 @@ database:
# Redis 配置 (用于缓存和集群消息广播)
# ==========================================
redis:
- addr: "101.43.12.11:6379"
- password: "redis_xeePNa"
+# addr: "101.43.12.11:6379"
+# password: "redis_xeePNa"
+ addr: "127.0.0.1:6379"
+ password: "redis_5E4HwH"
# 数据库索引 (0-15)
db: 0
+# ==========================================
+# IP2Region 离线 IP 库
+# ==========================================
+ip2region:
+ db_path: "./ip2region.xdb"
+
# ==========================================
# JWT 配置
# ==========================================
+# 安全提示:以下为本地开发默认值。生产环境请用环境变量覆盖,切勿提交真实密钥:
+# JWT_SECRET 覆盖 jwt.secret
+# DATABASE_DSN 覆盖 database.dsn
+# REDIS_PASSWORD 覆盖 redis.password
+# TURN_SHARED_SECRET 覆盖 turn.shared_secret
jwt:
- # JWT签名密钥(生产环境应使用强随机密钥)
+ # JWT签名密钥(生产环境应使用强随机密钥,建议用 JWT_SECRET 环境变量注入)
secret: "xk-websocket-jwt-secret-key-2025-change-in-production"
# ==========================================
diff --git a/go.mod b/go.mod
index 7e9e1d5..cf3f5cd 100644
--- a/go.mod
+++ b/go.mod
@@ -8,13 +8,16 @@ require (
github.com/go-redis/redis/v8 v8.11.5
github.com/golang-jwt/jwt/v5 v5.3.0
github.com/gorilla/websocket v1.5.3
+ github.com/lionsoul2014/ip2region/binding/golang v0.0.0-20260630140118-a99989343ebd
github.com/panjf2000/ants/v2 v2.11.3
github.com/pion/interceptor v0.1.29
github.com/pion/rtp v1.8.7
github.com/pion/turn/v2 v2.1.6
github.com/pion/webrtc/v3 v3.3.6
+ github.com/skip2/go-qrcode v0.0.0-20200617195104-da1b6568686e
github.com/spf13/viper v1.21.0
golang.org/x/crypto v0.45.0
+ golang.org/x/net v0.47.0
gorm.io/driver/mysql v1.6.0
gorm.io/gorm v1.31.1
)
@@ -44,7 +47,6 @@ require (
github.com/json-iterator/go v1.1.12 // indirect
github.com/klauspost/cpuid/v2 v2.3.0 // indirect
github.com/leodido/go-urn v1.4.0 // indirect
- github.com/lionsoul2014/ip2region/binding/golang v0.0.0-20260630140118-a99989343ebd // indirect
github.com/mattn/go-isatty v0.0.20 // indirect
github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // indirect
github.com/modern-go/reflect2 v1.0.2 // indirect
@@ -77,7 +79,6 @@ require (
go.uber.org/mock v0.6.0 // indirect
go.yaml.in/yaml/v3 v3.0.4 // indirect
golang.org/x/arch v0.23.0 // indirect
- golang.org/x/net v0.47.0 // indirect
golang.org/x/sync v0.18.0 // indirect
golang.org/x/sys v0.38.0 // indirect
golang.org/x/text v0.31.0 // indirect
diff --git a/go.sum b/go.sum
index 92b1fa4..601312d 100644
--- a/go.sum
+++ b/go.sum
@@ -143,6 +143,8 @@ github.com/rogpeppe/go-internal v1.10.0 h1:TMyTOH3F/DB16zRVcYyreMH6GnZZrwQVAoYjR
github.com/rogpeppe/go-internal v1.10.0/go.mod h1:UQnix2H7Ngw/k4C5ijL5+65zddjncjaFoBhdsK/akog=
github.com/sagikazarmark/locafero v0.12.0 h1:/NQhBAkUb4+fH1jivKHWusDYFjMOOKU88eegjfxfHb4=
github.com/sagikazarmark/locafero v0.12.0/go.mod h1:sZh36u/YSZ918v0Io+U9ogLYQJ9tLLBmM4eneO6WwsI=
+github.com/skip2/go-qrcode v0.0.0-20200617195104-da1b6568686e h1:MRM5ITcdelLK2j1vwZ3Je0FKVCfqOLp5zO6trqMLYs0=
+github.com/skip2/go-qrcode v0.0.0-20200617195104-da1b6568686e/go.mod h1:XV66xRDqSt+GTGFMVlhk3ULuV0y9ZmzeVGR4mloJI3M=
github.com/spf13/afero v1.15.0 h1:b/YBCLWAJdFWJTN9cLhiXXcD7mzKn9Dm86dNnfyQw1I=
github.com/spf13/afero v1.15.0/go.mod h1:NC2ByUVxtQs4b3sIUphxK0NioZnmxgyCrfzeuq8lxMg=
github.com/spf13/cast v1.10.0 h1:h2x0u2shc1QuLHfxi+cTJvs30+ZAHOGRic8uyGTDWxY=
diff --git a/interaction-audit.html b/interaction-audit.html
new file mode 100644
index 0000000..ac88ff3
--- /dev/null
+++ b/interaction-audit.html
@@ -0,0 +1,680 @@
+
+
+
+
+
+NL-IM 全系统交互审查报告
+
+
+
+
+
NL-IM 全系统交互审查报告
+
+ 来源:9 个并行审查子代理的静态代码审查 · 2026-08-13 · 范围:nl-um-vue-ts(PC 端 11 视图 + 40 组件 + 数据层)与 nl-im-uniapp(移动端 25 页面 + 34 组件 + 数据层),只读审查未改动任何代码;行号基于当日快照,修复前建议先复现确认。
+
+
+
+
+
+
最优先处理的三个结论
+
1. 移动端 H5 音视频通话目前实际不可用:srcObject 绑在 uni-video 上不生效导致画面全黑,语音通话没有媒体元素承载远端流全程静音;叠加两端同款的"群通话全群响铃""主叫取消后来电层永久卡死",通话链路建议整体修复后回归。
+
2. 两端 WebSocket 连接生命周期都有洞:重连次数耗尽后永久离线、PC 无心跳/移动端心跳无 ack 判定造成"假在线"、重连成功后都不补拉离线消息——组合起来就是"看着在线、实际丢消息"。
+
3. 一批"点了没反应"的失效交互源自组件事件契约错配:wd-action-sheet 负载取错(多个在用页面菜单失效)、wot confirm 回调被覆盖(退出登录无效)、MentionPicker 即开即关(两端 @ 失效)、群通知页监听不存在的事件(分页失效)。
+
+
+
发现分布(按功能域 × 严重程度)
+
+
+
+
横轴:发现条数 · 两端同款问题按 1 条计(已在明细中标注涉及的两端文件)
+
+
+
+
+
+
问题明细
+
+
+
+
+
+
+
+
+
diff --git a/internal/api/ai_handler.go b/internal/api/ai_handler.go
new file mode 100644
index 0000000..0c4778e
--- /dev/null
+++ b/internal/api/ai_handler.go
@@ -0,0 +1,249 @@
+/**
+ * package api
+ * 作用:AI 机器人相关 HTTP 接口。
+ *
+ * 权限模型:
+ * - AI 配置读写、机器人增删改:仅管理员(用户 id=1);
+ * - 机器人列表查询:所有登录用户(好友列表要展示机器人)。
+ */
+package api
+
+import (
+ "context"
+ "strconv"
+ "strings"
+ "time"
+
+ "xk-websocket-v2/internal/service"
+ "xk-websocket-v2/internal/utils"
+
+ "github.com/gin-gonic/gin"
+)
+
+/**
+ * requireAdmin
+ * 作用:校验当前登录用户是否为管理员(id=1)。
+ * 返回 false 时已写入 403 响应,调用方直接 return 即可
+ */
+func requireAdmin(c *gin.Context) bool {
+ userID, exists := c.Get("user_id")
+ if !exists || userID.(string) != service.AdminUserID {
+ utils.Error(c, utils.CodeForbidden, "仅管理员可操作")
+ return false
+ }
+ return true
+}
+
+/**
+ * maskAPIKey
+ * 作用:API 密钥脱敏(保留前4后4位),配置回显时避免明文泄露
+ */
+func maskAPIKey(key string) string {
+ if key == "" {
+ return ""
+ }
+ if len(key) <= 8 {
+ return "********"
+ }
+ return key[:4] + strings.Repeat("*", 8) + key[len(key)-4:]
+}
+
+/**
+ * GetAIConfigHandler
+ * 功能:获取全局 AI 配置(密钥脱敏返回)
+ * 路径:GET /api/ai/config (管理员)
+ */
+func GetAIConfigHandler(c *gin.Context) {
+ if !requireAdmin(c) {
+ return
+ }
+ cfg, err := service.AIBotSvc.GetAIConfig()
+ if err != nil {
+ utils.InternalError(c, "配置读取失败")
+ return
+ }
+ utils.SuccessWithData(c, gin.H{
+ "provider": cfg.Provider,
+ "base_url": cfg.BaseURL,
+ "api_key": maskAPIKey(cfg.APIKey),
+ "model": cfg.Model,
+ "enabled": cfg.Enabled,
+ // 前端据此判断"密钥是否已配置过"(脱敏值不能用于判断)
+ "has_key": cfg.APIKey != "",
+ }, "获取成功")
+}
+
+// saveAIConfigReq 保存 AI 配置请求体
+type saveAIConfigReq struct {
+ Provider string `json:"provider" binding:"required"`
+ BaseURL string `json:"base_url"`
+ // 为空表示不修改已保存的密钥(前端回显的是脱敏值)
+ APIKey string `json:"api_key"`
+ Model string `json:"model" binding:"required"`
+ Enabled bool `json:"enabled"`
+}
+
+/**
+ * SaveAIConfigHandler
+ * 功能:保存全局 AI 配置
+ * 路径:POST /api/ai/config (管理员)
+ */
+func SaveAIConfigHandler(c *gin.Context) {
+ if !requireAdmin(c) {
+ return
+ }
+ var req saveAIConfigReq
+ if err := c.ShouldBindJSON(&req); err != nil {
+ utils.BadRequest(c, "参数错误: "+err.Error())
+ return
+ }
+ // 前端可能把脱敏值原样提交回来,含 **** 的一律视为"未修改"
+ if strings.Contains(req.APIKey, "****") {
+ req.APIKey = ""
+ }
+ cfg, err := service.AIBotSvc.SaveAIConfig(req.Provider, req.BaseURL, req.APIKey, req.Model, req.Enabled)
+ if err != nil {
+ utils.InternalError(c, "保存失败: "+err.Error())
+ return
+ }
+ utils.SuccessWithData(c, gin.H{
+ "provider": cfg.Provider,
+ "base_url": cfg.BaseURL,
+ "api_key": maskAPIKey(cfg.APIKey),
+ "model": cfg.Model,
+ "enabled": cfg.Enabled,
+ "has_key": cfg.APIKey != "",
+ }, "保存成功")
+}
+
+/**
+ * TestAIConfigHandler
+ * 功能:测试当前 AI 配置连通性(发一条固定问候语,成功返回模型回复)
+ * 路径:POST /api/ai/config/test (管理员)
+ */
+func TestAIConfigHandler(c *gin.Context) {
+ if !requireAdmin(c) {
+ return
+ }
+ cfg, err := service.AIBotSvc.GetAIConfig()
+ if err != nil {
+ utils.InternalError(c, "配置读取失败")
+ return
+ }
+ provider, err := service.NewAIProvider(cfg)
+ if err != nil {
+ utils.BadRequest(c, err.Error())
+ return
+ }
+ ctx, cancel := context.WithTimeout(c.Request.Context(), 30*time.Second)
+ defer cancel()
+ reply, err := provider.Chat(ctx, []service.AIMessage{
+ {Role: "user", Content: "你好,请用一句话介绍你自己"},
+ })
+ if err != nil {
+ utils.BadRequest(c, "测试失败: "+err.Error())
+ return
+ }
+ utils.SuccessWithData(c, gin.H{"reply": reply}, "测试成功")
+}
+
+/**
+ * ListAIBotsHandler
+ * 功能:机器人列表。管理员传 all=1 返回全部(含停用),普通用户只见启用中的
+ * 路径:GET /api/ai/bots (登录用户)
+ */
+func ListAIBotsHandler(c *gin.Context) {
+ userID, _ := c.Get("user_id")
+ onlyEnabled := true
+ if c.Query("all") == "1" && userID != nil && userID.(string) == service.AdminUserID {
+ onlyEnabled = false
+ }
+ bots, err := service.AIBotSvc.ListBots(onlyEnabled)
+ if err != nil {
+ utils.InternalError(c, "查询失败")
+ return
+ }
+ utils.SuccessWithData(c, bots, "获取成功")
+}
+
+// botReq 创建/更新机器人请求体
+type botReq struct {
+ Name string `json:"name" binding:"required"`
+ Avatar string `json:"avatar"`
+ RolePrompt string `json:"role_prompt"`
+ Enabled *bool `json:"enabled"`
+}
+
+/**
+ * CreateAIBotHandler
+ * 功能:创建机器人(同时生成 users 虚拟用户)
+ * 路径:POST /api/ai/bots (管理员)
+ */
+func CreateAIBotHandler(c *gin.Context) {
+ if !requireAdmin(c) {
+ return
+ }
+ var req botReq
+ if err := c.ShouldBindJSON(&req); err != nil {
+ utils.BadRequest(c, "参数错误: "+err.Error())
+ return
+ }
+ bot, err := service.AIBotSvc.CreateBot(req.Name, req.Avatar, req.RolePrompt)
+ if err != nil {
+ utils.InternalError(c, "创建失败: "+err.Error())
+ return
+ }
+ utils.SuccessWithData(c, bot, "创建成功")
+}
+
+/**
+ * UpdateAIBotHandler
+ * 功能:更新机器人(名称/头像/角色设定/启停)
+ * 路径:POST /api/ai/bots/update/:id (管理员)
+ */
+func UpdateAIBotHandler(c *gin.Context) {
+ if !requireAdmin(c) {
+ return
+ }
+ id, err := strconv.ParseUint(c.Param("id"), 10, 64)
+ if err != nil || id == 0 {
+ utils.BadRequest(c, "机器人ID不合法")
+ return
+ }
+ var req botReq
+ if err := c.ShouldBindJSON(&req); err != nil {
+ utils.BadRequest(c, "参数错误: "+err.Error())
+ return
+ }
+ enabled := true
+ if req.Enabled != nil {
+ enabled = *req.Enabled
+ }
+ bot, err := service.AIBotSvc.UpdateBot(uint(id), req.Name, req.Avatar, req.RolePrompt, enabled)
+ if err != nil {
+ utils.InternalError(c, "更新失败: "+err.Error())
+ return
+ }
+ utils.SuccessWithData(c, bot, "更新成功")
+}
+
+/**
+ * DeleteAIBotHandler
+ * 功能:删除机器人(连带删除虚拟用户与群成员关系)
+ * 路径:POST /api/ai/bots/delete/:id (管理员)
+ */
+func DeleteAIBotHandler(c *gin.Context) {
+ if !requireAdmin(c) {
+ return
+ }
+ id, err := strconv.ParseUint(c.Param("id"), 10, 64)
+ if err != nil || id == 0 {
+ utils.BadRequest(c, "机器人ID不合法")
+ return
+ }
+ if err := service.AIBotSvc.DeleteBot(uint(id)); err != nil {
+ utils.InternalError(c, "删除失败: "+err.Error())
+ return
+ }
+ utils.Success(c, "删除成功")
+}
diff --git a/internal/api/auth_handler.go b/internal/api/auth_handler.go
index b8161eb..6c37705 100644
--- a/internal/api/auth_handler.go
+++ b/internal/api/auth_handler.go
@@ -10,8 +10,16 @@ import (
"xk-websocket-v2/internal/utils"
"github.com/gin-gonic/gin"
+ "github.com/spf13/viper"
)
+// isProdEnv 判断是否为生产环境(app.env 为 production/prod)。
+// 用于决定是否把验证码等敏感信息返回给客户端:仅非生产环境返回,便于本地调试。
+func isProdEnv() bool {
+ env := viper.GetString("app.env")
+ return env == "production" || env == "prod"
+}
+
/**
* LoginHandler
* 功能:用户登录
@@ -69,19 +77,20 @@ func RegisterHandler(c *gin.Context) {
return
}
- // 如果提供了验证码,验证验证码
+ // 如果提供了验证码,验证验证码(只校验不消费——注册还可能失败)
+ codeTarget, codeType := "", ""
if req.Code != "" {
// 根据邮箱或手机号确定类型
- codeType := "email"
+ codeType = "email"
if len(req.Phone) > 0 {
codeType = "sms"
}
- target := req.Email
+ codeTarget = req.Email
if codeType == "sms" {
- target = req.Phone
+ codeTarget = req.Phone
}
- valid, err := service.AuthSvc.VerifyCode(target, req.Code, codeType)
+ valid, err := service.AuthSvc.VerifyCode(codeTarget, req.Code, codeType)
if err != nil || !valid {
utils.BadRequest(c, "验证码无效或已过期")
return
@@ -91,6 +100,7 @@ func RegisterHandler(c *gin.Context) {
// 调用认证服务注册
user, err := service.AuthSvc.Register(&req)
if err != nil {
+ // 注册失败时验证码未被消费,用户可用同一验证码修正参数后重试
utils.BadRequest(c, err.Error())
return
}
@@ -102,6 +112,12 @@ func RegisterHandler(c *gin.Context) {
return
}
+ // 全流程成功(注册+Token)后才消费验证码,防止同码重放注册;
+ // 中途任何失败都保留验证码,把"可重试"的窗口留到最后一刻
+ if req.Code != "" {
+ service.AuthSvc.ConsumeCode(codeTarget, req.Code, codeType)
+ }
+
utils.SuccessWithData(c, model.RegisterResponse{
Token: token,
User: *user,
@@ -170,10 +186,12 @@ func SendEmailCodeHandler(c *gin.Context) {
return
}
- // 开发环境返回验证码,生产环境不应返回
- utils.SuccessWithData(c, gin.H{
- "code": code, // 仅开发环境,生产环境应移除
- }, "验证码已发送")
+ // 生产环境不返回验证码,避免直接绕过邮件校验;仅非生产环境返回以便调试
+ data := gin.H{}
+ if !isProdEnv() {
+ data["code"] = code
+ }
+ utils.SuccessWithData(c, data, "验证码已发送")
}
/**
@@ -199,8 +217,10 @@ func SendSmsCodeHandler(c *gin.Context) {
return
}
- // 开发环境返回验证码,生产环境不应返回
- utils.SuccessWithData(c, gin.H{
- "code": code, // 仅开发环境,生产环境应移除
- }, "验证码已发送")
+ // 生产环境不返回验证码,避免直接绕过短信校验;仅非生产环境返回以便调试
+ data := gin.H{}
+ if !isProdEnv() {
+ data["code"] = code
+ }
+ utils.SuccessWithData(c, data, "验证码已发送")
}
diff --git a/internal/api/call_handler.go b/internal/api/call_handler.go
index f69b7e0..b5049cd 100644
--- a/internal/api/call_handler.go
+++ b/internal/api/call_handler.go
@@ -103,6 +103,17 @@ type RoomInfoResponse struct {
// ========== API 处理函数 ==========
+// currentCallUserID 从 JWT 上下文获取当前登录用户ID。
+// 通话相关接口必须以此为准,忽略请求体里的 user_id,防止冒充他人进房/推流/信令。
+func currentCallUserID(c *gin.Context) string {
+ if uid, ok := c.Get("user_id"); ok {
+ if s, ok2 := uid.(string); ok2 {
+ return s
+ }
+ }
+ return ""
+}
+
// CreateCallRoomHandler 创建通话房间
// POST /api/call/room
func CreateCallRoomHandler(c *gin.Context) {
@@ -139,6 +150,14 @@ func JoinCallRoomHandler(c *gin.Context) {
return
}
+ // 身份以 JWT 为准,忽略请求体中的 user_id,防止冒充他人加入通话
+ if uid := currentCallUserID(c); uid != "" {
+ req.UserID = uid
+ } else {
+ utils.Unauthorized(c, "未认证")
+ return
+ }
+
ms := mediaserver.GetServer()
if ms.GetConfig() == nil || !ms.GetConfig().Enabled {
utils.Error(c, 503, "媒体服务未启用")
@@ -281,6 +300,14 @@ func LeaveCallRoomHandler(c *gin.Context) {
return
}
+ // 身份以 JWT 为准,防止冒充他人离开通话房间
+ if uid := currentCallUserID(c); uid != "" {
+ req.UserID = uid
+ } else {
+ utils.Unauthorized(c, "未认证")
+ return
+ }
+
ms := mediaserver.GetServer()
room := ms.GetRoom(req.RoomID)
if room == nil {
@@ -359,6 +386,14 @@ func WebRTCOfferHandler(c *gin.Context) {
return
}
+ // 身份以 JWT 为准,防止冒充他人提交 SDP Offer
+ if uid := currentCallUserID(c); uid != "" {
+ req.UserID = uid
+ } else {
+ utils.Unauthorized(c, "未认证")
+ return
+ }
+
ms := mediaserver.GetServer()
sfu := ms.GetSFU()
if sfu == nil {
@@ -387,6 +422,14 @@ func WebRTCICEHandler(c *gin.Context) {
return
}
+ // 身份以 JWT 为准,防止冒充他人提交 ICE 候选
+ if uid := currentCallUserID(c); uid != "" {
+ req.UserID = uid
+ } else {
+ utils.Unauthorized(c, "未认证")
+ return
+ }
+
ms := mediaserver.GetServer()
sfu := ms.GetSFU()
if sfu == nil {
@@ -406,9 +449,11 @@ func WebRTCICEHandler(c *gin.Context) {
// GetICEServersHandler 获取 ICE 服务器配置
// GET /api/call/ice-servers
func GetICEServersHandler(c *gin.Context) {
- userID := c.Query("user_id")
+ // 凭证绑定当前登录用户,取自 JWT,避免为任意 user_id 签发 TURN 凭证被滥用
+ userID := currentCallUserID(c)
if userID == "" {
- userID = "anonymous"
+ utils.Unauthorized(c, "未认证")
+ return
}
servers := getICEServers(userID)
diff --git a/internal/api/contact_handler.go b/internal/api/contact_handler.go
index a92521e..73cf5ca 100644
--- a/internal/api/contact_handler.go
+++ b/internal/api/contact_handler.go
@@ -232,11 +232,19 @@ func CreateGroupHandler(c *gin.Context) {
*/
func UpdateGroupHandler(c *gin.Context) {
userID, _ := c.Get("user_id")
- groupID, _ := strconv.ParseUint(c.Param("id"), 10, 32)
+ // 校验 :id 解析错误:原实现忽略 error,非法ID会静默变成 0 去更新不存在的分组
+ groupID, err := strconv.ParseUint(c.Param("id"), 10, 32)
+ if err != nil {
+ utils.BadRequest(c, "分组ID不合法")
+ return
+ }
+ // 使用指针字段表达"字段是否出现"语义:
+ // 原实现用 !="" / >0 判断,导致无法把排序值置 0;
+ // 指针为 nil 表示前端没传该字段(跳过),非 nil 表示要更新(含零值)
var req struct {
- GroupName string `json:"group_name"`
- SortOrder int `json:"sort_order"`
+ GroupName *string `json:"group_name"`
+ SortOrder *int `json:"sort_order"`
}
if err := c.ShouldBindJSON(&req); err != nil {
@@ -245,11 +253,16 @@ func UpdateGroupHandler(c *gin.Context) {
}
updates := make(map[string]interface{})
- if req.GroupName != "" {
- updates["group_name"] = req.GroupName
+ if req.GroupName != nil {
+ // 分组名不允许被清空(业务约束),显式传空串视为非法
+ if *req.GroupName == "" {
+ utils.BadRequest(c, "分组名不能为空")
+ return
+ }
+ updates["group_name"] = *req.GroupName
}
- if req.SortOrder > 0 {
- updates["sort_order"] = req.SortOrder
+ if req.SortOrder != nil {
+ updates["sort_order"] = *req.SortOrder
}
if err := service.ContactSvc.UpdateGroup(uint(groupID), userID.(string), updates); err != nil {
@@ -267,7 +280,12 @@ func UpdateGroupHandler(c *gin.Context) {
*/
func DeleteGroupHandler(c *gin.Context) {
userID, _ := c.Get("user_id")
- groupID, _ := strconv.ParseUint(c.Param("id"), 10, 32)
+ // 校验 :id 解析错误:非法ID静默变 0 会误删/误匹配
+ groupID, err := strconv.ParseUint(c.Param("id"), 10, 32)
+ if err != nil {
+ utils.BadRequest(c, "分组ID不合法")
+ return
+ }
if err := service.ContactSvc.DeleteGroup(uint(groupID), userID.(string)); err != nil {
utils.BadRequest(c, "删除失败: "+err.Error())
@@ -305,12 +323,16 @@ func UpdateContactHandler(c *gin.Context) {
userID, _ := c.Get("user_id")
contactID := c.Param("id")
+ // 全部使用指针字段表达"字段是否出现"语义:
+ // 原实现 remark_name!="" / group_id>0 才写入,导致无法清空备注、
+ // 无法把联系人移回默认分组(0);指针为 nil 表示未传(跳过),非 nil 表示更新(含零值)
var req struct {
- RemarkName string `json:"remark_name"`
- GroupID uint `json:"group_id"`
- IsTop *bool `json:"is_top"`
- IsMuted *bool `json:"is_muted"`
- IsBlocked *bool `json:"is_blocked"`
+ RemarkName *string `json:"remark_name"`
+ GroupID *uint `json:"group_id"`
+ IsTop *bool `json:"is_top"`
+ IsMuted *bool `json:"is_muted"`
+ IsBlocked *bool `json:"is_blocked"`
+ IsSpecialCare *bool `json:"is_special_care"` // 特别关心(此前漏了该字段,PC端开关一直不生效)
}
if err := c.ShouldBindJSON(&req); err != nil {
@@ -319,11 +341,11 @@ func UpdateContactHandler(c *gin.Context) {
}
updates := make(map[string]interface{})
- if req.RemarkName != "" {
- updates["remark_name"] = req.RemarkName
+ if req.RemarkName != nil {
+ updates["remark_name"] = *req.RemarkName
}
- if req.GroupID > 0 {
- updates["group_id"] = req.GroupID
+ if req.GroupID != nil {
+ updates["group_id"] = *req.GroupID
}
if req.IsTop != nil {
updates["is_top"] = *req.IsTop
@@ -334,6 +356,9 @@ func UpdateContactHandler(c *gin.Context) {
if req.IsBlocked != nil {
updates["is_blocked"] = *req.IsBlocked
}
+ if req.IsSpecialCare != nil {
+ updates["is_special_care"] = *req.IsSpecialCare
+ }
if err := service.ContactSvc.UpdateContact(userID.(string), contactID, updates); err != nil {
utils.BadRequest(c, "更新失败: "+err.Error())
diff --git a/internal/api/handler.go b/internal/api/handler.go
index 719f530..d7fb83c 100644
--- a/internal/api/handler.go
+++ b/internal/api/handler.go
@@ -38,12 +38,14 @@ func SendHandler(c *gin.Context) {
return
}
- // 2. 获取发送者 ID (从 Header 中获取,模拟鉴权)
- // 生产环境应从 JWT Token 中解析 UserID
- senderID := c.GetHeader("X-User-ID")
- if senderID == "" {
- senderID = "system" // 默认为系统消息
+ // 2. 获取发送者 ID:强制取自 JWT(由认证中间件注入 user_id),
+ // 禁止客户端通过 Header/Body 伪造他人身份发送消息或信令
+ uid, ok := c.Get("user_id")
+ if !ok {
+ utils.Unauthorized(c, "未认证")
+ return
}
+ senderID := uid.(string)
// 3. [关键修复] 获取发送端的 WebSocket ClientID
// 前端在调用此接口时,必须带上自己的 socket_client_id
@@ -62,9 +64,16 @@ func SendHandler(c *gin.Context) {
}
// 5. 调用核心业务逻辑处理消息,返回持久化后的消息体供前端对齐 ID
- savedMsg := service.ChatSvc.HandleUserMessage(mockClient, &req)
+ savedMsg, err := service.ChatSvc.HandleUserMessage(mockClient, &req)
- // 6. 返回成功响应(含服务端消息 ID,避免前端重复展示)
+ // 6. 校验被拒/落库失败必须返回错误:原实现只看消息指针,
+ // 被拒消息也响应"消息已发送",发送方界面显示成功但对方永远收不到
+ if err != nil {
+ utils.Error(c, utils.CodeForbidden, err.Error())
+ return
+ }
+
+ // 7. 返回成功响应(含服务端消息 ID,避免前端重复展示)
if savedMsg != nil {
utils.SuccessWithData(c, savedMsg, "消息已发送")
return
@@ -98,6 +107,14 @@ func BindHandler(c *gin.Context) {
utils.BadRequest(c, "参数错误")
return
}
+ // 绑定的用户ID强制取自 JWT,忽略请求体中的 user_id,
+ // 防止攻击者把任意 user_id 绑定到自己的 client_id 从而劫持他人的消息推送
+ uid, ok := c.Get("user_id")
+ if !ok {
+ utils.Unauthorized(c, "未认证")
+ return
+ }
+ req.UserID = uid.(string)
// 调用服务层进行绑定
service.ChatSvc.BindUserByClientID(req.ClientID, req.UserID)
utils.Success(c, "绑定成功")
@@ -131,6 +148,17 @@ func HistoryHandler(c *gin.Context) {
return
}
+ // 鉴权:仅房间成员可拉取历史消息,防止任意登录用户凭 room_id 窥探他人聊天记录
+ uid, ok := c.Get("user_id")
+ if !ok {
+ utils.Unauthorized(c, "未认证")
+ return
+ }
+ if !service.RoomSvc.IsRoomMember(roomID, uid.(string)) {
+ utils.Forbidden(c, "无权查看该房间消息")
+ return
+ }
+
// 分页参数
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", "50"))
@@ -145,12 +173,22 @@ func HistoryHandler(c *gin.Context) {
var msgs []model.ChatMessage
var total int64
- // 获取总数
- service.ChatSvc.DB.Model(&model.ChatMessage{}).Where("room_id = ?", roomID).Count(&total)
+ // 排除当前用户已删除的消息("删除仅对我生效",见 MessageDeletion 模型注释)。
+ // 用 NOT IN 子查询而非 JOIN:删除记录量级远小于消息量级,且保持分页语义简单
+ deletedSubQuery := service.ChatSvc.DB.Model(&model.MessageDeletion{}).
+ Select("message_id").
+ Where("user_id = ?", uid.(string))
+
+ // 获取总数(同样排除已删除,保证 total 与实际可见条数一致)
+ service.ChatSvc.DB.Model(&model.ChatMessage{}).
+ Where("room_id = ?", roomID).
+ Where("id NOT IN (?)", deletedSubQuery).
+ Count(&total)
// 分页查询
offset := (page - 1) * pageSize
result := service.ChatSvc.DB.Where("room_id = ?", roomID).
+ Where("id NOT IN (?)", deletedSubQuery).
Order("created_at desc").
Offset(offset).
Limit(pageSize).
@@ -222,6 +260,54 @@ func RecallMessageHandler(c *gin.Context) {
utils.SuccessWithData(c, msg, "撤回成功")
}
+/**
+ * DeleteMessageHandler
+ * 功能:删除消息(仅对操作者自己生效,对方仍可见,区别于撤回)。
+ * 逻辑:不删除消息本体,仅在 message_deletions 表插入删除记录,
+ * 历史消息查询时排除,避免"本地删了、刷新又回来"的假删除问题。
+ * 路径:POST /api/messages/delete
+ */
+func DeleteMessageHandler(c *gin.Context) {
+ userID, exists := c.Get("user_id")
+ if !exists {
+ utils.Unauthorized(c, "未认证")
+ return
+ }
+
+ var req struct {
+ MessageID uint `json:"message_id" binding:"required"`
+ }
+ if err := c.ShouldBindJSON(&req); err != nil {
+ utils.BadRequest(c, "参数错误")
+ return
+ }
+
+ // 校验消息存在且操作者是所在房间成员(防止删除无权查看的消息记录)
+ var msg model.ChatMessage
+ if err := service.ChatSvc.DB.First(&msg, req.MessageID).Error; err != nil {
+ utils.NotFound(c, "消息不存在")
+ return
+ }
+ if !service.RoomSvc.IsRoomMember(msg.RoomID, userID.(string)) {
+ utils.Forbidden(c, "无权删除该消息")
+ return
+ }
+
+ // 幂等插入:重复删除同一条消息时唯一索引冲突,视为成功
+ deletion := model.MessageDeletion{
+ MessageID: req.MessageID,
+ UserID: userID.(string),
+ }
+ if err := service.ChatSvc.DB.Create(&deletion).Error; err != nil {
+ if !utils.IsDuplicateEntryError(err) {
+ utils.InternalError(c, "删除失败")
+ return
+ }
+ }
+
+ utils.Success(c, "删除成功")
+}
+
/**
* SyncMessagesHandler
* 功能:全量同步消息 (V1 兼容)。
@@ -255,7 +341,13 @@ func HealthHandler(c *gin.Context) {
* 用途:WebRTC 前端在建立 PeerConnection 前需调用此接口。
*/
func ICEHandler(c *gin.Context) {
- userID := c.Query("user_id")
+ // 凭证绑定当前登录用户,取自 JWT,避免匿名为任意 user_id 签发 TURN 凭证被滥用
+ uid, ok := c.Get("user_id")
+ if !ok {
+ utils.Unauthorized(c, "未认证")
+ return
+ }
+ userID := uid.(string)
// 生成临时凭证 (HMAC-SHA1)
username, credential := turnserver.GenerateCredentials(userID)
diff --git a/internal/api/qrcode_handler.go b/internal/api/qrcode_handler.go
new file mode 100644
index 0000000..61a4567
--- /dev/null
+++ b/internal/api/qrcode_handler.go
@@ -0,0 +1,200 @@
+/**
+ * package api
+ * 作用:App 扫码登录 PC 的 HTTP 接口。
+ *
+ * 流程:
+ * 1. PC(未登录)POST /qrcode/generate → 返回 qr_id + 二维码图片(base64)
+ * 2. PC 轮询 GET /qrcode/status?qr_id= → pending/scanned/confirmed/cancelled/expired
+ * 3. App(已登录)POST /qrcode/scan → 标记已扫码
+ * 4. App POST /qrcode/confirm → 签发 token,PC 下次轮询取走并登录
+ * App POST /qrcode/cancel → 取消本次登录
+ *
+ * 鉴权划分:generate/status 必须公开(PC 此时未登录),
+ * scan/confirm/cancel 必须走 JWT 认证(确认人身份来自 token,不可伪造)。
+ */
+package api
+
+import (
+ "encoding/base64"
+ "time"
+
+ "xk-websocket-v2/internal/service"
+ "xk-websocket-v2/internal/utils"
+
+ "github.com/gin-gonic/gin"
+ qrcode "github.com/skip2/go-qrcode"
+)
+
+// 二维码内容前缀:App 扫码后据此识别"这是本系统的登录码",其他内容一律不处理
+const qrLoginContentPrefix = "nlim-login:"
+
+// qrCodeReq scan/confirm/cancel 的公共请求体
+type qrCodeReq struct {
+ QRID string `json:"qr_id" binding:"required"`
+}
+
+/**
+ * GenerateQRCodeHandler
+ * 功能:生成扫码登录二维码(PC 登录页调用,无需鉴权)
+ * 路径:POST /api/qrcode/generate
+ * 返回:qr_id、二维码 PNG 的 dataURL、有效秒数
+ * 说明:图片由服务端生成,前端无需引入二维码库;qr_id 为 32 字节加密随机数,
+ * TTL 2 分钟自动过期,被刷接口时 Redis 只会堆积少量短命 key,风险可控
+ */
+func GenerateQRCodeHandler(c *gin.Context) {
+ qrID, err := service.QRCodeLoginSvc.Generate()
+ if err != nil {
+ utils.InternalError(c, err.Error())
+ return
+ }
+
+ // 生成二维码 PNG(256px,中等纠错——内容短,M 级足够且图案更稀疏易扫)
+ png, err := qrcode.Encode(qrLoginContentPrefix+qrID, qrcode.Medium, 256)
+ if err != nil {
+ utils.InternalError(c, "二维码图片生成失败")
+ return
+ }
+
+ utils.SuccessWithData(c, gin.H{
+ "qr_id": qrID,
+ "qr_image": "data:image/png;base64," + base64.StdEncoding.EncodeToString(png),
+ "expires_in": 120,
+ }, "生成成功")
+}
+
+/**
+ * QRCodeStatusHandler
+ * 功能:PC 长轮询二维码状态(无需鉴权,qr_id 即凭证)
+ * 路径:GET /api/qrcode/status?qr_id=xxx&known=pending
+ * 说明:携带 known(客户端已知状态)时服务端最多 hold 25 秒,
+ * 状态一变化立即返回——相比 2 秒短轮询,请求数降一个量级且感知更即时。
+ * 25 秒小于前端 axios 超时(35s)与常见网关超时(60s),不会被中间层掐断。
+ * 不带 known 时立即返回当前状态(兼容行为)。
+ * 返回:status 为 confirmed 时附带 token 与用户信息(一次性,取走即删)
+ */
+func QRCodeStatusHandler(c *gin.Context) {
+ qrID := c.Query("qr_id")
+ if qrID == "" {
+ utils.BadRequest(c, "缺少 qr_id 参数")
+ return
+ }
+
+ // 长轮询等待阶段(不消费 token):hold 到状态变化/过期/超时/客户端断开
+ if known := c.Query("known"); known != "" {
+ session, expired, err := service.QRCodeLoginSvc.WaitForChange(c.Request.Context(), qrID, known, 25*time.Second)
+ if err != nil {
+ utils.InternalError(c, err.Error())
+ return
+ }
+ if expired {
+ utils.SuccessWithData(c, gin.H{"status": "expired"}, "二维码已过期")
+ return
+ }
+ // 非 confirmed 直接返回;confirmed 落到下方的原子消费逻辑
+ // (必须重新走 DEL 仲裁,并发轮询时保证 token 只交付一次)
+ if session.Status != service.QRStatusConfirmed {
+ utils.SuccessWithData(c, gin.H{"status": session.Status}, "查询成功")
+ return
+ }
+ }
+
+ // 窥视 → 取用户 → 原子消费 三步走:
+ // 用户查询可能失败(DB 抖动),若先消费(DEL)再查询、查询失败,
+ // 一次性 token 已被删除,手机确认成功而 PC 永远登录不上。
+ // 把可失败的准备工作全部放在消费之前,失败时 token 未消费,PC 长轮询下一轮重试
+ session, err := service.QRCodeLoginSvc.GetStatus(qrID)
+ if err != nil {
+ utils.InternalError(c, err.Error())
+ return
+ }
+ if session == nil {
+ utils.SuccessWithData(c, gin.H{"status": "expired"}, "二维码已过期")
+ return
+ }
+ if session.Status != service.QRStatusConfirmed {
+ utils.SuccessWithData(c, gin.H{"status": session.Status}, "查询成功")
+ return
+ }
+
+ // 确认成功:先取用户信息(与账号密码登录的响应结构对齐,
+ // GetUserByID 内部已清除密码字段),此时 token 尚未消费、失败可重试
+ user, uErr := service.UserSvc.GetUserByID(session.ScannerID)
+ if uErr != nil {
+ utils.InternalError(c, "获取用户信息失败")
+ return
+ }
+
+ // 最后一步才原子消费:token 唯一交付仍由 DEL 仲裁保证,
+ // 被并发轮询抢先(deleted==0)按已过期返回,不会重复交付
+ consumed, expired, err := service.QRCodeLoginSvc.GetStatusAndConsume(qrID)
+ if err != nil {
+ utils.InternalError(c, err.Error())
+ return
+ }
+ if expired || consumed == nil {
+ utils.SuccessWithData(c, gin.H{"status": "expired"}, "二维码已过期")
+ return
+ }
+ utils.SuccessWithData(c, gin.H{
+ "status": consumed.Status,
+ "token": consumed.Token,
+ "user": user,
+ }, "查询成功")
+}
+
+/**
+ * ScanQRCodeHandler
+ * 功能:App 扫码上报(需鉴权),PC 端随即显示"请在手机上确认"
+ * 路径:POST /api/qrcode/scan
+ */
+func ScanQRCodeHandler(c *gin.Context) {
+ var req qrCodeReq
+ if err := c.ShouldBindJSON(&req); err != nil {
+ utils.BadRequest(c, "参数错误")
+ return
+ }
+ uid, _ := c.Get("user_id")
+ if err := service.QRCodeLoginSvc.Scan(req.QRID, uid.(string)); err != nil {
+ utils.BadRequest(c, err.Error())
+ return
+ }
+ utils.Success(c, "扫码成功")
+}
+
+/**
+ * ConfirmQRCodeHandler
+ * 功能:App 确认登录(需鉴权且必须是扫码本人),为 PC 签发登录 token
+ * 路径:POST /api/qrcode/confirm
+ */
+func ConfirmQRCodeHandler(c *gin.Context) {
+ var req qrCodeReq
+ if err := c.ShouldBindJSON(&req); err != nil {
+ utils.BadRequest(c, "参数错误")
+ return
+ }
+ uid, _ := c.Get("user_id")
+ if err := service.QRCodeLoginSvc.Confirm(req.QRID, uid.(string), utils.GenerateToken); err != nil {
+ utils.BadRequest(c, err.Error())
+ return
+ }
+ utils.Success(c, "确认成功")
+}
+
+/**
+ * CancelQRCodeHandler
+ * 功能:App 取消本次扫码登录(需鉴权),PC 端提示后刷新二维码
+ * 路径:POST /api/qrcode/cancel
+ */
+func CancelQRCodeHandler(c *gin.Context) {
+ var req qrCodeReq
+ if err := c.ShouldBindJSON(&req); err != nil {
+ utils.BadRequest(c, "参数错误")
+ return
+ }
+ uid, _ := c.Get("user_id")
+ if err := service.QRCodeLoginSvc.Cancel(req.QRID, uid.(string)); err != nil {
+ utils.BadRequest(c, err.Error())
+ return
+ }
+ utils.Success(c, "已取消")
+}
diff --git a/internal/api/room_handler.go b/internal/api/room_handler.go
index bb2ba65..5d1b4ee 100644
--- a/internal/api/room_handler.go
+++ b/internal/api/room_handler.go
@@ -67,6 +67,11 @@ func CreateRoomHandler(c *gin.Context) {
func GetRoomHandler(c *gin.Context) {
roomID := c.Param("id")
+ // 仅房间成员可查看房间信息,防止非成员凭 room_id 窥探
+ if !requireRoomMember(c, roomID) {
+ return
+ }
+
room, err := service.RoomSvc.GetRoom(roomID)
if err != nil {
utils.NotFound(c, "房间不存在")
@@ -76,6 +81,25 @@ func GetRoomHandler(c *gin.Context) {
utils.SuccessWithData(c, room, "获取成功")
}
+/**
+ * requireRoomMember
+ * 功能:校验当前登录用户是否为该房间/群成员,不是则写 403 并返回 false。
+ * 为什么这样写:群信息、成员列表、公告、设置等读接口原先不校验调用者身份,
+ * 任意登录用户凭 room_id 即可窥探任意群数据,这里统一收口做成员校验。
+ */
+func requireRoomMember(c *gin.Context, roomID string) bool {
+ uid, ok := c.Get("user_id")
+ if !ok {
+ utils.Unauthorized(c, "未认证")
+ return false
+ }
+ if !service.RoomSvc.IsRoomMember(roomID, uid.(string)) {
+ utils.Forbidden(c, "无权访问该群")
+ return false
+ }
+ return true
+}
+
/**
* CreateChatGroupHandler
* 功能:创建群聊
@@ -270,6 +294,11 @@ func ListUserGroupsHandler(c *gin.Context) {
func GetGroupInfoHandler(c *gin.Context) {
roomID := c.Param("room_id")
+ // 仅群成员可查看群信息
+ if !requireRoomMember(c, roomID) {
+ return
+ }
+
room, err := service.RoomSvc.GetRoom(roomID)
if err != nil {
if err == gorm.ErrRecordNotFound {
@@ -311,6 +340,11 @@ func ListGroupMembersHandler(c *gin.Context) {
roomID := c.Param("room_id")
keyword := c.Query("keyword")
+ // 仅群成员可查看成员列表
+ if !requireRoomMember(c, roomID) {
+ return
+ }
+
var members []model.RoomMember
var err error
@@ -388,33 +422,8 @@ func AddGroupMembersHandler(c *gin.Context) {
}
notifExtraJSON, _ := json.Marshal(notifExtra)
- // 为所有成员发送群通知(包括新加入的成员)
- for _, uid := range memberIDs {
- // 群通知消息
- notifMsg := model.ChatMessage{
- RoomID: roomID,
- SenderUserID: userID.(string),
- ReceiverUserID: roomID,
- MessageType: model.MessageTypeGroupNotif,
- Content: notifContent,
- Extra: string(notifExtraJSON),
- }
-
- if err := service.ChatSvc.DB.Create(¬ifMsg).Error; err == nil {
- // 更新会话
- if service.ConversationSvc != nil {
- _ = service.ConversationSvc.UpsertConversationOnMessage(uid, roomID, roomID, notifMsg, false)
- }
-
- // 推送消息
- pushMsg := model.WsPayload{
- RequestType: "receive_message",
- Data: notifMsg,
- }
- msgBytes, _ := json.Marshal(pushMsg)
- service.ChatSvc.DispatchMessage(uid, msgBytes)
- }
- }
+ // 为所有成员发送群通知(包括新加入的成员):房间级落库一次,按成员推送
+ service.ChatSvc.BroadcastGroupNotification(userID.(string), roomID, notifContent, string(notifExtraJSON), memberIDs)
// 为新加入的成员发送系统消息
for _, memberID := range req.MemberIDs {
@@ -516,37 +525,14 @@ func RemoveGroupMemberHandler(c *gin.Context) {
}
systemExtraJSON, _ := json.Marshal(systemExtra)
- // 为所有成员发送群通知
+ // 为所有成员发送群通知(被移除的成员不接收):房间级落库一次,按成员推送
+ remainIDs := make([]string, 0, len(memberIDs))
for _, uid := range memberIDs {
- if uid == memberID {
- continue // 被移除的成员不接收群通知
- }
-
- // 群通知消息
- notifMsg := model.ChatMessage{
- RoomID: roomID,
- SenderUserID: operatorID.(string),
- ReceiverUserID: roomID,
- MessageType: model.MessageTypeGroupNotif,
- Content: notifContent,
- Extra: string(notifExtraJSON),
- }
-
- if err := service.ChatSvc.DB.Create(¬ifMsg).Error; err == nil {
- // 更新会话
- if service.ConversationSvc != nil {
- _ = service.ConversationSvc.UpsertConversationOnMessage(uid, roomID, roomID, notifMsg, false)
- }
-
- // 推送消息
- pushMsg := model.WsPayload{
- RequestType: "receive_message",
- Data: notifMsg,
- }
- msgBytes, _ := json.Marshal(pushMsg)
- service.ChatSvc.DispatchMessage(uid, msgBytes)
+ if uid != memberID {
+ remainIDs = append(remainIDs, uid)
}
}
+ service.ChatSvc.BroadcastGroupNotification(operatorID.(string), roomID, notifContent, string(notifExtraJSON), remainIDs)
// 为被移除的成员发送系统消息
systemMsg := model.ChatMessage{
@@ -718,32 +704,8 @@ func QuitGroupHandler(c *gin.Context) {
}
notifExtraJSON, _ := json.Marshal(notifExtra)
- // 为所有剩余成员发送群通知
- for _, uid := range memberIDs {
- notifMsg := model.ChatMessage{
- RoomID: roomID,
- SenderUserID: userID.(string),
- ReceiverUserID: roomID,
- MessageType: model.MessageTypeGroupNotif,
- Content: notifContent,
- Extra: string(notifExtraJSON),
- }
-
- if err := service.ChatSvc.DB.Create(¬ifMsg).Error; err == nil {
- // 更新会话
- if service.ConversationSvc != nil {
- _ = service.ConversationSvc.UpsertConversationOnMessage(uid, roomID, roomID, notifMsg, false)
- }
-
- // 推送消息
- pushMsg := model.WsPayload{
- RequestType: "receive_message",
- Data: notifMsg,
- }
- msgBytes, _ := json.Marshal(pushMsg)
- service.ChatSvc.DispatchMessage(uid, msgBytes)
- }
- }
+ // 为所有剩余成员发送群通知:房间级落库一次,按成员推送
+ service.ChatSvc.BroadcastGroupNotification(userID.(string), roomID, notifContent, string(notifExtraJSON), memberIDs)
}
}
}
@@ -836,6 +798,11 @@ func GetGroupNotificationsHandler(c *gin.Context) {
func GetGroupAnnouncementHandler(c *gin.Context) {
roomID := c.Param("room_id")
+ // 仅群成员可查看群公告
+ if !requireRoomMember(c, roomID) {
+ return
+ }
+
announcement, err := service.RoomSvc.GetGroupAnnouncement(roomID)
if err != nil {
if err == gorm.ErrRecordNotFound {
@@ -877,6 +844,40 @@ func UpdateGroupAnnouncementHandler(c *gin.Context) {
utils.Success(c, "更新成功")
}
+/**
+ * UpdateMyNicknameHandler
+ * 功能:更新"我在本群的昵称"(群名片)
+ * 路径:POST /api/groups/:room_id/nickname
+ * 说明:操作者从 JWT 取,只能改自己的群名片;此前该接口缺失导致移动端改群昵称功能 100% 失败
+ */
+func UpdateMyNicknameHandler(c *gin.Context) {
+ userID, _ := c.Get("user_id")
+ roomID := c.Param("room_id")
+
+ var req struct {
+ // 允许空串(清除群名片),因此不加 required
+ Nickname string `json:"nickname"`
+ }
+
+ if err := c.ShouldBindJSON(&req); err != nil {
+ utils.BadRequest(c, "参数错误: "+err.Error())
+ return
+ }
+
+ // 群名片长度限制,与数据库 varchar(100) 对齐(按字符数收紧到 30,避免界面溢出)
+ if len([]rune(req.Nickname)) > 30 {
+ utils.BadRequest(c, "群昵称不能超过30个字符")
+ return
+ }
+
+ if err := service.RoomSvc.UpdateMemberNickname(roomID, userID.(string), req.Nickname); err != nil {
+ utils.BadRequest(c, err.Error())
+ return
+ }
+
+ utils.Success(c, "修改成功")
+}
+
/**
* MuteGroupMemberHandler
* 功能:禁言群成员
@@ -919,6 +920,12 @@ func GetMemberMuteStatusHandler(c *gin.Context) {
roomID := c.Param("room_id")
memberID := c.Param("user_id")
+ // 鉴权:仅本群成员可查询,防止任意登录用户越权读取任意群成员的禁言状态
+ // (与同文件其他群接口的 requireRoomMember 校验保持一致)
+ if !requireRoomMember(c, roomID) {
+ return
+ }
+
isMuted, mutedUntil, err := service.RoomSvc.IsGroupMemberMuted(roomID, memberID)
if err != nil {
if err == gorm.ErrRecordNotFound {
@@ -942,6 +949,12 @@ func GetMemberMuteStatusHandler(c *gin.Context) {
*/
func GetGroupSettingsHandler(c *gin.Context) {
roomID := c.Param("room_id")
+
+ // 仅群成员可查看群设置
+ if !requireRoomMember(c, roomID) {
+ return
+ }
+
room, err := service.RoomSvc.GetRoom(roomID)
if err != nil {
utils.NotFound(c, "群聊不存在")
diff --git a/internal/api/settings_handler.go b/internal/api/settings_handler.go
index 73a71c5..205e506 100644
--- a/internal/api/settings_handler.go
+++ b/internal/api/settings_handler.go
@@ -13,6 +13,8 @@ import (
"xk-websocket-v2/internal/utils"
"github.com/gin-gonic/gin"
+ "gorm.io/gorm"
+ "gorm.io/gorm/clause"
)
/**
@@ -58,16 +60,49 @@ func UpdateUserSettingsHandler(c *gin.Context) {
return
}
+ // 上限校验:设置项数量与 key/value 长度设置合理上限,
+ // 防止单次提交超大 map 触发大量写库(滥用/误用防护)
+ const (
+ maxSettingsPerRequest = 100
+ maxKeyLen = 100
+ maxValueLen = 4000
+ )
+ if len(req.Settings) > maxSettingsPerRequest {
+ utils.BadRequest(c, "设置项过多")
+ return
+ }
+ for key, val := range req.Settings {
+ if key == "" || len(key) > maxKeyLen {
+ utils.BadRequest(c, "设置键不合法")
+ return
+ }
+ if len(val) > maxValueLen {
+ utils.BadRequest(c, "设置值过长")
+ return
+ }
+ }
+
uid := userID.(string)
now := time.Now()
- for key, val := range req.Settings {
- setting := model.UserSetting{
- UserID: uid,
- SettingKey: key,
- SettingValue: val,
- UpdatedAt: now,
+ // 事务包裹批量写入:原实现循环逐条 Save,中途某条失败时前面的 key 已落库、
+ // 接口却返回错误,用户设置停留在"新旧混合"的中间状态;
+ // 改为同一事务内全部成功才提交,任一失败整体回滚,保证批量更新的原子性
+ if err := service.ChatSvc.DB.Transaction(func(tx *gorm.DB) error {
+ for key, val := range req.Settings {
+ setting := model.UserSetting{
+ UserID: uid,
+ SettingKey: key,
+ SettingValue: val,
+ UpdatedAt: now,
+ }
+ if err := tx.Save(&setting).Error; err != nil {
+ return err
+ }
}
- service.ChatSvc.DB.Save(&setting)
+ return nil
+ }); err != nil {
+ utils.InternalError(c, "保存设置失败")
+ return
}
GetUserSettingsHandler(c)
@@ -94,31 +129,78 @@ func MarkMessagesReadHandler(c *gin.Context) {
uid := userID.(string)
now := time.Now()
- readIDs := make([]uint, 0, len(req.MessageIDs))
- for _, mid := range req.MessageIDs {
- receipt := model.MessageReadReceipt{
- MessageID: mid,
+
+ // 批量查询消息:一条 IN 查询取回,
+ // 原实现每条消息 FirstOrCreate + First 各一次(N+1),一次标记50条已读要打100+次库
+ var msgs []model.ChatMessage
+ if err := service.ChatSvc.DB.Where("id IN ?", req.MessageIDs).Find(&msgs).Error; err != nil {
+ utils.InternalError(c, "db error")
+ return
+ }
+
+ // 越权校验:只允许标记"自己所在房间"的消息为已读(口径与 HistoryHandler 的
+ // IsRoomMember 鉴权对称)。消息ID由客户端任意提交,不校验成员身份的话,
+ // 任何登录用户都能对别人房间的消息伪造已读回执,发送者会收到虚假的 messages_read 推送。
+ // 同批消息通常集中在少数房间,按房间去重后每个房间只查一次成员关系,避免逐条 N 次查库;
+ // 非成员房间的消息静默剔除(批量接口部分生效语义),readIDs 只返回真正标记成功的部分。
+ roomAllowed := make(map[string]bool)
+ visibleMsgs := msgs[:0]
+ for _, msg := range msgs {
+ allowed, checked := roomAllowed[msg.RoomID]
+ if !checked {
+ allowed = service.RoomSvc.IsRoomMember(msg.RoomID, uid)
+ roomAllowed[msg.RoomID] = allowed
+ }
+ if allowed {
+ visibleMsgs = append(visibleMsgs, msg)
+ }
+ }
+ msgs = visibleMsgs
+
+ // 批量插入已读回执:靠 (message_id, user_id) 唯一索引 + OnConflict DoNothing
+ // 实现幂等,重复标记已读不报错也不更新(保留首次已读时间)
+ receipts := make([]model.MessageReadReceipt, 0, len(msgs))
+ readIDs := make([]uint, 0, len(msgs))
+ for _, msg := range msgs {
+ receipts = append(receipts, model.MessageReadReceipt{
+ MessageID: msg.ID,
UserID: uid,
ReadAt: now,
+ })
+ readIDs = append(readIDs, msg.ID)
+ }
+ if len(receipts) > 0 {
+ if err := service.ChatSvc.DB.Clauses(clause.OnConflict{DoNothing: true}).Create(&receipts).Error; err != nil {
+ utils.InternalError(c, "标记已读失败")
+ return
}
- service.ChatSvc.DB.Where("message_id = ? AND user_id = ?", mid, uid).
- Assign(receipt).FirstOrCreate(&receipt)
- readIDs = append(readIDs, mid)
+ }
- // 通知消息发送者(已读回执)
- var msg model.ChatMessage
- if err := service.ChatSvc.DB.First(&msg, mid).Error; err == nil && msg.SenderUserID != "" && msg.SenderUserID != uid {
- payload := model.WsPayload{
- RequestType: "messages_read",
- Data: map[string]interface{}{
- "message_ids": []uint{mid},
- "reader_id": uid,
- "room_id": msg.RoomID,
- },
- }
- if bytes, err := json.Marshal(payload); err == nil {
- service.ChatSvc.DispatchMessage(msg.SenderUserID, bytes)
- }
+ // 按"发送者+房间"分组推送已读回执:同一发送者在同一房间的多条消息合并成一条推送,
+ // 原实现每条消息推一次,批量已读时会向发送者刷屏式推送
+ type senderRoomKey struct {
+ sender string
+ room string
+ }
+ grouped := make(map[senderRoomKey][]uint)
+ for _, msg := range msgs {
+ if msg.SenderUserID == "" || msg.SenderUserID == uid {
+ continue
+ }
+ key := senderRoomKey{sender: msg.SenderUserID, room: msg.RoomID}
+ grouped[key] = append(grouped[key], msg.ID)
+ }
+ for key, ids := range grouped {
+ payload := model.WsPayload{
+ RequestType: "messages_read",
+ Data: map[string]interface{}{
+ "message_ids": ids,
+ "reader_id": uid,
+ "room_id": key.room,
+ },
+ }
+ if bytes, err := json.Marshal(payload); err == nil {
+ service.ChatSvc.DispatchMessage(key.sender, bytes)
}
}
diff --git a/internal/api/user_handler.go b/internal/api/user_handler.go
index 56d6f01..8e8a505 100644
--- a/internal/api/user_handler.go
+++ b/internal/api/user_handler.go
@@ -102,22 +102,9 @@ func GetUserListHandler(c *gin.Context) {
* 路径:POST /api/user/create
*/
func CreateUserHandler(c *gin.Context) {
- var user model.User
- if err := c.ShouldBindJSON(&user); err != nil {
- utils.BadRequest(c, "参数错误: "+err.Error())
- return
- }
-
- // 生成用户ID
- user.ID = generateUserID()
-
- if err := service.UserSvc.CreateUser(&user); err != nil {
- utils.BadRequest(c, "创建用户失败: "+err.Error())
- return
- }
-
- user.Password = "" // 清除密码
- utils.SuccessWithData(c, user, "创建成功")
+ // 该接口本质是管理员功能,但当前系统没有角色体系,任何登录用户都能创建任意账号(普通注册应走 /api/register)。
+ // 在缺少管理员鉴权的前提下,直接禁用以关闭越权创建用户的风险;待引入角色系统后再按 admin 放开。
+ utils.Forbidden(c, "无权限:创建用户需要管理员权限")
}
/**
@@ -126,8 +113,16 @@ func CreateUserHandler(c *gin.Context) {
* 路径:POST /api/user/update
*/
func UpdateUserHandler(c *gin.Context) {
+ // 强制只能修改当前登录用户自己的资料:忽略请求体中的 id,防止越权修改他人
+ uid, ok := c.Get("user_id")
+ if !ok {
+ utils.Unauthorized(c, "未认证")
+ return
+ }
+ currentUserID := uid.(string)
+
var req struct {
- ID string `json:"id" binding:"required"`
+ ID string `json:"id"`
Updates map[string]interface{} `json:"updates" binding:"required"`
}
@@ -136,17 +131,34 @@ func UpdateUserHandler(c *gin.Context) {
return
}
- if err := service.UserSvc.UpdateUser(req.ID, req.Updates); err != nil {
+ // 字段白名单:仅允许更新以下资料字段。
+ // 过滤掉 password/id/email 等敏感或身份字段,防止通过此接口改密码或伪造身份。
+ // phone 保留:PC 端"修改手机号"功能走本接口更新自己的手机号(越权风险已由
+ // currentUserID 强制"只能改自己"消除;users.phone 有唯一索引,占用冲突由数据库兜底报错)
+ // moment_cover:朋友圈顶部封面图,用户可自行更换
+ allowed := map[string]bool{"name": true, "avatar": true, "desc": true, "region": true, "phone": true, "moment_cover": true}
+ updates := make(map[string]interface{})
+ for k, v := range req.Updates {
+ if allowed[k] {
+ updates[k] = v
+ }
+ }
+ if len(updates) == 0 {
+ utils.BadRequest(c, "没有可更新的字段")
+ return
+ }
+
+ if err := service.UserSvc.UpdateUser(currentUserID, updates); err != nil {
utils.BadRequest(c, "更新失败: "+err.Error())
return
}
// 如果更新了头像或名称,通知好友
if service.ContactSvc != nil && service.ChatSvc != nil {
- if _, hasName := req.Updates["name"]; hasName {
- go notifyFriendsOfProfileUpdate(req.ID, req.Updates)
- } else if _, hasAvatar := req.Updates["avatar"]; hasAvatar {
- go notifyFriendsOfProfileUpdate(req.ID, req.Updates)
+ if _, hasName := updates["name"]; hasName {
+ go notifyFriendsOfProfileUpdate(currentUserID, updates)
+ } else if _, hasAvatar := updates["avatar"]; hasAvatar {
+ go notifyFriendsOfProfileUpdate(currentUserID, updates)
}
}
@@ -211,8 +223,16 @@ func notifyFriendsOfProfileUpdate(userID string, updates map[string]interface{})
* 路径:POST /api/user/delete
*/
func DeleteUserHandler(c *gin.Context) {
+ // 仅允许注销当前登录用户自己,禁止删除他人(原实现可传任意 id 越权删除)
+ uid, ok := c.Get("user_id")
+ if !ok {
+ utils.Unauthorized(c, "未认证")
+ return
+ }
+ currentUserID := uid.(string)
+
var req struct {
- ID string `json:"id" binding:"required"`
+ ID string `json:"id"`
}
if err := c.ShouldBindJSON(&req); err != nil {
@@ -220,7 +240,12 @@ func DeleteUserHandler(c *gin.Context) {
return
}
- if err := service.UserSvc.DeleteUser(req.ID); err != nil {
+ if req.ID != "" && req.ID != currentUserID {
+ utils.Forbidden(c, "无权删除其他用户")
+ return
+ }
+
+ if err := service.UserSvc.DeleteUser(currentUserID); err != nil {
utils.BadRequest(c, "删除失败: "+err.Error())
return
}
diff --git a/internal/manager/client_manager.go b/internal/manager/client_manager.go
index 77a0916..2bcda50 100644
--- a/internal/manager/client_manager.go
+++ b/internal/manager/client_manager.go
@@ -23,6 +23,7 @@ type Client struct {
RemoteIP string // WebSocket 连接来源 IP
Conn *websocket.Conn // 底层 WebSocket 连接对象
SendQueue chan []byte // 发送缓冲队列 (防止网络阻塞导致写协程卡死)
+ closeOnce sync.Once // 保证 SendQueue/Conn 只被关闭一次,防止并发注销时 double-close panic
}
/**
@@ -34,6 +35,10 @@ type ClientManager struct {
clients sync.Map
// 用户反向索引: map[string][]string (Key: UserID, Value: [ClientID1, ClientID2])
userClients sync.Map
+ // userMu 保护 userClients 中切片的读-改-写复合操作。
+ // 为什么需要:sync.Map 只保证单次 Load/Store 原子,Bind/Remove 是"读出切片→修改→写回",
+ // 并发时会互相覆盖丢失设备,必须用互斥锁串行化。
+ userMu sync.Mutex
}
// Manager 全局单例
@@ -57,19 +62,36 @@ func (m *ClientManager) Register(c *Client) {
* @param c *Client 客户端对象
*/
func (m *ClientManager) Unregister(c *Client) {
- if _, ok := m.clients.Load(c.ID); ok {
- m.clients.Delete(c.ID)
- close(c.SendQueue) // 关闭通道,退出 writePump
- c.Conn.Close()
+ // LoadAndDelete 原子操作:并发调用时只有一个 goroutine 拿到 ok=true,
+ // 再配合 closeOnce 双保险,彻底避免重复 close(chan) 导致的 panic
+ if _, ok := m.clients.LoadAndDelete(c.ID); ok {
+ c.closeOnce.Do(func() {
+ close(c.SendQueue) // 关闭通道,退出 writePump
+ c.Conn.Close()
+ })
- // 如果已绑定用户,清理反向索引
+ // 读取 c.UserID 与清理反向索引都放在同一把锁内,
+ // 与 BindUser 对 c.UserID 的写入互斥,避免数据竞争
+ m.userMu.Lock()
if c.UserID != "" {
- m.removeUserClient(c.UserID, c.ID)
+ m.removeUserClientLocked(c.UserID, c.ID)
}
+ m.userMu.Unlock()
log.Printf("🔌 [ClientManager] 客户端注销: %s", c.ID)
}
}
+/**
+ * GetClientUserID
+ * 功能:在锁保护下读取某连接当前绑定的 UserID。
+ * 为什么:c.UserID 会被 BindUser 并发写入,裸读构成数据竞争,需与写入同锁。
+ */
+func (m *ClientManager) GetClientUserID(c *Client) string {
+ m.userMu.Lock()
+ defer m.userMu.Unlock()
+ return c.UserID
+}
+
/**
* BindUser
* 功能:将 ClientID 与 UserID 进行绑定,支持多端登录。
@@ -82,10 +104,24 @@ func (m *ClientManager) BindUser(clientID, userID string) {
return
}
client := clientInterface.(*Client)
- client.UserID = userID
- // 更新反向索引 (UserID -> []ClientID)
- // 使用 LoadOrStore 初始化切片
+ // 加锁保护:既保护 userClients 的"读出切片→追加→写回"复合操作,
+ // 也保护 client.UserID 的写入(Unregister/GetClientUserID 会并发读取该字段)
+ m.userMu.Lock()
+ defer m.userMu.Unlock()
+
+ // 仅在实际变更时写入 client.UserID:
+ // 1) 重绑到不同用户时,先把 clientID 从旧用户的反向索引里摘掉,
+ // 否则旧用户的 SendToUser 仍会把消息发到这条连接,注销时也只按新 UserID 清理、
+ // 在旧用户列表里永久残留悬空 clientID;
+ // 2) 重复 self-bind(值相同)时跳过写入,避免冗余写与业务路径的读构成竞争
+ if client.UserID != userID {
+ if client.UserID != "" {
+ m.removeUserClientLocked(client.UserID, clientID)
+ }
+ client.UserID = userID
+ }
+
actual, _ := m.userClients.LoadOrStore(userID, make([]string, 0))
ids := actual.([]string)
@@ -96,23 +132,28 @@ func (m *ClientManager) BindUser(clientID, userID string) {
}
}
- ids = append(ids, clientID)
- m.userClients.Store(userID, ids)
+ // 拷贝后追加,避免直接修改共享切片底层数组引起数据竞争
+ newIds := make([]string, len(ids), len(ids)+1)
+ copy(newIds, ids)
+ newIds = append(newIds, clientID)
+ m.userClients.Store(userID, newIds)
log.Printf("🔗 [ClientManager] 绑定成功: %s -> %s", clientID, userID)
}
/**
- * removeUserClient
- * 功能:从用户的客户端列表中移除指定的 ClientID。
+ * removeUserClientLocked
+ * 功能:从用户的客户端列表中移除指定 ClientID(无锁内核,调用方须已持有 m.userMu)。
+ * 为什么是"Locked"内核:BindUser/Unregister 已在锁内需要清理旧索引,
+ * 若再走一个自己加锁的版本会重复加锁导致死锁,故统一用这个不加锁的内核。
*/
-func (m *ClientManager) removeUserClient(userID, clientID string) {
+func (m *ClientManager) removeUserClientLocked(userID, clientID string) {
val, ok := m.userClients.Load(userID)
if !ok {
return
}
ids := val.([]string)
- newIds := make([]string, 0)
+ newIds := make([]string, 0, len(ids))
for _, id := range ids {
if id != clientID {
newIds = append(newIds, id)
@@ -126,6 +167,18 @@ func (m *ClientManager) removeUserClient(userID, clientID string) {
}
}
+/**
+ * UserHasClients
+ * 功能:判断某用户在本节点是否仍有在线连接。
+ * 场景:客户端下线时,用于决定是否需要清理 Redis 中该用户的节点路由 key。
+ */
+func (m *ClientManager) UserHasClients(userID string) bool {
+ if val, ok := m.userClients.Load(userID); ok {
+ return len(val.([]string)) > 0
+ }
+ return false
+}
+
/**
* SendToClient
* 功能:向指定客户端发送消息 (非阻塞模式)。
diff --git a/internal/mediaserver/ffmpeg_rtmp_receiver.go b/internal/mediaserver/ffmpeg_rtmp_receiver.go
index a8cddea..c7dadb0 100644
--- a/internal/mediaserver/ffmpeg_rtmp_receiver.go
+++ b/internal/mediaserver/ffmpeg_rtmp_receiver.go
@@ -41,14 +41,14 @@ type FFmpegRTMPReceiver struct {
// FFmpegRTMPReceiverPool 接收器池,管理多个 FFmpeg RTMP 接收器
type FFmpegRTMPReceiverPool struct {
- receivers map[string]*FFmpegRTMPReceiver
- mu sync.RWMutex
- rtmpServer *RTMPServer
- ffmpegPath string
- basePort int
- nextPort int
- maxPort int
- usedPorts map[int]bool
+ receivers map[string]*FFmpegRTMPReceiver
+ mu sync.RWMutex
+ rtmpServer *RTMPServer
+ ffmpegPath string
+ basePort int
+ nextPort int
+ maxPort int
+ usedPorts map[int]bool
}
// NewFFmpegRTMPReceiverPool 创建接收器池
@@ -165,14 +165,14 @@ func (r *FFmpegRTMPReceiver) Start() error {
// 构建 FFmpeg 命令
// ffmpeg -listen 1 -i rtmp://0.0.0.0:{port}/live/{stream_id} -c copy -f flv pipe:1
args := []string{
- "-listen", "1", // 作为服务器监听
- "-timeout", "30000000", // 超时时间 30 秒(微秒)
+ "-listen", "1", // 作为服务器监听
+ "-timeout", "30000000", // 超时时间 30 秒(微秒)
"-i", fmt.Sprintf("rtmp://0.0.0.0:%d/live/%s", r.port, r.streamID), // 监听地址
- "-c", "copy", // 直接复制,不重新编码
- "-fflags", "nobuffer", // 禁用输入缓冲
- "-flags", "low_delay", // 低延迟模式
- "-f", "flv", // 输出格式
- "pipe:1", // 输出到 stdout
+ "-c", "copy", // 直接复制,不重新编码
+ "-fflags", "nobuffer", // 禁用输入缓冲
+ "-flags", "low_delay", // 低延迟模式
+ "-f", "flv", // 输出格式
+ "pipe:1", // 输出到 stdout
}
r.cmd = exec.Command(r.ffmpegPath, args...)
@@ -215,6 +215,10 @@ func (r *FFmpegRTMPReceiver) Start() error {
}
// Stop 停止接收器
+// 与 FFmpegTranscoder.Stop 相同的并发安全约束:
+// 1. 不调用 cmd.Wait()——monitor 协程是唯一的 Wait 回收方,重复 Wait 行为未定义;
+// 2. 不 close(outputChan)——readOutput(唯一发送方)可能仍在发送,
+// 向已关闭通道发送会 panic 崩溃进程;通道由 readOutput 退出时统一关闭。
func (r *FFmpegRTMPReceiver) Stop() error {
r.mu.Lock()
defer r.mu.Unlock()
@@ -226,15 +230,12 @@ func (r *FFmpegRTMPReceiver) Stop() error {
r.running = false
close(r.stopChan)
- // 关闭进程
+ // 杀掉进程即可:stdout 管道随之 EOF,readOutput 退出并关闭 outputChan;
+ // 僵尸进程由 monitor 协程的 cmd.Wait() 统一回收
if r.cmd != nil && r.cmd.Process != nil {
r.cmd.Process.Kill()
- r.cmd.Wait()
}
- // 关闭输出通道
- close(r.outputChan)
-
duration := time.Since(r.startTime)
log.Printf("🛑 [FFmpegRTMPReceiver] 已停止 | Stream:%s Port:%d Duration:%v Output:%d bytes",
r.streamID, r.port, duration, r.outputBytes)
@@ -255,7 +256,10 @@ func (r *FFmpegRTMPReceiver) Output() <-chan []byte {
}
// readOutput 读取 FFmpeg FLV 输出
+// 本协程是 outputChan 的唯一发送方,由它在退出时关闭通道:
+// 既保证 range Output() 的消费者能正常结束,又杜绝"向已关闭通道发送"的 panic
func (r *FFmpegRTMPReceiver) readOutput() {
+ defer close(r.outputChan)
log.Printf("▶️ [FFmpegRTMPReceiver] 开始读取 FLV 输出 | Stream:%s", r.streamID)
buf := make([]byte, 64*1024) // 64KB 缓冲区
@@ -378,4 +382,3 @@ func (r *FFmpegRTMPReceiver) monitor() {
}
}
}
-
diff --git a/internal/mediaserver/ffmpeg_transcoder.go b/internal/mediaserver/ffmpeg_transcoder.go
index 6fdbdb6..ff08d89 100644
--- a/internal/mediaserver/ffmpeg_transcoder.go
+++ b/internal/mediaserver/ffmpeg_transcoder.go
@@ -21,12 +21,12 @@ import (
type TranscoderConfig struct {
// 模式配置
CopyMode bool // 是否使用 copy 模式(H.264 输入时使用,只重封装不重新编码)
-
+
// 视频配置
VideoCodec string // 输出视频编解码器,默认 libx264
VideoPreset string // x264 预设,默认 ultrafast
VideoBitrate string // 视频码率,默认 500k
-
+
// 音频配置
AudioCodec string // 输出音频编解码器,默认 aac
AudioBitrate string // 音频码率,默认 64k
@@ -51,18 +51,18 @@ func DefaultTranscoderConfig() *TranscoderConfig {
type FFmpegTranscoder struct {
config *TranscoderConfig
ffmpegPath string
-
+
cmd *exec.Cmd
stdin io.WriteCloser
stdout io.ReadCloser
stderr io.ReadCloser
-
+
outputChan chan []byte
stopChan chan struct{}
-
- mu sync.Mutex
- running bool
-
+
+ mu sync.Mutex
+ running bool
+
// 统计信息
inputBytes int64
outputBytes int64
@@ -74,13 +74,13 @@ func NewFFmpegTranscoder(config *TranscoderConfig) (*FFmpegTranscoder, error) {
if config == nil {
config = DefaultTranscoderConfig()
}
-
+
// 确保 FFmpeg 可用
ffmpegPath, err := EnsureFFmpeg()
if err != nil {
return nil, fmt.Errorf("FFmpeg 不可用: %w", err)
}
-
+
return &FFmpegTranscoder{
config: config,
ffmpegPath: ffmpegPath,
@@ -93,35 +93,35 @@ func NewFFmpegTranscoder(config *TranscoderConfig) (*FFmpegTranscoder, error) {
func (t *FFmpegTranscoder) Start() error {
t.mu.Lock()
defer t.mu.Unlock()
-
+
if t.running {
return fmt.Errorf("转码器已在运行")
}
-
+
// 构建 FFmpeg 命令
var args []string
-
+
if t.config.CopyMode {
// H.264 copy 模式:只重封装,不重新编码(低延迟)
// 添加低延迟参数:禁用缓冲、减少探测时间
args = []string{
// 低延迟输入参数
- "-fflags", "nobuffer", // 禁用输入缓冲
- "-flags", "low_delay", // 低延迟模式
- "-probesize", "32", // 减少探测大小(字节)
- "-analyzeduration", "0", // 禁用分析时长
- "-f", "webm", // 输入格式
- "-i", "pipe:0", // 从 stdin 读取
+ "-fflags", "nobuffer", // 禁用输入缓冲
+ "-flags", "low_delay", // 低延迟模式
+ "-probesize", "32", // 减少探测大小(字节)
+ "-analyzeduration", "0", // 禁用分析时长
+ "-f", "webm", // 输入格式
+ "-i", "pipe:0", // 从 stdin 读取
// 输出参数
- "-c:v", "copy", // 视频直接复制,不重新编码
+ "-c:v", "copy", // 视频直接复制,不重新编码
"-c:a", t.config.AudioCodec,
"-b:a", t.config.AudioBitrate,
"-ar", fmt.Sprintf("%d", t.config.AudioSampleRate),
"-ac", fmt.Sprintf("%d", t.config.AudioChannels),
// 低延迟输出参数
- "-fflags", "+genpts", // 生成时间戳
- "-f", "flv", // 输出格式
- "pipe:1", // 输出到 stdout
+ "-fflags", "+genpts", // 生成时间戳
+ "-f", "flv", // 输出格式
+ "pipe:1", // 输出到 stdout
}
log.Printf("🎬 [FFmpegTranscoder] 使用 copy 模式(H.264 重封装,低延迟)")
} else {
@@ -129,106 +129,109 @@ func (t *FFmpegTranscoder) Start() error {
// 添加低延迟参数
args = []string{
// 低延迟输入参数
- "-fflags", "nobuffer", // 禁用输入缓冲
- "-flags", "low_delay", // 低延迟模式
- "-probesize", "32", // 减少探测大小(字节)
- "-analyzeduration", "0", // 禁用分析时长
- "-f", "webm", // 输入格式
- "-i", "pipe:0", // 从 stdin 读取
+ "-fflags", "nobuffer", // 禁用输入缓冲
+ "-flags", "low_delay", // 低延迟模式
+ "-probesize", "32", // 减少探测大小(字节)
+ "-analyzeduration", "0", // 禁用分析时长
+ "-f", "webm", // 输入格式
+ "-i", "pipe:0", // 从 stdin 读取
// 视频编码参数
"-c:v", t.config.VideoCodec,
"-preset", t.config.VideoPreset,
- "-tune", "zerolatency", // 零延迟调优
+ "-tune", "zerolatency", // 零延迟调优
"-b:v", t.config.VideoBitrate,
- "-g", "30", // GOP 大小(减少关键帧间隔)
- "-keyint_min", "15", // 最小关键帧间隔
+ "-g", "30", // GOP 大小(减少关键帧间隔)
+ "-keyint_min", "15", // 最小关键帧间隔
// 音频编码参数
"-c:a", t.config.AudioCodec,
"-b:a", t.config.AudioBitrate,
"-ar", fmt.Sprintf("%d", t.config.AudioSampleRate),
"-ac", fmt.Sprintf("%d", t.config.AudioChannels),
// 低延迟输出参数
- "-fflags", "+genpts", // 生成时间戳
- "-f", "flv", // 输出格式
- "pipe:1", // 输出到 stdout
+ "-fflags", "+genpts", // 生成时间戳
+ "-f", "flv", // 输出格式
+ "pipe:1", // 输出到 stdout
}
log.Printf("🎬 [FFmpegTranscoder] 使用转码模式(VP8/VP9 → H.264,低延迟)")
}
-
+
t.cmd = exec.Command(t.ffmpegPath, args...)
-
+
var err error
-
+
// 获取 stdin
t.stdin, err = t.cmd.StdinPipe()
if err != nil {
return fmt.Errorf("获取 stdin 失败: %w", err)
}
-
+
// 获取 stdout
t.stdout, err = t.cmd.StdoutPipe()
if err != nil {
return fmt.Errorf("获取 stdout 失败: %w", err)
}
-
+
// 获取 stderr
t.stderr, err = t.cmd.StderrPipe()
if err != nil {
return fmt.Errorf("获取 stderr 失败: %w", err)
}
-
+
// 启动 FFmpeg 进程
if err := t.cmd.Start(); err != nil {
return fmt.Errorf("启动 FFmpeg 失败: %w", err)
}
-
+
t.running = true
t.startTime = time.Now()
-
+
log.Printf("🎬 [FFmpegTranscoder] 转码器已启动, PID=%d", t.cmd.Process.Pid)
-
+
// 启动输出读取协程
go t.readOutput()
-
+
// 启动 stderr 读取协程(用于日志)
go t.readStderr()
-
+
// 启动进程监控协程
go t.monitor()
-
+
return nil
}
// Stop 停止转码器
+// 注意两个并发安全约束:
+// 1. 不在此处调用 cmd.Wait()——monitor 协程是唯一的 Wait 回收方,
+// os/exec 不允许对同一 Cmd 重复 Wait(第二次行为未定义);
+// 2. 不在此处 close(outputChan)——readOutput 协程(唯一发送方)可能仍在
+// 向通道发送数据,向已关闭通道发送会 panic 并直接崩溃整个进程;
+// 通道统一由 readOutput 退出时关闭("由发送方关闭"的 Go 惯例)。
func (t *FFmpegTranscoder) Stop() error {
t.mu.Lock()
defer t.mu.Unlock()
-
+
if !t.running {
return nil
}
-
+
t.running = false
close(t.stopChan)
-
+
// 关闭 stdin,通知 FFmpeg 输入结束
if t.stdin != nil {
t.stdin.Close()
}
-
- // 等待进程结束
+
+ // 杀掉进程即可:stdout 管道随之 EOF,readOutput 退出并关闭 outputChan;
+ // 僵尸进程由 monitor 协程的 cmd.Wait() 统一回收
if t.cmd != nil && t.cmd.Process != nil {
t.cmd.Process.Kill()
- t.cmd.Wait()
}
-
- // 关闭输出通道
- close(t.outputChan)
-
+
duration := time.Since(t.startTime)
log.Printf("🛑 [FFmpegTranscoder] 转码器已停止, 运行时间=%v, 输入=%d bytes, 输出=%d bytes",
duration, t.inputBytes, t.outputBytes)
-
+
return nil
}
@@ -241,12 +244,12 @@ func (t *FFmpegTranscoder) Write(data []byte) (int, error) {
}
stdin := t.stdin
t.mu.Unlock()
-
+
n, err := stdin.Write(data)
if err != nil {
return n, err
}
-
+
t.inputBytes += int64(n)
return n, nil
}
@@ -257,16 +260,19 @@ func (t *FFmpegTranscoder) Output() <-chan []byte {
}
// readOutput 读取 FFmpeg 输出
+// 本协程是 outputChan 的唯一发送方,因此由它在退出时关闭通道:
+// 既保证 range Output() 的消费者能正常结束,又杜绝"向已关闭通道发送"的 panic
func (t *FFmpegTranscoder) readOutput() {
+ defer close(t.outputChan)
buf := make([]byte, 64*1024) // 64KB 缓冲区
-
+
for {
select {
case <-t.stopChan:
return
default:
}
-
+
n, err := t.stdout.Read(buf)
if err != nil {
if err != io.EOF {
@@ -274,14 +280,14 @@ func (t *FFmpegTranscoder) readOutput() {
}
return
}
-
+
if n > 0 {
t.outputBytes += int64(n)
-
+
// 复制数据并发送到通道
data := make([]byte, n)
copy(data, buf[:n])
-
+
select {
case t.outputChan <- data:
default:
@@ -296,7 +302,7 @@ func (t *FFmpegTranscoder) readOutput() {
func (t *FFmpegTranscoder) readStderr() {
buf := new(bytes.Buffer)
io.Copy(buf, t.stderr)
-
+
if buf.Len() > 0 {
// 只在有错误时输出
output := buf.String()
@@ -310,12 +316,12 @@ func (t *FFmpegTranscoder) readStderr() {
// monitor 监控 FFmpeg 进程
func (t *FFmpegTranscoder) monitor() {
err := t.cmd.Wait()
-
+
t.mu.Lock()
wasRunning := t.running
t.running = false
t.mu.Unlock()
-
+
if wasRunning {
if err != nil {
log.Printf("⚠️ [FFmpegTranscoder] FFmpeg 进程异常退出: %v", err)
@@ -361,21 +367,21 @@ func NewTranscoderPool(config *TranscoderConfig) *TranscoderPool {
func (p *TranscoderPool) Get(streamID string) (*FFmpegTranscoder, error) {
p.mu.Lock()
defer p.mu.Unlock()
-
+
if t, exists := p.transcoders[streamID]; exists && t.IsRunning() {
return t, nil
}
-
+
// 创建新的转码器
t, err := NewFFmpegTranscoder(p.config)
if err != nil {
return nil, err
}
-
+
if err := t.Start(); err != nil {
return nil, err
}
-
+
p.transcoders[streamID] = t
return t, nil
}
@@ -384,7 +390,7 @@ func (p *TranscoderPool) Get(streamID string) (*FFmpegTranscoder, error) {
func (p *TranscoderPool) Release(streamID string) {
p.mu.Lock()
defer p.mu.Unlock()
-
+
if t, exists := p.transcoders[streamID]; exists {
t.Stop()
delete(p.transcoders, streamID)
@@ -395,7 +401,7 @@ func (p *TranscoderPool) Release(streamID string) {
func (p *TranscoderPool) ReleaseAll() {
p.mu.Lock()
defer p.mu.Unlock()
-
+
for id, t := range p.transcoders {
t.Stop()
delete(p.transcoders, id)
@@ -408,5 +414,3 @@ func (p *TranscoderPool) Count() int {
defer p.mu.RUnlock()
return len(p.transcoders)
}
-
-
diff --git a/internal/mediaserver/room.go b/internal/mediaserver/room.go
index d3fc532..1ced2ec 100644
--- a/internal/mediaserver/room.go
+++ b/internal/mediaserver/room.go
@@ -41,6 +41,27 @@ type Participant struct {
StreamID string `json:"stream_id,omitempty"` // 流ID
}
+// peerConnectionCloser 抽象 PeerConnection 的关闭能力。
+// 为什么这样写:Participant.PeerConnection 为避免循环依赖声明成 interface{},
+// room.go 不便直接引用 *webrtc.PeerConnection,这里用最小接口做类型断言即可调用 Close。
+type peerConnectionCloser interface {
+ Close() error
+}
+
+// closeParticipantPC 关闭参与者持有的 PeerConnection 并置空。
+// 必须在移除参与者/关闭房间时调用,否则底层 ICE/DTLS/UDP 端口等资源会永久泄漏。
+func closeParticipantPC(roomID string, p *Participant) {
+ if p == nil || p.PeerConnection == nil {
+ return
+ }
+ if closer, ok := p.PeerConnection.(peerConnectionCloser); ok {
+ if err := closer.Close(); err != nil {
+ log.Printf("⚠️ [Room:%s] 关闭用户 %s 的 PeerConnection 失败: %v", roomID, p.UserID, err)
+ }
+ }
+ p.PeerConnection = nil
+}
+
// Room 通话房间
type Room struct {
ID string `json:"id"`
@@ -109,6 +130,9 @@ func (r *Room) AddParticipant(userID string, pType ParticipantType) (*Participan
if p, exists := r.Participants[userID]; exists {
// 用户重新加入,清除旧的流信息以便重新生成
oldStreamID := p.StreamID
+ // 先关闭旧 PeerConnection:重连/重新协商会创建新连接,
+ // 旧连接若不 Close,其 ICE/DTLS/端口资源将永久泄漏(每次重进泄漏一条)
+ closeParticipantPC(r.ID, p)
p.PushURL = ""
p.PullURL = ""
p.FLVURL = ""
@@ -142,10 +166,8 @@ func (r *Room) RemoveParticipant(userID string) {
defer r.mu.Unlock()
if p, exists := r.Participants[userID]; exists {
- // 清理 PeerConnection
- if p.PeerConnection != nil {
- // TODO: 关闭 PeerConnection
- }
+ // 清理 PeerConnection:离房必须关闭,释放 ICE/DTLS/端口资源(原TODO一直没关,存在资源泄漏)
+ closeParticipantPC(r.ID, p)
delete(r.Participants, userID)
log.Printf("🎥 [Room:%s] 用户 %s 离开", r.ID, userID)
}
@@ -277,11 +299,9 @@ func (r *Room) Close() {
r.closed = true
close(r.closeChan)
- // 清理所有参与者
+ // 清理所有参与者(逐个关闭 PeerConnection,防止关房后连接资源泄漏)
for userID, p := range r.Participants {
- if p.PeerConnection != nil {
- // TODO: 关闭 PeerConnection
- }
+ closeParticipantPC(r.ID, p)
delete(r.Participants, userID)
}
diff --git a/internal/mediaserver/rtmp.go b/internal/mediaserver/rtmp.go
index 9e743ec..79652bd 100644
--- a/internal/mediaserver/rtmp.go
+++ b/internal/mediaserver/rtmp.go
@@ -20,6 +20,7 @@ import (
"log"
"net"
"net/http"
+ "strconv"
"strings"
"sync"
"time"
@@ -277,7 +278,7 @@ func (h *ConnectionHandler) handleSetChunkSize(msg *RTMPMessage) error {
}
newSize := binary.BigEndian.Uint32(msg.Data)
-
+
// 关键验证:chunkSize 必须在合理范围内 (1 到 16MB)
// 参考 RTMP 规范:最大值通常不超过 16MB,常见值为 128-65536
// 如果值异常,可能是协议解析错误,忽略此消息
@@ -285,7 +286,7 @@ func (h *ConnectionHandler) handleSetChunkSize(msg *RTMPMessage) error {
log.Printf("⚠️ [RTMP] SetChunkSize 值异常,忽略: %d (remote: %s)", newSize, h.remoteAddr)
return nil
}
-
+
log.Printf("📝 [RTMP] SetChunkSize: %d -> %d (remote: %s)", h.reader.GetChunkSize(), newSize, h.remoteAddr)
// 关键:立即更新 chunk size
@@ -1265,15 +1266,25 @@ func (r *RTMPServer) cleanupExpiredStreams() {
}
}
-// generateStreamToken 生成流鉴权 Token
-func generateStreamToken(streamID string) string {
+// generateStreamTokenWithExp 按指定过期时间生成流鉴权 Token
+// Token 结构:.
+// 为什么这样写:原实现把"当前时间+24h"放进 HMAC 却不随 Token 下发,
+// 校验时重新取当前时间重算,两次时间戳必然不同,导致 Token 永远校验失败(等于没有鉴权)。
+// 现在把过期时间戳作为 Token 的一部分下发,校验时用同一时间戳重算签名,并检查是否过期。
+func generateStreamTokenWithExp(streamID string, exp int64) string {
secret := viper.GetString("turn.shared_secret")
- timestamp := time.Now().Add(24 * time.Hour).Unix()
- data := fmt.Sprintf("%s:%d", streamID, timestamp)
+ data := fmt.Sprintf("%s:%d", streamID, exp)
mac := hmac.New(sha1.New, []byte(secret))
mac.Write([]byte(data))
- return base64.URLEncoding.EncodeToString(mac.Sum(nil))
+ sig := base64.URLEncoding.EncodeToString(mac.Sum(nil))
+ return fmt.Sprintf("%d.%s", exp, sig)
+}
+
+// generateStreamToken 生成流鉴权 Token(默认24小时有效)
+func generateStreamToken(streamID string) string {
+ exp := time.Now().Add(24 * time.Hour).Unix()
+ return generateStreamTokenWithExp(streamID, exp)
}
// GenerateStreamToken 生成流鉴权 Token (公开接口)
@@ -1281,8 +1292,20 @@ func GenerateStreamToken(streamID string) string {
return generateStreamToken(streamID)
}
-// ValidateStreamToken 验证流鉴权 Token
+// ValidateStreamToken 验证流鉴权 Token:解析出过期时间戳,检查未过期且签名一致
func ValidateStreamToken(streamID, token string) bool {
- expectedToken := generateStreamToken(streamID)
- return token == expectedToken
+ parts := strings.SplitN(token, ".", 2)
+ if len(parts) != 2 {
+ return false
+ }
+ exp, err := strconv.ParseInt(parts[0], 10, 64)
+ if err != nil {
+ return false
+ }
+ if time.Now().Unix() > exp {
+ return false // Token 已过期
+ }
+ expected := generateStreamTokenWithExp(streamID, exp)
+ // 使用恒定时间比较,避免时序侧信道
+ return hmac.Equal([]byte(expected), []byte(token))
}
diff --git a/internal/mediaserver/rtmp_amf0.go b/internal/mediaserver/rtmp_amf0.go
index 729fdcd..3e84071 100644
--- a/internal/mediaserver/rtmp_amf0.go
+++ b/internal/mediaserver/rtmp_amf0.go
@@ -15,24 +15,24 @@ import (
// AMF0 数据类型标记
const (
- AMF0_NUMBER = 0x00 // 8 bytes double
- AMF0_BOOLEAN = 0x01 // 1 byte
- AMF0_STRING = 0x02 // 2 bytes length + data
- AMF0_OBJECT = 0x03 // key-value pairs
- AMF0_MOVIECLIP = 0x04 // reserved
- AMF0_NULL = 0x05 // no data
- AMF0_UNDEFINED = 0x06 // no data
- AMF0_REFERENCE = 0x07 // 2 bytes
- AMF0_ECMA_ARRAY = 0x08 // associative array
- AMF0_OBJECT_END = 0x09 // object end marker
- AMF0_STRICT_ARRAY = 0x0A // strict array
- AMF0_DATE = 0x0B // 8 bytes double + 2 bytes timezone
- AMF0_LONG_STRING = 0x0C // 4 bytes length + data
- AMF0_UNSUPPORTED = 0x0D
- AMF0_RECORDSET = 0x0E // reserved
- AMF0_XML_DOCUMENT = 0x0F
- AMF0_TYPED_OBJECT = 0x10
- AMF0_AVMPLUS = 0x11 // switch to AMF3
+ AMF0_NUMBER = 0x00 // 8 bytes double
+ AMF0_BOOLEAN = 0x01 // 1 byte
+ AMF0_STRING = 0x02 // 2 bytes length + data
+ AMF0_OBJECT = 0x03 // key-value pairs
+ AMF0_MOVIECLIP = 0x04 // reserved
+ AMF0_NULL = 0x05 // no data
+ AMF0_UNDEFINED = 0x06 // no data
+ AMF0_REFERENCE = 0x07 // 2 bytes
+ AMF0_ECMA_ARRAY = 0x08 // associative array
+ AMF0_OBJECT_END = 0x09 // object end marker
+ AMF0_STRICT_ARRAY = 0x0A // strict array
+ AMF0_DATE = 0x0B // 8 bytes double + 2 bytes timezone
+ AMF0_LONG_STRING = 0x0C // 4 bytes length + data
+ AMF0_UNSUPPORTED = 0x0D
+ AMF0_RECORDSET = 0x0E // reserved
+ AMF0_XML_DOCUMENT = 0x0F
+ AMF0_TYPED_OBJECT = 0x10
+ AMF0_AVMPLUS = 0x11 // switch to AMF3
)
// AMF0Object 表示一个 AMF0 对象
@@ -327,6 +327,14 @@ func (d *AMF0Decoder) decodeStrictArray() ([]interface{}, error) {
length := int(binary.BigEndian.Uint32(d.data[d.offset : d.offset+4]))
d.offset += 4
+ // 长度合法性校验:length 来自报文声明,恶意/异常 RTMP 报文可声明超大长度
+ // (最大 2^32-1),若直接 make 预分配会导致巨额内存分配甚至 OOM。
+ // AMF0 每个元素至少占 1 字节(类型标记),因此声明长度不可能超过剩余字节数,
+ // 以此为上界拒绝非法报文。
+ if length < 0 || length > d.Remaining() {
+ return nil, fmt.Errorf("严格数组长度非法: %d (剩余 %d 字节)", length, d.Remaining())
+ }
+
arr := make([]interface{}, length)
for i := 0; i < length; i++ {
val, err := d.DecodeValue()
@@ -423,5 +431,3 @@ func EncodeOnBWDone() []byte {
encoder.EncodeNull()
return encoder.Bytes()
}
-
-
diff --git a/internal/mediaserver/sfu.go b/internal/mediaserver/sfu.go
index 30ea7ab..f334ba9 100644
--- a/internal/mediaserver/sfu.go
+++ b/internal/mediaserver/sfu.go
@@ -28,11 +28,11 @@ import (
// SFUServer WebRTC SFU 服务器
type SFUServer struct {
- config *MediaServerConfig
- api *webrtc.API
- mu sync.RWMutex
- running bool
-
+ config *MediaServerConfig
+ api *webrtc.API
+ mu sync.RWMutex
+ running bool
+
// Track 管理
trackLocals map[string]*webrtc.TrackLocalStaticRTP // trackID -> localTrack
}
@@ -54,36 +54,46 @@ func (s *SFUServer) Start() {
}
s.running = true
s.mu.Unlock()
-
+
+ // 初始化失败时必须复位 running:
+ // 原实现失败直接 return,running 仍为 true 但 api 为 nil,
+ // SFU 处于"假启动"状态且因为 running=true 无法再次 Start,只能重启进程
+ fail := func(format string, args ...interface{}) {
+ log.Printf(format, args...)
+ s.mu.Lock()
+ s.running = false
+ s.mu.Unlock()
+ }
+
// 创建 MediaEngine
m := &webrtc.MediaEngine{}
-
+
// 注册默认编解码器
if err := m.RegisterDefaultCodecs(); err != nil {
- log.Printf("❌ [SFU] 注册编解码器失败: %v", err)
+ fail("❌ [SFU] 注册编解码器失败: %v", err)
return
}
-
+
// 创建拦截器注册表
i := &interceptor.Registry{}
-
+
// 注册 PLI 拦截器(用于请求关键帧)
intervalPliFactory, err := intervalpli.NewReceiverInterceptor()
if err != nil {
- log.Printf("❌ [SFU] 创建 PLI 拦截器失败: %v", err)
+ fail("❌ [SFU] 创建 PLI 拦截器失败: %v", err)
return
}
i.Add(intervalPliFactory)
-
+
// 使用拦截器
if err := webrtc.RegisterDefaultInterceptors(m, i); err != nil {
- log.Printf("❌ [SFU] 注册拦截器失败: %v", err)
+ fail("❌ [SFU] 注册拦截器失败: %v", err)
return
}
-
+
// 创建 API
s.api = webrtc.NewAPI(webrtc.WithMediaEngine(m), webrtc.WithInterceptorRegistry(i))
-
+
log.Printf("🚀 [SFU] WebRTC SFU 已启动 | 端口: %d", s.config.WebRTCPort)
}
@@ -91,7 +101,7 @@ func (s *SFUServer) Start() {
func (s *SFUServer) Stop() {
s.mu.Lock()
defer s.mu.Unlock()
-
+
s.running = false
s.trackLocals = make(map[string]*webrtc.TrackLocalStaticRTP)
log.Println("🛑 [SFU] 已停止")
@@ -105,43 +115,45 @@ func (s *SFUServer) CreatePeerConnection(roomID, userID string) (*webrtc.PeerCon
return nil, ErrSFUNotReady
}
s.mu.RUnlock()
-
+
// 获取 ICE 服务器配置
iceServers := s.getICEServers(userID)
-
+
// 创建 PeerConnection 配置
config := webrtc.Configuration{
ICEServers: iceServers,
}
-
+
// 创建 PeerConnection
pc, err := s.api.NewPeerConnection(config)
if err != nil {
return nil, fmt.Errorf("create peer connection failed: %w", err)
}
-
+
// 监听 ICE 连接状态
pc.OnICEConnectionStateChange(func(state webrtc.ICEConnectionState) {
log.Printf("🔗 [SFU] Room:%s User:%s ICE状态: %s", roomID, userID, state.String())
-
- if state == webrtc.ICEConnectionStateFailed || state == webrtc.ICEConnectionStateDisconnected {
- // 连接断开,清理资源
+
+ // 仅在 Failed / Closed 时清理资源。
+ // Disconnected 是可自愈的瞬时状态(网络抖动、WiFi切换等,ICE 会自动恢复),
+ // 原实现把 Disconnected 也当成断线立即踢出房间,导致弱网用户被误踢
+ if state == webrtc.ICEConnectionStateFailed || state == webrtc.ICEConnectionStateClosed {
s.handleDisconnect(roomID, userID)
}
})
-
+
// 监听轨道
pc.OnTrack(func(remoteTrack *webrtc.TrackRemote, receiver *webrtc.RTPReceiver) {
s.handleTrack(roomID, userID, remoteTrack, receiver)
})
-
+
return pc, nil
}
// getICEServers 获取 ICE 服务器配置
func (s *SFUServer) getICEServers(userID string) []webrtc.ICEServer {
servers := []webrtc.ICEServer{}
-
+
// 添加 STUN 服务器
turnPublicIP := viper.GetString("turn.public_ip")
turnPort := viper.GetInt("turn.listen_port")
@@ -150,7 +162,7 @@ func (s *SFUServer) getICEServers(userID string) []webrtc.ICEServer {
servers = append(servers, webrtc.ICEServer{
URLs: []string{stunURL},
})
-
+
// 添加 TURN 服务器(如果启用)
if viper.GetBool("turn.enabled") {
turnURL := fmt.Sprintf("turn:%s:%d", turnPublicIP, turnPort)
@@ -163,17 +175,17 @@ func (s *SFUServer) getICEServers(userID string) []webrtc.ICEServer {
})
}
}
-
+
return servers
}
// handleTrack 处理远程轨道
func (s *SFUServer) handleTrack(roomID, userID string, remoteTrack *webrtc.TrackRemote, receiver *webrtc.RTPReceiver) {
trackID := fmt.Sprintf("%s_%s_%s", roomID, userID, remoteTrack.Kind().String())
-
+
log.Printf("🎬 [SFU] 收到轨道 | Room:%s User:%s Kind:%s ID:%s",
roomID, userID, remoteTrack.Kind().String(), trackID)
-
+
// 创建本地轨道用于转发
localTrack, err := webrtc.NewTrackLocalStaticRTP(
remoteTrack.Codec().RTPCodecCapability,
@@ -184,12 +196,12 @@ func (s *SFUServer) handleTrack(roomID, userID string, remoteTrack *webrtc.Track
log.Printf("❌ [SFU] 创建本地轨道失败: %v", err)
return
}
-
+
// 保存轨道
s.mu.Lock()
s.trackLocals[trackID] = localTrack
s.mu.Unlock()
-
+
// 转发 RTP 包
go func() {
buf := make([]byte, 1500)
@@ -199,20 +211,20 @@ func (s *SFUServer) handleTrack(roomID, userID string, remoteTrack *webrtc.Track
log.Printf("⚠️ [SFU] 读取轨道失败 | Track:%s Error:%v", trackID, readErr)
break
}
-
+
// 写入本地轨道(会自动转发给所有订阅者)
if _, writeErr := localTrack.Write(buf[:n]); writeErr != nil {
log.Printf("⚠️ [SFU] 写入轨道失败 | Track:%s Error:%v", trackID, writeErr)
break
}
}
-
+
// 清理轨道
s.mu.Lock()
delete(s.trackLocals, trackID)
s.mu.Unlock()
}()
-
+
// 通知房间内其他用户有新轨道
s.notifyNewTrack(roomID, userID, localTrack)
}
@@ -224,22 +236,22 @@ func (s *SFUServer) notifyNewTrack(roomID, senderUserID string, track *webrtc.Tr
if room == nil {
return
}
-
+
// 获取其他参与者
for _, p := range room.GetOtherParticipants(senderUserID) {
if p.Type != ParticipantTypeWebRTC {
continue // 只处理 WebRTC 用户
}
-
+
if p.PeerConnection == nil {
continue
}
-
+
pc, ok := p.PeerConnection.(*webrtc.PeerConnection)
if !ok {
continue
}
-
+
// 添加轨道到对方的 PeerConnection
if _, err := pc.AddTrack(track); err != nil {
log.Printf("⚠️ [SFU] 添加轨道到用户 %s 失败: %v", p.UserID, err)
@@ -254,9 +266,9 @@ func (s *SFUServer) handleDisconnect(roomID, userID string) {
if room == nil {
return
}
-
+
room.RemoveParticipant(userID)
-
+
// 清理该用户的所有轨道
s.mu.Lock()
for trackID := range s.trackLocals {
@@ -267,7 +279,7 @@ func (s *SFUServer) handleDisconnect(roomID, userID string) {
}
}
s.mu.Unlock()
-
+
// 如果房间空了,移除房间
if room.IsEmpty() {
ms.RemoveRoom(roomID)
@@ -285,46 +297,65 @@ func (s *SFUServer) GetTrackCount() int {
func (s *SFUServer) HandleOffer(roomID, userID string, offerSDP string) (string, error) {
ms := GetServer()
room := ms.GetOrCreateRoom(roomID)
-
- // 添加参与者
+
+ // 添加参与者(重复加入时 AddParticipant 内部会先 Close 掉旧的 PeerConnection,防止泄漏)
participant, err := room.AddParticipant(userID, ParticipantTypeWebRTC)
if err != nil {
return "", err
}
-
+
// 创建 PeerConnection
pc, err := s.CreatePeerConnection(roomID, userID)
if err != nil {
return "", err
}
participant.PeerConnection = pc
-
+
+ // 协商中途失败时关闭刚建的 PeerConnection,避免半初始化的连接泄漏
+ cleanup := func() {
+ _ = pc.Close()
+ participant.PeerConnection = nil
+ }
+
// 设置远程描述
offer := webrtc.SessionDescription{
Type: webrtc.SDPTypeOffer,
SDP: offerSDP,
}
if err := pc.SetRemoteDescription(offer); err != nil {
+ cleanup()
return "", fmt.Errorf("set remote description failed: %w", err)
}
-
+
// 添加房间内其他用户的轨道
s.addExistingTracks(pc, roomID, userID)
-
+
// 创建 Answer
answer, err := pc.CreateAnswer(nil)
if err != nil {
+ cleanup()
return "", fmt.Errorf("create answer failed: %w", err)
}
-
+
+ // 在 SetLocalDescription 之前创建"收集完成"信号(SetLocalDescription 才会触发 ICE 收集)
+ gatherComplete := webrtc.GatheringCompletePromise(pc)
+
// 设置本地描述
if err := pc.SetLocalDescription(answer); err != nil {
+ cleanup()
return "", fmt.Errorf("set local description failed: %w", err)
}
-
- // 等待 ICE 收集完成
- <-webrtc.GatheringCompletePromise(pc)
-
+
+ // 等待 ICE 收集完成,最多等 10 秒。
+ // 原实现裸 <-channel 无超时:网络异常时收集可能永远不结束,
+ // 处理该请求的 HTTP worker 会被永久阻塞,积累多了服务就没有可用协程了。
+ // 超时后用当前已收集到的候选返回,客户端仍可通过 trickle ICE 补充。
+ select {
+ case <-gatherComplete:
+ case <-time.After(10 * time.Second):
+ log.Printf("⚠️ [SFU] ICE 收集超时(10s),返回已收集候选 | Room:%s User:%s", roomID, userID)
+ }
+
return pc.LocalDescription().SDP, nil
}
@@ -332,10 +363,10 @@ func (s *SFUServer) HandleOffer(roomID, userID string, offerSDP string) (string,
func (s *SFUServer) addExistingTracks(pc *webrtc.PeerConnection, roomID, excludeUserID string) {
s.mu.RLock()
defer s.mu.RUnlock()
-
+
prefix := fmt.Sprintf("%s_", roomID)
excludePrefix := fmt.Sprintf("%s_%s_", roomID, excludeUserID)
-
+
for trackID, track := range s.trackLocals {
// 检查是否是同房间的轨道,且不是自己的
if strings.HasPrefix(trackID, prefix) && !strings.HasPrefix(trackID, excludePrefix) {
@@ -353,26 +384,26 @@ func (s *SFUServer) HandleICECandidate(roomID, userID string, candidateJSON stri
if room == nil {
return ErrRoomNotFound
}
-
+
participant := room.GetParticipant(userID)
if participant == nil {
return ErrParticipantNotFound
}
-
+
pc, ok := participant.PeerConnection.(*webrtc.PeerConnection)
if !ok || pc == nil {
return ErrSFUNotReady
}
-
+
var candidate webrtc.ICECandidateInit
if err := json.Unmarshal([]byte(candidateJSON), &candidate); err != nil {
return fmt.Errorf("parse ICE candidate failed: %w", err)
}
-
+
if err := pc.AddICECandidate(candidate); err != nil {
return fmt.Errorf("add ICE candidate failed: %w", err)
}
-
+
return nil
}
@@ -381,11 +412,11 @@ func generateTURNCredentials(userID string) (string, string) {
timestamp := time.Now().Add(24 * time.Hour).Unix()
username := fmt.Sprintf("%d:%s", timestamp, userID)
secret := viper.GetString("turn.shared_secret")
-
+
mac := hmac.New(sha1.New, []byte(secret))
mac.Write([]byte(username))
password := base64.StdEncoding.EncodeToString(mac.Sum(nil))
-
+
return username, password
}
diff --git a/internal/mediaserver/ws_rtmp_proxy.go b/internal/mediaserver/ws_rtmp_proxy.go
index 90cbed1..f416e19 100644
--- a/internal/mediaserver/ws_rtmp_proxy.go
+++ b/internal/mediaserver/ws_rtmp_proxy.go
@@ -67,7 +67,7 @@ type ProxySession struct {
// 编解码器信息
codecInfo *CodecInfo
codecDetected bool
-
+
// FFmpeg 转码器(VP8/VP9 需要转码)
transcoder *FFmpegTranscoder
needTranscode bool
@@ -86,17 +86,17 @@ type WebMHeader struct {
// WebMSegment WebM Segment
type WebMSegment struct {
- Info WebMSegmentInfo `ebml:"Info"`
- Tracks WebMTracks `ebml:"Tracks"`
- Cluster []WebMCluster `ebml:"Cluster"`
+ Info WebMSegmentInfo `ebml:"Info"`
+ Tracks WebMTracks `ebml:"Tracks"`
+ Cluster []WebMCluster `ebml:"Cluster"`
}
// WebMSegmentInfo Segment 信息
type WebMSegmentInfo struct {
- TimecodeScale uint64 `ebml:"TimecodeScale"`
+ TimecodeScale uint64 `ebml:"TimecodeScale"`
Duration float64 `ebml:"Duration,omitempty"`
- MuxingApp string `ebml:"MuxingApp,omitempty"`
- WritingApp string `ebml:"WritingApp,omitempty"`
+ MuxingApp string `ebml:"MuxingApp,omitempty"`
+ WritingApp string `ebml:"WritingApp,omitempty"`
}
// WebMTracks 轨道信息
@@ -106,9 +106,9 @@ type WebMTracks struct {
// WebMTrackEntry 轨道条目
type WebMTrackEntry struct {
- TrackNumber uint64 `ebml:"TrackNumber"`
- TrackType uint64 `ebml:"TrackType"` // 1=video, 2=audio
- CodecID string `ebml:"CodecID"`
+ TrackNumber uint64 `ebml:"TrackNumber"`
+ TrackType uint64 `ebml:"TrackType"` // 1=video, 2=audio
+ CodecID string `ebml:"CodecID"`
Video *WebMVideoTrack `ebml:"Video,omitempty"`
Audio *WebMAudioTrack `ebml:"Audio,omitempty"`
}
@@ -128,8 +128,8 @@ type WebMAudioTrack struct {
// WebMCluster WebM Cluster
type WebMCluster struct {
- Timecode uint64 `ebml:"Timecode"`
- SimpleBlock []ebml.Block `ebml:"SimpleBlock,omitempty"`
+ Timecode uint64 `ebml:"Timecode"`
+ SimpleBlock []ebml.Block `ebml:"SimpleBlock,omitempty"`
}
// NewWebMToRTMPProxy 创建代理服务
@@ -173,9 +173,11 @@ func (p *WebMToRTMPProxy) HandleWebSocket(w http.ResponseWriter, r *http.Request
return
}
- // 验证 token(简化版,生产环境需要更严格的验证)
- if token == "" {
- log.Printf("⚠️ [WSProxy] 缺少 token: stream=%s", streamID)
+ // 验证流 Token:签名+过期时间校验,防止任意人拿到 stream_id 就能推流
+ if !ValidateStreamToken(streamID, token) {
+ log.Printf("❌ [WSProxy] token 校验失败: stream=%s", streamID)
+ http.Error(w, "Invalid stream token", http.StatusUnauthorized)
+ return
}
// 升级到 WebSocket
@@ -205,16 +207,16 @@ func (p *WebMToRTMPProxy) HandleWebSocket(w http.ResponseWriter, r *http.Request
// 创建会话
session := &ProxySession{
- ID: fmt.Sprintf("%s_%d", userID, time.Now().UnixNano()),
- StreamID: streamID,
- UserID: userID,
- RoomID: roomID,
- Conn: conn,
- Stream: stream,
- StopChan: make(chan struct{}),
- StartTime: time.Now(),
- webmBuffer: bytes.NewBuffer(nil),
- headerParsed: false,
+ ID: fmt.Sprintf("%s_%d", userID, time.Now().UnixNano()),
+ StreamID: streamID,
+ UserID: userID,
+ RoomID: roomID,
+ Conn: conn,
+ Stream: stream,
+ StopChan: make(chan struct{}),
+ StartTime: time.Now(),
+ webmBuffer: bytes.NewBuffer(nil),
+ headerParsed: false,
}
// 注册会话
@@ -243,7 +245,7 @@ func (p *WebMToRTMPProxy) handleSession(session *ProxySession) {
if session.transcoder != nil {
session.transcoder.Stop()
}
-
+
// 清理
p.sessionsMu.Lock()
delete(p.sessions, session.ID)
@@ -296,7 +298,7 @@ func (p *WebMToRTMPProxy) handleSession(session *ProxySession) {
// 每 100 条消息记录一次统计
if messageCount%100 == 0 {
- log.Printf("📊 [WSProxy] 消息统计 | Stream:%s Count:%d TotalBytes:%d NeedTranscode:%v",
+ log.Printf("📊 [WSProxy] 消息统计 | Stream:%s Count:%d TotalBytes:%d NeedTranscode:%v",
session.StreamID, messageCount, totalBytes, session.needTranscode)
}
@@ -353,7 +355,7 @@ func (p *WebMToRTMPProxy) handleTextMessage(session *ProxySession, data []byte)
// H.264: 使用 FFmpeg copy 模式,只重封装为 FLV,不重新编码
session.needTranscode = true
log.Printf("✅ [WSProxy] H.264 编码,使用 FFmpeg copy 模式重封装为 FLV")
-
+
if err := p.initTranscoderForH264(session); err != nil {
log.Printf("⚠️ [WSProxy] 初始化 H.264 转码器失败: %v,尝试直接解析", err)
session.needTranscode = false
@@ -362,7 +364,7 @@ func (p *WebMToRTMPProxy) handleTextMessage(session *ProxySession, data []byte)
// VP8/VP9: 需要 FFmpeg 转码为 H.264
session.needTranscode = true
log.Printf("🔄 [WSProxy] %s 编码,需要 FFmpeg 转码为 H.264", codecInfo.VideoCodec)
-
+
if err := p.initTranscoder(session); err != nil {
log.Printf("⚠️ [WSProxy] 初始化转码器失败: %v,将使用非标准 FLV 格式", err)
session.needTranscode = false
@@ -399,7 +401,7 @@ func (p *WebMToRTMPProxy) initTranscoder(session *ProxySession) error {
// H.264 视频只需要重封装,不需要重新编码,使用 copy 模式可以大大减少延迟
func (p *WebMToRTMPProxy) initTranscoderForH264(session *ProxySession) error {
config := &TranscoderConfig{
- CopyMode: true, // 使用 copy 模式
+ CopyMode: true, // 使用 copy 模式
AudioCodec: "aac",
AudioBitrate: "64k",
AudioSampleRate: 44100,
@@ -437,7 +439,7 @@ func (p *WebMToRTMPProxy) readTranscodedOutput(session *ProxySession) {
// 解析 FLV 数据
for {
bufData := flvBuffer.Bytes()
-
+
// 首先解析 FLV 头(9 字节)+ PreviousTagSize(4 字节)
if !flvHeaderParsed {
if len(bufData) < 13 {
@@ -457,7 +459,7 @@ func (p *WebMToRTMPProxy) readTranscodedOutput(session *ProxySession) {
// Tag 头
tagType := bufData[0]
dataSize := int(bufData[1])<<16 | int(bufData[2])<<8 | int(bufData[3])
-
+
// 检查是否有完整的 tag(11 字节头 + dataSize + 4 字节 PreviousTagSize)
totalTagSize := 11 + dataSize + 4
if len(bufData) < totalTagSize {
@@ -525,7 +527,7 @@ func (p *WebMToRTMPProxy) parseAndConvert(session *ProxySession) error {
// parseWebMHeader 解析 WebM 头部
func (p *WebMToRTMPProxy) parseWebMHeader(session *ProxySession, data []byte) error {
reader := bytes.NewReader(data)
-
+
// 解析 EBML 头
var header struct {
EBML struct {
@@ -533,7 +535,7 @@ func (p *WebMToRTMPProxy) parseWebMHeader(session *ProxySession, data []byte) er
DocType string `ebml:"DocType"`
} `ebml:"EBML"`
}
-
+
if err := ebml.Unmarshal(reader, &header); err != nil {
return err
}
@@ -546,7 +548,7 @@ func (p *WebMToRTMPProxy) parseWebMHeader(session *ProxySession, data []byte) er
func (p *WebMToRTMPProxy) parseWebMClusters(session *ProxySession) error {
// 简化处理:直接将 WebM 数据转换为 FLV
// 实际实现需要完整解析 WebM 的 Cluster/SimpleBlock
-
+
bufData := session.webmBuffer.Bytes()
if len(bufData) < 100 {
return nil
@@ -554,7 +556,7 @@ func (p *WebMToRTMPProxy) parseWebMClusters(session *ProxySession) error {
// 寻找 Cluster 标记 (0x1F43B675)
clusterMarker := []byte{0x1F, 0x43, 0xB6, 0x75}
-
+
for {
idx := bytes.Index(bufData, clusterMarker)
if idx == -1 || idx+12 > len(bufData) {
@@ -576,7 +578,7 @@ func (p *WebMToRTMPProxy) parseWebMClusters(session *ProxySession) error {
// 提取 Cluster 数据
clusterData := bufData[idx:totalSize]
-
+
// 转换为 FLV 并广播
if err := p.convertClusterToFLV(session, clusterData); err != nil {
log.Printf("⚠️ [WSProxy] 转换 Cluster 失败: %v", err)
@@ -649,7 +651,7 @@ func (p *WebMToRTMPProxy) convertClusterToFLV(session *ProxySession, clusterData
// 解析 Cluster 时间戳
timecodeMarker := []byte{0xE7} // Timecode element ID
timecodeIdx := bytes.Index(clusterData, timecodeMarker)
-
+
var timestamp uint32 = session.lastTimestamp
if timecodeIdx != -1 && timecodeIdx+1 < len(clusterData) {
tcSize, bytesRead := readVarInt(clusterData[timecodeIdx+1:])
@@ -664,7 +666,7 @@ func (p *WebMToRTMPProxy) convertClusterToFLV(session *ProxySession, clusterData
// 解析 SimpleBlock
simpleBlockMarker := []byte{0xA3} // SimpleBlock element ID
blockData := clusterData
-
+
for {
blockIdx := bytes.Index(blockData, simpleBlockMarker)
if blockIdx == -1 || blockIdx+1 >= len(blockData) {
@@ -704,7 +706,7 @@ func (p *WebMToRTMPProxy) convertClusterToFLV(session *ProxySession, clusterData
continue
}
relativeTimestamp := binary.BigEndian.Uint16(block[trackBytes : trackBytes+2])
-
+
// 解析 flags (1 byte)
flagsIdx := trackBytes + 2
if flagsIdx >= len(block) {
@@ -712,7 +714,7 @@ func (p *WebMToRTMPProxy) convertClusterToFLV(session *ProxySession, clusterData
continue
}
flags := block[flagsIdx]
-
+
// 帧数据
frameData := block[flagsIdx+1:]
if len(frameData) == 0 {
@@ -752,24 +754,24 @@ func (p *WebMToRTMPProxy) createFLVVideoTag(timestamp uint32, data []byte, isKey
// FLV Video Tag Header:
// FrameType (4 bits): 1=keyframe, 2=inter frame
// CodecID (4 bits): 7=AVC (H.264), 12=VP8 (非标准)
-
+
frameType := byte(2) // inter frame
if isKeyframe {
frameType = 1 // keyframe
}
-
+
// 默认使用 VP8(非标准),如果是 H.264 则使用标准格式
// 注意:VP8 不能被标准 FLV 播放器播放,需要转码
// 这里保留 VP8 封装作为降级方案
codecID := byte(12) // 自定义:VP8
-
+
header := (frameType << 4) | codecID
-
+
// 构建完整数据
videoData := make([]byte, 1+len(data))
videoData[0] = header
copy(videoData[1:], data)
-
+
return p.createFLVTag(9, timestamp, videoData) // 9 = video
}
@@ -779,15 +781,15 @@ func (p *WebMToRTMPProxy) createFLVVideoTagH264(timestamp uint32, data []byte, i
// FLV Video Tag Header:
// FrameType (4 bits): 1=keyframe, 2=inter frame
// CodecID (4 bits): 7=AVC (H.264)
-
+
frameType := byte(2) // inter frame
if isKeyframe {
frameType = 1 // keyframe
}
-
+
codecID := byte(7) // AVC (H.264)
header := (frameType << 4) | codecID
-
+
// AVC 数据需要额外的封装
// AVCPacketType: 0=AVC sequence header, 1=AVC NALU
// CompositionTime: 3 bytes (通常为 0)
@@ -796,7 +798,7 @@ func (p *WebMToRTMPProxy) createFLVVideoTagH264(timestamp uint32, data []byte, i
// 关键帧可能需要先发送 sequence header
// 这里简化处理,假设数据已经是正确格式
}
-
+
// 构建 AVC 数据
// 1 byte header + 1 byte AVCPacketType + 3 bytes CompositionTime + data
videoData := make([]byte, 5+len(data))
@@ -806,7 +808,7 @@ func (p *WebMToRTMPProxy) createFLVVideoTagH264(timestamp uint32, data []byte, i
videoData[3] = 0
videoData[4] = 0
copy(videoData[5:], data)
-
+
return p.createFLVTag(9, timestamp, videoData) // 9 = video
}
@@ -815,26 +817,26 @@ func (p *WebMToRTMPProxy) createFLVVideoTagH264(timestamp uint32, data []byte, i
func (p *WebMToRTMPProxy) createFLVAudioTag(timestamp uint32, data []byte) []byte {
// 注意:Opus 不能直接封装到 FLV,需要转码为 AAC
// 这里使用简化方案
-
+
// FLV Audio Tag Header:
// SoundFormat (4 bits): 10=AAC, 13=Opus (非标准)
// SoundRate (2 bits): 3=44kHz
// SoundSize (1 bit): 1=16-bit
// SoundType (1 bit): 1=stereo
-
+
// 使用自定义格式标记 Opus
soundFormat := byte(13) // 自定义:Opus
soundRate := byte(3) // 44kHz
soundSize := byte(1) // 16-bit
soundType := byte(1) // stereo
-
+
header := (soundFormat << 4) | (soundRate << 2) | (soundSize << 1) | soundType
-
+
// 构建完整数据
audioData := make([]byte, 1+len(data))
audioData[0] = header
copy(audioData[1:], data)
-
+
return p.createFLVTag(8, timestamp, audioData) // 8 = audio
}
@@ -970,4 +972,3 @@ func (p *WebMToRTMPProxy) CloseAllSessions() {
// 确保导入被使用
var _ io.Reader = (*bytes.Reader)(nil)
-
diff --git a/internal/middleware/auth.go b/internal/middleware/auth.go
index b77fced..e6c08f4 100644
--- a/internal/middleware/auth.go
+++ b/internal/middleware/auth.go
@@ -12,22 +12,35 @@ import (
"github.com/gin-gonic/gin"
)
+/**
+ * AbortUnauthorized
+ * 功能:返回统一格式的401响应并终止请求。
+ * 说明:保持 HTTP 401 状态码(前端请求层按状态码识别认证失败并触发登出),
+ * 响应体经 utils.ResponseWithStatus 组装——与业务接口完全同构,
+ * 含 interface_info(result_time/ecs),不再手工构造 ApiResponse。
+ * 导出原因:/ws 握手鉴权(main.go)等非中间件场景也需要同样的统一 401 响应,复用避免格式漂移
+ */
+func AbortUnauthorized(c *gin.Context, message string) {
+ utils.ResponseWithStatus(c, http.StatusUnauthorized, utils.CodeUnauthorized, message, nil, utils.TypeError)
+ c.Abort()
+}
+
/**
* JWTAuthMiddleware
- *
+ *
* 功能:JWT认证中间件
- *
+ *
* 作用:
* 1. 从HTTP请求头或查询参数中提取JWT Token
* 2. 验证Token的有效性和过期时间
* 3. 从Token中解析出用户ID
* 4. 将用户ID注入到Gin Context中,供后续处理器使用
* 5. 如果Token无效或缺失,返回401未授权错误
- *
+ *
* 使用场景:
* - 需要用户登录才能访问的API接口
* - 需要在处理器中获取当前用户信息的接口
- *
+ *
* @returns gin.HandlerFunc 中间件处理函数
*/
func JWTAuthMiddleware() gin.HandlerFunc {
@@ -41,8 +54,7 @@ func JWTAuthMiddleware() gin.HandlerFunc {
// 步骤3: 如果仍然没有Token,返回401未授权错误
if token == "" {
- c.JSON(http.StatusUnauthorized, gin.H{"error": "缺少认证Token"})
- c.Abort() // 终止请求处理
+ AbortUnauthorized(c, "缺少认证Token")
return
}
@@ -56,8 +68,7 @@ func JWTAuthMiddleware() gin.HandlerFunc {
userID, err := utils.ValidateToken(token)
if err != nil {
// Token无效或已过期,返回401错误
- c.JSON(http.StatusUnauthorized, gin.H{"error": "Token无效或已过期: " + err.Error()})
- c.Abort()
+ AbortUnauthorized(c, "Token无效或已过期: "+err.Error())
return
}
@@ -73,18 +84,18 @@ func JWTAuthMiddleware() gin.HandlerFunc {
/**
* OptionalJWTAuthMiddleware
- *
+ *
* 功能:可选的JWT认证中间件(不强制要求认证)
- *
+ *
* 作用:
* 1. 如果请求中提供了Token,则验证Token并注入用户ID
* 2. 如果没有提供Token,则继续执行,不返回错误
* 3. 适用于既支持登录用户访问,也支持匿名用户访问的接口
- *
+ *
* 使用场景:
* - 公开接口,但登录用户可以获取更多信息
* - 兼容旧代码,不强制要求认证
- *
+ *
* @returns gin.HandlerFunc 中间件处理函数
*/
func OptionalJWTAuthMiddleware() gin.HandlerFunc {
@@ -113,4 +124,3 @@ func OptionalJWTAuthMiddleware() gin.HandlerFunc {
c.Next()
}
}
-
diff --git a/internal/middleware/request_log.go b/internal/middleware/request_log.go
index c08f082..8332230 100644
--- a/internal/middleware/request_log.go
+++ b/internal/middleware/request_log.go
@@ -10,6 +10,7 @@ import (
"encoding/json"
"fmt"
"io"
+ "regexp"
"strings"
"time"
"xk-websocket-v2/internal/model"
@@ -19,6 +20,47 @@ import (
"gorm.io/gorm"
)
+// sensitiveKeyPattern 匹配需要脱敏的字段名(大小写不敏感):密码、验证码、Token、密钥等
+var sensitiveKeyPattern = regexp.MustCompile(`(?i)(password|pwd|passwd|code|token|secret|authorization)`)
+
+// sanitizeParams 对请求体中的敏感字段做脱敏,避免明文密码/验证码/Token 落库。
+// 为什么这样写:请求日志会把原始 body 写入数据库,登录/注册/验证码接口的密码与验证码若明文入库风险很高,
+// 因此优先按 JSON 逐字段递归脱敏,非 JSON 时退化为整体截断。
+func sanitizeParams(body []byte) string {
+ if len(body) == 0 {
+ return ""
+ }
+ var m map[string]interface{}
+ if err := json.Unmarshal(body, &m); err == nil {
+ redactSensitiveMap(m)
+ if b, err := json.Marshal(m); err == nil {
+ return truncateString(string(b), 5000)
+ }
+ }
+ return truncateString(string(body), 5000)
+}
+
+// redactSensitiveMap 递归地把 map 中的敏感字段值替换为 ***
+func redactSensitiveMap(m map[string]interface{}) {
+ for k, v := range m {
+ if sensitiveKeyPattern.MatchString(k) {
+ m[k] = "***"
+ continue
+ }
+ if child, ok := v.(map[string]interface{}); ok {
+ redactSensitiveMap(child)
+ }
+ }
+}
+
+// truncateString 超长字符串截断,避免日志过大
+func truncateString(s string, n int) string {
+ if len(s) > n {
+ return s[:n] + "...(truncated)"
+ }
+ return s
+}
+
// 二进制 Content-Type 前缀列表
var binaryContentTypes = []string{
"image/",
@@ -40,10 +82,12 @@ var skipResponseBodyRoutes = []string{
"/static/",
}
-// 需要完全跳过中间件的路由(如 WebSocket)
+// 需要完全跳过中间件的路由(如 WebSocket、高频轮询)
var skipMiddlewareRoutes = []string{
"/ws",
"/api/call/ws-push",
+ // 扫码登录状态轮询:PC 登录页每 2 秒一次,入库毫无审计价值还会刷爆日志表
+ "/api/qrcode/status",
}
// isBinaryContentType 检测是否为二进制 Content-Type
@@ -74,25 +118,30 @@ func isBinaryData(data []byte) bool {
}
}
// 检查是否以常见的二进制文件头开始
+ // 注意:GIF/PDF/ZIP 的魔数以可打印 ASCII 开头("GIF8"/"%PDF"/"PK"),
+ // 若只比对前两字节会把 "GI..."/"%P..." 开头的普通文本误判为二进制而漏记日志,
+ // 因此必须校验完整魔数
if len(data) >= 2 {
- // JPEG: FF D8
+ // JPEG: FF D8(非 ASCII 前缀,两字节即可判定)
if data[0] == 0xFF && data[1] == 0xD8 {
return true
}
- // PNG: 89 50
+ // PNG: 89 50(首字节非 ASCII,两字节即可判定)
if data[0] == 0x89 && data[1] == 0x50 {
return true
}
- // GIF: 47 49
- if data[0] == 0x47 && data[1] == 0x49 {
+ }
+ if len(data) >= 4 {
+ // GIF: "GIF8"(GIF87a / GIF89a)
+ if data[0] == 'G' && data[1] == 'I' && data[2] == 'F' && data[3] == '8' {
return true
}
- // PDF: 25 50
- if data[0] == 0x25 && data[1] == 0x50 {
+ // PDF: "%PDF"
+ if data[0] == '%' && data[1] == 'P' && data[2] == 'D' && data[3] == 'F' {
return true
}
- // ZIP/DOCX/XLSX: 50 4B
- if data[0] == 0x50 && data[1] == 0x4B {
+ // ZIP/DOCX/XLSX: "PK" + 0x03/0x05/0x07(本地文件头/空档案尾/分卷标记)
+ if data[0] == 'P' && data[1] == 'K' && (data[2] == 0x03 || data[2] == 0x05 || data[2] == 0x07) {
return true
}
}
@@ -160,14 +209,6 @@ func (m *RequestLogMiddleware) Handler() gin.HandlerFunc {
c.Request.Body = io.NopCloser(bytes.NewBuffer(requestBody))
}
- // 获取用户ID(从Context中获取,未登录为"0")
- userID := "0"
- if uid, exists := c.Get("user_id"); exists {
- if uidStr, ok := uid.(string); ok {
- userID = uidStr
- }
- }
-
// 记录开始时间
startTime := time.Now()
@@ -184,13 +225,44 @@ func (m *RequestLogMiddleware) Handler() gin.HandlerFunc {
// 计算请求时间
duration := time.Since(startTime)
+ // 获取用户ID(从Context中获取,未登录为"0")
+ // 注意必须在 c.Next() 之后读取:JWT 认证中间件在 c.Next() 内部才执行 c.Set("user_id"),
+ // 若在 c.Next() 之前读取则永远是 "0"
+ userID := "0"
+ if uid, exists := c.Get("user_id"); exists {
+ if uidStr, ok := uid.(string); ok {
+ userID = uidStr
+ }
+ }
+
+ // Gin 通过 sync.Pool 复用 Context,本函数返回后 c 可能被重置并服务其他请求;
+ // 若在异步协程里继续读 c,会产生数据竞争,甚至把别的请求的 Method/Path/Header 写进日志。
+ // 因此进入协程前把所有需要的字段快照成值,协程内不再触碰 c。
+ snap := requestLogSnapshot{
+ route: c.FullPath(),
+ method: c.Request.Method,
+ path: c.Request.URL.Path,
+ requestContentType: c.GetHeader("Content-Type"),
+ responseContentType: writer.Header().Get("Content-Type"),
+ }
+
// 异步记录日志(避免影响性能)
- go m.logRequest(c, ip, userID, requestBody, writer.body.Bytes(), writer.status, duration)
+ go m.logRequest(snap, ip, userID, requestBody, writer.body.Bytes(), writer.status, duration)
}
}
-// logRequest 记录请求日志
-func (m *RequestLogMiddleware) logRequest(c *gin.Context, ip, userID string, requestBody, responseBody []byte, httpStatus int, duration time.Duration) {
+// requestLogSnapshot 请求上下文快照
+// 在请求协程内取值、传给异步日志协程使用,避免跨协程访问被复用的 *gin.Context
+type requestLogSnapshot struct {
+ route string // 注册的路由模板(如 /api/user/:id)
+ method string // HTTP 方法
+ path string // 实际请求路径
+ requestContentType string // 请求体 Content-Type
+ responseContentType string // 响应体 Content-Type
+}
+
+// logRequest 记录请求日志(运行在独立协程,只使用快照数据,不访问 gin.Context)
+func (m *RequestLogMiddleware) logRequest(snap requestLogSnapshot, ip, userID string, requestBody, responseBody []byte, httpStatus int, duration time.Duration) {
// 获取IP归属地
location := utils.GetIPLocation(ip)
@@ -206,30 +278,25 @@ func (m *RequestLogMiddleware) logRequest(c *gin.Context, ip, userID string, req
// 处理请求参数
var requestParams string
- requestContentType := c.GetHeader("Content-Type")
- if isBinaryContentType(requestContentType) || isBinaryData(requestBody) {
+ if isBinaryContentType(snap.requestContentType) || isBinaryData(requestBody) {
// 二进制请求体,只记录大小
requestParams = fmt.Sprintf("[binary data: %d bytes]", len(requestBody))
} else {
- requestParams = string(requestBody)
- if len(requestParams) > 5000 {
- requestParams = requestParams[:5000] + "...(truncated)"
- }
+ // 文本请求体:对密码/验证码/Token 等敏感字段脱敏后再记录
+ requestParams = sanitizeParams(requestBody)
}
// 处理响应参数
var responseParams string
- responseContentType := c.Writer.Header().Get("Content-Type")
- requestPath := c.Request.URL.Path
-
- if shouldSkipResponseBody(requestPath) {
+ if shouldSkipResponseBody(snap.path) {
// 静态文件路由,跳过响应体记录
responseParams = fmt.Sprintf("[static file: %d bytes]", len(responseBody))
- } else if isBinaryContentType(responseContentType) || isBinaryData(responseBody) {
+ } else if isBinaryContentType(snap.responseContentType) || isBinaryData(responseBody) {
// 二进制响应体,只记录大小
responseParams = fmt.Sprintf("[binary data: %d bytes]", len(responseBody))
} else {
- responseParams = string(responseBody)
+ // 文本响应体:与请求体一致做敏感字段脱敏(登录/注册响应中的 token 等),再截断长度
+ responseParams = sanitizeParams(responseBody)
if len(responseParams) > 5000 {
responseParams = responseParams[:5000] + "...(truncated)"
}
@@ -237,11 +304,11 @@ func (m *RequestLogMiddleware) logRequest(c *gin.Context, ip, userID string, req
// 创建日志记录
log := model.ApiRequestLog{
- Route: c.FullPath(),
+ Route: snap.route,
IP: ip,
IPLocation: location,
UserID: userID,
- Method: c.Request.Method,
+ Method: snap.method,
RequestParams: requestParams,
ResponseParams: responseParams,
ResponseCode: responseCode,
diff --git a/internal/model/types.go b/internal/model/types.go
index 7d369bd..ffd4df6 100644
--- a/internal/model/types.go
+++ b/internal/model/types.go
@@ -5,7 +5,10 @@
*/
package model
-import "time"
+import (
+ "encoding/json"
+ "time"
+)
// ==========================================
// 数据库实体 (PO - Persistent Object)
@@ -24,8 +27,10 @@ type ChatMessage struct {
// 发送者用户ID
SenderUserID string `gorm:"type:varchar(100);index;comment:发送者用户ID" json:"sender_user_id"`
// 发送者 IP
- SenderIP string `gorm:"type:varchar(50);comment:发送者IP" json:"sender_ip"`
- // IP 归属地
+ // json:"-":原始IP属于敏感隐私,仅入库供风控/审计使用,
+ // 禁止随消息历史/推送下发给聊天对端(前端也从未使用该字段)
+ SenderIP string `gorm:"type:varchar(50);comment:发送者IP" json:"-"`
+ // IP 归属地(粗粒度地域信息,前端消息气泡会展示"IP归属地",属于产品功能,保留下发)
IPLocation string `gorm:"type:varchar(255);comment:IP归属地" json:"ip_location"`
// 接收者用户ID
ReceiverUserID string `gorm:"type:varchar(100);index;comment:接收者用户ID" json:"receiver_user_id"`
@@ -43,6 +48,10 @@ type ChatMessage struct {
CallStatus string `gorm:"type:varchar(50);comment:通话状态" json:"call_status"`
// 创建时间,自动生成
CreatedAt time.Time `gorm:"autoCreateTime;comment:创建时间" json:"created_at"`
+ // 发送端 WebSocket 客户端ID(gorm:"-" 不入库,仅随推送下发)
+ // 为什么需要:消息会回推给发送者的所有在线设备以实现多端同步,
+ // 发送端自身可根据该字段识别"自己的回声"并忽略,避免重复渲染同一条消息
+ SenderClientID string `gorm:"-" json:"sender_client_id,omitempty"`
}
// TableName 指定表名和注释
@@ -50,6 +59,30 @@ func (ChatMessage) TableName() string {
return "chat_messages"
}
+/**
+ * MessageDeletion
+ * 作用:消息删除记录("删除仅对我生效"语义)。
+ * 逻辑:用户删除某条消息时插入一条 (message_id, user_id) 记录,
+ * 历史消息查询时排除当前用户已删除的消息;消息本体不动,
+ * 对方仍然可见(区别于"撤回"是双方都不可见)。
+ * 表结构见 migrations/20260811_message_deletions.sql(项目不使用 AutoMigrate)
+ */
+type MessageDeletion struct {
+ // 自增主键
+ ID uint `gorm:"primaryKey" json:"id"`
+ // 被删除的消息ID
+ MessageID uint `gorm:"index:uk_message_user,unique;comment:被删除的消息ID" json:"message_id"`
+ // 执行删除的用户ID(删除仅对该用户生效)
+ UserID string `gorm:"type:varchar(100);index:uk_message_user,unique;comment:执行删除的用户ID" json:"user_id"`
+ // 删除时间
+ CreatedAt time.Time `gorm:"autoCreateTime" json:"created_at"`
+}
+
+// TableName 指定表名
+func (MessageDeletion) TableName() string {
+ return "message_deletions"
+}
+
// ==========================================
// 交互数据传输对象 (DTO - Data Transfer Object)
// ==========================================
@@ -156,6 +189,10 @@ type User struct {
Desc string `gorm:"type:varchar(500);comment:用户描述或签名" json:"desc"`
// 地区
Region string `gorm:"type:varchar(100);comment:地区" json:"region"`
+ // 朋友圈封面图URL(空则前端展示默认封面)
+ MomentCover string `gorm:"type:varchar(500);comment:朋友圈封面图URL" json:"moment_cover"`
+ // 是否为AI机器人虚拟用户(机器人复用 users 表,消息/群成员/会话链路零改造)
+ IsBot bool `gorm:"default:false;comment:是否AI机器人" json:"is_bot"`
// 创建时间
CreatedAt time.Time `gorm:"autoCreateTime;comment:创建时间" json:"created_at"`
// 更新时间
@@ -167,6 +204,67 @@ func (User) TableName() string {
return "users"
}
+/**
+ * AIConfig
+ * 对应数据库表:ai_configs
+ * 作用:AI 提供商配置(全局一份,仅管理员 id=1 可维护)。
+ * 各机器人共用该配置调用大模型,工厂函数按 provider 创建对应 Provider 实例
+ */
+type AIConfig struct {
+ // 主键ID(全局仅一条记录,id 固定为 1)
+ ID uint `gorm:"primaryKey;comment:主键ID" json:"id"`
+ // 提供商标识:openai/deepseek/dashscope/moonshot/ollama/custom
+ Provider string `gorm:"type:varchar(50);comment:提供商标识" json:"provider"`
+ // API 基础地址(空则用 provider 的默认地址;custom 必填)
+ BaseURL string `gorm:"type:varchar(500);comment:API基础地址" json:"base_url"`
+ // API 密钥(返回给前端时须脱敏)
+ APIKey string `gorm:"type:varchar(500);comment:API密钥" json:"api_key"`
+ // 模型名称,如 gpt-4o-mini / deepseek-chat / qwen-plus
+ Model string `gorm:"type:varchar(100);comment:模型名称" json:"model"`
+ // 是否启用(关闭后所有机器人停止应答)
+ Enabled bool `gorm:"default:false;comment:是否启用" json:"enabled"`
+ // 创建时间
+ CreatedAt time.Time `gorm:"autoCreateTime;comment:创建时间" json:"created_at"`
+ // 更新时间
+ UpdatedAt time.Time `gorm:"autoUpdateTime;comment:更新时间" json:"updated_at"`
+}
+
+// TableName 指定表名
+func (AIConfig) TableName() string {
+ return "ai_configs"
+}
+
+/**
+ * AIBot
+ * 对应数据库表:ai_bots
+ * 作用:AI 机器人定义。每个机器人在 users 表有一个 is_bot=1 的虚拟用户,
+ * 使其能作为"联系人/群成员"复用全部现有消息链路;
+ * role_prompt 为角色设定,作为对话的 system 提示词
+ */
+type AIBot struct {
+ // 主键ID
+ ID uint `gorm:"primaryKey;comment:主键ID" json:"id"`
+ // 关联的虚拟用户ID(users.id,bot_ 前缀)
+ UserID string `gorm:"type:varchar(100);uniqueIndex;comment:虚拟用户ID" json:"user_id"`
+ // 机器人名称(群聊中 @名称 触发应答)
+ Name string `gorm:"type:varchar(100);comment:机器人名称" json:"name"`
+ // 机器人头像
+ Avatar string `gorm:"type:varchar(500);comment:机器人头像" json:"avatar"`
+ // 角色设定(system 提示词),如"你是一个温柔的客服"
+ RolePrompt string `gorm:"type:text;comment:角色设定" json:"role_prompt"`
+ // 是否启用(停用后不应答、不出现在好友列表)
+ Enabled bool `gorm:"default:true;comment:是否启用" json:"enabled"`
+ // 创建时间
+ CreatedAt time.Time `gorm:"autoCreateTime;comment:创建时间" json:"created_at"`
+ // 更新时间
+ UpdatedAt time.Time `gorm:"autoUpdateTime;comment:更新时间" json:"updated_at"`
+}
+
+// TableName 指定表名
+func (AIBot) TableName() string {
+ return "ai_bots"
+}
+
/**
* UserContact
* 对应数据库表:user_contacts
@@ -288,8 +386,10 @@ type RoomMember struct {
MutedUntil *time.Time `gorm:"type:datetime;comment:禁言到期时间" json:"muted_until,omitempty"`
// 用户信息(关联查询,不存储在数据库中)
User *User `gorm:"foreignKey:UserID;references:ID" json:"user,omitempty"`
- // 群名片(可选,用于群聊中显示的自定义昵称)
- Nickname string `gorm:"-" json:"nickname,omitempty"`
+ // 群名片(群内显示的自定义昵称)。
+ // 原先是 gorm:"-" 非持久化字段且后端从未赋值,前端"我在群里的昵称"功能整体失效,
+ // 现改为实体列(见 migrations/20260811_room_member_nickname.sql)
+ Nickname string `gorm:"type:varchar(100);default:'';comment:群名片(群内显示昵称)" json:"nickname,omitempty"`
}
// IsMuted 判断成员是否被禁言
@@ -511,8 +611,11 @@ type ClusterMessage struct {
SenderNodeID string `json:"sender_node_id"`
// 目标用户ID
TargetUserID string `json:"target_user_id"`
- // 原始消息体 (WsPayload)
- Payload interface{} `json:"payload"`
+ // 原始消息体 (WsPayload 的 JSON 字节)。
+ // 为什么用 json.RawMessage:原来的 interface{} 装 []byte 时会被 encoding/json 编码成 base64 字符串,
+ // 订阅端再 Marshal 一次得到的是带引号的 base64 文本,跨节点转发后客户端收到乱码;
+ // RawMessage 序列化时原样内嵌 JSON,反序列化时保留原始字节,两端零转换。
+ Payload json.RawMessage `json:"payload"`
}
// ==========================================
@@ -644,10 +747,11 @@ func (Moment) TableName() string {
type MomentLike struct {
// 主键ID
ID uint `gorm:"primaryKey;comment:主键ID" json:"id"`
- // 动态ID
- MomentID uint `gorm:"type:bigint;index;comment:动态ID" json:"moment_id"`
+ // 动态ID。与 UserID 组成唯一索引:并发点赞时靠数据库唯一约束兜底,
+ // 防止"先查后插"竞态导致同一用户对同一动态插入两条点赞记录
+ MomentID uint `gorm:"type:bigint;uniqueIndex:uk_moment_user;comment:动态ID" json:"moment_id"`
// 点赞用户ID
- UserID string `gorm:"type:varchar(100);index;comment:点赞用户ID" json:"user_id"`
+ UserID string `gorm:"type:varchar(100);uniqueIndex:uk_moment_user;comment:点赞用户ID" json:"user_id"`
// 创建时间
CreatedAt time.Time `gorm:"autoCreateTime;comment:创建时间" json:"created_at"`
// 关联字段(不存储在数据库中)
@@ -859,4 +963,8 @@ const (
MessageTypeGroupNotif = 7 // 群通知
MessageTypeFile = 8 // 文件
MessageTypeMoments = 9 // 朋友圈通知(预留)
+
+ // 会话类型常量(chat_conversations.type / chat_rooms.type)
+ ConversationTypeP2P = 1 // 私聊
+ ConversationTypeGroup = 2 // 群聊
)
diff --git a/internal/service/ai_bot_service.go b/internal/service/ai_bot_service.go
new file mode 100644
index 0000000..1823323
--- /dev/null
+++ b/internal/service/ai_bot_service.go
@@ -0,0 +1,529 @@
+/**
+ * package service
+ * 作用:AI 机器人业务服务。
+ *
+ * 职责:
+ * 1. AI 配置的读写(全局一条,管理员维护);
+ * 2. 机器人 CRUD(每个机器人在 users 表建 is_bot=1 的虚拟用户,
+ * 让机器人天然复用消息/群成员/会话等全部现有链路);
+ * 3. 消息应答:私聊发给机器人的消息、群聊 @机器人 的消息,
+ * 取最近聊天记录作为上下文调用大模型,以机器人身份回复。
+ *
+ * 能力边界(按需求约束):机器人只做聊天——上下文仅限"与它的私聊记录"
+ * 或"所在群的聊天记录",不赋予任何系统操作能力。
+ */
+package service
+
+import (
+ "context"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "log"
+ "strings"
+ "sync"
+ "time"
+
+ "xk-websocket-v2/internal/model"
+ "xk-websocket-v2/internal/utils"
+
+ "gorm.io/gorm"
+)
+
+// AIBotService AI 机器人服务
+type AIBotService struct {
+ DB *gorm.DB
+}
+
+// AIBotSvc 全局单例
+var AIBotSvc *AIBotService
+
+// 上下文取的历史消息条数:足够体现对话语境,又不至于撑爆 token
+const botContextMessageCount = 20
+
+// 管理员用户ID:按需求约定 id=1 的用户为系统管理员
+const AdminUserID = "1"
+
+// ---- 应答限流:AI 调用是计费的外部请求,必须限流防刷 ----
+// 全局并发上限:同时最多 N 个 goroutine 在等大模型返回(单次最长60s),
+// 防止恶意刷消息导致 goroutine 堆积与账单放大
+var botReplySem = make(chan struct{}, 3)
+
+// 每房间冷却:同一会话冷却期内只应答一次,连续刷屏的消息直接丢弃
+var (
+ botReplyMu sync.Mutex
+ botLastReplyAt = make(map[string]time.Time)
+)
+
+const botReplyCooldown = 3 * time.Second
+
+/**
+ * acquireBotReplySlot
+ * 作用:应答前的限流闸门——先过冷却(key=房间+机器人,互不干扰),再抢全局并发额度。
+ * 必须在"确认要触发机器人"之后调用,否则普通消息也会写冷却记录,
+ * 误伤"先发普通消息、紧接着 @机器人"的正常用法。
+ * 返回 false 表示该消息被限流丢弃(不回复提示,避免限流提示本身刷屏)
+ */
+func acquireBotReplySlot(cooldownKey string) bool {
+ botReplyMu.Lock()
+ if last, ok := botLastReplyAt[cooldownKey]; ok && time.Since(last) < botReplyCooldown {
+ botReplyMu.Unlock()
+ return false
+ }
+ botLastReplyAt[cooldownKey] = time.Now()
+ botReplyMu.Unlock()
+
+ select {
+ case botReplySem <- struct{}{}:
+ return true
+ default:
+ // 并发额度已满:直接丢弃而非排队,防止请求积压。
+ // 必须回滚上面写入的冷却记录——本次并未真正应答,
+ // 不回滚的话消息被丢弃、房间却进入 3 秒冷却,后续消息也被误挡
+ botReplyMu.Lock()
+ delete(botLastReplyAt, cooldownKey)
+ botReplyMu.Unlock()
+ log.Printf("⚠️ [AIBot] 应答并发已满,丢弃 key=%s", cooldownKey)
+ return false
+ }
+}
+
+// releaseBotReplySlot 释放全局并发额度
+func releaseBotReplySlot() {
+ <-botReplySem
+}
+
+/**
+ * InitAIBotService
+ * 作用:初始化 AI 机器人服务(main.go 启动时调用)
+ */
+func InitAIBotService(db *gorm.DB) {
+ AIBotSvc = &AIBotService{DB: db}
+}
+
+// ==================== AI 配置 ====================
+
+/**
+ * GetAIConfig
+ * 作用:读取全局 AI 配置(不存在时返回空配置而非报错,便于前端首次渲染表单)
+ */
+func (s *AIBotService) GetAIConfig() (*model.AIConfig, error) {
+ var cfg model.AIConfig
+ err := s.DB.First(&cfg, 1).Error
+ if errors.Is(err, gorm.ErrRecordNotFound) {
+ return &model.AIConfig{ID: 1}, nil
+ }
+ if err != nil {
+ return nil, err
+ }
+ return &cfg, nil
+}
+
+/**
+ * SaveAIConfig
+ * 作用:保存全局 AI 配置(固定 id=1 的单条记录,存在则更新)。
+ * apiKey 为空串表示"不修改密钥"(前端回显的是脱敏值,原样提交时不能覆盖真实 key)
+ */
+func (s *AIBotService) SaveAIConfig(provider, baseURL, apiKey, modelName string, enabled bool) (*model.AIConfig, error) {
+ cfg, err := s.GetAIConfig()
+ if err != nil {
+ return nil, err
+ }
+ cfg.ID = 1
+ cfg.Provider = provider
+ cfg.BaseURL = baseURL
+ cfg.Model = modelName
+ cfg.Enabled = enabled
+ if apiKey != "" {
+ cfg.APIKey = apiKey
+ }
+ if err := s.DB.Save(cfg).Error; err != nil {
+ return nil, err
+ }
+ return cfg, nil
+}
+
+// ==================== 机器人 CRUD ====================
+
+/**
+ * ListBots
+ * 作用:机器人列表。onlyEnabled=true 时仅返回启用中的(好友列表/应答判定用),
+ * false 返回全部(管理页用)
+ */
+func (s *AIBotService) ListBots(onlyEnabled bool) ([]model.AIBot, error) {
+ var bots []model.AIBot
+ query := s.DB.Order("id ASC")
+ if onlyEnabled {
+ query = query.Where("enabled = ?", true)
+ }
+ if err := query.Find(&bots).Error; err != nil {
+ return nil, err
+ }
+ return bots, nil
+}
+
+/**
+ * CreateBot
+ * 作用:创建机器人——事务内同时创建 users 虚拟用户(is_bot=1)与 ai_bots 记录。
+ * 虚拟用户让机器人能作为联系人展示、被拉进群、收发消息,全链路零特判
+ */
+func (s *AIBotService) CreateBot(name, avatar, rolePrompt string) (*model.AIBot, error) {
+ if strings.TrimSpace(name) == "" {
+ return nil, errors.New("机器人名称不能为空")
+ }
+ snowID, err := utils.NextID()
+ if err != nil {
+ return nil, errors.New("生成机器人ID失败")
+ }
+ botUserID := fmt.Sprintf("bot_%d", snowID)
+ if avatar == "" {
+ avatar = "🤖"
+ }
+
+ bot := &model.AIBot{
+ UserID: botUserID,
+ Name: name,
+ Avatar: avatar,
+ RolePrompt: rolePrompt,
+ Enabled: true,
+ }
+ err = s.DB.Transaction(func(tx *gorm.DB) error {
+ // 虚拟用户:email/phone 置为 bot 专属占位(这两列是唯一索引,不能为空串重复)。
+ // phone 列是 varchar(20),"bot_"+雪花共23字符会超长,只存19位纯雪花数字
+ //(11位真实手机号不可能与19位数字撞车)
+ user := &model.User{
+ ID: botUserID,
+ Email: botUserID + "@bot.local",
+ Phone: fmt.Sprintf("%d", snowID),
+ Name: name,
+ Avatar: avatar,
+ Desc: "AI机器人",
+ IsBot: true,
+ }
+ if err := tx.Create(user).Error; err != nil {
+ return err
+ }
+ return tx.Create(bot).Error
+ })
+ if err != nil {
+ return nil, err
+ }
+ return bot, nil
+}
+
+/**
+ * UpdateBot
+ * 作用:更新机器人信息,并同步 users 虚拟用户的名称/头像
+ * (群成员列表、消息气泡展示的是 users 表数据,不同步会出现两处名字不一致)
+ */
+func (s *AIBotService) UpdateBot(id uint, name, avatar, rolePrompt string, enabled bool) (*model.AIBot, error) {
+ var bot model.AIBot
+ if err := s.DB.First(&bot, id).Error; err != nil {
+ return nil, errors.New("机器人不存在")
+ }
+ if strings.TrimSpace(name) == "" {
+ return nil, errors.New("机器人名称不能为空")
+ }
+ bot.Name = name
+ if avatar != "" {
+ bot.Avatar = avatar
+ }
+ bot.RolePrompt = rolePrompt
+ bot.Enabled = enabled
+
+ err := s.DB.Transaction(func(tx *gorm.DB) error {
+ if err := tx.Save(&bot).Error; err != nil {
+ return err
+ }
+ return tx.Model(&model.User{}).Where("id = ?", bot.UserID).
+ Updates(map[string]interface{}{"name": bot.Name, "avatar": bot.Avatar}).Error
+ })
+ if err != nil {
+ return nil, err
+ }
+ return &bot, nil
+}
+
+/**
+ * DeleteBot
+ * 作用:删除机器人——事务内删除 ai_bots 记录、users 虚拟用户、群成员关系。
+ * 历史聊天记录保留(消息里的 sender 头像/名字查不到时前端按未知用户兜底展示)
+ */
+func (s *AIBotService) DeleteBot(id uint) error {
+ var bot model.AIBot
+ if err := s.DB.First(&bot, id).Error; err != nil {
+ return errors.New("机器人不存在")
+ }
+ return s.DB.Transaction(func(tx *gorm.DB) error {
+ if err := tx.Delete(&model.AIBot{}, id).Error; err != nil {
+ return err
+ }
+ if err := tx.Where("id = ?", bot.UserID).Delete(&model.User{}).Error; err != nil {
+ return err
+ }
+ // 移出所有群,避免群成员列表出现"幽灵成员"
+ return tx.Where("user_id = ?", bot.UserID).Delete(&model.RoomMember{}).Error
+ })
+}
+
+/**
+ * GetBotByUserID
+ * 作用:按虚拟用户ID查启用中的机器人(应答判定用)
+ */
+func (s *AIBotService) GetBotByUserID(userID string) (*model.AIBot, error) {
+ var bot model.AIBot
+ err := s.DB.Where("user_id = ? AND enabled = ?", userID, true).First(&bot).Error
+ if err != nil {
+ return nil, err
+ }
+ return &bot, nil
+}
+
+// ==================== 消息应答 ====================
+
+/**
+ * MaybeReply
+ * 作用:消息落库后的机器人应答钩子(chat_service 持久化成功后调用)。
+ * 判定规则:
+ * - 私聊:接收方是启用中的机器人 → 应答;
+ * - 群聊:消息文本中 @了群内某个启用中的机器人 → 应答。
+ * 整个过程在独立 goroutine 执行,不阻塞消息主链路;
+ * 仅处理文本消息,防止机器人对图片/语音等无法理解的内容乱答
+ */
+func (s *AIBotService) MaybeReply(msg *model.ChatMessage, isGroup bool) {
+ if msg == nil || msg.MessageType != model.MessageTypeText {
+ return
+ }
+ // 机器人发的消息不再触发应答,防止两个机器人互相@形成死循环
+ if strings.HasPrefix(msg.SenderUserID, "bot_") {
+ return
+ }
+
+ go func() {
+ defer func() {
+ // AI 调用链路涉及外部服务,任何 panic 不能拖垮进程
+ if r := recover(); r != nil {
+ log.Printf("❌ [AIBot] 应答处理 panic: %v", r)
+ }
+ }()
+
+ if isGroup {
+ s.replyInGroup(msg)
+ } else {
+ s.replyInP2P(msg)
+ }
+ }()
+}
+
+/**
+ * replyInP2P
+ * 作用:私聊应答——接收方是机器人时,取双方最近对话作上下文回复
+ */
+func (s *AIBotService) replyInP2P(msg *model.ChatMessage) {
+ bot, err := s.GetBotByUserID(msg.ReceiverUserID)
+ if err != nil {
+ return // 接收方不是启用中的机器人
+ }
+ // 限流:确认命中机器人后才进闸门(计费 API 防刷),被限流的消息静默丢弃
+ if !acquireBotReplySlot(msg.RoomID + "|" + bot.UserID) {
+ return
+ }
+ defer releaseBotReplySlot()
+
+ reply, err := s.callAI(bot, msg.RoomID, false)
+ if err != nil {
+ log.Printf("❌ [AIBot] 私聊应答失败 bot=%s: %v", bot.Name, err)
+ reply = "抱歉,我暂时无法回复(" + err.Error() + ")"
+ }
+ s.sendBotMessage(bot, msg.RoomID, msg.SenderUserID, reply, false)
+}
+
+/**
+ * replyInGroup
+ * 作用:群聊应答——文本中 @了群内启用中的机器人时触发。
+ * @判定:消息含 "@机器人名"(前端 @ 选择器与手输均适用)
+ */
+func (s *AIBotService) replyInGroup(msg *model.ChatMessage) {
+ if !strings.Contains(msg.Content, "@") {
+ return
+ }
+ bots, err := s.ListBots(true)
+ if err != nil || len(bots) == 0 {
+ return
+ }
+ for i := range bots {
+ bot := &bots[i]
+ if !strings.Contains(msg.Content, "@"+bot.Name) {
+ continue
+ }
+ // 机器人必须真的在这个群里(防止跨群喊话)
+ if RoomSvc == nil || !RoomSvc.IsRoomMember(msg.RoomID, bot.UserID) {
+ continue
+ }
+ s.replyToGroupBot(bot, msg)
+ }
+}
+
+/**
+ * replyToGroupBot
+ * 作用:群聊中单个机器人的应答(独立函数便于 defer 释放并发额度)。
+ * 冷却 key 含机器人ID:一条消息同时 @多个机器人时互不挤占冷却
+ */
+func (s *AIBotService) replyToGroupBot(bot *model.AIBot, msg *model.ChatMessage) {
+ if !acquireBotReplySlot(msg.RoomID + "|" + bot.UserID) {
+ return
+ }
+ defer releaseBotReplySlot()
+
+ reply, aiErr := s.callAI(bot, msg.RoomID, true)
+ if aiErr != nil {
+ log.Printf("❌ [AIBot] 群聊应答失败 bot=%s: %v", bot.Name, aiErr)
+ reply = "抱歉,我暂时无法回复(" + aiErr.Error() + ")"
+ }
+ // 群聊里带上 @提问人 前缀,明确这条回复对应谁的问题
+ senderName := msg.SenderUserID
+ if UserSvc != nil {
+ if u, uErr := UserSvc.GetUserByID(msg.SenderUserID); uErr == nil && u.Name != "" {
+ senderName = u.Name
+ }
+ }
+ s.sendBotMessage(bot, msg.RoomID, "", "@"+senderName+" "+reply, true)
+}
+
+/**
+ * callAI
+ * 作用:构建上下文并调用大模型。
+ * 上下文 = 角色设定(system) + 最近 N 条聊天记录(仅文本消息)。
+ * system 里显式限定能力边界:只基于聊天内容作答,不执行任何操作
+ */
+func (s *AIBotService) callAI(bot *model.AIBot, roomID string, isGroup bool) (string, error) {
+ cfg, err := s.GetAIConfig()
+ if err != nil {
+ return "", errors.New("AI配置读取失败")
+ }
+ if !cfg.Enabled {
+ return "", errors.New("管理员未启用AI服务")
+ }
+ provider, err := NewAIProvider(cfg)
+ if err != nil {
+ return "", err
+ }
+
+ // 取该房间最近 N 条文本消息(含刚落库的触发消息),按时间正序组装
+ var history []model.ChatMessage
+ if err := s.DB.Where("room_id = ? AND message_type = ?", roomID, model.MessageTypeText).
+ Order("id DESC").Limit(botContextMessageCount).Find(&history).Error; err != nil {
+ return "", errors.New("聊天记录读取失败")
+ }
+
+ // system 提示词:角色设定 + 能力边界约束
+ scene := "私聊"
+ if isGroup {
+ scene = "群聊"
+ }
+ systemPrompt := bot.RolePrompt
+ if strings.TrimSpace(systemPrompt) == "" {
+ systemPrompt = "你是一个友好的聊天助手。"
+ }
+ systemPrompt += fmt.Sprintf(
+ "\n\n约束:你叫\"%s\",正在参与一个%s。你只能基于下面的聊天内容进行回答,"+
+ "不能执行任何系统操作、不能声称自己能查询或修改任何数据。回答保持简洁自然,使用中文。",
+ bot.Name, scene)
+
+ messages := make([]AIMessage, 0, len(history)+1)
+ messages = append(messages, AIMessage{Role: "system", Content: systemPrompt})
+
+ // 历史倒序查出,反转为正序;批量查发送者名字用于群聊语境标注
+ nameCache := map[string]string{}
+ for i := len(history) - 1; i >= 0; i-- {
+ m := history[i]
+ if m.SenderUserID == bot.UserID {
+ messages = append(messages, AIMessage{Role: "assistant", Content: m.Content})
+ continue
+ }
+ content := m.Content
+ if isGroup {
+ // 群聊标注说话人,模型才能区分"谁在说什么"
+ name, ok := nameCache[m.SenderUserID]
+ if !ok {
+ name = m.SenderUserID
+ if UserSvc != nil {
+ if u, uErr := UserSvc.GetUserByID(m.SenderUserID); uErr == nil && u.Name != "" {
+ name = u.Name
+ }
+ }
+ nameCache[m.SenderUserID] = name
+ }
+ content = "[" + name + "]: " + content
+ }
+ messages = append(messages, AIMessage{Role: "user", Content: content})
+ }
+
+ ctx, cancel := context.WithTimeout(context.Background(), 60*time.Second)
+ defer cancel()
+ return provider.Chat(ctx, messages)
+}
+
+/**
+ * sendBotMessage
+ * 作用:以机器人身份发消息——落库、更新会话列表、WS 推送,
+ * 流程与 HandleUserMessage 的常规消息一致(复用同一套分发),
+ * 保证接收端体验与真人消息无差别
+ */
+func (s *AIBotService) sendBotMessage(bot *model.AIBot, roomID, receiverUserID, content string, isGroup bool) {
+ if ChatSvc == nil {
+ return
+ }
+ receiver := receiverUserID
+ if isGroup {
+ receiver = roomID
+ }
+ botMsg := model.ChatMessage{
+ RoomID: roomID,
+ SenderUserID: bot.UserID,
+ ReceiverUserID: receiver,
+ MessageType: model.MessageTypeText,
+ Content: content,
+ }
+ if err := s.DB.Create(&botMsg).Error; err != nil {
+ log.Printf("❌ [AIBot] 回复落库失败: %v", err)
+ return
+ }
+
+ // 更新会话列表(未读数、最后一条消息),与常规消息路径一致
+ if ConversationSvc != nil {
+ if isGroup {
+ if memberIDs, mErr := RoomSvc.GetRoomMembers(roomID); mErr == nil {
+ for _, memberID := range memberIDs {
+ if memberID != bot.UserID {
+ _ = ConversationSvc.UpsertConversationOnMessage(memberID, roomID, roomID, botMsg, false)
+ }
+ }
+ }
+ } else if receiverUserID != "" {
+ _ = ConversationSvc.UpsertConversationOnMessage(receiverUserID, bot.UserID, roomID, botMsg, false)
+ }
+ }
+
+ // WS 推送
+ pushMsg := model.WsPayload{
+ RequestType: "receive_message",
+ Data: botMsg,
+ }
+ msgBytes, err := json.Marshal(pushMsg)
+ if err != nil {
+ return
+ }
+ if isGroup {
+ if memberIDs, mErr := RoomSvc.GetRoomMembers(roomID); mErr == nil {
+ for _, uid := range memberIDs {
+ if uid != bot.UserID {
+ ChatSvc.DispatchMessage(uid, msgBytes)
+ }
+ }
+ }
+ } else if receiverUserID != "" {
+ ChatSvc.DispatchMessage(receiverUserID, msgBytes)
+ }
+ log.Printf("🤖 [AIBot] %s 已回复 room=%s", bot.Name, roomID)
+}
diff --git a/internal/service/ai_provider.go b/internal/service/ai_provider.go
new file mode 100644
index 0000000..fc68daa
--- /dev/null
+++ b/internal/service/ai_provider.go
@@ -0,0 +1,170 @@
+/**
+ * package service
+ * 作用:AI 大模型提供商抽象(工厂模式)。
+ *
+ * 设计:
+ * AIProvider 接口 ← openAICompatProvider(OpenAI 兼容协议实现)
+ * 覆盖 openai / deepseek / dashscope / moonshot / ollama / custom
+ * NewAIProvider(cfg) 为工厂函数:按 cfg.Provider 创建对应实例,
+ * 未来接入非兼容协议的厂商(如 anthropic / gemini)时新增实现类即可,调用方零改动。
+ *
+ * 为什么不引第三方 SDK:机器人只需要 chat/completions 一个端点,
+ * 直接用 net/http 实现,避免引入体积大、依赖多的 SDK。
+ */
+package service
+
+import (
+ "bytes"
+ "context"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "io"
+ "net/http"
+ "strings"
+ "time"
+
+ "xk-websocket-v2/internal/model"
+)
+
+// AIMessage 对话消息(role: system/user/assistant)
+type AIMessage struct {
+ Role string `json:"role"`
+ Content string `json:"content"`
+}
+
+// AIProvider AI 提供商接口:工厂模式的产品抽象
+type AIProvider interface {
+ // Chat 发起一轮对话补全,返回助手回复文本
+ Chat(ctx context.Context, messages []AIMessage) (string, error)
+ // Name 提供商名称(日志用)
+ Name() string
+}
+
+// 各提供商的默认 API 基础地址(OpenAI 兼容协议端点)
+var providerDefaultBaseURL = map[string]string{
+ "openai": "https://api.openai.com/v1",
+ "deepseek": "https://api.deepseek.com/v1",
+ "dashscope": "https://dashscope.aliyuncs.com/compatible-mode/v1",
+ "moonshot": "https://api.moonshot.cn/v1",
+ "ollama": "http://127.0.0.1:11434/v1",
+}
+
+/**
+ * NewAIProvider
+ * 作用:工厂函数——按配置创建 AIProvider 实例。
+ * 当前所有支持的提供商都走 OpenAI 兼容协议(同一实现、不同 base_url),
+ * 新增非兼容厂商时在 switch 里挂新实现即可
+ */
+func NewAIProvider(cfg *model.AIConfig) (AIProvider, error) {
+ if cfg == nil {
+ return nil, errors.New("AI 配置为空")
+ }
+ switch cfg.Provider {
+ case "openai", "deepseek", "dashscope", "moonshot", "ollama", "custom":
+ baseURL := strings.TrimRight(cfg.BaseURL, "/")
+ if baseURL == "" {
+ baseURL = providerDefaultBaseURL[cfg.Provider]
+ }
+ if baseURL == "" {
+ return nil, errors.New("custom 提供商必须填写 API 基础地址")
+ }
+ if cfg.Model == "" {
+ return nil, errors.New("模型名称未配置")
+ }
+ return &openAICompatProvider{
+ provider: cfg.Provider,
+ baseURL: baseURL,
+ apiKey: cfg.APIKey,
+ model: cfg.Model,
+ client: &http.Client{Timeout: 60 * time.Second},
+ }, nil
+ default:
+ return nil, fmt.Errorf("不支持的AI提供商: %s", cfg.Provider)
+ }
+}
+
+// openAICompatProvider OpenAI 兼容协议实现(chat/completions)
+type openAICompatProvider struct {
+ provider string
+ baseURL string
+ apiKey string
+ model string
+ client *http.Client
+}
+
+func (p *openAICompatProvider) Name() string {
+ return p.provider
+}
+
+// chatCompletionReq /chat/completions 请求体
+type chatCompletionReq struct {
+ Model string `json:"model"`
+ Messages []AIMessage `json:"messages"`
+ // 机器人用于聊天场景,限制回复长度避免刷屏与超额计费
+ MaxTokens int `json:"max_tokens,omitempty"`
+}
+
+// chatCompletionResp /chat/completions 响应体(只解析需要的字段)
+type chatCompletionResp struct {
+ Choices []struct {
+ Message struct {
+ Content string `json:"content"`
+ } `json:"message"`
+ } `json:"choices"`
+ Error *struct {
+ Message string `json:"message"`
+ } `json:"error"`
+}
+
+/**
+ * Chat
+ * 作用:调用 OpenAI 兼容的 chat/completions 端点完成一轮对话
+ */
+func (p *openAICompatProvider) Chat(ctx context.Context, messages []AIMessage) (string, error) {
+ body, err := json.Marshal(chatCompletionReq{
+ Model: p.model,
+ Messages: messages,
+ MaxTokens: 1024,
+ })
+ if err != nil {
+ return "", err
+ }
+
+ req, err := http.NewRequestWithContext(ctx, http.MethodPost, p.baseURL+"/chat/completions", bytes.NewReader(body))
+ if err != nil {
+ return "", err
+ }
+ req.Header.Set("Content-Type", "application/json")
+ // ollama 本地部署可以无鉴权,其余提供商必须带 key
+ if p.apiKey != "" {
+ req.Header.Set("Authorization", "Bearer "+p.apiKey)
+ }
+
+ resp, err := p.client.Do(req)
+ if err != nil {
+ return "", fmt.Errorf("AI请求失败: %w", err)
+ }
+ defer resp.Body.Close()
+
+ // 限制响应体大小,防御异常服务返回超大内容拖垮内存
+ respBody, err := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
+ if err != nil {
+ return "", err
+ }
+
+ var result chatCompletionResp
+ if err := json.Unmarshal(respBody, &result); err != nil {
+ return "", fmt.Errorf("AI响应解析失败(HTTP %d)", resp.StatusCode)
+ }
+ if result.Error != nil && result.Error.Message != "" {
+ return "", fmt.Errorf("AI服务错误: %s", result.Error.Message)
+ }
+ if resp.StatusCode != http.StatusOK {
+ return "", fmt.Errorf("AI服务返回 HTTP %d", resp.StatusCode)
+ }
+ if len(result.Choices) == 0 {
+ return "", errors.New("AI未返回任何回复")
+ }
+ return strings.TrimSpace(result.Choices[0].Message.Content), nil
+}
diff --git a/internal/service/attachment_service.go b/internal/service/attachment_service.go
index 15e2ea0..45ab438 100644
--- a/internal/service/attachment_service.go
+++ b/internal/service/attachment_service.go
@@ -71,7 +71,7 @@ func (s *AttachmentService) UploadFile(userID, fileName, fileType string, fileSi
// 获取文件扩展名
ext := strings.ToLower(filepath.Ext(fileName))
-
+
// 验证文件扩展名(file类型不限制扩展名)
if fileType != "file" {
allowedExts := s.getAllowedExtensions(fileType)
@@ -80,10 +80,10 @@ func (s *AttachmentService) UploadFile(userID, fileName, fileType string, fileSi
}
}
- // 生成唯一文件名
+ // 生成唯一且不可猜测的文件名。
+ // 原实现的随机段由同一时间戳取模得到,可被推算;改用加密随机串,避免他人通过枚举 URL 下载私有文件。
timestamp := time.Now().UnixNano()
- randomStr := fmt.Sprintf("%d", timestamp%1000000)
- newFileName := fmt.Sprintf("%d_%s%s", timestamp, randomStr, ext)
+ newFileName := fmt.Sprintf("%d_%s%s", timestamp, generateRandomString(16), ext)
// 确定保存路径
var saveDir string
@@ -245,4 +245,3 @@ func contains(slice []string, item string) bool {
}
return false
}
-
diff --git a/internal/service/auth_service.go b/internal/service/auth_service.go
index 29fe0ca..c69c2a3 100644
--- a/internal/service/auth_service.go
+++ b/internal/service/auth_service.go
@@ -9,6 +9,7 @@ import (
"encoding/hex"
"errors"
"fmt"
+ "log"
"math/big"
"strings"
"time"
@@ -45,7 +46,7 @@ func InitAuthService(db *gorm.DB, rdb *redis.Client) {
*/
func (s *AuthService) Login(account, password string) (*model.User, error) {
var user model.User
-
+
// 根据邮箱或手机号查询用户
result := s.DB.Where("email = ? OR phone = ?", account, account).First(&user)
if result.Error != nil {
@@ -135,7 +136,7 @@ func (s *AuthService) Register(req *model.RegisterReq) (*model.User, error) {
func (s *AuthService) SendEmailCode(email string) (string, error) {
// 生成6位验证码
code := generateCode(6)
-
+
// 保存验证码到数据库(5分钟过期)
vc := model.VerificationCode{
Target: email,
@@ -162,7 +163,7 @@ func (s *AuthService) SendEmailCode(email string) (string, error) {
func (s *AuthService) SendSmsCode(phone string) (string, error) {
// 生成6位验证码
code := generateCode(6)
-
+
// 保存验证码到数据库(5分钟过期)
vc := model.VerificationCode{
Target: phone,
@@ -182,7 +183,11 @@ func (s *AuthService) SendSmsCode(phone string) (string, error) {
/**
* VerifyCode
- * 功能:验证验证码
+ * 功能:验证验证码(只读校验,不消费)。
+ * 为什么不在这里删除:调用方(注册)在校验之后还有可失败的步骤(落库、
+ * 邮箱/手机号重复等),校验即删会导致注册失败后同一验证码无法重试,
+ * 用户被迫重新获取。消费动作由 ConsumeCode 在业务成功后显式执行——
+ * 与扫码登录"窥视→动作→原子消费"同一模式
* @param target 邮箱或手机号
* @param code 验证码
* @param codeType 类型(email/sms)
@@ -190,13 +195,13 @@ func (s *AuthService) SendSmsCode(phone string) (string, error) {
*/
func (s *AuthService) VerifyCode(target, code, codeType string) (bool, error) {
var vc model.VerificationCode
-
+
// 查询验证码
- result := s.DB.Where("target = ? AND code = ? AND type = ? AND expires_at > ?",
+ result := s.DB.Where("target = ? AND code = ? AND type = ? AND expires_at > ?",
target, code, codeType, time.Now()).
Order("created_at DESC").
First(&vc)
-
+
if result.Error != nil {
if errors.Is(result.Error, gorm.ErrRecordNotFound) {
return false, errors.New("验证码无效或已过期")
@@ -207,18 +212,49 @@ func (s *AuthService) VerifyCode(target, code, codeType string) (bool, error) {
return true, nil
}
+/**
+ * ConsumeCode
+ * 功能:消费(删除)验证码,业务成功后调用,防止同一验证码被重放使用。
+ * 说明:按 (target, code, type) 删除全部匹配记录,重复调用幂等;
+ * 校验到消费之间的并发重放窗口由注册本身的邮箱/手机唯一约束兜底,无实际危害
+ */
+func (s *AuthService) ConsumeCode(target, code, codeType string) {
+ if err := s.DB.Where("target = ? AND code = ? AND type = ?", target, code, codeType).
+ Delete(&model.VerificationCode{}).Error; err != nil {
+ // 删除失败只记日志不阻断业务:验证码有 5 分钟 TTL 兜底过期
+ log.Printf("⚠️ [Auth] 消费验证码失败 target=%s: %v", target, err)
+ }
+}
+
// 辅助函数:生成随机字符串
+// rand.Read 错误必须处理:crypto/rand 失败时缓冲区是全零字节,
+// 生成的"随机"串完全可预测(用于用户ID后缀等场景会产生碰撞/被猜测风险),
+// 此时降级为纳秒时间戳兜底,保证仍有基本的唯一性
func generateRandomString(length int) string {
b := make([]byte, length)
- rand.Read(b)
+ if _, err := rand.Read(b); err != nil {
+ log.Printf("⚠️ crypto/rand 读取失败,降级为时间戳随机: %v", err)
+ ts := fmt.Sprintf("%x", time.Now().UnixNano())
+ for len(ts) < length {
+ ts += ts
+ }
+ return ts[:length]
+ }
return hex.EncodeToString(b)[:length]
}
// 辅助函数:生成验证码
+// rand.Int 错误必须处理:忽略错误时 n 为 nil 会直接 panic(n.String() 空指针),
+// 失败位降级为纳秒时间戳取模,保证验证码仍是完整位数的数字
func generateCode(length int) string {
code := ""
for i := 0; i < length; i++ {
- n, _ := rand.Int(rand.Reader, big.NewInt(10))
+ n, err := rand.Int(rand.Reader, big.NewInt(10))
+ if err != nil {
+ log.Printf("⚠️ crypto/rand 生成验证码位失败,降级为时间戳随机: %v", err)
+ code += fmt.Sprintf("%d", time.Now().UnixNano()%10)
+ continue
+ }
code += n.String()
}
return code
@@ -232,4 +268,3 @@ func generateAvatar(userID string) string {
}
return "U"
}
-
diff --git a/internal/service/chat_service.go b/internal/service/chat_service.go
index 63ff2a6..67441b3 100644
--- a/internal/service/chat_service.go
+++ b/internal/service/chat_service.go
@@ -6,15 +6,20 @@ package service
import (
"context"
+ "crypto/sha1"
+ "encoding/hex"
"encoding/json"
+ "errors"
"fmt"
+ "html"
"io"
"log"
+ "net"
"net/http"
+ "net/url"
"regexp"
"strings"
"time"
- "unicode/utf8"
"xk-websocket-v2/internal/manager"
"xk-websocket-v2/internal/model"
@@ -23,6 +28,7 @@ import (
"github.com/go-redis/redis/v8"
"github.com/spf13/viper"
+ "golang.org/x/net/html/charset"
"gorm.io/gorm"
)
@@ -77,6 +83,52 @@ func (s *ChatService) BindUserByClientID(clientID, userID string) {
s.Redis.Set(ctx, key, nodeID, 24*time.Hour)
}
+// delRouteIfOwned 原子化的"值匹配才删除"脚本(compare-and-delete)。
+// 为什么需要:先 GET 确认路由指向本节点、再 DEL 是两步操作,间隙中用户可能已在
+// 其他节点重连(BindUser 已把路由 SET 成新节点 ID),继续 DEL 会误删新节点的路由,
+// DispatchMessage 完全依赖该路由 key,误删后用户"在线却收不到任何推送"。
+// Lua 脚本在 Redis 内单线程执行,读取、比较、删除一步完成,跨节点竞态被彻底消除。
+var delRouteIfOwned = redis.NewScript(`
+if redis.call("GET", KEYS[1]) == ARGV[1] then
+ return redis.call("DEL", KEYS[1])
+end
+return 0`)
+
+/**
+ * HandleClientOffline
+ * 功能:客户端断开连接时的统一清理入口。
+ * 为什么需要:原实现只注销本地连接,不清 Redis 路由 key(24小时过期),
+ * 导致用户下线后长时间被误判"在线",且跨节点消息会被继续投递到已无连接的节点。
+ * 逻辑:注销本地连接后,若该用户在本节点已无任何设备在线、且路由仍指向本节点,
+ * 则用 Lua 脚本原子化地"值匹配才删除"路由 key;
+ * 若路由已指向其他节点(用户在别的节点重连),脚本比较不通过、不会删除。
+ */
+func (s *ChatService) HandleClientOffline(client *manager.Client) {
+ // 在锁保护下读取 UserID,避免与 BindUser 的并发写入构成数据竞争
+ userID := manager.Manager.GetClientUserID(client)
+ manager.Manager.Unregister(client)
+
+ if userID == "" {
+ return
+ }
+ // 本节点还有该用户的其他设备在线,不清理
+ if manager.Manager.UserHasClients(userID) {
+ return
+ }
+
+ ctx := context.Background()
+ key := model.KeyUserNodeMap + userID
+ nodeID := viper.GetString("app.node_id")
+ // 原子化删除:仅当路由值仍等于本节点 ID 时才删除(跨节点重连后值已变,脚本不删)
+ if err := delRouteIfOwned.Run(ctx, s.Redis, []string{key}, nodeID).Err(); err == nil {
+ // 自愈复查:同节点重连时新旧值相同(都是本节点 ID),脚本无法区分仍会删。
+ // 删除后发现本节点又有该用户的连接,说明删的是新路由,立即补写回去
+ if manager.Manager.UserHasClients(userID) {
+ s.Redis.Set(ctx, key, nodeID, 24*time.Hour)
+ }
+ }
+}
+
/**
* IsUserOnline
* 功能:检查用户是否在线 (查询 Redis)。
@@ -91,34 +143,37 @@ func (s *ChatService) IsUserOnline(userID string) bool {
/**
* HandleUserMessage
* 功能:处理用户发来的消息(入口函数)。包含多端同步、URL 抓取、持久化和转发逻辑。
- * @param senderClient 发送消息的客户端(用于识别来源设备)
+ * @param senderClient 发送消息的客户端(用于识别发送者的 UserID 和来源设备)
* @param req 消息请求体
- * @returns 持久化后的消息(信令/群通知/被拦截时返回 nil)
+ * @returns (持久化后的消息, 错误)。
+ * - 常规消息成功:返回 (&msg, nil)
+ * - 信令消息(无需持久化):返回 (nil, nil)
+ * - 校验拒绝/落库失败:返回 (nil, err)。
+ * 为什么必须返回 error:HTTP /send 路径原先只看返回的消息指针,信令正常返回 nil
+ * 与"被拒绝"无法区分,被拒消息也会响应"消息已发送",发送方完全无感知
*/
-func (s *ChatService) HandleUserMessage(senderClient *manager.Client, req *model.SendMessageReq) *model.ChatMessage {
+func (s *ChatService) HandleUserMessage(senderClient *manager.Client, req *model.SendMessageReq) (*model.ChatMessage, error) {
+ // 0. 群通知(7)只能由服务端生成:客户端上行一律拒绝。
+ // 原先只是"不持久化",但仍会掉进下方信令分支被实时广播,
+ // 伪造的"xxx被移出群聊"通知照样能刷到全群屏幕上(刷新后才消失)
+ if req.MessageType == model.MessageTypeGroupNotif {
+ return nil, errors.New("非法的消息类型")
+ }
+
// 1. [关键] 拦截通话信令:处理多端同步逻辑
if req.CallStatus == "accepted" {
s.NotifyOtherDevices(senderClient.UserID, senderClient.ID, req.CallID)
}
- // [FIXED] 核心修复:直接使用前端传递的 Extra
+ // 直接使用前端传递的 Extra。
+ // URL 预览抓取不在发送路径上同步执行(最坏阻塞 5 秒,体感"发送卡住"):
+ // 消息先落库分发,分发完成后由 maybeScrapeURLAsync 异步抓取、
+ // 更新 extra 并推送 message_extra_update 补出卡片(仿微信:消息秒发、卡片随后出现)
var extraData string = req.Extra
- // 2. URL 识别与抓取逻辑
- if req.MessageType == 0 {
- url := extractURL(req.Content)
- if url != "" {
- meta := s.scrapeURL(url)
- if meta != nil {
- metaJson, _ := json.Marshal(meta)
- extraData = string(metaJson)
- log.Printf("🌐 抓取成功: %s", meta.Title)
- }
- }
- }
-
// 3. 消息持久化 (MySQL)
- // 信令消息(6) 和 群通知(1000+) 不持久化
+ // 信令消息(6)不持久化;群通知(7)只能由服务端通过 BroadcastGroupNotification 落库,
+ // 客户端上行的 type=7 一律不持久化,防止伪造"xxx 被移出群聊"之类的系统通知
if req.MessageType != model.MessageTypeSignal && req.MessageType != model.MessageTypeGroupNotif {
// ... (省略常规消息持久化逻辑,保持原样) ...
isGroupMessage := false
@@ -130,7 +185,7 @@ func (s *ChatService) HandleUserMessage(senderClient *manager.Client, req *model
// 验证群成员身份
memberIDs, mErr := RoomSvc.GetRoomMembers(req.RoomID)
if mErr != nil {
- return nil
+ return nil, errors.New("群成员信息查询失败,请重试")
}
isMember := false
@@ -152,7 +207,26 @@ func (s *ChatService) HandleUserMessage(senderClient *manager.Client, req *model
}
errorBytes, _ := json.Marshal(errorMsg)
s.DispatchMessage(senderClient.UserID, errorBytes)
- return nil
+ return nil, errors.New("您已被移出群聊,无法发送消息")
+ }
+
+ // 群禁言校验:原实现只在接口上能设置禁言,发消息路径不校验,禁言形同虚设。
+ // 此处在发言前检查 room_members.muted_until,禁言期内直接拒绝并回错误提示。
+ if muted, mutedUntil, muteErr := RoomSvc.IsGroupMemberMuted(req.RoomID, senderClient.UserID); muteErr == nil && muted {
+ tip := "您已被禁言,暂时无法发言"
+ if mutedUntil != nil {
+ tip = fmt.Sprintf("您已被禁言至 %s,暂时无法发言", mutedUntil.Format("2006-01-02 15:04"))
+ }
+ errorMsg := model.WsPayload{
+ RequestType: "error",
+ Data: map[string]interface{}{
+ "message": tip,
+ "code": "MUTED",
+ },
+ }
+ errorBytes, _ := json.Marshal(errorMsg)
+ s.DispatchMessage(senderClient.UserID, errorBytes)
+ return nil, errors.New(tip)
}
}
}
@@ -162,6 +236,65 @@ func (s *ChatService) HandleUserMessage(senderClient *manager.Client, req *model
receiverUserID = req.RoomID
}
+ // 私聊安全校验(与历史读取端 HistoryHandler 的 IsRoomMember 鉴权对称):
+ // room_id 与 receiver_user_id 均由客户端提供,此前发送路径不做任何校验,
+ // 可把消息写进任意他人的 P2P 房间(污染他人历史记录)或伪造房间归属。
+ // 合法判定(满足其一):
+ // 1) room_id 等于 GenerateP2PRoomID(发送者, 接收者) 按规则生成的ID(当前统一规则,前端合成一致,零额外查询);
+ // 2) 兼容存量数据:历史版本曾用随机雪花ID给 P2P 房间命名(见 AcceptFriendRequest 注释),
+ // 这类房间只要真实存在、类型为 p2p、且发送者与接收者都是房间成员,同样放行——
+ // 与读取端 IsRoomMember 的口径一致,攻击者不是目标房间成员时依然会被拒绝。
+ if !isGroupMessage {
+ legal := false
+ if req.ReceiverUserID != "" && req.RoomID != "" {
+ if req.RoomID == GenerateP2PRoomID(senderClient.UserID, req.ReceiverUserID) {
+ legal = true
+ } else if room, rErr := RoomSvc.GetRoom(req.RoomID); rErr == nil && room.RoomType == "p2p" {
+ legal = RoomSvc.IsRoomMember(req.RoomID, senderClient.UserID) &&
+ RoomSvc.IsRoomMember(req.RoomID, req.ReceiverUserID)
+ }
+ }
+ if !legal {
+ log.Printf("🚫 [Chat] 拒绝非法私聊消息: sender=%s receiver=%s room=%s",
+ senderClient.UserID, req.ReceiverUserID, req.RoomID)
+ errorMsg := model.WsPayload{
+ RequestType: "error",
+ Data: map[string]interface{}{
+ "message": "会话不合法,消息发送失败",
+ "code": "INVALID_ROOM",
+ },
+ }
+ errorBytes, _ := json.Marshal(errorMsg)
+ s.DispatchMessage(senderClient.UserID, errorBytes)
+ return nil, errors.New("会话不合法,消息发送失败")
+ }
+ }
+
+ // 私聊好友关系实时校验:删除好友只清理了 user_contacts/会话,房间与成员可能残留,
+ // 上面的房间合法性只证明"房间归属这两个人",不代表"当前仍是好友",
+ // 不校验的话删好友后双方仍能通过残留房间继续互发私聊。
+ // 豁免场景:AI 机器人(bot_ 前缀)是全局虚拟好友,没有 user_contacts 行;
+ // 自己给自己发消息(文件传输助手类场景)也不做好友校验
+ if !isGroupMessage && ContactSvc != nil &&
+ senderClient.UserID != req.ReceiverUserID &&
+ !strings.HasPrefix(senderClient.UserID, "bot_") &&
+ !strings.HasPrefix(req.ReceiverUserID, "bot_") {
+ if !ContactSvc.IsFriend(senderClient.UserID, req.ReceiverUserID) {
+ log.Printf("🚫 [Chat] 拒绝非好友私聊消息: sender=%s receiver=%s room=%s",
+ senderClient.UserID, req.ReceiverUserID, req.RoomID)
+ errorMsg := model.WsPayload{
+ RequestType: "error",
+ Data: map[string]interface{}{
+ "message": "你们还不是好友,无法发送消息",
+ "code": "NOT_FRIEND",
+ },
+ }
+ errorBytes, _ := json.Marshal(errorMsg)
+ s.DispatchMessage(senderClient.UserID, errorBytes)
+ return nil, errors.New("你们还不是好友,无法发送消息")
+ }
+ }
+
// 私聊:接收方已拉黑发送方则拒绝
if !isGroupMessage && req.ReceiverUserID != "" && ContactSvc != nil {
if ContactSvc.IsBlocked(req.ReceiverUserID, senderClient.UserID) {
@@ -174,7 +307,7 @@ func (s *ChatService) HandleUserMessage(senderClient *manager.Client, req *model
}
errorBytes, _ := json.Marshal(errorMsg)
s.DispatchMessage(senderClient.UserID, errorBytes)
- return nil
+ return nil, errors.New("消息已发出,但被对方拒收")
}
}
@@ -193,7 +326,29 @@ func (s *ChatService) HandleUserMessage(senderClient *manager.Client, req *model
}
if err := s.DB.Create(&msg).Error; err != nil {
+ // 落库失败必须终止推送:否则接收方看到了消息、但双方历史记录里都查不到,
+ // 造成"消息闪现后消失"的不一致。这里给发送者回错误提示并直接返回。
log.Printf("❌ 消息持久化失败: %v", err)
+ errorMsg := model.WsPayload{
+ RequestType: "error",
+ Data: map[string]interface{}{
+ "message": "消息发送失败,请重试",
+ "code": "PERSIST_FAILED",
+ },
+ }
+ errorBytes, _ := json.Marshal(errorMsg)
+ s.DispatchMessage(senderClient.UserID, errorBytes)
+ return nil, errors.New("消息发送失败,请重试")
+ }
+
+ // 标记来源设备(不入库):消息会回推给发送者的所有设备做多端同步,
+ // 发送端根据 sender_client_id 识别自己的回声并忽略,避免重复渲染
+ msg.SenderClientID = senderClient.ID
+
+ // AI 机器人应答钩子:私聊发给机器人 / 群聊 @机器人 时异步生成回复。
+ // 放在落库成功之后,保证机器人取上下文时能读到触发它的这条消息
+ if AIBotSvc != nil {
+ AIBotSvc.MaybeReply(&msg, isGroupMessage)
}
// 更新会话列表逻辑 (保持原样)
@@ -230,7 +385,9 @@ func (s *ChatService) HandleUserMessage(senderClient *manager.Client, req *model
}
msgBytes, _ := json.Marshal(pushMsg)
- // 分发消息 (保持原样)
+ // 分发消息:群聊推给所有成员,私聊推给接收方;
+ // 同时回推给发送者本人(覆盖其所有在线设备),保证发送者的其他端也能实时同步,
+ // 发送端自身通过 sender_client_id 忽略回声(原实现只有撤回会回推,普通消息多端不同步)
if req.RoomID != "" {
room, err := RoomSvc.GetRoom(req.RoomID)
if err == nil && room.RoomType == "group" {
@@ -242,15 +399,25 @@ func (s *ChatService) HandleUserMessage(senderClient *manager.Client, req *model
s.DispatchMessage(uid, msgBytes)
}
}
+ // 回推发送者(多端同步)
+ s.DispatchMessage(senderID, msgBytes)
// 处理@通知
s.handleMentionNotification(senderClient.UserID, req.RoomID, extraData, msg)
- return &msg
+ // 消息已送达,异步抓取 URL 预览并补推卡片(不阻塞发送路径)
+ s.maybeScrapeURLAsync(&msg)
+ return &msg, nil
}
}
}
s.DispatchMessage(req.ReceiverUserID, msgBytes)
- return &msg
+ // 私聊也回推发送者(多端同步);给自己发消息时上面已推过,避免重复
+ if req.ReceiverUserID != senderClient.UserID {
+ s.DispatchMessage(senderClient.UserID, msgBytes)
+ }
+ // 消息已送达,异步抓取 URL 预览并补推卡片(不阻塞发送路径)
+ s.maybeScrapeURLAsync(&msg)
+ return &msg, nil
} else {
// ==========================================
@@ -296,7 +463,8 @@ func (s *ChatService) HandleUserMessage(senderClient *manager.Client, req *model
log.Printf("📡 [WebRTC] 定向信令: To=%s Action=%s", req.ReceiverUserID, req.CallStatus)
}
}
- return nil
+ // 信令消息无需持久化,正常结束(区别于上方"被拒绝"的 error 返回)
+ return nil, nil
}
// 辅助:正则提取第一个 URL
@@ -305,11 +473,184 @@ func extractURL(text string) string {
return re.FindString(text)
}
-// 辅助:优化后的网页抓取 (Title, Desc, Image)
-// 解决乱码和图片路径问题
-func (s *ChatService) scrapeURL(url string) *model.UrlMeta {
- client := &http.Client{Timeout: 5 * time.Second}
- resp, err := client.Get(url)
+// isPrivateOrReservedIP 判断 IP 是否为内网/保留地址(用于防 SSRF)
+func isPrivateOrReservedIP(ip net.IP) bool {
+ if ip == nil {
+ return true
+ }
+ if ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() || ip.IsUnspecified() {
+ return true
+ }
+ // 云厂商元数据地址,必须禁止
+ if ip.String() == "169.254.169.254" {
+ return true
+ }
+ return false
+}
+
+// isSafeExternalURL 校验 URL 是否指向安全的公网地址。
+// 为什么这样写:URL 预览会由服务端主动请求用户消息里的链接,若不校验会被用于 SSRF 探测内网/元数据服务,
+// 因此这里限制仅允许 http(s) 协议,并解析域名对应的所有 IP,任一为内网/保留地址即拒绝。
+func isSafeExternalURL(rawURL string) bool {
+ u, err := url.Parse(rawURL)
+ if err != nil {
+ return false
+ }
+ if u.Scheme != "http" && u.Scheme != "https" {
+ return false
+ }
+ host := u.Hostname()
+ if host == "" {
+ return false
+ }
+ ips, err := net.LookupIP(host)
+ if err != nil || len(ips) == 0 {
+ return false
+ }
+ for _, ip := range ips {
+ if isPrivateOrReservedIP(ip) {
+ return false
+ }
+ }
+ return true
+}
+
+// safeDialContext 抓取专用拨号器:真正建连后校验实际远端 IP。
+// 为什么需要:isSafeExternalURL 的 LookupIP 校验与 http.Client 发请求是两次独立的 DNS 解析,
+// 攻击者域名可在校验通过后把解析切到内网(DNS rebinding)。
+// 这里把"校验"绑定到建连拿到的 RemoteAddr 上,与实际请求同一连接,无法绕过
+func safeDialContext(ctx context.Context, network, addr string) (net.Conn, error) {
+ dialer := &net.Dialer{Timeout: 5 * time.Second}
+ conn, err := dialer.DialContext(ctx, network, addr)
+ if err != nil {
+ return nil, err
+ }
+ if tcpAddr, ok := conn.RemoteAddr().(*net.TCPAddr); ok && isPrivateOrReservedIP(tcpAddr.IP) {
+ conn.Close()
+ return nil, fmt.Errorf("禁止访问内网地址")
+ }
+ return conn, nil
+}
+
+// ---- URL 预览缓存 ----
+// key 用 URL 的 SHA1(URL 可能超长且含特殊字符,不宜直接做 Redis key)。
+// 成功结果缓存 24h(页面元数据变化不频繁,同一链接重复发送秒出卡片);
+// 失败缓存 10 分钟空标记,防止对故障/超时站点反复发起抓取
+const (
+ urlPreviewCachePrefix = "url_preview:"
+ urlPreviewCacheTTL = 24 * time.Hour
+ urlPreviewNegativeTTL = 10 * time.Minute
+ urlPreviewNegativeValue = "-" // 空标记:表示"抓过且失败",区别于"没抓过"
+)
+
+// urlPreviewCacheKey 计算 URL 对应的缓存 key
+func urlPreviewCacheKey(rawURL string) string {
+ sum := sha1.Sum([]byte(rawURL))
+ return urlPreviewCachePrefix + hex.EncodeToString(sum[:])
+}
+
+// titleTagRe 提取 文本((?is) 跨行 + 忽略大小写;容忍 等属性)
+var titleTagRe = regexp.MustCompile(`(?is)]*>(.*?)`)
+
+// extractMetaContent 提取 标签的 content 值,兼容两种常见写法:
+// property/name 在前()
+// 与 content 在前()。
+// 每次调用动态编译正则的成本可接受:抓取在异步路径且有缓存挡重复
+func extractMetaContent(htmlText, key string) string {
+ quoted := regexp.QuoteMeta(key)
+ re1 := regexp.MustCompile(`(?is)]+(?:property|name)\s*=\s*["']` + quoted + `["'][^>]*?content\s*=\s*["'](.*?)["']`)
+ if m := re1.FindStringSubmatch(htmlText); len(m) > 1 {
+ return strings.TrimSpace(m[1])
+ }
+ re2 := regexp.MustCompile(`(?is)]+content\s*=\s*["'](.*?)["'][^>]*?(?:property|name)\s*=\s*["']` + quoted + `["']`)
+ if m := re2.FindStringSubmatch(htmlText); len(m) > 1 {
+ return strings.TrimSpace(m[1])
+ }
+ return ""
+}
+
+// cleanMetaText 元数据文本清洗:HTML 实体解码(& 等不再原样显示)、
+// 空白规整(换行/连续空格合并为单空格)、按字符数截断(防超长撑爆 extra 与 UI)
+func cleanMetaText(s string, maxRunes int) string {
+ s = html.UnescapeString(s)
+ s = strings.Join(strings.Fields(s), " ")
+ r := []rune(s)
+ if len(r) > maxRunes {
+ return string(r[:maxRunes])
+ }
+ return s
+}
+
+// firstNonEmpty 返回第一个非空字符串(用于 og 标签 → 传统标签的回退链)
+func firstNonEmpty(values ...string) string {
+ for _, v := range values {
+ if v != "" {
+ return v
+ }
+ }
+ return ""
+}
+
+/**
+ * scrapeURL
+ * 作用:抓取 URL 的预览元数据(标题/描述/缩略图),带 Redis 缓存。
+ * 缓存命中(含失败空标记)时不发起网络请求;未命中走 fetchURLMeta 实际抓取并回填缓存
+ */
+func (s *ChatService) scrapeURL(rawURL string) *model.UrlMeta {
+ ctx := context.Background()
+ cacheKey := urlPreviewCacheKey(rawURL)
+ if cached, err := s.Redis.Get(ctx, cacheKey).Result(); err == nil {
+ if cached == urlPreviewNegativeValue {
+ return nil
+ }
+ meta := &model.UrlMeta{}
+ if json.Unmarshal([]byte(cached), meta) == nil && meta.Title != "" {
+ return meta
+ }
+ }
+
+ meta := s.fetchURLMeta(rawURL)
+ if meta == nil {
+ s.Redis.Set(ctx, cacheKey, urlPreviewNegativeValue, urlPreviewNegativeTTL)
+ return nil
+ }
+ if data, err := json.Marshal(meta); err == nil {
+ s.Redis.Set(ctx, cacheKey, data, urlPreviewCacheTTL)
+ }
+ return meta
+}
+
+/**
+ * fetchURLMeta
+ * 作用:实际执行网页抓取与解析(不含缓存)。
+ * SSRF 防护三道闸保持不变:入口校验 + 建连校验实际远端 IP + 重定向逐跳校验。
+ * 解析增强点:
+ * - charset.NewReader 按 Content-Type / HTML meta 自动探测编码并转 UTF-8(GBK/GB2312 中文站不再直接放弃)
+ * - og:title/og:description 优先于 /meta description(og 内容通常更准确)
+ * - 实体解码 + 空白规整 + 长度截断
+ * - 图片相对路径用 ResolveReference 解析(修复协议相对路径 //cdn.xx/a.png 被拼错的问题)
+ */
+func (s *ChatService) fetchURLMeta(rawURL string) *model.UrlMeta {
+ // 防 SSRF 第一道滤网:仅允许公网 http(s) 地址(快速拒绝明显非法目标)
+ if !isSafeExternalURL(rawURL) {
+ return nil
+ }
+ client := &http.Client{
+ Timeout: 5 * time.Second,
+ // 防 SSRF 第二道(关键):建连时校验实际远端 IP,堵死 DNS rebinding
+ Transport: &http.Transport{DialContext: safeDialContext},
+ // 跟随重定向时对每一跳都做安全校验,防止通过重定向绕过 SSRF 检查
+ CheckRedirect: func(req *http.Request, via []*http.Request) error {
+ if len(via) >= 5 {
+ return fmt.Errorf("重定向次数过多")
+ }
+ if !isSafeExternalURL(req.URL.String()) {
+ return fmt.Errorf("不允许重定向到非公网地址")
+ }
+ return nil
+ },
+ }
+ resp, err := client.Get(rawURL)
if err != nil {
return nil
}
@@ -318,59 +659,196 @@ func (s *ChatService) scrapeURL(url string) *model.UrlMeta {
if resp.StatusCode != 200 {
return nil
}
-
- // 读取更多内容 (100KB) 以确保 meta 标签完整
- buf := make([]byte, 100*1024)
- n, _ := io.ReadFull(resp.Body, buf)
-
- // 简单的编码修复:如果不是有效UTF-8,直接忽略或返回空(生产环境应用 charset decoder)
- html := string(buf[:n])
- if !utf8.ValidString(html) {
- // 尝试作为 GBK 处理(这里简化处理,直接返回 nil 或 url 本身)
+ // 非 HTML 响应(图片/视频/文件直链)没有 meta 标签可解析,直接放弃
+ contentType := resp.Header.Get("Content-Type")
+ if contentType != "" && !strings.Contains(contentType, "text/html") && !strings.Contains(contentType, "application/xhtml") {
return nil
}
- meta := &model.UrlMeta{Url: url}
-
- // 优化正则:支持单双引号,支持换行
- titleRe := regexp.MustCompile(`(?i)(.*?)`)
- matches := titleRe.FindStringSubmatch(html)
- if len(matches) > 1 {
- meta.Title = strings.TrimSpace(matches[1])
+ // 只读前 100KB(meta 标签都在文档头部),并经 charset 转码为 UTF-8
+ reader, err := charset.NewReader(io.LimitReader(resp.Body, 100*1024), contentType)
+ if err != nil {
+ return nil
}
-
- // 优化 og:image 正则
- imgRe := regexp.MustCompile(`(?i) 1 {
- meta.Image = imgMatches[1]
+ body, err := io.ReadAll(reader)
+ if len(body) == 0 {
+ return nil
}
+ _ = err // 读到部分内容也继续解析(LimitReader 截断不算失败)
+ htmlText := string(body)
- // 优化 description 正则
- descRe := regexp.MustCompile(`(?i) 1 {
- meta.Description = strings.TrimSpace(descMatches[1])
- }
+ meta := &model.UrlMeta{Url: rawURL}
+ meta.Title = cleanMetaText(firstNonEmpty(
+ extractMetaContent(htmlText, "og:title"),
+ func() string {
+ if m := titleTagRe.FindStringSubmatch(htmlText); len(m) > 1 {
+ return m[1]
+ }
+ return ""
+ }(),
+ ), 200)
+ meta.Description = cleanMetaText(firstNonEmpty(
+ extractMetaContent(htmlText, "og:description"),
+ extractMetaContent(htmlText, "description"),
+ ), 300)
+ meta.Image = html.UnescapeString(extractMetaContent(htmlText, "og:image"))
+ // 没有标题的卡片没有展示价值,视为抓取失败
if meta.Title == "" {
return nil
}
- // 相对路径处理
- if meta.Image != "" && !strings.HasPrefix(meta.Image, "http") {
- // 简单拼接,不严谨但够用
- if strings.HasPrefix(meta.Image, "/") {
- // 提取 domain
- domainRe := regexp.MustCompile(`(https?://[^/]+)`)
- domain := domainRe.FindString(url)
- meta.Image = domain + meta.Image
+ // 图片相对路径解析:以最终响应地址为基准(重定向后页面内的相对路径相对最终地址),
+ // ResolveReference 正确处理 /a.png、a.png 与协议相对 //cdn.xx/a.png 三种形态
+ if meta.Image != "" {
+ if imgRef, iErr := url.Parse(strings.TrimSpace(meta.Image)); iErr == nil && resp.Request != nil && resp.Request.URL != nil {
+ meta.Image = resp.Request.URL.ResolveReference(imgRef).String()
}
}
return meta
}
+/**
+ * maybeScrapeURLAsync
+ * 作用:文本消息含 URL 时,在消息分发完成后异步抓取预览并补推卡片。
+ * 流程:抓取(带缓存)→ 合并进消息 extra → 更新 DB(拉历史能看到卡片)
+ * → 推送 message_extra_update 给消息相关方(在线端实时刷出卡片)。
+ * 为什么异步:同步抓取最坏阻塞 5 秒,带链接的消息体感"发送卡住";
+ * goroutine 数量与"消息中带链接"的频率相当且有缓存挡重复,无需额外并发限制
+ */
+func (s *ChatService) maybeScrapeURLAsync(msg *model.ChatMessage) {
+ if msg == nil || msg.ID == 0 || msg.MessageType != model.MessageTypeText {
+ return
+ }
+ pageURL := extractURL(msg.Content)
+ if pageURL == "" {
+ return
+ }
+ go func() {
+ meta := s.scrapeURL(pageURL)
+ if meta == nil {
+ // 抓取失败:消息保持纯文本展示,与"从未有卡片"一致,无需通知任何端
+ return
+ }
+
+ // 合并而非覆盖 extra:保留前端传入的其他字段(如 @提及列表)
+ merged := map[string]interface{}{}
+ if msg.Extra != "" {
+ _ = json.Unmarshal([]byte(msg.Extra), &merged) // 解析失败按空 map 处理
+ }
+ merged["url"] = meta.Url
+ merged["title"] = meta.Title
+ if meta.Image != "" {
+ merged["image"] = meta.Image
+ }
+ if meta.Description != "" {
+ merged["description"] = meta.Description
+ }
+ extraJSON, mErr := json.Marshal(merged)
+ if mErr != nil {
+ return
+ }
+
+ // 先落库再推送:保证收到推送的端刷新历史时结果一致
+ if err := s.DB.Model(&model.ChatMessage{}).Where("id = ?", msg.ID).
+ Update("extra", string(extraJSON)).Error; err != nil {
+ log.Printf("❌ [URL预览] 更新消息 extra 失败 (id=%d): %v", msg.ID, err)
+ return
+ }
+
+ payload := model.WsPayload{
+ RequestType: "message_extra_update",
+ Data: map[string]interface{}{
+ "id": msg.ID,
+ "room_id": msg.RoomID,
+ "extra": string(extraJSON),
+ },
+ }
+ payloadBytes, pErr := json.Marshal(payload)
+ if pErr != nil {
+ return
+ }
+ s.dispatchToMessageParticipants(msg, payloadBytes)
+ log.Printf("🌐 [URL预览] 卡片已补推: id=%d title=%s", msg.ID, meta.Title)
+ }()
+}
+
+/**
+ * dispatchToMessageParticipants
+ * 作用:把 payload 推送给一条消息的所有相关方,口径与 receive_message 的分发一致:
+ * 群聊推给除发送者外的全体成员再回推发送者(覆盖其多端);
+ * 私聊推给接收方,发送者非本人时再回推发送者
+ */
+func (s *ChatService) dispatchToMessageParticipants(msg *model.ChatMessage, payloadBytes []byte) {
+ if msg.RoomID != "" {
+ if room, err := RoomSvc.GetRoom(msg.RoomID); err == nil && room.RoomType == "group" {
+ if memberIDs, mErr := RoomSvc.GetRoomMembers(msg.RoomID); mErr == nil {
+ for _, uid := range memberIDs {
+ if uid != msg.SenderUserID {
+ s.DispatchMessage(uid, payloadBytes)
+ }
+ }
+ s.DispatchMessage(msg.SenderUserID, payloadBytes)
+ return
+ }
+ }
+ }
+ // 私聊(群消息的 ReceiverUserID 是房间ID,不会走到这里)
+ s.DispatchMessage(msg.ReceiverUserID, payloadBytes)
+ if msg.ReceiverUserID != msg.SenderUserID {
+ s.DispatchMessage(msg.SenderUserID, payloadBytes)
+ }
+}
+
+/**
+ * BroadcastGroupNotification
+ * 功能:向群成员广播一条"群通知"(入群、踢人、退群、群主转让、公告更新等)。
+ * 为什么这样写:原实现在多处 handler 里复制同一段代码,且为每个成员各落库一条相同的通知
+ * (N 个成员产生 N 条重复记录,历史记录里同一通知重复出现);
+ * 现统一收口:房间级只落库一次,再针对每个成员更新会话摘要并推送 WebSocket。
+ * @param operatorID 触发通知的操作者ID
+ * @param roomID 群房间ID
+ * @param content 通知文案
+ * @param extraJSON 通知扩展JSON(前端据此渲染富通知)
+ * @param memberIDs 需要接收该通知的成员ID列表(调用方可先过滤掉不应接收的人)
+ * @returns 已落库的通知消息;落库失败返回 nil(此时不推送,保证库内与推送一致)
+ */
+func (s *ChatService) BroadcastGroupNotification(operatorID, roomID, content, extraJSON string, memberIDs []string) *model.ChatMessage {
+ notifMsg := model.ChatMessage{
+ RoomID: roomID,
+ SenderUserID: operatorID,
+ ReceiverUserID: roomID,
+ MessageType: model.MessageTypeGroupNotif,
+ Content: content,
+ Extra: extraJSON,
+ }
+
+ // 房间级仅落库一次;失败则不推送
+ if err := s.DB.Create(¬ifMsg).Error; err != nil {
+ log.Printf("❌ 群通知落库失败 (room: %s): %v", roomID, err)
+ return nil
+ }
+
+ pushMsg := model.WsPayload{
+ RequestType: "receive_message",
+ Data: notifMsg,
+ }
+ msgBytes, _ := json.Marshal(pushMsg)
+
+ for _, uid := range memberIDs {
+ if uid == "" {
+ continue
+ }
+ // 更新每个成员的会话摘要;操作者本人视为发送方,不增加其未读数
+ if ConversationSvc != nil {
+ _ = ConversationSvc.UpsertConversationOnMessage(uid, roomID, roomID, notifMsg, uid == operatorID)
+ }
+ s.DispatchMessage(uid, msgBytes)
+ }
+ return ¬ifMsg
+}
+
/**
* NotifyOtherDevices
* 功能:通知同一用户的其他设备“电话已被接听”,实现多端协同。
@@ -389,8 +867,8 @@ func (s *ChatService) NotifyOtherDevices(userID, currentClientID, callID string)
cancelMsg := model.WsPayload{
RequestType: "receive_message",
Data: model.ChatMessage{
- MessageType: 6, // 视频通话类型
- CallStatus: "answered_elsewhere", // 特殊状态码
+ MessageType: model.MessageTypeSignal, // 信令消息
+ CallStatus: "answered_elsewhere", // 特殊状态码
CallID: callID,
Content: "通话已在其他设备接听",
},
@@ -492,11 +970,8 @@ func (s *ChatService) SubscribeClusterMessages() {
continue
}
- // 修正:Payload 在 JSON 传输后变成了 interface{},需要重新序列化才能发给 WebSocket
- payloadBytes, _ := json.Marshal(clusterMsg.Payload)
-
- // 尝试发送给本地用户
- manager.Manager.SendToUser(clusterMsg.TargetUserID, payloadBytes)
+ // Payload 为 json.RawMessage(原始 JSON 字节),直接投递,无需二次序列化
+ manager.Manager.SendToUser(clusterMsg.TargetUserID, clusterMsg.Payload)
}
}
diff --git a/internal/service/contact_service.go b/internal/service/contact_service.go
index c70d84a..3a3d72e 100644
--- a/internal/service/contact_service.go
+++ b/internal/service/contact_service.go
@@ -7,10 +7,9 @@ package service
import (
"encoding/json"
"errors"
- "fmt"
+ "strings"
"time"
"xk-websocket-v2/internal/model"
- "xk-websocket-v2/internal/utils"
"gorm.io/gorm"
)
@@ -44,6 +43,27 @@ func (s *ContactService) SearchUsers(keyword string, limit int) ([]model.User, e
* 功能:发送好友申请
*/
func (s *ContactService) AddFriend(fromUserID, toUserID, message string) error {
+ // 边界校验1:不允许添加自己为好友
+ // 否则会产生"自己是自己好友"的脏数据(申请通过后 UserContact 双向落库)
+ if fromUserID == toUserID {
+ return errors.New("不能添加自己为好友")
+ }
+
+ // AI 机器人是全局固定联系人(好友列表自动展示),无需也不允许发好友申请
+ if strings.HasPrefix(toUserID, "bot_") {
+ return errors.New("AI机器人已自动出现在你的好友列表,无需添加")
+ }
+
+ // 边界校验2:目标用户必须真实存在
+ // 否则可对任意/不存在的ID落一条 pending 申请,形成垃圾数据
+ var targetCount int64
+ if err := s.DB.Model(&model.User{}).Where("id = ?", toUserID).Count(&targetCount).Error; err != nil {
+ return err
+ }
+ if targetCount == 0 {
+ return errors.New("目标用户不存在")
+ }
+
// 检查是否已经是好友
var existingContact model.UserContact
result := s.DB.Where("user_id = ? AND contact_id = ?", fromUserID, toUserID).First(&existingContact)
@@ -116,88 +136,75 @@ func (s *ContactService) AcceptFriendRequest(requestID uint, userID string) erro
return err
}
- // 使用雪花ID生成唯一的房间ID
- roomID, err := utils.NextIDString()
- if err != nil {
- tx.Rollback()
- return fmt.Errorf("生成房间ID失败: %v", err)
+ // 房间ID统一用 GenerateP2PRoomID(两个用户ID按序拼接)。
+ // 历史Bug:这里曾用随机雪花ID建房,而消息发送路径的 GetOrCreateP2PRoom 用的是
+ // userA_userB 规则,同一对好友产生两套房间号,消息与会话被割裂到两个房间里。
+ roomID := GenerateP2PRoomID(request.FromUserID, request.ToUserID)
+
+ // 创建房间(p2p类型)。若房间已存在(如曾经是好友删除后重加、或之前直接聊过)则复用
+ var existingRoom model.ChatRoom
+ if err := tx.Where("room_id = ?", roomID).First(&existingRoom).Error; err != nil {
+ if err != gorm.ErrRecordNotFound {
+ tx.Rollback()
+ return err
+ }
+ room := model.ChatRoom{
+ RoomID: roomID,
+ RoomType: "p2p",
+ OwnerID: "0", // 单聊群主为0
+ CreatorID: request.ToUserID, // 接受者为创建者
+ }
+ if err := tx.Create(&room).Error; err != nil {
+ tx.Rollback()
+ return err
+ }
}
- // 创建房间(p2p类型)
- room := model.ChatRoom{
- RoomID: roomID,
- RoomType: "p2p",
- OwnerID: "0", // 单聊群主为0
- CreatorID: request.ToUserID, // 接受者为创建者
- }
- if err := tx.Create(&room).Error; err != nil {
- tx.Rollback()
- return err
+ // 创建房间成员记录(两个用户),FirstOrCreate 幂等:已存在则跳过
+ for _, uid := range []string{request.FromUserID, request.ToUserID} {
+ member := model.RoomMember{RoomID: roomID, UserID: uid}
+ if err := tx.Where("room_id = ? AND user_id = ?", roomID, uid).
+ FirstOrCreate(&member).Error; err != nil {
+ tx.Rollback()
+ return err
+ }
}
- // 创建房间成员记录(两个用户)
- member1 := model.RoomMember{
- RoomID: roomID,
- UserID: request.FromUserID,
- }
- member2 := model.RoomMember{
- RoomID: roomID,
- UserID: request.ToUserID,
- }
- if err := tx.Create(&member1).Error; err != nil {
- tx.Rollback()
- return err
- }
- if err := tx.Create(&member2).Error; err != nil {
- tx.Rollback()
- return err
+ // 创建双向好友关系并绑定房间ID,FirstOrCreate 幂等防重复好友记录
+ for _, pair := range [][2]string{
+ {request.FromUserID, request.ToUserID},
+ {request.ToUserID, request.FromUserID},
+ } {
+ contact := model.UserContact{
+ UserID: pair[0],
+ ContactID: pair[1],
+ RoomID: roomID,
+ }
+ if err := tx.Where("user_id = ? AND contact_id = ?", pair[0], pair[1]).
+ FirstOrCreate(&contact).Error; err != nil {
+ tx.Rollback()
+ return err
+ }
}
- // 创建双向好友关系,并绑定房间ID
- contact1 := model.UserContact{
- UserID: request.FromUserID,
- ContactID: request.ToUserID,
- RoomID: roomID,
- }
- contact2 := model.UserContact{
- UserID: request.ToUserID,
- ContactID: request.FromUserID,
- RoomID: roomID,
- }
-
- if err := tx.Create(&contact1).Error; err != nil {
- tx.Rollback()
- return err
- }
- if err := tx.Create(&contact2).Error; err != nil {
- tx.Rollback()
- return err
- }
-
- // 为双方创建会话记录
+ // 为双方创建会话记录,FirstOrCreate 幂等:曾聊过时会话已存在,避免重复插入
now := time.Now()
- conv1 := model.ChatConversation{
- UserID: request.FromUserID,
- TargetID: request.ToUserID,
- RoomID: roomID,
- Type: 1, // 私聊
- LastTime: now,
- }
- conv2 := model.ChatConversation{
- UserID: request.ToUserID,
- TargetID: request.FromUserID,
- RoomID: roomID,
- Type: 1, // 私聊
- LastTime: now,
- }
-
- if err := tx.Create(&conv1).Error; err != nil {
- tx.Rollback()
- return err
- }
- if err := tx.Create(&conv2).Error; err != nil {
- tx.Rollback()
- return err
+ for _, pair := range [][2]string{
+ {request.FromUserID, request.ToUserID},
+ {request.ToUserID, request.FromUserID},
+ } {
+ conv := model.ChatConversation{
+ UserID: pair[0],
+ TargetID: pair[1],
+ RoomID: roomID,
+ Type: model.ConversationTypeP2P,
+ LastTime: now,
+ }
+ if err := tx.Where("user_id = ? AND target_id = ? AND type = 1", pair[0], pair[1]).
+ FirstOrCreate(&conv).Error; err != nil {
+ tx.Rollback()
+ return err
+ }
}
// 提交事务
@@ -221,7 +228,7 @@ func (s *ContactService) sendWelcomeMessage(senderID, receiverID, roomID string)
RoomID: roomID,
SenderUserID: senderID,
ReceiverUserID: receiverID,
- MessageType: 0, // 文本消息
+ MessageType: model.MessageTypeText,
Content: "我已经通过了你的好友申请,开始和我聊天吧~",
}
@@ -303,8 +310,24 @@ func (s *ContactService) UpdateContact(userID, contactID string, updates map[str
* 功能:删除好友
*/
func (s *ContactService) DeleteContact(userID, contactID string) error {
- // 删除双向好友关系
+ // AI 机器人是全局固定联系人,任何人都不能从好友列表删除(需求约束);
+ // 机器人管理走管理员专用的 /ai/bots/delete 接口
+ if strings.HasPrefix(contactID, "bot_") {
+ return errors.New("AI机器人无法删除")
+ }
+ // 删除双向好友关系(事务内一并清理双方会话,避免残留"幽灵会话"还能继续发消息)
tx := s.DB.Begin()
+ defer func() {
+ if r := recover(); r != nil {
+ tx.Rollback()
+ }
+ }()
+
+ // 删除前先取好友关系行上存储的 RoomID:历史版本曾用随机雪花ID给 P2P 房间命名,
+ // 与按 GenerateP2PRoomID 规则重算的标准房间ID可能不同,两者都要纳入成员清理范围
+ var contactRow model.UserContact
+ hasContactRow := tx.Where("user_id = ? AND contact_id = ?", userID, contactID).
+ First(&contactRow).Error == nil
if err := tx.Where("user_id = ? AND contact_id = ?", userID, contactID).Delete(&model.UserContact{}).Error; err != nil {
tx.Rollback()
@@ -316,6 +339,30 @@ func (s *ContactService) DeleteContact(userID, contactID string) error {
return err
}
+ // 清理双方的私聊会话记录:好友关系是双向删除的,会话也应双向清理,
+ // 否则任一方点开残留会话仍可向已删除的"好友"发消息
+ if err := tx.Where("(user_id = ? AND target_id = ?) OR (user_id = ? AND target_id = ?)",
+ userID, contactID, contactID, userID).
+ Where("type = 1").
+ Delete(&model.ChatConversation{}).Error; err != nil {
+ tx.Rollback()
+ return err
+ }
+
+ // 撤销双方在私聊房间的成员资格:好友关系与会话删了但 room_members 残留的话,
+ // 发送路径针对旧式随机ID房间的 IsRoomMember 兜底仍会放行,删好友后照样能互发消息。
+ // 标准房间ID按 GenerateP2PRoomID 规则重算,旧式房间ID从好友关系行上取;
+ // 重新加好友时 AcceptFriendRequest 用 FirstOrCreate 幂等恢复成员记录,这里删除是安全的
+ roomIDs := []string{GenerateP2PRoomID(userID, contactID)}
+ if hasContactRow && contactRow.RoomID != "" && contactRow.RoomID != roomIDs[0] {
+ roomIDs = append(roomIDs, contactRow.RoomID)
+ }
+ if err := tx.Where("room_id IN ? AND user_id IN ?", roomIDs, []string{userID, contactID}).
+ Delete(&model.RoomMember{}).Error; err != nil {
+ tx.Rollback()
+ return err
+ }
+
return tx.Commit().Error
}
@@ -377,15 +424,30 @@ func (s *ContactService) DeleteGroup(groupID uint, userID string) error {
return err
}
+ // 事务保证"移动联系人到默认分组"与"删除分组"要么都成功要么都不做,
+ // 否则中途失败会出现联系人已被移出、分组却还在(或反之)的脏状态
+ tx := s.DB.Begin()
+ defer func() {
+ if r := recover(); r != nil {
+ tx.Rollback()
+ }
+ }()
+
// 将该分组下的联系人移到默认分组(group_id = 0)
- if err := s.DB.Model(&model.UserContact{}).
+ if err := tx.Model(&model.UserContact{}).
Where("user_id = ? AND group_id = ?", userID, groupID).
Update("group_id", 0).Error; err != nil {
+ tx.Rollback()
return err
}
// 删除分组
- return s.DB.Delete(&group).Error
+ if err := tx.Delete(&group).Error; err != nil {
+ tx.Rollback()
+ return err
+ }
+
+ return tx.Commit().Error
}
/**
@@ -400,14 +462,31 @@ func (s *ContactService) GetContactsWithUserInfo(userID string) ([]map[string]in
return nil, err
}
+ // 批量查询联系人用户信息:一条 IN 查询代替每个联系人单查一次(N+1),
+ // 好友数量多时显著减少数据库往返
+ contactIDs := make([]string, 0, len(contacts))
+ for _, contact := range contacts {
+ contactIDs = append(contactIDs, contact.ContactID)
+ }
+ userMap := make(map[string]model.User, len(contactIDs))
+ if len(contactIDs) > 0 {
+ var users []model.User
+ if err := s.DB.Where("id IN ?", contactIDs).Find(&users).Error; err != nil {
+ return nil, err
+ }
+ for _, u := range users {
+ u.Password = ""
+ userMap[u.ID] = u
+ }
+ }
+
var result []map[string]interface{}
for _, contact := range contacts {
- // 获取联系人用户信息
- var user model.User
- if err := s.DB.Where("id = ?", contact.ContactID).First(&user).Error; err != nil {
+ // 用户记录不存在(账号已注销等)时跳过,与原实现行为一致
+ user, ok := userMap[contact.ContactID]
+ if !ok {
continue
}
- user.Password = ""
// 组合数据,包含 user 对象
item := map[string]interface{}{
@@ -431,6 +510,51 @@ func (s *ContactService) GetContactsWithUserInfo(userID string) ([]map[string]in
result = append(result, item)
}
+ // 追加 AI 机器人虚拟联系人:机器人不写 user_contacts 表(好友关系不可删,
+ // 用真实记录会被删好友/清库脚本误伤),改为查询时动态拼接。
+ // room_id 按统一 P2P 规则生成,消息发送校验天然放行,点开即聊
+ if AIBotSvc != nil {
+ if bots, botErr := AIBotSvc.ListBots(true); botErr == nil && len(bots) > 0 {
+ // 已有联系人(异常情况下 bot 被加成真实好友)去重
+ existing := make(map[string]bool, len(contacts))
+ for _, contact := range contacts {
+ existing[contact.ContactID] = true
+ }
+ for _, bot := range bots {
+ if existing[bot.UserID] {
+ continue
+ }
+ botUser := model.User{
+ ID: bot.UserID,
+ Name: bot.Name,
+ Avatar: bot.Avatar,
+ Desc: "AI机器人",
+ IsBot: true,
+ }
+ result = append(result, map[string]interface{}{
+ "id": bot.UserID,
+ "user_id": bot.UserID,
+ "contact_user_id": bot.UserID,
+ "user": botUser,
+ "remark_name": "",
+ "room_id": GenerateP2PRoomID(userID, bot.UserID),
+ "group_id": 0,
+ "is_top": false,
+ "is_muted": false,
+ "is_special_care": false,
+ "is_blocked": false,
+ "last_chat_time": nil,
+ "last_message": "",
+ "last_msg": "",
+ "unread_count": 0,
+ "unread": 0,
+ // 前端据此隐藏"删除好友/拉黑"入口并展示机器人标识
+ "is_bot": true,
+ })
+ }
+ }
+ }
+
return result, nil
}
@@ -523,6 +647,17 @@ func (s *ContactService) GetUserDetailWithFriendStatus(currentUserID, targetUser
isFriend = true
}
+ // AI 机器人:没有真实 user_contacts 记录,但按需求它是"必定存在的好友",
+ // 构造虚拟关系让详情页直接走好友分支(备注/置顶等本地开关对它无意义,前端会隐藏)
+ if !isFriend && user.IsBot {
+ isFriend = true
+ contact = &model.UserContact{
+ UserID: currentUserID,
+ ContactID: targetUserID,
+ RoomID: GenerateP2PRoomID(currentUserID, targetUserID),
+ }
+ }
+
return map[string]interface{}{
"is_friend": isFriend,
"contact": contact,
@@ -575,3 +710,16 @@ func (s *ContactService) IsBlocked(ownerID, targetID string) bool {
First(&contact).Error
return err == nil
}
+
+/**
+ * IsFriend
+ * 功能:检查 owner 的好友列表里是否仍存在 target(user_contacts 存在 owner→target 行)
+ * 说明:好友关系在接受申请时双向创建、删除好友时双向删除,正常数据下查任一方向即可判定;
+ * 私聊发送路径用它做实时好友校验,堵住"删除好友后房间残留仍可互发消息"的漏洞
+ */
+func (s *ContactService) IsFriend(ownerID, targetID string) bool {
+ var contact model.UserContact
+ err := s.DB.Where("user_id = ? AND contact_id = ?", ownerID, targetID).
+ First(&contact).Error
+ return err == nil
+}
diff --git a/internal/service/conversation_service.go b/internal/service/conversation_service.go
index 36b0df9..2f6768b 100644
--- a/internal/service/conversation_service.go
+++ b/internal/service/conversation_service.go
@@ -38,25 +38,55 @@ func (s *ConversationService) GetConversations(userID string) ([]model.ChatConve
if err != nil {
return list, err
}
-
- // 手动加载关联信息
+
+ // 批量加载关联信息:先收集私聊目标用户ID与群聊房间ID,
+ // 各用一条 IN 查询取回后按ID回填。原实现每个会话单独查一次库(N+1),
+ // 会话越多接口越慢
+ userIDs := make([]string, 0, len(list))
+ roomIDs := make([]string, 0, len(list))
for i := range list {
- if list[i].Type == 1 { // 私聊:加载目标用户信息
- var user model.User
- if err := s.DB.Where("id = ?", list[i].TargetID).First(&user).Error; err == nil {
- list[i].TargetUser = &user
- }
- } else if list[i].Type == 2 { // 群聊:加载群信息
- var room model.ChatRoom
- if err := s.DB.Where("room_id = ?", list[i].RoomID).First(&room).Error; err == nil {
- list[i].Room = &room
- } else {
- // 记录警告:群聊信息查询失败
- log.Printf("⚠️ 群聊信息查询失败 (room_id: %s): %v", list[i].RoomID, err)
- }
+ if list[i].Type == model.ConversationTypeP2P {
+ userIDs = append(userIDs, list[i].TargetID)
+ } else if list[i].Type == model.ConversationTypeGroup {
+ roomIDs = append(roomIDs, list[i].RoomID)
}
}
-
+
+ // 私聊目标用户批量查询(Password 字段 json:"-" 不会随响应序列化)
+ userMap := make(map[string]*model.User, len(userIDs))
+ if len(userIDs) > 0 {
+ var users []model.User
+ if err := s.DB.Where("id IN ?", userIDs).Find(&users).Error; err == nil {
+ for i := range users {
+ userMap[users[i].ID] = &users[i]
+ }
+ } else {
+ log.Printf("⚠️ 会话目标用户批量查询失败: %v", err)
+ }
+ }
+
+ // 群聊房间批量查询
+ roomMap := make(map[string]*model.ChatRoom, len(roomIDs))
+ if len(roomIDs) > 0 {
+ var rooms []model.ChatRoom
+ if err := s.DB.Where("room_id IN ?", roomIDs).Find(&rooms).Error; err == nil {
+ for i := range rooms {
+ roomMap[rooms[i].RoomID] = &rooms[i]
+ }
+ } else {
+ log.Printf("⚠️ 会话群聊信息批量查询失败: %v", err)
+ }
+ }
+
+ // 回填关联信息
+ for i := range list {
+ if list[i].Type == model.ConversationTypeP2P {
+ list[i].TargetUser = userMap[list[i].TargetID]
+ } else if list[i].Type == model.ConversationTypeGroup {
+ list[i].Room = roomMap[list[i].RoomID]
+ }
+ }
+
return list, nil
}
@@ -75,21 +105,21 @@ func (s *ConversationService) UpsertConversationOnMessage(userID, targetID, room
}
// 判断是否为群聊:优先通过查询数据库确认房间类型,其次通过 roomID 格式判断
- conversationType := 1 // 默认私聊
+ conversationType := model.ConversationTypeP2P // 默认私聊
if roomID != "" {
// 情况1:targetID == roomID,说明是群聊(群聊时 targetID 就是 roomID)
if targetID == roomID {
- conversationType = 2 // 群聊
+ conversationType = model.ConversationTypeGroup
} else {
// 情况2:查询数据库确认房间类型(用于兼容旧数据或特殊情况)
var room model.ChatRoom
if err := s.DB.Where("room_id = ?", roomID).First(&room).Error; err == nil {
if room.RoomType == "group" {
- conversationType = 2 // 群聊
+ conversationType = model.ConversationTypeGroup
}
} else if len(roomID) > 6 && roomID[:6] == "group_" {
// 如果查询失败,回退到格式判断
- conversationType = 2 // 群聊
+ conversationType = model.ConversationTypeGroup
}
}
}
@@ -99,7 +129,7 @@ func (s *ConversationService) UpsertConversationOnMessage(userID, targetID, room
now := time.Now()
// 计算摘要:群聊时包含发送者信息
- summary := buildMessageSummary(msg, conversationType == 2, msg.SenderUserID)
+ summary := buildMessageSummary(msg, conversationType == model.ConversationTypeGroup, msg.SenderUserID)
if tx.Error != nil {
if tx.Error == gorm.ErrRecordNotFound {
@@ -116,15 +146,10 @@ func (s *ConversationService) UpsertConversationOnMessage(userID, targetID, room
if !isSender {
conv.UnreadCount = 1
}
-
- // 群聊时,target_id 是群ID,不是用户ID,需要临时禁用外键检查
- if conversationType == 2 {
- s.DB.Exec("SET FOREIGN_KEY_CHECKS = 0")
- err := s.DB.Create(&conv).Error
- s.DB.Exec("SET FOREIGN_KEY_CHECKS = 1")
- return err
- }
-
+
+ // 说明:chat_conversations 表没有外键约束(模型中 TargetUser 已用 gorm:"-" 禁用),
+ // 历史遗留的 SET FOREIGN_KEY_CHECKS=0 开关已删除——在连接池下 Exec 可能跑在
+ // 与 Create 不同的连接上,既不生效还会把"关闭外键检查"的状态泄漏给其他请求
return s.DB.Create(&conv).Error
}
return tx.Error
@@ -137,7 +162,9 @@ func (s *ConversationService) UpsertConversationOnMessage(userID, targetID, room
"last_time": now,
}
if !isSender {
- updates["unread_count"] = conv.UnreadCount + 1
+ // 原子自增:并发收多条消息时 conv.UnreadCount+1 会互相覆盖丢未读,
+ // 改用 SQL 表达式在数据库侧自增,保证并发安全
+ updates["unread_count"] = gorm.Expr("unread_count + 1")
}
return s.DB.Model(&model.ChatConversation{}).
@@ -205,16 +232,13 @@ func (s *ConversationService) EnsureGroupConversations(roomID string) error {
// 检查是否已有会话
var existingConv model.ChatConversation
err := s.DB.Where("user_id = ? AND target_id = ? AND type = ?", memberID, roomID, 2).First(&existingConv).Error
-
+
if err != nil {
if err == gorm.ErrRecordNotFound {
- // 不存在,创建新会话
- // 临时禁用外键检查,因为 target_id 是群ID,不是用户ID
- s.DB.Exec("SET FOREIGN_KEY_CHECKS = 0")
+ // 不存在,创建新会话(表无外键约束,无需再动 FOREIGN_KEY_CHECKS 开关)
insertSQL := `INSERT INTO chat_conversations (user_id, target_id, room_id, type, is_top, is_muted, is_special_care, unread_count, last_message, last_time, created_at, updated_at)
VALUES (?, ?, ?, 2, false, false, false, 0, '', ?, ?, ?)`
createErr := s.DB.Exec(insertSQL, memberID, roomID, roomID, now, now, now).Error
- s.DB.Exec("SET FOREIGN_KEY_CHECKS = 1")
if createErr != nil {
log.Printf("❌ 创建群聊会话失败 (userID: %s, roomID: %s): %v", memberID, roomID, createErr)
}
@@ -284,21 +308,21 @@ func (s *ConversationService) GetOrCreateConversationByRoom(userID, roomID strin
// 先查询房间信息(必须存在才能创建会话)
var room model.ChatRoom
roomErr := s.DB.Where("room_id = ?", roomID).First(&room).Error
-
+
if roomErr != nil {
// 房间不存在,返回错误
return nil, gorm.ErrRecordNotFound
}
-
+
var conversationType int
var targetID string
-
+
// 根据房间类型判断会话类型
if room.RoomType == "group" {
- conversationType = 2 // 群聊
- targetID = roomID // 群聊时 targetID = roomID
+ conversationType = model.ConversationTypeGroup
+ targetID = roomID // 群聊时 targetID = roomID
} else {
- conversationType = 1 // 私聊
+ conversationType = model.ConversationTypeP2P
// 私聊需要确定 targetID(对方用户ID)
// 从 room_members 中查找另一个成员
var members []model.RoomMember
@@ -318,7 +342,7 @@ func (s *ConversationService) GetOrCreateConversationByRoom(userID, roomID strin
// 查询会话
var conv model.ChatConversation
tx := s.DB.Where("user_id = ? AND target_id = ? AND type = ?", userID, targetID, conversationType).First(&conv)
-
+
if tx.Error != nil {
if tx.Error == gorm.ErrRecordNotFound {
// 会话不存在,根据房间信息创建新会话
@@ -331,19 +355,10 @@ func (s *ConversationService) GetOrCreateConversationByRoom(userID, roomID strin
LastTime: time.Now(),
UnreadCount: 0,
}
-
- // 群聊时,target_id 是群ID,不是用户ID,需要临时禁用外键检查
- if conversationType == 2 {
- s.DB.Exec("SET FOREIGN_KEY_CHECKS = 0")
- err := s.DB.Create(&conv).Error
- s.DB.Exec("SET FOREIGN_KEY_CHECKS = 1")
- if err != nil {
- return nil, err
- }
- } else {
- if err := s.DB.Create(&conv).Error; err != nil {
- return nil, err
- }
+
+ // 直接创建(表无外键约束,无需再动 FOREIGN_KEY_CHECKS 开关)
+ if err := s.DB.Create(&conv).Error; err != nil {
+ return nil, err
}
} else {
return nil, tx.Error
@@ -351,18 +366,16 @@ func (s *ConversationService) GetOrCreateConversationByRoom(userID, roomID strin
}
// 加载关联信息
- if conv.Type == 1 { // 私聊:加载目标用户信息
+ if conv.Type == model.ConversationTypeP2P { // 私聊:加载目标用户信息
var user model.User
if err := s.DB.Where("id = ?", conv.TargetID).First(&user).Error; err == nil {
conv.TargetUser = &user
} else {
log.Printf("⚠️ 用户信息查询失败 (user_id: %s): %v", conv.TargetID, err)
}
- } else if conv.Type == 2 { // 群聊:加载群信息
+ } else if conv.Type == model.ConversationTypeGroup { // 群聊:加载群信息
conv.Room = &room
}
return &conv, nil
}
-
-
diff --git a/internal/service/moment_service.go b/internal/service/moment_service.go
index 781ee56..562fd6b 100644
--- a/internal/service/moment_service.go
+++ b/internal/service/moment_service.go
@@ -10,10 +10,14 @@ import (
"log"
"time"
"xk-websocket-v2/internal/model"
+ "xk-websocket-v2/internal/utils"
"gorm.io/gorm"
)
+// 每条动态在列表中最多附带的点赞数(完整点赞列表走 GetMomentLikes 单独接口)
+const maxLikesPerMoment = 10
+
// MomentService 朋友圈服务结构体
type MomentService struct {
DB *gorm.DB
@@ -147,16 +151,14 @@ func (s *MomentService) GetFriendsMoments(userID string, page, pageSize int) ([]
return nil, 0, err
}
- // 过滤可见性并加载关联数据
+ // 过滤可见性后批量加载关联数据(点赞/评论/用户信息各一条 IN 查询,避免 N+1)
var visibleMoments []model.Moment
for i := range moments {
if s.checkVisibility(&moments[i], userID) {
- s.loadMomentUser(&moments[i])
- s.loadMomentLikes(&moments[i], userID)
- s.loadMomentComments(&moments[i])
visibleMoments = append(visibleMoments, moments[i])
}
}
+ s.loadMomentsRelations(visibleMoments, userID)
return visibleMoments, total, nil
}
@@ -189,16 +191,14 @@ func (s *MomentService) GetUserMoments(viewerID, targetUserID string, page, page
return nil, 0, err
}
- // 过滤可见性并加载关联数据
+ // 过滤可见性后批量加载关联数据(点赞/评论/用户信息各一条 IN 查询,避免 N+1)
var visibleMoments []model.Moment
for i := range moments {
if s.checkVisibility(&moments[i], viewerID) {
- s.loadMomentUser(&moments[i])
- s.loadMomentLikes(&moments[i], viewerID)
- s.loadMomentComments(&moments[i])
visibleMoments = append(visibleMoments, moments[i])
}
}
+ s.loadMomentsRelations(visibleMoments, viewerID)
return visibleMoments, total, nil
}
@@ -210,6 +210,9 @@ func (s *MomentService) GetUserMoments(viewerID, targetUserID string, page, page
/**
* LikeMoment
* 功能:点赞动态
+ * 并发安全:moment_likes 有 (moment_id, user_id) 唯一索引,
+ * "先查后插"竞态下第二个请求插入会触发唯一冲突,按"已点赞"幂等返回,
+ * 避免重复记录和 like_count 多加
*/
func (s *MomentService) LikeMoment(userID string, momentID uint) error {
// 检查动态是否存在
@@ -218,7 +221,7 @@ func (s *MomentService) LikeMoment(userID string, momentID uint) error {
return errors.New("动态不存在")
}
- // 检查是否已点赞
+ // 检查是否已点赞(快速路径,减少无谓的事务开销)
var existingLike model.MomentLike
if err := s.DB.Where("moment_id = ? AND user_id = ?", momentID, userID).First(&existingLike).Error; err == nil {
return errors.New("已点赞过该动态")
@@ -234,6 +237,10 @@ func (s *MomentService) LikeMoment(userID string, momentID uint) error {
}
if err := tx.Create(&like).Error; err != nil {
tx.Rollback()
+ // 唯一索引冲突说明并发点赞已成功一次,幂等处理
+ if utils.IsDuplicateEntryError(err) {
+ return errors.New("已点赞过该动态")
+ }
return err
}
@@ -244,7 +251,10 @@ func (s *MomentService) LikeMoment(userID string, momentID uint) error {
return err
}
- tx.Commit()
+ // Commit 失败必须返回错误且不发通知:否则点赞实际没落库却通知了对方
+ if err := tx.Commit().Error; err != nil {
+ return err
+ }
// 发送通知(异步)
if moment.UserID != userID {
@@ -261,16 +271,21 @@ func (s *MomentService) LikeMoment(userID string, momentID uint) error {
func (s *MomentService) UnlikeMoment(userID string, momentID uint) error {
tx := s.DB.Begin()
- // 删除点赞记录
+ // 删除点赞记录(先判错误再判影响行数,删除报错时不应误报"未点赞过")
result := tx.Where("moment_id = ? AND user_id = ?", momentID, userID).Delete(&model.MomentLike{})
+ if result.Error != nil {
+ tx.Rollback()
+ return result.Error
+ }
if result.RowsAffected == 0 {
tx.Rollback()
return errors.New("未点赞过该动态")
}
- // 更新动态点赞数
+ // 更新动态点赞数。GREATEST 兜底防止减到负数:
+ // 历史脏数据(重复点赞记录已清理但计数未修正等)可能导致计数与记录不一致
if err := tx.Model(&model.Moment{}).Where("id = ?", momentID).
- UpdateColumn("like_count", gorm.Expr("like_count - 1")).Error; err != nil {
+ UpdateColumn("like_count", gorm.Expr("GREATEST(like_count - 1, 0)")).Error; err != nil {
tx.Rollback()
return err
}
@@ -290,12 +305,22 @@ func (s *MomentService) GetMomentLikes(momentID uint) ([]model.MomentLike, error
return nil, err
}
- // 加载用户信息
- for i := range likes {
- var user model.User
- if err := s.DB.Where("id = ?", likes[i].UserID).First(&user).Error; err == nil {
- user.Password = ""
- likes[i].User = &user
+ // 批量加载点赞用户信息:一条 IN 查询代替逐条查询(N+1)
+ if len(likes) > 0 {
+ userIDs := make([]string, 0, len(likes))
+ for i := range likes {
+ userIDs = append(userIDs, likes[i].UserID)
+ }
+ var users []model.User
+ if err := s.DB.Where("id IN ?", userIDs).Find(&users).Error; err == nil {
+ userMap := make(map[string]*model.User, len(users))
+ for i := range users {
+ users[i].Password = ""
+ userMap[users[i].ID] = &users[i]
+ }
+ for i := range likes {
+ likes[i].User = userMap[likes[i].UserID]
+ }
}
}
@@ -350,7 +375,10 @@ func (s *MomentService) CreateComment(userID string, momentID uint, req *model.C
return nil, err
}
- tx.Commit()
+ // Commit 失败必须返回错误且不发通知:否则评论实际没落库却通知了对方
+ if err := tx.Commit().Error; err != nil {
+ return nil, err
+ }
// 发送通知(异步)
go func() {
@@ -457,14 +485,86 @@ func (s *MomentService) GetNotifications(userID string, page, pageSize int) ([]m
return nil, 0, err
}
- // 加载关联数据
- for i := range notifications {
- s.loadNotificationData(¬ifications[i])
- }
+ // 批量加载关联数据(发送者/动态/评论各一条 IN 查询,避免逐条 N+1)
+ s.loadNotificationsRelations(notifications)
return notifications, total, nil
}
+/**
+ * loadNotificationsRelations
+ * 功能:批量加载通知列表的关联数据(发送者、动态、评论)。
+ * 为什么需要:原实现每条通知分别查发送者/动态/评论(每条通知3次查询),
+ * 改为按整页收集ID后各一条 IN 查询回填
+ */
+func (s *MomentService) loadNotificationsRelations(notifications []model.MomentNotification) {
+ if len(notifications) == 0 {
+ return
+ }
+
+ userIDSet := make(map[string]struct{}, len(notifications))
+ momentIDSet := make(map[uint]struct{}, len(notifications))
+ commentIDs := make([]uint, 0, len(notifications))
+ for i := range notifications {
+ userIDSet[notifications[i].FromUserID] = struct{}{}
+ momentIDSet[notifications[i].MomentID] = struct{}{}
+ if notifications[i].CommentID != nil {
+ commentIDs = append(commentIDs, *notifications[i].CommentID)
+ }
+ }
+
+ // 批量查询发送者
+ userIDs := make([]string, 0, len(userIDSet))
+ for id := range userIDSet {
+ userIDs = append(userIDs, id)
+ }
+ userMap := make(map[string]*model.User, len(userIDs))
+ if len(userIDs) > 0 {
+ var users []model.User
+ if err := s.DB.Where("id IN ?", userIDs).Find(&users).Error; err == nil {
+ for i := range users {
+ users[i].Password = ""
+ userMap[users[i].ID] = &users[i]
+ }
+ }
+ }
+
+ // 批量查询动态
+ momentIDs := make([]uint, 0, len(momentIDSet))
+ for id := range momentIDSet {
+ momentIDs = append(momentIDs, id)
+ }
+ momentMap := make(map[uint]*model.Moment, len(momentIDs))
+ if len(momentIDs) > 0 {
+ var moments []model.Moment
+ if err := s.DB.Where("id IN ?", momentIDs).Find(&moments).Error; err == nil {
+ for i := range moments {
+ momentMap[moments[i].ID] = &moments[i]
+ }
+ }
+ }
+
+ // 批量查询评论
+ commentMap := make(map[uint]*model.MomentComment, len(commentIDs))
+ if len(commentIDs) > 0 {
+ var comments []model.MomentComment
+ if err := s.DB.Where("id IN ?", commentIDs).Find(&comments).Error; err == nil {
+ for i := range comments {
+ commentMap[comments[i].ID] = &comments[i]
+ }
+ }
+ }
+
+ // 回填
+ for i := range notifications {
+ notifications[i].FromUser = userMap[notifications[i].FromUserID]
+ notifications[i].Moment = momentMap[notifications[i].MomentID]
+ if notifications[i].CommentID != nil {
+ notifications[i].Comment = commentMap[*notifications[i].CommentID]
+ }
+ }
+}
+
/**
* MarkNotificationsRead
* 功能:标记通知为已读
@@ -577,6 +677,104 @@ func (s *MomentService) getFriendIDs(userID string) []string {
return ids
}
+/**
+ * loadMomentsRelations
+ * 功能:批量加载一页动态的关联数据(发布者、点赞、评论及涉及的用户信息)。
+ * 为什么需要:原实现每条动态分别查发布者、点赞、评论,点赞和评论又逐条查用户,
+ * 一页 20 条动态可能产生上百次数据库往返(N+1)。这里按整页收集 ID 后,
+ * 点赞/评论/用户各用一条 IN 查询取回,再在内存中分组回填,
+ * 数据库往返次数固定为 3 次,与页大小无关
+ */
+func (s *MomentService) loadMomentsRelations(moments []model.Moment, viewerID string) {
+ if len(moments) == 0 {
+ return
+ }
+
+ momentIDs := make([]uint, 0, len(moments))
+ for i := range moments {
+ momentIDs = append(momentIDs, moments[i].ID)
+ }
+
+ // 1. 批量查询点赞:按动态分组,每条动态截取最近 maxLikesPerMoment 条;
+ // 同时顺带得出当前用户对哪些动态点过赞(无需再逐条 Count)
+ var likes []model.MomentLike
+ s.DB.Where("moment_id IN ?", momentIDs).
+ Order("created_at DESC").
+ Find(&likes)
+ likesByMoment := make(map[uint][]model.MomentLike, len(momentIDs))
+ likedByViewer := make(map[uint]bool)
+ for _, like := range likes {
+ if viewerID != "" && like.UserID == viewerID {
+ likedByViewer[like.MomentID] = true
+ }
+ if len(likesByMoment[like.MomentID]) < maxLikesPerMoment {
+ likesByMoment[like.MomentID] = append(likesByMoment[like.MomentID], like)
+ }
+ }
+
+ // 2. 批量查询评论并按动态分组
+ var comments []model.MomentComment
+ s.DB.Where("moment_id IN ? AND is_deleted = ?", momentIDs, false).
+ Order("created_at ASC").
+ Find(&comments)
+ commentsByMoment := make(map[uint][]model.MomentComment, len(momentIDs))
+ for _, comment := range comments {
+ commentsByMoment[comment.MomentID] = append(commentsByMoment[comment.MomentID], comment)
+ }
+
+ // 3. 收集所有涉及的用户ID(发布者/展示的点赞者/评论者/被回复者),一条 IN 查询取回
+ userIDSet := make(map[string]struct{})
+ for i := range moments {
+ userIDSet[moments[i].UserID] = struct{}{}
+ }
+ for _, momentLikes := range likesByMoment {
+ for _, like := range momentLikes {
+ userIDSet[like.UserID] = struct{}{}
+ }
+ }
+ for _, comment := range comments {
+ userIDSet[comment.UserID] = struct{}{}
+ if comment.ReplyToUserID != "" {
+ userIDSet[comment.ReplyToUserID] = struct{}{}
+ }
+ }
+ userIDs := make([]string, 0, len(userIDSet))
+ for id := range userIDSet {
+ userIDs = append(userIDs, id)
+ }
+ userMap := make(map[string]*model.User, len(userIDs))
+ if len(userIDs) > 0 {
+ var users []model.User
+ if err := s.DB.Where("id IN ?", userIDs).Find(&users).Error; err == nil {
+ for i := range users {
+ users[i].Password = ""
+ userMap[users[i].ID] = &users[i]
+ }
+ }
+ }
+
+ // 4. 回填关联数据
+ for i := range moments {
+ moments[i].User = userMap[moments[i].UserID]
+ moments[i].IsLiked = likedByViewer[moments[i].ID]
+
+ momentLikes := likesByMoment[moments[i].ID]
+ for j := range momentLikes {
+ momentLikes[j].User = userMap[momentLikes[j].UserID]
+ }
+ moments[i].Likes = momentLikes
+
+ momentComments := commentsByMoment[moments[i].ID]
+ for j := range momentComments {
+ momentComments[j].User = userMap[momentComments[j].UserID]
+ if momentComments[j].ReplyToUserID != "" {
+ momentComments[j].ReplyToUser = userMap[momentComments[j].ReplyToUserID]
+ }
+ }
+ moments[i].Comments = momentComments
+ }
+}
+
/**
* loadMomentUser
* 功能:加载动态的发布者信息
@@ -597,7 +795,7 @@ func (s *MomentService) loadMomentLikes(moment *model.Moment, viewerID string) {
var likes []model.MomentLike
s.DB.Where("moment_id = ?", moment.ID).
Order("created_at DESC").
- Limit(10). // 只加载最近10个点赞
+ Limit(maxLikesPerMoment). // 只加载最近N个点赞,完整列表走 GetMomentLikes
Find(&likes)
for i := range likes {
diff --git a/internal/service/qrcode_login_service.go b/internal/service/qrcode_login_service.go
new file mode 100644
index 0000000..d434bc0
--- /dev/null
+++ b/internal/service/qrcode_login_service.go
@@ -0,0 +1,288 @@
+/**
+ * package service
+ * 作用:App 扫码登录 PC 的会话状态管理。
+ *
+ * 状态机(Redis 存储,TTL 自然过期兜底):
+ * pending(PC 生成二维码)
+ * → scanned(App 扫码,PC 显示"请在手机上确认")
+ * → confirmed(App 确认,写入登录 token,PC 轮询取走后立即删除)
+ * → cancelled(App 取消)
+ * 任意状态超时未推进 → key 过期,PC 轮询报 expired,提示刷新二维码
+ *
+ * 为什么用 Redis 而不是内存 map:服务多节点部署时(node_id 配置已支持),
+ * PC 轮询与 App 确认可能打到不同节点,必须走共享存储。
+ */
+package service
+
+import (
+ "context"
+ "crypto/rand"
+ "encoding/hex"
+ "encoding/json"
+ "errors"
+ "time"
+
+ "github.com/go-redis/redis/v8"
+)
+
+// 扫码登录会话状态常量
+const (
+ QRStatusPending = "pending" // 等待扫码
+ QRStatusScanned = "scanned" // 已扫码待确认
+ QRStatusConfirmed = "confirmed" // 已确认(token 已写入)
+ QRStatusCancelled = "cancelled" // 手机端取消
+)
+
+const (
+ // Redis key 前缀
+ qrLoginKeyPrefix = "qrcode:login:"
+ // 二维码有效期:过短用户来不及扫,过长增加被截屏盗用的窗口
+ qrLoginPendingTTL = 2 * time.Minute
+ // 扫码后确认阶段的有效期(扫码时续期,给用户留够阅读确认页的时间)
+ qrLoginConfirmTTL = 2 * time.Minute
+)
+
+// QRLoginSession 扫码登录会话(整体 JSON 序列化存入 Redis)
+type QRLoginSession struct {
+ Status string `json:"status"` // 当前状态(见 QRStatus* 常量)
+ ScannerID string `json:"scanner_id"` // 扫码用户ID(scan 时写入,confirm 时校验必须同一人)
+ Token string `json:"token"` // confirmed 后签发的 JWT(PC 取走即删,一次性)
+}
+
+// QRCodeLoginService 扫码登录服务
+type QRCodeLoginService struct {
+ Redis *redis.Client
+}
+
+// QRCodeLoginSvc 全局单例
+var QRCodeLoginSvc *QRCodeLoginService
+
+/**
+ * InitQRCodeLoginService
+ * 作用:初始化扫码登录服务(main.go 启动时调用)
+ */
+func InitQRCodeLoginService(rdb *redis.Client) {
+ QRCodeLoginSvc = &QRCodeLoginService{Redis: rdb}
+}
+
+/**
+ * generateQRLoginID
+ * 作用:生成不可猜测的二维码会话ID。
+ * 为什么用 crypto/rand:qr_id 同时是 PC 轮询的唯一凭证(status 接口无需登录),
+ * 必须保证无法被枚举,否则攻击者可劫持他人的确认结果拿到 token
+ */
+func generateQRLoginID() (string, error) {
+ b := make([]byte, 32)
+ if _, err := rand.Read(b); err != nil {
+ return "", err
+ }
+ return hex.EncodeToString(b), nil
+}
+
+/**
+ * Generate
+ * 作用:创建一个 pending 状态的扫码登录会话,返回 qr_id
+ */
+func (s *QRCodeLoginService) Generate() (string, error) {
+ qrID, err := generateQRLoginID()
+ if err != nil {
+ return "", errors.New("生成二维码会话失败")
+ }
+ session := QRLoginSession{Status: QRStatusPending}
+ data, _ := json.Marshal(session)
+ ctx := context.Background()
+ if err := s.Redis.Set(ctx, qrLoginKeyPrefix+qrID, data, qrLoginPendingTTL).Err(); err != nil {
+ return "", errors.New("二维码会话存储失败")
+ }
+ return qrID, nil
+}
+
+/**
+ * getSession
+ * 作用:读取会话;key 不存在返回 (nil, nil),由调用方解释为"已过期"
+ */
+func (s *QRCodeLoginService) getSession(qrID string) (*QRLoginSession, error) {
+ ctx := context.Background()
+ data, err := s.Redis.Get(ctx, qrLoginKeyPrefix+qrID).Result()
+ if err == redis.Nil {
+ return nil, nil
+ }
+ if err != nil {
+ return nil, err
+ }
+ var session QRLoginSession
+ if err := json.Unmarshal([]byte(data), &session); err != nil {
+ return nil, err
+ }
+ return &session, nil
+}
+
+/**
+ * saveSession
+ * 作用:写回会话并设置 TTL
+ */
+func (s *QRCodeLoginService) saveSession(qrID string, session *QRLoginSession, ttl time.Duration) error {
+ data, _ := json.Marshal(session)
+ return s.Redis.Set(context.Background(), qrLoginKeyPrefix+qrID, data, ttl).Err()
+}
+
+/**
+ * WaitForChange
+ * 作用:长轮询等待会话状态变化(不消费 token)。
+ * 服务端 hold 住请求,内部以短间隔轮询 Redis:
+ * - 状态与客户端已知状态(known)不同 → 立即返回;
+ * - key 不存在(过期)→ 立即返回 expired;
+ * - 超过 maxWait → 返回当前状态(客户端拿到后立刻发起下一轮);
+ * - 客户端断开(ctx.Done) → 提前返回,不再空转。
+ * 为什么轮 Redis 而不用 pub/sub:GET 每 800ms 一次成本可忽略,
+ * 且多节点部署下无需额外的订阅通道,实现最简单可靠
+ */
+func (s *QRCodeLoginService) WaitForChange(ctx context.Context, qrID, known string, maxWait time.Duration) (*QRLoginSession, bool, error) {
+ deadline := time.Now().Add(maxWait)
+ for {
+ session, err := s.getSession(qrID)
+ if err != nil {
+ return nil, false, errors.New("查询二维码状态失败")
+ }
+ if session == nil {
+ return nil, true, nil
+ }
+ if session.Status != known || time.Now().After(deadline) {
+ return session, false, nil
+ }
+ select {
+ case <-ctx.Done():
+ // 客户端已断开(关页面/取消请求),立即结束 hold
+ return session, false, nil
+ case <-time.After(800 * time.Millisecond):
+ }
+ }
+}
+
+/**
+ * GetStatus
+ * 作用:只读窥视会话状态,不消费 token。
+ * 返回 (session, error):session 为 nil 表示二维码已过期/不存在。
+ * 供 handler 在消费前先做可失败的准备工作(如查询用户信息)——
+ * 若先消费再准备、准备失败,一次性 token 已被删除且无法重试
+ */
+func (s *QRCodeLoginService) GetStatus(qrID string) (*QRLoginSession, error) {
+ session, err := s.getSession(qrID)
+ if err != nil {
+ return nil, errors.New("查询二维码状态失败")
+ }
+ return session, nil
+}
+
+/**
+ * GetStatusAndConsume
+ * 作用:读取会话状态,confirmed 时原子消费 token。
+ * 返回 (session, expired, error):expired=true 表示二维码已过期/不存在。
+ * 关键安全设计:confirmed 状态的 token 一次性交付。
+ * 用 DEL 的返回值做唯一性仲裁——两个客户端并发轮询同一 qr_id 时,
+ * Redis 只会让其中一个 DEL 返回 1,另一个返回 0 按已过期处理,
+ * 保证登录凭证绝不会被交付两次(单纯 GET 后 DEL 存在竞态窗口)
+ */
+func (s *QRCodeLoginService) GetStatusAndConsume(qrID string) (*QRLoginSession, bool, error) {
+ session, err := s.getSession(qrID)
+ if err != nil {
+ return nil, false, errors.New("查询二维码状态失败")
+ }
+ if session == nil {
+ return nil, true, nil
+ }
+ if session.Status == QRStatusConfirmed {
+ deleted, dErr := s.Redis.Del(context.Background(), qrLoginKeyPrefix+qrID).Result()
+ if dErr != nil {
+ return nil, false, errors.New("查询二维码状态失败")
+ }
+ if deleted == 0 {
+ // key 已被并发请求消费(或恰好过期):本请求视为过期,不重复交付 token
+ return nil, true, nil
+ }
+ }
+ return session, false, nil
+}
+
+/**
+ * Scan
+ * 作用:App 扫码,pending → scanned,记录扫码人并续期(给确认页留时间)
+ */
+func (s *QRCodeLoginService) Scan(qrID, userID string) error {
+ session, err := s.getSession(qrID)
+ if err != nil {
+ return errors.New("查询二维码状态失败")
+ }
+ if session == nil {
+ return errors.New("二维码已过期,请刷新后重新扫描")
+ }
+ if session.Status != QRStatusPending {
+ return errors.New("二维码已被使用")
+ }
+ session.Status = QRStatusScanned
+ session.ScannerID = userID
+ return s.saveSession(qrID, session, qrLoginConfirmTTL)
+}
+
+/**
+ * Confirm
+ * 作用:App 确认登录,scanned → confirmed,签发 PC 端登录 token。
+ * @param signToken 签发函数(由调用方传入 utils.GenerateToken,避免 service 反向依赖 utils/jwt 的初始化时序)
+ */
+func (s *QRCodeLoginService) Confirm(qrID, userID string, signToken func(string) (string, error)) error {
+ session, err := s.getSession(qrID)
+ if err != nil {
+ return errors.New("查询二维码状态失败")
+ }
+ if session == nil {
+ return errors.New("二维码已过期,请重新扫描")
+ }
+ if session.Status != QRStatusScanned {
+ return errors.New("二维码状态异常,请重新扫描")
+ }
+ // 必须由扫码本人确认:防止 qr_id 在扫码后泄漏,被其他账号"替身确认"
+ if session.ScannerID != userID {
+ return errors.New("确认用户与扫码用户不一致")
+ }
+ token, err := signToken(userID)
+ if err != nil {
+ return errors.New("生成登录凭证失败")
+ }
+ session.Status = QRStatusConfirmed
+ session.Token = token
+ return s.saveSession(qrID, session, qrLoginConfirmTTL)
+}
+
+/**
+ * Cancel
+ * 作用:App 取消登录,scanned → cancelled(PC 轮询到后提示"已取消"并刷新二维码)
+ * 状态机校验:只允许 scanned 状态由扫码本人取消。为什么必须收紧:
+ * - pending 阶段无人扫码、ScannerID 为空,原实现的权限判断直接放行,
+ * 任意登录用户拿到 qr_id 就能把他人的二维码置为已取消(拒绝服务);
+ * - confirmed 阶段登录 token 已写入会话,再取消会覆盖状态、丢弃 token,
+ * PC 轮询永远拿不到登录凭证(手机上显示确认成功、PC 却登录失败)。
+ */
+func (s *QRCodeLoginService) Cancel(qrID, userID string) error {
+ session, err := s.getSession(qrID)
+ if err != nil {
+ return errors.New("查询二维码状态失败")
+ }
+ if session == nil {
+ // 已过期视为取消成功(幂等)
+ return nil
+ }
+ // 重复取消视为成功(幂等,避免 App 端重复点击报错)
+ if session.Status == QRStatusCancelled {
+ return nil
+ }
+ // 仅 scanned 可取消:pending 无人可取消,confirmed 不可回退
+ if session.Status != QRStatusScanned {
+ return errors.New("二维码状态已变更,无法取消")
+ }
+ // 仅扫码本人可取消
+ if session.ScannerID != userID {
+ return errors.New("无权操作该二维码")
+ }
+ session.Status = QRStatusCancelled
+ return s.saveSession(qrID, session, qrLoginConfirmTTL)
+}
diff --git a/internal/service/room_service.go b/internal/service/room_service.go
index 275a822..f43521d 100644
--- a/internal/service/room_service.go
+++ b/internal/service/room_service.go
@@ -13,6 +13,7 @@ import (
"xk-websocket-v2/internal/utils"
"gorm.io/gorm"
+ "gorm.io/gorm/clause"
)
// RoomService 房间服务结构体
@@ -65,12 +66,29 @@ func (s *RoomService) CreateRoom(roomType string, members []string, creatorID st
var roomID string
var ownerID string = "0" // 单聊默认为0
+ // 强制校验创建者必须在成员列表中:
+ // p2p 房间若创建者不在成员里,等于替别人建两人房间(越权);
+ // 群聊则自动把创建者补进成员,防止建出一个自己不在内的群
+ creatorInMembers := false
+ for _, m := range members {
+ if m == creatorID {
+ creatorInMembers = true
+ break
+ }
+ }
+
if roomType == "p2p" {
if len(members) != 2 {
return nil, fmt.Errorf("点对点房间需要2个成员")
}
+ if !creatorInMembers {
+ return nil, fmt.Errorf("点对点房间创建者必须是成员之一")
+ }
roomID = GenerateP2PRoomID(members[0], members[1])
} else {
+ if !creatorInMembers {
+ members = append(members, creatorID)
+ }
roomID = GenerateGroupRoomID()
// 群聊时,创建者就是群主
ownerID = creatorID
@@ -201,6 +219,43 @@ func (s *RoomService) GetRoomMember(roomID, userID string) (*model.RoomMember, e
return &member, nil
}
+/**
+ * IsRoomMember
+ * 功能:判断用户是否为房间成员(用于接口越权校验)
+ * 为什么这样写:
+ * 1. 优先查 room_members 表,命中即为成员;
+ * 2. 点对点房间的 room_id 形如 "a_b"(两个用户ID按序拼接),
+ * 为兼容历史数据可能缺失成员记录的情况,回退用房间ID中是否包含该用户ID判断,
+ * 避免误拦截合法的单聊双方。
+ */
+func (s *RoomService) IsRoomMember(roomID, userID string) bool {
+ if roomID == "" || userID == "" {
+ return false
+ }
+ var count int64
+ s.DB.Model(&model.RoomMember{}).Where("room_id = ? AND user_id = ?", roomID, userID).Count(&count)
+ if count > 0 {
+ return true
+ }
+ // 点对点房间回退判断:room_id 形如 "a_b"(两个用户ID按序拼接)。
+ // 不能按 Split 段数判断:AI 机器人的用户ID自带下划线(bot_xxx),
+ // "1001_bot_xxx" 会切出3段导致兜底失效,机器人私聊拉历史直接 403。
+ // 改为前/后缀匹配,并用统一生成规则重组校验,防止"ID 恰好是房间号子串"的误判
+ if strings.HasPrefix(roomID, userID+"_") {
+ other := strings.TrimPrefix(roomID, userID+"_")
+ if other != "" && GenerateP2PRoomID(userID, other) == roomID {
+ return true
+ }
+ }
+ if strings.HasSuffix(roomID, "_"+userID) {
+ other := strings.TrimSuffix(roomID, "_"+userID)
+ if other != "" && GenerateP2PRoomID(userID, other) == roomID {
+ return true
+ }
+ }
+ return false
+}
+
/**
* AddRoomMember
* 功能:添加房间成员
@@ -231,37 +286,21 @@ func (s *RoomService) CreateGroupRoom(creatorID string, memberIDs []string, name
return nil, fmt.Errorf("成员列表不能为空")
}
- // 1. 先通过通用的 CreateRoom 创建群聊和基础成员关系(角色默认都是 0)
- room, err := s.CreateRoom("group", memberIDs, creatorID)
- if err != nil {
- return nil, err
- }
-
- // 2. 更新群名称和头像(如果提供)
- updates := map[string]interface{}{}
- if name != "" {
- updates["room_name"] = name
- }
- if avatar != "" {
- updates["room_avatar"] = avatar
- }
- if len(updates) > 0 {
- if err := s.DB.Model(room).Updates(updates).Error; err != nil {
- return nil, err
+ // 成员去重,并强制把创建者纳入成员列表(不能建一个自己不在内的群)
+ memberSet := make(map[string]struct{}, len(memberIDs)+1)
+ orderedMembers := make([]string, 0, len(memberIDs)+1)
+ for _, id := range append([]string{creatorID}, memberIDs...) {
+ if id == "" {
+ continue
}
+ if _, ok := memberSet[id]; ok {
+ continue
+ }
+ memberSet[id] = struct{}{}
+ orderedMembers = append(orderedMembers, id)
}
- // 3. 为成员设置角色:群主=2 管理员=1 成员=0
- roomID := room.RoomID
-
- // 创建者作为群主
- if err := s.DB.Model(&model.RoomMember{}).
- Where("room_id = ? AND user_id = ?", roomID, creatorID).
- Update("role", 2).Error; err != nil {
- return nil, err
- }
-
- // 管理员列表去重并排除群主
+ // 管理员列表去重并排除群主(群主角色固定为2,不能被降级成管理员)
adminSet := map[string]struct{}{}
for _, id := range adminIDs {
if id == "" || id == creatorID {
@@ -270,17 +309,54 @@ func (s *RoomService) CreateGroupRoom(creatorID string, memberIDs []string, name
adminSet[id] = struct{}{}
}
- if len(adminSet) > 0 {
- for adminID := range adminSet {
- if err := s.DB.Model(&model.RoomMember{}).
- Where("room_id = ? AND user_id = ?", roomID, adminID).
- Update("role", 1).Error; err != nil {
- return nil, err
- }
+ roomName := name
+ if roomName == "" {
+ roomName = "群聊"
+ }
+
+ // 单事务完成"建房 + 写成员(带角色)":
+ // 原实现分四步独立写库(建房、改名/头像、设群主、逐个设管理员),
+ // 任一步失败都会留下"群已存在但名称/角色不完整"的脏数据,这里改为要么全部成功要么全部回滚
+ tx := s.DB.Begin()
+ defer func() {
+ if r := recover(); r != nil {
+ tx.Rollback()
+ }
+ }()
+
+ room := model.ChatRoom{
+ RoomID: GenerateGroupRoomID(),
+ RoomType: "group",
+ RoomName: roomName,
+ RoomAvatar: avatar,
+ OwnerID: creatorID,
+ CreatorID: creatorID,
+ }
+ if err := tx.Create(&room).Error; err != nil {
+ tx.Rollback()
+ return nil, err
+ }
+
+ // 写入成员记录,角色在插入时一次确定:群主=2 管理员=1 成员=0
+ for _, uid := range orderedMembers {
+ var role int8 = 0
+ if uid == creatorID {
+ role = 2
+ } else if _, ok := adminSet[uid]; ok {
+ role = 1
+ }
+ member := model.RoomMember{RoomID: room.RoomID, UserID: uid, Role: role}
+ if err := tx.Create(&member).Error; err != nil {
+ tx.Rollback()
+ return nil, err
}
}
- return room, nil
+ if err := tx.Commit().Error; err != nil {
+ return nil, err
+ }
+
+ return &room, nil
}
/**
@@ -387,20 +463,29 @@ func (s *RoomService) AddGroupMembers(roomID, operatorID string, memberIDs []str
return fmt.Errorf("无权限添加成员")
}
+ // 组装批量插入的数据(过滤空ID)
+ members := make([]model.RoomMember, 0, len(memberIDs))
for _, uid := range memberIDs {
if uid == "" {
continue
}
- member := model.RoomMember{
+ members = append(members, model.RoomMember{
RoomID: roomID,
UserID: uid,
Role: 0, // 默认成员
- }
- if err := s.DB.Create(&member).Error; err != nil {
- return err
- }
+ })
}
- return nil
+ if len(members) == 0 {
+ return nil
+ }
+
+ // 事务 + OnConflict DoNothing:
+ // 1) 原逻辑逐条 Create 且不在事务内,某个成员已在群时 (RoomID,UserID) 复合主键冲突
+ // 会中途报错返回——前面的已插入、后面的没插入(非原子);
+ // 2) DoNothing 让"重复加入同一群"幂等跳过而不是整体失败
+ return s.DB.Transaction(func(tx *gorm.DB) error {
+ return tx.Clauses(clause.OnConflict{DoNothing: true}).Create(&members).Error
+ })
}
/**
@@ -462,18 +547,36 @@ func (s *RoomService) DissolveGroup(roomID, ownerID string) error {
return fmt.Errorf("只有群主可以解散群聊")
}
+ // 事务保证"删成员、删房间、清会话"原子执行:
+ // 原实现分两次独立删除且不清会话,中途失败会出现"房间还在但成员全没了",
+ // 且所有成员的会话列表里永远残留一个点不开的死群
+ tx := s.DB.Begin()
+ defer func() {
+ if r := recover(); r != nil {
+ tx.Rollback()
+ }
+ }()
+
// 删除所有成员记录
- if err := s.DB.Where("room_id = ?", roomID).Delete(&model.RoomMember{}).Error; err != nil {
+ if err := tx.Where("room_id = ?", roomID).Delete(&model.RoomMember{}).Error; err != nil {
+ tx.Rollback()
return fmt.Errorf("删除成员记录失败: %v", err)
}
- // 可选:标记房间为已解散(或直接删除房间)
- // 这里选择删除房间记录
- if err := s.DB.Where("room_id = ?", roomID).Delete(&model.ChatRoom{}).Error; err != nil {
+ // 删除房间记录
+ if err := tx.Where("room_id = ?", roomID).Delete(&model.ChatRoom{}).Error; err != nil {
+ tx.Rollback()
return fmt.Errorf("删除房间记录失败: %v", err)
}
- return nil
+ // 清理所有成员的该群会话(type=2 群聊会话,target_id 即群ID)
+ if err := tx.Where("target_id = ? AND type = 2", roomID).
+ Delete(&model.ChatConversation{}).Error; err != nil {
+ tx.Rollback()
+ return fmt.Errorf("清理群会话失败: %v", err)
+ }
+
+ return tx.Commit().Error
}
/**
@@ -501,30 +604,56 @@ func (s *RoomService) ChangeMemberRole(roomID, operatorID, memberID string, role
return fmt.Errorf("不能修改群主自身角色")
}
- // 获取原角色
+ // 获取原角色:查询错误必须处理,查不到说明目标不是群成员,
+ // 原实现忽略错误会把 oldRole 当 0 继续执行,给非成员"改角色"
var oldMember model.RoomMember
- s.DB.Where("room_id = ? AND user_id = ?", roomID, memberID).First(&oldMember)
+ if err := s.DB.Where("room_id = ? AND user_id = ?", roomID, memberID).First(&oldMember).Error; err != nil {
+ if err == gorm.ErrRecordNotFound {
+ return fmt.Errorf("目标用户不是群成员")
+ }
+ return err
+ }
oldRole := oldMember.Role
- if err := s.DB.Model(&model.RoomMember{}).
+ // 事务保证角色变更原子性:转让群主涉及"新群主升级、原群主降级、房间owner_id更换"三步,
+ // 原实现三次独立 Update 且忽略后两步错误,中途失败会出现双群主/无群主的脏数据
+ tx := s.DB.Begin()
+ defer func() {
+ if r := recover(); r != nil {
+ tx.Rollback()
+ }
+ }()
+
+ if err := tx.Model(&model.RoomMember{}).
Where("room_id = ? AND user_id = ?", roomID, memberID).
Update("role", role).Error; err != nil {
+ tx.Rollback()
return err
}
- // 如果是转让群主,需要把原群主降级
+ // 如果是转让群主,需要把原群主降级并更换房间 owner_id
if role == 2 && memberID != room.OwnerID {
// 将原群主降级为管理员
- s.DB.Model(&model.RoomMember{}).
+ if err := tx.Model(&model.RoomMember{}).
Where("room_id = ? AND user_id = ?", roomID, operatorID).
- Update("role", 1)
+ Update("role", 1).Error; err != nil {
+ tx.Rollback()
+ return err
+ }
// 更新房间的 owner_id
- s.DB.Model(&model.ChatRoom{}).
+ if err := tx.Model(&model.ChatRoom{}).
Where("room_id = ?", roomID).
- Update("owner_id", memberID)
+ Update("owner_id", memberID).Error; err != nil {
+ tx.Rollback()
+ return err
+ }
}
- // 发送角色变更通知(异步)
+ if err := tx.Commit().Error; err != nil {
+ return err
+ }
+
+ // 发送角色变更通知(异步,放在事务提交成功之后,避免回滚了却已发通知)
go s.sendRoleChangeNotification(roomID, operatorID, memberID, oldRole, role)
return nil
@@ -605,6 +734,29 @@ func (s *RoomService) UpdateGroupAnnouncement(roomID, userID, announcement strin
Update("announcement", announcement).Error
}
+/**
+ * UpdateMemberNickname
+ * 功能:更新"我在本群的昵称"(群名片)
+ * 说明:只允许修改自己的群名片,因此无需管理员权限,仅校验成员身份;
+ * 昵称允许为空字符串(表示清除群名片,恢复显示账号昵称)
+ */
+func (s *RoomService) UpdateMemberNickname(roomID, userID, nickname string) error {
+ // 必须是群成员才能设置群名片
+ var count int64
+ if err := s.DB.Model(&model.RoomMember{}).
+ Where("room_id = ? AND user_id = ?", roomID, userID).
+ Count(&count).Error; err != nil {
+ return err
+ }
+ if count == 0 {
+ return fmt.Errorf("您不是该群成员")
+ }
+
+ return s.DB.Model(&model.RoomMember{}).
+ Where("room_id = ? AND user_id = ?", roomID, userID).
+ Update("nickname", nickname).Error
+}
+
/**
* MuteGroupMember
* 功能:禁言群成员(仅群主或管理员可操作)
@@ -736,32 +888,8 @@ func (s *RoomService) sendMuteNotification(roomID, operatorID, memberID string,
}
notifExtraJSON, _ := json.Marshal(notifExtra)
- // 为所有成员发送群通知
- for _, uid := range memberIDs {
- notifMsg := model.ChatMessage{
- RoomID: roomID,
- SenderUserID: operatorID,
- ReceiverUserID: roomID,
- MessageType: model.MessageTypeGroupNotif,
- Content: notifContent,
- Extra: string(notifExtraJSON),
- }
-
- if err := ChatSvc.DB.Create(¬ifMsg).Error; err == nil {
- // 更新会话
- if ConversationSvc != nil {
- _ = ConversationSvc.UpsertConversationOnMessage(uid, roomID, roomID, notifMsg, false)
- }
-
- // 推送消息
- pushMsg := model.WsPayload{
- RequestType: "receive_message",
- Data: notifMsg,
- }
- msgBytes, _ := json.Marshal(pushMsg)
- ChatSvc.DispatchMessage(uid, msgBytes)
- }
- }
+ // 为所有成员发送群通知:房间级落库一次,按成员推送(复用统一广播方法)
+ ChatSvc.BroadcastGroupNotification(operatorID, roomID, notifContent, string(notifExtraJSON), memberIDs)
}
/**
@@ -832,30 +960,6 @@ func (s *RoomService) sendRoleChangeNotification(roomID, operatorID, memberID st
}
notifExtraJSON, _ := json.Marshal(notifExtra)
- // 为所有成员发送群通知
- for _, uid := range memberIDs {
- notifMsg := model.ChatMessage{
- RoomID: roomID,
- SenderUserID: operatorID,
- ReceiverUserID: roomID,
- MessageType: model.MessageTypeGroupNotif,
- Content: notifContent,
- Extra: string(notifExtraJSON),
- }
-
- if err := ChatSvc.DB.Create(¬ifMsg).Error; err == nil {
- // 更新会话
- if ConversationSvc != nil {
- _ = ConversationSvc.UpsertConversationOnMessage(uid, roomID, roomID, notifMsg, false)
- }
-
- // 推送消息
- pushMsg := model.WsPayload{
- RequestType: "receive_message",
- Data: notifMsg,
- }
- msgBytes, _ := json.Marshal(pushMsg)
- ChatSvc.DispatchMessage(uid, msgBytes)
- }
- }
+ // 为所有成员发送群通知:房间级落库一次,按成员推送(复用统一广播方法)
+ ChatSvc.BroadcastGroupNotification(operatorID, roomID, notifContent, string(notifExtraJSON), memberIDs)
}
diff --git a/internal/utils/db.go b/internal/utils/db.go
new file mode 100644
index 0000000..ec21641
--- /dev/null
+++ b/internal/utils/db.go
@@ -0,0 +1,29 @@
+/**
+ * package utils
+ * 作用:数据库相关的通用工具函数
+ */
+package utils
+
+import (
+ "errors"
+ "strings"
+
+ "gorm.io/gorm"
+)
+
+/**
+ * IsDuplicateEntryError
+ * 功能:判断是否为唯一索引冲突错误(MySQL 错误码 1062)。
+ * 为什么需要:依赖唯一索引做并发兜底的写入(如点赞、消息删除记录)
+ * 需要把"重复插入"当作幂等成功处理。GORM 的 ErrDuplicatedKey 依赖驱动
+ * Translate 支持,部分配置下返回的是原始 mysql 错误,需按错误文本兜底判断
+ */
+func IsDuplicateEntryError(err error) bool {
+ if err == nil {
+ return false
+ }
+ if errors.Is(err, gorm.ErrDuplicatedKey) {
+ return true
+ }
+ return strings.Contains(err.Error(), "Duplicate entry") || strings.Contains(err.Error(), "1062")
+}
diff --git a/internal/utils/jwt.go b/internal/utils/jwt.go
index 348be8b..c85de89 100644
--- a/internal/utils/jwt.go
+++ b/internal/utils/jwt.go
@@ -1,13 +1,13 @@
/**
* package utils
- *
+ *
* JWT Token生成和验证工具包
- *
+ *
* 功能概述:
* 1. 生成JWT Token(包含用户ID和过期时间)
* 2. 解析JWT Token(验证签名和过期时间)
* 3. 验证Token有效性(提取用户ID)
- *
+ *
* 使用场景:
* - 用户登录后生成Token
* - API请求时验证Token
@@ -20,64 +20,64 @@ import (
"time"
"github.com/golang-jwt/jwt/v5"
- "github.com/spf13/viper"
)
-// jwtSecret JWT签名密钥(从配置文件读取)
-var jwtSecret []byte
+// jwtSecret JWT签名密钥
+// 说明:默认赋一个本地开发用的弱密钥兜底,避免未初始化时为 nil;
+// 生产环境必须通过 InitJWT 从配置/环境变量注入强随机密钥覆盖。
+var jwtSecret = []byte("xk-websocket-secret-key-2025")
/**
- * init
- *
- * 功能:初始化JWT签名密钥
- *
- * 步骤:
- * 1. 从配置文件读取JWT密钥
- * 2. 如果配置文件中没有,使用默认密钥(仅用于开发环境)
- * 3. 将密钥转换为字节数组存储
- *
- * 注意:生产环境必须使用配置文件中的强随机密钥
+ * InitJWT
+ *
+ * 功能:初始化(覆盖)JWT签名密钥
+ *
+ * 为什么这样写:
+ * 原实现用包级 init() 读取 viper 配置,但 init() 会早于 main 中的 initConfig()(viper.ReadInConfig)执行,
+ * 导致读到空串而永远落到默认密钥,配置文件里的 jwt.secret 从未生效。
+ * 因此改为显式初始化函数,由 main 在配置加载完成后调用,确保真正使用配置/环境变量中的密钥。
+ *
+ * @param secret 签名密钥(为空时保留默认弱密钥,仅适用于本地开发)
*/
-func init() {
- secret := viper.GetString("jwt.secret")
+func InitJWT(secret string) {
if secret == "" {
- secret = "xk-websocket-secret-key-2025" // 默认密钥,生产环境应使用配置
+ return // 保留默认密钥,避免线上误清空导致签名不一致
}
jwtSecret = []byte(secret)
}
/**
* Claims
- *
+ *
* JWT Token的载荷结构
- *
+ *
* 字段说明:
* - UserID: 用户ID(业务数据)
* - RegisteredClaims: JWT标准声明(过期时间、签发时间等)
*/
type Claims struct {
- UserID string `json:"user_id"` // 用户ID
- jwt.RegisteredClaims // JWT标准声明
+ UserID string `json:"user_id"` // 用户ID
+ jwt.RegisteredClaims // JWT标准声明
}
/**
* GenerateToken
- *
+ *
* 功能:生成JWT Token
- *
+ *
* 步骤:
* 1. 设置Token过期时间(默认7天)
* 2. 创建Claims对象,包含用户ID和标准声明
* 3. 使用HS256算法签名Token
* 4. 返回Token字符串
- *
+ *
* @param userID 用户ID
* @returns token字符串和错误
*/
func GenerateToken(userID string) (string, error) {
// 步骤1: 设置Token过期时间(7天后过期)
expirationTime := time.Now().Add(7 * 24 * time.Hour)
-
+
// 步骤2: 创建Claims对象,包含用户ID和标准声明
claims := &Claims{
UserID: userID, // 业务数据:用户ID
@@ -101,22 +101,22 @@ func GenerateToken(userID string) (string, error) {
/**
* ParseToken
- *
+ *
* 功能:解析JWT Token
- *
+ *
* 步骤:
* 1. 创建空的Claims对象
* 2. 使用密钥解析Token并验证签名
* 3. 检查Token是否有效(签名正确、未过期)
* 4. 返回Claims对象
- *
+ *
* @param tokenString token字符串
* @returns Claims和错误
*/
func ParseToken(tokenString string) (*Claims, error) {
// 步骤1: 创建空的Claims对象
claims := &Claims{}
-
+
// 步骤2: 解析Token并验证签名
// 使用密钥验证Token的签名是否有效
token, err := jwt.ParseWithClaims(tokenString, claims, func(token *jwt.Token) (interface{}, error) {
@@ -138,18 +138,18 @@ func ParseToken(tokenString string) (*Claims, error) {
/**
* ValidateToken
- *
+ *
* 功能:验证Token有效性并提取用户ID
- *
+ *
* 步骤:
* 1. 调用ParseToken解析Token
* 2. 如果解析成功,从Claims中提取用户ID
* 3. 返回用户ID
- *
+ *
* 使用场景:
* - 中间件中验证Token
* - API处理器中获取当前用户ID
- *
+ *
* @param tokenString token字符串
* @returns 用户ID和错误
*/
@@ -159,8 +159,7 @@ func ValidateToken(tokenString string) (string, error) {
if err != nil {
return "", err // Token无效或已过期
}
-
+
// 步骤2: 从Claims中提取用户ID
return claims.UserID, nil
}
-
diff --git a/internal/utils/response.go b/internal/utils/response.go
index 50e3db5..5774c99 100644
--- a/internal/utils/response.go
+++ b/internal/utils/response.go
@@ -17,12 +17,12 @@ import (
// 业务状态码常量
const (
- CodeSuccess = 0 // 成功
- CodeBadRequest = 400 // 参数错误
- CodeUnauthorized = 401 // 未认证
- CodeForbidden = 403 // 无权限
- CodeNotFound = 404 // 资源不存在
- CodeInternalError = 500 // 服务器错误
+ CodeSuccess = 0 // 成功
+ CodeBadRequest = 400 // 参数错误
+ CodeUnauthorized = 401 // 未认证
+ CodeForbidden = 403 // 无权限
+ CodeNotFound = 404 // 资源不存在
+ CodeInternalError = 500 // 服务器错误
)
// 响应类型常量
@@ -53,11 +53,14 @@ func formatDuration(d time.Duration) string {
}
/**
- * Response
- * 作用:统一的响应函数,所有响应都通过此函数返回
- * 说明:统一返回HTTP 200状态码
+ * ResponseWithStatus
+ * 作用:统一响应的单一组装口——构造含 interface_info(result_time/ecs)的
+ * ApiResponse 并按指定 HTTP 状态码返回。
+ * 说明:绝大多数业务接口经 Response 走 HTTP 200;少数需要真实 HTTP 状态码的
+ * 场景(如认证中间件的 401,前端请求层按状态码识别登出)直接调用本函数,
+ * 保证所有出口的响应体结构完全一致,不再出现缺 interface_info 的手工构造
*/
-func Response(c *gin.Context, code int, message string, result interface{}, responseType string) {
+func ResponseWithStatus(c *gin.Context, httpStatus int, code int, message string, result interface{}, responseType string) {
// 从Context获取请求开始时间
startTime, exists := c.Get("request_start_time")
var duration time.Duration
@@ -78,8 +81,16 @@ func Response(c *gin.Context, code int, message string, result interface{}, resp
},
}
- // 统一返回HTTP 200状态码
- c.JSON(http.StatusOK, response)
+ c.JSON(httpStatus, response)
+}
+
+/**
+ * Response
+ * 作用:统一的响应函数,所有业务响应都通过此函数返回
+ * 说明:统一返回HTTP 200状态码(错误经 code 字段标识),组装逻辑收口到 ResponseWithStatus
+ */
+func Response(c *gin.Context, code int, message string, result interface{}, responseType string) {
+ ResponseWithStatus(c, http.StatusOK, code, message, result, responseType)
}
/**
@@ -169,4 +180,3 @@ func InternalError(c *gin.Context, message string) {
}
Response(c, CodeInternalError, message, nil, TypeError)
}
-
diff --git a/internal/ws/worker.go b/internal/ws/worker.go
index b54194b..1b07d87 100644
--- a/internal/ws/worker.go
+++ b/internal/ws/worker.go
@@ -128,12 +128,21 @@ func handleMessage(client *manager.Client, message []byte) {
// 处理用户绑定 (bind)
// 格式: {"request_type": "bind", "data": {"user_id": "1001"}}
var bindData model.BindReq
- if err := json.Unmarshal(jsonBytes(req.Data), &bindData); err == nil {
- // 调用 Service 层逻辑
- service.ChatSvc.BindUser(client, bindData.UserID)
- } else {
+ if err := json.Unmarshal(jsonBytes(req.Data), &bindData); err != nil {
log.Printf("⚠️ [Worker] 绑定参数错误: %v", err)
+ break
}
+ // 安全校验:连接建立时已用 JWT 解析出的身份绑定(见 main.go 步骤7),
+ // 这里的 bind 只允许绑定“自己”。否则任意登录用户发
+ // {"request_type":"bind","data":{"user_id":"<受害者ID>"}} 就能把自己的连接
+ // 登记进受害者的投递列表、并把 Redis 路由改写到自己节点,从而劫持/错投受害者的消息(严重越权)。
+ // 通过 GetClientUserID 在锁下读取,避免与 BindUser 的写入构成数据竞争
+ boundUserID := manager.Manager.GetClientUserID(client)
+ if boundUserID != "" && bindData.UserID != boundUserID {
+ log.Printf("🚫 [Worker] 拒绝越权 bind:client=%s(已绑定 %s) 试图绑定 %s", client.ID, boundUserID, bindData.UserID)
+ break
+ }
+ service.ChatSvc.BindUser(client, bindData.UserID)
case "send_message":
// 处理消息发送 (send_message)
@@ -142,7 +151,10 @@ func handleMessage(client *manager.Client, message []byte) {
if err := json.Unmarshal(jsonBytes(req.Data), &msgData); err == nil {
// 核心:调用 ChatService 处理消息(持久化、转发、多端同步)
// 注意:这里传入的是 client 指针,以便 Service 层获取发送者的 UserID 和 ClientID
- service.ChatSvc.HandleUserMessage(client, &msgData)
+ // WS 路径下被拒绝时 Service 内部已向发送者推送 request_type=error,这里只记日志
+ if _, hErr := service.ChatSvc.HandleUserMessage(client, &msgData); hErr != nil {
+ log.Printf("🚫 [Worker] 消息被拒绝: user=%s err=%v", client.UserID, hErr)
+ }
} else {
log.Printf("⚠️ [Worker] 消息参数错误: %v", err)
}
diff --git a/migrations/01_20260707_upgrade_p1_p2.sql b/migrations/01_20260707_upgrade_p1_p2.sql
new file mode 100644
index 0000000..03f7bf6
--- /dev/null
+++ b/migrations/01_20260707_upgrade_p1_p2.sql
@@ -0,0 +1,104 @@
+/*
+ * NL-IM 增量数据库迁移脚本
+ * 文件: migrations/01_20260707_upgrade_p1_p2.sql
+ * 版本: P1 (ip2region 消息归属地) + P2 (已读回执 / 用户设置)
+ *
+ * 适用 MySQL: 5.7+ / 8.0+(推荐 8.0.12+)
+ * 目标库名: nl_im_plus(请按实际环境修改下方 USE 语句)
+ *
+ * ----------------------------------------------------------------
+ * 执行方式:
+ * mysql -u root -p nl_im_plus < migrations/01_20260707_upgrade_p1_p2.sql
+ *
+ * 或在客户端中:
+ * USE nl_im_plus;
+ * SOURCE /path/to/migrations/01_20260707_upgrade_p1_p2.sql;
+ *
+ * 本脚本设计为「可重复执行」:列/表已存在时会跳过,不会报错。
+ *
+ * 变更内容:
+ * 1. chat_messages 增加 sender_ip、ip_location
+ * 2. 新建 message_read_receipts(消息已读回执)
+ * 3. 新建 user_settings(用户个性化设置)
+ *
+ * 回滚说明(谨慎操作,生产环境请先备份):
+ * -- ALTER TABLE chat_messages DROP COLUMN ip_location;
+ * -- ALTER TABLE chat_messages DROP COLUMN sender_ip;
+ * -- DROP TABLE IF EXISTS message_read_receipts;
+ * -- DROP TABLE IF EXISTS user_settings;
+ *
+ * 与 GORM AutoMigrate 关系:
+ * 新环境可直接启动 nl-im-service 自动迁移;已有库建议用本脚本做可控升级。
+ * ----------------------------------------------------------------
+ */
+
+SET NAMES utf8mb4;
+SET FOREIGN_KEY_CHECKS = 0;
+
+-- 请按实际库名修改
+USE `nl_im_plus`;
+
+-- ============================================================
+-- 1. chat_messages 增加 IP 归属地字段(兼容 MySQL 5.7 / 8.0)
+-- ============================================================
+
+-- sender_ip
+SET @col_exists := (
+ SELECT COUNT(*) FROM information_schema.COLUMNS
+ WHERE TABLE_SCHEMA = DATABASE()
+ AND TABLE_NAME = 'chat_messages'
+ AND COLUMN_NAME = 'sender_ip'
+);
+SET @sql := IF(
+ @col_exists = 0,
+ 'ALTER TABLE `chat_messages` ADD COLUMN `sender_ip` VARCHAR(50) NULL COMMENT ''发送者IP'' AFTER `sender_user_id`',
+ 'SELECT ''skip: chat_messages.sender_ip already exists'' AS migration_note'
+);
+PREPARE stmt FROM @sql;
+EXECUTE stmt;
+DEALLOCATE PREPARE stmt;
+
+-- ip_location
+SET @col_exists := (
+ SELECT COUNT(*) FROM information_schema.COLUMNS
+ WHERE TABLE_SCHEMA = DATABASE()
+ AND TABLE_NAME = 'chat_messages'
+ AND COLUMN_NAME = 'ip_location'
+);
+SET @sql := IF(
+ @col_exists = 0,
+ 'ALTER TABLE `chat_messages` ADD COLUMN `ip_location` VARCHAR(255) NULL COMMENT ''IP归属地'' AFTER `sender_ip`',
+ 'SELECT ''skip: chat_messages.ip_location already exists'' AS migration_note'
+);
+PREPARE stmt FROM @sql;
+EXECUTE stmt;
+DEALLOCATE PREPARE stmt;
+
+-- ============================================================
+-- 2. message_read_receipts 消息已读回执表
+-- ============================================================
+CREATE TABLE IF NOT EXISTS `message_read_receipts` (
+ `id` bigint UNSIGNED NOT NULL AUTO_INCREMENT COMMENT '主键ID',
+ `message_id` bigint UNSIGNED NOT NULL COMMENT '消息ID',
+ `user_id` varchar(100) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NOT NULL COMMENT '已读用户ID',
+ `read_at` datetime(3) NOT NULL COMMENT '已读时间',
+ PRIMARY KEY (`id`) USING BTREE,
+ UNIQUE INDEX `uk_msg_user`(`message_id` ASC, `user_id` ASC) USING BTREE,
+ INDEX `idx_message_read_receipts_user_id`(`user_id` ASC) USING BTREE
+) ENGINE = InnoDB CHARACTER SET = utf8mb4 COLLATE = utf8mb4_0900_ai_ci COMMENT = '消息已读回执表' ROW_FORMAT = Dynamic;
+
+-- ============================================================
+-- 3. user_settings 用户个性化设置表
+-- ============================================================
+CREATE TABLE IF NOT EXISTS `user_settings` (
+ `user_id` varchar(100) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NOT NULL COMMENT '用户ID',
+ `setting_key` varchar(50) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NOT NULL COMMENT '设置键',
+ `setting_value` text CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NULL COMMENT '设置值',
+ `updated_at` datetime(3) NULL DEFAULT NULL COMMENT '更新时间',
+ PRIMARY KEY (`user_id`, `setting_key`) USING BTREE
+) ENGINE = InnoDB CHARACTER SET = utf8mb4 COLLATE = utf8mb4_0900_ai_ci COMMENT = '用户个性化设置表' ROW_FORMAT = Dynamic;
+
+SET FOREIGN_KEY_CHECKS = 1;
+
+-- 迁移完成
+SELECT '20260707_upgrade_p1_p2 migration finished' AS status;
diff --git a/migrations/02_20260811_fix_fk_and_room_members.sql b/migrations/02_20260811_fix_fk_and_room_members.sql
new file mode 100644
index 0000000..9bde332
--- /dev/null
+++ b/migrations/02_20260811_fix_fk_and_room_members.sql
@@ -0,0 +1,79 @@
+/*
+ * ----------------------------------------------------------------
+ * 迁移脚本: 02_20260811_fix_fk_and_room_members.sql
+ * 目的:
+ * 1. 确保 room_members 表存在(历史上该表未纳入 AutoMigrate,
+ * 仅靠全量 SQL 建库;老环境如果缺表会导致群成员相关功能全部失败)。
+ * 2. 删除 chat_conversations 上遗留的外键约束(如果存在)。
+ * 背景: target_id 既可能是用户ID(私聊)也可能是群ID(群聊),
+ * 任何指向 users 表的外键都是错误设计。代码层历史上靠
+ * "SET FOREIGN_KEY_CHECKS=0" 运行时绕过,该做法在连接池下不可靠
+ * 且会把关闭外键检查的状态泄漏给其他请求,现已从代码中移除,
+ * 因此必须保证数据库中也没有这条外键。
+ *
+ * 执行方式:
+ * mysql -u root -p nl_im_plus < migrations/02_20260811_fix_fk_and_room_members.sql
+ *
+ * 回滚方式:
+ * 本脚本只删错误约束/补缺失表,无需回滚。
+ * ----------------------------------------------------------------
+ */
+
+SET NAMES utf8mb4;
+
+-- 请按实际库名修改
+USE `nl_im_plus`;
+
+-- ============================================================
+-- 1. room_members 房间成员表(不存在则创建,结构与 nl_im_plus.sql 一致)
+-- ============================================================
+CREATE TABLE IF NOT EXISTS `room_members` (
+ `room_id` varchar(100) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NOT NULL COMMENT '房间ID',
+ `user_id` varchar(100) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NOT NULL COMMENT '用户ID',
+ `role` tinyint(1) NOT NULL DEFAULT 0 COMMENT '成员角色(0=成员,1=管理员,2=群主)',
+ `joined_at` datetime(3) NULL DEFAULT NULL COMMENT '加入时间',
+ `muted_until` datetime NULL DEFAULT NULL COMMENT '禁言到期时间',
+ PRIMARY KEY (`room_id`, `user_id`) USING BTREE,
+ INDEX `idx_room_members_user_id`(`user_id` ASC) USING BTREE
+) ENGINE = InnoDB CHARACTER SET = utf8mb4 COLLATE = utf8mb4_0900_ai_ci COMMENT = '房间成员表,存储房间与用户的关联关系' ROW_FORMAT = Dynamic;
+
+-- ============================================================
+-- 2. 删除 chat_conversations 上所有遗留外键(MySQL 无 DROP FOREIGN KEY IF EXISTS,
+-- 这里通过 information_schema 动态查找后逐条删除)
+-- ============================================================
+DROP PROCEDURE IF EXISTS `drop_chat_conversations_fks`;
+
+DELIMITER $$
+CREATE PROCEDURE `drop_chat_conversations_fks`()
+BEGIN
+ DECLARE done INT DEFAULT FALSE;
+ DECLARE fk_name VARCHAR(255);
+ -- 查出 chat_conversations 表上的全部外键约束名
+ DECLARE fk_cursor CURSOR FOR
+ SELECT CONSTRAINT_NAME
+ FROM information_schema.TABLE_CONSTRAINTS
+ WHERE TABLE_SCHEMA = DATABASE()
+ AND TABLE_NAME = 'chat_conversations'
+ AND CONSTRAINT_TYPE = 'FOREIGN KEY';
+ DECLARE CONTINUE HANDLER FOR NOT FOUND SET done = TRUE;
+
+ OPEN fk_cursor;
+ read_loop: LOOP
+ FETCH fk_cursor INTO fk_name;
+ IF done THEN
+ LEAVE read_loop;
+ END IF;
+ SET @sql = CONCAT('ALTER TABLE `chat_conversations` DROP FOREIGN KEY `', fk_name, '`');
+ PREPARE stmt FROM @sql;
+ EXECUTE stmt;
+ DEALLOCATE PREPARE stmt;
+ END LOOP;
+ CLOSE fk_cursor;
+END$$
+DELIMITER ;
+
+CALL `drop_chat_conversations_fks`();
+DROP PROCEDURE IF EXISTS `drop_chat_conversations_fks`;
+
+-- 迁移完成
+SELECT '20260811_fix_fk_and_room_members migration finished' AS status;
diff --git a/migrations/03_20260811_room_member_nickname.sql b/migrations/03_20260811_room_member_nickname.sql
new file mode 100644
index 0000000..cb1114c
--- /dev/null
+++ b/migrations/03_20260811_room_member_nickname.sql
@@ -0,0 +1,43 @@
+/*
+ * ----------------------------------------------------------------
+ * 迁移脚本: 03_20260811_room_member_nickname.sql
+ * 目的:
+ * 为 room_members 表补充 nickname 列(群名片/群内显示昵称)。
+ * 背景: 模型中 Nickname 原先是 gorm:"-" 非持久化字段,后端从未写入,
+ * 前端"我在群里的昵称"功能(POST /groups/:room_id/nickname)一直不可用。
+ *
+ * 执行方式:
+ * mysql -u root -p nl_im_plus < migrations/03_20260811_room_member_nickname.sql
+ *
+ * 回滚方式:
+ * ALTER TABLE `room_members` DROP COLUMN `nickname`;
+ * ----------------------------------------------------------------
+ */
+
+SET NAMES utf8mb4;
+
+-- 请按实际库名修改
+USE `nl_im_plus`;
+
+-- MySQL 不支持 ADD COLUMN IF NOT EXISTS,通过 information_schema 判断后再执行,
+-- 保证脚本可重复运行
+SET @col_exists = (
+ SELECT COUNT(*)
+ FROM information_schema.COLUMNS
+ WHERE TABLE_SCHEMA = DATABASE()
+ AND TABLE_NAME = 'room_members'
+ AND COLUMN_NAME = 'nickname'
+);
+
+SET @sql = IF(
+ @col_exists = 0,
+ 'ALTER TABLE `room_members` ADD COLUMN `nickname` varchar(100) NOT NULL DEFAULT '''' COMMENT ''群名片(群内显示昵称)'' AFTER `muted_until`',
+ 'SELECT ''nickname column already exists'' AS notice'
+);
+
+PREPARE stmt FROM @sql;
+EXECUTE stmt;
+DEALLOCATE PREPARE stmt;
+
+-- 迁移完成
+SELECT '20260811_room_member_nickname migration finished' AS status;
diff --git a/migrations/04_20260811_message_deletions.sql b/migrations/04_20260811_message_deletions.sql
new file mode 100644
index 0000000..3cef0d4
--- /dev/null
+++ b/migrations/04_20260811_message_deletions.sql
@@ -0,0 +1,36 @@
+/*
+ * ----------------------------------------------------------------
+ * 迁移脚本: 04_20260811_message_deletions.sql
+ * 目的:
+ * 新增 message_deletions 消息删除记录表,支持"删除仅对我生效"语义。
+ * 背景: 移动端聊天页的"删除"操作此前只把消息从本地数组移除,
+ * 下次拉取历史消息时又会重新出现(假删除)。参照主流 IM 语义,
+ * 删除应只影响操作者自己的可见性(撤回才影响双方),因此按
+ * (message_id, user_id) 记录删除关系,历史消息查询时排除。
+ *
+ * 执行方式:
+ * mysql -u root -p nl_im_plus < migrations/04_20260811_message_deletions.sql
+ *
+ * 回滚方式:
+ * DROP TABLE IF EXISTS `message_deletions`;
+ * ----------------------------------------------------------------
+ */
+
+SET NAMES utf8mb4;
+
+-- 请按实际库名修改
+USE `nl_im_plus`;
+
+CREATE TABLE IF NOT EXISTS `message_deletions` (
+ `id` bigint UNSIGNED NOT NULL AUTO_INCREMENT COMMENT '自增主键',
+ `message_id` bigint UNSIGNED NOT NULL COMMENT '被删除的消息ID',
+ `user_id` varchar(100) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NOT NULL COMMENT '执行删除的用户ID(删除仅对该用户生效)',
+ `created_at` datetime(3) NULL DEFAULT NULL COMMENT '删除时间',
+ PRIMARY KEY (`id`) USING BTREE,
+ -- 唯一索引防止同一用户对同一消息重复插入删除记录
+ UNIQUE INDEX `uk_message_user`(`message_id` ASC, `user_id` ASC) USING BTREE,
+ INDEX `idx_message_deletions_user_id`(`user_id` ASC) USING BTREE
+) ENGINE = InnoDB CHARACTER SET = utf8mb4 COLLATE = utf8mb4_0900_ai_ci COMMENT = '消息删除记录表,记录哪个用户删除了哪条消息(仅影响本人可见性)' ROW_FORMAT = Dynamic;
+
+-- 迁移完成
+SELECT '20260811_message_deletions migration finished' AS status;
diff --git a/migrations/05_20260811_moment_like_unique.sql b/migrations/05_20260811_moment_like_unique.sql
new file mode 100644
index 0000000..0af0779
--- /dev/null
+++ b/migrations/05_20260811_moment_like_unique.sql
@@ -0,0 +1,51 @@
+/*
+ * ----------------------------------------------------------------
+ * 迁移脚本: 05_20260811_moment_like_unique.sql
+ * 目的:
+ * 为 moment_likes 表补充 (moment_id, user_id) 唯一索引。
+ * 背景: 点赞接口是"先查是否已赞再插入",并发双击时两个请求都能
+ * 通过检查,导致同一用户对同一动态插入两条点赞记录、like_count 多加。
+ * 加唯一索引后由数据库兜底,代码层把唯一冲突当作"已点赞"幂等处理。
+ *
+ * 执行方式:
+ * mysql -u root -p nl_im_plus < migrations/05_20260811_moment_like_unique.sql
+ *
+ * 回滚方式:
+ * ALTER TABLE `moment_likes` DROP INDEX `uk_moment_user`;
+ * ----------------------------------------------------------------
+ */
+
+SET NAMES utf8mb4;
+
+-- 请按实际库名修改
+USE `nl_im_plus`;
+
+-- 1. 先清理历史重复数据(保留每组最早的一条),否则加唯一索引会失败
+DELETE ml FROM `moment_likes` ml
+INNER JOIN `moment_likes` ml2
+ ON ml.moment_id = ml2.moment_id
+ AND ml.user_id = ml2.user_id
+ AND ml.id > ml2.id;
+
+-- 2. 修正因重复点赞导致虚高的 like_count(以实际点赞记录数为准)
+UPDATE `moments` m
+SET m.like_count = (
+ SELECT COUNT(*) FROM `moment_likes` ml WHERE ml.moment_id = m.id
+);
+
+-- 3. 添加唯一索引(幂等:已存在则跳过)
+SET @index_exists = (
+ SELECT COUNT(*) FROM information_schema.STATISTICS
+ WHERE TABLE_SCHEMA = DATABASE()
+ AND TABLE_NAME = 'moment_likes'
+ AND INDEX_NAME = 'uk_moment_user'
+);
+SET @sql = IF(@index_exists = 0,
+ 'ALTER TABLE `moment_likes` ADD UNIQUE INDEX `uk_moment_user`(`moment_id`, `user_id`)',
+ 'SELECT ''uk_moment_user already exists'' AS notice');
+PREPARE stmt FROM @sql;
+EXECUTE stmt;
+DEALLOCATE PREPARE stmt;
+
+-- 迁移完成
+SELECT '20260811_moment_like_unique migration finished' AS status;
diff --git a/migrations/06_20260812_user_moment_cover.sql b/migrations/06_20260812_user_moment_cover.sql
new file mode 100644
index 0000000..2de4ae2
--- /dev/null
+++ b/migrations/06_20260812_user_moment_cover.sql
@@ -0,0 +1,44 @@
+/*
+ * ----------------------------------------------------------------
+ * 迁移脚本: 06_20260812_user_moment_cover.sql
+ * 目的:
+ * 为 users 表补充 moment_cover 列(朋友圈顶部封面图 URL)。
+ * 背景: 朋友圈页顶部封面原先是前端硬编码的外链图片,用户无法自行更换。
+ * 现改为存储在用户资料中,通过 POST /user/update 更新(白名单已放行),
+ * 查看他人朋友圈时通过 GET /user/:id 获取对方封面。
+ *
+ * 执行方式:
+ * mysql -u root -p nl_im_plus < migrations/06_20260812_user_moment_cover.sql
+ *
+ * 回滚方式:
+ * ALTER TABLE `users` DROP COLUMN `moment_cover`;
+ * ----------------------------------------------------------------
+ */
+
+SET NAMES utf8mb4;
+
+-- 请按实际库名修改
+USE `nl_im_plus`;
+
+-- MySQL 不支持 ADD COLUMN IF NOT EXISTS,通过 information_schema 判断后再执行,
+-- 保证脚本可重复运行
+SET @col_exists = (
+ SELECT COUNT(*)
+ FROM information_schema.COLUMNS
+ WHERE TABLE_SCHEMA = DATABASE()
+ AND TABLE_NAME = 'users'
+ AND COLUMN_NAME = 'moment_cover'
+);
+
+SET @sql = IF(
+ @col_exists = 0,
+ 'ALTER TABLE `users` ADD COLUMN `moment_cover` varchar(500) NOT NULL DEFAULT '''' COMMENT ''朋友圈封面图URL'' AFTER `region`',
+ 'SELECT ''moment_cover column already exists'' AS notice'
+);
+
+PREPARE stmt FROM @sql;
+EXECUTE stmt;
+DEALLOCATE PREPARE stmt;
+
+-- 迁移完成
+SELECT '20260812_user_moment_cover migration finished' AS status;
diff --git a/migrations/07_20260812_ai_bots.sql b/migrations/07_20260812_ai_bots.sql
new file mode 100644
index 0000000..b152669
--- /dev/null
+++ b/migrations/07_20260812_ai_bots.sql
@@ -0,0 +1,77 @@
+/*
+ * ----------------------------------------------------------------
+ * 迁移脚本: 07_20260812_ai_bots.sql
+ * 目的:
+ * AI 机器人功能:
+ * 1. users 表增加 is_bot 列(机器人复用 users 表作为虚拟用户,
+ * 消息/群成员/会话等现有链路零改造);
+ * 2. 新建 ai_configs 表(AI 提供商配置,全局一条,仅管理员 id=1 维护);
+ * 3. 新建 ai_bots 表(机器人定义:名称/头像/角色设定,关联虚拟用户ID)。
+ *
+ * 执行方式:
+ * mysql -u root -p nl_im_plus < migrations/07_20260812_ai_bots.sql
+ *
+ * 回滚方式:
+ * ALTER TABLE `users` DROP COLUMN `is_bot`;
+ * DROP TABLE IF EXISTS `ai_configs`;
+ * DROP TABLE IF EXISTS `ai_bots`;
+ *
+ * 本脚本可重复执行:列/表已存在时跳过,不会报错。
+ * ----------------------------------------------------------------
+ */
+
+SET NAMES utf8mb4;
+
+-- 请按实际库名修改
+USE `nl_im_plus`;
+
+-- ============================================================
+-- 1. users 表增加 is_bot 列(已存在则跳过)
+-- ============================================================
+SET @col_exists = (
+ SELECT COUNT(*) FROM information_schema.COLUMNS
+ WHERE TABLE_SCHEMA = DATABASE()
+ AND TABLE_NAME = 'users'
+ AND COLUMN_NAME = 'is_bot'
+);
+SET @sql = IF(
+ @col_exists = 0,
+ 'ALTER TABLE `users` ADD COLUMN `is_bot` tinyint(1) NOT NULL DEFAULT 0 COMMENT ''是否AI机器人'' AFTER `moment_cover`',
+ 'SELECT ''is_bot column already exists'' AS notice'
+);
+PREPARE stmt FROM @sql;
+EXECUTE stmt;
+DEALLOCATE PREPARE stmt;
+
+-- ============================================================
+-- 2. ai_configs AI提供商配置表
+-- ============================================================
+CREATE TABLE IF NOT EXISTS `ai_configs` (
+ `id` bigint UNSIGNED NOT NULL AUTO_INCREMENT COMMENT '主键ID',
+ `provider` varchar(50) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NULL DEFAULT NULL COMMENT '提供商标识(openai/deepseek/dashscope/moonshot/ollama/custom)',
+ `base_url` varchar(500) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NULL DEFAULT NULL COMMENT 'API基础地址(空用默认)',
+ `api_key` varchar(500) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NULL DEFAULT NULL COMMENT 'API密钥',
+ `model` varchar(100) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NULL DEFAULT NULL COMMENT '模型名称',
+ `enabled` tinyint(1) NOT NULL DEFAULT 0 COMMENT '是否启用',
+ `created_at` datetime(3) NULL DEFAULT NULL COMMENT '创建时间',
+ `updated_at` datetime(3) NULL DEFAULT NULL COMMENT '更新时间',
+ PRIMARY KEY (`id`) USING BTREE
+) ENGINE = InnoDB CHARACTER SET = utf8mb4 COLLATE = utf8mb4_0900_ai_ci COMMENT = 'AI提供商配置表(全局一条,仅管理员维护)' ROW_FORMAT = Dynamic;
+
+-- ============================================================
+-- 3. ai_bots 机器人定义表
+-- ============================================================
+CREATE TABLE IF NOT EXISTS `ai_bots` (
+ `id` bigint UNSIGNED NOT NULL AUTO_INCREMENT COMMENT '主键ID',
+ `user_id` varchar(100) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NOT NULL COMMENT '关联虚拟用户ID(users.id, bot_前缀)',
+ `name` varchar(100) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NULL DEFAULT NULL COMMENT '机器人名称(群聊@名称触发)',
+ `avatar` varchar(500) CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NULL DEFAULT NULL COMMENT '机器人头像',
+ `role_prompt` text CHARACTER SET utf8mb4 COLLATE utf8mb4_0900_ai_ci NULL COMMENT '角色设定(system提示词)',
+ `enabled` tinyint(1) NOT NULL DEFAULT 1 COMMENT '是否启用',
+ `created_at` datetime(3) NULL DEFAULT NULL COMMENT '创建时间',
+ `updated_at` datetime(3) NULL DEFAULT NULL COMMENT '更新时间',
+ PRIMARY KEY (`id`) USING BTREE,
+ UNIQUE INDEX `uk_ai_bots_user_id`(`user_id` ASC) USING BTREE
+) ENGINE = InnoDB CHARACTER SET = utf8mb4 COLLATE = utf8mb4_0900_ai_ci COMMENT = 'AI机器人定义表' ROW_FORMAT = Dynamic;
+
+SELECT '07_20260812_ai_bots migration finished' AS result;