package db import ( "bytes" "code.hoteas.com/golang/hotime/cache" . "code.hoteas.com/golang/hotime/common" "code.hoteas.com/golang/hotime/log" "database/sql" "strings" "sync" "sync/atomic" _ "gitee.com/chunanyong/dm" _ "github.com/go-sql-driver/mysql" _ "github.com/mattn/go-sqlite3" ) // DBError 最后一次 SQL 操作的结构化快照 // 嵌入框架 Error 类型,与 Obj 等保持一致的错误 API(HasError/GetError/Unwrap) type DBError struct { Error Query string Data []interface{} } // HoTimeDB 数据库操作核心结构体 type HoTimeDB struct { *sql.DB ContextBase DBName string *cache.HoTimeCache Log *log.Logger Type string // 数据库类型: mysql, sqlite3, postgres Prefix string ConnectFunc func() (*sql.DB, *sql.DB) lastError atomic.Pointer[DBError] // 无锁读写,始终记录最后一次 SQL 操作 limit Slice *sql.Tx //事务对象 SlaveDB *sql.DB // 从数据库 limitMu sync.Mutex Dialect Dialect // 数据库方言适配器 testTx *sql.Tx // 测试事务:设置后所有操作都在此事务内,测试结束回滚 testMu *sync.Mutex // 保护 testTx 单连接不被并发访问 testLogBuf *bytes.Buffer // 测试模式:SQL 日志写入此 buffer 而非直接打印,失败时才输出 testLogBufMu sync.Mutex // 保护 testLogBuf 并发写入 testSQLErrCount int32 // 测试模式下 SQL 错误计数(atomic) txFailed *bool // 事务内是否有SQL出错,出错则事务必须回滚(指针确保按值传递 HoTimeDB 时状态共享) } // GetLast 返回最后一次 SQL 操作的完整快照(Query + Data + Error) func (that *HoTimeDB) GetLast() *DBError { return that.lastError.Load() } // GetLastError 返回最后一次 SQL 的错误(nil = 成功) func (that *HoTimeDB) GetLastError() error { if e := that.lastError.Load(); e != nil { return e.GetError() } return nil } // GetLastQuery 返回最后一次执行的 SQL 语句 func (that *HoTimeDB) GetLastQuery() string { if e := that.lastError.Load(); e != nil { return e.Query } return "" } // GetLastData 返回最后一次 SQL 的参数 func (that *HoTimeDB) GetLastData() []interface{} { if e := that.lastError.Load(); e != nil { return e.Data } return nil } // SetConnect 设置数据库配置连接 func (that *HoTimeDB) SetConnect(connect func() (master, slave *sql.DB)) { that.ConnectFunc = connect that.InitDb() } // InitDb 初始化数据库连接 func (that *HoTimeDB) InitDb() { that.DB, that.SlaveDB = that.ConnectFunc() if that.DB == nil { return } e := that.DB.Ping() if e != nil { dbErr := &DBError{} dbErr.SetError(e) that.lastError.Store(dbErr) if that.Log != nil { that.Log.Error().Err(e).Msg("数据库连接失败") } } if that.SlaveDB != nil { e := that.SlaveDB.Ping() if e != nil { dbErr := &DBError{} dbErr.SetError(e) that.lastError.Store(dbErr) if that.Log != nil { that.Log.Error().Err(e).Msg("从数据库连接失败") } } } if that.Dialect == nil { that.initDialect() } } // initDialect 根据数据库类型初始化方言 func (that *HoTimeDB) initDialect() { switch that.Type { case "postgres", "postgresql": that.Dialect = &PostgreSQLDialect{} case "sqlite3", "sqlite": that.Dialect = &SQLiteDialect{} case "dm", "dameng": that.Dialect = &DMDialect{} default: that.Dialect = &MySQLDialect{} } } // GetDialect 获取当前方言适配器 func (that *HoTimeDB) GetDialect() Dialect { if that.Dialect == nil { that.initDialect() } return that.Dialect } // SetDialect 设置方言适配器 func (that *HoTimeDB) SetDialect(dialect Dialect) { that.Dialect = dialect } // GetType 获取数据库类型 func (that *HoTimeDB) GetType() string { return that.Type } // GetPrefix 获取表前缀 func (that *HoTimeDB) GetPrefix() string { return that.Prefix } // GetProcessor 获取标识符处理器 // 用于处理表名、字段名的前缀添加和引号转换 func (that *HoTimeDB) GetProcessor() *IdentifierProcessor { return NewIdentifierProcessor(that.GetDialect(), that.Prefix) } // T 辅助方法:获取带前缀和引号的表名 // 用于手动构建 SQL 时使用 // 示例: db.T("order") 返回 "`app_order`" (MySQL) 或 "\"app_order\"" (PostgreSQL) func (that *HoTimeDB) T(table string) string { return that.GetProcessor().ProcessTableName(table) } // C 辅助方法:获取带前缀和引号的 table.column // 支持两种调用方式: // - db.C("order", "name") 返回 "`app_order`.`name`" // - db.C("order.name") 返回 "`app_order`.`name`" func (that *HoTimeDB) C(args ...string) string { if len(args) == 0 { return "" } if len(args) == 1 { return that.GetProcessor().ProcessColumn(args[0]) } // 两个参数: table, column dialect := that.GetDialect() table := args[0] column := args[1] // 去除已有引号 table = trimQuotes(table) column = trimQuotes(column) return dialect.QuoteIdentifier(that.Prefix+table) + "." + dialect.QuoteIdentifier(column) } // trimQuotes 去除字符串两端的引号 func trimQuotes(s string) string { s = strings.TrimSpace(s) if len(s) >= 2 { if (s[0] == '`' && s[len(s)-1] == '`') || (s[0] == '"' && s[len(s)-1] == '"') { return s[1 : len(s)-1] } } return s } // SetTestLogBuffer 设置测试 SQL 日志缓冲区,非 nil 时 SQL 日志写入 buffer 而非直接打印 func (that *HoTimeDB) SetTestLogBuffer(buf *bytes.Buffer) { that.testLogBufMu.Lock() that.testLogBuf = buf that.testLogBufMu.Unlock() } // FlushTestLog 取出并清空缓冲区中的 SQL 日志,返回累积的日志文本 func (that *HoTimeDB) FlushTestLog() string { that.testLogBufMu.Lock() defer that.testLogBufMu.Unlock() if that.testLogBuf == nil { return "" } s := that.testLogBuf.String() that.testLogBuf.Reset() return s } // TestSQLErrorCount 返回测试模式下的 SQL 错误计数 func (that *HoTimeDB) TestSQLErrorCount() int32 { return atomic.LoadInt32(&that.testSQLErrCount) } // ResetTestSQLErrorCount 重置测试 SQL 错误计数 func (that *HoTimeDB) ResetTestSQLErrorCount() { atomic.StoreInt32(&that.testSQLErrCount, 0) } // BeginTestTx 开启测试事务,设置后所有 Query/Exec/Action 都在此事务内执行 func (that *HoTimeDB) BeginTestTx() error { tx, err := that.DB.Begin() if err != nil { return err } that.testTx = tx that.testMu = &sync.Mutex{} return nil } // RollbackTestTx 回滚测试事务,清除 testTx 状态 func (that *HoTimeDB) RollbackTestTx() error { if that.testTx != nil { err := that.testTx.Rollback() that.testTx = nil return err } return nil }