diff --git a/app/dns/script.go b/app/dns/script.go index e1b7ce3ac..160f703ae 100644 --- a/app/dns/script.go +++ b/app/dns/script.go @@ -6,6 +6,7 @@ import ( "github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/geodata" + "github.com/xtls/xray-core/common/log" luamgr "github.com/xtls/xray-core/common/lua" "github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/features/dns" @@ -30,6 +31,7 @@ func newScriptEngine(path string, server *DNS) (*scriptEngine, error) { defer cancel() L, err := program.NewState(initCtx, func(L *lua.LState) { geodata.RegisterLua(L) + log.RegisterLua(L) server.RegisterLua(L) }) if err != nil { diff --git a/app/dns/script_test.go b/app/dns/script_test.go index 835fe493e..a3849deb1 100644 --- a/app/dns/script_test.go +++ b/app/dns/script_test.go @@ -142,9 +142,14 @@ func TestDNSScriptHookErrorAndFakeDNSOption(t *testing.T) { path := filepath.Join(t.TempDir(), "script.lua") script := ` local server = require("xray.dns").servers[1] +local log = require("xray.log") +log.info("DNS script loaded") function handleDNSQuery(q) + log.debug("DNS query: ", q.domain) if q.domain == "bad.example" then error("script failure") end - return server:query(q) + local answer = server:query(q) + if answer.error then log.error("DNS failed: ", answer.error) end + return answer end ` if err := os.WriteFile(path, []byte(script), 0o600); err != nil { diff --git a/common/log/lua.go b/common/log/lua.go new file mode 100644 index 000000000..52bc8abff --- /dev/null +++ b/common/log/lua.go @@ -0,0 +1,44 @@ +package log + +import ( + "strings" + + lua "github.com/yuin/gopher-lua" +) + +// RegisterLua makes xray.log available to require in an LState. +// Logging functions concatenate arguments using Lua's tostring semantics, +// except Go errors in userdata use Error(). Messages go through the current +// log handler. +func RegisterLua(L *lua.LState) { + L.PreloadModule("xray.log", func(L *lua.LState) int { + module := L.NewTable() + for name, severity := range map[string]Severity{ + "debug": Severity_Debug, + "info": Severity_Info, + "warning": Severity_Warning, + "error": Severity_Error, + } { + module.RawSetString(name, L.NewFunction(func(L *lua.LState) int { + var content strings.Builder + for i := 1; i <= L.GetTop(); i++ { + value := L.Get(i) + if ud, ok := value.(*lua.LUserData); ok { + if err, ok := ud.Value.(error); ok { + content.WriteString(err.Error()) + continue + } + } + content.WriteString(L.ToStringMeta(value).String()) + } + Record(&GeneralMessage{ + Severity: severity, + Content: content.String(), + }) + return 0 + })) + } + L.Push(module) + return 1 + }) +} diff --git a/common/log/lua_test.go b/common/log/lua_test.go new file mode 100644 index 000000000..10550fcba --- /dev/null +++ b/common/log/lua_test.go @@ -0,0 +1,75 @@ +package log + +import ( + "errors" + "fmt" + "testing" + + lua "github.com/yuin/gopher-lua" +) + +type luaLogHandler struct { + messages []Message +} + +func (h *luaLogHandler) Handle(msg Message) { + h.messages = append(h.messages, msg) +} + +func TestLuaLog(t *testing.T) { + logHandler.RLock() + previous := logHandler.Handler + logHandler.RUnlock() + t.Cleanup(func() { RegisterHandler(previous) }) + handler := &luaLogHandler{} + RegisterHandler(handler) + + L := lua.NewState() + defer L.Close() + RegisterLua(L) + nativeError := L.NewUserData() + nativeError.Value = fmt.Errorf("lookup failed: %w", errors.New("upstream timeout")) + L.SetGlobal("nativeError", nativeError) + if err := L.DoString(` + local log = require("xray.log") + assert(log == require("xray.log")) + log.debug("query: ", "example.com") + log.info("count=", 42, ", enabled=", true, ", value=", nil) + log.warning(setmetatable({}, { + __tostring = function() return "fallback" end + })) + assert(select("#", log.error("failed")) == 0) + log.error("DNS failed: ", nativeError) + log.warning(nativeError) + local ok, err = pcall(function() error("Lua failure", 0) end) + assert(not ok) + log.error(err) + `); err != nil { + t.Fatal(err) + } + + want := []struct { + severity Severity + message string + }{ + {Severity_Debug, "[Debug] query: example.com"}, + {Severity_Info, "[Info] count=42, enabled=true, value=nil"}, + {Severity_Warning, "[Warning] fallback"}, + {Severity_Error, "[Error] failed"}, + {Severity_Error, "[Error] DNS failed: lookup failed: upstream timeout"}, + {Severity_Warning, "[Warning] lookup failed: upstream timeout"}, + {Severity_Error, "[Error] Lua failure"}, + } + if len(handler.messages) != len(want) { + t.Fatalf("logged %d messages, want %d", len(handler.messages), len(want)) + } + for i, expected := range want { + msg, ok := handler.messages[i].(*GeneralMessage) + if !ok { + t.Fatalf("message %d has type %T, want *GeneralMessage", i, handler.messages[i]) + } + if msg.Severity != expected.severity || msg.String() != expected.message { + t.Errorf("message %d = %q with severity %v, want %q with severity %v", i, msg.String(), msg.Severity, expected.message, expected.severity) + } + } +}