diff --git a/dtmsvr/svr.go b/dtmsvr/svr.go index 288bec2..46df9ba 100644 --- a/dtmsvr/svr.go +++ b/dtmsvr/svr.go @@ -22,6 +22,8 @@ import ( "github.com/dtm-labs/dtm/dtmutil" "github.com/dtm-labs/dtmdriver" "google.golang.org/grpc" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" ) // StartSvr StartSvr @@ -56,7 +58,7 @@ func StartSvr() *gin.Engine { // start grpc server lis, err := net.Listen("tcp", fmt.Sprintf(":%d", conf.GrpcPort)) logger.FatalIfError(err) - s := grpc.NewServer(grpc.ChainUnaryInterceptor(grpcMetrics, dtmgimp.GrpcServerLog)) + s := grpc.NewServer(grpc.ChainUnaryInterceptor(grpcRecover, grpcMetrics, dtmgimp.GrpcServerLog)) dtmgpb.RegisterDtmServer(s, &dtmServer{}) logger.Infof("grpc listening at %v", lis.Addr()) go func() { @@ -136,3 +138,13 @@ func updateBranchAsync() { flushBranchs() } } + +func grpcRecover(ctx context.Context, req interface{}, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (res interface{}, rerr error) { + defer func() { + if x := recover(); x != nil { + rerr = status.Errorf(codes.Internal, "%v", x) + } + }() + res, rerr = handler(ctx, req) + return +} diff --git a/test/workflow_grpc_test.go b/test/workflow_grpc_test.go index a75f844..726f83b 100644 --- a/test/workflow_grpc_test.go +++ b/test/workflow_grpc_test.go @@ -19,7 +19,7 @@ import ( ) func TestWorkflowGrpcSimple(t *testing.T) { - workflow.SetProtocolForTest(dtmimp.ProtocolHTTP) + workflow.SetProtocolForTest(dtmimp.ProtocolGRPC) req := &busi.ReqGrpc{Amount: 30, TransInResult: "FAILURE"} gid := dtmimp.GetFuncName() workflow.Register(gid, func(wf *workflow.Workflow, data []byte) error { @@ -33,7 +33,7 @@ func TestWorkflowGrpcSimple(t *testing.T) { return err }) err := workflow.Execute(gid, gid, dtmgimp.MustProtoMarshal(req)) - assert.Error(t, err, dtmcli.ErrFailure) + assert.Error(t, err) assert.Equal(t, StatusFailed, getTransStatus(gid)) }