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") }, } } 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") } req, err := http.NewRequestWithContext(ctx, http.MethodGet, d.JWKSURI, nil) if err != nil { return oidcClaims{}, err } response, err := s.oidcHTTPClient().Do(req) if err != nil { return oidcClaims{}, err } defer response.Body.Close() payload, _ := io.ReadAll(io.LimitReader(response.Body, oidcMaxResponse+1)) var keys struct { Keys []struct{ Kid, Kty, N, E string } } if response.StatusCode/100 != 2 || len(payload) > oidcMaxResponse || json.Unmarshal(payload, &keys) != nil { return oidcClaims{}, errors.New("jwks") } 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 }