seccomp: Optimize common 32-bit matchers away from disjunctions.

This turns applies the same logic as `extractRepeatedMatchers` to each half
of all `splitMatcher`s. This allows common 32-bit matchers in disjunctions
to be extracted out of them.

This is useful for almost all `EqualTo` rules, because they tend to look for
values that fit in the lower 32 bits. As such, the check for the higher 32
bits (that must be equal to zero in all cases of the disjunction) can be moved
out of the `Or`.

`splitMatcher` is updated to more efficiently handle the cases where either
of its branches are set to `halfAnyValue{}`.

```
              │   before    │                  after                  │
              │   sec/op    │   sec/op     vs base                    │
SentrySystrap   71.31n ± 9%   71.88n ± 5%       ~ (p=0.744 n=145+146)
SentryKVM       58.23n ± 8%   59.17n ± 5%       ~ (p=0.930 n=145)
NVProxyIoctl    92.92n ± 1%   86.14n ± 1%  -7.30% (n=145)

              │    before    │              after               │
              │  build-sec   │  build-sec   vs base             │
SentrySystrap    74.15m ± 0%   42.34m ± 0%  -42.90% (n=145+146)
SentryKVM       419.51m ± 0%   54.67m ± 0%  -86.97% (n=145)
NVProxyIoctl    1139.7m ± 0%   116.9m ± 0%  -89.75% (n=145)

              │      before       │                     after                      │
              │ compression-ratio │ compression-ratio  vs base                     │
SentrySystrap          3.443 ± 0%          3.872 ± 0%  +12.46% (p=0.000 n=145+146)
SentryKVM              3.501 ± 0%          3.983 ± 0%  +13.77% (p=0.000 n=145)
NVProxyIoctl           3.504 ± 0%          4.026 ± 0%  +14.90% (p=0.000 n=145)

              │   before    │              after              │
              │  gen-instr  │  gen-instr   vs base            │
SentrySystrap   1.618k ± 0%   1.545k ± 0%  -4.51% (n=145+146)
SentryKVM       1.733k ± 0%   1.677k ± 0%  -3.23% (n=145)
NVProxyIoctl    2.355k ± 0%   2.351k ± 0%  -0.17% (n=145)

              │   before   │              after              │
              │ opt-instr  │ opt-instr   vs base             │
SentrySystrap   470.0 ± 0%   399.0 ± 0%  -15.11% (n=145+146)
SentryKVM       495.0 ± 0%   421.0 ± 0%  -14.95% (n=145)
NVProxyIoctl    672.0 ± 0%   584.0 ± 0%  -13.10% (n=145)

              │    before    │              after               │
              │   opt-sec    │   opt-sec    vs base             │
SentrySystrap   109.70m ± 1%   91.88m ± 1%  -16.25% (n=145+146)
SentryKVM       103.24m ± 1%   90.06m ± 1%  -12.77% (n=145)
NVProxyIoctl     264.9m ± 1%   205.0m ± 1%  -22.61% (n=145)
```

PiperOrigin-RevId: 586505148
This commit is contained in:
Etienne Perot
2023-11-29 18:08:54 -08:00
committed by gVisor bot
parent 69e0c7643d
commit a880da69f2
3 changed files with 328 additions and 80 deletions
+235 -64
View File
@@ -279,16 +279,30 @@ func anySplitMatchersToAnyValue(rule SyscallRule) (SyscallRule, bool) {
// invalidValueMatcher is a stand-in `ValueMatcher` with a unique
// representation that doesn't look like any legitimate `ValueMatcher`.
// Calling any method other than `Repr` will fail.
// Calling any method other than `Repr` will panic.
// It is used as an intermediate step for some optimizers.
type invalidValueMatcher struct {
ValueMatcher
}
// Repr implements `ValueMatcher.Repr`.
func (m invalidValueMatcher) Repr() string {
func (invalidValueMatcher) Repr() string {
return "invalidValueMatcher"
}
// invalidHalfValueMatcher is a stand-in `HalfValueMatcher` with a unique
// representation that doesn't look like any legitimate `HalfValueMatcher`.
// Calling any method other than `Repr` will panic.
// It is used as an intermediate step for some optimizers.
type invalidHalfValueMatcher struct {
halfValueMatcher
}
// Repr implements `HalfValueMatcher.Repr`.
func (invalidHalfValueMatcher) Repr() string {
return "invalidHalfValueMatcher"
}
// sameStringSet returns whether the given string sets are equal.
func sameStringSet(m1, m2 map[string]struct{}) bool {
if len(m1) != len(m2) {
@@ -391,76 +405,230 @@ func extractRepeatedMatchers(rule SyscallRule) (SyscallRule, bool) {
}
}
// extractData is the result of extracting a matcher at `argNum`.
type extractData struct {
// extractedMatcher is the extracted matcher that should be AND'd
// with the rest.
extractedMatcher ValueMatcher
// otherMatchers represents the rest of the matchers after
// `extractedMatcher` is extracted from a `PerArg`.
// The matcher that was extracted should be replaced with something
// that matches any value (i.e. either `AnyValue` or `halfAnyValue`).
otherMatchers PerArg
// otherMatchersSig represents the signature of other matchers, with
// the extracted matcher being replaced with an "invalid" matcher.
// The "invalid" matcher acts as a token that is equal across all
// instances of `otherMatchersSig` for the other `PerArg` rules of the
// `Or` expression.
// `otherMatchersSig` isn't the same as `otherMatchers.Signature()`,
// as `otherMatchers` does not contain this "invalid" matcher (it
// contains a matcher that matches any value instead).
otherMatchersSig string
// extractedMatcherIsAnyValue is true iff `extractedMatcher` would
// match any value thrown at it.
// If this is the case across all branches of the `Or` expression,
// the optimization is skipped.
extractedMatcherIsAnyValue bool
// otherMatchersAreAllAnyValue is true iff all matchers in
// `otherMatchers` would match any value thrown at them.
// If this is the case across all branches of the `Or` expression,
// the optimization is skipped.
otherMatchersAreAllAnyValue bool
}
allOtherMatchersSigs := make(map[string]struct{}, len(orRule))
argExprToOtherMatchersSigs := make(map[string]map[string]struct{}, len(orRule))
for argNum := 0; argNum < len(orRule[0].(PerArg)); argNum++ {
// Check if this argNum is always AnyValue,
// or if all other arguments are always AnyValue.
// If either of that is true, there is nothing for this filter to do.
allArgNumMatchersAreAnyValue := true
allOtherMatchersAreAnyValue := true
for _, subRule := range orRule {
perArg := subRule.(PerArg)
for i, valueMatcher := range perArg {
_, isAnyValue := valueMatcher.(AnyValue)
if i == argNum {
allArgNumMatchersAreAnyValue = allArgNumMatchersAreAnyValue && isAnyValue
} else {
allOtherMatchersAreAnyValue = allOtherMatchersAreAnyValue && isAnyValue
}
}
}
if allArgNumMatchersAreAnyValue || allOtherMatchersAreAnyValue {
// Cannot optimize.
continue
}
// Check if `argNum` takes on a set of matchers common for all
// combinations of all other matchers.
clear(allOtherMatchersSigs)
clear(argExprToOtherMatchersSigs)
for _, subRule := range orRule {
perArg := subRule.(PerArg)
repr := perArg[argNum].Repr()
otherMatchers := perArg.clone()
otherMatchers[argNum] = invalidValueMatcher{}
otherMatchersSig := otherMatchers.signature()
allOtherMatchersSigs[otherMatchersSig] = struct{}{}
if _, reprSeen := argExprToOtherMatchersSigs[repr]; !reprSeen {
argExprToOtherMatchersSigs[repr] = make(map[string]struct{}, len(orRule))
// We try to extract a common matcher by three ways, which we
// iterate over here.
// Each of them returns the result of their extraction attempt,
// along with a boolean representing whether extraction was
// possible at all.
// To "extract" a matcher means to replace it with an "invalid"
// matcher in the PerArg expression, and checking if their set of
// signatures is identical for each unique `Repr()` of the extracted
// matcher. For splittable matcher, we try each half as well.
// Conceptually (simplify PerArg to 3 arguments for simplicity),
// if we have:
//
// Or{
// PerArg{A, B, C},
// PerArg{D, E, F},
// }
//
// ... then first, we will try:
//
// Or{
// PerArg{invalid, B, C}
// PerArg{invalid, E, F}
// }
//
// ... then, assuming both A and D are `splitMatcher`s:
// we will try:
//
// Or{
// PerArg{splitMatcher{invalid, A.lowMatcher}, B, C}
// PerArg{splitMatcher{invalid, D.lowMatcher}, E, F}
// }
//
// ... and finally we will try:
//
// Or{
// PerArg{splitMatcher{A.highMatcher, invalid}, B, C}
// PerArg{splitMatcher{D.highMatcher, invalid}, E, F}
// }
for _, extractFn := range []func(PerArg) (extractData, bool){
// Return whole ValueMatcher at a time:
func(pa PerArg) (extractData, bool) {
extractedMatcher := pa[argNum]
_, extractedMatcherIsAnyValue := extractedMatcher.(AnyValue)
otherMatchers := pa.clone()
otherMatchers[argNum] = invalidValueMatcher{}
otherMatchersSig := otherMatchers.signature()
otherMatchers[argNum] = AnyValue{}
otherMatchersAreAllAnyValue := true
for _, valueMatcher := range otherMatchers {
if _, isAnyValue := valueMatcher.(AnyValue); !isAnyValue {
otherMatchersAreAllAnyValue = false
break
}
}
return extractData{
extractedMatcher: extractedMatcher,
otherMatchers: otherMatchers,
otherMatchersSig: otherMatchersSig,
extractedMatcherIsAnyValue: extractedMatcherIsAnyValue,
otherMatchersAreAllAnyValue: otherMatchersAreAllAnyValue,
}, true
},
// Extract a matcher for the high bits only:
func(pa PerArg) (extractData, bool) {
split, isSplit := pa[argNum].(splitMatcher)
if !isSplit {
return extractData{}, false
}
_, extractedMatcherIsAnyValue := split.highMatcher.(halfAnyValue)
_, lowMatcherIsAnyValue := split.lowMatcher.(halfAnyValue)
extractedMatcher := high32BitsMatch(split.highMatcher)
otherMatchers := pa.clone()
otherMatchers[argNum] = splitMatcher{
highMatcher: invalidHalfValueMatcher{},
lowMatcher: split.lowMatcher,
}
otherMatchersSig := otherMatchers.signature()
otherMatchers[argNum] = low32BitsMatch(split.lowMatcher)
otherMatchersAreAllAnyValue := lowMatcherIsAnyValue
for i, valueMatcher := range otherMatchers {
if i == argNum {
continue
}
if _, isAnyValue := valueMatcher.(AnyValue); !isAnyValue {
otherMatchersAreAllAnyValue = false
break
}
}
return extractData{
extractedMatcher: extractedMatcher,
otherMatchers: otherMatchers,
otherMatchersSig: otherMatchersSig,
extractedMatcherIsAnyValue: extractedMatcherIsAnyValue,
otherMatchersAreAllAnyValue: otherMatchersAreAllAnyValue,
}, true
},
// Extract a matcher for the low bits only:
func(pa PerArg) (extractData, bool) {
split, isSplit := pa[argNum].(splitMatcher)
if !isSplit {
return extractData{}, false
}
_, extractedMatcherIsAnyValue := split.lowMatcher.(halfAnyValue)
_, highMatcherIsAnyValue := split.highMatcher.(halfAnyValue)
extractedMatcher := low32BitsMatch(split.lowMatcher)
otherMatchers := pa.clone()
otherMatchers[argNum] = splitMatcher{
highMatcher: split.highMatcher,
lowMatcher: invalidHalfValueMatcher{},
}
otherMatchersSig := otherMatchers.signature()
otherMatchers[argNum] = high32BitsMatch(split.highMatcher)
otherMatchersAreAllAnyValue := highMatcherIsAnyValue
for i, valueMatcher := range otherMatchers {
if i == argNum {
continue
}
if _, isAnyValue := valueMatcher.(AnyValue); !isAnyValue {
otherMatchersAreAllAnyValue = false
break
}
}
return extractData{
extractedMatcher: extractedMatcher,
otherMatchers: otherMatchers,
otherMatchersSig: otherMatchersSig,
extractedMatcherIsAnyValue: extractedMatcherIsAnyValue,
otherMatchersAreAllAnyValue: otherMatchersAreAllAnyValue,
}, true
},
} {
clear(allOtherMatchersSigs)
clear(argExprToOtherMatchersSigs)
allExtractable := true
allArgNumMatchersAreAnyValue := true
allOtherMatchersAreAnyValue := true
for _, subRule := range orRule {
ed, extractable := extractFn(subRule.(PerArg))
if allExtractable = allExtractable && extractable; !allExtractable {
break
}
allArgNumMatchersAreAnyValue = allArgNumMatchersAreAnyValue && ed.extractedMatcherIsAnyValue
allOtherMatchersAreAnyValue = allOtherMatchersAreAnyValue && ed.otherMatchersAreAllAnyValue
repr := ed.extractedMatcher.Repr()
allOtherMatchersSigs[ed.otherMatchersSig] = struct{}{}
if _, reprSeen := argExprToOtherMatchersSigs[repr]; !reprSeen {
argExprToOtherMatchersSigs[repr] = make(map[string]struct{}, len(orRule))
}
argExprToOtherMatchersSigs[repr][ed.otherMatchersSig] = struct{}{}
}
argExprToOtherMatchersSigs[repr][otherMatchersSig] = struct{}{}
}
// Now check if each possible repr of `argNum` got the same set of
// signatures for other matchers as `allOtherMatchersSigs`.
sameOtherMatchers := true
for _, omsigs := range argExprToOtherMatchersSigs {
if !sameStringSet(omsigs, allOtherMatchersSigs) {
sameOtherMatchers = false
break
if !allExtractable || allArgNumMatchersAreAnyValue || allOtherMatchersAreAnyValue {
// Cannot optimize.
continue
}
// Now check if each possible repr of `argNum` got the same set of
// signatures for other matchers as `allOtherMatchersSigs`.
sameOtherMatchers := true
for _, omsigs := range argExprToOtherMatchersSigs {
if !sameStringSet(omsigs, allOtherMatchersSigs) {
sameOtherMatchers = false
break
}
}
if !sameOtherMatchers {
continue
}
// We can simplify the rule by extracting `argNum` out.
// Create two copies of `orRule`: One with only `argNum`,
// and the other one with all arguments except `argNum`.
// This will likely contain many duplicates but that's OK,
// they'll be optimized out by `deduplicatePerArgs`.
argNumMatch := Or(make([]SyscallRule, len(orRule)))
otherArgsMatch := Or(make([]SyscallRule, len(orRule)))
for i, subRule := range orRule {
ed, _ := extractFn(subRule.(PerArg))
onlyArg := PerArg{AnyValue{}, AnyValue{}, AnyValue{}, AnyValue{}, AnyValue{}, AnyValue{}, AnyValue{}}
onlyArg[argNum] = ed.extractedMatcher
argNumMatch[i] = onlyArg
otherArgsMatch[i] = ed.otherMatchers
}
// Attempt to optimize the "other" arguments:
otherArgsMatchOpt, _ := extractRepeatedMatchers(otherArgsMatch)
return And{argNumMatch, otherArgsMatchOpt}, true
}
if !sameOtherMatchers {
continue
}
// We can simplify the rule by extracting `argNum` out.
// Create two copies of `orRule`: One with only `argNum`,
// and the other one with all arguments except `argNum`.
// This will likely contain many duplicates but that's OK,
// they'll be optimized out by `deduplicatePerArgs`.
argNumMatch := Or(make([]SyscallRule, len(orRule)))
otherArgsMatch := Or(make([]SyscallRule, len(orRule)))
for i, subRule := range orRule {
perArg := subRule.(PerArg)
onlyArg := PerArg{AnyValue{}, AnyValue{}, AnyValue{}, AnyValue{}, AnyValue{}, AnyValue{}, AnyValue{}}
onlyArg[argNum] = perArg[argNum]
allExceptArg := perArg.clone()
allExceptArg[argNum] = AnyValue{}
argNumMatch[i] = onlyArg
otherArgsMatch[i] = allExceptArg
}
// Attempt to optimize the "other" arguments:
otherArgsMatchOpt, _ := extractRepeatedMatchers(otherArgsMatch)
return And{argNumMatch, otherArgsMatchOpt}, true
}
return rule, false
}
@@ -554,6 +722,9 @@ func optimizeSyscallRule(rule SyscallRule) SyscallRule {
convertUselessPerArgToMatchAll,
// Replace `ValueMatcher`s that are splittable into their split version.
// Like `nilInPerArgToAnyValue`, this isn't so much an optimization,
// but allows the matchers below (which are `splitMatcher`-aware) to not
// have to carry logic to split the matchers they encounter.
splitMatchers,
// Replace `halfValueMatcher`s with their simplified version.
+51 -4
View File
@@ -108,7 +108,7 @@ type halfEqualTo uint32
// Repr implements `halfValueMatcher.Repr`.
func (heq halfEqualTo) Repr() string {
return fmt.Sprintf("halfEq(%d)", uint32(heq))
return fmt.Sprintf("halfEq(%#x)", uint32(heq))
}
// HalfRender implements `halfValueMatcher.HalfRender`.
@@ -123,7 +123,7 @@ type halfNotSet uint32
// Repr implements `halfValueMatcher.Repr`.
func (hns halfNotSet) Repr() string {
return fmt.Sprintf("halfNotSet(%x)", uint32(hns))
return fmt.Sprintf("halfNotSet(%#x)", uint32(hns))
}
// HalfRender implements `halfValueMatcher.HalfRender`.
@@ -141,7 +141,7 @@ type halfMaskedEqual struct {
// Repr implements `halfValueMatcher.Repr`.
func (hmeq halfMaskedEqual) Repr() string {
return fmt.Sprintf("halfMaskedEqual(%x, %x)", hmeq.mask, hmeq.value)
return fmt.Sprintf("halfMaskedEqual(%#x, %#x)", hmeq.mask, hmeq.value)
}
// HalfRender implements `halfValueMatcher.HalfRender`.
@@ -172,16 +172,45 @@ func (sm splitMatcher) String() string {
// Repr implements `ValueMatcher.Repr`.
func (sm splitMatcher) Repr() string {
if sm.repr == "" {
_, highIsAnyValue := sm.highMatcher.(halfAnyValue)
_, lowIsAnyValue := sm.lowMatcher.(halfAnyValue)
if highIsAnyValue && lowIsAnyValue {
return "split(*)"
}
if highIsAnyValue {
return fmt.Sprintf("low=%s", sm.lowMatcher.Repr())
}
if lowIsAnyValue {
return fmt.Sprintf("high=%s", sm.highMatcher.Repr())
}
return fmt.Sprintf("(high=%s && low=%s)", sm.highMatcher.Repr(), sm.lowMatcher.Repr())
}
return sm.repr
}
// Render implements `ValueMatcher.Render`.
func (sm splitMatcher) Render(program *syscallProgram, labelSet *labelSet, value matchedValue) {
_, highIsAny := sm.highMatcher.(halfAnyValue)
_, lowIsAny := sm.lowMatcher.(halfAnyValue)
if highIsAny && lowIsAny {
program.JumpTo(labelSet.Matched())
return
}
if highIsAny {
value.LoadLow32Bits()
sm.lowMatcher.HalfRender(program, labelSet)
return
}
if lowIsAny {
value.LoadHigh32Bits()
sm.highMatcher.HalfRender(program, labelSet)
return
}
// We render the "low" bits first on the assumption that most syscall
// arguments fit within 32-bits, and those rules actually only care
// about the value of the low 32 bits. This way, we only check the
// high 32 bits if the low 32 bits have already matched.
lowLabels := labelSet.Push("low", labelSet.NewLabel(), labelSet.Mismatched())
lowFrag := program.Record()
value.LoadLow32Bits()
@@ -195,6 +224,24 @@ func (sm splitMatcher) Render(program *syscallProgram, labelSet *labelSet, value
highFrag.MustHaveJumpedTo(labelSet.Matched(), labelSet.Mismatched())
}
// high32BitsMatch returns a `splitMatcher` that only matches the high 32 bits
// of a 64-bit value.
func high32BitsMatch(hvm halfValueMatcher) splitMatcher {
return splitMatcher{
highMatcher: hvm,
lowMatcher: halfAnyValue{},
}
}
// low32BitsMatch returns a `splitMatcher` that only matches the low 32 bits
// of a 64-bit value.
func low32BitsMatch(hvm halfValueMatcher) splitMatcher {
return splitMatcher{
highMatcher: halfAnyValue{},
lowMatcher: hvm,
}
}
// splittableValueMatcher should be implemented by `ValueMatcher` that can
// be expressed as a `splitMatcher`.
type splittableValueMatcher interface {
+42 -12
View File
@@ -1237,6 +1237,7 @@ func TestOptimizeSyscallRule(t *testing.T) {
c2 = 0xC2C2C2C2C2C2C2C2
c3 = 0xC3C3C3C3C3C3C3C3
d0 = 0xD0D0D0D0D0D0D0D0
d1 = 0xD1D1D1D1D1D1D1D1
)
// split replaces a `splittableValueMatcher` rule with its `splitMatcher`
@@ -1245,6 +1246,13 @@ func TestOptimizeSyscallRule(t *testing.T) {
return matcher.split()
}
highEq := func(val uintptr) splitMatcher {
return high32BitsMatch(halfEqualTo(val))
}
lowEq := func(val uintptr) splitMatcher {
return low32BitsMatch(halfEqualTo(val))
}
for _, test := range []struct {
name string
rule SyscallRule
@@ -1260,25 +1268,25 @@ func TestOptimizeSyscallRule(t *testing.T) {
name: "flatten Or rule",
rule: Or{
Or{
PerArg{EqualTo(0x11)},
PerArg{EqualTo(a1)},
Or{
PerArg{EqualTo(0x22)},
PerArg{EqualTo(0x33)},
PerArg{EqualTo(b1)},
PerArg{EqualTo(b2)},
},
PerArg{EqualTo(0x44)},
PerArg{EqualTo(c1)},
},
Or{
PerArg{EqualTo(0x55)},
PerArg{EqualTo(0x66)},
PerArg{EqualTo(d0)},
PerArg{EqualTo(d1)},
},
},
want: Or{
PerArg{s(EqualTo(0x11)), av, av, av, av, av, av},
PerArg{s(EqualTo(0x22)), av, av, av, av, av, av},
PerArg{s(EqualTo(0x33)), av, av, av, av, av, av},
PerArg{s(EqualTo(0x44)), av, av, av, av, av, av},
PerArg{s(EqualTo(0x55)), av, av, av, av, av, av},
PerArg{s(EqualTo(0x66)), av, av, av, av, av, av},
PerArg{s(EqualTo(a1)), av, av, av, av, av, av},
PerArg{s(EqualTo(b1)), av, av, av, av, av, av},
PerArg{s(EqualTo(b2)), av, av, av, av, av, av},
PerArg{s(EqualTo(c1)), av, av, av, av, av, av},
PerArg{s(EqualTo(d0)), av, av, av, av, av, av},
PerArg{s(EqualTo(d1)), av, av, av, av, av, av},
},
},
{
@@ -1475,6 +1483,28 @@ func TestOptimizeSyscallRule(t *testing.T) {
},
},
},
{
name: "Common halfValueMatchers in splitMatchers are extracted",
rule: Or{
PerArg{EqualTo(0xA1), EqualTo(0xB1), EqualTo(c1), EqualTo(d0)},
PerArg{EqualTo(0xA1), EqualTo(0xB1), EqualTo(c2), EqualTo(d0)},
PerArg{EqualTo(0xA2), EqualTo(0xB2), EqualTo(c1), EqualTo(d0)},
PerArg{EqualTo(0xA2), EqualTo(0xB2), EqualTo(c2), EqualTo(d0)},
},
want: And{
PerArg{highEq(0), av, av, av, av, av, av},
PerArg{av, highEq(0), av, av, av, av, av},
Or{
PerArg{av, av, s(EqualTo(c1)), av, av, av, av},
PerArg{av, av, s(EqualTo(c2)), av, av, av, av},
},
PerArg{av, av, av, s(EqualTo(d0)), av, av, av},
Or{
PerArg{lowEq(0xA1), lowEq(0xB1), av, av, av, av, av},
PerArg{lowEq(0xA2), lowEq(0xB2), av, av, av, av, av},
},
},
},
} {
t.Run(test.name, func(t *testing.T) {
var got SyscallRule