- 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
368 lines
9.7 KiB
Go
368 lines
9.7 KiB
Go
// 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
|
||
}
|