373 lines
11 KiB
Go
373 lines
11 KiB
Go
|
|
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)
|
||
|
|
}
|