Files
SeeyonFileSystem/server/service/sso_service.go

373 lines
11 KiB
Go
Raw Normal View History

2026-07-03 15:58:29 +08:00
package service
import (
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/url"
"strings"
"time"
"seeyon-filesystem/config"
"seeyon-filesystem/model"
"seeyon-filesystem/repository"
"seeyon-filesystem/utils"
)
// SSOService SSO单点登录服务
// 支持 OAuth2.0 / OIDC 协议
type SSOService struct {
ssoRepo *repository.SSORepository
userRepo *repository.UserRepository
roleRepo *repository.RoleRepository
userRoleRepo *repository.UserRoleRepository
}
// NewSSOService 创建SSO服务实例
func NewSSOService() *SSOService {
return &SSOService{
ssoRepo: repository.NewSSORepository(),
userRepo: repository.NewUserRepository(),
roleRepo: repository.NewRoleRepository(),
userRoleRepo: repository.NewUserRoleRepository(),
}
}
// GetAuthURL 获取SSO授权跳转URL
// 用户点击SSO登录按钮时调用, 返回第三方授权页面地址
func (s *SSOService) GetAuthURL(providerID uint64, state string) (string, error) {
provider, err := s.ssoRepo.FindProviderByID(providerID)
if err != nil {
return "", errors.New("SSO提供商不存在")
}
if provider.Status != 1 {
return "", errors.New("SSO提供商已禁用")
}
cfg := provider.Config
switch provider.Type {
case "oauth2", "oidc":
// 构建OAuth2授权URL
params := url.Values{}
params.Set("client_id", cfg.ClientID)
params.Set("redirect_uri", cfg.RedirectURI)
params.Set("response_type", "code")
params.Set("scope", cfg.Scopes)
params.Set("state", state)
return cfg.AuthURL + "?" + params.Encode(), nil
case "cas":
// CAS协议: 重定向到CAS登录页
serviceURL := url.QueryEscape(cfg.RedirectURI)
return cfg.CasServerURL + "/login?service=" + serviceURL, nil
default:
return "", fmt.Errorf("不支持的SSO类型: %s", provider.Type)
}
}
// HandleCallback 处理SSO回调
// 用户在第三方授权后回调此接口, 用授权码换取用户信息并完成登录
func (s *SSOService) HandleCallback(providerID uint64, code string) (*LoginResult, error) {
provider, err := s.ssoRepo.FindProviderByID(providerID)
if err != nil {
return nil, errors.New("SSO提供商不存在")
}
// 1. 用授权码换取Token和用户信息
ssoUser, accessToken, refreshToken, expiresAt, err := s.exchangeToken(provider, code)
if err != nil {
return nil, fmt.Errorf("获取SSO用户信息失败: %w", err)
}
// 2. 查找或创建本地用户
user, err := s.findOrCreateUser(provider, ssoUser)
if err != nil {
return nil, fmt.Errorf("创建本地用户失败: %w", err)
}
// 3. 检查用户状态
if user.Status != 1 {
return nil, errors.New("账号已被禁用")
}
// 4. 保存/更新SSO会话
s.saveSession(user.ID, providerID, ssoUser["id"], accessToken, refreshToken, expiresAt)
// 5. 生成本地JWT Token
token, err := utils.GenerateToken(
user.ID,
user.Username,
user.Role,
config.AppConfig.JWT.Secret,
config.AppConfig.JWT.Expiration,
)
if err != nil {
return nil, errors.New("生成Token失败")
}
return &LoginResult{
Token: token,
Username: user.Username,
Nickname: user.Nickname,
Role: user.Role,
}, nil
}
// LoginResult SSO登录结果
type LoginResult struct {
Token string `json:"token"`
Username string `json:"username"`
Nickname string `json:"nickname"`
Role string `json:"role"`
}
// ssoUser SSO用户信息
type ssoUser map[string]string
// exchangeToken 用授权码换取Token和用户信息
func (s *SSOService) exchangeToken(provider *model.SSOProvider, code string) (ssoUser, string, string, time.Time, error) {
cfg := provider.Config
switch provider.Type {
case "oauth2", "oidc":
return s.exchangeOAuth2Token(cfg, code)
case "cas":
return s.exchangeCASTicket(cfg, code)
default:
return nil, "", "", time.Time{}, fmt.Errorf("不支持的SSO类型: %s", provider.Type)
}
}
// exchangeOAuth2Token OAuth2.0 授权码换取Token
func (s *SSOService) exchangeOAuth2Token(cfg model.SSOConfig, code string) (ssoUser, string, string, time.Time, error) {
// 1. 用code换取access_token
data := url.Values{}
data.Set("grant_type", "authorization_code")
data.Set("code", code)
data.Set("client_id", cfg.ClientID)
data.Set("client_secret", cfg.ClientSecret)
data.Set("redirect_uri", cfg.RedirectURI)
resp, err := http.Post(cfg.TokenURL, "application/x-www-form-urlencoded", strings.NewReader(data.Encode()))
if err != nil {
return nil, "", "", time.Time{}, fmt.Errorf("请求Token失败: %w", err)
}
defer resp.Body.Close()
body, _ := io.ReadAll(resp.Body)
if resp.StatusCode != 200 {
return nil, "", "", time.Time{}, fmt.Errorf("获取Token失败: HTTP %d, %s", resp.StatusCode, string(body))
}
var tokenResp struct {
AccessToken string `json:"access_token"`
RefreshToken string `json:"refresh_token"`
ExpiresIn int `json:"expires_in"`
}
if err := json.Unmarshal(body, &tokenResp); err != nil {
return nil, "", "", time.Time{}, fmt.Errorf("解析Token响应失败: %w", err)
}
// 2. 用access_token获取用户信息
userInfo, err := s.fetchOAuth2UserInfo(cfg.UserInfoURL, tokenResp.AccessToken)
if err != nil {
return nil, "", "", time.Time{}, err
}
expiresAt := time.Now().Add(time.Duration(tokenResp.ExpiresIn) * time.Second)
return userInfo, tokenResp.AccessToken, tokenResp.RefreshToken, expiresAt, nil
}
// fetchOAuth2UserInfo 获取OAuth2用户信息
func (s *SSOService) fetchOAuth2UserInfo(userInfoURL, accessToken string) (ssoUser, error) {
req, _ := http.NewRequest("GET", userInfoURL, nil)
req.Header.Set("Authorization", "Bearer "+accessToken)
resp, err := http.DefaultClient.Do(req)
if err != nil {
return nil, fmt.Errorf("请求用户信息失败: %w", err)
}
defer resp.Body.Close()
body, _ := io.ReadAll(resp.Body)
if resp.StatusCode != 200 {
return nil, fmt.Errorf("获取用户信息失败: HTTP %d", resp.StatusCode)
}
var raw map[string]interface{}
if err := json.Unmarshal(body, &raw); err != nil {
return nil, fmt.Errorf("解析用户信息失败: %w", err)
}
// 将嵌套结构展平为 map[string]string
user := make(ssoUser)
for k, v := range raw {
user[k] = fmt.Sprintf("%v", v)
}
return user, nil
}
// exchangeCASTicket CAS协议票据验证
func (s *SSOService) exchangeCASTicket(cfg model.SSOConfig, ticket string) (ssoUser, string, string, time.Time, error) {
// CAS票据验证
validateURL := fmt.Sprintf("%s/serviceValidate?ticket=%s&service=%s&format=JSON",
cfg.CasServerURL, ticket, url.QueryEscape(cfg.RedirectURI))
resp, err := http.Get(validateURL)
if err != nil {
return nil, "", "", time.Time{}, fmt.Errorf("CAS票据验证失败: %w", err)
}
defer resp.Body.Close()
body, _ := io.ReadAll(resp.Body)
var casResp struct {
ServiceResponse struct {
AuthenticationSuccess struct {
User string `json:"user"`
Attributes map[string]interface{} `json:"attributes"`
} `json:"authenticationSuccess"`
AuthenticationFailure *struct {
Code string `json:"code"`
} `json:"authenticationFailure,omitempty"`
} `json:"serviceResponse"`
}
if err := json.Unmarshal(body, &casResp); err != nil {
return nil, "", "", time.Time{}, fmt.Errorf("解析CAS响应失败: %w", err)
}
if casResp.ServiceResponse.AuthenticationFailure != nil {
return nil, "", "", time.Time{}, fmt.Errorf("CAS认证失败: %s", casResp.ServiceResponse.AuthenticationFailure.Code)
}
success := casResp.ServiceResponse.AuthenticationSuccess
user := make(ssoUser)
user["id"] = success.User
user["username"] = success.User
// 提取属性
if email, ok := success.Attributes["email"]; ok {
user["email"] = fmt.Sprintf("%v", email)
}
if name, ok := success.Attributes["name"]; ok {
user["nickname"] = fmt.Sprintf("%v", name)
}
return user, ticket, "", time.Now().Add(24 * time.Hour), nil
}
// findOrCreateUser 查找或创建本地用户
func (s *SSOService) findOrCreateUser(provider *model.SSOProvider, ssoUserData ssoUser) (*model.User, error) {
externalID := ssoUserData["id"]
if externalID == "" {
externalID = ssoUserData["username"]
}
// 1. 通过SSO会话查找已有用户
session, _ := s.ssoRepo.FindSessionByExternalID(provider.ID, externalID)
if session != nil {
return s.userRepo.FindByID(session.UserID)
}
// 2. 通过用户名查找已有用户
username := ssoUserData["username"]
if username == "" {
username = externalID
}
user, _ := s.userRepo.FindByUsername(username)
if user != nil {
return user, nil
}
// 3. 自动创建用户
if provider.AutoCreate != 1 {
return nil, errors.New("用户不存在且未启用自动创建")
}
// 构建用户名: 前缀 + 外部ID, 避免冲突
localUsername := fmt.Sprintf("sso_%s_%s", provider.Type, username)
newUser := &model.User{
Username: localUsername,
Password: "_SSO_NO_PASSWORD_", // SSO用户无本地密码
Email: ssoUserData["email"],
Nickname: ssoUserData["nickname"],
Role: provider.DefaultRole,
Status: 1,
}
if newUser.Nickname == "" {
newUser.Nickname = username
}
if err := s.userRepo.Create(newUser); err != nil {
return nil, fmt.Errorf("创建用户失败: %w", err)
}
// 分配默认角色
role, _ := s.roleRepo.FindByName(provider.DefaultRole)
if role != nil {
s.userRoleRepo.SetUserRoles(newUser.ID, []uint64{role.ID})
}
return newUser, nil
}
// saveSession 保存SSO会话
func (s *SSOService) saveSession(userID, providerID uint64, externalID, accessToken, refreshToken string, expiresAt time.Time) {
session := &model.SSOSession{
UserID: userID,
ProviderID: providerID,
ExternalID: externalID,
AccessToken: accessToken,
RefreshToken: refreshToken,
ExpiresAt: expiresAt,
LastLoginAt: time.Now(),
}
existing, _ := s.ssoRepo.FindSessionByExternalID(providerID, externalID)
if existing != nil {
existing.AccessToken = accessToken
existing.RefreshToken = refreshToken
existing.ExpiresAt = expiresAt
existing.LastLoginAt = time.Now()
s.ssoRepo.UpdateSession(existing)
} else {
s.ssoRepo.CreateSession(session)
}
}
// ========== 管理接口 ==========
// CreateProvider 创建SSO提供商
func (s *SSOService) CreateProvider(provider *model.SSOProvider) error {
return s.ssoRepo.CreateProvider(provider)
}
// ListProviders 获取所有提供商
func (s *SSOService) ListProviders() ([]model.SSOProvider, error) {
return s.ssoRepo.ListAllProviders()
}
// ListActiveProviders 获取启用的提供商(登录页展示)
func (s *SSOService) ListActiveProviders() ([]model.SSOProvider, error) {
return s.ssoRepo.ListProviders()
}
// UpdateProvider 更新提供商
func (s *SSOService) UpdateProvider(provider *model.SSOProvider) error {
return s.ssoRepo.UpdateProvider(provider)
}
// DeleteProvider 删除提供商
func (s *SSOService) DeleteProvider(id uint64) error {
return s.ssoRepo.DeleteProvider(id)
}