// Package gormkit contains reusable GORM integration helpers. package gormkit import ( "context" "database/sql/driver" "errors" "fmt" "io" "log/slog" "time" "gorm.io/gorm" "gorm.io/gorm/logger" "gorm.io/gorm/schema" ) // JSON is a portable JSON value for GORM models. It selects a native JSON // column where the driver supports one and a textual fallback elsewhere. type JSON []byte func (JSON) GormDataType() string { return "json" } func (JSON) GormDBDataType(db *gorm.DB, _ *schema.Field) string { switch db.Dialector.Name() { case "mysql", "sqlite": return "JSON" case "postgres": return "JSONB" case "sqlserver": return "NVARCHAR(MAX)" case "oracle": return "CLOB" default: return "TEXT" } } func (j JSON) Value() (driver.Value, error) { if len(j) == 0 { return nil, nil } return string(j), nil } func (j *JSON) Scan(value any) error { switch raw := value.(type) { case nil: *j = nil case []byte: *j = append((*j)[:0], raw...) case string: *j = append((*j)[:0], raw...) default: return fmt.Errorf("cannot scan JSON from %T", value) } return nil } // Logger forwards GORM diagnostics through an application slog logger. type Logger struct { logger *slog.Logger slowThreshold time.Duration level logger.LogLevel } func NewLogger(log *slog.Logger, level logger.LogLevel) *Logger { if log == nil { log = slog.New(slog.NewTextHandler(io.Discard, nil)) } return &Logger{logger: log, slowThreshold: 200 * time.Millisecond, level: level} } func (g *Logger) LogMode(level logger.LogLevel) logger.Interface { next := *g next.level = level return &next } func (g *Logger) Info(ctx context.Context, message string, args ...any) { g.logger.InfoContext(ctx, fmt.Sprintf(message, args...), "mod", "sql", "gorm_logger", true) } func (g *Logger) Warn(ctx context.Context, message string, args ...any) { g.logger.WarnContext(ctx, fmt.Sprintf(message, args...), "mod", "sql", "gorm_logger", true) } func (g *Logger) Error(ctx context.Context, message string, args ...any) { g.logger.ErrorContext(ctx, fmt.Sprintf(message, args...), "mod", "sql", "gorm_logger", true) } func (g *Logger) Trace(ctx context.Context, begin time.Time, query func() (string, int64), queryErr error) { if g.level <= logger.Silent { return } elapsed := time.Since(begin) sql, rows := query() fields := []any{"mod", "sql", "gorm_logger", true, "sql", sql, "rows", rows, "elapsed_ms", elapsed.Milliseconds()} switch { case queryErr != nil && g.level >= logger.Error && !errors.Is(queryErr, logger.ErrRecordNotFound): fields = append(fields, "error", queryErr) g.logger.ErrorContext(ctx, "SQL 执行错误", fields...) case elapsed > g.slowThreshold && g.level >= logger.Warn: g.logger.WarnContext(ctx, "SQL 慢查询", fields...) case g.level >= logger.Info: g.logger.InfoContext(ctx, "SQL", fields...) } }