293 lines
8.2 KiB
Go
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, ",")
|
|
}
|