SQLTemplate: Make Ident only work for identifiers (not any string) (#92387)
This commit is contained in:
@@ -1,6 +1,7 @@
|
||||
package sqltemplate
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"strconv"
|
||||
"strings"
|
||||
@@ -9,6 +10,7 @@ import (
|
||||
// Dialect-agnostic errors.
|
||||
var (
|
||||
ErrEmptyIdent = errors.New("empty identifier")
|
||||
ErrInvalidIdentInput = errors.New("identifier contains invalid characters")
|
||||
ErrInvalidRowLockingClause = errors.New("invalid row-locking clause")
|
||||
)
|
||||
|
||||
@@ -40,7 +42,8 @@ type Dialect interface {
|
||||
|
||||
// Ident returns the given string quoted in a way that is suitable to be
|
||||
// used as an identifier. Database names, schema names, table names, column
|
||||
// names are all examples of identifiers.
|
||||
// names are all examples of identifiers. When the value includes a "."
|
||||
// each part side of the separator will be escaped: (eg: `db`.`table`)
|
||||
Ident(string) (string, error)
|
||||
|
||||
// ArgPlaceholder returns a safe argument suitable to be used in a SQL
|
||||
@@ -126,11 +129,34 @@ var rowLockingClauseAll = rowLockingClauseMap{
|
||||
// standardIdent provides standard SQL escaping of identifiers.
|
||||
type standardIdent struct{}
|
||||
|
||||
func (standardIdent) Ident(s string) (string, error) {
|
||||
func escapeIdentity(s string, quote rune, clean func(string) string) (string, error) {
|
||||
if s == "" {
|
||||
return "", ErrEmptyIdent
|
||||
}
|
||||
return `"` + strings.ReplaceAll(s, `"`, `""`) + `"`, nil
|
||||
var buffer bytes.Buffer
|
||||
for i, part := range strings.Split(s, ".") {
|
||||
// We may want to check that the identifier is simple alphanumeric
|
||||
// var alphanumeric = regexp.MustCompile("^[a-zA-Z0-9_]*$")
|
||||
|
||||
if i > 1 {
|
||||
return "", ErrInvalidIdentInput
|
||||
}
|
||||
if i > 0 {
|
||||
_, _ = buffer.WriteRune('.')
|
||||
}
|
||||
_, _ = buffer.WriteRune(quote)
|
||||
_, _ = buffer.WriteString(clean(part))
|
||||
_, _ = buffer.WriteRune(quote)
|
||||
}
|
||||
return buffer.String(), nil
|
||||
}
|
||||
|
||||
func (standardIdent) Ident(s string) (string, error) {
|
||||
return escapeIdentity(s, '"', func(s string) string {
|
||||
// not sure we should support escaping quotes in table/column names,
|
||||
// but it is valid so we will support it for now
|
||||
return strings.ReplaceAll(s, `"`, `""`)
|
||||
})
|
||||
}
|
||||
|
||||
type argPlaceholderFunc func(int) string
|
||||
|
||||
@@ -1,10 +1,11 @@
|
||||
package sqltemplate
|
||||
|
||||
// MySQL is the default implementation of Dialect for the MySQL DMBS, currently
|
||||
// supporting MySQL-8.x. It relies on having ANSI_QUOTES SQL Mode enabled. For
|
||||
// more information about ANSI_QUOTES and SQL Modes see:
|
||||
//
|
||||
// https://dev.mysql.com/doc/refman/8.4/en/sql-mode.html#sqlmode_ansi_quotes
|
||||
import (
|
||||
"strings"
|
||||
)
|
||||
|
||||
// MySQL is the default implementation of Dialect for the MySQL DMBS,
|
||||
// currently supporting MySQL-8.x.
|
||||
var MySQL = mysql{
|
||||
rowLockingClauseMap: rowLockingClauseAll,
|
||||
argPlaceholderFunc: argFmtSQL92,
|
||||
@@ -20,19 +21,15 @@ type mysql struct {
|
||||
name
|
||||
}
|
||||
|
||||
// standardIdent provides standard SQL escaping of identifiers.
|
||||
// MySQL always supports backticks for identifiers
|
||||
// https://dev.mysql.com/doc/refman/8.4/en/identifiers.html
|
||||
type backtickIdent struct{}
|
||||
|
||||
var standardFallback = standardIdent{}
|
||||
|
||||
func (backtickIdent) Ident(s string) (string, error) {
|
||||
switch s {
|
||||
// Internal identifiers require backticks to work properly
|
||||
case "user":
|
||||
return "`" + s + "`", nil
|
||||
case "":
|
||||
return "", ErrEmptyIdent
|
||||
if strings.ContainsRune(s, '`') {
|
||||
return "", ErrInvalidIdentInput
|
||||
}
|
||||
// standard
|
||||
return standardFallback.Ident(s)
|
||||
return escapeIdentity(s, '`', func(s string) string {
|
||||
return s
|
||||
})
|
||||
}
|
||||
|
||||
@@ -152,5 +152,5 @@ func Example() {
|
||||
fmt.Println(query)
|
||||
|
||||
// Output:
|
||||
// SELECT "id", "type", "name" FROM "users" WHERE "id" = ?;
|
||||
// SELECT `id`, `type`, `name` FROM `users` WHERE `id` = ?;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user