From 9aa05f01e038a361a91c13af898e2936aef57aad Mon Sep 17 00:00:00 2001 From: Etienne Perot Date: Thu, 5 Oct 2023 11:50:29 -0700 Subject: [PATCH] BPF program builder: Add support for recording/analyzing fragment outcomes. This adds a new `Record` function to `bpf.ProgramBuilder`, which returns a function to stop recording that returns the "fragment" of the program made of the instructions that were added between the time `Record` was called and the time the stop function was called. This fragment can in turn be interrogated for which `Outcomes` may happen from executing it: returning a value, jumping to a label, jumping away from the fragment, falling through. This is useful while building complex BPF programs with nested rules. By recording instructions added by a possibly-nested set of rules (the final outcome of which is to jump to a known set of labels), we can now actually verify the assertion that the instructions that were added indeed end up jumping to one of the expected labels, and nothing else. This is useful not just for safety but also optimization purposes. In an upcoming refactor to argument matching code, I plan to add a "value matcher" interface that renders rules that verify the value of the `A` register. Some matchers may need to modify the `A` register in order to work, but others don't. By checking whether the set of instructions modifies `A` or not, the higher-level code can determine whether or not it needs to add code to reload the value of the `A` register or not before moving on to the next matcher. PiperOrigin-RevId: 571087694 --- pkg/bpf/bpf.go | 24 +++ pkg/bpf/program_builder.go | 178 ++++++++++++++++++++-- pkg/bpf/program_builder_test.go | 258 ++++++++++++++++++++++++++++++++ 3 files changed, 446 insertions(+), 14 deletions(-) diff --git a/pkg/bpf/bpf.go b/pkg/bpf/bpf.go index 8d8346114..107f4935f 100644 --- a/pkg/bpf/bpf.go +++ b/pkg/bpf/bpf.go @@ -163,3 +163,27 @@ func (ins Instruction) IsConditionalJump() bool { func (ins Instruction) IsUnconditionalJump() bool { return ins.IsJump() && ins.OpCode&jmpMask == Ja } + +// JumpOffset is a possible jump offset that an instruction may jump to. +type JumpOffset struct { + // Type is the type of jump that an instruction may execute. + Type JumpType + + // Offset is the number of instructions that the jump skips over. + Offset uint32 +} + +// JumpOffsets returns the set of instruction offsets that this instruction +// may jump to. Returns a nil slice if this is not a jump instruction. +func (ins Instruction) JumpOffsets() []JumpOffset { + if !ins.IsJump() { + return nil + } + if ins.IsConditionalJump() { + return []JumpOffset{ + {JumpTrue, uint32(ins.JumpIfTrue)}, + {JumpFalse, uint32(ins.JumpIfFalse)}, + } + } + return []JumpOffset{{JumpDirect, ins.K}} +} diff --git a/pkg/bpf/program_builder.go b/pkg/bpf/program_builder.go index 8839c61cf..584cbef90 100644 --- a/pkg/bpf/program_builder.go +++ b/pkg/bpf/program_builder.go @@ -17,6 +17,8 @@ package bpf import ( "fmt" "math" + "sort" + "strings" ) const ( @@ -56,12 +58,14 @@ type label struct { target int } -type jmpType int +// JumpType is the type of jump target that an instruction may use. +type JumpType int +// Types of jump that an instruction may use. const ( - jDirect jmpType = iota - jTrue - jFalse + JumpDirect JumpType = iota + JumpTrue + JumpFalse ) // source contains information about a single reference to a label. @@ -71,7 +75,7 @@ type source struct { // True if label reference is in the 'jump if true' part of the jump. // False if label reference is in the 'jump if false' part of the jump. - jt jmpType + jt JumpType } // AddStmt adds a new statement to the program. @@ -86,26 +90,26 @@ func (b *ProgramBuilder) AddJump(code uint16, k uint32, jt, jf uint8) { // AddDirectJumpLabel adds a new jump to the program where is labelled. func (b *ProgramBuilder) AddDirectJumpLabel(labelName string) { - b.addLabelSource(labelName, jDirect) + b.addLabelSource(labelName, JumpDirect) b.AddJump(Jmp|Ja, labelDirectTarget, 0, 0) } // AddJumpTrueLabel adds a new jump to the program where 'jump if true' is a label. func (b *ProgramBuilder) AddJumpTrueLabel(code uint16, k uint32, jtLabel string, jf uint8) { - b.addLabelSource(jtLabel, jTrue) + b.addLabelSource(jtLabel, JumpTrue) b.AddJump(code, k, labelTarget, jf) } // AddJumpFalseLabel adds a new jump to the program where 'jump if false' is a label. func (b *ProgramBuilder) AddJumpFalseLabel(code uint16, k uint32, jt uint8, jfLabel string) { - b.addLabelSource(jfLabel, jFalse) + b.addLabelSource(jfLabel, JumpFalse) b.AddJump(code, k, jt, labelTarget) } // AddJumpLabels adds a new jump to the program where both jump targets are labels. func (b *ProgramBuilder) AddJumpLabels(code uint16, k uint32, jtLabel, jfLabel string) { - b.addLabelSource(jtLabel, jTrue) - b.addLabelSource(jfLabel, jFalse) + b.addLabelSource(jtLabel, JumpTrue) + b.addLabelSource(jfLabel, JumpFalse) b.AddJump(code, k, labelTarget, labelTarget) } @@ -139,7 +143,7 @@ func (b *ProgramBuilder) Instructions() ([]Instruction, error) { return b.instructions, nil } -func (b *ProgramBuilder) addLabelSource(labelName string, t jmpType) { +func (b *ProgramBuilder) addLabelSource(labelName string, t JumpType) { l, ok := b.labels[labelName] if !ok { l = &label{sources: make([]source, 0), target: -1} @@ -170,7 +174,7 @@ func (b *ProgramBuilder) resolveLabels() error { offset := v.target - s.line - 1 // Sets offset into jump instruction. switch s.jt { - case jDirect: + case JumpDirect: if offset > labelDirectTarget { return fmt.Errorf("jump offset to label '%v' is too large: %v, inst: %v, lineno: %v", key, offset, inst, s.line) } @@ -178,7 +182,7 @@ func (b *ProgramBuilder) resolveLabels() error { return fmt.Errorf("jump target is not a label") } inst.K = uint32(offset) - case jTrue: + case JumpTrue: if offset > labelTarget { return fmt.Errorf("jump offset to label '%v' is too large: %v, inst: %v, lineno: %v", key, offset, inst, s.line) } @@ -186,7 +190,7 @@ func (b *ProgramBuilder) resolveLabels() error { return fmt.Errorf("jump target is not a label") } inst.JumpIfTrue = uint8(offset) - case jFalse: + case JumpFalse: if offset > labelTarget { return fmt.Errorf("jump offset to label '%v' is too large: %v, inst: %v, lineno: %v", key, offset, inst, s.line) } @@ -202,3 +206,149 @@ func (b *ProgramBuilder) resolveLabels() error { b.labels = map[string]*label{} return nil } + +// ProgramFragment is a set of not-compiled instructions that were added to +// a ProgramBuilder from the moment the `Record` function was called on it. +type ProgramFragment struct { + // b is a reference to the ProgramBuilder that this is a fragment from. + b *ProgramBuilder + + // fromPC is the index of the first instruction that was recorded. + // If no instruction was recorded, this index will be equal to `toPC`. + fromPC int + + // toPC is the index *after* the last instruction that was recorded. + // This means that right after recording, the program will not have + // any instruction at index `toPC`. + toPC int +} + +// Record starts recording the instructions being added to the ProgramBuilder +// until the returned function is called. +// The returned function returns a ProgramFragment which represents the +// recorded instructions. It may be called repeatedly. +func (b *ProgramBuilder) Record() func() ProgramFragment { + currentPC := len(b.instructions) + return func() ProgramFragment { + return ProgramFragment{ + b: b, + fromPC: currentPC, + toPC: len(b.instructions), + } + } +} + +// String returns a string version of the fragment. +func (f ProgramFragment) String() string { + return fmt.Sprintf("fromPC=%d toPC=%d", f.fromPC, f.toPC) +} + +// FragmentOutcomes represents the set of outcomes that a ProgramFragment +// execution may result into. +type FragmentOutcomes struct { + // MayFallThrough is true if executing the fragment may cause it to start + // executing the program instruction that comes right after the last + // instruction in this fragment (i.e. at `Fragment.toPC`). + MayFallThrough bool + + // MayJumpToKnownOffsetBeyondFragment is true if executing the fragment may + // jump to a fixed offset (or resolved label) that is not within the range + // of the fragment itself, nor does it point to the instruction that would + // come right after this fragment. + // If the fragment jumps to an unresolved label, this will instead be + // indicated in `MayJumpToUnresolvedLabels`. + MayJumpToKnownOffsetBeyondFragment bool + + // MayJumpToUnresolvedLabels is the set of named labels that have not yet + // been added to the program (the labels are not resolvable) but that the + // fragment may jump to. + MayJumpToUnresolvedLabels map[string]struct{} + + // MayReturn is true if executing the fragment may cause a return statement + // to be executed. + MayReturn bool +} + +// String returns a list of possible human-readable outcomes. +func (o FragmentOutcomes) String() string { + var s []string + if o.MayJumpToKnownOffsetBeyondFragment { + s = append(s, "may jump to known offset beyond fragment") + } + if o.MayJumpToUnresolvedLabels != nil { + sortedLabels := make([]string, 0, len(o.MayJumpToUnresolvedLabels)) + for lbl := range o.MayJumpToUnresolvedLabels { + sortedLabels = append(sortedLabels, lbl) + } + sort.Strings(sortedLabels) + for _, lbl := range sortedLabels { + s = append(s, fmt.Sprintf("may jump to unresolved label %q", lbl)) + } + } + if o.MayFallThrough { + s = append(s, "may fall through") + } + if o.MayReturn { + s = append(s, "may return") + } + if len(s) == 0 { + return "no outcomes (this should never happen)" + } + return strings.Join(s, ", ") +} + +// Outcomes returns the set of possible outcomes that executing this fragment +// may result into. +func (f ProgramFragment) Outcomes() FragmentOutcomes { + if f.fromPC == f.toPC { + // No instructions, this just falls through. + return FragmentOutcomes{ + MayFallThrough: true, + } + } + outcomes := FragmentOutcomes{ + MayJumpToUnresolvedLabels: make(map[string]struct{}), + } + for pc := f.fromPC; pc < f.toPC; pc++ { + ins := f.b.instructions[pc] + isLastInstruction := pc == f.toPC-1 + switch ins.OpCode & instructionClassMask { + case Ret: + outcomes.MayReturn = true + case Jmp: + for _, offset := range ins.JumpOffsets() { + var foundLabelName string + var foundLabel *label + for labelName, label := range f.b.labels { + for _, s := range label.sources { + if s.jt == offset.Type && s.line == pc { + foundLabelName = labelName + foundLabel = label + break + } + } + } + if foundLabel != nil && foundLabel.target == -1 { + outcomes.MayJumpToUnresolvedLabels[foundLabelName] = struct{}{} + continue + } + var target int + if foundLabel != nil { + target = foundLabel.target + } else { + target = pc + int(offset.Offset) + 1 + } + if target == f.toPC { + outcomes.MayFallThrough = true + } else if target > f.toPC { + outcomes.MayJumpToKnownOffsetBeyondFragment = true + } + } + default: + if isLastInstruction { + outcomes.MayFallThrough = true + } + } + } + return outcomes +} diff --git a/pkg/bpf/program_builder_test.go b/pkg/bpf/program_builder_test.go index cb5bb90a6..77583834b 100644 --- a/pkg/bpf/program_builder_test.go +++ b/pkg/bpf/program_builder_test.go @@ -16,6 +16,7 @@ package bpf import ( "fmt" + "reflect" "testing" ) @@ -181,3 +182,260 @@ func TestProgramBuilderJumpBackwards(t *testing.T) { t.Errorf("Instructions() should have failed") } } + +func TestProgramBuilderOutcomes(t *testing.T) { + p := NewProgramBuilder() + getOverallFragment := p.Record() + fixup := func(f FragmentOutcomes) FragmentOutcomes { + if f.MayJumpToUnresolvedLabels == nil { + f.MayJumpToUnresolvedLabels = map[string]struct{}{} + } + return f + } + for _, test := range []struct { + // Name of the sub-test. + name string + + // Function that adds statements to `p`. + build func() + + // Expected outcomes from recording the instructions added + // by `build` alone. + wantLocal FragmentOutcomes + + // Expected outcomes from recording the instructions added + // to the program since the test began. + wantOverall FragmentOutcomes + }{ + { + name: "empty program", + build: func() {}, + wantLocal: FragmentOutcomes{MayFallThrough: true}, + wantOverall: FragmentOutcomes{MayFallThrough: true}, + }, + { + name: "simple instruction", + build: func() { + p.AddStmt(Ld|Abs|W, 10) + }, + wantLocal: FragmentOutcomes{MayFallThrough: true}, + wantOverall: FragmentOutcomes{MayFallThrough: true}, + }, + { + name: "jump to unresolved label", + build: func() { + p.AddDirectJumpLabel("label1") + }, + wantLocal: FragmentOutcomes{ + MayJumpToUnresolvedLabels: map[string]struct{}{ + "label1": struct{}{}, + }, + }, + wantOverall: FragmentOutcomes{ + MayJumpToUnresolvedLabels: map[string]struct{}{ + "label1": struct{}{}, + }, + }, + }, + { + name: "another simple load so may fall through again", + build: func() { + p.AddStmt(Ld|Abs|W, 10) + }, + wantLocal: FragmentOutcomes{ + MayFallThrough: true, + }, + wantOverall: FragmentOutcomes{ + MayJumpToUnresolvedLabels: map[string]struct{}{ + "label1": struct{}{}, + }, + MayFallThrough: true, + }, + }, + { + name: "resolve label1", + build: func() { + p.AddLabel("label1") + }, + wantLocal: FragmentOutcomes{ + MayFallThrough: true, + }, + wantOverall: FragmentOutcomes{ + MayFallThrough: true, + }, + }, + { + name: "populate instruction at label1", + build: func() { + p.AddStmt(Ld|Abs|W, 10) + }, + wantLocal: FragmentOutcomes{ + MayFallThrough: true, + }, + wantOverall: FragmentOutcomes{ + MayFallThrough: true, + }, + }, + { + name: "conditional jump to two unresolved labels", + build: func() { + p.AddJumpLabels(Jmp|Jeq|K, 1337, "truelabel", "falselabel") + }, + wantLocal: FragmentOutcomes{ + MayJumpToUnresolvedLabels: map[string]struct{}{ + "truelabel": struct{}{}, + "falselabel": struct{}{}, + }, + }, + wantOverall: FragmentOutcomes{ + MayJumpToUnresolvedLabels: map[string]struct{}{ + "truelabel": struct{}{}, + "falselabel": struct{}{}, + }, + }, + }, + { + name: "resolve truelabel only", + build: func() { + p.AddLabel("truelabel") + }, + wantLocal: FragmentOutcomes{ + MayFallThrough: true, + }, + wantOverall: FragmentOutcomes{ + MayJumpToUnresolvedLabels: map[string]struct{}{ + "falselabel": struct{}{}, + }, + MayFallThrough: true, + }, + }, + { + name: "jump one beyond end of program", + build: func() { + p.AddJump(Jmp|Ja, 1, 0, 0) + }, + wantLocal: FragmentOutcomes{ + MayJumpToKnownOffsetBeyondFragment: true, + }, + wantOverall: FragmentOutcomes{ + MayJumpToUnresolvedLabels: map[string]struct{}{ + "falselabel": struct{}{}, + }, + MayJumpToKnownOffsetBeyondFragment: true, + }, + }, + { + name: "add return", + build: func() { + p.AddStmt(Ret|K, 1337) + }, + wantLocal: FragmentOutcomes{ + MayReturn: true, + }, + wantOverall: FragmentOutcomes{ + MayJumpToUnresolvedLabels: map[string]struct{}{ + "falselabel": struct{}{}, + }, + MayFallThrough: true, // From jump in previous test. + MayReturn: true, + }, + }, + { + name: "add another instruction after return", + build: func() { + p.AddStmt(Ld|Abs|W, 10) + }, + wantLocal: FragmentOutcomes{ + MayFallThrough: true, + }, + wantOverall: FragmentOutcomes{ + MayJumpToUnresolvedLabels: map[string]struct{}{ + "falselabel": struct{}{}, + }, + MayReturn: true, + MayFallThrough: true, + }, + }, + { + name: "zero-instruction jump counts as fallthrough", + build: func() { + p.AddJump(Jmp|Ja, 0, 0, 0) + }, + wantLocal: FragmentOutcomes{ + MayFallThrough: true, + }, + wantOverall: FragmentOutcomes{ + MayJumpToUnresolvedLabels: map[string]struct{}{ + "falselabel": struct{}{}, + }, + MayReturn: true, + MayFallThrough: true, + }, + }, + { + name: "non-zero-instruction jumps that points to end of fragment also counts as fallthrough", + build: func() { + p.AddJump(Jmp|Jeq|K, 42, 3, 1) + p.AddJump(Jmp|Ja, 2, 0, 0) + p.AddStmt(Ld|Abs|W, 11) + p.AddStmt(Ld|Abs|W, 12) + }, + wantLocal: FragmentOutcomes{ + MayFallThrough: true, + }, + wantOverall: FragmentOutcomes{ + MayJumpToUnresolvedLabels: map[string]struct{}{ + "falselabel": struct{}{}, + }, + MayReturn: true, + MayFallThrough: true, + }, + }, + { + name: "jump forward beyond fragment", + build: func() { + p.AddJumpFalseLabel(Jmp|Jeq|K, 1337, 123, "falselabel") + }, + wantLocal: FragmentOutcomes{ + MayJumpToUnresolvedLabels: map[string]struct{}{ + "falselabel": struct{}{}, + }, + MayJumpToKnownOffsetBeyondFragment: true, + }, + wantOverall: FragmentOutcomes{ + MayJumpToKnownOffsetBeyondFragment: true, + MayJumpToUnresolvedLabels: map[string]struct{}{ + "falselabel": struct{}{}, + }, + MayReturn: true, + }, + }, + { + name: "resolve falselabel", + build: func() { + p.AddLabel("falselabel") + }, + wantLocal: FragmentOutcomes{ + MayFallThrough: true, + }, + wantOverall: FragmentOutcomes{ + MayJumpToKnownOffsetBeyondFragment: true, + MayReturn: true, + MayFallThrough: true, + }, + }, + } { + t.Run(test.name, func(t *testing.T) { + getLocalFragment := p.Record() + test.build() + localFragment := getLocalFragment() + if localOutcomes := localFragment.Outcomes(); !reflect.DeepEqual(fixup(localOutcomes), fixup(test.wantLocal)) { + t.Errorf("local fragment %v: got outcomes %v want %v", localFragment, localOutcomes, test.wantLocal) + } + overallFragment := getOverallFragment() + if overallOutcomes := overallFragment.Outcomes(); !reflect.DeepEqual(fixup(overallOutcomes), fixup(test.wantOverall)) { + t.Errorf("overall fragment %v: got outcomes %v want %v", overallFragment, overallOutcomes, test.wantOverall) + } + }) + } +}