diff --git a/docker-compose.yml b/docker-compose.yml index 2480d030..65530fb0 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -1,10 +1,10 @@ services: wavelet: - build: - context: . - dockerfile: docker/Dockerfile - container_name: wavelet-app +# build: +# context: . +# dockerfile: docker/Dockerfile + image: ghcr.io/rain-kl/wavelet:v0.1.0 restart: unless-stopped env_file: .env environment: @@ -22,7 +22,6 @@ services: postgres: image: postgres:17-alpine - container_name: wavelet-postgres restart: unless-stopped environment: POSTGRES_DB: ${POSTGRES_DB:-wavelet} @@ -42,7 +41,6 @@ services: redis: image: redis:7-alpine - container_name: wavelet-redis restart: unless-stopped command: ["redis-server", "--appendonly", "yes"] ports: @@ -55,27 +53,26 @@ services: timeout: 5s retries: 5 start_period: 5s - - clickhouse: - image: clickhouse/clickhouse-server:25.3-alpine - container_name: wavelet-clickhouse - restart: unless-stopped - profiles: - - clickhouse - environment: - CLICKHOUSE_DB: ${CLICKHOUSE_DB:-wavelet} - CLICKHOUSE_USER: ${CLICKHOUSE_USER:-default} - CLICKHOUSE_PASSWORD: ${CLICKHOUSE_PASSWORD:-123456} - CLICKHOUSE_DEFAULT_ACCESS_MANAGEMENT: 1 - TZ: ${TZ:-Asia/Shanghai} - ports: - - "${CLICKHOUSE_HTTP_PORT:-8123}:8123" - - "${CLICKHOUSE_NATIVE_PORT:-9000}:9000" - volumes: - - ./data/clickhouse_data:/var/lib/clickhouse - healthcheck: - test: ["CMD", "clickhouse-client", "--query", "SELECT 1"] - interval: 10s - timeout: 5s - retries: 5 - start_period: 15s +# +# clickhouse: +# image: clickhouse/clickhouse-server:25.3-alpine +# restart: unless-stopped +# profiles: +# - clickhouse +# environment: +# CLICKHOUSE_DB: ${CLICKHOUSE_DB:-wavelet} +# CLICKHOUSE_USER: ${CLICKHOUSE_USER:-default} +# CLICKHOUSE_PASSWORD: ${CLICKHOUSE_PASSWORD:-123456} +# CLICKHOUSE_DEFAULT_ACCESS_MANAGEMENT: 1 +# TZ: ${TZ:-Asia/Shanghai} +# ports: +# - "${CLICKHOUSE_HTTP_PORT:-8123}:8123" +# - "${CLICKHOUSE_NATIVE_PORT:-9000}:9000" +# volumes: +# - ./data/clickhouse_data:/var/lib/clickhouse +# healthcheck: +# test: ["CMD", "clickhouse-client", "--query", "SELECT 1"] +# interval: 10s +# timeout: 5s +# retries: 5 +# start_period: 15s diff --git a/internal/apps/oauth/sources.go b/internal/apps/oauth/sources.go index b76062b3..7a645907 100644 --- a/internal/apps/oauth/sources.go +++ b/internal/apps/oauth/sources.go @@ -49,13 +49,15 @@ type AuthSourceView struct { ClientSecretConfigured bool `json:"client_secret_configured"` } -// AuthorizeResponse 授权 URL 响应 +// OAuthAuthorizeResponse 授权 URL 响应 +// //nolint:revive // OAuth 前缀保持包内语义清晰 type OAuthAuthorizeResponse struct { AuthorizeURL string `json:"authorize_url"` } -// CallbackResult 回调处理结果 +// OAuthCallbackResult 回调处理结果 +// //nolint:revive // OAuth 前缀保持包内语义清晰 type OAuthCallbackResult struct { Status string `json:"status"` diff --git a/internal/apps/risk_control/logics.go b/internal/apps/risk_control/logics.go new file mode 100644 index 00000000..56bf445f --- /dev/null +++ b/internal/apps/risk_control/logics.go @@ -0,0 +1,130 @@ +/* +Copyright 2026 Arctel.net + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package risk_control + +import ( + "context" + "time" + + "github.com/Rain-kl/Wavelet/internal/config" + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/logger" +) + +var logChan chan *UserAccessLog + +const ( + defaultQueueSize = 10000 + maxBatchSize = 1000 + flushInterval = 1 * time.Second +) + +// InitLogWriter 初始化日志写入通道和后台写入协程 +func InitLogWriter() { + if !config.Config.ClickHouse.Enabled { + return + } + + logChan = make(chan *UserAccessLog, defaultQueueSize) + go startBatchWorker() +} + +// IsBufferFull 检查当前本地缓冲队列是否已满 +// 如果没有启用 ClickHouse,默认返回 false,不触发限流 +func IsBufferFull() bool { + if !config.Config.ClickHouse.Enabled || logChan == nil { + return false + } + return len(logChan) >= cap(logChan) +} + +// QueueAccessLog 异步非阻塞地将日志推入缓冲队列 +func QueueAccessLog(logItem *UserAccessLog) { + if !config.Config.ClickHouse.Enabled || logChan == nil { + return + } + + select { + case logChan <- logItem: + default: + logger.WarnF(context.Background(), "[RiskControl] Log queue full, dropping log item for path: %s", logItem.Path) + } +} + +func startBatchWorker() { + ticker := time.NewTicker(flushInterval) + defer ticker.Stop() + + var batch []*UserAccessLog + + flush := func() { + if len(batch) == 0 { + return + } + if db.ChConn == nil { + batch = nil + return + } + + ctx := context.Background() + b, err := db.ChConn.PrepareBatch(ctx, "INSERT INTO user_access_logs (id, user_id, path, method, ip, user_agent, headers, status, latency, created_at)") + if err != nil { + logger.ErrorF(ctx, "[RiskControl] Prepare ClickHouse batch failed: %v", err) + batch = nil + return + } + + for _, item := range batch { + err = b.Append( + item.ID, + item.UserID, + item.Path, + item.Method, + item.IP, + item.UserAgent, + item.Headers, + item.Status, + item.Latency, + item.CreatedAt, + ) + if err != nil { + logger.ErrorF(ctx, "[RiskControl] Append item to ClickHouse batch failed: %v", err) + } + } + + if err := b.Send(); err != nil { + logger.ErrorF(ctx, "[RiskControl] Send ClickHouse batch failed: %v", err) + } + batch = nil + } + + for { + select { + case item, ok := <-logChan: + if !ok { + flush() + return + } + batch = append(batch, item) + if len(batch) >= maxBatchSize { + flush() + } + case <-ticker.C: + flush() + } + } +} diff --git a/internal/apps/risk_control/model.go b/internal/apps/risk_control/model.go index 13a0701a..5ca2ff19 100644 --- a/internal/apps/risk_control/model.go +++ b/internal/apps/risk_control/model.go @@ -17,12 +17,7 @@ limitations under the License. package risk_control import ( - "context" "time" - - "github.com/Rain-kl/Wavelet/internal/config" - "github.com/Rain-kl/Wavelet/internal/db" - "github.com/Rain-kl/Wavelet/internal/logger" ) // UserAccessLog 用户访问记录 @@ -38,111 +33,3 @@ type UserAccessLog struct { Latency int64 `json:"latency"` // 耗时毫秒 CreatedAt time.Time `json:"created_at"` } - -var ( - logChan chan *UserAccessLog -) - -const ( - defaultQueueSize = 10000 - maxBatchSize = 1000 - flushInterval = 1 * time.Second -) - -// InitLogWriter 初始化日志写入通道和后台写入协程 -func InitLogWriter() { - if !config.Config.ClickHouse.Enabled { - return - } - - logChan = make(chan *UserAccessLog, defaultQueueSize) - go startBatchWorker() -} - -// IsBufferFull 检查当前本地缓冲队列是否已满 -// 如果没有启用 ClickHouse,默认返回 false,不触发限流 -func IsBufferFull() bool { - if !config.Config.ClickHouse.Enabled || logChan == nil { - return false - } - return len(logChan) >= cap(logChan) -} - -// QueueAccessLog 异步非阻塞地将日志推入缓冲队列 -func QueueAccessLog(logItem *UserAccessLog) { - if !config.Config.ClickHouse.Enabled || logChan == nil { - return - } - - select { - case logChan <- logItem: - default: - // 如果在极端并发下仍然写满了,这里做非阻塞丢弃,防止卡死 - logger.WarnF(context.Background(), "[RiskControl] Log queue full, dropping log item for path: %s", logItem.Path) - } -} - -// startBatchWorker 后台批量写入 ClickHouse 的工作协程 -func startBatchWorker() { - ticker := time.NewTicker(flushInterval) - defer ticker.Stop() - - var batch []*UserAccessLog - - flush := func() { - if len(batch) == 0 { - return - } - if db.ChConn == nil { - batch = nil - return - } - - ctx := context.Background() - b, err := db.ChConn.PrepareBatch(ctx, "INSERT INTO user_access_logs (id, user_id, path, method, ip, user_agent, headers, status, latency, created_at)") - if err != nil { - logger.ErrorF(ctx, "[RiskControl] Prepare ClickHouse batch failed: %v", err) - batch = nil - return - } - - for _, item := range batch { - err = b.Append( - item.ID, - item.UserID, - item.Path, - item.Method, - item.IP, - item.UserAgent, - item.Headers, - item.Status, - item.Latency, - item.CreatedAt, - ) - if err != nil { - logger.ErrorF(ctx, "[RiskControl] Append item to ClickHouse batch failed: %v", err) - } - } - - if err := b.Send(); err != nil { - logger.ErrorF(ctx, "[RiskControl] Send ClickHouse batch failed: %v", err) - } - batch = nil - } - - for { - select { - case item, ok := <-logChan: - if !ok { - flush() - return - } - batch = append(batch, item) - if len(batch) >= maxBatchSize { - flush() - } - case <-ticker.C: - flush() - } - } -} diff --git a/internal/apps/upload/constants.go b/internal/apps/upload/constants.go new file mode 100644 index 00000000..370aeae0 --- /dev/null +++ b/internal/apps/upload/constants.go @@ -0,0 +1,24 @@ +/* +Copyright 2026 Arctel.net + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package upload + +const ( + maxUploadSize = 32 * 1024 * 1024 // 32MB + detectContentBytes = 512 // http.DetectContentType 需要的最小字节数 + uploadDirPerm = 0755 // 上传目录权限 + uploadFilePerm = 0644 // 上传文件权限 +) diff --git a/internal/apps/upload/errs.go b/internal/apps/upload/errs.go index 4a963fb3..1863bd27 100644 --- a/internal/apps/upload/errs.go +++ b/internal/apps/upload/errs.go @@ -18,37 +18,32 @@ limitations under the License. // Package upload 提供文件上传与下载功能 package upload -// 上传模块错误消息常量 +// 文件管理常量 const ( - ErrNoFileSelected = "请选择要上传的文件" - ErrInvalidUploadType = "无效的上传类型" - ErrFileTooLarge = "图片大小不能超过 2MB" - ErrUnsupportedFormat = "只支持 JPG、PNG、WEBP 格式的图片" - ErrInvalidImage = "无效的图片文件" - ErrUploadExtensionsNotConfigured = "上传扩展名未配置" - ErrProcessFileFailed = "处理文件失败" - ErrSaveFileFailed = "保存文件失败" - ErrOpenFileFailed = "打开文件失败" - ErrInvalidFilePath = "非法文件路径" - ErrSaveUploadRecordFailed = "保存上传记录失败" - ErrQueryHistoryUploadFailed = "查询历史上传记录失败" - ErrGenericFileTooLarge = "文件大小不能超过 32MB" - ErrFileContentExtensionMismatch = "文件内容与扩展名不匹配,可能包含安全风险" - ErrFileValidationFailed = "文件校验失败" - ErrInvalidMetadataJSON = "元数据 JSON 格式不合法" - ErrInvalidFileID = "无效的文件 ID" - ErrQueryUploadRecordFailed = "查询文件记录失败" - ErrInvalidBatchDownloadRequest = "参数绑定失败,请传入有效的文件 ID 数组" - ErrInvalidIDValueFormat = "无效的 ID 值: %s" - ErrRetrieveUploadRecordsFailed = "检索文件记录失败" - ErrNoValidFilesForArchive = "没有找到任何有效的文件记录进行打包" - ErrInvalidParams = "参数错误" - ErrQueryFileCountFailed = "查询文件数量失败" - ErrQueryFileListFailed = "查询文件列表失败" - ErrDeleteFileFailed = "删除文件失败" - ErrS3KeyRequired = "s3 key must not be empty" - ErrS3KeyTooLongFormat = "s3 key exceeds maximum length of %d" - ErrS3KeyStartsWithSlash = "s3 key must not start with /" - ErrS3KeyContainsNullBytes = "s3 key must not contain null bytes" - ErrQueryUnusedUploadsFailed = "查询未使用的上传文件失败: %w" + ErrNoFileSelected = "请选择要上传的文件" + ErrUnsupportedFormat = "只支持 JPG、PNG、WEBP 格式的图片" + ErrProcessFileFailed = "处理文件失败" + ErrSaveFileFailed = "保存文件失败" + ErrOpenFileFailed = "打开文件失败" + ErrSaveUploadRecordFailed = "保存上传记录失败" + ErrGenericFileTooLarge = "文件大小不能超过 32MB" + ErrFileContentExtensionMismatch = "文件内容与扩展名不匹配,可能包含安全风险" + ErrFileValidationFailed = "文件校验失败" + ErrInvalidMetadataJSON = "元数据 JSON 格式不合法" + ErrInvalidFileID = "无效的文件 ID" + ErrQueryUploadRecordFailed = "查询文件记录失败" + ErrInvalidBatchDownloadRequest = "参数绑定失败,请传入有效的文件 ID 数组" + ErrInvalidIDValueFormat = "无效的 ID 值: %s" + ErrRetrieveUploadRecordsFailed = "检索文件记录失败" + ErrNoValidFilesForArchive = "没有找到任何有效的文件记录进行打包" + ErrInvalidParams = "参数错误" + ErrQueryFileCountFailed = "查询文件数量失败" + ErrQueryFileListFailed = "查询文件列表失败" + ErrDeleteFileFailed = "删除文件失败" + ErrS3KeyRequired = "s3 key must not be empty" + ErrS3KeyTooLongFormat = "s3 key exceeds maximum length of %d" + ErrS3KeyStartsWithSlash = "s3 key must not start with /" + ErrS3KeyContainsNullBytes = "s3 key must not contain null bytes" + // + ErrQueryUnusedUploadsFailed = "查询未使用的上传文件失败: %w" ) diff --git a/internal/apps/upload/routers.go b/internal/apps/upload/routers.go index eec890c8..b8292525 100644 --- a/internal/apps/upload/routers.go +++ b/internal/apps/upload/routers.go @@ -48,13 +48,6 @@ import ( "gorm.io/gorm" ) -const ( - maxUploadSize = 32 * 1024 * 1024 // 32MB - detectContentBytes = 512 // http.DetectContentType 需要的最小字节数 - uploadDirPerm = 0755 // 上传目录权限 - uploadFilePerm = 0644 // 上传文件权限 -) - type batchDownloadRequest struct { IDs []string `json:"ids" binding:"required,min=1"` } diff --git a/internal/apps/user/logics.go b/internal/apps/user/logics.go new file mode 100644 index 00000000..d51cfbb1 --- /dev/null +++ b/internal/apps/user/logics.go @@ -0,0 +1,365 @@ +/* +Copyright 2026 Arctel.net + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package user + +import ( + "context" + "crypto/rand" + "encoding/json" + "errors" + "fmt" + "math/big" + "net/http" + "strings" + + "github.com/gin-contrib/sessions" + + "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/task" + "github.com/Rain-kl/Wavelet/internal/util" + "github.com/gin-gonic/gin" +) + +type sendEmailCodeRequest struct { + Email string `json:"email" binding:"required,email"` + Scene string `json:"scene" binding:"required"` +} + +func isEmailLoginVerificationEnabled() bool { + enabled, err := model.GetBoolByKey(context.Background(), model.ConfigKeyEmailLoginVerificationEnabled) + if err != nil { + return false + } + return enabled +} + +func isEmailRegisterVerificationEnabled() bool { + enabled, err := model.GetBoolByKey(context.Background(), model.ConfigKeyEmailRegisterVerificationEnabled) + if err != nil { + return false + } + return enabled +} + +func isSMTPConfigured(ctx context.Context) bool { + var sc model.SystemConfig + var host, port, username string + + if err := sc.GetByKey(ctx, model.ConfigKeySMTPHost); err == nil { + host = sc.Value + } + if err := sc.GetByKey(ctx, model.ConfigKeySMTPPort); err == nil { + port = sc.Value + } + if err := sc.GetByKey(ctx, model.ConfigKeySMTPUsername); err == nil { + username = sc.Value + } + + return host != "" && port != "" && username != "" +} + +func generateVerificationCode() (string, error) { + n, err := rand.Int(rand.Reader, big.NewInt(verificationCodeRange)) + if err != nil { + return "", err + } + return fmt.Sprintf("%06d", n.Int64()+verificationCodeOffset), nil +} + +func getEmailCodeKey(scene, email string) string { + return fmt.Sprintf("email_code:%s:%s", scene, email) +} + +func getEmailCooldownKey(scene, email string) string { + return fmt.Sprintf("email_code:cooldown:%s:%s", scene, email) +} + +func sendEmailVerificationCode(ctx context.Context, email, scene, templateName string) error { + if !isSMTPConfigured(ctx) { + return errors.New(errSMTPConfigIncomplete) + } + + code, err := generateVerificationCode() + if err != nil { + return errors.New(errGenerateEmailCodeFailed) + } + codeKey := getEmailCodeKey(scene, email) + cooldownKey := getEmailCooldownKey(scene, email) + + emailSubject, emailBody, err := model.RenderTemplate( + ctx, + templateName, + map[string]any{"Code": code}, + ) + if err != nil { + return fmt.Errorf(errRenderEmailTemplateFailed, err) + } + + if err := db.SetJSON(ctx, codeKey, code, emailCodeExpiry); err != nil { + return errors.New(errGenerateEmailCodeFailed) + } + _ = db.SetJSON(ctx, cooldownKey, "1", emailCodeCooldown) + + payload := SendEmailPayload{ + To: email, + Subject: emailSubject, + Body: emailBody, + } + payloadBytes, _ := json.Marshal(payload) + _, err = task.DispatchTask(ctx, task.TaskTypeSendEmail, payloadBytes, "system") + if err != nil { + return errors.New(errDispatchEmailTaskFailed) + } + return nil +} + +func verifyEmailCode(ctx context.Context, email, scene, code string) bool { + codeKey := getEmailCodeKey(scene, email) + var storedCode string + if err := db.GetJSON(ctx, codeKey, &storedCode); err != nil { + return false + } + if storedCode != code { + return false + } + _ = db.Redis.Del(ctx, db.PrefixedKey(codeKey)).Err() + return true +} + +func handleLoginEmailVerification(ctx context.Context, c *gin.Context, req *loginRequest, user *model.User) error { + if user.Email == "" { + c.JSON(http.StatusOK, util.Err(errLoginEmailMissing)) + return errors.New("handled") + } + + if req.Code == "" { + cooldownKey := getEmailCooldownKey("login", user.Email) + var temp string + err := db.GetJSON(ctx, cooldownKey, &temp) + if err != nil { + if err := sendEmailVerificationCode(ctx, user.Email, "login", "login_email"); err != nil { + c.JSON(http.StatusOK, util.Err(err.Error())) + return errors.New("handled") + } + } + + maskedEmail := util.MaskEmail(user.Email) + c.JSON(http.StatusOK, util.Err(errNeedEmailCodePrefix+maskedEmail)) + return errors.New("handled") + } + + if !verifyEmailCode(ctx, user.Email, "login", req.Code) { + c.JSON(http.StatusOK, util.Err(errEmailCodeInvalidOrExpired)) + return errors.New("handled") + } + return nil +} + +// SendEmailCode 发送邮箱验证码 +// @Summary 发送邮箱验证码 +// @Description 向指定邮箱发送验证码(用于注册场景) +// @Tags user +// @Accept json +// @Produce json +// @Param request body user.sendEmailCodeRequest true "发送验证码请求参数" +// @Success 200 {object} util.ResponseAny "发送成功" +// @Failure 400 {object} util.ResponseAny "参数错误" +// @Router /api/v1/user/send-email-code [post] +func SendEmailCode(c *gin.Context) { + var req sendEmailCodeRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, util.Err(err.Error())) + return + } + + req.Email = strings.TrimSpace(req.Email) + if req.Email == "" { + c.JSON(http.StatusOK, util.Err(errEmailRequired)) + return + } + + if req.Scene != "register" { + c.JSON(http.StatusOK, util.Err(errUnsupportedEmailScene)) + return + } + + ctx := c.Request.Context() + + var count int64 + if err := db.DB(ctx).Model(&model.User{}).Where("email = ?", req.Email).Count(&count).Error; err != nil { + c.JSON(http.StatusOK, util.Err(err.Error())) + return + } + if count > 0 { + c.JSON(http.StatusOK, util.Err(errEmailAlreadyRegistered)) + return + } + + cooldownKey := getEmailCooldownKey("register", req.Email) + var temp string + err := db.GetJSON(ctx, cooldownKey, &temp) + if err == nil { + c.JSON(http.StatusOK, util.Err(errEmailCodeCooldown)) + return + } + + if err := sendEmailVerificationCode(ctx, req.Email, "register", "register_email"); err != nil { + c.JSON(http.StatusOK, util.Err(err.Error())) + return + } + + c.JSON(http.StatusOK, util.OKNil()) +} + +func validateRegisterEmailVerification(ctx context.Context, req *registerRequest) error { + if !isEmailRegisterVerificationEnabled() { + return nil + } + if req.Email == "" || req.Code == "" { + return errors.New(errEmailOrCodeRequired) + } + if !verifyEmailCode(ctx, req.Email, "register", req.Code) { + return errors.New(errEmailCodeInvalidOrExpired) + } + return nil +} + +// completePendingOAuthBinding 完成登录后的 OAuth 待绑定绑定流程 +func completePendingOAuthBinding(session sessions.Session, user *model.User) { + pendingSourceID := session.Get(oauth.PendingOAuthSourceIDKey) + pendingExternalID := session.Get(oauth.PendingOAuthExternalIDKey) + pendingExternalUsername := session.Get(oauth.PendingOAuthExternalUsernameKey) + pendingEmail := session.Get(oauth.PendingOAuthEmailKey) + + if pendingSourceID == nil || pendingExternalID == nil { + return + } + + var sourceID uint64 + switch v := pendingSourceID.(type) { + case uint64: + sourceID = v + case int: + sourceID = uint64(v) + case float64: + sourceID = uint64(v) + } + externalID, _ := pendingExternalID.(string) + externalUsername, _ := pendingExternalUsername.(string) + email, _ := pendingEmail.(string) + + if sourceID != 0 && externalID != "" { + _ = model.BindExternalAccount(&model.ExternalAccount{ + AuthSourceID: sourceID, + UserID: user.ID, + ExternalID: externalID, + ExternalUsername: externalUsername, + Email: email, + }) + } + + session.Delete(oauth.PendingOAuthSourceIDKey) + session.Delete(oauth.PendingOAuthExternalIDKey) + session.Delete(oauth.PendingOAuthExternalUsernameKey) + session.Delete(oauth.PendingOAuthEmailKey) + _ = session.Save() +} + +type updateProfileRequest struct { + Nickname string `json:"nickname"` + Email string `json:"email"` + AvatarURL string `json:"avatar_url"` + Bio string `json:"bio"` + Phone string `json:"phone"` + Gender string `json:"gender"` + Website string `json:"website"` + Location string `json:"location"` +} + +// UpdateProfile 修改当前登录用户的个人资料 +// @Summary 修改当前登录用户的个人资料 +// @Description 修改当前登录用户的昵称、邮箱、头像、简介、电话、性别、个人网站和所在地。 +// @Tags user +// @Accept json +// @Produce json +// @Param request body user.updateProfileRequest true "更新请求参数" +// @Success 200 {object} util.ResponseAny{data=oauth.BasicUserInfo} "修改成功,返回更新后的用户信息" +// @Failure 400 {object} util.ResponseAny "邮箱已被占用或参数错误" +// @Failure 401 {object} util.ResponseAny "未登录" +// @Router /api/v1/user/profile [put] +func UpdateProfile(c *gin.Context) { + var req updateProfileRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusBadRequest, util.Err(err.Error())) + return + } + + userObj, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey) + if userObj == nil { + c.JSON(http.StatusUnauthorized, util.Err(errLoginRequired)) + return + } + + ctx := c.Request.Context() + var dbUser model.User + if err := db.DB(ctx).Where("id = ?", userObj.ID).First(&dbUser).Error; err != nil { + c.JSON(http.StatusOK, util.Err(errUserNotFound)) + return + } + + req.Email = strings.TrimSpace(req.Email) + if req.Email != "" && req.Email != dbUser.Email { + if !strings.Contains(req.Email, "@") || !strings.Contains(req.Email, ".") { + c.JSON(http.StatusOK, util.Err(errEmailFormatInvalid)) + return + } + + var count int64 + if err := db.DB(ctx).Model(&model.User{}).Where("email = ? AND id != ?", req.Email, dbUser.ID).Count(&count).Error; err != nil { + c.JSON(http.StatusOK, util.Err(err.Error())) + return + } + if count > 0 { + c.JSON(http.StatusOK, util.Err(errEmailAlreadyBound)) + return + } + } + + dbUser.Nickname = strings.TrimSpace(req.Nickname) + if dbUser.Nickname == "" { + dbUser.Nickname = dbUser.Username + } + dbUser.Email = req.Email + dbUser.AvatarURL = req.AvatarURL + dbUser.Bio = req.Bio + dbUser.Phone = strings.TrimSpace(req.Phone) + dbUser.Gender = strings.TrimSpace(req.Gender) + dbUser.Website = strings.TrimSpace(req.Website) + dbUser.Location = strings.TrimSpace(req.Location) + + if err := db.DB(ctx).Save(&dbUser).Error; err != nil { + c.JSON(http.StatusOK, util.Err(err.Error())) + return + } + + session := sessions.Default(c) + needChange := session.Get("need_change_password") == true + + c.JSON(http.StatusOK, util.OK(oauth.BuildBasicUserInfo(&dbUser, needChange))) +} diff --git a/internal/apps/user/routers.go b/internal/apps/user/routers.go index ce188f64..fe3d10c7 100644 --- a/internal/apps/user/routers.go +++ b/internal/apps/user/routers.go @@ -18,11 +18,6 @@ package user import ( "context" - "crypto/rand" - "encoding/json" - "errors" - "fmt" - "math/big" "net/http" "strings" "time" @@ -32,7 +27,6 @@ import ( "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/task" "github.com/Rain-kl/Wavelet/internal/util" "github.com/gin-contrib/sessions" "github.com/gin-gonic/gin" @@ -53,112 +47,6 @@ type registerRequest struct { Code string `json:"code"` } -type sendEmailCodeRequest struct { - Email string `json:"email" binding:"required,email"` - Scene string `json:"scene" binding:"required"` -} - -func isEmailLoginVerificationEnabled() bool { - enabled, err := model.GetBoolByKey(context.Background(), model.ConfigKeyEmailLoginVerificationEnabled) - if err != nil { - return false - } - return enabled -} - -func isEmailRegisterVerificationEnabled() bool { - enabled, err := model.GetBoolByKey(context.Background(), model.ConfigKeyEmailRegisterVerificationEnabled) - if err != nil { - return false - } - return enabled -} - -func isSMTPConfigured(ctx context.Context) bool { - var sc model.SystemConfig - var host, port, username string - - if err := sc.GetByKey(ctx, model.ConfigKeySMTPHost); err == nil { - host = sc.Value - } - if err := sc.GetByKey(ctx, model.ConfigKeySMTPPort); err == nil { - port = sc.Value - } - if err := sc.GetByKey(ctx, model.ConfigKeySMTPUsername); err == nil { - username = sc.Value - } - - return host != "" && port != "" && username != "" -} - -func generateVerificationCode() string { - n, _ := rand.Int(rand.Reader, big.NewInt(verificationCodeRange)) - return fmt.Sprintf("%06d", n.Int64()+verificationCodeOffset) -} - -func getEmailCodeKey(scene, email string) string { - return fmt.Sprintf("email_code:%s:%s", scene, email) -} - -func getEmailCooldownKey(email string) string { - return fmt.Sprintf("email_code:cooldown:%s", email) -} - -func sendEmailVerificationCode(ctx context.Context, email, scene, templateName string) error { - // 校验 SMTP 配置是否完整 - if !isSMTPConfigured(ctx) { - return errors.New(errSMTPConfigIncomplete) - } - - code := generateVerificationCode() - codeKey := getEmailCodeKey(scene, email) - cooldownKey := getEmailCooldownKey(email) - - // 使用模板管理获取并渲染邮件标题和正文。模板缺失或渲染失败时不发送验证码。 - emailSubject, emailBody, err := model.RenderTemplate( - ctx, - templateName, - map[string]any{"Code": code}, - ) - if err != nil { - return fmt.Errorf(errRenderEmailTemplateFailed, err) - } - - // 存验证码,5分钟有效 - if err := db.SetJSON(ctx, codeKey, code, emailCodeExpiry); err != nil { - return errors.New(errGenerateEmailCodeFailed) - } - // 存冷却,60秒有效 - _ = db.SetJSON(ctx, cooldownKey, "1", emailCodeCooldown) - - // 构建异步邮件发送任务 - payload := SendEmailPayload{ - To: email, - Subject: emailSubject, - Body: emailBody, - } - payloadBytes, _ := json.Marshal(payload) - _, err = task.DispatchTask(ctx, task.TaskTypeSendEmail, payloadBytes, "system") - if err != nil { - return errors.New(errDispatchEmailTaskFailed) - } - return nil -} - -func verifyEmailCode(ctx context.Context, email, scene, code string) bool { - codeKey := getEmailCodeKey(scene, email) - var storedCode string - if err := db.GetJSON(ctx, codeKey, &storedCode); err != nil { - return false - } - if storedCode != code { - return false - } - // 验证成功,删除验证码 - _ = db.Redis.Del(ctx, db.PrefixedKey(codeKey)).Err() - return true -} - func isPasswordLoginEnabled() bool { enabled, err := model.GetBoolByKey(context.Background(), model.ConfigKeyPasswordLoginEnabled) if err != nil { @@ -193,40 +81,6 @@ func setLoginSession(c *gin.Context, user *model.User) error { return nil } -// handleLoginEmailVerification 处理登录时的邮箱验证码校验流程 -func handleLoginEmailVerification(ctx context.Context, c *gin.Context, req *loginRequest, user *model.User) error { - if user.Email == "" { - c.JSON(http.StatusOK, util.Err(errLoginEmailMissing)) - return errors.New("handled") - } - - if req.Code == "" { - // 校验 Redis 发送冷却时间 - cooldownKey := getEmailCooldownKey(user.Email) - var temp string - err := db.GetJSON(ctx, cooldownKey, &temp) - if err != nil { - // 没有冷却,触发验证码发送 - if err := sendEmailVerificationCode(ctx, user.Email, "login", "login_email"); err != nil { - c.JSON(http.StatusOK, util.Err(err.Error())) - return errors.New("handled") - } - } - - // 脱敏邮箱并返回错误,提示前端需要输入验证码 - maskedEmail := util.MaskEmail(user.Email) - c.JSON(http.StatusOK, util.Err(errNeedEmailCodePrefix+maskedEmail)) - return errors.New("handled") - } - - // 校验验证码 - if !verifyEmailCode(ctx, user.Email, "login", req.Code) { - c.JSON(http.StatusOK, util.Err(errEmailCodeInvalidOrExpired)) - return errors.New("handled") - } - return nil -} - // Login 用户密码登录 // @Summary 用户密码登录 // @Description 使用用户名和密码登录,登录成功后建立 Session。若管理员已关闭密码登录功能则返回错误。 @@ -477,202 +331,3 @@ func ChangePassword(c *gin.Context) { c.JSON(http.StatusOK, util.OK("密码修改成功")) } - -// SendEmailCode 发送邮箱验证码 -// @Summary 发送邮箱验证码 -// @Description 向指定邮箱发送验证码(用于注册场景) -// @Tags user -// @Accept json -// @Produce json -// @Param request body user.sendEmailCodeRequest true "发送验证码请求参数" -// @Success 200 {object} util.ResponseAny "发送成功" -// @Failure 400 {object} util.ResponseAny "参数错误" -// @Router /api/v1/user/send-email-code [post] -func SendEmailCode(c *gin.Context) { - var req sendEmailCodeRequest - if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, util.Err(err.Error())) - return - } - - req.Email = strings.TrimSpace(req.Email) - if req.Email == "" { - c.JSON(http.StatusOK, util.Err(errEmailRequired)) - return - } - - if req.Scene != "register" { - c.JSON(http.StatusOK, util.Err(errUnsupportedEmailScene)) - return - } - - ctx := c.Request.Context() - - // 1. 检查邮箱是否已被注册 - var count int64 - if err := db.DB(ctx).Model(&model.User{}).Where("email = ?", req.Email).Count(&count).Error; err != nil { - c.JSON(http.StatusOK, util.Err(err.Error())) - return - } - if count > 0 { - c.JSON(http.StatusOK, util.Err(errEmailAlreadyRegistered)) - return - } - - // 2. 校验 Redis 发送冷却时间 - cooldownKey := getEmailCooldownKey(req.Email) - var temp string - err := db.GetJSON(ctx, cooldownKey, &temp) - if err == nil { - c.JSON(http.StatusOK, util.Err(errEmailCodeCooldown)) - return - } - - // 3. 发送验证码 - if err := sendEmailVerificationCode(ctx, req.Email, "register", "register_email"); err != nil { - c.JSON(http.StatusOK, util.Err(err.Error())) - return - } - - c.JSON(http.StatusOK, util.OKNil()) -} - -type updateProfileRequest struct { - Nickname string `json:"nickname"` - Email string `json:"email"` - AvatarURL string `json:"avatar_url"` - Bio string `json:"bio"` - Phone string `json:"phone"` - Gender string `json:"gender"` - Website string `json:"website"` - Location string `json:"location"` -} - -// UpdateProfile 修改当前登录用户的个人资料 -// @Summary 修改当前登录用户的个人资料 -// @Description 修改当前登录用户的昵称、邮箱、头像、简介、电话、性别、个人网站和所在地。 -// @Tags user -// @Accept json -// @Produce json -// @Param request body user.updateProfileRequest true "更新请求参数" -// @Success 200 {object} util.ResponseAny{data=oauth.BasicUserInfo} "修改成功,返回更新后的用户信息" -// @Failure 400 {object} util.ResponseAny "邮箱已被占用或参数错误" -// @Failure 401 {object} util.ResponseAny "未登录" -// @Router /api/v1/user/profile [put] -func UpdateProfile(c *gin.Context) { - var req updateProfileRequest - if err := c.ShouldBindJSON(&req); err != nil { - c.JSON(http.StatusBadRequest, util.Err(err.Error())) - return - } - - userObj, _ := util.GetFromContext[*model.User](c, oauth.UserObjKey) - if userObj == nil { - c.JSON(http.StatusUnauthorized, util.Err(errLoginRequired)) - return - } - - ctx := c.Request.Context() - var dbUser model.User - if err := db.DB(ctx).Where("id = ?", userObj.ID).First(&dbUser).Error; err != nil { - c.JSON(http.StatusOK, util.Err(errUserNotFound)) - return - } - - // 校验邮箱格式与唯一性 - req.Email = strings.TrimSpace(req.Email) - if req.Email != "" && req.Email != dbUser.Email { - if !strings.Contains(req.Email, "@") || !strings.Contains(req.Email, ".") { - c.JSON(http.StatusOK, util.Err(errEmailFormatInvalid)) - return - } - - var count int64 - if err := db.DB(ctx).Model(&model.User{}).Where("email = ? AND id != ?", req.Email, dbUser.ID).Count(&count).Error; err != nil { - c.JSON(http.StatusOK, util.Err(err.Error())) - return - } - if count > 0 { - c.JSON(http.StatusOK, util.Err(errEmailAlreadyBound)) - return - } - } - - // 更新字段 - dbUser.Nickname = strings.TrimSpace(req.Nickname) - if dbUser.Nickname == "" { - dbUser.Nickname = dbUser.Username - } - dbUser.Email = req.Email - dbUser.AvatarURL = req.AvatarURL - dbUser.Bio = req.Bio - dbUser.Phone = strings.TrimSpace(req.Phone) - dbUser.Gender = strings.TrimSpace(req.Gender) - dbUser.Website = strings.TrimSpace(req.Website) - dbUser.Location = strings.TrimSpace(req.Location) - - if err := db.DB(ctx).Save(&dbUser).Error; err != nil { - c.JSON(http.StatusOK, util.Err(err.Error())) - return - } - - session := sessions.Default(c) - needChange := session.Get("need_change_password") == true - - c.JSON(http.StatusOK, util.OK(oauth.BuildBasicUserInfo(&dbUser, needChange))) -} - -// validateRegisterEmailVerification 校验注册时的邮箱验证码 -func validateRegisterEmailVerification(ctx context.Context, req *registerRequest) error { - if !isEmailRegisterVerificationEnabled() { - return nil - } - if req.Email == "" || req.Code == "" { - return errors.New(errEmailOrCodeRequired) - } - if !verifyEmailCode(ctx, req.Email, "register", req.Code) { - return errors.New(errEmailCodeInvalidOrExpired) - } - return nil -} - -// completePendingOAuthBinding 完成登录后的 OAuth 待绑定绑定流程 -func completePendingOAuthBinding(session sessions.Session, user *model.User) { - pendingSourceID := session.Get(oauth.PendingOAuthSourceIDKey) - pendingExternalID := session.Get(oauth.PendingOAuthExternalIDKey) - pendingExternalUsername := session.Get(oauth.PendingOAuthExternalUsernameKey) - pendingEmail := session.Get(oauth.PendingOAuthEmailKey) - - if pendingSourceID == nil || pendingExternalID == nil { - return - } - - var sourceID uint64 - switch v := pendingSourceID.(type) { - case uint64: - sourceID = v - case int: - sourceID = uint64(v) - case float64: - sourceID = uint64(v) - } - externalID, _ := pendingExternalID.(string) - externalUsername, _ := pendingExternalUsername.(string) - email, _ := pendingEmail.(string) - - if sourceID != 0 && externalID != "" { - _ = model.BindExternalAccount(&model.ExternalAccount{ - AuthSourceID: sourceID, - UserID: user.ID, - ExternalID: externalID, - ExternalUsername: externalUsername, - Email: email, - }) - } - // 清除 pending 信息 - session.Delete(oauth.PendingOAuthSourceIDKey) - session.Delete(oauth.PendingOAuthExternalIDKey) - session.Delete(oauth.PendingOAuthExternalUsernameKey) - session.Delete(oauth.PendingOAuthEmailKey) - _ = session.Save() -} diff --git a/internal/apps/user/routers_test.go b/internal/apps/user/routers_test.go index bd96ef18..363a112e 100644 --- a/internal/apps/user/routers_test.go +++ b/internal/apps/user/routers_test.go @@ -117,6 +117,34 @@ func basicUserInfoFromResponse(t *testing.T, w *httptest.ResponseRecorder) oauth return info } +func TestEmailCooldownKeyIncludesScene(t *testing.T) { + email := "user@example.com" + + loginKey := getEmailCooldownKey("login", email) + registerKey := getEmailCooldownKey("register", email) + if loginKey == registerKey { + t.Errorf("getEmailCooldownKey(%q, %q) = %q, want different key from register scene", "login", email, loginKey) + } + if want := "email_code:cooldown:login:user@example.com"; loginKey != want { + t.Errorf("getEmailCooldownKey(%q, %q) = %q, want %q", "login", email, loginKey, want) + } +} + +func TestGenerateVerificationCode(t *testing.T) { + code, err := generateVerificationCode() + if err != nil { + t.Fatalf("generateVerificationCode() error = %v, want nil", err) + } + if len(code) != 6 { + t.Fatalf("generateVerificationCode() length = %d, want 6. Code: %q", len(code), code) + } + for _, r := range code { + if r < '0' || r > '9' { + t.Fatalf("generateVerificationCode() = %q, want only digits", code) + } + } +} + func TestRegisterCreatesAuthenticatedEncryptedUser(t *testing.T) { dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) defer cleanup()