From e6979cb4d68c36fe6ae85e6c9b7445b7407cd33c Mon Sep 17 00:00:00 2001 From: Etienne Perot Date: Wed, 15 Nov 2023 14:29:35 -0800 Subject: [PATCH] `seccomp`: Reorder generated syscall rules for better efficiency. MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit This subdivides all `RuleSet`s into single-syscall rulesets, and then classifies them depending on: - Whether they are "trivial" or not, where "trivial" means that the syscall rules do not perform any verification of the syscall arguments or RIP. - Whether they are marked "hot" or not, where "hot" means "expected to be frequently called". It then orders the program as follows: - All hot non-trivial rules go first. This makes it so that the host kernel can clear the syscall faster for frequently-called syscalls. These are checked linearly, as they tend to follow a Pareto distribution in terms of frequency. If they need a vsyscall check, that check is added individually. - All cold rules go next, and form a BST. This mimics the structure of the BST construction that existed prior to this change. - Lastly, all the trivial syscalls are added as a last BST. This speeds up rule evaluation because it maximizes the use of Linux's seccomp cache for trivial syscalls. These are therefore only ever checked once, so they can stay at the "bottom" of the program. All remaining (non-trivial) syscalls are ordered such that hot syscalls are checked first, and then cold syscalls are checked with a BST. This is a complex and security-sensitive change, but fuzz testing with full branch coverage has shown that this has the exact same behavior as a BPF program taken from before any of my recent seccomp/BPF changes (other than the one adding non-negative FD checks to all `ioctl(2)` system calls). Some benchmark results ("orig" is the state before this change): ``` │ orig │ reordered │ │ sec/op │ sec/op vs base │ SentrySystrap/futex 79.44n ± 2% 73.93n ± 2% -6.93% (n=729+722) SentrySystrap/nanosleep 112.3n ± 12% 107.2n ± 12% ~ (p=0.505 n=482+477) SentrySystrap/sendmmsg 88.50n ± 1% 81.62n ± 1% -7.78% (n=729+722) SentrySystrap/fstat 30.80n ± 2% 30.63n ± 3% ~ (p=0.903 n=722+712) [...] SentrySystrap/Postgres-48 64.30n ± 5% 61.74n ± 6% -3.97% (p=0.039 n=376+377) ``` PiperOrigin-RevId: 582808055 --- pkg/seccomp/BUILD | 1 + pkg/seccomp/seccomp.go | 571 +++++++++++++++++++++++++++++------ pkg/seccomp/seccomp_test.go | 577 ++++++++++++++++++++++++++++++++++++ 3 files changed, 1056 insertions(+), 93 deletions(-) diff --git a/pkg/seccomp/BUILD b/pkg/seccomp/BUILD index 0c721654e..667428cc0 100644 --- a/pkg/seccomp/BUILD +++ b/pkg/seccomp/BUILD @@ -38,5 +38,6 @@ go_test( deps = [ "//pkg/abi/linux", "//pkg/bpf", + "@org_golang_x_sys//unix:go_default_library", ], ) diff --git a/pkg/seccomp/seccomp.go b/pkg/seccomp/seccomp.go index 81ddd1203..60c1047ad 100644 --- a/pkg/seccomp/seccomp.go +++ b/pkg/seccomp/seccomp.go @@ -19,6 +19,8 @@ package seccomp import ( "fmt" "sort" + "strconv" + "strings" "time" "gvisor.dev/gvisor/pkg/abi/linux" @@ -306,6 +308,15 @@ type ProgramOptions struct { // BadArchAction is the action returned when the architecture of the // syscall structure input doesn't match the one the program expects. BadArchAction linux.BPFAction + + // HotSyscalls is the set of syscall numbers that are the hottest, + // where "hotness" refers to frequency (regardless of the amount of + // computation that the kernel will do handling them, and regardless of + // the complexity of the syscall rule for this). + // It should only contain very hot syscalls (i.e. any syscall that is + // called >10% of the time out of all syscalls made). + // It is ordered from most frequent to least frequent. + HotSyscalls []uintptr } // DefaultProgramOptions returns the default program options. @@ -345,17 +356,17 @@ func BuildProgram(rules []RuleSet, options ProgramOptions) ([]bpf.Instruction, B start := time.Now() // Make a copy of the syscall rules and optimize them. - rulesCopy := make([]RuleSet, len(rules)) - for i, rs := range rules { - rulesMap := MakeSyscallRules(make(map[uintptr]SyscallRule, len(rs.Rules.rules))) - for sysno, rule := range rs.Rules.rules { - rulesMap.rules[sysno] = optimizeSyscallRule(rule) - } - rs.Rules = rulesMap - rulesCopy[i] = rs + ors, err := orderRuleSets(rules, options) + if err != nil { + return nil, BuildStats{}, err } ruleOptimizeDuration := time.Since(start) + possibleActions := make(map[linux.BPFAction]struct{}) + for _, ruleSet := range rules { + possibleActions[ruleSet.Action] = struct{}{} + } + program := &syscallProgram{ program: bpf.NewProgramBuilder(), } @@ -367,9 +378,11 @@ func BuildProgram(rules []RuleSet, options ProgramOptions) ([]bpf.Instruction, B badArchLabel := label("badarch") program.Stmt(bpf.Ld|bpf.Abs|bpf.W, seccompDataOffsetArch) program.IfNot(bpf.Jmp|bpf.Jeq|bpf.K, LINUX_AUDIT_ARCH, badArchLabel) - if err := buildIndex(rulesCopy, program); err != nil { + orsFrag := program.Record() + if err := ors.render(program); err != nil { return nil, BuildStats{}, err } + orsFrag.MustHaveJumpedToOrReturned([]label{defaultLabel}, possibleActions) // Default label if none of the rules matched: program.Label(defaultLabel) @@ -398,53 +411,417 @@ func BuildProgram(rules []RuleSet, options ProgramOptions) ([]bpf.Instruction, B }, nil } -// buildIndex builds a BST to quickly search through all syscalls. -func buildIndex(rules []RuleSet, program *syscallProgram) error { - // Do nothing if rules is empty. - if len(rules) == 0 { - return nil - } +// singleSyscallRuleSet represents what to do for a single syscall. +// It is used inside `orderedRules`. +type singleSyscallRuleSet struct { + sysno uintptr + rules []syscallRuleAction + vsyscall bool +} - // Build a list of all application system calls, across all given rule - // sets. We have a simple BST, but may dispatch individual matchers - // with different actions. The matchers are evaluated linearly. - requiredSyscalls := make(map[uintptr]struct{}) +// Render renders the ruleset for this syscall. +func (ssrs singleSyscallRuleSet) Render(program *syscallProgram, ls *labelSet, noMatch label) { + frag := program.Record() + if ssrs.vsyscall { + // Emit a vsyscall check. + // This rule ensures that the top bit is set in the + // instruction pointer, which is where the vsyscall page + // will be mapped. + program.Stmt(bpf.Ld|bpf.Abs|bpf.W, seccompDataOffsetIPHigh) + program.IfNot(bpf.Jmp|bpf.Jset|bpf.K, vsyscallPageIPMask, noMatch) + } + var nextRule label + actions := make(map[linux.BPFAction]struct{}) + for i, ra := range ssrs.rules { + actions[ra.action] = struct{}{} + + // Render the rule. + nextRule = ls.NewLabel() + ruleLabels := ls.Push(fmt.Sprintf("sysno%d_rule%d", ssrs.sysno, i), ls.NewLabel(), nextRule) + ruleFrag := program.Record() + ra.rule.Render(program, ruleLabels) + program.Label(ruleLabels.Matched()) + program.Ret(ra.action) + ruleFrag.MustHaveJumpedToOrReturned( + []label{nextRule}, + map[linux.BPFAction]struct{}{ + ra.action: struct{}{}, + }) + program.Label(nextRule) + } + program.JumpTo(noMatch) + frag.MustHaveJumpedToOrReturned([]label{noMatch}, actions) +} + +// String returns a human-friendly representation of the +// `singleSyscallRuleSet`. +func (ssrs singleSyscallRuleSet) String() string { + var sb strings.Builder + sb.WriteString("sysno=") + sb.WriteString(strconv.Itoa(int(ssrs.sysno))) + if ssrs.vsyscall { + sb.WriteString("[vsyscall]") + } + sb.WriteString(": ") + if len(ssrs.rules) == 0 { + sb.WriteString("(no rules)") + } else { + sb.WriteRune('{') + for i, r := range ssrs.rules { + if i != 0 { + sb.WriteString(", ") + } + sb.WriteString(r.String()) + } + sb.WriteRune('}') + } + return sb.String() +} + +// syscallRuleAction groups a `SyscallRule` and an action that should be +// returned if the rule matches. +type syscallRuleAction struct { + rule SyscallRule + action linux.BPFAction +} + +// String returns a human-friendly representation of the `syscallRuleAction`. +func (sra syscallRuleAction) String() string { + return fmt.Sprintf("(%v) => %v", sra.rule.String(), sra.action) +} + +// orderedRules contains an ordering of syscall rules used to render a +// program. It is derived from a list of `RuleSet`s and `ProgramOptions`. +// Its fields represent the order in which rulesets are rendered. +// There are three categorization criteria: +// - "Hot" vs "cold": hot syscalls go first and are checked linearly, cold +// syscalls go later. +// - "Trivial" vs "non-trivial": A "trivial" syscall rule means one that +// does not require checking any argument or RIP data. This basically +// means a syscall mapped to `MatchAll{}`. +// If a syscall shows up in multiple RuleSets where any of them is +// non-trivial, the whole syscall is considered non-trivial. +// - "vsyscall" vs "non-vsyscall": A syscall that needs vsyscall checking +// checks that the function is dispatched from the vsyscall page by +// checking RIP. This inherently makes it non-trivial. All trivial +// rules are non-vsyscall, but not all non-vsyscall rules are trivial. +type orderedRuleSets struct { + // hotNonTrivial is the set of hot syscalls that are non-trivial + // and may or may not require vsyscall checking. + // They come first and are checked linearly using `hotNonTrivialOrder`. + hotNonTrivial map[uintptr]singleSyscallRuleSet + + // hotNonTrivial is the set of hot syscalls that are non-trivial + // and may or may not require vsyscall checking. + // They come first and are checked linearly using `hotNonTrivialOrder`. + hotNonTrivialOrder []uintptr + + // coldNonTrivial is the set of non-hot syscalls that are non-trivial. + // They may or may not require vsyscall checking. + // They come second. + coldNonTrivial map[uintptr]singleSyscallRuleSet + + // trivial is the set of syscalls that are trivial. They may or may not be + // hot, but they may not require vsyscall checking (otherwise they would + // be non-trivial). + // They come last. This is because the host kernel will cache the results + // of these system calls, and will never execute them on the hot path. + trivial map[uintptr]singleSyscallRuleSet +} + +// orderRuleSets converts a set of `RuleSet`s into an `orderedRuleSets`. +func orderRuleSets(rules []RuleSet, options ProgramOptions) (orderedRuleSets, error) { + // Do a pass to determine if vsyscall is consistent across syscall numbers. + vsyscallBySysno := make(map[uintptr]bool) for _, rs := range rules { for sysno := range rs.Rules.rules { - requiredSyscalls[sysno] = struct{}{} - } - } - syscalls := make([]uintptr, 0, len(requiredSyscalls)) - for sysno := range requiredSyscalls { - syscalls = append(syscalls, sysno) - } - sort.Slice(syscalls, func(i, j int) bool { return syscalls[i] < syscalls[j] }) - for _, sysno := range syscalls { - for _, rs := range rules { - // Print only if there is a corresponding set of rules. - if r, ok := rs.Rules.rules[sysno]; ok { - log.Debugf("syscall filter %v: %s => 0x%x", SyscallName(sysno), r, rs.Action) + if prevVsyscall, ok := vsyscallBySysno[sysno]; ok { + if prevVsyscall != rs.Vsyscall { + return orderedRuleSets{}, fmt.Errorf("syscall %d has conflicting vsyscall checking rules", sysno) + } + } else { + vsyscallBySysno[sysno] = rs.Vsyscall } } } - root := createBST(syscalls) - root.root = true + // Build a single map of per-syscall syscallRuleActions. + // We will split this map up later. + allSyscallRuleActions := make(map[uintptr][]syscallRuleAction) + for _, rs := range rules { + for sysno, rule := range rs.Rules.rules { + existing, found := allSyscallRuleActions[sysno] + if !found { + allSyscallRuleActions[sysno] = []syscallRuleAction{{ + rule: rule, + action: rs.Action, + }} + continue + } + if existing[len(existing)-1].action == rs.Action { + // If the last action for this syscall is the same, union the rules. + existing[len(existing)-1].rule = Or{existing[len(existing)-1].rule, rule} + } else { + // Otherwise, add it as a new ruleset. + existing = append(existing, syscallRuleAction{ + rule: rule, + action: rs.Action, + }) + } + allSyscallRuleActions[sysno] = existing + } + } - // Load syscall number into A and run through BST. - // - // A = seccomp_data.nr + // Optimize all rules. + for _, ruleActions := range allSyscallRuleActions { + for i, ra := range ruleActions { + ra.rule = optimizeSyscallRule(ra.rule) + ruleActions[i] = ra + } + } + + // Do a pass that checks which syscall numbers are trivial. + isTrivial := make(map[uintptr]bool) + for sysno, ruleActions := range allSyscallRuleActions { + for _, ra := range ruleActions { + _, isMatchAll := ra.rule.(MatchAll) + isVsyscall := vsyscallBySysno[sysno] + trivial := isMatchAll && !isVsyscall + if prevTrivial, ok := isTrivial[sysno]; ok { + isTrivial[sysno] = prevTrivial && trivial + } else { + isTrivial[sysno] = trivial + } + } + } + + // Compute the set of non-trivial hot syscalls and their order. + hotNonTrivialSyscallsIndex := make(map[uintptr]int, len(options.HotSyscalls)) + for i, sysno := range options.HotSyscalls { + if _, hasRule := allSyscallRuleActions[sysno]; !hasRule { + continue + } + if isTrivial[sysno] { + continue + } + if _, ok := hotNonTrivialSyscallsIndex[sysno]; ok { + continue + } + hotNonTrivialSyscallsIndex[sysno] = i + } + hotNonTrivialOrder := make([]uintptr, 0, len(hotNonTrivialSyscallsIndex)) + for sysno := range hotNonTrivialSyscallsIndex { + hotNonTrivialOrder = append(hotNonTrivialOrder, sysno) + } + sort.Slice(hotNonTrivialOrder, func(i, j int) bool { + return hotNonTrivialSyscallsIndex[hotNonTrivialOrder[i]] < hotNonTrivialSyscallsIndex[hotNonTrivialOrder[j]] + }) + + // Now split up the map and build the `orderedRuleSets`. + ors := orderedRuleSets{ + hotNonTrivial: make(map[uintptr]singleSyscallRuleSet), + hotNonTrivialOrder: hotNonTrivialOrder, + coldNonTrivial: make(map[uintptr]singleSyscallRuleSet), + trivial: make(map[uintptr]singleSyscallRuleSet), + } + for sysno, ruleActions := range allSyscallRuleActions { + _, hot := hotNonTrivialSyscallsIndex[sysno] + trivial := isTrivial[sysno] + var subMap map[uintptr]singleSyscallRuleSet + switch { + case trivial: + subMap = ors.trivial + case hot: + subMap = ors.hotNonTrivial + default: + subMap = ors.coldNonTrivial + } + subMap[sysno] = singleSyscallRuleSet{ + sysno: sysno, + vsyscall: vsyscallBySysno[sysno], + rules: ruleActions, + } + } + + // Log our findings. + if log.IsLogging(log.Debug) { + ors.log(log.Debugf) + } + + return ors, nil +} + +// log logs the set of seccomp rules to the given logger. +func (ors orderedRuleSets) log(logFn func(string, ...any)) { + logFn("Ordered seccomp rules:") + for _, sm := range []struct { + name string + m map[uintptr]singleSyscallRuleSet + }{ + {"Hot non-trivial", ors.hotNonTrivial}, + {"Cold non-trivial", ors.coldNonTrivial}, + {"Trivial", ors.trivial}, + } { + if len(sm.m) == 0 { + logFn(" %s syscalls: None.", sm.name) + continue + } + logFn(" %s syscalls:", sm.name) + orderedSysnos := make([]int, 0, len(sm.m)) + for sysno := range sm.m { + orderedSysnos = append(orderedSysnos, int(sysno)) + } + sort.Ints(orderedSysnos) + for _, sysno := range orderedSysnos { + logFn(" - %s", sm.m[uintptr(sysno)].String()) + } + } + logFn("End of ordered seccomp rules.") +} + +// render renders all rulesets in the given program. +func (ors orderedRuleSets) render(program *syscallProgram) error { + ls := &labelSet{prefix: string("ors")} + + // totalFrag wraps the entire output of the `render` function. + totalFrag := program.Record() + + // Load syscall number into register A. program.Stmt(bpf.Ld|bpf.Abs|bpf.W, seccompDataOffsetNR) - return root.traverse( + + // Keep track of which syscalls we've already looked for. + sysnosChecked := make(map[uintptr]struct{}) + + // First render hot syscalls linearly. + if len(ors.hotNonTrivialOrder) > 0 { + notHotLabel := ls.NewLabel() + // hotFrag wraps the "hot syscalls" part of the program. + // It must either return one of `hotActions`, or jump to `defaultLabel` if + // the syscall number matched but the vsyscall match failed, or + // `notHotLabel` if none of the hot syscall numbers matched. + hotFrag := program.Record() + possibleActions := ors.renderLinear(program, ls, sysnosChecked, ors.hotNonTrivial, ors.hotNonTrivialOrder, notHotLabel) + hotFrag.MustHaveJumpedToOrReturned([]label{notHotLabel, defaultLabel}, possibleActions) + program.Label(notHotLabel) + } + + // Now render the cold non-trivial rules as a binary search tree: + if len(ors.coldNonTrivial) > 0 { + frag := program.Record() + noSycallNumberMatch := ls.NewLabel() + possibleActions, err := ors.renderBST(program, ls, sysnosChecked, ors.coldNonTrivial, noSycallNumberMatch) + if err != nil { + return err + } + frag.MustHaveJumpedToOrReturned([]label{noSycallNumberMatch, defaultLabel}, possibleActions) + program.Label(noSycallNumberMatch) + } + + // Finally render the trivial rules as a binary search tree: + if len(ors.trivial) > 0 { + frag := program.Record() + noSycallNumberMatch := ls.NewLabel() + possibleActions, err := ors.renderBST(program, ls, sysnosChecked, ors.trivial, noSycallNumberMatch) + if err != nil { + return err + } + frag.MustHaveJumpedToOrReturned([]label{noSycallNumberMatch, defaultLabel}, possibleActions) + program.Label(noSycallNumberMatch) + } + program.JumpTo(defaultLabel) + + // Reached the end of the program. + // Independently verify the set of all possible actions. + allPossibleActions := make(map[linux.BPFAction]struct{}) + for _, mapping := range []map[uintptr]singleSyscallRuleSet{ + ors.hotNonTrivial, + ors.coldNonTrivial, + ors.trivial, + } { + for _, ssrs := range mapping { + for _, ra := range ssrs.rules { + allPossibleActions[ra.action] = struct{}{} + } + } + } + totalFrag.MustHaveJumpedToOrReturned([]label{defaultLabel}, allPossibleActions) + return nil +} + +// renderLinear renders linear search code that searches for syscall matches +// in the given order. It assumes the syscall number is loaded into register +// A. Rulesets for all syscall numbers in `order` must exist in `syscallMap`. +// It returns the list of possible actions the generated code may return. +// `alreadyChecked` will be updated with the syscalls that have been checked. +func (ors orderedRuleSets) renderLinear(program *syscallProgram, ls *labelSet, alreadyChecked map[uintptr]struct{}, syscallMap map[uintptr]singleSyscallRuleSet, order []uintptr, noSycallNumberMatch label) map[linux.BPFAction]struct{} { + allActions := make(map[linux.BPFAction]struct{}) + for _, sysno := range order { + ssrs, found := syscallMap[sysno] + if !found { + panic(fmt.Sprintf("syscall %d found in linear order but not map", sysno)) + } + nextSyscall := ls.NewLabel() + // sysnoFrag wraps the "statements about this syscall number" part of + // the program. It must either return one of the actions specified in + // that syscall number's rules (`sysnoActions`), or jump to + // `nextSyscall`. + sysnoFrag := program.Record() + sysnoActions := make(map[linux.BPFAction]struct{}) + for _, ra := range ssrs.rules { + sysnoActions[ra.action] = struct{}{} + allActions[ra.action] = struct{}{} + } + program.IfNot(bpf.Jmp|bpf.Jeq|bpf.K, uint32(ssrs.sysno), nextSyscall) + ssrs.Render(program, ls, defaultLabel) + sysnoFrag.MustHaveJumpedToOrReturned([]label{nextSyscall, defaultLabel}, sysnoActions) + program.Label(nextSyscall) + } + program.JumpTo(noSycallNumberMatch) + for _, sysno := range order { + alreadyChecked[sysno] = struct{}{} + } + return allActions +} + +// renderBST renders a binary search tree that searches the given map of +// syscalls. It assumes the syscall number is loaded into register A. +// It returns the list of possible actions the generated code may return. +// `alreadyChecked` will be updated with the syscalls that the BST has +// searched. +func (ors orderedRuleSets) renderBST(program *syscallProgram, ls *labelSet, alreadyChecked map[uintptr]struct{}, syscallMap map[uintptr]singleSyscallRuleSet, noSycallNumberMatch label) (map[linux.BPFAction]struct{}, error) { + possibleActions := make(map[linux.BPFAction]struct{}) + orderedSysnos := make([]uintptr, 0, len(syscallMap)) + for sysno, ruleActions := range syscallMap { + orderedSysnos = append(orderedSysnos, sysno) + for _, ra := range ruleActions.rules { + possibleActions[ra.action] = struct{}{} + } + } + sort.Slice(orderedSysnos, func(i, j int) bool { + return orderedSysnos[i] < orderedSysnos[j] + }) + frag := program.Record() + root := createBST(orderedSysnos) + root.root = true + if err := root.traverse( buildBSTProgram, knownRange{ lowerBoundExclusive: -1, // sysno fits in 32 bits, so this is definitely out of bounds: upperBoundExclusive: 1 << 32, + previouslyChecked: alreadyChecked, }, - rules, + syscallMap, program, - ) + noSycallNumberMatch, + ); err != nil { + return nil, err + } + frag.MustHaveJumpedToOrReturned([]label{noSycallNumberMatch, defaultLabel}, possibleActions) + for sysno := range syscallMap { + alreadyChecked[sysno] = struct{}{} + } + return possibleActions, nil } // createBST converts sorted syscall slice into a balanced BST. @@ -499,7 +876,7 @@ func createBST(syscalls []uintptr) *node { // index_50: // SYS_LISTEN(50), leaf // (A == 50) ? continue : goto defaultLabel // (args OK for SYS_LISTEN) ? return action : goto defaultLabel -func buildBSTProgram(n *node, rng knownRange, rules []RuleSet, program *syscallProgram) error { +func buildBSTProgram(n *node, rng knownRange, syscallMap map[uintptr]singleSyscallRuleSet, program *syscallProgram, searchFailed label) error { // Root node is never referenced by label, skip it. if !n.root { program.Label(n.label()) @@ -520,53 +897,23 @@ func buildBSTProgram(n *node, rng knownRange, rules []RuleSet, program *syscallP // If the previous BST nodes we traversed haven't fully established // that the current node's syscall value is exactly `sysno`, we still // need to verify it. - program.IfNot(bpf.Jmp|bpf.Jeq|bpf.K, uint32(sysno), defaultLabel) + program.IfNot(bpf.Jmp|bpf.Jeq|bpf.K, uint32(sysno), searchFailed) } program.JumpTo(checkArgsLabel) - nodeFrag.MustHaveJumpedTo(n.left.label(), n.right.label(), checkArgsLabel, defaultLabel) + nodeFrag.MustHaveJumpedTo(n.left.label(), n.right.label(), checkArgsLabel, searchFailed) program.Label(checkArgsLabel) ruleSetsFrag := program.Record() possibleActions := make(map[linux.BPFAction]struct{}) - for ruleSetIdx, rs := range rules { - rule, ok := rs.Rules.rules[sysno] - if !ok { - continue - } - ruleSetLabelSet := nodeLabelSet.Push(fmt.Sprintf("rs[%d]", ruleSetIdx), nodeLabelSet.NewLabel(), nodeLabelSet.NewLabel()) - ruleSetFrag := program.Record() - - // Emit a vsyscall check if this rule requires a - // Vsyscall match. This rule ensures that the top bit - // is set in the instruction pointer, which is where - // the vsyscall page will be mapped. - if rs.Vsyscall { - program.Stmt(bpf.Ld|bpf.Abs|bpf.W, seccompDataOffsetIPHigh) - program.IfNot(bpf.Jmp|bpf.Jset|bpf.K, vsyscallPageIPMask, ruleSetLabelSet.Mismatched()) - } - - // Add an argument check for these particular - // arguments. This will continue execution and - // check the next rule set. We need to ensure - // that at the very end, we insert a direct - // jump label for the unmatched case. - rule.Render(program, ruleSetLabelSet) - ruleSetFrag.MustHaveJumpedTo(ruleSetLabelSet.Matched(), ruleSetLabelSet.Mismatched()) - program.Label(ruleSetLabelSet.Matched()) - program.Ret(rs.Action) - possibleActions[rs.Action] = struct{}{} - ruleSetFrag.MustHaveJumpedToOrReturned( - []label{ruleSetLabelSet.Mismatched()}, // Either the ruleset mismatched... - map[linux.BPFAction]struct{}{ - rs.Action: struct{}{}, // ... or it returned its defined action. - }, - ) - program.Label(ruleSetLabelSet.Mismatched()) + for _, ra := range syscallMap[sysno].rules { + possibleActions[ra.action] = struct{}{} } - program.JumpTo(defaultLabel) + syscallMap[sysno].Render(program, nodeLabelSet, defaultLabel) ruleSetsFrag.MustHaveJumpedToOrReturned( - []label{defaultLabel}, // Either we jumped to the default label... - possibleActions, // ... or we returned one of the actions of the rulesets. + []label{ + defaultLabel, // Either we jumped to the default label (if the rules didn't match)... + }, + possibleActions, // ... or we returned one of the actions of the rulesets. ) return nil } @@ -594,36 +941,74 @@ func (n *node) label() label { type knownRange struct { lowerBoundExclusive int upperBoundExclusive int + + // alreadyChecked is a set of node values that were already checked + // earlier in the program (prior to the BST being built). + // It is *not* updated during BST traversal. + previouslyChecked map[uintptr]struct{} } -type traverseFunc func(*node, knownRange, []RuleSet, *syscallProgram) error +// withLowerBoundExclusive returns an updated `knownRange` with the given +// new exclusive lower bound. The actual exclusive lower bound of the +// returned `knownRange` may be higher, in case `previouslyChecked` covers +// more numbers. +func (kr knownRange) withLowerBoundExclusive(newLowerBoundExclusive int) knownRange { + nkr := knownRange{ + lowerBoundExclusive: newLowerBoundExclusive, + upperBoundExclusive: kr.upperBoundExclusive, + previouslyChecked: kr.previouslyChecked, + } + for ; nkr.lowerBoundExclusive < nkr.upperBoundExclusive; nkr.lowerBoundExclusive++ { + if _, ok := nkr.previouslyChecked[uintptr(nkr.lowerBoundExclusive+1)]; !ok { + break + } + } + return nkr +} -func (n *node) traverse(fn traverseFunc, rng knownRange, rules []RuleSet, program *syscallProgram) error { +// withUpperBoundExclusive returns an updated `knownRange` with the given +// new exclusive upper bound. The actual exclusive upper bound of the +// returned `knownRange` may be lower, in case `previouslyChecked` covers +// more numbers. +func (kr knownRange) withUpperBoundExclusive(newUpperBoundExclusive int) knownRange { + nkr := knownRange{ + lowerBoundExclusive: kr.lowerBoundExclusive, + upperBoundExclusive: newUpperBoundExclusive, + previouslyChecked: kr.previouslyChecked, + } + for ; nkr.lowerBoundExclusive < nkr.upperBoundExclusive; nkr.upperBoundExclusive-- { + if _, ok := nkr.previouslyChecked[uintptr(nkr.upperBoundExclusive-1)]; !ok { + break + } + } + return nkr +} + +// traverseFunc is called as the BST is traversed. +type traverseFunc func(*node, knownRange, map[uintptr]singleSyscallRuleSet, *syscallProgram, label) error + +func (n *node) traverse(fn traverseFunc, kr knownRange, syscallMap map[uintptr]singleSyscallRuleSet, program *syscallProgram, searchFailed label) error { if n == nil { return nil } - if err := fn(n, rng, rules, program); err != nil { + if err := fn(n, kr, syscallMap, program, searchFailed); err != nil { return err } if err := n.left.traverse( fn, - knownRange{ - lowerBoundExclusive: rng.lowerBoundExclusive, - upperBoundExclusive: int(n.value), - }, - rules, + kr.withUpperBoundExclusive(int(n.value)), + syscallMap, program, + searchFailed, ); err != nil { return err } return n.right.traverse( fn, - knownRange{ - lowerBoundExclusive: int(n.value), - upperBoundExclusive: rng.upperBoundExclusive, - }, - rules, + kr.withLowerBoundExclusive(int(n.value)), + syscallMap, program, + searchFailed, ) } diff --git a/pkg/seccomp/seccomp_test.go b/pkg/seccomp/seccomp_test.go index e3bb4e4a1..79f7e8324 100644 --- a/pkg/seccomp/seccomp_test.go +++ b/pkg/seccomp/seccomp_test.go @@ -29,6 +29,7 @@ import ( "testing" "time" + "golang.org/x/sys/unix" "gvisor.dev/gvisor/pkg/abi/linux" "gvisor.dev/gvisor/pkg/bpf" ) @@ -1377,3 +1378,579 @@ func TestOptimizeSyscallRule(t *testing.T) { }) } } + +func TestOrderRuleSets(t *testing.T) { + for _, test := range []struct { + name string + ruleSets []RuleSet + options ProgramOptions + want orderedRuleSets + wantErr bool + }{ + { + name: "no RuleSets", + options: DefaultProgramOptions(), + want: orderedRuleSets{}, + }, + { + name: "inconsistent vsyscall", + options: DefaultProgramOptions(), + ruleSets: []RuleSet{ + { + Rules: NewSyscallRules().Add(unix.SYS_READ, MatchAll{}), + Action: linux.SECCOMP_RET_TRACE, + Vsyscall: false, + }, + { + Rules: NewSyscallRules().Add(unix.SYS_READ, MatchAll{}), + Action: linux.SECCOMP_RET_TRACE, + Vsyscall: true, + }, + }, + wantErr: true, + }, + { + name: "single trivial rule", + ruleSets: []RuleSet{ + { + Rules: NewSyscallRules().Add(unix.SYS_READ, MatchAll{}), + Action: linux.SECCOMP_RET_TRACE, + Vsyscall: false, + }, + }, + options: DefaultProgramOptions(), + want: orderedRuleSets{ + trivial: map[uintptr]singleSyscallRuleSet{ + unix.SYS_READ: { + sysno: unix.SYS_READ, + rules: []syscallRuleAction{ + { + rule: MatchAll{}, + action: linux.SECCOMP_RET_TRACE, + }, + }, + vsyscall: false, + }, + }, + }, + }, + { + name: "hot single trivial rule still ends up as trivial", + ruleSets: []RuleSet{ + { + Rules: NewSyscallRules().Add(unix.SYS_READ, MatchAll{}), + Action: linux.SECCOMP_RET_TRACE, + Vsyscall: false, + }, + }, + options: ProgramOptions{ + HotSyscalls: []uintptr{unix.SYS_READ}, + }, + want: orderedRuleSets{ + trivial: map[uintptr]singleSyscallRuleSet{ + unix.SYS_READ: { + sysno: unix.SYS_READ, + rules: []syscallRuleAction{ + { + rule: MatchAll{}, + action: linux.SECCOMP_RET_TRACE, + }, + }, + vsyscall: false, + }, + }, + }, + }, + { + name: "hot single non-trivial rule", + ruleSets: []RuleSet{ + { + Rules: NewSyscallRules().Add( + unix.SYS_READ, PerArg{EqualTo(0)}, + ), + Action: linux.SECCOMP_RET_TRACE, + Vsyscall: false, + }, + }, + options: ProgramOptions{ + HotSyscalls: []uintptr{unix.SYS_READ}, + }, + want: orderedRuleSets{ + hotNonTrivial: map[uintptr]singleSyscallRuleSet{ + unix.SYS_READ: { + sysno: unix.SYS_READ, + rules: []syscallRuleAction{ + { + rule: PerArg{EqualTo(0)}, + action: linux.SECCOMP_RET_TRACE, + }, + }, + vsyscall: false, + }, + }, + hotNonTrivialOrder: []uintptr{unix.SYS_READ}, + }, + }, + { + name: "hot rule ordering", + ruleSets: []RuleSet{ + { + Rules: NewSyscallRules().Add( + unix.SYS_FLOCK, PerArg{EqualTo(0)}, + ), + Action: linux.SECCOMP_RET_TRACE, + Vsyscall: false, + }, + { + Rules: NewSyscallRules().Add( + unix.SYS_WRITE, PerArg{EqualTo(1)}, + ).Add( + unix.SYS_READ, PerArg{EqualTo(2)}, + ), + Action: linux.SECCOMP_RET_TRACE, + Vsyscall: false, + }, + }, + options: ProgramOptions{ + HotSyscalls: []uintptr{ + unix.SYS_READ, + unix.SYS_WRITE, + unix.SYS_FLOCK, + }, + }, + want: orderedRuleSets{ + hotNonTrivial: map[uintptr]singleSyscallRuleSet{ + unix.SYS_READ: { + sysno: unix.SYS_READ, + rules: []syscallRuleAction{ + { + rule: PerArg{EqualTo(2)}, + action: linux.SECCOMP_RET_TRACE, + }, + }, + vsyscall: false, + }, + unix.SYS_WRITE: { + sysno: unix.SYS_WRITE, + rules: []syscallRuleAction{ + { + rule: PerArg{EqualTo(1)}, + action: linux.SECCOMP_RET_TRACE, + }, + }, + vsyscall: false, + }, + unix.SYS_FLOCK: { + sysno: unix.SYS_FLOCK, + rules: []syscallRuleAction{ + { + rule: PerArg{EqualTo(0)}, + action: linux.SECCOMP_RET_TRACE, + }, + }, + vsyscall: false, + }, + }, + hotNonTrivialOrder: []uintptr{ + unix.SYS_READ, + unix.SYS_WRITE, + unix.SYS_FLOCK, + }, + }, + }, + { + name: "cold single non-trivial non-vsyscall rule", + ruleSets: []RuleSet{ + { + Rules: NewSyscallRules().Add( + unix.SYS_READ, PerArg{EqualTo(0)}, + ), + Action: linux.SECCOMP_RET_TRACE, + Vsyscall: false, + }, + }, + options: ProgramOptions{ + HotSyscalls: []uintptr{unix.SYS_WRITE}, + }, + want: orderedRuleSets{ + coldNonTrivial: map[uintptr]singleSyscallRuleSet{ + unix.SYS_READ: { + sysno: unix.SYS_READ, + rules: []syscallRuleAction{ + { + rule: PerArg{EqualTo(0)}, + action: linux.SECCOMP_RET_TRACE, + }, + }, + vsyscall: false, + }, + }, + hotNonTrivialOrder: nil, // Empty + }, + }, + { + name: "cold single non-trivial yes-vsyscall rule", + ruleSets: []RuleSet{ + { + Rules: NewSyscallRules().Add( + unix.SYS_READ, PerArg{EqualTo(0)}, + ), + Action: linux.SECCOMP_RET_TRACE, + Vsyscall: true, + }, + }, + options: ProgramOptions{ + HotSyscalls: []uintptr{unix.SYS_WRITE}, + }, + want: orderedRuleSets{ + coldNonTrivial: map[uintptr]singleSyscallRuleSet{ + unix.SYS_READ: { + sysno: unix.SYS_READ, + rules: []syscallRuleAction{ + { + rule: PerArg{EqualTo(0)}, + action: linux.SECCOMP_RET_TRACE, + }, + }, + vsyscall: true, + }, + }, + hotNonTrivialOrder: nil, // Empty + }, + }, + { + name: "all rule types at once", + ruleSets: []RuleSet{ + { + Rules: NewSyscallRules().Add( + // Hot, non-vsyscall, trivial + unix.SYS_READ, MatchAll{}, + ).Add( + // Hot, non-vsyscall, non-trivial + unix.SYS_WRITE, PerArg{EqualTo(4)}, + ).Add( + // Hot, non-vsyscall, non-trivial + unix.SYS_FLOCK, PerArg{EqualTo(131)}, + ), + Action: linux.SECCOMP_RET_TRACE, + Vsyscall: false, + }, + { + Rules: NewSyscallRules().Add( + // Hot, vsyscall, trivial (after optimizations) + unix.SYS_FUTEX, PerArg{AnyValue{}}, + ).Add( + // Hot, vsyscall, non-trivial + unix.SYS_GETPID, PerArg{EqualTo(20)}, + ).Add( + // Cold, vsyscall, trivial (after optimizations) + unix.SYS_GETTIMEOFDAY, PerArg{AnyValue{}}, + ).Add( + // Cold, vsyscall, non-trivial + unix.SYS_MINCORE, PerArg{EqualTo(78)}, + ), + Action: linux.SECCOMP_RET_ERRNO, + Vsyscall: true, + }, + { + Rules: NewSyscallRules().Add( + // Hot, non-vsyscall, trivial (after optimizations) + unix.SYS_CLOSE, PerArg{AnyValue{}}, + ).Add( + // Hot, non-vsyscall, non-trivial + unix.SYS_OPENAT, PerArg{EqualTo(463)}, + ).Add( + // Cold, non-vsyscall, non-trivial + unix.SYS_LINKAT, PerArg{EqualTo(471)}, + ).Add( + // Cold, non-vsyscall, trivial here but + // a later RuleSet will make it non-trivial + // overall. + unix.SYS_UNLINKAT, MatchAll{}, + ).Add( + // Cold, non-vsyscall, trivial and another + // later RuleSet will keep it trivial later. + unix.SYS_CHDIR, MatchAll{}, + ), + Action: linux.SECCOMP_RET_KILL_THREAD, + Vsyscall: false, + }, + { + Rules: NewSyscallRules().Add( + // Hot, non-vsyscall, non-trivial + unix.SYS_OPENAT, PerArg{EqualTo(463463)}, + ).Add( + // Cold, non-vsyscall, non-trivial + unix.SYS_LINKAT, PerArg{EqualTo(471471)}, + ).Add( + // Cold, non-vsyscall, no longer trivial. + unix.SYS_UNLINKAT, PerArg{EqualTo(472)}, + ).Add( + // Cold, non-vsyscall, remains trivial. + unix.SYS_FCHDIR, PerArg{AnyValue{}}, + ), + Action: linux.SECCOMP_RET_KILL_PROCESS, + Vsyscall: false, + }, + { + Rules: NewSyscallRules().Add( + // Cold, non-vsyscall, adds to previous rule + // after optimizations because it has the + // same action. + unix.SYS_UNLINKAT, PerArg{EqualTo(472472)}, + ), + Action: linux.SECCOMP_RET_KILL_PROCESS, + Vsyscall: false, + }, + { + Rules: NewSyscallRules().Add( + // Cold, non-vsyscall, does not add to + // previous rule because it has a different + // action. + unix.SYS_UNLINKAT, PerArg{EqualTo(472472472)}, + ), + Action: linux.SECCOMP_RET_TRAP, + Vsyscall: false, + }, + }, + options: ProgramOptions{ + HotSyscalls: []uintptr{ + unix.SYS_READ, + unix.SYS_WRITE, + unix.SYS_FLOCK, + unix.SYS_CLOSE, + unix.SYS_OPENAT, + unix.SYS_GETPID, + // Note: not used in any RuleSet, so it should not + // appear in `want`: + unix.SYS_GETTID, + }, + }, + want: orderedRuleSets{ + hotNonTrivial: map[uintptr]singleSyscallRuleSet{ + unix.SYS_WRITE: { + sysno: unix.SYS_WRITE, + rules: []syscallRuleAction{ + { + rule: PerArg{EqualTo(4)}, + action: linux.SECCOMP_RET_TRACE, + }, + }, + vsyscall: false, + }, + unix.SYS_FLOCK: { + sysno: unix.SYS_FLOCK, + rules: []syscallRuleAction{ + { + rule: PerArg{EqualTo(131)}, + action: linux.SECCOMP_RET_TRACE, + }, + }, + vsyscall: false, + }, + unix.SYS_GETPID: { + sysno: unix.SYS_GETPID, + rules: []syscallRuleAction{ + { + rule: PerArg{EqualTo(20)}, + action: linux.SECCOMP_RET_ERRNO, + }, + }, + vsyscall: true, + }, + unix.SYS_OPENAT: { + sysno: unix.SYS_OPENAT, + rules: []syscallRuleAction{ + { + rule: PerArg{EqualTo(463)}, + action: linux.SECCOMP_RET_KILL_THREAD, + }, + { + rule: PerArg{EqualTo(463463)}, + action: linux.SECCOMP_RET_KILL_PROCESS, + }, + }, + vsyscall: false, + }, + }, + coldNonTrivial: map[uintptr]singleSyscallRuleSet{ + unix.SYS_LINKAT: { + sysno: unix.SYS_LINKAT, + rules: []syscallRuleAction{ + { + rule: PerArg{EqualTo(471)}, + action: linux.SECCOMP_RET_KILL_THREAD, + }, + { + rule: PerArg{EqualTo(471471)}, + action: linux.SECCOMP_RET_KILL_PROCESS, + }, + }, + vsyscall: false, + }, + unix.SYS_UNLINKAT: { + sysno: unix.SYS_UNLINKAT, + rules: []syscallRuleAction{ + { + rule: MatchAll{}, + action: linux.SECCOMP_RET_KILL_THREAD, + }, + { + rule: Or{ + PerArg{EqualTo(472)}, + PerArg{EqualTo(472472)}, + }, + action: linux.SECCOMP_RET_KILL_PROCESS, + }, + { + rule: PerArg{EqualTo(472472472)}, + action: linux.SECCOMP_RET_TRAP, + }, + }, + vsyscall: false, + }, + unix.SYS_MINCORE: { + sysno: unix.SYS_MINCORE, + rules: []syscallRuleAction{ + { + rule: PerArg{EqualTo(78)}, + action: linux.SECCOMP_RET_ERRNO, + }, + }, + vsyscall: true, + }, + unix.SYS_FUTEX: { + sysno: unix.SYS_FUTEX, + rules: []syscallRuleAction{ + { + rule: MatchAll{}, + action: linux.SECCOMP_RET_ERRNO, + }, + }, + vsyscall: true, + }, + unix.SYS_GETTIMEOFDAY: { + sysno: unix.SYS_GETTIMEOFDAY, + rules: []syscallRuleAction{ + { + rule: MatchAll{}, + action: linux.SECCOMP_RET_ERRNO, + }, + }, + vsyscall: true, + }, + }, + trivial: map[uintptr]singleSyscallRuleSet{ + unix.SYS_READ: { + sysno: unix.SYS_READ, + rules: []syscallRuleAction{ + { + rule: MatchAll{}, + action: linux.SECCOMP_RET_TRACE, + }, + }, + vsyscall: false, + }, + unix.SYS_CHDIR: { + sysno: unix.SYS_CHDIR, + rules: []syscallRuleAction{ + { + rule: MatchAll{}, + action: linux.SECCOMP_RET_KILL_THREAD, + }, + }, + vsyscall: false, + }, + unix.SYS_CLOSE: { + sysno: unix.SYS_CLOSE, + rules: []syscallRuleAction{ + { + rule: MatchAll{}, + action: linux.SECCOMP_RET_KILL_THREAD, + }, + }, + vsyscall: false, + }, + unix.SYS_FCHDIR: { + sysno: unix.SYS_FCHDIR, + rules: []syscallRuleAction{ + { + rule: MatchAll{}, + action: linux.SECCOMP_RET_KILL_PROCESS, + }, + }, + vsyscall: false, + }, + }, + hotNonTrivialOrder: []uintptr{ + // SYS_READ becomes trivial so it does not show up here. + unix.SYS_WRITE, + unix.SYS_FLOCK, + // SYS_CLOSE also becomes trivial. + unix.SYS_OPENAT, + unix.SYS_GETPID, + // SYS_GETTID does not show up in any RuleSet so it + // also does not show up here. + }, + }, + }, + } { + t.Run(test.name, func(t *testing.T) { + got, gotErr := orderRuleSets(test.ruleSets, test.options) + if (gotErr != nil) != test.wantErr { + t.Errorf("got error: %v, want error: %v", gotErr, test.wantErr) + } + if gotErr != nil || t.Failed() { + return + } + // Replace empty maps with nil for simpler comparison output. + for _, ors := range []*orderedRuleSets{&got, &test.want} { + if len(ors.hotNonTrivial) == 0 { + ors.hotNonTrivial = nil + } + if len(ors.hotNonTrivialOrder) == 0 { + ors.hotNonTrivialOrder = nil + } + if len(ors.coldNonTrivial) == 0 { + ors.coldNonTrivial = nil + } + if len(ors.trivial) == 0 { + ors.trivial = nil + } + } + // Run optimizers on all rules of `test.want`. + for _, m := range []map[uintptr]singleSyscallRuleSet{ + test.want.hotNonTrivial, + test.want.coldNonTrivial, + test.want.trivial, + } { + for sysno, ssrs := range m { + for i, r := range ssrs.rules { + ssrs.rules[i].rule = optimizeSyscallRule(r.rule) + } + m[sysno] = ssrs + } + } + if !reflect.DeepEqual(got, test.want) { + t.Errorf("got orderedRuleSets:\n%v\nwant orderedRuleSets:\n%v\n", got, test.want) + t.Error("got:") + got.log(t.Errorf) + t.Errorf("want:") + test.want.log(t.Errorf) + return + } + // Log the orderedRuleSet, if only to make sure the log function + // actually works and doesn't panic: + t.Log("orderedRuleSets matched expectations:") + got.log(t.Logf) + + // Attempt to render the `orderedRuleSets`. + // We don't check the instructions it generates, but the rendering + // code checks itself and panics if it fails an assertion. + program := &syscallProgram{bpf.NewProgramBuilder()} + if err := got.render(program); err != nil { + t.Fatalf("got unexpected error while rendering: %v", err) + } + }) + } +}