package testdb
import (
"bytes"
"context"
"database/sql"
"fmt"
"log"
"net/url"
"os"
"strings"
"testing"
"time"
"github.com/google/trillian/testonly"
"golang.org/x/sys/unix"
"k8s.io/klog/v2"
_ "github.com/go-sql-driver/mysql"
_ "github.com/lib/pq"
)
const (
MySQLURIEnv = "TEST_MYSQL_URI"
defaultTestMySQLURI = "root@tcp(127.0.0.1)/"
CockroachDBURIEnv = "TEST_COCKROACHDB_URI"
defaultTestCockroachDBURI = "postgres://root@localhost:26257/?sslmode=disable"
)
type storageDriverInfo struct {
sqlDriverName string
schema string
uriFunc func(paths ...string) string
}
var (
trillianMySQLSchema = testonly.RelativeToPackage("../mysql/schema/storage.sql")
trillianCRDBSchema = testonly.RelativeToPackage("../crdb/schema/storage.sql")
)
type DriverName string
const (
DriverMySQL DriverName = "mysql"
DriverCockroachDB DriverName = "cockroachdb"
)
var driverMapping = map[DriverName]storageDriverInfo{
DriverMySQL: {
sqlDriverName: "mysql",
schema: trillianMySQLSchema,
uriFunc: mysqlURI,
},
DriverCockroachDB: {
sqlDriverName: "postgres",
schema: trillianCRDBSchema,
uriFunc: crdbURI,
},
}
func mysqlURI(dbRef ...string) string {
var stringurl string
if e := os.Getenv(MySQLURIEnv); len(e) > 0 {
stringurl = e
} else {
stringurl = defaultTestMySQLURI
}
for _, ref := range dbRef {
separator := "/"
if strings.HasSuffix(stringurl, "/") {
separator = ""
}
stringurl = strings.Join([]string{stringurl, ref}, separator)
}
return stringurl
}
func crdbURI(dbRef ...string) string {
var uri *url.URL
if e := os.Getenv(CockroachDBURIEnv); len(e) > 0 {
uri = getURL(e)
} else {
uri = getURL(defaultTestCockroachDBURI)
}
return addPathToURI(uri, dbRef...)
}
func addPathToURI(uri *url.URL, paths ...string) string {
if len(paths) > 0 {
for _, ref := range paths {
currentPaths := uri.Path
if currentPaths == "/" {
currentPaths = ""
}
uri.Path = strings.Join([]string{currentPaths, ref}, "/")
}
}
return uri.String()
}
func getURL(unparsedurl string) *url.URL {
u, _ := url.Parse(unparsedurl)
return u
}
func MySQLAvailable() bool {
return dbAvailable(DriverMySQL)
}
func CockroachDBAvailable() bool {
return dbAvailable(DriverCockroachDB)
}
func dbAvailable(driver DriverName) bool {
driverName := driverMapping[driver].sqlDriverName
uri := driverMapping[driver].uriFunc()
db, err := sql.Open(driverName, uri)
if err != nil {
log.Printf("sql.Open(): %v", err)
return false
}
defer func() {
if err := db.Close(); err != nil {
log.Printf("db.Close(): %v", err)
}
}()
if err := db.Ping(); err != nil {
log.Printf("db.Ping(): %v", err)
return false
}
return true
}
func SetFDLimit(uLimit uint64) error {
var rLimit unix.Rlimit
if err := unix.Getrlimit(unix.RLIMIT_NOFILE, &rLimit); err != nil {
return err
}
if uLimit > rLimit.Max {
return fmt.Errorf("Could not set FD limit to %v. Must be less than the hard limit %v", uLimit, rLimit.Max)
}
rLimit.Cur = uLimit
return unix.Setrlimit(unix.RLIMIT_NOFILE, &rLimit)
}
func newEmptyDB(ctx context.Context, driver DriverName) (*sql.DB, func(context.Context), error) {
if err := SetFDLimit(2048); err != nil {
return nil, nil, err
}
inf, gotinf := driverMapping[driver]
if !gotinf {
return nil, nil, fmt.Errorf("unknown driver %q", driver)
}
db, err := sql.Open(inf.sqlDriverName, inf.uriFunc())
if err != nil {
return nil, nil, err
}
name := fmt.Sprintf("trl_%v", time.Now().UnixNano())
stmt := fmt.Sprintf("CREATE DATABASE %v", name)
if _, err := db.ExecContext(ctx, stmt); err != nil {
return nil, nil, fmt.Errorf("error running statement %q: %v", stmt, err)
}
if err := db.Close(); err != nil {
return nil, nil, fmt.Errorf("failed to close DB: %v", err)
}
uri := inf.uriFunc(name)
db, err = sql.Open(inf.sqlDriverName, uri)
if err != nil {
return nil, nil, err
}
done := func(ctx context.Context) {
defer func() {
if err := db.Close(); err != nil {
klog.Errorf("db.Close(): %v", err)
}
}()
if _, err := db.ExecContext(ctx, fmt.Sprintf("DROP DATABASE %v", name)); err != nil {
klog.Warningf("Failed to drop test database %q: %v", name, err)
}
}
return db, done, db.Ping()
}
func NewTrillianDB(ctx context.Context, driver DriverName) (*sql.DB, func(context.Context), error) {
db, done, err := newEmptyDB(ctx, driver)
if err != nil {
return nil, nil, err
}
schema := driverMapping[driver].schema
sqlBytes, err := os.ReadFile(schema)
if err != nil {
return nil, nil, err
}
for _, stmt := range strings.Split(sanitize(string(sqlBytes)), ";") {
stmt = strings.TrimSpace(stmt)
if stmt == "" {
continue
}
if _, err := db.ExecContext(ctx, stmt); err != nil {
return nil, nil, fmt.Errorf("error running statement %q: %v", stmt, err)
}
}
return db, done, nil
}
func sanitize(script string) string {
buf := &bytes.Buffer{}
for _, line := range strings.Split(string(script), "\n") {
line = strings.TrimSpace(line)
if line == "" || line[0] == '#' || strings.Index(line, "--") == 0 {
continue
}
buf.WriteString(line)
buf.WriteString("\n")
}
return buf.String()
}
func SkipIfNoMySQL(t *testing.T) {
t.Helper()
if !MySQLAvailable() {
t.Skip("Skipping test as MySQL not available")
}
t.Logf("Test MySQL available at %q", mysqlURI())
}
func SkipIfNoCockroachDB(t *testing.T) {
t.Helper()
if !CockroachDBAvailable() {
t.Skip("Skipping test as CockroachDB not available")
}
t.Logf("Test CockroachDB available at %q", crdbURI())
}