From afb4f9125810084559c0c5252e6b662ad03d3585 Mon Sep 17 00:00:00 2001 From: jonny rhea Date: Fri, 2 Jan 2026 15:35:19 -0600 Subject: [PATCH] - ensure span.End() is called on error - ensure runMethod() doesn't record spans if parentSpan isn't recording - ensure unsubscribe() isn't recorded - add tests to verify that subscribe/unsubscribe don't record --- rpc/handler.go | 69 +++++++++++++++++++++++++-------------------- rpc/tracing_test.go | 31 +++++++++++++++++--- 2 files changed, 66 insertions(+), 34 deletions(-) diff --git a/rpc/handler.go b/rpc/handler.go index 004dae44a1..5ee6bf34d1 100644 --- a/rpc/handler.go +++ b/rpc/handler.go @@ -503,26 +503,30 @@ func (h *handler) handleCallMsg(ctx *callProc, msg *jsonrpcMessage) *jsonrpcMess // handleCall processes method calls. func (h *handler) handleCall(cp *callProc, msg *jsonrpcMessage) *jsonrpcMessage { - // Start root span for the request. - ctx, rootSpan := h.startSpan(cp.ctx, msg, "rpc.handleCall", cp.isBatch) - defer rootSpan.End() if msg.isSubscribe() { return h.handleSubscribe(cp, msg) } - var callb *callback if msg.isUnsubscribe() { - callb = h.unsubscribeCb - } else { - // Check method name length - if len(msg.Method) > maxMethodNameLength { - return msg.errorResponse(&invalidRequestError{fmt.Sprintf("method name too long: %d > %d", len(msg.Method), maxMethodNameLength)}) + args, err := parsePositionalArguments(msg.Params, h.unsubscribeCb.argTypes) + if err != nil { + return msg.errorResponse(&invalidParamsError{err.Error()}) } - callb = h.reg.callback(msg.Method) + return h.runMethod(cp.ctx, msg, h.unsubscribeCb, args, cp.isBatch) } + + // Check method name length + if len(msg.Method) > maxMethodNameLength { + return msg.errorResponse(&invalidRequestError{fmt.Sprintf("method name too long: %d > %d", len(msg.Method), maxMethodNameLength)}) + } + callb := h.reg.callback(msg.Method) if callb == nil { return msg.errorResponse(&methodNotFoundError{method: msg.Method}) } + // Start root span for the request. + ctx, rootSpan := h.startSpan(cp.ctx, msg, "rpc.handleCall", cp.isBatch) + defer rootSpan.End() + // Start tracing span before parsing arguments. _, pspan := h.startSpan(ctx, msg, "rpc.parsePositionalArguments", cp.isBatch) args, err := parsePositionalArguments(msg.Params, callb.argTypes) @@ -530,6 +534,7 @@ func (h *handler) handleCall(cp *callProc, msg *jsonrpcMessage) *jsonrpcMessage pspan.RecordError(err) pspan.SetStatus(codes.Error, err.Error()) rootSpan.SetStatus(codes.Error, err.Error()) + pspan.End() return msg.errorResponse(&invalidParamsError{err.Error()}) } pspan.End() @@ -545,17 +550,14 @@ func (h *handler) handleCall(cp *callProc, msg *jsonrpcMessage) *jsonrpcMessage rspan.End() // Collect the statistics for RPC calls if metrics is enabled. - // We only care about pure rpc call. Filter out subscription. - if callb != h.unsubscribeCb { - rpcRequestGauge.Inc(1) - if answer.Error != nil { - failedRequestGauge.Inc(1) - } else { - successfulRequestGauge.Inc(1) - } - rpcServingTimer.UpdateSince(start) - updateServeTimeHistogram(msg.Method, answer.Error == nil, time.Since(start)) + rpcRequestGauge.Inc(1) + if answer.Error != nil { + failedRequestGauge.Inc(1) + } else { + successfulRequestGauge.Inc(1) } + rpcServingTimer.UpdateSince(start) + updateServeTimeHistogram(msg.Method, answer.Error == nil, time.Since(start)) return answer } @@ -632,22 +634,29 @@ func (h *handler) tracer() trace.Tracer { // runMethod runs the Go callback for an RPC method. func (h *handler) runMethod(ctx context.Context, msg *jsonrpcMessage, callb *callback, args []reflect.Value, isBatch bool) *jsonrpcMessage { - result, err := callb.call(ctx, msg.Method, args) parentSpan := trace.SpanFromContext(ctx) + result, err := callb.call(ctx, msg.Method, args) if err != nil { parentSpan.SetStatus(codes.Error, err.Error()) return msg.errorResponse(err) } - _, span := h.startSpan(ctx, msg, "rpc.response", isBatch) - response := msg.response(result) - if response.Error != nil { - err := errors.New(response.Error.Message) - span.RecordError(err) - span.SetStatus(codes.Error, err.Error()) - parentSpan.SetStatus(codes.Error, err.Error()) + + // If parent span is recording, start a span for the response. + // Note: This prevents msg.response spans from being created when + // the parent span is not recording (e.g. subscription tracing disabled). + if parentSpan.IsRecording() { + _, span := h.startSpan(ctx, msg, "rpc.msg.response", isBatch) + defer span.End() + response := msg.response(result) + if response.Error != nil { + err := errors.New(response.Error.Message) + span.RecordError(errors.New(response.Error.Message)) + span.SetStatus(codes.Error, err.Error()) + parentSpan.SetStatus(codes.Error, err.Error()) + } + return response } - span.End() - return response + return msg.response(result) } // unsubscribe is the callback function for all *_unsubscribe calls. diff --git a/rpc/tracing_test.go b/rpc/tracing_test.go index ad7ad8c115..1f72c73bb1 100644 --- a/rpc/tracing_test.go +++ b/rpc/tracing_test.go @@ -63,10 +63,8 @@ func newTracingServer(t *testing.T) (*Server, *sdktrace.TracerProvider, *tracete func TestTracingHTTP(t *testing.T) { t.Parallel() server, tracer, exporter := newTracingServer(t) - httpsrv := httptest.NewServer(server) t.Cleanup(httpsrv.Close) - client, err := DialHTTP(httpsrv.URL) if err != nil { t.Fatalf("failed to dial: %v", err) @@ -114,10 +112,8 @@ func TestTracingHTTP(t *testing.T) { func TestTracingBatchHTTP(t *testing.T) { t.Parallel() server, tracer, exporter := newTracingServer(t) - httpsrv := httptest.NewServer(server) t.Cleanup(httpsrv.Close) - client, err := DialHTTP(httpsrv.URL) if err != nil { t.Fatalf("failed to dial: %v", err) @@ -164,3 +160,30 @@ func TestTracingBatchHTTP(t *testing.T) { t.Fatalf("expected %d matching batch spans, got %d", len(batch), found) } } + +// TestTracingSubscribeUnsubscribe verifies that subscribe and unsubscribe calls +// do not emit any spans. +func TestTracingSubscribeUnsubscribe(t *testing.T) { + t.Parallel() + server, tracer, exporter := newTracingServer(t) + client := DialInProc(server) + t.Cleanup(client.Close) + + // Subscribe to notifications. + sub, err := client.Subscribe(context.Background(), "nftest", make(chan int), "someSubscription", 1, 1) + if err != nil { + t.Fatalf("subscribe failed: %v", err) + } + + // Unsubscribe. + sub.Unsubscribe() + + // Flush and check that no spans were emitted. + if err := tracer.ForceFlush(context.Background()); err != nil { + t.Fatalf("failed to flush: %v", err) + } + spans := exporter.GetSpans() + if len(spans) != 0 { + t.Errorf("expected no spans for subscribe/unsubscribe, got %d", len(spans)) + } +}