diff --git a/.gitea/workflows/test-sqlite.yml b/.gitea/workflows/test-sqlite.yml index 5b64f02..a0af711 100644 --- a/.gitea/workflows/test-sqlite.yml +++ b/.gitea/workflows/test-sqlite.yml @@ -19,6 +19,8 @@ jobs: run: make vet - name: format check run: make fmt-check + - name: lint + run: make lint - name: unit test run: make test - name: test sqlite3 diff --git a/.golangci.yml b/.golangci.yml index 7b91f22..8c282cb 100644 --- a/.golangci.yml +++ b/.golangci.yml @@ -1,24 +1,198 @@ +version: "2" +output: + sort-order: + - file linters: + default: none enable: - - gosimple - - deadcode - - typecheck - - govet - - errcheck - - staticcheck - - unused - - structcheck - - varcheck - - dupl - #- gocyclo # The cyclomatic complexety of a lot of functions is too high, we should refactor those another time. - - gofmt - - misspell - - gocritic - bidichk + - bodyclose + - depguard + - dupl + - errcheck + - forbidigo + - gocheckcompilerdirectives + - gocritic + - govet - ineffassign - enable-all: false - disable-all: true - fast: false + - mirror + - modernize + - nakedret + - nolintlint + - perfsprint + - revive + - staticcheck + - unconvert + - unparam + - unused + - usestdlibvars + - usetesting + - wastedassign + settings: + depguard: + rules: + main: + deny: + - pkg: github.com/unknwon/com + desc: use gitea's util and replacements + - pkg: io/ioutil + desc: use os or io instead + - pkg: golang.org/x/exp + desc: it's experimental and unreliable + - pkg: gopkg.in/ini.v1 + desc: do not use the ini package, use gitea's config system instead + - pkg: github.com/pkg/errors + desc: use builtin errors package instead + - pkg: github.com/go-ap/errors + desc: use builtin errors package instead + nolintlint: + allow-unused: false + require-explanation: true + require-specific: true + gocritic: + enabled-checks: + - equalFold + disabled-checks: + - ifElseChain + - singleCaseSwitch # Every time this occurred in the code, there was no other way. + - deprecatedComment # conflicts with go-swagger comments + revive: + severity: error + rules: + - name: atomic + - name: bare-return + - name: blank-imports + - name: constant-logical-expr + - name: context-as-argument + disabled: true + - name: context-keys-type + - name: dot-imports + - name: duplicated-imports + - name: empty-lines + - name: error-naming + - name: error-return + - name: error-strings + - name: errorf + - name: exported + disabled: true + - name: identical-branches + - name: if-return + - name: increment-decrement + - name: indent-error-flow + - name: modifies-value-receiver + - name: package-comments + - name: range + - name: receiver-naming + - name: redefines-builtin-id + - name: string-of-int + - name: superfluous-else + - name: time-naming + - name: unconditional-recursion + - name: unexported-return + disabled: true + - name: unreachable-code + - name: var-declaration + - name: var-naming + disabled: true + arguments: + - [] # AllowList - do not remove as args for the rule are positional and won't work without lists first + - [] # DenyList + - - skip-package-name-checks: true # supress errors from underscore in migration packages + staticcheck: + checks: + - all + - -ST1003 + - -ST1005 + - -QF1001 + - -QF1006 + - -QF1008 + testifylint: + disable: + - go-require + - require-error + usetesting: + os-temp-dir: true + perfsprint: + concat-loop: false + govet: + enable: + - nilness + - unusedwrite + exclusions: + generated: lax + presets: + - comments + - common-false-positives + - legacy + - std-error-handling + rules: + - linters: + - dupl + - errcheck + - gocyclo + - gosec + - staticcheck + - unparam + path: _test\.go + - linters: + - dupl + - errcheck + - gocyclo + - gosec + path: models/migrations/v + - linters: + - forbidigo + path: cmd + - linters: + - dupl + text: (?i)webhook + - linters: + - dupl + path: dialects/(dameng|gbase8s|oracle)\.go + - linters: + - gocritic + text: (?i)`ID' should not be capitalized + - linters: + - deadcode + - unused + text: (?i)swagger + - linters: + - staticcheck + text: (?i)argument x is overwritten before first use + - linters: + - gocritic + text: '(?i)commentFormatting: put a space between `//` and comment text' + - linters: + - gocritic + text: '(?i)exitAfterDefer:' + paths: + - node_modules + - .venv + - public + - web_src + - third_party$ + - builtin$ + - examples$ +issues: + max-issues-per-linter: 0 + max-same-issues: 0 +formatters: + enable: + - gofmt + - gofumpt + settings: + gofumpt: + extra-rules: true + exclusions: + generated: lax + paths: + - node_modules + - .venv + - public + - web_src + - third_party$ + - builtin$ + - examples$ run: - timeout: 3m \ No newline at end of file + timeout: 10m diff --git a/AGENTS.md b/AGENTS.md new file mode 100644 index 0000000..8b74dc2 --- /dev/null +++ b/AGENTS.md @@ -0,0 +1,8 @@ +# Instructions for agents + +- Use `make help` to find available development targets +- Before committing `.go` changes, run `make fmt` to format, and run `make lint` to lint +- Before committing `go.mod` changes, run `go mod tidy` +- Before committing new `.go` files, add the current year into the copyright header +- Before committing any files, remove all trailing whitespace from source code lines +- Add test code for newly added functions or new logic of old functions \ No newline at end of file diff --git a/Makefile b/Makefile index 5dc67e0..2f1a5c6 100644 --- a/Makefile +++ b/Makefile @@ -6,11 +6,12 @@ GOFMT ?= gofmt -s TAGS ?= SED_INPLACE := sed -i -GO_DIRS := caches contexts integrations core dialects internal log migrate names schemas tags +GO_DIRS := caches contexts convert core dialects internal log migrate names schemas tags tests GOFILES := $(wildcard *.go) GOFILES += $(shell find $(GO_DIRS) -name "*.go" -type f) INTEGRATION_PACKAGES := xorm.io/xorm/tests PACKAGES ?= $(filter-out $(INTEGRATION_PACKAGES),$(shell $(GO) list ./...)) +GOLANGCI_LINT_PACKAGE ?= github.com/golangci/golangci-lint/v2/cmd/golangci-lint@v2.8.0 TEST_COCKROACH_HOST ?= cockroach:26257 TEST_COCKROACH_SCHEMA ?= @@ -72,7 +73,7 @@ coverage: .PHONY: fmt fmt: - $(GOFMT) -w $(GOFILES) + $(GO) fmt ./... .PHONY: fmt-check fmt-check: @@ -111,25 +112,8 @@ help: @echo " - vet examines Go source code and reports suspicious constructs" .PHONY: lint -lint: golangci-lint - -.PHONY: golangci-lint -golangci-lint: golangci-lint-check - golangci-lint run --timeout 10m - -.PHONY: golangci-lint-check -golangci-lint-check: - $(eval GOLANGCI_LINT_VERSION := $(shell printf "%03d%03d%03d" $(shell golangci-lint --version | grep -Eo '[0-9]+\.[0-9.]+' | tr '.' ' ');)) - $(eval MIN_GOLANGCI_LINT_VER_FMT := $(shell printf "%g.%g.%g" $(shell echo $(MIN_GOLANGCI_LINT_VERSION) | grep -o ...))) - @hash golangci-lint > /dev/null 2>&1; if [ $$? -ne 0 ]; then \ - echo "Downloading golangci-lint v${MIN_GOLANGCI_LINT_VER_FMT}"; \ - export BINARY="golangci-lint"; \ - curl -sfL "https://raw.githubusercontent.com/golangci/golangci-lint/v${MIN_GOLANGCI_LINT_VER_FMT}/install.sh" | sh -s -- -b $(GOPATH)/bin v$(MIN_GOLANGCI_LINT_VER_FMT); \ - elif [ "$(GOLANGCI_LINT_VERSION)" -lt "$(MIN_GOLANGCI_LINT_VERSION)" ]; then \ - echo "Downloading newer version of golangci-lint v${MIN_GOLANGCI_LINT_VER_FMT}"; \ - export BINARY="golangci-lint"; \ - curl -sfL "https://raw.githubusercontent.com/golangci/golangci-lint/v${MIN_GOLANGCI_LINT_VER_FMT}/install.sh" | sh -s -- -b $(GOPATH)/bin v$(MIN_GOLANGCI_LINT_VER_FMT); \ - fi +lint: + $(GO) run $(GOLANGCI_LINT_PACKAGE) run .PHONY: test test: go-check diff --git a/caches/cache.go b/caches/cache.go index 7b80eb8..2983778 100644 --- a/caches/cache.go +++ b/caches/cache.go @@ -39,20 +39,20 @@ type CacheStore interface { // key is primary key or composite primary key // value is struct's pointer // key format : -p--... - Put(key string, value interface{}) error - Get(key string) (interface{}, error) + Put(key string, value any) error + Get(key string) (any, error) Del(key string) error } // Cacher is an interface to provide cache // id format : u--... type Cacher interface { - GetIds(tableName, sql string) interface{} - GetBean(tableName string, id string) interface{} - PutIds(tableName, sql string, ids interface{}) - PutBean(tableName string, id string, obj interface{}) + GetIds(tableName, sql string) any + GetBean(tableName, id string) any + PutIds(tableName, sql string, ids any) + PutBean(tableName, id string, obj any) DelIds(tableName, sql string) - DelBean(tableName string, id string) + DelBean(tableName, id string) ClearIds(tableName string) ClearBeans(tableName string) } @@ -75,7 +75,7 @@ func decodeIds(s string) ([]schemas.PK, error) { } // GetCacheSql returns cacher PKs via SQL -func GetCacheSql(m Cacher, tableName, sql string, args interface{}) ([]schemas.PK, error) { +func GetCacheSql(m Cacher, tableName, sql string, args any) ([]schemas.PK, error) { bytes := m.GetIds(tableName, GenSqlKey(sql, args)) if bytes == nil { return nil, errors.New("Not Exist") @@ -84,7 +84,7 @@ func GetCacheSql(m Cacher, tableName, sql string, args interface{}) ([]schemas.P } // PutCacheSql puts cacher SQL and PKs -func PutCacheSql(m Cacher, ids []schemas.PK, tableName, sql string, args interface{}) error { +func PutCacheSql(m Cacher, ids []schemas.PK, tableName, sql string, args any) error { bytes, err := encodeIds(ids) if err != nil { return err @@ -94,6 +94,6 @@ func PutCacheSql(m Cacher, ids []schemas.PK, tableName, sql string, args interfa } // GenSqlKey generates cache key -func GenSqlKey(sql string, args interface{}) string { +func GenSqlKey(sql string, args any) string { return fmt.Sprintf("%v-%v", sql, args) } diff --git a/caches/encode.go b/caches/encode.go index 8659668..332199f 100644 --- a/caches/encode.go +++ b/caches/encode.go @@ -8,8 +8,8 @@ import ( "bytes" "crypto/md5" "encoding/gob" + "encoding/hex" "encoding/json" - "fmt" "io" ) @@ -17,23 +17,23 @@ import ( func Md5(str string) string { m := md5.New() _, _ = io.WriteString(m, str) - return fmt.Sprintf("%x", m.Sum(nil)) + return hex.EncodeToString(m.Sum(nil)) } // Encode Encode data -func Encode(data interface{}) ([]byte, error) { +func Encode(data any) ([]byte, error) { // return JsonEncode(data) return GobEncode(data) } // Decode decode data -func Decode(data []byte, to interface{}) error { +func Decode(data []byte, to any) error { // return JsonDecode(data, to) return GobDecode(data, to) } // GobEncode encode data with gob -func GobEncode(data interface{}) ([]byte, error) { +func GobEncode(data any) ([]byte, error) { var buf bytes.Buffer enc := gob.NewEncoder(&buf) err := enc.Encode(&data) @@ -44,14 +44,14 @@ func GobEncode(data interface{}) ([]byte, error) { } // GobDecode decode data with gob -func GobDecode(data []byte, to interface{}) error { +func GobDecode(data []byte, to any) error { buf := bytes.NewBuffer(data) dec := gob.NewDecoder(buf) return dec.Decode(to) } // JsonEncode encode data with json -func JsonEncode(data interface{}) ([]byte, error) { +func JsonEncode(data any) ([]byte, error) { val, err := json.Marshal(data) if err != nil { return nil, err @@ -60,6 +60,6 @@ func JsonEncode(data interface{}) ([]byte, error) { } // JsonDecode decode data with json -func JsonDecode(data []byte, to interface{}) error { +func JsonDecode(data []byte, to any) error { return json.Unmarshal(data, to) } diff --git a/caches/leveldb.go b/caches/leveldb.go index f2f71d8..04d3cfd 100644 --- a/caches/leveldb.go +++ b/caches/leveldb.go @@ -14,7 +14,7 @@ import ( type LevelDBStore struct { store *leveldb.DB Debug bool - v interface{} + v any } var _ CacheStore = &LevelDBStore{} @@ -31,7 +31,7 @@ func NewLevelDBStore(dbfile string) (*LevelDBStore, error) { } // Put implements CacheStore -func (s *LevelDBStore) Put(key string, value interface{}) error { +func (s *LevelDBStore) Put(key string, value any) error { val, err := Encode(value) if err != nil { if s.Debug { @@ -53,7 +53,7 @@ func (s *LevelDBStore) Put(key string, value interface{}) error { } // Get implements CacheStore -func (s *LevelDBStore) Get(key string) (interface{}, error) { +func (s *LevelDBStore) Get(key string) (any, error) { data, err := s.store.Get([]byte(key), nil) if err != nil { if s.Debug { diff --git a/caches/leveldb_test.go b/caches/leveldb_test.go index 35981db..b85ea8b 100644 --- a/caches/leveldb_test.go +++ b/caches/leveldb_test.go @@ -14,7 +14,7 @@ func TestLevelDBStore(t *testing.T) { store, err := NewLevelDBStore("./level.db") assert.NoError(t, err) - var kvs = map[string]interface{}{ + kvs := map[string]any{ "a": "b", } for k, v := range kvs { diff --git a/caches/lru.go b/caches/lru.go index 885f02d..9e0f565 100644 --- a/caches/lru.go +++ b/caches/lru.go @@ -31,7 +31,8 @@ func NewLRUCacher(store CacheStore, maxElementSize int) *LRUCacher { // NewLRUCacher2 creates a cache include different params func NewLRUCacher2(store CacheStore, expired time.Duration, maxElementSize int) *LRUCacher { - cacher := &LRUCacher{store: store, idList: list.New(), + cacher := &LRUCacher{ + store: store, idList: list.New(), sqlList: list.New(), Expired: expired, GcInterval: CacheGcInterval, MaxElementSize: maxElementSize, sqlIndex: make(map[string]map[string]*list.Element), @@ -83,7 +84,7 @@ func (m *LRUCacher) GC() { } // GetIds returns all bean's ids according to sql and parameter from cache -func (m *LRUCacher) GetIds(tableName, sql string) interface{} { +func (m *LRUCacher) GetIds(tableName, sql string) any { m.mutex.Lock() defer m.mutex.Unlock() if _, ok := m.sqlIndex[tableName]; !ok { @@ -111,7 +112,7 @@ func (m *LRUCacher) GetIds(tableName, sql string) interface{} { } // GetBean returns bean according tableName and id from cache -func (m *LRUCacher) GetBean(tableName string, id string) interface{} { +func (m *LRUCacher) GetBean(tableName, id string) any { m.mutex.Lock() defer m.mutex.Unlock() if _, ok := m.idIndex[tableName]; !ok { @@ -177,7 +178,7 @@ func (m *LRUCacher) ClearBeans(tableName string) { } // PutIds pus ids into table -func (m *LRUCacher) PutIds(tableName, sql string, ids interface{}) { +func (m *LRUCacher) PutIds(tableName, sql string, ids any) { m.mutex.Lock() if _, ok := m.sqlIndex[tableName]; !ok { m.sqlIndex[tableName] = make(map[string]*list.Element) @@ -198,7 +199,7 @@ func (m *LRUCacher) PutIds(tableName, sql string, ids interface{}) { } // PutBean puts beans into table -func (m *LRUCacher) PutBean(tableName string, id string, obj interface{}) { +func (m *LRUCacher) PutBean(tableName, id string, obj any) { m.mutex.Lock() var el *list.Element var ok bool @@ -236,7 +237,7 @@ func (m *LRUCacher) DelIds(tableName, sql string) { m.mutex.Unlock() } -func (m *LRUCacher) delBean(tableName string, id string) { +func (m *LRUCacher) delBean(tableName, id string) { tid := genID(tableName, id) if el, ok := m.idIndex[tableName][id]; ok { delete(m.idIndex[tableName], id) @@ -247,7 +248,7 @@ func (m *LRUCacher) delBean(tableName string, id string) { } // DelBean deletes beans in some table -func (m *LRUCacher) DelBean(tableName string, id string) { +func (m *LRUCacher) DelBean(tableName, id string) { m.mutex.Lock() m.delBean(tableName, id) m.mutex.Unlock() @@ -265,11 +266,11 @@ type sqlNode struct { lastVisit time.Time } -func genID(prefix string, id string) string { +func genID(prefix, id string) string { return fmt.Sprintf("%s-%s", prefix, id) } -func newIDNode(tbName string, id string) *idNode { +func newIDNode(tbName, id string) *idNode { return &idNode{tbName, id, time.Now()} } diff --git a/caches/lru_test.go b/caches/lru_test.go index 771b924..0d9aa86 100644 --- a/caches/lru_test.go +++ b/caches/lru_test.go @@ -40,7 +40,7 @@ func TestLRUCache(t *testing.T) { obj2 := cacher.GetBean(tableName, sid) assert.Nil(t, obj2) - var obj = new(CacheObject1) + obj := new(CacheObject1) cacher.PutBean(tableName, sid, obj) obj3 := cacher.GetBean(tableName, sid) assert.EqualValues(t, obj, obj3) diff --git a/caches/memory_store.go b/caches/memory_store.go index f16254d..bf2fc74 100644 --- a/caches/memory_store.go +++ b/caches/memory_store.go @@ -12,17 +12,17 @@ var _ CacheStore = NewMemoryStore() // MemoryStore represents in-memory store type MemoryStore struct { - store map[interface{}]interface{} + store map[any]any mutex sync.RWMutex } // NewMemoryStore creates a new store in memory func NewMemoryStore() *MemoryStore { - return &MemoryStore{store: make(map[interface{}]interface{})} + return &MemoryStore{store: make(map[any]any)} } // Put puts object into store -func (s *MemoryStore) Put(key string, value interface{}) error { +func (s *MemoryStore) Put(key string, value any) error { s.mutex.Lock() defer s.mutex.Unlock() s.store[key] = value @@ -30,7 +30,7 @@ func (s *MemoryStore) Put(key string, value interface{}) error { } // Get gets object from store -func (s *MemoryStore) Get(key string) (interface{}, error) { +func (s *MemoryStore) Get(key string) (any, error) { s.mutex.RLock() defer s.mutex.RUnlock() if v, ok := s.store[key]; ok { diff --git a/caches/memory_store_test.go b/caches/memory_store_test.go index 12db4ea..b439ec8 100644 --- a/caches/memory_store_test.go +++ b/caches/memory_store_test.go @@ -12,7 +12,7 @@ import ( func TestMemoryStore(t *testing.T) { store := NewMemoryStore() - var kvs = map[string]interface{}{ + kvs := map[string]any{ "a": "b", } for k, v := range kvs { diff --git a/contexts/context_cache.go b/contexts/context_cache.go index 0d0f0f0..f7d24df 100644 --- a/contexts/context_cache.go +++ b/contexts/context_cache.go @@ -7,24 +7,24 @@ package contexts // ContextCache is the interface that operates the cache data. type ContextCache interface { // Put puts value into cache with key. - Put(key string, val interface{}) + Put(key string, val any) // Get gets cached value by given key. - Get(key string) interface{} + Get(key string) any } -type memoryContextCache map[string]interface{} +type memoryContextCache map[string]any // NewMemoryContextCache return memoryContextCache func NewMemoryContextCache() memoryContextCache { - return make(map[string]interface{}) + return make(map[string]any) } // Put puts value into cache with key. -func (m memoryContextCache) Put(key string, val interface{}) { +func (m memoryContextCache) Put(key string, val any) { m[key] = val } // Get gets cached value by given key. -func (m memoryContextCache) Get(key string) interface{} { +func (m memoryContextCache) Get(key string) any { return m[key] } diff --git a/contexts/hook.go b/contexts/hook.go index f6d86cf..a6d2225 100644 --- a/contexts/hook.go +++ b/contexts/hook.go @@ -14,15 +14,15 @@ import ( type ContextHook struct { start time.Time Ctx context.Context - SQL string // log content or SQL - Args []interface{} // if it's a SQL, it's the arguments + SQL string // log content or SQL + Args []any // if it's a SQL, it's the arguments Result sql.Result ExecuteTime time.Duration Err error // SQL executed error } // NewContextHook return context for hook -func NewContextHook(ctx context.Context, sql string, args []interface{}) *ContextHook { +func NewContextHook(ctx context.Context, sql string, args []any) *ContextHook { return &ContextHook{ start: time.Now(), Ctx: ctx, diff --git a/convert/bool.go b/convert/bool.go index 58b23f4..e8ed163 100644 --- a/convert/bool.go +++ b/convert/bool.go @@ -11,7 +11,7 @@ import ( ) // AsBool convert interface as bool -func AsBool(src interface{}) (bool, error) { +func AsBool(src any) (bool, error) { switch v := src.(type) { case bool: return v, nil @@ -33,12 +33,14 @@ func AsBool(src interface{}) (bool, error) { if len(v) == 0 { return false, nil } - if v[0] == 0x00 { + switch v[0] { + case 0x00: return false, nil - } else if v[0] == 0x01 { + case 0x01: return true, nil + default: + return strconv.ParseBool(string(v)) } - return strconv.ParseBool(string(v)) case string: return strconv.ParseBool(v) case *sql.NullInt64: diff --git a/convert/conversion.go b/convert/conversion.go index ffe5fd9..a6fcfea 100644 --- a/convert/conversion.go +++ b/convert/conversion.go @@ -56,10 +56,10 @@ func cloneBytes(b []byte) []byte { // Assign copies to dest the value in src, converting it if possible. // An error is returned if the copy would result in loss of information. // dest should be a pointer type. -func Assign(dest, src interface{}, originalLocation *time.Location, convertedLocation *time.Location) error { +func Assign(dest, src any, originalLocation, convertedLocation *time.Location) error { // Common cases, without reflect. switch s := src.(type) { - case *interface{}: + case *any: return Assign(dest, *s, originalLocation, convertedLocation) case string: switch d := dest.(type) { @@ -84,7 +84,7 @@ func Assign(dest, src interface{}, originalLocation *time.Location, convertedLoc } *d = string(s) return nil - case *interface{}: + case *any: if d == nil { return ErrNilPtr } @@ -111,7 +111,7 @@ func Assign(dest, src interface{}, originalLocation *time.Location, convertedLoc } case nil: switch d := dest.(type) { - case *interface{}: + case *any: if d == nil { return ErrNilPtr } @@ -296,7 +296,7 @@ func Assign(dest, src interface{}, originalLocation *time.Location, convertedLoc switch d := dest.(type) { case *string: - var sv = reflect.ValueOf(src) + sv := reflect.ValueOf(src) switch sv.Kind() { case reflect.Bool, reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64, @@ -316,7 +316,7 @@ func Assign(dest, src interface{}, originalLocation *time.Location, convertedLoc *d = bv.(bool) } return err - case *interface{}: + case *any: *d = src return nil } @@ -330,11 +330,11 @@ var ( ) // AssignValue assign src as dv -func AssignValue(dv reflect.Value, src interface{}) error { +func AssignValue(dv reflect.Value, src any) error { if src == nil { return nil } - if v, ok := src.(*interface{}); ok { + if v, ok := src.(*any); ok { return AssignValue(dv, *v) } diff --git a/convert/float.go b/convert/float.go index 51b441c..e8df38a 100644 --- a/convert/float.go +++ b/convert/float.go @@ -13,7 +13,7 @@ import ( ) // AsFloat64 convets interface as float64 -func AsFloat64(src interface{}) (float64, error) { +func AsFloat64(src any) (float64, error) { switch v := src.(type) { case int: return float64(v), nil @@ -64,7 +64,7 @@ func AsFloat64(src interface{}) (float64, error) { } // AsBigFloat converts interface as big.Float -func AsBigFloat(src interface{}) (*big.Float, error) { +func AsBigFloat(src any) (*big.Float, error) { res := big.NewFloat(0) switch v := src.(type) { case int: @@ -80,7 +80,7 @@ func AsBigFloat(src interface{}) (*big.Float, error) { res.SetInt64(int64(v)) return res, nil case int64: - res.SetInt64(int64(v)) + res.SetInt64(v) return res, nil case uint: res.SetUint64(uint64(v)) @@ -95,7 +95,7 @@ func AsBigFloat(src interface{}) (*big.Float, error) { res.SetUint64(uint64(v)) return res, nil case uint64: - res.SetUint64(uint64(v)) + res.SetUint64(v) return res, nil case []byte: res.SetString(string(v)) @@ -117,7 +117,7 @@ func AsBigFloat(src interface{}) (*big.Float, error) { return nil, nil case *sql.NullInt64: if v.Valid { - res.SetInt64(int64(v.Int64)) + res.SetInt64(v.Int64) return res, nil } return nil, nil diff --git a/convert/int.go b/convert/int.go index 0399477..3689f57 100644 --- a/convert/int.go +++ b/convert/int.go @@ -13,7 +13,7 @@ import ( ) // AsInt64 converts interface as int64 -func AsInt64(src interface{}) (int64, error) { +func AsInt64(src any) (int64, error) { switch v := src.(type) { case int: return int64(v), nil @@ -94,7 +94,7 @@ func AsInt64(src interface{}) (int64, error) { case *sql.NullInt32: return int64(v.Int32), nil case *sql.NullInt64: - return int64(v.Int64), nil + return v.Int64, nil } rv := reflect.ValueOf(src) @@ -112,7 +112,7 @@ func AsInt64(src interface{}) (int64, error) { } // AsUint64 converts interface as uint64 -func AsUint64(src interface{}) (uint64, error) { +func AsUint64(src any) (uint64, error) { switch v := src.(type) { case int: return uint64(v), nil @@ -151,7 +151,7 @@ func AsUint64(src interface{}) (uint64, error) { case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64: return uint64(rv.Int()), nil case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64: - return uint64(rv.Uint()), nil + return rv.Uint(), nil case reflect.Float64, reflect.Float32: return uint64(rv.Float()), nil case reflect.String: @@ -171,7 +171,7 @@ type NullUint64 struct { } // Scan implements the Scanner interface. -func (n *NullUint64) Scan(value interface{}) error { +func (n *NullUint64) Scan(value any) error { if value == nil { n.Uint64, n.Valid = 0, false return nil @@ -201,7 +201,7 @@ type NullUint32 struct { } // Scan implements the Scanner interface. -func (n *NullUint32) Scan(value interface{}) error { +func (n *NullUint32) Scan(value any) error { if value == nil { n.Uint32, n.Valid = 0, false return nil diff --git a/convert/interface.go b/convert/interface.go index 2cc8d9f..58347b8 100644 --- a/convert/interface.go +++ b/convert/interface.go @@ -11,7 +11,7 @@ import ( ) // Interface2Interface converts interface of pointer as interface of value -func Interface2Interface(userLocation *time.Location, v interface{}) (interface{}, error) { +func Interface2Interface(userLocation *time.Location, v any) (any, error) { if v == nil { return nil, nil } diff --git a/convert/scanner.go b/convert/scanner.go index 505d3be..7ea47a8 100644 --- a/convert/scanner.go +++ b/convert/scanner.go @@ -6,14 +6,12 @@ package convert import "database/sql" -var ( - _ sql.Scanner = &EmptyScanner{} -) +var _ sql.Scanner = &EmptyScanner{} // EmptyScanner represents an empty scanner which will ignore the scan type EmptyScanner struct{} // Scan implements sql.Scanner -func (EmptyScanner) Scan(value interface{}) error { +func (EmptyScanner) Scan(value any) error { return nil } diff --git a/convert/string.go b/convert/string.go index de11fa0..dc034fa 100644 --- a/convert/string.go +++ b/convert/string.go @@ -12,7 +12,7 @@ import ( ) // AsString converts interface as string -func AsString(src interface{}) string { +func AsString(src any) string { switch v := src.(type) { case string: return v @@ -21,9 +21,9 @@ func AsString(src interface{}) string { case *sql.NullString: return v.String case *sql.NullInt32: - return fmt.Sprintf("%d", v.Int32) + return strconv.Itoa(int(v.Int32)) case *sql.NullInt64: - return fmt.Sprintf("%d", v.Int64) + return strconv.FormatInt(v.Int64, 10) } rv := reflect.ValueOf(src) switch rv.Kind() { @@ -42,7 +42,7 @@ func AsString(src interface{}) string { } // AsBytes converts interface as bytes -func AsBytes(src interface{}) ([]byte, bool) { +func AsBytes(src any) ([]byte, bool) { switch t := src.(type) { case []byte: return t, true diff --git a/convert/time.go b/convert/time.go index 8447214..5f35bdc 100644 --- a/convert/time.go +++ b/convert/time.go @@ -16,7 +16,7 @@ import ( // String2Time converts a string to time with original location // be aware for time strings (HH:mm:ss) returns zero year (LMT) for converted location -func String2Time(s string, originalLocation *time.Location, convertedLocation *time.Location) (*time.Time, error) { +func String2Time(s string, originalLocation, convertedLocation *time.Location) (*time.Time, error) { if len(s) == 19 { if s == utils.ZeroTime0 || s == utils.ZeroTime1 { return &time.Time{}, nil @@ -85,25 +85,24 @@ func String2Time(s string, originalLocation *time.Location, convertedLocation *t if err != nil { return nil, err } - dt = dt.AddDate(2006, 01, 02).In(convertedLocation) + dt = dt.AddDate(2006, 0o1, 0o2).In(convertedLocation) // back to zero year - dt = dt.AddDate(-2006, -01, -02) + dt = dt.AddDate(-2006, -0o1, -0o2) return &dt, nil - } else { - i, err := strconv.ParseInt(s, 10, 64) - if err == nil { - if i == 0 { - return &time.Time{}, nil - } - tm := time.Unix(i, 0).In(convertedLocation) - return &tm, nil + } + i, err := strconv.ParseInt(s, 10, 64) + if err == nil { + if i == 0 { + return &time.Time{}, nil } + tm := time.Unix(i, 0).In(convertedLocation) + return &tm, nil } return nil, fmt.Errorf("unsupported conversion from %s to time", s) } // AsTime converts interface as time -func AsTime(src interface{}, dbLoc *time.Location, uiLoc *time.Location) (*time.Time, error) { +func AsTime(src any, dbLoc, uiLoc *time.Location) (*time.Time, error) { switch t := src.(type) { case string: return String2Time(t, dbLoc, uiLoc) diff --git a/core/db.go b/core/db.go index b476ef9..8abf3b7 100644 --- a/core/db.go +++ b/core/db.go @@ -18,19 +18,17 @@ import ( "xorm.io/xorm/names" ) -var ( - // DefaultCacheSize sets the default cache size - DefaultCacheSize = 200 -) +// DefaultCacheSize sets the default cache size +var DefaultCacheSize = 200 // MapToSlice map query and struct as sql and args -func MapToSlice(query string, mp interface{}) (string, []interface{}, error) { +func MapToSlice(query string, mp any) (string, []any, error) { vv := reflect.ValueOf(mp) if vv.Kind() != reflect.Ptr || vv.Elem().Kind() != reflect.Map { - return "", []interface{}{}, ErrNoMapPointer + return "", []any{}, ErrNoMapPointer } - args := make([]interface{}, 0, len(vv.Elem().MapKeys())) + args := make([]any, 0, len(vv.Elem().MapKeys())) var err error query = re.ReplaceAllStringFunc(query, func(src string) string { v := vv.Elem().MapIndex(reflect.ValueOf(src[1:])) @@ -46,13 +44,13 @@ func MapToSlice(query string, mp interface{}) (string, []interface{}, error) { } // StructToSlice converts a query and struct as sql and args -func StructToSlice(query string, st interface{}) (string, []interface{}, error) { +func StructToSlice(query string, st any) (string, []any, error) { vv := reflect.ValueOf(st) if vv.Kind() != reflect.Ptr || vv.Elem().Kind() != reflect.Struct { - return "", []interface{}{}, ErrNoStructPointer + return "", []any{}, ErrNoStructPointer } - args := make([]interface{}, 0) + args := make([]any, 0) var err error query = re.ReplaceAllStringFunc(query, func(src string) string { fv := vv.Elem().FieldByName(src[1:]).Interface() @@ -69,7 +67,7 @@ func StructToSlice(query string, st interface{}) (string, []interface{}, error) return "?" }) if err != nil { - return "", []interface{}{}, err + return "", []any{}, err } return query, args, nil } @@ -79,9 +77,7 @@ type cacheStruct struct { idx int } -var ( - _ QueryExecuter = &DB{} -) +var _ QueryExecuter = &DB{} // DB is a wrap of sql.DB with extra contents type DB struct { @@ -142,7 +138,7 @@ func (db *DB) reflectNew(typ reflect.Type) reflect.Value { } // QueryContext overwrites sql.DB.QueryContext -func (db *DB) QueryContext(ctx context.Context, query string, args ...interface{}) (*Rows, error) { +func (db *DB) QueryContext(ctx context.Context, query string, args ...any) (*Rows, error) { hookCtx := contexts.NewContextHook(ctx, query, args) ctx, err := db.beforeProcess(hookCtx) if err != nil { @@ -160,12 +156,12 @@ func (db *DB) QueryContext(ctx context.Context, query string, args ...interface{ } // Query overwrites sql.DB.Query -func (db *DB) Query(query string, args ...interface{}) (*Rows, error) { +func (db *DB) Query(query string, args ...any) (*Rows, error) { return db.QueryContext(context.Background(), query, args...) } // QueryMapContext executes query with parameters via map and context -func (db *DB) QueryMapContext(ctx context.Context, query string, mp interface{}) (*Rows, error) { +func (db *DB) QueryMapContext(ctx context.Context, query string, mp any) (*Rows, error) { query, args, err := MapToSlice(query, mp) if err != nil { return nil, err @@ -174,12 +170,12 @@ func (db *DB) QueryMapContext(ctx context.Context, query string, mp interface{}) } // QueryMap executes query with parameters via map -func (db *DB) QueryMap(query string, mp interface{}) (*Rows, error) { +func (db *DB) QueryMap(query string, mp any) (*Rows, error) { return db.QueryMapContext(context.Background(), query, mp) } // QueryStructContext query rows with struct -func (db *DB) QueryStructContext(ctx context.Context, query string, st interface{}) (*Rows, error) { +func (db *DB) QueryStructContext(ctx context.Context, query string, st any) (*Rows, error) { query, args, err := StructToSlice(query, st) if err != nil { return nil, err @@ -188,12 +184,12 @@ func (db *DB) QueryStructContext(ctx context.Context, query string, st interface } // QueryStruct query rows with struct -func (db *DB) QueryStruct(query string, st interface{}) (*Rows, error) { +func (db *DB) QueryStruct(query string, st any) (*Rows, error) { return db.QueryStructContext(context.Background(), query, st) } // QueryRowContext query row with args -func (db *DB) QueryRowContext(ctx context.Context, query string, args ...interface{}) *Row { +func (db *DB) QueryRowContext(ctx context.Context, query string, args ...any) *Row { rows, err := db.QueryContext(ctx, query, args...) if err != nil { return &Row{nil, err} @@ -202,12 +198,12 @@ func (db *DB) QueryRowContext(ctx context.Context, query string, args ...interfa } // QueryRow query row with args -func (db *DB) QueryRow(query string, args ...interface{}) *Row { +func (db *DB) QueryRow(query string, args ...any) *Row { return db.QueryRowContext(context.Background(), query, args...) } // QueryRowMapContext query row with map -func (db *DB) QueryRowMapContext(ctx context.Context, query string, mp interface{}) *Row { +func (db *DB) QueryRowMapContext(ctx context.Context, query string, mp any) *Row { query, args, err := MapToSlice(query, mp) if err != nil { return &Row{nil, err} @@ -216,12 +212,12 @@ func (db *DB) QueryRowMapContext(ctx context.Context, query string, mp interface } // QueryRowMap query row with map -func (db *DB) QueryRowMap(query string, mp interface{}) *Row { +func (db *DB) QueryRowMap(query string, mp any) *Row { return db.QueryRowMapContext(context.Background(), query, mp) } // QueryRowStructContext query row with struct -func (db *DB) QueryRowStructContext(ctx context.Context, query string, st interface{}) *Row { +func (db *DB) QueryRowStructContext(ctx context.Context, query string, st any) *Row { query, args, err := StructToSlice(query, st) if err != nil { return &Row{nil, err} @@ -230,18 +226,16 @@ func (db *DB) QueryRowStructContext(ctx context.Context, query string, st interf } // QueryRowStruct query row with struct -func (db *DB) QueryRowStruct(query string, st interface{}) *Row { +func (db *DB) QueryRowStruct(query string, st any) *Row { return db.QueryRowStructContext(context.Background(), query, st) } -var ( - re = regexp.MustCompile(`[?](\w+)`) -) +var re = regexp.MustCompile(`[?](\w+)`) // ExecMapContext exec map with context.ContextHook // insert into (name) values (?) // insert into (name) values (?name) -func (db *DB) ExecMapContext(ctx context.Context, query string, mp interface{}) (sql.Result, error) { +func (db *DB) ExecMapContext(ctx context.Context, query string, mp any) (sql.Result, error) { query, args, err := MapToSlice(query, mp) if err != nil { return nil, err @@ -250,12 +244,12 @@ func (db *DB) ExecMapContext(ctx context.Context, query string, mp interface{}) } // ExecMap exec query with map -func (db *DB) ExecMap(query string, mp interface{}) (sql.Result, error) { +func (db *DB) ExecMap(query string, mp any) (sql.Result, error) { return db.ExecMapContext(context.Background(), query, mp) } // ExecStructContext exec query with map -func (db *DB) ExecStructContext(ctx context.Context, query string, st interface{}) (sql.Result, error) { +func (db *DB) ExecStructContext(ctx context.Context, query string, st any) (sql.Result, error) { query, args, err := StructToSlice(query, st) if err != nil { return nil, err @@ -264,7 +258,7 @@ func (db *DB) ExecStructContext(ctx context.Context, query string, st interface{ } // ExecContext exec query with args -func (db *DB) ExecContext(ctx context.Context, query string, args ...interface{}) (sql.Result, error) { +func (db *DB) ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error) { hookCtx := contexts.NewContextHook(ctx, query, args) ctx, err := db.beforeProcess(hookCtx) if err != nil { @@ -279,7 +273,7 @@ func (db *DB) ExecContext(ctx context.Context, query string, args ...interface{} } // ExecStruct exec query with struct -func (db *DB) ExecStruct(query string, st interface{}) (sql.Result, error) { +func (db *DB) ExecStruct(query string, st any) (sql.Result, error) { return db.ExecStructContext(context.Background(), query, st) } diff --git a/core/db_test.go b/core/db_test.go index a9c1939..9369186 100644 --- a/core/db_test.go +++ b/core/db_test.go @@ -239,7 +239,7 @@ func BenchmarkSliceInterfaceQuery(b *testing.B) { } for rows.Next() { - slice := make([]interface{}, len(cols)) + slice := make([]any, len(cols)) err = rows.ScanSlice(&slice) if err != nil { b.Error(err) @@ -394,7 +394,7 @@ func BenchmarkMapInterfaceQuery(b *testing.B) { } for rows.Next() { - m := make(map[string]interface{}) + m := make(map[string]any) err = rows.ScanMap(&m) if err != nil { b.Error(err) @@ -551,7 +551,7 @@ func BenchmarkExecMap(b *testing.B) { b.StartTimer() - mp := map[string]interface{}{ + mp := map[string]any{ "name": "xlw", "title": "tester", "age": 1.2, @@ -582,7 +582,7 @@ func TestExecMap(t *testing.T) { t.Error(err) } - mp := map[string]interface{}{ + mp := map[string]any{ "name": "xlw", "title": "tester", "age": 1.2, @@ -625,7 +625,8 @@ func TestExecStruct(t *testing.T) { t.Error(err) } - user := User{Name: "xlw", + user := User{ + Name: "xlw", Title: "tester", Age: 1.2, Alias: "lunny", @@ -670,7 +671,8 @@ func BenchmarkExecStruct(b *testing.B) { b.StartTimer() - user := User{Name: "xlw", + user := User{ + Name: "xlw", Title: "tester", Age: 1.2, Alias: "lunny", diff --git a/core/interface.go b/core/interface.go index a5c8e4e..d5ebb89 100644 --- a/core/interface.go +++ b/core/interface.go @@ -7,12 +7,12 @@ import ( // Queryer represents an interface to query a SQL to get data from database type Queryer interface { - QueryContext(ctx context.Context, query string, args ...interface{}) (*Rows, error) + QueryContext(ctx context.Context, query string, args ...any) (*Rows, error) } // Executer represents an interface to execute a SQL type Executer interface { - ExecContext(ctx context.Context, query string, args ...interface{}) (sql.Result, error) + ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error) } // QueryExecuter combines the Queryer and Executer diff --git a/core/rows.go b/core/rows.go index 75d6ebf..42ab762 100644 --- a/core/rows.go +++ b/core/rows.go @@ -24,9 +24,9 @@ func (rs *Rows) ToMapString() ([]map[string]string, error) { return nil, err } - var results = make([]map[string]string, 0, 10) + results := make([]map[string]string, 0, 10) for rs.Next() { - var record = make(map[string]string, len(cols)) + record := make(map[string]string, len(cols)) err = rs.ScanMap(&record) if err != nil { return nil, err @@ -37,7 +37,7 @@ func (rs *Rows) ToMapString() ([]map[string]string, error) { } // ScanStructByIndex scan data to a struct's pointer according field index -func (rs *Rows) ScanStructByIndex(dest ...interface{}) error { +func (rs *Rows) ScanStructByIndex(dest ...any) error { if len(dest) == 0 { return errors.New("at least one struct") } @@ -56,9 +56,9 @@ func (rs *Rows) ScanStructByIndex(dest ...interface{}) error { if err != nil { return err } - newDest := make([]interface{}, len(cols)) + newDest := make([]any, len(cols)) - var i = 0 + i := 0 for _, vvv := range vvvs { for j := 0; j < vvv.NumField(); j++ { newDest[i] = vvv.Field(j).Addr().Interface() @@ -97,7 +97,7 @@ func fieldByName(v reflect.Value, name string) reflect.Value { } // ScanStructByName scan data to a struct's pointer according field name -func (rs *Rows) ScanStructByName(dest interface{}) error { +func (rs *Rows) ScanStructByName(dest any) error { vv := reflect.ValueOf(dest) if vv.Kind() != reflect.Ptr || vv.Elem().Kind() != reflect.Struct { return errors.New("dest should be a struct's pointer") @@ -108,7 +108,7 @@ func (rs *Rows) ScanStructByName(dest interface{}) error { return err } - newDest := make([]interface{}, len(cols)) + newDest := make([]any, len(cols)) var v EmptyScanner for j, name := range cols { f := fieldByName(vv.Elem(), rs.db.Mapper.Table2Obj(name)) @@ -123,7 +123,7 @@ func (rs *Rows) ScanStructByName(dest interface{}) error { } // ScanSlice scan data to a slice's pointer, slice's length should equal to columns' number -func (rs *Rows) ScanSlice(dest interface{}) error { +func (rs *Rows) ScanSlice(dest any) error { vv := reflect.ValueOf(dest) if vv.Kind() != reflect.Ptr || vv.Elem().Kind() != reflect.Slice { return errors.New("dest should be a slice's pointer") @@ -135,7 +135,7 @@ func (rs *Rows) ScanSlice(dest interface{}) error { return err } - newDest := make([]interface{}, len(cols)) + newDest := make([]any, len(cols)) for j := 0; j < len(cols); j++ { if j >= vvv.Len() { @@ -158,7 +158,7 @@ func (rs *Rows) ScanSlice(dest interface{}) error { } // ScanMap scan data to a map's pointer -func (rs *Rows) ScanMap(dest interface{}) error { +func (rs *Rows) ScanMap(dest any) error { vv := reflect.ValueOf(dest) if vv.Kind() != reflect.Ptr || vv.Elem().Kind() != reflect.Map { return errors.New("dest should be a map's pointer") @@ -169,7 +169,7 @@ func (rs *Rows) ScanMap(dest interface{}) error { return err } - newDest := make([]interface{}, len(cols)) + newDest := make([]any, len(cols)) vvv := vv.Elem() for i := range cols { @@ -217,7 +217,7 @@ func (row *Row) Columns() ([]string, error) { } // Scan retrieves all row column values -func (row *Row) Scan(dest ...interface{}) error { +func (row *Row) Scan(dest ...any) error { if row.err != nil { return row.err } @@ -244,7 +244,7 @@ func (row *Row) Scan(dest ...interface{}) error { } // ScanStructByName retrieves all row column values into a struct -func (row *Row) ScanStructByName(dest interface{}) error { +func (row *Row) ScanStructByName(dest any) error { if row.err != nil { return row.err } @@ -265,7 +265,7 @@ func (row *Row) ScanStructByName(dest interface{}) error { } // ScanStructByIndex retrieves all row column values into a struct -func (row *Row) ScanStructByIndex(dest interface{}) error { +func (row *Row) ScanStructByIndex(dest any) error { if row.err != nil { return row.err } @@ -286,7 +286,7 @@ func (row *Row) ScanStructByIndex(dest interface{}) error { } // ScanSlice scan data to a slice's pointer, slice's length should equal to columns' number -func (row *Row) ScanSlice(dest interface{}) error { +func (row *Row) ScanSlice(dest any) error { if row.err != nil { return row.err } @@ -308,7 +308,7 @@ func (row *Row) ScanSlice(dest interface{}) error { } // ScanMap scan data to a map's pointer -func (row *Row) ScanMap(dest interface{}) error { +func (row *Row) ScanMap(dest any) error { if row.err != nil { return row.err } @@ -336,7 +336,7 @@ func (row *Row) ToMapString() (map[string]string, error) { return nil, err } - var record = make(map[string]string, len(cols)) + record := make(map[string]string, len(cols)) err = row.ScanMap(&record) if err != nil { return nil, err diff --git a/core/scan.go b/core/scan.go index 1e7e452..7ba3909 100644 --- a/core/scan.go +++ b/core/scan.go @@ -13,12 +13,10 @@ import ( // NullTime defines a customize type NullTime type NullTime time.Time -var ( - _ driver.Valuer = NullTime{} -) +var _ driver.Valuer = NullTime{} // Scan implements driver.Valuer -func (ns *NullTime) Scan(value interface{}) error { +func (ns *NullTime) Scan(value any) error { if value == nil { return nil } @@ -33,7 +31,7 @@ func (ns NullTime) Value() (driver.Value, error) { return (time.Time)(ns).Format("2006-01-02 15:04:05"), nil } -func convertTime(dest *NullTime, src interface{}) error { +func convertTime(dest *NullTime, src any) error { // Common cases, without reflect. switch s := src.(type) { case string: @@ -61,10 +59,9 @@ func convertTime(dest *NullTime, src interface{}) error { } // EmptyScanner represents an empty scanner -type EmptyScanner struct { -} +type EmptyScanner struct{} // Scan implements -func (EmptyScanner) Scan(src interface{}) error { +func (EmptyScanner) Scan(src any) error { return nil } diff --git a/core/stmt.go b/core/stmt.go index 3247efe..8f5a46e 100644 --- a/core/stmt.go +++ b/core/stmt.go @@ -49,13 +49,13 @@ func (db *DB) Prepare(query string) (*Stmt, error) { } // ExecMapContext execute with map -func (s *Stmt) ExecMapContext(ctx context.Context, mp interface{}) (sql.Result, error) { +func (s *Stmt) ExecMapContext(ctx context.Context, mp any) (sql.Result, error) { vv := reflect.ValueOf(mp) if vv.Kind() != reflect.Ptr || vv.Elem().Kind() != reflect.Map { return nil, errors.New("mp should be a map's pointer") } - args := make([]interface{}, len(s.names)) + args := make([]any, len(s.names)) for k, i := range s.names { args[i] = vv.Elem().MapIndex(reflect.ValueOf(k)).Interface() } @@ -63,18 +63,18 @@ func (s *Stmt) ExecMapContext(ctx context.Context, mp interface{}) (sql.Result, } // ExecMap executes with map -func (s *Stmt) ExecMap(mp interface{}) (sql.Result, error) { +func (s *Stmt) ExecMap(mp any) (sql.Result, error) { return s.ExecMapContext(context.Background(), mp) } // ExecStructContext executes with struct -func (s *Stmt) ExecStructContext(ctx context.Context, st interface{}) (sql.Result, error) { +func (s *Stmt) ExecStructContext(ctx context.Context, st any) (sql.Result, error) { vv := reflect.ValueOf(st) if vv.Kind() != reflect.Ptr || vv.Elem().Kind() != reflect.Struct { return nil, errors.New("mp should be a map's pointer") } - args := make([]interface{}, len(s.names)) + args := make([]any, len(s.names)) for k, i := range s.names { args[i] = vv.Elem().FieldByName(k).Interface() } @@ -82,12 +82,12 @@ func (s *Stmt) ExecStructContext(ctx context.Context, st interface{}) (sql.Resul } // ExecStruct executes with struct -func (s *Stmt) ExecStruct(st interface{}) (sql.Result, error) { +func (s *Stmt) ExecStruct(st any) (sql.Result, error) { return s.ExecStructContext(context.Background(), st) } // ExecContext with args -func (s *Stmt) ExecContext(ctx context.Context, args ...interface{}) (sql.Result, error) { +func (s *Stmt) ExecContext(ctx context.Context, args ...any) (sql.Result, error) { hookCtx := contexts.NewContextHook(ctx, s.query, args) ctx, err := s.db.beforeProcess(hookCtx) if err != nil { @@ -102,7 +102,7 @@ func (s *Stmt) ExecContext(ctx context.Context, args ...interface{}) (sql.Result } // QueryContext query with args -func (s *Stmt) QueryContext(ctx context.Context, args ...interface{}) (*Rows, error) { +func (s *Stmt) QueryContext(ctx context.Context, args ...any) (*Rows, error) { hookCtx := contexts.NewContextHook(ctx, s.query, args) ctx, err := s.db.beforeProcess(hookCtx) if err != nil { @@ -117,18 +117,18 @@ func (s *Stmt) QueryContext(ctx context.Context, args ...interface{}) (*Rows, er } // Query query with args -func (s *Stmt) Query(args ...interface{}) (*Rows, error) { +func (s *Stmt) Query(args ...any) (*Rows, error) { return s.QueryContext(context.Background(), args...) } // QueryMapContext query with map -func (s *Stmt) QueryMapContext(ctx context.Context, mp interface{}) (*Rows, error) { +func (s *Stmt) QueryMapContext(ctx context.Context, mp any) (*Rows, error) { vv := reflect.ValueOf(mp) if vv.Kind() != reflect.Ptr || vv.Elem().Kind() != reflect.Map { return nil, errors.New("mp should be a map's pointer") } - args := make([]interface{}, len(s.names)) + args := make([]any, len(s.names)) for k, i := range s.names { args[i] = vv.Elem().MapIndex(reflect.ValueOf(k)).Interface() } @@ -137,18 +137,18 @@ func (s *Stmt) QueryMapContext(ctx context.Context, mp interface{}) (*Rows, erro } // QueryMap query with map -func (s *Stmt) QueryMap(mp interface{}) (*Rows, error) { +func (s *Stmt) QueryMap(mp any) (*Rows, error) { return s.QueryMapContext(context.Background(), mp) } // QueryStructContext query with struct -func (s *Stmt) QueryStructContext(ctx context.Context, st interface{}) (*Rows, error) { +func (s *Stmt) QueryStructContext(ctx context.Context, st any) (*Rows, error) { vv := reflect.ValueOf(st) if vv.Kind() != reflect.Ptr || vv.Elem().Kind() != reflect.Struct { return nil, errors.New("mp should be a map's pointer") } - args := make([]interface{}, len(s.names)) + args := make([]any, len(s.names)) for k, i := range s.names { args[i] = vv.Elem().FieldByName(k).Interface() } @@ -157,29 +157,29 @@ func (s *Stmt) QueryStructContext(ctx context.Context, st interface{}) (*Rows, e } // QueryStruct query with struct -func (s *Stmt) QueryStruct(st interface{}) (*Rows, error) { +func (s *Stmt) QueryStruct(st any) (*Rows, error) { return s.QueryStructContext(context.Background(), st) } // QueryRowContext query row with args -func (s *Stmt) QueryRowContext(ctx context.Context, args ...interface{}) *Row { +func (s *Stmt) QueryRowContext(ctx context.Context, args ...any) *Row { rows, err := s.QueryContext(ctx, args...) return &Row{rows, err} } // QueryRow query row with args -func (s *Stmt) QueryRow(args ...interface{}) *Row { +func (s *Stmt) QueryRow(args ...any) *Row { return s.QueryRowContext(context.Background(), args...) } // QueryRowMapContext query row with map -func (s *Stmt) QueryRowMapContext(ctx context.Context, mp interface{}) *Row { +func (s *Stmt) QueryRowMapContext(ctx context.Context, mp any) *Row { vv := reflect.ValueOf(mp) if vv.Kind() != reflect.Ptr || vv.Elem().Kind() != reflect.Map { return &Row{nil, errors.New("mp should be a map's pointer")} } - args := make([]interface{}, len(s.names)) + args := make([]any, len(s.names)) for k, i := range s.names { args[i] = vv.Elem().MapIndex(reflect.ValueOf(k)).Interface() } @@ -188,18 +188,18 @@ func (s *Stmt) QueryRowMapContext(ctx context.Context, mp interface{}) *Row { } // QueryRowMap query row with map -func (s *Stmt) QueryRowMap(mp interface{}) *Row { +func (s *Stmt) QueryRowMap(mp any) *Row { return s.QueryRowMapContext(context.Background(), mp) } // QueryRowStructContext query row with struct -func (s *Stmt) QueryRowStructContext(ctx context.Context, st interface{}) *Row { +func (s *Stmt) QueryRowStructContext(ctx context.Context, st any) *Row { vv := reflect.ValueOf(st) if vv.Kind() != reflect.Ptr || vv.Elem().Kind() != reflect.Struct { return &Row{nil, errors.New("st should be a struct's pointer")} } - args := make([]interface{}, len(s.names)) + args := make([]any, len(s.names)) for k, i := range s.names { args[i] = vv.Elem().FieldByName(k).Interface() } @@ -208,6 +208,6 @@ func (s *Stmt) QueryRowStructContext(ctx context.Context, st interface{}) *Row { } // QueryRowStruct query row with struct -func (s *Stmt) QueryRowStruct(st interface{}) *Row { +func (s *Stmt) QueryRowStruct(st any) *Row { return s.QueryRowStructContext(context.Background(), st) } diff --git a/core/tx.go b/core/tx.go index a2f745f..954be89 100644 --- a/core/tx.go +++ b/core/tx.go @@ -11,9 +11,7 @@ import ( "xorm.io/xorm/contexts" ) -var ( - _ QueryExecuter = &Tx{} -) +var _ QueryExecuter = &Tx{} // Tx represents a transaction type Tx struct { @@ -105,7 +103,7 @@ func (tx *Tx) Stmt(stmt *Stmt) *Stmt { } // ExecMapContext executes query with args in a map -func (tx *Tx) ExecMapContext(ctx context.Context, query string, mp interface{}) (sql.Result, error) { +func (tx *Tx) ExecMapContext(ctx context.Context, query string, mp any) (sql.Result, error) { query, args, err := MapToSlice(query, mp) if err != nil { return nil, err @@ -114,12 +112,12 @@ func (tx *Tx) ExecMapContext(ctx context.Context, query string, mp interface{}) } // ExecMap executes query with args in a map -func (tx *Tx) ExecMap(query string, mp interface{}) (sql.Result, error) { +func (tx *Tx) ExecMap(query string, mp any) (sql.Result, error) { return tx.ExecMapContext(context.Background(), query, mp) } // ExecStructContext executes query with args in a struct -func (tx *Tx) ExecStructContext(ctx context.Context, query string, st interface{}) (sql.Result, error) { +func (tx *Tx) ExecStructContext(ctx context.Context, query string, st any) (sql.Result, error) { query, args, err := StructToSlice(query, st) if err != nil { return nil, err @@ -128,7 +126,7 @@ func (tx *Tx) ExecStructContext(ctx context.Context, query string, st interface{ } // ExecContext executes a query with args -func (tx *Tx) ExecContext(ctx context.Context, query string, args ...interface{}) (sql.Result, error) { +func (tx *Tx) ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error) { hookCtx := contexts.NewContextHook(ctx, query, args) ctx, err := tx.db.beforeProcess(hookCtx) if err != nil { @@ -143,12 +141,12 @@ func (tx *Tx) ExecContext(ctx context.Context, query string, args ...interface{} } // ExecStruct executes query with args in a struct -func (tx *Tx) ExecStruct(query string, st interface{}) (sql.Result, error) { +func (tx *Tx) ExecStruct(query string, st any) (sql.Result, error) { return tx.ExecStructContext(context.Background(), query, st) } // QueryContext query with args -func (tx *Tx) QueryContext(ctx context.Context, query string, args ...interface{}) (*Rows, error) { +func (tx *Tx) QueryContext(ctx context.Context, query string, args ...any) (*Rows, error) { hookCtx := contexts.NewContextHook(ctx, query, args) ctx, err := tx.db.beforeProcess(hookCtx) if err != nil { @@ -166,12 +164,12 @@ func (tx *Tx) QueryContext(ctx context.Context, query string, args ...interface{ } // Query query with args -func (tx *Tx) Query(query string, args ...interface{}) (*Rows, error) { +func (tx *Tx) Query(query string, args ...any) (*Rows, error) { return tx.QueryContext(context.Background(), query, args...) } // QueryMapContext query with args in a map -func (tx *Tx) QueryMapContext(ctx context.Context, query string, mp interface{}) (*Rows, error) { +func (tx *Tx) QueryMapContext(ctx context.Context, query string, mp any) (*Rows, error) { query, args, err := MapToSlice(query, mp) if err != nil { return nil, err @@ -180,12 +178,12 @@ func (tx *Tx) QueryMapContext(ctx context.Context, query string, mp interface{}) } // QueryMap query with args in a map -func (tx *Tx) QueryMap(query string, mp interface{}) (*Rows, error) { +func (tx *Tx) QueryMap(query string, mp any) (*Rows, error) { return tx.QueryMapContext(context.Background(), query, mp) } // QueryStructContext query with args in struct -func (tx *Tx) QueryStructContext(ctx context.Context, query string, st interface{}) (*Rows, error) { +func (tx *Tx) QueryStructContext(ctx context.Context, query string, st any) (*Rows, error) { query, args, err := StructToSlice(query, st) if err != nil { return nil, err @@ -194,23 +192,23 @@ func (tx *Tx) QueryStructContext(ctx context.Context, query string, st interface } // QueryStruct query with args in struct -func (tx *Tx) QueryStruct(query string, st interface{}) (*Rows, error) { +func (tx *Tx) QueryStruct(query string, st any) (*Rows, error) { return tx.QueryStructContext(context.Background(), query, st) } // QueryRowContext query one row with args -func (tx *Tx) QueryRowContext(ctx context.Context, query string, args ...interface{}) *Row { +func (tx *Tx) QueryRowContext(ctx context.Context, query string, args ...any) *Row { rows, err := tx.QueryContext(ctx, query, args...) return &Row{rows, err} } // QueryRow query one row with args -func (tx *Tx) QueryRow(query string, args ...interface{}) *Row { +func (tx *Tx) QueryRow(query string, args ...any) *Row { return tx.QueryRowContext(context.Background(), query, args...) } // QueryRowMapContext query one row with args in a map -func (tx *Tx) QueryRowMapContext(ctx context.Context, query string, mp interface{}) *Row { +func (tx *Tx) QueryRowMapContext(ctx context.Context, query string, mp any) *Row { query, args, err := MapToSlice(query, mp) if err != nil { return &Row{nil, err} @@ -219,12 +217,12 @@ func (tx *Tx) QueryRowMapContext(ctx context.Context, query string, mp interface } // QueryRowMap query one row with args in a map -func (tx *Tx) QueryRowMap(query string, mp interface{}) *Row { +func (tx *Tx) QueryRowMap(query string, mp any) *Row { return tx.QueryRowMapContext(context.Background(), query, mp) } // QueryRowStructContext query one row with args in struct -func (tx *Tx) QueryRowStructContext(ctx context.Context, query string, st interface{}) *Row { +func (tx *Tx) QueryRowStructContext(ctx context.Context, query string, st any) *Row { query, args, err := StructToSlice(query, st) if err != nil { return &Row{nil, err} @@ -233,6 +231,6 @@ func (tx *Tx) QueryRowStructContext(ctx context.Context, query string, st interf } // QueryRowStruct query one row with args in struct -func (tx *Tx) QueryRowStruct(query string, st interface{}) *Row { +func (tx *Tx) QueryRowStruct(query string, st any) *Row { return tx.QueryRowStructContext(context.Background(), query, st) } diff --git a/dialects/dameng.go b/dialects/dameng.go index d1120fe..0558e49 100644 --- a/dialects/dameng.go +++ b/dialects/dameng.go @@ -654,7 +654,7 @@ func (db *dameng) IsReserved(name string) bool { } func (db *dameng) DropTableSQL(tableName string) (string, bool) { - return fmt.Sprintf("DROP TABLE %s", db.quoter.Quote(tableName)), false + return "DROP TABLE " + db.quoter.Quote(tableName), false } // ModifyColumnSQL returns a SQL to modify SQL @@ -685,9 +685,10 @@ func (db *dameng) CreateTableSQL(ctx context.Context, queryer core.Queryer, tabl for i, colName := range table.ColumnsSeq() { col := table.GetColumn(colName) if col.SQLType.IsBool() && !col.DefaultIsEmpty { - if col.Default == "true" { + switch col.Default { + case "true": col.Default = "1" - } else if col.Default == "false" { + case "false": col.Default = "0" } } @@ -749,8 +750,8 @@ func (db *dameng) SetQuotePolicy(quotePolicy QuotePolicy) { } } -func (db *dameng) IndexCheckSQL(tableName, idxName string) (string, []interface{}) { - args := []interface{}{tableName, idxName} +func (db *dameng) IndexCheckSQL(tableName, idxName string) (string, []any) { + args := []any{tableName, idxName} return `SELECT INDEX_NAME FROM USER_INDEXES ` + `WHERE TABLE_NAME = ? AND INDEX_NAME = ?`, args } @@ -780,7 +781,7 @@ func (db *dameng) IsSequenceExist(ctx context.Context, queryer core.Queryer, seq } func (db *dameng) IsColumnExist(queryer core.Queryer, ctx context.Context, tableName, colName string) (bool, error) { - args := []interface{}{tableName, colName} + args := []any{tableName, colName} query := "SELECT column_name FROM USER_TAB_COLUMNS WHERE table_name = ?" + " AND column_name = ?" return db.HasRecords(queryer, ctx, query, args...) @@ -800,7 +801,7 @@ type dmClobObject interface { // var _ dmClobObject = &dm.DmClob{} -func (d *dmClobScanner) Scan(data interface{}) error { +func (d *dmClobScanner) Scan(data any) error { if data == nil { return nil } @@ -832,10 +833,10 @@ func (d *dmClobScanner) Scan(data interface{}) error { d.valid = true return nil case string: - if len(t) <= 0 { + if len(t) == 0 { return nil } - d.data = string(t) + d.data = t d.valid = true return nil default: @@ -1011,7 +1012,7 @@ func (db *dameng) GetColumns(queryer core.Queryer, ctx context.Context, tableNam func (db *dameng) GetTables(queryer core.Queryer, ctx context.Context) ([]*schemas.Table, error) { s := "SELECT table_name FROM user_tables WHERE temporary = 'N' AND table_name NOT LIKE ?" - args := []interface{}{strings.ToUpper(db.uri.User), "%$%"} + args := []any{strings.ToUpper(db.uri.User), "%$%"} rows, err := queryer.QueryContext(ctx, s, args...) if err != nil { @@ -1036,7 +1037,7 @@ func (db *dameng) GetTables(queryer core.Queryer, ctx context.Context) ([]*schem } func (db *dameng) GetIndexes(queryer core.Queryer, ctx context.Context, tableName string) (map[string]*schemas.Index, error) { - args := []interface{}{tableName, tableName} + args := []any{tableName, tableName} s := "SELECT t.column_name,i.uniqueness,i.index_name FROM user_ind_columns t,user_indexes i " + "WHERE t.index_name = i.index_name and t.table_name = i.table_name and t.table_name =?" + " AND t.index_name not in (SELECT index_name FROM ALL_CONSTRAINTS WHERE CONSTRAINT_TYPE='P' AND table_name = ?)" @@ -1092,9 +1093,7 @@ func (db *dameng) Filters() []Filter { return []Filter{} } -type damengDriver struct { - baseDriver -} +type damengDriver struct{} // Features return features func (d *damengDriver) Features() *DriverFeatures { @@ -1127,7 +1126,7 @@ func (d *damengDriver) Parse(driverName, dataSourceName string) (*URI, error) { }, nil } -func (d *damengDriver) GenScanResult(colType string) (interface{}, error) { +func (d *damengDriver) GenScanResult(colType string) (any, error) { switch colType { case "CHAR", "NCHAR", "VARCHAR", "VARCHAR2", "NVARCHAR2", "LONG", "CLOB", "NCLOB": var s sql.NullString @@ -1156,13 +1155,13 @@ func (d *damengDriver) GenScanResult(colType string) (interface{}, error) { } } -func (d *damengDriver) Scan(ctx *ScanContext, rows *core.Rows, types []*sql.ColumnType, vv ...interface{}) error { - scanResults := make([]interface{}, 0, len(types)) +func (d *damengDriver) Scan(ctx *ScanContext, rows *core.Rows, types []*sql.ColumnType, vv ...any) error { + scanResults := make([]any, 0, len(types)) replaces := make([]bool, 0, len(types)) var err error for i, v := range vv { var replaced bool - var scanResult interface{} + var scanResult any switch types[i].DatabaseTypeName() { case "CLOB", "TEXT": scanResult = &dmClobScanner{} @@ -1186,7 +1185,7 @@ func (d *damengDriver) Scan(ctx *ScanContext, rows *core.Rows, types []*sql.Colu if replaced { switch t := scanResults[i].(type) { case *dmClobScanner: - var d interface{} + var d any if t.valid { d = t.data } else { @@ -1204,7 +1203,7 @@ func (d *damengDriver) Scan(ctx *ScanContext, rows *core.Rows, types []*sql.Colu } s := ns.String fields := strings.Split(s, "+") - if err := convert.Assign(vv[i], strings.Replace(fields[0], "T", " ", -1), ctx.DBLocation, ctx.UserLocation); err != nil { + if err := convert.Assign(vv[i], strings.ReplaceAll(fields[0], "T", " "), ctx.DBLocation, ctx.UserLocation); err != nil { return err } default: diff --git a/dialects/dialect.go b/dialects/dialect.go index f907c50..9a93df0 100644 --- a/dialects/dialect.go +++ b/dialects/dialect.go @@ -6,6 +6,7 @@ package dialects import ( "context" + "errors" "fmt" "strings" "time" @@ -67,7 +68,7 @@ type Dialect interface { AutoIncrStr() string GetIndexes(queryer core.Queryer, ctx context.Context, tableName string) (map[string]*schemas.Index, error) - IndexCheckSQL(tableName, idxName string) (string, []interface{}) + IndexCheckSQL(tableName, idxName string) (string, []any) CreateIndexSQL(tableName string, index *schemas.Index) string DropIndexSQL(tableName string, index *schemas.Index) string @@ -81,7 +82,7 @@ type Dialect interface { DropSequenceSQL(seqName string) (string, error) GetColumns(queryer core.Queryer, ctx context.Context, tableName string) ([]string, map[string]*schemas.Column, error) - IsColumnExist(queryer core.Queryer, ctx context.Context, tableName string, colName string) (bool, error) + IsColumnExist(queryer core.Queryer, ctx context.Context, tableName, colName string) (bool, error) AddColumnSQL(tableName string, col *schemas.Column) string ModifyColumnSQL(tableName string, col *schemas.Column) string @@ -163,21 +164,21 @@ func (db *Base) CreateSequenceSQL(ctx context.Context, queryer core.Queryer, seq } func (db *Base) IsSequenceExist(ctx context.Context, queryer core.Queryer, seqName string) (bool, error) { - return false, fmt.Errorf("unsupported sequence feature") + return false, errors.New("unsupported sequence feature") } func (db *Base) DropSequenceSQL(seqName string) (string, error) { - return fmt.Sprintf("DROP SEQUENCE %s", seqName), nil + return "DROP SEQUENCE " + seqName, nil } // DropTableSQL returns drop table SQL func (db *Base) DropTableSQL(tableName string) (string, bool) { quote := db.dialect.Quoter().Quote - return fmt.Sprintf("DROP TABLE IF EXISTS %s", quote(tableName)), true + return "DROP TABLE IF EXISTS " + quote(tableName), true } // HasRecords returns true if the SQL has records returned -func (db *Base) HasRecords(queryer core.Queryer, ctx context.Context, query string, args ...interface{}) (bool, error) { +func (db *Base) HasRecords(queryer core.Queryer, ctx context.Context, query string, args ...any) (bool, error) { rows, err := queryer.QueryContext(ctx, query, args...) if err != nil { return false, err diff --git a/dialects/driver.go b/dialects/driver.go index c63dbfa..2b7e3d1 100644 --- a/dialects/driver.go +++ b/dialects/driver.go @@ -27,13 +27,11 @@ type DriverFeatures struct { type Driver interface { Parse(string, string) (*URI, error) Features() *DriverFeatures - GenScanResult(string) (interface{}, error) // according given column type generating a suitable scan interface - Scan(*ScanContext, *core.Rows, []*sql.ColumnType, ...interface{}) error + GenScanResult(string) (any, error) // according given column type generating a suitable scan interface + Scan(*ScanContext, *core.Rows, []*sql.ColumnType, ...any) error } -var ( - drivers = map[string]Driver{} -) +var drivers = map[string]Driver{} // RegisterDriver register a driver func RegisterDriver(driverName string, driver Driver) { @@ -73,13 +71,15 @@ func OpenDialect(driverName, connstr string) (Dialect, error) { return nil, fmt.Errorf("unsupported dialect type: %v", uri.DBType) } - dialect.Init(uri) + if err := dialect.Init(uri); err != nil { + return nil, err + } return dialect, nil } type baseDriver struct{} -func (b *baseDriver) Scan(ctx *ScanContext, rows *core.Rows, types []*sql.ColumnType, v ...interface{}) error { +func (b *baseDriver) Scan(ctx *ScanContext, rows *core.Rows, types []*sql.ColumnType, v ...any) error { return rows.Scan(v...) } diff --git a/dialects/gbase8s.go b/dialects/gbase8s.go index d3a8fd2..59bdec5 100644 --- a/dialects/gbase8s.go +++ b/dialects/gbase8s.go @@ -596,7 +596,7 @@ func (db *gbase8s) SQLType(c *schemas.Column) string { res = "NUMERIC(1,0)" case schemas.Char, schemas.NChar, schemas.Uuid: res = "CHAR" - case schemas.Varchar, schemas.NVarchar, schemas.VARCHAR2, schemas.NVarchar: + case schemas.Varchar, schemas.NVarchar, schemas.VARCHAR2: res = "VARCHAR" case schemas.Enum, schemas.Set: res = "VARCHAR(255)" @@ -615,9 +615,8 @@ func (db *gbase8s) SQLType(c *schemas.Column) string { case schemas.DateTime, schemas.SmallDateTime, schemas.TimeStamp, schemas.TimeStampz, schemas.Time: if c.Length >= 1 && c.Length <= 5 { return fmt.Sprintf("DATETIME YEAR TO FRACTION(%d)", c.Length) - } else { - return "DATETIME YEAR TO FRACTION(5)" } + return "DATETIME YEAR TO FRACTION(5)" default: res = t } @@ -745,8 +744,8 @@ func (db *gbase8s) GetIndexes(queryer core.Queryer, ctx context.Context, tableNa return indexes, nil } -func (db *gbase8s) IndexCheckSQL(tableName, idxName string) (string, []interface{}) { - args := []interface{}{tableName, idxName} +func (db *gbase8s) IndexCheckSQL(tableName, idxName string) (string, []any) { + args := []any{tableName, idxName} sql := `SELECT idx.idxname FROM sysindexes idx JOIN systables tab ON idx.tabid = tab.tabid JOIN syscolumns col ON col.tabid = tab.tabid @@ -765,7 +764,7 @@ func (db *gbase8s) DropIndexSQL(tableName string, index *schemas.Index) string { } else { name = index.Name } - return fmt.Sprintf("DROP INDEX %s", db.quoter.Quote(name)) + return "DROP INDEX " + db.quoter.Quote(name) } func (db *gbase8s) GetTables(queryer core.Queryer, ctx context.Context) ([]*schemas.Table, error) { @@ -808,7 +807,7 @@ func (db *gbase8s) CreateTableSQL(ctx context.Context, queryer core.Queryer, tab s, _ := ColumnString(db, col, false, false) sql += s if len(col.Comment) > 0 { - sql += fmt.Sprintf(" COMMENT '%s'", col.Comment) + sql += " COMMENT '" + col.Comment + "'" } sql = strings.TrimSpace(sql) sql += ", " @@ -823,7 +822,7 @@ func (db *gbase8s) CreateTableSQL(ctx context.Context, queryer core.Queryer, tab } func (db *gbase8s) DropTableSQL(tableName string) (string, bool) { - return fmt.Sprintf("DROP TABLE %s", db.quoter.Quote(tableName)), false + return "DROP TABLE " + db.quoter.Quote(tableName), false } func (db *gbase8s) GetColumns(queryer core.Queryer, ctx context.Context, tableName string) ([]string, map[string]*schemas.Column, error) { @@ -959,7 +958,7 @@ func (db *gbase8s) ModifyColumnSQL(tableName string, col *schemas.Column) string } func (db *gbase8s) IsColumnExist(queryer core.Queryer, ctx context.Context, tableName, colName string) (bool, error) { - args := []interface{}{tableName, colName} + args := []any{tableName, colName} query := "SELECT colname FROM syscolumnsext c, systables t WHERE c.tabid = t.tabid and tabname = :1 AND colname = :2" return db.HasRecords(queryer, ctx, query, args...) } @@ -978,7 +977,7 @@ func (g *gbase8sDriver) Features() *DriverFeatures { } } -func (g *gbase8sDriver) GenScanResult(colType string) (interface{}, error) { +func (g *gbase8sDriver) GenScanResult(colType string) (any, error) { colType = strings.Replace(colType, "SQLT_", "", 1) switch colType { case "CHAR", "NCHAR", "VARCHAR", "VARCHAR2", "NVARCHAR2", "AFC": @@ -1024,7 +1023,7 @@ func (g *gbase8sDriver) GenScanResult(colType string) (interface{}, error) { // dataSourceName=user/password@ipv4:port/dbname // gbase8s://user:password@ip:port/dbname?param2=1¶m2=2 -func (o *gbase8sDriver) Parse(driverName, dataSourceName string) (*URI, error) { +func (g *gbase8sDriver) Parse(driverName, dataSourceName string) (*URI, error) { db := &URI{DBType: schemas.GBASE8S} dsnPattern := regexp.MustCompile( `^(?P.*):(?P.*)@` + // user:password@ diff --git a/dialects/mssql.go b/dialects/mssql.go index a7a8da1..b077ace 100644 --- a/dialects/mssql.go +++ b/dialects/mssql.go @@ -438,8 +438,8 @@ func (db *mssql) ModifyColumnSQL(tableName string, col *schemas.Column) string { return fmt.Sprintf("ALTER TABLE %s ALTER COLUMN %s", db.quoter.Quote(tableName), s) } -func (db *mssql) IndexCheckSQL(tableName, idxName string) (string, []interface{}) { - args := []interface{}{idxName} +func (db *mssql) IndexCheckSQL(tableName, idxName string) (string, []any) { + args := []any{idxName} sql := "select name from sysindexes where id=object_id('" + tableName + "') and name=?" return sql, args } @@ -456,7 +456,7 @@ func (db *mssql) IsTableExist(queryer core.Queryer, ctx context.Context, tableNa } func (db *mssql) GetColumns(queryer core.Queryer, ctx context.Context, tableName string) ([]string, map[string]*schemas.Column, error) { - args := []interface{}{} + args := []any{} s := `select a.name as name, b.name as ctype,a.max_length,a.precision,a.scale,a.is_nullable as nullable, "default_is_null" = (CASE WHEN c.text is null THEN 1 ELSE 0 END), replace(replace(isnull(c.text,''),'(',''),')','') as vdefault, @@ -550,7 +550,7 @@ func (db *mssql) GetColumns(queryer core.Queryer, ctx context.Context, tableName } func (db *mssql) GetTables(queryer core.Queryer, ctx context.Context) ([]*schemas.Table, error) { - args := []interface{}{} + args := []any{} s := `select name from sysobjects where xtype ='U'` rows, err := queryer.QueryContext(ctx, s, args...) @@ -577,7 +577,7 @@ func (db *mssql) GetTables(queryer core.Queryer, ctx context.Context) ([]*schema } func (db *mssql) GetIndexes(queryer core.Queryer, ctx context.Context, tableName string) (map[string]*schemas.Index, error) { - args := []interface{}{tableName} + args := []any{tableName} s := `SELECT IXS.NAME AS [INDEX_NAME], C.NAME AS [COLUMN_NAME], @@ -649,9 +649,13 @@ func (db *mssql) CreateTableSQL(ctx context.Context, queryer core.Queryer, table quoter := db.dialect.Quoter() var b strings.Builder b.WriteString("IF NOT EXISTS (SELECT [name] FROM sys.tables WHERE [name] = '") - quoter.QuoteTo(&b, tableName) + if err := quoter.QuoteTo(&b, tableName); err != nil { + return "", false, err + } b.WriteString("' ) CREATE TABLE ") - quoter.QuoteTo(&b, tableName) + if err := quoter.QuoteTo(&b, tableName); err != nil { + return "", false, err + } b.WriteString(" (") for i, colName := range table.ColumnsSeq() { @@ -703,7 +707,7 @@ func (p *odbcDriver) Parse(driverName, dataSourceName string) (*URI, error) { for _, c := range kv { vv := strings.Split(strings.TrimSpace(c), "=") if len(vv) == 2 { - if strings.ToLower(vv[0]) == "database" { + if strings.EqualFold(vv[0], "database") { dbName = vv[1] } } @@ -715,7 +719,7 @@ func (p *odbcDriver) Parse(driverName, dataSourceName string) (*URI, error) { return &URI{DBName: dbName, DBType: schemas.MSSQL}, nil } -func (p *odbcDriver) GenScanResult(colType string) (interface{}, error) { +func (p *odbcDriver) GenScanResult(colType string) (any, error) { switch colType { case "VARCHAR", "TEXT", "CHAR", "NVARCHAR", "NCHAR", "NTEXT": fallthrough diff --git a/dialects/mysql.go b/dialects/mysql.go index 424807d..f4dc6ca 100644 --- a/dialects/mysql.go +++ b/dialects/mysql.go @@ -281,22 +281,18 @@ func (db *mysql) SQLType(c *schemas.Column) string { c.Length = 64 case schemas.Enum: // mysql enum res = schemas.Enum - res += "(" - opts := "" + opts := make([]string, 0, len(c.EnumOptions)) for v := range c.EnumOptions { - opts += fmt.Sprintf(",'%v'", v) + opts = append(opts, fmt.Sprintf("'%v'", v)) } - res += strings.TrimLeft(opts, ",") - res += ")" + res += "(" + strings.Join(opts, ",") + ")" case schemas.Set: // mysql set res = schemas.Set - res += "(" - opts := "" + opts := make([]string, 0, len(c.SetOptions)) for v := range c.SetOptions { - opts += fmt.Sprintf(",'%v'", v) + opts = append(opts, fmt.Sprintf("'%v'", v)) } - res += strings.TrimLeft(opts, ",") - res += ")" + res += "(" + strings.Join(opts, ",") + ")" case schemas.NVarchar: res = schemas.Varchar case schemas.Uuid: @@ -371,8 +367,8 @@ func (db *mysql) AutoIncrStr() string { return "AUTO_INCREMENT" } -func (db *mysql) IndexCheckSQL(tableName, idxName string) (string, []interface{}) { - args := []interface{}{db.uri.DBName, tableName, idxName} +func (db *mysql) IndexCheckSQL(tableName, idxName string) (string, []any) { + args := []any{db.uri.DBName, tableName, idxName} sql := "SELECT `INDEX_NAME` FROM `INFORMATION_SCHEMA`.`STATISTICS`" sql += " WHERE `TABLE_SCHEMA` = ? AND `TABLE_NAME` = ? AND `INDEX_NAME`=?" return sql, args @@ -388,7 +384,7 @@ func (db *mysql) AddColumnSQL(tableName string, col *schemas.Column) string { s, _ := ColumnString(db, col, true, true) var b strings.Builder b.WriteString("ALTER TABLE ") - quoter.QuoteTo(&b, tableName) + b.WriteString(quoter.Quote(tableName)) b.WriteString(" ADD ") b.WriteString(s) if len(col.Comment) > 0 { @@ -406,13 +402,13 @@ func (db *mysql) ModifyColumnSQL(tableName string, col *schemas.Column) string { s += " " + db.AutoIncrStr() } if col.Comment != "" { - s += fmt.Sprintf(" COMMENT '%s'", col.Comment) + s += " COMMENT '" + col.Comment + "'" } return fmt.Sprintf("ALTER TABLE %s MODIFY COLUMN %s", db.quoter.Quote(tableName), s) } func (db *mysql) GetColumns(queryer core.Queryer, ctx context.Context, tableName string) ([]string, map[string]*schemas.Column, error) { - args := []interface{}{db.uri.DBName, tableName} + args := []any{db.uri.DBName, tableName} alreadyQuoted := "(INSTR(VERSION(), 'maria') > 0 && " + "(SUBSTRING_INDEX(VERSION(), '.', 1) > 10 || " + "(SUBSTRING_INDEX(VERSION(), '.', 1) = 10 && " + @@ -549,7 +545,7 @@ func (db *mysql) GetColumns(queryer core.Queryer, ctx context.Context, tableName } func (db *mysql) GetTables(queryer core.Queryer, ctx context.Context) ([]*schemas.Table, error) { - args := []interface{}{db.uri.DBName} + args := []any{db.uri.DBName} s := "SELECT `TABLE_NAME`, `ENGINE`, `AUTO_INCREMENT`, `TABLE_COMMENT`, `TABLE_COLLATION` from " + "`INFORMATION_SCHEMA`.`TABLES` WHERE `TABLE_SCHEMA`=? AND (`ENGINE`='MyISAM' OR `ENGINE` = 'InnoDB' OR `ENGINE` = 'TokuDB')" @@ -601,7 +597,7 @@ func (db *mysql) SetQuotePolicy(quotePolicy QuotePolicy) { } func (db *mysql) GetIndexes(queryer core.Queryer, ctx context.Context, tableName string) (map[string]*schemas.Index, error) { - args := []interface{}{db.uri.DBName, tableName} + args := []any{db.uri.DBName, tableName} s := "SELECT `INDEX_NAME`, `NON_UNIQUE`, `COLUMN_NAME` FROM `INFORMATION_SCHEMA`.`STATISTICS` WHERE `TABLE_SCHEMA` = ? AND `TABLE_NAME` = ? ORDER BY `SEQ_IN_INDEX`" rows, err := queryer.QueryContext(ctx, s, args...) @@ -661,7 +657,7 @@ func (db *mysql) CreateTableSQL(ctx context.Context, queryer core.Queryer, table quoter := db.dialect.Quoter() var b strings.Builder b.WriteString("CREATE TABLE IF NOT EXISTS ") - quoter.QuoteTo(&b, tableName) + b.WriteString(quoter.Quote(tableName)) b.WriteString(" (") for i, colName := range table.ColumnsSeq() { @@ -763,8 +759,8 @@ func (p *mysqlDriver) Parse(driverName, dataSourceName string) (*URI, error) { return uri, nil } -func (p *mysqlDriver) GenScanResult(colType string) (interface{}, error) { - colType = strings.Replace(colType, "UNSIGNED ", "", -1) +func (p *mysqlDriver) GenScanResult(colType string) (any, error) { + colType = strings.ReplaceAll(colType, "UNSIGNED ", "") switch colType { case "CHAR", "VARCHAR", "TINYTEXT", "TEXT", "MEDIUMTEXT", "LONGTEXT", "ENUM", "SET", "JSON": var s sql.NullString diff --git a/dialects/oracle.go b/dialects/oracle.go index 0170bca..c8fd451 100644 --- a/dialects/oracle.go +++ b/dialects/oracle.go @@ -561,9 +561,10 @@ func (db *oracle) SQLType(c *schemas.Column) string { var res string switch t := c.SQLType.Name; t { case schemas.Bool: - if c.Default == "true" { + switch c.Default { + case "true": c.Default = "1" - } else if c.Default == "false" { + case "false": c.Default = "0" } res = "NUMBER(1,0)" @@ -678,8 +679,8 @@ func (db *oracle) SetQuotePolicy(quotePolicy QuotePolicy) { } } -func (db *oracle) IndexCheckSQL(tableName, idxName string) (string, []interface{}) { - args := []interface{}{tableName, idxName} +func (db *oracle) IndexCheckSQL(tableName, idxName string) (string, []any) { + args := []any{tableName, idxName} return `SELECT INDEX_NAME FROM USER_INDEXES ` + `WHERE TABLE_NAME = :1 AND INDEX_NAME = :2`, args } @@ -700,14 +701,14 @@ func (db *oracle) IsTableExist(queryer core.Queryer, ctx context.Context, tableN } func (db *oracle) IsColumnExist(queryer core.Queryer, ctx context.Context, tableName, colName string) (bool, error) { - args := []interface{}{tableName, colName} + args := []any{tableName, colName} query := "SELECT column_name FROM USER_TAB_COLUMNS WHERE table_name = :1" + " AND column_name = :2" return db.HasRecords(queryer, ctx, query, args...) } func (db *oracle) GetColumns(queryer core.Queryer, ctx context.Context, tableName string) ([]string, map[string]*schemas.Column, error) { - args := []interface{}{tableName} + args := []any{tableName} s := "SELECT column_name,data_default,data_type,data_length,data_precision,data_scale," + "nullable FROM USER_TAB_COLUMNS WHERE table_name = :1" @@ -807,7 +808,7 @@ func (db *oracle) GetColumns(queryer core.Queryer, ctx context.Context, tableNam } func (db *oracle) GetTables(queryer core.Queryer, ctx context.Context) ([]*schemas.Table, error) { - args := []interface{}{} + args := []any{} s := "SELECT table_name FROM user_tables" rows, err := queryer.QueryContext(ctx, s, args...) @@ -833,7 +834,7 @@ func (db *oracle) GetTables(queryer core.Queryer, ctx context.Context) ([]*schem } func (db *oracle) GetIndexes(queryer core.Queryer, ctx context.Context, tableName string) (map[string]*schemas.Index, error) { - args := []interface{}{tableName} + args := []any{tableName} s := "SELECT t.column_name,i.uniqueness,i.index_name FROM user_ind_columns t,user_indexes i " + "WHERE t.index_name = i.index_name and t.table_name = i.table_name and t.table_name =:1" @@ -922,7 +923,7 @@ func (g *godrorDriver) Parse(driverName, dataSourceName string) (*URI, error) { return db, nil } -func (g *godrorDriver) GenScanResult(colType string) (interface{}, error) { +func (g *godrorDriver) GenScanResult(colType string) (any, error) { switch colType { case "CHAR", "NCHAR", "VARCHAR", "VARCHAR2", "NVARCHAR2", "LONG", "CLOB", "NCLOB": var s sql.NullString diff --git a/dialects/postgres.go b/dialects/postgres.go index 6e37958..a2ee7d4 100644 --- a/dialects/postgres.go +++ b/dialects/postgres.go @@ -838,14 +838,13 @@ func (db *postgres) Version(ctx context.Context, queryer core.Queryer) (*schemas Level: versions[1], Edition: "PostgreSQL", }, nil - } else { - versions := strings.Split(strings.TrimPrefix(version, "PostgreSQL "), ",") - return &schemas.Version{ - Number: versions[0], - Level: versions[1], - Edition: "PostgreSQL", - }, nil } + versions := strings.Split(strings.TrimPrefix(version, "PostgreSQL "), ",") + return &schemas.Version{ + Number: versions[0], + Level: versions[1], + Edition: "PostgreSQL", + }, nil } else if strings.HasPrefix(version, "KingbaseES") { if strings.Contains(version, " on ") { versions := strings.Split(strings.TrimPrefix(version, "KingbaseES "), " on ") @@ -854,14 +853,13 @@ func (db *postgres) Version(ctx context.Context, queryer core.Queryer) (*schemas Level: versions[1], Edition: "KingbaseES", }, nil - } else { - versions := strings.Split(strings.TrimPrefix(version, "KingbaseES "), ",") - return &schemas.Version{ - Number: versions[0], - Level: versions[1], - Edition: "KingbaseES", - }, nil } + versions := strings.Split(strings.TrimPrefix(version, "KingbaseES "), ",") + return &schemas.Version{ + Number: versions[0], + Level: versions[1], + Edition: "KingbaseES", + }, nil } return nil, errors.New("unknow database version") @@ -998,13 +996,13 @@ func (db *postgres) AutoIncrStr() string { return "" } -func (db *postgres) IndexCheckSQL(tableName, idxName string) (string, []interface{}) { +func (db *postgres) IndexCheckSQL(tableName, idxName string) (string, []any) { if len(db.getSchema()) == 0 { - args := []interface{}{tableName, idxName} + args := []any{tableName, idxName} return `SELECT indexname FROM pg_indexes WHERE tablename = ? AND indexname = ?`, args } - args := []interface{}{db.getSchema(), tableName, idxName} + args := []any{db.getSchema(), tableName, idxName} return `SELECT indexname FROM pg_indexes ` + `WHERE schemaname = ? AND tablename = ? AND indexname = ?`, args } @@ -1022,31 +1020,29 @@ func (db *postgres) AddColumnSQL(tableName string, col *schemas.Column) string { s, _ := ColumnString(db.dialect, col, true, false) quoter := db.dialect.Quoter() - addColumnSQL := "" commentSQL := "; " if len(db.getSchema()) == 0 || strings.Contains(tableName, ".") { - addColumnSQL = fmt.Sprintf("ALTER TABLE %s ADD %s", quoter.Quote(tableName), s) + addColumnSQL := fmt.Sprintf("ALTER TABLE %s ADD %s", quoter.Quote(tableName), s) commentSQL += fmt.Sprintf("COMMENT ON COLUMN %s.%s IS '%s'", quoter.Quote(tableName), quoter.Quote(col.Name), col.Comment) return addColumnSQL + commentSQL } - addColumnSQL = fmt.Sprintf("ALTER TABLE %s.%s ADD %s", quoter.Quote(db.getSchema()), quoter.Quote(tableName), s) + addColumnSQL := fmt.Sprintf("ALTER TABLE %s.%s ADD %s", quoter.Quote(db.getSchema()), quoter.Quote(tableName), s) commentSQL += fmt.Sprintf("COMMENT ON COLUMN %s.%s.%s IS '%s'", quoter.Quote(db.getSchema()), quoter.Quote(tableName), quoter.Quote(col.Name), col.Comment) return addColumnSQL + commentSQL } func (db *postgres) ModifyColumnSQL(tableName string, col *schemas.Column) string { quoter := db.dialect.Quoter() - modifyColumnSQL := "" commentSQL := "; " if len(db.getSchema()) == 0 || strings.Contains(tableName, ".") { - modifyColumnSQL = fmt.Sprintf("ALTER TABLE %s ALTER COLUMN %s TYPE %s", quoter.Quote(tableName), quoter.Quote(col.Name), db.SQLType(col)) + modifyColumnSQL := fmt.Sprintf("ALTER TABLE %s ALTER COLUMN %s TYPE %s", quoter.Quote(tableName), quoter.Quote(col.Name), db.SQLType(col)) commentSQL += fmt.Sprintf("COMMENT ON COLUMN %s.%s IS '%s'", quoter.Quote(tableName), quoter.Quote(col.Name), col.Comment) return modifyColumnSQL + commentSQL } - modifyColumnSQL = fmt.Sprintf("ALTER TABLE %s.%s ALTER COLUMN %s TYPE %s", quoter.Quote(db.getSchema()), quoter.Quote(tableName), quoter.Quote(col.Name), db.SQLType(col)) + modifyColumnSQL := fmt.Sprintf("ALTER TABLE %s.%s ALTER COLUMN %s TYPE %s", quoter.Quote(db.getSchema()), quoter.Quote(tableName), quoter.Quote(col.Name), db.SQLType(col)) commentSQL += fmt.Sprintf("COMMENT ON COLUMN %s.%s.%s IS '%s'", quoter.Quote(db.getSchema()), quoter.Quote(tableName), quoter.Quote(col.Name), col.Comment) return modifyColumnSQL + commentSQL } @@ -1054,7 +1050,7 @@ func (db *postgres) ModifyColumnSQL(tableName string, col *schemas.Column) strin func (db *postgres) DropIndexSQL(tableName string, index *schemas.Index) string { idxName := index.Name - tableParts := strings.Split(strings.Replace(tableName, `"`, "", -1), ".") + tableParts := strings.Split(strings.ReplaceAll(tableName, `"`, ""), ".") tableName = tableParts[len(tableParts)-1] if index.IsRegular { @@ -1071,11 +1067,11 @@ func (db *postgres) DropIndexSQL(tableName string, index *schemas.Index) string } func (db *postgres) IsColumnExist(queryer core.Queryer, ctx context.Context, tableName, colName string) (bool, error) { - args := []interface{}{db.getSchema(), tableName, colName} + args := []any{db.getSchema(), tableName, colName} query := "SELECT column_name FROM INFORMATION_SCHEMA.COLUMNS WHERE table_schema = $1 AND table_name = $2" + " AND column_name = $3" if len(db.getSchema()) == 0 { - args = []interface{}{tableName, colName} + args = []any{tableName, colName} query = "SELECT column_name FROM INFORMATION_SCHEMA.COLUMNS WHERE table_name = $1" + " AND column_name = $2" } @@ -1093,7 +1089,7 @@ func (db *postgres) IsColumnExist(queryer core.Queryer, ctx context.Context, tab } func (db *postgres) GetColumns(queryer core.Queryer, ctx context.Context, tableName string) ([]string, map[string]*schemas.Column, error) { - args := []interface{}{tableName} + args := []any{tableName} s := `SELECT column_name, column_default, is_nullable, data_type, character_maximum_length, description, CASE WHEN p.contype = 'p' THEN true ELSE false END AS primarykey, CASE WHEN p.contype = 'u' THEN true ELSE false END AS uniquekey @@ -1226,8 +1222,8 @@ WHERE n.nspname= s.table_schema AND c.relkind = 'r' AND c.relname = $1%s AND f.a if !col.DefaultIsEmpty { if col.SQLType.IsText() { - if strings.HasSuffix(col.Default, "::character varying") { - col.Default = strings.TrimSuffix(col.Default, "::character varying") + if trimmed, ok := strings.CutSuffix(col.Default, "::character varying"); ok { + col.Default = trimmed } else if !strings.HasPrefix(col.Default, "'") { col.Default = "'" + col.Default + "'" } @@ -1246,12 +1242,12 @@ WHERE n.nspname= s.table_schema AND c.relkind = 'r' AND c.relname = $1%s AND f.a } func (db *postgres) GetTables(queryer core.Queryer, ctx context.Context) ([]*schemas.Table, error) { - args := []interface{}{} + args := []any{} s := "SELECT tablename FROM pg_tables" schema := db.getSchema() if schema != "" { args = append(args, schema) - s = s + " WHERE schemaname = $1" + s += " WHERE schemaname = $1" } rows, err := queryer.QueryContext(ctx, s, args...) @@ -1289,7 +1285,7 @@ func getIndexColName(indexdef string) []string { } func (db *postgres) GetIndexes(queryer core.Queryer, ctx context.Context, tableName string) (map[string]*schemas.Index, error) { - args := []interface{}{tableName} + args := []any{tableName} s := "SELECT indexname, indexdef FROM pg_indexes WHERE tablename=$1" if len(db.getSchema()) != 0 { args = append(args, db.getSchema()) @@ -1343,7 +1339,7 @@ func (db *postgres) GetIndexes(queryer core.Queryer, ctx context.Context, tableN index := &schemas.Index{Name: indexName, Type: indexType, Cols: make([]string, 0)} for _, colName := range colNames { - col := strings.TrimSpace(strings.Replace(colName, `"`, "", -1)) + col := strings.TrimSpace(strings.ReplaceAll(colName, `"`, "")) fields := strings.Split(col, " ") index.Cols = append(index.Cols, fields[0]) } @@ -1367,21 +1363,26 @@ func (db *postgres) CreateTableSQL(ctx context.Context, queryer core.Queryer, ta return "", ok, err } - commentSQL := "; " + var commentSQL strings.Builder + commentSQL.WriteString("; ") if table.Comment != "" { // support schema.table -> "schema"."table" - commentSQL += fmt.Sprintf("COMMENT ON TABLE %s IS '%s'; ", quoter.Quote(tableName), table.Comment) + if _, err := fmt.Fprintf(&commentSQL, "COMMENT ON TABLE %s IS '%s'; ", quoter.Quote(tableName), table.Comment); err != nil { + return "", ok, err + } } for _, colName := range table.ColumnsSeq() { col := table.GetColumn(colName) if len(col.Comment) > 0 { - commentSQL += fmt.Sprintf("COMMENT ON COLUMN %s.%s IS '%s'; ", quoter.Quote(tableName), quoter.Quote(col.Name), col.Comment) + if _, err := fmt.Fprintf(&commentSQL, "COMMENT ON COLUMN %s.%s IS '%s'; ", quoter.Quote(tableName), quoter.Quote(col.Name), col.Comment); err != nil { + return "", ok, err + } } } - return createTableSQL + commentSQL, true, nil + return createTableSQL + commentSQL.String(), true, nil } func (db *postgres) Filters() []Filter { @@ -1460,10 +1461,11 @@ func parseOpts(urlStr string, o values) error { state = 3 } default: - if state == 3 { + switch state { + case 3: state = 2 start = i - } else if state == 1 { + case 1: state = 0 start = i } @@ -1522,7 +1524,7 @@ func (p *pqDriver) Parse(driverName, dataSourceName string) (*URI, error) { return db, nil } -func (p *pqDriver) GenScanResult(colType string) (interface{}, error) { +func (p *pqDriver) GenScanResult(colType string) (any, error) { switch colType { case "VARCHAR", "TEXT": var s sql.NullString diff --git a/dialects/sqlite3.go b/dialects/sqlite3.go index ece360e..37b2d80 100644 --- a/dialects/sqlite3.go +++ b/dialects/sqlite3.go @@ -211,9 +211,10 @@ func (db *sqlite3) SetQuotePolicy(quotePolicy QuotePolicy) { func (db *sqlite3) SQLType(c *schemas.Column) string { switch t := c.SQLType.Name; t { case schemas.Bool: - if c.Default == "true" { + switch c.Default { + case "true": c.Default = "1" - } else if c.Default == "false" { + case "false": c.Default = "0" } return schemas.Integer @@ -268,8 +269,8 @@ func (db *sqlite3) AutoIncrStr() string { return "AUTOINCREMENT" } -func (db *sqlite3) IndexCheckSQL(tableName, idxName string) (string, []interface{}) { - args := []interface{}{idxName} +func (db *sqlite3) IndexCheckSQL(tableName, idxName string) (string, []any) { + args := []any{idxName} return "SELECT name FROM sqlite_master WHERE type='index' and name = ?", args } @@ -342,7 +343,7 @@ func splitColStr(colStr string) []string { return results } -func parseString(colStr string) (*schemas.Column, error) { +func parseString(colStr string) *schemas.Column { fields := splitColStr(colStr) col := new(schemas.Column) col.Indexes = make(map[string]int) @@ -373,11 +374,11 @@ func parseString(colStr string) (*schemas.Column, error) { col.DefaultIsEmpty = false } } - return col, nil + return col } func (db *sqlite3) GetColumns(queryer core.Queryer, ctx context.Context, tableName string) ([]string, map[string]*schemas.Column, error) { - args := []interface{}{tableName} + args := []any{tableName} s := "SELECT sql FROM sqlite_master WHERE type='table' and name = ?" rows, err := queryer.QueryContext(ctx, s, args...) @@ -426,11 +427,7 @@ func (db *sqlite3) GetColumns(queryer core.Queryer, ctx context.Context, tableNa continue } - col, err := parseString(colStr) - if err != nil { - return colSeq, cols, err - } - + col := parseString(colStr) cols[col.Name] = col colSeq = append(colSeq, col.Name) } @@ -438,7 +435,7 @@ func (db *sqlite3) GetColumns(queryer core.Queryer, ctx context.Context, tableNa } func (db *sqlite3) GetTables(queryer core.Queryer, ctx context.Context) ([]*schemas.Table, error) { - args := []interface{}{} + args := []any{} s := "SELECT name FROM sqlite_master WHERE type='table'" rows, err := queryer.QueryContext(ctx, s, args...) @@ -466,7 +463,7 @@ func (db *sqlite3) GetTables(queryer core.Queryer, ctx context.Context) ([]*sche } func (db *sqlite3) GetIndexes(queryer core.Queryer, ctx context.Context, tableName string) (map[string]*schemas.Index, error) { - args := []interface{}{tableName} + args := []any{tableName} s := "SELECT sql FROM sqlite_master WHERE type='index' and tbl_name = ?" rows, err := queryer.QueryContext(ctx, s, args...) @@ -543,14 +540,14 @@ func (p *sqlite3Driver) Features() *DriverFeatures { } func (p *sqlite3Driver) Parse(driverName, dataSourceName string) (*URI, error) { - if strings.Contains(dataSourceName, "?") { - dataSourceName = dataSourceName[:strings.Index(dataSourceName, "?")] + if base, _, ok := strings.Cut(dataSourceName, "?"); ok { + dataSourceName = base } return &URI{DBType: schemas.SQLITE, DBName: dataSourceName}, nil } -func (p *sqlite3Driver) GenScanResult(colType string) (interface{}, error) { +func (p *sqlite3Driver) GenScanResult(colType string) (any, error) { switch colType { case "TEXT": var s sql.NullString diff --git a/dialects/table_name.go b/dialects/table_name.go index 8a0baea..fef4b74 100644 --- a/dialects/table_name.go +++ b/dialects/table_name.go @@ -25,7 +25,7 @@ func TableNameWithSchema(dialect Dialect, tableName string) string { } // TableNameNoSchema returns table name with given tableName -func TableNameNoSchema(dialect Dialect, mapper names.Mapper, tableName interface{}) string { +func TableNameNoSchema(dialect Dialect, mapper names.Mapper, tableName any) string { quote := dialect.Quoter().Quote switch tt := tableName.(type) { case []string: @@ -37,16 +37,15 @@ func TableNameNoSchema(dialect Dialect, mapper names.Mapper, tableName interface } else if len(tt) == 1 { return quote(tt[0]) } - case []interface{}: + case []any: l := len(tt) var table string if l > 0 { - f := tt[0] - switch f.(type) { + switch f := tt[0].(type) { case string: - table = f.(string) + table = f case names.TableName: - table = f.(names.TableName).TableName() + table = f.TableName() default: v := utils.ReflectValue(f) t := v.Type() @@ -66,11 +65,11 @@ func TableNameNoSchema(dialect Dialect, mapper names.Mapper, tableName interface return quote(table) } case names.TableName: - return tableName.(names.TableName).TableName() + return tt.TableName() case string: - return tableName.(string) + return tt case reflect.Value: - v := tableName.(reflect.Value) + v := tt return names.GetTableName(mapper, v) default: v := utils.ReflectValue(tableName) @@ -84,7 +83,7 @@ func TableNameNoSchema(dialect Dialect, mapper names.Mapper, tableName interface } // FullTableName returns table name with quote and schema according parameter -func FullTableName(dialect Dialect, mapper names.Mapper, bean interface{}, includeSchema ...bool) string { +func FullTableName(dialect Dialect, mapper names.Mapper, bean any, includeSchema ...bool) string { tbName := TableNameNoSchema(dialect, mapper, bean) if len(includeSchema) > 0 && includeSchema[0] && !utils.IsSubQuery(tbName) { tbName = TableNameWithSchema(dialect, tbName) diff --git a/dialects/time.go b/dialects/time.go index 4a6beb7..e7d3f81 100644 --- a/dialects/time.go +++ b/dialects/time.go @@ -7,13 +7,13 @@ package dialects import ( "strings" "time" - "xorm.io/xorm/internal/utils" + "xorm.io/xorm/internal/utils" "xorm.io/xorm/schemas" ) // FormatColumnTime format column time -func FormatColumnTime(dialect Dialect, dbLocation *time.Location, col *schemas.Column, t time.Time) (interface{}, error) { +func FormatColumnTime(dialect Dialect, dbLocation *time.Location, col *schemas.Column, t time.Time) (any, error) { if utils.IsTimeZero(t) { if col.Nullable { return nil, nil @@ -55,9 +55,8 @@ func FormatColumnTime(dialect Dialect, dbLocation *time.Location, col *schemas.C case schemas.TimeStampz: if dialect.URI().DBType == schemas.MSSQL { return t.Format("2006-01-02T15:04:05.9999999Z07:00"), nil - } else { - return t.Format(time.RFC3339Nano), nil } + return t.Format(time.RFC3339Nano), nil case schemas.BigInt, schemas.Int: return t.Unix(), nil default: diff --git a/dialects/time_test.go b/dialects/time_test.go index 670207c..4eb4216 100644 --- a/dialects/time_test.go +++ b/dialects/time_test.go @@ -31,7 +31,7 @@ func TestFormatColumnTime(t *testing.T) { location *time.Location column *schemas.Column time time.Time - wantRes interface{} + wantRes any wantErr error }{ { diff --git a/doc.go b/doc.go index f88f537..6e22fe9 100644 --- a/doc.go +++ b/doc.go @@ -32,7 +32,7 @@ XORM supports raw SQL execution: results, err := engine.QueryString("select * from user") -3. query with a SQL string, the returned results is []map[string]interface{} +3. query with a SQL string, the returned results is []map[string]any results, err := engine.QueryInterface("select * from user") @@ -88,7 +88,7 @@ There are 8 major ORM methods and many helpful methods to use to operate databas 4. Query multiple records and record by record handle, there two methods, one is Iterate, another is Rows - err := engine.Iterate(new(User), func(i int, bean interface{}) error { + err := engine.Iterate(new(User), func(i int, bean any) error { // do something }) // SELECT * FROM user diff --git a/engine.go b/engine.go index 57c7073..73b10ad 100644 --- a/engine.go +++ b/engine.go @@ -51,7 +51,7 @@ type Engine struct { // NewEngine new a db manager according to the parameter. Currently support four // drivers -func NewEngine(driverName string, dataSourceName string, driverOptions ...func(db *sql.DB) error) (*Engine, error) { +func NewEngine(driverName, dataSourceName string, driverOptions ...func(db *sql.DB) error) (*Engine, error) { dialect, err := dialects.OpenDialect(driverName, dataSourceName) if err != nil { return nil, err @@ -107,7 +107,7 @@ func newEngine(driverName, dataSourceName string, dialect dialects.Dialect, db * } // NewEngineWithParams new a db manager with params. The params will be passed to dialects. -func NewEngineWithParams(driverName string, dataSourceName string, params map[string]string) (*Engine, error) { +func NewEngineWithParams(driverName, dataSourceName string, params map[string]string) (*Engine, error) { engine, err := NewEngine(driverName, dataSourceName) if err != nil { return nil, err @@ -117,7 +117,7 @@ func NewEngineWithParams(driverName string, dataSourceName string, params map[st } // NewEngineWithDB new a db manager with db. The params will be passed to db. -func NewEngineWithDB(driverName string, dataSourceName string, db *core.DB) (*Engine, error) { +func NewEngineWithDB(driverName, dataSourceName string, db *core.DB) (*Engine, error) { dialect, err := dialects.OpenDialect(driverName, dataSourceName) if err != nil { return nil, err @@ -172,7 +172,7 @@ func (engine *Engine) Logger() log.ContextLogger { } // SetLogger set the new logger -func (engine *Engine) SetLogger(logger interface{}) { +func (engine *Engine) SetLogger(logger any) { var realLogger log.ContextLogger switch t := logger.(type) { case log.ContextLogger: @@ -250,7 +250,7 @@ func (engine *Engine) QuoteTo(buf *strings.Builder, value string) { if value == "" { return } - engine.dialect.Quoter().QuoteTo(buf, value) + buf.WriteString(engine.dialect.Quoter().Quote(value)) } // SQLType A simple wrapper to dialect's core.SqlType method @@ -304,7 +304,7 @@ func (engine *Engine) NoCascade() *Session { } // MapCacher Set a table use a special cacher -func (engine *Engine) MapCacher(bean interface{}, cacher caches.Cacher) error { +func (engine *Engine) MapCacher(bean any, cacher caches.Cacher) error { engine.SetCacher(dialects.FullTableName(engine.dialect, engine.GetTableMapper(), bean, true), cacher) return nil } @@ -347,7 +347,7 @@ func (engine *Engine) Ping() error { // engine.SQL("select * from user").Find(&users) // // This code will execute "select * from user" and set the records to users -func (engine *Engine) SQL(query interface{}, args ...interface{}) *Session { +func (engine *Engine) SQL(query any, args ...any) *Session { session := engine.NewSession() session.isAutoClose = true return session.SQL(query, args...) @@ -388,9 +388,10 @@ func (engine *Engine) loadTableInfo(ctx context.Context, table *schemas.Table) e for _, name := range index.Cols { parts := strings.Split(strings.TrimSpace(name), " ") if len(parts) > 1 { - if parts[1] == "DESC" { + switch parts[1] { + case "DESC": seq = 1 - } else if parts[1] == "ASC" { + case "ASC": seq = 0 } } @@ -466,6 +467,42 @@ func formatBool(s bool, dstDialect dialects.Dialect) string { var controlCharactersRe = regexp.MustCompile(`[\x00-\x1f\x7f]+`) +func writeConcatWithChar(w io.Writer, value string) error { + if _, err := io.WriteString(w, "CONCAT("); err != nil { + return err + } + toCheck := strings.ReplaceAll(value, "'", "''") + for len(toCheck) > 0 { + loc := controlCharactersRe.FindStringIndex(toCheck) + if loc == nil { + _, err := io.WriteString(w, "'"+toCheck+"')") + return err + } + if loc[0] > 0 { + if _, err := io.WriteString(w, "'"+toCheck[:loc[0]]+"', "); err != nil { + return err + } + } + for i := loc[0]; i < loc[1]-1; i++ { + if _, err := io.WriteString(w, "CHAR("+strconv.Itoa(int(toCheck[i]))+"), "); err != nil { + return err + } + } + char := toCheck[loc[1]-1] + toCheck = toCheck[loc[1]:] + if len(toCheck) > 0 { + if _, err := io.WriteString(w, "CHAR("+strconv.Itoa(int(char))+"), "); err != nil { + return err + } + } else { + if _, err := io.WriteString(w, "CHAR("+strconv.Itoa(int(char))+"))"); err != nil { + return err + } + } + } + return nil +} + // dumpTables dump database all table structs and data to w with specify db type func (engine *Engine) dumpTables(ctx context.Context, tables []*schemas.Table, w io.Writer, tp ...schemas.DBType) error { var dstDialect dialects.Dialect @@ -678,40 +715,9 @@ func (engine *Engine) dumpTables(ctx context.Context, tables []*schemas.Table, w return err } } else { - if _, err := io.WriteString(w, "CONCAT("); err != nil { + if err := writeConcatWithChar(w, s.String); err != nil { return err } - toCheck := strings.ReplaceAll(s.String, "'", "''") - for len(toCheck) > 0 { - loc := controlCharactersRe.FindStringIndex(toCheck) - if loc == nil { - if _, err := io.WriteString(w, "'"+toCheck+"')"); err != nil { - return err - } - break - } - if loc[0] > 0 { - if _, err := io.WriteString(w, "'"+toCheck[:loc[0]]+"', "); err != nil { - return err - } - } - for i := loc[0]; i < loc[1]-1; i++ { - if _, err := io.WriteString(w, "CHAR("+strconv.Itoa(int(toCheck[i]))+"), "); err != nil { - return err - } - } - char := toCheck[loc[1]-1] - toCheck = toCheck[loc[1]:] - if len(toCheck) > 0 { - if _, err := io.WriteString(w, "CHAR("+strconv.Itoa(int(char))+"), "); err != nil { - return err - } - } else { - if _, err = io.WriteString(w, "CHAR("+strconv.Itoa(int(char))+"))"); err != nil { - return err - } - } - } } } else if dstDialect.URI().DBType == schemas.SQLITE { if dstTable.Columns()[i].SQLType.IsBlob() { @@ -755,40 +761,9 @@ func (engine *Engine) dumpTables(ctx context.Context, tables []*schemas.Table, w } else { // ORACLE/DAMENG concatentates strings in multiple ways but uses CHAR and has CONCAT // (NOTE: a NUL byte in a text segment will fail) - if _, err := io.WriteString(w, "CONCAT("); err != nil { + if err := writeConcatWithChar(w, s.String); err != nil { return err } - toCheck := strings.ReplaceAll(s.String, "'", "''") - for len(toCheck) > 0 { - loc := controlCharactersRe.FindStringIndex(toCheck) - if loc == nil { - if _, err := io.WriteString(w, "'"+toCheck+"')"); err != nil { - return err - } - break - } - if loc[0] > 0 { - if _, err := io.WriteString(w, "'"+toCheck[:loc[0]]+"', "); err != nil { - return err - } - } - for i := loc[0]; i < loc[1]-1; i++ { - if _, err := io.WriteString(w, "CHAR("+strconv.Itoa(int(toCheck[i]))+"), "); err != nil { - return err - } - } - char := toCheck[loc[1]-1] - toCheck = toCheck[loc[1]:] - if len(toCheck) > 0 { - if _, err := io.WriteString(w, "CHAR("+strconv.Itoa(int(char))+"), "); err != nil { - return err - } - } else { - if _, err = io.WriteString(w, "CHAR("+strconv.Itoa(int(char))+"))"); err != nil { - return err - } - } - } } } else if dstDialect.URI().DBType == schemas.MSSQL { if dstTable.Columns()[i].SQLType.IsBlob() { @@ -860,28 +835,28 @@ func (engine *Engine) Cascade(trueOrFalse ...bool) *Session { } // Where method provide a condition query -func (engine *Engine) Where(query interface{}, args ...interface{}) *Session { +func (engine *Engine) Where(query any, args ...any) *Session { session := engine.NewSession() session.isAutoClose = true return session.Where(query, args...) } // ID method provoide a condition as (id) = ? -func (engine *Engine) ID(id interface{}) *Session { +func (engine *Engine) ID(id any) *Session { session := engine.NewSession() session.isAutoClose = true return session.ID(id) } // Before apply before Processor, affected bean is passed to closure arg -func (engine *Engine) Before(closures func(interface{})) *Session { +func (engine *Engine) Before(closures func(any)) *Session { session := engine.NewSession() session.isAutoClose = true return session.Before(closures) } // After apply after insert Processor, affected bean is passed to closure arg -func (engine *Engine) After(closures func(interface{})) *Session { +func (engine *Engine) After(closures func(any)) *Session { session := engine.NewSession() session.isAutoClose = true return session.After(closures) @@ -964,42 +939,42 @@ func (engine *Engine) Nullable(columns ...string) *Session { } // In will generate "column IN (?, ?)" -func (engine *Engine) In(column string, args ...interface{}) *Session { +func (engine *Engine) In(column string, args ...any) *Session { session := engine.NewSession() session.isAutoClose = true return session.In(column, args...) } // NotIn will generate "column NOT IN (?, ?)" -func (engine *Engine) NotIn(column string, args ...interface{}) *Session { +func (engine *Engine) NotIn(column string, args ...any) *Session { session := engine.NewSession() session.isAutoClose = true return session.NotIn(column, args...) } // Incr provides a update string like "column = column + ?" -func (engine *Engine) Incr(column string, arg ...interface{}) *Session { +func (engine *Engine) Incr(column string, arg ...any) *Session { session := engine.NewSession() session.isAutoClose = true return session.Incr(column, arg...) } // Decr provides a update string like "column = column - ?" -func (engine *Engine) Decr(column string, arg ...interface{}) *Session { +func (engine *Engine) Decr(column string, arg ...any) *Session { session := engine.NewSession() session.isAutoClose = true return session.Decr(column, arg...) } // SetExpr provides a update string like "column = {expression}" -func (engine *Engine) SetExpr(column string, expression interface{}) *Session { +func (engine *Engine) SetExpr(column string, expression any) *Session { session := engine.NewSession() session.isAutoClose = true return session.SetExpr(column, expression) } // Table temporarily change the Get, Find, Update's table -func (engine *Engine) Table(tableNameOrBean interface{}) *Session { +func (engine *Engine) Table(tableNameOrBean any) *Session { session := engine.NewSession() session.isAutoClose = true return session.Table(tableNameOrBean) @@ -1038,7 +1013,7 @@ func (engine *Engine) Asc(colNames ...string) *Session { } // OrderBy will generate "ORDER BY order" -func (engine *Engine) OrderBy(order interface{}, args ...interface{}) *Session { +func (engine *Engine) OrderBy(order any, args ...any) *Session { session := engine.NewSession() session.isAutoClose = true return session.OrderBy(order, args...) @@ -1052,7 +1027,7 @@ func (engine *Engine) Prepare() *Session { } // Join the join_operator should be one of INNER, LEFT OUTER, CROSS etc - this will be prepended to JOIN -func (engine *Engine) Join(joinOperator string, tablename interface{}, condition interface{}, args ...interface{}) *Session { +func (engine *Engine) Join(joinOperator string, tablename, condition any, args ...any) *Session { session := engine.NewSession() session.isAutoClose = true return session.Join(joinOperator, tablename, condition, args...) @@ -1078,46 +1053,46 @@ func (engine *Engine) DBVersion() (*schemas.Version, error) { } // TableInfo get table info according to bean's content -func (engine *Engine) TableInfo(bean interface{}) (*schemas.Table, error) { +func (engine *Engine) TableInfo(bean any) (*schemas.Table, error) { v := utils.ReflectValue(bean) return engine.tagParser.ParseWithCache(v) } // IsTableEmpty if a table has any reocrd -func (engine *Engine) IsTableEmpty(bean interface{}) (bool, error) { +func (engine *Engine) IsTableEmpty(bean any) (bool, error) { session := engine.NewSession() defer session.Close() return session.IsTableEmpty(bean) } // IsTableExist if a table is exist -func (engine *Engine) IsTableExist(beanOrTableName interface{}) (bool, error) { +func (engine *Engine) IsTableExist(beanOrTableName any) (bool, error) { session := engine.NewSession() defer session.Close() return session.IsTableExist(beanOrTableName) } // TableName returns table name with schema prefix if has -func (engine *Engine) TableName(bean interface{}, includeSchema ...bool) string { +func (engine *Engine) TableName(bean any, includeSchema ...bool) string { return dialects.FullTableName(engine.dialect, engine.GetTableMapper(), bean, includeSchema...) } // CreateIndexes create indexes -func (engine *Engine) CreateIndexes(bean interface{}) error { +func (engine *Engine) CreateIndexes(bean any) error { session := engine.NewSession() defer session.Close() return session.CreateIndexes(bean) } // CreateUniques create uniques -func (engine *Engine) CreateUniques(bean interface{}) error { +func (engine *Engine) CreateUniques(bean any) error { session := engine.NewSession() defer session.Close() return session.CreateUniques(bean) } // ClearCacheBean if enabled cache, clear the cache bean -func (engine *Engine) ClearCacheBean(bean interface{}, id string) error { +func (engine *Engine) ClearCacheBean(bean any, id string) error { tableName := dialects.FullTableName(engine.dialect, engine.GetTableMapper(), bean) cacher := engine.GetCacher(tableName) if cacher != nil { @@ -1128,7 +1103,7 @@ func (engine *Engine) ClearCacheBean(bean interface{}, id string) error { } // ClearCache if enabled cache, clear some tables' cache -func (engine *Engine) ClearCache(beans ...interface{}) error { +func (engine *Engine) ClearCache(beans ...any) error { for _, bean := range beans { tableName := dialects.FullTableName(engine.dialect, engine.GetTableMapper(), bean) cacher := engine.GetCacher(tableName) @@ -1146,7 +1121,7 @@ func (engine *Engine) UnMapType(t reflect.Type) { } // CreateTables create tabls according bean -func (engine *Engine) CreateTables(beans ...interface{}) error { +func (engine *Engine) CreateTables(beans ...any) error { session := engine.NewSession() defer session.Close() @@ -1166,7 +1141,7 @@ func (engine *Engine) CreateTables(beans ...interface{}) error { } // DropTables drop specify tables -func (engine *Engine) DropTables(beans ...interface{}) error { +func (engine *Engine) DropTables(beans ...any) error { session := engine.NewSession() defer session.Close() @@ -1186,49 +1161,49 @@ func (engine *Engine) DropTables(beans ...interface{}) error { } // DropIndexes drop indexes of a table -func (engine *Engine) DropIndexes(bean interface{}) error { +func (engine *Engine) DropIndexes(bean any) error { session := engine.NewSession() defer session.Close() return session.DropIndexes(bean) } // Exec raw sql -func (engine *Engine) Exec(sqlOrArgs ...interface{}) (sql.Result, error) { +func (engine *Engine) Exec(sqlOrArgs ...any) (sql.Result, error) { session := engine.NewSession() defer session.Close() return session.Exec(sqlOrArgs...) } // Query a raw sql and return records as []map[string][]byte -func (engine *Engine) Query(sqlOrArgs ...interface{}) (resultsSlice []map[string][]byte, err error) { +func (engine *Engine) Query(sqlOrArgs ...any) (resultsSlice []map[string][]byte, err error) { session := engine.NewSession() defer session.Close() return session.Query(sqlOrArgs...) } // QueryString runs a raw sql and return records as []map[string]string -func (engine *Engine) QueryString(sqlOrArgs ...interface{}) ([]map[string]string, error) { +func (engine *Engine) QueryString(sqlOrArgs ...any) ([]map[string]string, error) { session := engine.NewSession() defer session.Close() return session.QueryString(sqlOrArgs...) } -// QueryInterface runs a raw sql and return records as []map[string]interface{} -func (engine *Engine) QueryInterface(sqlOrArgs ...interface{}) ([]map[string]interface{}, error) { +// QueryInterface runs a raw sql and return records as []map[string]any +func (engine *Engine) QueryInterface(sqlOrArgs ...any) ([]map[string]any, error) { session := engine.NewSession() defer session.Close() return session.QueryInterface(sqlOrArgs...) } // Insert one or more records -func (engine *Engine) Insert(beans ...interface{}) (int64, error) { +func (engine *Engine) Insert(beans ...any) (int64, error) { session := engine.NewSession() defer session.Close() return session.Insert(beans...) } // InsertOne insert only one record -func (engine *Engine) InsertOne(bean interface{}) (int64, error) { +func (engine *Engine) InsertOne(bean any) (int64, error) { session := engine.NewSession() defer session.Close() return session.Insert(bean) @@ -1241,7 +1216,7 @@ func (engine *Engine) InsertOne(bean interface{}) (int64, error) { // 1.bool will defaultly be updated content nor conditions // You should call UseBool if you have bool to use. // 2.float32 & float64 may be not inexact as conditions -func (engine *Engine) Update(bean interface{}, condiBeans ...interface{}) (int64, error) { +func (engine *Engine) Update(bean any, condiBeans ...any) (int64, error) { session := engine.NewSession() defer session.Close() return session.Update(bean, condiBeans...) @@ -1249,7 +1224,7 @@ func (engine *Engine) Update(bean interface{}, condiBeans ...interface{}) (int64 // Delete records, bean's non-empty fields are conditions // At least one condition must be set. -func (engine *Engine) Delete(beans ...interface{}) (int64, error) { +func (engine *Engine) Delete(beans ...any) (int64, error) { session := engine.NewSession() defer session.Close() return session.Delete(beans...) @@ -1257,7 +1232,7 @@ func (engine *Engine) Delete(beans ...interface{}) (int64, error) { // Truncate records, bean's non-empty fields are conditions // In contrast to Delete this method allows deletes without conditions. -func (engine *Engine) Truncate(beans ...interface{}) (int64, error) { +func (engine *Engine) Truncate(beans ...any) (int64, error) { session := engine.NewSession() defer session.Close() return session.Truncate(beans...) @@ -1265,14 +1240,14 @@ func (engine *Engine) Truncate(beans ...interface{}) (int64, error) { // Get retrieve one record from table, bean's non-empty fields // are conditions -func (engine *Engine) Get(beans ...interface{}) (bool, error) { +func (engine *Engine) Get(beans ...any) (bool, error) { session := engine.NewSession() defer session.Close() return session.Get(beans...) } // Exist returns true if the record exist otherwise return false -func (engine *Engine) Exist(bean ...interface{}) (bool, error) { +func (engine *Engine) Exist(bean ...any) (bool, error) { session := engine.NewSession() defer session.Close() return session.Exist(bean...) @@ -1281,14 +1256,14 @@ func (engine *Engine) Exist(bean ...interface{}) (bool, error) { // Find retrieve records from table, condiBeans's non-empty fields // are conditions. beans could be []Struct, []*Struct, map[int64]Struct // map[int64]*Struct -func (engine *Engine) Find(beans interface{}, condiBeans ...interface{}) error { +func (engine *Engine) Find(beans any, condiBeans ...any) error { session := engine.NewSession() defer session.Close() return session.Find(beans, condiBeans...) } // FindAndCount find the results and also return the counts -func (engine *Engine) FindAndCount(rowsSlicePtr interface{}, condiBean ...interface{}) (int64, error) { +func (engine *Engine) FindAndCount(rowsSlicePtr any, condiBean ...any) (int64, error) { session := engine.NewSession() defer session.Close() return session.FindAndCount(rowsSlicePtr, condiBean...) @@ -1296,7 +1271,7 @@ func (engine *Engine) FindAndCount(rowsSlicePtr interface{}, condiBean ...interf // Iterate record by record handle records from table, bean's non-empty fields // are conditions. -func (engine *Engine) Iterate(bean interface{}, fun IterFunc) error { +func (engine *Engine) Iterate(bean any, fun IterFunc) error { session := engine.NewSession() defer session.Close() return session.Iterate(bean, fun) @@ -1304,41 +1279,41 @@ func (engine *Engine) Iterate(bean interface{}, fun IterFunc) error { // Rows return sql.Rows compatible Rows obj, as a forward Iterator object for iterating record by record, bean's non-empty fields // are conditions. -func (engine *Engine) Rows(bean interface{}) (*Rows, error) { +func (engine *Engine) Rows(bean any) (*Rows, error) { session := engine.NewSession() return session.Rows(bean) } // Count counts the records. bean's non-empty fields are conditions. -func (engine *Engine) Count(bean ...interface{}) (int64, error) { +func (engine *Engine) Count(bean ...any) (int64, error) { session := engine.NewSession() defer session.Close() return session.Count(bean...) } // Sum sum the records by some column. bean's non-empty fields are conditions. -func (engine *Engine) Sum(bean interface{}, colName string) (float64, error) { +func (engine *Engine) Sum(bean any, colName string) (float64, error) { session := engine.NewSession() defer session.Close() return session.Sum(bean, colName) } // SumInt sum the records by some column. bean's non-empty fields are conditions. -func (engine *Engine) SumInt(bean interface{}, colName string) (int64, error) { +func (engine *Engine) SumInt(bean any, colName string) (int64, error) { session := engine.NewSession() defer session.Close() return session.SumInt(bean, colName) } // Sums sum the records by some columns. bean's non-empty fields are conditions. -func (engine *Engine) Sums(bean interface{}, colNames ...string) ([]float64, error) { +func (engine *Engine) Sums(bean any, colNames ...string) ([]float64, error) { session := engine.NewSession() defer session.Close() return session.Sums(bean, colNames...) } // SumsInt like Sums but return slice of int64 instead of float64. -func (engine *Engine) SumsInt(bean interface{}, colNames ...string) ([]int64, error) { +func (engine *Engine) SumsInt(bean any, colNames ...string) ([]int64, error) { session := engine.NewSession() defer session.Close() return session.SumsInt(bean, colNames...) @@ -1359,7 +1334,7 @@ func (engine *Engine) Import(r io.Reader) ([]sql.Result, error) { } // nowTime return current time -func (engine *Engine) nowTime(col *schemas.Column) (interface{}, time.Time, error) { +func (engine *Engine) nowTime(col *schemas.Column) (any, time.Time, error) { t := time.Now() result, err := dialects.FormatColumnTime(engine.dialect, engine.DatabaseTZ, col, t) if err != nil { @@ -1439,7 +1414,7 @@ func (engine *Engine) PingContext(ctx context.Context) error { } // Transaction Execute sql wrapped in a transaction(abbr as tx), tx will automatic commit if no errors occurred -func (engine *Engine) Transaction(f func(*Session) (interface{}, error)) (interface{}, error) { +func (engine *Engine) Transaction(f func(*Session) (any, error)) (any, error) { session := engine.NewSession() defer session.Close() diff --git a/engine_group.go b/engine_group.go index f2fe913..6d8b300 100644 --- a/engine_group.go +++ b/engine_group.go @@ -23,7 +23,7 @@ type EngineGroup struct { } // NewEngineGroup creates a new engine group -func NewEngineGroup(args1 interface{}, args2 interface{}, policies ...GroupPolicy) (*EngineGroup, error) { +func NewEngineGroup(args1, args2 any, policies ...GroupPolicy) (*EngineGroup, error) { var eg EngineGroup if len(policies) > 0 { eg.policy = policies[0] @@ -137,7 +137,7 @@ func (eg *EngineGroup) SetDefaultCacher(cacher caches.Cacher) { } // SetLogger set the new logger -func (eg *EngineGroup) SetLogger(logger interface{}) { +func (eg *EngineGroup) SetLogger(logger any) { eg.Engine.SetLogger(logger) for i := 0; i < len(eg.slaves); i++ { eg.slaves[i].SetLogger(logger) @@ -239,28 +239,28 @@ func (eg *EngineGroup) Slaves() []*Engine { } // Query execcute a select SQL and return the result -func (eg *EngineGroup) Query(sqlOrArgs ...interface{}) (resultsSlice []map[string][]byte, err error) { +func (eg *EngineGroup) Query(sqlOrArgs ...any) (resultsSlice []map[string][]byte, err error) { sess := eg.NewSession() sess.isAutoClose = true return sess.Query(sqlOrArgs...) } // QueryInterface execcute a select SQL and return the result -func (eg *EngineGroup) QueryInterface(sqlOrArgs ...interface{}) ([]map[string]interface{}, error) { +func (eg *EngineGroup) QueryInterface(sqlOrArgs ...any) ([]map[string]any, error) { sess := eg.NewSession() sess.isAutoClose = true return sess.QueryInterface(sqlOrArgs...) } // QueryString execcute a select SQL and return the result -func (eg *EngineGroup) QueryString(sqlOrArgs ...interface{}) ([]map[string]string, error) { +func (eg *EngineGroup) QueryString(sqlOrArgs ...any) ([]map[string]string, error) { sess := eg.NewSession() sess.isAutoClose = true return sess.QueryString(sqlOrArgs...) } // Rows execcute a select SQL and return the result -func (eg *EngineGroup) Rows(bean interface{}) (*Rows, error) { +func (eg *EngineGroup) Rows(bean any) (*Rows, error) { sess := eg.NewSession() sess.isAutoClose = true return sess.Rows(bean) diff --git a/engine_group_policy.go b/engine_group_policy.go index 1def8ce..c0a9010 100644 --- a/engine_group_policy.go +++ b/engine_group_policy.go @@ -25,7 +25,7 @@ func (h GroupPolicyHandler) Slave(eg *EngineGroup) *Engine { // RandomPolicy implmentes randomly chose the slave of slaves func RandomPolicy() GroupPolicyHandler { - var r = rand.New(rand.NewSource(time.Now().UnixNano())) + r := rand.New(rand.NewSource(time.Now().UnixNano())) return func(g *EngineGroup) *Engine { return g.Slaves()[r.Intn(len(g.Slaves()))] } @@ -33,16 +33,16 @@ func RandomPolicy() GroupPolicyHandler { // WeightRandomPolicy implmentes randomly chose the slave of slaves func WeightRandomPolicy(weights []int) GroupPolicyHandler { - var rands = make([]int, 0, len(weights)) + rands := make([]int, 0, len(weights)) for i := 0; i < len(weights); i++ { for n := 0; n < weights[i]; n++ { rands = append(rands, i) } } - var r = rand.New(rand.NewSource(time.Now().UnixNano())) + r := rand.New(rand.NewSource(time.Now().UnixNano())) return func(g *EngineGroup) *Engine { - var slaves = g.Slaves() + slaves := g.Slaves() idx := rands[r.Intn(len(rands))] if idx >= len(slaves) { idx = len(slaves) - 1 @@ -53,10 +53,10 @@ func WeightRandomPolicy(weights []int) GroupPolicyHandler { // RoundRobinPolicy returns a group policy handler func RoundRobinPolicy() GroupPolicyHandler { - var pos = -1 + pos := -1 var lock sync.Mutex return func(g *EngineGroup) *Engine { - var slaves = g.Slaves() + slaves := g.Slaves() lock.Lock() defer lock.Unlock() @@ -71,17 +71,17 @@ func RoundRobinPolicy() GroupPolicyHandler { // WeightRoundRobinPolicy returns a group policy handler func WeightRoundRobinPolicy(weights []int) GroupPolicyHandler { - var rands = make([]int, 0, len(weights)) + rands := make([]int, 0, len(weights)) for i := 0; i < len(weights); i++ { for n := 0; n < weights[i]; n++ { rands = append(rands, i) } } - var pos = -1 + pos := -1 var lock sync.Mutex return func(g *EngineGroup) *Engine { - var slaves = g.Slaves() + slaves := g.Slaves() lock.Lock() defer lock.Unlock() pos++ @@ -100,7 +100,7 @@ func WeightRoundRobinPolicy(weights []int) GroupPolicyHandler { // LeastConnPolicy implements GroupPolicy, every time will get the least connections slave func LeastConnPolicy() GroupPolicyHandler { return func(g *EngineGroup) *Engine { - var slaves = g.Slaves() + slaves := g.Slaves() connections := 0 idx := 0 for i := 0; i < len(slaves); i++ { diff --git a/interface.go b/interface.go index 03dfd23..5d05451 100644 --- a/interface.go +++ b/interface.go @@ -25,70 +25,70 @@ type Interface interface { Asc(colNames ...string) *Session BufferSize(size int) *Session Cols(columns ...string) *Session - Count(...interface{}) (int64, error) - CreateIndexes(bean interface{}) error - CreateUniques(bean interface{}) error - Decr(column string, arg ...interface{}) *Session + Count(...any) (int64, error) + CreateIndexes(bean any) error + CreateUniques(bean any) error + Decr(column string, arg ...any) *Session Desc(...string) *Session - Delete(...interface{}) (int64, error) - Truncate(...interface{}) (int64, error) + Delete(...any) (int64, error) + Truncate(...any) (int64, error) Distinct(columns ...string) *Session - DropIndexes(bean interface{}) error - Exec(sqlOrArgs ...interface{}) (sql.Result, error) - Exist(bean ...interface{}) (bool, error) - Find(interface{}, ...interface{}) error - FindAndCount(interface{}, ...interface{}) (int64, error) - Get(...interface{}) (bool, error) + DropIndexes(bean any) error + Exec(sqlOrArgs ...any) (sql.Result, error) + Exist(bean ...any) (bool, error) + Find(any, ...any) error + FindAndCount(any, ...any) (int64, error) + Get(...any) (bool, error) GroupBy(keys string) *Session - ID(interface{}) *Session - In(string, ...interface{}) *Session - Incr(column string, arg ...interface{}) *Session - Insert(...interface{}) (int64, error) - InsertOne(interface{}) (int64, error) - IsTableEmpty(bean interface{}) (bool, error) - IsTableExist(beanOrTableName interface{}) (bool, error) - Iterate(interface{}, IterFunc) error + ID(any) *Session + In(string, ...any) *Session + Incr(column string, arg ...any) *Session + Insert(...any) (int64, error) + InsertOne(any) (int64, error) + IsTableEmpty(bean any) (bool, error) + IsTableExist(beanOrTableName any) (bool, error) + Iterate(any, IterFunc) error Limit(int, ...int) *Session MustCols(columns ...string) *Session NoAutoCondition(...bool) *Session - NotIn(string, ...interface{}) *Session + NotIn(string, ...any) *Session Nullable(...string) *Session - Join(joinOperator string, tablename interface{}, condition interface{}, args ...interface{}) *Session + Join(joinOperator string, tablename, condition any, args ...any) *Session Omit(columns ...string) *Session - OrderBy(order interface{}, args ...interface{}) *Session + OrderBy(order any, args ...any) *Session Ping() error - Query(sqlOrArgs ...interface{}) (resultsSlice []map[string][]byte, err error) - QueryInterface(sqlOrArgs ...interface{}) ([]map[string]interface{}, error) - QueryString(sqlOrArgs ...interface{}) ([]map[string]string, error) - Rows(bean interface{}) (*Rows, error) - SetExpr(string, interface{}) *Session + Query(sqlOrArgs ...any) (resultsSlice []map[string][]byte, err error) + QueryInterface(sqlOrArgs ...any) ([]map[string]any, error) + QueryString(sqlOrArgs ...any) ([]map[string]string, error) + Rows(bean any) (*Rows, error) + SetExpr(string, any) *Session Select(string) *Session - SQL(interface{}, ...interface{}) *Session - Sum(bean interface{}, colName string) (float64, error) - SumInt(bean interface{}, colName string) (int64, error) - Sums(bean interface{}, colNames ...string) ([]float64, error) - SumsInt(bean interface{}, colNames ...string) ([]int64, error) - Table(tableNameOrBean interface{}) *Session + SQL(any, ...any) *Session + Sum(bean any, colName string) (float64, error) + SumInt(bean any, colName string) (int64, error) + Sums(bean any, colNames ...string) ([]float64, error) + SumsInt(bean any, colNames ...string) ([]int64, error) + Table(tableNameOrBean any) *Session Unscoped() *Session - Update(bean interface{}, condiBeans ...interface{}) (int64, error) + Update(bean any, condiBeans ...any) (int64, error) UseBool(...string) *Session - Where(interface{}, ...interface{}) *Session + Where(any, ...any) *Session } // EngineInterface defines the interface which Engine, EngineGroup will implementate. type EngineInterface interface { Interface - Before(func(interface{})) *Session + Before(func(any)) *Session Charset(charset string) *Session - ClearCache(...interface{}) error + ClearCache(...any) error Context(context.Context) *Session - CreateTables(...interface{}) error + CreateTables(...any) error DBMetas() ([]*schemas.Table, error) DBVersion() (*schemas.Version, error) Dialect() dialects.Dialect DriverName() string - DropTables(...interface{}) error + DropTables(...any) error DumpAllToFile(fp string, tp ...schemas.DBType) error GetCacher(string) caches.Cacher GetColumnMapper() names.Mapper @@ -97,7 +97,7 @@ type EngineInterface interface { GetTZDatabase() *time.Location GetTZLocation() *time.Location ImportFile(fp string) ([]sql.Result, error) - MapCacher(interface{}, caches.Cacher) error + MapCacher(any, caches.Cacher) error NewSession() *Session NoAutoTime() *Session Prepare() *Session @@ -107,7 +107,7 @@ type EngineInterface interface { SetColumnMapper(names.Mapper) SetTagIdentifier(string) SetDefaultCacher(caches.Cacher) - SetLogger(logger interface{}) + SetLogger(logger any) SetLogLevel(log.LogLevel) SetMapper(names.Mapper) SetMaxOpenConns(int) @@ -119,12 +119,12 @@ type EngineInterface interface { SetTZLocation(tz *time.Location) AddHook(hook contexts.Hook) ShowSQL(show ...bool) - Sync(...interface{}) error - Sync2(...interface{}) error - SyncWithOptions(SyncOptions, ...interface{}) (*SyncResult, error) + Sync(...any) error + Sync2(...any) error + SyncWithOptions(SyncOptions, ...any) (*SyncResult, error) StoreEngine(storeEngine string) *Session - TableInfo(bean interface{}) (*schemas.Table, error) - TableName(interface{}, ...bool) string + TableInfo(bean any) (*schemas.Table, error) + TableName(any, ...bool) string UnMapType(reflect.Type) EnableSessionID(bool) } diff --git a/internal/json/gojson.go b/internal/json/gojson.go index 9bfa5c2..7009e09 100644 --- a/internal/json/gojson.go +++ b/internal/json/gojson.go @@ -3,7 +3,6 @@ // license that can be found in the LICENSE file. //go:build gojson -// +build gojson package json @@ -19,11 +18,11 @@ func init() { type GOjson struct{} // Marshal implements JSONInterface -func (GOjson) Marshal(v interface{}) ([]byte, error) { +func (GOjson) Marshal(v any) ([]byte, error) { return gojson.Marshal(v) } // Unmarshal implements JSONInterface -func (GOjson) Unmarshal(data []byte, v interface{}) error { +func (GOjson) Unmarshal(data []byte, v any) error { return gojson.Unmarshal(data, v) } diff --git a/internal/json/json.go b/internal/json/json.go index ef52f51..02b2adc 100644 --- a/internal/json/json.go +++ b/internal/json/json.go @@ -8,24 +8,22 @@ import "encoding/json" // Interface represents an interface to handle json data type Interface interface { - Marshal(v interface{}) ([]byte, error) - Unmarshal(data []byte, v interface{}) error + Marshal(v any) ([]byte, error) + Unmarshal(data []byte, v any) error } -var ( - // DefaultJSONHandler default json handler - DefaultJSONHandler Interface = StdJSON{} -) +// DefaultJSONHandler default json handler +var DefaultJSONHandler Interface = StdJSON{} // StdJSON implements JSONInterface via encoding/json type StdJSON struct{} // Marshal implements JSONInterface -func (StdJSON) Marshal(v interface{}) ([]byte, error) { +func (StdJSON) Marshal(v any) ([]byte, error) { return json.Marshal(v) } // Unmarshal implements JSONInterface -func (StdJSON) Unmarshal(data []byte, v interface{}) error { +func (StdJSON) Unmarshal(data []byte, v any) error { return json.Unmarshal(data, v) } diff --git a/internal/statements/args.go b/internal/statements/args.go index ce11dba..92ddede 100644 --- a/internal/statements/args.go +++ b/internal/statements/args.go @@ -23,7 +23,7 @@ func (n DateTimeString) Value() (driver.Value, error) { } // WriteArg writes an arg -func (statement *Statement) WriteArg(w *builder.BytesWriter, arg interface{}) error { +func (statement *Statement) WriteArg(w *builder.BytesWriter, arg any) error { switch argv := arg.(type) { case *builder.Builder: if _, err := w.WriteString("("); err != nil { @@ -64,7 +64,7 @@ func (statement *Statement) WriteArg(w *builder.BytesWriter, arg interface{}) er } // WriteArgs writes args -func (statement *Statement) WriteArgs(w *builder.BytesWriter, args []interface{}) error { +func (statement *Statement) WriteArgs(w *builder.BytesWriter, args []any) error { for i, arg := range args { if err := statement.WriteArg(w, arg); err != nil { return err diff --git a/internal/statements/column_map.go b/internal/statements/column_map.go index bb764b4..d87bafe 100644 --- a/internal/statements/column_map.go +++ b/internal/statements/column_map.go @@ -46,7 +46,7 @@ func (m *columnMap) Add(colName string) bool { return true } -func getFlagForColumn(m map[string]bool, col *schemas.Column) (val bool, has bool) { +func getFlagForColumn(m map[string]bool, col *schemas.Column) (val, has bool) { if len(m) == 0 { return false, false } diff --git a/internal/statements/cond.go b/internal/statements/cond.go index dfc6c20..77851b7 100644 --- a/internal/statements/cond.go +++ b/internal/statements/cond.go @@ -27,17 +27,17 @@ func (statement *Statement) QuoteReplacer(w *builder.BytesWriter) *QuoteReplacer } // Where add Where statement -func (statement *Statement) Where(query interface{}, args ...interface{}) *Statement { +func (statement *Statement) Where(query any, args ...any) *Statement { return statement.And(query, args...) } // And add Where & and statement -func (statement *Statement) And(query interface{}, args ...interface{}) *Statement { +func (statement *Statement) And(query any, args ...any) *Statement { switch qr := query.(type) { case string: cond := builder.Expr(qr, args...) statement.cond = statement.cond.And(cond) - case map[string]interface{}: + case map[string]any: cond := make(builder.Eq) for k, v := range qr { cond[statement.quote(k)] = v @@ -58,12 +58,12 @@ func (statement *Statement) And(query interface{}, args ...interface{}) *Stateme } // Or add Where & Or statement -func (statement *Statement) Or(query interface{}, args ...interface{}) *Statement { +func (statement *Statement) Or(query any, args ...any) *Statement { switch qr := query.(type) { case string: cond := builder.Expr(qr, args...) statement.cond = statement.cond.Or(cond) - case map[string]interface{}: + case map[string]any: cond := make(builder.Eq) for k, v := range qr { cond[statement.quote(k)] = v @@ -83,14 +83,14 @@ func (statement *Statement) Or(query interface{}, args ...interface{}) *Statemen } // In generate "Where column IN (?) " statement -func (statement *Statement) In(column string, args ...interface{}) *Statement { +func (statement *Statement) In(column string, args ...any) *Statement { in := builder.In(statement.quote(column), args...) statement.cond = statement.cond.And(in) return statement } // NotIn generate "Where column NOT IN (?) " statement -func (statement *Statement) NotIn(column string, args ...interface{}) *Statement { +func (statement *Statement) NotIn(column string, args ...any) *Statement { notIn := builder.NotIn(statement.quote(column), args...) statement.cond = statement.cond.And(notIn) return statement diff --git a/internal/statements/delete.go b/internal/statements/delete.go index 6e85939..78e23ee 100644 --- a/internal/statements/delete.go +++ b/internal/statements/delete.go @@ -21,7 +21,7 @@ func (statement *Statement) writeDeleteOrder(w *builder.BytesWriter) error { if statement.LimitN != nil && *statement.LimitN > 0 { if statement.Start > 0 { - return fmt.Errorf("Delete with Limit start is unsupported") + return errors.New("Delete with Limit start is unsupported") } limitNValue := *statement.LimitN if _, err := fmt.Fprintf(w, " LIMIT %d", limitNValue); err != nil { @@ -84,7 +84,7 @@ func (statement *Statement) writeOrderCond(orderCondWriter *builder.BytesWriter, } } -func (statement *Statement) WriteDelete(realSQLWriter, deleteSQLWriter *builder.BytesWriter, nowTime func(*schemas.Column) (interface{}, time.Time, error)) error { +func (statement *Statement) WriteDelete(realSQLWriter, deleteSQLWriter *builder.BytesWriter, nowTime func(*schemas.Column) (any, time.Time, error)) error { tableNameNoQuote := statement.TableName() tableName := statement.dialect.Quoter().Quote(tableNameNoQuote) table := statement.RefTable diff --git a/internal/statements/expr.go b/internal/statements/expr.go index c2a2e1c..58f73b3 100644 --- a/internal/statements/expr.go +++ b/internal/statements/expr.go @@ -24,7 +24,7 @@ func (err ErrUnsupportedExprType) Error() string { // Expr represents an SQL express type Expr struct { ColName string - Arg interface{} + Arg any } // WriteArgs writes args to the writer @@ -44,7 +44,7 @@ func (expr *Expr) WriteArgs(w *builder.BytesWriter) error { if arg == "" { arg = "''" } - if _, err := w.WriteString(fmt.Sprintf("%v", arg)); err != nil { + if _, err := w.WriteString(arg); err != nil { return err } default: @@ -59,14 +59,14 @@ func (expr *Expr) WriteArgs(w *builder.BytesWriter) error { type exprParams []Expr func (exprs exprParams) ColNames() []string { - var cols = make([]string, 0, len(exprs)) + cols := make([]string, 0, len(exprs)) for _, expr := range exprs { cols = append(cols, expr.ColName) } return cols } -func (exprs *exprParams) Add(name string, arg interface{}) { +func (exprs *exprParams) Add(name string, arg any) { *exprs = append(*exprs, Expr{name, arg}) } diff --git a/internal/statements/insert.go b/internal/statements/insert.go index aa39643..30d1858 100644 --- a/internal/statements/insert.go +++ b/internal/statements/insert.go @@ -27,7 +27,7 @@ func (statement *Statement) writeInsertOutput(buf *strings.Builder, table *schem } // GenInsertSQL generates insert beans SQL -func (statement *Statement) GenInsertSQL(colNames []string, args []interface{}) (string, []interface{}, error) { +func (statement *Statement) GenInsertSQL(colNames []string, args []any) (string, []any, error) { var ( buf = builder.NewWriter() exprs = statement.ExprColumns @@ -177,14 +177,14 @@ func (statement *Statement) GenInsertSQL(colNames []string, args []interface{}) } // GenInsertMapSQL generates insert map SQL -func (statement *Statement) GenInsertMapSQL(columns []string, args []interface{}) (string, []interface{}, error) { +func (statement *Statement) GenInsertMapSQL(columns []string, args []any) (string, []any, error) { var ( buf = builder.NewWriter() exprs = statement.ExprColumns tableName = statement.TableName() ) - if _, err := buf.WriteString(fmt.Sprintf("INSERT INTO %s (", statement.quote(tableName))); err != nil { + if _, err := fmt.Fprintf(buf, "INSERT INTO %s (", statement.quote(tableName)); err != nil { return "", nil, err } @@ -211,7 +211,7 @@ func (statement *Statement) GenInsertMapSQL(columns []string, args []interface{} } } - if _, err := buf.WriteString(fmt.Sprintf(" FROM %s WHERE ", statement.quote(tableName))); err != nil { + if _, err := fmt.Fprintf(buf, " FROM %s WHERE ", statement.quote(tableName)); err != nil { return "", nil, err } @@ -242,14 +242,14 @@ func (statement *Statement) GenInsertMapSQL(columns []string, args []interface{} return buf.String(), buf.Args(), nil } -func (statement *Statement) GenInsertMultipleMapSQL(columns []string, argss [][]interface{}) (string, []interface{}, error) { +func (statement *Statement) GenInsertMultipleMapSQL(columns []string, argss [][]any) (string, []any, error) { var ( buf = builder.NewWriter() exprs = statement.ExprColumns tableName = statement.TableName() ) - if _, err := buf.WriteString(fmt.Sprintf("INSERT INTO %s (", statement.quote(tableName))); err != nil { + if _, err := fmt.Fprintf(buf, "INSERT INTO %s (", statement.quote(tableName)); err != nil { return "", nil, err } @@ -308,21 +308,7 @@ func (statement *Statement) writeColumns(w *builder.BytesWriter, slice []string) return nil } -func (statement *Statement) writeQuestions(w *builder.BytesWriter, length int) error { - for i := 0; i < length; i++ { - if i > 0 { - if _, err := fmt.Fprint(w, ","); err != nil { - return err - } - } - if _, err := fmt.Fprint(w, "?"); err != nil { - return err - } - } - return nil -} - -func (statement *Statement) oracleWriteInsertMultiple(w *builder.BytesWriter, tableName string, colNames []string, colMultiPlaces []string) error { +func (statement *Statement) oracleWriteInsertMultiple(w *builder.BytesWriter, tableName string, colNames, colMultiPlaces []string) error { if _, err := fmt.Fprint(w, "INSERT ALL"); err != nil { return err } @@ -354,14 +340,14 @@ func (statement *Statement) oracleWriteInsertMultiple(w *builder.BytesWriter, ta return nil } -func (statement *Statement) WriteInsertMultiple(w *builder.BytesWriter, tableName string, colNames []string, colMultiPlaces []string) error { +func (statement *Statement) WriteInsertMultiple(w *builder.BytesWriter, tableName string, colNames, colMultiPlaces []string) error { if statement.dialect.URI().DBType == schemas.ORACLE { return statement.oracleWriteInsertMultiple(w, tableName, colNames, colMultiPlaces) } return statement.plainWriteInsertMultiple(w, tableName, colNames, colMultiPlaces) } -func (statement *Statement) plainWriteInsertMultiple(w *builder.BytesWriter, tableName string, colNames []string, colMultiPlaces []string) error { +func (statement *Statement) plainWriteInsertMultiple(w *builder.BytesWriter, tableName string, colNames, colMultiPlaces []string) error { if _, err := fmt.Fprint(w, "INSERT INTO "); err != nil { return err } diff --git a/internal/statements/join.go b/internal/statements/join.go index 6fa42ea..555427e 100644 --- a/internal/statements/join.go +++ b/internal/statements/join.go @@ -15,7 +15,7 @@ import ( ) // Join The joinOP should be one of INNER, LEFT OUTER, CROSS etc - this will be prepended to JOIN -func (statement *Statement) Join(joinOP string, joinTable interface{}, condition interface{}, args ...interface{}) *Statement { +func (statement *Statement) Join(joinOP string, joinTable, condition any, args ...any) *Statement { statement.joins = append(statement.joins, join{ op: joinOP, table: joinTable, diff --git a/internal/statements/order_by.go b/internal/statements/order_by.go index 04197fb..6f7129c 100644 --- a/internal/statements/order_by.go +++ b/internal/statements/order_by.go @@ -12,28 +12,28 @@ import ( ) type orderBy struct { - orderStr interface{} - orderArgs []interface{} + orderStr any + orderArgs []any direction string // ASC, DESC or "", "" means raw orderStr } func (ob orderBy) CheckValid() error { if ob.orderStr == nil { - return fmt.Errorf("order by string is nil") + return errors.New("order by string is nil") } switch t := ob.orderStr.(type) { case string: if t == "" { - return fmt.Errorf("order by string is empty") + return errors.New("order by string is empty") } return nil case *builder.Expression: if t.Content() == "" { - return fmt.Errorf("order by string is empty") + return errors.New("order by string is empty") } return nil default: - return fmt.Errorf("order by string is not string or builder.Expression") + return errors.New("order by string is not string or builder.Expression") } } @@ -97,7 +97,7 @@ func (statement *Statement) writeOrderBys(w *builder.BytesWriter) error { } // OrderBy generate "Order By order" statement -func (statement *Statement) OrderBy(order interface{}, args ...interface{}) *Statement { +func (statement *Statement) OrderBy(order any, args ...any) *Statement { ob := orderBy{order, args, ""} if err := ob.CheckValid(); err != nil { statement.LastError = err diff --git a/internal/statements/pk.go b/internal/statements/pk.go index 59da89c..96f2e78 100644 --- a/internal/statements/pk.go +++ b/internal/statements/pk.go @@ -36,7 +36,7 @@ func IsIDConditionWithNoTableErr(err error) bool { } // ID generate "where id = ? " statement or for composite key "where key1 = ? and key2 = ?" -func (statement *Statement) ID(id interface{}) *Statement { +func (statement *Statement) ID(id any) *Statement { switch t := id.(type) { case *schemas.PK: statement.idParam = *t @@ -91,7 +91,7 @@ func (statement *Statement) ProcessIDParam() error { } for i, col := range statement.RefTable.PKColumns() { - var colName = statement.colName(col, statement.TableName()) + colName := statement.colName(col, statement.TableName()) statement.cond = statement.cond.And(builder.Eq{colName: statement.idParam[i]}) } return nil diff --git a/internal/statements/query.go b/internal/statements/query.go index e817403..36dbabc 100644 --- a/internal/statements/query.go +++ b/internal/statements/query.go @@ -15,7 +15,7 @@ import ( ) // GenQuerySQL generate query SQL -func (statement *Statement) GenQuerySQL(sqlOrArgs ...interface{}) (string, []interface{}, error) { +func (statement *Statement) GenQuerySQL(sqlOrArgs ...any) (string, []any, error) { if len(sqlOrArgs) > 0 { return statement.ConvertSQLOrArgs(sqlOrArgs...) } @@ -24,7 +24,7 @@ func (statement *Statement) GenQuerySQL(sqlOrArgs ...interface{}) (string, []int return statement.GenRawSQL(), statement.RawParams, nil } - if len(statement.TableName()) <= 0 { + if len(statement.TableName()) == 0 { return "", nil, ErrTableNotFound } @@ -40,7 +40,7 @@ func (statement *Statement) GenQuerySQL(sqlOrArgs ...interface{}) (string, []int } // GenSumSQL generates sum SQL -func (statement *Statement) GenSumSQL(bean interface{}, columns ...string) (string, []interface{}, error) { +func (statement *Statement) GenSumSQL(bean any, columns ...string) (string, []any, error) { if statement.RawSQL != "" { return statement.GenRawSQL(), statement.RawParams, nil } @@ -71,7 +71,7 @@ func (statement *Statement) GenSumSQL(bean interface{}, columns ...string) (stri } // GenGetSQL generates Get SQL -func (statement *Statement) GenGetSQL(bean interface{}) (string, []interface{}, error) { +func (statement *Statement) GenGetSQL(bean any) (string, []any, error) { var isStruct bool if bean != nil { v := rValue(bean) @@ -127,7 +127,7 @@ func (statement *Statement) GenGetSQL(bean interface{}) (string, []interface{}, } // GenCountSQL generates the SQL for counting -func (statement *Statement) GenCountSQL(beans ...interface{}) (string, []interface{}, error) { +func (statement *Statement) GenCountSQL(beans ...any) (string, []any, error) { if statement.RawSQL != "" { return statement.GenRawSQL(), statement.RawParams, nil } @@ -142,7 +142,7 @@ func (statement *Statement) GenCountSQL(beans ...interface{}) (string, []interfa } selectSQL := statement.SelectStr - if len(selectSQL) <= 0 { + if len(selectSQL) == 0 { if statement.IsDistinct { selectSQL = fmt.Sprintf("count(DISTINCT %s)", statement.ColumnStr()) } else if statement.ColumnStr() != "" { @@ -263,18 +263,18 @@ func (statement *Statement) writeSelect(buf *builder.BytesWriter, columnStr stri // ORDER BY is mandatory to use OFFSET and FETCH clause (only in sqlserver) if statement.LimitN == nil && statement.Start == 0 { // no need to add - return + return nil } if statement.IsDistinct || len(statement.GroupByStr) > 0 || isCounting { // the order-by column should be one of distincts or group-bys // order by the first column _, err = bw.WriteString(" ORDER BY 1 ASC") - return + return err } if statement.RefTable == nil || len(statement.RefTable.PrimaryKeys) != 1 { // no primary key, order by the first column _, err = bw.WriteString(" ORDER BY 1 ASC") - return + return err } // order by primary key statement.orderBy = []orderBy{{orderStr: statement.colName(statement.RefTable.GetColumn(statement.RefTable.PrimaryKeys[0]), statement.TableName()), direction: "ASC"}} @@ -287,12 +287,12 @@ func (statement *Statement) writeSelect(buf *builder.BytesWriter, columnStr stri } // GenExistSQL generates Exist SQL -func (statement *Statement) GenExistSQL(bean ...interface{}) (string, []interface{}, error) { +func (statement *Statement) GenExistSQL(bean ...any) (string, []any, error) { if statement.RawSQL != "" { return statement.GenRawSQL(), statement.RawParams, nil } - var b interface{} + var b any if len(bean) > 0 { b = bean[0] beanValue := reflect.ValueOf(bean[0]) @@ -307,7 +307,7 @@ func (statement *Statement) GenExistSQL(bean ...interface{}) (string, []interfac } } tableName := statement.TableName() - if len(tableName) <= 0 { + if len(tableName) == 0 { return "", nil, ErrTableNotFound } if statement.RefTable != nil { @@ -394,12 +394,12 @@ func (statement *Statement) genSelectColumnStr() string { } // GenFindSQL generates Find SQL -func (statement *Statement) GenFindSQL(autoCond builder.Cond) (string, []interface{}, error) { +func (statement *Statement) GenFindSQL(autoCond builder.Cond) (string, []any, error) { if statement.RawSQL != "" { return statement.GenRawSQL(), statement.RawParams, nil } - if len(statement.TableName()) <= 0 { + if len(statement.TableName()) == 0 { return "", nil, ErrTableNotFound } diff --git a/internal/statements/select.go b/internal/statements/select.go index 59161d7..2d23150 100644 --- a/internal/statements/select.go +++ b/internal/statements/select.go @@ -20,8 +20,8 @@ func (statement *Statement) Select(str string) *Statement { func col2NewCols(columns ...string) []string { newColumns := make([]string, 0, len(columns)) for _, col := range columns { - col = strings.Replace(col, "`", "", -1) - col = strings.Replace(col, `"`, "", -1) + col = strings.ReplaceAll(col, "`", "") + col = strings.ReplaceAll(col, `"`, "") ccols := strings.Split(col, ",") for _, c := range ccols { newColumns = append(newColumns, strings.TrimSpace(c)) @@ -112,7 +112,7 @@ func (statement *Statement) genColumnStr() string { buf.WriteString(".") } - statement.dialect.Quoter().QuoteTo(&buf, col.Name) + buf.WriteString(statement.dialect.Quoter().Quote(col.Name)) } return buf.String() diff --git a/internal/statements/statement.go b/internal/statements/statement.go index bbed50e..dd8d981 100644 --- a/internal/statements/statement.go +++ b/internal/statements/statement.go @@ -36,9 +36,9 @@ var ( type join struct { op string - table interface{} - condition interface{} - args []interface{} + table any + condition any + args []any } type indexHint struct { @@ -65,7 +65,7 @@ type Statement struct { AltTableName string tableName string RawSQL string - RawParams []interface{} + RawParams []any UseCascade bool UseAutoJoin bool StoreEngine string @@ -144,7 +144,7 @@ func (statement *Statement) Reset() { statement.tableName = "" statement.idParam = nil statement.RawSQL = "" - statement.RawParams = make([]interface{}, 0) + statement.RawParams = make([]any, 0) statement.UseCache = true statement.UseAutoTime = true statement.NoAutoCondition = false @@ -168,7 +168,7 @@ func (statement *Statement) Reset() { } // SQL adds raw sql statement -func (statement *Statement) SQL(query interface{}, args ...interface{}) *Statement { +func (statement *Statement) SQL(query any, args ...any) *Statement { switch t := query.(type) { case *builder.Builder: var err error @@ -201,12 +201,12 @@ func (statement *Statement) SetRefValue(v reflect.Value) error { return nil } -func rValue(bean interface{}) reflect.Value { +func rValue(bean any) reflect.Value { return reflect.Indirect(reflect.ValueOf(bean)) } // SetRefBean set ref bean -func (statement *Statement) SetRefBean(bean interface{}) error { +func (statement *Statement) SetRefBean(bean any) error { var err error statement.RefTable, err = statement.tagParser.ParseWithCache(rValue(bean)) if err != nil { @@ -221,7 +221,7 @@ func (statement *Statement) NeedTableName() bool { } // Incr Generate "Update ... Set column = column + arg" statement -func (statement *Statement) Incr(column string, arg ...interface{}) *Statement { +func (statement *Statement) Incr(column string, arg ...any) *Statement { if len(arg) > 0 { statement.IncrColumns.Add(column, arg[0]) } else { @@ -231,7 +231,7 @@ func (statement *Statement) Incr(column string, arg ...interface{}) *Statement { } // Decr Generate "Update ... Set column = column - arg" statement -func (statement *Statement) Decr(column string, arg ...interface{}) *Statement { +func (statement *Statement) Decr(column string, arg ...any) *Statement { if len(arg) > 0 { statement.DecrColumns.Add(column, arg[0]) } else { @@ -241,7 +241,7 @@ func (statement *Statement) Decr(column string, arg ...interface{}) *Statement { } // SetExpr Generate "Update ... Set column = {expression}" statement -func (statement *Statement) SetExpr(column string, expression interface{}) *Statement { +func (statement *Statement) SetExpr(column string, expression any) *Statement { if e, ok := expression.(string); ok { statement.ExprColumns.Add(column, statement.dialect.Quoter().Replace(e)) } else { @@ -280,7 +280,7 @@ func (statement *Statement) Limit(limit int, start ...int) *Statement { } // SetTable tempororily set table name, the parameter could be a string or a pointer of struct -func (statement *Statement) SetTable(tableNameOrBean interface{}) error { +func (statement *Statement) SetTable(tableNameOrBean any) error { v := rValue(tableNameOrBean) t := v.Type() if t.Kind() == reflect.Struct { @@ -374,7 +374,7 @@ func (statement *Statement) GenDelIndexSQL() []string { return sqls } -func (statement *Statement) asDBCond(fieldValue reflect.Value, fieldType reflect.Type, col *schemas.Column, allUseBool, requiredField bool) (interface{}, bool, error) { +func (statement *Statement) asDBCond(fieldValue reflect.Value, fieldType reflect.Type, col *schemas.Column, allUseBool, requiredField bool) (any, bool, error) { switch fieldType.Kind() { case reflect.Ptr: if fieldValue.IsNil() { @@ -443,40 +443,39 @@ func (statement *Statement) asDBCond(fieldValue reflect.Value, fieldType reflect } } return val, true, nil - } else { - if col.IsJSON { - if col.SQLType.IsText() { - bytes, err := json.DefaultJSONHandler.Marshal(fieldValue.Interface()) - if err != nil { - return nil, false, err - } - return string(bytes), true, nil - } else if col.SQLType.IsBlob() { - var bytes []byte - var err error - bytes, err = json.DefaultJSONHandler.Marshal(fieldValue.Interface()) - if err != nil { - return nil, false, err - } - return bytes, true, nil - } - } else { - table, err := statement.tagParser.ParseWithCache(fieldValue) + } + if col.IsJSON { + if col.SQLType.IsText() { + bytes, err := json.DefaultJSONHandler.Marshal(fieldValue.Interface()) if err != nil { - return fieldValue.Interface(), true, nil + return nil, false, err } - - if len(table.PrimaryKeys) == 1 { - pkField := reflect.Indirect(fieldValue).FieldByName(table.PKColumns()[0].FieldName) - // fix non-int pk issues - // if pkField.Int() != 0 { - if pkField.IsValid() && !utils.IsZero(pkField.Interface()) { - return pkField.Interface(), true, nil - } - return nil, false, nil + return string(bytes), true, nil + } else if col.SQLType.IsBlob() { + var bytes []byte + var err error + bytes, err = json.DefaultJSONHandler.Marshal(fieldValue.Interface()) + if err != nil { + return nil, false, err } - return nil, false, fmt.Errorf("not supported %v as %v", fieldValue.Interface(), table.PrimaryKeys) + return bytes, true, nil } + } else { + table, err := statement.tagParser.ParseWithCache(fieldValue) + if err != nil { + return fieldValue.Interface(), true, nil + } + + if len(table.PrimaryKeys) == 1 { + pkField := reflect.Indirect(fieldValue).FieldByName(table.PKColumns()[0].FieldName) + // fix non-int pk issues + // if pkField.Int() != 0 { + if pkField.IsValid() && !utils.IsZero(pkField.Interface()) { + return pkField.Interface(), true, nil + } + return nil, false, nil + } + return nil, false, fmt.Errorf("not supported %v as %v", fieldValue.Interface(), table.PrimaryKeys) } case reflect.Array: return nil, false, nil @@ -515,9 +514,8 @@ func (statement *Statement) asDBCond(fieldValue reflect.Value, fieldType reflect return fieldValue.Interface(), true, nil } -func (statement *Statement) buildConds2(table *schemas.Table, bean interface{}, - includeVersion bool, includeUpdated bool, includeNil bool, - includeAutoIncr bool, allUseBool bool, useAllCols bool, unscoped bool, +func (statement *Statement) buildConds2(table *schemas.Table, bean any, + includeVersion, includeUpdated, includeNil, includeAutoIncr, allUseBool, useAllCols, unscoped bool, mustColumnMap map[string]bool, tableName, aliasName string, addedTableName bool, ) (builder.Cond, error) { var conds []builder.Cond @@ -588,14 +586,14 @@ func (statement *Statement) buildConds2(table *schemas.Table, bean interface{}, conds = append(conds, builder.Eq{colName: nil}) } continue - } else if !fieldValue.IsValid() { - continue - } else { - // dereference ptr type to instance type - fieldValue = fieldValue.Elem() - fieldType = reflect.TypeOf(fieldValue.Interface()) - requiredField = true } + if !fieldValue.IsValid() { + continue + } + // dereference ptr type to instance type + fieldValue = fieldValue.Elem() + fieldType = reflect.TypeOf(fieldValue.Interface()) + requiredField = true } val, ok, err := statement.asDBCond(fieldValue, fieldType, col, allUseBool, requiredField) @@ -613,13 +611,13 @@ func (statement *Statement) buildConds2(table *schemas.Table, bean interface{}, } // BuildConds builds condition -func (statement *Statement) BuildConds(table *schemas.Table, bean interface{}, includeVersion bool, includeUpdated bool, includeNil bool, includeAutoIncr bool, addedTableName bool) (builder.Cond, error) { +func (statement *Statement) BuildConds(table *schemas.Table, bean any, includeVersion, includeUpdated, includeNil, includeAutoIncr, addedTableName bool) (builder.Cond, error) { return statement.buildConds2(table, bean, includeVersion, includeUpdated, includeNil, includeAutoIncr, statement.allUseBool, statement.useAllCols, statement.unscoped, statement.MustColumnMap, statement.TableName(), statement.TableAlias, addedTableName) } // MergeConds merge conditions from bean and id -func (statement *Statement) MergeConds(bean interface{}) error { +func (statement *Statement) MergeConds(bean any) error { if !statement.NoAutoCondition && statement.RefTable != nil { addedTableName := len(statement.joins) > 0 autoCond, err := statement.BuildConds(statement.RefTable, bean, true, true, false, true, addedTableName) @@ -638,7 +636,7 @@ func (statement *Statement) quoteColumnStr(columnStr string) string { } // ConvertSQLOrArgs converts sql or args -func (statement *Statement) ConvertSQLOrArgs(sqlOrArgs ...interface{}) (string, []interface{}, error) { +func (statement *Statement) ConvertSQLOrArgs(sqlOrArgs ...any) (string, []any, error) { sql, args, err := statement.convertSQLOrArgs(sqlOrArgs...) if err != nil { return "", nil, err @@ -646,11 +644,11 @@ func (statement *Statement) ConvertSQLOrArgs(sqlOrArgs ...interface{}) (string, return statement.ReplaceQuote(sql), args, nil } -func (statement *Statement) convertSQLOrArgs(sqlOrArgs ...interface{}) (string, []interface{}, error) { - switch sqlOrArgs[0].(type) { +func (statement *Statement) convertSQLOrArgs(sqlOrArgs ...any) (string, []any, error) { + switch arg := sqlOrArgs[0].(type) { case string: if len(sqlOrArgs) > 1 { - newArgs := make([]interface{}, 0, len(sqlOrArgs)-1) + newArgs := make([]any, 0, len(sqlOrArgs)-1) for _, arg := range sqlOrArgs[1:] { if v, ok := arg.(time.Time); ok { newArgs = append(newArgs, v.In(statement.defaultTimeZone).Format("2006-01-02 15:04:05")) @@ -677,14 +675,13 @@ func (statement *Statement) convertSQLOrArgs(sqlOrArgs ...interface{}) (string, newArgs = append(newArgs, arg) } } - return sqlOrArgs[0].(string), newArgs, nil + return arg, newArgs, nil } - return sqlOrArgs[0].(string), sqlOrArgs[1:], nil + return arg, sqlOrArgs[1:], nil case *builder.Builder: - return sqlOrArgs[0].(*builder.Builder).ToSQL() + return arg.ToSQL() case builder.Builder: - bd := sqlOrArgs[0].(builder.Builder) - return bd.ToSQL() + return arg.ToSQL() } return "", nil, ErrUnSupportedType @@ -715,7 +712,7 @@ func (statement *Statement) CondDeleted(col *schemas.Column) builder.Cond { } colName = statement.quote(prefix) + "." + statement.quote(col.Name) } - cond := builder.NewCond() + var cond builder.Cond if col.SQLType.IsNumeric() { cond = builder.Eq{colName: 0} } else if col.SQLType.Name == schemas.TimeStamp || col.SQLType.Name == schemas.TimeStampz { diff --git a/internal/statements/statement_test.go b/internal/statements/statement_test.go index f4fc2df..e59b5e5 100644 --- a/internal/statements/statement_test.go +++ b/internal/statements/statement_test.go @@ -89,7 +89,7 @@ func TestConvertSQLOrArgs(t *testing.T) { // ID int // del *time.Time `xorm:"deleted"` // } - args := []interface{}{ + args := []any{ "INSERT `table` (`id`, `del`) VALUES (?, ?)", 1, (*time.Time)(nil), } // before fix, here will panic diff --git a/internal/statements/update.go b/internal/statements/update.go index dd490ae..fefee6e 100644 --- a/internal/statements/update.go +++ b/internal/statements/update.go @@ -19,55 +19,55 @@ import ( "xorm.io/xorm/schemas" ) -func (statement *Statement) ifAddColUpdate(col *schemas.Column, includeVersion, includeUpdated, includeNil, - includeAutoIncr, update bool, -) (bool, error) { +func (statement *Statement) ifAddColUpdate(col *schemas.Column, includeVersion, includeUpdated, + includeAutoIncr bool, +) bool { columnMap := statement.ColumnMap omitColumnMap := statement.OmitColumnMap unscoped := statement.unscoped if !includeVersion && col.IsVersion { - return false, nil + return false } if col.IsCreated && !columnMap.Contain(col.Name) { - return false, nil + return false } if !includeUpdated && col.IsUpdated { - return false, nil + return false } if !includeAutoIncr && col.IsAutoIncrement { - return false, nil + return false } if col.IsDeleted && !unscoped { - return false, nil + return false } if omitColumnMap.Contain(col.Name) { - return false, nil + return false } if len(columnMap) > 0 && !columnMap.Contain(col.Name) { - return false, nil + return false } if col.MapType == schemas.ONLYFROMDB { - return false, nil + return false } if statement.IncrColumns.IsColExist(col.Name) { - return false, nil + return false } else if statement.DecrColumns.IsColExist(col.Name) { - return false, nil + return false } else if statement.ExprColumns.IsColExist(col.Name) { - return false, nil + return false } - return true, nil + return true } // BuildUpdates auto generating update columnes and values according a struct func (statement *Statement) BuildUpdates(tableValue reflect.Value, includeVersion, includeUpdated, includeNil, includeAutoIncr, update bool, -) ([]string, []interface{}, error) { +) ([]string, []any, error) { table := statement.RefTable allUseBool := statement.allUseBool useAllCols := statement.useAllCols @@ -75,14 +75,10 @@ func (statement *Statement) BuildUpdates(tableValue reflect.Value, nullableMap := statement.NullableMap colNames := make([]string, 0) - args := make([]interface{}, 0) + args := make([]any, 0) for _, col := range table.Columns() { - ok, err := statement.ifAddColUpdate(col, includeVersion, includeUpdated, includeNil, - includeAutoIncr, update) - if err != nil { - return nil, nil, err - } + ok := statement.ifAddColUpdate(col, includeVersion, includeUpdated, includeAutoIncr) if !ok { continue } @@ -122,7 +118,7 @@ func (statement *Statement) BuildUpdates(tableValue reflect.Value, } } - var val interface{} + var val any if fieldValue.CanAddr() { if structConvert, ok := fieldValue.Addr().Interface().(convert.Conversion); ok { @@ -164,14 +160,14 @@ func (statement *Statement) BuildUpdates(tableValue reflect.Value, colNames = append(colNames, fmt.Sprintf("%v=?", statement.quote(col.Name))) } continue - } else if !fieldValue.IsValid() { - continue - } else { - // dereference ptr type to instance type - fieldValue = fieldValue.Elem() - fieldType = reflect.TypeOf(fieldValue.Interface()) - requiredField = true } + if !fieldValue.IsValid() { + continue + } + // dereference ptr type to instance type + fieldValue = fieldValue.Elem() + fieldType = reflect.TypeOf(fieldValue.Interface()) + requiredField = true } switch fieldType.Kind() { @@ -461,13 +457,13 @@ func (statement *Statement) writeUpdateLimit(updateWriter *builder.BytesWriter, } return nil default: // TODO: Oracle support needed - return fmt.Errorf("not implemented") + return errors.New("not implemented") } } -func (statement *Statement) GenConditionsFromMap(m interface{}) ([]builder.Cond, error) { +func (statement *Statement) GenConditionsFromMap(m any) ([]builder.Cond, error) { switch t := m.(type) { - case map[string]interface{}: + case map[string]any: conds := []builder.Cond{} for k, v := range t { conds = append(conds, builder.Eq{k: v}) @@ -582,7 +578,7 @@ func (statement *Statement) writeExprSets(w *builder.BytesWriter, hasPreviousSet return nil } -func (statement *Statement) writeSetColumns(colNames []string, args []interface{}) func(w *builder.BytesWriter) error { +func (statement *Statement) writeSetColumns(colNames []string, args []any) func(w *builder.BytesWriter) error { return func(w *builder.BytesWriter) error { if len(colNames) == 0 { return nil @@ -615,7 +611,7 @@ func (statement *Statement) writeSetColumns(colNames []string, args []interface{ } } -func (statement *Statement) writeUpdateSets(w *builder.BytesWriter, v reflect.Value, colNames []string, args []interface{}) error { +func (statement *Statement) writeUpdateSets(w *builder.BytesWriter, v reflect.Value, colNames []string, args []any) error { // write set if _, err := fmt.Fprint(w, " SET "); err != nil { return err @@ -655,7 +651,7 @@ func (statement *Statement) writeUpdateSets(w *builder.BytesWriter, v reflect.Va var ErrNoColumnsTobeUpdated = errors.New("no columns found to be updated") -func (statement *Statement) WriteUpdate(updateWriter *builder.BytesWriter, cond builder.Cond, v reflect.Value, colNames []string, args []interface{}) error { +func (statement *Statement) WriteUpdate(updateWriter *builder.BytesWriter, cond builder.Cond, v reflect.Value, colNames []string, args []any) error { switch statement.dialect.URI().DBType { case schemas.MYSQL: return statement.writeUpdateMySQL(updateWriter, cond, v, colNames, args) @@ -666,7 +662,7 @@ func (statement *Statement) WriteUpdate(updateWriter *builder.BytesWriter, cond } } -func (statement *Statement) writeUpdateMySQL(updateWriter *builder.BytesWriter, cond builder.Cond, v reflect.Value, colNames []string, args []interface{}) error { +func (statement *Statement) writeUpdateMySQL(updateWriter *builder.BytesWriter, cond builder.Cond, v reflect.Value, colNames []string, args []any) error { if _, err := fmt.Fprintf(updateWriter, "UPDATE"); err != nil { return err } @@ -689,7 +685,7 @@ func (statement *Statement) writeUpdateMySQL(updateWriter *builder.BytesWriter, return statement.writeUpdateLimit(updateWriter, cond) } -func (statement *Statement) writeUpdateMSSQL(updateWriter *builder.BytesWriter, cond builder.Cond, v reflect.Value, colNames []string, args []interface{}) error { +func (statement *Statement) writeUpdateMSSQL(updateWriter *builder.BytesWriter, cond builder.Cond, v reflect.Value, colNames []string, args []any) error { if _, err := fmt.Fprintf(updateWriter, "UPDATE"); err != nil { return err } @@ -725,7 +721,7 @@ func (statement *Statement) writeUpdateMSSQL(updateWriter *builder.BytesWriter, } // writeUpdateCommon write update sql for non mysql && non mssql -func (statement *Statement) writeUpdateCommon(updateWriter *builder.BytesWriter, cond builder.Cond, v reflect.Value, colNames []string, args []interface{}) error { +func (statement *Statement) writeUpdateCommon(updateWriter *builder.BytesWriter, cond builder.Cond, v reflect.Value, colNames []string, args []any) error { if _, err := fmt.Fprintf(updateWriter, "UPDATE"); err != nil { return err } diff --git a/internal/statements/values.go b/internal/statements/values.go index f27caed..dcaea75 100644 --- a/internal/statements/values.go +++ b/internal/statements/values.go @@ -24,7 +24,7 @@ var ( ) // Value2Interface convert a field value of a struct to interface for putting into database -func (statement *Statement) Value2Interface(col *schemas.Column, fieldValue reflect.Value) (interface{}, error) { +func (statement *Statement) Value2Interface(col *schemas.Column, fieldValue reflect.Value) (any, error) { if fieldValue.CanAddr() { if fieldConvert, ok := fieldValue.Addr().Interface().(convert.Conversion); ok { data, err := fieldConvert.ToDB() @@ -69,14 +69,14 @@ func (statement *Statement) Value2Interface(col *schemas.Column, fieldValue refl if k == reflect.Ptr { if fieldValue.IsNil() { return nil, nil - } else if !fieldValue.IsValid() { - return nil, nil - } else { - // !nashtsai! deference pointer type to instance type - fieldValue = fieldValue.Elem() - fieldType = fieldValue.Type() - k = fieldType.Kind() } + if !fieldValue.IsValid() { + return nil, nil + } + // !nashtsai! deference pointer type to instance type + fieldValue = fieldValue.Elem() + fieldType = fieldValue.Type() + k = fieldType.Kind() } switch k { @@ -99,9 +99,8 @@ func (statement *Statement) Value2Interface(col *schemas.Column, fieldValue refl layout = "yyyy-MM-dd HH24:mi:ss" } return &DateTimeString{Layout: layout, Str: val}, err - } else { - return tf, err } + return tf, err } else if fieldType.ConvertibleTo(nullFloatType) { t := fieldValue.Convert(nullFloatType).Interface().(sql.NullFloat64) if !t.Valid { diff --git a/internal/statements/writer.go b/internal/statements/writer.go index b4ca804..073601e 100644 --- a/internal/statements/writer.go +++ b/internal/statements/writer.go @@ -30,8 +30,8 @@ func (statement *Statement) groupWriteFns(writeFuncs ...func(*builder.BytesWrite func (statement *Statement) writeMultiple(buf *builder.BytesWriter, writeFuncs ...func(*builder.BytesWriter) error) (err error) { for _, fn := range writeFuncs { if err = fn(buf); err != nil { - return + return err } } - return + return nil } diff --git a/internal/utils/builder.go b/internal/utils/builder.go index bc97526..da7a640 100644 --- a/internal/utils/builder.go +++ b/internal/utils/builder.go @@ -12,7 +12,7 @@ import ( type BuildReader interface { String() string - Args() []interface{} + Args() []any } // WriteBuilder writes writers to one diff --git a/internal/utils/new.go b/internal/utils/new.go index e3b4eae..44da464 100644 --- a/internal/utils/new.go +++ b/internal/utils/new.go @@ -7,15 +7,15 @@ package utils import "reflect" // New creates a value according type -func New(tp reflect.Type, length, cap int) reflect.Value { +func New(tp reflect.Type, length, capacity int) reflect.Value { switch tp.Kind() { case reflect.Slice: - slice := reflect.MakeSlice(tp, length, cap) + slice := reflect.MakeSlice(tp, length, capacity) x := reflect.New(slice.Type()) x.Elem().Set(slice) return x case reflect.Map: - mp := reflect.MakeMapWithSize(tp, cap) + mp := reflect.MakeMapWithSize(tp, capacity) x := reflect.New(mp.Type()) x.Elem().Set(mp) return x diff --git a/internal/utils/reflect.go b/internal/utils/reflect.go index 7973d4d..d2b4b3b 100644 --- a/internal/utils/reflect.go +++ b/internal/utils/reflect.go @@ -9,6 +9,6 @@ import ( ) // ReflectValue returns value of a bean -func ReflectValue(bean interface{}) reflect.Value { +func ReflectValue(bean any) reflect.Value { return reflect.Indirect(reflect.ValueOf(bean)) } diff --git a/internal/utils/zero.go b/internal/utils/zero.go index 591ae24..d105f4f 100644 --- a/internal/utils/zero.go +++ b/internal/utils/zero.go @@ -17,7 +17,7 @@ type Zeroable interface { var nilTime *time.Time // IsZero returns false if k is nil or has a zero value -func IsZero(k interface{}) bool { +func IsZero(k any) bool { if k == nil { return true } @@ -55,10 +55,11 @@ func IsZero(k interface{}) bool { return t == nilTime || IsTimeZero(*t) case time.Time: return IsTimeZero(t) - case Zeroable: - return k.(Zeroable) == nil || k.(Zeroable).IsZero() case reflect.Value: // for go version less than 1.13 because reflect.Value has no method IsZero return IsValueZero(k.(reflect.Value)) + case Zeroable: + zeroable := k.(Zeroable) + return zeroable == nil || zeroable.IsZero() } return IsValueZero(reflect.ValueOf(k)) diff --git a/internal/utils/zero_test.go b/internal/utils/zero_test.go index c8c5ef1..b5486ef 100644 --- a/internal/utils/zero_test.go +++ b/internal/utils/zero_test.go @@ -13,11 +13,13 @@ import ( "github.com/stretchr/testify/assert" ) -type MyInt int -type ZeroStruct struct{} +type ( + MyInt int + ZeroStruct struct{} +) func TestZero(t *testing.T) { - var zeroValues = []interface{}{ + zeroValues := []any{ int8(0), int16(0), int(0), @@ -95,7 +97,7 @@ func NewNonZeroStruct() NonZeroStruct { func TestNoZero(t *testing.T) { now := time.Now() nonZeroStruct := NewNonZeroStruct() - var nonZeroValues = []interface{}{ + nonZeroValues := []any{ int8(1), int16(1), int(1), @@ -121,7 +123,7 @@ func TestNoZero(t *testing.T) { } func TestIsValueZero(t *testing.T) { - var zeroReflectValues = []reflect.Value{ + zeroReflectValues := []reflect.Value{ reflect.ValueOf(int8(0)), reflect.ValueOf(int16(0)), reflect.ValueOf(int(0)), diff --git a/log/logger.go b/log/logger.go index b8798c3..f90f74e 100644 --- a/log/logger.go +++ b/log/logger.go @@ -33,14 +33,14 @@ const ( // Logger is a logger interface type Logger interface { - Debug(v ...interface{}) - Debugf(format string, v ...interface{}) - Error(v ...interface{}) - Errorf(format string, v ...interface{}) - Info(v ...interface{}) - Infof(format string, v ...interface{}) - Warn(v ...interface{}) - Warnf(format string, v ...interface{}) + Debug(v ...any) + Debugf(format string, v ...any) + Error(v ...any) + Errorf(format string, v ...any) + Info(v ...any) + Infof(format string, v ...any) + Warn(v ...any) + Warnf(format string, v ...any) Level() LogLevel SetLevel(l LogLevel) @@ -55,28 +55,28 @@ var _ Logger = DiscardLogger{} type DiscardLogger struct{} // Debug empty implementation -func (DiscardLogger) Debug(v ...interface{}) {} +func (DiscardLogger) Debug(v ...any) {} // Debugf empty implementation -func (DiscardLogger) Debugf(format string, v ...interface{}) {} +func (DiscardLogger) Debugf(format string, v ...any) {} // Error empty implementation -func (DiscardLogger) Error(v ...interface{}) {} +func (DiscardLogger) Error(v ...any) {} // Errorf empty implementation -func (DiscardLogger) Errorf(format string, v ...interface{}) {} +func (DiscardLogger) Errorf(format string, v ...any) {} // Info empty implementation -func (DiscardLogger) Info(v ...interface{}) {} +func (DiscardLogger) Info(v ...any) {} // Infof empty implementation -func (DiscardLogger) Infof(format string, v ...interface{}) {} +func (DiscardLogger) Infof(format string, v ...any) {} // Warn empty implementation -func (DiscardLogger) Warn(v ...interface{}) {} +func (DiscardLogger) Warn(v ...any) {} // Warnf empty implementation -func (DiscardLogger) Warnf(format string, v ...interface{}) {} +func (DiscardLogger) Warnf(format string, v ...any) {} // Level empty implementation func (DiscardLogger) Level() LogLevel { @@ -119,65 +119,65 @@ func NewSimpleLogger2(out io.Writer, prefix string, flag int) *SimpleLogger { // NewSimpleLogger3 let you customrize your logger prefix and flag and logLevel func NewSimpleLogger3(out io.Writer, prefix string, flag int, l LogLevel) *SimpleLogger { return &SimpleLogger{ - DEBUG: log.New(out, fmt.Sprintf("%s [debug] ", prefix), flag), - ERR: log.New(out, fmt.Sprintf("%s [error] ", prefix), flag), - INFO: log.New(out, fmt.Sprintf("%s [info] ", prefix), flag), - WARN: log.New(out, fmt.Sprintf("%s [warn] ", prefix), flag), + DEBUG: log.New(out, prefix+" [debug] ", flag), + ERR: log.New(out, prefix+" [error] ", flag), + INFO: log.New(out, prefix+" [info] ", flag), + WARN: log.New(out, prefix+" [warn] ", flag), level: l, } } // Error implement ILogger -func (s *SimpleLogger) Error(v ...interface{}) { +func (s *SimpleLogger) Error(v ...any) { if s.level <= LOG_ERR { _ = s.ERR.Output(2, fmt.Sprintln(v...)) } } // Errorf implement ILogger -func (s *SimpleLogger) Errorf(format string, v ...interface{}) { +func (s *SimpleLogger) Errorf(format string, v ...any) { if s.level <= LOG_ERR { _ = s.ERR.Output(2, fmt.Sprintf(format, v...)) } } // Debug implement ILogger -func (s *SimpleLogger) Debug(v ...interface{}) { +func (s *SimpleLogger) Debug(v ...any) { if s.level <= LOG_DEBUG { _ = s.DEBUG.Output(2, fmt.Sprintln(v...)) } } // Debugf implement ILogger -func (s *SimpleLogger) Debugf(format string, v ...interface{}) { +func (s *SimpleLogger) Debugf(format string, v ...any) { if s.level <= LOG_DEBUG { _ = s.DEBUG.Output(2, fmt.Sprintf(format, v...)) } } // Info implement ILogger -func (s *SimpleLogger) Info(v ...interface{}) { +func (s *SimpleLogger) Info(v ...any) { if s.level <= LOG_INFO { _ = s.INFO.Output(2, fmt.Sprintln(v...)) } } // Infof implement ILogger -func (s *SimpleLogger) Infof(format string, v ...interface{}) { +func (s *SimpleLogger) Infof(format string, v ...any) { if s.level <= LOG_INFO { _ = s.INFO.Output(2, fmt.Sprintf(format, v...)) } } // Warn implement ILogger -func (s *SimpleLogger) Warn(v ...interface{}) { +func (s *SimpleLogger) Warn(v ...any) { if s.level <= LOG_WARNING { _ = s.WARN.Output(2, fmt.Sprintln(v...)) } } // Warnf implement ILogger -func (s *SimpleLogger) Warnf(format string, v ...interface{}) { +func (s *SimpleLogger) Warnf(format string, v ...any) { if s.level <= LOG_WARNING { _ = s.WARN.Output(2, fmt.Sprintf(format, v...)) } diff --git a/log/logger_context.go b/log/logger_context.go index 4680257..5213bea 100644 --- a/log/logger_context.go +++ b/log/logger_context.go @@ -23,10 +23,10 @@ type SQLLogger interface { type ContextLogger interface { SQLLogger - Debugf(format string, v ...interface{}) - Errorf(format string, v ...interface{}) - Infof(format string, v ...interface{}) - Warnf(format string, v ...interface{}) + Debugf(format string, v ...any) + Errorf(format string, v ...any) + Infof(format string, v ...any) + Warnf(format string, v ...any) Level() LogLevel SetLevel(l LogLevel) @@ -35,15 +35,15 @@ type ContextLogger interface { IsShowSQL() bool } -var ( - _ ContextLogger = &LoggerAdapter{} -) +var _ ContextLogger = &LoggerAdapter{} // enumerate all the context keys -var ( - SessionIDKey = "__xorm_session_id" - SessionKey = "__xorm_session_key" - SessionShowSQLKey = "__xorm_show_sql" +type sessionContextKey string + +const ( + SessionIDKey sessionContextKey = "__xorm_session_id" + SessionKey sessionContextKey = "__xorm_session_key" + SessionShowSQLKey sessionContextKey = "__xorm_show_sql" ) // LoggerAdapter wraps a Logger interface as LoggerContext interface @@ -76,22 +76,22 @@ func (l *LoggerAdapter) AfterSQL(ctx LogContext) { } // Debugf implements ContextLogger -func (l *LoggerAdapter) Debugf(format string, v ...interface{}) { +func (l *LoggerAdapter) Debugf(format string, v ...any) { l.logger.Debugf(format, v...) } // Errorf implements ContextLogger -func (l *LoggerAdapter) Errorf(format string, v ...interface{}) { +func (l *LoggerAdapter) Errorf(format string, v ...any) { l.logger.Errorf(format, v...) } // Infof implements ContextLogger -func (l *LoggerAdapter) Infof(format string, v ...interface{}) { +func (l *LoggerAdapter) Infof(format string, v ...any) { l.logger.Infof(format, v...) } // Warnf implements ContextLogger -func (l *LoggerAdapter) Warnf(format string, v ...interface{}) { +func (l *LoggerAdapter) Warnf(format string, v ...any) { l.logger.Warnf(format, v...) } diff --git a/log/syslogger.go b/log/syslogger.go index 4427258..ca6bb84 100644 --- a/log/syslogger.go +++ b/log/syslogger.go @@ -3,7 +3,6 @@ // license that can be found in the LICENSE file. //go:build !windows && !nacl && !plan9 -// +build !windows,!nacl,!plan9 package log @@ -26,42 +25,42 @@ func NewSyslogLogger(w *syslog.Writer) *SyslogLogger { } // Debug log content as Debug -func (s *SyslogLogger) Debug(v ...interface{}) { +func (s *SyslogLogger) Debug(v ...any) { _ = s.w.Debug(fmt.Sprint(v...)) } // Debugf log content as Debug and format -func (s *SyslogLogger) Debugf(format string, v ...interface{}) { +func (s *SyslogLogger) Debugf(format string, v ...any) { _ = s.w.Debug(fmt.Sprintf(format, v...)) } // Error log content as Error -func (s *SyslogLogger) Error(v ...interface{}) { +func (s *SyslogLogger) Error(v ...any) { _ = s.w.Err(fmt.Sprint(v...)) } // Errorf log content as Errorf and format -func (s *SyslogLogger) Errorf(format string, v ...interface{}) { +func (s *SyslogLogger) Errorf(format string, v ...any) { _ = s.w.Err(fmt.Sprintf(format, v...)) } // Info log content as Info -func (s *SyslogLogger) Info(v ...interface{}) { +func (s *SyslogLogger) Info(v ...any) { _ = s.w.Info(fmt.Sprint(v...)) } // Infof log content as Infof and format -func (s *SyslogLogger) Infof(format string, v ...interface{}) { +func (s *SyslogLogger) Infof(format string, v ...any) { _ = s.w.Info(fmt.Sprintf(format, v...)) } // Warn log content as Warn -func (s *SyslogLogger) Warn(v ...interface{}) { +func (s *SyslogLogger) Warn(v ...any) { _ = s.w.Warning(fmt.Sprint(v...)) } // Warnf log content as Warnf and format -func (s *SyslogLogger) Warnf(format string, v ...interface{}) { +func (s *SyslogLogger) Warnf(format string, v ...any) { _ = s.w.Warning(fmt.Sprintf(format, v...)) } diff --git a/migrate/migrate.go b/migrate/migrate.go index 2e5a6e3..14cdccb 100644 --- a/migrate/migrate.go +++ b/migrate/migrate.go @@ -218,20 +218,20 @@ func (m *Migrate) createMigrationTableIfNotExists() error { func (m *Migrate) migrationDidRun(mig *Migration) (bool, error) { tableName := m.db.TableName(m.options.TableName, true) - count, err := m.db.SQL(fmt.Sprintf("SELECT COUNT(*) FROM %s WHERE %s = ?", tableName, m.options.IDColumnName), mig.ID).Count() + count, err := m.db.SQL("SELECT COUNT(*) FROM "+tableName+" WHERE "+m.options.IDColumnName+" = ?", mig.ID).Count() return count > 0, err } func (m *Migrate) isFirstRun() (bool, error) { var count int tableName := m.db.TableName(m.options.TableName, true) - _, err := m.db.SQL(fmt.Sprintf("SELECT COUNT(*) FROM %s", tableName)).Get(&count) + _, err := m.db.SQL("SELECT COUNT(*) FROM " + tableName).Get(&count) return count == 0, err } func (m *Migrate) insertMigration(id string) error { tableName := m.db.TableName(m.options.TableName, true) - sql := fmt.Sprintf("INSERT INTO %s (%s) VALUES (?)", tableName, m.options.IDColumnName) + sql := "INSERT INTO " + tableName + " (" + m.options.IDColumnName + ") VALUES (?)" _, err := m.db.Exec(sql, id) return err } diff --git a/migrate/migrate_test.go b/migrate/migrate_test.go index 5e8ae7b..074d0a5 100644 --- a/migrate/migrate_test.go +++ b/migrate/migrate_test.go @@ -1,7 +1,6 @@ package migrate import ( - "fmt" "log" "os" "testing" @@ -139,6 +138,6 @@ func TestMissingID(t *testing.T) { } func tableCount(db *xorm.Engine, tableName string) (count int) { - _, _ = db.SQL(fmt.Sprintf("SELECT COUNT(*) FROM %s", tableName)).Get(&count) - return + _, _ = db.SQL("SELECT COUNT(*) FROM " + tableName).Get(&count) + return count } diff --git a/names/mapper.go b/names/mapper.go index 4f5910e..16b235d 100644 --- a/names/mapper.go +++ b/names/mapper.go @@ -27,7 +27,8 @@ type CacheMapper struct { // NewCacheMapper creates a cache mapper func NewCacheMapper(mapper Mapper) *CacheMapper { - return &CacheMapper{oriMapper: mapper, obj2tableCache: make(map[string]string), + return &CacheMapper{ + oriMapper: mapper, obj2tableCache: make(map[string]string), table2objCache: make(map[string]string), } } @@ -66,8 +67,7 @@ func (m *CacheMapper) Table2Obj(t string) string { // SameMapper implements Mapper and provides same name between struct and // database table -type SameMapper struct { -} +type SameMapper struct{} // Obj2Table implements Mapper func (m SameMapper) Obj2Table(o string) string { @@ -81,8 +81,7 @@ func (m SameMapper) Table2Obj(t string) string { // SnakeMapper implements IMapper and provides name translation between // struct and database table -type SnakeMapper struct { -} +type SnakeMapper struct{} func b2s(b []byte) string { return *(*string)(unsafe.Pointer(&b)) diff --git a/names/table_name_test.go b/names/table_name_test.go index 76da413..972e9b2 100644 --- a/names/table_name_test.go +++ b/names/table_name_test.go @@ -51,7 +51,7 @@ func (t *TestTableNameStruct) TableName() string { } func TestGetTableName(t *testing.T) { - var kases = []struct { + kases := []struct { mapper Mapper v reflect.Value expectedTableName string @@ -108,8 +108,7 @@ func TestGetTableName(t *testing.T) { } } -type OAuth2Application struct { -} +type OAuth2Application struct{} // TableName sets the table name to `oauth2_application` func (app *OAuth2Application) TableName() string { diff --git a/processors.go b/processors.go index 8697e30..938c078 100644 --- a/processors.go +++ b/processors.go @@ -54,12 +54,12 @@ type AfterLoadSessionProcessor interface { AfterLoad(*Session) } -type executedProcessorFunc func(*Session, interface{}) error +type executedProcessorFunc func(*Session, any) error type executedProcessor struct { fun executedProcessorFunc session *Session - bean interface{} + bean any } func (executor *executedProcessor) execute() error { @@ -77,13 +77,13 @@ func (session *Session) executeProcessors() error { return nil } -func cleanupProcessorsClosures(slices *[]func(interface{})) { +func cleanupProcessorsClosures(slices *[]func(any)) { if len(*slices) > 0 { - *slices = make([]func(interface{}), 0) + *slices = make([]func(any), 0) } } -func executeBeforeClosures(session *Session, bean interface{}) { +func executeBeforeClosures(session *Session, bean any) { // handle before delete processors for _, closure := range session.beforeClosures { closure(bean) @@ -91,27 +91,27 @@ func executeBeforeClosures(session *Session, bean interface{}) { cleanupProcessorsClosures(&session.beforeClosures) } -func executeBeforeSet(bean interface{}, fields []string, scanResults []interface{}) { +func executeBeforeSet(bean any, fields []string, scanResults []any) { if b, hasBeforeSet := bean.(BeforeSetProcessor); hasBeforeSet { for ii, key := range fields { - b.BeforeSet(key, Cell(scanResults[ii].(*interface{}))) + b.BeforeSet(key, Cell(scanResults[ii].(*any))) } } } -func executeAfterSet(bean interface{}, fields []string, scanResults []interface{}) { +func executeAfterSet(bean any, fields []string, scanResults []any) { if b, hasAfterSet := bean.(AfterSetProcessor); hasAfterSet { for ii, key := range fields { - b.AfterSet(key, Cell(scanResults[ii].(*interface{}))) + b.AfterSet(key, Cell(scanResults[ii].(*any))) } } } -func buildAfterProcessors(session *Session, bean interface{}) { +func buildAfterProcessors(session *Session, bean any) { // handle afterClosures for _, closure := range session.afterClosures { session.afterProcessors = append(session.afterProcessors, executedProcessor{ - fun: func(sess *Session, bean interface{}) error { + fun: func(sess *Session, bean any) error { closure(bean) return nil }, @@ -122,7 +122,7 @@ func buildAfterProcessors(session *Session, bean interface{}) { if a, has := bean.(AfterLoadProcessor); has { session.afterProcessors = append(session.afterProcessors, executedProcessor{ - fun: func(sess *Session, bean interface{}) error { + fun: func(sess *Session, bean any) error { a.AfterLoad() return nil }, @@ -133,7 +133,7 @@ func buildAfterProcessors(session *Session, bean interface{}) { if a, has := bean.(AfterLoadSessionProcessor); has { session.afterProcessors = append(session.afterProcessors, executedProcessor{ - fun: func(sess *Session, bean interface{}) error { + fun: func(sess *Session, bean any) error { a.AfterLoad(sess) return nil }, diff --git a/rows.go b/rows.go index 7393f57..59ffce8 100644 --- a/rows.go +++ b/rows.go @@ -20,13 +20,13 @@ type Rows struct { beanType reflect.Type } -func newRows(session *Session, bean interface{}) (*Rows, error) { +func newRows(session *Session, bean any) (*Rows, error) { rows := new(Rows) rows.session = session rows.beanType = reflect.Indirect(reflect.ValueOf(bean)).Type() var sqlStr string - var args []interface{} + var args []any var err error beanValue := reflect.ValueOf(bean) @@ -98,7 +98,7 @@ func (rows *Rows) Err() error { } // Scan row record to bean properties -func (rows *Rows) Scan(beans ...interface{}) error { +func (rows *Rows) Scan(beans ...any) error { if rows.Err() != nil { return rows.Err() } diff --git a/scan.go b/scan.go index 00cee4d..a3bdea1 100644 --- a/scan.go +++ b/scan.go @@ -18,9 +18,9 @@ import ( ) // genScanResultsByBeanNullabale generates scan result -func genScanResultsByBeanNullable(bean interface{}) (interface{}, bool, error) { +func genScanResultsByBeanNullable(bean any) (any, bool, error) { switch t := bean.(type) { - case *interface{}: + case *any: return t, false, nil case *sql.NullInt64, *sql.NullBool, *sql.NullFloat64, *sql.NullString, *sql.RawBytes, *[]byte: return t, false, nil @@ -71,9 +71,9 @@ func genScanResultsByBeanNullable(bean interface{}) (interface{}, bool, error) { } } -func genScanResultsByBean(bean interface{}) (interface{}, bool, error) { +func genScanResultsByBean(bean any) (any, bool, error) { switch t := bean.(type) { - case *interface{}: + case *any: return t, false, nil case *sql.NullInt64, *sql.NullBool, *sql.NullFloat64, *sql.NullString, *sql.RawBytes, @@ -129,8 +129,8 @@ func genScanResultsByBean(bean interface{}) (interface{}, bool, error) { } } -func (engine *Engine) scanStringInterface(rows *core.Rows, fields []string, types []*sql.ColumnType) ([]interface{}, error) { - scanResults := make([]interface{}, len(types)) +func (engine *Engine) scanStringInterface(rows *core.Rows, fields []string, types []*sql.ColumnType) ([]any, error) { + scanResults := make([]any, len(types)) for i := 0; i < len(types); i++ { var s sql.NullString scanResults[i] = &s @@ -143,13 +143,13 @@ func (engine *Engine) scanStringInterface(rows *core.Rows, fields []string, type } // scan is a wrap of driver.Scan but will automatically change the input values according requirements -func (engine *Engine) scan(rows *core.Rows, fields []string, types []*sql.ColumnType, vv ...interface{}) error { - scanResults := make([]interface{}, 0, len(types)) +func (engine *Engine) scan(rows *core.Rows, _ []string, types []*sql.ColumnType, vv ...any) error { + scanResults := make([]any, 0, len(types)) replaces := make([]bool, 0, len(types)) var err error for _, v := range vv { var replaced bool - var scanResult interface{} + var scanResult any switch t := v.(type) { case *big.Float, *time.Time, *sql.NullTime: scanResult = &sql.NullString{} @@ -193,8 +193,8 @@ func (engine *Engine) scan(rows *core.Rows, fields []string, types []*sql.Column return nil } -func (engine *Engine) scanInterfaces(rows *core.Rows, fields []string, types []*sql.ColumnType) ([]interface{}, error) { - scanResultContainers := make([]interface{}, len(types)) +func (engine *Engine) scanInterfaces(rows *core.Rows, fields []string, types []*sql.ColumnType) ([]any, error) { + scanResultContainers := make([]any, len(types)) for i := 0; i < len(types); i++ { scanResult, err := engine.driver.GenScanResult(types[i].DatabaseTypeName()) if err != nil { @@ -209,11 +209,11 @@ func (engine *Engine) scanInterfaces(rows *core.Rows, fields []string, types []* } //////////////////// -// row -> map[string]interface{} +// row -> map[string]any -func (engine *Engine) row2mapInterface(rows *core.Rows, types []*sql.ColumnType, fields []string) (map[string]interface{}, error) { - resultsMap := make(map[string]interface{}, len(fields)) - scanResultContainers := make([]interface{}, len(fields)) +func (engine *Engine) row2mapInterface(rows *core.Rows, types []*sql.ColumnType, fields []string) (map[string]any, error) { + resultsMap := make(map[string]any, len(fields)) + scanResultContainers := make([]any, len(fields)) for i := 0; i < len(fields); i++ { scanResult, err := engine.driver.GenScanResult(types[i].DatabaseTypeName()) if err != nil { @@ -236,7 +236,7 @@ func (engine *Engine) row2mapInterface(rows *core.Rows, types []*sql.ColumnType, } // ScanInterfaceMap scan result from *core.Rows and return a map -func (engine *Engine) ScanInterfaceMap(rows *core.Rows) (map[string]interface{}, error) { +func (engine *Engine) ScanInterfaceMap(rows *core.Rows) (map[string]any, error) { fields, err := rows.Columns() if err != nil { return nil, err @@ -250,7 +250,7 @@ func (engine *Engine) ScanInterfaceMap(rows *core.Rows) (map[string]interface{}, } // ScanInterfaceMaps scan results from *core.Rows and return a slice of map -func (engine *Engine) ScanInterfaceMaps(rows *core.Rows) (resultsSlice []map[string]interface{}, err error) { +func (engine *Engine) ScanInterfaceMaps(rows *core.Rows) (resultsSlice []map[string]any, err error) { fields, err := rows.Columns() if err != nil { return nil, err @@ -277,7 +277,7 @@ func (engine *Engine) ScanInterfaceMaps(rows *core.Rows) (resultsSlice []map[str // row -> map[string]string func (engine *Engine) row2mapStr(rows *core.Rows, types []*sql.ColumnType, fields []string) (map[string]string, error) { - scanResults := make([]interface{}, len(fields)) + scanResults := make([]any, len(fields)) for i := 0; i < len(fields); i++ { var s sql.NullString scanResults[i] = &s diff --git a/schemas/column.go b/schemas/column.go index 2da18ee..5683d13 100644 --- a/schemas/column.go +++ b/schemas/column.go @@ -77,7 +77,7 @@ func NewColumn(name, fieldName string, sqlType SQLType, len1, len2 int64, nullab } // ValueOf returns column's filed of struct's value -func (col *Column) ValueOf(bean interface{}) (*reflect.Value, error) { +func (col *Column) ValueOf(bean any) (*reflect.Value, error) { dataStruct := reflect.Indirect(reflect.ValueOf(bean)) return col.ValueOfV(&dataStruct) } @@ -100,7 +100,7 @@ func (col *Column) ValueOfV(dataStruct *reflect.Value) (*reflect.Value, error) { } // ConvertID converts id content to suitable type according column type -func (col *Column) ConvertID(sid string) (interface{}, error) { +func (col *Column) ConvertID(sid string) (any, error) { if col.SQLType.IsNumeric() { n, err := strconv.ParseInt(sid, 10, 64) if err != nil { diff --git a/schemas/pk.go b/schemas/pk.go index da3c789..6ed825e 100644 --- a/schemas/pk.go +++ b/schemas/pk.go @@ -12,10 +12,10 @@ import ( ) // PK represents primary key values -type PK []interface{} +type PK []any // NewPK creates primay keys -func NewPK(pks ...interface{}) *PK { +func NewPK(pks ...any) *PK { p := PK(pks) return &p } diff --git a/schemas/table.go b/schemas/table.go index 5c38cc7..be19127 100644 --- a/schemas/table.go +++ b/schemas/table.go @@ -156,7 +156,7 @@ func (table *Table) AddIndex(index *Index) { // IDOfV get id from one value of struct func (table *Table) IDOfV(rv reflect.Value) (PK, error) { v := reflect.Indirect(rv) - pk := make([]interface{}, len(table.PrimaryKeys)) + pk := make([]any, len(table.PrimaryKeys)) for i, col := range table.PKColumns() { var err error diff --git a/schemas/table_test.go b/schemas/table_test.go index f352675..c28fddc 100644 --- a/schemas/table_test.go +++ b/schemas/table_test.go @@ -71,9 +71,8 @@ func BenchmarkGetColumnIdxWithToLower(b *testing.B) { if c, ok := table.columnsMap[strings.ToLower(test.name)]; ok { if test.idx < len(c) { continue - } else { - b.Errorf("Bad idx in: %s, %d", test.name, test.idx) } + b.Errorf("Bad idx in: %s, %d", test.name, test.idx) } else { b.Errorf("Column not found: %s, %d", test.name, test.idx) } diff --git a/schemas/type.go b/schemas/type.go index aa0fce2..13c1247 100644 --- a/schemas/type.go +++ b/schemas/type.go @@ -324,7 +324,7 @@ func Type2SQLType(t reflect.Type) (st SQLType) { default: st = SQLType{Text, 0, 0} } - return + return st } // SQLType2Type convert default sql type change to go types diff --git a/session.go b/session.go index 4288d6a..7c7912f 100644 --- a/session.go +++ b/session.go @@ -15,6 +15,7 @@ import ( "io" "reflect" "strconv" + "xorm.io/xorm/contexts" "xorm.io/xorm/convert" "xorm.io/xorm/core" @@ -67,20 +68,20 @@ type Session struct { autoResetStatement bool // !nashtsai! storing these beans due to yet committed tx - afterInsertBeans map[interface{}]*[]func(interface{}) - afterUpdateBeans map[interface{}]*[]func(interface{}) - afterDeleteBeans map[interface{}]*[]func(interface{}) + afterInsertBeans map[any]*[]func(any) + afterUpdateBeans map[any]*[]func(any) + afterDeleteBeans map[any]*[]func(any) // -- - beforeClosures []func(interface{}) - afterClosures []func(interface{}) + beforeClosures []func(any) + afterClosures []func(any) afterProcessors []executedProcessor stmtCache map[string]*core.Stmt // key: full SQL string txStmtCache map[string]*core.Stmt // for tx statement lastSQL string - lastSQLArgs []interface{} + lastSQLArgs []any ctx context.Context sessionType sessionType @@ -121,17 +122,17 @@ func newSession(engine *Engine) *Session { autoResetStatement: true, prepareStmt: false, - afterInsertBeans: make(map[interface{}]*[]func(interface{})), - afterUpdateBeans: make(map[interface{}]*[]func(interface{})), - afterDeleteBeans: make(map[interface{}]*[]func(interface{})), - beforeClosures: make([]func(interface{}), 0), - afterClosures: make([]func(interface{}), 0), + afterInsertBeans: make(map[any]*[]func(any)), + afterUpdateBeans: make(map[any]*[]func(any)), + afterDeleteBeans: make(map[any]*[]func(any)), + beforeClosures: make([]func(any), 0), + afterClosures: make([]func(any), 0), afterProcessors: make([]executedProcessor, 0), stmtCache: make(map[string]*core.Stmt), txStmtCache: make(map[string]*core.Stmt), lastSQL: "", - lastSQLArgs: make([]interface{}, 0), + lastSQLArgs: make([]any, 0), sessionType: engineSession, } @@ -217,7 +218,7 @@ func (session *Session) Prepare() *Session { } // Before Apply before Processor, affected bean is passed to closure arg -func (session *Session) Before(closures func(interface{})) *Session { +func (session *Session) Before(closures func(any)) *Session { if closures != nil { session.beforeClosures = append(session.beforeClosures, closures) } @@ -225,7 +226,7 @@ func (session *Session) Before(closures func(interface{})) *Session { } // After Apply after Processor, affected bean is passed to closure arg -func (session *Session) After(closures func(interface{})) *Session { +func (session *Session) After(closures func(any)) *Session { if closures != nil { session.afterClosures = append(session.afterClosures, closures) } @@ -233,7 +234,7 @@ func (session *Session) After(closures func(interface{})) *Session { } // Table can input a string or pointer to struct for special a table to operate. -func (session *Session) Table(tableNameOrBean interface{}) *Session { +func (session *Session) Table(tableNameOrBean any) *Session { if err := session.statement.SetTable(tableNameOrBean); err != nil { session.statement.LastError = err } @@ -272,7 +273,7 @@ func (session *Session) Limit(limit int, start ...int) *Session { // OrderBy provide order by query condition, the input parameter is the content // after order by on a sql statement. -func (session *Session) OrderBy(order interface{}, args ...interface{}) *Session { +func (session *Session) OrderBy(order any, args ...any) *Session { session.statement.OrderBy(order, args...) return session } @@ -327,7 +328,7 @@ func (session *Session) NoCache() *Session { } // Join join_operator should be one of INNER, LEFT OUTER, CROSS etc - this will be prepended to JOIN -func (session *Session) Join(joinOperator string, tablename interface{}, condition interface{}, args ...interface{}) *Session { +func (session *Session) Join(joinOperator string, tablename, condition any, args ...any) *Session { session.statement.Join(joinOperator, tablename, condition, args...) return session } @@ -409,7 +410,7 @@ func getField(dataStruct *reflect.Value, table *schemas.Table, field *QueryedFie } // Cell cell is a result of one column field -type Cell *interface{} +type Cell *any func (session *Session) rows2Beans(rows *core.Rows, columnsSchema *ColumnsSchema, fields []string, types []*sql.ColumnType, table *schemas.Table, newElemFunc func([]string) reflect.Value, @@ -430,7 +431,7 @@ func (session *Session) rows2Beans(rows *core.Rows, columnsSchema *ColumnsSchema return err } session.afterProcessors = append(session.afterProcessors, executedProcessor{ - fun: func(*Session, interface{}) error { + fun: func(*Session, any) error { return sliceValueSetFunc(&newValue, pk) }, session: session, @@ -440,14 +441,14 @@ func (session *Session) rows2Beans(rows *core.Rows, columnsSchema *ColumnsSchema return rows.Err() } -func (session *Session) row2Slice(rows *core.Rows, fields []string, types []*sql.ColumnType, bean interface{}) ([]interface{}, error) { +func (session *Session) row2Slice(rows *core.Rows, fields []string, types []*sql.ColumnType, bean any) ([]any, error) { for _, closure := range session.beforeClosures { closure(bean) } - scanResults := make([]interface{}, len(fields)) + scanResults := make([]any, len(fields)) for i := 0; i < len(fields); i++ { - var cell interface{} + var cell any scanResults[i] = &cell } if err := session.engine.scan(rows, fields, types, scanResults...); err != nil { @@ -459,7 +460,7 @@ func (session *Session) row2Slice(rows *core.Rows, fields []string, types []*sql return scanResults, nil } -func setJSON(fieldValue *reflect.Value, fieldType reflect.Type, scanResult interface{}) error { +func setJSON(fieldValue *reflect.Value, fieldType reflect.Type, scanResult any) error { bs, ok := convert.AsBytes(scanResult) if !ok { return fmt.Errorf("unsupported database data type: %#v", scanResult) @@ -489,7 +490,7 @@ func setJSON(fieldValue *reflect.Value, fieldType reflect.Type, scanResult inter return nil } -func asKind(vv reflect.Value, tp reflect.Type) (interface{}, error) { +func asKind(vv reflect.Value, tp reflect.Type) (any, error) { switch tp.Kind() { case reflect.Ptr: return asKind(vv.Elem(), tp.Elem()) @@ -529,10 +530,8 @@ func asKind(vv reflect.Value, tp reflect.Type) (interface{}, error) { var uint8ZeroValue = reflect.ValueOf(uint8(0)) -func (session *Session) convertBeanField(col *schemas.Column, fieldValue *reflect.Value, - scanResult interface{}, table *schemas.Table, -) error { - v, ok := scanResult.(*interface{}) +func (session *Session) convertBeanField(col *schemas.Column, fieldValue *reflect.Value, scanResult any) error { + v, ok := scanResult.(*any) if ok { scanResult = *v } @@ -584,7 +583,7 @@ func (session *Session) convertBeanField(col *schemas.Column, fieldValue *reflec } else { e = fieldValue.Elem() } - if err := session.convertBeanField(col, &e, scanResult, table); err != nil { + if err := session.convertBeanField(col, &e, scanResult); err != nil { return err } if fieldValue.IsNil() { @@ -698,7 +697,7 @@ func (session *Session) convertBeanField(col *schemas.Column, fieldValue *reflec return convert.AssignValue(fieldValue.Addr(), scanResult) } -func (session *Session) slice2Bean(scanResults []interface{}, columnsSchema *ColumnsSchema, fields []string, bean interface{}, dataStruct *reflect.Value, table *schemas.Table) (schemas.PK, error) { +func (session *Session) slice2Bean(scanResults []any, columnsSchema *ColumnsSchema, fields []string, bean any, dataStruct *reflect.Value, table *schemas.Table) (schemas.PK, error) { defer func() { executeAfterSet(bean, fields, scanResults) }() @@ -718,7 +717,7 @@ func (session *Session) slice2Bean(scanResults []interface{}, columnsSchema *Col continue } - if err := session.convertBeanField(col, fieldValue, scanResults[i], table); err != nil { + if err := session.convertBeanField(col, fieldValue, scanResults[i]); err != nil { return nil, err } if col.IsPrimaryKey { @@ -729,13 +728,13 @@ func (session *Session) slice2Bean(scanResults []interface{}, columnsSchema *Col } // saveLastSQL stores executed query information -func (session *Session) saveLastSQL(sql string, args ...interface{}) { +func (session *Session) saveLastSQL(sql string, args ...any) { session.lastSQL = sql session.lastSQLArgs = args } // LastSQL returns last query information -func (session *Session) LastSQL() (string, []interface{}) { +func (session *Session) LastSQL() (string, []any) { return session.lastSQL, session.lastSQLArgs } diff --git a/session_cols.go b/session_cols.go index ca3589a..8ef7ca2 100644 --- a/session_cols.go +++ b/session_cols.go @@ -12,7 +12,7 @@ import ( "xorm.io/xorm/schemas" ) -func setColumnInt(bean interface{}, col *schemas.Column, t int64) { +func setColumnInt(bean any, col *schemas.Column, t int64) { v, err := col.ValueOf(bean) if err != nil { return @@ -27,7 +27,7 @@ func setColumnInt(bean interface{}, col *schemas.Column, t int64) { } } -func setColumnTime(bean interface{}, col *schemas.Column, t time.Time) { +func setColumnTime(bean any, col *schemas.Column, t time.Time) { v, err := col.ValueOf(bean) if err != nil { return @@ -44,9 +44,9 @@ func setColumnTime(bean interface{}, col *schemas.Column, t time.Time) { } } -func getFlagForColumn(m map[string]bool, col *schemas.Column) (val bool, has bool) { +func getFlagForColumn(m map[string]bool, col *schemas.Column) bool { if len(m) == 0 { - return false, false + return false } n := len(col.Name) @@ -56,27 +56,27 @@ func getFlagForColumn(m map[string]bool, col *schemas.Column) (val bool, has boo continue } if strings.EqualFold(mk, col.Name) { - return m[mk], true + return true } } - return false, false + return false } // Incr provides a query string like "count = count + 1" -func (session *Session) Incr(column string, arg ...interface{}) *Session { +func (session *Session) Incr(column string, arg ...any) *Session { session.statement.Incr(column, arg...) return session } // Decr provides a query string like "count = count - 1" -func (session *Session) Decr(column string, arg ...interface{}) *Session { +func (session *Session) Decr(column string, arg ...any) *Session { session.statement.Decr(column, arg...) return session } // SetExpr provides a query string like "column = {expression}" -func (session *Session) SetExpr(column string, expression interface{}) *Session { +func (session *Session) SetExpr(column string, expression any) *Session { session.statement.SetExpr(column, expression) return session } diff --git a/session_cond.go b/session_cond.go index 25d1714..6688abf 100644 --- a/session_cond.go +++ b/session_cond.go @@ -8,43 +8,43 @@ import "xorm.io/builder" // SQL provides raw sql input parameter. When you have a complex SQL statement // and cannot use Where, Id, In and etc. Methods to describe, you can use SQL. -func (session *Session) SQL(query interface{}, args ...interface{}) *Session { +func (session *Session) SQL(query any, args ...any) *Session { session.statement.SQL(query, args...) return session } // Where provides custom query condition. -func (session *Session) Where(query interface{}, args ...interface{}) *Session { +func (session *Session) Where(query any, args ...any) *Session { session.statement.Where(query, args...) return session } // And provides custom query condition. -func (session *Session) And(query interface{}, args ...interface{}) *Session { +func (session *Session) And(query any, args ...any) *Session { session.statement.And(query, args...) return session } // Or provides custom query condition. -func (session *Session) Or(query interface{}, args ...interface{}) *Session { +func (session *Session) Or(query any, args ...any) *Session { session.statement.Or(query, args...) return session } // ID provides converting id as a query condition -func (session *Session) ID(id interface{}) *Session { +func (session *Session) ID(id any) *Session { session.statement.ID(id) return session } // In provides a query string like "id in (1, 2, 3)" -func (session *Session) In(column string, args ...interface{}) *Session { +func (session *Session) In(column string, args ...any) *Session { session.statement.In(column, args...) return session } // NotIn provides a query string like "id in (1, 2, 3)" -func (session *Session) NotIn(column string, args ...interface{}) *Session { +func (session *Session) NotIn(column string, args ...any) *Session { session.statement.NotIn(column, args...) return session } diff --git a/session_delete.go b/session_delete.go index 7336040..3f7586c 100644 --- a/session_delete.go +++ b/session_delete.go @@ -16,7 +16,7 @@ import ( // ErrNeedDeletedCond delete needs less one condition error var ErrNeedDeletedCond = errors.New("Delete action needs at least one condition") -func (session *Session) cacheDelete(table *schemas.Table, tableName, sqlStr string, args ...interface{}) error { +func (session *Session) cacheDelete(table *schemas.Table, tableName, sqlStr string, args ...any) error { if table == nil || session.tx != nil { return ErrCacheFailed @@ -49,7 +49,7 @@ func (session *Session) cacheDelete(table *schemas.Table, tableName, sqlStr stri if len(resultsSlice) > 0 { for _, data := range resultsSlice { var id int64 - var pk schemas.PK = make([]interface{}, 0) + var pk schemas.PK = make([]any, 0) for _, col := range pkColumns { if v, ok := data[col.Name]; !ok { return errors.New("no id") @@ -85,28 +85,30 @@ func (session *Session) cacheDelete(table *schemas.Table, tableName, sqlStr stri // Delete records, bean's non-empty fields are conditions // At least one condition must be set. -func (session *Session) Delete(beans ...interface{}) (int64, error) { +func (session *Session) Delete(beans ...any) (int64, error) { return session.delete(beans, true) } // Truncate records, bean's non-empty fields are conditions // In contrast to Delete this method allows deletes without conditions. -func (session *Session) Truncate(beans ...interface{}) (int64, error) { +func (session *Session) Truncate(beans ...any) (int64, error) { return session.delete(beans, false) } -func (session *Session) delete(beans []interface{}, mustHaveConditions bool) (int64, error) { +func (session *Session) delete(beans []any, mustHaveConditions bool) (int64, error) { if session.isAutoClose { defer session.Close() } + defer cleanupProcessorsClosures(&session.afterClosures) + if session.statement.LastError != nil { return 0, session.statement.LastError } var ( err error - bean interface{} + bean any ) if len(beans) > 0 { bean = beans[0] @@ -116,7 +118,7 @@ func (session *Session) delete(beans []interface{}, mustHaveConditions bool) (in executeBeforeClosures(session, bean) - if processor, ok := interface{}(bean).(BeforeDeleteProcessor); ok { + if processor, ok := bean.(BeforeDeleteProcessor); ok { processor.BeforeDelete() } @@ -148,13 +150,13 @@ func (session *Session) delete(beans []interface{}, mustHaveConditions bool) (in } colName := deletedColumn.Name - session.afterClosures = append(session.afterClosures, func(bean interface{}) { + session.afterClosures = append(session.afterClosures, func(bean any) { col := table.GetColumn(colName) setColumnTime(bean, col, t) }) } - argsForCache := make([]interface{}, 0, len(deleteSQLWriter.Args())*2) + argsForCache := make([]any, 0, len(deleteSQLWriter.Args())*2) copy(argsForCache, deleteSQLWriter.Args()) argsForCache = append(deleteSQLWriter.Args(), argsForCache...) @@ -174,7 +176,7 @@ func (session *Session) delete(beans []interface{}, mustHaveConditions bool) (in for _, closure := range session.afterClosures { closure(bean) } - if processor, ok := interface{}(bean).(AfterDeleteProcessor); ok { + if processor, ok := bean.(AfterDeleteProcessor); ok { processor.AfterDelete() } } else { @@ -183,12 +185,12 @@ func (session *Session) delete(beans []interface{}, mustHaveConditions bool) (in if value, has := session.afterDeleteBeans[beans[0]]; has && value != nil { *value = append(*value, session.afterClosures...) } else { - afterClosures := make([]func(interface{}), lenAfterClosures) + afterClosures := make([]func(any), lenAfterClosures) copy(afterClosures, session.afterClosures) session.afterDeleteBeans[bean] = &afterClosures } } else { - if _, ok := interface{}(bean).(AfterDeleteProcessor); ok { + if _, ok := bean.(AfterDeleteProcessor); ok { session.afterDeleteBeans[bean] = nil } } diff --git a/session_exist.go b/session_exist.go index b5e4a65..eef95b2 100644 --- a/session_exist.go +++ b/session_exist.go @@ -5,7 +5,7 @@ package xorm // Exist returns true if the record exist otherwise return false -func (session *Session) Exist(bean ...interface{}) (bool, error) { +func (session *Session) Exist(bean ...any) (bool, error) { if session.isAutoClose { defer session.Close() } diff --git a/session_find.go b/session_find.go index 1026910..e833753 100644 --- a/session_find.go +++ b/session_find.go @@ -26,7 +26,7 @@ const ( // Find retrieve records from table, condiBeans's non-empty fields // are conditions. beans could be []Struct, []*Struct, map[int64]Struct // map[int64]*Struct -func (session *Session) Find(rowsSlicePtr interface{}, condiBean ...interface{}) error { +func (session *Session) Find(rowsSlicePtr any, condiBean ...any) error { if session.isAutoClose { defer session.Close() } @@ -34,7 +34,7 @@ func (session *Session) Find(rowsSlicePtr interface{}, condiBean ...interface{}) } // FindAndCount find the results and also return the counts -func (session *Session) FindAndCount(rowsSlicePtr interface{}, condiBean ...interface{}) (int64, error) { +func (session *Session) FindAndCount(rowsSlicePtr any, condiBean ...any) (int64, error) { if session.isAutoClose { defer session.Close() } @@ -78,7 +78,7 @@ func (session *Session) FindAndCount(rowsSlicePtr interface{}, condiBean ...inte return session.Unscoped().Count() } -func (session *Session) find(rowsSlicePtr interface{}, condiBean ...interface{}) error { +func (session *Session) find(rowsSlicePtr any, condiBean ...any) error { defer session.resetStatement() if session.statement.LastError != nil { return session.statement.LastError @@ -221,7 +221,7 @@ func ParseColumnsSchema(fieldNames []string, types []*sql.ColumnType, table *sch return &columnsSchema } -func (session *Session) noCacheFind(table *schemas.Table, containerValue reflect.Value, sqlStr string, args ...interface{}) error { +func (session *Session) noCacheFind(table *schemas.Table, containerValue reflect.Value, sqlStr string, args ...any) error { elemType := containerValue.Type().Elem() var isPointer bool if elemType.Kind() == reflect.Ptr { @@ -332,7 +332,7 @@ func (session *Session) noCacheFind(table *schemas.Table, containerValue reflect return rows.Err() } -func (session *Session) cacheFind(t reflect.Type, sqlStr string, rowsSlicePtr interface{}, args ...interface{}) (err error) { +func (session *Session) cacheFind(t reflect.Type, sqlStr string, rowsSlicePtr any, args ...any) (err error) { if !session.canCache() || utils.IndexNoCase(sqlStr, "having") != -1 || utils.IndexNoCase(sqlStr, "group by") != -1 { @@ -376,7 +376,7 @@ func (session *Session) cacheFind(t reflect.Type, sqlStr string, rowsSlicePtr in if err != nil { return err } - var pk schemas.PK = make([]interface{}, len(table.PrimaryKeys)) + var pk schemas.PK = make([]any, len(table.PrimaryKeys)) for i, col := range table.PKColumns() { pk[i], err = col.ConvertID(res[i]) if err != nil { @@ -403,7 +403,7 @@ func (session *Session) cacheFind(t reflect.Type, sqlStr string, rowsSlicePtr in ididxes := make(map[string]int) var ides []schemas.PK - temps := make([]interface{}, len(ids)) + temps := make([]any, len(ids)) for idx, id := range ids { sid, err := id.ToString() @@ -413,17 +413,15 @@ func (session *Session) cacheFind(t reflect.Type, sqlStr string, rowsSlicePtr in bean := cacher.GetBean(tableName, sid) // fix issue #894 - isHit := func() (ht bool) { + isHit := func() bool { if bean == nil { - ht = false - return + return false } ckb := reflect.ValueOf(bean).Elem().Type() - ht = ckb == t - if !ht && t.Kind() == reflect.Ptr { - ht = t.Elem() == ckb + if ckb == t { + return true } - return + return t.Kind() == reflect.Ptr && t.Elem() == ckb } if !isHit() { ides = append(ides, id) @@ -460,7 +458,7 @@ func (session *Session) cacheFind(t reflect.Type, sqlStr string, rowsSlicePtr in session.engine.DatabaseTZ, ) if len(table.PrimaryKeys) == 1 { - ff := make([]interface{}, 0, len(ides)) + ff := make([]any, 0, len(ides)) for _, ie := range ides { ff = append(ff, ie[0]) } @@ -521,7 +519,7 @@ func (session *Session) cacheFind(t reflect.Type, sqlStr string, rowsSlicePtr in key := ids[j] keyType := sliceValue.Type().Key() keyValue := reflect.New(keyType) - var ikey interface{} + var ikey any if len(key) == 1 { if err := convert.AssignValue(keyValue, key[0]); err != nil { return err diff --git a/session_get.go b/session_get.go index 76ccaf9..deb7e2d 100644 --- a/session_get.go +++ b/session_get.go @@ -20,21 +20,19 @@ import ( "xorm.io/xorm/schemas" ) -var ( - // ErrObjectIsNil return error of object is nil - ErrObjectIsNil = errors.New("object should not be nil") -) +// ErrObjectIsNil return error of object is nil +var ErrObjectIsNil = errors.New("object should not be nil") // Get retrieve one record from database, bean's non-empty fields // will be as conditions -func (session *Session) Get(beans ...interface{}) (bool, error) { +func (session *Session) Get(beans ...any) (bool, error) { if session.isAutoClose { defer session.Close() } return session.get(beans...) } -func isPtrOfTime(v interface{}) bool { +func isPtrOfTime(v any) bool { if _, ok := v.(*time.Time); ok { return true } @@ -47,7 +45,7 @@ func isPtrOfTime(v interface{}) bool { return el.Type().ConvertibleTo(schemas.TimeType) } -func (session *Session) get(beans ...interface{}) (bool, error) { +func (session *Session) get(beans ...any) (bool, error) { defer session.resetStatement() if session.statement.LastError != nil { @@ -66,7 +64,7 @@ func (session *Session) get(beans ...interface{}) (bool, error) { return false, ErrObjectIsNil } - var isStruct = beanValue.Elem().Kind() == reflect.Struct && !isPtrOfTime(beans[0]) + isStruct := beanValue.Elem().Kind() == reflect.Struct && !isPtrOfTime(beans[0]) if isStruct { if err := session.statement.SetRefBean(beans[0]); err != nil { return false, err @@ -74,7 +72,7 @@ func (session *Session) get(beans ...interface{}) (bool, error) { } var sqlStr string - var args []interface{} + var args []any var err error if session.statement.RawSQL == "" { @@ -129,7 +127,7 @@ func (session *Session) get(beans ...interface{}) (bool, error) { return true, nil } -func isScannableStruct(bean interface{}, typeLen int) bool { +func isScannableStruct(bean any, typeLen int) bool { switch bean.(type) { case *time.Time: return false @@ -143,7 +141,7 @@ func isScannableStruct(bean interface{}, typeLen int) bool { return true } -func (session *Session) nocacheGet(beanKind reflect.Kind, table *schemas.Table, beans []interface{}, sqlStr string, args ...interface{}) (bool, error) { +func (session *Session) nocacheGet(beanKind reflect.Kind, table *schemas.Table, beans []any, sqlStr string, args ...any) (bool, error) { rows, err := session.queryRows(sqlStr, args...) if err != nil { return false, err @@ -174,7 +172,7 @@ func (session *Session) nocacheGet(beanKind reflect.Kind, table *schemas.Table, return true, session.executeProcessors() } -func (session *Session) scan(rows *core.Rows, table *schemas.Table, firstBeanKind reflect.Kind, beans []interface{}, columnsSchema *ColumnsSchema, types []*sql.ColumnType, fields []string) error { +func (session *Session) scan(rows *core.Rows, table *schemas.Table, firstBeanKind reflect.Kind, beans []any, columnsSchema *ColumnsSchema, types []*sql.ColumnType, fields []string) error { if len(beans) == 1 { bean := beans[0] switch firstBeanKind { @@ -204,7 +202,7 @@ func (session *Session) scan(rows *core.Rows, table *schemas.Table, firstBeanKin return session.engine.scan(rows, fields, types, beans...) } -func (session *Session) getSlice(rows *core.Rows, types []*sql.ColumnType, fields []string, bean interface{}) error { +func (session *Session) getSlice(rows *core.Rows, types []*sql.ColumnType, fields []string, bean any) error { switch t := bean.(type) { case *[]string: res, err := session.engine.scanStringInterface(rows, fields, types) @@ -212,7 +210,7 @@ func (session *Session) getSlice(rows *core.Rows, types []*sql.ColumnType, field return err } - var needAppend = len(*t) == 0 // both support slice is empty or has been initlized + needAppend := len(*t) == 0 // both support slice is empty or has been initlized for i, r := range res { if needAppend { *t = append(*t, r.(*sql.NullString).String) @@ -221,12 +219,12 @@ func (session *Session) getSlice(rows *core.Rows, types []*sql.ColumnType, field } } return nil - case *[]interface{}: + case *[]any: scanResults, err := session.engine.scanInterfaces(rows, fields, types) if err != nil { return err } - var needAppend = len(*t) == 0 + needAppend := len(*t) == 0 for ii := range fields { s, err := convert.Interface2Interface(session.engine.DatabaseTZ, scanResults[ii]) if err != nil { @@ -244,7 +242,7 @@ func (session *Session) getSlice(rows *core.Rows, types []*sql.ColumnType, field } } -func (session *Session) getMap(rows *core.Rows, types []*sql.ColumnType, fields []string, bean interface{}) error { +func (session *Session) getMap(rows *core.Rows, types []*sql.ColumnType, fields []string, bean any) error { switch t := bean.(type) { case *map[string]string: scanResults, err := session.engine.scanStringInterface(rows, fields, types) @@ -255,7 +253,7 @@ func (session *Session) getMap(rows *core.Rows, types []*sql.ColumnType, fields (*t)[key] = scanResults[ii].(*sql.NullString).String } return nil - case *map[string]interface{}: + case *map[string]any: scanResults, err := session.engine.scanInterfaces(rows, fields, types) if err != nil { return err @@ -273,7 +271,7 @@ func (session *Session) getMap(rows *core.Rows, types []*sql.ColumnType, fields } } -func (session *Session) cacheGet(bean interface{}, sqlStr string, args ...interface{}) (has bool, err error) { +func (session *Session) cacheGet(bean any, sqlStr string, args ...any) (has bool, err error) { // if has no reftable, then don't use cache currently if !session.canCache() { return false, ErrCacheFailed @@ -294,7 +292,7 @@ func (session *Session) cacheGet(bean interface{}, sqlStr string, args ...interf table := session.statement.RefTable ids, err := caches.GetCacheSql(cacher, tableName, newsql, args) if err != nil { - var res = make([]string, len(table.PrimaryKeys)) + res := make([]string, len(table.PrimaryKeys)) rows, err := session.NoCache().queryRows(newsql, args...) if err != nil { return false, err @@ -313,7 +311,7 @@ func (session *Session) cacheGet(bean interface{}, sqlStr string, args ...interf return false, ErrCacheFailed } - var pk schemas.PK = make([]interface{}, len(table.PrimaryKeys)) + var pk schemas.PK = make([]any, len(table.PrimaryKeys)) for i, col := range table.PKColumns() { if col.SQLType.IsText() { pk[i] = res[i] @@ -349,7 +347,7 @@ func (session *Session) cacheGet(bean interface{}, sqlStr string, args ...interf cacheBean := cacher.GetBean(tableName, sid) if cacheBean == nil { cacheBean = bean - has, err = session.nocacheGet(reflect.Struct, table, []interface{}{cacheBean}, sqlStr, args...) + has, err = session.nocacheGet(reflect.Struct, table, []any{cacheBean}, sqlStr, args...) if err != nil || !has { return has, err } diff --git a/session_insert.go b/session_insert.go index 7cc1524..fa94839 100644 --- a/session_insert.go +++ b/session_insert.go @@ -23,7 +23,7 @@ import ( var ErrNoElementsOnSlice = errors.New("no element on slice when insert") // Insert insert one or more beans -func (session *Session) Insert(beans ...interface{}) (int64, error) { +func (session *Session) Insert(beans ...any) (int64, error) { var affected int64 var err error @@ -41,9 +41,9 @@ func (session *Session) Insert(beans ...interface{}) (int64, error) { var cnt int64 var err error switch v := bean.(type) { - case map[string]interface{}: + case map[string]any: cnt, err = session.insertMapInterface(v) - case []map[string]interface{}: + case []map[string]any: cnt, err = session.insertMultipleMapInterface(v) case map[string]string: cnt, err = session.insertMapString(v) @@ -66,7 +66,9 @@ func (session *Session) Insert(beans ...interface{}) (int64, error) { return affected, err } -func (session *Session) insertMultipleStruct(rowsSlicePtr interface{}) (int64, error) { +func (session *Session) insertMultipleStruct(rowsSlicePtr any) (int64, error) { + defer cleanupProcessorsClosures(&session.afterClosures) + sliceValue := reflect.Indirect(reflect.ValueOf(rowsSlicePtr)) if sliceValue.Kind() != reflect.Slice { return 0, errors.New("needs a pointer to a slice") @@ -90,7 +92,7 @@ func (session *Session) insertMultipleStruct(rowsSlicePtr interface{}) (int64, e size = sliceValue.Len() colNames []string colMultiPlaces []string - args []interface{} + args []any ) for i := 0; i < size; i++ { @@ -111,7 +113,7 @@ func (session *Session) insertMultipleStruct(rowsSlicePtr interface{}) (int64, e closure(elemValue) } - if processor, ok := interface{}(elemValue).(BeforeInsertProcessor); ok { + if processor, ok := elemValue.(BeforeInsertProcessor); ok { processor.BeforeInsert() } // -- @@ -144,7 +146,7 @@ func (session *Session) insertMultipleStruct(rowsSlicePtr interface{}) (int64, e continue } // !satorunooshie! set fieldValue as nil when column is nullable and zero-value - if _, ok := getFlagForColumn(session.statement.NullableMap, col); ok { + if getFlagForColumn(session.statement.NullableMap, col) { if col.Nullable && utils.IsValueZero(fieldValue) { var nilValue *int fieldValue = reflect.ValueOf(nilValue) @@ -158,14 +160,14 @@ func (session *Session) insertMultipleStruct(rowsSlicePtr interface{}) (int64, e args = append(args, val) colName := col.Name - session.afterClosures = append(session.afterClosures, func(bean interface{}) { + session.afterClosures = append(session.afterClosures, func(bean any) { col := table.GetColumn(colName) setColumnTime(bean, col, t) }) } else if col.IsVersion && session.statement.CheckVersion { args = append(args, 1) colName := col.Name - session.afterClosures = append(session.afterClosures, func(bean interface{}) { + session.afterClosures = append(session.afterClosures, func(bean any) { col := table.GetColumn(colName) setColumnInt(bean, col, 1) }) @@ -197,7 +199,7 @@ func (session *Session) insertMultipleStruct(rowsSlicePtr interface{}) (int64, e return 0, err } - _ = session.cacheInsert(tableName) + session.cacheInsert(tableName) lenAfterClosures := len(session.afterClosures) for i := 0; i < size; i++ { @@ -217,7 +219,7 @@ func (session *Session) insertMultipleStruct(rowsSlicePtr interface{}) (int64, e if value, has := session.afterInsertBeans[elemValue]; has && value != nil { *value = append(*value, session.afterClosures...) } else { - afterClosures := make([]func(interface{}), lenAfterClosures) + afterClosures := make([]func(any), lenAfterClosures) copy(afterClosures, session.afterClosures) session.afterInsertBeans[elemValue] = &afterClosures } @@ -234,7 +236,7 @@ func (session *Session) insertMultipleStruct(rowsSlicePtr interface{}) (int64, e } // InsertMulti insert multiple records -func (session *Session) InsertMulti(rowsSlicePtr interface{}) (int64, error) { +func (session *Session) InsertMulti(rowsSlicePtr any) (int64, error) { if session.isAutoClose { defer session.Close() } @@ -247,7 +249,9 @@ func (session *Session) InsertMulti(rowsSlicePtr interface{}) (int64, error) { return session.insertMultipleStruct(rowsSlicePtr) } -func (session *Session) insertStruct(bean interface{}) (int64, error) { +func (session *Session) insertStruct(bean any) (int64, error) { + defer cleanupProcessorsClosures(&session.afterClosures) + if err := session.statement.SetRefBean(bean); err != nil { return 0, err } @@ -261,7 +265,7 @@ func (session *Session) insertStruct(bean interface{}) (int64, error) { } cleanupProcessorsClosures(&session.beforeClosures) // cleanup after used - if processor, ok := interface{}(bean).(BeforeInsertProcessor); ok { + if processor, ok := bean.(BeforeInsertProcessor); ok { processor.BeforeInsert() } @@ -279,12 +283,12 @@ func (session *Session) insertStruct(bean interface{}) (int64, error) { } sqlStr = session.engine.dialect.Quoter().Replace(sqlStr) - handleAfterInsertProcessorFunc := func(bean interface{}) { + handleAfterInsertProcessorFunc := func(bean any) { if session.isAutoCommit { for _, closure := range session.afterClosures { closure(bean) } - if processor, ok := interface{}(bean).(AfterInsertProcessor); ok { + if processor, ok := bean.(AfterInsertProcessor); ok { processor.AfterInsert() } } else { @@ -293,12 +297,12 @@ func (session *Session) insertStruct(bean interface{}) (int64, error) { if value, has := session.afterInsertBeans[bean]; has && value != nil { *value = append(*value, session.afterClosures...) } else { - afterClosures := make([]func(interface{}), lenAfterClosures) + afterClosures := make([]func(any), lenAfterClosures) copy(afterClosures, session.afterClosures) session.afterInsertBeans[bean] = &afterClosures } } else { - if _, ok := interface{}(bean).(AfterInsertProcessor); ok { + if _, ok := bean.(AfterInsertProcessor); ok { session.afterInsertBeans[bean] = nil } } @@ -309,7 +313,7 @@ func (session *Session) insertStruct(bean interface{}) (int64, error) { // if there is auto increment column and driver don't support return it if len(table.AutoIncrement) > 0 && !session.engine.driver.Features().SupportReturnInsertedID { var sql string - var newArgs []interface{} + var newArgs []any var needCommit bool var id int64 if session.engine.dialect.URI().DBType == schemas.ORACLE || session.engine.dialect.URI().DBType == schemas.DAMENG { @@ -354,7 +358,7 @@ func (session *Session) insertStruct(bean interface{}) (int64, error) { defer handleAfterInsertProcessorFunc(bean) - _ = session.cacheInsert(tableName) + session.cacheInsert(tableName) if table.Version != "" && session.statement.CheckVersion { verValue, err := table.VersionColumn().ValueOf(bean) @@ -384,7 +388,7 @@ func (session *Session) insertStruct(bean interface{}) (int64, error) { defer handleAfterInsertProcessorFunc(bean) - _ = session.cacheInsert(tableName) + session.cacheInsert(tableName) if table.Version != "" && session.statement.CheckVersion { verValue, err := table.VersionColumn().ValueOf(bean) @@ -425,7 +429,7 @@ func (session *Session) insertStruct(bean interface{}) (int64, error) { // The in parameter bean must a struct or a point to struct. The return // parameter is inserted and error // Deprecated: Please use Insert directly -func (session *Session) InsertOne(bean interface{}) (int64, error) { +func (session *Session) InsertOne(bean any) (int64, error) { if session.isAutoClose { defer session.Close() } @@ -433,24 +437,23 @@ func (session *Session) InsertOne(bean interface{}) (int64, error) { return session.insertStruct(bean) } -func (session *Session) cacheInsert(table string) error { +func (session *Session) cacheInsert(table string) { if !session.statement.UseCache { - return nil + return } cacher := session.engine.cacherMgr.GetCacher(table) if cacher == nil { - return nil + return } session.engine.logger.Debugf("[cache] clear SQL: %v", table) cacher.ClearIds(table) - return nil } // genInsertColumns generates insert needed columns -func (session *Session) genInsertColumns(bean interface{}) ([]string, []interface{}, error) { +func (session *Session) genInsertColumns(bean any) ([]string, []any, error) { table := session.statement.RefTable colNames := make([]string, 0, len(table.ColumnsSeq())) - args := make([]interface{}, 0, len(table.ColumnsSeq())) + args := make([]any, 0, len(table.ColumnsSeq())) for _, col := range table.Columns() { if col.MapType == schemas.ONLYFROMDB { @@ -492,7 +495,7 @@ func (session *Session) genInsertColumns(bean interface{}) ([]string, []interfac } // !evalphobia! set fieldValue as nil when column is nullable and zero-value - if _, ok := getFlagForColumn(session.statement.NullableMap, col); ok { + if getFlagForColumn(session.statement.NullableMap, col) { if col.Nullable && utils.IsValueZero(fieldValue) { var nilValue *int fieldValue = reflect.ValueOf(nilValue) @@ -508,7 +511,7 @@ func (session *Session) genInsertColumns(bean interface{}) ([]string, []interfac args = append(args, val) colName := col.Name - session.afterClosures = append(session.afterClosures, func(bean interface{}) { + session.afterClosures = append(session.afterClosures, func(bean any) { col := table.GetColumn(colName) setColumnTime(bean, col, t) }) @@ -527,7 +530,7 @@ func (session *Session) genInsertColumns(bean interface{}) ([]string, []interfac return colNames, args, nil } -func (session *Session) insertMapInterface(m map[string]interface{}) (int64, error) { +func (session *Session) insertMapInterface(m map[string]any) (int64, error) { if len(m) == 0 { return 0, ErrParamsType } @@ -546,7 +549,7 @@ func (session *Session) insertMapInterface(m map[string]interface{}) (int64, err } sort.Strings(columns) - args := make([]interface{}, 0, len(m)) + args := make([]any, 0, len(m)) for _, colName := range columns { args = append(args, m[colName]) } @@ -554,7 +557,7 @@ func (session *Session) insertMapInterface(m map[string]interface{}) (int64, err return session.insertMap(columns, args) } -func (session *Session) insertMultipleMapInterface(maps []map[string]interface{}) (int64, error) { +func (session *Session) insertMultipleMapInterface(maps []map[string]any) (int64, error) { if len(maps) == 0 { return 0, ErrNoElementsOnSlice } @@ -573,9 +576,9 @@ func (session *Session) insertMultipleMapInterface(maps []map[string]interface{} } sort.Strings(columns) - argss := make([][]interface{}, 0, len(maps)) + argss := make([][]any, 0, len(maps)) for _, m := range maps { - args := make([]interface{}, 0, len(m)) + args := make([]any, 0, len(m)) for _, colName := range columns { args = append(args, m[colName]) } @@ -605,7 +608,7 @@ func (session *Session) insertMapString(m map[string]string) (int64, error) { sort.Strings(columns) - args := make([]interface{}, 0, len(m)) + args := make([]any, 0, len(m)) for _, colName := range columns { args = append(args, m[colName]) } @@ -632,9 +635,9 @@ func (session *Session) insertMultipleMapString(maps []map[string]string) (int64 } sort.Strings(columns) - argss := make([][]interface{}, 0, len(maps)) + argss := make([][]any, 0, len(maps)) for _, m := range maps { - args := make([]interface{}, 0, len(m)) + args := make([]any, 0, len(m)) for _, colName := range columns { args = append(args, m[colName]) } @@ -644,7 +647,7 @@ func (session *Session) insertMultipleMapString(maps []map[string]string) (int64 return session.insertMultipleMap(columns, argss) } -func (session *Session) insertMap(columns []string, args []interface{}) (int64, error) { +func (session *Session) insertMap(columns []string, args []any) (int64, error) { tableName := session.statement.TableName() if len(tableName) == 0 { return 0, ErrTableNotFound @@ -656,9 +659,7 @@ func (session *Session) insertMap(columns []string, args []interface{}) (int64, } sql = session.engine.dialect.Quoter().Replace(sql) - if err := session.cacheInsert(tableName); err != nil { - return 0, err - } + session.cacheInsert(tableName) res, err := session.exec(sql, args...) if err != nil { @@ -671,7 +672,7 @@ func (session *Session) insertMap(columns []string, args []interface{}) (int64, return affected, nil } -func (session *Session) insertMultipleMap(columns []string, argss [][]interface{}) (int64, error) { +func (session *Session) insertMultipleMap(columns []string, argss [][]any) (int64, error) { tableName := session.statement.TableName() if len(tableName) == 0 { return 0, ErrTableNotFound @@ -683,9 +684,7 @@ func (session *Session) insertMultipleMap(columns []string, argss [][]interface{ } sql = session.engine.dialect.Quoter().Replace(sql) - if err := session.cacheInsert(tableName); err != nil { - return 0, err - } + session.cacheInsert(tableName) res, err := session.exec(sql, args...) if err != nil { diff --git a/session_iterate.go b/session_iterate.go index ff42555..8d765f4 100644 --- a/session_iterate.go +++ b/session_iterate.go @@ -11,18 +11,18 @@ import ( ) // IterFunc only use by Iterate -type IterFunc func(idx int, bean interface{}) error +type IterFunc func(idx int, bean any) error // Rows return sql.Rows compatible Rows obj, as a forward Iterator object for iterating record by record, bean's non-empty fields // are conditions. -func (session *Session) Rows(bean interface{}) (*Rows, error) { +func (session *Session) Rows(bean any) (*Rows, error) { return newRows(session, bean) } // Iterate record by record handle records from table, condiBeans's non-empty fields // are conditions. beans could be []Struct, []*Struct, map[int64]Struct // map[int64]*Struct -func (session *Session) Iterate(bean interface{}, fun IterFunc) error { +func (session *Session) Iterate(bean any, fun IterFunc) error { if session.isAutoClose { defer session.Close() } @@ -69,7 +69,7 @@ func (session *Session) BufferSize(size int) *Session { return session } -func (session *Session) bufferIterate(bean interface{}, fun IterFunc) error { +func (session *Session) bufferIterate(bean any, fun IterFunc) error { bufferSize := session.statement.BufferSize pLimitN := session.statement.LimitN if pLimitN != nil && bufferSize > *pLimitN { diff --git a/session_raw.go b/session_raw.go index 99f6be9..06ede41 100644 --- a/session_raw.go +++ b/session_raw.go @@ -11,7 +11,7 @@ import ( "xorm.io/xorm/core" ) -func (session *Session) queryPreprocess(sqlStr *string, paramStr ...interface{}) { +func (session *Session) queryPreprocess(sqlStr *string, paramStr ...any) { for _, filter := range session.engine.dialect.Filters() { *sqlStr = filter.Do(session.ctx, *sqlStr) } @@ -20,7 +20,7 @@ func (session *Session) queryPreprocess(sqlStr *string, paramStr ...interface{}) session.lastSQLArgs = paramStr } -func (session *Session) queryRows(sqlStr string, args ...interface{}) (*core.Rows, error) { +func (session *Session) queryRows(sqlStr string, args ...any) (*core.Rows, error) { defer session.resetStatement() if session.statement.LastError != nil { return nil, session.statement.LastError @@ -64,12 +64,12 @@ func (session *Session) queryRows(sqlStr string, args ...interface{}) (*core.Row return session.tx.QueryContext(session.ctx, sqlStr, args...) } -func (session *Session) queryRow(sqlStr string, args ...interface{}) *core.Row { +func (session *Session) queryRow(sqlStr string, args ...any) *core.Row { return core.NewRow(session.queryRows(sqlStr, args...)) } // Query runs a raw sql and return records as []map[string][]byte -func (session *Session) Query(sqlOrArgs ...interface{}) ([]map[string][]byte, error) { +func (session *Session) Query(sqlOrArgs ...any) ([]map[string][]byte, error) { if session.isAutoClose { defer session.Close() } @@ -89,7 +89,7 @@ func (session *Session) Query(sqlOrArgs ...interface{}) ([]map[string][]byte, er } // QueryString runs a raw sql and return records as []map[string]string -func (session *Session) QueryString(sqlOrArgs ...interface{}) ([]map[string]string, error) { +func (session *Session) QueryString(sqlOrArgs ...any) ([]map[string]string, error) { if session.isAutoClose { defer session.Close() } @@ -109,7 +109,7 @@ func (session *Session) QueryString(sqlOrArgs ...interface{}) ([]map[string]stri } // QuerySliceString runs a raw sql and return records as [][]string -func (session *Session) QuerySliceString(sqlOrArgs ...interface{}) ([][]string, error) { +func (session *Session) QuerySliceString(sqlOrArgs ...any) ([][]string, error) { if session.isAutoClose { defer session.Close() } @@ -128,8 +128,8 @@ func (session *Session) QuerySliceString(sqlOrArgs ...interface{}) ([][]string, return session.engine.ScanStringSlices(rows) } -// QueryInterface runs a raw sql and return records as []map[string]interface{} -func (session *Session) QueryInterface(sqlOrArgs ...interface{}) ([]map[string]interface{}, error) { +// QueryInterface runs a raw sql and return records as []map[string]any +func (session *Session) QueryInterface(sqlOrArgs ...any) ([]map[string]any, error) { if session.isAutoClose { defer session.Close() } @@ -148,7 +148,7 @@ func (session *Session) QueryInterface(sqlOrArgs ...interface{}) ([]map[string]i return session.engine.ScanInterfaceMaps(rows) } -func (session *Session) exec(sqlStr string, args ...interface{}) (sql.Result, error) { +func (session *Session) exec(sqlStr string, args ...any) (sql.Result, error) { defer session.resetStatement() session.queryPreprocess(&sqlStr, args...) @@ -179,7 +179,7 @@ func (session *Session) exec(sqlStr string, args ...interface{}) (sql.Result, er } // Exec raw sql -func (session *Session) Exec(sqlOrArgs ...interface{}) (sql.Result, error) { +func (session *Session) Exec(sqlOrArgs ...any) (sql.Result, error) { if session.isAutoClose { defer session.Close() } diff --git a/session_schema.go b/session_schema.go index 4bb0b85..0c8644b 100644 --- a/session_schema.go +++ b/session_schema.go @@ -8,7 +8,6 @@ import ( "bufio" "context" "database/sql" - "fmt" "io" "os" "strings" @@ -28,7 +27,7 @@ func (session *Session) Ping() error { } // CreateTable create a table according a bean -func (session *Session) CreateTable(bean interface{}) error { +func (session *Session) CreateTable(bean any) error { if session.isAutoClose { defer session.Close() } @@ -36,7 +35,7 @@ func (session *Session) CreateTable(bean interface{}) error { return session.createTable(bean) } -func (session *Session) createTable(bean interface{}) error { +func (session *Session) createTable(bean any) error { if err := session.statement.SetRefBean(bean); err != nil { return err } @@ -67,7 +66,7 @@ func (session *Session) createTable(bean interface{}) error { } // CreateIndexes create indexes -func (session *Session) CreateIndexes(bean interface{}) error { +func (session *Session) CreateIndexes(bean any) error { if session.isAutoClose { defer session.Close() } @@ -75,7 +74,7 @@ func (session *Session) CreateIndexes(bean interface{}) error { return session.createIndexes(bean) } -func (session *Session) createIndexes(bean interface{}) error { +func (session *Session) createIndexes(bean any) error { if err := session.statement.SetRefBean(bean); err != nil { return err } @@ -91,14 +90,14 @@ func (session *Session) createIndexes(bean interface{}) error { } // CreateUniques create uniques -func (session *Session) CreateUniques(bean interface{}) error { +func (session *Session) CreateUniques(bean any) error { if session.isAutoClose { defer session.Close() } return session.createUniques(bean) } -func (session *Session) createUniques(bean interface{}) error { +func (session *Session) createUniques(bean any) error { if err := session.statement.SetRefBean(bean); err != nil { return err } @@ -114,7 +113,7 @@ func (session *Session) createUniques(bean interface{}) error { } // DropIndexes drop indexes -func (session *Session) DropIndexes(bean interface{}) error { +func (session *Session) DropIndexes(bean any) error { if session.isAutoClose { defer session.Close() } @@ -122,7 +121,7 @@ func (session *Session) DropIndexes(bean interface{}) error { return session.dropIndexes(bean) } -func (session *Session) dropIndexes(bean interface{}) error { +func (session *Session) dropIndexes(bean any) error { if err := session.statement.SetRefBean(bean); err != nil { return err } @@ -138,7 +137,7 @@ func (session *Session) dropIndexes(bean interface{}) error { } // DropTable drop table will drop table if exist, if drop failed, it will return error -func (session *Session) DropTable(beanOrTableName interface{}) error { +func (session *Session) DropTable(beanOrTableName any) error { if session.isAutoClose { defer session.Close() } @@ -146,7 +145,7 @@ func (session *Session) DropTable(beanOrTableName interface{}) error { return session.dropTable(beanOrTableName) } -func (session *Session) dropTable(beanOrTableName interface{}) error { +func (session *Session) dropTable(beanOrTableName any) error { tableName := session.engine.TableName(beanOrTableName) sqlStr, checkIfExist := session.engine.dialect.DropTableSQL(session.engine.TableName(tableName, true)) if !checkIfExist { @@ -186,7 +185,7 @@ func (session *Session) dropTable(beanOrTableName interface{}) error { } // IsTableExist if a table is exist -func (session *Session) IsTableExist(beanOrTableName interface{}) (bool, error) { +func (session *Session) IsTableExist(beanOrTableName any) (bool, error) { if session.isAutoClose { defer session.Close() } @@ -201,7 +200,7 @@ func (session *Session) isTableExist(tableName string) (bool, error) { } // IsTableEmpty if table have any records -func (session *Session) IsTableEmpty(bean interface{}) (bool, error) { +func (session *Session) IsTableEmpty(bean any) (bool, error) { if session.isAutoClose { defer session.Close() } @@ -210,7 +209,7 @@ func (session *Session) IsTableEmpty(bean interface{}) (bool, error) { func (session *Session) isTableEmpty(tableName string) (bool, error) { var total int64 - sqlStr := fmt.Sprintf("select count(*) from %s", session.engine.Quote(session.engine.TableName(tableName, true))) + sqlStr := "select count(*) from " + session.engine.Quote(session.engine.TableName(tableName, true)) err := session.queryRow(sqlStr).Scan(&total) if err != nil { if err == sql.ErrNoRows { @@ -313,6 +312,8 @@ func (session *Session) Import(r io.Reader) ([]sql.Result, error) { } func (session *Session) IndexHint(op, forType, indexerOrColName string) *Session { - session.statement.IndexHint(op, forType, indexerOrColName) + if err := session.statement.IndexHint(op, forType, indexerOrColName); err != nil { + session.statement.LastError = err + } return session } diff --git a/session_stats.go b/session_stats.go index be98e46..f194cde 100644 --- a/session_stats.go +++ b/session_stats.go @@ -12,7 +12,7 @@ import ( // Count counts the records. bean's non-empty fields // are conditions. -func (session *Session) Count(bean ...interface{}) (int64, error) { +func (session *Session) Count(bean ...any) (int64, error) { if session.isAutoClose { defer session.Close() } @@ -32,7 +32,7 @@ func (session *Session) Count(bean ...interface{}) (int64, error) { } // sum call sum some column. bean's non-empty fields are conditions. -func (session *Session) sum(res interface{}, bean interface{}, columnNames ...string) error { +func (session *Session) sum(res, bean any, columnNames ...string) error { if session.isAutoClose { defer session.Close() } @@ -59,23 +59,23 @@ func (session *Session) sum(res interface{}, bean interface{}, columnNames ...st } // Sum call sum some column. bean's non-empty fields are conditions. -func (session *Session) Sum(bean interface{}, columnName string) (res float64, err error) { +func (session *Session) Sum(bean any, columnName string) (res float64, err error) { return res, session.sum(&res, bean, columnName) } // SumInt call sum some column. bean's non-empty fields are conditions. -func (session *Session) SumInt(bean interface{}, columnName string) (res int64, err error) { +func (session *Session) SumInt(bean any, columnName string) (res int64, err error) { return res, session.sum(&res, bean, columnName) } // Sums call sum some columns. bean's non-empty fields are conditions. -func (session *Session) Sums(bean interface{}, columnNames ...string) ([]float64, error) { +func (session *Session) Sums(bean any, columnNames ...string) ([]float64, error) { res := make([]float64, len(columnNames)) return res, session.sum(&res, bean, columnNames...) } // SumsInt sum specify columns and return as []int64 instead of []float64 -func (session *Session) SumsInt(bean interface{}, columnNames ...string) ([]int64, error) { +func (session *Session) SumsInt(bean any, columnNames ...string) ([]int64, error) { res := make([]int64, len(columnNames)) return res, session.sum(&res, bean, columnNames...) } diff --git a/session_tx.go b/session_tx.go index 4fa5689..30bd15e 100644 --- a/session_tx.go +++ b/session_tx.go @@ -44,7 +44,7 @@ func (session *Session) Commit() error { } // handle processors after tx committed - closureCallFunc := func(closuresPtr *[]func(interface{}), bean interface{}) { + closureCallFunc := func(closuresPtr *[]func(any), bean any) { if closuresPtr != nil { for _, closure := range *closuresPtr { closure(bean) @@ -55,27 +55,27 @@ func (session *Session) Commit() error { for bean, closuresPtr := range session.afterInsertBeans { closureCallFunc(closuresPtr, bean) - if processor, ok := interface{}(bean).(AfterInsertProcessor); ok { + if processor, ok := bean.(AfterInsertProcessor); ok { processor.AfterInsert() } } for bean, closuresPtr := range session.afterUpdateBeans { closureCallFunc(closuresPtr, bean) - if processor, ok := interface{}(bean).(AfterUpdateProcessor); ok { + if processor, ok := bean.(AfterUpdateProcessor); ok { processor.AfterUpdate() } } for bean, closuresPtr := range session.afterDeleteBeans { closureCallFunc(closuresPtr, bean) - if processor, ok := interface{}(bean).(AfterDeleteProcessor); ok { + if processor, ok := bean.(AfterDeleteProcessor); ok { processor.AfterDelete() } } - cleanUpFunc := func(slices *map[interface{}]*[]func(interface{})) { + cleanUpFunc := func(slices *map[any]*[]func(any)) { if len(*slices) > 0 { - *slices = make(map[interface{}]*[]func(interface{})) + *slices = make(map[any]*[]func(any)) } } cleanUpFunc(&session.afterInsertBeans) diff --git a/session_update.go b/session_update.go index b3640ad..70db35d 100644 --- a/session_update.go +++ b/session_update.go @@ -18,12 +18,12 @@ var ( ErrNoColumnsTobeUpdated = statements.ErrNoColumnsTobeUpdated ) -func (session *Session) genAutoCond(condiBean interface{}) (builder.Cond, error) { +func (session *Session) genAutoCond(condiBean any) (builder.Cond, error) { if session.statement.NoAutoCondition { return builder.NewCond(), nil } - if c, ok := condiBean.(map[string]interface{}); ok { + if c, ok := condiBean.(map[string]any); ok { eq := make(builder.Eq) for k, v := range c { eq[session.engine.Quote(k)] = v @@ -60,11 +60,13 @@ func (session *Session) genAutoCond(condiBean interface{}) (builder.Cond, error) // 1.bool will defaultly be updated content nor conditions // You should call UseBool if you have bool to use. // 2.float32 & float64 may be not inexact as conditions -func (session *Session) Update(bean interface{}, condiBean ...interface{}) (int64, error) { +func (session *Session) Update(bean any, condiBean ...any) (int64, error) { if session.isAutoClose { defer session.Close() } + defer cleanupProcessorsClosures(&session.afterClosures) + defer session.resetStatement() if session.statement.LastError != nil { @@ -79,13 +81,13 @@ func (session *Session) Update(bean interface{}, condiBean ...interface{}) (int6 closure(bean) } cleanupProcessorsClosures(&session.beforeClosures) // cleanup after used - if processor, ok := interface{}(bean).(BeforeUpdateProcessor); ok { + if processor, ok := bean.(BeforeUpdateProcessor); ok { processor.BeforeUpdate() } // -- var colNames []string - var args []interface{} + var args []any var err error isMap := t.Kind() == reflect.Map isStruct := t.Kind() == reflect.Struct @@ -109,7 +111,7 @@ func (session *Session) Update(bean interface{}, condiBean ...interface{}) (int6 } } else if isMap { colNames = make([]string, 0) - args = make([]interface{}, 0) + args = make([]any, 0) bValue := reflect.Indirect(reflect.ValueOf(bean)) for _, v := range bValue.MapKeys() { @@ -139,7 +141,7 @@ func (session *Session) Update(bean interface{}, condiBean ...interface{}) (int6 colName := col.Name if isStruct { - session.afterClosures = append(session.afterClosures, func(bean interface{}) { + session.afterClosures = append(session.afterClosures, func(bean any) { col := table.GetColumn(colName) setColumnTime(bean, col, t) }) @@ -159,13 +161,7 @@ func (session *Session) Update(bean interface{}, condiBean ...interface{}) (int6 } } else if table != nil { if col := table.DeletedColumn(); col != nil && !session.statement.GetUnscoped() { // tag "deleted" is enabled - autoCond1 := session.statement.CondDeleted(col) - - if autoCond == nil { - autoCond = autoCond1 - } else { - autoCond = autoCond.And(autoCond1) - } + autoCond = session.statement.CondDeleted(col) } } @@ -213,7 +209,7 @@ func (session *Session) Update(bean interface{}, condiBean ...interface{}) (int6 for _, closure := range session.afterClosures { closure(bean) } - if processor, ok := interface{}(bean).(AfterUpdateProcessor); ok { + if processor, ok := bean.(AfterUpdateProcessor); ok { session.engine.logger.Debugf("[event] %v has after update processor", tableName) processor.AfterUpdate() } @@ -223,13 +219,13 @@ func (session *Session) Update(bean interface{}, condiBean ...interface{}) (int6 if value, has := session.afterUpdateBeans[bean]; has && value != nil { *value = append(*value, session.afterClosures...) } else { - afterClosures := make([]func(interface{}), lenAfterClosures) + afterClosures := make([]func(any), lenAfterClosures) copy(afterClosures, session.afterClosures) // FIXME: if bean is a map type, it will panic because map cannot be as map key session.afterUpdateBeans[bean] = &afterClosures } } else { - if _, ok := interface{}(bean).(AfterUpdateProcessor); ok { + if _, ok := bean.(AfterUpdateProcessor); ok { session.afterUpdateBeans[bean] = nil } } @@ -240,10 +236,10 @@ func (session *Session) Update(bean interface{}, condiBean ...interface{}) (int6 return res.RowsAffected() } -func (session *Session) genUpdateColumns(bean interface{}) ([]string, []interface{}, error) { +func (session *Session) genUpdateColumns(bean any) ([]string, []any, error) { table := session.statement.RefTable colNames := make([]string, 0, len(table.ColumnsSeq())) - args := make([]interface{}, 0, len(table.ColumnsSeq())) + args := make([]any, 0, len(table.ColumnsSeq())) for _, col := range table.Columns() { if !col.IsVersion && !col.IsCreated && !col.IsUpdated { @@ -283,7 +279,7 @@ func (session *Session) genUpdateColumns(bean interface{}) ([]string, []interfac } // !evalphobia! set fieldValue as nil when column is nullable and zero-value - if _, ok := getFlagForColumn(session.statement.NullableMap, col); ok { + if getFlagForColumn(session.statement.NullableMap, col) { if col.Nullable && utils.IsValueZero(fieldValue) { var nilValue *int fieldValue = reflect.ValueOf(nilValue) @@ -299,7 +295,7 @@ func (session *Session) genUpdateColumns(bean interface{}) ([]string, []interfac args = append(args, val) colName := col.Name - session.afterClosures = append(session.afterClosures, func(bean interface{}) { + session.afterClosures = append(session.afterClosures, func(bean any) { col := table.GetColumn(colName) setColumnTime(bean, col, t) }) diff --git a/sync.go b/sync.go index b8b827d..c1755bf 100644 --- a/sync.go +++ b/sync.go @@ -26,14 +26,14 @@ type SyncResult struct{} // Sync the new struct changes to database, this method will automatically add // table, column, index, unique. but will not delete or change anything. // If you change some field, you should change the database manually. -func (engine *Engine) Sync(beans ...interface{}) error { +func (engine *Engine) Sync(beans ...any) error { session := engine.NewSession() defer session.Close() return session.Sync(beans...) } // SyncWithOptions sync the database schemas according options and table structs -func (engine *Engine) SyncWithOptions(opts SyncOptions, beans ...interface{}) (*SyncResult, error) { +func (engine *Engine) SyncWithOptions(opts SyncOptions, beans ...any) (*SyncResult, error) { session := engine.NewSession() defer session.Close() return session.SyncWithOptions(opts, beans...) @@ -41,18 +41,18 @@ func (engine *Engine) SyncWithOptions(opts SyncOptions, beans ...interface{}) (* // Sync2 synchronize structs to database tables // Depricated -func (engine *Engine) Sync2(beans ...interface{}) error { +func (engine *Engine) Sync2(beans ...any) error { return engine.Sync(beans...) } // Sync2 synchronize structs to database tables // Depricated -func (session *Session) Sync2(beans ...interface{}) error { +func (session *Session) Sync2(beans ...any) error { return session.Sync(beans...) } // Sync synchronize structs to database tables -func (session *Session) Sync(beans ...interface{}) error { +func (session *Session) Sync(beans ...any) error { _, err := session.SyncWithOptions(SyncOptions{ WarnIfDatabaseColumnMissed: false, IgnoreConstrains: false, @@ -62,7 +62,7 @@ func (session *Session) Sync(beans ...interface{}) error { return err } -func (session *Session) SyncWithOptions(opts SyncOptions, beans ...interface{}) (*SyncResult, error) { +func (session *Session) SyncWithOptions(opts SyncOptions, beans ...any) (*SyncResult, error) { engine := session.engine if session.isAutoClose { diff --git a/tags/parser.go b/tags/parser.go index 8b97160..63abe58 100644 --- a/tags/parser.go +++ b/tags/parser.go @@ -7,7 +7,6 @@ package tags import ( "encoding/gob" "errors" - "fmt" "reflect" "strings" "sync" @@ -151,7 +150,7 @@ func (parser *Parser) getSQLTypeByType(t reflect.Type) (schemas.SQLType, error) return pkCols[0].SQLType, nil } if len(pkCols) > 1 { - return schemas.SQLType{}, fmt.Errorf("unsupported mulitiple primary key on cascade") + return schemas.SQLType{}, errors.New("unsupported mulitiple primary key on cascade") } } } @@ -160,12 +159,11 @@ func (parser *Parser) getSQLTypeByType(t reflect.Type) (schemas.SQLType, error) func (parser *Parser) parseFieldWithNoTag(fieldIndex int, field reflect.StructField, fieldValue reflect.Value) (*schemas.Column, error) { var sqlType schemas.SQLType + hasAddrConversion := false if fieldValue.CanAddr() { - if _, ok := fieldValue.Addr().Interface().(convert.Conversion); ok { - sqlType = schemas.SQLType{Name: schemas.Text} - } + _, hasAddrConversion = fieldValue.Addr().Interface().(convert.Conversion) } - if _, ok := fieldValue.Interface().(convert.Conversion); ok { + if _, ok := fieldValue.Interface().(convert.Conversion); ok || hasAddrConversion { sqlType = schemas.SQLType{Name: schemas.Text} } else { var err error @@ -179,7 +177,8 @@ func (parser *Parser) parseFieldWithNoTag(fieldIndex int, field reflect.StructFi sqlType.DefaultLength2, true) col.FieldIndex = []int{fieldIndex} - if field.Type.Kind() == reflect.Int64 && (strings.ToUpper(col.FieldName) == "ID" || strings.HasSuffix(strings.ToUpper(col.FieldName), ".ID")) { + fieldNameUpper := strings.ToUpper(col.FieldName) + if field.Type.Kind() == reflect.Int64 && (strings.EqualFold(col.FieldName, "ID") || strings.HasSuffix(fieldNameUpper, ".ID")) { col.IsAutoIncrement = true col.IsPrimaryKey = true col.Nullable = false diff --git a/tags/parser_test.go b/tags/parser_test.go index be6ea6f..75e02d1 100644 --- a/tags/parser_test.go +++ b/tags/parser_test.go @@ -104,6 +104,7 @@ func TestUnexportField(t *testing.T) { private int // unexported fields will be ignored Public int } + _ = VanilaStruct{}.private table, err := parser.Parse(reflect.ValueOf(new(VanilaStruct))) assert.NoError(t, err) assert.EqualValues(t, "vanila_struct", table.Name) @@ -118,6 +119,7 @@ func TestUnexportField(t *testing.T) { private int `xorm:"private"` // unexported fields will be ignored Public int `xorm:"-"` } + _ = TaggedStruct{}.private table, err = parser.Parse(reflect.ValueOf(new(TaggedStruct))) assert.NoError(t, err) assert.EqualValues(t, "tagged_struct", table.Name) diff --git a/tags/tag.go b/tags/tag.go index cfe35c8..cf4bdd1 100644 --- a/tags/tag.go +++ b/tags/tag.go @@ -169,7 +169,7 @@ func PKTagHandler(ctx *Context) error { // NULLTagHandler describes null tag handler func NULLTagHandler(ctx *Context) error { - ctx.col.Nullable = strings.ToUpper(ctx.preTag) != "NOT" + ctx.col.Nullable = !strings.EqualFold(ctx.preTag, "NOT") return nil } @@ -378,10 +378,9 @@ func ExtendsTagHandler(ctx *Context) error { col.FieldName = fmt.Sprintf("%v.%v", ctx.col.FieldName, col.FieldName) col.FieldIndex = append(ctx.col.FieldIndex, col.FieldIndex...) - tagPrefix := ctx.col.FieldName if len(ctx.params) > 0 { col.Nullable = isPtr - tagPrefix = strings.Trim(ctx.params[0], "'") + tagPrefix := strings.Trim(ctx.params[0], "'") if col.IsPrimaryKey { col.Name = ctx.col.FieldName col.IsPrimaryKey = false diff --git a/tests/engine_dm_test.go b/tests/engine_dm_test.go index 5b25af2..3961e73 100644 --- a/tests/engine_dm_test.go +++ b/tests/engine_dm_test.go @@ -3,7 +3,6 @@ // license that can be found in the LICENSE file. //go:build dm -// +build dm package tests diff --git a/tests/engine_test.go b/tests/engine_test.go index 0bbf191..1da929b 100644 --- a/tests/engine_test.go +++ b/tests/engine_test.go @@ -6,6 +6,7 @@ package tests import ( "context" + "errors" "fmt" "os" "testing" @@ -58,7 +59,7 @@ func TestAutoTransaction(t *testing.T) { engine := testEngine.(*xorm.Engine) // will success - _, err := engine.Transaction(func(session *xorm.Session) (interface{}, error) { + _, err := engine.Transaction(func(session *xorm.Session) (any, error) { _, err := session.Insert(TestTx{Msg: "hi"}) assert.NoError(t, err) @@ -71,11 +72,11 @@ func TestAutoTransaction(t *testing.T) { assert.EqualValues(t, true, has) // will rollback - _, err = engine.Transaction(func(session *xorm.Session) (interface{}, error) { + _, err = engine.Transaction(func(session *xorm.Session) (any, error) { _, err := session.Insert(TestTx{Msg: "hello"}) assert.NoError(t, err) - return nil, fmt.Errorf("rollback") + return nil, errors.New("rollback") }) assert.Error(t, err) @@ -84,7 +85,7 @@ func TestAutoTransaction(t *testing.T) { assert.EqualValues(t, false, has) } -func assertSync(t *testing.T, beans ...interface{}) { +func assertSync(t *testing.T, beans ...any) { for _, bean := range beans { t.Run(testEngine.TableName(bean, true), func(t *testing.T) { assert.NoError(t, testEngine.DropTables(bean)) @@ -252,7 +253,7 @@ func TestDBVersion(t *testing.T) { version, err := testEngine.DBVersion() assert.NoError(t, err) - fmt.Println(testEngine.Dialect().URI().DBType, "version is", version) + t.Log(testEngine.Dialect().URI().DBType, "version is", version) } func TestGetColumnsComment(t *testing.T) { diff --git a/tests/processors_test.go b/tests/processors_test.go index af2866e..f915dec 100644 --- a/tests/processors_test.go +++ b/tests/processors_test.go @@ -32,7 +32,7 @@ func TestBefore_Get(t *testing.T) { assert.EqualValues(t, 1, cnt) var be BeforeTable - has, err := testEngine.Before(func(bean interface{}) { + has, err := testEngine.Before(func(bean any) { bean.(*BeforeTable).Val = "val" }).Get(&be) assert.NoError(t, err) @@ -60,7 +60,7 @@ func TestBefore_Find(t *testing.T) { assert.EqualValues(t, 2, cnt) var be []BeforeTable2 - err = testEngine.Before(func(bean interface{}) { + err = testEngine.Before(func(bean any) { bean.(*BeforeTable2).Val = "val" }).Find(&be) assert.NoError(t, err) @@ -133,7 +133,7 @@ func TestProcessors(t *testing.T) { err = testEngine.CreateTables(&ProcessorsStruct{}) assert.NoError(t, err) - b4InsertFunc := func(bean interface{}) { + b4InsertFunc := func(bean any) { if v, ok := (bean).(*ProcessorsStruct); ok { v.B4InsertViaExt = 1 } else { @@ -141,7 +141,7 @@ func TestProcessors(t *testing.T) { } } - afterInsertFunc := func(bean interface{}) { + afterInsertFunc := func(bean any) { if v, ok := (bean).(*ProcessorsStruct); ok { v.AfterInsertedViaExt = 1 } else { @@ -234,7 +234,7 @@ func TestProcessors(t *testing.T) { // -- // test update processors - b4UpdateFunc := func(bean interface{}) { + b4UpdateFunc := func(bean any) { if v, ok := (bean).(*ProcessorsStruct); ok { v.B4UpdateViaExt = 1 } else { @@ -242,7 +242,7 @@ func TestProcessors(t *testing.T) { } } - afterUpdateFunc := func(bean interface{}) { + afterUpdateFunc := func(bean any) { if v, ok := (bean).(*ProcessorsStruct); ok { v.AfterUpdatedViaExt = 1 } else { @@ -274,7 +274,7 @@ func TestProcessors(t *testing.T) { // -- // test delete processors - b4DeleteFunc := func(bean interface{}) { + b4DeleteFunc := func(bean any) { if v, ok := (bean).(*ProcessorsStruct); ok { v.B4DeleteViaExt = 1 } else { @@ -282,7 +282,7 @@ func TestProcessors(t *testing.T) { } } - afterDeleteFunc := func(bean interface{}) { + afterDeleteFunc := func(bean any) { if v, ok := (bean).(*ProcessorsStruct); ok { v.AfterDeletedViaExt = 1 } else { @@ -374,7 +374,7 @@ func TestProcessorsTx(t *testing.T) { assert.NoError(t, err) p := &ProcessorsStruct{} - b4InsertFunc := func(bean interface{}) { + b4InsertFunc := func(bean any) { if v, ok := (bean).(*ProcessorsStruct); ok { v.B4InsertViaExt = 1 } else { @@ -382,7 +382,7 @@ func TestProcessorsTx(t *testing.T) { } } - afterInsertFunc := func(bean interface{}) { + afterInsertFunc := func(bean any) { if v, ok := (bean).(*ProcessorsStruct); ok { v.AfterInsertedViaExt = 1 } else { @@ -497,7 +497,7 @@ func TestProcessorsTx(t *testing.T) { err = session.Begin() assert.NoError(t, err) - b4UpdateFunc := func(bean interface{}) { + b4UpdateFunc := func(bean any) { if v, ok := (bean).(*ProcessorsStruct); ok { v.B4UpdateViaExt = 1 } else { @@ -505,7 +505,7 @@ func TestProcessorsTx(t *testing.T) { } } - afterUpdateFunc := func(bean interface{}) { + afterUpdateFunc := func(bean any) { if v, ok := (bean).(*ProcessorsStruct); ok { v.AfterUpdatedViaExt = 1 } else { @@ -671,7 +671,7 @@ func TestProcessorsTx(t *testing.T) { err = session.Begin() assert.NoError(t, err) - b4DeleteFunc := func(bean interface{}) { + b4DeleteFunc := func(bean any) { if v, ok := (bean).(*ProcessorsStruct); ok { v.B4DeleteViaExt = 1 } else { @@ -679,7 +679,7 @@ func TestProcessorsTx(t *testing.T) { } } - afterDeleteFunc := func(bean interface{}) { + afterDeleteFunc := func(bean any) { if v, ok := (bean).(*ProcessorsStruct); ok { v.AfterDeletedViaExt = 1 } else { diff --git a/tests/schema_test.go b/tests/schema_test.go index f248717..4613bd6 100644 --- a/tests/schema_test.go +++ b/tests/schema_test.go @@ -333,7 +333,7 @@ func TestIsTableEmpty(t *testing.T) { Created time.Time `xorm:"created"` ILike int PageView int - From_url string // nolint + From_url string Pre_url string `xorm:"unique"` // pre view image's url Uid int64 } @@ -726,7 +726,7 @@ func TestSyncWithOptions(t *testing.T) { assert.ElementsMatch(t, getKeysFromMap(tableInfoFromStruct.Indexes), getKeysFromMap(getIndicesOfBeanFromDB(t, &SyncWithOpts1{}))) } -func getIndicesOfBeanFromDB(t *testing.T, bean interface{}) map[string]*schemas.Index { +func getIndicesOfBeanFromDB(t *testing.T, bean any) map[string]*schemas.Index { dbm, err := testEngine.DBMetas() assert.NoError(t, err) @@ -779,13 +779,13 @@ func TestSync2_3(t *testing.T) { tables, err := testEngine.DBMetas() assert.NoError(t, err) tableInfo, err := testEngine.TableInfo(new(SyncTestUser2)) + assert.NoError(t, err) assert.EqualValues(t, tables[0].GetColumn("id").IsAutoIncrement, tableInfo.GetColumn("id").IsAutoIncrement) assert.EqualValues(t, tables[0].GetColumn("id").Name, tableInfo.GetColumn("id").Name) assert.EqualValues(t, tables[0].GetColumn("id").SQLType.Name, tableInfo.GetColumn("id").SQLType.Name) assert.EqualValues(t, tables[0].GetColumn("id").Nullable, tableInfo.GetColumn("id").Nullable) assert.EqualValues(t, tables[0].GetColumn("id").Comment, tableInfo.GetColumn("id").Comment) - } } diff --git a/tests/session_cond_test.go b/tests/session_cond_test.go index cbcd0cb..ce7cd03 100644 --- a/tests/session_cond_test.go +++ b/tests/session_cond_test.go @@ -155,7 +155,7 @@ func TestIn(t *testing.T) { } users = make([]Userinfo, 0) - var idsInterface []interface{} + var idsInterface []any for _, id := range ids { idsInterface = append(idsInterface, id) } diff --git a/tests/session_find_test.go b/tests/session_find_test.go index c800775..439d303 100644 --- a/tests/session_find_test.go +++ b/tests/session_find_test.go @@ -147,7 +147,7 @@ func TestFind3(t *testing.T) { err = testEngine.Cols("`team`.`id`"). Where("`team_user`.`org_id`=?", 1). And("`team_user`.`uid`=?", 2). - Join("INNER", []interface{}{teamUser}, "`team_user`.`team_id`=`team`.`id`"). + Join("INNER", []any{teamUser}, "`team_user`.`team_id`=`team`.`id`"). Find(&teams) assert.NoError(t, err) @@ -163,7 +163,7 @@ func TestFind3(t *testing.T) { err = testEngine.Cols("`team`.`id`"). Where("`tu`.`org_id`=?", 1). And("`tu`.`uid`=?", 2). - Join("INNER", []interface{}{"team_user", "tu"}, "`tu`.`team_id`=`team`.`id`"). + Join("INNER", []any{"team_user", "tu"}, "`tu`.`team_id`=`team`.`id`"). Find(&teams) assert.NoError(t, err) @@ -171,7 +171,7 @@ func TestFind3(t *testing.T) { err = testEngine.Cols("`team`.`id`"). Where("`tu`.`org_id`=?", 1). And("`tu`.`uid`=?", 2). - Join("INNER", []interface{}{teamUser, "tu"}, "`tu`.`team_id`=`team`.`id`"). + Join("INNER", []any{teamUser, "tu"}, "`tu`.`team_id`=`team`.`id`"). Find(&teams) assert.NoError(t, err) } @@ -368,7 +368,7 @@ func TestFindInterface(t *testing.T) { userinfo := testEngine.GetTableMapper().Obj2Table("Userinfo") username := testEngine.GetColumnMapper().Obj2Table("Username") - var idsInterface []interface{} + var idsInterface []any err := testEngine.Table(userinfo).Cols(username).Desc("id").Find(&idsInterface) assert.NoError(t, err) } diff --git a/tests/session_get_test.go b/tests/session_get_test.go index 15293cd..1a55028 100644 --- a/tests/session_get_test.go +++ b/tests/session_get_test.go @@ -186,9 +186,9 @@ func TestGetVar(t *testing.T) { assert.Equal(t, "28", valuesString["age"]) assert.Equal(t, "1.5", valuesString["money"]) - // for mymysql driver, interface{} will be []byte, so ignore it currently + // for mymysql driver, any will be []byte, so ignore it currently if testEngine.DriverName() != "mymysql" { - valuesInter := make(map[string]interface{}) + valuesInter := make(map[string]any) has, err = testEngine.Table("get_var").Where("`id` = ?", 1).Select("*").Get(&valuesInter) assert.NoError(t, err) assert.Equal(t, true, has) @@ -208,7 +208,7 @@ func TestGetVar(t *testing.T) { assert.Equal(t, "28", valuesSliceString[2]) assert.Equal(t, "1.5", valuesSliceString[3]) - valuesSliceInter := make([]interface{}, 5) + valuesSliceInter := make([]any, 5) has, err = testEngine.Table("get_var").Get(&valuesSliceInter) assert.NoError(t, err) assert.Equal(t, true, has) @@ -729,7 +729,7 @@ func TestGetViaMapCond(t *testing.T) { r GetViaMapCond platformStr = colMapper.Obj2Table("Platform") indexStr = colMapper.Obj2Table("Index") - query = map[string]interface{}{ + query = map[string]any{ platformStr: 1, indexStr: 1, } diff --git a/tests/session_insert_test.go b/tests/session_insert_test.go index cf8a3ce..2b6e3f1 100644 --- a/tests/session_insert_test.go +++ b/tests/session_insert_test.go @@ -5,6 +5,7 @@ package tests import ( + "errors" "fmt" "reflect" "testing" @@ -47,22 +48,22 @@ func TestInsertMulti(t *testing.T) { assert.EqualValues(t, 3, num) } -func insertMultiDatas(step int, datas interface{}) (num int64, err error) { +func insertMultiDatas(step int, datas any) (num int64, err error) { sliceValue := reflect.Indirect(reflect.ValueOf(datas)) var iLen int64 if sliceValue.Kind() != reflect.Slice { - return 0, fmt.Errorf("not silce") + return 0, errors.New("not silce") } iLen = int64(sliceValue.Len()) if iLen == 0 { - return + return 0, nil } session := testEngine.NewSession() defer session.Close() if err = callbackLooper(datas, step, - func(innerDatas interface{}) error { + func(innerDatas any) error { n, e := session.InsertMulti(innerDatas) if e != nil { return e @@ -74,19 +75,19 @@ func insertMultiDatas(step int, datas interface{}) (num int64, err error) { } else if num != iLen { return 0, fmt.Errorf("num error: %d - %d", num, iLen) } - return + return num, nil } -func callbackLooper(datas interface{}, step int, actionFunc func(interface{}) error) (err error) { +func callbackLooper(datas any, step int, actionFunc func(any) error) (err error) { sliceValue := reflect.Indirect(reflect.ValueOf(datas)) if sliceValue.Kind() != reflect.Slice { - return fmt.Errorf("not slice") + return errors.New("not slice") } if sliceValue.Len() <= 0 { - return + return nil } - tempLen := 0 + var tempLen int processedLen := sliceValue.Len() for i := 0; i < sliceValue.Len(); i += step { if processedLen > step { @@ -94,16 +95,16 @@ func callbackLooper(datas interface{}, step int, actionFunc func(interface{}) er } else { tempLen = sliceValue.Len() } - var tempInterface []interface{} + var tempInterface []any for j := i; j < tempLen; j++ { tempInterface = append(tempInterface, sliceValue.Index(j).Interface()) } if err = actionFunc(tempInterface); err != nil { - return + return err } processedLen -= step } - return + return nil } func TestInsertOneIfPkIsPoint(t *testing.T) { @@ -482,7 +483,7 @@ func TestInsertMulti2Interface(t *testing.T) { assertSync(t, new(Userinfo)) - users := []interface{}{ + users := []any{ Userinfo{Username: "xlw", Departname: "dev", Alias: "lunny2", Created: time.Now()}, Userinfo{Username: "xlw2", Departname: "dev", Alias: "lunny3", Created: time.Now()}, Userinfo{Username: "xlw11", Departname: "dev", Alias: "lunny2", Created: time.Now()}, @@ -496,7 +497,7 @@ func TestInsertMulti2Interface(t *testing.T) { } assert.EqualValues(t, len(users), cnt) - users2 := []interface{}{ + users2 := []any{ &Userinfo{Username: "1xlw", Departname: "dev", Alias: "lunny2", Created: time.Now()}, &Userinfo{Username: "1xlw2", Departname: "dev", Alias: "lunny3", Created: time.Now()}, &Userinfo{Username: "1xlw11", Departname: "dev", Alias: "lunny2", Created: time.Now()}, @@ -683,7 +684,7 @@ func TestInsertMap(t *testing.T) { assert.NoError(t, PrepareEngine()) assertSync(t, new(InsertMap)) - cnt, err := testEngine.Table(new(InsertMap)).Insert(map[string]interface{}{ + cnt, err := testEngine.Table(new(InsertMap)).Insert(map[string]any{ "width": 20, "height": 10, "name": "lunny", @@ -699,7 +700,7 @@ func TestInsertMap(t *testing.T) { assert.EqualValues(t, 10, im.Height) assert.EqualValues(t, "lunny", im.Name) - cnt, err = testEngine.Table("insert_map").Insert(map[string]interface{}{ + cnt, err = testEngine.Table("insert_map").Insert(map[string]any{ "width": 30, "height": 10, "name": "lunny", @@ -718,7 +719,7 @@ func TestInsertMap(t *testing.T) { assert.EqualValues(t, 10, ims[1].Height) assert.EqualValues(t, "lunny", ims[1].Name) - cnt, err = testEngine.Table("insert_map").Insert([]map[string]interface{}{ + cnt, err = testEngine.Table("insert_map").Insert([]map[string]any{ { "width": 40, "height": 10, @@ -797,7 +798,7 @@ func TestInsertWhere(t *testing.T) { inserted, err = testEngine.Table(new(InsertWhere)).Where("`repo_id`=?", 1). SetExpr("`index`", "coalesce(MAX(`index`),0)+1"). - Insert(map[string]interface{}{ + Insert(map[string]any{ "repo_id": 1, "width": 20, "height": 40, @@ -834,7 +835,7 @@ func TestInsertWhere(t *testing.T) { inserted, err = testEngine.Table(new(InsertWhere)).Where("`repo_id`=?", 1). SetExpr("`index`", "coalesce(MAX(`index`),0)+1"). - Insert(map[string]interface{}{ + Insert(map[string]any{ "repo_id": 1, "name": "10';delete * from insert_where; --", }) @@ -850,7 +851,7 @@ func TestInsertWhere(t *testing.T) { inserted, err = testEngine.Table(new(InsertWhere)).Where("`repo_id`=?", 1). SetExpr("`index`", "coalesce(MAX(`index`),0)+1"). - Insert(map[string]interface{}{ + Insert(map[string]any{ "repo_id": 1, "name": "10\\';delete * from insert_where; --", }) @@ -910,7 +911,7 @@ func TestInsertExpr2(t *testing.T) { SetExpr("is_draft", true). SetExpr("num_commits", 0). SetExpr("sha1", ""). - Insert(map[string]interface{}{ + Insert(map[string]any{ "repo_id": 1, "is_tag": true, }) @@ -947,7 +948,7 @@ func TestMultipleInsertTableName(t *testing.T) { err := trans.Begin() assert.NoError(t, err) - rtArr := []interface{}{ + rtArr := []any{ []*NightlyRate{ {ID: 1}, {ID: 2}, @@ -978,7 +979,7 @@ func TestInsertMultiWithOmit(t *testing.T) { assert.NoError(t, testEngine.Sync(new(TestMultiOmit))) - l := []interface{}{ + l := []any{ TestMultiOmit{Id: 1, Name: "1", Omitted: "1"}, TestMultiOmit{Id: 2, Name: "1", Omitted: "2"}, TestMultiOmit{Id: 3, Name: "1", Omitted: "3"}, @@ -1142,7 +1143,7 @@ func TestInsertMultipleMap(t *testing.T) { assert.NoError(t, PrepareEngine()) assertSync(t, new(InsertMultipleMap)) - cnt, err := testEngine.Table(new(InsertMultipleMap)).Insert([]map[string]interface{}{ + cnt, err := testEngine.Table(new(InsertMultipleMap)).Insert([]map[string]any{ { "width": 20, "height": 10, @@ -1298,7 +1299,6 @@ func TestInsertNotDeletedTimeStamp(t *testing.T) { if testEngine.Dialect().URI().DBType == schemas.MSSQL || testEngine.Dialect().URI().DBType == schemas.SQLITE || testEngine.Dialect().URI().DBType == schemas.POSTGRES { - type TestInsertNotDeletedTimeStampStructNotRight struct { ID uint64 `xorm:"'ID' pk autoincr"` DeletedAt time.Time `xorm:"'DELETED_AT' deleted notnull TIMESTAMP"` diff --git a/tests/session_iterate_test.go b/tests/session_iterate_test.go index f42f724..c6fd284 100644 --- a/tests/session_iterate_test.go +++ b/tests/session_iterate_test.go @@ -33,7 +33,7 @@ func TestIterate(t *testing.T) { assert.EqualValues(t, 1, cnt) cnt = 0 - err = testEngine.Iterate(new(UserIterate), func(i int, bean interface{}) error { + err = testEngine.Iterate(new(UserIterate), func(i int, bean any) error { user := bean.(*UserIterate) if cnt == 0 { assert.EqualValues(t, 1, user.Id) @@ -69,7 +69,7 @@ func TestBufferIterate(t *testing.T) { } cnt := 0 - err := testEngine.BufferSize(9).Iterate(new(UserBufferIterate), func(i int, bean interface{}) error { + err := testEngine.BufferSize(9).Iterate(new(UserBufferIterate), func(i int, bean any) error { user := bean.(*UserBufferIterate) assert.EqualValues(t, cnt+1, user.Id) assert.EqualValues(t, true, user.IsMan) @@ -80,7 +80,7 @@ func TestBufferIterate(t *testing.T) { assert.EqualValues(t, size, cnt) cnt = 0 - err = testEngine.Limit(20).BufferSize(9).Iterate(new(UserBufferIterate), func(i int, bean interface{}) error { + err = testEngine.Limit(20).BufferSize(9).Iterate(new(UserBufferIterate), func(i int, bean any) error { user := bean.(*UserBufferIterate) assert.EqualValues(t, cnt+1, user.Id) assert.EqualValues(t, true, user.IsMan) @@ -91,7 +91,7 @@ func TestBufferIterate(t *testing.T) { assert.EqualValues(t, size, cnt) cnt = 0 - err = testEngine.Limit(7).BufferSize(9).Iterate(new(UserBufferIterate), func(i int, bean interface{}) error { + err = testEngine.Limit(7).BufferSize(9).Iterate(new(UserBufferIterate), func(i int, bean any) error { user := bean.(*UserBufferIterate) assert.EqualValues(t, cnt+1, user.Id) assert.EqualValues(t, true, user.IsMan) @@ -102,7 +102,7 @@ func TestBufferIterate(t *testing.T) { assert.EqualValues(t, 7, cnt) cnt = 0 - err = testEngine.Where("`id` <= 10").BufferSize(2).Iterate(new(UserBufferIterate), func(i int, bean interface{}) error { + err = testEngine.Where("`id` <= 10").BufferSize(2).Iterate(new(UserBufferIterate), func(i int, bean any) error { user := bean.(*UserBufferIterate) assert.EqualValues(t, cnt+1, user.Id) assert.EqualValues(t, true, user.IsMan) @@ -117,7 +117,7 @@ func TestBufferIterate(t *testing.T) { sess := testEngine.NewSession() defer sess.Close() // generate: SELECT `id`, `is_man` FROM `user_buffer_iterate` WHERE (`id` <= 10) LIMIT 2 OFFSET 10 [] - err = sess.Where("`id` <= 10").BufferSize(2).Iterate(new(UserBufferIterate), func(i int, bean interface{}) error { + err = sess.Where("`id` <= 10").BufferSize(2).Iterate(new(UserBufferIterate), func(i int, bean any) error { user := bean.(*UserBufferIterate) assert.EqualValues(t, cnt+1, user.Id) assert.EqualValues(t, true, user.IsMan) diff --git a/tests/session_query_test.go b/tests/session_query_test.go index 726b19e..00aec8f 100644 --- a/tests/session_query_test.go +++ b/tests/session_query_test.go @@ -72,7 +72,7 @@ func TestQueryString2(t *testing.T) { assert.True(t, "0" == records[0]["msg"] || "false" == records[0]["msg"]) } -func toBool(i interface{}) bool { +func toBool(i any) bool { switch t := i.(type) { case int32: return t > 0 @@ -422,7 +422,7 @@ func TestQueryBLOBInMySQL(t *testing.T) { } { - arr := make([][]interface{}, 0) + arr := make([][]any, 0) err = testEngine.Table(testEngine.Quote(testEngine.TableName("avatar", true))).Cols("avatar").Find(&arr) assert.NoError(t, err) for i, record := range arr { @@ -433,7 +433,7 @@ func TestQueryBLOBInMySQL(t *testing.T) { } { - arr := make([]map[string]interface{}, 0) + arr := make([]map[string]any, 0) err = testEngine.Table(testEngine.Quote(testEngine.TableName("avatar", true))).Cols("avatar").Find(&arr) assert.NoError(t, err) for i, record := range arr { @@ -493,4 +493,4 @@ func TestRowsReset(t *testing.T) { assert.EqualValues(t, "4", rrs[0].Name) assert.EqualValues(t, "5", rrs[1].Name) assert.EqualValues(t, "6", rrs[2].Name) -} \ No newline at end of file +} diff --git a/tests/session_tx_test.go b/tests/session_tx_test.go index c9db40b..8553c89 100644 --- a/tests/session_tx_test.go +++ b/tests/session_tx_test.go @@ -5,7 +5,6 @@ package tests import ( - "fmt" "testing" "time" @@ -54,7 +53,7 @@ func TestCombineTransaction(t *testing.T) { counter := func() { total, err := testEngine.Count(&Userinfo{}) assert.NoError(t, err) - fmt.Printf("----now total %v records\n", total) + t.Logf("----now total %v records", total) } counter() @@ -96,7 +95,7 @@ func TestCombineTransactionSameMapper(t *testing.T) { counter := func() { total, err := testEngine.Count(&Userinfo{}) assert.NoError(t, err) - fmt.Printf("----now total %v records\n", total) + t.Logf("----now total %v records", total) } counter() @@ -201,7 +200,7 @@ func TestInsertMulti2InterfaceTransaction(t *testing.T) { err := session.Begin() assert.NoError(t, err) - users := []interface{}{ + users := []any{ &Multi2InterfaceTransaction{Name: "a", Alias: "A"}, &Multi2InterfaceTransaction{Name: "b", Alias: "B"}, &Multi2InterfaceTransaction{Name: "c", Alias: "C"}, diff --git a/tests/session_update_test.go b/tests/session_update_test.go index 350f813..0a5792c 100644 --- a/tests/session_update_test.go +++ b/tests/session_update_test.go @@ -35,14 +35,14 @@ func TestUpdateMap(t *testing.T) { _, err := testEngine.Insert(&tb) assert.NoError(t, err) - cnt, err := testEngine.Table("update_table").Where("`id` = ?", tb.Id).Update(map[string]interface{}{ + cnt, err := testEngine.Table("update_table").Where("`id` = ?", tb.Id).Update(map[string]any{ "name": "test2", "age": 36, }) assert.NoError(t, err) assert.EqualValues(t, 1, cnt) - cnt, err = testEngine.Table("update_table").ID(tb.Id).Update(map[string]interface{}{ + cnt, err = testEngine.Table("update_table").ID(tb.Id).Update(map[string]any{ "name": "test2", "age": 36, }) @@ -50,7 +50,7 @@ func TestUpdateMap(t *testing.T) { assert.True(t, statements.IsIDConditionWithNoTableErr(err)) assert.EqualValues(t, 0, cnt) - cnt, err = testEngine.Table("update_table").Update(map[string]interface{}{ + cnt, err = testEngine.Table("update_table").Update(map[string]any{ "name": "test2", "age": 36, }, &UpdateTable{ @@ -269,7 +269,7 @@ func TestWithIn(t *testing.T) { assert.EqualValues(t, 3, cnt) } -type Condi map[string]interface{} +type Condi map[string]any type UpdateAllCols struct { Id int64 @@ -304,7 +304,7 @@ func TestUpdateMap2(t *testing.T) { assert.NoError(t, PrepareEngine()) assertSync(t, new(UpdateMustCols)) - _, err := testEngine.Table("update_must_cols").Where("`id` =?", 1).Update(map[string]interface{}{ + _, err := testEngine.Table("update_must_cols").Where("`id` =?", 1).Update(map[string]any{ "bool": true, }) assert.NoError(t, err) @@ -751,7 +751,7 @@ func TestUseBool(t *testing.T) { cnt2, err := testEngine.UseBool().Update(&Userinfo{IsMan: true}) assert.NoError(t, err) if fNumber != cnt2 { - fmt.Println("cnt1", cnt1, "fNumber", fNumber, "cnt2", cnt2) + t.Log("cnt1", cnt1, "fNumber", fNumber, "cnt2", cnt2) /*err = errors.New("Updated number is not corrected.") t.Error(err) panic(err)*/ @@ -809,17 +809,17 @@ func TestNewUpdate(t *testing.T) { assert.NoError(t, PrepareEngine()) type TbUserInfo struct { - Id int64 `xorm:"pk autoincr unique BIGINT" json:"id"` - Phone string `xorm:"not null unique VARCHAR(20)" json:"phone"` - UserName string `xorm:"VARCHAR(20)" json:"user_name"` - Gender int `xorm:"default 0 INTEGER" json:"gender"` - Pw string `xorm:"VARCHAR(100)" json:"pw"` - Token string `xorm:"TEXT" json:"token"` - Avatar string `xorm:"TEXT" json:"avatar"` - Extras interface{} `xorm:"JSON" json:"extras"` - Created time.Time `xorm:"DATETIME created"` - Updated time.Time `xorm:"DATETIME updated"` - Deleted time.Time `xorm:"DATETIME deleted"` + Id int64 `xorm:"pk autoincr unique BIGINT" json:"id"` + Phone string `xorm:"not null unique VARCHAR(20)" json:"phone"` + UserName string `xorm:"VARCHAR(20)" json:"user_name"` + Gender int `xorm:"default 0 INTEGER" json:"gender"` + Pw string `xorm:"VARCHAR(100)" json:"pw"` + Token string `xorm:"TEXT" json:"token"` + Avatar string `xorm:"TEXT" json:"avatar"` + Extras any `xorm:"JSON" json:"extras"` + Created time.Time `xorm:"DATETIME created"` + Updated time.Time `xorm:"DATETIME updated"` + Deleted time.Time `xorm:"DATETIME deleted"` } assertSync(t, new(TbUserInfo)) @@ -964,7 +964,7 @@ func TestUpdateMapCondition(t *testing.T) { cnt, err := testEngine.Update(&UpdateMapCondition{ String: "string1", - }, map[string]interface{}{ + }, map[string]any{ "id": c.Id, }) assert.NoError(t, err) @@ -1000,7 +1000,7 @@ func TestUpdateMapContent(t *testing.T) { assert.NoError(t, err) assert.EqualValues(t, 18, c.Age) - cnt, err := testEngine.Table(new(UpdateMapContent)).ID(c.Id).Update(map[string]interface{}{"age": 0}) + cnt, err := testEngine.Table(new(UpdateMapContent)).ID(c.Id).Update(map[string]any{"age": 0}) assert.NoError(t, err) assert.EqualValues(t, 1, cnt) @@ -1010,7 +1010,7 @@ func TestUpdateMapContent(t *testing.T) { assert.True(t, has) assert.EqualValues(t, 0, c1.Age) - cnt, err = testEngine.Table(new(UpdateMapContent)).ID(c.Id).Update(map[string]interface{}{ + cnt, err = testEngine.Table(new(UpdateMapContent)).ID(c.Id).Update(map[string]any{ "age": 16, "is_man": false, "gender": 2, @@ -1026,7 +1026,7 @@ func TestUpdateMapContent(t *testing.T) { assert.EqualValues(t, false, c2.IsMan) assert.EqualValues(t, 2, c2.Gender) - cnt, err = testEngine.Table(new(UpdateMapContent)).ID(c.Id).Update(map[string]interface{}{ + cnt, err = testEngine.Table(new(UpdateMapContent)).ID(c.Id).Update(map[string]any{ "age": 15, "is_man": true, "gender": 1, @@ -1143,7 +1143,7 @@ func TestUpdateDeleted(t *testing.T) { assert.NoError(t, err) assert.EqualValues(t, 0, cnt) - cnt, err = testEngine.Table(&UpdateDeletedStruct{}).ID(s.Id).Update(map[string]interface{}{ + cnt, err = testEngine.Table(&UpdateDeletedStruct{}).ID(s.Id).Update(map[string]any{ "name": "test1", }) assert.NoError(t, err) @@ -1287,13 +1287,13 @@ func TestUpdateMap3(t *testing.T) { assertSync(t, new(UpdateMapUser)) - _, err := testEngine.Table(new(UpdateMapUser)).Insert(map[string]interface{}{ + _, err := testEngine.Table(new(UpdateMapUser)).Insert(map[string]any{ "Fname": "first user name", "Fver": 1, }) assert.NoError(t, err) - update := map[string]interface{}{ + update := map[string]any{ "Fname": "user name", "Fver": 1, } @@ -1301,7 +1301,7 @@ func TestUpdateMap3(t *testing.T) { assert.NoError(t, err) assert.EqualValues(t, 1, rows) - update = map[string]interface{}{ + update = map[string]any{ "Name": "user name", "Ver": 1, } @@ -1375,10 +1375,10 @@ func TestUpdateMultiplePK(t *testing.T) { assert.EqualValues(t, 1, num) test.Value = "4" - _, err = testEngine.ID([]interface{}{test.Id, test.Name}).Update(test) + _, err = testEngine.ID([]any{test.Id, test.Name}).Update(test) assert.NoError(t, err) - type MySlice []interface{} + type MySlice []any test.Value = "5" _, err = testEngine.ID(&MySlice{test.Id, test.Name}).Update(test) assert.NoError(t, err) diff --git a/tests/tests.go b/tests/tests.go index dfc04d0..d9c55a0 100644 --- a/tests/tests.go +++ b/tests/tests.go @@ -58,7 +58,7 @@ func createEngine(dbType, connStr string) error { } createDBSQL := "If(db_id(N'xorm_test') IS NULL) BEGIN CREATE DATABASE xorm_test" if collation != nil && *collation != "" { - createDBSQL += fmt.Sprintf(" COLLATE %s", *collation) + createDBSQL += " COLLATE " + *collation } createDBSQL += "; END;" if _, err = db.Exec(createDBSQL); err != nil { @@ -147,11 +147,12 @@ func createEngine(dbType, connStr string) error { } } - if *quotePolicyStr == "none" { + switch *quotePolicyStr { + case "none": testEngine.SetQuotePolicy(dialects.QuotePolicyNone) - } else if *quotePolicyStr == "reserved" { + case "reserved": testEngine.SetQuotePolicy(dialects.QuotePolicyReserved) - } else { + default: testEngine.SetQuotePolicy(dialects.QuotePolicyAlways) } @@ -168,7 +169,7 @@ func createEngine(dbType, connStr string) error { if err != nil { return err } - tableNames := make([]interface{}, 0, len(tables)) + tableNames := make([]any, 0, len(tables)) for _, table := range tables { tableNames = append(tableNames, table.Name) } @@ -185,21 +186,22 @@ func MainTest(m *testing.M) { flag.Parse() dbType = *db - if *db == "sqlite3" { + switch *db { + case "sqlite3": if ptrConnStr == nil { connString = "./test_sqlite3.db?cache=shared&mode=rwc" } else { connString = *ptrConnStr } - } else if *db == "sqlite" { + case "sqlite": if ptrConnStr == nil { connString = "./test_sqlite.db?cache=shared&mode=rwc" } else { connString = *ptrConnStr } - } else { + default: if ptrConnStr == nil { - fmt.Println("you should indicate conn string") + fmt.Fprintln(os.Stderr, "you should indicate conn string") return } connString = *ptrConnStr @@ -213,10 +215,10 @@ func MainTest(m *testing.M) { dbType = dbs[i] connString = conns[i] testEngine = nil - fmt.Println("testing", dbType, connString) + fmt.Fprintln(os.Stderr, "testing", dbType, connString) if err := PrepareEngine(); err != nil { - fmt.Println(err) + fmt.Fprintln(os.Stderr, err) os.Exit(1) return } diff --git a/tests/time_test.go b/tests/time_test.go index 13b9ed1..a769ea2 100644 --- a/tests/time_test.go +++ b/tests/time_test.go @@ -5,7 +5,6 @@ package tests import ( - "fmt" "strconv" "strings" "testing" @@ -41,7 +40,7 @@ func TestTimeUserTime(t *testing.T) { OperTime: time.Now(), } - fmt.Println("user", user.OperTime) + t.Log("user", user.OperTime) cnt, err := testEngine.Insert(&user) assert.NoError(t, err) @@ -53,7 +52,7 @@ func TestTimeUserTime(t *testing.T) { assert.True(t, has) assert.EqualValues(t, user.OperTime.Unix(), user2.OperTime.Unix()) assert.EqualValues(t, formatTime(user.OperTime), formatTime(user2.OperTime)) - fmt.Println("user2", user2.OperTime) + t.Log("user2", user2.OperTime) } func TestTimeUserTimeDiffLoc(t *testing.T) { @@ -86,7 +85,7 @@ func TestTimeUserTimeDiffLoc(t *testing.T) { OperTime: time.Now(), } - fmt.Println("user", user.OperTime) + t.Log("user", user.OperTime) cnt, err := testEngine.Insert(&user) assert.NoError(t, err) @@ -98,7 +97,7 @@ func TestTimeUserTimeDiffLoc(t *testing.T) { assert.True(t, has) assert.EqualValues(t, user.OperTime.Unix(), user2.OperTime.Unix()) assert.EqualValues(t, formatTime(user.OperTime.In(loc)), formatTime(user2.OperTime)) - fmt.Println("user2", user2.OperTime) + t.Log("user2", user2.OperTime) } func TestTimeUserCreated(t *testing.T) { @@ -115,7 +114,7 @@ func TestTimeUserCreated(t *testing.T) { Id: "lunny", } - fmt.Println("user", user.CreatedAt) + t.Log("user", user.CreatedAt) cnt, err := testEngine.Insert(&user) assert.NoError(t, err) @@ -127,7 +126,7 @@ func TestTimeUserCreated(t *testing.T) { assert.True(t, has) assert.EqualValues(t, user.CreatedAt.Unix(), user2.CreatedAt.Unix()) assert.EqualValues(t, formatTime(user.CreatedAt), formatTime(user2.CreatedAt)) - fmt.Println("user2", user2.CreatedAt) + t.Log("user2", user2.CreatedAt) } func TestTimeUserCreatedDiffLoc(t *testing.T) { @@ -159,7 +158,7 @@ func TestTimeUserCreatedDiffLoc(t *testing.T) { Id: "lunny", } - fmt.Println("user", user.CreatedAt) + t.Log("user", user.CreatedAt) cnt, err := testEngine.Insert(&user) assert.NoError(t, err) @@ -171,7 +170,7 @@ func TestTimeUserCreatedDiffLoc(t *testing.T) { assert.True(t, has) assert.EqualValues(t, user.CreatedAt.Unix(), user2.CreatedAt.Unix()) assert.EqualValues(t, formatTime(user.CreatedAt), formatTime(user2.CreatedAt)) - fmt.Println("user2", user2.CreatedAt) + t.Log("user2", user2.CreatedAt) } func TestTimeUserUpdated(t *testing.T) { @@ -189,7 +188,7 @@ func TestTimeUserUpdated(t *testing.T) { Id: "lunny", } - fmt.Println("user", user.CreatedAt, user.UpdatedAt) + t.Log("user", user.CreatedAt, user.UpdatedAt) cnt, err := testEngine.Insert(&user) assert.NoError(t, err) @@ -203,7 +202,7 @@ func TestTimeUserUpdated(t *testing.T) { assert.EqualValues(t, formatTime(user.CreatedAt), formatTime(user2.CreatedAt)) assert.EqualValues(t, user.UpdatedAt.Unix(), user2.UpdatedAt.Unix()) assert.EqualValues(t, formatTime(user.UpdatedAt), formatTime(user2.UpdatedAt)) - fmt.Println("user2", user2.CreatedAt, user2.UpdatedAt) + t.Log("user2", user2.CreatedAt, user2.UpdatedAt) user3 := UserUpdated{ Id: "lunny2", @@ -222,7 +221,7 @@ func TestTimeUserUpdated(t *testing.T) { assert.EqualValues(t, formatTime(user.CreatedAt), formatTime(user4.CreatedAt)) assert.EqualValues(t, user3.UpdatedAt.Unix(), user4.UpdatedAt.Unix()) assert.EqualValues(t, formatTime(user3.UpdatedAt), formatTime(user4.UpdatedAt)) - fmt.Println("user3", user.CreatedAt, user4.UpdatedAt) + t.Log("user3", user.CreatedAt, user4.UpdatedAt) } func TestTimeUserUpdatedDiffLoc(t *testing.T) { @@ -255,7 +254,7 @@ func TestTimeUserUpdatedDiffLoc(t *testing.T) { Id: "lunny", } - fmt.Println("user", user.CreatedAt, user.UpdatedAt) + t.Log("user", user.CreatedAt, user.UpdatedAt) cnt, err := testEngine.Insert(&user) assert.NoError(t, err) @@ -269,7 +268,7 @@ func TestTimeUserUpdatedDiffLoc(t *testing.T) { assert.EqualValues(t, formatTime(user.CreatedAt), formatTime(user2.CreatedAt)) assert.EqualValues(t, user.UpdatedAt.Unix(), user2.UpdatedAt.Unix()) assert.EqualValues(t, formatTime(user.UpdatedAt), formatTime(user2.UpdatedAt)) - fmt.Println("user2", user2.CreatedAt, user2.UpdatedAt) + t.Log("user2", user2.CreatedAt, user2.UpdatedAt) user3 := UserUpdated2{ Id: "lunny2", @@ -288,7 +287,7 @@ func TestTimeUserUpdatedDiffLoc(t *testing.T) { assert.EqualValues(t, formatTime(user.CreatedAt), formatTime(user4.CreatedAt)) assert.EqualValues(t, user3.UpdatedAt.Unix(), user4.UpdatedAt.Unix()) assert.EqualValues(t, formatTime(user3.UpdatedAt), formatTime(user4.UpdatedAt)) - fmt.Println("user3", user.CreatedAt, user4.UpdatedAt) + t.Log("user3", user.CreatedAt, user4.UpdatedAt) } func TestTimeUserDeleted(t *testing.T) { @@ -312,7 +311,7 @@ func TestTimeUserDeleted(t *testing.T) { cnt, err := testEngine.Insert(&user) assert.NoError(t, err) assert.EqualValues(t, 1, cnt) - fmt.Println("user", user.CreatedAt, user.UpdatedAt, user.DeletedAt) + t.Log("user", user.CreatedAt, user.UpdatedAt, user.DeletedAt) var user2 UserDeleted has, err := testEngine.Get(&user2) @@ -323,8 +322,8 @@ func TestTimeUserDeleted(t *testing.T) { assert.EqualValues(t, user.UpdatedAt.Unix(), user2.UpdatedAt.Unix()) assert.EqualValues(t, formatTime(user.UpdatedAt), formatTime(user2.UpdatedAt)) assert.True(t, utils.IsTimeZero(user2.DeletedAt)) - fmt.Println("user2", user2.CreatedAt, user2.UpdatedAt, user2.DeletedAt) - fmt.Println("user2 str", user2.CreatedAtStr, user2.UpdatedAtStr) + t.Log("user2", user2.CreatedAt, user2.UpdatedAt, user2.DeletedAt) + t.Log("user2 str", user2.CreatedAtStr, user2.UpdatedAtStr) var user3 UserDeleted cnt, err = testEngine.Where("`id` = ?", "lunny").Delete(&user3) @@ -338,7 +337,7 @@ func TestTimeUserDeleted(t *testing.T) { assert.True(t, has) assert.EqualValues(t, user3.DeletedAt.Unix(), user4.DeletedAt.Unix()) assert.EqualValues(t, formatTime(user3.DeletedAt), formatTime(user4.DeletedAt)) - fmt.Println("user3", user3.DeletedAt, user4.DeletedAt) + t.Log("user3", user3.DeletedAt, user4.DeletedAt) } func TestTimeUserDeletedDiffLoc(t *testing.T) { @@ -375,7 +374,7 @@ func TestTimeUserDeletedDiffLoc(t *testing.T) { cnt, err := testEngine.Insert(&user) assert.NoError(t, err) assert.EqualValues(t, 1, cnt) - fmt.Println("user", user.CreatedAt, user.UpdatedAt, user.DeletedAt) + t.Log("user", user.CreatedAt, user.UpdatedAt, user.DeletedAt) var user2 UserDeleted2 has, err := testEngine.Get(&user2) @@ -386,7 +385,7 @@ func TestTimeUserDeletedDiffLoc(t *testing.T) { assert.EqualValues(t, user.UpdatedAt.Unix(), user2.UpdatedAt.Unix()) assert.EqualValues(t, formatTime(user.UpdatedAt), formatTime(user2.UpdatedAt)) assert.True(t, utils.IsTimeZero(user2.DeletedAt)) - fmt.Println("user2", user2.CreatedAt, user2.UpdatedAt, user2.DeletedAt) + t.Log("user2", user2.CreatedAt, user2.UpdatedAt, user2.DeletedAt) var user3 UserDeleted2 cnt, err = testEngine.Where("`id` = ?", "lunny").Delete(&user3) @@ -400,7 +399,7 @@ func TestTimeUserDeletedDiffLoc(t *testing.T) { assert.True(t, has) assert.EqualValues(t, user3.DeletedAt.Unix(), user4.DeletedAt.Unix()) assert.EqualValues(t, formatTime(user3.DeletedAt), formatTime(user4.DeletedAt)) - fmt.Println("user3", user3.DeletedAt, user4.DeletedAt) + t.Log("user3", user3.DeletedAt, user4.DeletedAt) } type JSONDate time.Time @@ -446,7 +445,7 @@ func TestCustomTimeUserDeleted(t *testing.T) { cnt, err := testEngine.Insert(&user) assert.NoError(t, err) assert.EqualValues(t, 1, cnt) - fmt.Println("user", user.CreatedAt, user.UpdatedAt, user.DeletedAt) + t.Log("user", user.CreatedAt, user.UpdatedAt, user.DeletedAt) var user2 UserDeleted3 has, err := testEngine.Get(&user2) @@ -457,7 +456,7 @@ func TestCustomTimeUserDeleted(t *testing.T) { assert.EqualValues(t, user.UpdatedAt.Unix(), user2.UpdatedAt.Unix()) assert.EqualValues(t, formatTime(time.Time(user.UpdatedAt)), formatTime(time.Time(user2.UpdatedAt))) assert.True(t, utils.IsTimeZero(time.Time(user2.DeletedAt))) - fmt.Println("user2", user2.CreatedAt, user2.UpdatedAt, user2.DeletedAt) + t.Log("user2", user2.CreatedAt, user2.UpdatedAt, user2.DeletedAt) var user3 UserDeleted3 cnt, err = testEngine.Where("`id` = ?", "lunny").Delete(&user3) @@ -471,7 +470,7 @@ func TestCustomTimeUserDeleted(t *testing.T) { assert.True(t, has) assert.EqualValues(t, user3.DeletedAt.Unix(), user4.DeletedAt.Unix()) assert.EqualValues(t, formatTime(time.Time(user3.DeletedAt)), formatTime(time.Time(user4.DeletedAt))) - fmt.Println("user3", user3.DeletedAt, user4.DeletedAt) + t.Log("user3", user3.DeletedAt, user4.DeletedAt) } func TestCustomTimeUserDeletedDiffLoc(t *testing.T) { @@ -508,7 +507,7 @@ func TestCustomTimeUserDeletedDiffLoc(t *testing.T) { cnt, err := testEngine.Insert(&user) assert.NoError(t, err) assert.EqualValues(t, 1, cnt) - fmt.Println("user", user.CreatedAt, user.UpdatedAt, user.DeletedAt) + t.Log("user", user.CreatedAt, user.UpdatedAt, user.DeletedAt) var user2 UserDeleted4 has, err := testEngine.Get(&user2) @@ -519,7 +518,7 @@ func TestCustomTimeUserDeletedDiffLoc(t *testing.T) { assert.EqualValues(t, user.UpdatedAt.Unix(), user2.UpdatedAt.Unix()) assert.EqualValues(t, formatTime(time.Time(user.UpdatedAt)), formatTime(time.Time(user2.UpdatedAt))) assert.True(t, utils.IsTimeZero(time.Time(user2.DeletedAt))) - fmt.Println("user2", user2.CreatedAt, user2.UpdatedAt, user2.DeletedAt) + t.Log("user2", user2.CreatedAt, user2.UpdatedAt, user2.DeletedAt) var user3 UserDeleted4 cnt, err = testEngine.Where("`id` = ?", "lunny").Delete(&user3) @@ -533,7 +532,7 @@ func TestCustomTimeUserDeletedDiffLoc(t *testing.T) { assert.True(t, has) assert.EqualValues(t, user3.DeletedAt.Unix(), user4.DeletedAt.Unix()) assert.EqualValues(t, formatTime(time.Time(user3.DeletedAt)), formatTime(time.Time(user4.DeletedAt))) - fmt.Println("user3", user3.DeletedAt, user4.DeletedAt) + t.Log("user3", user3.DeletedAt, user4.DeletedAt) } func TestDeletedInt64(t *testing.T) { diff --git a/tests/types_null_test.go b/tests/types_null_test.go index d4fa250..e285ae2 100644 --- a/tests/types_null_test.go +++ b/tests/types_null_test.go @@ -35,7 +35,7 @@ func (CustomStruct) String() string { return "CustomStruct" } -func (m *CustomStruct) Scan(value interface{}) error { +func (m *CustomStruct) Scan(value any) error { if value == nil { m.Year, m.Month, m.Day = 0, 0, 0 return nil @@ -98,7 +98,7 @@ func TestNullStructInsert(t *testing.T) { items := []NullStruct{} for i := 0; i < 5; i++ { item := NullStruct{ - Name: sql.NullString{String: "haolei_" + fmt.Sprint(i+1), Valid: true}, + Name: sql.NullString{String: "haolei_" + strconv.Itoa(i+1), Valid: true}, Age: sql.NullInt64{Int64: 30 + int64(i), Valid: true}, Height: sql.NullFloat64{Float64: 1.5 + 1.1*float64(i), Valid: true}, IsMan: sql.NullBool{Bool: true, Valid: true}, @@ -261,9 +261,9 @@ func TestNullStructIterate(t *testing.T) { if true { err := testEngine.Where("`age` IS NOT NULL").OrderBy("age").Iterate(new(NullStruct), - func(i int, bean interface{}) error { + func(i int, bean any) error { nultype := bean.(*NullStruct) - fmt.Println(i, nultype) + t.Log(i, nultype) return nil }) assert.NoError(t, err) diff --git a/tests/types_test.go b/tests/types_test.go index c6749ce..a8ba23b 100644 --- a/tests/types_test.go +++ b/tests/types_test.go @@ -124,7 +124,6 @@ type ConvConfig struct { func (s *ConvConfig) FromDB(data []byte) error { if data == nil { - s = nil return nil } return json.DefaultJSONHandler.Unmarshal(data, s) @@ -313,7 +312,7 @@ func TestCustomType1(t *testing.T) { return } - fmt.Println(i) + t.Log(i) i.NameArray = []string{} i.MSS = map[string]string{} i.F = 0 @@ -409,14 +408,14 @@ func TestCustomType2(t *testing.T) { assert.NoError(t, err) assert.True(t, exist) - fmt.Println(user) + t.Log(user) users := make([]UserCus, 0) err = testEngine.Where("`"+testEngine.GetColumnMapper().Obj2Table("Status")+"` = ?", "Registered").Find(&users) assert.NoError(t, err) assert.EqualValues(t, 1, len(users)) - fmt.Println(users) + t.Log(users) } func TestUnsignedUint64(t *testing.T) { @@ -523,7 +522,7 @@ func TestUnsignedTinyInt(t *testing.T) { case schemas.MSSQL: assert.EqualValues(t, "INT", tables[0].Columns()[0].SQLType.Name) default: - assert.False(t, true, fmt.Sprintf("Unsigned is not implemented, returned %s", tables[0].Columns()[0].SQLType.Name)) + assert.False(t, true, "Unsigned is not implemented, returned "+tables[0].Columns()[0].SQLType.Name) } cnt, err := testEngine.Insert(&MyUnsignedTinyIntStruct{ @@ -544,15 +543,14 @@ type MyDecimal big.Int func (d *MyDecimal) FromDB(data []byte) error { i, _ := strconv.ParseInt(string(data), 10, 64) if d == nil { - d = (*MyDecimal)(big.NewInt(i)) - } else { - (*big.Int)(d).SetInt64(i) + return nil } + (*big.Int)(d).SetInt64(i) return nil } func (d *MyDecimal) ToDB() ([]byte, error) { - return []byte(fmt.Sprintf("%d", (*big.Int)(d).Int64())), nil + return fmt.Appendf(nil, "%d", (*big.Int)(d).Int64()), nil } func (d *MyDecimal) AsBigInt() *big.Int { @@ -636,7 +634,7 @@ func (d ZDecimal) ToDB() ([]byte, error) { if d.value == nil { return []byte("0"), nil } - return []byte(fmt.Sprintf("%d", (d.value).Int64())), nil + return fmt.Appendf(nil, "%d", (d.value).Int64()), nil } func (d ZDecimal) IsZero() bool {