Files
contracts/base/resource_hooks.go
T
2026-05-15 09:04:12 +08:00

427 lines
11 KiB
Go

package base
import (
"database/sql"
"encoding/json"
"fmt"
"log"
"reflect"
"strconv"
"strings"
"sync"
"git.fsdpf.net/go/db"
"git.fsdpf.net/go/db/exp"
"git.fsdpf.net/go/db/schema"
"git.fsdpf.net/go/req"
"github.com/samber/lo"
)
var changeLogOnce sync.Map // key: Conn,每个数据库连接只建一次 _change_log 表
// isLocalDB 判断是否为本地文件型数据库(LastInsertId 返回最后一条而非第一条)
func isLocalDB(dialect string) bool {
return dialect == "sqlite3" || dialect == "duckdb"
}
type ResourceHooks struct {
res Resource
u req.User
tx *db.TxDatabase
snapshots []map[string]any // UPDATE/DELETE 前预查的行数据
}
func (rh *ResourceHooks) isTracked() bool {
if len(rh.res.HistoryRoles) == 0 || rh.res.IsVirtual() {
return false
}
if lo.Contains(rh.res.HistoryRoles, "00000000-0000-0000-0000-000000000000") {
return true
}
return lo.Some(rh.res.HistoryRoles, rh.u.Roles())
}
// autoCreateChangeLog 在当前连接上建 _change_log 表,每个连接只执行一次
func (rh *ResourceHooks) autoCreateChangeLog() error {
conn := rh.res.DB()
v, _ := changeLogOnce.LoadOrStore(rh.res.Conn, &sync.Once{})
var err error
v.(*sync.Once).Do(func() {
sb := schema.New(conn)
err = sb.Create("_change_log", func(bp *schema.Blueprint) {
bp.Comment = "资源数据变更记录"
bp.BigIncrements("id").AutoIncrement().Comment("ID")
bp.Boolean("enabled").Default("1").Comment("是否有效")
bp.Char("created_user", 36).Default("00000000-0000-0000-0000-000000000000").Comment("操作用户")
bp.Char("owned_user", 36).Default("00000000-0000-0000-0000-000000000000").Comment("拥有者")
bp.Timestamp("created_at").UseCurrent().Comment("操作时间")
bp.Timestamp("updated_at").UseCurrent().Default(db.L("ON UPDATE CURRENT_TIMESTAMP")).Comment("更新时间")
bp.DateTime("deleted_at").Nullable().Comment("删除时间")
bp.Char("resource_uuid", 36).Default("").Comment("资源UUID")
bp.String("category", 10).Default("").Comment("操作类型 INSERT/UPDATE/DELETE")
bp.Char("trace_id", 36).Default("").Comment("请求追踪ID")
bp.Integer("row_id").Default("0").Comment("受影响行ID")
bp.Json("snapshot").Nullable().Comment("操作前数据快照")
})
})
return err
}
// captureSnapshot 在写操作执行前,查出受影响的行存入 snapshots
func (rh *ResourceHooks) captureSnapshot(where exp.ExpressionList) {
sd := rh.res.DB().From(rh.res.GetTable())
if where != nil && len(where.Expressions()) > 0 {
sd = sd.Where(where.Expressions()...)
}
rows, err := sd.Executor().GetRecords()
if err != nil {
log.Printf("_change_log captureSnapshot err: %v", err)
return
}
rh.snapshots = rows
}
// writeChangeLogs 将变更写入 _change_log
func (rh *ResourceHooks) writeChangeLogs(category string, result sql.Result) {
if err := rh.autoCreateChangeLog(); err != nil {
log.Printf("_change_log init err: %v", err)
return
}
traceId := rh.u.Runtime().TraceId()
opUser := rh.u.Uuid()
var rows []any
switch category {
case "INSERT":
lastId, _ := result.LastInsertId()
count, _ := result.RowsAffected()
for i := int64(0); i < count; i++ {
rowId := lastId + i // MySQL: lastId 是第一条
if isLocalDB(rh.res.DB().Dialect()) {
rowId = lastId - count + 1 + i // SQLite: lastId 是最后一条
}
rows = append(rows, map[string]any{
"enabled": true,
"created_user": opUser,
"owned_user": opUser,
"resource_uuid": rh.res.Uuid,
"category": category,
"trace_id": traceId,
"row_id": rowId,
"snapshot": nil,
})
}
case "UPDATE", "DELETE":
for _, snap := range rh.snapshots {
rowId := int64(0)
if id, ok := snap["id"]; ok {
rowId = toChangeLogInt64(id)
}
snapJSON, _ := json.Marshal(snap)
rows = append(rows, map[string]any{
"enabled": true,
"created_user": opUser,
"owned_user": opUser,
"resource_uuid": rh.res.Uuid,
"category": category,
"trace_id": traceId,
"row_id": rowId,
"snapshot": string(snapJSON),
})
}
}
if len(rows) == 0 {
return
}
inserter := rh.res.DB().Insert("_change_log")
if rh.tx != nil {
inserter = rh.tx.Insert("_change_log")
}
if _, err := inserter.Rows(rows...).Executor().Exec(); err != nil {
log.Printf("_change_log write err: %v", err)
}
}
func toChangeLogInt64(v any) int64 {
switch val := v.(type) {
case int64:
return val
case int:
return int64(val)
case int32:
return int64(val)
case float64:
return int64(val)
case string:
n, _ := strconv.ParseInt(val, 10, 64)
return n
}
return 0
}
func (rh *ResourceHooks) Before(dataset interface{}) error {
switch d := dataset.(type) {
case *db.SelectDataset:
return rh.beforeSelectDataset(d)
case *db.InsertDataset:
return rh.beforeInsertDataset(d)
case *db.UpdateDataset:
return rh.beforeUpdateDataset(d)
case *db.DeleteDataset:
return rh.beforeDeleteDataset(d)
}
return nil
}
func (rh *ResourceHooks) After(dataset interface{}, result interface{}) error {
if rh.res.IsVirtual() {
// 返回鉴权后的 DB Builder
return nil
}
// 用户事件
// if rh.res.IsHistoryRecord {
// rh.res.onUserEvent(builder, user)
// }
switch dataset.(type) {
case *db.SelectDataset:
case *db.InsertDataset:
r := result.(sql.Result)
rh.res.onResEvent("INSERT", r)
if rh.isTracked() {
rh.writeChangeLogs("INSERT", r)
}
case *db.UpdateDataset:
r := result.(sql.Result)
rh.res.onResEvent("UPDATE", r)
if rh.isTracked() {
rh.writeChangeLogs("UPDATE", r)
}
case *db.DeleteDataset:
r := result.(sql.Result)
rh.res.onResEvent("DELETE", r)
if rh.isTracked() {
rh.writeChangeLogs("DELETE", r)
}
}
return nil
}
func (rh *ResourceHooks) beforeInsertDataset(id *db.InsertDataset) error {
switch true {
case id.GetClauses().HasRows():
return rh.beforeInsertRows(id)
case id.GetClauses().HasVals():
return rh.beforeInsertColsVals(id)
case id.GetClauses().HasFrom():
return rh.beforeInsertFromQuery(id)
}
return nil
}
func (rh *ResourceHooks) beforeInsertRows(id *db.InsertDataset) error {
rows := id.GetClauses().Rows()
for i := range rows {
rowValue := reflect.ValueOf(rows[i])
if rowValue.Kind() == reflect.Ptr {
rowValue = rowValue.Elem()
}
if rowValue.Kind() == reflect.Struct {
if row, err := exp.NewRecordFromStruct(rowValue.Interface(), true, false); err != nil {
return err
} else {
rows[i] = row
}
}
}
for i := range rows {
row := rows[i].(map[string]any)
// 格式化保存数据
if err := rh.normalizeSaveValue(row); err != nil {
return err
}
// 填充默认数据
if err := rh.applyDefaultValue(row, true, false); err != nil {
return err
}
}
return nil
}
func (rh *ResourceHooks) beforeInsertColsVals(id *db.InsertDataset) error {
cols := id.GetClauses().Cols()
vals := id.GetClauses().Vals()
colsName := []string{}
for _, col := range cols.Columns() {
colsName = append(colsName, col.(exp.IdentifierExpression).GetCol().(string))
}
colsVal := []any{}
for k, v := range lo.OmitByKeys(rh.getFieldsDefaultValue(true, false), colsName) {
cols = cols.Append(db.C(k))
colsVal = append(colsVal, v)
}
for i := 0; i < len(vals); i++ {
for j := 0; j < len(vals[i]); j++ {
if filed, ok := rh.res.GetField(colsName[j]); ok {
vals[i][j] = filed.ToValue(vals[i][j])
}
}
vals[i] = append(vals[i], colsVal...)
}
*id = *id.Cols(cols)
return nil
}
func (rh *ResourceHooks) beforeInsertFromQuery(id *db.InsertDataset) error {
return nil
}
// 规范化保存数据
// 1. 移除系统缺省字段, 如果是 insert 需要调用 applyDefaultValue
// 2. 移除系统中没有的字段
func (rh *ResourceHooks) normalizeSaveValue(row db.Record) error {
for k, v := range row {
if k == "id" || k == "created_user" || k == "created_at" || k == "deleted_at" || k == "updated_at" {
delete(row, k)
} else if val, ok := v.(db.Expression); ok {
row[k] = val
} else if field, ok := rh.res.GetField(k); ok {
row[k] = field.ToValue(v)
} else {
delete(row, k)
}
}
return nil
}
// 填充默认数据
func (rh *ResourceHooks) applyDefaultValue(row db.Record, forInsert, forUpdate bool) error {
// 填充默认数据
for k, v := range rh.getFieldsDefaultValue(forInsert, forUpdate) {
if _, ok := row[k]; !ok {
row[k] = v
}
}
return nil
}
func (rh *ResourceHooks) getFieldsDefaultValue(forInsert, forUpdate bool) map[string]db.Expression {
defs := map[string]db.Expression{}
if forInsert {
for _, item := range rh.res.Fields {
if item.GetCode() == "updated_at" || item.GetCode() == "created_user" || item.GetCode() == "owned_user" {
continue
}
if len(item.Default) > 4 && strings.ToLower(item.Default[0:4]) == "sql:" {
defs[item.GetCode()] = item.GetRawDefault()
} else if item.DataType == req.ResJson {
defs[item.GetCode()] = item.GetRawDefault()
}
}
if _, ok := defs["owned_user"]; !ok {
defs["owned_user"] = db.V(rh.u.Uuid())
}
defs["created_user"] = db.V(rh.u.Uuid())
}
if forUpdate && isLocalDB(rh.res.DB().Dialect()) {
defs["updated_at"] = db.L("CURRENT_TIMESTAMP")
}
return defs
}
func (rh *ResourceHooks) beforeSelectDataset(sd *db.SelectDataset) error {
sub, ex := rh.res.GetRolesCondition(rh.u)
if ex != nil {
*sd = *sd.Where(ex)
} else if sub != nil {
table := sd.GetClauses().From().Columns()[0]
if alias, ok := table.(exp.AliasedExpression); ok {
*sd = *sd.From(db.V(sub.Expression()).As(alias.GetAs()))
} else {
*sd = *sd.From(db.V(sub.Expression()).As(table))
}
}
return nil
}
func (rh *ResourceHooks) beforeUpdateDataset(ud *db.UpdateDataset) error {
// 虚拟资源暂时不考虑鉴权
if rh.res.IsVirtual() {
return nil
}
sub, ex := rh.res.GetRolesCondition(rh.u)
if ex != nil {
*ud = *ud.Where(ex)
} else if sub != nil {
table := ud.GetClauses().Table()
if alias, ok := table.(exp.AliasedExpression); ok {
*ud = *ud.Where(alias.GetAs().Col("id").In(sub.Select(db.T(rh.res.GetCode()).Col("id"))))
} else {
*ud = *ud.Where(db.L("?.id", table).In(sub.Select(db.T(rh.res.GetCode()).Col("id"))))
}
}
// 记录变更前快照(role 条件已应用,WHERE 与实际 UPDATE 一致)
if rh.isTracked() {
rh.captureSnapshot(ud.GetClauses().Where())
}
// 格式化保存数据
if ud.GetClauses().HasSetValues() {
udv := ud.GetClauses().SetValues()
switch data := udv.(type) {
case map[string]any:
return rh.normalizeSaveValue(data)
case db.Record:
return rh.normalizeSaveValue(data)
default:
return fmt.Errorf("unsupported type: %T", udv)
}
}
return nil
}
func (rh *ResourceHooks) beforeDeleteDataset(dd *db.DeleteDataset) error {
// 虚拟资源暂时不考虑鉴权
if rh.res.IsVirtual() {
return nil
}
sub, ex := rh.res.GetRolesCondition(rh.u)
if ex != nil {
*dd = *dd.Where(ex)
} else if sub != nil {
table := dd.GetClauses().From().Columns()[0]
if alias, ok := table.(exp.AliasedExpression); ok {
*dd = *dd.Where(alias.GetAs().Col("id").In(sub.Select(db.T(rh.res.GetCode()).Col("id"))))
} else {
*dd = *dd.Where(db.L("?.id", table).In(sub.Select(db.T(rh.res.GetCode()).Col("id"))))
}
}
// 记录变更前快照(role 条件已应用,WHERE 与实际 DELETE 一致)
if rh.isTracked() {
rh.captureSnapshot(dd.GetClauses().Where())
}
return nil
}