131 lines
3.0 KiB
Go
131 lines
3.0 KiB
Go
package store
|
|
|
|
import (
|
|
"database/sql"
|
|
"fmt"
|
|
"net/url"
|
|
"strconv"
|
|
"strings"
|
|
"time"
|
|
|
|
_ "github.com/jackc/pgx/v5/stdlib"
|
|
)
|
|
|
|
const (
|
|
dialectSQLite = "sqlite"
|
|
dialectPostgres = "postgres"
|
|
)
|
|
|
|
func rebind(query string) string {
|
|
n := 0
|
|
var b strings.Builder
|
|
for i := 0; i < len(query); i++ {
|
|
if query[i] == '?' {
|
|
n++
|
|
b.WriteByte('$')
|
|
b.WriteString(strconv.Itoa(n))
|
|
continue
|
|
}
|
|
b.WriteByte(query[i])
|
|
}
|
|
return b.String()
|
|
}
|
|
|
|
func (s *Store) q(query string) string {
|
|
if s.dialect == dialectPostgres {
|
|
return rebind(query)
|
|
}
|
|
return query
|
|
}
|
|
|
|
func applySchema(db *sql.DB, schema string) error {
|
|
for _, stmt := range strings.Split(schema, ";") {
|
|
stmt = strings.TrimSpace(stmt)
|
|
if stmt == "" {
|
|
continue
|
|
}
|
|
upper := strings.ToUpper(stmt)
|
|
if strings.HasPrefix(upper, "PRAGMA") {
|
|
continue
|
|
}
|
|
if _, err := db.Exec(stmt); err != nil {
|
|
return fmt.Errorf("%w: %s", err, stmt)
|
|
}
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func postgresDSN(raw string) (string, error) {
|
|
u, err := url.Parse(raw)
|
|
if err != nil {
|
|
return "", fmt.Errorf("DATABASE_URL: %w", err)
|
|
}
|
|
switch u.Scheme {
|
|
case "postgres", "postgresql":
|
|
default:
|
|
return "", fmt.Errorf("DATABASE_URL must be a postgres URL")
|
|
}
|
|
q := u.Query()
|
|
if strings.EqualFold(q.Get("sslrootcert"), "system") {
|
|
q.Del("sslrootcert")
|
|
}
|
|
q.Del("sslnegotiation")
|
|
if q.Get("sslmode") == "" {
|
|
q.Set("sslmode", "verify-full")
|
|
}
|
|
u.RawQuery = q.Encode()
|
|
return u.String(), nil
|
|
}
|
|
|
|
func OpenPostgres(databaseURL, schema string) (*Store, error) {
|
|
return openPostgres(databaseURL, schema, 5*time.Minute)
|
|
}
|
|
|
|
// OpenPostgresWithoutSessionCleanup opens Postgres without a session cleanup goroutine (for tests).
|
|
func OpenPostgresWithoutSessionCleanup(databaseURL, schema string) (*Store, error) {
|
|
return openPostgres(databaseURL, schema, 0)
|
|
}
|
|
|
|
func openPostgres(databaseURL, schema string, sessionCleanup time.Duration) (*Store, error) {
|
|
dsn, err := postgresDSN(databaseURL)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
db, err := sql.Open("pgx", dsn)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
db.SetMaxOpenConns(20)
|
|
db.SetMaxIdleConns(5)
|
|
if err := db.Ping(); err != nil {
|
|
_ = db.Close()
|
|
return nil, fmt.Errorf("postgres ping: %w", err)
|
|
}
|
|
if err := applySchema(db, schema); err != nil {
|
|
_ = db.Close()
|
|
return nil, fmt.Errorf("apply schema: %w", err)
|
|
}
|
|
if err := applySessionsSchema(db, dialectPostgres); err != nil {
|
|
_ = db.Close()
|
|
return nil, fmt.Errorf("apply sessions schema: %w", err)
|
|
}
|
|
if err := migrateUserProfileColumns(db, dialectPostgres); err != nil {
|
|
_ = db.Close()
|
|
return nil, fmt.Errorf("migrate profile columns: %w", err)
|
|
}
|
|
st := &Store{db: db, dialect: dialectPostgres}
|
|
st.initSessionStore(sessionCleanup)
|
|
return st, nil
|
|
}
|
|
|
|
// Connect uses PlanetScale Postgres when DATABASE_URL is set, otherwise SQLite.
|
|
func Connect(databaseURL, sqlitePath, schema string) (*Store, error) {
|
|
if strings.TrimSpace(databaseURL) != "" {
|
|
return OpenPostgres(databaseURL, schema)
|
|
}
|
|
if sqlitePath == "" {
|
|
sqlitePath = "data.db"
|
|
}
|
|
return Open(sqlitePath, schema)
|
|
}
|