Files
opencatd-open/backend/internal/passkey/passkey.go
T
Sakurasan a376ac0722 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
2026-09-01 23:03:39 +08:00

368 lines
9.7 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// Package passkey 封装 WebAuthn(passkey)注册与登录。
package passkey
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"net/http/httptest"
"strconv"
"sync"
"time"
"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 时使用内存存储
}
// 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
}
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) {
wa, err := webauthn.New(&webauthn.Config{
RPDisplayName: cfg.Name,
RPID: cfg.RPID,
RPOrigins: cfg.Origins,
})
if err != nil {
return nil, err
}
// 根据配置选择存储后端
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 接口。
type webUser struct {
id uint64
name string
displayName string
credentials []webauthn.Credential
}
func (u *webUser) WebAuthnID() []byte { return []byte(strconv.FormatUint(u.id, 10)) }
func (u *webUser) WebAuthnName() string { return u.name }
func (u *webUser) WebAuthnDisplayName() string { return u.displayName }
func (u *webUser) WebAuthnIcon() string { return "" }
func (u *webUser) WebAuthnCredentials() []webauthn.Credential { return u.credentials }
func (s *Service) loadWebUser(u *store.User) (*webUser, error) {
var pks []store.Passkey
s.db.Where("user_id = ?", u.ID).Find(&pks)
creds := make([]webauthn.Credential, 0, len(pks))
for _, pk := range pks {
var c webauthn.Credential
if err := json.Unmarshal(pk.Credential, &c); err == nil {
creds = append(creds, c)
}
}
return &webUser{id: u.ID, name: u.Username, displayName: u.Username, credentials: creds}, nil
}
// GetUserByUsername 通过用户名或邮箱查找用户
func (s *Service) GetUserByUsername(username string) (*store.User, error) {
var u store.User
if err := s.db.Where("username = ? OR email = ?", username, username).First(&u).Error; err != nil {
return nil, err
}
return &u, nil
}
// GetUserByID 通过 ID 查找用户
func (s *Service) GetUserByID(id uint64) (*store.User, error) {
var u store.User
if err := s.db.First(&u, id).Error; err != nil {
return nil, err
}
return &u, nil
}
// ---------------------------------------------------------------------------
// 注册
// BeginRegistration 生成注册选项并暂存 challenge。
func (s *Service) BeginRegistration(u *store.User) (*protocol.CredentialCreation, error) {
wu, err := s.loadWebUser(u)
if err != nil {
return nil, err
}
creation, session, err := s.wa.BeginRegistration(wu)
if err != nil {
return nil, err
}
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, 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)
if err != nil {
return err
}
raw, _ := json.Marshal(cred)
nm := name
if nm == "" {
nm = "passkey"
}
return s.db.Create(&store.Passkey{
UserID: u.ID, Name: nm, CredentialID: cred.ID, Credential: raw,
}).Error
}
// ---------------------------------------------------------------------------
// 登录
// BeginLogin 已知用户(按用户名)发起断言。
func (s *Service) BeginLogin(u *store.User) (*protocol.CredentialAssertion, error) {
wu, err := s.loadWebUser(u)
if err != nil {
return nil, err
}
assertion, session, err := s.wa.BeginLogin(wu)
if err != nil {
return nil, err
}
if err := s.sessions.Set(context.Background(), session); err != nil {
return nil, err
}
return assertion, nil
}
// BeginDiscoverableLogin 无用户名(使用平台/漫游器上的可发现凭据)。
func (s *Service) BeginDiscoverableLogin() (*protocol.CredentialAssertion, error) {
assertion, session, err := s.wa.BeginDiscoverableLogin()
if err != nil {
return nil, err
}
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, 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)
if err != nil {
return err
}
return s.updateCredential(u.ID, cred)
}
// FinishDiscoverableLogin 通过凭据定位用户并校验断言。
func (s *Service) FinishDiscoverableLogin(challenge string, body []byte) (*store.User, error) {
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 反查用户
var pks []store.Passkey
if err := s.db.Find(&pks).Error; err != nil {
return nil, err
}
// 遍历所有 passkey 找到匹配的
for _, pk := range pks {
var c webauthn.Credential
if err := json.Unmarshal(pk.Credential, &c); err != nil {
continue
}
// 尝试用这个用户的凭据进行登录
var u store.User
if err := s.db.First(&u, pk.UserID).Error; err != nil {
continue
}
wu, err := s.loadWebUser(&u)
if err != nil {
continue
}
cred, err := s.wa.FinishLogin(wu, *session, req)
if err != nil {
continue
}
_ = s.updateCredential(u.ID, cred)
return &u, nil
}
return nil, errors.New("no matching passkey found")
}
// ---------------------------------------------------------------------------
// 管理
// List 列出用户的 passkey。
func (s *Service) List(userID uint64) ([]store.Passkey, error) {
var pks []store.Passkey
err := s.db.Where("user_id = ?", userID).Order("id DESC").Find(&pks).Error
return pks, err
}
// Delete 删除用户的 passkey。
func (s *Service) Delete(userID, id uint64) error {
res := s.db.Where("id = ? AND user_id = ?", id, userID).Delete(&store.Passkey{})
if res.Error != nil {
return res.Error
}
if res.RowsAffected == 0 {
return gorm.ErrRecordNotFound
}
return nil
}
func (s *Service) updateCredential(userID uint64, cred *webauthn.Credential) error {
raw, _ := json.Marshal(cred)
return s.db.Model(&store.Passkey{}).
Where("user_id = ? AND credential_id = ?", userID, cred.ID).
Update("credential", raw).Error
}