From 0885c45f34d49d84157edbf1473979f040fc4603 Mon Sep 17 00:00:00 2001 From: Qiu Jian Date: Thu, 16 Apr 2020 22:19:10 +0800 Subject: [PATCH] fix: cas auto create project fetch failure --- pkg/keystone/driver/cas/cas.go | 22 +++++++++++------- pkg/keystone/driver/cas/cas_test.go | 29 +++++++++++++++++++----- pkg/keystone/models/identity_provider.go | 2 +- pkg/keystone/models/projects.go | 12 +++++----- 4 files changed, 44 insertions(+), 21 deletions(-) diff --git a/pkg/keystone/driver/cas/cas.go b/pkg/keystone/driver/cas/cas.go index 6cdc883be6..45f0b6fa41 100644 --- a/pkg/keystone/driver/cas/cas.go +++ b/pkg/keystone/driver/cas/cas.go @@ -27,6 +27,7 @@ import ( "yunion.io/x/pkg/errors" api "yunion.io/x/onecloud/pkg/apis/identity" + "yunion.io/x/onecloud/pkg/cloudcommon/consts" "yunion.io/x/onecloud/pkg/httperrors" "yunion.io/x/onecloud/pkg/keystone/driver" "yunion.io/x/onecloud/pkg/keystone/models" @@ -148,22 +149,27 @@ func (self *SCASDriver) Authenticate(ctx context.Context, ident mcclient.SAuthen return nil, errors.Wrap(err, "models.UserManager.FetchUserExtended") } - self.userTryJoinProject(ctx, usr, domain, resp) + self.userTryJoinProject(ctx, usr, domain.Id, resp) return extUser, nil } -func (self *SCASDriver) userTryJoinProject(ctx context.Context, usr *models.SUser, domain *models.SDomain, resp []byte) { +func (self *SCASDriver) userTryJoinProject(ctx context.Context, usr *models.SUser, domainId string, resp []byte) { var err error var targetProject *models.SProject + log.Debugf("userTryJoinProject resp %s proj %s", string(resp), self.casConfig.CasProjectAttribute) + if !consts.GetNonDefaultDomainProjects() { + domainId = api.DEFAULT_DOMAIN_ID + } if len(self.casConfig.CasProjectAttribute) > 0 { - projName := fetchAttribute(resp, self.casConfig.CasProjectAttribute) - if len(projName) > 0 { - targetProject, err = models.ProjectManager.FetchProject("", projName, domain.Id, domain.Name) + casProjName := fetchAttribute(resp, self.casConfig.CasProjectAttribute) + if len(casProjName) > 0 { + projName := models.NormalizeProjectName(casProjName) + targetProject, err = models.ProjectManager.FetchProject("", projName, domainId, "") if err != nil { log.Errorf("fetch project %s fail %s", projName, err) if errors.Cause(err) == sql.ErrNoRows && self.casConfig.AutoCreateCasProject.IsTrue() { - targetProject, err = models.ProjectManager.NewProject(ctx, projName, "cas project", domain) + targetProject, err = models.ProjectManager.NewProject(ctx, projName, "cas project", domainId) if err != nil { log.Errorf("auto create project %s fail %s", projName, err) } @@ -183,7 +189,7 @@ func (self *SCASDriver) userTryJoinProject(ctx context.Context, usr *models.SUse if len(self.casConfig.CasRoleAttribute) > 0 { roleName := fetchAttribute(resp, self.casConfig.CasRoleAttribute) if len(roleName) > 0 { - targetRole, err = models.RoleManager.FetchRole("", roleName, domain.Id, domain.Name) + targetRole, err = models.RoleManager.FetchRole("", roleName, domainId, "") if err != nil { log.Errorf("fetch role %s fail %s", roleName, err) } @@ -205,7 +211,7 @@ func (self *SCASDriver) userTryJoinProject(ctx context.Context, usr *models.SUse } func fetchAttribute(heystack []byte, name string) string { - pattern := regexp.MustCompile(fmt.Sprintf(`<%s>(\w+)`, name, name)) + pattern := regexp.MustCompile(fmt.Sprintf(`<%s>([^<]*)`, name, name)) result := pattern.FindAllStringSubmatch(string(heystack), -1) if len(result) > 0 && len(result[0]) > 1 { return strings.TrimSpace(result[0][1]) diff --git a/pkg/keystone/driver/cas/cas_test.go b/pkg/keystone/driver/cas/cas_test.go index 96b9a865e7..3f4fbe355d 100644 --- a/pkg/keystone/driver/cas/cas_test.go +++ b/pkg/keystone/driver/cas/cas_test.go @@ -48,15 +48,32 @@ func TestXmlUnmarshal(t *testing.T) { } func TestFetchAttribute(t *testing.T) { - xmlstr := ` + cases := []struct { + Xml string + Key string + Want string + }{ + { + Xml: ` casuser casproj -` - got := fetchAttribute([]byte(xmlstr), "cas:proj") - want := "casproj" - if got != want { - t.Errorf("want %s got %s", want, got) +`, + Key: "cas:proj", + Want: "casproj", + }, + { + Xml: ` + lcftest0416周凌测试无线公司1112342`, + Key: "cas:proj", + Want: "周凌测试无线公司1112342", + }, + } + for _, c := range cases { + got := fetchAttribute([]byte(c.Xml), c.Key) + if got != c.Want { + t.Errorf("want %s got %s", c.Want, got) + } } } diff --git a/pkg/keystone/models/identity_provider.go b/pkg/keystone/models/identity_provider.go index 8d82d2a375..cd3fc3a581 100644 --- a/pkg/keystone/models/identity_provider.go +++ b/pkg/keystone/models/identity_provider.go @@ -748,7 +748,7 @@ func (self *SIdentityProvider) SyncOrCreateDomain(ctx context.Context, extId str _, err := ProjectManager.NewProject(ctx, fmt.Sprintf("%s_default_project", extName), fmt.Sprintf("Default project for domain %s", extName), - domain, + domain.Id, ) if err != nil { log.Errorf("ProjectManager.NewProject fail %s", err) diff --git a/pkg/keystone/models/projects.go b/pkg/keystone/models/projects.go index 218dfbf475..8f44b1f7c9 100644 --- a/pkg/keystone/models/projects.go +++ b/pkg/keystone/models/projects.go @@ -530,16 +530,16 @@ func (project *SProject) PerformLeave( return nil, nil } -func (manager *SProjectManager) NewProject(ctx context.Context, name string, desc string, domain *SDomain) (*SProject, error) { - lockman.LockClass(ctx, manager, domain.Id) - defer lockman.ReleaseClass(ctx, manager, domain.Id) +func (manager *SProjectManager) NewProject(ctx context.Context, name string, desc string, domainId string) (*SProject, error) { + lockman.LockClass(ctx, manager, domainId) + defer lockman.ReleaseClass(ctx, manager, domainId) project := &SProject{} project.SetModelManager(ProjectManager, project) projectName := NormalizeProjectName(name) ownerId := &db.SOwnerId{} if manager.NamespaceScope() == rbacutils.ScopeDomain { - ownerId.DomainId = domain.Id + ownerId.DomainId = domainId } newName, err := db.GenerateName(ProjectManager, ownerId, projectName) if err != nil { @@ -548,10 +548,10 @@ func (manager *SProjectManager) NewProject(ctx context.Context, name string, des newName = projectName } project.Name = newName - project.DomainId = domain.Id + project.DomainId = domainId project.Description = desc project.IsDomain = tristate.False - project.ParentId = domain.Id + project.ParentId = domainId err = ProjectManager.TableSpec().Insert(project) if err != nil { return nil, errors.Wrap(err, "Insert")