299 lines
9.7 KiB
Go
299 lines
9.7 KiB
Go
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)
|
|
}
|
|
}
|