package database import ( "context" "errors" "fmt" "strings" "time" appLogging "hfb_sys/backend/internal/logging" "go.uber.org/zap" "gorm.io/gorm" gormLogger "gorm.io/gorm/logger" ) type structuredGormLogger struct { logger *zap.Logger level gormLogger.LogLevel slowThreshold time.Duration } func newGormLogger(logLevel string, logger *zap.Logger) gormLogger.Interface { return newGormLoggerWithThreshold(logLevel, logger, 500*time.Millisecond) } func newGormLoggerWithThreshold(logLevel string, logger *zap.Logger, slowThreshold time.Duration) gormLogger.Interface { level := gormLogger.Warn if strings.EqualFold(strings.TrimSpace(logLevel), "debug") { level = gormLogger.Info } if logger == nil { logger = zap.L() } if slowThreshold <= 0 { slowThreshold = 500 * time.Millisecond } return &structuredGormLogger{ logger: logger.With(zap.String("module", "database")), level: level, slowThreshold: slowThreshold, } } func (l *structuredGormLogger) LogMode(level gormLogger.LogLevel) gormLogger.Interface { cloned := *l cloned.level = level return &cloned } func (l *structuredGormLogger) Info(ctx context.Context, message string, args ...any) { if l.level >= gormLogger.Info { l.withContext(ctx).Debug("数据库信息", zap.String("detail", fmt.Sprintf(message, args...))) } } func (l *structuredGormLogger) Warn(ctx context.Context, message string, args ...any) { if l.level >= gormLogger.Warn { l.withContext(ctx).Warn("数据库警告", zap.String("detail", fmt.Sprintf(message, args...))) } } func (l *structuredGormLogger) Error(ctx context.Context, message string, args ...any) { if l.level >= gormLogger.Error { l.withContext(ctx).Error("数据库错误", zap.String("detail", fmt.Sprintf(message, args...))) } } func (l *structuredGormLogger) Trace(ctx context.Context, begin time.Time, query func() (string, int64), err error) { if l.level == gormLogger.Silent { return } elapsed := time.Since(begin) switch { case err != nil && !errors.Is(err, gorm.ErrRecordNotFound) && l.level >= gormLogger.Error: sql, rows := query() l.withContext(ctx).Error("数据库查询失败", queryFields(sql, rows, elapsed, err)...) case elapsed >= l.slowThreshold && l.level >= gormLogger.Warn: sql, rows := query() l.withContext(ctx).Warn("数据库慢查询", queryFields(sql, rows, elapsed, nil)...) case l.level == gormLogger.Info: sql, rows := query() l.withContext(ctx).Debug("数据库查询", queryFields(sql, rows, elapsed, nil)...) } } // ParamsFilter 让 GORM 保留 SQL 占位符,避免查询参数进入日志。 func (l *structuredGormLogger) ParamsFilter(_ context.Context, sql string, _ ...any) (string, []any) { return sql, nil } func (l *structuredGormLogger) withContext(ctx context.Context) *zap.Logger { fields := make([]zap.Field, 0, 2) if requestID := appLogging.RequestIDFromContext(ctx); requestID != "" { fields = append(fields, zap.String("request_id", requestID)) } if adminID := appLogging.AdminIDFromContext(ctx); adminID != 0 { fields = append(fields, zap.Uint64("admin_id", adminID)) } return l.logger.With(fields...) } func queryFields(sql string, rows int64, elapsed time.Duration, err error) []zap.Field { fields := []zap.Field{ zap.Float64("duration_ms", float64(elapsed.Microseconds())/1000), zap.Int64("rows", rows), zap.String("sql", sql), } if err != nil { fields = append(fields, zap.Error(err)) } return fields }