feat: harden loose-offline sync for user JWT, schema, and console ops
Enable Binding-scoped agent push/pull, empty-table schema ensure, SyncPage inspect/drop-table, default module import, and agent-bound publish docs from the 宇恒联调意见. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
@@ -11,7 +11,7 @@ import (
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
// Binding 本机库 ↔ 线上库映射(P1;不挡推送,供登记/查询)。
|
||||
// Binding 本机库 ↔ 线上库映射(登记/查询;用户自助 push 时按 user_id+online_db_id 鉴权)。
|
||||
type Binding struct {
|
||||
ID string `json:"id"`
|
||||
TenantID int64 `json:"tenant_id"`
|
||||
@@ -19,9 +19,12 @@ type Binding struct {
|
||||
LocalDatabaseID string `json:"local_database_id"`
|
||||
OnlineDBID string `json:"online_db_id"`
|
||||
ChannelID string `json:"channel_id,omitempty"`
|
||||
Note string `json:"note,omitempty"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
// 可读名(与宇恒 database_name / display_name 对齐;控制台优先展示)
|
||||
DatabaseName string `json:"database_name,omitempty"` // 本地库可读名
|
||||
DisplayName string `json:"display_name,omitempty"` // 线上库可读名
|
||||
Note string `json:"note,omitempty"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
func (s *FileStore) bindingPath() string {
|
||||
@@ -37,6 +40,8 @@ func (s *FileStore) EnsureBinding(b Binding) (Binding, error) {
|
||||
}
|
||||
b.LocalDatabaseID = strings.TrimSpace(b.LocalDatabaseID)
|
||||
b.OnlineDBID = strings.TrimSpace(b.OnlineDBID)
|
||||
b.DatabaseName = strings.TrimSpace(b.DatabaseName)
|
||||
b.DisplayName = strings.TrimSpace(b.DisplayName)
|
||||
if b.TenantID <= 0 {
|
||||
return b, fmt.Errorf("tenant_id required")
|
||||
}
|
||||
@@ -53,6 +58,12 @@ func (s *FileStore) EnsureBinding(b Binding) (Binding, error) {
|
||||
if b.ChannelID != "" {
|
||||
list[i].ChannelID = b.ChannelID
|
||||
}
|
||||
if b.DatabaseName != "" {
|
||||
list[i].DatabaseName = b.DatabaseName
|
||||
}
|
||||
if b.DisplayName != "" {
|
||||
list[i].DisplayName = b.DisplayName
|
||||
}
|
||||
if b.Note != "" {
|
||||
list[i].Note = b.Note
|
||||
}
|
||||
@@ -79,6 +90,11 @@ func (s *FileStore) EnsureBinding(b Binding) (Binding, error) {
|
||||
}
|
||||
|
||||
func (s *FileStore) ListBindings(tenantID int64, localDatabaseID string) ([]Binding, error) {
|
||||
return s.ListBindingsFiltered(tenantID, 0, localDatabaseID)
|
||||
}
|
||||
|
||||
// ListBindingsFiltered 按租户列出;userID>0 时仅返回该用户的 Binding。
|
||||
func (s *FileStore) ListBindingsFiltered(tenantID, userID int64, localDatabaseID string) ([]Binding, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
list, err := s.readBindingsUnlocked()
|
||||
@@ -90,6 +106,9 @@ func (s *FileStore) ListBindings(tenantID int64, localDatabaseID string) ([]Bind
|
||||
if b.TenantID != tenantID {
|
||||
continue
|
||||
}
|
||||
if userID > 0 && b.UserID != userID {
|
||||
continue
|
||||
}
|
||||
if localDatabaseID != "" && b.LocalDatabaseID != localDatabaseID {
|
||||
continue
|
||||
}
|
||||
@@ -98,6 +117,47 @@ func (s *FileStore) ListBindings(tenantID int64, localDatabaseID string) ([]Bind
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// UserOwnsOnlineDB 用户是否登记了该 online_db_id(可选限定 channel)。
|
||||
func (s *FileStore) UserOwnsOnlineDB(tenantID, userID int64, channelID, onlineDBID string) bool {
|
||||
if tenantID <= 0 || userID <= 0 || strings.TrimSpace(onlineDBID) == "" {
|
||||
return false
|
||||
}
|
||||
list, err := s.ListBindingsFiltered(tenantID, userID, "")
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
online := strings.TrimSpace(onlineDBID)
|
||||
ch := strings.TrimSpace(channelID)
|
||||
for _, b := range list {
|
||||
if strings.TrimSpace(b.OnlineDBID) != online {
|
||||
continue
|
||||
}
|
||||
if ch != "" && b.ChannelID != "" && b.ChannelID != ch {
|
||||
continue
|
||||
}
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// UserCanAccessChannel 用户是否登记了指向该通道的 Binding(channel_id 空视为未限定通道)。
|
||||
func (s *FileStore) UserCanAccessChannel(tenantID, userID int64, channelID string) bool {
|
||||
if tenantID <= 0 || userID <= 0 || strings.TrimSpace(channelID) == "" {
|
||||
return false
|
||||
}
|
||||
list, err := s.ListBindingsFiltered(tenantID, userID, "")
|
||||
if err != nil || len(list) == 0 {
|
||||
return false
|
||||
}
|
||||
ch := strings.TrimSpace(channelID)
|
||||
for _, b := range list {
|
||||
if b.ChannelID == "" || b.ChannelID == ch {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (s *FileStore) GetBinding(tenantID int64, localDatabaseID string) (*Binding, error) {
|
||||
list, err := s.ListBindings(tenantID, localDatabaseID)
|
||||
if err != nil {
|
||||
|
||||
@@ -16,16 +16,19 @@ func TestEnsureBindingUpsert(t *testing.T) {
|
||||
LocalDatabaseID: "local-a",
|
||||
OnlineDBID: "online-1",
|
||||
ChannelID: "ch1",
|
||||
DatabaseName: "本地演示库",
|
||||
DisplayName: "AI建站智能体API",
|
||||
})
|
||||
if err != nil || b.ID == "" {
|
||||
if err != nil || b.ID == "" || b.DisplayName != "AI建站智能体API" {
|
||||
t.Fatalf("ensure: %+v err=%v", b, err)
|
||||
}
|
||||
b2, err := st.EnsureBinding(Binding{
|
||||
TenantID: 1,
|
||||
LocalDatabaseID: "local-a",
|
||||
OnlineDBID: "online-2",
|
||||
DisplayName: "线上库B",
|
||||
})
|
||||
if err != nil || b2.OnlineDBID != "online-2" || b2.ID != b.ID {
|
||||
if err != nil || b2.OnlineDBID != "online-2" || b2.ID != b.ID || b2.DisplayName != "线上库B" {
|
||||
t.Fatalf("upsert: %+v err=%v", b2, err)
|
||||
}
|
||||
list, err := st.ListBindings(1, "local-a")
|
||||
@@ -33,3 +36,42 @@ func TestEnsureBindingUpsert(t *testing.T) {
|
||||
t.Fatalf("list=%d err=%v", len(list), err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUserBindingScope(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
st, err := NewFileStore(filepath.Join(dir, "dbsync"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, err = st.EnsureBinding(Binding{
|
||||
TenantID: 1, UserID: 10,
|
||||
LocalDatabaseID: "loc-a", OnlineDBID: "on-a", ChannelID: "ch1",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, err = st.EnsureBinding(Binding{
|
||||
TenantID: 1, UserID: 20,
|
||||
LocalDatabaseID: "loc-b", OnlineDBID: "on-b", ChannelID: "ch1",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
mine, err := st.ListBindingsFiltered(1, 10, "")
|
||||
if err != nil || len(mine) != 1 || mine[0].OnlineDBID != "on-a" {
|
||||
t.Fatalf("filtered: %+v err=%v", mine, err)
|
||||
}
|
||||
if !st.UserOwnsOnlineDB(1, 10, "ch1", "on-a") {
|
||||
t.Fatal("should own on-a")
|
||||
}
|
||||
if st.UserOwnsOnlineDB(1, 10, "ch1", "on-b") {
|
||||
t.Fatal("must not own other's online_db_id")
|
||||
}
|
||||
if !st.UserCanAccessChannel(1, 10, "ch1") {
|
||||
t.Fatal("should access ch1")
|
||||
}
|
||||
if st.UserCanAccessChannel(1, 10, "ch-other") {
|
||||
t.Fatal("must not access unbound channel")
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -19,7 +19,7 @@ func Open(driver Driver, dsn string) (*sql.DB, error) {
|
||||
if dsn == "" {
|
||||
return nil, fmt.Errorf("sqlite dsn required")
|
||||
}
|
||||
// modernc.org/sqlite 注册名为 sqlite
|
||||
dsn = normalizeSQLiteDSN(dsn)
|
||||
case DriverMySQL:
|
||||
name = "mysql"
|
||||
case DriverPostgres:
|
||||
@@ -31,8 +31,14 @@ func Open(driver Driver, dsn string) (*sql.DB, error) {
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
db.SetMaxOpenConns(5)
|
||||
db.SetMaxIdleConns(2)
|
||||
if driver == DriverSQLite {
|
||||
// SQLite 文件锁:单连接更稳,避免 Windows 上并发 open 卡死/EOF
|
||||
db.SetMaxOpenConns(1)
|
||||
db.SetMaxIdleConns(1)
|
||||
} else {
|
||||
db.SetMaxOpenConns(5)
|
||||
db.SetMaxIdleConns(2)
|
||||
}
|
||||
if err := db.Ping(); err != nil {
|
||||
_ = db.Close()
|
||||
return nil, err
|
||||
@@ -40,6 +46,33 @@ func Open(driver Driver, dsn string) (*sql.DB, error) {
|
||||
return db, nil
|
||||
}
|
||||
|
||||
// normalizeSQLiteDSN:Windows 反斜杠 → 正斜杠;裸盘符路径补 file:;补 busy_timeout。
|
||||
func normalizeSQLiteDSN(dsn string) string {
|
||||
dsn = strings.TrimSpace(dsn)
|
||||
lower := strings.ToLower(dsn)
|
||||
if strings.HasPrefix(lower, "file:") {
|
||||
rest := dsn[5:]
|
||||
path, query, hasQ := strings.Cut(rest, "?")
|
||||
path = strings.ReplaceAll(path, `\`, `/`)
|
||||
if hasQ {
|
||||
dsn = "file:" + path + "?" + query
|
||||
} else {
|
||||
dsn = "file:" + path
|
||||
}
|
||||
} else if strings.Contains(dsn, `\`) || (len(dsn) >= 2 && dsn[1] == ':') {
|
||||
path := strings.ReplaceAll(dsn, `\`, `/`)
|
||||
dsn = "file:" + path
|
||||
}
|
||||
if !strings.Contains(strings.ToLower(dsn), "busy_timeout") {
|
||||
if strings.Contains(dsn, "?") {
|
||||
dsn += "&_pragma=busy_timeout(5000)"
|
||||
} else {
|
||||
dsn += "?_pragma=busy_timeout(5000)"
|
||||
}
|
||||
}
|
||||
return dsn
|
||||
}
|
||||
|
||||
func quoteIdent(driver Driver, name string) string {
|
||||
name = strings.ReplaceAll(name, "`", "")
|
||||
name = strings.ReplaceAll(name, `"`, "")
|
||||
|
||||
31
platform/internal/dbsync/dialect_test.go
Normal file
31
platform/internal/dbsync/dialect_test.go
Normal file
@@ -0,0 +1,31 @@
|
||||
package dbsync
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestNormalizeSQLiteDSN(t *testing.T) {
|
||||
got := normalizeSQLiteDSN(`file:E:\data\remote.db`)
|
||||
if got != "file:E:/data/remote.db?_pragma=busy_timeout(5000)" {
|
||||
t.Fatalf("got %q", got)
|
||||
}
|
||||
got = normalizeSQLiteDSN(`E:\data\remote.db`)
|
||||
if got != "file:E:/data/remote.db?_pragma=busy_timeout(5000)" {
|
||||
t.Fatalf("bare path got %q", got)
|
||||
}
|
||||
got = normalizeSQLiteDSN(`file:E:/data/remote.db?_pragma=foreign_keys(1)`)
|
||||
if got != "file:E:/data/remote.db?_pragma=foreign_keys(1)&_pragma=busy_timeout(5000)" {
|
||||
t.Fatalf("keep query got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsRetryable(t *testing.T) {
|
||||
if !IsRetryable(Retryable("x")) {
|
||||
t.Fatal("expected retryable")
|
||||
}
|
||||
if IsRetryable(fmtError("plain")) {
|
||||
t.Fatal("plain should not")
|
||||
}
|
||||
}
|
||||
|
||||
type fmtError string
|
||||
|
||||
func (e fmtError) Error() string { return string(e) }
|
||||
87
platform/internal/dbsync/drop_table.go
Normal file
87
platform/internal/dbsync/drop_table.go
Normal file
@@ -0,0 +1,87 @@
|
||||
package dbsync
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// DropTableResult 控制台删表结果。
|
||||
type DropTableResult struct {
|
||||
OK bool `json:"ok"`
|
||||
Side string `json:"side"`
|
||||
Table string `json:"table"`
|
||||
Dropped bool `json:"dropped"`
|
||||
MetaCleared int64 `json:"meta_cleared,omitempty"`
|
||||
Message string `json:"message,omitempty"`
|
||||
}
|
||||
|
||||
// DropTableOnEndpoint 删除业务表(禁止 _ajz_*);并清理 _ajz_sync_meta 中该表残留。
|
||||
// 注意:仅作用于当前 endpoint(线上 A 或本机 B),不向另一侧同步 DDL。
|
||||
func DropTableOnEndpoint(ctx context.Context, ep Endpoint, side, table string) (*DropTableResult, error) {
|
||||
table = strings.TrimSpace(table)
|
||||
side = strings.TrimSpace(side)
|
||||
if side == "" {
|
||||
side = "remote"
|
||||
}
|
||||
if table == "" {
|
||||
return nil, fmt.Errorf("table required")
|
||||
}
|
||||
if strings.HasPrefix(table, "_ajz_") {
|
||||
return nil, fmt.Errorf("同步系统表不可删: %s", table)
|
||||
}
|
||||
if strings.EqualFold(table, "sqlite_master") || strings.HasPrefix(strings.ToLower(table), "sqlite_") {
|
||||
return nil, fmt.Errorf("系统表不可删: %s", table)
|
||||
}
|
||||
|
||||
db, owned, err := openInspectDB(ep)
|
||||
if err != nil {
|
||||
return &DropTableResult{OK: false, Side: side, Table: table, Message: err.Error()}, err
|
||||
}
|
||||
if owned {
|
||||
defer db.Close()
|
||||
}
|
||||
|
||||
exists, err := tableExists(ctx, db, ep.Driver, table)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !exists {
|
||||
return &DropTableResult{
|
||||
OK: true,
|
||||
Side: side,
|
||||
Table: table,
|
||||
Dropped: false,
|
||||
Message: "表不存在(可能已删)",
|
||||
}, nil
|
||||
}
|
||||
|
||||
ddl := fmt.Sprintf(`DROP TABLE IF EXISTS %s`, quoteIdent(ep.Driver, table))
|
||||
if _, err := db.ExecContext(ctx, ddl); err != nil {
|
||||
return nil, fmt.Errorf("drop table: %w", err)
|
||||
}
|
||||
|
||||
var metaCleared int64
|
||||
if metaExists, _ := tableExists(ctx, db, ep.Driver, MetaTable); metaExists {
|
||||
q := fmt.Sprintf(`DELETE FROM %s WHERE table_name = ?`, quoteIdent(ep.Driver, MetaTable))
|
||||
if ep.Driver == DriverPostgres {
|
||||
q = fmt.Sprintf(`DELETE FROM %s WHERE table_name = $1`, quoteIdent(ep.Driver, MetaTable))
|
||||
}
|
||||
res, err := db.ExecContext(ctx, q, table)
|
||||
if err == nil && res != nil {
|
||||
metaCleared, _ = res.RowsAffected()
|
||||
}
|
||||
}
|
||||
|
||||
// 若该 DSN 在 remote 池里,丢掉缓存连接,避免旧 schema 缓存感。
|
||||
InvalidateRemote(ep.Driver, ep.DSN)
|
||||
|
||||
return &DropTableResult{
|
||||
OK: true,
|
||||
Side: side,
|
||||
Table: table,
|
||||
Dropped: true,
|
||||
MetaCleared: metaCleared,
|
||||
Message: "已删除;若另一侧仍有同名表,不会自动同步删除,需分别处理或避免再次 push/ensure",
|
||||
}, nil
|
||||
}
|
||||
56
platform/internal/dbsync/drop_table_test.go
Normal file
56
platform/internal/dbsync/drop_table_test.go
Normal file
@@ -0,0 +1,56 @@
|
||||
package dbsync
|
||||
|
||||
import (
|
||||
"context"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestDropTableOnEndpoint(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
dsn := "file:" + filepath.ToSlash(filepath.Join(dir, "drop.db")) + "?_pragma=busy_timeout(5000)"
|
||||
db, err := Open(DriverSQLite, dsn)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ctx := context.Background()
|
||||
if err := EnsureTableFromColumns(ctx, db, DriverSQLite, "模块·测2·item", "id", []string{"id", "before_day"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := EnsureMeta(ctx, db, DriverSQLite); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_ = UpsertMeta(ctx, db, DriverSQLite, "模块·测2·item", "pk1", 1)
|
||||
_ = db.Close()
|
||||
|
||||
ep := Endpoint{Driver: DriverSQLite, DSN: dsn}
|
||||
res, err := DropTableOnEndpoint(ctx, ep, "remote", "模块·测2·item")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !res.OK || !res.Dropped {
|
||||
t.Fatalf("unexpected: %+v", res)
|
||||
}
|
||||
if res.MetaCleared < 1 {
|
||||
t.Fatalf("want meta cleared: %+v", res)
|
||||
}
|
||||
|
||||
db2, err := Open(DriverSQLite, dsn)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer db2.Close()
|
||||
exists, err := tableExists(ctx, db2, DriverSQLite, "模块·测2·item")
|
||||
if err != nil || exists {
|
||||
t.Fatalf("table should be gone exists=%v err=%v", exists, err)
|
||||
}
|
||||
|
||||
// idempotent
|
||||
res2, err := DropTableOnEndpoint(ctx, ep, "remote", "模块·测2·item")
|
||||
if err != nil || !res2.OK || res2.Dropped {
|
||||
t.Fatalf("idempotent: %+v err=%v", res2, err)
|
||||
}
|
||||
if _, err := DropTableOnEndpoint(ctx, ep, "remote", "_ajz_sync_meta"); err == nil {
|
||||
t.Fatal("expected reject system table")
|
||||
}
|
||||
}
|
||||
105
platform/internal/dbsync/ensure_table.go
Normal file
105
platform/internal/dbsync/ensure_table.go
Normal file
@@ -0,0 +1,105 @@
|
||||
package dbsync
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// EnsureTableFromRow 表不存在时按行字段自动建表(TEXT 列 + PK),满足 Z4「按 Binding 接受任意表」。
|
||||
func EnsureTableFromRow(ctx context.Context, db *sql.DB, driver Driver, table, pkCol string, row map[string]any) error {
|
||||
table = strings.TrimSpace(table)
|
||||
pkCol = strings.TrimSpace(pkCol)
|
||||
if table == "" || pkCol == "" {
|
||||
return fmt.Errorf("table and pk required")
|
||||
}
|
||||
if row == nil {
|
||||
row = map[string]any{pkCol: ""}
|
||||
}
|
||||
cols := make([]string, 0, len(row))
|
||||
seen := map[string]struct{}{}
|
||||
if _, ok := row[pkCol]; !ok {
|
||||
cols = append(cols, pkCol)
|
||||
seen[pkCol] = struct{}{}
|
||||
}
|
||||
for k := range row {
|
||||
k = strings.TrimSpace(k)
|
||||
if k == "" {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[k]; ok {
|
||||
continue
|
||||
}
|
||||
seen[k] = struct{}{}
|
||||
cols = append(cols, k)
|
||||
}
|
||||
return EnsureTableFromColumns(ctx, db, driver, table, pkCol, cols)
|
||||
}
|
||||
|
||||
// EnsureTableFromColumns 表不存在时按列名建空表(全部 TEXT,指定 PK)。已存在则幂等跳过。
|
||||
// 用于空表结构同步:本机有空表 → 线上也建同名空表(无需 outbox 行)。
|
||||
func EnsureTableFromColumns(ctx context.Context, db *sql.DB, driver Driver, table, pkCol string, columns []string) error {
|
||||
table = strings.TrimSpace(table)
|
||||
pkCol = strings.TrimSpace(pkCol)
|
||||
if table == "" {
|
||||
return fmt.Errorf("table required")
|
||||
}
|
||||
if strings.HasPrefix(table, "_ajz_") {
|
||||
return fmt.Errorf("sync system table not allowed: %s", table)
|
||||
}
|
||||
if pkCol == "" {
|
||||
pkCol = "id"
|
||||
}
|
||||
exists, err := tableExists(ctx, db, driver, table)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if exists {
|
||||
return nil
|
||||
}
|
||||
seen := map[string]struct{}{}
|
||||
defs := make([]string, 0, len(columns)+1)
|
||||
defs = append(defs, fmt.Sprintf("%s TEXT PRIMARY KEY", quoteIdent(driver, pkCol)))
|
||||
seen[pkCol] = struct{}{}
|
||||
for _, c := range columns {
|
||||
c = strings.TrimSpace(c)
|
||||
if c == "" {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[c]; ok {
|
||||
continue
|
||||
}
|
||||
seen[c] = struct{}{}
|
||||
defs = append(defs, fmt.Sprintf("%s TEXT", quoteIdent(driver, c)))
|
||||
}
|
||||
if len(defs) == 0 {
|
||||
return fmt.Errorf("columns required")
|
||||
}
|
||||
ddl := fmt.Sprintf(`CREATE TABLE IF NOT EXISTS %s (%s)`, quoteIdent(driver, table), strings.Join(defs, ", "))
|
||||
_, err = db.ExecContext(ctx, ddl)
|
||||
return err
|
||||
}
|
||||
|
||||
func tableExists(ctx context.Context, db *sql.DB, driver Driver, table string) (bool, error) {
|
||||
var q string
|
||||
switch driver {
|
||||
case DriverSQLite:
|
||||
q = `SELECT 1 FROM sqlite_master WHERE type='table' AND name=? LIMIT 1`
|
||||
case DriverMySQL:
|
||||
q = `SELECT 1 FROM information_schema.tables WHERE table_schema=DATABASE() AND table_name=? LIMIT 1`
|
||||
case DriverPostgres:
|
||||
q = `SELECT 1 FROM information_schema.tables WHERE table_schema=current_schema() AND table_name=$1 LIMIT 1`
|
||||
default:
|
||||
return false, fmt.Errorf("unsupported driver")
|
||||
}
|
||||
var n int
|
||||
err := db.QueryRowContext(ctx, q, table).Scan(&n)
|
||||
if err == sql.ErrNoRows {
|
||||
return false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return true, nil
|
||||
}
|
||||
29
platform/internal/dbsync/ensure_table_test.go
Normal file
29
platform/internal/dbsync/ensure_table_test.go
Normal file
@@ -0,0 +1,29 @@
|
||||
package dbsync
|
||||
|
||||
import (
|
||||
"context"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestEnsureTableFromRow(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
dsn := "file:" + filepath.ToSlash(filepath.Join(dir, "t.db")) + "?_pragma=busy_timeout(5000)"
|
||||
db, err := Open(DriverSQLite, dsn)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer db.Close()
|
||||
ctx := context.Background()
|
||||
row := map[string]any{"id": "u1", "title": "x"}
|
||||
if err := EnsureTableFromRow(ctx, db, DriverSQLite, "orders", "id", row); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := EnsureTableFromRow(ctx, db, DriverSQLite, "orders", "id", row); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var n int
|
||||
if err := db.QueryRowContext(ctx, `SELECT COUNT(1) FROM sqlite_master WHERE type='table' AND name='orders'`).Scan(&n); err != nil || n != 1 {
|
||||
t.Fatalf("table missing n=%d err=%v", n, err)
|
||||
}
|
||||
}
|
||||
277
platform/internal/dbsync/inspect.go
Normal file
277
platform/internal/dbsync/inspect.go
Normal file
@@ -0,0 +1,277 @@
|
||||
package dbsync
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// TableInspect 线上/本机库一张表的摘要(控制台验同步用)。
|
||||
type TableInspect struct {
|
||||
Name string `json:"name"`
|
||||
RowCount int64 `json:"row_count"`
|
||||
Columns []string `json:"columns"`
|
||||
ColumnCount int `json:"column_count"`
|
||||
}
|
||||
|
||||
// InspectResult 库体检结果。
|
||||
type InspectResult struct {
|
||||
OK bool `json:"ok"`
|
||||
Side string `json:"side"` // remote | local
|
||||
Driver string `json:"driver"`
|
||||
DSNHint string `json:"dsn_hint,omitempty"`
|
||||
Tables []TableInspect `json:"tables"`
|
||||
Message string `json:"message,omitempty"`
|
||||
}
|
||||
|
||||
// PreviewResult 单表行预览。
|
||||
type PreviewResult struct {
|
||||
OK bool `json:"ok"`
|
||||
Table string `json:"table"`
|
||||
Columns []string `json:"columns"`
|
||||
Rows []map[string]any `json:"rows"`
|
||||
Total int64 `json:"total"`
|
||||
Limit int `json:"limit"`
|
||||
Message string `json:"message,omitempty"`
|
||||
}
|
||||
|
||||
// InspectEndpoint 列出业务表 + 行数 + 列名(含 _ajz_ 同步系统表,便于对照)。
|
||||
func InspectEndpoint(ctx context.Context, ep Endpoint, side string, includeSyncMeta bool) (*InspectResult, error) {
|
||||
db, owned, err := openInspectDB(ep)
|
||||
if err != nil {
|
||||
return &InspectResult{OK: false, Side: side, Driver: string(ep.Driver), Message: err.Error()}, err
|
||||
}
|
||||
if owned {
|
||||
defer db.Close()
|
||||
}
|
||||
names, err := listTablesForInspect(ctx, db, ep.Driver, includeSyncMeta)
|
||||
if err != nil {
|
||||
return &InspectResult{OK: false, Side: side, Driver: string(ep.Driver), Message: err.Error()}, err
|
||||
}
|
||||
out := &InspectResult{
|
||||
OK: true,
|
||||
Side: side,
|
||||
Driver: string(ep.Driver),
|
||||
DSNHint: dsnFileHint(ep.DSN),
|
||||
Tables: make([]TableInspect, 0, len(names)),
|
||||
}
|
||||
for _, name := range names {
|
||||
ti := TableInspect{Name: name}
|
||||
cols, _ := listColumns(ctx, db, ep.Driver, name)
|
||||
ti.Columns = cols
|
||||
ti.ColumnCount = len(cols)
|
||||
n, _ := countRows(ctx, db, ep.Driver, name)
|
||||
ti.RowCount = n
|
||||
out.Tables = append(out.Tables, ti)
|
||||
}
|
||||
if len(out.Tables) == 0 {
|
||||
out.Message = "库可连接,但尚无业务表(可能还未 push)"
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// PreviewTable 预览表前 limit 行(默认 50,最大 200)。
|
||||
func PreviewTable(ctx context.Context, ep Endpoint, table string, limit int) (*PreviewResult, error) {
|
||||
table = strings.TrimSpace(table)
|
||||
if table == "" {
|
||||
return nil, fmt.Errorf("table required")
|
||||
}
|
||||
if limit <= 0 {
|
||||
limit = 50
|
||||
}
|
||||
if limit > 200 {
|
||||
limit = 200
|
||||
}
|
||||
db, owned, err := openInspectDB(ep)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if owned {
|
||||
defer db.Close()
|
||||
}
|
||||
cols, err := listColumns(ctx, db, ep.Driver, table)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
total, _ := countRows(ctx, db, ep.Driver, table)
|
||||
rows, err := fetchPreviewRows(ctx, db, ep.Driver, table, cols, limit)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &PreviewResult{
|
||||
OK: true,
|
||||
Table: table,
|
||||
Columns: cols,
|
||||
Rows: rows,
|
||||
Total: total,
|
||||
Limit: limit,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func openInspectDB(ep Endpoint) (*sql.DB, bool, error) {
|
||||
// 控制台验库用短连接,不进 remote 池,避免 SQLite 文件被长期占用。
|
||||
db, err := Open(ep.Driver, ep.DSN)
|
||||
if err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
return db, true, nil
|
||||
}
|
||||
|
||||
func listTablesForInspect(ctx context.Context, db *sql.DB, driver Driver, includeSyncMeta bool) ([]string, error) {
|
||||
var q string
|
||||
switch driver {
|
||||
case DriverSQLite:
|
||||
if includeSyncMeta {
|
||||
q = `SELECT name FROM sqlite_master WHERE type='table' AND name NOT LIKE 'sqlite_%' ORDER BY name`
|
||||
} else {
|
||||
q = `SELECT name FROM sqlite_master WHERE type='table' AND name NOT LIKE 'sqlite_%' AND name NOT LIKE '_ajz_%' ORDER BY name`
|
||||
}
|
||||
case DriverMySQL:
|
||||
if includeSyncMeta {
|
||||
q = `SELECT table_name FROM information_schema.tables WHERE table_schema = DATABASE() ORDER BY table_name`
|
||||
} else {
|
||||
q = `SELECT table_name FROM information_schema.tables WHERE table_schema = DATABASE() AND table_name NOT LIKE '\_ajz\_%' ORDER BY table_name`
|
||||
}
|
||||
case DriverPostgres:
|
||||
if includeSyncMeta {
|
||||
q = `SELECT tablename FROM pg_tables WHERE schemaname='public' ORDER BY tablename`
|
||||
} else {
|
||||
q = `SELECT tablename FROM pg_tables WHERE schemaname='public' AND tablename NOT LIKE '\_ajz\_%' ORDER BY tablename`
|
||||
}
|
||||
default:
|
||||
return nil, fmt.Errorf("unsupported driver")
|
||||
}
|
||||
rows, err := db.QueryContext(ctx, q)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
var out []string
|
||||
for rows.Next() {
|
||||
var n string
|
||||
if err := rows.Scan(&n); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, n)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
func listColumns(ctx context.Context, db *sql.DB, driver Driver, table string) ([]string, error) {
|
||||
var (
|
||||
rows *sql.Rows
|
||||
err error
|
||||
)
|
||||
switch driver {
|
||||
case DriverSQLite:
|
||||
rows, err = db.QueryContext(ctx, fmt.Sprintf(`PRAGMA table_info(%s)`, quoteIdent(driver, table)))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
var cols []string
|
||||
for rows.Next() {
|
||||
var cid int
|
||||
var name, typ string
|
||||
var notnull, pk int
|
||||
var dflt sql.NullString
|
||||
if err := rows.Scan(&cid, &name, &typ, ¬null, &dflt, &pk); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cols = append(cols, name)
|
||||
}
|
||||
return cols, rows.Err()
|
||||
case DriverMySQL:
|
||||
rows, err = db.QueryContext(ctx, `
|
||||
SELECT COLUMN_NAME FROM information_schema.columns
|
||||
WHERE table_schema = DATABASE() AND table_name = ? ORDER BY ORDINAL_POSITION`, table)
|
||||
case DriverPostgres:
|
||||
rows, err = db.QueryContext(ctx, `
|
||||
SELECT column_name FROM information_schema.columns
|
||||
WHERE table_schema = 'public' AND table_name = $1 ORDER BY ordinal_position`, table)
|
||||
default:
|
||||
return nil, fmt.Errorf("unsupported")
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
var cols []string
|
||||
for rows.Next() {
|
||||
var n string
|
||||
if err := rows.Scan(&n); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cols = append(cols, n)
|
||||
}
|
||||
return cols, rows.Err()
|
||||
}
|
||||
|
||||
func countRows(ctx context.Context, db *sql.DB, driver Driver, table string) (int64, error) {
|
||||
q := fmt.Sprintf(`SELECT COUNT(*) FROM %s`, quoteIdent(driver, table))
|
||||
var n int64
|
||||
err := db.QueryRowContext(ctx, q).Scan(&n)
|
||||
return n, err
|
||||
}
|
||||
|
||||
func fetchPreviewRows(ctx context.Context, db *sql.DB, driver Driver, table string, cols []string, limit int) ([]map[string]any, error) {
|
||||
if len(cols) == 0 {
|
||||
return []map[string]any{}, nil
|
||||
}
|
||||
q := fmt.Sprintf(`SELECT * FROM %s LIMIT %d`, quoteIdent(driver, table), limit)
|
||||
rows, err := db.QueryContext(ctx, q)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
colNames, err := rows.Columns()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var out []map[string]any
|
||||
for rows.Next() {
|
||||
raw := make([]any, len(colNames))
|
||||
ptrs := make([]any, len(colNames))
|
||||
for i := range raw {
|
||||
ptrs[i] = &raw[i]
|
||||
}
|
||||
if err := rows.Scan(ptrs...); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
m := make(map[string]any, len(colNames))
|
||||
for i, c := range colNames {
|
||||
m[c] = normalizeSQLValue(raw[i])
|
||||
}
|
||||
out = append(out, m)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
func normalizeSQLValue(v any) any {
|
||||
switch x := v.(type) {
|
||||
case nil:
|
||||
return nil
|
||||
case []byte:
|
||||
return string(x)
|
||||
default:
|
||||
return x
|
||||
}
|
||||
}
|
||||
|
||||
func dsnFileHint(dsn string) string {
|
||||
dsn = strings.TrimSpace(dsn)
|
||||
if i := strings.LastIndexAny(dsn, `/\`); i >= 0 {
|
||||
rest := dsn[i+1:]
|
||||
if j := strings.IndexAny(rest, "?#"); j >= 0 {
|
||||
rest = rest[:j]
|
||||
}
|
||||
if strings.HasSuffix(strings.ToLower(rest), ".db") {
|
||||
return rest
|
||||
}
|
||||
}
|
||||
if len(dsn) > 64 {
|
||||
return dsn[:64] + "…"
|
||||
}
|
||||
return dsn
|
||||
}
|
||||
34
platform/internal/dbsync/inspect_test.go
Normal file
34
platform/internal/dbsync/inspect_test.go
Normal file
@@ -0,0 +1,34 @@
|
||||
package dbsync
|
||||
|
||||
import (
|
||||
"context"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestInspectAndPreviewSQLite(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
dsn := "file:" + filepath.ToSlash(filepath.Join(dir, "inspect.db"))
|
||||
db, err := Open(DriverSQLite, dsn)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := db.Exec(`CREATE TABLE orders (id TEXT PRIMARY KEY, title TEXT); INSERT INTO orders VALUES ('a','one'),('b','two')`); err != nil {
|
||||
_ = db.Close()
|
||||
t.Fatal(err)
|
||||
}
|
||||
_ = db.Close()
|
||||
|
||||
ep := Endpoint{Driver: DriverSQLite, DSN: dsn}
|
||||
ins, err := InspectEndpoint(context.Background(), ep, "remote", false)
|
||||
if err != nil || !ins.OK {
|
||||
t.Fatalf("inspect: %+v %v", ins, err)
|
||||
}
|
||||
if len(ins.Tables) != 1 || ins.Tables[0].Name != "orders" || ins.Tables[0].RowCount != 2 || ins.Tables[0].ColumnCount != 2 {
|
||||
t.Fatalf("tables=%+v", ins.Tables)
|
||||
}
|
||||
prev, err := PreviewTable(context.Background(), ep, "orders", 10)
|
||||
if err != nil || !prev.OK || prev.Total != 2 || len(prev.Rows) != 2 {
|
||||
t.Fatalf("preview: %+v %v", prev, err)
|
||||
}
|
||||
}
|
||||
@@ -55,6 +55,12 @@ func (m *Manager) StartChannel(id string) error {
|
||||
c.Enabled = true
|
||||
c.LastError = ""
|
||||
})
|
||||
// 预热 remote,避免首个 agent push 冷开 SQLite 触发超时
|
||||
go func(driver Driver, dsn string) {
|
||||
if err := WarmRemote(driver, dsn); err != nil {
|
||||
log.Printf("dbsync warm remote %s: %v", id, err)
|
||||
}
|
||||
}(ch.Remote.Driver, ch.Remote.DSN)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -65,6 +71,9 @@ func (m *Manager) StopChannel(id string) {
|
||||
cancel()
|
||||
delete(m.runners, id)
|
||||
}
|
||||
if ch, err := m.store.GetChannel(id); err == nil {
|
||||
InvalidateRemote(ch.Remote.Driver, ch.Remote.DSN)
|
||||
}
|
||||
_ = m.store.PatchStats(id, func(c *Channel) { c.Enabled = false })
|
||||
}
|
||||
|
||||
|
||||
85
platform/internal/dbsync/pool.go
Normal file
85
platform/internal/dbsync/pool.go
Normal file
@@ -0,0 +1,85 @@
|
||||
package dbsync
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// remotePool 复用线上 A 连接,避免每次 push 冷开 SQLite(Windows 上常 >3s → 旧默认超时 503)。
|
||||
type remotePool struct {
|
||||
mu sync.Mutex
|
||||
dbs map[string]*pooledDB
|
||||
}
|
||||
|
||||
type pooledDB struct {
|
||||
db *sql.DB
|
||||
driver Driver
|
||||
lastUsed time.Time
|
||||
}
|
||||
|
||||
var sharedRemotePool = &remotePool{dbs: map[string]*pooledDB{}}
|
||||
|
||||
func poolKey(driver Driver, dsn string) string {
|
||||
if driver == DriverSQLite {
|
||||
dsn = normalizeSQLiteDSN(dsn)
|
||||
}
|
||||
return string(driver) + "|" + dsn
|
||||
}
|
||||
|
||||
// AcquireRemote 获取(或新建)remote 连接;调用方不要 Close。
|
||||
func AcquireRemote(driver Driver, dsn string) (*sql.DB, error) {
|
||||
return sharedRemotePool.acquire(driver, dsn)
|
||||
}
|
||||
|
||||
func (p *remotePool) acquire(driver Driver, dsn string) (*sql.DB, error) {
|
||||
key := poolKey(driver, dsn)
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
if e, ok := p.dbs[key]; ok && e.db != nil {
|
||||
if err := e.db.Ping(); err == nil {
|
||||
e.lastUsed = time.Now()
|
||||
return e.db, nil
|
||||
}
|
||||
_ = e.db.Close()
|
||||
delete(p.dbs, key)
|
||||
}
|
||||
db, err := Open(driver, dsn)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
p.dbs[key] = &pooledDB{db: db, driver: driver, lastUsed: time.Now()}
|
||||
return db, nil
|
||||
}
|
||||
|
||||
// WarmRemote 通道启动时预热,把冷开成本挪出请求路径。
|
||||
func WarmRemote(driver Driver, dsn string) error {
|
||||
if strings.TrimSpace(dsn) == "" {
|
||||
return fmt.Errorf("dsn required")
|
||||
}
|
||||
_, err := AcquireRemote(driver, dsn)
|
||||
return err
|
||||
}
|
||||
|
||||
// InvalidateRemote 通道 DSN 变更或停用时丢掉缓存连接。
|
||||
func InvalidateRemote(driver Driver, dsn string) {
|
||||
key := poolKey(driver, dsn)
|
||||
sharedRemotePool.mu.Lock()
|
||||
defer sharedRemotePool.mu.Unlock()
|
||||
if e, ok := sharedRemotePool.dbs[key]; ok {
|
||||
_ = e.db.Close()
|
||||
delete(sharedRemotePool.dbs, key)
|
||||
}
|
||||
}
|
||||
|
||||
// CloseAllRemotes 测试或进程退出时关闭池内连接。
|
||||
func CloseAllRemotes() {
|
||||
sharedRemotePool.mu.Lock()
|
||||
defer sharedRemotePool.mu.Unlock()
|
||||
for k, e := range sharedRemotePool.dbs {
|
||||
_ = e.db.Close()
|
||||
delete(sharedRemotePool.dbs, k)
|
||||
}
|
||||
}
|
||||
254
platform/internal/dbsync/pull.go
Normal file
254
platform/internal/dbsync/pull.go
Normal file
@@ -0,0 +1,254 @@
|
||||
package dbsync
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
PullModeBootstrap = "bootstrap" // 全量分页灌库
|
||||
PullModeRows = "rows" // 按 row_pks 取行
|
||||
PullModePKs = "pks" // 只列主键(客户端本地 diff)
|
||||
)
|
||||
|
||||
const DefaultPullLimit = 200
|
||||
const MaxPullLimit = 500
|
||||
|
||||
// PullRequest 形态 B:平台只读线上 A,把行/主键返回给本机 agent 写入 B。
|
||||
type PullRequest struct {
|
||||
Mode string `json:"mode"` // bootstrap | rows | pks;空则 bootstrap
|
||||
Table string `json:"table"`
|
||||
RowPKs []string `json:"row_pks"`
|
||||
AfterPK string `json:"after_pk"` // 分页游标(字典序)
|
||||
Limit int `json:"limit"`
|
||||
OnlineDBID string `json:"online_db_id"`
|
||||
}
|
||||
|
||||
// PullItem 下行一条(客户端对本机 B upsert,建议 WithApplying 防回声)。
|
||||
type PullItem struct {
|
||||
Table string `json:"table"`
|
||||
Op string `json:"op"` // upsert
|
||||
RowPK string `json:"row_pk"`
|
||||
Row map[string]any `json:"row,omitempty"`
|
||||
Version int64 `json:"version"`
|
||||
}
|
||||
|
||||
// PullResult 下行响应。
|
||||
type PullResult struct {
|
||||
OK bool `json:"ok"`
|
||||
Mode string `json:"mode"`
|
||||
Table string `json:"table"`
|
||||
PKColumn string `json:"pk_column"`
|
||||
Columns []string `json:"columns,omitempty"` // 表结构;空表时也返回,便于本机建空表
|
||||
Items []PullItem `json:"items,omitempty"`
|
||||
PKs []string `json:"pks,omitempty"`
|
||||
NextAfterPK string `json:"next_after_pk,omitempty"`
|
||||
HasMore bool `json:"has_more"`
|
||||
Message string `json:"message,omitempty"`
|
||||
}
|
||||
|
||||
// PullFromRemote 从通道 remote(线上 A)读出数据,供本机 agent 写入 B。
|
||||
func PullFromRemote(ctx context.Context, ch *Channel, store *FileStore, req PullRequest) (*PullResult, error) {
|
||||
if ch == nil {
|
||||
return nil, fmt.Errorf("channel is nil")
|
||||
}
|
||||
mode := strings.ToLower(strings.TrimSpace(req.Mode))
|
||||
if mode == "" {
|
||||
mode = PullModeBootstrap
|
||||
}
|
||||
switch mode {
|
||||
case PullModeBootstrap, PullModeRows, PullModePKs:
|
||||
default:
|
||||
return nil, fmt.Errorf("unsupported mode: %s (use bootstrap|rows|pks)", req.Mode)
|
||||
}
|
||||
|
||||
table := strings.TrimSpace(req.Table)
|
||||
if table == "" {
|
||||
return nil, fmt.Errorf("table required")
|
||||
}
|
||||
if !tableInChannel(ch, table) {
|
||||
return nil, fmt.Errorf("table %s 不在通道白名单", table)
|
||||
}
|
||||
pkCol := "id"
|
||||
if ch.PKColumns != nil && strings.TrimSpace(ch.PKColumns[table]) != "" {
|
||||
pkCol = ch.PKColumns[table]
|
||||
}
|
||||
|
||||
limit := req.Limit
|
||||
if limit <= 0 {
|
||||
limit = DefaultPullLimit
|
||||
}
|
||||
if limit > MaxPullLimit {
|
||||
limit = MaxPullLimit
|
||||
}
|
||||
|
||||
db, err := AcquireRemote(ch.Remote.Driver, ch.Remote.DSN)
|
||||
if err != nil {
|
||||
return nil, wrapOpenRemote(err)
|
||||
}
|
||||
|
||||
if err := EnsureMeta(ctx, db, ch.Remote.Driver); err != nil {
|
||||
return nil, Retryablef("ensure meta: %v", err)
|
||||
}
|
||||
|
||||
exists, err := tableExists(ctx, db, ch.Remote.Driver, table)
|
||||
if err != nil {
|
||||
return nil, Retryablef("table exists: %v", err)
|
||||
}
|
||||
if !exists {
|
||||
return nil, fmt.Errorf("table %s not found on online A; ensure-schema or push first", table)
|
||||
}
|
||||
cols, err := listColumns(ctx, db, ch.Remote.Driver, table)
|
||||
if err != nil {
|
||||
return nil, Retryablef("columns: %v", err)
|
||||
}
|
||||
|
||||
res := &PullResult{
|
||||
OK: true,
|
||||
Mode: mode,
|
||||
Table: table,
|
||||
PKColumn: pkCol,
|
||||
Columns: cols,
|
||||
}
|
||||
|
||||
switch mode {
|
||||
case PullModePKs:
|
||||
pks, next, more, err := listPKsPage(ctx, db, ch.Remote.Driver, table, pkCol, strings.TrimSpace(req.AfterPK), limit)
|
||||
if err != nil {
|
||||
return nil, Retryablef("list pks: %v", err)
|
||||
}
|
||||
res.PKs = pks
|
||||
res.NextAfterPK = next
|
||||
res.HasMore = more
|
||||
res.Message = "pks from online A; client diffs then mode=rows"
|
||||
case PullModeRows:
|
||||
if len(req.RowPKs) == 0 {
|
||||
return nil, fmt.Errorf("row_pks required for mode=rows")
|
||||
}
|
||||
if len(req.RowPKs) > MaxPullLimit {
|
||||
return nil, fmt.Errorf("row_pks limit %d", MaxPullLimit)
|
||||
}
|
||||
items, err := fetchRowsByPKs(ctx, db, ch.Remote.Driver, table, pkCol, req.RowPKs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
res.Items = items
|
||||
res.Message = "rows from online A; apply to local B with applying flag"
|
||||
default: // bootstrap
|
||||
items, next, more, err := fetchRowsPage(ctx, db, ch.Remote.Driver, table, pkCol, strings.TrimSpace(req.AfterPK), limit)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
res.Items = items
|
||||
res.NextAfterPK = next
|
||||
res.HasMore = more
|
||||
if len(items) == 0 && afterEmpty(req.AfterPK) {
|
||||
res.Message = "empty table on online A; use columns to CREATE TABLE IF NOT EXISTS on local B"
|
||||
} else {
|
||||
res.Message = "bootstrap page from online A; loop until has_more=false; create local table from columns if missing"
|
||||
}
|
||||
}
|
||||
|
||||
if store != nil && (len(res.Items) > 0 || len(res.PKs) > 0) {
|
||||
n := int64(len(res.Items))
|
||||
if n == 0 {
|
||||
n = int64(len(res.PKs))
|
||||
}
|
||||
_ = store.PatchStats(ch.ID, func(c *Channel) {
|
||||
c.Stats.PulledOK += n
|
||||
})
|
||||
}
|
||||
return res, nil
|
||||
}
|
||||
|
||||
func listPKsPage(ctx context.Context, db *sql.DB, driver Driver, table, pkCol, afterPK string, limit int) (pks []string, next string, more bool, err error) {
|
||||
q, args := buildPKPageQuery(driver, table, pkCol, afterPK, limit+1)
|
||||
rows, err := db.QueryContext(ctx, q, args...)
|
||||
if err != nil {
|
||||
return nil, "", false, err
|
||||
}
|
||||
defer rows.Close()
|
||||
for rows.Next() {
|
||||
var v any
|
||||
if err := rows.Scan(&v); err != nil {
|
||||
return nil, "", false, err
|
||||
}
|
||||
pks = append(pks, fmt.Sprint(v))
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, "", false, err
|
||||
}
|
||||
if len(pks) > limit {
|
||||
more = true
|
||||
pks = pks[:limit]
|
||||
}
|
||||
if len(pks) > 0 {
|
||||
next = pks[len(pks)-1]
|
||||
}
|
||||
return pks, next, more, nil
|
||||
}
|
||||
|
||||
func fetchRowsPage(ctx context.Context, db *sql.DB, driver Driver, table, pkCol, afterPK string, limit int) (items []PullItem, next string, more bool, err error) {
|
||||
pks, next, more, err := listPKsPage(ctx, db, driver, table, pkCol, afterPK, limit)
|
||||
if err != nil {
|
||||
return nil, "", false, Retryablef("list pks: %v", err)
|
||||
}
|
||||
items, err = fetchRowsByPKs(ctx, db, driver, table, pkCol, pks)
|
||||
if err != nil {
|
||||
return nil, "", false, err
|
||||
}
|
||||
return items, next, more, nil
|
||||
}
|
||||
|
||||
func fetchRowsByPKs(ctx context.Context, db *sql.DB, driver Driver, table, pkCol string, pks []string) ([]PullItem, error) {
|
||||
items := make([]PullItem, 0, len(pks))
|
||||
for _, pk := range pks {
|
||||
pk = strings.TrimSpace(pk)
|
||||
if pk == "" {
|
||||
continue
|
||||
}
|
||||
_, row, err := FetchRowJSON(ctx, db, driver, table, pkCol, pk)
|
||||
if err == sql.ErrNoRows {
|
||||
continue
|
||||
}
|
||||
if err != nil {
|
||||
return nil, Retryablef("fetch row: %v", err)
|
||||
}
|
||||
ver, has, _ := GetMetaVersion(ctx, db, driver, table, pk)
|
||||
if !has || ver <= 0 {
|
||||
ver = time.Now().UnixNano()
|
||||
}
|
||||
items = append(items, PullItem{
|
||||
Table: table,
|
||||
Op: "upsert",
|
||||
RowPK: pk,
|
||||
Row: row,
|
||||
Version: ver,
|
||||
})
|
||||
}
|
||||
return items, nil
|
||||
}
|
||||
|
||||
func afterEmpty(afterPK string) bool {
|
||||
return strings.TrimSpace(afterPK) == ""
|
||||
}
|
||||
|
||||
func buildPKPageQuery(driver Driver, table, pkCol, afterPK string, limit int) (string, []any) {
|
||||
t := quoteIdent(driver, table)
|
||||
p := quoteIdent(driver, pkCol)
|
||||
switch driver {
|
||||
case DriverPostgres:
|
||||
if afterPK != "" {
|
||||
return fmt.Sprintf(`SELECT %s FROM %s WHERE %s > $1 ORDER BY %s ASC LIMIT $2`, p, t, p, p), []any{afterPK, limit}
|
||||
}
|
||||
return fmt.Sprintf(`SELECT %s FROM %s ORDER BY %s ASC LIMIT $1`, p, t, p), []any{limit}
|
||||
default:
|
||||
if afterPK != "" {
|
||||
return fmt.Sprintf(`SELECT %s FROM %s WHERE %s > ? ORDER BY %s ASC LIMIT ?`, p, t, p, p), []any{afterPK, limit}
|
||||
}
|
||||
return fmt.Sprintf(`SELECT %s FROM %s ORDER BY %s ASC LIMIT ?`, p, t, p), []any{limit}
|
||||
}
|
||||
}
|
||||
73
platform/internal/dbsync/pull_test.go
Normal file
73
platform/internal/dbsync/pull_test.go
Normal file
@@ -0,0 +1,73 @@
|
||||
package dbsync
|
||||
|
||||
import (
|
||||
"context"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestPullFromRemoteBootstrapAndRows(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
dbPath := filepath.Join(dir, "remote.db")
|
||||
dsn := "file:" + filepath.ToSlash(dbPath) + "?_pragma=foreign_keys(1)"
|
||||
|
||||
db, err := Open(DriverSQLite, dsn)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer db.Close()
|
||||
ctx := context.Background()
|
||||
if _, err := db.ExecContext(ctx, `CREATE TABLE orders (id TEXT PRIMARY KEY, title TEXT)`); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := EnsureMeta(ctx, db, DriverSQLite); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, id := range []string{"a-1", "a-2", "b-1"} {
|
||||
if _, err := db.ExecContext(ctx, `INSERT INTO orders(id,title) VALUES(?,?)`, id, "t-"+id); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_ = UpsertMeta(ctx, db, DriverSQLite, "orders", id, time.Now().UnixNano())
|
||||
}
|
||||
_ = db.Close()
|
||||
|
||||
st, err := NewFileStore(filepath.Join(dir, "dbsync"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ch := &Channel{
|
||||
ID: "ch1",
|
||||
TenantID: 1,
|
||||
Remote: Endpoint{
|
||||
Driver: DriverSQLite,
|
||||
DSN: dsn,
|
||||
Tables: []string{"orders"},
|
||||
},
|
||||
PKColumns: map[string]string{"orders": "id"},
|
||||
}
|
||||
|
||||
boot, err := PullFromRemote(ctx, ch, st, PullRequest{Mode: PullModeBootstrap, Table: "orders", Limit: 2})
|
||||
if err != nil || !boot.OK || len(boot.Items) != 2 || !boot.HasMore {
|
||||
t.Fatalf("bootstrap page1: %+v err=%v", boot, err)
|
||||
}
|
||||
boot2, err := PullFromRemote(ctx, ch, st, PullRequest{
|
||||
Mode: PullModeBootstrap, Table: "orders", Limit: 2, AfterPK: boot.NextAfterPK,
|
||||
})
|
||||
if err != nil || len(boot2.Items) != 1 || boot2.HasMore {
|
||||
t.Fatalf("bootstrap page2: %+v err=%v", boot2, err)
|
||||
}
|
||||
|
||||
pks, err := PullFromRemote(ctx, ch, st, PullRequest{Mode: PullModePKs, Table: "orders", Limit: 10})
|
||||
if err != nil || len(pks.PKs) != 3 {
|
||||
t.Fatalf("pks: %+v err=%v", pks, err)
|
||||
}
|
||||
|
||||
rows, err := PullFromRemote(ctx, ch, st, PullRequest{
|
||||
Mode: PullModeRows, Table: "orders", RowPKs: []string{"a-2", "missing"},
|
||||
})
|
||||
if err != nil || len(rows.Items) != 1 || rows.Items[0].RowPK != "a-2" {
|
||||
t.Fatalf("rows: %+v err=%v", rows, err)
|
||||
}
|
||||
CloseAllRemotes()
|
||||
}
|
||||
@@ -40,9 +40,7 @@ func PushToRemote(ctx context.Context, ch *Channel, store *FileStore, item PushI
|
||||
if table == "" {
|
||||
return nil, fmt.Errorf("table required")
|
||||
}
|
||||
if !tableInChannel(ch, table) {
|
||||
return nil, fmt.Errorf("table %s 不在通道白名单", table)
|
||||
}
|
||||
// Z4 冻结:通道表白名单不再作为拒收依据;Binding + JWT 才是权限源。
|
||||
pkCol := "id"
|
||||
if ch.PKColumns != nil && strings.TrimSpace(ch.PKColumns[table]) != "" {
|
||||
pkCol = ch.PKColumns[table]
|
||||
@@ -57,14 +55,21 @@ func PushToRemote(ctx context.Context, ch *Channel, store *FileStore, item PushI
|
||||
version = time.Now().UnixNano()
|
||||
}
|
||||
|
||||
db, err := Open(ch.Remote.Driver, ch.Remote.DSN)
|
||||
db, err := AcquireRemote(ch.Remote.Driver, ch.Remote.DSN)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("open remote: %w", err)
|
||||
return nil, wrapOpenRemote(err)
|
||||
}
|
||||
defer db.Close()
|
||||
|
||||
if err := EnsureMeta(ctx, db, ch.Remote.Driver); err != nil {
|
||||
return nil, fmt.Errorf("ensure meta: %w", err)
|
||||
return nil, Retryablef("ensure meta: %v", err)
|
||||
}
|
||||
|
||||
if op != "delete" {
|
||||
var row map[string]any
|
||||
_ = json.Unmarshal([]byte(payload), &row)
|
||||
if err := EnsureTableFromRow(ctx, db, ch.Remote.Driver, table, pkCol, row); err != nil {
|
||||
return nil, Retryablef("ensure table: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
res := &PushResult{
|
||||
@@ -74,13 +79,16 @@ func PushToRemote(ctx context.Context, ch *Channel, store *FileStore, item PushI
|
||||
|
||||
tgtVer, has, err := GetMetaVersion(ctx, db, ch.Remote.Driver, table, rowPK)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil, Retryablef("meta version: %v", err)
|
||||
}
|
||||
// 幂等:同 version 已落地 → 跳过
|
||||
if has && tgtVer == version {
|
||||
res.OK = true
|
||||
res.Skipped = true
|
||||
res.Message = "already applied (same version)"
|
||||
if store != nil {
|
||||
_ = store.PatchStats(ch.ID, func(c *Channel) { c.Stats.PushedSkipped++ })
|
||||
}
|
||||
return res, nil
|
||||
}
|
||||
if has && tgtVer > version {
|
||||
@@ -109,11 +117,14 @@ func PushToRemote(ctx context.Context, ch *Channel, store *FileStore, item PushI
|
||||
res.OK = true
|
||||
res.Skipped = true
|
||||
res.Message = "target newer; kept (lww_target)"
|
||||
if store != nil {
|
||||
_ = store.PatchStats(ch.ID, func(c *Channel) { c.Stats.PushedSkipped++ })
|
||||
}
|
||||
return res, nil
|
||||
case PolicyLWWSource:
|
||||
loser := SnapshotTargetRow(ctx, db, ch.Remote.Driver, table, pkCol, rowPK)
|
||||
if err := ApplyChange(ctx, db, ch.Remote.Driver, table, pkCol, op, payload, version); err != nil {
|
||||
return nil, err
|
||||
return nil, Retryablef("apply: %v", err)
|
||||
}
|
||||
RecordLwwOverride(store, LwwOverride{
|
||||
TenantID: ch.TenantID,
|
||||
@@ -130,7 +141,10 @@ func PushToRemote(ctx context.Context, ch *Channel, store *FileStore, item PushI
|
||||
SourceVer: version,
|
||||
})
|
||||
if store != nil {
|
||||
_ = store.PatchStats(ch.ID, func(c *Channel) { c.Stats.PushedOK++ })
|
||||
_ = store.PatchStats(ch.ID, func(c *Channel) {
|
||||
c.Stats.PushedOK++
|
||||
c.Stats.PushedApplied++
|
||||
})
|
||||
}
|
||||
res.OK = true
|
||||
res.Applied = true
|
||||
@@ -150,7 +164,10 @@ func PushToRemote(ctx context.Context, ch *Channel, store *FileStore, item PushI
|
||||
SourceVer: version,
|
||||
Message: "target version newer than agent push",
|
||||
})
|
||||
_ = store.PatchStats(ch.ID, func(c *Channel) { c.Stats.Conflicts++ })
|
||||
_ = store.PatchStats(ch.ID, func(c *Channel) {
|
||||
c.Stats.Conflicts++
|
||||
c.Stats.PushedSkipped++
|
||||
})
|
||||
}
|
||||
res.OK = true
|
||||
res.Skipped = true
|
||||
@@ -161,10 +178,13 @@ func PushToRemote(ctx context.Context, ch *Channel, store *FileStore, item PushI
|
||||
}
|
||||
|
||||
if err := ApplyChange(ctx, db, ch.Remote.Driver, table, pkCol, op, payload, version); err != nil {
|
||||
return nil, err
|
||||
return nil, Retryablef("apply: %v", err)
|
||||
}
|
||||
if store != nil {
|
||||
_ = store.PatchStats(ch.ID, func(c *Channel) { c.Stats.PushedOK++ })
|
||||
_ = store.PatchStats(ch.ID, func(c *Channel) {
|
||||
c.Stats.PushedOK++
|
||||
c.Stats.PushedApplied++
|
||||
})
|
||||
}
|
||||
res.OK = true
|
||||
res.Applied = true
|
||||
@@ -192,12 +212,9 @@ func PushBatchToRemote(ctx context.Context, ch *Channel, store *FileStore, items
|
||||
}
|
||||
|
||||
func tableInChannel(ch *Channel, table string) bool {
|
||||
for _, t := range uniqueTables(ch.Local.Tables, ch.Remote.Tables) {
|
||||
if strings.EqualFold(strings.TrimSpace(t), table) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
// Z4 冻结:通道 tables 仅历史兼容展示,push/pull 不再据此拒收。
|
||||
_ = ch
|
||||
return strings.TrimSpace(table) != ""
|
||||
}
|
||||
|
||||
func normalizePushPayload(item PushItem, pkCol string) (op string, payload string, rowPK string, err error) {
|
||||
|
||||
@@ -37,10 +37,11 @@ func TestTableInChannel(t *testing.T) {
|
||||
Local: Endpoint{Tables: []string{"orders"}},
|
||||
Remote: Endpoint{Tables: []string{"orders"}},
|
||||
}
|
||||
if !tableInChannel(ch, "orders") {
|
||||
t.Fatal("expected in")
|
||||
// Z4:通道表白名单不再拒收
|
||||
if !tableInChannel(ch, "orders") || !tableInChannel(ch, "other") {
|
||||
t.Fatal("expected any non-empty table accepted")
|
||||
}
|
||||
if tableInChannel(ch, "other") {
|
||||
t.Fatal("expected out")
|
||||
if tableInChannel(ch, "") {
|
||||
t.Fatal("empty table rejected")
|
||||
}
|
||||
}
|
||||
|
||||
51
platform/internal/dbsync/retryable.go
Normal file
51
platform/internal/dbsync/retryable.go
Normal file
@@ -0,0 +1,51 @@
|
||||
package dbsync
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// RetryableError 上游 IO / remote 暂不可达,agent 可稍后重试(不换 UUID)。
|
||||
type RetryableError struct {
|
||||
Msg string
|
||||
}
|
||||
|
||||
func (e *RetryableError) Error() string {
|
||||
if e == nil {
|
||||
return "retryable error"
|
||||
}
|
||||
return e.Msg
|
||||
}
|
||||
|
||||
func Retryable(msg string) error {
|
||||
return &RetryableError{Msg: msg}
|
||||
}
|
||||
|
||||
func Retryablef(format string, args ...any) error {
|
||||
return &RetryableError{Msg: fmt.Sprintf(format, args...)}
|
||||
}
|
||||
|
||||
func IsRetryable(err error) bool {
|
||||
var r *RetryableError
|
||||
return errors.As(err, &r)
|
||||
}
|
||||
|
||||
func wrapOpenRemote(err error) error {
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
msg := err.Error()
|
||||
lower := strings.ToLower(msg)
|
||||
if strings.Contains(lower, "open remote") ||
|
||||
strings.Contains(lower, "unable to open") ||
|
||||
strings.Contains(lower, "no such file") ||
|
||||
strings.Contains(lower, "locked") ||
|
||||
strings.Contains(lower, "busy") ||
|
||||
strings.Contains(lower, "connection refused") ||
|
||||
strings.Contains(lower, "timeout") ||
|
||||
strings.Contains(lower, "i/o") {
|
||||
return Retryablef("open remote: %v", err)
|
||||
}
|
||||
return Retryablef("open remote: %v", err)
|
||||
}
|
||||
184
platform/internal/dbsync/schema.go
Normal file
184
platform/internal/dbsync/schema.go
Normal file
@@ -0,0 +1,184 @@
|
||||
package dbsync
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// SchemaTableSpec 本机 → 线上:空表/表结构 ensure 请求中的一张表。
|
||||
type SchemaTableSpec struct {
|
||||
Table string `json:"table"`
|
||||
PKColumn string `json:"pk_column,omitempty"`
|
||||
Columns []string `json:"columns"`
|
||||
OnlineDBID string `json:"online_db_id,omitempty"`
|
||||
}
|
||||
|
||||
// SchemaEnsureRequest agent 批量 ensure 空表到线上 A。
|
||||
type SchemaEnsureRequest struct {
|
||||
OnlineDBID string `json:"online_db_id"`
|
||||
Tables []SchemaTableSpec `json:"tables"`
|
||||
}
|
||||
|
||||
// SchemaEnsureItemResult 单表 ensure 结果。
|
||||
type SchemaEnsureItemResult struct {
|
||||
Table string `json:"table"`
|
||||
OK bool `json:"ok"`
|
||||
Created bool `json:"created"` // true=本次新建;false=已存在跳过
|
||||
Message string `json:"message,omitempty"`
|
||||
}
|
||||
|
||||
// SchemaEnsureResult 批量 ensure 响应。
|
||||
type SchemaEnsureResult struct {
|
||||
OK bool `json:"ok"`
|
||||
Results []SchemaEnsureItemResult `json:"results"`
|
||||
Message string `json:"message,omitempty"`
|
||||
}
|
||||
|
||||
// SchemaDescribeRequest 从线上 A 拉取表结构(供本机建空表)。
|
||||
type SchemaDescribeRequest struct {
|
||||
OnlineDBID string `json:"online_db_id"`
|
||||
Tables []string `json:"tables,omitempty"` // 空=全部业务表
|
||||
}
|
||||
|
||||
// SchemaTableDesc 线上表结构描述。
|
||||
type SchemaTableDesc struct {
|
||||
Name string `json:"name"`
|
||||
PKColumn string `json:"pk_column"`
|
||||
Columns []string `json:"columns"`
|
||||
RowCount int64 `json:"row_count"`
|
||||
}
|
||||
|
||||
// SchemaDescribeResult 表结构列表。
|
||||
type SchemaDescribeResult struct {
|
||||
OK bool `json:"ok"`
|
||||
Tables []SchemaTableDesc `json:"tables"`
|
||||
Message string `json:"message,omitempty"`
|
||||
}
|
||||
|
||||
const MaxSchemaEnsureTables = 100
|
||||
|
||||
// EnsureSchemasOnRemote 在线上 A 为每张表 CREATE IF NOT EXISTS(空表也建)。
|
||||
func EnsureSchemasOnRemote(ctx context.Context, ch *Channel, req SchemaEnsureRequest) (*SchemaEnsureResult, error) {
|
||||
if ch == nil {
|
||||
return nil, fmt.Errorf("channel is nil")
|
||||
}
|
||||
if len(req.Tables) == 0 {
|
||||
return nil, fmt.Errorf("tables required")
|
||||
}
|
||||
if len(req.Tables) > MaxSchemaEnsureTables {
|
||||
return nil, fmt.Errorf("tables limit %d", MaxSchemaEnsureTables)
|
||||
}
|
||||
db, err := AcquireRemote(ch.Remote.Driver, ch.Remote.DSN)
|
||||
if err != nil {
|
||||
return nil, wrapOpenRemote(err)
|
||||
}
|
||||
if err := EnsureMeta(ctx, db, ch.Remote.Driver); err != nil {
|
||||
return nil, Retryablef("ensure meta: %v", err)
|
||||
}
|
||||
|
||||
out := &SchemaEnsureResult{OK: true, Results: make([]SchemaEnsureItemResult, 0, len(req.Tables))}
|
||||
for _, spec := range req.Tables {
|
||||
table := strings.TrimSpace(spec.Table)
|
||||
item := SchemaEnsureItemResult{Table: table}
|
||||
if table == "" {
|
||||
item.Message = "table required"
|
||||
out.OK = false
|
||||
out.Results = append(out.Results, item)
|
||||
continue
|
||||
}
|
||||
if !tableInChannel(ch, table) {
|
||||
item.Message = "table rejected"
|
||||
out.OK = false
|
||||
out.Results = append(out.Results, item)
|
||||
continue
|
||||
}
|
||||
pkCol := strings.TrimSpace(spec.PKColumn)
|
||||
if pkCol == "" && ch.PKColumns != nil && strings.TrimSpace(ch.PKColumns[table]) != "" {
|
||||
pkCol = ch.PKColumns[table]
|
||||
}
|
||||
if pkCol == "" {
|
||||
pkCol = "id"
|
||||
}
|
||||
existed, err := tableExists(ctx, db, ch.Remote.Driver, table)
|
||||
if err != nil {
|
||||
return nil, Retryablef("table exists: %v", err)
|
||||
}
|
||||
if err := EnsureTableFromColumns(ctx, db, ch.Remote.Driver, table, pkCol, spec.Columns); err != nil {
|
||||
item.Message = err.Error()
|
||||
out.OK = false
|
||||
out.Results = append(out.Results, item)
|
||||
continue
|
||||
}
|
||||
item.OK = true
|
||||
item.Created = !existed
|
||||
if existed {
|
||||
item.Message = "already exists"
|
||||
} else {
|
||||
item.Message = "created"
|
||||
}
|
||||
out.Results = append(out.Results, item)
|
||||
}
|
||||
out.Message = "empty tables ensured on online A; client should also create missing local tables from schema describe"
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// DescribeSchemasFromRemote 列出线上 A 业务表结构(含空表),供本机 CREATE IF NOT EXISTS。
|
||||
func DescribeSchemasFromRemote(ctx context.Context, ch *Channel, req SchemaDescribeRequest) (*SchemaDescribeResult, error) {
|
||||
if ch == nil {
|
||||
return nil, fmt.Errorf("channel is nil")
|
||||
}
|
||||
db, err := AcquireRemote(ch.Remote.Driver, ch.Remote.DSN)
|
||||
if err != nil {
|
||||
return nil, wrapOpenRemote(err)
|
||||
}
|
||||
if err := EnsureMeta(ctx, db, ch.Remote.Driver); err != nil {
|
||||
return nil, Retryablef("ensure meta: %v", err)
|
||||
}
|
||||
|
||||
want := map[string]struct{}{}
|
||||
for _, t := range req.Tables {
|
||||
t = strings.TrimSpace(t)
|
||||
if t != "" {
|
||||
want[t] = struct{}{}
|
||||
}
|
||||
}
|
||||
|
||||
names, err := ListTables(ctx, db, ch.Remote.Driver)
|
||||
if err != nil {
|
||||
return nil, Retryablef("list tables: %v", err)
|
||||
}
|
||||
|
||||
out := &SchemaDescribeResult{OK: true, Tables: make([]SchemaTableDesc, 0, len(names))}
|
||||
for _, name := range names {
|
||||
if len(want) > 0 {
|
||||
if _, ok := want[name]; !ok {
|
||||
continue
|
||||
}
|
||||
}
|
||||
if !tableInChannel(ch, name) {
|
||||
continue
|
||||
}
|
||||
pkCol := "id"
|
||||
if ch.PKColumns != nil && strings.TrimSpace(ch.PKColumns[name]) != "" {
|
||||
pkCol = ch.PKColumns[name]
|
||||
}
|
||||
cols, err := listColumns(ctx, db, ch.Remote.Driver, name)
|
||||
if err != nil {
|
||||
return nil, Retryablef("columns %s: %v", name, err)
|
||||
}
|
||||
n, _ := countRows(ctx, db, ch.Remote.Driver, name)
|
||||
out.Tables = append(out.Tables, SchemaTableDesc{
|
||||
Name: name,
|
||||
PKColumn: pkCol,
|
||||
Columns: cols,
|
||||
RowCount: n,
|
||||
})
|
||||
}
|
||||
if len(out.Tables) == 0 {
|
||||
out.Message = "online A has no business tables yet; push or ensure-schema first"
|
||||
} else {
|
||||
out.Message = "use columns to CREATE TABLE IF NOT EXISTS on local B (including empty tables)"
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
94
platform/internal/dbsync/schema_test.go
Normal file
94
platform/internal/dbsync/schema_test.go
Normal file
@@ -0,0 +1,94 @@
|
||||
package dbsync
|
||||
|
||||
import (
|
||||
"context"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestEnsureTableFromColumnsEmpty(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
dsn := "file:" + filepath.ToSlash(filepath.Join(dir, "empty.db")) + "?_pragma=busy_timeout(5000)"
|
||||
db, err := Open(DriverSQLite, dsn)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer db.Close()
|
||||
ctx := context.Background()
|
||||
cols := []string{"id", "username", "password", "email"}
|
||||
if err := EnsureTableFromColumns(ctx, db, DriverSQLite, "accounts", "id", cols); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var n int64
|
||||
if err := db.QueryRowContext(ctx, `SELECT COUNT(*) FROM accounts`).Scan(&n); err != nil || n != 0 {
|
||||
t.Fatalf("want empty table n=%d err=%v", n, err)
|
||||
}
|
||||
// idempotent
|
||||
if err := EnsureTableFromColumns(ctx, db, DriverSQLite, "accounts", "id", cols); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnsureSchemasOnRemoteCreatesEmpty(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
dsn := "file:" + filepath.ToSlash(filepath.Join(dir, "remote.db")) + "?_pragma=busy_timeout(5000)"
|
||||
ch := &Channel{
|
||||
ID: "ch1",
|
||||
Remote: Endpoint{
|
||||
Driver: DriverSQLite,
|
||||
DSN: dsn,
|
||||
},
|
||||
}
|
||||
ctx := context.Background()
|
||||
defer InvalidateRemote(DriverSQLite, dsn)
|
||||
res, err := EnsureSchemasOnRemote(ctx, ch, SchemaEnsureRequest{
|
||||
Tables: []SchemaTableSpec{{
|
||||
Table: "accounts",
|
||||
PKColumn: "id",
|
||||
Columns: []string{"id", "username", "email"},
|
||||
}},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !res.OK || len(res.Results) != 1 || !res.Results[0].Created {
|
||||
t.Fatalf("unexpected: %+v", res)
|
||||
}
|
||||
desc, err := DescribeSchemasFromRemote(ctx, ch, SchemaDescribeRequest{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(desc.Tables) != 1 || desc.Tables[0].Name != "accounts" || desc.Tables[0].RowCount != 0 {
|
||||
t.Fatalf("describe: %+v", desc)
|
||||
}
|
||||
if len(desc.Tables[0].Columns) < 2 {
|
||||
t.Fatalf("columns: %v", desc.Tables[0].Columns)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPullEmptyTableReturnsColumns(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
dsn := "file:" + filepath.ToSlash(filepath.Join(dir, "pull.db")) + "?_pragma=busy_timeout(5000)"
|
||||
db, err := Open(DriverSQLite, dsn)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ctx := context.Background()
|
||||
if err := EnsureTableFromColumns(ctx, db, DriverSQLite, "accounts", "id", []string{"id", "username"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_ = db.Close()
|
||||
|
||||
ch := &Channel{ID: "ch1", Remote: Endpoint{Driver: DriverSQLite, DSN: dsn}}
|
||||
defer InvalidateRemote(DriverSQLite, dsn)
|
||||
res, err := PullFromRemote(ctx, ch, nil, PullRequest{Mode: PullModeBootstrap, Table: "accounts"})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(res.Items) != 0 {
|
||||
t.Fatalf("want empty items: %+v", res.Items)
|
||||
}
|
||||
if len(res.Columns) < 2 {
|
||||
t.Fatalf("want columns for empty table: %+v", res)
|
||||
}
|
||||
}
|
||||
@@ -48,20 +48,25 @@ type Channel struct {
|
||||
Local Endpoint `json:"local"`
|
||||
Remote Endpoint `json:"remote"`
|
||||
PKColumns map[string]string `json:"pk_columns"` // table -> pk col,默认 id
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
LastError string `json:"last_error,omitempty"`
|
||||
LastSyncAt *time.Time `json:"last_sync_at,omitempty"`
|
||||
LastReconcileAt *time.Time `json:"last_reconcile_at,omitempty"`
|
||||
Stats ChannelStats `json:"stats"`
|
||||
// AgentID / AppSlug:把通道挂到某个智能体及其模块,便于「模块数据进该智能体库」对照。
|
||||
AgentID int64 `json:"agent_id,omitempty"`
|
||||
AppSlug string `json:"app_slug,omitempty"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
LastError string `json:"last_error,omitempty"`
|
||||
LastSyncAt *time.Time `json:"last_sync_at,omitempty"`
|
||||
LastReconcileAt *time.Time `json:"last_reconcile_at,omitempty"`
|
||||
Stats ChannelStats `json:"stats"`
|
||||
}
|
||||
|
||||
type ChannelStats struct {
|
||||
PushedOK int64 `json:"pushed_ok"`
|
||||
PulledOK int64 `json:"pulled_ok"`
|
||||
Conflicts int64 `json:"conflicts"`
|
||||
Retries int64 `json:"retries"`
|
||||
LastBatch int `json:"last_batch"`
|
||||
PushedOK int64 `json:"pushed_ok"` // 兼容:成功 apply 次数
|
||||
PushedApplied int64 `json:"pushed_applied,omitempty"` // agent/drain 实际写入
|
||||
PushedSkipped int64 `json:"pushed_skipped,omitempty"` // 同 version 幂等跳过
|
||||
PulledOK int64 `json:"pulled_ok"`
|
||||
Conflicts int64 `json:"conflicts"`
|
||||
Retries int64 `json:"retries"`
|
||||
LastBatch int `json:"last_batch"`
|
||||
}
|
||||
|
||||
type Conflict struct {
|
||||
|
||||
Reference in New Issue
Block a user