Files
oci-portal/internal/service/totp_oauth_test.go
T
2026-07-22 16:51:23 +08:00

380 lines
12 KiB
Go

package service
import (
"context"
"errors"
"testing"
"time"
"github.com/glebarez/sqlite"
"github.com/pquerna/otp/totp"
"gorm.io/gorm"
"gorm.io/gorm/logger"
"oci-portal/internal/crypto"
"oci-portal/internal/model"
)
// newTotpEnv 建含 User/UserIdentity/Setting 表的环境,预置 admin 并注入 cipher。
func newTotpEnv(t *testing.T) (*AuthService, *gorm.DB) {
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)
}
sqlDB, err := db.DB()
if err != nil {
t.Fatalf("db handle: %v", err)
}
sqlDB.SetMaxOpenConns(1)
if err := db.AutoMigrate(&model.User{}, &model.UserIdentity{}, &model.Setting{}); err != nil {
t.Fatalf("auto migrate: %v", err)
}
auth := NewAuthService(db, "test-secret")
cipher, err := crypto.NewCipher("test-key")
if err != nil {
t.Fatalf("new cipher: %v", err)
}
auth.SetCipher(cipher)
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
t.Fatalf("ensure admin: %v", err)
}
return auth, db
}
// enableTotp 走完整 setup→activate 流程,返回明文密钥供测试生成验证码。
func enableTotp(t *testing.T, auth *AuthService) string {
t.Helper()
ctx := context.Background()
secret, uri, err := auth.SetupTotp(ctx, "admin")
if err != nil {
t.Fatalf("SetupTotp: %v", err)
}
if secret == "" || uri == "" {
t.Fatalf("setup 返回空 secret/uri")
}
code, err := totp.GenerateCode(secret, time.Now())
if err != nil {
t.Fatalf("generate code: %v", err)
}
if err := auth.ActivateTotp(ctx, "admin", code); err != nil {
t.Fatalf("ActivateTotp: %v", err)
}
return secret
}
func TestTotpLifecycle(t *testing.T) {
auth, db := newTotpEnv(t)
ctx := context.Background()
if on, _ := auth.TotpStatus(ctx, "admin"); on {
t.Fatal("初始不应启用")
}
secret := enableTotp(t, auth)
if on, _ := auth.TotpStatus(ctx, "admin"); !on {
t.Fatal("激活后应为启用")
}
// 密钥必须密文落库
var user model.User
if err := db.Where("username = ?", "admin").First(&user).Error; err != nil {
t.Fatalf("load user: %v", err)
}
if user.TotpSecretEnc == secret || user.TotpSecretEnc == "" {
t.Errorf("totp 密钥未加密落库")
}
// 重复 setup 拒绝
if _, _, err := auth.SetupTotp(ctx, "admin"); !errors.Is(err, ErrTotpAlreadyOn) {
t.Errorf("重复 setup err = %v, want ErrTotpAlreadyOn", err)
}
// 停用:无凭证拒绝,密码通过
if err := auth.DisableTotp(ctx, "admin", "", ""); !errors.Is(err, ErrTotpConfirm) {
t.Errorf("空凭证停用 err = %v, want ErrTotpConfirm", err)
}
if err := auth.DisableTotp(ctx, "admin", "pass123", ""); err != nil {
t.Fatalf("密码停用: %v", err)
}
if on, _ := auth.TotpStatus(ctx, "admin"); on {
t.Fatal("停用后应为关闭")
}
}
func TestActivateTotpRejects(t *testing.T) {
auth, _ := newTotpEnv(t)
ctx := context.Background()
// 未 setup 直接激活
if err := auth.ActivateTotp(ctx, "admin", "123456"); !errors.Is(err, ErrTotpNotSetup) {
t.Errorf("err = %v, want ErrTotpNotSetup", err)
}
// setup 后错误验证码
if _, _, err := auth.SetupTotp(ctx, "admin"); err != nil {
t.Fatalf("SetupTotp: %v", err)
}
if err := auth.ActivateTotp(ctx, "admin", "000000"); !errors.Is(err, ErrTotpInvalid) {
t.Errorf("err = %v, want ErrTotpInvalid", err)
}
}
func TestLoginWithTotp(t *testing.T) {
auth, _ := newTotpEnv(t)
ctx := context.Background()
secret := enableTotp(t, auth)
// 缺验证码:密码对也返回 ErrTotpRequired(不计失败)
if _, _, err := auth.Login(ctx, "admin", "pass123", "127.0.0.1", ""); !errors.Is(err, ErrTotpRequired) {
t.Fatalf("err = %v, want ErrTotpRequired", err)
}
// 错误验证码:按失败处理
if _, _, err := auth.Login(ctx, "admin", "pass123", "127.0.0.1", "000000"); !errors.Is(err, ErrInvalidCredentials) {
t.Fatalf("err = %v, want ErrInvalidCredentials", err)
}
// 正确验证码:登录成功
code, err := totp.GenerateCode(secret, time.Now())
if err != nil {
t.Fatalf("generate code: %v", err)
}
token, _, err := auth.Login(ctx, "admin", "pass123", "127.0.0.1", code)
if err != nil || token == "" {
t.Fatalf("带验证码登录失败: %v", err)
}
}
// fakeBindIdentity 直接把外部身份写入待测服务(绕过真实 OAuth flow)。
func fakeBindIdentity(t *testing.T, o *OAuthService, username, provider, subject, display string) {
t.Helper()
if err := o.bind(context.Background(), username, provider, externalIdentity{Subject: subject, Display: display}); err != nil {
t.Fatalf("bind: %v", err)
}
}
func newOAuthEnv(t *testing.T) (*OAuthService, *AuthService) {
t.Helper()
auth, db := newTotpEnv(t)
cipher, err := crypto.NewCipher("test-key")
if err != nil {
t.Fatalf("new cipher: %v", err)
}
settings := NewSettingService(db, cipher)
return NewOAuthService(db, settings, auth), auth
}
func TestOAuthBindLoginUnbind(t *testing.T) {
o, auth := newOAuthEnv(t)
ctx := context.Background()
fakeBindIdentity(t, o, "admin", "github", "10086", "octocat")
// 重复绑定同一身份拒绝
if err := o.bind(ctx, "admin", "github", externalIdentity{Subject: "10086", Display: "octocat"}); !errors.Is(err, ErrOAuthBound) {
t.Errorf("重复绑定 err = %v, want ErrOAuthBound", err)
}
// 已绑定身份可登录并拿到有效 JWT
token, loginUser, err := o.loginByIdentity(ctx, "github", externalIdentity{Subject: "10086", Display: "octocat"})
if err != nil || token == "" {
t.Fatalf("loginByIdentity: %v", err)
}
if loginUser != "admin" {
t.Errorf("loginUser = %q, want admin", loginUser)
}
if username, err := auth.ParseToken(context.Background(), token); err != nil || username != "admin" {
t.Errorf("token 应属 admin, got %q (%v)", username, err)
}
// 未绑定身份拒绝登录
if _, _, err := o.loginByIdentity(ctx, "github", externalIdentity{Subject: "999"}); !errors.Is(err, ErrOAuthNotBound) {
t.Errorf("未绑定登录 err = %v, want ErrOAuthNotBound", err)
}
// 列表与解绑
items, err := o.Identities(ctx, "admin")
if err != nil || len(items) != 1 {
t.Fatalf("identities = %d (%v), want 1", len(items), err)
}
if err := o.Unbind(ctx, "admin", items[0].ID); err != nil {
t.Fatalf("Unbind: %v", err)
}
if items, _ = o.Identities(ctx, "admin"); len(items) != 0 {
t.Errorf("解绑后仍有 %d 条", len(items))
}
}
func TestOAuthStateOneShot(t *testing.T) {
o, _ := newOAuthEnv(t)
o.mu.Lock()
o.pending["st1"] = oauthPending{provider: "github", mode: "login", expires: time.Now().Add(time.Minute)}
o.pending["st2"] = oauthPending{provider: "github", mode: "login", expires: time.Now().Add(-time.Minute)}
o.mu.Unlock()
if _, err := o.takeState("github", "st1"); err != nil {
t.Fatalf("首次消费: %v", err)
}
if _, err := o.takeState("github", "st1"); !errors.Is(err, ErrOAuthState) {
t.Errorf("二次消费 err = %v, want ErrOAuthState", err)
}
if _, err := o.takeState("github", "st2"); !errors.Is(err, ErrOAuthState) {
t.Errorf("过期 state err = %v, want ErrOAuthState", err)
}
if _, err := o.takeState("oidc", "no-such"); !errors.Is(err, ErrOAuthState) {
t.Errorf("未知 state err = %v, want ErrOAuthState", err)
}
}
func TestOAuthProvidersListsConfigured(t *testing.T) {
o, _ := newOAuthEnv(t)
ctx := context.Background()
if got := o.Providers(ctx); len(got) != 0 {
t.Fatalf("未配置时 providers = %v", got)
}
secret := "gh-secret"
err := o.settings.UpdateOAuth(ctx, UpdateOAuthInput{GithubClientID: strPtr("cid"), GithubClientSecret: &secret})
if err != nil {
t.Fatalf("UpdateOAuth: %v", err)
}
got := o.Providers(ctx)
if len(got) != 1 || got[0].Provider != "github" {
t.Errorf("providers = %v, want [github]", got)
}
// 视图不回明文且标记已设置
view, err := o.settings.OAuthView(ctx)
if err != nil || !view.GithubSecretSet || view.GithubClientID != "cid" {
t.Errorf("view = %+v (%v)", view, err)
}
}
// 未配置 provider 与未设面板地址是两类错误,文案分别指路。
func TestOAuthAuthorizeConfigErrors(t *testing.T) {
o, _ := newOAuthEnv(t)
ctx := context.Background()
// clientID 缺失
if _, err := o.AuthorizeURL(ctx, "github", "login", ""); !errors.Is(err, ErrOAuthNotConfigured) {
t.Errorf("无 clientID err = %v, want ErrOAuthNotConfigured", err)
}
// clientID 已配但面板地址未设置
if err := o.settings.UpdateOAuth(ctx, UpdateOAuthInput{GithubClientID: strPtr("Iv1.test")}); err != nil {
t.Fatalf("UpdateOAuth: %v", err)
}
if _, err := o.AuthorizeURL(ctx, "github", "login", ""); !errors.Is(err, ErrOAuthNoAppURL) {
t.Errorf("无面板地址 err = %v, want ErrOAuthNoAppURL", err)
}
// 面板地址就绪后正常返回授权 URL
o.settings.SetEnvPublicURL("https://demo.example.com")
url, err := o.AuthorizeURL(ctx, "github", "login", "")
if err != nil || url == "" {
t.Fatalf("AuthorizeURL: %q, %v", url, err)
}
}
func TestOAuthProvidersAndDisabled(t *testing.T) {
o, _ := newOAuthEnv(t)
ctx := context.Background()
o.settings.SetEnvPublicURL("https://demo.example.com")
// 未配置任何 provider → 空列表
if got := o.Providers(ctx); len(got) != 0 {
t.Fatalf("Providers = %v, want empty", got)
}
// 配置 github(无显示名称)→ 默认名 GitHub
in := UpdateOAuthInput{GithubClientID: strPtr("Iv1.test")}
if err := o.settings.UpdateOAuth(ctx, in); err != nil {
t.Fatalf("UpdateOAuth: %v", err)
}
got := o.Providers(ctx)
if len(got) != 1 || got[0].Provider != "github" || got[0].DisplayName != "GitHub" {
t.Fatalf("Providers = %+v, want [github/GitHub]", got)
}
// 自定义显示名称
in.GithubDisplayName = strPtr("公司账号")
if err := o.settings.UpdateOAuth(ctx, in); err != nil {
t.Fatalf("UpdateOAuth: %v", err)
}
if got = o.Providers(ctx); got[0].DisplayName != "公司账号" {
t.Fatalf("DisplayName = %q, want 公司账号", got[0].DisplayName)
}
// 禁用后:列表隐藏,login 模式 409,bind 模式仍可发起
in.GithubDisabled = boolPtr(true)
if err := o.settings.UpdateOAuth(ctx, in); err != nil {
t.Fatalf("UpdateOAuth: %v", err)
}
if got = o.Providers(ctx); len(got) != 0 {
t.Fatalf("禁用后 Providers = %v, want empty", got)
}
if _, err := o.AuthorizeURL(ctx, "github", "login", ""); !errors.Is(err, ErrOAuthDisabled) {
t.Errorf("禁用 login err = %v, want ErrOAuthDisabled", err)
}
if url, err := o.AuthorizeURL(ctx, "github", "bind", "admin"); err != nil || url == "" {
t.Errorf("禁用 bind = %q, %v, want 正常返回", url, err)
}
// view 回读禁用态与显示名称
view, err := o.settings.OAuthView(ctx)
if err != nil || !view.GithubDisabled || view.GithubDisplayName != "公司账号" {
t.Fatalf("OAuthView = %+v, %v", view, err)
}
}
// TestUpdateOAuthPartialPatch 锁定 provider 配置的部分更新语义:
// 只写出现字段,另一 provider 与未出现字段不回滚,secret nil 沿用空串清除。
func TestUpdateOAuthPartialPatch(t *testing.T) {
o, _ := newOAuthEnv(t)
ctx := context.Background()
secret := "gh-secret"
seed := UpdateOAuthInput{GithubClientID: strPtr("cid"), GithubClientSecret: &secret, GithubDisplayName: strPtr("公司账号")}
if err := o.settings.UpdateOAuth(ctx, seed); err != nil {
t.Fatalf("seed: %v", err)
}
steps := []struct {
name string
patch UpdateOAuthInput
check func(v OAuthProvidersView) string
}{
{
name: "只更新 oidc 不动 github",
patch: UpdateOAuthInput{OidcIssuer: strPtr("https://sso.example.com/"), OidcClientID: strPtr("oidc-cid")},
check: func(v OAuthProvidersView) string {
if v.OidcIssuer != "https://sso.example.com" || v.OidcClientID != "oidc-cid" {
return "oidc 字段未生效或未规范化"
}
if v.GithubClientID != "cid" || !v.GithubSecretSet || v.GithubDisplayName != "公司账号" {
return "github 字段被回滚"
}
return ""
},
},
{
name: "单开关启停,secret nil 沿用",
patch: UpdateOAuthInput{GithubDisabled: boolPtr(true)},
check: func(v OAuthProvidersView) string {
if !v.GithubDisabled || v.GithubClientID != "cid" || !v.GithubSecretSet {
return "启停外字段被动到或 secret 丢失"
}
return ""
},
},
{
name: "secret 空串清除",
patch: UpdateOAuthInput{GithubClientSecret: strPtr("")},
check: func(v OAuthProvidersView) string {
if v.GithubSecretSet {
return "空串未清除 secret"
}
if v.GithubClientID != "cid" {
return "clientID 被动到"
}
return ""
},
},
}
for _, st := range steps {
t.Run(st.name, func(t *testing.T) {
if err := o.settings.UpdateOAuth(ctx, st.patch); err != nil {
t.Fatalf("UpdateOAuth: %v", err)
}
view, err := o.settings.OAuthView(ctx)
if err != nil {
t.Fatalf("OAuthView: %v", err)
}
if msg := st.check(view); msg != "" {
t.Error(msg)
}
})
}
}