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 }