package cloudspanner
import (
"bytes"
"context"
"errors"
"fmt"
"sync"
"time"
"cloud.google.com/go/spanner"
"github.com/google/trillian"
"github.com/google/trillian/storage"
"github.com/google/trillian/storage/cache"
"github.com/google/trillian/storage/cloudspanner/spannerpb"
"github.com/google/trillian/storage/storagepb"
"github.com/google/trillian/storage/tree"
"github.com/transparency-dev/merkle/compact"
"golang.org/x/sync/errgroup"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
"google.golang.org/protobuf/proto"
"k8s.io/klog/v2"
)
var (
ErrNotFound = status.Errorf(codes.NotFound, "not found")
ErrNotImplemented = errors.New("not implemented")
ErrTransactionClosed = errors.New("transaction is closed")
ErrWrongTXType = errors.New("mutating method called on read-only transaction")
)
const (
subtreeTbl = "SubtreeData"
colSubtree = "Subtree"
colSubtreeID = "SubtreeID"
colTreeID = "TreeID"
colRevision = "Revision"
)
type treeStorage struct {
admin storage.AdminStorage
opts TreeStorageOptions
client *spanner.Client
}
type TreeStorageOptions struct {
ReadOnlyStaleness time.Duration
}
func newTreeStorageWithOpts(client *spanner.Client, opts TreeStorageOptions) *treeStorage {
return &treeStorage{client: client, admin: nil, opts: opts}
}
type spanRead interface {
Query(context.Context, spanner.Statement) *spanner.RowIterator
Read(ctx context.Context, table string, keys spanner.KeySet, columns []string) *spanner.RowIterator
ReadUsingIndex(ctx context.Context, table, index string, keys spanner.KeySet, columns []string) *spanner.RowIterator
ReadRow(ctx context.Context, table string, key spanner.Key, columns []string) (*spanner.Row, error)
ReadWithOptions(ctx context.Context, table string, keys spanner.KeySet, columns []string, opts *spanner.ReadOptions) (ri *spanner.RowIterator)
}
func (t *treeStorage) latestSTH(ctx context.Context, stx spanRead, treeID int64) (*spannerpb.TreeHead, error) {
query := spanner.NewStatement(
"SELECT TreeID, TimestampNanos, TreeSize, RootHash, RootSignature, TreeRevision, TreeMetadata FROM TreeHeads" +
" WHERE TreeID = @tree_id" +
" ORDER BY TreeRevision DESC " +
" LIMIT 1")
query.Params["tree_id"] = treeID
var th *spannerpb.TreeHead
rows := stx.Query(ctx, query)
defer rows.Stop()
err := rows.Do(func(r *spanner.Row) error {
tth := &spannerpb.TreeHead{}
if err := r.Columns(&tth.TreeId, &tth.TsNanos, &tth.TreeSize, &tth.RootHash, &tth.Signature, &tth.TreeRevision, &tth.Metadata); err != nil {
return err
}
th = tth
return nil
})
if err != nil {
return nil, err
}
if th == nil {
klog.Warningf("no head found for treeID %v", treeID)
return nil, storage.ErrTreeNeedsInit
}
return th, nil
}
type newCacheFn func(*trillian.Tree) (*cache.SubtreeCache, error)
func (t *treeStorage) getTreeAndConfig(ctx context.Context, tree *trillian.Tree) (*trillian.Tree, proto.Message, error) {
config, err := unmarshalSettings(tree)
if err != nil {
return nil, nil, err
}
return tree, config, nil
}
func (t *treeStorage) begin(ctx context.Context, tree *trillian.Tree, newCache newCacheFn, stx spanRead) (*treeTX, error) {
tree, config, err := t.getTreeAndConfig(ctx, tree)
if err != nil {
return nil, err
}
subtreeCache, err := newCache(tree)
if err != nil {
return nil, err
}
treeTX := &treeTX{
treeID: tree.TreeId,
treeType: tree.TreeType,
ts: t,
stx: stx,
cache: subtreeCache,
config: config,
_writeRev: -1,
}
return treeTX, nil
}
func (t *treeTX) getLatestRoot(ctx context.Context) error {
t.getLatestRootOnce.Do(func() {
t._currentSTH, t._currentSTHErr = t.ts.latestSTH(ctx, t.stx, t.treeID)
if t._currentSTH != nil {
t._writeRev = t._currentSTH.TreeRevision + 1
}
})
return t._currentSTHErr
}
type treeTX struct {
treeID int64
treeType trillian.TreeType
ts *treeStorage
mu sync.RWMutex
stx spanRead
config proto.Message
_currentSTH *spannerpb.TreeHead
_currentSTHErr error
_writeRev int64
cache *cache.SubtreeCache
getLatestRootOnce sync.Once
}
func (t *treeTX) currentSTH(ctx context.Context) (*spannerpb.TreeHead, error) {
if err := t.getLatestRoot(ctx); err != nil {
return nil, err
}
return t._currentSTH, nil
}
func (t *treeTX) writeRev(ctx context.Context) (int64, error) {
if err := t.getLatestRoot(ctx); err == storage.ErrTreeNeedsInit {
return 0, nil
} else if err != nil {
return -1, fmt.Errorf("writeRev(): %v", err)
}
return t._writeRev, nil
}
func (t *treeTX) storeSubtrees(ctx context.Context, sts []*storagepb.SubtreeProto) error {
stx, ok := t.stx.(*spanner.ReadWriteTransaction)
if !ok {
return ErrWrongTXType
}
for _, st := range sts {
if st == nil {
continue
}
stBytes, err := proto.Marshal(st)
if err != nil {
return err
}
m := spanner.Insert(
subtreeTbl,
[]string{colTreeID, colSubtreeID, colRevision, colSubtree},
[]interface{}{t.treeID, st.Prefix, t._writeRev, stBytes},
)
if err := stx.BufferWrite([]*spanner.Mutation{m}); err != nil {
return err
}
}
return nil
}
func (t *treeTX) flushSubtrees(ctx context.Context) error {
tiles, err := t.cache.UpdatedTiles()
if err != nil {
return err
}
return t.storeSubtrees(ctx, tiles)
}
func (t *treeTX) Commit(ctx context.Context) error {
t.mu.Lock()
defer func() {
t.stx = nil
t.mu.Unlock()
}()
if t.stx == nil {
return ErrTransactionClosed
}
switch stx := t.stx.(type) {
case *spanner.ReadOnlyTransaction:
klog.V(1).Infof("Closed readonly tx %p", stx)
stx.Close()
return nil
case *spanner.ReadWriteTransaction:
return t.flushSubtrees(ctx)
default:
return fmt.Errorf("internal error: unknown transaction type %T", stx)
}
}
func (t *treeTX) Close() error {
t.mu.Lock()
defer t.mu.Unlock()
if t.stx == nil {
return ErrTransactionClosed
}
if stx, ok := t.stx.(*spanner.ReadOnlyTransaction); ok {
klog.V(1).Infof("Closed snapshot %p", stx)
stx.Close()
}
return nil
}
func (t *treeTX) readRevision(ctx context.Context) (int64, error) {
sth, err := t.currentSTH(ctx)
if err != nil {
return -1, err
}
return sth.TreeRevision, nil
}
func (t *treeTX) getSubtree(ctx context.Context, rev int64, id []byte) (p *storagepb.SubtreeProto, e error) {
var ret *storagepb.SubtreeProto
stmt := spanner.NewStatement(
"SELECT Revision, Subtree FROM SubtreeData" +
" WHERE TreeID = @tree_id" +
" AND SubtreeID = @subtree_id" +
" AND Revision <= @revision" +
" ORDER BY Revision DESC" +
" LIMIT 1")
stmt.Params["tree_id"] = t.treeID
stmt.Params["subtree_id"] = id
stmt.Params["revision"] = rev
rows := t.stx.Query(ctx, stmt)
err := rows.Do(func(r *spanner.Row) error {
if ret != nil {
return nil
}
var rRev int64
var st storagepb.SubtreeProto
stBytes := make([]byte, 1<<20)
if err := r.Columns(&rRev, &stBytes); err != nil {
return err
}
if err := proto.Unmarshal(stBytes, &st); err != nil {
return err
}
if rRev > rev {
return fmt.Errorf("got subtree with too new a revision %d, want %d", rRev, rev)
}
if got, want := id, st.Prefix; !bytes.Equal(got, want) {
return fmt.Errorf("got subtree with prefix %v, wanted %v", got, want)
}
if got, want := rRev, rev; got > rev {
return fmt.Errorf("got subtree rev %d, wanted <= %d", got, want)
}
ret = &st
if st.Prefix == nil && len(id) == 0 {
st.Prefix = []byte{}
}
return nil
})
return ret, err
}
func (t *treeTX) GetMerkleNodes(ctx context.Context, ids []compact.NodeID) ([]tree.Node, error) {
t.mu.RLock()
defer t.mu.RUnlock()
if t.stx == nil {
return nil, ErrTransactionClosed
}
rev, err := t.readRevision(ctx)
if err != nil {
return nil, fmt.Errorf("failed to get read revision: %v", err)
}
return t.cache.GetNodes(ids, t.getSubtreesAtRev(ctx, rev))
}
func (t *treeTX) getSubtreesAtRev(ctx context.Context, rev int64) cache.GetSubtreesFunc {
return func(ids [][]byte) ([]*storagepb.SubtreeProto, error) {
c := make(chan *storagepb.SubtreeProto, len(ids))
g, gctx := errgroup.WithContext(ctx)
for _, id := range ids {
id := id
g.Go(func() error {
st, err := t.getSubtree(gctx, rev, id)
if err != nil {
return err
}
c <- st
return nil
})
}
if err := g.Wait(); err != nil {
return nil, err
}
close(c)
ret := make([]*storagepb.SubtreeProto, 0, len(ids))
for st := range c {
if st != nil {
ret = append(ret, st)
}
}
return ret, nil
}
}
func (t *treeTX) SetMerkleNodes(ctx context.Context, nodes []tree.Node) error {
t.mu.RLock()
defer t.mu.RUnlock()
if t.stx == nil {
return ErrTransactionClosed
}
writeRev, err := t.writeRev(ctx)
if err != nil {
return err
}
return t.cache.SetNodes(nodes, t.getSubtreesAtRev(ctx, writeRev-1))
}
func checkDatabaseAccessible(ctx context.Context, client *spanner.Client) error {
stmt := spanner.NewStatement("SELECT 1")
rows := client.Single().Query(ctx, stmt)
defer rows.Stop()
return rows.Do(func(row *spanner.Row) error { return nil })
}