From 969afee28467a8507bc79016949544d07ddc34e3 Mon Sep 17 00:00:00 2001 From: Yvan <8574526@qq,com> Date: Mon, 24 Aug 2026 21:50:33 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BC=98=E5=8C=96=E7=BB=93=E6=9E=84?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- cmd/wire_gen.go | 123 ++++++++----- internal/biz/system/authentication.go | 16 +- internal/biz/system/authority.go | 5 + internal/biz/system/security.go | 10 + internal/data/config_store.go | 24 ++- internal/data/data.go | 93 ++++++---- internal/data/data_scope_test.go | 2 +- internal/data/database.go | 28 ++- internal/data/integration_config_test.go | 4 +- internal/data/migrations_test.go | 2 +- internal/data/runtime_clients.go | 172 +++++++++++------- internal/data/system/seed.go | 34 +++- internal/server/handler/provider.go | 6 +- internal/server/handler/set.go | 4 - internal/server/middleware/access_log.go | 39 ++-- internal/server/middleware/audit.go | 49 +++-- internal/server/middleware/auth.go | 20 +- .../server/middleware/payment_audit_test.go | 8 +- internal/worker/task_executor.go | 81 ++++++++- internal/worker/task_scheduler.go | 21 +-- 20 files changed, 497 insertions(+), 244 deletions(-) diff --git a/cmd/wire_gen.go b/cmd/wire_gen.go index 49948d3..c2dc84a 100644 --- a/cmd/wire_gen.go +++ b/cmd/wire_gen.go @@ -59,7 +59,7 @@ func wireApp(configServer *config.Server, store *config.Store, logger *slog.Logg authorityAccessRepo := system.NewAuthorityAccessRepo(dataData) apiRepo := system.NewAPIRepo(dataData) accessControlUsecase := system2.NewAccessControlUsecase(authorityAccessRepo, apiRepo) - v := system3.NewAccessControlService(accessControlUsecase) + accessControlService := system3.NewAccessControlService(accessControlUsecase) userRepo := system.NewUserRepo(dataData) userUsecase := system2.NewUserUsecase(userRepo) securityRepo := system.NewSecurityRepo(dataData) @@ -72,44 +72,44 @@ func wireApp(configServer *config.Server, store *config.Store, logger *slog.Logg auditRecordRepo := system.NewAuditRecorderRepo(dataData) authenticationUsecase := system2.NewAuthenticationUsecase(userUsecase, securityUsecase, tokenIssuer, auditRecordRepo) authService := system3.NewAuthService(authenticationUsecase) - v2 := system3.NewSecurityService(securityUsecase) + securityService := system3.NewSecurityService(securityUsecase) auditRecorderUsecase := system2.NewAuditRecorderUsecase(auditRecordRepo) - v3 := system3.NewAuditRecorder(auditRecorderUsecase) + auditRecorder := system3.NewAuditRecorder(auditRecorderUsecase) authorityUsecase := system2.NewAuthorityUsecase(authorityAccessRepo) - v4 := system3.NewAuthorityService(authorityUsecase) - authority := handler.NewAuthority(v4) + authorityService := system3.NewAuthorityService(authorityUsecase) + authority := handler.NewAuthority(authorityService) menuRepo := system.NewMenuRepo(dataData) menuUsecase := system2.NewMenuUsecase(menuRepo) - v5 := system3.NewMenuService(menuUsecase) - menu := handler.NewMenu(v5) + menuService := system3.NewMenuService(menuUsecase) + menu := handler.NewMenu(menuService) apiUsecase := system2.NewAPIUsecase(apiRepo) - v6 := system3.NewAPIService(apiUsecase, runtimeSettings) - api := handler.NewAPI(v6) + apiService := system3.NewAPIService(apiUsecase, runtimeSettings) + api := handler.NewAPI(apiService) permissionRepo := system.NewPermissionRepo(dataData) permissionUsecase := system2.NewPermissionUsecase(permissionRepo) - v7 := system3.NewPermissionService(permissionUsecase) - permission := handler.NewPermission(v7) + permissionService := system3.NewPermissionService(permissionUsecase) + permission := handler.NewPermission(permissionService) departmentRepo := system.NewDepartmentRepo(dataData) departmentUsecase := system2.NewDepartmentUsecase(departmentRepo) - v8 := system3.NewDepartmentService(departmentUsecase) + departmentService := system3.NewDepartmentService(departmentUsecase) positionRepo := system.NewPositionRepo(dataData) positionUsecase := system2.NewPositionUsecase(positionRepo) - v9 := system3.NewPositionService(positionUsecase) - organization := handler.NewOrganization(v8, v9) + positionService := system3.NewPositionService(positionUsecase) + organization := handler.NewOrganization(departmentService, positionService) announcementRepo := system.NewAnnouncementRepo(dataData) announcementUsecase := system2.NewAnnouncementUsecase(announcementRepo) - v10 := system3.NewAnnouncementService(announcementUsecase) - announcement := handler.NewAnnouncement(v10) + announcementService := system3.NewAnnouncementService(announcementUsecase) + announcement := handler.NewAnnouncement(announcementService) emailRepo := email.NewEmailRepo(store) emailUsecase := system2.NewEmailUsecase(emailRepo) - v11 := system3.NewEmailService(emailUsecase) - handlerEmail := handler.NewEmail(v11) + emailService := system3.NewEmailService(emailUsecase) + handlerEmail := handler.NewEmail(emailService) paymentConfigReader := integration.NewPaymentConfigReader(dataData) paymentRepo := payment.NewPaymentRepo(dataData, paymentConfigReader) paymentOrderRepo := payment.NewPaymentOrderRepo(dataData) paymentUsecase := payment2.NewPaymentUsecase(paymentRepo, paymentOrderRepo, logger) - v12 := payment3.NewPaymentService(paymentUsecase) - handlerPayment := handler.NewPayment(v12) + paymentService := payment3.NewPaymentService(paymentUsecase) + handlerPayment := handler.NewPayment(paymentService) taskRepo := task.NewTaskRepo(dataData) registry := app.TaskRegistry(catalog) taskUsecase := task2.NewTaskUsecaseWithRegistry(taskRepo, registry) @@ -117,55 +117,78 @@ func wireApp(configServer *config.Server, store *config.Store, logger *slog.Logg taskScheduler := worker.NewTaskScheduler(taskUsecase, authorityUsecase, taskExecutor, logger) taskRuntime := worker.NewTaskRuntime(taskScheduler) taskApplicationUsecase := task2.NewTaskApplicationUsecase(taskUsecase, taskRuntime) - v13 := task3.NewTaskService(taskApplicationUsecase) - handlerTask := handler.NewTask(v13) + taskService := task3.NewTaskService(taskApplicationUsecase) + handlerTask := handler.NewTask(taskService) mediaRepo := system.NewMediaRepo(dataData) mediaUsecase := system2.NewMediaUsecase(mediaRepo, reloadable, runtimeSettings) - v14 := system3.NewMediaService(mediaUsecase, runtimeSettings) - media := handler.NewMedia(v14) + mediaService := system3.NewMediaService(mediaUsecase, runtimeSettings) + media := handler.NewMedia(mediaService) auditQueryRepo := system.NewAuditRepo(dataData) auditUsecase := system2.NewAuditUsecase(auditQueryRepo) - v15 := system3.NewAuditService(auditUsecase) + auditService := system3.NewAuditService(auditUsecase) logFileRepo := system.NewLogFileRepo(dataData) logViewerUsecase := system2.NewLogViewerUsecase(logFileRepo) - v16 := system3.NewLogViewerService(logViewerUsecase) - audit := handler.NewAudit(v15, v3, v16, logger) + logViewerService := system3.NewLogViewerService(logViewerUsecase) + audit := handler.NewAudit(auditService, auditRecorder, logViewerService, logger) exportRepo := system.NewExportRepo(dataData) exportUsecase := system2.NewExportUsecase(exportRepo) - v17 := system3.NewExportService(exportUsecase, systemCache) - export := handler.NewExport(v17) + exportService := system3.NewExportService(exportUsecase, systemCache) + export := handler.NewExport(exportService) versionRepo := system.NewVersionRepo(dataData) versionUsecase := system2.NewVersionUsecase(versionRepo) - v18 := system3.NewVersionService(versionUsecase) - version := handler.NewVersion(v18) + versionService := system3.NewVersionService(versionUsecase) + version := handler.NewVersion(versionService) dictionaryRepo := system.NewDictionaryRepo(dataData) dictionaryUsecase := system2.NewDictionaryUsecase(dictionaryRepo) - v19 := system3.NewDictionaryService(dictionaryUsecase) - dictionary := handler.NewDictionary(v19) + dictionaryService := system3.NewDictionaryService(dictionaryUsecase) + dictionary := handler.NewDictionary(dictionaryService) parameterRepo := system.NewParameterRepo(dataData) parameterUsecase := system2.NewParameterUsecase(parameterRepo) - v20 := system3.NewParameterService(parameterUsecase) - parameter := handler.NewParameter(v20) - v21 := system3.NewTokenService(tokenUsecase, tokenIssuer) - apiToken := handler.NewAPIToken(v21) + parameterService := system3.NewParameterService(parameterUsecase) + parameter := handler.NewParameter(parameterService) + tokenService := system3.NewTokenService(tokenUsecase, tokenIssuer) + apiToken := handler.NewAPIToken(tokenService) initializationRepo := initialize.NewRepo(dataData, catalog) taskReloader := worker.NewTaskReloader(taskScheduler) systemConfigUsecase := system2.NewSystemConfigUsecase(initializationRepo, taskReloader) - v22 := system3.NewSystemConfigService(systemConfigUsecase, runtimeSettings) - systemConfig := handler.NewSystemConfig(v22, v2) - public := handler.NewPublic(authService, v22, v2) - v23 := system3.NewUserService(userUsecase, v2) - user := handler.NewUser(v23, authService) - navigation := handler.NewNavigation(v23) - session := handler.NewSession(v21) + systemConfigService := system3.NewSystemConfigService(systemConfigUsecase, runtimeSettings) + systemConfig := handler.NewSystemConfig(systemConfigService, securityService) + public := handler.NewPublic(authService, systemConfigService, securityService) + userService := system3.NewUserService(userUsecase, securityService) + user := handler.NewUser(userService, authService) + navigation := handler.NewNavigation(userService) + session := handler.NewSession(tokenService) integrationConfigRepo := integration.NewIntegrationConfigRepo(dataData) runtimeconfigStore := data.NewIntegrationRuntime(dataData) connectivityTester := integration2.NewConnectivityTester(runtimeconfigStore) integrationConfigUsecase := integration3.NewIntegrationConfigUsecase(integrationConfigRepo, connectivityTester) - v24 := integration4.NewIntegrationConfigService(integrationConfigUsecase) - integrationConfig := handler.NewIntegrationConfig(v24) - v25 := handler.NewSet(authority, menu, api, permission, organization, announcement, handlerEmail, handlerPayment, handlerTask, media, audit, export, version, dictionary, parameter, apiToken, systemConfig, public, user, navigation, session, integrationConfig) - routes := router.NewRoutes(v25) + integrationConfigService := integration4.NewIntegrationConfigService(integrationConfigUsecase) + integrationConfig := handler.NewIntegrationConfig(integrationConfigService) + set := &handler.Set{ + Authority: authority, + Menu: menu, + API: api, + Permission: permission, + Organization: organization, + Announcement: announcement, + Email: handlerEmail, + Payment: handlerPayment, + Task: handlerTask, + Media: media, + Audit: audit, + Export: export, + Version: version, + Dictionary: dictionary, + Parameter: parameter, + APIToken: apiToken, + SystemConfig: systemConfig, + Public: public, + User: user, + Navigation: navigation, + Session: session, + IntegrationConfig: integrationConfig, + } + routes := router.NewRoutes(set) maintenanceRepo := system.NewMaintenanceRepo(dataData) maintenanceUsecase := system2.NewMaintenanceUsecase(maintenanceRepo) taskMethods := worker.NewTaskMethods(taskUsecase, maintenanceUsecase, mediaUsecase, store) @@ -176,7 +199,7 @@ func wireApp(configServer *config.Server, store *config.Store, logger *slog.Logg cleanup() return nil, nil, err } - engine := server.NewGinEngineWithRuntime(store, v, authService, v2, v3, logger, string2, runtime, websocketServer) + engine := server.NewGinEngineWithRuntime(store, accessControlService, authService, securityService, auditRecorder, logger, string2, runtime, websocketServer) httpServer := server.NewGinServer(configServer, engine) mqReloadable, cleanup3, err := mq.New(runtimeconfigStore, logger) if err != nil { @@ -185,7 +208,7 @@ func wireApp(configServer *config.Server, store *config.Store, logger *slog.Logg return nil, nil, err } resourceRegistry := installGlobalResources(logger, dataData, reloadable, mqReloadable, websocketServer, taskScheduler) - kratosApp := newApp(logger, httpServer, taskScheduler, v3, reloadableLogger, mqReloadable, resourceRegistry) + kratosApp := newApp(logger, httpServer, taskScheduler, auditRecorder, reloadableLogger, mqReloadable, resourceRegistry) return kratosApp, func() { cleanup3() cleanup2() diff --git a/internal/biz/system/authentication.go b/internal/biz/system/authentication.go index e0dd2dc..e64c3c8 100644 --- a/internal/biz/system/authentication.go +++ b/internal/biz/system/authentication.go @@ -56,6 +56,16 @@ func NewAuthenticationUsecase(users *UserUsecase, security *SecurityUsecase, iss return &AuthenticationUsecase{users: users, security: security, issuer: issuer, audit: audit} } +// passwordExpired reports whether the security policy has aged out this user's +// password. Both the credential login and the per-request token check consult +// it, so the rule lives in one place. +func passwordExpired(config *SecurityConfig, user *User) bool { + if config == nil || user == nil || !config.PwdExpireEnable || config.PwdExpireDays <= 0 || user.PasswordUpdatedAt == nil { + return false + } + return time.Now().After(user.PasswordUpdatedAt.AddDate(0, 0, config.PwdExpireDays)) +} + func (uc *AuthenticationUsecase) recordLogin(ctx context.Context, attempt *LoginAttempt, status bool, message string, userID uint) { if uc.audit == nil || attempt == nil { return @@ -119,7 +129,7 @@ func (uc *AuthenticationUsecase) Login(ctx context.Context, attempt *LoginAttemp } uc.security.ClearLoginState(ctx, attempt.Username) - if config.PwdExpireEnable && config.PwdExpireDays > 0 && user.PasswordUpdatedAt != nil && time.Now().After((*user.PasswordUpdatedAt).AddDate(0, 0, config.PwdExpireDays)) { + if passwordExpired(config, user) { user.MustChangePassword = true } issued, err := uc.issuer.IssueToken(user, user.AuthorityID, user.MustChangePassword, 0) @@ -185,7 +195,7 @@ func (uc *AuthenticationUsecase) currentTokenUser(ctx context.Context, claims *A if err != nil || config == nil { return nil, ErrTokenDisabled } - if config.PwdExpireEnable && config.PwdExpireDays > 0 && user.PasswordUpdatedAt != nil && time.Now().After(user.PasswordUpdatedAt.AddDate(0, 0, config.PwdExpireDays)) { + if passwordExpired(config, user) { user.MustChangePassword = true } return user, nil @@ -219,7 +229,7 @@ func (uc *AuthenticationUsecase) AuthenticateToken(ctx context.Context, token st // this makes a revoked-but-expired token report the revocation reason (and // not the generic expiry message), which the frontend uses to decide whether // to clear a session. - disabled, err := uc.security.tokens.IsTokenDisabled(ctx, token) + disabled, err := uc.security.TokenDisabled(ctx, token) if err != nil || disabled { return nil, ErrTokenDisabled } diff --git a/internal/biz/system/authority.go b/internal/biz/system/authority.go index 81c925a..bcc228f 100644 --- a/internal/biz/system/authority.go +++ b/internal/biz/system/authority.go @@ -5,6 +5,11 @@ import ( "time" ) +// SuperAdminAuthorityID is the seeded super-administrator role created by the +// bootstrap data in internal/data/system/seed.go. It is the fallback authority +// for new users and the recipient group for system-wide alerts. +const SuperAdminAuthorityID uint = 888 + type AuthorityAccessRepo interface { CreateAuthority(context.Context, *Authority) error CopyAuthority(context.Context, uint, *Authority) error diff --git a/internal/biz/system/security.go b/internal/biz/system/security.go index b04db8d..d63bfd9 100644 --- a/internal/biz/system/security.go +++ b/internal/biz/system/security.go @@ -270,6 +270,16 @@ func (uc *SecurityUsecase) RotateActiveToken(ctx context.Context, username, oldT func (uc *SecurityUsecase) UseMultipoint() bool { return uc.settings.UseMultipoint() } +// TokenDisabled reports whether the token sits on the revocation blacklist. The +// blacklist belongs to the token usecase; exposing it here keeps authentication +// from reaching through this usecase into another one's dependencies. +func (uc *SecurityUsecase) TokenDisabled(ctx context.Context, token string) (bool, error) { + if uc == nil || uc.tokens == nil { + return false, nil + } + return uc.tokens.IsTokenDisabled(ctx, token) +} + func (uc *SecurityUsecase) CaptchaRuntimeSettings() CaptchaSettings { return uc.settings.CaptchaSettings() } diff --git a/internal/data/config_store.go b/internal/data/config_store.go index 17ff6db..5d32405 100644 --- a/internal/data/config_store.go +++ b/internal/data/config_store.go @@ -303,7 +303,7 @@ func (d *Data) reloadConfig(ctx context.Context) error { if configPath == "" { return fmt.Errorf("configuration path is not set") } - next, err := readBootstrap(configPath) + next, err := config.Load(configPath) if err != nil { return err } @@ -406,7 +406,7 @@ func (d *Data) reloadConfig(ctx context.Context) error { closeDatabaseList(candidateDBList) } }() - integrationConfigs, err := readIntegrationRuntime(candidateDB) + integrationConfigs, err := readIntegrationRuntime(candidateDB.WithContext(ctx)) if err != nil { return fmt.Errorf("reload integration runtime: %w", err) } @@ -417,8 +417,20 @@ func (d *Data) reloadConfig(ctx context.Context) error { registerDataScopeCallbacks(item, d.enqueueDataScopeAudit) } d.replaceDatabaseList(candidateDBList) - d.redis.replace(candidateRedis) - d.replaceRedisList(candidateRedisList) + // openRedis and openRedisList return nil when a ping fails, so replacing + // unconditionally would let a transient Redis outage during an unrelated + // configuration reload retire the still-healthy clients and silently downgrade + // the process to the in-memory cache until the next reload. + if candidateRedis != nil || !useRedis || !redisConnectionConfigured(next.Data.Redis) { + d.redis.replace(candidateRedis) + } else { + d.logger().Warn("keeping the previous redis client because the reloaded configuration failed to connect", "mod", "redis") + } + if candidateRedisList != nil || !useRedisList || len(next.Data.RedisList) == 0 { + d.replaceRedisList(candidateRedisList) + } else { + d.logger().Warn("keeping the previous redis list because the reloaded configuration failed to connect", "mod", "redis") + } if mongoErr == nil { d.mongo.replace(candidateMongo) mongoAccepted = true @@ -437,7 +449,3 @@ func (d *Data) reloadConfig(ctx context.Context) error { d.notifyResources() return nil } - -func readBootstrap(configPath string) (*config.Config, error) { - return config.Load(configPath) -} diff --git a/internal/data/data.go b/internal/data/data.go index ad8cbb8..0ad6db4 100644 --- a/internal/data/data.go +++ b/internal/data/data.go @@ -58,10 +58,14 @@ type Data struct { storage *storage.Reloadable dbListMu sync.RWMutex dbList map[string]*gorm.DB - appLogger *slog.Logger - auditLog *dataScopeAuditWriter - catalog module.Catalog - resourceHook func() + // Handles replaced by a hot reload of the named lists. They follow the same + // grace period as the primary pool so in-flight queries are not cut off. + retiredDBList retiredSet[*gorm.DB] + retiredRedisList retiredSet[redis.UniversalClient] + appLogger *slog.Logger + auditLog *dataScopeAuditWriter + catalog module.Catalog + resourceHook func() } // DB exposes the active primary database to narrowly scoped data submodules. @@ -100,7 +104,16 @@ func (d *Data) IntegrationRuntime() *runtimeconfig.Store { // Database resolves the primary or a named database for repositories such as // the system export module. func (d *Data) Database(name string) (*gorm.DB, error) { - return d.database(name) + if name == "" { + return d.gormDB.DB(), nil + } + d.dbListMu.RLock() + db := d.dbList[name] + d.dbListMu.RUnlock() + if db == nil { + return nil, fmt.Errorf("database %q not found", name) + } + return db, nil } func (d *Data) RedisClient() redis.UniversalClient { @@ -190,13 +203,13 @@ func (d *Data) logger() *slog.Logger { return slog.Default() } -func openDatabaseList(configs []*config.Database, appLogger ...*slog.Logger) (map[string]*gorm.DB, error) { +func openDatabaseList(configs []*config.Database, appLogger *slog.Logger) (map[string]*gorm.DB, error) { items := make(map[string]*gorm.DB) for _, config := range configs { if config == nil || config.Disable || config.AliasName == "" { continue } - db, err := openDatabase(config, false, "", appLogger...) + db, err := openDatabase(config, false, "", appLogger) if err != nil { for _, opened := range items { if sqlDB, dbErr := opened.DB(); dbErr == nil { @@ -212,31 +225,26 @@ func openDatabaseList(configs []*config.Database, appLogger ...*slog.Logger) (ma func closeDatabaseList(items map[string]*gorm.DB) { for _, db := range items { - if sqlDB, err := db.DB(); err == nil { - _ = sqlDB.Close() + if db != nil { + closeGormDB(db) } } } +// replaceDatabaseList swaps the named database handles. Superseded handles are +// retired instead of closed on the spot: a request that already resolved a +// handle would otherwise fail with "sql: database is closed" mid-flight. func (d *Data) replaceDatabaseList(items map[string]*gorm.DB) { d.dbListMu.Lock() old := d.dbList d.dbList = items d.dbListMu.Unlock() - closeDatabaseList(old) -} - -func (d *Data) database(name string) (*gorm.DB, error) { - if name == "" { - return d.gormDB.DB(), nil + for name, db := range old { + if db == nil || db == items[name] { + continue + } + d.retiredDBList.retire(db, closeGormDB) } - d.dbListMu.RLock() - db := d.dbList[name] - d.dbListMu.RUnlock() - if db == nil { - return nil, fmt.Errorf("database %q not found", name) - } - return db, nil } func NewData(runtime *config.Store, appLogger *slog.Logger, storageManager *storage.Reloadable, catalog module.Catalog) (*Data, func(), error) { @@ -264,10 +272,16 @@ func NewData(runtime *config.Store, appLogger *slog.Logger, storageManager *stor d.gormDB.close() } closeDatabaseList(d.dbList) + for _, db := range d.retiredDBList.drain() { + closeGormDB(db) + } if d.redis != nil { d.redis.close() } closeRedisList(d.redisList) + for _, client := range d.retiredRedisList.drain() { + closeRedisClient(client) + } if d.mongo != nil { d.mongo.close() } @@ -363,10 +377,20 @@ func NewData(runtime *config.Store, appLogger *slog.Logger, storageManager *stor return d, cleanup, nil } -func openRedis(config *config.Redis, enabled bool, appLogger ...*slog.Logger) redis.UniversalClient { - if !enabled || config == nil || (config.Addr == "" && len(config.ClusterAddrs) == 0) { +// redisConnectionConfigured reports whether a Redis block names an endpoint. +// The reload path shares it with openRedis so "no endpoint configured" is never +// confused with "configured endpoint failed to answer". +func redisConnectionConfigured(config *config.Redis) bool { + return config != nil && (config.Addr != "" || len(config.ClusterAddrs) > 0) +} + +func openRedis(config *config.Redis, enabled bool, appLogger *slog.Logger) redis.UniversalClient { + if !enabled || !redisConnectionConfigured(config) { return nil } + if appLogger == nil { + appLogger = slog.Default() + } var candidate redis.UniversalClient if config.UseCluster { addresses := config.ClusterAddrs @@ -387,18 +411,14 @@ func openRedis(config *config.Redis, enabled bool, appLogger ...*slog.Logger) re pingCtx, cancel := context.WithTimeout(context.Background(), 800*time.Millisecond) defer cancel() if err := candidate.Ping(pingCtx).Err(); err != nil { - log := slog.Default() - if len(appLogger) > 0 && appLogger[0] != nil { - log = appLogger[0] - } - log.Warn("redis unavailable, using in-memory cache", "mod", "redis", "error", err) + appLogger.Warn("redis unavailable, using in-memory cache", "mod", "redis", "error", err) _ = candidate.Close() return nil } return candidate } -func openRedisList(configs []*config.Redis, enabled bool, appLogger ...*slog.Logger) map[string]redis.UniversalClient { +func openRedisList(configs []*config.Redis, enabled bool, appLogger *slog.Logger) map[string]redis.UniversalClient { if !enabled || len(configs) == 0 { return nil } @@ -407,7 +427,7 @@ func openRedisList(configs []*config.Redis, enabled bool, appLogger ...*slog.Log if item == nil || item.Name == "" { continue } - if client := openRedis(item, true, appLogger...); client != nil { + if client := openRedis(item, true, appLogger); client != nil { clients[item.Name] = client } } @@ -420,11 +440,13 @@ func openRedisList(configs []*config.Redis, enabled bool, appLogger ...*slog.Log func closeRedisList(clients map[string]redis.UniversalClient) { for _, client := range clients { if client != nil { - _ = client.Close() + closeRedisClient(client) } } } +// replaceRedisList mirrors replaceDatabaseList: superseded clients stay open for +// the retire grace period so in-flight commands are not aborted. func (d *Data) replaceRedisList(clients map[string]redis.UniversalClient) { if d == nil { return @@ -433,7 +455,12 @@ func (d *Data) replaceRedisList(clients map[string]redis.UniversalClient) { old := d.redisList d.redisList = clients d.redisListMu.Unlock() - closeRedisList(old) + for name, client := range old { + if client == nil || client == clients[name] { + continue + } + d.retiredRedisList.retire(client, closeRedisClient) + } } func (d *Data) activateDatabase(db *gorm.DB, config *config.Database) { diff --git a/internal/data/data_scope_test.go b/internal/data/data_scope_test.go index b9427b3..fddd96c 100644 --- a/internal/data/data_scope_test.go +++ b/internal/data/data_scope_test.go @@ -20,7 +20,7 @@ func (dataScopeRecord) TableName() string { return "business_scope_records" } func newDataScopeTestDB(t *testing.T) *gorm.DB { t.Helper() - db, err := openWithDriver("sqlite", "file:"+t.Name()+"?mode=memory&cache=shared") + db, err := openWithDriver("sqlite", "file:"+t.Name()+"?mode=memory&cache=shared", nil) if err != nil { t.Fatal(err) } diff --git a/internal/data/database.go b/internal/data/database.go index db3ec04..55aebac 100644 --- a/internal/data/database.go +++ b/internal/data/database.go @@ -118,7 +118,7 @@ func databaseDSN(c *config.Database, name string) (string, error) { return "", fmt.Errorf("unsupported database driver %q", c.Driver) } -func gormConfig(config *config.Database, appLogger ...*slog.Logger) *gorm.Config { +func gormConfig(config *config.Database, appLogger *slog.Logger) *gorm.Config { level := logger.Info switch strings.ToLower(config.LogMode) { case "silent": @@ -128,19 +128,15 @@ func gormConfig(config *config.Database, appLogger ...*slog.Logger) *gorm.Config case "warn": level = logger.Warn } - var log *slog.Logger - if len(appLogger) > 0 { - log = appLogger[0] - } - return &gorm.Config{Logger: gormkit.NewLogger(log, level), NamingStrategy: schema.NamingStrategy{TablePrefix: config.Prefix, SingularTable: config.Singular}} + return &gorm.Config{Logger: gormkit.NewLogger(appLogger, level), NamingStrategy: schema.NamingStrategy{TablePrefix: config.Prefix, SingularTable: config.Singular}} } -func openWithDriver(driver, dsn string, appLogger ...*slog.Logger) (*gorm.DB, error) { - return openWithDriverConfig(driver, dsn, &config.Database{Driver: driver}, appLogger...) +func openWithDriver(driver, dsn string, appLogger *slog.Logger) (*gorm.DB, error) { + return openWithDriverConfig(driver, dsn, &config.Database{Driver: driver}, appLogger) } -func openWithDriverConfig(driver, dsn string, config *config.Database, appLogger ...*slog.Logger) (*gorm.DB, error) { - gormConfig := gormConfig(config, appLogger...) +func openWithDriverConfig(driver, dsn string, config *config.Database, appLogger *slog.Logger) (*gorm.DB, error) { + gormConfig := gormConfig(config, appLogger) var db *gorm.DB var err error switch normalizedDriver(driver) { @@ -179,7 +175,7 @@ func openWithDriverConfig(driver, dsn string, config *config.Database, appLogger return db, nil } -func openDatabase(c *config.Database, create bool, template string, appLogger ...*slog.Logger) (*gorm.DB, error) { +func openDatabase(c *config.Database, create bool, template string, appLogger *slog.Logger) (*gorm.DB, error) { driver := normalizedDriver(c.Driver) if driver == "" { return nil, fmt.Errorf("unsupported database driver %q", c.Driver) @@ -192,7 +188,7 @@ func openDatabase(c *config.Database, create bool, template string, appLogger .. if err = os.MkdirAll(filepath.Dir(dsn), 0o755); err != nil { return nil, err } - return openWithDriverConfig(driver, dsn, c, appLogger...) + return openWithDriverConfig(driver, dsn, c, appLogger) } if create && driver != "oracle" { if !databaseNamePattern.MatchString(c.Name) { @@ -209,7 +205,7 @@ func openDatabase(c *config.Database, create bool, template string, appLogger .. if err != nil { return nil, err } - adminDB, err := openWithDriverConfig(driver, dsn, c, appLogger...) + adminDB, err := openWithDriverConfig(driver, dsn, c, appLogger) if err != nil { return nil, fmt.Errorf("connect database server: %w", err) } @@ -245,9 +241,9 @@ func openDatabase(c *config.Database, create bool, template string, appLogger .. if err != nil { return nil, err } - return openWithDriverConfig(driver, dsn, c, appLogger...) + return openWithDriverConfig(driver, dsn, c, appLogger) } -func openFallbackDatabase(appLogger ...*slog.Logger) (*gorm.DB, error) { - return openWithDriver("sqlite", "file:kra-bootstrap?mode=memory&cache=shared", appLogger...) +func openFallbackDatabase(appLogger *slog.Logger) (*gorm.DB, error) { + return openWithDriver("sqlite", "file:kra-bootstrap?mode=memory&cache=shared", appLogger) } diff --git a/internal/data/integration_config_test.go b/internal/data/integration_config_test.go index 4b9c628..1e740e6 100644 --- a/internal/data/integration_config_test.go +++ b/internal/data/integration_config_test.go @@ -18,7 +18,7 @@ import ( func openIntegrationConfigTestDB(t *testing.T) *gorm.DB { t.Helper() - db, err := openWithDriver("sqlite", "file:"+t.Name()+"?mode=memory&cache=shared") + db, err := openWithDriver("sqlite", "file:"+t.Name()+"?mode=memory&cache=shared", nil) if err != nil { t.Fatal(err) } @@ -34,7 +34,7 @@ func openIntegrationConfigTestDB(t *testing.T) *gorm.DB { } func TestMigrateAllCreatesIntegrationConfigTable(t *testing.T) { - db, err := openWithDriver("sqlite", "file:"+t.Name()+"?mode=memory&cache=shared") + db, err := openWithDriver("sqlite", "file:"+t.Name()+"?mode=memory&cache=shared", nil) if err != nil { t.Fatal(err) } diff --git a/internal/data/migrations_test.go b/internal/data/migrations_test.go index 6268c4a..3661dd0 100644 --- a/internal/data/migrations_test.go +++ b/internal/data/migrations_test.go @@ -15,7 +15,7 @@ func testCatalog() platformmodule.Catalog { } func TestMigrateAllRunsModuleSchemasWithoutBootstrapSeed(t *testing.T) { - db, err := openWithDriver("sqlite", "file:"+t.Name()+"?mode=memory&cache=shared") + db, err := openWithDriver("sqlite", "file:"+t.Name()+"?mode=memory&cache=shared", nil) if err != nil { t.Fatal(err) } diff --git a/internal/data/runtime_clients.go b/internal/data/runtime_clients.go index fd5babd..0fed2b5 100644 --- a/internal/data/runtime_clients.go +++ b/internal/data/runtime_clients.go @@ -4,60 +4,63 @@ import ( "context" "sync" "sync/atomic" + "time" "github.com/redis/go-redis/v9" "go.mongodb.org/mongo-driver/mongo" "gorm.io/gorm" ) -// reloadableDB makes the pointer swap atomic. Replaced pools are retained -// until application shutdown so in-flight GORM operations remain valid. -type reloadableDB struct { - current atomic.Pointer[gorm.DB] - mu sync.Mutex - retired []*gorm.DB +// retireGrace bounds how long a client replaced by a hot reload stays open. +// In-flight operations still hold the old handle, so it cannot be closed at +// swap time; keeping it until shutdown would leak one full pool per reload. +const retireGrace = 5 * time.Minute + +// retiredSet holds clients replaced by a hot reload and closes each of them +// once the grace period expires. Whatever is still pending at shutdown is +// drained by the owner's close. +type retiredSet[T comparable] struct { + mu sync.Mutex + items []T } -type reloadableMongo struct { - mu sync.RWMutex - current *mongo.Client - retired []*mongo.Client +func (s *retiredSet[T]) retire(item T, closeItem func(T)) { + s.mu.Lock() + s.items = append(s.items, item) + s.mu.Unlock() + time.AfterFunc(retireGrace, func() { + if s.take(item) { + closeItem(item) + } + }) } -func newReloadableMongo(client *mongo.Client) *reloadableMongo { - return &reloadableMongo{current: client} -} -func (r *reloadableMongo) replace(client *mongo.Client) { - r.mu.Lock() - old := r.current - r.current = client - if old != nil && old != client { - r.retired = append(r.retired, old) - } - r.mu.Unlock() -} - -func (r *reloadableMongo) load() *mongo.Client { - if r == nil { - return nil - } - r.mu.RLock() - client := r.current - r.mu.RUnlock() - return client -} - -func (r *reloadableMongo) close() { - r.mu.Lock() - all := append([]*mongo.Client{r.current}, r.retired...) - r.current = nil - r.retired = nil - r.mu.Unlock() - for _, client := range all { - if client != nil { - _ = client.Disconnect(context.Background()) +// take removes item and reports whether this caller now owns closing it. +func (s *retiredSet[T]) take(item T) bool { + s.mu.Lock() + defer s.mu.Unlock() + for index, existing := range s.items { + if existing == item { + s.items = append(s.items[:index], s.items[index+1:]...) + return true } } + return false +} + +func (s *retiredSet[T]) drain() []T { + s.mu.Lock() + defer s.mu.Unlock() + items := s.items + s.items = nil + return items +} + +// reloadableDB makes the pointer swap atomic. Replaced pools stay open for +// retireGrace so in-flight GORM operations remain valid. +type reloadableDB struct { + current atomic.Pointer[gorm.DB] + retired retiredSet[*gorm.DB] } func newReloadableDB(db *gorm.DB, enqueue dataScopeAuditEnqueue) *reloadableDB { @@ -77,18 +80,12 @@ func (r *reloadableDB) replace(db *gorm.DB, enqueue dataScopeAuditEnqueue) { registerDataScopeCallbacks(db, enqueue) old := r.current.Swap(db) if old != nil && old != db { - r.mu.Lock() - r.retired = append(r.retired, old) - r.mu.Unlock() + r.retired.retire(old, closeGormDB) } } func (r *reloadableDB) close() { - current := r.current.Load() - r.mu.Lock() - all := append([]*gorm.DB{current}, r.retired...) - r.retired = nil - r.mu.Unlock() + all := append([]*gorm.DB{r.current.Load()}, r.retired.drain()...) seen := map[*gorm.DB]struct{}{} for _, db := range all { if db == nil { @@ -98,16 +95,65 @@ func (r *reloadableDB) close() { continue } seen[db] = struct{}{} - if sqlDB, err := db.DB(); err == nil { - _ = sqlDB.Close() + closeGormDB(db) + } +} + +func closeGormDB(db *gorm.DB) { + if sqlDB, err := db.DB(); err == nil { + _ = sqlDB.Close() + } +} + +type reloadableMongo struct { + mu sync.RWMutex + current *mongo.Client + retired retiredSet[*mongo.Client] +} + +func newReloadableMongo(client *mongo.Client) *reloadableMongo { + return &reloadableMongo{current: client} +} + +func (r *reloadableMongo) replace(client *mongo.Client) { + r.mu.Lock() + old := r.current + r.current = client + r.mu.Unlock() + if old != nil && old != client { + r.retired.retire(old, closeMongoClient) + } +} + +func (r *reloadableMongo) load() *mongo.Client { + if r == nil { + return nil + } + r.mu.RLock() + defer r.mu.RUnlock() + return r.current +} + +func (r *reloadableMongo) close() { + r.mu.Lock() + current := r.current + r.current = nil + r.mu.Unlock() + for _, client := range append([]*mongo.Client{current}, r.retired.drain()...) { + if client != nil { + closeMongoClient(client) } } } +func closeMongoClient(client *mongo.Client) { + _ = client.Disconnect(context.Background()) +} + type reloadableRedis struct { mu sync.RWMutex current redis.UniversalClient - retired []redis.UniversalClient + retired retiredSet[redis.UniversalClient] } func newReloadableRedis(client redis.UniversalClient) *reloadableRedis { @@ -124,22 +170,24 @@ func (r *reloadableRedis) replace(client redis.UniversalClient) { r.mu.Lock() old := r.current r.current = client - if old != nil && old != client { - r.retired = append(r.retired, old) - } r.mu.Unlock() + if old != nil && old != client { + r.retired.retire(old, closeRedisClient) + } } func (r *reloadableRedis) close() { r.mu.Lock() - all := append([]redis.UniversalClient{r.current}, r.retired...) + current := r.current r.current = nil - r.retired = nil r.mu.Unlock() - for _, client := range all { - if client == nil { - continue + for _, client := range append([]redis.UniversalClient{current}, r.retired.drain()...) { + if client != nil { + closeRedisClient(client) } - _ = client.Close() } } + +func closeRedisClient(client redis.UniversalClient) { + _ = client.Close() +} diff --git a/internal/data/system/seed.go b/internal/data/system/seed.go index 45db9f2..63efe11 100644 --- a/internal/data/system/seed.go +++ b/internal/data/system/seed.go @@ -208,5 +208,37 @@ func defaultMenus() []menuPO { value.KeepAlive = true return value } - return []menuPO{{Path: "dashboard", Name: "dashboard", Component: "view/dashboard/index.vue", Title: "仪表盘", Icon: "odometer", Sort: 1}, root("permission", "permission", "权限管理", "perm-kra", 2), root("org", "org", "组织管理", "share", 3), root("systemConfig", "systemConfig", "系统设置", "config-kra", 4), root("monitor", "monitor", "运维监控", "monitor-kra", 5), root("media", "media", "媒体管理", "folder-opened", 6), root("extensions", "extensions", "扩展功能", "cherry", 10), {Path: "person", Name: "person", Component: "view/person/person.vue", Title: "个人信息", Icon: "postcard", Hidden: true, Sort: 13}, child("permission", "authority", "authority", "view/superAdmin/authority/authority.vue", "角色管理", "role-kra", 1), cachedChild("permission", "menu", "menu", "view/superAdmin/menu/menu.vue", "菜单管理", "tickets", 2), cachedChild("permission", "api", "api", "view/superAdmin/api/api.vue", "api管理", "api-kra", 3), child("permission", "apiToken", "apiToken", "view/systemTools/apiToken/index.vue", "API Token", "key", 4), child("org", "user", "user", "view/superAdmin/user/user.vue", "用户管理", "user", 1), child("org", "department", "department", "view/superAdmin/department/department.vue", "部门管理", "office-building", 2), child("org", "position", "position", "view/superAdmin/position/position.vue", "岗位管理", "postcard", 3), child("systemConfig", "system", "system", "view/systemTools/system/system.vue", "配置文件", "config-file-kra", 1), child("systemConfig", "dictionary", "dictionary", "view/superAdmin/dictionary/sysDictionary.vue", "字典管理", "notebook", 2), child("systemConfig", "sysParams", "sysParams", "view/superAdmin/params/sysParams.vue", "参数管理", "set-up", 3), child("systemConfig", "security", "security", "view/system/security/index.vue", "安全配置", "security-kra", 4), child("monitor", "operation", "operation", "view/superAdmin/operation/sysOperationRecord.vue", "操作历史", "document", 1), child("monitor", "loginLog", "loginLog", "view/systemTools/loginLog/index.vue", "登录日志", "clock", 2), child("monitor", "sysError", "sysError", "view/systemTools/sysError/sysError.vue", "错误日志", "error-kra", 3), child("monitor", "sysVersion", "sysVersion", "view/systemTools/version/version.vue", "版本管理", "version-kra", 4), child("monitor", "state", "state", "view/system/state.vue", "服务器状态", "server", 5), child("monitor", "dataAccessLog", "dataAccessLog", "view/superAdmin/dataAccessLog/dataAccessLog.vue", "数据权限审计", "warning", 6), child("monitor", "timedTask", "timedTask", "view/systemTools/timedTask/index.vue", "定时任务", "timer", 7), child("monitor", "logViewer", "logViewer", "view/systemTools/logViewer/index.vue", "文件日志", "document", 8), child("media", "upload", "upload", "view/media/upload.vue", "媒体库(上传下载)", "upload", 1), child("media", "chunkUpload", "chunkUpload", "view/media/chunkUpload.vue", "大文件上传", "folder-add", 2), child("extensions", "email", "email", "modules/email/view/index.vue", "邮件发送", "message", 4), child("extensions", "anInfo", "anInfo", "modules/announcement/view/info.vue", "公告管理", "bell", 5)} + return []menuPO{ + {Path: "dashboard", Name: "dashboard", Component: "view/dashboard/index.vue", Title: "仪表盘", Icon: "odometer", Sort: 1}, + root("permission", "permission", "权限管理", "perm-kra", 2), + root("org", "org", "组织管理", "share", 3), + root("systemConfig", "systemConfig", "系统设置", "config-kra", 4), + root("monitor", "monitor", "运维监控", "monitor-kra", 5), + root("media", "media", "媒体管理", "folder-opened", 6), + root("extensions", "extensions", "扩展功能", "cherry", 10), + {Path: "person", Name: "person", Component: "view/person/person.vue", Title: "个人信息", Icon: "postcard", Hidden: true, Sort: 13}, + child("permission", "authority", "authority", "view/superAdmin/authority/authority.vue", "角色管理", "role-kra", 1), + cachedChild("permission", "menu", "menu", "view/superAdmin/menu/menu.vue", "菜单管理", "tickets", 2), + cachedChild("permission", "api", "api", "view/superAdmin/api/api.vue", "api管理", "api-kra", 3), + child("permission", "apiToken", "apiToken", "view/systemTools/apiToken/index.vue", "API Token", "key", 4), + child("org", "user", "user", "view/superAdmin/user/user.vue", "用户管理", "user", 1), + child("org", "department", "department", "view/superAdmin/department/department.vue", "部门管理", "office-building", 2), + child("org", "position", "position", "view/superAdmin/position/position.vue", "岗位管理", "postcard", 3), + child("systemConfig", "system", "system", "view/systemTools/system/system.vue", "配置文件", "config-file-kra", 1), + child("systemConfig", "dictionary", "dictionary", "view/superAdmin/dictionary/sysDictionary.vue", "字典管理", "notebook", 2), + child("systemConfig", "sysParams", "sysParams", "view/superAdmin/params/sysParams.vue", "参数管理", "set-up", 3), + child("systemConfig", "security", "security", "view/system/security/index.vue", "安全配置", "security-kra", 4), + child("monitor", "operation", "operation", "view/superAdmin/operation/sysOperationRecord.vue", "操作历史", "document", 1), + child("monitor", "loginLog", "loginLog", "view/systemTools/loginLog/index.vue", "登录日志", "clock", 2), + child("monitor", "sysError", "sysError", "view/systemTools/sysError/sysError.vue", "错误日志", "error-kra", 3), + child("monitor", "sysVersion", "sysVersion", "view/systemTools/version/version.vue", "版本管理", "version-kra", 4), + child("monitor", "state", "state", "view/system/state.vue", "服务器状态", "server", 5), + child("monitor", "dataAccessLog", "dataAccessLog", "view/superAdmin/dataAccessLog/dataAccessLog.vue", "数据权限审计", "warning", 6), + child("monitor", "timedTask", "timedTask", "view/systemTools/timedTask/index.vue", "定时任务", "timer", 7), + child("monitor", "logViewer", "logViewer", "view/systemTools/logViewer/index.vue", "文件日志", "document", 8), + child("media", "upload", "upload", "view/media/upload.vue", "媒体库(上传下载)", "upload", 1), + child("media", "chunkUpload", "chunkUpload", "view/media/chunkUpload.vue", "大文件上传", "folder-add", 2), + child("extensions", "email", "email", "modules/email/view/index.vue", "邮件发送", "message", 4), + child("extensions", "anInfo", "anInfo", "modules/announcement/view/info.vue", "公告管理", "bell", 5), + } } diff --git a/internal/server/handler/provider.go b/internal/server/handler/provider.go index 7eaa591..c00712d 100644 --- a/internal/server/handler/provider.go +++ b/internal/server/handler/provider.go @@ -2,11 +2,13 @@ package handler import "github.com/google/wire" -// ProviderSet wires the system HTTP handlers. +// ProviderSet wires the system HTTP handlers. Set itself is filled field by +// field so adding a handler only means adding its constructor here. var ProviderSet = wire.NewSet( NewAuthority, NewMenu, NewAPI, NewPermission, NewOrganization, NewAnnouncement, NewEmail, NewPayment, NewTask, NewMedia, NewAudit, NewExport, NewVersion, NewDictionary, NewParameter, NewAPIToken, - NewSystemConfig, NewPublic, NewUser, NewNavigation, NewSession, NewSet, + NewSystemConfig, NewPublic, NewUser, NewNavigation, NewSession, NewIntegrationConfig, + wire.Struct(new(Set), "*"), ) diff --git a/internal/server/handler/set.go b/internal/server/handler/set.go index cb02879..d892c7b 100644 --- a/internal/server/handler/set.go +++ b/internal/server/handler/set.go @@ -24,7 +24,3 @@ type Set struct { Session *Session IntegrationConfig *IntegrationConfig } - -func NewSet(authority *Authority, menu *Menu, api *API, permission *Permission, organization *Organization, announcement *Announcement, email *Email, payment *Payment, task *Task, media *Media, audit *Audit, export *Export, version *Version, dictionary *Dictionary, parameter *Parameter, apiToken *APIToken, systemConfig *SystemConfig, public *Public, user *User, navigation *Navigation, session *Session, integrationConfig *IntegrationConfig) *Set { - return &Set{Authority: authority, Menu: menu, API: api, Permission: permission, Organization: organization, Announcement: announcement, Email: email, Payment: payment, Task: task, Media: media, Audit: audit, Export: export, Version: version, Dictionary: dictionary, Parameter: parameter, APIToken: apiToken, SystemConfig: systemConfig, Public: public, User: user, Navigation: navigation, Session: session, IntegrationConfig: integrationConfig} -} diff --git a/internal/server/middleware/access_log.go b/internal/server/middleware/access_log.go index 2532fec..05b4432 100644 --- a/internal/server/middleware/access_log.go +++ b/internal/server/middleware/access_log.go @@ -142,17 +142,28 @@ func AccessLog(runtime *config.Store, logger *slog.Logger, version string) gin.H "req_query", redactQuery(c.Request.URL.RawQuery), } } - if !paymentCallback && admin != nil && admin.Zap != nil && admin.Zap.AccessReqHeaders { - attributes = append(attributes, "req_headers", redactHeaders(c.Request.Header)) - } - if !paymentCallback && admin != nil && admin.Zap != nil && admin.Zap.AccessReqBody { - attributes = append(attributes, "req_body", requestText) - } - if !paymentCallback && admin != nil && admin.Zap != nil && admin.Zap.AccessRespData { - attributes = append(attributes, "resp_data", responseText) - } - if !paymentCallback && privateErrors != "" { - attributes = append(attributes, "error_msg", privateErrors) + // A payment callback is logged through its own attribute set: neither the + // configurable request/response capture nor the private error text may leak + // provider payloads into the generic access log. + if !paymentCallback { + var zap *config.Zap + if admin != nil { + zap = admin.Zap + } + if zap != nil { + if zap.AccessReqHeaders { + attributes = append(attributes, "req_headers", redactHeaders(c.Request.Header)) + } + if zap.AccessReqBody { + attributes = append(attributes, "req_body", requestText) + } + if zap.AccessRespData { + attributes = append(attributes, "resp_data", responseText) + } + } + if privateErrors != "" { + attributes = append(attributes, "error_msg", privateErrors) + } } logger.InfoContext(c.Request.Context(), "请求完成", attributes...) } @@ -173,12 +184,6 @@ func paymentCallbackProvider(path string) string { return "unknown" } -// isPaymentIntegrationConfigWrite is kept as a compatibility seam for -// middleware tests; policy ownership lives in routecatalog. -func isPaymentIntegrationConfigWrite(method, path string) bool { - return routecatalog.BodyPolicyFor(method, path) == routecatalog.BodyPolicyPaymentConfig -} - func paymentCallbackSummary(body []byte, contentType string) string { mediaType, _, err := mime.ParseMediaType(contentType) if err != nil || mediaType == "" { diff --git a/internal/server/middleware/audit.go b/internal/server/middleware/audit.go index 931b0d2..fd30291 100644 --- a/internal/server/middleware/audit.go +++ b/internal/server/middleware/audit.go @@ -135,12 +135,24 @@ func operationRequestBody(raw []byte, contentType string, limit int) string { return text } +// operationSecretKeys holds the normalized JSON keys whose values never reach +// the operation record. Keys are compared after lowercasing and stripping +// separators, so "new_password" and "newPassword" both match "newpassword". +var operationSecretKeys = map[string]struct{}{ + "password": {}, "newpassword": {}, "oldpassword": {}, "confirmpassword": {}, + "passwd": {}, "pwd": {}, "token": {}, "accesstoken": {}, "refreshtoken": {}, + "secret": {}, "clientsecret": {}, "apikey": {}, "privatekey": {}, "idcard": {}, + "appkey": {}, "mchkey": {}, "apiv3key": {}, "clientcert": {}, "clientkey": {}, + "platformcert": {}, "platformserialno": {}, "credentialcode": {}, "certfile": {}, + "keyfile": {}, "publickey": {}, "rootcert": {}, "appcert": {}, "webhookid": {}, +} + func maskOperationBody(value any) { switch current := value.(type) { case map[string]any: for key, item := range current { normalized := strings.ToLower(strings.ReplaceAll(strings.ReplaceAll(key, "_", ""), "-", "")) - if normalized == "password" || normalized == "newpassword" || normalized == "oldpassword" || normalized == "confirmpassword" || normalized == "passwd" || normalized == "pwd" || normalized == "token" || normalized == "accesstoken" || normalized == "refreshtoken" || normalized == "secret" || normalized == "clientsecret" || normalized == "apikey" || normalized == "privatekey" || normalized == "idcard" || normalized == "appkey" || normalized == "mchkey" || normalized == "apiv3key" || normalized == "clientcert" || normalized == "clientkey" || normalized == "platformcert" || normalized == "platformserialno" || normalized == "credentialcode" || normalized == "certfile" || normalized == "keyfile" || normalized == "publickey" || normalized == "rootcert" || normalized == "appcert" || normalized == "webhookid" { + if _, secret := operationSecretKeys[normalized]; secret { current[key] = "***" continue } @@ -153,21 +165,26 @@ func maskOperationBody(value any) { } } -func isDownloadResponse(c *gin.Context) bool { - header := c.Writer.Header() - return strings.Contains(header.Get("Pragma"), "public") || - strings.Contains(header.Get("Expires"), "0") || - strings.Contains(header.Get("Cache-Control"), "must-revalidate, post-check=0, pre-check=0") || - strings.Contains(header.Get("Content-Type"), "application/force-download") || - strings.Contains(header.Get("Content-Type"), "application/octet-stream") || - strings.Contains(header.Get("Content-Type"), "application/vnd.ms-excel") || - strings.Contains(header.Get("Content-Type"), "application/download") || - strings.Contains(header.Get("Content-Disposition"), "attachment") || - strings.Contains(header.Get("Content-Transfer-Encoding"), "binary") +// operationDownloadHeaders lists the response headers whose presence marks a +// download: the value is the substring that identifies it. +var operationDownloadHeaders = [...]struct{ header, marker string }{ + {"Pragma", "public"}, + {"Expires", "0"}, + {"Cache-Control", "must-revalidate, post-check=0, pre-check=0"}, + {"Content-Type", "application/force-download"}, + {"Content-Type", "application/octet-stream"}, + {"Content-Type", "application/vnd.ms-excel"}, + {"Content-Type", "application/download"}, + {"Content-Disposition", "attachment"}, + {"Content-Transfer-Encoding", "binary"}, } -// recordsOperation remains as a small compatibility helper for tests and -// custom middleware chains that do not have a Gin context. -func recordsOperation(method, path string) bool { - return routecatalog.ShouldAudit(method, path) +func isDownloadResponse(c *gin.Context) bool { + header := c.Writer.Header() + for _, candidate := range operationDownloadHeaders { + if strings.Contains(header.Get(candidate.header), candidate.marker) { + return true + } + } + return false } diff --git a/internal/server/middleware/auth.go b/internal/server/middleware/auth.go index ffac3b6..1cc57c3 100644 --- a/internal/server/middleware/auth.go +++ b/internal/server/middleware/auth.go @@ -7,6 +7,7 @@ import ( "net/http" "strconv" "strings" + "time" "github.com/gin-gonic/gin" "golang.org/x/sync/singleflight" @@ -14,6 +15,10 @@ import ( const claimsKey = "admin_claims" +// tokenAuthTimeout bounds a shared token authentication flight. It replaces the +// leader request's own deadline, which followers must not inherit. +const tokenAuthTimeout = 10 * time.Second + var refreshTokens singleflight.Group type TokenAuthenticator interface { @@ -36,7 +41,7 @@ func AuthenticateWebSocket(c *gin.Context, auth TokenAuthenticator) bool { } func authenticate(c *gin.Context, auth TokenAuthenticator, allowQueryToken bool) bool { - token := requestToken(c, allowQueryToken) + token := RequestToken(c, allowQueryToken) if token == "" { NoAuth(c, "未登录或非法访问,请登录") return false @@ -45,8 +50,15 @@ func authenticate(c *gin.Context, auth TokenAuthenticator, allowQueryToken bool) NoAuth(c, "认证服务不可用") return false } + // Followers of a singleflight flight must not inherit the leader's request + // context: if the leader's client disconnects mid-flight, its cancellation + // would surface as an auth failure for every follower and force-log-out + // unrelated sessions. Detach cancellation and bound the shared call with an + // explicit timeout instead. value, err, _ := refreshTokens.Do(token, func() (any, error) { - return auth.AuthenticateToken(c.Request.Context(), token) + ctx, cancel := context.WithTimeout(context.WithoutCancel(c.Request.Context()), tokenAuthTimeout) + defer cancel() + return auth.AuthenticateToken(ctx, token) }) if err != nil { SetTokenCookie(c, "", -1) @@ -93,10 +105,6 @@ func RequestToken(c *gin.Context, allowQueryToken bool) string { return token } -func requestToken(c *gin.Context, allowQueryToken bool) string { - return RequestToken(c, allowQueryToken) -} - func tokenErrorMessage(err error) string { message := "无法处理此token" switch { diff --git a/internal/server/middleware/payment_audit_test.go b/internal/server/middleware/payment_audit_test.go index 60b7012..b995d4b 100644 --- a/internal/server/middleware/payment_audit_test.go +++ b/internal/server/middleware/payment_audit_test.go @@ -4,6 +4,8 @@ import ( "encoding/json" "strings" "testing" + + "kra/internal/routecatalog" ) func TestPaymentIntegrationSecretsAreRedacted(t *testing.T) { @@ -35,7 +37,7 @@ func TestPaymentOperationsAreAuditedWithRouterPrefix(t *testing.T) { {method: "POST", path: "/api/payment/fulfill"}, {method: "POST", path: "/api/payment/providers/alipay/test"}, } { - if !recordsOperation(route.method, route.path) { + if !routecatalog.ShouldAudit(route.method, route.path) { t.Fatalf("payment operation was not audited: %s %s", route.method, route.path) } } @@ -47,10 +49,10 @@ func TestPaymentIntegrationConfigUsesRouteLevelSummary(t *testing.T) { if strings.Contains(summary, "secret") || strings.Contains(summary, "certificate") { t.Fatalf("payment configuration summary leaked payload: %s", summary) } - if !isPaymentIntegrationConfigWrite("PUT", "/api/integration/configs/payment/saobei") { + if routecatalog.BodyPolicyFor("PUT", "/api/integration/configs/payment/saobei") != routecatalog.BodyPolicyPaymentConfig { t.Fatal("payment configuration write route was not recognized") } - if isPaymentIntegrationConfigWrite("PUT", "/api/integration/configs/mq/emqx") { + if routecatalog.BodyPolicyFor("PUT", "/api/integration/configs/mq/emqx") == routecatalog.BodyPolicyPaymentConfig { t.Fatal("non-payment integration was treated as payment configuration") } } diff --git a/internal/worker/task_executor.go b/internal/worker/task_executor.go index 3c533a0..97bb8ea 100644 --- a/internal/worker/task_executor.go +++ b/internal/worker/task_executor.go @@ -7,12 +7,15 @@ import ( "errors" "fmt" "io" + "log/slog" "net" "net/http" "net/url" "strings" + "sync" "syscall" "time" + "unicode/utf8" taskbiz "kra/internal/biz/task" ) @@ -20,6 +23,10 @@ import ( type TaskExecutor struct { tasks *taskbiz.TaskUsecase methods taskbiz.TaskMethodRegistry + // orphans tracks method goroutines that outlived their execution timeout. + // A method cannot be killed, so the task stays claimed until it returns. + orphanMu sync.Mutex + orphans map[uint]struct{} } func NewTaskExecutorWithRegistry(tasks *taskbiz.TaskUsecase, methods taskbiz.TaskMethodRegistry) *TaskExecutor { @@ -30,7 +37,21 @@ func privateIP(ip net.IP) bool { return ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalMulticast() || ip.IsLinkLocalUnicast() || ip.IsUnspecified() } +// Task HTTP clients are built once and shared: a per-request client would +// discard every pooled connection and never release its transport. +var ( + taskHTTPClientPublic = newTaskHTTPClient(false) + taskHTTPClientPrivate = newTaskHTTPClient(true) +) + func taskHTTPClient(allowPrivate bool) *http.Client { + if allowPrivate { + return taskHTTPClientPrivate + } + return taskHTTPClientPublic +} + +func newTaskHTTPClient(allowPrivate bool) *http.Client { dialer := &net.Dialer{Timeout: 10 * time.Second, Control: func(_ string, address string, _ syscall.RawConn) error { if allowPrivate { return nil @@ -48,7 +69,8 @@ func taskHTTPClient(allowPrivate bool) *http.Client { } return nil }} - return &http.Client{Timeout: 30 * time.Second, Transport: &http.Transport{Proxy: nil, DialContext: dialer.DialContext}} + transport := &http.Transport{Proxy: nil, DialContext: dialer.DialContext, MaxIdleConnsPerHost: 4, IdleConnTimeout: 90 * time.Second} + return &http.Client{Timeout: 30 * time.Second, Transport: transport} } func (e *TaskExecutor) runHTTP(ctx context.Context, task *taskbiz.TimedTask) (string, error) { @@ -104,22 +126,67 @@ func (e *TaskExecutor) runHTTP(ctx context.Context, task *taskbiz.TimedTask) (st return output, nil } -var errTaskTimeout = errors.New("任务执行超时") +var ( + errTaskTimeout = errors.New("任务执行超时") + errTaskOrphaned = errors.New("上一次执行超时后仍在运行, 本次执行已跳过") +) func truncateTaskText(value string) string { const limit = 4000 if len(value) <= limit { return value } - return value[:limit] + "...(截断)" + // Cut on a rune boundary: slicing raw bytes would split multi-byte text + // (Chinese output in particular) into an invalid UTF-8 sequence. + cut := limit + for cut > 0 && !utf8.RuneStart(value[cut]) { + cut-- + } + return value[:cut] + "...(截断)" } func (e *TaskExecutor) recordTaskLog(ctx context.Context, log *taskbiz.TimedTaskLog) { - defer func() { _ = recover() }() if e == nil || e.tasks == nil || log == nil { return } - _ = e.tasks.RecordTaskLog(ctx, log) + // This runs from Run's deferred block, so a panic here would replace the + // task's own result with a log-writing panic. Recover, but report it: a + // silently swallowed panic hides the fact that no execution log was stored. + defer func() { + if recovered := recover(); recovered != nil { + slog.Default().Error("recording the timed task log panicked", "mod", "timedTask", "task_id", log.TaskID, "task_name", log.TaskName, "panic", recovered) + } + }() + if err := e.tasks.RecordTaskLog(ctx, log); err != nil { + slog.Default().Error("recording the timed task log failed", "mod", "timedTask", "task_id", log.TaskID, "task_name", log.TaskName, "error", err) + } +} + +// methodOrphaned reports whether a previous invocation of the task's method is +// still running after its execution timeout elapsed. +func (e *TaskExecutor) methodOrphaned(id uint) bool { + e.orphanMu.Lock() + defer e.orphanMu.Unlock() + _, orphaned := e.orphans[id] + return orphaned +} + +// markMethodOrphan keeps the task claimed until the abandoned goroutine +// returns. Without it the scheduler's running-set is cleared as soon as the +// timeout fires, so the next trigger would run the same method concurrently. +func (e *TaskExecutor) markMethodOrphan(id uint, done <-chan error) { + e.orphanMu.Lock() + if e.orphans == nil { + e.orphans = map[uint]struct{}{} + } + e.orphans[id] = struct{}{} + e.orphanMu.Unlock() + go func() { + <-done + e.orphanMu.Lock() + delete(e.orphans, id) + e.orphanMu.Unlock() + }() } func (e *TaskExecutor) runMethod(ctx context.Context, task *taskbiz.TimedTask) error { @@ -140,6 +207,9 @@ func (e *TaskExecutor) runMethod(ctx context.Context, task *taskbiz.TimedTask) e if !ok { return fmt.Errorf("方法 %s 未注册(需通过 platform/task.Registry 注册)", task.MethodName) } + if e.methodOrphaned(task.ID) { + return errTaskOrphaned + } runCtx, cancel := context.WithTimeout(ctx, 5*time.Minute) defer cancel() done := make(chan error, 1) @@ -158,6 +228,7 @@ func (e *TaskExecutor) runMethod(ctx context.Context, task *taskbiz.TimedTask) e } return err case <-runCtx.Done(): + e.markMethodOrphan(task.ID, done) if errors.Is(runCtx.Err(), context.DeadlineExceeded) { return errTaskTimeout } diff --git a/internal/worker/task_scheduler.go b/internal/worker/task_scheduler.go index f419e67..92d17af 100644 --- a/internal/worker/task_scheduler.go +++ b/internal/worker/task_scheduler.go @@ -120,14 +120,7 @@ func (s *TaskScheduler) Stop(ctx context.Context) error { } // Keep the lifecycle lock while closing subscriber channels so a new SSE // subscription cannot arrive between the shutdown flag and the close. - s.subMu.Lock() - for userID, subscribers := range s.subscribers { - for ch := range subscribers { - close(ch) - } - delete(s.subscribers, userID) - } - s.subMu.Unlock() + s.closeSubscribers() s.runMu.Unlock() s.ctxMu.Lock() cancel := s.cancel @@ -323,7 +316,7 @@ func (s *TaskScheduler) run(ctx context.Context, task *taskbiz.TimedTask, trigge } alertCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 3*time.Second) defer cancel() - ids, err := s.authorities.AuthorityUserIDs(alertCtx, 888) + ids, err := s.authorities.AuthorityUserIDs(alertCtx, system.SuperAdminAuthorityID) if err != nil { logger.Error("query timed task alert recipients failed", "error", err) return @@ -368,10 +361,10 @@ func parseTaskSchedule(task *taskbiz.TimedTask) (cron.Schedule, error) { } func cloneTimedTask(task *taskbiz.TimedTask) *taskbiz.TimedTask { - copy := *task - copy.Params = append([]byte(nil), task.Params...) - copy.HTTPHeader = append([]byte(nil), task.HTTPHeader...) - return © + clone := *task + clone.Params = append([]byte(nil), task.Params...) + clone.HTTPHeader = append([]byte(nil), task.HTTPHeader...) + return &clone } func (s *TaskScheduler) scheduleLocked(task *taskbiz.TimedTask, schedule cron.Schedule) { @@ -472,10 +465,10 @@ func (s *TaskScheduler) Subscribe(userID uint) chan []byte { s.runMu.Unlock() return ch } + s.subMu.Lock() if s.subscribers == nil { s.subscribers = make(map[uint]map[chan []byte]struct{}) } - s.subMu.Lock() if s.subscribers[userID] == nil { s.subscribers[userID] = map[chan []byte]struct{}{} }