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: "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: "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: "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 = "公司账号" 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 = 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) } }