kra-new/pkg/database/gormkit/logger.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...)
}
}