From 1bd0dcfca6ef60ae8777159d446b59a0bac0aa75 Mon Sep 17 00:00:00 2001 From: Christian Schwarz Date: Wed, 15 Jan 2020 17:37:23 +0100 Subject: [PATCH] WIP request ids + trace context --- daemon/job/passive.go | 33 ++++++++++- daemon/logging/build_logging.go | 5 ++ daemon/logging/logging_formatters.go | 3 +- endpoint/context.go | 4 ++ endpoint/endpoint.go | 6 +- go.mod | 1 + go.sum | 1 + replication/driver/replication_driver.go | 9 ++- replication/logic/replication_logic.go | 11 ++-- rpc/dataconn/dataconn_client.go | 2 +- rpc/dataconn/dataconn_server.go | 33 +++++++---- .../authlistener_grpc_adaptor.go | 11 +++- .../authlistener_grpc_adaptor_wrapper.go | 48 +++++++++++++-- rpc/rpc_client.go | 10 +++- rpc/rpc_requestid.go | 58 +++++++++++++++++++ rpc/rpc_server.go | 29 ++++++++-- 16 files changed, 227 insertions(+), 37 deletions(-) create mode 100644 rpc/rpc_requestid.go diff --git a/daemon/job/passive.go b/daemon/job/passive.go index ce3bafe..6c9e14e 100644 --- a/daemon/job/passive.go +++ b/daemon/job/passive.go @@ -4,6 +4,7 @@ import ( "context" "fmt" + "github.com/golang/protobuf/proto" "github.com/pkg/errors" "github.com/prometheus/client_golang/prometheus" @@ -180,11 +181,39 @@ func (j *PassiveSide) Run(ctx context.Context) { } ctxInterceptor := func(handlerCtx context.Context) context.Context { - return logging.WithSubsystemLoggers(handlerCtx, log) + reqLog := log.WithField(logging.ReqIDField, rpc.MustGetRequestID(handlerCtx)) + return logging.WithSubsystemLoggers(handlerCtx, reqLog) + } + + // ctxInterceptor already injected request loggers (with request id) + pre := func(ctx context.Context, method string, req interface{}) { + l := endpoint.GetLogger(ctx).WithField("method", method).WithField("reqT", fmt.Sprintf("%T", req)) + if p, ok := req.(proto.Message); ok { + l = l.WithField("req", proto.CompactTextString(p)) + } + ci, ok := ctx.Value(endpoint.ClientIdentityKey).(string) + if !ok { + // like endpoint, we must assume that rpc.Server sets it + panic(method) + } + l = l.WithField("clientIdentity", ci) + l.Debug("incoming") + } + post := func(ctx context.Context, res interface{}, err error) { + l := endpoint.GetLogger(ctx). + WithField("responseT", fmt.Sprintf("%T", res)). + WithField("errT", fmt.Sprintf("%T", err)) + if s, ok := res.(fmt.Stringer); ok { + l = l.WithField("response", s.String()) + } + if err != nil { + l = l.WithError(err) + } + l.Debug("handler done") } rpcLoggers := rpc.GetLoggersOrPanic(ctx) // WithSubsystemLoggers above - server := rpc.NewServer(handler, rpcLoggers, ctxInterceptor) + server := rpc.NewServer(handler, rpcLoggers, ctxInterceptor, pre, post) listener, err := j.listen() if err != nil { diff --git a/daemon/logging/build_logging.go b/daemon/logging/build_logging.go index 5021bed..1ab35ea 100644 --- a/daemon/logging/build_logging.go +++ b/daemon/logging/build_logging.go @@ -96,6 +96,11 @@ func WithSubsystemLoggers(ctx context.Context, log logger.Logger) context.Contex Control: log.WithField(SubsysField, SubsysRPCControl), Data: log.WithField(SubsysField, SubsysRPCData), }, + // rpc.Loggers{ + // General: logger.NewNullLogger(), + // Control: logger.NewNullLogger(), + // Data: logger.NewNullLogger(), + // }, ) return ctx } diff --git a/daemon/logging/logging_formatters.go b/daemon/logging/logging_formatters.go index 1ee8260..9169246 100644 --- a/daemon/logging/logging_formatters.go +++ b/daemon/logging/logging_formatters.go @@ -22,6 +22,7 @@ const ( const ( JobField string = "job" SubsysField string = "subsystem" + ReqIDField string = "reqid" ) type MetadataFlags int64 @@ -85,7 +86,7 @@ func (f *HumanFormatter) Format(e *logger.Entry) (out []byte, err error) { fmt.Fprintf(&line, "[%s]", col.Sprint(e.Level.Short())) } - prefixFields := []string{JobField, SubsysField} + prefixFields := []string{JobField, SubsysField, ReqIDField} prefixed := make(map[string]bool, len(prefixFields)+2) for _, field := range prefixFields { val, ok := e.Fields[field] diff --git a/endpoint/context.go b/endpoint/context.go index 700d7f6..7a81700 100644 --- a/endpoint/context.go +++ b/endpoint/context.go @@ -19,6 +19,10 @@ func WithLogger(ctx context.Context, log Logger) context.Context { return context.WithValue(ctx, contextKeyLogger, log) } +func GetLogger(ctx context.Context) Logger { + return getLogger(ctx) +} + func getLogger(ctx context.Context) Logger { if l, ok := ctx.Value(contextKeyLogger).(Logger); ok { return l diff --git a/endpoint/endpoint.go b/endpoint/endpoint.go index 50bb409..5e23b1d 100644 --- a/endpoint/endpoint.go +++ b/endpoint/endpoint.go @@ -159,6 +159,7 @@ func sendArgsFromPDUAndValidateExists(ctx context.Context, fs string, fsv *pdu.F } func (s *Sender) Send(ctx context.Context, r *pdu.SendReq) (*pdu.SendRes, zfs.StreamCopier, error) { + getLogger(ctx).WithField("req", r.String()).Debug("incoming Send request") _, err := s.filterCheckFS(r.Filesystem) if err != nil { @@ -259,6 +260,7 @@ func (s *Sender) Send(ctx context.Context, r *pdu.SendReq) (*pdu.SendRes, zfs.St } func (p *Sender) SendCompleted(ctx context.Context, r *pdu.SendCompletedReq) (*pdu.SendCompletedRes, error) { + getLogger(ctx).WithField("req", r.String()).Debug("incoming SendCompleted request") orig := r.GetOriginalReq() // may be nil, always use proto getters fs := orig.GetFilesystem() @@ -510,7 +512,7 @@ func (s *Receiver) ListFilesystems(ctx context.Context, req *pdu.ListFilesystemR // present filesystem without the root_fs prefix fss := make([]*pdu.Filesystem, 0, len(filtered)) for _, a := range filtered { - l := getLogger(ctx).WithField("fs", a) + l := getLogger(ctx).WithField("fs", a.ToString()) ph, err := zfs.ZFSGetFilesystemPlaceholderState(a) if err != nil { l.WithError(err).Error("error getting placeholder state") @@ -597,7 +599,7 @@ func (s *Receiver) Send(ctx context.Context, req *pdu.SendReq) (*pdu.SendRes, zf var maxConcurrentZFSRecvSemaphore = semaphore.New(envconst.Int64("ZREPL_ENDPOINT_MAX_CONCURRENT_RECV", 10)) func (s *Receiver) Receive(ctx context.Context, req *pdu.ReceiveReq, receive zfs.StreamCopier) (*pdu.ReceiveRes, error) { - getLogger(ctx).Debug("incoming Receive") + getLogger(ctx).WithField("rr", req.String()).Debug("incoming Receive") defer receive.Close() root := s.clientRootFromCtx(ctx) diff --git a/go.mod b/go.mod index 2bbbea1..df6907d 100644 --- a/go.mod +++ b/go.mod @@ -7,6 +7,7 @@ require ( github.com/gdamore/tcell v1.2.0 github.com/go-logfmt/logfmt v0.4.0 github.com/go-sql-driver/mysql v1.4.1-0.20190907122137-b2c03bcae3d4 + github.com/gogo/protobuf v1.1.1 github.com/golang/protobuf v1.3.2 github.com/google/uuid v1.1.1 github.com/jinzhu/copier v0.0.0-20170922082739-db4671f3a9b8 diff --git a/go.sum b/go.sum index 4b55d9a..b39c815 100644 --- a/go.sum +++ b/go.sum @@ -68,6 +68,7 @@ github.com/go-toolsmith/strparse v1.0.0/go.mod h1:YI2nUKP9YGZnL/L1/DLFBfixrcjslW github.com/go-toolsmith/typep v0.0.0-20181030061450-d63dc7650676/go.mod h1:JSQCQMUPdRlMZFswiq3TGpNp1GMktqkR2Ns5AIQkATU= github.com/go-toolsmith/typep v1.0.0/go.mod h1:JSQCQMUPdRlMZFswiq3TGpNp1GMktqkR2Ns5AIQkATU= github.com/gobwas/glob v0.2.3/go.mod h1:d3Ez4x06l9bZtSvzIay5+Yzi0fmZzPgnTbPcKjJAkT8= +github.com/gogo/protobuf v1.1.1 h1:72R+M5VuhED/KujmZVcIquuo8mBgX4oVda//DQb3PXo= github.com/gogo/protobuf v1.1.1/go.mod h1:r8qH/GZQm5c6nD/R0oafs1akxWv10x8SbQlK7atdtwQ= github.com/gogo/protobuf v1.2.1/go.mod h1:hp+jE20tsWTFYpLwKvXlhS1hjn+gTNwPg2I6zVXpSg4= github.com/golang/glog v0.0.0-20160126235308-23def4e6c14b h1:VKtxabqXZkF25pY9ekfRL6a582T4P37/31XEstQ5p58= diff --git a/replication/driver/replication_driver.go b/replication/driver/replication_driver.go index a9c2ce7..eecd499 100644 --- a/replication/driver/replication_driver.go +++ b/replication/driver/replication_driver.go @@ -5,6 +5,7 @@ import ( "errors" "fmt" "net" + "os" "sort" "strings" "sync" @@ -14,6 +15,7 @@ import ( "google.golang.org/grpc/status" "github.com/zrepl/zrepl/replication/report" + "github.com/zrepl/zrepl/tracing" "github.com/zrepl/zrepl/util/chainlock" "github.com/zrepl/zrepl/util/envconst" ) @@ -206,6 +208,7 @@ func Do(ctx context.Context, planner Planner) (ReportFunc, WaitFunc) { } run.attempts = append(run.attempts, cur) run.l.DropWhile(func() { + ctx := tracing.Child(ctx, fmt.Sprintf("attempt#%d", ano)) // shadow cur.do(ctx, prev) }) prev = cur @@ -281,7 +284,7 @@ func Do(ctx context.Context, planner Planner) (ReportFunc, WaitFunc) { } func (a *attempt) do(ctx context.Context, prev *attempt) { - pfss, err := a.planner.Plan(ctx) + pfss, err := a.planner.Plan(tracing.Child(ctx, "plan")) errTime := time.Now() defer a.l.Lock().Unlock() if err != nil { @@ -361,7 +364,7 @@ func (a *attempt) do(ctx context.Context, prev *attempt) { fssesDone.Add(1) go func(f *fs) { defer fssesDone.Done() - f.do(ctx, stepQueue, prevs[f]) + f.do(tracing.Child(ctx, f.fs.ReportInfo().Name), stepQueue, prevs[f]) }(f) } a.l.DropWhile(func() { @@ -372,6 +375,8 @@ func (a *attempt) do(ctx context.Context, prev *attempt) { func (fs *fs) do(ctx context.Context, pq *stepQueue, prev *fs) { + fmt.Fprintf(os.Stderr, "CHILD STACK: %v\n", tracing.GetStack(ctx)) + defer fs.l.Lock().Unlock() // get planned steps from replication logic diff --git a/replication/logic/replication_logic.go b/replication/logic/replication_logic.go index 2617217..548250b 100644 --- a/replication/logic/replication_logic.go +++ b/replication/logic/replication_logic.go @@ -606,7 +606,7 @@ func (s *Step) doReplication(ctx context.Context) error { log := getLogger(ctx) sr := s.buildSendRequest(false) - log.Debug("initiate send request") + log.WithField("sr", sr.String()).Debug("initiate send request") sres, sstreamCopier, err := s.sender.Send(ctx, sr) if err != nil { log.WithError(err).Error("send request failed") @@ -633,7 +633,7 @@ func (s *Step) doReplication(ctx context.Context) error { To: sr.GetTo(), ClearResumeToken: !sres.UsedResumeToken, } - log.Debug("initiate receive request") + log.WithField("rr", rr.String()).Debug("initiate receive request") _, err = s.receiver.Receive(ctx, rr, byteCountingStream) if err != nil { log. @@ -649,10 +649,11 @@ func (s *Step) doReplication(ctx context.Context) error { } log.Debug("receive finished") - log.Debug("tell sender replication completed") - _, err = s.sender.SendCompleted(ctx, &pdu.SendCompletedReq{ + scr := &pdu.SendCompletedReq{ OriginalReq: sr, - }) + } + log.WithField("scr", scr.String()).Debug("tell sender replication completed") + _, err = s.sender.SendCompleted(ctx, scr) if err != nil { log.WithError(err).Error("error telling sender that replication completed successfully") return err diff --git a/rpc/dataconn/dataconn_client.go b/rpc/dataconn/dataconn_client.go index 6cc049b..f655a63 100644 --- a/rpc/dataconn/dataconn_client.go +++ b/rpc/dataconn/dataconn_client.go @@ -75,7 +75,7 @@ func (c *Client) recv(ctx context.Context, conn *stream.Conn, res proto.Message) if err != nil { return err } - header := string(headerBuf) + header := string(headerBuf) // FIXME if strings.HasPrefix(header, responseHeaderHandlerErrorPrefix) { // FIXME distinguishable error type return &RemoteHandlerError{strings.TrimPrefix(header, responseHeaderHandlerErrorPrefix)} diff --git a/rpc/dataconn/dataconn_server.go b/rpc/dataconn/dataconn_server.go index fbd1a95..200b08d 100644 --- a/rpc/dataconn/dataconn_server.go +++ b/rpc/dataconn/dataconn_server.go @@ -14,8 +14,8 @@ import ( "github.com/zrepl/zrepl/zfs" ) -// WireInterceptor has a chance to exchange the context and connection on each client connection. -type WireInterceptor func(ctx context.Context, rawConn *transport.AuthConn) (context.Context, *transport.AuthConn) +// ReqInterceptor has a chance to exchange the context and connection on each request. +type ReqInterceptor func(ctx context.Context, rawConn *transport.AuthConn) (context.Context, *transport.AuthConn) // Handler implements the functionality that is exposed by Server to the Client. type Handler interface { @@ -30,19 +30,27 @@ type Handler interface { PingDataconn(ctx context.Context, r *pdu.PingReq) (*pdu.PingRes, error) } +type PreHandlerInspector = func(ctx context.Context, endpoint string, req interface{}) + +type PostHandlerInspector = func(ctx context.Context, response interface{}, err error) + type Logger = logger.Logger type Server struct { - h Handler - wi WireInterceptor - log Logger + h Handler + ri ReqInterceptor + log Logger + pre PreHandlerInspector + post PostHandlerInspector } -func NewServer(wi WireInterceptor, logger Logger, handler Handler) *Server { +func NewServer(ri ReqInterceptor, logger Logger, pre PreHandlerInspector, handler Handler, post PostHandlerInspector) *Server { return &Server{ h: handler, - wi: wi, + ri: ri, log: logger, + pre: pre, + post: post, } } @@ -83,8 +91,8 @@ func (s *Server) serveConn(nc *transport.AuthConn) { defer s.log.Debug("serveConn done") ctx := context.Background() - if s.wi != nil { - ctx, nc = s.wi(ctx, nc) + if s.ri != nil { + ctx, nc = s.ri(ctx, nc) } c := stream.Wrap(nc, HeartbeatInterval, HeartbeatPeerTimeout) @@ -100,7 +108,7 @@ func (s *Server) serveConn(nc *transport.AuthConn) { s.log.WithError(err).Error("error reading structured part") return } - endpoint := string(header) + endpoint := string(header) // FIXME reqStructured, err := c.ReadStreamedMessage(ctx, RequestStructuredMaxSize, ReqStructured) if err != nil { @@ -120,6 +128,7 @@ func (s *Server) serveConn(nc *transport.AuthConn) { s.log.WithError(err).Error("cannot unmarshal send request") return } + s.pre(ctx, endpoint, &req) res, sendStream, handlerErr = s.h.Send(ctx, &req) // SHADOWING case EndpointRecv: var req pdu.ReceiveReq @@ -127,6 +136,7 @@ func (s *Server) serveConn(nc *transport.AuthConn) { s.log.WithError(err).Error("cannot unmarshal receive request") return } + s.pre(ctx, endpoint, &req) res, handlerErr = s.h.Receive(ctx, &req, &streamCopier{streamConn: c, closeStreamOnClose: false}) // SHADOWING case EndpointPing: var req pdu.PingReq @@ -134,6 +144,7 @@ func (s *Server) serveConn(nc *transport.AuthConn) { s.log.WithError(err).Error("cannot unmarshal ping request") return } + s.pre(ctx, endpoint, &req) res, handlerErr = s.h.PingDataconn(ctx, &req) // SHADOWING default: s.log.WithField("endpoint", endpoint).Error("unknown endpoint") @@ -142,6 +153,8 @@ func (s *Server) serveConn(nc *transport.AuthConn) { s.log.WithField("endpoint", endpoint).WithField("errType", fmt.Sprintf("%T", handlerErr)).Debug("handler returned") + s.post(ctx, res, handlerErr) + // prepare protobuf now to return the protobuf error in the header // if marshaling fails. We consider failed marshaling a handler error var protobuf *bytes.Buffer diff --git a/rpc/grpcclientidentity/authlistener_grpc_adaptor.go b/rpc/grpcclientidentity/authlistener_grpc_adaptor.go index ec7d5a1..bd0ba27 100644 --- a/rpc/grpcclientidentity/authlistener_grpc_adaptor.go +++ b/rpc/grpcclientidentity/authlistener_grpc_adaptor.go @@ -101,7 +101,11 @@ func (*transportCredentials) OverrideServerName(string) error { type ContextInterceptor = func(ctx context.Context) context.Context -func NewInterceptors(logger Logger, clientIdentityKey interface{}, ctxInterceptor ContextInterceptor) (unary grpc.UnaryServerInterceptor, stream grpc.StreamServerInterceptor) { +type PreHandlerInspector func(ctx context.Context, endpoint string, req interface{}) + +type PostHandlerInspector func(ctx context.Context, response interface{}, err error) + +func NewInterceptors(logger Logger, clientIdentityKey interface{}, ctxInterceptor ContextInterceptor, pre PreHandlerInspector, post PostHandlerInspector) (unary grpc.UnaryServerInterceptor, stream grpc.StreamServerInterceptor) { unary = func(ctx context.Context, req interface{}, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (resp interface{}, err error) { logger.WithField("fullMethod", info.FullMethod).Debug("request") p, ok := peer.FromContext(ctx) @@ -118,7 +122,10 @@ func NewInterceptors(logger Logger, clientIdentityKey interface{}, ctxIntercepto if ctxInterceptor != nil { ctx = ctxInterceptor(ctx) } - return handler(ctx, req) + pre(ctx, info.FullMethod, req) + res, err := handler(ctx, req) + post(ctx, res, err) + return res, err } stream = func(srv interface{}, ss grpc.ServerStream, info *grpc.StreamServerInfo, handler grpc.StreamHandler) error { panic("unimplemented") diff --git a/rpc/grpcclientidentity/grpchelper/authlistener_grpc_adaptor_wrapper.go b/rpc/grpcclientidentity/grpchelper/authlistener_grpc_adaptor_wrapper.go index 85486de..8c3b329 100644 --- a/rpc/grpcclientidentity/grpchelper/authlistener_grpc_adaptor_wrapper.go +++ b/rpc/grpcclientidentity/grpchelper/authlistener_grpc_adaptor_wrapper.go @@ -4,14 +4,18 @@ package grpchelper import ( "context" + "fmt" + "strings" "time" "google.golang.org/grpc" "google.golang.org/grpc/keepalive" + "google.golang.org/grpc/metadata" "github.com/zrepl/zrepl/logger" "github.com/zrepl/zrepl/rpc/grpcclientidentity" "github.com/zrepl/zrepl/rpc/netadaptor" + "github.com/zrepl/zrepl/tracing" "github.com/zrepl/zrepl/transport" ) @@ -28,15 +32,40 @@ type Logger = logger.Logger // ClientConn is an easy-to-use wrapper around the Dialer and TransportCredentials interface // to produce a grpc.ClientConn -func ClientConn(cn transport.Connecter, log Logger) *grpc.ClientConn { +func ClientConn(cn transport.Connecter, log Logger, genReqID func() string) *grpc.ClientConn { ka := grpc.WithKeepaliveParams(keepalive.ClientParameters{ Time: StartKeepalivesAfterInactivityDuration, Timeout: KeepalivePeerTimeout, PermitWithoutStream: true, }) + unaryIntcpt := grpc.WithUnaryInterceptor(func(ctx context.Context, method string, req, reply interface{}, cc *grpc.ClientConn, invoker grpc.UnaryInvoker, opts ...grpc.CallOption) error { + + traceStack := tracing.GetStack(ctx) + if len(traceStack) == 0 { + panic("implementation error: expecting trace stack") + } + reqId := genReqID() + " " + strings.Join(traceStack, " <> ") + + l := log + l = l.WithField("reqID", reqId) + l = l.WithField("reqT", fmt.Sprintf("%T", req)) + l = l.WithField("repT", fmt.Sprintf("%T", reply)) + l = l.WithField("ctxT", fmt.Sprintf("%T", ctx)) + l.Debug("unary intercepted") + + ctx = metadata.AppendToOutgoingContext(ctx, "zrepl-req-id", reqId) + + err := invoker(ctx, method, req, reply, cc, opts...) + + l = log.WithField("repT", fmt.Sprintf("%T", reply)) + l = log.WithField("errT", fmt.Sprintf("%T", err)) + l.Debug("reply intercepted") + + return err + }) dialerOption := grpc.WithDialer(grpcclientidentity.NewDialer(log, cn)) cred := grpc.WithTransportCredentials(grpcclientidentity.NewTransportCredentials(log)) - cc, err := grpc.DialContext(context.Background(), "doesn't matter done by dialer", dialerOption, cred, ka) + cc, err := grpc.DialContext(context.Background(), "doesn't matter done by dialer", dialerOption, cred, ka, unaryIntcpt) if err != nil { log.WithError(err).Error("cannot create gRPC client conn (non-blocking)") // It's ok to panic here: the we call grpc.DialContext without the @@ -50,7 +79,7 @@ func ClientConn(cn transport.Connecter, log Logger) *grpc.ClientConn { } // NewServer is a convenience interface around the TransportCredentials and Interceptors interface. -func NewServer(authListener transport.AuthenticatedListener, clientIdentityKey interface{}, logger grpcclientidentity.Logger, ctxInterceptor grpcclientidentity.ContextInterceptor) (srv *grpc.Server, serve func() error) { +func NewServer(authListener transport.AuthenticatedListener, clientIdentityKey interface{}, logger grpcclientidentity.Logger, interceptor grpcclientidentity.ContextInterceptor, pre grpcclientidentity.PreHandlerInspector, post grpcclientidentity.PostHandlerInspector) (srv *grpc.Server, serve func() error) { ka := grpc.KeepaliveParams(keepalive.ServerParameters{ Time: StartKeepalivesAfterInactivityDuration, Timeout: KeepalivePeerTimeout, @@ -60,8 +89,17 @@ func NewServer(authListener transport.AuthenticatedListener, clientIdentityKey i PermitWithoutStream: true, }) tcs := grpcclientidentity.NewTransportCredentials(logger) - unary, stream := grpcclientidentity.NewInterceptors(logger, clientIdentityKey, ctxInterceptor) - srv = grpc.NewServer(grpc.Creds(tcs), grpc.UnaryInterceptor(unary), grpc.StreamInterceptor(stream), ka, ep) + unary, stream := grpcclientidentity.NewInterceptors(logger, clientIdentityKey, interceptor, pre, post) + unary2 := func(ctx context.Context, req interface{}, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (resp interface{}, err error) { + md, ok := metadata.FromIncomingContext(ctx) + if !ok { + logger.Error("no request id in incoming request") + } else { + logger.Info(fmt.Sprintf("req-id = %v", md.Get("zrepl-req-id"))) + } + return unary(ctx, req, info, handler) + } + srv = grpc.NewServer(grpc.Creds(tcs), grpc.UnaryInterceptor(unary2), grpc.StreamInterceptor(stream), ka, ep) serve = func() error { if err := srv.Serve(netadaptor.New(authListener, logger)); err != nil { diff --git a/rpc/rpc_client.go b/rpc/rpc_client.go index bdc81de..969c58f 100644 --- a/rpc/rpc_client.go +++ b/rpc/rpc_client.go @@ -26,6 +26,7 @@ import ( // Client implements the active side of a replication setup. // It satisfies the Endpoint, Sender and Receiver interface defined by package replication. type Client struct { + reqIDGen *requestIDGenerator dataClient *dataconn.Client controlClient pdu.ReplicationClient // this the grpc client instance, see constructor controlConn *grpc.ClientConn @@ -47,10 +48,13 @@ func NewClient(cn transport.Connecter, loggers Loggers) *Client { muxedConnecter := mux(cn) c := &Client{ - loggers: loggers, - closed: make(chan struct{}), + loggers: loggers, + closed: make(chan struct{}), + reqIDGen: newRequestIDGenerator(), } - grpcConn := grpchelper.ClientConn(muxedConnecter.control, loggers.Control) + grpcConn := grpchelper.ClientConn(muxedConnecter.control, loggers.Control, func() string { + return c.reqIDGen.newID().String() + }) go func() { ctx, cancel := context.WithCancel(context.Background()) diff --git a/rpc/rpc_requestid.go b/rpc/rpc_requestid.go new file mode 100644 index 0000000..e44a7ef --- /dev/null +++ b/rpc/rpc_requestid.go @@ -0,0 +1,58 @@ +package rpc + +import ( + "context" + "encoding/base64" + "strings" + + "github.com/google/uuid" +) + +type requestIDGenerator struct{} + +func newRequestIDGenerator() *requestIDGenerator { + return &requestIDGenerator{} +} + +type RequestID struct{ s string } + +func (r RequestID) String() string { return r.s } + +type requestIDContextKey int + +const ( + requestIDContextKeyRequestID requestIDContextKey = 1 + iota +) + +func (g *requestIDGenerator) newID() RequestID { + id := uuid.New() + var buf strings.Builder + enc := base64.NewEncoder(base64.RawStdEncoding, &buf) + n, err := enc.Write(id[:]) + if err != nil { + panic(err) + } else if n != len(id) { + panic(n) + } + if err := enc.Close(); err != nil { + panic(err) + } + return RequestID{buf.String()} +} + +func (g *requestIDGenerator) inject(ctx context.Context) context.Context { + return context.WithValue(ctx, requestIDContextKeyRequestID, g.newID()) +} + +func GetRequestID(ctx context.Context) (id RequestID, ok bool) { + id, ok = ctx.Value(requestIDContextKeyRequestID).(RequestID) + return id, ok +} + +func MustGetRequestID(ctx context.Context) (id RequestID) { + id, ok := GetRequestID(ctx) + if !ok { + panic("calling context expectes request id to bet set") + } + return id +} diff --git a/rpc/rpc_server.go b/rpc/rpc_server.go index 69e6335..744ce16 100644 --- a/rpc/rpc_server.go +++ b/rpc/rpc_server.go @@ -7,6 +7,7 @@ import ( "github.com/zrepl/zrepl/endpoint" "github.com/zrepl/zrepl/replication/logic/pdu" "github.com/zrepl/zrepl/rpc/dataconn" + "github.com/zrepl/zrepl/rpc/grpcclientidentity" "github.com/zrepl/zrepl/rpc/grpcclientidentity/grpchelper" "github.com/zrepl/zrepl/rpc/versionhandshake" "github.com/zrepl/zrepl/transport" @@ -30,15 +31,35 @@ type Server struct { dataServerServe serveFunc } -type HandlerContextInterceptor func(ctx context.Context) context.Context +type CtxInterceptor = grpcclientidentity.ContextInterceptor + +func chainInterceptors(interceptors []CtxInterceptor) CtxInterceptor { + return func(ctx context.Context) (_ context.Context) { + for _, i := range interceptors { + ctx = i(ctx) + } + return ctx + } +} + +type PreHandlerInspector = func(ctx context.Context, endpoint string, req interface{}) + +type PostHandlerInspector = func(ctx context.Context, response interface{}, err error) // config must be valid (use its Validate function). -func NewServer(handler Handler, loggers Loggers, ctxInterceptor HandlerContextInterceptor) *Server { +func NewServer(handler Handler, loggers Loggers, ctxInterceptor CtxInterceptor, pre PreHandlerInspector, post PostHandlerInspector) *Server { + + reqIdGen := newRequestIDGenerator() + + ctxInterceptor = chainInterceptors([]CtxInterceptor{ + reqIdGen.inject, + ctxInterceptor, + }) // setup control server controlServerServe := func(ctx context.Context, controlListener transport.AuthenticatedListener, errOut chan<- error) { - controlServer, serve := grpchelper.NewServer(controlListener, endpoint.ClientIdentityKey, loggers.Control, ctxInterceptor) + controlServer, serve := grpchelper.NewServer(controlListener, endpoint.ClientIdentityKey, loggers.Control, ctxInterceptor, pre, post) pdu.RegisterReplicationServer(controlServer, handler) // give time for graceful stop until deadline expires, then hard stop @@ -63,7 +84,7 @@ func NewServer(handler Handler, loggers Loggers, ctxInterceptor HandlerContextIn } return ctx, wire } - dataServer := dataconn.NewServer(dataServerClientIdentitySetter, loggers.Data, handler) + dataServer := dataconn.NewServer(dataServerClientIdentitySetter, loggers.Data, pre, handler, post) dataServerServe := func(ctx context.Context, dataListener transport.AuthenticatedListener, errOut chan<- error) { dataServer.Serve(ctx, dataListener) errOut <- nil // TODO bad design of dataServer?