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
}