Viewing:
package db import ( "database/sql" _ "embed" "errors" "fmt" "log/slog" "strconv" "strings" "sync" "time" _ "gopkg.awl.red/phb" _ "modernc.org/sqlite" "gopkg.awl.red/bizdex/auth" "gopkg.awl.red/bizdex/config" "gopkg.awl.red/bizdex/log" ) type ConfigCacher interface { GetConfig() *config.Config } var ( ErrNoAdminRole = errors.New("could not find admin role") ) const ( RoleBusinessOwner = "site_admin" RoleNewPersonApprover = "new_person_approver" RoleAdmin = "site_admin" ) //go:embed migrations/001_initial.up.sql var m_001 []byte var migrations = [][]byte{ nil, m_001, } type DB struct { db *sql.DB muCfg sync.RWMutex cfg *config.Config chAddrChange chan config.AddrChange } func New(dsn 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(); err != nil { return nil, err } if err := db.getConfig(); err != nil { return nil, err } go db.reloadConfigLoop() return db, nil } func (d *DB) migrate() error { var version int var dirty bool row := d.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(migrations) == version+1 { slog.Info("no migrations to run", "version", version) return nil } if version+1 > len(migrations)-1 { return fmt.Errorf("unknown future schema version: observed: %d, last_in_code: %d", version, len(migrations)-1) } for i, migration := range migrations[version+1:] { slog.Info("running migration", "migration_char_len", len(migration), "version", i+1) _, err := d.db.Exec(string(migration)) if err != nil { _, errDirty := d.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:") fmt.Println(downMigration) } return errors.Join(err, errDirty) } } return nil } 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 (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 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 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 (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 (d *DB) getConfig() error { siteConfigSql := `SELECT title, support_email, support_phone, purpose, disclaimer, max_default_businesses, listen_address, port, signup_needs_approve, locale_map_pic, 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, person_style_text_color, person_style_text_dark_color, person_style_background_color, person_style_background_dark_color, 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.LocaleMapPic, &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.PersonStyleTextColor, &cfg.PersonStyleTextDarkColor, &cfg.PersonStyleBackgroundColor, &cfg.PersonStyleBackgroundDarkColor, &cfg.ConfigReloadSeconds, ) if err != nil { return err } 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) 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 } d.muCfg.Lock() defer d.muCfg.Unlock() d.cfg = cfg return nil } func mkPl(start int) func() string { i := start return func() string { pl := "$" + strconv.Itoa(i) i += 1 return pl } } func (d *DB) InsertUser(req *auth.VettedSignupRequest, cfg *config.Config) (int, error) { personCount := 0 personCountRow := d.db.QueryRow(`SELECT count(*) FROM person;`) err := personCountRow.Scan(&personCount) if err != nil { return 0, err } cols := []string{ "username", "firstname", "lastname", } vals := []any{ req.Username, req.Firstname, req.Lastname, } userPlaceholders := []string{"$1", "$2", "$3"} nextPl := mkPl(4) if !cfg.SignupNeedsApprove || personCount == 0 { cols = append(cols, "approved") vals = append(vals, true) userPlaceholders = append(userPlaceholders, nextPl()) } if req.Hashed != "" { cols = append(cols, "password") vals = append(vals, req.Hashed) userPlaceholders = append(userPlaceholders, nextPl()) } insertUserSql := fmt.Sprintf(`INSERT INTO person (%s) VALUES (%s) RETURNING id;`, strings.Join(cols, ", "), strings.Join(userPlaceholders, ", "), ) tx, err := d.db.Begin() if err != nil { return 0, err } defer tx.Rollback() // slog.Info("signup", "sql", insertUserSql) insertUserRow := tx.QueryRow(insertUserSql, vals...) newId := 0 err = insertUserRow.Scan(&newId) if err != nil { return 0, err } roleMap := map[string]int{} rows, err := d.db.Query(`SELECT name, id FROM role;`) if err != nil { return 0, err } for rows.Next() { var k string var v int err = rows.Scan(&k, &v) if err != nil { return 0, err } roleMap[k] = v } if personCount == 0 { insertAdminRoleSql := `INSERT INTO person_role (person_id, role_id) VALUES ($1, $2);` adminRoleId, ok := roleMap[RoleAdmin] if !ok { return 0, ErrNoAdminRole } _, err := tx.Exec(insertAdminRoleSql, newId, adminRoleId) if err != nil { return 0, err } // first user becomes admin, and does not go through person_request process return newId, nil } personRequestCols := []string{ "person_id", "request_type", "response_status", } personRequestPlaceholders := []string{ "$1", "$2", "$3", } personRequestVals := []any{ newId, auth.PersonRequestTypeUserApprove, "requested", } if req.RequestText != "" { personRequestCols = append(personRequestCols, "request_text") personRequestPlaceholders = append(personRequestPlaceholders, "$4") personRequestVals = append(personRequestVals, req.RequestText) } personRequestSql := fmt.Sprintf( `INSERT INTO person_request (%s) VALUES (%s);`, strings.Join(personRequestCols, ", "), strings.Join(personRequestPlaceholders, ", "), ) _, err = tx.Exec(personRequestSql, personRequestVals...) if err != nil { return 0, err } return newId, tx.Commit() } 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, expire_at, token) VALUES($1,$2,$3);` _, err := d.db.Exec(sessionSql, userId, expireAt, token) return token, expireAt, err }