初始化
This commit is contained in:
94
server/repository/sso_repo.go
Normal file
94
server/repository/sso_repo.go
Normal file
@@ -0,0 +1,94 @@
|
||||
package repository
|
||||
|
||||
import (
|
||||
"seeyon-filesystem/config"
|
||||
"seeyon-filesystem/model"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// SSORepository SSO数据访问
|
||||
type SSORepository struct{}
|
||||
|
||||
func NewSSORepository() *SSORepository {
|
||||
return &SSORepository{}
|
||||
}
|
||||
|
||||
// ========== SSO 提供商 ==========
|
||||
|
||||
// CreateProvider 创建SSO提供商
|
||||
func (r *SSORepository) CreateProvider(provider *model.SSOProvider) error {
|
||||
return config.DB.Create(provider).Error
|
||||
}
|
||||
|
||||
// FindProviderByID 根据ID查找提供商
|
||||
func (r *SSORepository) FindProviderByID(id uint64) (*model.SSOProvider, error) {
|
||||
var provider model.SSOProvider
|
||||
err := config.DB.First(&provider, id).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &provider, nil
|
||||
}
|
||||
|
||||
// ListProviders 查询所有启用的提供商
|
||||
func (r *SSORepository) ListProviders() ([]model.SSOProvider, error) {
|
||||
var providers []model.SSOProvider
|
||||
err := config.DB.Where("status = 1").Order("sort_order ASC, id ASC").Find(&providers).Error
|
||||
return providers, err
|
||||
}
|
||||
|
||||
// ListAllProviders 查询所有提供商(含禁用)
|
||||
func (r *SSORepository) ListAllProviders() ([]model.SSOProvider, error) {
|
||||
var providers []model.SSOProvider
|
||||
err := config.DB.Order("sort_order ASC, id ASC").Find(&providers).Error
|
||||
return providers, err
|
||||
}
|
||||
|
||||
// UpdateProvider 更新提供商
|
||||
func (r *SSORepository) UpdateProvider(provider *model.SSOProvider) error {
|
||||
return config.DB.Save(provider).Error
|
||||
}
|
||||
|
||||
// DeleteProvider 删除提供商
|
||||
func (r *SSORepository) DeleteProvider(id uint64) error {
|
||||
return config.DB.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Where("provider_id = ?", id).Delete(&model.SSOSession{}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Delete(&model.SSOProvider{}, id).Error
|
||||
})
|
||||
}
|
||||
|
||||
// ========== SSO 会话 ==========
|
||||
|
||||
// CreateSession 创建SSO会话
|
||||
func (r *SSORepository) CreateSession(session *model.SSOSession) error {
|
||||
return config.DB.Create(session).Error
|
||||
}
|
||||
|
||||
// FindSessionByExternalID 根据外部用户ID查找会话
|
||||
func (r *SSORepository) FindSessionByExternalID(providerID uint64, externalID string) (*model.SSOSession, error) {
|
||||
var session model.SSOSession
|
||||
err := config.DB.Where("provider_id = ? AND external_id = ?", providerID, externalID).
|
||||
First(&session).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &session, nil
|
||||
}
|
||||
|
||||
// FindSessionByUserID 根据本地用户ID查找SSO会话
|
||||
func (r *SSORepository) FindSessionByUserID(userID uint64) (*model.SSOSession, error) {
|
||||
var session model.SSOSession
|
||||
err := config.DB.Where("user_id = ?", userID).Order("last_login_at DESC").First(&session).Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &session, nil
|
||||
}
|
||||
|
||||
// UpdateSession 更新会话
|
||||
func (r *SSORepository) UpdateSession(session *model.SSOSession) error {
|
||||
return config.DB.Save(session).Error
|
||||
}
|
||||
Reference in New Issue
Block a user