AI Gateway Go 0.10.0 源码快照 + 旗舰版需求规划报告

M0-M7 已完成:核心网关(身份/RBAC/TOTP/OIDC/SAML/Provider/配额/路由/内容策略/审计/定价)+ 资源市场(MCP/Skills/数字员工)。
含 22 个 PostgreSQL 迁移、管理端/门户端前端源码、OpenAPI 契约、部署 compose。

Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
ben
2026-08-12 11:45:54 +08:00
commit 5759c1862e
807 changed files with 114727 additions and 0 deletions
+818
View File
@@ -0,0 +1,818 @@
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
}