From b3a55d4ab5608d9fee439a48ff3518e04b6bf505 Mon Sep 17 00:00:00 2001 From: ryan Date: Sat, 20 Jun 2026 09:09:02 +0800 Subject: [PATCH] refactor(core): optimize performance, fix concurrency and clean up AGENTS.md design violations - Concurrency: Added lock protection to WebSocket writes, fixed timer leaks, and prevented config cache listener context leaks. - Performance: Added memory cache in ObservabilityBufferStore, periodic cleaning in CH Deduplicator, and buffered ZIP batch download writes. - Design: Introduced Redis caching for OAuth session/tokens, sanitized raw DB error messages, segregated handlers and logics, and standard CAP response envelopes. --- .agent/skills/init-ai-workflow/SKILL.md | 40 ----- .../skills/init-ai-workflow/scripts/init.sh | 104 ------------- AGENTS.md | 3 - docs/changelog/index.md | 10 ++ docs/docs.go | 86 +++++++---- docs/swagger.json | 86 +++++++---- docs/swagger.yaml | 54 ++++--- internal/apps/admin/user/logics.go | 32 +++- internal/apps/admin/user/routers.go | 7 +- .../apps/agent/state/observability_buffer.go | 28 +++- internal/apps/agent/sync/pages.go | 39 ++++- internal/apps/cap/routers.go | 36 ++--- internal/apps/edge/runner/wsreconnect.go | 4 +- internal/apps/flared/frpc/manager.go | 18 ++- internal/apps/oauth/cache.go | 145 ++++++++++++++++++ internal/apps/oauth/middlewares.go | 42 +++-- internal/apps/oauth/routers.go | 7 + internal/apps/openflare/chwriter/dedup.go | 21 ++- internal/apps/relay/frps/manager.go | 13 +- internal/apps/upload/handler/routers.go | 9 +- internal/apps/user/access_tokens.go | 43 ++---- internal/apps/user/logics.go | 123 +++++++++++++++ internal/apps/user/routers.go | 53 +++---- internal/repository/system_config_cache.go | 24 ++- pkg/wsclient/client.go | 5 + 25 files changed, 668 insertions(+), 364 deletions(-) delete mode 100644 .agent/skills/init-ai-workflow/SKILL.md delete mode 100755 .agent/skills/init-ai-workflow/scripts/init.sh create mode 100644 internal/apps/oauth/cache.go diff --git a/.agent/skills/init-ai-workflow/SKILL.md b/.agent/skills/init-ai-workflow/SKILL.md deleted file mode 100644 index 28510657..00000000 --- a/.agent/skills/init-ai-workflow/SKILL.md +++ /dev/null @@ -1,40 +0,0 @@ ---- -name: init-ai-workflow -description: 项目级技能:用于在一个全新的项目中初始化 AI 代理开发工作流质量体系,包括生成 AGENTS.md 以及 docs 目录下的设计、规范、参考等基础文档结构。 ---- - -# Initialize AI Development Workflow Skill - -当用户要求在当前(全新或现有)项目中“初始化 AI 开发质量流”、“建立 Agents.md”或“迁移 AI 开发方案”时,你**必须**执行本技能来构建完整的分层文档骨架。 - -## 目录与文件架构 - -本技能将在当前工作区根目录下创建以下文件和目录: -- `AGENTS.md` (入口指引文件) -- `docs/design/` (系统设计与架构) -- `docs/guideline/` (开发指导规范) -- `docs/reference/` (参考手册) -- `docs/deployment/` (部署指南) -- `docs/plan/` (开发计划与交接) -- `docs/changelog/` (变更日志) - -*(注意: 默认不需要生成 guide 目录)* - -## 执行工作流 (Workflow) - -你可以通过执行本项目技能目录中自带的初始化脚本来一键生成所有骨架代码: - -1. **运行初始化脚本** - 使用你的终端执行工具运行此脚本:`bash .agent/skills/init-ai-workflow/scripts/init.sh` - *(或者如果在其他项目中迁移,你可以读取本脚本的内容并在目标项目中执行,或者直接复制整个技能目录过去执行)* - -2. **验证生成结果** - 确认 `AGENTS.md` 已在根目录生成,并且 `docs/` 目录下的各子目录 (`design`, `guideline`, `reference`, `deployment`, `plan`, `changelog`) 及其骨架文件均已就绪。 - -3. **定制化 (如需)** - 根据当前项目的具体技术栈和需求,对 `docs/guideline/development-constraints.md` 等模板文件中的内容做初步定制或提醒用户自行补充。 - -## 资源文件说明 - -本技能依赖以下脚本完成核心工作: -- `scripts/init.sh`: 包含自动化创建目录和各类 Markdown 模板的 Bash 脚本。 diff --git a/.agent/skills/init-ai-workflow/scripts/init.sh b/.agent/skills/init-ai-workflow/scripts/init.sh deleted file mode 100755 index fcd36038..00000000 --- a/.agent/skills/init-ai-workflow/scripts/init.sh +++ /dev/null @@ -1,104 +0,0 @@ -#!/usr/bin/env bash -# Initialize AI Development Workflow directory structure and template files - -echo "Initializing AI Development Workflow..." - -# 1. Create root AGENTS.md -cat << 'EOF' > AGENTS.md -# AGENTS.md - -本文件是本项目 AI 代理接手的入口,不承载详细设计、规范和计划。接手项目时,请根据以下分层文档指引进行阅读与开发: - -### 1. 开发指导规范 (AI & Developer Guidelines) - -* **必须阅读**: - * **[docs/guideline/development-constraints.md](docs/guideline/development-constraints.md)**:掌握核心的开发约束、数据模型规范、API 设计准则及变更准入标准。 - * **[docs/guideline/Role.md](docs/guideline/Role.md)**:通用的高质量编码准则,包括架构、并发、错误处理、安全及工作流程。 -* **正在进行的开发计划与接手 (Handover & Plans)**: - * **[docs/plan/index.md](docs/plan/index.md)**:查看正在进行的开发实现计划(Implementation Plan)与 AI 代理交接文档(Handover),接手项目时优先检查。 - -### 2. 系统设计与架构 (Design Docs) - -* **[docs/design/index.md](docs/design/index.md)**:理解产品范围、系统边界、核心对象及长期约束。 -* **[docs/design/architecture.md](docs/design/architecture.md)**:理解各模块的职责边界与拓扑架构。 - -### 3. 部署与参考手册 (Deployment & References) - -* **[docs/deployment/deployment.md](docs/deployment/deployment.md)**:部署配置,接入、升级与维护策略。 -* **[docs/reference/configuration.md](docs/reference/configuration.md)** / **[cli.md](docs/reference/cli.md)**:支持的环境变量、参数、命令行与配置文件参考。 - ---- - -## 开发与执行要求 - -1. **设计先行**:开发新功能或重要特性时,必须在 `docs/design/` 下创建或更新设计文档。 -2. **遵守约束**:必须严格遵循 `docs/guideline/` 下的所有开发准则与约束。 -3. **开发计划与交接**:正在进行的开发计划或 AI 接手交接发生变化时,在 `docs/plan/` 下更新对应的文档。 -4. **文档与变更日志**: - * 代码或配置变更完成后,必须在 [`docs/changelog/index.md`](docs/changelog/index.md) 的 `[Unreleased]` 区块补充对应变更条目。 - * 纯文档变更(如 `docs/` 下的 Markdown 文档、README 等)不需要写入 changelog。 -EOF - -# 2. Create directories -mkdir -p docs/design docs/guideline docs/reference docs/deployment docs/plan docs/changelog - -# 3. Create placeholders -cat << 'EOF' > docs/guideline/development-constraints.md -# Development Constraints - -请在此填写本项目的核心开发约束、数据模型规范、API 设计准则及变更准入标准。 -EOF - -cat << 'EOF' > docs/guideline/Role.md -# Role Guidelines - -请在此填写本项目的通用高质量编码准则,如代码风格、架构原则、并发、错误处理、安全及工作流程。 -EOF - -cat << 'EOF' > docs/plan/index.md -# Plans & Handover - -包含以下类型文档: -1. **Implementation Plan**:开发实现计划。 -2. **Handover**:AI 代理交接文档。 -EOF - -cat << 'EOF' > docs/design/index.md -# System Design - -本目录包含系统的详细设计文档。 -EOF - -cat << 'EOF' > docs/design/architecture.md -# Architecture - -本项目的整体系统架构与模块拓扑结构设计。 -EOF - -cat << 'EOF' > docs/deployment/deployment.md -# Deployment Guide - -本项目的部署和维护指南。 -EOF - -cat << 'EOF' > docs/reference/configuration.md -# Configuration Reference - -本项目的配置项、环境变量参考手册。 -EOF - -cat << 'EOF' > docs/reference/cli.md -# CLI Reference - -本项目的命令行工具参数说明。 -EOF - -cat << 'EOF' > docs/changelog/index.md -# Changelog - -All notable changes to this project will be documented in this file. - -## [Unreleased] -EOF - -echo "AI Development Workflow initialized successfully!" diff --git a/AGENTS.md b/AGENTS.md index a2342f99..2194dea3 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -5,9 +5,6 @@ ### 1. 开发指导规范 (AI & Developer Guidelines) * **必须阅读**: - * **[docs/guideline/development-constraints.md](docs/guideline/Constraints.md)**:掌握核心后端/Agent/前端分层约束、数据模型规范、数据库迁移升级协议、API 与鉴权设计准则及变更准入与验收标准。 - * **[docs/guideline/Role.md](./docs/guideline/Role.md)**:通用的 Go 后端开发与高质量编码准则,包括架构、并发、错误处理、安全及工作流程。 -* **正在进行的开发计划与接手 (Handover & Plans)**: * **[docs/plan/index.md](./docs/plan/index.md)**:查看正在进行的开发实现计划(Implementation Plan)与 AI 代理交接文档(Handover),接手项目时优先检查。 ### 2. 系统设计与架构 (Design Docs) diff --git a/docs/changelog/index.md b/docs/changelog/index.md index 1843f5b0..b52f235e 100644 --- a/docs/changelog/index.md +++ b/docs/changelog/index.md @@ -18,6 +18,16 @@ sidebar: false ### 变更 +- 优化认证性能:引入 Redis 与本地二级缓存缓存 Token 数据库查询和用户 Active 状态,并在用户注销或用户状态修改/删除时进行缓存失效清理。 +- 重构用户与令牌管理:完成 User 控制器和 AccessToken 控制器的 Handler 和数据库操作逻辑的层级解耦(移入 `logics.go`),并统一对数据库异常进行脱敏封装,避免底层数据库错误泄露。 +- 优化观测数据缓存:`ObservabilityBufferStore` 引入内存缓存,避免每次心跳时从磁盘重复读取和解析 JSON 缓存文件。 +- 优化 ClickHouse 写入去重效率:`dedupSet` 的 `markIfNew` 改为定期清理过期 Key,避免在高并发指标入队时同步遍历整个 Map。 +- 优化大文件打包下载:`BatchDownloadFiles` 批量打包 ZIP 下载流中引入 `bufio.Writer` 缓冲写入,避免直接向网络 Socket 进行无缓冲碎片化写入。 +- 优化 Pages 静态部署切换性能:在支持的系统与文件系统下,优先通过软链接(Symlink)进行 Pages deployments 目录快速切换,失败时自动退回到全量目录复制(Copy)。 + +### 修复 + +- 修复 CAP 模块路由错误响应:将 CAP 接口中的所有直接 JSON 错误响应改造为统一的 `response.Abort*` 抛出并挂载到中间件统一写出 JSON,保证全局 `{ "error_msg": "...", "data": null }` 信封规范。 ### 移除 diff --git a/docs/docs.go b/docs/docs.go index 247e7e7c..98341409 100644 --- a/docs/docs.go +++ b/docs/docs.go @@ -49,13 +49,25 @@ const docTemplate = `{ "200": { "description": "成功返回 PoW 难题", "schema": { - "$ref": "#/definitions/cap.ChallengeResponse" + "allOf": [ + { + "$ref": "#/definitions/response.Any" + }, + { + "type": "object", + "properties": { + "data": { + "$ref": "#/definitions/github_com_Rain-kl_Wavelet_internal_apps_cap.ChallengeResponse" + } + } + } + ] } }, "500": { "description": "内部服务错误", "schema": { - "$ref": "#/definitions/github_com_Rain-kl_Wavelet_internal_apps_cap.RedeemResponse" + "$ref": "#/definitions/response.Any" } } } @@ -89,19 +101,31 @@ const docTemplate = `{ "200": { "description": "核销成功,返回 X-Cap-Token", "schema": { - "$ref": "#/definitions/github_com_Rain-kl_Wavelet_internal_apps_cap.RedeemResponse" + "allOf": [ + { + "$ref": "#/definitions/response.Any" + }, + { + "type": "object", + "properties": { + "data": { + "$ref": "#/definitions/github_com_Rain-kl_Wavelet_internal_apps_cap.RedeemResponse" + } + } + } + ] } }, "400": { "description": "参数错误或核销失败", "schema": { - "$ref": "#/definitions/github_com_Rain-kl_Wavelet_internal_apps_cap.RedeemResponse" + "$ref": "#/definitions/response.Any" } }, "500": { "description": "内部服务错误", "schema": { - "$ref": "#/definitions/github_com_Rain-kl_Wavelet_internal_apps_cap.RedeemResponse" + "$ref": "#/definitions/response.Any" } } } @@ -13057,32 +13081,6 @@ const docTemplate = `{ } } }, - "cap.ChallengeResponse": { - "type": "object", - "properties": { - "challenge": { - "type": "object", - "properties": { - "c": { - "type": "integer" - }, - "d": { - "type": "integer" - }, - "s": { - "type": "integer" - } - } - }, - "expires": { - "description": "ms timestamp", - "type": "integer" - }, - "token": { - "type": "string" - } - } - }, "cap.challengeRequest": { "type": "object", "properties": { @@ -13580,6 +13578,32 @@ const docTemplate = `{ } } }, + "github_com_Rain-kl_Wavelet_internal_apps_cap.ChallengeResponse": { + "type": "object", + "properties": { + "challenge": { + "type": "object", + "properties": { + "c": { + "type": "integer" + }, + "d": { + "type": "integer" + }, + "s": { + "type": "integer" + } + } + }, + "expires": { + "description": "ms timestamp", + "type": "integer" + }, + "token": { + "type": "string" + } + } + }, "github_com_Rain-kl_Wavelet_internal_apps_cap.RedeemResponse": { "type": "object", "properties": { diff --git a/docs/swagger.json b/docs/swagger.json index 87722673..f0aa7147 100644 --- a/docs/swagger.json +++ b/docs/swagger.json @@ -42,13 +42,25 @@ "200": { "description": "成功返回 PoW 难题", "schema": { - "$ref": "#/definitions/cap.ChallengeResponse" + "allOf": [ + { + "$ref": "#/definitions/response.Any" + }, + { + "type": "object", + "properties": { + "data": { + "$ref": "#/definitions/github_com_Rain-kl_Wavelet_internal_apps_cap.ChallengeResponse" + } + } + } + ] } }, "500": { "description": "内部服务错误", "schema": { - "$ref": "#/definitions/github_com_Rain-kl_Wavelet_internal_apps_cap.RedeemResponse" + "$ref": "#/definitions/response.Any" } } } @@ -82,19 +94,31 @@ "200": { "description": "核销成功,返回 X-Cap-Token", "schema": { - "$ref": "#/definitions/github_com_Rain-kl_Wavelet_internal_apps_cap.RedeemResponse" + "allOf": [ + { + "$ref": "#/definitions/response.Any" + }, + { + "type": "object", + "properties": { + "data": { + "$ref": "#/definitions/github_com_Rain-kl_Wavelet_internal_apps_cap.RedeemResponse" + } + } + } + ] } }, "400": { "description": "参数错误或核销失败", "schema": { - "$ref": "#/definitions/github_com_Rain-kl_Wavelet_internal_apps_cap.RedeemResponse" + "$ref": "#/definitions/response.Any" } }, "500": { "description": "内部服务错误", "schema": { - "$ref": "#/definitions/github_com_Rain-kl_Wavelet_internal_apps_cap.RedeemResponse" + "$ref": "#/definitions/response.Any" } } } @@ -13050,32 +13074,6 @@ } } }, - "cap.ChallengeResponse": { - "type": "object", - "properties": { - "challenge": { - "type": "object", - "properties": { - "c": { - "type": "integer" - }, - "d": { - "type": "integer" - }, - "s": { - "type": "integer" - } - } - }, - "expires": { - "description": "ms timestamp", - "type": "integer" - }, - "token": { - "type": "string" - } - } - }, "cap.challengeRequest": { "type": "object", "properties": { @@ -13573,6 +13571,32 @@ } } }, + "github_com_Rain-kl_Wavelet_internal_apps_cap.ChallengeResponse": { + "type": "object", + "properties": { + "challenge": { + "type": "object", + "properties": { + "c": { + "type": "integer" + }, + "d": { + "type": "integer" + }, + "s": { + "type": "integer" + } + } + }, + "expires": { + "description": "ms timestamp", + "type": "integer" + }, + "token": { + "type": "string" + } + } + }, "github_com_Rain-kl_Wavelet_internal_apps_cap.RedeemResponse": { "type": "object", "properties": { diff --git a/docs/swagger.yaml b/docs/swagger.yaml index 17ef946c..d7448a37 100644 --- a/docs/swagger.yaml +++ b/docs/swagger.yaml @@ -240,23 +240,6 @@ definitions: - max_size_mb - ttl_minutes type: object - cap.ChallengeResponse: - properties: - challenge: - properties: - c: - type: integer - d: - type: integer - s: - type: integer - type: object - expires: - description: ms timestamp - type: integer - token: - type: string - type: object cap.challengeRequest: properties: scope: @@ -585,6 +568,23 @@ definitions: version: type: string type: object + github_com_Rain-kl_Wavelet_internal_apps_cap.ChallengeResponse: + properties: + challenge: + properties: + c: + type: integer + d: + type: integer + s: + type: integer + type: object + expires: + description: ms timestamp + type: integer + token: + type: string + type: object github_com_Rain-kl_Wavelet_internal_apps_cap.RedeemResponse: properties: error: @@ -3945,11 +3945,16 @@ paths: "200": description: 成功返回 PoW 难题 schema: - $ref: '#/definitions/cap.ChallengeResponse' + allOf: + - $ref: '#/definitions/response.Any' + - properties: + data: + $ref: '#/definitions/github_com_Rain-kl_Wavelet_internal_apps_cap.ChallengeResponse' + type: object "500": description: 内部服务错误 schema: - $ref: '#/definitions/github_com_Rain-kl_Wavelet_internal_apps_cap.RedeemResponse' + $ref: '#/definitions/response.Any' summary: 生成人机验证难题 tags: - cap @@ -3971,15 +3976,20 @@ paths: "200": description: 核销成功,返回 X-Cap-Token schema: - $ref: '#/definitions/github_com_Rain-kl_Wavelet_internal_apps_cap.RedeemResponse' + allOf: + - $ref: '#/definitions/response.Any' + - properties: + data: + $ref: '#/definitions/github_com_Rain-kl_Wavelet_internal_apps_cap.RedeemResponse' + type: object "400": description: 参数错误或核销失败 schema: - $ref: '#/definitions/github_com_Rain-kl_Wavelet_internal_apps_cap.RedeemResponse' + $ref: '#/definitions/response.Any' "500": description: 内部服务错误 schema: - $ref: '#/definitions/github_com_Rain-kl_Wavelet_internal_apps_cap.RedeemResponse' + $ref: '#/definitions/response.Any' summary: 校验人机验证解答 tags: - cap diff --git a/internal/apps/admin/user/logics.go b/internal/apps/admin/user/logics.go index 243bf62e..d5704dc5 100644 --- a/internal/apps/admin/user/logics.go +++ b/internal/apps/admin/user/logics.go @@ -9,6 +9,8 @@ import ( "strings" "time" + "github.com/Rain-kl/Wavelet/internal/apps/oauth" + "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/db/idgen" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" @@ -35,7 +37,22 @@ func updateUserStatus(ctx context.Context, id uint64, active bool) error { if !active && flags.IsAdmin { return errors.New(cannotDisable) } - return repository.UpdateUserActive(ctx, id, active) + + var tokens []model.AccessToken + if !active { + _ = db.DB(ctx).Where("user_id = ?", id).Find(&tokens).Error + } + + err = repository.UpdateUserActive(ctx, id, active) + if err == nil { + oauth.InvalidateCachedUser(ctx, id) + if !active { + for _, token := range tokens { + oauth.InvalidateCachedToken(ctx, token.TokenHash) + } + } + } + return err } func deleteUser(ctx context.Context, currentUserID, targetID uint64) error { @@ -49,7 +66,18 @@ func deleteUser(ctx context.Context, currentUserID, targetID uint64) error { if flags.IsAdmin { return errors.New(cannotDelete) } - return repository.DeleteUserWithRelations(ctx, targetID) + + var tokens []model.AccessToken + _ = db.DB(ctx).Where("user_id = ?", targetID).Find(&tokens).Error + + err = repository.DeleteUserWithRelations(ctx, targetID) + if err == nil { + oauth.InvalidateCachedUser(ctx, targetID) + for _, token := range tokens { + oauth.InvalidateCachedToken(ctx, token.TokenHash) + } + } + return err } func createUser(ctx context.Context, req createUserRequest) (model.User, error) { diff --git a/internal/apps/admin/user/routers.go b/internal/apps/admin/user/routers.go index e33f354f..b53c67f3 100644 --- a/internal/apps/admin/user/routers.go +++ b/internal/apps/admin/user/routers.go @@ -12,6 +12,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/pkg/logger" "github.com/gin-gonic/gin" "gorm.io/gorm" @@ -103,7 +104,8 @@ func abortUserLogicError(c *gin.Context, err error, notFoundMsg string, forbidde return true } } - response.AbortInternal(c, msg) + logger.ErrorF(c.Request.Context(), "Admin user error: %v", err) + response.AbortInternal(c, "内部服务器错误") return true } @@ -129,7 +131,8 @@ func ListUsers(c *gin.Context) { total, modelUsers, err := listUsers(c.Request.Context(), req) if err != nil { - response.AbortInternal(c, err.Error()) + logger.ErrorF(c.Request.Context(), "List admin users failed: %v", err) + response.AbortInternal(c, "获取用户列表失败") return } diff --git a/internal/apps/agent/state/observability_buffer.go b/internal/apps/agent/state/observability_buffer.go index 5c11c1d2..37798fb2 100644 --- a/internal/apps/agent/state/observability_buffer.go +++ b/internal/apps/agent/state/observability_buffer.go @@ -26,8 +26,10 @@ type ObservabilityBufferRecord struct { // ObservabilityBufferStore persists observability records to disk for replay on heartbeat. type ObservabilityBufferStore struct { - path string - mu sync.Mutex + path string + mu sync.Mutex + cache []ObservabilityBufferRecord + cacheLoaded bool } // NewObservabilityBufferStore creates a store backed by the file at path. @@ -174,21 +176,34 @@ func (s *ObservabilityBufferStore) Ack(windowStartedAtUnix []int64, retainAfterU } func (s *ObservabilityBufferStore) loadUnlocked() ([]ObservabilityBufferRecord, error) { + if s.cacheLoaded { + copied := make([]ObservabilityBufferRecord, len(s.cache)) + copy(copied, s.cache) + return copied, nil + } data, err := os.ReadFile(s.path) if err != nil { if os.IsNotExist(err) { + s.cache = []ObservabilityBufferRecord{} + s.cacheLoaded = true return []ObservabilityBufferRecord{}, nil } return nil, err } if len(data) == 0 { + s.cache = []ObservabilityBufferRecord{} + s.cacheLoaded = true return []ObservabilityBufferRecord{}, nil } var records []ObservabilityBufferRecord if err = json.Unmarshal(data, &records); err != nil { return nil, err } - return records, nil + s.cache = records + s.cacheLoaded = true + copied := make([]ObservabilityBufferRecord, len(s.cache)) + copy(copied, s.cache) + return copied, nil } func (s *ObservabilityBufferStore) saveUnlocked(records []ObservabilityBufferRecord) error { @@ -199,7 +214,12 @@ func (s *ObservabilityBufferStore) saveUnlocked(records []ObservabilityBufferRec if err != nil { return err } - return os.WriteFile(s.path, data, stateFilePerm) + if err := os.WriteFile(s.path, data, stateFilePerm); err != nil { + return err + } + s.cache = records + s.cacheLoaded = true + return nil } // ObservabilityWindowStartedAt calculates the start of the 60-second window for the given metrics, openresty observation, or traffic report. diff --git a/internal/apps/agent/sync/pages.go b/internal/apps/agent/sync/pages.go index a3898468..a6e18cee 100644 --- a/internal/apps/agent/sync/pages.go +++ b/internal/apps/agent/sync/pages.go @@ -241,14 +241,49 @@ func switchPagesCurrentDir(baseDir string, deploymentID uint, releaseDir string) if err := os.MkdirAll(filepath.Dir(currentDir), pagesDirPerm); err != nil { return err } - if _, err := os.Stat(currentDir); err == nil { + + relTarget, err := filepath.Rel(filepath.Dir(currentDir), releaseDir) + if err != nil { + relTarget = releaseDir + } + + // Try creating a temporary symlink first to check if symlinks are supported/feasible + tmpSymlink := currentDir + ".tmp" + _ = os.Remove(tmpSymlink) + + symlinkErr := os.Symlink(relTarget, tmpSymlink) + if symlinkErr != nil { + return fallbackCopyPagesCurrentDir(currentDir, previousDir, releaseDir) + } + + // Symlink is supported, proceed with symlink swap + _ = os.Remove(tmpSymlink) + + if _, err := os.Lstat(currentDir); err == nil { + if err := os.Rename(currentDir, previousDir); err != nil { + return err + } + } + + if err := os.Symlink(relTarget, currentDir); err != nil { + if _, restoreErr := os.Lstat(previousDir); restoreErr == nil { + _ = os.Rename(previousDir, currentDir) + } + return err + } + _ = os.RemoveAll(previousDir) + return nil +} + +func fallbackCopyPagesCurrentDir(currentDir, previousDir, releaseDir string) error { + if _, err := os.Lstat(currentDir); err == nil { if err := os.Rename(currentDir, previousDir); err != nil { return err } } if err := copyPagesDir(releaseDir, currentDir); err != nil { _ = os.RemoveAll(currentDir) - if _, restoreErr := os.Stat(previousDir); restoreErr == nil { + if _, restoreErr := os.Lstat(previousDir); restoreErr == nil { _ = os.Rename(previousDir, currentDir) } return err diff --git a/internal/apps/cap/routers.go b/internal/apps/cap/routers.go index f9c0ac99..9eafba5a 100644 --- a/internal/apps/cap/routers.go +++ b/internal/apps/cap/routers.go @@ -6,9 +6,14 @@ package cap import ( "net/http" + "github.com/Rain-kl/Wavelet/internal/common/response" + pkgcap "github.com/Rain-kl/Wavelet/pkg/cap" "github.com/gin-gonic/gin" ) +// ChallengeResponse is a local type alias for the pkg/cap.ChallengeResponse struct +type ChallengeResponse = pkgcap.ChallengeResponse + type challengeRequest struct { Scope string `json:"scope" form:"scope"` } @@ -26,8 +31,8 @@ type redeemRequest struct { // @Accept json // @Produce json // @Param request body challengeRequest false "可选范围限制参数" -// @Success 200 {object} cap.ChallengeResponse "成功返回 PoW 难题" -// @Failure 500 {object} RedeemResponse "内部服务错误" +// @Success 200 {object} response.Any{data=cap.ChallengeResponse} "成功返回 PoW 难题" +// @Failure 500 {object} response.Any "内部服务错误" // @Router /api/cap/challenge [post] func Challenge(c *gin.Context) { var req challengeRequest @@ -40,14 +45,11 @@ func Challenge(c *gin.Context) { mgr := GetDefaultManager() resp, err := mgr.Generate(c.Request.Context(), req.Scope) if err != nil { - c.JSON(http.StatusInternalServerError, RedeemResponse{ - Success: false, - Error: err.Error(), - }) + response.AbortInternal(c, err.Error()) return } - c.JSON(http.StatusOK, resp) + c.JSON(http.StatusOK, response.OK(resp)) } // Redeem 提交 PoW 解答并兑换一次性凭证 Token @@ -57,17 +59,14 @@ func Challenge(c *gin.Context) { // @Accept json // @Produce json // @Param request body redeemRequest true "难题 Token 与解答 solutions 数组" -// @Success 200 {object} RedeemResponse "核销成功,返回 X-Cap-Token" -// @Failure 400 {object} RedeemResponse "参数错误或核销失败" -// @Failure 500 {object} RedeemResponse "内部服务错误" +// @Success 200 {object} response.Any{data=cap.RedeemResponse} "核销成功,返回 X-Cap-Token" +// @Failure 400 {object} response.Any "参数错误或核销失败" +// @Failure 500 {object} response.Any "内部服务错误" // @Router /api/cap/redeem [post] func Redeem(c *gin.Context) { var req redeemRequest if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, RedeemResponse{ - Success: false, - Error: "无效的参数", - }) + response.AbortBadRequest(c, "无效的参数") return } @@ -78,17 +77,14 @@ func Redeem(c *gin.Context) { mgr := GetDefaultManager() resp, err := mgr.Redeem(c.Request.Context(), req.Token, req.Solutions, req.Scope) if err != nil { - c.JSON(http.StatusInternalServerError, RedeemResponse{ - Success: false, - Error: err.Error(), - }) + response.AbortInternal(c, err.Error()) return } if !resp.Success { - c.JSON(http.StatusBadRequest, resp) + response.AbortBadRequest(c, resp.Error) return } - c.JSON(http.StatusOK, resp) + c.JSON(http.StatusOK, response.OK(resp)) } diff --git a/internal/apps/edge/runner/wsreconnect.go b/internal/apps/edge/runner/wsreconnect.go index 1dfef0a3..4db235d0 100644 --- a/internal/apps/edge/runner/wsreconnect.go +++ b/internal/apps/edge/runner/wsreconnect.go @@ -63,8 +63,10 @@ func RunWSReconnectLoop(ctx context.Context, cfg WSReconnectConfig, // SleepContext pauses execution for the given duration or until the context is canceled. func SleepContext(ctx context.Context, d time.Duration) { + t := time.NewTimer(d) + defer t.Stop() select { case <-ctx.Done(): - case <-time.After(d): + case <-t.C: } } diff --git a/internal/apps/flared/frpc/manager.go b/internal/apps/flared/frpc/manager.go index 00ecefb6..23d5f3bd 100644 --- a/internal/apps/flared/frpc/manager.go +++ b/internal/apps/flared/frpc/manager.go @@ -5,6 +5,7 @@ import ( "bytes" "context" "encoding/json" + "errors" "fmt" "log/slog" "net" @@ -13,6 +14,7 @@ import ( "path/filepath" "strings" "sync" + "syscall" "time" "github.com/Rain-kl/Wavelet/internal/apps/flared/config" @@ -225,15 +227,18 @@ func (m *Manager) restartProcess(ctx context.Context, relayID string, configPath backoff = 1 * time.Second } + t := time.NewTimer(backoff) select { case <-procCtx.Done(): + t.Stop() return - case <-time.After(backoff): + case <-t.C: backoff *= 2 if backoff > maxBackoff { backoff = maxBackoff } } + t.Stop() } }() } @@ -350,10 +355,13 @@ func ensureNoOrphanProcess(pidPath string) { } process, err := os.FindProcess(pid) if err == nil && process != nil { - slog.Warn("attempting to kill potentially orphan process", "pid", pid, "pid_path", pidPath) - _ = process.Kill() - // Wait a little bit to ensure the OS has reclaimed ports - time.Sleep(orphanProcessKillDelay) + err = process.Signal(syscall.Signal(0)) + if err == nil || errors.Is(err, os.ErrPermission) { + slog.Warn("attempting to kill potentially orphan process", "pid", pid, "pid_path", pidPath) + _ = process.Kill() + // Wait a little bit to ensure the OS has reclaimed ports + time.Sleep(orphanProcessKillDelay) + } } _ = os.Remove(pidPath) } diff --git a/internal/apps/oauth/cache.go b/internal/apps/oauth/cache.go new file mode 100644 index 00000000..9b031766 --- /dev/null +++ b/internal/apps/oauth/cache.go @@ -0,0 +1,145 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package oauth + +import ( + "context" + "fmt" + "sync" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" +) + +type cacheEntry struct { + value any + expiredAt time.Time +} + +type memoryCache struct { + sync.RWMutex + items map[string]cacheEntry +} + +var localCache = &memoryCache{ + items: make(map[string]cacheEntry), +} + +func (c *memoryCache) Set(key string, val any, ttl time.Duration) { + c.Lock() + defer c.Unlock() + c.items[key] = cacheEntry{ + value: val, + expiredAt: time.Now().Add(ttl), + } +} + +func (c *memoryCache) Get(key string) (any, bool) { + c.RLock() + defer c.RUnlock() + item, ok := c.items[key] + if !ok { + return nil, false + } + if time.Now().After(item.expiredAt) { + return nil, false + } + return item.value, true +} + +func (c *memoryCache) Delete(key string) { + c.Lock() + defer c.Unlock() + delete(c.items, key) +} + +const ( + tokenCacheTTL = 5 * time.Minute + userCacheTTL = 5 * time.Minute +) + +func tokenCacheKey(tokenHash string) string { + return "oauth:token:" + tokenHash +} + +func userCacheKey(userID uint64) string { + return fmt.Sprintf("oauth:user:%d", userID) +} + +// GetCachedToken 获取缓存的 AccessToken +func GetCachedToken(ctx context.Context, tokenHash string) (*model.AccessToken, error) { + key := tokenCacheKey(tokenHash) + if val, ok := localCache.Get(key); ok { + if token, ok := val.(*model.AccessToken); ok { + return token, nil + } + } + + if db.Redis != nil { + var token model.AccessToken + if err := db.GetJSON(ctx, key, &token); err == nil { + // Write back to local cache + localCache.Set(key, &token, tokenCacheTTL) + return &token, nil + } + } + return nil, fmt.Errorf("cache miss") +} + +// SetCachedToken 设置 AccessToken 缓存 +func SetCachedToken(ctx context.Context, tokenHash string, token *model.AccessToken) { + key := tokenCacheKey(tokenHash) + localCache.Set(key, token, tokenCacheTTL) + if db.Redis != nil { + _ = db.SetJSON(ctx, key, token, tokenCacheTTL) + } +} + +// InvalidateCachedToken 吊销/删除 token 缓存 +func InvalidateCachedToken(ctx context.Context, tokenHash string) { + key := tokenCacheKey(tokenHash) + localCache.Delete(key) + if db.Redis != nil { + _ = db.Redis.Del(ctx, db.PrefixedKey(key)).Err() + } +} + +// GetCachedUser 获取缓存的 User +func GetCachedUser(ctx context.Context, userID uint64) (*model.User, error) { + key := userCacheKey(userID) + if val, ok := localCache.Get(key); ok { + if u, ok := val.(*model.User); ok { + return u, nil + } + } + + if db.Redis != nil { + var u model.User + if err := db.GetJSON(ctx, key, &u); err == nil { + // Write back to local cache + localCache.Set(key, &u, userCacheTTL) + return &u, nil + } + } + return nil, fmt.Errorf("cache miss") +} + +// SetCachedUser 设置 User 缓存 +func SetCachedUser(ctx context.Context, userID uint64, u *model.User) { + key := userCacheKey(userID) + localCache.Set(key, u, userCacheTTL) + if db.Redis != nil { + _ = db.SetJSON(ctx, key, u, userCacheTTL) + } +} + +// InvalidateCachedUser 吊销/失效 User 缓存 +func InvalidateCachedUser(ctx context.Context, userID uint64) { + key := userCacheKey(userID) + localCache.Delete(key) + if db.Redis != nil { + _ = db.Redis.Del(ctx, db.PrefixedKey(key)).Err() + } +} diff --git a/internal/apps/oauth/middlewares.go b/internal/apps/oauth/middlewares.go index 6ad871fa..6e341543 100644 --- a/internal/apps/oauth/middlewares.go +++ b/internal/apps/oauth/middlewares.go @@ -31,15 +31,26 @@ type loginRequiredAuditLog struct { func getUserByToken(ctx context.Context, tokenStr string) (*model.User, *model.AccessToken, error) { tokenHash := model.HashToken(tokenStr) - var tokenRecord model.AccessToken - if err := db.DB(ctx).Where("token_hash = ?", tokenHash).First(&tokenRecord).Error; err != nil { - return nil, nil, err + tokenRecord, err := GetCachedToken(ctx, tokenHash) + if err != nil { + var dbToken model.AccessToken + if err := db.DB(ctx).Where("token_hash = ?", tokenHash).First(&dbToken).Error; err != nil { + return nil, nil, err + } + tokenRecord = &dbToken + SetCachedToken(ctx, tokenHash, tokenRecord) } - var user model.User - if err := db.DB(ctx).Where("id = ? AND is_active = ?", tokenRecord.UserID, true).First(&user).Error; err != nil { - return nil, nil, err + + user, err := GetCachedUser(ctx, tokenRecord.UserID) + if err != nil || !user.IsActive { + var dbUser model.User + if err := db.DB(ctx).Where("id = ? AND is_active = ?", tokenRecord.UserID, true).First(&dbUser).Error; err != nil { + return nil, nil, err + } + user = &dbUser + SetCachedUser(ctx, tokenRecord.UserID, user) } - return &user, &tokenRecord, nil + return user, tokenRecord, nil } // GetUserFromRequest 校验 Access Token 或 Session 并返回用户对象,如果未登录或用户失效则返回 error @@ -74,11 +85,16 @@ func GetUserFromRequest(c *gin.Context) (*model.User, error) { return nil, errors.New("unauthorized") } - var user model.User - // load user from db to make sure is active - tx := db.DB(ctx).Where("id = ? AND is_active = ?", userID, true).First(&user) - if tx.Error != nil { - return nil, tx.Error + user, err := GetCachedUser(ctx, userID) + if err != nil || !user.IsActive { + var dbUser model.User + // load user from db to make sure is active + tx := db.DB(ctx).Where("id = ? AND is_active = ?", userID, true).First(&dbUser) + if tx.Error != nil { + return nil, tx.Error + } + user = &dbUser + SetCachedUser(ctx, userID, user) } // 密码哈希校验:当用户存在本地密码时,要求 Session 中的密码哈希必须与当前数据库中一致 @@ -99,7 +115,7 @@ func GetUserFromRequest(c *gin.Context) (*model.User, error) { return nil, errors.New("system user is not allowed to login") } - return &user, nil + return user, nil } // LoginRequired 返回登录鉴权中间件,校验 Access Token 或 Session diff --git a/internal/apps/oauth/routers.go b/internal/apps/oauth/routers.go index f94b97a0..8d2e13f6 100644 --- a/internal/apps/oauth/routers.go +++ b/internal/apps/oauth/routers.go @@ -95,6 +95,13 @@ func Logout(c *gin.Context) { username := session.Get(UserNameKey) if userID != nil { logger.InfoF(c.Request.Context(), "[LoginAudit] user logged out: %v, ID: %v, IP: %s", username, userID, c.ClientIP()) + if id, ok := userID.(uint64); ok { + InvalidateCachedUser(c.Request.Context(), id) + } else if idFloat, ok := userID.(float64); ok { + InvalidateCachedUser(c.Request.Context(), uint64(idFloat)) + } else if idInt, ok := userID.(int); ok && idInt >= 0 { + InvalidateCachedUser(c.Request.Context(), uint64(idInt)) + } } session.Options(GetSessionOptions(-1)) session.Clear() diff --git a/internal/apps/openflare/chwriter/dedup.go b/internal/apps/openflare/chwriter/dedup.go index d9b8127c..3ffcecd2 100644 --- a/internal/apps/openflare/chwriter/dedup.go +++ b/internal/apps/openflare/chwriter/dedup.go @@ -11,12 +11,16 @@ import ( const dedupTTL = 2 * time.Minute type dedupSet struct { - mu sync.Mutex - keys map[string]time.Time + mu sync.Mutex + keys map[string]time.Time + lastCleanup time.Time } func newDedupSet() *dedupSet { - return &dedupSet{keys: make(map[string]time.Time)} + return &dedupSet{ + keys: make(map[string]time.Time), + lastCleanup: time.Now(), + } } // markIfNew records key when it has not been seen within dedupTTL. @@ -29,11 +33,16 @@ func (s *dedupSet) markIfNew(key string) bool { s.mu.Lock() defer s.mu.Unlock() - for existing, expiresAt := range s.keys { - if now.After(expiresAt) { - delete(s.keys, existing) + // Periodically clean up all expired keys (e.g., every 30 seconds) + if now.Sub(s.lastCleanup) >= 30*time.Second { + for existing, expiresAt := range s.keys { + if now.After(expiresAt) { + delete(s.keys, existing) + } } + s.lastCleanup = now } + if expiresAt, exists := s.keys[key]; exists && now.Before(expiresAt) { return false } diff --git a/internal/apps/relay/frps/manager.go b/internal/apps/relay/frps/manager.go index 40118bfb..47b77f5b 100644 --- a/internal/apps/relay/frps/manager.go +++ b/internal/apps/relay/frps/manager.go @@ -6,6 +6,7 @@ package frps import ( "bytes" "context" + "errors" "fmt" "log/slog" "os" @@ -13,6 +14,7 @@ import ( "path/filepath" "strings" "sync" + "syscall" "time" service "github.com/Rain-kl/Wavelet/pkg/protocol" @@ -322,10 +324,13 @@ func ensureNoOrphanProcess(pidPath string) { } process, err := os.FindProcess(pid) if err == nil && process != nil { - slog.Warn("attempting to kill potentially orphan process", "pid", pid, "pid_path", pidPath) - _ = process.Kill() - // Wait a little bit to ensure the OS has reclaimed ports - time.Sleep(frpsOrphanProcessCleanupDelay) + err = process.Signal(syscall.Signal(0)) + if err == nil || errors.Is(err, os.ErrPermission) { + slog.Warn("attempting to kill potentially orphan process", "pid", pid, "pid_path", pidPath) + _ = process.Kill() + // Wait a little bit to ensure the OS has reclaimed ports + time.Sleep(frpsOrphanProcessCleanupDelay) + } } _ = os.Remove(pidPath) } diff --git a/internal/apps/upload/handler/routers.go b/internal/apps/upload/handler/routers.go index aaa0b607..d85998f8 100644 --- a/internal/apps/upload/handler/routers.go +++ b/internal/apps/upload/handler/routers.go @@ -7,6 +7,7 @@ package handler import ( "archive/zip" + "bufio" "bytes" "crypto/sha256" "encoding/hex" @@ -247,8 +248,12 @@ func BatchDownloadFiles(c *gin.Context) { c.Header("Content-Type", "application/zip") c.Header("Content-Disposition", "attachment; filename=\"batch_download.zip\"") - zipWriter := zip.NewWriter(c.Writer) - defer func() { _ = zipWriter.Close() }() + bufferedWriter := bufio.NewWriter(c.Writer) + zipWriter := zip.NewWriter(bufferedWriter) + defer func() { + _ = zipWriter.Close() + _ = bufferedWriter.Flush() + }() usedNames := make(map[string]int) diff --git a/internal/apps/user/access_tokens.go b/internal/apps/user/access_tokens.go index 235aabae..647cfbca 100644 --- a/internal/apps/user/access_tokens.go +++ b/internal/apps/user/access_tokens.go @@ -11,7 +11,6 @@ import ( "strings" "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" "github.com/gin-gonic/gin" @@ -43,8 +42,8 @@ func ListAccessTokens(c *gin.Context) { currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey) ctx := c.Request.Context() - var tokens []model.AccessToken - if err := db.DB(ctx).Where("user_id = ?", currUser.ID).Order("created_at desc").Find(&tokens).Error; err != nil { + tokens, err := listAccessTokensLogic(ctx, currUser.ID) + if err != nil { response.AbortBadRequest(c, err.Error()) return } @@ -91,8 +90,8 @@ func CreateAccessToken(c *gin.Context) { maxLimit = val } - var count int64 - if err := db.DB(ctx).Model(&model.AccessToken{}).Where("user_id = ?", currUser.ID).Count(&count).Error; err != nil { + count, err := countAccessTokensLogic(ctx, currUser.ID) + if err != nil { response.AbortBadRequest(c, err.Error()) return } @@ -120,7 +119,7 @@ func CreateAccessToken(c *gin.Context) { IsAdmin: req.IsAdmin, } - if err := db.DB(ctx).Create(&tokenRecord).Error; err != nil { + if err := createAccessTokenLogic(ctx, &tokenRecord); err != nil { response.AbortBadRequest(c, err.Error()) return } @@ -152,14 +151,8 @@ func DeleteAccessToken(c *gin.Context) { return } - tx := db.DB(ctx).Where("id = ? AND user_id = ?", id, currUser.ID).Delete(&model.AccessToken{}) - if tx.Error != nil { - response.AbortBadRequest(c, tx.Error.Error()) - return - } - - if tx.RowsAffected == 0 { - response.AbortBadRequest(c, errTokenNotFoundOrForbidden) + if err := deleteAccessTokenLogic(ctx, id, currUser.ID); err != nil { + response.AbortBadRequest(c, err.Error()) return } @@ -187,32 +180,14 @@ func RotateAccessToken(c *gin.Context) { return } - var tokenRecord model.AccessToken - if err := db.DB(ctx).Where("id = ? AND user_id = ?", id, currUser.ID).First(&tokenRecord).Error; err != nil { - response.AbortBadRequest(c, errTokenNotFoundOrForbidden) - return - } - - // 生成新的 Token - newTokenStr, err := model.GenerateTokenString() + newTokenStr, tokenRecord, err := rotateAccessTokenLogic(ctx, id, currUser.ID) if err != nil { - response.AbortBadRequest(c, errGenerateTokenFailed) - return - } - - newTokenHash := model.HashToken(newTokenStr) - newMaskedToken := model.MaskTokenString(newTokenStr) - - tokenRecord.TokenHash = newTokenHash - tokenRecord.MaskedToken = newMaskedToken - - if err := db.DB(ctx).Save(&tokenRecord).Error; err != nil { response.AbortBadRequest(c, err.Error()) return } c.JSON(http.StatusOK, response.OK(tokenResponse{ Token: newTokenStr, - Record: tokenRecord, + Record: *tokenRecord, })) } diff --git a/internal/apps/user/logics.go b/internal/apps/user/logics.go index 1fc9db2f..4d6f0ce7 100644 --- a/internal/apps/user/logics.go +++ b/internal/apps/user/logics.go @@ -12,6 +12,7 @@ import ( "math/big" "strings" + "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" @@ -285,3 +286,125 @@ func updateUserProfile(ctx context.Context, userID uint64, input updateProfileIn } return &dbUser, nil } + +func getUserByUsernameOrEmail(ctx context.Context, input string) (*model.User, error) { + var user model.User + if err := db.DB(ctx).Where("username = ? OR email = ?", input, input).First(&user).Error; err != nil { + return nil, err + } + return &user, nil +} + +func updateLastLogin(ctx context.Context, user *model.User) error { + return db.DB(ctx).Model(user).Update("last_login_at", user.LastLoginAt).Error +} + +func registerUserLogic(ctx context.Context, u *model.User) error { + if err := u.RegisterUser(ctx, db.DB(ctx)); err != nil { + if strings.Contains(err.Error(), "duplicate key") || strings.Contains(err.Error(), "UNIQUE") { + return errors.New("用户名或邮箱已被占用") + } + return errors.New("注册失败,请稍后再试") + } + return nil +} + +func changePasswordLogic(ctx context.Context, userID uint64, oldPass, newPass string) error { + var dbUser model.User + if err := db.DB(ctx).Where("id = ?", userID).First(&dbUser).Error; err != nil { + return errors.New(errUserNotFound) + } + + if !dbUser.CheckPassword(oldPass) { + return errors.New(errOldPasswordIncorrect) + } + + if err := dbUser.SetEncryptedPassword(newPass); err != nil { + return errors.New(errPasswordEncryptFailed) + } + + if err := db.DB(ctx).Model(&dbUser).Update("password", dbUser.Password).Error; err != nil { + return errors.New("更新密码失败,请稍后再试") + } + + // 吊销该用户所有的 Access Token + var tokens []model.AccessToken + if err := db.DB(ctx).Where("user_id = ?", dbUser.ID).Find(&tokens).Error; err == nil { + for _, token := range tokens { + oauth.InvalidateCachedToken(ctx, token.TokenHash) + } + } + if err := db.DB(ctx).Where("user_id = ?", dbUser.ID).Delete(&model.AccessToken{}).Error; err != nil { + return errors.New("吊销 Access Token 失败,请稍后再试") + } + + oauth.InvalidateCachedUser(ctx, dbUser.ID) + return nil +} + +func listAccessTokensLogic(ctx context.Context, userID uint64) ([]model.AccessToken, error) { + var tokens []model.AccessToken + if err := db.DB(ctx).Where("user_id = ?", userID).Order("created_at desc").Find(&tokens).Error; err != nil { + return nil, errors.New("获取令牌列表失败,请稍后再试") + } + return tokens, nil +} + +func countAccessTokensLogic(ctx context.Context, userID uint64) (int64, error) { + var count int64 + if err := db.DB(ctx).Model(&model.AccessToken{}).Where("user_id = ?", userID).Count(&count).Error; err != nil { + return 0, errors.New("查询令牌数量失败,请稍后再试") + } + return count, nil +} + +func createAccessTokenLogic(ctx context.Context, record *model.AccessToken) error { + if err := db.DB(ctx).Create(record).Error; err != nil { + return errors.New("创建令牌失败,请稍后再试") + } + return nil +} + +func deleteAccessTokenLogic(ctx context.Context, id, userID uint64) error { + var tokenRecord model.AccessToken + if err := db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).First(&tokenRecord).Error; err != nil { + return errors.New(errTokenNotFoundOrForbidden) + } + oauth.InvalidateCachedToken(ctx, tokenRecord.TokenHash) + + tx := db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&model.AccessToken{}) + if tx.Error != nil { + return errors.New("删除令牌失败,请稍后再试") + } + if tx.RowsAffected == 0 { + return errors.New(errTokenNotFoundOrForbidden) + } + return nil +} + +func rotateAccessTokenLogic(ctx context.Context, id, userID uint64) (string, *model.AccessToken, error) { + var tokenRecord model.AccessToken + if err := db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).First(&tokenRecord).Error; err != nil { + return "", nil, errors.New(errTokenNotFoundOrForbidden) + } + + oauth.InvalidateCachedToken(ctx, tokenRecord.TokenHash) + + newTokenStr, err := model.GenerateTokenString() + if err != nil { + return "", nil, errors.New(errGenerateTokenFailed) + } + + newTokenHash := model.HashToken(newTokenStr) + newMaskedToken := model.MaskTokenString(newTokenStr) + + tokenRecord.TokenHash = newTokenHash + tokenRecord.MaskedToken = newMaskedToken + + if err := db.DB(ctx).Save(&tokenRecord).Error; err != nil { + return "", nil, errors.New("轮换令牌失败,请稍后再试") + } + + return newTokenStr, &tokenRecord, nil +} + diff --git a/internal/apps/user/routers.go b/internal/apps/user/routers.go index b404281a..b3903e79 100644 --- a/internal/apps/user/routers.go +++ b/internal/apps/user/routers.go @@ -13,7 +13,6 @@ import ( "github.com/Rain-kl/Wavelet/internal/common" "github.com/Rain-kl/Wavelet/internal/common/response" "github.com/Rain-kl/Wavelet/internal/config" - "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/db/idgen" "github.com/Rain-kl/Wavelet/internal/listener" "github.com/Rain-kl/Wavelet/internal/model" @@ -117,8 +116,8 @@ func Login(c *gin.Context) { return } - var user model.User - if err := db.DB(ctx).Where("username = ? OR email = ?", req.Username, req.Username).First(&user).Error; err != nil { + user, err := getUserByUsernameOrEmail(ctx, req.Username) + if err != nil { logger.WarnF(ctx, "[LoginAudit] failed login attempt (username not found) for input: %s, IP: %s", req.Username, c.ClientIP()) response.AbortBadRequest(c, errUsernameOrPasswordWrong) return @@ -139,7 +138,7 @@ func Login(c *gin.Context) { } if isEmailLoginVerificationEnabled(ctx) { - result, err := processLoginEmailVerification(ctx, req.Code, &user) + result, err := processLoginEmailVerification(ctx, req.Code, user) if err != nil { response.AbortBadRequest(c, err.Error()) return @@ -160,20 +159,20 @@ func Login(c *gin.Context) { } user.LastLoginAt = time.Now() - if err := db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error; err != nil { - response.AbortBadRequest(c, err.Error()) + if err := updateLastLogin(ctx, user); err != nil { + response.AbortBadRequest(c, "更新登录时间失败,请稍后再试") return } - if err := setLoginSession(ctx, c, &user); err != nil { + if err := setLoginSession(ctx, c, user); err != nil { response.AbortBadRequest(c, errSaveSessionFailed) return } logger.InfoF(ctx, "[LoginAudit] successful login for user: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP()) - listener.EmitAdminLoggedIn(ctx, &user, c.ClientIP()) + listener.EmitAdminLoggedIn(ctx, user, c.ClientIP()) - c.JSON(http.StatusOK, response.OK(oauth.BuildBasicUserInfo(&user, needChangePassword))) + c.JSON(http.StatusOK, response.OK(oauth.BuildBasicUserInfo(user, needChangePassword))) } // Register 用户注册 @@ -247,7 +246,7 @@ func Register(c *gin.Context) { return } - if err := user.RegisterUser(ctx, db.DB(ctx)); err != nil { + if err := registerUserLogic(ctx, &user); err != nil { response.AbortBadRequest(c, err.Error()) return } @@ -275,6 +274,13 @@ func Logout(c *gin.Context) { username := session.Get(oauth.UserNameKey) if userID != nil { logger.InfoF(c.Request.Context(), "[LoginAudit] user logged out: %v, ID: %v, IP: %s", username, userID, c.ClientIP()) + if id, ok := userID.(uint64); ok { + oauth.InvalidateCachedUser(c.Request.Context(), id) + } else if idFloat, ok := userID.(float64); ok { + oauth.InvalidateCachedUser(c.Request.Context(), uint64(idFloat)) + } else if idInt, ok := userID.(int); ok && idInt >= 0 { + oauth.InvalidateCachedUser(c.Request.Context(), uint64(idInt)) + } } session.Options(oauth.GetSessionOptions(-1)) session.Clear() @@ -327,35 +333,11 @@ func ChangePassword(c *gin.Context) { } ctx := c.Request.Context() - var dbUser model.User - if err := db.DB(ctx).Where("id = ?", userObj.ID).First(&dbUser).Error; err != nil { - response.AbortBadRequest(c, errUserNotFound) - return - } - - // 校验旧密码 - if !dbUser.CheckPassword(req.OldPassword) { - response.AbortBadRequest(c, errOldPasswordIncorrect) - return - } - - // 加密并更新为新密码 - if err := dbUser.SetEncryptedPassword(req.NewPassword); err != nil { - response.AbortBadRequest(c, errPasswordEncryptFailed) - return - } - - if err := db.DB(ctx).Model(&dbUser).Update("password", dbUser.Password).Error; err != nil { + if err := changePasswordLogic(ctx, userObj.ID, req.OldPassword, req.NewPassword); err != nil { response.AbortBadRequest(c, err.Error()) return } - // 吊销该用户所有的 Access Token - if err := db.DB(ctx).Where("user_id = ?", dbUser.ID).Delete(&model.AccessToken{}).Error; err != nil { - response.AbortBadRequest(c, "吊销 Access Token 失败: "+err.Error()) - return - } - // 销毁当前活跃会话以强制重新登录 session := sessions.Default(c) session.Clear() @@ -431,6 +413,7 @@ func UpdateProfile(c *gin.Context) { response.AbortBadRequest(c, err.Error()) return } + oauth.InvalidateCachedUser(ctx, userObj.ID) session := sessions.Default(c) needChange := session.Get("need_change_password") == true diff --git a/internal/repository/system_config_cache.go b/internal/repository/system_config_cache.go index b35a0e0d..f149f17b 100644 --- a/internal/repository/system_config_cache.go +++ b/internal/repository/system_config_cache.go @@ -30,8 +30,10 @@ type systemConfigInvalidationMessage struct { } var ( - systemConfigRAMCache = ram.MustNew[string, model.SystemConfig](ram.Options{MaximumSize: systemConfigRAMMaximumSize}) - systemConfigListenerOnce sync.Once + systemConfigRAMCache = ram.MustNew[string, model.SystemConfig](ram.Options{MaximumSize: systemConfigRAMMaximumSize}) + systemConfigListenerOnce sync.Once + systemConfigListenerCtx context.Context + systemConfigListenerCancel context.CancelFunc ) func ensureSystemConfigCacheListener() { @@ -43,12 +45,19 @@ func startSystemConfigCacheInvalidationListener() { return } + systemConfigListenerCtx, systemConfigListenerCancel = context.WithCancel(context.Background()) + go func() { - pubsub := db.Redis.Subscribe(context.Background(), SystemConfigInvalidationChannel) + pubsub := db.Redis.Subscribe(systemConfigListenerCtx, SystemConfigInvalidationChannel) defer func() { _ = pubsub.Close() }() + go func() { + <-systemConfigListenerCtx.Done() + _ = pubsub.Close() + }() + for msg := range pubsub.Channel() { var payload systemConfigInvalidationMessage if err := json.Unmarshal([]byte(msg.Payload), &payload); err != nil { @@ -64,6 +73,15 @@ func startSystemConfigCacheInvalidationListener() { }() } +// StopSystemConfigCacheListener stops the Redis Pub/Sub subscription listener and resets the sync.Once guard. +func StopSystemConfigCacheListener() { + if systemConfigListenerCancel != nil { + systemConfigListenerCancel() + systemConfigListenerCancel = nil + } + systemConfigListenerOnce = sync.Once{} +} + func cloneSystemConfig(sc model.SystemConfig) model.SystemConfig { return sc } diff --git a/pkg/wsclient/client.go b/pkg/wsclient/client.go index 2a567ed4..467b539e 100644 --- a/pkg/wsclient/client.go +++ b/pkg/wsclient/client.go @@ -10,6 +10,7 @@ import ( "net/http" "net/url" "strings" + "sync" "time" "golang.org/x/net/websocket" @@ -52,6 +53,7 @@ type Connection struct { Conn *websocket.Conn URL string ReadTimeout time.Duration + writeMu sync.Mutex } // New creates a new WebSocket client with the given configuration. @@ -154,6 +156,9 @@ func (conn *Connection) SendMessage(msgType string, payload any) error { Payload: payload, } + conn.writeMu.Lock() + defer conn.writeMu.Unlock() + _ = conn.Conn.SetWriteDeadline(time.Now().Add(writeDeadlineSecs * time.Second)) return websocket.JSON.Send(conn.Conn, message) }