aboutsummaryrefslogtreecommitdiffstats
path: root/vendor/github.com/lunny/nodb/tx.go
blob: 5ce99db57ab779593b62edce57ed697b469f0513 (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
package nodb

import (
	"errors"
	"fmt"

	"github.com/lunny/nodb/store"
)

var (
	ErrNestTx = errors.New("nest transaction not supported")
	ErrTxDone = errors.New("Transaction has already been committed or rolled back")
)

type Tx struct {
	*DB

	tx *store.Tx

	logs [][]byte
}

func (db *DB) IsTransaction() bool {
	return db.status == DBInTransaction
}

// Begin a transaction, it will block all other write operations before calling Commit or Rollback.
// You must be very careful to prevent long-time transaction.
func (db *DB) Begin() (*Tx, error) {
	if db.IsTransaction() {
		return nil, ErrNestTx
	}

	tx := new(Tx)

	tx.DB = new(DB)
	tx.DB.l = db.l

	tx.l.wLock.Lock()

	tx.DB.sdb = db.sdb

	var err error
	tx.tx, err = db.sdb.Begin()
	if err != nil {
		tx.l.wLock.Unlock()
		return nil, err
	}

	tx.DB.bucket = tx.tx

	tx.DB.status = DBInTransaction

	tx.DB.index = db.index

	tx.DB.kvBatch = tx.newBatch()
	tx.DB.listBatch = tx.newBatch()
	tx.DB.hashBatch = tx.newBatch()
	tx.DB.zsetBatch = tx.newBatch()
	tx.DB.binBatch = tx.newBatch()
	tx.DB.setBatch = tx.newBatch()

	return tx, nil
}

func (tx *Tx) Commit() error {
	if tx.tx == nil {
		return ErrTxDone
	}

	tx.l.commitLock.Lock()
	err := tx.tx.Commit()
	tx.tx = nil

	if len(tx.logs) > 0 {
		tx.l.binlog.Log(tx.logs...)
	}

	tx.l.commitLock.Unlock()

	tx.l.wLock.Unlock()

	tx.DB.bucket = nil

	return err
}

func (tx *Tx) Rollback() error {
	if tx.tx == nil {
		return ErrTxDone
	}

	err := tx.tx.Rollback()
	tx.tx = nil

	tx.l.wLock.Unlock()
	tx.DB.bucket = nil

	return err
}

func (tx *Tx) newBatch() *batch {
	return tx.l.newBatch(tx.tx.NewWriteBatch(), &txBatchLocker{}, tx)
}

func (tx *Tx) Select(index int) error {
	if index < 0 || index >= int(MaxDBNumber) {
		return fmt.Errorf("invalid db index %d", index)
	}

	tx.DB.index = uint8(index)
	return nil
}