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.UserPasskey{}, &model.UserSession{}, &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 走完整 login→setup→activate 流程,返回明文密钥与接续令牌。 func enableTotp(t *testing.T, auth *AuthService) (string, string) { t.Helper() ctx := context.Background() oldToken := loginSession(t, auth, "127.0.0.1", "totp-test") 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) } token, _, err := auth.ActivateTotp( ctx, "admin", code, oldToken, SessionMeta{ClientIP: "127.0.0.1"}, mustProof(t, auth, oldToken)) if err != nil { t.Fatalf("ActivateTotp: %v", err) } return secret, token } func TestTotpLifecycle(t *testing.T) { auth, db := newTotpEnv(t) ctx := context.Background() if on, _ := auth.TotpStatus(ctx, "admin"); on { t.Fatal("初始不应启用") } secret, token := 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", "", "", token, SessionMeta{}, mustProof(t, auth, token)); !errors.Is(err, ErrTotpConfirm) { t.Errorf("空凭证停用 err = %v, want ErrTotpConfirm", err) } if _, _, err := auth.DisableTotp( ctx, "admin", "pass123", "", token, SessionMeta{}, mustProof(t, auth, token)); 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", "", SessionMeta{}, TokenProof{}); !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", "", SessionMeta{}, TokenProof{}); !errors.Is(err, ErrTotpInvalid) { t.Errorf("err = %v, want ErrTotpInvalid", err) } } func TestActivateTotpRejectsRevokedProof(t *testing.T) { auth, _ := newTotpEnv(t) ctx := context.Background() staleToken := loginSession(t, auth, "127.0.0.1", "stale") staleProof := mustProof(t, auth, staleToken) secret, _, err := auth.SetupTotp(ctx, "admin") if err != nil { t.Fatalf("SetupTotp: %v", err) } code, err := totp.GenerateCode(secret, time.Now()) if err != nil { t.Fatalf("GenerateCode: %v", err) } current := loginSession(t, auth, "127.0.0.2", "current") if _, _, err := auth.RevokeSessions( ctx, "admin", current, SessionMeta{ClientIP: "127.0.0.2"}, mustProof(t, auth, current)); err != nil { t.Fatalf("RevokeSessions: %v", err) } if _, _, err := auth.ActivateTotp( ctx, "admin", code, staleToken, SessionMeta{ClientIP: "127.0.0.1"}, staleProof, ); !errors.Is(err, ErrTokenStale) { t.Fatalf("ActivateTotp err = %v, want ErrTokenStale", err) } if on, err := auth.TotpStatus(ctx, "admin"); err != nil || on { t.Fatalf("TotpStatus = %v, %v; want disabled", on, err) } } func TestLoginWithTotp(t *testing.T) { auth, _ := newTotpEnv(t) ctx := context.Background() secret, _ := enableTotp(t, auth) // 缺验证码:密码对也返回 ErrTotpRequired(不计失败) if _, _, err := auth.Login(ctx, "admin", "pass123", "", SessionMeta{ClientIP: "127.0.0.1"}); !errors.Is(err, ErrTotpRequired) { t.Fatalf("err = %v, want ErrTotpRequired", err) } // 错误验证码:按失败处理 if _, _, err := auth.Login(ctx, "admin", "pass123", "000000", SessionMeta{ClientIP: "127.0.0.1"}); !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", code, SessionMeta{ClientIP: "127.0.0.1"}) 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() p := fakePending(t, o, username, provider) if _, err := o.bind(context.Background(), p, externalIdentity{Subject: subject, Display: display}, SessionMeta{}); err != nil { t.Fatalf("bind: %v", err) } } // fakePending 构造与当前令牌版本一致的 bind 流程上下文(跳过外部授权码交换)。 func fakePending(t *testing.T, o *OAuthService, username, provider string) oauthPending { t.Helper() var user model.User if err := o.db.Where("username = ?", username).First(&user).Error; err != nil { t.Fatalf("find user: %v", err) } return oauthPending{provider: provider, mode: "bind", username: username, proof: TokenProof{Ver: user.TokenVersion}} } // TestOAuthBindStaleToken 验证 bind 回调复验发起时令牌:撤销全部(版本递增)后, // 已登记的绑定流程随之作废,不再触发外部换码。 func TestOAuthBindStaleToken(t *testing.T) { o, auth := newOAuthEnv(t) ctx := context.Background() o.settings.SetEnvPublicURL("https://demo.example.com") secret := "secret" if err := o.settings.UpdateOAuth(ctx, UpdateOAuthInput{ GithubClientID: strPtr("cid"), GithubClientSecret: &secret, }); err != nil { t.Fatalf("UpdateOAuth: %v", err) } token, _, err := auth.Login(ctx, "admin", "pass123", "", SessionMeta{ClientIP: "10.0.0.1"}) if err != nil { t.Fatalf("login: %v", err) } if _, err := o.AuthorizeURL(ctx, "github", "bind", "admin", token); err != nil { t.Fatalf("AuthorizeURL(bind): %v", err) } var state string o.mu.Lock() for k := range o.pending { state = k } o.mu.Unlock() if err := auth.bumpTokenVersion(ctx, "admin"); err != nil { t.Fatalf("bump token version: %v", err) } if _, _, _, err := o.HandleCallback(ctx, "github", state, "code", SessionMeta{ClientIP: "10.0.0.1"}); !errors.Is(err, ErrOAuthState) { t.Fatalf("stale-token callback err = %v, want ErrOAuthState", err) } } // TestUpdateOAuthKeepsLogin 验证 provider 配置防自锁:密码登录禁用期间, // 禁用或清空最后可用的登录方式被拒;有通行密钥兜底后放行。 func TestUpdateOAuthKeepsLogin(t *testing.T) { o, auth := newOAuthEnv(t) ctx := context.Background() o.settings.SetEnvPublicURL("https://demo.example.com") secret := "gh-secret" if err := o.settings.UpdateOAuth(ctx, UpdateOAuthInput{GithubClientID: strPtr("cid"), GithubClientSecret: &secret}); err != nil { t.Fatalf("UpdateOAuth: %v", err) } fakeBindIdentity(t, o, "admin", "github", "77", "octo") if err := saveSettingTx(auth.db, settingSecPasswordLoginOff, "1"); err != nil { t.Fatalf("disable password login: %v", err) } off := true if err := o.settings.UpdateOAuth(ctx, UpdateOAuthInput{GithubDisabled: &off}); !errors.Is(err, ErrProviderLastLogin) { t.Fatalf("disable last provider err = %v, want ErrProviderLastLogin", err) } if err := o.settings.UpdateOAuth(ctx, UpdateOAuthInput{GithubClientID: strPtr("")}); !errors.Is(err, ErrProviderLastLogin) { t.Fatalf("clear last clientID err = %v, want ErrProviderLastLogin", err) } pk := model.UserPasskey{ UserID: 1, Name: "k", CredentialID: "c", CredentialIDHash: "h", Credential: "{}", Origin: "https://demo.example.com", } if err := auth.db.Create(&pk).Error; err != nil { t.Fatalf("seed passkey: %v", err) } if err := o.settings.UpdateOAuth(ctx, UpdateOAuthInput{GithubDisabled: &off}); err != nil { t.Fatalf("disable provider with passkey fallback: %v", err) } } // TestUpdateSecurityAppURLKeepsLogin 验证面板地址防自锁:密码禁用期间 // 清空地址一律拒绝;域名变更须有钱包身份兜底(通行密钥随 RP ID 失效)。 func TestUpdateSecurityAppURLKeepsLogin(t *testing.T) { o, auth := newOAuthEnv(t) ctx := context.Background() if err := o.settings.UpdateSecurity(ctx, SecurityPatch{AppURL: strPtr("https://a.example.com")}); err != nil { t.Fatalf("seed app url: %v", err) } pk := model.UserPasskey{ UserID: 1, Name: "k", CredentialID: "c", CredentialIDHash: "h", Credential: "{}", Origin: "https://a.example.com", } if err := auth.db.Create(&pk).Error; err != nil { t.Fatalf("seed passkey: %v", err) } if err := saveSettingTx(auth.db, settingSecPasswordLoginOff, "1"); err != nil { t.Fatalf("disable password login: %v", err) } if err := o.settings.UpdateSecurity(ctx, SecurityPatch{AppURL: strPtr("")}); !errors.Is(err, ErrProviderLastLogin) { t.Fatalf("clear app url err = %v, want ErrProviderLastLogin", err) } if err := o.settings.UpdateSecurity(ctx, SecurityPatch{AppURL: strPtr("https://b.example.com")}); !errors.Is(err, ErrProviderLastLogin) { t.Fatalf("change host err = %v, want ErrProviderLastLogin", err) } ident := model.UserIdentity{UserID: 1, Provider: "wallet", Subject: "0xaa", Display: "w"} if err := auth.db.Create(&ident).Error; err != nil { t.Fatalf("seed wallet identity: %v", err) } if err := o.settings.UpdateSecurity(ctx, SecurityPatch{AppURL: strPtr("https://b.example.com")}); err != nil { t.Fatalf("change host with wallet fallback: %v", err) } } // TestPasswordDisableStalePasskeyOrigin 验证逆序自锁防线:地址 A 注册的 // 通行密钥在改到地址 B 后不再计入「可用免密方式」,禁用密码被拒。 func TestPasswordDisableStalePasskeyOrigin(t *testing.T) { o, auth := newOAuthEnv(t) ctx := context.Background() auth.SetNotifier(nil, o.settings) if err := o.settings.UpdateSecurity(ctx, SecurityPatch{AppURL: strPtr("https://a.example.com")}); err != nil { t.Fatalf("seed app url: %v", err) } pk := model.UserPasskey{UserID: 1, Name: "k", CredentialID: "c", CredentialIDHash: "h", Origin: "https://a.example.com", Credential: "{}"} if err := auth.db.Create(&pk).Error; err != nil { t.Fatalf("seed passkey: %v", err) } // 密码未禁用:改地址不受限 if err := o.settings.UpdateSecurity(ctx, SecurityPatch{AppURL: strPtr("https://b.example.com")}); err != nil { t.Fatalf("change app url: %v", err) } // 旧地址的通行密钥不计入可用方式:禁用密码被拒,不再自锁 if err := auth.SetPasswordLoginDisabled(ctx, "admin", true, proofOf(t, auth.db, "admin")); !errors.Is(err, ErrNeedIdentity) { t.Fatalf("disable with stale-origin passkey err = %v, want ErrNeedIdentity", err) } // 当前地址重新注册(origin 一致)后可禁用 pk2 := model.UserPasskey{UserID: 1, Name: "k2", CredentialID: "c2", CredentialIDHash: "h2", Origin: "https://b.example.com", Credential: "{}"} if err := auth.db.Create(&pk2).Error; err != nil { t.Fatalf("seed passkey2: %v", err) } if err := auth.SetPasswordLoginDisabled(ctx, "admin", true, proofOf(t, auth.db, "admin")); err != nil { t.Fatalf("disable with current-origin passkey: %v", err) } } func TestAuthenticatedSettingsRejectRevokedProof(t *testing.T) { tests := []struct { name string apply func(*OAuthService, *AuthService, TokenProof) error }{ { name: "security patch", apply: func(o *OAuthService, auth *AuthService, proof TokenProof) error { return o.settings.UpdateSecurityAuthenticated( context.Background(), SecurityPatch{LoginFailLimit: intPtr(9)}, auth, "admin", proof) }, }, { name: "oauth patch", apply: func(o *OAuthService, auth *AuthService, proof TokenProof) error { return o.settings.UpdateOAuthAuthenticated( context.Background(), UpdateOAuthInput{GithubDisplayName: strPtr("late")}, auth, "admin", proof) }, }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { o, auth := newOAuthEnv(t) proof := revokeSettingsToken(t, auth) if err := tt.apply(o, auth, proof); !errors.Is(err, ErrTokenStale) { t.Fatalf("patch err = %v, want ErrTokenStale", err) } }) } } func revokeSettingsToken(t *testing.T, auth *AuthService) TokenProof { t.Helper() stale := loginSession(t, auth, "10.20.0.1", "stale") current := loginSession(t, auth, "10.20.0.2", "current") proof := mustProof(t, auth, stale) for _, item := range mustSessions(t, auth, current) { if !item.Current { if err := auth.RevokeSession(context.Background(), "admin", current, item.ID); err != nil { t.Fatalf("RevokeSession: %v", err) } } } return proof } 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 TestOAuthLoginRechecksIdentityAndProvider(t *testing.T) { tests := []struct { name string mutate func(*testing.T, *OAuthService, *model.UserIdentity) want error }{ { name: "identity unbound", mutate: func(t *testing.T, o *OAuthService, row *model.UserIdentity) { if err := o.db.Delete(row).Error; err != nil { t.Fatalf("delete identity: %v", err) } }, want: ErrOAuthNotBound, }, { name: "provider disabled", mutate: func(t *testing.T, o *OAuthService, _ *model.UserIdentity) { patch := UpdateOAuthInput{GithubDisabled: boolPtr(true)} if err := o.settings.UpdateOAuth(context.Background(), patch); err != nil { t.Fatalf("disable provider: %v", err) } }, want: ErrOAuthDisabled, }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { testOAuthLoginRecheck(t, tt.mutate, tt.want) }) } } func testOAuthLoginRecheck( t *testing.T, mutate func(*testing.T, *OAuthService, *model.UserIdentity), want error, ) { t.Helper() o, _ := newOAuthEnv(t) o.settings.SetEnvPublicURL("https://demo.example.com") secret := "secret" if err := o.settings.UpdateOAuth(context.Background(), UpdateOAuthInput{ GithubClientID: strPtr("cid"), GithubClientSecret: &secret, }); err != nil { t.Fatalf("UpdateOAuth: %v", err) } fakeBindIdentity(t, o, "admin", "github", "10086", "octocat") row, err := o.findIdentity(context.Background(), "github", "10086") if err != nil { t.Fatalf("findIdentity: %v", err) } mutate(t, o, row) _, _, err = o.loginIdentityRow( context.Background(), row, "github", "10086", SessionMeta{ClientIP: "127.0.0.1"}) if !errors.Is(err, want) { t.Fatalf("login err = %v, want %v", err, want) } } func TestOAuthBindLoginUnbind(t *testing.T) { o, auth := newOAuthEnv(t) ctx := context.Background() o.settings.SetEnvPublicURL("https://demo.example.com") secret := "secret" if err := o.settings.UpdateOAuth(ctx, UpdateOAuthInput{ GithubClientID: strPtr("cid"), GithubClientSecret: &secret, }); err != nil { t.Fatalf("UpdateOAuth: %v", err) } fakeBindIdentity(t, o, "admin", "github", "10086", "octocat") // 重复绑定同一身份拒绝 if _, err := o.bind(ctx, fakePending(t, o, "admin", "github"), externalIdentity{Subject: "10086", Display: "octocat"}, SessionMeta{}); !errors.Is(err, ErrOAuthBound) { t.Errorf("重复绑定 err = %v, want ErrOAuthBound", err) } // 发起后版本已变(改密/撤销全部):事务内比对拒绝,身份不落库 stale := fakePending(t, o, "admin", "github") stale.proof.Ver-- if _, err := o.bind(ctx, stale, externalIdentity{Subject: "20250", Display: "x"}, SessionMeta{}); !errors.Is(err, ErrOAuthState) { t.Errorf("stale-version bind err = %v, want ErrOAuthState", err) } // 已绑定身份可登录并拿到有效 JWT token, loginUser, err := o.loginByIdentity(ctx, "github", externalIdentity{Subject: "10086", Display: "octocat"}, SessionMeta{ClientIP: "127.0.0.1"}) 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"}, SessionMeta{ClientIP: "127.0.0.1"}); !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, proofOf(t, o.db, "admin")); 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) } // 面板地址缺失时回调无从拼接:半配置不暴露必败入口 if got := o.Providers(ctx); len(got) != 0 { t.Fatalf("无面板地址 providers = %v, want empty", got) } o.settings.SetEnvPublicURL("https://demo.example.com") 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 的半配置仍按未配置拒绝(不能生成必败授权 URL) 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, ErrOAuthNotConfigured) { t.Errorf("无 secret err = %v, want ErrOAuthNotConfigured", err) } // clientID + secret 已配但面板地址未设置 secret := "secret" if err := o.settings.UpdateOAuth(ctx, UpdateOAuthInput{GithubClientSecret: &secret}); err != nil { t.Fatalf("UpdateOAuth secret: %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, auth := 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;仅 clientID 的半配置不暴露 in := UpdateOAuthInput{GithubClientID: strPtr("Iv1.test")} if err := o.settings.UpdateOAuth(ctx, in); err != nil { t.Fatalf("UpdateOAuth: %v", err) } if got := o.Providers(ctx); len(got) != 0 { t.Fatalf("仅 clientID 半配置 Providers = %v, want empty", got) } ghSecret := "gh-secret" in.GithubClientSecret = &ghSecret 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) } // bind 模式须携带有效令牌(发起即验),但不受 provider 禁用影响 token, _, err := auth.Login(ctx, "admin", "pass123", "", SessionMeta{ClientIP: "10.0.0.9"}) if err != nil { t.Fatalf("login: %v", err) } if url, err := o.AuthorizeURL(ctx, "github", "bind", "admin", token); err != nil || url == "" { t.Errorf("禁用 bind = %q, %v, want 正常返回", url, err) } if _, err := o.AuthorizeURL(ctx, "github", "bind", "admin", ""); !errors.Is(err, ErrOAuthState) { t.Errorf("bind 无令牌 err = %v, want ErrOAuthState", 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) } }) } }