Hotfix/qj password login wrong idp (#7511)

* fix: role assignment code to include group info for group users

* fix: password authentication may choose wrong idp backend

Co-authored-by: Qiu Jian <qiujian@yunionyun.com>
This commit is contained in:
Jian Qiu
2020-08-07 14:19:34 +08:00
committed by GitHub
parent e8e5c3dbe5
commit b0505ba5fd
11 changed files with 254 additions and 142 deletions
+3
View File
@@ -182,9 +182,12 @@ func (t SAuthToken) GetAuthCookie(token mcclient.TokenCredential) string {
info := jsonutils.NewDict()
info.Add(jsonutils.NewTimeString(token.GetExpires()), "exp")
info.Add(jsonutils.NewString(sid), "session")
info.Add(jsonutils.NewBool(t.verifyTotp), "totp_verified") // 用户totp验证通过
info.Add(jsonutils.NewBool(t.initTotp), "totp_init") // 是否初始化TOTP密钥
info.Add(jsonutils.NewBool(t.enableTotp), "totp_on") // 用户totp 开启状态。 True(已开启)|False(未开启)
info.Add(jsonutils.NewBool(options.Options.EnableTotp), "system_totp_on") // 全局totp 开启状态。 True(已开启)|False(未开启)
info.Add(jsonutils.NewString(token.GetUserId()), "user_id")
info.Add(jsonutils.NewString(token.GetUserName()), "user")
return info.String()
}
+2 -2
View File
@@ -36,6 +36,6 @@ type SUserExtended struct {
DomainName string
DomainEnabled bool
IsLocal bool
IdpId string
IdpName string
// IdpId string
// IdpName string
}
+93 -77
View File
@@ -461,21 +461,55 @@ func roleAssignmentHandler(ctx context.Context, w http.ResponseWriter, r *http.R
}
func (manager *SAssignmentManager) queryAll(userId, groupId, roleId, domainId, projectId string) *sqlchemy.SQuery {
q := manager.Query("type", "actor_id", "target_id", "role_id")
assigments := manager.Query().SubQuery()
q := assigments.Query(
assigments.Field("type"),
sqlchemy.NewFunction(
sqlchemy.NewCase().When(sqlchemy.OR(
sqlchemy.Equals(assigments.Field("type"), sqlchemy.NewStringField(api.AssignmentUserProject)),
sqlchemy.Equals(assigments.Field("type"), sqlchemy.NewStringField(api.AssignmentUserDomain)),
), assigments.Field("actor_id")).Else(sqlchemy.NewStringField("")),
"user_id",
),
sqlchemy.NewFunction(
sqlchemy.NewCase().When(sqlchemy.OR(
sqlchemy.Equals(assigments.Field("type"), sqlchemy.NewStringField(api.AssignmentGroupProject)),
sqlchemy.Equals(assigments.Field("type"), sqlchemy.NewStringField(api.AssignmentGroupDomain)),
), assigments.Field("actor_id")).Else(sqlchemy.NewStringField("")),
"group_id",
),
sqlchemy.NewFunction(
sqlchemy.NewCase().When(sqlchemy.OR(
sqlchemy.Equals(assigments.Field("type"), sqlchemy.NewStringField(api.AssignmentUserDomain)),
sqlchemy.Equals(assigments.Field("type"), sqlchemy.NewStringField(api.AssignmentGroupDomain)),
), assigments.Field("target_id")).Else(sqlchemy.NewStringField("")),
"domain_id",
),
sqlchemy.NewFunction(
sqlchemy.NewCase().When(sqlchemy.OR(
sqlchemy.Equals(assigments.Field("type"), sqlchemy.NewStringField(api.AssignmentUserProject)),
sqlchemy.Equals(assigments.Field("type"), sqlchemy.NewStringField(api.AssignmentGroupProject)),
), assigments.Field("target_id")).Else(sqlchemy.NewStringField("")),
"project_id",
),
assigments.Field("role_id"),
)
// here use subquery.query to produce a effective reference to case function fields
q = q.SubQuery().Query()
if len(userId) > 0 {
q = q.In("type", []string{api.AssignmentUserProject, api.AssignmentUserDomain}).Equals("actor_id", userId)
q = q.In("type", []string{api.AssignmentUserProject, api.AssignmentUserDomain}).Equals("user_id", userId)
}
if len(groupId) > 0 {
q = q.In("type", []string{api.AssignmentGroupProject, api.AssignmentGroupDomain}).Equals("actor_id", groupId)
q = q.In("type", []string{api.AssignmentGroupProject, api.AssignmentGroupDomain}).Equals("group_id", groupId)
}
if len(roleId) > 0 {
q = q.Equals("role_id", roleId)
}
if len(projectId) > 0 {
q = q.Equals("target_id", projectId).In("type", []string{api.AssignmentUserProject, api.AssignmentGroupProject})
q = q.Equals("project_id", projectId).In("type", []string{api.AssignmentUserProject, api.AssignmentGroupProject})
}
if len(domainId) > 0 {
q = q.Equals("target_id", domainId).In("type", []string{api.AssignmentUserDomain, api.AssignmentGroupDomain})
q = q.Equals("domain_id", domainId).In("type", []string{api.AssignmentUserDomain, api.AssignmentGroupDomain})
}
return q
}
@@ -486,51 +520,44 @@ func fetchRoleAssignmentPolicies(ra *api.SRoleAssignment) {
ra.Policies.System = policy.PolicyManager.MatchedPolicyNames(rbacutils.ScopeSystem, ra)
}
func (assign *SAssignment) getRoleAssignment(domains, projects, groups, users, roles map[string]api.SFetchDomainObject, fetchPolicies bool) api.SRoleAssignment {
type sAssignmentInternal struct {
Type string `json:"type"`
UserId string `json:"user_id"`
GroupId string `json:"group_id"`
DomainId string `json:"domain_id"`
ProjectId string `json:"project_id"`
RoleId string `json:"role_id"`
}
func (assign *sAssignmentInternal) getRoleAssignment(domains, projects, groups, users, roles map[string]api.SFetchDomainObject, fetchPolicies bool) api.SRoleAssignment {
ra := api.SRoleAssignment{}
ra.Role.Id = assign.RoleId
ra.Role.Name = roles[assign.RoleId].Name
ra.Role.Domain.Id = roles[assign.RoleId].DomainId
ra.Role.Domain.Name = roles[assign.RoleId].Domain
switch assign.Type {
case api.AssignmentUserDomain:
ra.Scope.Domain.Id = assign.TargetId
ra.Scope.Domain.Name = domains[assign.TargetId].Name
ra.User.Id = assign.ActorId
ra.User.Name = users[assign.ActorId].Name
ra.User.Domain.Id = users[assign.ActorId].DomainId
ra.User.Domain.Name = users[assign.ActorId].Domain
case api.AssignmentUserProject:
ra.Scope.Project.Id = assign.TargetId
ra.Scope.Project.Name = projects[assign.TargetId].Name
ra.Scope.Project.Domain.Id = projects[assign.TargetId].DomainId
ra.Scope.Project.Domain.Name = projects[assign.TargetId].Domain
ra.User.Id = assign.ActorId
ra.User.Name = users[assign.ActorId].Name
ra.User.Domain.Id = users[assign.ActorId].DomainId
ra.User.Domain.Name = users[assign.ActorId].Domain
if fetchPolicies {
fetchRoleAssignmentPolicies(&ra)
}
case api.AssignmentGroupDomain:
ra.Scope.Domain.Id = assign.TargetId
ra.Scope.Domain.Name = domains[assign.TargetId].Name
ra.Group.Id = assign.ActorId
ra.Group.Name = groups[assign.ActorId].Name
ra.Group.Domain.Id = groups[assign.ActorId].DomainId
ra.Group.Domain.Name = groups[assign.ActorId].Domain
case api.AssignmentGroupProject:
ra.Scope.Project.Id = assign.TargetId
ra.Scope.Project.Name = projects[assign.TargetId].Name
ra.Scope.Project.Domain.Id = projects[assign.TargetId].DomainId
ra.Scope.Project.Domain.Name = projects[assign.TargetId].Domain
ra.Group.Id = assign.ActorId
ra.Group.Name = groups[assign.ActorId].Name
ra.Group.Domain.Id = groups[assign.ActorId].DomainId
ra.Group.Domain.Name = groups[assign.ActorId].Domain
if len(assign.UserId) > 0 {
ra.User.Id = assign.UserId
ra.User.Name = users[assign.UserId].Name
ra.User.Domain.Id = users[assign.UserId].DomainId
ra.User.Domain.Name = users[assign.UserId].Domain
}
if len(assign.GroupId) > 0 {
ra.Group.Id = assign.GroupId
ra.Group.Name = groups[assign.GroupId].Name
ra.Group.Domain.Id = groups[assign.GroupId].DomainId
ra.Group.Domain.Name = groups[assign.GroupId].Domain
}
if len(assign.ProjectId) > 0 {
ra.Scope.Project.Id = assign.ProjectId
ra.Scope.Project.Name = projects[assign.ProjectId].Name
ra.Scope.Project.Domain.Id = projects[assign.ProjectId].DomainId
ra.Scope.Project.Domain.Name = projects[assign.ProjectId].Domain
if fetchPolicies {
fetchRoleAssignmentPolicies(&ra)
}
} else if len(assign.DomainId) > 0 {
ra.Scope.Domain.Id = assign.DomainId
ra.Scope.Domain.Name = domains[assign.DomainId].Name
}
return ra
}
@@ -542,37 +569,28 @@ func (manager *SAssignmentManager) FetchAll(userId, groupId, roleId, domainId, p
memberships := UsergroupManager.Query("user_id", "group_id").SubQuery()
grpproj := manager.queryAll("", groupId, roleId, domainId, projectId).Equals("type", api.AssignmentGroupProject).SubQuery()
q2 := grpproj.Query(sqlchemy.NewStringField(api.AssignmentUserProject).Label("type"),
memberships.Field("user_id", "actor_id"),
grpproj.Field("target_id"), grpproj.Field("role_id"))
q2 = q2.Join(memberships, sqlchemy.Equals(grpproj.Field("actor_id"), memberships.Field("group_id")))
q2 = q2.Filter(sqlchemy.Equals(grpproj.Field("type"), api.AssignmentGroupProject))
grpproj := manager.queryAll("", groupId, roleId, domainId, projectId).In("type", []string{api.AssignmentGroupProject, api.AssignmentGroupDomain}).SubQuery()
q2 := grpproj.Query(
grpproj.Field("type"),
memberships.Field("user_id"),
grpproj.Field("group_id"),
grpproj.Field("domain_id"),
grpproj.Field("project_id"),
grpproj.Field("role_id"),
)
q2 = q2.Join(memberships, sqlchemy.Equals(grpproj.Field("group_id"), memberships.Field("group_id")))
if len(userId) > 0 {
q2 = q2.Filter(sqlchemy.Equals(memberships.Field("user_id"), userId))
}
grpdom := manager.queryAll("", groupId, roleId, domainId, projectId).Equals("type", api.AssignmentGroupDomain).SubQuery()
q3 := grpdom.Query(sqlchemy.NewStringField(api.AssignmentUserDomain).Label("type"),
memberships.Field("user_id", "actor_id"),
grpdom.Field("target_id"), grpdom.Field("role_id"))
q3 = q3.Join(memberships, sqlchemy.Equals(grpdom.Field("actor_id"), memberships.Field("group_id")))
q3 = q3.Filter(sqlchemy.Equals(grpdom.Field("type"), api.AssignmentGroupDomain))
if len(userId) > 0 {
q3 = q3.Filter(sqlchemy.Equals(memberships.Field("user_id"), userId))
}
q = sqlchemy.Union(usrq, q2, q3).Query().Distinct()
q = sqlchemy.Union(usrq, q2).Query().Distinct()
} else {
q = manager.queryAll(userId, groupId, roleId, domainId, projectId).Distinct()
}
if !includeSystem {
users := UserManager.Query().SubQuery()
q = q.LeftJoin(users, sqlchemy.AND(
sqlchemy.Equals(q.Field("actor_id"), users.Field("id")),
sqlchemy.In(q.Field("type"), []string{api.AssignmentUserProject, api.AssignmentUserDomain}),
))
q = q.LeftJoin(users, sqlchemy.Equals(q.Field("user_id"), users.Field("id")))
q = q.Filter(sqlchemy.OR(
sqlchemy.IsFalse(users.Field("is_system_account")),
sqlchemy.IsNull(users.Field("is_system_account")),
@@ -591,7 +609,7 @@ func (manager *SAssignmentManager) FetchAll(userId, groupId, roleId, domainId, p
q = q.Offset(offset)
}
assigns := make([]SAssignment, 0)
assigns := make([]sAssignmentInternal, 0)
err = q.All(&assigns)
if err != nil && err != sql.ErrNoRows {
return nil, -1, httperrors.NewInternalServerError("query error %s", err)
@@ -604,19 +622,17 @@ func (manager *SAssignmentManager) FetchAll(userId, groupId, roleId, domainId, p
roleIds := stringutils2.SSortedStrings{}
for i := range assigns {
switch assigns[i].Type {
case api.AssignmentGroupProject:
projectIds = stringutils2.Append(projectIds, assigns[i].TargetId)
groupIds = stringutils2.Append(groupIds, assigns[i].ActorId)
case api.AssignmentGroupDomain:
domainIds = stringutils2.Append(domainIds, assigns[i].TargetId)
groupIds = stringutils2.Append(groupIds, assigns[i].ActorId)
case api.AssignmentUserProject:
projectIds = stringutils2.Append(projectIds, assigns[i].TargetId)
userIds = stringutils2.Append(userIds, assigns[i].ActorId)
case api.AssignmentUserDomain:
domainIds = stringutils2.Append(domainIds, assigns[i].TargetId)
userIds = stringutils2.Append(userIds, assigns[i].ActorId)
if len(assigns[i].UserId) > 0 {
userIds = stringutils2.Append(userIds, assigns[i].UserId)
}
if len(assigns[i].GroupId) > 0 {
groupIds = stringutils2.Append(groupIds, assigns[i].GroupId)
}
if len(assigns[i].DomainId) > 0 {
domainIds = stringutils2.Append(domainIds, assigns[i].DomainId)
}
if len(assigns[i].ProjectId) > 0 {
projectIds = stringutils2.Append(projectIds, assigns[i].ProjectId)
}
roleIds = stringutils2.Append(roleIds, assigns[i].RoleId)
}
+1 -1
View File
@@ -151,7 +151,7 @@ func (manager *SIdmappingManager) FetchEntities(idStr string, entType string) ([
q := manager.Query().Equals("public_id", idStr).Equals("entity_type", entType)
idMaps := make([]SIdmapping, 0)
err := db.FetchModelObjects(manager, q, &idMaps)
if err != nil {
if err != nil && errors.Cause(err) != sql.ErrNoRows {
return nil, errors.Wrap(err, "FetchModelObjects")
} else {
return idMaps, nil
+19 -2
View File
@@ -113,7 +113,7 @@ func (manager *SIdentityProviderManager) initializeAutoCreateUser() error {
if errors.Cause(err) == sql.ErrNoRows {
return nil
} else {
return errors.Wrap(err, "FetchModelObjeccts")
return errors.Wrap(err, "FetchModelObjects")
}
}
for i := range idps {
@@ -141,7 +141,7 @@ func (manager *SIdentityProviderManager) initializeIcon() error {
if errors.Cause(err) == sql.ErrNoRows {
return nil
} else {
return errors.Wrap(err, "FetchModelObjeccts")
return errors.Wrap(err, "FetchModelObjects")
}
}
for i := range idps {
@@ -1363,3 +1363,20 @@ func (idp *SIdentityProvider) SyncOrCreateDomainAndUser(ctx context.Context, ext
}
return domain, usr, nil
}
func (manager *SIdentityProviderManager) FetchIdentityProvidersByUserId(uid string, drivers []string) ([]SIdentityProvider, error) {
idps := make([]SIdentityProvider, 0)
idmappings := IdmappingManager.Query().SubQuery()
q := manager.Query()
q = q.Join(idmappings, sqlchemy.Equals(q.Field("id"), idmappings.Field("domain_id")))
q = q.Filter(sqlchemy.Equals(idmappings.Field("entity_type"), api.IdMappingEntityUser))
q = q.Filter(sqlchemy.Equals(idmappings.Field("public_id"), uid))
if len(drivers) > 0 {
q = q.Filter(sqlchemy.In(q.Field("driver"), drivers))
}
err := db.FetchModelObjects(manager, q, &idps)
if err != nil && errors.Cause(err) != sql.ErrNoRows {
return nil, errors.Wrap(err, "FetchModelObjects")
}
return idps, nil
}
+7 -6
View File
@@ -116,7 +116,7 @@ func (manager *SUserManager) InitializeData() error {
}
name := extUser.LocalName
if len(name) == 0 {
name = extUser.IdpName
name = extUser.DomainName
}
var desc, email, mobile, dispName string
if users[i].Extra != nil {
@@ -233,7 +233,7 @@ func (manager *SUserManager) FetchUserExtended(userId, userName, domainId, domai
// nonlocalUsers := NonlocalUserManager.Query().SubQuery()
users := UserManager.Query().SubQuery()
domains := DomainManager.Query().SubQuery()
idmappings := IdmappingManager.Query().SubQuery()
// idmappings := IdmappingManager.Query().SubQuery()
q := users.Query(
users.Field("id"),
@@ -251,13 +251,13 @@ func (manager *SUserManager) FetchUserExtended(userId, userName, domainId, domai
localUsers.Field("name", "local_name"),
domains.Field("name", "domain_name"),
domains.Field("enabled", "domain_enabled"),
idmappings.Field("domain_id", "idp_id"),
idmappings.Field("local_id", "idp_name"),
// idmappings.Field("domain_id", "idp_id"),
// idmappings.Field("local_id", "idp_name"),
)
q = q.Join(domains, sqlchemy.Equals(users.Field("domain_id"), domains.Field("id")))
q = q.LeftJoin(localUsers, sqlchemy.Equals(localUsers.Field("user_id"), users.Field("id")))
q = q.LeftJoin(idmappings, sqlchemy.Equals(users.Field("id"), idmappings.Field("public_id")))
// q = q.LeftJoin(idmappings, sqlchemy.Equals(users.Field("id"), idmappings.Field("public_id")))
if len(userId) > 0 {
q = q.Filter(sqlchemy.Equals(users.Field("id"), userId))
@@ -279,7 +279,8 @@ func (manager *SUserManager) FetchUserExtended(userId, userName, domainId, domai
return nil, err
}
if len(extUser.IdpName) > 0 {
idMaps, err := IdmappingManager.FetchEntities(extUser.Id, api.IdMappingEntityUser)
if len(idMaps) > 0 {
extUser.IsLocal = false
} else {
extUser.IsLocal = true
+20 -18
View File
@@ -71,6 +71,7 @@ func authUserByIdentity(ctx context.Context, ident mcclient.SAuthenticationIdent
return nil, ErrEmptyAuth
}
if len(ident.Password.User.Name) > 0 && len(ident.Password.User.Id) == 0 && len(ident.Password.User.Domain.Id) == 0 && len(ident.Password.User.Domain.Name) == 0 {
// no use domain specified, try to find use domain
users := models.UserManager.Query().SubQuery()
idMappings := models.IdmappingManager.Query().SubQuery()
q := users.Query()
@@ -103,25 +104,16 @@ func authUserByIdentity(ctx context.Context, ident mcclient.SAuthenticationIdent
return nil, errors.Wrap(err, "Query user")
}
ident.Password.User.Domain.Id = usr.DomainId
idmaps, err := models.IdmappingManager.FetchEntities(usr.Id, api.IdMappingEntityUser)
if err != nil && err != sql.ErrNoRows {
return nil, errors.Wrap(err, "IdmappingManager.FetchEntity")
idps, err := models.IdentityProviderManager.FetchIdentityProvidersByUserId(usr.Id, api.PASSWORD_PROTECTED_IDPS)
if err != nil {
return nil, errors.Wrap(err, "IdentityProviderManager.FetchIdentityProvidersByUserId")
}
var idmap *models.SIdmapping
for i := range idmaps {
idp, err := models.IdentityProviderManager.FetchIdentityProviderById(idmaps[i].IdpId)
if err != nil {
return nil, errors.Wrap(err, "IdentityProviderManager.FetchIdentityProviderById")
}
if idp.Driver == api.IdentityDriverLDAP {
idmap = &idmaps[i]
break
}
}
if idmap == nil { // sql
if len(idps) == 0 {
idpId = api.DEFAULT_IDP_ID
} else if len(idps) == 1 {
idpId = idps[0].Id
} else {
idpId = idmap.IdpId
return nil, sqlchemy.ErrDuplicateEntry
}
}
} else {
@@ -144,7 +136,17 @@ func authUserByIdentity(ctx context.Context, ident mcclient.SAuthenticationIdent
idpId = mapping.IdpId
} else {
// user exists, query user's idp
idpId = usrExt.IdpId
idps, err := models.IdentityProviderManager.FetchIdentityProvidersByUserId(usrExt.Id, api.PASSWORD_PROTECTED_IDPS)
if err != nil {
return nil, errors.Wrap(err, "IdentityProviderManager.FetchIdentityProvidersByUserId")
}
if len(idps) == 0 {
idpId = api.DEFAULT_IDP_ID
} else if len(idps) == 1 {
idpId = idps[0].Id
} else {
return nil, sqlchemy.ErrDuplicateEntry
}
}
}
@@ -429,7 +431,7 @@ func AuthenticateV3(ctx context.Context, input mcclient.SAuthenticationInputV3)
return nil, errors.Wrap(err, "authUserByOAuth2")
}
default:
// auth by other methods, password, openid, saml, etc...
// auth by other methods, e.g. password , etc...
user, err = authUserByIdentityV3(ctx, input)
if err != nil {
return nil, errors.Wrap(err, "authUserByIdentityV3")
+68
View File
@@ -0,0 +1,68 @@
// Copyright 2019 Yunion
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package sqlchemy
import (
"bytes"
)
type sCaseFieldBranch struct {
whenCondition ICondition
thenField IQueryField
}
type SCaseFunction struct {
branches []sCaseFieldBranch
elseField IQueryField
}
func NewFunction(ifunc IFunction, name string) IQueryField {
return &SFunctionFieldBase{
IFunction: ifunc,
alias: name,
}
}
func (cf *SCaseFunction) Else(field IQueryField) *SCaseFunction {
cf.elseField = field
return cf
}
func (cf *SCaseFunction) When(when ICondition, then IQueryField) *SCaseFunction {
cf.branches = append(cf.branches, sCaseFieldBranch{
whenCondition: when,
thenField: then,
})
return cf
}
func NewCase() *SCaseFunction {
return &SCaseFunction{}
}
func (cf *SCaseFunction) expression() string {
var buf bytes.Buffer
buf.WriteString("CASE ")
for i := range cf.branches {
buf.WriteString("WHEN ")
buf.WriteString(cf.branches[i].whenCondition.WhereClause())
buf.WriteString(" THEN ")
buf.WriteString(cf.branches[i].thenField.Reference())
}
buf.WriteString(" ELSE ")
buf.WriteString(cf.elseField.Reference())
buf.WriteString(" END")
return buf.String()
}
+1 -1
View File
@@ -20,8 +20,8 @@ import (
"reflect"
"yunion.io/x/log"
"yunion.io/x/pkg/util/reflectutils"
"yunion.io/x/pkg/errors"
"yunion.io/x/pkg/util/reflectutils"
)
/*
+36 -35
View File
@@ -16,45 +16,58 @@ package sqlchemy
import (
"fmt"
"log"
"strconv"
"strings"
)
type SFunctionField struct {
fields []IQueryField
function string
alias string
type IFunction interface {
expression() string
}
func (ff *SFunctionField) Expression() string {
fieldRefs := make([]interface{}, 0)
for _, f := range ff.fields {
fieldRefs = append(fieldRefs, f.Reference())
type SFunctionFieldBase struct {
IFunction
alias string
}
func (ff *SFunctionFieldBase) Reference() string {
if len(ff.alias) == 0 {
log.Fatalf("reference a function field without alias! %s", ff.expression())
}
return fmt.Sprintf("%s AS `%s`", fmt.Sprintf(ff.function, fieldRefs...), ff.Name())
return fmt.Sprintf("`%s`", ff.alias)
}
func (ff *SFunctionField) Name() string {
return ff.alias
func (ff *SFunctionFieldBase) Expression() string {
if len(ff.alias) > 0 {
// add alias
return fmt.Sprintf("%s AS `%s`", ff.expression(), ff.alias)
} else {
// no alias
return ff.expression()
}
}
func (ff *SFunctionField) Reference() string {
return ff.alias
func (ff *SFunctionFieldBase) Name() string {
if len(ff.alias) > 0 {
return ff.alias
} else {
return ff.expression()
}
}
func (ff *SFunctionField) Label(label string) IQueryField {
func (ff *SFunctionFieldBase) Label(label string) IQueryField {
if len(label) > 0 && label != ff.alias {
ff.alias = label
}
return ff
}
type SFunctionFieldWithoutAlias struct {
type SExprFunction struct {
fields []IQueryField
function string
}
func (ff *SFunctionFieldWithoutAlias) Expression() string {
func (ff *SExprFunction) expression() string {
fieldRefs := make([]interface{}, 0)
for _, f := range ff.fields {
fieldRefs = append(fieldRefs, f.Reference())
@@ -62,26 +75,14 @@ func (ff *SFunctionFieldWithoutAlias) Expression() string {
return fmt.Sprintf(ff.function, fieldRefs...)
}
func (ff *SFunctionFieldWithoutAlias) Name() string {
return ff.Expression()
}
func (ff *SFunctionFieldWithoutAlias) Reference() string {
return ff.Expression()
}
func (ff *SFunctionFieldWithoutAlias) Label(label string) IQueryField {
if len(label) > 0 {
return &SFunctionField{ff.fields, ff.function, label}
}
return ff
}
func NewFunctionField(name string, funcexp string, fields ...IQueryField) IQueryField {
if len(name) > 0 {
return &SFunctionField{function: funcexp, alias: name, fields: fields}
} else {
return &SFunctionFieldWithoutAlias{fields: fields, function: funcexp}
funcBase := &SExprFunction{
fields: fields,
function: funcexp,
}
return &SFunctionFieldBase{
IFunction: funcBase,
alias: name,
}
}
+4
View File
@@ -516,6 +516,10 @@ func (tq *SQuery) findField(name string) IQueryField {
func (tq *SQuery) internalFindField(name string) IQueryField {
for _, f := range tq.fields {
if f.Name() == name {
switch f.(type) {
case *SFunctionFieldBase:
log.Errorf("cannot directly reference a function alias, should use Subquery() to enclose the query")
}
return f
}
}