From cb15328e5d73bb45605d35b09950a278a5326301 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=9D=8E=E7=90=A6?= Date: Mon, 24 Aug 2026 15:29:53 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BF=AE=E5=A4=8D=E4=BA=86=E4=B8=80=E4=BA=9B?= =?UTF-8?q?=E9=97=AE=E9=A2=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .idea/go.imports.xml | 11 + cmd/server/main.go | 389 +++++++++- configs/config.yaml | 27 +- go.mod | 5 +- go.sum | 2 + interaction-audit.html | 680 ++++++++++++++++++ internal/api/ai_handler.go | 249 +++++++ internal/api/auth_handler.go | 46 +- internal/api/call_handler.go | 49 +- internal/api/contact_handler.go | 59 +- internal/api/handler.go | 112 ++- internal/api/qrcode_handler.go | 200 ++++++ internal/api/room_handler.go | 175 ++--- internal/api/settings_handler.go | 138 +++- internal/api/user_handler.go | 73 +- internal/manager/client_manager.go | 83 ++- internal/mediaserver/ffmpeg_rtmp_receiver.go | 45 +- internal/mediaserver/ffmpeg_transcoder.go | 156 ++-- internal/mediaserver/room.go | 36 +- internal/mediaserver/rtmp.go | 43 +- internal/mediaserver/rtmp_amf0.go | 46 +- internal/mediaserver/sfu.go | 155 ++-- internal/mediaserver/ws_rtmp_proxy.go | 115 +-- internal/middleware/auth.go | 36 +- internal/middleware/request_log.go | 135 +++- internal/model/types.go | 128 +++- internal/service/ai_bot_service.go | 529 ++++++++++++++ internal/service/ai_provider.go | 170 +++++ internal/service/attachment_service.go | 9 +- internal/service/auth_service.go | 55 +- internal/service/chat_service.go | 623 ++++++++++++++-- internal/service/contact_service.go | 318 +++++--- internal/service/conversation_service.go | 133 ++-- internal/service/moment_service.go | 248 ++++++- internal/service/qrcode_login_service.go | 288 ++++++++ internal/service/room_service.go | 320 ++++++--- internal/utils/db.go | 29 + internal/utils/jwt.go | 73 +- internal/utils/response.go | 36 +- internal/ws/worker.go | 22 +- migrations/01_20260707_upgrade_p1_p2.sql | 104 +++ .../02_20260811_fix_fk_and_room_members.sql | 79 ++ .../03_20260811_room_member_nickname.sql | 43 ++ migrations/04_20260811_message_deletions.sql | 36 + migrations/05_20260811_moment_like_unique.sql | 51 ++ migrations/06_20260812_user_moment_cover.sql | 44 ++ migrations/07_20260812_ai_bots.sql | 77 ++ 47 files changed, 5514 insertions(+), 966 deletions(-) create mode 100644 .idea/go.imports.xml create mode 100644 interaction-audit.html create mode 100644 internal/api/ai_handler.go create mode 100644 internal/api/qrcode_handler.go create mode 100644 internal/service/ai_bot_service.go create mode 100644 internal/service/ai_provider.go create mode 100644 internal/service/qrcode_login_service.go create mode 100644 internal/utils/db.go create mode 100644 migrations/01_20260707_upgrade_p1_p2.sql create mode 100644 migrations/02_20260811_fix_fk_and_room_members.sql create mode 100644 migrations/03_20260811_room_member_nickname.sql create mode 100644 migrations/04_20260811_message_deletions.sql create mode 100644 migrations/05_20260811_moment_like_unique.sql create mode 100644 migrations/06_20260812_user_moment_cover.sql create mode 100644 migrations/07_20260812_ai_bots.sql 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) 跨行 + 忽略大小写;容忍 <title lang="..."> 等属性) +var titleTagRe = regexp.MustCompile(`(?is)<title[^>]*>(.*?)`) + +// 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)<title>(.*?)`) - 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;