Align GoChat with Chatwoot frontend contracts

This commit is contained in:
2026-06-13 22:13:32 +08:00
parent 71abf58636
commit d884fdda0a
162 changed files with 9825 additions and 528 deletions
+47 -4
View File
@@ -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 {
+33 -1
View File
@@ -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)
}
+6 -3
View File
@@ -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))
}
}
+21 -1
View File
@@ -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)
}
}
+4 -4
View File
@@ -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,
}
}
}
+46 -2
View File
@@ -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
}
}
+56
View File
@@ -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)
}
+24 -8
View File
@@ -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()
}
}
}