2025-12-16 12:05:05 +09:00

210 lines
5.0 KiB
Go

package umsql
import (
"database/sql"
"errors"
"fmt"
"strconv"
"strings"
"time"
"AngkorWalletScanning/model"
"AngkorWalletScanning/umlog"
_ "github.com/go-sql-driver/mysql"
)
// type DB_handler struct {
// conn *sql.DB
// }
func DbConnect(host string, port string, user string, password string, dbname string) *sql.DB {
conn, connect_err := sql.Open("mysql", user+":"+password+"@tcp("+host+":"+port+")/"+dbname)
if connect_err != nil {
umlog.Error(connect_err.Error())
panic(connect_err.Error())
}
conn.SetConnMaxLifetime(time.Minute * 3)
conn.SetMaxOpenConns(1000)
conn.SetMaxIdleConns(1000)
return conn
}
func DbDisconnect(connA *sql.DB) {
umlog.Debug("DataBase DISCONNECTION")
err := connA.Close()
if err != nil {
umlog.Warn("DataBase DISCONNECTION Error A: " + err.Error())
}
}
func DbBegin(conn *sql.DB) (*sql.Tx, error) {
if conn != nil {
return conn.Begin()
} else {
return nil, errors.New("db connecter is nil")
}
}
func DbRollback(tx *sql.Tx) error {
return tx.Rollback()
}
func DbCommit(tx *sql.Tx) error {
return tx.Commit()
}
func SqlSelect(conn *sql.DB, query string, args ...any) (*sql.Rows, []string, error) {
var rows *sql.Rows
var err error
if conn != nil {
if len(args) == 0 {
rows, err = conn.Query(query)
} else {
rows, err = conn.Query(query, args...)
}
if err != nil {
umlog.Error(err.Error())
return nil, nil, err
}
cols, err := rows.Columns()
if err != nil {
umlog.Error(err.Error())
return nil, nil, err
}
return rows, cols, nil
}
umlog.Error("db connecter is nil")
return nil, nil, errors.New("db connecter is nil")
}
func SqlTxExecute(tx *sql.Tx, query string, args ...any) (int64, error) {
var result sql.Result
var err error
if tx != nil {
if len(args) == 0 {
result, err = tx.Exec(query)
} else {
result, err = tx.Exec(query, args...)
}
if err != nil {
umlog.Error(err.Error())
return -1, err
}
n, err := result.RowsAffected()
if err != nil {
umlog.Error(err.Error())
return -1, err
}
return n, nil
} else {
umlog.Error("tx is nil")
return -1, errors.New("tx is nil")
}
}
func SqlExecute(conn *sql.DB, query string, args ...any) (int64, error) {
var result sql.Result
var err error
if conn != nil {
if len(args) == 0 {
result, err = conn.Exec(query)
} else {
result, err = conn.Exec(query, args...)
}
if err != nil {
umlog.Error(err.Error())
return -1, err
}
n, err := result.RowsAffected()
if err != nil {
umlog.Error(err.Error())
return -1, err
}
return n, nil
}
umlog.Error("db connecter is nil")
return -1, errors.New("db connecter is nil")
}
func SqlStatement(conn *sql.DB, sql string, param [][]string) error {
stmt, err := conn.Prepare(sql)
if err != nil {
umlog.Error(err.Error())
return err
}
defer stmt.Close()
for i, sub_val := range param {
fmt.Println(sub_val)
err = sqlStatementExec(stmt, param[i], len(param[i]))
if err != nil {
umlog.Error(err.Error())
return err
}
}
return nil
}
func sqlStatementExec(stmt *sql.Stmt, param []string, param_count int) error {
args := make([]interface{}, param_count)
for i, v := range param {
args[i] = v
}
_, err := stmt.Exec(args...)
return err
}
func SqlTransaction(conn *sql.DB, maxRetries int, txFunc func(*sql.Tx) error) model.DefaultErrorModel {
var err error
var value string
for i := 0; i <= maxRetries; i++ {
tx, beginErr := conn.Begin()
if beginErr != nil {
return model.DefaultErrorModel{Code: 4001201, Message: fmt.Sprintf("failed to begin transaction: %v", beginErr)}
}
err = txFunc(tx)
if err != nil {
// 트랜잭션 롤백
_ = tx.Rollback()
// 1205 오류인지 확인
if isLockWaitTimeout(err) {
umlog.Info(value + ": Transaction lock timeout (1205), retrying... (" + strconv.Itoa(i+1) + "/" + strconv.Itoa(maxRetries) + ")")
time.Sleep(time.Duration(100*(i+1)) * time.Millisecond) // simple backoff
continue
}
return model.DefaultErrorModel{Code: 4001201, Message: fmt.Sprintf("transaction failed: %v", err)}
}
// 커밋 시도
commitErr := tx.Commit()
if commitErr != nil {
// 커밋 중 에러가 1205일 수도 있음
if isLockWaitTimeout(commitErr) {
umlog.Info(value + ": Commit failed due to lock timeout, retrying... (" + strconv.Itoa(i+1) + "/" + strconv.Itoa(maxRetries) + ")")
time.Sleep(time.Duration(100*(i+1)) * time.Millisecond)
continue
}
umlog.Error(value + ": commit error: " + commitErr.Error())
return model.DefaultErrorModel{Code: 4001201, Message: fmt.Sprintf("commit error: %v", commitErr)}
}
// 성공
return model.DefaultErrorModel{Code: 200, Message: ""}
}
umlog.Error(value + ": Transaction failed after " + strconv.Itoa(maxRetries) + " retries: " + err.Error())
return model.DefaultErrorModel{Code: 4001201, Message: fmt.Sprintf("%s: Transaction failed after %d retries: %w", value, maxRetries, err)}
}
func isLockWaitTimeout(err error) bool {
// mysql 드라이버의 오류 메시지에서 1205 감지
return err != nil && (strings.Contains(err.Error(), "Error 1205") || strings.Contains(err.Error(), "Lock wait timeout"))
}