Appearance
拦截器与中间件模式
拦截器(Interceptor)是 gRPC 的横切关注点(cross-cutting concern)机制,类似 HTTP 框架中的中间件。本篇不再重复拦截器的基础语法,而是深入拦截器的实现原理、四种拦截器的完整实现、grpc-middleware 的链式模式,以及生产环境最常用的六类拦截器(日志、认证、恢复、请求 ID、限流、指标采集)的实战代码。
一、拦截器机制原理
gRPC 的拦截器本质上是一个函数包装器:在真正的 handler 前后插入自定义逻辑。框架在调度 RPC 时,会先经过拦截器链,再到达业务 handler。
客户端调用
│
▼
客户端拦截器链 (UnaryInterceptor / StreamInterceptor)
│
▼ (HTTP/2 传输)
│
服务端拦截器链 (UnaryInterceptor / StreamInterceptor)
│
▼
业务 handlergRPC 有四种拦截器,对应服务端/客户端 × 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-Id和x-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-middleware:
ChainUnaryServer递归构建洋葱链,顺序为注册顺序从外到内。 - 六大实战拦截器:
- 日志:结构化记录方法、耗时、请求/响应摘要。
- 认证: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 的使用。