package mysql
import (
"context"
"database/sql"
"encoding/base64"
"fmt"
"runtime/debug"
"strings"
"sync"
"github.com/google/trillian"
"github.com/google/trillian/storage/cache"
"github.com/google/trillian/storage/mysql/mysqlpb"
"github.com/google/trillian/storage/storagepb"
"github.com/google/trillian/storage/tree"
"google.golang.org/protobuf/proto"
"google.golang.org/protobuf/types/known/anypb"
"k8s.io/klog/v2"
)
const (
insertSubtreeMultiSQL = `INSERT INTO Subtree(TreeId, SubtreeId, Nodes, SubtreeRevision) ` + placeholderSQL + ` ON DUPLICATE KEY UPDATE Nodes=VALUES(Nodes)`
insertTreeHeadSQL = `INSERT INTO TreeHead(TreeId,TreeHeadTimestamp,TreeSize,RootHash,TreeRevision,RootSignature)
VALUES(?,?,?,?,?,?)`
selectSubtreeSQL = `
SELECT x.SubtreeId, Subtree.Nodes
FROM (
SELECT n.TreeId, n.SubtreeId, max(n.SubtreeRevision) AS MaxRevision
FROM Subtree n
WHERE n.SubtreeId IN (` + placeholderSQL + `) AND
n.TreeId = ? AND n.SubtreeRevision <= ?
GROUP BY n.TreeId, n.SubtreeId
) AS x
INNER JOIN Subtree
ON Subtree.SubtreeId = x.SubtreeId
AND Subtree.SubtreeRevision = x.MaxRevision
AND Subtree.TreeId = x.TreeId
AND Subtree.TreeId = ?`
selectSubtreeSQLNoRev = `
SELECT SubtreeId, Subtree.Nodes
FROM Subtree
WHERE Subtree.TreeId = ?
AND SubtreeId IN (` + placeholderSQL + `)`
placeholderSQL = "<placeholder>"
)
type mySQLTreeStorage struct {
db *sql.DB
statementMutex sync.Mutex
statements map[string]map[int]*sql.Stmt
}
func OpenDB(dbURL string) (*sql.DB, error) {
db, err := sql.Open("mysql", dbURL)
if err != nil {
klog.Warningf("Could not open MySQL database, check config: %s", err)
return nil, err
}
if _, err := db.ExecContext(context.TODO(), "SET sql_mode = 'STRICT_ALL_TABLES'"); err != nil {
klog.Warningf("Failed to set strict mode on mysql db: %s", err)
return nil, err
}
return db, nil
}
func newTreeStorage(db *sql.DB) *mySQLTreeStorage {
return &mySQLTreeStorage{
db: db,
statements: make(map[string]map[int]*sql.Stmt),
}
}
func expandPlaceholderSQL(sql string, num int, first, rest string) string {
if num <= 0 {
panic(fmt.Errorf("trying to expand SQL placeholder with <= 0 parameters: %s", sql))
}
parameters := first + strings.Repeat(","+rest, num-1)
return strings.Replace(sql, placeholderSQL, parameters, 1)
}
func (m *mySQLTreeStorage) getStmt(ctx context.Context, statement string, num int, first, rest string) (*sql.Stmt, error) {
m.statementMutex.Lock()
defer m.statementMutex.Unlock()
if m.statements[statement] != nil {
if m.statements[statement][num] != nil {
return m.statements[statement][num], nil
}
} else {
m.statements[statement] = make(map[int]*sql.Stmt)
}
s, err := m.db.PrepareContext(ctx, expandPlaceholderSQL(statement, num, first, rest))
if err != nil {
klog.Warningf("Failed to prepare statement %d: %s", num, err)
return nil, err
}
m.statements[statement][num] = s
return s, nil
}
func (m *mySQLTreeStorage) getSubtreeStmt(ctx context.Context, subtreeRevs bool, num int) (*sql.Stmt, error) {
if subtreeRevs {
return m.getStmt(ctx, selectSubtreeSQL, num, "?", "?")
} else {
return m.getStmt(ctx, selectSubtreeSQLNoRev, num, "?", "?")
}
}
func (m *mySQLTreeStorage) setSubtreeStmt(ctx context.Context, num int) (*sql.Stmt, error) {
return m.getStmt(ctx, insertSubtreeMultiSQL, num, "VALUES(?, ?, ?, ?)", "(?, ?, ?, ?)")
}
func (m *mySQLTreeStorage) beginTreeTx(ctx context.Context, tree *trillian.Tree, hashSizeBytes int, subtreeCache *cache.SubtreeCache) (treeTX, error) {
t, err := m.db.BeginTx(ctx, nil )
if err != nil {
klog.Warningf("Could not start tree TX: %s", err)
return treeTX{}, err
}
var subtreeRevisions bool
o := &mysqlpb.StorageOptions{}
if err := anypb.UnmarshalTo(tree.StorageSettings, o, proto.UnmarshalOptions{}); err != nil {
return treeTX{}, fmt.Errorf("failed to unmarshal StorageSettings: %v", err)
}
subtreeRevisions = o.SubtreeRevisions
return treeTX{
tx: t,
mu: &sync.Mutex{},
ts: m,
treeID: tree.TreeId,
treeType: tree.TreeType,
hashSizeBytes: hashSizeBytes,
subtreeCache: subtreeCache,
writeRevision: -1,
subtreeRevs: subtreeRevisions,
}, nil
}
type treeTX struct {
mu *sync.Mutex
closed bool
tx *sql.Tx
ts *mySQLTreeStorage
treeID int64
treeType trillian.TreeType
hashSizeBytes int
subtreeCache *cache.SubtreeCache
writeRevision int64
subtreeRevs bool
}
func (t *treeTX) getSubtrees(ctx context.Context, treeRevision int64, ids [][]byte) ([]*storagepb.SubtreeProto, error) {
klog.V(2).Infof("getSubtrees(len(ids)=%d)", len(ids))
klog.V(4).Infof("getSubtrees(")
if len(ids) == 0 {
return nil, nil
}
tmpl, err := t.ts.getSubtreeStmt(ctx, t.subtreeRevs, len(ids))
if err != nil {
return nil, err
}
stx := t.tx.StmtContext(ctx, tmpl)
defer func() {
if err := stx.Close(); err != nil {
klog.Errorf("stx.Close(): %v", err)
}
}()
var args []interface{}
if t.subtreeRevs {
args = make([]interface{}, 0, len(ids)+3)
for _, id := range ids {
klog.V(4).Infof(" id: %x", id)
args = append(args, id)
}
args = append(args, t.treeID)
args = append(args, treeRevision)
args = append(args, t.treeID)
} else {
args = make([]interface{}, 0, len(ids)+1)
args = append(args, t.treeID)
for _, id := range ids {
klog.V(4).Infof(" id: %x", id)
args = append(args, id)
}
}
rows, err := stx.QueryContext(ctx, args...)
if err != nil {
klog.Warningf("Failed to get merkle subtrees: %s", err)
return nil, err
}
defer func() {
if err := rows.Close(); err != nil {
klog.Errorf("rows.Close(): %v", err)
}
}()
if rows.Err() != nil {
klog.Warningf("Nothing from DB: %s", rows.Err())
return nil, rows.Err()
}
ret := make([]*storagepb.SubtreeProto, 0, len(ids))
for rows.Next() {
var subtreeIDBytes []byte
var nodesRaw []byte
if err := rows.Scan(&subtreeIDBytes, &nodesRaw); err != nil {
klog.Warningf("Failed to scan merkle subtree: %s", err)
return nil, err
}
var subtree storagepb.SubtreeProto
if err := proto.Unmarshal(nodesRaw, &subtree); err != nil {
klog.Warningf("Failed to unmarshal SubtreeProto: %s", err)
return nil, err
}
if subtree.Prefix == nil {
subtree.Prefix = []byte{}
}
ret = append(ret, &subtree)
if klog.V(4).Enabled() {
klog.Infof(" subtree: NID: %x, prefix: %x, depth: %d",
subtreeIDBytes, subtree.Prefix, subtree.Depth)
for k, v := range subtree.Leaves {
b, err := base64.StdEncoding.DecodeString(k)
if err != nil {
klog.Errorf("base64.DecodeString(%v): %v", k, err)
}
klog.Infof(" %x: %x", b, v)
}
}
}
if err := rows.Err(); err != nil {
return nil, err
}
return ret, nil
}
func (t *treeTX) storeSubtrees(ctx context.Context, subtrees []*storagepb.SubtreeProto) error {
klog.V(2).Infof("storeSubtrees(len(subtrees)=%d)", len(subtrees))
if klog.V(4).Enabled() {
klog.Infof("storeSubtrees(")
for _, s := range subtrees {
klog.Infof(" prefix: %x, depth: %d", s.Prefix, s.Depth)
for k, v := range s.Leaves {
b, err := base64.StdEncoding.DecodeString(k)
if err != nil {
klog.Errorf("base64.DecodeString(%v): %v", k, err)
}
klog.Infof(" %x: %x", b, v)
}
}
}
if len(subtrees) == 0 {
return nil
}
args := make([]interface{}, 0, len(subtrees))
var subtreeRev int64
if t.subtreeRevs {
subtreeRev = t.writeRevision
}
for _, s := range subtrees {
s := s
if s.Prefix == nil {
panic(fmt.Errorf("nil prefix on %v", s))
}
subtreeBytes, err := proto.Marshal(s)
if err != nil {
return err
}
args = append(args, t.treeID)
args = append(args, s.Prefix)
args = append(args, subtreeBytes)
args = append(args, subtreeRev)
}
tmpl, err := t.ts.setSubtreeStmt(ctx, len(subtrees))
if err != nil {
return err
}
stx := t.tx.StmtContext(ctx, tmpl)
defer func() {
if err := stx.Close(); err != nil {
klog.Errorf("stx.Close(): %v", err)
}
}()
r, err := stx.ExecContext(ctx, args...)
if err != nil {
klog.Warningf("Failed to set merkle subtrees: %s", err)
return err
}
_, _ = r.RowsAffected()
return nil
}
func checkResultOkAndRowCountIs(res sql.Result, err error, count int64) error {
if err != nil {
return mysqlToGRPC(err)
}
rowsAffected, rowsError := res.RowsAffected()
if rowsError != nil {
return mysqlToGRPC(rowsError)
}
if rowsAffected != count {
return fmt.Errorf("expected %d row(s) to be affected but saw: %d", count,
rowsAffected)
}
return nil
}
func (t *treeTX) getSubtreesAtRev(ctx context.Context, rev int64) cache.GetSubtreesFunc {
return func(ids [][]byte) ([]*storagepb.SubtreeProto, error) {
return t.getSubtrees(ctx, rev, ids)
}
}
func (t *treeTX) SetMerkleNodes(ctx context.Context, nodes []tree.Node) error {
t.mu.Lock()
defer t.mu.Unlock()
rev := t.writeRevision - 1
return t.subtreeCache.SetNodes(nodes, t.getSubtreesAtRev(ctx, rev))
}
func (t *treeTX) Commit(ctx context.Context) error {
t.mu.Lock()
defer t.mu.Unlock()
if t.writeRevision > -1 {
tiles, err := t.subtreeCache.UpdatedTiles()
if err != nil {
klog.Warningf("SubtreeCache updated tiles error: %v", err)
return err
}
if err := t.storeSubtrees(ctx, tiles); err != nil {
klog.Warningf("TX commit flush error: %v", err)
return err
}
}
t.closed = true
if err := t.tx.Commit(); err != nil {
klog.Warningf("TX commit error: %s, stack:\n%s", err, string(debug.Stack()))
return err
}
return nil
}
func (t *treeTX) rollbackInternal() error {
t.closed = true
if err := t.tx.Rollback(); err != nil {
klog.Warningf("TX rollback error: %s, stack:\n%s", err, string(debug.Stack()))
return err
}
return nil
}
func (t *treeTX) Close() error {
t.mu.Lock()
defer t.mu.Unlock()
if t.closed {
return nil
}
err := t.rollbackInternal()
if err != nil {
klog.Warningf("Rollback error on Close(): %v", err)
}
return err
}