Align GoChat with Chatwoot frontend contracts
This commit is contained in:
@@ -39,9 +39,14 @@ func AccountScope() gin.HandlerFunc {
|
||||
return
|
||||
}
|
||||
|
||||
// Step 2: Determine account_id
|
||||
// Priority: X-Account-ID header > JWT claims account_id
|
||||
accountID := getAccountID(c)
|
||||
// Step 2: Determine account_id.
|
||||
// Account-scoped Chatwoot routes carry :account_id in the URL. The token/header
|
||||
// account context must match that URL account instead of silently allowing a
|
||||
// token scoped to one account to read another account's route.
|
||||
accountID, ok := resolveScopedAccountID(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if accountID == 0 {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest,
|
||||
"Account ID required — provide via X-Account-ID header or JWT claims")
|
||||
@@ -107,7 +112,10 @@ func AccountScopeWithService(lookup RBACLookup) gin.HandlerFunc {
|
||||
return
|
||||
}
|
||||
|
||||
accountID := getAccountID(c)
|
||||
accountID, ok := resolveScopedAccountID(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if accountID == 0 {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest,
|
||||
"Account ID required — provide via X-Account-ID header or JWT claims")
|
||||
@@ -154,6 +162,41 @@ func AccountScopeWithService(lookup RBACLookup) gin.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
func resolveScopedAccountID(c *gin.Context) (uint, bool) {
|
||||
contextAccountID := getAccountID(c)
|
||||
routeAccountID, hasRouteAccountID, routeOK := routeAccountID(c)
|
||||
if !routeOK {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "invalid account_id")
|
||||
return 0, false
|
||||
}
|
||||
if !hasRouteAccountID {
|
||||
return contextAccountID, true
|
||||
}
|
||||
if contextAccountID == 0 {
|
||||
return routeAccountID, true
|
||||
}
|
||||
if contextAccountID != routeAccountID {
|
||||
response.AbortWithStatusError(c, http.StatusForbidden, response.ErrForbidden, "User does not belong to this account")
|
||||
return 0, false
|
||||
}
|
||||
return routeAccountID, true
|
||||
}
|
||||
|
||||
func routeAccountID(c *gin.Context) (uint, bool, bool) {
|
||||
for _, param := range []string{"account_id", "id"} {
|
||||
raw := c.Param(param)
|
||||
if raw == "" {
|
||||
continue
|
||||
}
|
||||
id, err := strconv.ParseUint(raw, 10, 32)
|
||||
if err != nil || id == 0 {
|
||||
return 0, true, false
|
||||
}
|
||||
return uint(id), true, true
|
||||
}
|
||||
return 0, false, true
|
||||
}
|
||||
|
||||
// getAccountID extracts account ID from the request.
|
||||
// Priority: X-Account-ID header > JWT claims account_id
|
||||
func getAccountID(c *gin.Context) uint {
|
||||
|
||||
@@ -102,4 +102,36 @@ func TestAccountScope_InvalidHeaderAccountID(t *testing.T) {
|
||||
req.Header.Set("X-Account-ID", "abc")
|
||||
r.ServeHTTP(w, req)
|
||||
assert.Equal(t, 400, w.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAccountScope_RouteAccountMustMatchContext(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
r := gin.New()
|
||||
r.Use(func(c *gin.Context) { c.Set("user_id", uint(1)); c.Set("account_id", uint(2)); c.Next() })
|
||||
r.Use(AccountScope())
|
||||
r.GET("/api/v2/accounts/:account_id/live_reports/conversation_metrics", func(c *gin.Context) {
|
||||
accountID, _ := c.Get("account_id")
|
||||
c.JSON(200, gin.H{"account_id": accountID})
|
||||
})
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/v2/accounts/3/live_reports/conversation_metrics", nil)
|
||||
r.ServeHTTP(w, req)
|
||||
assert.Equal(t, http.StatusForbidden, w.Code)
|
||||
}
|
||||
|
||||
func TestAccountScope_RouteAccountMatchesContext(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
r := gin.New()
|
||||
r.Use(func(c *gin.Context) { c.Set("user_id", uint(1)); c.Set("account_id", uint(2)); c.Next() })
|
||||
r.Use(AccountScope())
|
||||
r.GET("/api/v2/accounts/:account_id/live_reports/conversation_metrics", func(c *gin.Context) {
|
||||
accountID, _ := c.Get("account_id")
|
||||
c.JSON(200, gin.H{"account_id": accountID})
|
||||
})
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/v2/accounts/2/live_reports/conversation_metrics", nil)
|
||||
r.ServeHTTP(w, req)
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
}
|
||||
|
||||
@@ -32,9 +32,12 @@ func AuthMiddleware(cfg *config.JWTConfig) gin.HandlerFunc {
|
||||
func AuthMiddlewareWithService(jwtSvc *auth.JWTService) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
authHeader := c.GetHeader("Authorization")
|
||||
if authHeader != "" {
|
||||
chatwootAccessToken := strings.TrimSpace(c.GetHeader("access-token"))
|
||||
if authHeader != "" || chatwootAccessToken != "" {
|
||||
tokenString := strings.TrimPrefix(authHeader, "Bearer ")
|
||||
if tokenString == authHeader {
|
||||
if authHeader == "" {
|
||||
tokenString = chatwootAccessToken
|
||||
} else if tokenString == authHeader {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "bearer token required"})
|
||||
return
|
||||
}
|
||||
@@ -105,4 +108,4 @@ func GenerateToken(cfg *config.JWTConfig, userID uint, accountID uint, role stri
|
||||
|
||||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
|
||||
return token.SignedString([]byte(cfg.Secret))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -98,6 +98,26 @@ func TestAuthMiddleware_ValidToken(t *testing.T) {
|
||||
assert.Equal(t, 200, w.Code)
|
||||
}
|
||||
|
||||
func TestAuthMiddleware_ChatwootAccessTokenHeader(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
cfg := makeJWTConfig()
|
||||
r := gin.New()
|
||||
r.Use(AuthMiddleware(cfg))
|
||||
r.GET("/test", func(c *gin.Context) {
|
||||
userID, _ := c.Get("user_id")
|
||||
accountID, _ := c.Get("account_id")
|
||||
c.JSON(200, gin.H{"user_id": userID, "account_id": accountID})
|
||||
})
|
||||
|
||||
token := makeValidAccessToken(cfg)
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/test", nil)
|
||||
req.Header.Set("access-token", token)
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, 200, w.Code)
|
||||
}
|
||||
|
||||
func TestAuthMiddleware_FallbackHeaders(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
cfg := makeJWTConfig()
|
||||
@@ -155,4 +175,4 @@ func TestGenerateToken(t *testing.T) {
|
||||
tokenString, err := GenerateToken(cfg, 1, 2, "agent")
|
||||
assert.NoError(t, err)
|
||||
assert.NotEmpty(t, tokenString)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -18,14 +18,14 @@ type CORSConfig struct {
|
||||
AllowedHeaders []string
|
||||
ExposeHeaders []string
|
||||
AllowCredentials bool
|
||||
MaxAge int // seconds
|
||||
MaxAge int // seconds
|
||||
DevMode bool // when true and AllowedOrigins is empty, fall back to Allow-Origin: *
|
||||
}
|
||||
|
||||
// Default CORS values for production.
|
||||
var defaultCORSMethods = []string{"GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS"}
|
||||
var defaultCORSHeaders = []string{"Origin", "Content-Type", "Accept", "Authorization", "X-Account-ID"}
|
||||
var defaultCORSExposeHeaders = []string{"Content-Length"}
|
||||
var defaultCORSHeaders = []string{"Origin", "Content-Type", "Accept", "Authorization", "X-Account-ID", "access-token", "client", "uid", "token-type", "expiry"}
|
||||
var defaultCORSExposeHeaders = []string{"Content-Length", "access-token", "client", "uid", "token-type", "expiry"}
|
||||
var defaultCORSMaxAge = 86400 // 24 hours
|
||||
|
||||
// CORS adds Cross-Origin Resource Sharing headers.
|
||||
@@ -164,4 +164,4 @@ func CORSConfigFromAppConfig(cfg *config.Config) CORSConfig {
|
||||
MaxAge: cfg.Server.CORS.MaxAge,
|
||||
DevMode: devMode,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -32,10 +32,54 @@ func TestCORS_DevMode_AllOrigins(t *testing.T) {
|
||||
assert.Contains(t, w.Header().Get("Access-Control-Allow-Methods"), "GET")
|
||||
assert.Contains(t, w.Header().Get("Access-Control-Allow-Methods"), "POST")
|
||||
assert.Contains(t, w.Header().Get("Access-Control-Allow-Headers"), "Authorization")
|
||||
assert.Equal(t, "Content-Length", w.Header().Get("Access-Control-Expose-Headers"))
|
||||
assert.Contains(t, w.Header().Get("Access-Control-Allow-Headers"), "access-token")
|
||||
assert.Contains(t, w.Header().Get("Access-Control-Allow-Headers"), "client")
|
||||
assert.Contains(t, w.Header().Get("Access-Control-Allow-Headers"), "uid")
|
||||
assert.Contains(t, w.Header().Get("Access-Control-Allow-Headers"), "token-type")
|
||||
assert.Contains(t, w.Header().Get("Access-Control-Expose-Headers"), "Content-Length")
|
||||
assert.Contains(t, w.Header().Get("Access-Control-Expose-Headers"), "access-token")
|
||||
assert.Contains(t, w.Header().Get("Access-Control-Expose-Headers"), "client")
|
||||
assert.Contains(t, w.Header().Get("Access-Control-Expose-Headers"), "uid")
|
||||
assert.Contains(t, w.Header().Get("Access-Control-Expose-Headers"), "token-type")
|
||||
assert.Contains(t, w.Header().Get("Access-Control-Expose-Headers"), "expiry")
|
||||
assert.Equal(t, "86400", w.Header().Get("Access-Control-Max-Age"))
|
||||
}
|
||||
|
||||
func TestCORS_DefaultAllowedHeadersIncludeChatwootAuthTokens(t *testing.T) {
|
||||
cfg := CORSConfig{DevMode: true}
|
||||
router := gin.New()
|
||||
router.Use(CORS(cfg))
|
||||
router.GET("/auth/validate_token", func(c *gin.Context) { c.Status(200) })
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
req, _ := http.NewRequest("OPTIONS", "/auth/validate_token", nil)
|
||||
req.Header.Set("Origin", "http://localhost:3037")
|
||||
req.Header.Set("Access-Control-Request-Headers", "access-token, client, uid, token-type")
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
allowedHeaders := w.Header().Get("Access-Control-Allow-Headers")
|
||||
for _, header := range []string{"access-token", "client", "uid", "token-type", "expiry"} {
|
||||
assert.Contains(t, allowedHeaders, header)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCORS_DefaultExposeHeadersIncludeChatwootAuthTokens(t *testing.T) {
|
||||
cfg := CORSConfig{DevMode: true}
|
||||
router := gin.New()
|
||||
router.Use(CORS(cfg))
|
||||
router.POST("/auth/sign_in", func(c *gin.Context) { c.Status(200) })
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
req, _ := http.NewRequest("POST", "/auth/sign_in", nil)
|
||||
req.Header.Set("Origin", "http://localhost:3037")
|
||||
router.ServeHTTP(w, req)
|
||||
|
||||
exposedHeaders := w.Header().Get("Access-Control-Expose-Headers")
|
||||
for _, header := range []string{"access-token", "client", "uid", "token-type", "expiry"} {
|
||||
assert.Contains(t, exposedHeaders, header)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCORS_DevMode_WithWhitelist(t *testing.T) {
|
||||
cfg := CORSConfig{
|
||||
AllowedOrigins: []string{"https://app.example.com"},
|
||||
@@ -314,4 +358,4 @@ func TestCORSConfigFromAppConfig_EmptyCORS(t *testing.T) {
|
||||
assert.False(t, mwCfg.DevMode)
|
||||
assert.Empty(t, mwCfg.AllowedOrigins)
|
||||
assert.Empty(t, mwCfg.AllowedMethods) // defaults applied in CORS() middleware, not here
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestCSRFSkipsChatwootAuthRoutes(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
router := gin.New()
|
||||
router.Use(CSRF(CSRFConfig{
|
||||
Enabled: true,
|
||||
Secret: "test-secret",
|
||||
CookieName: "_gochat_csrf",
|
||||
HeaderName: "X-CSRF-Token",
|
||||
TokenLength: 32,
|
||||
SafeMethods: []string{"GET", "HEAD", "OPTIONS"},
|
||||
SkipPaths: []string{"/auth/"},
|
||||
CookiePath: "/",
|
||||
CookieSameSite: "Lax",
|
||||
}))
|
||||
router.POST("/auth/sign_in", func(c *gin.Context) { c.Status(http.StatusOK) })
|
||||
|
||||
recorder := httptest.NewRecorder()
|
||||
request := httptest.NewRequest(http.MethodPost, "/auth/sign_in", nil)
|
||||
router.ServeHTTP(recorder, request)
|
||||
|
||||
require.Equal(t, http.StatusOK, recorder.Code)
|
||||
}
|
||||
|
||||
func TestCSRFSkipsTokenAuthenticatedAPIRoutes(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
router := gin.New()
|
||||
router.Use(CSRF(CSRFConfig{
|
||||
Enabled: true,
|
||||
Secret: "test-secret",
|
||||
CookieName: "_gochat_csrf",
|
||||
HeaderName: "X-CSRF-Token",
|
||||
TokenLength: 32,
|
||||
SafeMethods: []string{"GET", "HEAD", "OPTIONS"},
|
||||
SkipPaths: []string{"/api/v1/"},
|
||||
CookiePath: "/",
|
||||
CookieSameSite: "Lax",
|
||||
}))
|
||||
router.POST("/api/v1/accounts/:account_id/conversations/:conversation_id/messages", func(c *gin.Context) { c.Status(http.StatusOK) })
|
||||
|
||||
recorder := httptest.NewRecorder()
|
||||
request := httptest.NewRequest(http.MethodPost, "/api/v1/accounts/1/conversations/1/messages", nil)
|
||||
router.ServeHTTP(recorder, request)
|
||||
|
||||
require.Equal(t, http.StatusOK, recorder.Code)
|
||||
}
|
||||
@@ -26,9 +26,9 @@ const (
|
||||
// Falls back to in-memory limiting when Redis is unavailable.
|
||||
// Reference: Chatwoot's Rack::Attack throttle configuration.
|
||||
type slidingWindowLimiter struct {
|
||||
redis *redis.Client
|
||||
cfg *config.RateLimitConfig
|
||||
fallback *inMemoryLimiter
|
||||
redis *redis.Client
|
||||
cfg *config.RateLimitConfig
|
||||
fallback *inMemoryLimiter
|
||||
redisAvailable atomic.Bool
|
||||
}
|
||||
|
||||
@@ -287,6 +287,11 @@ func RateLimit(cfg *config.Config, rdb *redis.Client) gin.HandlerFunc {
|
||||
}
|
||||
|
||||
return func(c *gin.Context) {
|
||||
if isRateLimitExemptPath(c.Request.URL.Path) {
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
|
||||
ip := c.ClientIP()
|
||||
key := "global:" + ip
|
||||
|
||||
@@ -317,14 +322,24 @@ func RateLimit(cfg *config.Config, rdb *redis.Client) gin.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
func isRateLimitExemptPath(path string) bool {
|
||||
switch path {
|
||||
case "/health", "/metrics":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// PerRouteLimit creates a per-route rate limiting middleware.
|
||||
// This allows different rate limits for different API endpoints, matching
|
||||
// Chatwoot's Rack::Attack per-route throttle configuration.
|
||||
//
|
||||
// Usage:
|
||||
// router.GET("/conversations", PerRouteLimit("conversations_list", 60), listConversations)
|
||||
// router.POST("/messages", PerRouteLimit("messages_create", 30), createMessage)
|
||||
// router.GET("/reports", PerRouteLimit("reports_read", 10), viewReports)
|
||||
//
|
||||
// router.GET("/conversations", PerRouteLimit("conversations_list", 60), listConversations)
|
||||
// router.POST("/messages", PerRouteLimit("messages_create", 30), createMessage)
|
||||
// router.GET("/reports", PerRouteLimit("reports_read", 10), viewReports)
|
||||
//
|
||||
// The route identifier is used as a key namespace so that limits on one route
|
||||
// don't affect limits on another route for the same IP.
|
||||
@@ -410,7 +425,8 @@ func PerRouteLimit(route string, requestsPerMin int) gin.HandlerFunc {
|
||||
// endpoints where multiple users may share the same IP (e.g., office networks).
|
||||
//
|
||||
// Usage:
|
||||
// router.POST("/api/v1/conversations", AuthRequired(jwtSvc), PerUserLimit("conversations_create", 30), createConversation)
|
||||
//
|
||||
// router.POST("/api/v1/conversations", AuthRequired(jwtSvc), PerUserLimit("conversations_create", 30), createConversation)
|
||||
func PerUserLimit(route string, requestsPerMin int) gin.HandlerFunc {
|
||||
type visitor struct {
|
||||
count int
|
||||
@@ -526,4 +542,4 @@ func PerUserLimit(route string, requestsPerMin int) gin.HandlerFunc {
|
||||
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user