diff --git a/api.go b/api.go index 22479c61b..7c2da4893 100644 --- a/api.go +++ b/api.go @@ -1585,30 +1585,30 @@ func (api *API) PrimaryReplicaNodeURL() url.URL { return node.URI.URL() } -func (api *API) StartTransaction(ctx context.Context, id string, timeout time.Duration, exclusive bool, remote bool) (Transaction, error) { +func (api *API) StartTransaction(ctx context.Context, id string, timeout time.Duration, exclusive bool, remote bool) (*Transaction, error) { if err := api.validate(apiStartTransaction); err != nil { - return Transaction{}, errors.Wrap(err, "validating api method") + return nil, errors.Wrap(err, "validating api method") } return api.server.StartTransaction(ctx, id, timeout, exclusive, remote) } -func (api *API) FinishTransaction(ctx context.Context, id string, remote bool) (Transaction, error) { +func (api *API) FinishTransaction(ctx context.Context, id string, remote bool) (*Transaction, error) { if err := api.validate(apiFinishTransaction); err != nil { - return Transaction{}, errors.Wrap(err, "validating api method") + return nil, errors.Wrap(err, "validating api method") } return api.server.FinishTransaction(ctx, id, remote) } -func (api *API) Transactions(ctx context.Context) (map[string]Transaction, error) { +func (api *API) Transactions(ctx context.Context) (map[string]*Transaction, error) { if err := api.validate(apiTransactions); err != nil { return nil, errors.Wrap(err, "validating api method") } return api.server.Transactions(ctx) } -func (api *API) GetTransaction(ctx context.Context, id string, remote bool) (Transaction, error) { +func (api *API) GetTransaction(ctx context.Context, id string, remote bool) (*Transaction, error) { if err := api.validate(apiGetTransaction); err != nil { - return Transaction{}, errors.Wrap(err, "validating api method") + return nil, errors.Wrap(err, "validating api method") } return api.server.GetTransaction(ctx, id, remote) } diff --git a/client.go b/client.go index 1fb893ee8..01579b1e8 100644 --- a/client.go +++ b/client.go @@ -75,10 +75,10 @@ type InternalClient interface { ImportRoaring(ctx context.Context, uri *URI, index, field string, shard uint64, remote bool, req *ImportRoaringRequest) error ImportColumnAttrs(ctx context.Context, uri *URI, index string, req *ImportColumnAttrsRequest) error - StartTransaction(ctx context.Context, id string, timeout time.Duration, exclusive bool) (Transaction, error) - FinishTransaction(ctx context.Context, id string) (Transaction, error) - Transactions(ctx context.Context) (map[string]Transaction, error) - GetTransaction(ctx context.Context, id string) (Transaction, error) + StartTransaction(ctx context.Context, id string, timeout time.Duration, exclusive bool) (*Transaction, error) + FinishTransaction(ctx context.Context, id string) (*Transaction, error) + Transactions(ctx context.Context) (map[string]*Transaction, error) + GetTransaction(ctx context.Context, id string) (*Transaction, error) } //=============== @@ -211,15 +211,15 @@ func (n nopInternalClient) RetrieveTranslatePartitionFromURI(ctx context.Context return nil, nil } -func (n nopInternalClient) StartTransaction(ctx context.Context, id string, timeout time.Duration, exclusive bool) (Transaction, error) { - return Transaction{}, nil -} -func (n nopInternalClient) FinishTransaction(ctx context.Context, id string) (Transaction, error) { - return Transaction{}, nil -} -func (n nopInternalClient) Transactions(ctx context.Context) (map[string]Transaction, error) { +func (n nopInternalClient) StartTransaction(ctx context.Context, id string, timeout time.Duration, exclusive bool) (*Transaction, error) { return nil, nil } -func (n nopInternalClient) GetTransaction(ctx context.Context, id string) (Transaction, error) { - return Transaction{}, nil +func (n nopInternalClient) FinishTransaction(ctx context.Context, id string) (*Transaction, error) { + return nil, nil +} +func (n nopInternalClient) Transactions(ctx context.Context) (map[string]*Transaction, error) { + return nil, nil +} +func (n nopInternalClient) GetTransaction(ctx context.Context, id string) (*Transaction, error) { + return nil, nil } diff --git a/cluster.go b/cluster.go index d76365943..ae7c2106c 100644 --- a/cluster.go +++ b/cluster.go @@ -2598,6 +2598,6 @@ const ( ) type TransactionMessage struct { - Transaction Transaction + Transaction *Transaction Action string } diff --git a/encoding/proto/proto.go b/encoding/proto/proto.go index a08e3de28..678b2d04c 100644 --- a/encoding/proto/proto.go +++ b/encoding/proto/proto.go @@ -847,7 +847,10 @@ func encodeTransactionMessage(msg *pilosa.TransactionMessage) *internal.Transact } } -func encodeTransaction(trns pilosa.Transaction) *internal.Transaction { +func encodeTransaction(trns *pilosa.Transaction) *internal.Transaction { + if trns == nil { + return nil + } return &internal.Transaction{ ID: trns.ID, Active: trns.Active, @@ -1240,10 +1243,17 @@ func decodeTranslateIDsResponse(pb *internal.TranslateIDsResponse, m *pilosa.Tra func decodeTransactionMessage(pb *internal.TransactionMessage, m *pilosa.TransactionMessage) { m.Action = pb.Action - decodeTransaction(pb.Transaction, &m.Transaction) + if pb.Transaction == nil { + m.Transaction = nil + return + } else if m.Transaction == nil { + m.Transaction = &pilosa.Transaction{} + } + decodeTransaction(pb.Transaction, m.Transaction) } func decodeTransaction(pb *internal.Transaction, trns *pilosa.Transaction) { + trns.ID = pb.ID trns.Active = pb.Active trns.Exclusive = pb.Exclusive diff --git a/holder.go b/holder.go index 261e5fbe7..ad52ee082 100644 --- a/holder.go +++ b/holder.go @@ -104,19 +104,19 @@ type Holder struct { opening bool } -func (h *Holder) StartTransaction(ctx context.Context, id string, timeout time.Duration, exclusive bool) (Transaction, error) { +func (h *Holder) StartTransaction(ctx context.Context, id string, timeout time.Duration, exclusive bool) (*Transaction, error) { return h.transactionManager.Start(ctx, id, timeout, exclusive) } -func (h *Holder) FinishTransaction(ctx context.Context, id string) (Transaction, error) { +func (h *Holder) FinishTransaction(ctx context.Context, id string) (*Transaction, error) { return h.transactionManager.Finish(ctx, id) } -func (h *Holder) Transactions(ctx context.Context) (map[string]Transaction, error) { +func (h *Holder) Transactions(ctx context.Context) (map[string]*Transaction, error) { return h.transactionManager.List(ctx) } -func (h *Holder) GetTransaction(ctx context.Context, id string) (Transaction, error) { +func (h *Holder) GetTransaction(ctx context.Context, id string) (*Transaction, error) { return h.transactionManager.Get(ctx, id) } diff --git a/http/client.go b/http/client.go index aeac62485..2b112e9c5 100644 --- a/http/client.go +++ b/http/client.go @@ -1235,49 +1235,41 @@ func (c *InternalClient) TranslateIDsNode(ctx context.Context, uri *pilosa.URI, return tkresp.Keys, nil } -func (c *InternalClient) Transactions(ctx context.Context) (map[string]pilosa.Transaction, error) { +func (c *InternalClient) Transactions(ctx context.Context) (map[string]*pilosa.Transaction, error) { span, ctx := tracing.StartSpanFromContext(ctx, "InternalClient.Transactions") defer span.Finish() - trnsMap := make(map[string]pilosa.Transaction) - u := uriPathToURL(c.defaultURI, "/transactions") req, err := http.NewRequest("GET", u.String(), nil) if err != nil { - return trnsMap, errors.Wrap(err, "creating transactions request") + return nil, errors.Wrap(err, "creating transactions request") } req.Header.Set("Accept", "application/json") req.Header.Set("User-Agent", "pilosa/"+pilosa.Version) resp, err := c.executeRequest(req.WithContext(ctx)) if err != nil { - return trnsMap, errors.Wrap(err, "executing request") + return nil, errors.Wrap(err, "executing request") } defer func() { _, _ = io.Copy(ioutil.Discard, resp.Body) _ = resp.Body.Close() }() - tmpTrnsMap := make(map[string]*pilosa.Transaction) - err = json.NewDecoder(resp.Body).Decode(&tmpTrnsMap) - - for id, trnsp := range tmpTrnsMap { - trnsMap[id] = *trnsp - } - + trnsMap := make(map[string]*pilosa.Transaction) + err = json.NewDecoder(resp.Body).Decode(&trnsMap) return trnsMap, errors.Wrap(err, "json decoding") } -func (c *InternalClient) StartTransaction(ctx context.Context, id string, timeout time.Duration, exclusive bool) (pilosa.Transaction, error) { +func (c *InternalClient) StartTransaction(ctx context.Context, id string, timeout time.Duration, exclusive bool) (*pilosa.Transaction, error) { span, ctx := tracing.StartSpanFromContext(ctx, "InternalClient.StartTransaction") defer span.Finish() - tr := &TransactionResponse{Transaction: &pilosa.Transaction{}} buf, err := json.Marshal(&pilosa.Transaction{ ID: id, Timeout: timeout, Exclusive: exclusive, }) if err != nil { - return pilosa.Transaction{}, errors.Wrap(err, "marshalling payload") + return nil, errors.Wrap(err, "marshalling payload") } // We're using the defaultURI here because this is only used by // tests, and we want to test requests against all hosts. A robust @@ -1286,7 +1278,7 @@ func (c *InternalClient) StartTransaction(ctx context.Context, id string, timeou u := uriPathToURL(c.defaultURI, "/transaction/"+id) req, err := http.NewRequest("POST", u.String(), bytes.NewReader(buf)) if err != nil { - return pilosa.Transaction{}, errors.Wrap(err, "creating post transaction request") + return nil, errors.Wrap(err, "creating post transaction request") } req.Header.Set("Content-Length", strconv.Itoa(len(buf))) req.Header.Set("Content-Type", "application/json") @@ -1295,32 +1287,33 @@ func (c *InternalClient) StartTransaction(ctx context.Context, id string, timeou resp, err := c.executeRequest(req.WithContext(ctx), giveRawResponse(true)) if err != nil { - return pilosa.Transaction{}, errors.Wrap(err, "executing request") + return nil, errors.Wrap(err, "executing request") } defer func() { _, _ = io.Copy(ioutil.Discard, resp.Body) _ = resp.Body.Close() }() - err = json.NewDecoder(resp.Body).Decode(&tr) + tr := &TransactionResponse{} + err = json.NewDecoder(resp.Body).Decode(tr) if err != nil { - return pilosa.Transaction{}, errors.Wrap(err, "decoding response") + return nil, errors.Wrap(err, "decoding response") } if resp.StatusCode == 409 { err = pilosa.ErrTransactionExclusive } else if tr.Error != "" { err = errors.New(tr.Error) } - return *tr.Transaction, err + return tr.Transaction, err } -func (c *InternalClient) FinishTransaction(ctx context.Context, id string) (pilosa.Transaction, error) { +func (c *InternalClient) FinishTransaction(ctx context.Context, id string) (*pilosa.Transaction, error) { span, ctx := tracing.StartSpanFromContext(ctx, "InternalClient.FinishTransaction") defer span.Finish() u := uriPathToURL(c.defaultURI, "/transaction/"+id+"/finish") req, err := http.NewRequest("POST", u.String(), nil) if err != nil { - return pilosa.Transaction{}, errors.Wrap(err, "creating finish transaction request") + return nil, errors.Wrap(err, "creating finish transaction request") } req.Header.Set("Accept", "application/json") @@ -1328,25 +1321,25 @@ func (c *InternalClient) FinishTransaction(ctx context.Context, id string) (pilo resp, err := c.executeRequest(req.WithContext(ctx), giveRawResponse(true)) if err != nil { - return pilosa.Transaction{}, errors.Wrap(err, "executing request") + return nil, errors.Wrap(err, "executing request") } defer func() { _, _ = io.Copy(ioutil.Discard, resp.Body) _ = resp.Body.Close() }() - tr := &TransactionResponse{Transaction: &pilosa.Transaction{}} - err = json.NewDecoder(resp.Body).Decode(&tr) + tr := &TransactionResponse{} + err = json.NewDecoder(resp.Body).Decode(tr) if err != nil { - return pilosa.Transaction{}, errors.Wrap(err, "decoding response") + return nil, errors.Wrap(err, "decoding response") } if tr.Error != "" { err = errors.New(tr.Error) } - return *tr.Transaction, err + return tr.Transaction, err } -func (c *InternalClient) GetTransaction(ctx context.Context, id string) (pilosa.Transaction, error) { +func (c *InternalClient) GetTransaction(ctx context.Context, id string) (*pilosa.Transaction, error) { span, ctx := tracing.StartSpanFromContext(ctx, "InternalClient.GetTransaction") defer span.Finish() @@ -1357,29 +1350,29 @@ func (c *InternalClient) GetTransaction(ctx context.Context, id string) (pilosa. u := uriPathToURL(c.defaultURI, "/transaction/"+id) req, err := http.NewRequest("GET", u.String(), nil) if err != nil { - return pilosa.Transaction{}, errors.Wrap(err, "creating get transaction request") + return nil, errors.Wrap(err, "creating get transaction request") } req.Header.Set("Accept", "application/json") req.Header.Set("User-Agent", "pilosa/"+pilosa.Version) resp, err := c.executeRequest(req.WithContext(ctx), giveRawResponse(true)) if err != nil { - return pilosa.Transaction{}, errors.Wrap(err, "executing request") + return nil, errors.Wrap(err, "executing request") } defer func() { _, _ = io.Copy(ioutil.Discard, resp.Body) _ = resp.Body.Close() }() - tr := &TransactionResponse{Transaction: &pilosa.Transaction{}} - err = json.NewDecoder(resp.Body).Decode(&tr) + tr := &TransactionResponse{} + err = json.NewDecoder(resp.Body).Decode(tr) if err != nil { - return pilosa.Transaction{}, errors.Wrap(err, "decoding response") + return nil, errors.Wrap(err, "decoding response") } if tr.Error != "" { err = errors.New(tr.Error) } - return *tr.Transaction, err + return tr.Transaction, err } type executeOpts struct { diff --git a/http/client_test.go b/http/client_test.go index 7924106c9..fb5568658 100644 --- a/http/client_test.go +++ b/http/client_test.go @@ -1258,7 +1258,7 @@ func TestClientTransactions(t *testing.T) { } else { expDeadline = time.Now().Add(time.Minute) test.CompareTransactions(t, - pilosa.Transaction{ID: "blah", Timeout: time.Minute, Active: true, Deadline: expDeadline}, + &pilosa.Transaction{ID: "blah", Timeout: time.Minute, Active: true, Deadline: expDeadline}, trns) } @@ -1269,7 +1269,7 @@ func TestClientTransactions(t *testing.T) { t.Errorf("unexpected trnsMap: %+v", trnsMap) } test.CompareTransactions(t, - pilosa.Transaction{ID: "blah", Timeout: time.Minute, Active: true, Deadline: expDeadline}, + &pilosa.Transaction{ID: "blah", Timeout: time.Minute, Active: true, Deadline: expDeadline}, trnsMap["blah"]) } @@ -1277,7 +1277,7 @@ func TestClientTransactions(t *testing.T) { t.Fatalf("error getting transaction: %v", err) } else { test.CompareTransactions(t, - pilosa.Transaction{ID: "blah", Timeout: time.Minute, Active: true, Deadline: expDeadline}, + &pilosa.Transaction{ID: "blah", Timeout: time.Minute, Active: true, Deadline: expDeadline}, trns) } @@ -1285,7 +1285,7 @@ func TestClientTransactions(t *testing.T) { t.Fatalf("error finishing transaction: %v", err) } else { test.CompareTransactions(t, - pilosa.Transaction{ID: "blah", Timeout: time.Minute, Active: true, Deadline: expDeadline}, + &pilosa.Transaction{ID: "blah", Timeout: time.Minute, Active: true, Deadline: expDeadline}, trns) } @@ -1295,7 +1295,7 @@ func TestClientTransactions(t *testing.T) { } else { expDeadline = time.Now().Add(time.Minute) test.CompareTransactions(t, - pilosa.Transaction{ID: "blahe", Timeout: time.Minute, Active: true, Exclusive: true, Deadline: expDeadline}, + &pilosa.Transaction{ID: "blahe", Timeout: time.Minute, Active: true, Exclusive: true, Deadline: expDeadline}, trns) } @@ -1304,7 +1304,7 @@ func TestClientTransactions(t *testing.T) { t.Fatalf("shouldn't be able to start transaction while an exclusive is running, but got: %+v, %v", trns, err) } else { test.CompareTransactions(t, - pilosa.Transaction{ID: "blahe", Timeout: time.Minute, Active: true, Exclusive: true, Deadline: expDeadline}, + &pilosa.Transaction{ID: "blahe", Timeout: time.Minute, Active: true, Exclusive: true, Deadline: expDeadline}, trns) } @@ -1313,7 +1313,7 @@ func TestClientTransactions(t *testing.T) { t.Fatalf("error finishing transaction: %v", err) } else { test.CompareTransactions(t, - pilosa.Transaction{ID: "blahe", Timeout: time.Minute, Active: true, Exclusive: true, Deadline: expDeadline}, + &pilosa.Transaction{ID: "blahe", Timeout: time.Minute, Active: true, Exclusive: true, Deadline: expDeadline}, trns) } @@ -1323,7 +1323,7 @@ func TestClientTransactions(t *testing.T) { } else { expDeadline = time.Now().Add(time.Minute) test.CompareTransactions(t, - pilosa.Transaction{ID: "blah", Timeout: time.Minute, Active: true, Deadline: expDeadline}, + &pilosa.Transaction{ID: "blah", Timeout: time.Minute, Active: true, Deadline: expDeadline}, trns) } @@ -1333,7 +1333,7 @@ func TestClientTransactions(t *testing.T) { t.Fatalf("expected ErrTransactionExists, but got: %v", err) } else { test.CompareTransactions(t, - pilosa.Transaction{ID: "blah", Timeout: time.Minute, Active: true, Deadline: expDeadline}, + &pilosa.Transaction{ID: "blah", Timeout: time.Minute, Active: true, Deadline: expDeadline}, trns) } @@ -1343,7 +1343,7 @@ func TestClientTransactions(t *testing.T) { } else { expDeadline = time.Now().Add(time.Minute) test.CompareTransactions(t, - pilosa.Transaction{ID: "blahe", Timeout: time.Minute, Active: false, Exclusive: true, Deadline: expDeadline}, + &pilosa.Transaction{ID: "blahe", Timeout: time.Minute, Active: false, Exclusive: true, Deadline: expDeadline}, trns) } @@ -1352,7 +1352,7 @@ func TestClientTransactions(t *testing.T) { t.Fatalf("error finishing transaction: %v", err) } else { test.CompareTransactions(t, - pilosa.Transaction{ID: "blahe", Timeout: time.Minute, Active: false, Exclusive: true, Deadline: expDeadline}, + &pilosa.Transaction{ID: "blahe", Timeout: time.Minute, Active: false, Exclusive: true, Deadline: expDeadline}, trns) } @@ -1362,7 +1362,7 @@ func TestClientTransactions(t *testing.T) { t.Fatalf("unexpected error finishing nonexistent transaction: %v", err) } else { test.CompareTransactions(t, - pilosa.Transaction{}, + nil, trns) } @@ -1372,7 +1372,7 @@ func TestClientTransactions(t *testing.T) { t.Fatalf("unexpected error getting nonexistent transaction: %v", err) } else { test.CompareTransactions(t, - pilosa.Transaction{}, + nil, trns) } @@ -1382,7 +1382,7 @@ func TestClientTransactions(t *testing.T) { t.Fatalf("unexpected error starting on non-coordinator: %v", err) } else { test.CompareTransactions(t, - pilosa.Transaction{}, + nil, trns) } @@ -1395,7 +1395,7 @@ func TestClientTransactions(t *testing.T) { t.Errorf("expected generated UUID, but got '%s'", trns.ID) } test.CompareTransactions(t, - pilosa.Transaction{ID: trns.ID, Timeout: time.Minute, Active: true, Deadline: expDeadline}, + &pilosa.Transaction{ID: trns.ID, Timeout: time.Minute, Active: true, Deadline: expDeadline}, trns) } diff --git a/http/handler.go b/http/handler.go index 3c5033dda..c46f3914f 100644 --- a/http/handler.go +++ b/http/handler.go @@ -1057,15 +1057,7 @@ func (h *Handler) handleGetTransactions(w http.ResponseWriter, r *http.Request) return } - // JSON marshalling bullshit. Maybe we should just use - // *Transaction everywhere. - tmapP := make(map[string]*pilosa.Transaction) - for id, trns := range trnsMap { - trns := trns - tmapP[id] = &trns - } - - if err := json.NewEncoder(w).Encode(tmapP); err != nil { + if err := json.NewEncoder(w).Encode(trnsMap); err != nil { h.logger.Printf("encoding GetTransactions response: %s", err) } } @@ -1075,7 +1067,7 @@ type TransactionResponse struct { Error string `json:"error,omitempty"` } -func (h *Handler) doTransactionResponse(w http.ResponseWriter, err error, trns pilosa.Transaction) { +func (h *Handler) doTransactionResponse(w http.ResponseWriter, err error, trns *pilosa.Transaction) { if err != nil { switch errors.Cause(err) { case pilosa.ErrNodeNotCoordinator, pilosa.ErrTransactionExists: @@ -1094,7 +1086,7 @@ func (h *Handler) doTransactionResponse(w http.ResponseWriter, err error, trns p errString = err.Error() } err = json.NewEncoder(w).Encode( - TransactionResponse{Error: errString, Transaction: &trns}) + TransactionResponse{Error: errString, Transaction: trns}) if err != nil { h.logger.Printf("encoding transaction response: %v", err) } diff --git a/http/handler_test.go b/http/handler_test.go index d275d450e..fa9b5fed8 100644 --- a/http/handler_test.go +++ b/http/handler_test.go @@ -15,11 +15,13 @@ package http_test import ( + "encoding/json" "net" "testing" "github.com/pilosa/pilosa/v2" "github.com/pilosa/pilosa/v2/http" + "github.com/pilosa/pilosa/v2/test" ) func TestHandlerOptions(t *testing.T) { @@ -40,3 +42,36 @@ func TestHandlerOptions(t *testing.T) { t.Fatalf("expected error making handler without options, got nil") } } + +func TestMarshalUnmarshalTransactionResponse(t *testing.T) { + tests := []struct { + name string + tr *http.TransactionResponse + }{ + { + name: "nil transaction", + tr: &http.TransactionResponse{}, + }, + { + name: "empty transaction", + tr: &http.TransactionResponse{Transaction: &pilosa.Transaction{}}, + }, + } + + for _, tst := range tests { + t.Run(tst.name, func(t *testing.T) { + data, err := json.Marshal(tst.tr) + if err != nil { + t.Fatalf("marshaling: %v", err) + } + + mytr := &http.TransactionResponse{} + json.Unmarshal(data, mytr) + + if mytr.Error != tst.tr.Error { + t.Errorf("errors mismatch:exp/got \n%v\n%v", tst.tr.Error, mytr.Error) + } + test.CompareTransactions(t, tst.tr.Transaction, mytr.Transaction) + }) + } +} diff --git a/server.go b/server.go index 9621450d3..7e0690c30 100644 --- a/server.go +++ b/server.go @@ -1024,13 +1024,13 @@ func (s *Server) monitorRuntime() { } } -func (srv *Server) StartTransaction(ctx context.Context, id string, timeout time.Duration, exclusive bool, remote bool) (Transaction, error) { +func (srv *Server) StartTransaction(ctx context.Context, id string, timeout time.Duration, exclusive bool, remote bool) (*Transaction, error) { node := srv.node() if !remote && !node.IsCoordinator && len(srv.cluster.Nodes()) > 1 { - return Transaction{}, ErrNodeNotCoordinator + return nil, ErrNodeNotCoordinator } if remote && (node.IsCoordinator || len(srv.cluster.Nodes()) == 1) { - return Transaction{}, errors.New("got a remote start call to coordinator or single node cluster... shouldn't ever happen") + return nil, errors.New("got a remote start call to coordinator or single node cluster... shouldn't ever happen") } if remote { @@ -1070,13 +1070,13 @@ func (srv *Server) StartTransaction(ctx context.Context, id string, timeout time return trns, nil } -func (srv *Server) FinishTransaction(ctx context.Context, id string, remote bool) (Transaction, error) { +func (srv *Server) FinishTransaction(ctx context.Context, id string, remote bool) (*Transaction, error) { node := srv.node() if !remote && !node.IsCoordinator && len(srv.cluster.Nodes()) > 1 { - return Transaction{}, ErrNodeNotCoordinator + return nil, ErrNodeNotCoordinator } if remote && (node.IsCoordinator || len(srv.cluster.Nodes()) == 1) { - return Transaction{}, errors.New("got a remote finish call to coordinator or single node cluster... shouldn't ever happen") + return nil, errors.New("got a remote finish call to coordinator or single node cluster... shouldn't ever happen") } if remote { @@ -1099,7 +1099,7 @@ func (srv *Server) FinishTransaction(ctx context.Context, id string, remote bool return trns, nil } -func (srv *Server) Transactions(ctx context.Context) (map[string]Transaction, error) { +func (srv *Server) Transactions(ctx context.Context) (map[string]*Transaction, error) { node := srv.node() if !node.IsCoordinator && len(srv.cluster.Nodes()) > 1 { return nil, ErrNodeNotCoordinator @@ -1108,19 +1108,19 @@ func (srv *Server) Transactions(ctx context.Context) (map[string]Transaction, er return srv.holder.Transactions(ctx) } -func (srv *Server) GetTransaction(ctx context.Context, id string, remote bool) (Transaction, error) { +func (srv *Server) GetTransaction(ctx context.Context, id string, remote bool) (*Transaction, error) { node := srv.node() if !remote && !node.IsCoordinator && len(srv.cluster.Nodes()) > 1 { - return Transaction{}, ErrNodeNotCoordinator + return nil, ErrNodeNotCoordinator } if remote && (node.IsCoordinator || len(srv.cluster.Nodes()) == 1) { - return Transaction{}, errors.New("got a remote finish call to coordinator or single node cluster... shouldn't ever happen") + return nil, errors.New("got a remote finish call to coordinator or single node cluster... shouldn't ever happen") } trns, err := srv.holder.GetTransaction(ctx, id) if err != nil { - return Transaction{}, errors.Wrap(err, "getting transaction") + return nil, errors.Wrap(err, "getting transaction") } // The way a client would find out that the exclusive transaction @@ -1138,7 +1138,7 @@ func (srv *Server) GetTransaction(ctx context.Context, id string, remote bool) ( }, ) if err != nil { - return Transaction{}, errors.Wrap(err, "contacting remote hosts") + return nil, errors.Wrap(err, "contacting remote hosts") } return trns, nil } diff --git a/server/server_test.go b/server/server_test.go index 2eca8bca9..d84dbe50c 100644 --- a/server/server_test.go +++ b/server/server_test.go @@ -399,14 +399,14 @@ func TestTransactionsAPI(t *testing.T) { if trns, err := api0.StartTransaction(ctx, "a", time.Minute, false, false); err != nil { t.Errorf("couldn't start transaction: %v", err) } else { - test.CompareTransactions(t, pilosa.Transaction{ID: "a", Active: true, Timeout: time.Minute, Deadline: time.Now().Add(time.Minute)}, trns) + test.CompareTransactions(t, &pilosa.Transaction{ID: "a", Active: true, Timeout: time.Minute, Deadline: time.Now().Add(time.Minute)}, trns) } // can retrieve transaction from other nodes with remote=true if trns, err := api1.GetTransaction(ctx, "a", true); err != nil { t.Errorf("couldn't fetch transaction from other node with remote=true: %v", err) } else { - test.CompareTransactions(t, pilosa.Transaction{ID: "a", Active: true, Timeout: time.Minute, Deadline: time.Now().Add(time.Minute)}, trns) + test.CompareTransactions(t, &pilosa.Transaction{ID: "a", Active: true, Timeout: time.Minute, Deadline: time.Now().Add(time.Minute)}, trns) } // can start transaction with blank id and get uuid back @@ -418,7 +418,7 @@ func TestTransactionsAPI(t *testing.T) { if len(id) != 36 { // UUID t.Errorf("unexpected generated ID: %s", id) } - test.CompareTransactions(t, pilosa.Transaction{ID: id, Active: true, Timeout: time.Minute, Deadline: time.Now().Add(time.Minute)}, trns) + test.CompareTransactions(t, &pilosa.Transaction{ID: id, Active: true, Timeout: time.Minute, Deadline: time.Now().Add(time.Minute)}, trns) } // can't finish transaction on non-coordinator @@ -452,7 +452,7 @@ func TestTransactionsAPI(t *testing.T) { if trns, err := api0.StartTransaction(ctx, "a", time.Minute, false, false); err != nil { t.Errorf("couldn't start transaction: %v", err) } else { - test.CompareTransactions(t, pilosa.Transaction{ID: "a", Active: true, Timeout: time.Minute, Deadline: time.Now().Add(time.Minute)}, trns) + test.CompareTransactions(t, &pilosa.Transaction{ID: "a", Active: true, Timeout: time.Minute, Deadline: time.Now().Add(time.Minute)}, trns) } // can start exclusive transaction and is not immediately active @@ -471,14 +471,14 @@ func TestTransactionsAPI(t *testing.T) { if trns, err := api0.GetTransaction(ctx, "exc", false); err != nil { t.Errorf("couldn't poll exclusive transaction: %v", err) } else { - test.CompareTransactions(t, pilosa.Transaction{ID: "exc", Active: true, Exclusive: true, Timeout: time.Minute, Deadline: time.Now().Add(time.Minute)}, trns) + test.CompareTransactions(t, &pilosa.Transaction{ID: "exc", Active: true, Exclusive: true, Timeout: time.Minute, Deadline: time.Now().Add(time.Minute)}, trns) } // transaction is active on other nodes with remote=true if trns, err := api1.GetTransaction(ctx, "exc", true); err != nil { t.Errorf("couldn't poll exclusive transaction: %v", err) } else { - test.CompareTransactions(t, pilosa.Transaction{ID: "exc", Active: true, Exclusive: true, Timeout: time.Minute, Deadline: time.Now().Add(time.Minute)}, trns) + test.CompareTransactions(t, &pilosa.Transaction{ID: "exc", Active: true, Exclusive: true, Timeout: time.Minute, Deadline: time.Now().Add(time.Minute)}, trns) } // LATER, test deadline extension on non-coordinator blocks active, exclusive transaction being returned diff --git a/test/transaction.go b/test/transaction.go index 53dadfce1..90addbeae 100644 --- a/test/transaction.go +++ b/test/transaction.go @@ -26,11 +26,14 @@ const deadlineSkew = time.Millisecond * 10 // CompareTransactions errors describing how the // transactions differ (if at all). The deadlines need only be close // (within deadlineSkew). -func CompareTransactions(t *testing.T, trns1, trns2 pilosa.Transaction) { +func CompareTransactions(t *testing.T, trns1, trns2 *pilosa.Transaction) { t.Helper() if err := pilosa.CompareTransactions(trns1, trns2); err != nil { t.Errorf("%v", err) } + if trns1 == nil || trns2 == nil { + return + } diff := trns1.Deadline.Sub(trns2.Deadline) diff --git a/transaction.go b/transaction.go index 1ad3077de..3fa2e2a42 100644 --- a/transaction.go +++ b/transaction.go @@ -88,13 +88,13 @@ func NewTransactionManager(store TransactionStore) *TransactionManager { // is returned—this is primarily so that the caller can discover if an // exclusive transaction has been made immediately active or if they // need to poll. -func (tm *TransactionManager) Start(ctx context.Context, id string, timeout time.Duration, exclusive bool) (Transaction, error) { +func (tm *TransactionManager) Start(ctx context.Context, id string, timeout time.Duration, exclusive bool) (*Transaction, error) { tm.mu.Lock() defer tm.mu.Unlock() trnsMap, err := tm.store.List() if err != nil { - return Transaction{}, errors.Wrap(err, "listing transactions in Start") + return nil, errors.Wrap(err, "listing transactions in Start") } // check for an exclusive transaction @@ -123,7 +123,7 @@ func (tm *TransactionManager) Start(ctx context.Context, id string, timeout time // set deadline according to timeout deadline := time.Now().Add(timeout) - trns := Transaction{ + trns := &Transaction{ ID: id, Active: active, Exclusive: exclusive, @@ -131,7 +131,7 @@ func (tm *TransactionManager) Start(ctx context.Context, id string, timeout time Deadline: deadline, } if err = tm.store.Put(trns); err != nil { - return trns, errors.Wrap(err, "adding to store") + return nil, errors.Wrap(err, "adding to store") } // we won't check deadlines unless there's actually a @@ -145,22 +145,17 @@ func (tm *TransactionManager) Start(ctx context.Context, id string, timeout time // Finish completes and removes a transaction, returning the completed // transaction (so that the caller can e.g. view the Stats) -func (tm *TransactionManager) Finish(ctx context.Context, id string) (Transaction, error) { +func (tm *TransactionManager) Finish(ctx context.Context, id string) (*Transaction, error) { tm.mu.Lock() defer tm.mu.Unlock() return tm.finish(id) } // finish is the unprotected implementation of Finish -func (tm *TransactionManager) finish(id string) (Transaction, error) { - // sanity check - if trns, err := tm.store.Get(id); err != nil { - return trns, err - } - +func (tm *TransactionManager) finish(id string) (*Transaction, error) { trns, err := tm.store.Remove(id) if err != nil { - return trns, err + return nil, err } // After removing, check to see if we need to activate an exclusive transaction @@ -190,7 +185,7 @@ func (tm *TransactionManager) finish(id string) (Transaction, error) { // Get retrieves the transaction with the given ID. Returns ErrTransactionNotFound // if there isn't one. -func (tm *TransactionManager) Get(ctx context.Context, id string) (Transaction, error) { +func (tm *TransactionManager) Get(ctx context.Context, id string) (*Transaction, error) { tm.mu.RLock() defer tm.mu.RUnlock() @@ -199,7 +194,7 @@ func (tm *TransactionManager) Get(ctx context.Context, id string) (Transaction, // List returns map of all transactions by their ID. It is a copy and // so may be retained and modified by the caller. -func (tm *TransactionManager) List(ctx context.Context) (map[string]Transaction, error) { +func (tm *TransactionManager) List(ctx context.Context) (map[string]*Transaction, error) { tm.mu.RLock() defer tm.mu.RUnlock() return tm.store.List() @@ -208,12 +203,12 @@ func (tm *TransactionManager) List(ctx context.Context) (map[string]Transaction, // ResetDeadline updates the deadline for the transaction with the // given ID to be equal to the current time plus the transaction's // timeout. -func (tm *TransactionManager) ResetDeadline(ctx context.Context, id string) (Transaction, error) { +func (tm *TransactionManager) ResetDeadline(ctx context.Context, id string) (*Transaction, error) { tm.mu.Lock() defer tm.mu.Unlock() trns, err := tm.store.Get(id) if err != nil { - return trns, errors.Wrap(err, "getting transaction") + return nil, errors.Wrap(err, "getting transaction") } trns.Deadline = time.Now().Add(trns.Timeout) @@ -272,13 +267,10 @@ func (tm *TransactionManager) checkDeadlines() time.Duration { // track the time interval to next deadline nextInterval := time.Duration(0) for id, trns := range trnsMap { - // fmt.Printf("trns: %v", id) if !trns.Active { - // fmt.Printf(" not active\n") continue } if !now.Before(trns.Deadline) { - // fmt.Printf(" finishing\n") trnsF, err := tm.finish(id) if err != nil { tm.log().Printf("error finishing expired transaction '%s': %+v: %v", id, trnsF, err) @@ -287,7 +279,6 @@ func (tm *TransactionManager) checkDeadlines() time.Duration { } } else { interval := trns.Deadline.Sub(now) - // fmt.Printf(" getting new interval: %v, next: %v\n", interval, nextInterval) if nextInterval == 0 || interval < nextInterval { nextInterval = interval } @@ -307,13 +298,13 @@ func (tm *TransactionManager) log() logger.Logger { // Pilosa transactions must implement. type TransactionStore interface { // Put stores a new transaction or replaces an existing transaction with the given one. - Put(trns Transaction) error + Put(trns *Transaction) error // Get retrieves the transaction at id or returns ErrTransactionNotFound if there isn't one. - Get(id string) (Transaction, error) + Get(id string) (*Transaction, error) // List returns a map of all transactions by ID. The map must be safe to modify by the caller. - List() (map[string]Transaction, error) + List() (map[string]*Transaction, error) // Remove deletes the transaction from the store. It must return ErrTransactionNotFound if there isn't one. - Remove(id string) (Transaction, error) + Remove(id string) (*Transaction, error) } type OpenTransactionStoreFunc func(path string) (TransactionStore, error) @@ -326,16 +317,16 @@ func OpenInMemTransactionStore(path string) (TransactionStore, error) { // useful for testing. type InMemTransactionStore struct { mu sync.RWMutex - tmap map[string]Transaction + tmap map[string]*Transaction } func NewInMemTransactionStore() *InMemTransactionStore { return &InMemTransactionStore{ - tmap: make(map[string]Transaction), + tmap: make(map[string]*Transaction), } } -func (s *InMemTransactionStore) Put(trns Transaction) error { +func (s *InMemTransactionStore) Put(trns *Transaction) error { s.mu.Lock() defer s.mu.Unlock() @@ -343,25 +334,25 @@ func (s *InMemTransactionStore) Put(trns Transaction) error { return nil } -func (s *InMemTransactionStore) Get(id string) (Transaction, error) { +func (s *InMemTransactionStore) Get(id string) (*Transaction, error) { s.mu.RLock() defer s.mu.RUnlock() if trns, ok := s.tmap[id]; ok { return trns, nil } - return Transaction{}, ErrTransactionNotFound + return nil, ErrTransactionNotFound } -func (s *InMemTransactionStore) List() (map[string]Transaction, error) { - cp := make(map[string]Transaction) +func (s *InMemTransactionStore) List() (map[string]*Transaction, error) { + cp := make(map[string]*Transaction) for id, trns := range s.tmap { cp[id] = trns } return cp, nil } -func (s *InMemTransactionStore) Remove(id string) (Transaction, error) { +func (s *InMemTransactionStore) Remove(id string) (*Transaction, error) { s.mu.Lock() defer s.mu.Unlock() @@ -369,7 +360,7 @@ func (s *InMemTransactionStore) Remove(id string) (Transaction, error) { delete(s.tmap, id) return trns, nil } - return Transaction{}, ErrTransactionNotFound + return nil, ErrTransactionNotFound } type Error string @@ -380,7 +371,13 @@ const ErrTransactionNotFound = Error("transaction not found") const ErrTransactionExclusive = Error("there is an exclusive transaction, try later") const ErrTransactionExists = Error("transaction with the given id already exists") -func CompareTransactions(t1, t2 Transaction) error { +func CompareTransactions(t1, t2 *Transaction) error { + if t1 == nil && t2 == nil { + return nil + } + if t1 == nil || t2 == nil { + return errors.Errorf("transactions are not equal: %+v %+v", t1, t2) + } if t1.ID != t2.ID { return errors.Errorf("transaction IDs not equal: %+v %+v", t1, t2) } diff --git a/transaction_test.go b/transaction_test.go index 704288a40..a05c220fb 100644 --- a/transaction_test.go +++ b/transaction_test.go @@ -37,11 +37,11 @@ func TestTransactionManager(t *testing.T) { // can add a non-exclusive transaction trns1 := mustStart(t, tm, "a", time.Microsecond, false) - test.CompareTransactions(t, pilosa.Transaction{ID: "a", Active: true, Timeout: time.Microsecond, Deadline: time.Now()}, trns1) + test.CompareTransactions(t, &pilosa.Transaction{ID: "a", Active: true, Timeout: time.Microsecond, Deadline: time.Now()}, trns1) // can have two non exclusive transactions trns2 := mustStart(t, tm, "b", time.Microsecond, false) - test.CompareTransactions(t, pilosa.Transaction{ID: "b", Active: true, Timeout: time.Microsecond, Deadline: time.Now()}, trns2) + test.CompareTransactions(t, &pilosa.Transaction{ID: "b", Active: true, Timeout: time.Microsecond, Deadline: time.Now()}, trns2) // trying to start a transaction with same name errors and returns previous transaction t3, err := tm.Start(ctx, "a", time.Second, true) @@ -64,7 +64,7 @@ func TestTransactionManager(t *testing.T) { // can submit an exclusive transaction trnsE := mustStart(t, tm, "ce", time.Millisecond*5, true) - test.CompareTransactions(t, pilosa.Transaction{ID: "ce", Active: false, Exclusive: true, Timeout: time.Millisecond * 5, Deadline: time.Now().Add(time.Millisecond * 5)}, trnsE) + test.CompareTransactions(t, &pilosa.Transaction{ID: "ce", Active: false, Exclusive: true, Timeout: time.Millisecond * 5, Deadline: time.Now().Add(time.Millisecond * 5)}, trnsE) // can't start new transactions while an exclusive transaction is pending if _, err := tm.Start(ctx, "d", time.Millisecond, false); err != pilosa.ErrTransactionExclusive { @@ -118,7 +118,7 @@ func TestTransactionManager(t *testing.T) { // can start a new exclusive transaction and it's immediately active trnsHE := mustStart(t, tm, "he", time.Hour, true) - test.CompareTransactions(t, pilosa.Transaction{ID: "he", Active: true, Exclusive: true, Timeout: time.Hour, Deadline: time.Now().Add(time.Hour)}, trnsHE) + test.CompareTransactions(t, &pilosa.Transaction{ID: "he", Active: true, Exclusive: true, Timeout: time.Hour, Deadline: time.Now().Add(time.Hour)}, trnsHE) // can't start new transactions while an exclusive transaction is active if _, err := tm.Start(ctx, "i", time.Millisecond, false); err != pilosa.ErrTransactionExclusive { @@ -131,7 +131,7 @@ func TestTransactionManager(t *testing.T) { // can start normal transaction after finishing exclusive transaction trnsJ := mustStart(t, tm, "j", time.Hour, false) - test.CompareTransactions(t, pilosa.Transaction{ID: "j", Active: true, Timeout: time.Hour, Deadline: time.Now().Add(time.Hour)}, trnsJ) + test.CompareTransactions(t, &pilosa.Transaction{ID: "j", Active: true, Timeout: time.Hour, Deadline: time.Now().Add(time.Hour)}, trnsJ) // can finish normal transaction trnsJ_finish := mustFinish(t, tm, "j") @@ -139,11 +139,11 @@ func TestTransactionManager(t *testing.T) { // can start normal transaction after finishing normal transaction trnsK := mustStart(t, tm, "k", time.Hour, false) - test.CompareTransactions(t, pilosa.Transaction{ID: "k", Active: true, Timeout: time.Hour, Deadline: time.Now().Add(time.Hour)}, trnsK) + test.CompareTransactions(t, &pilosa.Transaction{ID: "k", Active: true, Timeout: time.Hour, Deadline: time.Now().Add(time.Hour)}, trnsK) // can start new exclusive transaction, but not immediately active trnsLE := mustStart(t, tm, "le", time.Hour, true) - test.CompareTransactions(t, pilosa.Transaction{ID: "le", Exclusive: true, Timeout: time.Hour, Deadline: time.Now().Add(time.Hour)}, trnsLE) + test.CompareTransactions(t, &pilosa.Transaction{ID: "le", Exclusive: true, Timeout: time.Hour, Deadline: time.Now().Add(time.Hour)}, trnsLE) // finishing k should activate le trnsK_finish := mustFinish(t, tm, "k") @@ -156,11 +156,11 @@ func TestTransactionManager(t *testing.T) { // can start normal transaction to test deadline reset trnsM := mustStart(t, tm, "m", time.Millisecond*4, false) - test.CompareTransactions(t, pilosa.Transaction{ID: "m", Active: true, Timeout: time.Millisecond * 4, Deadline: time.Now().Add(time.Millisecond * 4)}, trnsM) + test.CompareTransactions(t, &pilosa.Transaction{ID: "m", Active: true, Timeout: time.Millisecond * 4, Deadline: time.Now().Add(time.Millisecond * 4)}, trnsM) // start new exclusive transaction to trigger deadline check trnsNE := mustStart(t, tm, "ne", time.Hour, true) - test.CompareTransactions(t, pilosa.Transaction{ID: "ne", Exclusive: true, Timeout: time.Hour, Deadline: time.Now().Add(time.Hour)}, trnsNE) + test.CompareTransactions(t, &pilosa.Transaction{ID: "ne", Exclusive: true, Timeout: time.Hour, Deadline: time.Now().Add(time.Hour)}, trnsNE) // sleep for most of the deadline time.Sleep(time.Millisecond * 3) @@ -182,7 +182,7 @@ func TestTransactionManager(t *testing.T) { } -func mustStart(t *testing.T, tm *pilosa.TransactionManager, id string, timeout time.Duration, exclusive bool) pilosa.Transaction { +func mustStart(t *testing.T, tm *pilosa.TransactionManager, id string, timeout time.Duration, exclusive bool) *pilosa.Transaction { t.Helper() trns, err := tm.Start(context.Background(), id, timeout, exclusive) if err != nil { @@ -191,7 +191,7 @@ func mustStart(t *testing.T, tm *pilosa.TransactionManager, id string, timeout t return trns } -func mustFinish(t *testing.T, tm *pilosa.TransactionManager, id string) pilosa.Transaction { +func mustFinish(t *testing.T, tm *pilosa.TransactionManager, id string) *pilosa.Transaction { t.Helper() trns, err := tm.Finish(context.Background(), id) if err != nil { @@ -200,7 +200,7 @@ func mustFinish(t *testing.T, tm *pilosa.TransactionManager, id string) pilosa.T return trns } -func mustGet(t *testing.T, tm *pilosa.TransactionManager, id string) pilosa.Transaction { +func mustGet(t *testing.T, tm *pilosa.TransactionManager, id string) *pilosa.Transaction { t.Helper() trns, err := tm.Get(context.Background(), id) if err != nil { @@ -209,7 +209,7 @@ func mustGet(t *testing.T, tm *pilosa.TransactionManager, id string) pilosa.Tran return trns } -func mustList(t *testing.T, tm *pilosa.TransactionManager) map[string]pilosa.Transaction { +func mustList(t *testing.T, tm *pilosa.TransactionManager) map[string]*pilosa.Transaction { t.Helper() trnsMap, err := tm.List(context.Background()) if err != nil { @@ -221,7 +221,7 @@ func mustList(t *testing.T, tm *pilosa.TransactionManager) map[string]pilosa.Tra func TestInMemTransactionStore(t *testing.T) { ims := pilosa.NewInMemTransactionStore() - err := ims.Put(pilosa.Transaction{ID: "blah", Timeout: time.Second}) + err := ims.Put(&pilosa.Transaction{ID: "blah", Timeout: time.Second}) if err != nil { t.Fatalf("adding blah: %v", err) } @@ -255,14 +255,15 @@ func TestInMemTransactionStore(t *testing.T) { func TestMarshalUnmarshalTransaction(t *testing.T) { tests := []struct { name string - transaction pilosa.Transaction + transaction *pilosa.Transaction }{ { - name: "empty", + name: "empty", + transaction: &pilosa.Transaction{}, }, { name: "basic", - transaction: pilosa.Transaction{ + transaction: &pilosa.Transaction{ ID: "blah", Active: true, Exclusive: true, @@ -274,7 +275,7 @@ func TestMarshalUnmarshalTransaction(t *testing.T) { for _, tst := range tests { t.Run(tst.name, func(t *testing.T) { - bytes, err := json.Marshal(&tst.transaction) + bytes, err := json.Marshal(tst.transaction) if err != nil { t.Errorf("marshalling: %v", err) } @@ -285,7 +286,7 @@ func TestMarshalUnmarshalTransaction(t *testing.T) { t.Fatalf("unmarshalling: %v", err) } - test.CompareTransactions(t, tst.transaction, *nt) + test.CompareTransactions(t, tst.transaction, nt) }) } } @@ -294,21 +295,22 @@ func TestUnmarshalTransaction(t *testing.T) { tests := []struct { name string transactionJSON string - exp pilosa.Transaction + exp *pilosa.Transaction }{ { name: "empty", transactionJSON: `{}`, + exp: &pilosa.Transaction{}, }, { name: "basicPost", transactionJSON: `{"id": "blah", "exclusive": false, "timeout": "1m"}`, - exp: pilosa.Transaction{ID: "blah", Timeout: time.Minute}, + exp: &pilosa.Transaction{ID: "blah", Timeout: time.Minute}, }, { name: "basicPostFloatTimeout", transactionJSON: `{"id": "blah", "exclusive": false, "timeout": 10.5}`, - exp: pilosa.Transaction{ID: "blah", Timeout: time.Second*10 + time.Second/2}, + exp: &pilosa.Transaction{ID: "blah", Timeout: time.Second*10 + time.Second/2}, }, } @@ -320,7 +322,7 @@ func TestUnmarshalTransaction(t *testing.T) { t.Fatalf("unmarshalling: %v", err) } - test.CompareTransactions(t, tst.exp, *nt) + test.CompareTransactions(t, tst.exp, nt) }) } }