Files
SeeyonFileSystem/server/repository/sso_repo.go
2026-07-03 15:58:29 +08:00

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
}