@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user