From 7ac07a446366dd499562f667c7dab18704c66896 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=B1=88=E8=BD=A9?= Date: Fri, 2 Nov 2018 13:29:04 +0800 Subject: [PATCH] =?UTF-8?q?=E6=9B=B4=E6=94=B9validate=20rule=20data=20?= =?UTF-8?q?=E6=96=B9=E5=BC=8F?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- pkg/compute/models/secgrouprules.go | 93 ++++++++++----- .../x/pkg/util/filterclause/filterclause.go | 6 +- .../yunion.io/x/pkg/util/regutils/regutils.go | 35 ++++-- .../yunion.io/x/pkg/util/secrules/secrules.go | 109 ++++++++++++++---- 4 files changed, 179 insertions(+), 64 deletions(-) diff --git a/pkg/compute/models/secgrouprules.go b/pkg/compute/models/secgrouprules.go index 8e596e11fc..fdb1f21fcc 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" @@ -122,42 +123,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.GetProjectId(), 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.GetProjectId(), 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/vendor/yunion.io/x/pkg/util/filterclause/filterclause.go b/vendor/yunion.io/x/pkg/util/filterclause/filterclause.go index 6f60b3343d..f820920bb8 100644 --- a/vendor/yunion.io/x/pkg/util/filterclause/filterclause.go +++ b/vendor/yunion.io/x/pkg/util/filterclause/filterclause.go @@ -56,11 +56,11 @@ func (fc *SFilterClause) QueryCondition(q *sqlchemy.SQuery) sqlchemy.ICondition case "like": return sqlchemy.Like(field, fc.params[0]) case "contains": - return sqlchemy.Like(field, fmt.Sprintf("%%%s%%", fc.params[0])) + return sqlchemy.Contains(field, fc.params[0]) case "startswith": - return sqlchemy.Like(field, fmt.Sprint("%s%%", fc.params[0])) + return sqlchemy.Startswith(field, fc.params[0]) case "endswith": - return sqlchemy.Like(field, fmt.Sprintf("%%%s", fc.params[0])) + return sqlchemy.Endswith(field, fc.params[0]) case "equals": return sqlchemy.Equals(field, fc.params[0]) case "notequals": diff --git a/vendor/yunion.io/x/pkg/util/regutils/regutils.go b/vendor/yunion.io/x/pkg/util/regutils/regutils.go index 5ae4abb988..b196f2772f 100644 --- a/vendor/yunion.io/x/pkg/util/regutils/regutils.go +++ b/vendor/yunion.io/x/pkg/util/regutils/regutils.go @@ -1,7 +1,9 @@ package regutils import ( + "net" "regexp" + "strings" ) var FUNCTION_REG *regexp.Regexp @@ -11,9 +13,6 @@ var INTEGER_REG *regexp.Regexp var FLOAT_REG *regexp.Regexp var MACADDR_REG *regexp.Regexp var COMPACT_MACADDR_REG *regexp.Regexp -var IPADDR_REG_PATTERN *regexp.Regexp -var IP6ADDR_REG *regexp.Regexp -var CIDR_REG_PATTERN *regexp.Regexp var NSPTR_REG *regexp.Regexp var NAME_REG *regexp.Regexp var DOMAINNAME_REG *regexp.Regexp @@ -32,6 +31,8 @@ var RFC2882_TIME_REG *regexp.Regexp var EMAIL_REG *regexp.Regexp var CHINA_MOBILE_REG *regexp.Regexp var FS_FORMAT_REG *regexp.Regexp +var US_CURRENCY_REG *regexp.Regexp +var EU_CURRENCY_REG *regexp.Regexp func init() { FUNCTION_REG = regexp.MustCompile(`^\w+\(.*\)$`) @@ -41,9 +42,6 @@ func init() { FLOAT_REG = regexp.MustCompile(`^\d+(\.\d*)?$`) MACADDR_REG = regexp.MustCompile(`^([0-9a-fA-F]{2}:){5}[0-9a-fA-F]{2}$`) COMPACT_MACADDR_REG = regexp.MustCompile(`^[0-9a-fA-F]{12}$`) - IPADDR_REG_PATTERN = regexp.MustCompile(`^\d{1,3}\.\d{1,3}\.\d{1,3}\.\d{1,3}$`) - CIDR_REG_PATTERN = regexp.MustCompile(`^\d{1,3}\.\d{1,3}\.\d{1,3}\.\d{1,3}(\/\d{1,2})?$`) - IP6ADDR_REG = regexp.MustCompile(`^\s*((([0-9A-Fa-f]{1,4}:){7}([0-9A-Fa-f]{1,4}|:))|(([0-9A-Fa-f]{1,4}:){6}(:[0-9A-Fa-f]{1,4}|((25[0-5]|2[0-4][0-9]|1[0-9][0-9]|[1-9]?[0-9])(.(25[0-5]|2[0-4][0-9]|1[0-9][0-9]|[1-9]?[0-9])){3})|:))|(([0-9A-Fa-f]{1,4}:){5}(((:[0-9A-Fa-f]{1,4}){1,2})|:((25[0-5]|2[0-4][0-9]|1[0-9][0-9]|[1-9]?[0-9])(.(25[0-5]|2[0-4][0-9]|1[0-9][0-9]|[1-9]?[0-9])){3})|:))|(([0-9A-Fa-f]{1,4}:){4}(((:[0-9A-Fa-f]{1,4}){1,3})|((:[0-9A-Fa-f]{1,4})?:((25[0-5]|2[0-4][0-9]|1[0-9][0-9]|[1-9]?[0-9])(.(25[0-5]|2[0-4][0-9]|1[0-9][0-9]|[1-9]?[0-9])){3}))|:))|(([0-9A-Fa-f]{1,4}:){3}(((:[0-9A-Fa-f]{1,4}){1,4})|((:[0-9A-Fa-f]{1,4}){0,2}:((25[0-5]|2[0-4][0-9]|1[0-9][0-9]|[1-9]?[0-9])(.(25[0-5]|2[0-4][0-9]|1[0-9][0-9]|[1-9]?[0-9])){3}))|:))|(([0-9A-Fa-f]{1,4}:){2}(((:[0-9A-Fa-f]{1,4}){1,5})|((:[0-9A-Fa-f]{1,4}){0,3}:((25[0-5]|2[0-4][0-9]|1[0-9][0-9]|[1-9]?[0-9])(.(25[0-5]|2[0-4][0-9]|1[0-9][0-9]|[1-9]?[0-9])){3}))|:))|(([0-9A-Fa-f]{1,4}:){1}(((:[0-9A-Fa-f]{1,4}){1,6})|((:[0-9A-Fa-f]{1,4}){0,4}:((25[0-5]|2[0-4][0-9]|1[0-9][0-9]|[1-9]?[0-9])(.(25[0-5]|2[0-4][0-9]|1[0-9][0-9]|[1-9]?[0-9])){3}))|:))|(:(((:[0-9A-Fa-f]{1,4}){1,7})|((:[0-9A-Fa-f]{1,4}){0,5}:((25[0-5]|2[0-4][0-9]|1[0-9][0-9]|[1-9]?[0-9])(.(25[0-5]|2[0-4][0-9]|1[0-9][0-9]|[1-9]?[0-9])){3}))|:)))(%.+)?\s*$`) NSPTR_REG = regexp.MustCompile(`^\d{1,3}\.\d{1,3}\.\d{1,3}\.\d{1,3}\.in-addr\.arpa$`) NAME_REG = regexp.MustCompile(`^[a-zA-Z][a-zA-Z0-9._@-]*$`) DOMAINNAME_REG = regexp.MustCompile(`^[a-zA-Z0-9-.]+$`) @@ -62,6 +60,8 @@ func init() { EMAIL_REG = regexp.MustCompile(`^[A-Za-z0-9._%+-]+@[A-Za-z0-9.-]+\.[A-Za-z]{2,4}$`) CHINA_MOBILE_REG = regexp.MustCompile(`^1[0-9-]{10}$`) FS_FORMAT_REG = regexp.MustCompile(`^(ext|fat|hfs|xfs|swap|ntfs|reiserfs|ufs|btrfs)`) + US_CURRENCY_REG = regexp.MustCompile(`^(\d{0,3}|((\d{1,3},)+\d{3}))(\.\d*)?$`) + EU_CURRENCY_REG = regexp.MustCompile(`^(\d{0,3}|((\d{1,3}\.)+\d{3}))(,\d*)?$`) } func MatchFunction(str string) bool { @@ -93,19 +93,26 @@ func MatchCompactMacAddr(str string) bool { } func MatchIP4Addr(str string) bool { - return IPADDR_REG_PATTERN.MatchString(str) + ip := net.ParseIP(str) + return ip != nil && !strings.Contains(str, ":") } func MatchCIDR(str string) bool { - return CIDR_REG_PATTERN.MatchString(str) + ip, _, err := net.ParseCIDR(str) + if err != nil { + return false + } + return ip != nil && !strings.Contains(str, ":") } func MatchIP6Addr(str string) bool { - return IP6ADDR_REG.MatchString(str) + ip := net.ParseIP(str) + return ip != nil && strings.Contains(str, ":") } func MatchIPAddr(str string) bool { - return MatchIP4Addr(str) || MatchIP6Addr(str) + ip := net.ParseIP(str) + return ip != nil } func MatchPtr(str string) bool { @@ -179,3 +186,11 @@ func MatchMobile(str string) bool { func MatchFS(str string) bool { return FS_FORMAT_REG.MatchString(str) } + +func MatchUSCurrency(str string) bool { + return US_CURRENCY_REG.MatchString(str) +} + +func MatchEUCurrency(str string) bool { + return EU_CURRENCY_REG.MatchString(str) +} diff --git a/vendor/yunion.io/x/pkg/util/secrules/secrules.go b/vendor/yunion.io/x/pkg/util/secrules/secrules.go index 4f249fdcd6..037bb62d45 100644 --- a/vendor/yunion.io/x/pkg/util/secrules/secrules.go +++ b/vendor/yunion.io/x/pkg/util/secrules/secrules.go @@ -7,7 +7,9 @@ import ( "strconv" "strings" + "yunion.io/x/log" "yunion.io/x/pkg/util/regutils" + "yunion.io/x/pkg/utils" ) type TSecurityRuleDirection string @@ -54,13 +56,16 @@ const PROTO_UDP = "udp" const PROTO_ICMP = "icmp" var ( - ErrInvalidDirection = errors.New("invalid direction") - ErrInvalidAction = errors.New("invalid action") - ErrInvalidNet = errors.New("invalid net") - ErrInvalidIPAddr = errors.New("invalid ip address") - ErrInvalidProtocol = errors.New("invalid protocol") - ErrInvalidPortRange = errors.New("invalid port range") - ErrInvalidPort = errors.New("invalid port") + ErrInvalidProtocolAny = errors.New("invalid protocol any with port option") + ErrInvalidProtocolICMP = errors.New("invalid protocol icmp with port option") + ErrInvalidPriority = errors.New("invalid priority") + ErrInvalidDirection = errors.New("invalid direction") + ErrInvalidAction = errors.New("invalid action") + ErrInvalidNet = errors.New("invalid net") + ErrInvalidIPAddr = errors.New("invalid ip address") + ErrInvalidProtocol = errors.New("invalid protocol") + ErrInvalidPortRange = errors.New("invalid port range") + ErrInvalidPort = errors.New("invalid port") ) type SecurityRuleSet []SecurityRule @@ -93,7 +98,7 @@ func parsePortString(ps string) (int, error) { func ParseSecurityRule(pattern string) (*SecurityRule, error) { rule := &SecurityRule{} for _, direction := range []TSecurityRuleDirection{SecurityRuleIngress, SecurityRuleEgress} { - if pattern[:len(direction)+1] == string(direction)+":" { + if len(pattern) > len(direction)+1 && pattern[:len(direction)+1] == string(direction)+":" { rule.Direction, pattern = direction, strings.Replace(pattern, string(direction)+":", "", -1) break } @@ -205,6 +210,76 @@ func (rule *SecurityRule) IsWildMatch() bool { rule.PortEnd == 0 } +func (rule *SecurityRule) ValidateRule() error { + if !utils.IsInStringArray(string(rule.Direction), []string{string(DIR_IN), string(DIR_OUT)}) { + return ErrInvalidDirection + } + if !utils.IsInStringArray(string(rule.Action), []string{string(SecurityRuleAllow), string(SecurityRuleDeny)}) { + return ErrInvalidAction + } + if !utils.IsInStringArray(rule.Protocol, []string{PROTO_ANY, PROTO_ICMP, PROTO_TCP, PROTO_UDP}) { + return ErrInvalidProtocol + } + + if rule.Protocol == PROTO_ICMP { + if len(rule.Ports) > 0 || rule.PortStart > 0 || rule.PortEnd > 0 { + return ErrInvalidProtocolICMP + } + } + + if rule.Protocol == PROTO_ANY { + if len(rule.Ports) > 0 || rule.PortStart > 0 || rule.PortEnd > 0 { + return ErrInvalidProtocolAny + } + } + + if len(rule.Ports) > 0 { + for i := 0; i < len(rule.Ports); i++ { + if rule.Ports[i] < 1 || rule.Ports[i] > 65535 { + return ErrInvalidPort + } + } + } + if rule.PortStart > 0 || rule.PortEnd > 0 { + if rule.PortStart < 1 { + return ErrInvalidPortRange + } + + if rule.PortStart > rule.PortEnd { + return ErrInvalidPortRange + } + if rule.PortStart > 65535 || rule.PortEnd > 65535 { + return ErrInvalidPortRange + } + } + if rule.Priority < 1 || rule.Priority > 100 { + return ErrInvalidPriority + } + return nil +} + +func (rule *SecurityRule) getPort() string { + if rule.PortStart > 0 && rule.PortEnd > 0 { + if rule.PortStart < rule.PortEnd { + return fmt.Sprintf("%d-%d", rule.PortStart, rule.PortEnd) + } + if rule.PortStart == rule.PortEnd { + return fmt.Sprintf("%d", rule.PortStart) + } + // panic on this badness + log.Errorf("invalid port range %d-%d", rule.PortStart, rule.PortEnd) + return "" + } + if len(rule.Ports) > 0 { + ps := []string{} + for _, p := range rule.Ports { + ps = append(ps, fmt.Sprintf("%d", p)) + } + return strings.Join(ps, ",") + } + return "" +} + func (rule *SecurityRule) String() (result string) { s := []string{} s = append(s, string(rule.Direction)+":"+string(rule.Action)) @@ -216,22 +291,12 @@ func (rule *SecurityRule) String() (result string) { s = append(s, rule.IPNet.IP.String()) } } + s = append(s, rule.Protocol) if rule.Protocol == PROTO_TCP || rule.Protocol == PROTO_UDP { - if rule.PortStart > 0 && rule.PortEnd > 0 { - if rule.PortStart < rule.PortEnd { - s = append(s, fmt.Sprintf("%d-%d", rule.PortStart, rule.PortEnd)) - } else if rule.PortStart == rule.PortEnd { - s = append(s, fmt.Sprintf("%d", rule.PortStart)) - } else { - // panic on this badness - } - } else if len(rule.Ports) > 0 { - ps := []string{} - for _, p := range rule.Ports { - ps = append(ps, fmt.Sprintf("%d", p)) - } - s = append(s, strings.Join(ps, ",")) + port := rule.getPort() + if len(port) > 0 { + s = append(s, port) } } return strings.Join(s, " ")