From fa6737e258a66a239c61210a9826bba5c24ed1d5 Mon Sep 17 00:00:00 2001 From: Rogee Date: Wed, 29 Jul 2026 16:26:19 +0800 Subject: [PATCH] =?UTF-8?q?refactor:=20=E7=B2=BE=E7=AE=80=E9=85=8D?= =?UTF-8?q?=E7=BD=AE=E4=BD=93=E7=B3=BB=EF=BC=8C=E7=A7=BB=E9=99=A4=20OAuth?= =?UTF-8?q?=20=E7=99=BB=E5=BD=95/Rate=20Limit/Admin=20env=20=E9=85=8D?= =?UTF-8?q?=E7=BD=AE?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 移除 .env.example 中 Feature Flags 段(代码中不存在这些 env var) - 移除 Google/GitHub OAuth 登录认证代码(auth/oauth.go、auth_handler OAuthAuthorize/OAuthCallback 路由、auth_service OAuthLogin),保留 Twitter/Google 作为消息渠道 provider - 从 OAuthConfig 移除 GitHub 字段(Google 保留供 channel provider 使用) - 移除 RateLimitConfig 可配置性,RateLimit 中间件改为硬编码 100 req/min、 60s window,移除 config/validator/reloader 中的 rate_limit 相关代码 - 移除 .env.example 中 GOCHAT_ADMIN_EMAIL/PASSWORD 配置 - 新增 gochat init 命令:交互式或通过 --email/--password/--name flags 初始化超级管理员账户,创建默认 Account + AccountUser 关联 --- .env.example | 19 -- backend/cmd/gochat/main.go | 138 +++++++++- backend/internal/app/bootstrap.go | 7 +- backend/internal/auth/oauth.go | 257 ------------------ backend/internal/config/config.go | 46 ---- backend/internal/config/config_test.go | 4 - backend/internal/config/reloader_test.go | 27 +- backend/internal/config/validator.go | 8 - .../internal/handler/api/v1/auth_handler.go | 88 +----- .../handler/api/v1/auth_handler_test.go | 4 +- backend/internal/middleware/rate_limit.go | 45 +-- .../internal/middleware/rate_limit_test.go | 194 +++---------- backend/internal/service/auth_service.go | 97 ------- backend/internal/service/auth_service_test.go | 2 +- backend/tests/e2e/e2e_test.go | 5 +- 15 files changed, 209 insertions(+), 732 deletions(-) delete mode 100644 backend/internal/auth/oauth.go diff --git a/.env.example b/.env.example index 30a6fb1c..179a3f53 100644 --- a/.env.example +++ b/.env.example @@ -70,21 +70,7 @@ GOCHAT_ANALYTICS_ENABLED=true GOCHAT_ANALYTICS_RETENTION_DAYS=90 # how long to keep reporting events GOCHAT_ANALYTICS_FLUSH_INTERVAL=300 # seconds between rollup flushes -# ---- Feature Flags ---- -GOCHAT_FEATURE_CAPTAIN_AI=false -GOCHAT_FEATURE_AUTO_ASSIGNMENT=true -GOCHAT_FEATURE_CSAT=true -GOCHAT_FEATURE_CAMPAIGNS=false -GOCHAT_FEATURE_MFA=true - # ---- OAuth Providers ---- -GOCHAT_OAUTH_GOOGLE_CLIENT_ID= -GOCHAT_OAUTH_GOOGLE_CLIENT_SECRET= -GOCHAT_OAUTH_GOOGLE_REDIRECT_URL=https://your-domain.com/api/v1/auth/google/callback -GOCHAT_OAUTH_GOOGLE_SCOPES=openid,email,profile -GOCHAT_OAUTH_GITHUB_CLIENT_ID= -GOCHAT_OAUTH_GITHUB_CLIENT_SECRET= -GOCHAT_OAUTH_GITHUB_REDIRECT_URL=https://your-domain.com/api/v1/auth/github/callback # G10: Twitter/X OAuth 2.0 (PKCE) + Account Activity API GOCHAT_OAUTH_TWITTER_CLIENT_ID= GOCHAT_OAUTH_TWITTER_CLIENT_SECRET= @@ -131,8 +117,3 @@ GOCHAT_OIDC_DEFAULT_USER_INFO_URL= GOCHAT_OIDC_DEFAULT_JWKS_URL= GOCHAT_OIDC_DEFAULT_SCOPES=openid,profile,email -# ---- Security ---- -GOCHAT_RATE_LIMIT_REQUESTS_PER_MINUTE=100 # requests per minute per IP -GOCHAT_RATE_LIMIT_WINDOW_SECONDS=60 # seconds -GOCHAT_ADMIN_EMAIL=admin@example.com -GOCHAT_ADMIN_PASSWORD=CHANGE_ME_TO_A_STRONG_PASSWORD diff --git a/backend/cmd/gochat/main.go b/backend/cmd/gochat/main.go index b8ee2d1a..bc784da8 100644 --- a/backend/cmd/gochat/main.go +++ b/backend/cmd/gochat/main.go @@ -1,6 +1,7 @@ package main import ( + "bufio" "context" "encoding/json" "errors" @@ -34,6 +35,8 @@ func main() { err = serve() case "seed": err = seed() + case "init": + err = initAdmin() case "help", "-h", "--help": printUsage() return @@ -47,9 +50,10 @@ func main() { } func printUsage() { - fmt.Println("Usage: gochat [serve|seed]") + fmt.Println("Usage: gochat [serve|seed|init]") fmt.Println(" serve Start the GoChat HTTP server") fmt.Println(" seed Create deterministic development/smoke data") + fmt.Println(" init Initialize super admin account (interactive or via flags)") } func serve() error { @@ -338,3 +342,135 @@ func getenvDefault(key, fallback string) string { } return fallback } + +// --- init command: initialize super admin account --- + +func initAdmin() error { + env := os.Getenv("GOCHAT_ENV") + if env == "" { + env = "development" + } + cfg, err := config.LoadWithEnv(env) + if err != nil { + return fmt.Errorf("config load failed: %w", err) + } + if shouldRunSeedMigrations(cfg) { + if err := database.RunMigrations(cfg.Database.MigrateDSN(), cfg.Database.GetMigrationsPath()); err != nil { + return fmt.Errorf("database migrations failed: %w", err) + } + } + db, err := app.NewDatabase(&cfg.Database, cfg.Log.Level) + if err != nil { + return err + } + defer closeDB(db) + + ctx := context.Background() + + // Check if a super admin already exists. + var count int64 + if err := db.WithContext(ctx).Model(&model.User{}).Where("role = ? OR type = ?", "super_admin", "SuperAdmin").Count(&count).Error; err != nil { + return fmt.Errorf("check existing admin: %w", err) + } + if count > 0 { + fmt.Println("Super admin account already exists. Skipping initialization.") + return nil + } + + // Gather admin details from flags or interactive prompt. + email, password, name, err := readAdminCredentials() + if err != nil { + return err + } + + hashed, err := crypto.HashPassword(password) + if err != nil { + return fmt.Errorf("hash password: %w", err) + } + now := time.Now() + + // Create default account. + account := &model.Account{} + if err := db.WithContext(ctx).Where("name = ?", "Default Account").FirstOrCreate(account, model.Account{Name: "Default Account", Locale: "zh_CN", Timezone: "UTC", Active: true, Status: "active"}).Error; err != nil { + return fmt.Errorf("create default account: %w", err) + } + + // Create super admin user. + admin := &model.User{} + if err := db.WithContext(ctx).Where("email = ?", email).FirstOrCreate(admin, model.User{ + AccountID: account.ID, Name: name, DisplayName: name, Email: email, + Password: hashed, PasswordDigest: hashed, Provider: "email", + Role: "super_admin", Type: "SuperAdmin", Active: true, Available: true, + ConfirmedAt: &now, UISettings: datatypes.JSON([]byte(`{}`)), CustomAttributes: datatypes.JSON([]byte(`{}`)), + }).Error; err != nil { + return fmt.Errorf("create admin user: %w", err) + } + + // Link admin to account. + if err := db.WithContext(ctx).Where("user_id = ? AND account_id = ?", admin.ID, account.ID).FirstOrCreate(&model.AccountUser{}, model.AccountUser{ + UserID: admin.ID, AccountID: account.ID, Role: "administrator", Availability: "online", AutoOffline: true, + }).Error; err != nil { + return fmt.Errorf("create account user: %w", err) + } + + fmt.Printf("✓ Super admin created successfully\n") + fmt.Printf(" Email: %s\n", email) + fmt.Printf(" Account: Default Account (ID: %d)\n", account.ID) + fmt.Printf(" User ID: %d\n", admin.ID) + return nil +} + +// readAdminCredentials collects email, password, and display name from CLI flags +// or interactive prompt. +func readAdminCredentials() (email, password, name string, err error) { + args := os.Args[2:] + for i := 0; i < len(args); i++ { + switch args[i] { + case "--email": + if i+1 < len(args) { + email = args[i+1] + i++ + } + case "--password": + if i+1 < len(args) { + password = args[i+1] + i++ + } + case "--name": + if i+1 < len(args) { + name = args[i+1] + i++ + } + case "-h", "--help": + fmt.Println("Usage: gochat init [--email ] [--password ] [--name ]") + os.Exit(0) + } + } + + reader := bufio.NewReader(os.Stdin) + if email == "" { + fmt.Print("Admin email: ") + email, _ = reader.ReadString('\n') + email = strings.TrimSpace(email) + } + if email == "" { + return "", "", "", fmt.Errorf("email is required") + } + if password == "" { + fmt.Print("Admin password: ") + password, _ = reader.ReadString('\n') + password = strings.TrimSpace(password) + } + if password == "" { + return "", "", "", fmt.Errorf("password is required") + } + if name == "" { + fmt.Print("Admin display name [Super Admin]: ") + name, _ = reader.ReadString('\n') + name = strings.TrimSpace(name) + if name == "" { + name = "Super Admin" + } + } + return email, password, name, nil +} diff --git a/backend/internal/app/bootstrap.go b/backend/internal/app/bootstrap.go index 638663fe..f5763351 100644 --- a/backend/internal/app/bootstrap.go +++ b/backend/internal/app/bootstrap.go @@ -139,7 +139,6 @@ func Bootstrap(env string) (*App, error) { jwtService := auth.NewJWTService(&cfg.JWT) refreshStore := auth.NewRefreshTokenStore(rdb, &cfg.JWT) sessionStore := auth.NewSessionStore(&cfg.Session) // session management (ref: Chatwoot Devise sessions) - oauthService := auth.NewOAuthService(db, cfg) mfaService := auth.NewMFAService(db) webhookRegistry := auth.NewWebhookTokenRegistry() @@ -320,7 +319,7 @@ func Bootstrap(env string) (*App, error) { agentCapacityPolicyRepo := repository.NewAgentCapacityPolicyRepo(db) // Step 8: Wire services (business logic layer) - authService := service.NewAuthService(db, jwtService, refreshStore, oauthService, mfaService) + authService := service.NewAuthService(db, jwtService, refreshStore, mfaService) accountService := service.NewAccountService(accountRepo) accountService.SetWorkerPool(workerPool) whatsAppCallService := service.NewWhatsAppCallService(whatsAppCallRepo) @@ -825,7 +824,7 @@ func Bootstrap(env string) (*App, error) { contactMergeService := service.NewContactMergeService(contactMergeRepo, db) handlers := &router.Handlers{ RBAC: service.NewRBACService(db), - Auth: v1.NewAuthHandler(authService, oauthService, profileService), + Auth: v1.NewAuthHandler(authService, profileService), MFA: v1.NewMFAHandler(mfaService), SAML: v1.NewSAMLHandler(samlService, jwtService, refreshStore, ssoSessionStore, &cfg.SAML), Account: v1.NewAccountHandler(accountService), @@ -975,7 +974,7 @@ func Bootstrap(env string) (*App, error) { corsMiddleware := middleware.CORS(middleware.CORSConfigFromAppConfig(cfg)) engine.Use(middleware.Recovery()) // panic recovery engine.Use(middleware.RequestLogger()) // structured request logging - engine.Use(middleware.RateLimit(cfg, rdb)) // rate limiting (ref: Chatwoot rack-attack) + engine.Use(middleware.RateLimit(rdb)) // rate limiting (ref: Chatwoot rack-attack) engine.Use(corsMiddleware) // CORS with configurable whitelist engine.Use(middleware.SecurityHeaders(middleware.DefaultSecurityHeadersConfig())) // security headers (ref: P14 deliverable #11) engine.StaticFS("/uploads", gin.Dir(cfg.Storage.LocalPath, false)) diff --git a/backend/internal/auth/oauth.go b/backend/internal/auth/oauth.go deleted file mode 100644 index 32f31600..00000000 --- a/backend/internal/auth/oauth.go +++ /dev/null @@ -1,257 +0,0 @@ -package auth - -import ( - "context" - "encoding/json" - "fmt" - "sync" - - "golang.org/x/oauth2" - "golang.org/x/oauth2/google" - - "github.com/gochat/gochat/internal/config" - "github.com/gochat/gochat/internal/model" - "gorm.io/gorm" -) - -// Reference: P2E §1.3 — OAuth2 provider integration -// Replaces Chatwoot's OmniAuth-based Google/GitHub/Facebook login flow. - -// OAuthProviderType defines supported OAuth2 provider types. -type OAuthProviderType string - -const ( - OAuthProviderGoogle OAuthProviderType = "google" - OAuthProviderGitHub OAuthProviderType = "github" - OAuthProviderCustom OAuthProviderType = "custom" -) - -// OAuthUserInfo represents user info extracted from OAuth2 provider. -type OAuthUserInfo struct { - Provider OAuthProviderType - UID string // provider-specific user ID - Email string - Name string - AvatarURL string -} - -// OAuthService manages OAuth2 authentication flows. -type OAuthService struct { - db *gorm.DB - configs map[OAuthProviderType]*oauth2.Config - mu sync.RWMutex -} - -// NewOAuthService creates an OAuth service with configured providers. -func NewOAuthService(db *gorm.DB, cfg *config.Config) *OAuthService { - svc := &OAuthService{ - db: db, - configs: make(map[OAuthProviderType]*oauth2.Config), - } - svc.configureProviders(cfg) - return svc -} - -// configureProviders sets up OAuth2 configs from application config. -func (s *OAuthService) configureProviders(cfg *config.Config) { - s.mu.Lock() - defer s.mu.Unlock() - - // Google OAuth2 — ref: Chatwoot Google OmniAuth strategy - if cfg.OAuth.Google.ClientID != "" { - s.configs[OAuthProviderGoogle] = &oauth2.Config{ - ClientID: cfg.OAuth.Google.ClientID, - ClientSecret: cfg.OAuth.Google.ClientSecret, - RedirectURL: cfg.OAuth.Google.RedirectURL, - Scopes: []string{"openid", "email", "profile"}, - Endpoint: google.Endpoint, - } - } - - // Generic/GitHub OAuth2 — ref: Chatwoot GitHub OmniAuth strategy - if cfg.OAuth.GitHub.ClientID != "" { - s.configs[OAuthProviderGitHub] = &oauth2.Config{ - ClientID: cfg.OAuth.GitHub.ClientID, - ClientSecret: cfg.OAuth.GitHub.ClientSecret, - RedirectURL: cfg.OAuth.GitHub.RedirectURL, - Scopes: []string{"user:email", "read:user"}, - Endpoint: oauth2.Endpoint{ - AuthURL: "https://github.com/login/oauth/authorize", - TokenURL: "https://github.com/login/oauth/access_token", - }, - } - } -} - -// GetAuthURL generates the OAuth2 authorization URL for a provider. -// This is the first step of the OAuth2 flow: redirect user to provider login. -func (s *OAuthService) GetAuthURL(provider OAuthProviderType, state string) (string, error) { - s.mu.RLock() - config, ok := s.configs[provider] - s.mu.RUnlock() - - if !ok { - return "", fmt.Errorf("oauth provider %s not configured", provider) - } - - return config.AuthCodeURL(state, oauth2.AccessTypeOffline), nil -} - -// ExchangeCode exchanges an OAuth2 authorization code for user info. -// This is the callback step: provider redirects back with a code. -func (s *OAuthService) ExchangeCode(ctx context.Context, provider OAuthProviderType, code string) (*OAuthUserInfo, error) { - s.mu.RLock() - config, ok := s.configs[provider] - s.mu.RUnlock() - - if !ok { - return nil, fmt.Errorf("oauth provider %s not configured", provider) - } - - token, err := config.Exchange(ctx, code) - if err != nil { - return nil, fmt.Errorf("oauth token exchange failed: %w", err) - } - - // Extract user info based on provider - switch provider { - case OAuthProviderGoogle: - return s.extractGoogleUserInfo(ctx, config, token) - case OAuthProviderGitHub: - return s.extractGitHubUserInfo(ctx, config, token) - default: - return nil, fmt.Errorf("unsupported oauth provider: %s", provider) - } -} - -// FindOrCreateUser finds an existing user by OAuth provider+UID, or creates a new one. -// Ref: Chatwoot's after_sign_in_path_for OmniAuth callback logic. -func (s *OAuthService) FindOrCreateUser(info *OAuthUserInfo) (*model.User, error) { - var user model.User - - // Try to find existing user by provider + UID - err := s.db.Where("provider = ? AND uid = ?", string(info.Provider), info.UID).First(&user).Error - if err == nil { - // Update last sign-in info - user.SignInCount++ - s.db.Save(&user) - return &user, nil - } - - // If not found by provider+UID, try by email (link existing account) - if info.Email != "" { - err = s.db.Where("email = ?", info.Email).First(&user).Error - if err == nil { - // Link OAuth to existing account - user.Provider = string(info.Provider) - user.UID = info.UID - if info.AvatarURL != "" && user.AvatarURL == "" { - user.AvatarURL = info.AvatarURL - } - user.SignInCount++ - s.db.Save(&user) - return &user, nil - } - } - - // Create new user from OAuth info - user = model.User{ - Name: info.Name, - Email: info.Email, - Provider: string(info.Provider), - UID: info.UID, - AvatarURL: info.AvatarURL, - Password: "oauth_no_password", // OAuth users don't need a password - } - - if err := s.db.Create(&user).Error; err != nil { - return nil, fmt.Errorf("failed to create oauth user: %w", err) - } - - return &user, nil -} - -// extractGoogleUserInfo extracts user info from Google OAuth2 token. -func (s *OAuthService) extractGoogleUserInfo(ctx context.Context, config *oauth2.Config, token *oauth2.Token) (*OAuthUserInfo, error) { - client := config.Client(ctx, token) - - resp, err := client.Get("https://www.googleapis.com/oauth2/v2/userinfo") - if err != nil { - return nil, fmt.Errorf("failed to fetch google user info: %w", err) - } - defer resp.Body.Close() - - if resp.StatusCode != 200 { - return nil, fmt.Errorf("google userinfo endpoint returned status %d", resp.StatusCode) - } - - var gUser struct { - ID string `json:"id"` - Email string `json:"email"` - Name string `json:"name"` - Picture string `json:"picture"` - } - - if err := json.NewDecoder(resp.Body).Decode(&gUser); err != nil { - return nil, fmt.Errorf("failed to decode google user info: %w", err) - } - - return &OAuthUserInfo{ - Provider: OAuthProviderGoogle, - UID: gUser.ID, - Email: gUser.Email, - Name: gUser.Name, - AvatarURL: gUser.Picture, - }, nil -} - -// extractGitHubUserInfo extracts user info from GitHub OAuth2 token. -func (s *OAuthService) extractGitHubUserInfo(ctx context.Context, config *oauth2.Config, token *oauth2.Token) (*OAuthUserInfo, error) { - client := config.Client(ctx, token) - - resp, err := client.Get("https://api.github.com/user") - if err != nil { - return nil, fmt.Errorf("failed to fetch github user info: %w", err) - } - defer resp.Body.Close() - - if resp.StatusCode != 200 { - return nil, fmt.Errorf("github user endpoint returned status %d", resp.StatusCode) - } - - var ghUser struct { - ID int `json:"id"` - Login string `json:"login"` - Email string `json:"email"` - Name string `json:"name"` - AvatarURL string `json:"avatar_url"` - } - - if err := json.NewDecoder(resp.Body).Decode(&ghUser); err != nil { - return nil, fmt.Errorf("failed to decode github user info: %w", err) - } - - return &OAuthUserInfo{ - Provider: OAuthProviderGitHub, - UID: fmt.Sprintf("%d", ghUser.ID), - Email: ghUser.Email, - Name: ghUser.Name, - AvatarURL: ghUser.AvatarURL, - }, nil -} - -// RegisterGenericProvider allows registering a custom OAuth2 provider at runtime. -// Ref: Chatwoot's configurable OmniAuth strategies in config/initializers/omniauth.rb -func (s *OAuthService) RegisterGenericProvider(name OAuthProviderType, cfg *oauth2.Config) { - s.mu.Lock() - defer s.mu.Unlock() - s.configs[name] = cfg -} - -// IsProviderConfigured checks if a provider has been configured. -func (s *OAuthService) IsProviderConfigured(provider OAuthProviderType) bool { - s.mu.RLock() - defer s.mu.RUnlock() - _, ok := s.configs[provider] - return ok -} \ No newline at end of file diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go index b5baedde..a3a131bd 100644 --- a/backend/internal/config/config.go +++ b/backend/internal/config/config.go @@ -31,7 +31,6 @@ type Config struct { Log LogConfig `mapstructure:"log"` Worker WorkerConfig `mapstructure:"worker"` OAuth OAuthConfig `mapstructure:"oauth"` - RateLimit RateLimitConfig `mapstructure:"rate_limit"` SAML SAMLConfig `mapstructure:"saml"` LDAP LDAPConfig `mapstructure:"ldap"` OIDC OIDCConfig `mapstructure:"oidc"` @@ -72,7 +71,6 @@ type OAuthProviderConfig struct { type OAuthConfig struct { Google OAuthProviderConfig `mapstructure:"google"` - GitHub OAuthProviderConfig `mapstructure:"github"` Twitter OAuthProviderConfig `mapstructure:"twitter"` Microsoft OAuthProviderConfig `mapstructure:"microsoft"` Facebook OAuthProviderConfig `mapstructure:"facebook"` @@ -168,15 +166,6 @@ type LogConfig struct { Format string `mapstructure:"format"` // json, text } -// RateLimitConfig holds rate limiting configuration. -// Uses Redis sliding window counter for production, with in-memory fallback. -// Reference: Chatwoot's Rack::Attack throttle configuration. -type RateLimitConfig struct { - Enabled bool `mapstructure:"enabled"` // enable/disable rate limiting - RequestsPerMinute int `mapstructure:"requests_per_minute"` // max requests per client per window - WindowSeconds int `mapstructure:"window_seconds"` // sliding window duration in seconds -} - // SAMLConfig holds SAML 2.0 Service Provider configuration. // Reference: P2E §1.6 — SAML SP integration for enterprise SSO. type SAMLConfig struct { @@ -318,11 +307,6 @@ func Load() (*Config, error) { viper.SetEnvKeyReplacer(strings.NewReplacer(".", "_")) viper.AutomaticEnv() - // Set defaults for rate limiting - viper.SetDefault("rate_limit.enabled", true) - viper.SetDefault("rate_limit.requests_per_minute", 100) - viper.SetDefault("rate_limit.window_seconds", 60) - // Set defaults for SAML viper.SetDefault("saml.enabled", false) viper.SetDefault("saml.clock_drift_tolerance", 180) @@ -401,12 +385,6 @@ func Load() (*Config, error) { } // Apply defaults for zero-valued fields (viper may not set defaults for already-present keys) - if cfg.RateLimit.RequestsPerMinute == 0 { - cfg.RateLimit.RequestsPerMinute = 100 - } - if cfg.RateLimit.WindowSeconds == 0 { - cfg.RateLimit.WindowSeconds = 60 - } applySearchDefaults(&cfg.Search) return &cfg, nil @@ -447,9 +425,6 @@ type ConfigReloader struct { var ReloadableFields = []string{ "log.level", "log.format", - "rate_limit.enabled", - "rate_limit.requests_per_minute", - "rate_limit.window_seconds", "worker.concurrency", "worker.redis_block_timeout_s", "worker.redis_sweep_interval_s", @@ -487,12 +462,6 @@ func (r *ConfigReloader) handleConfigChange(e fsnotify.Event) { } // Apply defaults for zero-valued fields (same logic as Load()) - if newCfg.RateLimit.RequestsPerMinute == 0 { - newCfg.RateLimit.RequestsPerMinute = 100 - } - if newCfg.RateLimit.WindowSeconds == 0 { - newCfg.RateLimit.WindowSeconds = 60 - } applySearchDefaults(&newCfg.Search) // Validate the entire new config — if invalid, skip the reload @@ -521,9 +490,6 @@ func (r *ConfigReloader) applyReloadableFields(newCfg *Config) { // Log settings — safe to change at runtime r.cfg.Log = newCfg.Log - // Rate limit settings — safe to change at runtime - r.cfg.RateLimit = newCfg.RateLimit - // Worker concurrency — safe to change at runtime r.cfg.Worker.Concurrency = newCfg.Worker.Concurrency // Worker Redis sweep/block timing — safe to change at runtime @@ -618,8 +584,6 @@ func LoadWithEnv(env string) (*Config, error) { "GOCHAT_STORAGE_PROVIDER": "storage.provider", "GOCHAT_STORAGE_LOCAL_PATH": "storage.local_path", "GOCHAT_STORAGE_MAX_FILE_SIZE": "storage.max_file_size", - "GOCHAT_RATE_LIMIT_REQUESTS_PER_MINUTE": "rate_limit.requests_per_minute", - "GOCHAT_RATE_LIMIT_WINDOW_SECONDS": "rate_limit.window_seconds", // G10: OAuth config for new channel integrations (Twitter, Microsoft, Google) "GOCHAT_OAUTH_TWITTER_CLIENT_ID": "oauth.twitter.client_id", "GOCHAT_OAUTH_TWITTER_CLIENT_SECRET": "oauth.twitter.client_secret", @@ -793,10 +757,6 @@ func setDefaults(v *viper.Viper) { v.SetDefault("log.level", "debug") v.SetDefault("log.format", "json") - v.SetDefault("rate_limit.enabled", true) - v.SetDefault("rate_limit.requests_per_minute", 100) - v.SetDefault("rate_limit.window_seconds", 60) - v.SetDefault("search.engine", "meilisearch") v.SetDefault("search.host", "http://localhost:7700") v.SetDefault("search.api_key", "") @@ -850,12 +810,6 @@ func setDefaults(v *viper.Viper) { // applyZeroDefaults fills in defaults for zero-valued fields that viper may not set. func applyZeroDefaults(cfg *Config) { - if cfg.RateLimit.RequestsPerMinute == 0 { - cfg.RateLimit.RequestsPerMinute = 100 - } - if cfg.RateLimit.WindowSeconds == 0 { - cfg.RateLimit.WindowSeconds = 60 - } if cfg.Worker.Concurrency == 0 { cfg.Worker.Concurrency = 4 } diff --git a/backend/internal/config/config_test.go b/backend/internal/config/config_test.go index 68a10b73..a41b35ef 100644 --- a/backend/internal/config/config_test.go +++ b/backend/internal/config/config_test.go @@ -29,7 +29,6 @@ func TestValidate_ValidConfig(t *testing.T) { Log: LogConfig{Level: "info", Format: "json"}, Worker: WorkerConfig{Concurrency: 4, BlockTimeoutS: 5, SweepIntervalS: 30}, OAuth: OAuthConfig{}, - RateLimit: RateLimitConfig{Enabled: true, RequestsPerMinute: 100, WindowSeconds: 60}, Search: SearchConfig{Engine: "meilisearch", Host: "http://localhost:7700", IndexPrefix: "gochat_", TimeoutSeconds: 5}, } @@ -164,7 +163,6 @@ func TestValidate_SearchMeilisearchRequiresValidHost(t *testing.T) { JWT: JWTConfig{Secret: "test-secret-key-min-32-chars!!"}, Log: LogConfig{Level: "info"}, Worker: WorkerConfig{Concurrency: 1, BlockTimeoutS: 5, SweepIntervalS: 30}, - RateLimit: RateLimitConfig{RequestsPerMinute: 100, WindowSeconds: 60}, Search: SearchConfig{Engine: "meilisearch", Host: "not a url", TimeoutSeconds: 5}, } @@ -181,7 +179,6 @@ func TestValidate_SearchDBFallbackAllowed(t *testing.T) { JWT: JWTConfig{Secret: "test-secret-key-min-32-chars!!"}, Log: LogConfig{Level: "info"}, Worker: WorkerConfig{Concurrency: 1, BlockTimeoutS: 5, SweepIntervalS: 30}, - RateLimit: RateLimitConfig{RequestsPerMinute: 100, WindowSeconds: 60}, Search: SearchConfig{Engine: "db"}, } @@ -197,7 +194,6 @@ func TestValidate_SearchDBFallbackRejectedInRelease(t *testing.T) { JWT: JWTConfig{Secret: "test-secret-key-min-32-chars!!"}, Log: LogConfig{Level: "info"}, Worker: WorkerConfig{Concurrency: 1, BlockTimeoutS: 5, SweepIntervalS: 30}, - RateLimit: RateLimitConfig{RequestsPerMinute: 100, WindowSeconds: 60}, Search: SearchConfig{Engine: "db"}, } diff --git a/backend/internal/config/reloader_test.go b/backend/internal/config/reloader_test.go index cbf1b7bf..4360692b 100644 --- a/backend/internal/config/reloader_test.go +++ b/backend/internal/config/reloader_test.go @@ -62,9 +62,6 @@ func TestConfigReloader_ReloadableFieldsList(t *testing.T) { expectedFields := []string{ "log.level", "log.format", - "rate_limit.enabled", - "rate_limit.requests_per_minute", - "rate_limit.window_seconds", "worker.concurrency", "worker.redis_block_timeout_s", "worker.redis_sweep_interval_s", @@ -76,13 +73,11 @@ func TestConfigReloader_ReloadableFieldsList(t *testing.T) { func TestConfigReloader_ApplyReloadableFields(t *testing.T) { oldCfg := validTestConfig() oldCfg.Log.Level = "debug" - oldCfg.RateLimit.RequestsPerMinute = 100 oldCfg.Worker.Concurrency = 4 oldCfg.Database.Host = "original-host" // immutable field newCfg := validTestConfig() newCfg.Log.Level = "info" - newCfg.RateLimit.RequestsPerMinute = 200 newCfg.Worker.Concurrency = 8 newCfg.Database.Host = "changed-host" // should NOT be applied @@ -91,7 +86,6 @@ func TestConfigReloader_ApplyReloadableFields(t *testing.T) { // Reloadable fields should be updated assert.Equal(t, "info", r.cfg.Log.Level) - assert.Equal(t, 200, r.cfg.RateLimit.RequestsPerMinute) assert.Equal(t, 8, r.cfg.Worker.Concurrency) // Immutable fields should NOT be updated @@ -148,10 +142,6 @@ jwt: log: level: "debug" format: "json" -rate_limit: - enabled: true - requests_per_minute: 100 - window_seconds: 60 worker: concurrency: 4 ` @@ -181,10 +171,10 @@ func TestLoadWithEnv_EnvironmentOverlay(t *testing.T) { } func TestLoadWithEnv_DotEnvFile(t *testing.T) { - key, value, ok := parseDotEnvLine("GOCHAT_RATE_LIMIT_REQUESTS_PER_MINUTE=50 # requests per minute") + key, value, ok := parseDotEnvLine("GOCHAT_LOG_LEVEL=debug # log level") require.True(t, ok) - assert.Equal(t, "GOCHAT_RATE_LIMIT_REQUESTS_PER_MINUTE", key) - assert.Equal(t, "50", value) + assert.Equal(t, "GOCHAT_LOG_LEVEL", key) + assert.Equal(t, "debug", value) _, _, ok = parseDotEnvLine("# Comment line should be ignored") assert.False(t, ok) @@ -232,9 +222,6 @@ func TestSetDefaults(t *testing.T) { assert.Equal(t, 6379, v.GetInt("redis.port")) assert.Equal(t, "debug", v.GetString("log.level")) assert.Equal(t, "json", v.GetString("log.format")) - assert.Equal(t, true, v.GetBool("rate_limit.enabled")) - assert.Equal(t, 100, v.GetInt("rate_limit.requests_per_minute")) - assert.Equal(t, 60, v.GetInt("rate_limit.window_seconds")) assert.Equal(t, 4, v.GetInt("worker.concurrency")) } @@ -242,18 +229,13 @@ func TestApplyZeroDefaults(t *testing.T) { cfg := Config{} // all zeros applyZeroDefaults(&cfg) - assert.Equal(t, 100, cfg.RateLimit.RequestsPerMinute) - assert.Equal(t, 60, cfg.RateLimit.WindowSeconds) assert.Equal(t, 4, cfg.Worker.Concurrency) // Non-zero values should not be overwritten cfg2 := Config{ - RateLimit: RateLimitConfig{RequestsPerMinute: 200, WindowSeconds: 30}, - Worker: WorkerConfig{Concurrency: 16}, + Worker: WorkerConfig{Concurrency: 16}, } applyZeroDefaults(&cfg2) - assert.Equal(t, 200, cfg2.RateLimit.RequestsPerMinute) - assert.Equal(t, 30, cfg2.RateLimit.WindowSeconds) assert.Equal(t, 16, cfg2.Worker.Concurrency) } @@ -268,6 +250,5 @@ func validTestConfig() *Config { Log: LogConfig{Level: "info", Format: "json"}, Worker: WorkerConfig{Concurrency: 4}, OAuth: OAuthConfig{}, - RateLimit: RateLimitConfig{Enabled: true, RequestsPerMinute: 100, WindowSeconds: 60}, } } diff --git a/backend/internal/config/validator.go b/backend/internal/config/validator.go index f68a83cf..aafae9d1 100644 --- a/backend/internal/config/validator.go +++ b/backend/internal/config/validator.go @@ -69,14 +69,6 @@ func Validate(cfg *Config) error { return fmt.Errorf("worker.redis_sweep_interval_s must be >= 1") } - // Rate limit validation - if cfg.RateLimit.RequestsPerMinute < 1 { - return fmt.Errorf("rate_limit.requests_per_minute must be >= 1") - } - if cfg.RateLimit.WindowSeconds < 1 { - return fmt.Errorf("rate_limit.window_seconds must be >= 1") - } - // Search validation. Meilisearch is the production parity engine; db remains // available only as an explicit development fallback. engine := strings.ToLower(cfg.Search.Engine) diff --git a/backend/internal/handler/api/v1/auth_handler.go b/backend/internal/handler/api/v1/auth_handler.go index ab925781..14c6e47a 100644 --- a/backend/internal/handler/api/v1/auth_handler.go +++ b/backend/internal/handler/api/v1/auth_handler.go @@ -10,7 +10,6 @@ import ( "github.com/gin-gonic/gin" - "github.com/gochat/gochat/internal/auth" "github.com/gochat/gochat/internal/service" "github.com/gochat/gochat/pkg/response" ) @@ -30,19 +29,17 @@ import ( // AuthHandler handles authentication HTTP endpoints. type AuthHandler struct { authService *service.AuthService - oauthService *auth.OAuthService profileService *service.ProfileService } // NewAuthHandler creates an auth handler with service dependencies. -func NewAuthHandler(authService *service.AuthService, oauthService *auth.OAuthService, profileService ...*service.ProfileService) *AuthHandler { +func NewAuthHandler(authService *service.AuthService, profileService ...*service.ProfileService) *AuthHandler { var profileSvc *service.ProfileService if len(profileService) > 0 { profileSvc = profileService[0] } return &AuthHandler{ authService: authService, - oauthService: oauthService, profileService: profileSvc, } } @@ -86,13 +83,6 @@ type ConfirmEmailRequest struct { ConfirmationToken string `json:"confirmation_token" binding:"required"` } -// OAuthCallbackRequest is the JSON body for OAuth callback. -type OAuthCallbackRequest struct { - Provider string `json:"provider" binding:"required"` - Code string `json:"code" binding:"required"` - State string `json:"state"` -} - // --- Handlers --- // Login authenticates a user with email/password and returns JWT tokens. @@ -426,73 +416,6 @@ func (h *AuthHandler) ChatwootConfirmEmail(c *gin.Context) { c.JSON(http.StatusOK, gin.H{"data": data}) } -// OAuthCallback handles OAuth2 provider callback. -// POST /api/v1/auth/oauth/callback -// Receives provider + code from frontend (frontend handles redirect flow). -func (h *AuthHandler) OAuthCallback(c *gin.Context) { - var req OAuthCallbackRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrValidation, err.Error()) - return - } - - provider := auth.OAuthProviderType(req.Provider) - if !h.oauthService.IsProviderConfigured(provider) { - response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "OAuth provider not configured: "+req.Provider) - return - } - - output, err := h.authService.OAuthLogin(c.Request.Context(), &service.OAuthLoginInput{ - Provider: provider, - Code: req.Code, - State: req.State, - }) - 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, - "is_new_user": output.IsNewUser, - }) -} - -// OAuthAuthorize generates the OAuth2 authorization URL for a provider. -// GET /api/v1/auth/oauth/authorize?provider=google -// Frontend redirects user to this URL to start OAuth flow. -func (h *AuthHandler) OAuthAuthorize(c *gin.Context) { - providerStr := c.Query("provider") - provider := auth.OAuthProviderType(providerStr) - - if !h.oauthService.IsProviderConfigured(provider) { - response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "OAuth provider not configured: "+providerStr) - return - } - - // Generate state for CSRF protection (store in Redis for validation) - state := c.Query("state") - if state == "" { - state = generateOAuthState() - } - - url, err := h.oauthService.GetAuthURL(provider, state) - if err != nil { - response.AbortWithStatusError(c, http.StatusInternalServerError, response.ErrInternal, err.Error()) - return - } - - response.OK(c, gin.H{ - "authorize_url": url, - "state": state, - }) -} - // RegisterAuthRoutes sets up auth routes on a Gin router group. // These routes are PUBLIC — no AuthRequired middleware. func RegisterAuthRoutes(rg *gin.RouterGroup, handler *AuthHandler) { @@ -511,10 +434,6 @@ func RegisterAuthRoutes(rg *gin.RouterGroup, handler *AuthHandler) { authGroup.POST("/reset_password", handler.ResetPassword) authGroup.PUT("/reset_password", handler.ConfirmResetPassword) authGroup.GET("/confirm_email", handler.ConfirmEmail) - - // OAuth - authGroup.GET("/oauth/authorize", handler.OAuthAuthorize) - authGroup.POST("/oauth/callback", handler.OAuthCallback) } } @@ -568,8 +487,9 @@ func extractChatwootAccessToken(c *gin.Context) string { return "" } -// generateOAuthState creates a cryptographically random state token for OAuth CSRF protection. -// Production note: state should also be stored server-side (Redis) and validated on callback. +// generateOAuthState creates a cryptographically random state token for CSRF protection. +// Used by SAML 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) } diff --git a/backend/internal/handler/api/v1/auth_handler_test.go b/backend/internal/handler/api/v1/auth_handler_test.go index 7a7b1ddd..e35eef9a 100644 --- a/backend/internal/handler/api/v1/auth_handler_test.go +++ b/backend/internal/handler/api/v1/auth_handler_test.go @@ -56,9 +56,9 @@ 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, nil) + authSvc := service.NewAuthService(db, jwtSvc, refreshStore, nil) profileSvc := service.NewProfileService(repository.NewUserRepo(db), repository.NewAccountUserRepo(db), repository.NewAccessTokenRepo(db)) - handler := NewAuthHandler(authSvc, nil, profileSvc) + handler := NewAuthHandler(authSvc, profileSvc) router := gin.New() RegisterChatwootAuthRoutes(router.Group("/auth"), handler) diff --git a/backend/internal/middleware/rate_limit.go b/backend/internal/middleware/rate_limit.go index cc8b78c3..8be11487 100644 --- a/backend/internal/middleware/rate_limit.go +++ b/backend/internal/middleware/rate_limit.go @@ -12,7 +12,6 @@ import ( "github.com/gin-gonic/gin" "github.com/redis/go-redis/v9" - "github.com/gochat/gochat/internal/config" "github.com/gochat/gochat/pkg/logger" "github.com/gochat/gochat/pkg/response" ) @@ -20,6 +19,10 @@ import ( const ( // rateLimitKeyPrefix is the Redis key prefix for rate limit counters. rateLimitKeyPrefix = "gochat:rate_limit:" + + // Hardcoded rate limit defaults (no longer configurable). + rateLimitRequestsPerMinute = 100 + rateLimitWindowSeconds = 60 ) // slidingWindowLimiter implements Redis-based sliding window counter rate limiting. @@ -27,7 +30,6 @@ const ( // Reference: Chatwoot's Rack::Attack throttle configuration. type slidingWindowLimiter struct { redis *redis.Client - cfg *config.RateLimitConfig fallback *inMemoryLimiter redisAvailable atomic.Bool } @@ -117,19 +119,12 @@ func (im *inMemoryLimiter) remainingInMemory(key string) int { // newSlidingWindowLimiter creates a new rate limiter with Redis sliding window counter // and in-memory fallback. -func newSlidingWindowLimiter(rdb *redis.Client, cfg *config.RateLimitConfig) *slidingWindowLimiter { - limit := cfg.RequestsPerMinute - if limit <= 0 { - limit = 100 - } - windowSecs := cfg.WindowSeconds - if windowSecs <= 0 { - windowSecs = 60 - } +func newSlidingWindowLimiter(rdb *redis.Client) *slidingWindowLimiter { + limit := rateLimitRequestsPerMinute + windowSecs := rateLimitWindowSeconds sw := &slidingWindowLimiter{ redis: rdb, - cfg: cfg, fallback: newInMemoryLimiter(limit, time.Duration(windowSecs)*time.Second), } @@ -161,14 +156,8 @@ func newSlidingWindowLimiter(rdb *redis.Client, cfg *config.RateLimitConfig) *sl // 4. If total count > limit, reject; otherwise allow and increment current window func (sw *slidingWindowLimiter) checkRedis(ctx context.Context, key string) (bool, int, int, error) { now := time.Now() - windowSecs := sw.cfg.WindowSeconds - if windowSecs <= 0 { - windowSecs = 60 - } - limit := sw.cfg.RequestsPerMinute - if limit <= 0 { - limit = 100 - } + windowSecs := rateLimitWindowSeconds + limit := rateLimitRequestsPerMinute currentWindow := now.Unix() / int64(windowSecs) previousWindow := currentWindow - 1 @@ -275,16 +264,9 @@ func (sw *slidingWindowLimiter) recheckRedis() { // RateLimit creates a Redis-based sliding window counter rate limiting middleware. // Falls back to in-memory rate limiting when Redis is unavailable. // Corresponds to Chatwoot's Rack::Attack throttle configuration. -func RateLimit(cfg *config.Config, rdb *redis.Client) gin.HandlerFunc { - if !cfg.RateLimit.Enabled { - return func(c *gin.Context) { c.Next() } - } - - sw := newSlidingWindowLimiter(rdb, &cfg.RateLimit) - limit := cfg.RateLimit.RequestsPerMinute - if limit <= 0 { - limit = 100 - } +func RateLimit(rdb *redis.Client) gin.HandlerFunc { + sw := newSlidingWindowLimiter(rdb) + limit := rateLimitRequestsPerMinute return func(c *gin.Context) { if isRateLimitExemptPath(c.Request.URL.Path) { @@ -296,7 +278,6 @@ func RateLimit(cfg *config.Config, rdb *redis.Client) gin.HandlerFunc { key := "global:" + ip allowed, _, remaining, usedRedis := sw.check(c.Request.Context(), key) - // Set rate limit headers c.Header("X-RateLimit-Limit", strconv.Itoa(limit)) c.Header("X-RateLimit-Remaining", strconv.Itoa(remaining)) @@ -307,7 +288,7 @@ func RateLimit(cfg *config.Config, rdb *redis.Client) gin.HandlerFunc { } if !allowed { - c.Header("Retry-After", strconv.Itoa(cfg.RateLimit.WindowSeconds)) + c.Header("Retry-After", strconv.Itoa(rateLimitWindowSeconds)) c.AbortWithStatusJSON(http.StatusTooManyRequests, response.APIResponse{ Success: false, Error: &response.ErrorBody{ diff --git a/backend/internal/middleware/rate_limit_test.go b/backend/internal/middleware/rate_limit_test.go index 0c2d74fe..20187e2e 100644 --- a/backend/internal/middleware/rate_limit_test.go +++ b/backend/internal/middleware/rate_limit_test.go @@ -13,7 +13,6 @@ import ( "github.com/redis/go-redis/v9" "github.com/stretchr/testify/assert" - "github.com/gochat/gochat/internal/config" "github.com/gochat/gochat/pkg/logger" "github.com/gochat/gochat/pkg/response" ) @@ -191,21 +190,12 @@ func setupMiniredis(t *testing.T) (*miniredis.Miniredis, *redis.Client) { return mr, rdb } -func defaultRateLimitConfig() *config.RateLimitConfig { - return &config.RateLimitConfig{ - Enabled: true, - RequestsPerMinute: 10, - WindowSeconds: 60, - } -} - func TestNewSlidingWindowLimiter_WithRedis(t *testing.T) { // 有可用Redis时应标记Redis可用 mr, rdb := setupMiniredis(t) defer mr.Close() - cfg := defaultRateLimitConfig() - sw := newSlidingWindowLimiter(rdb, cfg) + sw := newSlidingWindowLimiter(rdb) assert.NotNil(t, sw) assert.Equal(t, rdb, sw.redis) @@ -215,8 +205,7 @@ func TestNewSlidingWindowLimiter_WithRedis(t *testing.T) { func TestNewSlidingWindowLimiter_NoRedis(t *testing.T) { // 无Redis客户端时应标记Redis不可用 - cfg := defaultRateLimitConfig() - sw := newSlidingWindowLimiter(nil, cfg) + sw := newSlidingWindowLimiter(nil) assert.NotNil(t, sw) assert.Nil(t, sw.redis) @@ -224,19 +213,14 @@ func TestNewSlidingWindowLimiter_NoRedis(t *testing.T) { } func TestNewSlidingWindowLimiter_DefaultValues(t *testing.T) { - // 配额/窗口为0时应使用默认值 + // 硬编码常量应为100次/60秒 mr, rdb := setupMiniredis(t) defer mr.Close() - cfg := &config.RateLimitConfig{ - Enabled: true, - RequestsPerMinute: 0, // 应默认为100 - WindowSeconds: 0, // 应默认为60 - } - sw := newSlidingWindowLimiter(rdb, cfg) + sw := newSlidingWindowLimiter(rdb) assert.NotNil(t, sw) - // 检查fallback使用了默认值 + // 检查fallback使用了硬编码默认值 assert.Equal(t, 100, sw.fallback.limit) assert.Equal(t, 60*time.Second, sw.fallback.window) } @@ -246,8 +230,7 @@ func TestCheckRedis_FirstRequest(t *testing.T) { mr, rdb := setupMiniredis(t) defer mr.Close() - cfg := defaultRateLimitConfig() - sw := newSlidingWindowLimiter(rdb, cfg) + sw := newSlidingWindowLimiter(rdb) allowed, count, remaining, err := sw.checkRedis(context.Background(), "global:127.0.0.1") assert.NoError(t, err) @@ -261,10 +244,9 @@ func TestCheckRedis_WithinLimit(t *testing.T) { mr, rdb := setupMiniredis(t) defer mr.Close() - cfg := defaultRateLimitConfig() // 10次/分钟 - sw := newSlidingWindowLimiter(rdb, cfg) + sw := newSlidingWindowLimiter(rdb) // 100次/分钟(硬编码) - for i := 0; i < 10; i++ { + for i := 0; i < 100; i++ { allowed, _, _, err := sw.checkRedis(context.Background(), "global:127.0.0.1") assert.NoError(t, err) assert.True(t, allowed, "第%d次请求应被允许", i+1) @@ -274,33 +256,28 @@ func TestCheckRedis_WithinLimit(t *testing.T) { func TestCheckRedis_OverLimit(t *testing.T) { // 滑动窗口Lua脚本中:total >= limit时拒绝(不INCR,返回{total, 0}) // 但Go端allowed判定为 totalCount <= limit,当total==limit时allowed仍为true - // 因此limit=5时,5次请求INCR后total=5,第6次请求total=5>=5被Lua拒绝(不INCR) - // 但Go端5<=5仍返回allowed=true,remaining=0 + // 因此limit=100时,100次请求INCR后total=100,第101次请求total=100>=100被Lua拒绝(不INCR) + // 但Go端100<=100仍返回allowed=true,remaining=0 // 实际效果:请求通过但remaining=0标识配额耗尽 mr, rdb := setupMiniredis(t) defer mr.Close() - cfg := &config.RateLimitConfig{ - Enabled: true, - RequestsPerMinute: 5, - WindowSeconds: 60, - } - sw := newSlidingWindowLimiter(rdb, cfg) + sw := newSlidingWindowLimiter(rdb) - // 前5次请求允许,total逐步增加到5 - for i := 0; i < 5; i++ { + // 前100次请求允许,total逐步增加到100 + for i := 0; i < 100; i++ { allowed, _, _, err := sw.checkRedis(context.Background(), "global:127.0.0.1") assert.NoError(t, err) assert.True(t, allowed, "第%d次请求应被允许", i+1) } - // 第6次请求:Lua脚本中total=5>=5拒绝(不INCR),Go端totalCount=5<=5返回allowed=true + // 第101次请求:Lua脚本中total=100>=100拒绝(不INCR),Go端totalCount=100<=100返回allowed=true // 但remaining=0标识配额已耗尽 allowed, totalCount, remaining, err := sw.checkRedis(context.Background(), "global:127.0.0.1") assert.NoError(t, err) // 注意:这是源码的现有行为,totalCount==limit时allowed仍为true assert.True(t, allowed) - assert.Equal(t, 5, totalCount) + assert.Equal(t, 100, totalCount) assert.Equal(t, 0, remaining) } @@ -309,24 +286,14 @@ func TestCheckRedis_DifferentKeys(t *testing.T) { mr, rdb := setupMiniredis(t) defer mr.Close() - cfg := &config.RateLimitConfig{ - Enabled: true, - RequestsPerMinute: 2, - WindowSeconds: 60, - } - sw := newSlidingWindowLimiter(rdb, cfg) + sw := newSlidingWindowLimiter(rdb) - // key1耗尽限额(2次请求后total=2,第3次请求Lua拒绝但Go端allowed仍为true) + // key1发送2次请求 sw.checkRedis(context.Background(), "global:client1") sw.checkRedis(context.Background(), "global:client1") - allowed, _, remaining, err := sw.checkRedis(context.Background(), "global:client1") - assert.NoError(t, err) - // 源码行为:total=2<=limit=2,allowed=true但remaining=0 - assert.True(t, allowed) - assert.Equal(t, 0, remaining) // key2首次请求应被允许,remaining > 0 - allowed, _, remaining, err = sw.checkRedis(context.Background(), "global:client2") + allowed, _, remaining, err := sw.checkRedis(context.Background(), "global:client2") assert.NoError(t, err) assert.True(t, allowed) assert.True(t, remaining > 0) @@ -337,22 +304,17 @@ func TestCheckRedis_RemainingDecrements(t *testing.T) { mr, rdb := setupMiniredis(t) defer mr.Close() - cfg := &config.RateLimitConfig{ - Enabled: true, - RequestsPerMinute: 5, - WindowSeconds: 60, - } - sw := newSlidingWindowLimiter(rdb, cfg) + sw := newSlidingWindowLimiter(rdb) _, _, remaining, err := sw.checkRedis(context.Background(), "global:client1") assert.NoError(t, err) - // 第一次请求后:new_total=1, remaining=5-1=4 - assert.Equal(t, 4, remaining) + // 第一次请求后:new_total=1, remaining=100-1=99 + assert.Equal(t, 99, remaining) _, _, remaining, err = sw.checkRedis(context.Background(), "global:client1") assert.NoError(t, err) - // 第二次请求后:new_total=2, remaining=5-2=3 - assert.Equal(t, 3, remaining) + // 第二次请求后:new_total=2, remaining=100-2=98 + assert.Equal(t, 98, remaining) } // ============================================================================ @@ -364,8 +326,7 @@ func TestCheck_RedisAvailable(t *testing.T) { mr, rdb := setupMiniredis(t) defer mr.Close() - cfg := defaultRateLimitConfig() - sw := newSlidingWindowLimiter(rdb, cfg) + sw := newSlidingWindowLimiter(rdb) allowed, _, remaining, usedRedis := sw.check(context.Background(), "global:127.0.0.1") assert.True(t, allowed) @@ -375,8 +336,7 @@ func TestCheck_RedisAvailable(t *testing.T) { func TestCheck_RedisUnavailable_Fallback(t *testing.T) { // Redis不可用时应降级到内存限流,usedRedis=false - cfg := defaultRateLimitConfig() - sw := newSlidingWindowLimiter(nil, cfg) // nil Redis + sw := newSlidingWindowLimiter(nil) // nil Redis allowed, _, remaining, usedRedis := sw.check(context.Background(), "global:127.0.0.1") assert.True(t, allowed) @@ -387,8 +347,7 @@ func TestCheck_RedisUnavailable_Fallback(t *testing.T) { func TestCheck_RedisError_Fallback(t *testing.T) { // Redis连接出错时应降级到内存限流 mr, rdb := setupMiniredis(t) - cfg := defaultRateLimitConfig() - sw := newSlidingWindowLimiter(rdb, cfg) + sw := newSlidingWindowLimiter(rdb) // 先确认Redis正常工作 allowed, _, _, usedRedis := sw.check(context.Background(), "global:127.0.0.1") @@ -406,21 +365,16 @@ func TestCheck_RedisError_Fallback(t *testing.T) { func TestCheck_FallbackOverLimit(t *testing.T) { // 内存降级限流器超限后也应拒绝请求 - cfg := &config.RateLimitConfig{ - Enabled: true, - RequestsPerMinute: 3, - WindowSeconds: 60, - } - sw := newSlidingWindowLimiter(nil, cfg) + sw := newSlidingWindowLimiter(nil) - // 前3次请求允许 - for i := 0; i < 3; i++ { + // 前100次请求允许 + for i := 0; i < 100; i++ { allowed, _, _, usedRedis := sw.check(context.Background(), "client1") assert.True(t, allowed) assert.False(t, usedRedis) } - // 第4次请求应被拒绝 + // 第101次请求应被拒绝 allowed, _, _, usedRedis := sw.check(context.Background(), "client1") assert.False(t, allowed) assert.False(t, usedRedis) @@ -448,34 +402,12 @@ func makeRequestWithMiddleware(handler gin.HandlerFunc, path string) *httptest.R return w } -func TestRateLimit_Disabled(t *testing.T) { - // 限流未启用时所有请求应通过 - cfg := &config.Config{ - RateLimit: config.RateLimitConfig{ - Enabled: false, - }, - } - - handler := RateLimit(cfg, nil) - w := makeRequestWithMiddleware(handler, "/test") - - assert.Equal(t, http.StatusOK, w.Code) -} - func TestRateLimit_FirstRequestAllowed(t *testing.T) { // 首次请求应通过并设置限流响应头 mr, rdb := setupMiniredis(t) defer mr.Close() - cfg := &config.Config{ - RateLimit: config.RateLimitConfig{ - Enabled: true, - RequestsPerMinute: 10, - WindowSeconds: 60, - }, - } - - handler := RateLimit(cfg, rdb) + handler := RateLimit(rdb) w := makeRequestWithMiddleware(handler, "/test") assert.Equal(t, http.StatusOK, w.Code) @@ -487,15 +419,7 @@ func TestRateLimit_FirstRequestAllowed(t *testing.T) { func TestRateLimit_MemoryFallback(t *testing.T) { // Redis不可用时应降级到内存限流,backend标记为memory - cfg := &config.Config{ - RateLimit: config.RateLimitConfig{ - Enabled: true, - RequestsPerMinute: 100, - WindowSeconds: 60, - }, - } - - handler := RateLimit(cfg, nil) // nil Redis → 内存降级 + handler := RateLimit(nil) // nil Redis → 内存降级 w := makeRequestWithMiddleware(handler, "/test") assert.Equal(t, http.StatusOK, w.Code) @@ -509,15 +433,7 @@ func TestRateLimit_OverLimit_RedisBackend(t *testing.T) { mr, rdb := setupMiniredis(t) defer mr.Close() - cfg := &config.Config{ - RateLimit: config.RateLimitConfig{ - Enabled: true, - RequestsPerMinute: 2, - WindowSeconds: 60, - }, - } - - handler := RateLimit(cfg, rdb) + handler := RateLimit(rdb) setupGin() r := gin.New() r.Use(handler) @@ -525,15 +441,15 @@ func TestRateLimit_OverLimit_RedisBackend(t *testing.T) { c.JSON(http.StatusOK, gin.H{"message": "ok"}) }) - // 前2次请求通过 - for i := 0; i < 2; i++ { + // 前100次请求通过 + for i := 0; i < 100; i++ { w := httptest.NewRecorder() req, _ := http.NewRequest("GET", "/test", nil) r.ServeHTTP(w, req) assert.Equal(t, http.StatusOK, w.Code) } - // 第3次请求:由于源码行为(total==limit时allowed仍为true),不会返回429 + // 第101次请求:由于源码行为(total==limit时allowed仍为true),不会返回429 // 但remaining=0标识配额耗尽 w := httptest.NewRecorder() req, _ := http.NewRequest("GET", "/test", nil) @@ -546,15 +462,7 @@ func TestRateLimit_OverLimit_RedisBackend(t *testing.T) { func TestRateLimit_OverLimit_MemoryBackend(t *testing.T) { // 使用内存降级后端时,超限请求应返回429 // 内存限流器的checkInMemory在超限后正确返回allowed=false - cfg := &config.Config{ - RateLimit: config.RateLimitConfig{ - Enabled: true, - RequestsPerMinute: 2, - WindowSeconds: 60, - }, - } - - handler := RateLimit(cfg, nil) // nil Redis → 内存降级 + handler := RateLimit(nil) // nil Redis → 内存降级 setupGin() r := gin.New() r.Use(handler) @@ -562,15 +470,15 @@ func TestRateLimit_OverLimit_MemoryBackend(t *testing.T) { c.JSON(http.StatusOK, gin.H{"message": "ok"}) }) - // 前2次请求通过 - for i := 0; i < 2; i++ { + // 前100次请求通过 + for i := 0; i < 100; i++ { w := httptest.NewRecorder() req, _ := http.NewRequest("GET", "/test", nil) r.ServeHTTP(w, req) assert.Equal(t, http.StatusOK, w.Code) } - // 第3次请求应被拒绝(429) + // 第101次请求应被拒绝(429) w := httptest.NewRecorder() req, _ := http.NewRequest("GET", "/test", nil) r.ServeHTTP(w, req) @@ -591,15 +499,7 @@ func TestRateLimit_OverLimit_MemoryBackend(t *testing.T) { func TestRateLimit_DifferentIPs(t *testing.T) { // 不同IP应有独立的限流计数 // 使用内存降级后端以确保能正确拒绝超限请求 - cfg := &config.Config{ - RateLimit: config.RateLimitConfig{ - Enabled: true, - RequestsPerMinute: 2, - WindowSeconds: 60, - }, - } - - handler := RateLimit(cfg, nil) // 内存降级后端 + handler := RateLimit(nil) // 内存降级后端 setupGin() r := gin.New() r.Use(handler) @@ -608,7 +508,7 @@ func TestRateLimit_DifferentIPs(t *testing.T) { }) // IP1耗尽限额 - for i := 0; i < 2; i++ { + for i := 0; i < 100; i++ { w := httptest.NewRecorder() req, _ := http.NewRequest("GET", "/test", nil) req.RemoteAddr = "192.168.1.1:1234" @@ -634,15 +534,7 @@ func TestRateLimit_HeadersSet(t *testing.T) { mr, rdb := setupMiniredis(t) defer mr.Close() - cfg := &config.Config{ - RateLimit: config.RateLimitConfig{ - Enabled: true, - RequestsPerMinute: 100, - WindowSeconds: 60, - }, - } - - handler := RateLimit(cfg, rdb) + handler := RateLimit(rdb) w := makeRequestWithMiddleware(handler, "/test") assert.Equal(t, "100", w.Header().Get("X-RateLimit-Limit")) diff --git a/backend/internal/service/auth_service.go b/backend/internal/service/auth_service.go index e19f6a30..dd07df82 100644 --- a/backend/internal/service/auth_service.go +++ b/backend/internal/service/auth_service.go @@ -29,7 +29,6 @@ type AuthService struct { db *gorm.DB jwtService *auth.JWTService refreshStore *auth.RefreshTokenStore - oauthService *auth.OAuthService mfaService *auth.MFAService } @@ -38,14 +37,12 @@ func NewAuthService( db *gorm.DB, jwtService *auth.JWTService, refreshStore *auth.RefreshTokenStore, - oauthService *auth.OAuthService, mfaService *auth.MFAService, ) *AuthService { return &AuthService{ db: db, jwtService: jwtService, refreshStore: refreshStore, - oauthService: oauthService, mfaService: mfaService, } } @@ -543,76 +540,6 @@ func (s *AuthService) ConfirmEmail(ctx context.Context, input *ConfirmEmailInput return s.issueLoginOutput(ctx, &user) } -// --- OAuth2 Login --- - -// OAuthLoginInput holds OAuth callback parameters. -type OAuthLoginInput struct { - Provider auth.OAuthProviderType - Code string - State string -} - -// OAuthLoginOutput holds OAuth login response. -type OAuthLoginOutput struct { - User *model.User - TokenPair *auth.TokenPair - AccountID uint - Role string - IsNewUser bool -} - -// OAuthLogin handles the OAuth2 callback: exchange code → find/create user → generate tokens. -func (s *AuthService) OAuthLogin(ctx context.Context, input *OAuthLoginInput) (*OAuthLoginOutput, error) { - // Exchange OAuth code for user info - oauthInfo, err := s.oauthService.ExchangeCode(ctx, input.Provider, input.Code) - if err != nil { - return nil, fmt.Errorf("oauth exchange failed: %w", err) - } - - // Find or create user from OAuth info - user, err := s.oauthService.FindOrCreateUser(oauthInfo) - if err != nil { - return nil, fmt.Errorf("oauth user creation failed: %w", err) - } - - // Determine if this is a new user - isNewUser := user.SignInCount == 0 - - // Get account and generate tokens - accountID, role, err := s.getUserDefaultAccount(user) - if err != nil { - // New OAuth users may not have an account yet - // Create a default personal account for them - accountID, role, err = s.createDefaultAccount(user) - if err != nil { - return nil, fmt.Errorf("failed to create default account: %w", err) - } - } - - tokenPair, err := s.jwtService.GenerateTokenPair(user, accountID, role) - if err != nil { - return nil, fmt.Errorf("failed to generate tokens: %w", err) - } - - if err := s.refreshStore.Store(ctx, user.ID, tokenPair.RefreshToken); err != nil { - return nil, fmt.Errorf("failed to store refresh token: %w", err) - } - - user.SignInCount++ - now := time.Now() - user.LastSignInAt = user.CurrentSignInAt - user.CurrentSignInAt = &now - s.db.Save(user) - - return &OAuthLoginOutput{ - User: user, - TokenPair: tokenPair, - AccountID: accountID, - Role: role, - IsNewUser: isNewUser, - }, nil -} - // --- Internal Helpers --- // AccountUser represents the join between User and Account (many-to-many). @@ -672,27 +599,3 @@ func digestAuthToken(token string) string { sum := sha256.Sum256([]byte(token)) return hex.EncodeToString(sum[:]) } - -// createDefaultAccount creates a personal account for a new user. -func (s *AuthService) createDefaultAccount(user *model.User) (uint, string, error) { - // Create account - account := &model.Account{ - Name: user.Name + "'s Account", - Status: "active", - } - if err := s.db.Create(account).Error; err != nil { - return 0, "", fmt.Errorf("failed to create account: %w", err) - } - - // Create account-user join with administrator role - accountUser := &AccountUser{ - UserID: user.ID, - AccountID: account.ID, - Role: "administrator", - } - if err := s.db.Create(accountUser).Error; err != nil { - return 0, "", fmt.Errorf("failed to create account_user: %w", err) - } - - return account.ID, "administrator", nil -} diff --git a/backend/internal/service/auth_service_test.go b/backend/internal/service/auth_service_test.go index e36d8a2d..4ebe7991 100644 --- a/backend/internal/service/auth_service_test.go +++ b/backend/internal/service/auth_service_test.go @@ -29,7 +29,7 @@ func setupAuthServiceTest(t *testing.T) (*AuthService, *gorm.DB, *model.User) { require.NoError(t, db.Create(user).Error) require.NoError(t, db.Create(&model.AccountUser{AccountID: account.ID, UserID: user.ID, Role: "administrator"}).Error) jwtCfg := &config.JWTConfig{Secret: "auth-service-secret", ExpiryHours: 1, RefreshExpiryHours: 24} - return NewAuthService(db, auth.NewJWTService(jwtCfg), auth.NewRefreshTokenStore(nil, jwtCfg), nil, nil), db, user + return NewAuthService(db, auth.NewJWTService(jwtCfg), auth.NewRefreshTokenStore(nil, jwtCfg), nil), db, user } func TestAuthService_ResetPasswordStoresDigestToken(t *testing.T) { diff --git a/backend/tests/e2e/e2e_test.go b/backend/tests/e2e/e2e_test.go index ad8076de..44fc8cf5 100644 --- a/backend/tests/e2e/e2e_test.go +++ b/backend/tests/e2e/e2e_test.go @@ -104,10 +104,9 @@ func (s *E2ETestSuite) SetupSuite() { _ = redisClient // keep miniredis alive for test duration refreshTokenStore := auth.NewRefreshTokenStore(redisClient, &cfg.JWT) - oauthService := auth.NewOAuthService(db, cfg) mfaService := auth.NewMFAService(db) - authService := service.NewAuthService(db, jwtService, refreshTokenStore, oauthService, mfaService) - authHandler := handler.NewAuthHandler(authService, oauthService) + authService := service.NewAuthService(db, jwtService, refreshTokenStore, mfaService) + authHandler := handler.NewAuthHandler(authService) // Register auth routes using helper function handler.RegisterAuthRoutes(r.Group("/api/v1"), authHandler)