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:
Rogee
2026-07-29 19:03:04 +08:00
parent 09f274e965
commit 851ca7e372
66 changed files with 154 additions and 10590 deletions
@@ -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)
}