store.go 8.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326
  1. package store
  2. import (
  3. "context"
  4. "database/sql"
  5. "errors"
  6. "fmt"
  7. "os"
  8. "path/filepath"
  9. "strings"
  10. "time"
  11. _ "modernc.org/sqlite"
  12. )
  13. const schemaVersion = 6
  14. var ErrNotFound = errors.New("store: not found")
  15. // Store owns the SQLite connection used by the process.
  16. type Store struct {
  17. db *sql.DB
  18. }
  19. type Admin struct {
  20. ID int64
  21. Username string
  22. PasswordHash []byte
  23. CreatedAt time.Time
  24. UpdatedAt time.Time
  25. }
  26. type Session struct {
  27. TokenHash []byte
  28. CSRFHash []byte
  29. ExpiresAt time.Time
  30. CreatedAt time.Time
  31. Admin Admin
  32. }
  33. // Open creates the parent directory, opens SQLite, applies safety pragmas and
  34. // runs the built-in schema migration.
  35. func Open(ctx context.Context, path string) (*Store, error) {
  36. if err := prepareDatabasePath(path); err != nil {
  37. return nil, err
  38. }
  39. db, err := sql.Open("sqlite", path)
  40. if err != nil {
  41. return nil, fmt.Errorf("open sqlite: %w", err)
  42. }
  43. db.SetMaxOpenConns(1)
  44. db.SetMaxIdleConns(1)
  45. closeOnError := func(err error) (*Store, error) {
  46. _ = db.Close()
  47. return nil, err
  48. }
  49. if err := db.PingContext(ctx); err != nil {
  50. return closeOnError(fmt.Errorf("ping sqlite: %w", err))
  51. }
  52. for _, pragma := range []string{
  53. "PRAGMA foreign_keys = ON",
  54. "PRAGMA busy_timeout = 5000",
  55. "PRAGMA journal_mode = WAL",
  56. } {
  57. if _, err := db.ExecContext(ctx, pragma); err != nil {
  58. return closeOnError(fmt.Errorf("%s: %w", pragma, err))
  59. }
  60. }
  61. if err := migrate(ctx, db); err != nil {
  62. return closeOnError(err)
  63. }
  64. if isFilesystemPath(path) {
  65. if err := os.Chmod(path, 0o600); err != nil {
  66. return closeOnError(fmt.Errorf("secure sqlite file: %w", err))
  67. }
  68. }
  69. return &Store{db: db}, nil
  70. }
  71. func prepareDatabasePath(path string) error {
  72. if !isFilesystemPath(path) {
  73. return nil
  74. }
  75. parent := filepath.Dir(path)
  76. if parent == "." {
  77. return nil
  78. }
  79. if err := os.MkdirAll(parent, 0o750); err != nil {
  80. return fmt.Errorf("create sqlite directory: %w", err)
  81. }
  82. return nil
  83. }
  84. func isFilesystemPath(path string) bool {
  85. return path != ":memory:" && !strings.HasPrefix(path, "file:")
  86. }
  87. func migrate(ctx context.Context, db *sql.DB) error {
  88. var version int
  89. if err := db.QueryRowContext(ctx, "PRAGMA user_version").Scan(&version); err != nil {
  90. return fmt.Errorf("read sqlite schema version: %w", err)
  91. }
  92. if version > schemaVersion {
  93. return fmt.Errorf("sqlite schema version %d is newer than supported version %d", version, schemaVersion)
  94. }
  95. if version == schemaVersion {
  96. return nil
  97. }
  98. for nextVersion := version + 1; nextVersion <= schemaVersion; nextVersion++ {
  99. tx, err := db.BeginTx(ctx, nil)
  100. if err != nil {
  101. return fmt.Errorf("begin sqlite migration %d: %w", nextVersion, err)
  102. }
  103. for _, statement := range migrationStatements(nextVersion) {
  104. if _, err := tx.ExecContext(ctx, statement); err != nil {
  105. _ = tx.Rollback()
  106. return fmt.Errorf("apply sqlite migration %d: %w", nextVersion, err)
  107. }
  108. }
  109. if _, err := tx.ExecContext(
  110. ctx,
  111. fmt.Sprintf("PRAGMA user_version = %d", nextVersion),
  112. ); err != nil {
  113. _ = tx.Rollback()
  114. return fmt.Errorf("record sqlite migration %d: %w", nextVersion, err)
  115. }
  116. if err := tx.Commit(); err != nil {
  117. return fmt.Errorf("commit sqlite migration %d: %w", nextVersion, err)
  118. }
  119. }
  120. return nil
  121. }
  122. func (s *Store) Close() error {
  123. return s.db.Close()
  124. }
  125. func (s *Store) Ready(ctx context.Context) error {
  126. if err := s.db.PingContext(ctx); err != nil {
  127. return err
  128. }
  129. var one int
  130. if err := s.db.QueryRowContext(ctx, "SELECT 1").Scan(&one); err != nil {
  131. return err
  132. }
  133. if one != 1 {
  134. return errors.New("sqlite readiness query returned an unexpected value")
  135. }
  136. return nil
  137. }
  138. func (s *Store) CurrentAdmin(ctx context.Context) (Admin, error) {
  139. return scanAdmin(s.db.QueryRowContext(ctx, `
  140. SELECT id, username, password_hash, created_at, updated_at
  141. FROM admins
  142. WHERE id = 1
  143. `))
  144. }
  145. func (s *Store) AdminByUsername(ctx context.Context, username string) (Admin, error) {
  146. return scanAdmin(s.db.QueryRowContext(ctx, `
  147. SELECT id, username, password_hash, created_at, updated_at
  148. FROM admins
  149. WHERE username = ?
  150. `, username))
  151. }
  152. type rowScanner interface {
  153. Scan(dest ...any) error
  154. }
  155. func scanAdmin(row rowScanner) (Admin, error) {
  156. var admin Admin
  157. var createdAt int64
  158. var updatedAt int64
  159. err := row.Scan(
  160. &admin.ID,
  161. &admin.Username,
  162. &admin.PasswordHash,
  163. &createdAt,
  164. &updatedAt,
  165. )
  166. if errors.Is(err, sql.ErrNoRows) {
  167. return Admin{}, ErrNotFound
  168. }
  169. if err != nil {
  170. return Admin{}, err
  171. }
  172. admin.CreatedAt = time.Unix(createdAt, 0).UTC()
  173. admin.UpdatedAt = time.Unix(updatedAt, 0).UTC()
  174. return admin, nil
  175. }
  176. // SetAdmin inserts or replaces the single configured administrator and
  177. // atomically revokes all existing sessions.
  178. func (s *Store) SetAdmin(ctx context.Context, username string, passwordHash []byte) error {
  179. tx, err := s.db.BeginTx(ctx, nil)
  180. if err != nil {
  181. return fmt.Errorf("begin admin update: %w", err)
  182. }
  183. defer tx.Rollback()
  184. now := time.Now().UTC().Unix()
  185. _, err = tx.ExecContext(ctx, `
  186. INSERT INTO admins (id, username, password_hash, created_at, updated_at)
  187. VALUES (1, ?, ?, ?, ?)
  188. ON CONFLICT(id) DO UPDATE SET
  189. username = excluded.username,
  190. password_hash = excluded.password_hash,
  191. updated_at = excluded.updated_at
  192. `, username, passwordHash, now, now)
  193. if err != nil {
  194. return fmt.Errorf("set admin: %w", err)
  195. }
  196. if _, err := tx.ExecContext(ctx, "DELETE FROM sessions"); err != nil {
  197. return fmt.Errorf("revoke sessions after admin update: %w", err)
  198. }
  199. if err := tx.Commit(); err != nil {
  200. return fmt.Errorf("commit admin update: %w", err)
  201. }
  202. return nil
  203. }
  204. func (s *Store) DeleteAllSessions(ctx context.Context) error {
  205. if _, err := s.db.ExecContext(ctx, "DELETE FROM sessions"); err != nil {
  206. return fmt.Errorf("delete all sessions: %w", err)
  207. }
  208. return nil
  209. }
  210. func (s *Store) CreateSession(
  211. ctx context.Context,
  212. adminID int64,
  213. tokenHash []byte,
  214. csrfHash []byte,
  215. expiresAt time.Time,
  216. ) error {
  217. now := time.Now().UTC().Unix()
  218. _, err := s.db.ExecContext(ctx, `
  219. INSERT INTO sessions (token_hash, admin_id, csrf_hash, expires_at, created_at)
  220. VALUES (?, ?, ?, ?, ?)
  221. `, tokenHash, adminID, csrfHash, expiresAt.UTC().Unix(), now)
  222. if err != nil {
  223. return fmt.Errorf("create session: %w", err)
  224. }
  225. return nil
  226. }
  227. func (s *Store) SessionByTokenHash(ctx context.Context, tokenHash []byte) (Session, error) {
  228. var session Session
  229. var expiresAt int64
  230. var createdAt int64
  231. var adminCreatedAt int64
  232. var adminUpdatedAt int64
  233. err := s.db.QueryRowContext(ctx, `
  234. SELECT
  235. s.token_hash,
  236. s.csrf_hash,
  237. s.expires_at,
  238. s.created_at,
  239. a.id,
  240. a.username,
  241. a.created_at,
  242. a.updated_at
  243. FROM sessions s
  244. JOIN admins a ON a.id = s.admin_id
  245. WHERE s.token_hash = ?
  246. `, tokenHash).Scan(
  247. &session.TokenHash,
  248. &session.CSRFHash,
  249. &expiresAt,
  250. &createdAt,
  251. &session.Admin.ID,
  252. &session.Admin.Username,
  253. &adminCreatedAt,
  254. &adminUpdatedAt,
  255. )
  256. if errors.Is(err, sql.ErrNoRows) {
  257. return Session{}, ErrNotFound
  258. }
  259. if err != nil {
  260. return Session{}, err
  261. }
  262. session.ExpiresAt = time.Unix(expiresAt, 0).UTC()
  263. session.CreatedAt = time.Unix(createdAt, 0).UTC()
  264. session.Admin.CreatedAt = time.Unix(adminCreatedAt, 0).UTC()
  265. session.Admin.UpdatedAt = time.Unix(adminUpdatedAt, 0).UTC()
  266. return session, nil
  267. }
  268. func (s *Store) UpdateSessionCSRF(ctx context.Context, tokenHash []byte, csrfHash []byte) error {
  269. result, err := s.db.ExecContext(ctx, `
  270. UPDATE sessions
  271. SET csrf_hash = ?
  272. WHERE token_hash = ?
  273. `, csrfHash, tokenHash)
  274. if err != nil {
  275. return fmt.Errorf("update session csrf: %w", err)
  276. }
  277. affected, err := result.RowsAffected()
  278. if err != nil {
  279. return fmt.Errorf("read session update result: %w", err)
  280. }
  281. if affected == 0 {
  282. return ErrNotFound
  283. }
  284. return nil
  285. }
  286. func (s *Store) DeleteSession(ctx context.Context, tokenHash []byte) error {
  287. if _, err := s.db.ExecContext(ctx, "DELETE FROM sessions WHERE token_hash = ?", tokenHash); err != nil {
  288. return fmt.Errorf("delete session: %w", err)
  289. }
  290. return nil
  291. }
  292. func (s *Store) DeleteExpiredSessions(ctx context.Context, now time.Time) error {
  293. if _, err := s.db.ExecContext(ctx, "DELETE FROM sessions WHERE expires_at <= ?", now.UTC().Unix()); err != nil {
  294. return fmt.Errorf("delete expired sessions: %w", err)
  295. }
  296. return nil
  297. }