diff --git a/client/dtmcli/trans_saga.go b/client/dtmcli/trans_saga.go index 04c8124..db79b4f 100644 --- a/client/dtmcli/trans_saga.go +++ b/client/dtmcli/trans_saga.go @@ -7,6 +7,7 @@ package dtmcli import ( + "context" "github.com/dtm-labs/dtm/client/dtmcli/dtmimp" ) @@ -21,6 +22,13 @@ func NewSaga(server string, gid string) *Saga { return &Saga{TransBase: *dtmimp.NewTransBase(gid, "saga", server, ""), orders: map[int][]int{}} } +// NewSagaWithContext create a saga with context +func NewSagaWithContext(ctx context.Context, server string, gid string) *Saga { + saga := NewSaga(server, gid) + saga.TransBase.Context = ctx + return saga +} + // Add add a saga step func (s *Saga) Add(action string, compensate string, postData interface{}) *Saga { s.Steps = append(s.Steps, map[string]string{"action": action, "compensate": compensate}) diff --git a/client/dtmgrpc/options_test.go b/client/dtmgrpc/options_test.go index 011c36b..6a4ddf1 100644 --- a/client/dtmgrpc/options_test.go +++ b/client/dtmgrpc/options_test.go @@ -1,6 +1,7 @@ package dtmgrpc import ( + "context" "reflect" "testing" @@ -102,3 +103,52 @@ func TestNewSagaGrpc(t *testing.T) { }) } } + +// TestNewSagaGrpcWithContext ut for NewSagaGrpcWithContext +func TestNewSagaGrpcWithContext(t *testing.T) { + var ( + ctx = context.Background() + server = "dmt_server_address" + gidNoOptions = "msg_no_setup_options" + gidTraceIDXXX = "msg_setup_options_trace_id_xxx" + sagaWithTraceIDXXX = &SagaGrpc{Saga: *dtmcli.NewSagaWithContext(ctx, server, gidTraceIDXXX)} + traceIDHeaders = map[string]string{ + "x-trace-id": "xxx", + } + ) + sagaWithTraceIDXXX.BranchHeaders = traceIDHeaders + type args struct { + gid string + opts []TransBaseOption + } + tests := []struct { + name string + args args + want *SagaGrpc + }{ + { + name: "no setup options", + args: args{gid: gidNoOptions}, + want: &SagaGrpc{Saga: *dtmcli.NewSaga(server, gidNoOptions)}, + }, + { + name: "msg with trace_id", + args: args{ + gid: gidTraceIDXXX, + opts: []TransBaseOption{ + WithBranchHeaders(traceIDHeaders), + }, + }, + want: sagaWithTraceIDXXX, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := NewSagaGrpcWithContext(ctx, server, tt.args.gid, tt.args.opts...) + t.Logf("TestNewSagaGrpc %s got %+v\n", tt.name, got) + if !reflect.DeepEqual(got, tt.want) { + t.Errorf("NewSagaGrpc() = %v, want %v", got, tt.want) + } + }) + } +} diff --git a/client/dtmgrpc/saga.go b/client/dtmgrpc/saga.go index cc13df8..59766ef 100644 --- a/client/dtmgrpc/saga.go +++ b/client/dtmgrpc/saga.go @@ -7,6 +7,7 @@ package dtmgrpc import ( + "context" "github.com/dtm-labs/dtm/client/dtmcli" "github.com/dtm-labs/dtm/client/dtmgrpc/dtmgimp" "google.golang.org/protobuf/proto" @@ -28,6 +29,17 @@ func NewSagaGrpc(server string, gid string, opts ...TransBaseOption) *SagaGrpc { return sg } +// NewSagaGrpcWithContext create a saga with context +func NewSagaGrpcWithContext(ctx context.Context, server string, gid string, opts ...TransBaseOption) *SagaGrpc { + sg := &SagaGrpc{Saga: *dtmcli.NewSagaWithContext(ctx, server, gid)} + + for _, opt := range opts { + opt(&sg.TransBase) + } + + return sg +} + // Add add a saga step func (s *SagaGrpc) Add(action string, compensate string, payload proto.Message) *SagaGrpc { s.Steps = append(s.Steps, map[string]string{"action": action, "compensate": compensate})