mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-07 08:06:37 +08:00
重构
This commit is contained in:
@@ -91,19 +91,19 @@ func (source *AuthSource) Normalize() {
|
||||
func (source *AuthSource) Validate() error {
|
||||
source.Normalize()
|
||||
if source.Name == "" {
|
||||
return errors.New("认证源名称不能为空")
|
||||
return errors.New(errAuthSourceNameRequired)
|
||||
}
|
||||
if !authSourceNamePattern.MatchString(source.Name) {
|
||||
return errors.New("认证源名称只能包含字母、数字、短横线或下划线,且必须以字母或数字开头")
|
||||
return errors.New(errAuthSourceNameInvalid)
|
||||
}
|
||||
if source.Type != AuthSourceTypeOIDC {
|
||||
return errors.New("认证源类型仅支持 oidc")
|
||||
return errors.New(errAuthSourceTypeUnsupported)
|
||||
}
|
||||
if source.OpenIDDiscoveryURL == "" {
|
||||
return errors.New("OIDC 认证源必须配置 Discovery URL")
|
||||
return errors.New(errAuthSourceDiscoveryURLRequired)
|
||||
}
|
||||
if source.IsActive && (source.ClientID == "" || source.ClientSecret == "") {
|
||||
return errors.New("启用认证源前必须配置 Client ID 和 Client Secret")
|
||||
return errors.New(errAuthSourceClientCredentialsRequired)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -137,7 +137,7 @@ func GetActiveAuthSources() ([]AuthSource, error) {
|
||||
|
||||
func GetAuthSourceByID(id uint64) (*AuthSource, error) {
|
||||
if id == 0 {
|
||||
return nil, errors.New("认证源 ID 不能为空")
|
||||
return nil, errors.New(errAuthSourceIDRequired)
|
||||
}
|
||||
var source AuthSource
|
||||
if err := db.DB(context.Background()).First(&source, "id = ?", id).Error; err != nil {
|
||||
@@ -150,7 +150,7 @@ func GetAuthSourceByID(id uint64) (*AuthSource, error) {
|
||||
func GetAuthSourceByName(name string) (*AuthSource, error) {
|
||||
name = strings.TrimSpace(name)
|
||||
if name == "" {
|
||||
return nil, errors.New("认证源名称不能为空")
|
||||
return nil, errors.New(errAuthSourceNameRequired)
|
||||
}
|
||||
var source AuthSource
|
||||
if err := db.DB(context.Background()).First(&source, "name = ?", name).Error; err != nil {
|
||||
@@ -169,7 +169,7 @@ func CreateAuthSource(source *AuthSource) error {
|
||||
|
||||
func UpdateAuthSource(source *AuthSource, keepSecret bool) error {
|
||||
if source.ID == 0 {
|
||||
return errors.New("认证源 ID 不能为空")
|
||||
return errors.New(errAuthSourceIDRequired)
|
||||
}
|
||||
var current AuthSource
|
||||
if err := db.DB(context.Background()).First(¤t, "id = ?", source.ID).Error; err != nil {
|
||||
@@ -208,7 +208,7 @@ func ToggleAuthSource(id uint64, isActive bool) error {
|
||||
|
||||
func DeleteAuthSource(id uint64) error {
|
||||
if id == 0 {
|
||||
return errors.New("认证源 ID 不能为空")
|
||||
return errors.New(errAuthSourceIDRequired)
|
||||
}
|
||||
return db.DB(context.Background()).Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("auth_source_id = ?", id).Delete(&ExternalAccount{}).Error; err != nil {
|
||||
@@ -228,7 +228,7 @@ func FindExternalAccount(sourceID uint64, externalID string) (*ExternalAccount,
|
||||
|
||||
func BindExternalAccount(account *ExternalAccount) error {
|
||||
if account.UserID == 0 || strings.TrimSpace(account.ExternalID) == "" {
|
||||
return errors.New("外部账号绑定信息不完整")
|
||||
return errors.New(errExternalAccountBindingIncomplete)
|
||||
}
|
||||
account.ExternalID = strings.TrimSpace(account.ExternalID)
|
||||
account.ExternalUsername = strings.TrimSpace(account.ExternalUsername)
|
||||
@@ -239,7 +239,7 @@ func BindExternalAccount(account *ExternalAccount) error {
|
||||
err := tx.Where("auth_source_id = ? AND external_id = ?", account.AuthSourceID, account.ExternalID).First(¤t).Error
|
||||
if err == nil {
|
||||
if current.UserID != account.UserID {
|
||||
return errors.New("该外部账号已绑定到其他用户")
|
||||
return errors.New(errExternalAccountAlreadyBoundToAnother)
|
||||
}
|
||||
return tx.Model(¤t).Updates(map[string]any{
|
||||
"external_username": account.ExternalUsername,
|
||||
@@ -255,7 +255,7 @@ func BindExternalAccount(account *ExternalAccount) error {
|
||||
|
||||
func ListExternalAccountsByUserID(userID uint64) ([]ExternalAccountView, error) {
|
||||
if userID == 0 {
|
||||
return nil, errors.New("用户 ID 不能为空")
|
||||
return nil, errors.New(errUserIDRequired)
|
||||
}
|
||||
var accounts []ExternalAccount
|
||||
if err := db.DB(context.Background()).Where("user_id = ?", userID).Order("id asc").Find(&accounts).Error; err != nil {
|
||||
@@ -296,7 +296,7 @@ func ListExternalAccountsByUserID(userID uint64) ([]ExternalAccountView, error)
|
||||
|
||||
func DeleteExternalAccountForUser(id uint64, userID uint64) error {
|
||||
if id == 0 || userID == 0 {
|
||||
return errors.New("绑定记录 ID 不能为空")
|
||||
return errors.New(errExternalAccountBindingIDRequired)
|
||||
}
|
||||
return db.DB(context.Background()).Where("id = ? AND user_id = ?", id, userID).Delete(&ExternalAccount{}).Error
|
||||
}
|
||||
|
||||
@@ -0,0 +1,43 @@
|
||||
/*
|
||||
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 model
|
||||
|
||||
const (
|
||||
errRegistrationDisabled = "注册已关闭"
|
||||
errDatabaseNotInitialized = "database not initialized"
|
||||
errUsernameExists = "用户名已存在"
|
||||
errEmailAlreadyBound = "该邮箱已被其他账号绑定"
|
||||
errConfigIntParseFailed = "配置 %s 的值 '%s' 无法转换为整数: %w"
|
||||
errConfigDecimalParseFailed = "配置 %s 的值 '%s' 无法转换为decimal: %w"
|
||||
errConfigBoolParseFailed = "配置 %s 的值 '%s' 无法转换为布尔值: %w"
|
||||
errParseMenuDisplayConfigFailed = "解析目录显示配置失败: %w"
|
||||
errTemplateKeyRequired = "模板标识符不能为空"
|
||||
errTemplateNameRequired = "模板名称不能为空"
|
||||
errTemplateContentRequired = "模板内容不能为空"
|
||||
errTemplateUnavailable = "模板 %s 不存在或不可用: %w"
|
||||
errTemplateRenderFailed = "模板 %s 渲染失败: %w"
|
||||
errAuthSourceNameRequired = "认证源名称不能为空"
|
||||
errAuthSourceNameInvalid = "认证源名称只能包含字母、数字、短横线或下划线,且必须以字母或数字开头"
|
||||
errAuthSourceTypeUnsupported = "认证源类型仅支持 oidc"
|
||||
errAuthSourceDiscoveryURLRequired = "OIDC 认证源必须配置 Discovery URL"
|
||||
errAuthSourceClientCredentialsRequired = "启用认证源前必须配置 Client ID 和 Client Secret"
|
||||
errAuthSourceIDRequired = "认证源 ID 不能为空"
|
||||
errExternalAccountBindingIncomplete = "外部账号绑定信息不完整"
|
||||
errExternalAccountAlreadyBoundToAnother = "该外部账号已绑定到其他用户"
|
||||
errUserIDRequired = "用户 ID 不能为空"
|
||||
errExternalAccountBindingIDRequired = "绑定记录 ID 不能为空"
|
||||
)
|
||||
@@ -85,7 +85,7 @@ func (sc *SystemConfig) GetByKey(ctx context.Context, key string) error {
|
||||
// 查数据库
|
||||
database := db.DB(ctx)
|
||||
if database == nil {
|
||||
return errors.New("database not initialized")
|
||||
return errors.New(errDatabaseNotInitialized)
|
||||
}
|
||||
|
||||
if err := database.Where("key = ?", key).First(sc).Error; err != nil {
|
||||
@@ -109,7 +109,7 @@ func GetIntByKey(ctx context.Context, key string) (int, error) {
|
||||
|
||||
value, err := strconv.Atoi(sc.Value)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("配置 %s 的值 '%s' 无法转换为整数: %w", key, sc.Value, err)
|
||||
return 0, fmt.Errorf(errConfigIntParseFailed, key, sc.Value, err)
|
||||
}
|
||||
|
||||
return value, nil
|
||||
@@ -125,7 +125,7 @@ func GetDecimalByKey(ctx context.Context, key string, precision int32) (decimal.
|
||||
|
||||
value, err := decimal.NewFromString(sc.Value)
|
||||
if err != nil {
|
||||
return decimal.Zero, fmt.Errorf("配置 %s 的值 '%s' 无法转换为decimal: %w", key, sc.Value, err)
|
||||
return decimal.Zero, fmt.Errorf(errConfigDecimalParseFailed, key, sc.Value, err)
|
||||
}
|
||||
|
||||
// 裁剪到指定小数位数
|
||||
@@ -141,7 +141,7 @@ func GetBoolByKey(ctx context.Context, key string) (bool, error) {
|
||||
|
||||
value, err := strconv.ParseBool(sc.Value)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("配置 %s 的值 '%s' 无法转换为布尔值: %w", key, sc.Value, err)
|
||||
return false, fmt.Errorf(errConfigBoolParseFailed, key, sc.Value, err)
|
||||
}
|
||||
|
||||
return value, nil
|
||||
@@ -160,7 +160,7 @@ func GetMenuDisplayConfig(ctx context.Context) (map[string]bool, error) {
|
||||
}
|
||||
|
||||
if err := json.Unmarshal([]byte(sc.Value), &config); err != nil {
|
||||
return nil, fmt.Errorf("解析目录显示配置失败: %w", err)
|
||||
return nil, fmt.Errorf(errParseMenuDisplayConfigFailed, err)
|
||||
}
|
||||
|
||||
return config, nil
|
||||
|
||||
+12
-21
@@ -20,6 +20,7 @@ import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"text/template"
|
||||
"time"
|
||||
@@ -55,13 +56,13 @@ func (t *Template) Normalize() {
|
||||
func (t *Template) Validate() error {
|
||||
t.Normalize()
|
||||
if t.Key == "" {
|
||||
return errors.New("模板标识符不能为空")
|
||||
return errors.New(errTemplateKeyRequired)
|
||||
}
|
||||
if t.Name == "" {
|
||||
return errors.New("模板名称不能为空")
|
||||
return errors.New(errTemplateNameRequired)
|
||||
}
|
||||
if t.Content == "" {
|
||||
return errors.New("模板内容不能为空")
|
||||
return errors.New(errTemplateContentRequired)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -94,26 +95,16 @@ func (t *Template) Render(data any) (string, string, error) {
|
||||
return subject, bodyBuf.String(), nil
|
||||
}
|
||||
|
||||
// RenderTemplate 渲染模板的高级包装。如果读取或渲染失败,将使用 fallbackSubject 和 fallbackBody 进行解析和返回。
|
||||
func RenderTemplate(ctx context.Context, key string, data any, fallbackSubject, fallbackBody string) (string, string) {
|
||||
// RenderTemplate 渲染指定模板。模板不存在或渲染失败时返回错误,由调用方决定如何处理。
|
||||
func RenderTemplate(ctx context.Context, key string, data any) (string, string, error) {
|
||||
var t Template
|
||||
if err := db.DB(ctx).Where("key = ?", key).First(&t).Error; err == nil {
|
||||
subject, body, err := t.Render(data)
|
||||
if err == nil {
|
||||
return subject, body
|
||||
}
|
||||
if err := db.DB(ctx).Where("key = ?", key).First(&t).Error; err != nil {
|
||||
return "", "", fmt.Errorf(errTemplateUnavailable, key, err)
|
||||
}
|
||||
|
||||
// 降级使用传入的默认模板内容渲染
|
||||
tFallback := Template{
|
||||
Key: key + "_fallback",
|
||||
Subject: fallbackSubject,
|
||||
Content: fallbackBody,
|
||||
subject, body, err := t.Render(data)
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf(errTemplateRenderFailed, key, err)
|
||||
}
|
||||
subject, body, err := tFallback.Render(data)
|
||||
if err == nil {
|
||||
return subject, body
|
||||
}
|
||||
|
||||
return fallbackSubject, fallbackBody
|
||||
return subject, body, nil
|
||||
}
|
||||
|
||||
@@ -92,12 +92,15 @@ func (u *User) SetEncryptedPassword(password string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (u *User) IsPasswordEncrypted() bool {
|
||||
return strings.HasPrefix(u.Password, "$2a$") || strings.HasPrefix(u.Password, "$2b$") || strings.HasPrefix(u.Password, "$2y$")
|
||||
}
|
||||
|
||||
func (u *User) CheckPassword(password string) bool {
|
||||
if u.Password == "" || password == "" {
|
||||
return false
|
||||
}
|
||||
isBcrypt := strings.HasPrefix(u.Password, "$2a$") || strings.HasPrefix(u.Password, "$2b$") || strings.HasPrefix(u.Password, "$2y$")
|
||||
if isBcrypt {
|
||||
if u.IsPasswordEncrypted() {
|
||||
return util.CheckPasswordHash(u.Password, password)
|
||||
}
|
||||
return u.Password == password
|
||||
@@ -132,7 +135,7 @@ func (u *User) CheckActive() error {
|
||||
func (u *User) CreateUser(ctx context.Context, tx *gorm.DB, oauthInfo *OAuthUserInfo) error {
|
||||
enabled, err := GetBoolByKey(ctx, ConfigKeyRegistrationEnabled)
|
||||
if err == nil && !enabled {
|
||||
return errors.New("注册已关闭")
|
||||
return errors.New(errRegistrationDisabled)
|
||||
}
|
||||
|
||||
now := time.Now()
|
||||
@@ -158,7 +161,7 @@ func (u *User) CreateUser(ctx context.Context, tx *gorm.DB, oauthInfo *OAuthUser
|
||||
func (u *User) RegisterUser(ctx context.Context, tx *gorm.DB) error {
|
||||
enabled, err := GetBoolByKey(ctx, ConfigKeyRegistrationEnabled)
|
||||
if err == nil && !enabled {
|
||||
return errors.New("注册已关闭")
|
||||
return errors.New(errRegistrationDisabled)
|
||||
}
|
||||
|
||||
// 检查用户名冲突
|
||||
@@ -167,7 +170,7 @@ func (u *User) RegisterUser(ctx context.Context, tx *gorm.DB) error {
|
||||
return err
|
||||
}
|
||||
if count > 0 {
|
||||
return errors.New("用户名已存在")
|
||||
return errors.New(errUsernameExists)
|
||||
}
|
||||
|
||||
// 检查邮箱冲突
|
||||
@@ -177,7 +180,7 @@ func (u *User) RegisterUser(ctx context.Context, tx *gorm.DB) error {
|
||||
return err
|
||||
}
|
||||
if emailCount > 0 {
|
||||
return errors.New("该邮箱已被其他账号绑定")
|
||||
return errors.New(errEmailAlreadyBound)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user