dnsrecords: move name,value checks to a single place

This commit is contained in:
Yousong Zhou
2018-10-25 09:37:43 +00:00
parent 0031b533b7
commit 3de4afe4fc
+47 -28
View File
@@ -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
}