From c3c86c36e58c2721a2b519bc60f3f7d96da3f3d3 Mon Sep 17 00:00:00 2001 From: TangBin Date: Mon, 21 Jan 2019 20:27:52 +0800 Subject: [PATCH] update aws secrule sync --- Gopkg.lock | 16 +- pkg/util/aws/utils.go | 29 ++ pkg/util/aws/vpc.go | 1 + vendor/yunion.io/x/jsonutils/marshal.go | 164 +++-------- vendor/yunion.io/x/jsonutils/unmarshal.go | 10 +- vendor/yunion.io/x/pkg/gotypes/gotypes.go | 68 ++--- .../yunion.io/x/pkg/util/netutils/netutils.go | 104 +++++++ .../x/pkg/util/reflectutils/jsonfield.go | 177 ++++++++++++ .../x/pkg/util/reflectutils/reflectutils.go | 124 +++++---- vendor/yunion.io/x/pkg/util/secrules/cut.go | 260 ++++++++++++++++++ vendor/yunion.io/x/pkg/util/secrules/port.go | 217 +++++++++++++++ .../yunion.io/x/pkg/util/secrules/secrules.go | 147 ++++++---- .../x/pkg/util/secrules/secruleset.go | 181 ++++++++++++ .../x/pkg/util/stringutils/stringutils.go | 14 + vendor/yunion.io/x/sqlchemy/column.go | 32 +-- vendor/yunion.io/x/sqlchemy/const.go | 1 - vendor/yunion.io/x/sqlchemy/field_update.go | 50 ++-- vendor/yunion.io/x/sqlchemy/insert.go | 20 +- vendor/yunion.io/x/sqlchemy/parser.go | 36 ++- vendor/yunion.io/x/sqlchemy/query.go | 4 +- vendor/yunion.io/x/sqlchemy/table.go | 5 +- vendor/yunion.io/x/sqlchemy/update.go | 36 +-- vendor/yunion.io/x/structarg/structarg.go | 16 +- 23 files changed, 1312 insertions(+), 400 deletions(-) create mode 100644 vendor/yunion.io/x/pkg/util/reflectutils/jsonfield.go create mode 100644 vendor/yunion.io/x/pkg/util/secrules/cut.go create mode 100644 vendor/yunion.io/x/pkg/util/secrules/port.go create mode 100644 vendor/yunion.io/x/pkg/util/secrules/secruleset.go diff --git a/Gopkg.lock b/Gopkg.lock index 1907e6d158..d1b3f2c613 100644 --- a/Gopkg.lock +++ b/Gopkg.lock @@ -1348,11 +1348,11 @@ [[projects]] branch = "master" - digest = "1:9c30ecaff00e904c298f8bf113f52942dc0067385eaebaec62f09b96301eb8b2" + digest = "1:9a92b34083e218c123d21a4b0d74228df6b186de9dcd6489e5b4cf68f4e90104" name = "yunion.io/x/jsonutils" packages = ["."] pruneopts = "UT" - revision = "77c73104d4d728207b7f389a5f4ef480e97c2d44" + revision = "c40afc81cccc4618168f722c64deba7cfbb1f8bd" [[projects]] branch = "master" @@ -1367,7 +1367,7 @@ [[projects]] branch = "master" - digest = "1:3415345332af51147fa83d11c5e1ec32bcdf5f153a947891d5cd19ef98ea93ae" + digest = "1:f7e6b859716c65473233390b3c69c919d0d6c2d9bd25ee599c29d0e8e3419d35" name = "yunion.io/x/pkg" packages = [ "gotypes", @@ -1400,23 +1400,23 @@ "utils", ] pruneopts = "UT" - revision = "6918b664c6710b0736ad1c7272f6a24ff15df327" + revision = "4d8dfcd77e8817d5dc37eacf4b6b098429c7d259" [[projects]] branch = "master" - digest = "1:42e452d5282f9b5e95d8806f601dba1e9846401c0fb1c328f2fd787195fe40e0" + digest = "1:0e3c8da76b7b7ba0f67ef0737d361e117b65bd328bfbcb28fb030268df587ca9" name = "yunion.io/x/sqlchemy" packages = ["."] pruneopts = "UT" - revision = "6fdedc07ce571d9789f21bc63d35479a5c2df717" + revision = "0b1ca973f3e5140dee7e773fc51cfc8c07ce0f69" [[projects]] branch = "master" - digest = "1:428cc9d6e84fe526b606d57ca4c18aa0433e310ab25c6c1106607c0c64c9c4c2" + digest = "1:6ea9ac8f317e7dd79309bb1633642bc9d1822011c5175bd241552534d216ac6d" name = "yunion.io/x/structarg" packages = ["."] pruneopts = "UT" - revision = "28def21ba4844dc7a92bfc8833a3f9af157a6c52" + revision = "e4f0f5201b4acad185ebfb76e65905fe5ae7c1b1" [solve-meta] analyzer-name = "dep" diff --git a/pkg/util/aws/utils.go b/pkg/util/aws/utils.go index bf52f0174a..45e842107f 100644 --- a/pkg/util/aws/utils.go +++ b/pkg/util/aws/utils.go @@ -5,6 +5,7 @@ import ( "net" "reflect" "regexp" + "sort" "strings" "github.com/aws/aws-sdk-go/service/ec2" @@ -157,6 +158,34 @@ func IntVal(s *int64) int64 { return 0 } +// SecurityRuleSet to allow list +// 将安全组规则全部转换为等价的allow规则 +func SecurityRuleSetToAllowSet(srs secrules.SecurityRuleSet) secrules.SecurityRuleSet { + inRuleSet := secrules.SecurityRuleSet{} + outRuleSet := secrules.SecurityRuleSet{} + + for _, rule := range srs { + if rule.Direction == secrules.SecurityRuleIngress { + inRuleSet = append(inRuleSet, rule) + } + + if rule.Direction == secrules.SecurityRuleEgress { + outRuleSet = append(outRuleSet, rule) + } + } + + sort.Sort(inRuleSet) + sort.Sort(outRuleSet) + + inRuleSet = inRuleSet.AllowList() + outRuleSet = outRuleSet.AllowList() + + ret := secrules.SecurityRuleSet{} + ret = append(ret, inRuleSet...) + ret = append(ret, outRuleSet...) + return ret +} + func isAwsPermissionAllPorts(p ec2.IpPermission) bool { if p.FromPort == nil || p.ToPort == nil { return false diff --git a/pkg/util/aws/vpc.go b/pkg/util/aws/vpc.go index e4d18ca8ab..c12dbc2500 100644 --- a/pkg/util/aws/vpc.go +++ b/pkg/util/aws/vpc.go @@ -146,6 +146,7 @@ func (self *SRegion) SyncSecurityGroup(secgroupId string, vpcId string, name str secgroupId = fmt.Sprintf("%s-%s", vpcId, secgroupId) } + rules = SecurityRuleSetToAllowSet(rules) if secgroup, err := self.getSecurityGroupById(vpcId, secgroupId); err != nil { if len(desc) == 0 { desc = fmt.Sprintf("security group %s for vpc %s", name, vpcId) diff --git a/vendor/yunion.io/x/jsonutils/marshal.go b/vendor/yunion.io/x/jsonutils/marshal.go index f1e60f0388..db0d942cd8 100644 --- a/vendor/yunion.io/x/jsonutils/marshal.go +++ b/vendor/yunion.io/x/jsonutils/marshal.go @@ -12,16 +12,15 @@ import ( "reflect" "time" - "strings" "yunion.io/x/log" "yunion.io/x/pkg/gotypes" "yunion.io/x/pkg/tristate" + "yunion.io/x/pkg/util/reflectutils" "yunion.io/x/pkg/util/timeutils" - "yunion.io/x/pkg/utils" ) -func marshalSlice(val reflect.Value, info *jsonMarshalInfo) JSONObject { - if val.Len() == 0 && info != nil && info.omitEmpty { +func marshalSlice(val reflect.Value, info *reflectutils.SStructFieldInfo) JSONObject { + if val.Len() == 0 && info != nil && info.OmitEmpty { return JSONNull } objs := make([]JSONObject, val.Len()) @@ -29,16 +28,16 @@ func marshalSlice(val reflect.Value, info *jsonMarshalInfo) JSONObject { objs[i] = marshalValue(val.Index(i), nil) } arr := NewArray(objs...) - if info != nil && info.forceString { + if info != nil && info.ForceString { return NewString(arr.String()) } else { return arr } } -func marshalMap(val reflect.Value, info *jsonMarshalInfo) JSONObject { +func marshalMap(val reflect.Value, info *reflectutils.SStructFieldInfo) JSONObject { keys := val.MapKeys() - if len(keys) == 0 && info != nil && info.omitEmpty { + if len(keys) == 0 && info != nil && info.OmitEmpty { return JSONNull } objPairs := make([]JSONPair, 0) @@ -50,125 +49,36 @@ func marshalMap(val reflect.Value, info *jsonMarshalInfo) JSONObject { } } dict := NewDict(objPairs...) - if info != nil && info.forceString { + if info != nil && info.ForceString { return NewString(dict.String()) } else { return dict } } -func marshalStruct(val reflect.Value, info *jsonMarshalInfo) JSONObject { +func marshalStruct(val reflect.Value, info *reflectutils.SStructFieldInfo) JSONObject { objPairs := struct2JSONPairs(val) - if len(objPairs) == 0 && info != nil && info.omitEmpty { + if len(objPairs) == 0 && info != nil && info.OmitEmpty { return JSONNull } dict := NewDict(objPairs...) - if info != nil && info.forceString { + if info != nil && info.ForceString { return NewString(dict.String()) } else { return dict } } -type jsonMarshalInfo struct { - ignore bool - omitEmpty bool - omitFalse bool - omitZero bool - name string - forceString bool -} - -func parseJsonMarshalInfo(fieldTag reflect.StructTag) jsonMarshalInfo { - info := jsonMarshalInfo{} - info.omitEmpty = true - info.omitZero = false - info.omitFalse = false - - tags := utils.TagMap(fieldTag) - if val, ok := tags["json"]; ok { - keys := strings.Split(val, ",") - if len(keys) > 0 { - if keys[0] == "-" { - if len(keys) > 1 { - info.name = keys[0] - } else { - info.ignore = true - } - } else { - info.name = keys[0] - } - } - if len(keys) > 1 { - for _, k := range keys[1:] { - switch k { - case "omitempty": - info.omitEmpty = true - case "allowempty": - info.omitEmpty = false - case "omitzero": - info.omitZero = true - case "allowzero": - info.omitZero = false - case "omitfalse": - info.omitFalse = true - case "allowfalse": - info.omitFalse = false - case "string": - info.forceString = true - } - } - } - } - if val, ok := tags["name"]; ok { - info.name = val - } - return info -} - func struct2JSONPairs(val reflect.Value) []JSONPair { - structType := val.Type() objPairs := make([]JSONPair, 0) - for i := 0; i < structType.NumField(); i += 1 { - sf := structType.Field(i) - - // ignore unexported field altogether - if !gotypes.IsFieldExportable(sf.Name) { + fields := reflectutils.FetchStructFieldValueSet(val) + for i := 0; i < len(fields); i += 1 { + jsonInfo := fields[i].Info + if jsonInfo.Ignore { continue } - - if sf.Anonymous { - fv := val.Field(i) - - // T, *T - switch fv.Kind() { - case reflect.Ptr, reflect.Interface: - // ignore nil values completely - if !fv.IsValid() || fv.IsNil() { - continue - } - fv = fv.Elem() - } - // note that we regard anonymous interface field the - // same as with anonymous struct field. This is - // different from how encoding/json handles struct - // field of interface type. - if fv.Kind() == reflect.Struct { - newPairs := struct2JSONPairs(fv) - objPairs = append(objPairs, newPairs...) - continue - } - } - - jsonInfo := parseJsonMarshalInfo(sf.Tag) - if jsonInfo.ignore { - continue - } - key := jsonInfo.name - if len(key) == 0 { - key = utils.CamelSplit(sf.Name, "_") - } - val := marshalValue(val.Field(i), &jsonInfo) + key := jsonInfo.MarshalName() + val := marshalValue(fields[i].Value, &jsonInfo) if val != nil && val != JSONNull { objPair := JSONPair{key: key, val: val} objPairs = append(objPairs, objPair) @@ -177,30 +87,30 @@ func struct2JSONPairs(val reflect.Value) []JSONPair { return objPairs } -func marshalInt64(val int64, info *jsonMarshalInfo) JSONObject { - if val == 0 && info != nil && info.omitZero { +func marshalInt64(val int64, info *reflectutils.SStructFieldInfo) JSONObject { + if val == 0 && info != nil && info.OmitZero { return JSONNull - } else if info != nil && info.forceString { + } else if info != nil && info.ForceString { return NewString(fmt.Sprintf("%d", val)) } else { return NewInt(val) } } -func marshalFloat64(val float64, info *jsonMarshalInfo) JSONObject { - if val == 0.0 && info != nil && info.omitZero { +func marshalFloat64(val float64, info *reflectutils.SStructFieldInfo) JSONObject { + if val == 0.0 && info != nil && info.OmitZero { return JSONNull - } else if info != nil && info.forceString { + } else if info != nil && info.ForceString { return NewString(fmt.Sprintf("%f", val)) } else { return NewFloat(val) } } -func marshalBoolean(val bool, info *jsonMarshalInfo) JSONObject { - if !val && info != nil && info.omitFalse { +func marshalBoolean(val bool, info *reflectutils.SStructFieldInfo) JSONObject { + if !val && info != nil && info.OmitFalse { return JSONNull - } else if info != nil && info.forceString { + } else if info != nil && info.ForceString { return NewString(fmt.Sprintf("%v", val)) } else { if val { @@ -211,7 +121,7 @@ func marshalBoolean(val bool, info *jsonMarshalInfo) JSONObject { } } -func marshalTristate(val tristate.TriState, info *jsonMarshalInfo) JSONObject { +func marshalTristate(val tristate.TriState, info *reflectutils.SStructFieldInfo) JSONObject { if val.IsTrue() { return JSONTrue } else if val.IsFalse() { @@ -221,17 +131,17 @@ func marshalTristate(val tristate.TriState, info *jsonMarshalInfo) JSONObject { } } -func marshalString(val string, info *jsonMarshalInfo) JSONObject { - if len(val) == 0 && info != nil && info.omitEmpty { +func marshalString(val string, info *reflectutils.SStructFieldInfo) JSONObject { + if len(val) == 0 && info != nil && info.OmitEmpty { return JSONNull } else { return NewString(val) } } -func marshalTime(val time.Time, info *jsonMarshalInfo) JSONObject { +func marshalTime(val time.Time, info *reflectutils.SStructFieldInfo) JSONObject { if val.IsZero() { - if info != nil && info.omitEmpty { + if info != nil && info.OmitEmpty { return JSONNull } return NewString("") @@ -248,7 +158,7 @@ func Marshal(obj interface{}) JSONObject { return marshalValue(objValue, nil) } -func marshalValue(objValue reflect.Value, info *jsonMarshalInfo) JSONObject { +func marshalValue(objValue reflect.Value, info *reflectutils.SStructFieldInfo) JSONObject { switch objValue.Type() { case JSONDictPtrType, JSONArrayPtrType, JSONBoolPtrType, JSONIntPtrType, JSONFloatPtrType, JSONStringPtrType, JSONObjectType: if objValue.IsNil() { @@ -258,7 +168,7 @@ func marshalValue(objValue reflect.Value, info *jsonMarshalInfo) JSONObject { case JSONDictType: json, ok := objValue.Interface().(JSONDict) if ok { - if len(json.data) == 0 && info != nil && info.omitEmpty { + if len(json.data) == 0 && info != nil && info.OmitEmpty { return JSONNull } else { return &json @@ -269,7 +179,7 @@ func marshalValue(objValue reflect.Value, info *jsonMarshalInfo) JSONObject { case JSONArrayType: json, ok := objValue.Interface().(JSONArray) if ok { - if len(json.data) == 0 && info != nil && info.omitEmpty { + if len(json.data) == 0 && info != nil && info.OmitEmpty { return JSONNull } else { return &json @@ -280,7 +190,7 @@ func marshalValue(objValue reflect.Value, info *jsonMarshalInfo) JSONObject { case JSONBoolType: json, ok := objValue.Interface().(JSONBool) if ok { - if !json.data && info != nil && info.omitEmpty { + if !json.data && info != nil && info.OmitEmpty { return JSONNull } else { return &json @@ -291,7 +201,7 @@ func marshalValue(objValue reflect.Value, info *jsonMarshalInfo) JSONObject { case JSONIntType: json, ok := objValue.Interface().(JSONInt) if ok { - if json.data == 0 && info != nil && info.omitEmpty { + if json.data == 0 && info != nil && info.OmitEmpty { return JSONNull } else { return &json @@ -302,7 +212,7 @@ func marshalValue(objValue reflect.Value, info *jsonMarshalInfo) JSONObject { case JSONFloatType: json, ok := objValue.Interface().(JSONFloat) if ok { - if json.data == 0.0 && info != nil && info.omitEmpty { + if json.data == 0.0 && info != nil && info.OmitEmpty { return JSONNull } else { return &json @@ -313,7 +223,7 @@ func marshalValue(objValue reflect.Value, info *jsonMarshalInfo) JSONObject { case JSONStringType: json, ok := objValue.Interface().(JSONString) if ok { - if len(json.data) == 0 && info != nil && info.omitEmpty { + if len(json.data) == 0 && info != nil && info.OmitEmpty { return JSONNull } else { return &json diff --git a/vendor/yunion.io/x/jsonutils/unmarshal.go b/vendor/yunion.io/x/jsonutils/unmarshal.go index 9a1015f3da..4023cec956 100644 --- a/vendor/yunion.io/x/jsonutils/unmarshal.go +++ b/vendor/yunion.io/x/jsonutils/unmarshal.go @@ -459,14 +459,10 @@ func (this *JSONDict) unmarshalMap(val reflect.Value) error { } func (this *JSONDict) unmarshalStruct(val reflect.Value) error { - fieldValues := reflectutils.FetchStructFieldNameValues(val) + fieldValues := reflectutils.FetchStructFieldValueSetForWrite(val) for k, v := range this.data { - fieldValue, ok := fieldValues[k] // first try original key - if !ok { // try kebab - k = utils.CamelSplit(k, "_") - fieldValue, ok = fieldValues[k] - } - if ok { + fieldValue, find := fieldValues.GetValue(k) + if find { err := v.unmarshalValue(fieldValue) if err != nil { log.Debugf("unmarshalStruct field %s error %s", k, err) diff --git a/vendor/yunion.io/x/pkg/gotypes/gotypes.go b/vendor/yunion.io/x/pkg/gotypes/gotypes.go index 615499fa4a..e510878a3e 100644 --- a/vendor/yunion.io/x/pkg/gotypes/gotypes.go +++ b/vendor/yunion.io/x/pkg/gotypes/gotypes.go @@ -8,6 +8,7 @@ import ( "time" "yunion.io/x/pkg/util/timeutils" + "yunion.io/x/pkg/utils" ) const ( @@ -129,6 +130,17 @@ func ParseValue(val string, tp reflect.Type) (reflect.Value, error) { rvv := reflect.New(tpElem) rvv.Elem().Set(rv) return rvv, nil + case reflect.Slice, reflect.Array: + values := utils.FindWords([]byte(val), 0) + sliceVal := reflect.MakeSlice(reflect.SliceOf(tp.Elem()), len(values), len(values)) + for i, vv := range values { + vvv, err := ParseValue(vv, tp.Elem()) + if err != nil { + return sliceVal, fmt.Errorf("Cannot parse %s to %s", vv, tp.Elem()) + } + sliceVal.Index(i).Set(vvv) + } + return sliceVal, nil default: if tp == TimeType { tm, e := timeutils.ParseTimeStr(val) @@ -143,55 +155,15 @@ func SetValue(value reflect.Value, valStr string) error { if !value.CanSet() { return fmt.Errorf("Value is not settable") } - switch value.Type() { - case BoolType: - val_bool, e := strconv.ParseBool(valStr) - if e != nil { - return e - } - value.SetBool(val_bool) - case IntType, Int8Type, Int16Type, Int32Type, Int64Type: - val_int, e := strconv.ParseInt(valStr, 10, 64) - if e != nil { - return e - } - value.SetInt(val_int) - case UintType, Uint8Type, Uint16Type, Uint32Type, Uint64Type: - val_uint, e := strconv.ParseUint(valStr, 10, 64) - if e != nil { - return e - } - value.SetUint(val_uint) - case Float32Type, Float64Type: - val_float, e := strconv.ParseFloat(valStr, 64) - if e != nil { - return e - } - value.SetFloat(val_float) - case StringType: - value.SetString(valStr) - case TimeType: - tm, e := timeutils.ParseTimeStr(valStr) - if e != nil { - return e - } - value.Set(reflect.ValueOf(tm)) - case BoolSliceType, IntSliceType, Int8SliceType, Int16SliceType, - Int32SliceType, Int64SliceType, UintSliceType, Uint8SliceType, - Uint16SliceType, Uint32SliceType, Uint64SliceType, - Float32SliceType, Float64SliceType, StringSliceType: - reflect.Append(value, reflect.ValueOf(valStr)) + parseValue, err := ParseValue(valStr, value.Type()) + if err != nil { + return err + } + switch value.Kind() { + case reflect.Slice, reflect.Array: + value.Set(reflect.AppendSlice(value, parseValue)) default: - if value.Kind() == reflect.Ptr && value.Elem().Kind() != reflect.Slice { - newVal := reflect.New(value.Type().Elem()) - newValElem := newVal.Elem() - if err := SetValue(newValElem, valStr); err != nil { - return err - } - value.Set(newVal) - } else { - return fmt.Errorf("Unsupported type: %v", value.Type()) - } + value.Set(parseValue) } return nil } diff --git a/vendor/yunion.io/x/pkg/util/netutils/netutils.go b/vendor/yunion.io/x/pkg/util/netutils/netutils.go index 4d36184266..e7e30273d7 100644 --- a/vendor/yunion.io/x/pkg/util/netutils/netutils.go +++ b/vendor/yunion.io/x/pkg/util/netutils/netutils.go @@ -3,8 +3,10 @@ package netutils import ( "fmt" "math/rand" + "net" "strconv" "strings" + "yunion.io/x/pkg/util/regutils" ) @@ -148,6 +150,15 @@ func NewIPV4AddrRange(ip1 IPV4Addr, ip2 IPV4Addr) IPV4AddrRange { } } +// n.IP and n.Mask must be ipv4 type. n.Mask must be canonical +func NewIPV4AddrRangeFromIPNet(n *net.IPNet) IPV4AddrRange { + pref, err := NewIPV4Prefix(n.String()) + if err != nil { + panic("unexpected IPNet: " + n.String()) + } + return pref.ToIPRange() +} + func (ar IPV4AddrRange) Contains(ip IPV4Addr) bool { return (ip >= ar.start) && (ip <= ar.end) } @@ -184,6 +195,99 @@ func (ar IPV4AddrRange) IsOverlap(ar2 IPV4AddrRange) bool { } } +func (ar IPV4AddrRange) ToIPNets() []*net.IPNet { + r := []*net.IPNet{} + mms := ar.ToMaskMatches() + for _, mm := range mms { + a := mm[0] + m := mm[1] + addr := net.IPv4(byte((a>>24)&0xff), byte((a>>16)&0xff), byte((a>>8)&0xff), byte(a&0xff)) + mask := net.IPv4Mask(byte((m>>24)&0xff), byte((m>>16)&0xff), byte((m>>8)&0xff), byte(m&0xff)) + r = append(r, &net.IPNet{ + IP: addr, + Mask: mask, + }) + } + return r +} + +func (ar IPV4AddrRange) ToMaskMatches() [][2]uint32 { + r := [][2]uint32{} + s := uint32(ar.start) + e := uint32(ar.end) + if s == e { + r = append(r, [2]uint32{s, ^uint32(0)}) + return r + } + sp, ep := uint64(s), uint64(e) + ep = ep + 1 + for sp < ep { + b := uint64(1) + for (sp+b) <= ep && (sp&(b-1)) == 0 { + b <<= 1 + } + b >>= 1 + r = append(r, [2]uint32{uint32(sp), uint32(^(b - 1))}) + sp = sp + b + } + return r +} + +func (ar IPV4AddrRange) Substract(ar2 IPV4AddrRange) (lefts []IPV4AddrRange, sub *IPV4AddrRange) { + lefts = []IPV4AddrRange{} + // no intersection, no substract + if ar.end < ar2.start || ar.start > ar2.end { + lefts = append(lefts, ar) + return + } + + // ar contains ar2 + if ar.ContainsRange(ar2) { + nns := [][2]int64{ + [2]int64{int64(ar.start), int64(ar2.start) - 1}, + [2]int64{int64(ar2.end) + 1, int64(ar.end)}, + } + for _, nn := range nns { + if nn[0] <= nn[1] { + lefts = append(lefts, NewIPV4AddrRange(IPV4Addr(nn[0]), IPV4Addr(nn[1]))) + } + } + ar2_ := ar2 + sub = &ar2_ + return + } + + // ar contained by ar2 + if ar2.ContainsRange(ar) { + ar_ := ar + sub = &ar_ + return + } + + // intersect, ar on the left + if ar.start < ar2.start && ar.end >= ar2.start { + lefts = append(lefts, NewIPV4AddrRange(ar.start, ar2.start-1)) + sub_ := NewIPV4AddrRange(ar2.start, ar.end) + sub = &sub_ + return + } + + // intersect, ar on the right + if ar.start <= ar2.end && ar.end > ar2.end { + lefts = append(lefts, NewIPV4AddrRange(ar2.end+1, ar.end)) + sub_ := NewIPV4AddrRange(ar.start, ar2.end) + sub = &sub_ + return + } + + // no intersection + return +} + +func (ar IPV4AddrRange) equals(ar2 IPV4AddrRange) bool { + return ar.start == ar2.start && ar.end == ar2.end +} + func Masklen2Mask(maskLen int8) IPV4Addr { var mask uint32 = 0 for i := 0; i < int(maskLen); i += 1 { diff --git a/vendor/yunion.io/x/pkg/util/reflectutils/jsonfield.go b/vendor/yunion.io/x/pkg/util/reflectutils/jsonfield.go new file mode 100644 index 0000000000..44e8390ab7 --- /dev/null +++ b/vendor/yunion.io/x/pkg/util/reflectutils/jsonfield.go @@ -0,0 +1,177 @@ +package reflectutils + +import ( + "reflect" + "strings" + + "yunion.io/x/pkg/gotypes" + "yunion.io/x/pkg/utils" +) + +type SStructFieldInfo struct { + Ignore bool + OmitEmpty bool + OmitFalse bool + OmitZero bool + Name string + FieldName string + ForceString bool + Tags map[string]string +} + +func ParseStructFieldJsonInfo(sf reflect.StructField) SStructFieldInfo { + info := SStructFieldInfo{} + info.FieldName = sf.Name + info.OmitEmpty = true + info.OmitZero = false + info.OmitFalse = false + + info.Tags = utils.TagMap(sf.Tag) + if val, ok := info.Tags["json"]; ok { + keys := strings.Split(val, ",") + if len(keys) > 0 { + if keys[0] == "-" { + if len(keys) > 1 { + info.Name = keys[0] + } else { + info.Ignore = true + } + } else { + info.Name = keys[0] + } + } + if len(keys) > 1 { + for _, k := range keys[1:] { + switch strings.ToLower(k) { + case "omitempty": + info.OmitEmpty = true + case "allowempty": + info.OmitEmpty = false + case "omitzero": + info.OmitZero = true + case "allowzero": + info.OmitZero = false + case "omitfalse": + info.OmitFalse = true + case "allowfalse": + info.OmitFalse = false + case "string": + info.ForceString = true + } + } + } + } + if val, ok := info.Tags["name"]; ok { + info.Name = val + } + return info +} + +func (info *SStructFieldInfo) MarshalName() string { + if len(info.Name) > 0 { + return info.Name + } + return utils.CamelSplit(info.FieldName, "_") +} + +type SStructFieldValue struct { + Info SStructFieldInfo + Value reflect.Value +} + +type SStructFieldValueSet []SStructFieldValue + +func FetchStructFieldValueSet(dataValue reflect.Value) SStructFieldValueSet { + return fetchStructFieldValueSet(dataValue, false) +} + +func FetchStructFieldValueSetForWrite(dataValue reflect.Value) SStructFieldValueSet { + return fetchStructFieldValueSet(dataValue, true) +} + +func fetchStructFieldValueSet(dataValue reflect.Value, allocatePtr bool) SStructFieldValueSet { + fields := SStructFieldValueSet{} + dataType := dataValue.Type() + for i := 0; i < dataType.NumField(); i += 1 { + sf := dataType.Field(i) + + // ignore unexported field altogether + if !gotypes.IsFieldExportable(sf.Name) { + continue + } + fv := dataValue.Field(i) + if !fv.IsValid() { + continue + } + if sf.Anonymous { + // T, *T + switch fv.Kind() { + case reflect.Ptr, reflect.Interface: + if !fv.IsValid() { + continue + } + if fv.IsNil() { + if fv.Kind() == reflect.Ptr && allocatePtr { + fv.Set(reflect.New(fv.Type().Elem())) + } else { + continue + } + } + fv = fv.Elem() + } + // note that we regard anonymous interface field the + // same as with anonymous struct field. This is + // different from how encoding/json handles struct + // field of interface type. + if fv.Kind() == reflect.Struct && sf.Type != gotypes.TimeType { + subfields := fetchStructFieldValueSet(fv, allocatePtr) + fields = append(fields, subfields...) + continue + } + } + jsonInfo := ParseStructFieldJsonInfo(sf) + fields = append(fields, SStructFieldValue{ + Info: jsonInfo, + Value: fv, + }) + } + return fields +} + +func (set SStructFieldValueSet) GetStructFieldIndex(name string) int { + for i := 0; i < len(set); i += 1 { + jsonInfo := set[i].Info + if jsonInfo.MarshalName() == name { + return i + } + if utils.CamelSplit(jsonInfo.FieldName, "_") == utils.CamelSplit(name, "_") { + return i + } + if jsonInfo.FieldName == name { + return i + } + if jsonInfo.FieldName == utils.Capitalize(name) { + return i + } + } + return -1 +} + +func (set SStructFieldValueSet) GetValue(name string) (reflect.Value, bool) { + idx := set.GetStructFieldIndex(name) + if idx < 0 { + return reflect.Value{}, false + } + return set[idx].Value, true +} + +func (set SStructFieldValueSet) GetInterface(name string) (interface{}, bool) { + idx := set.GetStructFieldIndex(name) + if idx < 0 { + return nil, false + } + if set[idx].Value.CanInterface() { + return set[idx].Value.Interface(), true + } + return nil, false +} diff --git a/vendor/yunion.io/x/pkg/util/reflectutils/reflectutils.go b/vendor/yunion.io/x/pkg/util/reflectutils/reflectutils.go index 55d9019d8e..d33f9438e0 100644 --- a/vendor/yunion.io/x/pkg/util/reflectutils/reflectutils.go +++ b/vendor/yunion.io/x/pkg/util/reflectutils/reflectutils.go @@ -1,13 +1,13 @@ package reflectutils import ( + "fmt" "reflect" - "strings" - "yunion.io/x/pkg/gotypes" - "yunion.io/x/pkg/utils" + "yunion.io/x/log" ) +/* func GetStructFieldName(field *reflect.StructField) string { tagMap := utils.TagMap(field.Tag) // var name string @@ -80,49 +80,20 @@ func fetchStructFieldNameValues(dataType reflect.Type, dataValue reflect.Value, } } } +*/ func FindStructFieldValue(dataValue reflect.Value, name string) (reflect.Value, bool) { - dataType := dataValue.Type() - for i := 0; i < dataType.NumField(); i += 1 { - fieldType := dataType.Field(i) - if gotypes.IsFieldExportable(fieldType.Name) { - fieldValue := dataValue.Field(i) - if fieldType.Type.Kind() == reflect.Struct && fieldType.Type != gotypes.TimeType { - val, find := FindStructFieldValue(fieldValue, name) - if find { - return val, find - } - } else if fieldValue.CanSet() { - fName := GetStructFieldName(&fieldType) - if fName == name { - return fieldValue, true - } - } - } + set := FetchStructFieldValueSet(dataValue) + val, find := set.GetValue(name) + if find && val.CanSet() { + return val, true } return reflect.Value{}, false } func FindStructFieldInterface(dataValue reflect.Value, name string) (interface{}, bool) { - dataType := dataValue.Type() - for i := 0; i < dataType.NumField(); i += 1 { - fieldType := dataType.Field(i) - if gotypes.IsFieldExportable(fieldType.Name) { - fieldValue := dataValue.Field(i) - if fieldType.Type.Kind() == reflect.Struct && fieldType.Type != gotypes.TimeType { - val, find := FindStructFieldInterface(fieldValue, name) - if find { - return val, find - } - } else if fieldValue.CanInterface() { - fName := GetStructFieldName(&fieldType) - if fName == name { - return fieldValue.Interface(), true - } - } - } - } - return nil, false + set := FetchStructFieldValueSet(dataValue) + return set.GetInterface(name) } func FillEmbededStructValue(container reflect.Value, embed reflect.Value) bool { @@ -147,19 +118,16 @@ func FillEmbededStructValue(container reflect.Value, embed reflect.Value) bool { } func SetStructFieldValue(structValue reflect.Value, fieldName string, val reflect.Value) bool { - dataType := structValue.Type() - for i := 0; i < dataType.NumField(); i += 1 { - fieldType := dataType.Field(i) - if gotypes.IsFieldExportable(fieldType.Name) { - fName := GetStructFieldName(&fieldType) - if fName == fieldName { - fieldValue := structValue.Field(i) - fieldValue.Set(val) - return true - } - } + set := FetchStructFieldValueSet(structValue) + target, find := set.GetValue(fieldName) + if !find { + return false } - return false + if !target.CanSet() { + return false + } + target.Set(val) + return true } func ExpandInterface(val interface{}) []interface{} { @@ -174,3 +142,57 @@ func ExpandInterface(val interface{}) []interface{} { return []interface{}{val} } } + +// tagetType must not be a pointer +func getAnonymouStructPointer(structValue reflect.Value, targetType reflect.Type) interface{} { + structType := structValue.Type() + for i := 0; i < structValue.NumField(); i += 1 { + fieldType := structType.Field(i) + if fieldType.Type == targetType { + val := structValue.Field(i) // val is not a pointer + return val.Addr().Interface() + } + if fieldType.Anonymous && fieldType.Type.Kind() == reflect.Struct { + ptr := getAnonymouStructPointer(structValue.Field(i), targetType) + if ptr != nil { + return ptr + } + } + } + return nil +} + +func FindAnonymouStructPointer(data interface{}, targetPtr interface{}) error { + targetValue := reflect.ValueOf(targetPtr).Elem() + if targetValue.Kind() != reflect.Ptr { + return fmt.Errorf("target must be a pointer to pointer") + } + targetType := targetValue.Type().Elem() + structValue := reflect.Indirect(reflect.ValueOf(data)) + ptr := getAnonymouStructPointer(structValue, targetType) + if ptr == nil { + return fmt.Errorf("no anonymous struct found") + } + targetValue.Set(reflect.ValueOf(ptr)) + return nil +} + +func StructContains(type1 reflect.Type, type2 reflect.Type) bool { + if type1.Kind() != reflect.Struct || type2.Kind() != reflect.Struct { + log.Errorf("types should be struct!") + return false + } + if type1 == type2 { + return true + } + for i := 0; i < type1.NumField(); i += 1 { + field := type1.Field(i) + if field.Anonymous && field.Type.Kind() == reflect.Struct { + contains := StructContains(field.Type, type2) + if contains { + return true + } + } + } + return false +} diff --git a/vendor/yunion.io/x/pkg/util/secrules/cut.go b/vendor/yunion.io/x/pkg/util/secrules/cut.go new file mode 100644 index 0000000000..3217d07800 --- /dev/null +++ b/vendor/yunion.io/x/pkg/util/secrules/cut.go @@ -0,0 +1,260 @@ +package secrules + +import ( + "bytes" + "fmt" + "net" + "sort" + + "yunion.io/x/pkg/util/netutils" +) + +type securityRuleCut struct { + r SecurityRule + protocolCut bool + netCut bool + portCut bool +} + +func (src *securityRuleCut) String() string { + s := fmt.Sprintf("[%s;protocolCut=%v;netCut=%v;portCut=%v]", + src.r.String(), src.protocolCut, src.netCut, src.portCut) + return s +} + +func (src *securityRuleCut) isCut() bool { + return src.protocolCut && src.netCut && src.portCut +} + +type securityRuleCuts []securityRuleCut + +func newSecurityRuleSetCuts(srs SecurityRuleSet) securityRuleCuts { + srcs := make(securityRuleCuts, len(srs)) + for i := range srcs { + srcs[i].r = srs[i] + } + return srcs +} + +func (srcs securityRuleCuts) String() string { + buf := bytes.Buffer{} + for i := range srcs { + s := srcs[i].String() + buf.WriteString(s) + buf.WriteString("\n") + } + return buf.String() +} + +func (srcs securityRuleCuts) securityRuleSet() SecurityRuleSet { + srs := SecurityRuleSet{} + for i := range srcs { + src := &srcs[i] + if src.isCut() { + continue + } + srs = append(srs, src.r) + } + return srs +} + +func (srcs securityRuleCuts) cutOutProtocol(protocol string) securityRuleCuts { + r := securityRuleCuts{} + for _, src := range srcs { + sr := src.r + if sr.Protocol == protocol { + // cut + src.protocolCut = true + r = append(r, src) + } else if sr.Protocol == PROTO_ANY { + for _, p := range protocolsSupported { + src_ := src + src_.r.Protocol = p + if p == protocol { + src_.protocolCut = true + } + r = append(r, src_) + } + } else if protocol == PROTO_ANY { + // cut + src.protocolCut = true + r = append(r, src) + } else { + // retain + r = append(r, src) + } + } + return r +} + +func (srcs securityRuleCuts) cutOutIPNet(n *net.IPNet) securityRuleCuts { + r := securityRuleCuts{} + ar2 := netutils.NewIPV4AddrRangeFromIPNet(n) + for _, src := range srcs { + sr := src.r + ar := netutils.NewIPV4AddrRangeFromIPNet(sr.IPNet) + left, subs := ar.Substract(ar2) + for _, l := range left { + // retain + nets := l.ToIPNets() + for _, net_ := range nets { + src_ := src + src_.r.IPNet = net_ + r = append(r, src_) + } + } + if subs != nil { + // cut + nets := subs.ToIPNets() + for _, net_ := range nets { + src_ := src + src_.r.IPNet = net_ + src_.netCut = true + r = append(r, src_) + } + } + } + return r +} + +func (srcs securityRuleCuts) cutOutPortRange(portStart, portEnd uint16) securityRuleCuts { + pr1 := &portRange{ + start: portStart, + end: portEnd, + } + r := securityRuleCuts{} + for _, src := range srcs { + sr := src.r + if len(sr.Ports) > 0 { + ps := newPortsFromInts(sr.Ports...) + left, sub := ps.substractPortRange(pr1) + if len(left) > 0 { + src_ := src + src_.r.Ports = left.IntSlice() + r = append(r, src_) + } + if len(sub) > 0 { + src_ := src + src_.r.Ports = left.IntSlice() + src_.portCut = true + r = append(r, src_) + } + } else if sr.PortStart > 0 && sr.PortEnd > 0 { + pr := newPortRange(uint16(sr.PortStart), uint16(sr.PortEnd)) + left, sub := pr.substractPortRange(pr1) + for _, l := range left { + src_ := src + src_.r.PortStart = int(l.start) + src_.r.PortEnd = int(l.end) + r = append(r, src_) + } + if sub != nil && sub.count() > 0 { + src_ := src + src_.r.PortStart = int(sub.start) + src_.r.PortEnd = int(sub.end) + src_.portCut = true + r = append(r, src_) + } + } else { + { + nns := [][2]int32{ + [2]int32{1, int32(portStart) - 1}, + [2]int32{int32(portEnd) + 1, 65535}, + } + for _, nn := range nns { + if nn[0] <= nn[1] { + src_ := src + src_.r.PortStart = int(nn[0]) + src_.r.PortEnd = int(nn[1]) + r = append(r, src_) + } + } + } + { + src_ := src + src_.r.PortStart = int(portStart) + src_.r.PortEnd = int(portEnd) + src_.portCut = true + r = append(r, src_) + } + } + } + return r +} + +func (srcs securityRuleCuts) cutOutPorts(ps1 []uint16) securityRuleCuts { + r := securityRuleCuts{} + for _, src := range srcs { + sr := src.r + if len(sr.Ports) > 0 { + ps0 := newPortsFromInts(sr.Ports...) + left, sub := ps0.substractPorts(ps1) + if len(left) > 0 { + src_ := src + src_.r.Ports = left.IntSlice() + r = append(r, src_) + } + if len(sub) > 0 { + src_ := src + src_.r.Ports = sub.IntSlice() + src_.portCut = true + r = append(r, src_) + } + } else if sr.PortStart > 0 && sr.PortEnd > 0 { + pr := newPortRange(uint16(sr.PortStart), uint16(sr.PortEnd)) + ps := ports(ps1) + left, sub := pr.substractPorts(ps) + for _, l := range left { + src_ := src + src_.r.PortStart = int(l.start) + src_.r.PortEnd = int(l.end) + r = append(r, src_) + } + if len(sub) > 0 { + src_ := src + src_.r.Ports = sub.IntSlice() + src_.r.PortStart = 0 + src_.r.PortEnd = 0 + src_.portCut = true + r = append(r, src_) + } + } else { + sort.Slice(ps1, func(i, j int) bool { + return ps1[i] < ps1[j] + }) + add := func(s, e uint16) { + src_ := src + src_.r.PortStart = int(s) + src_.r.PortEnd = int(e) + r = append(r, src_) + } + s := uint16(1) + for _, p := range ps1 { + if s <= p-1 { + add(s, p-1) + s = p + 1 + } + } + if s != 0 && s <= 65535 { + add(s, 65535) + } + { + src_ := src + src_.r.Ports = ports(ps1).IntSlice() + src_.portCut = true + r = append(r, src_) + } + } + } + return r +} + +func (srcs securityRuleCuts) cutOutPortsAll() securityRuleCuts { + r := securityRuleCuts{} + for _, src := range srcs { + src_ := src + src_.portCut = true + r = append(r, src_) + } + return r +} diff --git a/vendor/yunion.io/x/pkg/util/secrules/port.go b/vendor/yunion.io/x/pkg/util/secrules/port.go new file mode 100644 index 0000000000..f7fea85eed --- /dev/null +++ b/vendor/yunion.io/x/pkg/util/secrules/port.go @@ -0,0 +1,217 @@ +package secrules + +import ( + "sort" +) + +type ports []uint16 + +func newPortsFromInts(ints ...int) ports { + ps := make(ports, len(ints)) + for i := range ints { + // panic on invalid value + ps[i] = uint16(ints[i]) + } + return ps +} + +func (ps ports) IntSlice() []int { + ints := make([]int, len(ps)) + for i := range ps { + ints[i] = int(ps[i]) + } + return ints +} + +func (ps ports) Len() int { + return len(ps) +} + +func (ps ports) Swap(i, j int) { + ps[i], ps[j] = ps[j], ps[i] +} + +func (ps ports) Less(i, j int) bool { + return ps[i] < ps[j] +} + +func (ps ports) contains(p uint16) bool { + for _, p0 := range ps { + if p0 == p { + return true + } + } + return false +} + +func (ps ports) containsPorts(ps1 ports) bool { + for _, p := range ps1 { + if !ps.contains(p) { + return false + } + } + return true +} + +func (ps ports) dedup() ports { + ps1 := make(ports, len(ps)) + copy(ps1, ps) + sort.Sort(ps1) + for i := len(ps1) - 1; i > 0; i-- { + if ps1[i] == ps1[i-1] { + ps1 = append(ps1[:i], ps1[i+1:]...) + } + } + return ps1 +} + +func (ps ports) sameAs(ps1 ports) bool { + if len(ps) != len(ps1) { + return false + } + for i := range ps { + if ps[i] != ps1[i] { + return false + } + } + return true +} + +func (ps ports) substractPortRange(pr *portRange) (left, subs ports) { + left = ports{} + subs = ports{} + for _, p := range ps { + if pr.contains(p) { + subs = append(subs, p) + } else { + left = append(left, p) + } + } + return +} + +func (ps ports) substractPorts(ps1 ports) (left, subs ports) { + left = ports{} + subs = ports{} + for _, p0 := range ps { + if ps1.contains(p0) { + subs = append(subs, p0) + } else { + left = append(left, p0) + } + } + return +} + +type portRange struct { + start uint16 + end uint16 +} + +func newPortRange(s, e uint16) *portRange { + // panic on s > e + return &portRange{ + start: s, + end: e, + } +} + +func (pr *portRange) equals(pr1 *portRange) bool { + return pr.start == pr1.start && pr.end == pr1.end +} + +func (pr *portRange) contains(p uint16) bool { + return p >= pr.start && p <= pr.end +} + +func (pr *portRange) containsRange(pr1 *portRange) bool { + return pr.start <= pr1.start && pr.end >= pr1.end +} + +func (pr *portRange) count() uint16 { + return pr.end - pr.start + 1 +} + +func (pr *portRange) substractPortRange(pr1 *portRange) (lefts []*portRange, sub *portRange) { + // no intersection, no substract + if pr.end < pr1.start || pr.start > pr1.end { + l := *pr + lefts = []*portRange{&l} + return + } + + // pr contains pr1 + if pr.containsRange(pr1) { + nns := [][2]int32{ + [2]int32{int32(pr.start), int32(pr1.start) - 1}, + [2]int32{int32(pr1.end) + 1, int32(pr.end)}, + } + lefts = []*portRange{} + for _, nn := range nns { + if nn[0] <= nn[1] { + lefts = append(lefts, &portRange{ + start: uint16(nn[0]), + end: uint16(nn[1]), + }) + } + } + s := *pr1 + sub = &s + return + } + + // pr contained by pr1 + if pr1.containsRange(pr) { + s := *pr + sub = &s + return + } + + // intersect, pr on the left + if pr.start < pr1.start && pr.end >= pr1.start { + lefts = []*portRange{&portRange{start: pr.start, end: pr1.start - 1}} + sub = &portRange{pr1.start, pr.end} + return + } + + // intersect, pr on the right + if pr.start <= pr1.end && pr.end > pr1.end { + lefts = []*portRange{&portRange{pr1.end + 1, pr.end}} + sub = &portRange{pr.start, pr1.end} + return + } + + // no intersection + return +} + +func (pr *portRange) substractPorts(ps1 ports) (lefts []*portRange, subs ports) { + // no duplicate + // then ordered + ps2 := make(ports, len(ps1)) + copy(ps2, ps1) + sort.Sort(ps2) + + lefts = []*portRange{} + subs = ports{} + s := pr.start + for _, p := range ps2 { + if pr.contains(p) { + if p > s { + lefts = append(lefts, &portRange{ + start: s, + end: p - 1, + }) + } + subs = append(subs, p) + s = p + 1 + } + } + if s != 0 && s <= pr.end { + lefts = append(lefts, &portRange{ + start: s, + end: pr.end, + }) + } + return +} diff --git a/vendor/yunion.io/x/pkg/util/secrules/secrules.go b/vendor/yunion.io/x/pkg/util/secrules/secrules.go index 6437ff1ae9..7770029a1e 100644 --- a/vendor/yunion.io/x/pkg/util/secrules/secrules.go +++ b/vendor/yunion.io/x/pkg/util/secrules/secrules.go @@ -55,6 +55,13 @@ const PROTO_TCP = "tcp" const PROTO_UDP = "udp" const PROTO_ICMP = "icmp" +// non-wild protocols +var protocolsSupported = []string{ + PROTO_TCP, + PROTO_UDP, + PROTO_ICMP, +} + var ( ErrInvalidProtocolAny = errors.New("invalid protocol any with port option") ErrInvalidProtocolICMP = errors.New("invalid protocol icmp with port option") @@ -68,25 +75,6 @@ var ( ErrInvalidPort = errors.New("invalid port") ) -type SecurityRuleSet []SecurityRule - -func (v SecurityRuleSet) Len() int { - return len(v) -} - -func (v SecurityRuleSet) Swap(i, j int) { - v[i], v[j] = v[j], v[i] -} - -func (v SecurityRuleSet) Less(i, j int) bool { - if v[i].Priority > v[j].Priority { - return true - } else if v[i].Priority == v[j].Priority { - return strings.Compare(v[i].String(), v[j].String()) <= 0 - } - return false -} - func parsePortString(ps string) (int, error) { p, err := strconv.ParseUint(ps, 10, 16) if err != nil || p == 0 { @@ -95,6 +83,15 @@ func parsePortString(ps string) (int, error) { return int(p), nil } +func MustParseSecurityRule(s string) *SecurityRule { + r, err := ParseSecurityRule(s) + if err != nil { + msg := fmt.Sprintf("parse security rule %q: %v", s, err) + panic(msg) + } + return r +} + func ParseSecurityRule(pattern string) (*SecurityRule, error) { rule := &SecurityRule{} for _, direction := range []TSecurityRuleDirection{SecurityRuleIngress, SecurityRuleEgress} { @@ -152,42 +149,11 @@ func ParseSecurityRule(pattern string) (*SecurityRule, error) { } rule.Protocol = seg } else if status == SEG_PORT { - if len(seg) == 0 { - status = SEG_END - } else if idx := strings.Index(seg, "-"); idx > -1 { - segs := strings.SplitN(seg, "-", 2) - var ps, pe int - var err error - if ps, err = parsePortString(segs[0]); err != nil { - return nil, ErrInvalidPortRange - } - if pe, err = parsePortString(segs[1]); err != nil { - return nil, ErrInvalidPortRange - } - if ps > pe { - ps, pe = pe, ps - } - rule.PortStart = ps - rule.PortEnd = pe - } else if idx := strings.Index(seg, ","); idx > -1 { - ports := make([]int, 0) - segs := strings.Split(seg, ",") - for _, seg := range segs { - p, err := parsePortString(seg) - if err != nil { - return nil, err - } - ports = append(ports, p) - } - rule.Ports = ports - } else { - p, err := parsePortString(seg) - if err != nil { - return nil, err - } - rule.PortStart, rule.PortEnd = p, p - } status = SEG_END + if err := rule.ParsePorts(seg); err != nil { + return nil, err + } + return rule, nil } } return rule, nil @@ -201,6 +167,48 @@ func (rule *SecurityRule) IsWildMatch() bool { rule.PortEnd == 0 } +func (rule *SecurityRule) ParsePorts(seg string) error { + if len(seg) == 0 { + rule.Ports = []int{} + rule.PortStart = -1 + rule.PortEnd = -1 + return nil + } else if idx := strings.Index(seg, "-"); idx > -1 { + segs := strings.SplitN(seg, "-", 2) + var ps, pe int + var err error + if ps, err = parsePortString(segs[0]); err != nil { + return ErrInvalidPortRange + } + if pe, err = parsePortString(segs[1]); err != nil { + return ErrInvalidPortRange + } + if ps > pe { + ps, pe = pe, ps + } + rule.PortStart = ps + rule.PortEnd = pe + } else if idx := strings.Index(seg, ","); idx > -1 { + ports := make([]int, 0) + segs := strings.Split(seg, ",") + for _, seg := range segs { + p, err := parsePortString(seg) + if err != nil { + return err + } + ports = append(ports, p) + } + rule.Ports = ports + } else { + p, err := parsePortString(seg) + if err != nil { + return err + } + rule.PortStart, rule.PortEnd = p, p + } + return nil +} + func (rule *SecurityRule) ValidateRule() error { if !utils.IsInStringArray(string(rule.Direction), []string{string(DIR_IN), string(DIR_OUT)}) { return ErrInvalidDirection @@ -292,3 +300,34 @@ func (rule *SecurityRule) String() (result string) { } return strings.Join(s, " ") } + +func (rule *SecurityRule) equals(r *SecurityRule) bool { + // essence of String, bom + s0 := rule.String() + s1 := r.String() + return s0 == s1 +} + +func (rule *SecurityRule) netEquals(r *SecurityRule) bool { + net0 := rule.IPNet.String() + net1 := r.IPNet.String() + return net0 == net1 +} + +func (rule *SecurityRule) cutOut(r *SecurityRule) SecurityRuleSet { + srcs := securityRuleCuts{securityRuleCut{r: *rule}} + //a := srcs + srcs = srcs.cutOutProtocol(r.Protocol) + srcs = srcs.cutOutIPNet(r.IPNet) + if len(r.Ports) > 0 { + srcs = srcs.cutOutPorts([]uint16(newPortsFromInts(r.Ports...))) + } else if r.PortStart > 0 && r.PortEnd > 0 { + srcs = srcs.cutOutPortRange(uint16(r.PortStart), uint16(r.PortEnd)) + } else { + srcs = srcs.cutOutPortsAll() + } + //fmt.Printf("a %s\n", a) + //fmt.Printf("b %s\n", srcs) + srs := srcs.securityRuleSet() + return srs +} diff --git a/vendor/yunion.io/x/pkg/util/secrules/secruleset.go b/vendor/yunion.io/x/pkg/util/secrules/secruleset.go new file mode 100644 index 0000000000..7ff2370e15 --- /dev/null +++ b/vendor/yunion.io/x/pkg/util/secrules/secruleset.go @@ -0,0 +1,181 @@ +package secrules + +import ( + "bytes" + "sort" +) + +type SecurityRuleSet []SecurityRule + +func (srs SecurityRuleSet) Len() int { + return len(srs) +} + +func (srs SecurityRuleSet) Swap(i, j int) { + srs[i], srs[j] = srs[j], srs[i] +} + +func (srs SecurityRuleSet) Less(i, j int) bool { + if srs[i].Priority > srs[j].Priority { + return true + } else if srs[i].Priority == srs[j].Priority { + return srs[i].String() < srs[j].String() + } + return false +} + +func (srs SecurityRuleSet) stringList() []string { + r := make([]string, len(srs)) + for i := range srs { + r = append(r, srs[i].String()) + } + return r +} + +func (srs SecurityRuleSet) String() string { + buf := bytes.Buffer{} + for i := range srs { + buf.WriteString(srs[i].String()) + buf.WriteString(";") + } + return buf.String() +} + +func (srs SecurityRuleSet) equals(srs1 SecurityRuleSet) bool { + if len(srs) != len(srs1) { + return false + } + for i := range srs { + if !srs[i].equals(&srs1[i]) { + return false + } + } + return true +} + +// convert to pure allow list +// +// requirements on srs +// +// - ordered by priority +// - same direction +// +func (srs SecurityRuleSet) AllowList() SecurityRuleSet { + r := SecurityRuleSet{} + wq := make(SecurityRuleSet, len(srs)) + copy(wq, srs) + + for len(wq) > 0 { + sr := wq[0] + if sr.Action == SecurityRuleAllow { + r = append(r, sr) + wq = wq[1:] + continue + } + wq = wq.cutOutFirst() + } + r = r.collapse() + return r +} + +func (srs SecurityRuleSet) cutOutFirst() SecurityRuleSet { + r := SecurityRuleSet{} + if len(srs) == 0 { + return r + } + sr := &srs[0] + srs_ := srs[1:] + + for _, sr_ := range srs_ { + if sr.Action == sr_.Action { + r = append(r, sr_) + continue + } + cut := sr_.cutOut(sr) + r = append(r, cut...) + } + return r +} + +// collapse result of AllowList +// +// - same direction +// - same action +// +// As they share the same action, priority's influence on order of rules can be ignored +// +func (srs SecurityRuleSet) collapse() SecurityRuleSet { + srs1 := make(SecurityRuleSet, len(srs)) + copy(srs1, srs) + for i := range srs1 { + sr := &srs1[i] + if len(sr.Ports) > 0 { + sort.Slice(sr.Ports, func(i, j int) bool { + return sr.Ports[i] < sr.Ports[j] + }) + } + } + sort.Slice(srs1, func(i, j int) bool { + sr0 := &srs1[i] + sr1 := &srs1[j] + if sr0.Protocol != sr1.Protocol { + return sr0.Protocol < sr1.Protocol + } + net0 := sr0.IPNet.String() + net1 := sr1.IPNet.String() + if net0 != net1 { + return net0 < net1 + } + if sr0.PortStart > 0 && sr0.PortEnd > 0 { + if sr1.PortStart > 0 && sr1.PortEnd > 0 { + return sr0.PortStart < sr1.PortStart + } + // port range comes first + return true + } else if len(sr0.Ports) > 0 { + if sr1.PortStart > 0 && sr1.PortEnd > 0 { + return false + } else if len(sr1.Ports) > 0 { + sr0l := len(sr0.Ports) + sr1l := len(sr1.Ports) + for i := 0; i < sr0l && i < sr1l; i++ { + if sr0.Ports[i] != sr1.Ports[i] { + return sr0.Ports[i] < sr1.Ports[i] + } + } + return sr0l < sr1l + } + } + return sr0.Priority < sr1.Priority + }) + // merge ports + for i := len(srs1) - 1; i > 0; i-- { + sr0 := &srs1[i-1] + sr1 := &srs1[i] + if sr0.Protocol != sr1.Protocol { + continue + } + if !sr0.netEquals(sr1) { + continue + } + if len(sr0.Ports) > 0 && len(sr1.Ports) > 0 { + ps := newPortsFromInts(sr0.Ports...) + ps = append(ps, newPortsFromInts(sr1.Ports...)...) + ps = ps.dedup() + sr0.Ports = ps.IntSlice() + srs1 = append(srs1[:i], srs1[i+1:]...) + } else if sr0.PortStart > 0 && sr1.PortStart > 0 && sr0.PortEnd > 0 && sr1.PortEnd > 0 { + if sr0.PortEnd == sr1.PortStart-1 { + sr0.PortEnd = sr1.PortEnd + srs1 = append(srs1[:i], srs1[i+1:]...) + } else if sr0.PortStart-1 == sr1.PortEnd { + sr0.PortStart = sr1.PortStart + srs1 = append(srs1[:i], srs1[i+1:]...) + } else if sr0.PortStart == sr1.PortStart && sr0.PortEnd == sr1.PortEnd { + srs1 = append(srs1[:i], srs1[i+1:]...) + } + // save that contains, intersects + } + } + return srs1 +} diff --git a/vendor/yunion.io/x/pkg/util/stringutils/stringutils.go b/vendor/yunion.io/x/pkg/util/stringutils/stringutils.go index bc75956156..038f54e7d0 100644 --- a/vendor/yunion.io/x/pkg/util/stringutils/stringutils.go +++ b/vendor/yunion.io/x/pkg/util/stringutils/stringutils.go @@ -57,3 +57,17 @@ func Interface2String(val interface{}) string { return json.String() } } + +func SplitKeyValue(line string) (string, string) { + return SplitKeyValueBySep(line, ":") +} + +func SplitKeyValueBySep(line string, sep string) (string, string) { + pos := strings.Index(line, sep) + if pos > 0 { + key := strings.TrimSpace(line[:pos]) + val := strings.TrimSpace(line[pos+1:]) + return key, val + } + return "", "" +} diff --git a/vendor/yunion.io/x/sqlchemy/column.go b/vendor/yunion.io/x/sqlchemy/column.go index e2d367f12a..277412a07a 100644 --- a/vendor/yunion.io/x/sqlchemy/column.go +++ b/vendor/yunion.io/x/sqlchemy/column.go @@ -21,7 +21,6 @@ type IColumnSpec interface { IsSupportDefault() bool IsNullable() bool IsPrimary() bool - IsKeyIndex() bool IsUnique() bool IsIndex() bool ExtraDefs() string @@ -50,7 +49,6 @@ type SBaseColumn struct { isPointer bool isNullable bool isPrimary bool - isKeyIndex bool isUnique bool isIndex bool tags map[string]string @@ -92,10 +90,6 @@ func (c *SBaseColumn) IsPrimary() bool { return c.isPrimary } -func (c *SBaseColumn) IsKeyIndex() bool { - return c.isKeyIndex -} - func (c *SBaseColumn) IsUnique() bool { return c.isUnique } @@ -201,11 +195,6 @@ func NewBaseColumn(name string, sqltype string, tagmap map[string]string) SBaseC if ok { isPrimary = utils.ToBool(val) } - isKeyIndex := false - tagmap, val, ok = utils.TagPop(tagmap, TAG_KEY_INDEX) - if ok { - isKeyIndex = utils.ToBool(val) - } isUnique := false tagmap, val, ok = utils.TagPop(tagmap, TAG_UNIQUE) if ok { @@ -226,7 +215,6 @@ func NewBaseColumn(name string, sqltype string, tagmap map[string]string) SBaseC defaultString: defStr, isNullable: isNullable, isPrimary: isPrimary, - isKeyIndex: isKeyIndex, isUnique: isUnique, isIndex: isIndex, tags: tagmap, @@ -275,10 +263,22 @@ func (c *SBooleanColumn) ConvertFromString(str string) string { } func (c *SBooleanColumn) ConvertFromValue(val interface{}) interface{} { - bVal := val.(bool) - if bVal { - return 1 - } else { + switch bVal := val.(type) { + case bool: + if bVal { + return 1 + } else { + return 0 + } + case *bool: + if gotypes.IsNil(bVal) { + return 0 + } else if *bVal { + return 1 + } else { + return 0 + } + default: return 0 } } diff --git a/vendor/yunion.io/x/sqlchemy/const.go b/vendor/yunion.io/x/sqlchemy/const.go index e39de066c9..3bf5293677 100644 --- a/vendor/yunion.io/x/sqlchemy/const.go +++ b/vendor/yunion.io/x/sqlchemy/const.go @@ -31,5 +31,4 @@ const ( TAG_AUTOVERSION = "auto_version" TAG_UPDATE_TIMESTAMP = "updated_at" TAG_CREATE_TIMESTAMP = "created_at" - TAG_KEY_INDEX = "key_index" ) diff --git a/vendor/yunion.io/x/sqlchemy/field_update.go b/vendor/yunion.io/x/sqlchemy/field_update.go index 517bccbaae..28df69bc15 100644 --- a/vendor/yunion.io/x/sqlchemy/field_update.go +++ b/vendor/yunion.io/x/sqlchemy/field_update.go @@ -2,15 +2,16 @@ package sqlchemy import ( "bytes" - "errors" + // "errors" "fmt" "reflect" "yunion.io/x/log" - "yunion.io/x/pkg/gotypes" + // "yunion.io/x/pkg/gotypes" "yunion.io/x/pkg/util/reflectutils" ) +/* func (ts *STableSpec) GetUpdateColumnValue(dataType reflect.Type, dataValue reflect.Value, cv map[string]interface{}, fields map[string]interface{}) error { for i := 0; i < dataType.NumField(); i++ { fieldType := dataType.Field(i) @@ -35,6 +36,7 @@ func (ts *STableSpec) GetUpdateColumnValue(dataType reflect.Type, dataValue refl } return nil } +*/ func (ts *STableSpec) UpdateFields(dt interface{}, fields map[string]interface{}) error { return ts.updateFields(dt, fields, false) @@ -45,37 +47,30 @@ func (ts *STableSpec) UpdateFields(dt interface{}, fields map[string]interface{} // find fields correlatively columns // joint sql and executed func (ts *STableSpec) updateFields(dt interface{}, fields map[string]interface{}, debug bool) error { - dataValue := reflect.ValueOf(dt) - if dataValue.Kind() == reflect.Ptr { - dataValue = dataValue.Elem() - } + dataValue := reflect.Indirect(reflect.ValueOf(dt)) - //cv: {"column name": "update value"} - cv := make(map[string]interface{}, 0) - dataType := dataValue.Type() - ts.GetUpdateColumnValue(dataType, dataValue, cv, fields) - if len(cv) == 0 { - log.Infof("Nothing update") - return nil - } + // cv: {"column name": "update value"} + cv := make(map[string]interface{}) + // dataType := dataValue.Type() + // ts.GetUpdateColumnValue(dataType, dataValue, cv, fields) + // if len(cv) == 0 { + // log.Infof("Nothing update") + // return nil + // } - fullFields := reflectutils.FetchStructFieldNameValueInterfaces(dataValue) + fullFields := reflectutils.FetchStructFieldValueSet(dataValue) versionFields := make([]string, 0) updatedFields := make([]string, 0) primaryCols := make(map[string]interface{}, 0) - indexCols := make(map[string]interface{}, 0) for _, col := range ts.Columns() { name := col.Name() - colValue, ok := fullFields[name] + colValue, ok := fullFields.GetInterface(name) if !ok { continue } if col.IsPrimary() && !col.IsZero(colValue) { primaryCols[name] = colValue continue - } else if col.IsKeyIndex() && !col.IsZero(colValue) { - indexCols[name] = colValue - continue } intCol, ok := col.(*SIntegerColumn) if ok && intCol.IsAutoVersion { @@ -87,8 +82,8 @@ func (ts *STableSpec) updateFields(dt interface{}, fields map[string]interface{} updatedFields = append(updatedFields, name) continue } - if _, exist := cv[name]; exist { - cv[name] = col.ConvertFromValue(cv[name]) + if _, exist := fields[name]; exist { + cv[name] = col.ConvertFromValue(fields[name]) } } @@ -113,16 +108,11 @@ func (ts *STableSpec) updateFields(dt interface{}, fields map[string]interface{} } buf.WriteString(" WHERE ") first = true - var indexFilter map[string]interface{} - if len(primaryCols) > 0 { - indexFilter = primaryCols - } else if len(indexCols) > 0 { - indexFilter = indexCols - } else { - return fmt.Errorf("neither primary key nor key indexes empty???") + if len(primaryCols) == 0 { + return fmt.Errorf("primary key empty???") } - for k, v := range indexFilter { + for k, v := range primaryCols { if first { first = false } else { diff --git a/vendor/yunion.io/x/sqlchemy/insert.go b/vendor/yunion.io/x/sqlchemy/insert.go index e19d141a42..43aa7780d4 100644 --- a/vendor/yunion.io/x/sqlchemy/insert.go +++ b/vendor/yunion.io/x/sqlchemy/insert.go @@ -6,6 +6,7 @@ import ( "strings" "yunion.io/x/log" + "yunion.io/x/pkg/gotypes" "yunion.io/x/pkg/util/reflectutils" ) @@ -13,7 +14,7 @@ func (t *STableSpec) Insert(dt interface{}) error { return t.insert(dt, false) } -func (t *STableSpec) insertSqlPrep(dataFields map[string]interface{}) (string, []interface{}, error) { +func (t *STableSpec) insertSqlPrep(dataFields reflectutils.SStructFieldValueSet) (string, []interface{}, error) { var autoIncField string createdAtFields := make([]string, 0) @@ -31,18 +32,22 @@ func (t *STableSpec) insertSqlPrep(dataFields map[string]interface{}) (string, [ k := c.Name() dtc, ok := c.(*SDateTimeColumn) - ov := dataFields[k] + ov, find := dataFields.GetInterface(k) + + if !find { + continue + } if ok && (dtc.IsCreatedAt || dtc.IsUpdatedAt) { createdAtFields = append(createdAtFields, k) names = append(names, fmt.Sprintf("`%s`", k)) format = append(format, "UTC_TIMESTAMP()") - } else if c.IsSupportDefault() && len(c.Default()) > 0 && ov != nil && c.IsZero(ov) { // empty text value + } else if c.IsSupportDefault() && len(c.Default()) > 0 && !gotypes.IsNil(ov) && c.IsZero(ov) { // empty text value val := c.ConvertFromString(c.Default()) values = append(values, val) names = append(names, fmt.Sprintf("`%s`", k)) format = append(format, "?") - } else if ov != nil && (!c.IsZero(ov) || (!c.IsPointer() && !c.IsText())) && !isAutoInc { + } else if !gotypes.IsNil(ov) && (!c.IsZero(ov) || (!c.IsPointer() && !c.IsText())) && !isAutoInc { v := c.ConvertFromValue(ov) values = append(values, v) names = append(names, fmt.Sprintf("`%s`", k)) @@ -73,7 +78,7 @@ func (t *STableSpec) insert(data interface{}, debug bool) error { } dataValue := reflect.ValueOf(data).Elem() - dataFields := reflectutils.FetchStructFieldNameValueInterfaces(dataValue) + dataFields := reflectutils.FetchStructFieldValueSet(dataValue) insertSql, values, err := t.insertSqlPrep(dataFields) if err != nil { return err @@ -122,7 +127,10 @@ func (t *STableSpec) insert(data interface{}, debug bool) error { q = q.Equals(c.Name(), lastId) } } else { - q = q.Equals(c.Name(), dataFields[c.Name()]) + priVal, _ := dataFields.GetInterface(c.Name()) + if !gotypes.IsNil(priVal) { + q = q.Equals(c.Name(), priVal) + } } } } diff --git a/vendor/yunion.io/x/sqlchemy/parser.go b/vendor/yunion.io/x/sqlchemy/parser.go index b4af1a67dd..93b87336d0 100644 --- a/vendor/yunion.io/x/sqlchemy/parser.go +++ b/vendor/yunion.io/x/sqlchemy/parser.go @@ -6,24 +6,24 @@ import ( "yunion.io/x/pkg/gotypes" "yunion.io/x/pkg/tristate" "yunion.io/x/pkg/util/reflectutils" - "yunion.io/x/pkg/utils" ) -func structField2ColumnSpec(field *reflect.StructField) IColumnSpec { - fieldname := reflectutils.GetStructFieldName(field) - tagmap := utils.TagMap(field.Tag) +func structField2ColumnSpec(field *reflectutils.SStructFieldValue) IColumnSpec { + fieldname := field.Info.MarshalName() + tagmap := field.Info.Tags if _, ok := tagmap[TAG_IGNORE]; ok { return nil } - var retCol = getFiledTypeCol(field.Type, fieldname, tagmap) - if retCol == nil && field.Type.Kind() == reflect.Ptr { - retCol = getFiledTypeCol(field.Type.Elem(), fieldname, tagmap) + fieldType := field.Value.Type() + var retCol = getFiledTypeCol(fieldType, fieldname, tagmap) + if retCol == nil && fieldType.Kind() == reflect.Ptr { + retCol = getFiledTypeCol(fieldType.Elem(), fieldname, tagmap) if retCol != nil { retCol.SetIsPointer() } } if retCol == nil { - panic("not supported type %s" + field.Type.Name()) + panic("unsupported colume data type %s" + fieldType.Name()) } return retCol } @@ -97,19 +97,15 @@ func getFiledTypeCol(fieldType reflect.Type, fieldname string, tagmap map[string return nil } -func struct2TableSpec(st reflect.Type, table *STableSpec) { - for i := 0; i < st.NumField(); i++ { - f := st.Field(i) - if f.Type.Kind() == reflect.Struct && f.Type != gotypes.TimeType { - struct2TableSpec(f.Type, table) - } else { - column := structField2ColumnSpec(&f) - if column != nil { - if column.IsIndex() { - table.AddIndex(column.IsUnique(), column.Name()) - } - table.columns = append(table.columns, column) +func struct2TableSpec(sv reflect.Value, table *STableSpec) { + fields := reflectutils.FetchStructFieldValueSet(sv) + for i := 0; i < len(fields); i += 1 { + column := structField2ColumnSpec(&fields[i]) + if column != nil { + if column.IsIndex() { + table.AddIndex(column.IsUnique(), column.Name()) } + table.columns = append(table.columns, column) } } } diff --git a/vendor/yunion.io/x/sqlchemy/query.go b/vendor/yunion.io/x/sqlchemy/query.go index 60b2cd7cd5..0aa5349444 100644 --- a/vendor/yunion.io/x/sqlchemy/query.go +++ b/vendor/yunion.io/x/sqlchemy/query.go @@ -463,11 +463,11 @@ func (q *SQuery) AllStringMap() ([]map[string]string, error) { } func mapString2Struct(mapResult map[string]string, destValue reflect.Value) error { - destFields := reflectutils.FetchStructFieldNameValues(destValue) + destFields := reflectutils.FetchStructFieldValueSet(destValue) var err error for k, v := range mapResult { if len(v) > 0 { - fieldValue, ok := destFields[k] + fieldValue, ok := destFields.GetValue(k) if ok { err = setValueBySQLString(fieldValue, v) if err != nil { diff --git a/vendor/yunion.io/x/sqlchemy/table.go b/vendor/yunion.io/x/sqlchemy/table.go index b49e71f379..c43cee7733 100644 --- a/vendor/yunion.io/x/sqlchemy/table.go +++ b/vendor/yunion.io/x/sqlchemy/table.go @@ -28,7 +28,8 @@ type STableField struct { } func NewTableSpecFromStruct(s interface{}, name string) *STableSpec { - st := reflect.TypeOf(s) + val := reflect.Indirect(reflect.ValueOf(s)) + st := val.Type() if st.Kind() != reflect.Struct { panic("expect Struct kind") } @@ -37,7 +38,7 @@ func NewTableSpecFromStruct(s interface{}, name string) *STableSpec { name: name, structType: st, } - struct2TableSpec(st, table) + struct2TableSpec(val, table) return table } diff --git a/vendor/yunion.io/x/sqlchemy/update.go b/vendor/yunion.io/x/sqlchemy/update.go index 3a6512d2fb..7dd3b2454e 100644 --- a/vendor/yunion.io/x/sqlchemy/update.go +++ b/vendor/yunion.io/x/sqlchemy/update.go @@ -22,26 +22,23 @@ func (ts *STableSpec) prepareUpdate(dt interface{}) (*SUpdateSession, error) { return nil, fmt.Errorf("Update input must be a Pointer") } dataValue := reflect.ValueOf(dt).Elem() - fields := reflectutils.FetchStructFieldNameValueInterfaces(dataValue) // fetchStructFieldNameValue(dataType, dataValue) + fields := reflectutils.FetchStructFieldValueSet(dataValue) // fetchStructFieldNameValue(dataType, dataValue) zeroPrimary := make([]string, 0) - zeroKeyIndex := make([]string, 0) for _, c := range ts.columns { k := c.Name() - ov, ok := fields[k] + ov, ok := fields.GetInterface(k) if !ok { continue } if c.IsPrimary() && c.IsZero(ov) { zeroPrimary = append(zeroPrimary, k) - } else if c.IsKeyIndex() && c.IsZero(ov) { - zeroKeyIndex = append(zeroKeyIndex, k) } } - if len(zeroPrimary) > 0 && len(zeroKeyIndex) > 0 { - return nil, fmt.Errorf("not a valid data, primary key %s and key index %s are empty", - strings.Join(zeroPrimary, ","), strings.Join(zeroKeyIndex, ",")) + if len(zeroPrimary) > 0 { + return nil, fmt.Errorf("not a valid data, primary key %s empty", + strings.Join(zeroPrimary, ",")) } originValue := gotypes.DeepCopyRv(dataValue) @@ -73,25 +70,21 @@ func (us *SUpdateSession) saveUpdate(dt interface{}) (map[string]SUpdateDiff, er // dataType := reflect.TypeOf(dt).Elem() dataValue := reflect.ValueOf(dt).Elem() - ofields := reflectutils.FetchStructFieldNameValueInterfaces(us.oValue) - fields := reflectutils.FetchStructFieldNameValueInterfaces(dataValue) + ofields := reflectutils.FetchStructFieldValueSet(us.oValue) + fields := reflectutils.FetchStructFieldValueSet(dataValue) versionFields := make([]string, 0) updatedFields := make([]string, 0) primaries := make(map[string]interface{}) - keyIndexes := make(map[string]interface{}) setters := make(map[string]SUpdateDiff) for _, c := range us.tableSpec.columns { k := c.Name() - of := ofields[k] - nf := fields[k] + of, _ := ofields.GetInterface(k) + nf, _ := fields.GetInterface(k) if !gotypes.IsNil(of) { if c.IsPrimary() && !c.IsZero(of) { // skip update primary key primaries[k] = of continue - } else if c.IsKeyIndex() && !c.IsZero(of) { - keyIndexes[k] = of - continue } } nc, ok := c.(*SIntegerColumn) @@ -142,15 +135,10 @@ func (us *SUpdateSession) saveUpdate(dt interface{}) (map[string]SUpdateDiff, er } buf.WriteString(" WHERE ") first = true - var indexFields map[string]interface{} - if len(primaries) > 0 { - indexFields = primaries - } else if len(keyIndexes) > 0 { - indexFields = keyIndexes - } else { - return nil, fmt.Errorf("neither primary key nor key indexes empty???") + if len(primaries) == 0 { + return nil, fmt.Errorf("primary key empty???") } - for k, v := range indexFields { + for k, v := range primaries { if first { first = false } else { diff --git a/vendor/yunion.io/x/structarg/structarg.go b/vendor/yunion.io/x/structarg/structarg.go index 48ce3b3c38..f5e924de1d 100644 --- a/vendor/yunion.io/x/structarg/structarg.go +++ b/vendor/yunion.io/x/structarg/structarg.go @@ -11,6 +11,7 @@ import ( "yunion.io/x/log" "yunion.io/x/pkg/gotypes" + "yunion.io/x/pkg/util/reflectutils" "yunion.io/x/pkg/utils" ) @@ -192,11 +193,12 @@ func (this *ArgumentParser) addStructArgument(prefix string, tp reflect.Type, va } func (this *ArgumentParser) addArgument(prefix string, f reflect.StructField, v reflect.Value) error { - tagMap := utils.TagMap(f.Tag) + info := reflectutils.ParseStructFieldJsonInfo(f) + tagMap := info.Tags help := tagMap[TAG_HELP] token, ok := tagMap[TAG_TOKEN] if !ok { - token = f.Name + token = info.MarshalName() } token = prefix + token shorttoken := tagMap[TAG_SHORT_TOKEN] @@ -283,7 +285,9 @@ func (this *ArgumentParser) addArgument(prefix string, f reflect.StructField, v var arg Argument = nil ovalue := reflect.New(v.Type()).Elem() ovalue.Set(v) - sarg := SingleArgument{token: token, shortToken: shorttoken, + sarg := SingleArgument{ + token: token, + shortToken: shorttoken, positional: positional, required: required, metavar: metavar, @@ -294,7 +298,8 @@ func (this *ArgumentParser) addArgument(prefix string, f reflect.StructField, v defValue: defval_t, value: v, ovalue: ovalue, - parser: this} + parser: this, + } // fmt.Println(token, f.Type, f.Type.Kind()) if subcommand { arg = &SubcommandArgument{SingleArgument: sarg, @@ -946,6 +951,9 @@ func (this *ArgumentParser) ParseFile(filepath string) error { line = strings.TrimSpace(removeComments(line)) // line = removeCharacters(line, `"'`) if len(line) > 0 { + if line[0] == '[' { + continue + } key, val, e := line2KeyValue(line) if e == nil { this.parseKeyValue(key, val)