otel: Segregate client and server RPCInfo used for metrics and traces (#9081)
Fixes #9053
Client and Server metrics and traces handling code during RPC was using
the same RPCInfo. This lead to overwriting of context when an
application was acting as server and then initiating other grpc call as
a client. This PR segregates the rpcInfoKey context key into
clientRPCInfoKey and serverRPCInfoKey so that it is not
reused/overwritten between server and client code.
RELEASE NOTES:
* otel: Segregate client and server RPCInfo used for metrics and traces.
diff --git a/stats/opentelemetry/client_metrics.go b/stats/opentelemetry/client_metrics.go
index 420482b..91bda9a 100644
--- a/stats/opentelemetry/client_metrics.go
+++ b/stats/opentelemetry/client_metrics.go
@@ -76,7 +76,7 @@
target: cc.CanonicalTarget(),
method: determineMethod(method, opts...),
}
- ctx = setCallInfo(ctx, ci)
+ ctx = context.WithValue(ctx, callInfoKey{}, ci)
}
return ctx, ci
}
@@ -157,17 +157,6 @@
// HandleConn exists to satisfy stats.Handler.
func (h *clientMetricsHandler) HandleConn(context.Context, stats.ConnStats) {}
-// getOrCreateRPCAttemptInfo retrieves or creates an rpc attemptInfo object
-// and ensures it is set in the context along with the rpcInfo.
-func getOrCreateRPCAttemptInfo(ctx context.Context) (context.Context, *attemptInfo) {
- ri := getRPCInfo(ctx)
- if ri != nil {
- return ctx, ri.ai
- }
- ri = &rpcInfo{ai: &attemptInfo{}}
- return setRPCInfo(ctx, ri), ri.ai
-}
-
// TagRPC implements per RPC attempt context management for metrics.
func (h *clientMetricsHandler) TagRPC(ctx context.Context, info *stats.RPCTagInfo) context.Context {
// Numerous stats handlers can be used for the same channel. The cluster
@@ -187,17 +176,18 @@
}
ctx = istats.SetLabels(ctx, labels)
}
- ctx, ai := getOrCreateRPCAttemptInfo(ctx)
+ ctx, ri := getOrCreateClientRPCInfo(ctx)
+ ai := ri.ai
ai.startTime = time.Now()
ai.xdsLabels = labels.TelemetryLabels
ai.method = removeLeadingSlash(info.FullMethodName)
- return setRPCInfo(ctx, &rpcInfo{ai: ai})
+ return ctx
}
// HandleRPC handles per RPC stats implementation.
func (h *clientMetricsHandler) HandleRPC(ctx context.Context, rs stats.RPCStats) {
- ri := getRPCInfo(ctx)
+ ri := clientRPCInfo(ctx)
if ri == nil {
logger.Error("ctx passed into client side stats handler metrics event handling has no client attempt data present")
return
diff --git a/stats/opentelemetry/client_tracing.go b/stats/opentelemetry/client_tracing.go
index 868d6a2..718a634 100644
--- a/stats/opentelemetry/client_tracing.go
+++ b/stats/opentelemetry/client_tracing.go
@@ -120,14 +120,14 @@
// TagRPC implements per RPC attempt context management for traces.
func (h *clientTracingHandler) TagRPC(ctx context.Context, info *stats.RPCTagInfo) context.Context {
- ctx, ai := getOrCreateRPCAttemptInfo(ctx)
- ctx, ai = h.traceTagRPC(ctx, ai, info.NameResolutionDelay)
- return setRPCInfo(ctx, &rpcInfo{ai: ai})
+ ctx, ri := getOrCreateClientRPCInfo(ctx)
+ ctx, _ = h.traceTagRPC(ctx, ri.ai, info.NameResolutionDelay)
+ return ctx
}
// HandleRPC handles per RPC tracing implementation.
func (h *clientTracingHandler) HandleRPC(ctx context.Context, rs stats.RPCStats) {
- ri := getRPCInfo(ctx)
+ ri := clientRPCInfo(ctx)
if ri == nil {
logger.Error("ctx passed into client side tracing handler trace event handling has no client attempt data present")
return
diff --git a/stats/opentelemetry/e2e_test.go b/stats/opentelemetry/e2e_test.go
index 7c0ecf4..dda883f 100644
--- a/stats/opentelemetry/e2e_test.go
+++ b/stats/opentelemetry/e2e_test.go
@@ -2204,3 +2204,146 @@
t.Fatalf("Metric verification failed for case %s: %v", name, err)
}
}
+
+// TestRelayContextCollisionMetrics verifies that when an application acts as
+// both a server and a client using the same context, the client metrics do not
+// inherit or overwrite the server's telemetry metadata (e.g., grpc.method).
+func (s) TestRelayContextCollisionMetrics(t *testing.T) {
+ backendMetricsOpts, _ := defaultMetricsOptions(t, nil)
+ backendServer := setupStubServer(t, backendMetricsOpts, nil)
+ backendServer.EmptyCallF = func(_ context.Context, _ *testpb.Empty) (*testpb.Empty, error) {
+ return nil, status.Error(codes.Unimplemented, "EmptyCall not implemented")
+ }
+ defer backendServer.Stop()
+
+ relayMetricsOpts, relayMetricsReader := defaultMetricsOptions(t, nil)
+ otelOpts := opentelemetry.Options{MetricsOptions: *relayMetricsOpts}
+
+ relayServer := &stubserver.StubServer{
+ UnaryCallF: func(ctx context.Context, _ *testpb.SimpleRequest) (*testpb.SimpleResponse, error) {
+ relayCC, err := grpc.NewClient(
+ backendServer.Address,
+ grpc.WithTransportCredentials(insecure.NewCredentials()),
+ opentelemetry.DialOption(otelOpts),
+ )
+ if err != nil {
+ return nil, fmt.Errorf("failed to create relay client: %v", err)
+ }
+ defer relayCC.Close()
+ client := testpb.NewTestServiceClient(relayCC)
+ _, err = client.EmptyCall(ctx, &testpb.Empty{})
+ if status.Code(err) != codes.Unimplemented {
+ t.Errorf("Expected Unimplemented error, got: %v", err)
+ }
+ return &testpb.SimpleResponse{}, nil
+ },
+ }
+ if err := relayServer.Start([]grpc.ServerOption{opentelemetry.ServerOption(otelOpts)}, opentelemetry.DialOption(otelOpts)); err != nil {
+ t.Fatalf("Failed to start relay server: %v", err)
+ }
+ defer relayServer.Stop()
+
+ ctx, cancel := context.WithTimeout(context.Background(), defaultTestTimeout)
+ defer cancel()
+
+ if _, err := relayServer.Client.UnaryCall(ctx, &testpb.SimpleRequest{}); err != nil {
+ t.Fatalf("Unexpected UnaryCall error: %v", err)
+ }
+
+ // Verify Server Metric Identity is retained.
+ if err := checkMetricWithMethod(ctx, relayMetricsReader, "grpc.server.call.started", "grpc.testing.TestService/UnaryCall"); err != nil {
+ t.Fatal(err)
+ }
+
+ // Verify Client Metric Identity correctly resolved to "grpc.testing.TestService/EmptyCall".
+ if err := checkMetricWithMethod(ctx, relayMetricsReader, "grpc.client.attempt.started", "grpc.testing.TestService/EmptyCall"); err != nil {
+ t.Fatal(err)
+ }
+}
+
+// TestRelayContextCollisionTracing verifies that span context is correctly
+// propagated from incoming server requests to outgoing client requests without
+// the client span accidentally adopting the server's identity or breaking the
+// trace chain.
+func (s) TestRelayContextCollisionTracing(t *testing.T) {
+ backendTraceOpts, _ := defaultTraceOptions(t)
+ backendServer := setupStubServer(t, nil, backendTraceOpts)
+ backendServer.EmptyCallF = func(_ context.Context, _ *testpb.Empty) (*testpb.Empty, error) {
+ return nil, status.Error(codes.Unimplemented, "EmptyCall not implemented")
+ }
+ defer backendServer.Stop()
+
+ relayTraceOpts, relayTraceExporter := defaultTraceOptions(t)
+ otelOpts := opentelemetry.Options{TraceOptions: *relayTraceOpts}
+
+ relayServer := &stubserver.StubServer{
+ UnaryCallF: func(ctx context.Context, _ *testpb.SimpleRequest) (*testpb.SimpleResponse, error) {
+ relayCC, err := grpc.NewClient(
+ backendServer.Address,
+ grpc.WithTransportCredentials(insecure.NewCredentials()),
+ opentelemetry.DialOption(otelOpts),
+ )
+ if err != nil {
+ return nil, fmt.Errorf("failed to create relay client: %v", err)
+ }
+ defer relayCC.Close()
+ client := testpb.NewTestServiceClient(relayCC)
+ _, err = client.EmptyCall(ctx, &testpb.Empty{})
+ if status.Code(err) != codes.Unimplemented {
+ t.Errorf("Expected Unimplemented error, got: %v", err)
+ }
+ return &testpb.SimpleResponse{}, nil
+ },
+ }
+ if err := relayServer.Start([]grpc.ServerOption{opentelemetry.ServerOption(otelOpts)}, opentelemetry.DialOption(otelOpts)); err != nil {
+ t.Fatalf("Failed to start relay server: %v", err)
+ }
+ defer relayServer.Stop()
+
+ ctx, cancel := context.WithTimeout(context.Background(), defaultTestTimeout)
+ defer cancel()
+
+ _, _ = relayServer.Client.UnaryCall(ctx, &testpb.SimpleRequest{})
+
+ wantSpans := []traceSpanInfo{
+ {name: "Recv.", spanKind: "server"},
+ {name: "Sent.grpc.testing.TestService.EmptyCall", spanKind: "client"},
+ }
+ spans, err := waitForTraceSpans(ctx, relayTraceExporter, wantSpans)
+ if err != nil {
+ t.Fatalf("Failed to wait for spans: %v", err)
+ }
+
+ var srvTraceID, cliTraceID oteltrace.TraceID
+ for _, span := range spans {
+ if span.Name == "Recv." && span.SpanKind == oteltrace.SpanKindServer {
+ srvTraceID = span.SpanContext.TraceID()
+ }
+ if span.Name == "Sent.grpc.testing.TestService.EmptyCall" && span.SpanKind == oteltrace.SpanKindClient {
+ cliTraceID = span.SpanContext.TraceID()
+ }
+ }
+ if !srvTraceID.IsValid() || !cliTraceID.IsValid() {
+ t.Fatalf("Invalid trace IDs found. Server: %s, Client: %s", srvTraceID, cliTraceID)
+ }
+
+ if srvTraceID != cliTraceID {
+ t.Errorf("Trace continuity broken: Server TraceID %s != Client TraceID %s", srvTraceID, cliTraceID)
+ }
+}
+
+// checkMetricWithMethod verifies that a metric with the specified name contains
+// a data point matching the target grpc.method. It does not poll.
+func checkMetricWithMethod(ctx context.Context, reader *metric.ManualReader, metricName, method string) error {
+ metrics := metricsDataFromReader(ctx, reader)
+ if m, ok := metrics[metricName]; ok {
+ if sum, ok := m.Data.(metricdata.Sum[int64]); ok {
+ for _, dp := range sum.DataPoints {
+ if val, ok := dp.Attributes.Value("grpc.method"); ok && val.AsString() == method {
+ return nil
+ }
+ }
+ }
+ }
+ return fmt.Errorf("metric %q with method %q not found", metricName, method)
+}
diff --git a/stats/opentelemetry/opentelemetry.go b/stats/opentelemetry/opentelemetry.go
index 1031e9f..d976e60 100644
--- a/stats/opentelemetry/opentelemetry.go
+++ b/stats/opentelemetry/opentelemetry.go
@@ -183,10 +183,6 @@
type callInfoKey struct{}
-func setCallInfo(ctx context.Context, ci *callInfo) context.Context {
- return context.WithValue(ctx, callInfoKey{}, ci)
-}
-
// getCallInfo returns the callInfo stored in the context, or nil
// if there isn't one.
func getCallInfo(ctx context.Context) *callInfo {
@@ -200,19 +196,41 @@
ai *attemptInfo
}
-type rpcInfoKey struct{}
+type clientRPCInfoKey struct{}
+type serverRPCInfoKey struct{}
-func setRPCInfo(ctx context.Context, ri *rpcInfo) context.Context {
- return context.WithValue(ctx, rpcInfoKey{}, ri)
+// clientRPCInfo returns the rpcInfo stored in the context for client, or nil
+// if there isn't one.
+func clientRPCInfo(ctx context.Context) *rpcInfo {
+ ri, _ := ctx.Value(clientRPCInfoKey{}).(*rpcInfo)
+ return ri
}
-// getRPCInfo returns the rpcInfo stored in the context, or nil
+// serverRPCInfo returns the rpcInfo stored in the context for server, or nil
// if there isn't one.
-func getRPCInfo(ctx context.Context) *rpcInfo {
- ri, _ := ctx.Value(rpcInfoKey{}).(*rpcInfo)
+func serverRPCInfo(ctx context.Context) *rpcInfo {
+ ri, _ := ctx.Value(serverRPCInfoKey{}).(*rpcInfo)
return ri
}
+func getOrCreateClientRPCInfo(ctx context.Context) (context.Context, *rpcInfo) {
+ ri := clientRPCInfo(ctx)
+ if ri != nil {
+ return ctx, ri
+ }
+ ri = &rpcInfo{ai: &attemptInfo{}}
+ return context.WithValue(ctx, clientRPCInfoKey{}, ri), ri
+}
+
+func getOrCreateServerRPCInfo(ctx context.Context) (context.Context, *rpcInfo) {
+ ri := serverRPCInfo(ctx)
+ if ri != nil {
+ return ctx, ri
+ }
+ ri = &rpcInfo{ai: &attemptInfo{}}
+ return context.WithValue(ctx, serverRPCInfoKey{}, ri), ri
+}
+
func removeLeadingSlash(mn string) string {
return strings.TrimLeft(mn, "/")
}
diff --git a/stats/opentelemetry/server_metrics.go b/stats/opentelemetry/server_metrics.go
index 75d922e..4db7978 100644
--- a/stats/opentelemetry/server_metrics.go
+++ b/stats/opentelemetry/server_metrics.go
@@ -196,16 +196,17 @@
method = "other"
}
}
- ctx, ai := getOrCreateRPCAttemptInfo(ctx)
+ ctx, ri := getOrCreateServerRPCInfo(ctx)
+ ai := ri.ai
ai.startTime = time.Now()
ai.method = removeLeadingSlash(method)
- return setRPCInfo(ctx, &rpcInfo{ai: ai})
+ return ctx
}
// HandleRPC handles per RPC stats implementation.
func (h *serverMetricsHandler) HandleRPC(ctx context.Context, rs stats.RPCStats) {
- ri := getRPCInfo(ctx)
+ ri := serverRPCInfo(ctx)
if ri == nil {
logger.Error("ctx passed into server side stats handler metrics event handling has no server call data present")
return
diff --git a/stats/opentelemetry/server_tracing.go b/stats/opentelemetry/server_tracing.go
index 0e2181b..c267ba1 100644
--- a/stats/opentelemetry/server_tracing.go
+++ b/stats/opentelemetry/server_tracing.go
@@ -40,9 +40,9 @@
// TagRPC implements per RPC attempt context management for traces.
func (h *serverTracingHandler) TagRPC(ctx context.Context, _ *stats.RPCTagInfo) context.Context {
- ctx, ai := getOrCreateRPCAttemptInfo(ctx)
- ctx, ai = h.traceTagRPC(ctx, ai)
- return setRPCInfo(ctx, &rpcInfo{ai: ai})
+ ctx, ri := getOrCreateServerRPCInfo(ctx)
+ ctx, _ = h.traceTagRPC(ctx, ri.ai)
+ return ctx
}
// traceTagRPC populates context with new span data using the TextMapPropagator
@@ -67,7 +67,7 @@
// HandleRPC handles per RPC tracing implementation.
func (h *serverTracingHandler) HandleRPC(ctx context.Context, rs stats.RPCStats) {
- ri := getRPCInfo(ctx)
+ ri := serverRPCInfo(ctx)
if ri == nil {
logger.Error("ctx passed into server side tracing handler trace event handling has no server call data present")
return