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
+17 -15
View File
@@ -126,31 +126,32 @@ func newLuaClientQuery(L *lua.LState, client featureDNS.Client) *lua.LFunction {
})
}
// callLuaHook invokes HandleDNSQuery in the supplied state.
// Returned slices and IP bytes may share storage with DNS caches or matcher inputs.
func (s *DNS) callLuaHook(L *lua.LState, domain string, option featureDNS.IPOption) ([]net.IP, uint32, error) {
top := L.GetTop()
defer L.SetTop(top)
// callLuaQuery runs HandleDNSQuery and leaves (ips, ttl, err) on the stack.
func callLuaQuery(L *lua.LState, domain string, option featureDNS.IPOption) error {
fn := L.GetGlobal("HandleDNSQuery")
if fn.Type() != lua.LTFunction {
return nil, 0, errors.New("DNS script must define HandleDNSQuery(...)")
return errors.New("DNS script must define HandleDNSQuery(...)")
}
if err := L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true},
lua.LString(strings.ToLower(domain)), lua.LBool(option.IPv4Enable),
lua.LBool(option.IPv6Enable), lua.LBool(option.FakeEnable)); err != nil {
return nil, 0, err
}
return readLuaDNSResult(L.Get(-3), L.Get(-2), L.Get(-1))
return L.CallByParam(lua.P{Fn: fn, NRet: 3, Protect: true},
lua.LString(strings.ToLower(domain)),
lua.LBool(option.IPv4Enable),
lua.LBool(option.IPv6Enable),
lua.LBool(option.FakeEnable))
}
func readLuaDNSResult(addresses, ttlValue, errorValue lua.LValue) ([]net.IP, uint32, error) {
if err := xlua.ReadError(errorValue, "DNS script error must be an error or string"); err != nil {
// readLuaQueryResult reads (ips, ttl, err) from the stack without copying the IPs.
func readLuaQueryResult(L *lua.LState) ([]net.IP, uint32, error) {
if err := xlua.ReadError(L.Get(-1), "DNS script error must be an error or string"); err != nil {
return nil, 0, err
}
ttl, err := xlua.ReadUint32(ttlValue, "DNS script returned invalid TTL")
ttl, err := xlua.ReadUint32(L.Get(-2), "DNS script returned invalid TTL")
if err != nil {
return nil, 0, err
}
addresses := L.Get(-3)
if addresses == lua.LNil {
return nil, 0, featureDNS.ErrEmptyResponse
}
@@ -161,5 +162,6 @@ func readLuaDNSResult(addresses, ttlValue, errorValue lua.LValue) ([]net.IP, uin
if len(ips) == 0 {
return nil, 0, featureDNS.ErrEmptyResponse
}
return ips, ttl, nil
}