+125
-10
@@ -4,29 +4,86 @@ import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
_ "github.com/go-sql-driver/mysql"
|
||||
)
|
||||
|
||||
func Open(ctx context.Context, url string) (*sql.DB, error) {
|
||||
d, err := sql.Open("mysql", url)
|
||||
func Open(ctx context.Context, dsn string) (*sql.DB, error) {
|
||||
dsn = normalizeDSN(dsn)
|
||||
d, err := sql.Open("mysql", dsn)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
d.SetMaxOpenConns(25)
|
||||
d.SetMaxIdleConns(10)
|
||||
if err := d.PingContext(ctx); err != nil {
|
||||
return nil, err
|
||||
|
||||
// MariaDB/MySQL can close idle TCP connections. Keep pooled connections fresh
|
||||
// so users do not randomly hit "[mysql] invalid connection" after startup or idle time.
|
||||
d.SetMaxOpenConns(20)
|
||||
d.SetMaxIdleConns(5)
|
||||
d.SetConnMaxLifetime(3 * time.Minute)
|
||||
d.SetConnMaxIdleTime(90 * time.Second)
|
||||
|
||||
deadline := time.Now().Add(90 * time.Second)
|
||||
var lastErr error
|
||||
for {
|
||||
pingCtx, cancel := context.WithTimeout(ctx, 5*time.Second)
|
||||
lastErr = d.PingContext(pingCtx)
|
||||
cancel()
|
||||
if lastErr == nil {
|
||||
return d, nil
|
||||
}
|
||||
if time.Now().After(deadline) {
|
||||
d.Close()
|
||||
return nil, fmt.Errorf("database not ready: %w", lastErr)
|
||||
}
|
||||
time.Sleep(2 * time.Second)
|
||||
}
|
||||
return d, nil
|
||||
}
|
||||
|
||||
func normalizeDSN(dsn string) string {
|
||||
parts := strings.SplitN(dsn, "?", 2)
|
||||
base := parts[0]
|
||||
q := url.Values{}
|
||||
if len(parts) == 2 {
|
||||
parsed, err := url.ParseQuery(parts[1])
|
||||
if err == nil {
|
||||
q = parsed
|
||||
}
|
||||
}
|
||||
defaults := map[string]string{
|
||||
"parseTime": "true",
|
||||
"charset": "utf8mb4",
|
||||
"collation": "utf8mb4_unicode_ci",
|
||||
"loc": "Local",
|
||||
"timeout": "10s",
|
||||
"readTimeout": "30s",
|
||||
"writeTimeout": "30s",
|
||||
}
|
||||
for k, v := range defaults {
|
||||
if q.Get(k) == "" {
|
||||
q.Set(k, v)
|
||||
}
|
||||
}
|
||||
// Migrations are executed statement-by-statement, so multiStatements is not needed.
|
||||
q.Del("multiStatements")
|
||||
return base + "?" + q.Encode()
|
||||
}
|
||||
|
||||
func Migrate(ctx context.Context, d *sql.DB) error {
|
||||
d.ExecContext(ctx, `SELECT GET_LOCK('trading_tool_migrate', 30)`)
|
||||
defer d.ExecContext(ctx, `SELECT RELEASE_LOCK('trading_tool_migrate')`)
|
||||
var locked int
|
||||
if err := d.QueryRowContext(ctx, `SELECT GET_LOCK('trading_tool_migrate', 30)`).Scan(&locked); err != nil {
|
||||
return err
|
||||
}
|
||||
if locked != 1 {
|
||||
return fmt.Errorf("could not acquire migration lock")
|
||||
}
|
||||
defer d.ExecContext(context.Background(), `SELECT RELEASE_LOCK('trading_tool_migrate')`)
|
||||
|
||||
if _, err := d.ExecContext(ctx, `CREATE TABLE IF NOT EXISTS schema_migrations(version VARCHAR(255) PRIMARY KEY, applied_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP)`); err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -47,7 +104,7 @@ func Migrate(ctx context.Context, d *sql.DB) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := d.ExecContext(ctx, string(b)); err != nil {
|
||||
if err := execSQLFile(ctx, d, string(b)); err != nil {
|
||||
return fmt.Errorf("%s: %w", f, err)
|
||||
}
|
||||
if _, err := d.ExecContext(ctx, `INSERT INTO schema_migrations(version) VALUES(?)`, f); err != nil {
|
||||
@@ -56,3 +113,61 @@ func Migrate(ctx context.Context, d *sql.DB) error {
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func execSQLFile(ctx context.Context, d *sql.DB, content string) error {
|
||||
for _, stmt := range splitSQLStatements(content) {
|
||||
stmt = strings.TrimSpace(stmt)
|
||||
if stmt == "" {
|
||||
continue
|
||||
}
|
||||
if _, err := d.ExecContext(ctx, stmt); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func splitSQLStatements(s string) []string {
|
||||
var out []string
|
||||
var b strings.Builder
|
||||
inSingle, inDouble, inBacktick := false, false, false
|
||||
escaped := false
|
||||
for i := 0; i < len(s); i++ {
|
||||
c := s[i]
|
||||
if escaped {
|
||||
b.WriteByte(c)
|
||||
escaped = false
|
||||
continue
|
||||
}
|
||||
if c == '\\' && (inSingle || inDouble) {
|
||||
b.WriteByte(c)
|
||||
escaped = true
|
||||
continue
|
||||
}
|
||||
switch c {
|
||||
case '\'':
|
||||
if !inDouble && !inBacktick {
|
||||
inSingle = !inSingle
|
||||
}
|
||||
case '"':
|
||||
if !inSingle && !inBacktick {
|
||||
inDouble = !inDouble
|
||||
}
|
||||
case '`':
|
||||
if !inSingle && !inDouble {
|
||||
inBacktick = !inBacktick
|
||||
}
|
||||
case ';':
|
||||
if !inSingle && !inDouble && !inBacktick {
|
||||
out = append(out, b.String())
|
||||
b.Reset()
|
||||
continue
|
||||
}
|
||||
}
|
||||
b.WriteByte(c)
|
||||
}
|
||||
if strings.TrimSpace(b.String()) != "" {
|
||||
out = append(out, b.String())
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user