package mysql
import (
"bytes"
"database/sql"
"encoding/gob"
"fmt"
"time"
"github.com/google/trillian"
"github.com/google/trillian/storage/mysql/mysqlpb"
"google.golang.org/protobuf/types/known/anypb"
"google.golang.org/protobuf/types/known/durationpb"
"google.golang.org/protobuf/types/known/timestamppb"
)
func toMillisSinceEpoch(t time.Time) int64 {
return t.UnixNano() / 1000000
}
func fromMillisSinceEpoch(ts int64) time.Time {
return time.Unix(0, ts*1000000)
}
func setNullStringIfValid(src sql.NullString, dest *string) {
if src.Valid {
*dest = src.String
}
}
type row interface {
Scan(dest ...interface{}) error
}
func readTree(r row) (*trillian.Tree, error) {
tree := &trillian.Tree{}
var treeState, treeType, hashStrategy, hashAlgorithm, signatureAlgorithm string
var createMillis, updateMillis, maxRootDurationMillis int64
var displayName, description sql.NullString
var privateKey, publicKey []byte
var deleted sql.NullBool
var deleteMillis sql.NullInt64
err := r.Scan(
&tree.TreeId,
&treeState,
&treeType,
&hashStrategy,
&hashAlgorithm,
&signatureAlgorithm,
&displayName,
&description,
&createMillis,
&updateMillis,
&privateKey,
&publicKey,
&maxRootDurationMillis,
&deleted,
&deleteMillis,
)
if err != nil {
return nil, err
}
setNullStringIfValid(displayName, &tree.DisplayName)
setNullStringIfValid(description, &tree.Description)
if ts, ok := trillian.TreeState_value[treeState]; ok {
tree.TreeState = trillian.TreeState(ts)
} else {
return nil, fmt.Errorf("unknown TreeState: %v", treeState)
}
if tt, ok := trillian.TreeType_value[treeType]; ok {
tree.TreeType = trillian.TreeType(tt)
} else {
return nil, fmt.Errorf("unknown TreeType: %v", treeType)
}
if hashStrategy != "RFC6962_SHA256" {
return nil, fmt.Errorf("unknown HashStrategy: %v", hashStrategy)
}
ok := tree.TreeState.String() == treeState &&
tree.TreeType.String() == treeType
if !ok {
return nil, fmt.Errorf(
"mismatched enum: tree = %v, enums = [%v, %v, %v, %v, %v]",
tree,
treeState, treeType, hashStrategy, hashAlgorithm, signatureAlgorithm)
}
tree.CreateTime = timestamppb.New(fromMillisSinceEpoch(createMillis))
if err := tree.CreateTime.CheckValid(); err != nil {
return nil, fmt.Errorf("failed to parse create time: %w", err)
}
tree.UpdateTime = timestamppb.New(fromMillisSinceEpoch(updateMillis))
if err := tree.UpdateTime.CheckValid(); err != nil {
return nil, fmt.Errorf("failed to parse update time: %w", err)
}
tree.MaxRootDuration = durationpb.New(time.Duration(maxRootDurationMillis * int64(time.Millisecond)))
tree.Deleted = deleted.Valid && deleted.Bool
if tree.Deleted && deleteMillis.Valid {
tree.DeleteTime = timestamppb.New(fromMillisSinceEpoch(deleteMillis.Int64))
if err := tree.DeleteTime.CheckValid(); err != nil {
return nil, fmt.Errorf("failed to parse delete time: %w", err)
}
}
buff := bytes.NewBuffer(publicKey)
dec := gob.NewDecoder(buff)
ss := &storageSettings{}
var o *mysqlpb.StorageOptions
if err := dec.Decode(ss); err != nil {
o = &mysqlpb.StorageOptions{
SubtreeRevisions: true,
}
} else {
o = &mysqlpb.StorageOptions{
SubtreeRevisions: ss.Revisioned,
}
}
tree.StorageSettings, err = anypb.New(o)
if err != nil {
return nil, fmt.Errorf("failed to put StorageSettings into tree: %w", err)
}
return tree, nil
}