88 lines
2.0 KiB
Go
88 lines
2.0 KiB
Go
package database
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"errors"
|
|
"net/url"
|
|
"os"
|
|
"path/filepath"
|
|
"strings"
|
|
|
|
_ "github.com/mattn/go-sqlite3"
|
|
)
|
|
|
|
const busyTimeoutMilliseconds = 5000
|
|
|
|
type safeError struct {
|
|
message string
|
|
cause error
|
|
}
|
|
|
|
func (e safeError) Error() string {
|
|
return e.message
|
|
}
|
|
|
|
func (e safeError) Unwrap() error {
|
|
return e.cause
|
|
}
|
|
|
|
func Open(ctx context.Context, path string) (*sql.DB, error) {
|
|
absolutePath, err := filepath.Abs(path)
|
|
if err != nil {
|
|
return nil, errors.New("resolve SQLite path")
|
|
}
|
|
if err := os.MkdirAll(filepath.Dir(absolutePath), 0o700); err != nil {
|
|
return nil, errors.New("create SQLite directory")
|
|
}
|
|
if info, err := os.Stat(absolutePath); err == nil && info.IsDir() {
|
|
return nil, errors.New("SQLite path is a directory")
|
|
} else if err != nil && !errors.Is(err, os.ErrNotExist) {
|
|
return nil, errors.New("inspect SQLite path")
|
|
}
|
|
|
|
db, err := sql.Open("sqlite3", dataSourceName(absolutePath))
|
|
if err != nil {
|
|
return nil, errors.New("initialize SQLite driver")
|
|
}
|
|
db.SetMaxOpenConns(1)
|
|
db.SetMaxIdleConns(1)
|
|
|
|
if err := db.PingContext(ctx); err != nil {
|
|
_ = db.Close()
|
|
return nil, safeError{
|
|
message: "connect to SQLite",
|
|
cause: err,
|
|
}
|
|
}
|
|
return db, nil
|
|
}
|
|
|
|
func dataSourceName(absolutePath string) string {
|
|
slashPath := filepath.ToSlash(absolutePath)
|
|
if filepath.VolumeName(absolutePath) != "" &&
|
|
!strings.HasPrefix(slashPath, "/") {
|
|
slashPath = "/" + slashPath
|
|
}
|
|
dsnURL := &url.URL{
|
|
Scheme: "file",
|
|
Path: slashPath,
|
|
}
|
|
query := dsnURL.Query()
|
|
query.Set("_busy_timeout", "5000")
|
|
query.Set("_foreign_keys", "on")
|
|
query.Set("_journal_mode", "WAL")
|
|
query.Set("_txlock", "immediate")
|
|
dsnURL.RawQuery = query.Encode()
|
|
return dsnURL.String()
|
|
}
|
|
|
|
func isSafeDataSourceName(value string) bool {
|
|
lowerValue := strings.ToLower(value)
|
|
return strings.HasPrefix(lowerValue, "file:") &&
|
|
strings.Contains(lowerValue, "_busy_timeout=") &&
|
|
strings.Contains(lowerValue, "_foreign_keys=") &&
|
|
strings.Contains(lowerValue, "_journal_mode=wal") &&
|
|
strings.Contains(lowerValue, "_txlock=immediate")
|
|
}
|