package mysql
import (
"bytes"
"context"
"database/sql"
"encoding/gob"
"fmt"
"testing"
"github.com/google/trillian"
"github.com/google/trillian/storage"
"github.com/google/trillian/storage/mysql/mysqlpb"
"github.com/google/trillian/storage/testonly"
"google.golang.org/protobuf/proto"
"google.golang.org/protobuf/types/known/anypb"
)
const selectTreeControlByID = "SELECT SigningEnabled, SequencingEnabled, SequenceIntervalSeconds FROM TreeControl WHERE TreeId = ?"
func TestMysqlAdminStorage(t *testing.T) {
tester := &testonly.AdminStorageTester{NewAdminStorage: func() storage.AdminStorage {
cleanTestDB(DB)
return NewAdminStorage(DB)
}}
tester.RunAllTests(t)
}
func TestAdminTX_CreateTree_InitializesStorageStructures(t *testing.T) {
cleanTestDB(DB)
s := NewAdminStorage(DB)
ctx := context.Background()
tree, err := storage.CreateTree(ctx, s, testonly.LogTree)
if err != nil {
t.Fatalf("CreateTree() failed: %v", err)
}
var signingEnabled, sequencingEnabled bool
var sequenceIntervalSeconds int
if err := DB.QueryRowContext(ctx, selectTreeControlByID, tree.TreeId).Scan(&signingEnabled, &sequencingEnabled, &sequenceIntervalSeconds); err != nil {
t.Fatalf("Failed to read TreeControl: %v", err)
}
if sequenceIntervalSeconds <= 0 {
t.Errorf("sequenceIntervalSeconds = %v, want > 0", sequenceIntervalSeconds)
}
}
func TestCreateTreeInvalidStates(t *testing.T) {
cleanTestDB(DB)
s := NewAdminStorage(DB)
ctx := context.Background()
states := []trillian.TreeState{trillian.TreeState_DRAINING, trillian.TreeState_FROZEN}
for _, state := range states {
inTree := proto.Clone(testonly.LogTree).(*trillian.Tree)
inTree.TreeState = state
if _, err := storage.CreateTree(ctx, s, inTree); err == nil {
t.Errorf("CreateTree() state: %v got: nil want: err", state)
}
}
}
func TestAdminTX_TreeWithNulls(t *testing.T) {
cleanTestDB(DB)
s := NewAdminStorage(DB)
ctx := context.Background()
tree, err := storage.CreateTree(ctx, s, testonly.LogTree)
if err != nil {
t.Fatalf("CreateTree() failed: %v", err)
}
treeID := tree.TreeId
if err := setNulls(ctx, DB, treeID); err != nil {
t.Fatalf("setNulls() = %v, want = nil", err)
}
tests := []struct {
desc string
fn storage.AdminTXFunc
}{
{
desc: "GetTree",
fn: func(ctx context.Context, tx storage.AdminTX) error {
_, err := tx.GetTree(ctx, treeID)
return err
},
},
{
desc: "ListTrees",
fn: func(ctx context.Context, tx storage.AdminTX) error {
trees, err := tx.ListTrees(ctx, false )
if err != nil {
return err
}
for _, tree := range trees {
if tree.TreeId == treeID {
return nil
}
}
return fmt.Errorf("ID not found: %v", treeID)
},
},
}
for _, test := range tests {
if err := s.ReadWriteTransaction(ctx, test.fn); err != nil {
t.Errorf("%v: err = %v, want = nil", test.desc, err)
}
}
}
func TestAdminTX_StorageSettings(t *testing.T) {
cleanTestDB(DB)
s := NewAdminStorage(DB)
ctx := context.Background()
badSettings, err := anypb.New(&trillian.Tree{})
if err != nil {
t.Fatalf("Error marshaling proto: %v", err)
}
goodSettings, err := anypb.New(&mysqlpb.StorageOptions{})
if err != nil {
t.Fatalf("Error marshaling proto: %v", err)
}
tests := []struct {
desc string
fn func(storage.AdminStorage) error
wantErr bool
}{
{
desc: "CreateTree Bad Settings",
fn: func(s storage.AdminStorage) error {
tree := proto.Clone(testonly.LogTree).(*trillian.Tree)
tree.StorageSettings = badSettings
_, err := storage.CreateTree(ctx, s, tree)
return err
},
wantErr: true,
},
{
desc: "CreateTree nil Settings",
fn: func(s storage.AdminStorage) error {
tree := proto.Clone(testonly.LogTree).(*trillian.Tree)
tree.StorageSettings = nil
_, err := storage.CreateTree(ctx, s, tree)
return err
},
wantErr: false,
},
{
desc: "CreateTree StorageOptions Settings",
fn: func(s storage.AdminStorage) error {
tree := proto.Clone(testonly.LogTree).(*trillian.Tree)
tree.StorageSettings = goodSettings
_, err := storage.CreateTree(ctx, s, tree)
return err
},
wantErr: false,
},
{
desc: "UpdateTree",
fn: func(s storage.AdminStorage) error {
tree, err := storage.CreateTree(ctx, s, testonly.LogTree)
if err != nil {
t.Fatalf("CreateTree() failed with err = %v", err)
}
_, err = storage.UpdateTree(ctx, s, tree.TreeId, func(tree *trillian.Tree) { tree.StorageSettings = badSettings })
return err
},
wantErr: true,
},
}
for _, test := range tests {
if err := test.fn(s); (err != nil) != test.wantErr {
t.Errorf("err: %v, wantErr = %v", err, test.wantErr)
}
}
}
func TestAdminTX_GetTreeLegacies(t *testing.T) {
cleanTestDB(DB)
s := NewAdminStorage(DB)
ctx := context.Background()
serializedStorageSettings := func(revisioned bool) []byte {
ss := storageSettings{
Revisioned: revisioned,
}
buff := &bytes.Buffer{}
enc := gob.NewEncoder(buff)
if err := enc.Encode(ss); err != nil {
t.Fatalf("failed to encode storageSettings: %v", err)
}
return buff.Bytes()
}
tests := []struct {
desc string
key []byte
wantRevisioned bool
}{
{
desc: "No data",
key: []byte{},
wantRevisioned: true,
},
{
desc: "Public key",
key: []byte("trustmethatthisisapublickey"),
wantRevisioned: true,
},
{
desc: "StorageOptions revisioned",
key: serializedStorageSettings(true),
wantRevisioned: true,
},
{
desc: "StorageOptions revisionless",
key: serializedStorageSettings(false),
wantRevisioned: false,
},
}
for _, tC := range tests {
tree, err := storage.CreateTree(ctx, s, testonly.LogTree)
if err != nil {
t.Fatal(err)
}
tx, err := s.db.BeginTx(ctx, nil )
if err != nil {
t.Fatal(err)
}
if _, err := tx.Exec("UPDATE Trees SET PublicKey = ? WHERE TreeId = ?", tC.key, tree.TreeId); err != nil {
t.Fatal(err)
}
if err := tx.Commit(); err != nil {
t.Fatal(err)
}
readTree, err := storage.GetTree(ctx, s, tree.TreeId)
if err != nil {
t.Fatal(err)
}
o := &mysqlpb.StorageOptions{}
if err := anypb.UnmarshalTo(readTree.StorageSettings, o, proto.UnmarshalOptions{}); err != nil {
t.Fatal(err)
}
if got, want := o.SubtreeRevisions, tC.wantRevisioned; got != want {
t.Errorf("%s SubtreeRevisions: got %t, wanted %t", tC.desc, got, want)
}
}
}
func TestAdminTX_HardDeleteTree(t *testing.T) {
cleanTestDB(DB)
s := NewAdminStorage(DB)
ctx := context.Background()
tree, err := storage.CreateTree(ctx, s, testonly.LogTree)
if err != nil {
t.Fatalf("CreateTree() returned err = %v", err)
}
if err := s.ReadWriteTransaction(ctx, func(ctx context.Context, tx storage.AdminTX) error {
if _, err := tx.SoftDeleteTree(ctx, tree.TreeId); err != nil {
return err
}
return tx.HardDeleteTree(ctx, tree.TreeId)
}); err != nil {
t.Fatalf("ReadWriteTransaction() returned err = %v", err)
}
var name string
if err := DB.QueryRowContext(ctx, "SELECT DisplayName FROM Trees WHERE TreeId = ?", tree.TreeId).Scan(&name); err != sql.ErrNoRows {
t.Errorf("QueryRowContext() returned err = %v, want = %v", err, sql.ErrNoRows)
}
}
func TestCheckDatabaseAccessible_Fails(t *testing.T) {
ctx := context.Background()
db, done := openTestDBOrDie()
cleanTestDB(db)
s := NewAdminStorage(db)
done(ctx)
if err := s.CheckDatabaseAccessible(ctx); err == nil {
t.Error("TestCheckDatabaseAccessible_Fails got: nil, want: err")
}
}
func TestCheckDatabaseAccessible_OK(t *testing.T) {
cleanTestDB(DB)
s := NewAdminStorage(DB)
ctx := context.Background()
if err := s.CheckDatabaseAccessible(ctx); err != nil {
t.Errorf("TestCheckDatabaseAccessible_OK got: %v, want: nil", err)
}
}
func setNulls(ctx context.Context, db *sql.DB, treeID int64) error {
stmt, err := db.PrepareContext(ctx, "UPDATE Trees SET DisplayName = NULL, Description = NULL WHERE TreeId = ?")
if err != nil {
return err
}
defer func() { _ = stmt.Close() }()
_, err = stmt.ExecContext(ctx, treeID)
return err
}