feat: add Redis support for distributed passkey sessions

- Implement SessionStore interface for pluggable storage backends
- Add redisStore for distributed environment
- Add memoryStore as fallback for single instance
- Update router to initialize Redis client when configured
- Update .env.example with Redis documentation
This commit is contained in:
Sakurasan
2026-09-01 23:03:39 +08:00
parent 9733b3c20b
commit a376ac0722
5 changed files with 157 additions and 46 deletions
+137 -45
View File
@@ -3,8 +3,10 @@ package passkey
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"net/http/httptest"
"strconv"
@@ -13,23 +15,112 @@ import (
"github.com/go-webauthn/webauthn/protocol"
"github.com/go-webauthn/webauthn/webauthn"
"github.com/redis/go-redis/v9"
"opencatd-open/internal/store"
"gorm.io/gorm"
)
const (
sessionPrefix = "passkey:session:"
sessionTTL = 5 * time.Minute
)
type Config struct {
RPID string
Origins []string
Name string
Redis *redis.Client // 可选,nil 时使用内存存储
}
// Service WebAuthn 服务:凭据存储 + challenge 会话(内存)。
type Service struct {
wa *webauthn.WebAuthn
db *gorm.DB
// SessionStore challenge 会话存储接口
type SessionStore interface {
Set(ctx context.Context, session *webauthn.SessionData) error
Get(ctx context.Context, challenge string) (*webauthn.SessionData, bool, error)
Delete(ctx context.Context, challenge string) error
}
// memoryStore 内存存储(单实例)
type memoryStore struct {
mu sync.Mutex
sessions map[string]webauthn.SessionData // keyed by challenge
sessions map[string]webauthn.SessionData
}
func newMemoryStore() *memoryStore {
return &memoryStore{sessions: make(map[string]webauthn.SessionData)}
}
func (m *memoryStore) Set(_ context.Context, session *webauthn.SessionData) error {
m.mu.Lock()
m.sessions[session.Challenge] = *session
m.mu.Unlock()
return nil
}
func (m *memoryStore) Get(_ context.Context, challenge string) (*webauthn.SessionData, bool, error) {
m.mu.Lock()
sess, ok := m.sessions[challenge]
m.mu.Unlock()
if !ok {
return nil, false, nil
}
// 检查过期
if !sess.Expires.IsZero() && time.Now().After(sess.Expires) {
return nil, false, nil
}
return &sess, true, nil
}
func (m *memoryStore) Delete(_ context.Context, challenge string) error {
m.mu.Lock()
delete(m.sessions, challenge)
m.mu.Unlock()
return nil
}
// redisStore Redis 存储(分布式)
type redisStore struct {
rdb *redis.Client
}
func newRedisStore(rdb *redis.Client) *redisStore {
return &redisStore{rdb: rdb}
}
func (r *redisStore) Set(ctx context.Context, session *webauthn.SessionData) error {
data, err := json.Marshal(session)
if err != nil {
return fmt.Errorf("marshal session: %w", err)
}
key := sessionPrefix + session.Challenge
return r.rdb.Set(ctx, key, data, sessionTTL).Err()
}
func (r *redisStore) Get(ctx context.Context, challenge string) (*webauthn.SessionData, bool, error) {
key := sessionPrefix + challenge
data, err := r.rdb.Get(ctx, key).Bytes()
if err == redis.Nil {
return nil, false, nil
}
if err != nil {
return nil, false, fmt.Errorf("redis get: %w", err)
}
var sess webauthn.SessionData
if err := json.Unmarshal(data, &sess); err != nil {
return nil, false, fmt.Errorf("unmarshal session: %w", err)
}
return &sess, true, nil
}
func (r *redisStore) Delete(ctx context.Context, challenge string) error {
key := sessionPrefix + challenge
return r.rdb.Del(ctx, key).Err()
}
// Service WebAuthn 服务:凭据存储 + challenge 会话。
type Service struct {
wa *webauthn.WebAuthn
db *gorm.DB
sessions SessionStore
}
func New(db *gorm.DB, cfg Config) (*Service, error) {
@@ -41,7 +132,16 @@ func New(db *gorm.DB, cfg Config) (*Service, error) {
if err != nil {
return nil, err
}
return &Service{wa: wa, db: db, sessions: map[string]webauthn.SessionData{}}, nil
// 根据配置选择存储后端
var store SessionStore
if cfg.Redis != nil {
store = newRedisStore(cfg.Redis)
} else {
store = newMemoryStore()
}
return &Service{wa: wa, db: db, sessions: store}, nil
}
// webUser 实现 go-webauthn 的 User 接口。
@@ -102,22 +202,30 @@ func (s *Service) BeginRegistration(u *store.User) (*protocol.CredentialCreation
if err != nil {
return nil, err
}
s.storeSession(session)
if err := s.sessions.Set(context.Background(), session); err != nil {
return nil, err
}
return creation, nil
}
// FinishRegistration 校验浏览器返回的凭据并落库。
func (s *Service) FinishRegistration(u *store.User, challenge string, body []byte, name string) error {
session, ok := s.takeSession(challenge)
session, ok, err := s.sessions.Get(context.Background(), challenge)
if err != nil {
return err
}
if !ok {
return errors.New("challenge 已过期或不存在")
}
// 删除已使用的 challenge
_ = s.sessions.Delete(context.Background(), challenge)
wu, err := s.loadWebUser(u)
if err != nil {
return err
}
req := httptest.NewRequest(http.MethodPost, "/", bytes.NewReader(body))
cred, err := s.wa.FinishRegistration(wu, session, req)
cred, err := s.wa.FinishRegistration(wu, *session, req)
if err != nil {
return err
}
@@ -144,7 +252,9 @@ func (s *Service) BeginLogin(u *store.User) (*protocol.CredentialAssertion, erro
if err != nil {
return nil, err
}
s.storeSession(session)
if err := s.sessions.Set(context.Background(), session); err != nil {
return nil, err
}
return assertion, nil
}
@@ -154,22 +264,29 @@ func (s *Service) BeginDiscoverableLogin() (*protocol.CredentialAssertion, error
if err != nil {
return nil, err
}
s.storeSession(session)
if err := s.sessions.Set(context.Background(), session); err != nil {
return nil, err
}
return assertion, nil
}
// FinishLogin 校验断言并更新签名计数。
func (s *Service) FinishLogin(u *store.User, challenge string, body []byte) error {
session, ok := s.takeSession(challenge)
session, ok, err := s.sessions.Get(context.Background(), challenge)
if err != nil {
return err
}
if !ok {
return errors.New("challenge 已过期或不存在")
}
_ = s.sessions.Delete(context.Background(), challenge)
wu, err := s.loadWebUser(u)
if err != nil {
return err
}
req := httptest.NewRequest(http.MethodPost, "/", bytes.NewReader(body))
cred, err := s.wa.FinishLogin(wu, session, req)
cred, err := s.wa.FinishLogin(wu, *session, req)
if err != nil {
return err
}
@@ -178,10 +295,15 @@ func (s *Service) FinishLogin(u *store.User, challenge string, body []byte) erro
// FinishDiscoverableLogin 通过凭据定位用户并校验断言。
func (s *Service) FinishDiscoverableLogin(challenge string, body []byte) (*store.User, error) {
session, ok := s.takeSession(challenge)
session, ok, err := s.sessions.Get(context.Background(), challenge)
if err != nil {
return nil, err
}
if !ok {
return nil, errors.New("challenge 已过期或不存在")
}
_ = s.sessions.Delete(context.Background(), challenge)
req := httptest.NewRequest(http.MethodPost, "/", bytes.NewReader(body))
// 先通过 credential_id 反查用户
@@ -205,13 +327,8 @@ func (s *Service) FinishDiscoverableLogin(challenge string, body []byte) (*store
if err != nil {
continue
}
cred, err := s.wa.FinishLogin(wu, session, req)
cred, err := s.wa.FinishLogin(wu, *session, req)
if err != nil {
// 不是这个用户的 passkey,继续尝试
session, _ = s.takeSession(challenge)
if !ok {
return nil, errors.New("challenge 已过期或不存在")
}
continue
}
_ = s.updateCredential(u.ID, cred)
@@ -248,28 +365,3 @@ func (s *Service) updateCredential(userID uint64, cred *webauthn.Credential) err
Where("user_id = ? AND credential_id = ?", userID, cred.ID).
Update("credential", raw).Error
}
// ---------------------------------------------------------------------------
// challenge 会话
func (s *Service) storeSession(session *webauthn.SessionData) {
s.mu.Lock()
s.sessions[session.Challenge] = *session
s.mu.Unlock()
}
func (s *Service) takeSession(challenge string) (webauthn.SessionData, bool) {
s.mu.Lock()
sess, ok := s.sessions[challenge]
if ok {
delete(s.sessions, challenge)
}
s.mu.Unlock()
// Expires 可能为零值:go-webauthn 默认 Enforce=false 不设过期时间。
// 零值时间恒早于 now,直接 After 会把每个 challenge 都判为过期,
// 与库内部一致,仅当显式设置了过期时间才做校验。
if ok && !sess.Expires.IsZero() && time.Now().After(sess.Expires) {
return webauthn.SessionData{}, false
}
return sess, ok
}