diff --git a/internal/metadata/meta.go b/internal/metadata/meta.go index d3b00c507b..f51db0c763 100644 --- a/internal/metadata/meta.go +++ b/internal/metadata/meta.go @@ -46,29 +46,36 @@ func Parse(t string, commentStyle CommentSyntax) (string, string, error) { if !commentStyle.Dash { continue } - prefix = "-- name:" + prefix = "--" } if strings.HasPrefix(line, "/*") { if !commentStyle.SlashStar { continue } - prefix = "/* name:" + prefix = "/*" } if strings.HasPrefix(line, "#") { if !commentStyle.Hash { continue } - prefix = "# name:" + prefix = "#" } if prefix == "" { continue } - if !strings.HasPrefix(line, prefix) { + rest := line[len(prefix):] + if !strings.HasPrefix(strings.TrimSpace(rest), "name") { continue } + if !strings.Contains(rest, ":") { + continue + } + if !strings.HasPrefix(rest, " name: ") { + return "", "", fmt.Errorf("invalid metadata: %s", line) + } part := strings.Split(strings.TrimSpace(line), " ") - if strings.HasPrefix(line, "/*") { + if prefix == "/*" { part = part[:len(part)-1] // removes the trailing "*/" element } if len(part) == 2 { diff --git a/internal/metadata/meta_test.go b/internal/metadata/meta_test.go index 37e99307c6..e2f7905cba 100644 --- a/internal/metadata/meta_test.go +++ b/internal/metadata/meta_test.go @@ -3,6 +3,7 @@ package metadata import "testing" func TestParseMetadata(t *testing.T) { + for _, query := range []string{ `-- name: CreateFoo, :one`, `-- name: 9Foo_, :one`, @@ -10,9 +11,37 @@ func TestParseMetadata(t *testing.T) { `-- name: CreateFoo`, `-- name: CreateFoo :one something`, `-- name: `, + `--name: CreateFoo :one`, + `--name CreateFoo :one`, + `--name: CreateFoo :two`, + "-- name:CreateFoo", + `--name:CreateFoo :two`, } { if _, _, err := Parse(query, CommentSyntax{Dash: true}); err == nil { t.Errorf("expected invalid metadata: %q", query) } } + + for _, query := range []string{ + `-- some comment`, + `-- name comment`, + `--name comment`, + } { + if _, _, err := Parse(query, CommentSyntax{Dash: true}); err != nil { + t.Errorf("expected valid comment: %q", query) + } + } + + query := `-- name: CreateFoo :one` + queryName, queryType, err := Parse(query, CommentSyntax{Dash: true}) + if err != nil { + t.Errorf("expected valid metadata: %q", query) + } + if queryName != "CreateFoo" { + t.Errorf("incorrect queryName parsed: %q", query) + } + if queryType != CmdOne { + t.Errorf("incorrect queryType parsed: %q", query) + } + }