package middleware import ( "fmt" "net/http" "net/http/httptest" "testing" "github.com/gin-gonic/gin" "github.com/stretchr/testify/assert" "gorm.io/driver/sqlite" "gorm.io/gorm" "github.com/gochat/gochat/internal/auth" "github.com/gochat/gochat/internal/config" "github.com/gochat/gochat/internal/model" ) func makeJWTConfig() *config.JWTConfig { return &config.JWTConfig{ Secret: "test-secret-key-for-middleware-test", ExpiryHours: 24, RefreshExpiryHours: 48, AccessExpiryMinutes: 30, Audience: "gochat-test", Issuer: "gochat-test", } } func makeValidAccessToken(cfg *config.JWTConfig) string { svc := auth.NewJWTService(cfg) user := &model.User{Base: model.Base{ID: 1}, Provider: "email", Email: "test@test.com"} pair, _ := svc.GenerateTokenPair(user, 2, "agent") return pair.AccessToken } func TestAuthMiddleware_NoAuthHeader(t *testing.T) { gin.SetMode(gin.TestMode) cfg := makeJWTConfig() r := gin.New() r.Use(AuthMiddleware(cfg)) r.GET("/test", func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) }) w := httptest.NewRecorder() req := httptest.NewRequest(http.MethodGet, "/test", nil) r.ServeHTTP(w, req) assert.Equal(t, 401, w.Code) } func TestAuthMiddleware_BearerTokenRequired(t *testing.T) { gin.SetMode(gin.TestMode) cfg := makeJWTConfig() r := gin.New() r.Use(AuthMiddleware(cfg)) r.GET("/test", func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) }) w := httptest.NewRecorder() req := httptest.NewRequest(http.MethodGet, "/test", nil) req.Header.Set("Authorization", "Basic abc123") r.ServeHTTP(w, req) assert.Equal(t, 401, w.Code) } func TestAuthMiddleware_InvalidToken(t *testing.T) { gin.SetMode(gin.TestMode) cfg := makeJWTConfig() r := gin.New() r.Use(AuthMiddleware(cfg)) r.GET("/test", func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) }) w := httptest.NewRecorder() req := httptest.NewRequest(http.MethodGet, "/test", nil) req.Header.Set("Authorization", "Bearer invalid-token-string") r.ServeHTTP(w, req) assert.Equal(t, 401, w.Code) } func TestAuthMiddleware_ValidToken(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("Authorization", fmt.Sprintf("Bearer %s", token)) r.ServeHTTP(w, req) 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_RejectsRevokedChatwootSession(t *testing.T) { ginsvc := gin.New() cfg := makeJWTConfig() db, err := gorm.Open(sqlite.Open("file:"+t.Name()+"?mode=memory&cache=private"), &gorm.Config{}) assert.NoError(t, err) assert.NoError(t, db.AutoMigrate(&model.User{}, &model.UserSession{})) jwtService := auth.NewJWTService(cfg) user := &model.User{Base: model.Base{ID: 1}, Provider: "email", Email: "session@example.com"} assert.NoError(t, db.Create(user).Error) pair, err := jwtService.GenerateTokenPairForClient(user, 2, "agent", "client-1") assert.NoError(t, err) ginsvc.Use(AuthMiddlewareWithServiceAndDB(jwtService, db)) ginsvc.GET("/test", func(c *gin.Context) { c.Status(http.StatusOK) }) request := func() int { w := httptest.NewRecorder() req := httptest.NewRequest(http.MethodGet, "/test", nil) req.Header.Set("access-token", pair.AccessToken) ginsvc.ServeHTTP(w, req) return w.Code } assert.Equal(t, http.StatusUnauthorized, request()) assert.NoError(t, db.Create(&model.UserSession{UserID: user.ID, ClientID: "client-1"}).Error) assert.Equal(t, http.StatusOK, request()) assert.NoError(t, db.Where("user_id = ? AND client_id = ?", user.ID, "client-1").Delete(&model.UserSession{}).Error) assert.Equal(t, http.StatusUnauthorized, request()) } func TestAuthMiddleware_AllowsPlatformAdminThroughSuperAdminGuard(t *testing.T) { gin.SetMode(gin.TestMode) cfg := makeJWTConfig() jwtService := auth.NewJWTService(cfg) user := &model.User{ Base: model.Base{ID: 1}, Provider: "email", Role: "super_admin", Type: "User", } pair, err := jwtService.GenerateTokenPair(user, 2, "administrator") assert.NoError(t, err) r := gin.New() r.Use(AuthMiddleware(cfg), SuperAdmin()) r.GET("/test", func(c *gin.Context) { c.JSON(http.StatusOK, gin.H{"ok": true}) }) w := httptest.NewRecorder() req := httptest.NewRequest(http.MethodGet, "/test", nil) req.Header.Set("access-token", pair.AccessToken) r.ServeHTTP(w, req) assert.Equal(t, http.StatusOK, w.Code) } func TestAuthMiddleware_FallbackHeaders(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}) }) w := httptest.NewRecorder() req := httptest.NewRequest(http.MethodGet, "/test", nil) req.Header.Set("X-User-ID", "5") req.Header.Set("X-Account-ID", "10") r.ServeHTTP(w, req) assert.Equal(t, 200, w.Code) } func TestAuthMiddleware_FallbackInvalidUserID(t *testing.T) { gin.SetMode(gin.TestMode) cfg := makeJWTConfig() r := gin.New() r.Use(AuthMiddleware(cfg)) r.GET("/test", func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) }) w := httptest.NewRecorder() req := httptest.NewRequest(http.MethodGet, "/test", nil) req.Header.Set("X-User-ID", "abc") req.Header.Set("X-Account-ID", "10") r.ServeHTTP(w, req) assert.Equal(t, 401, w.Code) } func TestAuthMiddleware_FallbackMissingUserID(t *testing.T) { gin.SetMode(gin.TestMode) cfg := makeJWTConfig() r := gin.New() r.Use(AuthMiddleware(cfg)) r.GET("/test", func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) }) w := httptest.NewRecorder() req := httptest.NewRequest(http.MethodGet, "/test", nil) req.Header.Set("X-Account-ID", "10") r.ServeHTTP(w, req) assert.Equal(t, 401, w.Code) } func TestGenerateToken(t *testing.T) { cfg := makeJWTConfig() tokenString, err := GenerateToken(cfg, 1, 2, "agent") assert.NoError(t, err) assert.NotEmpty(t, tokenString) }