Files
junhong_cmp_fiber/pkg/logger/middleware.go
break 17782d5f8e
All checks were successful
构建并部署到测试环境(无 SSH) / build-and-deploy (push) Successful in 9m20s
实现七月迭代公共技术基础
2026-07-23 17:52:48 +09:00

240 lines
7.2 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 logger
import (
"context"
"crypto/sha256"
"encoding/hex"
"net/url"
"strings"
"time"
"github.com/break/junhong_cmp_fiber/pkg/constants"
"github.com/bytedance/sonic"
"github.com/gofiber/fiber/v2"
"go.uber.org/zap"
)
const (
// MaxBodyLogSize 限制记录的请求/响应 body 大小为 50KB
MaxBodyLogSize = 50 * 1024
redactedValue = "[已脱敏]"
)
// BodyPolicy 表示访问日志正文记录策略。
type BodyPolicy int
const (
// BodyPolicyJSON 表示只允许记录脱敏后的 JSON 正文。
BodyPolicyJSON BodyPolicy = iota + 1
// BodyPolicySummary 表示只记录不可逆安全摘要。
BodyPolicySummary
)
// SanitizedContent 表示访问日志可安全记录的正文或查询摘要。
type SanitizedContent struct {
Content string
Size int
SHA256 string
Truncated bool
}
// truncateBody 截断 body 到指定大小
func truncateBody(body []byte, maxSize int) (string, bool) {
if len(body) == 0 {
return "", false
}
if len(body) <= maxSize {
return string(body), false
}
return string(body[:maxSize]), true
}
// maskSensitiveValue 按字段名判断并脱敏访问日志中的敏感值
func maskSensitiveValue(key, value string) string {
if value == "" {
return value
}
if shouldMaskField(key) {
return redactedValue
}
return value
}
// sanitizeQuery 脱敏 query 中的密码、签名、nonce、token 等字段
func sanitizeQuery(rawQuery string) SanitizedContent {
if rawQuery == "" {
return SanitizedContent{}
}
values, err := url.ParseQuery(rawQuery)
if err != nil {
return summarize([]byte(rawQuery))
}
for key, items := range values {
for index, item := range items {
items[index] = maskSensitiveValue(key, item)
}
values[key] = items
}
content, truncated := truncateBody([]byte(values.Encode()), MaxBodyLogSize)
return SanitizedContent{Content: content, Size: len(rawQuery), SHA256: digest([]byte(rawQuery)), Truncated: truncated}
}
// sanitizeBody 按策略脱敏正文后再写入访问日志。
func sanitizeBody(rawBody []byte, policy BodyPolicy) SanitizedContent {
if len(rawBody) == 0 {
return SanitizedContent{}
}
if policy == BodyPolicySummary {
return summarize(rawBody)
}
var payload any
if err := sonic.Unmarshal(rawBody, &payload); err != nil {
return summarize(rawBody)
}
sanitizeJSONValue(payload)
data, err := sonic.Marshal(payload)
if err != nil {
return summarize(rawBody)
}
content, truncated := truncateBody(data, MaxBodyLogSize)
return SanitizedContent{Content: content, Size: len(rawBody), SHA256: digest(rawBody), Truncated: truncated}
}
func summarize(raw []byte) SanitizedContent {
return SanitizedContent{Content: "[仅记录安全摘要]", Size: len(raw), SHA256: digest(raw), Truncated: len(raw) > MaxBodyLogSize}
}
func digest(raw []byte) string {
sum := sha256.Sum256(raw)
return hex.EncodeToString(sum[:])
}
// sanitizeJSONValue 递归脱敏 JSON 对象中的敏感字段
func sanitizeJSONValue(value any) {
switch typed := value.(type) {
case map[string]any:
for key, item := range typed {
if masked, ok := item.(string); ok {
typed[key] = maskSensitiveValue(key, masked)
continue
}
if shouldMaskField(key) {
typed[key] = redactedValue
continue
}
sanitizeJSONValue(item)
}
case []any:
for _, item := range typed {
sanitizeJSONValue(item)
}
}
}
// shouldMaskField 判断字段名是否属于访问日志敏感字段
func shouldMaskField(key string) bool {
normalized := strings.ToLower(key)
return strings.Contains(normalized, "password") ||
strings.Contains(normalized, "passwd") ||
strings.Contains(normalized, "credential") ||
strings.Contains(normalized, "authorization") ||
strings.Contains(normalized, "cookie") ||
strings.Contains(normalized, "key") ||
strings.Contains(normalized, "url") ||
strings.Contains(normalized, "qr_content") ||
strings.Contains(normalized, "verification_code") ||
strings.Contains(normalized, "sign") ||
strings.Contains(normalized, "nonce") ||
strings.Contains(normalized, "token") ||
strings.Contains(normalized, "secret")
}
// Middleware 创建 Fiber 日志中间件
// 记录所有 HTTP 请求到访问日志(包括请求和响应 body
func Middleware() fiber.Handler {
return MiddlewareWithLogger(GetAccessLogger())
}
// MiddlewareWithLogger 创建可注入访问日志器的 Fiber 中间件。
// 生产环境使用 Middleware该入口让集成测试捕获最终 JSON 日志而无需修改全局状态。
func MiddlewareWithLogger(accessLogger *zap.Logger) fiber.Handler {
return func(c *fiber.Ctx) error {
// 记录请求开始时间
startTime := time.Now()
c.Locals(constants.ContextKeyStartTime, startTime)
// 注入请求上下文,供 Service 层审计日志复用
ctx := c.UserContext()
if rid := c.Locals(constants.ContextKeyRequestID); rid != nil {
if requestID, ok := rid.(string); ok && requestID != "" {
ctx = context.WithValue(ctx, constants.ContextKeyRequestID, requestID)
}
}
ctx = context.WithValue(ctx, constants.ContextKeyIP, c.IP())
ctx = context.WithValue(ctx, constants.ContextKeyUserAgent, c.Get("User-Agent"))
ctx = context.WithValue(ctx, constants.ContextKeyRequestPath, c.Path())
ctx = context.WithValue(ctx, constants.ContextKeyRequestMethod, c.Method())
c.SetUserContext(ctx)
// 获取请求 body在 c.Next() 之前读取)
policy := policyForPath(c.Path())
requestBody := sanitizeWithPolicy(c.Body(), c.Get("Content-Type"), policy)
// 获取 query 参数
queryParams := sanitizeQuery(string(c.Request().URI().QueryString()))
if policy.Sensitive && queryParams.Size > 0 {
queryParams = summarize([]byte(c.Request().URI().QueryString()))
}
// 处理请求
err := c.Next()
// 计算请求持续时间
duration := time.Since(startTime)
// 获取请求 ID由 requestid 中间件设置)
requestID := ""
if rid := c.Locals(constants.ContextKeyRequestID); rid != nil {
requestID = rid.(string)
}
// 获取用户 ID由 auth 中间件设置)
var userID uint
if uid := c.Locals(constants.ContextKeyUserID); uid != nil {
if id, ok := uid.(uint); ok {
userID = id
}
}
// 获取响应 body
responseBody := sanitizeWithPolicy(c.Response().Body(), string(c.Response().Header.ContentType()), policy)
// 记录访问日志
accessLogger.Info("",
zap.String("method", c.Method()),
zap.String("path", c.Path()),
zap.String("body_policy", policy.Name),
zap.String("query", queryParams.Content),
zap.Bool("query_truncated", queryParams.Truncated),
zap.Int("status", c.Response().StatusCode()),
zap.Float64("duration_ms", float64(duration.Microseconds())/1000.0),
zap.String("request_id", requestID),
zap.String("ip", c.IP()),
zap.String("user_agent", c.Get("User-Agent")),
zap.Uint("user_id", userID),
zap.String("request_body", requestBody.Content),
zap.Int("request_body_size", requestBody.Size),
zap.String("request_body_sha256", requestBody.SHA256),
zap.Bool("request_body_truncated", requestBody.Truncated),
zap.String("response_body", responseBody.Content),
zap.Int("response_body_size", responseBody.Size),
zap.String("response_body_sha256", responseBody.SHA256),
zap.Bool("response_body_truncated", responseBody.Truncated),
)
return err
}
}