2026-07-03 15:58:29 +08:00
|
|
|
package controller
|
|
|
|
|
|
|
|
|
|
import (
|
|
|
|
|
"crypto/rand"
|
|
|
|
|
"encoding/hex"
|
|
|
|
|
"fmt"
|
|
|
|
|
"net/http"
|
|
|
|
|
"strconv"
|
|
|
|
|
|
|
|
|
|
"seeyon-filesystem/dto"
|
|
|
|
|
"seeyon-filesystem/model"
|
|
|
|
|
"seeyon-filesystem/service"
|
|
|
|
|
"seeyon-filesystem/utils"
|
|
|
|
|
|
|
|
|
|
"github.com/gin-gonic/gin"
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
// SSOController SSO控制器
|
|
|
|
|
type SSOController struct {
|
|
|
|
|
ssoService *service.SSOService
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func NewSSOController() *SSOController {
|
|
|
|
|
return &SSOController{
|
|
|
|
|
ssoService: service.NewSSOService(),
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// ========== 公开接口 (无需登录) ==========
|
|
|
|
|
|
|
|
|
|
// GetProviders 获取启用的SSO提供商列表(登录页展示)
|
|
|
|
|
// GET /api/sso/providers
|
|
|
|
|
func (c *SSOController) GetProviders(ctx *gin.Context) {
|
|
|
|
|
providers, err := c.ssoService.ListActiveProviders()
|
|
|
|
|
if err != nil {
|
|
|
|
|
utils.ServerError(ctx, "获取SSO提供商失败")
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
vos := make([]dto.SSOProviderVO, 0, len(providers))
|
|
|
|
|
for _, p := range providers {
|
|
|
|
|
vos = append(vos, dto.SSOProviderVO{
|
|
|
|
|
ID: p.ID,
|
|
|
|
|
Name: p.Name,
|
|
|
|
|
DisplayName: p.DisplayName,
|
|
|
|
|
Icon: p.Icon,
|
|
|
|
|
Type: p.Type,
|
|
|
|
|
})
|
|
|
|
|
}
|
|
|
|
|
utils.Success(ctx, vos)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// GetAuthURL 获取SSO授权跳转URL
|
|
|
|
|
// GET /api/sso/auth/url?provider_id=1
|
|
|
|
|
func (c *SSOController) GetAuthURL(ctx *gin.Context) {
|
|
|
|
|
var req dto.SSOAuthURLRequest
|
|
|
|
|
if err := ctx.ShouldBindQuery(&req); err != nil {
|
|
|
|
|
utils.BadRequest(ctx, "参数错误: "+err.Error())
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
state, _ := generateState()
|
|
|
|
|
authURL, err := c.ssoService.GetAuthURL(req.ProviderID, state)
|
|
|
|
|
if err != nil {
|
|
|
|
|
utils.BadRequest(ctx, err.Error())
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
utils.Success(ctx, gin.H{
|
|
|
|
|
"auth_url": authURL,
|
|
|
|
|
"state": state,
|
|
|
|
|
})
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// Callback 处理SSO回调(重定向到前端)
|
|
|
|
|
// GET /api/sso/callback?provider_id=1&code=xxx&state=xxx
|
|
|
|
|
func (c *SSOController) Callback(ctx *gin.Context) {
|
|
|
|
|
var req dto.SSOCallbackRequest
|
|
|
|
|
if err := ctx.ShouldBindQuery(&req); err != nil {
|
|
|
|
|
utils.BadRequest(ctx, "参数错误: "+err.Error())
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
result, err := c.ssoService.HandleCallback(req.ProviderID, req.Code)
|
|
|
|
|
if err != nil {
|
|
|
|
|
// 回调失败, 重定向到登录页并携带错误信息
|
|
|
|
|
ctx.Redirect(http.StatusFound, "/login?sso_error="+err.Error())
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
ctx.Redirect(http.StatusFound, fmt.Sprintf("/login?sso=success&token=%s&username=%s", result.Token, result.Username))
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// CallbackJSON 处理SSO回调(JSON版本, 供前端AJAX调用)
|
|
|
|
|
// POST /api/sso/callback
|
|
|
|
|
func (c *SSOController) CallbackJSON(ctx *gin.Context) {
|
|
|
|
|
var req dto.SSOCallbackRequest
|
|
|
|
|
if err := ctx.ShouldBindJSON(&req); err != nil {
|
|
|
|
|
utils.BadRequest(ctx, "参数错误: "+err.Error())
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
result, err := c.ssoService.HandleCallback(req.ProviderID, req.Code)
|
|
|
|
|
if err != nil {
|
|
|
|
|
utils.Unauthorized(ctx, err.Error())
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
utils.Success(ctx, result)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// ========== 管理员接口 (需要登录 + admin权限) ==========
|
|
|
|
|
|
|
|
|
|
// AdminListProviders 管理员获取所有SSO提供商
|
|
|
|
|
// GET /api/admin/sso/providers
|
|
|
|
|
func (c *SSOController) AdminListProviders(ctx *gin.Context) {
|
|
|
|
|
providers, err := c.ssoService.ListProviders()
|
|
|
|
|
if err != nil {
|
|
|
|
|
utils.ServerError(ctx, "获取提供商列表失败")
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
utils.Success(ctx, providers)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// AdminGetProvider 获取单个SSO提供商详情
|
|
|
|
|
// GET /api/admin/sso/providers/:id
|
|
|
|
|
func (c *SSOController) AdminGetProvider(ctx *gin.Context) {
|
|
|
|
|
id, err := strconv.ParseUint(ctx.Param("id"), 10, 64)
|
|
|
|
|
if err != nil {
|
|
|
|
|
utils.BadRequest(ctx, "无效的提供商ID")
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
providers, _ := c.ssoService.ListProviders()
|
|
|
|
|
for _, p := range providers {
|
|
|
|
|
if p.ID == id {
|
|
|
|
|
utils.Success(ctx, p)
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
utils.NotFound(ctx, "提供商不存在")
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// AdminCreateProvider 管理员创建SSO提供商
|
|
|
|
|
// POST /api/admin/sso/providers
|
|
|
|
|
func (c *SSOController) AdminCreateProvider(ctx *gin.Context) {
|
|
|
|
|
var req dto.SSOProviderCreateRequest
|
|
|
|
|
if err := ctx.ShouldBindJSON(&req); err != nil {
|
|
|
|
|
utils.BadRequest(ctx, "参数错误: "+err.Error())
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
provider := &model.SSOProvider{
|
2026-07-18 10:52:32 +08:00
|
|
|
ID: utils.GenID(),
|
2026-07-03 15:58:29 +08:00
|
|
|
Name: req.Name,
|
|
|
|
|
Type: req.Type,
|
|
|
|
|
DisplayName: req.DisplayName,
|
|
|
|
|
Icon: req.Icon,
|
|
|
|
|
AutoCreate: req.AutoCreate,
|
|
|
|
|
DefaultRole: req.DefaultRole,
|
|
|
|
|
Status: 1,
|
|
|
|
|
Config: model.SSOConfig{
|
|
|
|
|
ClientID: req.Config.ClientID,
|
|
|
|
|
ClientSecret: req.Config.ClientSecret,
|
|
|
|
|
AuthURL: req.Config.AuthURL,
|
|
|
|
|
TokenURL: req.Config.TokenURL,
|
|
|
|
|
UserInfoURL: req.Config.UserInfoURL,
|
|
|
|
|
RedirectURI: req.Config.RedirectURI,
|
|
|
|
|
Scopes: req.Config.Scopes,
|
|
|
|
|
CasServerURL: req.Config.CasServerURL,
|
|
|
|
|
LDAPServer: req.Config.LDAPServer,
|
|
|
|
|
LDAPPort: req.Config.LDAPPort,
|
|
|
|
|
LDAPBaseDN: req.Config.LDAPBaseDN,
|
|
|
|
|
LDAPBindUser: req.Config.LDAPBindUser,
|
|
|
|
|
LDAPBindPass: req.Config.LDAPBindPass,
|
|
|
|
|
LDAPFilter: req.Config.LDAPFilter,
|
|
|
|
|
MappingUsername: req.Config.MappingUsername,
|
|
|
|
|
MappingEmail: req.Config.MappingEmail,
|
|
|
|
|
MappingNickname: req.Config.MappingNickname,
|
|
|
|
|
},
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
if err := c.ssoService.CreateProvider(provider); err != nil {
|
|
|
|
|
utils.ServerError(ctx, err.Error())
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
utils.SuccessWithMessage(ctx, "创建成功", nil)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// AdminUpdateProvider 管理员更新SSO提供商
|
|
|
|
|
// PUT /api/admin/sso/providers/:id
|
|
|
|
|
func (c *SSOController) AdminUpdateProvider(ctx *gin.Context) {
|
|
|
|
|
id, err := strconv.ParseUint(ctx.Param("id"), 10, 64)
|
|
|
|
|
if err != nil {
|
|
|
|
|
utils.BadRequest(ctx, "无效的提供商ID")
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// 获取现有提供商
|
|
|
|
|
providers, _ := c.ssoService.ListProviders()
|
|
|
|
|
var existing *model.SSOProvider
|
|
|
|
|
for i, p := range providers {
|
|
|
|
|
if p.ID == id {
|
|
|
|
|
existing = &providers[i]
|
|
|
|
|
break
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
if existing == nil {
|
|
|
|
|
utils.NotFound(ctx, "提供商不存在")
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
var req dto.SSOProviderCreateRequest
|
|
|
|
|
if err := ctx.ShouldBindJSON(&req); err != nil {
|
|
|
|
|
utils.BadRequest(ctx, "参数错误: "+err.Error())
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// 更新字段
|
|
|
|
|
if req.Name != "" { existing.Name = req.Name }
|
|
|
|
|
if req.DisplayName != "" { existing.DisplayName = req.DisplayName }
|
|
|
|
|
if req.Icon != "" { existing.Icon = req.Icon }
|
|
|
|
|
if req.DefaultRole != "" { existing.DefaultRole = req.DefaultRole }
|
|
|
|
|
if req.AutoCreate != 0 { existing.AutoCreate = req.AutoCreate }
|
|
|
|
|
|
|
|
|
|
// 更新配置
|
|
|
|
|
if req.Config.ClientID != "" { existing.Config.ClientID = req.Config.ClientID }
|
|
|
|
|
if req.Config.ClientSecret != "" { existing.Config.ClientSecret = req.Config.ClientSecret }
|
|
|
|
|
if req.Config.AuthURL != "" { existing.Config.AuthURL = req.Config.AuthURL }
|
|
|
|
|
if req.Config.TokenURL != "" { existing.Config.TokenURL = req.Config.TokenURL }
|
|
|
|
|
if req.Config.UserInfoURL != "" { existing.Config.UserInfoURL = req.Config.UserInfoURL }
|
|
|
|
|
if req.Config.RedirectURI != "" { existing.Config.RedirectURI = req.Config.RedirectURI }
|
|
|
|
|
if req.Config.Scopes != "" { existing.Config.Scopes = req.Config.Scopes }
|
|
|
|
|
if req.Config.CasServerURL != "" { existing.Config.CasServerURL = req.Config.CasServerURL }
|
|
|
|
|
|
|
|
|
|
if err := c.ssoService.UpdateProvider(existing); err != nil {
|
|
|
|
|
utils.ServerError(ctx, err.Error())
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
utils.SuccessWithMessage(ctx, "更新成功", nil)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// AdminDeleteProvider 管理员删除SSO提供商
|
|
|
|
|
// DELETE /api/admin/sso/providers/:id
|
|
|
|
|
func (c *SSOController) AdminDeleteProvider(ctx *gin.Context) {
|
|
|
|
|
id, err := strconv.ParseUint(ctx.Param("id"), 10, 64)
|
|
|
|
|
if err != nil {
|
|
|
|
|
utils.BadRequest(ctx, "无效的提供商ID")
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
if err := c.ssoService.DeleteProvider(id); err != nil {
|
|
|
|
|
utils.ServerError(ctx, err.Error())
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
utils.SuccessWithMessage(ctx, "删除成功", nil)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// AdminToggleProvider 启用/禁用SSO提供商
|
|
|
|
|
// PUT /api/admin/sso/providers/:id/toggle
|
|
|
|
|
func (c *SSOController) AdminToggleProvider(ctx *gin.Context) {
|
|
|
|
|
id, err := strconv.ParseUint(ctx.Param("id"), 10, 64)
|
|
|
|
|
if err != nil {
|
|
|
|
|
utils.BadRequest(ctx, "无效的提供商ID")
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
var req struct {
|
|
|
|
|
Status int8 `json:"status"`
|
|
|
|
|
}
|
|
|
|
|
if err := ctx.ShouldBindJSON(&req); err != nil {
|
|
|
|
|
utils.BadRequest(ctx, "参数错误")
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
providers, _ := c.ssoService.ListProviders()
|
|
|
|
|
for i, p := range providers {
|
|
|
|
|
if p.ID == id {
|
|
|
|
|
providers[i].Status = req.Status
|
|
|
|
|
if err := c.ssoService.UpdateProvider(&providers[i]); err != nil {
|
|
|
|
|
utils.ServerError(ctx, err.Error())
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
utils.SuccessWithMessage(ctx, "操作成功", nil)
|
|
|
|
|
return
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
utils.NotFound(ctx, "提供商不存在")
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// generateState 生成随机state参数(防CSRF)
|
|
|
|
|
func generateState() (string, error) {
|
|
|
|
|
bytes := make([]byte, 16)
|
|
|
|
|
if _, err := rand.Read(bytes); err != nil {
|
|
|
|
|
return "", err
|
|
|
|
|
}
|
|
|
|
|
return hex.EncodeToString(bytes), nil
|
|
|
|
|
}
|