From 3de4afe4fc9e7f7330134ba7958c8e09d95ae8b8 Mon Sep 17 00:00:00 2001 From: Yousong Zhou Date: Thu, 25 Oct 2018 09:37:43 +0000 Subject: [PATCH] dnsrecords: move name,value checks to a single place --- pkg/compute/models/dnsrecords.go | 75 ++++++++++++++++++++------------ 1 file changed, 47 insertions(+), 28 deletions(-) diff --git a/pkg/compute/models/dnsrecords.go b/pkg/compute/models/dnsrecords.go index fea9752bfb..3038e1089b 100644 --- a/pkg/compute/models/dnsrecords.go +++ b/pkg/compute/models/dnsrecords.go @@ -64,8 +64,8 @@ func (man *SDnsRecordManager) ParseInputInfo(data *jsonutils.JSONDict) ([]string if err != nil { return nil, err } - if (typ == "A" && !regutils.MatchIP4Addr(addr)) || (typ == "AAAA" && !regutils.MatchIP6Addr(addr)) { - return nil, httperrors.NewNotAcceptableError("Invalid type %s address: %s", typ, addr) + if err := man.checkRecordValue(typ, addr); err != nil { + return nil, err } records = append(records, fmt.Sprintf("%s:%s", typ, addr)) } @@ -82,12 +82,8 @@ func (man *SDnsRecordManager) ParseInputInfo(data *jsonutils.JSONDict) ([]string return "", httperrors.NewNotAcceptableError("SRV: insufficient param: %s", s) } host := parts[0] - if !regutils.MatchDomainName(host) && - !regutils.MatchIPAddr(host) { - return "", httperrors.NewNotAcceptableError("SRV: target is not valid domain: %s", host) - } - if regutils.MatchIPAddr(host) { - return "", httperrors.NewNotAcceptableError("SRV: target cannot be an IP address: %s", host) + if err := man.checkRecordValue("SRV", host); err != nil { + return "", err } port, err := strconv.Atoi(parts[1]) if err != nil || port <= 0 || port >= 65536 { @@ -162,10 +158,8 @@ func (man *SDnsRecordManager) ParseInputInfo(data *jsonutils.JSONDict) ([]string } if cname, err := data.GetString("CNAME"); err != nil { return nil, err - } else if !regutils.MatchDomainName(cname) { - return nil, httperrors.NewNotAcceptableError("CNAME: record value must be valid domain name: %s", cname) - } else if regutils.MatchIPAddr(cname) { - return nil, httperrors.NewNotAcceptableError("CNAME: record value cannot be ip address: %s", cname) + } else if err := man.checkRecordValue("CNAME", cname); err != nil { + return nil, err } else { records = []string{fmt.Sprintf("%s:%s", "CNAME", cname)} } @@ -179,8 +173,8 @@ func (man *SDnsRecordManager) ParseInputInfo(data *jsonutils.JSONDict) ([]string if err != nil { return nil, err } - if !regutils.MatchPtr(name) { - return nil, httperrors.NewNotAcceptableError(fmt.Sprintf("Invalid ptr %s", name)) + if err := man.checkRecordName("PTR", name); err != nil { + return nil, err } } domainName, err := data.GetString("PTR") @@ -188,11 +182,8 @@ func (man *SDnsRecordManager) ParseInputInfo(data *jsonutils.JSONDict) ([]string if err != nil { return nil, err } - if !regutils.MatchDomainName(domainName) { - return nil, httperrors.NewNotAcceptableError("ptr: record value must be valid domain name: %s", domainName) - } - if regutils.MatchIPAddr(domainName) { - return nil, httperrors.NewNotAcceptableError("ptr: record value cannot be ip address: %s", domainName) + if err := man.checkRecordValue("PTR", domainName); err != nil { + return nil, err } } records = []string{fmt.Sprintf("%s:%s", "PTR", domainName)} @@ -216,24 +207,52 @@ func (man *SDnsRecordManager) GetRecordsType(recs []string) string { return "" } -func (man *SDnsRecordManager) CheckNameForDnsType(name, recType string) error { - if regutils.MatchIPAddr(name) { - return httperrors.NewNotAcceptableError("domain name cannot be ip address: %s", name) - } - switch recType { +func (man *SDnsRecordManager) checkRecordName(typ, name string) error { + switch typ { case "A", "CNAME": if !regutils.MatchDomainName(name) { - return httperrors.NewNotAcceptableError("Invalid domain name %s for type %s", name, recType) + return httperrors.NewNotAcceptableError("%s: invalid domain name: %s", typ, name) } case "SRV": if !regutils.MatchDomainSRV(name) { - return httperrors.NewNotAcceptableError("Invalid SRV name %s for type %s", name, recType) + return httperrors.NewNotAcceptableError("SRV: invalid srv record name: %s", typ, name) } case "PTR": if !regutils.MatchPtr(name) { - return httperrors.NewNotAcceptableError("Invalid ptr name %s", name) + return httperrors.NewNotAcceptableError("PTR: invalid ptr record name: %s", typ, name) } } + if regutils.MatchIPAddr(name) { + return httperrors.NewNotAcceptableError("%s: name cannot be ip address: %s", typ, name) + } + return nil +} + +func (man *SDnsRecordManager) checkRecordValue(typ, val string) error { + switch typ { + case "A": + if !regutils.MatchIP4Addr(val) { + return httperrors.NewNotAcceptableError("A: record value must be ipv4 address: %s", val) + } + case "AAAA": + if !regutils.MatchIP6Addr(val) { + return httperrors.NewNotAcceptableError("AAAA: record value must be ipv6 address: %s", val) + } + case "CNAME", "PTR", "SRV": + fieldMsg := "record value" + if typ == "SRV" { + fieldMsg = "target" + } + if !regutils.MatchDomainName(val) { + return httperrors.NewNotAcceptableError("%s: %s must be domain name: %s", typ, fieldMsg, val) + } + if regutils.MatchIPAddr(val) { + return httperrors.NewNotAcceptableError("%s: %s cannot be ip address: %s", typ, fieldMsg, val) + } + default: + // internal error + return httperrors.NewNotAcceptableError("%s: unknown record type", typ) + } return nil } @@ -256,7 +275,7 @@ func (man *SDnsRecordManager) validateModelData( if err != nil { return nil, err } - err = man.CheckNameForDnsType(name, recType) + err = man.checkRecordName(recType, name) if err != nil { return nil, err }