package identity import ( "context" "encoding/json" "errors" "strings" "time" platformid "aigateway.local/core/internal/platform/id" "github.com/jackc/pgx/v5" ) // SocialProvider 是内置扫码登录身份源(企微/钉钉/飞书)。 // 复用 identity_providers 表:client_id 存平台 AppID(企微为 corp_id), // encrypted_credentials 加密存放 AppSecret,agent_id 等平台特有参数放 config jsonb。 type SocialProvider struct { ID string `json:"id"` Code string `json:"code"` Kind string `json:"kind"` DisplayName string `json:"display_name"` ClientID string `json:"client_id"` AgentID string `json:"agent_id,omitempty"` EncryptedCredentials []byte `json:"-"` CredentialKEKVersion int `json:"-"` RedirectURI string `json:"redirect_uri"` PortalReturnURL string `json:"portal_return_url"` 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 socialCredentials struct { Secret string `json:"secret"` } // ProviderBinding 是门户账号与企微/钉钉/飞书账号的绑定关系。 type ProviderBinding struct { Kind string `json:"kind"` ProviderUID string `json:"provider_uid"` CreatedAt time.Time `json:"created_at"` } const socialKinds = "('wecom','dingtalk','feishu')" func scanSocialProvider(row pgx.Row) (SocialProvider, error) { var p SocialProvider var config []byte var defaultDepartment *string err := row.Scan(&p.ID, &p.Code, &p.DisplayName, &p.Kind, &p.ClientID, &config, &p.EncryptedCredentials, &p.CredentialKEKVersion, &p.RedirectURI, &p.PortalReturnURL, &p.AutoProvision, &defaultDepartment, &p.Enabled, &p.Revision, &p.CreatedAt, &p.UpdatedAt) if errors.Is(err, pgx.ErrNoRows) { return p, ErrNotFound } if err != nil { return p, mapRepositoryError(err) } p.DefaultDepartmentID = defaultDepartment var values map[string]string if json.Unmarshal(config, &values) == nil { p.AgentID = values["agent_id"] } return p, nil } // ListSocialProviders 返回全部扫码登录身份源。 func (r *Repository) ListSocialProviders(ctx context.Context) ([]SocialProvider, error) { rows, err := r.pool.Query(ctx, `SELECT id::text,code,display_name,kind,client_id,config,encrypted_credentials,credential_kek_version,redirect_uri,portal_return_url,auto_provision,default_department_id::text,enabled,revision,created_at,updated_at FROM gateway.identity_providers WHERE kind IN `+socialKinds+` ORDER BY kind,code`) if err != nil { return nil, mapRepositoryError(err) } defer rows.Close() items := []SocialProvider{} for rows.Next() { p, err := scanSocialProvider(rows) if err != nil { return nil, err } items = append(items, p) } return items, mapRepositoryError(rows.Err()) } // GetSocialProviderByKind 按 kind 返回扫码登录身份源。 func (r *Repository) GetSocialProviderByKind(ctx context.Context, kind string) (SocialProvider, error) { return scanSocialProvider(r.pool.QueryRow(ctx, `SELECT id::text,code,display_name,kind,client_id,config,encrypted_credentials,credential_kek_version,redirect_uri,portal_return_url,auto_provision,default_department_id::text,enabled,revision,created_at,updated_at FROM gateway.identity_providers WHERE kind=$1`, strings.ToLower(strings.TrimSpace(kind)))) } // GetSocialProviderByCode 按 SSO 代码返回扫码登录身份源(start/callback 分发用)。 func (r *Repository) GetSocialProviderByCode(ctx context.Context, code string) (SocialProvider, error) { return scanSocialProvider(r.pool.QueryRow(ctx, `SELECT id::text,code,display_name,kind,client_id,config,encrypted_credentials,credential_kek_version,redirect_uri,portal_return_url,auto_provision,default_department_id::text,enabled,revision,created_at,updated_at FROM gateway.identity_providers WHERE code=$1 AND kind IN `+socialKinds, strings.ToLower(strings.TrimSpace(code)))) } // SaveSocialProvider 创建/更新扫码登录身份源;replaceSecret=false 时保留原 Secret。 func (r *Repository) SaveSocialProvider(ctx context.Context, p SocialProvider, actorID string, creating, replaceSecret bool) (SocialProvider, error) { tx, err := r.pool.Begin(ctx) if err != nil { return p, ErrUnavailable } defer func() { _ = tx.Rollback(ctx) }() if creating { id, err := platformid.NewUUID() if err != nil { return p, err } p.ID = id err = tx.QueryRow(ctx, `INSERT INTO gateway.identity_providers(id,code,kind,display_name,client_id,encrypted_credentials,credential_kek_version,redirect_uri,portal_return_url,auto_provision,default_department_id,enabled,config) VALUES($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13) RETURNING revision,created_at,updated_at`, p.ID, p.Code, p.Kind, p.DisplayName, p.ClientID, p.EncryptedCredentials, p.CredentialKEKVersion, p.RedirectURI, p.PortalReturnURL, p.AutoProvision, p.DefaultDepartmentID, p.Enabled, agentConfigJSON(p.AgentID)).Scan(&p.Revision, &p.CreatedAt, &p.UpdatedAt) } else { err = tx.QueryRow(ctx, `UPDATE gateway.identity_providers SET code=$2,display_name=$3,client_id=$4,encrypted_credentials=CASE WHEN $13 THEN $5 ELSE encrypted_credentials END,credential_kek_version=CASE WHEN $13 THEN $6 ELSE credential_kek_version END,redirect_uri=$7,portal_return_url=$8,auto_provision=$9,default_department_id=$10,enabled=$11,config=$12,revision=revision+1,updated_at=clock_timestamp() WHERE id=$1 AND kind IN `+socialKinds+` RETURNING encrypted_credentials,credential_kek_version,revision,created_at,updated_at`, p.ID, p.Code, p.DisplayName, p.ClientID, p.EncryptedCredentials, p.CredentialKEKVersion, p.RedirectURI, p.PortalReturnURL, p.AutoProvision, p.DefaultDepartmentID, p.Enabled, agentConfigJSON(p.AgentID), replaceSecret).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": actorID}) 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 } // DeleteSocialProvider 删除扫码登录身份源及其全部绑定。 func (r *Repository) DeleteSocialProvider(ctx context.Context, kind string) error { kind = strings.ToLower(strings.TrimSpace(kind)) tx, err := r.pool.Begin(ctx) if err != nil { return ErrUnavailable } defer func() { _ = tx.Rollback(ctx) }() var id string if err = tx.QueryRow(ctx, `DELETE FROM gateway.identity_providers WHERE kind=$1 RETURNING id::text`, kind).Scan(&id); err != nil { if errors.Is(err, pgx.ErrNoRows) { return ErrNotFound } return mapManagementError(err) } if _, err = tx.Exec(ctx, `DELETE FROM gateway.portal_user_provider_bindings WHERE provider_kind=$1`, kind); err != nil { return ErrUnavailable } eventID, _ := platformid.NewUUID() if _, err = tx.Exec(ctx, `INSERT INTO gateway.outbox_events(event_id,event_type,event_version,aggregate_type,aggregate_id,payload) VALUES($1,'identity_provider.deleted',1,'identity_provider',$2,$3)`, eventID, id, `{"identity_provider_id":"`+id+`"}`); err != nil { return ErrUnavailable } return mapManagementError(tx.Commit(ctx)) } func agentConfigJSON(agentID string) []byte { if strings.TrimSpace(agentID) == "" { return []byte(`{}`) } raw, _ := json.Marshal(map[string]string{"agent_id": strings.TrimSpace(agentID)}) return raw } // FindProviderBinding 按 (kind, uid) 反查门户账号;未绑定返回 ErrNotFound。 func (r *Repository) FindProviderBinding(ctx context.Context, kind, uid string) (string, error) { var id string err := r.pool.QueryRow(ctx, `SELECT portal_user_id::text FROM gateway.portal_user_provider_bindings WHERE provider_kind=$1 AND provider_uid=$2`, strings.ToLower(strings.TrimSpace(kind)), uid).Scan(&id) if errors.Is(err, pgx.ErrNoRows) { return "", ErrNotFound } return id, mapRepositoryError(err) } // BindProvider 建立绑定。kind+uid 冲突(已被他人绑定)返回错误,账号重复绑定同一 // 平台(unique)冲突时先解绑旧绑定再写入,保证一个账号每平台至多一个绑定。 func (r *Repository) BindProvider(ctx context.Context, portalUserID, kind, uid string) error { kind = strings.ToLower(strings.TrimSpace(kind)) tx, err := r.pool.Begin(ctx) if err != nil { return ErrUnavailable } defer func() { _ = tx.Rollback(ctx) }() if _, err = tx.Exec(ctx, `DELETE FROM gateway.portal_user_provider_bindings WHERE portal_user_id=$1 AND provider_kind=$2`, portalUserID, kind); err != nil { return ErrUnavailable } tag, err := tx.Exec(ctx, `INSERT INTO gateway.portal_user_provider_bindings(portal_user_id,provider_kind,provider_uid) VALUES($1,$2,$3) ON CONFLICT(provider_kind,provider_uid) DO NOTHING`, portalUserID, kind, uid) if err != nil { return ErrUnavailable } if tag.RowsAffected() == 0 { return errors.New("该平台账号已被其他本系统账号绑定") } return mapManagementError(tx.Commit(ctx)) } // UnbindProvider 解除绑定(仅本人)。 func (r *Repository) UnbindProvider(ctx context.Context, portalUserID, kind string) error { _, err := r.pool.Exec(ctx, `DELETE FROM gateway.portal_user_provider_bindings WHERE portal_user_id=$1 AND provider_kind=$2`, portalUserID, strings.ToLower(strings.TrimSpace(kind))) return mapRepositoryError(err) } // ListProviderBindings 返回账号的全部扫码绑定。 func (r *Repository) ListProviderBindings(ctx context.Context, portalUserID string) ([]ProviderBinding, error) { rows, err := r.pool.Query(ctx, `SELECT provider_kind,provider_uid,created_at FROM gateway.portal_user_provider_bindings WHERE portal_user_id=$1 ORDER BY provider_kind`, portalUserID) if err != nil { return nil, mapRepositoryError(err) } defer rows.Close() items := []ProviderBinding{} for rows.Next() { var item ProviderBinding if err := rows.Scan(&item.Kind, &item.ProviderUID, &item.CreatedAt); err != nil { return nil, err } items = append(items, item) } return items, mapRepositoryError(rows.Err()) }