dns: preserve Lua states after recoverable query errors & refactor

This commit is contained in:
Meo597
2026-10-05 06:15:49 +08:00
parent 54610a39a5
commit 248cb666f3
4 changed files with 207 additions and 133 deletions
+80 -95
View File
@@ -3,7 +3,6 @@ package dns
import (
"context"
go_errors "errors"
"math"
"strings"
"testing"
"time"
@@ -15,70 +14,74 @@ import (
lua "github.com/yuin/gopher-lua"
)
func TestReadLuaDNSResult(t *testing.T) {
L := lua.NewState()
defer L.Close()
want := []net.IP{net.ParseIP("8.8.8.8"), {127, 0, 0, 1}, net.ParseIP("::1")}
addresses := L.NewUserData()
addresses.Value = want
ips, ttl, err := readLuaDNSResult(addresses, lua.LNumber(45), lua.LNil)
if err != nil || ttl != 45 || len(ips) != len(want) {
t.Fatalf("readLuaDNSResult() = %v, TTL %d, %v", ips, ttl, err)
}
for i := range want {
if !ips[i].Equal(want[i]) {
t.Fatalf("IP %d = %v, want %v", i, ips[i], want[i])
}
}
}
func TestReadLuaDNSResultValidation(t *testing.T) {
L := lua.NewState()
defer L.Close()
func TestReadLuaQueryResult(t *testing.T) {
wantIPs := []net.IP{net.ParseIP("8.8.8.8"), {127, 0, 0, 1}, net.ParseIP("::1")}
nativeErr := go_errors.New("upstream failed")
for _, tc := range []struct {
name string
change func(*[3]lua.LValue)
want string
name, values string
wantIPs []net.IP
wantTTL uint32
wantErr error
wantMessage string
}{
{"fractional TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(1.5) }, "invalid TTL"},
{"oversized TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(4294967296) }, "invalid TTL"},
{"negative TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(-1) }, "invalid TTL"},
{"NaN TTL", func(v *[3]lua.LValue) { v[1] = lua.LNumber(math.NaN()) }, "invalid TTL"},
{"missing TTL", func(v *[3]lua.LValue) { v[1] = lua.LNil }, "invalid TTL"},
{"string IPs", func(v *[3]lua.LValue) { v[0] = lua.LString("127.0.0.1") }, "native IP slice"},
{"wrong userdata", func(v *[3]lua.LValue) { v[0].(*lua.LUserData).Value = net.ParseIP("127.0.0.1") }, "native IP slice"},
{"script error", func(v *[3]lua.LValue) { v[2] = lua.LString("blocked by script") }, "blocked by script"},
{"invalid error", func(v *[3]lua.LValue) { v[2] = lua.LTrue }, "error or string"},
{name: "IPs", values: `ips, 45`, wantIPs: wantIPs, wantTTL: 45},
{name: "nil IPs", values: `nil, 0`, wantErr: featureDNS.ErrEmptyResponse},
{name: "empty IPs", values: `emptyIPs, 0`, wantErr: featureDNS.ErrEmptyResponse},
{name: "native error", values: `nil, nil, nativeError`, wantErr: nativeErr},
{name: "string error", values: `nil, nil, "blocked"`, wantMessage: "blocked"},
{name: "fractional TTL", values: `ips, 1.5`, wantMessage: "invalid TTL"},
{name: "oversized TTL", values: `ips, 4294967296`, wantMessage: "invalid TTL"},
{name: "negative TTL", values: `ips, -1`, wantMessage: "invalid TTL"},
{name: "NaN TTL", values: `ips, 0/0`, wantMessage: "invalid TTL"},
{name: "missing TTL", values: `ips`, wantMessage: "invalid TTL"},
{name: "string IPs", values: `"127.0.0.1", 60`, wantMessage: "native IP slice"},
{name: "wrong userdata", values: `ip, 60`, wantMessage: "native IP slice"},
{name: "invalid error", values: `ips, 60, false`, wantMessage: "error or string"},
} {
t.Run(tc.name, func(t *testing.T) {
addresses := L.NewUserData()
addresses.Value = []net.IP{net.ParseIP("127.0.0.1")}
values := [3]lua.LValue{addresses, lua.LNumber(60), lua.LNil}
tc.change(&values)
_, _, err := readLuaDNSResult(values[0], values[1], values[2])
if err == nil || !strings.Contains(err.Error(), tc.want) {
t.Fatalf("readLuaDNSResult error = %v, want %q", err, tc.want)
L := lua.NewState()
defer L.Close()
for name, value := range map[string]any{"ips": wantIPs, "ip": wantIPs[0], "emptyIPs": []net.IP(nil), "nativeError": nativeErr} {
ud := L.NewUserData()
ud.Value = value
L.SetGlobal(name, ud)
}
fn, err := L.LoadString("return " + tc.values)
if err != nil {
t.Fatal(err)
}
if err := L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true}); err != nil {
t.Fatal(err)
}
ips, ttl, err := readLuaQueryResult(L)
switch {
case tc.wantErr != nil:
if err != tc.wantErr {
t.Fatalf("error = %v, want original error %v", err, tc.wantErr)
}
case tc.wantMessage != "":
if err == nil || !strings.Contains(err.Error(), tc.wantMessage) {
t.Fatalf("error = %v, want %q", err, tc.wantMessage)
}
case err != nil:
t.Fatal(err)
}
if ttl != tc.wantTTL || len(ips) != len(tc.wantIPs) {
t.Fatalf("result = %v, TTL %d; want %v, TTL %d", ips, ttl, tc.wantIPs, tc.wantTTL)
}
for i := range ips {
if !ips[i].Equal(tc.wantIPs[i]) {
t.Fatalf("IP %d = %v, want %v", i, ips[i], tc.wantIPs[i])
}
}
if len(ips) != 0 && &ips[0] != &tc.wantIPs[0] {
t.Fatal("result copied the IP slice")
}
})
}
addresses := L.NewUserData()
addresses.Value = []net.IP(nil)
for _, empty := range []lua.LValue{addresses, lua.LNil} {
if _, _, err := readLuaDNSResult(empty, lua.LNumber(0), lua.LNil); !go_errors.Is(err, featureDNS.ErrEmptyResponse) {
t.Fatalf("empty result error = %v, want ErrEmptyResponse", err)
}
}
wantErr := go_errors.New("upstream failed")
errorValue := L.NewUserData()
errorValue.Value = wantErr
if _, _, err := readLuaDNSResult(lua.LNil, lua.LNil, errorValue); err != wantErr {
t.Fatalf("upstream error = %v, want original error %v", err, wantErr)
}
}
func TestCallLuaHookCancellation(t *testing.T) {
func TestCallLuaQueryCancellation(t *testing.T) {
L := lua.NewState()
defer L.Close()
if err := L.DoString(`function HandleDNSQuery(domain, ipv4, ipv6, fake) while true do end end`); err != nil {
@@ -87,16 +90,16 @@ func TestCallLuaHookCancellation(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
cancel()
L.SetContext(ctx)
_, _, err := (&DNS{}).callLuaHook(L, "example.com", featureDNS.IPOption{IPv4Enable: true})
err := callLuaQuery(L, "example.com", featureDNS.IPOption{IPv4Enable: true})
if err == nil {
t.Fatal("CallLuaHook did not stop after context cancellation")
t.Fatal("callLuaQuery did not stop after context cancellation")
}
if L.Context() != ctx {
t.Fatal("CallLuaHook changed the Lua state's context")
t.Fatal("callLuaQuery changed the Lua state's context")
}
}
func TestCallLuaHookNormalizesDomain(t *testing.T) {
func TestCallLuaQuery(t *testing.T) {
L := lua.NewState()
defer L.Close()
addresses := L.NewUserData()
@@ -111,39 +114,11 @@ func TestCallLuaHookNormalizesDomain(t *testing.T) {
`); err != nil {
t.Fatal(err)
}
s := &DNS{}
if _, _, err := s.callLuaHook(L, "ExAmPlE.CoM", featureDNS.IPOption{IPv4Enable: true}); err != nil {
if err := callLuaQuery(L, "ExAmPlE.CoM", featureDNS.IPOption{IPv4Enable: true}); err != nil {
t.Fatal(err)
}
}
func TestCallLuaHookRestoresStack(t *testing.T) {
for _, tc := range []struct {
name string
body string
wantErr bool
}{
{"success", `return ips, 60`, false},
{"error", `error("failed")`, true},
} {
t.Run(tc.name, func(t *testing.T) {
L := lua.NewState()
defer L.Close()
addresses := L.NewUserData()
addresses.Value = []net.IP{net.ParseIP("127.0.0.1")}
L.SetGlobal("ips", addresses)
if err := L.DoString("function HandleDNSQuery() " + tc.body + " end"); err != nil {
t.Fatal(err)
}
L.Push(lua.LTrue)
_, _, err := (&DNS{}).callLuaHook(L, "example.com", featureDNS.IPOption{IPv4Enable: true})
if (err != nil) != tc.wantErr {
t.Fatalf("hook error = %v, want error %t", err, tc.wantErr)
}
if L.GetTop() != 1 || L.Get(1) != lua.LTrue {
t.Fatal("hook did not restore the stack")
}
})
if L.GetTop() != 3 || L.Get(1) != addresses || L.Get(2) != lua.LNumber(60) || L.Get(3) != lua.LNil {
t.Fatal("callLuaQuery did not leave the three query results on the stack")
}
}
@@ -169,7 +144,10 @@ end
t.Fatal(err)
}
L.SetContext(context.Background())
got, ttl, err := server.callLuaHook(L, "example.com", option)
if err := callLuaQuery(L, "example.com", option); err != nil {
t.Fatal(err)
}
got, ttl, err := readLuaQueryResult(L)
if err != nil || ttl != 60 || len(got) != 1 || !got[0].Equal(ips[0]) {
t.Fatalf("server query = %v, TTL %d, %v", got, ttl, err)
}
@@ -244,9 +222,9 @@ func (s *benchmarkLuaNameServer) QueryIP(context.Context, string, featureDNS.IPO
return s.ips, 60, nil
}
// BenchmarkLuaDNSHookCall isolates a preloaded Lua hook and its server:Query bridge.
// BenchmarkLuaDNSQuery measures a preloaded DNS script using server:Query.
// The direct case measures the same DNS client without Lua.
func BenchmarkLuaDNSHookCall(b *testing.B) {
func BenchmarkLuaDNSQuery(b *testing.B) {
option := featureDNS.IPOption{IPv4Enable: true}
ip := net.ParseIP("127.0.0.1")
upstream := &benchmarkLuaNameServer{ips: []net.IP{ip}}
@@ -271,7 +249,14 @@ end
query func() ([]net.IP, uint32, error)
}{
{"direct", func() ([]net.IP, uint32, error) { return client.QueryIP(ctx, "example.com", option) }},
{"lua_hook", func() ([]net.IP, uint32, error) { return server.callLuaHook(L, "example.com", option) }},
{"lua_script", func() ([]net.IP, uint32, error) {
if err := callLuaQuery(L, "example.com", option); err != nil {
return nil, 0, err
}
ips, ttl, err := readLuaQueryResult(L)
L.Pop(3)
return ips, ttl, err
}},
} {
b.Run(bench.name, func(b *testing.B) {
b.ReportAllocs()