Files
rill/internal/database/logger.go
T
2026-09-19 16:02:30 +08:00

82 lines
2.0 KiB
Go

package database
import (
"context"
"errors"
"fmt"
"log/slog"
"time"
gormlogger "gorm.io/gorm/logger"
)
// slowThreshold 超过该耗时的 SQL 记为慢查询。
const slowThreshold = 500 * time.Millisecond
// slogLogger 将 GORM 日志转接到 slog。
type slogLogger struct {
logger *slog.Logger
level gormlogger.LogLevel
}
func newLogger(level slog.Level) gormlogger.Interface {
return &slogLogger{
logger: slog.Default(),
level: gormLogLevel(level),
}
}
func gormLogLevel(level slog.Level) gormlogger.LogLevel {
switch {
case level <= slog.LevelDebug:
return gormlogger.Info
case level <= slog.LevelWarn:
return gormlogger.Warn
default:
return gormlogger.Error
}
}
func (l *slogLogger) LogMode(level gormlogger.LogLevel) gormlogger.Interface {
next := *l
next.level = level
return &next
}
func (l *slogLogger) Info(ctx context.Context, msg string, args ...any) {
if l.level >= gormlogger.Info {
l.logger.InfoContext(ctx, fmt.Sprintf(msg, args...))
}
}
func (l *slogLogger) Warn(ctx context.Context, msg string, args ...any) {
if l.level >= gormlogger.Warn {
l.logger.WarnContext(ctx, fmt.Sprintf(msg, args...))
}
}
func (l *slogLogger) Error(ctx context.Context, msg string, args ...any) {
if l.level >= gormlogger.Error {
l.logger.ErrorContext(ctx, fmt.Sprintf(msg, args...))
}
}
func (l *slogLogger) Trace(ctx context.Context, begin time.Time, fc func() (sql string, rowsAffected int64), err error) {
if l.level <= gormlogger.Silent {
return
}
elapsed := time.Since(begin)
sql, rows := fc()
attrs := []any{"elapsed", elapsed, "rows", rows, "sql", sql}
switch {
case err != nil && l.level >= gormlogger.Error && !errors.Is(err, gormlogger.ErrRecordNotFound):
l.logger.ErrorContext(ctx, "数据库执行出错", append(attrs, "err", err)...)
case elapsed > slowThreshold && l.level >= gormlogger.Warn:
l.logger.WarnContext(ctx, "慢查询", attrs...)
case l.level >= gormlogger.Info:
l.logger.DebugContext(ctx, "数据库执行", attrs...)
}
}