Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 15 additions & 5 deletions cmd/dialect-import/main_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package main

import (
"os"
"path/filepath"
"testing"

"github.com/stretchr/testify/require"
Expand Down Expand Up @@ -44,18 +45,27 @@ const testDialect = `<?xml version="1.0"?>
`

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)
}
150 changes: 105 additions & 45 deletions pkg/conversion/conversion.go
Original file line number Diff line number Diff line change
Expand Up @@ -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))
Expand Down Expand Up @@ -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) {
Expand All @@ -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
}
Expand All @@ -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
}
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -645,19 +688,36 @@ 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)
}

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
}
Expand Down
Loading