package admin
import (
"context"
"net"
"sort"
"testing"
"time"
"github.com/google/go-cmp/cmp"
"github.com/google/trillian"
"github.com/google/trillian/server/interceptor"
"github.com/google/trillian/storage"
"github.com/google/trillian/storage/testdb"
"github.com/google/trillian/storage/testonly"
"github.com/google/trillian/testonly/integration"
"google.golang.org/genproto/protobuf/field_mask"
"google.golang.org/grpc"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/credentials/insecure"
"google.golang.org/grpc/status"
"google.golang.org/protobuf/proto"
"google.golang.org/protobuf/types/known/timestamppb"
"k8s.io/klog/v2"
sa "github.com/google/trillian/server/admin"
grpc_middleware "github.com/grpc-ecosystem/go-grpc-middleware"
)
func TestAdminServer_CreateTree(t *testing.T) {
ctx := context.Background()
ts, err := setupAdminServer(ctx, t)
if err != nil {
t.Fatalf("setupAdminServer() failed: %v", err)
}
defer ts.closeAll()
invalidTree := proto.Clone(testonly.LogTree).(*trillian.Tree)
invalidTree.TreeState = trillian.TreeState_UNKNOWN_TREE_STATE
timestamp := timestamppb.New(time.Unix(1000, 0))
generatedFieldsTree := proto.Clone(testonly.LogTree).(*trillian.Tree)
generatedFieldsTree.TreeId = 10
generatedFieldsTree.CreateTime = timestamp
generatedFieldsTree.UpdateTime = timestamp
generatedFieldsTree.UpdateTime.Seconds++
generatedFieldsTree.Deleted = true
generatedFieldsTree.DeleteTime = timestamp
generatedFieldsTree.DeleteTime.Seconds++
tests := []struct {
desc string
req *trillian.CreateTreeRequest
wantCode codes.Code
}{
{
desc: "validTree",
req: &trillian.CreateTreeRequest{Tree: testonly.LogTree},
},
{
desc: "generatedFieldsTree",
req: &trillian.CreateTreeRequest{Tree: generatedFieldsTree},
},
{
desc: "nilTree",
req: &trillian.CreateTreeRequest{},
wantCode: codes.InvalidArgument,
},
{
desc: "invalidTree",
req: &trillian.CreateTreeRequest{Tree: invalidTree},
wantCode: codes.InvalidArgument,
},
}
for _, test := range tests {
createdTree, err := ts.adminClient.CreateTree(ctx, test.req)
if s, ok := status.FromError(err); !ok || s.Code() != test.wantCode {
t.Errorf("%v: CreateTree() = (_, %v), wantCode = %v", test.desc, err, test.wantCode)
continue
} else if err != nil {
continue
}
if createdTree.TreeId == 0 {
t.Errorf("%v: createdTree.TreeId = 0", test.desc)
}
if createdTree.CreateTime == nil {
t.Errorf("%v: createdTree.CreateTime = nil", test.desc)
}
if !proto.Equal(createdTree.CreateTime, createdTree.UpdateTime) {
t.Errorf("%v: createdTree.UpdateTime = %+v, want = %+v", test.desc, createdTree.UpdateTime, createdTree.CreateTime)
}
if createdTree.Deleted {
t.Errorf("%v: createdTree.Deleted = true", test.desc)
}
if createdTree.DeleteTime != nil {
t.Errorf("%v: createdTree.DeleteTime is non-nil", test.desc)
}
storedTree, err := ts.adminClient.GetTree(ctx, &trillian.GetTreeRequest{TreeId: createdTree.TreeId})
if err != nil {
t.Errorf("%v: GetTree() = (_, %v), want = (_, nil)", test.desc, err)
continue
}
if diff := cmp.Diff(storedTree, createdTree, cmp.Comparer(proto.Equal)); diff != "" {
t.Errorf("%v: post-CreateTree diff (-stored +created):\n%v", test.desc, diff)
}
}
}
func TestAdminServer_UpdateTree(t *testing.T) {
ctx := context.Background()
ts, err := setupAdminServer(ctx, t)
if err != nil {
t.Fatalf("setupAdminServer() failed: %v", err)
}
defer ts.closeAll()
baseTree := proto.Clone(testonly.LogTree).(*trillian.Tree)
successTree := &trillian.Tree{
TreeState: trillian.TreeState_FROZEN,
DisplayName: "Brand New Tree Name",
Description: "Brand New Tree Desc",
}
successMask := &field_mask.FieldMask{Paths: []string{"tree_state", "display_name", "description"}}
successWant := proto.Clone(baseTree).(*trillian.Tree)
successWant.TreeState = successTree.TreeState
successWant.DisplayName = successTree.DisplayName
successWant.Description = successTree.Description
tests := []struct {
desc string
createTree, wantTree *trillian.Tree
req *trillian.UpdateTreeRequest
wantCode codes.Code
}{
{
desc: "success",
createTree: baseTree,
wantTree: successWant,
req: &trillian.UpdateTreeRequest{Tree: successTree, UpdateMask: successMask},
},
{
desc: "notFound",
req: &trillian.UpdateTreeRequest{
Tree: &trillian.Tree{TreeId: 12345, DisplayName: "New Name"},
UpdateMask: &field_mask.FieldMask{Paths: []string{"display_name"}},
},
wantCode: codes.NotFound,
},
{
desc: "readonlyField",
createTree: baseTree,
req: &trillian.UpdateTreeRequest{
Tree: successTree,
UpdateMask: &field_mask.FieldMask{Paths: []string{"tree_type"}},
},
wantCode: codes.InvalidArgument,
},
{
desc: "invalidUpdate",
createTree: baseTree,
req: &trillian.UpdateTreeRequest{
Tree: &trillian.Tree{},
UpdateMask: &field_mask.FieldMask{Paths: []string{"tree_state"}},
},
wantCode: codes.InvalidArgument,
},
}
for _, test := range tests {
if test.createTree != nil {
tree, err := ts.adminClient.CreateTree(ctx, &trillian.CreateTreeRequest{Tree: test.createTree})
if err != nil {
t.Errorf("%v: CreateTree() returned err = %v", test.desc, err)
continue
}
test.req.Tree.TreeId = tree.TreeId
}
tree, err := ts.adminClient.UpdateTree(ctx, test.req)
if s, ok := status.FromError(err); !ok || s.Code() != test.wantCode {
t.Errorf("%v: UpdateTree() returned err = %q, wantCode = %v", test.desc, err, test.wantCode)
continue
} else if err != nil {
continue
}
created := tree.CreateTime.AsTime()
updated := tree.UpdateTime.AsTime()
if created.After(updated) {
t.Errorf("%v: CreateTime > UpdateTime (%v > %v)", test.desc, tree.CreateTime, tree.UpdateTime)
}
want := proto.Clone(test.wantTree).(*trillian.Tree)
want.TreeId = tree.TreeId
want.CreateTime = tree.CreateTime
want.UpdateTime = tree.UpdateTime
want.StorageSettings = tree.StorageSettings
if !proto.Equal(tree, want) {
diff := cmp.Diff(tree, want, cmp.Comparer(proto.Equal))
t.Errorf("%v: post-UpdateTree diff:\n%v", test.desc, diff)
}
}
}
func TestAdminServer_GetTree(t *testing.T) {
ctx := context.Background()
ts, err := setupAdminServer(ctx, t)
if err != nil {
t.Fatalf("setupAdminServer() failed: %v", err)
}
defer ts.closeAll()
tests := []struct {
desc string
treeID int64
wantCode codes.Code
}{
{
desc: "negativeTreeID",
treeID: -1,
wantCode: codes.NotFound,
},
{
desc: "notFound",
treeID: 12345,
wantCode: codes.NotFound,
},
}
for _, test := range tests {
_, err := ts.adminClient.GetTree(ctx, &trillian.GetTreeRequest{TreeId: test.treeID})
if s, ok := status.FromError(err); !ok || s.Code() != test.wantCode {
t.Errorf("%v: GetTree() = (_, %v), wantCode = %v", test.desc, err, test.wantCode)
}
}
}
func TestAdminServer_ListTrees(t *testing.T) {
ctx := context.Background()
ts, err := setupAdminServer(ctx, t)
if err != nil {
t.Fatalf("setupAdminServer() failed: %v", err)
}
defer ts.closeAll()
tests := []struct {
desc string
numTrees int
}{
{desc: "empty"},
{desc: "oneTree", numTrees: 1},
{desc: "threeTrees", numTrees: 3},
}
var createdTrees []*trillian.Tree
for _, test := range tests {
t.Run(test.desc, func(t *testing.T) {
if l := len(createdTrees); l > test.numTrees {
t.Fatalf("numTrees = %v, but we already have %v stored trees", test.numTrees, l)
} else if l < test.numTrees {
for i := l; i < test.numTrees; i++ {
tree := proto.Clone(testonly.LogTree).(*trillian.Tree)
req := &trillian.CreateTreeRequest{Tree: tree}
resp, err := ts.adminClient.CreateTree(ctx, req)
if err != nil {
t.Fatalf("CreateTree(_, %v) = (_, %q), want = (_, nil)", req, err)
}
createdTrees = append(createdTrees, resp)
}
sortByTreeID(createdTrees)
}
resp, err := ts.adminClient.ListTrees(ctx, &trillian.ListTreesRequest{})
if err != nil {
t.Errorf("%v: ListTrees() = (_, %q), want = (_, nil)", test.desc, err)
return
}
got := resp.Tree
sortByTreeID(got)
if diff := cmp.Diff(got, createdTrees, cmp.Comparer(proto.Equal)); diff != "" {
t.Errorf("post-ListTrees diff:\n%v", diff)
}
})
}
}
func sortByTreeID(s []*trillian.Tree) {
less := func(i, j int) bool {
return s[i].TreeId < s[j].TreeId
}
sort.Slice(s, less)
}
func TestAdminServer_DeleteTree(t *testing.T) {
ctx := context.Background()
ts, err := setupAdminServer(ctx, t)
if err != nil {
t.Fatalf("setupAdminServer() failed: %v", err)
}
defer ts.closeAll()
tests := []struct {
desc string
baseTree *trillian.Tree
}{
{desc: "logTree", baseTree: testonly.LogTree},
}
for _, test := range tests {
createdTree, err := ts.adminClient.CreateTree(ctx, &trillian.CreateTreeRequest{Tree: test.baseTree})
if err != nil {
t.Fatalf("%v: CreateTree() returned err = %v", test.desc, err)
}
deletedTree, err := ts.adminClient.DeleteTree(ctx, &trillian.DeleteTreeRequest{TreeId: createdTree.TreeId})
if err != nil {
t.Errorf("%v: DeleteTree() returned err = %v", test.desc, err)
continue
}
if deletedTree.DeleteTime == nil {
t.Errorf("%v: tree.DeleteTime = nil, want non-nil", test.desc)
}
want := proto.Clone(createdTree).(*trillian.Tree)
want.Deleted = true
want.DeleteTime = deletedTree.DeleteTime
if got := deletedTree; !proto.Equal(got, want) {
diff := cmp.Diff(got, want, cmp.Comparer(proto.Equal))
t.Errorf("%v: post-DeleteTree() diff (-got +want):\n%v", test.desc, diff)
}
storedTree, err := ts.adminClient.GetTree(ctx, &trillian.GetTreeRequest{TreeId: deletedTree.TreeId})
if err != nil {
t.Fatalf("%v: GetTree() returned err = %v", test.desc, err)
}
if got, want := storedTree, deletedTree; !proto.Equal(got, want) {
diff := cmp.Diff(got, want, cmp.Comparer(proto.Equal))
t.Errorf("%v: post-GetTree() diff (-got +want):\n%v", test.desc, diff)
}
}
}
func TestAdminServer_DeleteTreeErrors(t *testing.T) {
ctx := context.Background()
ts, err := setupAdminServer(ctx, t)
if err != nil {
t.Fatalf("setupAdminServer() failed: %v", err)
}
defer ts.closeAll()
createdTree, err := ts.adminClient.CreateTree(ctx, &trillian.CreateTreeRequest{Tree: testonly.LogTree})
if err != nil {
t.Fatalf("CreateTree() returned err = %v", err)
}
deletedTree, err := ts.adminClient.DeleteTree(ctx, &trillian.DeleteTreeRequest{TreeId: createdTree.TreeId})
if err != nil {
t.Fatalf("DeleteTree() returned err = %v", err)
}
tests := []struct {
desc string
req *trillian.DeleteTreeRequest
wantCode codes.Code
}{
{
desc: "unknownTree",
req: &trillian.DeleteTreeRequest{TreeId: 12345},
wantCode: codes.NotFound,
},
{
desc: "alreadyDeleted",
req: &trillian.DeleteTreeRequest{TreeId: deletedTree.TreeId},
wantCode: codes.FailedPrecondition,
},
}
for _, test := range tests {
_, err := ts.adminClient.DeleteTree(ctx, test.req)
if s, ok := status.FromError(err); !ok || s.Code() != test.wantCode {
t.Errorf("%v: DeleteTree() returned err = %v, wantCode = %s", test.desc, err, test.wantCode)
}
}
}
func TestAdminServer_UndeleteTree(t *testing.T) {
ctx := context.Background()
ts, err := setupAdminServer(ctx, t)
if err != nil {
t.Fatalf("setupAdminServer() failed: %v", err)
}
defer ts.closeAll()
tests := []struct {
desc string
baseTree *trillian.Tree
}{
{desc: "logTree", baseTree: testonly.LogTree},
}
for _, test := range tests {
createdTree, err := ts.adminClient.CreateTree(ctx, &trillian.CreateTreeRequest{Tree: test.baseTree})
if err != nil {
t.Fatalf("%v: CreateTree() returned err = %v", test.desc, err)
}
deletedTree, err := ts.adminClient.DeleteTree(ctx, &trillian.DeleteTreeRequest{TreeId: createdTree.TreeId})
if err != nil {
t.Fatalf("%v: DeleteTree() returned err = %v", test.desc, err)
}
undeletedTree, err := ts.adminClient.UndeleteTree(ctx, &trillian.UndeleteTreeRequest{TreeId: deletedTree.TreeId})
if err != nil {
t.Errorf("%v: UndeleteTree() returned err = %v", test.desc, err)
continue
}
if got, want := undeletedTree, createdTree; !proto.Equal(got, want) {
diff := cmp.Diff(got, want, cmp.Comparer(proto.Equal))
t.Errorf("%v: post-UndeleteTree() diff (-got +want):\n%v", test.desc, diff)
}
storedTree, err := ts.adminClient.GetTree(ctx, &trillian.GetTreeRequest{TreeId: deletedTree.TreeId})
if err != nil {
t.Fatalf("%v: GetTree() returned err = %v", test.desc, err)
}
if got, want := storedTree, createdTree; !proto.Equal(got, want) {
diff := cmp.Diff(got, want, cmp.Comparer(proto.Equal))
t.Errorf("%v: post-GetTree() diff (-got +want):\n%v", test.desc, diff)
}
}
}
func TestAdminServer_UndeleteTreeErrors(t *testing.T) {
ctx := context.Background()
ts, err := setupAdminServer(ctx, t)
if err != nil {
t.Fatalf("setupAdminServer() failed: %v", err)
}
defer ts.closeAll()
tree, err := ts.adminClient.CreateTree(ctx, &trillian.CreateTreeRequest{Tree: testonly.LogTree})
if err != nil {
t.Fatalf("CreateTree() returned err = %v", err)
}
tests := []struct {
desc string
req *trillian.UndeleteTreeRequest
wantCode codes.Code
}{
{
desc: "unknownTree",
req: &trillian.UndeleteTreeRequest{TreeId: 12345},
wantCode: codes.NotFound,
},
{
desc: "notDeleted",
req: &trillian.UndeleteTreeRequest{TreeId: tree.TreeId},
wantCode: codes.FailedPrecondition,
},
}
for _, test := range tests {
_, err := ts.adminClient.UndeleteTree(ctx, test.req)
if s, ok := status.FromError(err); !ok || s.Code() != test.wantCode {
t.Errorf("%v: UndeleteTree() returned err = %v, wantCode = %s", test.desc, err, test.wantCode)
}
}
}
func TestAdminServer_TreeGC(t *testing.T) {
ctx := context.Background()
ts, err := setupAdminServer(ctx, t)
if err != nil {
t.Fatalf("setupAdminServer() failed: %v", err)
}
defer ts.closeAll()
tree, err := ts.adminClient.CreateTree(ctx, &trillian.CreateTreeRequest{Tree: testonly.LogTree})
if err != nil {
t.Fatalf("CreateTree() returned err = %v", err)
}
if _, err := ts.adminClient.DeleteTree(ctx, &trillian.DeleteTreeRequest{TreeId: tree.TreeId}); err != nil {
t.Fatalf("DeleteTree() returned err = %v", err)
}
treeGC := sa.NewDeletedTreeGC(
ts.adminStorage, 1*time.Second , 1*time.Second , nil )
success := false
const attempts = 3
for i := 0; i < attempts; i++ {
_, err := treeGC.RunOnce(ctx)
if err != nil {
t.Error(err)
}
_, err = ts.adminClient.GetTree(ctx, &trillian.GetTreeRequest{TreeId: tree.TreeId})
if s, ok := status.FromError(err); ok && s.Code() == codes.NotFound {
success = true
break
}
time.Sleep(1 * time.Second)
}
if !success {
t.Errorf("Tree %v not hard-deleted after max attempts", tree.TreeId)
}
}
type testServer struct {
adminClient trillian.TrillianAdminClient
adminStorage storage.AdminStorage
lis net.Listener
server *grpc.Server
conn *grpc.ClientConn
cleanup func(context.Context)
}
func (ts *testServer) closeAll() {
if ts.conn != nil {
if err := ts.conn.Close(); err != nil {
klog.Errorf("testServer: conn.Close()=%v", err)
}
}
if ts.server != nil {
ts.server.GracefulStop()
}
if ts.lis != nil {
if err := ts.lis.Close(); err != nil {
klog.Errorf("testServer: lis.Close()=%v", err)
}
}
if ts.cleanup != nil {
ts.cleanup(context.TODO())
}
}
func setupAdminServer(ctx context.Context, t *testing.T) (*testServer, error) {
t.Helper()
testdb.SkipIfNoMySQL(t)
ts := &testServer{}
var err error
ts.lis, err = net.Listen("tcp", "127.0.0.1:0")
if err != nil {
return nil, err
}
registry, done, err := integration.NewRegistryForTests(ctx, testdb.DriverMySQL)
if err != nil {
ts.closeAll()
return nil, err
}
ts.adminStorage = registry.AdminStorage
ts.cleanup = done
ti := interceptor.New(
registry.AdminStorage, registry.QuotaManager, false , registry.MetricFactory)
ts.server = grpc.NewServer(
grpc.UnaryInterceptor(grpc_middleware.ChainUnaryServer(
interceptor.ErrorWrapper,
ti.UnaryInterceptor,
)),
)
trillian.RegisterTrillianAdminServer(ts.server, sa.New(registry, nil ))
go func() {
if err := ts.server.Serve(ts.lis); err != nil {
klog.Errorf("server.Serve()=%v", err)
}
}()
ts.conn, err = grpc.Dial(ts.lis.Addr().String(), grpc.WithTransportCredentials(insecure.NewCredentials()))
if err != nil {
ts.closeAll()
return nil, err
}
ts.adminClient = trillian.NewTrillianAdminClient(ts.conn)
return ts, nil
}