diff --git a/openflare_server/model/main.go b/openflare_server/model/main.go index ce651a02..60bf8306 100644 --- a/openflare_server/model/main.go +++ b/openflare_server/model/main.go @@ -320,19 +320,88 @@ func renameLegacyObservabilityShardTables(db *gorm.DB) error { if err := db.Migrator().RenameTable(table, legacyTable); err != nil { return fmt.Errorf("rename sharded table %s to %s failed: %w", table, legacyTable, err) } + if err := dropLegacyObservabilitySecondaryIndexes(db, legacyTable); err != nil { + return err + } + } + } + return nil +} + +func dropLegacyObservabilitySecondaryIndexes(db *gorm.DB, table string) error { + db = sessionIgnoringSharding(db) + if db == nil { + return fmt.Errorf("database handle is nil") + } + backend := baseDialector(db).Name() + indexes := make([]string, 0) + switch backend { + case "sqlite": + if err := db.Raw( + `SELECT name FROM sqlite_master WHERE type = 'index' AND tbl_name = ? AND name LIKE 'idx_%'`, + table, + ).Scan(&indexes).Error; err != nil { + return fmt.Errorf("list indexes for %s failed: %w", table, err) + } + case "postgres": + if err := db.Raw( + `SELECT indexname FROM pg_indexes WHERE schemaname = current_schema() AND tablename = ? AND indexname LIKE 'idx_%'`, + table, + ).Scan(&indexes).Error; err != nil { + return fmt.Errorf("list indexes for %s failed: %w", table, err) + } + default: + return fmt.Errorf("unsupported database backend %s", backend) + } + for _, indexName := range indexes { + if err := db.Exec(fmt.Sprintf(`DROP INDEX IF EXISTS "%s"`, indexName)).Error; err != nil { + return fmt.Errorf("drop legacy index %s failed: %w", indexName, err) + } + } + return nil +} + +func autoMigrateObservabilityShardTables(db *gorm.DB) error { + db = sessionIgnoringSharding(db) + if db == nil { + return fmt.Errorf("database handle is nil") + } + dialector := baseDialector(db) + if dialector == nil { + return fmt.Errorf("database dialector is nil") + } + type shardedTable struct { + model any + base string + } + tables := []shardedTable{ + {model: &NodeMetricSnapshot{}, base: "node_metric_snapshots"}, + {model: &NodeRequestReport{}, base: "node_request_reports"}, + {model: &NodeAccessLog{}, base: "node_access_logs"}, + } + for _, item := range tables { + for _, table := range observabilityShardTables(item.base) { + tx := db.Table(table) + if err := dialector.Migrator(tx).AutoMigrate(item.model); err != nil { + return fmt.Errorf("auto migrate sharded table %s failed: %w", table, err) + } } } return nil } func dropLegacyObservabilityShardTables(db *gorm.DB) error { + db = sessionIgnoringSharding(db) + if db == nil { + return fmt.Errorf("database handle is nil") + } for _, baseTable := range shardedObservabilityBaseTables() { for _, table := range observabilityShardTables(baseTable) { legacyTable := legacyObservabilityShardTableName(table) if !db.Migrator().HasTable(legacyTable) { continue } - if err := db.Migrator().DropTable(legacyTable); err != nil { + if err := db.Exec(fmt.Sprintf(`DROP TABLE IF EXISTS "%s"`, legacyTable)).Error; err != nil { return fmt.Errorf("drop legacy sharded table %s failed: %w", legacyTable, err) } } @@ -461,10 +530,11 @@ func migrateObservabilityShardsToID(db *gorm.DB, backend string) error { if db == nil { return fmt.Errorf("database handle is nil") } + _ = backend if err := renameLegacyObservabilityShardTables(db); err != nil { return err } - if err := applyCurrentSchema(db, backend); err != nil { + if err := autoMigrateObservabilityShardTables(db); err != nil { return err } if err := migrateLegacyNodeMetricSnapshots(db); err != nil { diff --git a/openflare_server/model/sharding.go b/openflare_server/model/sharding.go index 32915cb3..704ba45a 100644 --- a/openflare_server/model/sharding.go +++ b/openflare_server/model/sharding.go @@ -3,6 +3,7 @@ package model import ( "fmt" "sort" + "strconv" "strings" "sync" @@ -25,8 +26,14 @@ func registerSharding(db *gorm.DB, backend string) error { } _ = backend if err := db.Use(sharding.Register(sharding.Config{ - ShardingKey: "id", - NumberOfShards: observabilityShardCount, + ShardingKey: "id", + NumberOfShards: observabilityShardCount, + ShardingAlgorithm: func(value any) (string, error) { + return observabilityShardSuffixForValue(value) + }, + ShardingAlgorithmByPrimaryKey: func(id int64) string { + return observabilityShardSuffixForInt64(id) + }, PrimaryKeyGenerator: sharding.PKCustom, PrimaryKeyGeneratorFn: func(tableIdx int64) int64 { return 0 @@ -82,6 +89,46 @@ func observabilityShardSuffixForID(id uint) string { return fmt.Sprintf("_%02d", uint64(id)%uint64(observabilityShardCount)) } +func observabilityShardSuffixForInt64(id int64) string { + if id < 0 { + id = -id + } + return fmt.Sprintf("_%02d", uint64(id)%uint64(observabilityShardCount)) +} + +func observabilityShardSuffixForValue(value any) (string, error) { + switch typed := value.(type) { + case int: + return observabilityShardSuffixForInt64(int64(typed)), nil + case int8: + return observabilityShardSuffixForInt64(int64(typed)), nil + case int16: + return observabilityShardSuffixForInt64(int64(typed)), nil + case int32: + return observabilityShardSuffixForInt64(int64(typed)), nil + case int64: + return observabilityShardSuffixForInt64(typed), nil + case uint: + return observabilityShardSuffixForID(typed), nil + case uint8: + return observabilityShardSuffixForID(uint(typed)), nil + case uint16: + return observabilityShardSuffixForID(uint(typed)), nil + case uint32: + return observabilityShardSuffixForID(uint(typed)), nil + case uint64: + return fmt.Sprintf("_%02d", typed%uint64(observabilityShardCount)), nil + case string: + id, err := strconv.ParseUint(strings.TrimSpace(typed), 10, 64) + if err != nil { + return "", fmt.Errorf("invalid sharding id %q", typed) + } + return fmt.Sprintf("_%02d", id%uint64(observabilityShardCount)), nil + default: + return "", fmt.Errorf("unsupported observability sharding value type %T", value) + } +} + func observabilityShardTableForID(baseTable string, id uint) string { return baseTable + observabilityShardSuffixForID(id) } @@ -97,6 +144,24 @@ func normalizeShardedDB(db *gorm.DB) *gorm.DB { return DB } +func sessionIgnoringSharding(db *gorm.DB) *gorm.DB { + db = normalizeShardedDB(db) + if db == nil { + return nil + } + return db.Session(&gorm.Session{}).Set(sharding.ShardingIgnoreStoreKey, true) +} + +func baseDialector(db *gorm.DB) gorm.Dialector { + if db == nil { + return nil + } + if dialector, ok := db.Dialector.(sharding.ShardingDialector); ok { + return dialector.Dialector + } + return db.Dialector +} + func nextObservabilityID() (uint, error) { observabilityIDNodeOnce.Do(func() { observabilityIDNode, observabilityIDNodeErr = snowflake.NewNode(0)