package db
import (
"database/sql"
_ "embed"
"errors"
"fmt"
"log/slog"
"strconv"
"strings"
"sync"
"time"
_ "modernc.org/sqlite"
"gopkg.awl.red/lx/internal/auth"
"gopkg.awl.red/lx/internal/config"
"gopkg.awl.red/lx/internal/log"
)
type DB struct {
db *sql.DB
muCfg sync.RWMutex
cfg *config.Config
chAddrChange chan config.AddrChange
// analytics db
adb *sql.DB
}
type ConfigCacher interface {
GetConfig() *config.Config
}
var (
ErrNoAdminRole = errors.New("could not find admin role")
ErrNotBusinessOwner = errors.New("could not do action: user does not own the business")
ErrMismatchedColsVals = errors.New("mismatched number of columns and values to update on table")
ErrBackHalfExpired = errors.New("backhalf expired")
ErrBackHalf404 = errors.New("link not found")
)
//go:embed migrations/001_initial.up.sql
var m_001 []byte
var migrations = [][]byte{
nil,
m_001,
}
//go:embed analyticsmigrations/001_initial.up.sql
var am_001 []byte
var amigrations = [][]byte{
nil,
am_001,
}
func New(dsn string, adsn string, chAddrChange chan config.AddrChange) (*DB, error) {
_ = log.New(slog.LevelInfo, false, "lvl", "D|I|W|E", "unix", "@", true)
sqlDb, err := sql.Open("sqlite", dsn)
if err != nil {
return nil, err
}
db := &DB{db: sqlDb, chAddrChange: chAddrChange}
if err := db.migrate(db.db, migrations); err != nil {
return nil, err
}
if err := db.getConfig(); err != nil {
return nil, err
}
go db.reloadConfigLoop()
cfg := db.GetConfig()
if !cfg.AnalyticsDisable {
sqlDb, err := sql.Open("sqlite", adsn)
if err != nil {
return nil, err
}
db.adb = sqlDb
if err := db.migrate(db.adb, amigrations); err != nil {
return nil, err
}
}
return db, nil
}
func (d *DB) migrate(db *sql.DB, steps [][]byte) error {
var version int
var dirty bool
row := db.QueryRow("SELECT version, dirty FROM version;")
err := row.Scan(&version, &dirty)
if err != nil {
slog.Debug("could not get initial migration version", "err", err)
}
if dirty {
return fmt.Errorf("migration dirty: on version %d", version)
}
// e.g. len is 2 if on version 1, and there is nothing more to do
if len(steps) == version+1 {
slog.Info("no steps to run", "version", version)
return nil
}
if version+1 > len(steps)-1 {
return fmt.Errorf("unknown future schema version: observed: %d, last_in_code: %d", version, len(steps)-1)
}
for i, migration := range steps[version+1:] {
slog.Info("running migration", "migration_char_len", len(migration), "version", i+1)
_, err := db.Exec(string(migration))
if err != nil {
_, errDirty := db.Exec(`UPDATE version SET version=$1, dirty=true`, i+1)
if downMigration, ok := getDownMigration(string(migration)); ok {
slog.Error("failed to run migration, you can run a down migration by manually executing the following on the db:", "err", err, "err_dirty", errDirty)
fmt.Println(downMigration)
}
return errors.Join(err, errDirty)
}
}
return nil
}
func (d *DB) getConfig() error {
siteConfigSql := `SELECT
title,
support_email,
support_phone,
purpose,
disclaimer,
max_default_businesses,
listen_address,
port,
signup_needs_approve,
session_length_max,
email_confirm,
email_required,
email_regex,
password_salt_length,
password_key_length,
password_rounds,
password_hash_strategy,
password_attempts,
password_min,
password_regex,
password_help,
password_required,
log_add_source,
log_level,
log_level_key,
log_level_overrides,
log_time_format,
log_time_key,
analytics_disable,
analytics_write_timeout,
analytics_read_timeout,
person_style_text_color,
person_style_text_dark_color,
person_style_background_color,
person_style_background_dark_color,
merchant_default,
config_reload_seconds
FROM site_config;`
row := d.db.QueryRow(siteConfigSql)
cfg := &config.Config{}
err := row.Scan(
&cfg.Title, &cfg.SupportEmail, &cfg.SupportPhone, &cfg.Purpose, &cfg.Disclaimer,
&cfg.MaxDefaultBusinesses, &cfg.ListenAddress, &cfg.Port, &cfg.SignupNeedsApprove,
&cfg.SessionLengthMax,
&cfg.EmailConfirm, &cfg.EmailRequired, &cfg.EmailRegex,
&cfg.PasswordSaltLength, &cfg.PasswordKeyLength, &cfg.PasswordRounds, &cfg.PasswordHashStrategy,
&cfg.PasswordAttempts, &cfg.PasswordMin, &cfg.PasswordRegex, &cfg.PasswordHelp, &cfg.PasswordRequired,
&cfg.LogAddSource, &cfg.LogLevel, &cfg.LogLevelKey, &cfg.LogLevelOverrides, &cfg.LogTimeFormat, &cfg.LogTimeKey,
&cfg.AnalyticsDisable, &cfg.AnalyticsWriteTimeout, &cfg.AnalyticsReadTimeout,
&cfg.PersonStyleTextColor,
&cfg.PersonStyleTextDarkColor,
&cfg.PersonStyleBackgroundColor,
&cfg.PersonStyleBackgroundDarkColor,
&cfg.MerchantDefault,
&cfg.ConfigReloadSeconds,
)
if err != nil {
return err
}
businessTiersSql := `SELECT name, display_name, feature_level,
backhalves_anonymous, backhalves_named, folders,
analytics_retention, analytics_max_granularity, analytics_view_limit,
price_by_year, price_by_month
FROM business_tier;`
businessTiers := map[string]*config.BusinessTier{}
rows, err := d.db.Query(businessTiersSql)
if err != nil {
return err
}
for rows.Next() {
b := &config.BusinessTier{}
err = rows.Scan(
&b.Name, &b.DisplayName, &b.FeatureLevel,
&b.BackhalvesAnonymous, &b.BackhalvesNamed, &b.Folders,
&b.AnalyticsRetention, &b.AnalyticsMaxGranularity, &b.AnalyticsViewLimit,
&b.PriceByYear, &b.PriceByMonth,
)
if err != nil {
return err
}
businessTiers[b.Name] = b
}
err = rows.Err()
if err != nil {
return err
}
cfg.BusinessTiers = businessTiers
merchantSql := `SELECT m.name, m.api_url,
group_concat(mbt.merchant_business_tier_id_string),
group_concat(bt.name)
FROM merchant m
LEFT JOIN merchant_business_tier mbt ON m.id = mbt.merchant_id
LEFT JOIN business_tier bt ON bt.id = mbt.business_tier_id
;`
merchants := map[string]*config.Merchant{}
rows, err = d.db.Query(merchantSql)
if err != nil {
return err
}
for rows.Next() {
m := &config.Merchant{
TierPlanIds: map[string]string{},
PlanIdTiers: map[string]string{},
}
var planIdsStr string
var tiersStr string
err = rows.Scan(
&m.Name, &m.ApiUrl, &planIdsStr, &tiersStr,
)
if err != nil {
return err
}
tiers := strings.Split(tiersStr, ",")
planIds := strings.Split(planIdsStr, ",")
for i := range tiers {
m.TierPlanIds[tiers[i]] = planIds[i]
m.PlanIdTiers[planIds[i]] = tiers[i]
}
merchants[m.Name] = m
}
err = rows.Err()
if err != nil {
return err
}
cfg.Merchants = merchants
// siteUsageLimitsSql := `SELECT name, noun, verb, subject, extra, extra_type, windows FROM site_usage_limits;`
// usagePolicies := map[string]*config.UsagePolicy{}
// rows, err := d.db.Query(siteUsageLimitsSql)
// if err != nil {
// return err
// }
// for rows.Next() {
// up := &config.UsagePolicy{}
// err = rows.Scan(&up.Name, &up.Noun, &up.Verb, &up.Subject, &up.Extra, &up.ExtraType, &up.Windows)
// if err != nil {
// return err
// }
// usagePolicies[up.Name] = up
// }
// err = rows.Err()
// if err != nil {
// return err
// }
// cfg.UsagePolicies = usagePolicies
d.muCfg.Lock()
defer d.muCfg.Unlock()
d.cfg = cfg
return nil
}
func (d *DB) GetConfig() *config.Config {
d.muCfg.RLock()
if d.cfg != nil {
defer d.muCfg.RUnlock()
return d.cfg
}
d.muCfg.RUnlock()
err := d.getConfig()
if err != nil {
slog.Error("could not get config", "err", err)
}
d.muCfg.RLock()
defer d.muCfg.RUnlock()
return d.cfg
}
func (d *DB) reloadConfigLoop() {
cfgOld := d.GetConfig()
i := 0
for {
var cfg *config.Config
if i == 0 {
cfg = cfgOld
} else {
d.getConfig()
cfg = d.GetConfig()
if cfgOld != nil && cfg != nil {
doLogChange(cfgOld, cfg)
doListenChange(cfgOld, cfg, d.chAddrChange)
}
cfgOld = cfg
}
if cfg == nil || cfg.ConfigReloadSeconds <= 0 {
slog.Info("config reloader exited, config nil or config_reload_seconds not positive")
return
}
time.Sleep(time.Second * time.Duration(cfg.ConfigReloadSeconds))
i += 1
}
}
func getDownMigration(migration string) (string, bool) {
_, downComment, hasDownMigration := strings.Cut(string(migration), "\n-- DOWN:")
if !hasDownMigration {
return "", false
}
downMigration := strings.Join(strings.Split(downComment, "\n-- "), "\n")
return downMigration, true
}
func doLogChange(a *config.Config, b *config.Config) {
logDelta := []any{}
if a.LogAddSource != b.LogAddSource {
logDelta = append(logDelta,
"add_source_old", a.LogAddSource,
"add_source_new", b.LogAddSource,
)
}
if a.LogLevel != b.LogLevel {
logDelta = append(logDelta,
"level_old", a.LogLevel,
"level_new", b.LogLevel,
)
}
if a.LogLevelKey != b.LogLevelKey {
logDelta = append(logDelta,
"level_key_old", a.LogLevelKey,
"level_key_new", b.LogLevelKey,
)
}
if a.LogLevelOverrides != b.LogLevelOverrides {
logDelta = append(logDelta,
"level_overrides_old", a.LogLevelOverrides,
"level_overrides_new", b.LogLevelOverrides,
)
}
if a.LogTimeFormat != b.LogTimeFormat {
logDelta = append(logDelta,
"time_format_old", a.LogTimeFormat,
"time_format_new", b.LogTimeFormat,
)
}
if a.LogTimeKey != b.LogTimeKey {
logDelta = append(logDelta,
"time_key_old", a.LogTimeKey,
"time_key_new", b.LogTimeKey,
)
}
if len(logDelta) > 0 {
log.SetFromConfig(b)
slog.Info("logger changed", logDelta...)
}
}
func doListenChange(a *config.Config, b *config.Config, chAddrChange chan config.AddrChange) {
logDelta := []any{}
if a.ListenAddress != b.ListenAddress {
logDelta = append(logDelta,
"listen_address_old", a.ListenAddress,
"listen_address_new", b.ListenAddress,
)
}
if a.Port != b.Port {
logDelta = append(logDelta,
"port_old", a.Port,
"port_new", b.Port,
)
}
if len(logDelta) > 0 {
// TODO: signal server to stop listening
chAddrChange <- config.AddrChange{Address: b.ListenAddress, Port: b.Port}
slog.Info("server listen changed", logDelta...)
}
}
func (d *DB) GetSession(token string, cfg *config.Config) (*auth.Session, error) {
sessionSql := `SELECT s.person_id, s.expires_at, p.approved, r.name
FROM person_session s
JOIN person_role pr ON s.person_id = pr.person_id
JOIN person p ON s.person_id = p.id
JOIN role r ON pr.role_id = r.id
WHERE s.token = $1;
`
rows, err := d.db.Query(sessionSql, token)
if err != nil {
return nil, err
}
var personId int
var expires int
session := &auth.Session{
Roles: map[string]struct{}{},
}
var approved bool
for rows.Next() {
var roleName string
err = rows.Scan(&personId, &expires, &approved, &roleName)
if err != nil {
return nil, err
}
if err := rows.Err(); err != nil {
return nil, err
}
if !approved {
return nil, auth.ErrUserNotApproved
}
_, isRecognizedRole := auth.Roles[roleName]
if isRecognizedRole {
session.Roles[roleName] = struct{}{}
} else {
slog.Warn("unrecognized role", "role", roleName, "person_id", personId)
}
}
session.ExpiresAt = expires
session.PersonId = personId
expired := time.Now().UTC().After(time.Unix(int64(expires), 0).UTC())
if expired {
return nil, auth.ErrSessionExpired
}
return session, nil
}
type BackHalf struct {
Id int
BusinessId int
Redirect string
}
func (d *DB) GetBackHalf(cfg *config.Config, domain, name string) (BackHalf, error) {
bhSql := `SELECT b.id, b.business_id, b.redirect, b.expire_at
FROM back_half b
JOIN domain d ON d.name = $1
WHERE b.name = $2;
`
var expire *time.Time
slog.Info("to get domain", "sql", bhSql, "domain", domain, "name", name)
row := d.db.QueryRow(bhSql, domain, name)
bh := BackHalf{}
err := row.Scan(&bh.Id, &bh.BusinessId, &bh.Redirect, &expire)
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return bh, ErrBackHalf404
}
}
if expire != nil && expire.Before(time.Now()) {
return bh, ErrBackHalfExpired
}
return bh, err
}
func (d *DB) SetConfig(cols []string, vals []any) error {
if len(cols) != len(vals) {
return ErrMismatchedColsVals
}
if len(cols) == 0 {
return nil
}
pl := mkPl(1)
var updateSqlB strings.Builder
updateSqlB.WriteString(`UPDATE site_config SET `)
for i, c := range cols {
if i != 0 {
updateSqlB.WriteRune(',')
}
updateSqlB.WriteString(c)
updateSqlB.WriteRune('=')
updateSqlB.WriteString(pl())
}
updateSqlB.WriteRune(';')
updateSql := updateSqlB.String()
_, err := d.db.Exec(updateSql, vals...)
if err != nil {
return err
}
old := d.GetConfig()
err = d.getConfig()
if err != nil {
return err
}
newCfg := d.GetConfig()
doLogChange(old, newCfg)
doListenChange(old, newCfg, d.chAddrChange)
return nil
}
func mkPl(start int) func() string {
i := start
return func() string {
pl := "$" + strconv.Itoa(i)
i += 1
return pl
}
}
func (d *DB) DoLogin(username, pw, email string, cfg *config.Config) (string, *auth.Session, error) {
personSql := `SELECT p.id, p.username, p.password, p.email,
p.email_confirmed, p.approved, p.password_attempts_remaining
FROM person p WHERE username=$1 OR email=$2;
`
rows, err := d.db.Query(personSql, username, email)
if err != nil {
return "", nil, err
}
var personId int
var theUser, hashed, theEmail string
var confirmed, approved bool
var remainingAttempts *int
i := 0
for rows.Next() {
if i > 0 {
return "", nil, auth.ErrTooManyUsers
}
err := rows.Scan(&personId, &theUser, &hashed, &theEmail, &confirmed, &approved, &remainingAttempts)
if err != nil {
return "", nil, err
}
if err := rows.Err(); err != nil {
return "", nil, err
}
i += 1
}
if remainingAttempts != nil && *remainingAttempts < 0 {
return "", nil, auth.ErrUserPasswordLocked
}
err = auth.PwMatchesHash(pw, hashed)
if err != nil {
if errors.Is(err, auth.ErrHashMismatch) && remainingAttempts != nil {
_, _ = d.db.Exec(`UPDATE person p SET p.password_attempts_remaining = p.password_attempts_remaining - 1 WHERE p.id = $1;`)
}
return "", nil, err
}
token, _, err := d.InsertSession(personId, cfg)
if err != nil {
return "", nil, err
}
sess, err := d.GetSession(token, cfg)
if err != nil {
return "", nil, err
}
return token, sess, nil
}
func (d *DB) InsertSession(userId int, cfg *config.Config) (string, int, error) {
expireAt := int(time.Now().UTC().Add(time.Duration(cfg.SessionLengthMax) * time.Second).Unix())
token := auth.NewToken()
sessionSql := `INSERT INTO person_session (person_id, expires_at, token) VALUES($1,$2,$3);`
_, err := d.db.Exec(sessionSql, userId, expireAt, token)
return token, expireAt, err
}
func (d *DB) DeleteSession(token string) error {
_, err := d.db.Exec(`DELETE FROM person_session WHERE token = $1;`, token)
return err
}