diff --git a/cmd/dialect-import/main_test.go b/cmd/dialect-import/main_test.go
index 6285f4975..92497a58c 100644
--- a/cmd/dialect-import/main_test.go
+++ b/cmd/dialect-import/main_test.go
@@ -2,6 +2,7 @@ package main
import (
"os"
+ "path/filepath"
"testing"
"github.com/stretchr/testify/require"
@@ -44,18 +45,27 @@ const testDialect = `
`
func TestRun(t *testing.T) {
- dir, err := os.MkdirTemp("", "gomavlib")
+ dir := t.TempDir()
+ t.Chdir(dir)
+
+ err := os.WriteFile("testdialect.xml", []byte(testDialect), 0o644)
require.NoError(t, err)
- defer os.RemoveAll(dir)
- os.Chdir(dir)
+ err = run([]string{"testdialect.xml"})
+ require.NoError(t, err)
- err = os.WriteFile("testdialect.xml", []byte(testDialect), 0o644)
+ _, err = os.Stat("testdialect/message_a_message.go")
require.NoError(t, err)
+}
- err = run([]string{"testdialect.xml"})
+func TestRunAbsolutePath(t *testing.T) {
+ dir := t.TempDir()
+ t.Chdir(dir)
+ err := os.WriteFile("testdialect.xml", []byte(testDialect), 0o644)
require.NoError(t, err)
+ err = run([]string{filepath.Join(dir, "testdialect.xml")})
+ require.NoError(t, err)
_, err = os.Stat("testdialect/message_a_message.go")
require.NoError(t, err)
}
diff --git a/pkg/conversion/conversion.go b/pkg/conversion/conversion.go
index eca2c4e44..6fb76fa88 100644
--- a/pkg/conversion/conversion.go
+++ b/pkg/conversion/conversion.go
@@ -235,11 +235,13 @@ var dialectTypeToGo = map[string]string{
func defAddrToName(pa string) string {
var b string
- u, err := url.ParseRequestURI(pa)
- if err == nil {
- b = path.Base(u.Path)
+ if strings.HasPrefix(pa, "http://") || strings.HasPrefix(pa, "https://") {
+ u, err := url.Parse(pa)
+ if err == nil {
+ b = path.Base(u.Path)
+ }
} else {
- b = path.Base(pa)
+ b = filepath.Base(pa)
}
b = strings.TrimSuffix(b, path.Ext(b))
@@ -324,9 +326,52 @@ type outDefinition struct {
Messages []*outMessage
}
+func download(addr string) ([]byte, error) {
+ res, err := http.Get(addr)
+ if err != nil {
+ return nil, err
+ }
+ defer res.Body.Close()
+
+ if res.StatusCode != http.StatusOK {
+ return nil, fmt.Errorf("bad return code: %v", res.StatusCode)
+ }
+
+ byt, err := io.ReadAll(&customLimitReader{res.Body, maxInboundDialectSize})
+ if err != nil {
+ return nil, err
+ }
+ return byt, nil
+}
+
+func readLocalDefinition(root *os.Root, processedFiles *[]os.FileInfo, defAddr string) ([]byte, bool, error) {
+ file, err := root.Open(defAddr)
+ if err != nil {
+ return nil, false, fmt.Errorf("unable to open %q: %w", defAddr, err)
+ }
+ defer file.Close()
+
+ info, err := file.Stat()
+ if err != nil {
+ return nil, false, err
+ }
+
+ for _, previous := range *processedFiles {
+ if os.SameFile(previous, info) {
+ return nil, true, nil
+ }
+ }
+ *processedFiles = append(*processedFiles, info)
+
+ content, err := io.ReadAll(file)
+ return content, false, err
+}
+
func processDefinition(
version *string,
processedDefs map[string]struct{},
+ processedFiles *[]os.FileInfo,
+ root *os.Root,
isRemote bool,
defAddr string,
) ([]*outDefinition, error) {
@@ -338,7 +383,20 @@ func processDefinition(
fmt.Fprintf(os.Stderr, "processing definition %s\n", defAddr)
- content, err := getDefinition(isRemote, defAddr)
+ var content []byte
+ var err error
+ if isRemote {
+ content, err = download(defAddr)
+ if err != nil {
+ return nil, fmt.Errorf("unable to download: %w", err)
+ }
+ } else {
+ var alreadyProcessed bool
+ content, alreadyProcessed, err = readLocalDefinition(root, processedFiles, defAddr)
+ if alreadyProcessed {
+ return nil, nil
+ }
+ }
if err != nil {
return nil, err
}
@@ -354,12 +412,31 @@ func processDefinition(
// includes
for _, subDefAddr := range def.Includes {
- // prepend url to remote address
+ include := subDefAddr
+
+ // resolve remote URLs or local paths relative to the current definition
if isRemote {
subDefAddr = addrPath + subDefAddr
+ } else {
+ if filepath.IsAbs(subDefAddr) {
+ subDefAddr, err = filepath.Rel(root.Name(), subDefAddr)
+ if err != nil {
+ return nil, fmt.Errorf("invalid include %q: %w", include, err)
+ }
+ } else {
+ if filepath.VolumeName(subDefAddr) != "" || strings.HasPrefix(subDefAddr, string(filepath.Separator)) {
+ return nil, fmt.Errorf("invalid include %q: outside dialect root", subDefAddr)
+ }
+ subDefAddr = filepath.Join(addrPath, subDefAddr)
+ }
+
+ if !filepath.IsLocal(subDefAddr) {
+ return nil, fmt.Errorf("invalid include %q: outside dialect root", include)
+ }
}
+
var subDefs []*outDefinition
- subDefs, err = processDefinition(version, processedDefs, isRemote, subDefAddr)
+ subDefs, err = processDefinition(version, processedDefs, processedFiles, root, isRemote, subDefAddr)
if err != nil {
return nil, err
}
@@ -454,40 +531,6 @@ func processDefinition(
return outDefs, nil
}
-func getDefinition(isRemote bool, defAddr string) ([]byte, error) {
- if isRemote {
- byt, err := download(defAddr)
- if err != nil {
- return nil, fmt.Errorf("unable to download: %w", err)
- }
- return byt, nil
- }
-
- byt, err := os.ReadFile(defAddr)
- if err != nil {
- return nil, fmt.Errorf("unable to open: %w", err)
- }
- return byt, nil
-}
-
-func download(addr string) ([]byte, error) {
- res, err := http.Get(addr)
- if err != nil {
- return nil, err
- }
- defer res.Body.Close()
-
- if res.StatusCode != http.StatusOK {
- return nil, fmt.Errorf("bad return code: %v", res.StatusCode)
- }
-
- byt, err := io.ReadAll(&customLimitReader{res.Body, maxInboundDialectSize})
- if err != nil {
- return nil, err
- }
- return byt, nil
-}
-
func processMessage(defName string, msgDef *definitionMessage) (*outMessage, error) {
if m := reMsgName.FindStringSubmatch(msgDef.Name); m == nil {
return nil, fmt.Errorf("unsupported message name: %s", msgDef.Name)
@@ -645,11 +688,28 @@ func writeMessage(
func Convert(path string, link bool) error {
version := ""
processedDefs := make(map[string]struct{})
- _, err := url.ParseRequestURI(path)
- isRemote := (err == nil)
+ var processedFiles []os.FileInfo
defName := defAddrToName(path)
- _, err = os.Stat(defName)
+ var root *os.Root
+ isRemote := strings.HasPrefix(path, "http://") || strings.HasPrefix(path, "https://")
+
+ if !isRemote {
+ absPath, err := filepath.Abs(path)
+ if err != nil {
+ return err
+ }
+
+ root, err = os.OpenRoot(filepath.Dir(absPath))
+ if err != nil {
+ return err
+ }
+ defer root.Close() //nolint:errcheck
+
+ path = filepath.Base(absPath)
+ }
+
+ _, err := os.Stat(defName)
if !os.IsNotExist(err) {
return fmt.Errorf("directory '%s' already exists", defName)
}
@@ -657,7 +717,7 @@ func Convert(path string, link bool) error {
os.Mkdir(defName, 0o755)
// parse all definitions recursively
- outDefs, err := processDefinition(&version, processedDefs, isRemote, path)
+ outDefs, err := processDefinition(&version, processedDefs, &processedFiles, root, isRemote, path)
if err != nil {
return err
}
diff --git a/pkg/conversion/conversion_test.go b/pkg/conversion/conversion_test.go
index 9ff47a4a2..dfed805fc 100644
--- a/pkg/conversion/conversion_test.go
+++ b/pkg/conversion/conversion_test.go
@@ -1,7 +1,11 @@
package conversion_test
import (
+ "net/http"
+ "net/http/httptest"
"os"
+ "path/filepath"
+ "strings"
"testing"
"github.com/stretchr/testify/require"
@@ -167,13 +171,10 @@ func (e A_TYPE) String() string {
`
func TestConversion(t *testing.T) {
- dir, err := os.MkdirTemp("", "gomavlib")
- require.NoError(t, err)
- defer os.RemoveAll(dir)
-
- os.Chdir(dir)
+ dir := t.TempDir()
+ t.Chdir(dir)
- err = os.WriteFile("testdialect.xml", []byte(testDialect), 0o644)
+ err := os.WriteFile("testdialect.xml", []byte(testDialect), 0o644)
require.NoError(t, err)
err = conversion.Convert("testdialect.xml", true)
@@ -187,3 +188,186 @@ func TestConversion(t *testing.T) {
require.NoError(t, err)
require.Equal(t, testEnumGo, string(buf))
}
+
+func TestConversionRelativeIncludes(t *testing.T) {
+ dir := t.TempDir()
+ t.Chdir(dir)
+ require.NoError(t, os.MkdirAll("dialects/sub", 0o755))
+ require.NoError(t, os.WriteFile("dialects/main.xml", []byte(`
+
+ sub/child.xml
+
+`), 0o644))
+ require.NoError(t, os.WriteFile("dialects/sub/child.xml", []byte(`
+
+ sibling.xml
+
+`), 0o644))
+ require.NoError(t, os.WriteFile("dialects/sub/sibling.xml", []byte(`
+
+
+
+
+
+`), 0o644))
+
+ require.NoError(t, conversion.Convert("dialects/main.xml", true))
+ _, err := os.Stat("main/message_sibling_message.go")
+ require.NoError(t, err)
+}
+
+func TestConversionRemoteIncludes(t *testing.T) {
+ server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ switch r.URL.Path {
+ case "/main.xml":
+ _, _ = w.Write([]byte(`child.xml`))
+ case "/child.xml":
+ _, _ = w.Write([]byte(``))
+ default:
+ http.NotFound(w, r)
+ }
+ }))
+ defer server.Close()
+
+ dir := t.TempDir()
+ t.Chdir(dir)
+ require.NoError(t, conversion.Convert(server.URL+"/main.xml", false))
+ _, err := os.Stat("main/message_child_message.go")
+ require.NoError(t, err)
+}
+
+func TestConversionIncludesWithinRoot(t *testing.T) {
+ dir := t.TempDir()
+ t.Chdir(dir)
+ require.NoError(t, os.MkdirAll("dialects/sub", 0o755))
+ require.NoError(t, os.WriteFile("dialects/main.xml", []byte(`
+ sub/child.xml
+`), 0o644))
+ require.NoError(t, os.WriteFile("dialects/sub/child.xml", []byte(`
+ ../sibling.xml
+`), 0o644))
+ require.NoError(t, os.WriteFile("dialects/sibling.xml", []byte(`
+
+`), 0o644))
+
+ require.NoError(t, conversion.Convert("./dialects/main.xml", false))
+ _, err := os.Stat("main/message_sibling_message.go")
+ require.NoError(t, err)
+}
+
+func TestConversionAbsoluteIncludesWithinRoot(t *testing.T) {
+ dir := t.TempDir()
+ t.Chdir(dir)
+ require.NoError(t, os.Mkdir("dialects", 0o755))
+ include := filepath.Join(dir, "dialects", "sibling.xml")
+ require.NoError(t, os.WriteFile("dialects/main.xml",
+ []byte(""+include+""), 0o644))
+ require.NoError(t, os.WriteFile(include, []byte(`
+
+`), 0o644))
+
+ require.NoError(t, conversion.Convert(filepath.Join(dir, "dialects", "main.xml"), false))
+ _, err := os.Stat("main/message_sibling_message.go")
+ require.NoError(t, err)
+}
+
+func TestConversionRejectsEscapingIncludes(t *testing.T) {
+ for _, ca := range []struct {
+ name string
+ include func(string) string
+ }{
+ {"parent", func(string) string { return "../outside.xml" }},
+ {"sibling prefix", func(string) string { return "../dialects-other/outside.xml" }},
+ {"absolute", func(dir string) string { return filepath.Join(dir, "outside.xml") }},
+ } {
+ t.Run(ca.name, func(t *testing.T) {
+ dir := t.TempDir()
+ t.Chdir(dir)
+ require.NoError(t, os.Mkdir("dialects", 0o755))
+ require.NoError(t, os.Mkdir("dialects-other", 0o755))
+ require.NoError(t, os.WriteFile("outside.xml", []byte(""), 0o644))
+ require.NoError(t, os.WriteFile("dialects-other/outside.xml", []byte(""), 0o644))
+ require.NoError(t, os.WriteFile("dialects/main.xml",
+ []byte(""+ca.include(dir)+""), 0o644))
+ require.ErrorContains(t, conversion.Convert("dialects/main.xml", false), "outside dialect root")
+ })
+ }
+}
+
+func TestConversionRejectsSymlinkEscape(t *testing.T) {
+ dir := t.TempDir()
+ t.Chdir(dir)
+ require.NoError(t, os.Mkdir("dialects", 0o755))
+ require.NoError(t, os.WriteFile("outside.xml", []byte(""), 0o644))
+ if err := os.Symlink(filepath.Join(dir, "outside.xml"), "dialects/escape.xml"); err != nil {
+ t.Skipf("symlinks unavailable: %v", err)
+ }
+ require.NoError(t, os.WriteFile("dialects/main.xml",
+ []byte("escape.xml"), 0o644))
+ require.Error(t, conversion.Convert("dialects/main.xml", false))
+}
+
+func TestConversionAllowsInRootSymlink(t *testing.T) {
+ dir := t.TempDir()
+ t.Chdir(dir)
+ require.NoError(t, os.Mkdir("dialects", 0o755))
+ if err := os.Symlink("sibling.xml", "dialects/alias.xml"); err != nil {
+ t.Skipf("symlinks unavailable: %v", err)
+ }
+ require.NoError(t, os.WriteFile("dialects/main.xml", []byte("alias.xml"), 0o644))
+ require.NoError(t, os.WriteFile("dialects/sibling.xml", []byte(`
+
+`), 0o644))
+ require.NoError(t, conversion.Convert("dialects/main.xml", false))
+ _, err := os.Stat("main/message_sibling_message.go")
+ require.NoError(t, err)
+}
+
+func TestConversionCycle(t *testing.T) {
+ dir := t.TempDir()
+ t.Chdir(dir)
+ require.NoError(t, os.MkdirAll("dialects/sub", 0o755))
+ require.NoError(t, os.WriteFile("dialects/main.xml", []byte(`
+ sub/child.xml
+
+`), 0o644))
+ require.NoError(t, os.WriteFile("dialects/sub/child.xml", []byte(`
+ ../main.xml
+`), 0o644))
+
+ require.NoError(t, conversion.Convert("./dialects/main.xml", false))
+ buf, err := os.ReadFile("main/dialect.go")
+ require.NoError(t, err)
+ require.Equal(t, 1, strings.Count(string(buf), "&MessageRootMessage{}"))
+}
+
+func TestConversionRejectsDriveRelativeInclude(t *testing.T) {
+ if filepath.Separator != '\\' {
+ t.Skip("drive-relative paths are Windows-specific")
+ }
+ dir := t.TempDir()
+ t.Chdir(dir)
+ require.NoError(t, os.Mkdir("dialects", 0o755))
+ require.NoError(t, os.WriteFile("dialects/main.xml",
+ []byte(`C:outside.xml`), 0o644))
+ require.ErrorContains(t, conversion.Convert("dialects/main.xml", false),
+ "outside dialect root")
+}
+
+func TestConversionSymlinkCycle(t *testing.T) {
+ dir := t.TempDir()
+ t.Chdir(dir)
+ require.NoError(t, os.Mkdir("dialects", 0o755))
+ if err := os.Symlink("main.xml", "dialects/alias.xml"); err != nil {
+ t.Skipf("symlinks unavailable: %v", err)
+ }
+ require.NoError(t, os.WriteFile("dialects/main.xml", []byte(`
+ alias.xml
+
+`), 0o644))
+
+ require.NoError(t, conversion.Convert("dialects/main.xml", false))
+ buf, err := os.ReadFile("main/dialect.go")
+ require.NoError(t, err)
+ require.Equal(t, 1, strings.Count(string(buf), "&MessageRootMessage{}"))
+}