refactor: 移除 SAML/LDAP/MFA 登录方式,仅保留本地账号密码和 OIDC
后端移除: - SAML: auth/saml.go, handler/saml_handler.go, account_saml_settings_handler.go, model/account_saml_settings.go, model/saml_idp_config.go, repo/*.go - LDAP: auth/ldap.go, handler/ldap_handler.go, model/account_ldap_settings.go, repo/account_ldap_settings_repo.go - MFA: auth/mfa.go, handler/mfa_handler.go - auth_service: 移除 mfaService 依赖、MFARequired 字段、LoginWithMFA 方法 - auth_handler: 移除 LoginMFA handler、MFA 分支逻辑 - bootstrap: 移除 SAML/LDAP/MFA service 初始化和 handler 注册 - sso_middleware: 精简为仅支持 OIDC provider - router: 移除 SAML/LDAP/MFA 路由注册 - config: 移除 SAMLConfig/LDAPConfig struct 和 defaults 前端移除: - v3/login: 移除 MFA 验证流程和 SAML 登录入口 - v3/api/auth: 移除 MFA 响应处理 - v3/routes: 移除 SSO login 路由 - dashboard: 移除 MFA 设置页面、SAML 安全设置页面 - i18n: 移除 mfa.json - featureFlags: 移除 SAML feature flag .env.example / .env: 移除 SAML/LDAP 配置段
This commit is contained in:
@@ -1,462 +0,0 @@
|
||||
package v1
|
||||
|
||||
// Reference: M13 §2 — Account-scoped SAML config admin API
|
||||
// CRUD endpoints for enterprise administrators to configure SAML SSO for their accounts.
|
||||
// Pattern follows Chatwoot AccountSamlSettings API (super_admin scoped).
|
||||
|
||||
import (
|
||||
"crypto/sha1"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/gochat/gochat/internal/model"
|
||||
"github.com/gochat/gochat/internal/repository"
|
||||
applogger "github.com/gochat/gochat/pkg/logger"
|
||||
"github.com/gochat/gochat/pkg/response"
|
||||
)
|
||||
|
||||
// AccountSamlSettingsHandler handles account-scoped SAML configuration endpoints.
|
||||
// Only accessible to account administrators (role: administrator or super_admin).
|
||||
type AccountSamlSettingsHandler struct {
|
||||
repo *repository.AccountSamlSettingsRepo
|
||||
}
|
||||
|
||||
// NewAccountSamlSettingsHandler creates a new AccountSamlSettings handler.
|
||||
func NewAccountSamlSettingsHandler(repo *repository.AccountSamlSettingsRepo) *AccountSamlSettingsHandler {
|
||||
return &AccountSamlSettingsHandler{repo: repo}
|
||||
}
|
||||
|
||||
// Get retrieves SAML settings for an account.
|
||||
// GET /api/v1/accounts/:account_id/saml_settings
|
||||
func (h *AccountSamlSettingsHandler) Get(c *gin.Context) {
|
||||
accountID, err := strconv.ParseUint(c.Param("account_id"), 10, 64)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrBadRequest,
|
||||
Message: "Invalid account ID",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
settings, err := h.repo.GetByAccount(uint(accountID))
|
||||
if err != nil {
|
||||
applogger.L().Errorf("Get SAML settings for account %d: %v", accountID, err)
|
||||
response.AbortWithStatusError(c, http.StatusInternalServerError, response.ErrInternal, "failed to get SAML settings")
|
||||
return
|
||||
}
|
||||
|
||||
if settings == nil {
|
||||
c.JSON(http.StatusNotFound, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrNotFound,
|
||||
Message: "SAML settings not found for this account",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, serializeSamlSettings(settings))
|
||||
}
|
||||
|
||||
// Create creates SAML settings for an account.
|
||||
// POST /api/v1/accounts/:account_id/saml_settings
|
||||
// Body: JSON with IdP configuration fields.
|
||||
func (h *AccountSamlSettingsHandler) Create(c *gin.Context) {
|
||||
accountID, err := strconv.ParseUint(c.Param("account_id"), 10, 64)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrBadRequest,
|
||||
Message: "Invalid account ID",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
req, err := bindCreateSamlSettingsRequest(c)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrBadRequest,
|
||||
Message: "Invalid request body",
|
||||
Detail: err.Error(),
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Check if settings already exist for this account
|
||||
existing, err := h.repo.GetByAccount(uint(accountID))
|
||||
if err != nil {
|
||||
applogger.L().Errorf("Check existing SAML settings for account %d: %v", accountID, err)
|
||||
response.AbortWithStatusError(c, http.StatusInternalServerError, response.ErrInternal, "failed to check existing settings")
|
||||
return
|
||||
}
|
||||
if existing != nil {
|
||||
c.JSON(http.StatusConflict, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrBadRequest,
|
||||
Message: "SAML settings already exist for this account",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
settings := req.ToModel(uint(accountID))
|
||||
if err := h.repo.Create(&settings); err != nil {
|
||||
applogger.L().Errorf("Create SAML settings for account %d: %v", accountID, err)
|
||||
response.AbortWithStatusError(c, http.StatusInternalServerError, response.ErrInternal, "failed to create SAML settings")
|
||||
return
|
||||
}
|
||||
|
||||
applogger.L().Infof("SAML settings created for account %d", accountID)
|
||||
c.JSON(http.StatusCreated, serializeSamlSettings(&settings))
|
||||
}
|
||||
|
||||
// Update updates SAML settings for an account.
|
||||
// PUT /api/v1/accounts/:account_id/saml_settings
|
||||
// Body: JSON with fields to update (partial update supported).
|
||||
func (h *AccountSamlSettingsHandler) Update(c *gin.Context) {
|
||||
accountID, err := strconv.ParseUint(c.Param("account_id"), 10, 64)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrBadRequest,
|
||||
Message: "Invalid account ID",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
req, err := bindUpdateSamlSettingsRequest(c)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrBadRequest,
|
||||
Message: "Invalid request body",
|
||||
Detail: err.Error(),
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Verify settings exist
|
||||
existing, err := h.repo.GetByAccount(uint(accountID))
|
||||
if err != nil {
|
||||
applogger.L().Errorf("Get SAML settings for account %d: %v", accountID, err)
|
||||
response.AbortWithStatusError(c, http.StatusInternalServerError, response.ErrInternal, "failed to get SAML settings")
|
||||
return
|
||||
}
|
||||
if existing == nil {
|
||||
c.JSON(http.StatusNotFound, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrNotFound,
|
||||
Message: "SAML settings not found for this account",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Build updates map
|
||||
updates := req.ToUpdatesMap()
|
||||
if len(updates) == 0 {
|
||||
c.JSON(http.StatusBadRequest, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrBadRequest,
|
||||
Message: "No fields to update",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.repo.UpdateFields(uint(accountID), updates); err != nil {
|
||||
applogger.L().Errorf("Update SAML settings for account %d: %v", accountID, err)
|
||||
response.AbortWithStatusError(c, http.StatusInternalServerError, response.ErrInternal, "failed to update SAML settings")
|
||||
return
|
||||
}
|
||||
|
||||
// Return updated settings
|
||||
settings, _ := h.repo.GetByAccount(uint(accountID))
|
||||
applogger.L().Infof("SAML settings updated for account %d", accountID)
|
||||
c.JSON(http.StatusOK, serializeSamlSettings(settings))
|
||||
}
|
||||
|
||||
// Delete removes SAML settings for an account.
|
||||
// DELETE /api/v1/accounts/:account_id/saml_settings
|
||||
func (h *AccountSamlSettingsHandler) Delete(c *gin.Context) {
|
||||
accountID, err := strconv.ParseUint(c.Param("account_id"), 10, 64)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrBadRequest,
|
||||
Message: "Invalid account ID",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.repo.Delete(uint(accountID)); err != nil {
|
||||
applogger.L().Errorf("Delete SAML settings for account %d: %v", accountID, err)
|
||||
response.AbortWithStatusError(c, http.StatusInternalServerError, response.ErrInternal, "failed to delete SAML settings")
|
||||
return
|
||||
}
|
||||
|
||||
applogger.L().Infof("SAML settings deleted for account %d", accountID)
|
||||
c.JSON(http.StatusOK, response.APIResponse{
|
||||
Success: true,
|
||||
Data: map[string]string{
|
||||
"message": "SAML settings deleted",
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// ToggleActive enables or disables SAML SSO for an account.
|
||||
// POST /api/v1/accounts/:account_id/saml_settings/toggle_active
|
||||
// Body: { "active": true/false }
|
||||
func (h *AccountSamlSettingsHandler) ToggleActive(c *gin.Context) {
|
||||
accountID, err := strconv.ParseUint(c.Param("account_id"), 10, 64)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrBadRequest,
|
||||
Message: "Invalid account ID",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
var req ToggleActiveRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrBadRequest,
|
||||
Message: "Invalid request body",
|
||||
Detail: err.Error(),
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.repo.SetActive(uint(accountID), req.Active); err != nil {
|
||||
applogger.L().Errorf("Toggle SAML active for account %d: %v", accountID, err)
|
||||
response.AbortWithStatusError(c, http.StatusInternalServerError, response.ErrInternal, "failed to toggle SAML active status")
|
||||
return
|
||||
}
|
||||
|
||||
status := "disabled"
|
||||
if req.Active {
|
||||
status = "enabled"
|
||||
}
|
||||
applogger.L().Infof("SAML SSO %s for account %d", status, accountID)
|
||||
c.JSON(http.StatusOK, response.APIResponse{
|
||||
Success: true,
|
||||
Data: map[string]interface{}{
|
||||
"account_id": accountID,
|
||||
"active": req.Active,
|
||||
"status": status,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// --- Request/Response types ---
|
||||
|
||||
// CreateSamlSettingsRequest is the request body for creating SAML settings.
|
||||
type CreateSamlSettingsRequest struct {
|
||||
SsoURL string `json:"sso_url"`
|
||||
Certificate string `json:"certificate"`
|
||||
IdpEntityID string `json:"idp_entity_id" binding:"required"`
|
||||
IdpSsoTargetURL string `json:"idp_sso_target_url" binding:"required"`
|
||||
IdpSloTargetURL string `json:"idp_slo_target_url"`
|
||||
IdpCertificate string `json:"idp_certificate" binding:"required"`
|
||||
SpEntityID string `json:"sp_entity_id"`
|
||||
SpX509Certificate string `json:"sp_x509_certificate"`
|
||||
SpPrivateKey string `json:"sp_private_key"`
|
||||
RoleMappings json.RawMessage `json:"role_mappings"`
|
||||
Active bool `json:"active"`
|
||||
}
|
||||
|
||||
// ToModel converts a CreateSamlSettingsRequest to an AccountSamlSettings model.
|
||||
func (req *CreateSamlSettingsRequest) ToModel(accountID uint) model.AccountSamlSettings {
|
||||
req.normalizeChatwootAliases()
|
||||
return model.AccountSamlSettings{
|
||||
AccountID: accountID,
|
||||
IdpEntityID: req.IdpEntityID,
|
||||
IdpSsoTargetURL: req.IdpSsoTargetURL,
|
||||
IdpSloTargetURL: req.IdpSloTargetURL,
|
||||
IdpCertificate: req.IdpCertificate,
|
||||
SpEntityID: req.SpEntityID,
|
||||
SpX509Certificate: req.SpX509Certificate,
|
||||
SpPrivateKey: req.SpPrivateKey,
|
||||
RoleMappings: req.RoleMappings,
|
||||
Active: req.Active,
|
||||
}
|
||||
}
|
||||
|
||||
func (req *CreateSamlSettingsRequest) normalizeChatwootAliases() {
|
||||
if req.IdpSsoTargetURL == "" {
|
||||
req.IdpSsoTargetURL = req.SsoURL
|
||||
}
|
||||
if req.IdpCertificate == "" {
|
||||
req.IdpCertificate = req.Certificate
|
||||
}
|
||||
}
|
||||
|
||||
// UpdateSamlSettingsRequest is the request body for updating SAML settings.
|
||||
// All fields are optional — only non-nil/non-zero fields will be updated.
|
||||
type UpdateSamlSettingsRequest struct {
|
||||
SsoURL *string `json:"sso_url"`
|
||||
Certificate *string `json:"certificate"`
|
||||
IdpEntityID *string `json:"idp_entity_id"`
|
||||
IdpSsoTargetURL *string `json:"idp_sso_target_url"`
|
||||
IdpSloTargetURL *string `json:"idp_slo_target_url"`
|
||||
IdpCertificate *string `json:"idp_certificate"`
|
||||
SpEntityID *string `json:"sp_entity_id"`
|
||||
SpX509Certificate *string `json:"sp_x509_certificate"`
|
||||
SpPrivateKey *string `json:"sp_private_key"`
|
||||
RoleMappings json.RawMessage `json:"role_mappings"`
|
||||
Active *bool `json:"active"`
|
||||
}
|
||||
|
||||
// ToUpdatesMap converts an UpdateSamlSettingsRequest to a map of fields to update.
|
||||
func (req *UpdateSamlSettingsRequest) ToUpdatesMap() map[string]interface{} {
|
||||
updates := make(map[string]interface{})
|
||||
if req.IdpSsoTargetURL == nil {
|
||||
req.IdpSsoTargetURL = req.SsoURL
|
||||
}
|
||||
if req.IdpCertificate == nil {
|
||||
req.IdpCertificate = req.Certificate
|
||||
}
|
||||
if req.IdpEntityID != nil {
|
||||
updates["idp_entity_id"] = *req.IdpEntityID
|
||||
}
|
||||
if req.IdpSsoTargetURL != nil {
|
||||
updates["idp_sso_target_url"] = *req.IdpSsoTargetURL
|
||||
}
|
||||
if req.IdpSloTargetURL != nil {
|
||||
updates["idp_slo_target_url"] = *req.IdpSloTargetURL
|
||||
}
|
||||
if req.IdpCertificate != nil {
|
||||
updates["idp_certificate"] = *req.IdpCertificate
|
||||
}
|
||||
if req.SpEntityID != nil {
|
||||
updates["sp_entity_id"] = *req.SpEntityID
|
||||
}
|
||||
if req.SpX509Certificate != nil {
|
||||
updates["sp_x509_certificate"] = *req.SpX509Certificate
|
||||
}
|
||||
if req.SpPrivateKey != nil {
|
||||
updates["sp_private_key"] = *req.SpPrivateKey
|
||||
}
|
||||
if req.RoleMappings != nil {
|
||||
updates["role_mappings"] = req.RoleMappings
|
||||
}
|
||||
if req.Active != nil {
|
||||
updates["active"] = *req.Active
|
||||
}
|
||||
return updates
|
||||
}
|
||||
|
||||
// ToggleActiveRequest toggles the active status of SAML settings.
|
||||
type ToggleActiveRequest struct {
|
||||
Active bool `json:"active"`
|
||||
}
|
||||
|
||||
// RegisterAccountSamlSettingsRoutes maps account-scoped SAML config admin routes.
|
||||
// Only accessible to account administrators (enforced by AccountScope middleware in router).
|
||||
// Reference: Chatwoot AccountSamlSettings API — enterprise SSO configuration
|
||||
func RegisterAccountSamlSettingsRoutes(g *gin.RouterGroup, h *AccountSamlSettingsHandler) {
|
||||
g.GET("", h.Get) // Chatwoot frontend: fetch account SAML config
|
||||
g.POST("", h.Create) // Chatwoot frontend: create account SAML config
|
||||
g.PUT("", h.Update) // Chatwoot frontend: update account SAML config
|
||||
g.DELETE("", h.Delete) // Chatwoot frontend: delete account SAML config
|
||||
g.POST("/", h.Create) // Backward compatibility for trailing slash clients
|
||||
g.GET("/:id", h.Get) // Backward compatibility for id-based callers
|
||||
g.PUT("/:id", h.Update) // Backward compatibility for id-based callers
|
||||
g.DELETE("/:id", h.Delete) // Backward compatibility for id-based callers
|
||||
g.POST("/:id/toggle_active", h.ToggleActive) // Enable/disable SAML for a config
|
||||
}
|
||||
|
||||
func bindCreateSamlSettingsRequest(c *gin.Context) (CreateSamlSettingsRequest, error) {
|
||||
req := CreateSamlSettingsRequest{}
|
||||
if err := bindSamlSettingsPayload(c, &req); err != nil {
|
||||
return req, err
|
||||
}
|
||||
req.normalizeChatwootAliases()
|
||||
if req.IdpEntityID == "" || req.IdpSsoTargetURL == "" || req.IdpCertificate == "" {
|
||||
return req, errors.New("idp_entity_id, idp_sso_target_url, and idp_certificate are required")
|
||||
}
|
||||
return req, nil
|
||||
}
|
||||
|
||||
func bindUpdateSamlSettingsRequest(c *gin.Context) (UpdateSamlSettingsRequest, error) {
|
||||
req := UpdateSamlSettingsRequest{}
|
||||
return req, bindSamlSettingsPayload(c, &req)
|
||||
}
|
||||
|
||||
func bindSamlSettingsPayload(c *gin.Context, req interface{}) error {
|
||||
var payload map[string]json.RawMessage
|
||||
if err := json.NewDecoder(c.Request.Body).Decode(&payload); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if nested, ok := payload["saml_settings"]; ok {
|
||||
return json.Unmarshal(nested, req)
|
||||
}
|
||||
|
||||
flat, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return json.Unmarshal(flat, req)
|
||||
}
|
||||
|
||||
func serializeSamlSettings(settings *model.AccountSamlSettings) gin.H {
|
||||
if settings == nil {
|
||||
return gin.H{}
|
||||
}
|
||||
|
||||
return gin.H{
|
||||
"id": settings.ID,
|
||||
"account_id": settings.AccountID,
|
||||
"idp_entity_id": settings.IdpEntityID,
|
||||
"idp_sso_target_url": settings.IdpSsoTargetURL,
|
||||
"idp_slo_target_url": settings.IdpSloTargetURL,
|
||||
"idp_certificate": settings.IdpCertificate,
|
||||
"sso_url": settings.IdpSsoTargetURL,
|
||||
"certificate": settings.IdpCertificate,
|
||||
"sp_entity_id": settings.SpEntityID,
|
||||
"sp_x509_certificate": settings.SpX509Certificate,
|
||||
"role_mappings": settings.RoleMappings,
|
||||
"active": settings.Active,
|
||||
"fingerprint": samlCertificateFingerprint(settings.IdpCertificate),
|
||||
"created_at": settings.CreatedAt,
|
||||
"updated_at": settings.UpdatedAt,
|
||||
}
|
||||
}
|
||||
|
||||
func samlCertificateFingerprint(certificate string) string {
|
||||
if certificate == "" {
|
||||
return ""
|
||||
}
|
||||
sum := sha1.Sum([]byte(certificate))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
@@ -1,258 +0,0 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/suite"
|
||||
"gorm.io/driver/sqlite"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/gochat/gochat/internal/model"
|
||||
"github.com/gochat/gochat/internal/repository"
|
||||
)
|
||||
|
||||
type AccountSamlSettingsHandlerTestSuite struct {
|
||||
suite.Suite
|
||||
db *gorm.DB
|
||||
handler *AccountSamlSettingsHandler
|
||||
router *gin.Engine
|
||||
account *model.Account
|
||||
}
|
||||
|
||||
func (s *AccountSamlSettingsHandlerTestSuite) SetupSuite() {
|
||||
s.db, _ = gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
s.db.AutoMigrate(&model.Account{}, &model.AccountSamlSettings{})
|
||||
|
||||
repo := repository.NewAccountSamlSettingsRepo(s.db)
|
||||
s.handler = NewAccountSamlSettingsHandler(repo)
|
||||
|
||||
gin.SetMode(gin.TestMode)
|
||||
r := gin.New()
|
||||
RegisterAccountSamlSettingsRoutes(r.Group("/api/v1/accounts/:account_id/saml_settings"), s.handler)
|
||||
s.router = r
|
||||
|
||||
s.account = &model.Account{Name: "TestAccount"}
|
||||
s.db.Create(s.account)
|
||||
}
|
||||
|
||||
func (s *AccountSamlSettingsHandlerTestSuite) SetupTest() {
|
||||
s.db.Exec("DELETE FROM account_saml_settings")
|
||||
}
|
||||
|
||||
func TestAccountSamlSettingsHandlerTestSuite(t *testing.T) {
|
||||
suite.Run(t, new(AccountSamlSettingsHandlerTestSuite))
|
||||
}
|
||||
|
||||
func (s *AccountSamlSettingsHandlerTestSuite) TestGet_InvalidAccountID() {
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/v1/accounts/abc/saml_settings", nil)
|
||||
s.router.ServeHTTP(w, req)
|
||||
s.Equal(http.StatusBadRequest, w.Code)
|
||||
}
|
||||
|
||||
func (s *AccountSamlSettingsHandlerTestSuite) TestGet_NotFound() {
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, fmt.Sprintf("/api/v1/accounts/%d/saml_settings", s.account.ID), nil)
|
||||
s.router.ServeHTTP(w, req)
|
||||
s.Equal(http.StatusNotFound, w.Code)
|
||||
}
|
||||
|
||||
func (s *AccountSamlSettingsHandlerTestSuite) TestGet_Success() {
|
||||
settings := &model.AccountSamlSettings{
|
||||
AccountID: s.account.ID,
|
||||
IdpEntityID: "entity-id",
|
||||
IdpSsoTargetURL: "https://sso.example.com",
|
||||
}
|
||||
s.db.Create(settings)
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, fmt.Sprintf("/api/v1/accounts/%d/saml_settings", s.account.ID), nil)
|
||||
s.router.ServeHTTP(w, req)
|
||||
s.Equal(http.StatusOK, w.Code)
|
||||
|
||||
var payload map[string]interface{}
|
||||
s.Require().NoError(json.Unmarshal(w.Body.Bytes(), &payload))
|
||||
s.Equal("https://sso.example.com", payload["sso_url"])
|
||||
s.Equal("https://sso.example.com", payload["idp_sso_target_url"])
|
||||
s.Equal("entity-id", payload["idp_entity_id"])
|
||||
s.NotContains(payload, "data")
|
||||
}
|
||||
|
||||
func (s *AccountSamlSettingsHandlerTestSuite) TestChatwootFrontendCollectionCRUDPayloads() {
|
||||
createBody := bytes.NewBufferString(`{"saml_settings":{"sso_url":"https://idp.example.com/saml","certificate":"-----BEGIN CERTIFICATE-----chatwoot-----END CERTIFICATE-----","idp_entity_id":"chatwoot-idp","role_mappings":{}}}`)
|
||||
createReq := httptest.NewRequest(http.MethodPost, fmt.Sprintf("/api/v1/accounts/%d/saml_settings", s.account.ID), createBody)
|
||||
createReq.Header.Set("Content-Type", "application/json")
|
||||
createRecorder := httptest.NewRecorder()
|
||||
s.router.ServeHTTP(createRecorder, createReq)
|
||||
s.Equal(http.StatusCreated, createRecorder.Code)
|
||||
|
||||
var created map[string]interface{}
|
||||
s.Require().NoError(json.Unmarshal(createRecorder.Body.Bytes(), &created))
|
||||
s.NotZero(created["id"])
|
||||
s.Equal("https://idp.example.com/saml", created["sso_url"])
|
||||
s.Equal("https://idp.example.com/saml", created["idp_sso_target_url"])
|
||||
s.Equal("-----BEGIN CERTIFICATE-----chatwoot-----END CERTIFICATE-----", created["certificate"])
|
||||
s.Equal("-----BEGIN CERTIFICATE-----chatwoot-----END CERTIFICATE-----", created["idp_certificate"])
|
||||
s.Equal("chatwoot-idp", created["idp_entity_id"])
|
||||
s.NotEmpty(created["fingerprint"])
|
||||
s.NotContains(created, "data")
|
||||
s.NotContains(created, "success")
|
||||
|
||||
getReq := httptest.NewRequest(http.MethodGet, fmt.Sprintf("/api/v1/accounts/%d/saml_settings", s.account.ID), nil)
|
||||
getRecorder := httptest.NewRecorder()
|
||||
s.router.ServeHTTP(getRecorder, getReq)
|
||||
s.Equal(http.StatusOK, getRecorder.Code)
|
||||
|
||||
var fetched map[string]interface{}
|
||||
s.Require().NoError(json.Unmarshal(getRecorder.Body.Bytes(), &fetched))
|
||||
s.Equal(created["id"], fetched["id"])
|
||||
s.Equal("https://idp.example.com/saml", fetched["sso_url"])
|
||||
|
||||
updateBody := bytes.NewBufferString(`{"saml_settings":{"sso_url":"https://idp.example.com/updated","certificate":"updated-certificate","idp_entity_id":"updated-idp"}}`)
|
||||
updateReq := httptest.NewRequest(http.MethodPut, fmt.Sprintf("/api/v1/accounts/%d/saml_settings", s.account.ID), updateBody)
|
||||
updateReq.Header.Set("Content-Type", "application/json")
|
||||
updateRecorder := httptest.NewRecorder()
|
||||
s.router.ServeHTTP(updateRecorder, updateReq)
|
||||
s.Equal(http.StatusOK, updateRecorder.Code)
|
||||
|
||||
var updated map[string]interface{}
|
||||
s.Require().NoError(json.Unmarshal(updateRecorder.Body.Bytes(), &updated))
|
||||
s.Equal(created["id"], updated["id"])
|
||||
s.Equal("https://idp.example.com/updated", updated["sso_url"])
|
||||
s.Equal("https://idp.example.com/updated", updated["idp_sso_target_url"])
|
||||
s.Equal("updated-certificate", updated["certificate"])
|
||||
s.Equal("updated-certificate", updated["idp_certificate"])
|
||||
s.Equal("updated-idp", updated["idp_entity_id"])
|
||||
s.NotContains(updated, "data")
|
||||
|
||||
deleteReq := httptest.NewRequest(http.MethodDelete, fmt.Sprintf("/api/v1/accounts/%d/saml_settings", s.account.ID), nil)
|
||||
deleteRecorder := httptest.NewRecorder()
|
||||
s.router.ServeHTTP(deleteRecorder, deleteReq)
|
||||
s.Equal(http.StatusOK, deleteRecorder.Code)
|
||||
|
||||
afterDeleteReq := httptest.NewRequest(http.MethodGet, fmt.Sprintf("/api/v1/accounts/%d/saml_settings", s.account.ID), nil)
|
||||
afterDeleteRecorder := httptest.NewRecorder()
|
||||
s.router.ServeHTTP(afterDeleteRecorder, afterDeleteReq)
|
||||
s.Equal(http.StatusNotFound, afterDeleteRecorder.Code)
|
||||
}
|
||||
|
||||
func (s *AccountSamlSettingsHandlerTestSuite) TestCreate_InvalidAccountID() {
|
||||
w := httptest.NewRecorder()
|
||||
body := bytes.NewBufferString(`{"idp_entity_id":"eid","idp_sso_target_url":"https://sso.example.com"}`)
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/accounts/abc/saml_settings", body)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
s.router.ServeHTTP(w, req)
|
||||
s.Equal(http.StatusBadRequest, w.Code)
|
||||
}
|
||||
|
||||
func (s *AccountSamlSettingsHandlerTestSuite) TestCreate_Success() {
|
||||
w := httptest.NewRecorder()
|
||||
body := bytes.NewBufferString(fmt.Sprintf(`{"account_id":%d,"idp_entity_id":"eid","idp_sso_target_url":"https://sso.example.com","idp_certificate":"cert"}`, s.account.ID))
|
||||
req := httptest.NewRequest(http.MethodPost, fmt.Sprintf("/api/v1/accounts/%d/saml_settings", s.account.ID), body)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
s.router.ServeHTTP(w, req)
|
||||
s.Equal(http.StatusCreated, w.Code)
|
||||
}
|
||||
|
||||
func (s *AccountSamlSettingsHandlerTestSuite) TestCreate_MissingRequiredFields() {
|
||||
w := httptest.NewRecorder()
|
||||
body := bytes.NewBufferString(`{"account_id":0}`)
|
||||
req := httptest.NewRequest(http.MethodPost, fmt.Sprintf("/api/v1/accounts/%d/saml_settings", s.account.ID), body)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
s.router.ServeHTTP(w, req)
|
||||
s.True(w.Code == http.StatusBadRequest || w.Code == http.StatusUnprocessableEntity, "expected 400 or 500, got %d", w.Code)
|
||||
}
|
||||
|
||||
func (s *AccountSamlSettingsHandlerTestSuite) TestUpdate_InvalidAccountID() {
|
||||
w := httptest.NewRecorder()
|
||||
body := bytes.NewBufferString(`{"idp_entity_id":"updated-eid"}`)
|
||||
req := httptest.NewRequest(http.MethodPut, "/api/v1/accounts/abc/saml_settings", body)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
s.router.ServeHTTP(w, req)
|
||||
s.Equal(http.StatusBadRequest, w.Code)
|
||||
}
|
||||
|
||||
func (s *AccountSamlSettingsHandlerTestSuite) TestUpdate_Success() {
|
||||
settings := &model.AccountSamlSettings{
|
||||
AccountID: s.account.ID,
|
||||
IdpEntityID: "entity-id",
|
||||
IdpSsoTargetURL: "https://sso.example.com",
|
||||
}
|
||||
s.db.Create(settings)
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
body := bytes.NewBufferString(`{"idp_entity_id":"updated-eid"}`)
|
||||
req := httptest.NewRequest(http.MethodPut, fmt.Sprintf("/api/v1/accounts/%d/saml_settings/%d", s.account.ID, settings.ID), body)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
s.router.ServeHTTP(w, req)
|
||||
s.Equal(http.StatusOK, w.Code)
|
||||
|
||||
var payload map[string]interface{}
|
||||
s.Require().NoError(json.Unmarshal(w.Body.Bytes(), &payload))
|
||||
s.Equal("updated-eid", payload["idp_entity_id"])
|
||||
s.NotContains(payload, "data")
|
||||
}
|
||||
|
||||
func (s *AccountSamlSettingsHandlerTestSuite) TestUpdate_NotFound() {
|
||||
// account_id valid, but no settings exist for this account → repo returns "record not found" → 404
|
||||
w := httptest.NewRecorder()
|
||||
body := bytes.NewBufferString(`{"idp_entity_id":"updated-eid"}`)
|
||||
req := httptest.NewRequest(http.MethodPut, fmt.Sprintf("/api/v1/accounts/%d/saml_settings/999", s.account.ID), body)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
s.router.ServeHTTP(w, req)
|
||||
s.Equal(http.StatusNotFound, w.Code)
|
||||
}
|
||||
|
||||
func (s *AccountSamlSettingsHandlerTestSuite) TestDelete_InvalidAccountID() {
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodDelete, "/api/v1/accounts/abc/saml_settings", nil)
|
||||
s.router.ServeHTTP(w, req)
|
||||
s.Equal(http.StatusBadRequest, w.Code)
|
||||
}
|
||||
|
||||
func (s *AccountSamlSettingsHandlerTestSuite) TestDelete_Success() {
|
||||
settings := &model.AccountSamlSettings{
|
||||
AccountID: s.account.ID,
|
||||
IdpEntityID: "entity-id",
|
||||
IdpSsoTargetURL: "https://sso.example.com",
|
||||
}
|
||||
s.db.Create(settings)
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodDelete, fmt.Sprintf("/api/v1/accounts/%d/saml_settings/%d", s.account.ID, settings.ID), nil)
|
||||
s.router.ServeHTTP(w, req)
|
||||
// Handler returns 200 OK with message, not 204 NoContent
|
||||
s.Equal(http.StatusOK, w.Code)
|
||||
}
|
||||
|
||||
func (s *AccountSamlSettingsHandlerTestSuite) TestToggleActive_InvalidAccountID() {
|
||||
w := httptest.NewRecorder()
|
||||
body := bytes.NewBufferString(`{"active":true}`)
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/accounts/abc/saml_settings/1/toggle_active", body)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
s.router.ServeHTTP(w, req)
|
||||
s.Equal(http.StatusBadRequest, w.Code)
|
||||
}
|
||||
|
||||
func (s *AccountSamlSettingsHandlerTestSuite) TestToggleActive_Success() {
|
||||
settings := &model.AccountSamlSettings{
|
||||
AccountID: s.account.ID,
|
||||
IdpEntityID: "entity-id",
|
||||
IdpSsoTargetURL: "https://sso.example.com",
|
||||
Active: false,
|
||||
}
|
||||
s.db.Create(settings)
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
body := bytes.NewBufferString(`{"active":true}`)
|
||||
req := httptest.NewRequest(http.MethodPost, fmt.Sprintf("/api/v1/accounts/%d/saml_settings/%d/toggle_active", s.account.ID, settings.ID), body)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
s.router.ServeHTTP(w, req)
|
||||
s.Equal(http.StatusOK, w.Code)
|
||||
}
|
||||
@@ -52,12 +52,6 @@ type LoginRequest struct {
|
||||
Password string `json:"password" binding:"required,min=6"`
|
||||
}
|
||||
|
||||
// LoginMFAResquest is the JSON body for MFA login verification.
|
||||
type LoginMFAResquest struct {
|
||||
UserID uint `json:"user_id" binding:"required"`
|
||||
TOTPCode string `json:"totp_code" binding:"required"`
|
||||
}
|
||||
|
||||
// RefreshRequest is the JSON body for refresh endpoint.
|
||||
type RefreshRequest struct {
|
||||
RefreshToken string `json:"refresh_token" binding:"required"`
|
||||
@@ -87,7 +81,6 @@ type ConfirmEmailRequest struct {
|
||||
|
||||
// Login authenticates a user with email/password and returns JWT tokens.
|
||||
// POST /api/v1/auth/login
|
||||
// If MFA is enabled, returns mfa_required=true with user_id for TOTP verification.
|
||||
func (h *AuthHandler) Login(c *gin.Context) {
|
||||
var req LoginRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
@@ -104,15 +97,6 @@ func (h *AuthHandler) Login(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
if output.MFARequired {
|
||||
response.OK(c, gin.H{
|
||||
"mfa_required": true,
|
||||
"user_id": output.User.ID,
|
||||
"message": "MFA verification required, please provide TOTP code",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
response.OK(c, gin.H{
|
||||
"user": output.User,
|
||||
"access_token": output.TokenPair.AccessToken,
|
||||
@@ -141,13 +125,6 @@ func (h *AuthHandler) ChatwootSignIn(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
if output.MFARequired {
|
||||
c.JSON(http.StatusPartialContent, gin.H{
|
||||
"mfa_required": true,
|
||||
"mfa_token": strconv.FormatUint(uint64(output.User.ID), 10),
|
||||
})
|
||||
return
|
||||
}
|
||||
if err := h.trackChatwootSession(c, output); err != nil {
|
||||
response.AbortWithStatusError(c, http.StatusInternalServerError, response.ErrInternal, "failed to create session")
|
||||
return
|
||||
@@ -213,31 +190,6 @@ func (h *AuthHandler) ChatwootSignOut(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, gin.H{"success": true})
|
||||
}
|
||||
|
||||
// LoginMFA completes login after MFA TOTP code verification.
|
||||
// POST /api/v1/auth/login/mfa
|
||||
func (h *AuthHandler) LoginMFA(c *gin.Context) {
|
||||
var req LoginMFAResquest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrValidation, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
output, err := h.authService.LoginWithMFA(c.Request.Context(), req.UserID, req.TOTPCode)
|
||||
if err != nil {
|
||||
response.AbortWithStatusError(c, http.StatusUnauthorized, response.ErrUnauthorized, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
response.OK(c, gin.H{
|
||||
"user": output.User,
|
||||
"access_token": output.TokenPair.AccessToken,
|
||||
"refresh_token": output.TokenPair.RefreshToken,
|
||||
"expires_at": output.TokenPair.ExpiresAt,
|
||||
"account_id": output.AccountID,
|
||||
"role": output.Role,
|
||||
})
|
||||
}
|
||||
|
||||
// Refresh rotates a refresh token and returns new JWT pair.
|
||||
// POST /api/v1/auth/refresh
|
||||
// Implements refresh token rotation per P2E §1.4 security requirement.
|
||||
@@ -423,7 +375,6 @@ func RegisterAuthRoutes(rg *gin.RouterGroup, handler *AuthHandler) {
|
||||
{
|
||||
// Core auth endpoints
|
||||
authGroup.POST("/login", handler.Login)
|
||||
authGroup.POST("/login/mfa", handler.LoginMFA)
|
||||
authGroup.POST("/refresh", handler.Refresh)
|
||||
authGroup.DELETE("/logout", handler.Logout)
|
||||
|
||||
@@ -488,7 +439,7 @@ func extractChatwootAccessToken(c *gin.Context) string {
|
||||
}
|
||||
|
||||
// generateOAuthState creates a cryptographically random state token for CSRF protection.
|
||||
// Used by SAML and other auth flows. Production note: state should also be stored
|
||||
// Used by OIDC and other auth flows. Production note: state should also be stored
|
||||
// server-side (Redis) and validated on callback.
|
||||
func generateOAuthState() string {
|
||||
return "gochat_oauth_" + randomHex(16)
|
||||
|
||||
@@ -56,7 +56,7 @@ func setupChatwootAuthTest(t *testing.T) (*gin.Engine, *gorm.DB, *model.User) {
|
||||
jwtCfg := &config.JWTConfig{Secret: "auth-test-secret", ExpiryHours: 1, RefreshExpiryHours: 24}
|
||||
jwtSvc := auth.NewJWTService(jwtCfg)
|
||||
refreshStore := auth.NewRefreshTokenStore(nil, jwtCfg)
|
||||
authSvc := service.NewAuthService(db, jwtSvc, refreshStore, nil)
|
||||
authSvc := service.NewAuthService(db, jwtSvc, refreshStore)
|
||||
profileSvc := service.NewProfileService(repository.NewUserRepo(db), repository.NewAccountUserRepo(db), repository.NewAccessTokenRepo(db))
|
||||
handler := NewAuthHandler(authSvc, profileSvc)
|
||||
|
||||
|
||||
@@ -1,580 +0,0 @@
|
||||
package v1
|
||||
|
||||
// Reference: M13 §4.4 — LDAP HTTP endpoints for login, config, and connectivity testing
|
||||
// Provides four endpoints for LDAP/Active Directory integration:
|
||||
// - POST /api/v1/ldap/login → LDAP bind authentication + JWT issuance
|
||||
// - POST /api/v1/ldap/test → Admin-only LDAP connectivity test
|
||||
// - GET /api/v1/ldap/config → Admin-only LDAP settings retrieval
|
||||
// - PUT /api/v1/ldap/config → Admin-only LDAP settings update
|
||||
//
|
||||
// Enterprise feature: GoChat extends beyond Chatwoot's SAML-only SSO by adding
|
||||
// LDAP support for traditional enterprise AD/LDAP environments.
|
||||
//
|
||||
// Login flow:
|
||||
// 1. Client sends {username, password, account_id} to /api/v1/ldap/login
|
||||
// 2. Handler delegates to SSOMiddleware.AuthenticateLDAP for unified SSO processing
|
||||
// 3. SSOMiddleware routes to LDAPService.Authenticate (Bind + search + group extraction)
|
||||
// 4. On success: SSOMiddleware auto-provisions user, maps groups→roles, creates SSO session
|
||||
// 5. Handler issues JWT token pair via JWTService.GenerateTokenPair
|
||||
// 6. Store refresh token, return access+refresh tokens to client
|
||||
//
|
||||
// Config management (admin-only):
|
||||
// - Account administrators can configure per-account LDAP settings
|
||||
// - TestConnection validates LDAP bind connectivity before saving config
|
||||
// - GetConfig/UpdateConfig manage per-account LDAP settings in DB
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/gochat/gochat/internal/auth"
|
||||
"github.com/gochat/gochat/internal/config"
|
||||
"github.com/gochat/gochat/internal/model"
|
||||
"github.com/gochat/gochat/pkg/response"
|
||||
applogger "github.com/gochat/gochat/pkg/logger"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// LDAPHandler handles LDAP authentication HTTP endpoints.
|
||||
type LDAPHandler struct {
|
||||
ldapService *auth.LDAPService
|
||||
ssoMiddleware *auth.SSOMiddleware
|
||||
jwtService *auth.JWTService
|
||||
refreshStore *auth.RefreshTokenStore
|
||||
ssoSessionStore *auth.SSOSessionStore
|
||||
ldapCfg *config.LDAPConfig
|
||||
db *gorm.DB
|
||||
}
|
||||
|
||||
// NewLDAPHandler creates an LDAP handler with service dependencies.
|
||||
func NewLDAPHandler(
|
||||
ldapService *auth.LDAPService,
|
||||
ssoMiddleware *auth.SSOMiddleware,
|
||||
jwtService *auth.JWTService,
|
||||
refreshStore *auth.RefreshTokenStore,
|
||||
ssoSessionStore *auth.SSOSessionStore,
|
||||
ldapCfg *config.LDAPConfig,
|
||||
db *gorm.DB,
|
||||
) *LDAPHandler {
|
||||
return &LDAPHandler{
|
||||
ldapService: ldapService,
|
||||
ssoMiddleware: ssoMiddleware,
|
||||
jwtService: jwtService,
|
||||
refreshStore: refreshStore,
|
||||
ssoSessionStore: ssoSessionStore,
|
||||
ldapCfg: ldapCfg,
|
||||
db: db,
|
||||
}
|
||||
}
|
||||
|
||||
// ldapLoginRequest is the JSON body for POST /api/v1/ldap/login.
|
||||
type ldapLoginRequest struct {
|
||||
Username string `json:"username" binding:"required"`
|
||||
Password string `json:"password" binding:"required"`
|
||||
AccountID uint `json:"account_id" binding:"required"`
|
||||
}
|
||||
|
||||
// Login authenticates a user via LDAP bind and issues a JWT token pair.
|
||||
// POST /api/v1/ldap/login
|
||||
// This endpoint is PUBLIC — no AuthMiddleware required (LDAP login is the entry point).
|
||||
func (h *LDAPHandler) Login(c *gin.Context) {
|
||||
if !h.ldapCfg.Enabled {
|
||||
c.JSON(http.StatusNotFound, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrNotFound,
|
||||
Message: "LDAP is not enabled",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
var req ldapLoginRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrBadRequest,
|
||||
Message: "Invalid request body",
|
||||
Detail: err.Error(),
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Delegate to SSO middleware for unified authentication flow
|
||||
// (auto-provision, group→role mapping, SSO session creation)
|
||||
result, err := h.ssoMiddleware.AuthenticateLDAP(c.Request.Context(), req.AccountID, req.Username, req.Password)
|
||||
if err != nil {
|
||||
applogger.L().Errorf("LDAP authentication failed (account=%d, username=%s): %v", req.AccountID, req.Username, err)
|
||||
|
||||
statusCode := http.StatusInternalServerError
|
||||
errCode := response.ErrInternal
|
||||
if errors.Is(err, auth.ErrLDAPDisabled) {
|
||||
statusCode = http.StatusNotFound
|
||||
errCode = response.ErrNotFound
|
||||
} else if errors.Is(err, auth.ErrLDAPInvalidConfig) {
|
||||
statusCode = http.StatusBadRequest
|
||||
errCode = response.ErrBadRequest
|
||||
} else if errors.Is(err, auth.ErrLDAPConnection) {
|
||||
statusCode = http.StatusServiceUnavailable
|
||||
errCode = response.ErrServiceUnavail
|
||||
} else if errors.Is(err, auth.ErrLDAPUserNotFound) || errors.Is(err, auth.ErrLDAPBindFailed) {
|
||||
statusCode = http.StatusUnauthorized
|
||||
errCode = response.ErrUnauthorized
|
||||
}
|
||||
|
||||
c.JSON(statusCode, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: errCode,
|
||||
Message: "LDAP authentication failed",
|
||||
Detail: err.Error(),
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Issue JWT token pair using JWTService
|
||||
// SSO middleware already created the user and mapped roles
|
||||
tokenPair, err := h.jwtService.GenerateTokenPair(
|
||||
&model.User{
|
||||
Base: model.Base{ID: result.UserID},
|
||||
Email: result.Email,
|
||||
Name: result.Name,
|
||||
Provider: "ldap",
|
||||
UID: result.Subject,
|
||||
Role: result.Role,
|
||||
},
|
||||
result.AccountID,
|
||||
result.Role,
|
||||
)
|
||||
if err != nil {
|
||||
applogger.L().Errorf("Failed to generate JWT for LDAP user (account=%d, user=%d): %v", result.AccountID, result.UserID, err)
|
||||
c.JSON(http.StatusUnprocessableEntity, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrInternal,
|
||||
Message: "Failed to generate authentication tokens",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Store refresh token
|
||||
if h.refreshStore != nil {
|
||||
if err := h.refreshStore.Store(c.Request.Context(), result.UserID, tokenPair.RefreshToken); err != nil {
|
||||
applogger.L().Warnf("Failed to store refresh token for LDAP user %d: %v", result.UserID, err)
|
||||
// Non-fatal: access token is still valid, refresh just won't work until re-login
|
||||
}
|
||||
}
|
||||
|
||||
// Create SSO session in Redis (for session tracking and SLO)
|
||||
if h.ssoSessionStore != nil {
|
||||
sessionData := &auth.SSOSessionData{
|
||||
UserID: result.UserID,
|
||||
Provider: "ldap",
|
||||
IdPEntityID: fmt.Sprintf("ldap-account-%d", result.AccountID), // LDAP server as IdP identifier
|
||||
NameID: result.Subject, // LDAP DN as NameID
|
||||
AccountID: result.AccountID,
|
||||
Role: result.Role,
|
||||
CreatedAt: time.Now().Unix(),
|
||||
ExpiresAt: time.Now().Add(h.ssoSessionStore.SessionTTL()).Unix(),
|
||||
}
|
||||
sessionID, err := h.ssoSessionStore.Create(c.Request.Context(), sessionData)
|
||||
if err != nil {
|
||||
applogger.L().Warnf("Failed to create SSO session for LDAP user %d: %v", result.UserID, err)
|
||||
// Non-fatal: JWT tokens are still valid, SSO session is for tracking/SLO only
|
||||
} else {
|
||||
applogger.L().Infof("SSO session %s created for LDAP user %d (account=%d)", sessionID, result.UserID, result.AccountID)
|
||||
}
|
||||
}
|
||||
|
||||
// Return successful auth response (same format as SAML ACS and regular login)
|
||||
response.OK(c, gin.H{
|
||||
"user": gin.H{
|
||||
"id": result.UserID,
|
||||
"email": result.Email,
|
||||
"name": result.Name,
|
||||
"provider": "ldap",
|
||||
"uid": result.Subject,
|
||||
"role": result.Role,
|
||||
},
|
||||
"access_token": tokenPair.AccessToken,
|
||||
"refresh_token": tokenPair.RefreshToken,
|
||||
"expires_at": tokenPair.ExpiresAt,
|
||||
})
|
||||
}
|
||||
|
||||
// ldapTestRequest is the JSON body for POST /api/v1/ldap/test.
|
||||
type ldapTestRequest struct {
|
||||
AccountID uint `json:"account_id" binding:"required"`
|
||||
}
|
||||
|
||||
// getAccountSettingsFromDB loads per-account LDAP configuration from DB directly.
|
||||
// This is needed because LDAPService.getAccountSettings is unexported.
|
||||
func (h *LDAPHandler) getAccountSettingsFromDB(accountID uint) (*model.AccountLDAPSettings, error) {
|
||||
var settings model.AccountLDAPSettings
|
||||
err := h.db.Where("account_id = ? AND active = ?", accountID, true).First(&settings).Error
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
// No per-account settings — check if global defaults exist
|
||||
if h.ldapCfg.DefaultHost == "" {
|
||||
return nil, nil // LDAP not configured for this account
|
||||
}
|
||||
// Use global defaults as a fallback
|
||||
settings = model.AccountLDAPSettings{
|
||||
AccountID: accountID,
|
||||
Host: h.ldapCfg.DefaultHost,
|
||||
Port: h.ldapCfg.DefaultPort,
|
||||
UseTLS: h.ldapCfg.DefaultUseTLS,
|
||||
BaseDN: h.ldapCfg.DefaultBaseDN,
|
||||
BindDN: h.ldapCfg.DefaultBindDN,
|
||||
BindPassword: h.ldapCfg.DefaultBindPassword,
|
||||
UserFilter: h.ldapCfg.DefaultUserFilter,
|
||||
EmailAttribute: h.ldapCfg.DefaultEmailAttribute,
|
||||
NameAttribute: h.ldapCfg.DefaultNameAttribute,
|
||||
GroupAttribute: h.ldapCfg.DefaultGroupAttribute,
|
||||
AutoProvision: true,
|
||||
Active: true,
|
||||
}
|
||||
return &settings, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
return &settings, nil
|
||||
}
|
||||
|
||||
// testLDAPConnectivity tests LDAP bind connectivity for an account's configuration.
|
||||
// Connects to the LDAP server and attempts a bind with the service account to verify
|
||||
// that the configuration is correct before saving.
|
||||
func (h *LDAPHandler) testLDAPConnectivity(settings *model.AccountLDAPSettings) error {
|
||||
// Use LDAPService.Authenticate with a dummy test to verify connectivity.
|
||||
// The LDAPService handles connection + bind internally, so we use it to validate.
|
||||
// We attempt a bind-only test by calling Authenticate with empty credentials
|
||||
// and catching the specific error pattern.
|
||||
// However, since Authenticate requires a real username/password, we instead
|
||||
// try to directly connect and bind using the service account credentials.
|
||||
//
|
||||
// For simplicity, we delegate to ldapService.Authenticate with a test username.
|
||||
// If the connection itself fails, we get ErrLDAPConnection.
|
||||
// If the service account bind fails, we get an appropriate error.
|
||||
// If the user search fails (expected for test username), we know connectivity works.
|
||||
ctx := context.Background()
|
||||
_, err := h.ldapService.Authenticate(ctx, settings.AccountID, "__ldap_connectivity_test__", "__invalid_test_password__")
|
||||
if err == nil {
|
||||
// Unexpected: test credentials actually worked. Still means connectivity is good.
|
||||
return nil
|
||||
}
|
||||
// If connection failed, return that error
|
||||
if errors.Is(err, auth.ErrLDAPConnection) || errors.Is(err, auth.ErrLDAPDisabled) || errors.Is(err, auth.ErrLDAPInvalidConfig) {
|
||||
return err
|
||||
}
|
||||
// If we got ErrLDAPUserNotFound or ErrLDAPBindFailed, that means the connection
|
||||
// and service account bind succeeded — only the test user lookup/bind failed,
|
||||
// which is expected. Connectivity is confirmed.
|
||||
return nil
|
||||
}
|
||||
|
||||
// TestConnection tests LDAP bind connectivity for an account's configuration.
|
||||
// POST /api/v1/ldap/test
|
||||
// Admin-only: requires AuthMiddleware + admin role (enforced at router level).
|
||||
func (h *LDAPHandler) TestConnection(c *gin.Context) {
|
||||
if !h.ldapCfg.Enabled {
|
||||
c.JSON(http.StatusNotFound, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrNotFound,
|
||||
Message: "LDAP is not enabled",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
var req ldapTestRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrBadRequest,
|
||||
Message: "Invalid request body",
|
||||
Detail: err.Error(),
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Load per-account LDAP settings from DB
|
||||
settings, err := h.getAccountSettingsFromDB(req.AccountID)
|
||||
if err != nil {
|
||||
applogger.L().Errorf("Failed to load LDAP settings for account %d: %v", req.AccountID, err)
|
||||
c.JSON(http.StatusUnprocessableEntity, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrInternal,
|
||||
Message: "Failed to load LDAP settings",
|
||||
Detail: err.Error(),
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
if settings == nil {
|
||||
c.JSON(http.StatusNotFound, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrNotFound,
|
||||
Message: "No LDAP configuration found for this account",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Test LDAP connectivity
|
||||
err = h.testLDAPConnectivity(settings)
|
||||
if err != nil {
|
||||
applogger.L().Errorf("LDAP connectivity test failed (account=%d, host=%s:%d): %v", req.AccountID, settings.Host, settings.Port, err)
|
||||
c.JSON(http.StatusBadRequest, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrBadRequest,
|
||||
Message: "LDAP connectivity test failed",
|
||||
Detail: err.Error(),
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
response.OK(c, gin.H{
|
||||
"account_id": req.AccountID,
|
||||
"host": settings.Host,
|
||||
"port": settings.Port,
|
||||
"connected": true,
|
||||
})
|
||||
}
|
||||
|
||||
// GetConfig retrieves LDAP settings for an account.
|
||||
// GET /api/v1/ldap/config?account_id=123
|
||||
// Admin-only: requires AuthMiddleware + admin role (enforced at router level).
|
||||
func (h *LDAPHandler) GetConfig(c *gin.Context) {
|
||||
if !h.ldapCfg.Enabled {
|
||||
c.JSON(http.StatusNotFound, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrNotFound,
|
||||
Message: "LDAP is not enabled",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
accountIDStr := c.Query("account_id")
|
||||
if accountIDStr == "" {
|
||||
c.JSON(http.StatusBadRequest, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrBadRequest,
|
||||
Message: "account_id query parameter is required",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
accountID, err := strconv.ParseUint(accountIDStr, 10, 64)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrBadRequest,
|
||||
Message: "Invalid account_id",
|
||||
Detail: err.Error(),
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
settings, err := h.getAccountSettingsFromDB(uint(accountID))
|
||||
if err != nil {
|
||||
applogger.L().Errorf("Failed to get LDAP config for account %d: %v", accountID, err)
|
||||
c.JSON(http.StatusUnprocessableEntity, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrInternal,
|
||||
Message: "Failed to retrieve LDAP configuration",
|
||||
Detail: err.Error(),
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
if settings == nil {
|
||||
c.JSON(http.StatusNotFound, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrNotFound,
|
||||
Message: "No LDAP configuration found for this account",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
response.OK(c, settings)
|
||||
}
|
||||
|
||||
// ldapUpdateConfigRequest is the JSON body for PUT /api/v1/ldap/config.
|
||||
type ldapUpdateConfigRequest struct {
|
||||
AccountID uint `json:"account_id" binding:"required"`
|
||||
Host string `json:"host" binding:"required"`
|
||||
Port int `json:"port"`
|
||||
UseTLS bool `json:"use_tls"`
|
||||
BaseDN string `json:"base_dn" binding:"required"`
|
||||
BindDN string `json:"bind_dn,omitempty"`
|
||||
BindPassword string `json:"bind_password,omitempty"`
|
||||
UserFilter string `json:"user_filter"`
|
||||
EmailAttribute string `json:"email_attribute"`
|
||||
NameAttribute string `json:"name_attribute"`
|
||||
FirstNameAttribute string `json:"first_name_attribute"`
|
||||
LastNameAttribute string `json:"last_name_attribute"`
|
||||
GroupAttribute string `json:"group_attribute"`
|
||||
GroupFilter string `json:"group_filter"`
|
||||
RoleMappings json.RawMessage `json:"role_mappings"`
|
||||
AutoProvision bool `json:"auto_provision"`
|
||||
SyncInterval int `json:"sync_interval"`
|
||||
Active bool `json:"active"`
|
||||
}
|
||||
|
||||
// UpdateConfig updates per-account LDAP settings.
|
||||
// PUT /api/v1/ldap/config
|
||||
// Admin-only: requires AuthMiddleware + admin role (enforced at router level).
|
||||
func (h *LDAPHandler) UpdateConfig(c *gin.Context) {
|
||||
if !h.ldapCfg.Enabled {
|
||||
c.JSON(http.StatusNotFound, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrNotFound,
|
||||
Message: "LDAP is not enabled",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
var req ldapUpdateConfigRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrBadRequest,
|
||||
Message: "Invalid request body",
|
||||
Detail: err.Error(),
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Default port values
|
||||
if req.Port == 0 {
|
||||
if req.UseTLS {
|
||||
req.Port = 636
|
||||
} else {
|
||||
req.Port = 389
|
||||
}
|
||||
}
|
||||
|
||||
// Load existing settings or create new
|
||||
var settings model.AccountLDAPSettings
|
||||
err := h.db.Where("account_id = ?", req.AccountID).First(&settings).Error
|
||||
if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
applogger.L().Errorf("Failed to check existing LDAP settings for account %d: %v", req.AccountID, err)
|
||||
c.JSON(http.StatusUnprocessableEntity, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrInternal,
|
||||
Message: "Failed to check existing LDAP configuration",
|
||||
Detail: err.Error(),
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
// Create new settings
|
||||
settings = model.AccountLDAPSettings{
|
||||
AccountID: req.AccountID,
|
||||
Host: req.Host,
|
||||
Port: req.Port,
|
||||
UseTLS: req.UseTLS,
|
||||
BaseDN: req.BaseDN,
|
||||
BindDN: req.BindDN,
|
||||
BindPassword: req.BindPassword,
|
||||
UserFilter: req.UserFilter,
|
||||
EmailAttribute: req.EmailAttribute,
|
||||
NameAttribute: req.NameAttribute,
|
||||
FirstNameAttribute: req.FirstNameAttribute,
|
||||
LastNameAttribute: req.LastNameAttribute,
|
||||
GroupAttribute: req.GroupAttribute,
|
||||
GroupFilter: req.GroupFilter,
|
||||
RoleMappings: req.RoleMappings,
|
||||
AutoProvision: req.AutoProvision,
|
||||
SyncInterval: req.SyncInterval,
|
||||
Active: req.Active,
|
||||
}
|
||||
} else {
|
||||
// Update existing settings
|
||||
settings.Host = req.Host
|
||||
settings.Port = req.Port
|
||||
settings.UseTLS = req.UseTLS
|
||||
settings.BaseDN = req.BaseDN
|
||||
settings.BindDN = req.BindDN
|
||||
settings.BindPassword = req.BindPassword
|
||||
settings.UserFilter = req.UserFilter
|
||||
settings.EmailAttribute = req.EmailAttribute
|
||||
settings.NameAttribute = req.NameAttribute
|
||||
settings.FirstNameAttribute = req.FirstNameAttribute
|
||||
settings.LastNameAttribute = req.LastNameAttribute
|
||||
settings.GroupAttribute = req.GroupAttribute
|
||||
settings.GroupFilter = req.GroupFilter
|
||||
settings.RoleMappings = req.RoleMappings
|
||||
settings.AutoProvision = req.AutoProvision
|
||||
settings.SyncInterval = req.SyncInterval
|
||||
settings.Active = req.Active
|
||||
}
|
||||
|
||||
// Save to DB
|
||||
if err := h.db.Save(&settings).Error; err != nil {
|
||||
applogger.L().Errorf("Failed to save LDAP settings for account %d: %v", req.AccountID, err)
|
||||
c.JSON(http.StatusUnprocessableEntity, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrInternal,
|
||||
Message: "Failed to save LDAP configuration",
|
||||
Detail: err.Error(),
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
applogger.L().Infof("LDAP settings updated for account %d (host=%s, port=%d, active=%v)", req.AccountID, req.Host, req.Port, req.Active)
|
||||
response.OK(c, settings)
|
||||
}
|
||||
|
||||
// RegisterLDAPRoutes sets up LDAP routes on a Gin router group.
|
||||
// Login route is PUBLIC — no AuthRequired middleware (LDAP login doesn't require existing JWT).
|
||||
// Config management routes require AuthMiddleware + admin role (enforced at router level).
|
||||
func RegisterLDAPRoutes(rg *gin.RouterGroup, handler *LDAPHandler) {
|
||||
ldapGroup := rg.Group("/ldap")
|
||||
{
|
||||
// Public route: LDAP login (no AuthRequired middleware)
|
||||
ldapGroup.POST("/login", handler.Login)
|
||||
|
||||
// Admin-only routes: config management + connectivity test
|
||||
// These are wired into the authenticated + admin router group externally,
|
||||
// so AuthMiddleware + admin role check is enforced at the router level.
|
||||
ldapGroup.POST("/test", handler.TestConnection)
|
||||
ldapGroup.GET("/config", handler.GetConfig)
|
||||
ldapGroup.PUT("/config", handler.UpdateConfig)
|
||||
}
|
||||
}
|
||||
@@ -1,327 +0,0 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/gochat/gochat/internal/auth"
|
||||
"github.com/gochat/gochat/pkg/response"
|
||||
)
|
||||
|
||||
// Reference: P2E §1.5 — MFA HTTP handlers
|
||||
// Maps to Chatwoot enterprise TwoFactorAuthController:
|
||||
// - enable → POST /api/v1/auth/mfa/enable (generates secret + QR URI)
|
||||
// - verify → POST /api/v1/auth/mfa/verify (validates TOTP code, enables MFA)
|
||||
// - disable → POST /api/v1/auth/mfa/disable (disables MFA after code verification)
|
||||
|
||||
// MFAHandler handles MFA (TOTP) HTTP endpoints.
|
||||
type MFAHandler struct {
|
||||
mfaService *auth.MFAService
|
||||
}
|
||||
|
||||
// NewMFAHandler creates a MFA handler with service dependency.
|
||||
func NewMFAHandler(mfaService *auth.MFAService) *MFAHandler {
|
||||
return &MFAHandler{
|
||||
mfaService: mfaService,
|
||||
}
|
||||
}
|
||||
|
||||
// --- Request/Response structs ---
|
||||
|
||||
// EnableMFARequest is the JSON body for MFA enablement initiation.
|
||||
type EnableMFARequest struct {
|
||||
// No body required — user_id comes from auth context
|
||||
}
|
||||
|
||||
// EnableMFAResponse is the JSON response for MFA enablement initiation.
|
||||
type EnableMFAResponse struct {
|
||||
TOTPSecret string `json:"totp_secret"` // base32 secret for manual entry
|
||||
QRURI string `json:"qr_uri"` // otpauth:// URI for QR code generation
|
||||
Message string `json:"message"`
|
||||
}
|
||||
|
||||
// VerifyMFARequest is the JSON body for MFA TOTP verification.
|
||||
type VerifyMFARequest struct {
|
||||
TOTPSecret string `json:"totp_secret" binding:"required"` // secret from enable step
|
||||
TOTPCode string `json:"totp_code" binding:"required"` // 6-digit code from authenticator app
|
||||
}
|
||||
|
||||
// DisableMFARequest is the JSON body for MFA disablement.
|
||||
type DisableMFARequest struct {
|
||||
TOTPCode string `json:"totp_code" binding:"required"` // current TOTP code for verification
|
||||
}
|
||||
|
||||
type profileMFAVerifyRequest struct {
|
||||
OTPCode string `json:"otp_code"`
|
||||
TOTPCode string `json:"totp_code"`
|
||||
}
|
||||
|
||||
type profileMFADisableRequest struct {
|
||||
Password string `json:"password"`
|
||||
OTPCode string `json:"otp_code"`
|
||||
BackupCode string `json:"backup_code"`
|
||||
}
|
||||
|
||||
// --- Handlers ---
|
||||
|
||||
// EnableMFA initiates MFA setup: generates a TOTP secret and QR URI.
|
||||
// POST /api/v1/auth/mfa/enable
|
||||
// Requires authentication — uses user_id from JWT context.
|
||||
// The user must verify a TOTP code before MFA is actually activated.
|
||||
func (h *MFAHandler) EnableMFA(c *gin.Context) {
|
||||
userID := c.GetUint("user_id")
|
||||
if userID == 0 {
|
||||
response.AbortWithStatusError(c, http.StatusUnauthorized, response.ErrUnauthorized, "Authentication required")
|
||||
return
|
||||
}
|
||||
|
||||
// Check if MFA is already enabled
|
||||
enabled, err := h.mfaService.IsMFAEnabled(userID)
|
||||
if err != nil {
|
||||
response.AbortWithStatusError(c, http.StatusUnprocessableEntity, response.ErrInternal, err.Error())
|
||||
return
|
||||
}
|
||||
if enabled {
|
||||
response.AbortWithStatusError(c, http.StatusConflict, response.ErrConflict, "MFA is already enabled for this user")
|
||||
return
|
||||
}
|
||||
|
||||
// Generate new TOTP secret + QR URI
|
||||
secret, qrURI, err := h.mfaService.GenerateTOTPSecret(userID)
|
||||
if err != nil {
|
||||
response.AbortWithStatusError(c, http.StatusUnprocessableEntity, response.ErrInternal, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
response.OK(c, EnableMFAResponse{
|
||||
TOTPSecret: secret,
|
||||
QRURI: qrURI,
|
||||
Message: "Scan QR code with your authenticator app, then verify with a TOTP code",
|
||||
})
|
||||
}
|
||||
|
||||
// VerifyMFA completes MFA setup: verifies TOTP code and enables MFA on the user.
|
||||
// POST /api/v1/auth/mfa/verify
|
||||
// Requires authentication — uses user_id from JWT context.
|
||||
// This is the second step: user provides secret + code from authenticator app.
|
||||
func (h *MFAHandler) VerifyMFA(c *gin.Context) {
|
||||
userID := c.GetUint("user_id")
|
||||
if userID == 0 {
|
||||
response.AbortWithStatusError(c, http.StatusUnauthorized, response.ErrUnauthorized, "Authentication required")
|
||||
return
|
||||
}
|
||||
|
||||
var req VerifyMFARequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrValidation, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// Validate the TOTP code against the provided secret
|
||||
cfg := auth.DefaultTOTPConfig()
|
||||
if !auth.ValidateTOTPCode(req.TOTPSecret, req.TOTPCode, cfg) {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "Invalid TOTP code, please try again")
|
||||
return
|
||||
}
|
||||
|
||||
// Enable TOTP on the user (stores secret in DB)
|
||||
if err := h.mfaService.EnableTOTP(userID, req.TOTPSecret); err != nil {
|
||||
response.AbortWithStatusError(c, http.StatusInternalServerError, response.ErrInternal, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
response.OK(c, gin.H{
|
||||
"message": "MFA enabled successfully",
|
||||
"mfa_enabled": true,
|
||||
})
|
||||
}
|
||||
|
||||
// DisableMFA disables MFA after verifying the current TOTP code.
|
||||
// POST /api/v1/auth/mfa/disable
|
||||
// Requires authentication — uses user_id from JWT context.
|
||||
func (h *MFAHandler) DisableMFA(c *gin.Context) {
|
||||
userID := c.GetUint("user_id")
|
||||
if userID == 0 {
|
||||
response.AbortWithStatusError(c, http.StatusUnauthorized, response.ErrUnauthorized, "Authentication required")
|
||||
return
|
||||
}
|
||||
|
||||
var req DisableMFARequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrValidation, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// Disable TOTP (requires valid current code for security)
|
||||
if err := h.mfaService.DisableTOTP(userID, req.TOTPCode); err != nil {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
response.OK(c, gin.H{
|
||||
"message": "MFA disabled successfully",
|
||||
"mfa_enabled": false,
|
||||
})
|
||||
}
|
||||
|
||||
// MFAStatus returns the current MFA status for the authenticated user.
|
||||
// GET /api/v1/auth/mfa/status
|
||||
func (h *MFAHandler) MFAStatus(c *gin.Context) {
|
||||
userID := c.GetUint("user_id")
|
||||
if userID == 0 {
|
||||
response.AbortWithStatusError(c, http.StatusUnauthorized, response.ErrUnauthorized, "Authentication required")
|
||||
return
|
||||
}
|
||||
|
||||
enabled, err := h.mfaService.IsMFAEnabled(userID)
|
||||
if err != nil {
|
||||
response.AbortWithStatusError(c, http.StatusInternalServerError, response.ErrInternal, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
response.OK(c, gin.H{
|
||||
"mfa_enabled": enabled,
|
||||
})
|
||||
}
|
||||
|
||||
// RegisterMFARoutes sets up MFA routes on a Gin router group.
|
||||
// These routes require authentication (AuthRequired middleware).
|
||||
func RegisterMFARoutes(rg *gin.RouterGroup, handler *MFAHandler) {
|
||||
mfaGroup := rg.Group("/auth/mfa")
|
||||
{
|
||||
mfaGroup.POST("/enable", handler.EnableMFA)
|
||||
mfaGroup.POST("/verify", handler.VerifyMFA)
|
||||
mfaGroup.POST("/disable", handler.DisableMFA)
|
||||
mfaGroup.GET("/status", handler.MFAStatus)
|
||||
mfaGroup.POST("/backup_codes", handler.BackupCodes)
|
||||
}
|
||||
}
|
||||
|
||||
// ProfileMFAStatus matches Chatwoot Profile::MfaController#show.
|
||||
func (h *MFAHandler) ProfileMFAStatus(c *gin.Context) {
|
||||
userID := getUserID(c)
|
||||
if userID == 0 {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "user not authenticated"})
|
||||
return
|
||||
}
|
||||
enabled, err := h.mfaService.IsMFAEnabled(userID)
|
||||
if err != nil {
|
||||
c.AbortWithStatusJSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
backupCodesGenerated, err := h.mfaService.BackupCodesGenerated(userID)
|
||||
if err != nil {
|
||||
c.AbortWithStatusJSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"feature_available": true,
|
||||
"enabled": enabled,
|
||||
"backup_codes_generated": backupCodesGenerated,
|
||||
})
|
||||
}
|
||||
|
||||
// ProfileEnableMFA matches Chatwoot Profile::MfaController#create.
|
||||
func (h *MFAHandler) ProfileEnableMFA(c *gin.Context) {
|
||||
userID := getUserID(c)
|
||||
if userID == 0 {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "user not authenticated"})
|
||||
return
|
||||
}
|
||||
enabled, err := h.mfaService.IsMFAEnabled(userID)
|
||||
if err != nil {
|
||||
c.AbortWithStatusJSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if enabled {
|
||||
c.AbortWithStatusJSON(http.StatusUnprocessableEntity, gin.H{"error": "MFA is already enabled"})
|
||||
return
|
||||
}
|
||||
secret, uri, err := h.mfaService.BeginTOTPSetup(userID)
|
||||
if err != nil {
|
||||
c.AbortWithStatusJSON(http.StatusUnprocessableEntity, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"provisioning_url": uri, "secret": secret})
|
||||
}
|
||||
|
||||
// ProfileVerifyMFA matches Chatwoot Profile::MfaController#verify.
|
||||
func (h *MFAHandler) ProfileVerifyMFA(c *gin.Context) {
|
||||
userID := getUserID(c)
|
||||
if userID == 0 {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "user not authenticated"})
|
||||
return
|
||||
}
|
||||
var req profileMFAVerifyRequest
|
||||
_ = c.ShouldBindJSON(&req)
|
||||
code := req.OTPCode
|
||||
if code == "" {
|
||||
code = req.TOTPCode
|
||||
}
|
||||
backupCodes, err := h.mfaService.VerifyAndActivateTOTP(userID, code)
|
||||
if err != nil {
|
||||
c.AbortWithStatusJSON(http.StatusUnprocessableEntity, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"enabled": true, "backup_codes": backupCodes})
|
||||
}
|
||||
|
||||
// ProfileDisableMFA matches Chatwoot Profile::MfaController#destroy.
|
||||
func (h *MFAHandler) ProfileDisableMFA(c *gin.Context) {
|
||||
userID := getUserID(c)
|
||||
if userID == 0 {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "user not authenticated"})
|
||||
return
|
||||
}
|
||||
var req profileMFADisableRequest
|
||||
_ = c.ShouldBindJSON(&req)
|
||||
if err := h.mfaService.DisableTOTPWithPassword(userID, req.Password, req.OTPCode, req.BackupCode); err != nil {
|
||||
c.AbortWithStatusJSON(http.StatusUnprocessableEntity, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"enabled": false})
|
||||
}
|
||||
|
||||
// ProfileBackupCodes matches Chatwoot Profile::MfaController#backup_codes.
|
||||
func (h *MFAHandler) ProfileBackupCodes(c *gin.Context) {
|
||||
userID := getUserID(c)
|
||||
if userID == 0 {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "user not authenticated"})
|
||||
return
|
||||
}
|
||||
var req profileMFAVerifyRequest
|
||||
_ = c.ShouldBindJSON(&req)
|
||||
code := req.OTPCode
|
||||
if code == "" {
|
||||
code = req.TOTPCode
|
||||
}
|
||||
valid, err := h.mfaService.VerifyTOTPCode(userID, code)
|
||||
if err != nil || !valid {
|
||||
c.AbortWithStatusJSON(http.StatusUnprocessableEntity, gin.H{"error": "invalid totp code"})
|
||||
return
|
||||
}
|
||||
codes, err := h.mfaService.GenerateBackupCodes(userID)
|
||||
if err != nil {
|
||||
c.AbortWithStatusJSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"backup_codes": codes})
|
||||
}
|
||||
|
||||
// BackupCodes generates one-time MFA backup codes.
|
||||
// POST /api/v1/profile/mfa/backup_codes or /api/v1/auth/mfa/backup_codes
|
||||
// Reference: Chatwoot MfaController#backup_codes
|
||||
func (h *MFAHandler) BackupCodes(c *gin.Context) {
|
||||
userID := getUserID(c)
|
||||
if userID == 0 {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "user not authenticated"})
|
||||
return
|
||||
}
|
||||
codes, err := h.mfaService.GenerateBackupCodes(userID)
|
||||
if err != nil {
|
||||
c.AbortWithStatusJSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"backup_codes": codes})
|
||||
}
|
||||
@@ -1,723 +0,0 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto/hmac"
|
||||
"crypto/sha1"
|
||||
"encoding/base32"
|
||||
"encoding/binary"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"math"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/suite"
|
||||
"gorm.io/driver/sqlite"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/logger"
|
||||
|
||||
"github.com/gochat/gochat/internal/auth"
|
||||
"github.com/gochat/gochat/internal/model"
|
||||
pkgcrypto "github.com/gochat/gochat/pkg/crypto"
|
||||
"github.com/gochat/gochat/pkg/response"
|
||||
)
|
||||
|
||||
// --- MFA Handler Test Suite ---
|
||||
// Uses real SQLite DB + real MFAService + httptest.
|
||||
|
||||
type MFAHandlerTestSuite struct {
|
||||
suite.Suite
|
||||
|
||||
db *gorm.DB
|
||||
router *gin.Engine
|
||||
handler *MFAHandler
|
||||
mfaService *auth.MFAService
|
||||
|
||||
user *model.User
|
||||
account *model.Account
|
||||
|
||||
userID uint
|
||||
accountID uint
|
||||
}
|
||||
|
||||
func TestMFAHandlerSuite(t *testing.T) {
|
||||
suite.Run(t, new(MFAHandlerTestSuite))
|
||||
}
|
||||
|
||||
func (s *MFAHandlerTestSuite) SetupSuite() {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
// Create in-memory SQLite DB
|
||||
db, err := gorm.Open(sqlite.Open("file::memory:"), &gorm.Config{
|
||||
Logger: logger.Default.LogMode(logger.Silent),
|
||||
})
|
||||
s.Require().NoError(err, "failed to open SQLite test database")
|
||||
|
||||
// Migrate models needed for MFA operations
|
||||
s.Require().NoError(db.AutoMigrate(
|
||||
&model.Account{},
|
||||
&model.User{},
|
||||
))
|
||||
|
||||
s.db = db
|
||||
|
||||
// Create real MFA service backed by the test DB
|
||||
s.mfaService = auth.NewMFAService(db)
|
||||
s.handler = NewMFAHandler(s.mfaService)
|
||||
|
||||
// Create test account
|
||||
account := &model.Account{Name: "MFATestAccount"}
|
||||
s.Require().NoError(db.Create(account).Error)
|
||||
s.account = account
|
||||
s.accountID = account.ID
|
||||
|
||||
// Create test user belonging to the account
|
||||
user := &model.User{
|
||||
AccountID: account.ID,
|
||||
Name: "MFA Test User",
|
||||
Email: "mfatest@example.com",
|
||||
Password: "hashedpassword123",
|
||||
Provider: "email",
|
||||
Role: "agent",
|
||||
Active: true,
|
||||
}
|
||||
s.Require().NoError(db.Create(user).Error)
|
||||
s.user = user
|
||||
s.userID = user.ID
|
||||
|
||||
// Setup router with middleware that injects user_id into context
|
||||
s.setupRouter(s.userID)
|
||||
}
|
||||
|
||||
func (s *MFAHandlerTestSuite) setupRouter(userID uint) {
|
||||
r := gin.New()
|
||||
r.Use(func(c *gin.Context) {
|
||||
c.Set("user_id", userID)
|
||||
c.Next()
|
||||
})
|
||||
|
||||
mfaGroup := r.Group("/api/v1/auth/mfa")
|
||||
{
|
||||
mfaGroup.POST("/enable", s.handler.EnableMFA)
|
||||
mfaGroup.POST("/verify", s.handler.VerifyMFA)
|
||||
mfaGroup.POST("/disable", s.handler.DisableMFA)
|
||||
}
|
||||
r.GET("/api/v1/profile/mfa", s.handler.ProfileMFAStatus)
|
||||
r.POST("/api/v1/profile/mfa", s.handler.ProfileEnableMFA)
|
||||
r.DELETE("/api/v1/profile/mfa", s.handler.ProfileDisableMFA)
|
||||
r.POST("/api/v1/profile/mfa/verify", s.handler.ProfileVerifyMFA)
|
||||
r.POST("/api/v1/profile/mfa/backup_codes", s.handler.ProfileBackupCodes)
|
||||
|
||||
s.router = r
|
||||
}
|
||||
|
||||
func (s *MFAHandlerTestSuite) SetupTest() {
|
||||
// Hard cleanup: delete all users and accounts, then recreate
|
||||
s.db.Exec("DELETE FROM users")
|
||||
s.db.Exec("DELETE FROM accounts")
|
||||
|
||||
// Recreate test data
|
||||
account := &model.Account{Name: "MFATestAccount"}
|
||||
s.Require().NoError(s.db.Create(account).Error)
|
||||
s.account = account
|
||||
s.accountID = account.ID
|
||||
|
||||
user := &model.User{
|
||||
AccountID: account.ID,
|
||||
Name: "MFA Test User",
|
||||
Email: "mfatest@example.com",
|
||||
Password: "hashedpassword123",
|
||||
Provider: "email",
|
||||
Role: "agent",
|
||||
Active: true,
|
||||
}
|
||||
s.Require().NoError(s.db.Create(user).Error)
|
||||
s.user = user
|
||||
s.userID = user.ID
|
||||
|
||||
// Re-setup router with the new user ID
|
||||
s.setupRouter(s.userID)
|
||||
}
|
||||
|
||||
func (s *MFAHandlerTestSuite) TearDownSuite() {
|
||||
if s.db != nil {
|
||||
sqlDB, err := s.db.DB()
|
||||
if err == nil {
|
||||
sqlDB.Close()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// --- Helper to make requests and parse responses ---
|
||||
|
||||
func (s *MFAHandlerTestSuite) doRequest(method, path, body string) *httptest.ResponseRecorder {
|
||||
var reqBody *bytes.Buffer
|
||||
if body != "" {
|
||||
reqBody = bytes.NewBufferString(body)
|
||||
} else {
|
||||
reqBody = bytes.NewBufferString("")
|
||||
}
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(method, path, reqBody)
|
||||
if body != "" {
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
}
|
||||
s.router.ServeHTTP(w, req)
|
||||
return w
|
||||
}
|
||||
|
||||
func (s *MFAHandlerTestSuite) parseResponse(w *httptest.ResponseRecorder) response.APIResponse {
|
||||
var resp response.APIResponse
|
||||
s.NoError(json.Unmarshal(w.Body.Bytes(), &resp))
|
||||
return resp
|
||||
}
|
||||
|
||||
// --- Helper to generate a valid TOTP code for a secret ---
|
||||
// Uses the same algorithm as auth.validateTOTP/generateTOTP to compute a valid code.
|
||||
|
||||
func (s *MFAHandlerTestSuite) generateValidTOTPCode(secret string) string {
|
||||
cfg := auth.DefaultTOTPConfig()
|
||||
return computeTOTPCode(secret, cfg)
|
||||
}
|
||||
|
||||
func computeTOTPCode(secret string, cfg auth.TOTPConfig) string {
|
||||
key, err := decodeBase32NoPad(secret)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
|
||||
now := time.Now().Unix()
|
||||
period := int64(cfg.Period)
|
||||
timeCounter := now / period
|
||||
|
||||
return generateTOTPFromKey(key, timeCounter, cfg)
|
||||
}
|
||||
|
||||
func decodeBase32NoPad(secret string) ([]byte, error) {
|
||||
return base32.StdEncoding.WithPadding(base32.NoPadding).DecodeString(strings.ToUpper(secret))
|
||||
}
|
||||
|
||||
func generateTOTPFromKey(key []byte, timeCounter int64, cfg auth.TOTPConfig) string {
|
||||
// Encode time counter as 8-byte big-endian
|
||||
buf := make([]byte, 8)
|
||||
binary.BigEndian.PutUint64(buf, uint64(timeCounter))
|
||||
|
||||
// HMAC-SHA1
|
||||
h := hmac.New(sha1.New, key)
|
||||
h.Write(buf)
|
||||
hash := h.Sum(nil)
|
||||
|
||||
// Dynamic truncation per RFC 4226
|
||||
offset := hash[len(hash)-1] & 0x0f
|
||||
truncated := (int32(hash[offset]&0x7f) << 24) |
|
||||
(int32(hash[offset+1]&0xff) << 16) |
|
||||
(int32(hash[offset+2]&0xff) << 8) |
|
||||
(int32(hash[offset+3] & 0xff))
|
||||
|
||||
// Modulo 10^digits
|
||||
mod := int32(math.Pow10(cfg.Digits))
|
||||
code := truncated % mod
|
||||
|
||||
// Format with leading zeros
|
||||
return fmt.Sprintf("%0*d", cfg.Digits, code)
|
||||
}
|
||||
|
||||
func jsonBody(data map[string]interface{}) string {
|
||||
b, err := json.Marshal(data)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return string(b)
|
||||
}
|
||||
|
||||
// ============================================================
|
||||
// EnableMFA tests
|
||||
// ============================================================
|
||||
|
||||
func (s *MFAHandlerTestSuite) TestEnable_Success() {
|
||||
w := s.doRequest(http.MethodPost, "/api/v1/auth/mfa/enable", "{}")
|
||||
|
||||
s.Equal(http.StatusOK, w.Code)
|
||||
|
||||
resp := s.parseResponse(w)
|
||||
s.True(resp.Success)
|
||||
|
||||
dataMap, ok := resp.Data.(map[string]interface{})
|
||||
s.True(ok)
|
||||
|
||||
// Response should contain totp_secret and qr_uri
|
||||
s.NotEmpty(dataMap["totp_secret"])
|
||||
s.NotEmpty(dataMap["qr_uri"])
|
||||
s.Contains(dataMap["qr_uri"], "otpauth://totp/")
|
||||
s.Contains(dataMap["qr_uri"], dataMap["totp_secret"])
|
||||
}
|
||||
|
||||
func (s *MFAHandlerTestSuite) TestEnable_NoBody() {
|
||||
// EnableMFARequest has no required fields, empty body should still work
|
||||
// since the handler doesn't even call ShouldBindJSON
|
||||
w := s.doRequest(http.MethodPost, "/api/v1/auth/mfa/enable", "")
|
||||
|
||||
// With empty body and no Content-Type, the handler doesn't bind JSON,
|
||||
// so it just uses user_id from context → should succeed
|
||||
s.Equal(http.StatusOK, w.Code)
|
||||
|
||||
resp := s.parseResponse(w)
|
||||
s.True(resp.Success)
|
||||
}
|
||||
|
||||
func (s *MFAHandlerTestSuite) TestEnable_InvalidJSON() {
|
||||
// Enable handler does NOT call ShouldBindJSON at all — it only uses
|
||||
// c.GetUint("user_id") and service calls. So invalid JSON in the body
|
||||
// won't cause a binding error. With a valid user_id in context,
|
||||
// this should succeed regardless of body content.
|
||||
w := s.doRequest(http.MethodPost, "/api/v1/auth/mfa/enable", "{invalid}")
|
||||
|
||||
// The handler ignores the body entirely, so with valid user_id it succeeds
|
||||
s.Equal(http.StatusOK, w.Code)
|
||||
}
|
||||
|
||||
func (s *MFAHandlerTestSuite) TestEnable_Unauthorized_NoUserID() {
|
||||
// Create router without user_id middleware — user_id will be 0
|
||||
r := gin.New()
|
||||
r.POST("/api/v1/auth/mfa/enable", s.handler.EnableMFA)
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/auth/mfa/enable", bytes.NewBufferString("{}"))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
s.Equal(http.StatusUnauthorized, w.Code)
|
||||
|
||||
var respStruct response.APIResponse
|
||||
s.NoError(json.Unmarshal(w.Body.Bytes(), &respStruct))
|
||||
s.False(respStruct.Success)
|
||||
s.NotNil(respStruct.Error)
|
||||
s.Equal(response.ErrUnauthorized, respStruct.Error.Code)
|
||||
}
|
||||
|
||||
func (s *MFAHandlerTestSuite) TestEnable_UserNotFound() {
|
||||
// Router that sets a non-existent user ID
|
||||
r := gin.New()
|
||||
r.Use(func(c *gin.Context) {
|
||||
c.Set("user_id", uint(9999)) // non-existent user
|
||||
c.Next()
|
||||
})
|
||||
r.POST("/api/v1/auth/mfa/enable", s.handler.EnableMFA)
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/auth/mfa/enable", bytes.NewBufferString("{}"))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
// IsMFAEnabled will fail to find the user → 500 Internal Server Error
|
||||
s.Equal(http.StatusUnprocessableEntity, w.Code)
|
||||
|
||||
var respStruct response.APIResponse
|
||||
s.NoError(json.Unmarshal(w.Body.Bytes(), &respStruct))
|
||||
s.False(respStruct.Success)
|
||||
s.NotNil(respStruct.Error)
|
||||
s.Equal(response.ErrInternal, respStruct.Error.Code)
|
||||
}
|
||||
|
||||
func (s *MFAHandlerTestSuite) TestEnable_AlreadyEnabled() {
|
||||
// First enable MFA for the user
|
||||
user := s.user
|
||||
user.TOTPSecret = "JBSWY3DPEHPK3PXP"
|
||||
user.TOTPEnabled = true
|
||||
s.Require().NoError(s.db.Save(user).Error)
|
||||
|
||||
w := s.doRequest(http.MethodPost, "/api/v1/auth/mfa/enable", "{}")
|
||||
|
||||
s.Equal(http.StatusConflict, w.Code)
|
||||
|
||||
resp := s.parseResponse(w)
|
||||
s.False(resp.Success)
|
||||
s.NotNil(resp.Error)
|
||||
s.Equal(response.ErrConflict, resp.Error.Code)
|
||||
}
|
||||
|
||||
// ============================================================
|
||||
// VerifyMFA tests
|
||||
// ============================================================
|
||||
|
||||
func (s *MFAHandlerTestSuite) TestVerify_Success() {
|
||||
// Step 1: Generate secret
|
||||
secret, _, err := s.mfaService.GenerateTOTPSecret(s.userID)
|
||||
s.Require().NoError(err)
|
||||
|
||||
// Step 2: Compute a valid TOTP code for the secret
|
||||
code := s.generateValidTOTPCode(secret)
|
||||
|
||||
// Step 3: Verify with secret + code
|
||||
body := jsonBody(map[string]interface{}{
|
||||
"totp_secret": secret,
|
||||
"totp_code": code,
|
||||
})
|
||||
|
||||
w := s.doRequest(http.MethodPost, "/api/v1/auth/mfa/verify", body)
|
||||
|
||||
s.Equal(http.StatusOK, w.Code)
|
||||
|
||||
resp := s.parseResponse(w)
|
||||
s.True(resp.Success)
|
||||
|
||||
dataMap, ok := resp.Data.(map[string]interface{})
|
||||
s.True(ok)
|
||||
s.Equal("MFA enabled successfully", dataMap["message"])
|
||||
s.Equal(true, dataMap["mfa_enabled"])
|
||||
|
||||
// Verify that TOTPEnabled is now true in DB
|
||||
var updatedUser model.User
|
||||
s.Require().NoError(s.db.First(&updatedUser, s.userID).Error)
|
||||
s.True(updatedUser.TOTPEnabled)
|
||||
s.Equal(secret, updatedUser.TOTPSecret)
|
||||
}
|
||||
|
||||
func (s *MFAHandlerTestSuite) TestVerify_InvalidJSON() {
|
||||
w := s.doRequest(http.MethodPost, "/api/v1/auth/mfa/verify", "{invalid}")
|
||||
|
||||
s.Equal(http.StatusBadRequest, w.Code)
|
||||
|
||||
resp := s.parseResponse(w)
|
||||
s.False(resp.Success)
|
||||
s.NotNil(resp.Error)
|
||||
s.Equal(response.ErrValidation, resp.Error.Code)
|
||||
}
|
||||
|
||||
func (s *MFAHandlerTestSuite) TestVerify_MissingTOTPSecret() {
|
||||
// When totp_secret is missing from JSON, ShouldBindJSON fails with
|
||||
// binding:"required" validation error → handler returns ErrValidation
|
||||
body := jsonBody(map[string]interface{}{
|
||||
"totp_code": "123456",
|
||||
})
|
||||
|
||||
w := s.doRequest(http.MethodPost, "/api/v1/auth/mfa/verify", body)
|
||||
|
||||
// Missing required binding field → ShouldBindJSON error → 400 VALIDATION_ERROR
|
||||
s.Equal(http.StatusBadRequest, w.Code)
|
||||
|
||||
resp := s.parseResponse(w)
|
||||
s.False(resp.Success)
|
||||
s.NotNil(resp.Error)
|
||||
s.Equal(response.ErrValidation, resp.Error.Code)
|
||||
}
|
||||
|
||||
func (s *MFAHandlerTestSuite) TestVerify_MissingTOTPCode() {
|
||||
secret, _, err := s.mfaService.GenerateTOTPSecret(s.userID)
|
||||
s.Require().NoError(err)
|
||||
|
||||
// When totp_code is missing from JSON, ShouldBindJSON fails with
|
||||
// binding:"required" validation error → handler returns ErrValidation
|
||||
body := jsonBody(map[string]interface{}{
|
||||
"totp_secret": secret,
|
||||
})
|
||||
|
||||
w := s.doRequest(http.MethodPost, "/api/v1/auth/mfa/verify", body)
|
||||
|
||||
// Missing required binding field → ShouldBindJSON error → 400 VALIDATION_ERROR
|
||||
s.Equal(http.StatusBadRequest, w.Code)
|
||||
|
||||
resp := s.parseResponse(w)
|
||||
s.False(resp.Success)
|
||||
s.NotNil(resp.Error)
|
||||
s.Equal(response.ErrValidation, resp.Error.Code)
|
||||
}
|
||||
|
||||
func (s *MFAHandlerTestSuite) TestVerify_InvalidTOTPCode() {
|
||||
secret, _, err := s.mfaService.GenerateTOTPSecret(s.userID)
|
||||
s.Require().NoError(err)
|
||||
|
||||
body := jsonBody(map[string]interface{}{
|
||||
"totp_secret": secret,
|
||||
"totp_code": "000000", // definitely wrong code
|
||||
})
|
||||
|
||||
w := s.doRequest(http.MethodPost, "/api/v1/auth/mfa/verify", body)
|
||||
|
||||
// ValidateTOTPCode returns false → 400 Bad Request
|
||||
s.Equal(http.StatusBadRequest, w.Code)
|
||||
|
||||
resp := s.parseResponse(w)
|
||||
s.False(resp.Success)
|
||||
s.NotNil(resp.Error)
|
||||
s.Equal(response.ErrBadRequest, resp.Error.Code)
|
||||
}
|
||||
|
||||
func (s *MFAHandlerTestSuite) TestVerify_Unauthorized_NoUserID() {
|
||||
r := gin.New()
|
||||
r.POST("/api/v1/auth/mfa/verify", s.handler.VerifyMFA)
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/auth/mfa/verify", bytes.NewBufferString(`{"totp_secret":"abc","totp_code":"123456"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
s.Equal(http.StatusUnauthorized, w.Code)
|
||||
|
||||
var respStruct response.APIResponse
|
||||
s.NoError(json.Unmarshal(w.Body.Bytes(), &respStruct))
|
||||
s.False(respStruct.Success)
|
||||
s.Equal(response.ErrUnauthorized, respStruct.Error.Code)
|
||||
}
|
||||
|
||||
// ============================================================
|
||||
// DisableMFA tests
|
||||
// ============================================================
|
||||
|
||||
func (s *MFAHandlerTestSuite) TestDisable_Success() {
|
||||
// First, enable MFA for the user so we can disable it
|
||||
secret, _, err := s.mfaService.GenerateTOTPSecret(s.userID)
|
||||
s.Require().NoError(err)
|
||||
|
||||
// Enable TOTP via service directly
|
||||
s.Require().NoError(s.mfaService.EnableTOTP(s.userID, secret))
|
||||
|
||||
// Now the user has TOTP enabled. Generate a current valid code for disable.
|
||||
disableCode := s.generateValidTOTPCode(secret)
|
||||
|
||||
body := jsonBody(map[string]interface{}{
|
||||
"totp_code": disableCode,
|
||||
})
|
||||
|
||||
w := s.doRequest(http.MethodPost, "/api/v1/auth/mfa/disable", body)
|
||||
|
||||
s.Equal(http.StatusOK, w.Code)
|
||||
|
||||
resp := s.parseResponse(w)
|
||||
s.True(resp.Success)
|
||||
|
||||
dataMap, ok := resp.Data.(map[string]interface{})
|
||||
s.True(ok)
|
||||
s.Equal("MFA disabled successfully", dataMap["message"])
|
||||
s.Equal(false, dataMap["mfa_enabled"])
|
||||
|
||||
// Verify that TOTPEnabled is now false in DB
|
||||
var updatedUser model.User
|
||||
s.Require().NoError(s.db.First(&updatedUser, s.userID).Error)
|
||||
s.False(updatedUser.TOTPEnabled)
|
||||
s.Empty(updatedUser.TOTPSecret)
|
||||
}
|
||||
|
||||
func (s *MFAHandlerTestSuite) TestDisable_InvalidJSON() {
|
||||
w := s.doRequest(http.MethodPost, "/api/v1/auth/mfa/disable", "{invalid}")
|
||||
|
||||
s.Equal(http.StatusBadRequest, w.Code)
|
||||
|
||||
resp := s.parseResponse(w)
|
||||
s.False(resp.Success)
|
||||
s.NotNil(resp.Error)
|
||||
s.Equal(response.ErrValidation, resp.Error.Code)
|
||||
}
|
||||
|
||||
func (s *MFAHandlerTestSuite) TestDisable_MissingTOTPCode() {
|
||||
// First enable MFA
|
||||
secret, _, err := s.mfaService.GenerateTOTPSecret(s.userID)
|
||||
s.Require().NoError(err)
|
||||
s.Require().NoError(s.mfaService.EnableTOTP(s.userID, secret))
|
||||
|
||||
// When totp_code is missing from JSON, ShouldBindJSON fails with
|
||||
// binding:"required" validation error → handler returns ErrValidation
|
||||
body := jsonBody(map[string]interface{}{})
|
||||
|
||||
w := s.doRequest(http.MethodPost, "/api/v1/auth/mfa/disable", body)
|
||||
|
||||
// Missing required binding field → ShouldBindJSON error → 400 VALIDATION_ERROR
|
||||
s.Equal(http.StatusBadRequest, w.Code)
|
||||
|
||||
resp := s.parseResponse(w)
|
||||
s.False(resp.Success)
|
||||
s.NotNil(resp.Error)
|
||||
s.Equal(response.ErrValidation, resp.Error.Code)
|
||||
}
|
||||
|
||||
func (s *MFAHandlerTestSuite) TestDisable_InvalidTOTPCode() {
|
||||
// First enable MFA
|
||||
secret, _, err := s.mfaService.GenerateTOTPSecret(s.userID)
|
||||
s.Require().NoError(err)
|
||||
s.Require().NoError(s.mfaService.EnableTOTP(s.userID, secret))
|
||||
|
||||
body := jsonBody(map[string]interface{}{
|
||||
"totp_code": "000000", // definitely wrong
|
||||
})
|
||||
|
||||
w := s.doRequest(http.MethodPost, "/api/v1/auth/mfa/disable", body)
|
||||
|
||||
// DisableTOTP → VerifyTOTPCode → invalid code → error → 400 Bad Request
|
||||
s.Equal(http.StatusBadRequest, w.Code)
|
||||
|
||||
resp := s.parseResponse(w)
|
||||
s.False(resp.Success)
|
||||
s.NotNil(resp.Error)
|
||||
s.Equal(response.ErrBadRequest, resp.Error.Code)
|
||||
}
|
||||
|
||||
func (s *MFAHandlerTestSuite) TestDisable_Unauthorized_NoUserID() {
|
||||
r := gin.New()
|
||||
r.POST("/api/v1/auth/mfa/disable", s.handler.DisableMFA)
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/v1/auth/mfa/disable", bytes.NewBufferString(`{"totp_code":"123456"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
s.Equal(http.StatusUnauthorized, w.Code)
|
||||
|
||||
var respStruct response.APIResponse
|
||||
s.NoError(json.Unmarshal(w.Body.Bytes(), &respStruct))
|
||||
s.False(respStruct.Success)
|
||||
s.Equal(response.ErrUnauthorized, respStruct.Error.Code)
|
||||
}
|
||||
|
||||
func (s *MFAHandlerTestSuite) TestDisable_MFANotEnabled() {
|
||||
// User does not have MFA enabled — DisableTOTP calls VerifyTOTPCode
|
||||
// which checks user.TOTPEnabled == false → error
|
||||
body := jsonBody(map[string]interface{}{
|
||||
"totp_code": "123456",
|
||||
})
|
||||
|
||||
w := s.doRequest(http.MethodPost, "/api/v1/auth/mfa/disable", body)
|
||||
|
||||
// VerifyTOTPCode will return error "mfa not enabled for user" → 400 Bad Request
|
||||
s.Equal(http.StatusBadRequest, w.Code)
|
||||
|
||||
resp := s.parseResponse(w)
|
||||
s.False(resp.Success)
|
||||
s.NotNil(resp.Error)
|
||||
s.Equal(response.ErrBadRequest, resp.Error.Code)
|
||||
}
|
||||
|
||||
// ============================================================
|
||||
// MFAStatus tests (bonus coverage for the status endpoint)
|
||||
// ============================================================
|
||||
|
||||
func (s *MFAHandlerTestSuite) TestStatus_MFADisabled() {
|
||||
// Setup router with status route
|
||||
r := gin.New()
|
||||
r.Use(func(c *gin.Context) {
|
||||
c.Set("user_id", s.userID)
|
||||
c.Next()
|
||||
})
|
||||
r.GET("/api/v1/auth/mfa/status", s.handler.MFAStatus)
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/v1/auth/mfa/status", nil)
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
s.Equal(http.StatusOK, w.Code)
|
||||
|
||||
var respStruct response.APIResponse
|
||||
s.NoError(json.Unmarshal(w.Body.Bytes(), &respStruct))
|
||||
s.True(respStruct.Success)
|
||||
|
||||
dataMap, ok := respStruct.Data.(map[string]interface{})
|
||||
s.True(ok)
|
||||
s.Equal(false, dataMap["mfa_enabled"])
|
||||
}
|
||||
|
||||
func (s *MFAHandlerTestSuite) TestStatus_MFAEnabled() {
|
||||
// Enable MFA first
|
||||
secret, _, err := s.mfaService.GenerateTOTPSecret(s.userID)
|
||||
s.Require().NoError(err)
|
||||
s.Require().NoError(s.mfaService.EnableTOTP(s.userID, secret))
|
||||
|
||||
// Setup router with status route
|
||||
r := gin.New()
|
||||
r.Use(func(c *gin.Context) {
|
||||
c.Set("user_id", s.userID)
|
||||
c.Next()
|
||||
})
|
||||
r.GET("/api/v1/auth/mfa/status", s.handler.MFAStatus)
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/v1/auth/mfa/status", nil)
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
s.Equal(http.StatusOK, w.Code)
|
||||
|
||||
var respStruct response.APIResponse
|
||||
s.NoError(json.Unmarshal(w.Body.Bytes(), &respStruct))
|
||||
s.True(respStruct.Success)
|
||||
|
||||
dataMap, ok := respStruct.Data.(map[string]interface{})
|
||||
s.True(ok)
|
||||
s.Equal(true, dataMap["mfa_enabled"])
|
||||
}
|
||||
|
||||
func (s *MFAHandlerTestSuite) TestStatus_Unauthorized_NoUserID() {
|
||||
r := gin.New()
|
||||
r.GET("/api/v1/auth/mfa/status", s.handler.MFAStatus)
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/v1/auth/mfa/status", nil)
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
s.Equal(http.StatusUnauthorized, w.Code)
|
||||
|
||||
var respStruct response.APIResponse
|
||||
s.NoError(json.Unmarshal(w.Body.Bytes(), &respStruct))
|
||||
s.False(respStruct.Success)
|
||||
s.Equal(response.ErrUnauthorized, respStruct.Error.Code)
|
||||
}
|
||||
|
||||
func (s *MFAHandlerTestSuite) TestProfileMFA_StatusUsesChatwootRawPayload() {
|
||||
w := s.doRequest(http.MethodGet, "/api/v1/profile/mfa", "")
|
||||
s.Equal(http.StatusOK, w.Code)
|
||||
|
||||
var payload map[string]interface{}
|
||||
s.Require().NoError(json.Unmarshal(w.Body.Bytes(), &payload))
|
||||
s.Equal(true, payload["feature_available"])
|
||||
s.Equal(false, payload["enabled"])
|
||||
s.Equal(false, payload["backup_codes_generated"])
|
||||
s.NotContains(payload, "data")
|
||||
}
|
||||
|
||||
func (s *MFAHandlerTestSuite) TestProfileMFA_EnableVerifyBackupAndDisableUseFrontendPayloads() {
|
||||
passwordHash, err := pkgcrypto.HashPassword("current-password")
|
||||
s.Require().NoError(err)
|
||||
s.Require().NoError(s.db.Model(&model.User{}).Where("id = ?", s.userID).Updates(map[string]interface{}{
|
||||
"password": passwordHash,
|
||||
"password_digest": passwordHash,
|
||||
}).Error)
|
||||
|
||||
enableRec := s.doRequest(http.MethodPost, "/api/v1/profile/mfa", "")
|
||||
s.Equal(http.StatusOK, enableRec.Code)
|
||||
var enablePayload map[string]string
|
||||
s.Require().NoError(json.Unmarshal(enableRec.Body.Bytes(), &enablePayload))
|
||||
s.NotEmpty(enablePayload["secret"])
|
||||
s.Contains(enablePayload["provisioning_url"], "otpauth://totp/")
|
||||
|
||||
code := s.generateValidTOTPCode(enablePayload["secret"])
|
||||
verifyRec := s.doRequest(http.MethodPost, "/api/v1/profile/mfa/verify", jsonBody(map[string]interface{}{"otp_code": code}))
|
||||
s.Equal(http.StatusOK, verifyRec.Code)
|
||||
var verifyPayload struct {
|
||||
Enabled bool `json:"enabled"`
|
||||
BackupCodes []string `json:"backup_codes"`
|
||||
}
|
||||
s.Require().NoError(json.Unmarshal(verifyRec.Body.Bytes(), &verifyPayload))
|
||||
s.True(verifyPayload.Enabled)
|
||||
s.Len(verifyPayload.BackupCodes, 10)
|
||||
|
||||
backupRec := s.doRequest(http.MethodPost, "/api/v1/profile/mfa/backup_codes", jsonBody(map[string]interface{}{"otp_code": code}))
|
||||
s.Equal(http.StatusOK, backupRec.Code)
|
||||
var backupPayload struct {
|
||||
BackupCodes []string `json:"backup_codes"`
|
||||
}
|
||||
s.Require().NoError(json.Unmarshal(backupRec.Body.Bytes(), &backupPayload))
|
||||
s.Len(backupPayload.BackupCodes, 10)
|
||||
|
||||
disableRec := s.doRequest(http.MethodDelete, "/api/v1/profile/mfa", jsonBody(map[string]interface{}{"password": "current-password", "otp_code": code}))
|
||||
s.Equal(http.StatusOK, disableRec.Code)
|
||||
var disablePayload map[string]bool
|
||||
s.Require().NoError(json.Unmarshal(disableRec.Body.Bytes(), &disablePayload))
|
||||
s.False(disablePayload["enabled"])
|
||||
}
|
||||
|
||||
// Ensure unused import warning doesn't cause issues
|
||||
var _ = assert.Equal
|
||||
@@ -1,431 +0,0 @@
|
||||
package v1
|
||||
|
||||
// Reference: P2E §1.6 — SAML 2.0 SP HTTP handlers
|
||||
// Provides three endpoints for SAML SSO integration:
|
||||
// - GET /api/v1/saml/metadata → SP metadata (for IdP config import)
|
||||
// - GET /api/v1/saml/login → Initiate SP-initiated SSO (redirect to IdP)
|
||||
// - POST /api/v1/saml/acs → ACS endpoint (process IdP response, issue JWT)
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/gochat/gochat/internal/auth"
|
||||
"github.com/gochat/gochat/internal/config"
|
||||
"github.com/gochat/gochat/pkg/response"
|
||||
applogger "github.com/gochat/gochat/pkg/logger"
|
||||
)
|
||||
|
||||
// SAMLHandler handles SAML 2.0 authentication HTTP endpoints.
|
||||
type SAMLHandler struct {
|
||||
samlService *auth.SAMLService
|
||||
jwtService *auth.JWTService
|
||||
refreshStore *auth.RefreshTokenStore
|
||||
ssoSessionStore *auth.SSOSessionStore
|
||||
samlCfg *config.SAMLConfig
|
||||
}
|
||||
|
||||
// NewSAMLHandler creates a SAML handler with service dependencies.
|
||||
func NewSAMLHandler(
|
||||
samlService *auth.SAMLService,
|
||||
jwtService *auth.JWTService,
|
||||
refreshStore *auth.RefreshTokenStore,
|
||||
ssoSessionStore *auth.SSOSessionStore,
|
||||
samlCfg *config.SAMLConfig,
|
||||
) *SAMLHandler {
|
||||
return &SAMLHandler{
|
||||
samlService: samlService,
|
||||
jwtService: jwtService,
|
||||
refreshStore: refreshStore,
|
||||
ssoSessionStore: ssoSessionStore,
|
||||
samlCfg: samlCfg,
|
||||
}
|
||||
}
|
||||
|
||||
// Metadata returns the SP XML metadata for IdP administrators to import.
|
||||
// GET /api/v1/saml/metadata
|
||||
func (h *SAMLHandler) Metadata(c *gin.Context) {
|
||||
if !h.samlCfg.Enabled {
|
||||
c.JSON(http.StatusNotFound, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrNotFound,
|
||||
Message: "SAML is not enabled",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
xml, err := h.samlService.GetSPMetadata()
|
||||
if err != nil {
|
||||
applogger.L().Errorf("Failed to generate SAML SP metadata: %v", err)
|
||||
c.JSON(http.StatusUnprocessableEntity, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrInternal,
|
||||
Message: "Failed to generate SP metadata",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Return raw XML with appropriate content type
|
||||
c.Data(http.StatusOK, "application/samlmetadata+xml", xml)
|
||||
}
|
||||
|
||||
// Login initiates SP-initiated SSO by redirecting to the IdP.
|
||||
// GET /api/v1/saml/login
|
||||
func (h *SAMLHandler) Login(c *gin.Context) {
|
||||
if !h.samlCfg.Enabled {
|
||||
c.JSON(http.StatusNotFound, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrNotFound,
|
||||
Message: "SAML is not enabled",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Generate state token for CSRF protection (same pattern as OAuth)
|
||||
state := generateOAuthState()
|
||||
|
||||
redirectURL, err := h.samlService.InitiateLogin(state)
|
||||
if err != nil {
|
||||
applogger.L().Errorf("Failed to initiate SAML login: %v", err)
|
||||
c.JSON(http.StatusUnprocessableEntity, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrInternal,
|
||||
Message: "Failed to initiate SAML login",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Redirect user to IdP
|
||||
c.Redirect(http.StatusFound, redirectURL)
|
||||
}
|
||||
|
||||
// ACS (Assertion Consumer Service) processes the SAML Response from the IdP.
|
||||
// POST /api/v1/saml/acs
|
||||
// The IdP posts a base64-encoded SAMLResponse + RelayState to this endpoint.
|
||||
// On success: validates assertion, finds/creates user, issues JWT token pair.
|
||||
func (h *SAMLHandler) ACS(c *gin.Context) {
|
||||
if !h.samlCfg.Enabled {
|
||||
c.JSON(http.StatusNotFound, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrNotFound,
|
||||
Message: "SAML is not enabled",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Extract SAMLResponse from form POST (IdP sends as base64-encoded form param)
|
||||
samlResponse := c.PostForm("SAMLResponse")
|
||||
if samlResponse == "" {
|
||||
c.JSON(http.StatusBadRequest, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrBadRequest,
|
||||
Message: "Missing SAMLResponse parameter",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Process and validate the SAML response
|
||||
userInfo, err := h.samlService.ProcessResponse(samlResponse)
|
||||
if err != nil {
|
||||
applogger.L().Errorf("SAML ACS validation failed: %v", err)
|
||||
statusCode := http.StatusUnauthorized
|
||||
errCode := response.ErrUnauthorized
|
||||
if errors.Is(err, auth.ErrSAMLReplay) {
|
||||
statusCode = http.StatusForbidden
|
||||
errCode = response.ErrForbidden
|
||||
} else if errors.Is(err, auth.ErrSAMLInvalidResponse) {
|
||||
statusCode = http.StatusUnauthorized
|
||||
}
|
||||
c.JSON(statusCode, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: errCode,
|
||||
Message: "SAML authentication failed",
|
||||
Detail: err.Error(),
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Find or create user in GoChat
|
||||
user, err := h.samlService.FindOrCreateUser(userInfo)
|
||||
if err != nil {
|
||||
applogger.L().Errorf("SAML user lookup/creation failed: %v", err)
|
||||
c.JSON(http.StatusUnprocessableEntity, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrInternal,
|
||||
Message: "Failed to process SAML user",
|
||||
Detail: err.Error(),
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Issue JWT token pair (same pattern as login flow)
|
||||
tokenPair, err := h.jwtService.GenerateTokenPair(user, user.AccountID, user.Role)
|
||||
if err != nil {
|
||||
applogger.L().Errorf("Failed to generate JWT for SAML user: %v", err)
|
||||
c.JSON(http.StatusUnprocessableEntity, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrInternal,
|
||||
Message: "Failed to generate authentication tokens",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Store refresh token
|
||||
if h.refreshStore != nil {
|
||||
if err := h.refreshStore.Store(c.Request.Context(), user.ID, tokenPair.RefreshToken); err != nil {
|
||||
applogger.L().Warnf("Failed to store refresh token for SAML user: %v", err)
|
||||
// Non-fatal: access token is still valid
|
||||
}
|
||||
}
|
||||
|
||||
// Create SSO session in Redis (for SLO and session tracking)
|
||||
// Reference: M13 §4 — SSO session creation in ACS flow
|
||||
if h.ssoSessionStore != nil {
|
||||
idpEntityID := h.samlService.GetIdPEntityID()
|
||||
sessionData := &auth.SSOSessionData{
|
||||
UserID: user.ID,
|
||||
Provider: "saml",
|
||||
IdPEntityID: idpEntityID,
|
||||
NameID: userInfo.NameID,
|
||||
AccountID: user.AccountID,
|
||||
Role: user.Role,
|
||||
CreatedAt: time.Now().Unix(),
|
||||
ExpiresAt: time.Now().Add(h.ssoSessionStore.SessionTTL()).Unix(),
|
||||
}
|
||||
sessionID, err := h.ssoSessionStore.Create(c.Request.Context(), sessionData)
|
||||
if err != nil {
|
||||
applogger.L().Warnf("Failed to create SSO session for SAML user: %v", err)
|
||||
// Non-fatal: JWT tokens are still valid, SSO session is for tracking/SLO only
|
||||
} else {
|
||||
applogger.L().Infof("SSO session %s created for SAML user %d via IdP %s", sessionID, user.ID, idpEntityID)
|
||||
}
|
||||
}
|
||||
|
||||
// Return successful auth response (same format as login endpoint)
|
||||
response.OK(c, gin.H{
|
||||
"user": user,
|
||||
"access_token": tokenPair.AccessToken,
|
||||
"refresh_token": tokenPair.RefreshToken,
|
||||
"expires_at": tokenPair.ExpiresAt,
|
||||
})
|
||||
}
|
||||
|
||||
// RegisterSAMLRoutes sets up SAML routes on a Gin router group.
|
||||
// These routes are PUBLIC — no AuthRequired middleware (SAML flow is external).
|
||||
func RegisterSAMLRoutes(rg *gin.RouterGroup, handler *SAMLHandler) {
|
||||
samlGroup := rg.Group("/saml")
|
||||
{
|
||||
samlGroup.GET("/metadata", handler.Metadata)
|
||||
samlGroup.GET("/login", handler.Login)
|
||||
samlGroup.POST("/acs", handler.ACS)
|
||||
// SLO (Single Logout) endpoints — M13 §1
|
||||
samlGroup.GET("/slo", handler.SPInitiatedSLO) // SP-initiated: redirect user to IdP for logout
|
||||
samlGroup.POST("/slo", handler.IdPInitiatedSLO) // IdP-initiated: IdP sends LogoutRequest to us
|
||||
}
|
||||
}
|
||||
|
||||
// --- SAML Single Logout (SLO) Handlers ---
|
||||
// Reference: M13 §1 — SAML 2.0 Single Logout (SLO)
|
||||
|
||||
// SPInitiatedSLO handles SP-initiated Single Logout (HTTP-Redirect binding).
|
||||
// GET /api/v1/saml/slo
|
||||
// The user clicks logout in GoChat → we generate a SAML LogoutRequest
|
||||
// and redirect to the IdP's SLO endpoint. The IdP then propagates
|
||||
// logout to all SPs in the session.
|
||||
// Query params:
|
||||
// - session_id: the SSO session ID to terminate
|
||||
// - state: optional RelayState for post-logout redirect
|
||||
func (h *SAMLHandler) SPInitiatedSLO(c *gin.Context) {
|
||||
if !h.samlCfg.Enabled {
|
||||
c.JSON(http.StatusNotFound, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrNotFound,
|
||||
Message: "SAML is not enabled",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
sessionID := c.Query("session_id")
|
||||
if sessionID == "" {
|
||||
c.JSON(http.StatusBadRequest, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrBadRequest,
|
||||
Message: "Missing session_id parameter",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
state := c.Query("state")
|
||||
if state == "" {
|
||||
state = "/" // default redirect to home after logout
|
||||
}
|
||||
|
||||
// In a real flow, we'd look up the SSO session to get NameID + SessionIndex
|
||||
// from the DB/Redis. For now, use the session_id as both.
|
||||
// Production note: session store should be backed by Redis/DB for SLO validation.
|
||||
// Current implementation passes empty IDs as placeholders until SSOSessionRepo is wired.
|
||||
redirectURL, err := h.samlService.InitiateLogout(sessionID, sessionID, "", state)
|
||||
if err != nil {
|
||||
applogger.L().Errorf("SAML SLO initiation failed: %v", err)
|
||||
statusCode := http.StatusInternalServerError
|
||||
errCode := response.ErrInternal
|
||||
if errors.Is(err, auth.ErrSAMLEnabled) {
|
||||
statusCode = http.StatusNotFound
|
||||
errCode = response.ErrNotFound
|
||||
}
|
||||
c.JSON(statusCode, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: errCode,
|
||||
Message: "Failed to initiate SAML logout",
|
||||
Detail: err.Error(),
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Redirect user to IdP SLO endpoint
|
||||
c.Redirect(http.StatusFound, redirectURL)
|
||||
}
|
||||
|
||||
// IdPInitiatedSLO handles IdP-initiated Single Logout.
|
||||
// POST /api/v1/saml/slo
|
||||
// The IdP sends a base64-encoded SAML LogoutRequest to this endpoint.
|
||||
// We validate the request, terminate all matching SSO sessions,
|
||||
// and return a LogoutResponse.
|
||||
func (h *SAMLHandler) IdPInitiatedSLO(c *gin.Context) {
|
||||
if !h.samlCfg.Enabled {
|
||||
c.JSON(http.StatusNotFound, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrNotFound,
|
||||
Message: "SAML is not enabled",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
samlRequest := c.PostForm("SAMLRequest")
|
||||
if samlRequest == "" {
|
||||
c.JSON(http.StatusBadRequest, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrBadRequest,
|
||||
Message: "Missing SAMLRequest parameter",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Process the LogoutRequest from the IdP
|
||||
logoutResponse, err := h.samlService.ProcessLogoutRequest(samlRequest)
|
||||
if err != nil {
|
||||
applogger.L().Errorf("SAML IdP-initiated SLO failed: %v", err)
|
||||
statusCode := http.StatusInternalServerError
|
||||
errCode := response.ErrInternal
|
||||
if errors.Is(err, auth.ErrSAMLEnabled) {
|
||||
statusCode = http.StatusNotFound
|
||||
errCode = response.ErrNotFound
|
||||
}
|
||||
c.JSON(statusCode, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: errCode,
|
||||
Message: "Failed to process SAML logout request",
|
||||
Detail: err.Error(),
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// Return the LogoutResponse for the IdP (base64-encoded XML)
|
||||
c.JSON(http.StatusOK, response.APIResponse{
|
||||
Success: true,
|
||||
Data: map[string]string{
|
||||
"logout_response": logoutResponse,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// SLOResponse handles the IdP's LogoutResponse for SP-initiated SLO.
|
||||
// GET /api/v1/saml/slo/response
|
||||
// After the IdP processes our LogoutRequest, it redirects the user back
|
||||
// to this endpoint with a SAMLResponse (LogoutResponse) + RelayState.
|
||||
func (h *SAMLHandler) SLOResponse(c *gin.Context) {
|
||||
if !h.samlCfg.Enabled {
|
||||
c.JSON(http.StatusNotFound, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrNotFound,
|
||||
Message: "SAML is not enabled",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
samlResponse := c.Query("SAMLResponse")
|
||||
if samlResponse == "" {
|
||||
c.JSON(http.StatusBadRequest, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: response.ErrBadRequest,
|
||||
Message: "Missing SAMLResponse parameter",
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
relayState := c.Query("RelayState")
|
||||
|
||||
err := h.samlService.ProcessLogoutResponse(samlResponse, relayState)
|
||||
if err != nil {
|
||||
applogger.L().Errorf("SAML SLO response validation failed: %v", err)
|
||||
statusCode := http.StatusInternalServerError
|
||||
errCode := response.ErrInternal
|
||||
if errors.Is(err, auth.ErrSAMLEnabled) {
|
||||
statusCode = http.StatusNotFound
|
||||
errCode = response.ErrNotFound
|
||||
}
|
||||
c.JSON(statusCode, response.APIResponse{
|
||||
Success: false,
|
||||
Error: &response.ErrorBody{
|
||||
Code: errCode,
|
||||
Message: "SAML logout response validation failed",
|
||||
Detail: err.Error(),
|
||||
},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
// SLO successful — redirect user to the RelayState URL (or home)
|
||||
redirectURL := relayState
|
||||
if redirectURL == "" {
|
||||
redirectURL = "/"
|
||||
}
|
||||
c.Redirect(http.StatusFound, redirectURL)
|
||||
}
|
||||
Reference in New Issue
Block a user