Files
loggerx/middleware/traceId_test.go
T
2026-09-14 00:14:30 +08:00

155 lines
4.3 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package middleware_test
import (
"context"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"testing"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
"github.com/yuninks/loggerx"
"github.com/yuninks/loggerx/middleware"
)
func newTestEngine(log *loggerx.Logger) *gin.Engine {
gin.SetMode(gin.TestMode)
g := gin.New()
g.Use(middleware.SetGinTraceId(log))
return g
}
// trace id 应当写回响应头,并在 handler 里可读
func TestGinTraceIdSetsResponseHeader(t *testing.T) {
log := loggerx.NewLogger(context.Background(), loggerx.SetDir(t.TempDir()))
defer log.Close()
var seen string
g := newTestEngine(log)
g.GET("/ping", func(c *gin.Context) {
seen = middleware.GetTraceId(c.Request.Context(), log.GetTraceField())
c.Status(200)
})
w := httptest.NewRecorder()
g.ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/ping", nil))
header := w.Header().Get(middleware.TraceHeader)
if header == "" {
t.Fatalf("响应头 %s 没有写回,跨服务排障拿不到同一个 id", middleware.TraceHeader)
}
if _, err := uuid.Parse(header); err != nil {
t.Errorf("生成的 trace id 不是合法 uuid: %q", header)
}
if seen != header {
t.Errorf("handler 里读到的 trace id (%q) 与响应头 (%q) 不一致", seen, header)
}
}
// 上游带 trace id 时必须透传,不能重新生成
func TestGinTraceIdPropagatesFromHeader(t *testing.T) {
log := loggerx.NewLogger(context.Background(), loggerx.SetDir(t.TempDir()))
defer log.Close()
const upstream = "8f3d1c22-1111-4222-8333-444455556666"
var seen string
g := newTestEngine(log)
g.GET("/ping", func(c *gin.Context) {
seen = middleware.GetTraceId(c.Request.Context(), log.GetTraceField())
c.Status(200)
})
req := httptest.NewRequest(http.MethodGet, "/ping", nil)
req.Header.Set(middleware.TraceHeader, upstream)
w := httptest.NewRecorder()
g.ServeHTTP(w, req)
if seen != upstream {
t.Errorf("上游 trace id 未透传: 期望 %q 实际 %q", upstream, seen)
}
if got := w.Header().Get(middleware.TraceHeader); got != upstream {
t.Errorf("响应头应回传上游 trace id,实际 %q", got)
}
}
// 日志里应当带上 middleware 注入的 trace idgin.Context 与 request ctx 两条路径)
func TestTraceIdAppearsInLog(t *testing.T) {
dir := t.TempDir()
log := loggerx.NewLogger(context.Background(), loggerx.SetDir(dir))
defer log.Close()
g := newTestEngine(log)
g.GET("/ping", func(c *gin.Context) {
log.Info(c, "via-gin-ctx")
log.Info(c.Request.Context(), "via-request-ctx")
c.Status(200)
})
w := httptest.NewRecorder()
g.ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/ping", nil))
traceId := w.Header().Get(middleware.TraceHeader)
if traceId == "" {
t.Fatal("没有 trace id")
}
if err := log.MustSync(); err != nil {
t.Fatal(err)
}
content := readAll(t, dir)
if !contains(content, traceId) {
t.Errorf("日志里没有出现 trace id %q:\n%s", traceId, content)
}
if !contains(content, "via-gin-ctx") || !contains(content, "via-request-ctx") {
t.Errorf("两条日志都应当写入:\n%s", content)
}
}
// SetTraceId 打标:已存在时不覆盖,取不到时返回空串而不是 panic
func TestSetTraceIdByKeySemantics(t *testing.T) {
log := loggerx.NewLogger(context.Background(), loggerx.SetDir(t.TempDir()))
defer log.Close()
ctx := middleware.SetTraceId(context.Background(), log)
first := middleware.GetTraceId(ctx, log.GetTraceField())
if first == "" {
t.Fatal("SetTraceId 应当生成一个 id")
}
again := middleware.SetTraceId(ctx, log)
if got := middleware.GetTraceId(again, log.GetTraceField()); got != first {
t.Errorf("已有 trace id 时不该覆盖: %q -> %q", first, got)
}
if got := middleware.GetTraceId(nil, "trace_id"); got != "" {
t.Errorf("nil ctx 应返回空串,实际 %q", got)
}
if got := middleware.GetTraceId(context.Background(), "trace_id"); got != "" {
t.Errorf("没有值时该返回空串,实际 %q", got)
}
}
func readAll(t *testing.T, dir string) string {
t.Helper()
var sb []byte
files, _ := filepath.Glob(filepath.Join(dir, "*.log"))
for _, f := range files {
b, err := os.ReadFile(f)
if err != nil {
t.Fatalf("读取 %s: %v", f, err)
}
sb = append(sb, b...)
}
return string(sb)
}
func contains(s, sub string) bool {
for i := 0; i+len(sub) <= len(s); i++ {
if s[i:i+len(sub)] == sub {
return true
}
}
return false
}