mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-04 23:16:37 +08:00
OIDC
This commit is contained in:
@@ -0,0 +1,218 @@
|
||||
package model
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"regexp"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const (
|
||||
AuthSourceTypeGitHub = "github"
|
||||
AuthSourceTypeOIDC = "oidc"
|
||||
)
|
||||
|
||||
var authSourceNamePattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9_-]{0,79}$`)
|
||||
|
||||
type AuthSource struct {
|
||||
ID uint `json:"id"`
|
||||
Name string `json:"name" gorm:"uniqueIndex;size:80;not null"`
|
||||
Type string `json:"type" gorm:"index;size:20;not null"`
|
||||
DisplayName string `json:"display_name" gorm:"size:100"`
|
||||
IsActive bool `json:"is_active" gorm:"index;not null;default:false"`
|
||||
ClientID string `json:"client_id" gorm:"column:client_id;size:255"`
|
||||
ClientSecret string `json:"-" gorm:"column:client_secret;size:1024"`
|
||||
OpenIDDiscoveryURL string `json:"openid_discovery_url" gorm:"column:openid_discovery_url;size:1024"`
|
||||
Scopes string `json:"scopes" gorm:"size:255"`
|
||||
IconURL string `json:"icon_url" gorm:"size:1024"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
ClientSecretConfigured bool `json:"client_secret_configured" gorm:"-"`
|
||||
}
|
||||
|
||||
type ExternalAccount struct {
|
||||
ID uint `json:"id"`
|
||||
AuthSourceID uint `json:"auth_source_id" gorm:"uniqueIndex:idx_external_account_source_external;index;not null"`
|
||||
UserID int `json:"user_id" gorm:"index;not null"`
|
||||
ExternalID string `json:"external_id" gorm:"uniqueIndex:idx_external_account_source_external;size:255;not null"`
|
||||
ExternalUsername string `json:"external_username" gorm:"size:255"`
|
||||
Email string `json:"email" gorm:"size:255"`
|
||||
AuthSource AuthSource `json:"-" gorm:"constraint:OnDelete:CASCADE"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
func (source *AuthSource) Normalize() {
|
||||
source.Name = strings.TrimSpace(source.Name)
|
||||
source.Type = strings.TrimSpace(strings.ToLower(source.Type))
|
||||
source.DisplayName = strings.TrimSpace(source.DisplayName)
|
||||
source.ClientID = strings.TrimSpace(source.ClientID)
|
||||
source.ClientSecret = strings.TrimSpace(source.ClientSecret)
|
||||
source.OpenIDDiscoveryURL = strings.TrimSpace(source.OpenIDDiscoveryURL)
|
||||
source.Scopes = strings.TrimSpace(source.Scopes)
|
||||
source.IconURL = strings.TrimSpace(source.IconURL)
|
||||
if source.DisplayName == "" {
|
||||
source.DisplayName = source.Name
|
||||
}
|
||||
if source.Type == AuthSourceTypeOIDC && source.Scopes == "" {
|
||||
source.Scopes = "openid profile email"
|
||||
}
|
||||
if source.Type == AuthSourceTypeGitHub && source.Scopes == "" {
|
||||
source.Scopes = "user:email"
|
||||
}
|
||||
}
|
||||
|
||||
func (source *AuthSource) Validate() error {
|
||||
source.Normalize()
|
||||
if source.Name == "" {
|
||||
return errors.New("认证源名称不能为空")
|
||||
}
|
||||
if !authSourceNamePattern.MatchString(source.Name) {
|
||||
return errors.New("认证源名称只能包含字母、数字、短横线或下划线,且必须以字母或数字开头")
|
||||
}
|
||||
switch source.Type {
|
||||
case AuthSourceTypeGitHub:
|
||||
case AuthSourceTypeOIDC:
|
||||
if source.OpenIDDiscoveryURL == "" {
|
||||
return errors.New("OIDC 认证源必须配置 Discovery URL")
|
||||
}
|
||||
default:
|
||||
return errors.New("认证源类型仅支持 github 或 oidc")
|
||||
}
|
||||
if source.IsActive {
|
||||
if source.ClientID == "" || source.ClientSecret == "" {
|
||||
return errors.New("启用认证源前必须配置 Client ID 和 Client Secret")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (source *AuthSource) Sanitize() {
|
||||
source.ClientSecretConfigured = source.ClientSecret != ""
|
||||
source.ClientSecret = ""
|
||||
}
|
||||
|
||||
func GetAuthSources() ([]AuthSource, error) {
|
||||
var sources []AuthSource
|
||||
err := DB.Order("id asc").Find(&sources).Error
|
||||
for index := range sources {
|
||||
sources[index].Sanitize()
|
||||
}
|
||||
return sources, err
|
||||
}
|
||||
|
||||
func GetActiveAuthSources() ([]AuthSource, error) {
|
||||
var sources []AuthSource
|
||||
err := DB.Where("is_active = ?", true).Order("id asc").Find(&sources).Error
|
||||
for index := range sources {
|
||||
sources[index].Sanitize()
|
||||
}
|
||||
return sources, err
|
||||
}
|
||||
|
||||
func GetAuthSourceByID(id uint) (*AuthSource, error) {
|
||||
if id == 0 {
|
||||
return nil, errors.New("认证源 ID 不能为空")
|
||||
}
|
||||
var source AuthSource
|
||||
if err := DB.First(&source, "id = ?", id).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
source.ClientSecretConfigured = source.ClientSecret != ""
|
||||
return &source, nil
|
||||
}
|
||||
|
||||
func GetAuthSourceByName(name string) (*AuthSource, error) {
|
||||
name = strings.TrimSpace(name)
|
||||
if name == "" {
|
||||
return nil, errors.New("认证源名称不能为空")
|
||||
}
|
||||
var source AuthSource
|
||||
if err := DB.First(&source, "name = ?", name).Error; err != nil {
|
||||
return nil, err
|
||||
}
|
||||
source.ClientSecretConfigured = source.ClientSecret != ""
|
||||
return &source, nil
|
||||
}
|
||||
|
||||
func CreateAuthSource(source *AuthSource) error {
|
||||
if err := source.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
return DB.Create(source).Error
|
||||
}
|
||||
|
||||
func UpdateAuthSource(source *AuthSource, keepSecret bool) error {
|
||||
if source.ID == 0 {
|
||||
return errors.New("认证源 ID 不能为空")
|
||||
}
|
||||
var current AuthSource
|
||||
if err := DB.First(¤t, "id = ?", source.ID).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if keepSecret {
|
||||
source.ClientSecret = current.ClientSecret
|
||||
}
|
||||
if err := source.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
return DB.Model(¤t).Updates(map[string]any{
|
||||
"name": source.Name,
|
||||
"type": source.Type,
|
||||
"display_name": source.DisplayName,
|
||||
"is_active": source.IsActive,
|
||||
"client_id": source.ClientID,
|
||||
"client_secret": source.ClientSecret,
|
||||
"openid_discovery_url": source.OpenIDDiscoveryURL,
|
||||
"scopes": source.Scopes,
|
||||
"icon_url": source.IconURL,
|
||||
}).Error
|
||||
}
|
||||
|
||||
func ToggleAuthSource(id uint, isActive bool) error {
|
||||
source, err := GetAuthSourceByID(id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
source.IsActive = isActive
|
||||
if err := source.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
return DB.Model(&AuthSource{}).Where("id = ?", id).Update("is_active", isActive).Error
|
||||
}
|
||||
|
||||
func DeleteAuthSource(id uint) error {
|
||||
if id == 0 {
|
||||
return errors.New("认证源 ID 不能为空")
|
||||
}
|
||||
return DB.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("auth_source_id = ?", id).Delete(&ExternalAccount{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Delete(&AuthSource{}, "id = ?", id).Error
|
||||
})
|
||||
}
|
||||
|
||||
func FindExternalAccount(sourceID uint, externalID string) (*ExternalAccount, error) {
|
||||
var account ExternalAccount
|
||||
err := DB.Where("auth_source_id = ? AND external_id = ?", sourceID, externalID).First(&account).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &account, nil
|
||||
}
|
||||
|
||||
func LinkExternalAccount(account *ExternalAccount) error {
|
||||
if account.AuthSourceID == 0 || account.UserID == 0 || strings.TrimSpace(account.ExternalID) == "" {
|
||||
return errors.New("外部账号绑定信息不完整")
|
||||
}
|
||||
account.ExternalID = strings.TrimSpace(account.ExternalID)
|
||||
account.ExternalUsername = strings.TrimSpace(account.ExternalUsername)
|
||||
account.Email = strings.TrimSpace(account.Email)
|
||||
return DB.Where(ExternalAccount{
|
||||
AuthSourceID: account.AuthSourceID,
|
||||
ExternalID: account.ExternalID,
|
||||
}).FirstOrCreate(account).Error
|
||||
}
|
||||
@@ -4,7 +4,7 @@ import "time"
|
||||
|
||||
const (
|
||||
legacyDatabaseSchemaVersion = 1
|
||||
currentDatabaseSchemaVersion = 9
|
||||
currentDatabaseSchemaVersion = 10
|
||||
databaseSchemaVersionRowID = 1
|
||||
)
|
||||
|
||||
|
||||
@@ -26,6 +26,8 @@ func registeredModels() []any {
|
||||
return []any{
|
||||
&File{},
|
||||
&User{},
|
||||
&AuthSource{},
|
||||
&ExternalAccount{},
|
||||
&Option{},
|
||||
&Origin{},
|
||||
&ProxyRoute{},
|
||||
|
||||
@@ -1216,6 +1216,108 @@ func migrateV9(db *gorm.DB, backend string) error {
|
||||
return backfillProxyRouteDomainCertificateFields(db)
|
||||
}
|
||||
|
||||
func ensureDefaultGitHubAuthSource(db *gorm.DB) error {
|
||||
if db == nil || !db.Migrator().HasTable(&AuthSource{}) || !db.Migrator().HasTable(&ExternalAccount{}) {
|
||||
return nil
|
||||
}
|
||||
|
||||
var githubUserCount int64
|
||||
if db.Migrator().HasColumn(&User{}, "github_id") {
|
||||
if err := db.Model(&User{}).Where("github_id <> ''").Count(&githubUserCount).Error; err != nil {
|
||||
return fmt.Errorf("count legacy github users failed: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
optionMap := map[string]string{}
|
||||
if db.Migrator().HasTable(&Option{}) {
|
||||
var options []Option
|
||||
if err := db.Find(&options).Error; err != nil {
|
||||
return fmt.Errorf("query options for github auth source migration failed: %w", err)
|
||||
}
|
||||
for _, option := range options {
|
||||
optionMap[option.Key] = option.Value
|
||||
}
|
||||
}
|
||||
|
||||
clientID := strings.TrimSpace(optionMap["GitHubClientId"])
|
||||
clientSecret := strings.TrimSpace(optionMap["GitHubClientSecret"])
|
||||
enabled := optionMap["GitHubOAuthEnabled"] == "true" && clientID != "" && clientSecret != ""
|
||||
if githubUserCount == 0 && clientID == "" && clientSecret == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
source := AuthSource{}
|
||||
err := db.Where("type = ? AND name = ?", AuthSourceTypeGitHub, "GitHub").First(&source).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
source = AuthSource{
|
||||
Name: "GitHub",
|
||||
Type: AuthSourceTypeGitHub,
|
||||
DisplayName: "GitHub",
|
||||
IsActive: enabled,
|
||||
ClientID: clientID,
|
||||
ClientSecret: clientSecret,
|
||||
Scopes: "user:email",
|
||||
}
|
||||
if err := db.Create(&source).Error; err != nil {
|
||||
return fmt.Errorf("create default github auth source failed: %w", err)
|
||||
}
|
||||
} else if err != nil {
|
||||
return fmt.Errorf("query default github auth source failed: %w", err)
|
||||
} else {
|
||||
updates := map[string]any{}
|
||||
if source.ClientID == "" && clientID != "" {
|
||||
updates["client_id"] = clientID
|
||||
}
|
||||
if source.ClientSecret == "" && clientSecret != "" {
|
||||
updates["client_secret"] = clientSecret
|
||||
}
|
||||
if source.Scopes == "" {
|
||||
updates["scopes"] = "user:email"
|
||||
}
|
||||
if enabled && !source.IsActive {
|
||||
updates["is_active"] = true
|
||||
}
|
||||
if len(updates) > 0 {
|
||||
if err := db.Model(&source).Updates(updates).Error; err != nil {
|
||||
return fmt.Errorf("update default github auth source failed: %w", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if githubUserCount == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
var users []User
|
||||
if err := db.Select("id", "github_id", "username", "email").Where("github_id <> ''").Find(&users).Error; err != nil {
|
||||
return fmt.Errorf("query legacy github users failed: %w", err)
|
||||
}
|
||||
for _, user := range users {
|
||||
account := ExternalAccount{
|
||||
AuthSourceID: source.ID,
|
||||
UserID: user.Id,
|
||||
ExternalID: user.GitHubId,
|
||||
ExternalUsername: user.GitHubId,
|
||||
Email: user.Email,
|
||||
}
|
||||
if err := db.Where(ExternalAccount{
|
||||
AuthSourceID: source.ID,
|
||||
ExternalID: user.GitHubId,
|
||||
}).FirstOrCreate(&account).Error; err != nil {
|
||||
return fmt.Errorf("migrate github external account for user %d failed: %w", user.Id, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// migrateV10 adds configurable auth sources and external account bindings.
|
||||
func migrateV10(db *gorm.DB, backend string) error {
|
||||
if err := applyCurrentSchema(db, backend); err != nil {
|
||||
return err
|
||||
}
|
||||
return ensureDefaultGitHubAuthSource(db)
|
||||
}
|
||||
|
||||
func validateDatabaseSchemaV9(db *gorm.DB, backend string) error {
|
||||
if err := validateDatabaseSchemaV8(db, backend); err != nil {
|
||||
return err
|
||||
@@ -1229,6 +1331,19 @@ func validateDatabaseSchemaV9(db *gorm.DB, backend string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateDatabaseSchemaV10(db *gorm.DB, backend string) error {
|
||||
if err := validateDatabaseSchemaV9(db, backend); err != nil {
|
||||
return err
|
||||
}
|
||||
if !db.Migrator().HasTable(&AuthSource{}) {
|
||||
return fmt.Errorf("table auth_sources is missing")
|
||||
}
|
||||
if !db.Migrator().HasTable(&ExternalAccount{}) {
|
||||
return fmt.Errorf("table external_accounts is missing")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func databaseSchemaMigrations() []databaseSchemaMigration {
|
||||
return []databaseSchemaMigration{
|
||||
{fromVersion: 1, toVersion: 2, migrate: migrateV2, validate: validateDatabaseSchemaV2},
|
||||
@@ -1239,6 +1354,7 @@ func databaseSchemaMigrations() []databaseSchemaMigration {
|
||||
{fromVersion: 6, toVersion: 7, migrate: migrateV7, validate: validateDatabaseSchemaV7},
|
||||
{fromVersion: 7, toVersion: 8, migrate: migrateV8, validate: validateDatabaseSchemaV8},
|
||||
{fromVersion: 8, toVersion: 9, migrate: migrateV9, validate: validateDatabaseSchemaV9},
|
||||
{fromVersion: 9, toVersion: 10, migrate: migrateV10, validate: validateDatabaseSchemaV10},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1321,7 +1437,10 @@ func initializeFreshDatabaseSchema(db *gorm.DB, backend string) error {
|
||||
if err := backfillProxyRouteDomainCertificateFields(db); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := validateDatabaseSchemaV9(db, backend); err != nil {
|
||||
if err := ensureDefaultGitHubAuthSource(db); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := validateDatabaseSchemaV10(db, backend); err != nil {
|
||||
return err
|
||||
}
|
||||
return saveDatabaseSchemaVersion(db, currentDatabaseSchemaVersion)
|
||||
|
||||
Reference in New Issue
Block a user