diff --git a/openflare_server/model/database_schema_version.go b/openflare_server/model/database_schema_version.go index c8738253..62876a94 100644 --- a/openflare_server/model/database_schema_version.go +++ b/openflare_server/model/database_schema_version.go @@ -4,7 +4,7 @@ import "time" const ( legacyDatabaseSchemaVersion = 1 - currentDatabaseSchemaVersion = 2 + currentDatabaseSchemaVersion = 3 databaseSchemaVersionRowID = 1 ) diff --git a/openflare_server/model/main.go b/openflare_server/model/main.go index c1307415..ce651a02 100644 --- a/openflare_server/model/main.go +++ b/openflare_server/model/main.go @@ -292,6 +292,193 @@ func validateDatabaseSchemaV2(db *gorm.DB, backend string) error { return nil } +func validateDatabaseSchemaV3(db *gorm.DB, backend string) error { + if err := validateDatabaseSchemaV2(db, backend); err != nil { + return err + } + for _, baseTable := range shardedObservabilityBaseTables() { + for _, table := range observabilityShardTables(baseTable) { + legacyTable := legacyObservabilityShardTableName(table) + if db.Migrator().HasTable(legacyTable) { + return fmt.Errorf("legacy sharded table %s still exists", legacyTable) + } + } + } + return nil +} + +func renameLegacyObservabilityShardTables(db *gorm.DB) error { + for _, baseTable := range shardedObservabilityBaseTables() { + for _, table := range observabilityShardTables(baseTable) { + legacyTable := legacyObservabilityShardTableName(table) + if db.Migrator().HasTable(legacyTable) { + return fmt.Errorf("legacy sharded table %s already exists", legacyTable) + } + if !db.Migrator().HasTable(table) { + continue + } + if err := db.Migrator().RenameTable(table, legacyTable); err != nil { + return fmt.Errorf("rename sharded table %s to %s failed: %w", table, legacyTable, err) + } + } + } + return nil +} + +func dropLegacyObservabilityShardTables(db *gorm.DB) error { + 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 { + return fmt.Errorf("drop legacy sharded table %s failed: %w", legacyTable, err) + } + } + } + return nil +} + +func migrateLegacyNodeMetricSnapshots(db *gorm.DB) error { + for _, table := range observabilityShardTables("node_metric_snapshots") { + legacyTable := legacyObservabilityShardTableName(table) + if !db.Migrator().HasTable(legacyTable) { + continue + } + var lastSeenID uint + for { + var rows []NodeMetricSnapshot + query := db.Table(legacyTable).Order("id ASC").Limit(500) + if lastSeenID > 0 { + query = query.Where("id > ?", lastSeenID) + } + if err := query.Find(&rows).Error; err != nil { + return fmt.Errorf("query legacy sharded table %s failed: %w", legacyTable, err) + } + if len(rows) == 0 { + break + } + lastSeenID = rows[len(rows)-1].ID + grouped := make(map[string][]NodeMetricSnapshot, observabilityShardCount) + for index := range rows { + rows[index].ID = 0 + if err := assignObservabilityID(&rows[index].ID); err != nil { + return err + } + targetTable := observabilityShardTableForID("node_metric_snapshots", rows[index].ID) + grouped[targetTable] = append(grouped[targetTable], rows[index]) + } + for targetTable, batch := range grouped { + if err := db.Table(targetTable).Create(&batch).Error; err != nil { + return fmt.Errorf("write migrated rows into %s failed: %w", targetTable, err) + } + } + } + } + return nil +} + +func migrateLegacyNodeRequestReports(db *gorm.DB) error { + for _, table := range observabilityShardTables("node_request_reports") { + legacyTable := legacyObservabilityShardTableName(table) + if !db.Migrator().HasTable(legacyTable) { + continue + } + var lastSeenID uint + for { + var rows []NodeRequestReport + query := db.Table(legacyTable).Order("id ASC").Limit(500) + if lastSeenID > 0 { + query = query.Where("id > ?", lastSeenID) + } + if err := query.Find(&rows).Error; err != nil { + return fmt.Errorf("query legacy sharded table %s failed: %w", legacyTable, err) + } + if len(rows) == 0 { + break + } + lastSeenID = rows[len(rows)-1].ID + grouped := make(map[string][]NodeRequestReport, observabilityShardCount) + for index := range rows { + rows[index].ID = 0 + if err := assignObservabilityID(&rows[index].ID); err != nil { + return err + } + targetTable := observabilityShardTableForID("node_request_reports", rows[index].ID) + grouped[targetTable] = append(grouped[targetTable], rows[index]) + } + for targetTable, batch := range grouped { + if err := db.Table(targetTable).Create(&batch).Error; err != nil { + return fmt.Errorf("write migrated rows into %s failed: %w", targetTable, err) + } + } + } + } + return nil +} + +func migrateLegacyNodeAccessLogs(db *gorm.DB) error { + for _, table := range observabilityShardTables("node_access_logs") { + legacyTable := legacyObservabilityShardTableName(table) + if !db.Migrator().HasTable(legacyTable) { + continue + } + var lastSeenID uint + for { + var rows []NodeAccessLog + query := db.Table(legacyTable).Order("id ASC").Limit(500) + if lastSeenID > 0 { + query = query.Where("id > ?", lastSeenID) + } + if err := query.Find(&rows).Error; err != nil { + return fmt.Errorf("query legacy sharded table %s failed: %w", legacyTable, err) + } + if len(rows) == 0 { + break + } + lastSeenID = rows[len(rows)-1].ID + grouped := make(map[string][]NodeAccessLog, observabilityShardCount) + for index := range rows { + rows[index].ID = 0 + if err := assignObservabilityID(&rows[index].ID); err != nil { + return err + } + targetTable := observabilityShardTableForID("node_access_logs", rows[index].ID) + grouped[targetTable] = append(grouped[targetTable], rows[index]) + } + for targetTable, batch := range grouped { + if err := db.Table(targetTable).Create(&batch).Error; err != nil { + return fmt.Errorf("write migrated rows into %s failed: %w", targetTable, err) + } + } + } + } + return nil +} + +func migrateObservabilityShardsToID(db *gorm.DB, backend string) error { + if db == nil { + return fmt.Errorf("database handle is nil") + } + if err := renameLegacyObservabilityShardTables(db); err != nil { + return err + } + if err := applyCurrentSchema(db, backend); err != nil { + return err + } + if err := migrateLegacyNodeMetricSnapshots(db); err != nil { + return err + } + if err := migrateLegacyNodeRequestReports(db); err != nil { + return err + } + if err := migrateLegacyNodeAccessLogs(db); err != nil { + return err + } + return dropLegacyObservabilityShardTables(db) +} + func databaseSchemaMigrations() []databaseSchemaMigration { return []databaseSchemaMigration{ { @@ -300,6 +487,12 @@ func databaseSchemaMigrations() []databaseSchemaMigration { migrate: applyCurrentSchema, validate: validateDatabaseSchemaV2, }, + { + fromVersion: 2, + toVersion: 3, + migrate: migrateObservabilityShardsToID, + validate: validateDatabaseSchemaV3, + }, } } @@ -354,7 +547,7 @@ func initializeFreshDatabaseSchema(db *gorm.DB, backend string) error { if err := migrateSQLiteDataIfNeeded(db, backend); err != nil { return err } - if err := validateDatabaseSchemaV2(db, backend); err != nil { + if err := validateDatabaseSchemaV3(db, backend); err != nil { return err } return saveDatabaseSchemaVersion(db, currentDatabaseSchemaVersion) diff --git a/openflare_server/model/main_test.go b/openflare_server/model/main_test.go index dd5d6041..088bc210 100644 --- a/openflare_server/model/main_test.go +++ b/openflare_server/model/main_test.go @@ -256,6 +256,159 @@ func TestEnsureDatabaseSchemaUpToDateUpgradesLegacyDatabase(t *testing.T) { } } +func TestEnsureDatabaseSchemaUpToDateMigratesObservabilityShardsToID(t *testing.T) { + db := openBareTestSQLiteDB(t, "legacy-observability-shards.db") + if err := registerSharding(db, "sqlite"); err != nil { + t.Fatalf("register sharding: %v", err) + } + if err := autoMigrateAll(db); err != nil { + t.Fatalf("auto migrate db: %v", err) + } + if err := autoMigrateSchemaMetadata(db); err != nil { + t.Fatalf("auto migrate schema metadata: %v", err) + } + + now := time.Now().UTC() + if err := db.Table("node_metric_snapshots_00").Create(&NodeMetricSnapshot{ + ID: 1, + NodeID: "node-a", + CapturedAt: now.Add(-2 * time.Minute), + CPUUsagePercent: 22, + MemoryUsedBytes: 2, + MemoryTotalBytes: 8, + }).Error; err != nil { + t.Fatalf("seed metric snapshot shard 00: %v", err) + } + if err := db.Table("node_metric_snapshots_01").Create(&NodeMetricSnapshot{ + ID: 1, + NodeID: "node-b", + CapturedAt: now.Add(-time.Minute), + CPUUsagePercent: 44, + MemoryUsedBytes: 4, + MemoryTotalBytes: 8, + }).Error; err != nil { + t.Fatalf("seed metric snapshot shard 01: %v", err) + } + if err := db.Table("node_request_reports_00").Create(&NodeRequestReport{ + ID: 1, + NodeID: "node-a", + WindowStartedAt: now.Add(-3 * time.Minute), + WindowEndedAt: now.Add(-2 * time.Minute), + RequestCount: 12, + ErrorCount: 1, + UniqueVisitorCount: 6, + }).Error; err != nil { + t.Fatalf("seed request report shard 00: %v", err) + } + if err := db.Table("node_request_reports_01").Create(&NodeRequestReport{ + ID: 1, + NodeID: "node-b", + WindowStartedAt: now.Add(-2 * time.Minute), + WindowEndedAt: now.Add(-time.Minute), + RequestCount: 21, + ErrorCount: 2, + UniqueVisitorCount: 9, + }).Error; err != nil { + t.Fatalf("seed request report shard 01: %v", err) + } + if err := db.Table("node_access_logs_00").Create(&NodeAccessLog{ + ID: 1, + NodeID: "node-a", + LoggedAt: now.Add(-90 * time.Second), + RemoteAddr: "203.0.113.10", + Host: "a.example.com", + Path: "/alpha", + StatusCode: 200, + }).Error; err != nil { + t.Fatalf("seed access log shard 00: %v", err) + } + if err := db.Table("node_access_logs_01").Create(&NodeAccessLog{ + ID: 1, + NodeID: "node-b", + LoggedAt: now.Add(-60 * time.Second), + RemoteAddr: "203.0.113.11", + Host: "b.example.com", + Path: "/beta", + StatusCode: 502, + }).Error; err != nil { + t.Fatalf("seed access log shard 01: %v", err) + } + if err := saveDatabaseSchemaVersion(db, 2); err != nil { + t.Fatalf("save schema version: %v", err) + } + + previousDB := DB + DB = db + t.Cleanup(func() { + DB = previousDB + }) + + if err := ensureDatabaseSchemaUpToDate(db, "sqlite"); err != nil { + t.Fatalf("ensureDatabaseSchemaUpToDate: %v", err) + } + + version, exists, err := loadDatabaseSchemaVersion(db) + if err != nil { + t.Fatalf("loadDatabaseSchemaVersion: %v", err) + } + if !exists { + t.Fatal("expected migrated database to keep schema version record") + } + if version != currentDatabaseSchemaVersion { + t.Fatalf("unexpected schema version: got %d want %d", version, currentDatabaseSchemaVersion) + } + + for _, baseTable := range shardedObservabilityBaseTables() { + for _, table := range observabilityShardTables(baseTable) { + legacyTable := legacyObservabilityShardTableName(table) + if db.Migrator().HasTable(legacyTable) { + t.Fatalf("expected legacy shard table %s to be removed", legacyTable) + } + } + } + + snapshots, err := ListMetricSnapshotsSince(time.Time{}) + if err != nil { + t.Fatalf("ListMetricSnapshotsSince failed: %v", err) + } + if len(snapshots) != 2 { + t.Fatalf("expected 2 migrated metric snapshots, got %+v", snapshots) + } + reports, err := ListRequestReportsSince(time.Time{}) + if err != nil { + t.Fatalf("ListRequestReportsSince failed: %v", err) + } + if len(reports) != 2 { + t.Fatalf("expected 2 migrated request reports, got %+v", reports) + } + logs, err := ListNodeAccessLogs(NodeAccessLogQuery{Page: 0, PageSize: 10}) + if err != nil { + t.Fatalf("ListNodeAccessLogs failed: %v", err) + } + if len(logs) != 2 { + t.Fatalf("expected 2 migrated access logs, got %+v", logs) + } + + seenSnapshotIDs := make(map[uint]struct{}, len(snapshots)) + for _, item := range snapshots { + if item == nil || item.ID == 0 { + t.Fatalf("expected migrated metric snapshot to have a new non-zero id: %+v", item) + } + if _, exists := seenSnapshotIDs[item.ID]; exists { + t.Fatalf("expected migrated metric snapshot ids to be unique, got duplicate %d", item.ID) + } + seenSnapshotIDs[item.ID] = struct{}{} + targetTable := observabilityShardTableForID("node_metric_snapshots", item.ID) + var count int64 + if err := db.Table(targetTable).Where("id = ?", item.ID).Count(&count).Error; err != nil { + t.Fatalf("count migrated metric snapshot in target shard: %v", err) + } + if count != 1 { + t.Fatalf("expected migrated metric snapshot id %d to be stored in %s", item.ID, targetTable) + } + } +} + func TestRunDatabaseSchemaMigrationDoesNotAdvanceVersionWhenValidationFails(t *testing.T) { db := openBareTestSQLiteDB(t, "failed-validation.db") diff --git a/openflare_server/model/node_access_log.go b/openflare_server/model/node_access_log.go index 480e3e27..f5b79401 100644 --- a/openflare_server/model/node_access_log.go +++ b/openflare_server/model/node_access_log.go @@ -92,6 +92,10 @@ type NodeAccessLogTrendPointRow struct { RequestCount int64 `json:"request_count"` } +func (log *NodeAccessLog) BeforeCreate(tx *gorm.DB) error { + return assignObservabilityID(&log.ID) +} + func ListNodeAccessLogs(query NodeAccessLogQuery) (logs []*NodeAccessLog, err error) { all, err := listNodeAccessLogsAcrossShards(query) if err != nil { @@ -246,6 +250,46 @@ func DeleteNodeAccessLogsBefore(before time.Time) (deleted int64, err error) { return deleted, nil } +func NodeAccessLogExists(db *gorm.DB, record *NodeAccessLog) (bool, error) { + if record == nil { + return false, nil + } + db = normalizeShardedDB(db) + for _, table := range observabilityShardTables("node_access_logs") { + var count int64 + if err := db.Table(table). + Where( + "node_id = ? AND logged_at = ? AND remote_addr = ? AND host = ? AND path = ? AND status_code = ?", + record.NodeID, + record.LoggedAt, + record.RemoteAddr, + record.Host, + record.Path, + record.StatusCode, + ). + Limit(1). + Count(&count).Error; err != nil { + return false, err + } + if count > 0 { + return true, nil + } + } + return false, nil +} + +func DeleteNodeAccessLogsByNodeBefore(db *gorm.DB, nodeID string, before time.Time) (deleted int64, err error) { + db = normalizeShardedDB(db) + for _, table := range observabilityShardTables("node_access_logs") { + result := db.Table(table).Where("node_id = ? AND logged_at < ?", nodeID, before).Delete(&NodeAccessLog{}) + if result.Error != nil { + return deleted, result.Error + } + deleted += result.RowsAffected + } + return deleted, nil +} + func buildNodeAccessLogQuery(db *gorm.DB, query NodeAccessLogQuery) *gorm.DB { if db == nil { db = DB.Model(&NodeAccessLog{}) diff --git a/openflare_server/model/node_metric_snapshot.go b/openflare_server/model/node_metric_snapshot.go index 7f3aac48..434c4b57 100644 --- a/openflare_server/model/node_metric_snapshot.go +++ b/openflare_server/model/node_metric_snapshot.go @@ -26,20 +26,42 @@ type NodeMetricSnapshot struct { CreatedAt time.Time `json:"created_at"` } +func (snapshot *NodeMetricSnapshot) BeforeCreate(tx *gorm.DB) error { + return assignObservabilityID(&snapshot.ID) +} + func (snapshot *NodeMetricSnapshot) Insert() error { return DB.Create(snapshot).Error } func ListNodeMetricSnapshots(nodeID string, since time.Time, limit int) (snapshots []*NodeMetricSnapshot, err error) { - query := DB.Where("node_id = ?", nodeID).Order("captured_at desc") - if !since.IsZero() { - query = query.Where("captured_at >= ?", since) + rows, err := queryAcrossShards("node_metric_snapshots", func(tx *gorm.DB) ([]*NodeMetricSnapshot, error) { + var shardRows []*NodeMetricSnapshot + query := tx.Order("captured_at desc, id desc") + if nodeID != "" { + query = query.Where("node_id = ?", nodeID) + } + if !since.IsZero() { + query = query.Where("captured_at >= ?", since) + } + if err := query.Find(&shardRows).Error; err != nil { + return nil, err + } + return shardRows, nil + }) + if err != nil { + return nil, err } - if limit > 0 { - query = query.Limit(limit) + sort.Slice(rows, func(i int, j int) bool { + if rows[i].CapturedAt.Equal(rows[j].CapturedAt) { + return rows[i].ID > rows[j].ID + } + return rows[i].CapturedAt.After(rows[j].CapturedAt) + }) + if limit > 0 && len(rows) > limit { + rows = rows[:limit] } - err = query.Find(&snapshots).Error - return snapshots, err + return rows, nil } func ListMetricSnapshotsSince(since time.Time) (snapshots []*NodeMetricSnapshot, err error) { @@ -65,3 +87,20 @@ func ListMetricSnapshotsSince(since time.Time) (snapshots []*NodeMetricSnapshot, }) return rows, nil } + +func NodeMetricSnapshotExists(db *gorm.DB, nodeID string, capturedAt time.Time) (bool, error) { + db = normalizeShardedDB(db) + for _, table := range observabilityShardTables("node_metric_snapshots") { + var count int64 + if err := db.Table(table). + Where("node_id = ? AND captured_at = ?", nodeID, capturedAt). + Limit(1). + Count(&count).Error; err != nil { + return false, err + } + if count > 0 { + return true, nil + } + } + return false, nil +} diff --git a/openflare_server/model/node_request_report.go b/openflare_server/model/node_request_report.go index fb5d17bf..1ed9d564 100644 --- a/openflare_server/model/node_request_report.go +++ b/openflare_server/model/node_request_report.go @@ -21,20 +21,42 @@ type NodeRequestReport struct { CreatedAt time.Time `json:"created_at"` } +func (report *NodeRequestReport) BeforeCreate(tx *gorm.DB) error { + return assignObservabilityID(&report.ID) +} + func (report *NodeRequestReport) Insert() error { return DB.Create(report).Error } func ListNodeRequestReports(nodeID string, since time.Time, limit int) (reports []*NodeRequestReport, err error) { - query := DB.Where("node_id = ?", nodeID).Order("window_ended_at desc") - if !since.IsZero() { - query = query.Where("window_ended_at >= ?", since) + rows, err := queryAcrossShards("node_request_reports", func(tx *gorm.DB) ([]*NodeRequestReport, error) { + var shardRows []*NodeRequestReport + query := tx.Order("window_ended_at desc, id desc") + if nodeID != "" { + query = query.Where("node_id = ?", nodeID) + } + if !since.IsZero() { + query = query.Where("window_ended_at >= ?", since) + } + if err := query.Find(&shardRows).Error; err != nil { + return nil, err + } + return shardRows, nil + }) + if err != nil { + return nil, err } - if limit > 0 { - query = query.Limit(limit) + sort.Slice(rows, func(i int, j int) bool { + if rows[i].WindowEndedAt.Equal(rows[j].WindowEndedAt) { + return rows[i].ID > rows[j].ID + } + return rows[i].WindowEndedAt.After(rows[j].WindowEndedAt) + }) + if limit > 0 && len(rows) > limit { + rows = rows[:limit] } - err = query.Find(&reports).Error - return reports, err + return rows, nil } func ListRequestReportsSince(since time.Time) (reports []*NodeRequestReport, err error) { @@ -60,3 +82,20 @@ func ListRequestReportsSince(since time.Time) (reports []*NodeRequestReport, err }) return rows, nil } + +func NodeRequestReportExists(db *gorm.DB, nodeID string, windowStartedAt time.Time, windowEndedAt time.Time) (bool, error) { + db = normalizeShardedDB(db) + for _, table := range observabilityShardTables("node_request_reports") { + var count int64 + if err := db.Table(table). + Where("node_id = ? AND window_started_at = ? AND window_ended_at = ?", nodeID, windowStartedAt, windowEndedAt). + Limit(1). + Count(&count).Error; err != nil { + return false, err + } + if count > 0 { + return true, nil + } + } + return false, nil +} diff --git a/openflare_server/model/sharding.go b/openflare_server/model/sharding.go index 91548e54..32915cb3 100644 --- a/openflare_server/model/sharding.go +++ b/openflare_server/model/sharding.go @@ -4,20 +4,28 @@ import ( "fmt" "sort" "strings" + "sync" + "github.com/bwmarrin/snowflake" "gorm.io/gorm" "gorm.io/sharding" ) const observabilityShardCount = 10 +var ( + observabilityIDNode *snowflake.Node + observabilityIDNodeErr error + observabilityIDNodeOnce sync.Once +) + func registerSharding(db *gorm.DB, backend string) error { if db == nil { return nil } _ = backend if err := db.Use(sharding.Register(sharding.Config{ - ShardingKey: "node_id", + ShardingKey: "id", NumberOfShards: observabilityShardCount, PrimaryKeyGenerator: sharding.PKCustom, PrimaryKeyGeneratorFn: func(tableIdx int64) int64 { @@ -37,6 +45,14 @@ func shardedObservabilityTables() []any { } } +func shardedObservabilityBaseTables() []string { + return []string{ + "node_metric_snapshots", + "node_request_reports", + "node_access_logs", + } +} + func isShardedObservabilityTable(tableName string) bool { switch strings.TrimSpace(tableName) { case "node_metric_snapshots", "node_request_reports", "node_access_logs": @@ -62,10 +78,60 @@ func observabilityShardSuffixes() []string { return suffixes } +func observabilityShardSuffixForID(id uint) string { + return fmt.Sprintf("_%02d", uint64(id)%uint64(observabilityShardCount)) +} + +func observabilityShardTableForID(baseTable string, id uint) string { + return baseTable + observabilityShardSuffixForID(id) +} + +func legacyObservabilityShardTableName(tableName string) string { + return tableName + "_legacy_v2_to_v3" +} + +func normalizeShardedDB(db *gorm.DB) *gorm.DB { + if db != nil { + return db + } + return DB +} + +func nextObservabilityID() (uint, error) { + observabilityIDNodeOnce.Do(func() { + observabilityIDNode, observabilityIDNodeErr = snowflake.NewNode(0) + }) + if observabilityIDNodeErr != nil { + return 0, observabilityIDNodeErr + } + id := observabilityIDNode.Generate().Int64() + if id <= 0 { + return 0, fmt.Errorf("generated invalid observability id %d", id) + } + return uint(id), nil +} + +func assignObservabilityID(id *uint) error { + if id == nil || *id != 0 { + return nil + } + generated, err := nextObservabilityID() + if err != nil { + return err + } + *id = generated + return nil +} + func queryAcrossShards[T any](baseTable string, query func(tx *gorm.DB) ([]T, error)) ([]T, error) { + return queryAcrossShardsWithDB(DB, baseTable, query) +} + +func queryAcrossShardsWithDB[T any](db *gorm.DB, baseTable string, query func(tx *gorm.DB) ([]T, error)) ([]T, error) { items := make([]T, 0) + db = normalizeShardedDB(db) for _, table := range observabilityShardTables(baseTable) { - rows, err := query(DB.Table(table)) + rows, err := query(db.Table(table)) if err != nil { return nil, err } diff --git a/openflare_server/service/node_update_test.go b/openflare_server/service/node_update_test.go index a839aa9c..344310f9 100644 --- a/openflare_server/service/node_update_test.go +++ b/openflare_server/service/node_update_test.go @@ -882,6 +882,15 @@ func TestHeartbeatNodePersistsBufferedObservabilityPayload(t *testing.T) { TopDomains: map[string]int64{"edge.example.com": 40}, SourceCountries: map[string]int64{"CN": 20}, }, + AccessLogs: []AgentNodeAccessLog{ + { + LoggedAtUnix: now.Add(-110 * time.Second).Unix(), + RemoteAddr: "203.0.113.21", + Host: "edge.example.com", + Path: "/buffered", + StatusCode: 200, + }, + }, }, }, }) @@ -903,6 +912,18 @@ func TestHeartbeatNodePersistsBufferedObservabilityPayload(t *testing.T) { if len(reports) != 2 { t.Fatalf("expected replay dedupe to keep report count stable, got %+v", reports) } + accessLogs, err = model.ListNodeAccessLogs(model.NodeAccessLogQuery{ + NodeID: node.NodeID, + Since: time.Time{}, + Page: 0, + PageSize: 10, + }) + if err != nil { + t.Fatalf("expected node access logs query to succeed after replay: %v", err) + } + if len(accessLogs) != 1 { + t.Fatalf("expected replay dedupe to keep access log count stable, got %+v", accessLogs) + } } func TestListAccessLogsUsesPagination(t *testing.T) { diff --git a/openflare_server/service/observability.go b/openflare_server/service/observability.go index 48544bf1..ccee23d8 100644 --- a/openflare_server/service/observability.go +++ b/openflare_server/service/observability.go @@ -176,7 +176,14 @@ func persistNodeMetricSnapshot(tx *gorm.DB, nodeID string, snapshot *AgentNodeMe OpenrestyTxBytes: snapshot.OpenrestyTxBytes, OpenrestyConnections: snapshot.OpenrestyConnections, } - return tx.Where("node_id = ? AND captured_at = ?", nodeID, record.CapturedAt).FirstOrCreate(record).Error + exists, err := model.NodeMetricSnapshotExists(tx, nodeID, record.CapturedAt) + if err != nil { + return err + } + if exists { + return nil + } + return tx.Create(record).Error } func persistNodeTrafficReport(tx *gorm.DB, nodeID string, report *AgentNodeTrafficReport, reportedAt time.Time) error { @@ -197,7 +204,14 @@ func persistNodeTrafficReport(tx *gorm.DB, nodeID string, report *AgentNodeTraff TopDomainsJSON: marshalJSON(report.TopDomains), SourceCountriesJSON: marshalJSON(report.SourceCountries), } - return tx.Where("node_id = ? AND window_started_at = ? AND window_ended_at = ?", nodeID, record.WindowStartedAt, record.WindowEndedAt).FirstOrCreate(record).Error + exists, err := model.NodeRequestReportExists(tx, nodeID, record.WindowStartedAt, record.WindowEndedAt) + if err != nil { + return err + } + if exists { + return nil + } + return tx.Create(record).Error } func persistNodeAccessLogs(tx *gorm.DB, nodeID string, logs []AgentNodeAccessLog, reportedAt time.Time) error { @@ -224,19 +238,19 @@ func persistNodeAccessLogs(tx *gorm.DB, nodeID string, logs []AgentNodeAccessLog if resolver != nil { record.Region = resolver.Resolve(record.RemoteAddr) } - if err := tx.Where( - "node_id = ? AND logged_at = ? AND remote_addr = ? AND host = ? AND path = ? AND status_code = ?", - nodeID, - record.LoggedAt, - record.RemoteAddr, - record.Host, - record.Path, - record.StatusCode, - ).FirstOrCreate(record).Error; err != nil { + exists, err := model.NodeAccessLogExists(tx, record) + if err != nil { + return err + } + if exists { + continue + } + if err := tx.Create(record).Error; err != nil { return err } } - return tx.Where("node_id = ? AND logged_at < ?", nodeID, reportedAt.Add(-nodeAccessLogRetentionWindow)).Delete(&model.NodeAccessLog{}).Error + _, err = model.DeleteNodeAccessLogsByNodeBefore(tx, nodeID, reportedAt.Add(-nodeAccessLogRetentionWindow)) + return err } func reconcileNodeHealthEvents(tx *gorm.DB, nodeID string, events []AgentNodeHealthEvent, reportedAt time.Time) error {