Merge pull request #951 in YUNIONIO/onecloud from ~TANGBIN/onecloud:bugfix/tb-aws-update-secrule to release/2.4.0

* commit 'c3c86c36e58c2721a2b519bc60f3f7d96da3f3d3':
  update aws secrule sync
This commit is contained in:
唐斌
2019-01-23 11:45:29 +08:00
23 changed files with 1312 additions and 400 deletions
Generated
+8 -8
View File
@@ -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"
+29
View File
@@ -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
+1
View File
@@ -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)
+37 -127
View File
@@ -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
+3 -7
View File
@@ -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)
+20 -48
View File
@@ -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
}
+104
View File
@@ -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 {
+177
View File
@@ -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
}
+73 -51
View File
@@ -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
}
+260
View File
@@ -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
}
+217
View File
@@ -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
}
+93 -54
View File
@@ -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
}
+181
View File
@@ -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
}
+14
View File
@@ -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 "", ""
}
+16 -16
View File
@@ -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
}
}
-1
View File
@@ -31,5 +31,4 @@ const (
TAG_AUTOVERSION = "auto_version"
TAG_UPDATE_TIMESTAMP = "updated_at"
TAG_CREATE_TIMESTAMP = "created_at"
TAG_KEY_INDEX = "key_index"
)
+20 -30
View File
@@ -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 {
+14 -6
View File
@@ -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)
}
}
}
}
+16 -20
View File
@@ -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)
}
}
}
+2 -2
View File
@@ -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 {
+3 -2
View File
@@ -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
}
+12 -24
View File
@@ -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 {
+12 -4
View File
@@ -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)