feat(middleware): clean up logging middleware, add span attributes

Co-authored-by: sindrerh2 <sindre.rodseth.hansen@nav.no>
This commit is contained in:
Trong Huu Nguyen
2025-01-30 14:03:29 +01:00
co-authored by sindrerh2
parent 98cc534806
commit 10360958c0
5 changed files with 80 additions and 71 deletions
+1 -1
View File
@@ -10,7 +10,7 @@ local: fmt
--bind-address=127.0.0.1:3000 \
--upstream-host=localhost:4000 \
--redis.uri=redis://localhost:6379 \
--log-level=debug \
--log-level=info \
--log-format=text
test: fmt
+1
View File
@@ -47,6 +47,7 @@ func SetupLogger(level, format string) error {
log.FatalLevel,
log.ErrorLevel,
log.WarnLevel,
log.InfoLevel,
)))
return nil
+75 -67
View File
@@ -10,28 +10,45 @@ import (
"github.com/go-chi/chi/v5/middleware"
httpinternal "github.com/nais/wonderwall/internal/http"
log "github.com/sirupsen/logrus"
"go.opentelemetry.io/otel/attribute"
"go.opentelemetry.io/otel/trace"
"github.com/nais/wonderwall/pkg/cookie"
"github.com/nais/wonderwall/pkg/router/paths"
)
var logger *requestLogger
type LogEntryMiddleware struct{}
// LogEntry is copied verbatim from httplog package to replace with our own requestLogger implementation.
func LogEntry(provider string) LogEntryMiddleware {
logger = &requestLogger{Logger: log.StandardLogger(), Provider: provider}
return LogEntryMiddleware{}
type logger struct {
Logger *log.Logger
Provider string
}
func (l *LogEntryMiddleware) Handler(next http.Handler) http.Handler {
// Logger provides a middleware that logs requests and responses.
func Logger(provider string) logger {
return logger{
Logger: log.StandardLogger(),
Provider: provider,
}
}
// LogEntryFrom returns a log entry from the request context.
func LogEntryFrom(r *http.Request) *log.Entry {
ctx := r.Context()
entry, ok := ctx.Value(middleware.LogEntryCtxKey).(*logEntryAdapter)
if ok {
return entry.Logger
}
return log.NewEntry(log.StandardLogger()).
WithFields(requestFields(r)).
WithFields(traceFields(r))
}
func (l *logger) Handler(next http.Handler) http.Handler {
fn := func(w http.ResponseWriter, r *http.Request) {
entry := logger.NewLogEntry(r)
entry := l.newLogEntry(r)
ww := middleware.NewWrapResponseWriter(w, r.ProtoMajor)
if !strings.HasSuffix(r.URL.Path, paths.Ping) {
entry.Logger.Debugf("request start: %s - %s", r.Method, r.URL.Path)
t1 := time.Now()
defer func() {
entry.Write(ww.Status(), ww.BytesWritten(), ww.Header(), time.Since(t1), nil)
@@ -43,28 +60,46 @@ func (l *LogEntryMiddleware) Handler(next http.Handler) http.Handler {
return http.HandlerFunc(fn)
}
func LogEntryFrom(r *http.Request) *log.Entry {
ctx := r.Context()
val := ctx.Value(middleware.LogEntryCtxKey)
entry, ok := val.(*requestLoggerEntry)
if ok {
return entry.Logger
func (l *logger) newLogEntry(r *http.Request) *logEntryAdapter {
return &logEntryAdapter{
requestFields: requestFields(r),
Logger: l.Logger.WithContext(r.Context()).
WithField("provider", l.Provider).
WithFields(traceFields(r)),
}
}
// logEntryAdapter implements [middleware.LogEntry]
type logEntryAdapter struct {
Logger *log.Entry
requestFields log.Fields
}
func (l *logEntryAdapter) Write(status, bytes int, _ http.Header, elapsed time.Duration, _ any) {
responseFields := log.Fields{
"response_status": status,
"response_bytes": bytes,
"response_elapsed_ms": float64(elapsed.Nanoseconds()) / 1000000.0, // in milliseconds, with fractional
}
entry = logger.NewLogEntry(r)
return entry.Logger
l.Logger.WithFields(l.requestFields).
WithFields(responseFields).
Debugf("response: %d %s", status, http.StatusText(status))
}
type requestLogger struct {
Logger *log.Logger
Provider string
}
func (l *logEntryAdapter) Panic(v interface{}, _ []byte) {
stacktrace := "#"
func (l *requestLogger) NewLogEntry(r *http.Request) *requestLoggerEntry {
entry := &requestLoggerEntry{}
fields := log.Fields{
"correlation_id": middleware.GetReqID(r.Context()),
"provider": l.Provider,
"stacktrace": stacktrace,
"error": fmt.Sprintf("%+v", v),
}
l.Logger = l.Logger.WithFields(fields)
}
func requestFields(r *http.Request) log.Fields {
fields := log.Fields{
"request_cookies": nonEmptyRequestCookies(r),
"request_host": r.Host,
"request_is_navigational": httpinternal.IsNavigationRequest(r),
@@ -78,52 +113,25 @@ func (l *requestLogger) NewLogEntry(r *http.Request) *requestLoggerEntry {
"request_user_agent": r.UserAgent(),
}
entry.Logger = l.Logger.
WithContext(r.Context()).
WithFields(fields)
return entry
}
type requestLoggerEntry struct {
Logger *log.Entry
}
func (l *requestLoggerEntry) Write(status, bytes int, _ http.Header, elapsed time.Duration, _ any) {
msg := fmt.Sprintf("request end: HTTP %d (%s)", status, statusLabel(status))
fields := log.Fields{
"response_status": status,
"response_bytes": bytes,
"response_elapsed_ms": float64(elapsed.Nanoseconds()) / 1000000.0, // in milliseconds, with fractional
span := trace.SpanFromContext(r.Context())
for k, v := range fields {
attrKey := "wonderwall." + k
span.SetAttributes(attribute.String(attrKey, fmt.Sprint(v)))
}
entry := l.Logger.WithFields(fields)
entry.Debugf(msg)
return fields
}
func (l *requestLoggerEntry) Panic(v interface{}, _ []byte) {
stacktrace := "#"
fields := log.Fields{
"stacktrace": stacktrace,
"error": fmt.Sprintf("%+v", v),
func traceFields(r *http.Request) log.Fields {
fields := log.Fields{}
span := trace.SpanFromContext(r.Context())
if span.SpanContext().HasTraceID() {
fields["trace_id"] = span.SpanContext().TraceID().String()
} else {
fields["correlation_id"] = middleware.GetReqID(r.Context())
}
l.Logger = l.Logger.WithFields(fields)
}
func statusLabel(status int) string {
switch {
case status >= 100 && status < 300:
return "OK"
case status >= 300 && status < 400:
return "Redirect"
case status >= 400 && status < 500:
return "Client Error"
case status >= 500:
return "Server Error"
default:
return "Unknown"
}
return fields
}
func nonEmptyRequestCookies(r *http.Request) string {
+1 -1
View File
@@ -19,7 +19,7 @@ func NewGetRequest(target string, ingresses *ingress.Ingresses) *http.Request {
req = mw.RequestWithIngress(req, ing)
}
mw.LogEntry("test")
mw.Logger("test")
return req
}
+2 -2
View File
@@ -52,7 +52,7 @@ func New(src Source, cfg *config.Config) chi.Router {
providerName := string(cfg.OpenID.Provider)
ingressMw := middleware.Ingress(src)
prometheus := middleware.Prometheus(providerName)
logentry := middleware.LogEntry(providerName)
logger := middleware.Logger(providerName)
r := chi.NewRouter()
if cfg.OpenTelemetry.Enabled {
@@ -64,6 +64,7 @@ func New(src Source, cfg *config.Config) chi.Router {
r.Use(middleware.CorrelationIDHandler)
r.Use(chi_middleware.Recoverer)
r.Use(ingressMw.Handler)
r.Use(logger.Handler)
prefixes := src.GetIngresses().Paths()
@@ -72,7 +73,6 @@ func New(src Source, cfg *config.Config) chi.Router {
}
r.Group(func(r chi.Router) {
r.Use(logentry.Handler)
r.Use(prometheus.Handler)
r.Use(chi_middleware.NoCache)