新增活跃会话管理与通行密钥、钱包登录
CI / test (push) Successful in 55s

This commit is contained in:
2026-07-30 12:23:05 +08:00
parent f1914880ac
commit 109c345e5e
49 changed files with 6468 additions and 336 deletions
+298
View File
@@ -0,0 +1,298 @@
package service
import (
"context"
"encoding/base64"
"encoding/json"
"errors"
"strings"
"testing"
"time"
"github.com/glebarez/sqlite"
"github.com/go-webauthn/webauthn/webauthn"
"gorm.io/gorm"
"gorm.io/gorm/logger"
"oci-portal/internal/crypto"
"oci-portal/internal/model"
)
// newTestPasskey 组装内存库上的 PasskeyService(admin 账号已建,面板地址经环境回退注入)。
func newTestPasskey(t *testing.T, appURL string) *PasskeyService {
t.Helper()
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
})
if err != nil {
t.Fatalf("open in-memory sqlite: %v", err)
}
if err := db.AutoMigrate(&model.User{}, &model.UserPasskey{}, &model.UserIdentity{}, &model.UserSession{}, &model.Setting{}); err != nil {
t.Fatalf("auto migrate: %v", err)
}
auth := NewAuthService(db, "test-jwt-secret")
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
t.Fatalf("ensure admin: %v", err)
}
cipher, err := crypto.NewCipher("test-key")
if err != nil {
t.Fatalf("new cipher: %v", err)
}
settings := NewSettingService(db, cipher)
settings.SetEnvPublicURL(appURL)
return NewPasskeyService(db, settings, auth)
}
func TestPasskeyRPDerivation(t *testing.T) {
tests := []struct {
name string
appURL string
wantErr error
wantRPID string
wantOrigin string
}{
{name: "无面板地址", appURL: "", wantErr: ErrPasskeyNoAppURL},
{name: "https 域名", appURL: "https://demo.example.com", wantRPID: "demo.example.com", wantOrigin: "https://demo.example.com"},
{name: "带端口", appURL: "https://panel.example.com:8443", wantRPID: "panel.example.com", wantOrigin: "https://panel.example.com:8443"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
p := newTestPasskey(t, tt.appURL)
w, err := p.rp()
if tt.wantErr != nil {
if !errors.Is(err, tt.wantErr) {
t.Fatalf("rp() err = %v, want %v", err, tt.wantErr)
}
return
}
if err != nil {
t.Fatalf("rp(): %v", err)
}
if w.Config.RPID != tt.wantRPID {
t.Errorf("RPID = %q, want %q", w.Config.RPID, tt.wantRPID)
}
if len(w.Config.RPOrigins) != 1 || w.Config.RPOrigins[0] != tt.wantOrigin {
t.Errorf("RPOrigins = %v, want [%s]", w.Config.RPOrigins, tt.wantOrigin)
}
})
}
}
func TestPasskeyPendingLifecycle(t *testing.T) {
p := newTestPasskey(t, "https://demo.example.com")
session := webauthn.SessionData{Challenge: "challenge-1"}
id, err := p.putPending(session, "admin")
if err != nil {
t.Fatalf("putPending: %v", err)
}
got, err := p.takePending(id)
if err != nil {
t.Fatalf("takePending: %v", err)
}
if got.session.Challenge != "challenge-1" || got.username != "admin" {
t.Errorf("pending = %+v, want challenge-1/admin", got)
}
// 一次性:再次消费同一 sessionId 必须失效
if _, err := p.takePending(id); !errors.Is(err, ErrPasskeySession) {
t.Errorf("second take err = %v, want ErrPasskeySession", err)
}
// 过期条目视为无效
expiredID, _ := p.putPending(session, "")
p.mu.Lock()
e := p.pending[expiredID]
e.expires = time.Now().Add(-time.Second)
p.pending[expiredID] = e
p.mu.Unlock()
if _, err := p.takePending(expiredID); !errors.Is(err, ErrPasskeySession) {
t.Errorf("expired take err = %v, want ErrPasskeySession", err)
}
}
func TestPasskeyBeginRegister(t *testing.T) {
tests := []struct {
name string
appURL string
seedKeys int
wantErr error
}{
{name: "正常发起", appURL: "https://demo.example.com"},
{name: "无面板地址", appURL: "", wantErr: ErrPasskeyNoAppURL},
{name: "数量达上限", appURL: "https://demo.example.com", seedKeys: passkeyMaxPerUser, wantErr: ErrPasskeyLimit},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
p := newTestPasskey(t, tt.appURL)
for i := 0; i < tt.seedKeys; i++ {
seedPasskeyRow(t, p.db, uint(i+1))
}
sid, opts, err := p.BeginRegister(context.Background(), "admin")
if tt.wantErr != nil {
if !errors.Is(err, tt.wantErr) {
t.Fatalf("BeginRegister err = %v, want %v", err, tt.wantErr)
}
return
}
if err != nil {
t.Fatalf("BeginRegister: %v", err)
}
if sid == "" || opts == nil || opts.Response.Challenge.String() == "" {
t.Errorf("BeginRegister returned empty session/options")
}
})
}
}
// seedPasskeyRow 直插一行合法凭据(JSON 与 CredentialID 对应)。
func seedPasskeyRow(t *testing.T, db *gorm.DB, seq uint) {
t.Helper()
credID := []byte{byte(seq), 2, 3, 4}
cred := webauthn.Credential{ID: credID, PublicKey: []byte{5, 6}}
row := model.UserPasskey{
UserID: 1,
Name: "key",
CredentialID: base64.RawURLEncoding.EncodeToString(credID),
CredentialIDHash: passkeyCredHash(credID),
Credential: mustCredJSON(t, cred),
Origin: "https://demo.example.com",
}
if err := db.Create(&row).Error; err != nil {
t.Fatalf("seed passkey: %v", err)
}
}
func mustCredJSON(t *testing.T, cred webauthn.Credential) string {
t.Helper()
raw, err := json.Marshal(cred)
if err != nil {
t.Fatalf("marshal credential: %v", err)
}
return string(raw)
}
func TestPasskeyFinishRegisterSessionChecks(t *testing.T) {
tests := []struct {
name string
sessionUser string
finishUser string
body string
wantErr error
}{
{name: "会话不存在", sessionUser: "-", finishUser: "admin", wantErr: ErrPasskeySession},
{name: "会话归属不符", sessionUser: "other", finishUser: "admin", wantErr: ErrPasskeySession},
{name: "登录会话不可注册", sessionUser: "", finishUser: "admin", wantErr: ErrPasskeySession},
{name: "凭据体不合法", sessionUser: "admin", finishUser: "admin", body: "{}", wantErr: ErrPasskeyVerify},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
p := newTestPasskey(t, "https://demo.example.com")
sid := "missing"
if tt.sessionUser != "-" {
var err error
sid, err = p.putPending(webauthn.SessionData{Challenge: "c"}, tt.sessionUser)
if err != nil {
t.Fatalf("putPending: %v", err)
}
}
err := p.FinishRegister(context.Background(), tt.finishUser, sid, "名称", strings.NewReader(tt.body), proofOf(t, p.db, "admin"))
if !errors.Is(err, tt.wantErr) {
t.Fatalf("FinishRegister err = %v, want %v", err, tt.wantErr)
}
})
}
}
func TestPasskeyCredentialRoundTrip(t *testing.T) {
p := newTestPasskey(t, "https://demo.example.com")
cred := &webauthn.Credential{ID: []byte{9, 9, 9}, PublicKey: []byte{1, 2, 3}}
cred.Authenticator.SignCount = 7
if err := saveCredentialTx(p.db, 1, "", "https://demo.example.com", cred); err != nil {
t.Fatalf("saveCredentialTx: %v", err)
}
u, err := p.loadUser(context.Background(), "admin")
if err != nil {
t.Fatalf("loadUser: %v", err)
}
creds := u.WebAuthnCredentials()
if len(creds) != 1 || creds[0].Authenticator.SignCount != 7 {
t.Fatalf("credentials = %+v, want 1 item signCount 7", creds)
}
if u.keys[0].Name != "通行密钥" {
t.Errorf("default name = %q, want 通行密钥", u.keys[0].Name)
}
// userHandle 往返:8 字节大端 ID 找回同一账号
pu, err := p.userByHandle(context.Background(), passkeyUserHandle(u.user.ID))
if err != nil || pu.user.Username != "admin" {
t.Errorf("userByHandle = (%+v, %v), want admin", pu.user, err)
}
if _, err := p.userByHandle(context.Background(), []byte{1, 2}); !errors.Is(err, ErrPasskeyVerify) {
t.Errorf("short handle err = %v, want ErrPasskeyVerify", err)
}
}
func TestPasskeyRemoveAndVersionBump(t *testing.T) {
p := newTestPasskey(t, "https://demo.example.com")
seedPasskeyRow(t, p.db, 1)
var before model.User
if err := p.db.First(&before, 1).Error; err != nil {
t.Fatalf("load user: %v", err)
}
if !p.HasAny(context.Background()) {
t.Fatal("HasAny = false, want true after seed")
}
if err := p.Remove(context.Background(), "admin", 1, proofOf(t, p.db, "admin")); err != nil {
t.Fatalf("Remove: %v", err)
}
if err := p.Remove(context.Background(), "admin", 1, proofOf(t, p.db, "admin")); !errors.Is(err, gorm.ErrRecordNotFound) {
t.Errorf("second Remove err = %v, want ErrRecordNotFound", err)
}
if p.HasAny(context.Background()) {
t.Error("HasAny = true, want false after remove")
}
var after model.User
if err := p.db.First(&after, 1).Error; err != nil {
t.Fatalf("reload user: %v", err)
}
if after.TokenVersion != before.TokenVersion+1 {
t.Errorf("TokenVersion = %d, want %d", after.TokenVersion, before.TokenVersion+1)
}
}
func TestPasskeyHasAnyRequiresCurrentOrigin(t *testing.T) {
p := newTestPasskey(t, "https://demo.example.com")
seedPasskeyRow(t, p.db, 1)
tests := []struct {
name string
appURL string
want bool
}{
{name: "registered origin", appURL: "https://demo.example.com", want: true},
{name: "missing app url", appURL: ""},
{name: "different origin", appURL: "https://other.example.com"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
p.settings.SetEnvPublicURL(tt.appURL)
if got := p.HasAny(context.Background()); got != tt.want {
t.Fatalf("HasAny = %v, want %v", got, tt.want)
}
})
}
}
func TestPasskeyFinishLoginGuard(t *testing.T) {
p := newTestPasskey(t, "https://demo.example.com")
ctx := context.Background()
// 无效 sessionId 反复失败:达到默认阈值后转锁定
var lastErr error
for i := 0; i < securityDefaults.LoginFailLimit+1; i++ {
_, _, _, lastErr = p.FinishLogin(ctx, "bad-session", SessionMeta{ClientIP: "10.0.0.9"}, strings.NewReader("{}"))
}
if !errors.Is(lastErr, ErrLoginLocked) {
t.Fatalf("after %d failures err = %v, want ErrLoginLocked", securityDefaults.LoginFailLimit+1, lastErr)
}
// 其他 IP 不受连坐
if _, _, _, err := p.FinishLogin(ctx, "bad-session", SessionMeta{ClientIP: "10.0.0.10"}, strings.NewReader("{}")); errors.Is(err, ErrLoginLocked) {
t.Errorf("different IP got locked prematurely: %v", err)
}
}