Skip to content

拦截器与中间件模式

拦截器(Interceptor)是 gRPC 的横切关注点(cross-cutting concern)机制,类似 HTTP 框架中的中间件。本篇不再重复拦截器的基础语法,而是深入拦截器的实现原理、四种拦截器的完整实现、grpc-middleware 的链式模式,以及生产环境最常用的六类拦截器(日志、认证、恢复、请求 ID、限流、指标采集)的实战代码。

一、拦截器机制原理

gRPC 的拦截器本质上是一个函数包装器:在真正的 handler 前后插入自定义逻辑。框架在调度 RPC 时,会先经过拦截器链,再到达业务 handler。

客户端调用


客户端拦截器链 (UnaryInterceptor / StreamInterceptor)

    ▼ (HTTP/2 传输)

服务端拦截器链 (UnaryInterceptor / StreamInterceptor)


业务 handler

gRPC 有四种拦截器,对应服务端/客户端 × Unary/Stream:

拦截器类型接口签名关键点
服务端 Unary(ctx, req, *UnaryServerInfo, handler) -> (resp, err)
服务端 Stream(srv, ServerStream, *StreamServerInfo, handler) -> err
客户端 Unary(ctx, method, req, reply, *ClientConn, invoker, opts...) -> err
客户端 Stream(ctx, *StreamDesc, *ClientConn, method, streamer, opts...) -> (ClientStream, err)

二、服务端一元拦截器实现

go
package main

import (
	"context"
	"fmt"
	"log"
	"runtime/debug"
	"time"

	"google.golang.org/grpc"
	"google.golang.org/grpc/codes"
	"google.golang.org/grpc/status"
)

// === Recovery 拦截器:捕获 panic,防止进程崩溃 ===
func RecoveryUnaryInterceptor(ctx context.Context, req interface{},
	info *grpc.UnaryServerInfo, handler grpc.UnaryHandler,
) (resp interface{}, err error) {
	defer func() {
		if r := recover(); r != nil {
			log.Printf("[recovery] panic in %s: %v\n%s", info.FullMethod, r, debug.Stack())
			err = status.Errorf(codes.Internal, "internal error: %v", r)
		}
	}()
	return handler(ctx, req)
}

// === 日志拦截器:记录请求耗时和结果 ===
func LoggingUnaryInterceptor(ctx context.Context, req interface{},
	info *grpc.UnaryServerInfo, handler grpc.UnaryHandler,
) (interface{}, error) {
	start := time.Now()
	resp, err := handler(ctx, req)
	log.Printf("[log] %s cost=%v req=%T resp=%T err=%v",
		info.FullMethod, time.Since(start), req, resp, err)
	return resp, err
}

// === 认证拦截器:从 metadata 提取 token 并验证 ===
func AuthUnaryInterceptor(ctx context.Context, req interface{},
	info *grpc.UnaryServerInfo, handler grpc.UnaryHandler,
) (interface{}, error) {
	// 跳过不需要认证的方法
	if info.FullMethod == "/auth.AuthService/Login" {
		return handler(ctx, req)
	}

	// 从 context 中提取认证信息(metadata 传递见后文)
	token := extractTokenFromContext(ctx)
	if token == "" {
		return nil, status.Error(codes.Unauthenticated, "missing token")
	}
	userID, err := validateToken(token)
	if err != nil {
		return nil, status.Errorf(codes.Unauthenticated, "invalid token: %v", err)
	}

	// 将用户信息存入 context,传递给后续 handler
	ctx = context.WithValue(ctx, userIDKey{}, userID)
	return handler(ctx, req)
}

// === 辅助类型和函数 ===
type userIDKey struct{}

func extractTokenFromContext(ctx context.Context) string {
	// 简化:实际从 metadata.FromIncomingContext 提取
	return ""
}

func validateToken(token string) (string, error) {
	if token == "" {
		return "", fmt.Errorf("empty token")
	}
	return "user-123", nil
}

func main() {
	// 组装拦截器链:recovery(最外层兜底)-> logging -> auth(最内层靠近业务)
	server := grpc.NewServer(grpc.ChainUnaryInterceptor(
		RecoveryUnaryInterceptor,
		LoggingUnaryInterceptor,
		AuthUnaryInterceptor,
	))
	fmt.Printf("server with unary interceptor chain: %T\n", server)
}

三、服务端流拦截器实现

流拦截器的核心在于包装 grpc.ServerStream,从而拦截每一次 SendMsg/RecvMsg

go
package main

import (
	"context"
	"fmt"
	"log"
	"runtime/debug"
	"time"

	"google.golang.org/grpc"
	"google.golang.org/grpc/codes"
	"google.golang.org/grpc/status"
)

// 通用 stream 包装器
type wrappedServerStream struct {
	grpc.ServerStream
	method  string
	recvCnt int
	sendCnt int
}

func (w *wrappedServerStream) RecvMsg(m interface{}) error {
	err := w.ServerStream.RecvMsg(m)
	if err == nil {
		w.recvCnt++
	}
	return err
}

func (w *wrappedServerStream) SendMsg(m interface{}) error {
	err := w.ServerStream.SendMsg(m)
	if err == nil {
		w.sendCnt++
	}
	return err
}

// Recovery 流拦截器
func RecoveryStreamInterceptor(srv interface{}, ss grpc.ServerStream,
	info *grpc.StreamServerInfo, handler grpc.StreamHandler,
) (err error) {
	defer func() {
		if r := recover(); r != nil {
			log.Printf("[recovery-stream] panic in %s: %v\n%s", info.FullMethod, r, debug.Stack())
			err = status.Errorf(codes.Internal, "internal error: %v", r)
		}
	}()
	return handler(srv, ss)
}

// 日志流拦截器
func LoggingStreamInterceptor(srv interface{}, ss grpc.ServerStream,
	info *grpc.StreamServerInfo, handler grpc.StreamHandler,
) error {
	start := time.Now()
	wrapped := &wrappedServerStream{ServerStream: ss, method: info.FullMethod}
	err := handler(srv, wrapped)
	log.Printf("[log-stream] %s cost=%v recv=%d send=%d err=%v",
		info.FullMethod, time.Since(start), wrapped.recvCnt, wrapped.sendCnt, err)
	return err
}

func main() {
	server := grpc.NewServer(grpc.ChainStreamInterceptor(
		RecoveryStreamInterceptor,
		LoggingStreamInterceptor,
	))
	fmt.Printf("server with stream interceptor chain: %T\n", server)
	_ = context.Background
}

四、客户端一元拦截器实现

go
package main

import (
	"context"
	"fmt"
	"log"
	"time"

	"google.golang.org/grpc"
)

// 客户端日志拦截器
func ClientLoggingInterceptor(ctx context.Context, method string,
	req, reply interface{}, cc *grpc.ClientConn,
	invoker grpc.UnaryInvoker, opts ...grpc.CallOption,
) error {
	start := time.Now()
	err := invoker(ctx, method, req, reply, cc, opts...)
	log.Printf("[client-log] %s cost=%v err=%v", method, time.Since(start), err)
	return err
}

// 客户端 request-id 注入拦截器
func ClientRequestIDInterceptor(ctx context.Context, method string,
	req, reply interface{}, cc *grpc.ClientConn,
	invoker grpc.UnaryInvoker, opts ...grpc.CallOption,
) error {
	// 生成 request-id 并注入 metadata
	reqID := generateRequestID()
	// 实际实现用 metadata.AppendToOutgoingContext
	fmt.Printf("[client-reqid] %s request-id=%s\n", method, reqID)
	return invoker(ctx, method, req, reply, cc, opts...)
}

func generateRequestID() string {
	return fmt.Sprintf("req-%d", time.Now().UnixNano())
}

func main() {
	// 链式客户端拦截器
	conn, _ := grpc.NewClient("127.0.0.1:50051",
		grpc.WithChainUnaryInterceptor(
			ClientRequestIDInterceptor,
			ClientLoggingInterceptor,
		),
	)
	defer conn.Close()
	fmt.Printf("client with interceptor chain: %T\n", conn)
}

五、客户端流拦截器实现

go
package main

import (
	"context"
	"fmt"
	"log"
	"time"

	"google.golang.org/grpc"
)

type wrappedClientStream struct {
	grpc.ClientStream
	method  string
	recvCnt int
	sendCnt int
}

func (w *wrappedClientStream) RecvMsg(m interface{}) error {
	err := w.ClientStream.RecvMsg(m)
	if err == nil {
		w.recvCnt++
		log.Printf("[client-stream %s] recv #%d", w.method, w.recvCnt)
	}
	return err
}

func (w *wrappedClientStream) SendMsg(m interface{}) error {
	err := w.ClientStream.SendMsg(m)
	if err == nil {
		w.sendCnt++
		log.Printf("[client-stream %s] send #%d", w.method, w.sendCnt)
	}
	return err
}

func ClientStreamLoggingInterceptor(ctx context.Context, desc *grpc.StreamDesc,
	cc *grpc.ClientConn, method string, streamer grpc.Streamer,
	opts ...grpc.CallOption,
) (grpc.ClientStream, error) {
	start := time.Now()
	s, err := streamer(ctx, desc, cc, method, opts...)
	if err != nil {
		return nil, err
	}
	wrapped := &wrappedClientStream{ClientStream: s, method: method}

	// 用 finalizer 记录总耗时
	go func() {
		<-ctx.Done()
		log.Printf("[client-stream %s] total cost=%v recv=%d send=%d",
			method, time.Since(start), wrapped.recvCnt, wrapped.sendCnt)
	}()

	return wrapped, nil
}

func main() {
	conn, _ := grpc.NewClient("127.0.0.1:50051",
		grpc.WithStreamInterceptor(ClientStreamLoggingInterceptor),
	)
	defer conn.Close()
	fmt.Printf("client with stream interceptor: %T\n", conn)
}

六、拦截器链模式:grpc-middleware

grpc-middleware 是 gRPC 生态中最流行的拦截器工具库,提供了 ChainUnaryServer / ChainStreamServer 等便捷函数,以及一批开箱即用的拦截器。

go
package main

import (
	"context"
	"fmt"
	"log"
	"runtime/debug"
	"time"

	"google.golang.org/grpc"
	"google.golang.org/grpc/codes"
	"google.golang.org/grpc/status"
)

// grpc-middleware 的 ChainUnaryServer 等价实现(简化版)
// 官方库的实现在 github.com/grpc-ecosystem/go-grpc-middleware/v2/interceptors

func chainUnaryServer(interceptors []grpc.UnaryServerInterceptor) grpc.UnaryServerInterceptor {
	n := len(interceptors)
	if n == 0 {
		return func(ctx context.Context, req interface{},
			info *grpc.UnaryServerInfo, handler grpc.UnaryHandler,
		) (interface{}, error) {
			return handler(ctx, req)
		}
	}
	if n == 1 {
		return interceptors[0]
	}
	// 递归构建洋葱链
	return func(ctx context.Context, req interface{},
		info *grpc.UnaryServerInfo, handler grpc.UnaryHandler,
	) (interface{}, error) {
		chainer := func(currentInter grpc.UnaryServerInterceptor, currentHandler grpc.UnaryHandler) grpc.UnaryHandler {
			return func(currentCtx context.Context, currentReq interface{}) (interface{}, error) {
				return currentInter(currentCtx, currentReq, info, currentHandler)
			}
		}
		chi := handler
		for i := n - 1; i >= 0; i-- {
			chi = chainer(interceptors[i], chi)
		}
		return chi(ctx, req)
	}
}

func main() {
	interceptors := []grpc.UnaryServerInterceptor{
		recoveryInterceptor,
		loggingInterceptor,
		authInterceptor,
	}
	combined := chainUnaryServer(interceptors)

	// 模拟一次调用
	handler := func(ctx context.Context, req interface{}) (interface{}, error) {
		fmt.Println("  [handler] executing")
		return "ok", nil
	}

	fmt.Println("Executing through chained interceptors:")
	_, err := combined(context.Background(), "request",
		&grpc.UnaryServerInfo{FullMethod: "/svc.Method"},
		handler,
	)
	fmt.Printf("result err=%v\n", err)
}

func recoveryInterceptor(ctx context.Context, req interface{},
	info *grpc.UnaryServerInfo, handler grpc.UnaryHandler,
) (resp interface{}, err error) {
	defer func() {
		if r := recover(); r != nil {
			log.Printf("[recovery] panic: %v\n%s", r, debug.Stack())
			err = status.Error(codes.Internal, "panic recovered")
		}
	}()
	fmt.Println("  [recovery] before")
	resp, err = handler(ctx, req)
	fmt.Println("  [recovery] after")
	return
}

func loggingInterceptor(ctx context.Context, req interface{},
	info *grpc.UnaryServerInfo, handler grpc.UnaryHandler,
) (interface{}, error) {
	fmt.Println("  [logging] before")
	start := time.Now()
	resp, err := handler(ctx, req)
	fmt.Printf("  [logging] after (cost=%v)\n", time.Since(start))
	return resp, err
}

func authInterceptor(ctx context.Context, req interface{},
	info *grpc.UnaryServerInfo, handler grpc.UnaryHandler,
) (interface{}, error) {
	fmt.Println("  [auth] before")
	resp, err := handler(ctx, req)
	fmt.Println("  [auth] after")
	return resp, err
}

七、常用拦截器实战

1. 日志拦截器:记录请求/响应

go
package main

import (
	"context"
	"encoding/json"
	"fmt"
	"log"
	"time"

	"google.golang.org/grpc"
	"google.golang.org/grpc/metadata"
)

// 结构化日志拦截器:记录方法、耗时、请求摘要、响应摘要
type RPCLog struct {
	Method    string        `json:"method"`
	Duration  string        `json:"duration"`
	Request   string        `json:"request"`
	Response  string        `json:"response"`
	Error     string        `json:"error,omitempty"`
	RequestID string        `json:"request_id"`
}

func StructuredLoggingInterceptor(ctx context.Context, req interface{},
	info *grpc.UnaryServerInfo, handler grpc.UnaryHandler,
) (interface{}, error) {
	start := time.Now()

	// 提取 request-id
	reqID := ""
	if md, ok := metadata.FromIncomingContext(ctx); ok {
		if vals := md.Get("x-request-id"); len(vals) > 0 {
			reqID = vals[0]
		}
	}

	resp, err := handler(ctx, req)

	// 序列化请求/响应摘要(注意控制大小,避免日志爆炸)
	reqSummary := summarize(req, 256)
	respSummary := summarize(resp, 256)
	errStr := ""
	if err != nil {
		errStr = err.Error()
	}

	rpcLog := RPCLog{
		Method:    info.FullMethod,
		Duration:  time.Since(start).String(),
		Request:   reqSummary,
		Response:  respSummary,
		Error:     errStr,
		RequestID: reqID,
	}
	if b, e := json.Marshal(rpcLog); e == nil {
		log.Println(string(b))
	}

	return resp, err
}

func summarize(v interface{}, maxLen int) string {
	b, err := json.Marshal(v)
	if err != nil {
		return fmt.Sprintf("%T", v)
	}
	s := string(b)
	if len(s) > maxLen {
		return s[:maxLen] + "..."
	}
	return s
}

func main() {
	// 模拟日志拦截器效果
	handler := func(ctx context.Context, req interface{}) (interface{}, error) {
		return map[string]string{"status": "ok"}, nil
	}

	_, _ = StructuredLoggingInterceptor(
		metadata.NewIncomingContext(context.Background(),
			metadata.Pairs("x-request-id", "req-abc-123")),
		map[string]string{"user_id": "42"},
		&grpc.UnaryServerInfo{FullMethod: "/shop.v1.OrderService/GetOrder"},
		handler,
	)
}

2. 认证拦截器:JWT 验证

go
package main

import (
	"context"
	"fmt"
	"log"
	"strings"
	"time"

	"google.golang.org/grpc"
	"google.golang.org/grpc/codes"
	"google.golang.org/grpc/metadata"
	"google.golang.org/grpc/status"
)

// 模拟 JWT 解析(实际项目用 github.com/golang-jwt/jwt/v5)
type Claims struct {
	UserID string
	Role   string
	Exp    int64
}

func parseJWT(token string) (*Claims, error) {
	// 简化演示:真实项目需验证签名、过期时间
	if token == "" {
		return nil, fmt.Errorf("empty token")
	}
	// 模拟解析
	if strings.HasPrefix(token, "Bearer invalid") {
		return nil, fmt.Errorf("invalid signature")
	}
	return &Claims{UserID: "user-42", Role: "admin", Exp: time.Now().Add(time.Hour).Unix()}, nil
}

type claimsKey struct{}

// 不需要认证的方法白名单
var authWhitelist = map[string]bool{
	"/auth.AuthService/Login":    true,
	"/auth.AuthService/Register": true,
	"/grpc.health.v1.Health/Check": true,
}

func JWTAuthInterceptor(ctx context.Context, req interface{},
	info *grpc.UnaryServerInfo, handler grpc.UnaryHandler,
) (interface{}, error) {
	// 白名单跳过
	if authWhitelist[info.FullMethod] {
		return handler(ctx, req)
	}

	// 从 metadata 提取 authorization 头
	md, ok := metadata.FromIncomingContext(ctx)
	if !ok {
		return nil, status.Error(codes.Unauthenticated, "no metadata")
	}
	values := md.Get("authorization")
	if len(values) == 0 {
		return nil, status.Error(codes.Unauthenticated, "no authorization header")
	}

	token := values[0]
	token = strings.TrimPrefix(token, "Bearer ")

	claims, err := parseJWT(token)
	if err != nil {
		return nil, status.Errorf(codes.Unauthenticated, "invalid token: %v", err)
	}

	// 检查过期
	if claims.Exp < time.Now().Unix() {
		return nil, status.Error(codes.Unauthenticated, "token expired")
	}

	// 将 claims 存入 context 供后续使用
	ctx = context.WithValue(ctx, claimsKey{}, claims)
	return handler(ctx, req)
}

// 从 context 获取 claims 的辅助函数
func ClaimsFromContext(ctx context.Context) (*Claims, bool) {
	c, ok := ctx.Value(claimsKey{}).(*Claims)
	return c, ok
}

func main() {
	server := grpc.NewServer(grpc.ChainUnaryInterceptor(JWTAuthInterceptor))
	fmt.Printf("server with JWT auth: %T\n", server)

	// 演示白名单和认证流程
	ctx := metadata.NewIncomingContext(context.Background(),
		metadata.Pairs("authorization", "Bearer valid-token"))
	claims, err := parseJWT("valid-token")
	fmt.Printf("claims: %+v, err: %v\n", claims, err)
	_ = ctx
	_ = log.Printf
}

3. 恢复拦截器:panic recovery

go
package main

import (
	"context"
	"fmt"
	"log"
	"runtime/debug"

	"google.golang.org/grpc"
	"google.golang.org/grpc/codes"
	"google.golang.org/grpc/status"
)

func RecoveryInterceptor(ctx context.Context, req interface{},
	info *grpc.UnaryServerInfo, handler grpc.UnaryHandler,
) (resp interface{}, err error) {
	defer func() {
		if r := recover(); r != nil {
			// 记录完整堆栈
			log.Printf("[PANIC] method=%s recover=%v\n%s",
				info.FullMethod, r, debug.Stack())
			// 对外返回 Internal 错误,不泄露堆栈
			err = status.Error(codes.Internal, "internal server error")
		}
	}()
	return handler(ctx, req)
}

func main() {
	// 模拟一个会 panic 的 handler
	handler := func(ctx context.Context, req interface{}) (interface{}, error) {
		// 模拟空指针 panic
		var m map[string]string
		m["key"] = "value" // panic: assignment to entry in nil map
		return m, nil
	}

	// 经过 recovery 拦截器
	resp, err := RecoveryInterceptor(
		context.Background(),
		"request",
		&grpc.UnaryServerInfo{FullMethod: "/svc.Method"},
		handler,
	)
	fmt.Printf("resp=%v, err=%v\n", resp, err)
	// 进程不会崩溃,err 是 Internal 错误
}

4. 请求 ID 拦截器

go
package main

import (
	"context"
	"fmt"
	"log"
	"math/rand"
	"time"

	"google.golang.org/grpc"
	"google.golang.org/grpc/metadata"
)

type requestIDKey struct{}

func RequestIDInterceptor(ctx context.Context, req interface{},
	info *grpc.UnaryServerInfo, handler grpc.UnaryHandler,
) (interface{}, error) {
	// 优先复用客户端传入的 request-id
	reqID := ""
	if md, ok := metadata.FromIncomingContext(ctx); ok {
		if vals := md.Get("x-request-id"); len(vals) > 0 {
			reqID = vals[0]
		}
	}
	// 没有则生成新的
	if reqID == "" {
		reqID = generateID()
	}

	// 存入 context 供业务代码使用
	ctx = context.WithValue(ctx, requestIDKey{}, reqID)

	// 注入到 outgoing metadata(如果服务端再调下游,会自动传播)
	ctx = metadata.AppendToOutgoingContext(ctx, "x-request-id", reqID)

	log.Printf("[reqid] %s request-id=%s", info.FullMethod, reqID)

	return handler(ctx, req)
}

func generateID() string {
	r := rand.New(rand.NewSource(time.Now().UnixNano()))
	return fmt.Sprintf("req-%d-%06d", time.Now().Unix(), r.Intn(1000000))
}

func RequestIDFromContext(ctx context.Context) string {
	if v, ok := ctx.Value(requestIDKey{}).(string); ok {
		return v
	}
	return ""
}

func main() {
	server := grpc.NewServer(grpc.ChainUnaryInterceptor(RequestIDInterceptor))
	fmt.Printf("server with request-id: %T\n", server)

	// 演示
	ctx := context.Background()
	handler := func(ctx context.Context, req interface{}) (interface{}, error) {
		fmt.Printf("handler got request-id: %s\n", RequestIDFromContext(ctx))
		return "ok", nil
	}
	RequestIDInterceptor(ctx, "req", &grpc.UnaryServerInfo{FullMethod: "/svc.Test"}, handler)
}

5. 限流拦截器

go
package main

import (
	"context"
	"fmt"
	"log"
	"sync"
	"time"

	"google.golang.org/grpc"
	"google.golang.org/grpc/codes"
	"google.golang.org/grpc/status"
)

// === 令牌桶限流器 ===
type tokenBucket struct {
	mu       sync.Mutex
	tokens   int
	maxToken int
	rate     int // tokens per second
	lastTime time.Time
}

func newTokenBucket(maxToken, rate int) *tokenBucket {
	return &tokenBucket{
		tokens:   maxToken,
		maxToken: maxToken,
		rate:     rate,
		lastTime: time.Now(),
	}
}

func (tb *tokenBucket) allow() bool {
	tb.mu.Lock()
	defer tb.mu.Unlock()

	now := time.Now()
	elapsed := now.Sub(tb.lastTime).Seconds()
	tb.tokens += int(elapsed * float64(tb.rate))
	if tb.tokens > tb.maxToken {
		tb.tokens = tb.maxToken
	}
	tb.lastTime = now

	if tb.tokens > 0 {
		tb.tokens--
		return true
	}
	return false
}

// === 限流拦截器 ===
var globalLimiter = newTokenBucket(100, 50) // 容量 100,每秒补充 50

func RateLimitInterceptor(ctx context.Context, req interface{},
	info *grpc.UnaryServerInfo, handler grpc.UnaryHandler,
) (interface{}, error) {
	if !globalLimiter.allow() {
		log.Printf("[ratelimit] %s rejected (too many requests)", info.FullMethod)
		return nil, status.Error(codes.ResourceExhausted, "rate limit exceeded")
	}
	return handler(ctx, req)
}

func main() {
	server := grpc.NewServer(grpc.ChainUnaryInterceptor(RateLimitInterceptor))
	fmt.Printf("server with rate limit: %T\n", server)

	// 演示限流效果
	tb := newTokenBucket(5, 2) // 容量 5,每秒补充 2
	for i := 0; i < 10; i++ {
		if tb.allow() {
			fmt.Printf("request %d: allowed\n", i)
		} else {
			fmt.Printf("request %d: rejected\n", i)
		}
	}
}

6. 指标采集拦截器

go
package main

import (
	"context"
	"fmt"
	"log"
	"sync"
	"time"

	"google.golang.org/grpc"
	"google.golang.org/grpc/codes"
	"google.golang.org/grpc/status"
)

// 简化的指标收集器(实际项目用 Prometheus client)
type MetricsCollector struct {
	mu       sync.Mutex
	counters map[string]int64
	latencies map[string][]time.Duration
	errors   map[string]int64
}

func NewMetricsCollector() *MetricsCollector {
	return &MetricsCollector{
		counters:  make(map[string]int64),
		latencies: make(map[string][]time.Duration),
		errors:    make(map[string]int64),
	}
}

func (m *MetricsCollector) Record(method string, duration time.Duration, err error) {
	m.mu.Lock()
	defer m.mu.Unlock()
	m.counters[method]++
	m.latencies[method] = append(m.latencies[method], duration)
	if err != nil {
		m.errors[method]++
	}
}

func (m *MetricsCollector) Print() {
	m.mu.Lock()
	defer m.mu.Unlock()
	for method, count := range m.counters {
		var total time.Duration
		for _, d := range m.latencies[method] {
			total += d
		}
		avg := total / time.Duration(len(m.latencies[method]))
		fmt.Printf("[metrics] %s: count=%d avg=%v errors=%d\n",
			method, count, avg, m.errors[method])
	}
}

var metrics = NewMetricsCollector()

func MetricsInterceptor(ctx context.Context, req interface{},
	info *grpc.UnaryServerInfo, handler grpc.UnaryHandler,
) (interface{}, error) {
	start := time.Now()
	resp, err := handler(ctx, req)
	metrics.Record(info.FullMethod, time.Since(start), err)

	// 按状态码分类
	if err != nil {
		if st, ok := status.FromError(err); ok {
			switch st.Code() {
			case codes.Unavailable:
				log.Printf("[metrics] %s unavailable (service health issue)", info.FullMethod)
			case codes.DeadlineExceeded:
				log.Printf("[metrics] %s timeout (slow handler)", info.FullMethod)
			case codes.ResourceExhausted:
				log.Printf("[metrics] %s rate limited", info.FullMethod)
			}
		}
	}
	return resp, err
}

func main() {
	server := grpc.NewServer(grpc.ChainUnaryInterceptor(MetricsInterceptor))
	fmt.Printf("server with metrics: %T\n", server)

	// 模拟多次调用
	handler := func(ctx context.Context, req interface{}) (interface{}, error) {
		time.Sleep(time.Millisecond * 10)
		return "ok", nil
	}

	for i := 0; i < 5; i++ {
		MetricsInterceptor(context.Background(), "req",
			&grpc.UnaryServerInfo{FullMethod: "/svc.GetOrder"}, handler)
	}
	metrics.Print()
}

八、拦截器中的 metadata 传递

metadata 是 gRPC 的「HTTP 头」,用于在拦截器之间、客户端与服务端之间传递键值对信息。

go
package main

import (
	"context"
	"fmt"
	"log"

	"google.golang.org/grpc"
	"google.golang.org/grpc/metadata"
)

// === 客户端:注入 metadata ===
func clientInjectMetadata() {
	ctx := metadata.AppendToOutgoingContext(context.Background(),
		"x-request-id", "req-123",
		"x-trace-id", "trace-abc",
		"authorization", "Bearer my-token",
	)

	// 提取并打印(模拟发送)
	md, _ := metadata.FromOutgoingContext(ctx)
	fmt.Printf("client outgoing metadata: %v\n", md)

	// 调用时 gRPC 自动将 outgoing metadata 放入 HTTP/2 头部
	// resp, err := client.GetOrder(ctx, req)
	_ = ctx
}

// === 服务端:提取 metadata ===
func serverExtractMetadata() {
	// 模拟服务端收到的 incoming metadata
	ctx := metadata.NewIncomingContext(context.Background(),
		metadata.Pairs(
			"x-request-id", "req-123",
			"x-trace-id", "trace-abc",
			"authorization", "Bearer my-token",
		))

	md, ok := metadata.FromIncomingContext(ctx)
	if !ok {
		log.Println("no metadata")
		return
	}

	// 单值获取
	if vals := md.Get("x-request-id"); len(vals) > 0 {
		fmt.Printf("request-id: %s\n", vals[0])
	}

	// 多值获取(同一 key 可以有多个值)
	if vals := md.Get("authorization"); len(vals) > 0 {
		fmt.Printf("auth: %s\n", vals[0])
	}

	// 遍历所有 metadata
	for k, v := range md {
		fmt.Printf("  md[%s] = %v\n", k, v)
	}
}

// === metadata 在拦截器链中的传播 ===
func serverInterceptorWithMetadata(ctx context.Context, req interface{},
	info *grpc.UnaryServerInfo, handler grpc.UnaryHandler,
) (interface{}, error) {
	// 提取 incoming metadata
	md, _ := metadata.FromIncomingContext(ctx)
	reqID := ""
	if vals := md.Get("x-request-id"); len(vals) > 0 {
		reqID = vals[0]
	}

	// 将 request-id 存入 context
	ctx = context.WithValue(ctx, "requestID", reqID)

	// 如果服务端需要调用下游,将 request-id 注入 outgoing metadata
	ctx = metadata.AppendToOutgoingContext(ctx, "x-request-id", reqID)

	return handler(ctx, req)
}

func main() {
	fmt.Println("=== client inject ===")
	clientInjectMetadata()

	fmt.Println("\n=== server extract ===")
	serverExtractMetadata()

	_ = grpc.NewServer
}

metadata 注意事项

  • metadata 的 key 自动转小写,X-Request-Idx-request-id 是同一个 key。
  • 二进制 key(以 -bin 结尾)用于传递二进制数据,如 grpc-trace-bin
  • metadata 大小受 HTTP/2 头部限制(默认 8KB),不要放大数据。

九、完整示例:生产级拦截器组合

go
package main

import (
	"context"
	"fmt"
	"log"
	"math/rand"
	"runtime/debug"
	"strings"
	"sync"
	"time"

	"google.golang.org/grpc"
	"google.golang.org/grpc/codes"
	"google.golang.org/grpc/metadata"
	"google.golang.org/grpc/status"
)

// === 1. Recovery 拦截器 ===
func RecoveryInterceptor(ctx context.Context, req interface{},
	info *grpc.UnaryServerInfo, handler grpc.UnaryHandler,
) (resp interface{}, err error) {
	defer func() {
		if r := recover(); r != nil {
			log.Printf("[PANIC] %s: %v\n%s", info.FullMethod, r, debug.Stack())
			err = status.Error(codes.Internal, "internal error")
		}
	}()
	return handler(ctx, req)
}

// === 2. Request-ID 拦截器 ===
type ctxKey string

const reqIDKey ctxKey = "requestID"

func RequestIDInterceptor(ctx context.Context, req interface{},
	info *grpc.UnaryServerInfo, handler grpc.UnaryHandler,
) (interface{}, error) {
	reqID := ""
	if md, ok := metadata.FromIncomingContext(ctx); ok {
		if v := md.Get("x-request-id"); len(v) > 0 {
			reqID = v[0]
		}
	}
	if reqID == "" {
		reqID = fmt.Sprintf("req-%d%06d", time.Now().Unix(), rand.Intn(1000000))
	}
	ctx = context.WithValue(ctx, reqIDKey, reqID)
	ctx = metadata.AppendToOutgoingContext(ctx, "x-request-id", reqID)
	return handler(ctx, req)
}

// === 3. 日志拦截器 ===
func LoggingInterceptor(ctx context.Context, req interface{},
	info *grpc.UnaryServerInfo, handler grpc.UnaryHandler,
) (interface{}, error) {
	start := time.Now()
	reqID, _ := ctx.Value(reqIDKey).(string)
	resp, err := handler(ctx, req)
	log.Printf("[LOG] reqID=%s method=%s cost=%v err=%v",
		reqID, info.FullMethod, time.Since(start), err)
	return resp, err
}

// === 4. 认证拦截器 ===
var authWhitelist = map[string]bool{
	"/auth.AuthService/Login": true,
}

func AuthInterceptor(ctx context.Context, req interface{},
	info *grpc.UnaryServerInfo, handler grpc.UnaryHandler,
) (interface{}, error) {
	if authWhitelist[info.FullMethod] {
		return handler(ctx, req)
	}
	md, ok := metadata.FromIncomingContext(ctx)
	if !ok {
		return nil, status.Error(codes.Unauthenticated, "no metadata")
	}
	tokens := md.Get("authorization")
	if len(tokens) == 0 || !strings.HasPrefix(tokens[0], "Bearer ") {
		return nil, status.Error(codes.Unauthenticated, "no valid token")
	}
	// 模拟 token 验证
	token := strings.TrimPrefix(tokens[0], "Bearer ")
	if token == "" {
		return nil, status.Error(codes.Unauthenticated, "invalid token")
	}
	ctx = context.WithValue(ctx, ctxKey("userID"), "user-42")
	return handler(ctx, req)
}

// === 5. 限流拦截器 ===
type rateLimiter struct {
	mu        sync.Mutex
	tokens    int
	maxTokens int
	rate      float64
	last      time.Time
}

func newRateLimiter(max, rate int) *rateLimiter {
	return &rateLimiter{tokens: max, maxTokens: max, rate: float64(rate), last: time.Now()}
}

func (rl *rateLimiter) allow() bool {
	rl.mu.Lock()
	defer rl.mu.Unlock()
	now := time.Now()
	rl.tokens += int(now.Sub(rl.last).Seconds() * rl.rate)
	if rl.tokens > rl.maxTokens {
		rl.tokens = rl.maxTokens
	}
	rl.last = now
	if rl.tokens > 0 {
		rl.tokens--
		return true
	}
	return false
}

var limiter = newRateLimiter(100, 50)

func RateLimitInterceptor(ctx context.Context, req interface{},
	info *grpc.UnaryServerInfo, handler grpc.UnaryHandler,
) (interface{}, error) {
	if !limiter.allow() {
		return nil, status.Error(codes.ResourceExhausted, "rate limit exceeded")
	}
	return handler(ctx, req)
}

// === 6. 指标拦截器 ===
type metricsStore struct {
	mu     sync.Mutex
	stats  map[string]*methodStat
}

type methodStat struct {
	count   int64
	errors  int64
	totalNs int64
}

var store = &metricsStore{stats: make(map[string]*methodStat)}

func MetricsInterceptor(ctx context.Context, req interface{},
	info *grpc.UnaryServerInfo, handler grpc.UnaryHandler,
) (interface{}, error) {
	start := time.Now()
	resp, err := handler(ctx, req)
	duration := time.Since(start)

	store.mu.Lock()
	stat, ok := store.stats[info.FullMethod]
	if !ok {
		stat = &methodStat{}
		store.stats[info.FullMethod] = stat
	}
	stat.count++
	stat.totalNs += duration.Nanoseconds()
	if err != nil {
		stat.errors++
	}
	store.mu.Unlock()

	return resp, err
}

func (s *metricsStore) Print() {
	s.mu.Lock()
	defer s.mu.Unlock()
	for method, stat := range s.stats {
		avg := time.Duration(stat.totalNs/stat.count) * time.Nanosecond
		fmt.Printf("[METRICS] %s: count=%d errors=%d avg=%v\n",
			method, stat.count, stat.errors, avg)
	}
}

// === 组装生产级拦截器链 ===
func newProductionServer() *grpc.Server {
	return grpc.NewServer(grpc.ChainUnaryInterceptor(
		RecoveryInterceptor,      // 1. 最外层:兜底 panic
		RequestIDInterceptor,     // 2. 注入 request-id
		LoggingInterceptor,       // 3. 记录日志
		MetricsInterceptor,       // 4. 采集指标
		RateLimitInterceptor,     // 5. 限流
		AuthInterceptor,          // 6. 最内层:认证(最靠近业务)
	))
}

func main() {
	server := newProductionServer()
	fmt.Printf("production server: %T\n", server)

	// 模拟一次完整调用链
	ctx := metadata.NewIncomingContext(context.Background(),
		metadata.Pairs("x-request-id", "req-demo-001", "authorization", "Bearer valid-jwt"))

	handler := func(ctx context.Context, req interface{}) (interface{}, error) {
		reqID, _ := ctx.Value(reqIDKey).(string)
		fmt.Printf("  [handler] processing, reqID=%s\n", reqID)
		return "order-12345", nil
	}

	// 手动串联拦截器链(模拟 grpc.ChainUnaryInterceptor 的效果)
	chain := []grpc.UnaryServerInterceptor{
		RecoveryInterceptor, RequestIDInterceptor, LoggingInterceptor,
		MetricsInterceptor, RateLimitInterceptor, AuthInterceptor,
	}
	finalHandler := handler
	for i := len(chain) - 1; i >= 0; i-- {
		interceptor := chain[i]
		next := finalHandler
		finalHandler = func(ctx context.Context, req interface{}) (interface{}, error) {
			return interceptor(ctx, req, &grpc.UnaryServerInfo{FullMethod: "/shop.v1.OrderService/GetOrder"}, next)
		}
	}

	resp, err := finalHandler(ctx, "get-order-req")
	fmt.Printf("\nresult: resp=%v err=%v\n", resp, err)
	store.Print()
}

这个示例展示了六个拦截器如何按洋葱模型协作:recovery 兜底、request-id 贯穿全链路、日志和指标并行采集、限流和认证守卫业务入口。实际生产中只需将 newProductionServer() 的拦截器列表替换为你的组合即可。

十、小结

本篇深入 gRPC 拦截器与中间件模式:

  • 机制原理:四种拦截器(服务端/客户端 × Unary/Stream),通过包装 handler 或 stream 实现横切逻辑。
  • 服务端拦截器:Unary 通过包装 handler,Stream 通过包装 ServerStream 拦截每次 SendMsg/RecvMsg
  • 客户端拦截器:通过 invoker 调用真正的 RPC,可注入 metadata、记录耗时。
  • grpc-middlewareChainUnaryServer 递归构建洋葱链,顺序为注册顺序从外到内。
  • 六大实战拦截器
    • 日志:结构化记录方法、耗时、请求/响应摘要。
    • 认证:JWT 验证 + 白名单 + claims 注入 context。
    • 恢复:recover panic + 堆栈日志 + Internal 错误返回。
    • 请求 ID:复用或生成 + context + outgoing metadata 传播。
    • 限流:令牌桶算法 + ResourceExhausted 状态码。
    • 指标:按方法统计 QPS、延迟、错误率。
  • metadata 传递:incoming/outgoing metadata 的提取与注入,key 自动小写,二进制 key 以 -bin 结尾。
  • 生产级组合:recovery(外)-> request-id -> logging -> metrics -> ratelimit -> auth(内)的推荐顺序。

下一篇将深入 gRPC 的错误处理与状态码体系,讲解 rich error 模式和 errdetails 的使用。