mirror of
https://github.com/XTLS/Xray-core.git
synced 2026-10-11 00:38:24 +03:00
Xray-core: Add Lua script for dns and routing (#6823)
https://github.com/XTLS/Xray-core/pull/6823#issuecomment-5843754759 https://github.com/XTLS/Xray-core/pull/6823#issuecomment-5861456450 https://github.com/XTLS/Xray-core/pull/6823#issuecomment-6093596069
This commit is contained in:
@@ -0,0 +1,3 @@
|
||||
// Package lua provides shared GopherLua programs, state management, and value
|
||||
// conversion and validation helpers for Xray scripts.
|
||||
package lua
|
||||
@@ -0,0 +1,65 @@
|
||||
package lua
|
||||
|
||||
import (
|
||||
glua "github.com/yuin/gopher-lua"
|
||||
luar "layeh.com/gopher-luar"
|
||||
)
|
||||
|
||||
// NewSlicePusher captures luar's slice metatable during state initialization.
|
||||
// The returned function wraps slices without reflection or metatable lookup,
|
||||
// and pushes nil for nil slices. Use it with this state or its coroutines.
|
||||
func NewSlicePusher[T any](L *glua.LState) func(*glua.LState, []T) {
|
||||
metatable := luar.New(L, []T{}).(*glua.LUserData).Metatable
|
||||
return func(L *glua.LState, values []T) {
|
||||
if values == nil {
|
||||
L.Push(glua.LNil)
|
||||
return
|
||||
}
|
||||
userdata := L.NewUserData()
|
||||
userdata.Value = values
|
||||
userdata.Metatable = metatable
|
||||
L.Push(userdata)
|
||||
}
|
||||
}
|
||||
|
||||
// DirectMethod handles a Lua call without luar's reflected method invocation.
|
||||
// It returns the result count and whether it handled the arguments. On false,
|
||||
// it must leave the stack unchanged for the original luar wrapper.
|
||||
type DirectMethod func(L *glua.LState) (nresults int, handled bool)
|
||||
|
||||
// PushWithDirectMethods pushes a luar userdata with typed Go method bindings.
|
||||
// Handled calls bypass luar's argument conversion and reflect.Call; method lookup
|
||||
// uses the methods table directly instead of luar's reflected __index handler.
|
||||
// value must expose methods only. Bindings and their closures are installed once
|
||||
// per Go type per LState, outside the method-call hot path.
|
||||
func PushWithDirectMethods(L *glua.LState, value any, directMethods map[string]DirectMethod) {
|
||||
userdata := luar.New(L, value).(*glua.LUserData)
|
||||
metatable := userdata.Metatable.(*glua.LTable)
|
||||
methods := metatable.RawGetString("methods").(*glua.LTable)
|
||||
if metatable.RawGetString("__index") != methods {
|
||||
for name, direct := range directMethods {
|
||||
original := methods.RawGetString(name)
|
||||
fn := L.NewFunction(func(L *glua.LState) int {
|
||||
if nresults, handled := direct(L); handled {
|
||||
return nresults
|
||||
}
|
||||
return callLuarMethod(L, original)
|
||||
})
|
||||
// Keep luar's method aliases on the same direct binding.
|
||||
for key, method := methods.Next(glua.LNil); key != glua.LNil; key, method = methods.Next(key) {
|
||||
if method == original {
|
||||
methods.RawSet(key, fn)
|
||||
}
|
||||
}
|
||||
}
|
||||
metatable.RawSetString("__index", methods)
|
||||
}
|
||||
L.Push(userdata)
|
||||
}
|
||||
|
||||
func callLuarMethod(L *glua.LState, method glua.LValue) int {
|
||||
nargs := L.GetTop()
|
||||
L.Insert(method, 1)
|
||||
L.Call(nargs, glua.MultRet)
|
||||
return L.GetTop()
|
||||
}
|
||||
@@ -0,0 +1,84 @@
|
||||
package lua
|
||||
|
||||
import (
|
||||
"net"
|
||||
"testing"
|
||||
|
||||
glua "github.com/yuin/gopher-lua"
|
||||
luar "layeh.com/gopher-luar"
|
||||
)
|
||||
|
||||
func TestSlicePusher(t *testing.T) {
|
||||
L := glua.NewState()
|
||||
defer L.Close()
|
||||
push := NewSlicePusher[int](L)
|
||||
values := []int{3, 5}
|
||||
L.SetGlobal("getValues", L.NewFunction(func(L *glua.LState) int {
|
||||
push(L, values)
|
||||
return 1
|
||||
}))
|
||||
if err := L.DoString(`
|
||||
local values = getValues()
|
||||
assert(#values == 2 and values[1] == 3 and values[2] == 5)
|
||||
values[2] = 7
|
||||
local co = coroutine.create(function()
|
||||
local values = getValues()
|
||||
assert(#values == 2 and values[1] == 3 and values[2] == 7)
|
||||
return true
|
||||
end)
|
||||
local ok, result = coroutine.resume(co)
|
||||
assert(ok and result == true)
|
||||
`); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if values[1] != 7 {
|
||||
t.Fatal("slice storage was copied")
|
||||
}
|
||||
push(L, nil)
|
||||
if L.Get(-1) != glua.LNil {
|
||||
t.Fatal("nil slice must push Lua nil")
|
||||
}
|
||||
L.Pop(1)
|
||||
push(L, []int{})
|
||||
L.SetGlobal("empty", L.Get(-1))
|
||||
L.Pop(1)
|
||||
if err := L.DoString(`assert(type(empty) == "userdata" and #empty == 0)`); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSlicePusherMetatablePerState(t *testing.T) {
|
||||
first := glua.NewState()
|
||||
defer first.Close()
|
||||
second := glua.NewState()
|
||||
defer second.Close()
|
||||
NewSlicePusher[int](first)(first, []int{1})
|
||||
NewSlicePusher[int](second)(second, []int{1})
|
||||
if first.Get(-1).(*glua.LUserData).Metatable == second.Get(-1).(*glua.LUserData).Metatable {
|
||||
t.Fatal("independent states share a slice metatable")
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkSlicePusher(b *testing.B) {
|
||||
L := glua.NewState()
|
||||
defer L.Close()
|
||||
ips := []net.IP{net.ParseIP("127.0.0.1")}
|
||||
pushIPs := NewSlicePusher[net.IP](L)
|
||||
for _, benchmark := range []struct {
|
||||
name string
|
||||
push func(*glua.LState, []net.IP)
|
||||
}{
|
||||
{"bare", func(L *glua.LState, ips []net.IP) { PushUserData(L, ips) }},
|
||||
{"luar", func(L *glua.LState, ips []net.IP) { L.Push(luar.New(L, ips)) }},
|
||||
{"cached", pushIPs},
|
||||
} {
|
||||
b.Run(benchmark.name, func(b *testing.B) {
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
benchmark.push(L, ips)
|
||||
L.Pop(1)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,151 @@
|
||||
package lua
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
glua "github.com/yuin/gopher-lua"
|
||||
)
|
||||
|
||||
const maxIdleStates = 16
|
||||
|
||||
// Pool lends each state to one caller at a time. It grows on contention and
|
||||
// keeps up to maxIdleStates idle states until Close. Acquire/Release callers
|
||||
// decide reusability; WithState uses its callback's error.
|
||||
type Pool struct {
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
timeout time.Duration
|
||||
|
||||
factory LStateFactory
|
||||
idle []*glua.LState
|
||||
top int
|
||||
|
||||
mu sync.Mutex
|
||||
active sync.WaitGroup
|
||||
closed bool
|
||||
}
|
||||
|
||||
// NewPool tests the factory by creating one state during initialization.
|
||||
func NewPool(ctx context.Context, timeout time.Duration, factory LStateFactory) (*Pool, error) {
|
||||
if timeout <= 0 {
|
||||
return nil, errors.New("Lua pool timeout must be positive")
|
||||
}
|
||||
|
||||
poolCtx, cancel := context.WithCancel(ctx)
|
||||
|
||||
state, err := factory(poolCtx)
|
||||
if err != nil {
|
||||
cancel()
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &Pool{ctx: poolCtx, cancel: cancel, timeout: timeout, factory: factory, idle: []*glua.LState{state}, top: state.GetTop()}, nil
|
||||
}
|
||||
|
||||
// Acquire returns an initialized exclusive state, growing the pool if necessary.
|
||||
// ctx is passed to the factory for state creation; nil uses the pool context.
|
||||
func (p *Pool) Acquire(ctx context.Context) (*glua.LState, error) {
|
||||
p.mu.Lock()
|
||||
if p.closed {
|
||||
p.mu.Unlock()
|
||||
return nil, errors.New("Lua pool is closed")
|
||||
}
|
||||
if err := p.ctx.Err(); err != nil {
|
||||
p.mu.Unlock()
|
||||
return nil, err
|
||||
}
|
||||
if ctx == nil {
|
||||
ctx = p.ctx
|
||||
} else if err := ctx.Err(); err != nil {
|
||||
p.mu.Unlock()
|
||||
return nil, err
|
||||
}
|
||||
|
||||
p.active.Add(1)
|
||||
|
||||
n := len(p.idle)
|
||||
if n != 0 {
|
||||
state := p.idle[n-1]
|
||||
p.idle[n-1] = nil
|
||||
p.idle = p.idle[:n-1]
|
||||
p.mu.Unlock()
|
||||
return state, nil
|
||||
}
|
||||
p.mu.Unlock()
|
||||
|
||||
// TODO: Limit the total number of states. When the limit is reached, wait
|
||||
// for a Release instead of creating another state; allow the wait to be
|
||||
// cancelled by the caller or by Close.
|
||||
state, err := p.factory(ctx)
|
||||
if err != nil {
|
||||
p.active.Done()
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return state, nil
|
||||
}
|
||||
|
||||
// WithState runs work on an exclusive state and releases it afterward.
|
||||
// Nil ctx and zero timeout use pool defaults. The timeout starts after acquisition.
|
||||
func (p *Pool) WithState(ctx context.Context, timeout time.Duration, work func(*glua.LState) error) error {
|
||||
state, err := p.Acquire(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if ctx == nil {
|
||||
ctx = p.ctx
|
||||
}
|
||||
if timeout == 0 {
|
||||
timeout = p.timeout
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(ctx, timeout)
|
||||
state.SetContext(ctx)
|
||||
reusable := false
|
||||
defer func() {
|
||||
cancel()
|
||||
p.Release(state, reusable)
|
||||
}()
|
||||
err = work(state)
|
||||
reusable = err == nil
|
||||
return err
|
||||
}
|
||||
|
||||
// Release resets a state for reuse or closes it.
|
||||
func (p *Pool) Release(state *glua.LState, reusable bool) {
|
||||
if reusable {
|
||||
state.RemoveContext()
|
||||
state.SetTop(p.top)
|
||||
p.mu.Lock()
|
||||
if !p.closed && p.ctx.Err() == nil && len(p.idle) < maxIdleStates {
|
||||
p.idle = append(p.idle, state)
|
||||
} else {
|
||||
reusable = false
|
||||
}
|
||||
p.mu.Unlock()
|
||||
}
|
||||
|
||||
if !reusable {
|
||||
state.Close()
|
||||
}
|
||||
|
||||
p.active.Done()
|
||||
}
|
||||
|
||||
// Close cancels the pool context, closes idle states, and waits for borrowed states.
|
||||
func (p *Pool) Close() {
|
||||
p.mu.Lock()
|
||||
if !p.closed {
|
||||
p.closed = true
|
||||
p.cancel()
|
||||
for _, state := range p.idle {
|
||||
state.Close()
|
||||
}
|
||||
p.idle = nil
|
||||
}
|
||||
p.mu.Unlock()
|
||||
|
||||
p.active.Wait()
|
||||
}
|
||||
@@ -0,0 +1,466 @@
|
||||
package lua
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
glua "github.com/yuin/gopher-lua"
|
||||
)
|
||||
|
||||
func newTestPool(t testing.TB, ctx context.Context, timeout time.Duration, factory LStateFactory) *Pool {
|
||||
t.Helper()
|
||||
pool, err := NewPool(ctx, timeout, factory)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(pool.Close)
|
||||
return pool
|
||||
}
|
||||
|
||||
func assertPoolCloseBlocked(t *testing.T, done <-chan struct{}) {
|
||||
t.Helper()
|
||||
select {
|
||||
case <-done:
|
||||
t.Fatal("Close returned while work was still active")
|
||||
case <-time.After(20 * time.Millisecond):
|
||||
}
|
||||
}
|
||||
|
||||
func TestPoolTimeoutValidation(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
timeout time.Duration
|
||||
wantErr bool
|
||||
}{
|
||||
{"zero", 0, true},
|
||||
{"negative", -time.Nanosecond, true},
|
||||
{"positive", time.Nanosecond, false},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
called := false
|
||||
pool, err := NewPool(context.Background(), tc.timeout, func(context.Context) (*glua.LState, error) {
|
||||
called = true
|
||||
return glua.NewState(), nil
|
||||
})
|
||||
if pool != nil {
|
||||
t.Cleanup(pool.Close)
|
||||
}
|
||||
if (err != nil) != tc.wantErr {
|
||||
t.Fatalf("NewPool error = %v, want error %t", err, tc.wantErr)
|
||||
}
|
||||
if tc.wantErr && (pool != nil || called) {
|
||||
t.Fatal("invalid timeout created a pool or called the factory")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPoolFactoryFailure(t *testing.T) {
|
||||
failure := errors.New("factory failed")
|
||||
_, err := NewPool(context.Background(), time.Second, func(context.Context) (*glua.LState, error) {
|
||||
return nil, failure
|
||||
})
|
||||
if !errors.Is(err, failure) {
|
||||
t.Fatalf("NewPool error = %v, want original factory error", err)
|
||||
}
|
||||
|
||||
calls := 0
|
||||
pool := newTestPool(t, context.Background(), time.Second, func(context.Context) (*glua.LState, error) {
|
||||
calls++
|
||||
if calls == 1 {
|
||||
return glua.NewState(), nil
|
||||
}
|
||||
return nil, failure
|
||||
})
|
||||
state, err := pool.Acquire(nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer pool.Release(state, true)
|
||||
err = pool.WithState(nil, 0, func(*glua.LState) error {
|
||||
t.Error("work ran after factory failure")
|
||||
return nil
|
||||
})
|
||||
if !errors.Is(err, failure) {
|
||||
t.Fatalf("WithState error = %v, want original factory error", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPoolReusesStatesAndLimitsIdle(t *testing.T) {
|
||||
created := 0
|
||||
pool := newTestPool(t, context.Background(), time.Second, func(context.Context) (*glua.LState, error) {
|
||||
created++
|
||||
return glua.NewState(), nil
|
||||
})
|
||||
var borrowed []*glua.LState
|
||||
defer func() {
|
||||
for _, state := range borrowed {
|
||||
pool.Release(state, false)
|
||||
}
|
||||
}()
|
||||
for range maxIdleStates + 3 {
|
||||
state, err := pool.Acquire(nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
borrowed = append(borrowed, state)
|
||||
state.SetContext(context.Background())
|
||||
}
|
||||
states := borrowed
|
||||
for _, state := range states {
|
||||
pool.Release(state, true)
|
||||
}
|
||||
borrowed = nil
|
||||
open := 0
|
||||
for _, state := range states {
|
||||
if !state.IsClosed() {
|
||||
if state.Context() != nil {
|
||||
t.Fatal("Release left a context on a reusable state")
|
||||
}
|
||||
open++
|
||||
}
|
||||
}
|
||||
if open != maxIdleStates {
|
||||
t.Fatalf("retained %d states, want %d", open, maxIdleStates)
|
||||
}
|
||||
if err := pool.WithState(nil, 0, func(*glua.LState) error { return nil }); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if created != len(states) {
|
||||
t.Fatalf("created %d states, want %d", created, len(states))
|
||||
}
|
||||
pool.Close()
|
||||
for _, state := range states {
|
||||
if !state.IsClosed() {
|
||||
t.Fatal("Close left an idle state open")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPoolWithStateOptions(t *testing.T) {
|
||||
key := struct{}{}
|
||||
parent := context.WithValue(context.Background(), key, "pool")
|
||||
caller := context.WithValue(context.Background(), key, "caller")
|
||||
pool := newTestPool(t, parent, time.Second, func(context.Context) (*glua.LState, error) {
|
||||
return glua.NewState(), nil
|
||||
})
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
ctx context.Context
|
||||
timeout time.Duration
|
||||
wantValue string
|
||||
wantTimeout time.Duration
|
||||
}{
|
||||
{"defaults", nil, 0, "pool", time.Second},
|
||||
{"context", caller, 0, "caller", time.Second},
|
||||
{"timeout", nil, 2 * time.Second, "pool", 2 * time.Second},
|
||||
{"both", caller, 2 * time.Second, "caller", 2 * time.Second},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
started := time.Now()
|
||||
err := pool.WithState(tc.ctx, tc.timeout, func(L *glua.LState) error {
|
||||
ctx := L.Context()
|
||||
if ctx.Value(key) != tc.wantValue {
|
||||
t.Errorf("context value = %v, want %q", ctx.Value(key), tc.wantValue)
|
||||
}
|
||||
deadline, ok := ctx.Deadline()
|
||||
if !ok || deadline.Before(started.Add(tc.wantTimeout)) || deadline.After(time.Now().Add(tc.wantTimeout)) {
|
||||
t.Errorf("deadline = %v, want timeout %v", deadline, tc.wantTimeout)
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPoolFactoryContext(t *testing.T) {
|
||||
caller, cancel := context.WithTimeout(context.Background(), time.Minute)
|
||||
defer cancel()
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
ctx context.Context
|
||||
}{
|
||||
{"default", nil},
|
||||
{"caller", caller},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
var contexts []context.Context
|
||||
pool := newTestPool(t, context.Background(), time.Second, func(ctx context.Context) (*glua.LState, error) {
|
||||
contexts = append(contexts, ctx)
|
||||
return glua.NewState(), nil
|
||||
})
|
||||
state, err := pool.Acquire(nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer pool.Release(state, true)
|
||||
if err := pool.WithState(tc.ctx, 2*time.Second, func(*glua.LState) error { return nil }); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
want := tc.ctx
|
||||
if want == nil {
|
||||
want = pool.ctx
|
||||
}
|
||||
if len(contexts) != 2 || contexts[0] != pool.ctx || contexts[1] != want {
|
||||
t.Fatal("factory did not receive the initialization and acquisition contexts unchanged")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPoolWithStateLifecycle(t *testing.T) {
|
||||
failure := errors.New("work failed")
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
work func(*glua.LState, context.CancelFunc) error
|
||||
reusable bool
|
||||
wantPanic bool
|
||||
wantErr error
|
||||
}{
|
||||
{"success", func(*glua.LState, context.CancelFunc) error { return nil }, true, false, nil},
|
||||
{"canceled success", func(_ *glua.LState, cancel context.CancelFunc) error {
|
||||
cancel()
|
||||
return nil
|
||||
}, true, false, nil},
|
||||
{"error", func(*glua.LState, context.CancelFunc) error { return failure }, false, false, failure},
|
||||
{"timeout", func(L *glua.LState, _ context.CancelFunc) error { return L.DoString("while true do end") }, false, false, nil},
|
||||
{"panic", func(*glua.LState, context.CancelFunc) error { panic(failure) }, false, true, nil},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
pool := newTestPool(t, context.Background(), 10*time.Millisecond, func(context.Context) (*glua.LState, error) {
|
||||
state := glua.NewState()
|
||||
state.Push(glua.LTrue)
|
||||
return state, nil
|
||||
})
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
var state *glua.LState
|
||||
var workCtx context.Context
|
||||
var recovered any
|
||||
err := func() (err error) {
|
||||
defer func() { recovered = recover() }()
|
||||
return pool.WithState(ctx, 0, func(L *glua.LState) error {
|
||||
state, workCtx = L, L.Context()
|
||||
L.Push(glua.LFalse)
|
||||
return tc.work(L, cancel)
|
||||
})
|
||||
}()
|
||||
if tc.wantPanic {
|
||||
if recovered != failure {
|
||||
t.Fatalf("panic = %v, want original panic", recovered)
|
||||
}
|
||||
} else {
|
||||
if recovered != nil || (err == nil) != tc.reusable {
|
||||
t.Fatalf("WithState error = %v, panic = %v", err, recovered)
|
||||
}
|
||||
if tc.wantErr != nil && !errors.Is(err, tc.wantErr) {
|
||||
t.Fatalf("WithState error = %v, want %v", err, tc.wantErr)
|
||||
}
|
||||
}
|
||||
if workCtx.Err() == nil {
|
||||
t.Fatal("WithState did not cancel the execution context")
|
||||
}
|
||||
if closed := state.IsClosed(); closed == tc.reusable {
|
||||
t.Fatalf("state closed = %t, want %t", closed, !tc.reusable)
|
||||
}
|
||||
if tc.reusable && (state.Context() != nil || state.GetTop() != 1 || state.Get(1) != glua.LTrue) {
|
||||
t.Fatal("WithState did not reset the state for reuse")
|
||||
}
|
||||
if err := pool.WithState(nil, 0, func(L *glua.LState) error {
|
||||
if (L == state) != tc.reusable {
|
||||
t.Error("unexpected state reuse")
|
||||
}
|
||||
return nil
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPoolClose(t *testing.T) {
|
||||
pool := newTestPool(t, context.Background(), time.Second, func(context.Context) (*glua.LState, error) {
|
||||
return glua.NewState(), nil
|
||||
})
|
||||
finishCtx, finish := context.WithCancel(context.Background())
|
||||
t.Cleanup(finish)
|
||||
started, done := make(chan *glua.LState, 1), make(chan error, 1)
|
||||
var workCtx context.Context
|
||||
go func() {
|
||||
done <- pool.WithState(nil, 0, func(L *glua.LState) error {
|
||||
workCtx = L.Context()
|
||||
started <- L
|
||||
<-finishCtx.Done()
|
||||
return nil
|
||||
})
|
||||
}()
|
||||
var state *glua.LState
|
||||
select {
|
||||
case state = <-started:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("WithState did not start")
|
||||
}
|
||||
closed := make(chan struct{})
|
||||
go func() {
|
||||
pool.Close()
|
||||
close(closed)
|
||||
}()
|
||||
select {
|
||||
case <-workCtx.Done():
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("Close did not cancel work using the pool context")
|
||||
}
|
||||
if !errors.Is(workCtx.Err(), context.Canceled) {
|
||||
t.Fatalf("work context error = %v, want context.Canceled", workCtx.Err())
|
||||
}
|
||||
assertPoolCloseBlocked(t, closed)
|
||||
finish()
|
||||
select {
|
||||
case err := <-done:
|
||||
if err != nil {
|
||||
t.Fatalf("successful work returned an error: %v", err)
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("WithState did not finish")
|
||||
}
|
||||
select {
|
||||
case <-closed:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("Close did not finish after WithState")
|
||||
}
|
||||
if !state.IsClosed() {
|
||||
t.Fatal("Release returned a state to a closed pool")
|
||||
}
|
||||
if state, err := pool.Acquire(nil); state != nil || err == nil || errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("Acquire after Close = %v, %v; want closed pool error", state, err)
|
||||
}
|
||||
pool.Close()
|
||||
}
|
||||
|
||||
func TestPoolCloseWaitsForFactory(t *testing.T) {
|
||||
finishCtx, finish := context.WithCancel(context.Background())
|
||||
started, canceled := make(chan struct{}), make(chan struct{})
|
||||
first := true
|
||||
pool := newTestPool(t, context.Background(), time.Second, func(ctx context.Context) (*glua.LState, error) {
|
||||
if first {
|
||||
first = false
|
||||
return glua.NewState(), nil
|
||||
}
|
||||
close(started)
|
||||
<-ctx.Done()
|
||||
close(canceled)
|
||||
<-finishCtx.Done()
|
||||
return nil, ctx.Err()
|
||||
})
|
||||
t.Cleanup(finish)
|
||||
state, err := pool.Acquire(nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
pool.Release(state, false)
|
||||
acquireDone := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := pool.Acquire(nil)
|
||||
acquireDone <- err
|
||||
}()
|
||||
select {
|
||||
case <-started:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("state creation did not start")
|
||||
}
|
||||
closed := make(chan struct{})
|
||||
go func() {
|
||||
pool.Close()
|
||||
close(closed)
|
||||
}()
|
||||
select {
|
||||
case <-canceled:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("Close did not cancel state creation")
|
||||
}
|
||||
assertPoolCloseBlocked(t, closed)
|
||||
finish()
|
||||
select {
|
||||
case err := <-acquireDone:
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("Acquire error = %v, want context.Canceled", err)
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("state creation did not finish")
|
||||
}
|
||||
select {
|
||||
case <-closed:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("Close did not finish after state creation")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPoolCloseWaitsForCallerContext(t *testing.T) {
|
||||
pool := newTestPool(t, context.Background(), time.Minute, func(context.Context) (*glua.LState, error) {
|
||||
return glua.NewState(), nil
|
||||
})
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
t.Cleanup(cancel)
|
||||
started, done := make(chan context.Context, 1), make(chan error, 1)
|
||||
go func() {
|
||||
done <- pool.WithState(ctx, 0, func(L *glua.LState) error {
|
||||
started <- L.Context()
|
||||
<-L.Context().Done()
|
||||
return L.Context().Err()
|
||||
})
|
||||
}()
|
||||
var workCtx context.Context
|
||||
select {
|
||||
case workCtx = <-started:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("WithState did not start")
|
||||
}
|
||||
closed := make(chan struct{})
|
||||
go func() {
|
||||
pool.Close()
|
||||
close(closed)
|
||||
}()
|
||||
select {
|
||||
case <-pool.ctx.Done():
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("Close did not cancel the pool context")
|
||||
}
|
||||
assertPoolCloseBlocked(t, closed)
|
||||
if workCtx.Err() != nil || ctx.Err() != nil {
|
||||
t.Fatal("Close canceled the caller's execution context")
|
||||
}
|
||||
cancel()
|
||||
select {
|
||||
case err := <-done:
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("WithState error = %v, want context.Canceled", err)
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("WithState did not stop after caller cancellation")
|
||||
}
|
||||
select {
|
||||
case <-closed:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("Close did not finish after WithState")
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkPoolAcquireRelease(b *testing.B) {
|
||||
pool := newTestPool(b, context.Background(), time.Second, func(context.Context) (*glua.LState, error) {
|
||||
return glua.NewState(), nil
|
||||
})
|
||||
b.ReportAllocs()
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
state, err := pool.Acquire(nil)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
pool.Release(state, true)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,77 @@
|
||||
package lua
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"os"
|
||||
"time"
|
||||
|
||||
glua "github.com/yuin/gopher-lua"
|
||||
"github.com/yuin/gopher-lua/parse"
|
||||
)
|
||||
|
||||
// Program holds immutable bytecode that can be run by independent LStates.
|
||||
type Program struct {
|
||||
proto *glua.FunctionProto
|
||||
}
|
||||
|
||||
// LStateFactory returns a fully initialized state or nil and an error.
|
||||
// Implementations must close partial states on failure; callers own successful states.
|
||||
type LStateFactory func(context.Context) (*glua.LState, error)
|
||||
|
||||
// CompileFile reads and compiles a Lua file once.
|
||||
func CompileFile(path string) (*Program, error) {
|
||||
f, err := os.Open(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer f.Close()
|
||||
chunk, err := parse.Parse(bufio.NewReader(f), path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
proto, err := glua.Compile(chunk, path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &Program{proto: proto}, nil
|
||||
}
|
||||
|
||||
// NewState creates a state, runs register, executes the program under ctx, and
|
||||
// runs validate. It removes the initialization context before returning a state
|
||||
// owned by the caller.
|
||||
func (p *Program) NewState(ctx context.Context, register func(*glua.LState), validate func(*glua.LState) error) (*glua.LState, error) {
|
||||
L := glua.NewState()
|
||||
valid := false
|
||||
defer func() {
|
||||
if !valid {
|
||||
L.Close()
|
||||
}
|
||||
}()
|
||||
L.SetContext(ctx)
|
||||
defer L.RemoveContext()
|
||||
if register != nil {
|
||||
register(L)
|
||||
}
|
||||
L.Push(L.NewFunctionFromProto(p.proto))
|
||||
// Execute the Lua script's top level.
|
||||
if err := L.PCall(0, 0, nil); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if validate != nil {
|
||||
if err := validate(L); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
valid = true
|
||||
return L, nil
|
||||
}
|
||||
|
||||
// NewStateFactory returns a factory that gives each state an initialization timeout.
|
||||
func (p *Program) NewStateFactory(initTimeout time.Duration, register func(*glua.LState), validate func(*glua.LState) error) LStateFactory {
|
||||
return func(ctx context.Context) (*glua.LState, error) {
|
||||
initCtx, cancel := context.WithTimeout(ctx, initTimeout)
|
||||
defer cancel()
|
||||
return p.NewState(initCtx, register, validate)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,76 @@
|
||||
package lua
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
glua "github.com/yuin/gopher-lua"
|
||||
)
|
||||
|
||||
func TestProgramStatesAreIndependent(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "state.lua")
|
||||
if err := os.WriteFile(path, []byte("value = (value or 0) + 1"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
program, err := CompileFile(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
first, err := program.NewState(context.Background(), nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer first.Close()
|
||||
first.SetGlobal("value", glua.LNumber(42))
|
||||
second, err := program.NewState(context.Background(), nil, nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer second.Close()
|
||||
if got := second.GetGlobal("value"); got != glua.LNumber(1) {
|
||||
t.Fatalf("second state value = %v, want 1", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProgramInitializationObservesCancellation(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "loop.lua")
|
||||
if err := os.WriteFile(path, []byte("while true do end"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
program, err := CompileFile(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
state, err := program.NewState(ctx, nil, nil)
|
||||
if err == nil || state != nil {
|
||||
if state != nil {
|
||||
state.Close()
|
||||
}
|
||||
t.Fatalf("NewState with canceled context = %v, %v; want nil state and error", state, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewStateClosesFailedValidation(t *testing.T) {
|
||||
path := filepath.Join(t.TempDir(), "state.lua")
|
||||
if err := os.WriteFile(path, []byte("value = 1"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
program, err := CompileFile(path)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
wantErr := errors.New("invalid script")
|
||||
var checked *glua.LState
|
||||
L, err := program.NewState(context.Background(), nil, func(L *glua.LState) error {
|
||||
checked = L
|
||||
return wantErr
|
||||
})
|
||||
if L != nil || !errors.Is(err, wantErr) || checked == nil || !checked.IsClosed() {
|
||||
t.Fatalf("state = %v, error = %v, checked state closed = %t", L, err, checked != nil && checked.IsClosed())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,95 @@
|
||||
package lua
|
||||
|
||||
import (
|
||||
"math"
|
||||
|
||||
"github.com/xtls/xray-core/common/errors"
|
||||
glua "github.com/yuin/gopher-lua"
|
||||
)
|
||||
|
||||
type number interface {
|
||||
~int | ~int8 | ~int16 | ~int32 | ~int64 |
|
||||
~uint | ~uint8 | ~uint16 | ~uint32 | ~uint64 | ~uintptr |
|
||||
~float32 | ~float64
|
||||
}
|
||||
|
||||
// PushNumber converts a Go number to a Lua number and pushes it.
|
||||
func PushNumber[T number](L *glua.LState, value T) {
|
||||
L.Push(glua.LNumber(value))
|
||||
}
|
||||
|
||||
// PushString converts a Go string to a Lua string and pushes it.
|
||||
func PushString(L *glua.LState, value string) {
|
||||
L.Push(glua.LString(value))
|
||||
}
|
||||
|
||||
// PushNil pushes Lua nil.
|
||||
func PushNil(L *glua.LState) {
|
||||
L.Push(glua.LNil)
|
||||
}
|
||||
|
||||
// PushUserData pushes a native Go value without copying it.
|
||||
func PushUserData(L *glua.LState, value any) {
|
||||
ud := L.NewUserData()
|
||||
ud.Value = value
|
||||
L.Push(ud)
|
||||
}
|
||||
|
||||
// PushError pushes nil or the original Go error as userdata.
|
||||
func PushError(L *glua.LState, err error) {
|
||||
if err == nil {
|
||||
L.Push(glua.LNil)
|
||||
return
|
||||
}
|
||||
PushUserData(L, err)
|
||||
}
|
||||
|
||||
// ReadUserData reads a native Go value of type T without copying it.
|
||||
// Other Lua values or userdata containing a different type return invalidMessage.
|
||||
func ReadUserData[T any](value glua.LValue, invalidMessage string) (T, error) {
|
||||
if ud, ok := value.(*glua.LUserData); ok {
|
||||
if result, ok := ud.Value.(T); ok {
|
||||
return result, nil
|
||||
}
|
||||
}
|
||||
var zero T
|
||||
return zero, errors.New(invalidMessage)
|
||||
}
|
||||
|
||||
// ReadError accepts nil, a native Go error, or a Lua string.
|
||||
// Native errors retain their identity; other values return invalidMessage.
|
||||
func ReadError(value glua.LValue, invalidMessage string) error {
|
||||
if value == glua.LNil {
|
||||
return nil
|
||||
}
|
||||
if ud, ok := value.(*glua.LUserData); ok {
|
||||
if err, ok := ud.Value.(error); ok {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if message, ok := value.(glua.LString); ok {
|
||||
return errors.New(string(message))
|
||||
}
|
||||
return errors.New(invalidMessage)
|
||||
}
|
||||
|
||||
// ReadUint32 accepts only integral Lua numbers in the uint32 range.
|
||||
func ReadUint32(value glua.LValue, invalidMessage string) (uint32, error) {
|
||||
number, ok := value.(glua.LNumber)
|
||||
if !ok || number < 0 || number > math.MaxUint32 || math.Trunc(float64(number)) != float64(number) {
|
||||
return 0, errors.New(invalidMessage)
|
||||
}
|
||||
return uint32(number), nil
|
||||
}
|
||||
|
||||
// ReadOptionalString accepts a Lua string or nil, which becomes an empty string.
|
||||
// It does not coerce other values to strings.
|
||||
func ReadOptionalString(value glua.LValue, invalidMessage string) (string, error) {
|
||||
if value == glua.LNil {
|
||||
return "", nil
|
||||
}
|
||||
if result, ok := value.(glua.LString); ok {
|
||||
return string(result), nil
|
||||
}
|
||||
return "", errors.New(invalidMessage)
|
||||
}
|
||||
@@ -0,0 +1,121 @@
|
||||
package lua
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"math"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
glua "github.com/yuin/gopher-lua"
|
||||
)
|
||||
|
||||
func TestReadUint32(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
value glua.LValue
|
||||
want uint32
|
||||
wantErr bool
|
||||
}{
|
||||
{name: "zero", value: glua.LNumber(0)},
|
||||
{name: "integer", value: glua.LNumber(45), want: 45},
|
||||
{name: "maximum", value: glua.LNumber(math.MaxUint32), want: math.MaxUint32},
|
||||
{name: "fraction", value: glua.LNumber(1.5), wantErr: true},
|
||||
{name: "negative", value: glua.LNumber(-1), wantErr: true},
|
||||
{name: "overflow", value: glua.LNumber(math.MaxUint32 + 1), wantErr: true},
|
||||
{name: "NaN", value: glua.LNumber(math.NaN()), wantErr: true},
|
||||
{name: "positive infinity", value: glua.LNumber(math.Inf(1)), wantErr: true},
|
||||
{name: "negative infinity", value: glua.LNumber(math.Inf(-1)), wantErr: true},
|
||||
{name: "nil", value: glua.LNil, wantErr: true},
|
||||
{name: "numeric string", value: glua.LString("45"), wantErr: true},
|
||||
{name: "boolean", value: glua.LTrue, wantErr: true},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
got, err := ReadUint32(tc.value, "invalid number")
|
||||
if got != tc.want || (err != nil) != tc.wantErr {
|
||||
t.Fatalf("ReadUint32() = %d, %v; want %d, error %t", got, err, tc.want, tc.wantErr)
|
||||
}
|
||||
if err != nil && !strings.Contains(err.Error(), "invalid number") {
|
||||
t.Fatalf("error = %v, want invalid number", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadOptionalString(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
value glua.LValue
|
||||
want string
|
||||
wantErr bool
|
||||
}{
|
||||
{name: "nil", value: glua.LNil},
|
||||
{name: "empty", value: glua.LString("")},
|
||||
{name: "string", value: glua.LString("out"), want: "out"},
|
||||
{name: "number", value: glua.LNumber(1), wantErr: true},
|
||||
{name: "boolean", value: glua.LFalse, wantErr: true},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
got, err := ReadOptionalString(tc.value, "invalid string")
|
||||
if got != tc.want || (err != nil) != tc.wantErr {
|
||||
t.Fatalf("ReadOptionalString() = %q, %v; want %q, error %t", got, err, tc.want, tc.wantErr)
|
||||
}
|
||||
if err != nil && !strings.Contains(err.Error(), "invalid string") {
|
||||
t.Fatalf("error = %v, want invalid string", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestUserDataRoundTrip(t *testing.T) {
|
||||
L := glua.NewState()
|
||||
defer L.Close()
|
||||
want := []int{1, 2}
|
||||
PushUserData(L, want)
|
||||
if L.GetTop() != 1 {
|
||||
t.Fatalf("stack top = %d, want 1", L.GetTop())
|
||||
}
|
||||
got, err := ReadUserData[[]int](L.Get(-1), "invalid userdata")
|
||||
if err != nil || len(got) != len(want) || &got[0] != &want[0] {
|
||||
t.Fatalf("userdata = %v, %v; want original slice", got, err)
|
||||
}
|
||||
PushUserData(L, []int(nil))
|
||||
if got, err := ReadUserData[[]int](L.Get(-1), "invalid userdata"); err != nil || got != nil {
|
||||
t.Fatalf("nil slice userdata = %v, %v", got, err)
|
||||
}
|
||||
for _, value := range []glua.LValue{glua.LNil, glua.LString("1"), L.NewTable(), L.Get(1)} {
|
||||
if got, err := ReadUserData[int](value, "invalid userdata"); got != 0 || err == nil || !strings.Contains(err.Error(), "invalid userdata") {
|
||||
t.Fatalf("ReadUserData(%v) = %d, %v; want invalid userdata", value, got, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestErrorRoundTrip(t *testing.T) {
|
||||
L := glua.NewState()
|
||||
defer L.Close()
|
||||
want := errors.New("upstream failed")
|
||||
for _, err := range []error{nil, want} {
|
||||
PushError(L, err)
|
||||
if L.GetTop() != 1 {
|
||||
t.Fatalf("stack top = %d, want 1", L.GetTop())
|
||||
}
|
||||
if err == nil && L.Get(-1) != glua.LNil {
|
||||
t.Fatalf("nil error pushed as %v", L.Get(-1))
|
||||
}
|
||||
if got := ReadError(L.Get(-1), "invalid error"); got != err {
|
||||
t.Fatalf("ReadError() = %v, want original error %v", got, err)
|
||||
}
|
||||
L.Pop(1)
|
||||
}
|
||||
for _, message := range []string{"script failed", ""} {
|
||||
if err := ReadError(glua.LString(message), "invalid error"); err == nil || !strings.Contains(err.Error(), message) {
|
||||
t.Fatalf("string error = %v, want %q", err, message)
|
||||
}
|
||||
}
|
||||
wrong := L.NewUserData()
|
||||
wrong.Value = "not a native error"
|
||||
for _, value := range []glua.LValue{glua.LTrue, glua.LNumber(1), L.NewTable(), wrong, L.NewUserData()} {
|
||||
if err := ReadError(value, "invalid error"); err == nil || !strings.Contains(err.Error(), "invalid error") {
|
||||
t.Fatalf("ReadError(%v) = %v, want invalid error", value, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user