u
This commit is contained in:
1 parent
33113f279c
commit
c2c5cc92af
16 files changed
+864
-352
No files matched your search
@@ -69,7 +69,7 @@ func loginDefault(c *router.Context) error {
|
||||
"account", account,
|
||||
"password", password,
|
||||
)
|
||||
user, err := db.FindOne("user_info", userQuery)
|
||||
user, err := db.SelectOne("user_info", db.BuildWhereString(userQuery), "")
|
||||
if err != nil {
|
||||
return response.Success(c).Message("查询用户失败,请稍后重试").Send()
|
||||
}
|
||||
@@ -127,7 +127,7 @@ func registerDefault(c *router.Context) error {
|
||||
// 构造查询条件,检查账号是否已存在
|
||||
accountQuery := mapx.New().SetKV("account", account)
|
||||
|
||||
existingUsers, err := db.FindOne("user_info", accountQuery)
|
||||
existingUsers, err := db.SelectOne("user_info", db.BuildWhereString(accountQuery), "")
|
||||
if err != nil {
|
||||
return c.AddError(err)
|
||||
}
|
||||
|
||||
@@ -8,10 +8,12 @@ import (
|
||||
"base-go-v2/internal/utils/mapx"
|
||||
"fmt"
|
||||
"github.com/shopspring/decimal"
|
||||
"time"
|
||||
)
|
||||
|
||||
func InitMoneyRouter(r *router.Router) {
|
||||
group := r.Group("money")
|
||||
group.POST("/query", queryMoneyRecords)
|
||||
group.GET("/category", getUserCategory)
|
||||
group.POST("/insert", insertMoneyRecord)
|
||||
group.POST("/update", updateMoneyRecord)
|
||||
@@ -20,10 +22,8 @@ func InitMoneyRouter(r *router.Router) {
|
||||
|
||||
// InitDefaultCategories 初始化系统默认分类到用户分类表
|
||||
func InitDefaultCategories(userID int64) error {
|
||||
fmt.Println("进入 InitDefaultCategories ===")
|
||||
|
||||
// 1. 查询系统分类表
|
||||
sysCategories, err := db.FindAll("money_sys_category", "sort_number")
|
||||
sysCategories, err := db.SelectList("money_sys_category", "", "sort_number")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -31,7 +31,6 @@ func InitDefaultCategories(userID int64) error {
|
||||
fmt.Printf("系统分类数量: %d\n", len(sysCategories))
|
||||
|
||||
// 2. 构建批量插入数据
|
||||
//userCategories := make([]mapx.M, 0, len(sysCategories))
|
||||
userCategories := mapx.NewArr(len(sysCategories))
|
||||
|
||||
for i, row := range sysCategories {
|
||||
@@ -52,15 +51,124 @@ func InitDefaultCategories(userID int64) error {
|
||||
return err
|
||||
}
|
||||
|
||||
fmt.Println("初始化用户分类成功!")
|
||||
return nil
|
||||
}
|
||||
|
||||
// 根据
|
||||
func queryMoneyRecords(c *router.Context) error {
|
||||
bodyData, err := c.GetBodyWithRequired("record_time") // "YYYY-MM"
|
||||
if err != nil {
|
||||
return c.AddError(err)
|
||||
}
|
||||
|
||||
recordMonth := bodyData.GetString("record_time")
|
||||
|
||||
// 解析年月
|
||||
startTime, err := time.Parse("2006-01", recordMonth)
|
||||
if err != nil {
|
||||
return c.AddError(fmt.Errorf("invalid record_time format: %v", err))
|
||||
}
|
||||
|
||||
// 获取当月最后一天
|
||||
endTime := startTime.AddDate(0, 1, 0).Add(-time.Nanosecond)
|
||||
|
||||
// SQL: money_record LEFT JOIN money_user_category
|
||||
sqlStr := fmt.Sprintf(`
|
||||
SELECT
|
||||
mr.*,
|
||||
muc.name AS category_name,
|
||||
muc.icon AS category_icon,
|
||||
muc.type AS category_type
|
||||
FROM money_record mr
|
||||
LEFT JOIN money_user_category muc ON mr.category_id = muc.id
|
||||
WHERE mr.create_time BETWEEN '%s' AND '%s'
|
||||
ORDER BY mr.create_time DESC
|
||||
`, startTime.Format("2006-01-02 15:04:05"), endTime.Format("2006-01-02 15:04:05"))
|
||||
|
||||
// 查询
|
||||
list, err := db.SelectBySql(sqlStr)
|
||||
if err != nil {
|
||||
return c.AddError(err)
|
||||
}
|
||||
|
||||
// 分组数据
|
||||
grouped := mapx.New() // mapx.M
|
||||
totalIncome := decimal.NewFromInt(0)
|
||||
totalExpense := decimal.NewFromInt(0)
|
||||
|
||||
for _, r := range list {
|
||||
row := mapx.M(r)
|
||||
|
||||
createTime, ok := row["create_time"].(time.Time)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
day := createTime.Format("2006-01-02")
|
||||
|
||||
// 获取当天数据
|
||||
dayDataI := grouped.Get(day)
|
||||
var dayData mapx.M
|
||||
if dayDataI == nil {
|
||||
dayData = mapx.New().
|
||||
Set("records", mapx.NewArr()).
|
||||
Set("income", decimal.NewFromInt(0).String()).
|
||||
Set("expense", decimal.NewFromInt(0).String())
|
||||
} else {
|
||||
dayData = dayDataI.(mapx.M)
|
||||
}
|
||||
|
||||
// 追加记录
|
||||
records := dayData.GetArray("records")
|
||||
records = append(records, row)
|
||||
dayData.Set("records", records)
|
||||
|
||||
// 计算当天收入支出
|
||||
categoryType := row.GetInt("category_type")
|
||||
amount, err := decimal.NewFromString(row.GetString("amount"))
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
|
||||
dayIncome, err := decimal.NewFromString(dayData.GetString("income"))
|
||||
if err != nil {
|
||||
dayIncome = decimal.NewFromInt(0)
|
||||
}
|
||||
|
||||
dayExpense, err := decimal.NewFromString(dayData.GetString("expense"))
|
||||
if err != nil {
|
||||
dayExpense = decimal.NewFromInt(0)
|
||||
}
|
||||
|
||||
if categoryType == 0 {
|
||||
dayIncome = dayIncome.Add(amount)
|
||||
totalIncome = totalIncome.Add(amount)
|
||||
} else if categoryType == 1 {
|
||||
dayExpense = dayExpense.Add(amount)
|
||||
totalExpense = totalExpense.Add(amount)
|
||||
}
|
||||
|
||||
dayData.Set("income", dayIncome.String())
|
||||
dayData.Set("expense", dayExpense.String())
|
||||
|
||||
grouped.Set(day, dayData)
|
||||
}
|
||||
|
||||
result := mapx.New().
|
||||
Set("records", grouped).
|
||||
Set("total_income", totalIncome.String()).
|
||||
Set("total_expense", totalExpense.String())
|
||||
|
||||
return response.Success(c).Data(result).Send()
|
||||
}
|
||||
|
||||
func getUserCategory(c *router.Context) error {
|
||||
userId := c.GetUserID()
|
||||
|
||||
// 查询当前用户的分类
|
||||
userCategories, err := db.Find("money_user_category", mapx.New().Set("user_id", userId), "sort_number ASC")
|
||||
userCategories, err := db.SelectList("money_user_category",
|
||||
db.BuildWhereString(mapx.New().Set("user_id", userId)),
|
||||
"sort_number ASC")
|
||||
|
||||
if err != nil {
|
||||
return c.AddError(err)
|
||||
}
|
||||
@@ -126,7 +234,7 @@ func updateMoneyRecord(c *router.Context) error {
|
||||
|
||||
recordId := bodyData.GetInt64("id")
|
||||
|
||||
record, err := db.GetOne("money_record", "id", recordId)
|
||||
record, err := db.SelectById("money_record", "id", recordId)
|
||||
if err != nil {
|
||||
return c.AddError(err)
|
||||
}
|
||||
|
||||
@@ -24,9 +24,11 @@ var defaultPoolConfig = struct {
|
||||
func InitDb() error {
|
||||
d := config.App.Db
|
||||
|
||||
// 构造 DSN
|
||||
dsn := fmt.Sprintf("host=%s port=%d user=%s password=%s dbname=%s sslmode=disable",
|
||||
d.Host, d.Port, d.User, d.Password, d.Dbname)
|
||||
// 构造 DSN,指定时区为 Asia/Shanghai
|
||||
dsn := fmt.Sprintf(
|
||||
"host=%s port=%d user=%s password=%s dbname=%s sslmode=disable TimeZone=Asia/Shanghai",
|
||||
d.Host, d.Port, d.User, d.Password, d.Dbname,
|
||||
)
|
||||
|
||||
db, err := sqlx.Connect("postgres", dsn)
|
||||
if err != nil {
|
||||
|
||||
@@ -8,7 +8,6 @@ import (
|
||||
|
||||
// ---------------- 通用查询方法 ----------------
|
||||
|
||||
// 执行查询,返回多条记录
|
||||
func queryMaps(query string, args []interface{}) ([]mapx.M, error) {
|
||||
rows, err := DB.Queryx(query, args...)
|
||||
if err != nil {
|
||||
@@ -22,13 +21,20 @@ func queryMaps(query string, args []interface{}) ([]mapx.M, error) {
|
||||
if err := rows.MapScan(row); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 统一处理 []byte 类型,直接转字符串,保留原始小数点
|
||||
for k, v := range row {
|
||||
if b, ok := v.([]byte); ok {
|
||||
row[k] = string(b)
|
||||
}
|
||||
}
|
||||
|
||||
result = append(result, row)
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// 执行查询,返回单条记录
|
||||
func queryMap(query string, args []interface{}) (mapx.M, error) {
|
||||
rows, err := DB.Queryx(query, args...)
|
||||
if err != nil {
|
||||
@@ -41,93 +47,73 @@ func queryMap(query string, args []interface{}) (mapx.M, error) {
|
||||
if err := rows.MapScan(row); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 统一处理 []byte 类型,直接转字符串
|
||||
for k, v := range row {
|
||||
if b, ok := v.([]byte); ok {
|
||||
row[k] = string(b)
|
||||
}
|
||||
}
|
||||
|
||||
return row, nil
|
||||
}
|
||||
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// 构建 WHERE 条件
|
||||
func buildWhere(conditions map[string]interface{}) (string, []interface{}) {
|
||||
var whereParts []string
|
||||
var values []interface{}
|
||||
i := 1
|
||||
for k, v := range conditions {
|
||||
whereParts = append(whereParts, fmt.Sprintf("%s=$%d", k, i))
|
||||
values = append(values, v)
|
||||
i++
|
||||
// BuildWhereString 将 map 条件转换为 SQL WHERE 字符串
|
||||
func BuildWhereString(conds mapx.M) string {
|
||||
if len(conds) == 0 {
|
||||
return ""
|
||||
}
|
||||
return strings.Join(whereParts, " AND "), values
|
||||
|
||||
var parts []string
|
||||
for k, v := range conds {
|
||||
parts = append(parts, fmt.Sprintf("%s='%v'", k, v))
|
||||
}
|
||||
|
||||
return strings.Join(parts, " AND ")
|
||||
}
|
||||
|
||||
// ---------------- 公共查询方法 ----------------
|
||||
// ---------------- 统一 select 系列方法 ----------------
|
||||
|
||||
// GetOne 根据主键查询单条记录
|
||||
func GetOne(table string, pkColumn string, pkValue interface{}) (mapx.M, error) {
|
||||
query := fmt.Sprintf("SELECT * FROM %s WHERE %s=$1", table, pkColumn)
|
||||
// SelectById 根据主键查询单条记录
|
||||
func SelectById(tableName, pkColumn string, pkValue interface{}) (mapx.M, error) {
|
||||
query := fmt.Sprintf("SELECT * FROM %s WHERE %s=$1", tableName, pkColumn)
|
||||
return queryMap(query, []interface{}{pkValue})
|
||||
}
|
||||
|
||||
// GetBatch 根据主键批量查询
|
||||
func GetBatch(table string, pkColumn string, pkValues []interface{}) ([]mapx.M, error) {
|
||||
if len(pkValues) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
placeholders := make([]string, len(pkValues))
|
||||
for i := range pkValues {
|
||||
placeholders[i] = fmt.Sprintf("$%d", i+1)
|
||||
}
|
||||
|
||||
query := fmt.Sprintf("SELECT * FROM %s WHERE %s IN (%s)", table, pkColumn, strings.Join(placeholders, ","))
|
||||
return queryMaps(query, pkValues)
|
||||
// SelectBySql 根据任意 SQL 查询多条记录
|
||||
func SelectBySql(sql string, args ...interface{}) ([]mapx.M, error) {
|
||||
return queryMaps(sql, args)
|
||||
}
|
||||
|
||||
// Find 根据条件查询多条记录,可选排序
|
||||
func Find(table string, conditions map[string]interface{}, orderBy ...string) ([]mapx.M, error) {
|
||||
if len(conditions) == 0 {
|
||||
return nil, fmt.Errorf("查询条件不能为空")
|
||||
// SelectList 查询列表,可传 where 和 order,where 可以为空
|
||||
func SelectList(tableName string, where string, order string) ([]mapx.M, error) {
|
||||
query := fmt.Sprintf("SELECT * FROM %s", tableName)
|
||||
if strings.TrimSpace(where) != "" {
|
||||
query += " WHERE " + where
|
||||
}
|
||||
|
||||
where, values := buildWhere(conditions)
|
||||
query := fmt.Sprintf("SELECT * FROM %s WHERE %s", table, where)
|
||||
|
||||
order := "id DESC"
|
||||
if len(orderBy) > 0 && orderBy[0] != "" {
|
||||
order = orderBy[0]
|
||||
}
|
||||
query += " ORDER BY " + order
|
||||
|
||||
return queryMaps(query, values)
|
||||
}
|
||||
|
||||
// FindAll 查询整个表,可选排序
|
||||
func FindAll(table string, orderBy ...string) ([]mapx.M, error) {
|
||||
query := fmt.Sprintf("SELECT * FROM %s", table)
|
||||
|
||||
order := "id DESC"
|
||||
if len(orderBy) > 0 && orderBy[0] != "" {
|
||||
order = orderBy[0]
|
||||
if strings.TrimSpace(order) == "" {
|
||||
order = "id DESC"
|
||||
}
|
||||
query += " ORDER BY " + order
|
||||
|
||||
return queryMaps(query, nil)
|
||||
}
|
||||
|
||||
// FindOne 根据条件查询单条记录,可选排序
|
||||
func FindOne(table string, conditions map[string]interface{}, orderBy ...string) (mapx.M, error) {
|
||||
if len(conditions) == 0 {
|
||||
return nil, fmt.Errorf("查询条件不能为空")
|
||||
// SelectOne 查询单条记录,可传 where 和 order(order 可为空)
|
||||
func SelectOne(tableName string, where string, order string) (mapx.M, error) {
|
||||
query := fmt.Sprintf("SELECT * FROM %s", tableName)
|
||||
if strings.TrimSpace(where) != "" {
|
||||
query += " WHERE " + where
|
||||
}
|
||||
|
||||
where, values := buildWhere(conditions)
|
||||
query := fmt.Sprintf("SELECT * FROM %s WHERE %s", table, where)
|
||||
|
||||
order := "id DESC"
|
||||
if len(orderBy) > 0 && orderBy[0] != "" {
|
||||
order = orderBy[0]
|
||||
if strings.TrimSpace(order) == "" {
|
||||
order = "id DESC" // 默认排序
|
||||
}
|
||||
query += " ORDER BY " + order + " LIMIT 1"
|
||||
|
||||
return queryMap(query, values)
|
||||
return queryMap(query, nil)
|
||||
}
|
||||
@@ -86,6 +86,16 @@ func (m M) GetString(key string) string {
|
||||
return ""
|
||||
}
|
||||
|
||||
// GetArray 获取 map 中的 []M,如果不存在或类型不对返回空切片
|
||||
func (m M) GetArray(key string) []M {
|
||||
if v, ok := m[key]; ok {
|
||||
if arr, ok := v.([]M); ok {
|
||||
return arr
|
||||
}
|
||||
}
|
||||
return NewArr()
|
||||
}
|
||||
|
||||
// 核心获取整数的方法,返回 int64
|
||||
func (m M) getIntValue(key string) int64 {
|
||||
if v, ok := m[key]; ok && v != nil {
|
||||
|
||||
Reference in new issue
Block a user