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 }