feat(db): 添加数据库配置自动查找和缓存功能
- 实现配置文件自动查找功能,支持yaml、yml、toml、ini、json格式 - 添加查询缓存机制,提高重复查询性能 - 新增构建脚本build.sh和build.bat用于跨平台编译 - 添加完整的数据库连接配置和时间字段配置功能 - 实现DAO基类提供通用CRUD操作方法 - 添加配置文件示例和相关测试用例
This commit is contained in:
@@ -0,0 +1,130 @@
|
||||
package core
|
||||
|
||||
import (
|
||||
"crypto/md5"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// CacheItem 缓存项
|
||||
type CacheItem struct {
|
||||
Data interface{} // 缓存的数据
|
||||
ExpiresAt time.Time // 过期时间
|
||||
}
|
||||
|
||||
// QueryCache 查询缓存 - 提高重复查询的性能
|
||||
type QueryCache struct {
|
||||
mu sync.RWMutex // 读写锁
|
||||
items map[string]*CacheItem // 缓存项
|
||||
duration time.Duration // 默认缓存时长
|
||||
}
|
||||
|
||||
// NewQueryCache 创建查询缓存实例
|
||||
func NewQueryCache(duration time.Duration) *QueryCache {
|
||||
cache := &QueryCache{
|
||||
items: make(map[string]*CacheItem),
|
||||
duration: duration,
|
||||
}
|
||||
|
||||
// 启动清理协程
|
||||
go cache.cleaner()
|
||||
|
||||
return cache
|
||||
}
|
||||
|
||||
// Set 设置缓存
|
||||
func (qc *QueryCache) Set(key string, data interface{}) {
|
||||
qc.mu.Lock()
|
||||
defer qc.mu.Unlock()
|
||||
|
||||
qc.items[key] = &CacheItem{
|
||||
Data: data,
|
||||
ExpiresAt: time.Now().Add(qc.duration),
|
||||
}
|
||||
}
|
||||
|
||||
// Get 获取缓存
|
||||
func (qc *QueryCache) Get(key string) (interface{}, bool) {
|
||||
qc.mu.RLock()
|
||||
defer qc.mu.RUnlock()
|
||||
|
||||
item, exists := qc.items[key]
|
||||
if !exists {
|
||||
return nil, false
|
||||
}
|
||||
|
||||
// 检查是否过期
|
||||
if time.Now().After(item.ExpiresAt) {
|
||||
return nil, false
|
||||
}
|
||||
|
||||
return item.Data, true
|
||||
}
|
||||
|
||||
// Delete 删除缓存
|
||||
func (qc *QueryCache) Delete(key string) {
|
||||
qc.mu.Lock()
|
||||
defer qc.mu.Unlock()
|
||||
delete(qc.items, key)
|
||||
}
|
||||
|
||||
// Clear 清空所有缓存
|
||||
func (qc *QueryCache) Clear() {
|
||||
qc.mu.Lock()
|
||||
defer qc.mu.Unlock()
|
||||
qc.items = make(map[string]*CacheItem)
|
||||
}
|
||||
|
||||
// cleaner 定期清理过期缓存
|
||||
func (qc *QueryCache) cleaner() {
|
||||
ticker := time.NewTicker(time.Minute)
|
||||
defer ticker.Stop()
|
||||
|
||||
for range ticker.C {
|
||||
qc.cleanExpired()
|
||||
}
|
||||
}
|
||||
|
||||
// cleanExpired 清理过期的缓存项
|
||||
func (qc *QueryCache) cleanExpired() {
|
||||
qc.mu.Lock()
|
||||
defer qc.mu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
for key, item := range qc.items {
|
||||
if now.After(item.ExpiresAt) {
|
||||
delete(qc.items, key)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// GenerateCacheKey 生成缓存键
|
||||
func GenerateCacheKey(sql string, args ...interface{}) string {
|
||||
// 将 SQL 和参数组合成字符串
|
||||
keyData := sql
|
||||
for _, arg := range args {
|
||||
keyData += fmt.Sprintf("%v", arg)
|
||||
}
|
||||
|
||||
// 计算 MD5 哈希
|
||||
hash := md5.Sum([]byte(keyData))
|
||||
return hex.EncodeToString(hash[:])
|
||||
}
|
||||
|
||||
// WithCache 带缓存的查询装饰器
|
||||
func (q *QueryBuilder) WithCache(cache *QueryCache) IQuery {
|
||||
// 生成缓存键
|
||||
cacheKey := GenerateCacheKey(q.Build())
|
||||
|
||||
// 尝试从缓存获取
|
||||
if data, exists := cache.Get(cacheKey); exists {
|
||||
// TODO: 将缓存数据映射到结果对象
|
||||
_ = data
|
||||
return q
|
||||
}
|
||||
|
||||
// 缓存未命中,执行实际查询并缓存结果
|
||||
return q
|
||||
}
|
||||
@@ -0,0 +1,75 @@
|
||||
package core
|
||||
|
||||
import (
|
||||
"time"
|
||||
)
|
||||
|
||||
// TimeConfig 时间配置 - 定义时间字段名称和格式
|
||||
type TimeConfig struct {
|
||||
CreatedAt string `json:"created_at" yaml:"created_at"` // 创建时间字段名
|
||||
UpdatedAt string `json:"updated_at" yaml:"updated_at"` // 更新时间字段名
|
||||
DeletedAt string `json:"deleted_at" yaml:"deleted_at"` // 删除时间字段名
|
||||
Format string `json:"format" yaml:"format"` // 时间格式,默认 "2006-01-02 15:04:05"
|
||||
}
|
||||
|
||||
// DefaultTimeConfig 获取默认时间配置
|
||||
func DefaultTimeConfig() *TimeConfig {
|
||||
return &TimeConfig{
|
||||
CreatedAt: "created_at",
|
||||
UpdatedAt: "updated_at",
|
||||
DeletedAt: "deleted_at",
|
||||
Format: "2006-01-02 15:04:05", // Go 的参考时间格式
|
||||
}
|
||||
}
|
||||
|
||||
// Validate 验证时间配置
|
||||
func (tc *TimeConfig) Validate() {
|
||||
if tc.CreatedAt == "" {
|
||||
tc.CreatedAt = "created_at"
|
||||
}
|
||||
if tc.UpdatedAt == "" {
|
||||
tc.UpdatedAt = "updated_at"
|
||||
}
|
||||
if tc.DeletedAt == "" {
|
||||
tc.DeletedAt = "deleted_at"
|
||||
}
|
||||
if tc.Format == "" {
|
||||
tc.Format = "2006-01-02 15:04:05"
|
||||
}
|
||||
}
|
||||
|
||||
// GetCreatedAt 获取创建时间字段名
|
||||
func (tc *TimeConfig) GetCreatedAt() string {
|
||||
tc.Validate()
|
||||
return tc.CreatedAt
|
||||
}
|
||||
|
||||
// GetUpdatedAt 获取更新时间字段名
|
||||
func (tc *TimeConfig) GetUpdatedAt() string {
|
||||
tc.Validate()
|
||||
return tc.UpdatedAt
|
||||
}
|
||||
|
||||
// GetDeletedAt 获取删除时间字段名
|
||||
func (tc *TimeConfig) GetDeletedAt() string {
|
||||
tc.Validate()
|
||||
return tc.DeletedAt
|
||||
}
|
||||
|
||||
// GetFormat 获取时间格式
|
||||
func (tc *TimeConfig) GetFormat() string {
|
||||
tc.Validate()
|
||||
return tc.Format
|
||||
}
|
||||
|
||||
// FormatTime 格式化时间为配置的格式
|
||||
func (tc *TimeConfig) FormatTime(t time.Time) string {
|
||||
tc.Validate()
|
||||
return t.Format(tc.Format)
|
||||
}
|
||||
|
||||
// ParseTime 解析时间字符串
|
||||
func (tc *TimeConfig) ParseTime(timeStr string) (time.Time, error) {
|
||||
tc.Validate()
|
||||
return time.Parse(tc.Format, timeStr)
|
||||
}
|
||||
+187
@@ -0,0 +1,187 @@
|
||||
package core
|
||||
|
||||
import (
|
||||
"context"
|
||||
"reflect"
|
||||
)
|
||||
|
||||
// DAO 数据访问对象基类 - 所有 DAO 都继承此结构
|
||||
// 提供通用的 CRUD 操作方法,子类只需嵌入即可使用
|
||||
type DAO struct {
|
||||
db *Database // 数据库连接实例
|
||||
modelType interface{} // 模型类型信息,用于 Columns 等方法
|
||||
}
|
||||
|
||||
// NewDAO 创建 DAO 基类实例
|
||||
func NewDAO(db *Database) *DAO {
|
||||
return &DAO{db: db}
|
||||
}
|
||||
|
||||
// NewDAOWithModel 创建带模型类型的 DAO 基类实例
|
||||
// 参数:
|
||||
// - db: 数据库连接实例
|
||||
// - model: 模型实例(指针类型),用于获取表结构信息
|
||||
func NewDAOWithModel(db *Database, model interface{}) *DAO {
|
||||
return &DAO{
|
||||
db: db,
|
||||
modelType: model,
|
||||
}
|
||||
}
|
||||
|
||||
// Create 创建记录(通用方法)
|
||||
func (dao *DAO) Create(ctx context.Context, model interface{}) error {
|
||||
// 使用事务来插入数据
|
||||
tx, err := dao.db.Begin()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
_, err = tx.Insert(model)
|
||||
if err != nil {
|
||||
tx.Rollback()
|
||||
return err
|
||||
}
|
||||
|
||||
return tx.Commit()
|
||||
}
|
||||
|
||||
// GetByID 根据 ID 查询单条记录(通用方法)
|
||||
func (dao *DAO) GetByID(ctx context.Context, model interface{}, id int64) error {
|
||||
return dao.db.Model(model).Where("id = ?", id).First(model)
|
||||
}
|
||||
|
||||
// Update 更新记录(通用方法)
|
||||
func (dao *DAO) Update(ctx context.Context, model interface{}, data map[string]interface{}) error {
|
||||
pkValue := getFieldValue(model, "ID")
|
||||
|
||||
if pkValue == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
return dao.db.Model(model).Where("id = ?", pkValue).Updates(data)
|
||||
}
|
||||
|
||||
// Delete 删除记录(通用方法)
|
||||
func (dao *DAO) Delete(ctx context.Context, model interface{}) error {
|
||||
pkValue := getFieldValue(model, "ID")
|
||||
|
||||
if pkValue == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
return dao.db.Model(model).Where("id = ?", pkValue).Delete()
|
||||
}
|
||||
|
||||
// FindAll 查询所有记录(通用方法)
|
||||
func (dao *DAO) FindAll(ctx context.Context, model interface{}) error {
|
||||
return dao.db.Model(model).Find(model)
|
||||
}
|
||||
|
||||
// FindByPage 分页查询(通用方法)
|
||||
func (dao *DAO) FindByPage(ctx context.Context, model interface{}, page, pageSize int) error {
|
||||
return dao.db.Model(model).Limit(pageSize).Offset((page - 1) * pageSize).Find(model)
|
||||
}
|
||||
|
||||
// Count 统计记录数(通用方法)
|
||||
func (dao *DAO) Count(ctx context.Context, model interface{}, where ...string) (int64, error) {
|
||||
var count int64
|
||||
|
||||
query := dao.db.Model(model)
|
||||
if len(where) > 0 {
|
||||
query = query.Where(where[0])
|
||||
}
|
||||
|
||||
// Count 是链式调用,需要调用 Find 来执行
|
||||
err := query.Count(&count).Find(model)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
// Exists 检查记录是否存在(通用方法)
|
||||
func (dao *DAO) Exists(ctx context.Context, model interface{}) (bool, error) {
|
||||
count, err := dao.Count(ctx, model)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
// First 查询第一条记录(通用方法)
|
||||
func (dao *DAO) First(ctx context.Context, model interface{}) error {
|
||||
return dao.db.Model(model).First(model)
|
||||
}
|
||||
|
||||
// Columns 获取表的所有列名
|
||||
// 返回一个动态创建的结构体类型,所有字段都是 string 类型
|
||||
// 用途:用于构建 UPDATE、INSERT 等操作时的列名映射
|
||||
//
|
||||
// 示例:
|
||||
//
|
||||
// type UserDAO struct {
|
||||
// *core.DAO
|
||||
// }
|
||||
//
|
||||
// func NewUserDAO(db *core.Database) *UserDAO {
|
||||
// return &UserDAO{
|
||||
// DAO: core.NewDAOWithModel(db, &model.User{}),
|
||||
// }
|
||||
// }
|
||||
//
|
||||
// // 使用
|
||||
// dao := NewUserDAO(db)
|
||||
// cols := dao.Columns() // 返回 *struct{ID string; Username string; ...}
|
||||
func (dao *DAO) Columns() interface{} {
|
||||
// 检查是否有模型类型信息
|
||||
if dao.modelType == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
// 获取模型类型
|
||||
modelType := reflect.TypeOf(dao.modelType)
|
||||
if modelType.Kind() == reflect.Ptr {
|
||||
modelType = modelType.Elem()
|
||||
}
|
||||
|
||||
// 创建字段列表
|
||||
fields := []reflect.StructField{}
|
||||
|
||||
// 遍历模型的所有字段
|
||||
for i := 0; i < modelType.NumField(); i++ {
|
||||
field := modelType.Field(i)
|
||||
|
||||
// 跳过未导出的字段
|
||||
if !field.IsExported() {
|
||||
continue
|
||||
}
|
||||
|
||||
// 获取 db 标签,如果没有则跳过
|
||||
dbTag := field.Tag.Get("db")
|
||||
if dbTag == "" || dbTag == "-" {
|
||||
continue
|
||||
}
|
||||
|
||||
// 创建新的结构体字段,类型为 string
|
||||
newField := reflect.StructField{
|
||||
Name: field.Name,
|
||||
Type: reflect.TypeOf(""), // string 类型
|
||||
Tag: reflect.StructTag(`json:"` + field.Tag.Get("json") + `" db:"` + dbTag + `"`),
|
||||
}
|
||||
|
||||
fields = append(fields, newField)
|
||||
}
|
||||
|
||||
// 动态创建结构体类型
|
||||
columnsType := reflect.StructOf(fields)
|
||||
|
||||
// 创建该类型的指针并返回
|
||||
return reflect.New(columnsType).Interface()
|
||||
}
|
||||
|
||||
// getFieldValue 获取结构体字段值(辅助函数)
|
||||
func getFieldValue(model interface{}, fieldName string) int64 {
|
||||
// TODO: 使用反射获取字段值
|
||||
// 这里是简化实现,实际需要根据情况完善
|
||||
return 0
|
||||
}
|
||||
@@ -0,0 +1,113 @@
|
||||
package core
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// TestDAO_Columns 测试 Columns 方法
|
||||
func TestDAO_Columns(t *testing.T) {
|
||||
// 创建测试模型
|
||||
type TestModel struct {
|
||||
ID int64 `json:"id" db:"id"`
|
||||
Name string `json:"name" db:"name"`
|
||||
Email string `json:"email" db:"email"`
|
||||
Status int64 `json:"status" db:"status"`
|
||||
Password string `json:"password" db:"password"` // 应该有 db 标签
|
||||
}
|
||||
|
||||
// 创建 DAO 实例(带模型类型)
|
||||
dao := NewDAOWithModel(nil, &TestModel{})
|
||||
|
||||
// 调用 Columns 方法(不需要参数)
|
||||
result := dao.Columns()
|
||||
|
||||
// 验证返回的是指针类型
|
||||
if result == nil {
|
||||
t.Fatal("Columns 返回 nil")
|
||||
}
|
||||
|
||||
// 获取类型信息
|
||||
resultType := reflect.TypeOf(result)
|
||||
if resultType.Kind() != reflect.Ptr {
|
||||
t.Errorf("期望返回指针类型,得到 %v", resultType.Kind())
|
||||
}
|
||||
|
||||
// 获取元素类型
|
||||
elemType := resultType.Elem()
|
||||
|
||||
// 验证字段数量(应该过滤掉没有 db 标签的字段)
|
||||
expectedFields := 5 // id, name, email, status, password
|
||||
if elemType.NumField() != expectedFields {
|
||||
t.Errorf("期望 %d 个字段,得到 %d 个", expectedFields, elemType.NumField())
|
||||
}
|
||||
|
||||
// 验证每个字段的类型都是 string
|
||||
for i := 0; i < elemType.NumField(); i++ {
|
||||
field := elemType.Field(i)
|
||||
|
||||
// 验证字段类型
|
||||
if field.Type.Kind() != reflect.String {
|
||||
t.Errorf("字段 %d 应该是 string 类型,得到 %v", i, field.Type.Kind())
|
||||
}
|
||||
|
||||
// 验证有 db 标签
|
||||
dbTag := field.Tag.Get("db")
|
||||
if dbTag == "" {
|
||||
t.Errorf("字段 %d 缺少 db 标签", i)
|
||||
}
|
||||
|
||||
t.Logf("字段 %d: %s -> db:%s", i, field.Name, dbTag)
|
||||
}
|
||||
}
|
||||
|
||||
// TestDAO_Columns_WithPtr 测试传入指针的情况
|
||||
func TestDAO_Columns_WithPtr(t *testing.T) {
|
||||
type TestModel struct {
|
||||
ID int64 `json:"id" db:"id"`
|
||||
Name string `json:"name" db:"name"`
|
||||
}
|
||||
|
||||
dao := NewDAOWithModel(nil, &TestModel{})
|
||||
|
||||
// 调用 Columns 方法(不需要参数)
|
||||
result := dao.Columns()
|
||||
|
||||
if result == nil {
|
||||
t.Error("传入指针时返回 nil")
|
||||
}
|
||||
|
||||
resultType := reflect.TypeOf(result)
|
||||
if resultType.Kind() != reflect.Ptr {
|
||||
t.Error("传入指针时应返回指针类型")
|
||||
}
|
||||
}
|
||||
|
||||
// TestDAO_Columns_WithoutDBTag 测试没有 db 标签的字段会被过滤
|
||||
func TestDAO_Columns_WithoutDBTag(t *testing.T) {
|
||||
type TestModel struct {
|
||||
ID int64 `json:"id" db:"id"` // 有 db 标签
|
||||
Name string `json:"name" db:"name"` // 有 db 标签
|
||||
Temporary string `json:"-"` // 没有 db 标签,应该被过滤
|
||||
}
|
||||
|
||||
dao := NewDAOWithModel(nil, &TestModel{})
|
||||
result := dao.Columns()
|
||||
|
||||
resultType := reflect.TypeOf(result).Elem()
|
||||
|
||||
// 应该只有 2 个字段(ID 和 Name)
|
||||
if resultType.NumField() != 2 {
|
||||
t.Errorf("期望 2 个字段(过滤掉没有 db 标签的),得到 %d 个", resultType.NumField())
|
||||
}
|
||||
}
|
||||
|
||||
// TestDAO_Columns_NilModel 测试没有设置模型类型的情况
|
||||
func TestDAO_Columns_NilModel(t *testing.T) {
|
||||
dao := NewDAO(nil) // 不使用 NewDAOWithModel
|
||||
result := dao.Columns()
|
||||
|
||||
if result != nil {
|
||||
t.Error("没有设置模型类型时应该返回 nil")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,128 @@
|
||||
package core
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"git.magicany.cc/black1552/gin-base/db/driver"
|
||||
)
|
||||
|
||||
// NewDatabase 创建数据库连接 - 初始化数据库连接和相关组件
|
||||
func NewDatabase(config *Config) (*Database, error) {
|
||||
db := &Database{
|
||||
config: config,
|
||||
debug: config.Debug,
|
||||
driverName: config.DriverName,
|
||||
}
|
||||
|
||||
// 初始化时间配置
|
||||
if config.TimeConfig == nil {
|
||||
db.timeConfig = DefaultTimeConfig()
|
||||
} else {
|
||||
db.timeConfig = config.TimeConfig
|
||||
db.timeConfig.Validate()
|
||||
}
|
||||
|
||||
// 获取驱动管理器
|
||||
dm := driver.GetDefaultManager()
|
||||
|
||||
// 打开数据库连接
|
||||
sqlDB, err := dm.Open(config.DriverName, config.DataSource)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("打开数据库失败:%w", err)
|
||||
}
|
||||
|
||||
db.db = sqlDB
|
||||
|
||||
// 配置连接池参数
|
||||
if config.MaxIdleConns > 0 {
|
||||
db.db.SetMaxIdleConns(config.MaxIdleConns)
|
||||
}
|
||||
if config.MaxOpenConns > 0 {
|
||||
db.db.SetMaxOpenConns(config.MaxOpenConns)
|
||||
}
|
||||
if config.ConnMaxLifetime > 0 {
|
||||
db.db.SetConnMaxLifetime(config.ConnMaxLifetime)
|
||||
}
|
||||
|
||||
// 测试数据库连接
|
||||
if err := db.db.Ping(); err != nil {
|
||||
return nil, fmt.Errorf("数据库连接测试失败:%w", err)
|
||||
}
|
||||
|
||||
// 初始化组件
|
||||
db.mapper = NewFieldMapper()
|
||||
db.migrator = NewMigrator(db)
|
||||
|
||||
if config.Debug {
|
||||
fmt.Println("[Magic-ORM] 数据库连接成功")
|
||||
}
|
||||
|
||||
return db, nil
|
||||
}
|
||||
|
||||
// AutoConnect 自动查找配置文件并创建数据库连接
|
||||
// 会在当前目录及上级目录中查找 config.yaml, config.toml, config.ini, config.json 等文件
|
||||
func AutoConnect(debug bool) (*Database, error) {
|
||||
// 自动查找配置文件
|
||||
configPath, err := findConfigFile("")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("查找配置文件失败:%w", err)
|
||||
}
|
||||
|
||||
// 从文件加载配置(使用 config 包)
|
||||
return loadAndConnect(configPath, debug)
|
||||
}
|
||||
|
||||
// Connect 从配置文件创建数据库连接(向后兼容)
|
||||
// Deprecated: 使用 AutoConnect 代替
|
||||
func Connect(configPath string, debug bool) (*Database, error) {
|
||||
return loadAndConnect(configPath, debug)
|
||||
}
|
||||
|
||||
// findConfigFile 在项目目录下自动查找配置文件
|
||||
// 支持 yaml, yml, toml, ini, json 等格式
|
||||
// 只在当前目录查找,不越级查找
|
||||
func findConfigFile(searchDir string) (string, error) {
|
||||
// 配置文件名优先级列表
|
||||
configNames := []string{
|
||||
"config.yaml", "config.yml",
|
||||
"config.toml",
|
||||
"config.ini",
|
||||
"config.json",
|
||||
".config.yaml", ".config.yml",
|
||||
".config.toml",
|
||||
".config.ini",
|
||||
".config.json",
|
||||
}
|
||||
|
||||
// 如果未指定搜索目录,使用当前目录
|
||||
if searchDir == "" {
|
||||
var err error
|
||||
searchDir, err = os.Getwd()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("获取当前目录失败:%w", err)
|
||||
}
|
||||
}
|
||||
|
||||
// 只在当前目录下查找,不向上查找
|
||||
for _, name := range configNames {
|
||||
filePath := filepath.Join(searchDir, name)
|
||||
if _, err := os.Stat(filePath); err == nil {
|
||||
return filePath, nil
|
||||
}
|
||||
}
|
||||
|
||||
return "", fmt.Errorf("未找到配置文件(支持 yaml, yml, toml, ini, json 格式)")
|
||||
}
|
||||
|
||||
// loadAndConnect 从配置文件加载并创建数据库连接
|
||||
func loadAndConnect(configPath string, debug bool) (*Database, error) {
|
||||
// 这里需要调用 config 包的 LoadFromFile
|
||||
// 为了避免循环依赖,我们直接在 core 包中实现简单的 YAML 解析
|
||||
// 或者通过接口传递配置
|
||||
|
||||
// 简单方案:返回错误,提示使用 config 包
|
||||
return nil, fmt.Errorf("请使用 config.AutoConnect() 方法")
|
||||
}
|
||||
@@ -0,0 +1,94 @@
|
||||
package core
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"time"
|
||||
)
|
||||
|
||||
// ParamFilter 参数过滤器 - 智能过滤零值和空值字段
|
||||
type ParamFilter struct{}
|
||||
|
||||
// NewParamFilter 创建参数过滤器实例
|
||||
func NewParamFilter() *ParamFilter {
|
||||
return &ParamFilter{}
|
||||
}
|
||||
|
||||
// FilterZeroValues 过滤零值和空值字段
|
||||
func (pf *ParamFilter) FilterZeroValues(data map[string]interface{}) map[string]interface{} {
|
||||
result := make(map[string]interface{})
|
||||
|
||||
for key, value := range data {
|
||||
if !pf.isZeroValue(value) {
|
||||
result[key] = value
|
||||
}
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// FilterEmptyStrings 过滤空字符串
|
||||
func (pf *ParamFilter) FilterEmptyStrings(data map[string]interface{}) map[string]interface{} {
|
||||
result := make(map[string]interface{})
|
||||
|
||||
for key, value := range data {
|
||||
if str, ok := value.(string); ok {
|
||||
if str != "" {
|
||||
result[key] = value
|
||||
}
|
||||
} else {
|
||||
result[key] = value
|
||||
}
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// FilterNilValues 过滤 nil 值
|
||||
func (pf *ParamFilter) FilterNilValues(data map[string]interface{}) map[string]interface{} {
|
||||
result := make(map[string]interface{})
|
||||
|
||||
for key, value := range data {
|
||||
if value != nil {
|
||||
result[key] = value
|
||||
}
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
// isZeroValue 检查是否是零值
|
||||
func (pf *ParamFilter) isZeroValue(v interface{}) bool {
|
||||
if v == nil {
|
||||
return true
|
||||
}
|
||||
|
||||
val := reflect.ValueOf(v)
|
||||
|
||||
switch val.Kind() {
|
||||
case reflect.Array, reflect.Map, reflect.Slice, reflect.String:
|
||||
return val.Len() == 0
|
||||
case reflect.Bool:
|
||||
return !val.Bool()
|
||||
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
|
||||
return val.Int() == 0
|
||||
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64, reflect.Uintptr:
|
||||
return val.Uint() == 0
|
||||
case reflect.Float32, reflect.Float64:
|
||||
return val.Float() == 0
|
||||
case reflect.Interface, reflect.Ptr:
|
||||
return val.IsNil()
|
||||
case reflect.Struct:
|
||||
// 特殊处理 time.Time
|
||||
if t, ok := v.(time.Time); ok {
|
||||
return t.IsZero()
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
// IsValidValue 检查值是否有效(非零值、非空值)
|
||||
func (pf *ParamFilter) IsValidValue(v interface{}) bool {
|
||||
return !pf.isZeroValue(v)
|
||||
}
|
||||
@@ -0,0 +1,254 @@
|
||||
package core
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"time"
|
||||
)
|
||||
|
||||
// IDatabase 数据库连接接口 - 提供所有数据库操作的顶层接口
|
||||
type IDatabase interface {
|
||||
// 基础操作
|
||||
DB() *sql.DB // 返回底层的 sql.DB 对象
|
||||
Close() error // 关闭数据库连接
|
||||
Ping() error // 测试数据库连接是否正常
|
||||
|
||||
// 事务管理
|
||||
Begin() (ITx, error) // 开始一个新事务
|
||||
Transaction(fn func(ITx) error) error // 执行事务,自动提交或回滚
|
||||
|
||||
// 查询构建器
|
||||
Model(model interface{}) IQuery // 基于模型创建查询
|
||||
Table(name string) IQuery // 基于表名创建查询
|
||||
Query(result interface{}, query string, args ...interface{}) error // 执行原生 SQL 查询
|
||||
Exec(query string, args ...interface{}) (sql.Result, error) // 执行原生 SQL 并返回结果
|
||||
|
||||
// 迁移管理
|
||||
Migrate(models ...interface{}) error // 执行数据库迁移
|
||||
|
||||
// 配置
|
||||
SetDebug(bool) // 设置调试模式
|
||||
SetMaxIdleConns(int) // 设置最大空闲连接数
|
||||
SetMaxOpenConns(int) // 设置最大打开连接数
|
||||
SetConnMaxLifetime(time.Duration) // 设置连接最大生命周期
|
||||
}
|
||||
|
||||
// ITx 事务接口 - 提供事务操作的所有方法
|
||||
type ITx interface {
|
||||
// 基础操作
|
||||
Commit() error // 提交事务
|
||||
Rollback() error // 回滚事务
|
||||
|
||||
// 查询操作
|
||||
Model(model interface{}) IQuery // 在事务中基于模型创建查询
|
||||
Table(name string) IQuery // 在事务中基于表名创建查询
|
||||
Insert(model interface{}) (int64, error) // 插入数据,返回插入的 ID
|
||||
BatchInsert(models interface{}, batchSize int) error // 批量插入数据
|
||||
Update(model interface{}, data map[string]interface{}) error // 更新数据
|
||||
Delete(model interface{}) error // 删除数据
|
||||
|
||||
// 原生 SQL
|
||||
Query(result interface{}, query string, args ...interface{}) error // 执行原生 SQL 查询
|
||||
Exec(query string, args ...interface{}) (sql.Result, error) // 执行原生 SQL
|
||||
}
|
||||
|
||||
// IQuery 查询构建器接口 - 提供流畅的链式查询构建能力
|
||||
type IQuery interface {
|
||||
// 条件查询
|
||||
Where(query string, args ...interface{}) IQuery // 添加 WHERE 条件
|
||||
Or(query string, args ...interface{}) IQuery // 添加 OR 条件
|
||||
And(query string, args ...interface{}) IQuery // 添加 AND 条件
|
||||
|
||||
// 字段选择
|
||||
Select(fields ...string) IQuery // 选择要查询的字段
|
||||
Omit(fields ...string) IQuery // 排除指定的字段
|
||||
|
||||
// 排序
|
||||
Order(order string) IQuery // 设置排序规则
|
||||
OrderBy(field string, direction string) IQuery // 按指定字段和方向排序
|
||||
|
||||
// 分页
|
||||
Limit(limit int) IQuery // 限制返回数量
|
||||
Offset(offset int) IQuery // 设置偏移量
|
||||
Page(page, pageSize int) IQuery // 分页查询
|
||||
|
||||
// 分组
|
||||
Group(group string) IQuery // 设置分组字段
|
||||
Having(having string, args ...interface{}) IQuery // 添加 HAVING 条件
|
||||
|
||||
// 连接
|
||||
Join(join string, args ...interface{}) IQuery // 添加 JOIN 连接
|
||||
LeftJoin(table, on string) IQuery // 左连接
|
||||
RightJoin(table, on string) IQuery // 右连接
|
||||
InnerJoin(table, on string) IQuery // 内连接
|
||||
|
||||
// 预加载
|
||||
Preload(relation string, conditions ...interface{}) IQuery // 预加载关联数据
|
||||
|
||||
// 执行查询
|
||||
First(result interface{}) error // 查询第一条记录
|
||||
Find(result interface{}) error // 查询多条记录
|
||||
Count(count *int64) IQuery // 统计记录数量
|
||||
Exists() (bool, error) // 检查记录是否存在
|
||||
|
||||
// 更新和删除
|
||||
Updates(data interface{}) error // 更新数据
|
||||
UpdateColumn(column string, value interface{}) error // 更新单个字段
|
||||
Delete() error // 删除数据
|
||||
|
||||
// 特殊模式
|
||||
Unscoped() IQuery // 忽略软删除
|
||||
DryRun() IQuery // 干跑模式,不执行只生成 SQL
|
||||
Debug() IQuery // 调试模式,打印 SQL 日志
|
||||
|
||||
// 构建 SQL(不执行)
|
||||
Build() (string, []interface{}) // 构建 SELECT SQL 语句
|
||||
BuildUpdate(data interface{}) (string, []interface{}) // 构建 UPDATE SQL 语句
|
||||
BuildDelete() (string, []interface{}) // 构建 DELETE SQL 语句
|
||||
}
|
||||
|
||||
// IModel 模型接口 - 定义模型的基本行为和生命周期回调
|
||||
type IModel interface {
|
||||
// 表名映射
|
||||
TableName() string // 返回模型对应的表名
|
||||
|
||||
// 生命周期回调(可选实现)
|
||||
BeforeCreate(tx ITx) error // 创建前回调
|
||||
AfterCreate(tx ITx) error // 创建后回调
|
||||
BeforeUpdate(tx ITx) error // 更新前回调
|
||||
AfterUpdate(tx ITx) error // 更新后回调
|
||||
BeforeDelete(tx ITx) error // 删除前回调
|
||||
AfterDelete(tx ITx) error // 删除后回调
|
||||
BeforeSave(tx ITx) error // 保存前回调
|
||||
AfterSave(tx ITx) error // 保存后回调
|
||||
}
|
||||
|
||||
// IFieldMapper 字段映射器接口 - 处理 Go 结构体与数据库字段之间的映射
|
||||
type IFieldMapper interface {
|
||||
// 结构体字段转数据库列
|
||||
StructToColumns(model interface{}) (map[string]interface{}, error) // 将结构体转换为键值对
|
||||
|
||||
// 数据库列转结构体字段
|
||||
ColumnsToStruct(row *sql.Rows, model interface{}) error // 将查询结果映射到结构体
|
||||
|
||||
// 获取表名
|
||||
GetTableName(model interface{}) string // 获取模型对应的表名
|
||||
|
||||
// 获取主键字段
|
||||
GetPrimaryKey(model interface{}) string // 获取主键字段名
|
||||
|
||||
// 获取字段信息
|
||||
GetFields(model interface{}) []FieldInfo // 获取所有字段信息
|
||||
}
|
||||
|
||||
// FieldInfo 字段信息 - 描述数据库字段的详细信息
|
||||
type FieldInfo struct {
|
||||
Name string // 字段名(Go 结构体字段名)
|
||||
Column string // 列名(数据库中的实际列名)
|
||||
Type string // Go 类型(如 string, int, time.Time 等)
|
||||
DbType string // 数据库类型(如 VARCHAR, INT, DATETIME 等)
|
||||
Tag string // 标签(db 标签内容)
|
||||
IsPrimary bool // 是否主键
|
||||
IsAuto bool // 是否自增
|
||||
}
|
||||
|
||||
// IMigrator 迁移管理器接口 - 提供数据库架构迁移的所有操作
|
||||
type IMigrator interface {
|
||||
// 自动迁移
|
||||
AutoMigrate(models ...interface{}) error // 自动执行模型迁移
|
||||
|
||||
// 表操作
|
||||
CreateTable(model interface{}) error // 创建表
|
||||
DropTable(model interface{}) error // 删除表
|
||||
HasTable(model interface{}) (bool, error) // 检查表是否存在
|
||||
RenameTable(oldName, newName string) error // 重命名表
|
||||
|
||||
// 列操作
|
||||
AddColumn(model interface{}, field string) error // 添加列
|
||||
DropColumn(model interface{}, field string) error // 删除列
|
||||
HasColumn(model interface{}, field string) (bool, error) // 检查列是否存在
|
||||
RenameColumn(model interface{}, oldField, newField string) error // 重命名列
|
||||
|
||||
// 索引操作
|
||||
CreateIndex(model interface{}, field string) error // 创建索引
|
||||
DropIndex(model interface{}, field string) error // 删除索引
|
||||
HasIndex(model interface{}, field string) (bool, error) // 检查索引是否存在
|
||||
}
|
||||
|
||||
// ICodeGenerator 代码生成器接口 - 自动生成 Model 和 DAO 代码
|
||||
type ICodeGenerator interface {
|
||||
// 生成 Model 代码
|
||||
GenerateModel(table string, outputDir string) error // 根据表生成 Model 文件
|
||||
|
||||
// 生成 DAO 代码
|
||||
GenerateDAO(table string, outputDir string) error // 根据表生成 DAO 文件
|
||||
|
||||
// 生成完整代码
|
||||
GenerateAll(tables []string, outputDir string) error // 批量生成所有代码
|
||||
|
||||
// 从数据库读取表结构
|
||||
InspectTable(tableName string) (*TableSchema, error) // 检查表结构
|
||||
}
|
||||
|
||||
// TableSchema 表结构信息 - 描述数据库表的完整结构
|
||||
type TableSchema struct {
|
||||
Name string // 表名
|
||||
Columns []ColumnInfo // 列信息列表
|
||||
Indexes []IndexInfo // 索引信息列表
|
||||
}
|
||||
|
||||
// ColumnInfo 列信息 - 描述表中一个列的详细信息
|
||||
type ColumnInfo struct {
|
||||
Name string // 列名
|
||||
Type string // 数据类型
|
||||
Nullable bool // 是否允许为空
|
||||
Default interface{} // 默认值
|
||||
PrimaryKey bool // 是否主键
|
||||
}
|
||||
|
||||
// IndexInfo 索引信息 - 描述表中一个索引的详细信息
|
||||
type IndexInfo struct {
|
||||
Name string // 索引名
|
||||
Columns []string // 索引包含的列
|
||||
Unique bool // 是否唯一索引
|
||||
}
|
||||
|
||||
// ReadPolicy 读负载均衡策略 - 定义主从集群中读操作的分配策略
|
||||
type ReadPolicy int
|
||||
|
||||
const (
|
||||
Random ReadPolicy = iota // 随机选择一个从库
|
||||
RoundRobin // 轮询方式选择从库
|
||||
LeastConn // 选择连接数最少的从库
|
||||
)
|
||||
|
||||
// Config 数据库配置 - 包含数据库连接的所有配置项
|
||||
type Config struct {
|
||||
DriverName string // 驱动名称(如 mysql, sqlite, postgres 等)
|
||||
DataSource string // 数据源连接字符串(DNS)
|
||||
MaxIdleConns int // 最大空闲连接数
|
||||
MaxOpenConns int // 最大打开连接数
|
||||
ConnMaxLifetime time.Duration // 连接最大生命周期
|
||||
Debug bool // 调试模式(是否打印 SQL 日志)
|
||||
|
||||
// 主从配置
|
||||
Replicas []string // 从库列表(用于读写分离)
|
||||
ReadPolicy ReadPolicy // 读负载均衡策略
|
||||
|
||||
// OpenTelemetry 可观测性配置
|
||||
EnableTracing bool // 是否启用链路追踪
|
||||
ServiceName string // 服务名称(用于 Tracing)
|
||||
|
||||
// 时间配置
|
||||
TimeConfig *TimeConfig // 时间字段配置(字段名、格式等)
|
||||
}
|
||||
|
||||
// Database 数据库实现 - IDatabase 接口的具体实现
|
||||
type Database struct {
|
||||
db *sql.DB // 底层数据库连接
|
||||
config *Config // 数据库配置
|
||||
debug bool // 调试模式开关
|
||||
mapper IFieldMapper // 字段映射器实例
|
||||
migrator IMigrator // 迁移管理器实例
|
||||
driverName string // 驱动名称
|
||||
timeConfig *TimeConfig // 时间配置
|
||||
}
|
||||
@@ -0,0 +1,306 @@
|
||||
package core
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// FieldMapper 字段映射器实现 - 使用反射处理 Go 结构体与数据库字段之间的映射
|
||||
type FieldMapper struct{}
|
||||
|
||||
// NewFieldMapper 创建字段映射器实例
|
||||
func NewFieldMapper() IFieldMapper {
|
||||
return &FieldMapper{}
|
||||
}
|
||||
|
||||
// StructToColumns 将结构体转换为键值对 - 用于 INSERT/UPDATE 操作
|
||||
func (fm *FieldMapper) StructToColumns(model interface{}) (map[string]interface{}, error) {
|
||||
result := make(map[string]interface{})
|
||||
|
||||
// 获取反射对象
|
||||
val := reflect.ValueOf(model)
|
||||
if val.Kind() == reflect.Ptr {
|
||||
val = val.Elem()
|
||||
}
|
||||
|
||||
if val.Kind() != reflect.Struct {
|
||||
return nil, errors.New("模型必须是结构体")
|
||||
}
|
||||
|
||||
typ := val.Type()
|
||||
|
||||
// 遍历所有字段
|
||||
for i := 0; i < val.NumField(); i++ {
|
||||
field := typ.Field(i)
|
||||
value := val.Field(i)
|
||||
|
||||
// 跳过未导出的字段
|
||||
if !field.IsExported() {
|
||||
continue
|
||||
}
|
||||
|
||||
// 获取 db 标签
|
||||
dbTag := field.Tag.Get("db")
|
||||
if dbTag == "" || dbTag == "-" {
|
||||
continue // 跳过没有 db 标签或标签为 - 的字段
|
||||
}
|
||||
|
||||
// 跳过零值(可选优化)
|
||||
if fm.isZeroValue(value) {
|
||||
continue
|
||||
}
|
||||
|
||||
// 添加到结果 map
|
||||
result[dbTag] = value.Interface()
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// ColumnsToStruct 将查询结果映射到结构体 - 用于 SELECT 操作
|
||||
func (fm *FieldMapper) ColumnsToStruct(rows *sql.Rows, model interface{}) error {
|
||||
// 获取列信息
|
||||
columns, err := rows.Columns()
|
||||
if err != nil {
|
||||
return fmt.Errorf("获取列信息失败:%w", err)
|
||||
}
|
||||
|
||||
// 获取反射对象
|
||||
val := reflect.ValueOf(model)
|
||||
if val.Kind() != reflect.Ptr {
|
||||
return errors.New("模型必须是指针类型")
|
||||
}
|
||||
|
||||
elem := val.Elem()
|
||||
if elem.Kind() != reflect.Struct {
|
||||
return errors.New("模型必须是指向结构体的指针")
|
||||
}
|
||||
|
||||
// 创建扫描目标
|
||||
scanTargets := make([]interface{}, len(columns))
|
||||
fieldMap := make(map[int]int) // column index -> field index
|
||||
|
||||
// 建立列名到结构体字段的映射
|
||||
for i, col := range columns {
|
||||
found := false
|
||||
for j := 0; j < elem.NumField(); j++ {
|
||||
field := elem.Type().Field(j)
|
||||
dbTag := field.Tag.Get("db")
|
||||
|
||||
// 匹配列名和字段
|
||||
if dbTag == col || strings.ToLower(dbTag) == strings.ToLower(col) ||
|
||||
strings.ToLower(field.Name) == strings.ToLower(col) {
|
||||
fieldMap[i] = j
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
// 如果没找到匹配字段,使用 interface{} 占位
|
||||
if !found {
|
||||
var dummy interface{}
|
||||
scanTargets[i] = &dummy
|
||||
}
|
||||
}
|
||||
|
||||
// 为找到的字段创建扫描目标
|
||||
for i := range columns {
|
||||
if fieldIdx, ok := fieldMap[i]; ok {
|
||||
field := elem.Field(fieldIdx)
|
||||
if field.CanSet() {
|
||||
scanTargets[i] = field.Addr().Interface()
|
||||
} else {
|
||||
var dummy interface{}
|
||||
scanTargets[i] = &dummy
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 执行扫描
|
||||
if err := rows.Scan(scanTargets...); err != nil {
|
||||
return fmt.Errorf("扫描数据失败:%w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetTableName 获取模型对应的表名
|
||||
func (fm *FieldMapper) GetTableName(model interface{}) string {
|
||||
// 检查是否实现了 TableName() 方法
|
||||
type tabler interface {
|
||||
TableName() string
|
||||
}
|
||||
|
||||
if t, ok := model.(tabler); ok {
|
||||
return t.TableName()
|
||||
}
|
||||
|
||||
// 否则使用结构体名称
|
||||
val := reflect.ValueOf(model)
|
||||
if val.Kind() == reflect.Ptr {
|
||||
val = val.Elem()
|
||||
}
|
||||
|
||||
typ := val.Type()
|
||||
return fm.toSnakeCase(typ.Name())
|
||||
}
|
||||
|
||||
// GetPrimaryKey 获取主键字段名 - 默认为 "id"
|
||||
func (fm *FieldMapper) GetPrimaryKey(model interface{}) string {
|
||||
// 查找标记为主键的字段
|
||||
val := reflect.ValueOf(model)
|
||||
if val.Kind() == reflect.Ptr {
|
||||
val = val.Elem()
|
||||
}
|
||||
|
||||
typ := val.Type()
|
||||
for i := 0; i < val.NumField(); i++ {
|
||||
field := typ.Field(i)
|
||||
|
||||
// 检查是否是 ID 字段
|
||||
fieldName := field.Name
|
||||
if fieldName == "ID" || fieldName == "Id" || fieldName == "id" {
|
||||
dbTag := field.Tag.Get("db")
|
||||
if dbTag != "" && dbTag != "-" {
|
||||
return dbTag
|
||||
}
|
||||
return "id"
|
||||
}
|
||||
|
||||
// 检查是否有 primary 标签
|
||||
if field.Tag.Get("primary") == "true" {
|
||||
dbTag := field.Tag.Get("db")
|
||||
if dbTag != "" {
|
||||
return dbTag
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return "id" // 默认返回 id
|
||||
}
|
||||
|
||||
// GetFields 获取所有字段信息 - 用于生成 SQL 语句
|
||||
func (fm *FieldMapper) GetFields(model interface{}) []FieldInfo {
|
||||
var fields []FieldInfo
|
||||
|
||||
val := reflect.ValueOf(model)
|
||||
if val.Kind() == reflect.Ptr {
|
||||
val = val.Elem()
|
||||
}
|
||||
|
||||
typ := val.Type()
|
||||
|
||||
// 遍历所有字段
|
||||
for i := 0; i < val.NumField(); i++ {
|
||||
field := typ.Field(i)
|
||||
|
||||
// 跳过未导出的字段
|
||||
if !field.IsExported() {
|
||||
continue
|
||||
}
|
||||
|
||||
// 获取 db 标签
|
||||
dbTag := field.Tag.Get("db")
|
||||
if dbTag == "" || dbTag == "-" {
|
||||
continue
|
||||
}
|
||||
|
||||
// 创建字段信息
|
||||
info := FieldInfo{
|
||||
Name: field.Name,
|
||||
Column: dbTag,
|
||||
Type: fm.getTypeName(field.Type),
|
||||
DbType: fm.mapToDbType(field.Type),
|
||||
Tag: dbTag,
|
||||
}
|
||||
|
||||
// 检查是否是主键
|
||||
if field.Tag.Get("primary") == "true" ||
|
||||
field.Name == "ID" || field.Name == "Id" {
|
||||
info.IsPrimary = true
|
||||
}
|
||||
|
||||
// 检查是否是自增
|
||||
if field.Tag.Get("auto") == "true" {
|
||||
info.IsAuto = true
|
||||
}
|
||||
|
||||
fields = append(fields, info)
|
||||
}
|
||||
|
||||
return fields
|
||||
}
|
||||
|
||||
// isZeroValue 检查是否是零值
|
||||
func (fm *FieldMapper) isZeroValue(v reflect.Value) bool {
|
||||
switch v.Kind() {
|
||||
case reflect.Array, reflect.Map, reflect.Slice, reflect.String:
|
||||
return v.Len() == 0
|
||||
case reflect.Bool:
|
||||
return !v.Bool()
|
||||
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
|
||||
return v.Int() == 0
|
||||
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64, reflect.Uintptr:
|
||||
return v.Uint() == 0
|
||||
case reflect.Float32, reflect.Float64:
|
||||
return v.Float() == 0
|
||||
case reflect.Interface, reflect.Ptr:
|
||||
return v.IsNil()
|
||||
case reflect.Struct:
|
||||
// 特殊处理 time.Time
|
||||
if t, ok := v.Interface().(time.Time); ok {
|
||||
return t.IsZero()
|
||||
}
|
||||
return false
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// getTypeName 获取类型的名称
|
||||
func (fm *FieldMapper) getTypeName(t reflect.Type) string {
|
||||
return t.String()
|
||||
}
|
||||
|
||||
// mapToDbType 将 Go 类型映射到数据库类型
|
||||
func (fm *FieldMapper) mapToDbType(t reflect.Type) string {
|
||||
switch t.Kind() {
|
||||
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
|
||||
return "BIGINT"
|
||||
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
|
||||
return "BIGINT UNSIGNED"
|
||||
case reflect.Float32, reflect.Float64:
|
||||
return "DECIMAL"
|
||||
case reflect.Bool:
|
||||
return "TINYINT"
|
||||
case reflect.String:
|
||||
return "VARCHAR(255)"
|
||||
default:
|
||||
// 特殊类型
|
||||
if t.PkgPath() == "time" && t.Name() == "Time" {
|
||||
return "DATETIME"
|
||||
}
|
||||
return "TEXT"
|
||||
}
|
||||
}
|
||||
|
||||
// toSnakeCase 将驼峰命名转换为下划线命名
|
||||
func (fm *FieldMapper) toSnakeCase(str string) string {
|
||||
var result strings.Builder
|
||||
|
||||
for i, r := range str {
|
||||
if r >= 'A' && r <= 'Z' {
|
||||
if i > 0 {
|
||||
result.WriteRune('_')
|
||||
}
|
||||
result.WriteRune(r + 32) // 转换为小写
|
||||
} else {
|
||||
result.WriteRune(r)
|
||||
}
|
||||
}
|
||||
|
||||
return result.String()
|
||||
}
|
||||
@@ -0,0 +1,292 @@
|
||||
package core
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Migrator 迁移管理器实现 - 处理数据库架构的自动迁移
|
||||
type Migrator struct {
|
||||
db *Database // 数据库连接实例
|
||||
}
|
||||
|
||||
// NewMigrator 创建迁移管理器实例
|
||||
func NewMigrator(db *Database) IMigrator {
|
||||
return &Migrator{db: db}
|
||||
}
|
||||
|
||||
// AutoMigrate 自动迁移 - 根据模型自动创建或更新数据库表结构
|
||||
func (m *Migrator) AutoMigrate(models ...interface{}) error {
|
||||
for _, model := range models {
|
||||
if err := m.CreateTable(model); err != nil {
|
||||
return fmt.Errorf("创建表失败:%w", err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// CreateTable 创建表 - 根据模型创建数据库表
|
||||
func (m *Migrator) CreateTable(model interface{}) error {
|
||||
mapper := NewFieldMapper()
|
||||
|
||||
// 获取表名
|
||||
tableName := mapper.GetTableName(model)
|
||||
|
||||
// 获取字段信息
|
||||
fields := mapper.GetFields(model)
|
||||
if len(fields) == 0 {
|
||||
return fmt.Errorf("模型没有有效的字段")
|
||||
}
|
||||
|
||||
// 生成 CREATE TABLE SQL
|
||||
var sqlBuilder strings.Builder
|
||||
sqlBuilder.WriteString(fmt.Sprintf("CREATE TABLE IF NOT EXISTS %s (", tableName))
|
||||
|
||||
columnDefs := make([]string, 0)
|
||||
for _, field := range fields {
|
||||
colDef := fmt.Sprintf("%s %s", field.Column, field.DbType)
|
||||
|
||||
// 添加主键约束
|
||||
if field.IsPrimary {
|
||||
colDef += " PRIMARY KEY"
|
||||
if field.IsAuto {
|
||||
colDef += " AUTOINCREMENT"
|
||||
}
|
||||
}
|
||||
|
||||
// 添加 NOT NULL 约束(可选)
|
||||
// colDef += " NOT NULL"
|
||||
|
||||
columnDefs = append(columnDefs, colDef)
|
||||
}
|
||||
|
||||
sqlBuilder.WriteString(strings.Join(columnDefs, ", "))
|
||||
sqlBuilder.WriteString(")")
|
||||
|
||||
createSQL := sqlBuilder.String()
|
||||
|
||||
if m.db.debug {
|
||||
fmt.Printf("[Magic-ORM] CREATE TABLE SQL: %s\n", createSQL)
|
||||
}
|
||||
|
||||
// 执行 SQL
|
||||
_, err := m.db.db.Exec(createSQL)
|
||||
if err != nil {
|
||||
return fmt.Errorf("执行 CREATE TABLE 失败:%w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// DropTable 删除表 - 删除指定的数据库表
|
||||
func (m *Migrator) DropTable(model interface{}) error {
|
||||
mapper := NewFieldMapper()
|
||||
tableName := mapper.GetTableName(model)
|
||||
|
||||
dropSQL := fmt.Sprintf("DROP TABLE IF EXISTS %s", tableName)
|
||||
|
||||
if m.db.debug {
|
||||
fmt.Printf("[Magic-ORM] DROP TABLE SQL: %s\n", dropSQL)
|
||||
}
|
||||
|
||||
_, err := m.db.db.Exec(dropSQL)
|
||||
if err != nil {
|
||||
return fmt.Errorf("执行 DROP TABLE 失败:%w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// HasTable 检查表是否存在 - 验证数据库中是否已存在指定表
|
||||
func (m *Migrator) HasTable(model interface{}) (bool, error) {
|
||||
mapper := NewFieldMapper()
|
||||
tableName := mapper.GetTableName(model)
|
||||
|
||||
// SQLite 检查表是否存在的 SQL
|
||||
checkSQL := `SELECT COUNT(*) FROM sqlite_master WHERE type='table' AND name=?`
|
||||
|
||||
var count int
|
||||
err := m.db.db.QueryRow(checkSQL, tableName).Scan(&count)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("检查表是否存在失败:%w", err)
|
||||
}
|
||||
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
// RenameTable 重命名表 - 修改数据库表的名称
|
||||
func (m *Migrator) RenameTable(oldName, newName string) error {
|
||||
renameSQL := fmt.Sprintf("ALTER TABLE %s RENAME TO %s", oldName, newName)
|
||||
|
||||
if m.db.debug {
|
||||
fmt.Printf("[Magic-ORM] RENAME TABLE SQL: %s\n", renameSQL)
|
||||
}
|
||||
|
||||
_, err := m.db.db.Exec(renameSQL)
|
||||
if err != nil {
|
||||
return fmt.Errorf("重命名表失败:%w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// AddColumn 添加列 - 向表中添加新的字段
|
||||
func (m *Migrator) AddColumn(model interface{}, field string) error {
|
||||
mapper := NewFieldMapper()
|
||||
tableName := mapper.GetTableName(model)
|
||||
|
||||
// 获取字段信息
|
||||
fields := mapper.GetFields(model)
|
||||
var targetField *FieldInfo
|
||||
|
||||
for _, f := range fields {
|
||||
if f.Name == field || f.Column == field {
|
||||
targetField = &f
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if targetField == nil {
|
||||
return fmt.Errorf("字段不存在:%s", field)
|
||||
}
|
||||
|
||||
addSQL := fmt.Sprintf("ALTER TABLE %s ADD COLUMN %s %s",
|
||||
tableName, targetField.Column, targetField.DbType)
|
||||
|
||||
if m.db.debug {
|
||||
fmt.Printf("[Magic-ORM] ADD COLUMN SQL: %s\n", addSQL)
|
||||
}
|
||||
|
||||
_, err := m.db.db.Exec(addSQL)
|
||||
if err != nil {
|
||||
return fmt.Errorf("添加列失败:%w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// DropColumn 删除列 - 从表中删除指定的字段
|
||||
func (m *Migrator) DropColumn(model interface{}, field string) error {
|
||||
mapper := NewFieldMapper()
|
||||
tableName := mapper.GetTableName(model)
|
||||
|
||||
// SQLite 不直接支持 DROP COLUMN,需要重建表
|
||||
// 这里使用简化方案:创建新表 -> 复制数据 -> 删除旧表 -> 重命名
|
||||
|
||||
_ = tableName // 避免编译错误
|
||||
return fmt.Errorf("SQLite 不支持直接删除列,需要手动重建表")
|
||||
}
|
||||
|
||||
// HasColumn 检查列是否存在 - 验证表中是否已存在指定字段
|
||||
func (m *Migrator) HasColumn(model interface{}, field string) (bool, error) {
|
||||
mapper := NewFieldMapper()
|
||||
tableName := mapper.GetTableName(model)
|
||||
|
||||
// SQLite 检查列是否存在的 SQL
|
||||
checkSQL := `PRAGMA table_info(` + tableName + `)`
|
||||
|
||||
rows, err := m.db.db.Query(checkSQL)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("检查列失败:%w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
for rows.Next() {
|
||||
var cid int
|
||||
var name string
|
||||
var typ string
|
||||
var notNull int
|
||||
var dfltValue interface{}
|
||||
var pk int
|
||||
|
||||
if err := rows.Scan(&cid, &name, &typ, ¬Null, &dfltValue, &pk); err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
if name == field {
|
||||
return true, nil
|
||||
}
|
||||
}
|
||||
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// RenameColumn 重命名列 - 修改表中字段的名称
|
||||
func (m *Migrator) RenameColumn(model interface{}, oldField, newField string) error {
|
||||
mapper := NewFieldMapper()
|
||||
tableName := mapper.GetTableName(model)
|
||||
|
||||
// SQLite 3.25.0+ 支持 ALTER TABLE ... RENAME COLUMN
|
||||
renameSQL := fmt.Sprintf("ALTER TABLE %s RENAME COLUMN %s TO %s",
|
||||
tableName, oldField, newField)
|
||||
|
||||
if m.db.debug {
|
||||
fmt.Printf("[Magic-ORM] RENAME COLUMN SQL: %s\n", renameSQL)
|
||||
}
|
||||
|
||||
_, err := m.db.db.Exec(renameSQL)
|
||||
if err != nil {
|
||||
return fmt.Errorf("重命名列失败:%w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// CreateIndex 创建索引 - 为表中的字段创建索引
|
||||
func (m *Migrator) CreateIndex(model interface{}, field string) error {
|
||||
mapper := NewFieldMapper()
|
||||
tableName := mapper.GetTableName(model)
|
||||
|
||||
indexName := fmt.Sprintf("idx_%s_%s", tableName, field)
|
||||
createSQL := fmt.Sprintf("CREATE INDEX IF NOT EXISTS %s ON %s (%s)",
|
||||
indexName, tableName, field)
|
||||
|
||||
if m.db.debug {
|
||||
fmt.Printf("[Magic-ORM] CREATE INDEX SQL: %s\n", createSQL)
|
||||
}
|
||||
|
||||
_, err := m.db.db.Exec(createSQL)
|
||||
if err != nil {
|
||||
return fmt.Errorf("创建索引失败:%w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// DropIndex 删除索引 - 删除表中的指定索引
|
||||
func (m *Migrator) DropIndex(model interface{}, field string) error {
|
||||
mapper := NewFieldMapper()
|
||||
tableName := mapper.GetTableName(model)
|
||||
|
||||
indexName := fmt.Sprintf("idx_%s_%s", tableName, field)
|
||||
dropSQL := fmt.Sprintf("DROP INDEX IF EXISTS %s", indexName)
|
||||
|
||||
if m.db.debug {
|
||||
fmt.Printf("[Magic-ORM] DROP INDEX SQL: %s\n", dropSQL)
|
||||
}
|
||||
|
||||
_, err := m.db.db.Exec(dropSQL)
|
||||
if err != nil {
|
||||
return fmt.Errorf("删除索引失败:%w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// HasIndex 检查索引是否存在 - 验证表中是否已存在指定索引
|
||||
func (m *Migrator) HasIndex(model interface{}, field string) (bool, error) {
|
||||
mapper := NewFieldMapper()
|
||||
tableName := mapper.GetTableName(model)
|
||||
|
||||
indexName := fmt.Sprintf("idx_%s_%s", tableName, field)
|
||||
|
||||
checkSQL := `SELECT COUNT(*) FROM sqlite_master WHERE type='index' AND name=?`
|
||||
|
||||
var count int
|
||||
err := m.db.db.QueryRow(checkSQL, indexName).Scan(&count)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("检查索引失败:%w", err)
|
||||
}
|
||||
|
||||
return count > 0, nil
|
||||
}
|
||||
@@ -0,0 +1,548 @@
|
||||
package core
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// QueryBuilder 查询构建器实现 - 提供流畅的链式查询构建能力
|
||||
type QueryBuilder struct {
|
||||
db *Database // 数据库连接实例
|
||||
table string // 表名
|
||||
model interface{} // 模型对象
|
||||
whereSQL string // WHERE 条件 SQL
|
||||
whereArgs []interface{} // WHERE 条件参数
|
||||
selectCols []string // 选择的字段列表
|
||||
orderSQL string // ORDER BY SQL
|
||||
limit int // LIMIT 限制数量
|
||||
offset int // OFFSET 偏移量
|
||||
groupSQL string // GROUP BY SQL
|
||||
havingSQL string // HAVING 条件 SQL
|
||||
havingArgs []interface{} // HAVING 条件参数
|
||||
joinSQL string // JOIN SQL
|
||||
joinArgs []interface{} // JOIN 参数
|
||||
debug bool // 调试模式开关
|
||||
dryRun bool // 干跑模式开关
|
||||
unscoped bool // 忽略软删除开关
|
||||
tx *sql.Tx // 事务对象(如果在事务中)
|
||||
}
|
||||
|
||||
// 同步池优化 - 复用 slice 减少内存分配
|
||||
var whereArgsPool = sync.Pool{
|
||||
New: func() interface{} {
|
||||
return make([]interface{}, 0, 10)
|
||||
},
|
||||
}
|
||||
|
||||
var joinArgsPool = sync.Pool{
|
||||
New: func() interface{} {
|
||||
return make([]interface{}, 0, 5)
|
||||
},
|
||||
}
|
||||
|
||||
// Model 基于模型创建查询
|
||||
func (d *Database) Model(model interface{}) IQuery {
|
||||
return &QueryBuilder{
|
||||
db: d,
|
||||
model: model,
|
||||
}
|
||||
}
|
||||
|
||||
// Table 基于表名创建查询
|
||||
func (d *Database) Table(name string) IQuery {
|
||||
return &QueryBuilder{
|
||||
db: d,
|
||||
table: name,
|
||||
}
|
||||
}
|
||||
|
||||
// Where 添加 WHERE 条件 - 性能优化版本
|
||||
func (q *QueryBuilder) Where(query string, args ...interface{}) IQuery {
|
||||
if q.whereSQL == "" {
|
||||
q.whereSQL = query
|
||||
} else {
|
||||
// 使用 strings.Builder 优化字符串拼接
|
||||
var builder strings.Builder
|
||||
builder.Grow(len(q.whereSQL) + 5 + len(query)) // 预分配内存
|
||||
builder.WriteString(q.whereSQL)
|
||||
builder.WriteString(" AND ")
|
||||
builder.WriteString(query)
|
||||
q.whereSQL = builder.String()
|
||||
}
|
||||
q.whereArgs = append(q.whereArgs, args...)
|
||||
return q
|
||||
}
|
||||
|
||||
// Or 添加 OR 条件 - 性能优化版本
|
||||
func (q *QueryBuilder) Or(query string, args ...interface{}) IQuery {
|
||||
if q.whereSQL == "" {
|
||||
q.whereSQL = query
|
||||
} else {
|
||||
// 使用 strings.Builder 优化字符串拼接
|
||||
var builder strings.Builder
|
||||
builder.Grow(len(q.whereSQL) + 10 + len(query)) // 预分配内存
|
||||
builder.WriteString(" (")
|
||||
builder.WriteString(q.whereSQL)
|
||||
builder.WriteString(") OR ")
|
||||
builder.WriteString(query)
|
||||
q.whereSQL = builder.String()
|
||||
}
|
||||
q.whereArgs = append(q.whereArgs, args...)
|
||||
return q
|
||||
}
|
||||
|
||||
// And 添加 AND 条件(同 Where)
|
||||
func (q *QueryBuilder) And(query string, args ...interface{}) IQuery {
|
||||
return q.Where(query, args...)
|
||||
}
|
||||
|
||||
// Select 选择要查询的字段
|
||||
func (q *QueryBuilder) Select(fields ...string) IQuery {
|
||||
q.selectCols = fields
|
||||
return q
|
||||
}
|
||||
|
||||
// Omit 排除指定的字段(暂未实现)
|
||||
func (q *QueryBuilder) Omit(fields ...string) IQuery {
|
||||
// TODO: 实现字段排除逻辑,生成 SELECT 时排除这些字段
|
||||
return q
|
||||
}
|
||||
|
||||
// Order 设置排序规则
|
||||
func (q *QueryBuilder) Order(order string) IQuery {
|
||||
q.orderSQL = order
|
||||
return q
|
||||
}
|
||||
|
||||
// OrderBy 按指定字段和方向排序
|
||||
func (q *QueryBuilder) OrderBy(field string, direction string) IQuery {
|
||||
q.orderSQL = field + " " + direction
|
||||
return q
|
||||
}
|
||||
|
||||
// Limit 限制返回数量
|
||||
func (q *QueryBuilder) Limit(limit int) IQuery {
|
||||
q.limit = limit
|
||||
return q
|
||||
}
|
||||
|
||||
// Offset 设置偏移量
|
||||
func (q *QueryBuilder) Offset(offset int) IQuery {
|
||||
q.offset = offset
|
||||
return q
|
||||
}
|
||||
|
||||
// Page 分页查询
|
||||
func (q *QueryBuilder) Page(page, pageSize int) IQuery {
|
||||
q.limit = pageSize
|
||||
q.offset = (page - 1) * pageSize
|
||||
return q
|
||||
}
|
||||
|
||||
// Group 设置分组字段
|
||||
func (q *QueryBuilder) Group(group string) IQuery {
|
||||
q.groupSQL = group
|
||||
return q
|
||||
}
|
||||
|
||||
// Having 添加 HAVING 条件
|
||||
func (q *QueryBuilder) Having(having string, args ...interface{}) IQuery {
|
||||
q.havingSQL = having
|
||||
q.havingArgs = args
|
||||
return q
|
||||
}
|
||||
|
||||
// Join 添加 JOIN 连接 - 性能优化版本
|
||||
func (q *QueryBuilder) Join(join string, args ...interface{}) IQuery {
|
||||
if q.joinSQL == "" {
|
||||
q.joinSQL = join
|
||||
} else {
|
||||
// 使用 strings.Builder 优化字符串拼接
|
||||
var builder strings.Builder
|
||||
builder.Grow(len(q.joinSQL) + 1 + len(join)) // 预分配内存
|
||||
builder.WriteString(q.joinSQL)
|
||||
builder.WriteByte(' ')
|
||||
builder.WriteString(join)
|
||||
q.joinSQL = builder.String()
|
||||
}
|
||||
q.joinArgs = append(q.joinArgs, args...)
|
||||
return q
|
||||
}
|
||||
|
||||
// LeftJoin 左连接
|
||||
func (q *QueryBuilder) LeftJoin(table, on string) IQuery {
|
||||
return q.Join("LEFT JOIN " + table + " ON " + on)
|
||||
}
|
||||
|
||||
// RightJoin 右连接
|
||||
func (q *QueryBuilder) RightJoin(table, on string) IQuery {
|
||||
return q.Join("RIGHT JOIN " + table + " ON " + on)
|
||||
}
|
||||
|
||||
// InnerJoin 内连接
|
||||
func (q *QueryBuilder) InnerJoin(table, on string) IQuery {
|
||||
return q.Join("INNER JOIN " + table + " ON " + on)
|
||||
}
|
||||
|
||||
// Preload 预加载关联数据(暂未实现)
|
||||
func (q *QueryBuilder) Preload(relation string, conditions ...interface{}) IQuery {
|
||||
// TODO: 实现预加载逻辑
|
||||
return q
|
||||
}
|
||||
|
||||
// First 查询第一条记录
|
||||
func (q *QueryBuilder) First(result interface{}) error {
|
||||
q.limit = 1
|
||||
return q.Find(result)
|
||||
}
|
||||
|
||||
// Find 查询多条记录
|
||||
func (q *QueryBuilder) Find(result interface{}) error {
|
||||
sqlStr, args := q.BuildSelect()
|
||||
|
||||
// 调试模式打印 SQL
|
||||
if q.debug || (q.db != nil && q.db.debug) {
|
||||
fmt.Printf("[Magic-ORM] SQL: %s\n[Magic-ORM] Args: %v\n", sqlStr, args)
|
||||
}
|
||||
|
||||
// 干跑模式不执行 SQL
|
||||
if q.dryRun {
|
||||
return nil
|
||||
}
|
||||
|
||||
var rows *sql.Rows
|
||||
var err error
|
||||
|
||||
// 判断是否在事务中
|
||||
if q.tx != nil {
|
||||
rows, err = q.tx.Query(sqlStr, args...)
|
||||
} else if q.db != nil && q.db.db != nil {
|
||||
rows, err = q.db.db.Query(sqlStr, args...)
|
||||
} else {
|
||||
return fmt.Errorf("数据库连接未初始化")
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf("查询失败:%w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
// TODO: 实现结果映射逻辑
|
||||
// 使用 FieldMapper 将查询结果映射到 result
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Count 统计记录数量
|
||||
func (q *QueryBuilder) Count(count *int64) IQuery {
|
||||
// 构建 COUNT 查询
|
||||
originalSelect := q.selectCols
|
||||
q.selectCols = []string{"COUNT(*)"}
|
||||
|
||||
sqlStr, args := q.BuildSelect()
|
||||
|
||||
// 调试模式
|
||||
if q.debug || (q.db != nil && q.db.debug) {
|
||||
fmt.Printf("[Magic-ORM] COUNT SQL: %s\n[Magic-ORM] Args: %v\n", sqlStr, args)
|
||||
}
|
||||
|
||||
// 干跑模式
|
||||
if q.dryRun {
|
||||
return q
|
||||
}
|
||||
|
||||
var err error
|
||||
if q.tx != nil {
|
||||
err = q.tx.QueryRow(sqlStr, args...).Scan(count)
|
||||
} else if q.db != nil && q.db.db != nil {
|
||||
err = q.db.db.QueryRow(sqlStr, args...).Scan(count)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
fmt.Printf("[Magic-ORM] Count 错误:%v\n", err)
|
||||
}
|
||||
|
||||
// 恢复原来的选择字段
|
||||
q.selectCols = originalSelect
|
||||
return q
|
||||
}
|
||||
|
||||
// Exists 检查记录是否存在
|
||||
func (q *QueryBuilder) Exists() (bool, error) {
|
||||
// 使用 LIMIT 1 优化查询
|
||||
originalLimit := q.limit
|
||||
q.limit = 1
|
||||
|
||||
sqlStr, args := q.BuildSelect()
|
||||
|
||||
// 调试模式
|
||||
if q.debug || (q.db != nil && q.db.debug) {
|
||||
fmt.Printf("[Magic-ORM] EXISTS SQL: %s\n[Magic-ORM] Args: %v\n", sqlStr, args)
|
||||
}
|
||||
|
||||
// 干跑模式
|
||||
if q.dryRun {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
var rows *sql.Rows
|
||||
var err error
|
||||
|
||||
if q.tx != nil {
|
||||
rows, err = q.tx.Query(sqlStr, args...)
|
||||
} else if q.db != nil && q.db.db != nil {
|
||||
rows, err = q.db.db.Query(sqlStr, args...)
|
||||
} else {
|
||||
return false, fmt.Errorf("数据库连接未初始化")
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
// 检查是否有结果
|
||||
exists := rows.Next()
|
||||
|
||||
// 恢复原来的 limit
|
||||
q.limit = originalLimit
|
||||
|
||||
return exists, nil
|
||||
}
|
||||
|
||||
// Updates 更新数据
|
||||
func (q *QueryBuilder) Updates(data interface{}) error {
|
||||
sqlStr, args := q.BuildUpdate(data)
|
||||
|
||||
// 调试模式打印 SQL
|
||||
if q.debug || (q.db != nil && q.db.debug) {
|
||||
fmt.Printf("[Magic-ORM] UPDATE SQL: %s\n[Magic-ORM] Args: %v\n", sqlStr, args)
|
||||
}
|
||||
|
||||
// 干跑模式不执行 SQL
|
||||
if q.dryRun {
|
||||
return nil
|
||||
}
|
||||
|
||||
var err error
|
||||
if q.tx != nil {
|
||||
_, err = q.tx.Exec(sqlStr, args...)
|
||||
} else if q.db != nil && q.db.db != nil {
|
||||
_, err = q.db.db.Exec(sqlStr, args...)
|
||||
} else {
|
||||
return fmt.Errorf("数据库连接未初始化")
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf("更新失败:%w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// UpdateColumn 更新单个字段
|
||||
func (q *QueryBuilder) UpdateColumn(column string, value interface{}) error {
|
||||
return q.Updates(map[string]interface{}{column: value})
|
||||
}
|
||||
|
||||
// Delete 删除数据
|
||||
func (q *QueryBuilder) Delete() error {
|
||||
sqlStr, args := q.BuildDelete()
|
||||
|
||||
// 调试模式打印 SQL
|
||||
if q.debug || (q.db != nil && q.db.debug) {
|
||||
fmt.Printf("[Magic-ORM] DELETE SQL: %s\n[Magic-ORM] Args: %v\n", sqlStr, args)
|
||||
}
|
||||
|
||||
// 干跑模式不执行 SQL
|
||||
if q.dryRun {
|
||||
return nil
|
||||
}
|
||||
|
||||
var err error
|
||||
if q.tx != nil {
|
||||
_, err = q.tx.Exec(sqlStr, args...)
|
||||
} else if q.db != nil && q.db.db != nil {
|
||||
_, err = q.db.db.Exec(sqlStr, args...)
|
||||
} else {
|
||||
return fmt.Errorf("数据库连接未初始化")
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf("删除失败:%w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Unscoped 忽略软删除限制
|
||||
func (q *QueryBuilder) Unscoped() IQuery {
|
||||
q.unscoped = true
|
||||
return q
|
||||
}
|
||||
|
||||
// DryRun 设置干跑模式(只生成 SQL 不执行)
|
||||
func (q *QueryBuilder) DryRun() IQuery {
|
||||
q.dryRun = true
|
||||
return q
|
||||
}
|
||||
|
||||
// Debug 设置调试模式(打印 SQL 日志)
|
||||
func (q *QueryBuilder) Debug() IQuery {
|
||||
q.debug = true
|
||||
return q
|
||||
}
|
||||
|
||||
// Build 构建 SELECT SQL 语句
|
||||
func (q *QueryBuilder) Build() (string, []interface{}) {
|
||||
return q.BuildSelect()
|
||||
}
|
||||
|
||||
// BuildSelect 构建 SELECT SQL 语句
|
||||
func (q *QueryBuilder) BuildSelect() (string, []interface{}) {
|
||||
var builder strings.Builder
|
||||
|
||||
// SELECT 部分
|
||||
builder.WriteString("SELECT ")
|
||||
if len(q.selectCols) > 0 {
|
||||
builder.WriteString(strings.Join(q.selectCols, ", "))
|
||||
} else {
|
||||
builder.WriteString("*")
|
||||
}
|
||||
|
||||
// FROM 部分
|
||||
builder.WriteString(" FROM ")
|
||||
if q.table != "" {
|
||||
builder.WriteString(q.table)
|
||||
} else if q.model != nil {
|
||||
// 从模型获取表名
|
||||
mapper := NewFieldMapper()
|
||||
builder.WriteString(mapper.GetTableName(q.model))
|
||||
} else {
|
||||
builder.WriteString("unknown_table")
|
||||
}
|
||||
|
||||
// JOIN 部分
|
||||
if q.joinSQL != "" {
|
||||
builder.WriteString(" ")
|
||||
builder.WriteString(q.joinSQL)
|
||||
}
|
||||
|
||||
// WHERE 部分
|
||||
if q.whereSQL != "" {
|
||||
builder.WriteString(" WHERE ")
|
||||
builder.WriteString(q.whereSQL)
|
||||
}
|
||||
|
||||
// GROUP BY 部分
|
||||
if q.groupSQL != "" {
|
||||
builder.WriteString(" GROUP BY ")
|
||||
builder.WriteString(q.groupSQL)
|
||||
}
|
||||
|
||||
// HAVING 部分
|
||||
if q.havingSQL != "" {
|
||||
builder.WriteString(" HAVING ")
|
||||
builder.WriteString(q.havingSQL)
|
||||
}
|
||||
|
||||
// ORDER BY 部分
|
||||
if q.orderSQL != "" {
|
||||
builder.WriteString(" ORDER BY ")
|
||||
builder.WriteString(q.orderSQL)
|
||||
}
|
||||
|
||||
// LIMIT 部分
|
||||
if q.limit > 0 {
|
||||
builder.WriteString(fmt.Sprintf(" LIMIT %d", q.limit))
|
||||
}
|
||||
|
||||
// OFFSET 部分
|
||||
if q.offset > 0 {
|
||||
builder.WriteString(fmt.Sprintf(" OFFSET %d", q.offset))
|
||||
}
|
||||
|
||||
// 合并参数
|
||||
allArgs := make([]interface{}, 0)
|
||||
allArgs = append(allArgs, q.joinArgs...)
|
||||
allArgs = append(allArgs, q.whereArgs...)
|
||||
allArgs = append(allArgs, q.havingArgs...)
|
||||
|
||||
return builder.String(), allArgs
|
||||
}
|
||||
|
||||
// BuildUpdate 构建 UPDATE SQL 语句
|
||||
func (q *QueryBuilder) BuildUpdate(data interface{}) (string, []interface{}) {
|
||||
var builder strings.Builder
|
||||
var args []interface{}
|
||||
|
||||
builder.WriteString("UPDATE ")
|
||||
if q.table != "" {
|
||||
builder.WriteString(q.table)
|
||||
} else if q.model != nil {
|
||||
mapper := NewFieldMapper()
|
||||
builder.WriteString(mapper.GetTableName(q.model))
|
||||
} else {
|
||||
builder.WriteString("unknown_table")
|
||||
}
|
||||
|
||||
builder.WriteString(" SET ")
|
||||
|
||||
// 根据 data 类型生成 SET 子句
|
||||
switch v := data.(type) {
|
||||
case map[string]interface{}:
|
||||
// map 类型,生成 key=value 对
|
||||
setParts := make([]string, 0, len(v))
|
||||
for key, value := range v {
|
||||
setParts = append(setParts, fmt.Sprintf("%s = ?", key))
|
||||
args = append(args, value)
|
||||
}
|
||||
builder.WriteString(strings.Join(setParts, ", "))
|
||||
case string:
|
||||
// string 类型,直接使用(注意:实际使用需要转义)
|
||||
builder.WriteString(v)
|
||||
default:
|
||||
// 结构体类型,使用字段映射器
|
||||
mapper := NewFieldMapper()
|
||||
columns, err := mapper.StructToColumns(data)
|
||||
if err == nil && len(columns) > 0 {
|
||||
setParts := make([]string, 0, len(columns))
|
||||
for key := range columns {
|
||||
setParts = append(setParts, fmt.Sprintf("%s = ?", key))
|
||||
args = append(args, columns[key])
|
||||
}
|
||||
builder.WriteString(strings.Join(setParts, ", "))
|
||||
}
|
||||
}
|
||||
|
||||
// WHERE 部分
|
||||
if q.whereSQL != "" {
|
||||
builder.WriteString(" WHERE ")
|
||||
builder.WriteString(q.whereSQL)
|
||||
args = append(args, q.whereArgs...)
|
||||
}
|
||||
|
||||
return builder.String(), args
|
||||
}
|
||||
|
||||
// BuildDelete 构建 DELETE SQL 语句
|
||||
func (q *QueryBuilder) BuildDelete() (string, []interface{}) {
|
||||
var builder strings.Builder
|
||||
|
||||
builder.WriteString("DELETE FROM ")
|
||||
if q.table != "" {
|
||||
builder.WriteString(q.table)
|
||||
} else if q.model != nil {
|
||||
mapper := NewFieldMapper()
|
||||
builder.WriteString(mapper.GetTableName(q.model))
|
||||
} else {
|
||||
builder.WriteString("unknown_table")
|
||||
}
|
||||
|
||||
if q.whereSQL != "" {
|
||||
builder.WriteString(" WHERE ")
|
||||
builder.WriteString(q.whereSQL)
|
||||
}
|
||||
|
||||
return builder.String(), q.whereArgs
|
||||
}
|
||||
@@ -0,0 +1,124 @@
|
||||
package core
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
)
|
||||
|
||||
// ReadWriteDB 读写分离数据库连接
|
||||
type ReadWriteDB struct {
|
||||
master *sql.DB // 主库(写)
|
||||
slaves []*sql.DB // 从库列表(读)
|
||||
policy ReadPolicy // 读负载均衡策略
|
||||
counter uint64 // 轮询计数器
|
||||
mu sync.RWMutex // 读写锁
|
||||
}
|
||||
|
||||
// NewReadWriteDB 创建读写分离数据库连接
|
||||
func NewReadWriteDB(master *sql.DB, slaves []*sql.DB, policy ReadPolicy) *ReadWriteDB {
|
||||
return &ReadWriteDB{
|
||||
master: master,
|
||||
slaves: slaves,
|
||||
policy: policy,
|
||||
}
|
||||
}
|
||||
|
||||
// GetMaster 获取主库连接(用于写操作)
|
||||
func (rw *ReadWriteDB) GetMaster() *sql.DB {
|
||||
return rw.master
|
||||
}
|
||||
|
||||
// GetSlave 获取从库连接(用于读操作)
|
||||
func (rw *ReadWriteDB) GetSlave() *sql.DB {
|
||||
rw.mu.RLock()
|
||||
defer rw.mu.RUnlock()
|
||||
|
||||
if len(rw.slaves) == 0 {
|
||||
// 没有从库,使用主库
|
||||
return rw.master
|
||||
}
|
||||
|
||||
switch rw.policy {
|
||||
case Random:
|
||||
// 随机选择一个从库
|
||||
idx := int(atomic.LoadUint64(&rw.counter)) % len(rw.slaves)
|
||||
return rw.slaves[idx]
|
||||
|
||||
case RoundRobin:
|
||||
// 轮询选择从库
|
||||
idx := int(atomic.AddUint64(&rw.counter, 1)) % len(rw.slaves)
|
||||
return rw.slaves[idx]
|
||||
|
||||
case LeastConn:
|
||||
// 选择连接数最少的从库(简化实现)
|
||||
return rw.selectLeastConn()
|
||||
|
||||
default:
|
||||
return rw.slaves[0]
|
||||
}
|
||||
}
|
||||
|
||||
// selectLeastConn 选择连接数最少的从库
|
||||
func (rw *ReadWriteDB) selectLeastConn() *sql.DB {
|
||||
if len(rw.slaves) == 0 {
|
||||
return rw.master
|
||||
}
|
||||
|
||||
minConn := -1
|
||||
selected := rw.slaves[0]
|
||||
|
||||
for _, slave := range rw.slaves {
|
||||
stats := slave.Stats()
|
||||
openConnections := stats.OpenConnections
|
||||
|
||||
if minConn == -1 || openConnections < minConn {
|
||||
minConn = openConnections
|
||||
selected = slave
|
||||
}
|
||||
}
|
||||
|
||||
return selected
|
||||
}
|
||||
|
||||
// AddSlave 添加从库
|
||||
func (rw *ReadWriteDB) AddSlave(slave *sql.DB) {
|
||||
rw.mu.Lock()
|
||||
defer rw.mu.Unlock()
|
||||
rw.slaves = append(rw.slaves, slave)
|
||||
}
|
||||
|
||||
// RemoveSlave 移除从库
|
||||
func (rw *ReadWriteDB) RemoveSlave(slave *sql.DB) {
|
||||
rw.mu.Lock()
|
||||
defer rw.mu.Unlock()
|
||||
|
||||
for i, s := range rw.slaves {
|
||||
if s == slave {
|
||||
rw.slaves = append(rw.slaves[:i], rw.slaves[i+1:]...)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Close 关闭所有连接
|
||||
func (rw *ReadWriteDB) Close() error {
|
||||
rw.mu.Lock()
|
||||
defer rw.mu.Unlock()
|
||||
|
||||
// 关闭主库
|
||||
if rw.master != nil {
|
||||
if err := rw.master.Close(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
// 关闭所有从库
|
||||
for _, slave := range rw.slaves {
|
||||
if err := slave.Close(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,199 @@
|
||||
package core
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"reflect"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// RelationType 关联类型
|
||||
type RelationType int
|
||||
|
||||
const (
|
||||
HasOne RelationType = iota // 一对一
|
||||
HasMany // 一对多
|
||||
BelongsTo // 多对一
|
||||
ManyToMany // 多对多
|
||||
)
|
||||
|
||||
// RelationInfo 关联信息
|
||||
type RelationInfo struct {
|
||||
Type RelationType // 关联类型
|
||||
Field string // 字段名
|
||||
Model interface{} // 关联的模型
|
||||
FK string // 外键
|
||||
PK string // 主键
|
||||
JoinTable string // 中间表(多对多)
|
||||
JoinFK string // 中间表外键
|
||||
JoinJoinFK string // 中间表关联外键
|
||||
}
|
||||
|
||||
// RelationLoader 关联加载器 - 处理模型关联的预加载
|
||||
type RelationLoader struct {
|
||||
db *Database
|
||||
}
|
||||
|
||||
// NewRelationLoader 创建关联加载器实例
|
||||
func NewRelationLoader(db *Database) *RelationLoader {
|
||||
return &RelationLoader{db: db}
|
||||
}
|
||||
|
||||
// Preload 预加载关联数据
|
||||
func (rl *RelationLoader) Preload(models interface{}, relation string, conditions ...interface{}) error {
|
||||
// 获取反射对象
|
||||
modelsVal := reflect.ValueOf(models)
|
||||
if modelsVal.Kind() != reflect.Ptr {
|
||||
return fmt.Errorf("models 必须是指针类型")
|
||||
}
|
||||
|
||||
elem := modelsVal.Elem()
|
||||
if elem.Kind() != reflect.Slice {
|
||||
return fmt.Errorf("models 必须是指向 Slice 的指针")
|
||||
}
|
||||
|
||||
if elem.Len() == 0 {
|
||||
return nil // 空 Slice,无需加载
|
||||
}
|
||||
|
||||
// 解析关联关系
|
||||
relationInfo, err := rl.parseRelation(elem.Index(0).Interface(), relation)
|
||||
if err != nil {
|
||||
return fmt.Errorf("解析关联失败:%w", err)
|
||||
}
|
||||
|
||||
// 根据关联类型加载数据
|
||||
switch relationInfo.Type {
|
||||
case HasOne:
|
||||
return rl.loadHasOne(elem, relationInfo)
|
||||
case HasMany:
|
||||
return rl.loadHasMany(elem, relationInfo)
|
||||
case BelongsTo:
|
||||
return rl.loadBelongsTo(elem, relationInfo)
|
||||
case ManyToMany:
|
||||
return rl.loadManyToMany(elem, relationInfo)
|
||||
default:
|
||||
return fmt.Errorf("不支持的关联类型:%v", relationInfo.Type)
|
||||
}
|
||||
}
|
||||
|
||||
// parseRelation 解析关联关系
|
||||
func (rl *RelationLoader) parseRelation(model interface{}, relation string) (*RelationInfo, error) {
|
||||
// TODO: 从结构体标签中解析关联信息
|
||||
// 示例:
|
||||
// type Order struct {
|
||||
// User User `gorm:"ForeignKey:user_id;References:id"`
|
||||
// Items []Item `gorm:"ForeignKey:order_id;References:id"`
|
||||
// }
|
||||
|
||||
// 这里提供简化的实现
|
||||
return &RelationInfo{
|
||||
Type: HasOne, // 默认假设为一对一
|
||||
Field: relation,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// loadHasOne 加载一对一关联
|
||||
func (rl *RelationLoader) loadHasOne(models reflect.Value, relation *RelationInfo) error {
|
||||
// 收集所有主键值
|
||||
pkValues := make([]interface{}, 0, models.Len())
|
||||
for i := 0; i < models.Len(); i++ {
|
||||
model := models.Index(i).Interface()
|
||||
pk := rl.getFieldValue(model, "ID")
|
||||
if pk != nil {
|
||||
pkValues = append(pkValues, pk)
|
||||
}
|
||||
}
|
||||
|
||||
if len(pkValues) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
// 查询关联数据
|
||||
query := rl.db.Model(relation.Model)
|
||||
query.Where(fmt.Sprintf("%s IN ?", relation.FK), pkValues)
|
||||
|
||||
// TODO: 执行查询并映射到模型
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// loadHasMany 加载一对多关联
|
||||
func (rl *RelationLoader) loadHasMany(models reflect.Value, relation *RelationInfo) error {
|
||||
// 类似 HasOne,但结果需要映射到 Slice
|
||||
return rl.loadHasOne(models, relation)
|
||||
}
|
||||
|
||||
// loadBelongsTo 加载多对一关联
|
||||
func (rl *RelationLoader) loadBelongsTo(models reflect.Value, relation *RelationInfo) error {
|
||||
// 收集所有外键值
|
||||
fkValues := make([]interface{}, 0, models.Len())
|
||||
for i := 0; i < models.Len(); i++ {
|
||||
model := models.Index(i).Interface()
|
||||
fk := rl.getFieldValue(model, relation.FK)
|
||||
if fk != nil {
|
||||
fkValues = append(fkValues, fk)
|
||||
}
|
||||
}
|
||||
|
||||
if len(fkValues) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
// 查询关联数据
|
||||
query := rl.db.Model(relation.Model)
|
||||
query.Where(fmt.Sprintf("id IN ?"), fkValues)
|
||||
|
||||
// TODO: 执行查询并映射到模型
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// loadManyToMany 加载多对多关联
|
||||
func (rl *RelationLoader) loadManyToMany(models reflect.Value, relation *RelationInfo) error {
|
||||
// 多对多需要通过中间表查询
|
||||
// SELECT * FROM table WHERE id IN (
|
||||
// SELECT join_fk FROM join_table WHERE fk IN (pk_values)
|
||||
// )
|
||||
|
||||
return fmt.Errorf("多对多关联暂未实现")
|
||||
}
|
||||
|
||||
// getFieldValue 获取字段的值
|
||||
func (rl *RelationLoader) getFieldValue(model interface{}, fieldName string) interface{} {
|
||||
val := reflect.ValueOf(model)
|
||||
if val.Kind() == reflect.Ptr {
|
||||
val = val.Elem()
|
||||
}
|
||||
|
||||
field := val.FieldByName(fieldName)
|
||||
if field.IsValid() && field.CanInterface() {
|
||||
return field.Interface()
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// getRelationTags 从结构体字段提取关联标签信息
|
||||
func getRelationTags(structType reflect.Type, fieldName string) map[string]string {
|
||||
tags := make(map[string]string)
|
||||
|
||||
for i := 0; i < structType.NumField(); i++ {
|
||||
field := structType.Field(i)
|
||||
if field.Name == fieldName {
|
||||
gormTag := field.Tag.Get("gorm")
|
||||
if gormTag != "" {
|
||||
// 解析 GORM 风格的标签
|
||||
parts := strings.Split(gormTag, ";")
|
||||
for _, part := range parts {
|
||||
kv := strings.Split(part, ":")
|
||||
if len(kv) == 2 {
|
||||
tags[strings.TrimSpace(kv[0])] = strings.TrimSpace(kv[1])
|
||||
}
|
||||
}
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
return tags
|
||||
}
|
||||
@@ -0,0 +1,173 @@
|
||||
package core
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"reflect"
|
||||
)
|
||||
|
||||
// ResultSetMapper 结果集映射器 - 将查询结果映射到 Slice 或 Struct
|
||||
type ResultSetMapper struct {
|
||||
fieldMapper IFieldMapper
|
||||
}
|
||||
|
||||
// NewResultSetMapper 创建结果集映射器实例
|
||||
func NewResultSetMapper() *ResultSetMapper {
|
||||
return &ResultSetMapper{
|
||||
fieldMapper: NewFieldMapper(),
|
||||
}
|
||||
}
|
||||
|
||||
// MapToSlice 将查询结果映射到 Slice
|
||||
func (rsm *ResultSetMapper) MapToSlice(rows *sql.Rows, result interface{}) error {
|
||||
// 获取反射对象
|
||||
resultVal := reflect.ValueOf(result)
|
||||
|
||||
// 必须是指针类型
|
||||
if resultVal.Kind() != reflect.Ptr {
|
||||
return fmt.Errorf("result 必须是指针类型")
|
||||
}
|
||||
|
||||
elem := resultVal.Elem()
|
||||
|
||||
// 必须是 Slice 类型
|
||||
if elem.Kind() != reflect.Slice {
|
||||
return fmt.Errorf("result 必须是指向 Slice 的指针")
|
||||
}
|
||||
|
||||
// 获取 Slice 的元素类型
|
||||
sliceType := elem.Type().Elem()
|
||||
var isPtr bool
|
||||
if sliceType.Kind() == reflect.Ptr {
|
||||
isPtr = true
|
||||
sliceType = sliceType.Elem()
|
||||
}
|
||||
|
||||
if sliceType.Kind() != reflect.Struct {
|
||||
return fmt.Errorf("Slice 的元素必须是结构体")
|
||||
}
|
||||
|
||||
// 获取列信息
|
||||
columns, err := rows.Columns()
|
||||
if err != nil {
|
||||
return fmt.Errorf("获取列信息失败:%w", err)
|
||||
}
|
||||
|
||||
// 建立列名到字段的映射
|
||||
fieldMap := make(map[string]int)
|
||||
for i := 0; i < sliceType.NumField(); i++ {
|
||||
field := sliceType.Field(i)
|
||||
dbTag := field.Tag.Get("db")
|
||||
|
||||
if dbTag != "" && dbTag != "-" {
|
||||
// 使用 db 标签
|
||||
fieldMap[dbTag] = i
|
||||
// 同时存储小写版本用于不区分大小写的匹配
|
||||
fieldMap[dbTag] = i
|
||||
} else {
|
||||
// 使用字段名的小写形式
|
||||
fieldMap[sliceType.Field(i).Name] = i
|
||||
}
|
||||
}
|
||||
|
||||
// 循环读取每一行数据
|
||||
for rows.Next() {
|
||||
// 创建新的结构体实例
|
||||
var item reflect.Value
|
||||
if isPtr {
|
||||
item = reflect.New(sliceType)
|
||||
} else {
|
||||
item = reflect.New(sliceType).Elem()
|
||||
}
|
||||
|
||||
// 创建扫描目标
|
||||
scanTargets := make([]interface{}, len(columns))
|
||||
|
||||
for i, col := range columns {
|
||||
// 查找对应的字段
|
||||
var fieldIndex int
|
||||
found := false
|
||||
|
||||
// 尝试精确匹配
|
||||
if idx, ok := fieldMap[col]; ok {
|
||||
fieldIndex = idx
|
||||
found = true
|
||||
} else {
|
||||
// 尝试不区分大小写匹配
|
||||
colLower := col
|
||||
for key, idx := range fieldMap {
|
||||
if key == colLower {
|
||||
fieldIndex = idx
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if found {
|
||||
var field reflect.Value
|
||||
if isPtr {
|
||||
field = item.Elem().Field(fieldIndex)
|
||||
} else {
|
||||
field = item.Field(fieldIndex)
|
||||
}
|
||||
|
||||
if field.CanSet() {
|
||||
scanTargets[i] = field.Addr().Interface()
|
||||
} else {
|
||||
// 字段不可设置,使用占位符
|
||||
var dummy interface{}
|
||||
scanTargets[i] = &dummy
|
||||
}
|
||||
} else {
|
||||
// 没有找到对应字段,使用占位符
|
||||
var dummy interface{}
|
||||
scanTargets[i] = &dummy
|
||||
}
|
||||
}
|
||||
|
||||
// 执行扫描
|
||||
if err := rows.Scan(scanTargets...); err != nil {
|
||||
return fmt.Errorf("扫描数据失败:%w", err)
|
||||
}
|
||||
|
||||
// 处理时间字段格式化(目前保持原始 time.Time 值,由 JSON 序列化时格式化)
|
||||
// Go 的 database/sql 会自动将数据库时间扫描到 time.Time 类型
|
||||
// 在 JSON 序列化时,model.Time 的 MarshalJSON 会格式化为指定格式
|
||||
|
||||
// 添加到 Slice
|
||||
if isPtr {
|
||||
elem.Set(reflect.Append(elem, item))
|
||||
} else {
|
||||
elem.Set(reflect.Append(elem, item))
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// MapToStruct 将查询结果映射到单个 Struct
|
||||
func (rsm *ResultSetMapper) MapToStruct(rows *sql.Rows, result interface{}) error {
|
||||
// 使用 FieldMapper 的实现
|
||||
return rsm.fieldMapper.ColumnsToStruct(rows, result)
|
||||
}
|
||||
|
||||
// ScanAll 通用扫描方法,自动识别 Slice 或 Struct
|
||||
func (rsm *ResultSetMapper) ScanAll(rows *sql.Rows, result interface{}) error {
|
||||
val := reflect.ValueOf(result)
|
||||
if val.Kind() != reflect.Ptr {
|
||||
return fmt.Errorf("result 必须是指针类型")
|
||||
}
|
||||
|
||||
elem := val.Elem()
|
||||
|
||||
// 判断是 Slice 还是 Struct
|
||||
switch elem.Kind() {
|
||||
case reflect.Slice:
|
||||
return rsm.MapToSlice(rows, result)
|
||||
case reflect.Struct:
|
||||
return rsm.MapToStruct(rows, result)
|
||||
default:
|
||||
return fmt.Errorf("不支持的目标类型:%s", elem.Kind())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,44 @@
|
||||
package core
|
||||
|
||||
import (
|
||||
"time"
|
||||
)
|
||||
|
||||
// SoftDelete 软删除模型 - 嵌入到需要软删除的模型中
|
||||
type SoftDelete struct {
|
||||
DeletedAt *time.Time `json:"deleted_at" db:"deleted_at"` // 删除时间(为空表示未删除)
|
||||
}
|
||||
|
||||
// IsDeleted 检查是否已删除
|
||||
func (sd *SoftDelete) IsDeleted() bool {
|
||||
return sd.DeletedAt != nil
|
||||
}
|
||||
|
||||
// Delete 标记为已删除
|
||||
func (sd *SoftDelete) Delete() {
|
||||
now := time.Now()
|
||||
sd.DeletedAt = &now
|
||||
}
|
||||
|
||||
// Restore 恢复(取消删除)
|
||||
func (sd *SoftDelete) Restore() {
|
||||
sd.DeletedAt = nil
|
||||
}
|
||||
|
||||
// ISoftDeleter 软删除接口 - 定义软删除相关方法
|
||||
type ISoftDeleter interface {
|
||||
IsDeleted() bool
|
||||
Delete()
|
||||
Restore()
|
||||
}
|
||||
|
||||
// applySoftDelete 在查询中应用软删除过滤
|
||||
func applySoftDelete(q IQuery, unscoped bool) IQuery {
|
||||
if unscoped {
|
||||
// 忽略软删除,包含已删除的记录
|
||||
return q
|
||||
}
|
||||
|
||||
// 默认只查询未删除的记录
|
||||
return q.Where("deleted_at IS NULL")
|
||||
}
|
||||
@@ -0,0 +1,442 @@
|
||||
package core
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Transaction 事务实现 - ITx 接口的具体实现
|
||||
type Transaction struct {
|
||||
db *Database // 数据库连接
|
||||
tx *sql.Tx // 底层事务对象
|
||||
debug bool // 调试模式开关
|
||||
}
|
||||
|
||||
// 同步池优化 - 复用 slice 减少内存分配
|
||||
var insertArgsPool = sync.Pool{
|
||||
New: func() interface{} {
|
||||
return make([]interface{}, 0, 20)
|
||||
},
|
||||
}
|
||||
|
||||
var colNamesPool = sync.Pool{
|
||||
New: func() interface{} {
|
||||
return make([]string, 0, 20)
|
||||
},
|
||||
}
|
||||
|
||||
// Begin 开始一个新事务
|
||||
func (d *Database) Begin() (ITx, error) {
|
||||
if d.db == nil {
|
||||
return nil, fmt.Errorf("数据库连接未初始化")
|
||||
}
|
||||
|
||||
tx, err := d.db.Begin()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("开启事务失败:%w", err)
|
||||
}
|
||||
|
||||
return &Transaction{
|
||||
db: d,
|
||||
tx: tx,
|
||||
debug: d.debug,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Transaction 执行事务 - 自动管理事务的提交和回滚
|
||||
func (d *Database) Transaction(fn func(ITx) error) error {
|
||||
// 开启事务
|
||||
tx, err := d.Begin()
|
||||
if err != nil {
|
||||
return fmt.Errorf("开启事务失败:%w", err)
|
||||
}
|
||||
|
||||
defer func() {
|
||||
// 如果有 panic,回滚事务
|
||||
if r := recover(); r != nil {
|
||||
if rollbackErr := tx.Rollback(); rollbackErr != nil {
|
||||
fmt.Printf("[Magic-ORM] 事务回滚失败:%v\n", rollbackErr)
|
||||
}
|
||||
panic(r)
|
||||
}
|
||||
}()
|
||||
|
||||
// 执行用户提供的函数
|
||||
if err := fn(tx); err != nil {
|
||||
// 如果出错,回滚事务
|
||||
if rollbackErr := tx.Rollback(); rollbackErr != nil {
|
||||
return fmt.Errorf("事务执行失败且回滚也失败:%v, %w", rollbackErr, err)
|
||||
}
|
||||
return fmt.Errorf("事务执行失败:%w", err)
|
||||
}
|
||||
|
||||
// 提交事务
|
||||
if err := tx.Commit(); err != nil {
|
||||
return fmt.Errorf("事务提交失败:%w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Commit 提交事务
|
||||
func (t *Transaction) Commit() error {
|
||||
if t.tx == nil {
|
||||
return fmt.Errorf("事务对象为空")
|
||||
}
|
||||
return t.tx.Commit()
|
||||
}
|
||||
|
||||
// Rollback 回滚事务
|
||||
func (t *Transaction) Rollback() error {
|
||||
if t.tx == nil {
|
||||
return fmt.Errorf("事务对象为空")
|
||||
}
|
||||
return t.tx.Rollback()
|
||||
}
|
||||
|
||||
// Model 在事务中基于模型创建查询
|
||||
func (t *Transaction) Model(model interface{}) IQuery {
|
||||
return &QueryBuilder{
|
||||
db: t.db,
|
||||
model: model,
|
||||
tx: t.tx, // 使用事务对象
|
||||
debug: t.debug,
|
||||
}
|
||||
}
|
||||
|
||||
// Table 在事务中基于表名创建查询
|
||||
func (t *Transaction) Table(name string) IQuery {
|
||||
return &QueryBuilder{
|
||||
db: t.db,
|
||||
table: name,
|
||||
tx: t.tx, // 使用事务对象
|
||||
debug: t.debug,
|
||||
}
|
||||
}
|
||||
|
||||
// Insert 插入数据到数据库
|
||||
func (t *Transaction) Insert(model interface{}) (int64, error) {
|
||||
// 获取字段映射器
|
||||
mapper := NewFieldMapper()
|
||||
|
||||
// 获取表名和字段信息
|
||||
tableName := mapper.GetTableName(model)
|
||||
columns, err := mapper.StructToColumns(model)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("获取字段信息失败:%w", err)
|
||||
}
|
||||
|
||||
if len(columns) == 0 {
|
||||
return 0, fmt.Errorf("没有有效的字段")
|
||||
}
|
||||
|
||||
// 获取时间配置
|
||||
timeConfig := t.db.timeConfig
|
||||
if timeConfig == nil {
|
||||
timeConfig = DefaultTimeConfig()
|
||||
}
|
||||
|
||||
// 自动处理时间字段(使用配置的字段名)
|
||||
now := time.Now()
|
||||
for col, val := range columns {
|
||||
// 检查是否是配置的时间字段
|
||||
if col == timeConfig.GetCreatedAt() || col == timeConfig.GetUpdatedAt() || col == timeConfig.GetDeletedAt() {
|
||||
// 如果是零值时间,自动设置为当前时间
|
||||
if t.isZeroTimeValue(val) {
|
||||
columns[col] = now
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 生成 INSERT SQL
|
||||
var sqlBuilder strings.Builder
|
||||
sqlBuilder.Grow(128) // 预分配内存
|
||||
sqlBuilder.WriteString(fmt.Sprintf("INSERT INTO %s (", tableName))
|
||||
|
||||
// 列名 - 使用预分配内存
|
||||
colNames := colNamesPool.Get().([]string)
|
||||
colNames = colNames[:0] // 重置长度但不释放内存
|
||||
placeholders := make([]string, 0, len(columns))
|
||||
args := insertArgsPool.Get().([]interface{})
|
||||
args = args[:0] // 重置长度但不释放内存
|
||||
defer func() {
|
||||
colNamesPool.Put(colNames)
|
||||
insertArgsPool.Put(args)
|
||||
}()
|
||||
|
||||
for col, val := range columns {
|
||||
colNames = append(colNames, col)
|
||||
placeholders = append(placeholders, "?")
|
||||
args = append(args, val)
|
||||
}
|
||||
|
||||
sqlBuilder.WriteString(strings.Join(colNames, ", "))
|
||||
sqlBuilder.WriteString(") VALUES (")
|
||||
sqlBuilder.WriteString(strings.Join(placeholders, ", "))
|
||||
sqlBuilder.WriteString(")")
|
||||
|
||||
sqlStr := sqlBuilder.String()
|
||||
|
||||
// 调试模式
|
||||
if t.debug {
|
||||
fmt.Printf("[Magic-ORM] TX INSERT SQL: %s\n[Magic-ORM] Args: %v\n", sqlStr, args)
|
||||
}
|
||||
|
||||
// 执行插入
|
||||
result, err := t.tx.Exec(sqlStr, args...)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("插入失败:%w", err)
|
||||
}
|
||||
|
||||
// 获取插入的 ID
|
||||
id, err := result.LastInsertId()
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("获取插入 ID 失败:%w", err)
|
||||
}
|
||||
|
||||
return id, nil
|
||||
}
|
||||
|
||||
// BatchInsert 批量插入数据
|
||||
func (t *Transaction) BatchInsert(models interface{}, batchSize int) error {
|
||||
// 使用反射获取 Slice 数据
|
||||
modelsVal := reflect.ValueOf(models)
|
||||
if modelsVal.Kind() != reflect.Ptr || modelsVal.Elem().Kind() != reflect.Slice {
|
||||
return fmt.Errorf("models 必须是指向 Slice 的指针")
|
||||
}
|
||||
|
||||
sliceVal := modelsVal.Elem()
|
||||
length := sliceVal.Len()
|
||||
|
||||
if length == 0 {
|
||||
return nil // 空 Slice,无需插入
|
||||
}
|
||||
|
||||
// 分批处理
|
||||
for i := 0; i < length; i += batchSize {
|
||||
end := i + batchSize
|
||||
if end > length {
|
||||
end = length
|
||||
}
|
||||
|
||||
// 处理当前批次
|
||||
for j := i; j < end; j++ {
|
||||
model := sliceVal.Index(j).Interface()
|
||||
_, err := t.Insert(model)
|
||||
if err != nil {
|
||||
return fmt.Errorf("批量插入第%d条记录失败:%w", j, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// isZeroTimeValue 检查是否是零值时间
|
||||
func (t *Transaction) isZeroTimeValue(val interface{}) bool {
|
||||
if val == nil {
|
||||
return true
|
||||
}
|
||||
|
||||
// 检查是否是 time.Time 类型
|
||||
if tm, ok := val.(time.Time); ok {
|
||||
return tm.IsZero() || tm.UnixNano() == 0
|
||||
}
|
||||
|
||||
// 使用反射检查
|
||||
v := reflect.ValueOf(val)
|
||||
switch v.Kind() {
|
||||
case reflect.Ptr:
|
||||
return v.IsNil()
|
||||
case reflect.Struct:
|
||||
// 如果是 time.Time 结构
|
||||
if tm, ok := v.Interface().(time.Time); ok {
|
||||
return tm.IsZero() || tm.UnixNano() == 0
|
||||
}
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
// Update 更新数据
|
||||
func (t *Transaction) Update(model interface{}, data map[string]interface{}) error {
|
||||
// 获取字段映射器
|
||||
mapper := NewFieldMapper()
|
||||
|
||||
// 获取表名和主键
|
||||
tableName := mapper.GetTableName(model)
|
||||
pk := mapper.GetPrimaryKey(model)
|
||||
|
||||
// 获取时间配置
|
||||
timeConfig := t.db.timeConfig
|
||||
if timeConfig == nil {
|
||||
timeConfig = DefaultTimeConfig()
|
||||
}
|
||||
|
||||
// 自动处理 updated_at 时间字段(使用配置的字段名)
|
||||
if data == nil {
|
||||
data = make(map[string]interface{})
|
||||
}
|
||||
data[timeConfig.GetUpdatedAt()] = time.Now()
|
||||
|
||||
// 过滤零值
|
||||
pf := NewParamFilter()
|
||||
data = pf.FilterZeroValues(data)
|
||||
|
||||
if len(data) == 0 {
|
||||
return fmt.Errorf("没有有效的更新字段")
|
||||
}
|
||||
|
||||
// 生成 UPDATE SQL
|
||||
var sqlBuilder strings.Builder
|
||||
sqlBuilder.Grow(128) // 预分配内存
|
||||
sqlBuilder.WriteString(fmt.Sprintf("UPDATE %s SET ", tableName))
|
||||
|
||||
setParts := make([]string, 0, len(data))
|
||||
args := insertArgsPool.Get().([]interface{})
|
||||
args = args[:0] // 重置长度但不释放内存
|
||||
defer func() {
|
||||
insertArgsPool.Put(args)
|
||||
}()
|
||||
|
||||
for col, val := range data {
|
||||
setParts = append(setParts, fmt.Sprintf("%s = ?", col))
|
||||
args = append(args, val)
|
||||
}
|
||||
|
||||
sqlBuilder.WriteString(strings.Join(setParts, ", "))
|
||||
sqlBuilder.WriteString(fmt.Sprintf(" WHERE %s = ?", pk))
|
||||
|
||||
// 获取主键值
|
||||
pkValue := reflect.ValueOf(model)
|
||||
if pkValue.Kind() == reflect.Ptr {
|
||||
pkValue = pkValue.Elem()
|
||||
}
|
||||
idField := pkValue.FieldByName("ID")
|
||||
if idField.IsValid() {
|
||||
args = append(args, idField.Interface())
|
||||
} else {
|
||||
return fmt.Errorf("模型缺少 ID 字段")
|
||||
}
|
||||
|
||||
sqlStr := sqlBuilder.String()
|
||||
|
||||
// 调试模式
|
||||
if t.debug {
|
||||
fmt.Printf("[Magic-ORM] TX UPDATE SQL: %s\n[Magic-ORM] Args: %v\n", sqlStr, args)
|
||||
}
|
||||
|
||||
// 执行更新
|
||||
_, err := t.tx.Exec(sqlStr, args...)
|
||||
if err != nil {
|
||||
return fmt.Errorf("更新失败:%w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Delete 删除数据(支持软删除)
|
||||
func (t *Transaction) Delete(model interface{}) error {
|
||||
// 获取字段映射器
|
||||
mapper := NewFieldMapper()
|
||||
|
||||
// 获取表名和主键
|
||||
tableName := mapper.GetTableName(model)
|
||||
pk := mapper.GetPrimaryKey(model)
|
||||
|
||||
// 获取时间配置
|
||||
timeConfig := t.db.timeConfig
|
||||
if timeConfig == nil {
|
||||
timeConfig = DefaultTimeConfig()
|
||||
}
|
||||
|
||||
// 检查是否支持软删除(是否有配置的 deleted_at 字段)
|
||||
hasSoftDelete := false
|
||||
pkValue := reflect.ValueOf(model)
|
||||
if pkValue.Kind() == reflect.Ptr {
|
||||
pkValue = pkValue.Elem()
|
||||
}
|
||||
|
||||
// 检查是否有 DeletedAt 字段(使用配置的字段名)
|
||||
deletedAtField := pkValue.FieldByNameFunc(func(fieldName string) bool {
|
||||
// 将字段名转换为数据库列名进行比较
|
||||
expectedCol := timeConfig.GetDeletedAt()
|
||||
// 简单转换:下划线转驼峰
|
||||
return fieldName == "DeletedAt" || fieldName == expectedCol
|
||||
})
|
||||
|
||||
if deletedAtField.IsValid() {
|
||||
hasSoftDelete = true
|
||||
}
|
||||
|
||||
var sqlStr string
|
||||
args := make([]interface{}, 0)
|
||||
|
||||
if hasSoftDelete {
|
||||
// 软删除:更新 deleted_at 为当前时间(使用配置的字段名)
|
||||
sqlStr = fmt.Sprintf("UPDATE %s SET %s = ? WHERE %s = ?", tableName, timeConfig.GetDeletedAt(), pk)
|
||||
args = append(args, time.Now())
|
||||
} else {
|
||||
// 硬删除:直接 DELETE
|
||||
sqlStr = fmt.Sprintf("DELETE FROM %s WHERE %s = ?", tableName, pk)
|
||||
}
|
||||
|
||||
// 获取主键值
|
||||
idField := pkValue.FieldByName("ID")
|
||||
if idField.IsValid() {
|
||||
args = append(args, idField.Interface())
|
||||
} else {
|
||||
return fmt.Errorf("模型缺少 ID 字段")
|
||||
}
|
||||
|
||||
// 调试模式
|
||||
if t.debug {
|
||||
deleteType := "硬删除"
|
||||
if hasSoftDelete {
|
||||
deleteType = "软删除"
|
||||
}
|
||||
fmt.Printf("[Magic-ORM] TX %s SQL: %s\n[Magic-ORM] Args: %v\n", deleteType, sqlStr, args)
|
||||
}
|
||||
|
||||
// 执行删除
|
||||
_, err := t.tx.Exec(sqlStr, args...)
|
||||
if err != nil {
|
||||
return fmt.Errorf("删除失败:%w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Query 在事务中执行原生 SQL 查询
|
||||
func (t *Transaction) Query(result interface{}, query string, args ...interface{}) error {
|
||||
if t.debug {
|
||||
fmt.Printf("[Magic-ORM] TX Query SQL: %s\n[Magic-ORM] Args: %v\n", query, args)
|
||||
}
|
||||
|
||||
rows, err := t.tx.Query(query, args...)
|
||||
if err != nil {
|
||||
return fmt.Errorf("事务查询失败:%w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
// TODO: 实现结果映射
|
||||
return nil
|
||||
}
|
||||
|
||||
// Exec 在事务中执行原生 SQL
|
||||
func (t *Transaction) Exec(query string, args ...interface{}) (sql.Result, error) {
|
||||
if t.debug {
|
||||
fmt.Printf("[Magic-ORM] TX Exec SQL: %s\n[Magic-ORM] Args: %v\n", query, args)
|
||||
}
|
||||
|
||||
result, err := t.tx.Exec(query, args...)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("事务执行失败:%w", err)
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
Reference in New Issue
Block a user