package admin
import (
"context"
"fmt"
"github.com/google/trillian"
"github.com/google/trillian/extension"
"github.com/google/trillian/storage"
"google.golang.org/genproto/protobuf/field_mask"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
"k8s.io/klog/v2"
)
type Server struct {
registry extension.Registry
allowedTreeTypes []trillian.TreeType
}
func New(registry extension.Registry, allowedTreeTypes []trillian.TreeType) *Server {
return &Server{
registry: registry,
allowedTreeTypes: allowedTreeTypes,
}
}
func (s *Server) IsHealthy() error {
return s.registry.AdminStorage.CheckDatabaseAccessible(context.Background())
}
func (s *Server) ListTrees(ctx context.Context, req *trillian.ListTreesRequest) (*trillian.ListTreesResponse, error) {
resp, err := storage.ListTrees(ctx, s.registry.AdminStorage, req.GetShowDeleted())
if err != nil {
return nil, err
}
return &trillian.ListTreesResponse{Tree: resp}, nil
}
func (s *Server) GetTree(ctx context.Context, req *trillian.GetTreeRequest) (*trillian.Tree, error) {
tree, err := storage.GetTree(ctx, s.registry.AdminStorage, req.GetTreeId())
if err != nil {
return nil, err
}
return tree, nil
}
func (s *Server) CreateTree(ctx context.Context, req *trillian.CreateTreeRequest) (*trillian.Tree, error) {
tree := req.GetTree()
if tree == nil {
return nil, status.Errorf(codes.InvalidArgument, "a tree is required")
}
if err := s.validateAllowedTreeType(tree.TreeType); err != nil {
return nil, status.Error(codes.InvalidArgument, err.Error())
}
if tree.TreeType != trillian.TreeType_LOG && tree.TreeType != trillian.TreeType_PREORDERED_LOG {
return nil, status.Errorf(codes.InvalidArgument, "invalid tree type: %v", tree.TreeType)
}
tree.TreeId = 0
tree.CreateTime = nil
tree.UpdateTime = nil
tree.Deleted = false
tree.DeleteTime = nil
createdTree, err := storage.CreateTree(ctx, s.registry.AdminStorage, tree)
if err != nil {
return nil, err
}
return createdTree, nil
}
func (s *Server) validateAllowedTreeType(tt trillian.TreeType) error {
if s.allowedTreeTypes == nil {
return nil
}
for _, allowedType := range s.allowedTreeTypes {
if tt == allowedType {
return nil
}
}
return fmt.Errorf("tree type %s not allowed by this server", tt)
}
func (s *Server) UpdateTree(ctx context.Context, req *trillian.UpdateTreeRequest) (*trillian.Tree, error) {
tree := req.GetTree()
mask := req.GetUpdateMask()
if tree == nil {
return nil, status.Errorf(codes.InvalidArgument, "a tree is required")
}
if err := applyUpdateMask(&trillian.Tree{}, &trillian.Tree{}, mask); err != nil {
return nil, err
}
updatedTree, err := storage.UpdateTree(ctx, s.registry.AdminStorage, tree.TreeId, func(other *trillian.Tree) {
if err := applyUpdateMask(tree, other, mask); err != nil {
klog.Errorf("Error applying mask on tree update: %v", err)
}
})
if err != nil {
return nil, err
}
return updatedTree, nil
}
func applyUpdateMask(from, to *trillian.Tree, mask *field_mask.FieldMask) error {
if mask == nil || len(mask.Paths) == 0 {
return status.Errorf(codes.InvalidArgument, "an update_mask is required")
}
for _, path := range mask.Paths {
switch path {
case "tree_state":
to.TreeState = from.TreeState
case "tree_type":
to.TreeType = from.TreeType
case "display_name":
to.DisplayName = from.DisplayName
case "description":
to.Description = from.Description
case "storage_settings":
to.StorageSettings = from.StorageSettings
case "max_root_duration":
to.MaxRootDuration = from.MaxRootDuration
default:
return status.Errorf(codes.InvalidArgument, "invalid update_mask path: %q", path)
}
}
return nil
}
func (s *Server) DeleteTree(ctx context.Context, req *trillian.DeleteTreeRequest) (*trillian.Tree, error) {
tree, err := storage.SoftDeleteTree(ctx, s.registry.AdminStorage, req.GetTreeId())
if err != nil {
return nil, err
}
return tree, nil
}
func (s *Server) UndeleteTree(ctx context.Context, req *trillian.UndeleteTreeRequest) (*trillian.Tree, error) {
tree, err := storage.UndeleteTree(ctx, s.registry.AdminStorage, req.GetTreeId())
if err != nil {
return nil, err
}
return tree, nil
}