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+)%s>`, name, name))
+ pattern := regexp.MustCompile(fmt.Sprintf(`<%s>([^<]*)%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")