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{ ID: utils.GenID(), 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 }