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