From 47a1b9de22d9372b93bd173edb36876860d171f8 Mon Sep 17 00:00:00 2001 From: lqs Date: Wed, 1 Jul 2020 19:54:03 +0800 Subject: [PATCH] support for multiple dialects --- database.go | 4 ++-- database_test.go | 1 + dialect.go | 30 ++++++++++++++++++++++++++++++ expression.go | 14 ++++++++++++++ field.go | 20 +++++++++++++++----- utils_test.go | 6 ++++-- 6 files changed, 66 insertions(+), 9 deletions(-) create mode 100644 dialect.go diff --git a/database.go b/database.go index ddbab55..8f7e06c 100644 --- a/database.go +++ b/database.go @@ -36,7 +36,7 @@ type database struct { db *sql.DB tx *sql.Tx logger func(sql string, durationNano int64) - dialect string + dialect dialect retryPolicy func(error) bool enableCallerInfo bool interceptor InterceptorFunc @@ -67,7 +67,7 @@ func Open(driverName string, dataSourceName string) (db Database, err error) { } } db = &database{ - dialect: driverName, + dialect: getDialectFromDriverName(driverName), db: sqlDB, } return diff --git a/database_test.go b/database_test.go index ccf7da6..15d3fd0 100644 --- a/database_test.go +++ b/database_test.go @@ -30,6 +30,7 @@ func newMockDatabase() Database { if err != nil { panic(err) } + db.(*database).dialect = dialectMySQL return db } diff --git a/dialect.go b/dialect.go new file mode 100644 index 0000000..66ac3c7 --- /dev/null +++ b/dialect.go @@ -0,0 +1,30 @@ +package sqlingo + +type dialect int + +const ( + dialectUnknown dialect = iota + dialectMySQL + dialectSqlite3 + dialectPostgres + dialectMSSQL + + dialectCount +) + +type dialectArray [dialectCount]string + +func getDialectFromDriverName(driverName string) dialect { + switch driverName { + case "mysql": + return dialectMySQL + case "sqlite3": + return dialectSqlite3 + case "postgres": + return dialectPostgres + case "sqlserver", "mssql": + return dialectMSSQL + default: + return dialectUnknown + } +} diff --git a/expression.go b/expression.go index 66804be..7003e9c 100644 --- a/expression.go +++ b/expression.go @@ -174,6 +174,20 @@ func (e expression) GetSQL(scope scope) (string, error) { return e.builder(scope) } +func quoteIdentifier(identifier string) (result dialectArray) { + for dialect := dialect(0); dialect < dialectCount; dialect++ { + switch dialect { + case dialectMySQL: + result[dialect] = "`" + identifier + "`" + case dialectMSSQL: + result[dialect] = "[" + identifier + "]" + default: + result[dialect] = "\"" + identifier + "\"" + } + } + return +} + func quoteString(s string) string { bytes := []byte(s) buf := make([]byte, len(s)*2+2) diff --git a/field.go b/field.go index 2018da7..7e75132 100644 --- a/field.go +++ b/field.go @@ -19,14 +19,24 @@ type StringField interface { } func newFieldExpression(tableName string, fieldName string) expression { - shortFieldNameSql := getSQLForName(fieldName) - fullFieldNameSql := getSQLForName(tableName) + "." + shortFieldNameSql + tableNameSqlArray := quoteIdentifier(tableName) + fieldNameSqlArray := quoteIdentifier(fieldName) + + var fullFieldNameSqlArray dialectArray + for dialect := dialect(0); dialect < dialectCount; dialect++ { + fullFieldNameSqlArray[dialect] = tableNameSqlArray[dialect] + "." + fieldNameSqlArray[dialect] + } + return expression{ builder: func(scope scope) (string, error) { - if len(scope.Tables) != 1 || scope.lastJoin != nil || scope.Tables[0].GetName() != tableName { - return fullFieldNameSql, nil + dialect := dialectUnknown + if scope.Database != nil { + dialect = scope.Database.dialect } - return shortFieldNameSql, nil + if len(scope.Tables) != 1 || scope.lastJoin != nil || scope.Tables[0].GetName() != tableName { + return fullFieldNameSqlArray[dialect], nil + } + return fieldNameSqlArray[dialect], nil }, } } diff --git a/utils_test.go b/utils_test.go index 952c51b..d21dbfa 100644 --- a/utils_test.go +++ b/utils_test.go @@ -2,9 +2,11 @@ package sqlingo import "testing" +var dummyMySQLScope = scope{Database: &database{dialect: dialectMySQL}} + func assertValue(t *testing.T, value interface{}, expectedSql string) { t.Helper() - if generatedSql, _, _ := getSQL(scope{}, value); generatedSql != expectedSql { + if generatedSql, _, _ := getSQL(dummyMySQLScope, value); generatedSql != expectedSql { t.Errorf("value [%v] generated [%s] expected [%s]", value, generatedSql, expectedSql) } } @@ -18,7 +20,7 @@ func assertLastSql(t *testing.T, expectedSql string) { func assertError(t *testing.T, value interface{}) { t.Helper() - if generatedSql, _, err := getSQL(scope{}, value); err == nil { + if generatedSql, _, err := getSQL(dummyMySQLScope, value); err == nil { t.Errorf("value [%v] generated [%s] expected error", value, generatedSql) } }