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:
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user