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.
This commit is contained in:
ryan
2026-06-20 09:09:02 +08:00
parent 5bed2bae9f
commit b3a55d4ab5
25 changed files with 668 additions and 364 deletions
-40
View File
@@ -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 脚本。
@@ -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!"
-3
View File
@@ -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)
+10
View File
@@ -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 }` 信封规范。
### 移除
+55 -31
View File
@@ -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": {
+55 -31
View File
@@ -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": {
+32 -22
View File
@@ -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
+30 -2
View File
@@ -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) {
+5 -2
View File
@@ -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
}
@@ -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.
+37 -2
View File
@@ -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
+16 -20
View File
@@ -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))
}
+3 -1
View File
@@ -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:
}
}
+13 -5
View File
@@ -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)
}
+145
View File
@@ -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()
}
}
+29 -13
View File
@@ -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
+7
View File
@@ -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()
+15 -6
View File
@@ -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
}
+9 -4
View File
@@ -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)
}
+7 -2
View File
@@ -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)
+9 -34
View File
@@ -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,
}))
}
+123
View File
@@ -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
}
+18 -35
View File
@@ -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
+21 -3
View File
@@ -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
}
+5
View File
@@ -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)
}