mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-29 14:06:36 +08:00
质量优化
This commit is contained in:
+27
-30
@@ -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
|
||||
|
||||
@@ -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"`
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 // 上传文件权限
|
||||
)
|
||||
@@ -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"
|
||||
)
|
||||
|
||||
@@ -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"`
|
||||
}
|
||||
|
||||
@@ -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)))
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user