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) }