GrpcServerInterceptor.java 7.17 KB
package com.tianting.infoloop.tracing;

import brave.Span;
import brave.Tracing;
import brave.baggage.BaggagePropagation;
import brave.propagation.Propagation;
import io.grpc.Context;
import io.grpc.ForwardingServerCall;
import io.grpc.Grpc;
import io.grpc.Metadata;
import io.grpc.ServerCall;
import io.grpc.ServerCallHandler;
import io.grpc.ServerInterceptor;
import io.grpc.Status;
import lombok.extern.slf4j.Slf4j;
import org.slf4j.MDC;
import org.springframework.context.annotation.Primary;
import org.springframework.stereotype.Component;

import java.util.Collections;
import java.util.LinkedHashMap;
import java.util.Map;
import java.util.Objects;

import static com.tianting.infoloop.constants.ConfigConstants.BRAVE_PROPAGATION_DEBUG_FIELD;
import static com.tianting.infoloop.constants.ConfigConstants.ENTERPRISE_ID_KEY;
import static com.tianting.infoloop.constants.ConfigConstants.LOGGING_PARENT_TRACING_ID;
import static com.tianting.infoloop.constants.ConfigConstants.LOGGING_TRACING_ID;
import static com.tianting.infoloop.constants.ConfigConstants.LOGGING_UNIQUE_ID;
import static com.tianting.infoloop.constants.ConfigConstants.SOURCE_KEY;
import static com.tianting.infoloop.constants.ConfigConstants.USER_ID_KEY;


@Component
@Primary
@Slf4j
public class GrpcServerInterceptor implements ServerInterceptor {
    private final Tracing tracing;
    private final Propagation<String> propagation;
    private final GrpcGetter grpcGetter;

    public GrpcServerInterceptor(final Tracing tracing) {
        this.tracing = tracing;
        this.propagation = tracing.propagation();
        this.grpcGetter = new GrpcGetter(nameToKey(tracing.propagation()));
    }

    private static Map<String, Metadata.Key<String>> nameToKey(Propagation<String> propagation) {
        final Map<String, Metadata.Key<String>> keyByName = new LinkedHashMap<>();
        for (String keyName : propagation.keys()) {
            keyByName.put(keyName, Metadata.Key.of(keyName, Metadata.ASCII_STRING_MARSHALLER));
        }
        for (String keyName : BaggagePropagation.allKeyNames(propagation)) {
            keyByName.put(keyName, Metadata.Key.of(keyName, Metadata.ASCII_STRING_MARSHALLER));
        }
        return Collections.unmodifiableMap(keyByName);
    }

    @Override
    public <ReqT, RespT> ServerCall.Listener<ReqT> interceptCall(final ServerCall<ReqT, RespT> call,
                                                                 final Metadata headers,
                                                                 final ServerCallHandler<ReqT, RespT> next) {
        final var extractor = propagation.extractor(grpcGetter);
        final var extractedContext = extractor.extract(headers);
        final var spanContext = extractedContext.context();
        if (spanContext == null) {
            throw new RuntimeException("rpc call should always have a context");
        }
        final var tracer = tracing.tracer();
        final var span = tracer.newChild(spanContext);
        final var remote = Objects.requireNonNull(call.getAttributes().get(Grpc.TRANSPORT_ATTR_REMOTE_ADDR)).toString().substring(1);
        final var remoteAddressAndPort = remote.split(":");
        final var fullMethodName = call.getMethodDescriptor().getFullMethodName();
        span.remoteIpAndPort(remoteAddressAndPort[0], Integer.parseInt(remoteAddressAndPort[1]));
        span.name(fullMethodName);
        span.start();
        MDC.put(LOGGING_TRACING_ID, span.context().traceIdString());
        MDC.put(LOGGING_UNIQUE_ID, span.context().spanIdString());
        MDC.put(LOGGING_PARENT_TRACING_ID, span.context().parentIdString());
        final var scope = tracing.currentTraceContext().newScope(span.context());

        final var listener = next.startCall(new ForwardingServerCall.SimpleForwardingServerCall<>(call) {

            @Override
            public void sendMessage(RespT message) {
                try (final var scope = tracing.currentTraceContext().maybeScope(span.context())) {
                    delegate().sendMessage(message);
                    if ("true".equals(BRAVE_PROPAGATION_DEBUG_FIELD.getValue(span.context()))) {
                        span.customizer().tag("grpc-response", message.toString());
                    }
                }
            }

            @Override
            public void request(int numMessages) {
                try (final var scope = tracing.currentTraceContext().maybeScope(span.context())) {
                    delegate().request(numMessages);
                }
            }

            @Override
            public void sendHeaders(Metadata headers) {
                try (final var scope = tracing.currentTraceContext().maybeScope(span.context())) {
                    delegate().sendHeaders(headers);
                }
            }

            @Override
            public void close(Status status, Metadata trailers) {
                try (final var scope = tracing.currentTraceContext().maybeScope(span.context())) {
                    delegate().close(status, trailers);
                }
                scope.close();
                span.finish();
                MDC.clear();
            }
        }, headers);
        Context context = setUserContextFromHeaders(headers, fullMethodName);
        Context previous = context.attach();
        try {
            return new MyContextualizedServerCallListener<>(
                    listener,
                    context,
                    tracing,
                    span,
                    scope,
                    fullMethodName);
        } finally {
            context.detach(previous);
        }
    }

    private Context setUserContextFromHeaders(Metadata headers, String methodName) {
        final var userIdKey = Metadata.Key.of(USER_ID_KEY, Metadata.ASCII_STRING_MARSHALLER);
        final var enterpriseIdKey = Metadata.Key.of(ENTERPRISE_ID_KEY, Metadata.ASCII_STRING_MARSHALLER);
        final var sourceKey = Metadata.Key.of(SOURCE_KEY, Metadata.ASCII_STRING_MARSHALLER);
        final var userId = headers.get(userIdKey) == null ? "0" : headers.get(userIdKey);
        final var enterpriseId = headers.get(enterpriseIdKey) == null ? "0" : headers.get(enterpriseIdKey);
        final var source = headers.get(sourceKey) == null ? "0" : headers.get(sourceKey);
        log.info("userContext; methodName:{}, userId:{}, enterpriseId:{}, source:{}", methodName, userId, enterpriseId, source);
        return Context.current();
//                .withValue(ConfigConstants.USER_ID, userId)
//                .withValue(ConfigConstants.ENTERPRISE_ID, enterpriseId)
//                .withValue(ConfigConstants.SOURCE, source);
    }

    private static class GrpcGetter implements Propagation.RemoteGetter<Metadata> {
        private final Map<String, Metadata.Key<String>> keyByName;

        GrpcGetter(final Map<String, Metadata.Key<String>> keyByName) {
            this.keyByName = keyByName;
        }

        @Override
        public Span.Kind spanKind() {
            return Span.Kind.SERVER;
        }

        @Override
        public String get(final Metadata request, final String fieldName) {
            final var key = this.keyByName.get(fieldName);
            if (key == null) {
                return null;
            }
            return request.get(key);
        }
    }
}