95 lines
2.7 KiB
Go
95 lines
2.7 KiB
Go
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
|
|
}
|