diff --git a/shared/driver/conn.go b/shared/driver/conn.go new file mode 100644 index 0000000000..4c7824bf1b --- /dev/null +++ b/shared/driver/conn.go @@ -0,0 +1,93 @@ +// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved. +// See LICENSE.txt for license information. + +package driver + +import ( + "context" + "database/sql" + "database/sql/driver" +) + +// Conn is a DB driver conn implementation +// which will just pass-through all queries to its +// underlying connection. +type Conn struct { + conn *sql.Conn +} + +// driverConn is a super-interface combining the basic +// driver.Conn interface with some new additions later. +type driverConn interface { + driver.Conn + driver.ConnBeginTx + driver.ConnPrepareContext + driver.ExecerContext + driver.QueryerContext + driver.Pinger +} + +var ( + // Compile-time check to ensure Conn implements the interface. + _ driverConn = &Conn{} +) + +func (c *Conn) Begin() (tx driver.Tx, err error) { + err = c.conn.Raw(func(innerConn interface{}) error { + tx, err = innerConn.(driver.Conn).Begin() //nolint:staticcheck + return err + }) + return tx, err +} + +func (c *Conn) BeginTx(ctx context.Context, opts driver.TxOptions) (tx driver.Tx, err error) { + err = c.conn.Raw(func(innerConn interface{}) error { + tx, err = innerConn.(driver.ConnBeginTx).BeginTx(ctx, opts) + return err + }) + return tx, err +} + +func (c *Conn) Prepare(q string) (stmt driver.Stmt, err error) { + err = c.conn.Raw(func(innerConn interface{}) error { + stmt, err = innerConn.(driver.Conn).Prepare(q) + return err + }) + return stmt, err +} + +func (c *Conn) PrepareContext(ctx context.Context, q string) (stmt driver.Stmt, err error) { + err = c.conn.Raw(func(innerConn interface{}) error { + stmt, err = innerConn.(driver.ConnPrepareContext).PrepareContext(ctx, q) + return err + }) + return stmt, err +} + +func (c *Conn) ExecContext(ctx context.Context, q string, args []driver.NamedValue) (res driver.Result, err error) { + err = c.conn.Raw(func(innerConn interface{}) error { + res, err = innerConn.(driver.ExecerContext).ExecContext(ctx, q, args) + return err + }) + return res, err +} + +func (c *Conn) QueryContext(ctx context.Context, q string, args []driver.NamedValue) (rows driver.Rows, err error) { + err = c.conn.Raw(func(innerConn interface{}) error { + rows, err = innerConn.(driver.QueryerContext).QueryContext(ctx, q, args) + return err + }) + return rows, err +} + +func (c *Conn) Ping(ctx context.Context) error { + return c.conn.Raw(func(innerConn interface{}) error { + return innerConn.(driver.Pinger).Ping(ctx) + }) +} + +func (c *Conn) Close() error { + return c.conn.Raw(func(innerConn interface{}) error { + return innerConn.(driver.Conn).Close() + }) +} diff --git a/shared/driver/driver.go b/shared/driver/driver.go new file mode 100644 index 0000000000..a38854a4b9 --- /dev/null +++ b/shared/driver/driver.go @@ -0,0 +1,57 @@ +// Copyright (c) 2015-present Mattermost, Inc. All Rights Reserved. +// See LICENSE.txt for license information. + +package driver + +import ( + "context" + "database/sql" + "database/sql/driver" +) + +var ( + // Compile-time check to ensure Connector implements the interface. + _ driver.Connector = &Connector{} +) + +// Connector is the DB connector which is used to +// initialize the underlying DB. +type Connector struct { + driverName string + dsn string + db *sql.DB +} + +func NewConnector(driverName, dsn string) (*Connector, error) { + db, err := sql.Open(driverName, dsn) + if err != nil { + return nil, err + } + return &Connector{ + driverName: driverName, + dsn: dsn, + db: db, + }, nil +} + +func (c *Connector) Connect(ctx context.Context) (driver.Conn, error) { + conn, err := c.db.Conn(ctx) + if err != nil { + return nil, err + } + + return &Conn{conn: conn}, nil +} + +func (c *Connector) Driver() driver.Driver { + return &Driver{c: c} +} + +// Driver is a DB driver implementation. +type Driver struct { + c *Connector +} + +func (d Driver) Open(name string) (driver.Conn, error) { + return d.c.Connect(context.Background()) +} diff --git a/store/sqlstore/store.go b/store/sqlstore/store.go index d44ec6922e..ac5ca67fa1 100644 --- a/store/sqlstore/store.go +++ b/store/sqlstore/store.go @@ -321,25 +321,39 @@ func setupConnection(connType string, dataSource string, settings *model.SqlSett db.SetConnMaxLifetime(time.Duration(*settings.ConnMaxLifetimeMilliseconds) * time.Millisecond) db.SetConnMaxIdleTime(time.Duration(*settings.ConnMaxIdleTimeMilliseconds) * time.Millisecond) - var dbmap *gorp.DbMap + dbMap := getDBMap(settings, db) + return dbMap +} + +func getDBMap(settings *model.SqlSettings, db *dbsql.DB) *gorp.DbMap { connectionTimeout := time.Duration(*settings.QueryTimeout) * time.Second - - if *settings.DriverName == model.DATABASE_DRIVER_MYSQL { - dbmap = &gorp.DbMap{Db: db, TypeConverter: mattermConverter{}, Dialect: gorp.MySQLDialect{Engine: "InnoDB", Encoding: "UTF8MB4"}, QueryTimeout: connectionTimeout} - } else if *settings.DriverName == model.DATABASE_DRIVER_POSTGRES { - dbmap = &gorp.DbMap{Db: db, TypeConverter: mattermConverter{}, Dialect: gorp.PostgresDialect{}, QueryTimeout: connectionTimeout} - } else { + var dbMap *gorp.DbMap + switch *settings.DriverName { + case model.DATABASE_DRIVER_MYSQL: + dbMap = &gorp.DbMap{ + Db: db, + TypeConverter: mattermConverter{}, + Dialect: gorp.MySQLDialect{Engine: "InnoDB", Encoding: "UTF8MB4"}, + QueryTimeout: connectionTimeout, + } + case model.DATABASE_DRIVER_POSTGRES: + dbMap = &gorp.DbMap{ + Db: db, + TypeConverter: mattermConverter{}, + Dialect: gorp.PostgresDialect{}, + QueryTimeout: connectionTimeout, + } + default: mlog.Critical("Failed to create dialect specific driver") time.Sleep(time.Second) os.Exit(ExitNoDriver) + return nil } - if settings.Trace != nil && *settings.Trace { - dbmap.TraceOn("sql-trace:", &TraceOnAdapter{}) + dbMap.TraceOn("sql-trace:", &TraceOnAdapter{}) } - - return dbmap + return dbMap } func (ss *SqlStore) SetContext(context context.Context) { diff --git a/store/sqlstore/store_test.go b/store/sqlstore/store_test.go index b31f3d550c..123172e497 100644 --- a/store/sqlstore/store_test.go +++ b/store/sqlstore/store_test.go @@ -4,6 +4,7 @@ package sqlstore import ( + "database/sql" "fmt" "os" "regexp" @@ -20,6 +21,7 @@ import ( "github.com/mattermost/mattermost-server/v5/einterfaces/mocks" "github.com/mattermost/mattermost-server/v5/model" + "github.com/mattermost/mattermost-server/v5/shared/driver" "github.com/mattermost/mattermost-server/v5/store" "github.com/mattermost/mattermost-server/v5/store/searchtest" "github.com/mattermost/mattermost-server/v5/store/storetest" @@ -95,6 +97,47 @@ func StoreTestWithSqlStore(t *testing.T, f func(*testing.T, store.Store, storete } } +func TestDBConnector(t *testing.T) { + testDrivers := []string{ + model.DATABASE_DRIVER_POSTGRES, + model.DATABASE_DRIVER_MYSQL, + } + + for _, dr := range testDrivers { + settings := makeSqlSettings(dr) + + store := &SqlStore{ + settings: settings, + } + connector, err := driver.NewConnector(*settings.DriverName, *settings.DataSource) + require.NoError(t, err) + db := sql.OpenDB(connector) + + store.master = getDBMap(settings, db) + store.stores.post = newSqlPostStore(store, nil) + store.stores.channel = newSqlChannelStore(store, nil) + store.stores.team = newSqlTeamStore(store) + store.stores.thread = newSqlThreadStore(store) + store.stores.user = newSqlUserStore(store, nil) + store.stores.preference = newSqlPreferenceStore(store) + store.stores.bot = newSqlBotStore(store, nil) + store.stores.scheme = newSqlSchemeStore(store) + err = store.GetMaster().CreateTablesIfNotExists() + require.NoError(t, err) + + store.stores.post.(*SqlPostStore).createIndexesIfNotExists() + store.stores.channel.(*SqlChannelStore).createIndexesIfNotExists() + + t.Run(dr, func(t *testing.T) { + // Just testing post store for now. + // This will eventually go away when it is replaced with RPC. + storetest.TestPostStore(t, store, store) + }) + + store.Close() + } +} + func initStores() { if testing.Short() { return @@ -109,8 +152,10 @@ func initStores() { storeTypes = append(storeTypes, newStoreType("PostgreSQL", model.DATABASE_DRIVER_POSTGRES)) } } else { - storeTypes = append(storeTypes, newStoreType("MySQL", model.DATABASE_DRIVER_MYSQL), - newStoreType("PostgreSQL", model.DATABASE_DRIVER_POSTGRES)) + storeTypes = append(storeTypes, + newStoreType("MySQL", model.DATABASE_DRIVER_MYSQL), + newStoreType("PostgreSQL", model.DATABASE_DRIVER_POSTGRES), + ) } defer func() {