diff --git a/README.md b/README.md index facadd8..89cafaf 100644 --- a/README.md +++ b/README.md @@ -8,6 +8,32 @@ go-lua is a port of the Lua 5.2 VM to pure Go. It is compatible with binary file The motivation is to enable simple scripting of Go applications. For example, it is used to describe flows in [Shopify's](http://www.shopify.com/) load generation tool, Genghis. +Security +-------- + +go-lua is a VM, not a sandbox. Two boundaries are worth knowing before you run +code you did not write: + +**Binary chunks are not verified.** `luac` output is checked for structural +consistency when it is loaded, but individual instructions are not validated +against the function's registers, constants or jump targets. A crafted binary +chunk can therefore make a Lua function misbehave within its own state. Load +binary chunks only from a trusted source; for untrusted input pass mode `"t"` +so only Lua source is accepted: + +```go +err := l.Load(reader, name, "t") +``` + +Note that `LoadString`, `LoadBuffer`, `DoString` and `DoFile` pass mode `""` by +default, which accepts both text and binary. + +**The `debug` library is not safe to expose.** `OpenLibraries` opens `debug` +along with everything else. It hands scripts access to the registry, upvalues +and hook machinery. If you are running untrusted code, use `Require` to open +only the libraries you need. + + Usage ----- diff --git a/auxiliary.go b/auxiliary.go index 78dc064..3afd334 100644 --- a/auxiliary.go +++ b/auxiliary.go @@ -522,8 +522,16 @@ func LoadFile(l *State, fileName, mode string) error { return err } +// LoadString loads the given string as a chunk, without running it. +// +// The chunk may be either Lua source or a precompiled binary chunk. Binary +// chunks are not verified instruction by instruction, so load them only from +// a trusted source. For untrusted input, use LoadBuffer with mode "t". func LoadString(l *State, s string) error { return LoadBuffer(l, s, s, "") } +// LoadBuffer loads the given string as a chunk named name, without running +// it. The mode is as described for (*State).Load: "t" permits only text +// chunks, "b" only binary chunks, and "" (or any other value) permits both. func LoadBuffer(l *State, b, name, mode string) error { return l.Load(strings.NewReader(b), name, mode) } @@ -573,6 +581,9 @@ func FileResult(l *State, err error, filename string) int { } // DoFile loads and runs the given file. +// +// The file may contain either Lua source or a precompiled binary chunk. See +// LoadString for the trust requirement that binary chunks carry. func DoFile(l *State, fileName string) error { if err := LoadFile(l, fileName, ""); err != nil { return err @@ -581,6 +592,9 @@ func DoFile(l *State, fileName string) error { } // DoString loads and runs the given string. +// +// The string may contain either Lua source or a precompiled binary chunk. See +// LoadString for the trust requirement that binary chunks carry. func DoString(l *State, s string) error { if err := LoadString(l, s); err != nil { return err diff --git a/debug.go b/debug.go index 3041bc5..c814658 100644 --- a/debug.go +++ b/debug.go @@ -481,9 +481,7 @@ var debugLibrary = []RegistryFunction{ l.PushString("external hook") } else { hookTable(l) - l1.PushThread() - // XMove(l1, l, 1) - panic("XMove not implemented yet") + l.apiPush(l1) l.RawGet(-2) l.Remove(-2) } @@ -542,9 +540,7 @@ var debugLibrary = []RegistryFunction{ l.PushValue(-1) l.SetMetaTable(-2) } - l1.PushThread() - // XMove(l1, l, 1) - panic("XMove not yet implemented") + l.apiPush(l1) l.PushValue(i + 1) l.RawSet(-3) SetDebugHook(l1, hook, mask, count) diff --git a/lua.go b/lua.go index 68514e8..1d1740f 100644 --- a/lua.go +++ b/lua.go @@ -421,6 +421,13 @@ func (l *State) ProtectedCallWithContinuation(argCount, resultCount, errorFuncti // pushes the compiled chunk as a Lua function on top of the stack. // Otherwise, it pushes an error message. // +// The mode controls which kinds of chunk are accepted: "t" permits only text +// chunks, "b" only binary chunks, and "" (or any other value) permits both. +// Binary chunks are checked for structural consistency but are not verified +// instruction by instruction, so a crafted chunk can still make a Lua +// function misbehave within its own state. Load binary chunks only from a +// trusted source, and pass "t" for input you do not control. +// // http://www.lua.org/manual/5.2/manual.html#lua_load func (l *State) Load(r io.Reader, chunkName string, mode string) error { if chunkName == "" { @@ -547,7 +554,8 @@ func (l *State) AbsIndex(index int) int { // SetTop accepts any index, or 0, and sets the stack top to index. If the // new top is larger than the old one, then the new elements are filled with -// nil. If index is 0, then all stack elements are removed. +// nil. If index is 0, then all stack elements are removed. Removed elements +// are cleared, so they are neither visible to Lua code nor kept alive. // // If index is negative, the stack will be decremented by that much. If // the decrement is larger than the stack, SetTop will panic(). @@ -567,7 +575,11 @@ func (l *State) SetTop(index int) { if apiCheck && -(index+1) > l.top-(f+1) { panic("invalid new top") } + old := l.top l.top += index + 1 // 'subtract' index (index is negative) + for i := l.top; i < old; i++ { + l.stack[i] = nil + } } } diff --git a/parser.go b/parser.go index f833569..dc3577f 100644 --- a/parser.go +++ b/parser.go @@ -681,7 +681,12 @@ func protectedParser(l *State, r io.Reader, name, chunkMode string) error { } else if c == Signature[0] { l.checkMode(chunkMode, "binary") b.UnreadByte() - closure, _ = l.undump(b, name) // TODO handle err + undumped, undumpErr := l.undump(b, name) + if undumpErr != nil { + l.push(undumpErr.Error()) + l.throw(SyntaxError) + } + closure = undumped } else { l.checkMode(chunkMode, "text") b.UnreadByte() diff --git a/security_test.go b/security_test.go new file mode 100644 index 0000000..a1741f3 --- /dev/null +++ b/security_test.go @@ -0,0 +1,244 @@ +package lua + +import ( + "bytes" + "encoding/binary" + "runtime" + "strings" + "testing" +) + +const functionPrefixSize = 11 + +func chunkPrefix(t *testing.T, parameterCount, maxStackSize byte) *bytes.Buffer { + t.Helper() + b := new(bytes.Buffer) + if err := binary.Write(b, endianness(), header); err != nil { + t.Fatal(err) + } + if err := binary.Write(b, endianness(), []int32{0, 0}); err != nil { + t.Fatal(err) + } + if err := binary.Write(b, endianness(), []byte{parameterCount, 1, maxStackSize}); err != nil { + t.Fatal(err) + } + return b +} + +func writeInts(t *testing.T, b *bytes.Buffer, values ...int32) { + t.Helper() + if err := binary.Write(b, endianness(), values); err != nil { + t.Fatal(err) + } +} + +func undumpChunk(t *testing.T, b *bytes.Buffer) error { + t.Helper() + _, err := NewState().undump(bytes.NewReader(b.Bytes()), "crafted") + return err +} + +func dumpedChunk(t *testing.T, source string) []byte { + t.Helper() + l := NewState() + if err := LoadString(l, source); err != nil { + t.Fatal(err) + } + var out bytes.Buffer + if err := l.Dump(&out); err != nil { + t.Fatal(err) + } + return out.Bytes() +} + +func TestUndumpBoundsHugeCodeCount(t *testing.T) { + b := chunkPrefix(t, 0, 2) + writeInts(t, b, 0x7fffffff) + + var before, after runtime.MemStats + runtime.GC() + runtime.ReadMemStats(&before) + err := undumpChunk(t, b) + runtime.ReadMemStats(&after) + + if err == nil { + t.Fatal("expected an error for a truncated chunk claiming 0x7fffffff instructions") + } + if allocated := after.TotalAlloc - before.TotalAlloc; allocated > 8<<20 { + t.Errorf("undump allocated %d bytes for a %d byte chunk; want under 8 MiB", allocated, b.Len()) + } +} + +func TestUndumpBoundsHugeStringLength(t *testing.T) { + b := chunkPrefix(t, 0, 2) + writeInts(t, b, 0, 1) + b.WriteByte(byte(TypeString)) + if err := binary.Write(b, endianness(), uint64(1<<40)); err != nil { + t.Fatal(err) + } + + var before, after runtime.MemStats + runtime.GC() + runtime.ReadMemStats(&before) + err := undumpChunk(t, b) + runtime.ReadMemStats(&after) + + if err == nil { + t.Fatal("expected an error for a truncated chunk claiming a 1 TiB string constant") + } + if allocated := after.TotalAlloc - before.TotalAlloc; allocated > 8<<20 { + t.Errorf("undump allocated %d bytes for a %d byte chunk; want under 8 MiB", allocated, b.Len()) + } +} + +func TestUndumpRejectsNegativeCounts(t *testing.T) { + for _, test := range []struct { + name string + counts []int32 + }{ + {"code", []int32{-1}}, + {"constants", []int32{0, -1}}, + {"prototypes", []int32{0, 0, -1}}, + {"upvalues", []int32{0, 0, 0, -1}}, + } { + t.Run(test.name, func(t *testing.T) { + b := chunkPrefix(t, 0, 2) + writeInts(t, b, test.counts...) + if err := undumpChunk(t, b); err != errCorrupted { + t.Errorf("expected errCorrupted for a negative %s count, got %v", test.name, err) + } + }) + } +} + +func TestUndumpRejectsMoreUpValueNamesThanUpValues(t *testing.T) { + b := chunkPrefix(t, 0, 2) + writeInts(t, b, 0, 0, 0, 0) + if err := binary.Write(b, endianness(), uint64(0)); err != nil { + t.Fatal(err) + } + writeInts(t, b, 0, 0, 1) + if err := undumpChunk(t, b); err != errCorrupted { + t.Errorf("expected errCorrupted for 1 upvalue name against 0 upvalues, got %v", err) + } +} + +func TestUndumpRejectsDeeplyNestedPrototypes(t *testing.T) { + b := chunkPrefix(t, 0, 2) + for i := 0; i < maxUndumpNesting+1; i++ { + writeInts(t, b, 0, 0, 1) + writeInts(t, b, 0, 0) + if err := binary.Write(b, endianness(), []byte{0, 1, 2}); err != nil { + t.Fatal(err) + } + } + if err := undumpChunk(t, b); err != errCorrupted { + t.Errorf("expected errCorrupted for %d nested prototypes, got %v", maxUndumpNesting+1, err) + } +} + +func TestLoadReportsUndumpErrorAsSyntaxError(t *testing.T) { + b := chunkPrefix(t, 0, 2) + writeInts(t, b, 0x7fffffff) + if err := LoadBuffer(NewState(), b.String(), "crafted", "b"); err != SyntaxError { + t.Errorf("expected SyntaxError for a truncated binary chunk, got %v", err) + } +} + +func TestCraftedReturnDoesNotExposePoppedHostValues(t *testing.T) { + const secret = "SECRET" + chunk := dumpedChunk(t, "return 1") + at := binary.Size(header) + functionPrefixSize + 4 + endianness().PutUint32(chunk[at:at+4], uint32(createABC(opReturn, 0, 2, 0))) + + l := NewState() + l.PushString("filler") + l.PushString(secret) + l.Pop(2) + + if err := l.Load(bytes.NewReader(chunk), "crafted", ""); err != nil { + return + } + if err := l.ProtectedCall(0, 1, 0); err != nil { + return + } + if s, ok := l.ToString(-1); ok && s == secret { + t.Fatalf("crafted chunk returned the popped host value %q", s) + } + if !l.IsNil(-1) { + t.Errorf("expected an unwritten register to read as nil, got %v", l.TypeOf(-1)) + } +} + +func TestDebugGetHookWithoutExternalHook(t *testing.T) { + l := NewState() + Require(l, "debug", DebugOpen, true) + l.Pop(1) + if err := DoString(l, "return debug.gethook()"); err != nil { + t.Fatalf("debug.gethook() on a state with no hook installed: %v", err) + } +} + +func TestDebugSetHookThenGetHook(t *testing.T) { + l := NewState() + OpenLibraries(l) + if err := DoString(l, ` + local calls = 0 + local h = function() calls = calls + 1 end + debug.sethook(h, "l") + local got, mask = debug.gethook() + local x = 1 + debug.sethook() + assert(got == h, "gethook did not return the installed hook") + assert(mask == "l", "gethook returned mask " .. tostring(mask)) + assert(calls > 0, "the line hook never fired") + `); err != nil { + t.Fatalf("debug.sethook/gethook round trip: %v", err) + } +} + +func TestDebugSetHookArgumentFormsDoNotEscapeProtectedCall(t *testing.T) { + for _, source := range []string{ + "debug.sethook()", + "debug.sethook(nil)", + `debug.sethook(function() end, "")`, + `debug.sethook(function() end, "", 1)`, + `debug.sethook(function() end, "c", 0)`, + `debug.sethook(function() end, "lcr", 5)`, + } { + l := NewState() + Require(l, "debug", DebugOpen, true) + l.Pop(1) + if err := DoString(l, source); err != nil { + t.Errorf("%s returned %v", source, err) + } + } +} + +func TestDebugSetHookCountHookFires(t *testing.T) { + l := NewState() + OpenLibraries(l) + if err := DoString(l, ` + local calls = 0 + debug.sethook(function() calls = calls + 1 end, "", 1) + local x = 0 + for i = 1, 50 do x = x + i end + debug.sethook() + assert(x == 1275, "loop computed " .. tostring(x)) + assert(calls > 0, "the count hook never fired") + `); err != nil { + t.Fatalf("count hook: %v", err) + } +} + +func TestProtectedCallContainsNonErrorPanic(t *testing.T) { + l := NewState() + l.PushGoFunction(func(*State) int { panic("boom") }) + err := l.ProtectedCall(0, 0, 0) + if err == nil { + t.Fatal("expected an error from a Go function that panicked with a string") + } + if !strings.Contains(err.Error(), "boom") { + t.Errorf("expected the panic payload in the error, got %q", err.Error()) + } +} diff --git a/stack.go b/stack.go index 2570a54..9ba15b8 100644 --- a/stack.go +++ b/stack.go @@ -1,6 +1,9 @@ package lua -import "log" +import ( + "fmt" + "log" +) func (l *State) push(v value) { l.stack[l.top] = v @@ -180,6 +183,11 @@ func (l *State) pushLuaFrame(function, base, resultCount int, p *prototype) *cal ci.resultCount = resultCount ci.callStatus = callStatusLua ci.frame = l.stack[base:ci.top] + first := base + p.parameterCount + if l.top > first { + first = l.top + } + clear(l.stack[first:ci.top]) l.callInfo = ci l.top = ci.top return ci @@ -358,7 +366,7 @@ func (l *State) postCall(firstResult int) bool { result++ } l.top = result - if l.hookMask&(MaskReturn|MaskLine) != 0 { + if l.hookMask&(MaskReturn|MaskLine) != 0 && l.callInfo.isLua() { l.oldPC = l.callInfo.savedPC // oldPC for caller function } return wanted != MultipleReturns @@ -406,7 +414,10 @@ func (l *State) protect(f func()) (err error) { nestedGoCallCount, protectFunction := l.nestedGoCallCount, l.protectFunction l.protectFunction = func() { if e := recover(); e != nil { - err = e.(error) + var ok bool + if err, ok = e.(error); !ok { + err = fmt.Errorf("%v", e) + } l.nestedGoCallCount, l.protectFunction = nestedGoCallCount, protectFunction } } diff --git a/undump.go b/undump.go index b08bc37..5e7a123 100644 --- a/undump.go +++ b/undump.go @@ -12,8 +12,14 @@ import ( type loadState struct { in io.Reader order binary.ByteOrder + depth int } +const ( + undumpBatchSize = 4096 + maxUndumpNesting = 200 +) + var header struct { Signature [4]byte Version, Format, Endianness, IntSize byte @@ -59,6 +65,20 @@ func (state *loadState) readBool() (bool, error) { return b != 0, err } +func (state *loadState) readCount() (int, error) { + n, err := state.readInt() + if err != nil { + return 0, err + } else if n < 0 { + return 0, errCorrupted + } + return int(n), nil +} + +func initialCapacity(n int) int { + return min(n, undumpBatchSize) +} + func (state *loadState) readString() (s string, err error) { // Feel my pain maxUint := ^uint(0) @@ -77,72 +97,98 @@ func (state *loadState) readString() (s string, err error) { if err != nil || size == 0 { return } - ba := make([]byte, size) - if err = state.read(ba); err == nil { - s = string(ba[:len(ba)-1]) + capacity := uintptr(undumpBatchSize) + if size < capacity { + capacity = size } + ba := make([]byte, 0, capacity) + for uintptr(len(ba)) < size { + start := len(ba) + batch := size - uintptr(start) + if batch > undumpBatchSize { + batch = undumpBatchSize + } + ba = append(ba, make([]byte, batch)...) + if err = state.read(ba[start:]); err != nil { + return "", err + } + } + s = string(ba[:len(ba)-1]) return } func (state *loadState) readCode() (code []instruction, err error) { - n, err := state.readInt() + n, err := state.readCount() if err != nil || n == 0 { return } - code = make([]instruction, n) - err = state.read(code) + code = make([]instruction, 0, initialCapacity(n)) + for len(code) < n { + start := len(code) + code = append(code, make([]instruction, min(undumpBatchSize, n-start))...) + if err = state.read(code[start:]); err != nil { + return nil, err + } + } return } func (state *loadState) readUpValues() (u []upValueDesc, err error) { - n, err := state.readInt() + n, err := state.readCount() if err != nil || n == 0 { return } - v := make([]struct{ IsLocal, Index byte }, n) - err = state.read(v) - if err != nil { - return - } - u = make([]upValueDesc, n) - for i := range v { - u[i].isLocal, u[i].index = v[i].IsLocal != 0, int(v[i].Index) + u = make([]upValueDesc, 0, initialCapacity(n)) + for i := 0; i < n; i++ { + var v struct{ IsLocal, Index byte } + if err = state.read(&v); err != nil { + return nil, err + } + u = append(u, upValueDesc{isLocal: v.IsLocal != 0, index: int(v.Index)}) } return } func (state *loadState) readLocalVariables() (localVariables []localVariable, err error) { - var n int32 - if n, err = state.readInt(); err != nil || n == 0 { + var n int + if n, err = state.readCount(); err != nil || n == 0 { return } - localVariables = make([]localVariable, n) - for i := range localVariables { - if localVariables[i].name, err = state.readString(); err != nil { - return + localVariables = make([]localVariable, 0, initialCapacity(n)) + for i := 0; i < n; i++ { + var v localVariable + if v.name, err = state.readString(); err != nil { + return nil, err } - if localVariables[i].startPC, err = state.readPC(); err != nil { - return + if v.startPC, err = state.readPC(); err != nil { + return nil, err } - if localVariables[i].endPC, err = state.readPC(); err != nil { - return + if v.endPC, err = state.readPC(); err != nil { + return nil, err } + localVariables = append(localVariables, v) } return } func (state *loadState) readLineInfo() (lineInfo []int32, err error) { - var n int32 - if n, err = state.readInt(); err != nil || n == 0 { + var n int + if n, err = state.readCount(); err != nil || n == 0 { return } - lineInfo = make([]int32, n) - err = state.read(lineInfo) + lineInfo = make([]int32, 0, initialCapacity(n)) + for len(lineInfo) < n { + start := len(lineInfo) + lineInfo = append(lineInfo, make([]int32, min(undumpBatchSize, n-start))...) + if err = state.read(lineInfo[start:]); err != nil { + return nil, err + } + } return } func (state *loadState) readDebug(p *prototype) (source string, lineInfo []int32, localVariables []localVariable, names []string, err error) { - var n int32 + var n int if source, err = state.readString(); err != nil { return } @@ -152,63 +198,76 @@ func (state *loadState) readDebug(p *prototype) (source string, lineInfo []int32 if localVariables, err = state.readLocalVariables(); err != nil { return } - if n, err = state.readInt(); err != nil { + if n, err = state.readCount(); err != nil { + return + } else if n > len(p.upValues) { + err = errCorrupted return } - names = make([]string, n) - for i := range names { - if names[i], err = state.readString(); err != nil { + names = make([]string, 0, initialCapacity(n)) + for i := 0; i < n; i++ { + var name string + if name, err = state.readString(); err != nil { return } + names = append(names, name) } return } func (state *loadState) readConstants() (constants []value, prototypes []prototype, err error) { - var n int32 - if n, err = state.readInt(); err != nil || n == 0 { + var n int + if n, err = state.readCount(); err != nil || n == 0 { return } - constants = make([]value, n) - for i := range constants { + constants = make([]value, 0, initialCapacity(n)) + for i := 0; i < n; i++ { + var c value var t byte switch t, err = state.readByte(); { case err != nil: return case t == byte(TypeNil): - constants[i] = nil + c = nil case t == byte(TypeBoolean): - constants[i], err = state.readBool() + c, err = state.readBool() case t == byte(TypeNumber): - constants[i], err = state.readNumber() + c, err = state.readNumber() case t == byte(TypeString): - constants[i], err = state.readString() + c, err = state.readString() default: err = errUnknownConstantType } if err != nil { return } + constants = append(constants, c) } return } func (state *loadState) readPrototypes() (prototypes []prototype, err error) { - var n int32 - if n, err = state.readInt(); err != nil || n == 0 { + var n int + if n, err = state.readCount(); err != nil || n == 0 { return } - prototypes = make([]prototype, n) - for i := range prototypes { - if prototypes[i], err = state.readFunction(); err != nil { - return + prototypes = make([]prototype, 0, initialCapacity(n)) + for i := 0; i < n; i++ { + var p prototype + if p, err = state.readFunction(); err != nil { + return nil, err } + prototypes = append(prototypes, p) } return } func (state *loadState) readFunction() (p prototype, err error) { + if state.depth++; state.depth > maxUndumpNesting { + return p, errCorrupted + } + defer func() { state.depth-- }() var n int32 if n, err = state.readInt(); err != nil { return @@ -231,9 +290,15 @@ func (state *loadState) readFunction() (p prototype, err error) { return } p.maxStackSize = int(b) + if p.maxStackSize < p.parameterCount { + return p, errCorrupted + } if p.code, err = state.readCode(); err != nil { return } + if len(p.code) > 0 && p.maxStackSize == 0 { + return p, errCorrupted + } if p.constants, p.prototypes, err = state.readConstants(); err != nil { return } @@ -310,7 +375,7 @@ func (l *State) undump(in io.Reader, name string) (c *luaClosure, err error) { name = "binary string" } // TODO assign name to p.source? - s := &loadState{in, endianness()} + s := &loadState{in: in, order: endianness()} var p prototype if err = s.checkHeader(); err != nil { return diff --git a/vm.go b/vm.go index cf4d57d..dd65111 100644 --- a/vm.go +++ b/vm.go @@ -233,9 +233,15 @@ func (l *State) traceExecution() { if mask&MaskLine != 0 { p := l.prototype(callInfo) npc := callInfo.savedPC - 1 - newline := p.lineInfo[npc] - if npc == 0 || callInfo.savedPC <= l.oldPC || newline != p.lineInfo[l.oldPC-1] { - l.hook(HookLine, int(newline)) + if npc < 0 { + npc = 0 + } + if int(npc) < len(p.lineInfo) { + newline := p.lineInfo[npc] + if npc == 0 || l.oldPC == 0 || callInfo.savedPC <= l.oldPC || + int(l.oldPC-1) >= len(p.lineInfo) || newline != p.lineInfo[l.oldPC-1] { + l.hook(HookLine, int(newline)) + } } } l.oldPC = callInfo.savedPC