service.go 8.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291
  1. package auth
  2. import (
  3. "context"
  4. "crypto/rand"
  5. "crypto/sha256"
  6. "crypto/subtle"
  7. "encoding/base64"
  8. "errors"
  9. "fmt"
  10. "strings"
  11. "time"
  12. "golang.org/x/crypto/bcrypt"
  13. "vocat/internal/store"
  14. )
  15. var (
  16. ErrInvalidCredentials = errors.New("invalid credentials")
  17. ErrUnauthorized = errors.New("unauthorized")
  18. ErrInvalidCSRF = errors.New("invalid csrf token")
  19. )
  20. type Options struct {
  21. SessionTTL time.Duration
  22. BcryptCost int
  23. }
  24. type Service struct {
  25. store *store.Store
  26. sessionTTL time.Duration
  27. bcryptCost int
  28. dummyHash []byte
  29. }
  30. type Principal struct {
  31. ID int64 `json:"-"`
  32. Username string `json:"username"`
  33. }
  34. type Credentials struct {
  35. SessionToken string
  36. CSRFToken string
  37. ExpiresAt time.Time
  38. Principal Principal
  39. }
  40. type AuthenticatedSession struct {
  41. Principal Principal
  42. ExpiresAt time.Time
  43. tokenHash []byte
  44. csrfHash []byte
  45. }
  46. func New(database *store.Store, options Options) (*Service, error) {
  47. if database == nil {
  48. return nil, errors.New("auth: store is required")
  49. }
  50. if options.SessionTTL <= 0 {
  51. return nil, errors.New("auth: session TTL must be positive")
  52. }
  53. if options.BcryptCost == 0 {
  54. options.BcryptCost = 12
  55. }
  56. if options.BcryptCost < bcrypt.MinCost || options.BcryptCost > bcrypt.MaxCost {
  57. return nil, errors.New("auth: bcrypt cost is out of range")
  58. }
  59. dummyHash, err := bcrypt.GenerateFromPassword([]byte("not-a-real-password"), options.BcryptCost)
  60. if err != nil {
  61. return nil, fmt.Errorf("auth: generate timing hash: %w", err)
  62. }
  63. return &Service{
  64. store: database,
  65. sessionTTL: options.SessionTTL,
  66. bcryptCost: options.BcryptCost,
  67. dummyHash: dummyHash,
  68. }, nil
  69. }
  70. // EnsureAdmin configures the single administrator. Existing sessions are
  71. // revoked only when the configured username or password changes.
  72. func (s *Service) EnsureAdmin(ctx context.Context, username string, password string) error {
  73. username = strings.TrimSpace(username)
  74. current, err := s.store.CurrentAdmin(ctx)
  75. if err == nil &&
  76. current.Username == username &&
  77. bcrypt.CompareHashAndPassword(current.PasswordHash, []byte(password)) == nil {
  78. return nil
  79. }
  80. if err != nil && !errors.Is(err, store.ErrNotFound) {
  81. return fmt.Errorf("auth: read configured admin: %w", err)
  82. }
  83. passwordHash, err := bcrypt.GenerateFromPassword([]byte(password), s.bcryptCost)
  84. if err != nil {
  85. return fmt.Errorf("auth: hash admin password: %w", err)
  86. }
  87. if err := s.store.SetAdmin(ctx, username, passwordHash); err != nil {
  88. return err
  89. }
  90. return nil
  91. }
  92. func (s *Service) Login(ctx context.Context, username string, password string) (Credentials, error) {
  93. admin, err := s.store.AdminByUsername(ctx, strings.TrimSpace(username))
  94. if errors.Is(err, store.ErrNotFound) {
  95. _ = bcrypt.CompareHashAndPassword(s.dummyHash, []byte(password))
  96. return Credentials{}, ErrInvalidCredentials
  97. }
  98. if err != nil {
  99. return Credentials{}, fmt.Errorf("auth: find admin: %w", err)
  100. }
  101. if bcrypt.CompareHashAndPassword(admin.PasswordHash, []byte(password)) != nil {
  102. return Credentials{}, ErrInvalidCredentials
  103. }
  104. if err := s.store.DeleteExpiredSessions(ctx, time.Now()); err != nil {
  105. return Credentials{}, err
  106. }
  107. sessionToken, err := randomToken()
  108. if err != nil {
  109. return Credentials{}, err
  110. }
  111. csrfToken, err := randomToken()
  112. if err != nil {
  113. return Credentials{}, err
  114. }
  115. expiresAt := time.Now().UTC().Add(s.sessionTTL)
  116. if err := s.store.CreateSession(
  117. ctx,
  118. admin.ID,
  119. hashToken(sessionToken),
  120. hashToken(csrfToken),
  121. expiresAt,
  122. ); err != nil {
  123. return Credentials{}, err
  124. }
  125. return Credentials{
  126. SessionToken: sessionToken,
  127. CSRFToken: csrfToken,
  128. ExpiresAt: expiresAt,
  129. Principal: Principal{
  130. ID: admin.ID,
  131. Username: admin.Username,
  132. },
  133. }, nil
  134. }
  135. func (s *Service) Authenticate(ctx context.Context, sessionToken string) (AuthenticatedSession, error) {
  136. if sessionToken == "" {
  137. return AuthenticatedSession{}, ErrUnauthorized
  138. }
  139. tokenHash := hashToken(sessionToken)
  140. session, err := s.store.SessionByTokenHash(ctx, tokenHash)
  141. if errors.Is(err, store.ErrNotFound) {
  142. return AuthenticatedSession{}, ErrUnauthorized
  143. }
  144. if err != nil {
  145. return AuthenticatedSession{}, fmt.Errorf("auth: load session: %w", err)
  146. }
  147. if !session.ExpiresAt.After(time.Now().UTC()) {
  148. _ = s.store.DeleteSession(ctx, tokenHash)
  149. return AuthenticatedSession{}, ErrUnauthorized
  150. }
  151. return AuthenticatedSession{
  152. Principal: Principal{
  153. ID: session.Admin.ID,
  154. Username: session.Admin.Username,
  155. },
  156. ExpiresAt: session.ExpiresAt,
  157. tokenHash: tokenHash,
  158. csrfHash: session.CSRFHash,
  159. }, nil
  160. }
  161. // RotateCSRF replaces the session-bound CSRF value and returns the new raw
  162. // token. Only its SHA-256 digest is persisted.
  163. func (s *Service) RotateCSRF(ctx context.Context, sessionToken string) (AuthenticatedSession, string, error) {
  164. return s.CSRFToken(ctx, sessionToken, "")
  165. }
  166. // CSRFToken reuses a valid CSRF cookie or rotates it when the cookie is absent
  167. // or stale. Reuse prevents one browser tab from invalidating another tab's
  168. // session-bound token.
  169. func (s *Service) CSRFToken(
  170. ctx context.Context,
  171. sessionToken string,
  172. existingToken string,
  173. ) (AuthenticatedSession, string, error) {
  174. session, err := s.Authenticate(ctx, sessionToken)
  175. if err != nil {
  176. return AuthenticatedSession{}, "", err
  177. }
  178. if existingToken != "" {
  179. existingHash := hashToken(existingToken)
  180. if subtle.ConstantTimeCompare(existingHash, session.csrfHash) == 1 {
  181. return session, existingToken, nil
  182. }
  183. }
  184. csrfToken, err := randomToken()
  185. if err != nil {
  186. return AuthenticatedSession{}, "", err
  187. }
  188. csrfHash := hashToken(csrfToken)
  189. if err := s.store.UpdateSessionCSRF(ctx, session.tokenHash, csrfHash); err != nil {
  190. if errors.Is(err, store.ErrNotFound) {
  191. return AuthenticatedSession{}, "", ErrUnauthorized
  192. }
  193. return AuthenticatedSession{}, "", err
  194. }
  195. session.csrfHash = csrfHash
  196. return session, csrfToken, nil
  197. }
  198. func (s *Service) ValidateCSRF(
  199. ctx context.Context,
  200. sessionToken string,
  201. csrfToken string,
  202. ) (AuthenticatedSession, error) {
  203. if csrfToken == "" {
  204. return AuthenticatedSession{}, ErrInvalidCSRF
  205. }
  206. session, err := s.Authenticate(ctx, sessionToken)
  207. if err != nil {
  208. return AuthenticatedSession{}, err
  209. }
  210. providedHash := hashToken(csrfToken)
  211. if subtle.ConstantTimeCompare(providedHash, session.csrfHash) != 1 {
  212. return AuthenticatedSession{}, ErrInvalidCSRF
  213. }
  214. return session, nil
  215. }
  216. func (s *Service) Logout(ctx context.Context, sessionToken string) error {
  217. if sessionToken == "" {
  218. return nil
  219. }
  220. if err := s.store.DeleteSession(ctx, hashToken(sessionToken)); err != nil {
  221. return err
  222. }
  223. return nil
  224. }
  225. // ChangePassword verifies the current password, replaces it with a fresh
  226. // bcrypt hash and revokes every session through Store.SetAdmin.
  227. func (s *Service) ChangePassword(
  228. ctx context.Context,
  229. username string,
  230. currentPassword string,
  231. newPassword string,
  232. ) error {
  233. if len(newPassword) < 12 || len(newPassword) > 1024 {
  234. return errors.New("new password must contain between 12 and 1024 characters")
  235. }
  236. admin, err := s.store.AdminByUsername(ctx, strings.TrimSpace(username))
  237. if errors.Is(err, store.ErrNotFound) {
  238. _ = bcrypt.CompareHashAndPassword(s.dummyHash, []byte(currentPassword))
  239. return ErrInvalidCredentials
  240. }
  241. if err != nil {
  242. return fmt.Errorf("auth: find admin: %w", err)
  243. }
  244. if bcrypt.CompareHashAndPassword(admin.PasswordHash, []byte(currentPassword)) != nil {
  245. return ErrInvalidCredentials
  246. }
  247. if bcrypt.CompareHashAndPassword(admin.PasswordHash, []byte(newPassword)) == nil {
  248. return errors.New("new password must differ from the current password")
  249. }
  250. passwordHash, err := bcrypt.GenerateFromPassword([]byte(newPassword), s.bcryptCost)
  251. if err != nil {
  252. return fmt.Errorf("auth: hash new password: %w", err)
  253. }
  254. if err := s.store.SetAdmin(ctx, admin.Username, passwordHash); err != nil {
  255. return fmt.Errorf("auth: save new password: %w", err)
  256. }
  257. return nil
  258. }
  259. func randomToken() (string, error) {
  260. buffer := make([]byte, 32)
  261. if _, err := rand.Read(buffer); err != nil {
  262. return "", fmt.Errorf("auth: generate random token: %w", err)
  263. }
  264. return base64.RawURLEncoding.EncodeToString(buffer), nil
  265. }
  266. func hashToken(token string) []byte {
  267. digest := sha256.Sum256([]byte(token))
  268. return digest[:]
  269. }