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,636 @@
|
||||
package identity
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/x509"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"encoding/xml"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"aigateway.local/core/internal/platform/apiresponse"
|
||||
platformid "aigateway.local/core/internal/platform/id"
|
||||
"github.com/beevik/etree"
|
||||
"github.com/crewjam/saml"
|
||||
"github.com/crewjam/saml/samlsp"
|
||||
dsig "github.com/russellhaering/goxmldsig"
|
||||
)
|
||||
|
||||
const samlMaxResponse = 4 << 20
|
||||
|
||||
type SAMLConfig struct {
|
||||
MetadataURL string `json:"metadata_url"`
|
||||
SPEntityID string `json:"sp_entity_id"`
|
||||
ACSURL string `json:"acs_url"`
|
||||
EmailAttribute string `json:"email_attribute"`
|
||||
NameAttribute string `json:"name_attribute"`
|
||||
}
|
||||
|
||||
type SAMLProvider struct {
|
||||
ID string
|
||||
Code string
|
||||
DisplayName string
|
||||
PortalReturnURL string
|
||||
AutoProvision bool
|
||||
DefaultDepartmentID *string
|
||||
Enabled bool
|
||||
Revision int64
|
||||
Config SAMLConfig
|
||||
CreatedAt time.Time
|
||||
UpdatedAt time.Time
|
||||
}
|
||||
|
||||
type samlProviderInput struct {
|
||||
Code string `json:"code"`
|
||||
DisplayName string `json:"display_name"`
|
||||
MetadataURL string `json:"metadata_url"`
|
||||
SPEntityID string `json:"sp_entity_id"`
|
||||
ACSURL string `json:"acs_url"`
|
||||
PortalReturnURL string `json:"portal_return_url"`
|
||||
EmailAttribute string `json:"email_attribute"`
|
||||
NameAttribute string `json:"name_attribute"`
|
||||
AutoProvision bool `json:"auto_provision"`
|
||||
DefaultDepartmentID *string `json:"default_department_id"`
|
||||
Enabled bool `json:"enabled"`
|
||||
}
|
||||
|
||||
type samlChallenge struct {
|
||||
ProviderID string `json:"provider_id"`
|
||||
RequestID string `json:"request_id"`
|
||||
}
|
||||
|
||||
type samlMetadataCacheEntry struct {
|
||||
Metadata *saml.EntityDescriptor
|
||||
ExpiresAt time.Time
|
||||
}
|
||||
|
||||
type PublicIdentityProvider struct {
|
||||
Code string `json:"code"`
|
||||
DisplayName string `json:"display_name"`
|
||||
Kind string `json:"kind"`
|
||||
}
|
||||
|
||||
func (h *ManagementHTTPHandler) registerSAML() {
|
||||
h.mux.HandleFunc("GET /api/v1/admin/saml-providers", h.listSAMLProviders)
|
||||
h.mux.HandleFunc("POST /api/v1/admin/saml-providers", h.createSAMLProvider)
|
||||
h.mux.HandleFunc("PUT /api/v1/admin/saml-providers/{provider_id}", h.updateSAMLProvider)
|
||||
}
|
||||
|
||||
func (h *ManagementHTTPHandler) listSAMLProviders(w http.ResponseWriter, r *http.Request) {
|
||||
if _, ok := h.requirePermission(w, r); !ok {
|
||||
return
|
||||
}
|
||||
providers, err := h.service.repository.ListSAMLProviders(r.Context())
|
||||
if err != nil {
|
||||
h.writeError(w, err)
|
||||
return
|
||||
}
|
||||
items := make([]map[string]any, 0, len(providers))
|
||||
for _, provider := range providers {
|
||||
items = append(items, samlView(provider))
|
||||
}
|
||||
apiresponse.OK(w, items)
|
||||
}
|
||||
|
||||
func (h *ManagementHTTPHandler) createSAMLProvider(w http.ResponseWriter, r *http.Request) {
|
||||
actor, ok := h.requirePermission(w, r)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
provider, ok := h.decodeSAMLProvider(w, r)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
created, err := h.service.repository.CreateSAMLProvider(r.Context(), provider, actor.ID)
|
||||
if err != nil {
|
||||
h.writeError(w, err)
|
||||
return
|
||||
}
|
||||
apiresponse.OK(w, samlView(created))
|
||||
}
|
||||
|
||||
func (h *ManagementHTTPHandler) updateSAMLProvider(w http.ResponseWriter, r *http.Request) {
|
||||
actor, ok := h.requirePermission(w, r)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
provider, ok := h.decodeSAMLProvider(w, r)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
provider.ID = r.PathValue("provider_id")
|
||||
updated, err := h.service.repository.UpdateSAMLProvider(r.Context(), provider, actor.ID)
|
||||
if err != nil {
|
||||
h.writeError(w, err)
|
||||
return
|
||||
}
|
||||
apiresponse.OK(w, samlView(updated))
|
||||
}
|
||||
|
||||
func (h *ManagementHTTPHandler) decodeSAMLProvider(w http.ResponseWriter, r *http.Request) (SAMLProvider, bool) {
|
||||
var input samlProviderInput
|
||||
decoder := json.NewDecoder(http.MaxBytesReader(w, r.Body, 1<<20))
|
||||
decoder.DisallowUnknownFields()
|
||||
if decoder.Decode(&input) != nil {
|
||||
apiresponse.Error(w, http.StatusBadRequest, "请求格式无效")
|
||||
return SAMLProvider{}, false
|
||||
}
|
||||
input.Code = strings.ToLower(strings.TrimSpace(input.Code))
|
||||
input.DisplayName = strings.TrimSpace(input.DisplayName)
|
||||
input.EmailAttribute = strings.TrimSpace(input.EmailAttribute)
|
||||
input.NameAttribute = strings.TrimSpace(input.NameAttribute)
|
||||
if !oidcCodePattern.MatchString(input.Code) || input.DisplayName == "" || len(input.DisplayName) > 128 {
|
||||
apiresponse.Error(w, http.StatusBadRequest, "身份源代码或名称无效")
|
||||
return SAMLProvider{}, false
|
||||
}
|
||||
metadataURL, err := validateOIDCURL(r.Context(), input.MetadataURL, h.service.allowPrivateIdentityProvider)
|
||||
if err != nil {
|
||||
apiresponse.Error(w, http.StatusBadRequest, "SAML metadata URL 无效")
|
||||
return SAMLProvider{}, false
|
||||
}
|
||||
entityID, err := validateSAMLEntityID(input.SPEntityID)
|
||||
if err != nil {
|
||||
apiresponse.Error(w, http.StatusBadRequest, "SAML SP Entity ID 无效")
|
||||
return SAMLProvider{}, false
|
||||
}
|
||||
acsURL, err := validateAbsoluteURL(input.ACSURL)
|
||||
acs, _ := url.Parse(acsURL)
|
||||
expectedACSSuffix := "/api/v1/portal/sso/" + input.Code + "/callback"
|
||||
if err != nil || acs.Fragment != "" || !strings.HasSuffix(strings.TrimRight(acs.Path, "/"), expectedACSSuffix) {
|
||||
apiresponse.Error(w, http.StatusBadRequest, "SAML ACS URL 无效")
|
||||
return SAMLProvider{}, false
|
||||
}
|
||||
returnURL, err := validateAbsoluteURL(input.PortalReturnURL)
|
||||
if err != nil {
|
||||
apiresponse.Error(w, http.StatusBadRequest, "门户返回 URL 无效")
|
||||
return SAMLProvider{}, false
|
||||
}
|
||||
if input.EmailAttribute == "" {
|
||||
input.EmailAttribute = "email"
|
||||
}
|
||||
if input.NameAttribute == "" {
|
||||
input.NameAttribute = "displayName"
|
||||
}
|
||||
if len(input.EmailAttribute) > 256 || len(input.NameAttribute) > 256 {
|
||||
apiresponse.Error(w, http.StatusBadRequest, "SAML 属性名过长")
|
||||
return SAMLProvider{}, 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, http.StatusBadRequest, "默认部门不存在或已停用")
|
||||
return SAMLProvider{}, false
|
||||
}
|
||||
}
|
||||
provider := SAMLProvider{
|
||||
Code: input.Code, DisplayName: input.DisplayName, PortalReturnURL: returnURL,
|
||||
AutoProvision: input.AutoProvision, DefaultDepartmentID: input.DefaultDepartmentID, Enabled: input.Enabled,
|
||||
Config: SAMLConfig{MetadataURL: metadataURL, SPEntityID: entityID, ACSURL: acsURL, EmailAttribute: input.EmailAttribute, NameAttribute: input.NameAttribute},
|
||||
}
|
||||
if input.Enabled {
|
||||
if _, err := h.service.samlServiceProvider(r.Context(), provider); err != nil {
|
||||
apiresponse.Error(w, http.StatusBadRequest, "SAML metadata 无法验证或缺少有效签名证书")
|
||||
return SAMLProvider{}, false
|
||||
}
|
||||
}
|
||||
return provider, true
|
||||
}
|
||||
|
||||
func samlView(provider SAMLProvider) map[string]any {
|
||||
return map[string]any{
|
||||
"id": provider.ID, "kind": "saml", "code": provider.Code, "display_name": provider.DisplayName,
|
||||
"metadata_url": provider.Config.MetadataURL, "sp_entity_id": provider.Config.SPEntityID,
|
||||
"acs_url": provider.Config.ACSURL, "portal_return_url": provider.PortalReturnURL,
|
||||
"email_attribute": provider.Config.EmailAttribute, "name_attribute": provider.Config.NameAttribute,
|
||||
"auto_provision": provider.AutoProvision, "default_department_id": provider.DefaultDepartmentID,
|
||||
"enabled": provider.Enabled, "revision": provider.Revision,
|
||||
}
|
||||
}
|
||||
|
||||
func (h *HTTPHandler) startSSO(w http.ResponseWriter, r *http.Request) {
|
||||
kind, err := h.service.repository.GetIdentityProviderKind(r.Context(), r.PathValue("provider_code"))
|
||||
if err != nil {
|
||||
http.NotFound(w, r)
|
||||
return
|
||||
}
|
||||
if kind == "saml" {
|
||||
h.startSAML(w, r)
|
||||
return
|
||||
}
|
||||
h.startOIDC(w, r)
|
||||
}
|
||||
|
||||
func (h *HTTPHandler) startSAML(w http.ResponseWriter, r *http.Request) {
|
||||
provider, err := h.service.repository.GetSAMLProviderByCode(r.Context(), r.PathValue("provider_code"))
|
||||
if err != nil || !provider.Enabled {
|
||||
http.NotFound(w, r)
|
||||
return
|
||||
}
|
||||
sp, err := h.service.samlServiceProvider(r.Context(), provider)
|
||||
if err != nil {
|
||||
apiresponse.Error(w, http.StatusBadGateway, "SAML metadata 获取失败")
|
||||
return
|
||||
}
|
||||
idpURL := sp.GetSSOBindingLocation(saml.HTTPRedirectBinding)
|
||||
if idpURL == "" {
|
||||
apiresponse.Error(w, http.StatusBadGateway, "SAML 身份源不支持 Redirect 登录")
|
||||
return
|
||||
}
|
||||
authnRequest, err := sp.MakeAuthenticationRequest(idpURL, saml.HTTPRedirectBinding, saml.HTTPPostBinding)
|
||||
if err != nil {
|
||||
apiresponse.Error(w, http.StatusBadGateway, "SAML 登录请求生成失败")
|
||||
return
|
||||
}
|
||||
relayState, err := h.service.sessions.StoreOneTime(r.Context(), "saml-state", samlChallenge{ProviderID: provider.ID, RequestID: authnRequest.ID}, 5*time.Minute)
|
||||
if err != nil {
|
||||
h.writeIdentityError(w, err)
|
||||
return
|
||||
}
|
||||
target, err := authnRequest.Redirect(relayState, sp)
|
||||
if err != nil {
|
||||
apiresponse.Error(w, http.StatusBadGateway, "SAML 登录请求生成失败")
|
||||
return
|
||||
}
|
||||
http.Redirect(w, r, target.String(), http.StatusFound)
|
||||
}
|
||||
|
||||
func (h *HTTPHandler) callbackSAML(w http.ResponseWriter, r *http.Request) {
|
||||
r.Body = http.MaxBytesReader(w, r.Body, samlMaxResponse)
|
||||
if err := r.ParseForm(); err != nil || r.PostForm.Get("SAMLResponse") == "" || r.PostForm.Get("RelayState") == "" || r.PostForm.Get("SAMLart") != "" {
|
||||
apiresponse.Error(w, http.StatusBadRequest, "SAML 回调格式无效")
|
||||
return
|
||||
}
|
||||
var challenge samlChallenge
|
||||
if err := h.service.sessions.ConsumeOneTime(r.Context(), "saml-state", r.PostForm.Get("RelayState"), &challenge); err != nil {
|
||||
apiresponse.Error(w, http.StatusUnauthorized, "SAML RelayState 无效或已使用")
|
||||
return
|
||||
}
|
||||
provider, err := h.service.repository.GetSAMLProviderByCode(r.Context(), r.PathValue("provider_code"))
|
||||
if err != nil || !provider.Enabled || provider.ID != challenge.ProviderID {
|
||||
apiresponse.Error(w, http.StatusUnauthorized, "SAML 身份源无效")
|
||||
return
|
||||
}
|
||||
sp, err := h.service.samlServiceProvider(r.Context(), provider)
|
||||
if err != nil {
|
||||
apiresponse.Error(w, http.StatusBadGateway, "SAML metadata 获取失败")
|
||||
return
|
||||
}
|
||||
assertion, err := sp.ParseResponse(r, []string{challenge.RequestID})
|
||||
if err != nil || assertion == nil || assertion.Subject == nil || assertion.Subject.NameID == nil || strings.TrimSpace(assertion.Subject.NameID.Value) == "" {
|
||||
apiresponse.Error(w, http.StatusUnauthorized, "SAML 断言校验失败")
|
||||
return
|
||||
}
|
||||
claimed, err := h.service.sessions.ClaimIdentifier(r.Context(), "saml-assertion", assertion.ID, 10*time.Minute)
|
||||
if err != nil {
|
||||
h.writeIdentityError(w, err)
|
||||
return
|
||||
}
|
||||
if !claimed {
|
||||
apiresponse.Error(w, http.StatusUnauthorized, "SAML 断言已使用")
|
||||
return
|
||||
}
|
||||
subject := strings.TrimSpace(assertion.Subject.NameID.Value)
|
||||
email := samlAttribute(assertion, provider.Config.EmailAttribute)
|
||||
name := samlAttribute(assertion, provider.Config.NameAttribute)
|
||||
account, err := h.service.repository.resolveExternalAccount(r.Context(), externalProvider{
|
||||
ID: provider.ID, Code: provider.Code, AuthSource: "saml", AutoProvision: provider.AutoProvision, DefaultDepartmentID: provider.DefaultDepartmentID,
|
||||
}, externalClaims{Subject: subject, Email: email, PreferredUsername: subject, Name: name})
|
||||
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(provider.PortalReturnURL)
|
||||
query := returnURL.Query()
|
||||
query.Set("sso_code", exchange)
|
||||
returnURL.RawQuery = query.Encode()
|
||||
http.Redirect(w, r, returnURL.String(), http.StatusFound)
|
||||
}
|
||||
|
||||
func (h *HTTPHandler) samlMetadata(w http.ResponseWriter, r *http.Request) {
|
||||
provider, err := h.service.repository.GetSAMLProviderByCode(r.Context(), r.PathValue("provider_code"))
|
||||
if err != nil {
|
||||
http.NotFound(w, r)
|
||||
return
|
||||
}
|
||||
sp, err := newSAMLServiceProvider(provider, nil, h.service.oidcHTTPClient())
|
||||
if err != nil {
|
||||
apiresponse.Error(w, http.StatusBadGateway, "SAML metadata 获取失败")
|
||||
return
|
||||
}
|
||||
payload, err := xml.MarshalIndent(sp.Metadata(), "", " ")
|
||||
if err != nil {
|
||||
apiresponse.Error(w, http.StatusInternalServerError, "SAML SP metadata 生成失败")
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/samlmetadata+xml; charset=utf-8")
|
||||
w.Header().Set("X-Content-Type-Options", "nosniff")
|
||||
_, _ = w.Write(append([]byte(xml.Header), payload...))
|
||||
}
|
||||
|
||||
func (s *Service) samlServiceProvider(ctx context.Context, provider SAMLProvider) (*saml.ServiceProvider, error) {
|
||||
metadata, err := s.fetchSAMLMetadata(ctx, provider.Config.MetadataURL)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return newSAMLServiceProvider(provider, metadata, s.oidcHTTPClient())
|
||||
}
|
||||
|
||||
func newSAMLServiceProvider(provider SAMLProvider, metadata *saml.EntityDescriptor, client *http.Client) (*saml.ServiceProvider, error) {
|
||||
acsURL, err := url.Parse(provider.Config.ACSURL)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
metadataURL := *acsURL
|
||||
metadataURL.Path = strings.TrimSuffix(metadataURL.Path, "/callback") + "/metadata"
|
||||
metadataURL.RawQuery = ""
|
||||
metadataURL.Fragment = ""
|
||||
return &saml.ServiceProvider{
|
||||
EntityID: provider.Config.SPEntityID, MetadataURL: metadataURL, AcsURL: *acsURL,
|
||||
IDPMetadata: metadata, HTTPClient: client, AllowIDPInitiated: false,
|
||||
SignatureVerifier: modernSAMLSignatureVerifier{},
|
||||
}, nil
|
||||
}
|
||||
|
||||
type modernSAMLSignatureVerifier struct{}
|
||||
|
||||
func (modernSAMLSignatureVerifier) VerifySignature(validationContext *dsig.ValidationContext, element *etree.Element) error {
|
||||
if err := validateSAMLSignatureAlgorithms(element); err != nil {
|
||||
return err
|
||||
}
|
||||
_, err := validationContext.Validate(element)
|
||||
return err
|
||||
}
|
||||
|
||||
func validateSAMLSignatureAlgorithms(element *etree.Element) error {
|
||||
method := element.FindElement("./Signature/SignedInfo/SignatureMethod")
|
||||
if method == nil {
|
||||
return errors.New("SAML signature method missing")
|
||||
}
|
||||
switch method.SelectAttrValue("Algorithm", "") {
|
||||
case dsig.RSASHA256SignatureMethod, dsig.RSASHA384SignatureMethod, dsig.RSASHA512SignatureMethod,
|
||||
dsig.ECDSASHA256SignatureMethod, dsig.ECDSASHA384SignatureMethod, dsig.ECDSASHA512SignatureMethod:
|
||||
default:
|
||||
return errors.New("legacy or unsupported SAML signature algorithm")
|
||||
}
|
||||
allowedDigests := map[string]bool{
|
||||
"http://www.w3.org/2001/04/xmlenc#sha256": true,
|
||||
"http://www.w3.org/2001/04/xmldsig-more#sha384": true,
|
||||
"http://www.w3.org/2001/04/xmlenc#sha512": true,
|
||||
}
|
||||
digests := element.FindElements("./Signature/SignedInfo/Reference/DigestMethod")
|
||||
if len(digests) == 0 {
|
||||
return errors.New("SAML digest method missing")
|
||||
}
|
||||
for _, digest := range digests {
|
||||
if !allowedDigests[digest.SelectAttrValue("Algorithm", "")] {
|
||||
return errors.New("legacy or unsupported SAML digest algorithm")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Service) fetchSAMLMetadata(ctx context.Context, rawURL string) (*saml.EntityDescriptor, error) {
|
||||
target, err := validateOIDCURL(ctx, rawURL, s.allowPrivateIdentityProvider)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
now := time.Now()
|
||||
s.samlMetadataMu.RLock()
|
||||
cached, found := s.samlMetadata[target]
|
||||
s.samlMetadataMu.RUnlock()
|
||||
if found && now.Before(cached.ExpiresAt) {
|
||||
return cached.Metadata, nil
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, target, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
response, err := s.oidcHTTPClient().Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer response.Body.Close()
|
||||
payload, err := io.ReadAll(io.LimitReader(response.Body, oidcMaxResponse+1))
|
||||
if err != nil || response.StatusCode/100 != 2 || len(payload) > oidcMaxResponse {
|
||||
return nil, errors.New("invalid SAML metadata response")
|
||||
}
|
||||
metadata, err := samlsp.ParseMetadata(payload)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := validateSAMLMetadata(ctx, metadata, s.allowPrivateIdentityProvider, now); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
expiresAt := now.Add(time.Minute)
|
||||
if !metadata.ValidUntil.IsZero() && metadata.ValidUntil.Before(expiresAt) {
|
||||
expiresAt = metadata.ValidUntil
|
||||
}
|
||||
s.samlMetadataMu.Lock()
|
||||
if s.samlMetadata == nil {
|
||||
s.samlMetadata = make(map[string]samlMetadataCacheEntry)
|
||||
}
|
||||
s.samlMetadata[target] = samlMetadataCacheEntry{Metadata: metadata, ExpiresAt: expiresAt}
|
||||
s.samlMetadataMu.Unlock()
|
||||
return metadata, nil
|
||||
}
|
||||
|
||||
func validateSAMLMetadata(ctx context.Context, metadata *saml.EntityDescriptor, allowPrivate bool, now time.Time) error {
|
||||
if metadata == nil || strings.TrimSpace(metadata.EntityID) == "" || len(metadata.IDPSSODescriptors) == 0 {
|
||||
return errors.New("metadata has no IDP descriptor")
|
||||
}
|
||||
if !metadata.ValidUntil.IsZero() && !metadata.ValidUntil.After(now) {
|
||||
return errors.New("metadata expired")
|
||||
}
|
||||
descriptor := metadata.IDPSSODescriptors[0]
|
||||
redirectFound := false
|
||||
for _, endpoint := range descriptor.SingleSignOnServices {
|
||||
if endpoint.Binding == saml.HTTPRedirectBinding {
|
||||
if _, err := validateOIDCURL(ctx, endpoint.Location, allowPrivate); err != nil {
|
||||
return err
|
||||
}
|
||||
redirectFound = true
|
||||
}
|
||||
}
|
||||
if !redirectFound {
|
||||
return errors.New("metadata has no redirect SSO endpoint")
|
||||
}
|
||||
validSigningCertificate := false
|
||||
for _, key := range descriptor.KeyDescriptors {
|
||||
if key.Use != "" && key.Use != "signing" {
|
||||
continue
|
||||
}
|
||||
for _, encoded := range key.KeyInfo.X509Data.X509Certificates {
|
||||
der, err := base64.StdEncoding.DecodeString(strings.Join(strings.Fields(encoded.Data), ""))
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
certificate, err := x509.ParseCertificate(der)
|
||||
if err == nil && !now.Before(certificate.NotBefore) && now.Before(certificate.NotAfter) {
|
||||
validSigningCertificate = true
|
||||
}
|
||||
}
|
||||
}
|
||||
if !validSigningCertificate {
|
||||
return errors.New("metadata has no currently valid signing certificate")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateSAMLEntityID(raw string) (string, error) {
|
||||
value := strings.TrimSpace(raw)
|
||||
if value == "" || len(value) > 512 {
|
||||
return "", errors.New("invalid entity ID")
|
||||
}
|
||||
parsed, err := url.Parse(value)
|
||||
if err != nil || parsed.Scheme == "" || parsed.Fragment != "" {
|
||||
return "", errors.New("invalid entity ID")
|
||||
}
|
||||
if (parsed.Scheme == "http" || parsed.Scheme == "https") && (parsed.Hostname() == "" || parsed.User != nil) {
|
||||
return "", errors.New("invalid entity ID")
|
||||
}
|
||||
return value, nil
|
||||
}
|
||||
|
||||
func samlAttribute(assertion *saml.Assertion, name string) string {
|
||||
for _, statement := range assertion.AttributeStatements {
|
||||
for _, attribute := range statement.Attributes {
|
||||
if attribute.Name != name && attribute.FriendlyName != name {
|
||||
continue
|
||||
}
|
||||
for _, value := range attribute.Values {
|
||||
if strings.TrimSpace(value.Value) != "" {
|
||||
return strings.TrimSpace(value.Value)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (r *Repository) GetIdentityProviderKind(ctx context.Context, code string) (string, error) {
|
||||
var kind string
|
||||
err := r.pool.QueryRow(ctx, `SELECT kind FROM gateway.identity_providers WHERE code=$1 AND enabled`, strings.ToLower(strings.TrimSpace(code))).Scan(&kind)
|
||||
return kind, mapRepositoryError(err)
|
||||
}
|
||||
|
||||
func (r *Repository) ListPublicIdentityProviders(ctx context.Context) ([]PublicIdentityProvider, error) {
|
||||
rows, err := r.pool.Query(ctx, `SELECT code,display_name,kind FROM gateway.identity_providers WHERE enabled ORDER BY code`)
|
||||
if err != nil {
|
||||
return nil, ErrUnavailable
|
||||
}
|
||||
defer rows.Close()
|
||||
providers := []PublicIdentityProvider{}
|
||||
for rows.Next() {
|
||||
var provider PublicIdentityProvider
|
||||
if err := rows.Scan(&provider.Code, &provider.DisplayName, &provider.Kind); err != nil {
|
||||
return nil, ErrUnavailable
|
||||
}
|
||||
providers = append(providers, provider)
|
||||
}
|
||||
return providers, mapRepositoryError(rows.Err())
|
||||
}
|
||||
|
||||
func (r *Repository) ListSAMLProviders(ctx context.Context) ([]SAMLProvider, error) {
|
||||
rows, err := r.pool.Query(ctx, `SELECT id::text,code,display_name,portal_return_url,auto_provision,default_department_id::text,enabled,revision,config,created_at,updated_at FROM gateway.identity_providers WHERE kind='saml' ORDER BY code`)
|
||||
if err != nil {
|
||||
return nil, ErrUnavailable
|
||||
}
|
||||
defer rows.Close()
|
||||
providers := []SAMLProvider{}
|
||||
for rows.Next() {
|
||||
provider, err := scanSAMLProvider(rows)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
providers = append(providers, provider)
|
||||
}
|
||||
return providers, mapRepositoryError(rows.Err())
|
||||
}
|
||||
|
||||
func (r *Repository) GetSAMLProviderByCode(ctx context.Context, code string) (SAMLProvider, error) {
|
||||
return scanSAMLProvider(r.pool.QueryRow(ctx, `SELECT id::text,code,display_name,portal_return_url,auto_provision,default_department_id::text,enabled,revision,config,created_at,updated_at FROM gateway.identity_providers WHERE kind='saml' AND code=$1`, strings.ToLower(strings.TrimSpace(code))))
|
||||
}
|
||||
|
||||
type rowScanner interface {
|
||||
Scan(dest ...any) error
|
||||
}
|
||||
|
||||
func scanSAMLProvider(row rowScanner) (SAMLProvider, error) {
|
||||
var provider SAMLProvider
|
||||
var configJSON []byte
|
||||
if err := row.Scan(&provider.ID, &provider.Code, &provider.DisplayName, &provider.PortalReturnURL, &provider.AutoProvision, &provider.DefaultDepartmentID, &provider.Enabled, &provider.Revision, &configJSON, &provider.CreatedAt, &provider.UpdatedAt); err != nil {
|
||||
return provider, mapRepositoryError(err)
|
||||
}
|
||||
if json.Unmarshal(configJSON, &provider.Config) != nil {
|
||||
return provider, ErrUnavailable
|
||||
}
|
||||
return provider, nil
|
||||
}
|
||||
|
||||
func (r *Repository) CreateSAMLProvider(ctx context.Context, provider SAMLProvider, actor string) (SAMLProvider, error) {
|
||||
id, err := platformid.NewUUID()
|
||||
if err != nil {
|
||||
return provider, ErrUnavailable
|
||||
}
|
||||
provider.ID = id
|
||||
return r.storeSAMLProvider(ctx, provider, actor, true)
|
||||
}
|
||||
|
||||
func (r *Repository) UpdateSAMLProvider(ctx context.Context, provider SAMLProvider, actor string) (SAMLProvider, error) {
|
||||
return r.storeSAMLProvider(ctx, provider, actor, false)
|
||||
}
|
||||
|
||||
func (r *Repository) storeSAMLProvider(ctx context.Context, provider SAMLProvider, actor string, creating bool) (SAMLProvider, error) {
|
||||
configJSON, err := json.Marshal(provider.Config)
|
||||
if err != nil {
|
||||
return provider, ErrUnavailable
|
||||
}
|
||||
tx, err := r.pool.Begin(ctx)
|
||||
if err != nil {
|
||||
return provider, ErrUnavailable
|
||||
}
|
||||
defer tx.Rollback(ctx)
|
||||
if creating {
|
||||
err = tx.QueryRow(ctx, `INSERT INTO gateway.identity_providers(id,code,kind,display_name,portal_return_url,auto_provision,default_department_id,enabled,config) VALUES($1,$2,'saml',$3,$4,$5,$6,$7,$8) RETURNING revision,created_at,updated_at`, provider.ID, provider.Code, provider.DisplayName, provider.PortalReturnURL, provider.AutoProvision, provider.DefaultDepartmentID, provider.Enabled, configJSON).Scan(&provider.Revision, &provider.CreatedAt, &provider.UpdatedAt)
|
||||
} else {
|
||||
err = tx.QueryRow(ctx, `UPDATE gateway.identity_providers SET code=$2,display_name=$3,portal_return_url=$4,auto_provision=$5,default_department_id=$6,enabled=$7,config=$8,revision=revision+1,updated_at=clock_timestamp() WHERE id=$1 AND kind='saml' RETURNING revision,created_at,updated_at`, provider.ID, provider.Code, provider.DisplayName, provider.PortalReturnURL, provider.AutoProvision, provider.DefaultDepartmentID, provider.Enabled, configJSON).Scan(&provider.Revision, &provider.CreatedAt, &provider.UpdatedAt)
|
||||
}
|
||||
if err != nil {
|
||||
return provider, mapManagementError(err)
|
||||
}
|
||||
eventID, err := platformid.NewUUID()
|
||||
if err != nil {
|
||||
return provider, ErrUnavailable
|
||||
}
|
||||
eventType := "identity_provider.updated"
|
||||
if creating {
|
||||
eventType = "identity_provider.created"
|
||||
}
|
||||
payload, _ := json.Marshal(map[string]any{"identity_provider_id": provider.ID, "kind": "saml", "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, provider.ID, payload); err != nil {
|
||||
return provider, ErrUnavailable
|
||||
}
|
||||
if err := tx.Commit(ctx); err != nil {
|
||||
return provider, ErrUnavailable
|
||||
}
|
||||
return provider, nil
|
||||
}
|
||||
Reference in New Issue
Block a user