diff --git a/common/platform/platform.go b/common/platform/platform.go index b60b8bd22..4c5cf2a0f 100644 --- a/common/platform/platform.go +++ b/common/platform/platform.go @@ -1,6 +1,8 @@ package platform // import "github.com/xtls/xray-core/common/platform" import ( + "errors" + "fmt" "os" "path/filepath" "strconv" @@ -90,3 +92,49 @@ func GetConfDirPath() string { configPath := NewEnvFlag(ConfdirLocation).GetValue(func() string { return "" }) return configPath } + +// ResolveLuaFile finds a local Lua script and returns its absolute path. +// Relative paths: XRAY_LOCATION_CONFDIR > XRAY_LOCATION_CONFIG > working dir > executable dir. +func ResolveLuaFile(path string) (string, error) { + if path == "" { + return "", errors.New("Lua file path is empty") + } + paths := []string{path} + if !filepath.IsAbs(path) { + paths = nil + for _, dir := range []string{ + GetConfDirPath(), + NewEnvFlag(ConfigLocation).GetValue(func() string { return "" }), + ".", + getExecutableDir(), + } { + if dir != "" { + paths = append(paths, filepath.Join(dir, path)) + } + } + } + return resolveFile(paths) +} + +func resolveFile(paths []string) (string, error) { + var tried []string + for _, path := range paths { + path, err := filepath.Abs(path) + if err != nil { + return "", fmt.Errorf("failed to resolve file path: %w", err) + } + tried = append(tried, path) + info, err := os.Stat(path) + if errors.Is(err, os.ErrNotExist) { + continue + } + if err != nil { + return "", fmt.Errorf("failed to inspect file %q: %w", path, err) + } + if !info.Mode().IsRegular() { + return "", fmt.Errorf("file is not a regular file: %s", path) + } + return path, nil + } + return "", fmt.Errorf("file not found; tried %q: %w", tried, os.ErrNotExist) +} diff --git a/common/platform/platform_test.go b/common/platform/platform_test.go index 854c397f6..95d5b6545 100644 --- a/common/platform/platform_test.go +++ b/common/platform/platform_test.go @@ -1,6 +1,7 @@ package platform_test import ( + "errors" "os" "path/filepath" "runtime" @@ -64,3 +65,53 @@ func TestGetAssetLocation(t *testing.T) { } } } + +func TestResolveLuaFile(t *testing.T) { + workingDir := t.TempDir() + t.Chdir(workingDir) + executable, err := os.Executable() + common.Must(err) + file, err := os.CreateTemp(filepath.Dir(executable), "lua-*.lua") + common.Must(err) + common.Must(file.Close()) + defer os.Remove(file.Name()) + + name := filepath.Base(file.Name()) + paths := []string{ + filepath.Join(t.TempDir(), name), + filepath.Join(t.TempDir(), name), + filepath.Join(workingDir, name), + file.Name(), + } + t.Setenv(ConfdirLocation, filepath.Dir(paths[0])) + t.Setenv(ConfigLocation, filepath.Dir(paths[1])) + for _, path := range paths[:3] { + common.Must(os.WriteFile(path, nil, 0o600)) + } + if got, err := ResolveLuaFile(paths[2]); err != nil || got != paths[2] { + t.Fatalf("absolute path = %q, %v; want %q", got, err, paths[2]) + } + for i, want := range paths { + if i == 2 { + t.Setenv(ConfdirLocation, "") + t.Setenv(ConfigLocation, "") + } + if got, err := ResolveLuaFile(name); err != nil || got != want { + t.Fatalf("resolved path = %q, %v; want %q", got, err, want) + } + common.Must(os.Remove(want)) + } + if _, err := ResolveLuaFile(name); !errors.Is(err, os.ErrNotExist) { + t.Fatalf("missing file error = %v", err) + } + + t.Setenv(ConfdirLocation, filepath.Dir(paths[0])) + t.Setenv(ConfigLocation, filepath.Dir(paths[1])) + common.Must(os.Mkdir(paths[0], 0o700)) + common.Must(os.WriteFile(paths[1], nil, 0o600)) + for _, path := range []string{"", name, filepath.Join(t.TempDir(), name)} { + if _, err := ResolveLuaFile(path); err == nil { + t.Fatalf("accepted invalid path %q", path) + } + } +} diff --git a/infra/conf/dns.go b/infra/conf/dns.go index a05814227..07b482742 100644 --- a/infra/conf/dns.go +++ b/infra/conf/dns.go @@ -285,16 +285,9 @@ func (c *DNSConfig) Build() (*dns.Config, error) { } if c.Script != "" { - path := c.Script - if !filepath.IsAbs(path) { - path = filepath.Join(platform.GetConfDirPath(), path) - } - info, err := os.Stat(path) + path, err := platform.ResolveLuaFile(c.Script) if err != nil { - return nil, errors.New("DNS script does not exist: ", path).Base(err) - } - if !info.Mode().IsRegular() { - return nil, errors.New("DNS script is not a regular file: ", path) + return nil, errors.New("failed to resolve DNS script: ", c.Script).Base(err) } config.Script = path }