Files
ai_site/platform/internal/dbsync/apply.go
2026-07-31 10:19:22 +08:00

293 lines
8.2 KiB
Go

package dbsync
import (
"context"
"database/sql"
"encoding/json"
"fmt"
"strings"
"time"
)
func PollOutbox(ctx context.Context, db *sql.DB, driver Driver, limit int) ([]OutboxRow, error) {
if limit <= 0 {
limit = 100
}
q := fmt.Sprintf(`SELECT id, table_name, row_pk, op, payload, version, created_at
FROM %s WHERE synced_at IS NULL ORDER BY id ASC LIMIT %d`, OutboxTable, limit)
rows, err := db.QueryContext(ctx, q)
if err != nil {
return nil, err
}
defer rows.Close()
var out []OutboxRow
for rows.Next() {
var r OutboxRow
var created any
if err := rows.Scan(&r.ID, &r.TableName, &r.RowPK, &r.Op, &r.Payload, &r.Version, &created); err != nil {
return nil, err
}
switch v := created.(type) {
case time.Time:
r.CreatedAt = v
case string:
t, _ := time.Parse("2006-01-02 15:04:05", v)
r.CreatedAt = t
case []byte:
t, _ := time.Parse("2006-01-02 15:04:05", string(v))
r.CreatedAt = t
}
out = append(out, r)
}
return out, rows.Err()
}
func MarkSynced(ctx context.Context, db *sql.DB, driver Driver, ids []int64) error {
if len(ids) == 0 {
return nil
}
nowExpr := "CURRENT_TIMESTAMP"
switch driver {
case DriverSQLite:
nowExpr = "datetime('now')"
case DriverMySQL:
nowExpr = "UTC_TIMESTAMP(3)"
case DriverPostgres:
nowExpr = "now()"
}
placeholders := make([]string, len(ids))
args := make([]any, len(ids))
for i, id := range ids {
placeholders[i] = "?"
if driver == DriverPostgres {
placeholders[i] = fmt.Sprintf("$%d", i+1)
}
args[i] = id
}
q := fmt.Sprintf(`UPDATE %s SET synced_at = %s WHERE id IN (%s)`, OutboxTable, nowExpr, strings.Join(placeholders, ","))
_, err := db.ExecContext(ctx, q, args...)
return err
}
func FetchRowJSON(ctx context.Context, db *sql.DB, driver Driver, table, pkCol, pkVal string) (string, map[string]any, error) {
q := fmt.Sprintf(`SELECT * FROM %s WHERE %s = ? LIMIT 1`, quoteIdent(driver, table), quoteIdent(driver, pkCol))
if driver == DriverPostgres {
q = fmt.Sprintf(`SELECT * FROM %s WHERE %s = $1 LIMIT 1`, quoteIdent(driver, table), quoteIdent(driver, pkCol))
}
rows, err := db.QueryContext(ctx, q, pkVal)
if err != nil {
return "", nil, err
}
defer rows.Close()
cols, err := rows.Columns()
if err != nil {
return "", nil, err
}
if !rows.Next() {
return "", nil, sql.ErrNoRows
}
raw := make([]any, len(cols))
ptrs := make([]any, len(cols))
for i := range raw {
ptrs[i] = &raw[i]
}
if err := rows.Scan(ptrs...); err != nil {
return "", nil, err
}
m := map[string]any{}
for i, c := range cols {
m[c] = normalizeValue(raw[i])
}
b, err := json.Marshal(m)
if err != nil {
return "", nil, err
}
return string(b), m, nil
}
func normalizeValue(v any) any {
switch x := v.(type) {
case nil:
return nil
case []byte:
return string(x)
case time.Time:
return x.UTC().Format(time.RFC3339Nano)
default:
return x
}
}
func GetMetaVersion(ctx context.Context, db *sql.DB, driver Driver, table, pk string) (int64, bool, error) {
q := fmt.Sprintf(`SELECT version FROM %s WHERE table_name = ? AND row_pk = ?`, MetaTable)
if driver == DriverPostgres {
q = fmt.Sprintf(`SELECT version FROM %s WHERE table_name = $1 AND row_pk = $2`, MetaTable)
}
var ver int64
err := db.QueryRowContext(ctx, q, table, pk).Scan(&ver)
if err == sql.ErrNoRows {
return 0, false, nil
}
if err != nil {
return 0, false, err
}
return ver, true, nil
}
func UpsertMeta(ctx context.Context, db *sql.DB, driver Driver, table, pk string, version int64) error {
now := time.Now().UTC().Format(time.RFC3339Nano)
switch driver {
case DriverSQLite:
_, err := db.ExecContext(ctx, `
INSERT INTO `+MetaTable+`(table_name,row_pk,version,updated_at) VALUES(?,?,?,?)
ON CONFLICT(table_name,row_pk) DO UPDATE SET version=excluded.version, updated_at=excluded.updated_at`,
table, pk, version, now)
return err
case DriverMySQL:
_, err := db.ExecContext(ctx, `
INSERT INTO `+MetaTable+`(table_name,row_pk,version,updated_at) VALUES(?,?,?,?)
ON DUPLICATE KEY UPDATE version=VALUES(version), updated_at=VALUES(updated_at)`,
table, pk, version, now)
return err
case DriverPostgres:
_, err := db.ExecContext(ctx, `
INSERT INTO `+MetaTable+`(table_name,row_pk,version,updated_at) VALUES($1,$2,$3,$4)
ON CONFLICT(table_name,row_pk) DO UPDATE SET version=EXCLUDED.version, updated_at=EXCLUDED.updated_at`,
table, pk, version, time.Now().UTC())
return err
default:
return fmt.Errorf("unsupported")
}
}
func ApplyChange(ctx context.Context, db *sql.DB, driver Driver, table, pkCol, op, payload string, version int64) error {
return WithApplying(ctx, db, driver, func() error {
return applyChangeInner(ctx, db, driver, table, pkCol, op, payload, version)
})
}
func applyChangeInner(ctx context.Context, db *sql.DB, driver Driver, table, pkCol, op, payload string, version int64) error {
switch op {
case "delete":
var row map[string]any
_ = json.Unmarshal([]byte(payload), &row)
pkVal := ""
if row != nil {
if v, ok := row[pkCol]; ok {
pkVal = fmt.Sprint(v)
}
}
if pkVal == "" {
return fmt.Errorf("delete missing pk")
}
q := fmt.Sprintf(`DELETE FROM %s WHERE %s = ?`, quoteIdent(driver, table), quoteIdent(driver, pkCol))
if driver == DriverPostgres {
q = fmt.Sprintf(`DELETE FROM %s WHERE %s = $1`, quoteIdent(driver, table), quoteIdent(driver, pkCol))
}
if _, err := db.ExecContext(ctx, q, pkVal); err != nil {
return err
}
return UpsertMeta(ctx, db, driver, table, pkVal, version)
default: // upsert
var row map[string]any
if err := json.Unmarshal([]byte(payload), &row); err != nil {
return err
}
pkVal := fmt.Sprint(row[pkCol])
if pkVal == "" || pkVal == "<nil>" {
return fmt.Errorf("upsert missing pk %s", pkCol)
}
cols := make([]string, 0, len(row))
vals := make([]any, 0, len(row))
for k, v := range row {
cols = append(cols, k)
vals = append(vals, v)
}
return upsertRow(ctx, db, driver, table, pkCol, cols, vals, version)
}
}
func upsertRow(ctx context.Context, db *sql.DB, driver Driver, table, pkCol string, cols []string, vals []any, version int64) error {
qcols := make([]string, len(cols))
ph := make([]string, len(cols))
for i, c := range cols {
qcols[i] = quoteIdent(driver, c)
if driver == DriverPostgres {
ph[i] = fmt.Sprintf("$%d", i+1)
} else {
ph[i] = "?"
}
}
pkVal := ""
for i, c := range cols {
if c == pkCol {
pkVal = fmt.Sprint(vals[i])
break
}
}
switch driver {
case DriverSQLite:
q := fmt.Sprintf(`INSERT INTO %s (%s) VALUES (%s) ON CONFLICT(%s) DO UPDATE SET %s`,
quoteIdent(driver, table),
strings.Join(qcols, ","),
strings.Join(ph, ","),
quoteIdent(driver, pkCol),
sqliteSetClause(cols, pkCol, driver),
)
if _, err := db.ExecContext(ctx, q, vals...); err != nil {
return err
}
case DriverMySQL:
sets := make([]string, 0, len(cols))
for _, c := range cols {
if c == pkCol {
continue
}
qi := quoteIdent(driver, c)
sets = append(sets, qi+"=VALUES("+qi+")")
}
if len(sets) == 0 {
sets = append(sets, quoteIdent(driver, pkCol)+"="+quoteIdent(driver, pkCol))
}
q := fmt.Sprintf(`INSERT INTO %s (%s) VALUES (%s) ON DUPLICATE KEY UPDATE %s`,
quoteIdent(driver, table), strings.Join(qcols, ","), strings.Join(ph, ","), strings.Join(sets, ","))
if _, err := db.ExecContext(ctx, q, vals...); err != nil {
return err
}
case DriverPostgres:
sets := make([]string, 0, len(cols))
for _, c := range cols {
if c == pkCol {
continue
}
qi := quoteIdent(driver, c)
sets = append(sets, qi+"=EXCLUDED."+qi)
}
if len(sets) == 0 {
sets = append(sets, quoteIdent(driver, pkCol)+"="+quoteIdent(driver, pkCol))
}
q := fmt.Sprintf(`INSERT INTO %s (%s) VALUES (%s) ON CONFLICT (%s) DO UPDATE SET %s`,
quoteIdent(driver, table), strings.Join(qcols, ","), strings.Join(ph, ","),
quoteIdent(driver, pkCol), strings.Join(sets, ","))
if _, err := db.ExecContext(ctx, q, vals...); err != nil {
return err
}
}
return UpsertMeta(ctx, db, driver, table, pkVal, version)
}
func sqliteSetClause(cols []string, pkCol string, driver Driver) string {
parts := make([]string, 0, len(cols))
for _, c := range cols {
if c == pkCol {
continue
}
qi := quoteIdent(driver, c)
parts = append(parts, qi+"=excluded."+qi)
}
if len(parts) == 0 {
return quoteIdent(driver, pkCol) + "=" + quoteIdent(driver, pkCol)
}
return strings.Join(parts, ",")
}