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 }