diff --git a/pkg/compute/models/secgrouprules.go b/pkg/compute/models/secgrouprules.go index 29dd59d98e..e107610536 100644 --- a/pkg/compute/models/secgrouprules.go +++ b/pkg/compute/models/secgrouprules.go @@ -13,6 +13,7 @@ import ( "yunion.io/x/onecloud/pkg/httperrors" "yunion.io/x/onecloud/pkg/mcclient" "yunion.io/x/pkg/util/compare" + "yunion.io/x/pkg/util/regutils" "yunion.io/x/pkg/util/secrules" "yunion.io/x/pkg/util/stringutils" "yunion.io/x/sqlchemy" @@ -129,42 +130,76 @@ func (self *SSecurityGroupRule) BeforeInsert() { } } -func (manager *SSecurityGroupRuleManager) ValidateCreateData( - ctx context.Context, - userCred mcclient.TokenCredential, - ownerProjId string, - query jsonutils.JSONObject, - data *jsonutils.JSONDict, -) (*jsonutils.JSONDict, error) { - if defsecgroup, _ := data.GetString("secgroup"); len(defsecgroup) > 0 { - if secgroup, _ := SecurityGroupManager.FetchByIdOrName(userCred, defsecgroup); secgroup != nil { - data.Set("secgroup_id", jsonutils.NewString(secgroup.GetId())) - } else { - return nil, httperrors.NewNotFoundError(fmt.Sprintf("Security Group %s not found", defsecgroup)) +func (manager *SSecurityGroupRuleManager) ValidateCreateData(ctx context.Context, userCred mcclient.TokenCredential, ownerProjId string, query jsonutils.JSONObject, data *jsonutils.JSONDict) (*jsonutils.JSONDict, error) { + defsecgroup, _ := data.GetString("secgroup") + if len(defsecgroup) == 0 { + return nil, httperrors.NewInputParameterError("Missing Security Group info") + } + secgroup, _ := SecurityGroupManager.FetchByIdOrName(userCred, defsecgroup) + if secgroup == nil { + return nil, httperrors.NewInputParameterError("Security Group %s not found", defsecgroup) + } + data.Add(jsonutils.NewString(secgroup.GetId()), "secgroup_id") + + priority, _ := data.Int("priority") + direction, _ := data.GetString("direction") + action, _ := data.GetString("action") + cidr, _ := data.GetString("cidr") + protocol, _ := data.GetString("protocol") + + if len(cidr) > 0 { + if !regutils.MatchCIDR(cidr) && !regutils.MatchIPAddr(cidr) { + return nil, httperrors.NewInputParameterError("invalid ip address: %s", cidr) } } else { - return nil, httperrors.NewInputParameterError("missing Security Group info") + data.Add(jsonutils.NewString("0.0.0.0/0"), "cidr") } - if _priority, _ := data.GetString("priority"); len(_priority) > 0 { - if priority, err := strconv.Atoi(_priority); err != nil { - return nil, httperrors.NewInputParameterError("UnSupport priority %s, only support 1-100", err.Error()) - } else if priority < 1 || priority > 100 { - return nil, httperrors.NewInputParameterError("UnSupport priority range, only support 1-100") - } + + rule := secrules.SecurityRule{ + Priority: int(priority), + Direction: secrules.TSecurityRuleDirection(direction), + Action: secrules.TSecurityRuleAction(action), + Protocol: protocol, + Ports: []int{}, + PortStart: -1, + PortEnd: -1, } - var fields []string - for _, field := range []string{"direction", "action", "cidr", "protocol", "ports"} { - if key, _ := data.GetString(field); len(key) > 0 { - if field == "direction" { - key += ":" + ports, _ := data.GetString("ports") + var err error + if len(ports) > 0 { + if strings.Index(ports, "-") > 0 { + portsInfo := strings.Split(ports, "-") + if len(portsInfo) != 2 { + return nil, httperrors.NewInputParameterError("invalid ports: %s", ports) } - fields = append(fields, key) - } else if field == "cidr" { - data.Add(jsonutils.NewString("0.0.0.0/0"), "cidr") + rule.PortStart, err = strconv.Atoi(portsInfo[0]) + if err != nil { + return nil, httperrors.NewInputParameterError("invalid port start: %s", portsInfo[0]) + } + rule.PortEnd, err = strconv.Atoi(portsInfo[1]) + if err != nil { + return nil, httperrors.NewInputParameterError("invalid port end: %s", portsInfo[1]) + } + } else if strings.Index(ports, ",") > 0 { + for _, port := range strings.Split(ports, ",") { + _port, err := strconv.Atoi(port) + if err != nil { + return nil, httperrors.NewInputParameterError("invalid port : %d", port) + } + rule.Ports = append(rule.Ports, _port) + } + } else { + port, err := strconv.Atoi(ports) + if err != nil { + return nil, httperrors.NewInputParameterError("invalid ports: %s", ports) + } + rule.Ports = append(rule.Ports, port) } } - if _, err := secrules.ParseSecurityRule(strings.Join(fields, " ")); err != nil { - return nil, err + + err = rule.ValidateRule() + if err != nil { + return nil, httperrors.NewInputParameterError(err.Error()) } return manager.SModelBaseManager.ValidateCreateData(ctx, userCred, ownerProjId, query, data) } diff --git a/pkg/scheduler/manager/task_queue.go b/pkg/scheduler/manager/task_queue.go index 0fffd0268d..0a413c8703 100644 --- a/pkg/scheduler/manager/task_queue.go +++ b/pkg/scheduler/manager/task_queue.go @@ -200,8 +200,8 @@ type TaskManager struct { func NewTaskManager(stopCh <-chan struct{}) *TaskManager { return &TaskManager{ taskExecutorQueueManager: NewTaskExecutorQueueManager(stopCh), - stopCh: stopCh, - lock: sync.Mutex{}, + stopCh: stopCh, + lock: sync.Mutex{}, } }