9501751792
三轮审查修复(60+ 项),相对远端 main(b536672)的关键变更:
- 安全: 数据面 SSRF 拨号防护(防 DNS rebinding)/上游凭据剥离/登录防枚举
与锁定态统一/可信代理(X-Forwarded-For)限流加固/会话版本失效机制/
撤销即时传播/弱密钥拒绝启动/脱敏字节级重写(保签名契约)
- 业务逻辑: 裸 body 上传 panic/bootstrap 审计管线卡死/定价通配符优先级/
全局工具可见性/调度器停机补跑/TOTP 挑战令牌消费顺序/熔断探针语义/
>4MB 响应 token 计量/管理员重置密码作废会话 等
- 前端: 新 logo(语枢 AI 网关主题)/Provider 凭据异常警示/删除入口/
后端错误消息透传/localStorage 敏感数据收敛
- 部署: CREDENTIAL_MASTER_KEY 持久化与弱值拒绝/Provider DELETE 接口/
nginx 安全头/worker 内存限制
- 新增迁移 000029(key_hash 索引)/000030(usage_daily 币种维度)
854 lines
32 KiB
Go
854 lines
32 KiB
Go
package identity
|
|
|
|
import (
|
|
"context"
|
|
"crypto"
|
|
"crypto/rand"
|
|
"crypto/rsa"
|
|
"crypto/sha256"
|
|
"encoding/base64"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"math/big"
|
|
"net"
|
|
"net/http"
|
|
"net/url"
|
|
"regexp"
|
|
"strings"
|
|
"time"
|
|
|
|
"aigateway.local/core/internal/platform/apiresponse"
|
|
platformid "aigateway.local/core/internal/platform/id"
|
|
"github.com/jackc/pgx/v5"
|
|
)
|
|
|
|
const oidcMaxResponse = 2 << 20
|
|
|
|
// oidcMaxTokenLifetimeSeconds bounds how far an ID token's exp may sit past
|
|
// its iat, preventing long-lived or replayed tokens from being accepted.
|
|
const oidcMaxTokenLifetimeSeconds = 24 * 3600
|
|
|
|
var oidcCodePattern = regexp.MustCompile(`^[a-z][a-z0-9_-]{1,63}$`)
|
|
|
|
type OIDCProvider struct {
|
|
ID string `json:"id"`
|
|
Code string `json:"code"`
|
|
DisplayName string `json:"display_name"`
|
|
IssuerURL string `json:"issuer_url"`
|
|
ClientID string `json:"client_id"`
|
|
EncryptedCredentials []byte `json:"-"`
|
|
CredentialKEKVersion int `json:"-"`
|
|
RedirectURI string `json:"redirect_uri"`
|
|
PortalReturnURL string `json:"portal_return_url"`
|
|
Scopes []string `json:"scopes"`
|
|
AutoProvision bool `json:"auto_provision"`
|
|
DefaultDepartmentID *string `json:"default_department_id,omitempty"`
|
|
Enabled bool `json:"enabled"`
|
|
Revision int64 `json:"revision"`
|
|
CreatedAt time.Time `json:"created_at"`
|
|
UpdatedAt time.Time `json:"updated_at"`
|
|
}
|
|
|
|
type oidcProviderInput struct {
|
|
Code string `json:"code"`
|
|
DisplayName string `json:"display_name"`
|
|
IssuerURL string `json:"issuer_url"`
|
|
ClientID string `json:"client_id"`
|
|
ClientSecret *string `json:"client_secret"`
|
|
RedirectURI string `json:"redirect_uri"`
|
|
PortalReturnURL string `json:"portal_return_url"`
|
|
Scopes []string `json:"scopes"`
|
|
AutoProvision bool `json:"auto_provision"`
|
|
DefaultDepartmentID *string `json:"default_department_id"`
|
|
Enabled bool `json:"enabled"`
|
|
}
|
|
|
|
type oidcCredentials struct {
|
|
ClientSecret string `json:"client_secret"`
|
|
}
|
|
type oidcChallenge struct{ ProviderID, Verifier, Nonce string }
|
|
type oidcExchange struct {
|
|
Token string `json:"token"`
|
|
}
|
|
type oidcDiscovery struct{ Issuer, AuthorizationEndpoint, TokenEndpoint, JWKSURI string }
|
|
|
|
func (h *ManagementHTTPHandler) registerOIDC() {
|
|
h.mux.HandleFunc("GET /api/v1/admin/identity-providers", h.listOIDCProviders)
|
|
h.mux.HandleFunc("POST /api/v1/admin/identity-providers", h.createOIDCProvider)
|
|
h.mux.HandleFunc("PUT /api/v1/admin/identity-providers/{provider_id}", h.updateOIDCProvider)
|
|
h.registerSAML()
|
|
}
|
|
|
|
func (h *HTTPHandler) registerOIDC() {
|
|
h.mux.HandleFunc("GET /api/v1/portal/sso/providers", h.listPublicOIDCProviders)
|
|
h.mux.HandleFunc("GET /api/v1/portal/sso/{provider_code}/start", h.startSSO)
|
|
h.mux.HandleFunc("GET /api/v1/portal/sso/{provider_code}/callback", h.callbackOIDC)
|
|
h.mux.HandleFunc("POST /api/v1/portal/sso/{provider_code}/callback", h.callbackSAML)
|
|
h.mux.HandleFunc("GET /api/v1/portal/sso/{provider_code}/metadata", h.samlMetadata)
|
|
h.mux.HandleFunc("POST /api/v1/portal/sso/exchange", h.exchangeOIDC)
|
|
}
|
|
|
|
func (h *ManagementHTTPHandler) listOIDCProviders(w http.ResponseWriter, r *http.Request) {
|
|
if _, ok := h.requirePermission(w, r); !ok {
|
|
return
|
|
}
|
|
records, err := h.service.repository.ListOIDCProviders(r.Context(), false)
|
|
if err != nil {
|
|
h.writeError(w, err)
|
|
return
|
|
}
|
|
items := make([]map[string]any, 0, len(records))
|
|
for _, record := range records {
|
|
items = append(items, h.oidcView(record))
|
|
}
|
|
apiresponse.OK(w, items)
|
|
}
|
|
|
|
func (h *ManagementHTTPHandler) createOIDCProvider(w http.ResponseWriter, r *http.Request) {
|
|
actor, ok := h.requirePermission(w, r)
|
|
if !ok {
|
|
return
|
|
}
|
|
input, record, ok := h.decodeOIDC(w, r, true)
|
|
if !ok {
|
|
return
|
|
}
|
|
if err := h.setOIDCCredentials(&record, strings.TrimSpace(*input.ClientSecret)); err != nil {
|
|
h.writeError(w, err)
|
|
return
|
|
}
|
|
created, err := h.service.repository.CreateOIDCProvider(r.Context(), record, actor.ID)
|
|
if err != nil {
|
|
h.writeError(w, err)
|
|
return
|
|
}
|
|
apiresponse.OK(w, h.oidcView(created))
|
|
}
|
|
|
|
func (h *ManagementHTTPHandler) updateOIDCProvider(w http.ResponseWriter, r *http.Request) {
|
|
actor, ok := h.requirePermission(w, r)
|
|
if !ok {
|
|
return
|
|
}
|
|
input, record, ok := h.decodeOIDC(w, r, false)
|
|
if !ok {
|
|
return
|
|
}
|
|
record.ID = r.PathValue("provider_id")
|
|
replace := input.ClientSecret != nil
|
|
if replace {
|
|
if err := h.setOIDCCredentials(&record, strings.TrimSpace(*input.ClientSecret)); err != nil {
|
|
h.writeError(w, err)
|
|
return
|
|
}
|
|
}
|
|
updated, err := h.service.repository.UpdateOIDCProvider(r.Context(), record, actor.ID, replace)
|
|
if err != nil {
|
|
h.writeError(w, err)
|
|
return
|
|
}
|
|
apiresponse.OK(w, h.oidcView(updated))
|
|
}
|
|
|
|
func (h *ManagementHTTPHandler) decodeOIDC(w http.ResponseWriter, r *http.Request, creating bool) (oidcProviderInput, OIDCProvider, bool) {
|
|
var input oidcProviderInput
|
|
decoder := json.NewDecoder(http.MaxBytesReader(w, r.Body, 1<<20))
|
|
decoder.DisallowUnknownFields()
|
|
if decoder.Decode(&input) != nil {
|
|
apiresponse.Error(w, 400, "请求格式无效")
|
|
return input, OIDCProvider{}, false
|
|
}
|
|
input.Code = strings.ToLower(strings.TrimSpace(input.Code))
|
|
input.DisplayName = strings.TrimSpace(input.DisplayName)
|
|
input.ClientID = strings.TrimSpace(input.ClientID)
|
|
if !oidcCodePattern.MatchString(input.Code) || input.DisplayName == "" || input.ClientID == "" ||
|
|
(creating && input.ClientSecret == nil) || (input.ClientSecret != nil && strings.TrimSpace(*input.ClientSecret) == "") {
|
|
apiresponse.Error(w, 400, "身份源代码、名称、Client ID 或 Client Secret 无效")
|
|
return input, OIDCProvider{}, false
|
|
}
|
|
issuer, err := validateOIDCURL(r.Context(), input.IssuerURL, h.service.allowPrivateIdentityProvider)
|
|
if err != nil {
|
|
apiresponse.Error(w, 400, "Issuer URL 无效")
|
|
return input, OIDCProvider{}, false
|
|
}
|
|
redirectURI, err := validateAbsoluteURL(input.RedirectURI)
|
|
redirectURL, _ := url.Parse(redirectURI)
|
|
if err != nil || redirectURL.Fragment != "" {
|
|
apiresponse.Error(w, 400, "回调 URL 无效")
|
|
return input, OIDCProvider{}, false
|
|
}
|
|
returnURL, err := validateAbsoluteURL(input.PortalReturnURL)
|
|
if err != nil {
|
|
apiresponse.Error(w, 400, "门户返回 URL 无效")
|
|
return input, OIDCProvider{}, false
|
|
}
|
|
if len(input.Scopes) == 0 {
|
|
input.Scopes = []string{"openid", "profile", "email"}
|
|
}
|
|
input.Scopes, err = normalizeOIDCScopes(input.Scopes)
|
|
if err != nil {
|
|
apiresponse.Error(w, 400, "OIDC scopes 无效")
|
|
return input, OIDCProvider{}, false
|
|
}
|
|
if !contains(input.Scopes, "openid") {
|
|
apiresponse.Error(w, 400, "OIDC scopes 必须包含 openid")
|
|
return input, OIDCProvider{}, false
|
|
}
|
|
if input.DefaultDepartmentID != nil && *input.DefaultDepartmentID != "" {
|
|
department, err := h.service.repository.GetDepartment(r.Context(), *input.DefaultDepartmentID)
|
|
if err != nil || !department.Active {
|
|
apiresponse.Error(w, 400, "默认部门不存在或已停用")
|
|
return input, OIDCProvider{}, false
|
|
}
|
|
}
|
|
return input, OIDCProvider{Code: input.Code, DisplayName: input.DisplayName, IssuerURL: issuer, ClientID: input.ClientID,
|
|
RedirectURI: redirectURI, PortalReturnURL: returnURL, Scopes: input.Scopes, AutoProvision: input.AutoProvision,
|
|
DefaultDepartmentID: input.DefaultDepartmentID, Enabled: input.Enabled}, true
|
|
}
|
|
|
|
func (h *ManagementHTTPHandler) setOIDCCredentials(record *OIDCProvider, secret string) error {
|
|
payload, _ := json.Marshal(oidcCredentials{ClientSecret: secret})
|
|
encrypted, version, err := h.service.idpCipher.Encrypt(payload)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
record.EncryptedCredentials, record.CredentialKEKVersion = encrypted, version
|
|
return nil
|
|
}
|
|
|
|
func (h *ManagementHTTPHandler) oidcView(record OIDCProvider) map[string]any {
|
|
configured := false
|
|
if plaintext, err := h.service.idpCipher.Decrypt(record.EncryptedCredentials, record.CredentialKEKVersion); err == nil {
|
|
var credentials oidcCredentials
|
|
configured = json.Unmarshal(plaintext, &credentials) == nil && credentials.ClientSecret != ""
|
|
}
|
|
return map[string]any{"id": record.ID, "code": record.Code, "display_name": record.DisplayName, "issuer_url": record.IssuerURL,
|
|
"client_id": record.ClientID, "secret_configured": configured, "redirect_uri": record.RedirectURI, "portal_return_url": record.PortalReturnURL,
|
|
"scopes": record.Scopes, "auto_provision": record.AutoProvision, "default_department_id": record.DefaultDepartmentID,
|
|
"enabled": record.Enabled, "revision": record.Revision, "credential_kek_version": record.CredentialKEKVersion}
|
|
}
|
|
|
|
func (h *HTTPHandler) listPublicOIDCProviders(w http.ResponseWriter, r *http.Request) {
|
|
records, err := h.service.repository.ListPublicIdentityProviders(r.Context())
|
|
if err != nil {
|
|
h.writeIdentityError(w, err)
|
|
return
|
|
}
|
|
apiresponse.OK(w, records)
|
|
}
|
|
|
|
func (h *HTTPHandler) startOIDC(w http.ResponseWriter, r *http.Request) {
|
|
p, err := h.service.repository.GetOIDCProviderByCode(r.Context(), r.PathValue("provider_code"))
|
|
if err != nil || !p.Enabled {
|
|
http.NotFound(w, r)
|
|
return
|
|
}
|
|
discovery, err := h.service.discoverOIDC(r.Context(), p)
|
|
if err != nil {
|
|
apiresponse.Error(w, 502, "身份源发现失败")
|
|
return
|
|
}
|
|
verifier, err := randomURLToken(48)
|
|
if err != nil {
|
|
h.writeIdentityError(w, ErrUnavailable)
|
|
return
|
|
}
|
|
nonce, err := randomURLToken(32)
|
|
if err != nil {
|
|
h.writeIdentityError(w, ErrUnavailable)
|
|
return
|
|
}
|
|
state, err := h.service.sessions.StoreOneTime(r.Context(), "oidc-state", oidcChallenge{ProviderID: p.ID, Verifier: verifier, Nonce: nonce}, 5*time.Minute)
|
|
if err != nil {
|
|
h.writeIdentityError(w, err)
|
|
return
|
|
}
|
|
challenge := sha256.Sum256([]byte(verifier))
|
|
target, _ := url.Parse(discovery.AuthorizationEndpoint)
|
|
query := target.Query()
|
|
query.Set("response_type", "code")
|
|
query.Set("client_id", p.ClientID)
|
|
query.Set("redirect_uri", p.RedirectURI)
|
|
query.Set("scope", strings.Join(p.Scopes, " "))
|
|
query.Set("state", state)
|
|
query.Set("nonce", nonce)
|
|
query.Set("code_challenge", base64.RawURLEncoding.EncodeToString(challenge[:]))
|
|
query.Set("code_challenge_method", "S256")
|
|
target.RawQuery = query.Encode()
|
|
http.Redirect(w, r, target.String(), http.StatusFound)
|
|
}
|
|
|
|
func (h *HTTPHandler) callbackOIDC(w http.ResponseWriter, r *http.Request) {
|
|
if r.URL.Query().Get("error") != "" {
|
|
apiresponse.Error(w, 401, "OIDC 登录被拒绝")
|
|
return
|
|
}
|
|
if r.URL.Query().Get("state") == "" || r.URL.Query().Get("code") == "" {
|
|
apiresponse.Error(w, 401, "OIDC 回调参数无效")
|
|
return
|
|
}
|
|
var challenge oidcChallenge
|
|
if err := h.service.sessions.ConsumeOneTime(r.Context(), "oidc-state", r.URL.Query().Get("state"), &challenge); err != nil {
|
|
apiresponse.Error(w, 401, "OIDC state 无效或已使用")
|
|
return
|
|
}
|
|
p, err := h.service.repository.GetOIDCProviderByCode(r.Context(), r.PathValue("provider_code"))
|
|
if err != nil || p.ID != challenge.ProviderID || !p.Enabled {
|
|
apiresponse.Error(w, 401, "OIDC 身份源无效")
|
|
return
|
|
}
|
|
discovery, err := h.service.discoverOIDC(r.Context(), p)
|
|
if err != nil {
|
|
apiresponse.Error(w, 502, "身份源发现失败")
|
|
return
|
|
}
|
|
credentials, err := h.service.oidcCredentials(p)
|
|
if err != nil {
|
|
apiresponse.Error(w, 503, "身份源凭据不可用")
|
|
return
|
|
}
|
|
form := url.Values{"grant_type": {"authorization_code"}, "code": {r.URL.Query().Get("code")}, "redirect_uri": {p.RedirectURI}, "code_verifier": {challenge.Verifier}}
|
|
req, err := http.NewRequestWithContext(r.Context(), http.MethodPost, discovery.TokenEndpoint, strings.NewReader(form.Encode()))
|
|
if err != nil {
|
|
apiresponse.Error(w, 502, "OIDC token 交换失败")
|
|
return
|
|
}
|
|
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
|
req.SetBasicAuth(p.ClientID, credentials.ClientSecret)
|
|
response, err := h.service.oidcHTTPClient().Do(req)
|
|
if err != nil {
|
|
apiresponse.Error(w, 502, "OIDC token 交换失败")
|
|
return
|
|
}
|
|
defer response.Body.Close()
|
|
payload, _ := io.ReadAll(io.LimitReader(response.Body, oidcMaxResponse+1))
|
|
if response.StatusCode/100 != 2 || len(payload) > oidcMaxResponse {
|
|
apiresponse.Error(w, 502, "OIDC token 交换失败")
|
|
return
|
|
}
|
|
var tokens struct {
|
|
IDToken string `json:"id_token"`
|
|
}
|
|
if json.Unmarshal(payload, &tokens) != nil || tokens.IDToken == "" {
|
|
apiresponse.Error(w, 502, "OIDC 响应缺少 ID Token")
|
|
return
|
|
}
|
|
claims, err := h.service.verifyIDToken(r.Context(), discovery, p.ClientID, challenge.Nonce, tokens.IDToken)
|
|
if err != nil {
|
|
apiresponse.Error(w, 401, "OIDC ID Token 校验失败")
|
|
return
|
|
}
|
|
account, err := h.service.repository.ResolveExternalAccount(r.Context(), p, claims)
|
|
if err != nil {
|
|
h.writeIdentityError(w, err)
|
|
return
|
|
}
|
|
if !account.Active {
|
|
h.writeIdentityError(w, ErrAccountDisabled)
|
|
return
|
|
}
|
|
token, err := h.service.sessions.Create(r.Context(), principalFor(account))
|
|
if err != nil {
|
|
h.writeIdentityError(w, err)
|
|
return
|
|
}
|
|
exchange, err := h.service.sessions.StoreOneTime(r.Context(), "oidc-exchange", oidcExchange{Token: token}, time.Minute)
|
|
if err != nil {
|
|
h.writeIdentityError(w, err)
|
|
return
|
|
}
|
|
returnURL, _ := url.Parse(p.PortalReturnURL)
|
|
query := returnURL.Query()
|
|
query.Set("sso_code", exchange)
|
|
returnURL.RawQuery = query.Encode()
|
|
http.Redirect(w, r, returnURL.String(), http.StatusFound)
|
|
}
|
|
|
|
func (h *HTTPHandler) exchangeOIDC(w http.ResponseWriter, r *http.Request) {
|
|
var input struct {
|
|
Code string `json:"code"`
|
|
}
|
|
if !decodeJSON(w, r, &input) {
|
|
apiresponse.Error(w, 400, "请求格式无效")
|
|
return
|
|
}
|
|
var exchange oidcExchange
|
|
if h.service.sessions.ConsumeOneTime(r.Context(), "oidc-exchange", input.Code, &exchange) != nil {
|
|
apiresponse.Error(w, 401, "SSO 交换码无效或已使用")
|
|
return
|
|
}
|
|
apiresponse.OK(w, map[string]string{"token": exchange.Token, "refreshToken": ""})
|
|
}
|
|
|
|
type oidcClaims struct {
|
|
Issuer string `json:"iss"`
|
|
Subject string `json:"sub"`
|
|
Audience json.RawMessage `json:"aud"`
|
|
AuthorizedParty string `json:"azp"`
|
|
ExpiresAt int64 `json:"exp"`
|
|
IssuedAt int64 `json:"iat"`
|
|
NotBefore int64 `json:"nbf"`
|
|
Nonce string `json:"nonce"`
|
|
Email string `json:"email"`
|
|
PreferredUsername string `json:"preferred_username"`
|
|
Name string `json:"name"`
|
|
}
|
|
|
|
type externalProvider struct {
|
|
ID string
|
|
Code string
|
|
AuthSource string
|
|
AutoProvision bool
|
|
DefaultDepartmentID *string
|
|
}
|
|
|
|
type externalClaims struct {
|
|
Subject string
|
|
Email string
|
|
PreferredUsername string
|
|
Name string
|
|
}
|
|
|
|
func (s *Service) discoverOIDC(ctx context.Context, p OIDCProvider) (oidcDiscovery, error) {
|
|
target := strings.TrimRight(p.IssuerURL, "/") + "/.well-known/openid-configuration"
|
|
req, err := http.NewRequestWithContext(ctx, http.MethodGet, target, nil)
|
|
if err != nil {
|
|
return oidcDiscovery{}, err
|
|
}
|
|
response, err := s.oidcHTTPClient().Do(req)
|
|
if err != nil {
|
|
return oidcDiscovery{}, err
|
|
}
|
|
defer response.Body.Close()
|
|
payload, _ := io.ReadAll(io.LimitReader(response.Body, oidcMaxResponse+1))
|
|
var raw struct {
|
|
Issuer string `json:"issuer"`
|
|
AuthorizationEndpoint string `json:"authorization_endpoint"`
|
|
TokenEndpoint string `json:"token_endpoint"`
|
|
JWKSURI string `json:"jwks_uri"`
|
|
}
|
|
if response.StatusCode/100 != 2 || len(payload) > oidcMaxResponse || json.Unmarshal(payload, &raw) != nil || raw.Issuer != p.IssuerURL {
|
|
return oidcDiscovery{}, errors.New("invalid discovery")
|
|
}
|
|
for _, value := range []string{raw.AuthorizationEndpoint, raw.TokenEndpoint, raw.JWKSURI} {
|
|
if _, err := validateOIDCURL(ctx, value, s.allowPrivateIdentityProvider); err != nil {
|
|
return oidcDiscovery{}, err
|
|
}
|
|
}
|
|
return oidcDiscovery{Issuer: raw.Issuer, AuthorizationEndpoint: raw.AuthorizationEndpoint, TokenEndpoint: raw.TokenEndpoint, JWKSURI: raw.JWKSURI}, nil
|
|
}
|
|
|
|
func (s *Service) oidcCredentials(p OIDCProvider) (oidcCredentials, error) {
|
|
plaintext, err := s.idpCipher.Decrypt(p.EncryptedCredentials, p.CredentialKEKVersion)
|
|
var c oidcCredentials
|
|
if err == nil {
|
|
err = json.Unmarshal(plaintext, &c)
|
|
}
|
|
return c, err
|
|
}
|
|
|
|
func (s *Service) oidcHTTPClient() *http.Client {
|
|
if s.oidcClient == nil {
|
|
s.oidcClient = newOIDCHTTPClient(s.allowPrivateIdentityProvider)
|
|
}
|
|
return s.oidcClient
|
|
}
|
|
|
|
func newOIDCHTTPClient(allowPrivate bool) *http.Client {
|
|
return &http.Client{
|
|
Timeout: 10 * time.Second,
|
|
Transport: &http.Transport{
|
|
DialContext: safeOIDCDial(allowPrivate),
|
|
TLSHandshakeTimeout: 5 * time.Second,
|
|
ResponseHeaderTimeout: 8 * time.Second,
|
|
MaxIdleConns: 32,
|
|
MaxIdleConnsPerHost: 8,
|
|
IdleConnTimeout: 90 * time.Second,
|
|
},
|
|
CheckRedirect: func(*http.Request, []*http.Request) error { return errors.New("redirect rejected") },
|
|
}
|
|
}
|
|
|
|
// jwksCacheTTL 是 JWKS 缓存有效期,以 IdP 最短键轮换周期为界。
|
|
const jwksCacheTTL = 5 * time.Minute
|
|
|
|
type jwksKey struct{ Kid, Kty, N, E string }
|
|
|
|
type jwksDocument struct {
|
|
Keys []jwksKey
|
|
}
|
|
|
|
type jwksCacheEntry struct {
|
|
doc jwksDocument
|
|
fetched time.Time
|
|
}
|
|
|
|
// jwksFor 返回 IdP 的 JWKS,带数分钟缓存(provider 数量有限,map 无需淘汰)。
|
|
func (s *Service) jwksFor(ctx context.Context, jwksURI string) (jwksDocument, error) {
|
|
s.jwksMu.Lock()
|
|
if s.jwks == nil {
|
|
s.jwks = make(map[string]jwksCacheEntry)
|
|
}
|
|
if entry, ok := s.jwks[jwksURI]; ok && time.Since(entry.fetched) < jwksCacheTTL {
|
|
s.jwksMu.Unlock()
|
|
return entry.doc, nil
|
|
}
|
|
s.jwksMu.Unlock()
|
|
req, err := http.NewRequestWithContext(ctx, http.MethodGet, jwksURI, nil)
|
|
if err != nil {
|
|
return jwksDocument{}, err
|
|
}
|
|
response, err := s.oidcHTTPClient().Do(req)
|
|
if err != nil {
|
|
return jwksDocument{}, err
|
|
}
|
|
defer response.Body.Close()
|
|
payload, _ := io.ReadAll(io.LimitReader(response.Body, oidcMaxResponse+1))
|
|
var doc jwksDocument
|
|
if response.StatusCode/100 != 2 || len(payload) > oidcMaxResponse || json.Unmarshal(payload, &doc) != nil {
|
|
return jwksDocument{}, errors.New("jwks")
|
|
}
|
|
s.jwksMu.Lock()
|
|
s.jwks[jwksURI] = jwksCacheEntry{doc: doc, fetched: time.Now()}
|
|
s.jwksMu.Unlock()
|
|
return doc, nil
|
|
}
|
|
|
|
func (s *Service) verifyIDToken(ctx context.Context, d oidcDiscovery, clientID, nonce, token string) (oidcClaims, error) {
|
|
parts := strings.Split(token, ".")
|
|
if len(parts) != 3 {
|
|
return oidcClaims{}, errors.New("jwt format")
|
|
}
|
|
headerBytes, err := base64.RawURLEncoding.DecodeString(parts[0])
|
|
if err != nil {
|
|
return oidcClaims{}, err
|
|
}
|
|
var header struct{ Alg, Kid string }
|
|
if json.Unmarshal(headerBytes, &header) != nil || header.Alg != "RS256" || header.Kid == "" {
|
|
return oidcClaims{}, errors.New("jwt header")
|
|
}
|
|
// JWKS 按 provider 缓存数分钟:每次登录都拉 discovery + JWKS 是 2-3 个
|
|
// 同步往返;缓存以 IdP 最短键轮换周期为界(默认 5 分钟)。
|
|
keys, err := s.jwksFor(ctx, d.JWKSURI)
|
|
if err != nil {
|
|
return oidcClaims{}, err
|
|
}
|
|
var key *rsa.PublicKey
|
|
for _, j := range keys.Keys {
|
|
if j.Kid == header.Kid && j.Kty == "RSA" {
|
|
nBytes, nErr := base64.RawURLEncoding.DecodeString(j.N)
|
|
eBytes, eErr := base64.RawURLEncoding.DecodeString(j.E)
|
|
e := 0
|
|
for _, b := range eBytes {
|
|
e = e<<8 + int(b)
|
|
}
|
|
if nErr == nil && eErr == nil && len(nBytes) >= 256 && e >= 3 && e%2 == 1 {
|
|
key = &rsa.PublicKey{N: new(big.Int).SetBytes(nBytes), E: e}
|
|
}
|
|
}
|
|
}
|
|
if key == nil {
|
|
return oidcClaims{}, errors.New("key")
|
|
}
|
|
signature, err := base64.RawURLEncoding.DecodeString(parts[2])
|
|
if err != nil {
|
|
return oidcClaims{}, err
|
|
}
|
|
digest := sha256.Sum256([]byte(parts[0] + "." + parts[1]))
|
|
if rsa.VerifyPKCS1v15(key, crypto.SHA256, digest[:], signature) != nil {
|
|
return oidcClaims{}, errors.New("signature")
|
|
}
|
|
claimsBytes, err := base64.RawURLEncoding.DecodeString(parts[1])
|
|
if err != nil {
|
|
return oidcClaims{}, err
|
|
}
|
|
var claims oidcClaims
|
|
now := time.Now().Unix()
|
|
if s.now != nil {
|
|
now = s.now().Unix()
|
|
}
|
|
if json.Unmarshal(claimsBytes, &claims) != nil || claims.Issuer != d.Issuer || claims.Subject == "" || claims.Nonce != nonce ||
|
|
claims.ExpiresAt <= now-30 || claims.IssuedAt == 0 || claims.IssuedAt > now+30 || claims.NotBefore > now+30 ||
|
|
claims.ExpiresAt-claims.IssuedAt > oidcMaxTokenLifetimeSeconds {
|
|
return oidcClaims{}, errors.New("claims")
|
|
}
|
|
audiences, ok := parseAudience(claims.Audience)
|
|
if !ok || !contains(audiences, clientID) || (len(audiences) > 1 && claims.AuthorizedParty != clientID) ||
|
|
(claims.AuthorizedParty != "" && claims.AuthorizedParty != clientID) {
|
|
return oidcClaims{}, errors.New("audience")
|
|
}
|
|
return claims, nil
|
|
}
|
|
|
|
func (r *Repository) ListOIDCProviders(ctx context.Context, enabledOnly bool) ([]OIDCProvider, error) {
|
|
query := `SELECT id::text,code,display_name,issuer_url,client_id,encrypted_credentials,credential_kek_version,redirect_uri,portal_return_url,scopes,auto_provision,default_department_id::text,enabled,revision,created_at,updated_at FROM gateway.identity_providers WHERE kind='oidc'`
|
|
if enabledOnly {
|
|
query += " AND enabled"
|
|
}
|
|
query += " ORDER BY code"
|
|
rows, err := r.pool.Query(ctx, query)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("%w: %v", ErrUnavailable, err)
|
|
}
|
|
defer rows.Close()
|
|
items := []OIDCProvider{}
|
|
for rows.Next() {
|
|
var p OIDCProvider
|
|
if err := rows.Scan(&p.ID, &p.Code, &p.DisplayName, &p.IssuerURL, &p.ClientID, &p.EncryptedCredentials, &p.CredentialKEKVersion, &p.RedirectURI, &p.PortalReturnURL, &p.Scopes, &p.AutoProvision, &p.DefaultDepartmentID, &p.Enabled, &p.Revision, &p.CreatedAt, &p.UpdatedAt); err != nil {
|
|
return nil, err
|
|
}
|
|
items = append(items, p)
|
|
}
|
|
return items, mapRepositoryError(rows.Err())
|
|
}
|
|
func (r *Repository) GetOIDCProviderByCode(ctx context.Context, code string) (OIDCProvider, error) {
|
|
var p OIDCProvider
|
|
err := r.pool.QueryRow(ctx, `SELECT id::text,code,display_name,issuer_url,client_id,encrypted_credentials,credential_kek_version,redirect_uri,portal_return_url,scopes,auto_provision,default_department_id::text,enabled,revision,created_at,updated_at FROM gateway.identity_providers WHERE code=$1 AND kind='oidc'`, strings.ToLower(code)).Scan(&p.ID, &p.Code, &p.DisplayName, &p.IssuerURL, &p.ClientID, &p.EncryptedCredentials, &p.CredentialKEKVersion, &p.RedirectURI, &p.PortalReturnURL, &p.Scopes, &p.AutoProvision, &p.DefaultDepartmentID, &p.Enabled, &p.Revision, &p.CreatedAt, &p.UpdatedAt)
|
|
return p, mapRepositoryError(err)
|
|
}
|
|
|
|
func (r *Repository) CreateOIDCProvider(ctx context.Context, p OIDCProvider, actor string) (OIDCProvider, error) {
|
|
id, _ := platformid.NewUUID()
|
|
p.ID = id
|
|
return r.storeOIDCProvider(ctx, p, actor, true, true)
|
|
}
|
|
func (r *Repository) UpdateOIDCProvider(ctx context.Context, p OIDCProvider, actor string, replace bool) (OIDCProvider, error) {
|
|
return r.storeOIDCProvider(ctx, p, actor, false, replace)
|
|
}
|
|
func (r *Repository) storeOIDCProvider(ctx context.Context, p OIDCProvider, actor string, creating, replace bool) (OIDCProvider, error) {
|
|
tx, err := r.pool.Begin(ctx)
|
|
if err != nil {
|
|
return p, ErrUnavailable
|
|
}
|
|
defer tx.Rollback(ctx)
|
|
if creating {
|
|
err = tx.QueryRow(ctx, `INSERT INTO gateway.identity_providers(id,code,kind,display_name,issuer_url,client_id,encrypted_credentials,credential_kek_version,redirect_uri,portal_return_url,scopes,auto_provision,default_department_id,enabled) VALUES($1,$2,'oidc',$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13) RETURNING revision,created_at,updated_at`, p.ID, p.Code, p.DisplayName, p.IssuerURL, p.ClientID, p.EncryptedCredentials, p.CredentialKEKVersion, p.RedirectURI, p.PortalReturnURL, p.Scopes, p.AutoProvision, p.DefaultDepartmentID, p.Enabled).Scan(&p.Revision, &p.CreatedAt, &p.UpdatedAt)
|
|
} else {
|
|
err = tx.QueryRow(ctx, `UPDATE gateway.identity_providers SET code=$2,display_name=$3,issuer_url=$4,client_id=$5,encrypted_credentials=CASE WHEN $14 THEN $6 ELSE encrypted_credentials END,credential_kek_version=CASE WHEN $14 THEN $7 ELSE credential_kek_version END,redirect_uri=$8,portal_return_url=$9,scopes=$10,auto_provision=$11,default_department_id=$12,enabled=$13,revision=revision+1,updated_at=clock_timestamp() WHERE id=$1 AND kind='oidc' RETURNING encrypted_credentials,credential_kek_version,revision,created_at,updated_at`, p.ID, p.Code, p.DisplayName, p.IssuerURL, p.ClientID, p.EncryptedCredentials, p.CredentialKEKVersion, p.RedirectURI, p.PortalReturnURL, p.Scopes, p.AutoProvision, p.DefaultDepartmentID, p.Enabled, replace).Scan(&p.EncryptedCredentials, &p.CredentialKEKVersion, &p.Revision, &p.CreatedAt, &p.UpdatedAt)
|
|
}
|
|
if err != nil {
|
|
return p, mapManagementError(err)
|
|
}
|
|
eventID, _ := platformid.NewUUID()
|
|
eventType := "identity_provider.updated"
|
|
if creating {
|
|
eventType = "identity_provider.created"
|
|
}
|
|
payload, _ := json.Marshal(map[string]any{"identity_provider_id": p.ID, "actor_id": actor})
|
|
if _, err = tx.Exec(ctx, `INSERT INTO gateway.outbox_events(event_id,event_type,event_version,aggregate_type,aggregate_id,payload) VALUES($1,$2,1,'identity_provider',$3,$4)`, eventID, eventType, p.ID, payload); err != nil {
|
|
return p, ErrUnavailable
|
|
}
|
|
if tx.Commit(ctx) != nil {
|
|
return p, ErrUnavailable
|
|
}
|
|
return p, nil
|
|
}
|
|
|
|
func (r *Repository) ResolveExternalAccount(ctx context.Context, p OIDCProvider, c oidcClaims) (Account, error) {
|
|
return r.resolveExternalAccount(ctx, externalProvider{
|
|
ID: p.ID, Code: p.Code, AuthSource: "oidc", AutoProvision: p.AutoProvision, DefaultDepartmentID: p.DefaultDepartmentID,
|
|
}, externalClaims{Subject: c.Subject, Email: c.Email, PreferredUsername: c.PreferredUsername, Name: c.Name})
|
|
}
|
|
|
|
func (r *Repository) resolveExternalAccount(ctx context.Context, p externalProvider, c externalClaims) (Account, error) {
|
|
var id string
|
|
err := r.pool.QueryRow(ctx, `SELECT id::text FROM gateway.portal_users WHERE identity_provider_id=$1 AND external_subject=$2`, p.ID, c.Subject).Scan(&id)
|
|
if err == nil {
|
|
return r.FindPortalByID(ctx, id)
|
|
}
|
|
if !errors.Is(err, pgx.ErrNoRows) {
|
|
return Account{}, ErrUnavailable
|
|
}
|
|
if !p.AutoProvision {
|
|
return Account{}, ErrNotFound
|
|
}
|
|
if p.DefaultDepartmentID != nil {
|
|
var active bool
|
|
if err := r.pool.QueryRow(ctx, `SELECT active FROM gateway.departments WHERE id=$1`, *p.DefaultDepartmentID).Scan(&active); err != nil || !active {
|
|
return Account{}, ErrUnavailable
|
|
}
|
|
}
|
|
login := strings.ToLower(strings.TrimSpace(c.Email))
|
|
if login == "" {
|
|
login = strings.ToLower(strings.TrimSpace(c.PreferredUsername))
|
|
}
|
|
if login == "" {
|
|
sum := sha256.Sum256([]byte(c.Subject))
|
|
login = p.Code + "_" + fmt.Sprintf("%x", sum[:6])
|
|
}
|
|
login = truncate(login, 128)
|
|
id, _ = platformid.NewUUID()
|
|
eventID, _ := platformid.NewUUID()
|
|
tx, err := r.pool.Begin(ctx)
|
|
if err != nil {
|
|
return Account{}, ErrUnavailable
|
|
}
|
|
defer tx.Rollback(ctx)
|
|
name := truncate(firstNonEmpty(c.Name, c.PreferredUsername, login), 64)
|
|
insert := func(candidate string) (bool, error) {
|
|
var insertedID string
|
|
err := tx.QueryRow(ctx, `INSERT INTO gateway.portal_users(id,account,name,role,permissions,password_hash,auth_source,external_subject,identity_provider_id,department_id,active) VALUES($1,$2,$3,'member','{}',NULL,$4,$5,$6,$7,true) ON CONFLICT DO NOTHING RETURNING id::text`, id, candidate, name, p.AuthSource, c.Subject, p.ID, p.DefaultDepartmentID).Scan(&insertedID)
|
|
if errors.Is(err, pgx.ErrNoRows) {
|
|
return false, nil
|
|
}
|
|
return err == nil, err
|
|
}
|
|
inserted, err := insert(login)
|
|
if err != nil {
|
|
return Account{}, ErrUnavailable
|
|
}
|
|
if !inserted {
|
|
var existingID string
|
|
lookupErr := tx.QueryRow(ctx, `SELECT id::text FROM gateway.portal_users WHERE identity_provider_id=$1 AND external_subject=$2`, p.ID, c.Subject).Scan(&existingID)
|
|
if lookupErr == nil {
|
|
_ = tx.Rollback(ctx)
|
|
return r.FindPortalByID(ctx, existingID)
|
|
}
|
|
if !errors.Is(lookupErr, pgx.ErrNoRows) {
|
|
return Account{}, ErrUnavailable
|
|
}
|
|
sum := sha256.Sum256([]byte(p.Code + "|" + c.Subject))
|
|
login = truncate(login, 110) + "-" + fmt.Sprintf("%x", sum[:8])
|
|
inserted, err = insert(login)
|
|
}
|
|
if err != nil || !inserted {
|
|
return Account{}, ErrUnavailable
|
|
}
|
|
payload, _ := json.Marshal(map[string]any{"identity_id": id, "identity_provider_id": p.ID})
|
|
if _, err = tx.Exec(ctx, `INSERT INTO gateway.outbox_events(event_id,event_type,event_version,aggregate_type,aggregate_id,payload) VALUES($1,'identity.external_provisioned',1,'identity',$2,$3)`, eventID, id, payload); err != nil {
|
|
return Account{}, ErrUnavailable
|
|
}
|
|
if tx.Commit(ctx) != nil {
|
|
return Account{}, ErrUnavailable
|
|
}
|
|
return r.FindPortalByID(ctx, id)
|
|
}
|
|
|
|
func validateAbsoluteURL(raw string) (string, error) {
|
|
u, err := url.Parse(strings.TrimSpace(raw))
|
|
if err != nil || (u.Scheme != "http" && u.Scheme != "https") || u.Hostname() == "" || u.User != nil {
|
|
return "", errors.New("invalid url")
|
|
}
|
|
return u.String(), nil
|
|
}
|
|
func validateOIDCURL(ctx context.Context, raw string, allowPrivate bool) (string, error) {
|
|
value, err := validateAbsoluteURL(raw)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
u, _ := url.Parse(value)
|
|
if u.RawQuery != "" || u.Fragment != "" || (!allowPrivate && u.Scheme != "https") {
|
|
return "", errors.New("invalid issuer url")
|
|
}
|
|
if !allowPrivate {
|
|
addresses, err := net.DefaultResolver.LookupIPAddr(ctx, u.Hostname())
|
|
if err != nil || len(addresses) == 0 {
|
|
return "", errors.New("resolve")
|
|
}
|
|
for _, a := range addresses {
|
|
if !publicIP(a.IP) {
|
|
return "", errors.New("blocked address")
|
|
}
|
|
}
|
|
}
|
|
u.Path = strings.TrimRight(u.Path, "/")
|
|
return u.String(), nil
|
|
}
|
|
func safeOIDCDial(allowPrivate bool) func(context.Context, string, string) (net.Conn, error) {
|
|
d := &net.Dialer{Timeout: 5 * time.Second, KeepAlive: 30 * time.Second}
|
|
if allowPrivate {
|
|
return d.DialContext
|
|
}
|
|
return func(ctx context.Context, network, address string) (net.Conn, error) {
|
|
host, port, err := net.SplitHostPort(address)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
addresses, err := net.DefaultResolver.LookupIPAddr(ctx, host)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
for _, a := range addresses {
|
|
if !publicIP(a.IP) {
|
|
return nil, errors.New("blocked address")
|
|
}
|
|
}
|
|
if len(addresses) == 0 {
|
|
return nil, errors.New("resolve")
|
|
}
|
|
return d.DialContext(ctx, network, net.JoinHostPort(addresses[0].IP.String(), port))
|
|
}
|
|
}
|
|
func publicIP(ip net.IP) bool {
|
|
return ip != nil && !ip.IsPrivate() && !ip.IsLoopback() && !ip.IsLinkLocalUnicast() && !ip.IsLinkLocalMulticast() && !ip.IsMulticast() && !ip.IsUnspecified()
|
|
}
|
|
func randomURLToken(size int) (string, error) {
|
|
b := make([]byte, size)
|
|
_, err := rand.Read(b)
|
|
return base64.RawURLEncoding.EncodeToString(b), err
|
|
}
|
|
func contains(values []string, target string) bool {
|
|
for _, v := range values {
|
|
if v == target {
|
|
return true
|
|
}
|
|
}
|
|
return false
|
|
}
|
|
func audienceContains(raw json.RawMessage, target string) bool {
|
|
values, ok := parseAudience(raw)
|
|
return ok && contains(values, target)
|
|
}
|
|
func parseAudience(raw json.RawMessage) ([]string, bool) {
|
|
var one string
|
|
if json.Unmarshal(raw, &one) == nil {
|
|
return []string{one}, one != ""
|
|
}
|
|
var many []string
|
|
if json.Unmarshal(raw, &many) != nil || len(many) == 0 {
|
|
return nil, false
|
|
}
|
|
for _, value := range many {
|
|
if value == "" {
|
|
return nil, false
|
|
}
|
|
}
|
|
return many, true
|
|
}
|
|
func normalizeOIDCScopes(values []string) ([]string, error) {
|
|
if len(values) > 16 {
|
|
return nil, errors.New("too many scopes")
|
|
}
|
|
seen := make(map[string]struct{}, len(values))
|
|
result := make([]string, 0, len(values))
|
|
for _, value := range values {
|
|
value = strings.TrimSpace(value)
|
|
if value == "" || len(value) > 64 || strings.ContainsAny(value, " \t\r\n\"") {
|
|
return nil, errors.New("invalid scope")
|
|
}
|
|
if _, exists := seen[value]; exists {
|
|
continue
|
|
}
|
|
seen[value] = struct{}{}
|
|
result = append(result, value)
|
|
}
|
|
return result, nil
|
|
}
|
|
func firstNonEmpty(values ...string) string {
|
|
for _, v := range values {
|
|
if strings.TrimSpace(v) != "" {
|
|
return strings.TrimSpace(v)
|
|
}
|
|
}
|
|
return "用户"
|
|
}
|
|
func truncate(value string, max int) string {
|
|
runes := []rune(value)
|
|
if len(runes) > max {
|
|
return string(runes[:max])
|
|
}
|
|
return value
|
|
}
|