88 lines
2.5 KiB
Go
88 lines
2.5 KiB
Go
package store
|
|
|
|
import (
|
|
"database/sql"
|
|
"errors"
|
|
"strconv"
|
|
)
|
|
|
|
func fmtSscanInt(s string, out *int64) (int, error) {
|
|
v, err := strconv.ParseInt(s, 10, 64)
|
|
if err != nil {
|
|
return 0, err
|
|
}
|
|
*out = v
|
|
return 1, nil
|
|
}
|
|
|
|
// Session 会话行(token 明文不入库,只存 SHA-256)。
|
|
type Session struct {
|
|
TokenHash string
|
|
CSRFToken string
|
|
CreatedAt int64
|
|
ExpiresAt int64
|
|
}
|
|
|
|
// CreateSession 插入会话行。
|
|
func (s *Store) CreateSession(tokenHash, csrfToken string, createdAt, expiresAt int64) error {
|
|
_, err := s.db.Exec(
|
|
`INSERT INTO sessions (token_hash, csrf_token, created_at, expires_at) VALUES (?, ?, ?, ?)`,
|
|
tokenHash, csrfToken, createdAt, expiresAt)
|
|
return err
|
|
}
|
|
|
|
// RotateSession 重建会话行:新 token、**csrf_token 保持不变**(§7.3-5)、新过期时间。
|
|
// 原子替换,避免窗口期内两行并存。
|
|
func (s *Store) RotateSession(oldTokenHash, newTokenHash, csrfToken string, createdAt, expiresAt int64) error {
|
|
tx, err := s.db.Begin()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer tx.Rollback()
|
|
if _, err := tx.Exec(`DELETE FROM sessions WHERE token_hash=?`, oldTokenHash); err != nil {
|
|
return err
|
|
}
|
|
if _, err := tx.Exec(
|
|
`INSERT INTO sessions (token_hash, csrf_token, created_at, expires_at) VALUES (?, ?, ?, ?)`,
|
|
newTokenHash, csrfToken, createdAt, expiresAt); err != nil {
|
|
return err
|
|
}
|
|
return tx.Commit()
|
|
}
|
|
|
|
// GetSession 按 token 摘要取会话(含已过期行;过期判定在调用方)。
|
|
func (s *Store) GetSession(tokenHash string) (*Session, error) {
|
|
row := s.db.QueryRow(
|
|
`SELECT token_hash, csrf_token, created_at, expires_at FROM sessions WHERE token_hash=?`, tokenHash)
|
|
var sess Session
|
|
err := row.Scan(&sess.TokenHash, &sess.CSRFToken, &sess.CreatedAt, &sess.ExpiresAt)
|
|
if errors.Is(err, sql.ErrNoRows) {
|
|
return nil, ErrNotFound
|
|
}
|
|
return &sess, err
|
|
}
|
|
|
|
// DeleteSession 删除会话行(登出/轮换)。
|
|
func (s *Store) DeleteSession(tokenHash string) error {
|
|
_, err := s.db.Exec(`DELETE FROM sessions WHERE token_hash=?`, tokenHash)
|
|
return err
|
|
}
|
|
|
|
// DeleteExpiredSessions 清理过期会话,返回删除行数。
|
|
func (s *Store) DeleteExpiredSessions(now int64) (int64, error) {
|
|
res, err := s.db.Exec(`DELETE FROM sessions WHERE expires_at < ?`, now)
|
|
if err != nil {
|
|
return 0, err
|
|
}
|
|
return res.RowsAffected()
|
|
}
|
|
|
|
// DeleteAllSessions 吊销全部会话(口令重置后强制所有端重新登录)。
|
|
func (s *Store) DeleteAllSessions() (int64, error) {
|
|
res, err := s.db.Exec(`DELETE FROM sessions`)
|
|
if err != nil {
|
|
return 0, err
|
|
}
|
|
return res.RowsAffected()
|
|
}
|