diff --git a/cmd/climc/shell/secgroups.go b/cmd/climc/shell/secgroups.go index 7d133892fd..e7fd9a15a9 100644 --- a/cmd/climc/shell/secgroups.go +++ b/cmd/climc/shell/secgroups.go @@ -27,6 +27,7 @@ import ( func init() { type SecGroupsListOptions struct { Equals string `help:"Secgroup ID or Name, filter secgroups whose rules equals the specified one"` + Server string `help:"Filter secgroups bound to specified server"` options.BaseListOptions } @@ -43,6 +44,9 @@ func init() { if len(args.Equals) > 0 { params.Add(jsonutils.NewString(args.Equals), "equals") } + if len(args.Server) > 0 { + params.Add(jsonutils.NewString(args.Server), "server") + } result, err := modules.SecGroups.List(s, params) if err != nil { return err diff --git a/docs/parameters/secgroup.yaml b/docs/parameters/secgroup.yaml index 760dd37f95..bc8b897e0a 100644 --- a/docs/parameters/secgroup.yaml +++ b/docs/parameters/secgroup.yaml @@ -23,3 +23,8 @@ shared_projects: description: 项目名称或ID description: 共享到的项目列表, 仅scope=project有效 +server: + name: server + in: query + type: string + description: Filter secgroups bound to specified server diff --git a/pkg/apis/compute/secgroup.go b/pkg/apis/compute/secgroup.go index a7dfae1d19..19032f4911 100644 --- a/pkg/apis/compute/secgroup.go +++ b/pkg/apis/compute/secgroup.go @@ -77,3 +77,8 @@ type SSecgroupCreateInput struct { Description string Rules []SSecgroupRuleCreateInput } + +type SSecgroupListFilterInput struct { + Equals string + Server string +} diff --git a/pkg/compute/models/secgroups.go b/pkg/compute/models/secgroups.go index 1ee0a85670..6695131ecb 100644 --- a/pkg/compute/models/secgroups.go +++ b/pkg/compute/models/secgroups.go @@ -71,11 +71,12 @@ type SSecurityGroup struct { } func (manager *SSecurityGroupManager) ListItemFilter(ctx context.Context, q *sqlchemy.SQuery, userCred mcclient.TokenCredential, query jsonutils.JSONObject) (*sqlchemy.SQuery, error) { - equalSecgroup, _ := query.GetString("equals") - if len(equalSecgroup) > 0 { - _secgroup, err := manager.FetchByIdOrName(userCred, equalSecgroup) + input := api.SSecgroupListFilterInput{} + query.Unmarshal(&input) + if len(input.Equals) > 0 { + _secgroup, err := manager.FetchByIdOrName(userCred, input.Equals) if err != nil { - return nil, httperrors.NewInputParameterError("Failed fetching secgroup %s", equalSecgroup) + return nil, httperrors.NewInputParameterError("Failed fetching secgroup %s", input.Equals) } secgroup := _secgroup.(*SSecurityGroup) sq := manager.Query().NotEquals("id", secgroup.Id) @@ -93,6 +94,22 @@ func (manager *SSecurityGroupManager) ListItemFilter(ctx context.Context, q *sql } q = q.In("id", secgroupIds) } + if len(input.Server) > 0 { + guest, err := GuestManager.FetchByIdOrName(userCred, input.Server) + if err != nil { + if err != sql.ErrNoRows { + return nil, httperrors.NewResourceNotFoundError("failed to found server %s", input.Server) + } + return nil, httperrors.NewGeneralError(err) + } + serverId := guest.GetId() + sq1 := GuestManager.Query("secgrp_id").Equals("id", serverId).SubQuery() + sq2 := GuestsecgroupManager.Query("secgroup_id").Equals("guest_id", serverId).SubQuery() + q = q.Filter(sqlchemy.OR( + sqlchemy.In(q.Field("id"), sq1), + sqlchemy.In(q.Field("id"), sq2), + )) + } return q, nil }