106 lines
2.8 KiB
Go
106 lines
2.8 KiB
Go
// 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...)
|
|
}
|
|
}
|