// 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 }