349 lines
11 KiB
Go
349 lines
11 KiB
Go
package service
|
|
|
|
import (
|
|
"context"
|
|
"fmt"
|
|
"strings"
|
|
|
|
"gorm.io/gorm"
|
|
)
|
|
|
|
// OAuth provider 配置键;client secret 以 AES-GCM 密文落库。
|
|
const (
|
|
settingOauthOidcIssuer = "oauth_oidc_issuer"
|
|
settingOauthOidcClientID = "oauth_oidc_client_id"
|
|
settingOauthOidcClientSecret = "oauth_oidc_client_secret"
|
|
settingOauthOidcDisplayName = "oauth_oidc_display_name"
|
|
settingOauthOidcDisabled = "oauth_oidc_disabled"
|
|
settingOauthGithubClientID = "oauth_github_client_id"
|
|
settingOauthGithubClientSecret = "oauth_github_client_secret"
|
|
settingOauthGithubDisplayName = "oauth_github_display_name"
|
|
settingOauthGithubDisabled = "oauth_github_disabled"
|
|
)
|
|
|
|
// OAuthProvidersView 是 OAuth provider 配置视图;绝不返回 secret 明文。
|
|
type OAuthProvidersView struct {
|
|
OidcIssuer string `json:"oidcIssuer"`
|
|
OidcClientID string `json:"oidcClientId"`
|
|
OidcSecretSet bool `json:"oidcSecretSet"`
|
|
OidcDisplayName string `json:"oidcDisplayName"`
|
|
OidcDisabled bool `json:"oidcDisabled"`
|
|
GithubClientID string `json:"githubClientId"`
|
|
GithubSecretSet bool `json:"githubSecretSet"`
|
|
GithubDisplayName string `json:"githubDisplayName"`
|
|
GithubDisabled bool `json:"githubDisabled"`
|
|
}
|
|
|
|
// UpdateOAuthInput 是 provider 配置的部分更新:所有字段 nil=沿用已存值,
|
|
// 只落库出现的字段;并发编辑不同 provider 因此互不回滚(2026-07-22 审查 #18)。
|
|
// secret 非 nil 覆盖(空串清除),绝不回读。
|
|
type UpdateOAuthInput struct {
|
|
OidcIssuer *string `json:"oidcIssuer"`
|
|
OidcClientID *string `json:"oidcClientId"`
|
|
OidcClientSecret *string `json:"oidcClientSecret"`
|
|
OidcDisplayName *string `json:"oidcDisplayName"`
|
|
OidcDisabled *bool `json:"oidcDisabled"`
|
|
GithubClientID *string `json:"githubClientId"`
|
|
GithubClientSecret *string `json:"githubClientSecret"`
|
|
GithubDisplayName *string `json:"githubDisplayName"`
|
|
GithubDisabled *bool `json:"githubDisabled"`
|
|
}
|
|
|
|
// OAuthView 返回脱敏后的 provider 配置。
|
|
func (s *SettingService) OAuthView(ctx context.Context) (OAuthProvidersView, error) {
|
|
var view OAuthProvidersView
|
|
vals, err := s.getMany(ctx,
|
|
settingOauthOidcIssuer, settingOauthOidcClientID, settingOauthOidcClientSecret,
|
|
settingOauthOidcDisplayName, settingOauthOidcDisabled,
|
|
settingOauthGithubClientID, settingOauthGithubClientSecret,
|
|
settingOauthGithubDisplayName, settingOauthGithubDisabled)
|
|
if err != nil {
|
|
return view, err
|
|
}
|
|
view.OidcIssuer = vals[settingOauthOidcIssuer]
|
|
view.OidcClientID = vals[settingOauthOidcClientID]
|
|
view.OidcSecretSet = vals[settingOauthOidcClientSecret] != ""
|
|
view.OidcDisplayName = vals[settingOauthOidcDisplayName]
|
|
view.OidcDisabled = vals[settingOauthOidcDisabled] == "1"
|
|
view.GithubClientID = vals[settingOauthGithubClientID]
|
|
view.GithubSecretSet = vals[settingOauthGithubClientSecret] != ""
|
|
view.GithubDisplayName = vals[settingOauthGithubDisplayName]
|
|
view.GithubDisabled = vals[settingOauthGithubDisabled] == "1"
|
|
return view, nil
|
|
}
|
|
|
|
// boolFlag 把开关序列化为 settings 存储值。
|
|
func boolFlag(on bool) string {
|
|
if on {
|
|
return "1"
|
|
}
|
|
return ""
|
|
}
|
|
|
|
// trimPtr / issuerPtr / flagPtr 把补丁字段规范化为存储值;nil 表示未出现不写。
|
|
func trimPtr(p *string) *string {
|
|
if p == nil {
|
|
return nil
|
|
}
|
|
v := strings.TrimSpace(*p)
|
|
return &v
|
|
}
|
|
|
|
func issuerPtr(p *string) *string {
|
|
if p == nil {
|
|
return nil
|
|
}
|
|
v := strings.TrimRight(strings.TrimSpace(*p), "/")
|
|
return &v
|
|
}
|
|
|
|
func flagPtr(p *bool) *string {
|
|
if p == nil {
|
|
return nil
|
|
}
|
|
v := boolFlag(*p)
|
|
return &v
|
|
}
|
|
|
|
// UpdateOAuth 部分更新 provider 配置:只落库非 nil 字段;issuer 规范化去尾斜杠,
|
|
// secret 加密落库(空串清除)。预检与写入在同一事务并持有认证变更共用的用户行锁:
|
|
// 密码登录禁用期间,禁止把最后可实际登录的方式禁用或清空(防自锁),
|
|
// 且不与「禁用密码 / 删除最后因子」并发交错。
|
|
func (s *SettingService) UpdateOAuth(ctx context.Context, in UpdateOAuthInput) error {
|
|
return s.updateOAuth(ctx, in, nil)
|
|
}
|
|
|
|
// UpdateOAuthAuthenticated 在写事务持锁后复核请求令牌,防撤销后的慢请求落库。
|
|
func (s *SettingService) UpdateOAuthAuthenticated(
|
|
ctx context.Context, in UpdateOAuthInput, auth *AuthService, username string, proof TokenProof,
|
|
) error {
|
|
check := &authenticatedMutation{auth: auth, username: username, proof: proof}
|
|
return s.updateOAuth(ctx, in, check)
|
|
}
|
|
|
|
func (s *SettingService) updateOAuth(ctx context.Context, in UpdateOAuthInput, check *authenticatedMutation) error {
|
|
secrets, err := s.encryptOAuthSecrets(in)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
|
if check != nil {
|
|
err = check.lockAndCheck(tx)
|
|
} else {
|
|
err = lockUsersForAuthChange(tx)
|
|
}
|
|
if err != nil {
|
|
return err
|
|
}
|
|
origin, err := effectiveOriginTx(tx, s)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if err := ensureLoginRemainsTx(tx, in, origin); err != nil {
|
|
return err
|
|
}
|
|
return writeOAuthTx(tx, in, secrets)
|
|
})
|
|
}
|
|
|
|
// encryptOAuthSecrets 预先加密补丁中的 secret(nil 沿用,空串清除),事务内直接落库。
|
|
func (s *SettingService) encryptOAuthSecrets(in UpdateOAuthInput) (map[string]*string, error) {
|
|
out := map[string]*string{}
|
|
for key, sec := range map[string]*string{
|
|
settingOauthOidcClientSecret: in.OidcClientSecret,
|
|
settingOauthGithubClientSecret: in.GithubClientSecret,
|
|
} {
|
|
if sec == nil {
|
|
continue
|
|
}
|
|
v := ""
|
|
if *sec != "" {
|
|
enc, err := s.cipher.EncryptString(*sec)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("encrypt oauth secret: %w", err)
|
|
}
|
|
v = enc
|
|
}
|
|
out[key] = &v
|
|
}
|
|
return out, nil
|
|
}
|
|
|
|
// writeOAuthTx 事务内落库补丁中出现的字段。
|
|
func writeOAuthTx(tx *gorm.DB, in UpdateOAuthInput, secrets map[string]*string) error {
|
|
writes := []struct {
|
|
key string
|
|
val *string
|
|
}{
|
|
{settingOauthOidcIssuer, issuerPtr(in.OidcIssuer)},
|
|
{settingOauthOidcClientID, trimPtr(in.OidcClientID)},
|
|
{settingOauthOidcClientSecret, secrets[settingOauthOidcClientSecret]},
|
|
{settingOauthOidcDisplayName, trimPtr(in.OidcDisplayName)},
|
|
{settingOauthOidcDisabled, flagPtr(in.OidcDisabled)},
|
|
{settingOauthGithubClientID, trimPtr(in.GithubClientID)},
|
|
{settingOauthGithubClientSecret, secrets[settingOauthGithubClientSecret]},
|
|
{settingOauthGithubDisplayName, trimPtr(in.GithubDisplayName)},
|
|
{settingOauthGithubDisabled, flagPtr(in.GithubDisabled)},
|
|
}
|
|
for _, w := range writes {
|
|
if w.val == nil {
|
|
continue
|
|
}
|
|
if err := saveSettingTx(tx, w.key, *w.val); err != nil {
|
|
return err
|
|
}
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// ensureLoginRemainsTx 事务内校验补丁生效后仍有可实际登录的方式:密码可登直接放行;
|
|
// 否则须有「可登录且已绑定身份」的 provider、任一通行密钥或钱包身份
|
|
// (单管理员面板,不区分账号统计)。开关读取失败按失败关闭处理。
|
|
func ensureLoginRemainsTx(tx *gorm.DB, in UpdateOAuthInput, origin string) error {
|
|
off, err := settingValueTx(tx, settingSecPasswordLoginOff)
|
|
if err != nil || off != "1" {
|
|
return err
|
|
}
|
|
if origin == "" {
|
|
return ErrProviderLastLogin
|
|
}
|
|
view, err := oauthViewTx(tx)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if ok, err := anyBoundUsableProviderTx(tx, patchedUsable(view, in)); err != nil || ok {
|
|
return err
|
|
}
|
|
// 通行密钥兜底须当前地址下可用(origin 一致);0 表示不排除任何行
|
|
if n, err := passkeyCountExcludingTx(tx, 0, 0, origin); err != nil || n > 0 {
|
|
return err
|
|
}
|
|
n, err := identityProviderCountTx(tx, "wallet")
|
|
if err != nil || n > 0 {
|
|
return err
|
|
}
|
|
return ErrProviderLastLogin
|
|
}
|
|
|
|
// oauthViewTx 事务内读 provider 配置视图(secret 只取「已设置」布尔)。
|
|
func oauthViewTx(tx *gorm.DB) (OAuthProvidersView, error) {
|
|
var view OAuthProvidersView
|
|
reads := []struct {
|
|
key string
|
|
set func(string)
|
|
}{
|
|
{settingOauthOidcIssuer, func(v string) { view.OidcIssuer = v }},
|
|
{settingOauthOidcClientID, func(v string) { view.OidcClientID = v }},
|
|
{settingOauthOidcClientSecret, func(v string) { view.OidcSecretSet = v != "" }},
|
|
{settingOauthOidcDisabled, func(v string) { view.OidcDisabled = v == "1" }},
|
|
{settingOauthGithubClientID, func(v string) { view.GithubClientID = v }},
|
|
{settingOauthGithubClientSecret, func(v string) { view.GithubSecretSet = v != "" }},
|
|
{settingOauthGithubDisabled, func(v string) { view.GithubDisabled = v == "1" }},
|
|
}
|
|
for _, r := range reads {
|
|
v, err := settingValueTx(tx, r.key)
|
|
if err != nil {
|
|
return view, err
|
|
}
|
|
r.set(v)
|
|
}
|
|
return view, nil
|
|
}
|
|
|
|
// patchedOidcUsable 计算补丁生效后 OIDC 是否可登录(clientID/secret/issuer 齐备且未禁用)。
|
|
func patchedOidcUsable(view OAuthProvidersView, in UpdateOAuthInput) bool {
|
|
id, sec, iss, off := view.OidcClientID, view.OidcSecretSet, view.OidcIssuer, view.OidcDisabled
|
|
if v := trimPtr(in.OidcClientID); v != nil {
|
|
id = *v
|
|
}
|
|
if in.OidcClientSecret != nil {
|
|
sec = *in.OidcClientSecret != ""
|
|
}
|
|
if v := issuerPtr(in.OidcIssuer); v != nil {
|
|
iss = *v
|
|
}
|
|
if in.OidcDisabled != nil {
|
|
off = *in.OidcDisabled
|
|
}
|
|
return id != "" && sec && iss != "" && !off
|
|
}
|
|
|
|
// patchedGithubUsable 计算补丁生效后 GitHub 是否可登录(clientID/secret 齐备且未禁用)。
|
|
func patchedGithubUsable(view OAuthProvidersView, in UpdateOAuthInput) bool {
|
|
id, sec, off := view.GithubClientID, view.GithubSecretSet, view.GithubDisabled
|
|
if v := trimPtr(in.GithubClientID); v != nil {
|
|
id = *v
|
|
}
|
|
if in.GithubClientSecret != nil {
|
|
sec = *in.GithubClientSecret != ""
|
|
}
|
|
if in.GithubDisabled != nil {
|
|
off = *in.GithubDisabled
|
|
}
|
|
return id != "" && sec && !off
|
|
}
|
|
|
|
// patchedUsable 汇总补丁生效后各 provider 的可登录性。
|
|
func patchedUsable(view OAuthProvidersView, in UpdateOAuthInput) map[string]bool {
|
|
return map[string]bool{
|
|
"oidc": patchedOidcUsable(view, in),
|
|
"github": patchedGithubUsable(view, in),
|
|
}
|
|
}
|
|
|
|
// anyBoundUsableProviderTx 事务内判断是否存在「可登录且已有绑定身份」的 provider。
|
|
func anyBoundUsableProviderTx(tx *gorm.DB, usable map[string]bool) (bool, error) {
|
|
for p, ok := range usable {
|
|
if !ok {
|
|
continue
|
|
}
|
|
n, err := identityProviderCountTx(tx, p)
|
|
if err != nil {
|
|
return false, err
|
|
}
|
|
if n > 0 {
|
|
return true, nil
|
|
}
|
|
}
|
|
return false, nil
|
|
}
|
|
|
|
// oauthClient 返回 provider 的 clientID/明文 secret/issuer(仅 oidc);未配置时 clientID 为空。
|
|
func (s *SettingService) oauthClient(ctx context.Context, provider string) (clientID, secret, issuer string, err error) {
|
|
idKey, secKey := settingOauthGithubClientID, settingOauthGithubClientSecret
|
|
if provider == "oidc" {
|
|
idKey, secKey = settingOauthOidcClientID, settingOauthOidcClientSecret
|
|
if issuer, err = s.get(ctx, settingOauthOidcIssuer); err != nil {
|
|
return
|
|
}
|
|
}
|
|
if clientID, err = s.get(ctx, idKey); err != nil {
|
|
return
|
|
}
|
|
enc, err := s.get(ctx, secKey)
|
|
if err != nil || enc == "" {
|
|
return
|
|
}
|
|
if secret, err = s.cipher.DecryptString(enc); err != nil {
|
|
err = fmt.Errorf("decrypt oauth secret: %w", err)
|
|
}
|
|
return
|
|
}
|
|
|
|
// oauthProviderMeta 返回 provider 的展示名(空则给默认名)与禁用态。
|
|
func (s *SettingService) oauthProviderMeta(ctx context.Context, provider string) (display string, disabled bool, err error) {
|
|
nameKey, offKey, def := settingOauthGithubDisplayName, settingOauthGithubDisabled, "GitHub"
|
|
if provider == "oidc" {
|
|
nameKey, offKey, def = settingOauthOidcDisplayName, settingOauthOidcDisabled, "OIDC 单点登录"
|
|
}
|
|
if display, err = s.get(ctx, nameKey); err != nil {
|
|
return
|
|
}
|
|
if display == "" {
|
|
display = def
|
|
}
|
|
off, err := s.get(ctx, offKey)
|
|
disabled = off == "1"
|
|
return
|
|
}
|